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
|
@@ -0,0 +1,183 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from dataclasses import dataclass
|
|
4
|
+
|
|
5
|
+
import sqlglot
|
|
6
|
+
from sqlglot import exp
|
|
7
|
+
|
|
8
|
+
from sql_safe_mcp.config import RuntimeConfig
|
|
9
|
+
from sql_safe_mcp.errors import DomainError, ErrorCode
|
|
10
|
+
from sql_safe_mcp.security.dialect import SQLSERVER, SqlDialect
|
|
11
|
+
from sql_safe_mcp.security.reasons import Reason
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
@dataclass(frozen=True, slots=True)
|
|
15
|
+
class ParserLimits:
|
|
16
|
+
max_sql_chars: int
|
|
17
|
+
max_ast_nodes: int
|
|
18
|
+
max_joins: int
|
|
19
|
+
max_in_list_items: int
|
|
20
|
+
|
|
21
|
+
@classmethod
|
|
22
|
+
def from_runtime(cls, runtime: RuntimeConfig) -> ParserLimits:
|
|
23
|
+
return cls(
|
|
24
|
+
runtime.max_sql_chars,
|
|
25
|
+
runtime.max_ast_nodes,
|
|
26
|
+
runtime.max_joins,
|
|
27
|
+
runtime.max_in_list_items,
|
|
28
|
+
)
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
REJECT_HINT = "Rewrite the operation as a single, simpler SELECT query."
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
def reject(reason: Reason, detail: str | None = None) -> DomainError:
|
|
35
|
+
"""Build the QUERY_REJECTED error; only a fixed Reason (plus an optional name) is accepted."""
|
|
36
|
+
if not isinstance(reason, Reason):
|
|
37
|
+
raise TypeError("reject() takes a Reason")
|
|
38
|
+
text = str(reason) if detail is None else f"{reason}: {detail}"
|
|
39
|
+
return DomainError(ErrorCode.QUERY_REJECTED, f"Query rejected: {text}.", REJECT_HINT)
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
def check_limits(query: exp.Expression, limits: ParserLimits) -> None:
|
|
43
|
+
"""Enforce AST size limits; also used again after transformations."""
|
|
44
|
+
for count, _ in enumerate(query.walk(), start=1):
|
|
45
|
+
if count > limits.max_ast_nodes:
|
|
46
|
+
raise reject(Reason.LIMIT_NODES)
|
|
47
|
+
if sum(1 for _ in query.find_all(exp.Join)) > limits.max_joins:
|
|
48
|
+
raise reject(Reason.LIMIT_JOINS)
|
|
49
|
+
for node in query.find_all(exp.In):
|
|
50
|
+
if len(node.expressions) > limits.max_in_list_items:
|
|
51
|
+
raise reject(Reason.LIMIT_IN_LIST)
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
_COMPARISONS = (exp.EQ, exp.NEQ, exp.GT, exp.GTE, exp.LT, exp.LTE)
|
|
55
|
+
|
|
56
|
+
# node type -> argument names that may be non-empty; everything else is rejected.
|
|
57
|
+
_ALLOWED_ARGS: dict[type[exp.Expr], frozenset[str]] = {
|
|
58
|
+
exp.Select: frozenset({"expressions", "from_", "joins", "where", "order", "limit"}),
|
|
59
|
+
exp.From: frozenset({"this"}),
|
|
60
|
+
exp.Table: frozenset({"this", "db", "alias"}),
|
|
61
|
+
exp.TableAlias: frozenset({"this"}),
|
|
62
|
+
exp.Identifier: frozenset({"this", "quoted"}),
|
|
63
|
+
exp.Column: frozenset({"this", "table"}),
|
|
64
|
+
exp.Star: frozenset(),
|
|
65
|
+
exp.Count: frozenset({"this", "big_int"}), # big_int: a flag the mysql parser sets
|
|
66
|
+
exp.Literal: frozenset({"this", "is_string"}),
|
|
67
|
+
exp.Null: frozenset(),
|
|
68
|
+
exp.Neg: frozenset({"this"}),
|
|
69
|
+
exp.Alias: frozenset({"this", "alias"}),
|
|
70
|
+
exp.Join: frozenset({"this", "on", "side", "kind"}),
|
|
71
|
+
exp.And: frozenset({"this", "expression"}),
|
|
72
|
+
exp.Or: frozenset({"this", "expression"}),
|
|
73
|
+
exp.Paren: frozenset({"this"}),
|
|
74
|
+
exp.In: frozenset({"this", "expressions"}),
|
|
75
|
+
exp.Is: frozenset({"this", "expression"}),
|
|
76
|
+
exp.Where: frozenset({"this"}),
|
|
77
|
+
exp.Order: frozenset({"expressions"}),
|
|
78
|
+
exp.Ordered: frozenset({"this", "desc", "nulls_first"}),
|
|
79
|
+
exp.Limit: frozenset({"expression"}),
|
|
80
|
+
**{comparison: frozenset({"this", "expression"}) for comparison in _COMPARISONS},
|
|
81
|
+
}
|
|
82
|
+
_PLACEHOLDER_ARGS = frozenset({"this"})
|
|
83
|
+
_PROJECTIONS = (exp.Column, exp.Star, exp.Literal, exp.Null, exp.Neg, exp.Alias, exp.Count)
|
|
84
|
+
|
|
85
|
+
|
|
86
|
+
def _is_numeric_literal(node: object) -> bool:
|
|
87
|
+
return isinstance(node, exp.Literal) and not node.is_string
|
|
88
|
+
|
|
89
|
+
|
|
90
|
+
def _is_empty(value: object) -> bool:
|
|
91
|
+
return value is None or value is False or value == []
|
|
92
|
+
|
|
93
|
+
|
|
94
|
+
def _check_node(node: exp.Expr, allow_placeholders: bool) -> None:
|
|
95
|
+
kind = type(node)
|
|
96
|
+
if kind is exp.Placeholder and allow_placeholders:
|
|
97
|
+
allowed_args = _PLACEHOLDER_ARGS
|
|
98
|
+
elif kind in _ALLOWED_ARGS:
|
|
99
|
+
allowed_args = _ALLOWED_ARGS[kind]
|
|
100
|
+
else:
|
|
101
|
+
raise reject(Reason.NODE_UNSUPPORTED, kind.__name__)
|
|
102
|
+
for name, value in node.args.items():
|
|
103
|
+
if name not in allowed_args and not _is_empty(value):
|
|
104
|
+
raise reject(Reason.OPTION_UNSUPPORTED, f"{kind.__name__}.{name}")
|
|
105
|
+
|
|
106
|
+
if isinstance(node, exp.Table):
|
|
107
|
+
name = node.this
|
|
108
|
+
if not isinstance(name, exp.Identifier) or str(name.this).startswith(("#", "@")):
|
|
109
|
+
raise reject(Reason.TABLE_UNSUPPORTED)
|
|
110
|
+
elif isinstance(node, exp.Column):
|
|
111
|
+
if not isinstance(node.this, (exp.Identifier, exp.Star)):
|
|
112
|
+
raise reject(Reason.COLUMN_REFERENCE_UNSUPPORTED)
|
|
113
|
+
elif isinstance(node, exp.Count):
|
|
114
|
+
if not isinstance(node.this, exp.Star):
|
|
115
|
+
raise reject(Reason.COUNT_STAR_ONLY)
|
|
116
|
+
elif isinstance(node, exp.Star):
|
|
117
|
+
if not isinstance(node.parent, (exp.Select, exp.Column, exp.Count)):
|
|
118
|
+
raise reject(Reason.STAR_PROJECTION_ONLY)
|
|
119
|
+
elif isinstance(node, exp.Join):
|
|
120
|
+
if not isinstance(node.this, exp.Table) or node.args.get("on") is None:
|
|
121
|
+
raise reject(Reason.JOIN_NEEDS_ON)
|
|
122
|
+
if node.side not in {"", "LEFT"} or node.kind not in {"", "INNER"}:
|
|
123
|
+
raise reject(Reason.JOIN_TYPE_UNSUPPORTED)
|
|
124
|
+
elif isinstance(node, exp.Is):
|
|
125
|
+
if not isinstance(node.expression, exp.Null):
|
|
126
|
+
raise reject(Reason.IS_NULL_ONLY)
|
|
127
|
+
elif isinstance(node, exp.Neg):
|
|
128
|
+
if not _is_numeric_literal(node.this):
|
|
129
|
+
raise reject(Reason.NEGATION_NUMERIC_ONLY)
|
|
130
|
+
elif isinstance(node, exp.Limit):
|
|
131
|
+
value = node.expression
|
|
132
|
+
if not _is_numeric_literal(value) or not str(value.this).isdigit():
|
|
133
|
+
raise reject(Reason.TOP_INTEGER_ONLY)
|
|
134
|
+
elif isinstance(node, exp.Ordered):
|
|
135
|
+
if not isinstance(node.this, exp.Column) or isinstance(node.this.this, exp.Star):
|
|
136
|
+
raise reject(Reason.ORDER_DIRECT_COLUMNS_ONLY)
|
|
137
|
+
elif isinstance(node, exp.In):
|
|
138
|
+
if not isinstance(node.this, exp.Column) or not all(
|
|
139
|
+
isinstance(item, (exp.Literal, exp.Neg, exp.Placeholder)) for item in node.expressions
|
|
140
|
+
):
|
|
141
|
+
raise reject(Reason.IN_LITERALS_ONLY)
|
|
142
|
+
elif isinstance(node, _COMPARISONS):
|
|
143
|
+
for side in (node.this, node.expression):
|
|
144
|
+
if isinstance(side, exp.Star):
|
|
145
|
+
raise reject(Reason.STAR_COMPARISON)
|
|
146
|
+
|
|
147
|
+
|
|
148
|
+
def validate_allowlist(query: exp.Expression, *, allow_placeholders: bool = False) -> None:
|
|
149
|
+
"""Fail closed on any node type, argument, or shape outside the supported subset."""
|
|
150
|
+
if not isinstance(query, exp.Select):
|
|
151
|
+
raise reject(Reason.SELECT_ONLY)
|
|
152
|
+
for node in query.walk():
|
|
153
|
+
_check_node(node, allow_placeholders)
|
|
154
|
+
for projection in query.expressions:
|
|
155
|
+
if not isinstance(projection, _PROJECTIONS):
|
|
156
|
+
raise reject(Reason.PROJECTION_UNSUPPORTED)
|
|
157
|
+
|
|
158
|
+
|
|
159
|
+
def strip_comments(query: exp.Expression) -> None:
|
|
160
|
+
"""Drop caller comments so no caller-controlled text is ever emitted into generated SQL."""
|
|
161
|
+
for node in query.walk():
|
|
162
|
+
node.pop_comments()
|
|
163
|
+
|
|
164
|
+
|
|
165
|
+
def parse_select(sql: str, limits: ParserLimits, dialect: SqlDialect = SQLSERVER) -> exp.Select:
|
|
166
|
+
"""Parse exactly one root SELECT; any failure becomes QUERY_REJECTED."""
|
|
167
|
+
if len(sql) > limits.max_sql_chars:
|
|
168
|
+
raise reject(Reason.LIMIT_CHARS)
|
|
169
|
+
try:
|
|
170
|
+
statements = [s for s in sqlglot.parse(sql, dialect=dialect.name) if s is not None]
|
|
171
|
+
if len(statements) != 1:
|
|
172
|
+
raise reject(Reason.ONE_STATEMENT)
|
|
173
|
+
query = statements[0]
|
|
174
|
+
if not isinstance(query, exp.Select):
|
|
175
|
+
raise reject(Reason.SELECT_ONLY)
|
|
176
|
+
check_limits(query, limits)
|
|
177
|
+
validate_allowlist(query)
|
|
178
|
+
strip_comments(query)
|
|
179
|
+
except DomainError:
|
|
180
|
+
raise
|
|
181
|
+
except Exception as exc:
|
|
182
|
+
raise reject(Reason.PARSE_FAILED) from exc
|
|
183
|
+
return query
|
|
@@ -0,0 +1,43 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from sql_safe_mcp.config import PiiConfig, RuntimeConfig
|
|
4
|
+
from sql_safe_mcp.security.dialect import SqlDialect
|
|
5
|
+
from sql_safe_mcp.security.lineage import analyze_query
|
|
6
|
+
from sql_safe_mcp.security.parser import ParserLimits, parse_select
|
|
7
|
+
from sql_safe_mcp.security.policy import PiiPolicy
|
|
8
|
+
from sql_safe_mcp.security.schema import TableCatalog, resolve_tables
|
|
9
|
+
from sql_safe_mcp.security.tokens import TokenCodec
|
|
10
|
+
from sql_safe_mcp.security.validated_query import ValidatedQuery, issue_validated_query
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def validate_sql(
|
|
14
|
+
sql: str,
|
|
15
|
+
*,
|
|
16
|
+
alias: str,
|
|
17
|
+
database: str,
|
|
18
|
+
catalog: TableCatalog,
|
|
19
|
+
pii_config: PiiConfig,
|
|
20
|
+
codec: TokenCodec,
|
|
21
|
+
runtime: RuntimeConfig,
|
|
22
|
+
max_rows: int,
|
|
23
|
+
dialect: SqlDialect,
|
|
24
|
+
) -> ValidatedQuery:
|
|
25
|
+
"""Run caller SQL through parse, allowlist, reflection, lineage, and PII policy.
|
|
26
|
+
|
|
27
|
+
Any rejection raises a DomainError before the catalog or database is used further.
|
|
28
|
+
"""
|
|
29
|
+
limits = ParserLimits.from_runtime(runtime)
|
|
30
|
+
query = parse_select(sql, limits, dialect)
|
|
31
|
+
tables = resolve_tables(query, catalog, dialect)
|
|
32
|
+
analyzed = analyze_query(query, tables, limits, dialect)
|
|
33
|
+
decision = PiiPolicy(database, pii_config, codec).evaluate(analyzed)
|
|
34
|
+
return issue_validated_query(
|
|
35
|
+
analyzed,
|
|
36
|
+
decision,
|
|
37
|
+
alias=alias,
|
|
38
|
+
database=database,
|
|
39
|
+
max_rows=max_rows,
|
|
40
|
+
runtime=runtime,
|
|
41
|
+
limits=limits,
|
|
42
|
+
dialect=dialect,
|
|
43
|
+
)
|
|
@@ -0,0 +1,101 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from dataclasses import dataclass, field
|
|
4
|
+
|
|
5
|
+
from sqlglot import exp
|
|
6
|
+
|
|
7
|
+
from sql_safe_mcp.config import PiiConfig
|
|
8
|
+
from sql_safe_mcp.security.lineage import AnalyzedQuery, ColumnRef, SourceColumn
|
|
9
|
+
from sql_safe_mcp.security.parser import reject
|
|
10
|
+
from sql_safe_mcp.security.reasons import Reason
|
|
11
|
+
from sql_safe_mcp.security.tokens import PREFIX, TokenCodec
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
@dataclass(frozen=True, slots=True)
|
|
15
|
+
class TokenSite:
|
|
16
|
+
"""A token literal in an allowed predicate position and its decrypted bind value."""
|
|
17
|
+
|
|
18
|
+
literal: exp.Literal
|
|
19
|
+
value: object = field(repr=False)
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
@dataclass(frozen=True, slots=True)
|
|
23
|
+
class PolicyDecision:
|
|
24
|
+
protected: tuple[bool, ...]
|
|
25
|
+
token_sites: tuple[TokenSite, ...]
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
class PiiPolicy:
|
|
29
|
+
"""Decides protected-column usage from source lineage and configured rules only."""
|
|
30
|
+
|
|
31
|
+
def __init__(self, database: str, config: PiiConfig, codec: TokenCodec) -> None:
|
|
32
|
+
self._database = database.casefold()
|
|
33
|
+
self._config = config
|
|
34
|
+
self._codec = codec
|
|
35
|
+
|
|
36
|
+
def is_protected(self, source: SourceColumn) -> bool:
|
|
37
|
+
schema = source.schema.casefold()
|
|
38
|
+
table = source.table.casefold()
|
|
39
|
+
column = source.column.casefold()
|
|
40
|
+
for rule in self._config.rules:
|
|
41
|
+
if rule.database != "*" and rule.database.casefold() != self._database:
|
|
42
|
+
continue
|
|
43
|
+
if rule.schema_ is not None and rule.schema_.casefold() != schema:
|
|
44
|
+
continue
|
|
45
|
+
if rule.table.casefold() == table and column in {c.casefold() for c in rule.columns}:
|
|
46
|
+
return True
|
|
47
|
+
return False
|
|
48
|
+
|
|
49
|
+
def evaluate(self, analyzed: AnalyzedQuery) -> PolicyDecision:
|
|
50
|
+
sites: list[TokenSite] = []
|
|
51
|
+
for ref in analyzed.references:
|
|
52
|
+
if self.is_protected(ref.source):
|
|
53
|
+
sites.extend(self._check_protected(ref))
|
|
54
|
+
consumed = {id(site.literal) for site in sites}
|
|
55
|
+
for literal in analyzed.query.find_all(exp.Literal):
|
|
56
|
+
if (
|
|
57
|
+
literal.is_string
|
|
58
|
+
and str(literal.this).startswith(PREFIX)
|
|
59
|
+
and id(literal) not in consumed
|
|
60
|
+
):
|
|
61
|
+
raise reject(Reason.TOKEN_OUTSIDE_PROTECTED)
|
|
62
|
+
protected = tuple(
|
|
63
|
+
output.source is not None and self.is_protected(output.source)
|
|
64
|
+
for output in analyzed.outputs
|
|
65
|
+
)
|
|
66
|
+
return PolicyDecision(protected, tuple(sites))
|
|
67
|
+
|
|
68
|
+
def _check_protected(self, ref: ColumnRef) -> list[TokenSite]:
|
|
69
|
+
column = ref.node
|
|
70
|
+
parent = column.parent
|
|
71
|
+
if isinstance(parent, exp.Select):
|
|
72
|
+
return []
|
|
73
|
+
if isinstance(parent, exp.Alias) and isinstance(parent.parent, exp.Select):
|
|
74
|
+
return []
|
|
75
|
+
if not _in_where(column):
|
|
76
|
+
raise reject(Reason.PROTECTED_POSITION)
|
|
77
|
+
if isinstance(parent, exp.EQ) and parent.this is column:
|
|
78
|
+
return [self._token_site(parent.expression)]
|
|
79
|
+
if isinstance(parent, exp.In) and parent.this is column and parent.expressions:
|
|
80
|
+
return [self._token_site(item) for item in parent.expressions]
|
|
81
|
+
raise reject(Reason.PROTECTED_POSITION)
|
|
82
|
+
|
|
83
|
+
def _token_site(self, node: exp.Expression) -> TokenSite:
|
|
84
|
+
if (
|
|
85
|
+
not isinstance(node, exp.Literal)
|
|
86
|
+
or not node.is_string
|
|
87
|
+
or not str(node.this).startswith(PREFIX)
|
|
88
|
+
):
|
|
89
|
+
raise reject(Reason.PROTECTED_NEEDS_TOKEN)
|
|
90
|
+
return TokenSite(node, self._codec.decrypt(str(node.this)))
|
|
91
|
+
|
|
92
|
+
|
|
93
|
+
def _in_where(node: exp.Expression) -> bool:
|
|
94
|
+
current = node.parent
|
|
95
|
+
while current is not None:
|
|
96
|
+
if isinstance(current, exp.Where):
|
|
97
|
+
return True
|
|
98
|
+
if isinstance(current, exp.Join | exp.Order | exp.Select):
|
|
99
|
+
return False
|
|
100
|
+
current = current.parent
|
|
101
|
+
return False
|
|
@@ -0,0 +1,43 @@
|
|
|
1
|
+
"""Fixed, reviewable reasons for QUERY_REJECTED. Free text never reaches reject()."""
|
|
2
|
+
|
|
3
|
+
from enum import StrEnum
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
class Reason(StrEnum):
|
|
7
|
+
TABLE_RESOLUTION_MISMATCH = "table resolution does not match the query"
|
|
8
|
+
DUPLICATE_BINDING = "table names and aliases must be unique"
|
|
9
|
+
QUALIFIER_UNKNOWN = "a column qualifier does not match a table in the query"
|
|
10
|
+
COLUMN_NOT_FOUND = "a referenced column was not found or is ambiguous"
|
|
11
|
+
COLUMN_AMBIGUOUS = "a referenced column was not found or is ambiguous; qualify it"
|
|
12
|
+
ORDER_OUTPUT_AMBIGUOUS = "ORDER BY refers to an ambiguous or non-column output"
|
|
13
|
+
QUALIFY_FAILED = "the query references unknown or ambiguous columns"
|
|
14
|
+
LIMIT_NODES = "the query exceeds the configured complexity limit"
|
|
15
|
+
LIMIT_JOINS = "the query has too many joins"
|
|
16
|
+
LIMIT_IN_LIST = "an IN list exceeds the configured item limit"
|
|
17
|
+
TABLE_UNSUPPORTED = "only base tables of the current database are supported"
|
|
18
|
+
COLUMN_REFERENCE_UNSUPPORTED = "unsupported column reference"
|
|
19
|
+
COUNT_STAR_ONLY = "only COUNT(*) is supported"
|
|
20
|
+
STAR_PROJECTION_ONLY = "* is supported only as a projection"
|
|
21
|
+
JOIN_NEEDS_ON = "JOIN requires a table and an explicit ON expression"
|
|
22
|
+
JOIN_TYPE_UNSUPPORTED = "only INNER JOIN and LEFT JOIN are supported"
|
|
23
|
+
IS_NULL_ONLY = "IS is supported only with NULL"
|
|
24
|
+
NEGATION_NUMERIC_ONLY = "negation is supported only for numeric literals"
|
|
25
|
+
TOP_INTEGER_ONLY = "TOP/LIMIT must be a non-negative integer literal"
|
|
26
|
+
ORDER_DIRECT_COLUMNS_ONLY = "ORDER BY supports only direct columns"
|
|
27
|
+
IN_LITERALS_ONLY = "IN supports a column and a list of literals"
|
|
28
|
+
STAR_COMPARISON = "* cannot be compared"
|
|
29
|
+
SELECT_ONLY = "only SELECT statements are supported"
|
|
30
|
+
PROJECTION_UNSUPPORTED = "unsupported projection"
|
|
31
|
+
LIMIT_CHARS = "the SQL exceeds the configured length limit"
|
|
32
|
+
ONE_STATEMENT = "exactly one statement is required"
|
|
33
|
+
PARSE_FAILED = "the SQL could not be parsed"
|
|
34
|
+
TOKEN_OUTSIDE_PROTECTED = "a PII token can only be compared with a protected column"
|
|
35
|
+
PROTECTED_POSITION = "a protected column can only be projected or compared with a token"
|
|
36
|
+
PROTECTED_NEEDS_TOKEN = "a protected column can only be compared with a PII token"
|
|
37
|
+
TABLE_GONE = "a referenced table no longer exists"
|
|
38
|
+
TABLE_NOT_FOUND = "a referenced table was not found or is ambiguous; qualify it with a schema"
|
|
39
|
+
TABLE_SCOPE = "only tables of the current database are supported"
|
|
40
|
+
BIND_UNSAFE = "bind parameters could not be generated safely"
|
|
41
|
+
TOKEN_NOT_REPLACED = "a PII token could not be replaced by a bind parameter"
|
|
42
|
+
NODE_UNSUPPORTED = "unsupported construct"
|
|
43
|
+
OPTION_UNSUPPORTED = "unsupported option"
|
|
@@ -0,0 +1,139 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import time
|
|
4
|
+
from collections import OrderedDict
|
|
5
|
+
from collections.abc import Callable, Hashable, Sequence
|
|
6
|
+
from dataclasses import dataclass
|
|
7
|
+
from threading import Lock
|
|
8
|
+
from typing import Protocol
|
|
9
|
+
|
|
10
|
+
from sqlalchemy import Connection, inspect
|
|
11
|
+
from sqlalchemy.exc import NoSuchTableError
|
|
12
|
+
from sqlglot import exp
|
|
13
|
+
|
|
14
|
+
from sql_safe_mcp.db.reflection import list_tables as reflect_tables
|
|
15
|
+
from sql_safe_mcp.security.dialect import SQLSERVER, SqlDialect
|
|
16
|
+
from sql_safe_mcp.security.parser import reject
|
|
17
|
+
from sql_safe_mcp.security.reasons import Reason
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
@dataclass(frozen=True, slots=True)
|
|
21
|
+
class TableSchema:
|
|
22
|
+
schema: str
|
|
23
|
+
name: str
|
|
24
|
+
columns: tuple[str, ...]
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
class TableCatalog(Protocol):
|
|
28
|
+
"""Reflected user base tables of one alias and database."""
|
|
29
|
+
|
|
30
|
+
def list_tables(self) -> Sequence[tuple[str, str]]: ...
|
|
31
|
+
|
|
32
|
+
def columns(self, schema: str, table: str) -> Sequence[str]: ...
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
class SchemaCache:
|
|
36
|
+
"""Small thread-safe LRU with a time-to-live for immutable reflection results."""
|
|
37
|
+
|
|
38
|
+
def __init__(
|
|
39
|
+
self,
|
|
40
|
+
max_entries: int,
|
|
41
|
+
ttl_seconds: float,
|
|
42
|
+
clock: Callable[[], float] = time.monotonic,
|
|
43
|
+
) -> None:
|
|
44
|
+
self._max_entries = max_entries
|
|
45
|
+
self._ttl_seconds = ttl_seconds
|
|
46
|
+
self._clock = clock
|
|
47
|
+
self._entries: OrderedDict[Hashable, tuple[float, tuple]] = OrderedDict()
|
|
48
|
+
self._lock = Lock()
|
|
49
|
+
|
|
50
|
+
def get(self, key: Hashable) -> tuple | None:
|
|
51
|
+
with self._lock:
|
|
52
|
+
entry = self._entries.get(key)
|
|
53
|
+
if entry is None:
|
|
54
|
+
return None
|
|
55
|
+
stored_at, value = entry
|
|
56
|
+
if self._clock() - stored_at >= self._ttl_seconds:
|
|
57
|
+
del self._entries[key]
|
|
58
|
+
return None
|
|
59
|
+
self._entries.move_to_end(key)
|
|
60
|
+
return value
|
|
61
|
+
|
|
62
|
+
def put(self, key: Hashable, value: tuple) -> None:
|
|
63
|
+
with self._lock:
|
|
64
|
+
self._entries[key] = (self._clock(), value)
|
|
65
|
+
self._entries.move_to_end(key)
|
|
66
|
+
while len(self._entries) > self._max_entries:
|
|
67
|
+
self._entries.popitem(last=False)
|
|
68
|
+
|
|
69
|
+
|
|
70
|
+
class ReflectedCatalog:
|
|
71
|
+
"""Inspector-backed catalog; cache keys always include server alias and database."""
|
|
72
|
+
|
|
73
|
+
def __init__(
|
|
74
|
+
self, connection: Connection, cache: SchemaCache, alias: str, database: str
|
|
75
|
+
) -> None:
|
|
76
|
+
self._connection = connection
|
|
77
|
+
self._cache = cache
|
|
78
|
+
self._scope = (alias, database)
|
|
79
|
+
|
|
80
|
+
def list_tables(self) -> list[tuple[str, str]]:
|
|
81
|
+
key = (*self._scope, "tables")
|
|
82
|
+
cached = self._cache.get(key)
|
|
83
|
+
if cached is None:
|
|
84
|
+
cached = tuple(
|
|
85
|
+
(table.schema_ or "", table.name) for table in reflect_tables(self._connection)
|
|
86
|
+
)
|
|
87
|
+
self._cache.put(key, cached)
|
|
88
|
+
return list(cached)
|
|
89
|
+
|
|
90
|
+
def columns(self, schema: str, table: str) -> tuple[str, ...]:
|
|
91
|
+
key = (*self._scope, "columns", schema, table)
|
|
92
|
+
cached = self._cache.get(key)
|
|
93
|
+
if cached is None:
|
|
94
|
+
try:
|
|
95
|
+
reflected = inspect(self._connection).get_columns(table, schema=schema or None)
|
|
96
|
+
except NoSuchTableError as exc:
|
|
97
|
+
raise reject(Reason.TABLE_GONE) from exc
|
|
98
|
+
cached = tuple(str(column["name"]) for column in reflected)
|
|
99
|
+
self._cache.put(key, cached)
|
|
100
|
+
return cached
|
|
101
|
+
|
|
102
|
+
|
|
103
|
+
def _match(tables: Sequence[tuple[str, str]], schema: str | None, name: str) -> tuple[str, str]:
|
|
104
|
+
folded_name = name.casefold()
|
|
105
|
+
folded_schema = None if schema is None else schema.casefold()
|
|
106
|
+
matches = [
|
|
107
|
+
(table_schema, table_name)
|
|
108
|
+
for table_schema, table_name in tables
|
|
109
|
+
if table_name.casefold() == folded_name
|
|
110
|
+
and (folded_schema is None or table_schema.casefold() == folded_schema)
|
|
111
|
+
]
|
|
112
|
+
if len(matches) != 1:
|
|
113
|
+
raise reject(Reason.TABLE_NOT_FOUND)
|
|
114
|
+
return matches[0]
|
|
115
|
+
|
|
116
|
+
|
|
117
|
+
def resolve_tables(
|
|
118
|
+
query: exp.Select, catalog: TableCatalog, dialect: SqlDialect = SQLSERVER
|
|
119
|
+
) -> tuple[TableSchema, ...]:
|
|
120
|
+
"""Resolve every raw-AST table against reflection and rewrite it to its reflected name.
|
|
121
|
+
|
|
122
|
+
Raises QUERY_REJECTED for missing, ambiguous, system, and cross-database objects.
|
|
123
|
+
"""
|
|
124
|
+
available = list(catalog.list_tables())
|
|
125
|
+
resolved: list[TableSchema] = []
|
|
126
|
+
for table in query.find_all(exp.Table):
|
|
127
|
+
if table.args.get("catalog") or not isinstance(table.this, exp.Identifier):
|
|
128
|
+
raise reject(Reason.TABLE_SCOPE)
|
|
129
|
+
if not dialect.has_schema and table.db:
|
|
130
|
+
raise reject(Reason.TABLE_SCOPE) # db.table is a cross-database reference
|
|
131
|
+
schema = table.db or None
|
|
132
|
+
schema_name, table_name = _match(available, schema, table.name)
|
|
133
|
+
resolved.append(
|
|
134
|
+
TableSchema(schema_name, table_name, tuple(catalog.columns(schema_name, table_name)))
|
|
135
|
+
)
|
|
136
|
+
table.set("this", exp.to_identifier(table_name, quoted=True))
|
|
137
|
+
if dialect.has_schema:
|
|
138
|
+
table.set("db", exp.to_identifier(schema_name, quoted=True))
|
|
139
|
+
return tuple(resolved)
|