Files
Mutual_Fund/nl2sql/streaming.py

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,
}