| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379 |
- """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}
|