Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 5 additions & 4 deletions internal/core/analyzer/analyzer.go
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,7 @@ func Prepare(cat *core.Catalog, stmt ast.Node) (core.PrepareResult, error) {
}
a := &analyzer{
cat: cat,
params: map[int]*core.Parameter{},
params: map[int]core.Parameter{},
}
switch s := stmt.(type) {
case *ast.SelectStmt:
Expand Down Expand Up @@ -46,7 +46,7 @@ type analyzer struct {
cat *core.Catalog
scope *scope
columns []core.Column
params map[int]*core.Parameter
params map[int]core.Parameter
command core.Command
}

Expand All @@ -58,7 +58,7 @@ func (a *analyzer) result() core.PrepareResult {
}
}

func orderedParams(m map[int]*core.Parameter) []core.Parameter {
func orderedParams(m map[int]core.Parameter) []core.Parameter {
if len(m) == 0 {
return nil
}
Expand All @@ -71,7 +71,7 @@ func orderedParams(m map[int]*core.Parameter) []core.Parameter {
out := make([]core.Parameter, 0, len(m))
for i := 1; i <= maxN; i++ {
if p, ok := m[i]; ok {
out = append(out, *p)
out = append(out, p)
}
}
return out
Expand Down Expand Up @@ -112,6 +112,7 @@ func (a *analyzer) analyzeSelect(s *ast.SelectStmt) error {
if targets == nil {
return fmt.Errorf("select: empty target list")
}
a.columns = make([]core.Column, 0, len(targets))
for _, t := range targets {
rt, ok := t.(*ast.ResTarget)
if !ok {
Expand Down
125 changes: 125 additions & 0 deletions internal/core/analyzer/bench_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,125 @@
package analyzer_test

import (
"strings"
"testing"

"github.com/sqlc-dev/sqlc/internal/config"
"github.com/sqlc-dev/sqlc/internal/core"
"github.com/sqlc-dev/sqlc/internal/core/analyzer"
coreschema "github.com/sqlc-dev/sqlc/internal/core/schema"
"github.com/sqlc-dev/sqlc/internal/engine/googlesql"
"github.com/sqlc-dev/sqlc/internal/sql/ast"
"github.com/sqlc-dev/sqlc/internal/sql/rewrite"
"github.com/sqlc-dev/sqlc/internal/sql/validate"
)

const benchSchema = `
CREATE TABLE users (
id INT64 NOT NULL,
name STRING NOT NULL,
bio STRING,
) PRIMARY KEY (id);

CREATE TABLE posts (
id INT64 NOT NULL,
user_id INT64 NOT NULL,
title STRING(255),
created TIMESTAMP NOT NULL,
) PRIMARY KEY (id);
`

const benchQueries = `
SELECT * FROM users;
SELECT COUNT(*) AS total FROM users;
SELECT u.name, p.* FROM users u JOIN posts p ON p.user_id = u.id WHERE u.name = @name;
SELECT id, name, bio FROM users WHERE id = @id AND name = @name;
SELECT p.title, p.created FROM posts p WHERE p.user_id = @uid AND p.title = @t;
`

func parseAll(t testing.TB, src string) []ast.Node {
p := googlesql.NewParser()
stmts, err := p.Parse(strings.NewReader(src))
if err != nil {
t.Fatal(err)
}
out := make([]ast.Node, 0, len(stmts))
for i := range stmts {
out = append(out, stmts[i].Raw)
}
return out
}

// parseQueries mirrors the compiler pipeline: parse, then rewrite named
// parameters into ParamRef nodes before handing the statement to the analyzer.
func parseQueries(t testing.TB, src string) []ast.Node {
out := make([]ast.Node, 0)
for _, n := range parseAll(t, src) {
raw, ok := n.(*ast.RawStmt)
if !ok {
t.Fatalf("not a raw statement: %T", n)
}
numbers, dollar, err := validate.ParamRef(raw)
if err != nil {
t.Fatal(err)
}
rewritten, _, _ := rewrite.NamedParameters(config.EngineGoogleSQL, raw, numbers, dollar)
out = append(out, rewritten)
}
return out
}

func newBenchCatalog(t testing.TB) (*core.Catalog, []ast.Node) {
cat, err := core.New(googlesql.Dialect())
if err != nil {
t.Fatal(err)
}
for _, n := range parseAll(t, benchSchema) {
if err := coreschema.Apply(cat, n); err != nil {
t.Fatal(err)
}
}
return cat, parseQueries(t, benchQueries)
}

// BenchmarkCatalogNew measures per-compile catalog setup: open the SQLite
// catalog, install the schema, and run the dialect seed.
func BenchmarkCatalogNew(b *testing.B) {
for b.Loop() {
cat, err := core.New(googlesql.Dialect())
if err != nil {
b.Fatal(err)
}
cat.Close()
}
}

// BenchmarkApplySchema measures loading user DDL into the catalog.
func BenchmarkApplySchema(b *testing.B) {
stmts := parseAll(b, benchSchema)
for b.Loop() {
cat, err := core.New(googlesql.Dialect())
if err != nil {
b.Fatal(err)
}
for _, n := range stmts {
if err := coreschema.Apply(cat, n); err != nil {
b.Fatal(err)
}
}
cat.Close()
}
}

// BenchmarkPrepare measures steady-state query analysis against a warm catalog.
func BenchmarkPrepare(b *testing.B) {
cat, queries := newBenchCatalog(b)
defer cat.Close()
for b.Loop() {
for _, q := range queries {
if _, err := analyzer.Prepare(cat, q); err != nil {
b.Fatal(err)
}
}
}
}
25 changes: 13 additions & 12 deletions internal/core/analyzer/dml.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ package analyzer
import (
"fmt"

"github.com/sqlc-dev/sqlc/internal/core"
"github.com/sqlc-dev/sqlc/internal/sql/ast"
)

Expand Down Expand Up @@ -89,12 +90,12 @@ func (a *analyzer) relationScope(relations, extra *ast.List) (*scope, error) {
return sc, nil
}

func insertTargets(rel scopeRel, cols *ast.List) ([]scopeCol, error) {
func insertTargets(rel scopeRel, cols *ast.List) ([]core.ClassColumn, error) {
items := listItems(cols)
if len(items) == 0 {
return rel.cols, nil
}
out := make([]scopeCol, 0, len(items))
out := make([]core.ClassColumn, 0, len(items))
for _, item := range items {
rt, ok := item.(*ast.ResTarget)
if !ok || rt.Name == nil {
Expand All @@ -109,7 +110,7 @@ func insertTargets(rel scopeRel, cols *ast.List) ([]scopeCol, error) {
return out, nil
}

func (a *analyzer) bindInsertValues(n ast.Node, rel scopeRel, targets []scopeCol) error {
func (a *analyzer) bindInsertValues(n ast.Node, rel scopeRel, targets []core.ClassColumn) error {
if n == nil {
return nil
}
Expand All @@ -126,7 +127,7 @@ func (a *analyzer) bindInsertValues(n ast.Node, rel scopeRel, targets []scopeCol
continue
}
for i, v := range values.Items {
var target *scopeCol
var target *core.ClassColumn
if i < len(targets) {
target = &targets[i]
}
Expand All @@ -138,7 +139,7 @@ func (a *analyzer) bindInsertValues(n ast.Node, rel scopeRel, targets []scopeCol
return nil
}

func (a *analyzer) bindValue(rel scopeRel, target *scopeCol, v ast.Node) error {
func (a *analyzer) bindValue(rel scopeRel, target *core.ClassColumn, v ast.Node) error {
if target != nil {
switch value := v.(type) {
case *ast.ParamRef:
Expand All @@ -165,21 +166,21 @@ func (a *analyzer) projectReturning(l *ast.List) error {
return nil
}

func findColumn(rel scopeRel, name string) (scopeCol, bool) {
func findColumn(rel scopeRel, name string) (core.ClassColumn, bool) {
for _, col := range rel.cols {
if col.name == name {
if col.Name == name {
return col, true
}
}
return scopeCol{}, false
return core.ClassColumn{}, false
}

func columnType(rel scopeRel, col scopeCol) exprType {
func columnType(rel scopeRel, col core.ClassColumn) exprType {
return exprType{
typeOID: col.typeOID,
nullable: !col.notNull,
typeOID: col.TypeOID,
nullable: !col.NotNull,
sourceClassOID: rel.classOID,
sourceAttributeOID: col.attOID,
sourceAttributeOID: col.AttOID,
sourceTableAlias: rel.alias,
}
}
12 changes: 6 additions & 6 deletions internal/core/analyzer/expr.go
Original file line number Diff line number Diff line change
Expand Up @@ -113,10 +113,10 @@ func (a *analyzer) typeColumnRef(c *ast.ColumnRef) (exprType, error) {
return exprType{}, fmt.Errorf("unknown column %q", column)
}
return exprType{
typeOID: col.typeOID,
nullable: !col.notNull,
typeOID: col.TypeOID,
nullable: !col.NotNull,
sourceClassOID: rel.classOID,
sourceAttributeOID: col.attOID,
sourceAttributeOID: col.AttOID,
sourceTableAlias: rel.alias,
}, nil
}
Expand All @@ -141,7 +141,7 @@ func flattenFields(fields *ast.List) []string {
func (a *analyzer) typeParamRef(p *ast.ParamRef) (exprType, error) {
cur, ok := a.params[p.Number]
if !ok {
cur = &core.Parameter{Number: p.Number}
cur = core.Parameter{Number: p.Number}
a.params[p.Number] = cur
}
return exprType{typeOID: cur.TypeOID, nullable: !cur.NotNull}, nil
Expand All @@ -150,8 +150,7 @@ func (a *analyzer) typeParamRef(p *ast.ParamRef) (exprType, error) {
func (a *analyzer) inferParam(number int, t exprType) {
cur, ok := a.params[number]
if !ok {
cur = &core.Parameter{Number: number}
a.params[number] = cur
cur = core.Parameter{Number: number}
}
if cur.TypeOID == 0 && t.typeOID != 0 {
cur.TypeOID = t.typeOID
Expand All @@ -171,6 +170,7 @@ func (a *analyzer) inferParam(number int, t exprType) {
}
}
}
a.params[number] = cur
}

func (a *analyzer) typeAExpr(e *ast.A_Expr) (exprType, error) {
Expand Down
48 changes: 26 additions & 22 deletions internal/core/analyzer/projection.go
Original file line number Diff line number Diff line change
@@ -1,14 +1,20 @@
package analyzer

import (
"slices"

"github.com/sqlc-dev/sqlc/internal/core"
"github.com/sqlc-dev/sqlc/internal/sql/ast"
)

func (a *analyzer) projectTarget(rt *ast.ResTarget) error {
// A column reference's field list is flattened once here and threaded
// through the star check, the star expansion and the output name.
var fields []string
if cr, ok := rt.Val.(*ast.ColumnRef); ok {
if isStarRef(cr) {
a.emitStar(cr)
fields = flattenFields(cr.Fields)
if isStar(fields) {
a.emitStar(fields)
return nil
}
}
Expand All @@ -18,7 +24,7 @@ func (a *analyzer) projectTarget(rt *ast.ResTarget) error {
return err
}
col := core.Column{
Name: targetName(rt),
Name: targetName(rt, fields),
TypeOID: t.typeOID,
NotNull: !t.nullable,
SourceClassOID: t.sourceClassOID,
Expand Down Expand Up @@ -56,15 +62,14 @@ func (a *analyzer) decorateSource(col *core.Column, attOID int64, tableAlias str
col.IsAutoIncrement = ad.AutoIncrement
}

func targetName(rt *ast.ResTarget) string {
// targetName picks the output name for a target. fields is the already
// flattened field list when rt.Val is a column reference, and nil otherwise.
func targetName(rt *ast.ResTarget, fields []string) string {
if rt.Name != nil && *rt.Name != "" {
return *rt.Name
}
if cr, ok := rt.Val.(*ast.ColumnRef); ok {
parts := flattenFields(cr.Fields)
if len(parts) > 0 {
return parts[len(parts)-1]
}
if len(fields) > 0 {
return fields[len(fields)-1]
}
if fc, ok := rt.Val.(*ast.FuncCall); ok {
if name := funcCallName(fc); name != "" {
Expand All @@ -74,33 +79,32 @@ func targetName(rt *ast.ResTarget) string {
return "?column?"
}

func isStarRef(c *ast.ColumnRef) bool {
parts := flattenFields(c.Fields)
return len(parts) > 0 && parts[len(parts)-1] == "*"
func isStar(fields []string) bool {
return len(fields) > 0 && fields[len(fields)-1] == "*"
}

func (a *analyzer) emitStar(cr *ast.ColumnRef) {
parts := flattenFields(cr.Fields)
func (a *analyzer) emitStar(fields []string) {
relName := ""
if len(parts) > 1 {
relName = parts[0]
if len(fields) > 1 {
relName = fields[0]
}
for _, rel := range a.scope.rels {
if relName != "" && rel.alias != relName {
continue
}
a.columns = slices.Grow(a.columns, len(rel.cols))
for _, c := range rel.cols {
col := core.Column{
Name: c.name,
TypeOID: c.typeOID,
NotNull: c.notNull,
Name: c.Name,
TypeOID: c.TypeOID,
NotNull: c.NotNull,
SourceClassOID: rel.classOID,
SourceAttributeOID: c.attOID,
SourceAttributeOID: c.AttOID,
}
if name, err := a.cat.TypeName(c.typeOID); err == nil {
if name, err := a.cat.TypeName(c.TypeOID); err == nil {
col.DataType = name
}
a.decorateSource(&col, c.attOID, rel.alias)
a.decorateSource(&col, c.AttOID, rel.alias)
a.columns = append(a.columns, col)
}
}
Expand Down
Loading
Loading