"""Default-deny, hash-only runtime governance for model and Agent invocations.""" from __future__ import annotations import hashlib import json import re import threading import unicodedata from dataclasses import dataclass from pathlib import Path from typing import Any from uuid import uuid4 from app.core.mcp.governed_invocation import ( InvocationContractError, normalize_mcp_invocation, ) class RuntimeDecisionError(ValueError): def __init__(self, code: str): self.code = code super().__init__(code) class DatasetLeakageError(ValueError): pass def load_engineering_fixture(manifest_path: str | Path) -> dict[str, Any]: """Read each fixture once and verify its declared SHA-256 before evaluation.""" manifest_file = Path(manifest_path) manifest = json.loads(manifest_file.read_text(encoding="utf-8")) if manifest.get("fixture_only") is not True or manifest.get("schema_version") != "1": raise DatasetLeakageError("fixture_manifest_invalid") files = manifest.get("files") if not isinstance(files, dict) or set(files) != {"dataset.json", "temporal-dataset.json"}: raise DatasetLeakageError("fixture_manifest_invalid") loaded: dict[str, Any] = {} for name, expected in files.items(): if not isinstance(expected, str) or not expected.startswith("sha256:"): raise DatasetLeakageError("fixture_manifest_invalid") content = (manifest_file.parent / name).read_bytes() if hashlib.sha256(content).hexdigest() != expected.removeprefix("sha256:"): raise DatasetLeakageError("fixture_digest_mismatch") loaded[name] = json.loads(content) train = set(loaded["dataset.json"].get("train_case_ids", [])) evaluation = {item.get("case_id") for item in loaded["dataset.json"].get("evaluation_cases", [])} if train & evaluation: raise DatasetLeakageError("dataset_leakage_detected") return loaded @dataclass(frozen=True) class ModelRoute: route_id: str provider: str model: str prompt_version: str generation: str tenant_id: str principal_id: str business_domain_uid: str environment: str canary_status: str @dataclass(frozen=True) class BudgetPolicy: token_limit: int cost_limit_micros: int tool_limit: int time_limit_ms: int concurrency_limit: int model_limits: dict[str, int] _SCOPE_FIELDS = ("tenant_id", "principal_id", "business_domain_uid", "environment") _REQUEST_KEYS = frozenset( { *_SCOPE_FIELDS, "provider", "model", "prompt_version", "generation", "interface_type", "tool_name", "action", "requested_capabilities", "input_text", "evidence_refs", "idempotency_key", "estimated_tokens", "estimated_cost_micros", "requested_time_ms", } ) _FORBIDDEN_KEYS = frozenset( { "url", "uri", "endpoint", "path", "file_path", "filename", "command", "shell", "script", "secret", "token", "password", "api_key", "raw_rows", "rows", "records", "payload", "arguments", "output", "raw_input", "raw_output", } ) _SAFE_ID = re.compile(r"^[A-Za-z0-9_.:-]{1,160}$") _INJECTION = re.compile( r"(?is)(ignore\s+(?:all\s+)?(?:previous|prior|system)\s+instructions?|" r"bypass\s+(?:the\s+)?(?:policy|guard|permission)|reveal.{0,50}(?:secret|token|password|api[_ -]?key)|" r"(?:忽略).{0,16}(?:指令|规则|系统)|(?:绕过|关闭|禁用).{0,16}(?:权限|策略|审计))" ) _LOCATION = re.compile(r"(?i)(?:https?://|file://|(?:^|\s)/[A-Za-z0-9_./-]+)") def _hash(value: str) -> str: return hashlib.sha256(value.encode("utf-8")).hexdigest() def _positive_int(value: Any, label: str, maximum: int) -> int: if isinstance(value, bool) or not isinstance(value, int) or not 0 < value <= maximum: raise RuntimeDecisionError(f"{label}_invalid") return value def _closed_request(value: Any) -> dict[str, Any]: if not isinstance(value, dict): raise RuntimeDecisionError("request_invalid") if len(value) > len(_REQUEST_KEYS): raise RuntimeDecisionError("request_invalid") for key in value: normalized = unicodedata.normalize("NFKC", str(key)) if normalized != key or normalized.casefold() in _FORBIDDEN_KEYS: raise RuntimeDecisionError("forbidden_contract_key") unknown = set(value) - _REQUEST_KEYS if unknown or set(value) != _REQUEST_KEYS: raise RuntimeDecisionError("request_schema_invalid") return dict(value) def _normalize_request(value: Any) -> dict[str, Any]: request = _closed_request(value) for key in _SCOPE_FIELDS + ("provider", "model", "prompt_version", "generation", "idempotency_key"): item = request[key] if not isinstance(item, str) or not _SAFE_ID.fullmatch(item): raise RuntimeDecisionError(f"{key}_invalid") text = request["input_text"] if not isinstance(text, str) or not text.strip() or len(text) > 8000: raise RuntimeDecisionError("input_invalid") normalized_text = unicodedata.normalize("NFKC", text).strip() if _INJECTION.search(normalized_text): raise RuntimeDecisionError("prompt_injection_detected") if _LOCATION.search(normalized_text): raise RuntimeDecisionError("input_location_denied") capabilities = request["requested_capabilities"] if not isinstance(capabilities, list) or not capabilities or len(capabilities) > 8: raise RuntimeDecisionError("capability_invalid") if any(not isinstance(item, str) or item not in {"network", "file", "data", "command", "resource", "time"} for item in capabilities): raise RuntimeDecisionError("capability_invalid") evidence = request["evidence_refs"] if not isinstance(evidence, list) or len(evidence) > 20: raise RuntimeDecisionError("citation_invalid") for item in evidence: if not isinstance(item, dict) or set(item) != {"evidence_id", "digest"}: raise RuntimeDecisionError("citation_invalid") if not isinstance(item["evidence_id"], str) or not _SAFE_ID.fullmatch(item["evidence_id"]): raise RuntimeDecisionError("citation_invalid") if not isinstance(item["digest"], str) or not re.fullmatch(r"[a-f0-9]{64}", item["digest"]): raise RuntimeDecisionError("citation_invalid") try: normalize_mcp_invocation( { "interface_type": request["interface_type"], "tool_name": request["tool_name"], "action": request["action"], "arguments_digest": _hash(normalized_text), "evidence_refs": evidence, } ) except InvocationContractError as error: raise RuntimeDecisionError("mcp_contract_invalid") from error return { **{key: request[key] for key in _SCOPE_FIELDS}, **{key: request[key] for key in ("provider", "model", "prompt_version", "generation", "interface_type", "tool_name", "action", "idempotency_key")}, "requested_capabilities": frozenset(request["requested_capabilities"]), "input_hash": _hash(normalized_text), "evidence_refs": tuple((item["evidence_id"], item["digest"]) for item in evidence), "estimated_tokens": _positive_int(request["estimated_tokens"], "estimated_tokens", 1_000_000), "estimated_cost_micros": _positive_int(request["estimated_cost_micros"], "estimated_cost_micros", 1_000_000_000), "requested_time_ms": _positive_int(request["requested_time_ms"], "requested_time_ms", 300_000), } class InMemoryRuntimeRepository: """Thread-safe reference repository used by unit contracts and local wiring.""" def __init__(self) -> None: self._lock = threading.RLock() self._routes: list[ModelRoute] = [] self._capabilities: set[tuple[str, str, str, str, str, str, frozenset[str]]] = set() self._budgets: dict[str, dict[str, Any]] = {} self._decisions: dict[tuple[str, str], dict[str, Any]] = {} self._audit: list[dict[str, Any]] = [] self._states: dict[str, dict[str, Any]] = {} self._generations: dict[tuple[str, str], dict[str, Any]] = {} def register_route(self, route: ModelRoute) -> None: with self._lock: self._routes.append(route) self._generations.setdefault((route.tenant_id, route.generation), {"status": route.canary_status, "fence": 0, "parent": None}) self._states.setdefault(route.tenant_id, {"state": "active", "incident_ref": None, "fence": 0}) 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: if risk_level != "low" or not allowed_capabilities <= {"data"}: raise ValueError("only declared low-risk data capability may be registered locally") with self._lock: self._capabilities.add((tenant_id, business_domain_uid, environment, tool_name, interface_type, action, frozenset(allowed_capabilities))) def set_budget(self, tenant_id: str, policy: BudgetPolicy) -> None: if min(policy.token_limit, policy.cost_limit_micros, policy.tool_limit, policy.time_limit_ms, policy.concurrency_limit) < 1: raise ValueError("budget limits must be positive") with self._lock: 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} def find_route(self, request: dict[str, Any]) -> ModelRoute | None: with self._lock: for route in self._routes: if all(getattr(route, field) == request[field] for field in _SCOPE_FIELDS + ("provider", "model", "prompt_version", "generation")): return route return None def generation_status(self, tenant_id: str, generation: str) -> str | None: with self._lock: entry = self._generations.get((tenant_id, generation)) return entry["status"] if entry else None def allowed_capability(self, request: dict[str, Any]) -> bool: key = (request["tenant_id"], request["business_domain_uid"], request["environment"], request["tool_name"], request["interface_type"], request["action"], request["requested_capabilities"]) with self._lock: return key in self._capabilities def existing(self, tenant_id: str, idempotency_key: str) -> dict[str, Any] | None: with self._lock: record = self._decisions.get((tenant_id, idempotency_key)) return dict(record) if record else None def reserve(self, request: dict[str, Any]) -> int: with self._lock: budget = self._budgets.get(request["tenant_id"]) 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)): raise RuntimeDecisionError("budget_exhausted") 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"][request["model"]] -= 1 budget["fence"] += 1 return budget["fence"] def write_authorized(self, request: dict[str, Any], route: ModelRoute, fence: int) -> dict[str, Any]: 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"]} with self._lock: self._decisions[(request["tenant_id"], request["idempotency_key"])] = record self._audit.append(dict(record)) return dict(record) def remaining_budget(self, tenant_id: str) -> dict[str, int]: with self._lock: budget = self._budgets[tenant_id] return {"tokens": budget["tokens"], "cost_micros": budget["cost"], "tools": budget["tools"], "time_ms": budget["time"], "concurrency": budget["concurrency"]} def audit_records(self) -> list[dict[str, Any]]: with self._lock: return [dict(item) for item in self._audit] def runtime_state(self, tenant_id: str) -> dict[str, Any]: with self._lock: return dict(self._states.get(tenant_id, {"state": "active", "incident_ref": None, "fence": 0})) def set_runtime_state(self, tenant_id: str, state: str, incident_ref: str | None, *, require_active: bool = False) -> dict[str, Any]: with self._lock: current = self._states.setdefault(tenant_id, {"state": "active", "incident_ref": None, "fence": 0}) if require_active and current["state"] != "paused": raise RuntimeDecisionError("runtime_state_conflict") current["state"] = state current["incident_ref"] = incident_ref or current["incident_ref"] current["fence"] += 1 return dict(current) def register_generation(self, tenant_id: str, generation: str, *, parent_generation: str, status: str) -> None: if status not in {"canary", "approved"}: raise ValueError("generation status invalid") with self._lock: self._generations[(tenant_id, generation)] = {"status": status, "parent": parent_generation, "fence": 0} def promote_generation(self, tenant_id: str, generation: str) -> int: with self._lock: current = self._generations.get((tenant_id, generation)) if not current or current["status"] != "canary": raise RuntimeDecisionError("canary_not_approved") current["status"] = "approved" current["fence"] += 1 return current["fence"] def rollback_generation(self, tenant_id: str, generation: str, fence: int) -> str: with self._lock: approved = [entry for (tenant, name), entry in self._generations.items() if tenant == tenant_id and entry["status"] == "approved"] if not approved or max(entry["fence"] for entry in approved) != fence: raise RuntimeDecisionError("generation_fence_conflict") target = self._generations.get((tenant_id, generation)) if not target: raise RuntimeDecisionError("route_scope_denied") return generation def replay(self, invocation_id: str) -> dict[str, Any]: with self._lock: for record in self._audit: if record["invocation_id"] == invocation_id: return dict(record) raise RuntimeDecisionError("invocation_not_found") class GovernedInvocationService: def __init__(self, repository: InMemoryRuntimeRepository): self.repository = repository def authorize(self, raw_request: dict[str, Any]) -> dict[str, Any]: request = _normalize_request(raw_request) replay = self.repository.existing(request["tenant_id"], request["idempotency_key"]) if replay: return replay state = self.repository.runtime_state(request["tenant_id"]) if state["state"] == "paused": raise RuntimeDecisionError("runtime_paused") route = self.repository.find_route(request) if route is None: if self.repository.generation_status(request["tenant_id"], request["generation"]) == "canary": raise RuntimeDecisionError("canary_not_approved") raise RuntimeDecisionError("route_scope_denied") if route.canary_status != "approved": raise RuntimeDecisionError("canary_not_approved") if not self.repository.allowed_capability(request): if request["tool_name"] == "knowledge.search": raise RuntimeDecisionError("capability_denied") raise RuntimeDecisionError("tool_contract_denied") fence = self.repository.reserve(request) return self.repository.write_authorized(request, route, fence) def record_anomaly(self, tenant_id: str, *, incident_ref: str, severity: str, actor: str) -> dict[str, Any]: if not _SAFE_ID.fullmatch(tenant_id) or not _SAFE_ID.fullmatch(incident_ref) or actor != "detector": raise RuntimeDecisionError("anomaly_contract_invalid") if severity not in {"medium", "high", "critical"}: raise RuntimeDecisionError("anomaly_severity_invalid") return self.repository.set_runtime_state(tenant_id, "paused" if severity == "critical" else "degraded", incident_ref) def recover(self, tenant_id: str, *, approval_ref: str | None, actor: str) -> dict[str, Any]: if actor != "operator" or not isinstance(approval_ref, str) or not _SAFE_ID.fullmatch(approval_ref): raise RuntimeDecisionError("recovery_approval_required") return self.repository.set_runtime_state(tenant_id, "active", None, require_active=True) def promote_canary(self, tenant_id: str, generation: str, *, approval_ref: str) -> int: if not _SAFE_ID.fullmatch(approval_ref): raise RuntimeDecisionError("recovery_approval_required") return self.repository.promote_generation(tenant_id, generation) def rollback_generation(self, tenant_id: str, generation: str, *, fence: int, approval_ref: str) -> str: if not _SAFE_ID.fullmatch(approval_ref): raise RuntimeDecisionError("recovery_approval_required") return self.repository.rollback_generation(tenant_id, generation, fence) def replay(self, invocation_id: str) -> dict[str, Any]: return self.repository.replay(invocation_id) def evaluate_governed_dataset(*, train_case_ids: set[str], evaluation_cases: list[dict[str, Any]]) -> dict[str, float]: case_ids = [str(case.get("case_id", "")) for case in evaluation_cases] if not case_ids or len(case_ids) != len(set(case_ids)) or train_case_ids.intersection(case_ids): raise DatasetLeakageError("dataset_leakage_detected") recalls, reciprocal_ranks, citations, violations, fresh = [], [], [], [], [] for case in evaluation_cases: expected = case.get("expected_evidence_id") retrieved = list(case.get("retrieved_evidence_ids") or []) cited = list(case.get("citation_evidence_ids") or []) if not isinstance(expected, str) or not expected: raise DatasetLeakageError("dataset_case_invalid") recalls.append(float(expected in retrieved)) reciprocal_ranks.append(1.0 / (retrieved.index(expected) + 1) if expected in retrieved else 0.0) citations.append(float(cited == [expected])) violations.append(float(not bool(case.get("authorized")))) fresh.append(float(bool(case.get("fresh")))) count = float(len(evaluation_cases)) 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}