diff --git a/docs/infrastructure/workload-credentials.md b/docs/infrastructure/workload-credentials.md index a2ab49f283..2b011e8c0f 100644 --- a/docs/infrastructure/workload-credentials.md +++ b/docs/infrastructure/workload-credentials.md @@ -361,6 +361,87 @@ The verifier has no shared-Valkey cache. JWKS is trusted verification material: the general Valkey is writable by multiple services, so accepting keys from it would allow those writers to forge workload identities. +## Middleman inference + +Middleman accepts version-2 Hawk workload JWTs on Anthropic messages/token +counting, OpenAI-compatible chat/raw completions/foreground Responses/input-token +counting/compaction, Gemini generate/stream/count operations, and filtered +OpenAI model listing. Each call requires the exact `middleman:call@models/...` +atom and the saved public-name/resolved-name/group/lab approval for the actual +model dispatched. For secret models, the saved resolved name is a versioned +HMAC-SHA256 binding to the effective public name, upstream target, group, and +lab; it does not expose the upstream target. Provision a stable +`MIDDLEMAN_MODEL_APPROVAL_BINDING_KEY` containing 32 cryptographically random +bytes encoded as exactly 64 hexadecimal characters in the existing provider +secret store, and share it across Middleman replicas and restarts. Missing or +malformed keys prevent secret-model approvals. Changing the key or bound +model identity requires fresh approvals; renewal does not update the saved +binding. See [Secret model approval bindings](middleman.md#secret-model-approval-bindings). + +Legacy unified inference, other catalog/query/control/admin endpoints, and file +uploads remain human-only. Workload Responses require `background` to be absent +or the JSON boolean `false`; other values fail before provider work. Batch +execution and background retrieval/polling/cancellation are unsupported. +Inline images/screenshots remain supported; excluding uploads does not remove +these inference-body formats. +Inspect's Google video/audio/document Files API workflow is not exposed by +Middleman. Provider configurations that require background execution cannot use +workload credentials. There is no automatic conversion or human-credential fallback. + +`MIDDLEMAN_WORKLOAD_AUTH` is the shared validation JSON object with `issuer`, +`audience`, and `jwks_uri`. Its audience is `/workload/services`. +Absent configuration leaves the deployment human-only. Invalid settings or a +workload issuer also configured as a human issuer fail startup. Middleman owns +one local JWKS cache and HTTP client per worker, fetches lazily, and never uses +general shared Valkey for verification keys. A key-fetch outage is a retryable +provider-shaped 503 with `Retry-After: 5`; invalid credentials are authentication +errors. Expiry is checked on admission, including the shared 60-second skew +allowance; an admitted stream is allowed to finish. + +Traffic audit fields retain the signed execution subject in `user_id` and the +execution/grant/job identifiers, with no human groups, admin privileges, or +inferred team membership. Signed usage attribution is recorded separately in +`usage_user_id` and `usage_teams` and supplies the user and team dimensions for +usage metrics. Verified no-team membership is `[]`; missing or null attribution +or teams is rejected during authentication. +Priority admission uses that accounting user, so one human's direct requests +and executions share a flow within each priority class and quota group. +Accounting metadata grants no permissions and does not change trace identity. +The mounted workload key cannot authenticate inference. Middleman receives no +signing permission or issuance/completion credential. + +Model-call approvals do not grant transcript access. Workloads read approved +inputs and write their own outputs through brokered S3 sessions. Human readers +need the inherited model and code permissions described under +[data sensitivity](#data-sensitivity-and-credential-lifetime). Hawk's human-facing +data APIs, including warehouse queries and presigning, reject workload JWTs. + +Clients must supply the current JWT for each request and use supported +foreground or streaming modes. See +[credential lifetime](#data-sensitivity-and-credential-lifetime) for renewal +behavior. Stopping new workload launches must preserve verification for existing +runs. + +```json +{ + "issuer": "https://api.example/workload", + "audience": "https://api.example/workload/services", + "jwks_uri": "https://api.example/.well-known/jwks.json" +} +``` + +| Route | Workload behavior | +| --- | --- | +| `/anthropic/v1/messages`, `/anthropic/v1/messages/count_tokens` | Exact authorized model | +| `/openai/v1/chat/completions`, `/openai/v1/completions` | Exact authorized model | +| `/openai/v1/responses` | Exact authorized model; foreground/streaming only | +| `/openai/v1/responses/input_tokens`, `/openai/v1/responses/compact` | Exact authorized model | +| `/gemini/v1beta1/publishers/{publisher}/models/{model}:{operation}` | Exact authorized model; generate/stream/count only | +| `/google-ai/{api_version}/models/{model}:{operation}` | Exact authorized model; generate/stream/count only | +| `/openai/v1/models` | Only live models matching saved approvals | +| `/openai/v1/files`; legacy unified/catalog/query/control/admin routes | Human-only | +| Provider batch/background follow-up routes | Not exposed | + ## Database authority The `hawk_api` permission role is `NOLOGIN`. It has `SELECT` and `INSERT` on the diff --git a/docs/user-guide/token-usage.md b/docs/user-guide/token-usage.md index c52163bc79..551f5b3edc 100644 --- a/docs/user-guide/token-usage.md +++ b/docs/user-guide/token-usage.md @@ -125,16 +125,18 @@ to a label). | Label | Grouping | Meaning | | --- | --- | --- | -| `unassigned` | team | The caller's token carried no team claim | +| `unassigned` | team | The caller's teams are empty or unknown | | `a+b` | team | The caller is in several teams; sorted and `+`-joined, counted once | -| `direct` | job, task | Traffic with no `x-hawk-job-id` header: coding agents, notebooks, scripts, anything not launched by Hawk | +| `direct` | job, task | Traffic with neither a signed workload job id nor an `x-hawk-job-id` header: coding agents, notebooks, scripts, anything not launched by Hawk | | `scan` | task | A scan or scan-resume job (scans have no Inspect task) | | `unknown` | task | A job that sent no task name, e.g. runners predating the `x-inspect-task-name` header | | `other` | job, task | Everything past the 50 highest-token labels for that model, plus any job that also used a model you may not see, summed | Job labels are the eval set id or scan id, the same value shown as the job id in -the viewer and by `hawk eval-set` / `hawk scan`. Task labels are the Inspect task -name (`Task.name`, e.g. `gaia`). A job that used any model outside your model +the viewer and by `hawk eval-set` / `hawk scan`. Signed workload job ids and types +take precedence over correlation headers; other traffic falls back to those +headers. Task labels are the Inspect task name (`Task.name`, e.g. `gaia`). +A job that used any model outside your model groups is not named: its tokens on the models you can see are counted under `other`, the same rule that hides the run itself. @@ -147,11 +149,25 @@ models offered in `/usage/history` are discovered from CloudWatch, which only lists metrics active in the last ~2 weeks, so a model idle for longer drops out of the history view even if the metric data itself is retained. +Workload tokens require signed usage attribution and charge the submitting user's +id and teams, so runner usage appears in that user's `hawk usage` totals. Traffic +logs keep the authenticated execution subject in `user_id`, with signed +accounting metadata in `usage_user_id` and `usage_teams`. Human requests keep +their existing `user_id` and `user_teams`. New workload tokens require explicit +teams, with `[]` for verified no-team membership. General and historical logs +can still have unknown teams (`null`); both map to `unassigned` in team metrics. Existing +CloudWatch usage series recorded under execution subjects remain unchanged; +this change cannot relabel historical metrics to their submitting users. + **Team member history API** (`group_by=user&team=…`) uses traffic logs with the same 31-day limit and top-50/`other` behavior as job/task history. It matches exact recorded membership: `a+b` differs from `a`; `unassigned` means an explicitly empty team list. Member totals use the same provider token accounting as the team overview; log coverage and recorded membership can still differ. +User grouping prefers `usage_user_id`, falling back to `user_id` for older logs +and human requests, then `unknown`. When `usage_user_id` is present and nonempty, +team filtering uses `usage_teams`; otherwise it uses the legacy `user_teams`. +Unknown or missing membership never matches the empty list for `unassigned`. The API returns recorded per-member cache counts and the applied `team_filter`. The usage page uses this response to identify members, then shows their ordinary user metrics for the selected window, including usage outside that recorded team. diff --git a/hawk/hawk/core/rate_limits/traffic_log_usage.py b/hawk/hawk/core/rate_limits/traffic_log_usage.py index ba0be02a57..69aaa4ef41 100644 --- a/hawk/hawk/core/rate_limits/traffic_log_usage.py +++ b/hawk/hawk/core/rate_limits/traffic_log_usage.py @@ -44,19 +44,19 @@ _QUERY_TIMEOUT_S = 60.0 LIVE_ALIGN_SECONDS = 60 -# OpenAI totals now include caches; do not reuse rows with the old token basis. -_CACHE_PREFIX = "hawk:usage:traffic:v4" +# Signed usage attribution takes precedence over legacy authentication fields. +_CACHE_PREFIX = "hawk:usage:traffic:v6" _CACHE_TTL_RECENT_S = 60 _CACHE_TTL_SETTLED_S = 15 * 60 _SETTLED_AFTER_S = 60 * 60 _LABEL_FIELDS = { - "user": 'coalesce(user_id, "unknown") as label', - "job": 'coalesce(`correlation.x-hawk-job-id`, "direct") as label', + "user": 'coalesce(usage_user_id, user_id, "unknown") as label', + "job": 'coalesce(workload_job_id, `correlation.x-hawk-job-id`, "direct") as label', "task": ( 'coalesce(`correlation.x-inspect-task-name`, "") as label,' - " `correlation.x-hawk-job-id` as job_id," - " `correlation.x-hawk-job-type` as job_type" + " coalesce(workload_job_id, `correlation.x-hawk-job-id`) as job_id," + " coalesce(workload_job_type, `correlation.x-hawk-job-type`) as job_type" ), } _JOB_COLUMNS = {"job": [], "task": ["job_id", "job_type"], "user": []} @@ -145,7 +145,9 @@ def build_query( team_json = json.dumps(teams, ensure_ascii=False, separators=(",", ":")) parts += [ "| fields jsonParse(@message) as traffic", - f"| filter jsonStringify(traffic.user_teams) = {json.dumps(team_json, ensure_ascii=False)}", + '| fields if(strlen(coalesce(usage_user_id, "")) > 0, ' + + "jsonStringify(traffic.usage_teams), jsonStringify(traffic.user_teams)) as recorded_teams", + f"| filter recorded_teams = {json.dumps(team_json, ensure_ascii=False)}", ] if period is not None: # Arithmetic rather than `bin()`, whose unit caps reject e.g. 90-minute periods. diff --git a/hawk/tests/core/rate_limits/test_traffic_log_usage.py b/hawk/tests/core/rate_limits/test_traffic_log_usage.py index 2a4e5f7c1e..3ce2889f35 100644 --- a/hawk/tests/core/rate_limits/test_traffic_log_usage.py +++ b/hawk/tests/core/rate_limits/test_traffic_log_usage.py @@ -134,14 +134,23 @@ class TestBuildQuery: def test_live_query_groups_by_label_without_bins(self): query = tl.build_query("job", start=1000, period=None) assert "bin_idx" not in query - assert 'coalesce(`correlation.x-hawk-job-id`, "direct") as label' in query + assert ( + 'coalesce(workload_job_id, `correlation.x-hawk-job-id`, "direct") as label' + in query + ) assert "by provider, public_name, label" in query assert f"limit {tl.ROW_LIMIT}" in query def test_task_query_carries_the_job_alongside_the_raw_task_name(self): query = tl.build_query("task", start=0, period=None) assert 'coalesce(`correlation.x-inspect-task-name`, "") as label' in query - assert "`correlation.x-hawk-job-id` as job_id" in query + assert ( + "coalesce(workload_job_id, `correlation.x-hawk-job-id`) as job_id" in query + ) + assert ( + "coalesce(workload_job_type, `correlation.x-hawk-job-type`) as job_type" + in query + ) assert "by provider, public_name, label, job_id, job_type" in query def test_history_query_bins_relative_to_start(self): @@ -601,7 +610,7 @@ def as_client(self) -> redis.asyncio.Redis: class TestFetchRows: async def test_caches_rows_and_serves_repeat_from_cache(self): redis_ = _FakeRedis() - redis_.store["hawk:usage:traffic:v3:g:job:0:300:0"] = json.dumps( + redis_.store["hawk:usage:traffic:v5:g:job:0:300:0"] = json.dumps( [ { "bin_idx": 0, @@ -641,7 +650,7 @@ async def test_caches_rows_and_serves_repeat_from_cache(self): 2, ) client.start_query.assert_awaited_once() - assert "hawk:usage:traffic:v4:g:job:0:300:0" in redis_.store + assert "hawk:usage:traffic:v6:g:job:0:300:0" in redis_.store @pytest.mark.parametrize( ("end", "now", "expected_ttl"), @@ -1012,14 +1021,19 @@ async def test_team_history_scopes_every_query_and_returns_people( assert client.start_query.await_count == 3 for call in client.start_query.await_args_list: query = call.kwargs["queryString"] - assert 'coalesce(user_id, "unknown") as label' in query + assert 'coalesce(usage_user_id, user_id, "unknown") as label' in query + assert ( + '| fields if(strlen(coalesce(usage_user_id, "")) > 0, ' + "jsonStringify(traffic.usage_teams), jsonStringify(traffic.user_teams)) as recorded_teams" + ) in query (team_filter,) = [ line for line in query.splitlines() - if "jsonStringify(traffic.user_teams)" in line + if line.startswith("| filter recorded_teams = ") ] assert json.loads(json.loads(team_filter.split(" = ", 1)[1])) == memberships assert "coalesce(traffic.user_teams" not in query + assert "coalesce(traffic.usage_teams" not in query async def test_bins_carry_labels_and_model_series_over_the_aligned_range(self): client = _history_client( diff --git a/infra/app.py b/infra/app.py index 8a17eaeabe..f886d62e8f 100644 --- a/infra/app.py +++ b/infra/app.py @@ -314,6 +314,7 @@ def deploy( task_cpu=config.middleman_task_cpu, task_memory=config.middleman_task_memory, enable_datadog=config.enable_datadog, + workload_validation=hawk.workload_validation, oidc_issuer=hawk.oidc_issuer, oidc_audience=hawk.oidc_audience, oidc_jwks_uri=hawk.oidc_jwks_uri, diff --git a/infra/core/middleman.py b/infra/core/middleman.py index 49f24e8a38..6792be0a23 100644 --- a/infra/core/middleman.py +++ b/infra/core/middleman.py @@ -85,6 +85,7 @@ def __init__( dd_api_key_secret_arn: pulumi.Input[str] | None = None, api_keys_secret_arn: pulumi.Input[str] | None = None, enable_datadog: bool = False, + workload_validation: pulumi.Input[dict[str, Any]] | None = None, oidc_issuer: pulumi.Input[str] = "", oidc_audience: pulumi.Input[str] = "", oidc_jwks_uri: pulumi.Input[str] = "", @@ -695,6 +696,9 @@ def _build_auth_providers_json(args: AuthProviderArgs) -> str: assert dd_api_key_secret_arn is not None task_def_inputs["dd_api_key_arn"] = dd_api_key_secret_arn + if workload_validation is not None: + task_def_inputs["workload_auth_json"] = pulumi.Output.from_input(workload_validation).apply(json.dumps) + task_def_inputs["traffic_log_bucket"] = self.traffic_log.bucket_name task_def_inputs["traffic_log_group"] = self.traffic_log.log_group_name task_def_inputs["metrics_log_group"] = metrics_log_group.name @@ -718,6 +722,8 @@ def _build_container_defs(args: dict[str, Any]) -> str: ] if args.get("auth_providers_json"): app_env.append({"name": "MIDDLEMAN_AUTH_PROVIDERS", "value": args["auth_providers_json"]}) + if args.get("workload_auth_json"): + app_env.append({"name": "MIDDLEMAN_WORKLOAD_AUTH", "value": args["workload_auth_json"]}) if args.get("anthropic_profiles_json"): # Per-model-group override of Anthropic auth. See middleman/src/middleman/credential_broker.py. app_env.append({"name": "MIDDLEMAN_ANTHROPIC_PROFILES", "value": args["anthropic_profiles_json"]}) diff --git a/infra/hawk/__init__.py b/infra/hawk/__init__.py index f907443875..2f6f196653 100644 --- a/infra/hawk/__init__.py +++ b/infra/hawk/__init__.py @@ -116,6 +116,8 @@ def _staging_plugin_image(images: dict[str, str] | None, viewer_id: str) -> str: class HawkStack(pulumi.ComponentResource): """Hawk platform: API, Lambdas, EventBridge, and Batch.""" + workload_validation: pulumi.Output[dict[str, Any]] + def __init__( self, name: str, @@ -518,6 +520,7 @@ def _resolve_job_token_issuer(url: str, require_job_token: bool) -> str: api_url=f"{'http' if config.skip_tls_certs else 'https'}://api{hawk_slug}.{hawk_base}", opts=child, ) + self.workload_validation = workload_identity.validation_settings hawk_api = HawkApi( "api", env=env, diff --git a/infra/hawk/workload_identity.py b/infra/hawk/workload_identity.py index 248d2f010a..80ce949e9d 100644 --- a/infra/hawk/workload_identity.py +++ b/infra/hawk/workload_identity.py @@ -51,6 +51,7 @@ def encode(value: int) -> str: class WorkloadIdentity(pulumi.ComponentResource): public_keys: pulumi.Output[dict[str, Any]] issuer_settings: pulumi.Output[dict[str, Any]] + validation_settings: pulumi.Output[dict[str, Any]] def __init__( self, @@ -101,14 +102,18 @@ def __init__( issuer = f"{api_url}/workload" # Match hawk.core.types.workload_auth.WorkloadIssuerSettings. Pulumi runs # from infra/, where this hawk package shadows the application package. + validation: dict[str, Any] = { + "issuer": issuer, + "audience": f"{issuer}/services", + "jwks_uri": f"{api_url}/.well-known/jwks.json", + } + self.validation_settings = pulumi.Output.from_input(validation) self.issuer_settings = self.public_keys.apply( lambda keys: { - "issuer": issuer, - "audience": f"{issuer}/services", - "jwks_uri": f"{api_url}/.well-known/jwks.json", + **validation, "kms_key_arn": keys[config.workload_active_signing_key_generation]["arn"], "kid": keys[config.workload_active_signing_key_generation]["kid"], "jwks": {"keys": [keys[generation]["jwk"] for generation in published]}, } ) - self.register_outputs({"public_keys": self.public_keys}) + self.register_outputs({"public_keys": self.public_keys, "validation_settings": self.validation_settings}) diff --git a/infra/tests/test_infra.py b/infra/tests/test_infra.py index 301baa28e2..fde2a0ba54 100644 --- a/infra/tests/test_infra.py +++ b/infra/tests/test_infra.py @@ -75,6 +75,13 @@ class _FakeHawkStack(pulumi.ComponentResource): def __init__(self, name: str, **_: Never) -> None: super().__init__("metr:hawk:HawkStack", name) + self.workload_validation: pulumi.Output[dict[str, Any]] = pulumi.Output.from_input( + { + "issuer": "https://api.example/workload", + "audience": "https://api.example/workload/services", + "jwks_uri": "https://api.example/.well-known/jwks.json", + } + ) def __getattr__(self, _: str) -> str: return "fake-hawk-output" @@ -253,6 +260,18 @@ def test_deploy_forwards_middleman_auth_settings(self, enabled: bool) -> None: ) ] assert len(sts_policies) == int(enabled) + task = next( + r + for r in mocks.created_resources + if r.typ == "aws:ecs/taskDefinition:TaskDefinition" and r.inputs["family"] == "staging-middleman" + ) + container = next(c for c in json.loads(task.inputs["containerDefinitions"]) if c["name"] == "middleman") + env = {entry["name"]: entry["value"] for entry in container["environment"]} + assert json.loads(env["MIDDLEMAN_WORKLOAD_AUTH"]) == { + "issuer": "https://api.example/workload", + "audience": "https://api.example/workload/services", + "jwks_uri": "https://api.example/.well-known/jwks.json", + } def test_deploy_exposes_its_configuration_and_source_inputs(self) -> None: """A consumer supplies the deployment configuration and optional Docker checkout.""" diff --git a/infra/tests/test_workload_identity.py b/infra/tests/test_workload_identity.py index 9a51156558..a8a45572fb 100644 --- a/infra/tests/test_workload_identity.py +++ b/infra/tests/test_workload_identity.py @@ -66,6 +66,8 @@ def test_generations_stay_managed_while_publication_and_activation_change( public_keys = _sync_await(identity.public_keys.future()) _sync_await(wait_for_rpcs()) assert settings is not None + validation = _sync_await(identity.validation_settings.future()) + assert validation == {key: settings[key] for key in ("issuer", "audience", "jwks_uri")} assert public_keys is not None keys = [r for r in mocks.created_resources if r.typ == "aws:kms/key:Key"] aliases = [r for r in mocks.created_resources if r.typ == "aws:kms/alias:Alias"] diff --git a/middleman/src/middleman/auth.py b/middleman/src/middleman/auth.py index 4d9b9056b7..da5097aeaa 100644 --- a/middleman/src/middleman/auth.py +++ b/middleman/src/middleman/auth.py @@ -17,14 +17,17 @@ import yaml from ddtrace.trace import tracer from fastapi import HTTPException +from hawk.core.types.workload_auth import WorkloadPrincipal from joserfc import jwk, jws, jwt from pydantic import BaseModel, model_validator +from middleman import workload_auth from middleman.observability.emf import team_name_error from middleman.observability.logging import get_logger from middleman.observability.metrics import record_auth_duration from middleman.request import get_client_session from middleman.traffic_log import context +from middleman.traffic_log.handle import TrafficLog ALGORITHMS = ["RS256"] ACCEPT_DEV_ADMIN = os.environ.get("MIDDLEMAN_ACCEPT_DEV_ADMIN", "false").lower() == "true" @@ -48,6 +51,39 @@ class UserInfo: teams: list[str] = field(default_factory=list) +type RequestPrincipal = UserInfo | WorkloadPrincipal + + +def principal_subject(principal: RequestPrincipal) -> str: + return principal.claims.sub if isinstance(principal, WorkloadPrincipal) else principal.id + + +def usage_attribution(principal: RequestPrincipal) -> tuple[str, list[str]]: + """Accounting identity; authorization and tracing retain the workload subject.""" + if isinstance(principal, WorkloadPrincipal): + attribution = principal.claims.workload.usage_attribution + return attribution.user_id, list(attribution.teams) + return principal.id, principal.teams + + +def record_principal(handle: TrafficLog, principal: RequestPrincipal) -> None: + if isinstance(principal, WorkloadPrincipal): + payload = principal.claims.workload + handle.set_user(principal.claims.sub, None, False) + user_id, teams = usage_attribution(principal) + handle.set_usage_attribution(user_id, teams) + handle.set_workload( + execution_id=payload.execution_id, + grant_id=payload.grant_id, + job_id=payload.job.id, + job_type=payload.job.type, + ) + else: + handle.set_user( + principal.id, principal.groups, principal.is_admin, teams=principal.teams, email=principal.email + ) + + def require_admin(user: UserInfo) -> None: """Raise 403 if user is not a middleman admin.""" if not user.is_admin: @@ -161,6 +197,8 @@ async def load_auth_providers() -> dict[str, _AuthProvider]: provider = _AuthProvider(**p) providers[provider.issuer] = provider + workload_auth.validate_human_issuers(providers) + return providers @@ -217,6 +255,8 @@ async def get_user_info(token: str) -> UserInfo: providers = await load_auth_providers() issuer: str = token_claims.get("iss", "") + if workload_auth.is_workload_issuer(issuer): + raise AuthError("invalid_token", "this endpoint requires human credentials") auth_provider = providers.get(issuer, None) if not auth_provider: logger.warning("auth.failed", reason="invalid_issuer", issuer=issuer) diff --git a/middleman/src/middleman/model_policy.py b/middleman/src/middleman/model_policy.py new file mode 100644 index 0000000000..e7e4a22d55 --- /dev/null +++ b/middleman/src/middleman/model_policy.py @@ -0,0 +1,49 @@ +from collections.abc import Iterable, Mapping + +from hawk.core.types.permission_atoms import format_atom +from hawk.core.types.workload_auth import WorkloadPrincipal +from pydantic import ValidationError + +from middleman.auth import RequestPrincipal +from middleman.model_approvals import ModelApprovalBindingUnavailableError, approval_for_model +from middleman.model_info import ModelInfo + + +class ModelNotAllowedError(Exception): + pass + + +def permitted_models(snapshot: Mapping[str, ModelInfo], principal: RequestPrincipal) -> dict[str, ModelInfo]: + if not isinstance(principal, WorkloadPrincipal): + groups = set(principal.groups) + return {name: model for name, model in snapshot.items() if model.group in groups} + payload = principal.claims.workload + permissions = set(payload.permissions) + approvals = {approval.public_name: approval for approval in payload.model_approvals} + permitted: dict[str, ModelInfo] = {} + for name, model in snapshot.items(): + saved = approvals.get(name) + if model.dead or saved is None or format_atom("model_call", name) not in permissions: + continue + try: + current = approval_for_model(model) + except (ValidationError, ModelApprovalBindingUnavailableError): + # Unverifiable models are denied without breaking the rest of the grant. + continue + if saved.model_dump() == current.model_dump(): + permitted[name] = model + return permitted + + +def select_models( + snapshot: Mapping[str, ModelInfo], + principal: RequestPrincipal, + names: Iterable[str], +) -> list[ModelInfo]: + permitted = permitted_models(snapshot, principal) + selected: list[ModelInfo] = [] + for name in names: + if name not in permitted: + raise ModelNotAllowedError("model not found") + selected.append(permitted[name]) + return selected diff --git a/middleman/src/middleman/observability/emf.py b/middleman/src/middleman/observability/emf.py index 555f25c52f..31cefb7700 100644 --- a/middleman/src/middleman/observability/emf.py +++ b/middleman/src/middleman/observability/emf.py @@ -7,14 +7,14 @@ stdout-EMF auto-extraction would never fire. Token metrics (input/output) are emitted on [provider, model], -[provider, model, user] (`user` = the caller's email from the auth provider's -`email_field` claim, falling back to the JWT sub for tokens without one or -with a non-ASCII one, which CloudWatch would reject as a dimension value), -[provider, model, channel] (`channel` = eval-set | scan | direct, from Hawk's -x-hawk-job-type correlation header — splits runner-driven usage from direct API -usage), [provider, model, team] (`team` = the user's IdP `teams` claim, -prefix-stripped and `+`-joined when the user is in several teams, so usage is -never double-counted; `unassigned` when absent), and +[provider, model, user] (`user` = signed workload accounting user, or the human +caller's ASCII email from the auth provider's `email_field` claim, falling back +to JWT sub), [provider, model, channel] (`channel` = eval-set | scan | direct, +from signed workload job identity or Hawk's x-hawk-job-type correlation header), +[provider, model, team] (`team` = signed workload accounting teams or the human +user's IdP `teams` claim, prefix-stripped and `+`-joined when the user is in +several teams, so usage is never double-counted; `unassigned` when empty or +unknown), and [provider, model, class] (`class` = low | medium | high from x-middleman-priority); cache metrics on [provider, model] only, to keep the high-cardinality user dimension cheap. diff --git a/middleman/src/middleman/passthrough.py b/middleman/src/middleman/passthrough.py index 1b29cd739b..609227eb64 100644 --- a/middleman/src/middleman/passthrough.py +++ b/middleman/src/middleman/passthrough.py @@ -15,11 +15,20 @@ from ddtrace.trace import tracer from fastapi import HTTPException, Request from fastapi.responses import JSONResponse, StreamingResponse +from hawk.core.auth.workload_jwks import WorkloadIssuerUnavailable +from hawk.core.types.workload_auth import WorkloadPrincipal from opentelemetry import trace as otel_trace -from middleman import anthropic_betas, apis, gcloud, models +from middleman import anthropic_betas, apis, gcloud, model_policy, models, workload_auth from middleman.apis import OpenaiLegacyCompletionsApi, api_to_class -from middleman.auth import UserInfo, get_user_info +from middleman.auth import ( + RequestPrincipal, + UserInfo, + get_user_info, + principal_subject, + record_principal, + usage_attribution, +) from middleman.classes import Priority, UpstreamStreamInterrupted from middleman.credential_broker import ( ApiKeyCredential, @@ -286,7 +295,9 @@ def _extract_bearer_token(auth_header: str) -> str: return parts[1] -async def _authenticate_request(request: Request, header: str, error_status_code: int) -> UserInfo: +async def _authenticate_request( + request: Request, header: str, error_status_code: int, *, allow_workload: bool = True +) -> RequestPrincipal: raw_value = request.headers.get(header) if not raw_value or not raw_value.strip(): context.mark_anonymous(getattr(request.state, "traffic_log", None)) @@ -299,31 +310,48 @@ async def _authenticate_request(request: Request, header: str, error_status_code raise PassthroughException(status_code=error_status_code, detail="invalid api key") from None try: - return await get_user_info(api_key) + principal = await workload_auth.maybe_authenticate(api_key) if allow_workload else None + if principal is None: + principal = await get_user_info(api_key) + handle = getattr(request.state, "traffic_log", None) + if handle is not None: + record_principal(handle, principal) + return principal + except WorkloadIssuerUnavailable: + raise PassthroughException( + status_code=503, + detail="Workload issuer unavailable", + headers={"Retry-After": "5"}, + ) from None except Exception: raise PassthroughException(status_code=error_status_code, detail="invalid api key") from None -async def authenticate_anthropic_request(request: Request) -> UserInfo: +async def authenticate_anthropic_request(request: Request) -> RequestPrincipal: return await _authenticate_request(request, "x-api-key", 401) -async def authenticate_openai_request(request: Request) -> UserInfo: +async def authenticate_openai_request(request: Request) -> RequestPrincipal: return await _authenticate_request(request, "authorization", 401) -async def authenticate_gemini_request(request: Request) -> UserInfo: +async def authenticate_gemini_request(request: Request) -> RequestPrincipal: return await _authenticate_request(request, "x-goog-api-key", 401) -async def validate_model_access(model_names: list[str], user_groups: list[str]) -> list[models.ModelInfo]: - permitted = models.get_current_models().get_permitted_models_by_public_name(user_groups) +async def authenticate_openai_human_request(request: Request) -> UserInfo: + principal = await _authenticate_request(request, "authorization", 401, allow_workload=False) + assert isinstance(principal, UserInfo) + return principal - model_infos: list[models.ModelInfo] = [] - for model_name in model_names: - if model_name not in permitted: - raise PassthroughException(status_code=404, detail="model not found") - model_info = permitted[model_name] + +async def validate_model_access(model_names: list[str], principal: RequestPrincipal) -> list[models.ModelInfo]: + snapshot = dict(models.get_current_models().models) + try: + model_infos = model_policy.select_models(snapshot, principal, model_names) + except model_policy.ModelNotAllowedError: + raise PassthroughException(status_code=404, detail="model not found") from None + for model_info in model_infos: # Checked here, the one place every passthrough route resolves its model, so a # model whose required_anthropic_betas is malformed is refused on every route, # not just the Anthropic one. After the access check, as for a missing flag. @@ -332,8 +360,6 @@ async def validate_model_access(model_names: list[str], user_groups: list[str]) status_code=500, detail=anthropic_betas.misconfigured_detail(model_info.public_name) ) - model_infos.append(model_info) - return model_infos @@ -342,10 +368,10 @@ async def validate_model_access(model_names: list[str], user_groups: list[str]) _CHANNEL_BY_JOB_TYPE = {"eval-set": "eval-set", "scan": "scan", "scan-resume": "scan"} -def request_channel(request: Request) -> str: - """Classify traffic for usage metrics: Hawk runner jobs send the - x-hawk-job-type correlation header (eval-set | scan | scan-resume); - anything else is direct API usage.""" +def request_channel(request: Request, principal: RequestPrincipal | None = None) -> str: + """Use the signed workload job type, or human requests' correlation header.""" + if isinstance(principal, WorkloadPrincipal): + return principal.claims.workload.job.type return _CHANNEL_BY_JOB_TYPE.get(request.headers.get("x-hawk-job-type", ""), "direct") @@ -371,7 +397,7 @@ async def make_post_request( provider_name: str = "unknown", public_name: str = "unknown", model_config: models.ModelInfo | None = None, - user: UserInfo | None = None, + user: RequestPrincipal | None = None, traffic_log: TrafficLog | None = None, channel: str = "direct", is_metadata_request: bool = False, @@ -395,9 +421,9 @@ async def make_post_request( otel_span.set_attribute("upstream.provider", provider_name) otel_span.set_attribute("peer.service", provider_name) otel_span.set_attribute("http.method", "POST") - if user: - otel_span.set_attribute("hawk.user.id", user.id) - if user.email: + if user is not None: + otel_span.set_attribute("hawk.user.id", principal_subject(user)) + if not isinstance(user, WorkloadPrincipal) and user.email: otel_span.set_attribute("hawk.user.email", user.email) start = time.monotonic() lab_response = await session.post(url, data=data, json=json, headers=headers, **kwargs) @@ -591,12 +617,15 @@ async def get_content() -> AsyncIterator[bytes]: cache_creation_5m=usage.cache_write_5m_tokens, cache_creation_1h=usage.cache_write_1h_tokens, ) + usage_user, usage_teams = ( + usage_attribution(user) if user is not None else ("unknown", None) + ) emf_emitter.record_usage( provider=provider_name, model=public_name, - user=_usage_user_label(user), + user=usage_user if isinstance(user, WorkloadPrincipal) else _usage_user_label(user), channel=channel, - team=team_label(user.teams) if user else UNASSIGNED_TEAM, + team=team_label(usage_teams) if usage_teams is not None else UNASSIGNED_TEAM, usage=usage, priority_class=priority_class, ) @@ -791,10 +820,10 @@ def _populate_traffic_log_entry( model_info: models.ModelInfo, request: Request, stream: bool | None, - user: UserInfo, + user: RequestPrincipal, ) -> None: """Populate traffic-log fields that are known at handler entry.""" - handle.set_user(user.id, user.groups, user.is_admin, teams=user.teams, email=user.email) + record_principal(handle, user) handle.set_provider(provider) handle.set_public_name(model_info.public_name) handle.set_model_lab(model_info.lab) @@ -820,7 +849,7 @@ async def _handle_anthropic_request( raise PassthroughException(status_code=400, detail="model field is required") try: - model_infos = await validate_model_access(model_names=[body["model"]], user_groups=user.groups) + model_infos = await validate_model_access(model_names=[body["model"]], principal=user) model_info = model_infos[0] if handle is not None: @@ -848,7 +877,7 @@ async def _handle_anthropic_request( decision = priority_runtime.decide( provider="anthropic", public_name=model_info.public_name, - principal=user.id, + principal=usage_attribution(user)[0], priority=priority, baseline_profile=model_info.anthropic_account, ) @@ -894,7 +923,7 @@ async def _handle_anthropic_request( model_config=model_info, user=user, traffic_log=handle, - channel=request_channel(request), + channel=request_channel(request, user), is_metadata_request=is_metadata_request, priority_decision=decision, priority_class=priority.value, @@ -936,7 +965,7 @@ async def handle_gemini_vertex_passthrough( raise PassthroughException(status_code=400, detail="invalid JSON body") from None try: - model_infos = await validate_model_access(model_names=[model], user_groups=user.groups) + model_infos = await validate_model_access(model_names=[model], principal=user) model_info = model_infos[0] if not model_info.lab.startswith("gemini-vertex-chat"): @@ -966,7 +995,7 @@ async def handle_gemini_vertex_passthrough( model_config=model_info, user=user, traffic_log=handle, - channel=request_channel(request), + channel=request_channel(request, user), priority_class=get_priority(request).value, ) if handle is not None: @@ -1008,7 +1037,7 @@ async def handle_gemini_developer_api_passthrough( raise PassthroughException(status_code=400, detail="invalid JSON body") from None try: - model_infos = await validate_model_access(model_names=[model], user_groups=user.groups) + model_infos = await validate_model_access(model_names=[model], principal=user) model_info = model_infos[0] if model_info.lab != "gemini-developer-api": @@ -1038,7 +1067,7 @@ async def handle_gemini_developer_api_passthrough( model_config=model_info, user=user, traffic_log=handle, - channel=request_channel(request), + channel=request_channel(request, user), priority_class=get_priority(request).value, ) if handle is not None: @@ -1109,7 +1138,7 @@ async def handle_openai_v1_chat_completions_and_responses(request: Request) -> P raise PassthroughException(status_code=400, detail="model field is required") try: - model_infos = await validate_model_access(model_names=[body["model"]], user_groups=user.groups) + model_infos = await validate_model_access(model_names=[body["model"]], principal=user) model_info = model_infos[0] lab_class = api_to_class.get(model_info.lab, None) @@ -1126,6 +1155,14 @@ async def handle_openai_v1_chat_completions_and_responses(request: Request) -> P if path not in _SUPPORTED_OPENAI_CHAT_COMPLETIONS_AND_RESPONSES_PATHS: raise PassthroughException(status_code=404, detail="not found") + if ( + isinstance(user, WorkloadPrincipal) + and path == "/responses" + and "background" in body + and body["background"] is not False + ): + raise PassthroughException(status_code=400, detail="background Responses are not supported for workloads") + if handle is not None: _populate_traffic_log_entry(handle, "openai", model_info, request, body.get("stream"), user) @@ -1147,7 +1184,7 @@ async def handle_openai_v1_chat_completions_and_responses(request: Request) -> P decision = priority_runtime.decide( provider="openai", public_name=model_info.public_name, - principal=user.id, + principal=usage_attribution(user)[0], priority=priority, baseline_profile=baseline_profile, ) @@ -1197,7 +1234,7 @@ async def handle_openai_v1_chat_completions_and_responses(request: Request) -> P model_config=model_info, user=user, traffic_log=handle, - channel=request_channel(request), + channel=request_channel(request, user), is_metadata_request=path == "/responses/input_tokens", priority_decision=decision, priority_class=priority.value, @@ -1232,7 +1269,7 @@ async def handle_openai_v1_completions(request: Request) -> PassthroughResult: raise PassthroughException(status_code=400, detail="model field is required") try: - model_infos = await validate_model_access(model_names=[body["model"]], user_groups=user.groups) + model_infos = await validate_model_access(model_names=[body["model"]], principal=user) model_info = model_infos[0] lab_class = api_to_class.get(model_info.lab, None) @@ -1265,7 +1302,7 @@ async def handle_openai_v1_completions(request: Request) -> PassthroughResult: model_config=model_info, user=user, traffic_log=handle, - channel=request_channel(request), + channel=request_channel(request, user), priority_class=priority.value, ) if handle is not None: @@ -1317,7 +1354,7 @@ async def _validate_file(user: UserInfo, file: BinaryIO) -> None: if not model_names: raise PassthroughException(status_code=400, detail="file contains no valid requests") - await validate_model_access(model_names, user.groups) + await validate_model_access(model_names, user) file.seek(0) @@ -1328,9 +1365,9 @@ async def handle_openai_v1_upload_file(request: Request) -> PassthroughResult: handle.set_provider("openai") handle.set_routing(method=request.method, endpoint=request.url.path) - user = await authenticate_openai_request(request) + user = await authenticate_openai_human_request(request) if handle is not None: - handle.set_user(user.id, user.groups, user.is_admin, teams=user.teams, email=user.email) + record_principal(handle, user) try: request_data = await request.form() @@ -1370,7 +1407,9 @@ async def handle_openai_v1_upload_file(request: Request) -> PassthroughResult: include_response_header=lambda header: header.startswith(("x-", "openai-")), provider_name="openai", public_name="batch-file-upload", + user=user, traffic_log=handle, + channel=request_channel(request, user), ) if handle is not None: handle.set_upstream( @@ -1397,6 +1436,8 @@ def get_anthropic_error_response(exc: PassthroughException) -> JSONResponse: error_type = "authentication_error" case 403: error_type = "permission_error" + case 503: + error_type = "api_error" case 404: error_type = "not_found_error" case 429: @@ -1412,10 +1453,12 @@ def get_anthropic_error_response(exc: PassthroughException) -> JSONResponse: def get_openai_error_response(exc: PassthroughException) -> JSONResponse: - error_type = "invalid_request_error" + error_type = "server_error" if exc.status_code == 503 else "invalid_request_error" match exc.status_code: case 401: code = "invalid_authentication" + case 503: + code = "service_unavailable" case 404: code = "model_not_found" case _: @@ -1437,6 +1480,8 @@ def get_gemini_error_response(exc: PassthroughException) -> JSONResponse: status = "INVALID_ARGUMENT" case 401 | 403: status = "PERMISSION_DENIED" + case 503: + status = "UNAVAILABLE" case 404: status = "NOT_FOUND" case _: @@ -1444,4 +1489,5 @@ def get_gemini_error_response(exc: PassthroughException) -> JSONResponse: return JSONResponse( {"error": {"code": exc.status_code, "message": exc.detail, "status": status}}, status_code=exc.status_code, + headers=exc.headers, ) diff --git a/middleman/src/middleman/server.py b/middleman/src/middleman/server.py index 6355e1722f..9e0af89c34 100644 --- a/middleman/src/middleman/server.py +++ b/middleman/src/middleman/server.py @@ -24,14 +24,14 @@ from fastapi.encoders import jsonable_encoder from fastapi.exceptions import RequestValidationError from fastapi.responses import JSONResponse, Response, StreamingResponse -from hawk.core.types import workload_auth +from hawk.core.types import workload_auth as workload_types from openai.types import Model as OpenAIModel from pydantic import BaseModel from starlette.exceptions import HTTPException as StarletteHTTPException from starlette.requests import ClientDisconnect from starlette.types import ASGIApp, Message, Receive, Scope, Send -from middleman import gcloud, model_approvals, models, otel_tracing, passthrough +from middleman import gcloud, model_approvals, model_policy, models, otel_tracing, passthrough, workload_auth from middleman.admin.models_router import router as models_router from middleman.admin.priority_router import router as priority_router from middleman.admin.secrets_router import router as secrets_router @@ -202,40 +202,41 @@ async def lifespan(_app: FastAPI) -> AsyncGenerator[None]: anthropic_credential_broker.load_profiles() openai_credential_broker.load_profiles() - current_models, _ = await asyncio.gather(models.init_models(), load_auth_providers()) - maybe_init_vertex_urls(m.lab for m in current_models.models.values()) + current_models, human_providers = await asyncio.gather(models.init_models(), load_auth_providers()) + async with workload_auth.lifespan(human_providers): + maybe_init_vertex_urls(m.lab for m in current_models.models.values()) - # Warm the gcloud token before serving so the first Gemini request reads it - # from cache instead of doing a blocking refresh on the loop. - await gcloud.refresh_gcloud_token() + # Warm the gcloud token before serving so the first Gemini request reads it + # from cache instead of doing a blocking refresh on the loop. + await gcloud.refresh_gcloud_token() - app_state.token_counter = get_default_token_counter() + app_state.token_counter = get_default_token_counter() - refresh_task = asyncio.create_task(_periodic_key_refresh()) - await cache_bus.start(_reload_all_caches) - await reload_priority_config() - await priority_runtime.start() - priority_refresh_task = asyncio.create_task(refresh_priority_config()) - await rate_limit_store.start() - await emf_emitter.start() - await inflight.start() - if traffic_log_emitter is not None: - await traffic_log_emitter.start() - logger.info("traffic_log_started") - try: - yield - finally: - await _cancel_task(refresh_task) - await _cancel_task(priority_refresh_task) - await priority_runtime.stop() - await cache_bus.stop() - await rate_limit_store.stop() - await emf_emitter.stop(drain_timeout_s=2.0) - await inflight.stop() + refresh_task = asyncio.create_task(_periodic_key_refresh()) + await cache_bus.start(_reload_all_caches) + await reload_priority_config() + await priority_runtime.start() + priority_refresh_task = asyncio.create_task(refresh_priority_config()) + await rate_limit_store.start() + await emf_emitter.start() + await inflight.start() if traffic_log_emitter is not None: - await traffic_log_emitter.stop(drain_timeout_s=2.0) - otel_tracing.shutdown() - await close_client_session() + await traffic_log_emitter.start() + logger.info("traffic_log_started") + try: + yield + finally: + await _cancel_task(refresh_task) + await _cancel_task(priority_refresh_task) + await priority_runtime.stop() + await cache_bus.stop() + await rate_limit_store.stop() + await emf_emitter.stop(drain_timeout_s=2.0) + await inflight.stop() + if traffic_log_emitter is not None: + await traffic_log_emitter.stop(drain_timeout_s=2.0) + otel_tracing.shutdown() + await close_client_session() # App-wide so no passthrough route can be added without it. @@ -623,15 +624,15 @@ async def get_model_groups( @app.post("/model_approvals/resolve") async def resolve_model_approvals_route( - request: workload_auth.ModelApprovalRequest, + request: workload_types.ModelApprovalRequest, current_models: Annotated[Models, Depends(get_models)], credentials: Annotated[fastapi.security.HTTPAuthorizationCredentials, Depends(http_bearer)], -) -> workload_auth.ModelApprovalResponse: +) -> workload_types.ModelApprovalResponse: user = await get_user_info(credentials.credentials) approvals = model_approvals.resolve_model_approvals(current_models.models, request.model_names, user.groups) - return workload_auth.ModelApprovalResponse( + return workload_types.ModelApprovalResponse( models=approvals, - usage_attribution=workload_auth.UsageAttribution(user_id=user.id, teams=tuple(user.teams)), + usage_attribution=workload_types.UsageAttribution(user_id=user.id, teams=tuple(user.teams)), ) @@ -910,7 +911,7 @@ async def openai_v1_models( # Answering here skips handle_http_exception, which used to attribute this. _record_exception_on_traffic_log(request, exc) return passthrough.get_openai_error_response(exc) - permitted = models.get_permitted_models_by_public_name(user.groups) + permitted = model_policy.permitted_models(dict(models.models), user) return OpenAIModelList( data=[ OpenAIModel( diff --git a/middleman/src/middleman/traffic_log/envelope.py b/middleman/src/middleman/traffic_log/envelope.py index 135ecab0a8..e7880a0d4c 100644 --- a/middleman/src/middleman/traffic_log/envelope.py +++ b/middleman/src/middleman/traffic_log/envelope.py @@ -34,6 +34,13 @@ class TrafficLogEnvelope(BaseModel): user_groups: list[str] | None = None user_teams: list[str] | None = None is_admin: bool | None = None + workload_execution_id: str | None = None + workload_grant_id: str | None = None + workload_job_id: str | None = None + workload_job_type: Literal["eval-set", "scan"] | None = None + # Signed accounting metadata; distinct from the authenticated workload subject. + usage_user_id: str | None = None + usage_teams: list[str] | None = None source_ip: str user_agent: str @@ -50,7 +57,7 @@ class TrafficLogEnvelope(BaseModel): endpoint: str | None = None upstream_url: str | None = None - # Admission policy observation (authenticated identity, never caller metadata). + # Admission policy observation (trusted accounting identity, never caller metadata). priority_policy_version: str | None = None priority_mode: Literal["off", "shadow", "enforce"] | None = None priority_requested: str | None = None diff --git a/middleman/src/middleman/traffic_log/handle.py b/middleman/src/middleman/traffic_log/handle.py index 1db809bea7..d5760cfff3 100644 --- a/middleman/src/middleman/traffic_log/handle.py +++ b/middleman/src/middleman/traffic_log/handle.py @@ -35,6 +35,25 @@ def set_user( self.fields["user_teams"] = list(teams) if teams is not None else None self.fields["is_admin"] = is_admin + def set_usage_attribution(self, user_id: str, teams: list[str]) -> None: + self.fields["usage_user_id"] = user_id + self.fields["usage_teams"] = list(teams) + + def set_workload( + self, + *, + execution_id: str, + grant_id: str, + job_id: str, + job_type: Literal["eval-set", "scan"], + ) -> None: + self.fields.update( + workload_execution_id=execution_id, + workload_grant_id=grant_id, + workload_job_id=job_id, + workload_job_type=job_type, + ) + def set_provider(self, provider: str | None) -> None: self.fields["provider"] = provider diff --git a/middleman/src/middleman/traffic_log/middleware.py b/middleman/src/middleman/traffic_log/middleware.py index c97fba00b5..1c62e4e6a1 100644 --- a/middleman/src/middleman/traffic_log/middleware.py +++ b/middleman/src/middleman/traffic_log/middleware.py @@ -4,6 +4,7 @@ import collections import datetime import json +import re import time from typing import TYPE_CHECKING, Any, cast @@ -50,6 +51,8 @@ def parse_body_cap(value: str | None, default: int) -> int: _SENSITIVE_HEADER_NAMES = {"authorization", "x-api-key", "x-goog-api-key", "cookie", "set-cookie"} +# Possessive matching avoids retaining backtracking state for large truncated bodies. +_JSON_FIELD = re.compile(r'"(?:[^"\\]|\\.)*+"\s*:') _EXCLUDED_PATH_PREFIXES = ("/health", "/admin") @@ -361,6 +364,17 @@ def _scrub_headers(headers: dict[str, str]) -> dict[str, str]: def _redact_api_key(body: Any) -> Any: if isinstance(body, dict) and "api_key" in body: return {**cast("dict[str, Any]", body), "api_key": "[REDACTED]"} + if isinstance(body, str): + # Malformed or capture-truncated JSON falls back to text. Once a sensitive + # field starts, omit the tail: its value may itself be incomplete, so we + # cannot safely determine where the credential ends. + for field in _JSON_FIELD.finditer(body): + try: + name = json.loads(field.group()[:-1]) + except json.JSONDecodeError: + continue + if name == "api_key": + return body[: field.start()] + '"api_key": "[REDACTED]"' return cast("Any", body) diff --git a/middleman/src/middleman/workload_auth.py b/middleman/src/middleman/workload_auth.py new file mode 100644 index 0000000000..30520b08cd --- /dev/null +++ b/middleman/src/middleman/workload_auth.py @@ -0,0 +1,96 @@ +from __future__ import annotations + +import json +import os +import time +from collections.abc import AsyncGenerator, Collection +from contextlib import asynccontextmanager +from dataclasses import dataclass +from typing import Any, cast + +import joserfc.errors +import joserfc.jws +from ddtrace.trace import tracer +from hawk.core.auth.workload_jwks import WorkloadJWKSCache +from hawk.core.auth.workload_jwt import InvalidWorkloadToken, validate_workload_jwt +from hawk.core.types.workload_auth import WorkloadPrincipal, WorkloadValidationConfig +from httpx import AsyncClient + +from middleman.observability.metrics import record_auth_duration + + +def load_settings() -> WorkloadValidationConfig | None: + raw = os.environ.get("MIDDLEMAN_WORKLOAD_AUTH") + if raw is None: + return None + try: + return WorkloadValidationConfig.model_validate_json(raw) + except ValueError: + raise ValueError("Invalid MIDDLEMAN_WORKLOAD_AUTH") from None + + +def validate_human_issuers(issuers: Collection[str]) -> None: + settings = load_settings() + if settings is not None and settings.issuer in issuers: + raise ValueError("Workload issuer is also configured for human authentication") + + +@dataclass(frozen=True) +class _Runtime: + settings: WorkloadValidationConfig + keys: WorkloadJWKSCache + + +_runtime: _Runtime | None = None + + +@asynccontextmanager +async def lifespan(human_issuers: Collection[str]) -> AsyncGenerator[None]: + global _runtime + settings = load_settings() + validate_human_issuers(human_issuers) + if settings is None: + yield + return + async with AsyncClient() as client: + _runtime = _Runtime(settings, WorkloadJWKSCache(settings, client=client)) + try: + yield + finally: + _runtime = None + + +def is_workload_issuer(issuer: object) -> bool: + return _runtime is not None and issuer == _runtime.settings.issuer + + +async def maybe_authenticate(token: str) -> WorkloadPrincipal | None: + if token.startswith("hawk_wl_v1."): + raise InvalidWorkloadToken("Invalid workload token") + runtime = _runtime + if runtime is None: + return None + try: + payload = json.loads(joserfc.jws.extract_compact(token.encode("utf-8"), None).payload) + if not isinstance(payload, dict): + raise ValueError + payload = cast(dict[str, Any], payload) + if not isinstance(payload.get("iss"), str): + raise ValueError + except (ValueError, TypeError, RecursionError, joserfc.errors.JoseError): + raise InvalidWorkloadToken("Invalid workload token") from None + if payload["iss"] != runtime.settings.issuer: + return None + started = time.monotonic() + try: + with tracer.trace("auth.validate_workload", service="middleman") as span: + principal = await validate_workload_jwt( + token, + config=runtime.settings, + keys=runtime.keys, + ) + span.set_tag("auth.user_id", principal.claims.sub) + span.set_tag("auth.issuer", runtime.settings.issuer) + return principal + finally: + record_auth_duration((time.monotonic() - started) * 1000) diff --git a/middleman/tests/AGENTS.md b/middleman/tests/AGENTS.md index ee307d3aa4..23dda5bb1b 100644 --- a/middleman/tests/AGENTS.md +++ b/middleman/tests/AGENTS.md @@ -18,7 +18,7 @@ uv run pytest -k "test_successful" # By name Abstract class for parameterized passthrough testing. One test method runs against all providers. **To add a new provider's passthrough tests**: -1. Create `RequestExecutor` subclass with `expected_outgoing_url()`, `expected_outgoing_auth_header` (property), `_build_request()` +1. Create `RequestExecutor` subclass with `expected_outgoing_url()`, `expected_outgoing_auth_header` (property), `build_request()` 2. Add to `@pytest.mark.parametrize` list in `TestPassthroughEndpointHandler` 3. Use `pytest.param(MyExecutor(), model, id=f"prefix_{model}")` for test IDs diff --git a/middleman/tests/conftest.py b/middleman/tests/conftest.py index 9bb645ac06..46581610f1 100644 --- a/middleman/tests/conftest.py +++ b/middleman/tests/conftest.py @@ -106,3 +106,18 @@ def fixture_mock_private_models(tmp_path: pathlib.Path, monkeypatch: pytest.Monk monkeypatch.setenv(file_env, str(file)) if module_attr: monkeypatch.setattr(models, module_attr, str(file)) + + +@pytest.fixture +async def workload_authority(monkeypatch: pytest.MonkeyPatch): + import httpx + + from middleman import workload_auth + from tests.workload_support import SETTINGS, Authority + + authority = Authority() + client = httpx.AsyncClient(transport=httpx.MockTransport(authority.endpoint)) + monkeypatch.setenv("MIDDLEMAN_WORKLOAD_AUTH", json.dumps(SETTINGS)) + monkeypatch.setattr(workload_auth, "AsyncClient", lambda: client) + async with workload_auth.lifespan({TEST_ISSUER}): + yield authority diff --git a/middleman/tests/test_model_policy.py b/middleman/tests/test_model_policy.py new file mode 100644 index 0000000000..7174be76dc --- /dev/null +++ b/middleman/tests/test_model_policy.py @@ -0,0 +1,67 @@ +import dataclasses +from typing import Any + +import pytest + +from middleman import model_policy, workload_auth +from middleman.auth import UserInfo +from middleman.model_info import ModelInfo +from tests.workload_support import Authority + + +@pytest.fixture +def model(): + return ModelInfo(public_name="allowed", danger_name="upstream-v1", lab="openai-chat", group="same-group") + + +@pytest.mark.parametrize( + "change", + [ + {"danger_name": "upstream-v2"}, + {"group": "changed"}, + {"lab": "deepseek"}, + {"dead": True}, + ], +) +async def test_changed_approval_or_dead_model_is_denied( + workload_authority: Authority, model: ModelInfo, change: dict[str, Any] +): + principal = await workload_auth.maybe_authenticate(workload_authority.token((model,))) + assert principal is not None + replacement = dataclasses.replace(model, **change) + with pytest.raises(model_policy.ModelNotAllowedError): + model_policy.select_models({model.public_name: replacement}, principal, [model.public_name]) + + +async def test_empty_invocation_grant_denies_public_models(workload_authority: Authority, model: ModelInfo): + principal = await workload_auth.maybe_authenticate(workload_authority.token()) + assert principal is not None + public = dataclasses.replace(model, group="model-access-public") + assert model_policy.permitted_models({public.public_name: public}, principal) == {} + + +def test_human_selection_preserves_group_and_dead_model_behavior(model: ModelInfo): + dead = dataclasses.replace(model, dead=True) + human = UserInfo(id="human", groups=[model.group]) + assert model_policy.select_models({model.public_name: dead}, human, [model.public_name]) == [dead] + + +async def test_unapproved_invalid_registry_name_does_not_break_grant(workload_authority: Authority, model: ModelInfo): + principal = await workload_auth.maybe_authenticate(workload_authority.token((model,))) + assert principal is not None + unrelated = dataclasses.replace(model, public_name="trailing-space ", group="unrelated") + snapshot = {entry.public_name: entry for entry in (model, unrelated)} + assert model_policy.permitted_models(snapshot, principal) == {model.public_name: model} + assert model_policy.select_models(snapshot, principal, [model.public_name]) == [model] + + +async def test_invalid_changed_approval_does_not_block_other_models(workload_authority: Authority, model: ModelInfo): + other = dataclasses.replace(model, public_name="other") + principal = await workload_auth.maybe_authenticate(workload_authority.token((model, other))) + assert principal is not None + changed = dataclasses.replace(other, danger_name="changed ") + snapshot = {model.public_name: model, other.public_name: changed} + assert model_policy.permitted_models(snapshot, principal) == {model.public_name: model} + assert model_policy.select_models(snapshot, principal, [model.public_name]) == [model] + with pytest.raises(model_policy.ModelNotAllowedError): + model_policy.select_models(snapshot, principal, [other.public_name]) diff --git a/middleman/tests/test_passthrough.py b/middleman/tests/test_passthrough.py index 3fd8df38c3..ad5a93d85e 100644 --- a/middleman/tests/test_passthrough.py +++ b/middleman/tests/test_passthrough.py @@ -127,7 +127,7 @@ async def shed(): class RequestExecutor(ABC): def execute(self, model: str, api_key: str | None, extra_headers: dict[str, str] | None = None) -> httpx.Response: - request = self._build_request(model, api_key) + request = self.build_request(model, api_key) test_client = fastapi.testclient.TestClient(server.app) return test_client.post(request.path, headers={**request.headers, **(extra_headers or {})}, json=request.body) @@ -140,7 +140,7 @@ def expected_outgoing_auth_header(self) -> str: raise NotImplementedError @abstractmethod - def _build_request(self, model: str, api_key: str | None) -> Request: + def build_request(self, model: str, api_key: str | None) -> Request: pass @@ -155,7 +155,7 @@ def expected_outgoing_auth_header(self) -> str: return "x-api-key" @override - def _build_request(self, model: str, api_key: str | None) -> Request: + def build_request(self, model: str, api_key: str | None) -> Request: return Request( path="/anthropic/v1/messages", headers={"x-api-key": api_key} if api_key else {}, body={"model": model} ) @@ -172,7 +172,7 @@ def expected_outgoing_auth_header(self) -> str: return "x-api-key" @override - def _build_request(self, model: str, api_key: str | None) -> Request: + def build_request(self, model: str, api_key: str | None) -> Request: return Request( path="/anthropic/v1/messages/count_tokens", headers={"x-api-key": api_key} if api_key else {}, @@ -194,7 +194,7 @@ def expected_outgoing_auth_header(self) -> str: return "authorization" @override - def _build_request(self, model: str, api_key: str | None) -> Request: + def build_request(self, model: str, api_key: str | None) -> Request: return Request( path=f"/gemini/v1beta1/publishers/google/models/{model}:{self.operation}?alt=sse", headers={"x-goog-api-key": api_key} if api_key else {}, @@ -216,7 +216,7 @@ def expected_outgoing_auth_header(self) -> str: return "x-goog-api-key" @override - def _build_request(self, model: str, api_key: str | None) -> Request: + def build_request(self, model: str, api_key: str | None) -> Request: return Request( path=f"/google-ai/v1beta/models/{model}:{self.operation}?alt=sse", headers={"x-goog-api-key": api_key} if api_key else {}, @@ -247,7 +247,7 @@ def expected_outgoing_auth_header(self) -> str: return "authorization" @override - def _build_request(self, model: str, api_key: str | None) -> Request: + def build_request(self, model: str, api_key: str | None) -> Request: return Request( path=f"/openai/v1/{self.operation}", headers={"authorization": f"Bearer {api_key}"} if api_key else {}, @@ -258,7 +258,7 @@ def _build_request(self, model: str, api_key: str | None) -> Request: @pytest.fixture def mock_auth(mocker: MockerFixture): mocked = mocker.patch("middleman.passthrough.get_user_info", autospec=True) - mocked.return_value.groups = ["test_permission"] + mocked.return_value = UserInfo(id="test-user", groups=["test_permission"]) return mocked @@ -1706,7 +1706,9 @@ async def test_anthropic_route_still_serves_every_anthropic_lab(model: str, exec async def test_anthropic_route_rejects_a_model_whose_lab_has_no_api_class(lab: str | None, mocker: MockerFixture): """ModelInfo.lab is unvalidated at runtime, so this needs no code change to occur.""" forward = mocker.spy(passthrough, "make_post_request") - real = await passthrough.validate_model_access(model_names=["test_model"], user_groups=["test_permission"]) + real = await passthrough.validate_model_access( + model_names=["test_model"], principal=UserInfo(id="test-user", groups=["test_permission"]) + ) mocker.patch.object( passthrough, "validate_model_access", @@ -3709,9 +3711,13 @@ async def test_validate_model_access_checks_access_before_misconfiguration(mocke _install_models(mocker, _required_betas_models(required_anthropic_betas=[])) with pytest.raises(passthrough.PassthroughException) as not_permitted: - await passthrough.validate_model_access(["claude-gated"], user_groups=["other_permission"]) + await passthrough.validate_model_access( + ["claude-gated"], principal=UserInfo(id="test-user", groups=["other_permission"]) + ) with pytest.raises(passthrough.PassthroughException) as permitted: - await passthrough.validate_model_access(["claude-gated"], user_groups=["test_permission"]) + await passthrough.validate_model_access( + ["claude-gated"], principal=UserInfo(id="test-user", groups=["test_permission"]) + ) assert (not_permitted.value.status_code, not_permitted.value.detail) == (404, "model not found") assert permitted.value.status_code == 500 diff --git a/middleman/tests/test_server.py b/middleman/tests/test_server.py index 67ae2f707c..64d15d517d 100644 --- a/middleman/tests/test_server.py +++ b/middleman/tests/test_server.py @@ -245,7 +245,9 @@ async def test_permitted_models_info_wire_shape( if path == "/public_models_info": assert client.post("/permitted_models", json={"api_key": "test"}).json() == ["private-model"] with pytest.raises(passthrough.PassthroughException): - await passthrough.validate_model_access([raw_model["public_name"]], groups) + await passthrough.validate_model_access( + [raw_model["public_name"]], principal=auth.UserInfo(id="test", groups=groups) + ) finally: models._current_models = None # pyright: ignore[reportPrivateUsage] diff --git a/middleman/tests/test_workload_auth.py b/middleman/tests/test_workload_auth.py new file mode 100644 index 0000000000..7415ad2a2c --- /dev/null +++ b/middleman/tests/test_workload_auth.py @@ -0,0 +1,27 @@ +import json + +import pytest + +from middleman import workload_auth +from tests.workload_support import ISSUER, SETTINGS + + +@pytest.mark.parametrize( + "raw", + [ + "", + json.dumps({key: value for key, value in SETTINGS.items() if key != "audience"}), + json.dumps({**SETTINGS, "audience": [SETTINGS["audience"]]}), + ], +) +def test_bad_config_fails_closed(monkeypatch: pytest.MonkeyPatch, raw: str): + monkeypatch.setenv("MIDDLEMAN_WORKLOAD_AUTH", raw) + with pytest.raises(ValueError, match="Invalid MIDDLEMAN_WORKLOAD_AUTH"): + workload_auth.load_settings() + + +async def test_human_issuer_overlap_is_rejected(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("MIDDLEMAN_WORKLOAD_AUTH", json.dumps(SETTINGS)) + with pytest.raises(ValueError, match="also configured for human authentication"): + async with workload_auth.lifespan({ISSUER}): + pytest.fail("overlapping issuers started") diff --git a/middleman/tests/test_workload_passthrough.py b/middleman/tests/test_workload_passthrough.py new file mode 100644 index 0000000000..344ad4baf6 --- /dev/null +++ b/middleman/tests/test_workload_passthrough.py @@ -0,0 +1,773 @@ +from __future__ import annotations + +import asyncio +import base64 +import dataclasses +import json +import time +from typing import Any, Literal +from unittest.mock import AsyncMock, MagicMock + +import aiohttp +import fastapi +import httpx +import pytest +from fastapi.responses import StreamingResponse +from hawk.core.auth import workload_jwt +from pytest_mock import MockerFixture + +from middleman import auth, models, passthrough, server +from middleman.auth import RequestPrincipal +from middleman.observability import emf +from middleman.priority_state import PriorityRuntime +from middleman.token_counter import TokenCounter +from middleman.traffic_log import emitter as traffic_emitter +from middleman.traffic_log import level, middleware +from tests import conftest, test_priority_state +from tests.test_passthrough import ( + AnthropicCountTokensRequestExecutor, + AnthropicRequestExecutor, + GeminiDeveloperApiRequestExecutor, + GeminiRequestExecutor, + OpenaiRequestExecutor, + RequestExecutor, +) +from tests.test_passthrough import mock_environment as mock_environment +from tests.workload_support import EXECUTION, Authority + +ROUTES = [ + (AnthropicRequestExecutor(), "anthropic-chat"), + (AnthropicCountTokensRequestExecutor(), "anthropic-chat"), + *[ + (OpenaiRequestExecutor("https://api.openai.com/v1", op), "openai-chat") + for op in ("chat/completions", "responses", "responses/input_tokens", "responses/compact") + ], + (OpenaiRequestExecutor("https://api.openai.com/v1", "completions"), "openai"), + *[ + (GeminiRequestExecutor(op), "gemini-vertex-chat") + for op in ("generateContent", "streamGenerateContent", "countTokens") + ], + *[ + (GeminiDeveloperApiRequestExecutor(op), "gemini-developer-api") + for op in ("generateContent", "streamGenerateContent", "countTokens") + ], + (OpenaiRequestExecutor("https://api.deepseek.com/v1", "chat/completions"), "deepseek"), +] + + +def install_models(monkeypatch: pytest.MonkeyPatch, lab: str = "openai-chat"): + store = models.Models( + models=[ + {"public_name": name, "danger_name": f"upstream-{name}", "lab": lab, "group": "same-group"} + for name in ("allowed", "denied") + ], + base_infos={}, + ) + monkeypatch.setattr(models, "_current_models", store) + monkeypatch.setattr(models, "_models_loaded_at", float("inf")) + return store + + +@pytest.fixture +async def client(): + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=server.app), base_url="http://middleman") as value: + yield value + + +@pytest.fixture +async def traffic_client(mocker: MockerFixture): + emitter = mocker.create_autospec(traffic_emitter.TrafficLogEmitter, instance=True) + emitted = asyncio.Event() + + def record_emit(**_kwargs: Any) -> None: + emitted.set() + + emitter.enqueue.side_effect = record_emit + + async def wait_for_emit(_response: httpx.Response) -> None: + await asyncio.wait_for(emitted.wait(), timeout=1) + emitted.clear() + + app = fastapi.FastAPI() + app.include_router(server.app.router) + app.add_middleware(middleware.TrafficLogMiddleware, env="test", level=level.Level.SUMMARY, emitter=emitter) + async with httpx.AsyncClient( + transport=httpx.ASGITransport(app=app), + base_url="http://middleman", + event_hooks={"response": [wait_for_emit]}, + ) as value: + yield value, emitter + + +@pytest.fixture +def upstream(mocker: MockerFixture): + async def forward(*_args: Any, **_kwargs: Any): + return StreamingResponse(iter([b'{"ok":true}']), media_type="application/json"), 1.0 + + return mocker.patch.object(passthrough, "make_post_request", autospec=True, side_effect=forward) + + +@pytest.mark.usefixtures("mock_environment") +@pytest.mark.parametrize("executor,lab", ROUTES) +async def test_registered_routes_enforce_exact_model( + executor: RequestExecutor, + lab: str, + workload_authority: Authority, + client: httpx.AsyncClient, + upstream: AsyncMock, + monkeypatch: pytest.MonkeyPatch, +): + store = install_models(monkeypatch, lab) + allowed = store.models["allowed"] + token = workload_authority.token( + (allowed,), + payload_overrides={"usage_attribution": {"user_id": "human-id", "teams": ["same-group"]}}, + ) + request = executor.build_request("allowed", token) + response = await client.post(request.path, headers=request.headers, json=request.body) + assert response.status_code == 200 + call = upstream.call_args + assert call is not None + assert call.kwargs["user"].claims.sub == f"workload:{EXECUTION}" + if isinstance(executor, (GeminiRequestExecutor, GeminiDeveloperApiRequestExecutor)): + assert "/models/upstream-allowed:" in call.args[0] + else: + assert call.kwargs["json"]["model"] == "upstream-allowed" + upstream.reset_mock() + request = executor.build_request("denied", token) + response = await client.post(request.path, headers=request.headers, json=request.body) + assert response.status_code == 404 + upstream.assert_not_called() + + +@pytest.mark.usefixtures("mock_environment") +@pytest.mark.parametrize( + "executor,lab,provider,baseline", + [ + (AnthropicRequestExecutor(), "anthropic-chat", "anthropic", None), + (OpenaiRequestExecutor("https://api.openai.com/v1", "responses"), "openai-chat", "openai", "OPENAI_API_KEY"), + ], +) +async def test_priority_flows_share_accounting_user_within_class_and_quota( + executor: RequestExecutor, + lab: str, + provider: Literal["anthropic", "openai"], + baseline: str | None, + workload_authority: Authority, + client: httpx.AsyncClient, + upstream: AsyncMock, + monkeypatch: pytest.MonkeyPatch, + mocker: MockerFixture, +): + store = install_models(monkeypatch, lab) + runtime = PriorityRuntime(clock=test_priority_state.Clock()) + config = test_priority_state.config( + provider=provider, dimension="requests", baseline_profile=baseline, aliases=["allowed"] + ) + runtime.configure( + config.model_copy( + update={ + "groups": [ + config.groups[0], + config.groups[0].model_copy(update={"name": "separate", "aliases": ["denied"]}), + ] + } + ) + ) + monkeypatch.setattr(passthrough, "priority_runtime", runtime) + mocker.patch.object( + passthrough, + "get_user_info", + autospec=True, + return_value=auth.UserInfo(id="alice", groups=["same-group"]), + ) + _, human_token = conftest.make_test_jwt(is_admin=False, sub="alice") + requests = [ + ("alice", None, "medium", "allowed"), + ("alice", "12345678-9abc-def0-1234-56789abcdef0", "medium", "allowed"), + ("alice", "12345678-9abc-def0-1234-56789abcdef5", "medium", "allowed"), + ("bob", "12345678-9abc-def0-1234-56789abcdef6", "medium", "allowed"), + ("alice", "12345678-9abc-def0-1234-56789abcdef0", "low", "allowed"), + ("alice", "12345678-9abc-def0-1234-56789abcdef0", "medium", "denied"), + ] + for user, execution, priority, model in requests: + token = ( + workload_authority.token( + tuple(store.models.values()), + overrides={"sub": f"workload:{execution}"}, + payload_overrides={"execution_id": execution, "usage_attribution": {"user_id": user, "teams": []}}, + ) + if execution + else human_token + ) + request = executor.build_request(model, token) + response = await client.post( + request.path, headers={**request.headers, "x-middleman-priority": priority}, json=request.body + ) + assert response.status_code == 200 + groups = runtime.worker_snapshot().groups + assert {key: value.arrivals for key, value in groups["shared"].items()} == { + "MEDIUM:alice": 3, + "MEDIUM:bob": 1, + "LOW:alice": 1, + } + assert {key: value.arrivals for key, value in groups["separate"].items()} == {"MEDIUM:alice": 1} + assert upstream.call_count == len(requests) + + +@pytest.mark.usefixtures("mock_environment") +@pytest.mark.parametrize("background", [True, None, 0]) +async def test_background_rejected_before_provider( + background: Any, + workload_authority: Authority, + client: httpx.AsyncClient, + upstream: AsyncMock, + monkeypatch: pytest.MonkeyPatch, +): + model = install_models(monkeypatch).models["allowed"] + response = await client.post( + "/openai/v1/responses", + headers={ + "authorization": f"Bearer {workload_authority.token((model,))}", + }, + json={"model": "allowed", "background": background, "stream": False, "input": "hello"}, + ) + assert response.status_code == 400 + assert "background" in response.json()["error"]["message"] + upstream.assert_not_called() + + +async def test_upload_rejected_before_parsing( + workload_authority: Authority, client: httpx.AsyncClient, mocker: MockerFixture +): + from starlette.requests import Request + + parse = mocker.patch.object(Request, "form", autospec=True) + response = await client.post( + "/openai/v1/files", + headers={ + "authorization": f"Bearer {workload_authority.token()}", + }, + files={"file": ("requests.jsonl", b"{}\n")}, + data={"purpose": "batch"}, + ) + assert response.status_code == 401 + parse.assert_not_called() + + +@pytest.mark.parametrize( + "path,body", + [ + ("/completions", {"model": "allowed", "prompt": "hello"}), + ("/embeddings", {"input": "hello"}), + ("/count_prompt_tokens", {"model": "allowed", "prompt": "hello"}), + ("/permitted_models", {}), + ("/permitted_models_info", {}), + ("/public_models_info", {}), + ("/is_public_model", {"model": "allowed"}), + ("/admin/reload-models", {}), + ], +) +async def test_body_key_routes_stay_human_only( + path: str, + body: dict[str, Any], + workload_authority: Authority, + client: httpx.AsyncClient, + monkeypatch: pytest.MonkeyPatch, + mocker: MockerFixture, +): + counter = mocker.create_autospec(TokenCounter, instance=True) + monkeypatch.setitem(server.app.dependency_overrides, server.get_token_counter, lambda: counter) + response = await client.post(path, json={**body, "api_key": workload_authority.token()}) + assert response.status_code == 401 + + +@pytest.mark.parametrize( + "method,path,body", + [ + ("GET", "/admin-check", None), + ("GET", "/model_groups?model=allowed", None), + ("GET", "/permitted_models_for_groups?group=same-group", None), + ("POST", "/model_approvals/resolve", {"model_names": ["allowed"]}), + ("GET", "/admin/models/discovery", None), + ("GET", "/admin/priority", None), + ], +) +async def test_bearer_control_routes_stay_human_only( + method: str, + path: str, + body: dict[str, Any] | None, + workload_authority: Authority, + client: httpx.AsyncClient, +): + response = await client.request( + method, + path, + headers={ + "authorization": f"Bearer {workload_authority.token()}", + }, + json=body, + ) + assert response.status_code == 401 + + +async def test_catalog_filters_approvals_and_redacts_secret_lab( + workload_authority: Authority, client: httpx.AsyncClient, monkeypatch: pytest.MonkeyPatch +): + monkeypatch.setenv("MIDDLEMAN_MODEL_APPROVAL_BINDING_KEY", "ab" * 32) + store = install_models(monkeypatch) + allowed = dataclasses.replace(store.models["allowed"], are_details_secret=True) + dead = dataclasses.replace(allowed, public_name="dead", dead=True) + store.models.update(allowed=allowed, dead=dead) + response = await client.get( + "/openai/v1/models", + headers={ + "authorization": f"Bearer {workload_authority.token((allowed, dead))}", + }, + ) + assert response.status_code == 200 + assert [(model["id"], model["owned_by"]) for model in response.json()["data"]] == [("allowed", "middleman")] + + +@pytest.mark.usefixtures("mock_environment") +@pytest.mark.parametrize( + "executor,lab,field,content", + [ + ( + OpenaiRequestExecutor("https://api.openai.com/v1", "responses"), + "openai-chat", + "input", + [{"role": "user", "content": [{"type": "input_image", "image_url": "data:image/png;base64,aW1hZ2U="}]}], + ), + ( + AnthropicRequestExecutor(), + "anthropic-chat", + "messages", + [ + { + "role": "user", + "content": [ + {"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": "aW1hZ2U="}} + ], + } + ], + ), + ( + GeminiDeveloperApiRequestExecutor("generateContent"), + "gemini-developer-api", + "contents", + [{"role": "user", "parts": [{"inlineData": {"mimeType": "image/png", "data": "aW1hZ2U="}}]}], + ), + ], +) +async def test_inline_images_preserved( + executor: RequestExecutor, + lab: str, + field: str, + content: list[dict[str, Any]], + workload_authority: Authority, + client: httpx.AsyncClient, + upstream: AsyncMock, + monkeypatch: pytest.MonkeyPatch, +): + model = install_models(monkeypatch, lab).models["allowed"] + request = executor.build_request("allowed", workload_authority.token((model,))) + response = await client.post(request.path, headers=request.headers, json={**request.body, field: content}) + assert response.status_code == 200 + assert upstream.call_args is not None + assert upstream.call_args.kwargs["json"][field] == content + + +@pytest.mark.usefixtures("mock_environment") +async def test_registry_reload_cannot_change_an_authorized_dispatch( + workload_authority: Authority, client: httpx.AsyncClient, upstream: AsyncMock, monkeypatch: pytest.MonkeyPatch +): + original = install_models(monkeypatch).models["allowed"] + select = passthrough.validate_model_access + + async def select_then_reload(model_names: list[str], principal: RequestPrincipal): + selected = await select(model_names, principal) + replacement = install_models(monkeypatch) + replacement.models["allowed"] = dataclasses.replace(original, danger_name="unapproved-target") + return selected + + monkeypatch.setattr(passthrough, "validate_model_access", select_then_reload) + response = await client.post( + "/openai/v1/chat/completions", + headers={ + "authorization": f"Bearer {workload_authority.token((original,))}", + }, + json={"model": "allowed", "messages": []}, + ) + assert response.status_code == 200 + assert upstream.call_args is not None + assert upstream.call_args.kwargs["json"]["model"] == "upstream-allowed" + + +@pytest.mark.usefixtures("mock_environment") +@pytest.mark.parametrize( + "executor,lab,error_type", + [ + (AnthropicRequestExecutor(), "anthropic-chat", "api_error"), + (OpenaiRequestExecutor("https://api.openai.com/v1", "responses"), "openai-chat", "server_error"), + (GeminiDeveloperApiRequestExecutor("generateContent"), "gemini-developer-api", "UNAVAILABLE"), + ], +) +async def test_issuer_outage_is_retryable_on_each_provider( + executor: RequestExecutor, + lab: str, + error_type: str, + workload_authority: Authority, + client: httpx.AsyncClient, + upstream: AsyncMock, + monkeypatch: pytest.MonkeyPatch, +): + model = install_models(monkeypatch, lab).models["allowed"] + token = workload_authority.token((model,)) + workload_authority.unavailable = True + request = executor.build_request("allowed", token) + response = await client.post(request.path, headers=request.headers, json=request.body) + assert response.status_code == 503 + assert response.headers["retry-after"] == "5" + assert error_type in response.json()["error"].values() + assert "Workload issuer unavailable" in response.text + assert token not in response.text + upstream.assert_not_called() + + +@pytest.mark.usefixtures("mock_environment") +async def test_expiry_is_not_extended_by_warm_keys( + workload_authority: Authority, client: httpx.AsyncClient, upstream: AsyncMock, monkeypatch: pytest.MonkeyPatch +): + now = int(time.time()) + model = install_models(monkeypatch).models["allowed"] + token = workload_authority.token((model,), overrides={"iat": now, "exp": now + 1800}) + headers = {"authorization": f"Bearer {token}"} + body = {"model": "allowed", "input": "hello"} + first = await client.post("/openai/v1/responses", headers=headers, json=body) + assert first.status_code == 200 + upstream.reset_mock() + monkeypatch.setattr(workload_jwt, "_now", lambda: now + 1860) + expired = await client.post("/openai/v1/responses", headers=headers, json=body) + assert expired.status_code == 401 + upstream.assert_not_called() + + +@pytest.mark.parametrize("problem", ["signature", "mounted_key", "audience", "null_attribution"]) +async def test_invalid_workload_credential_cannot_invoke( + problem: str, + workload_authority: Authority, + client: httpx.AsyncClient, + upstream: AsyncMock, + monkeypatch: pytest.MonkeyPatch, + mocker: MockerFixture, +): + model = install_models(monkeypatch).models["allowed"] + human_auth = mocker.patch.object( + passthrough, + "get_user_info", + autospec=True, + return_value=auth.UserInfo(id="human-id", groups=[model.group]), + ) + if problem == "signature": + header, payload, signature = workload_authority.token((model,)).split(".") + raw = bytearray(base64.urlsafe_b64decode(signature + "==")) + raw[0] ^= 1 + bad_signature = base64.urlsafe_b64encode(raw).decode().rstrip("=") + token = f"{header}.{payload}.{bad_signature}" + elif problem == "mounted_key": + token = "hawk_wl_v1.12345678-9abc-def0-1234-56789abcdef2." + "a" * 43 + elif problem == "audience": + token = workload_authority.token((model,), overrides={"aud": "foreign"}) + else: + token = workload_authority.token((model,), payload_overrides={"usage_attribution": None}) + response = await client.post( + "/openai/v1/responses", + headers={"authorization": f"Bearer {token}"}, + json={"model": "allowed", "input": "hello"}, + ) + assert response.status_code == 401 + assert token not in response.text + upstream.assert_not_called() + human_auth.assert_not_called() + + +@pytest.mark.usefixtures("mock_environment") +async def test_human_inference_survives_workload_issuer_outage( + workload_authority: Authority, + client: httpx.AsyncClient, + upstream: AsyncMock, + monkeypatch: pytest.MonkeyPatch, + mocker: MockerFixture, +): + from tests.conftest import TEST_AUDIENCE, TEST_ISSUER, make_test_jwt + + install_models(monkeypatch) + workload_authority.unavailable = True + key, token = make_test_jwt(is_admin=False) + monkeypatch.setenv( + "MIDDLEMAN_AUTH_PROVIDERS", + json.dumps( + [ + { + "issuer": TEST_ISSUER, + "audiences": [TEST_AUDIENCE], + "jwks_uri": f"{TEST_ISSUER}jwks.json", + "default_groups": ["same-group"], + } + ] + ), + ) + auth.load_auth_providers.cache_clear() + mocker.patch.object(auth, "_fetch_jwks", autospec=True, return_value={"keys": [key.as_dict(private=False)]}) + response = await client.post( + "/openai/v1/responses", + headers={"authorization": f"Bearer {token}"}, + json={"model": "allowed", "input": "hello", "background": True}, + ) + assert response.status_code == 200 + assert upstream.call_args is not None + assert upstream.call_args.kwargs["json"]["background"] is True + assert workload_authority.requests == [] + + +@pytest.mark.usefixtures("mock_environment") +@pytest.mark.parametrize( + "target_field,change", + [ + ("danger_name", "unchanged"), + ("danger_name", "retarget"), + ("private_name", "retarget"), + ("danger_name", "missing-key"), + ("danger_name", "malformed-key"), + ], +) +async def test_secret_model_approval_checked_before_provider_forwarding( + workload_authority: Authority, + client: httpx.AsyncClient, + upstream: AsyncMock, + monkeypatch: pytest.MonkeyPatch, + target_field: str, + change: str, +): + monkeypatch.setenv("MIDDLEMAN_MODEL_APPROVAL_BINDING_KEY", "ab" * 32) + definitions = [ + { + "public_name": "secret", + target_field: "secret-target-v1", + "lab": "openai-chat", + "group": "g", + "are_details_secret": True, + }, + {"public_name": "public", "danger_name": "public-target", "lab": "openai-chat", "group": "g"}, + ] + store = models.Models([dict(definition) for definition in definitions], {}) + token = workload_authority.token(tuple(store.models.values())) + if change == "retarget": + definitions[0][target_field] = "secret-target-v2" + store = models.Models(definitions, {}) + elif change == "missing-key": + monkeypatch.delenv("MIDDLEMAN_MODEL_APPROVAL_BINDING_KEY") + elif change == "malformed-key": + monkeypatch.setenv("MIDDLEMAN_MODEL_APPROVAL_BINDING_KEY", "invalid") + monkeypatch.setattr(models, "_current_models", store) + monkeypatch.setattr(models, "_models_loaded_at", float("inf")) + headers = {"authorization": f"Bearer {token}"} + response = await client.post( + "/openai/v1/chat/completions", headers=headers, json={"model": "secret", "messages": []} + ) + if change == "unchanged": + assert response.status_code == 200 + assert upstream.call_args is not None + assert upstream.call_args.kwargs["json"]["model"] == "secret-target-v1" + else: + assert response.status_code == 404 + assert "secret-target" not in response.text + upstream.assert_not_called() + upstream.reset_mock() + response = await client.post( + "/openai/v1/chat/completions", headers=headers, json={"model": "public", "messages": []} + ) + assert response.status_code == 200 + assert upstream.call_args is not None + assert upstream.call_args.kwargs["json"]["model"] == "public-target" + + +@pytest.mark.usefixtures("mock_environment") +async def test_attribution_uses_signed_execution_not_correlation_headers( + workload_authority: Authority, + traffic_client: tuple[httpx.AsyncClient, MagicMock], + upstream: AsyncMock, + monkeypatch: pytest.MonkeyPatch, +): + client, emitter = traffic_client + model = install_models(monkeypatch).models["allowed"] + token = workload_authority.token( + (model,), + payload_overrides={ + "usage_attribution": {"user_id": "human-id", "teams": ["alpha", "zeta"]}, + }, + ) + response = await client.post( + "/openai/v1/responses", + headers={ + "authorization": f"Bearer {token}", + "x-hawk-job-type": "forged-type", + "x-hawk-job-id": "forged-job", + "x-metr-user": "forged-user", + "x-metr-team": "administrators", + }, + json={"model": "allowed", "input": "hello", "background": False, "stream": True}, + ) + assert response.status_code == 200 + assert upstream.call_args is not None + call = upstream.call_args.kwargs + assert call["channel"] == "eval-set" + emitter.enqueue.assert_called_once() + fields = emitter.enqueue.call_args.kwargs["envelope"].model_dump() + assert fields["user_id"] == f"workload:{EXECUTION}" + assert fields["usage_user_id"] == "human-id" + assert auth.principal_subject(call["user"]) == f"workload:{EXECUTION}" + assert fields["workload_execution_id"] == EXECUTION + assert fields["workload_grant_id"] == "12345678-9abc-def0-1234-56789abcdef1" + assert fields["workload_job_id"] == "own-output" + assert fields["workload_job_type"] == "eval-set" + assert fields["is_admin"] is False + assert fields["user_groups"] is None + assert fields["user_teams"] is None + assert fields["usage_teams"] == ["alpha", "zeta"] + + +@pytest.mark.usefixtures("mock_environment") +@pytest.mark.parametrize("teams", [[], ["platform"]]) +async def test_human_traffic_keeps_auth_identity_and_teams( + traffic_client: tuple[httpx.AsyncClient, MagicMock], + upstream: AsyncMock, + monkeypatch: pytest.MonkeyPatch, + mocker: MockerFixture, + teams: list[str], +): + client, emitter = traffic_client + install_models(monkeypatch) + mocker.patch.object( + passthrough, + "get_user_info", + autospec=True, + return_value=auth.UserInfo(id="human-id", groups=["same-group"], teams=teams), + ) + response = await client.post( + "/openai/v1/responses", + headers={"authorization": "Bearer human-token"}, + json={"model": "allowed", "input": "hello"}, + ) + assert response.status_code == 200 + upstream.assert_called_once() + emitter.enqueue.assert_called_once() + fields = emitter.enqueue.call_args.kwargs["envelope"].model_dump() + assert fields["user_id"] == "human-id" + assert fields["user_groups"] == ["same-group"] + assert fields["user_teams"] == teams + assert fields["usage_user_id"] is None + assert fields["usage_teams"] is None + assert fields["workload_execution_id"] is None + + +@pytest.mark.usefixtures("mock_environment") +@pytest.mark.parametrize( + "job_type,forge_headers,teams,expected_team", + [ + pytest.param("eval-set", True, [], "unassigned", id="eval-forged-headers-no-teams"), + pytest.param("scan", False, ["alpha", "zeta"], "alpha+zeta", id="scan-no-headers-multiple-teams"), + ], +) +async def test_signed_usage_attribution_reaches_emf_without_changing_trace_identity( + workload_authority: Authority, + traffic_client: tuple[httpx.AsyncClient, MagicMock], + monkeypatch: pytest.MonkeyPatch, + mocker: MockerFixture, + teams: list[str], + expected_team: str, + job_type: str, + forge_headers: bool, +): + client, traffic = traffic_client + model = install_models(monkeypatch).models["allowed"] + provider_body = b'{"choices":[],"usage":{"prompt_tokens":3,"completion_tokens":5,"total_tokens":8}}' + + async def content(): + yield provider_body + + upstream_response = mocker.create_autospec(aiohttp.ClientResponse, instance=True) + upstream_response.status = 200 + upstream_response.headers = {"content-type": "application/json"} + upstream_response.content.iter_any.return_value = content() + upstream_response.release = mocker.async_stub() + session = mocker.create_autospec(aiohttp.ClientSession, instance=True) + session.post = mocker.AsyncMock(return_value=upstream_response) + mocker.patch.object(passthrough, "get_client_session", autospec=True, return_value=session) + mocker.patch.object(passthrough, "record_upstream_duration", autospec=True) + monkeypatch.setenv("MIDDLEMAN_METRICS_LOG_GROUP", "test-emf") + emitter = emf.EmfEmitter() + monkeypatch.setattr(passthrough, "emf_emitter", emitter) + tracer = mocker.patch("middleman.passthrough.otel_trace.get_tracer", autospec=True).return_value + token = workload_authority.token( + (model,), + payload_overrides={ + "usage_attribution": {"user_id": "human-id", "teams": teams}, + "job": {"type": job_type, "id": "own-output"}, + "permissions": [ + f"hawk:read@{'scan-runs' if job_type == 'scan' else 'eval-sets'}/own-output", + "middleman:call@models/allowed", + ], + }, + ) + response = await client.post( + "/openai/v1/chat/completions", + headers={ + "authorization": f"Bearer {token}", + **( + { + "x-metr-user": "forged-user", + "x-metr-team": "administrators", + "x-hawk-job-type": "scan", + "x-hawk-job-id": "forged-job", + } + if forge_headers + else {} + ), + }, + json={"model": "allowed", "messages": [], "stream": False}, + ) + assert response.status_code == 200 + assert response.content == provider_body + queue = emitter._queue # pyright: ignore[reportPrivateUsage] + records = [queue.get_nowait() for _ in range(queue.qsize())] + usage_records = [record for record in records if "InputTokens" in record] + assert len(usage_records) == 1 + usage = usage_records[0] + assert usage["user"] == "human-id" + assert usage["team"] == expected_team + assert usage["channel"] == job_type + assert usage["InputTokens"] == 3 + assert usage["OutputTokens"] == 5 + dimensions = usage["_aws"]["CloudWatchMetrics"][0]["Dimensions"] + assert ["provider", "model", "user"] in dimensions + assert ["provider", "model", "team"] in dimensions + traffic.enqueue.assert_called_once() + fields = traffic.enqueue.call_args.kwargs["envelope"].model_dump() + assert fields["input_tokens"] == 3 + assert fields["output_tokens"] == 5 + assert fields["user_id"] == f"workload:{EXECUTION}" + assert fields["workload_execution_id"] == EXECUTION + assert fields["workload_job_id"] == "own-output" + assert fields["workload_job_type"] == job_type + assert fields["usage_user_id"] == "human-id" + assert fields["usage_teams"] == teams + assert fields["user_groups"] is None + assert fields["is_admin"] is False + if not forge_headers: + assert fields["correlation"] == {} + tracer.start_as_current_span.return_value.__enter__.return_value.set_attribute.assert_any_call( + "hawk.user.id", + f"workload:{EXECUTION}", + ) diff --git a/middleman/tests/traffic_log/test_middleware.py b/middleman/tests/traffic_log/test_middleware.py index e123672836..2b24ddc961 100644 --- a/middleman/tests/traffic_log/test_middleware.py +++ b/middleman/tests/traffic_log/test_middleware.py @@ -1,6 +1,7 @@ from __future__ import annotations import asyncio +import json from collections.abc import AsyncIterator from typing import Any from unittest.mock import MagicMock @@ -787,3 +788,75 @@ async def send(_message: Any) -> None: assert pulled_before_app == [1], "prefill read past the cap — residency is no longer bounded" assert int(pulled) == 3, "the app must still receive the whole body" + + +def test_full_capture_redacts_workload_credentials_on_rejection(): + token = "distinctive-workload-credential" + mounted_key = "hawk_wl_v1.12345678-9abc-def0-1234-56789abcdef2." + "a" * 43 + emitter, enqueued = _mock_emitter() + app = fastapi.FastAPI() + app.add_middleware(TrafficLogMiddleware, env="test", level=Level.FULL, emitter=emitter) + + @app.post("/rejected") + async def rejected() -> fastapi.responses.JSONResponse: # pyright: ignore[reportUnusedFunction] + return fastapi.responses.JSONResponse({"detail": "invalid api key"}, status_code=401) + + response = fastapi.testclient.TestClient(app).post( + "/rejected", + headers={"authorization": f"Bearer {token}", "x-api-key": token, "x-goog-api-key": token}, + json={"api_key": mounted_key, "model": "allowed"}, + ) + assert response.status_code == 401 + captured = enqueued[0] + assert token not in json.dumps(captured["request_payload"]) + assert mounted_key not in json.dumps(captured["request_payload"]) + assert token not in captured["envelope"].model_dump_json() + assert mounted_key not in captured["envelope"].model_dump_json() + + +@pytest.mark.parametrize( + "capture,key_name", + [("malformed", '"api_key"'), ("after_key", r'"api_\u006bey"'), ("within_key", '"api_key"')], +) +def test_full_capture_redacts_credentials_in_incomplete_json(key_name: str, capture: str): + token = "test-secret-value" + prefix = '{"model":"allowed",' + key_name + ':"' + body = prefix + token + '","payload":"' + "x" * 200 + '"}' + if capture == "malformed": + body = prefix + token + '",' + cap = len(body) + 1 + elif capture == "after_key": + cap = len(prefix) + len(token) + 30 + else: + cap = len(prefix) + len(token) // 2 + emitter, enqueued = _mock_emitter() + app = fastapi.FastAPI() + app.add_middleware(TrafficLogMiddleware, env="test", level=Level.FULL, emitter=emitter, request_body_cap_bytes=cap) + + @app.post("/rejected") + async def rejected() -> fastapi.responses.JSONResponse: # pyright: ignore[reportUnusedFunction] + return fastapi.responses.JSONResponse({"detail": "invalid api key"}, status_code=401) + + response = fastapi.testclient.TestClient(app).post( + "/rejected", content=body, headers={"content-type": "application/json"} + ) + assert response.status_code == 401 + captured = enqueued[0]["request_payload"]["body"] + assert token[: len(token) // 2] not in captured + assert "[REDACTED]" in captured + assert '"model":"allowed"' in captured + + +def test_incomplete_json_redaction_uses_bounded_memory(): + import tracemalloc + + from middleman.traffic_log.middleware import _redact_api_key # pyright: ignore[reportPrivateUsage] + + body = '{"payload":"' + "x" * (2 * 1024 * 1024) + tracemalloc.start() + try: + assert _redact_api_key(body) == body + _, peak = tracemalloc.get_traced_memory() + assert peak < len(body) + finally: + tracemalloc.stop() diff --git a/middleman/tests/workload_support.py b/middleman/tests/workload_support.py new file mode 100644 index 0000000000..dbf2e9217b --- /dev/null +++ b/middleman/tests/workload_support.py @@ -0,0 +1,81 @@ +from __future__ import annotations + +import json +import time +from dataclasses import dataclass, field +from typing import Any + +import httpx +from hawk.core.types import permission_atoms +from joserfc import jwk, jwt + +from middleman.model_approvals import approval_for_model +from middleman.model_info import ModelInfo + +ISSUER = "https://hawk.example/workload" +KID = "12345678-9abc-def0-1234-56789abcdef3" +EXECUTION = "12345678-9abc-def0-1234-56789abcdef0" +SETTINGS: dict[str, Any] = { + "issuer": ISSUER, + "audience": f"{ISSUER}/services", + "jwks_uri": "https://hawk.example/.well-known/jwks.json", +} + + +@dataclass +class Authority: + key: jwk.RSAKey = field(default_factory=lambda: jwk.RSAKey.generate_key(2048, parameters={"kid": KID})) + requests: list[httpx.Request] = field(default_factory=list) + unavailable: bool = False + + def endpoint(self, request: httpx.Request) -> httpx.Response: + self.requests.append(request) + assert str(request.url) == SETTINGS["jwks_uri"] + if self.unavailable: + return httpx.Response(503) + return httpx.Response( + 200, + stream=httpx.ByteStream(json.dumps({"keys": [self.key.as_dict(private=False)]}).encode()), + ) + + def token( + self, + models: tuple[ModelInfo, ...] = (), + *, + overrides: dict[str, Any] | None = None, + payload_overrides: dict[str, Any] | None = None, + ) -> str: + now = int(time.time()) + payload: dict[str, Any] = { + "version": 2, + "execution_id": EXECUTION, + "grant_id": "12345678-9abc-def0-1234-56789abcdef1", + "job": {"type": "eval-set", "id": "own-output"}, + "kubernetes": { + "namespace_name": "runner-1234", + "service_account_name": "runner", + "kubernetes_job_name": "execution-1234", + }, + "permissions": ["hawk:read@eval-sets/own-output"] + + [permission_atoms.format_atom("model_call", model.public_name) for model in models], + "model_approvals": [approval_for_model(model).model_dump() for model in models], + "usage_attribution": {"user_id": "human-id", "teams": []}, + **(payload_overrides or {}), + } + claims = { + "iss": ISSUER, + "aud": SETTINGS["audience"], + "sub": f"workload:{EXECUTION}", + "client_id": "12345678-9abc-def0-1234-56789abcdef2", + "iat": now, + "exp": now + 1800, + "jti": "12345678-9abc-def0-1234-56789abcdef4", + "workload": payload, + **(overrides or {}), + } + return jwt.encode( + {"alg": "RS256", "typ": "at+jwt", "kid": KID}, + claims, + self.key, + algorithms=["RS256"], + )