diff --git a/cmd/benchmark/explain.go b/cmd/benchmark/explain.go index ca1334f7..ef062064 100644 --- a/cmd/benchmark/explain.go +++ b/cmd/benchmark/explain.go @@ -19,6 +19,7 @@ package main import ( "context" "fmt" + "maps" "github.com/specterops/dawgs/cypher/frontend" "github.com/specterops/dawgs/cypher/models/pgsql" @@ -50,7 +51,9 @@ func newPostgresExplainer(kindMapper pgsql.KindMapper, graphID int32) ExplainFun return nil, err } - result := tx.Raw("EXPLAIN (ANALYZE, BUFFERS) "+sqlQuery, translation.Parameters) + maps.Copy(translation.Parameters, sqlQuery.Parameters) + + result := tx.Raw("EXPLAIN (ANALYZE, BUFFERS) "+sqlQuery.Statement, translation.Parameters) defer result.Close() var plan []string @@ -67,8 +70,9 @@ func newPostgresExplainer(kindMapper pgsql.KindMapper, graphID int32) ExplainFun return nil, err } + // TODO: should this get the parameters as well? return &ExplainResult{ - SQL: sqlQuery, + SQL: sqlQuery.Statement, Plan: plan, Optimization: translation.Optimization, }, nil diff --git a/cmd/graphbench/postgres.go b/cmd/graphbench/postgres.go index 355b6bc3..cf6eeb50 100644 --- a/cmd/graphbench/postgres.go +++ b/cmd/graphbench/postgres.go @@ -19,6 +19,7 @@ package main import ( "context" "fmt" + "maps" "regexp" "strconv" "strings" @@ -187,9 +188,11 @@ func (s *postgresSQLRunner) explain(ctx context.Context, cypherQuery string, par return postgresExplain{}, err } + maps.Copy(translation.Parameters, sqlQuery.Parameters) + var plan []string if err := s.db.ReadTransaction(ctx, func(tx graph.Transaction) error { - result := tx.Raw("EXPLAIN (ANALYZE, BUFFERS, TIMING OFF) "+sqlQuery, translation.Parameters) + result := tx.Raw("EXPLAIN (ANALYZE, BUFFERS, TIMING OFF) "+sqlQuery.Statement, translation.Parameters) defer result.Close() for result.Next() { @@ -206,8 +209,9 @@ func (s *postgresSQLRunner) explain(ctx context.Context, cypherQuery string, par return postgresExplain{}, err } + // TODO: should this get the parameters as well? return postgresExplain{ - SQL: sqlQuery, + SQL: sqlQuery.Statement, Plan: plan, Metrics: parsePostgresPlanMetrics(plan), Optimization: translation.Optimization, diff --git a/cmd/plancorpus/capture.go b/cmd/plancorpus/capture.go index d05a7046..0ca73e6c 100644 --- a/cmd/plancorpus/capture.go +++ b/cmd/plancorpus/capture.go @@ -3,6 +3,7 @@ package main import ( "context" "fmt" + "maps" "net/url" "os" "path/filepath" @@ -280,9 +281,11 @@ func (s *backendCapture) capturePostgres(ctx context.Context, cypherQuery string return } + maps.Copy(translation.Parameters, sqlQuery.Parameters) + var plan []string if err := s.db.ReadTransaction(ctx, func(tx graph.Transaction) error { - result := tx.Raw("EXPLAIN "+sqlQuery, translation.Parameters) + result := tx.Raw("EXPLAIN "+sqlQuery.Statement, translation.Parameters) defer result.Close() for result.Next() { @@ -298,7 +301,8 @@ func (s *backendCapture) capturePostgres(ctx context.Context, cypherQuery string record.Error = err.Error() } - record.SQL = sqlQuery + // TODO: should this get the parameters as well? + record.SQL = sqlQuery.Statement record.PGPlan = plan record.PGOperators = postgresOperators(plan) record.PlannedLowerings = loweringNames(translation.Optimization.PlannedLowerings) diff --git a/cypher/models/pgsql/format/format.go b/cypher/models/pgsql/format/format.go index 9cbed49c..3755bb76 100644 --- a/cypher/models/pgsql/format/format.go +++ b/cypher/models/pgsql/format/format.go @@ -9,9 +9,9 @@ import ( ) type OutputBuilder struct { - MaterializeParameters bool - StripLiterals bool - parameters map[string]any + params map[string]any + materializeParameters bool + materializedParams map[string]any builder *strings.Builder } @@ -22,8 +22,8 @@ func NewOutputBuilder() *OutputBuilder { } func (s *OutputBuilder) WithMaterializedParameters(parameters map[string]any) *OutputBuilder { - s.MaterializeParameters = true - s.parameters = parameters + s.materializeParameters = true + s.materializedParams = parameters return s } @@ -47,8 +47,11 @@ func (s *OutputBuilder) Write(values ...any) { } } -func (s *OutputBuilder) Build() string { - return s.builder.String() +func (s *OutputBuilder) Build() Formatted { + return Formatted{ + Statement: s.builder.String(), + Parameters: s.params, + } } func formatSlice[T any, TS []T](builder *OutputBuilder, slice TS, dataType pgsql.DataType) error { @@ -546,8 +549,8 @@ func formatNode(builder *OutputBuilder, rootExpr pgsql.SyntaxNode) error { ) case pgsql.Parameter: - if builder.MaterializeParameters { - if parameterValue, hasParameter := builder.parameters[typedNextExpr.Identifier.String()]; !hasParameter { + if builder.materializeParameters { + if parameterValue, hasParameter := builder.materializedParams[typedNextExpr.Identifier.String()]; !hasParameter { return fmt.Errorf("invalid parameter %s", typedNextExpr.Identifier.String()) } else if parameterLiteral, err := pgsql.AsLiteral(parameterValue); err != nil { return fmt.Errorf("invalid parameter value for %s: %v", typedNextExpr.Identifier.String(), err) @@ -611,9 +614,9 @@ func formatNode(builder *OutputBuilder, rootExpr pgsql.SyntaxNode) error { return nil } -func Expression(expression pgsql.SyntaxNode, builder *OutputBuilder) (string, error) { +func Expression(expression pgsql.SyntaxNode, builder *OutputBuilder) (Formatted, error) { if err := formatNode(builder, expression); err != nil { - return "", err + return Formatted{}, err } return builder.Build(), nil @@ -1159,42 +1162,42 @@ func formatDeleteStatement(builder *OutputBuilder, sqlDelete pgsql.Delete) error return nil } -func Statement(statement pgsql.Statement, builder *OutputBuilder) (string, error) { +func Statement(statement pgsql.Statement, builder *OutputBuilder) (Formatted, error) { switch typedStatement := statement.(type) { case pgsql.Merge: if err := formatMergeStatement(builder, typedStatement); err != nil { - return "", err + return Formatted{}, err } case pgsql.Query: if err := formatSetExpression(builder, typedStatement); err != nil { - return "", err + return Formatted{}, err } case pgsql.Insert: if err := formatInsertStatement(builder, typedStatement); err != nil { - return "", err + return Formatted{}, err } case pgsql.Update: if err := formatUpdateStatement(builder, typedStatement); err != nil { - return "", err + return Formatted{}, err } case pgsql.Delete: if err := formatDeleteStatement(builder, typedStatement); err != nil { - return "", err + return Formatted{}, err } default: - return "", fmt.Errorf("unsupported PgSQL statement type: %T", statement) + return Formatted{}, fmt.Errorf("unsupported PgSQL statement type: %T", statement) } builder.Write(";") return builder.Build(), nil } -func SyntaxNode(node pgsql.SyntaxNode) (string, error) { +func SyntaxNode(node pgsql.SyntaxNode) (Formatted, error) { builder := NewOutputBuilder() switch typedNode := node.(type) { @@ -1205,7 +1208,7 @@ func SyntaxNode(node pgsql.SyntaxNode) (string, error) { return Expression(typedNode, builder) default: - return "", fmt.Errorf("unknown SQL AST type: %T", node) + return Formatted{}, fmt.Errorf("unknown SQL AST type: %T", node) } } diff --git a/cypher/models/pgsql/format/format_test.go b/cypher/models/pgsql/format/format_test.go index 86b9629d..0b733c75 100644 --- a/cypher/models/pgsql/format/format_test.go +++ b/cypher/models/pgsql/format/format_test.go @@ -23,7 +23,7 @@ func TestFormat_TypeCastedParenthetical(t *testing.T) { formattedQuery, err := format.Expression(typeCastedParenthetical, format.NewOutputBuilder()) require.Nil(t, err) - require.Equal(t, "('str')::text", formattedQuery) + require.Equal(t, "('str')::text", formattedQuery.Statement) } func TestFormat_Case(t *testing.T) { @@ -48,7 +48,7 @@ func TestFormat_Case(t *testing.T) { }, format.NewOutputBuilder()) require.NoError(t, err) - require.Equal(t, "case when s0.root_id != s0.next_id then true else shortest_path_self_endpoint_error(s0.root_id, s0.next_id) end", formattedQuery) + require.Equal(t, "case when s0.root_id != s0.next_id then true else shortest_path_self_endpoint_error(s0.root_id, s0.next_id) end", formattedQuery.Statement) } func TestFormat_SelectDistinct(t *testing.T) { @@ -67,7 +67,7 @@ func TestFormat_SelectDistinct(t *testing.T) { }, format.NewOutputBuilder()) require.Nil(t, err) - require.Equal(t, "select distinct id from node;", formattedQuery) + require.Equal(t, "select distinct id from node;", formattedQuery.Statement) } func TestFormat_LateralSubqueryJoin(t *testing.T) { @@ -115,7 +115,7 @@ func TestFormat_LateralSubqueryJoin(t *testing.T) { }, format.NewOutputBuilder()) require.Nil(t, err) - require.Equal(t, "select n.id, e.id from node n join lateral (select e.id from edge e where e.start_id = n.id offset 0) e on true;", formattedQuery) + require.Equal(t, "select n.id, e.id from node n join lateral (select e.id from edge e where e.start_id = n.id offset 0) e on true;", formattedQuery.Statement) } func TestFormat_Delete(t *testing.T) { @@ -132,7 +132,7 @@ func TestFormat_Delete(t *testing.T) { }, format.NewOutputBuilder()) require.Nil(t, err) - require.Equal(t, "delete from table t where t.col1 < 4;", formattedQuery) + require.Equal(t, "delete from table t where t.col1 < 4;", formattedQuery.Statement) } func TestFormat_Update(t *testing.T) { @@ -158,7 +158,7 @@ func TestFormat_Update(t *testing.T) { }, format.NewOutputBuilder()) require.Nil(t, err) - require.Equal(t, "update table t set col1 = 1, col2 = '12345' where t.col1 < 4;", formattedQuery) + require.Equal(t, "update table t set col1 = 1, col2 = '12345' where t.col1 < 4;", formattedQuery.Statement) } func TestFormat_Insert(t *testing.T) { @@ -177,7 +177,7 @@ func TestFormat_Insert(t *testing.T) { }, format.NewOutputBuilder()) require.Nil(t, err) - require.Equal(t, "insert into table (col1, col2, col3) values ('1', 1, false);", formattedQuery) + require.Equal(t, "insert into table (col1, col2, col3) values ('1', 1, false);", formattedQuery.Statement) formattedQuery, err = format.Statement(pgsql.Insert{ Table: pgsql.TableReference{ @@ -206,7 +206,7 @@ func TestFormat_Insert(t *testing.T) { }, format.NewOutputBuilder()) require.Nil(t, err) - require.Equal(t, "insert into table (col1, col2, col3) select * from other where other.col1 = '1234';", formattedQuery) + require.Equal(t, "insert into table (col1, col2, col3) select * from other where other.col1 = '1234';", formattedQuery.Statement) formattedQuery, err = format.Statement(pgsql.Insert{ Table: pgsql.TableReference{ @@ -238,7 +238,7 @@ func TestFormat_Insert(t *testing.T) { }, format.NewOutputBuilder()) require.Nil(t, err) - require.Equal(t, "insert into table (col1, col2, col3) select * from other where other.col1 = '1234' returning id;", formattedQuery) + require.Equal(t, "insert into table (col1, col2, col3) select * from other where other.col1 = '1234' returning id;", formattedQuery.Statement) formattedQuery, err = format.Statement(pgsql.Insert{ Table: pgsql.TableReference{ @@ -289,7 +289,7 @@ func TestFormat_Insert(t *testing.T) { }, format.NewOutputBuilder()) require.Nil(t, err) - require.Equal(t, "insert into table (col1, col2, col3) select * from other where other.col1 = '1234' on conflict on constraint other.hash_constraint do update set hit_count = hit_count + 1 where hit_count < 9999 returning id, hit_count;", formattedQuery) + require.Equal(t, "insert into table (col1, col2, col3) select * from other where other.col1 = '1234' on conflict on constraint other.hash_constraint do update set hit_count = hit_count + 1 where hit_count < 9999 returning id, hit_count;", formattedQuery.Statement) formattedQuery, err = format.Statement(pgsql.Insert{ Table: pgsql.TableReference{ @@ -339,7 +339,7 @@ func TestFormat_Insert(t *testing.T) { }, format.NewOutputBuilder()) require.Nil(t, err) - require.Equal(t, "insert into table (col1, col2, col3) select * from other where other.col1 = '1234' on conflict (hash) do update set hit_count = hit_count + 1 where hit_count < 9999;", formattedQuery) + require.Equal(t, "insert into table (col1, col2, col3) select * from other where other.col1 = '1234' on conflict (hash) do update set hit_count = hit_count + 1 where hit_count < 9999;", formattedQuery.Statement) } func TestFormat_Query(t *testing.T) { @@ -367,7 +367,7 @@ func TestFormat_Query(t *testing.T) { formattedQuery, err := format.Statement(query, format.NewOutputBuilder()) require.Nil(t, err) - require.Equal(t, "select * from table t where t.col1 > 1;", formattedQuery) + require.Equal(t, "select * from table t where t.col1 > 1;", formattedQuery.Statement) } func TestFormat_Merge(t *testing.T) { @@ -441,7 +441,7 @@ func TestFormat_Merge(t *testing.T) { }, format.NewOutputBuilder()) require.Nil(t, err) - require.Equal(t, "merge into table t using source s on t.source_id = s.id when matched and t.value > s.value then update set updated_at = now() when matched and t.value <= s.value then update set value = s.value, t.updated_at = now() when matched and t.value = s.value then delete when not matched and t.value = 0 then insert (hit_count) values (0);", formattedQuery) + require.Equal(t, "merge into table t using source s on t.source_id = s.id when matched and t.value > s.value then update set updated_at = now() when matched and t.value <= s.value then update set value = s.value, t.updated_at = now() when matched and t.value = s.value then delete when not matched and t.value = 0 then insert (hit_count) values (0);", formattedQuery.Statement) } func TestFormat_CTEs(t *testing.T) { @@ -661,7 +661,7 @@ func TestFormat_CTEs(t *testing.T) { }, format.NewOutputBuilder()) require.Nil(t, err) - require.Equal(t, "with recursive expansion_1(root_id, next_id, depth, stop, is_cycle, path) as materialized (select r.start_id, r.end_id, 1, false, r.start_id = r.end_id, array [r.id] from edge r join node a on a.id = r.start_id where a.kind_ids operator (pg_catalog.&&) array [23]::int2[] union all select expansion_1.root_id, r.end_id, expansion_1.depth + 1, b.kind_ids operator (pg_catalog.&&) array [24]::int2[], r.id = any(expansion_1.path), expansion_1.path || r.id from expansion_1 join edge r on r.start_id = expansion_1.next_id join node b on b.id = r.end_id where not expansion_1.is_cycle and not expansion_1.stop) select a.properties, b.properties from expansion_1 join node a on a.id = expansion_1.root_id join node b on b.id = expansion_1.next_id where not expansion_1.is_cycle and expansion_1.stop;", formattedQuery) + require.Equal(t, "with recursive expansion_1(root_id, next_id, depth, stop, is_cycle, path) as materialized (select r.start_id, r.end_id, 1, false, r.start_id = r.end_id, array [r.id] from edge r join node a on a.id = r.start_id where a.kind_ids operator (pg_catalog.&&) array [23]::int2[] union all select expansion_1.root_id, r.end_id, expansion_1.depth + 1, b.kind_ids operator (pg_catalog.&&) array [24]::int2[], r.id = any(expansion_1.path), expansion_1.path || r.id from expansion_1 join edge r on r.start_id = expansion_1.next_id join node b on b.id = r.end_id where not expansion_1.is_cycle and not expansion_1.stop) select a.properties, b.properties from expansion_1 join node a on a.id = expansion_1.root_id join node b on b.id = expansion_1.next_id where not expansion_1.is_cycle and expansion_1.stop;", formattedQuery.Statement) } func TestFormat_QueryInjection(t *testing.T) { @@ -689,5 +689,5 @@ func TestFormat_QueryInjection(t *testing.T) { formattedQuery, err := format.Statement(query, format.NewOutputBuilder()) require.Nil(t, err) - require.Equal(t, `select * from table t where t.col1 = 'alpha'' || select (''malicious'')';`, formattedQuery) + require.Equal(t, `select * from table t where t.col1 = 'alpha'' || select (''malicious'')';`, formattedQuery.Statement) } diff --git a/cypher/models/pgsql/test/testcase.go b/cypher/models/pgsql/test/testcase.go index 65dcf571..edc738e1 100644 --- a/cypher/models/pgsql/test/testcase.go +++ b/cypher/models/pgsql/test/testcase.go @@ -7,6 +7,7 @@ import ( "fmt" "io" "io/fs" + "maps" "os" "path/filepath" "strings" @@ -126,6 +127,7 @@ func (s *TranslationTestCase) WriteTo(output io.Writer, kindMapper pgsql.KindMap } else if formattedQuery, err := translate.Translated(translation); err != nil { return err } else { + maps.Copy(translation.Parameters, formattedQuery.Parameters) if len(translation.Parameters) > 0 { if encodedJSON, err := json.Marshal(translation.Parameters); err != nil { return err @@ -138,7 +140,7 @@ func (s *TranslationTestCase) WriteTo(output io.Writer, kindMapper pgsql.KindMap } } - if err := writeStrings(output, formattedQuery, "\n\n"); err != nil { + if err := writeStrings(output, formattedQuery.Statement, "\n\n"); err != nil { return err } } @@ -170,13 +172,14 @@ func (s *TranslationTestCase) Assert(t *testing.T, expectedSQL string, kindMappe t.Fatalf("Failed to format SQL translatedQuery: %v", err) } else { // Apply same whitespace normalization as expected SQL (from file parsing) - normalizedActual, err := regexp.ReplaceAll("\\s+", strings.TrimSpace(formattedQuery), " ") + normalizedActual, err := regexp.ReplaceAll("\\s+", strings.TrimSpace(formattedQuery.Statement), " ") if err != nil { t.Fatalf("error while attempting to collapse whitespace in actual query: %v", err) } require.Equalf(t, expectedSQL, normalizedActual, "Test case for cypher query: '%s' failed to match.", s.Cypher) if s.PgSQLParams != nil { + maps.Copy(translation.Parameters, formattedQuery.Parameters) require.Equal(t, s.PgSQLParams, translation.Parameters) } } @@ -210,7 +213,8 @@ func (s *TranslationTestCase) AssertLive(ctx context.Context, t *testing.T, driv } else if formattedQuery, err := translate.Translated(translation); err != nil { t.Fatalf("Failed to format SQL translatedQuery: %v", err) } else { - require.NoError(t, driver.Run(ctx, "explain "+formattedQuery, translation.Parameters)) + maps.Copy(translation.Parameters, formattedQuery.Parameters) + require.NoError(t, driver.Run(ctx, "explain "+formattedQuery.Statement, translation.Parameters)) } } } diff --git a/cypher/models/pgsql/translate/create_test.go b/cypher/models/pgsql/translate/create_test.go index 68f1eb90..8d6a37d7 100644 --- a/cypher/models/pgsql/translate/create_test.go +++ b/cypher/models/pgsql/translate/create_test.go @@ -25,6 +25,6 @@ func TestConsecutiveCreateClausesAreBuiltOnce(t *testing.T) { formatted, err := Translated(translation) require.NoError(t, err) - require.Equal(t, 2, strings.Count(formatted, "insert into node")) - require.Equal(t, 2, strings.Count(formatted, "nextval(pg_get_serial_sequence('node', 'id'))")) + require.Equal(t, 2, strings.Count(formatted.Statement, "insert into node")) + require.Equal(t, 2, strings.Count(formatted.Statement, "nextval(pg_get_serial_sequence('node', 'id'))")) } diff --git a/cypher/models/pgsql/translate/expansion.go b/cypher/models/pgsql/translate/expansion.go index c7d27587..62de0791 100644 --- a/cypher/models/pgsql/translate/expansion.go +++ b/cypher/models/pgsql/translate/expansion.go @@ -1936,13 +1936,15 @@ func (s *ExpansionBuilder) boundEndpointFilterParameters() ([]pgsql.Expression, if formattedFilter, err := format.Statement(pairFilterStatement, format.NewOutputBuilder().WithMaterializedParameters(s.queryParameters)); err != nil { return nil, err } else { - pairFilter = formattedFilter + // TODO: handle parameters + pairFilter = formattedFilter.Statement } } else if hasRootFilter { if formattedFilter, err := format.Statement(rootFilterStatement, format.NewOutputBuilder().WithMaterializedParameters(s.queryParameters)); err != nil { return nil, err } else { - rootFilter = formattedFilter + // TODO: handle parameters + rootFilter = formattedFilter.Statement } } @@ -1950,7 +1952,8 @@ func (s *ExpansionBuilder) boundEndpointFilterParameters() ([]pgsql.Expression, if formattedFilter, err := format.Statement(terminalFilterStatement, format.NewOutputBuilder().WithMaterializedParameters(s.queryParameters)); err != nil { return nil, err } else { - terminalFilter = formattedFilter + // TODO: handle parameters + terminalFilter = formattedFilter.Statement } } @@ -1970,9 +1973,11 @@ func (s *ExpansionBuilder) shortestPathsParameters(expansionModel *Expansion, fo var ( harnessParameters []pgsql.Expression formatFragment = func(query pgsql.SetExpression) (string, error) { - return format.Statement( + // TODO: handle parameters + stmt, err := format.Statement( nextFrontInsert(query), format.NewOutputBuilder().WithMaterializedParameters(s.queryParameters)) + return stmt.Statement, err } ) @@ -2014,9 +2019,11 @@ func (s *ExpansionBuilder) bidirectionalAllShortestPathsParameters(expansionMode var ( harnessParameters []pgsql.Expression formatFragment = func(query pgsql.SetExpression) (string, error) { - return format.Statement( + // TODO: handle parameters + stmt, err := format.Statement( nextFrontInsert(query), format.NewOutputBuilder().WithMaterializedParameters(s.queryParameters)) + return stmt.Statement, err } ) diff --git a/cypher/models/pgsql/translate/expansion_test.go b/cypher/models/pgsql/translate/expansion_test.go index 9075d55e..37772c43 100644 --- a/cypher/models/pgsql/translate/expansion_test.go +++ b/cypher/models/pgsql/translate/expansion_test.go @@ -95,24 +95,24 @@ func newShortestPathSeedTestBuilder(leftBound, rightBound bool) (*ExpansionBuild func TestShortestPathSelfEndpointGuardsUseCaseErrorHelper(t *testing.T) { projectionGuard, err := format.Expression(shortestPathSelfEndpointGuard(shortestPathSeedTestFrame), format.NewOutputBuilder()) require.NoError(t, err) - require.Equal(t, "case when s1.root_id != s1.next_id then true else shortest_path_self_endpoint_error(s1.root_id, s1.next_id) end", projectionGuard) - require.NotContains(t, projectionGuard, " / ") + require.Equal(t, "case when s1.root_id != s1.next_id then true else shortest_path_self_endpoint_error(s1.root_id, s1.next_id) end", projectionGuard.Statement) + require.NotContains(t, projectionGuard.Statement, " / ") terminalFilterGuard, err := format.Expression( shortestPathSeedSelfEndpointGuard(pgsql.CompoundIdentifier{shortestPathSeedTestEdge, pgsql.ColumnStartID}, false), format.NewOutputBuilder(), ) require.NoError(t, err) - require.Contains(t, terminalFilterGuard, "case when (select count(*)::int8 from traversal_terminal_filter where traversal_terminal_filter.id = e0.start_id) = 0 then true else shortest_path_self_endpoint_error(e0.start_id, e0.start_id) end") - require.NotContains(t, terminalFilterGuard, " / ") + require.Contains(t, terminalFilterGuard.Statement, "case when (select count(*)::int8 from traversal_terminal_filter where traversal_terminal_filter.id = e0.start_id) = 0 then true else shortest_path_self_endpoint_error(e0.start_id, e0.start_id) end") + require.NotContains(t, terminalFilterGuard.Statement, " / ") endpointPairFilterGuard, err := format.Expression( shortestPathSeedSelfEndpointGuard(pgsql.CompoundIdentifier{shortestPathSeedTestEdge, pgsql.ColumnStartID}, true), format.NewOutputBuilder(), ) require.NoError(t, err) - require.Contains(t, endpointPairFilterGuard, "case when (select count(*)::int8 from traversal_pair_filter where traversal_pair_filter.root_id = e0.start_id and traversal_pair_filter.terminal_id = e0.start_id) = 0 then true else shortest_path_self_endpoint_error(e0.start_id, e0.start_id) end") - require.NotContains(t, endpointPairFilterGuard, " / ") + require.Contains(t, endpointPairFilterGuard.Statement, "case when (select count(*)::int8 from traversal_pair_filter where traversal_pair_filter.root_id = e0.start_id and traversal_pair_filter.terminal_id = e0.start_id) = 0 then true else shortest_path_self_endpoint_error(e0.start_id, e0.start_id) end") + require.NotContains(t, endpointPairFilterGuard.Statement, " / ") } func TestBoundRootShortestPathPrimerKeepsOnlySeedLocalConstraints(t *testing.T) { @@ -134,8 +134,8 @@ func TestBoundRootShortestPathPrimerKeepsOnlySeedLocalConstraints(t *testing.T) formattedQuery, err := format.Statement(query, format.NewOutputBuilder()) require.NoError(t, err) - require.Contains(t, formattedQuery, "n0.id = (s0.x).id") - require.Contains(t, formattedQuery, "(s0.n0).id = s1.root_id") + require.Contains(t, formattedQuery.Statement, "n0.id = (s0.x).id") + require.Contains(t, formattedQuery.Statement, "(s0.n0).id = s1.root_id") } func TestBoundTerminalShortestPathPrimerKeepsOnlySeedLocalConstraints(t *testing.T) { @@ -159,8 +159,8 @@ func TestBoundTerminalShortestPathPrimerKeepsOnlySeedLocalConstraints(t *testing formattedQuery, err := format.Statement(query, format.NewOutputBuilder()) require.NoError(t, err) - require.Contains(t, formattedQuery, "n1.id = (s0.x).id") - require.Contains(t, formattedQuery, "(s0.n1).id = s1.next_id") + require.Contains(t, formattedQuery.Statement, "n1.id = (s0.x).id") + require.Contains(t, formattedQuery.Statement, "(s0.n1).id = s1.next_id") } func TestZeroDepthExpansionRejectsEdgeDependentTerminalSatisfaction(t *testing.T) { @@ -186,8 +186,8 @@ func TestZeroDepthExpansionRejectsEdgeDependentTerminalSatisfaction(t *testing.T formattedQuery, err := format.Statement(pgsql.Query{Body: zeroDepthSelect}, format.NewOutputBuilder()) require.NoError(t, err) - require.Contains(t, formattedQuery, "select s1_seed.root_id, s1_seed.root_id, 0, false, false") - require.NotContains(t, formattedQuery, "e0") + require.Contains(t, formattedQuery.Statement, "select s1_seed.root_id, s1_seed.root_id, 0, false, false") + require.NotContains(t, formattedQuery.Statement, "e0") } func TestZeroDepthExpansionBuildKeepsPrimerBranch(t *testing.T) { @@ -245,10 +245,10 @@ func TestZeroDepthExpansionBuildKeepsPrimerBranch(t *testing.T) { primerBranch := "select 1, 2, 1, true, e0.start_id = e0.end_id, array [7]" recursiveBranch := "select 1, 3, 2, true, false, array [8]" - require.Contains(t, formattedQuery, zeroDepthBranch) - require.Contains(t, formattedQuery, primerBranch) - require.Contains(t, formattedQuery, recursiveBranch) - require.Contains(t, formattedQuery, "where s1.depth > 0") - require.Less(t, strings.Index(formattedQuery, zeroDepthBranch), strings.Index(formattedQuery, primerBranch)) - require.Less(t, strings.Index(formattedQuery, primerBranch), strings.Index(formattedQuery, recursiveBranch)) + require.Contains(t, formattedQuery.Statement, zeroDepthBranch) + require.Contains(t, formattedQuery.Statement, primerBranch) + require.Contains(t, formattedQuery.Statement, recursiveBranch) + require.Contains(t, formattedQuery.Statement, "where s1.depth > 0") + require.Less(t, strings.Index(formattedQuery.Statement, zeroDepthBranch), strings.Index(formattedQuery.Statement, primerBranch)) + require.Less(t, strings.Index(formattedQuery.Statement, primerBranch), strings.Index(formattedQuery.Statement, recursiveBranch)) } diff --git a/cypher/models/pgsql/translate/expression_test.go b/cypher/models/pgsql/translate/expression_test.go index 9c9df618..c4d03fab 100644 --- a/cypher/models/pgsql/translate/expression_test.go +++ b/cypher/models/pgsql/translate/expression_test.go @@ -141,7 +141,7 @@ func TestInferExpressionType(t *testing.T) { if testName, err := format.Expression(nextCase.Expression, format.NewOutputBuilder()); err != nil { t.Fatalf("unable to format test case expression: %v", err) } else { - t.Run(testName, func(t *testing.T) { + t.Run(testName.Statement, func(t *testing.T) { inferredType, err := translate.InferExpressionType(nextCase.Expression) require.Nil(t, err) @@ -315,7 +315,8 @@ func TestPropertyLookupEqualityScalarRewrites(t *testing.T) { formatted, err := format.Expression(treeTranslator.PeekOperand(), format.NewOutputBuilder()) require.NoError(t, err) - return formatted + // TODO: does this need to handle Properties? + return formatted.Statement } testCases = []struct { Name string diff --git a/cypher/models/pgsql/translate/format.go b/cypher/models/pgsql/translate/format.go index f750b8c1..c82f6418 100644 --- a/cypher/models/pgsql/translate/format.go +++ b/cypher/models/pgsql/translate/format.go @@ -3,6 +3,7 @@ package translate import ( "bytes" "context" + "maps" "strings" "github.com/specterops/dawgs/cypher/models/cypher" @@ -11,7 +12,7 @@ import ( "github.com/specterops/dawgs/cypher/models/pgsql/format" ) -func Translated(translation Result) (string, error) { +func Translated(translation Result) (format.Formatted, error) { return format.Statement(translation.Statement, format.NewOutputBuilder()) } @@ -56,8 +57,9 @@ func FromCypher(ctx context.Context, regularQuery *cypher.RegularQuery, kindMapp } else if sqlQuery, err := format.Statement(translation.Statement, format.NewOutputBuilder()); err != nil { return format.Formatted{}, err } else { - output.WriteString(sqlQuery) + output.WriteString(sqlQuery.Statement) + maps.Copy(translation.Parameters, sqlQuery.Parameters) return format.Formatted{ Statement: output.String(), Parameters: translation.Parameters, diff --git a/cypher/models/pgsql/translate/function_test.go b/cypher/models/pgsql/translate/function_test.go index 066f0cf0..c6e5ee6f 100644 --- a/cypher/models/pgsql/translate/function_test.go +++ b/cypher/models/pgsql/translate/function_test.go @@ -23,8 +23,8 @@ func TestPathComponentFunctionsResolvePathAliases(t *testing.T) { formatted, err := Translated(translation) require.NoError(t, err) - require.Contains(t, formatted, ".nodes") - require.Contains(t, formatted, ".edges") + require.Contains(t, formatted.Statement, ".nodes") + require.Contains(t, formatted.Statement, ".edges") } func TestNodesFunctionTranslatesBoundPathToNodeArray(t *testing.T) { @@ -38,8 +38,8 @@ func TestNodesFunctionTranslatesBoundPathToNodeArray(t *testing.T) { formatted, err := Translated(translation) require.NoError(t, err) - require.Contains(t, formatted, "nodecomposite[]") - require.Contains(t, formatted, ".nodes") + require.Contains(t, formatted.Statement, "nodecomposite[]") + require.Contains(t, formatted.Statement, ".nodes") } func TestPathComponentFunctionsTranslateNullArguments(t *testing.T) { @@ -53,8 +53,8 @@ func TestPathComponentFunctionsTranslateNullArguments(t *testing.T) { formatted, err := Translated(translation) require.NoError(t, err) - require.Contains(t, formatted, "(null)::nodecomposite[]") - require.Contains(t, formatted, "(null)::edgecomposite[]") + require.Contains(t, formatted.Statement, "(null)::nodecomposite[]") + require.Contains(t, formatted.Statement, "(null)::edgecomposite[]") } func TestTailFunctionDoesNotDuplicatePathComponentExpression(t *testing.T) { @@ -68,8 +68,8 @@ func TestTailFunctionDoesNotDuplicatePathComponentExpression(t *testing.T) { formatted, err := Translated(translation) require.NoError(t, err) - require.Equal(t, 1, strings.Count(formatted, "ordered_edges_to_path"), formatted) - require.NotContains(t, formatted, "cardinality(((case when") + require.Equal(t, 1, strings.Count(formatted.Statement, "ordered_edges_to_path"), formatted) + require.NotContains(t, formatted.Statement, "cardinality(((case when") } func TestTailPredicateStagesPathComponentExpression(t *testing.T) { @@ -83,9 +83,9 @@ func TestTailPredicateStagesPathComponentExpression(t *testing.T) { formatted, err := Translated(translation) require.NoError(t, err) - require.Equal(t, 1, strings.Count(formatted, "ordered_edges_to_path")) - require.Contains(t, formatted, "lateral (select") - require.Contains(t, formatted, ".nodes") + require.Equal(t, 1, strings.Count(formatted.Statement, "ordered_edges_to_path")) + require.Contains(t, formatted.Statement, "lateral (select") + require.Contains(t, formatted.Statement, ".nodes") } func TestProjectionStagesPathBeforeReadingComponents(t *testing.T) { @@ -99,10 +99,10 @@ func TestProjectionStagesPathBeforeReadingComponents(t *testing.T) { formatted, err := Translated(translation) require.NoError(t, err) - require.Contains(t, formatted, "lateral (select") - require.Equal(t, 1, strings.Count(formatted, "ordered_edges_to_path"), formatted) - require.Contains(t, formatted, ".nodes") - require.Contains(t, formatted, ".edges") + require.Contains(t, formatted.Statement, "lateral (select") + require.Equal(t, 1, strings.Count(formatted.Statement, "ordered_edges_to_path"), formatted) + require.Contains(t, formatted.Statement, ".nodes") + require.Contains(t, formatted.Statement, ".edges") } func TestProjectionStagesRepeatedPathComponents(t *testing.T) { @@ -116,11 +116,11 @@ func TestProjectionStagesRepeatedPathComponents(t *testing.T) { formatted, err := Translated(translation) require.NoError(t, err) - require.Contains(t, formatted, "lateral (select") - require.Equal(t, 1, strings.Count(formatted, "ordered_edges_to_path"), formatted) - require.Equal(t, 1, strings.Count(formatted, "from unnest"), formatted) - require.Contains(t, formatted, ".nodes") - require.Contains(t, formatted, ".edges") + require.Contains(t, formatted.Statement, "lateral (select") + require.Equal(t, 1, strings.Count(formatted.Statement, "ordered_edges_to_path"), formatted) + require.Equal(t, 1, strings.Count(formatted.Statement, "from unnest"), formatted) + require.Contains(t, formatted.Statement, ".nodes") + require.Contains(t, formatted.Statement, ".edges") } func TestRelationshipEndpointFunctionsUseEdgeCompositeArguments(t *testing.T) { @@ -136,7 +136,7 @@ func TestRelationshipEndpointFunctionsUseEdgeCompositeArguments(t *testing.T) { formatted, err := Translated(translation) require.NoError(t, err) - normalized := strings.Join(strings.Fields(formatted), " ") + normalized := strings.Join(strings.Fields(formatted.Statement), " ") require.Contains(t, normalized, "start_node(((s0.e0).id, (s0.e0).start_id, (s0.e0).end_id, (s0.e0).kind_id, (s0.e0).properties)::edgecomposite)") require.Contains(t, normalized, "end_node(((s0.e0).id, (s0.e0).start_id, (s0.e0).end_id, (s0.e0).kind_id, (s0.e0).properties)::edgecomposite)") @@ -159,7 +159,7 @@ RETURN p formatted, err := Translated(translation) require.NoError(t, err) - normalized := strings.Join(strings.Fields(formatted), " ") + normalized := strings.Join(strings.Fields(formatted.Statement), " ") require.Contains(t, normalized, "from edge i0") require.Contains(t, normalized, "start_node((i0.id, i0.start_id, i0.end_id, i0.kind_id, i0.properties)::edgecomposite)") @@ -193,7 +193,7 @@ func TestCollectMembershipOnlyProjectionUsesIDs(t *testing.T) { formatted, err := Translated(translation) require.NoError(t, err) - normalized := strings.Join(strings.Fields(formatted), " ") + normalized := strings.Join(strings.Fields(formatted.Statement), " ") require.Contains(t, normalized, "array_agg((n0).id)") require.Contains(t, normalized, "array []::int8[]") @@ -215,7 +215,7 @@ func TestReturnedCollectNodeKeepsCompositeArray(t *testing.T) { formatted, err := Translated(translation) require.NoError(t, err) - normalized := strings.Join(strings.Fields(formatted), " ") + normalized := strings.Join(strings.Fields(formatted.Statement), " ") require.Contains(t, normalized, "array []::nodecomposite[]") require.NotContains(t, normalized, "array_agg((n0).id)") diff --git a/cypher/models/pgsql/translate/optimizer_safety_test.go b/cypher/models/pgsql/translate/optimizer_safety_test.go index 5e1786a2..0fab520a 100644 --- a/cypher/models/pgsql/translate/optimizer_safety_test.go +++ b/cypher/models/pgsql/translate/optimizer_safety_test.go @@ -63,7 +63,7 @@ func optimizerSafetySQL(t *testing.T, cypherQuery string) string { formattedQuery, err := Translated(translation) require.NoError(t, err) - return strings.Join(strings.Fields(formattedQuery), " ") + return strings.Join(strings.Fields(formattedQuery.Statement), " ") } func optimizerSafetyTranslation(t *testing.T, cypherQuery string) Result { @@ -210,7 +210,7 @@ func TestOptimizerSafetyCountStoreFastPathUsesBaseNodeCount(t *testing.T) { requirePlannedOptimizationLowering(t, translation.Optimization, optimize.LoweringCountStoreFastPath) requireOptimizationLowering(t, translation.Optimization, optimize.LoweringCountStoreFastPath) require.Empty(t, translation.Optimization.SkippedLowerings) - require.Equal(t, "select count(*)::int8 from node n0;", strings.Join(strings.Fields(formattedQuery), " ")) + require.Equal(t, "select count(*)::int8 from node n0;", strings.Join(strings.Fields(formattedQuery.Statement), " ")) } func TestOptimizerSafetyCountStoreFastPathKeepsKindConstraintAndAlias(t *testing.T) { @@ -222,7 +222,7 @@ func TestOptimizerSafetyCountStoreFastPathKeepsKindConstraintAndAlias(t *testing requirePlannedOptimizationLowering(t, translation.Optimization, optimize.LoweringCountStoreFastPath) requireOptimizationLowering(t, translation.Optimization, optimize.LoweringCountStoreFastPath) - require.Equal(t, "select count(*)::int8 as total from node n0 where n0.kind_ids operator (pg_catalog.@>) array [8]::int2[];", strings.Join(strings.Fields(formattedQuery), " ")) + require.Equal(t, "select count(*)::int8 as total from node n0 where n0.kind_ids operator (pg_catalog.@>) array [8]::int2[];", strings.Join(strings.Fields(formattedQuery.Statement), " ")) } func TestOptimizerSafetyCountStoreFastPathSupportsNodeCountStar(t *testing.T) { @@ -234,7 +234,7 @@ func TestOptimizerSafetyCountStoreFastPathSupportsNodeCountStar(t *testing.T) { requirePlannedOptimizationLowering(t, translation.Optimization, optimize.LoweringCountStoreFastPath) requireOptimizationLowering(t, translation.Optimization, optimize.LoweringCountStoreFastPath) - require.Equal(t, "select count(*)::int8 as total from node n0 where n0.kind_ids operator (pg_catalog.@>) array [8]::int2[];", strings.Join(strings.Fields(formattedQuery), " ")) + require.Equal(t, "select count(*)::int8 as total from node n0 where n0.kind_ids operator (pg_catalog.@>) array [8]::int2[];", strings.Join(strings.Fields(formattedQuery.Statement), " ")) } func TestOptimizerSafetyCountStoreFastPathUsesBaseEdgeCount(t *testing.T) { @@ -247,7 +247,7 @@ func TestOptimizerSafetyCountStoreFastPathUsesBaseEdgeCount(t *testing.T) { requirePlannedOptimizationLowering(t, translation.Optimization, optimize.LoweringCountStoreFastPath) requireOptimizationLowering(t, translation.Optimization, optimize.LoweringCountStoreFastPath) requireSkippedOptimizationLowering(t, translation.Optimization, optimize.LoweringProjectionPruning, "superseded by CountStoreFastPath") - require.Equal(t, "select count(*)::int8 from edge e0 join node n0 on n0.id = e0.start_id join node n1 on n1.id = e0.end_id where e0.kind_id = any (array [10]::int2[]);", strings.Join(strings.Fields(formattedQuery), " ")) + require.Equal(t, "select count(*)::int8 from edge e0 join node n0 on n0.id = e0.start_id join node n1 on n1.id = e0.end_id where e0.kind_id = any (array [10]::int2[]);", strings.Join(strings.Fields(formattedQuery.Statement), " ")) } func TestOptimizerSafetyCountStoreFastPathUsesSparseEdgeKindCount(t *testing.T) { @@ -256,7 +256,7 @@ func TestOptimizerSafetyCountStoreFastPathUsesSparseEdgeKindCount(t *testing.T) translation := optimizerSafetyTranslation(t, `MATCH ()-[r:Enroll]->() RETURN count(r)`) formattedQuery, err := Translated(translation) require.NoError(t, err) - normalizedQuery := strings.Join(strings.Fields(formattedQuery), " ") + normalizedQuery := strings.Join(strings.Fields(formattedQuery.Statement), " ") requirePlannedOptimizationLowering(t, translation.Optimization, optimize.LoweringCountStoreFastPath) requireOptimizationLowering(t, translation.Optimization, optimize.LoweringCountStoreFastPath) @@ -272,7 +272,7 @@ func TestOptimizerSafetyCountStoreFastPathUsesUntypedEdgeCount(t *testing.T) { translation := optimizerSafetyTranslation(t, `MATCH ()-[r]->() RETURN count(r)`) formattedQuery, err := Translated(translation) require.NoError(t, err) - normalizedQuery := strings.Join(strings.Fields(formattedQuery), " ") + normalizedQuery := strings.Join(strings.Fields(formattedQuery.Statement), " ") requirePlannedOptimizationLowering(t, translation.Optimization, optimize.LoweringCountStoreFastPath) requireOptimizationLowering(t, translation.Optimization, optimize.LoweringCountStoreFastPath) @@ -292,7 +292,7 @@ func TestOptimizerSafetyCountStoreFastPathSupportsEdgeCountStar(t *testing.T) { requirePlannedOptimizationLowering(t, translation.Optimization, optimize.LoweringCountStoreFastPath) requireOptimizationLowering(t, translation.Optimization, optimize.LoweringCountStoreFastPath) requireSkippedOptimizationLowering(t, translation.Optimization, optimize.LoweringProjectionPruning, "superseded by CountStoreFastPath") - require.Equal(t, "select count(*)::int8 from edge e0 join node n0 on n0.id = e0.start_id join node n1 on n1.id = e0.end_id where e0.kind_id = any (array [10]::int2[]);", strings.Join(strings.Fields(formattedQuery), " ")) + require.Equal(t, "select count(*)::int8 from edge e0 join node n0 on n0.id = e0.start_id join node n1 on n1.id = e0.end_id where e0.kind_id = any (array [10]::int2[]);", strings.Join(strings.Fields(formattedQuery.Statement), " ")) } func TestOptimizerSafetyADCSQueryPrunesExpansionEdgeCarry(t *testing.T) { @@ -301,7 +301,7 @@ func TestOptimizerSafetyADCSQueryPrunesExpansionEdgeCarry(t *testing.T) { translation := optimizerSafetyTranslation(t, optimizerADCSQuery) formattedQuery, err := Translated(translation) require.NoError(t, err) - normalizedQuery := strings.Join(strings.Fields(formattedQuery), " ") + normalizedQuery := strings.Join(strings.Fields(formattedQuery.Statement), " ") requirePlannedOptimizationLowering(t, translation.Optimization, "ExpansionSuffixPushdown") requirePlannedOptimizationLowering(t, translation.Optimization, "PredicatePlacement") @@ -422,7 +422,7 @@ RETURN p formattedQuery, err := Translated(translation) require.NoError(t, err) - normalizedQuery := strings.Join(strings.Fields(formattedQuery), " ") + normalizedQuery := strings.Join(strings.Fields(formattedQuery.Statement), " ") require.Contains(t, normalizedQuery, "(s1.n0).id = e0.start_id") require.Contains(t, normalizedQuery, "(s1.n1).id = e0.end_id") @@ -495,7 +495,7 @@ RETURN p formattedQuery, err := Translated(translation) require.NoError(t, err) - normalizedQuery := strings.Join(strings.Fields(formattedQuery), " ") + normalizedQuery := strings.Join(strings.Fields(formattedQuery.Statement), " ") requirePlannedOptimizationLowering(t, translation.Optimization, optimize.LoweringPredicatePlacement) requireOptimizationLowering(t, translation.Optimization, optimize.LoweringPredicatePlacement) @@ -517,7 +517,7 @@ RETURN dst formattedQuery, err := Translated(translation) require.NoError(t, err) - normalizedQuery := strings.Join(strings.Fields(formattedQuery), " ") + normalizedQuery := strings.Join(strings.Fields(formattedQuery.Statement), " ") requirePlannedOptimizationLowering(t, translation.Optimization, optimize.LoweringPredicatePlacement) requireOptimizationLowering(t, translation.Optimization, optimize.LoweringPredicatePlacement) @@ -540,7 +540,7 @@ RETURN s formattedQuery, err := Translated(translation) require.NoError(t, err) - normalizedQuery := strings.Join(strings.Fields(formattedQuery), " ") + normalizedQuery := strings.Join(strings.Fields(formattedQuery.Statement), " ") require.Contains(t, normalizedQuery, "not exists (select 1 from edge e0") requirePlannedOptimizationLowering(t, translation.Optimization, "PredicatePlacement") @@ -590,7 +590,7 @@ RETURN p `) formattedQuery, err := Translated(translation) require.NoError(t, err) - normalizedQuery := strings.Join(strings.Fields(formattedQuery), " ") + normalizedQuery := strings.Join(strings.Fields(formattedQuery.Statement), " ") requirePlannedOptimizationLowering(t, translation.Optimization, "TraversalDirectionSelection") requireOptimizationLowering(t, translation.Optimization, "TraversalDirectionSelection") @@ -609,7 +609,7 @@ RETURN p `) formattedQuery, err := Translated(translation) require.NoError(t, err) - normalizedQuery := strings.Join(strings.Fields(strings.ToLower(formattedQuery)), " ") + normalizedQuery := strings.Join(strings.Fields(strings.ToLower(formattedQuery.Statement)), " ") requirePlannedOptimizationLowering(t, translation.Optimization, optimize.LoweringExactRangeExpansion) requireOptimizationLowering(t, translation.Optimization, optimize.LoweringExactRangeExpansion) @@ -631,7 +631,7 @@ RETURN p `) formattedQuery, err := Translated(translation) require.NoError(t, err) - normalizedQuery := strings.Join(strings.Fields(strings.ToLower(formattedQuery)), " ") + normalizedQuery := strings.Join(strings.Fields(strings.ToLower(formattedQuery.Statement)), " ") requirePlannedOptimizationLowering(t, translation.Optimization, optimize.LoweringExactRangeExpansion) requireOptimizationLowering(t, translation.Optimization, optimize.LoweringExactRangeExpansion) @@ -652,7 +652,7 @@ RETURN a `) formattedQuery, err := Translated(translation) require.NoError(t, err) - normalizedQuery := strings.Join(strings.Fields(formattedQuery), " ") + normalizedQuery := strings.Join(strings.Fields(formattedQuery.Statement), " ") requirePlannedOptimizationLowering(t, translation.Optimization, optimize.LoweringExactRangeExpansion) requireOptimizationLowering(t, translation.Optimization, optimize.LoweringExactRangeExpansion) @@ -681,7 +681,7 @@ RETURN p `) formattedQuery, err := Translated(translation) require.NoError(t, err) - normalizedQuery := strings.Join(strings.Fields(strings.ToLower(formattedQuery)), " ") + normalizedQuery := strings.Join(strings.Fields(strings.ToLower(formattedQuery.Statement)), " ") requirePlannedOptimizationLowering(t, translation.Optimization, optimize.LoweringExactRangeExpansion) requireOptimizationLowering(t, translation.Optimization, optimize.LoweringExactRangeExpansion) @@ -716,7 +716,7 @@ RETURN p `) formattedQuery, err := Translated(translation) require.NoError(t, err) - normalizedQuery := strings.Join(strings.Fields(strings.ToLower(formattedQuery)), " ") + normalizedQuery := strings.Join(strings.Fields(strings.ToLower(formattedQuery.Statement)), " ") requirePlannedOptimizationLowering(t, translation.Optimization, optimize.LoweringPathRelationshipPredicate) requireOptimizationLowering(t, translation.Optimization, optimize.LoweringPathRelationshipPredicate) @@ -737,7 +737,7 @@ RETURN p `) formattedQuery, err := Translated(translation) require.NoError(t, err) - normalizedQuery := strings.Join(strings.Fields(strings.ToLower(formattedQuery)), " ") + normalizedQuery := strings.Join(strings.Fields(strings.ToLower(formattedQuery.Statement)), " ") requirePlannedOptimizationLowering(t, translation.Optimization, optimize.LoweringPathRelationshipPredicate) requireOptimizationLowering(t, translation.Optimization, optimize.LoweringPathRelationshipPredicate) @@ -762,7 +762,7 @@ LIMIT 100 formattedQuery, err := Translated(translation) require.NoError(t, err) var ( - normalizedQuery = strings.Join(strings.Fields(formattedQuery), " ") + normalizedQuery = strings.Join(strings.Fields(formattedQuery.Statement), " ") lowerQuery = strings.ToLower(normalizedQuery) ) @@ -848,7 +848,7 @@ LIMIT 100 `) formattedQuery, err := Translated(translation) require.NoError(t, err) - normalizedQuery := strings.Join(strings.Fields(strings.ToLower(formattedQuery)), " ") + normalizedQuery := strings.Join(strings.Fields(strings.ToLower(formattedQuery.Statement)), " ") requirePlannedOptimizationLowering(t, translation.Optimization, optimize.LoweringAggregateTraversalCount) requireOptimizationLowering(t, translation.Optimization, optimize.LoweringAggregateTraversalCount) @@ -870,7 +870,7 @@ LIMIT 100 `) formattedQuery, err := Translated(translation) require.NoError(t, err) - normalizedQuery := strings.Join(strings.Fields(strings.ToLower(formattedQuery)), " ") + normalizedQuery := strings.Join(strings.Fields(strings.ToLower(formattedQuery.Statement)), " ") requireOptimizationLowering(t, translation.Optimization, optimize.LoweringAggregateTraversalCount) require.Contains(t, normalizedQuery, "where traversal.depth < 4") @@ -891,7 +891,7 @@ LIMIT 100 `) formattedQuery, err := Translated(translation) require.NoError(t, err) - normalizedQuery := strings.Join(strings.Fields(strings.ToLower(formattedQuery)), " ") + normalizedQuery := strings.Join(strings.Fields(strings.ToLower(formattedQuery.Statement)), " ") requireOptimizationLowering(t, translation.Optimization, optimize.LoweringAggregateTraversalCount) require.Contains(t, normalizedQuery, "join edge e on e.end_id = candidate_sources.root_id") @@ -912,7 +912,7 @@ LIMIT 100 `) formattedQuery, err := Translated(translation) require.NoError(t, err) - normalizedQuery := strings.Join(strings.Fields(strings.ToLower(formattedQuery)), " ") + normalizedQuery := strings.Join(strings.Fields(strings.ToLower(formattedQuery.Statement)), " ") requireOptimizationLowering(t, translation.Optimization, optimize.LoweringAggregateTraversalCount) require.Contains(t, normalizedQuery, "(source_node.id, source_node.kind_ids, source_node.properties)::nodecomposite as user") @@ -935,7 +935,7 @@ LIMIT 100 `) formattedQuery, err := Translated(translation) require.NoError(t, err) - normalizedQuery := strings.Join(strings.Fields(strings.ToLower(formattedQuery)), " ") + normalizedQuery := strings.Join(strings.Fields(strings.ToLower(formattedQuery.Statement)), " ") requireOptimizationLowering(t, translation.Optimization, optimize.LoweringAggregateTraversalCount) require.Contains(t, normalizedQuery, "terminal_nodes(id) as materialized") @@ -961,7 +961,7 @@ LIMIT 100 }) formattedQuery, err := Translated(translation) require.NoError(t, err) - normalizedQuery := strings.Join(strings.Fields(strings.ToLower(formattedQuery)), " ") + normalizedQuery := strings.Join(strings.Fields(strings.ToLower(formattedQuery.Statement)), " ") parameterValues := make([]any, 0, len(translation.Parameters)) for _, value := range translation.Parameters { @@ -992,7 +992,7 @@ LIMIT 100 }) formattedQuery, err := Translated(translation) require.NoError(t, err) - normalizedQuery := strings.Join(strings.Fields(strings.ToLower(formattedQuery)), " ") + normalizedQuery := strings.Join(strings.Fields(strings.ToLower(formattedQuery.Statement)), " ") requireOptimizationLowering(t, translation.Optimization, optimize.LoweringAggregateTraversalCount) require.Contains(t, normalizedQuery, "source_node.properties -> 'enabled'") @@ -1134,7 +1134,7 @@ RETURN p formattedQuery, err := Translated(translation) require.NoError(t, err) - normalizedQuery := strings.Join(strings.Fields(formattedQuery), " ") + normalizedQuery := strings.Join(strings.Fields(formattedQuery.Statement), " ") require.Contains(t, normalizedQuery, "bidirectional_asp_harness") requirePlannedOptimizationLowering(t, translation.Optimization, "ShortestPathStrategySelection") @@ -1155,7 +1155,7 @@ RETURN p formattedQuery, err := Translated(translation) require.NoError(t, err) - normalizedQuery := strings.Join(strings.Fields(formattedQuery), " ") + normalizedQuery := strings.Join(strings.Fields(formattedQuery.Statement), " ") require.Contains(t, normalizedQuery, "unidirectional_sp_harness") require.Contains(t, normalizedQuery, "traversal_terminal_filter") @@ -1175,7 +1175,7 @@ LIMIT 1000 formattedQuery, err := Translated(translation) require.NoError(t, err) - normalizedQuery := strings.Join(strings.Fields(formattedQuery), " ") + normalizedQuery := strings.Join(strings.Fields(formattedQuery.Statement), " ") require.Contains(t, normalizedQuery, "unidirectional_sp_harness") require.Contains(t, normalizedQuery, "traversal_terminal_filter") @@ -1222,7 +1222,7 @@ func TestOptimizerSafetyShortestPathRootCarriesUnwindSources(t *testing.T) { formattedQuery, err := Translated(translation) require.NoError(t, err) - normalizedQuery := strings.Join(strings.Fields(formattedQuery), " ") + normalizedQuery := strings.Join(strings.Fields(formattedQuery.Statement), " ") require.Contains(t, normalizedQuery, "unidirectional_sp_harness") require.Contains(t, normalizedQuery, "unnest(array ['source']::text[]) as i0") @@ -1242,7 +1242,7 @@ func TestOptimizerSafetyShortestPathTerminalCarriesUnwindSources(t *testing.T) { formattedQuery, err := Translated(translation) require.NoError(t, err) - normalizedQuery := strings.Join(strings.Fields(formattedQuery), " ") + normalizedQuery := strings.Join(strings.Fields(formattedQuery.Statement), " ") require.Contains(t, normalizedQuery, "unidirectional_sp_harness") require.Contains(t, normalizedQuery, "unnest(array ['target']::text[]) as i0") @@ -1328,7 +1328,7 @@ RETURN p `) formattedQuery, err := Translated(translation) require.NoError(t, err) - normalizedQuery := strings.Join(strings.Fields(formattedQuery), " ") + normalizedQuery := strings.Join(strings.Fields(formattedQuery.Statement), " ") requirePlannedOptimizationLowering(t, translation.Optimization, "ExpansionSuffixPushdown") requireOptimizationLowering(t, translation.Optimization, "ExpansionSuffixPushdown") diff --git a/cypher/models/pgsql/translate/predicate_test.go b/cypher/models/pgsql/translate/predicate_test.go index 729712ba..902be276 100644 --- a/cypher/models/pgsql/translate/predicate_test.go +++ b/cypher/models/pgsql/translate/predicate_test.go @@ -76,7 +76,7 @@ func translatePredicateQuery(t *testing.T, cypherQuery string, parameters map[st formatted, err := Translated(translation) require.NoError(t, err) - return formatted + return formatted.Statement } func TestExclusiveDisjunctionTranslates(t *testing.T) { @@ -271,8 +271,8 @@ RETURN n`) // Extract the individual CTE bodies so each assertion is scoped to the CTE it // describes, rather than matching anywhere in the flattened query string. - s1Body := extractCTEBody(t, formatted, "s1") - s2Body := extractCTEBody(t, formatted, "s2") + s1Body := extractCTEBody(t, formatted.Statement, "s1") + s2Body := extractCTEBody(t, formatted.Statement, "s2") // The predicate root CTE (s1) must NOT have the outer MATCH frame (s0) as a // comma-joined FROM source. OmitPreviousFrameSource suppresses it so the subquery diff --git a/cypher/models/pgsql/visualization/visualizer.go b/cypher/models/pgsql/visualization/visualizer.go index 601bccab..17eca175 100644 --- a/cypher/models/pgsql/visualization/visualizer.go +++ b/cypher/models/pgsql/visualization/visualizer.go @@ -72,7 +72,8 @@ func SQLToDigraph(node pgsql.SyntaxNode) (Graph, error) { if title, err := format.SyntaxNode(node); err != nil { return Graph{}, err } else { - visualizer.Graph.Title = title + // TODO: do we need to use Parameters here somehow? + visualizer.Graph.Title = title.Statement } return visualizer.Graph, walk.PgSQL(node, visualizer) diff --git a/drivers/pg/transaction.go b/drivers/pg/transaction.go index 7bf4bbd7..dc44b487 100644 --- a/drivers/pg/transaction.go +++ b/drivers/pg/transaction.go @@ -3,6 +3,7 @@ package pg import ( "context" "fmt" + "maps" "github.com/specterops/dawgs/cypher/models/pgsql" "github.com/specterops/dawgs/cypher/models/pgsql/translate" @@ -284,7 +285,8 @@ func (s *transaction) Query(query string, parameters map[string]any) graph.Resul } else if sqlQuery, err := translate.Translated(translated); err != nil { return graph.NewErrorResult(err) } else { - return s.Raw(sqlQuery, translated.Parameters) + maps.Copy(translated.Parameters, sqlQuery.Parameters) + return s.Raw(sqlQuery.Statement, translated.Parameters) } } diff --git a/tools/dawgrun/pkg/commands/cypher.go b/tools/dawgrun/pkg/commands/cypher.go index 2d2e05ca..0c1fff8b 100644 --- a/tools/dawgrun/pkg/commands/cypher.go +++ b/tools/dawgrun/pkg/commands/cypher.go @@ -102,15 +102,21 @@ func translateToPsqlCmd() CommandDesc { return fmt.Errorf("could not format translated statement into a string query: %w", err) } - formattedQuery, err := sqlfmt.Format(sqlQuery, &sqlfmt.Options{ + formattedQuery, err := sqlfmt.Format(sqlQuery.Statement, &sqlfmt.Options{ Distance: 0, }) if err != nil { ctx.output.Warnf("could not format query: %s", err.Error()) - formattedQuery = sqlQuery + formattedQuery = sqlQuery.Statement } ctx.output.WriteHighlighted(formattedQuery, "postgres") + if len(sqlQuery.Parameters) > 0 { + fmt.Fprintf(ctx.output, "PARAMETERS\n\n") + ctx.output.WriteHighlighted(spew.Sdump(sqlQuery.Parameters), "golang") + fmt.Fprintf(ctx.output, "\n") + } + return nil }, } @@ -163,19 +169,24 @@ func explainAsPsqlCmd() CommandDesc { return fmt.Errorf("could not format translated statement into a string query: %w", err) } - formattedQuery, err := sqlfmt.Format(sqlQuery, &sqlfmt.Options{ + formattedQuery, err := sqlfmt.Format(sqlQuery.Statement, &sqlfmt.Options{ Distance: 2, }) if err != nil { ctx.output.Warnf("could not format query: %s", err.Error()) - formattedQuery = sqlQuery + formattedQuery = sqlQuery.Statement } explainSQLQuery := fmt.Sprintf("EXPLAIN %s", formattedQuery) ctx.output.WriteHighlighted(explainSQLQuery, "postgres") fmt.Fprint(ctx.output, "\n\n") + if len(sqlQuery.Parameters) > 0 { + fmt.Fprintf(ctx.output, "PARAMETERS\n\n") + ctx.output.WriteHighlighted(spew.Sdump(sqlQuery.Parameters), "golang") + fmt.Fprintf(ctx.output, "\n") + } err = conn.ReadTransaction(ctx, func(tx graph.Transaction) error { - result := tx.Raw(explainSQLQuery, nil) + result := tx.Raw(explainSQLQuery, sqlQuery.Parameters) if err := result.Error(); err != nil { return fmt.Errorf("error running raw query: '%s': %w", explainSQLQuery, err) }