nl2sql-engine 0.1.2__py3-none-any.whl → 0.2.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (184) hide show
  1. nl2sql/__init__.py +28 -4
  2. nl2sql/adapters/mssql/adapter.py +2 -2
  3. nl2sql/adapters/postgres/adapter.py +5 -3
  4. nl2sql/adapters/sqlalchemy_base/__init__.py +0 -4
  5. nl2sql/adapters/sqlalchemy_base/adapter.py +11 -6
  6. nl2sql/adapters/sqlalchemy_base/models.py +4 -17
  7. nl2sql/adapters/sqlite/adapter.py +48 -1
  8. nl2sql/aggregation/aggregator.py +3 -6
  9. nl2sql/aggregation/columns.py +110 -0
  10. nl2sql/aggregation/engines/polars_duckdb.py +96 -63
  11. nl2sql/api/benchmark_api.py +101 -61
  12. nl2sql/api/query_api.py +142 -9
  13. nl2sql/auth/rbac.py +28 -15
  14. nl2sql/cli/commands/benchmark.py +298 -16
  15. nl2sql/cli/commands/cache.py +32 -0
  16. nl2sql/cli/commands/demo.py +533 -0
  17. nl2sql/cli/commands/doctor.py +87 -2
  18. nl2sql/cli/commands/feedback.py +145 -0
  19. nl2sql/cli/commands/indexing.py +4 -125
  20. nl2sql/cli/commands/info.py +8 -1
  21. nl2sql/cli/commands/run.py +24 -7
  22. nl2sql/cli/commands/setup.py +77 -55
  23. nl2sql/cli/commands/trace.py +170 -0
  24. nl2sql/cli/common/indexing.py +121 -0
  25. nl2sql/cli/console.py +1 -1
  26. nl2sql/cli/demo/chinook.py +57 -0
  27. nl2sql/cli/demo/datasets.py +110 -0
  28. nl2sql/cli/demo/defaults.py +2 -117
  29. nl2sql/cli/demo/llm_config.py +98 -0
  30. nl2sql/cli/demo/manager.py +113 -139
  31. nl2sql/cli/demo/playground/__init__.py +1 -0
  32. nl2sql/cli/demo/playground/app.py +622 -0
  33. nl2sql/cli/demo/playground/assets/clips/ask-dark.jpg +0 -0
  34. nl2sql/cli/demo/playground/assets/clips/ask-dark.webm +0 -0
  35. nl2sql/cli/demo/playground/assets/clips/ask.jpg +0 -0
  36. nl2sql/cli/demo/playground/assets/clips/ask.webm +0 -0
  37. nl2sql/cli/demo/playground/assets/clips/pipeline-dark.jpg +0 -0
  38. nl2sql/cli/demo/playground/assets/clips/pipeline-dark.webm +0 -0
  39. nl2sql/cli/demo/playground/assets/clips/pipeline.jpg +0 -0
  40. nl2sql/cli/demo/playground/assets/clips/pipeline.webm +0 -0
  41. nl2sql/cli/demo/playground/assets/clips/retrieval-dark.jpg +0 -0
  42. nl2sql/cli/demo/playground/assets/clips/retrieval-dark.webm +0 -0
  43. nl2sql/cli/demo/playground/assets/clips/retrieval.jpg +0 -0
  44. nl2sql/cli/demo/playground/assets/clips/retrieval.webm +0 -0
  45. nl2sql/cli/demo/playground/assets/social-card.png +0 -0
  46. nl2sql/cli/demo/playground/hosted.py +443 -0
  47. nl2sql/cli/demo/playground/index_panel.py +133 -0
  48. nl2sql/cli/demo/playground/preview.py +101 -0
  49. nl2sql/cli/demo/playground/settings.py +399 -0
  50. nl2sql/cli/demo/playground/static/index.html +68 -0
  51. nl2sql/cli/demo/recordings/chinook.json +3817 -0
  52. nl2sql/cli/demo/stamp.py +66 -0
  53. nl2sql/cli/demo/support.py +45 -0
  54. nl2sql/cli/demo/webanalytics.py +41 -0
  55. nl2sql/cli/generators/env/templates.py +18 -1
  56. nl2sql/cli/generators/llm/generator.py +13 -2
  57. nl2sql/cli/main.py +275 -24
  58. nl2sql/cli/reporting.py +152 -332
  59. nl2sql/common/env_hint.py +32 -0
  60. nl2sql/common/errors.py +28 -15
  61. nl2sql/common/metrics.py +10 -5
  62. nl2sql/common/settings.py +106 -14
  63. nl2sql/configs/llm.py +23 -3
  64. nl2sql/configs/manager.py +19 -8
  65. nl2sql/context.py +23 -2
  66. nl2sql/datasets/README.md +108 -0
  67. nl2sql/datasets/__init__.py +21 -0
  68. nl2sql/datasets/chinook.sqlite +0 -0
  69. nl2sql/datasets/support.sqlite +0 -0
  70. nl2sql/datasets/webanalytics.sqlite +0 -0
  71. nl2sql/datasources/models.py +21 -3
  72. nl2sql/datasources/registry.py +31 -12
  73. nl2sql/evaluation/benchmark_runner.py +93 -291
  74. nl2sql/evaluation/datasets/chinook_gold.yaml +1079 -0
  75. nl2sql/evaluation/datasets/chinook_gold_plans.yaml +745 -0
  76. nl2sql/evaluation/evaluator.py +295 -94
  77. nl2sql/evaluation/faithfulness.py +134 -0
  78. nl2sql/evaluation/gold.py +148 -0
  79. nl2sql/evaluation/presets/__init__.py +79 -0
  80. nl2sql/evaluation/presets/claude-planner.yaml +20 -0
  81. nl2sql/evaluation/presets/gpt-5.4-mini-helpers.yaml +19 -0
  82. nl2sql/evaluation/presets/gpt-5.4.yaml +10 -0
  83. nl2sql/evaluation/prices.py +53 -0
  84. nl2sql/evaluation/records.py +613 -0
  85. nl2sql/evaluation/retrieval_recall.py +174 -0
  86. nl2sql/evaluation/tier1.py +140 -0
  87. nl2sql/evaluation/tier2.py +519 -0
  88. nl2sql/evaluation/types.py +10 -4
  89. nl2sql/execution/artifacts/store.py +5 -4
  90. nl2sql/{pipeline/nodes/global_planner/schemas.py → execution/dag.py} +4 -24
  91. nl2sql/execution/executor/sql_executor.py +1 -2
  92. nl2sql/feedback/__init__.py +16 -0
  93. nl2sql/feedback/drafts.py +65 -0
  94. nl2sql/feedback/record.py +90 -0
  95. nl2sql/feedback/stats.py +83 -0
  96. nl2sql/feedback/store.py +112 -0
  97. nl2sql/indexing/embeddings.py +12 -3
  98. nl2sql/indexing/enrichment_service.py +2 -3
  99. nl2sql/indexing/health.py +207 -0
  100. nl2sql/indexing/orchestrator.py +30 -8
  101. nl2sql/indexing/rebuild.py +144 -0
  102. nl2sql/indexing/retrieval_trace.py +144 -0
  103. nl2sql/indexing/vector_store.py +364 -127
  104. nl2sql/llm/failures.py +185 -0
  105. nl2sql/llm/providers.py +183 -0
  106. nl2sql/llm/registry.py +169 -38
  107. nl2sql/llm/replay.py +290 -0
  108. nl2sql/llm/request_key.py +182 -0
  109. nl2sql/llm/wires/__init__.py +33 -0
  110. nl2sql/llm/wires/anthropic.py +133 -0
  111. nl2sql/llm/wires/anthropic_client.py +63 -0
  112. nl2sql/llm/wires/base.py +87 -0
  113. nl2sql/llm/wires/openai.py +119 -0
  114. nl2sql/pipeline/graph.py +6 -8
  115. nl2sql/pipeline/graph_utils.py +85 -28
  116. nl2sql/pipeline/nodes/__init__.py +9 -24
  117. nl2sql/pipeline/nodes/aggregator/node.py +2 -7
  118. nl2sql/pipeline/nodes/aggregator/schemas.py +1 -20
  119. nl2sql/pipeline/nodes/answer_synthesizer/node.py +4 -4
  120. nl2sql/pipeline/nodes/ast_planner/__init__.py +8 -3
  121. nl2sql/pipeline/nodes/ast_planner/functions.py +68 -0
  122. nl2sql/pipeline/nodes/ast_planner/node.py +77 -10
  123. nl2sql/pipeline/nodes/ast_planner/prompts.py +84 -16
  124. nl2sql/pipeline/nodes/ast_planner/schemas.py +16 -13
  125. nl2sql/pipeline/nodes/datasource_resolver/node.py +167 -69
  126. nl2sql/pipeline/nodes/datasource_resolver/prompts.py +26 -0
  127. nl2sql/pipeline/nodes/datasource_resolver/schemas.py +10 -0
  128. nl2sql/pipeline/nodes/decomposer/dag.py +121 -0
  129. nl2sql/pipeline/nodes/decomposer/node.py +171 -18
  130. nl2sql/pipeline/nodes/decomposer/prompts.py +46 -32
  131. nl2sql/pipeline/nodes/decomposer/schemas.py +59 -30
  132. nl2sql/pipeline/nodes/executor/node.py +3 -13
  133. nl2sql/pipeline/nodes/generator/node.py +286 -63
  134. nl2sql/pipeline/nodes/refiner/node.py +8 -17
  135. nl2sql/pipeline/nodes/refiner/prompts.py +21 -12
  136. nl2sql/pipeline/nodes/schema_retriever/node.py +114 -9
  137. nl2sql/pipeline/nodes/schema_retriever/schema.py +43 -1
  138. nl2sql/pipeline/nodes/validator/node.py +347 -168
  139. nl2sql/pipeline/nodes/validator/schemas.py +14 -0
  140. nl2sql/pipeline/plan_cache.py +112 -0
  141. nl2sql/pipeline/routes.py +29 -21
  142. nl2sql/pipeline/runtime.py +165 -35
  143. nl2sql/pipeline/state.py +3 -3
  144. nl2sql/pipeline/steps.py +130 -0
  145. nl2sql/pipeline/subgraphs/schemas.py +5 -1
  146. nl2sql/pipeline/subgraphs/sql_agent.py +31 -21
  147. nl2sql/pipeline/timing.py +52 -0
  148. nl2sql/public_api.py +111 -5
  149. nl2sql/schema/in_memory_store.py +16 -0
  150. nl2sql/schema/protocol.py +16 -1
  151. nl2sql/schema/sqlite_store.py +105 -4
  152. nl2sql/schema/view.py +56 -0
  153. nl2sql/secrets/factory.py +4 -3
  154. nl2sql/secrets/manager.py +15 -2
  155. nl2sql/services/callbacks/monitor.py +10 -11
  156. nl2sql/services/callbacks/node_handlers.py +1 -2
  157. nl2sql/services/callbacks/node_metrics.py +0 -3
  158. nl2sql/services/callbacks/token_handler.py +249 -42
  159. nl2sql/testing/fake_llm.py +242 -0
  160. nl2sql/tracing/__init__.py +7 -0
  161. nl2sql/tracing/document.py +212 -0
  162. nl2sql/tracing/recorder.py +376 -0
  163. nl2sql/tracing/replay.py +318 -0
  164. nl2sql/tracing/trace.py +173 -0
  165. nl2sql_engine-0.2.0.dist-info/METADATA +226 -0
  166. nl2sql_engine-0.2.0.dist-info/RECORD +267 -0
  167. nl2sql_engine-0.2.0.dist-info/licenses/LICENSE +21 -0
  168. nl2sql_engine-0.2.0.dist-info/licenses/THIRD_PARTY_NOTICES.md +249 -0
  169. nl2sql/api/result_api.py +0 -24
  170. nl2sql/cli/common/prompts.py +0 -53
  171. nl2sql/cli/demo/data.py +0 -87
  172. nl2sql/cli/demo/factory.py +0 -289
  173. nl2sql/cli/demo/schemas.py +0 -348
  174. nl2sql/cli/demo/writers/docker.py +0 -194
  175. nl2sql/cli/demo/writers/sqlite.py +0 -88
  176. nl2sql/pipeline/nodes/aggregator/prompts.py +0 -20
  177. nl2sql/pipeline/nodes/global_planner/__init__.py +0 -4
  178. nl2sql/pipeline/nodes/global_planner/node.py +0 -186
  179. nl2sql_engine-0.1.2.dist-info/METADATA +0 -295
  180. nl2sql_engine-0.1.2.dist-info/RECORD +0 -193
  181. /nl2sql/{cli/demo/writers → testing}/__init__.py +0 -0
  182. {nl2sql_engine-0.1.2.dist-info → nl2sql_engine-0.2.0.dist-info}/WHEEL +0 -0
  183. {nl2sql_engine-0.1.2.dist-info → nl2sql_engine-0.2.0.dist-info}/entry_points.txt +0 -0
  184. {nl2sql_engine-0.1.2.dist-info → nl2sql_engine-0.2.0.dist-info}/top_level.txt +0 -0
