Files
group_xinghuo_jinrong/scripts/dev/sandbox_domain_test.py
T
zhanghongyu_0626 0fb7d34d7a feat(analyst): Implement audit logging for query denial and clarification
- Enhanced the `AnalystAgent` class to include an `_audit_terminal` method for logging query denials, clarifications, and errors, ensuring compliance and traceability.
- Updated error handling paths to call the new audit method, capturing relevant details such as question, user authentication, and SQL context.
- Introduced new validation checks in `sql_guard.py` to enforce ownership filters for sensitive queries, improving security measures.
- Added unit tests to verify the correct logging behavior and ownership filter enforcement, ensuring robust functionality.

This update significantly strengthens the auditing capabilities of the analyst agent, enhancing security and compliance in query handling.
2026-09-11 14:45:16 +08:00

174 lines
7.4 KiB
Python
Raw 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.
"""沙盘:问数线(/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())