Files
group_xinghuo_jinrong/scripts/dev/sandbox_domain_test.py
T

174 lines
7.4 KiB
Python
Raw Normal View History

"""沙盘:问数线(/api/analyst/chat)「表域」授权 + LLM 渲染验证。
真实 DeepSeek + 真实 MySQL(TestClient 进程内),跑 7 角色 × 域内/域外问题,
结构化断言返回,并附带确定性的 sql_guard 越权探针(不依赖 LLM)。
用法:
python scripts/dev/sandbox_domain_test.py
"""
from __future__ import annotations
import sys
from pathlib import Path
ROOT = Path(__file__).resolve().parents[2]
sys.path.insert(0, str(ROOT))
from fastapi.testclient import TestClient # noqa: E402
from app.api import analyst as analyst_api # noqa: E402
from app.main import app # noqa: E402
from app.service.auth_service import issue_dev_token # noqa: E402
from app.service.sql_guard import validate # noqa: E402
from app.service.analytics_repo import AnalyticsRepo # noqa: E402
client = TestClient(app)
analyst_api._agent = None # 由 get_agent() 构建真实 AnalystAgent(DeepSeek + MySQL)
PASS = WARN = FAIL = 0
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(",") if r.strip()],
token_type=token_type,
customer_id=customer_id,
)
def ask(tok: str | None, question: str) -> tuple[int, dict]:
headers = {"Authorization": f"Bearer {tok}"} if tok else {}
try:
r = client.post(
"/api/analyst/chat",
headers=headers,
json={"question": question, "interpret": True},
)
try:
body = r.json()
except Exception:
body = {"raw": r.text[:200]}
return r.status_code, body
except Exception as exc: # noqa: BLE001
return -1, {"_exception": repr(exc)}
def _fmt(v, n=90):
s = str(v).replace("\n", " ")
return s if len(s) <= n else s[: n - 1] + "…"
def judge_in(label: str, code: int, body: dict, need_customer: str | None = None) -> None:
"""域内:应 success/degrade 且 sql/answer 非空;被挡(SCOPE)记 WARN。"""
global PASS, WARN, FAIL
st = body.get("status")
sql = body.get("sql") or ""
ans = body.get("answer") or ""
detail = f"http={code} status={st} err={body.get('error_code')} rows={body.get('meta', {}).get('row_count')}"
if st in ("success", "degrade") and sql and ans:
verdict = "PASS"; PASS += 1
elif st == "deny" and body.get("error_code") == "AUTH_403_SCOPE" and "customer_id" not in sql.lower():
verdict = "WARN"; WARN += 1
detail += " [LLM 未注入 customer_id 过滤 → 被挡]"
elif st == "deny":
verdict = "FAIL"; FAIL += 1
else:
verdict = "FAIL"; FAIL += 1
print(f" [{verdict}] {label}")
print(f" {detail}")
print(f" sql = {_fmt(sql)}")
print(f" ans = {_fmt(ans)}")
def judge_deny(label: str, code: int, body: dict, expected: str) -> None:
"""域外:应 deny;error_code 与预期一致记 PASS,被挡但码不同记 WARN,未挡记 FAIL。"""
global PASS, WARN, FAIL
st = body.get("status")
ec = body.get("error_code")
sql = body.get("sql") or ""
detail = f"http={code} status={st} err={ec}"
if st == "deny" and ec == expected:
verdict = "PASS"; PASS += 1
elif st == "deny":
verdict = "WARN"; WARN += 1
detail += f" (预期 {expected})"
else:
verdict = "FAIL"; FAIL += 1
detail += f" (预期 deny/{expected})"
print(f" [{verdict}] {label}")
print(f" {detail}")
if sql:
print(f" sql = {_fmt(sql)}")
def probe(label: str, code: int, body: dict) -> None:
"""越权探针:仅记录,不判定。"""
st = body.get("status")
sql = body.get("sql") or ""
rows = (body.get("table") or {}).get("rows") or []
print(f" [PROBE] {label}")
print(f" http={code} status={st} err={body.get('error_code')} row_count={len(rows)}")
print(f" sql = {_fmt(sql, 160)}")
if rows:
print(f" rows(前2) = {rows[:2]}")
def main() -> int:
print("=== 问数线表域授权 + LLM 渲染 沙盘(真实 DeepSeek + 真实 MySQL)===\n")
# 1) 无 token
global PASS, FAIL
c, b = ask(None, "客户总数是多少")
if c == 401:
PASS += 1
print(f" [PASS] 无 token → 401 (http={c})")
else:
FAIL += 1
print(f" [FAIL] 无 token → 401 (http={c}, body={b})")
# 2) 域内(应命中并 LLM 渲染)
print("\n— 域内(应 success/degrade)—")
judge_in("customer 我的持仓有哪些", *ask(token("CUST-9527", "customer", token_type="customer", customer_id="CUST-9527"), "我的持仓有哪些?"))
judge_in("advisor 我名下客户的持仓总市值", *ask(token("STAFF-10086", "advisor"), "我名下客户的持仓总市值是多少?"))
judge_in("analyst 客户总数", *ask(token("STAFF-20001", "analyst"), "客户总数是多少?"))
judge_in("risk 待处理预警数量", *ask(token("STAFF-30001", "risk_officer"), "待处理预警有多少?"))
judge_in("ops 近30天申购金额", *ask(token("STAFF-50001", "ops"), "近30天申购金额是多少?"))
# 3) 域外(应 deny)
print("\n— 域外(应 deny)—")
judge_deny("customer 查 CUST-1001 持仓", *ask(token("CUST-9527", "customer", token_type="customer", customer_id="CUST-9527"), "查 CUST-1001 的持仓"), "AUTH_403_NOT_OWNER")
judge_deny("advisor 查 CUST-1010 持仓", *ask(token("STAFF-10086", "advisor"), "查 CUST-1010 的持仓"), "AUTH_403_NOT_ASSIGNED")
judge_deny("ops 查 CUST-9527 持仓明细", *ask(token("STAFF-50001", "ops"), "查 CUST-9527 的持仓明细"), "AUTH_403_SCOPE")
judge_deny("risk_manager 待处理预警", *ask(token("STAFF-31001", "risk_manager"), "待处理预警有多少?"), "AUTH_403_ROLE")
judge_deny("compliance 查所有客户", *ask(token("STAFF-40001", "compliance"), "查一下所有客户"), "AUTH_403_ROLE")
# 4) 越权探针(重点)
print("\n— 越权探针(重点,仅记录)—")
probe("advisor 列出所有客户预警台账(risk_alert 无客户过滤)", *ask(token("STAFF-10086", "advisor"), "列出所有客户的预警台账"))
probe("advisor 查 CUST-1004 客户画像 L2(profile 表+外客户)", *ask(token("STAFF-10086", "advisor"), "查 CUST-1004 的客户画像 L2"))
# 5) 确定性 sql_guard 探针(不依赖 LLM,直接证明代码层缺口)
print("\n— 确定性 sql_guard 探针(不依赖 LLM)—")
repo = AnalyticsRepo()
scope = repo.resolve_advisor_scope("STAFF-10086")
print(f" advisor(STAFF-10086) 名下 scope = {scope}")
for label, sql, domain in [
("risk_alert 全量(assigned)", "SELECT * FROM jinrong_agent.risk_alert", "assigned"),
("customer_profile_l3 全量(assigned)", "SELECT * FROM jinrong_agent.customer_profile_l3", "assigned"),
("core_holding 全量无过滤(assigned)", "SELECT * FROM core_holding", "assigned"),
("risk_alert 全量(aggregate/ops)", "SELECT * FROM jinrong_agent.risk_alert", "aggregate"),
]:
try:
res = validate(sql, domain, scope)
print(f" [{'ALLOWED' if res.allowed else 'DENIED'}] {label} -> allowed={res.allowed}")
except Exception as exc:
print(f" [DENIED] {label} -> {exc}")
print(f"\n=== 结果: {PASS} PASS, {WARN} WARN, {FAIL} FAIL ===")
return 1 if FAIL else 0
if __name__ == "__main__":
raise SystemExit(main())