|
| 1 | +"""Tests for extension management and the enable_extension tool.""" |
| 2 | + |
| 3 | +import pytest |
| 4 | +from _fakes import FakeDatabase, FakeDriver |
| 5 | +from mcp.shared.memory import create_connected_server_and_client_session |
| 6 | + |
| 7 | +from mcpg.config import load_settings |
| 8 | +from mcpg.extensions import EnableExtensionResult, ExtensionError, enable_extension |
| 9 | +from mcpg.server import create_server |
| 10 | + |
| 11 | +_UNRESTRICTED_DDL = load_settings( |
| 12 | + { |
| 13 | + "MCPG_DATABASE_URL": "postgresql://u:p@localhost/db", |
| 14 | + "MCPG_ACCESS_MODE": "unrestricted", |
| 15 | + "MCPG_ALLOW_DDL": "true", |
| 16 | + } |
| 17 | +) |
| 18 | +_READ_ONLY = load_settings({"MCPG_DATABASE_URL": "postgresql://u:p@localhost/db"}) |
| 19 | + |
| 20 | + |
| 21 | +async def test_enable_extension_runs_create_extension_for_an_allowlisted_name() -> None: |
| 22 | + driver = FakeDriver() |
| 23 | + |
| 24 | + result = await enable_extension(driver, "pg_trgm") |
| 25 | + |
| 26 | + assert result == EnableExtensionResult(name="pg_trgm", enabled=True) |
| 27 | + query, _params, force_readonly = driver.calls[0] |
| 28 | + assert query == 'CREATE EXTENSION IF NOT EXISTS "pg_trgm"' |
| 29 | + assert force_readonly is False |
| 30 | + |
| 31 | + |
| 32 | +async def test_enable_extension_rejects_a_name_not_on_the_allowlist() -> None: |
| 33 | + driver = FakeDriver() |
| 34 | + |
| 35 | + with pytest.raises(ExtensionError, match="allowlist"): |
| 36 | + await enable_extension(driver, "evil; DROP DATABASE postgres") |
| 37 | + # Rejection happens before any SQL is built. |
| 38 | + assert driver.calls == [] |
| 39 | + |
| 40 | + |
| 41 | +async def test_enable_extension_wraps_execution_failures() -> None: |
| 42 | + with pytest.raises(ExtensionError, match="execution failed"): |
| 43 | + await enable_extension(FakeDriver(fail=True), "pg_trgm") |
| 44 | + |
| 45 | + |
| 46 | +async def test_enable_extension_tool_is_callable_when_ddl_is_allowed() -> None: |
| 47 | + server = create_server(_UNRESTRICTED_DDL, database=FakeDatabase(FakeDriver())) # type: ignore[arg-type] |
| 48 | + |
| 49 | + async with create_connected_server_and_client_session(server) as client: |
| 50 | + result = await client.call_tool("enable_extension", {"name": "pg_trgm"}) |
| 51 | + |
| 52 | + assert result.isError is False |
| 53 | + |
| 54 | + |
| 55 | +async def test_enable_extension_tool_is_absent_without_ddl_opt_in() -> None: |
| 56 | + server = create_server(_READ_ONLY, database=FakeDatabase(FakeDriver())) # type: ignore[arg-type] |
| 57 | + |
| 58 | + async with create_connected_server_and_client_session(server) as client: |
| 59 | + names = {tool.name for tool in (await client.list_tools()).tools} |
| 60 | + |
| 61 | + assert "enable_extension" not in names |
0 commit comments