nl2sql/__init__.py CHANGED
@@ -1,7 +1,15 @@
1
1
  # nl2sql package
2
2
 
3
+ import importlib
4
+
3
5
  from .public_api import NL2SQL, QueryResult
4
6
 
7
+ # The shapes a client renders and the hooks an application entry point needs,
8
+ # so REST and SDK clients never import engine submodules.
9
+ from .api.query_api import SubQueryResult, RowSample
10
+ from .services.callbacks.token_handler import QuestionUsage
11
+ from .common.logger import configure_logging
12
+
5
13
  # Also expose individual API modules for more granular access
6
14
  from .api.query_api import QueryAPI
7
15
  from .api.datasource_api import DatasourceAPI
@@ -9,14 +17,26 @@ from .api.llm_api import LLM_API
9
17
  from .api.indexing_api import IndexingAPI
10
18
  from .api.auth_api import AuthAPI
11
19
  from .api.settings_api import SettingsAPI
12
- from .api.result_api import ResultAPI
13
20
  from .api.policy_api import PolicyAPI
14
- from .api.benchmark_api import BenchmarkAPI
15
21
 
16
22
  # Also expose core models and enums
17
23
  from .common.errors import ErrorSeverity, ErrorCode, PipelineError
24
+ from .common.cancellation import CancellationToken
18
25
  from .auth.models import UserContext
