from __future__ import annotations import argparse import hashlib import json import os import sys import time from concurrent.futures import ThreadPoolExecutor from dataclasses import dataclass from pathlib import Path from typing import Any, Iterable ROOT = Path(__file__).resolve().parents[2] if str(ROOT) not in sys.path: sys.path.insert(0, str(ROOT)) from app.api.analyst_auth_adapter import AnalystAuthContext # noqa: E402 from app.service import agent_service, customer_prompts, input_guard, rag_service, sql_guard, tool_service, visitor_service # noqa: E402 from app.service.analyst_agent import AnalystAgent # noqa: E402 from app.service.template_service import QueryTemplate, TemplateService # noqa: E402 from scripts.eval.sandbox import ( # noqa: E402 ResourceLedger, SandboxConfig, SandboxSafetyError, build_config, redact_input, write_json, write_jsonl, ) DEFAULT_CASE_DIR = ROOT / "scripts" / "eval" / "cases" DEFAULT_FILES = ( "common.jsonl", "adversarial.jsonl", "analyst.jsonl", "rag_access.jsonl", "visitor.jsonl", "visitor_fallback.jsonl", "customer.jsonl", "advisor.jsonl", "risk.jsonl", ) @dataclass class FakeLLM: sql: str = "SELECT 1" answers: list[str] | None = None calls: int = 0 def __post_init__(self) -> None: self.answers = list(self.answers or []) def complete(self, messages, temperature=0, max_tokens=2048): self.calls += 1 if self.calls == 1: return self.sql, {"prompt_tokens": 1, "completion_tokens": 1} answer = self.answers.pop(0) if self.answers else "无解读" return answer, {"prompt_tokens": 1, "completion_tokens": 1} class FakeRepo: def __init__( self, *, rows: list[list[Any]] | None = None, columns: list[str] | None = None, scope: list[str] | None = None, ) -> None: self.rows = list(rows or []) self.columns = list(columns or []) self.scope = list(scope or []) self.logged: list[dict[str, Any]] = [] def resolve_advisor_scope(self, subject_id: str) -> list[str]: return self.scope def execute_readonly(self, sql: str) -> dict[str, Any]: return {"columns": self.columns, "rows": self.rows} def get_data_as_of(self) -> str: return "2026-09-12" def log_query(self, **kwargs: Any) -> None: self.logged.append(kwargs) def log_audit(self, **kwargs: Any) -> None: self.logged.append({"audit": kwargs}) def list_published_templates(self) -> list[Any]: return [] def get_query_log_by_trace(self, trace_id: str) -> dict[str, Any] | None: return { "staff_id": "STAFF-A", "generated_sql": "SELECT COUNT(*) AS c FROM core_customer", "exec_status": "success", } class NoopCache: def permission_fingerprint(self, subject_id: str, domain: str, scope: list[str]) -> str: return f"{subject_id}:{domain}:{','.join(scope)}" def sql_hash(self, sql: str) -> str: return hashlib.sha256(sql.strip().encode("utf-8")).hexdigest()[:16] def get_result(self, *args: Any, **kwargs: Any) -> None: return None def set_result(self, *args: Any, **kwargs: Any) -> None: return None def load_cases(paths: Iterable[Path]) -> list[dict[str, Any]]: cases: list[dict[str, Any]] = [] for path in paths: with path.open(encoding="utf-8") as handle: for line_number, line in enumerate(handle, 1): if not line.strip() or line.lstrip().startswith("#"): continue try: case = json.loads(line) except json.JSONDecodeError as exc: raise ValueError(f"invalid case JSON: {path}:{line_number}: {exc}") from exc case.setdefault("source_file", str(path)) case.setdefault("case_id", f"{path.stem}-{line_number}") cases.append(case) return cases def _match(expected: Any, actual: Any) -> bool: if isinstance(expected, list): return actual in expected return expected == actual def _base_result(config: SandboxConfig, case: dict[str, Any]) -> dict[str, Any]: return { "run_id": config.run_id, "case_id": case["case_id"], "agent": case.get("agent"), "layer": case.get("layer", "l0"), "kind": case.get("kind"), "input": redact_input(str(case.get("message") or case.get("question") or case.get("sql") or "")), "status": "FAIL", "verdict": "FAIL", "observations": [], "error": None, } def evaluate_guard(config: SandboxConfig, case: dict[str, Any]) -> dict[str, Any]: result = _base_result(config, case) message = str(case.get("message", "")) verdict = input_guard.inspect_message(message, max_length=int(case.get("max_length", input_guard.MESSAGE_MAX_LENGTH))) expected = case.get("expected", {}) actual = {"blocked": verdict.blocked, "guard_type": verdict.guard_type, "reason": verdict.reason} result.update({"actual": actual, "expected": expected}) mismatches = [ field for field in ("blocked", "guard_type") if field in expected and not _match(expected[field], actual[field]) ] if mismatches: result["error"] = f"guard mismatch: {', '.join(mismatches)}" return result result["status"] = "PASS" result["verdict"] = "PASS" if expected.get("flags"): result["observations"].extend(expected["flags"]) return result def evaluate_intent(config: SandboxConfig, case: dict[str, Any]) -> dict[str, Any]: result = _base_result(config, case) actual = tool_service.match_intent(case.get("agent", ""), str(case.get("message", ""))) expected = case.get("expected", {}) result.update({"actual": {"tool": actual}, "expected": expected}) if "tool" in expected and not _match(expected["tool"], actual): result["error"] = f"expected tool {expected['tool']!r}, got {actual!r}" return result result["status"] = result["verdict"] = "PASS" result["observations"].extend(expected.get("flags", [])) return result def evaluate_tool_param(config: SandboxConfig, case: dict[str, Any]) -> dict[str, Any]: result = _base_result(config, case) spec = tool_service.get_registered_tool(case["tool_name"]) ok, error, cleaned = tool_service._normalize_params(spec, case.get("params")) actual = {"ok": ok, "error": error, "cleaned": cleaned} expected = case.get("expected", {}) result.update({"actual": actual, "expected": expected}) for field in ("ok", "error", "cleaned"): if field in expected and not _match(expected[field], actual[field]): result["error"] = f"param {field} mismatch: expected {expected[field]!r}, got {actual[field]!r}" return result result["status"] = result["verdict"] = "PASS" return result def evaluate_visitor_route(config: SandboxConfig, case: dict[str, Any]) -> dict[str, Any]: result = _base_result(config, case) routed = visitor_service.intent_classify({"message": str(case.get("message", ""))}) actual = { "intent": routed.get("intent"), "transfer_to_human": routed.get("intent") == "transfer_human", "reply": routed.get("reply", ""), } expected = case.get("expected", {}) result.update({"actual": actual, "expected": expected}) for field in ("intent", "transfer_to_human"): if field in expected and not _match(expected[field], actual[field]): result["error"] = f"visitor {field} mismatch: expected {expected[field]!r}, got {actual[field]!r}" return result if "reply_contains" in expected and expected["reply_contains"] not in actual["reply"]: result["error"] = f"visitor reply missing {expected['reply_contains']!r}" return result result["status"] = result["verdict"] = "PASS" return result def evaluate_rag_fallback(config: SandboxConfig, case: dict[str, Any]) -> dict[str, Any]: result = _base_result(config, case) context = str(case.get("rag_context", "")) actual = {"reply": visitor_service._degraded_reply_from_rag(context)} expected = case.get("expected", {}) result.update({"actual": actual, "expected": expected}) if "present" in expected and (actual["reply"] is not None) != bool(expected["present"]): result["error"] = ( f"RAG fallback presence mismatch: expected {expected['present']!r}, " f"got {actual['reply'] is not None!r}" ) return result if "contains" in expected: reply = actual["reply"] or "" if expected["contains"] not in reply: result["error"] = f"RAG fallback missing {expected['contains']!r}" return result result["status"] = result["verdict"] = "PASS" return result def evaluate_customer_route(config: SandboxConfig, case: dict[str, Any]) -> dict[str, Any]: result = _base_result(config, case) routed = customer_prompts.keyword_route(str(case.get("message", ""))) actual = { "intent": routed[0] if routed else None, "preset": routed[1] if routed else None, } expected = case.get("expected", {}) result.update({"actual": actual, "expected": expected}) for field in ("intent", "preset"): if field in expected and not _match(expected[field], actual[field]): result["error"] = f"customer {field} mismatch: expected {expected[field]!r}, got {actual[field]!r}" return result result["status"] = result["verdict"] = "PASS" return result def evaluate_disclaimer(config: SandboxConfig, case: dict[str, Any]) -> dict[str, Any]: result = _base_result(config, case) agent = str(case.get("agent", "")) actual = {"needs_disclaimer": agent_service.needs_disclaimer(agent)} expected = case.get("expected", {}) result.update({"actual": actual, "expected": expected}) if "needs_disclaimer" in expected and not _match( expected["needs_disclaimer"], actual["needs_disclaimer"] ): result["error"] = ( f"disclaimer mismatch: expected {expected['needs_disclaimer']!r}, " f"got {actual['needs_disclaimer']!r}" ) return result result["status"] = result["verdict"] = "PASS" return result def evaluate_sql(config: SandboxConfig, case: dict[str, Any]) -> dict[str, Any]: result = _base_result(config, case) expected = case.get("expected", {}) try: validation = sql_guard.validate(case["sql"], case.get("domain", "full"), case.get("scope")) actual = {"allowed": validation.allowed, "error_code": None, "tables": validation.tables} except sql_guard.SqlGuardError as exc: actual = {"allowed": False, "error_code": exc.error_code, "tables": []} result.update({"actual": actual, "expected": expected}) for field in ("allowed", "error_code"): if field in expected and not _match(expected[field], actual[field]): result["error"] = f"SQL {field} mismatch: expected {expected[field]!r}, got {actual[field]!r}" return result result["status"] = result["verdict"] = "PASS" result["observations"].extend(expected.get("flags", [])) if "sensitive_column_observation" in expected.get("flags", []): result["observations"].append("SENSITIVE_COLUMNS_NOT_ENFORCED_BY_VALIDATE") return result def _auth_from_case(case: dict[str, Any]) -> AnalystAuthContext: data = dict(case.get("auth") or {}) return AnalystAuthContext( subject_id=data.get("subject_id", "STAFF-A"), token_type=data.get("token_type", "staff"), roles=list(data.get("roles") or ["analyst"]), customer_id=data.get("customer_id"), ) def evaluate_analyst(config: SandboxConfig, case: dict[str, Any]) -> dict[str, Any]: result = _base_result(config, case) fake_llm = FakeLLM(sql=case.get("fake_sql", "SELECT 1"), answers=case.get("fake_answers")) repo = FakeRepo( rows=case.get("fake_rows"), columns=case.get("fake_columns"), scope=case.get("scope"), ) templates: TemplateService | None = None if case.get("template_sql"): templates = TemplateService( templates=[ QueryTemplate( template_key=case.get("template_key", "eval_template"), template_sql=case["template_sql"], params_schema={"match_all": ["客户", "总数"]}, ) ] ) agent = AnalystAgent( llm=fake_llm, repo=repo, cache=NoopCache(), templates=templates, ) response = agent.run( case["question"], _auth_from_case(case), session_id=f"{config.run_id}-{case['case_id']}", trace_id=f"{config.run_id}-{case['case_id']}", interpret=bool(case.get("interpret", False)), ) actual = { "status": response.status, "error_code": response.error_code, "row_count": len(response.table.rows) if response.table else 0, "template_hit": bool(response.meta and response.meta.template_hit), "llm_calls": fake_llm.calls, "answer": response.answer, "sql": response.sql, } expected = case.get("expected", {}) result.update({"actual": actual, "expected": expected, "audit_events": len(repo.logged)}) for field in ("status", "error_code", "row_count", "template_hit", "llm_calls"): if field in expected and not _match(expected[field], actual[field]): result["error"] = f"analyst {field} mismatch: expected {expected[field]!r}, got {actual[field]!r}" return result if "contains" in expected and expected["contains"] not in actual["answer"]: result["error"] = f"analyst answer missing {expected['contains']!r}" return result result["status"] = result["verdict"] = "PASS" return result def evaluate_rag_contract(config: SandboxConfig, case: dict[str, Any]) -> dict[str, Any]: result = _base_result(config, case) rows = list(case.get("results") or []) refs = rag_service._source_refs(rows) expected = case.get("expected", {}) actual = { "result_count": len(rows), "source_ref_count": len(refs), "source_refs": refs, "required_fields_present": all( all(ref.get(field) for field in ("source_doc_id", "source_version", "product_id", "product_name")) for ref in refs ), } result.update({"actual": actual, "expected": expected}) for field in ("result_count", "source_ref_count", "required_fields_present"): if field in expected and not _match(expected[field], actual[field]): result["error"] = f"RAG {field} mismatch: expected {expected[field]!r}, got {actual[field]!r}" return result result["status"] = result["verdict"] = "PASS" return result def evaluate_tool_access(config: SandboxConfig, case: dict[str, Any]) -> dict[str, Any]: result = _base_result(config, case) expected = case.get("expected", {}) class CoreStub: def is_advisor_assigned(self, actor_id: str, customer_id: str) -> bool: return bool(case.get("assigned", False)) actual = {"allowed": True, "error_code": None} try: tool_service.assert_tool_access( dict(case.get("actor") or {}), str(case.get("customer_id") or ""), CoreStub(), ) except Exception as exc: # PermissionDenied is intentionally kept behind the stable code contract. actual = {"allowed": False, "error_code": getattr(exc, "code", type(exc).__name__)} result.update({"actual": actual, "expected": expected}) for field in ("allowed", "error_code"): if field in expected and not _match(expected[field], actual[field]): result["error"] = f"tool access {field} mismatch: expected {expected[field]!r}, got {actual[field]!r}" return result result["status"] = result["verdict"] = "PASS" return result def evaluate_case(config: SandboxConfig, case: dict[str, Any]) -> dict[str, Any]: started = time.perf_counter() kind = case.get("kind") try: if kind == "guard": result = evaluate_guard(config, case) elif kind == "intent": result = evaluate_intent(config, case) elif kind == "tool_param": result = evaluate_tool_param(config, case) elif kind == "visitor_route": result = evaluate_visitor_route(config, case) elif kind == "customer_route": result = evaluate_customer_route(config, case) elif kind == "rag_fallback": result = evaluate_rag_fallback(config, case) elif kind == "disclaimer": result = evaluate_disclaimer(config, case) elif kind == "sql": result = evaluate_sql(config, case) elif kind == "analyst": result = evaluate_analyst(config, case) elif kind == "rag_contract": result = evaluate_rag_contract(config, case) elif kind == "tool_access": result = evaluate_tool_access(config, case) else: result = _base_result(config, case) result["status"] = result["verdict"] = "GAP" result["error"] = f"unsupported case kind: {kind}" except Exception as exc: # noqa: BLE001 result = _base_result(config, case) result["error"] = repr(exc) result["latency_ms"] = int((time.perf_counter() - started) * 1000) return result def run_cases( config: SandboxConfig, cases: list[dict[str, Any]], *, concurrency: int = 1, repeat: int = 1, ) -> tuple[list[dict[str, Any]], dict[str, Any]]: expanded = cases * max(1, repeat) started = time.perf_counter() if concurrency <= 1: results = [evaluate_case(config, case) for case in expanded] else: with ThreadPoolExecutor(max_workers=concurrency) as pool: results = list(pool.map(lambda item: evaluate_case(config, item), expanded)) elapsed_ms = int((time.perf_counter() - started) * 1000) latencies = [int(item.get("latency_ms", 0)) for item in results if item.get("latency_ms") is not None] passed = sum(item["verdict"] == "PASS" for item in results) failed = sum(item["verdict"] == "FAIL" for item in results) gaps = sum(item["verdict"] == "GAP" for item in results) summary = { "total": len(results), "passed": passed, "failed": failed, "gaps": gaps, "elapsed_ms": elapsed_ms, "throughput_per_second": round(len(results) / (elapsed_ms / 1000), 3) if elapsed_ms else None, "concurrency": concurrency, "repeat": repeat, "pressure": { "p50_ms": _percentile(latencies, 0.50), "p95_ms": _percentile(latencies, 0.95), "p99_ms": _percentile(latencies, 0.99), }, } return results, summary def _percentile(values: list[int], ratio: float) -> int | None: if not values: return None ordered = sorted(values) index = min(len(ordered) - 1, max(0, int((len(ordered) - 1) * ratio))) return ordered[index] def filter_cases( cases: list[dict[str, Any]], *, layer: str, agent: str | None = None, max_requests: int | None = None, ) -> list[dict[str, Any]]: selected = [case for case in cases if case.get("layer", "l0") == layer] if agent: selected = [case for case in selected if case.get("agent") == agent] if max_requests is not None and max_requests >= 0: selected = selected[:max_requests] return selected def _render_report_markdown( *, run_id: str, layer: str, summary: dict[str, Any], coverage_gaps: list[str], cleanup: dict[str, Any], ) -> str: return "\n".join( [ f"# Agent knowledge evaluation: `{run_id}`", "", f"- Layer: `{layer}`", f"- Total: `{summary['total']}`", f"- Passed: `{summary['passed']}`", f"- Failed: `{summary['failed']}`", f"- Coverage gaps: `{summary['gaps']}`", f"- Elapsed: `{summary['elapsed_ms']} ms`", f"- Throughput: `{summary['throughput_per_second']}` cases/s", f"- Cleanup: `{cleanup['status']}`", "", "## Pressure", "", f"- p50: `{summary['pressure']['p50_ms']} ms`", f"- p95: `{summary['pressure']['p95_ms']} ms`", f"- p99: `{summary['pressure']['p99_ms']} ms`", "", "## Coverage gaps", "", *[f"- {gap}" for gap in coverage_gaps], "", ] ) def _case_paths(case_dir: Path, names: list[str] | None) -> list[Path]: selected = names or list(DEFAULT_FILES) paths = [] for name in selected: candidate = Path(name) path = candidate if candidate.is_absolute() else case_dir / candidate path = path.resolve() if not path.is_file(): raise FileNotFoundError(f"missing case file: {path}") paths.append(path) return paths def parse_args(argv: list[str] | None = None) -> argparse.Namespace: parser = argparse.ArgumentParser(description="JinRong offline/isolated agent knowledge evaluation") parser.add_argument("--layer", choices=("l0", "l1", "l2", "l3", "l4"), default="l0") parser.add_argument("--case-dir", type=Path, default=DEFAULT_CASE_DIR) parser.add_argument("--case-file", action="append", dest="case_files") parser.add_argument("--agent") parser.add_argument("--max-requests", type=int) parser.add_argument("--timeout", type=float, default=30.0) parser.add_argument("--live-llm", action="store_true") parser.add_argument("--concurrency", type=int, default=1) parser.add_argument("--repeat", type=int, default=1) parser.add_argument("--run-id") parser.add_argument("--sandbox", action="store_true") parser.add_argument("--live", action="store_true") parser.add_argument("--keep", action="store_true") parser.add_argument("--mysql-database") parser.add_argument("--mysql-core-database") parser.add_argument("--redis-url") parser.add_argument("--milvus-uri") return parser.parse_args(argv) def main(argv: list[str] | None = None) -> int: args = parse_args(argv) resources = { "mysql_database": args.mysql_database or os.getenv("EVAL_MYSQL_DATABASE", ""), "mysql_core_database": args.mysql_core_database or os.getenv("EVAL_MYSQL_CORE_DATABASE", ""), "redis_url": args.redis_url or os.getenv("EVAL_REDIS_URL", ""), "milvus_uri": args.milvus_uri or os.getenv("EVAL_MILVUS_URI", ""), } try: config = build_config( ROOT, run_id=args.run_id, live=args.live, sandbox=args.sandbox, keep=args.keep, resources=resources, ) if args.layer in {"l2", "l3", "l4"} and not args.live: raise SandboxSafetyError(f"{args.layer} requires explicit --live --sandbox and JINRONG_EVAL=1") if args.live: raise SandboxSafetyError( "live adapters are not implemented; L0/L1 runner is offline-only" ) if args.live_llm: raise SandboxSafetyError( "--live-llm is unavailable until a live adapter is implemented" ) if args.layer not in {"l0", "l1"}: raise SandboxSafetyError( f"{args.layer} runner is not implemented yet; use L0/L1 offline or add a dedicated live adapter" ) paths = _case_paths(args.case_dir.resolve(), args.case_files) cases = filter_cases( load_cases(paths), layer=args.layer, agent=args.agent, max_requests=args.max_requests, ) if not cases: raise ValueError(f"no cases selected for layer={args.layer!r}, agent={args.agent!r}") ledger = ResourceLedger(config) ledger.save() results, summary = run_cases( config, cases, concurrency=max(1, args.concurrency), repeat=max(1, args.repeat), ) coverage_gaps = [ "advisor_independent_l2_and_draft_workflow_not_connected", "advisor_business_kb_not_connected", "knowledge_http_api_is_not_a_runtime_ingestion_boundary", "customer_visitor_source_refs_not_guaranteed_in_public_response", ] manifest = { "schema_version": "agent-kb-eval/v1", "run_id": config.run_id, "layer": args.layer, "agent": args.agent, "case_files": [ str(path.relative_to(ROOT)) if path.is_relative_to(ROOT) else str(path) for path in paths ], "selected_cases": len(cases), "concurrency": max(1, args.concurrency), "repeat": max(1, args.repeat), "timeout_seconds": args.timeout, "live_llm": args.live_llm, "environment": config.safe_environment(), } report = { "schema_version": "agent-kb-eval/v1", "run_id": config.run_id, "layer": args.layer, "environment": config.safe_environment(), "cases": results, "summary": summary, "coverage_gaps": coverage_gaps, "cleanup": {"status": "not_required", "remaining": []}, } failures = [item for item in results if item.get("verdict") == "FAIL"] write_json(config, "manifest.json", manifest) write_jsonl(config, "cases.jsonl", results) write_jsonl(config, "failures.jsonl", failures) write_json(config, "report.json", report) write_json(config, "summary.json", summary) cleanup = ledger.cleanup() report["cleanup"] = cleanup write_json(config, "cleanup.json", cleanup) write_json(config, "report.json", report) (config.artifact_dir / "report.md").write_text( _render_report_markdown( run_id=config.run_id, layer=args.layer, summary=summary, coverage_gaps=coverage_gaps, cleanup=cleanup, ), encoding="utf-8", ) print(json.dumps({"run_id": config.run_id, "report": str(config.artifact_dir / 'report.json'), **summary, "cleanup": cleanup["status"]}, ensure_ascii=False)) return 1 if summary["failed"] or cleanup["status"] == "failed" else 0 except (SandboxSafetyError, FileNotFoundError, ValueError) as exc: print(json.dumps({"status": "blocked_preflight", "error": str(exc)}, ensure_ascii=False), file=sys.stderr) return 2 if __name__ == "__main__": raise SystemExit(main())