#!/usr/bin/env python # -*- coding: utf-8 -*- """SQL Server 异步处理器(aioodbc)""" 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_aioodbc_pool from core.database_manager.handlers.transaction_mixin import HandlerTransactionMixin from core.database_manager.sql_utils import quote_identifier, quote_table, split_sql_statements logger = logging.getLogger(__name__) def _normalize_mssql_default(value: Optional[str]) -> Optional[str]: """Strip SQL Server INFORMATION_SCHEMA default wrappers, e.g. ((0)) -> 0.""" if not value: return None normalized = value.strip() while normalized.startswith("(") and normalized.endswith(")"): normalized = normalized[1:-1].strip() if not normalized or normalized.upper() == "NULL": return None return normalized _MSSQL_SYSTEM_SCHEMAS = ( "sys", "INFORMATION_SCHEMA", "guest", "db_owner", "db_accessadmin", "db_securityadmin", "db_ddladmin", "db_backupoperator", "db_datareader", "db_datawriter", "db_denydatareader", "db_denydatawriter", ) class SQLServerHandler(HandlerTransactionMixin): """SQL Server 异步处理器""" def __init__( self, host: str, port: int, user: str, password: str, database: str, extra_options: Optional[Dict[str, Any]] = None, ): self.host = host self.port = port self.user = user self.password = password self.database = database or "master" self.extra_options = extra_options or {} 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 SQL Server connection: %s", exc) self.conn = None self._pool = None async def connect(self, database: str = None) -> bool: db = database or self.database or "master" try: await self._release_connection() pool = await get_aioodbc_pool( self.host, self.port, self.user, self.password, db, self.extra_options, ) 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="sqlserver", 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 _execute_query(self, query: str, params: tuple = None) -> List[Dict[str, Any]]: async with self.conn.cursor() as cursor: await cursor.execute(query, params or ()) if not cursor.description: return [] columns = [col[0] for col in cursor.description] rows = await cursor.fetchall() return [dict(zip(columns, row)) for row in rows] async def _execute_command(self, command: str, params: tuple = None) -> int: async with self.conn.cursor() as cursor: await cursor.execute(command, params or ()) return cursor.rowcount async def _run_sql_query(self, sql: str) -> list: return await self._execute_query(sql) async def _run_sql_command(self, sql: str) -> int: return await self._execute_command(sql) async def _exec_tx_statement(self, sql: str) -> None: upper = sql.strip().upper() if upper == "BEGIN": await self._run_sql_command("BEGIN TRANSACTION") elif upper == "COMMIT": await self._run_sql_command("COMMIT TRANSACTION") elif upper == "ROLLBACK": await self._run_sql_command("ROLLBACK TRANSACTION") else: await self._run_sql_command(sql) async def get_databases(self) -> List[Dict[str, Any]]: try: if not await self.connect("master"): return [] query = """ SELECT d.name, SUSER_SNAME(d.owner_sid) AS owner, d.collation_name AS collation, CAST( (SELECT SUM(CAST(mf.size AS BIGINT)) * 8192 FROM sys.master_files mf WHERE mf.database_id = d.database_id) AS BIGINT ) AS size_bytes FROM sys.databases d WHERE d.state = 0 ORDER BY d.name """ databases = await self._execute_query(query) for db in databases: size_bytes = int(db.get("size_bytes") or 0) db["size"] = format_size(size_bytes) db["size_bytes"] = size_bytes return databases except Exception as e: logger.error("Failed to get databases: %s", e) return [] finally: await self.close() async def create_database(self, name: str, **kwargs) -> bool: try: if not await self.connect("master"): return False qname = quote_identifier(name, "sqlserver") await self._execute_command(f"CREATE DATABASE {qname}") return True except Exception as e: logger.error("Failed to create database %s: %s", name, e) raise finally: await self.close() async def drop_database(self, name: str) -> bool: try: if not await self.connect("master"): return False qname = quote_identifier(name, "sqlserver") await self._execute_command( f"ALTER DATABASE {qname} SET SINGLE_USER WITH ROLLBACK IMMEDIATE" ) await self._execute_command(f"DROP DATABASE {qname}") return True except Exception as e: logger.error("Failed to drop database %s: %s", name, e) raise finally: await self.close() async def rename_database(self, name: str, new_name: str) -> bool: try: if not await self.connect("master"): return False old_q = quote_identifier(name, "sqlserver") new_q = quote_identifier(new_name, "sqlserver") await self._execute_command(f"ALTER DATABASE {old_q} MODIFY NAME = {new_q}") return True except Exception as e: logger.error("Failed to rename database %s to %s: %s", name, 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, "sqlserver") query = f"CREATE SCHEMA {qname}" if owner: query += f" AUTHORIZATION {quote_identifier(owner, 'sqlserver')}" await self._execute_command(query) return True except Exception as e: logger.error("Failed to create schema %s: %s", 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, "sqlserver") await self._execute_command(f"DROP SCHEMA {qname}") return True except Exception as e: logger.error("Failed to drop schema %s: %s", name, e) raise finally: await self.close() async def rename_schema(self, name: str, new_name: str, database: str = None) -> bool: raise NotImplementedError("SQL Server 不支持直接重命名 Schema") async def get_schemas(self, database: str = None) -> List[Dict[str, Any]]: try: if not await self.connect(database): return [] placeholders = ", ".join(["?"] * len(_MSSQL_SYSTEM_SCHEMAS)) query = f""" SELECT s.name, USER_NAME(s.principal_id) AS owner, (SELECT COUNT(*) FROM sys.tables t WHERE t.schema_id = s.schema_id) AS tables_count FROM sys.schemas s WHERE s.name NOT IN ({placeholders}) AND s.name NOT LIKE 'db[_]%' ORDER BY s.name """ return await self._execute_query(query, _MSSQL_SYSTEM_SCHEMAS) except Exception as e: logger.error("Failed to get schemas: %s", e) return [] finally: await self.close() async def get_tables( self, database: str = None, schema_name: str = "dbo" ) -> List[Dict[str, Any]]: try: if not await self.connect(database): return [] query = """ SELECT t.TABLE_SCHEMA AS schema_name, t.TABLE_NAME AS table_name, t.TABLE_TYPE AS table_type, CAST(p.rows AS BIGINT) AS row_count FROM INFORMATION_SCHEMA.TABLES t LEFT JOIN sys.tables st ON st.name = t.TABLE_NAME AND SCHEMA_NAME(st.schema_id) = t.TABLE_SCHEMA LEFT JOIN sys.partitions p ON st.object_id = p.object_id AND p.index_id IN (0, 1) WHERE t.TABLE_SCHEMA = ? AND t.TABLE_TYPE = 'BASE TABLE' ORDER BY t.TABLE_NAME """ return await self._execute_query(query, (schema_name,)) except Exception as e: logger.error("Failed to get tables: %s", e) return [] finally: await self.close() async def get_table_columns( self, table_name: str, schema_name: str = "dbo", database: str = None ) -> List[Dict[str, Any]]: try: if not await self.connect(database): return [] query = """ SELECT c.COLUMN_NAME AS column_name, c.DATA_TYPE AS data_type, CASE WHEN c.IS_NULLABLE = 'YES' THEN 1 ELSE 0 END AS is_nullable, c.COLUMN_DEFAULT AS column_default, c.CHARACTER_MAXIMUM_LENGTH AS character_maximum_length, c.NUMERIC_PRECISION AS numeric_precision, c.NUMERIC_SCALE AS numeric_scale, c.ORDINAL_POSITION AS ordinal_position, CASE WHEN pk.COLUMN_NAME IS NOT NULL THEN 1 ELSE 0 END AS is_primary_key, CASE WHEN uq.COLUMN_NAME IS NOT NULL THEN 1 ELSE 0 END AS is_unique, ep.value AS description FROM INFORMATION_SCHEMA.COLUMNS c LEFT JOIN ( SELECT ku.TABLE_SCHEMA, ku.TABLE_NAME, ku.COLUMN_NAME FROM INFORMATION_SCHEMA.TABLE_CONSTRAINTS tc JOIN INFORMATION_SCHEMA.KEY_COLUMN_USAGE ku ON tc.CONSTRAINT_NAME = ku.CONSTRAINT_NAME AND tc.TABLE_SCHEMA = ku.TABLE_SCHEMA AND tc.TABLE_NAME = ku.TABLE_NAME WHERE tc.CONSTRAINT_TYPE = 'PRIMARY KEY' ) pk ON pk.TABLE_SCHEMA = c.TABLE_SCHEMA AND pk.TABLE_NAME = c.TABLE_NAME AND pk.COLUMN_NAME = c.COLUMN_NAME LEFT JOIN ( SELECT ku.TABLE_SCHEMA, ku.TABLE_NAME, ku.COLUMN_NAME FROM INFORMATION_SCHEMA.TABLE_CONSTRAINTS tc JOIN INFORMATION_SCHEMA.KEY_COLUMN_USAGE ku ON tc.CONSTRAINT_NAME = ku.CONSTRAINT_NAME AND tc.TABLE_SCHEMA = ku.TABLE_SCHEMA AND tc.TABLE_NAME = ku.TABLE_NAME WHERE tc.CONSTRAINT_TYPE = 'UNIQUE' ) uq ON uq.TABLE_SCHEMA = c.TABLE_SCHEMA AND uq.TABLE_NAME = c.TABLE_NAME AND uq.COLUMN_NAME = c.COLUMN_NAME LEFT JOIN sys.schemas ss ON ss.name = c.TABLE_SCHEMA LEFT JOIN sys.tables st ON st.schema_id = ss.schema_id AND st.name = c.TABLE_NAME LEFT JOIN sys.columns sc ON sc.object_id = st.object_id AND sc.name = c.COLUMN_NAME LEFT JOIN sys.extended_properties ep ON ep.major_id = sc.object_id AND ep.minor_id = sc.column_id AND ep.name = 'MS_Description' WHERE c.TABLE_SCHEMA = ? AND c.TABLE_NAME = ? ORDER BY c.ORDINAL_POSITION """ columns = await self._execute_query(query, (schema_name, table_name)) for col in columns: col["column_default"] = _normalize_mssql_default(col.get("column_default")) return columns except Exception as e: logger.error("Failed to get table columns: %s", e) return [] finally: await self.close() async def get_table_indexes( self, table_name: str, schema_name: str = "dbo", database: str = None ) -> List[Dict[str, Any]]: try: if not await self.connect(database): return [] query = """ SELECT i.name AS index_name, i.type_desc AS index_type, STUFF(( SELECT ', ' + c.name FROM sys.index_columns ic JOIN sys.columns c ON ic.object_id = c.object_id AND ic.column_id = c.column_id WHERE ic.object_id = i.object_id AND ic.index_id = i.index_id ORDER BY ic.key_ordinal FOR XML PATH('') ), 1, 2, '') AS columns, i.is_unique, i.is_primary_key AS is_primary, '' AS definition FROM sys.indexes i JOIN sys.tables t ON i.object_id = t.object_id JOIN sys.schemas s ON t.schema_id = s.schema_id WHERE s.name = ? AND t.name = ? AND i.name IS NOT NULL ORDER BY i.name """ indexes = await self._execute_query(query, (schema_name, table_name)) for idx in indexes: unique = "UNIQUE " if idx.get("is_unique") else "" idx["definition"] = f"CREATE {unique}INDEX {idx['index_name']} ON {schema_name}.{table_name} ({idx.get('columns', '')})" return indexes except Exception as e: logger.error("Failed to get table indexes: %s", e) return [] finally: await self.close() async def get_table_constraints( self, table_name: str, schema_name: str = "dbo", database: str = None ) -> List[Dict[str, Any]]: try: if not await self.connect(database): return [] query = """ SELECT tc.CONSTRAINT_NAME AS constraint_name, tc.CONSTRAINT_TYPE AS constraint_type, STUFF(( SELECT ', ' + kcu.COLUMN_NAME FROM INFORMATION_SCHEMA.KEY_COLUMN_USAGE kcu WHERE kcu.CONSTRAINT_NAME = tc.CONSTRAINT_NAME AND kcu.TABLE_SCHEMA = tc.TABLE_SCHEMA AND kcu.TABLE_NAME = tc.TABLE_NAME FOR XML PATH('') ), 1, 2, '') AS columns, kcu2.TABLE_NAME AS referenced_table, STUFF(( SELECT ', ' + fk.COLUMN_NAME FROM INFORMATION_SCHEMA.KEY_COLUMN_USAGE fk WHERE fk.CONSTRAINT_NAME = tc.CONSTRAINT_NAME AND fk.TABLE_SCHEMA = tc.TABLE_SCHEMA FOR XML PATH('') ), 1, 2, '') AS referenced_columns, '' AS definition FROM INFORMATION_SCHEMA.TABLE_CONSTRAINTS tc LEFT JOIN INFORMATION_SCHEMA.REFERENTIAL_CONSTRAINTS rc ON tc.CONSTRAINT_NAME = rc.CONSTRAINT_NAME AND tc.TABLE_SCHEMA = rc.CONSTRAINT_SCHEMA LEFT JOIN INFORMATION_SCHEMA.KEY_COLUMN_USAGE kcu2 ON rc.UNIQUE_CONSTRAINT_NAME = kcu2.CONSTRAINT_NAME WHERE tc.TABLE_SCHEMA = ? AND tc.TABLE_NAME = ? ORDER BY tc.CONSTRAINT_TYPE, tc.CONSTRAINT_NAME """ constraints = await self._execute_query(query, (schema_name, table_name)) for const in constraints: const["definition"] = f"{const['constraint_type']} ({const.get('columns', '')})" return constraints except Exception as e: logger.error("Failed to get table constraints: %s", e) return [] finally: await self.close() async def get_table_structure( self, table_name: str, database: str = None, schema_name: str = "dbo" ) -> 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 = "dbo", database: str = None ) -> str: try: if not await self.connect(database): return "-- 无法连接数据库" full_name = quote_table(schema_name, table_name, "sqlserver") rows = await self._execute_query( "SELECT OBJECT_DEFINITION(OBJECT_ID(?)) AS ddl", (f"{schema_name}.{table_name}",), ) if rows and rows[0].get("ddl"): return rows[0]["ddl"] columns = await self.get_table_columns(table_name, schema_name, database) if not columns: return f"-- 无法获取表 {full_name} 的DDL" ddl_lines = [f"CREATE TABLE {full_name} ("] col_defs = [] for col in columns: col_def = f" {quote_identifier(col['column_name'], 'sqlserver')} {col['data_type']}" if col.get("character_maximum_length"): col_def += f"({col['character_maximum_length']})" if not col.get("is_nullable"): col_def += " NOT NULL" col_defs.append(col_def) ddl_lines.append(",\n".join(col_defs)) ddl_lines.append(");") return "\n".join(ddl_lines) except Exception as e: return f"-- 获取DDL失败: {str(e)}" finally: await self.close() async def get_views( self, database: str = None, schema_name: str = "dbo" ) -> 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 AS view_definition, CASE WHEN IS_UPDATABLE = 'YES' THEN 1 ELSE 0 END AS is_updatable, CHECK_OPTION AS check_option, 'VIEW' AS view_type FROM INFORMATION_SCHEMA.VIEWS WHERE TABLE_SCHEMA = ? ORDER BY TABLE_NAME """ return await self._execute_query(query, (schema_name,)) except Exception as e: logger.error("Failed to get views: %s", e) return [] finally: await self.close() async def get_view_structure(self, view_name: str, schema_name: str = "dbo") -> 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 AS view_definition, CASE WHEN IS_UPDATABLE = 'YES' THEN 1 ELSE 0 END AS is_updatable, CHECK_OPTION AS check_option, 'VIEW' AS view_type FROM INFORMATION_SCHEMA.VIEWS WHERE TABLE_SCHEMA = ? AND TABLE_NAME = ? """ view_info = await self._execute_query(view_query, (schema_name, view_name)) if not view_info: raise ValueError(f"View {schema_name}.{view_name} not found") view_info = view_info[0] columns = await self.get_table_columns(view_name, schema_name) definition_sql = await self.get_view_definition(view_name, schema_name) dependencies = await self.get_view_dependencies(view_name, schema_name) return { "view_info": view_info, "columns": columns, "dependencies": dependencies, "definition_sql": definition_sql, } except Exception as e: logger.error("Failed to get view structure: %s", e) raise finally: await self.close() async def get_view_definition(self, view_name: str, schema_name: str = "dbo") -> str: try: if not await self.connect(): return "-- 无法连接数据库" rows = await self._execute_query( """ SELECT VIEW_DEFINITION AS definition FROM INFORMATION_SCHEMA.VIEWS WHERE TABLE_SCHEMA = ? AND TABLE_NAME = ? """, (schema_name, view_name), ) if rows and rows[0].get("definition"): full = quote_table(schema_name, view_name, "sqlserver") return f"CREATE VIEW {full} AS\n{rows[0]['definition']}" return f"-- 无法获取视图 {schema_name}.{view_name} 的定义" except Exception as e: return f"-- 获取视图定义失败: {str(e)}" finally: await self.close() async def get_view_dependencies(self, view_name: str, schema_name: str = "dbo") -> List[str]: try: if not await self.connect(): return [] rows = await self._execute_query( """ SELECT DISTINCT REFERENCED_ENTITY_NAME AS table_name FROM sys.sql_expression_dependencies d JOIN sys.views v ON d.referencing_id = v.object_id JOIN sys.schemas s ON v.schema_id = s.schema_id WHERE s.name = ? AND v.name = ? AND d.referenced_entity_name IS NOT NULL ORDER BY REFERENCED_ENTITY_NAME """, (schema_name, view_name), ) return [row["table_name"] for row in rows] except Exception as e: logger.error("Failed to get view dependencies: %s", 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 "dbo" full_table = quote_table(schema_name, table_name, "sqlserver") count_q = f"SELECT COUNT(*) AS total FROM {full_table}" if where: count_q += f" WHERE {where}" count_result = await self._execute_query(count_q) total = count_result[0]["total"] if count_result else 0 offset = (page - 1) * page_size data_q = f"SELECT * FROM {full_table}" if where: data_q += f" WHERE {where}" if order_by: data_q += f" ORDER BY {order_by}" data_q += f" OFFSET {offset} ROWS FETCH NEXT {page_size} ROWS ONLY" rows = await self._execute_query(data_q) rows_list = [serialize_row(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("Failed to query data: %s", e) return {"columns": [], "rows": [], "total": 0, "page": page, "page_size": page_size} finally: await self.close() 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 "dbo" full_table = quote_table(schema_name, table_name, "sqlserver") columns = list(data.keys()) quoted_columns = ", ".join( quote_identifier(col, "sqlserver") for col in columns ) placeholders = ", ".join(["?"] * len(columns)) query = f"INSERT INTO {full_table} ({quoted_columns}) VALUES ({placeholders})" affected_rows = await self._execute_command(query, tuple(data.values())) return {"success": True, "message": "插入成功", "affected_rows": affected_rows or 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 "dbo" full_table = quote_table(schema_name, table_name, "sqlserver") set_clause = ", ".join( f"{quote_identifier(k, 'sqlserver')} = ?" for k in data.keys() ) query = f"UPDATE {full_table} SET {set_clause} WHERE {where}" affected_rows = await self._execute_command(query, tuple(data.values())) 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 "dbo" full_table = quote_table(schema_name, table_name, "sqlserver") query = f"DELETE FROM {full_table} WHERE {where}" affected_rows = await self._execute_command(query) 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} statements = split_sql_statements(sql) if not statements: return {"success": False, "message": "无有效SQL语句", "affected_rows": 0} for index, statement in enumerate(statements, start=1): try: await self._execute_command(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()