sql-safe-mcp 1.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.
sql_safe_mcp/errors.py ADDED
@@ -0,0 +1,49 @@
1
+ from __future__ import annotations
2
+
3
+ from dataclasses import dataclass
4
+ from enum import StrEnum
5
+ from uuid import uuid4
6
+
7
+
8
+ class ErrorCode(StrEnum):
9
+ INVALID_ARGUMENT = "INVALID_ARGUMENT"
10
+ UNKNOWN_SERVER = "UNKNOWN_SERVER"
11
+ NOT_FOUND = "NOT_FOUND"
12
+ AMBIGUOUS_OBJECT = "AMBIGUOUS_OBJECT"
13
+ ACCESS_LEVEL_DENIED = "ACCESS_LEVEL_DENIED"
14
+ CONFIG_ERROR = "CONFIG_ERROR"
15
+ CONNECTION_FAILED = "CONNECTION_FAILED"
16
+ ACCESS_DENIED = "ACCESS_DENIED"
17
+ TIMEOUT = "TIMEOUT"
18
+ METADATA_UNAVAILABLE = "METADATA_UNAVAILABLE"
19
+ DEFINITION_TOO_LARGE = "DEFINITION_TOO_LARGE"
20
+ QUERY_REJECTED = "QUERY_REJECTED"
21
+ INVALID_PII_TOKEN = "INVALID_PII_TOKEN"
22
+ RESULT_LIMIT_EXCEEDED = "RESULT_LIMIT_EXCEEDED"
23
+ DATABASE_ERROR = "DATABASE_ERROR"
24
+
25
+
26
+ @dataclass(slots=True)
27
+ class DomainError(Exception):
28
+ code: ErrorCode
29
+ public_message: str
30
+ hint: str | None = None
31
+ retryable: bool = False
32
+ correlation_id: str | None = None
33
+
34
+ def __str__(self) -> str:
35
+ parts = [f"[{self.code}] {self.public_message}"]
36
+ if self.hint:
37
+ parts.append(f"Hint: {self.hint}")
38
+ if self.correlation_id:
39
+ parts.append(f"Reference: {self.correlation_id}")
40
+ return " ".join(parts)
41
+
42
+ @classmethod
43
+ def unexpected(cls) -> DomainError:
44
+ return cls(
45
+ ErrorCode.DATABASE_ERROR,
46
+ "The database operation failed unexpectedly.",
47
+ retryable=False,
48
+ correlation_id=uuid4().hex,
49
+ )
@@ -0,0 +1,140 @@
1
+ from __future__ import annotations
2
+
3
+ from collections.abc import AsyncIterator, Awaitable, Callable
4
+ from contextlib import asynccontextmanager
5
+ from dataclasses import dataclass
6
+
7
+ from mcp.server.mcpserver import Context, MCPServer
8
+ from mcp.server.mcpserver.exceptions import ToolError
9
+ from mcp_types import ToolAnnotations
10
+
11
+ from sql_safe_mcp.config import AppConfig
12
+ from sql_safe_mcp.db.registry import EngineRegistry
13
+ from sql_safe_mcp.errors import DomainError
14
+ from sql_safe_mcp.models import (
15
+ DatabaseList,
16
+ ServerList,
17
+ SqlResult,
18
+ StoredProcedureDefinition,
19
+ StoredProcedureList,
20
+ TableDefinition,
21
+ TableList,
22
+ )
23
+ from sql_safe_mcp.service import DatabaseService
24
+
25
+ READ_ONLY = ToolAnnotations(read_only_hint=True, open_world_hint=False)
26
+
27
+
28
+ @dataclass(slots=True)
29
+ class AppContext:
30
+ service: DatabaseService
31
+ registry: EngineRegistry
32
+
33
+
34
+ async def _domain_call[T](operation: Callable[[], Awaitable[T]]) -> T:
35
+ try:
36
+ return await operation()
37
+ except DomainError as exc:
38
+ raise ToolError(str(exc)) from exc
39
+
40
+
41
+ def create_server(config: AppConfig) -> MCPServer[AppContext]:
42
+ @asynccontextmanager
43
+ async def lifespan(_server: MCPServer[AppContext]) -> AsyncIterator[AppContext]:
44
+ registry = EngineRegistry(config)
45
+ try:
46
+ yield AppContext(service=DatabaseService(config, registry), registry=registry)
47
+ finally:
48
+ registry.dispose()
49
+
50
+ server: MCPServer[AppContext] = MCPServer(
51
+ "sql-safe-mcp",
52
+ description="Minimal, read-only, PII-safe SQL database navigation.",
53
+ version="1.2.0",
54
+ lifespan=lifespan,
55
+ )
56
+
57
+ @server.tool(annotations=READ_ONLY)
58
+ async def list_servers(ctx: Context[AppContext]) -> ServerList:
59
+ """List explicitly configured database server aliases."""
60
+ return ctx.request_context.lifespan_context.service.list_servers()
61
+
62
+ @server.tool(annotations=READ_ONLY)
63
+ async def list_databases(
64
+ server: str,
65
+ ctx: Context[AppContext],
66
+ name_contains: str | None = None,
67
+ ) -> DatabaseList:
68
+ """List databases visible to the configured credentials."""
69
+ service = ctx.request_context.lifespan_context.service
70
+ return await _domain_call(lambda: service.list_databases(server, name_contains))
71
+
72
+ @server.tool(annotations=READ_ONLY)
73
+ async def list_tables(
74
+ server: str,
75
+ database: str,
76
+ ctx: Context[AppContext],
77
+ schema: str | None = None,
78
+ name_contains: str | None = None,
79
+ ) -> TableList:
80
+ """List base tables, optionally filtering by schema and a name substring."""
81
+ service = ctx.request_context.lifespan_context.service
82
+ return await _domain_call(
83
+ lambda: service.list_tables(server, database, schema, name_contains)
84
+ )
85
+
86
+ @server.tool(annotations=READ_ONLY)
87
+ async def get_table_definition(
88
+ server: str,
89
+ database: str,
90
+ table: str,
91
+ ctx: Context[AppContext],
92
+ schema: str | None = None,
93
+ ) -> TableDefinition:
94
+ """Get columns, keys, constraints, and indexes for one table."""
95
+ service = ctx.request_context.lifespan_context.service
96
+ return await _domain_call(
97
+ lambda: service.get_table_definition(server, database, table, schema)
98
+ )
99
+
100
+ @server.tool(annotations=READ_ONLY)
101
+ async def list_stored_procedures(
102
+ server: str,
103
+ database: str,
104
+ ctx: Context[AppContext],
105
+ schema: str | None = None,
106
+ name_contains: str | None = None,
107
+ ) -> StoredProcedureList:
108
+ """List stored procedures without expanding their definitions."""
109
+ service = ctx.request_context.lifespan_context.service
110
+ return await _domain_call(
111
+ lambda: service.list_stored_procedures(server, database, schema, name_contains)
112
+ )
113
+
114
+ @server.tool(annotations=READ_ONLY)
115
+ async def get_stored_procedure(
116
+ server: str,
117
+ database: str,
118
+ name: str,
119
+ ctx: Context[AppContext],
120
+ schema: str | None = None,
121
+ ) -> StoredProcedureDefinition:
122
+ """Get the original definition of one stored procedure when visible."""
123
+ service = ctx.request_context.lifespan_context.service
124
+ return await _domain_call(
125
+ lambda: service.get_stored_procedure(server, database, name, schema)
126
+ )
127
+
128
+ @server.tool(annotations=READ_ONLY)
129
+ async def execute_sql(
130
+ server: str,
131
+ database: str,
132
+ sql: str,
133
+ ctx: Context[AppContext],
134
+ max_rows: int | None = None,
135
+ ) -> SqlResult:
136
+ """Run one restricted SELECT on a pii_safe server; protected columns return tokens."""
137
+ service = ctx.request_context.lifespan_context.service
138
+ return await _domain_call(lambda: service.execute_sql(server, database, sql, max_rows))
139
+
140
+ return server
sql_safe_mcp/models.py ADDED
@@ -0,0 +1,124 @@
1
+ from __future__ import annotations
2
+
3
+ from typing import Any, Literal
4
+
5
+ from pydantic import BaseModel, ConfigDict, Field
6
+
7
+
8
+ class OutputModel(BaseModel):
9
+ model_config = ConfigDict(extra="forbid", populate_by_name=True)
10
+
11
+
12
+ class ServerSummary(OutputModel):
13
+ name: str
14
+ engine: Literal["sqlserver", "mysql", "mariadb"]
15
+ access_level: Literal["metadata", "pii_safe"]
16
+
17
+
18
+ class ServerList(OutputModel):
19
+ servers: list[ServerSummary]
20
+
21
+
22
+ class DatabaseSummary(OutputModel):
23
+ name: str
24
+
25
+
26
+ class DatabaseList(OutputModel):
27
+ databases: list[DatabaseSummary]
28
+
29
+
30
+ class TableSummary(OutputModel):
31
+ schema_: str | None = Field(alias="schema")
32
+ name: str
33
+
34
+
35
+ class TableList(OutputModel):
36
+ tables: list[TableSummary]
37
+
38
+
39
+ class ColumnDefinition(OutputModel):
40
+ name: str
41
+ native_type: str
42
+ nullable: bool
43
+ length: int | None = None
44
+ precision: int | None = None
45
+ scale: int | None = None
46
+ autoincrement: bool | str | None = None
47
+ identity: dict[str, Any] | None = None
48
+ computed: dict[str, Any] | None = None
49
+ default: str | None = None
50
+
51
+
52
+ class PrimaryKeyDefinition(OutputModel):
53
+ name: str | None = None
54
+ columns: list[str]
55
+
56
+
57
+ class ForeignKeyDefinition(OutputModel):
58
+ name: str | None = None
59
+ columns: list[str]
60
+ referred_schema: str | None = None
61
+ referred_table: str
62
+ referred_columns: list[str]
63
+ options: dict[str, Any] = Field(default_factory=dict)
64
+
65
+
66
+ class UniqueConstraintDefinition(OutputModel):
67
+ name: str | None = None
68
+ columns: list[str]
69
+
70
+
71
+ class IndexDefinition(OutputModel):
72
+ name: str | None = None
73
+ columns: list[str | None]
74
+ unique: bool = False
75
+ expressions: list[str] = Field(default_factory=list)
76
+
77
+
78
+ class TableDefinition(OutputModel):
79
+ schema_: str | None = Field(alias="schema")
80
+ name: str
81
+ columns: list[ColumnDefinition]
82
+ primary_key: PrimaryKeyDefinition
83
+ foreign_keys: list[ForeignKeyDefinition]
84
+ unique_constraints: list[UniqueConstraintDefinition]
85
+ indexes: list[IndexDefinition]
86
+
87
+
88
+ class StoredProcedureSummary(OutputModel):
89
+ schema_: str | None = Field(alias="schema")
90
+ name: str
91
+
92
+
93
+ class StoredProcedureList(OutputModel):
94
+ stored_procedures: list[StoredProcedureSummary]
95
+
96
+
97
+ class StoredProcedureDefinition(OutputModel):
98
+ schema_: str | None = Field(alias="schema")
99
+ name: str
100
+ definition: str | None
101
+ definition_available: bool
102
+
103
+
104
+ ResultEncoding = Literal["json", "token", "decimal", "date", "time", "datetime", "uuid", "base64"]
105
+
106
+
107
+ class ColumnSource(OutputModel):
108
+ schema_: str | None = Field(alias="schema")
109
+ table: str
110
+ column: str
111
+
112
+
113
+ class SqlColumn(OutputModel):
114
+ name: str
115
+ source: ColumnSource | None
116
+ protected: bool
117
+ encoding: ResultEncoding
118
+
119
+
120
+ class SqlResult(OutputModel):
121
+ columns: list[SqlColumn]
122
+ rows: list[list[Any]]
123
+ row_count: int
124
+ truncated: bool
@@ -0,0 +1 @@
1
+ """SQL validation and PII protection pipeline for execute_sql."""
@@ -0,0 +1,29 @@
1
+ from __future__ import annotations
2
+
3
+ from dataclasses import dataclass
4
+
5
+
6
+ @dataclass(frozen=True, slots=True)
7
+ class SqlDialect:
8
+ """Everything the shared validation pipeline needs to know about one SQL dialect.
9
+
10
+ Policy, tokens, and ValidatedQuery are dialect-independent; only these facts differ. The
11
+ dialect always comes from the trusted server configuration, never from the caller.
12
+ """
13
+
14
+ name: str # SQLGlot dialect used to parse and to generate SQL
15
+ has_schema: bool # False when a database is the catalog and objects have no schema level
16
+ bind_marker: str # DBAPI positional marker written into generated SQL
17
+ escape_percent: bool # the driver applies %-formatting, so a literal % must be doubled
18
+
19
+
20
+ SQLSERVER = SqlDialect(name="tsql", has_schema=True, bind_marker="?", escape_percent=False)
21
+ MYSQL = SqlDialect(name="mysql", has_schema=False, bind_marker="%s", escape_percent=True)
22
+
23
+
24
+ def dialect_for(engine: str) -> SqlDialect:
25
+ if engine == "sqlserver":
26
+ return SQLSERVER
27
+ if engine in ("mysql", "mariadb"):
28
+ return MYSQL
29
+ raise ValueError(f"unsupported engine {engine!r}")
@@ -0,0 +1,129 @@
1
+ from __future__ import annotations
2
+
3
+ import base64
4
+ import math
5
+ from collections.abc import Sequence
6
+ from dataclasses import dataclass
7
+ from datetime import date, datetime, time, timedelta
8
+ from decimal import Decimal
9
+ from typing import Any, Protocol
10
+ from uuid import UUID
11
+
12
+ from sql_safe_mcp.errors import DomainError, ErrorCode
13
+ from sql_safe_mcp.models import ResultEncoding
14
+ from sql_safe_mcp.security.lineage import SourceColumn
15
+ from sql_safe_mcp.security.tokens import TokenCodec
16
+ from sql_safe_mcp.security.validated_query import ValidatedQuery
17
+
18
+
19
+ class _Result(Protocol):
20
+ def keys(self) -> Any: ...
21
+
22
+ def fetchmany(self, size: int) -> Sequence[Any]: ...
23
+
24
+ def close(self) -> None: ...
25
+
26
+
27
+ class SqlConnection(Protocol):
28
+ """The part of a SQLAlchemy Core connection the executor uses."""
29
+
30
+ def exec_driver_sql(self, statement: str, parameters: tuple[Any, ...]) -> _Result: ...
31
+
32
+
33
+ @dataclass(frozen=True, slots=True)
34
+ class ResultColumn:
35
+ name: str
36
+ source: SourceColumn | None
37
+ protected: bool
38
+ encoding: ResultEncoding
39
+
40
+
41
+ @dataclass(frozen=True, slots=True)
42
+ class ExecutionResult:
43
+ columns: tuple[ResultColumn, ...]
44
+ rows: tuple[tuple[Any, ...], ...]
45
+ truncated: bool
46
+
47
+
48
+ def _unsupported_result() -> DomainError:
49
+ return DomainError(
50
+ ErrorCode.DATABASE_ERROR,
51
+ "The result contains a value type that is not supported.",
52
+ "Select different columns.",
53
+ )
54
+
55
+
56
+ def _encode(value: Any) -> tuple[Any, ResultEncoding]:
57
+ """Return the JSON-safe form and encoding name of one non-null, unprotected value."""
58
+ kind = type(value)
59
+ if kind in (str, bool, int):
60
+ return value, "json"
61
+ if kind is float:
62
+ if not math.isfinite(value):
63
+ raise _unsupported_result()
64
+ return value, "json"
65
+ if kind is Decimal:
66
+ if not value.is_finite():
67
+ raise _unsupported_result()
68
+ return str(value), "decimal"
69
+ if kind is datetime:
70
+ return value.isoformat(), "datetime"
71
+ if kind is date:
72
+ return value.isoformat(), "date"
73
+ if kind is time:
74
+ return value.isoformat(), "time"
75
+ if kind is timedelta: # MySQL TIME; only a time of day fits the `time` encoding
76
+ if not timedelta(0) <= value < timedelta(days=1):
77
+ raise _unsupported_result()
78
+ return (datetime.min + value).time().isoformat(), "time"
79
+ if kind is UUID:
80
+ return str(value), "uuid"
81
+ if kind in (bytes, bytearray, memoryview):
82
+ return base64.b64encode(bytes(value)).decode("ascii"), "base64"
83
+ raise _unsupported_result()
84
+
85
+
86
+ def execute_validated(
87
+ connection: SqlConnection, query: ValidatedQuery, codec: TokenCodec
88
+ ) -> ExecutionResult:
89
+ """Run a ValidatedQuery: generated SQL and binds only, at most max_rows + 1 rows buffered."""
90
+ if not isinstance(query, ValidatedQuery):
91
+ raise TypeError("the executor accepts only ValidatedQuery")
92
+ if codec.alias != query.alias:
93
+ raise TypeError("the token codec belongs to a different server alias")
94
+
95
+ result = connection.exec_driver_sql(query.sql, query.parameters)
96
+ try:
97
+ if len(result.keys()) != len(query.outputs):
98
+ raise _unsupported_result()
99
+ fetched = result.fetchmany(query.max_rows + 1)
100
+ truncated = len(fetched) > query.max_rows
101
+ encodings: list[ResultEncoding | None] = [None] * len(query.outputs)
102
+ rows: list[tuple[Any, ...]] = []
103
+ for raw in fetched[: query.max_rows]:
104
+ cells: list[Any] = []
105
+ for index, (plan, value) in enumerate(zip(query.outputs, tuple(raw), strict=True)):
106
+ if value is None:
107
+ cells.append(None)
108
+ elif plan.protected:
109
+ cells.append(codec.encrypt(value))
110
+ else:
111
+ encoded, encoding = _encode(value)
112
+ if encodings[index] not in (None, encoding):
113
+ raise _unsupported_result()
114
+ encodings[index] = encoding
115
+ cells.append(encoded)
116
+ rows.append(tuple(cells))
117
+ finally:
118
+ result.close()
119
+
120
+ columns = tuple(
121
+ ResultColumn(
122
+ plan.label,
123
+ plan.source,
124
+ plan.protected,
125
+ "token" if plan.protected else (encodings[index] or "json"),
126
+ )
127
+ for index, plan in enumerate(query.outputs)
128
+ )
129
+ return ExecutionResult(columns, tuple(rows), truncated)
@@ -0,0 +1,224 @@
1
+ from __future__ import annotations
2
+
3
+ from collections.abc import Sequence
4
+ from dataclasses import dataclass
5
+ from typing import Literal, cast
6
+
7
+ from sqlglot import exp
8
+ from sqlglot.errors import OptimizeError, SqlglotError
9
+ from sqlglot.optimizer.qualify import qualify
10
+
11
+ from sql_safe_mcp.security.dialect import SqlDialect
12
+ from sql_safe_mcp.security.parser import (
13
+ ParserLimits,
14
+ check_limits,
15
+ reject,
16
+ validate_allowlist,
17
+ )
18
+ from sql_safe_mcp.security.reasons import Reason
19
+ from sql_safe_mcp.security.schema import TableSchema
20
+
21
+
22
+ @dataclass(frozen=True, slots=True)
23
+ class SourceColumn:
24
+ schema: str
25
+ table: str
26
+ column: str
27
+
28
+
29
+ @dataclass(frozen=True, slots=True)
30
+ class OutputColumn:
31
+ label: str
32
+ kind: Literal["column", "count", "literal"]
33
+ source: SourceColumn | None
34
+ binding: str | None = None # table alias/name the source column was read through
35
+
36
+
37
+ @dataclass(frozen=True, slots=True)
38
+ class ColumnRef:
39
+ node: exp.Column
40
+ source: SourceColumn
41
+
42
+
43
+ @dataclass(frozen=True, slots=True)
44
+ class AnalyzedQuery:
45
+ query: exp.Select
46
+ outputs: tuple[OutputColumn, ...]
47
+ references: tuple[ColumnRef, ...]
48
+
49
+
50
+ @dataclass(frozen=True, slots=True)
51
+ class _Binding:
52
+ name: str
53
+ table: TableSchema
54
+
55
+
56
+ def _identifier(name: str) -> exp.Identifier:
57
+ return exp.to_identifier(name, quoted=True)
58
+
59
+
60
+ def _collect(query: exp.Select, schemas: Sequence[TableSchema]) -> list[_Binding]:
61
+ nodes: list[exp.Table] = []
62
+ from_ = query.args.get("from_")
63
+ if from_ is not None:
64
+ nodes.append(from_.this)
65
+ nodes += [join.this for join in query.args.get("joins") or []]
66
+ if len(nodes) != len(schemas):
67
+ raise reject(Reason.TABLE_RESOLUTION_MISMATCH)
68
+ bindings: list[_Binding] = []
69
+ seen: set[str] = set()
70
+ for node, table in zip(nodes, schemas, strict=True):
71
+ name = node.alias or table.name
72
+ if name.casefold() in seen:
73
+ raise reject(Reason.DUPLICATE_BINDING)
74
+ seen.add(name.casefold())
75
+ bindings.append(_Binding(name, table))
76
+ return bindings
77
+
78
+
79
+ def _canonical_column(binding: _Binding, name: str) -> str | None:
80
+ matches = [column for column in binding.table.columns if column.casefold() == name.casefold()]
81
+ return matches[0] if len(matches) == 1 else None
82
+
83
+
84
+ def _resolve(column: exp.Column, bindings: Sequence[_Binding]) -> tuple[_Binding, str]:
85
+ name = column.name
86
+ qualifier = column.table
87
+ if qualifier:
88
+ matching = [b for b in bindings if b.name.casefold() == qualifier.casefold()]
89
+ if len(matching) != 1:
90
+ raise reject(Reason.QUALIFIER_UNKNOWN)
91
+ canonical = _canonical_column(matching[0], name)
92
+ if canonical is None:
93
+ raise reject(Reason.COLUMN_NOT_FOUND)
94
+ return matching[0], canonical
95
+ candidates = [
96
+ (binding, canonical)
97
+ for binding in bindings
98
+ if (canonical := _canonical_column(binding, name)) is not None
99
+ ]
100
+ if len(candidates) != 1:
101
+ raise reject(Reason.COLUMN_AMBIGUOUS)
102
+ return candidates[0]
103
+
104
+
105
+ def _rewrite(column: exp.Column, binding: _Binding, canonical: str) -> SourceColumn:
106
+ column.set("this", _identifier(canonical))
107
+ column.set("table", _identifier(binding.name))
108
+ return SourceColumn(binding.table.schema, binding.table.name, canonical)
109
+
110
+
111
+ def _expand_stars(query: exp.Select, bindings: Sequence[_Binding]) -> None:
112
+ expanded: list[exp.Expression] = []
113
+ for projection in query.expressions:
114
+ if isinstance(projection, exp.Star):
115
+ chosen = list(bindings)
116
+ elif isinstance(projection, exp.Column) and isinstance(projection.this, exp.Star):
117
+ chosen = [b for b in bindings if b.name.casefold() == projection.table.casefold()]
118
+ if len(chosen) != 1:
119
+ raise reject(Reason.QUALIFIER_UNKNOWN)
120
+ else:
121
+ expanded.append(projection)
122
+ continue
123
+ expanded.extend(
124
+ exp.Column(this=_identifier(name), table=_identifier(binding.name))
125
+ for binding in chosen
126
+ for name in binding.table.columns
127
+ )
128
+ query.set("expressions", expanded)
129
+
130
+
131
+ def _analyze_projections(
132
+ query: exp.Select, bindings: Sequence[_Binding], refs: list[ColumnRef]
133
+ ) -> list[OutputColumn]:
134
+ outputs: list[OutputColumn] = []
135
+ for index, projection in enumerate(query.expressions, start=1):
136
+ value = projection.this if isinstance(projection, exp.Alias) else projection
137
+ alias = projection.alias if isinstance(projection, exp.Alias) else ""
138
+ if isinstance(value, exp.Column):
139
+ binding, canonical = _resolve(value, bindings)
140
+ source = _rewrite(value, binding, canonical)
141
+ refs.append(ColumnRef(value, source))
142
+ outputs.append(OutputColumn(alias or canonical, "column", source, binding.name))
143
+ elif isinstance(value, exp.Count):
144
+ outputs.append(OutputColumn(alias or f"column_{index}", "count", None))
145
+ else:
146
+ outputs.append(OutputColumn(alias or f"column_{index}", "literal", None))
147
+ return outputs
148
+
149
+
150
+ def _resolve_order_column(
151
+ column: exp.Column,
152
+ bindings: Sequence[_Binding],
153
+ outputs: Sequence[OutputColumn],
154
+ ) -> SourceColumn:
155
+ if not column.table:
156
+ matches = [
157
+ output for output in outputs if output.label.casefold() == column.name.casefold()
158
+ ]
159
+ if matches:
160
+ # SQL Server resolves an unqualified ORDER BY name against output labels first.
161
+ if len(matches) != 1 or matches[0].source is None:
162
+ raise reject(Reason.ORDER_OUTPUT_AMBIGUOUS)
163
+ source = matches[0].source
164
+ binding = next(b for b in bindings if b.name == matches[0].binding)
165
+ return _rewrite(column, binding, source.column)
166
+ binding, canonical = _resolve(column, bindings)
167
+ return _rewrite(column, binding, canonical)
168
+
169
+
170
+ def _cross_check(query: exp.Select, schemas: Sequence[TableSchema], dialect: SqlDialect) -> None:
171
+ catalog: dict[str, dict] = {}
172
+ for table in schemas:
173
+ columns = dict.fromkeys(table.columns, "varchar")
174
+ if dialect.has_schema:
175
+ catalog.setdefault(table.schema, {})[table.name] = columns
176
+ else:
177
+ catalog[table.name] = columns
178
+ try:
179
+ qualify(
180
+ query.copy(),
181
+ schema=cast(dict[str, object], catalog),
182
+ dialect=f"{dialect.name}, normalization_strategy = case_sensitive",
183
+ validate_qualify_columns=True,
184
+ expand_stars=False,
185
+ )
186
+ except (OptimizeError, SqlglotError) as exc:
187
+ raise reject(Reason.QUALIFY_FAILED) from exc
188
+
189
+
190
+ def analyze_query(
191
+ query: exp.Select,
192
+ schemas: Sequence[TableSchema],
193
+ limits: ParserLimits,
194
+ dialect: SqlDialect,
195
+ ) -> AnalyzedQuery:
196
+ """Expand stars, resolve every column to a reflected source, and rebuild canonical names.
197
+
198
+ The input query is not modified. Unknown or ambiguous lineage raises QUERY_REJECTED.
199
+ """
200
+ working = query.copy()
201
+ bindings = _collect(working, schemas)
202
+ _expand_stars(working, bindings)
203
+ check_limits(working, limits)
204
+
205
+ refs: list[ColumnRef] = []
206
+ outputs = _analyze_projections(working, bindings, refs)
207
+
208
+ for join in working.args.get("joins") or []:
209
+ for column in join.args["on"].find_all(exp.Column):
210
+ binding, canonical = _resolve(column, bindings)
211
+ refs.append(ColumnRef(column, _rewrite(column, binding, canonical)))
212
+ where = working.args.get("where")
213
+ if where is not None:
214
+ for column in where.find_all(exp.Column):
215
+ binding, canonical = _resolve(column, bindings)
216
+ refs.append(ColumnRef(column, _rewrite(column, binding, canonical)))
217
+ order = working.args.get("order")
218
+ if order is not None:
219
+ for column in order.find_all(exp.Column):
220
+ refs.append(ColumnRef(column, _resolve_order_column(column, bindings, outputs)))
221
+
222
+ validate_allowlist(working)
223
+ _cross_check(working, schemas, dialect)
224
+ return AnalyzedQuery(working, tuple(outputs), tuple(refs))