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.
This commit is contained in:
2026-09-09 21:09:03 +08:00
parent bcd6175d4e
commit af9fb719e5
3 changed files with 183 additions and 3 deletions
+180
View File
@@ -0,0 +1,180 @@
"""数据分析 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())