Files

93 lines
2.6 KiB
Python
Raw Permalink Normal View History

2026-09-13 16:19:24 +08:00
"""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()