diff --git a/.github/workflows/cd.yml b/.github/workflows/cd.yml index 3895472..a9d8d8e 100644 --- a/.github/workflows/cd.yml +++ b/.github/workflows/cd.yml @@ -64,6 +64,11 @@ jobs: # One plan per track (model, adapter, policy), each naming what it # rolls back to and the criteria that trigger it. run: python -m llm_router.canary > canary-plan.json + - name: Render the governance plan + # What a sync would record in MLflow for this catalog: every model and + # adapter revision, its stage, and its benchmark evidence. Planned + # offline; syncing a live MLflow stays disabled with the rest of deploy. + run: python -m llm_router.governance plan > governance-plan.json - name: Validate the Kubernetes manifests run: | curl -sSLo kubeconform.tar.gz https://github.com/yannh/kubeconform/releases/download/v0.7.0/kubeconform-linux-amd64.tar.gz @@ -74,6 +79,7 @@ jobs: name: deployment-plan-${{ env.RELEASE_REF }} path: | canary-plan.json + governance-plan.json config/ray-serve.yaml deploy/kubernetes retention-days: 14 diff --git a/README.md b/README.md index e0ac7dd..6d97258 100644 --- a/README.md +++ b/README.md @@ -121,6 +121,37 @@ returns `502` with retry guidance, and `/readyz` fails while the engine is unhea Set `"stream": true` to receive OpenAI-compatible `text/event-stream` chunks. Streamed results are cached under the same eligibility rules and replayed as chunks on a hit. +### External providers + +The gateway never talks to a provider. An approved external model is reached through a +[LiteLLM](https://docs.litellm.ai/) proxy, configured in +[`config/litellm.yaml`](config/litellm.yaml), so provider credentials stay out of the gateway and +adding a provider does not change it. Every alias in that file must match a non-local model card +in the catalog; a test fails if they drift. The provider model named there is a placeholder: pick +the one your policy approves. + +A request reaches the proxy only when all of these hold: + +1. The operator has set `ROUTER_EXTERNAL_FALLBACK_ENABLED=true`. With `ROUTER_BACKEND=vllm` the + gateway refuses to start unless `ROUTER_EXTERNAL_BASE_URL` is also set. +2. The effective privacy class is `public`, after any tenant floor has been applied. +3. The request sets `routing.allow_external_fallback`, and the tenant does not forbid it. + +The rule is enforced twice. Routing never selects an external model for private or restricted +data, and the dispatch boundary refuses to send it even if routing were wrong, answering `500` +`policy_violation` and counting `router_rejections_total{reason="external_dispatch_refused"}`. + +An eligible request also falls back when the local engine fails with it: unreachable, out of +memory, or circuit open. The response names the model that answered, its route reason says +which model it fell back from, and `router_fallbacks_total` counts it by cause. A fallback +response is cached under the model that produced it. The local engine and the proxy have separate +circuits, so a failing engine does not close the path to the provider. Streamed requests do not +fall back; they fail as described under [Failure behaviour](#failure-behaviour). + +In the cluster the proxy is the only workload allowed to reach the internet, and only the gateway +may call it. This path has been tested against a stand-in transport, not against a running +LiteLLM proxy or a real provider. + ### Ray Serve deployment [`config/ray-serve.yaml`](config/ray-serve.yaml) is generated from the catalog, never @@ -169,9 +200,9 @@ evidence, and no claim is made about what either gains. ## Deployment topology [`deploy/kubernetes`](deploy/kubernetes) holds the namespaced manifests: gateway -Deployment and Service, GPU serving pool, Redis, KEDA autoscaling on queue depth and p95 -latency, a Prometheus `ServiceMonitor`, network policy, and credentials sourced from the -cluster secret manager. No secret material is committed. Unit tests enforce the contract: +Deployment and Service, GPU serving pool, Redis, the LiteLLM proxy, the MLflow server, KEDA +autoscaling on queue depth and p95 latency, a Prometheus `ServiceMonitor`, network policy, and +credentials sourced from the cluster secret manager. No secret material is committed. Unit tests enforce the contract: unprivileged workloads, digest-pinned images, bounded resources, real probes, GPU pool pinning, and `/metrics` reachable only from monitoring. @@ -187,7 +218,8 @@ in-process and correct for a single replica only. Install the client with the ex python -m pip install -e ".[redis]" ``` -CD renders the canary plans (one per track, each with its rollback target), verifies +CD renders the canary plans (one per track, each with its rollback target) and the governance +plan, verifies `config/ray-serve.yaml` against the catalog, and validates the manifests with kubeconform. Applying to a cluster stays disabled until a deployment destination is configured. @@ -216,6 +248,40 @@ be served. A request can never introduce a model path, revision, or adapter. Send `routing.domain` to request a domain adapter; the router applies the promoted adapter with the largest measured quality gain for that base revision and task, or none at all. +### Governance in MLflow + +The catalog decides what is served; [MLflow](https://mlflow.org/docs/latest/) keeps the record. +`llm_router.governance` syncs the catalog into an MLflow model registry and tracking store: + +| Catalog record | In MLflow | +|---|---| +| Model or adapter revision | One model version, tagged with its card (license, tier, hardware, limitations, base revision, dataset version), artifact location, and checksum. | +| Lifecycle stage | A version tag, plus a `staging` or `production` alias, so `models:/general-local@production` resolves. | +| Benchmark run | One run with its measurements as metrics and its dataset, workload, hardware, driver, and engine revision as parameters. | +| Stage change | One appended run in the promotions experiment: from, to, commit, and policy version. | + +```bash +python -m pip install -e ".[governance]" +python -m llm_router.governance plan # what a first sync would record +python -m llm_router.governance sync --tracking-uri "$MLFLOW_TRACKING_URI" +python -m llm_router.governance verify --tracking-uri "$MLFLOW_TRACKING_URI" # exit 3 on drift +python -m llm_router.governance history --name general-local --tracking-uri "$MLFLOW_TRACKING_URI" +``` + +- Governance flows one way. The gateway never reads MLflow, so a stage changed by hand there + changes nothing in production; `verify` reports it and the next `sync` puts it back. +- A revision is immutable. If a recorded revision turns up with a different checksum, the sync + is refused: a changed artifact needs a new revision. +- A revision the catalog replaces or removes is retired to `deprecated`, never deleted, so the + rollback target stays on record. +- Models registered in MLflow by anyone else are left alone. + +The sync records where an artifact belongs under `--artifact-root`; it does not upload weights. +CD renders the plan as an artifact. Running `sync` against a live MLflow is part of the deploy +step, which stays disabled until a destination is configured. The tests run against a real MLflow +on SQLite; the server deployment in [`deploy/kubernetes/mlflow.yaml`](deploy/kubernetes/mlflow.yaml) +has not been run on a cluster. + ## Failure behaviour | Condition | What the gateway does | @@ -231,8 +297,10 @@ An engine error body is inspected for an out-of-memory report and then discarded it can echo the prompt it rejected. `router_engine_circuit_open` reports 0 closed, 0.5 half-open, 1 open. -A request that fails is not retried on another model. Every local model shares the one engine a -gateway faces, so a retry would meet the same failure; the caller is told when to come back. +A request that fails is not retried on another local model. Every local model shares the one +engine a gateway faces, so a retry would meet the same failure; the caller is told when to come +back. The one exception is a request already entitled to the +[external provider](#external-providers), which falls back to it. Cold start is measured, not assumed (see `router_model_load_seconds` below). Tiers that keep a warm replica never pay it on the request path. The high-capability tier scales to zero, so its @@ -333,7 +401,8 @@ in-cluster scrapers can read it; restrict it with network policy rather than a b | `router_tokens_total` | Prompt and completion tokens per model. | | `router_inflight_requests` / `router_queued_requests` | Live capacity and queue depth. | | `router_routes_total` | Requests per route with task and privacy class. | -| `router_external_fallback_total` | Fallback frequency. | +| `router_external_fallback_total` | Requests answered by an external model. | +| `router_fallbacks_total` | Fallbacks after a local engine failure, by cause and by the models fallen back from and to. | | `router_queue_delay_prediction_error_ms` | Predicted versus observed queue delay. | | `router_rejections_total` | Quota, overload, and policy rejections. | | `router_cache_events_total` | Cache lookups by cache and result. | @@ -444,6 +513,8 @@ All settings use the `ROUTER_` prefix. | `ROUTER_ADMISSION_TIMEOUT_SECONDS` | `0.25` | Time allowed to wait for capacity. | | `ROUTER_QUOTA_REQUESTS_PER_MINUTE` | `120` | Per-token sliding-window quota. | | `ROUTER_EXTERNAL_FALLBACK_ENABLED` | `false` | Operator gate for external fallback. | +| `ROUTER_EXTERNAL_BASE_URL` | _(empty)_ | LiteLLM proxy address; required with `vllm` when external fallback is enabled. | +| `ROUTER_EXTERNAL_API_KEY` | _(empty)_ | Key the gateway presents to the proxy. | | `ROUTER_REDIS_URL` | _(empty)_ | Shared cache and quota state; in-process when empty. | | `ROUTER_TENANT_KEYS` | _(empty)_ | `tenant:key` bindings; bare `ROUTER_API_KEYS` keys use the default tenant. | | `ROUTER_OTLP_ENDPOINT` | _(empty)_ | OTLP/HTTP trace collector; tracing is a no-op when empty. | diff --git a/deploy/kubernetes/kustomization.yaml b/deploy/kubernetes/kustomization.yaml index 8e17a43..6503194 100644 --- a/deploy/kubernetes/kustomization.yaml +++ b/deploy/kubernetes/kustomization.yaml @@ -9,4 +9,5 @@ resources: - state.yaml - network-policy.yaml - external-provider.yaml + - mlflow.yaml - observability.yaml diff --git a/deploy/kubernetes/mlflow.yaml b/deploy/kubernetes/mlflow.yaml new file mode 100644 index 0000000..5321352 --- /dev/null +++ b/deploy/kubernetes/mlflow.yaml @@ -0,0 +1,181 @@ +# MLflow tracking server and model registry: the governance record of every +# model and adapter revision, its benchmark evidence, and its promotions. +# The gateway never reads it; the catalog shipped in the image decides what is +# served. Only the delivery pipeline writes here, with +# `python -m llm_router.governance sync`. +apiVersion: apps/v1 +kind: Deployment +metadata: + name: mlflow + namespace: llm-routing + labels: + app.kubernetes.io/name: mlflow + app.kubernetes.io/part-of: local-llm-router +spec: + replicas: 1 + selector: + matchLabels: + app.kubernetes.io/name: mlflow + template: + metadata: + labels: + app.kubernetes.io/name: mlflow + app.kubernetes.io/part-of: local-llm-router + spec: + securityContext: + runAsNonRoot: true + runAsUser: 10001 + seccompProfile: + type: RuntimeDefault + containers: + - name: mlflow + image: ghcr.io/mlflow/mlflow@sha256:REPLACE_ME + imagePullPolicy: IfNotPresent + command: [mlflow, server] + args: + - --host=0.0.0.0 + - --port=5000 + # MLflow rejects requests whose Host header it does not expect. + - --allowed-hosts=mlflow.llm-routing.svc.cluster.local,mlflow.llm-routing.svc.cluster.local:5000 + # Clients upload and download through the server, so object + # storage credentials stay here and are never handed to a client. + - --serve-artifacts + - --artifacts-destination=s3://llm-routing-artifacts + ports: + - name: http + containerPort: 5000 + env: + # The database address carries its password, so the whole value + # comes from the secret manager. + - name: MLFLOW_BACKEND_STORE_URI + valueFrom: + secretKeyRef: + name: mlflow-credentials + key: backend-store-uri + - name: MLFLOW_S3_ENDPOINT_URL + value: http://object-storage.storage.svc.cluster.local:9000 + - name: AWS_ACCESS_KEY_ID + valueFrom: + secretKeyRef: + name: mlflow-credentials + key: artifact-access-key-id + - name: AWS_SECRET_ACCESS_KEY + valueFrom: + secretKeyRef: + name: mlflow-credentials + key: artifact-secret-access-key + securityContext: + allowPrivilegeEscalation: false + readOnlyRootFilesystem: true + capabilities: + drop: [ALL] + resources: + requests: + cpu: 250m + memory: 512Mi + limits: + cpu: "1" + memory: 2Gi + livenessProbe: + httpGet: + path: /health + port: http + initialDelaySeconds: 20 + periodSeconds: 30 + readinessProbe: + httpGet: + path: /health + port: http + initialDelaySeconds: 10 + periodSeconds: 10 + volumeMounts: + - name: tmp + mountPath: /tmp + volumes: + - name: tmp + emptyDir: {} + terminationGracePeriodSeconds: 30 +--- +apiVersion: v1 +kind: Service +metadata: + name: mlflow + namespace: llm-routing + labels: + app.kubernetes.io/name: mlflow +spec: + selector: + app.kubernetes.io/name: mlflow + ports: + - name: http + port: 5000 + targetPort: http +--- +# Kept apart from the gateway's credentials: the gateway has no use for the +# governance database or the artifact store, so it is never given either. +apiVersion: external-secrets.io/v1 +kind: ExternalSecret +metadata: + name: mlflow-credentials + namespace: llm-routing +spec: + refreshInterval: 1h + secretStoreRef: + name: platform-secret-store + kind: ClusterSecretStore + target: + name: mlflow-credentials + data: + - secretKey: backend-store-uri + remoteRef: + key: llm-routing/mlflow + property: backend_store_uri + - secretKey: artifact-access-key-id + remoteRef: + key: llm-routing/mlflow + property: artifact_access_key_id + - secretKey: artifact-secret-access-key + remoteRef: + key: llm-routing/mlflow + property: artifact_secret_access_key +--- +apiVersion: networking.k8s.io/v1 +kind: NetworkPolicy +metadata: + name: mlflow + namespace: llm-routing +spec: + podSelector: + matchLabels: + app.kubernetes.io/name: mlflow + policyTypes: [Ingress, Egress] + ingress: + # The delivery pipeline syncs and verifies the catalog from here. No + # workload in this namespace, the gateway included, may reach MLflow. + - from: + - namespaceSelector: + matchLabels: + kubernetes.io/metadata.name: platform-delivery + ports: + - protocol: TCP + port: 5000 + egress: + - to: + - namespaceSelector: + matchLabels: + kubernetes.io/metadata.name: kube-system + ports: + - protocol: UDP + port: 53 + - protocol: TCP + port: 53 + # The backing database and the artifact object store, and nothing else. + - to: + - namespaceSelector: + matchLabels: + kubernetes.io/metadata.name: storage + ports: + - protocol: TCP + port: 5432 + - protocol: TCP + port: 9000 diff --git a/pyproject.toml b/pyproject.toml index 2b425d6..38a8fd0 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -22,17 +22,25 @@ dependencies = [ redis = [ "redis>=8.1.0,<9", ] +governance = [ + "alembic>=1.13,<2", + "mlflow-skinny>=3.5,<4", + "sqlalchemy>=2,<3", +] tracing = [ "opentelemetry-exporter-otlp-proto-http>=1.30,<2", "opentelemetry-sdk>=1.30,<2", ] dev = [ + "alembic>=1.13,<2", + "mlflow-skinny>=3.5,<4", "mypy>=2.3.1,<3", "opentelemetry-sdk>=1.30,<2", "pytest>=9.1.1,<10", "pytest-asyncio>=1.4.0,<2", "pytest-cov>=7.1.0,<8", "ruff>=0.16.9,<1", + "sqlalchemy>=2,<3", "types-PyYAML>=6.0.12.20260906,<7", ] @@ -43,6 +51,10 @@ packages = ["src/llm_router"] addopts = "-ra --strict-config --strict-markers" testpaths = ["tests"] pythonpath = ["."] +filterwarnings = [ + # Raised inside MLflow's own SQLAlchemy models; nothing here can act on it. + "ignore::sqlalchemy.exc.SADeprecationWarning", +] markers = [ "integration: tests that exercise multiple application components", ] @@ -75,6 +87,13 @@ module = ["redis.*"] ignore_missing_imports = true follow_imports = "skip" +[[tool.mypy.overrides]] +# MLflow ships in the optional governance extra and is imported lazily, only +# when a tracking URI is given; the client is used behind GovernanceStore. +module = ["mlflow.*"] +ignore_missing_imports = true +follow_imports = "skip" + [[tool.mypy.overrides]] # The OTLP exporter ships in the optional tracing extra and is imported lazily, # only when a collector endpoint is configured. diff --git a/src/llm_router/governance.py b/src/llm_router/governance.py new file mode 100644 index 0000000..f78a11f --- /dev/null +++ b/src/llm_router/governance.py @@ -0,0 +1,711 @@ +"""Model governance in MLflow (sections 6, 7.5 and 17). + +The catalog decides what may be served. MLflow keeps the record of it: every +model and adapter revision, where its artifact lives, the benchmark evidence +behind it, the stage it is in, and each promotion that moved it there. + +Governance flows one way. The catalog is reviewed and merged, then synced into +MLflow; nothing here reads MLflow to decide what the gateway serves. A stage +changed by hand in MLflow therefore changes nothing in production, and +``verify`` reports it as drift. +""" + +import json +import os +import re +import time +from collections.abc import Sequence +from enum import StrEnum +from typing import Any, Protocol + +from pydantic import BaseModel + +from llm_router.registry import ( + AdapterProfile, + BenchmarkRun, + LifecycleStage, + ModelCard, + Registry, +) + +DEFAULT_ARTIFACT_ROOT = "s3://llm-routing-artifacts" +DEFAULT_EXPERIMENT_PREFIX = "llm-routing" +# Stages a consumer may resolve by alias, as in models:/general-local@production. +ALIASED_STAGES = frozenset({LifecycleStage.STAGING, LifecycleStage.PRODUCTION}) +MANAGED_TAG = "catalog.managed" +REVISION_TAG = "catalog.revision" +STAGE_TAG = "catalog.stage" +CHECKSUM_TAG = "catalog.checksum" +BENCHMARK_TAG = "catalog.benchmark_id" +SUBJECT_TAG = "catalog.subject" + + +class GovernanceError(RuntimeError): + """Raised when the catalog cannot be recorded as it stands.""" + + +class ImmutableRevisionError(GovernanceError): + """A recorded revision now claims a different artifact.""" + + +class SubjectKind(StrEnum): + MODEL = "model" + ADAPTER = "adapter" + + +class GovernedVersion(BaseModel): + """One catalog revision as it should be recorded.""" + + name: str + kind: SubjectKind + revision: str + stage: LifecycleStage + source: str + checksum: str | None = None + description: str = "" + tags: dict[str, str] = {} + + +class RecordedVersion(BaseModel): + """One revision as the governance store holds it.""" + + name: str + revision: str + stage: LifecycleStage | None = None + checksum: str | None = None + tags: dict[str, str] = {} + + +class Promotion(BaseModel): + """One stage transition. Promotions are appended, never rewritten.""" + + name: str + revision: str + from_stage: LifecycleStage | None + to_stage: LifecycleStage + commit: str = "unknown" + policy_version: str = "" + reason: str = "" + recorded_at_ms: int = 0 + + +class ActionKind(StrEnum): + REGISTER = "register" + TRANSITION = "transition" + ANNOTATE = "annotate" + BENCHMARK = "benchmark" + + +class Action(BaseModel): + """One difference between the catalog and the governance store.""" + + kind: ActionKind + name: str + revision: str = "" + from_stage: LifecycleStage | None = None + to_stage: LifecycleStage | None = None + reason: str = "" + + def describe(self) -> str: + subject = f"{self.name}@{self.revision}" if self.revision else self.name + if self.kind is ActionKind.REGISTER: + return f"{subject} is not recorded; the catalog has it in {self.to_stage}" + if self.kind is ActionKind.TRANSITION: + recorded = self.from_stage or "unstaged" + return f"{subject} is recorded as {recorded}, expected {self.to_stage} ({self.reason})" + if self.kind is ActionKind.ANNOTATE: + return f"{subject} is recorded with out-of-date card metadata" + return f"benchmark {subject} is in the catalog but has no recorded run" + + +class GovernanceStore(Protocol): + """What governance needs from MLflow, and nothing else.""" + + def names(self) -> frozenset[str]: ... + + def versions(self, name: str) -> tuple[RecordedVersion, ...]: ... + + def register(self, subject: GovernedVersion) -> None: ... + + def annotate(self, subject: GovernedVersion) -> None: ... + + def set_stage(self, name: str, revision: str, stage: LifecycleStage) -> None: ... + + def benchmark_ids(self) -> frozenset[str]: ... + + def log_benchmark(self, run: BenchmarkRun) -> None: ... + + def log_promotion(self, promotion: Promotion) -> None: ... + + def promotions(self, name: str) -> tuple[Promotion, ...]: ... + + +def _slug(revision: str) -> str: + return re.sub(r"[^A-Za-z0-9._-]+", "-", revision).strip("-") + + +def _checksums(registry: Registry) -> tuple[dict[str, str], dict[str, str]]: + """Artifact checksums from the deployment that is in production.""" + + models: dict[str, str] = {} + adapters: dict[str, str] = {} + for deployment in registry.deployments: + if deployment.stage is LifecycleStage.PRODUCTION: + models.update(deployment.model_checksums) + adapters.update(deployment.adapter_checksums) + return models, adapters + + +def _model_version(card: ModelCard, artifact_root: str, checksum: str | None) -> GovernedVersion: + hardware = card.hardware + return GovernedVersion( + name=card.id, + kind=SubjectKind.MODEL, + revision=card.revision, + stage=card.stage, + # An external model has no artifact of ours; the record points at the + # proxy alias that stands for it. + source=( + f"{artifact_root}/models/{card.id}/{_slug(card.revision)}" + if card.local + else f"litellm://{card.id}" + ), + checksum=checksum, + description=card.intended_tasks, + tags={ + "catalog.kind": SubjectKind.MODEL.value, + "catalog.tier": card.tier.value, + "catalog.local": str(card.local).lower(), + "catalog.license": card.license, + "catalog.tokenizer": card.tokenizer, + "catalog.context_limit": str(card.context_limit), + "catalog.quantization": card.quantization.value, + "catalog.hardware": ( + f"{hardware.count}x {hardware.accelerator}, " + f"{hardware.minimum_memory_gb} GB, tensor parallel {hardware.tensor_parallel_size}" + ), + "catalog.supported_tasks": ",".join( + sorted(task.value for task in card.supported_tasks) + ), + "catalog.intended_tasks": card.intended_tasks, + "catalog.limitations": card.limitations, + "catalog.evaluation_references": ",".join(card.evaluation_references), + }, + ) + + +def _adapter_version( + adapter: AdapterProfile, artifact_root: str, checksum: str | None +) -> GovernedVersion: + return GovernedVersion( + name=adapter.id, + kind=SubjectKind.ADAPTER, + revision=adapter.adapter_revision, + stage=adapter.stage, + source=f"{artifact_root}/adapters/{adapter.id}/{_slug(adapter.adapter_revision)}", + checksum=checksum, + description=f"{adapter.domain} adapter for {adapter.base_model_id}", + tags={ + "catalog.kind": SubjectKind.ADAPTER.value, + "catalog.base_model_id": adapter.base_model_id, + "catalog.base_revision": adapter.base_revision, + "catalog.domain": adapter.domain, + "catalog.intended_tasks": ",".join( + sorted(task.value for task in adapter.intended_tasks) + ), + "catalog.dataset_version": adapter.dataset_version, + "catalog.quality_delta": f"{adapter.benchmark.quality_delta:g}", + "catalog.regressions": ",".join(adapter.benchmark.regressions), + "catalog.quantized": str(adapter.quantized).lower(), + }, + ) + + +def governed_versions( + registry: Registry, artifact_root: str = DEFAULT_ARTIFACT_ROOT +) -> tuple[GovernedVersion, ...]: + """Every model and adapter revision the catalog holds, as a governance record.""" + + root = artifact_root.rstrip("/") + model_checksums, adapter_checksums = _checksums(registry) + versions = [ + *(_model_version(card, root, model_checksums.get(card.id)) for card in registry.models), + *( + _adapter_version(adapter, root, adapter_checksums.get(adapter.id)) + for adapter in registry.adapters + ), + ] + names = [version.name for version in versions] + duplicates = sorted({name for name in names if names.count(name) > 1}) + if duplicates: + raise GovernanceError( + f"models and adapters share one namespace in MLflow; duplicated: {duplicates}" + ) + return tuple(versions) + + +def _recorded_tags(subject: GovernedVersion) -> dict[str, str]: + tags = {**subject.tags, MANAGED_TAG: "true", REVISION_TAG: subject.revision} + if subject.checksum: + tags[CHECKSUM_TAG] = subject.checksum + return tags + + +def _card_tags(tags: dict[str, str]) -> dict[str, str]: + """The catalog-owned tags that describe the card, leaving stage aside.""" + + return { + key: value for key, value in tags.items() if key.startswith("catalog.") and key != STAGE_TAG + } + + +def _retirement(recorded: RecordedVersion, reason: str) -> "Action": + return Action( + kind=ActionKind.TRANSITION, + name=recorded.name, + revision=recorded.revision, + from_stage=recorded.stage, + to_stage=LifecycleStage.DEPRECATED, + reason=reason, + ) + + +def plan( + registry: Registry, store: GovernanceStore, artifact_root: str = DEFAULT_ARTIFACT_ROOT +) -> tuple[Action, ...]: + """Everything that must change for the store to match the catalog. + + An empty plan means the store is in step. A revision is immutable: if the + store holds it with one checksum and the catalog now gives another, that + is a different artifact under an old name and nothing is planned at all. + """ + + actions: list[Action] = [] + wanted = governed_versions(registry, artifact_root) + for subject in wanted: + recorded = {item.revision: item for item in store.versions(subject.name)} + current = recorded.get(subject.revision) + if current is None: + actions.append( + Action( + kind=ActionKind.REGISTER, + name=subject.name, + revision=subject.revision, + to_stage=subject.stage, + reason="new revision in the catalog", + ) + ) + else: + if current.checksum and subject.checksum and current.checksum != subject.checksum: + raise ImmutableRevisionError( + f"{subject.name}@{subject.revision} is recorded with checksum " + f"{current.checksum} but the catalog gives {subject.checksum}; " + "a changed artifact needs a new revision" + ) + if current.stage is not subject.stage: + actions.append( + Action( + kind=ActionKind.TRANSITION, + name=subject.name, + revision=subject.revision, + from_stage=current.stage, + to_stage=subject.stage, + reason="stage differs from the catalog", + ) + ) + if _card_tags(current.tags) != _card_tags(_recorded_tags(subject)): + actions.append( + Action( + kind=ActionKind.ANNOTATE, + name=subject.name, + revision=subject.revision, + reason="card metadata changed in the catalog", + ) + ) + # A revision the catalog replaced is kept for rollback, but it may not + # go on claiming a live stage. + actions.extend( + _retirement(item, "superseded in the catalog") + for revision, item in sorted(recorded.items()) + if revision != subject.revision and item.stage is not LifecycleStage.DEPRECATED + ) + known = {subject.name for subject in wanted} + for name in sorted(store.names() - known): + actions.extend( + _retirement(item, "removed from the catalog") + for item in store.versions(name) + if item.stage is not LifecycleStage.DEPRECATED + ) + recorded_benchmarks = store.benchmark_ids() + actions.extend( + Action(kind=ActionKind.BENCHMARK, name=run.id, reason="evidence not yet recorded") + for run in registry.benchmarks + if run.id not in recorded_benchmarks + ) + return tuple(actions) + + +def apply( + actions: Sequence[Action], + registry: Registry, + store: GovernanceStore, + *, + artifact_root: str = DEFAULT_ARTIFACT_ROOT, + commit: str = "unknown", +) -> None: + """Carry out a plan, recording every stage change as a promotion.""" + + subjects = { + (subject.name, subject.revision): subject + for subject in governed_versions(registry, artifact_root) + } + benchmarks = {run.id: run for run in registry.benchmarks} + for action in actions: + if action.kind is ActionKind.BENCHMARK: + store.log_benchmark(benchmarks[action.name]) + continue + if action.kind is ActionKind.ANNOTATE: + store.annotate(subjects[action.name, action.revision]) + continue + if action.kind is ActionKind.REGISTER: + store.register(subjects[action.name, action.revision]) + stage = action.to_stage or LifecycleStage.DEVELOPMENT + store.set_stage(action.name, action.revision, stage) + store.log_promotion( + Promotion( + name=action.name, + revision=action.revision, + from_stage=action.from_stage, + to_stage=stage, + commit=commit, + policy_version=registry.policy.version, + reason=action.reason, + recorded_at_ms=int(time.time() * 1000), + ) + ) + + +def sync( + registry: Registry, + store: GovernanceStore, + *, + artifact_root: str = DEFAULT_ARTIFACT_ROOT, + commit: str = "unknown", +) -> tuple[Action, ...]: + """Bring the store in step with the catalog and return what was done.""" + + actions = plan(registry, store, artifact_root) + apply(actions, registry, store, artifact_root=artifact_root, commit=commit) + return actions + + +def verify( + registry: Registry, store: GovernanceStore, artifact_root: str = DEFAULT_ARTIFACT_ROOT +) -> tuple[str, ...]: + """Describe every way the store has drifted from the catalog.""" + + try: + return tuple(action.describe() for action in plan(registry, store, artifact_root)) + except ImmutableRevisionError as error: + return (str(error),) + + +class InMemoryStore: + """A governance store with no backing service. + + It lets a plan be rendered offline, where it shows what a first sync into + an empty MLflow would record. + """ + + def __init__(self) -> None: + self._versions: dict[str, dict[str, RecordedVersion]] = {} + self._benchmarks: dict[str, BenchmarkRun] = {} + self._promotions: list[Promotion] = [] + + def names(self) -> frozenset[str]: + return frozenset(self._versions) + + def versions(self, name: str) -> tuple[RecordedVersion, ...]: + return tuple(self._versions.get(name, {}).values()) + + def register(self, subject: GovernedVersion) -> None: + self._versions.setdefault(subject.name, {})[subject.revision] = RecordedVersion( + name=subject.name, + revision=subject.revision, + checksum=subject.checksum, + tags=_recorded_tags(subject), + ) + + def annotate(self, subject: GovernedVersion) -> None: + recorded = self._versions[subject.name][subject.revision] + recorded.tags = _recorded_tags(subject) + recorded.checksum = subject.checksum + + def set_stage(self, name: str, revision: str, stage: LifecycleStage) -> None: + self._versions[name][revision].stage = stage + + def benchmark_ids(self) -> frozenset[str]: + return frozenset(self._benchmarks) + + def log_benchmark(self, run: BenchmarkRun) -> None: + self._benchmarks[run.id] = run + + def log_promotion(self, promotion: Promotion) -> None: + self._promotions.append(promotion) + + def promotions(self, name: str) -> tuple[Promotion, ...]: + return tuple(item for item in self._promotions if item.name == name) + + +def benchmark_record(run: BenchmarkRun) -> tuple[dict[str, str], dict[str, float]]: + """Split a benchmark into what identifies it and what it measured.""" + + params = { + "dataset_version": run.dataset_version, + "workload_version": run.workload_version, + "hardware": run.hardware, + "driver": run.driver, + "container_digest": run.container_digest, + "engine_revision": run.engine_revision, + "model_revision": run.model_revision, + "adapter_revision": run.adapter_revision or "", + "task": run.task.value if run.task else "", + "concurrency": str(run.concurrency), + "prompt_tokens_p50": str(run.prompt_tokens_p50), + "prompt_tokens_p95": str(run.prompt_tokens_p95), + "engine_settings": json.dumps(run.engine_settings, sort_keys=True), + } + metrics = { + "quality_score": run.quality_score, + "latency_p95_ms": run.latency_p95_ms, + "throughput_rps": run.throughput_rps, + "gpu_seconds_per_request": run.gpu_seconds_per_request, + } + if run.gpu_memory_gb is not None: + metrics["gpu_memory_gb"] = run.gpu_memory_gb + if run.draft_acceptance_rate is not None: + metrics["draft_acceptance_rate"] = run.draft_acceptance_rate + return params, metrics + + +def _stage(value: str | None) -> LifecycleStage | None: + try: + return LifecycleStage(value) if value else None + except ValueError: + # A stage typed into MLflow by hand is not one of ours; report the + # version as unstaged so the plan puts it right. + return None + + +class MlflowStore: + """Governance records kept in an MLflow tracking server and model registry. + + Each catalog revision is one model version. Its stage is a version tag, + and staging and production are also aliases so a consumer can resolve + ``models:/@production``. Benchmarks and promotions are runs in two + experiments, which makes promotion history append-only. + """ + + def __init__(self, client: Any, experiment_prefix: str = DEFAULT_EXPERIMENT_PREFIX) -> None: + self.client = client + self.benchmark_experiment = f"{experiment_prefix}/benchmarks" + self.promotion_experiment = f"{experiment_prefix}/promotions" + + @classmethod + def from_uri( + cls, tracking_uri: str, experiment_prefix: str = DEFAULT_EXPERIMENT_PREFIX + ) -> "MlflowStore": + try: + from mlflow.tracking import MlflowClient + except ImportError as error: # pragma: no cover - only without the governance extra + raise GovernanceError( + 'MLflow is not installed; install it with pip install -e ".[governance]"' + ) from error + client = MlflowClient(tracking_uri=tracking_uri, registry_uri=tracking_uri) + return cls(client, experiment_prefix) + + def _model_versions(self, name: str) -> list[Any]: + return list(self.client.search_model_versions(f"name='{name}'")) + + def _find(self, name: str, revision: str) -> Any: + for version in self._model_versions(name): + if version.tags.get(REVISION_TAG) == revision: + return version + raise GovernanceError(f"{name}@{revision} is not recorded") + + def _experiment(self, name: str) -> str | None: + experiment = self.client.get_experiment_by_name(name) + return None if experiment is None else str(experiment.experiment_id) + + def _ensure_experiment(self, name: str) -> str: + return self._experiment(name) or str(self.client.create_experiment(name)) + + def _runs(self, experiment_name: str, filter_string: str = "") -> list[Any]: + experiment_id = self._experiment(experiment_name) + if experiment_id is None: + return [] + runs: list[Any] = [] + token: str | None = None + while True: + page = self.client.search_runs( + [experiment_id], filter_string=filter_string, page_token=token + ) + runs.extend(page) + token = getattr(page, "token", None) + if not token: + return runs + + def names(self) -> frozenset[str]: + return frozenset( + model.name + for model in self.client.search_registered_models( + filter_string=f"tags.`{MANAGED_TAG}` = 'true'" + ) + ) + + def versions(self, name: str) -> tuple[RecordedVersion, ...]: + return tuple( + RecordedVersion( + name=name, + revision=version.tags[REVISION_TAG], + stage=_stage(version.tags.get(STAGE_TAG)), + checksum=version.tags.get(CHECKSUM_TAG), + tags=dict(version.tags), + ) + for version in self._model_versions(name) + if REVISION_TAG in version.tags + ) + + def register(self, subject: GovernedVersion) -> None: + if not self.client.search_registered_models(filter_string=f"name = '{subject.name}'"): + self.client.create_registered_model( + subject.name, + tags={MANAGED_TAG: "true", "catalog.kind": subject.kind.value}, + description=subject.description, + ) + self.client.create_model_version( + subject.name, + source=subject.source, + tags=_recorded_tags(subject), + description=subject.description, + ) + + def annotate(self, subject: GovernedVersion) -> None: + version = self._find(subject.name, subject.revision) + wanted = _recorded_tags(subject) + for key, value in wanted.items(): + self.client.set_model_version_tag(subject.name, version.version, key, value) + for key in _card_tags(dict(version.tags)).keys() - wanted.keys(): + self.client.delete_model_version_tag(subject.name, version.version, key) + + def set_stage(self, name: str, revision: str, stage: LifecycleStage) -> None: + version = self._find(name, revision) + previous = _stage(version.tags.get(STAGE_TAG)) + self.client.set_model_version_tag(name, version.version, STAGE_TAG, stage.value) + if stage in ALIASED_STAGES: + self.client.set_registered_model_alias(name, stage.value, version.version) + if previous in ALIASED_STAGES and previous is not stage: + # The alias may already have moved on to a newer revision; only + # drop it while it still points here. + aliases = self.client.get_registered_model(name).aliases + if str(aliases.get(previous.value)) == str(version.version): + self.client.delete_registered_model_alias(name, previous.value) + + def benchmark_ids(self) -> frozenset[str]: + return frozenset( + run.data.tags[BENCHMARK_TAG] + for run in self._runs(self.benchmark_experiment) + if BENCHMARK_TAG in run.data.tags + ) + + def log_benchmark(self, run: BenchmarkRun) -> None: + params, metrics = benchmark_record(run) + created = self.client.create_run( + self._ensure_experiment(self.benchmark_experiment), + tags={BENCHMARK_TAG: run.id}, + run_name=run.id, + ) + run_id = created.info.run_id + for key, value in params.items(): + self.client.log_param(run_id, key, value) + for key, measured in metrics.items(): + self.client.log_metric(run_id, key, measured) + self.client.set_terminated(run_id) + + def log_promotion(self, promotion: Promotion) -> None: + created = self.client.create_run( + self._ensure_experiment(self.promotion_experiment), + tags={SUBJECT_TAG: promotion.name}, + run_name=f"{promotion.name}@{promotion.revision} -> {promotion.to_stage.value}", + ) + run_id = created.info.run_id + for key, value in promotion.model_dump(mode="json").items(): + self.client.log_param(run_id, key, "" if value is None else str(value)) + self.client.set_terminated(run_id) + + def promotions(self, name: str) -> tuple[Promotion, ...]: + found: list[Promotion] = [] + for run in self._runs(self.promotion_experiment, f"tags.`{SUBJECT_TAG}` = '{name}'"): + record = dict(run.data.params) + # A first registration has no stage to come from. + record["from_stage"] = record.get("from_stage") or None + found.append(Promotion.model_validate(record)) + return tuple(sorted(found, key=lambda item: item.recorded_at_ms)) + + +def main(argv: Sequence[str] | None = None) -> int: + """Plan, apply, or check the catalog's record in MLflow. + + ``plan`` prints what a sync would change; with no tracking URI it plans + against an empty store. ``sync`` applies it. ``verify`` exits 3 when the + store has drifted from the catalog, so a pipeline can fail on it. + ``history`` prints the promotions recorded for one model or adapter. + """ + + import argparse + + from llm_router.registry import load_registry + + parser = argparse.ArgumentParser(description=main.__doc__) + parser.add_argument("command", choices=["plan", "sync", "verify", "history"]) + parser.add_argument("--catalog", default="config/registry.yaml") + parser.add_argument("--tracking-uri", default=os.environ.get("MLFLOW_TRACKING_URI", "")) + parser.add_argument("--artifact-root", default=DEFAULT_ARTIFACT_ROOT) + parser.add_argument("--commit", default=os.environ.get("GITHUB_SHA", "unknown")) + parser.add_argument("--name", help="model or adapter whose promotion history to print") + arguments = parser.parse_args(argv) + + if arguments.command != "plan" and not arguments.tracking_uri: + parser.error(f"{arguments.command} needs --tracking-uri or MLFLOW_TRACKING_URI") + if arguments.command == "history" and not arguments.name: + parser.error("history needs --name") + + store: GovernanceStore = ( + MlflowStore.from_uri(arguments.tracking_uri) if arguments.tracking_uri else InMemoryStore() + ) + registry = load_registry(arguments.catalog) + + if arguments.command == "history": + history = store.promotions(arguments.name) + print(json.dumps([item.model_dump(mode="json") for item in history], indent=2)) + return 0 + if arguments.command == "verify": + drift = verify(registry, store, arguments.artifact_root) + print(json.dumps({"in_step": not drift, "drift": list(drift)}, indent=2)) + return 3 if drift else 0 + try: + if arguments.command == "sync": + actions = sync( + registry, store, artifact_root=arguments.artifact_root, commit=arguments.commit + ) + else: + actions = plan(registry, store, arguments.artifact_root) + except GovernanceError as error: + print(json.dumps({"error": str(error)}, indent=2)) + return 1 + print(json.dumps([action.model_dump(mode="json") for action in actions], indent=2)) + return 0 + + +if __name__ == "__main__": # pragma: no cover - command-line entry point + raise SystemExit(main()) diff --git a/tests/integration/test_governance_mlflow.py b/tests/integration/test_governance_mlflow.py new file mode 100644 index 0000000..62721e1 --- /dev/null +++ b/tests/integration/test_governance_mlflow.py @@ -0,0 +1,200 @@ +"""Governance against a real MLflow tracking store and model registry. + +The store is a SQLite file in a temporary directory, so this exercises the +MLflow client itself rather than a stand-in for it. +""" + +import json +import shutil +from pathlib import Path +from typing import Any + +import pytest + +from llm_router.governance import ( + BENCHMARK_TAG, + STAGE_TAG, + MlflowStore, + main, + plan, + sync, + verify, +) +from llm_router.registry import LifecycleStage, Registry, load_registry + +pytest.importorskip("mlflow") +pytestmark = pytest.mark.integration + +CATALOG = load_registry("config/registry.yaml") +STAGED = "claims-extraction-lora-next" + + +@pytest.fixture(scope="module") +def migrated_database(tmp_path_factory: pytest.TempPathFactory) -> Path: + """An empty MLflow database with the schema applied, built once.""" + + path = tmp_path_factory.mktemp("mlflow") / "template.db" + MlflowStore.from_uri(f"sqlite:///{path.as_posix()}").names() + return path + + +@pytest.fixture +def tracking_uri(tmp_path: Path, migrated_database: Path) -> str: + database = tmp_path / "mlflow.db" + shutil.copyfile(migrated_database, database) + return f"sqlite:///{database.as_posix()}" + + +@pytest.fixture +def store(tracking_uri: str) -> MlflowStore: + return MlflowStore.from_uri(tracking_uri) + + +def production_alias(store: MlflowStore, name: str) -> Any: + return store.client.get_model_version_by_alias(name, "production") + + +def test_a_sync_records_the_catalog_and_leaves_nothing_to_do(store: MlflowStore) -> None: + actions = sync(CATALOG, store, commit="abc123") + + assert len(actions) == len(CATALOG.models) + len(CATALOG.adapters) + len(CATALOG.benchmarks) + assert plan(CATALOG, store) == () + assert verify(CATALOG, store) == () + assert store.names() == {item.id for item in (*CATALOG.models, *CATALOG.adapters)} + + served = production_alias(store, "small-specialist") + assert served.tags["catalog.revision"] == "mock-small@sha256:dev" + assert served.tags["catalog.checksum"] == "sha256:mock-small" + assert served.tags["catalog.license"] == "apache-2.0" + assert served.source.endswith("/models/small-specialist/mock-small-sha256-dev") + staged = store.client.get_model_version_by_alias(STAGED, "staging") + assert staged.tags[STAGE_TAG] == "staging" + + +def test_benchmark_evidence_is_recorded_once_with_its_measurements(store: MlflowStore) -> None: + sync(CATALOG, store) + sync(CATALOG, store) + + experiment = store.client.get_experiment_by_name(store.benchmark_experiment) + runs = store.client.search_runs([experiment.experiment_id]) + assert [run.data.tags[BENCHMARK_TAG] for run in runs] == ["extraction-v3"] + assert runs[0].data.metrics["quality_score"] == 0.82 + assert runs[0].data.metrics["throughput_rps"] == 41.5 + assert runs[0].data.params["model_revision"] == "mock-small@sha256:dev" + assert runs[0].info.status == "FINISHED" + + +def test_a_new_revision_takes_the_production_alias_and_the_old_one_is_kept( + store: MlflowStore, +) -> None: + sync(CATALOG, store, commit="first") + replaced = CATALOG.model_copy( + update={ + "models": tuple( + card.model_copy(update={"revision": "mock-general@sha256:next"}) + if card.id == "general-local" + else card + for card in CATALOG.models + ) + } + ) + + sync(replaced, store, commit="second") + + assert production_alias(store, "general-local").tags["catalog.revision"] == ( + "mock-general@sha256:next" + ) + stages = {item.revision: item.stage for item in store.versions("general-local")} + assert stages["mock-general@sha256:dev"] is LifecycleStage.DEPRECATED + history = store.promotions("general-local") + assert [(item.revision, item.from_stage, item.to_stage, item.commit) for item in history] == [ + ("mock-general@sha256:dev", None, LifecycleStage.PRODUCTION, "first"), + ("mock-general@sha256:next", None, LifecycleStage.PRODUCTION, "second"), + ( + "mock-general@sha256:dev", + LifecycleStage.PRODUCTION, + LifecycleStage.DEPRECATED, + "second", + ), + ] + assert verify(replaced, store) == () + + +def promote(registry: Registry, adapter_id: str, stage: LifecycleStage) -> Registry: + return registry.model_copy( + update={ + "adapters": tuple( + item.model_copy(update={"stage": stage}) if item.id == adapter_id else item + for item in registry.adapters + ) + } + ) + + +def test_a_promotion_moves_the_alias_and_a_demotion_removes_it(store: MlflowStore) -> None: + from mlflow.exceptions import MlflowException + + sync(CATALOG, store) + + sync(promote(CATALOG, STAGED, LifecycleStage.PRODUCTION), store) + assert production_alias(store, STAGED).tags[STAGE_TAG] == "production" + with pytest.raises(MlflowException): + store.client.get_model_version_by_alias(STAGED, "staging") + + sync(promote(CATALOG, STAGED, LifecycleStage.DEPRECATED), store) + with pytest.raises(MlflowException): + production_alias(store, STAGED) + assert [item.to_stage for item in store.promotions(STAGED)] == [ + LifecycleStage.STAGING, + LifecycleStage.PRODUCTION, + LifecycleStage.DEPRECATED, + ] + + +def test_a_stage_edited_by_hand_in_mlflow_is_reported_and_repaired(store: MlflowStore) -> None: + sync(CATALOG, store) + version = production_alias(store, "high-capability").version + store.client.set_model_version_tag("high-capability", version, STAGE_TAG, "blessed") + + drift = verify(CATALOG, store) + assert len(drift) == 1 and "recorded as unstaged, expected production" in drift[0] + + sync(CATALOG, store) + assert verify(CATALOG, store) == () + + +def test_changed_card_metadata_replaces_the_recorded_tags(store: MlflowStore) -> None: + sync(CATALOG, store) + without_checksums = CATALOG.model_copy(update={"deployments": ()}) + + assert len(verify(without_checksums, store)) == 3 + sync(without_checksums, store) + + assert "catalog.checksum" not in production_alias(store, "small-specialist").tags + assert verify(without_checksums, store) == () + + +def test_models_registered_outside_the_catalog_are_left_alone(store: MlflowStore) -> None: + store.client.create_registered_model("someone-elses-experiment") + store.client.create_model_version("someone-elses-experiment", source="s3://elsewhere/model") + + sync(CATALOG, store) + + assert "someone-elses-experiment" not in store.names() + assert verify(CATALOG, store) == () + + +def test_the_command_line_fails_a_pipeline_on_drift( + tracking_uri: str, capsys: pytest.CaptureFixture[str] +) -> None: + assert main(["verify", "--tracking-uri", tracking_uri]) == 3 + assert json.loads(capsys.readouterr().out)["in_step"] is False + + assert main(["sync", "--tracking-uri", tracking_uri, "--commit", "abc123"]) == 0 + capsys.readouterr() + assert main(["verify", "--tracking-uri", tracking_uri]) == 0 + assert json.loads(capsys.readouterr().out) == {"in_step": True, "drift": []} + + assert main(["history", "--tracking-uri", tracking_uri, "--name", "general-local"]) == 0 + history = json.loads(capsys.readouterr().out) + assert [(item["to_stage"], item["commit"]) for item in history] == [("production", "abc123")] diff --git a/tests/unit/test_deployment_manifests.py b/tests/unit/test_deployment_manifests.py index b379884..a676160 100644 --- a/tests/unit/test_deployment_manifests.py +++ b/tests/unit/test_deployment_manifests.py @@ -46,6 +46,14 @@ def load_documents() -> list[dict[str, Any]]: ] +def external_secret_named(name: str) -> dict[str, Any]: + return next( + document + for document in DOCUMENTS + if document["kind"] == "ExternalSecret" and document["metadata"]["name"] == name + ) + + def pod_spec(workload: dict[str, Any]) -> dict[str, Any]: spec: dict[str, Any] = workload["spec"]["template"]["spec"] return spec @@ -104,9 +112,7 @@ def test_no_manifest_contains_secret_material() -> None: def test_gateway_reads_credentials_from_the_secret_manager() -> None: - external_secret = next( - document for document in DOCUMENTS if document["kind"] == "ExternalSecret" - ) + external_secret = external_secret_named("llm-gateway-credentials") gateway = next( document for document in WORKLOADS if document["metadata"]["name"] == "llm-gateway" ) @@ -248,9 +254,7 @@ def test_only_the_gateway_reaches_the_proxy_and_only_the_proxy_reaches_out() -> def test_the_proxy_reads_both_of_its_keys_from_the_secret_manager() -> None: - external_secret = next( - document for document in DOCUMENTS if document["kind"] == "ExternalSecret" - ) + external_secret = external_secret_named("llm-gateway-credentials") proxy = next( document for document in WORKLOADS if document["metadata"]["name"] == "litellm-proxy" ) @@ -259,3 +263,68 @@ def test_the_proxy_reads_both_of_its_keys_from_the_secret_manager() -> None: for variable in pod_spec(proxy)["containers"][0]["env"]: assert "value" not in variable assert variable["valueFrom"]["secretKeyRef"]["key"] in provided + + +def test_the_governance_store_is_reachable_only_by_the_delivery_pipeline() -> None: + policy = next( + document["spec"] + for document in DOCUMENTS + if document["kind"] == "NetworkPolicy" and document["metadata"]["name"] == "mlflow" + ) + + assert policy["ingress"] == [ + { + "from": [ + { + "namespaceSelector": { + "matchLabels": {"kubernetes.io/metadata.name": "platform-delivery"} + } + } + ], + "ports": [{"protocol": "TCP", "port": 5000}], + } + ] + reachable = { + rule["namespaceSelector"]["matchLabels"]["kubernetes.io/metadata.name"] + for entry in policy["egress"] + for rule in entry["to"] + } + assert reachable == {"kube-system", "storage"} + + +def test_the_governance_store_keeps_its_own_credentials_apart_from_the_gateway() -> None: + mlflow = next(document for document in WORKLOADS if document["metadata"]["name"] == "mlflow") + gateway = next( + document for document in WORKLOADS if document["metadata"]["name"] == "llm-gateway" + ) + provided = { + item["secretKey"] for item in external_secret_named("mlflow-credentials")["spec"]["data"] + } + + referenced = { + variable["valueFrom"]["secretKeyRef"]["key"] + for variable in pod_spec(mlflow)["containers"][0]["env"] + if "valueFrom" in variable + } + assert referenced == provided + # The database address embeds a password, so it may never be a literal. + backend = next( + variable + for variable in pod_spec(mlflow)["containers"][0]["env"] + if variable["name"] == "MLFLOW_BACKEND_STORE_URI" + ) + assert "value" not in backend + gateway_secrets = { + variable["valueFrom"]["secretKeyRef"]["name"] + for variable in pod_spec(gateway)["containers"][0]["env"] + if "valueFrom" in variable + } + assert gateway_secrets == {"llm-gateway-credentials"} + + +def test_artifacts_are_proxied_so_clients_never_hold_storage_credentials() -> None: + mlflow = next(document for document in WORKLOADS if document["metadata"]["name"] == "mlflow") + arguments = pod_spec(mlflow)["containers"][0]["args"] + + assert "--serve-artifacts" in arguments + assert any(item.startswith("--artifacts-destination=s3://") for item in arguments) diff --git a/tests/unit/test_governance.py b/tests/unit/test_governance.py new file mode 100644 index 0000000..590b703 --- /dev/null +++ b/tests/unit/test_governance.py @@ -0,0 +1,251 @@ +import json + +import pytest + +from llm_router.governance import ( + ActionKind, + GovernanceError, + ImmutableRevisionError, + InMemoryStore, + MlflowStore, + benchmark_record, + governed_versions, + main, + plan, + sync, + verify, +) +from llm_router.registry import LifecycleStage, Registry, load_registry + +CATALOG = load_registry("config/registry.yaml") +STAGED = "claims-extraction-lora-next" + + +def with_model(registry: Registry, model_id: str, **changes: object) -> Registry: + return registry.model_copy( + update={ + "models": tuple( + card.model_copy(update=changes) if card.id == model_id else card + for card in registry.models + ) + } + ) + + +def with_adapter(registry: Registry, adapter_id: str, **changes: object) -> Registry: + return registry.model_copy( + update={ + "adapters": tuple( + item.model_copy(update=changes) if item.id == adapter_id else item + for item in registry.adapters + ) + } + ) + + +def synced() -> InMemoryStore: + store = InMemoryStore() + sync(CATALOG, store, commit="first") + return store + + +def test_a_first_plan_records_every_model_adapter_and_benchmark() -> None: + actions = plan(CATALOG, InMemoryStore()) + + registered = {action.name for action in actions if action.kind is ActionKind.REGISTER} + benchmarks = {action.name for action in actions if action.kind is ActionKind.BENCHMARK} + assert registered == {card.id for card in CATALOG.models} | { + adapter.id for adapter in CATALOG.adapters + } + assert benchmarks == {run.id for run in CATALOG.benchmarks} + assert {action.kind for action in actions} == {ActionKind.REGISTER, ActionKind.BENCHMARK} + + +def test_a_synced_store_is_in_step_and_syncing_again_changes_nothing() -> None: + store = synced() + + assert plan(CATALOG, store) == () + assert verify(CATALOG, store) == () + assert sync(CATALOG, store) == () + # One promotion per registration, and no more after the second sync. + assert len(store.promotions("general-local")) == 1 + + +def test_records_carry_the_card_the_artifact_location_and_the_checksum() -> None: + versions = {item.name: item for item in governed_versions(CATALOG, "s3://bucket/")} + + local = versions["small-specialist"] + assert local.source == "s3://bucket/models/small-specialist/mock-small-sha256-dev" + assert local.checksum == "sha256:mock-small" + assert local.tags["catalog.license"] == "apache-2.0" + assert local.tags["catalog.hardware"] == "1x nvidia-l4, 24 GB, tensor parallel 1" + assert "Not evaluated for open-ended reasoning" in local.tags["catalog.limitations"] + # Nothing of ours to store for a provider model, only the alias that names it. + assert versions["approved-external-fallback"].source == "litellm://approved-external-fallback" + adapter = versions["claims-extraction-lora"] + assert adapter.source == "s3://bucket/adapters/claims-extraction-lora/claims-lora-sha256-dev" + assert adapter.tags["catalog.base_revision"] == "mock-small@sha256:dev" + assert adapter.tags["catalog.quality_delta"] == "0.06" + assert adapter.checksum == "sha256:mock-claims" + + +def test_a_promotion_in_the_catalog_is_applied_and_added_to_the_history() -> None: + store = synced() + promoted = with_adapter(CATALOG, STAGED, stage=LifecycleStage.PRODUCTION) + + actions = sync(promoted, store, commit="second") + + assert [(item.kind, item.name) for item in actions] == [(ActionKind.TRANSITION, STAGED)] + history = store.promotions(STAGED) + assert [(item.from_stage, item.to_stage, item.commit) for item in history] == [ + (None, LifecycleStage.STAGING, "first"), + (LifecycleStage.STAGING, LifecycleStage.PRODUCTION, "second"), + ] + assert history[-1].policy_version == CATALOG.policy.version + assert verify(promoted, store) == () + + +def test_a_new_revision_is_registered_and_the_one_it_replaces_is_kept_but_retired() -> None: + store = synced() + replaced = with_model(CATALOG, "general-local", revision="mock-general@sha256:next") + + actions = sync(replaced, store, commit="second") + + assert [(item.kind, item.revision, item.to_stage) for item in actions] == [ + (ActionKind.REGISTER, "mock-general@sha256:next", LifecycleStage.PRODUCTION), + (ActionKind.TRANSITION, "mock-general@sha256:dev", LifecycleStage.DEPRECATED), + ] + stages = {item.revision: item.stage for item in store.versions("general-local")} + assert stages == { + "mock-general@sha256:dev": LifecycleStage.DEPRECATED, + "mock-general@sha256:next": LifecycleStage.PRODUCTION, + } + assert store.promotions("general-local")[-1].reason == "superseded in the catalog" + + +def test_a_subject_removed_from_the_catalog_is_retired_not_deleted() -> None: + store = synced() + removed = CATALOG.model_copy( + update={"adapters": tuple(item for item in CATALOG.adapters if item.id != STAGED)} + ) + + actions = sync(removed, store) + + assert [(item.name, item.to_stage, item.reason) for item in actions] == [ + (STAGED, LifecycleStage.DEPRECATED, "removed from the catalog") + ] + assert [item.stage for item in store.versions(STAGED)] == [LifecycleStage.DEPRECATED] + assert sync(removed, store) == () + + +def test_a_recorded_revision_cannot_take_a_different_artifact() -> None: + store = synced() + deployments = tuple( + item.model_copy(update={"model_checksums": {"small-specialist": "sha256:swapped"}}) + if item.stage is LifecycleStage.PRODUCTION + else item + for item in CATALOG.deployments + ) + swapped = CATALOG.model_copy(update={"deployments": deployments}) + + with pytest.raises(ImmutableRevisionError, match="needs a new revision"): + plan(swapped, store) + drift = verify(swapped, store) + assert len(drift) == 1 and "sha256:swapped" in drift[0] + # Nothing was applied on the way to the refusal. + assert store.versions("small-specialist")[0].checksum == "sha256:mock-small" + + +def test_changed_card_metadata_is_updated_without_recording_a_promotion() -> None: + store = synced() + reworded = with_model(CATALOG, "general-local", limitations="Not for legal advice.") + + assert "out-of-date card metadata" in verify(reworded, store)[0] + actions = sync(reworded, store) + + assert [item.kind for item in actions] == [ActionKind.ANNOTATE] + recorded = store.versions("general-local")[0] + assert recorded.tags["catalog.limitations"] == "Not for legal advice." + assert len(store.promotions("general-local")) == 1 + + +def test_a_stage_edited_in_the_store_is_drift_and_a_sync_puts_it_back() -> None: + store = synced() + store.set_stage("high-capability", "mock-high@sha256:dev", LifecycleStage.STAGING) + + drift = verify(CATALOG, store) + assert drift == ( + "high-capability@mock-high@sha256:dev is recorded as staging, expected production " + "(stage differs from the catalog)", + ) + + sync(CATALOG, store, commit="repair") + assert verify(CATALOG, store) == () + assert store.promotions("high-capability")[-1].commit == "repair" + + +def test_models_and_adapters_may_not_share_a_name() -> None: + clash = with_adapter(CATALOG, STAGED, id="general-local") + + with pytest.raises(GovernanceError, match="general-local"): + governed_versions(clash) + + +def test_a_benchmark_is_split_into_what_identifies_it_and_what_it_measured() -> None: + run = CATALOG.benchmarks[0].model_copy( + update={"gpu_memory_gb": 21.5, "draft_acceptance_rate": 0.7} + ) + + params, metrics = benchmark_record(run) + + assert params["hardware"] == "nvidia-l4" and params["task"] == "extraction" + assert json.loads(params["engine_settings"])["max_num_seqs"] == 64 + assert metrics["quality_score"] == 0.82 and metrics["latency_p95_ms"] == 480 + assert metrics["gpu_memory_gb"] == 21.5 and metrics["draft_acceptance_rate"] == 0.7 + assert "gpu_memory_gb" not in benchmark_record(CATALOG.benchmarks[0])[1] + + +def test_the_offline_plan_shows_what_a_first_sync_would_record( + capsys: pytest.CaptureFixture[str], +) -> None: + assert main(["plan"]) == 0 + + printed = json.loads(capsys.readouterr().out) + assert {item["kind"] for item in printed} == {"register", "benchmark"} + assert {"name": "general-local", "to_stage": "production"}.items() <= next( + item for item in printed if item["name"] == "general-local" + ).items() + + +@pytest.mark.parametrize( + "arguments", [["sync"], ["verify"], ["history", "--tracking-uri", "sqlite:///unused.db"]] +) +def test_commands_that_need_a_store_or_a_subject_say_so( + arguments: list[str], monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.delenv("MLFLOW_TRACKING_URI", raising=False) + + with pytest.raises(SystemExit) as raised: + main(arguments) + assert raised.value.code == 2 + + +def test_the_command_line_reports_a_refused_plan( + monkeypatch: pytest.MonkeyPatch, capsys: pytest.CaptureFixture[str] +) -> None: + def refuse(*_: object, **__: object) -> None: + raise ImmutableRevisionError("a changed artifact needs a new revision") + + monkeypatch.setattr("llm_router.governance.plan", refuse) + + assert main(["plan"]) == 1 + assert "needs a new revision" in json.loads(capsys.readouterr().out)["error"] + + +def test_a_version_that_is_not_recorded_cannot_be_staged() -> None: + class Empty: + def search_model_versions(self, _: str) -> list[object]: + return [] + + with pytest.raises(GovernanceError, match="ghost@rev is not recorded"): + MlflowStore(Empty()).set_stage("ghost", "rev", LifecycleStage.PRODUCTION)