- Introduced new endpoints `/api/analyst/query/{trace_id}/sample` and `/api/analyst/escalate` for sampling query results and escalating issues to human analysts, respectively.
- Enhanced `AnalystAgent` to support sampling of SQL results based on trace ID and to handle escalation requests, improving user experience in error scenarios.
- Updated `analyst_schemas.py` to include `EscalateRequest` for structured escalation requests.
- Added corresponding frontend API calls and UI components to facilitate user interactions with the new features.
- Implemented unit tests to ensure the reliability of the new functionalities.
This update significantly enhances the analytical capabilities of the application, allowing users to retrieve detailed query samples and escalate issues effectively.
596 lines
23 KiB
Python
596 lines
23 KiB
Python
"""数据分析 Agent 编排:NL → 消歧 → 生成 SQL → 校验 → 执行 → 解读 → 护栏 → 留痕。
|
||
|
||
与架构说明书 §6 的节点一一对应。核心逻辑用可测试的类实现,
|
||
`build_graph()` 提供 LangGraph StateGraph 适配(架构对齐)。
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
import hashlib
|
||
import time
|
||
import uuid
|
||
from typing import Any
|
||
|
||
from app.model.analyst_schemas import (
|
||
CUSTOMER_AI_RISK_NOTE,
|
||
DISCLAIMER,
|
||
AnalystResponse,
|
||
InterpretRequest,
|
||
Meta,
|
||
TableData,
|
||
)
|
||
from app.api.analyst_auth_adapter import (
|
||
AnalystAuthContext,
|
||
AnalystAuthError,
|
||
assert_analyst_query_access,
|
||
resolve_analyst_scope,
|
||
)
|
||
from app.service.analytics_repo import AnalyticsRepo, classify_empty, _is_timeout_error
|
||
from app.service.cache_service import CacheService
|
||
from app.service.dict_service import Ambiguity, MetricRegistry, default_registry
|
||
from app.service.guardrail import GuardrailResult, verify
|
||
from app.config.settings import settings
|
||
from app.service.llm import DeepSeekLLM, estimate_cost, extract_sql
|
||
from app.service.schema_meta import SCHEMA_PROMPT
|
||
from app.service.sql_guard import SqlGuardError, validate
|
||
from app.service.template_service import TemplateService
|
||
|
||
SQL_GEN_SYSTEM = (
|
||
"你是金融数据查询助手。根据给定表结构与口径,把用户问题翻译成【一条】只读 SELECT SQL。"
|
||
"严格遵守:1) 只输出 SQL,不要代码块、不要解释、不要分号结尾外的多余内容;"
|
||
"2) 只能使用给定表;3) 不生成任何写操作;4) 涉及金额时用原始字段不要自行换算单位。"
|
||
)
|
||
|
||
ANSWER_SYSTEM = (
|
||
"你是金融数据分析助手。根据 SQL 与查询结果,用简洁人话解读,数字必须与结果完全一致,不得编造。"
|
||
"结尾无需重复免责声明(由系统统一附带)。若结果为空,如实说明。"
|
||
)
|
||
|
||
CUSTOMER_ANSWER_SYSTEM = (
|
||
"你是面向个人客户的财富数据解读助手。只能根据查询结果描述趋势、分布与数量,"
|
||
"不得给出投资建议、收益承诺、买卖时点或产品推荐。若用户问题涉及建议,只描述数据事实。"
|
||
"数字必须与结果完全一致,不得编造。"
|
||
)
|
||
|
||
_ESCALATE_HINTS = [
|
||
"可点击「转交人工分析」由分析师协助核对口径与 SQL。",
|
||
"或前往「数据分析 · 对话」联系值班分析师(演示环境可转顾问工作台)。",
|
||
]
|
||
|
||
|
||
def _build_sample_sql(sql: str, limit: int) -> str:
|
||
lim = max(1, min(int(limit), 20))
|
||
inner = sql.strip().rstrip(";")
|
||
return f"SELECT * FROM ({inner}) AS _n03_sample LIMIT {lim}"
|
||
|
||
|
||
class AnalystAgent:
|
||
def __init__(
|
||
self,
|
||
llm: DeepSeekLLM | None = None,
|
||
repo: AnalyticsRepo | None = None,
|
||
cache: CacheService | None = None,
|
||
registry: MetricRegistry | None = None,
|
||
templates: TemplateService | None = None,
|
||
) -> None:
|
||
self.llm = llm or DeepSeekLLM()
|
||
self.repo = repo or AnalyticsRepo()
|
||
self.cache = cache or CacheService.auto()
|
||
self.registry = registry or default_registry()
|
||
self.templates = templates if templates is not None else TemplateService(repo=self.repo)
|
||
|
||
# ---------- 主入口 ----------
|
||
def run(
|
||
self,
|
||
question: str,
|
||
auth: AnalystAuthContext,
|
||
session_id: str | None = None,
|
||
trace_id: str | None = None,
|
||
*,
|
||
interpret: bool = False,
|
||
) -> AnalystResponse:
|
||
trace_id = trace_id or f"trace-{uuid.uuid4().hex[:16]}"
|
||
session_id = session_id or f"sess-{uuid.uuid4().hex[:12]}"
|
||
auth.trace_id = trace_id
|
||
started = time.time()
|
||
cost_est = 0.0
|
||
|
||
try:
|
||
domain = assert_analyst_query_access(auth)
|
||
except AnalystAuthError as exc:
|
||
resp = self._deny(exc.error_code, exc.message, trace_id)
|
||
return self._audit_terminal(question, auth, session_id, trace_id, resp)
|
||
|
||
scope: list[str] = resolve_analyst_scope(auth, domain, self.repo)
|
||
|
||
# 1) 指标消歧(N-01)
|
||
amb = self._detect_ambiguity(question)
|
||
if amb is not None:
|
||
resp = self._clarify(amb, trace_id)
|
||
return self._audit_terminal(question, auth, session_id, trace_id, resp)
|
||
|
||
# 2) 模板填参(D-06)或 LLM 生成 SQL
|
||
template_key: str | None = None
|
||
template_hit = False
|
||
rendered = self.templates.try_render(question, domain, scope, auth.roles)
|
||
if rendered is not None:
|
||
sql_text, template_key = rendered
|
||
template_hit = True
|
||
usage: dict = {}
|
||
else:
|
||
sql_text, usage = self._generate_sql(question, domain, scope)
|
||
cost_est += estimate_cost(usage)
|
||
if not (sql_text or "").strip():
|
||
resp = self._escalate("未能生成有效 SQL", trace_id, error_code="LLM_EMPTY_SQL")
|
||
return self._audit_terminal(question, auth, session_id, trace_id, resp)
|
||
|
||
# 3) 校验(五层)
|
||
try:
|
||
vres = validate(sql_text, domain, scope)
|
||
except SqlGuardError as exc:
|
||
resp = self._deny(exc.error_code, exc.message, trace_id, domain)
|
||
return self._audit_terminal(
|
||
question, auth, session_id, trace_id, resp, sql=sql_text
|
||
)
|
||
|
||
# 4) 执行(缓存优先)
|
||
perm_fp = self.cache.permission_fingerprint(auth.subject_id, domain, scope)
|
||
sql_hash = self.cache.sql_hash(sql_text)
|
||
cached = self.cache.get_result(perm_fp, sql_text, vres.tables)
|
||
cache_hit = cached is not None
|
||
if cache_hit:
|
||
exec_result, data_as_of = cached
|
||
latency = 0
|
||
else:
|
||
t0 = time.time()
|
||
try:
|
||
exec_result = self.repo.execute_readonly(sql_text)
|
||
data_as_of = self.repo.get_data_as_of()
|
||
latency = int((time.time() - t0) * 1000)
|
||
self.cache.set_result(perm_fp, sql_text, (exec_result, data_as_of), vres.tables)
|
||
except Exception as exc: # noqa: BLE001
|
||
if _is_timeout_error(exc):
|
||
resp = self._escalate(
|
||
f"查询超时(>{settings.analyst_sql_timeout_s}s)",
|
||
trace_id,
|
||
error_code="EXEC_TIMEOUT",
|
||
)
|
||
else:
|
||
resp = self._escalate(f"SQL 执行失败:{exc}", trace_id, error_code="EXEC_ERROR")
|
||
return self._audit_terminal(
|
||
question, auth, session_id, trace_id, resp, sql=sql_text
|
||
)
|
||
|
||
table = TableData(columns=exec_result["columns"], rows=exec_result["rows"])
|
||
empty_state = classify_empty(exec_result["rows"], sql_text)
|
||
|
||
guard_result: GuardrailResult | None = None
|
||
answer = ""
|
||
status = "success"
|
||
if interpret:
|
||
try:
|
||
answer, guard_result, g_usage = self._generate_verified_answer(
|
||
question, sql_text, table, empty_state, domain=domain
|
||
)
|
||
cost_est += estimate_cost(g_usage)
|
||
except Exception as exc: # noqa: BLE001
|
||
resp = self._escalate(f"解读生成失败:{exc}", trace_id, error_code="INTERPRET_ERROR")
|
||
return self._audit_terminal(
|
||
question, auth, session_id, trace_id, resp, sql=sql_text
|
||
)
|
||
status = "degrade" if (guard_result is not None and not guard_result.passed) else "success"
|
||
if status == "degrade":
|
||
answer = "解读校验未通过,请以下方表格数据为准。"
|
||
|
||
disclaimer = DISCLAIMER
|
||
if domain == "self":
|
||
disclaimer = f"{DISCLAIMER} {CUSTOMER_AI_RISK_NOTE}"
|
||
if interpret and answer and CUSTOMER_AI_RISK_NOTE not in answer:
|
||
answer = f"{answer.rstrip()} {CUSTOMER_AI_RISK_NOTE}"
|
||
resp = AnalystResponse(
|
||
answer=answer,
|
||
table=table,
|
||
sql=sql_text,
|
||
meta=Meta(
|
||
exec_ms=latency,
|
||
row_count=len(exec_result["rows"]),
|
||
cache_hit=cache_hit,
|
||
template_hit=template_hit,
|
||
template_key=template_key,
|
||
data_as_of=data_as_of,
|
||
source="jinrong_core",
|
||
cost_est=round(cost_est, 6),
|
||
),
|
||
disclaimer=disclaimer,
|
||
status=status,
|
||
trace_id=trace_id,
|
||
)
|
||
|
||
# 7) 留痕(D-04)
|
||
self._persist(
|
||
question,
|
||
sql_text,
|
||
sql_hash,
|
||
resp,
|
||
auth,
|
||
session_id,
|
||
trace_id,
|
||
latency,
|
||
empty_state,
|
||
guard_result,
|
||
template_key=template_key,
|
||
)
|
||
return resp
|
||
|
||
def sample_by_trace(
|
||
self,
|
||
trace_id: str,
|
||
auth: AnalystAuthContext,
|
||
limit: int = 5,
|
||
) -> AnalystResponse:
|
||
"""N-03:按 trace 回放 SQL 并抽样明细行(与聚合同一 SQL 链路)。"""
|
||
trace_id = trace_id.strip()
|
||
try:
|
||
domain = assert_analyst_query_access(auth)
|
||
except AnalystAuthError as exc:
|
||
return self._deny(exc.error_code, exc.message, trace_id)
|
||
|
||
row = self.repo.get_query_log_by_trace(trace_id)
|
||
if not row:
|
||
return self._deny("NOT_FOUND", "未找到该 trace 的成功问数记录", trace_id, domain)
|
||
if row.get("staff_id") != auth.subject_id and "analyst" not in auth.roles:
|
||
return self._deny("AUTH_403", "无权查看他人问数溯源", trace_id, domain)
|
||
|
||
sql_text = (row.get("generated_sql") or "").strip()
|
||
if not sql_text:
|
||
return self._deny("NOT_FOUND", "该 trace 无可用 SQL", trace_id, domain)
|
||
|
||
scope: list[str] = resolve_analyst_scope(auth, domain, self.repo)
|
||
try:
|
||
vres = validate(sql_text, domain, scope)
|
||
except SqlGuardError as exc:
|
||
return self._deny(exc.error_code, exc.message, trace_id, domain)
|
||
|
||
sample_sql = _build_sample_sql(sql_text, limit)
|
||
try:
|
||
validate(sample_sql, domain, scope)
|
||
except SqlGuardError as exc:
|
||
return self._deny(exc.error_code, exc.message, trace_id, domain)
|
||
|
||
t0 = time.time()
|
||
try:
|
||
exec_result = self.repo.execute_readonly(sample_sql)
|
||
except Exception as exc: # noqa: BLE001
|
||
return self._escalate(f"抽样执行失败:{exc}", trace_id)
|
||
|
||
latency = int((time.time() - t0) * 1000)
|
||
table = TableData(columns=exec_result["columns"], rows=exec_result["rows"])
|
||
return AnalystResponse(
|
||
answer=f"已按原 SQL 抽样 {len(table.rows)} 行明细,供与聚合结果核对。",
|
||
table=table,
|
||
sql=sample_sql,
|
||
meta=Meta(
|
||
exec_ms=latency,
|
||
row_count=len(table.rows),
|
||
cache_hit=False,
|
||
data_as_of=self.repo.get_data_as_of(),
|
||
source="jinrong_core",
|
||
),
|
||
disclaimer=DISCLAIMER,
|
||
status="success",
|
||
trace_id=trace_id,
|
||
)
|
||
|
||
def escalate(
|
||
self,
|
||
trace_id: str,
|
||
question: str,
|
||
reason: str,
|
||
auth: AnalystAuthContext,
|
||
) -> dict[str, Any]:
|
||
"""N-07:写审计并返回转人工引导。"""
|
||
trace_id = trace_id.strip()
|
||
payload = {"question": question, "reason": reason}
|
||
try:
|
||
self.repo.log_audit(
|
||
trace_id=trace_id,
|
||
event_type="analyst_escalate",
|
||
actor_id=auth.subject_id,
|
||
decision="escalate",
|
||
input_summary=payload,
|
||
)
|
||
except Exception: # noqa: BLE001
|
||
pass
|
||
return {
|
||
"ok": True,
|
||
"trace_id": trace_id,
|
||
"message": "已登记转人工请求,分析师将按 trace 核对问数留痕。",
|
||
"hints": _ESCALATE_HINTS,
|
||
}
|
||
|
||
def interpret(
|
||
self,
|
||
req: InterpretRequest,
|
||
auth: AnalystAuthContext,
|
||
) -> AnalystResponse:
|
||
"""仅解读:上下文为客户端提交的上一次问数快照(无 Chat 历史)。"""
|
||
trace_id = req.trace_id or f"trace-{uuid.uuid4().hex[:16]}"
|
||
auth.trace_id = trace_id
|
||
try:
|
||
domain = assert_analyst_query_access(auth)
|
||
except AnalystAuthError as exc:
|
||
return self._deny(exc.error_code, exc.message, trace_id)
|
||
|
||
terminal = {"clarify", "deny", "error", "escalate"}
|
||
if req.status in terminal:
|
||
return AnalystResponse(
|
||
answer=req.answer,
|
||
table=req.table,
|
||
sql=req.sql,
|
||
meta=req.meta,
|
||
disclaimer=DISCLAIMER,
|
||
status=req.status,
|
||
trace_id=trace_id,
|
||
)
|
||
|
||
table = req.table
|
||
empty_state = classify_empty(table.rows, req.sql)
|
||
cost_est = float(req.meta.cost_est or 0.0)
|
||
try:
|
||
answer, guard_result, g_usage = self._generate_verified_answer(
|
||
req.question, req.sql, table, empty_state, domain=domain
|
||
)
|
||
cost_est += estimate_cost(g_usage)
|
||
except Exception as exc: # noqa: BLE001
|
||
return self._escalate(f"解读生成失败:{exc}", trace_id, error_code="INTERPRET_ERROR")
|
||
|
||
status = "degrade" if (guard_result is not None and not guard_result.passed) else "success"
|
||
if status == "degrade":
|
||
answer = "解读校验未通过,请以下方表格数据为准。"
|
||
disclaimer = DISCLAIMER
|
||
if domain == "self":
|
||
disclaimer = f"{DISCLAIMER} {CUSTOMER_AI_RISK_NOTE}"
|
||
if answer and CUSTOMER_AI_RISK_NOTE not in answer:
|
||
answer = f"{answer.rstrip()} {CUSTOMER_AI_RISK_NOTE}"
|
||
meta = req.meta.model_copy(update={"cost_est": round(cost_est, 6)})
|
||
return AnalystResponse(
|
||
answer=answer,
|
||
table=table,
|
||
sql=req.sql,
|
||
meta=meta,
|
||
disclaimer=disclaimer,
|
||
status=status,
|
||
trace_id=trace_id,
|
||
)
|
||
|
||
# ---------- 各步骤 ----------
|
||
def _detect_ambiguity(self, question: str) -> Ambiguity | None:
|
||
terms = self._metric_terms(question)
|
||
if not terms:
|
||
return None
|
||
# 只取最长(最具体)的指标词判定,避免"持仓规模"里的"规模"误触发歧义
|
||
resolved = self.registry.resolve(terms[0])
|
||
return resolved if isinstance(resolved, Ambiguity) else None
|
||
|
||
def _metric_terms(self, question: str) -> list[str]:
|
||
"""从问题里捞出可能的指标词(口径字典别名)。"""
|
||
terms: list[str] = []
|
||
for m in self.registry.all():
|
||
for n in [m.name, *m.aliases]:
|
||
if n and n in question:
|
||
terms.append(n)
|
||
return sorted(terms, key=len, reverse=True)
|
||
|
||
def _generate_sql(self, question: str, domain: str, scope: list[str]) -> tuple[str, dict]:
|
||
dict_hint = "\n".join(f"- {m.name}({m.key}):{m.definition}" for m in self.registry.all())
|
||
scope_hint = "无限制"
|
||
if domain == "assigned":
|
||
scope_hint = f"只能查询以下客户:customer_id IN ({', '.join(scope)});涉及客户的查询必须带此过滤"
|
||
elif domain == "self" and scope:
|
||
scope_hint = (
|
||
f"只能查询客户 {scope[0]} 本人的数据;涉及客户表必须带 customer_id = '{scope[0]}' 条件"
|
||
)
|
||
elif domain == "aggregate":
|
||
scope_hint = "只能输出聚合结果,禁止查单个客户或按 customer_id 分组/筛选"
|
||
prompt = (
|
||
f"{SQL_GEN_SYSTEM}\n\n表结构:\n{SCHEMA_PROMPT}\n\n口径字典:\n{dict_hint}\n\n"
|
||
f"权限约束:{scope_hint}\n\n用户问题:{question}\n\nSQL:"
|
||
)
|
||
text, usage = self.llm.complete(
|
||
[{"role": "system", "content": SQL_GEN_SYSTEM}, {"role": "user", "content": prompt}],
|
||
temperature=0.0,
|
||
max_tokens=800,
|
||
)
|
||
return extract_sql(text), usage
|
||
|
||
def _generate_verified_answer(
|
||
self, question: str, sql: str, table: TableData, empty_state: str, *, domain: str = "full"
|
||
) -> tuple[str, GuardrailResult | None, dict]:
|
||
summary = self._summarize(table)
|
||
empty_note = self._empty_note(empty_state)
|
||
answer, usage = self._generate_answer(question, sql, summary, empty_note, domain=domain)
|
||
guard = verify(answer, table, None)
|
||
if not guard.passed:
|
||
# 重试 1 次(加强提示)
|
||
answer2, usage2 = self._generate_answer(
|
||
question,
|
||
sql,
|
||
summary,
|
||
empty_note + " 特别注意:所有数字必须与结果逐字一致。",
|
||
domain=domain,
|
||
)
|
||
for k, v in usage2.items():
|
||
cur = usage.get(k)
|
||
if isinstance(cur, (int, float)) and isinstance(v, (int, float)):
|
||
usage[k] = cur + v
|
||
guard2 = verify(answer2, table, None)
|
||
if guard2.passed:
|
||
return answer2, guard2, usage
|
||
return answer2, guard2, usage
|
||
return answer, guard, usage
|
||
|
||
def _generate_answer(
|
||
self, question: str, sql: str, summary: str, note: str, *, domain: str = "full"
|
||
) -> tuple[str, dict]:
|
||
system = CUSTOMER_ANSWER_SYSTEM if domain == "self" else ANSWER_SYSTEM
|
||
prompt = (
|
||
f"问题:{question}\nSQL:{sql}\n查询结果摘要:{summary}\n空态说明:{note}\n"
|
||
f"请用 1~3 句人话解读:"
|
||
)
|
||
return self.llm.complete(
|
||
[{"role": "system", "content": system}, {"role": "user", "content": prompt}],
|
||
temperature=0.2,
|
||
max_tokens=500,
|
||
)
|
||
|
||
def _summarize(self, table: TableData, max_rows: int = 10) -> str:
|
||
if not table.rows:
|
||
return "(空)"
|
||
head = table.rows[:max_rows]
|
||
return f"列={table.columns} 行数={len(table.rows)} 前{len(head)}行={head}"
|
||
|
||
def _empty_note(self, empty_state: str) -> str:
|
||
return {
|
||
"zero": "结果为 0(确有数据,聚合值为 0)。",
|
||
"no_data": "源无此数据。",
|
||
"not_match": "查询条件未命中任何记录。",
|
||
"has_data": "",
|
||
}.get(empty_state, "")
|
||
|
||
def _clarify(self, amb: Ambiguity, trace_id: str) -> AnalystResponse:
|
||
cands = ";".join(f"{c.name}({c.definition})" for c in amb.candidates)
|
||
return AnalystResponse(
|
||
answer=f"“{amb.term}”有多个口径,请确认您指哪一个:{cands}",
|
||
status="clarify",
|
||
trace_id=trace_id,
|
||
)
|
||
|
||
def _deny(self, code: str, msg: str, trace_id: str, domain: str = "") -> AnalystResponse:
|
||
suggestions = {
|
||
"AUTH_403_NOT_ASSIGNED": ["你仅能查名下客户的数据,可尝试问自己名下客户的持仓、风险分布等。"],
|
||
"AUTH_403_NOT_OWNER": ["你仅能查本人数据,可尝试问我的持仓、近30日交易笔数、风险等级等。"],
|
||
"AUTH_403_SCOPE": ["你仅能查聚合数据,如近30天申购金额、各产品类型规模等。"],
|
||
}.get(code)
|
||
return AnalystResponse(
|
||
answer=f"无法执行:{msg}",
|
||
status="deny",
|
||
error_code=code,
|
||
suggestions=suggestions,
|
||
trace_id=trace_id,
|
||
)
|
||
|
||
def _escalate(self, msg: str, trace_id: str, *, error_code: str = "EXEC_ERROR") -> AnalystResponse:
|
||
return AnalystResponse(
|
||
answer=f"本次查询未能完成:{msg}。{' '.join(_ESCALATE_HINTS)}",
|
||
status="escalate",
|
||
error_code=error_code,
|
||
suggestions=_ESCALATE_HINTS,
|
||
trace_id=trace_id,
|
||
)
|
||
|
||
def _error(self, msg: str, trace_id: str) -> AnalystResponse:
|
||
return self._escalate(msg, trace_id)
|
||
|
||
def _persist(
|
||
self,
|
||
question,
|
||
sql,
|
||
sql_hash,
|
||
resp,
|
||
auth,
|
||
session_id,
|
||
trace_id,
|
||
latency,
|
||
empty_state,
|
||
guard_result,
|
||
*,
|
||
template_key: str | None = None,
|
||
) -> None:
|
||
if template_key:
|
||
sql_source = "template"
|
||
elif resp.meta.cache_hit:
|
||
sql_source = "cache"
|
||
else:
|
||
sql_source = "llm"
|
||
summary = {
|
||
"status": resp.status,
|
||
"empty_state": empty_state,
|
||
"guardrail": "passed" if (guard_result is None or guard_result.passed) else "degraded",
|
||
"source": sql_source,
|
||
"template_key": template_key,
|
||
}
|
||
try:
|
||
self.repo.log_query(
|
||
session_id=session_id, trace_id=trace_id, staff_id=auth.subject_id,
|
||
nl_question=question, generated_sql=sql, sql_hash=sql_hash,
|
||
row_count=resp.meta.row_count,
|
||
exec_status="success" if resp.status in ("success", "degrade") else "blocked",
|
||
result_summary=summary, exec_latency_ms=latency,
|
||
has_disclaimer=True,
|
||
)
|
||
self.repo.log_audit(
|
||
trace_id=trace_id, event_type="analyst_query", actor_id=auth.subject_id,
|
||
decision=resp.status, input_summary={"question": question},
|
||
)
|
||
except Exception: # noqa: BLE001
|
||
pass
|
||
|
||
def _audit_terminal(
|
||
self,
|
||
question: str,
|
||
auth: AnalystAuthContext,
|
||
session_id: str,
|
||
trace_id: str,
|
||
resp: AnalystResponse,
|
||
*,
|
||
sql: str = "",
|
||
) -> AnalystResponse:
|
||
"""阻断/clarify/error 路径留痕(TEST-AN-001 缺口 D)。"""
|
||
sql_text = (sql or "").strip()
|
||
sql_hash = self.cache.sql_hash(sql_text) if sql_text else "blocked"
|
||
summary = {
|
||
"status": resp.status,
|
||
"error_code": resp.error_code,
|
||
"source": "blocked",
|
||
}
|
||
try:
|
||
self.repo.log_query(
|
||
session_id=session_id,
|
||
trace_id=trace_id,
|
||
staff_id=auth.subject_id,
|
||
nl_question=question,
|
||
generated_sql=sql_text,
|
||
sql_hash=sql_hash or None,
|
||
row_count=0,
|
||
exec_status="blocked",
|
||
result_summary=summary,
|
||
exec_latency_ms=0,
|
||
has_disclaimer=False,
|
||
)
|
||
self.repo.log_audit(
|
||
trace_id=trace_id,
|
||
event_type="analyst_query",
|
||
actor_id=auth.subject_id,
|
||
decision=resp.status,
|
||
input_summary={"question": question, "error_code": resp.error_code},
|
||
)
|
||
except Exception: # noqa: BLE001
|
||
pass
|
||
return resp
|
||
|
||
|
||
def build_graph(agent: AnalystAgent):
|
||
"""LangGraph StateGraph 适配(架构对齐用;核心逻辑仍在 run())。"""
|
||
from langgraph.graph import END, StateGraph
|
||
|
||
def node_run(state: dict) -> dict:
|
||
resp = agent.run(
|
||
state["question"], state["auth"], state.get("session_id"), state.get("trace_id")
|
||
)
|
||
return {"response": resp}
|
||
|
||
g = StateGraph(dict)
|
||
g.add_node("run", node_run)
|
||
g.set_entry_point("run")
|
||
g.add_edge("run", END)
|
||
return g.compile()
|