Skip to content

Commit 673f56d

Browse files
authored
googlesql: type inference via the new analysis core (analyze only) (#4533)
1 parent af1016f commit 673f56d

26 files changed

Lines changed: 777 additions & 18 deletions

File tree

docs/howto/analyze.md

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -19,7 +19,7 @@ provided. The schema is always read from the `--schema` file.
1919
## Flags
2020

2121
- `--dialect`, `-d` - The SQL dialect to use. One of `postgresql`, `mysql`,
22-
`sqlite`, or `clickhouse`. Required.
22+
`sqlite`, `clickhouse`, or `googlesql`. Required.
2323
- `--schema`, `-s` - Path to the schema (DDL) file. Required.
2424
- `--ast` - Include each statement's AST in the output. Defaults to `false`.
2525

docs/howto/parse.md

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -20,7 +20,7 @@ provided.
2020
## Flags
2121

2222
- `--dialect`, `-d` - The SQL dialect to use. One of `postgresql`, `mysql`,
23-
`sqlite`, or `clickhouse`. Required.
23+
`sqlite`, `clickhouse`, or `googlesql`. Required.
2424

2525
## Examples
2626

internal/cmd/analyze.go

Lines changed: 8 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -40,6 +40,9 @@ Examples:
4040
# Analyze a ClickHouse query
4141
sqlc analyze --dialect clickhouse --schema schema.sql query.sql
4242
43+
# Analyze a GoogleSQL (BigQuery, Spanner) query
44+
sqlc analyze --dialect googlesql --schema schema.sql query.sql
45+
4346
# Analyze a query piped via stdin
4447
echo "-- name: GetAuthor :one
4548
SELECT * FROM authors WHERE id = $1;" | sqlc analyze --dialect postgresql --schema schema.sql
@@ -53,7 +56,7 @@ Examples:
5356
return err
5457
}
5558
if dialect == "" {
56-
return fmt.Errorf("--dialect flag is required (postgresql, mysql, sqlite, or clickhouse)")
59+
return fmt.Errorf("--dialect flag is required (postgresql, mysql, sqlite, clickhouse, or googlesql)")
5760
}
5861

5962
schemaPath, err := cmd.Flags().GetString("schema")
@@ -112,8 +115,10 @@ Examples:
112115
engine = config.EngineSQLite
113116
case "clickhouse":
114117
engine = config.EngineClickHouse
118+
case "googlesql":
119+
engine = config.EngineGoogleSQL
115120
default:
116-
return fmt.Errorf("unsupported dialect: %s (use postgresql, mysql, sqlite, or clickhouse)", dialect)
121+
return fmt.Errorf("unsupported dialect: %s (use postgresql, mysql, sqlite, clickhouse, or googlesql)", dialect)
117122
}
118123

119124
sql := config.SQL{
@@ -155,7 +160,7 @@ Examples:
155160
return nil
156161
},
157162
}
158-
cmd.Flags().StringP("dialect", "d", "", "SQL dialect to use (postgresql, mysql, sqlite, or clickhouse)")
163+
cmd.Flags().StringP("dialect", "d", "", "SQL dialect to use (postgresql, mysql, sqlite, clickhouse, or googlesql)")
159164
cmd.Flags().StringP("schema", "s", "", "path to the schema file")
160165
cmd.Flags().BoolP("ast", "", false, "include the statement AST in the output")
161166
return cmd

internal/compiler/engine.go

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@ import (
1010
"github.com/sqlc-dev/sqlc/internal/dbmanager"
1111
"github.com/sqlc-dev/sqlc/internal/engine/clickhouse"
1212
"github.com/sqlc-dev/sqlc/internal/engine/dolphin"
13+
"github.com/sqlc-dev/sqlc/internal/engine/googlesql"
1314
"github.com/sqlc-dev/sqlc/internal/engine/postgresql"
1415
pganalyze "github.com/sqlc-dev/sqlc/internal/engine/postgresql/analyzer"
1516
"github.com/sqlc-dev/sqlc/internal/engine/sqlite"
@@ -123,6 +124,14 @@ func NewCompiler(conf config.SQL, combo config.CombinedSettings, parserOpts opts
123124
return nil, fmt.Errorf("clickhouse: init catalog: %w", err)
124125
}
125126
c.coreCatalog = cat
127+
case config.EngineGoogleSQL:
128+
c.parser = googlesql.NewParser()
129+
c.selector = newDefaultSelector()
130+
cat, err := core.New(googlesql.Dialect())
131+
if err != nil {
132+
return nil, fmt.Errorf("googlesql: init catalog: %w", err)
133+
}
134+
c.coreCatalog = cat
126135
default:
127136
return nil, fmt.Errorf("unknown engine: %s", conf.Engine)
128137
}

