From 91e5e1b86fe3de36bfa394568386fadeb6654749 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Aybars=20Mete=20Kele=C5=9F?= Date: Tue, 4 Aug 2026 02:01:12 +0300 Subject: [PATCH 1/7] ci: add Go verification workflow --- .github/workflows/ci.yml | 19 +++++++++++++++++++ 1 file changed, 19 insertions(+) create mode 100644 .github/workflows/ci.yml diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml new file mode 100644 index 0000000..68bc650 --- /dev/null +++ b/.github/workflows/ci.yml @@ -0,0 +1,19 @@ +name: CI + +on: + push: + branches: [main] + pull_request: + +jobs: + test: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - uses: actions/setup-go@v5 + with: + go-version-file: go.mod + - name: Vet + run: go vet ./... + - name: Test (race) + run: go test -race ./... From dbbd7d2ea6caa6d26efd250ee37cac205aed0d7e Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Aybars=20Mete=20Kele=C5=9F?= Date: Tue, 4 Aug 2026 02:07:16 +0300 Subject: [PATCH 2/7] fix: type-check planned expressions --- internal/exec/eval.go | 6 +- internal/exec/eval_test.go | 17 +++++ internal/plan/plan.go | 108 ++------------------------------ internal/plan/typecheck.go | 101 +++++++++++++++++++++++++++++ internal/plan/typecheck_test.go | 82 ++++++++++++++++++++++++ 5 files changed, 210 insertions(+), 104 deletions(-) create mode 100644 internal/plan/typecheck.go create mode 100644 internal/plan/typecheck_test.go diff --git a/internal/exec/eval.go b/internal/exec/eval.go index 908fd01..9238872 100644 --- a/internal/exec/eval.go +++ b/internal/exec/eval.go @@ -100,7 +100,11 @@ func compareResult(op string, ord int) bool { func arithmetic(op string, l, r value.Value) (value.Value, error) { if l.IsNull() || r.IsNull() { - return value.NullOf(value.TFloat), nil + outType := value.TFloat + if op != "/" && l.Type == value.TInt && r.Type == value.TInt { + outType = value.TInt + } + return value.NullOf(outType), nil } lf, rf := toF(l), toF(r) var out float64 diff --git a/internal/exec/eval_test.go b/internal/exec/eval_test.go index 81cb945..59b8e1c 100644 --- a/internal/exec/eval_test.go +++ b/internal/exec/eval_test.go @@ -47,3 +47,20 @@ func TestEvalUnknownColumn(t *testing.T) { t.Fatal("expected error for unknown column") } } + +func TestEvalArithmeticNullKeepsInferredType(t *testing.T) { + tests := []struct { + op string + want value.Type + }{{"+", value.TInt}, {"/", value.TFloat}} + for _, tt := range tests { + e := &ast.BinaryExpr{Op: tt.op, Left: &ast.ColumnRef{Name: "age"}, Right: &ast.Literal{Val: value.Int64(2)}} + got, err := Eval(e, value.Row{value.NullOf(value.TInt), value.Text("x")}, testSchema()) + if err != nil { + t.Fatal(err) + } + if !got.IsNull() || got.Type != tt.want { + t.Fatalf("%s result = %#v, want NULL type %v", tt.op, got, tt.want) + } + } +} diff --git a/internal/plan/plan.go b/internal/plan/plan.go index aded448..e3532d7 100644 --- a/internal/plan/plan.go +++ b/internal/plan/plan.go @@ -9,7 +9,6 @@ import ( "github.com/aybavs/sql-query-engine/internal/catalog" "github.com/aybavs/sql-query-engine/internal/csv" "github.com/aybavs/sql-query-engine/internal/exec" - "github.com/aybavs/sql-query-engine/internal/value" ) // Build validates st against the catalog, loads the tables, and returns the @@ -39,7 +38,7 @@ func Build(st *ast.SelectStmt, cat *catalog.Catalog, dataDir string) (exec.Opera } if st.Where != nil { - if err := validate(st.Where, schema); err != nil { + if err := requireBool(st.Where, schema, "WHERE"); err != nil { return nil, nil, err } op = exec.NewFilter(op, st.Where) @@ -48,7 +47,7 @@ func Build(st *ast.SelectStmt, cat *catalog.Catalog, dataDir string) (exec.Opera if len(st.OrderBy) > 0 { keys := make([]exec.SortKey, 0, len(st.OrderBy)) for _, o := range st.OrderBy { - if err := validate(o.Expr, schema); err != nil { + if _, err := inferExprType(o.Expr, schema); err != nil { return nil, nil, err } keys = append(keys, exec.SortKey{Expr: o.Expr, Desc: o.Desc}) @@ -117,87 +116,6 @@ func splitJoinKeys(on ast.Expr, left, right exec.Schema) (leftKey, rightKey ast. return leftKey, rightKey, nil } -// inferExprType validates a join-key expression and reports its runtime type. -// It mirrors the expression forms supported by exec.Eval without adding SQL -// syntax: numeric arithmetic, boolean logic, comparisons, and IS NULL. -func inferExprType(e ast.Expr, s exec.Schema) (value.Type, error) { - switch n := e.(type) { - case *ast.Literal: - return n.Val.Type, nil - case *ast.ColumnRef: - i, err := s.Index(n.Table, n.Name) - if err != nil { - return 0, err - } - return s[i].Type, nil - case *ast.UnaryExpr: - t, err := inferExprType(n.Expr, s) - if err != nil { - return 0, err - } - switch n.Op { - case "-": - if !numericType(t) { - return 0, fmt.Errorf("operator - requires a numeric operand") - } - return t, nil - case "NOT": - if t != value.TBool { - return 0, fmt.Errorf("operator NOT requires a BOOL operand") - } - return value.TBool, nil - default: - return 0, fmt.Errorf("unsupported unary operator %q", n.Op) - } - case *ast.IsNull: - if _, err := inferExprType(n.Expr, s); err != nil { - return 0, err - } - return value.TBool, nil - case *ast.BinaryExpr: - leftType, err := inferExprType(n.Left, s) - if err != nil { - return 0, err - } - rightType, err := inferExprType(n.Right, s) - if err != nil { - return 0, err - } - switch n.Op { - case "+", "-", "*", "/": - if !numericType(leftType) || !numericType(rightType) { - return 0, fmt.Errorf("operator %s requires numeric operands", n.Op) - } - if n.Op == "/" || leftType == value.TFloat || rightType == value.TFloat { - return value.TFloat, nil - } - return value.TInt, nil - case "=", "<>", "<", "<=", ">", ">=": - if !comparableTypes(leftType, rightType) { - return 0, fmt.Errorf("operator %s requires compatible operands", n.Op) - } - return value.TBool, nil - case "AND", "OR": - if leftType != value.TBool || rightType != value.TBool { - return 0, fmt.Errorf("operator %s requires BOOL operands", n.Op) - } - return value.TBool, nil - default: - return 0, fmt.Errorf("unsupported binary operator %q", n.Op) - } - default: - return 0, fmt.Errorf("unsupported expression %T", e) - } -} - -func comparableTypes(left, right value.Type) bool { - return numericType(left) && numericType(right) || left == right -} - -func numericType(t value.Type) bool { - return t == value.TInt || t == value.TFloat -} - const invalidJoinOwner = -1 const ( @@ -266,11 +184,12 @@ func projections(st *ast.SelectStmt, s exec.Schema) ([]ast.Expr, exec.Schema, er } continue } - if err := validate(pr.Expr, s); err != nil { + t, err := inferExprType(pr.Expr, s) + if err != nil { return nil, nil, err } exprs = append(exprs, pr.Expr) - out = append(out, exec.Column{Name: exprName(pr.Expr), Type: exprType(pr.Expr, s)}) + out = append(out, exec.Column{Name: exprName(pr.Expr), Type: t}) } return exprs, out, nil } @@ -303,20 +222,3 @@ func exprName(e ast.Expr) string { } return "expr" } - -// exprType reports the declared type of a projected column. Column references -// carry their catalog type; computed expressions are labelled by their operand -// kind and carry their precise type at runtime on each Value. -func exprType(e ast.Expr, s exec.Schema) value.Type { - switch n := e.(type) { - case *ast.ColumnRef: - if i, err := s.Index(n.Table, n.Name); err == nil { - return s[i].Type - } - case *ast.Literal: - return n.Val.Type - case *ast.IsNull: - return value.TBool - } - return value.TFloat -} diff --git a/internal/plan/typecheck.go b/internal/plan/typecheck.go new file mode 100644 index 0000000..ba59504 --- /dev/null +++ b/internal/plan/typecheck.go @@ -0,0 +1,101 @@ +package plan + +import ( + "fmt" + + "github.com/aybavs/sql-query-engine/internal/ast" + "github.com/aybavs/sql-query-engine/internal/exec" + "github.com/aybavs/sql-query-engine/internal/value" +) + +// inferExprType validates an expression and reports its runtime type. +// It mirrors the expression forms supported by exec.Eval without adding SQL +// syntax: numeric arithmetic, boolean logic, comparisons, and IS NULL. +func inferExprType(e ast.Expr, s exec.Schema) (value.Type, error) { + switch n := e.(type) { + case *ast.Literal: + return n.Val.Type, nil + case *ast.ColumnRef: + i, err := s.Index(n.Table, n.Name) + if err != nil { + return 0, err + } + return s[i].Type, nil + case *ast.UnaryExpr: + t, err := inferExprType(n.Expr, s) + if err != nil { + return 0, err + } + switch n.Op { + case "-": + if !numericType(t) { + return 0, fmt.Errorf("operator - requires a numeric operand") + } + return t, nil + case "NOT": + if t != value.TBool { + return 0, fmt.Errorf("operator NOT requires a BOOL operand") + } + return value.TBool, nil + default: + return 0, fmt.Errorf("unsupported unary operator %q", n.Op) + } + case *ast.IsNull: + if _, err := inferExprType(n.Expr, s); err != nil { + return 0, err + } + return value.TBool, nil + case *ast.BinaryExpr: + leftType, err := inferExprType(n.Left, s) + if err != nil { + return 0, err + } + rightType, err := inferExprType(n.Right, s) + if err != nil { + return 0, err + } + switch n.Op { + case "+", "-", "*", "/": + if !numericType(leftType) || !numericType(rightType) { + return 0, fmt.Errorf("operator %s requires numeric operands", n.Op) + } + if n.Op == "/" || leftType == value.TFloat || rightType == value.TFloat { + return value.TFloat, nil + } + return value.TInt, nil + case "=", "<>", "<", "<=", ">", ">=": + if !comparableTypes(leftType, rightType) { + return 0, fmt.Errorf("operator %s requires compatible operands", n.Op) + } + return value.TBool, nil + case "AND", "OR": + if leftType != value.TBool || rightType != value.TBool { + return 0, fmt.Errorf("operator %s requires BOOL operands", n.Op) + } + return value.TBool, nil + default: + return 0, fmt.Errorf("unsupported binary operator %q", n.Op) + } + default: + return 0, fmt.Errorf("unsupported expression %T", e) + } +} + +func requireBool(e ast.Expr, schema exec.Schema, clause string) error { + t, err := inferExprType(e, schema) + if err != nil { + return err + } + if t != value.TBool { + return fmt.Errorf("%s requires BOOL, got %v", clause, t) + } + return nil +} + +func comparableTypes(left, right value.Type) bool { + return numericType(left) && numericType(right) || left == right +} + +func numericType(t value.Type) bool { + return t == value.TInt || t == value.TFloat +} diff --git a/internal/plan/typecheck_test.go b/internal/plan/typecheck_test.go new file mode 100644 index 0000000..b013526 --- /dev/null +++ b/internal/plan/typecheck_test.go @@ -0,0 +1,82 @@ +package plan + +import ( + "os" + "path/filepath" + "strings" + "testing" + + "github.com/aybavs/sql-query-engine/internal/ast" + "github.com/aybavs/sql-query-engine/internal/exec" + "github.com/aybavs/sql-query-engine/internal/lexer" + "github.com/aybavs/sql-query-engine/internal/parser" + "github.com/aybavs/sql-query-engine/internal/value" +) + +func TestInferExprType(t *testing.T) { + s := exec.Schema{{Table: "users", Name: "age", Type: value.TInt}, {Table: "users", Name: "name", Type: value.TText}} + tests := []struct { + name string + expr ast.Expr + want value.Type + }{ + {"integer arithmetic", &ast.BinaryExpr{Op: "+", Left: &ast.ColumnRef{Name: "age"}, Right: &ast.Literal{Val: value.Int64(1)}}, value.TInt}, + {"division", &ast.BinaryExpr{Op: "/", Left: &ast.ColumnRef{Name: "age"}, Right: &ast.Literal{Val: value.Int64(2)}}, value.TFloat}, + {"comparison", &ast.BinaryExpr{Op: ">", Left: &ast.ColumnRef{Name: "age"}, Right: &ast.Literal{Val: value.Int64(18)}}, value.TBool}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := inferExprType(tt.expr, s) + if err != nil || got != tt.want { + t.Fatalf("inferExprType() = %v, %v; want %v, nil", got, err, tt.want) + } + }) + } +} + +func TestBuildRejectsIllTypedExpressions(t *testing.T) { + dir := t.TempDir() + if err := os.WriteFile(filepath.Join(dir, "users.csv"), []byte("1,alice,30\n"), 0o644); err != nil { + t.Fatal(err) + } + for _, sql := range []string{ + "SELECT name FROM users WHERE age + 1", + "SELECT name + 1 FROM users", + "SELECT name FROM users WHERE name > age", + "SELECT name FROM users ORDER BY name AND TRUE", + } { + t.Run(sql, func(t *testing.T) { + toks, _ := lexer.Lex(sql) + st, _ := parser.New(toks).ParseSelect() + if _, _, err := Build(st, testCatalog(), dir); err == nil { + t.Fatal("Build succeeded; want a type error") + } + }) + } +} + +func TestBuildUsesPreciseComputedProjectionTypes(t *testing.T) { + dir := t.TempDir() + if err := os.WriteFile(filepath.Join(dir, "users.csv"), []byte("1,alice,30\n"), 0o644); err != nil { + t.Fatal(err) + } + toks, _ := lexer.Lex("SELECT age + 1, age / 2, age > 18 FROM users") + st, _ := parser.New(toks).ParseSelect() + _, schema, err := Build(st, testCatalog(), dir) + if err != nil { + t.Fatal(err) + } + want := []value.Type{value.TInt, value.TFloat, value.TBool} + for i := range want { + if schema[i].Type != want[i] { + t.Fatalf("schema[%d].Type = %v, want %v", i, schema[i].Type, want[i]) + } + } +} + +func TestRequireBoolNamesClause(t *testing.T) { + err := requireBool(&ast.ColumnRef{Name: "age"}, exec.Schema{{Name: "age", Type: value.TInt}}, "WHERE") + if err == nil || !strings.Contains(err.Error(), "WHERE requires BOOL") { + t.Fatalf("error = %v", err) + } +} From 34cb16bd1adedd3c74d3626a4edeb014a88d2eb7 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Aybars=20Mete=20Kele=C5=9F?= Date: Tue, 4 Aug 2026 02:11:58 +0300 Subject: [PATCH 3/7] feat: parse aggregate queries --- internal/ast/ast.go | 20 +++++++++--- internal/parser/aggregate_test.go | 49 ++++++++++++++++++++++++++++ internal/parser/parser.go | 53 +++++++++++++++++++++++++++++++ 3 files changed, 117 insertions(+), 5 deletions(-) create mode 100644 internal/parser/aggregate_test.go diff --git a/internal/ast/ast.go b/internal/ast/ast.go index 39bc142..ee74ee0 100644 --- a/internal/ast/ast.go +++ b/internal/ast/ast.go @@ -19,12 +19,20 @@ type IsNull struct { Expr Expr Negate bool } +type AggregateCall struct { + Name string + Arg Expr + Star bool +} +type SlotRef struct{ Index int } -func (*ColumnRef) isExpr() {} -func (*Literal) isExpr() {} -func (*BinaryExpr) isExpr() {} -func (*UnaryExpr) isExpr() {} -func (*IsNull) isExpr() {} +func (*ColumnRef) isExpr() {} +func (*Literal) isExpr() {} +func (*BinaryExpr) isExpr() {} +func (*UnaryExpr) isExpr() {} +func (*IsNull) isExpr() {} +func (*AggregateCall) isExpr() {} +func (*SlotRef) isExpr() {} // SelectStmt and friends are used by the statement parser. type SelectStmt struct { @@ -32,6 +40,8 @@ type SelectStmt struct { From string Joins []Join Where Expr + GroupBy []Expr + Having Expr OrderBy []OrderItem Limit *int } diff --git a/internal/parser/aggregate_test.go b/internal/parser/aggregate_test.go new file mode 100644 index 0000000..0d1528a --- /dev/null +++ b/internal/parser/aggregate_test.go @@ -0,0 +1,49 @@ +package parser + +import ( + "testing" + + "github.com/aybavs/sql-query-engine/internal/ast" + "github.com/aybavs/sql-query-engine/internal/lexer" +) + +func TestParseAggregateQuery(t *testing.T) { + st := parseSelect(t, "SELECT city, COUNT(*), avg(age) FROM users WHERE age > 0 GROUP BY city HAVING COUNT(*) >= 2 ORDER BY AVG(age) DESC LIMIT 3") + if len(st.GroupBy) != 1 || st.Having == nil || len(st.OrderBy) != 1 { + t.Fatalf("statement = %#v", st) + } + count, ok := st.Projections[1].Expr.(*ast.AggregateCall) + if !ok || count.Name != "COUNT" || !count.Star || count.Arg != nil { + t.Fatalf("count = %#v", st.Projections[1].Expr) + } + avg, ok := st.Projections[2].Expr.(*ast.AggregateCall) + if !ok || avg.Name != "AVG" || avg.Star || avg.Arg == nil { + t.Fatalf("avg = %#v", st.Projections[2].Expr) + } +} + +func TestParseMultiColumnGroupBy(t *testing.T) { + st := parseSelect(t, "SELECT city, age, COUNT(*) FROM users GROUP BY city, age") + if len(st.GroupBy) != 2 { + t.Fatalf("GroupBy length = %d, want 2", len(st.GroupBy)) + } +} + +func TestParseRejectsMalformedAggregateCalls(t *testing.T) { + for _, sql := range []string{ + "SELECT COUNT() FROM users", + "SELECT SUM(*) FROM users", + "SELECT AVG(age, id) FROM users", + "SELECT COUNT(age FROM users", + } { + t.Run(sql, func(t *testing.T) { + toks, err := lexer.Lex(sql) + if err == nil { + _, err = New(toks).ParseSelect() + } + if err == nil { + t.Fatal("parse succeeded; want an error") + } + }) + } +} diff --git a/internal/parser/parser.go b/internal/parser/parser.go index fe828ec..8527594 100644 --- a/internal/parser/parser.go +++ b/internal/parser/parser.go @@ -4,6 +4,7 @@ package parser import ( "fmt" "strconv" + "strings" "github.com/aybavs/sql-query-engine/internal/ast" "github.com/aybavs/sql-query-engine/internal/lexer" @@ -146,6 +147,31 @@ func (p *Parser) parsePrimary() (ast.Expr, error) { } return nil, fmt.Errorf("unexpected keyword %q at pos %d", t.Text, t.Pos) case lexer.Ident: + if p.peek().Kind == lexer.LParen { + p.next() + call := &ast.AggregateCall{Name: strings.ToUpper(t.Text)} + if p.peek().Kind == lexer.Star { + p.next() + call.Star = true + } else { + if p.peek().Kind == lexer.RParen { + return nil, fmt.Errorf("aggregate %s requires one argument", call.Name) + } + arg, err := p.ParseExpr() + if err != nil { + return nil, err + } + call.Arg = arg + } + if p.peek().Kind != lexer.RParen { + return nil, fmt.Errorf("aggregate %s requires exactly one argument", call.Name) + } + p.next() + if call.Star && call.Name != "COUNT" { + return nil, fmt.Errorf("only COUNT accepts *") + } + return call, nil + } if p.peek().Kind == lexer.Dot { p.next() col := p.next() @@ -226,6 +252,33 @@ func (p *Parser) ParseSelect() (*ast.SelectStmt, error) { st.Where = e } + if p.peek().Kind == lexer.Keyword && p.peek().Text == "GROUP" { + p.next() + if err := p.expectKeyword("BY"); err != nil { + return nil, err + } + for { + e, err := p.ParseExpr() + if err != nil { + return nil, err + } + st.GroupBy = append(st.GroupBy, e) + if p.peek().Kind != lexer.Comma { + break + } + p.next() + } + } + + if p.peek().Kind == lexer.Keyword && p.peek().Text == "HAVING" { + p.next() + e, err := p.ParseExpr() + if err != nil { + return nil, err + } + st.Having = e + } + if p.peek().Kind == lexer.Keyword && p.peek().Text == "ORDER" { p.next() if err := p.expectKeyword("BY"); err != nil { From 3e71bd0ff7e71c212cea4bb55526edc1bba63bdf Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Aybars=20Mete=20Kele=C5=9F?= Date: Tue, 4 Aug 2026 02:16:48 +0300 Subject: [PATCH 4/7] feat: add aggregate execution operator --- internal/exec/aggregate.go | 214 ++++++++++++++++++++++++++++++++ internal/exec/aggregate_test.go | 112 +++++++++++++++++ internal/exec/groupkey.go | 56 +++++++++ 3 files changed, 382 insertions(+) create mode 100644 internal/exec/aggregate.go create mode 100644 internal/exec/aggregate_test.go create mode 100644 internal/exec/groupkey.go diff --git a/internal/exec/aggregate.go b/internal/exec/aggregate.go new file mode 100644 index 0000000..5cae014 --- /dev/null +++ b/internal/exec/aggregate.go @@ -0,0 +1,214 @@ +package exec + +import ( + "github.com/aybavs/sql-query-engine/internal/ast" + "github.com/aybavs/sql-query-engine/internal/value" +) + +type AggregateKind int + +const ( + AggCount AggregateKind = iota + AggSum + AggAvg + AggMin + AggMax +) + +type AggregateSpec struct { + Kind AggregateKind + Expr ast.Expr + Star bool + OutType value.Type +} + +type aggregateState struct { + count int64 + intSum int64 + floatSum float64 + seen bool + best value.Value +} + +type aggregateGroup struct { + values value.Row + states []aggregateState +} + +// Aggregate is a blocking operator that groups all child rows before emitting +// finalized aggregate values. +type Aggregate struct { + child Operator + groupExprs []ast.Expr + specs []AggregateSpec + schema Schema + groups []*aggregateGroup + initialized bool + failed bool + index int +} + +func NewAggregate(child Operator, groupExprs []ast.Expr, specs []AggregateSpec, schema Schema) *Aggregate { + return &Aggregate{ + child: child, + groupExprs: groupExprs, + specs: specs, + schema: schema, + } +} + +func (a *Aggregate) Schema() Schema { return a.schema } + +func (a *Aggregate) build() { + a.initialized = true + byKey := make(map[string]*aggregateGroup) + if len(a.groupExprs) == 0 { + group := a.newGroup(nil) + a.groups = append(a.groups, group) + } + + for { + row, ok := a.child.Next() + if !ok { + return + } + + var group *aggregateGroup + if len(a.groupExprs) == 0 { + group = a.groups[0] + } else { + values := make(value.Row, len(a.groupExprs)) + for i, expr := range a.groupExprs { + v, err := Eval(expr, row, a.child.Schema()) + if err != nil { + a.failed = true + return + } + values[i] = v + } + key := groupKey(values) + group = byKey[key] + if group == nil { + group = a.newGroup(values) + byKey[key] = group + a.groups = append(a.groups, group) + } + } + + if err := a.step(group, row); err != nil { + a.failed = true + return + } + } +} + +func (a *Aggregate) newGroup(values value.Row) *aggregateGroup { + return &aggregateGroup{ + values: values, + states: make([]aggregateState, len(a.specs)), + } +} + +func (a *Aggregate) step(group *aggregateGroup, row value.Row) error { + for i, spec := range a.specs { + state := &group.states[i] + if spec.Kind == AggCount && spec.Star { + state.count++ + continue + } + v, err := Eval(spec.Expr, row, a.child.Schema()) + if err != nil { + return err + } + if v.IsNull() { + continue + } + if spec.Kind == AggCount { + state.count++ + continue + } + if v.IsNaN() { + continue + } + + switch spec.Kind { + case AggSum: + state.seen = true + if v.Type == value.TInt { + state.intSum += v.I + } else { + state.floatSum += v.F + } + case AggAvg: + state.count++ + if v.Type == value.TInt { + state.floatSum += float64(v.I) + } else { + state.floatSum += v.F + } + case AggMin: + if !state.seen { + state.seen = true + state.best = v + continue + } + if ord, known := value.Compare(v, state.best); known && ord < 0 { + state.best = v + } + case AggMax: + if !state.seen { + state.seen = true + state.best = v + continue + } + if ord, known := value.Compare(v, state.best); known && ord > 0 { + state.best = v + } + } + } + return nil +} + +func (a *Aggregate) Next() (value.Row, bool) { + if !a.initialized { + a.build() + } + if a.failed || a.index >= len(a.groups) { + return nil, false + } + + group := a.groups[a.index] + a.index++ + row := make(value.Row, 0, len(group.values)+len(a.specs)) + row = append(row, group.values...) + for i, spec := range a.specs { + row = append(row, finalizeAggregate(spec, group.states[i])) + } + return row, true +} + +func finalizeAggregate(spec AggregateSpec, state aggregateState) value.Value { + switch spec.Kind { + case AggCount: + return value.Int64(state.count) + case AggSum: + if !state.seen { + return value.NullOf(spec.OutType) + } + if spec.OutType == value.TInt { + return value.Int64(state.intSum) + } + return value.Float64(state.floatSum) + case AggAvg: + if state.count == 0 { + return value.NullOf(value.TFloat) + } + return value.Float64(state.floatSum / float64(state.count)) + case AggMin, AggMax: + if !state.seen { + return value.NullOf(spec.OutType) + } + return state.best + } + panic("unknown aggregate kind") +} diff --git a/internal/exec/aggregate_test.go b/internal/exec/aggregate_test.go new file mode 100644 index 0000000..49e1012 --- /dev/null +++ b/internal/exec/aggregate_test.go @@ -0,0 +1,112 @@ +package exec + +import ( + "math" + "testing" + + "github.com/aybavs/sql-query-engine/internal/ast" + "github.com/aybavs/sql-query-engine/internal/value" +) + +func drainRows(op Operator) []value.Row { + var rows []value.Row + for { + row, ok := op.Next() + if !ok { + return rows + } + rows = append(rows, row) + } +} + +func TestAggregateGlobalAllFunctions(t *testing.T) { + in := Schema{{Name: "n", Type: value.TInt}} + rows := []value.Row{{value.Int64(2)}, {value.NullOf(value.TInt)}, {value.Int64(4)}} + expr := &ast.ColumnRef{Name: "n"} + specs := []AggregateSpec{ + {Kind: AggCount, Star: true, OutType: value.TInt}, + {Kind: AggCount, Expr: expr, OutType: value.TInt}, + {Kind: AggSum, Expr: expr, OutType: value.TInt}, + {Kind: AggAvg, Expr: expr, OutType: value.TFloat}, + {Kind: AggMin, Expr: expr, OutType: value.TInt}, + {Kind: AggMax, Expr: expr, OutType: value.TInt}, + } + out := Schema{{Name: "count_star", Type: value.TInt}, {Name: "count_n", Type: value.TInt}, {Name: "sum_n", Type: value.TInt}, {Name: "avg_n", Type: value.TFloat}, {Name: "min_n", Type: value.TInt}, {Name: "max_n", Type: value.TInt}} + got := drainRows(NewAggregate(NewScan(in, rows), nil, specs, out)) + want := []string{"3", "2", "6", "3", "2", "4"} + if len(got) != 1 { + t.Fatalf("row count = %d, want 1", len(got)) + } + for i := range want { + if got[0][i].String() != want[i] { + t.Fatalf("cell %d = %s, want %s", i, got[0][i].String(), want[i]) + } + } +} + +func TestAggregateGroupsInFirstSeenOrder(t *testing.T) { + in := Schema{{Name: "city", Type: value.TText}, {Name: "n", Type: value.TInt}} + rows := []value.Row{{value.Text("b"), value.Int64(2)}, {value.Text("a"), value.Int64(3)}, {value.Text("b"), value.Int64(4)}} + out := Schema{{Name: "city", Type: value.TText}, {Name: "sum", Type: value.TInt}} + op := NewAggregate(NewScan(in, rows), []ast.Expr{&ast.ColumnRef{Name: "city"}}, []AggregateSpec{{Kind: AggSum, Expr: &ast.ColumnRef{Name: "n"}, OutType: value.TInt}}, out) + got := drainRows(op) + if len(got) != 2 || got[0][0].String() != "b" || got[0][1].String() != "6" || got[1][0].String() != "a" { + t.Fatalf("rows = %v", got) + } +} + +func TestAggregateEmptyInput(t *testing.T) { + in := Schema{{Name: "n", Type: value.TInt}} + global := drainRows(NewAggregate(NewScan(in, nil), nil, []AggregateSpec{{Kind: AggCount, Star: true, OutType: value.TInt}, {Kind: AggSum, Expr: &ast.ColumnRef{Name: "n"}, OutType: value.TInt}}, Schema{{Type: value.TInt}, {Type: value.TInt}})) + if len(global) != 1 || global[0][0].String() != "0" || !global[0][1].IsNull() || global[0][1].Type != value.TInt { + t.Fatalf("global = %#v", global) + } + grouped := drainRows(NewAggregate(NewScan(in, nil), []ast.Expr{&ast.ColumnRef{Name: "n"}}, nil, Schema{{Type: value.TInt}})) + if len(grouped) != 0 { + t.Fatalf("grouped rows = %d, want 0", len(grouped)) + } +} + +func TestAggregateNumericGroupingAndNaN(t *testing.T) { + in := Schema{{Name: "k", Type: value.TFloat}, {Name: "n", Type: value.TFloat}} + rows := []value.Row{{value.Float64(0), value.Float64(math.NaN())}, {value.Float64(math.Copysign(0, -1)), value.Float64(2)}, {value.Float64(math.NaN()), value.Float64(math.NaN())}, {value.Float64(math.NaN()), value.Float64(4)}} + specs := []AggregateSpec{{Kind: AggCount, Expr: &ast.ColumnRef{Name: "n"}, OutType: value.TInt}, {Kind: AggAvg, Expr: &ast.ColumnRef{Name: "n"}, OutType: value.TFloat}, {Kind: AggMin, Expr: &ast.ColumnRef{Name: "n"}, OutType: value.TFloat}, {Kind: AggMax, Expr: &ast.ColumnRef{Name: "n"}, OutType: value.TFloat}} + got := drainRows(NewAggregate(NewScan(in, rows), []ast.Expr{&ast.ColumnRef{Name: "k"}}, specs, Schema{{Type: value.TFloat}, {Type: value.TInt}, {Type: value.TFloat}, {Type: value.TFloat}, {Type: value.TFloat}})) + if len(got) != 2 || got[0][1].String() != "2" || got[0][2].String() != "2" || got[0][3].String() != "2" || got[0][4].String() != "2" || got[1][1].String() != "2" || got[1][2].String() != "4" || got[1][3].String() != "4" || got[1][4].String() != "4" { + t.Fatalf("rows = %#v", got) + } +} + +func TestGroupKeyNumericEquivalenceWithoutIntegerCollisions(t *testing.T) { + if groupKey([]value.Value{value.Int64(2)}) != groupKey([]value.Value{value.Float64(2)}) { + t.Fatal("INT 2 and FLOAT 2 must share a key") + } + if groupKey([]value.Value{value.Float64(0)}) != groupKey([]value.Value{value.Float64(math.Copysign(0, -1))}) { + t.Fatal("signed zero must share a key") + } + if groupKey([]value.Value{value.Float64(math.NaN())}) != groupKey([]value.Value{value.Float64(math.NaN())}) { + t.Fatal("NaN values must share a key") + } + if groupKey([]value.Value{value.Int64(9007199254740992)}) == groupKey([]value.Value{value.Int64(9007199254740993)}) { + t.Fatal("distinct large integers collided") + } + if groupKey([]value.Value{value.NullOf(value.TInt)}) != groupKey([]value.Value{value.NullOf(value.TInt)}) { + t.Fatal("NULL values must share a key") + } + if groupKey([]value.Value{value.Text("ab"), value.Text("c")}) == groupKey([]value.Value{value.Text("a"), value.Text("bc")}) { + t.Fatal("composite text boundaries collided") + } +} + +func TestAggregateMinMaxTextAndBool(t *testing.T) { + in := Schema{{Name: "text", Type: value.TText}, {Name: "flag", Type: value.TBool}} + rows := []value.Row{{value.Text("z"), value.Bool(true)}, {value.Text("a"), value.Bool(false)}} + specs := []AggregateSpec{{Kind: AggMin, Expr: &ast.ColumnRef{Name: "text"}, OutType: value.TText}, {Kind: AggMax, Expr: &ast.ColumnRef{Name: "text"}, OutType: value.TText}, {Kind: AggMin, Expr: &ast.ColumnRef{Name: "flag"}, OutType: value.TBool}, {Kind: AggMax, Expr: &ast.ColumnRef{Name: "flag"}, OutType: value.TBool}} + got := drainRows(NewAggregate(NewScan(in, rows), nil, specs, Schema{{Type: value.TText}, {Type: value.TText}, {Type: value.TBool}, {Type: value.TBool}})) + want := []string{"a", "z", "false", "true"} + for i := range want { + if got[0][i].String() != want[i] { + t.Fatalf("cell %d = %s, want %s", i, got[0][i].String(), want[i]) + } + } +} diff --git a/internal/exec/groupkey.go b/internal/exec/groupkey.go new file mode 100644 index 0000000..13ea5e0 --- /dev/null +++ b/internal/exec/groupkey.go @@ -0,0 +1,56 @@ +package exec + +import ( + "encoding/binary" + "math" + + "github.com/aybavs/sql-query-engine/internal/value" +) + +// groupKey returns a collision-safe, canonical encoding of a grouping tuple. +func groupKey(values []value.Value) string { + key := make([]byte, 0, len(values)*9) + for _, v := range values { + switch { + case v.IsNull(): + key = append(key, 'n') + key = binary.BigEndian.AppendUint64(key, uint64(v.Type)) + case v.Type == value.TInt || v.Type == value.TFloat: + domain, bits := canonicalNumber(v) + key = append(key, domain) + key = binary.BigEndian.AppendUint64(key, bits) + case v.Type == value.TText: + key = append(key, 's') + key = binary.BigEndian.AppendUint64(key, uint64(len(v.S))) + key = append(key, v.S...) + case v.Type == value.TBool: + key = append(key, 'b') + if v.B { + key = append(key, 1) + } else { + key = append(key, 0) + } + } + } + return string(key) +} + +func canonicalNumber(v value.Value) (domain byte, bits uint64) { + if v.Type == value.TInt { + return 'i', uint64(v.I) + } + f := v.F + if f == 0 { + return 'i', 0 + } + if math.IsNaN(f) { + return 'f', 0x7ff8000000000000 + } + if f >= math.MinInt64 && f < 9223372036854775808.0 { + i := int64(f) + if float64(i) == f { + return 'i', uint64(i) + } + } + return 'f', math.Float64bits(f) +} From 4ca95e398dddbf946f650cb291af42e9465868ee Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Aybars=20Mete=20Kele=C5=9F?= Date: Tue, 4 Aug 2026 02:28:06 +0300 Subject: [PATCH 5/7] feat: lower aggregate expressions in planner --- internal/plan/aggregate.go | 302 ++++++++++++++++++++++++++++++++ internal/plan/aggregate_test.go | 169 ++++++++++++++++++ internal/plan/typecheck.go | 5 + 3 files changed, 476 insertions(+) create mode 100644 internal/plan/aggregate.go create mode 100644 internal/plan/aggregate_test.go diff --git a/internal/plan/aggregate.go b/internal/plan/aggregate.go new file mode 100644 index 0000000..61d5b9d --- /dev/null +++ b/internal/plan/aggregate.go @@ -0,0 +1,302 @@ +package plan + +import ( + "fmt" + + "github.com/aybavs/sql-query-engine/internal/ast" + "github.com/aybavs/sql-query-engine/internal/exec" + "github.com/aybavs/sql-query-engine/internal/value" +) + +func containsAggregate(e ast.Expr) bool { + switch n := e.(type) { + case *ast.AggregateCall: + return true + case *ast.UnaryExpr: + return containsAggregate(n.Expr) + case *ast.BinaryExpr: + return containsAggregate(n.Left) || containsAggregate(n.Right) + case *ast.IsNull: + return containsAggregate(n.Expr) + default: + return false + } +} + +func isAggregateQuery(st *ast.SelectStmt) bool { + if len(st.GroupBy) > 0 { + return true + } + for _, projection := range st.Projections { + if containsAggregate(projection.Expr) { + return true + } + } + if containsAggregate(st.Having) { + return true + } + for _, item := range st.OrderBy { + if containsAggregate(item.Expr) { + return true + } + } + return false +} + +func buildAggregatePlan( + st *ast.SelectStmt, + input exec.Operator, + inputSchema exec.Schema, +) (exec.Operator, exec.Schema, []ast.Expr, ast.Expr, []ast.OrderItem, error) { + if containsAggregate(st.Where) { + return nil, nil, nil, nil, nil, fmt.Errorf("WHERE cannot contain aggregates") + } + if st.Having != nil && !isAggregateQuery(st) { + return nil, nil, nil, nil, nil, fmt.Errorf("HAVING requires an aggregate query") + } + for _, projection := range st.Projections { + if projection.Star { + return nil, nil, nil, nil, nil, fmt.Errorf("SELECT * is not supported in aggregate queries") + } + } + + groupExprs := make([]ast.Expr, 0, len(st.GroupBy)) + groupSlots := make(map[int]int, len(st.GroupBy)) + aggregateSchema := make(exec.Schema, 0, len(st.GroupBy)) + for _, expr := range st.GroupBy { + column, ok := expr.(*ast.ColumnRef) + if !ok { + return nil, nil, nil, nil, nil, fmt.Errorf("GROUP BY requires columns") + } + inputIndex, err := inputSchema.Index(column.Table, column.Name) + if err != nil { + return nil, nil, nil, nil, nil, err + } + if _, exists := groupSlots[inputIndex]; exists { + return nil, nil, nil, nil, nil, fmt.Errorf("duplicate GROUP BY column %s", column.Name) + } + groupSlots[inputIndex] = len(groupExprs) + groupExprs = append(groupExprs, column) + aggregateSchema = append(aggregateSchema, inputSchema[inputIndex]) + } + + collector := aggregateCollector{ + inputSchema: inputSchema, + slots: make(map[*ast.AggregateCall]int), + } + for _, projection := range st.Projections { + if err := collector.collect(projection.Expr); err != nil { + return nil, nil, nil, nil, nil, err + } + } + if err := collector.collect(st.Having); err != nil { + return nil, nil, nil, nil, nil, err + } + for _, item := range st.OrderBy { + if err := collector.collect(item.Expr); err != nil { + return nil, nil, nil, nil, nil, err + } + } + + for i, spec := range collector.specs { + collector.slots[collector.calls[i]] = len(aggregateSchema) + aggregateSchema = append(aggregateSchema, exec.Column{Name: "expr", Type: spec.OutType}) + } + + projections := make([]ast.Expr, 0, len(st.Projections)) + for _, projection := range st.Projections { + lowered, err := lowerAggregateExpr(projection.Expr, collector.slots, groupSlots, inputSchema) + if err != nil { + return nil, nil, nil, nil, nil, err + } + if _, err := inferExprType(lowered, aggregateSchema); err != nil { + return nil, nil, nil, nil, nil, err + } + projections = append(projections, lowered) + } + + var having ast.Expr + if st.Having != nil { + var err error + having, err = lowerAggregateExpr(st.Having, collector.slots, groupSlots, inputSchema) + if err != nil { + return nil, nil, nil, nil, nil, err + } + if err := requireBool(having, aggregateSchema, "HAVING"); err != nil { + return nil, nil, nil, nil, nil, err + } + } + + orderBy := make([]ast.OrderItem, 0, len(st.OrderBy)) + for _, item := range st.OrderBy { + lowered, err := lowerAggregateExpr(item.Expr, collector.slots, groupSlots, inputSchema) + if err != nil { + return nil, nil, nil, nil, nil, err + } + if _, err := inferExprType(lowered, aggregateSchema); err != nil { + return nil, nil, nil, nil, nil, err + } + orderBy = append(orderBy, ast.OrderItem{Expr: lowered, Desc: item.Desc}) + } + + aggregate := exec.NewAggregate(input, groupExprs, collector.specs, aggregateSchema) + return aggregate, aggregateSchema, projections, having, orderBy, nil +} + +type aggregateCollector struct { + inputSchema exec.Schema + calls []*ast.AggregateCall + specs []exec.AggregateSpec + slots map[*ast.AggregateCall]int +} + +func (c *aggregateCollector) collect(e ast.Expr) error { + switch n := e.(type) { + case nil, *ast.Literal, *ast.ColumnRef, *ast.SlotRef: + return nil + case *ast.AggregateCall: + if containsAggregate(n.Arg) { + return fmt.Errorf("nested aggregate") + } + spec, err := aggregateSpec(n, c.inputSchema) + if err != nil { + return err + } + c.calls = append(c.calls, n) + c.specs = append(c.specs, spec) + return nil + case *ast.UnaryExpr: + return c.collect(n.Expr) + case *ast.BinaryExpr: + if err := c.collect(n.Left); err != nil { + return err + } + return c.collect(n.Right) + case *ast.IsNull: + return c.collect(n.Expr) + default: + return fmt.Errorf("unsupported expression %T", e) + } +} + +func aggregateSpec(call *ast.AggregateCall, inputSchema exec.Schema) (exec.AggregateSpec, error) { + switch call.Name { + case "COUNT": + if call.Star { + return exec.AggregateSpec{Kind: exec.AggCount, Star: true, OutType: value.TInt}, nil + } + argType, err := inferExprType(call.Arg, inputSchema) + if err != nil { + return exec.AggregateSpec{}, err + } + _ = argType + return exec.AggregateSpec{Kind: exec.AggCount, Expr: call.Arg, OutType: value.TInt}, nil + case "SUM": + return numericAggregateSpec(exec.AggSum, call, inputSchema, false) + case "AVG": + return numericAggregateSpec(exec.AggAvg, call, inputSchema, true) + case "MIN": + return orderedAggregateSpec(exec.AggMin, call, inputSchema) + case "MAX": + return orderedAggregateSpec(exec.AggMax, call, inputSchema) + default: + return exec.AggregateSpec{}, fmt.Errorf("unknown aggregate %s", call.Name) + } +} + +func numericAggregateSpec( + kind exec.AggregateKind, + call *ast.AggregateCall, + inputSchema exec.Schema, + forceFloat bool, +) (exec.AggregateSpec, error) { + if call.Star { + return exec.AggregateSpec{}, fmt.Errorf("%s does not accept *", call.Name) + } + argType, err := inferExprType(call.Arg, inputSchema) + if err != nil { + return exec.AggregateSpec{}, err + } + if !numericType(argType) { + return exec.AggregateSpec{}, fmt.Errorf("%s requires numeric argument", call.Name) + } + outType := argType + if forceFloat { + outType = value.TFloat + } + return exec.AggregateSpec{Kind: kind, Expr: call.Arg, OutType: outType}, nil +} + +func orderedAggregateSpec( + kind exec.AggregateKind, + call *ast.AggregateCall, + inputSchema exec.Schema, +) (exec.AggregateSpec, error) { + if call.Star { + return exec.AggregateSpec{}, fmt.Errorf("%s does not accept *", call.Name) + } + argType, err := inferExprType(call.Arg, inputSchema) + if err != nil { + return exec.AggregateSpec{}, err + } + switch argType { + case value.TInt, value.TFloat, value.TText, value.TBool: + return exec.AggregateSpec{Kind: kind, Expr: call.Arg, OutType: argType}, nil + default: + return exec.AggregateSpec{}, fmt.Errorf("%s requires an ordered argument", call.Name) + } +} + +func lowerAggregateExpr( + e ast.Expr, + aggregateSlots map[*ast.AggregateCall]int, + groupSlots map[int]int, + inputSchema exec.Schema, +) (ast.Expr, error) { + switch n := e.(type) { + case *ast.Literal: + return n, nil + case *ast.SlotRef: + return n, nil + case *ast.AggregateCall: + slot, ok := aggregateSlots[n] + if !ok { + return nil, fmt.Errorf("aggregate %s has no output slot", n.Name) + } + return &ast.SlotRef{Index: slot}, nil + case *ast.ColumnRef: + inputIndex, err := inputSchema.Index(n.Table, n.Name) + if err != nil { + return nil, err + } + slot, ok := groupSlots[inputIndex] + if !ok { + return nil, fmt.Errorf("column %s must appear in GROUP BY", n.Name) + } + return &ast.SlotRef{Index: slot}, nil + case *ast.UnaryExpr: + expr, err := lowerAggregateExpr(n.Expr, aggregateSlots, groupSlots, inputSchema) + if err != nil { + return nil, err + } + return &ast.UnaryExpr{Op: n.Op, Expr: expr}, nil + case *ast.BinaryExpr: + left, err := lowerAggregateExpr(n.Left, aggregateSlots, groupSlots, inputSchema) + if err != nil { + return nil, err + } + right, err := lowerAggregateExpr(n.Right, aggregateSlots, groupSlots, inputSchema) + if err != nil { + return nil, err + } + return &ast.BinaryExpr{Op: n.Op, Left: left, Right: right}, nil + case *ast.IsNull: + expr, err := lowerAggregateExpr(n.Expr, aggregateSlots, groupSlots, inputSchema) + if err != nil { + return nil, err + } + return &ast.IsNull{Expr: expr, Negate: n.Negate}, nil + default: + return nil, fmt.Errorf("unsupported expression %T", e) + } +} diff --git a/internal/plan/aggregate_test.go b/internal/plan/aggregate_test.go new file mode 100644 index 0000000..f6d07b3 --- /dev/null +++ b/internal/plan/aggregate_test.go @@ -0,0 +1,169 @@ +package plan + +import ( + "strings" + "testing" + + "github.com/aybavs/sql-query-engine/internal/ast" + "github.com/aybavs/sql-query-engine/internal/exec" + "github.com/aybavs/sql-query-engine/internal/lexer" + "github.com/aybavs/sql-query-engine/internal/parser" + "github.com/aybavs/sql-query-engine/internal/value" +) + +func parsePlanSelect(t *testing.T, sql string) *ast.SelectStmt { + t.Helper() + tokens, err := lexer.Lex(sql) + if err != nil { + t.Fatal(err) + } + statement, err := parser.New(tokens).ParseSelect() + if err != nil { + t.Fatal(err) + } + return statement +} + +func TestIsAggregateQuery(t *testing.T) { + for _, sql := range []string{ + "SELECT COUNT(*) FROM users", + "SELECT city FROM users GROUP BY city", + "SELECT city FROM users ORDER BY MAX(age)", + } { + if !isAggregateQuery(parsePlanSelect(t, sql)) { + t.Fatalf("not classified as aggregate: %s", sql) + } + } + if isAggregateQuery(parsePlanSelect(t, "SELECT age FROM users ORDER BY age")) { + t.Fatal("ordinary query classified as aggregate") + } +} + +func TestBuildAggregatePlanLowersSlots(t *testing.T) { + st := parsePlanSelect(t, "SELECT city, COUNT(*) + 1 FROM users GROUP BY city HAVING COUNT(*) > 1 ORDER BY AVG(age) DESC") + inSchema := exec.Schema{{Table: "users", Name: "city", Type: value.TText}, {Table: "users", Name: "age", Type: value.TInt}} + op, aggSchema, projections, having, orderBy, err := buildAggregatePlan(st, exec.NewScan(inSchema, nil), inSchema) + if err != nil { + t.Fatal(err) + } + if op == nil || len(aggSchema) != 4 { + t.Fatalf("schema = %#v", aggSchema) + } + if _, ok := projections[0].(*ast.SlotRef); !ok { + t.Fatalf("group projection = %T", projections[0]) + } + if containsAggregate(projections[1]) || containsAggregate(having) || containsAggregate(orderBy[0].Expr) { + t.Fatal("post-aggregate expression was not fully lowered") + } + wantTypes := []value.Type{value.TText, value.TInt, value.TInt, value.TFloat} + for i, want := range wantTypes { + if aggSchema[i].Type != want { + t.Fatalf("schema[%d].Type = %v, want %v", i, aggSchema[i].Type, want) + } + } + assertInternalAggregateExpr(t, projections[0]) + assertInternalAggregateExpr(t, projections[1]) + assertInternalAggregateExpr(t, having) + assertInternalAggregateExpr(t, orderBy[0].Expr) + if projections[0].(*ast.SlotRef).Index != 0 || slotInBinary(t, projections[1], true) != 1 || + slotInBinary(t, having, true) != 2 || orderBy[0].Expr.(*ast.SlotRef).Index != 3 { + t.Fatalf("aggregate slots were not assigned in GROUP, SELECT, HAVING, ORDER BY order") + } +} + +func TestBuildAggregatePlanUsesResolvedGroupIdentity(t *testing.T) { + st := parsePlanSelect(t, "SELECT users.age, COUNT(*) FROM users GROUP BY age") + s := exec.Schema{{Table: "users", Name: "age", Type: value.TInt}} + _, _, projections, _, _, err := buildAggregatePlan(st, exec.NewScan(s, nil), s) + if err != nil { + t.Fatal(err) + } + if slot := projections[0].(*ast.SlotRef); slot.Index != 0 { + t.Fatalf("group slot = %d, want 0", slot.Index) + } + + st = parsePlanSelect(t, "SELECT COUNT(*) FROM users GROUP BY age, users.age") + _, _, _, _, _, err = buildAggregatePlan(st, exec.NewScan(s, nil), s) + if err == nil || !strings.Contains(err.Error(), "duplicate GROUP BY") { + t.Fatalf("duplicate error = %v", err) + } +} + +func TestBuildAggregatePlanPreservesAggregateOutputTypes(t *testing.T) { + st := parsePlanSelect(t, "SELECT COUNT(label), SUM(n), SUM(ratio), AVG(n), MIN(label), MAX(active) FROM users") + s := exec.Schema{ + {Table: "users", Name: "n", Type: value.TInt}, + {Table: "users", Name: "ratio", Type: value.TFloat}, + {Table: "users", Name: "label", Type: value.TText}, + {Table: "users", Name: "active", Type: value.TBool}, + } + _, aggregateSchema, _, _, _, err := buildAggregatePlan(st, exec.NewScan(s, nil), s) + if err != nil { + t.Fatal(err) + } + want := []value.Type{value.TInt, value.TInt, value.TFloat, value.TFloat, value.TText, value.TBool} + for i, wantType := range want { + if aggregateSchema[i].Type != wantType { + t.Fatalf("schema[%d].Type = %v, want %v", i, aggregateSchema[i].Type, wantType) + } + } +} + +func TestBuildAggregatePlanRejectsInvalidSQL(t *testing.T) { + tests := []struct{ sql, want string }{ + {"SELECT name, COUNT(*) FROM users", "must appear in GROUP BY"}, + {"SELECT * FROM users GROUP BY age", "SELECT *"}, + {"SELECT SUM(name) FROM users", "SUM requires numeric"}, + {"SELECT MIDDLE(age) FROM users", "unknown aggregate"}, + {"SELECT SUM(COUNT(*)) FROM users", "nested aggregate"}, + {"SELECT SUM(COUNT(missing)) FROM users", "nested aggregate"}, + {"SELECT age FROM users GROUP BY age + 1", "GROUP BY requires columns"}, + {"SELECT COUNT(*) FROM users GROUP BY COUNT(*)", "GROUP BY requires columns"}, + {"SELECT age FROM users WHERE COUNT(*) > 0", "WHERE cannot contain aggregates"}, + {"SELECT age FROM users HAVING age > 0", "HAVING requires an aggregate query"}, + } + for _, tt := range tests { + t.Run(tt.sql, func(t *testing.T) { + st := parsePlanSelect(t, tt.sql) + s := exec.Schema{{Table: "users", Name: "age", Type: value.TInt}, {Table: "users", Name: "name", Type: value.TText}} + _, _, _, _, _, err := buildAggregatePlan(st, exec.NewScan(s, nil), s) + if err == nil || !strings.Contains(err.Error(), tt.want) { + t.Fatalf("error = %v, want %q", err, tt.want) + } + }) + } +} + +func assertInternalAggregateExpr(t *testing.T, expr ast.Expr) { + t.Helper() + switch n := expr.(type) { + case *ast.Literal, *ast.SlotRef: + return + case *ast.UnaryExpr: + assertInternalAggregateExpr(t, n.Expr) + case *ast.BinaryExpr: + assertInternalAggregateExpr(t, n.Left) + assertInternalAggregateExpr(t, n.Right) + case *ast.IsNull: + assertInternalAggregateExpr(t, n.Expr) + default: + t.Fatalf("post-aggregate expression contains %T", expr) + } +} + +func slotInBinary(t *testing.T, expr ast.Expr, left bool) int { + t.Helper() + binary, ok := expr.(*ast.BinaryExpr) + if !ok { + t.Fatalf("expression = %T, want *ast.BinaryExpr", expr) + } + child := binary.Right + if left { + child = binary.Left + } + slot, ok := child.(*ast.SlotRef) + if !ok { + t.Fatalf("binary child = %T, want *ast.SlotRef", child) + } + return slot.Index +} diff --git a/internal/plan/typecheck.go b/internal/plan/typecheck.go index ba59504..2beb33d 100644 --- a/internal/plan/typecheck.go +++ b/internal/plan/typecheck.go @@ -21,6 +21,11 @@ func inferExprType(e ast.Expr, s exec.Schema) (value.Type, error) { return 0, err } return s[i].Type, nil + case *ast.SlotRef: + if n.Index < 0 || n.Index >= len(s) { + return 0, fmt.Errorf("slot %d out of range", n.Index) + } + return s[n.Index].Type, nil case *ast.UnaryExpr: t, err := inferExprType(n.Expr, s) if err != nil { From d571ea7eb6e311a4cfe421c13c7da0828e8ed7a8 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Aybars=20Mete=20Kele=C5=9F?= Date: Tue, 4 Aug 2026 02:37:40 +0300 Subject: [PATCH 6/7] feat: plan aggregate queries end to end --- internal/exec/eval.go | 5 ++ internal/plan/aggregate_integration_test.go | 97 +++++++++++++++++++++ internal/plan/plan.go | 39 +++++++++ internal/repl/repl_test.go | 29 ++++++ 4 files changed, 170 insertions(+) create mode 100644 internal/plan/aggregate_integration_test.go diff --git a/internal/exec/eval.go b/internal/exec/eval.go index 9238872..56c0a37 100644 --- a/internal/exec/eval.go +++ b/internal/exec/eval.go @@ -18,6 +18,11 @@ func Eval(e ast.Expr, row value.Row, s Schema) (value.Value, error) { return value.Value{}, err } return row[i], nil + case *ast.SlotRef: + if n.Index < 0 || n.Index >= len(row) { + return value.Value{}, fmt.Errorf("slot %d out of range", n.Index) + } + return row[n.Index], nil case *ast.IsNull: v, err := Eval(n.Expr, row, s) if err != nil { diff --git a/internal/plan/aggregate_integration_test.go b/internal/plan/aggregate_integration_test.go new file mode 100644 index 0000000..d3d1f5d --- /dev/null +++ b/internal/plan/aggregate_integration_test.go @@ -0,0 +1,97 @@ +package plan + +import ( + "fmt" + "os" + "path/filepath" + "reflect" + "strings" + "testing" + + "github.com/aybavs/sql-query-engine/internal/ast" + "github.com/aybavs/sql-query-engine/internal/catalog" + "github.com/aybavs/sql-query-engine/internal/exec" + "github.com/aybavs/sql-query-engine/internal/lexer" + "github.com/aybavs/sql-query-engine/internal/parser" + "github.com/aybavs/sql-query-engine/internal/value" +) + +func aggregateFixture(t *testing.T) (string, *catalog.Catalog) { + t.Helper() + dir := t.TempDir() + if err := os.WriteFile(filepath.Join(dir, "users.csv"), []byte("1,alice,30,ankara\n2,bob,15,izmir\n3,carol,40,ankara\n4,dave,,bursa\n"), 0o644); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(dir, "orders.csv"), []byte("1,1,10.5\n2,1,20\n3,2,7\n"), 0o644); err != nil { + t.Fatal(err) + } + cat := catalog.New() + cat.Add(&catalog.Table{Name: "users", File: "users.csv", Columns: []catalog.Column{{Name: "id", Type: value.TInt}, {Name: "name", Type: value.TText}, {Name: "age", Type: value.TInt}, {Name: "city", Type: value.TText}}}) + cat.Add(&catalog.Table{Name: "orders", File: "orders.csv", Columns: []catalog.Column{{Name: "id", Type: value.TInt}, {Name: "user_id", Type: value.TInt}, {Name: "total", Type: value.TFloat}}}) + return dir, cat +} + +func TestAggregateSlotRefEvaluationRejectsOutOfRangeIndexes(t *testing.T) { + row := value.Row{value.Int64(7)} + for _, index := range []int{-1, len(row)} { + _, err := exec.Eval(&ast.SlotRef{Index: index}, row, nil) + want := fmt.Sprintf("slot %d out of range", index) + if err == nil || err.Error() != want { + t.Fatalf("Eval slot %d error = %v, want %q", index, err, want) + } + } +} + +func TestAggregateQueriesEndToEnd(t *testing.T) { + dir, cat := aggregateFixture(t) + tests := []struct { + name, sql string + want [][]string + }{ + {"global", "SELECT COUNT(*), COUNT(age), SUM(age), AVG(age), MIN(age), MAX(age) FROM users", [][]string{{"4", "3", "85", "28.333333333333332", "15", "40"}}}, + {"aggregate argument expression", "SELECT SUM(age + 1) FROM users", [][]string{{"88"}}}, + {"group having order limit", "SELECT city, COUNT(*), AVG(age) FROM users GROUP BY city HAVING COUNT(*) >= 2 ORDER BY AVG(age) DESC LIMIT 1", [][]string{{"ankara", "2", "35"}}}, + {"multi column group and result expression", "SELECT city, age, COUNT(*) + 1 FROM users GROUP BY city, age ORDER BY city, age", [][]string{{"ankara", "30", "2"}, {"ankara", "40", "2"}, {"bursa", "NULL", "2"}, {"izmir", "15", "2"}}}, + {"hidden having aggregate", "SELECT city FROM users GROUP BY city HAVING COUNT(*) >= 2 ORDER BY city", [][]string{{"ankara"}}}, + {"hidden order aggregate", "SELECT city, SUM(age) FROM users GROUP BY city ORDER BY COUNT(*) DESC, city", [][]string{{"ankara", "70"}, {"bursa", "NULL"}, {"izmir", "15"}}}, + {"join aggregate", "SELECT users.name, SUM(orders.total) FROM users JOIN orders ON users.id = orders.user_id GROUP BY users.name ORDER BY SUM(orders.total) DESC", [][]string{{"alice", "30.5"}, {"bob", "7"}}}, + {"grouped empty", "SELECT city, COUNT(*) FROM users WHERE age > 100 GROUP BY city", nil}, + {"global empty", "SELECT COUNT(*), SUM(age), AVG(age), MIN(age), MAX(age) FROM users WHERE age > 100", [][]string{{"0", "NULL", "NULL", "NULL", "NULL"}}}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := buildAndRun(t, tt.sql, dir, cat) + if !reflect.DeepEqual(got, tt.want) { + t.Fatalf("result = %v, want %v", got, tt.want) + } + }) + } +} + +func TestAggregateQueriesRejectInvalidSQL(t *testing.T) { + dir, cat := aggregateFixture(t) + tests := []struct{ sql, want string }{ + {"SELECT age FROM users WHERE COUNT(*) > 0", "WHERE cannot contain aggregates"}, + {"SELECT COUNT(*) FROM users HAVING COUNT(*)", "HAVING requires BOOL"}, + {"SELECT name, COUNT(*) FROM users", "must appear in GROUP BY"}, + {"SELECT SUM(*) FROM users", "only COUNT accepts"}, + {"SELECT SUM(COUNT(*)) FROM users", "nested aggregate"}, + {"SELECT MIDDLE(age) FROM users", "unknown aggregate"}, + {"SELECT age FROM users HAVING age > 0", "HAVING requires an aggregate query"}, + } + for _, tt := range tests { + t.Run(tt.sql, func(t *testing.T) { + tokens, err := lexer.Lex(tt.sql) + if err == nil { + var statement *ast.SelectStmt + statement, err = parser.New(tokens).ParseSelect() + if err == nil { + _, _, err = Build(statement, cat, dir) + } + } + if err == nil || !strings.Contains(err.Error(), tt.want) { + t.Fatalf("error = %v, want %q", err, tt.want) + } + }) + } +} diff --git a/internal/plan/plan.go b/internal/plan/plan.go index e3532d7..b03ef7a 100644 --- a/internal/plan/plan.go +++ b/internal/plan/plan.go @@ -37,6 +37,14 @@ func Build(st *ast.SelectStmt, cat *catalog.Catalog, dataDir string) (exec.Opera schema = op.Schema() } + if containsAggregate(st.Where) { + return nil, nil, fmt.Errorf("WHERE cannot contain aggregates") + } + aggregateQuery := isAggregateQuery(st) + if st.Having != nil && !aggregateQuery { + return nil, nil, fmt.Errorf("HAVING requires an aggregate query") + } + if st.Where != nil { if err := requireBool(st.Where, schema, "WHERE"); err != nil { return nil, nil, err @@ -44,6 +52,37 @@ func Build(st *ast.SelectStmt, cat *catalog.Catalog, dataDir string) (exec.Opera op = exec.NewFilter(op, st.Where) } + if aggregateQuery { + aggOp, aggSchema, loweredProjections, loweredHaving, loweredOrderBy, err := buildAggregatePlan(st, op, schema) + if err != nil { + return nil, nil, err + } + op, schema = aggOp, aggSchema + if loweredHaving != nil { + op = exec.NewFilter(op, loweredHaving) + } + if len(loweredOrderBy) > 0 { + keys := make([]exec.SortKey, 0, len(loweredOrderBy)) + for _, item := range loweredOrderBy { + keys = append(keys, exec.SortKey{Expr: item.Expr, Desc: item.Desc}) + } + op = exec.NewSort(op, keys) + } + out := make(exec.Schema, len(loweredProjections)) + for i, e := range loweredProjections { + t, err := inferExprType(e, schema) + if err != nil { + return nil, nil, err + } + out[i] = exec.Column{Name: exprName(st.Projections[i].Expr), Type: t} + } + op = exec.NewProject(op, loweredProjections, out) + if st.Limit != nil { + op = exec.NewLimit(op, *st.Limit) + } + return op, out, nil + } + if len(st.OrderBy) > 0 { keys := make([]exec.SortKey, 0, len(st.OrderBy)) for _, o := range st.OrderBy { diff --git a/internal/repl/repl_test.go b/internal/repl/repl_test.go index 99bb963..7cc3eef 100644 --- a/internal/repl/repl_test.go +++ b/internal/repl/repl_test.go @@ -66,3 +66,32 @@ func TestLoadSchemaRejectsBadLine(t *testing.T) { t.Fatal("expected error for malformed schema line") } } + +func TestReplRunsAggregateQuery(t *testing.T) { + dir := t.TempDir() + if err := os.WriteFile(filepath.Join(dir, "users.csv"), []byte("1,alice,30\n2,bob,15\n"), 0o644); err != nil { + t.Fatal(err) + } + cat := catalog.New() + cat.Add(&catalog.Table{Name: "users", File: "users.csv", Columns: []catalog.Column{{Name: "id", Type: value.TInt}, {Name: "name", Type: value.TText}, {Name: "age", Type: value.TInt}}}) + in := strings.NewReader("SELECT COUNT(*), AVG(age) FROM users\n") + var out bytes.Buffer + Run(cat, dir, in, &out) + if !strings.Contains(out.String(), "2") || !strings.Contains(out.String(), "22.5") { + t.Fatalf("output = %q", out.String()) + } +} + +func TestReplReportsAggregatePlannerError(t *testing.T) { + dir := t.TempDir() + if err := os.WriteFile(filepath.Join(dir, "users.csv"), []byte("1,alice,30\n"), 0o644); err != nil { + t.Fatal(err) + } + cat := catalog.New() + cat.Add(&catalog.Table{Name: "users", File: "users.csv", Columns: []catalog.Column{{Name: "id", Type: value.TInt}, {Name: "name", Type: value.TText}, {Name: "age", Type: value.TInt}}}) + var out bytes.Buffer + Run(cat, dir, strings.NewReader("SELECT name, COUNT(*) FROM users\n"), &out) + if !strings.Contains(out.String(), "GROUP BY") { + t.Fatalf("output = %q", out.String()) + } +} From 0ac65fca1715711bd088c4bd89c1a519536a9815 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Aybars=20Mete=20Kele=C5=9F?= Date: Tue, 4 Aug 2026 03:00:15 +0300 Subject: [PATCH 7/7] fix: compare large numeric values exactly --- internal/exec/aggregate_test.go | 27 +++++++++++++++++ internal/value/logic.go | 53 +++++++++++++++++++++++++++++---- internal/value/logic_test.go | 51 +++++++++++++++++++++++++++---- 3 files changed, 120 insertions(+), 11 deletions(-) diff --git a/internal/exec/aggregate_test.go b/internal/exec/aggregate_test.go index 49e1012..256f7c0 100644 --- a/internal/exec/aggregate_test.go +++ b/internal/exec/aggregate_test.go @@ -44,6 +44,33 @@ func TestAggregateGlobalAllFunctions(t *testing.T) { } } +func TestAggregateMinMaxPreserveLargeIntegerOrder(t *testing.T) { + const ( + smaller int64 = 9007199254740992 + larger int64 = 9007199254740993 + ) + in := Schema{{Name: "min_n", Type: value.TInt}, {Name: "max_n", Type: value.TInt}} + rows := []value.Row{ + {value.Int64(larger), value.Int64(smaller)}, + {value.Int64(smaller), value.Int64(larger)}, + } + specs := []AggregateSpec{ + {Kind: AggMin, Expr: &ast.ColumnRef{Name: "min_n"}, OutType: value.TInt}, + {Kind: AggMax, Expr: &ast.ColumnRef{Name: "max_n"}, OutType: value.TInt}, + } + + got := drainRows(NewAggregate(NewScan(in, rows), nil, specs, Schema{{Type: value.TInt}, {Type: value.TInt}})) + if len(got) != 1 { + t.Fatalf("row count = %d, want 1", len(got)) + } + if got[0][0] != value.Int64(smaller) { + t.Errorf("MIN = %v, want %d", got[0][0], smaller) + } + if got[0][1] != value.Int64(larger) { + t.Errorf("MAX = %v, want %d", got[0][1], larger) + } +} + func TestAggregateGroupsInFirstSeenOrder(t *testing.T) { in := Schema{{Name: "city", Type: value.TText}, {Name: "n", Type: value.TInt}} rows := []value.Row{{value.Text("b"), value.Int64(2)}, {value.Text("a"), value.Int64(3)}, {value.Text("b"), value.Int64(4)}} diff --git a/internal/value/logic.go b/internal/value/logic.go index aa713e0..39d475a 100644 --- a/internal/value/logic.go +++ b/internal/value/logic.go @@ -10,7 +10,16 @@ func Compare(a, b Value) (ord int, known bool) { return 0, false } if isNumeric(a.Type) && isNumeric(b.Type) { - return cmpFloat(toFloat(a), toFloat(b)), true + switch { + case a.Type == TInt && b.Type == TInt: + return cmpInt(a.I, b.I), true + case a.Type == TFloat && b.Type == TFloat: + return cmpFloat(a.F, b.F), true + case a.Type == TInt: + return cmpIntFloat(a.I, b.F), true + default: + return -cmpIntFloat(b.I, a.F), true + } } switch a.Type { case TText: @@ -23,11 +32,45 @@ func Compare(a, b Value) (ord int, known bool) { func isNumeric(t Type) bool { return t == TInt || t == TFloat } -func toFloat(v Value) float64 { - if v.Type == TInt { - return float64(v.I) +func cmpInt(a, b int64) int { + switch { + case a < b: + return -1 + case a > b: + return 1 + default: + return 0 + } +} + +func cmpIntFloat(i int64, f float64) int { + const ( + minInt64Float = -1 << 63 + maxInt64FloatCutoff = 1 << 63 + ) + + // Screen values outside int64's range before converting f. These checks + // also order negative and positive infinity without a special case. + switch { + case f >= maxInt64FloatCutoff: + return -1 + case f < minInt64Float: + return 1 + } + + truncated := int64(f) + switch { + case i < truncated: + return -1 + case i > truncated: + return 1 + case f == float64(truncated): + return 0 + case f > 0: + return -1 + default: + return 1 } - return v.F } func cmpFloat(a, b float64) int { diff --git a/internal/value/logic_test.go b/internal/value/logic_test.go index 5e9c537..7af440c 100644 --- a/internal/value/logic_test.go +++ b/internal/value/logic_test.go @@ -6,17 +6,56 @@ import ( ) func TestCompareNumeric(t *testing.T) { - if ord, known := Compare(Int64(1), Int64(2)); !known || ord != -1 { - t.Fatalf("1 vs 2 = (%d,%v)", ord, known) + tests := []struct { + name string + a, b Value + want int + }{ + {name: "ordinary ints", a: Int64(1), b: Int64(2), want: -1}, + {name: "positive 2^53 adjacent ints", a: Int64(9007199254740993), b: Int64(9007199254740992), want: 1}, + {name: "negative 2^53 adjacent ints", a: Int64(-9007199254740993), b: Int64(-9007199254740992), want: -1}, + {name: "int64 max and predecessor", a: Int64(math.MaxInt64), b: Int64(math.MaxInt64 - 1), want: 1}, + {name: "int64 min and successor", a: Int64(math.MinInt64), b: Int64(math.MinInt64 + 1), want: -1}, + {name: "equal int and float", a: Int64(2), b: Float64(2), want: 0}, + {name: "equal at positive 2^53", a: Int64(9007199254740992), b: Float64(9007199254740992), want: 0}, + {name: "int above rounded positive float", a: Int64(9007199254740993), b: Float64(9007199254740992), want: 1}, + {name: "int below rounded negative float", a: Int64(-9007199254740993), b: Float64(-9007199254740992), want: -1}, + {name: "int64 max below 2^63 float", a: Int64(math.MaxInt64), b: Float64(1 << 63), want: -1}, + {name: "int64 min equals negative 2^63 float", a: Int64(math.MinInt64), b: Float64(-1 << 63), want: 0}, + {name: "int64 min successor above negative 2^63 float", a: Int64(math.MinInt64 + 1), b: Float64(-1 << 63), want: 1}, + {name: "positive fraction", a: Int64(1), b: Float64(1.5), want: -1}, + {name: "negative fraction", a: Int64(-1), b: Float64(-1.5), want: 1}, + {name: "float before int reverses order", a: Float64(1 << 63), b: Int64(math.MaxInt64), want: 1}, + {name: "positive infinity", a: Int64(math.MaxInt64), b: Float64(math.Inf(1)), want: -1}, + {name: "negative infinity", a: Int64(math.MinInt64), b: Float64(math.Inf(-1)), want: 1}, + {name: "int zero equals negative zero", a: Int64(0), b: Float64(math.Copysign(0, -1)), want: 0}, + {name: "negative zero equals int zero", a: Float64(math.Copysign(0, -1)), b: Int64(0), want: 0}, + {name: "float signed zeros compare equal", a: Float64(math.Copysign(0, -1)), b: Float64(0), want: 0}, } - if ord, known := Compare(Int64(2), Float64(2.0)); !known || ord != 0 { - t.Fatalf("int/float coercion failed: (%d,%v)", ord, known) + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if ord, known := Compare(tt.a, tt.b); !known || ord != tt.want { + t.Fatalf("Compare(%v, %v) = (%d, %v), want (%d, true)", tt.a, tt.b, ord, known, tt.want) + } + }) } } func TestCompareNullUnknown(t *testing.T) { - if _, known := Compare(Int64(1), NullOf(TInt)); known { - t.Fatal("comparison with NULL must be unknown") + tests := []struct { + name string + a, b Value + }{ + {name: "NULL on right", a: Int64(1), b: NullOf(TInt)}, + {name: "NULL on left", a: NullOf(TFloat), b: Int64(1)}, + {name: "both NULL", a: NullOf(TInt), b: NullOf(TFloat)}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if ord, known := Compare(tt.a, tt.b); known { + t.Fatalf("Compare() = (%d, true), want (_, false)", ord) + } + }) } }