Files
2026-06-08 18:14:59 +08:00

547 lines
29 KiB
Python

#!/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()