19
- from .evaluation.types import BenchmarkConfig
26
+
27
+ # The benchmark lives in nl2sql.evaluation, which the runtime never needs, so
28
+ # its names are imported only when first asked for.
29
+ _LAZY = {
30
+ "BenchmarkAPI": "nl2sql.api.benchmark_api",
31
+ "BenchmarkConfig": "nl2sql.evaluation.types",
32
+ }
33
+
34
+
35
+ def __getattr__(name):
36
+ if name in _LAZY:
37
+ return getattr(importlib.import_module(_LAZY[name]), name)
38
+ raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
39
+
20
40
 
21
41
  __all__ = [
22
42
  "NL2SQL",
@@ -27,12 +47,16 @@ __all__ = [
27
47
  "IndexingAPI",
28
48
  "AuthAPI",
29
49
  "SettingsAPI",
30
- "ResultAPI",
31
50
  "PolicyAPI",
32
51
  "BenchmarkAPI",
33
52
  "ErrorSeverity",
34
53
  "ErrorCode",
35
54
  "PipelineError",
55
+ "CancellationToken",
36
56
  "UserContext",
57
+ "SubQueryResult",
58
+ "RowSample",
59
+ "QuestionUsage",
60
+ "configure_logging",
37
61
  "BenchmarkConfig",
38
62
  ]
@@ -1,6 +1,5 @@
1
1
  from typing import Any, List, Dict
2
2
  from sqlalchemy import create_engine, text, inspect
3
- from sqlalchemy.dialects import mssql
4
3
  from nl2sql.adapters.sqlalchemy_base import (
5
4
  CostEstimate,
6
5
  DryRunResult,
@@ -88,7 +87,8 @@ class MssqlAdapter(BaseSQLAlchemyAdapter):
88
87
 
89
88
  def get_dialect(self) -> str:
90
89
  """MSSQL uses T-SQL dialect."""
91
- return mssql.dialect.name
90
+ # A sqlglot dialect name (SQLAlchemy calls it "mssql").
91
+ return "tsql"
92
92
 
93
93
  def cost_estimate(self, sql: str) -> CostEstimate:
94
94
  import re
@@ -1,6 +1,5 @@
1
1
  from typing import Dict, Any
2
2
  from sqlalchemy import create_engine, inspect, text
3
- from sqlalchemy.dialects import postgresql
4
3
  from nl2sql.adapters.sqlalchemy_base import (
5
4
  DryRunResult,
6
5
  QueryPlan,
@@ -54,7 +53,9 @@ class PostgresAdapter(BaseSQLAlchemyAdapter):
54
53
  from urllib.parse import urlencode
55
54
  query_str = "?" + urlencode(options)
56
55
 
57
- return f"postgresql://{creds}{netloc}/{database}{query_str}"
56
+ # Name the driver: the `postgres` extra installs psycopg2, and a bare
57
+ # `postgresql://` means psycopg (v3) from SQLAlchemy 2.1 onwards.
58
+ return f"postgresql+psycopg2://{creds}{netloc}/{database}{query_str}"
58
59
 
59
60
  def connect(self) -> None:
60
61
  """Postgres-specific connection with Native Server-Side Timeout."""
@@ -107,7 +108,8 @@ class PostgresAdapter(BaseSQLAlchemyAdapter):
107
108
  return CostEstimate(estimated_cost=0.0, estimated_rows=0)
108
109
 
109
110
  def get_dialect(self) -> str:
110
- return postgresql.dialect.name
111
+ # A sqlglot dialect name (SQLAlchemy calls it "postgresql").
112
+ return "postgres"
111
113
 
112
114
 
113
115
  @property
@@ -1,17 +1,13 @@
1
1
  from .adapter import BaseSQLAlchemyAdapter
2
2
  from .models import (
3
- QueryResult,
4
3
  CostEstimate,
5
4
  DryRunResult,
6
5
  QueryPlan,
7
- AdapterError,
8
6
  )
9
7
 
10
8
  __all__ = [
11
9
  "BaseSQLAlchemyAdapter",
12
- "QueryResult",
13
10
  "CostEstimate",
14
11
  "DryRunResult",
15
12
  "QueryPlan",
16
- "AdapterError",
17
13
  ]
@@ -101,12 +101,7 @@ class BaseSQLAlchemyAdapter:
101
101
 
102
102
  def capabilities(self) -> set[DatasourceCapability]:
103
103
  """Default capability set for SQL adapters."""
104
- return {
105
- DatasourceCapability.SUPPORTS_SQL,
106
- DatasourceCapability.SUPPORTS_SCHEMA_INTROSPECTION,
107
- DatasourceCapability.SUPPORTS_DRY_RUN,
108
- DatasourceCapability.SUPPORTS_COST_ESTIMATE,
109
- }
104
+ return {DatasourceCapability.SUPPORTS_SQL}
110
105
 
111
106
  def execute_sql(self, sql: str) -> ResultFrame:
112
107
  """Executes a SQL query against the datasource.
@@ -458,6 +453,16 @@ class BaseSQLAlchemyAdapter:
458
453
  def get_dialect(self) -> str:
459
454
  raise NotImplementedError(f"Adapter {self.__class__.__name__} must implement get_dialect")
460
455
 
456
+ def render_sql(self, expression: Any) -> str:
457
+ """Renders the engine's finished sqlglot tree as this database's SQL.
458
+
459
+ The default is sqlglot's own rendering for ``get_dialect()``. Override
460
+ it to rewrite what sqlglot cannot express for the database, keeping the
461
+ result types in the SDK contract (a date part is an integer, a
462
+ truncated date a ``YYYY-MM-DD`` string).
463
+ """
464
+ return expression.sql(dialect=self.get_dialect())
465
+
461
466
  def cost_estimate(self, sql: str) -> CostEstimate:
462
467
  raise NotImplementedError(f"Adapter {self.__class__.__name__} must implement cost_estimate")
463
468
 
@@ -1,15 +1,9 @@
1
- from pydantic import BaseModel, Field, ConfigDict
2
- from typing import List, Optional, Any
1
+ from pydantic import BaseModel
2
+ from typing import Optional, Any
3
3
 
4
+ # Results and errors use the SDK's ResultFrame and ResultError
5
+ # (nl2sql_adapter_sdk.contracts); these are the SQL-only extras.
4
6
 
5
- class QueryResult(BaseModel):
6
- """Normalized results from a datasource execution."""
7
- columns: List[str]
8
- rows: List[List[Any]]
9
- row_count: int
10
- raw: Optional[Any] = None
11
- execution_time_ms: Optional[float] = None
12
- bytes_returned: Optional[int] = None
13
7
 
14
8
  class DryRunResult(BaseModel):
15
9
  """Result of a query validation/dry-run."""
@@ -27,10 +21,3 @@ class CostEstimate(BaseModel):
27
21
  estimated_cost: float
28
22
  estimated_rows: int
29
23
  estimated_time_ms: Optional[float] = None
30
-
31
- class AdapterError(BaseModel):
32
- """Standardized error envelope for adapter failures."""
33
- code: str
34
- message: str
35
- retriable: bool
36
- raw: Optional[Any] = None
@@ -10,6 +10,41 @@ from nl2sql.adapters.sqlalchemy_base import (
10
10
 
11
11
  from pydantic import BaseModel, Field
12
12
  from typing import Optional
13
+ import sqlglot
14
+ from sqlglot import exp
15
+
16
+ # SQLite has no EXTRACT and no date truncation. Each template reads the date
17
+ # with STRFTIME/DATE; ``__x__`` stands for the operand. The engine wraps a
18
+ # date part in CAST(... AS INT) and a truncation in a YYYY-MM-DD format, so
19
+ # the types match every other adapter.
20
+ _DATE_PARTS = {
21
+ "YEAR": "STRFTIME('%Y', __x__)",
22
+ "QUARTER": "(CAST(STRFTIME('%m', __x__) AS INTEGER) + 2) / 3",
23
+ "MONTH": "STRFTIME('%m', __x__)",
24
+ "DAY": "STRFTIME('%d', __x__)",
25
+ }
26
+ _DATE_TRUNCS = {
27
+ "YEAR": "DATE(__x__, 'start of year')",
28
+ "QUARTER": "DATE(__x__, 'start of month', '-' || ((CAST(STRFTIME('%m', __x__) AS INTEGER) - 1) % 3) || ' months')",
29
+ "MONTH": "DATE(__x__, 'start of month')",
30
+ "DAY": "DATE(__x__)",
31
+ }
32
+
33
+
34
+ def _from_template(template: str, operand: exp.Expression) -> exp.Expression:
35
+ tree = sqlglot.parse_one(template, read="sqlite")
36
+ return tree.transform(
37
+ lambda node: operand.copy() if isinstance(node, exp.Column) and node.name == "__x__" else node
38
+ )
39
+
40
+
41
+ def _sqlite_dates(node: exp.Expression) -> exp.Expression:
42
+ """Rewrites the engine's portable date nodes with SQLite's date functions."""
43
+ if isinstance(node, exp.Extract) and node.name.upper() in _DATE_PARTS:
44
+ return _from_template(_DATE_PARTS[node.name.upper()], node.expression)
45
+ if isinstance(node, exp.TimestampTrunc) and node.unit and node.unit.name.upper() in _DATE_TRUNCS:
46
+ return _from_template(_DATE_TRUNCS[node.unit.name.upper()], node.this)
47
+ return node
13
48
 
14
49
  class SqliteConnectionConfig(BaseModel):
15
50
  """Strict configuration schema for SQLite adapter."""
@@ -24,16 +59,24 @@ class SqliteAdapter(BaseSQLAlchemyAdapter):
24
59
  def construct_uri(self, args: Dict[str, Any]) -> str:
25
60
  """Constructs the SQLite connection URI.
26
61
 
62
+ With ``options.read_only`` the file is opened through SQLite's own URI
63
+ form with ``mode=ro``, so the driver refuses every write before any SQL
64
+ is parsed. That is the last line under the policy and validator checks,
65
+ and the one a public demo needs: the sample databases are ours to show
66
+ and nobody's to change.
67
+
27
68
  Args:
28
69
  args: The raw connection arguments dictionary.
29
70
 
30
71
  Returns:
31
72
  str: The fully constructed SQLAlchemy connection URI.
32
-
73
+
33
74
  Raises:
34
75
  ValidationError: If the configuration is invalid.
35
76
  """
36
77
  config = SqliteConnectionConfig(**args)
78
+ if config.options.get("read_only"):
79
+ return f"sqlite:///file:{config.database}?mode=ro&uri=true"
37
80
  return f"sqlite:///{config.database}"
38
81
 
39
82
  def connect(self) -> None:
@@ -83,6 +126,10 @@ class SqliteAdapter(BaseSQLAlchemyAdapter):
83
126
  def get_dialect(self) -> str:
84
127
  return sqlite.dialect.name
85
128
 
129
+ def render_sql(self, expression: exp.Expression) -> str:
130
+ """sqlglot's SQLite rendering, with date parts and truncation rewritten."""
131
+ return expression.copy().transform(_sqlite_dates).sql(dialect=self.get_dialect())
132
+
86
133
  @property
87
134
  def exclude_schemas(self) -> set[str]:
88
135
  return set()
@@ -4,8 +4,9 @@ from typing import Dict, List, Tuple
4
4
 
5
5
  import polars as pl
6
6
  from nl2sql.execution.contracts import ArtifactRef
7
- from nl2sql.pipeline.nodes.global_planner.schemas import ExecutionDAG, LogicalNode, LogicalEdge
7
+ from nl2sql.execution.dag import ExecutionDAG, LogicalNode, LogicalEdge
8
8
 
9
+ from .columns import input_order
9
10
  from .engines.polars_duckdb import PolarsDuckdbEngine
10
11
 
11
12
 
@@ -90,9 +91,5 @@ class AggregationService:
90
91
  edges: List[LogicalEdge],
91
92
  computed: Dict[str, pl.DataFrame],
92
93
  ) -> List[Tuple[str, pl.DataFrame]]:
93
- def role_rank(role: str) -> int:
94
- order = {"left": 0, "base": 0, "primary": 0, "right": 1, "compare": 1, "secondary": 1}
95
- return order.get(role or "", 2)
96
-
97
- ordered = sorted(edges, key=lambda e: (role_rank(e.role), e.from_id))
94
+ ordered = sorted(edges, key=lambda e: input_order(e.role, e.from_id))
98
95
  return [(edge.role or "", computed[edge.from_id]) for edge in ordered if edge.from_id in computed]
@@ -0,0 +1,110 @@
1
+ """Which columns a combine produces, and which a post-combine op reads.
2
+
3
+ The aggregation engine combines sub-query results and runs post-combine ops on
4
+ them; a post-combine op that reads a column the combine does not produce fails
5
+ there, after every sub-query has run, and nothing upstream can retry it. The
6
+ decomposer asks these same functions first, from the decomposition alone, so
7
+ such a plan is rejected while the model can still be asked again.
8
+
9
+ A sub-query's result columns are its ``expected_schema`` names: the logical
10
+ validator rejects a plan whose select aliases differ from them.
11
+
12
+ Plain Python, no polars: :class:`~nl2sql.aggregation.engines.polars_duckdb.PolarsDuckdbEngine`
13
+ uses the same functions, so the prediction and the engine cannot drift apart.
14
+ """
15
+ from __future__ import annotations
16
+
17
+ from typing import Any, Dict, Iterable, List, Optional, Sequence, Tuple
18
+
19
+ # The order a combine's inputs are given to the engine: left-hand roles first.
20
+ _ROLE_RANK = {"left": 0, "base": 0, "primary": 0, "right": 1, "compare": 1, "secondary": 1}
21
+
22
+ _SIDE_PREFIXES = ("left", "right", "base", "compare", "primary", "secondary")
23
+
24
+ # What the engine appends to a right-hand column whose name the left side has.
25
+ RIGHT_SUFFIX = "_right"
26
+
27
+
28
+ def input_order(role: Optional[str], source_id: str) -> Tuple[int, str]:
29
+ """Sort key for a combine's inputs, as ``(role, source id)``."""
30
+ return _ROLE_RANK.get(role or "", 2), source_id
31
+
32
+
33
+ def resolve_key(columns: Sequence[str], key: str) -> str:
34
+ """A join key as one of ``columns``.
35
+
36
+ Models write keys qualified by side or sub-query (``right.customer``,
37
+ ``sq_2.customer``); results have plain column names. An unknown key is
38
+ returned unchanged.
39
+ """
40
+ if key in columns or not key or "." not in key:
41
+ return key
42
+ bare = key.rsplit(".", 1)[1]
43
+ return bare if bare in columns else key
44
+
45
+
46
+ def combined_columns(
47
+ operation: str,
48
+ inputs: Sequence[Sequence[str]],
49
+ join_keys: Iterable[Dict[str, Any]],
50
+ ) -> List[str]:
51
+ """The columns the engine's combine produces from inputs with these columns.
52
+
53
+ ``inputs`` are in engine order (:func:`input_order`). A ``join`` or
54
+ ``compare`` is an inner join of the first two on the join keys: the
55
+ right-hand keys are dropped and any other right-hand column the left side
56
+ also has is suffixed ``_right``.
57
+ """
58
+ if not inputs:
59
+ return []
60
+ left = list(inputs[0])
61
+ if operation in ("standalone", "union") or len(inputs) < 2:
62
+ return left
63
+ right = list(inputs[1])
64
+ right_keys = {resolve_key(right, (k or {}).get("right") or "") for k in join_keys}
65
+ out = list(left)
66
+ for column in right:
67
+ if column in right_keys:
68
+ continue
69
+ out.append(column + RIGHT_SUFFIX if column in left else column)
70
+ return out
71
+
72
+
73
+ def post_op_columns(operation: str, attributes: Dict[str, Any]) -> List[str]:
74
+ """Every column a post-combine op reads from the combined result."""
75
+ return [name for name in (
76
+ *((g or {}).get("attribute") for g in attributes.get("group_by") or []),
77
+ *((m or {}).get("name") for m in attributes.get("metrics") or []),
78
+ *((c or {}).get("name") for c in attributes.get("expected_schema") or []
79
+ if operation == "project"),
80
+ *((f or {}).get("attribute") for f in attributes.get("filters") or []),
81
+ *((o or {}).get("attribute") for o in attributes.get("order_by") or []),
82
+ ) if name]
83
+
84
+
85
+ def missing_columns_message(missing: Sequence[str], available: Sequence[str]) -> str:
86
+ """Why a post-combine op that reads ``missing`` cannot run on ``available``.
87
+
88
+ A side-qualified name (``right.customer``) is almost always one question:
89
+ "which customers bought jazz but never rock?". The plan language has no
90
+ anti-join -- every two-input combine is an inner join -- so the model
91
+ reaches for a right-hand column to negate against, and that column does not
92
+ survive the combine. Resolving the prefix away would filter the
93
+ *intersection* and answer "bought both", silently wrong.
94
+ """
95
+ detail = (
96
+ f"Post-combine operation references {', '.join(repr(m) for m in missing)}, "
97
+ f"which the combined result does not have. Available columns: {', '.join(available) or 'none'}."
98
+ )
99
+ if any("." in m and m.split(".", 1)[0].lower() in _SIDE_PREFIXES for m in missing):
100
+ detail += (
101
+ " A column qualified by side does not survive the combine: joining on"
102
+ " the shared column drops the right-hand copy, and the other right-hand"
103
+ " columns are suffixed '_right'. This is what a question of the form"
104
+ " 'has X but never Y' looks like here, and it cannot be expressed:"
105
+ " the plan language has no anti-join or set-difference operation"
106
+ " (combine is one of standalone, compare, join, union -- all inner),"
107
+ " so the question is refused rather than answered with the"
108
+ " intersection."
109
+ )
110
+ return detail
@@ -5,10 +5,77 @@ from typing import Any, Dict, List, Tuple
5
5
  import duckdb
6
6
  import polars as pl
7
7
 
8
+ from nl2sql.aggregation.columns import missing_columns_message, post_op_columns, resolve_key
8
9
  from nl2sql.execution.contracts import ArtifactRef
9
10
  from nl2sql.execution.artifacts import build_artifact_store
10
11
 
11
12
 
13
+ def _column(frame: pl.DataFrame, key: str) -> str:
14
+ """A join key as a column of ``frame``; see :func:`~nl2sql.aggregation.columns.resolve_key`."""
15
+ return resolve_key(frame.columns, key)
16
+
17
+
18
+ def _check_attributes(frame: pl.DataFrame, names: List[str]) -> None:
19
+ """Refuses a post-combine attribute the combined frame does not have.
20
+
21
+ polars reports this as ``unable to find column "right.customer"; valid
22
+ columns: ["customer"]``, which tells neither the caller nor the planner
23
+ anything. The decomposer runs the same check on the decomposition before
24
+ anything executes (``nl2sql.aggregation.columns``) and asks the model
25
+ again, so this is the last line of defence; the message explains why a
26
+ side-qualified column cannot work (see ``missing_columns_message``).
27
+ """
28
+ missing = [name for name in names if name and name not in frame.columns]
29
+ if missing:
30
+ raise ValueError(missing_columns_message(missing, frame.columns))
31
+
32
+
33
+ def _filter(rows: pl.DataFrame, filters: List[Dict[str, Any]]) -> pl.DataFrame:
34
+ for flt in filters:
35
+ col, op, val = pl.col(flt.get("attribute")), flt.get("operator"), flt.get("value")
36
+ if op == "=":
37
+ rows = rows.filter(col == val)
38
+ elif op == "!=":
39
+ rows = rows.filter(col != val)
40
+ elif op == ">":
41
+ rows = rows.filter(col > val)
42
+ elif op == ">=":
43
+ rows = rows.filter(col >= val)
44
+ elif op == "<":
45
+ rows = rows.filter(col < val)
46
+ elif op == "<=":
47
+ rows = rows.filter(col <= val)
48
+ elif op == "between" and isinstance(val, list) and len(val) == 2:
49
+ rows = rows.filter((col >= val[0]) & (col <= val[1]))
50
+ elif op == "in" and isinstance(val, list):
51
+ rows = rows.filter(col.is_in(val))
52
+ elif op == "contains":
53
+ rows = rows.filter(col.cast(pl.Utf8).str.contains(str(val)))
54
+ return rows
55
+
56
+
57
+ _AGGREGATIONS = {
58
+ "count": lambda c: c.count(),
59
+ "sum": lambda c: c.sum(),
60
+ "avg": lambda c: c.mean(),
61
+ "min": lambda c: c.min(),
62
+ "max": lambda c: c.max(),
63
+ }
64
+
65
+
66
+ def _aggregate(frame: pl.DataFrame, attributes: Dict[str, Any]) -> pl.DataFrame:
67
+ group_by = [g.get("attribute") for g in attributes.get("group_by", []) if g.get("attribute")]
68
+ exprs = [
69
+ _AGGREGATIONS[m.get("aggregation")](pl.col(m.get("name"))).alias(m.get("name"))
70
+ for m in attributes.get("metrics", [])
71
+ if m.get("aggregation") in _AGGREGATIONS
72
+ ]
73
+ if group_by:
74
+ # polars 1.x: group_by (groupby was removed).
75
+ return frame.group_by(group_by, maintain_order=True).agg(exprs)
76
+ return frame.select(exprs)
77
+
78
+
12
79
  class PolarsDuckdbEngine:
13
80
  def __init__(self):
14
81
  self.artifact_store = build_artifact_store()
@@ -34,16 +101,16 @@ class PolarsDuckdbEngine:
34
101
  return frames[0]
35
102
  left = frames[0]
36
103
  right = frames[1]
37
- left_on = [k.get("left") for k in join_keys]
38
- right_on = [k.get("right") for k in join_keys]
104
+ left_on = [_column(left, k.get("left")) for k in join_keys]
105
+ right_on = [_column(right, k.get("right")) for k in join_keys]
39
106
  return left.join(right, left_on=left_on, right_on=right_on, how="inner", suffix="_right")
40
107
  if operation == "compare":
41
108
  if len(frames) < 2:
42
109
  return frames[0]
43
110
  left = frames[0]
44
111
  right = frames[1]
45
- left_on = [k.get("left") for k in join_keys]
46
- right_on = [k.get("right") for k in join_keys]
112
+ left_on = [_column(left, k.get("left")) for k in join_keys]
113
+ right_on = [_column(right, k.get("right")) for k in join_keys]
47
114
  joined = left.join(right, left_on=left_on, right_on=right_on, how="inner", suffix="_right")
48
115
  diff_cols = []
49
116
  for col in left.columns:
@@ -59,67 +126,33 @@ class PolarsDuckdbEngine:
59
126
  raise ValueError(f"Unsupported combine operation '{operation}'.")
60
127
 
61
128
  def post_op(self, operation: str, frame: pl.DataFrame, attributes: Dict[str, Any]) -> pl.DataFrame:
62
- if operation == "filter":
63
- rows = frame
64
- for flt in attributes.get("filters", []):
65
- attr = flt.get("attribute")
66
- op = flt.get("operator")
67
- val = flt.get("value")
68
- col = pl.col(attr)
69
- if op == "=":
70
- rows = rows.filter(col == val)
71
- elif op == "!=":
72
- rows = rows.filter(col != val)
73
- elif op == ">":
74
- rows = rows.filter(col > val)
75
- elif op == ">=":
76
- rows = rows.filter(col >= val)
77
- elif op == "<":
78
- rows = rows.filter(col < val)
79
- elif op == "<=":
80
- rows = rows.filter(col <= val)
81
- elif op == "between" and isinstance(val, list) and len(val) == 2:
82
- rows = rows.filter((col >= val[0]) & (col <= val[1]))
83
- elif op == "in" and isinstance(val, list):
84
- rows = rows.filter(col.is_in(val))
85
- elif op == "contains":
86
- rows = rows.filter(col.cast(pl.Utf8).str.contains(str(val)))
87
- return rows
129
+ """Applies one post-combine op, every field it carries, in SQL's order.
130
+
131
+ ``operation`` picks the reshaping step (``aggregate`` or ``project``);
132
+ whatever the operation, the op's filters, order_by and limit are then
133
+ applied, in that order. A filter after an aggregate filters the
134
+ aggregated rows, as HAVING does. The decomposer's own example is a
135
+ ``filter`` op carrying ``order_by`` and ``limit``, so applying only the
136
+ named field dropped them.
137
+ """
138
+ if operation not in {"filter", "aggregate", "project", "sort", "limit"}:
139
+ raise ValueError(f"Unsupported post-combine operation '{operation}'.")
140
+ # Every column this op will read, checked against the combined frame
141
+ # before polars is asked for any of them. See `_check_attributes`.
142
+ _check_attributes(frame, post_op_columns(operation, attributes))
143
+ rows = frame
88
144
  if operation == "aggregate":
89
- group_by = [g.get("attribute") for g in attributes.get("group_by", []) if g.get("attribute")]
90
- metrics = attributes.get("metrics", [])
91
- agg_exprs = []
92
- for metric in metrics:
93
- name = metric.get("name")
94
- agg = metric.get("aggregation")
95
- col = pl.col(name)
96
- if agg == "count":
97
- agg_exprs.append(col.count().alias(name))
98
- elif agg == "sum":
99
- agg_exprs.append(col.sum().alias(name))
100
- elif agg == "avg":
101
- agg_exprs.append(col.mean().alias(name))
102
- elif agg == "min":
103
- agg_exprs.append(col.min().alias(name))
104
- elif agg == "max":
105
- agg_exprs.append(col.max().alias(name))
106
- if group_by:
107
- return frame.groupby(group_by).agg(agg_exprs)
108
- return frame.select(agg_exprs)
109
- if operation == "project":
145
+ rows = _aggregate(rows, attributes)
146
+ elif operation == "project":
110
147
  columns = [c.get("name") for c in attributes.get("expected_schema", []) if c.get("name")]
111
- return frame.select(columns) if columns else frame
112
- if operation == "sort":
113
- order_by = attributes.get("order_by", [])
114
- sort_cols = [o.get("attribute") for o in order_by if o.get("attribute")]
115
- descending = [o.get("direction") == "desc" for o in order_by if o.get("attribute")]
116
- if sort_cols:
117
- return frame.sort(sort_cols, descending=descending)
118
- return frame
119
- if operation == "limit":
120
- limit = attributes.get("limit")
121
- return frame.head(limit) if limit is not None else frame
122
- raise ValueError(f"Unsupported post-combine operation '{operation}'.")
148
+ rows = rows.select(columns) if columns else rows
149
+ rows = _filter(rows, attributes.get("filters", []))
150
+ order_by = [o for o in attributes.get("order_by", []) if o.get("attribute")]
151
+ if order_by:
152
+ rows = rows.sort([o["attribute"] for o in order_by],
153
+ descending=[o.get("direction") == "desc" for o in order_by])
154
+ limit = attributes.get("limit")
155
+ return rows.head(limit) if limit is not None else rows
123
156
 
124
157
  def to_rows(self, frame: pl.DataFrame) -> List[Dict[str, Any]]:
125
158
  return frame.to_dicts()