"""NL2SQL 运行中查询登记和管理员中止状态。""" from __future__ import annotations from dataclasses import asdict, dataclass from datetime import datetime, timezone from threading import RLock from sqlalchemy import text @dataclass class QueryRuntime: """单条运行中查询的最小运行信息。""" query_id: str user_id: int sql: str connection_id: int | None started_at: datetime status: str = "running" def to_dict(self) -> dict: data = asdict(self) data["started_at"] = self.started_at.isoformat() return data class QueryRuntimeRegistry: """进程内运行查询登记表,跨进程场景由网关保证路由到同一实例。""" def __init__(self): self._items: dict[str, QueryRuntime] = {} self._lock = RLock() async def register( self, *, query_id: str, user_id: int, sql: str, connection_id: int | None = None, ) -> QueryRuntime: item = QueryRuntime( query_id=query_id, user_id=user_id, sql=sql, connection_id=connection_id, started_at=datetime.now(timezone.utc), ) with self._lock: self._items[query_id] = item return item def get(self, query_id: str) -> QueryRuntime | None: with self._lock: return self._items.get(query_id) def list_active(self) -> list[QueryRuntime]: with self._lock: return [item for item in self._items.values() if item.status == "running"] async def complete(self, query_id: str, *, status: str = "completed") -> bool: with self._lock: item = self._items.get(query_id) if item is None: return False item.status = status self._items.pop(query_id, None) return True async def mark_killed(self, query_id: str) -> bool: with self._lock: item = self._items.get(query_id) if item is None or item.status != "running": return False item.status = "killed" return True async def kill_mysql_query(connection_id: int) -> bool: """通过独立连接执行 KILL QUERY,避免占用被中止的连接。""" from config.database.mysql import get_session_factory if int(connection_id) <= 0: return False async with get_session_factory()() as session: await session.execute(text(f"KILL QUERY {int(connection_id)}")) await session.commit() return True query_runtime_registry = QueryRuntimeRegistry()