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.
- nl2sql/__init__.py +28 -4
- nl2sql/adapters/mssql/adapter.py +2 -2
- nl2sql/adapters/postgres/adapter.py +5 -3
- nl2sql/adapters/sqlalchemy_base/__init__.py +0 -4
- nl2sql/adapters/sqlalchemy_base/adapter.py +11 -6
- nl2sql/adapters/sqlalchemy_base/models.py +4 -17
- nl2sql/adapters/sqlite/adapter.py +48 -1
- nl2sql/aggregation/aggregator.py +3 -6
- nl2sql/aggregation/columns.py +110 -0
- nl2sql/aggregation/engines/polars_duckdb.py +96 -63
- nl2sql/api/benchmark_api.py +101 -61
- nl2sql/api/query_api.py +142 -9
- nl2sql/auth/rbac.py +28 -15
- nl2sql/cli/commands/benchmark.py +298 -16
- nl2sql/cli/commands/cache.py +32 -0
- nl2sql/cli/commands/demo.py +533 -0
- nl2sql/cli/commands/doctor.py +87 -2
- nl2sql/cli/commands/feedback.py +145 -0
- nl2sql/cli/commands/indexing.py +4 -125
- nl2sql/cli/commands/info.py +8 -1
- nl2sql/cli/commands/run.py +24 -7
- nl2sql/cli/commands/setup.py +77 -55
- nl2sql/cli/commands/trace.py +170 -0
- nl2sql/cli/common/indexing.py +121 -0
- nl2sql/cli/console.py +1 -1
- nl2sql/cli/demo/chinook.py +57 -0
- nl2sql/cli/demo/datasets.py +110 -0
- nl2sql/cli/demo/defaults.py +2 -117
- nl2sql/cli/demo/llm_config.py +98 -0
- nl2sql/cli/demo/manager.py +113 -139
- nl2sql/cli/demo/playground/__init__.py +1 -0
- nl2sql/cli/demo/playground/app.py +622 -0
- nl2sql/cli/demo/playground/assets/clips/ask-dark.jpg +0 -0
- nl2sql/cli/demo/playground/assets/clips/ask-dark.webm +0 -0
- nl2sql/cli/demo/playground/assets/clips/ask.jpg +0 -0
- nl2sql/cli/demo/playground/assets/clips/ask.webm +0 -0
- nl2sql/cli/demo/playground/assets/clips/pipeline-dark.jpg +0 -0
- nl2sql/cli/demo/playground/assets/clips/pipeline-dark.webm +0 -0
- nl2sql/cli/demo/playground/assets/clips/pipeline.jpg +0 -0
- nl2sql/cli/demo/playground/assets/clips/pipeline.webm +0 -0
- nl2sql/cli/demo/playground/assets/clips/retrieval-dark.jpg +0 -0
- nl2sql/cli/demo/playground/assets/clips/retrieval-dark.webm +0 -0
- nl2sql/cli/demo/playground/assets/clips/retrieval.jpg +0 -0
- nl2sql/cli/demo/playground/assets/clips/retrieval.webm +0 -0
- nl2sql/cli/demo/playground/assets/social-card.png +0 -0
- nl2sql/cli/demo/playground/hosted.py +443 -0
- nl2sql/cli/demo/playground/index_panel.py +133 -0
- nl2sql/cli/demo/playground/preview.py +101 -0
- nl2sql/cli/demo/playground/settings.py +399 -0
- nl2sql/cli/demo/playground/static/index.html +68 -0
- nl2sql/cli/demo/recordings/chinook.json +3817 -0
- nl2sql/cli/demo/stamp.py +66 -0
- nl2sql/cli/demo/support.py +45 -0
- nl2sql/cli/demo/webanalytics.py +41 -0
- nl2sql/cli/generators/env/templates.py +18 -1
- nl2sql/cli/generators/llm/generator.py +13 -2
- nl2sql/cli/main.py +275 -24
- nl2sql/cli/reporting.py +152 -332
- nl2sql/common/env_hint.py +32 -0
- nl2sql/common/errors.py +28 -15
- nl2sql/common/metrics.py +10 -5
- nl2sql/common/settings.py +106 -14
- nl2sql/configs/llm.py +23 -3
- nl2sql/configs/manager.py +19 -8
- nl2sql/context.py +23 -2
- nl2sql/datasets/README.md +108 -0
- nl2sql/datasets/__init__.py +21 -0
- nl2sql/datasets/chinook.sqlite +0 -0
- nl2sql/datasets/support.sqlite +0 -0
- nl2sql/datasets/webanalytics.sqlite +0 -0
- nl2sql/datasources/models.py +21 -3
- nl2sql/datasources/registry.py +31 -12
- nl2sql/evaluation/benchmark_runner.py +93 -291
- nl2sql/evaluation/datasets/chinook_gold.yaml +1079 -0
- nl2sql/evaluation/datasets/chinook_gold_plans.yaml +745 -0
- nl2sql/evaluation/evaluator.py +295 -94
- nl2sql/evaluation/faithfulness.py +134 -0
- nl2sql/evaluation/gold.py +148 -0
- nl2sql/evaluation/presets/__init__.py +79 -0
- nl2sql/evaluation/presets/claude-planner.yaml +20 -0
- nl2sql/evaluation/presets/gpt-5.4-mini-helpers.yaml +19 -0
- nl2sql/evaluation/presets/gpt-5.4.yaml +10 -0
- nl2sql/evaluation/prices.py +53 -0
- nl2sql/evaluation/records.py +613 -0
- nl2sql/evaluation/retrieval_recall.py +174 -0
- nl2sql/evaluation/tier1.py +140 -0
- nl2sql/evaluation/tier2.py +519 -0
- nl2sql/evaluation/types.py +10 -4
- nl2sql/execution/artifacts/store.py +5 -4
- nl2sql/{pipeline/nodes/global_planner/schemas.py → execution/dag.py} +4 -24
- nl2sql/execution/executor/sql_executor.py +1 -2
- nl2sql/feedback/__init__.py +16 -0
- nl2sql/feedback/drafts.py +65 -0
- nl2sql/feedback/record.py +90 -0
- nl2sql/feedback/stats.py +83 -0
- nl2sql/feedback/store.py +112 -0
- nl2sql/indexing/embeddings.py +12 -3
- nl2sql/indexing/enrichment_service.py +2 -3
- nl2sql/indexing/health.py +207 -0
- nl2sql/indexing/orchestrator.py +30 -8
- nl2sql/indexing/rebuild.py +144 -0
- nl2sql/indexing/retrieval_trace.py +144 -0
- nl2sql/indexing/vector_store.py +364 -127
- nl2sql/llm/failures.py +185 -0
- nl2sql/llm/providers.py +183 -0
- nl2sql/llm/registry.py +169 -38
- nl2sql/llm/replay.py +290 -0
- nl2sql/llm/request_key.py +182 -0
- nl2sql/llm/wires/__init__.py +33 -0
- nl2sql/llm/wires/anthropic.py +133 -0
- nl2sql/llm/wires/anthropic_client.py +63 -0
- nl2sql/llm/wires/base.py +87 -0
- nl2sql/llm/wires/openai.py +119 -0
- nl2sql/pipeline/graph.py +6 -8
- nl2sql/pipeline/graph_utils.py +85 -28
- nl2sql/pipeline/nodes/__init__.py +9 -24
- nl2sql/pipeline/nodes/aggregator/node.py +2 -7
- nl2sql/pipeline/nodes/aggregator/schemas.py +1 -20
- nl2sql/pipeline/nodes/answer_synthesizer/node.py +4 -4
- nl2sql/pipeline/nodes/ast_planner/__init__.py +8 -3
- nl2sql/pipeline/nodes/ast_planner/functions.py +68 -0
- nl2sql/pipeline/nodes/ast_planner/node.py +77 -10
- nl2sql/pipeline/nodes/ast_planner/prompts.py +84 -16
- nl2sql/pipeline/nodes/ast_planner/schemas.py +16 -13
- nl2sql/pipeline/nodes/datasource_resolver/node.py +167 -69
- nl2sql/pipeline/nodes/datasource_resolver/prompts.py +26 -0
- nl2sql/pipeline/nodes/datasource_resolver/schemas.py +10 -0
- nl2sql/pipeline/nodes/decomposer/dag.py +121 -0
- nl2sql/pipeline/nodes/decomposer/node.py +171 -18
- nl2sql/pipeline/nodes/decomposer/prompts.py +46 -32
- nl2sql/pipeline/nodes/decomposer/schemas.py +59 -30
- nl2sql/pipeline/nodes/executor/node.py +3 -13
- nl2sql/pipeline/nodes/generator/node.py +286 -63
- nl2sql/pipeline/nodes/refiner/node.py +8 -17
- nl2sql/pipeline/nodes/refiner/prompts.py +21 -12
- nl2sql/pipeline/nodes/schema_retriever/node.py +114 -9
- nl2sql/pipeline/nodes/schema_retriever/schema.py +43 -1
- nl2sql/pipeline/nodes/validator/node.py +347 -168
- nl2sql/pipeline/nodes/validator/schemas.py +14 -0
- nl2sql/pipeline/plan_cache.py +112 -0
- nl2sql/pipeline/routes.py +29 -21
- nl2sql/pipeline/runtime.py +165 -35
- nl2sql/pipeline/state.py +3 -3
- nl2sql/pipeline/steps.py +130 -0
- nl2sql/pipeline/subgraphs/schemas.py +5 -1
- nl2sql/pipeline/subgraphs/sql_agent.py +31 -21
- nl2sql/pipeline/timing.py +52 -0
- nl2sql/public_api.py +111 -5
- nl2sql/schema/in_memory_store.py +16 -0
- nl2sql/schema/protocol.py +16 -1
- nl2sql/schema/sqlite_store.py +105 -4
- nl2sql/schema/view.py +56 -0
- nl2sql/secrets/factory.py +4 -3
- nl2sql/secrets/manager.py +15 -2
- nl2sql/services/callbacks/monitor.py +10 -11
- nl2sql/services/callbacks/node_handlers.py +1 -2
- nl2sql/services/callbacks/node_metrics.py +0 -3
- nl2sql/services/callbacks/token_handler.py +249 -42
- nl2sql/testing/fake_llm.py +242 -0
- nl2sql/tracing/__init__.py +7 -0
- nl2sql/tracing/document.py +212 -0
- nl2sql/tracing/recorder.py +376 -0
- nl2sql/tracing/replay.py +318 -0
- nl2sql/tracing/trace.py +173 -0
- nl2sql_engine-0.2.0.dist-info/METADATA +226 -0
- nl2sql_engine-0.2.0.dist-info/RECORD +267 -0
- nl2sql_engine-0.2.0.dist-info/licenses/LICENSE +21 -0
- nl2sql_engine-0.2.0.dist-info/licenses/THIRD_PARTY_NOTICES.md +249 -0
- nl2sql/api/result_api.py +0 -24
- nl2sql/cli/common/prompts.py +0 -53
- nl2sql/cli/demo/data.py +0 -87
- nl2sql/cli/demo/factory.py +0 -289
- nl2sql/cli/demo/schemas.py +0 -348
- nl2sql/cli/demo/writers/docker.py +0 -194
- nl2sql/cli/demo/writers/sqlite.py +0 -88
- nl2sql/pipeline/nodes/aggregator/prompts.py +0 -20
- nl2sql/pipeline/nodes/global_planner/__init__.py +0 -4
- nl2sql/pipeline/nodes/global_planner/node.py +0 -186
- nl2sql_engine-0.1.2.dist-info/METADATA +0 -295
- nl2sql_engine-0.1.2.dist-info/RECORD +0 -193
- /nl2sql/{cli/demo/writers → testing}/__init__.py +0 -0
- {nl2sql_engine-0.1.2.dist-info → nl2sql_engine-0.2.0.dist-info}/WHEEL +0 -0
- {nl2sql_engine-0.1.2.dist-info → nl2sql_engine-0.2.0.dist-info}/entry_points.txt +0 -0
- {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
|
-
|
|
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
|
]
|
nl2sql/adapters/mssql/adapter.py
CHANGED
|
@@ -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
|
-
|
|
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
|
-
|
|
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
|
-
|
|
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
|
|
2
|
-
from typing import
|
|
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()
|
nl2sql/aggregation/aggregator.py
CHANGED
|
@@ -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.
|
|
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
|
-
|
|
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
|
-
|
|
63
|
-
|
|
64
|
-
|
|
65
|
-
|
|
66
|
-
|
|
67
|
-
|
|
68
|
-
|
|
69
|
-
|
|
70
|
-
|
|
71
|
-
|
|
72
|
-
|
|
73
|
-
|
|
74
|
-
|
|
75
|
-
|
|
76
|
-
|
|
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
|
-
|
|
90
|
-
|
|
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
|
-
|
|
112
|
-
|
|
113
|
-
|
|
114
|
-
|
|
115
|
-
|
|
116
|
-
|
|
117
|
-
|
|
118
|
-
|
|
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()
|