feat:新增投顾agent和nl2sqlagent
This commit is contained in:
@@ -0,0 +1,41 @@
|
||||
"""NL2SQL 查询流式事件适配。"""
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import AsyncIterator, Awaitable, Callable
|
||||
from typing import Any
|
||||
|
||||
from nl2sql.contracts import DataQueryRequest, DataQueryResult
|
||||
|
||||
|
||||
async def stream_query_events(
|
||||
request: DataQueryRequest,
|
||||
*,
|
||||
query_runner: Callable[..., Awaitable[DataQueryResult]],
|
||||
**dependencies: Any,
|
||||
) -> AsyncIterator[dict[str, Any]]:
|
||||
"""复用查询服务并输出不泄露敏感内容的结构化事件。"""
|
||||
query_id = dependencies.get("query_id")
|
||||
if not query_id:
|
||||
raise ValueError("query_id must be provided")
|
||||
yield {
|
||||
"event": "started",
|
||||
"query_id": query_id,
|
||||
"trace_id": request.trace_id,
|
||||
}
|
||||
try:
|
||||
result = await query_runner(request, **dependencies)
|
||||
except Exception as exc: # noqa: BLE001 流式失败只输出异常类型
|
||||
yield {
|
||||
"event": "failed",
|
||||
"query_id": query_id,
|
||||
"trace_id": request.trace_id,
|
||||
"error_type": type(exc).__name__,
|
||||
}
|
||||
return
|
||||
yield {
|
||||
"event": "completed",
|
||||
"query_id": result.query_id,
|
||||
"trace_id": result.trace_id or request.trace_id,
|
||||
"row_count": result.row_count,
|
||||
"truncated": result.truncated,
|
||||
}
|
||||
Reference in New Issue
Block a user