From ae25bb6df108fd8689f42b7d1d732dddcbef4213 Mon Sep 17 00:00:00 2001 From: Caleb Mabry Date: Wed, 11 Feb 2026 11:07:11 -0500 Subject: [PATCH 01/12] Be able to get variables from config file --- .env.example | 12 ++++++++ pyproject.toml | 1 + src/postgres_mcp/config.py | 58 ++++++++++++++++++++++++++++++++++++++ uv.lock | 13 +++++++++ 4 files changed, 84 insertions(+) create mode 100644 .env.example create mode 100644 src/postgres_mcp/config.py diff --git a/.env.example b/.env.example new file mode 100644 index 0000000..a2b3d71 --- /dev/null +++ b/.env.example @@ -0,0 +1,12 @@ +# 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 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..9a10323 --- /dev/null +++ b/src/postgres_mcp/config.py @@ -0,0 +1,58 @@ +"""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("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}" + ) + + # 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("DEFAULT_PAGE_SIZE must be at least 1") + if self._default_page_size > self._max_page_size: + raise ValueError( + f"DEFAULT_PAGE_SIZE ({self._default_page_size}) cannot exceed " + f"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}" + ) + + @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 + + def reload(self): + """Reload configuration from environment variables.""" + self._load_config() + + +# Module-level configuration instance - import this directly +config = Config() 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" }, From 4ca9fbe6d70b5672841ed1601aff3b918a67ce1c Mon Sep 17 00:00:00 2001 From: Caleb Mabry Date: Wed, 11 Feb 2026 11:07:23 -0500 Subject: [PATCH 02/12] Allow for paginated queries --- src/postgres_mcp/server.py | 26 +++++++++++-- src/postgres_mcp/sql/safe_sql.py | 12 ++++++ src/postgres_mcp/sql/sql_driver.py | 61 +++++++++++++++++++++++++++--- 3 files changed, 90 insertions(+), 9 deletions(-) diff --git a/src/postgres_mcp/server.py b/src/postgres_mcp/server.py index c43bbc1..aa09872 100644 --- a/src/postgres_mcp/server.py +++ b/src/postgres_mcp/server.py @@ -20,6 +20,7 @@ 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 @@ -36,6 +37,9 @@ from .top_queries import TopQueriesCalc from .utils import sql_driver as sql_driver_module # Import the module to access global state from .utils.url import fix_connection_url +from dotenv import load_dotenv + +load_dotenv() # Initialize FastMCP with default settings mcp = FastMCP("postgres-mcp") @@ -270,11 +274,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"), + pageSize: 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=pageSize, + 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 +480,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..9b2125a 100644 --- a/src/postgres_mcp/sql/safe_sql.py +++ b/src/postgres_mcp/sql/safe_sql.py @@ -8,6 +8,8 @@ from typing import Optional import pglast + +from ..config import config from pglast.ast import A_ArrayExpr from pglast.ast import A_Const from pglast.ast import A_Expr @@ -982,8 +984,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 +1002,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 +1019,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..12f752f 100644 --- a/src/postgres_mcp/sql/sql_driver.py +++ b/src/postgres_mcp/sql/sql_driver.py @@ -14,9 +14,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 +196,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 +206,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 +223,20 @@ 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 +247,12 @@ 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.""" + 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 +261,29 @@ async def _execute_with_connection(self, connection, query, params, force_readon await cursor.execute("BEGIN TRANSACTION READ ONLY") transaction_started = True + paginated_query = query + if 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(): From 1a9491269d319d44a48305b0dbe2479aa7c1da11 Mon Sep 17 00:00:00 2001 From: Caleb Mabry Date: Wed, 11 Feb 2026 11:38:04 -0500 Subject: [PATCH 03/12] Add POSTGRES_MCP_MAX_PAYLOAD_SIZE_MB --- src/postgres_mcp/config.py | 24 ++++++++++++++++++---- src/postgres_mcp/server.py | 6 ++++-- src/postgres_mcp/sql/sql_driver.py | 32 +++++++++++++++++++++++++++++- 3 files changed, 55 insertions(+), 7 deletions(-) diff --git a/src/postgres_mcp/config.py b/src/postgres_mcp/config.py index 9a10323..8614628 100644 --- a/src/postgres_mcp/config.py +++ b/src/postgres_mcp/config.py @@ -17,22 +17,33 @@ def _load_config(self): try: self._max_page_size = int(max_page_size_str) if self._max_page_size < 1: - raise ValueError("MAX_PAGE_SIZE must be at least 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}" ) + # 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}" + ) + # 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("DEFAULT_PAGE_SIZE must be at least 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"DEFAULT_PAGE_SIZE ({self._default_page_size}) cannot exceed " - f"MAX_PAGE_SIZE ({self._max_page_size})" + f"POSTGRES_MCP_DEFAULT_PAGE_SIZE ({self._default_page_size}) cannot exceed " + f"POSTGRES_MCP_MAX_PAGE_SIZE ({self._max_page_size})" ) except ValueError as e: raise ValueError( @@ -49,6 +60,11 @@ 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 + def reload(self): """Reload configuration from environment variables.""" self._load_config() diff --git a/src/postgres_mcp/server.py b/src/postgres_mcp/server.py index aa09872..37847d2 100644 --- a/src/postgres_mcp/server.py +++ b/src/postgres_mcp/server.py @@ -1,4 +1,8 @@ # ruff: noqa: B008 +from dotenv import load_dotenv + +load_dotenv() + import argparse import asyncio import logging @@ -37,9 +41,7 @@ from .top_queries import TopQueriesCalc from .utils import sql_driver as sql_driver_module # Import the module to access global state from .utils.url import fix_connection_url -from dotenv import load_dotenv -load_dotenv() # Initialize FastMCP with default settings mcp = FastMCP("postgres-mcp") diff --git a/src/postgres_mcp/sql/sql_driver.py b/src/postgres_mcp/sql/sql_driver.py index 12f752f..ff17042 100644 --- a/src/postgres_mcp/sql/sql_driver.py +++ b/src/postgres_mcp/sql/sql_driver.py @@ -1,7 +1,11 @@ """SQL driver adapter for PostgreSQL connections.""" +from datetime import datetime, date +import io +import json import logging import re +import sys from dataclasses import dataclass from typing import Any from typing import Dict @@ -246,6 +250,20 @@ async def execute_query( self.conn = None raise e + + def get_wire_size(self, data: list[dict[str, Any]]) -> int: + def json_serial(obj: Any) -> str: + """JSON serializer for objects not handled by default json package""" + if isinstance(obj, (datetime, date)): + return obj.isoformat() + raise TypeError(f"Type {type(obj)} not serializable") + + """Calculates exact bytes of the JSON-serialized data including datetimes.""" + buffer = io.StringIO() + # Add 'default=json_serial' to handle the PostgreSQL timestamps + 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.""" @@ -307,7 +325,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 From dd6b2e3342d75ee832c03ff1cab1c42118e35279 Mon Sep 17 00:00:00 2001 From: Caleb Mabry Date: Wed, 11 Feb 2026 11:48:07 -0500 Subject: [PATCH 04/12] Ruff formatting --- src/postgres_mcp/config.py | 15 ++++----------- src/postgres_mcp/server.py | 2 +- src/postgres_mcp/sql/safe_sql.py | 2 +- src/postgres_mcp/sql/sql_driver.py | 29 +++++++++++++++++------------ 4 files changed, 23 insertions(+), 25 deletions(-) diff --git a/src/postgres_mcp/config.py b/src/postgres_mcp/config.py index 8614628..8922aa9 100644 --- a/src/postgres_mcp/config.py +++ b/src/postgres_mcp/config.py @@ -19,9 +19,7 @@ def _load_config(self): 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}" - ) + raise ValueError(f"Invalid POSTGRES_MCP_MAX_PAGE_SIZE value '{max_page_size_str}': {e}") # Maximum payload size in MB (default: 5) max_payload_size_mb_str = os.getenv("POSTGRES_MCP_MAX_PAYLOAD_SIZE_MB", "5") @@ -30,9 +28,7 @@ def _load_config(self): 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}" - ) + raise ValueError(f"Invalid POSTGRES_MCP_MAX_PAYLOAD_SIZE_MB value '{max_payload_size_mb_str}': {e}") # Default page size for queries (default: 100) default_page_size_str = os.getenv("POSTGRES_MCP_DEFAULT_PAGE_SIZE", "100") @@ -42,13 +38,10 @@ def _load_config(self): 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 " - f"POSTGRES_MCP_MAX_PAGE_SIZE ({self._max_page_size})" + 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}" - ) + raise ValueError(f"Invalid POSTGRES_MCP_DEFAULT_PAGE_SIZE value '{default_page_size_str}': {e}") @property def max_page_size(self) -> int: diff --git a/src/postgres_mcp/server.py b/src/postgres_mcp/server.py index 37847d2..33d141a 100644 --- a/src/postgres_mcp/server.py +++ b/src/postgres_mcp/server.py @@ -292,7 +292,7 @@ async def execute_sql( try: sql_driver = await sql_driver_module.get_sql_driver() rows = await sql_driver.execute_query( - sql, # type: ignore + sql, # type: ignore params=parameters if parameters else None, page_size=pageSize, offset=offset, diff --git a/src/postgres_mcp/sql/safe_sql.py b/src/postgres_mcp/sql/safe_sql.py index 9b2125a..5864031 100644 --- a/src/postgres_mcp/sql/safe_sql.py +++ b/src/postgres_mcp/sql/safe_sql.py @@ -991,7 +991,7 @@ async def execute_query( # 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 diff --git a/src/postgres_mcp/sql/sql_driver.py b/src/postgres_mcp/sql/sql_driver.py index ff17042..4d560c2 100644 --- a/src/postgres_mcp/sql/sql_driver.py +++ b/src/postgres_mcp/sql/sql_driver.py @@ -26,13 +26,13 @@ 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' + 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. @@ -228,7 +228,9 @@ async def execute_query( pool = await self.conn.pool_connect() async with pool.connection() as connection: return await self._execute_with_connection( - connection, query, params, + connection, + query, + params, force_readonly=force_readonly, page_size=page_size, offset=offset, @@ -236,7 +238,9 @@ async def execute_query( else: # Direct connection approach return await self._execute_with_connection( - self.conn, query, params, + self.conn, + query, + params, force_readonly=force_readonly, page_size=page_size, offset=offset, @@ -250,7 +254,7 @@ async def execute_query( self.conn = None raise e - + def get_wire_size(self, data: list[dict[str, Any]]) -> int: def json_serial(obj: Any) -> str: """JSON serializer for objects not handled by default json package""" @@ -264,8 +268,9 @@ def json_serial(obj: Any) -> str: 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]]: + 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: @@ -284,7 +289,7 @@ async def _execute_with_connection(self, connection, query, params, force_readon # 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 @@ -297,7 +302,7 @@ async def _execute_with_connection(self, connection, query, params, force_readon paginated_query = query_trimmed if had_semicolon: paginated_query += ";" - + if params: await cursor.execute(paginated_query, params) else: @@ -326,17 +331,17 @@ async def _execute_with_connection(self, connection, query, params, force_readon transaction_started = False 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: From fa337134c4bc173a5a0fb94a30935d7f7881f816 Mon Sep 17 00:00:00 2001 From: Caleb Mabry Date: Wed, 11 Feb 2026 11:56:08 -0500 Subject: [PATCH 05/12] Resolve ruff --- src/postgres_mcp/config.py | 6 +++--- src/postgres_mcp/server.py | 14 +++++++------- src/postgres_mcp/sql/safe_sql.py | 3 +-- src/postgres_mcp/sql/sql_driver.py | 4 ++-- 4 files changed, 13 insertions(+), 14 deletions(-) diff --git a/src/postgres_mcp/config.py b/src/postgres_mcp/config.py index 8922aa9..ad66fcd 100644 --- a/src/postgres_mcp/config.py +++ b/src/postgres_mcp/config.py @@ -19,7 +19,7 @@ def _load_config(self): 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}") + 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") @@ -28,7 +28,7 @@ def _load_config(self): 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}") + 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") @@ -41,7 +41,7 @@ def _load_config(self): 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}") + raise ValueError(f"Invalid POSTGRES_MCP_DEFAULT_PAGE_SIZE value '{default_page_size_str}': {e}") from e @property def max_page_size(self) -> int: diff --git a/src/postgres_mcp/server.py b/src/postgres_mcp/server.py index 33d141a..1a044b4 100644 --- a/src/postgres_mcp/server.py +++ b/src/postgres_mcp/server.py @@ -1,8 +1,4 @@ # ruff: noqa: B008 -from dotenv import load_dotenv - -load_dotenv() - import argparse import asyncio import logging @@ -15,11 +11,16 @@ from urllib.parse import urlparse import mcp.types as types +from dotenv import load_dotenv from mcp.server.fastmcp import FastMCP 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 @@ -42,7 +43,6 @@ from .utils import sql_driver as sql_driver_module # Import the module to access global state from .utils.url import fix_connection_url - # Initialize FastMCP with default settings mcp = FastMCP("postgres-mcp") @@ -276,7 +276,7 @@ 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"), - pageSize: int = Field( + page_size: int = Field( description=f"Number of rows to return (1-{config.max_page_size}", default=config.default_page_size, ge=1, @@ -294,7 +294,7 @@ async def execute_sql( rows = await sql_driver.execute_query( sql, # type: ignore params=parameters if parameters else None, - page_size=pageSize, + page_size=page_size, offset=offset, ) if rows is None: diff --git a/src/postgres_mcp/sql/safe_sql.py b/src/postgres_mcp/sql/safe_sql.py index 5864031..a797652 100644 --- a/src/postgres_mcp/sql/safe_sql.py +++ b/src/postgres_mcp/sql/safe_sql.py @@ -8,8 +8,6 @@ from typing import Optional import pglast - -from ..config import config from pglast.ast import A_ArrayExpr from pglast.ast import A_Const from pglast.ast import A_Expr @@ -82,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__) diff --git a/src/postgres_mcp/sql/sql_driver.py b/src/postgres_mcp/sql/sql_driver.py index 4d560c2..1be358d 100644 --- a/src/postgres_mcp/sql/sql_driver.py +++ b/src/postgres_mcp/sql/sql_driver.py @@ -1,12 +1,12 @@ """SQL driver adapter for PostgreSQL connections.""" -from datetime import datetime, date import io import json import logging import re -import sys from dataclasses import dataclass +from datetime import date +from datetime import datetime from typing import Any from typing import Dict from typing import List From 5b840de6a1804e6ebd35eeeef2880101065779c2 Mon Sep 17 00:00:00 2001 From: Caleb Mabry Date: Wed, 11 Feb 2026 12:11:09 -0500 Subject: [PATCH 06/12] Only limit if readonly. --- src/postgres_mcp/sql/sql_driver.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/src/postgres_mcp/sql/sql_driver.py b/src/postgres_mcp/sql/sql_driver.py index 1be358d..028b6e3 100644 --- a/src/postgres_mcp/sql/sql_driver.py +++ b/src/postgres_mcp/sql/sql_driver.py @@ -285,7 +285,8 @@ async def _execute_with_connection( transaction_started = True paginated_query = query - if page_size > 0: + # 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(";") From 5b1c7821b138214c3eef0c13678e273724d6834d Mon Sep 17 00:00:00 2001 From: Caleb Mabry Date: Wed, 11 Feb 2026 13:02:04 -0500 Subject: [PATCH 07/12] Resolve test failures and lint failures --- src/postgres_mcp/sql/sql_driver.py | 13 +++-- tests/unit/sql/test_safe_sql.py | 82 +++++++++++++++--------------- 2 files changed, 51 insertions(+), 44 deletions(-) diff --git a/src/postgres_mcp/sql/sql_driver.py b/src/postgres_mcp/sql/sql_driver.py index 028b6e3..2277e9d 100644 --- a/src/postgres_mcp/sql/sql_driver.py +++ b/src/postgres_mcp/sql/sql_driver.py @@ -256,15 +256,20 @@ async def execute_query( raise e 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 for objects not handled by default json package""" + """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() - raise TypeError(f"Type {type(obj)} not serializable") - """Calculates exact bytes of the JSON-serialized data including datetimes.""" + return str(obj) + buffer = io.StringIO() - # Add 'default=json_serial' to handle the PostgreSQL timestamps json.dump(data, buffer, default=json_serial) return len(buffer.getvalue().encode("utf-8")) 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) From d13b1345ddff804bacbc7d3bbe40f6348ef2f7bb Mon Sep 17 00:00:00 2001 From: Caleb Mabry Date: Mon, 23 Feb 2026 00:24:22 -0500 Subject: [PATCH 08/12] Load allowed hosts from MCP --- src/postgres_mcp/config.py | 11 +++++++++++ src/postgres_mcp/server.py | 6 +++++- 2 files changed, 16 insertions(+), 1 deletion(-) diff --git a/src/postgres_mcp/config.py b/src/postgres_mcp/config.py index ad66fcd..8c51271 100644 --- a/src/postgres_mcp/config.py +++ b/src/postgres_mcp/config.py @@ -42,6 +42,12 @@ def _load_config(self): ) 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,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: @@ -57,6 +63,11 @@ def default_page_size(self) -> int: 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.""" diff --git a/src/postgres_mcp/server.py b/src/postgres_mcp/server.py index 1a044b4..03b4cd9 100644 --- a/src/postgres_mcp/server.py +++ b/src/postgres_mcp/server.py @@ -13,6 +13,7 @@ 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 @@ -44,7 +45,10 @@ 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" From 2669dbf65d29bb74c83cdfe1db396aa876e769c0 Mon Sep 17 00:00:00 2001 From: Caleb Mabry Date: Mon, 23 Feb 2026 00:24:55 -0500 Subject: [PATCH 09/12] Include allowed_hosts for example --- .env.example | 3 +++ 1 file changed, 3 insertions(+) diff --git a/.env.example b/.env.example index a2b3d71..83726a1 100644 --- a/.env.example +++ b/.env.example @@ -10,3 +10,6 @@ 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, +ALLOWED_HOSTS=localhost,127.0.0.1 From 5eb35713f231e001c2d7ae9e27575e275306199d Mon Sep 17 00:00:00 2001 From: Caleb Mabry Date: Mon, 23 Feb 2026 00:31:20 -0500 Subject: [PATCH 10/12] Resolve to all all ports --- .env.example | 4 ++-- src/postgres_mcp/config.py | 2 +- src/postgres_mcp/server.py | 1 + 3 files changed, 4 insertions(+), 3 deletions(-) diff --git a/.env.example b/.env.example index 83726a1..f324fcc 100644 --- a/.env.example +++ b/.env.example @@ -11,5 +11,5 @@ 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, -ALLOWED_HOSTS=localhost,127.0.0.1 +# 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/src/postgres_mcp/config.py b/src/postgres_mcp/config.py index 8c51271..9b76e4a 100644 --- a/src/postgres_mcp/config.py +++ b/src/postgres_mcp/config.py @@ -43,7 +43,7 @@ def _load_config(self): 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,127.0.0.1") + 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: diff --git a/src/postgres_mcp/server.py b/src/postgres_mcp/server.py index 03b4cd9..4057c75 100644 --- a/src/postgres_mcp/server.py +++ b/src/postgres_mcp/server.py @@ -46,6 +46,7 @@ # Initialize FastMCP with default settings +print('FUCK YOU', config.allowed_hosts) mcp = FastMCP("postgres-mcp", transport_security=TransportSecuritySettings( allowed_hosts=config.allowed_hosts, )) From cad401dcd0f0445ad537208cbecab40b6a8ea895 Mon Sep 17 00:00:00 2001 From: Caleb Mabry Date: Mon, 23 Feb 2026 00:38:14 -0500 Subject: [PATCH 11/12] Lint --- src/postgres_mcp/config.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/postgres_mcp/config.py b/src/postgres_mcp/config.py index 9b76e4a..8e4aeeb 100644 --- a/src/postgres_mcp/config.py +++ b/src/postgres_mcp/config.py @@ -42,7 +42,7 @@ def _load_config(self): ) 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(",")] @@ -63,7 +63,7 @@ def default_page_size(self) -> int: 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.""" From e31216eff4d165c447bc5cf1e5f20c299d02d505 Mon Sep 17 00:00:00 2001 From: Caleb Mabry Date: Mon, 23 Feb 2026 00:43:25 -0500 Subject: [PATCH 12/12] Format --- src/postgres_mcp/server.py | 11 +++++++---- 1 file changed, 7 insertions(+), 4 deletions(-) diff --git a/src/postgres_mcp/server.py b/src/postgres_mcp/server.py index 4057c75..f190077 100644 --- a/src/postgres_mcp/server.py +++ b/src/postgres_mcp/server.py @@ -46,10 +46,13 @@ # Initialize FastMCP with default settings -print('FUCK YOU', config.allowed_hosts) -mcp = FastMCP("postgres-mcp", transport_security=TransportSecuritySettings( - allowed_hosts=config.allowed_hosts, -)) + +mcp = FastMCP( + "postgres-mcp", + transport_security=TransportSecuritySettings( + allowed_hosts=config.allowed_hosts, + ), +) # Constants PG_STAT_STATEMENTS = "pg_stat_statements"