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