93 lines
2.6 KiB
Python
93 lines
2.6 KiB
Python
"""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()
|