42 lines
1.3 KiB
Python
42 lines
1.3 KiB
Python
"""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,
|
||
|
|
}
|