internal/compiler/parse_core.go

Lines changed: 25 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,8 @@ import (
99
"github.com/sqlc-dev/sqlc/internal/metadata"
1010
"github.com/sqlc-dev/sqlc/internal/source"
1111
"github.com/sqlc-dev/sqlc/internal/sql/ast"
12+
"github.com/sqlc-dev/sqlc/internal/sql/named"
13+
"github.com/sqlc-dev/sqlc/internal/sql/rewrite"
1214
"github.com/sqlc-dev/sqlc/internal/sql/validate"
1315
)
1416

@@ -46,22 +48,33 @@ func (c *Compiler) parseQueryCore(stmt ast.Node, src string) (*Query, error) {
4648
return nil, err
4749
}
4850

51+
numbers, dollar, err := validate.ParamRef(raw)
52+
if err != nil {
53+
return nil, err
54+
}
55+
rewritten, namedParams, edits := rewrite.NamedParameters(c.conf.Engine, raw, numbers, dollar)
56+
expanded, err := source.Mutate(rawSQL, edits)
57+
if err != nil {
58+
return nil, err
59+
}
60+
4961
var cols []*Column
5062
var params []Parameter
51-
if _, ok := raw.Stmt.(*ast.SelectStmt); ok {
52-
res, err := coreanalyzer.Prepare(c.coreCatalog, raw)
63+
switch rewritten.Stmt.(type) {
64+
case *ast.SelectStmt, *ast.InsertStmt, *ast.UpdateStmt, *ast.DeleteStmt:
65+
res, err := coreanalyzer.Prepare(c.coreCatalog, rewritten)
5366
if err != nil {
5467
return nil, err
5568
}
5669
for _, col := range res.Columns {
5770
cols = append(cols, coreColumn(col))
5871
}
5972
for _, p := range res.Parameters {
60-
params = append(params, Parameter{Number: p.Number, Column: coreParamColumn(p)})
73+
params = append(params, Parameter{Number: p.Number, Column: coreParamColumn(p, namedParams)})
6174
}
6275
}
6376

64-
trimmed, comments, err := source.StripComments(rawSQL)
77+
trimmed, comments, err := source.StripComments(expanded)
6578
if err != nil {
6679
return nil, err
6780
}
@@ -100,7 +113,7 @@ func coreColumn(c core.Column) *Column {
100113
return col
101114
}
102115

103-
func coreParamColumn(p core.Parameter) *Column {
116+
func coreParamColumn(p core.Parameter, params *named.ParamSet) *Column {
104117
col := &Column{
105118
Name: p.Name,
106119
DataType: p.DataType,
@@ -110,5 +123,12 @@ func coreParamColumn(p core.Parameter) *Column {
110123
col.Table = &ast.TableName{Schema: p.Source.Schema, Name: p.Source.Table}
111124
col.OriginalName = p.Source.Column
112125
}
126+
if col.Name == "" {
127+
if name, ok := params.NameFor(p.Number); ok && name != "" {
128+
col.Name = name
129+
} else if p.Source != nil {
130+
col.Name = p.Source.Column
131+
}
132+
}
113133
return col
114134
}

internal/config/config.go

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -55,6 +55,7 @@ const (
5555
EnginePostgreSQL Engine = "postgresql"
5656
EngineSQLite Engine = "sqlite"
5757
EngineClickHouse Engine = "clickhouse"
58+
EngineGoogleSQL Engine = "googlesql"
5859
)
5960

6061
type Config struct {

internal/core/analyzer/analyzer.go

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -21,6 +21,21 @@ func Prepare(cat *core.Catalog, stmt ast.Node) (core.PrepareResult, error) {
2121
return core.PrepareResult{}, err
2222
}
2323
a.command = core.CommandSelect
24+
case *ast.InsertStmt:
25+
if err := a.analyzeInsert(s); err != nil {
26+
return core.PrepareResult{}, err
27+
}
28+
a.command = core.CommandInsert
29+
case *ast.UpdateStmt:
30+
if err := a.analyzeUpdate(s); err != nil {
31+
return core.PrepareResult{}, err
32+
}
33+
a.command = core.CommandUpdate
34+
case *ast.DeleteStmt:
35+
if err := a.analyzeDelete(s); err != nil {
36+
return core.PrepareResult{}, err
37+
}
38+
a.command = core.CommandDelete
2439
default:
2540
return core.PrepareResult{}, fmt.Errorf("analyzer: unsupported statement %T", stmt)
2641
}

internal/core/analyzer/dml.go

Lines changed: 185 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,185 @@
1+
package analyzer
2+
3+
import (
4+
"fmt"
5+
6+
"github.com/sqlc-dev/sqlc/internal/sql/ast"
7+
)
8+
9+
func (a *analyzer) analyzeInsert(s *ast.InsertStmt) error {
10+
if s.Relation == nil {
11+
return fmt.Errorf("insert: missing relation")
12+
}
13+
rel, err := a.bindRangeVar(s.Relation)
14+
if err != nil {
15+
return err
16+
}
17+
a.scope = &scope{rels: []scopeRel{rel}}
18+
19+
targets, err := insertTargets(rel, s.Cols)
20+
if err != nil {
21+
return err
22+
}
23+
if err := a.bindInsertValues(s.SelectStmt, rel, targets); err != nil {
24+
return err
25+
}
26+
return a.projectReturning(s.ReturningList)
27+
}
28+
29+
func (a *analyzer) analyzeUpdate(s *ast.UpdateStmt) error {
30+
sc, err := a.relationScope(s.Relations, s.FromClause)
31+
if err != nil {
32+
return err
33+
}
34+
a.scope = sc
35+
target := sc.rels[0]
36+
37+
for _, item := range listItems(s.TargetList) {
38+
rt, ok := item.(*ast.ResTarget)
39+
if !ok || rt.Name == nil {
40+
continue
41+
}
42+
col, ok := findColumn(target, *rt.Name)
43+
if !ok {
44+
return fmt.Errorf("unknown column %q", *rt.Name)
45+
}
46+
if err := a.bindValue(target, &col, rt.Val); err != nil {
47+
return fmt.Errorf("set %s: %w", *rt.Name, err)
48+
}
49+
}
50+
51+
if s.WhereClause != nil {
52+
if _, err := a.typeExpr(s.WhereClause); err != nil {
53+
return fmt.Errorf("where: %w", err)
54+
}
55+
}
56+
return a.projectReturning(s.ReturningList)
57+
}
58+
59+
func (a *analyzer) analyzeDelete(s *ast.DeleteStmt) error {
60+
sc, err := a.relationScope(s.Relations, s.UsingClause)
61+
if err != nil {
62+
return err
63+
}
64+
a.scope = sc
65+
66+
if s.WhereClause != nil {
67+
if _, err := a.typeExpr(s.WhereClause); err != nil {
68+
return fmt.Errorf("where: %w", err)
69+
}
70+
}
71+
return a.projectReturning(s.ReturningList)
72+
}
73+
74+
func (a *analyzer) relationScope(relations, extra *ast.List) (*scope, error) {
75+
sc := &scope{}
76+
for _, item := range listItems(relations) {
77+
if err := a.appendFromItem(sc, item); err != nil {
78+
return nil, err
79+
}
80+
}
81+
for _, item := range listItems(extra) {
82+
if err := a.appendFromItem(sc, item); err != nil {
83+
return nil, err
84+
}
85+
}
86+
if len(sc.rels) == 0 {
87+
return nil, fmt.Errorf("missing target relation")
88+
}
89+
return sc, nil
90+
}
91+
92+
func insertTargets(rel scopeRel, cols *ast.List) ([]scopeCol, error) {
93+
items := listItems(cols)
94+
if len(items) == 0 {
95+
return rel.cols, nil
96+
}
97+
out := make([]scopeCol, 0, len(items))
98+
for _, item := range items {
99+
rt, ok := item.(*ast.ResTarget)
100+
if !ok || rt.Name == nil {
101+
return nil, fmt.Errorf("insert: unsupported column target %T", item)
102+
}
103+
col, ok := findColumn(rel, *rt.Name)
104+
if !ok {
105+
return nil, fmt.Errorf("unknown column %q", *rt.Name)
106+
}
107+
out = append(out, col)
108+
}
109+
return out, nil
110+
}
111+
112+
func (a *analyzer) bindInsertValues(n ast.Node, rel scopeRel, targets []scopeCol) error {
113+
if n == nil {
114+
return nil
115+
}
116+
sel, ok := n.(*ast.SelectStmt)
117+
if !ok {
118+
return fmt.Errorf("insert: unsupported source %T", n)
119+
}
120+
if sel.ValuesLists == nil {
121+
return fmt.Errorf("insert: INSERT ... SELECT is not supported")
122+
}
123+
for _, row := range listItems(sel.ValuesLists) {
124+
values, ok := row.(*ast.List)
125+
if !ok {
126+
continue
127+
}
128+
for i, v := range values.Items {
129+
var target *scopeCol
130+
if i < len(targets) {
131+
target = &targets[i]
132+
}
133+
if err := a.bindValue(rel, target, v); err != nil {
134+
return err
135+
}
136+
}
137+
}
138+
return nil
139+
}
140+
141+
func (a *analyzer) bindValue(rel scopeRel, target *scopeCol, v ast.Node) error {
142+
if target != nil {
143+
switch value := v.(type) {
144+
case *ast.ParamRef:
145+
a.inferParam(value.Number, columnType(rel, *target))
146+
return nil
147+
case *ast.A_Const:
148+
return nil
149+
}
150+
}
151+
_, err := a.typeExpr(v)
152+
return err
153+
}
154+
155+
func (a *analyzer) projectReturning(l *ast.List) error {
156+
for _, item := range listItems(l) {
157+
rt, ok := item.(*ast.ResTarget)
158+
if !ok {
159+
continue
160+
}
161+
if err := a.projectTarget(rt); err != nil {
162+
return err
163+
}
164+
}
165+
return nil
166+
}
167+
168+
func findColumn(rel scopeRel, name string) (scopeCol, bool) {
169+
for _, col := range rel.cols {
170+
if col.name == name {
171+
return col, true
172+
}
173+
}
174+
return scopeCol{}, false
175+
}
176+
177+
func columnType(rel scopeRel, col scopeCol) exprType {
178+
return exprType{
179+
typeOID: col.typeOID,
180+
nullable: !col.notNull,
181+
sourceClassOID: rel.classOID,
182+
sourceAttributeOID: col.attOID,
183+
sourceTableAlias: rel.alias,
184+
}
185+
}

internal/core/analyzer/expr.go

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -276,6 +276,9 @@ func (a *analyzer) typeFuncCall(f *ast.FuncCall) (exprType, error) {
276276
}
277277

278278
if f.AggStar && (name == "count" || name == "count.*") {
279+
if overloads, err := a.cat.FindProcs("count", nil); err == nil && len(overloads) > 0 {
280+
return exprType{typeOID: overloads[0].ReturnTypeOID, nullable: overloads[0].ReturnNullable}, nil
281+
}
279282
oid, err := a.cat.TypeOID("int8")
280283
if err != nil {
281284
return exprType{}, err
Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,5 @@
1+
{
2+
"command": "analyze",
3+
"args": ["--dialect", "googlesql", "--schema", "schema.sql", "query.sql"],
4+
"contexts": ["base"]
5+
}

0 commit comments

Comments
 (0)