/
alexefan136
/
flowstack
Обзор
Документация
Войти
/
alexefan136
/
flowstack
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
CI/CD
Аналитика
Безопасность
main
core/engine/src/tools/mcp/postgres.py
1 844 строки
63 KB
Alexander Efanov
Обновление репозитория
15 июл 2026, 12:19
15 июл 2026, 12:19
76704c6
Код
Авторство
О чём код?
""" PostgreSQL MCP Server — интеграция с PostgreSQL через MCP протокол. Предоставляет tools для работы с PostgreSQL базой данных: - Выполнение SQL запросов (SELECT, INSERT, UPDATE, DELETE) - Read-only режим для безопасного доступа - Интроспекция схемы (schemas, tables, columns, indexes) - Статистика таблиц и БД - EXPLAIN для анализа планов выполнения Используемый драйвер: - asyncpg (async PostgreSQL driver): https://magicstack.github.io/asyncpg/ Аутентификация: - Connection string (postgres://user:password@host:port/database) - Отдельные параметры (host, port, user, password, database) Архитектурные принципы: - Чистые функции для преобразования данных - Type-safe dispatch через таблицу handlers - Явная обработка ошибок - Async-first с connection pooling - Read-only mode для безопасности - Prepared statements (защита от SQL injection) - Следование "You Might Not Need an Effect" Безопасность: - Read-only mode блокирует write операции (INSERT, UPDATE, DELETE, DROP, etc.) - Max rows limit для предотвращения OOM на больших таблицах - Query timeout для блокировки долгих запросов - Whitelist schemas для ограничения доступа - Prepared statements через asyncpg (защита от SQL injection) """ from __future__ import annotations import asyncio import logging import os import re from collections.abc import AsyncIterator from contextlib import asynccontextmanager from dataclasses import dataclass, field from datetime import datetime from typing import Any, TYPE_CHECKING, cast from urllib.parse import urlparse from src.primitives.context import MCPContext if TYPE_CHECKING: import asyncpg # type: ignore[import-not-found,import-untyped] logger = logging.getLogger(__name__) # ============================================================================ # Configuration # ============================================================================ @dataclass class PostgresConfig: """ Конфигурация подключения к PostgreSQL. Можно задать либо connection_string, либо отдельные параметры. """ # Connection string (приоритетнее отдельных параметров) connection_string: str | None = None # Отдельные параметры (используются если connection_string не задан) host: str = "localhost" port: int = 5432 user: str | None = None password: str | None = None database: str | None = None # Connection pool min_pool_size: int = 2 max_pool_size: int = 10 # Таймауты connect_timeout_seconds: float = 10.0 query_timeout_seconds: float = 30.0 # Безопасность read_only: bool = False # Если True — запрещены write операции max_rows: int = 1000 # Максимум строк для SELECT allowed_schemas: list[str] | None = None # None = все схемы разрешены blocked_keywords: list[str] = field(default_factory=lambda: [ "DROP", "TRUNCATE", "ALTER", "CREATE", "GRANT", "REVOKE", ]) # Retry max_retries: int = 3 retry_delay_seconds: float = 1.0 def get_dsn(self) -> str: """ Получить DSN (connection string) для asyncpg. Чистая функция — не мутирует состояние. """ if self.connection_string: return self.connection_string # Собираем DSN из отдельных параметров user_part = self.user or "" password_part = f":{self.password}" if self.password else "" auth = f"{user_part}{password_part}@" if user_part else "" db_part = f"/{self.database}" if self.database else "" return f"postgresql://{auth}{self.host}:{self.port}{db_part}" def get_database_name(self) -> str | None: """Извлечь имя базы из DSN.""" dsn = self.get_dsn() try: parsed = urlparse(dsn) path = parsed.path if path and path.startswith("/"): return path[1:] or None return None except Exception: return self.database def validate(self) -> list[str]: """Валидировать конфигурацию. Чистая функция.""" errors: list[str] = [] if not self.connection_string: if not self.host: errors.append("host is required when connection_string is not set") if self.port <= 0 or self.port > 65535: errors.append("port must be between 1 and 65535") if self.max_rows <= 0: errors.append("max_rows must be positive") if self.query_timeout_seconds <= 0: errors.append("query_timeout_seconds must be positive") if self.connect_timeout_seconds <= 0: errors.append("connect_timeout_seconds must be positive") if self.min_pool_size < 0: errors.append("min_pool_size must be non-negative") if self.max_pool_size <= 0: errors.append("max_pool_size must be positive") if self.min_pool_size > self.max_pool_size: errors.append("min_pool_size cannot be greater than max_pool_size") return errors @classmethod def from_env(cls, prefix: str = "POSTGRES_") -> PostgresConfig: """Создать из переменных окружения.""" port_str = os.getenv(f"{prefix}PORT", "5432") try: port = int(port_str) except ValueError: port = 5432 return cls( connection_string=os.getenv(f"{prefix}URL") or os.getenv(f"{prefix}DSN"), host=os.getenv(f"{prefix}HOST", "localhost"), port=port, user=os.getenv(f"{prefix}USER"), password=os.getenv(f"{prefix}PASSWORD"), database=os.getenv(f"{prefix}DATABASE") or os.getenv(f"{prefix}DB"), min_pool_size=int(os.getenv(f"{prefix}MIN_POOL_SIZE", "2")), max_pool_size=int(os.getenv(f"{prefix}MAX_POOL_SIZE", "10")), connect_timeout_seconds=float(os.getenv(f"{prefix}CONNECT_TIMEOUT", "10")), query_timeout_seconds=float(os.getenv(f"{prefix}QUERY_TIMEOUT", "30")), read_only=os.getenv(f"{prefix}READ_ONLY", "false").lower() == "true", max_rows=int(os.getenv(f"{prefix}MAX_ROWS", "1000")), ) # ============================================================================ # Exceptions # ============================================================================ class PostgresError(Exception): """Базовое исключение PostgreSQL.""" pass class PostgresConnectionError(PostgresError): """Ошибка подключения к БД.""" pass class PostgresAuthError(PostgresError): """Ошибка аутентификации.""" pass class PostgresQueryError(PostgresError): """Ошибка выполнения запроса.""" def __init__( self, message: str, sqlstate: str | None = None, detail: str | None = None, hint: str | None = None, ): super().__init__(message) self.sqlstate = sqlstate self.detail = detail self.hint = hint class PostgresReadOnlyError(PostgresError): """Попытка write операции в read-only режиме.""" pass class PostgresTimeoutError(PostgresError): """Превышен таймаут запроса.""" pass class PostgresSchemaBlockedError(PostgresError): """Доступ к схеме запрещен.""" pass class PostgresRowLimitError(PostgresError): """Превышен лимит строк.""" pass # ============================================================================ # Response Models # ============================================================================ @dataclass class PostgresColumn: """Колонка таблицы.""" name: str data_type: str udt_name: str = "" is_nullable: bool = True column_default: str | None = None ordinal_position: int = 0 character_maximum_length: int | None = None numeric_precision: int | None = None numeric_scale: int | None = None is_identity: bool = False is_generated: str = "NEVER" # NEVER | ALWAYS comment: str | None = None def to_dict(self) -> dict[str, Any]: """Чистая функция.""" return { "name": self.name, "data_type": self.data_type, "udt_name": self.udt_name, "is_nullable": self.is_nullable, "column_default": self.column_default, "ordinal_position": self.ordinal_position, "character_maximum_length": self.character_maximum_length, "numeric_precision": self.numeric_precision, "numeric_scale": self.numeric_scale, "is_identity": self.is_identity, "is_generated": self.is_generated, "comment": self.comment, } @classmethod def from_row(cls, row: dict[str, Any]) -> PostgresColumn: """Чистая функция.""" char_len = row.get("character_maximum_length") num_prec = row.get("numeric_precision") num_scale = row.get("numeric_scale") return cls( name=str(row.get("column_name", "")), data_type=str(row.get("data_type", "")), udt_name=str(row.get("udt_name", "")), is_nullable=str(row.get("is_nullable", "YES")) == "YES", column_default=row.get("column_default"), ordinal_position=int(row.get("ordinal_position", 0)), character_maximum_length=( int(char_len) if isinstance(char_len, (int, float)) else None ), numeric_precision=( int(num_prec) if isinstance(num_prec, (int, float)) else None ), numeric_scale=( int(num_scale) if isinstance(num_scale, (int, float)) else None ), is_identity=str(row.get("is_identity", "NO")) == "YES", is_generated=str(row.get("is_generated", "NEVER")), comment=row.get("comment"), ) @dataclass class PostgresIndex: """Индекс таблицы.""" name: str table_name: str schema_name: str = "public" is_unique: bool = False is_primary: bool = False index_type: str = "btree" # btree | hash | gist | gin | spgist | brin columns: list[str] = field(default_factory=list) definition: str = "" def to_dict(self) -> dict[str, Any]: """Чистая функция.""" return { "name": self.name, "table_name": self.table_name, "schema_name": self.schema_name, "is_unique": self.is_unique, "is_primary": self.is_primary, "index_type": self.index_type, "columns": self.columns, "definition": self.definition, } @classmethod def from_row(cls, row: dict[str, Any]) -> PostgresIndex: """Чистая функция.""" columns_raw = row.get("columns") if isinstance(columns_raw, (list, tuple)): columns = [str(c) for c in columns_raw] elif isinstance(columns_raw, str): # PostgreSQL возвращает array как строку "{a,b,c}" stripped = columns_raw.strip("{}") columns = [c.strip() for c in stripped.split(",") if c.strip()] else: columns = [] return cls( name=str(row.get("indexname", "")), table_name=str(row.get("tablename", "")), schema_name=str(row.get("schemaname", "public")), is_unique=bool(row.get("is_unique", False)), is_primary=bool(row.get("is_primary", False)), index_type=str(row.get("index_type", "btree")), columns=columns, definition=str(row.get("indexdef", "") or ""), ) @dataclass class PostgresTable: """Таблица в БД.""" schema_name: str table_name: str table_type: str = "BASE TABLE" # BASE TABLE | VIEW row_estimate: int = 0 total_size_bytes: int = 0 comment: str | None = None def to_dict(self) -> dict[str, Any]: """Чистая функция.""" return { "schema_name": self.schema_name, "table_name": self.table_name, "table_type": self.table_type, "row_estimate": self.row_estimate, "total_size_bytes": self.total_size_bytes, "total_size_human": _format_bytes(self.total_size_bytes), "comment": self.comment, } @classmethod def from_row(cls, row: dict[str, Any]) -> PostgresTable: """Чистая функция.""" row_estimate_raw = row.get("row_estimate") row_estimate = 0 if isinstance(row_estimate_raw, (int, float)): row_estimate = int(row_estimate_raw) size_raw = row.get("total_size_bytes") total_size_bytes = 0 if isinstance(size_raw, (int, float)): total_size_bytes = int(size_raw) return cls( schema_name=str(row.get("schema_name", "public")), table_name=str(row.get("table_name", "")), table_type=str(row.get("table_type", "BASE TABLE")), row_estimate=row_estimate, total_size_bytes=total_size_bytes, comment=row.get("comment"), ) @dataclass class PostgresSchema: """Схема в БД.""" name: str owner: str = "" is_system: bool = False comment: str | None = None def to_dict(self) -> dict[str, Any]: """Чистая функция.""" return { "name": self.name, "owner": self.owner, "is_system": self.is_system, "comment": self.comment, } @classmethod def from_row(cls, row: dict[str, Any]) -> PostgresSchema: """Чистая функция.""" name = str(row.get("schema_name", "")) # Системные схемы PostgreSQL system_schemas = { "pg_catalog", "information_schema", "pg_toast", "pg_temp_1", "pg_toast_temp_1", } is_system = name in system_schemas or name.startswith("pg_temp") or name.startswith("pg_toast_temp") return cls( name=name, owner=str(row.get("schema_owner", "") or ""), is_system=is_system, comment=row.get("comment"), ) @dataclass class PostgresDatabaseInfo: """Информация о базе данных.""" name: str version: str server_version_num: int = 0 size_bytes: int = 0 encoding: str = "" collation: str = "" connection_count: int = 0 max_connections: int = 0 uptime_seconds: float = 0.0 def to_dict(self) -> dict[str, Any]: """Чистая функция.""" return { "name": self.name, "version": self.version, "server_version_num": self.server_version_num, "size_bytes": self.size_bytes, "size_human": _format_bytes(self.size_bytes), "encoding": self.encoding, "collation": self.collation, "connection_count": self.connection_count, "max_connections": self.max_connections, "uptime_seconds": self.uptime_seconds, "uptime_human": _format_seconds(self.uptime_seconds), } @dataclass class PostgresQueryResult: """Результат выполнения запроса.""" rows: list[dict[str, Any]] = field(default_factory=list) row_count: int = 0 columns: list[str] = field(default_factory=list) command: str = "" # SELECT, INSERT, UPDATE, DELETE, etc. affected_rows: int | None = None # Для INSERT/UPDATE/DELETE execution_time_ms: float = 0.0 truncated: bool = False # True если результат обрезан из-за max_rows explain_output: str | None = None # Для EXPLAIN запросов def to_dict(self) -> dict[str, Any]: """Чистая функция.""" return { "rows": self.rows, "row_count": self.row_count, "columns": self.columns, "command": self.command, "affected_rows": self.affected_rows, "execution_time_ms": self.execution_time_ms, "truncated": self.truncated, "explain_output": self.explain_output, } @dataclass class PostgresExplainResult: """Результат EXPLAIN.""" query: str plan: str format: str = "text" # text | json | yaml | xml analyze: bool = False buffers: bool = False def to_dict(self) -> dict[str, Any]: """Чистая функция.""" return { "query": self.query, "plan": self.plan, "format": self.format, "analyze": self.analyze, "buffers": self.buffers, } # ============================================================================ # Helper Functions # ============================================================================ def _format_bytes(size_bytes: int) -> str: """Форматировать размер в человекочитаемый вид. Чистая функция.""" if size_bytes <= 0: return "0 B" units = ["B", "KB", "MB", "GB", "TB", "PB"] size = float(size_bytes) unit_index = 0 while size >= 1024.0 and unit_index < len(units) - 1: size /= 1024.0 unit_index += 1 return f"{size:.2f} {units[unit_index]}" def _format_seconds(seconds: float) -> str: """Форматировать секунды в человекочитаемый вид. Чистая функция.""" if seconds <= 0: return "0s" if seconds < 60: return f"{seconds:.1f}s" if seconds < 3600: return f"{seconds / 60:.1f}m" if seconds < 86400: return f"{seconds / 3600:.1f}h" return f"{seconds / 86400:.1f}d" _WRITE_KEYWORDS = [ r"\bINSERT\b", r"\bUPDATE\b", r"\bDELETE\b", r"\bDROP\b", r"\bTRUNCATE\b", r"\bALTER\b", r"\bCREATE\b", r"\bGRANT\b", r"\bREVOKE\b", r"\bMERGE\b", r"\bCOPY\b", r"\bVACUUM\b", r"\bREINDEX\b", r"\bCLUSTER\b", r"\bREFRESH\b", ] def _is_write_query(sql: str) -> bool: """ Проверить, является ли запрос write-операцией. Чистая функция — использует regex для поиска ключевых слов. Не идеально (может не учесть CTE с RETURNING), но достаточно для защиты. """ sql_upper = sql.upper() for pattern in _WRITE_KEYWORDS: if re.search(pattern, sql_upper): return True return False def _extract_schemas_from_sql(sql: str) -> list[str]: """ Извлечь имена схем из SQL запроса. Чистая функция — парсит паттерны schema.table и schema."table". """ # Паттерн: identifier.identifier pattern = r"\b([a-zA-Z_][a-zA-Z0-9_]*)\s*\.\s*[\"']?([a-zA-Z_][a-zA-Z0-9_]*)[\"']?" matches = re.findall(pattern, sql) # Исключаем PostgreSQL ключевые слова которые могут быть восприняты как схемы keywords = { "pg_catalog", "information_schema", "public", "SELECT", "FROM", "WHERE", "JOIN", "ON", "AND", "OR", "INSERT", "INTO", "VALUES", "UPDATE", "SET", "DELETE", } schemas = set() for match in matches: candidate = match[0] if candidate.lower() not in {k.lower() for k in keywords}: schemas.add(candidate) return list(schemas) def _row_to_dict(row: Any) -> dict[str, Any]: """ Преобразовать asyncpg.Record в dict. Обрабатывает специфичные PostgreSQL типы: - datetime → ISO строка - bytes → base64 строка - Decimal → float - UUID → строка """ import base64 import decimal import uuid from typing import Iterable, cast if hasattr(row, "keys") and callable(row.keys): # asyncpg.Record — имеет метод keys() result: dict[str, Any] = {} # Получаем keys и явно конвертируем в list для type safety keys_callable = row.keys if callable(keys_callable): keys_result = keys_callable() # Cast к Iterable чтобы Pylance знал что это итерируемо keys_iterable = cast(Iterable[Any], keys_result) keys_list: list[Any] = list(keys_iterable) for key in keys_list: key_str = str(key) value = row[key] result[key_str] = _serialize_value(value, base64, decimal, uuid) return result elif isinstance(row, dict): return { str(k): _serialize_value(v, base64, decimal, uuid) for k, v in row.items() } else: return {"value": _serialize_value(row, base64, decimal, uuid)} def _serialize_value( value: Any, base64_mod: Any, decimal_mod: Any, uuid_mod: Any, ) -> Any: """Сериализовать значение для JSON. Чистая функция.""" if value is None: return None if isinstance(value, datetime): return value.isoformat() if isinstance(value, bytes): return base64_mod.b64encode(value).decode("ascii") if isinstance(value, decimal_mod.Decimal): return float(value) if isinstance(value, uuid_mod.UUID): return str(value) if isinstance(value, (list, tuple)): return [_serialize_value(v, base64_mod, decimal_mod, uuid_mod) for v in value] if isinstance(value, dict): return { k: _serialize_value(v, base64_mod, decimal_mod, uuid_mod) for k, v in value.items() } if isinstance(value, (str, int, float, bool)): return value # Fallback return str(value) # ============================================================================ # Postgres Client # ============================================================================ class PostgresClient: """ Асинхронный клиент для PostgreSQL с connection pooling. Автоматически обрабатывает: - Connection pooling через asyncpg.create_pool - Prepared statements (защита от SQL injection) - Read-only mode - Max rows limit - Query timeouts - Schema whitelist """ def __init__(self, config: PostgresConfig): self.config = config self._pool: asyncpg.Pool | None = None async def __aenter__(self) -> PostgresClient: await self.connect() return self async def __aexit__(self, exc_type, exc_val, exc_tb) -> None: await self.close() async def connect(self) -> None: """Создать connection pool.""" if self._pool is not None: return try: import asyncpg # type: ignore[import-not-found,import-untyped] except ImportError as e: raise PostgresConnectionError( "asyncpg is required. Install with: pip install asyncpg" ) from e errors = self.config.validate() if errors: raise PostgresConnectionError( f"Invalid configuration: {', '.join(errors)}" ) dsn = self.config.get_dsn() try: self._pool = await asyncpg.create_pool( dsn=dsn, min_size=self.config.min_pool_size, max_size=self.config.max_pool_size, timeout=self.config.connect_timeout_seconds, command_timeout=self.config.query_timeout_seconds, ) logger.info(f"Connected to PostgreSQL: {self._mask_dsn(dsn)}") except asyncpg.InvalidPasswordError as e: raise PostgresAuthError(f"Authentication failed: {e}") from e except asyncpg.CannotConnectNowError as e: raise PostgresConnectionError(f"Cannot connect: {e}") from e except asyncpg.PostgresError as e: raise PostgresConnectionError(f"PostgreSQL error: {e}") from e except Exception as e: raise PostgresConnectionError(f"Failed to connect: {e}") from e async def close(self) -> None: """Закрыть connection pool.""" if self._pool is not None: await self._pool.close() self._pool = None def _require_pool(self) -> asyncpg.Pool: """Получить активный pool или выбросить ошибку.""" if self._pool is None: raise PostgresConnectionError("Client is not connected") return self._pool def _mask_dsn(self, dsn: str) -> str: """Замаскировать пароль в DSN для логирования. Чистая функция.""" try: parsed = urlparse(dsn) if parsed.password: # Заменяем пароль на *** masked = dsn.replace(f":{parsed.password}@", ":***@") return masked return dsn except Exception: return "postgresql://***" def _check_read_only(self, sql: str) -> None: """Проверить, разрешен ли запрос в read-only режиме.""" if not self.config.read_only: return if _is_write_query(sql): raise PostgresReadOnlyError( "Write operations are not allowed in read-only mode" ) def _check_schemas(self, sql: str) -> None: """Проверить, разрешены ли схемы в запросе.""" if self.config.allowed_schemas is None: return schemas_in_query = _extract_schemas_from_sql(sql) allowed_lower = {s.lower() for s in self.config.allowed_schemas} # Всегда разрешаем системные схемы allowed_lower.update({"pg_catalog", "information_schema", "public"}) for schema in schemas_in_query: if schema.lower() not in allowed_lower: raise PostgresSchemaBlockedError( f"Access to schema '{schema}' is not allowed. " f"Allowed schemas: {self.config.allowed_schemas}" ) async def execute_query( self, sql: str, params: list[Any] | None = None, max_rows: int | None = None, timeout_seconds: float | None = None, ) -> PostgresQueryResult: """ Выполнить SQL запрос. Args: sql: SQL запрос params: Параметры для prepared statement max_rows: Максимум строк (переопределяет config.max_rows) timeout_seconds: Таймаут (переопределяет config.query_timeout_seconds) Returns: PostgresQueryResult с результатами """ self._check_read_only(sql) self._check_schemas(sql) pool = self._require_pool() effective_max_rows = ( max_rows if max_rows is not None else self.config.max_rows ) effective_timeout = ( timeout_seconds if timeout_seconds is not None else self.config.query_timeout_seconds ) start_time = asyncio.get_event_loop().time() rows: list[dict[str, Any]] = [] columns: list[str] = [] command = "" affected_rows: int | None = None truncated = False try: async with pool.acquire() as connection: # Определяем тип запроса sql_stripped = sql.strip().upper() is_select = sql_stripped.startswith(( "SELECT", "WITH", "TABLE", "VALUES", "EXPLAIN", "SHOW", "DESCRIBE", )) if is_select: # Используем cursor для ограничения строк stmt = await connection.prepare(sql) # Извлекаем имена колонок if hasattr(stmt, "get_attributes"): attrs = stmt.get_attributes() columns = [ str(a.name) for a in attrs if hasattr(a, "name") ] else: columns = [] # Создаем cursor с явной типизацией через cast # asyncpg cursor factory — это async iterator, но Pylance # видит его как object из-за type: ignore на импорте if params: cursor_factory = stmt.cursor(*params) else: cursor_factory = stmt.cursor() # Cast к Any для корректной работы async for cursor_iter: Any = cursor_factory count = 0 async for record in cursor_iter: if count >= effective_max_rows: truncated = True break rows.append(_row_to_dict(record)) count += 1 command = "SELECT" affected_rows = None else: # Write операция — используем execute if params: result = await connection.execute( sql, *params, timeout=effective_timeout ) else: result = await connection.execute( sql, timeout=effective_timeout ) # result — это строка вида "INSERT 0 1" или "UPDATE 5" result_str = str(result) if result is not None else "UNKNOWN" parts = result_str.split() command = parts[0] if parts else "UNKNOWN" if len(parts) >= 2: try: affected_rows = int(parts[-1]) except ValueError: affected_rows = None execution_time_ms = (asyncio.get_event_loop().time() - start_time) * 1000 return PostgresQueryResult( rows=rows, row_count=len(rows), columns=columns, command=command, affected_rows=affected_rows, execution_time_ms=execution_time_ms, truncated=truncated, ) except asyncio.TimeoutError as e: raise PostgresTimeoutError( f"Query timed out after {effective_timeout}s" ) from e except Exception as e: # asyncpg ошибки error_type = type(e).__name__ sqlstate = getattr(e, "sqlstate", None) detail = getattr(e, "detail", None) hint = getattr(e, "hint", None) raise PostgresQueryError( f"Query failed ({error_type}): {e}", sqlstate=sqlstate, detail=detail, hint=hint, ) from e async def explain_query( self, sql: str, analyze: bool = False, buffers: bool = False, format: str = "text", ) -> PostgresExplainResult: """ Выполнить EXPLAIN для запроса. Args: sql: SQL запрос для анализа analyze: Фактически выполнить запрос (EXPLAIN ANALYZE) buffers: Показать информацию о буферах format: text | json | yaml | xml """ if analyze and self.config.read_only: # EXPLAIN ANALYZE выполняет запрос — запрещаем в read-only raise PostgresReadOnlyError( "EXPLAIN ANALYZE is not allowed in read-only mode" ) options: list[str] = [] if analyze: options.append("ANALYZE") if buffers: options.append("BUFFERS") if format != "text": options.append(f"FORMAT {format.upper()}") options_str = ", ".join(options) explain_sql = f"EXPLAIN ({options_str}) {sql}" if options_str else f"EXPLAIN {sql}" result = await self.execute_query(explain_sql) plan_lines = [] for row in result.rows: # EXPLAIN возвращает колонку "QUERY PLAN" plan_value = row.get("QUERY PLAN") if plan_value is None: # Fallback — берем первое значение for v in row.values(): plan_value = v break if plan_value is not None: plan_lines.append(str(plan_value)) return PostgresExplainResult( query=sql, plan="\n".join(plan_lines), format=format, analyze=analyze, buffers=buffers, ) async def list_schemas(self) -> list[PostgresSchema]: """Список схем в БД.""" sql = """ SELECT s.schema_name, s.schema_owner, obj_description((s.schema_name)::regnamespace, 'pg_namespace') as comment FROM information_schema.schemata s ORDER BY s.schema_name """ result = await self.execute_query(sql, max_rows=1000) return [PostgresSchema.from_row(r) for r in result.rows] async def list_tables( self, schema: str = "public", include_views: bool = True, ) -> list[PostgresTable]: """Список таблиц и view в схеме.""" self._check_schemas(f"SELECT * FROM {schema}.dummy") types_filter = "'BASE TABLE'" if include_views: types_filter = "'BASE TABLE', 'VIEW'" sql = f""" SELECT t.table_schema as schema_name, t.table_name, t.table_type, COALESCE(c.reltuples::bigint, 0) as row_estimate, COALESCE(pg_total_relation_size(c.oid), 0) as total_size_bytes, obj_description(c.oid) as comment FROM information_schema.tables t LEFT JOIN pg_class c ON c.relname = t.table_name LEFT JOIN pg_namespace n ON n.oid = c.relnamespace AND n.nspname = t.table_schema WHERE t.table_schema = $1 AND t.table_type IN ({types_filter}) ORDER BY t.table_name """ result = await self.execute_query(sql, params=[schema], max_rows=1000) return [PostgresTable.from_row(r) for r in result.rows] async def describe_table( self, table_name: str, schema: str = "public" ) -> list[PostgresColumn]: """Получить описание таблицы (колонки).""" self._check_schemas(f"SELECT * FROM {schema}.{table_name}") sql = """ SELECT c.column_name, c.data_type, c.udt_name, c.is_nullable, c.column_default, c.ordinal_position, c.character_maximum_length, c.numeric_precision, c.numeric_scale, c.is_identity, c.is_generated, pgd.description as comment FROM information_schema.columns c LEFT JOIN pg_catalog.pg_statio_all_tables st ON c.table_schema = st.schemaname AND c.table_name = st.relname LEFT JOIN pg_catalog.pg_description pgd ON pgd.objoid = st.relid AND pgd.objsubid = c.ordinal_position WHERE c.table_schema = $1 AND c.table_name = $2 ORDER BY c.ordinal_position """ result = await self.execute_query( sql, params=[schema, table_name], max_rows=1000 ) return [PostgresColumn.from_row(r) for r in result.rows] async def list_indexes( self, table_name: str, schema: str = "public" ) -> list[PostgresIndex]: """Список индексов таблицы.""" self._check_schemas(f"SELECT * FROM {schema}.{table_name}") sql = """ SELECT i.schemaname, i.tablename, i.indexname, i.indexdef, ix.indisunique as is_unique, ix.indisprimary as is_primary, am.amname as index_type, array_agg(a.attname ORDER BY array_position(ix.indkey, a.attnum)) as columns FROM pg_indexes i JOIN pg_class c ON c.relname = i.indexname JOIN pg_index ix ON ix.indexrelid = c.oid JOIN pg_am am ON am.oid = c.relam JOIN pg_attribute a ON a.attrelid = ix.indrelid AND a.attnum = ANY(ix.indkey) WHERE i.schemaname = $1 AND i.tablename = $2 GROUP BY i.schemaname, i.tablename, i.indexname, i.indexdef, ix.indisunique, ix.indisprimary, am.amname ORDER BY i.indexname """ result = await self.execute_query( sql, params=[schema, table_name], max_rows=1000 ) return [PostgresIndex.from_row(r) for r in result.rows] async def get_table_stats( self, table_name: str, schema: str = "public" ) -> dict[str, Any]: """Получить статистику таблицы.""" self._check_schemas(f"SELECT * FROM {schema}.{table_name}") sql = """ SELECT c.relname as table_name, n.nspname as schema_name, c.reltuples::bigint as row_estimate, pg_total_relation_size(c.oid) as total_size_bytes, pg_relation_size(c.oid) as table_size_bytes, pg_indexes_size(c.oid) as indexes_size_bytes, pg_stat_get_live_tuples(c.oid) as live_tuples, pg_stat_get_dead_tuples(c.oid) as dead_tuples, s.seq_scan, s.seq_tup_read, s.idx_scan, s.idx_tup_fetch, s.n_tup_ins, s.n_tup_upd, s.n_tup_del, s.last_vacuum, s.last_analyze, s.last_autovacuum, s.last_autoanalyze FROM pg_class c JOIN pg_namespace n ON n.oid = c.relnamespace LEFT JOIN pg_stat_user_tables s ON s.relid = c.oid WHERE n.nspname = $1 AND c.relname = $2 """ result = await self.execute_query(sql, params=[schema, table_name]) if not result.rows: raise PostgresQueryError( f"Table {schema}.{table_name} not found" ) row = result.rows[0] total_bytes = int(row.get("total_size_bytes", 0) or 0) table_bytes = int(row.get("table_size_bytes", 0) or 0) indexes_bytes = int(row.get("indexes_size_bytes", 0) or 0) def _to_int(v: Any) -> int: if isinstance(v, (int, float)): return int(v) return 0 def _to_str(v: Any) -> str | None: if isinstance(v, str): return v if isinstance(v, datetime): return v.isoformat() return None return { "table_name": str(row.get("table_name", "")), "schema_name": str(row.get("schema_name", "")), "row_estimate": _to_int(row.get("row_estimate")), "live_tuples": _to_int(row.get("live_tuples")), "dead_tuples": _to_int(row.get("dead_tuples")), "total_size_bytes": total_bytes, "total_size_human": _format_bytes(total_bytes), "table_size_bytes": table_bytes, "table_size_human": _format_bytes(table_bytes), "indexes_size_bytes": indexes_bytes, "indexes_size_human": _format_bytes(indexes_bytes), "seq_scan": _to_int(row.get("seq_scan")), "seq_tup_read": _to_int(row.get("seq_tup_read")), "idx_scan": _to_int(row.get("idx_scan")), "idx_tup_fetch": _to_int(row.get("idx_tup_fetch")), "n_tup_ins": _to_int(row.get("n_tup_ins")), "n_tup_upd": _to_int(row.get("n_tup_upd")), "n_tup_del": _to_int(row.get("n_tup_del")), "last_vacuum": _to_str(row.get("last_vacuum")), "last_analyze": _to_str(row.get("last_analyze")), "last_autovacuum": _to_str(row.get("last_autovacuum")), "last_autoanalyze": _to_str(row.get("last_autoanalyze")), } async def get_database_info(self) -> PostgresDatabaseInfo: """Получить информацию о базе данных.""" database = self.config.get_database_name() or "" # Версия PostgreSQL version_result = await self.execute_query( "SELECT version(), current_setting('server_version_num')::int as version_num" ) version = "" version_num = 0 if version_result.rows: row = version_result.rows[0] version = str(row.get("version", "")) vn = row.get("version_num") if isinstance(vn, (int, float)): version_num = int(vn) # Encoding и collation encoding_result = await self.execute_query(""" SELECT pg_encoding_to_char(encoding) as encoding, datcollate as collation FROM pg_database WHERE datname = current_database() """) encoding = "" collation = "" if encoding_result.rows: row = encoding_result.rows[0] encoding = str(row.get("encoding", "") or "") collation = str(row.get("collation", "") or "") # Размер БД size_result = await self.execute_query( "SELECT pg_database_size(current_database()) as size_bytes" ) size_bytes = 0 if size_result.rows: size_val = size_result.rows[0].get("size_bytes") if isinstance(size_val, (int, float)): size_bytes = int(size_val) # Connections conn_result = await self.execute_query(""" SELECT count(*)::int as connection_count, current_setting('max_connections')::int as max_connections FROM pg_stat_activity """) connection_count = 0 max_connections = 0 if conn_result.rows: row = conn_result.rows[0] cc = row.get("connection_count") mc = row.get("max_connections") if isinstance(cc, (int, float)): connection_count = int(cc) if isinstance(mc, (int, float)): max_connections = int(mc) # Uptime uptime_result = await self.execute_query(""" SELECT extract(epoch from (now() - pg_postmaster_start_time()))::float as uptime_seconds """) uptime_seconds = 0.0 if uptime_result.rows: u = uptime_result.rows[0].get("uptime_seconds") if isinstance(u, (int, float)): uptime_seconds = float(u) return PostgresDatabaseInfo( name=database, version=version, server_version_num=version_num, size_bytes=size_bytes, encoding=encoding, collation=collation, connection_count=connection_count, max_connections=max_connections, uptime_seconds=uptime_seconds, ) # ============================================================================ # MCP Tools Definition # ============================================================================ POSTGRES_TOOLS: list[dict[str, Any]] = [ { "name": "postgres_execute_query", "description": ( "Выполнить SQL запрос к PostgreSQL. " "Поддерживает SELECT, INSERT, UPDATE, DELETE. " "Использует prepared statements (параметры через $1, $2, ...)." ), "parameters": { "type": "object", "properties": { "sql": { "type": "string", "description": "SQL запрос. Параметры обозначаются как $1, $2, ...", }, "params": { "type": "array", "items": {}, "description": "Значения параметров для prepared statement", }, "max_rows": { "type": "integer", "description": "Максимум строк для SELECT (по умолчанию из конфига)", }, "timeout_seconds": { "type": "number", "description": "Таймаут выполнения в секундах", }, }, "required": ["sql"], }, }, { "name": "postgres_execute_readonly", "description": ( "Выполнить только SELECT запрос (безопасный режим). " "Write операции (INSERT/UPDATE/DELETE/DROP) запрещены." ), "parameters": { "type": "object", "properties": { "sql": { "type": "string", "description": "SELECT SQL запрос", }, "params": { "type": "array", "items": {}, "description": "Значения параметров", }, "max_rows": { "type": "integer", "default": 100, "description": "Максимум строк (по умолчанию 100)", }, }, "required": ["sql"], }, }, { "name": "postgres_explain_query", "description": ( "Показать план выполнения SQL запроса (EXPLAIN). " "Полезно для анализа производительности запросов." ), "parameters": { "type": "object", "properties": { "sql": {"type": "string", "description": "SQL запрос для анализа"}, "analyze": { "type": "boolean", "default": False, "description": "Фактически выполнить запрос (EXPLAIN ANALYZE)", }, "buffers": { "type": "boolean", "default": False, "description": "Показать информацию о буферах", }, "format": { "type": "string", "enum": ["text", "json", "yaml", "xml"], "default": "text", }, }, "required": ["sql"], }, }, { "name": "postgres_list_schemas", "description": "Получить список схем в базе данных", "parameters": {"type": "object", "properties": {}}, }, { "name": "postgres_list_tables", "description": "Получить список таблиц (и view) в схеме", "parameters": { "type": "object", "properties": { "schema": { "type": "string", "default": "public", "description": "Имя схемы", }, "include_views": { "type": "boolean", "default": True, "description": "Включать ли VIEW в результат", }, }, }, }, { "name": "postgres_describe_table", "description": "Получить описание таблицы (колонки, типы, nullable, defaults)", "parameters": { "type": "object", "properties": { "table_name": {"type": "string", "description": "Имя таблицы"}, "schema": { "type": "string", "default": "public", "description": "Имя схемы", }, }, "required": ["table_name"], }, }, { "name": "postgres_list_indexes", "description": "Получить список индексов таблицы", "parameters": { "type": "object", "properties": { "table_name": {"type": "string"}, "schema": {"type": "string", "default": "public"}, }, "required": ["table_name"], }, }, { "name": "postgres_get_table_stats", "description": ( "Получить статистику таблицы: размер, количество строк, " "dead tuples, время последнего vacuum/analyze" ), "parameters": { "type": "object", "properties": { "table_name": {"type": "string"}, "schema": {"type": "string", "default": "public"}, }, "required": ["table_name"], }, }, { "name": "postgres_get_database_info", "description": ( "Получить информацию о базе данных: версия PostgreSQL, " "размер, encoding, количество соединений, uptime" ), "parameters": {"type": "object", "properties": {}}, }, ] # ============================================================================ # Tool Handlers # ============================================================================ async def _handle_execute_query( client: PostgresClient, params: dict[str, Any] ) -> dict[str, Any]: sql_raw = params.get("sql") if not isinstance(sql_raw, str): raise PostgresQueryError("sql parameter must be a string") params_raw = params.get("params") query_params: list[Any] | None = None if isinstance(params_raw, list): query_params = params_raw max_rows_raw = params.get("max_rows") max_rows: int | None = None if isinstance(max_rows_raw, int): max_rows = max_rows_raw timeout_raw = params.get("timeout_seconds") timeout: float | None = None if isinstance(timeout_raw, (int, float)): timeout = float(timeout_raw) result = await client.execute_query( sql=sql_raw, params=query_params, max_rows=max_rows, timeout_seconds=timeout, ) return result.to_dict() async def _handle_execute_readonly( client: PostgresClient, params: dict[str, Any] ) -> dict[str, Any]: sql_raw = params.get("sql") if not isinstance(sql_raw, str): raise PostgresQueryError("sql parameter must be a string") # Строгая проверка — только SELECT sql_upper = sql_raw.strip().upper() if not sql_upper.startswith(("SELECT", "WITH", "TABLE", "VALUES")): raise PostgresReadOnlyError( "Only SELECT queries are allowed in readonly mode" ) params_raw = params.get("params") query_params: list[Any] | None = None if isinstance(params_raw, list): query_params = params_raw max_rows_raw = params.get("max_rows") max_rows = 100 if isinstance(max_rows_raw, int): max_rows = max_rows_raw result = await client.execute_query( sql=sql_raw, params=query_params, max_rows=max_rows, ) return result.to_dict() async def _handle_explain_query( client: PostgresClient, params: dict[str, Any] ) -> dict[str, Any]: sql_raw = params.get("sql") if not isinstance(sql_raw, str): raise PostgresQueryError("sql parameter must be a string") format_raw = params.get("format") format_str = "text" if isinstance(format_raw, str) and format_raw in ("text", "json", "yaml", "xml"): format_str = format_raw result = await client.explain_query( sql=sql_raw, analyze=bool(params.get("analyze", False)), buffers=bool(params.get("buffers", False)), format=format_str, ) return result.to_dict() async def _handle_list_schemas( client: PostgresClient, params: dict[str, Any] ) -> list[dict[str, Any]]: schemas = await client.list_schemas() return [s.to_dict() for s in schemas] async def _handle_list_tables( client: PostgresClient, params: dict[str, Any] ) -> list[dict[str, Any]]: schema_raw = params.get("schema") schema = "public" if isinstance(schema_raw, str): schema = schema_raw tables = await client.list_tables( schema=schema, include_views=bool(params.get("include_views", True)), ) return [t.to_dict() for t in tables] async def _handle_describe_table( client: PostgresClient, params: dict[str, Any] ) -> list[dict[str, Any]]: table_name_raw = params.get("table_name") if not isinstance(table_name_raw, str): raise PostgresQueryError("table_name must be a string") schema_raw = params.get("schema") schema = "public" if isinstance(schema_raw, str): schema = schema_raw columns = await client.describe_table( table_name=table_name_raw, schema=schema, ) return [c.to_dict() for c in columns] async def _handle_list_indexes( client: PostgresClient, params: dict[str, Any] ) -> list[dict[str, Any]]: table_name_raw = params.get("table_name") if not isinstance(table_name_raw, str): raise PostgresQueryError("table_name must be a string") schema_raw = params.get("schema") schema = "public" if isinstance(schema_raw, str): schema = schema_raw indexes = await client.list_indexes( table_name=table_name_raw, schema=schema, ) return [i.to_dict() for i in indexes] async def _handle_get_table_stats( client: PostgresClient, params: dict[str, Any] ) -> dict[str, Any]: table_name_raw = params.get("table_name") if not isinstance(table_name_raw, str): raise PostgresQueryError("table_name must be a string") schema_raw = params.get("schema") schema = "public" if isinstance(schema_raw, str): schema = schema_raw return await client.get_table_stats( table_name=table_name_raw, schema=schema, ) async def _handle_get_database_info( client: PostgresClient, params: dict[str, Any] ) -> dict[str, Any]: info = await client.get_database_info() return info.to_dict() # Таблица dispatch _TOOL_HANDLERS: dict[str, Any] = { "postgres_execute_query": _handle_execute_query, "postgres_execute_readonly": _handle_execute_readonly, "postgres_explain_query": _handle_explain_query, "postgres_list_schemas": _handle_list_schemas, "postgres_list_tables": _handle_list_tables, "postgres_describe_table": _handle_describe_table, "postgres_list_indexes": _handle_list_indexes, "postgres_get_table_stats": _handle_get_table_stats, "postgres_get_database_info": _handle_get_database_info, } # ============================================================================ # PostgreSQL MCP Server # ============================================================================ class PostgresMCPServer: """ MCP Server для PostgreSQL. Следует принципу "You Might Not Need an Effect": - Состояние клиента управляется через context manager - Dispatch по таблице вместо runtime reflection - Все преобразования — чистые функции - Явная обработка ошибок с типизированными исключениями - Prepared statements для безопасности """ def __init__(self, config: PostgresConfig | None = None): self.config = config or PostgresConfig.from_env() self._client: PostgresClient | None = None async def __aenter__(self) -> PostgresMCPServer: self._client = PostgresClient(self.config) await self._client.connect() return self async def __aexit__(self, exc_type, exc_val, exc_tb) -> None: if self._client is not None: await self._client.close() self._client = None def _require_client(self) -> PostgresClient: """Получить активный клиент или выбросить ошибку.""" if self._client is None: raise PostgresConnectionError( "Client is not connected. Use 'async with' to manage connection." ) return self._client def get_tools(self) -> list[dict[str, Any]]: """Получить список MCP tools. Чистая функция.""" return list(POSTGRES_TOOLS) async def call_tool( self, tool_name: str, parameters: dict[str, Any], context: MCPContext | None = None, ) -> dict[str, Any]: """ Вызвать tool по имени. Returns: {"success": bool, "result"|"error": ..., "error_type"?} """ try: client = self._require_client() except PostgresError as e: return { "success": False, "error": str(e), "error_type": "not_connected", } handler = _TOOL_HANDLERS.get(tool_name) if handler is None: return { "success": False, "error": f"Unknown tool: {tool_name}", "error_type": "unknown_tool", } try: result = await handler(client, parameters) if context is not None: context.add_tool_call() return { "success": True, "result": result, } except PostgresAuthError as e: return { "success": False, "error": f"Authentication failed: {e}", "error_type": "auth_error", } except PostgresConnectionError as e: return { "success": False, "error": f"Connection error: {e}", "error_type": "connection_error", } except PostgresReadOnlyError as e: return { "success": False, "error": str(e), "error_type": "read_only", } except PostgresTimeoutError as e: return { "success": False, "error": str(e), "error_type": "timeout", } except PostgresSchemaBlockedError as e: return { "success": False, "error": str(e), "error_type": "schema_blocked", } except PostgresRowLimitError as e: return { "success": False, "error": str(e), "error_type": "row_limit", } except PostgresQueryError as e: return { "success": False, "error": str(e), "error_type": "query_error", "sqlstate": e.sqlstate, "detail": e.detail, "hint": e.hint, } except PostgresError as e: return { "success": False, "error": str(e), "error_type": "postgres_error", } except Exception as e: logger.exception(f"Unexpected error in tool {tool_name}") return { "success": False, "error": f"Unexpected error: {e}", "error_type": "unexpected", } # ============================================================================ # Helper Functions (public API) # ============================================================================ @asynccontextmanager async def create_postgres_mcp( connection_string: str | None = None, host: str = "localhost", port: int = 5432, user: str | None = None, password: str | None = None, database: str | None = None, read_only: bool = False, max_rows: int = 1000, query_timeout_seconds: float = 30.0, ) -> AsyncIterator[PostgresMCPServer]: """ Создать PostgreSQL MCP Server с автоматическим управлением подключением. Usage: async with create_postgres_mcp( connection_string="postgresql://user:pass@localhost/mydb", read_only=True, ) as mcp: tools = mcp.get_tools() result = await mcp.call_tool( "postgres_execute_readonly", {"sql": "SELECT * FROM users WHERE id = $1", "params": [42]}, ) """ resolved_connection_string = ( connection_string or os.getenv("POSTGRES_URL") or os.getenv("POSTGRES_DSN") ) config = PostgresConfig( connection_string=resolved_connection_string, host=host, port=port, user=user or os.getenv("POSTGRES_USER"), password=password or os.getenv("POSTGRES_PASSWORD"), database=database or os.getenv("POSTGRES_DATABASE") or os.getenv("POSTGRES_DB"), read_only=read_only, max_rows=max_rows, query_timeout_seconds=query_timeout_seconds, ) async with PostgresMCPServer(config) as server: yield server # ============================================================================ # Exports # ============================================================================ __all__ = [ # Config "PostgresConfig", # Exceptions "PostgresError", "PostgresConnectionError", "PostgresAuthError", "PostgresQueryError", "PostgresReadOnlyError", "PostgresTimeoutError", "PostgresSchemaBlockedError", "PostgresRowLimitError", # Models "PostgresColumn", "PostgresIndex", "PostgresTable", "PostgresSchema", "PostgresDatabaseInfo", "PostgresQueryResult", "PostgresExplainResult", # Client & Server "PostgresClient", "PostgresMCPServer", "POSTGRES_TOOLS", # Helpers "create_postgres_mcp", ]