runtime_governance.py 19 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379
  1. """Default-deny, hash-only runtime governance for model and Agent invocations."""
  2. from __future__ import annotations
  3. import hashlib
  4. import json
  5. import re
  6. import threading
  7. import unicodedata
  8. from dataclasses import dataclass
  9. from pathlib import Path
  10. from typing import Any
  11. from uuid import uuid4
  12. from app.core.mcp.governed_invocation import (
  13. InvocationContractError,
  14. normalize_mcp_invocation,
  15. )
  16. class RuntimeDecisionError(ValueError):
  17. def __init__(self, code: str):
  18. self.code = code
  19. super().__init__(code)
  20. class DatasetLeakageError(ValueError):
  21. pass
  22. def load_engineering_fixture(manifest_path: str | Path) -> dict[str, Any]:
  23. """Read each fixture once and verify its declared SHA-256 before evaluation."""
  24. manifest_file = Path(manifest_path)
  25. manifest = json.loads(manifest_file.read_text(encoding="utf-8"))
  26. if manifest.get("fixture_only") is not True or manifest.get("schema_version") != "1":
  27. raise DatasetLeakageError("fixture_manifest_invalid")
  28. files = manifest.get("files")
  29. if not isinstance(files, dict) or set(files) != {"dataset.json", "temporal-dataset.json"}:
  30. raise DatasetLeakageError("fixture_manifest_invalid")
  31. loaded: dict[str, Any] = {}
  32. for name, expected in files.items():
  33. if not isinstance(expected, str) or not expected.startswith("sha256:"):
  34. raise DatasetLeakageError("fixture_manifest_invalid")
  35. content = (manifest_file.parent / name).read_bytes()
  36. if hashlib.sha256(content).hexdigest() != expected.removeprefix("sha256:"):
  37. raise DatasetLeakageError("fixture_digest_mismatch")
  38. loaded[name] = json.loads(content)
  39. train = set(loaded["dataset.json"].get("train_case_ids", []))
  40. evaluation = {item.get("case_id") for item in loaded["dataset.json"].get("evaluation_cases", [])}
  41. if train & evaluation:
  42. raise DatasetLeakageError("dataset_leakage_detected")
  43. return loaded
  44. @dataclass(frozen=True)
  45. class ModelRoute:
  46. route_id: str
  47. provider: str
  48. model: str
  49. prompt_version: str
  50. generation: str
  51. tenant_id: str
  52. principal_id: str
  53. business_domain_uid: str
  54. environment: str
  55. canary_status: str
  56. @dataclass(frozen=True)
  57. class BudgetPolicy:
  58. token_limit: int
  59. cost_limit_micros: int
  60. tool_limit: int
  61. time_limit_ms: int
  62. concurrency_limit: int
  63. model_limits: dict[str, int]
  64. _SCOPE_FIELDS = ("tenant_id", "principal_id", "business_domain_uid", "environment")
  65. _REQUEST_KEYS = frozenset(
  66. {
  67. *_SCOPE_FIELDS, "provider", "model", "prompt_version", "generation",
  68. "interface_type", "tool_name", "action", "requested_capabilities", "input_text",
  69. "evidence_refs", "idempotency_key", "estimated_tokens", "estimated_cost_micros",
  70. "requested_time_ms",
  71. }
  72. )
  73. _FORBIDDEN_KEYS = frozenset(
  74. {
  75. "url", "uri", "endpoint", "path", "file_path", "filename", "command", "shell",
  76. "script", "secret", "token", "password", "api_key", "raw_rows", "rows", "records",
  77. "payload", "arguments", "output", "raw_input", "raw_output",
  78. }
  79. )
  80. _SAFE_ID = re.compile(r"^[A-Za-z0-9_.:-]{1,160}$")
  81. _INJECTION = re.compile(
  82. r"(?is)(ignore\s+(?:all\s+)?(?:previous|prior|system)\s+instructions?|"
  83. r"bypass\s+(?:the\s+)?(?:policy|guard|permission)|reveal.{0,50}(?:secret|token|password|api[_ -]?key)|"
  84. r"(?:忽略).{0,16}(?:指令|规则|系统)|(?:绕过|关闭|禁用).{0,16}(?:权限|策略|审计))"
  85. )
  86. _LOCATION = re.compile(r"(?i)(?:https?://|file://|(?:^|\s)/[A-Za-z0-9_./-]+)")
  87. def _hash(value: str) -> str:
  88. return hashlib.sha256(value.encode("utf-8")).hexdigest()
  89. def _positive_int(value: Any, label: str, maximum: int) -> int:
  90. if isinstance(value, bool) or not isinstance(value, int) or not 0 < value <= maximum:
  91. raise RuntimeDecisionError(f"{label}_invalid")
  92. return value
  93. def _closed_request(value: Any) -> dict[str, Any]:
  94. if not isinstance(value, dict):
  95. raise RuntimeDecisionError("request_invalid")
  96. if len(value) > len(_REQUEST_KEYS):
  97. raise RuntimeDecisionError("request_invalid")
  98. for key in value:
  99. normalized = unicodedata.normalize("NFKC", str(key))
  100. if normalized != key or normalized.casefold() in _FORBIDDEN_KEYS:
  101. raise RuntimeDecisionError("forbidden_contract_key")
  102. unknown = set(value) - _REQUEST_KEYS
  103. if unknown or set(value) != _REQUEST_KEYS:
  104. raise RuntimeDecisionError("request_schema_invalid")
  105. return dict(value)
  106. def _normalize_request(value: Any) -> dict[str, Any]:
  107. request = _closed_request(value)
  108. for key in _SCOPE_FIELDS + ("provider", "model", "prompt_version", "generation", "idempotency_key"):
  109. item = request[key]
  110. if not isinstance(item, str) or not _SAFE_ID.fullmatch(item):
  111. raise RuntimeDecisionError(f"{key}_invalid")
  112. text = request["input_text"]
  113. if not isinstance(text, str) or not text.strip() or len(text) > 8000:
  114. raise RuntimeDecisionError("input_invalid")
  115. normalized_text = unicodedata.normalize("NFKC", text).strip()
  116. if _INJECTION.search(normalized_text):
  117. raise RuntimeDecisionError("prompt_injection_detected")
  118. if _LOCATION.search(normalized_text):
  119. raise RuntimeDecisionError("input_location_denied")
  120. capabilities = request["requested_capabilities"]
  121. if not isinstance(capabilities, list) or not capabilities or len(capabilities) > 8:
  122. raise RuntimeDecisionError("capability_invalid")
  123. if any(not isinstance(item, str) or item not in {"network", "file", "data", "command", "resource", "time"} for item in capabilities):
  124. raise RuntimeDecisionError("capability_invalid")
  125. evidence = request["evidence_refs"]
  126. if not isinstance(evidence, list) or len(evidence) > 20:
  127. raise RuntimeDecisionError("citation_invalid")
  128. for item in evidence:
  129. if not isinstance(item, dict) or set(item) != {"evidence_id", "digest"}:
  130. raise RuntimeDecisionError("citation_invalid")
  131. if not isinstance(item["evidence_id"], str) or not _SAFE_ID.fullmatch(item["evidence_id"]):
  132. raise RuntimeDecisionError("citation_invalid")
  133. if not isinstance(item["digest"], str) or not re.fullmatch(r"[a-f0-9]{64}", item["digest"]):
  134. raise RuntimeDecisionError("citation_invalid")
  135. try:
  136. normalize_mcp_invocation(
  137. {
  138. "interface_type": request["interface_type"], "tool_name": request["tool_name"],
  139. "action": request["action"], "arguments_digest": _hash(normalized_text),
  140. "evidence_refs": evidence,
  141. }
  142. )
  143. except InvocationContractError as error:
  144. raise RuntimeDecisionError("mcp_contract_invalid") from error
  145. return {
  146. **{key: request[key] for key in _SCOPE_FIELDS},
  147. **{key: request[key] for key in ("provider", "model", "prompt_version", "generation", "interface_type", "tool_name", "action", "idempotency_key")},
  148. "requested_capabilities": frozenset(request["requested_capabilities"]),
  149. "input_hash": _hash(normalized_text),
  150. "evidence_refs": tuple((item["evidence_id"], item["digest"]) for item in evidence),
  151. "estimated_tokens": _positive_int(request["estimated_tokens"], "estimated_tokens", 1_000_000),
  152. "estimated_cost_micros": _positive_int(request["estimated_cost_micros"], "estimated_cost_micros", 1_000_000_000),
  153. "requested_time_ms": _positive_int(request["requested_time_ms"], "requested_time_ms", 300_000),
  154. }
  155. class InMemoryRuntimeRepository:
  156. """Thread-safe reference repository used by unit contracts and local wiring."""
  157. def __init__(self) -> None:
  158. self._lock = threading.RLock()
  159. self._routes: list[ModelRoute] = []
  160. self._capabilities: set[tuple[str, str, str, str, str, str, frozenset[str]]] = set()
  161. self._budgets: dict[str, dict[str, Any]] = {}
  162. self._decisions: dict[tuple[str, str], dict[str, Any]] = {}
  163. self._audit: list[dict[str, Any]] = []
  164. self._states: dict[str, dict[str, Any]] = {}
  165. self._generations: dict[tuple[str, str], dict[str, Any]] = {}
  166. def register_route(self, route: ModelRoute) -> None:
  167. with self._lock:
  168. self._routes.append(route)
  169. self._generations.setdefault((route.tenant_id, route.generation), {"status": route.canary_status, "fence": 0, "parent": None})
  170. self._states.setdefault(route.tenant_id, {"state": "active", "incident_ref": None, "fence": 0})
  171. def register_capability(self, *, tenant_id: str, business_domain_uid: str, environment: str, tool_name: str, interface_type: str, action: str, risk_level: str, allowed_capabilities: set[str]) -> None:
  172. if risk_level != "low" or not allowed_capabilities <= {"data"}:
  173. raise ValueError("only declared low-risk data capability may be registered locally")
  174. with self._lock:
  175. self._capabilities.add((tenant_id, business_domain_uid, environment, tool_name, interface_type, action, frozenset(allowed_capabilities)))
  176. def set_budget(self, tenant_id: str, policy: BudgetPolicy) -> None:
  177. if min(policy.token_limit, policy.cost_limit_micros, policy.tool_limit, policy.time_limit_ms, policy.concurrency_limit) < 1:
  178. raise ValueError("budget limits must be positive")
  179. with self._lock:
  180. self._budgets[tenant_id] = {"policy": policy, "tokens": policy.token_limit, "cost": policy.cost_limit_micros, "tools": policy.tool_limit, "time": policy.time_limit_ms, "concurrency": policy.concurrency_limit, "models": dict(policy.model_limits), "fence": 0}
  181. def find_route(self, request: dict[str, Any]) -> ModelRoute | None:
  182. with self._lock:
  183. for route in self._routes:
  184. if all(getattr(route, field) == request[field] for field in _SCOPE_FIELDS + ("provider", "model", "prompt_version", "generation")):
  185. return route
  186. return None
  187. def generation_status(self, tenant_id: str, generation: str) -> str | None:
  188. with self._lock:
  189. entry = self._generations.get((tenant_id, generation))
  190. return entry["status"] if entry else None
  191. def allowed_capability(self, request: dict[str, Any]) -> bool:
  192. key = (request["tenant_id"], request["business_domain_uid"], request["environment"], request["tool_name"], request["interface_type"], request["action"], request["requested_capabilities"])
  193. with self._lock:
  194. return key in self._capabilities
  195. def existing(self, tenant_id: str, idempotency_key: str) -> dict[str, Any] | None:
  196. with self._lock:
  197. record = self._decisions.get((tenant_id, idempotency_key))
  198. return dict(record) if record else None
  199. def reserve(self, request: dict[str, Any]) -> int:
  200. with self._lock:
  201. budget = self._budgets.get(request["tenant_id"])
  202. if budget is None or any((budget["tokens"] < request["estimated_tokens"], budget["cost"] < request["estimated_cost_micros"], budget["tools"] < 1, budget["time"] < request["requested_time_ms"], budget["concurrency"] < 1, budget["models"].get(request["model"], 0) < 1)):
  203. raise RuntimeDecisionError("budget_exhausted")
  204. budget["tokens"] -= request["estimated_tokens"]
  205. budget["cost"] -= request["estimated_cost_micros"]
  206. budget["tools"] -= 1
  207. budget["time"] -= request["requested_time_ms"]
  208. budget["concurrency"] -= 1
  209. budget["models"][request["model"]] -= 1
  210. budget["fence"] += 1
  211. return budget["fence"]
  212. def write_authorized(self, request: dict[str, Any], route: ModelRoute, fence: int) -> dict[str, Any]:
  213. record = {"invocation_id": str(uuid4()), "decision": "authorized", "route_id": route.route_id, "generation": route.generation, "budget_fence": fence, "input_hash": request["input_hash"], "evidence_digests": tuple(digest for _, digest in request["evidence_refs"]), "tenant_id": request["tenant_id"], "principal_id": request["principal_id"], "business_domain_uid": request["business_domain_uid"], "environment": request["environment"]}
  214. with self._lock:
  215. self._decisions[(request["tenant_id"], request["idempotency_key"])] = record
  216. self._audit.append(dict(record))
  217. return dict(record)
  218. def remaining_budget(self, tenant_id: str) -> dict[str, int]:
  219. with self._lock:
  220. budget = self._budgets[tenant_id]
  221. return {"tokens": budget["tokens"], "cost_micros": budget["cost"], "tools": budget["tools"], "time_ms": budget["time"], "concurrency": budget["concurrency"]}
  222. def audit_records(self) -> list[dict[str, Any]]:
  223. with self._lock:
  224. return [dict(item) for item in self._audit]
  225. def runtime_state(self, tenant_id: str) -> dict[str, Any]:
  226. with self._lock:
  227. return dict(self._states.get(tenant_id, {"state": "active", "incident_ref": None, "fence": 0}))
  228. def set_runtime_state(self, tenant_id: str, state: str, incident_ref: str | None, *, require_active: bool = False) -> dict[str, Any]:
  229. with self._lock:
  230. current = self._states.setdefault(tenant_id, {"state": "active", "incident_ref": None, "fence": 0})
  231. if require_active and current["state"] != "paused":
  232. raise RuntimeDecisionError("runtime_state_conflict")
  233. current["state"] = state
  234. current["incident_ref"] = incident_ref or current["incident_ref"]
  235. current["fence"] += 1
  236. return dict(current)
  237. def register_generation(self, tenant_id: str, generation: str, *, parent_generation: str, status: str) -> None:
  238. if status not in {"canary", "approved"}:
  239. raise ValueError("generation status invalid")
  240. with self._lock:
  241. self._generations[(tenant_id, generation)] = {"status": status, "parent": parent_generation, "fence": 0}
  242. def promote_generation(self, tenant_id: str, generation: str) -> int:
  243. with self._lock:
  244. current = self._generations.get((tenant_id, generation))
  245. if not current or current["status"] != "canary":
  246. raise RuntimeDecisionError("canary_not_approved")
  247. current["status"] = "approved"
  248. current["fence"] += 1
  249. return current["fence"]
  250. def rollback_generation(self, tenant_id: str, generation: str, fence: int) -> str:
  251. with self._lock:
  252. approved = [entry for (tenant, name), entry in self._generations.items() if tenant == tenant_id and entry["status"] == "approved"]
  253. if not approved or max(entry["fence"] for entry in approved) != fence:
  254. raise RuntimeDecisionError("generation_fence_conflict")
  255. target = self._generations.get((tenant_id, generation))
  256. if not target:
  257. raise RuntimeDecisionError("route_scope_denied")
  258. return generation
  259. def replay(self, invocation_id: str) -> dict[str, Any]:
  260. with self._lock:
  261. for record in self._audit:
  262. if record["invocation_id"] == invocation_id:
  263. return dict(record)
  264. raise RuntimeDecisionError("invocation_not_found")
  265. class GovernedInvocationService:
  266. def __init__(self, repository: InMemoryRuntimeRepository):
  267. self.repository = repository
  268. def authorize(self, raw_request: dict[str, Any]) -> dict[str, Any]:
  269. request = _normalize_request(raw_request)
  270. replay = self.repository.existing(request["tenant_id"], request["idempotency_key"])
  271. if replay:
  272. return replay
  273. state = self.repository.runtime_state(request["tenant_id"])
  274. if state["state"] == "paused":
  275. raise RuntimeDecisionError("runtime_paused")
  276. route = self.repository.find_route(request)
  277. if route is None:
  278. if self.repository.generation_status(request["tenant_id"], request["generation"]) == "canary":
  279. raise RuntimeDecisionError("canary_not_approved")
  280. raise RuntimeDecisionError("route_scope_denied")
  281. if route.canary_status != "approved":
  282. raise RuntimeDecisionError("canary_not_approved")
  283. if not self.repository.allowed_capability(request):
  284. if request["tool_name"] == "knowledge.search":
  285. raise RuntimeDecisionError("capability_denied")
  286. raise RuntimeDecisionError("tool_contract_denied")
  287. fence = self.repository.reserve(request)
  288. return self.repository.write_authorized(request, route, fence)
  289. def record_anomaly(self, tenant_id: str, *, incident_ref: str, severity: str, actor: str) -> dict[str, Any]:
  290. if not _SAFE_ID.fullmatch(tenant_id) or not _SAFE_ID.fullmatch(incident_ref) or actor != "detector":
  291. raise RuntimeDecisionError("anomaly_contract_invalid")
  292. if severity not in {"medium", "high", "critical"}:
  293. raise RuntimeDecisionError("anomaly_severity_invalid")
  294. return self.repository.set_runtime_state(tenant_id, "paused" if severity == "critical" else "degraded", incident_ref)
  295. def recover(self, tenant_id: str, *, approval_ref: str | None, actor: str) -> dict[str, Any]:
  296. if actor != "operator" or not isinstance(approval_ref, str) or not _SAFE_ID.fullmatch(approval_ref):
  297. raise RuntimeDecisionError("recovery_approval_required")
  298. return self.repository.set_runtime_state(tenant_id, "active", None, require_active=True)
  299. def promote_canary(self, tenant_id: str, generation: str, *, approval_ref: str) -> int:
  300. if not _SAFE_ID.fullmatch(approval_ref):
  301. raise RuntimeDecisionError("recovery_approval_required")
  302. return self.repository.promote_generation(tenant_id, generation)
  303. def rollback_generation(self, tenant_id: str, generation: str, *, fence: int, approval_ref: str) -> str:
  304. if not _SAFE_ID.fullmatch(approval_ref):
  305. raise RuntimeDecisionError("recovery_approval_required")
  306. return self.repository.rollback_generation(tenant_id, generation, fence)
  307. def replay(self, invocation_id: str) -> dict[str, Any]:
  308. return self.repository.replay(invocation_id)
  309. def evaluate_governed_dataset(*, train_case_ids: set[str], evaluation_cases: list[dict[str, Any]]) -> dict[str, float]:
  310. case_ids = [str(case.get("case_id", "")) for case in evaluation_cases]
  311. if not case_ids or len(case_ids) != len(set(case_ids)) or train_case_ids.intersection(case_ids):
  312. raise DatasetLeakageError("dataset_leakage_detected")
  313. recalls, reciprocal_ranks, citations, violations, fresh = [], [], [], [], []
  314. for case in evaluation_cases:
  315. expected = case.get("expected_evidence_id")
  316. retrieved = list(case.get("retrieved_evidence_ids") or [])
  317. cited = list(case.get("citation_evidence_ids") or [])
  318. if not isinstance(expected, str) or not expected:
  319. raise DatasetLeakageError("dataset_case_invalid")
  320. recalls.append(float(expected in retrieved))
  321. reciprocal_ranks.append(1.0 / (retrieved.index(expected) + 1) if expected in retrieved else 0.0)
  322. citations.append(float(cited == [expected]))
  323. violations.append(float(not bool(case.get("authorized"))))
  324. fresh.append(float(bool(case.get("fresh"))))
  325. count = float(len(evaluation_cases))
  326. return {"recall": sum(recalls) / count, "mrr": sum(reciprocal_ranks) / count, "citation_correctness": sum(citations) / count, "privilege_violation_rate": sum(violations) / count, "freshness": sum(fresh) / count}