Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
57 changes: 44 additions & 13 deletions src/postgres_mcp/sql/safe_sql.py
Original file line number Diff line number Diff line change
Expand Up @@ -875,8 +875,13 @@ def __init__(self, sql_driver: SqlDriver, timeout: float | None = None):
self.sql_driver = sql_driver
self.timeout = timeout

def _validate_node(self, node: Node) -> None:
"""Recursively validate a node and all its children"""
def _validate_node(self, node: Node, unknown_funcs: set[str]) -> None:
"""Recursively validate a node and all its children.

Function calls outside ALLOWED_FUNCTIONS are collected into
unknown_funcs rather than rejected here; execute_query() then allows
them only if pg_proc proves every overload non-volatile.
"""
# Check if node type is allowed
if not isinstance(node, tuple(self.ALLOWED_NODE_TYPES)):
raise ValueError(f"Node type {type(node)} is not allowed")
Expand All @@ -900,17 +905,17 @@ def _validate_node(self, node: Node) -> None:
match = self.PG_CATALOG_PATTERN.match(func_name)
unqualified_name = match.group(1) if match else func_name
if unqualified_name not in self.ALLOWED_FUNCTIONS:
raise ValueError(f"Function {func_name} is not allowed")
unknown_funcs.add(func_name)

# Reject SELECT statements with locking clauses
if isinstance(node, SelectStmt) and getattr(node, "lockingClause", None):
raise ValueError("Locking clause on select is prohibited")

# Reject EXPLAIN ANALYZE statements
# EXPLAIN ANALYZE executes the inner statement, so only allow it over a SELECT
if isinstance(node, ExplainStmt):
for option in node.options or []:
if isinstance(option, DefElem) and option.defname == "analyze":
raise ValueError("EXPLAIN ANALYZE is not supported")
if isinstance(option, DefElem) and option.defname == "analyze" and not isinstance(node.query, SelectStmt):
raise ValueError("EXPLAIN ANALYZE is only supported for SELECT statements")

# Reject CREATE EXTENSION statements
if isinstance(node, CreateExtensionStmt):
Expand All @@ -933,20 +938,25 @@ def _validate_node(self, node: Node) -> None:
if isinstance(attr, list):
for item in attr:
if isinstance(item, Node):
self._validate_node(item)
self._validate_node(item, unknown_funcs)

# Handle tuples of nodes
elif isinstance(attr, tuple):
for item in attr:
if isinstance(item, Node):
self._validate_node(item)
self._validate_node(item, unknown_funcs)

# Handle single nodes
elif isinstance(attr, Node):
self._validate_node(attr)
self._validate_node(attr, unknown_funcs)

def _validate(self, query: str) -> None:
"""Validate query is safe to execute"""
def _validate(self, query: str) -> set[str]:
"""Validate query is safe to execute.

Returns the function names that are not on ALLOWED_FUNCTIONS; the
caller must clear them against the catalog before executing.
"""
unknown_funcs: set[str] = set()
try:
# Parse the SQL using pglast
parsed = pglast.parse_sql(query)
Expand All @@ -970,21 +980,42 @@ def _validate(self, query: str) -> None:
raise ValueError(
"Only SELECT, ANALYZE, VACUUM, EXPLAIN, SHOW and other read-only statements are allowed. Received: " + str(stmt)
)
self._validate_node(stmt)
self._validate_node(stmt, unknown_funcs)
except Exception as e:
raise ValueError(f"Error validating query: {query}") from e

except pglast.parser.ParseError as e:
raise ValueError("Failed to parse SQL statement") from e

return unknown_funcs

async def _reject_unproven_functions(self, func_names: set[str]) -> None:
"""Allow a function outside ALLOWED_FUNCTIONS only when pg_proc proves
every overload of its name non-volatile (IMMUTABLE/STABLE cannot write
or carry side effects). Volatile names and names absent from pg_proc
are rejected — fail closed."""
probe_names = sorted({name.rsplit(".", 1)[-1] for name in func_names})
name_list = SQL(", ").join(Literal(n) for n in probe_names)
probe = SafeSqlDriver.param_sql_to_query(
"SELECT proname, bool_and(provolatile <> 'v') AS all_stable FROM pg_proc WHERE proname IN ({}) GROUP BY proname",
[name_list],
)
rows = await self.sql_driver.execute_query(probe, force_readonly=True) # type: ignore
proven = {row.cells["proname"] for row in (rows or []) if row.cells["all_stable"]}
rejected = sorted(name for name in func_names if name.rsplit(".", 1)[-1] not in proven)
if rejected:
raise ValueError(f"Function {', '.join(rejected)} is not allowed")

async def execute_query(
self,
query: LiteralString,
params: list[Any] | None = None,
force_readonly: bool = True, # do not use value passed in
) -> Optional[list[SqlDriver.RowResult]]: # noqa: UP007
"""Execute a query after validating it is safe"""
self._validate(query)
unknown_funcs = self._validate(query)
if unknown_funcs:
await self._reject_unproven_functions(unknown_funcs)

# NOTE: Always force readonly=True in SafeSqlDriver regardless of what was passed
if self.timeout:
Expand Down
94 changes: 89 additions & 5 deletions tests/unit/sql/test_safe_sql.py
Original file line number Diff line number Diff line change
Expand Up @@ -213,12 +213,23 @@ async def test_explain_plan(safe_driver, mock_sql_driver):


