726 lines
30 KiB
Python
726 lines
30 KiB
Python
#!/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()
|