#!/usr/bin/env python # -*- coding: utf-8 -*- """PostgreSQL 异步处理器""" import logging import time from typing import Any, Dict, List, Optional from core.database_manager.handlers.common import ( format_connection_error, format_size, log_database_connect_failure, serialize_row, ) from core.database_manager.handlers.pools import get_asyncpg_pool from core.database_manager.handlers.transaction_mixin import HandlerTransactionMixin from core.database_manager.sql_utils import quote_identifier, split_sql_statements logger = logging.getLogger(__name__) class PostgreSQLHandler(HandlerTransactionMixin): """PostgreSQL异步处理器""" def __init__(self, host: str, port: int, user: str, password: str, database: str): self.host = host self.port = port self.user = user self.password = password self.database = database self.conn = None self._pool = None self._last_connect_error: Optional[str] = None self._in_transaction = False async def _release_connection(self) -> None: if self.conn and self._pool: try: await self._pool.release(self.conn) except Exception as exc: logger.warning("Error releasing PostgreSQL connection: %s", exc) self.conn = None self._pool = None async def connect(self, database: str = None) -> bool: if database: self.database = database db = (self.database or "").strip() or "postgres" try: await self._release_connection() pool = await get_asyncpg_pool( self.host, self.port, self.user, self.password, db, ) self._pool = pool self.conn = await pool.acquire() self._last_connect_error = None return True except Exception as e: self._last_connect_error = format_connection_error(e) log_database_connect_failure( db_type="postgresql", host=self.host, port=self.port, user=self.user, database=db, detail=self._last_connect_error, ) self.conn = None self._pool = None return False async def close(self): await self._release_connection() async def get_databases(self) -> List[Dict[str, Any]]: try: if not await self.connect(): return [] query = """ SELECT d.datname as name, pg_catalog.pg_get_userbyid(d.datdba) as owner, pg_catalog.pg_encoding_to_char(d.encoding) as encoding, d.datcollate as collation, pg_catalog.pg_size_pretty(pg_catalog.pg_database_size(d.datname)) as size, pg_catalog.pg_database_size(d.datname) as size_bytes, pg_catalog.shobj_description(d.oid, 'pg_database') as description FROM pg_catalog.pg_database d WHERE d.datistemplate = false ORDER BY d.datname """ rows = await self.conn.fetch(query) return [dict(row) for row in rows] except Exception as e: logger.error(f"Failed to get databases: {e}") return [] finally: await self.close() async def create_database(self, name: str, owner: str = None, encoding: str = "UTF8", template: str = "template0", **kwargs) -> bool: try: if not await self.connect("postgres"): return False qname = quote_identifier(name, "postgresql") query = f"CREATE DATABASE {qname} ENCODING '{encoding}' TEMPLATE {template}" if owner: query += f" OWNER {quote_identifier(owner, 'postgresql')}" await self.conn.execute(query) return True except Exception as e: logger.error(f"Failed to create database {name}: {e}") raise finally: await self.close() async def drop_database(self, name: str) -> bool: try: if not await self.connect("postgres"): return False qname = quote_identifier(name, "postgresql") await self.conn.execute( f"SELECT pg_terminate_backend(pid) FROM pg_stat_activity " f"WHERE datname = '{name.replace(chr(39), chr(39) + chr(39))}' " f"AND pid <> pg_backend_pid()" ) await self.conn.execute(f"DROP DATABASE {qname}") return True except Exception as e: logger.error(f"Failed to drop database {name}: {e}") raise finally: await self.close() async def rename_database(self, name: str, new_name: str) -> bool: try: if not await self.connect("postgres"): return False old_q = quote_identifier(name, "postgresql") new_q = quote_identifier(new_name, "postgresql") await self.conn.execute(f"ALTER DATABASE {old_q} RENAME TO {new_q}") return True except Exception as e: logger.error(f"Failed to rename database {name} to {new_name}: {e}") raise finally: await self.close() async def create_schema(self, name: str, database: str = None, owner: str = None) -> bool: try: if not await self.connect(database): return False qname = quote_identifier(name, "postgresql") query = f"CREATE SCHEMA {qname}" if owner: query += f" AUTHORIZATION {quote_identifier(owner, 'postgresql')}" await self.conn.execute(query) return True except Exception as e: logger.error(f"Failed to create schema {name}: {e}") raise finally: await self.close() async def drop_schema(self, name: str, database: str = None, cascade: bool = True) -> bool: try: if not await self.connect(database): return False qname = quote_identifier(name, "postgresql") cascade_clause = " CASCADE" if cascade else "" await self.conn.execute(f"DROP SCHEMA {qname}{cascade_clause}") return True except Exception as e: logger.error(f"Failed to drop schema {name}: {e}") raise finally: await self.close() async def rename_schema(self, name: str, new_name: str, database: str = None) -> bool: try: if not await self.connect(database): return False old_q = quote_identifier(name, "postgresql") new_q = quote_identifier(new_name, "postgresql") await self.conn.execute(f"ALTER SCHEMA {old_q} RENAME TO {new_q}") return True except Exception as e: logger.error(f"Failed to rename schema {name} to {new_name}: {e}") raise finally: await self.close() async def get_schemas(self, database: str = None) -> List[Dict[str, Any]]: try: if not await self.connect(database): return [] query = """ SELECT schema_name as name, schema_owner as owner, (SELECT count(*) FROM information_schema.tables WHERE table_schema = schema_name) as tables_count FROM information_schema.schemata WHERE schema_name NOT IN ('pg_catalog', 'information_schema', 'pg_toast') ORDER BY schema_name """ rows = await self.conn.fetch(query) return [dict(row) for row in rows] except Exception as e: logger.error(f"Failed to get schemas: {e}") return [] finally: await self.close() async def get_tables(self, database: str = None, schema_name: str = "public") -> List[Dict[str, Any]]: try: if not await self.connect(database): return [] query = """ SELECT schemaname as schema_name, tablename as table_name, 'BASE TABLE' as table_type, pg_size_pretty(pg_total_relation_size(schemaname || '.' || tablename)) as total_size, pg_total_relation_size(schemaname || '.' || tablename) as total_size_bytes, pg_size_pretty(pg_relation_size(schemaname || '.' || tablename)) as table_size, pg_relation_size(schemaname || '.' || tablename) as table_size_bytes, pg_size_pretty(pg_total_relation_size(schemaname || '.' || tablename) - pg_relation_size(schemaname || '.' || tablename)) as indexes_size, (pg_total_relation_size(schemaname || '.' || tablename) - pg_relation_size(schemaname || '.' || tablename)) as indexes_size_bytes, obj_description((schemaname || '.' || tablename)::regclass) as description FROM pg_catalog.pg_tables WHERE schemaname = $1 ORDER BY tablename """ rows = await self.conn.fetch(query, schema_name) tables = [dict(row) for row in rows] for table in tables: try: result = await self.conn.fetchval(f'SELECT count(*) FROM "{schema_name}"."{table["table_name"]}"') table['row_count'] = result except: table['row_count'] = None return tables except Exception as e: logger.error(f"Failed to get tables: {e}") return [] finally: await self.close() async def get_table_columns(self, table_name: str, schema_name: str = "public", database: str = None) -> List[Dict[str, Any]]: try: if not await self.connect(database): return [] query = """ SELECT c.column_name, c.data_type, c.is_nullable = 'YES' as is_nullable, c.column_default, c.character_maximum_length, c.numeric_precision, c.numeric_scale, c.ordinal_position, EXISTS(SELECT 1 FROM information_schema.table_constraints tc JOIN information_schema.key_column_usage kcu ON tc.constraint_name = kcu.constraint_name WHERE tc.table_schema = c.table_schema AND tc.table_name = c.table_name AND kcu.column_name = c.column_name AND tc.constraint_type = 'PRIMARY KEY') as is_primary_key, EXISTS(SELECT 1 FROM information_schema.table_constraints tc JOIN information_schema.key_column_usage kcu ON tc.constraint_name = kcu.constraint_name WHERE tc.table_schema = c.table_schema AND tc.table_name = c.table_name AND kcu.column_name = c.column_name AND tc.constraint_type = 'UNIQUE') as is_unique, col_description((c.table_schema || '.' || c.table_name)::regclass, c.ordinal_position) as description FROM information_schema.columns c WHERE c.table_schema = $1 AND c.table_name = $2 ORDER BY c.ordinal_position """ rows = await self.conn.fetch(query, schema_name, table_name) return [dict(row) for row in rows] except Exception as e: logger.error(f"Failed to get table columns: {e}") return [] finally: await self.close() async def get_table_indexes(self, table_name: str, schema_name: str = "public", database: str = None) -> List[Dict[str, Any]]: try: if not await self.connect(database): return [] query = """ SELECT i.indexname as index_name, am.amname as index_type, array_to_string(array_agg(a.attname ORDER BY k.ordinality), ', ') as columns, i.indexdef as definition, ix.indisunique as is_unique, ix.indisprimary as is_primary FROM pg_indexes i JOIN pg_class c ON c.relname = i.tablename AND c.relnamespace = (SELECT oid FROM pg_namespace WHERE nspname = i.schemaname) JOIN pg_index ix ON ix.indexrelid = (SELECT oid FROM pg_class WHERE relname = i.indexname AND relnamespace = (SELECT oid FROM pg_namespace WHERE nspname = i.schemaname)) JOIN pg_am am ON am.oid = (SELECT relam FROM pg_class WHERE relname = i.indexname AND relnamespace = (SELECT oid FROM pg_namespace WHERE nspname = i.schemaname)) CROSS JOIN LATERAL unnest(ix.indkey) WITH ORDINALITY AS k(attnum, ordinality) JOIN pg_attribute a ON a.attrelid = c.oid AND a.attnum = k.attnum WHERE i.schemaname = $1 AND i.tablename = $2 GROUP BY i.indexname, am.amname, i.indexdef, ix.indisunique, ix.indisprimary ORDER BY i.indexname """ rows = await self.conn.fetch(query, schema_name, table_name) return [dict(row) for row in rows] except Exception as e: logger.error(f"Failed to get table indexes: {e}") return [] finally: await self.close() async def get_table_constraints(self, table_name: str, schema_name: str = "public", database: str = None) -> List[Dict[str, Any]]: try: if not await self.connect(database): return [] query = """ SELECT c.conname AS constraint_name, CASE c.contype WHEN 'p' THEN 'PRIMARY KEY' WHEN 'f' THEN 'FOREIGN KEY' WHEN 'u' THEN 'UNIQUE' WHEN 'c' THEN 'CHECK' ELSE c.contype::text END AS constraint_type, ( SELECT string_agg(a.attname, ', ' ORDER BY u.ord) FROM unnest(c.conkey) WITH ORDINALITY AS u(attnum, ord) JOIN pg_attribute a ON a.attrelid = c.conrelid AND a.attnum = u.attnum AND a.attnum > 0 ) AS columns, pg_get_constraintdef(c.oid) AS definition, ref_cls.relname AS referenced_table, ( SELECT string_agg(a.attname, ', ' ORDER BY u.ord) FROM unnest(c.confkey) WITH ORDINALITY AS u(attnum, ord) JOIN pg_attribute a ON a.attrelid = c.confrelid AND a.attnum = u.attnum AND a.attnum > 0 ) AS referenced_columns FROM pg_constraint c JOIN pg_class cls ON cls.oid = c.conrelid JOIN pg_namespace n ON n.oid = cls.relnamespace LEFT JOIN pg_class ref_cls ON ref_cls.oid = c.confrelid WHERE n.nspname = $1 AND cls.relname = $2 AND c.contype IN ('p', 'f', 'u', 'c') ORDER BY c.contype, c.conname """ rows = await self.conn.fetch(query, schema_name, table_name) return [dict(row) for row in rows] except Exception as e: logger.error(f"Failed to get table constraints: {e}") return [] finally: await self.close() async def get_table_structure(self, table_name: str, database: str = None, schema_name: str = "public") -> Dict[str, Any]: tables = await self.get_tables(database=database, schema_name=schema_name) table_info = next((t for t in tables if t['table_name'] == table_name), None) if not table_info: raise ValueError(f"Table {schema_name}.{table_name} not found") columns = await self.get_table_columns(table_name, schema_name, database) indexes = await self.get_table_indexes(table_name, schema_name, database) constraints = await self.get_table_constraints(table_name, schema_name, database) return {"table_info": table_info, "columns": columns, "indexes": indexes, "constraints": constraints} async def get_table_ddl(self, table_name: str, schema_name: str = "public", database: str = None) -> str: try: columns = await self.get_table_columns(table_name, schema_name, database) indexes = await self.get_table_indexes(table_name, schema_name, database) constraints = await self.get_table_constraints(table_name, schema_name, database) ddl_lines = [f'CREATE TABLE "{schema_name}"."{table_name}" ('] col_defs = [] for col in columns: col_def = f' "{col["column_name"]}" {col["data_type"]}' if col.get("character_maximum_length"): col_def = f' "{col["column_name"]}" {col["data_type"]}({col["character_maximum_length"]})' if not col["is_nullable"]: col_def += " NOT NULL" if col.get("column_default"): col_def += f' DEFAULT {col["column_default"]}' col_defs.append(col_def) for pk in [c for c in constraints if c["constraint_type"] == "PRIMARY KEY"]: col_defs.append(f' PRIMARY KEY ({pk["columns"]})') ddl_lines.append(',\n'.join(col_defs)) ddl_lines.append(');') for idx in indexes: if not idx["is_primary"]: unique = "UNIQUE " if idx["is_unique"] else "" ddl_lines.append(f'\nCREATE {unique}INDEX "{idx["index_name"]}" ON "{schema_name}"."{table_name}" ({idx["columns"]});') return '\n'.join(ddl_lines) except Exception as e: return f"-- 获取DDL失败: {str(e)}" async def get_views(self, database: str = None, schema_name: str = "public") -> List[Dict[str, Any]]: try: if not await self.connect(database): return [] query = """ SELECT table_name as view_name, table_schema as schema_name, view_definition, is_updatable = 'YES' as is_updatable, check_option, 'VIEW' as view_type FROM information_schema.views WHERE table_schema = $1 UNION ALL SELECT matviewname as view_name, schemaname as schema_name, definition as view_definition, false as is_updatable, NULL as check_option, 'MATERIALIZED VIEW' as view_type FROM pg_matviews WHERE schemaname = $1 ORDER BY view_name """ rows = await self.conn.fetch(query, schema_name) return [dict(row) for row in rows] except Exception as e: logger.error(f"Failed to get views: {e}") return [] finally: await self.close() async def get_view_structure(self, view_name: str, schema_name: str = "public") -> Dict[str, Any]: try: if not await self.connect(): raise ValueError("Failed to connect") view_query = "SELECT table_name as view_name, table_schema as schema_name, view_definition, is_updatable = 'YES' as is_updatable, check_option, 'VIEW' as view_type FROM information_schema.views WHERE table_schema = $1 AND table_name = $2" view_info = await self.conn.fetchrow(view_query, schema_name, view_name) view_type = "VIEW" if not view_info: matview_query = "SELECT matviewname as view_name, schemaname as schema_name, definition as view_definition, false as is_updatable, NULL as check_option, 'MATERIALIZED VIEW' as view_type FROM pg_matviews WHERE schemaname = $1 AND matviewname = $2" view_info = await self.conn.fetchrow(matview_query, schema_name, view_name) view_type = "MATERIALIZED VIEW" if not view_info: raise ValueError(f"View {schema_name}.{view_name} not found") view_info = dict(view_info) columns_query = "SELECT column_name, data_type, is_nullable = 'YES' as is_nullable, ordinal_position, col_description((table_schema || '.' || table_name)::regclass::oid, ordinal_position) as description FROM information_schema.columns WHERE table_schema = $1 AND table_name = $2 ORDER BY ordinal_position" columns = [dict(c) for c in await self.conn.fetch(columns_query, schema_name, view_name)] deps_query = "SELECT DISTINCT ref_nsp.nspname || '.' || ref_cl.relname as table_name FROM pg_depend d JOIN pg_rewrite r ON r.oid = d.objid JOIN pg_class v ON v.oid = r.ev_class JOIN pg_namespace v_nsp ON v_nsp.oid = v.relnamespace JOIN pg_class ref_cl ON ref_cl.oid = d.refobjid JOIN pg_namespace ref_nsp ON ref_nsp.oid = ref_cl.relnamespace WHERE v.relkind IN ('v', 'm') AND v_nsp.nspname = $1 AND v.relname = $2 AND ref_cl.relkind IN ('r', 'v', 'm') AND d.deptype = 'n' ORDER BY table_name" dependencies = [row['table_name'] for row in await self.conn.fetch(deps_query, schema_name, view_name)] if view_type == "MATERIALIZED VIEW": def_result = await self.conn.fetchval("SELECT definition FROM pg_matviews WHERE schemaname = $1 AND matviewname = $2", schema_name, view_name) definition_sql = f'CREATE MATERIALIZED VIEW "{schema_name}"."{view_name}" AS\n{def_result}' if def_result else f"-- 无法获取视图定义" else: def_result = await self.conn.fetchval("SELECT pg_get_viewdef($1::regclass, true)", f"{schema_name}.{view_name}") definition_sql = f'CREATE OR REPLACE VIEW "{schema_name}"."{view_name}" AS\n{def_result}' if def_result else f"-- 无法获取视图定义" return {"view_info": view_info, "columns": columns, "dependencies": dependencies, "definition_sql": definition_sql} except Exception as e: logger.error(f"Failed to get view structure: {e}") raise finally: await self.close() async def get_view_definition(self, view_name: str, schema_name: str = "public") -> str: try: if not await self.connect(): return "-- 无法连接数据库" is_view = await self.conn.fetchval("SELECT 1 FROM information_schema.views WHERE table_schema = $1 AND table_name = $2", schema_name, view_name) if is_view: result = await self.conn.fetchval("SELECT pg_get_viewdef($1::regclass, true)", f"{schema_name}.{view_name}") definition_sql = f'CREATE OR REPLACE VIEW "{schema_name}"."{view_name}" AS\n{result}' if result else "-- 无法获取视图定义" else: result = await self.conn.fetchval("SELECT definition FROM pg_matviews WHERE schemaname = $1 AND matviewname = $2", schema_name, view_name) definition_sql = f'CREATE MATERIALIZED VIEW "{schema_name}"."{view_name}" AS\n{result}' if result else "-- 无法获取视图定义" return definition_sql except Exception as e: return f"-- 获取视图定义失败: {str(e)}" finally: await self.close() async def get_view_dependencies(self, view_name: str, schema_name: str = "public") -> List[str]: try: if not await self.connect(): return [] query = "SELECT DISTINCT ref_nsp.nspname || '.' || ref_cl.relname as table_name FROM pg_depend d JOIN pg_rewrite r ON r.oid = d.objid JOIN pg_class v ON v.oid = r.ev_class JOIN pg_namespace v_nsp ON v_nsp.oid = v.relnamespace JOIN pg_class ref_cl ON ref_cl.oid = d.refobjid JOIN pg_namespace ref_nsp ON ref_nsp.oid = ref_cl.relnamespace WHERE v.relkind IN ('v', 'm') AND v_nsp.nspname = $1 AND v.relname = $2 AND ref_cl.relkind IN ('r', 'v', 'm') AND d.deptype = 'n' ORDER BY table_name" rows = await self.conn.fetch(query, schema_name, view_name) return [row['table_name'] for row in rows] except Exception as e: return [] finally: await self.close() async def query_data(self, table_name: str, schema_name: str = None, page: int = 1, page_size: int = 20, where: str = None, order_by: str = None, database: str = None) -> Dict[str, Any]: try: if not await self.connect(database): return {"columns": [], "rows": [], "total": 0, "page": page, "page_size": page_size} schema_name = schema_name or "public" full_table = f'"{ schema_name}"."{table_name}"' count_q = f"SELECT COUNT(*) FROM {full_table}" + (f" WHERE {where}" if where else "") total = await self.conn.fetchval(count_q) offset = (page - 1) * page_size data_q = f"SELECT * FROM {full_table}" + (f" WHERE {where}" if where else "") + (f" ORDER BY {order_by}" if order_by else "") + f" LIMIT {page_size} OFFSET {offset}" rows = await self.conn.fetch(data_q) rows_list = [serialize_row(dict(row)) for row in rows] columns = list(rows_list[0].keys()) if rows_list else [] return {"columns": columns, "rows": rows_list, "total": total, "page": page, "page_size": page_size} except Exception as e: logger.error(f"Failed to query data: {e}") return {"columns": [], "rows": [], "total": 0, "page": page, "page_size": page_size} finally: await self.close() async def _run_sql_query(self, sql: str) -> list: return await self.conn.fetch(sql) async def _run_sql_command(self, sql: str) -> int: result = await self.conn.execute(sql) if result and str(result).split()[-1].isdigit(): return int(str(result).split()[-1]) return 0 async def insert_data(self, table_name: str, data: Dict[str, Any], schema_name: str = None) -> Dict[str, Any]: try: if not await self.connect(): return {"success": False, "message": "数据库连接失败", "affected_rows": 0} schema_name = schema_name or "public" columns = list(data.keys()) values = list(data.values()) placeholders = ', '.join([f'${i+1}' for i in range(len(values))]) query = f'INSERT INTO "{schema_name}"."{table_name}" ({", ".join(columns)}) VALUES ({placeholders})' await self.conn.execute(query, *values) return {"success": True, "message": "插入成功", "affected_rows": 1} except Exception as e: return {"success": False, "message": str(e), "affected_rows": 0} finally: await self.close() async def update_data(self, table_name: str, data: Dict[str, Any], where: str, schema_name: str = None) -> Dict[str, Any]: try: if not await self.connect(): return {"success": False, "message": "数据库连接失败", "affected_rows": 0} schema_name = schema_name or "public" set_clause = ', '.join([f'{k} = ${i+1}' for i, k in enumerate(data.keys())]) query = f'UPDATE "{schema_name}"."{table_name}" SET {set_clause} WHERE {where}' result = await self.conn.execute(query, *data.values()) affected_rows = int(result.split()[-1]) if result and result.split()[-1].isdigit() else 0 return {"success": True, "message": f"更新成功,影响 {affected_rows} 行", "affected_rows": affected_rows} except Exception as e: return {"success": False, "message": str(e), "affected_rows": 0} finally: await self.close() async def delete_data(self, table_name: str, where: str, schema_name: str = None) -> Dict[str, Any]: try: if not await self.connect(): return {"success": False, "message": "数据库连接失败", "affected_rows": 0} schema_name = schema_name or "public" query = f'DELETE FROM "{schema_name}"."{table_name}" WHERE {where}' result = await self.conn.execute(query) affected_rows = int(result.split()[-1]) if result and result.split()[-1].isdigit() else 0 return {"success": True, "message": f"删除成功,影响 {affected_rows} 行", "affected_rows": affected_rows} except Exception as e: return {"success": False, "message": str(e), "affected_rows": 0} finally: await self.close() async def execute_ddl(self, sql: str, database: str = None, schema_name: str = None) -> Dict[str, Any]: try: if not await self.connect(database): return {"success": False, "message": "数据库连接失败", "affected_rows": 0} if schema_name: await self.conn.execute( f"SET search_path TO {quote_identifier(schema_name, 'postgresql')}, public" ) statements = split_sql_statements(sql) if not statements: return {"success": False, "message": "无有效SQL语句", "affected_rows": 0} async with self.conn.transaction(): for index, statement in enumerate(statements, start=1): try: await self.conn.execute(statement) except Exception as exc: return { "success": False, "message": f"DDL执行失败 (第{index}条): {exc}", "affected_rows": 0, } return {"success": True, "message": "DDL执行成功", "affected_rows": 0} except Exception as e: return {"success": False, "message": f"DDL执行失败: {str(e)}", "affected_rows": 0} finally: await self.close()