feat:新增投顾agent和nl2sqlagent
This commit is contained in:
@@ -0,0 +1,92 @@
|
||||
"""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()
|
||||
Reference in New Issue
Block a user