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/__init__.py +3 -0
- sql_safe_mcp/__main__.py +43 -0
- sql_safe_mcp/config.py +201 -0
- sql_safe_mcp/db/__init__.py +1 -0
- sql_safe_mcp/db/extras.py +35 -0
- sql_safe_mcp/db/mysql.py +63 -0
- sql_safe_mcp/db/reflection.py +132 -0
- sql_safe_mcp/db/registry.py +123 -0
- sql_safe_mcp/db/sqlserver.py +65 -0
- sql_safe_mcp/errors.py +49 -0
- sql_safe_mcp/mcp_server.py +140 -0
- sql_safe_mcp/models.py +124 -0
- sql_safe_mcp/security/__init__.py +1 -0
- sql_safe_mcp/security/dialect.py +29 -0
- sql_safe_mcp/security/executor.py +129 -0
- sql_safe_mcp/security/lineage.py +224 -0
- sql_safe_mcp/security/parser.py +183 -0
- sql_safe_mcp/security/pipeline.py +43 -0
- sql_safe_mcp/security/policy.py +101 -0
- sql_safe_mcp/security/reasons.py +43 -0
- sql_safe_mcp/security/schema.py +139 -0
- sql_safe_mcp/security/tokens.py +207 -0
- sql_safe_mcp/security/validated_query.py +209 -0
- sql_safe_mcp/service.py +325 -0
- sql_safe_mcp-1.2.0.dist-info/METADATA +259 -0
- sql_safe_mcp-1.2.0.dist-info/RECORD +29 -0
- sql_safe_mcp-1.2.0.dist-info/WHEEL +4 -0
- sql_safe_mcp-1.2.0.dist-info/entry_points.txt +2 -0
- sql_safe_mcp-1.2.0.dist-info/licenses/LICENSE +21 -0
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))
|