diff --git a/CHANGELOG.md b/CHANGELOG.md index a60f81bd..f0a5bfb8 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -9,6 +9,8 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ### Added +- Added support for `LISTEN` and `NOTIFY` statements in PostgreSQL, including `bobgen-psql` parser support. Note: `LISTEN` requires a persistent connection and only registers the channel — receiving notifications must be implemented using your specific database driver. (thanks @manhrev) +- Added `bobgen-psql` and `bobgen-sqlite` parser support for top-level `VALUES` queries. MySQL support is partial. (thanks @manhrev) - Generated `dberrors` packages now include generic and per-table check-constraint errors for PostgreSQL, matched by constraint name for `pq` and `pgx` drivers. (thanks @keithbro-imx) ### Changed diff --git a/dialect/psql/dialect/listen.go b/dialect/psql/dialect/listen.go new file mode 100644 index 00000000..167984a4 --- /dev/null +++ b/dialect/psql/dialect/listen.go @@ -0,0 +1,20 @@ +package dialect + +import ( + "context" + "io" + + "github.com/stephenafamo/bob" +) + +// Trying to represent the listen query structure as documented in +// https://www.postgresql.org/docs/current/sql-listen.html +type ListenQuery struct { + Channel string +} + +func (l ListenQuery) WriteSQL(_ context.Context, w io.StringWriter, dl bob.Dialect, _ int) ([]any, error) { + w.WriteString("LISTEN ") + dl.WriteQuoted(w, l.Channel) + return nil, nil +} diff --git a/dialect/psql/dialect/notify.go b/dialect/psql/dialect/notify.go new file mode 100644 index 00000000..267a1e01 --- /dev/null +++ b/dialect/psql/dialect/notify.go @@ -0,0 +1,26 @@ +package dialect + +import ( + "context" + "io" + + "github.com/stephenafamo/bob" +) + +// Trying to represent the notify query structure as documented in +// https://www.postgresql.org/docs/current/sql-notify.html +type NotifyQuery struct { + Channel string + Payload string +} + +func (n NotifyQuery) WriteSQL(_ context.Context, w io.StringWriter, dl bob.Dialect, _ int) ([]any, error) { + w.WriteString("NOTIFY ") + dl.WriteQuoted(w, n.Channel) + if n.Payload != "" { + w.WriteString(", '") + w.WriteString(n.Payload) + w.WriteString("'") + } + return nil, nil +} diff --git a/dialect/psql/listen.go b/dialect/psql/listen.go new file mode 100644 index 00000000..b4f67188 --- /dev/null +++ b/dialect/psql/listen.go @@ -0,0 +1,19 @@ +package psql + +import ( + "github.com/stephenafamo/bob" + "github.com/stephenafamo/bob/dialect/psql/dialect" +) + +func Listen(mods ...bob.Mod[*dialect.ListenQuery]) bob.BaseQuery[*dialect.ListenQuery] { + q := &dialect.ListenQuery{} + for _, mod := range mods { + mod.Apply(q) + } + + return bob.BaseQuery[*dialect.ListenQuery]{ + Expression: q, + Dialect: dialect.Dialect, + QueryType: bob.QueryTypeListen, + } +} diff --git a/dialect/psql/listen_test.go b/dialect/psql/listen_test.go new file mode 100644 index 00000000..febaaa9b --- /dev/null +++ b/dialect/psql/listen_test.go @@ -0,0 +1,21 @@ +package psql_test + +import ( + "testing" + + "github.com/stephenafamo/bob/dialect/psql" + "github.com/stephenafamo/bob/dialect/psql/lm" + testutils "github.com/stephenafamo/bob/test/utils" +) + +func TestListen(t *testing.T) { + examples := testutils.Testcases{ + "simple": { + Query: psql.Listen(lm.Channel("my_channel")), + ExpectedSQL: `LISTEN "my_channel"`, + ExpectedArgs: nil, + }, + } + + testutils.RunTests(t, examples, formatter) +} diff --git a/dialect/psql/lm/qm.go b/dialect/psql/lm/qm.go new file mode 100644 index 00000000..67a8b6fc --- /dev/null +++ b/dialect/psql/lm/qm.go @@ -0,0 +1,12 @@ +package lm + +import ( + "github.com/stephenafamo/bob" + "github.com/stephenafamo/bob/dialect/psql/dialect" +) + +func Channel(name string) bob.Mod[*dialect.ListenQuery] { + return bob.ModFunc[*dialect.ListenQuery](func(q *dialect.ListenQuery) { + q.Channel = name + }) +} diff --git a/dialect/psql/nm/qm.go b/dialect/psql/nm/qm.go new file mode 100644 index 00000000..ec653525 --- /dev/null +++ b/dialect/psql/nm/qm.go @@ -0,0 +1,18 @@ +package nm + +import ( + "github.com/stephenafamo/bob" + "github.com/stephenafamo/bob/dialect/psql/dialect" +) + +func Channel(name string) bob.Mod[*dialect.NotifyQuery] { + return bob.ModFunc[*dialect.NotifyQuery](func(q *dialect.NotifyQuery) { + q.Channel = name + }) +} + +func Payload(p string) bob.Mod[*dialect.NotifyQuery] { + return bob.ModFunc[*dialect.NotifyQuery](func(q *dialect.NotifyQuery) { + q.Payload = p + }) +} diff --git a/dialect/psql/notify.go b/dialect/psql/notify.go new file mode 100644 index 00000000..21851ab5 --- /dev/null +++ b/dialect/psql/notify.go @@ -0,0 +1,19 @@ +package psql + +import ( + "github.com/stephenafamo/bob" + "github.com/stephenafamo/bob/dialect/psql/dialect" +) + +func Notify(mods ...bob.Mod[*dialect.NotifyQuery]) bob.BaseQuery[*dialect.NotifyQuery] { + q := &dialect.NotifyQuery{} + for _, mod := range mods { + mod.Apply(q) + } + + return bob.BaseQuery[*dialect.NotifyQuery]{ + Expression: q, + Dialect: dialect.Dialect, + QueryType: bob.QueryTypeNotify, + } +} diff --git a/dialect/psql/notify_test.go b/dialect/psql/notify_test.go new file mode 100644 index 00000000..385d2c66 --- /dev/null +++ b/dialect/psql/notify_test.go @@ -0,0 +1,26 @@ +package psql_test + +import ( + "testing" + + "github.com/stephenafamo/bob/dialect/psql" + "github.com/stephenafamo/bob/dialect/psql/nm" + testutils "github.com/stephenafamo/bob/test/utils" +) + +func TestNotify(t *testing.T) { + examples := testutils.Testcases{ + "simple": { + Query: psql.Notify(nm.Channel("my_channel")), + ExpectedSQL: `NOTIFY "my_channel"`, + ExpectedArgs: nil, + }, + "with payload": { + Query: psql.Notify(nm.Channel("my_channel"), nm.Payload("hello world")), + ExpectedSQL: `NOTIFY "my_channel", 'hello world'`, + ExpectedArgs: nil, + }, + } + + testutils.RunTests(t, examples, formatter) +} diff --git a/gen/bobgen-mysql/driver/parser/visitor.go b/gen/bobgen-mysql/driver/parser/visitor.go index f22ab268..42970d57 100644 --- a/gen/bobgen-mysql/driver/parser/visitor.go +++ b/gen/bobgen-mysql/driver/parser/visitor.go @@ -184,6 +184,8 @@ func (v *visitor) VisitSqlStatements(ctx *mysqlparser.SqlStatementsContext) any v.Err = fmt.Errorf("stmt %d: could not get columns in select statement, got %T", i, resp) return nil } + case *mysqlparser.ValuesStatementContext: + queryType = bob.QueryTypeValues } allresp = append(allresp, StmtInfo{ diff --git a/gen/bobgen-psql/driver/parser/mods_listen.go b/gen/bobgen-psql/driver/parser/mods_listen.go new file mode 100644 index 00000000..19fc3582 --- /dev/null +++ b/gen/bobgen-psql/driver/parser/mods_listen.go @@ -0,0 +1,11 @@ +package parser + +import ( + "fmt" + + pg "github.com/pganalyze/pg_query_go/v6" +) + +func (w *walker) modListenStatement(stmt *pg.Node_ListenStmt, _ nodeInfo) { + fmt.Fprintf(w.mods, "q.Channel = %q\n", stmt.ListenStmt.Conditionname) +} diff --git a/gen/bobgen-psql/driver/parser/mods_notify.go b/gen/bobgen-psql/driver/parser/mods_notify.go new file mode 100644 index 00000000..f60c9ad9 --- /dev/null +++ b/gen/bobgen-psql/driver/parser/mods_notify.go @@ -0,0 +1,14 @@ +package parser + +import ( + "fmt" + + pg "github.com/pganalyze/pg_query_go/v6" +) + +func (w *walker) modNotifyStatement(stmt *pg.Node_NotifyStmt, _ nodeInfo) { + fmt.Fprintf(w.mods, "q.Channel = %q\n", stmt.NotifyStmt.Conditionname) + if stmt.NotifyStmt.Payload != "" { + fmt.Fprintf(w.mods, "q.Payload = %q\n", stmt.NotifyStmt.Payload) + } +} diff --git a/gen/bobgen-psql/driver/parser/mods_values.go b/gen/bobgen-psql/driver/parser/mods_values.go new file mode 100644 index 00000000..9e4f57c5 --- /dev/null +++ b/gen/bobgen-psql/driver/parser/mods_values.go @@ -0,0 +1,49 @@ +package parser + +import ( + "fmt" + + pg "github.com/pganalyze/pg_query_go/v6" + "github.com/stephenafamo/bob/internal" +) + +func (w *walker) modValuesStatement(stmt *pg.Node_SelectStmt, info nodeInfo) { + if orderInfo, ok := info.children["SortClause"]; ok { + w.editRules = append(w.editRules, internal.RecordPoints( + int(orderInfo.start), + int(orderInfo.end)-1, + func(start, end int) error { + fmt.Fprintf(w.mods, "q.AppendOrder(EXPR.subExpr(%d, %d))\n", start, end) + return nil + }, + )...) + } + + if limitInfo, ok := info.children["LimitCount"]; ok { + w.editRules = append(w.editRules, internal.RecordPoints( + int(limitInfo.start), + int(limitInfo.end)-1, + func(start, end int) error { + switch stmt.SelectStmt.LimitOption { + case pg.LimitOption_LIMIT_OPTION_COUNT: + fmt.Fprintf(w.mods, "q.SetLimit(EXPR.subExpr(%d, %d))\n", start, end) + case pg.LimitOption_LIMIT_OPTION_WITH_TIES: + w.imports = append(w.imports, []string{"github.com/stephenafamo/bob/clause"}) + fmt.Fprintf(w.mods, "q.SetFetch(clause.Fetch{Count: EXPR.subExpr(%d, %d), WithTies: true})\n", start, end) + } + return nil + }, + )...) + } + + if offsetInfo, ok := info.children["LimitOffset"]; ok { + w.editRules = append(w.editRules, internal.RecordPoints( + int(offsetInfo.start), + int(offsetInfo.end)-1, + func(start, end int) error { + fmt.Fprintf(w.mods, "q.SetOffset(EXPR.subExpr(%d, %d))\n", start, end) + return nil + }, + )...) + } +} diff --git a/gen/bobgen-psql/driver/parser/parser.go b/gen/bobgen-psql/driver/parser/parser.go index fbcfe7b9..d87c2aac 100644 --- a/gen/bobgen-psql/driver/parser/parser.go +++ b/gen/bobgen-psql/driver/parser/parser.go @@ -89,6 +89,11 @@ func (p *Parser) ParseQuery(ctx context.Context, input string) (drivers.Query, e return drivers.Query{}, fmt.Errorf("expected 1 statement, got %d", len(parseResult.Stmts)) } + stmt := parseResult.Stmts[0] + qType := getQueryType(stmt.Stmt) + + var argTypes, resTypes []string + w := walker{ db: p.db, sharedSchema: p.sharedSchema, @@ -104,17 +109,16 @@ func (p *Parser) ParseQuery(ctx context.Context, input string) (drivers.Query, e paramIdxMap: make(map[int64]int64), } - stmt := parseResult.Stmts[0] info := w.walk(stmt.Stmt) switch node := stmt.Stmt.Node.(type) { case *pg.Node_SelectStmt: + info = info.children["SelectStmt"] if len(node.SelectStmt.ValuesLists) > 0 { - return drivers.Query{}, fmt.Errorf("VALUES statement is not supported") + w.modValuesStatement(node, info) + } else { + w.modSelectStatement(node, info) } - info = info.children["SelectStmt"] - w.modSelectStatement(node, info) - case *pg.Node_InsertStmt: info = info.children["InsertStmt"] w.modInsertStatement(node, info) @@ -130,6 +134,15 @@ func (p *Parser) ParseQuery(ctx context.Context, input string) (drivers.Query, e case *pg.Node_MergeStmt: info = info.children["MergeStmt"] w.modMergeStatement(node, info) + case *pg.Node_ListenStmt: + // pg.ListenStmt has no Location field; find the keyword token directly + info = w.findTokenAfter(0, pg.Token_LISTEN) + w.modListenStatement(node, info) + + case *pg.Node_NotifyStmt: + // pg.NotifyStmt has no Location field; find the keyword token directly + info = w.findTokenAfter(0, pg.Token_NOTIFY) + w.modNotifyStatement(node, info) } source := w.getSource(stmt.Stmt, info) @@ -143,9 +156,12 @@ func (p *Parser) ParseQuery(ctx context.Context, input string) (drivers.Query, e return drivers.Query{}, fmt.Errorf("format: %w", err) } - argTypes, resTypes, err := p.getArgsAndCols(ctx, formatted) - if err != nil { - return drivers.Query{}, fmt.Errorf("get args and cols: %w", err) + // LISTEN/NOTIFY cannot be PREPAREd; they have no args or result columns + if qType != bob.QueryTypeListen && qType != bob.QueryTypeNotify { + argTypes, resTypes, err = p.getArgsAndCols(ctx, formatted) + if err != nil { + return drivers.Query{}, fmt.Errorf("get args and cols: %w", err) + } } if len(source.columns) != len(resTypes) { @@ -223,6 +239,10 @@ func isReturningWithParseError(sql string, err error) bool { func getQueryType(stmt *pg.Node) bob.QueryType { switch stmt.Node.(type) { case *pg.Node_SelectStmt: + // VALUES (...) is parsed as SelectStmt with ValuesLists set; no separate node type exists + if len(stmt.Node.(*pg.Node_SelectStmt).SelectStmt.ValuesLists) > 0 { + return bob.QueryTypeValues + } return bob.QueryTypeSelect case *pg.Node_InsertStmt: return bob.QueryTypeInsert @@ -232,6 +252,10 @@ func getQueryType(stmt *pg.Node) bob.QueryType { return bob.QueryTypeDelete case *pg.Node_MergeStmt: return bob.QueryTypeMerge + case *pg.Node_ListenStmt: + return bob.QueryTypeListen + case *pg.Node_NotifyStmt: + return bob.QueryTypeNotify default: return bob.QueryTypeUnknown } diff --git a/gen/bobgen-sqlite/driver/parser/sources.go b/gen/bobgen-sqlite/driver/parser/sources.go index 74f4c85f..3bb35d83 100644 --- a/gen/bobgen-sqlite/driver/parser/sources.go +++ b/gen/bobgen-sqlite/driver/parser/sources.go @@ -16,14 +16,27 @@ func (v *visitor) addSourcesFromWithClause(ctx sqliteparser.IWith_clauseContext) } for _, cte := range ctx.AllCommon_table_expression() { - columns, ok := cte.Select_stmt().Accept(v).([]ReturnColumn) - if v.Err != nil { - v.Err = fmt.Errorf("CTE select stmt: %w", v.Err) - return - } - if !ok { - v.Err = fmt.Errorf("could not get stmt info") - return + var columns []ReturnColumn + + selectStmt := cte.Select_stmt() + if valClause := selectStmt.Select_core().Values_clause(); valClause != nil && len(selectStmt.AllCompound_select()) == 0 { + valClause.Accept(v) + if v.Err != nil { + v.Err = fmt.Errorf("CTE values stmt: %w", v.Err) + return + } + columns = v.getSourceFromValuesClause(valClause).Columns + } else { + var ok bool + columns, ok = selectStmt.Accept(v).([]ReturnColumn) + if v.Err != nil { + v.Err = fmt.Errorf("CTE select stmt: %w", v.Err) + return + } + if !ok { + v.Err = fmt.Errorf("could not get stmt info") + return + } } source := QuerySource{ @@ -100,7 +113,21 @@ func (v *visitor) getSourceFromTableOrSubQuery(ctx sqliteparser.ITable_or_subque return v.getSourceFromTable(ctx) case ctx.Select_stmt() != nil: - columns, ok := ctx.Select_stmt().Accept(v).([]ReturnColumn) + selectStmt := ctx.Select_stmt() + if valClause := selectStmt.Select_core().Values_clause(); valClause != nil && len(selectStmt.AllCompound_select()) == 0 { + valClause.Accept(v) + if v.Err != nil { + v.Err = fmt.Errorf("table values stmt: %w", v.Err) + return QuerySource{} + } + source := v.getSourceFromValuesClause(valClause) + return QuerySource{ + Name: getName(ctx.Table_alias()), + Columns: source.Columns, + } + } + + columns, ok := selectStmt.Accept(v).([]ReturnColumn) if v.Err != nil { v.Err = fmt.Errorf("table select stmt: %w", v.Err) return QuerySource{} @@ -124,6 +151,19 @@ func (v *visitor) getSourceFromTableOrSubQuery(ctx sqliteparser.ITable_or_subque } } +func (v *visitor) getSourceFromValuesClause(ctx sqliteparser.IValues_clauseContext) QuerySource { + rows := ctx.AllValue_row() + if len(rows) == 0 { + return QuerySource{} + } + exprs := rows[0].AllExpr() + columns := make([]ReturnColumn, len(exprs)) + for i := range exprs { + columns[i] = ReturnColumn{Name: fmt.Sprintf("column%d", i)} + } + return QuerySource{Columns: columns} +} + func (v *visitor) getSourceFromTable(ctx interface { Schema_name() sqliteparser.ISchema_nameContext Table_name() sqliteparser.ITable_nameContext diff --git a/gen/bobgen-sqlite/driver/parser/visitor.go b/gen/bobgen-sqlite/driver/parser/visitor.go index 232e90f0..ff5607f7 100644 --- a/gen/bobgen-sqlite/driver/parser/visitor.go +++ b/gen/bobgen-sqlite/driver/parser/visitor.go @@ -113,8 +113,12 @@ func (v *visitor) VisitSql_stmt_list(ctx *sqliteparser.Sql_stmt_listContext) any switch child := child.(type) { case *sqliteparser.Select_stmtContext: - queryType = bob.QueryTypeSelect - imports = v.modSelect_stmt(child, mods) + if child.Select_core().Values_clause() != nil { + queryType = bob.QueryTypeValues + } else { + queryType = bob.QueryTypeSelect + imports = v.modSelect_stmt(child, mods) + } case *sqliteparser.Insert_stmtContext: queryType = bob.QueryTypeInsert v.modInsert_stmt(child, mods) diff --git a/query.go b/query.go index 455c1707..e6c29eba 100644 --- a/query.go +++ b/query.go @@ -26,6 +26,8 @@ const ( QueryTypeDelete QueryTypeValues QueryTypeMerge + QueryTypeListen // Postgres only + QueryTypeNotify // Postgres only ) func (q QueryType) String() string { @@ -42,6 +44,10 @@ func (q QueryType) String() string { return "VALUES" case QueryTypeMerge: return "MERGE" + case QueryTypeListen: + return "LISTEN" + case QueryTypeNotify: + return "NOTIFY" default: return "UNKNOWN" }