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.
@@ -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)