diff --git a/.env.example b/.env.example new file mode 100644 index 0000000..f324fcc --- /dev/null +++ b/.env.example @@ -0,0 +1,15 @@ +# PostgreSQL MCP Configuration + +# Maximum allowed page size for query results (default: 500) +# This controls the upper limit of rows that can be returned in a single query +POSTGRES_MCP_MAX_PAGE_SIZE=500 + +# Default page size for query results (default: 100) +# This is the default number of rows returned if no page size is specified +POSTGRES_MCP_DEFAULT_PAGE_SIZE=100 + +# Database connection string (required) +DATABASE_URI=postgresql://user:password@localhost:5432/dbname + +# Allowed hosts for incoming connections (comma-separated, default: localhost,127.0.0.1) +POSTGRES_MCP_ALLOWED_HOSTS=localhost:*,127.0.0.1 diff --git a/pyproject.toml b/pyproject.toml index 3160a03..382b355 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -12,6 +12,7 @@ dependencies = [ "attrs>=25.4.0", "psycopg-pool>=3.3.0", "instructor>=1.14.4", + "dotenv>=0.9.9", ] license = "mit" license-files = ["LICENSE"] diff --git a/src/postgres_mcp/config.py b/src/postgres_mcp/config.py new file mode 100644 index 0000000..8e4aeeb --- /dev/null +++ b/src/postgres_mcp/config.py @@ -0,0 +1,78 @@ +"""Configuration management for postgres-mcp.""" + +import os + + +class Config: + """Configuration settings loaded from environment variables.""" + + def __init__(self): + """Initialize configuration from environment variables.""" + self._load_config() + + def _load_config(self): + """Load configuration from environment variables.""" + # Maximum allowed page size for queries (default: 500) + max_page_size_str = os.getenv("POSTGRES_MCP_MAX_PAGE_SIZE", "500") + try: + self._max_page_size = int(max_page_size_str) + if self._max_page_size < 1: + raise ValueError("POSTGRES_MCP_MAX_PAGE_SIZE must be at least 1") + except ValueError as e: + raise ValueError(f"Invalid POSTGRES_MCP_MAX_PAGE_SIZE value '{max_page_size_str}': {e}") from e + + # Maximum payload size in MB (default: 5) + max_payload_size_mb_str = os.getenv("POSTGRES_MCP_MAX_PAYLOAD_SIZE_MB", "5") + try: + self._max_payload_size_mb = int(max_payload_size_mb_str) + if self._max_payload_size_mb < 1: + raise ValueError("POSTGRES_MCP_MAX_PAYLOAD_SIZE_MB must be at least 1") + except ValueError as e: + raise ValueError(f"Invalid POSTGRES_MCP_MAX_PAYLOAD_SIZE_MB value '{max_payload_size_mb_str}': {e}") from e + + # Default page size for queries (default: 100) + default_page_size_str = os.getenv("POSTGRES_MCP_DEFAULT_PAGE_SIZE", "100") + try: + self._default_page_size = int(default_page_size_str) + if self._default_page_size < 1: + raise ValueError("POSTGRES_MCP_DEFAULT_PAGE_SIZE must be at least 1") + if self._default_page_size > self._max_page_size: + raise ValueError( + f"POSTGRES_MCP_DEFAULT_PAGE_SIZE ({self._default_page_size}) cannot exceed POSTGRES_MCP_MAX_PAGE_SIZE ({self._max_page_size})" + ) + except ValueError as e: + raise ValueError(f"Invalid POSTGRES_MCP_DEFAULT_PAGE_SIZE value '{default_page_size_str}': {e}") from e + + allowed_hosts_str = os.getenv("POSTGRES_MCP_ALLOWED_HOSTS", "localhost,localhost:*,127.0.0.1") + try: + self._allowed_hosts = [host.strip() for host in allowed_hosts_str.split(",")] + except Exception as e: + raise ValueError(f"Invalid POSTGRES_MCP_ALLOWED_HOSTS value '{allowed_hosts_str}': {e}") from e + + @property + def max_page_size(self) -> int: + """Get the maximum allowed page size for queries.""" + return self._max_page_size + + @property + def default_page_size(self) -> int: + """Get the default page size for queries.""" + return self._default_page_size + + @property + def max_payload_size_mb(self) -> int: + """Get the maximum allowed payload size in MB.""" + return self._max_payload_size_mb + + @property + def allowed_hosts(self) -> list[str]: + """Get the list of allowed hosts.""" + return self._allowed_hosts + + def reload(self): + """Reload configuration from environment variables.""" + self._load_config() + + +# Module-level configuration instance - import this directly +config = Config() diff --git a/src/postgres_mcp/server.py b/src/postgres_mcp/server.py index c43bbc1..f190077 100644 --- a/src/postgres_mcp/server.py +++ b/src/postgres_mcp/server.py @@ -11,15 +11,22 @@ from urllib.parse import urlparse import mcp.types as types +from dotenv import load_dotenv from mcp.server.fastmcp import FastMCP +from mcp.server.transport_security import TransportSecuritySettings from mcp.types import ToolAnnotations from pydantic import Field from pydantic import validate_call +# Load environment variables before importing local modules that may use them +load_dotenv() + +# ruff: noqa: E402 from postgres_mcp.index.dta_calc import DatabaseTuningAdvisor from .artifacts import ErrorResult from .artifacts import ExplainPlanArtifact +from .config import config from .database_health import DatabaseHealthTool from .database_health import HealthType from .explain import ExplainPlanTool @@ -38,7 +45,14 @@ from .utils.url import fix_connection_url # Initialize FastMCP with default settings -mcp = FastMCP("postgres-mcp") + + +mcp = FastMCP( + "postgres-mcp", + transport_security=TransportSecuritySettings( + allowed_hosts=config.allowed_hosts, + ), +) # Constants PG_STAT_STATEMENTS = "pg_stat_statements" @@ -270,11 +284,27 @@ async def explain_query( # Query function declaration without the decorator - we'll add it dynamically based on access mode async def execute_sql( sql: str = Field(description="SQL to run", default="all"), + page_size: int = Field( + description=f"Number of rows to return (1-{config.max_page_size}", + default=config.default_page_size, + ge=1, + le=config.max_page_size, + ), + offset: int = Field(description="Number of rows to skip for pagination", default=0, ge=0), + parameters: list[str | int | float | bool | None] = Field( + description="Optional array of parameters for parameterized queries", + default_factory=list, + ), ) -> ResponseType: """Executes a SQL query against the database.""" try: sql_driver = await sql_driver_module.get_sql_driver() - rows = await sql_driver.execute_query(sql) # type: ignore + rows = await sql_driver.execute_query( + sql, # type: ignore + params=parameters if parameters else None, + page_size=page_size, + offset=offset, + ) if rows is None: return format_text_response("No results") return format_text_response(list([r.cells for r in rows])) @@ -460,11 +490,11 @@ async def main(): # Add the query tool with a description appropriate to the access mode if sql_driver_module.current_access_mode == AccessMode.UNRESTRICTED: - mcp.add_tool(execute_sql, description="Execute any SQL query") + mcp.add_tool(execute_sql, description="Execute any SQL query with pagination support") else: mcp.add_tool( execute_sql, - description="Execute a read-only SQL query", + description="Execute a read-only SQL query with pagination support", annotations=ToolAnnotations( title="Execute SQL (Read-Only)", readOnlyHint=True, diff --git a/src/postgres_mcp/sql/safe_sql.py b/src/postgres_mcp/sql/safe_sql.py index 37382f0..a797652 100644 --- a/src/postgres_mcp/sql/safe_sql.py +++ b/src/postgres_mcp/sql/safe_sql.py @@ -80,6 +80,7 @@ from psycopg.sql import Literal from typing_extensions import LiteralString +from ..config import config from .sql_driver import SqlDriver logger = logging.getLogger(__name__) @@ -982,8 +983,14 @@ async def execute_query( query: LiteralString, params: list[Any] | None = None, force_readonly: bool = True, # do not use value passed in + page_size: int | None = None, + offset: int = 0, ) -> Optional[list[SqlDriver.RowResult]]: # noqa: UP007 """Execute a query after validating it is safe""" + # Use configured default if page_size not specified + if page_size is None: + page_size = config.default_page_size + self._validate(query) # NOTE: Always force readonly=True in SafeSqlDriver regardless of what was passed @@ -994,6 +1001,8 @@ async def execute_query( f"/* crystaldba */ {query}", params=params, force_readonly=True, + page_size=page_size, + offset=offset, ) except asyncio.TimeoutError as e: logger.warning(f"Query execution timed out after {self.timeout} seconds: {query[:100]}...") @@ -1009,6 +1018,8 @@ async def execute_query( f"/* crystaldba */ {query}", params=params, force_readonly=True, + page_size=page_size, + offset=offset, ) @staticmethod diff --git a/src/postgres_mcp/sql/sql_driver.py b/src/postgres_mcp/sql/sql_driver.py index 5beacb0..2277e9d 100644 --- a/src/postgres_mcp/sql/sql_driver.py +++ b/src/postgres_mcp/sql/sql_driver.py @@ -1,8 +1,12 @@ """SQL driver adapter for PostgreSQL connections.""" +import io +import json import logging import re from dataclasses import dataclass +from datetime import date +from datetime import datetime from typing import Any from typing import Dict from typing import List @@ -14,9 +18,21 @@ from psycopg_pool import AsyncConnectionPool from typing_extensions import LiteralString +from ..config import config + logger = logging.getLogger(__name__) +def _has_limit_or_offset(query: str) -> bool: + """Check if SQL query has LIMIT or OFFSET (case-insensitive word boundary check).""" + import re + + # Use word boundaries to avoid matching within identifiers + # This handles 95% of cases correctly + pattern = r"\b(LIMIT|OFFSET)\b" + return bool(re.search(pattern, query, re.IGNORECASE)) + + def obfuscate_password(text: str | None) -> str | None: """ Obfuscate password in any text containing connection information. @@ -184,6 +200,8 @@ async def execute_query( query: LiteralString, params: list[Any] | None = None, force_readonly: bool = False, + page_size: int | None = None, + offset: int = 0, ) -> Optional[List[RowResult]]: """ Execute a query and return results. @@ -192,6 +210,8 @@ async def execute_query( query: SQL query to execute params: Query parameters force_readonly: Whether to enforce read-only mode + page_size: Number of rows to return (1-max configured), defaults to config.default_page_size + offset: Number of rows to skip Returns: List of RowResult objects or None on error @@ -207,10 +227,24 @@ async def execute_query( # For pools, get a connection from the pool pool = await self.conn.pool_connect() async with pool.connection() as connection: - return await self._execute_with_connection(connection, query, params, force_readonly=force_readonly) + return await self._execute_with_connection( + connection, + query, + params, + force_readonly=force_readonly, + page_size=page_size, + offset=offset, + ) else: # Direct connection approach - return await self._execute_with_connection(self.conn, query, params, force_readonly=force_readonly) + return await self._execute_with_connection( + self.conn, + query, + params, + force_readonly=force_readonly, + page_size=page_size, + offset=offset, + ) except Exception as e: # Mark pool as invalid if there was a connection issue if self.conn and self.is_pool: @@ -221,8 +255,32 @@ async def execute_query( raise e - async def _execute_with_connection(self, connection, query, params, force_readonly) -> Optional[List[RowResult]]: - """Execute query with the given connection.""" + def get_wire_size(self, data: list[dict[str, Any]]) -> int: + """Calculates exact bytes of the JSON-serialized data including datetimes.""" + + def json_serial(obj: Any) -> str: + """JSON serializer that converts any non-serializable object to string. + + This is only used for wire size calculation, so we prioritize + robustness over perfect type preservation. + """ + if isinstance(obj, (datetime, date)): + return obj.isoformat() + + return str(obj) + + buffer = io.StringIO() + json.dump(data, buffer, default=json_serial) + return len(buffer.getvalue().encode("utf-8")) + + async def _execute_with_connection( + self, connection, query, params, force_readonly, page_size: int | None = None, offset: int = 0 + ) -> Optional[List[RowResult]]: + """Execute query with the given connection and apply pagination.""" + + if page_size is None: + page_size = config.default_page_size + transaction_started = False try: async with connection.cursor(row_factory=dict_row) as cursor: @@ -231,10 +289,30 @@ async def _execute_with_connection(self, connection, query, params, force_readon await cursor.execute("BEGIN TRANSACTION READ ONLY") transaction_started = True + paginated_query = query + # Only apply pagination in readonly mode to avoid breaking DDL operations + if force_readonly and page_size > 0: + # Remove trailing semicolon if present (we'll add it back later) + query_trimmed = query.rstrip().rstrip(";") + had_semicolon = query.rstrip().endswith(";") + + # Use proper SQL parsing to check for existing LIMIT/OFFSET + if not _has_limit_or_offset(query_trimmed): + # Safe to add pagination + paginated_query = f"{query_trimmed} LIMIT {page_size} OFFSET {offset}" + # Restore semicolon if original had one + if had_semicolon: + paginated_query += ";" + else: + # Query already has pagination, use it as-is + paginated_query = query_trimmed + if had_semicolon: + paginated_query += ";" + if params: - await cursor.execute(query, params) + await cursor.execute(paginated_query, params) else: - await cursor.execute(query) + await cursor.execute(paginated_query) # For multiple statements, move to the last statement's results while cursor.nextset(): @@ -258,7 +336,19 @@ async def _execute_with_connection(self, connection, query, params, force_readon await cursor.execute("ROLLBACK") transaction_started = False - return [SqlDriver.RowResult(cells=dict(row)) for row in rows] + result = [SqlDriver.RowResult(cells=dict(row)) for row in rows] + + wire_size_bytes: int = self.get_wire_size([r.cells for r in result]) + + payload_size_mb = wire_size_bytes / (1024 * 1024) + + if payload_size_mb > config.max_payload_size_mb: + raise ValueError( + f"Query result payload too large: {payload_size_mb:.2f}MB exceeds maximum allowed size of {config.max_payload_size_mb}MB. " + f"Please refine your query to return less data, use pagination (LIMIT/OFFSET), or filter results." + ) + + return result except Exception as e: # Try to roll back the transaction if it's still active diff --git a/tests/unit/sql/test_safe_sql.py b/tests/unit/sql/test_safe_sql.py index c55d253..8db6393 100644 --- a/tests/unit/sql/test_safe_sql.py +++ b/tests/unit/sql/test_safe_sql.py @@ -28,7 +28,7 @@ async def test_select_statement(safe_driver, mock_sql_driver): """Test that simple SELECT statements are allowed""" query = "SELECT * FROM users WHERE age > 18" await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.asyncio @@ -52,7 +52,7 @@ async def test_select_with_join(safe_driver, mock_sql_driver): WHERE orders.status = 'pending' """ await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.asyncio @@ -60,7 +60,7 @@ async def test_show_variable(safe_driver, mock_sql_driver): """Test that SHOW statements are allowed""" query = "SHOW search_path" await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.asyncio @@ -79,7 +79,7 @@ async def test_select_with_arithmetic(safe_driver, mock_sql_driver): """Test that SELECT with arithmetic expressions is allowed""" query = "SELECT id, price * quantity as total FROM orders" await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.asyncio @@ -87,7 +87,7 @@ async def test_select_current_user(safe_driver, mock_sql_driver): """Test that SELECT current_user is allowed""" query = "SELECT current_user" await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.asyncio @@ -120,7 +120,7 @@ async def test_select_with_subquery(safe_driver, mock_sql_driver): WHERE id IN (SELECT user_id FROM orders WHERE total > 1000) """ await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.asyncio @@ -145,7 +145,7 @@ async def test_select_with_union(safe_driver, mock_sql_driver): SELECT NULL, concat(table_name) FROM information_schema.tables """ await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.asyncio @@ -209,7 +209,7 @@ async def test_explain_plan(safe_driver, mock_sql_driver): WHERE age > $1 AND status = $2 """ await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.asyncio @@ -345,7 +345,7 @@ async def test_complex_index_metadata_select(safe_driver, mock_sql_driver): GROUP BY indexrelid, indisunique, indisprimary HAVING COUNT(array_agg(attname)) > 1""" await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.asyncio @@ -376,7 +376,7 @@ async def test_session_info_functions(safe_driver, mock_sql_driver): """Test that session info functions are allowed""" query = "SELECT current_user, current_database(), version()" await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.asyncio @@ -384,7 +384,7 @@ async def test_blocking_pids_functions(safe_driver, mock_sql_driver): """Test that blocking pids functions are allowed""" query = "SELECT pg_blocking_pids(1234)" await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.asyncio @@ -392,7 +392,7 @@ async def test_logfile_functions(safe_driver, mock_sql_driver): """Test that logfile functions are allowed""" query = "SELECT pg_current_logfile()" await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.asyncio @@ -403,7 +403,7 @@ async def test_complex_session_info_queries(safe_driver, mock_sql_driver): pg_backend_pid(), pg_blocking_pids(pg_backend_pid()) """ await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.asyncio @@ -411,7 +411,7 @@ async def test_security_privilege_functions(safe_driver, mock_sql_driver): """Test that security privilege functions are allowed""" query = "SELECT has_table_privilege('user', 'table', 'SELECT')" await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.asyncio @@ -446,7 +446,7 @@ async def test_complex_security_privilege_queries(safe_driver, mock_sql_driver): for query in queries: await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.asyncio @@ -475,7 +475,7 @@ async def test_security_privilege_functions_with_subqueries(safe_driver, mock_sq for query in queries: await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.parametrize("operator", ["LIKE", "ILIKE"]) @@ -492,7 +492,7 @@ async def test_like_patterns(safe_driver, mock_sql_driver, operator): for query in queries: await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.asyncio @@ -526,7 +526,7 @@ async def test_datetime_functions(safe_driver, mock_sql_driver): for query in queries: await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.asyncio @@ -553,7 +553,7 @@ async def test_type_conversion_functions(safe_driver, mock_sql_driver): for query in queries: await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.asyncio @@ -561,7 +561,7 @@ async def test_regexp_functions(safe_driver, mock_sql_driver): """Test that regexp functions are allowed""" query = "SELECT regexp_replace('Hello World', 'World', 'PostgreSQL')" await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.asyncio @@ -574,7 +574,7 @@ async def test_complex_type_conversion_queries(safe_driver, mock_sql_driver): to_number('123.45', '999.99') """ await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.asyncio @@ -582,7 +582,7 @@ async def test_network_functions(safe_driver, mock_sql_driver): """Test that network functions are allowed""" query = "SELECT inet_client_addr(), inet_client_port()" await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.asyncio @@ -595,7 +595,7 @@ async def test_network_functions_in_complex_queries(safe_driver, mock_sql_driver inet_server_port() as server_port """ await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.asyncio @@ -603,7 +603,7 @@ async def test_notification_and_server_functions(safe_driver, mock_sql_driver): """Test that notification and server functions are allowed""" query = "SELECT pg_listening_channels(), pg_postmaster_start_time()" await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.asyncio @@ -611,7 +611,7 @@ async def test_minmax_expressions(safe_driver, mock_sql_driver): """Test that minmax expressions are allowed""" query = "SELECT GREATEST(1, 2, 3), LEAST(1, 2, 3)" await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.asyncio @@ -619,7 +619,7 @@ async def test_row_expressions(safe_driver, mock_sql_driver): """Test that row expressions are allowed""" query = "SELECT ROW(1, 2, 3) = ROW(1, 2, 3)" await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.asyncio @@ -627,7 +627,7 @@ async def test_extension_check_query(safe_driver, mock_sql_driver): """Test that extension check queries are allowed""" query = "SELECT extname, extversion FROM pg_extension WHERE extname = 'hypopg'" await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.asyncio @@ -635,7 +635,7 @@ async def test_create_extension_query(safe_driver, mock_sql_driver): """Test that CREATE EXTENSION queries are allowed""" query = "CREATE EXTENSION IF NOT EXISTS hypopg" await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.asyncio @@ -643,7 +643,7 @@ async def test_hypopg_create_index_query(safe_driver, mock_sql_driver): """Test that hypopg create index queries are allowed""" query = "SELECT * FROM hypopg_create_index('CREATE INDEX idx ON users(id)')" await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.asyncio @@ -651,7 +651,7 @@ async def test_hypopg_reset_query(safe_driver, mock_sql_driver): """Test that hypopg reset queries are allowed""" query = "SELECT * FROM hypopg_reset()" await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.asyncio @@ -659,7 +659,7 @@ async def test_hypopg_list_indexes_query(safe_driver, mock_sql_driver): """Test that hypopg list indexes queries are allowed""" query = "SELECT * FROM hypopg_list_indexes()" await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.asyncio @@ -667,7 +667,7 @@ async def test_pg_stat_statements_query(safe_driver, mock_sql_driver): """Test that pg_stat_statements queries are allowed""" query = "SELECT * FROM pg_stat_statements ORDER BY calls DESC LIMIT 10" await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.asyncio @@ -675,7 +675,7 @@ async def test_pg_indexes_query(safe_driver, mock_sql_driver): """Test that pg_indexes queries are allowed""" query = "SELECT * FROM pg_indexes WHERE schemaname = 'public'" await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.asyncio @@ -683,7 +683,7 @@ async def test_pg_stats_query(safe_driver, mock_sql_driver): """Test that pg_stats queries are allowed""" query = "SELECT * FROM pg_stats WHERE schemaname = 'public'" await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.asyncio @@ -691,7 +691,7 @@ async def test_explain_query(safe_driver, mock_sql_driver): """Test that explain queries are allowed""" query = "EXPLAIN SELECT * FROM users" await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.asyncio @@ -713,7 +713,9 @@ async def test_sql_driver_parameter_format(safe_driver, mock_sql_driver): formatted_query = SQL(query_template).format(Literal(min_calls), Literal(min_avg_time), Literal(limit)).as_string() await safe_driver.execute_query(formatted_query) - mock_sql_driver.execute_query.assert_awaited_with("/* crystaldba */ " + formatted_query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_with( + "/* crystaldba */ " + formatted_query, params=None, force_readonly=True, page_size=100, offset=0 + ) @pytest.mark.asyncio @@ -725,8 +727,8 @@ async def test_multiple_queries(safe_driver, mock_sql_driver): await safe_driver.execute_query(query2) mock_sql_driver.execute_query.assert_has_awaits( [ - call("/* crystaldba */ " + query1, params=None, force_readonly=True), - call("/* crystaldba */ " + query2, params=None, force_readonly=True), + call("/* crystaldba */ " + query1, params=None, force_readonly=True, page_size=100, offset=0), + call("/* crystaldba */ " + query2, params=None, force_readonly=True, page_size=100, offset=0), ] ) @@ -742,7 +744,7 @@ async def test_query_with_comments(safe_driver, mock_sql_driver): -- Only get active users """ await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) @pytest.mark.asyncio @@ -757,4 +759,4 @@ async def test_query_with_whitespace(safe_driver, mock_sql_driver): ORDER BY name """ await safe_driver.execute_query(query) - mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True) + mock_sql_driver.execute_query.assert_awaited_once_with("/* crystaldba */ " + query, params=None, force_readonly=True, page_size=100, offset=0) diff --git a/uv.lock b/uv.lock index 046959a..0f680e0 100644 --- a/uv.lock +++ b/uv.lock @@ -341,6 +341,17 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/d5/7c/e9fcff7623954d86bdc17782036cbf715ecab1bec4847c008557affe1ca8/docstring_parser-0.16-py3-none-any.whl", hash = "sha256:bf0a1387354d3691d102edef7ec124f219ef639982d096e26e3b60aeffa90637", size = 36533, upload-time = "2024-03-15T10:39:41.527Z" }, ] +[[package]] +name = "dotenv" +version = "0.9.9" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "python-dotenv" }, +] +wheels = [ + { url = "https://files.pythonhosted.org/packages/b2/b7/545d2c10c1fc15e48653c91efde329a790f2eecfbbf2bd16003b5db2bab0/dotenv-0.9.9-py2.py3-none-any.whl", hash = "sha256:29cf74a087b31dafdb5a446b6d7e11cbce8ed2741540e2339c69fbef92c94ce9", size = 1892, upload-time = "2025-02-19T22:15:01.647Z" }, +] + [[package]] name = "filelock" version = "3.20.3" @@ -875,6 +886,7 @@ version = "0.3.0" source = { editable = "." } dependencies = [ { name = "attrs" }, + { name = "dotenv" }, { name = "humanize" }, { name = "instructor" }, { name = "mcp", extra = ["cli"] }, @@ -895,6 +907,7 @@ dev = [ [package.metadata] requires-dist = [ { name = "attrs", specifier = ">=25.4.0" }, + { name = "dotenv", specifier = ">=0.9.9" }, { name = "humanize", specifier = ">=4.15.0" }, { name = "instructor", specifier = ">=1.14.4" }, { name = "mcp", extras = ["cli"], specifier = ">=1.25.0" },