@pytest.mark.asyncio
async def test_explain_analyze_blocked(safe_driver):
"""Test that EXPLAIN ANALYZE is blocked"""
async def test_explain_analyze_select_allowed(safe_driver, mock_sql_driver):
"""EXPLAIN ANALYZE over a SELECT executes read-only work and is allowed"""
query = """
EXPLAIN ANALYZE
SELECT id, name FROM users
"""
await safe_driver.execute_query(query)
mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True)


@pytest.mark.asyncio
async def test_explain_analyze_non_select_blocked(safe_driver):
"""EXPLAIN ANALYZE over anything but a SELECT is still rejected"""
query = """
EXPLAIN ANALYZE
UPDATE users SET status = 'active'
"""
with pytest.raises(ValueError, match="Error validating query"):
await safe_driver.execute_query(query)

Expand Down Expand Up @@ -358,17 +369,90 @@ async def test_allowed_functions(safe_driver):


@pytest.mark.asyncio
async def test_disallowed_functions(safe_driver):
"""Test that disallowed functions are blocked"""
async def test_disallowed_functions(safe_driver, mock_sql_driver):
"""Volatile/unknown functions are blocked; only the pg_proc probe runs"""
queries = [
"SELECT pg_sleep(1);",
"SELECT pg_read_file('/etc/passwd');",
"SELECT lo_import('/etc/passwd');",
]

for query in queries:
with pytest.raises(ValueError, match="Error validating query"):
mock_sql_driver.execute_query.reset_mock()
with pytest.raises(ValueError, match="is not allowed"):
await safe_driver.execute_query(query)
assert mock_sql_driver.execute_query.await_count == 1
probe_sql = mock_sql_driver.execute_query.await_args.args[0]
assert "pg_proc" in probe_sql


@pytest.mark.asyncio
async def test_custom_function_proven_nonvolatile_allowed(safe_driver, mock_sql_driver):
"""A function off the static allowlist runs once pg_proc proves every overload non-volatile"""
mock_sql_driver.execute_query.side_effect = [
[SqlDriver.RowResult(cells={"proname": "shape_area", "all_stable": True})],
[],
]
query = "SELECT shape_area(shape) FROM leases"
await safe_driver.execute_query(query)
assert mock_sql_driver.execute_query.await_count == 2
probe_sql = mock_sql_driver.execute_query.await_args_list[0].args[0]
assert "pg_proc" in probe_sql and "'shape_area'" in probe_sql
assert mock_sql_driver.execute_query.await_args_list[1] == call("/* crystaldba */ " + query, params=None, force_readonly=True)


@pytest.mark.asyncio
async def test_custom_function_volatile_overload_blocked(safe_driver, mock_sql_driver):
"""A name with any volatile overload is rejected even though it exists"""
mock_sql_driver.execute_query.side_effect = [
[SqlDriver.RowResult(cells={"proname": "allocate_subnet_vni", "all_stable": False})],
]
with pytest.raises(ValueError, match="allocate_subnet_vni is not allowed"):
await safe_driver.execute_query("SELECT allocate_subnet_vni(7)")
assert mock_sql_driver.execute_query.await_count == 1


@pytest.mark.asyncio
async def test_schema_qualified_custom_function(safe_driver, mock_sql_driver):
"""Qualified names probe by their unqualified name but report the qualified one"""
mock_sql_driver.execute_query.side_effect = [[]]
with pytest.raises(ValueError, match=r"public\.no_such_func is not allowed"):
await safe_driver.execute_query("SELECT public.no_such_func(1)")
probe_sql = mock_sql_driver.execute_query.await_args_list[0].args[0]
assert "'no_such_func'" in probe_sql and "'public.no_such_func'" not in probe_sql


@pytest.mark.asyncio
async def test_mixed_known_and_unknown_functions_probe_only_unknown(safe_driver, mock_sql_driver):
"""Allowlisted functions never reach the probe; only unknown names do"""
mock_sql_driver.execute_query.side_effect = [
[SqlDriver.RowResult(cells={"proname": "shape_area", "all_stable": True})],
[],
]
await safe_driver.execute_query("SELECT count(*), shape_area(shape) FROM leases")
probe_sql = mock_sql_driver.execute_query.await_args_list[0].args[0]
assert "'shape_area'" in probe_sql and "'count'" not in probe_sql


@pytest.mark.asyncio
async def test_probe_returning_none_fails_closed(safe_driver, mock_sql_driver):
"""A probe that yields no rows at all still rejects every unknown name"""
mock_sql_driver.execute_query.side_effect = [None]
with pytest.raises(ValueError, match="shape_area is not allowed"):
await safe_driver.execute_query("SELECT shape_area(shape) FROM leases")


@pytest.mark.asyncio
async def test_partial_proof_rejects_only_unproven(safe_driver, mock_sql_driver):
"""With one proven and one volatile unknown, the rejection names only the volatile one"""
mock_sql_driver.execute_query.side_effect = [
[
SqlDriver.RowResult(cells={"proname": "shape_area", "all_stable": True}),
SqlDriver.RowResult(cells={"proname": "allocate_subnet_vni", "all_stable": False}),
],
]
with pytest.raises(ValueError, match=r"^Function allocate_subnet_vni is not allowed$"):
await safe_driver.execute_query("SELECT shape_area(shape), allocate_subnet_vni(7) FROM leases")


@pytest.mark.asyncio
Expand Down