Files
group_xinghuo_jinrong/tests/test_wave6_analyst_agent.py
T
zhanghongyu_0626 601ed258f3 feat(analyst): Enhance asset management and ambiguity detection features
- Added new endpoints to the analyst API for managing assets, including `GET /assets` to list assets and `POST /assets/{kind}/{asset_id}/publish` to publish assets.
- Introduced `DictAmbiguityCheckRequest` schema for checking metric ambiguities, enhancing the analyst's ability to clarify definitions and aliases.
- Implemented `detect_dict_ambiguity` function to analyze potential ambiguities in metrics, providing structured feedback for users.
- Updated `AnalystAgent` to support the new asset management functionalities and ambiguity detection logic, improving overall user experience.
- Enhanced existing schemas and services to accommodate new features, ensuring robust data handling and validation.

This update significantly improves the analyst API's capabilities, allowing for better asset management and clarity in metric definitions.
2026-09-12 15:14:01 +08:00

408 lines
16 KiB
Python

"""analyst_agent 编排测试(Wave 6)。"""
import unittest
import pytest
from app.api.analyst_auth_adapter import AnalystAuthContext
from app.service.analyst_agent import AnalystAgent, _nl_sql_hints
from app.service.template_service import QueryTemplate, TemplateService
class FakeLLM:
def __init__(self, sql, answers):
self.sql = sql
self.answers = list(answers)
self.calls = 0
def complete(self, messages, temperature=0, max_tokens=2048):
self.calls += 1
usage = {"prompt_tokens": 10, "completion_tokens": 10}
if self.calls == 1:
return self.sql, usage
ans = self.answers.pop(0) if self.answers else "无解读"
return ans, usage
class FakeRepo:
def __init__(self, rows=(), columns=(), scope=None):
self.rows = list(rows)
self.columns = list(columns)
self.scope = scope or []
self.logged = []
def resolve_advisor_scope(self, sid):
return self.scope
def execute_readonly(self, sql):
return {"columns": self.columns, "rows": self.rows}
def get_data_as_of(self):
return "2026-09-04"
def log_query(self, **kw):
self.logged.append(kw)
def log_audit(self, **kw):
pass
def list_published_templates(self):
return []
def list_published_metrics(self):
return []
def list_published_few_shots(self, limit=5):
return []
def get_query_log_by_trace(self, trace_id):
return {
"staff_id": "STAFF-A",
"generated_sql": "SELECT COUNT(*) AS c FROM core_customer",
"exec_status": "success",
}
def ctx(roles, subject="STAFF-A", *, token_type="staff", customer_id=None):
return AnalystAuthContext(
subject_id=subject,
token_type=token_type,
roles=roles,
customer_id=customer_id,
)
class TestAgentOrchestration(unittest.TestCase):
def test_clarify_trade_flow_ambiguity(self):
agent = AnalystAgent(llm=FakeLLM("SELECT 1", []), repo=FakeRepo())
resp = agent.run("近30天交易流水是多少", ctx(["analyst"]))
self.assertEqual(resp.status, "clarify")
self.assertIn("流水", resp.answer)
def test_clarify_ambiguity(self):
agent = AnalystAgent(llm=FakeLLM("SELECT 1", []), repo=FakeRepo())
resp = agent.run("我名下的规模是多少", ctx(["analyst"]))
self.assertEqual(resp.status, "clarify")
def test_success_query_only(self):
repo = FakeRepo(rows=[[33]], columns=["c"])
agent = AnalystAgent(
llm=FakeLLM("SELECT COUNT(*) AS c FROM core_customer", ["共 33 个客户"]),
repo=repo,
)
resp = agent.run("客户总数是多少", ctx(["analyst"]))
self.assertEqual(resp.status, "success")
self.assertEqual(resp.table.rows, [[33]])
self.assertEqual(resp.answer, "")
self.assertEqual(len(repo.logged), 1)
self.assertEqual(agent.llm.calls, 1)
def test_success_with_interpret_flag(self):
repo = FakeRepo(rows=[[33]], columns=["c"])
llm = FakeLLM("SELECT COUNT(*) AS c FROM core_customer", ["共 33 个客户"])
agent = AnalystAgent(llm=llm, repo=repo)
resp = agent.run("客户总数是多少", ctx(["analyst"]), interpret=True)
self.assertEqual(resp.status, "success")
self.assertIn("33", resp.answer)
self.assertEqual(llm.calls, 2)
def test_interpret_from_snapshot(self):
from app.model.analyst_schemas import InterpretRequest, Meta, TableData
class AnswerOnlyLLM:
def complete(self, messages, temperature=0, max_tokens=2048):
return "共 33 个客户", {"prompt_tokens": 1, "completion_tokens": 1}
repo = FakeRepo(rows=[[33]], columns=["c"])
agent = AnalystAgent(llm=AnswerOnlyLLM(), repo=repo)
req = InterpretRequest(
question="客户总数是多少",
status="success",
table=TableData(columns=["c"], rows=[[33]]),
sql="SELECT COUNT(*) AS c FROM core_customer",
meta=Meta(row_count=1),
)
resp = agent.interpret(req, ctx(["analyst"]))
self.assertEqual(resp.status, "success")
self.assertIn("33", resp.answer)
def test_interpret_echo_deny_without_llm(self):
from app.model.analyst_schemas import InterpretRequest
llm = FakeLLM("SELECT 1", [])
agent = AnalystAgent(llm=llm, repo=FakeRepo())
req = InterpretRequest(
question="查别人",
status="deny",
answer="无法执行:权限不足",
)
resp = agent.interpret(req, ctx(["customer"], "CUST-1", customer_id="CUST-1"))
self.assertEqual(resp.status, "deny")
self.assertEqual(resp.answer, "无法执行:权限不足")
self.assertEqual(llm.calls, 0)
def test_deny_bad_sql(self):
repo = FakeRepo()
agent = AnalystAgent(llm=FakeLLM("INSERT INTO core_customer VALUES (1)", []), repo=repo)
resp = agent.run("删库", ctx(["analyst"]))
self.assertEqual(resp.status, "deny")
self.assertEqual(resp.error_code, "SQL_NOT_SELECT")
self.assertEqual(len(repo.logged), 1)
self.assertEqual(repo.logged[0].get("exec_status"), "blocked")
def test_degrade_wrong_number(self):
repo = FakeRepo(rows=[[33]], columns=["c"])
agent = AnalystAgent(
llm=FakeLLM("SELECT COUNT(*) FROM core_customer", ["共 999 个客户", "共 999 个客户"]),
repo=repo,
)
resp = agent.run("客户总数", ctx(["analyst"]), interpret=True)
self.assertEqual(resp.status, "degrade")
def test_advisor_out_of_scope_deny(self):
repo = FakeRepo(scope=["CUST-1001"])
agent = AnalystAgent(
llm=FakeLLM("SELECT * FROM core_holding WHERE customer_id='CUST-1004'", []),
repo=repo,
)
resp = agent.run("查 CUST-1004 持仓", ctx(["advisor"], "STAFF-B"))
self.assertEqual(resp.status, "deny")
self.assertEqual(resp.error_code, "AUTH_403_NOT_ASSIGNED")
def test_customer_self_success(self):
repo = FakeRepo(rows=[[2]], columns=["cnt"])
agent = AnalystAgent(
llm=FakeLLM(
"SELECT COUNT(*) AS cnt FROM core_trade WHERE customer_id='CUST-9527'",
["近阶段共有 2 笔交易"],
),
repo=repo,
)
resp = agent.run(
"我有多少笔交易",
ctx(["customer"], "CUST-9527", token_type="customer", customer_id="CUST-9527"),
interpret=True,
)
self.assertEqual(resp.status, "success")
self.assertIn("AI 分析有风险", resp.answer)
def test_template_hit_skips_llm_sql(self):
tpl = QueryTemplate(
template_key="customer_total_count",
template_sql="SELECT COUNT(*) AS c FROM core_customer",
params_schema={"match_all": ["客户", "总数"]},
)
repo = FakeRepo(rows=[[33]], columns=["c"])
llm = FakeLLM("SELECT 1", ["共 33 个客户"])
agent = AnalystAgent(
llm=llm,
repo=repo,
templates=TemplateService(templates=[tpl]),
)
resp = agent.run("客户总数是多少", ctx(["analyst"]))
self.assertEqual(resp.status, "success")
self.assertTrue(resp.meta.template_hit)
self.assertEqual(resp.meta.template_key, "customer_total_count")
self.assertIn("COUNT(*)", resp.sql)
self.assertEqual(llm.calls, 0)
def test_dual_extreme_question_enriches_llm_prompt(self):
q = "最大和最低涨幅的产品分别是哪两个"
enriched = _nl_sql_hints(q)
self.assertIn("UNION ALL", enriched)
self.assertIn("禁止仅 LIMIT 1", enriched)
def test_generate_sql_uses_db_few_shots_only(self):
class CaptureLLM:
last_user = ""
def complete(self, messages, temperature=0.0, max_tokens=2048):
self.last_user = messages[-1]["content"]
return "SELECT 1 AS x", {"prompt_tokens": 1, "completion_tokens": 1}
class FewShotRepo(FakeRepo):
def list_published_few_shots(self, limit=5):
return [
{
"question": "平台持仓市值最大的客户是谁,其持有产品的净值走势如何",
"sql_text": "SELECT 1 FROM core_product_nav",
}
]
llm = CaptureLLM()
agent = AnalystAgent(llm=llm, repo=FewShotRepo())
agent._generate_sql("x", "full", [])
self.assertIn("平台持仓市值最大的客户", llm.last_user)
self.assertIn("core_product_nav", llm.last_user)
def test_generate_sql_empty_when_no_db_few_shots(self):
class CaptureLLM:
last_user = ""
def complete(self, messages, temperature=0.0, max_tokens=2048):
self.last_user = messages[-1]["content"]
return "SELECT 1 AS x", {"prompt_tokens": 1, "completion_tokens": 1}
llm = CaptureLLM()
agent = AnalystAgent(llm=llm, repo=FakeRepo())
agent._generate_sql("客户总数", "full", [])
self.assertNotIn("参考 few-shot", llm.last_user)
def test_clarify_includes_structured_payload(self):
agent = AnalystAgent(llm=FakeLLM("SELECT 1", []), repo=FakeRepo())
resp = agent.run("我名下的规模是多少", ctx(["analyst"]))
self.assertEqual(resp.status, "clarify")
self.assertIsNotNone(resp.clarify)
assert resp.clarify is not None
self.assertTrue(resp.clarify.options)
self.assertIn("规模", resp.clarify.term)
def test_generate_sql_prefers_db_few_shots(self):
class CaptureLLM:
last_user = ""
def complete(self, messages, temperature=0.0, max_tokens=2048):
self.last_user = messages[-1]["content"]
return "SELECT 1 AS x", {"prompt_tokens": 1, "completion_tokens": 1}
class FewShotRepo(FakeRepo):
def list_published_few_shots(self, limit=5):
return [{"question": "库内示例问法", "sql_text": "SELECT 2 AS y FROM core_product"}]
llm = CaptureLLM()
agent = AnalystAgent(llm=llm, repo=FewShotRepo())
agent._generate_sql("x", "full", [])
self.assertIn("库内示例问法", llm.last_user)
self.assertIn("core_product", llm.last_user)
def test_template_nav_max_min_skips_llm(self):
tpl = QueryTemplate(
template_key="product_nav_max_min_latest",
template_sql=(
"(SELECT 'max' AS extremum, n.product_id, n.daily_chg_pct FROM core_product_nav n "
"ORDER BY n.daily_chg_pct DESC LIMIT 1) "
"UNION ALL "
"(SELECT 'min' AS extremum, n.product_id, n.daily_chg_pct FROM core_product_nav n "
"ORDER BY n.daily_chg_pct ASC LIMIT 1)"
),
params_schema={"match_phrases": ["目前净值最高和最低的两个产品"]},
)
repo = FakeRepo(
rows=[["max", "P1", 1.2], ["min", "P2", -0.5]],
columns=["extremum", "product_id", "daily_chg_pct"],
)
llm = FakeLLM("SELECT 1 LIMIT 1", [])
agent = AnalystAgent(
llm=llm,
repo=repo,
templates=TemplateService(templates=[tpl]),
)
resp = agent.run("目前净值最高和最低的两个产品", ctx(["analyst"]))
self.assertEqual(resp.status, "success")
self.assertTrue(resp.meta.template_hit)
self.assertEqual(resp.meta.row_count, 2)
self.assertIn("UNION ALL", resp.sql.upper())
self.assertEqual(llm.calls, 0)
def test_analyze_text_only_json(self):
json_out = '{"answer":"产品净值整体平稳,未见异常波动。","chart_type":"none","title":"","x_field":null,"y_fields":[]}'
repo = FakeRepo(rows=[["P1", 1.1], ["P2", 1.2]], columns=["product_id", "nav"])
agent = AnalystAgent(llm=FakeLLM(json_out, []), repo=repo)
from app.model.analyst_schemas import AnalyzeRequest, Meta, TableData
req = AnalyzeRequest(
question="产品净值",
analysis_prompt="三句话总结",
status="success",
table=TableData(columns=["product_id", "nav"], rows=[["P1", 1.1], ["P2", 1.2]]),
sql="SELECT 1",
meta=Meta(row_count=2),
)
resp = agent.analyze(req, ctx(["analyst"]))
self.assertEqual(resp.status, "success")
self.assertEqual(resp.analysis_kind, "text")
self.assertIsNone(resp.chart)
def test_analyze_chart_only_json(self):
json_out = (
'{"answer":"","chart_type":"line","title":"趋势","reason":"日期+数值",'
'"x_field":"d","y_fields":["v"],"series_field":"s"}'
)
repo = FakeRepo(
rows=[["2026-01-01", "A", 1.0], ["2026-01-02", "A", 1.1]],
columns=["d", "s", "v"],
)
agent = AnalystAgent(llm=FakeLLM(json_out, []), repo=repo)
from app.model.analyst_schemas import AnalyzeRequest, Meta, TableData
req = AnalyzeRequest(
question="趋势",
analysis_prompt="画折线图",
status="success",
table=TableData(
columns=["d", "s", "v"],
rows=[["2026-01-01", "A", 1.0], ["2026-01-02", "A", 1.1]],
),
sql="SELECT 1",
meta=Meta(row_count=2),
)
resp = agent.analyze(req, ctx(["analyst"]))
self.assertEqual(resp.status, "success")
self.assertEqual(resp.analysis_kind, "chart")
self.assertIsNotNone(resp.chart)
self.assertEqual(resp.chart.chart_type, "line")
def test_sql_exec_failure_escalates(self):
class FailRepo(FakeRepo):
def execute_readonly(self, sql):
raise TimeoutError("read timed out")
class NoCache:
def permission_fingerprint(self, *a, **k):
return "fp"
def sql_hash(self, sql):
return "hash"
def get_result(self, *a, **k):
return None
def set_result(self, *a, **k):
pass
agent = AnalystAgent(
llm=FakeLLM("SELECT COUNT(*) AS c FROM core_customer", []),
repo=FailRepo(),
templates=TemplateService(templates=[]),
cache=NoCache(),
)
resp = agent.run("客户总数", ctx(["analyst"]))
self.assertEqual(resp.status, "escalate")
self.assertIn(resp.error_code, ("EXEC_TIMEOUT", "EXEC_ERROR"))
def test_sample_by_trace(self):
repo = FakeRepo(rows=[[1]], columns=["c"])
agent = AnalystAgent(llm=FakeLLM("SELECT 1", []), repo=repo)
resp = agent.sample_by_trace("trace-1", ctx(["analyst"], subject="STAFF-A"))
self.assertEqual(resp.status, "success")
self.assertIn("LIMIT", resp.sql)
def test_escalate_audit(self):
repo = FakeRepo()
agent = AnalystAgent(llm=FakeLLM("SELECT 1", []), repo=repo)
out = agent.escalate("trace-x", "q", "timeout", ctx(["analyst"]))
self.assertTrue(out["ok"])
@pytest.mark.integration
class TestAgentReal(unittest.TestCase):
@pytest.mark.skip(reason="需要真实 MySQL + DeepSeek Key")
def test_real_end_to_end(self):
from app.service.analytics_repo import AnalyticsRepo
from app.service.llm import DeepSeekLLM
agent = AnalystAgent(llm=DeepSeekLLM(), repo=AnalyticsRepo())
resp = agent.run("客户总数是多少", ctx(["analyst"]))
self.assertIn(resp.status, ("success", "degrade"))
self.assertTrue(resp.sql)
self.assertGreater(resp.meta.row_count, 0)