Files
zhanghongyu_0626 af9fb719e5 refactor(analyst): Update authentication context and add smoke test script
- Replaced `get_auth_context` with `get_platform_auth_context` in `analyst.py` to enhance authentication handling.
- Added a new smoke test script `smoke_analyst.py` for testing the data analysis agent with both fake and live LLM configurations.
- Updated TODO documentation to reflect the completion of Scope B smoke tests, ensuring clarity on testing status.

This update improves the authentication mechanism and introduces a comprehensive testing approach for the data analysis agent.
2026-09-09 21:09:03 +08:00

181 lines
6.4 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""数据分析 Agent Scope B 冒烟(四 Demo 角色 + dashboard + 归属拒答)。
默认用 FakeLLM(不依赖 DeepSeek);加 --live-llm 走真实 Key(需 .env DEEPSEEK_API_KEY 有效)。
"""
from __future__ import annotations
import argparse
import sys
from pathlib import Path
ROOT = Path(__file__).resolve().parents[2]
sys.path.insert(0, str(ROOT))
from fastapi.testclient import TestClient
from app.api import analyst as analyst_api
from app.main import app
from app.service.analyst_agent import AnalystAgent
from app.service.auth_service import issue_dev_token
client = TestClient(app)
PASS = FAIL = 0
class FakeLLM:
"""按问题关键词返回可过 sql_guard 的只读 SQL。"""
def __init__(self) -> None:
self._calls = 0
self._last_answer = "解读。"
def complete(self, messages, temperature=0.0, max_tokens=2048):
text = " ".join(str(m.get("content", "")) for m in messages)
usage = {"prompt_tokens": 10, "completion_tokens": 10}
if "CUST-9999" in text or "CUST-9999" in messages[-1].get("content", ""):
sql = "SELECT * FROM core_holding WHERE customer_id='CUST-9999'"
ans = "越权查询。"
elif "CUST-9527" in text or "我有多少笔交易" in text:
sql = (
"SELECT COUNT(*) AS cnt FROM core_trade "
"WHERE customer_id='CUST-9527'"
)
ans = "您共有 2 笔交易。"
elif "名下" in text and "客户" in text:
sql = "SELECT COUNT(DISTINCT customer_id) AS cnt FROM core_customer_advisor WHERE advisor_id='STAFF-10086' AND rel_status='active'"
ans = "名下有若干客户。"
elif "预警" in text:
sql = (
"SELECT COUNT(*) AS cnt FROM jinrong_agent.risk_alert "
"WHERE status='pending_review'"
)
ans = "有待处理预警。"
else:
sql = "SELECT COUNT(*) AS cnt FROM core_customer"
ans = "共 33 个客户。"
self._last_answer = ans
self._calls += 1
if self._calls == 1:
return sql, usage
return self._last_answer, usage
class FakeRepo:
def resolve_advisor_scope(self, advisor_id: str):
return ["CUST-9527", "CUST-1001"]
def execute_readonly(self, sql: str):
if "core_trade" in sql and "CUST-9527" in sql:
return {"columns": ["cnt"], "rows": [[2]]}
if "risk_alert" in sql:
return {"columns": ["cnt"], "rows": [[5]]}
if "core_customer_advisor" in sql:
return {"columns": ["cnt"], "rows": [[3]]}
return {"columns": ["cnt"], "rows": [[33]]}
def get_data_as_of(self):
return "2026-09-04"
def log_query(self, **kw):
pass
def log_audit(self, **kw):
pass
def ok(name: str, cond: bool, detail: str = "") -> None:
global PASS, FAIL
if cond:
PASS += 1
print(f" [PASS] {name}" + (f" — {detail}" if detail else ""))
else:
FAIL += 1
print(f" [FAIL] {name}" + (f" — {detail}" if detail else ""))
def token(sub: str, roles: str, *, token_type: str = "staff", customer_id: str | None = None) -> str:
return issue_dev_token(
sub=sub,
roles=[r.strip() for r in roles.split(",")],
token_type=token_type,
customer_id=customer_id,
)
def analyst_chat(tok: str, question: str) -> dict:
analyst_api._agent = AnalystAgent(llm=FakeLLM(), repo=FakeRepo())
r = client.post(
"/api/analyst/chat",
headers={"Authorization": f"Bearer {tok}"},
json={"question": question},
)
try:
body = r.json()
except Exception:
body = {"raw": r.text[:300]}
return {"status_code": r.status_code, "body": body}
def main() -> int:
parser = argparse.ArgumentParser()
parser.add_argument("--live-llm", action="store_true", help="使用真实 DeepSeek(需有效 Key)")
args = parser.parse_args()
if not args.live_llm:
mode = "FakeLLM + FakeRepo"
else:
analyst_api._agent = None
mode = "Live LLM + MySQL"
print(f"=== 数据分析 Agent Scope B 冒烟 ({mode}) ===\n")
h = client.get("/health")
ok("GET /health", h.status_code == 200)
ok("analyst routes", "/api/analyst/chat" in client.get("/openapi.json").json()["paths"])
ok("chat no token → 401", client.post("/api/analyst/chat", json={"question": "x"}).status_code == 401)
cases = [
("analyst", token("STAFF-20001", "analyst"), "客户总数是多少"),
("customer", token("CUST-9527", "customer", token_type="customer", customer_id="CUST-9527"), "我有多少笔交易"),
("advisor", token("STAFF-10086", "advisor"), "我名下有多少客户"),
("risk_officer", token("STAFF-30001", "risk_officer"), "待处理预警有多少"),
]
for role, tok, q in cases:
res = analyst_chat(tok, q)
body = res["body"]
ok(f"{role} chat → 200", res["status_code"] == 200, f"http={res['status_code']}")
ok(
f"{role} status success/degrade",
body.get("status") in ("success", "degrade"),
f"status={body.get('status')} err={body.get('error_code')}",
)
ok(f"{role} 四件套", bool(body.get("sql")) and "meta" in body and "table" in body)
if role == "customer":
ok(
"customer AI 风险尾注",
"AI 分析有风险" in (body.get("answer") or "") or "AI 分析有风险" in (body.get("disclaimer") or ""),
)
deny = analyst_chat(token("STAFF-10086", "advisor"), "查 CUST-9999 的持仓")
ok(
"advisor 越权 deny",
deny["status_code"] == 200 and deny["body"].get("status") == "deny",
deny["body"].get("error_code", ""),
)
for role, tok in [
("analyst", token("STAFF-20001", "analyst")),
("customer", token("CUST-9527", "customer", token_type="customer", customer_id="CUST-9527")),
]:
d = client.get("/api/analyst/dashboard", headers={"Authorization": f"Bearer {tok}"})
ok(f"{role} dashboard", d.status_code == 200 and "cards" in d.json())
print(f"\n=== 结果: {PASS} passed, {FAIL} failed ===")
if not args.live_llm:
print("提示: 真实 NL2SQL 请加 --live-llm(需有效 DEEPSEEK_API_KEY + MySQL)")
print("提示: 前端 dev 代理仍指向 :8000,请重启 uvicorn 加载 analyst 路由")
return 1 if FAIL else 0
if __name__ == "__main__":
raise SystemExit(main())