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,207 @@
1
+ from __future__ import annotations
2
+
3
+ import base64
4
+ import json
5
+ import os
6
+ from collections.abc import Callable, Mapping
7
+ from datetime import date, datetime, time
8
+ from decimal import Decimal
9
+ from typing import Any, overload
10
+ from uuid import UUID
11
+
12
+ from cryptography.hazmat.primitives.ciphers.aead import AESGCM
13
+
14
+ from sql_safe_mcp.config import AppConfig
15
+ from sql_safe_mcp.errors import DomainError, ErrorCode
16
+
17
+ PREFIX = "pii:v1:"
18
+ MAX_TOKEN_CHARS = 4096
19
+ _NONCE_BYTES = 12
20
+ _KEY_BYTES = 32
21
+ _PAYLOAD_VERSION = 1
22
+
23
+
24
+ def _invalid_token() -> DomainError:
25
+ return DomainError(
26
+ ErrorCode.INVALID_PII_TOKEN,
27
+ "The PII token is invalid for this server.",
28
+ "Use a token returned by execute_sql for the same server alias.",
29
+ )
30
+
31
+
32
+ def _unsupported_value() -> DomainError:
33
+ return DomainError(
34
+ ErrorCode.QUERY_REJECTED,
35
+ "Query rejected: a protected column contains a value that cannot be tokenized.",
36
+ "Do not select this protected column.",
37
+ )
38
+
39
+
40
+ def _reject_constant(_: str) -> Any:
41
+ raise ValueError("non-finite JSON number")
42
+
43
+
44
+ def _encode_float(value: float) -> float:
45
+ if value != value or value in (float("inf"), float("-inf")):
46
+ raise ValueError("non-finite float")
47
+ return value
48
+
49
+
50
+ def _identity_check(_value: Any, _data: Any) -> bool:
51
+ return True
52
+
53
+
54
+ def _bytes_check(value: bytes, data: str) -> bool:
55
+ return base64.b64encode(value).decode("ascii") == data
56
+
57
+
58
+ def _finite_decimal(value: Decimal, data: str) -> bool:
59
+ return value.is_finite() and str(value) == data
60
+
61
+
62
+ # tag -> (python type, to JSON data, from JSON data)
63
+ _CODECS: dict[type, tuple[str, Callable[[Any], Any]]] = {
64
+ str: ("str", lambda v: v),
65
+ bool: ("bool", lambda v: v),
66
+ int: ("int", lambda v: v),
67
+ float: ("float", _encode_float),
68
+ Decimal: ("decimal", str),
69
+ date: ("date", date.isoformat),
70
+ datetime: ("datetime", datetime.isoformat),
71
+ time: ("time", time.isoformat),
72
+ UUID: ("uuid", str),
73
+ bytes: ("bytes", lambda v: base64.b64encode(v).decode("ascii")),
74
+ }
75
+ # tag -> (JSON type of `d`, parser, canonical-form check); interpreted at decode time by
76
+ # _decode_value so every branch runs (and can be verified) per call.
77
+ _DECODERS: dict[str, tuple[type, Callable[[Any], Any], Callable[[Any, Any], bool]]] = {
78
+ "str": (str, lambda d: d, _identity_check),
79
+ "bool": (bool, lambda d: d, _identity_check),
80
+ "int": (int, lambda d: d, _identity_check),
81
+ "float": (float, lambda d: d, _identity_check),
82
+ "decimal": (str, Decimal, _finite_decimal),
83
+ "date": (str, date.fromisoformat, lambda v, d: v.isoformat() == d),
84
+ "datetime": (str, datetime.fromisoformat, lambda v, d: v.isoformat() == d),
85
+ "time": (str, time.fromisoformat, lambda v, d: v.isoformat() == d),
86
+ "uuid": (str, UUID, lambda v, d: str(v) == d),
87
+ "bytes": (str, lambda d: base64.b64decode(d, validate=True), _bytes_check),
88
+ }
89
+
90
+
91
+ def _decode_value(tag: str, data: Any) -> Any:
92
+ kind, parse, check = _DECODERS[tag]
93
+ if type(data) is not kind:
94
+ raise ValueError("wrong payload type")
95
+ value = parse(data)
96
+ if not check(value, data):
97
+ raise ValueError("non-canonical payload")
98
+ return value
99
+
100
+
101
+ class TokenCodec:
102
+ """AES-256-GCM tokens for one server alias; the alias is authenticated as AAD."""
103
+
104
+ __slots__ = ("_aad", "_aead", "_alias")
105
+
106
+ def __init__(self, alias: str, key: bytes) -> None:
107
+ if len(key) != _KEY_BYTES:
108
+ raise ValueError(f"PII key must be {_KEY_BYTES} bytes")
109
+ self._alias = alias
110
+ self._aead = AESGCM(key)
111
+ self._aad = f"pii:v1\x00{alias}".encode()
112
+
113
+ def __repr__(self) -> str:
114
+ return f"TokenCodec(alias={self._alias!r})"
115
+
116
+ @property
117
+ def alias(self) -> str:
118
+ return self._alias
119
+
120
+ @overload
121
+ def encrypt(self, value: None) -> None: ...
122
+
123
+ @overload
124
+ def encrypt(self, value: object) -> str: ...
125
+
126
+ def encrypt(self, value: object) -> str | None:
127
+ """Tokenize a scalar; NULL stays None. Unsupported values fail closed."""
128
+ if value is None:
129
+ return None
130
+ entry = _CODECS.get(type(value))
131
+ if entry is None:
132
+ raise _unsupported_value()
133
+ tag, to_data = entry
134
+ try:
135
+ payload = json.dumps(
136
+ {"v": _PAYLOAD_VERSION, "t": tag, "d": to_data(value)},
137
+ separators=(",", ":"),
138
+ ensure_ascii=True,
139
+ allow_nan=False,
140
+ ).encode("ascii")
141
+ except (ValueError, OverflowError) as exc:
142
+ raise _unsupported_value() from exc
143
+ nonce = os.urandom(_NONCE_BYTES)
144
+ sealed = nonce + self._aead.encrypt(nonce, payload, self._aad)
145
+ token = PREFIX + base64.urlsafe_b64encode(sealed).decode("ascii").rstrip("=")
146
+ if len(token) > MAX_TOKEN_CHARS:
147
+ raise _unsupported_value()
148
+ return token
149
+
150
+ def decrypt(self, token: str) -> object:
151
+ """Return the value; every failure raises the same INVALID_PII_TOKEN error."""
152
+ try:
153
+ return self._decrypt(token)
154
+ except Exception:
155
+ raise _invalid_token() from None
156
+
157
+ def _decrypt(self, token: str) -> object:
158
+ if len(token) > MAX_TOKEN_CHARS or not token.startswith(PREFIX):
159
+ raise ValueError
160
+ body = token[len(PREFIX) :]
161
+ if not body.isascii():
162
+ raise ValueError
163
+ raw = base64.urlsafe_b64decode(body + "=" * (-len(body) % 4))
164
+ if base64.urlsafe_b64encode(raw).decode("ascii").rstrip("=") != body:
165
+ raise ValueError
166
+ if len(raw) < _NONCE_BYTES + 16:
167
+ raise ValueError
168
+ nonce, sealed = raw[:_NONCE_BYTES], raw[_NONCE_BYTES:]
169
+ payload = self._aead.decrypt(nonce, sealed, self._aad)
170
+ document = json.loads(payload.decode("utf-8"), parse_constant=_reject_constant)
171
+ if not isinstance(document, dict) or set(document) != {"v", "t", "d"}:
172
+ raise ValueError
173
+ version = document["v"]
174
+ if type(version) is not int or version != _PAYLOAD_VERSION:
175
+ raise ValueError
176
+ return _decode_value(document["t"], document["d"])
177
+
178
+
179
+ class TokenKeyRegistry:
180
+ """One codec per pii_safe alias; a codec never sees another alias's key."""
181
+
182
+ __slots__ = ("_codecs",)
183
+
184
+ def __init__(self, codecs: Mapping[str, TokenCodec]) -> None:
185
+ self._codecs = dict(codecs)
186
+
187
+ @classmethod
188
+ def from_config(cls, config: AppConfig) -> TokenKeyRegistry:
189
+ codecs: dict[str, TokenCodec] = {}
190
+ for alias, server in config.servers.items():
191
+ key = server.key_bytes()
192
+ if key is not None:
193
+ codecs[alias] = TokenCodec(alias, key)
194
+ return cls(codecs)
195
+
196
+ def __repr__(self) -> str:
197
+ return f"TokenKeyRegistry(aliases={sorted(self._codecs)!r})"
198
+
199
+ def codec_for(self, alias: str) -> TokenCodec:
200
+ try:
201
+ return self._codecs[alias]
202
+ except KeyError:
203
+ raise DomainError(
204
+ ErrorCode.ACCESS_LEVEL_DENIED,
205
+ "PII tokens are not available for this server.",
206
+ "Use a server configured with access_level pii_safe.",
207
+ ) from None
@@ -0,0 +1,209 @@
1
+ from __future__ import annotations
2
+
3
+ import re
4
+ import secrets
5
+ from dataclasses import dataclass
6
+ from typing import Any, Literal, NoReturn
7
+
8
+ from sqlglot import exp
9
+
10
+ from sql_safe_mcp.config import RuntimeConfig
11
+ from sql_safe_mcp.errors import DomainError, ErrorCode
12
+ from sql_safe_mcp.security.dialect import SQLSERVER, SqlDialect
13
+ from sql_safe_mcp.security.lineage import AnalyzedQuery, SourceColumn
14
+ from sql_safe_mcp.security.parser import (
15
+ ParserLimits,
16
+ check_limits,
17
+ reject,
18
+ validate_allowlist,
19
+ )
20
+ from sql_safe_mcp.security.policy import PolicyDecision
21
+ from sql_safe_mcp.security.reasons import Reason
22
+ from sql_safe_mcp.security.tokens import PREFIX
23
+
24
+ _SEAL = object()
25
+
26
+
27
+ @dataclass(frozen=True, slots=True)
28
+ class OutputPlan:
29
+ label: str
30
+ kind: Literal["column", "count", "literal"]
31
+ source: SourceColumn | None
32
+ protected: bool
33
+
34
+
35
+ class ValidatedQuery:
36
+ """SQL produced by the validation pipeline; the only input the executor accepts.
37
+
38
+ `sql` uses the dialect bind markers (`?` or `%s`) and `parameters` holds the matching decrypted
39
+ values, so no bind value is ever part of the statement text.
40
+
41
+ Instances are issued by `issue_validated_query` after the final AST check. There is no public
42
+ constructor, and the object cannot be copied or serialized. `repr` never shows SQL or binds.
43
+ """
44
+
45
+ __slots__ = ("_alias", "_ast", "_database", "_max_rows", "_outputs", "_parameters", "_sql")
46
+
47
+ def __init__(
48
+ self,
49
+ seal: object,
50
+ *,
51
+ alias: str,
52
+ database: str,
53
+ ast: exp.Select,
54
+ sql: str,
55
+ parameters: tuple[Any, ...],
56
+ outputs: tuple[OutputPlan, ...],
57
+ max_rows: int,
58
+ ) -> None:
59
+ if seal is not _SEAL:
60
+ raise TypeError("ValidatedQuery can only be issued by the validation pipeline")
61
+ object.__setattr__(self, "_alias", alias)
62
+ object.__setattr__(self, "_database", database)
63
+ object.__setattr__(self, "_ast", ast)
64
+ object.__setattr__(self, "_sql", sql)
65
+ object.__setattr__(self, "_parameters", tuple(parameters))
66
+ object.__setattr__(self, "_outputs", outputs)
67
+ object.__setattr__(self, "_max_rows", max_rows)
68
+
69
+ def __setattr__(self, name: str, value: object) -> NoReturn:
70
+ raise AttributeError("ValidatedQuery is read-only")
71
+
72
+ def __delattr__(self, name: str) -> NoReturn:
73
+ raise AttributeError("ValidatedQuery is read-only")
74
+
75
+ def __copy__(self) -> NoReturn:
76
+ raise TypeError("ValidatedQuery cannot be copied")
77
+
78
+ def __deepcopy__(self, memo: object) -> NoReturn:
79
+ raise TypeError("ValidatedQuery cannot be copied")
80
+
81
+ def __reduce__(self) -> NoReturn:
82
+ raise TypeError("ValidatedQuery cannot be serialized")
83
+
84
+ def __repr__(self) -> str:
85
+ return (
86
+ f"ValidatedQuery(alias={self._alias!r}, database={self._database!r}, "
87
+ f"outputs={len(self._outputs)}, parameters={len(self._parameters)})"
88
+ )
89
+
90
+ @property
91
+ def alias(self) -> str:
92
+ return self._alias
93
+
94
+ @property
95
+ def database(self) -> str:
96
+ return self._database
97
+
98
+ @property
99
+ def sql(self) -> str:
100
+ return self._sql
101
+
102
+ @property
103
+ def parameters(self) -> tuple[Any, ...]:
104
+ """Bind values in the order of the bind markers in `sql`."""
105
+ return self._parameters
106
+
107
+ @property
108
+ def outputs(self) -> tuple[OutputPlan, ...]:
109
+ return self._outputs
110
+
111
+ @property
112
+ def max_rows(self) -> int:
113
+ return self._max_rows
114
+
115
+ @property
116
+ def ast(self) -> exp.Select:
117
+ return self._ast.copy()
118
+
119
+
120
+ def _check_row_limit(max_rows: int, runtime: RuntimeConfig) -> None:
121
+ if max_rows < 1:
122
+ raise DomainError(
123
+ ErrorCode.INVALID_ARGUMENT,
124
+ "max_rows must be at least 1.",
125
+ f"Use a value between 1 and {runtime.hard_max_rows}.",
126
+ )
127
+ if max_rows > runtime.hard_max_rows:
128
+ raise DomainError(
129
+ ErrorCode.RESULT_LIMIT_EXCEEDED,
130
+ f"max_rows exceeds the configured hard limit of {runtime.hard_max_rows}.",
131
+ f"Use a value between 1 and {runtime.hard_max_rows}.",
132
+ )
133
+
134
+
135
+ def _cap_rows(query: exp.Select, max_rows: int) -> None:
136
+ fetch = max_rows + 1 # one extra row proves truncation without unbounded buffering
137
+ limit = query.args.get("limit")
138
+ requested = int(str(limit.expression.this)) if limit is not None else fetch
139
+ query.set("limit", exp.Limit(expression=exp.Literal.number(min(requested, fetch))))
140
+
141
+
142
+ def _to_positional(
143
+ query: exp.Select, parameters: dict[str, Any], dialect: SqlDialect
144
+ ) -> tuple[str, tuple[Any, ...]]:
145
+ """Render the dialect bind markers and order the values by their position in the generated text.
146
+
147
+ Each placeholder is generated as a random one-time marker first, so user literals containing
148
+ colons or marker-like text cannot be mistaken for binds.
149
+ """
150
+ nonce = secrets.token_hex(8)
151
+ rendered = query.copy()
152
+ markers: dict[str, str] = {}
153
+ for placeholder in list(rendered.find_all(exp.Placeholder)):
154
+ marker = f"__bind_{nonce}_{len(markers)}__"
155
+ markers[marker] = str(placeholder.this)
156
+ placeholder.replace(exp.Var(this=marker))
157
+ sql = rendered.sql(dialect=dialect.name)
158
+ if dialect.escape_percent:
159
+ sql = sql.replace("%", "%%")
160
+ found = re.findall(rf"__bind_{nonce}_\d+__", sql)
161
+ if len(found) != len(markers) or set(found) != set(markers):
162
+ raise reject(Reason.BIND_UNSAFE)
163
+ for marker in found:
164
+ sql = sql.replace(marker, dialect.bind_marker, 1)
165
+ return sql, tuple(parameters[markers[marker]] for marker in found)
166
+
167
+
168
+ def issue_validated_query(
169
+ analyzed: AnalyzedQuery,
170
+ decision: PolicyDecision,
171
+ *,
172
+ alias: str,
173
+ database: str,
174
+ max_rows: int,
175
+ runtime: RuntimeConfig,
176
+ limits: ParserLimits,
177
+ dialect: SqlDialect = SQLSERVER,
178
+ ) -> ValidatedQuery:
179
+ """Rewrite tokens to binds, cap rows, re-validate, and generate SQL from the final AST.
180
+
181
+ Consumes `analyzed.query`: token literals are replaced in place.
182
+ """
183
+ _check_row_limit(max_rows, runtime)
184
+ query = analyzed.query
185
+ parameters: dict[str, Any] = {}
186
+ for site in decision.token_sites:
187
+ name = f"pii_{len(parameters)}"
188
+ parameters[name] = site.value
189
+ site.literal.replace(exp.Placeholder(this=name))
190
+ _cap_rows(query, max_rows)
191
+ check_limits(query, limits)
192
+ validate_allowlist(query, allow_placeholders=True)
193
+ sql, ordered = _to_positional(query, parameters, dialect)
194
+ if PREFIX in sql:
195
+ raise reject(Reason.TOKEN_NOT_REPLACED)
196
+ outputs = tuple(
197
+ OutputPlan(output.label, output.kind, output.source, protected)
198
+ for output, protected in zip(analyzed.outputs, decision.protected, strict=True)
199
+ )
200
+ return ValidatedQuery(
201
+ _SEAL,
202
+ alias=alias,
203
+ database=database,
204
+ ast=query,
205
+ sql=sql,
206
+ parameters=ordered,
207
+ outputs=outputs,
208
+ max_rows=max_rows,
209
+ )