diff --git a/.github/workflows/build.yml b/.github/workflows/build.yml index f1c491024..59de1cbf8 100644 --- a/.github/workflows/build.yml +++ b/.github/workflows/build.yml @@ -2183,7 +2183,7 @@ jobs: PYTHONPATH: '${{ env.PYTHONPATH }}:${{ github.workspace }}/test/python' IS_SKIP_MCP_TEST: 'true' if: success() && env.CI_IS_EXPRESS != 'true' && matrix.platform == 'linux/amd64' && env.BUILD_IMAGE_REQUIRED == 'true' && matrix.db_backend == 'sqlite' - timeout-minutes: ${{ vars.DEFAULT_STEP_TIMEOUT_MIN == '' && 20 || vars.DEFAULT_STEP_TIMEOUT_MIN }} + timeout-minutes: ${{ vars.DEFAULT_LONG_STEP_TIMEOUT_MIN == '' && 40 || vars.DEFAULT_LONG_STEP_TIMEOUT_MIN }} run: | sudo rm -rf test/tmp || true mkdir -p test/tmp diff --git a/.gitignore b/.gitignore index 87aa51f07..2f5031866 100644 --- a/.gitignore +++ b/.gitignore @@ -46,3 +46,5 @@ mcp-audit-*.log ## Generic sink default name (used by pkg/sink tests / non-MCP callers) sink_*.log + +/*.db diff --git a/docs/preview.md b/docs/preview.md index 1bb85758c..c3eb32a13 100644 --- a/docs/preview.md +++ b/docs/preview.md @@ -36,7 +36,7 @@ _googleProject="stackql-demo" && \ ### What streams, and what does not - Use `--output jsonl` or `--output otel`. Each row is written and flushed as it arrives. The default table output holds every row until the query completes. -- `ORDER BY`, `GROUP BY`, `HAVING`, `DISTINCT` and aggregates are refused on these relations, since each needs every row before it can emit one. +- `ORDER BY`, `GROUP BY`, `HAVING`, `DISTINCT` and aggregates are refused on these relations, since each needs every row before it can emit one, unless [staging](#staging) is enabled. - `LIMIT` without `ORDER BY` is pushed down and stops the requests early. - Supported joins are `INNER JOIN` and `LEFT JOIN` with `ON`. A condition in `ON` that feeds a required input of the joined relation becomes a request per row of the left side. @@ -155,3 +155,67 @@ from stackql_unstable_google.iam.service_accounts s inner join stackql_unstable_google.iam.service_account_keys k on k.serviceAccountsId = s.email where s.projectsId = '${_googleProject}' and k.projectsId = '${_googleProject}';" ``` + +## Staging + +Add `"staging":true` to `--preview` to run `ORDER BY`, `GROUP BY`, `HAVING`, `DISTINCT`, aggregates and `OFFSET` over `stackql_unstable_*` relations. `omnisdk` still runs the joins and filters; its final result is staged in one query-owned table in the SQL backend, which evaluates the rest of the statement. Queries without those clauses stream as before and stage nothing. See [the staging design note](/docs/technical/omnisdk_staging.md). + +Staged results are not streamed: the final result is read in full before it is returned. `SELECT *` and subqueries are refused when staging applies. + +### Live test + +Each query below has been run against the live GitHub API. + +```bash +./build/stackql exec "registry pull github v26.08.00448;" + +## Chuck these in ./cicd/vol/vendor-secrets/secrets.sh +## export STACKQL_GITHUB_USERNAME='' +## export STACKQL_GITHUB_PASSWORD='' + +source ./cicd/vol/vendor-secrets/secrets.sh + +## Ordering, LIMIT and OFFSET prior and preview +./build/stackql exec --auth '{ "github": { "credentialsenvvar": "STACKQL_GITHUB_TOKEN", "type": "api_key", "valuePrefix": "Bearer " } }' --output csv \ +"select name, stargazers_count + from github.repos.repos + where org = 'stackql' + order by stargazers_count desc limit 5 offset 1;" + + +./build/stackql exec --preview='{"unstable":true,"staging":true}' --auth '{ "github": { "credentialsenvvar": "STACKQL_GITHUB_TOKEN", "type": "api_key", "valuePrefix": "Bearer " } }' --output csv \ +"select name, stargazers_count + from stackql_unstable_github.repos.repos + where org = 'stackql' + order by stargazers_count desc limit 5 offset 1;" + +./build/stackql exec --preview='{"omni":"all","staging":true}' --auth '{ "github": { "credentialsenvvar": "STACKQL_GITHUB_TOKEN", "type": "api_key", "valuePrefix": "Bearer " } }' --output csv \ +"select name, stargazers_count + from github.repos.repos + where org = 'stackql' + order by stargazers_count desc limit 5 offset 1;" + + +## Grouping and aggregation. +./build/stackql exec --preview='{"unstable":true,"staging":true}' --output csv \ +"select language, count(*) as repo_count + from stackql_unstable_github.repos.repos + where org = 'stackql' + group by language order by repo_count desc;" + +## Aggregate with no column references. +./build/stackql exec --preview='{"unstable":true,"staging":true}' --output csv \ +"select count(*) as member_count + from stackql_unstable_github.orgs.members + where org = 'stackql';" +``` + +Without `"staging":true` the same queries are refused, for example with `ORDER BY cannot be applied to stackql_unstable_* relations`. + +Staged tables are dropped as soon as each result has been read. To confirm, add `--sqlBackend='{"dsn":"file:/tmp/stackql-staging.db"}'` to the commands above, then: + +```bash +sqlite3 /tmp/stackql-staging.db "select name from sqlite_master where name like '__iql__.queries.%';" +``` + +An empty result means every staged table was released. diff --git a/docs/technical/omnisdk_staging.md b/docs/technical/omnisdk_staging.md new file mode 100644 index 000000000..6c27bd5d6 --- /dev/null +++ b/docs/technical/omnisdk_staging.md @@ -0,0 +1,91 @@ + +# `omnisdk` staging + +`omnisdk` results are cursor streams. When the query plan can preserve SQL semantics incrementally, batches can flow directly to the end user. Otherwise, batches can be staged in an RDBMS for relational operations such as ordering, aggregation, and set operations. + +Unlike the eager, per-row RDBMS ingestion used by the `any-sdk` path, Omni staging is conditional and batched. The staging tablespace sits alongside StackQL-owned relations in the selected RDBMS, whether embedded SQLite or TCP-routed PostgreSQL. + +## High level details of staging + +Let $Q$ be a StackQL query with query ID $q$. Omni executes the joins and exchanges in its own query plan and produces a single final relation $R_Q$. An Omni plan may contain multiple exchanges, but exchange boundaries do not create staging tables or query IDs. + +The cursor partitions $R_Q$ into batches without changing its contents: + +$$ +R_Q = B_{1} \mathbin{\|} B_{2} \mathbin{\|} \cdots \mathbin{\|} B_{n} +$$ + +Here $\|$ denotes concatenation in cursor order. Batch size affects transport and insertion cost, not the relational result. The SQL operators that remain after Omni planning (for example global `ORDER BY`, aggregation, `DISTINCT`, set operations, or a `LIMIT` that depends on ordering) are evaluated over $R_Q$. If each of them can be evaluated incrementally within the chosen streaming-state bound, batches flow directly to the result consumer and no table is created. Otherwise $R_Q$ is staged, as a multiset, in exactly one query-owned table: + +$$ +T_{q} = \biguplus_{k=1}^{n} B_{k} +$$ + +The RDBMS then evaluates the remaining SQL over $T_q$, and its result cursor becomes the output stream. The planner, not the mere presence of an Omni exchange, determines whether materialization is needed. + +A blocking operation whose inputs Omni cannot combine into one relation (for example a set operation over two independent Omni plans) is not staged as several tables under one query ID. Pre-analysis instead splits $Q$ into child queries $q_1, \ldots, q_m$; each child owns at most one table $T_{q_j}$, and the parent evaluates the operation over those tables. Child IDs therefore correspond only to real pre-analysis splits, and the invariant is one table per query ID. + +Query IDs are allocated by the RDBMS itself (a PostgreSQL sequence, or an SQLite `AUTOINCREMENT` table), so they are collision free across concurrent queries and across processes sharing a backend. Rows from two queries never share a table, and cleanup is scoped to a query and its descendants. + +Reading a batch and inserting it form a back-pressured pipeline: the next batch is requested only after the previous one is written, rather than eagerly loading every result into application memory. Cursor exhaustion marks completion; query ownership tracks staged tables for cleanup on success, error, or cancellation. + +## Concrete descisions + +### Decision 1: Savage cut in transaction control counters + +The prior `any-sdk` implementation is supported by counters for: + +- `Generation ID`. +- `Sessions ID`. +- `Transaction ID`. +- `Insert ID`. + +The new, `omnisdk` implementation will require only one counter, `Query ID`. This is because staging happens only per query. We do reserve the right to split queries, so shall maintain a concurrency safe hierarchy store for `Query ID` parent-child associations. + +### Decision 2: New tablespace for omnisdk staging + + + +```bash +./build/stackql exec \ + --sqlBackend='{"dsn":"file:./stackql.db"}' \ + "SELECT + v.vpc_id, + s.subnet_id + FROM aws.ec2.vpcs AS v + INNER JOIN aws.ec2.subnets AS s + ON v.vpc_id = s.vpc_id + WHERE v.region = 'ap-southeast-2' + AND s.region = 'ap-southeast-2';" +``` + +**Figure MQ-1** Model query 1. A simple working query. + +--- + +**Table T-1**: Tablespace comparison for model query MQ-1. + +| sdk | RDBMS | tables | +|---|---|---| +| any-sdk | sqlite | `"aws.ec2.vpcs.generation_"`,
`"aws.ec2.subnets.generation_"` | +| any-sdk | postgres | `""."aws.ec2.vpcs.generation_"`,
`""."aws.ec2.subnets.generation_"` | +| omnisdk | sqlite | `"__iql__.queries."` | +| omnisdk | postgres | `"".""` | + +`` is configured with `schemata.querySchema` in `--sqlBackend` and defaults to `stackql_queries`. + +## Decision 3: GC simplification for omnisdk + +For omnisdk, any materializations needed will be created eagerly, **in the same place for prior** and all query tables can be marked for deletion immediately where no cache is in operation, or at whatever future time if cacheing is in effect. We will need a robust mechanism in place from day 1, default to no cache. + +This does imply a keyval store to look up tombstone times and find query ID by query plaintext. Intuitively, I favour a new GC mechanism with the old one phased out when `any-sdk` is decommissioned. + + +## Terse comparison prior any-sdk vs omnisdk + +| Aspect | any-sdk | omnisdk | Commment | +|----|----|----|----| +| RDBMS Ingestion | Per API response "record" | Per query. AOT configurable and runtime responsive batching | omnisdk clearly better performance best case | +| SQL translation of relation names | Recursive: global and per API call | Per query and flat | omnisdk has appealing simplicity and debug property of thin layer queries | +| SQL control counters | Mutiple counters and recursively applied: global and per API call | Per query and single counter only | | +| SQL secure storage | - | - | We have decided not to address security or obfuscation at this time. This will happen in future versions | diff --git a/internal/stackql/intrinsic/doc.go b/internal/stackql/intrinsic/doc.go index f18b2caff..e06141c98 100644 --- a/internal/stackql/intrinsic/doc.go +++ b/internal/stackql/intrinsic/doc.go @@ -28,17 +28,21 @@ const UnstablePrefix = aot.DefaultProviderPrefix // into. They are documents read straight from disk, with none of the registry's // curation behind them, so nothing exposes them until a caller asks. func IsUnstableEnabled() bool { - return previewCfg.getUnstableEnabled() + return previewCfg.getUnstableEnabled() || previewCfg.getOmniAll() } -// docProvider is the bundle behind an unstable provider name, or false. +// docProvider is the bundle behind an unstable provider name, or false. Once +// every provider is routed to omnisdk, an unprefixed name is one too. func docProvider(name string) (string, bool) { if !IsUnstableEnabled() { return "", false } trimmed := strings.TrimSpace(name) if !strings.HasPrefix(strings.ToLower(trimmed), UnstablePrefix) { - return "", false + if !previewCfg.getOmniAll() || trimmed == "" || strings.EqualFold(trimmed, ProviderName) { + return "", false + } + return trimmed, true } bundle := trimmed[len(UnstablePrefix):] if bundle == "" { @@ -159,6 +163,9 @@ func docSelectFunc( node *sqlparser.Select, currentProvider string, ) (func() internaldto.ExecutorOutput, bool) { + if previewCfg.getStagingEnabled() && needsStaging(node) { + return stagedSelectFunc(ctx, node, currentProvider), true + } translated, err := translateSelect(node, currentProvider) if err != nil { return refuse(err), true @@ -192,15 +199,16 @@ func docMutationFunc( const mutationSuccessMessage = "The operation was despatched successfully" -// runDocQuery describes each relation, resolves the query and runs it. -func runDocQuery(ctx queryContext, translated docQuery) internaldto.ExecutorOutput { +// openDocQuery describes each relation, resolves the query and opens its +// cursor, returning it with the alias of the relation it reports. +func openDocQuery(ctx queryContext, translated docQuery) (omnisdk.Rows, string, error) { registry := registryRoot(ctx) q := translated.getQuery() tables := make(map[string]omnisdk.Table, len(q.From())+1) for _, join := range q.From() { tbl, describeErr := omnisdk.DescribeTable(registry, join.Resource().Handle()) if describeErr != nil { - return internaldto.NewErroneousExecutorOutput(describeErr) + return nil, "", describeErr } tables[join.Resource().Alias()] = tbl } @@ -209,7 +217,7 @@ func runDocQuery(ctx queryContext, translated docQuery) internaldto.ExecutorOutp tbl, describeErr := omnisdk.DescribeMutation( registry, target.Resource().Handle(), target.Verb().String()) if describeErr != nil { - return internaldto.NewErroneousExecutorOutput(describeErr) + return nil, "", describeErr } tables[target.Resource().Alias()] = tbl relation = target.Resource().Alias() @@ -218,7 +226,7 @@ func runDocQuery(ctx queryContext, translated docQuery) internaldto.ExecutorOutp } res, resolveErr := omnisdk.Resolve(q, tables) if resolveErr != nil { - return internaldto.NewErroneousExecutorOutput(resolveErr) + return nil, "", resolveErr } // omnisdk takes one credential per run: the first relation's cloud - a // mutation's target - leaving the rest to the canonical environment @@ -227,12 +235,22 @@ func runDocQuery(ctx queryContext, translated docQuery) internaldto.ExecutorOutp args.Tuning.Limit = translated.getLimit() plan, planErr := omnisdk.NewGraphSelectQuery(registry, res.Graph(), args) if planErr != nil { - return internaldto.NewErroneousExecutorOutput(planErr) + return nil, "", planErr } rows, openErr := plan.Open(context.Background()) if openErr != nil { - return internaldto.NewErroneousExecutorOutput(openErr) + return nil, "", openErr } + return rows, relation, nil +} + +// runDocQuery runs the query and streams its rows back. +func runDocQuery(ctx queryContext, translated docQuery) internaldto.ExecutorOutput { + rows, relation, err := openDocQuery(ctx, translated) + if err != nil { + return internaldto.NewErroneousExecutorOutput(err) + } + q := translated.getQuery() if q.Target() != nil && len(translated.getOutputs()) == 0 { return drainMutation(rows) } diff --git a/internal/stackql/intrinsic/intrinsic.go b/internal/stackql/intrinsic/intrinsic.go index c23655ab3..716ff7741 100644 --- a/internal/stackql/intrinsic/intrinsic.go +++ b/internal/stackql/intrinsic/intrinsic.go @@ -6,6 +6,7 @@ import ( "github.com/stackql/any-sdk/pkg/dto" "github.com/stackql/any-sdk/public/formulation" + "github.com/stackql/any-sdk/public/sqlengine" "github.com/stackql/stackql/internal/stackql/internal_data_transfer/internaldto" "github.com/stackql/stackql/internal/stackql/typing" "github.com/stackql/stackql/internal/stackql/util" @@ -50,6 +51,8 @@ type queryContext interface { GetTypingConfig() typing.Config GetAuthContext(providerName string) (*dto.AuthCtx, error) GetRuntimeContext() dto.RuntimeCtx + GetSQLEngine() sqlengine.SQLEngine + GetASTFormatter() sqlparser.NodeFormatter } func GeneratePrimitiveFunc( diff --git a/internal/stackql/intrinsic/omnisdk.go b/internal/stackql/intrinsic/omnisdk.go index 5a3335cda..baeddd164 100644 --- a/internal/stackql/intrinsic/omnisdk.go +++ b/internal/stackql/intrinsic/omnisdk.go @@ -27,6 +27,8 @@ const defaultBatchSize = 100 const defaultFlushInterval = 50 * time.Millisecond +const omniAll = "all" + func relationName(path string) string { return strings.ReplaceAll(path, ".", "_") } @@ -611,6 +613,12 @@ func omnisdkAuth(authCtx *dto.AuthCtx) *omnisdk.Auth { UsernameEnvVar: authCtx.EnvVarUsername, PasswordEnvVar: authCtx.EnvVarPassword, } + if strings.EqualFold(authCtx.Type, "api_key") && auth.Name == "" { + auth.Name = "Authorization" + if auth.ValuePrefix == "" { + auth.ValuePrefix = "Bearer " + } + } if credentials, credErr := authCtx.GetCredentialsBytes(); credErr == nil { auth.SecretAccessKey = string(credentials) auth.Credentials = string(credentials) @@ -727,6 +735,8 @@ type backendInput interface { getFlushInterval() time.Duration getInsecureSkipTLSVerify() bool getUnstableEnabled() bool + getStagingEnabled() bool + getOmniAll() bool } type standardBackendInput struct { @@ -735,6 +745,8 @@ type standardBackendInput struct { flushInterval time.Duration insecureSkipTLSVerify bool unstableEnabled bool + stagingEnabled bool + omniAll bool } // previewCfg is the parsed --preview argument. Cobra binds the raw string in @@ -755,6 +767,11 @@ type previewCfgDTO struct { Endpoint json.RawMessage `json:"endpoint"` InsecureSkipTLSVerify bool `json:"insecureSkipTLSVerify"` Unstable bool `json:"unstable"` + // Staging opts SELECTs over document-driven relations into RDBMS staging + // for the SQL omnisdk leaves unapplied, instead of refusing them. + Staging bool `json:"staging"` + // Omni "all" routes every provider through omnisdk, never any-sdk. + Omni string `json:"omni"` } func (c previewCfgDTO) endpoint() string { @@ -786,6 +803,8 @@ func newBackendInput(cfg previewCfgDTO) backendInput { flushInterval: defaultFlushInterval, insecureSkipTLSVerify: cfg.InsecureSkipTLSVerify, unstableEnabled: cfg.Unstable, + stagingEnabled: cfg.Staging, + omniAll: strings.EqualFold(cfg.Omni, omniAll), } if cfg.BatchSize > 0 { rv.batchSize = cfg.BatchSize @@ -806,6 +825,10 @@ func (b *standardBackendInput) getInsecureSkipTLSVerify() bool { return b.insecu func (b *standardBackendInput) getUnstableEnabled() bool { return b.unstableEnabled } +func (b *standardBackendInput) getStagingEnabled() bool { return b.stagingEnabled } + +func (b *standardBackendInput) getOmniAll() bool { return b.omniAll } + // sourceKey is the row key a column reads from: its own name, unless an alias // renamed it. func (c column) sourceKey() string { diff --git a/internal/stackql/intrinsic/omnisdk_test.go b/internal/stackql/intrinsic/omnisdk_test.go index b06bb3e77..b2ac0ce17 100644 --- a/internal/stackql/intrinsic/omnisdk_test.go +++ b/internal/stackql/intrinsic/omnisdk_test.go @@ -303,3 +303,23 @@ func TestRowsReachOutputBeforeNextPage(t *testing.T) { type writerFunc func([]byte) (int, error) func (f writerFunc) Write(b []byte) (int, error) { return f(b) } + +func TestOmnisdkAuthAPIKeyDefaultsMatchCanonicalProviders(t *testing.T) { + for _, tc := range []struct { + in dto.AuthCtx + wantName string + wantPrefix string + }{ + {in: dto.AuthCtx{Type: "api_key", ValuePrefix: "token "}, wantName: "Authorization", wantPrefix: "token "}, + {in: dto.AuthCtx{Type: "api_key"}, wantName: "Authorization", wantPrefix: "Bearer "}, + {in: dto.AuthCtx{Type: "api_key", Name: "X-Api-Key"}, wantName: "X-Api-Key", wantPrefix: ""}, + {in: dto.AuthCtx{Type: "bearer"}, wantName: "", wantPrefix: ""}, + } { + authCtx := tc.in + got := omnisdkAuth(&authCtx) + if got.Name != tc.wantName || got.ValuePrefix != tc.wantPrefix { + t.Errorf("%+v: got name %q prefix %q, want %q %q", + tc.in, got.Name, got.ValuePrefix, tc.wantName, tc.wantPrefix) + } + } +} diff --git a/internal/stackql/intrinsic/staged.go b/internal/stackql/intrinsic/staged.go new file mode 100644 index 000000000..4e9e75bc4 --- /dev/null +++ b/internal/stackql/intrinsic/staged.go @@ -0,0 +1,523 @@ +package intrinsic + +// A SELECT over document-driven relations whose remaining SQL needs every row +// - ordering, grouping, aggregation, de-duplication, OFFSET - is split in two +// when staging is enabled: omnisdk runs the joins and filters, its final +// relation is staged in one query-owned table, and the RDBMS evaluates the +// rest of the statement over that table. See docs/technical/omnisdk_staging.md. + +import ( + "context" + "database/sql" + "errors" + "fmt" + "io" + "strconv" + "strings" + "sync" + + "github.com/stackql-labs/omnisdk/pkg/omnisdk" + "github.com/stackql-labs/omnisdk/pkg/query" + "github.com/stackql/any-sdk/pkg/dto" + "github.com/stackql/stackql/internal/stackql/internal_data_transfer/internaldto" + "github.com/stackql/stackql/internal/stackql/omnistaging" + "github.com/stackql/stackql/internal/stackql/util" + + "github.com/stackql/stackql-parser/go/vt/sqlparser" +) + +// stagedRelationPlaceholder stands in for the query table while the outer +// statement is formatted; the table only exists once a query ID is allocated. +const stagedRelationPlaceholder = "__omnistaging_relation__" + +// stagedColumnPrefix names the staged columns; it keeps them clear of the +// aliases the outer statement reports. +const stagedColumnPrefix = "_omni_" + +const ( + stagedTextType = "text" + stagedNumericType = "numeric" +) + +// stagingManager is created on first use from the SQL backend the session +// already holds, and shared by every query thereafter. +// +//nolint:gochecknoglobals // one manager per backend, created once +var ( + stagingOnce sync.Once + stagingManager omnistaging.Manager + errStaging error +) + +func getStagingManager(ctx queryContext) (omnistaging.Manager, error) { + stagingOnce.Do(func() { + raw := ctx.GetRuntimeContext().SQLBackendCfgRaw + sqlCfg, err := dto.GetSQLBackendCfg(raw) + if err != nil { + errStaging = err + return + } + cfg, err := omnistaging.NewConfigFromSQLBackend(raw) + if err != nil { + errStaging = err + return + } + dialect, err := omnistaging.NewDialect(sqlCfg.GetSQLDialect(), cfg) + if err != nil { + errStaging = err + return + } + db, err := ctx.GetSQLEngine().GetDB() + if err != nil { + errStaging = err + return + } + stagingManager, errStaging = omnistaging.NewManager(context.Background(), db, dialect) + }) + return stagingManager, errStaging +} + +// needsStaging reports whether a SELECT carries SQL omnisdk leaves unapplied. +func needsStaging(node *sqlparser.Select) bool { + if len(unsupportedDocClauses(node)) > 0 { + return true + } + if node.Limit != nil && node.Limit.Offset != nil { + return true + } + hasAggregate := false + //nolint:errcheck // the visitor returns no error + _ = sqlparser.Walk(func(n sqlparser.SQLNode) (bool, error) { + if fn, isFunc := n.(*sqlparser.FuncExpr); isFunc && fn.IsAggregate() { + hasAggregate = true + return false, nil + } + return true, nil + }, node.SelectExprs) + return hasAggregate +} + +// stagedSelect is a SELECT split at the staging boundary: the omnisdk query +// producing the staged relation, and the statement the RDBMS runs over it. +type stagedSelect interface { + getSource() docQuery + getColumns() []string + getOutputNames() []string + renderOuter(tableName string) string +} + +type standardStagedSelect struct { + source docQuery + columns []string + outputNames []string + outerTemplate string +} + +func newStagedSelect(source docQuery, columns, outputNames []string, outerTemplate string) stagedSelect { + return &standardStagedSelect{ + source: source, columns: columns, outputNames: outputNames, outerTemplate: outerTemplate, + } +} + +func (s *standardStagedSelect) getSource() docQuery { return s.source } + +func (s *standardStagedSelect) getColumns() []string { return s.columns } + +// getOutputNames is the reported name of each select-list item, in order. +func (s *standardStagedSelect) getOutputNames() []string { return s.outputNames } + +func (s *standardStagedSelect) renderOuter(tableName string) string { + quoted := `"` + stagedRelationPlaceholder + `"` + if strings.Contains(s.outerTemplate, quoted) { + return strings.ReplaceAll(s.outerTemplate, quoted, tableName) + } + return strings.ReplaceAll(s.outerTemplate, stagedRelationPlaceholder, tableName) +} + +func stagedSelectFunc( + ctx queryContext, + node *sqlparser.Select, + currentProvider string, +) func() internaldto.ExecutorOutput { + staged, err := planStagedSelect(node, currentProvider, ctx.GetASTFormatter()) + if err != nil { + return refuse(err) + } + return func() internaldto.ExecutorOutput { return runStagedSelect(ctx, staged) } +} + +// stagedRefs assigns a staged column to each distinct column reference. The +// statement is never modified: each reference occurrence is renamed as the +// outer statement is formatted. +type stagedRefs struct { + byKey map[string]string + renamed map[*sqlparser.ColName]string + outputs []query.Output + names []string +} + +func newStagedRefs() *stagedRefs { + return &stagedRefs{byKey: make(map[string]string), renamed: make(map[*sqlparser.ColName]string)} +} + +func (r *stagedRefs) collect(expr sqlparser.SQLNode, aliases map[string]bool) error { + return sqlparser.Walk(func(n sqlparser.SQLNode) (bool, error) { + switch node := n.(type) { + case *sqlparser.Subquery: + return false, fmt.Errorf("a subquery cannot be staged for %s relations", UnstablePrefix+"*") + case *sqlparser.ColName: + qualifier := node.Qualifier.Name.GetRawVal() + name := node.Name.GetRawVal() + if qualifier == "" && aliases[strings.ToLower(name)] { + return false, nil + } + key := qualifier + "." + name + staged, seen := r.byKey[key] + if !seen { + staged = stagedColumnPrefix + strconv.Itoa(len(r.names)) + r.byKey[key] = staged + r.outputs = append(r.outputs, query.NewOutput(staged, query.NewColumn(qualifier, name))) + r.names = append(r.names, staged) + } + r.renamed[node] = staged + return false, nil + } + return true, nil + }, expr) +} + +// formatter wraps the backend's formatter, rendering each collected +// reference as its staged column. +func (r *stagedRefs) formatter(inner sqlparser.NodeFormatter) sqlparser.NodeFormatter { + return func(buf *sqlparser.TrackedBuffer, node sqlparser.SQLNode) { + if col, isCol := node.(*sqlparser.ColName); isCol { + if staged, isRenamed := r.renamed[col]; isRenamed { + node = &sqlparser.ColName{Name: sqlparser.NewColIdent(staged)} + } + } + if inner == nil { + node.Format(buf) + return + } + inner(buf, node) + } +} + +// collectSelectExprs collects the select list's references and returns the +// outer select list with the reported name of each item and the aliases the +// rest of the statement may name. +func (r *stagedRefs) collectSelectExprs( + exprs sqlparser.SelectExprs, +) (sqlparser.SelectExprs, []string, map[string]bool, error) { + aliases := make(map[string]bool, len(exprs)) + outputNames := make([]string, 0, len(exprs)) + selectExprs := make(sqlparser.SelectExprs, 0, len(exprs)) + for i, expr := range exprs { + aliased, isAliased := expr.(*sqlparser.AliasedExpr) + if !isAliased { + return nil, nil, nil, fmt.Errorf("'%s' cannot be staged for %s relations; name the columns", + sqlparser.String(expr), UnstablePrefix+"*") + } + // Output names match the streamed path: an alias, else a bare + // column's own name, else the expression as written. The last is + // reported by position, so the outer statement aliases it safely. + outerExpr := &sqlparser.AliasedExpr{Expr: aliased.Expr, As: aliased.As} + name := aliased.As.GetRawVal() + if name == "" { + if col, isCol := aliased.Expr.(*sqlparser.ColName); isCol { + name = col.Name.GetRawVal() + outerExpr.As = sqlparser.NewColIdent(name) + } else { + name = sqlparser.String(aliased.Expr) + outerExpr.As = sqlparser.NewColIdent(stagedColumnPrefix + "out_" + strconv.Itoa(i)) + } + } + outputNames = append(outputNames, name) + selectExprs = append(selectExprs, outerExpr) + aliases[strings.ToLower(outerExpr.As.GetRawVal())] = true + if err := r.collect(aliased.Expr, nil); err != nil { + return nil, nil, nil, err + } + } + return selectExprs, outputNames, aliases, nil +} + +// planStagedSelect splits a SELECT: FROM and WHERE go to omnisdk, which +// outputs every column the rest of the statement references; the rest is +// formatted over the staged table. +func planStagedSelect( + node *sqlparser.Select, + currentProvider string, + formatter sqlparser.NodeFormatter, +) (stagedSelect, error) { + t, where, err := translateSource(node, currentProvider) + if err != nil { + return nil, err + } + refs := newStagedRefs() + selectExprs, outputNames, aliases, err := refs.collectSelectExprs(node.SelectExprs) + if err != nil { + return nil, err + } + // GROUP BY and ORDER BY may name a select-list alias; HAVING may not. + for _, expr := range node.GroupBy { + if err = refs.collect(expr, aliases); err != nil { + return nil, err + } + } + if node.Having != nil { + if err = refs.collect(node.Having.Expr, nil); err != nil { + return nil, err + } + } + for _, order := range node.OrderBy { + if err = refs.collect(order.Expr, aliases); err != nil { + return nil, err + } + } + // A statement referencing no column, such as count(*), still needs one + // staged row per source row. omnisdk outputs only columns, so the rows + // are taken whole and staged as a single null placeholder column. + if len(refs.outputs) == 0 { + refs.outputs = append(refs.outputs, query.NewOutput("", query.NewStar(""))) + refs.names = append(refs.names, stagedColumnPrefix+"0") + } + q, err := query.New(t.joins, where, refs.outputs) + if err != nil { + return nil, err + } + limit, err := stagedLimit(node.Limit) + if err != nil { + return nil, err + } + outer := &sqlparser.Select{ + Distinct: node.Distinct, + SelectExprs: selectExprs, + From: sqlparser.TableExprs{&sqlparser.AliasedTableExpr{ + Expr: sqlparser.TableName{Name: sqlparser.NewTableIdent(stagedRelationPlaceholder)}, + }}, + GroupBy: node.GroupBy, + Having: node.Having, + OrderBy: node.OrderBy, + } + buf := sqlparser.NewTrackedBuffer(refs.formatter(formatter)) + outer.Format(buf) + return newStagedSelect( + newDocQuery(q, refs.names, 0, t.bundles), refs.names, outputNames, buf.String()+limit), nil +} + +// stagedLimit renders LIMIT and OFFSET in the form both backends accept. +func stagedLimit(limit *sqlparser.Limit) (string, error) { + if limit == nil { + return "", nil + } + var b strings.Builder + if limit.Rowcount != nil { + n, err := rowCount(limit.Rowcount) + if err != nil { + return "", err + } + fmt.Fprintf(&b, " LIMIT %d", n) + } + if limit.Offset != nil { + n, err := rowCount(limit.Offset) + if err != nil { + return "", err + } + fmt.Fprintf(&b, " OFFSET %d", n) + } + return b.String(), nil +} + +func rowCount(expr sqlparser.Expr) (int, error) { + val, isVal := expr.(*sqlparser.SQLVal) + if !isVal || val.Type != sqlparser.IntVal { + return 0, fmt.Errorf("'%s' is not a row count", sqlparser.String(expr)) + } + n, err := strconv.Atoi(string(val.Val)) + if err != nil || n < 0 { + return 0, fmt.Errorf("'%s' is not a row count", string(val.Val)) + } + return n, nil +} + +// runStagedSelect stages the omnisdk relation, runs the outer statement over +// it and releases the table. There is no result cache, so the table is +// released as soon as the outer result has been read. +func runStagedSelect(ctx queryContext, staged stagedSelect) internaldto.ExecutorOutput { + manager, err := getStagingManager(ctx) + if err != nil { + return internaldto.NewErroneousExecutorOutput(fmt.Errorf("omnisdk staging unavailable: %w", err)) + } + bg := context.Background() + id, err := manager.Begin(bg) + if err != nil { + return internaldto.NewErroneousExecutorOutput(err) + } + output, runErr := stageAndQuery(ctx, manager, id, staged) + if releaseErr := manager.Release(bg, id); releaseErr != nil && runErr == nil { + runErr = releaseErr + } + if runErr != nil { + return internaldto.NewErroneousExecutorOutput(runErr) + } + return output +} + +func stageAndQuery( + ctx queryContext, + manager omnistaging.Manager, + id omnistaging.QueryID, + staged stagedSelect, +) (internaldto.ExecutorOutput, error) { + bg := context.Background() + rows, _, err := openDocQuery(ctx, staged.getSource()) + if err != nil { + return nil, err + } + defer rows.Close() + source := newOmniBatchSource(rows, staged.getColumns(), previewCfg.getBatchSize()) + // Column types come from the first batch, read before the table exists. + first, err := source.Next(bg) + if err != nil && !errors.Is(err, io.EOF) { + return nil, err + } + source.pending = first + if _, err = manager.Stage(bg, id, stagedColumns(staged.getColumns(), first), source); err != nil { + return nil, err + } + db, err := ctx.GetSQLEngine().GetDB() + if err != nil { + return nil, err + } + result, err := db.QueryContext(bg, staged.renderOuter(manager.TableName(id))) + if err != nil { + return nil, fmt.Errorf("staged query failed: %w", err) + } + defer result.Close() + return readStagedResult(ctx, result, staged.getOutputNames()) +} + +// stagedColumns types each column by its first non-null value. Numbers are +// staged as numeric so they order and aggregate as numbers; everything else +// is text, as the streamed path renders it. +func stagedColumns(names []string, first [][]any) []omnistaging.Column { + columns := make([]omnistaging.Column, 0, len(names)) + for i, name := range names { + relationalType := stagedTextType + for _, row := range first { + if row[i] == nil { + continue + } + if isNumber(row[i]) { + relationalType = stagedNumericType + } + break + } + columns = append(columns, omnistaging.NewColumn(name, relationalType)) + } + return columns +} + +func isNumber(value any) bool { + switch value.(type) { + case int64, float64: + return true + default: + return false + } +} + +// stagedValue keeps numbers as numbers and renders everything else as the +// streamed path would. +func stagedValue(value any) any { + switch typed := value.(type) { + case int: + return int64(typed) + case int32: + return int64(typed) + case int64, float64: + return typed + case float32: + return float64(typed) + default: + return textValue(value) + } +} + +// omniBatchSource reads the omnisdk cursor a batch at a time, only when the +// stager asks for one. +type omniBatchSource struct { + rows omnisdk.Rows + columns []string + size int + pending [][]any +} + +func newOmniBatchSource(rows omnisdk.Rows, columns []string, size int) *omniBatchSource { + if size < 1 { + size = defaultBatchSize + } + return &omniBatchSource{rows: rows, columns: columns, size: size} +} + +func (s *omniBatchSource) Next(context.Context) ([][]any, error) { + if s.pending != nil { + batch := s.pending + s.pending = nil + return batch, nil + } + batch := make([][]any, 0, s.size) + for len(batch) < s.size && s.rows.Next() { + row := s.rows.Row() + values := make([]any, 0, len(s.columns)) + for _, name := range s.columns { + values = append(values, stagedValue(row[name])) + } + batch = append(batch, values) + } + if len(batch) > 0 { + return batch, nil + } + if err := s.rows.Err(); err != nil { + return nil, err + } + return nil, io.EOF +} + +// readStagedResult reads the outer result under the reported output names; +// its columns are the select list, in order. +func readStagedResult( + ctx queryContext, + result *sql.Rows, + columnOrder []string, +) (internaldto.ExecutorOutput, error) { + var err error + rowMap := make(map[string]map[string]interface{}) + for i := 0; result.Next(); i++ { + values := make([]any, len(columnOrder)) + pointers := make([]any, len(columnOrder)) + for j := range values { + pointers[j] = &values[j] + } + if err = result.Scan(pointers...); err != nil { + return nil, err + } + row := make(map[string]interface{}, len(columnOrder)) + for j, name := range columnOrder { + row[name] = outputValue(values[j]) + } + rowMap[fmt.Sprintf("%012d", i)] = row + } + if err = result.Err(); err != nil { + return nil, err + } + return prepare(ctx, columnOrder, rowMap, util.DefaultRowSort), nil +} + +func outputValue(value any) any { + if raw, isBytes := value.([]byte); isBytes { + return string(raw) + } + return textValue(value) +} diff --git a/internal/stackql/intrinsic/staged_test.go b/internal/stackql/intrinsic/staged_test.go new file mode 100644 index 000000000..fe81d28ef --- /dev/null +++ b/internal/stackql/intrinsic/staged_test.go @@ -0,0 +1,261 @@ +package intrinsic //nolint:testpackage // tests unexported staging + +import ( + "bytes" + "encoding/json" + "errors" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "testing" + + "github.com/stackql/any-sdk/pkg/dto" + "github.com/stackql/any-sdk/public/sqlengine" + "github.com/stackql/stackql/internal/stackql/internal_data_transfer/internaldto" + "github.com/stackql/stackql/internal/stackql/output" + "github.com/stackql/stackql/internal/stackql/typing" + "github.com/stackql/stackql/pkg/astformat" + + "github.com/stackql/stackql-parser/go/vt/sqlparser" +) + +type stagingTestCtx struct { + engine sqlengine.SQLEngine + runtime dto.RuntimeCtx + typCfg typing.Config +} + +func (c *stagingTestCtx) GetCurrentProvider() string { return "" } +func (c *stagingTestCtx) SetCurrentProvider(string) {} +func (c *stagingTestCtx) GetTypingConfig() typing.Config { return c.typCfg } +func (c *stagingTestCtx) GetRuntimeContext() dto.RuntimeCtx { return c.runtime } +func (c *stagingTestCtx) GetSQLEngine() sqlengine.SQLEngine { return c.engine } +func (c *stagingTestCtx) GetASTFormatter() sqlparser.NodeFormatter { + return astformat.SQLiteSelectExprsFormatter +} +func (c *stagingTestCtx) GetAuthContext(string) (*dto.AuthCtx, error) { + return nil, errors.New("no auth configured") +} + +func TestNeedsStaging(t *testing.T) { + for sql, want := range map[string]bool{ + "select login from stackql_unstable_github.orgs.members": false, + "select login from stackql_unstable_github.orgs.members limit 3": false, + "select login from stackql_unstable_github.orgs.members order by login": true, + "select distinct login from stackql_unstable_github.orgs.members": true, + "select count(*) from stackql_unstable_github.orgs.members": true, + "select login from stackql_unstable_github.orgs.members limit 3 offset 1": true, + "select type from stackql_unstable_github.orgs.members group by type": true, + "select upper(login) from stackql_unstable_github.orgs.members": false, + "select login from stackql_unstable_github.orgs.members having login = 'x'": true, + } { + if got := needsStaging(parseSelect(t, sql)); got != want { + t.Errorf("%s: got %v want %v", sql, got, want) + } + } +} + +func TestPlanStagedSelectOuterStatement(t *testing.T) { + withUnstable(t, true) + sel := parseSelect(t, "select k.name as ring, upper(c.name), count(*) as n "+ + "from stackql_unstable_google.cloudkms.key_rings k "+ + "inner join stackql_unstable_google.cloudkms.crypto_keys c on c.keyRingsId = k.name "+ + "where k.projectsId = 'p' group by k.name, c.name having count(c.name) > 1 "+ + "order by n desc, ring limit 5 offset 2") + cases := []struct { + formatter sqlparser.NodeFormatter + table string + want string + }{ + { + formatter: astformat.SQLiteSelectExprsFormatter, + table: `"__iql__.queries.7"`, + want: `select _omni_0 as ring, upper(_omni_1) as _omni_out_1, count(*) as n ` + + `from "__iql__.queries.7" group by _omni_0, _omni_1 having count(_omni_1) > 1 ` + + `order by n desc, ring asc LIMIT 5 OFFSET 2`, + }, + { + formatter: astformat.PostgresSelectExprsFormatter, + table: `"stackql_queries"."7"`, + want: `select "_omni_0" as "ring", upper("_omni_1") as "_omni_out_1", count(*) as "n" ` + + `from "stackql_queries"."7" group by "_omni_0", "_omni_1" having count("_omni_1") > 1 ` + + `order by "n" desc, "ring" asc LIMIT 5 OFFSET 2`, + }, + } + for _, tc := range cases { + staged, err := planStagedSelect(sel, "", tc.formatter) + if err != nil { + t.Fatal(err) + } + if got := staged.renderOuter(tc.table); got != tc.want { + t.Errorf("outer:\n got %s\nwant %s", got, tc.want) + } + if got, want := staged.getOutputNames(), []string{"ring", `upper("c".name)`, "n"}; !equalStrings(got, want) { + t.Errorf("output names: got %v want %v", got, want) + } + if got, want := staged.getColumns(), []string{"_omni_0", "_omni_1"}; !equalStrings(got, want) { + t.Errorf("staged columns: got %v want %v", got, want) + } + if staged.getSource().getLimit() != 0 { + t.Errorf("limit must not be pushed below the staging boundary") + } + } + if _, err := planStagedSelect(parseSelect(t, + "select * from stackql_unstable_github.orgs.members order by login"), "", + astformat.SQLiteSelectExprsFormatter); err == nil || + err.Error() != "'*' cannot be staged for stackql_unstable_* relations; name the columns" { + t.Fatalf("star refusal: got %v", err) + } +} + +func equalStrings(a, b []string) bool { + if len(a) != len(b) { + return false + } + for i := range a { + if a[i] != b[i] { + return false + } + } + return true +} + +func newStagingTestCtx(t *testing.T, endpoint string) *stagingTestCtx { + t.Helper() + dir := t.TempDir() + registry, err := filepath.Abs(filepath.Join("testdata", "registry")) + if err != nil { + t.Fatal(err) + } + if err = os.Symlink(registry, filepath.Join(dir, "src")); err != nil { + t.Fatal(err) + } + sqlBackendRaw := `{"dsn":"file:` + filepath.ToSlash(filepath.Join(dir, "stackql.db")) + `"}` + sqlCfg, err := dto.GetSQLBackendCfg(sqlBackendRaw) + if err != nil { + t.Fatal(err) + } + engine, err := sqlengine.NewSQLEngine(sqlCfg, nil) + if err != nil { + t.Fatal(err) + } + db, err := engine.GetDB() + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = db.Close() }) + typCfg, err := typing.NewTypingConfig(sqlCfg.GetSQLDialect()) + if err != nil { + t.Fatal(err) + } + endpointJSON, err := json.Marshal(endpoint) + if err != nil { + t.Fatal(err) + } + previous := previewCfg + previewCfg = newBackendInput(previewCfgDTO{ + Unstable: true, Staging: true, Endpoint: endpointJSON, InsecureSkipTLSVerify: true, + }) + t.Cleanup(func() { previewCfg = previous }) + return &stagingTestCtx{ + engine: engine, + runtime: dto.RuntimeCtx{ + RegistryRaw: `{"url":"file:` + filepath.ToSlash(filepath.Join(dir, "registry")) + `"}`, + SQLBackendCfgRaw: sqlBackendRaw, + }, + typCfg: typCfg, + } +} + +func renderCSV(t *testing.T, sql string, out internaldto.ExecutorOutput) string { + t.Helper() + if err := out.GetError(); err != nil { + t.Fatalf("%s: %v", sql, err) + } + var buf, errBuf bytes.Buffer + writer, err := output.GetOutputWriter(&buf, &errBuf, internaldto.OutputContext{ + RuntimeContext: dto.RuntimeCtx{OutputFormat: "csv", Delimiter: ","}, + Result: out.GetSQLResult(), + }) + if err != nil { + t.Fatal(err) + } + if err = writer.Write(out.GetSQLResult()); err != nil { + t.Fatal(err) + } + return buf.String() +} + +// The staging manager is created once per process, so every end-to-end case +// shares one backend and runs in this one test. +func TestStagedSelectEndToEnd(t *testing.T) { + members := []map[string]any{ + {"login": "a", "id": 1, "type": "User"}, + {"login": "b", "id": 10, "type": "User"}, + {"login": "c", "id": 2, "type": "Bot"}, + {"login": "d", "id": 3, "type": "User"}, + } + srv := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/orgs/dummyorg/members" { + http.NotFound(w, r) + return + } + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(members) + })) + defer srv.Close() + ctx := newStagingTestCtx(t, srv.URL) + t.Setenv("FIXTURE_TOKEN", "test-token") + + const from = " from stackql_unstable_fixture.orgs.members where org = 'dummyorg'" + for _, tc := range []struct { + sql string + want string + }{ + { + // Numeric ordering: as text, 10 would sort before 2. + sql: "select login, id" + from + " order by id desc limit 2 offset 1", + want: "login,id\nd,3\nc,2\n", + }, + { + sql: "select type, count(*) as n" + from + " group by type having count(*) > 1 order by n desc", + want: "type,n\nUser,3\n", + }, + { + sql: "select count(*)" + from, + want: "count(*)\n4\n", + }, + { + sql: "select distinct type" + from + " order by type", + want: "type\nBot\nUser\n", + }, + { + sql: "select upper(login), id" + from + " order by login desc", + // Named as the streamed path names an unaliased expression. + want: "upper(`login`),id\nD,3\nC,2\nB,10\nA,1\n", + }, + } { + fn, claimed := selectFunc(ctx, parseSelect(t, tc.sql), "") + if !claimed { + t.Fatalf("%s: not claimed", tc.sql) + } + if got := renderCSV(t, tc.sql, fn()); got != tc.want { + t.Errorf("%s:\n got %q\nwant %q", tc.sql, got, tc.want) + } + } + + db, err := ctx.engine.GetDB() + if err != nil { + t.Fatal(err) + } + var remaining int + if err = db.QueryRow( + `SELECT count(*) FROM sqlite_master WHERE type = 'table' AND name LIKE '\_\_iql\_\_.queries.%' ESCAPE '\'`, + ).Scan(&remaining); err != nil { + t.Fatal(err) + } + if remaining != 0 { + t.Fatalf("staging tables left behind: %d", remaining) + } +} diff --git a/internal/stackql/intrinsic/translate.go b/internal/stackql/intrinsic/translate.go index 908e1b807..a082ce9ea 100644 --- a/internal/stackql/intrinsic/translate.go +++ b/internal/stackql/intrinsic/translate.go @@ -105,20 +105,10 @@ func translateSelect(node *sqlparser.Select, currentProvider string) (docQuery, if err != nil { return nil, err } - if len(node.From) != 1 { - return nil, fmt.Errorf("a comma-separated FROM cannot be applied to %s relations; use JOIN ... ON", - UnstablePrefix+"*") - } - t := &translator{currentProvider: currentProvider} - if err = t.from(node.From[0], query.Base, nil); err != nil { + t, where, err := translateSource(node, currentProvider) + if err != nil { return nil, err } - var where []query.Predicate - if node.Where != nil { - if where, err = conjuncts(node.Where.Expr); err != nil { - return nil, err - } - } outputs, names, err := selectOutputs(node.SelectExprs) if err != nil { return nil, err @@ -130,6 +120,27 @@ func translateSelect(node *sqlparser.Select, currentProvider string) (docQuery, return newDocQuery(q, names, limit, t.bundles), nil } +// translateSource translates a SELECT's FROM and WHERE: the joins omnisdk runs +// and the conjuncts it applies to them. +func translateSource(node *sqlparser.Select, currentProvider string) (*translator, []query.Predicate, error) { + if len(node.From) != 1 { + return nil, nil, fmt.Errorf("a comma-separated FROM cannot be applied to %s relations; use JOIN ... ON", + UnstablePrefix+"*") + } + t := &translator{currentProvider: currentProvider} + if err := t.from(node.From[0], query.Base, nil); err != nil { + return nil, nil, err + } + if node.Where == nil { + return t, nil, nil + } + where, err := conjuncts(node.Where.Expr) + if err != nil { + return nil, nil, err + } + return t, where, nil +} + // unsupportedDocClauses names what omnisdk returns unapplied: its row stream is // unordered, ungrouped and not de-duplicated. LIMIT is pushed down instead. func unsupportedDocClauses(node *sqlparser.Select) []string { diff --git a/internal/stackql/intrinsic/translate_test.go b/internal/stackql/intrinsic/translate_test.go index b6af24dbe..4556ea8a1 100644 --- a/internal/stackql/intrinsic/translate_test.go +++ b/internal/stackql/intrinsic/translate_test.go @@ -297,3 +297,28 @@ func TestMutationTablesIgnoresImplicitDual(t *testing.T) { t.Fatalf("got isDoc %v, err %v", isDoc, err) } } + +func TestDocProviderUnderOmniAll(t *testing.T) { + previous := previewCfg + t.Cleanup(func() { previewCfg = previous }) + for _, tc := range []struct { + cfg previewCfgDTO + name string + wantBundle string + wantDoc bool + }{ + {cfg: previewCfgDTO{}, name: "aws", wantDoc: false}, + {cfg: previewCfgDTO{Unstable: true}, name: "aws", wantDoc: false}, + {cfg: previewCfgDTO{Unstable: true}, name: "stackql_unstable_aws", wantBundle: "aws", wantDoc: true}, + {cfg: previewCfgDTO{Omni: "all"}, name: "aws", wantBundle: "aws", wantDoc: true}, + {cfg: previewCfgDTO{Omni: "all"}, name: "stackql_unstable_aws", wantBundle: "aws", wantDoc: true}, + {cfg: previewCfgDTO{Omni: "all"}, name: ProviderName, wantDoc: false}, + {cfg: previewCfgDTO{Omni: "all"}, name: "", wantDoc: false}, + } { + previewCfg = newBackendInput(tc.cfg) + bundle, isDoc := docProvider(tc.name) + if bundle != tc.wantBundle || isDoc != tc.wantDoc { + t.Errorf("%+v %q: got %q %v, want %q %v", tc.cfg, tc.name, bundle, isDoc, tc.wantBundle, tc.wantDoc) + } + } +} diff --git a/internal/stackql/omnistaging/dialect.go b/internal/stackql/omnistaging/dialect.go new file mode 100644 index 000000000..be32141c3 --- /dev/null +++ b/internal/stackql/omnistaging/dialect.go @@ -0,0 +1,186 @@ +package omnistaging + +import ( + "fmt" + "strings" + + "github.com/stackql/any-sdk/pkg/constants" +) + +const ( + sqliteQueryTablePrefix = "__iql__.queries." + sqliteQueryIDSeqTable = "__iql__.query_id_seq" + postgresQueryIDSequence = "query_id_seq" + // Per-statement bind parameter limits. + sqliteMaxBindParameters = 32766 + postgresMaxBindParameters = 65535 +) + +// Dialect renders the backend-specific SQL for query staging. +type Dialect interface { + SetupStatements() []string + NextQueryIDStatement() string + TableName(id QueryID) string + CreateTableStatement(id QueryID, columns []Column) (string, error) + InsertStatement(id QueryID, columns []Column, rowCount int) (string, error) + DropTableStatement(id QueryID) string + MaxBindParameters() int +} + +// NewDialect selects a dialect from an any-sdk SQL dialect name. +func NewDialect(sqlDialect string, cfg Config) (Dialect, error) { + switch sqlDialect { + case constants.SQLDialectSQLite3: + return newSQLiteDialect(), nil + case constants.SQLDialectPostgres: + return newPostgresDialect(cfg), nil + default: + return nil, fmt.Errorf("omnistaging: unsupported sql dialect %q", sqlDialect) + } +} + +func quoteIdentifier(name string) string { + return `"` + strings.ReplaceAll(name, `"`, `""`) + `"` +} + +func quoteLiteral(s string) string { + return `'` + strings.ReplaceAll(s, `'`, `''`) + `'` +} + +func createTableStatement(tableName string, columns []Column) (string, error) { + if len(columns) == 0 { + return "", fmt.Errorf("omnistaging: cannot create %s without columns", tableName) + } + defs := make([]string, 0, len(columns)) + for _, col := range columns { + if col.Name() == "" || col.RelationalType() == "" { + return "", fmt.Errorf("omnistaging: column requires name and relational type") + } + defs = append(defs, quoteIdentifier(col.Name())+" "+col.RelationalType()) + } + return fmt.Sprintf("CREATE TABLE IF NOT EXISTS %s (%s)", tableName, strings.Join(defs, ", ")), nil +} + +func insertStatement( + tableName string, + columns []Column, + rowCount int, + maxBindParameters int, + placeholder func(int) string, +) (string, error) { + if len(columns) == 0 || rowCount < 1 { + return "", fmt.Errorf("omnistaging: insert into %s requires columns and rows", tableName) + } + if len(columns)*rowCount > maxBindParameters { + return "", fmt.Errorf("omnistaging: insert into %s exceeds %d bind parameters", tableName, maxBindParameters) + } + names := make([]string, 0, len(columns)) + for _, col := range columns { + names = append(names, quoteIdentifier(col.Name())) + } + var b strings.Builder + fmt.Fprintf(&b, "INSERT INTO %s (%s) VALUES ", tableName, strings.Join(names, ", ")) + ordinal := 1 + for r := 0; r < rowCount; r++ { + if r > 0 { + b.WriteString(", ") + } + b.WriteByte('(') + for c := range columns { + if c > 0 { + b.WriteString(", ") + } + b.WriteString(placeholder(ordinal)) + ordinal++ + } + b.WriteByte(')') + } + return b.String(), nil +} + +type sqliteDialect struct{} + +func newSQLiteDialect() Dialect { + return &sqliteDialect{} +} + +func (d *sqliteDialect) SetupStatements() []string { + return []string{ + fmt.Sprintf( + "CREATE TABLE IF NOT EXISTS %s (id INTEGER PRIMARY KEY AUTOINCREMENT)", + quoteIdentifier(sqliteQueryIDSeqTable), + ), + } +} + +// AUTOINCREMENT never reuses a rowid, so IDs stay unique across processes +// sharing one database file. +func (d *sqliteDialect) NextQueryIDStatement() string { + return fmt.Sprintf("INSERT INTO %s DEFAULT VALUES RETURNING id", quoteIdentifier(sqliteQueryIDSeqTable)) +} + +func (d *sqliteDialect) TableName(id QueryID) string { + return quoteIdentifier(sqliteQueryTablePrefix + id.String()) +} + +func (d *sqliteDialect) CreateTableStatement(id QueryID, columns []Column) (string, error) { + return createTableStatement(d.TableName(id), columns) +} + +func (d *sqliteDialect) InsertStatement(id QueryID, columns []Column, rowCount int) (string, error) { + return insertStatement(d.TableName(id), columns, rowCount, d.MaxBindParameters(), func(int) string { return "?" }) +} + +func (d *sqliteDialect) DropTableStatement(id QueryID) string { + return "DROP TABLE IF EXISTS " + d.TableName(id) +} + +func (d *sqliteDialect) MaxBindParameters() int { + return sqliteMaxBindParameters +} + +type postgresDialect struct { + querySchema string +} + +func newPostgresDialect(cfg Config) Dialect { + return &postgresDialect{querySchema: cfg.QuerySchema()} +} + +func (d *postgresDialect) sequenceName() string { + return quoteIdentifier(d.querySchema) + "." + quoteIdentifier(postgresQueryIDSequence) +} + +func (d *postgresDialect) SetupStatements() []string { + return []string{ + "CREATE SCHEMA IF NOT EXISTS " + quoteIdentifier(d.querySchema), + "CREATE SEQUENCE IF NOT EXISTS " + d.sequenceName(), + } +} + +func (d *postgresDialect) NextQueryIDStatement() string { + return fmt.Sprintf("SELECT nextval(%s)", quoteLiteral(d.sequenceName())) +} + +func (d *postgresDialect) TableName(id QueryID) string { + return quoteIdentifier(d.querySchema) + "." + quoteIdentifier(id.String()) +} + +func (d *postgresDialect) CreateTableStatement(id QueryID, columns []Column) (string, error) { + return createTableStatement(d.TableName(id), columns) +} + +func (d *postgresDialect) InsertStatement(id QueryID, columns []Column, rowCount int) (string, error) { + return insertStatement( + d.TableName(id), columns, rowCount, d.MaxBindParameters(), + func(ordinal int) string { return fmt.Sprintf("$%d", ordinal) }, + ) +} + +func (d *postgresDialect) DropTableStatement(id QueryID) string { + return "DROP TABLE IF EXISTS " + d.TableName(id) +} + +func (d *postgresDialect) MaxBindParameters() int { + return postgresMaxBindParameters +} diff --git a/internal/stackql/omnistaging/manager.go b/internal/stackql/omnistaging/manager.go new file mode 100644 index 000000000..3c7559fc9 --- /dev/null +++ b/internal/stackql/omnistaging/manager.go @@ -0,0 +1,208 @@ +package omnistaging + +import ( + "context" + "errors" + "fmt" + "io" +) + +// Manager owns query identities and their staged tables. +// +// Release is idempotent: it drops every table owned by the query and its +// descendants, then forgets them. On error the registry is retained so the +// release can be retried. +type Manager interface { + Begin(ctx context.Context) (QueryID, error) + BeginChild(ctx context.Context, parent QueryID) (QueryID, error) + Parent(id QueryID) (QueryID, bool) + Children(id QueryID) []QueryID + TableName(id QueryID) string + // Stage creates the query table and writes every batch from source, + // returning the number of rows written. + Stage(ctx context.Context, id QueryID, columns []Column, source BatchSource) (int64, error) + // Owned returns the IDs with staged tables in the subtree rooted at id. + Owned(id QueryID) []QueryID + Release(ctx context.Context, id QueryID) error +} + +// writer renders and executes staging DDL and batch inserts. +type writer interface { + Create(ctx context.Context, id QueryID, columns []Column) error + Insert(ctx context.Context, id QueryID, columns []Column, rows [][]any) error + Drop(ctx context.Context, id QueryID) error +} + +type sqlWriter struct { + db Executor + dialect Dialect +} + +func newSQLWriter(db Executor, dialect Dialect) writer { + return &sqlWriter{db: db, dialect: dialect} +} + +func (w *sqlWriter) Create(ctx context.Context, id QueryID, columns []Column) error { + stmt, err := w.dialect.CreateTableStatement(id, columns) + if err != nil { + return err + } + if _, err = w.db.ExecContext(ctx, stmt); err != nil { + return fmt.Errorf("omnistaging: cannot create %s: %w", w.dialect.TableName(id), err) + } + return nil +} + +// Insert writes rows in chunks bounded by the dialect bind parameter limit. +func (w *sqlWriter) Insert(ctx context.Context, id QueryID, columns []Column, rows [][]any) error { + width := len(columns) + if width == 0 { + return fmt.Errorf("omnistaging: insert into %s requires columns", w.dialect.TableName(id)) + } + chunkSize := w.dialect.MaxBindParameters() / width + for start := 0; start < len(rows); start += chunkSize { + end := min(start+chunkSize, len(rows)) + chunk := rows[start:end] + args := make([]any, 0, len(chunk)*width) + for _, row := range chunk { + if len(row) != width { + return fmt.Errorf( + "omnistaging: row width %d does not match %d columns of %s", + len(row), width, w.dialect.TableName(id), + ) + } + args = append(args, row...) + } + stmt, err := w.dialect.InsertStatement(id, columns, len(chunk)) + if err != nil { + return err + } + if _, err = w.db.ExecContext(ctx, stmt, args...); err != nil { + return fmt.Errorf("omnistaging: cannot insert into %s: %w", w.dialect.TableName(id), err) + } + } + return nil +} + +func (w *sqlWriter) Drop(ctx context.Context, id QueryID) error { + if _, err := w.db.ExecContext(ctx, w.dialect.DropTableStatement(id)); err != nil { + return fmt.Errorf("omnistaging: cannot drop %s: %w", w.dialect.TableName(id), err) + } + return nil +} + +type standardManager struct { + dialect Dialect + counter counter + registry registry + writer writer +} + +// NewManager runs the dialect setup statements and returns a manager whose +// query IDs are allocated by the backend. +func NewManager(ctx context.Context, db Executor, dialect Dialect) (Manager, error) { + for _, stmt := range dialect.SetupStatements() { + if _, err := db.ExecContext(ctx, stmt); err != nil { + return nil, fmt.Errorf("omnistaging: setup failed: %w", err) + } + } + return &standardManager{ + dialect: dialect, + counter: newSQLCounter(db, dialect), + registry: newRegistry(), + writer: newSQLWriter(db, dialect), + }, nil +} + +func (m *standardManager) Begin(ctx context.Context) (QueryID, error) { + return m.begin(ctx, nil) +} + +func (m *standardManager) BeginChild(ctx context.Context, parent QueryID) (QueryID, error) { + if parent == nil { + return nil, fmt.Errorf("omnistaging: child query requires a parent") + } + return m.begin(ctx, parent) +} + +func (m *standardManager) begin(ctx context.Context, parent QueryID) (QueryID, error) { + id, err := m.counter.Next(ctx) + if err != nil { + return nil, err + } + if err = m.registry.Add(id, parent); err != nil { + return nil, err + } + return id, nil +} + +func (m *standardManager) Parent(id QueryID) (QueryID, bool) { + return m.registry.Parent(id) +} + +func (m *standardManager) Children(id QueryID) []QueryID { + return m.registry.Children(id) +} + +func (m *standardManager) TableName(id QueryID) string { + return m.dialect.TableName(id) +} + +func (m *standardManager) Stage( + ctx context.Context, + id QueryID, + columns []Column, + source BatchSource, +) (int64, error) { + if !m.registry.Contains(id) { + return 0, fmt.Errorf("omnistaging: query %s not registered", id) + } + // Ownership is recorded before DDL so a partial failure is still released. + if err := m.registry.MarkStaged(id); err != nil { + return 0, err + } + if err := m.writer.Create(ctx, id, columns); err != nil { + return 0, err + } + var written int64 + for { + if err := ctx.Err(); err != nil { + return written, err + } + rows, err := source.Next(ctx) + if errors.Is(err, io.EOF) { + return written, nil + } + if err != nil { + return written, fmt.Errorf("omnistaging: batch source failed for query %s: %w", id, err) + } + if err = m.writer.Insert(ctx, id, columns, rows); err != nil { + return written, err + } + written += int64(len(rows)) + } +} + +func (m *standardManager) Owned(id QueryID) []QueryID { + var owned []QueryID + for _, member := range m.registry.Subtree(id) { + if m.registry.IsStaged(member) { + owned = append(owned, member) + } + } + return owned +} + +func (m *standardManager) Release(ctx context.Context, id QueryID) error { + subtree := m.registry.Subtree(id) + for _, member := range subtree { + if !m.registry.IsStaged(member) { + continue + } + if err := m.writer.Drop(ctx, member); err != nil { + return err + } + } + m.registry.Remove(subtree) + return nil +} diff --git a/internal/stackql/omnistaging/omnistaging.go b/internal/stackql/omnistaging/omnistaging.go new file mode 100644 index 000000000..48b8bb387 --- /dev/null +++ b/internal/stackql/omnistaging/omnistaging.go @@ -0,0 +1,118 @@ +// Package omnistaging owns RDBMS staging for omnisdk-backed queries. +// +// Each top-level query receives one collision-free query ID. When remaining +// SQL work requires materialization, the final Omni relation for that query +// is staged in exactly one query-owned table. Child query IDs exist only for +// genuine pre-analysis splits and are released with their parent. +// +// See docs/technical/omnisdk_staging.md. +package omnistaging + +import ( + "context" + "database/sql" + "fmt" + "strconv" + + "gopkg.in/yaml.v2" +) + +// DefaultQuerySchema is the PostgreSQL schema holding query tables when +// none is configured. +const DefaultQuerySchema = "stackql_queries" + +// Executor is the subset of database access required for staging. +// *sql.DB and *sql.Tx satisfy it. +type Executor interface { + ExecContext(ctx context.Context, query string, args ...any) (sql.Result, error) + QueryRowContext(ctx context.Context, query string, args ...any) *sql.Row +} + +// QueryID identifies one query and its staging namespace. +type QueryID interface { + Value() int64 + String() string +} + +// Column describes one column of a staged relation. +type Column interface { + Name() string + RelationalType() string +} + +// BatchSource yields batches of rows in cursor order. Next returns io.EOF +// once the source is exhausted. It is pulled only as fast as batches are +// written, so the producer is back-pressured by staging. +type BatchSource interface { + Next(ctx context.Context) ([][]any, error) +} + +// Config carries staging configuration. +type Config interface { + QuerySchema() string +} + +type queryID struct { + value int64 +} + +func newQueryID(value int64) QueryID { + return &queryID{value: value} +} + +func (q *queryID) Value() int64 { + return q.value +} + +func (q *queryID) String() string { + return strconv.FormatInt(q.value, 10) +} + +type column struct { + name string + relationalType string +} + +// NewColumn constructs a staged relation column. +func NewColumn(name, relationalType string) Column { + return &column{name: name, relationalType: relationalType} +} + +func (c *column) Name() string { + return c.name +} + +func (c *column) RelationalType() string { + return c.relationalType +} + +type config struct { + querySchema string +} + +// NewConfig constructs staging configuration; an empty querySchema selects +// DefaultQuerySchema. +func NewConfig(querySchema string) Config { + if querySchema == "" { + querySchema = DefaultQuerySchema + } + return &config{querySchema: querySchema} +} + +// NewConfigFromSQLBackend reads staging configuration from the raw +// --sqlBackend string (JSON or YAML), key schemata.querySchema. +func NewConfigFromSQLBackend(raw string) (Config, error) { + var parsed struct { + Schemata struct { + QuerySchema string `yaml:"querySchema"` + } `yaml:"schemata"` + } + if err := yaml.Unmarshal([]byte(raw), &parsed); err != nil { + return nil, fmt.Errorf("omnistaging: cannot parse sql backend config: %w", err) + } + return NewConfig(parsed.Schemata.QuerySchema), nil +} + +func (c *config) QuerySchema() string { + return c.querySchema +} diff --git a/internal/stackql/omnistaging/omnistaging_test.go b/internal/stackql/omnistaging/omnistaging_test.go new file mode 100644 index 000000000..2d1dafee9 --- /dev/null +++ b/internal/stackql/omnistaging/omnistaging_test.go @@ -0,0 +1,381 @@ +package omnistaging_test + +import ( + "context" + "database/sql" + "errors" + "io" + "path/filepath" + "reflect" + "strconv" + "testing" + + "github.com/stackql/any-sdk/pkg/constants" + "github.com/stackql/any-sdk/pkg/dto" + "github.com/stackql/any-sdk/public/sqlengine" + + . "github.com/stackql/stackql/internal/stackql/omnistaging" //nolint:revive // test reads as package +) + +type fixedQueryID int64 + +func (q fixedQueryID) Value() int64 { return int64(q) } + +func (q fixedQueryID) String() string { return strconv.FormatInt(int64(q), 10) } + +func sqliteDialect(t *testing.T) Dialect { + t.Helper() + d, err := NewDialect(constants.SQLDialectSQLite3, NewConfig("")) + if err != nil { + t.Fatal(err) + } + return d +} + +func mustBegin(t *testing.T, begin func(context.Context) (QueryID, error)) QueryID { + t.Helper() + id, err := begin(context.Background()) + if err != nil { + t.Fatal(err) + } + return id +} + +func TestConfig(t *testing.T) { + cases := []struct { + raw string + want string + }{ + {raw: `{"dsn":"file:./stackql.db"}`, want: DefaultQuerySchema}, + {raw: ``, want: DefaultQuerySchema}, + {raw: `{"dbEngine":"postgres_tcp","schemata":{"tableSchema":"t","querySchema":"q"}}`, want: "q"}, + } + for _, tc := range cases { + cfg, err := NewConfigFromSQLBackend(tc.raw) + if err != nil { + t.Fatalf("unexpected error for %q: %v", tc.raw, err) + } + if got := cfg.QuerySchema(); got != tc.want { + t.Fatalf("query schema for %q: got %q want %q", tc.raw, got, tc.want) + } + } +} + +func TestDialectStatements(t *testing.T) { + columns := []Column{NewColumn("vpc_id", "text"), NewColumn(`we"ird`, "integer")} + id := fixedQueryID(42) + sqlite, err := NewDialect(constants.SQLDialectSQLite3, NewConfig("")) + if err != nil { + t.Fatal(err) + } + postgres, err := NewDialect(constants.SQLDialectPostgres, NewConfig("q")) + if err != nil { + t.Fatal(err) + } + mustString := func(s string, err error) string { + t.Helper() + if err != nil { + t.Fatal(err) + } + return s + } + got := []string{ + sqlite.TableName(id), + mustString(sqlite.CreateTableStatement(id, columns)), + mustString(sqlite.InsertStatement(id, columns, 2)), + sqlite.DropTableStatement(id), + sqlite.NextQueryIDStatement(), + postgres.TableName(id), + mustString(postgres.CreateTableStatement(id, columns)), + mustString(postgres.InsertStatement(id, columns, 2)), + postgres.DropTableStatement(id), + postgres.NextQueryIDStatement(), + } + want := []string{ + `"__iql__.queries.42"`, + `CREATE TABLE IF NOT EXISTS "__iql__.queries.42" ("vpc_id" text, "we""ird" integer)`, + `INSERT INTO "__iql__.queries.42" ("vpc_id", "we""ird") VALUES (?, ?), (?, ?)`, + `DROP TABLE IF EXISTS "__iql__.queries.42"`, + `INSERT INTO "__iql__.query_id_seq" DEFAULT VALUES RETURNING id`, + `"q"."42"`, + `CREATE TABLE IF NOT EXISTS "q"."42" ("vpc_id" text, "we""ird" integer)`, + `INSERT INTO "q"."42" ("vpc_id", "we""ird") VALUES ($1, $2), ($3, $4)`, + `DROP TABLE IF EXISTS "q"."42"`, + `SELECT nextval('"q"."query_id_seq"')`, + } + if !reflect.DeepEqual(got, want) { + t.Fatalf("statements:\n got %q\nwant %q", got, want) + } + if !reflect.DeepEqual(postgres.SetupStatements(), []string{ + `CREATE SCHEMA IF NOT EXISTS "q"`, + `CREATE SEQUENCE IF NOT EXISTS "q"."query_id_seq"`, + }) { + t.Fatalf("postgres setup: %q", postgres.SetupStatements()) + } + if _, err = sqlite.CreateTableStatement(id, nil); err == nil { + t.Fatal("expected error creating table without columns") + } + if _, err = NewDialect(constants.SQLDialectSnowflake, NewConfig("")); err == nil { + t.Fatal("expected unsupported dialect error") + } +} + +func openSQLite(t *testing.T, path string) *sql.DB { + t.Helper() + eng, err := sqlengine.NewSQLEngine(dto.SQLBackendCfg{ + DBEngine: constants.DBEngineSQLite3Embedded, + SQLSystem: constants.SQLDialectSQLite3, + DSN: "file:" + path, + }, nil) + if err != nil { + t.Fatalf("cannot open sqlite engine: %v", err) + } + db, err := eng.GetDB() + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = db.Close() }) + return db +} + +func newSQLiteManager(t *testing.T, db *sql.DB, dialect Dialect) Manager { + t.Helper() + m, err := NewManager(context.Background(), db, dialect) + if err != nil { + t.Fatal(err) + } + return m +} + +func stagingTables(t *testing.T, db *sql.DB) []string { + t.Helper() + rows, err := db.Query( + `SELECT name FROM sqlite_master WHERE type = 'table' AND name LIKE '\_\_iql\_\_.queries.%' ESCAPE '\' ORDER BY name`, + ) + if err != nil { + t.Fatal(err) + } + defer rows.Close() + var names []string + for rows.Next() { + var name string + if err = rows.Scan(&name); err != nil { + t.Fatal(err) + } + names = append(names, name) + } + if err = rows.Err(); err != nil { + t.Fatal(err) + } + return names +} + +func stagedPairs(t *testing.T, db *sql.DB, table string) [][2]sql.NullString { + t.Helper() + rows, err := db.Query("SELECT vpc_id, subnet_id FROM " + table + " ORDER BY rowid") + if err != nil { + t.Fatal(err) + } + defer rows.Close() + var got [][2]sql.NullString + for rows.Next() { + var r [2]sql.NullString + if err = rows.Scan(&r[0], &r[1]); err != nil { + t.Fatal(err) + } + got = append(got, r) + } + if err = rows.Err(); err != nil { + t.Fatal(err) + } + return got +} + +// sliceSource yields fixed batches and records how many rows had been +// written to the staging table each time a batch was requested. +type sliceSource struct { + batches [][][]any + next int + observe func() int64 + observed []int64 +} + +func (s *sliceSource) Next(context.Context) ([][]any, error) { + if s.observe != nil { + s.observed = append(s.observed, s.observe()) + } + if s.next == len(s.batches) { + return nil, io.EOF + } + b := s.batches[s.next] + s.next++ + return b, nil +} + +func countRows(t *testing.T, db *sql.DB, table string) int64 { + t.Helper() + var n int64 + if err := db.QueryRow("SELECT count(*) FROM " + table).Scan(&n); err != nil { + t.Fatal(err) + } + return n +} + +func TestQueryIDsCollisionFreeAcrossManagers(t *testing.T) { + path := filepath.Join(t.TempDir(), "stackql.db") + ctx := context.Background() + first := newSQLiteManager(t, openSQLite(t, path), sqliteDialect(t)) + second := newSQLiteManager(t, openSQLite(t, path), sqliteDialect(t)) + var got []int64 + for _, m := range []Manager{first, second, first, second} { + id, err := m.Begin(ctx) + if err != nil { + t.Fatal(err) + } + got = append(got, id.Value()) + } + if want := []int64{1, 2, 3, 4}; !reflect.DeepEqual(got, want) { + t.Fatalf("query ids: got %v want %v", got, want) + } +} + +func TestStageAndRelease(t *testing.T) { + ctx := context.Background() + db := openSQLite(t, filepath.Join(t.TempDir(), "stackql.db")) + m := newSQLiteManager(t, db, sqliteDialect(t)) + + streamOnly := mustBegin(t, m.Begin) + parent := mustBegin(t, m.Begin) + child := mustBegin(t, func(ctx context.Context) (QueryID, error) { return m.BeginChild(ctx, parent) }) + if p, ok := m.Parent(child); !ok || p.Value() != parent.Value() { + t.Fatalf("parent of %s: got %v", child, p) + } + + columns := []Column{NewColumn("vpc_id", "text"), NewColumn("subnet_id", "text")} + source := &sliceSource{ + batches: [][][]any{ + {{"vpc-1", "subnet-a"}, {"vpc-1", "subnet-a"}}, + {{"vpc-2", nil}}, + {}, + {{"vpc-3", "subnet-c"}}, + }, + observe: func() int64 { return countRows(t, db, m.TableName(parent)) }, + } + written, err := m.Stage(ctx, parent, columns, source) + if err != nil { + t.Fatal(err) + } + if written != 4 { + t.Fatalf("written: got %d want 4", written) + } + // Each batch is requested only after the previous one is inserted. + if want := []int64{0, 2, 3, 3, 4}; !reflect.DeepEqual(source.observed, want) { + t.Fatalf("rows staged at each pull: got %v want %v", source.observed, want) + } + if _, err = m.Stage(ctx, child, columns, &sliceSource{batches: [][][]any{{{"vpc-9", "subnet-z"}}}}); err != nil { + t.Fatal(err) + } + + got := stagedPairs(t, db, m.TableName(parent)) + ns := func(s string) sql.NullString { return sql.NullString{String: s, Valid: true} } + want := [][2]sql.NullString{ + {ns("vpc-1"), ns("subnet-a")}, + {ns("vpc-1"), ns("subnet-a")}, + {ns("vpc-2"), {}}, + {ns("vpc-3"), ns("subnet-c")}, + } + if !reflect.DeepEqual(got, want) { + t.Fatalf("staged rows: got %v want %v", got, want) + } + + if tables := stagingTables(t, db); !reflect.DeepEqual(tables, []string{ + "__iql__.queries." + parent.String(), + "__iql__.queries." + child.String(), + }) { + t.Fatalf("tables before release: %v", tables) + } + if owned := m.Owned(parent); len(owned) != 2 || owned[0].Value() != child.Value() || owned[1].Value() != parent.Value() { + t.Fatalf("owned by parent: %v", owned) + } + if owned := m.Owned(streamOnly); len(owned) != 0 { + t.Fatalf("stream-only query owns tables: %v", owned) + } + + if err = m.Release(ctx, parent); err != nil { + t.Fatal(err) + } + if err = m.Release(ctx, parent); err != nil { + t.Fatalf("second release must be a no-op: %v", err) + } + if err = m.Release(ctx, streamOnly); err != nil { + t.Fatal(err) + } + if tables := stagingTables(t, db); len(tables) != 0 { + t.Fatalf("tables after release: %v", tables) + } + if _, ok := m.Parent(child); ok { + t.Fatal("child must be released with parent") + } + if _, err = m.Stage(ctx, parent, columns, &sliceSource{}); err == nil { + t.Fatal("expected error staging a released query") + } +} + +type smallBindDialect struct { + Dialect +} + +func (d *smallBindDialect) MaxBindParameters() int { + return 4 +} + +func TestInsertChunksByBindParameterLimit(t *testing.T) { + ctx := context.Background() + db := openSQLite(t, filepath.Join(t.TempDir(), "stackql.db")) + m := newSQLiteManager(t, db, &smallBindDialect{Dialect: sqliteDialect(t)}) + id := mustBegin(t, m.Begin) + columns := []Column{NewColumn("a", "integer"), NewColumn("b", "integer")} + batch := [][]any{{1, 2}, {3, 4}, {5, 6}, {7, 8}, {9, 10}} + if _, err := m.Stage(ctx, id, columns, &sliceSource{batches: [][][]any{batch}}); err != nil { + t.Fatal(err) + } + var sum int64 + if err := db.QueryRow("SELECT sum(a) + sum(b) FROM " + m.TableName(id)).Scan(&sum); err != nil { + t.Fatal(err) + } + if countRows(t, db, m.TableName(id)) != 5 || sum != 55 { + t.Fatalf("chunked insert: rows %d sum %d", countRows(t, db, m.TableName(id)), sum) + } + if _, err := m.Stage(ctx, id, columns, &sliceSource{batches: [][][]any{{{1}}}}); err == nil { + t.Fatal("expected row width error") + } +} + +type failingSource struct { + cancel context.CancelFunc +} + +func (s *failingSource) Next(context.Context) ([][]any, error) { + s.cancel() + return [][]any{{"x"}}, nil +} + +func TestCancelledStageIsReleased(t *testing.T) { + db := openSQLite(t, filepath.Join(t.TempDir(), "stackql.db")) + m := newSQLiteManager(t, db, sqliteDialect(t)) + id := mustBegin(t, m.Begin) + ctx, cancel := context.WithCancel(context.Background()) + _, err := m.Stage(ctx, id, []Column{NewColumn("a", "text")}, &failingSource{cancel: cancel}) + if !errors.Is(err, context.Canceled) { + t.Fatalf("expected cancellation, got %v", err) + } + if len(m.Owned(id)) != 1 { + t.Fatal("cancelled stage must still be owned for release") + } + if err = m.Release(context.Background(), id); err != nil { + t.Fatal(err) + } + if tables := stagingTables(t, db); len(tables) != 0 { + t.Fatalf("tables after release: %v", tables) + } +} diff --git a/internal/stackql/omnistaging/registry.go b/internal/stackql/omnistaging/registry.go new file mode 100644 index 000000000..b2e05f171 --- /dev/null +++ b/internal/stackql/omnistaging/registry.go @@ -0,0 +1,169 @@ +package omnistaging + +import ( + "context" + "fmt" + "sync" +) + +// counter allocates collision-free query IDs. +type counter interface { + Next(ctx context.Context) (QueryID, error) +} + +// sqlCounter delegates allocation to the backend, so IDs are unique across +// every process sharing that backend. +type sqlCounter struct { + db Executor + dialect Dialect +} + +func newSQLCounter(db Executor, dialect Dialect) counter { + return &sqlCounter{db: db, dialect: dialect} +} + +func (c *sqlCounter) Next(ctx context.Context) (QueryID, error) { + var value int64 + if err := c.db.QueryRowContext(ctx, c.dialect.NextQueryIDStatement()).Scan(&value); err != nil { + return nil, fmt.Errorf("omnistaging: cannot allocate query id: %w", err) + } + return newQueryID(value), nil +} + +// registry is the concurrency-safe store of live queries, their parent/child +// associations and whether each owns a staged table. +type registry interface { + Add(id QueryID, parent QueryID) error + Contains(id QueryID) bool + Parent(id QueryID) (QueryID, bool) + Children(id QueryID) []QueryID + MarkStaged(id QueryID) error + // Subtree returns id and all descendants, descendants first. + Subtree(id QueryID) []QueryID + IsStaged(id QueryID) bool + Remove(ids []QueryID) +} + +type registryNode struct { + id QueryID + parent QueryID + children []QueryID + staged bool +} + +type standardRegistry struct { + mu sync.Mutex + nodes map[int64]*registryNode +} + +func newRegistry() registry { + return &standardRegistry{nodes: make(map[int64]*registryNode)} +} + +func (r *standardRegistry) Add(id QueryID, parent QueryID) error { + r.mu.Lock() + defer r.mu.Unlock() + if _, ok := r.nodes[id.Value()]; ok { + return fmt.Errorf("omnistaging: query %s already registered", id) + } + if parent != nil { + parentNode, ok := r.nodes[parent.Value()] + if !ok { + return fmt.Errorf("omnistaging: parent query %s not registered", parent) + } + parentNode.children = append(parentNode.children, id) + } + r.nodes[id.Value()] = ®istryNode{id: id, parent: parent} + return nil +} + +func (r *standardRegistry) Contains(id QueryID) bool { + r.mu.Lock() + defer r.mu.Unlock() + _, ok := r.nodes[id.Value()] + return ok +} + +func (r *standardRegistry) Parent(id QueryID) (QueryID, bool) { + r.mu.Lock() + defer r.mu.Unlock() + node, ok := r.nodes[id.Value()] + if !ok || node.parent == nil { + return nil, false + } + return node.parent, true +} + +func (r *standardRegistry) Children(id QueryID) []QueryID { + r.mu.Lock() + defer r.mu.Unlock() + node, ok := r.nodes[id.Value()] + if !ok { + return nil + } + return append([]QueryID(nil), node.children...) +} + +func (r *standardRegistry) MarkStaged(id QueryID) error { + r.mu.Lock() + defer r.mu.Unlock() + node, ok := r.nodes[id.Value()] + if !ok { + return fmt.Errorf("omnistaging: query %s not registered", id) + } + node.staged = true + return nil +} + +func (r *standardRegistry) IsStaged(id QueryID) bool { + r.mu.Lock() + defer r.mu.Unlock() + node, ok := r.nodes[id.Value()] + return ok && node.staged +} + +func (r *standardRegistry) Subtree(id QueryID) []QueryID { + r.mu.Lock() + defer r.mu.Unlock() + var out []QueryID + var walk func(QueryID) + walk = func(current QueryID) { + node, ok := r.nodes[current.Value()] + if !ok { + return + } + for _, child := range node.children { + walk(child) + } + out = append(out, node.id) + } + walk(id) + return out +} + +func (r *standardRegistry) Remove(ids []QueryID) { + r.mu.Lock() + defer r.mu.Unlock() + for _, id := range ids { + node, ok := r.nodes[id.Value()] + if !ok { + continue + } + if node.parent != nil { + if parentNode, parentOK := r.nodes[node.parent.Value()]; parentOK { + parentNode.children = removeQueryID(parentNode.children, id) + } + } + delete(r.nodes, id.Value()) + } +} + +func removeQueryID(ids []QueryID, target QueryID) []QueryID { + out := ids[:0] + for _, id := range ids { + if id.Value() != target.Value() { + out = append(out, id) + } + } + return out +} diff --git a/internal/stackql/writer/somefile b/internal/stackql/writer/somefile deleted file mode 100644 index e69de29bb..000000000 diff --git a/internal/stackql/writer/writer_test.go b/internal/stackql/writer/writer_test.go index 5c6b13825..54dae0cbb 100644 --- a/internal/stackql/writer/writer_test.go +++ b/internal/stackql/writer/writer_test.go @@ -3,59 +3,42 @@ package writer //nolint:testpackage // this violates another rule: var-naming: d import ( "io" "os" + "path/filepath" "testing" "github.com/stackql/any-sdk/pkg/dto" "github.com/stackql/stackql/internal/stackql/presentation" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) -type NopWriter struct{} - -func (nw NopWriter) Write(p []byte) (int, error) { - return len(p), nil -} - -func TestGetOutputWriter(t *testing.T) { - nopWriter := NopWriter{} - type args struct { - filename string - } - tests := []struct { - name string - args args - want io.Writer - }{ - { - "stdout", - args{"stdout"}, - os.Stdout, - }, - { - "stderr", - args{"stderr"}, - os.Stderr, - }, - { - "file", - args{"somefile"}, - nopWriter, - }, - } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - got, _ := GetOutputWriter(tt.args.filename) - - if tt.name == "file" { - assert.Implements(t, (*io.Writer)(nil), got) - defer os.Remove(tt.args.filename) - } else { - assert.Equal(t, got, tt.want) - } +func TestGetOutputWriterStdStreams(t *testing.T) { + for filename, want := range map[string]io.Writer{ + StdOutStr: os.Stdout, + StdErrStr: os.Stderr, + } { + t.Run(filename, func(t *testing.T) { + got, err := GetOutputWriter(filename) + require.NoError(t, err) + assert.Equal(t, want, got) }) } } +func TestGetOutputWriterFile(t *testing.T) { + path := filepath.Join(t.TempDir(), "out.txt") + got, err := GetOutputWriter(path) + require.NoError(t, err) + file, isFile := got.(*os.File) + require.True(t, isFile, "a file name must yield an *os.File") + _, err = io.WriteString(file, "some output") + require.NoError(t, err) + require.NoError(t, file.Close()) + content, err := os.ReadFile(path) + require.NoError(t, err) + assert.Equal(t, "some output", string(content)) +} + func TestGetDecoratedOutputWriter(t *testing.T) { type args struct { filename string diff --git a/test/robot/functional/stackql_mocked_from_cmd_line.robot b/test/robot/functional/stackql_mocked_from_cmd_line.robot index c23d8225f..ec64d074c 100644 --- a/test/robot/functional/stackql_mocked_from_cmd_line.robot +++ b/test/robot/functional/stackql_mocked_from_cmd_line.robot @@ -11245,6 +11245,83 @@ Unstable Github Org Members Filtered By In List Jsonl Row Set Matches Expectatio ... stdout=${CURDIR}${/}tmp${/}Unstable-Github-Org-Members-Filtered-By-In-List.tmp ... stderr=${CURDIR}${/}tmp${/}Unstable-Github-Org-Members-Filtered-By-In-List-stderr.tmp +Unstable Github Org Members Staged Order By With Offset Exact Match + [Documentation] With staging opted into, the ORDER BY, LIMIT and OFFSET + ... omnisdk leaves unapplied are evaluated by the SQL backend + ... over the staged result, so the row order is asserted. + [Teardown] Remove Preview Mock Environment + ${preview} = Catenate SEPARATOR= + ... {"endpoint":"https://${LOCAL_HOST_ALIAS}:${MOCKSERVER_PORT_GITHUB}", + ... "insecureSkipTLSVerify":true,"unstable":true,"staging":true} + ${query} = Catenate SEPARATOR=${SPACE} + ... select login, type from stackql_unstable_github.orgs.members + ... where org = 'dummyorg' order by login desc limit 3 offset 1; + Should StackQL Exec Inline Equal + ... ${STACKQL_EXE} + ... ${OKTA_SECRET_STR} + ... ${GITHUB_SECRET_STR} + ... ${K8S_SECRET_STR} + ... ${REGISTRY_NO_VERIFY_CFG_STR} + ... ${AUTH_CFG_STR} + ... ${SQL_BACKEND_CFG_STR_CANONICAL} + ... ${query} + ... login,type\nsome-jimbo-8,User\nsome-jimbo-7,User\nsome-jimbo-6,User + ... \-o\=csv + ... --preview\=${preview} + ... stdout=${CURDIR}${/}tmp${/}Unstable-Github-Org-Members-Staged-Order-By-With-Offset.tmp + ... stderr=${CURDIR}${/}tmp${/}Unstable-Github-Org-Members-Staged-Order-By-With-Offset-stderr.tmp + +Unstable Github Org Members Staged Group By Exact Match + [Documentation] With staging opted into, grouping and aggregation run in + ... the SQL backend over the staged result. + [Teardown] Remove Preview Mock Environment + ${preview} = Catenate SEPARATOR= + ... {"endpoint":"https://${LOCAL_HOST_ALIAS}:${MOCKSERVER_PORT_GITHUB}", + ... "insecureSkipTLSVerify":true,"unstable":true,"staging":true} + ${query} = Catenate SEPARATOR=${SPACE} + ... select type, count(*) as member_count from stackql_unstable_github.orgs.members + ... where org = 'dummyorg' group by type; + Should StackQL Exec Inline Equal + ... ${STACKQL_EXE} + ... ${OKTA_SECRET_STR} + ... ${GITHUB_SECRET_STR} + ... ${K8S_SECRET_STR} + ... ${REGISTRY_NO_VERIFY_CFG_STR} + ... ${AUTH_CFG_STR} + ... ${SQL_BACKEND_CFG_STR_CANONICAL} + ... ${query} + ... type,member_count\nUser,10 + ... \-o\=csv + ... --preview\=${preview} + ... stdout=${CURDIR}${/}tmp${/}Unstable-Github-Org-Members-Staged-Group-By.tmp + ... stderr=${CURDIR}${/}tmp${/}Unstable-Github-Org-Members-Staged-Group-By-stderr.tmp + +Omni All Github Org Members Staged Order By With Offset Exact Match + [Documentation] With omni set to all, a canonical provider name is routed + ... to omnisdk rather than any-sdk, and staging applies the + ... ORDER BY, LIMIT and OFFSET. + [Teardown] Remove Preview Mock Environment + ${preview} = Catenate SEPARATOR= + ... {"endpoint":"https://${LOCAL_HOST_ALIAS}:${MOCKSERVER_PORT_GITHUB}", + ... "insecureSkipTLSVerify":true,"omni":"all","staging":true} + ${query} = Catenate SEPARATOR=${SPACE} + ... select login, type from github.orgs.members + ... where org = 'dummyorg' order by login desc limit 3 offset 1; + Should StackQL Exec Inline Equal + ... ${STACKQL_EXE} + ... ${OKTA_SECRET_STR} + ... ${GITHUB_SECRET_STR} + ... ${K8S_SECRET_STR} + ... ${REGISTRY_NO_VERIFY_CFG_STR} + ... ${AUTH_CFG_STR} + ... ${SQL_BACKEND_CFG_STR_CANONICAL} + ... ${query} + ... login,type\nsome-jimbo-8,User\nsome-jimbo-7,User\nsome-jimbo-6,User + ... \-o\=csv + ... --preview\=${preview} + ... stdout=${CURDIR}${/}tmp${/}Omni-All-Github-Org-Members-Staged-Order-By-With-Offset.tmp + ... stderr=${CURDIR}${/}tmp${/}Omni-All-Github-Org-Members-Staged-Order-By-With-Offset-stderr.tmp + Unstable Github Org Update Reports Despatch [Documentation] A document-driven UPDATE without RETURNING: omnisdk sends ... the effect and stackql reports it. The mock refuses any