diff --git a/.repository-projection.json b/.repository-projection.json index 1672e75..9450f70 100644 --- a/.repository-projection.json +++ b/.repository-projection.json @@ -3,11 +3,11 @@ "projection": "deixic-python", "projectionSchemaVersion": 1, "sourceRepository": "dx-corp/mono", - "sourceSha": "ddaee8a19ef06ccd953aa67505182ec41e78f13c", + "sourceSha": "80416a066acccdb2394733a9db4adcae19d88f23", "destinationRepository": "dx-corp/deixic-python", - "priorProjectedBase": "334735fa42a46ba1f3bd05f768a31d3bcbcda394", - "definitionDigest": "01efa736c67353b6c6e353d7e2ab6886ecd21b68f31fe1a89234403247abcbdb", - "toolDigest": "0aae6000dbd0940f5a0af380463ccb5a83285eda", - "contentDigest": "14e4cd30768fa913750e0b17c4bca966269f33b72c226ac8924df966c045fbcf", + "priorProjectedBase": "be2ecd40236afbb15ee257b69e86ec6832226803", + "definitionDigest": "eedaf99642ef977b56b6e12ce151e1834670560f00294e89224f0628882937d9", + "toolDigest": "cef8579e533dbfa40cbd9070dd7581f65f800e5f", + "contentDigest": "13e942ba6333d3a42f22f7aeb9a324042fc4e8a56c52eaa28c67814b696f8a3b", "publicationEligible": true } diff --git a/CHANGELOG.md b/CHANGELOG.md index be02039..6fa4cd9 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -2,6 +2,8 @@ ## Unreleased +- Add `WorkloadFederationCredentialProvider` with GitHub Actions, file, and + environment-variable assertion sources for keyless workload authentication. - Report missing or removed model routes during setup instead of declaring the workspace accessible, and preserve plain-text account briefs when account names contain structured-output instruction text. Keep rejected credential diff --git a/README.md b/README.md index d4cf1be..1d5b87c 100644 --- a/README.md +++ b/README.md @@ -219,6 +219,77 @@ provide a `CredentialProvider`; an authentication replay is allowed only when the refreshed credential retains the same subject, tenant, and declared scopes. +## Workload federation + +A CI job or cloud workload can authenticate without a stored API key. An +organization admin first registers an issuer, a service account, and a rule in +Identity settings, as described in the +[workload federation guide](https://github.com/dx-corp/mono/blob/main/docs/services/identity/workload-federation.md). +The workload then exchanges a signed assertion from its platform for a Deixic +access token that lasts at most 300 seconds: + +```python +import os + +from deixic import ( + Deixic, + WorkloadFederationCredentialProvider, + github_actions_assertion_source, +) + +identity_url = os.environ["DEIXIC_IDENTITY_URL"] +credentials = WorkloadFederationCredentialProvider( + identity_url=identity_url, + assertion_source=github_actions_assertion_source( + audience=f"{identity_url}/v1/workload-federation/exchange", + ), +) +deixic = Deixic( + credential_provider=credentials, + organization_id="org_123", + workspace_id="ws_456", +) +``` + +The provider exchanges an assertion on the first request and caches the token. +It exchanges again 60 seconds before expiry (`refresh_margin`) and after an +HTTP 401 from Deixic. Identity accepts each assertion once, so every exchange +calls the assertion source for a new assertion. The provider refuses to send +an assertion it has already exchanged and raises `DeixicError` with code +`workload_assertion_reused`. When an early refresh cannot obtain a new +assertion and the cached token has not expired, the provider keeps using the +cached token. + +Assertion sources: + +- `github_actions_assertion_source(audience)` requests a new OIDC token from + GitHub Actions on each call. The job needs `permissions: id-token: write`. +- `file_assertion_source(path)` reads a file on each call, such as a Kubernetes + projected service-account token. The kubelet rewrites that file after 80% of + the token's `expirationSeconds`. A file token can therefore be exchanged once + per rotation. Set `expirationSeconds` so that rotation happens more often + than the 300-second Deixic token lifetime, or pass a callable that requests + a new token from the Kubernetes TokenRequest API. +- `environment_assertion_source(name)` reads an environment variable on each + call. The application must write a new value before each exchange. +- Any zero-argument callable that returns a new JWT string. + +Exchange failures raise `DeixicError`: + +| HTTP status | `kind` | `code` | Retried | +| --- | --- | --- | --- | +| 400 | `validation` | `workload_assertion_invalid` | No | +| 403 | `authorization` | `workload_federation_forbidden` | No | +| 409 | `conflict` | `workload_assertion_replayed` | No | +| 503 | `unavailable` | `workload_federation_unavailable` | Yes | +| Transport failure | `transport` | `workload_exchange_transport` | Yes | + +A retry waits `retry_delay` seconds, doubled on each attempt, up to +`max_attempts` total attempts (default 3). Each retry uses a new assertion. +A 403 means no active rule matched the assertion's issuer, audience, subject, +and claims. The provider never logs the assertion or the access token and +omits both from error messages. + ## Coding output readback `messages.send(coding_acceptance=contract)` sends the typed coding contract and diff --git a/src/deixic/__init__.py b/src/deixic/__init__.py index e22d5c4..fdf4a95 100644 --- a/src/deixic/__init__.py +++ b/src/deixic/__init__.py @@ -1,9 +1,17 @@ from .auth import Credential, CredentialProvider from .client import Deixic from .errors import DeixicError +from .federation import ( + AssertionSource, + WorkloadFederationCredentialProvider, + environment_assertion_source, + file_assertion_source, + github_actions_assertion_source, +) from .tasks import SetupCheck, Task, TaskCheckpoint, TaskResult __all__ = [ + "AssertionSource", "Credential", "CredentialProvider", "Deixic", @@ -12,4 +20,8 @@ "Task", "TaskCheckpoint", "TaskResult", + "WorkloadFederationCredentialProvider", + "environment_assertion_source", + "file_assertion_source", + "github_actions_assertion_source", ] diff --git a/src/deixic/federation.py b/src/deixic/federation.py new file mode 100644 index 0000000..500215c --- /dev/null +++ b/src/deixic/federation.py @@ -0,0 +1,466 @@ +"""Workload federation credentials for CI and cloud workloads. + +A workload presents a short-lived signed assertion (an OIDC JWT issued by its +CI system or cloud platform) to Identity's +``/v1/workload-federation/exchange`` endpoint and receives a tenant-bound +access token valid for at most 300 seconds. Identity accepts each assertion +once, so every exchange here obtains a new assertion from the configured +source. + +Assertions and access tokens are never logged or included in error messages. +""" + +from __future__ import annotations + +import hashlib +import json +import math +import os +import threading +import time +from collections import deque +from collections.abc import Callable, Mapping +from datetime import datetime, timezone +from pathlib import Path +from typing import Any +from urllib.parse import urlencode, urlsplit + +from .auth import Credential +from .errors import DeixicError, validation_error +from .transport import RequestsTransport, Response, Transport + +AssertionSource = Callable[[], str] +"""Zero-argument callable returning a fresh signed workload assertion.""" + +EXCHANGE_PATH = "/v1/workload-federation/exchange" + +_REMEMBERED_ASSERTIONS = 16 +_REUSED = "workload_assertion_reused" +_UNAVAILABLE = "workload_federation_unavailable" +_TRANSPORT = "workload_exchange_transport" +_EARLY_REFRESH_TOLERATED = frozenset({_REUSED, _UNAVAILABLE, _TRANSPORT}) + + +class WorkloadFederationCredentialProvider: + """Exchange workload assertions for short-lived Deixic access tokens. + + The first ``get_credential`` call performs an exchange. The token is + cached and exchanged again ``refresh_margin`` seconds before it expires, + or after half its lifetime when the token lives less than twice the + margin. Every exchange requests a new assertion from ``assertion_source``. + + Only HTTP 503 and transport failures are retried, each attempt with a new + assertion. A 400, 403, or 409 (replayed assertion) response is raised + immediately. + """ + + can_refresh = True + + def __init__( + self, + *, + identity_url: str, + assertion_source: AssertionSource, + refresh_margin: float = 60.0, + timeout: float = 10.0, + max_attempts: int = 3, + retry_delay: float = 0.5, + transport: Transport | None = None, + clock: Callable[[], float] = time.time, + sleep: Callable[[float], None] = time.sleep, + ) -> None: + self._exchange_url = _validated_identity_url(identity_url) + EXCHANGE_PATH + if not callable(assertion_source): + raise validation_error("assertion_source must be a callable") + self._assertion_source = assertion_source + self._refresh_margin = _non_negative(refresh_margin, "refresh_margin") + self._timeout = _positive(timeout, "timeout") + self._retry_delay = _non_negative(retry_delay, "retry_delay") + if ( + isinstance(max_attempts, bool) + or not isinstance(max_attempts, int) + or not 1 <= max_attempts <= 10 + ): + raise validation_error("max_attempts must be an integer between 1 and 10") + self._max_attempts = max_attempts + self._transport = transport or RequestsTransport() + self._clock = clock + self._sleep = sleep + self._lock = threading.Lock() + self._credential: Credential | None = None + self._refresh_at = 0.0 + self._expires_at = 0.0 + self._used_assertions: deque[str] = deque(maxlen=_REMEMBERED_ASSERTIONS) + + def __repr__(self) -> str: + return ( + f"WorkloadFederationCredentialProvider(exchange_url={self._exchange_url!r})" + ) + + def get_credential(self) -> Credential: + with self._lock: + cached = self._credential + now = self._clock() + if cached is not None and now < self._refresh_at: + return cached + try: + return self._exchange() + except DeixicError as error: + # An early refresh that fails transiently keeps the unexpired + # token. A file source that has not rotated yet lands here. + if ( + cached is not None + and now < self._expires_at + and error.code in _EARLY_REFRESH_TOLERATED + ): + return cached + raise + + def refresh_credential(self, current: Credential) -> Credential: + with self._lock: + cached = self._credential + if ( + cached is not None + and cached.access_token != current.access_token + and self._clock() < self._refresh_at + ): + # Another caller already replaced the rejected token. + return cached + return self._exchange() + + def _exchange(self) -> Credential: + last_error: DeixicError | None = None + for attempt in range(self._max_attempts): + if last_error is not None: + self._sleep(self._retry_delay * (2 ** (attempt - 1))) + try: + assertion = self._next_assertion() + except DeixicError as error: + # A retry whose source has no new assertion reports the + # original unavailable or transport failure. + if last_error is not None and error.code == _REUSED: + raise last_error from None + raise + try: + return self._exchange_once(assertion) + except DeixicError as error: + if error.code not in {_UNAVAILABLE, _TRANSPORT}: + raise + last_error = error + assert last_error is not None + raise last_error + + def _next_assertion(self) -> str: + try: + assertion = self._assertion_source() + except DeixicError: + raise + except Exception as exc: + raise _assertion_unavailable("workload assertion source failed") from exc + if not isinstance(assertion, str) or not assertion.strip(): + raise _assertion_unavailable( + "workload assertion source returned an empty assertion" + ) + assertion = assertion.strip() + digest = hashlib.sha256(assertion.encode("utf-8")).hexdigest() + if digest in self._used_assertions: + raise DeixicError( + "workload assertion source returned an assertion that was " + "already exchanged; Identity accepts each assertion once", + kind="authentication", + code=_REUSED, + ) + self._used_assertions.append(digest) + return assertion + + def _exchange_once(self, assertion: str) -> Credential: + body = json.dumps({"assertion": assertion}).encode("utf-8") + try: + response = self._transport.send( + "POST", + self._exchange_url, + headers={ + "Accept": "application/json", + "Content-Type": "application/json", + }, + body=body, + timeout=self._timeout, + ) + except DeixicError: + raise + except Exception as exc: + raise DeixicError( + "Identity workload exchange transport failed", + kind="transport", + code=_TRANSPORT, + ) from exc + try: + status = int(response.status_code) + if not 200 <= status < 300: + raise _exchange_error(status, response) + return self._credential_from(response) + finally: + response.close() + + def _credential_from(self, response: Response) -> Credential: + try: + payload = json.loads(bytes(response.content).decode("utf-8")) + except (UnicodeDecodeError, json.JSONDecodeError): + payload = None + if not isinstance(payload, Mapping): + raise _invalid_response("Identity returned an invalid exchange response") + access_token = _field(payload, "access_token") + token_type = _field(payload, "token_type") + principal_id = _field(payload, "principal_id") + organization_id = _field(payload, "organization_id") + workspace_id = _field(payload, "workspace_id") + expires_at = _parse_expiry(_field(payload, "expires_at")) + scopes = payload.get("scopes") + if not isinstance(scopes, list) or not all( + isinstance(scope, str) for scope in scopes + ): + raise _invalid_response("Identity exchange response has invalid scopes") + now = self._clock() + lifetime = max(0.0, expires_at - now) + self._refresh_at = now + max(lifetime - self._refresh_margin, lifetime / 2) + self._expires_at = now + lifetime + self._credential = Credential( + access_token=access_token, + token_type=token_type, + subject=principal_id, + organization_id=organization_id, + workspace_id=workspace_id, + scopes=tuple(scopes), + ) + return self._credential + + +def github_actions_assertion_source( + audience: str, + *, + transport: Transport | None = None, + timeout: float = 10.0, + environ: Mapping[str, str] | None = None, +) -> AssertionSource: + """Request a new GitHub Actions OIDC token on every call. + + The job needs ``permissions: id-token: write`` so that GitHub sets + ``ACTIONS_ID_TOKEN_REQUEST_URL`` and ``ACTIONS_ID_TOKEN_REQUEST_TOKEN``. + ``audience`` must equal the audience registered on the Identity issuer. + """ + + audience = audience.strip() if isinstance(audience, str) else "" + if not audience: + raise validation_error("audience is required") + timeout = _positive(timeout, "timeout") + client = transport or RequestsTransport() + env = os.environ if environ is None else environ + + def source() -> str: + request_url = (env.get("ACTIONS_ID_TOKEN_REQUEST_URL") or "").strip() + request_token = (env.get("ACTIONS_ID_TOKEN_REQUEST_TOKEN") or "").strip() + if not request_url or not request_token: + raise _assertion_unavailable( + "ACTIONS_ID_TOKEN_REQUEST_URL and ACTIONS_ID_TOKEN_REQUEST_TOKEN " + "are not set; grant the job the id-token: write permission" + ) + separator = "&" if urlsplit(request_url).query else "?" + url = f"{request_url}{separator}{urlencode({'audience': audience})}" + try: + response = client.send( + "GET", + url, + headers={ + "Accept": "application/json", + "Authorization": f"Bearer {request_token}", + }, + body=b"", + timeout=timeout, + ) + except DeixicError: + raise + except Exception: + raise _assertion_unavailable( + "GitHub Actions OIDC token request failed" + ) from None + try: + status = int(response.status_code) + if not 200 <= status < 300: + raise _assertion_unavailable( + f"GitHub Actions OIDC token request failed with HTTP {status}" + ) + try: + payload = json.loads(bytes(response.content).decode("utf-8")) + except (UnicodeDecodeError, json.JSONDecodeError): + payload = None + value = payload.get("value") if isinstance(payload, Mapping) else None + if not isinstance(value, str) or not value.strip(): + raise _assertion_unavailable( + "GitHub Actions OIDC token response has no value" + ) + return value.strip() + finally: + response.close() + + return source + + +def file_assertion_source(path: str | os.PathLike[str]) -> AssertionSource: + """Read the assertion from a file on every call. + + Use this for a Kubernetes projected service-account token or any file a + platform agent rewrites with a new token. + """ + + token_path = Path(path) + + def source() -> str: + try: + value = token_path.read_text(encoding="utf-8").strip() + except OSError: + raise _assertion_unavailable( + f"workload assertion file {str(token_path)!r} could not be read" + ) from None + if not value: + raise _assertion_unavailable( + f"workload assertion file {str(token_path)!r} is empty" + ) + return value + + return source + + +def environment_assertion_source( + name: str, *, environ: Mapping[str, str] | None = None +) -> AssertionSource: + """Read the assertion from an environment variable on every call.""" + + name = name.strip() if isinstance(name, str) else "" + if not name: + raise validation_error("environment variable name is required") + env = os.environ if environ is None else environ + + def source() -> str: + value = (env.get(name) or "").strip() + if not value: + raise _assertion_unavailable( + f"environment variable {name} does not contain a workload assertion" + ) + return value + + return source + + +def _exchange_error(status: int, response: Response) -> DeixicError: + kind, code, message = { + 400: ( + "validation", + "workload_assertion_invalid", + "Identity rejected the workload exchange request as malformed", + ), + 403: ( + "authorization", + "workload_federation_forbidden", + "Identity did not match the workload assertion to an active " + "federation rule", + ), + 409: ( + "conflict", + "workload_assertion_replayed", + "Identity already accepted this workload assertion", + ), + 503: ( + "unavailable", + _UNAVAILABLE, + "Identity workload federation is unavailable", + ), + }.get( + status, + ( + "protocol", + "workload_exchange_failed", + f"Identity workload exchange failed with HTTP {status}", + ), + ) + headers = getattr(response, "headers", {}) or {} + return DeixicError( + message, + kind=kind, + status_code=status, + code=code, + request_id=_header(headers, "x-request-id"), + traceparent=_header(headers, "traceparent"), + ) + + +def _assertion_unavailable(message: str) -> DeixicError: + return DeixicError( + message, kind="authentication", code="workload_assertion_unavailable" + ) + + +def _invalid_response(message: str) -> DeixicError: + return DeixicError(message, kind="protocol", code="workload_exchange_invalid") + + +def _field(payload: Mapping[str, Any], name: str) -> str: + value = payload.get(name) + if not isinstance(value, str) or not value.strip(): + raise _invalid_response(f"Identity exchange response is missing {name}") + return value.strip() + + +def _parse_expiry(value: str) -> float: + # Identity emits RFC 3339 UTC timestamps such as 2026-09-24T12:00:00Z. + # datetime.fromisoformat accepts a trailing "Z" only from Python 3.11. + text = value[:-1] + "+00:00" if value.endswith(("Z", "z")) else value + try: + parsed = datetime.fromisoformat(text) + except ValueError: + raise _invalid_response( + "Identity exchange response has an invalid expires_at" + ) from None + if parsed.tzinfo is None: + parsed = parsed.replace(tzinfo=timezone.utc) + return parsed.timestamp() + + +def _header(headers: Mapping[str, Any], name: str) -> str | None: + for key, value in headers.items(): + if str(key).lower() == name: + return str(value) + return None + + +def _validated_identity_url(value: str) -> str: + cleaned = value.strip().rstrip("/") if isinstance(value, str) else "" + parsed = urlsplit(cleaned) + if parsed.scheme not in {"http", "https"} or not parsed.netloc: + raise validation_error("identity_url must be an absolute HTTP(S) URL") + if parsed.username or parsed.password or parsed.query or parsed.fragment: + raise validation_error( + "identity_url must not contain credentials, a query, or a fragment" + ) + return cleaned + + +def _positive(value: float, field: str) -> float: + if ( + isinstance(value, bool) + or not isinstance(value, (int, float)) + or not math.isfinite(value) + or value <= 0 + ): + raise validation_error(f"{field} must be a positive finite number") + return float(value) + + +def _non_negative(value: float, field: str) -> float: + if ( + isinstance(value, bool) + or not isinstance(value, (int, float)) + or not math.isfinite(value) + or value < 0 + ): + raise validation_error(f"{field} must be a non-negative finite number") + return float(value) diff --git a/tests/test_workload_federation.py b/tests/test_workload_federation.py new file mode 100644 index 0000000..838e7ee --- /dev/null +++ b/tests/test_workload_federation.py @@ -0,0 +1,372 @@ +from __future__ import annotations + +import json +from collections.abc import Iterator, Mapping +from dataclasses import dataclass +from pathlib import Path +from typing import Any + +import pytest +from deixic import protocol as public_pb2 + +from deixic import ( + Deixic, + DeixicError, + WorkloadFederationCredentialProvider, + environment_assertion_source, + file_assertion_source, + github_actions_assertion_source, +) + +IDENTITY = "https://identity.deixic.test" +EXCHANGE = f"{IDENTITY}/v1/workload-federation/exchange" +NOW = 1_790_000_000.0 + + +@dataclass +class FakeResponse: + status_code: int + content: bytes = b"" + headers: Mapping[str, str] | None = None + closed: bool = False + + def json(self) -> Any: + return json.loads(self.content.decode("utf-8")) + + def iter_content(self, chunk_size: int | None = None) -> Iterator[bytes]: + return iter((self.content,)) + + def close(self) -> None: + self.closed = True + + def __post_init__(self) -> None: + if self.headers is None: + self.headers = {} + + +class FakeTransport: + def __init__(self, responses: list[FakeResponse | Exception]) -> None: + self.responses = responses + self.requests: list[dict[str, Any]] = [] + + def send( + self, + method: str, + url: str, + *, + headers: Mapping[str, str], + body: bytes, + timeout: float, + stream: bool = False, + ) -> FakeResponse: + self.requests.append( + { + "method": method, + "url": url, + "headers": dict(headers), + "body": body, + "timeout": timeout, + "stream": stream, + } + ) + item = self.responses.pop(0) + if isinstance(item, Exception): + raise item + return item + + +class Clock: + def __init__(self) -> None: + self.now = NOW + + def __call__(self) -> float: + return self.now + + +class Assertions: + def __init__(self) -> None: + self.calls = 0 + + def __call__(self) -> str: + self.calls += 1 + return f"assertion-{self.calls}" + + +def issued(token: str, *, lifetime: int = 300) -> FakeResponse: + from datetime import datetime, timezone + + expires = datetime.fromtimestamp(NOW + lifetime, tz=timezone.utc) + return FakeResponse( + 200, + json.dumps( + { + "access_token": token, + "token_type": "Bearer", + "expires_at": expires.strftime("%Y-%m-%dT%H:%M:%SZ"), + "organization_id": "org-a", + "workspace_id": "workspace-a", + "principal_id": "principal-a", + "rule_id": "rule-a", + "scopes": ["console:write"], + } + ).encode(), + {"Cache-Control": "no-store"}, + ) + + +def provider( + transport: FakeTransport, + source: Any, + clock: Clock | None = None, + sleeps: list[float] | None = None, +) -> WorkloadFederationCredentialProvider: + return WorkloadFederationCredentialProvider( + identity_url=IDENTITY + "/", + assertion_source=source, + transport=transport, + clock=clock or Clock(), + sleep=(sleeps.append if sleeps is not None else lambda _: None), + ) + + +def test_first_exchange_posts_assertion_and_caches_token() -> None: + transport = FakeTransport([issued("token-1")]) + source = Assertions() + credentials = provider(transport, source) + + credential = credentials.get_credential() + again = credentials.get_credential() + + assert again is credential + assert credential.access_token == "token-1" + assert credential.token_type == "Bearer" + assert credential.subject == "principal-a" + assert credential.organization_id == "org-a" + assert credential.workspace_id == "workspace-a" + assert credential.scopes == ("console:write",) + assert source.calls == 1 + assert len(transport.requests) == 1 + request = transport.requests[0] + assert request["method"] == "POST" + assert request["url"] == EXCHANGE + assert request["headers"]["Content-Type"] == "application/json" + assert "Authorization" not in request["headers"] + assert json.loads(request["body"]) == {"assertion": "assertion-1"} + + +def test_refresh_before_expiry_fetches_a_new_assertion() -> None: + clock = Clock() + transport = FakeTransport([issued("token-1"), issued("token-2", lifetime=540)]) + source = Assertions() + credentials = provider(transport, source, clock) + + assert credentials.get_credential().access_token == "token-1" + clock.now = NOW + 239 + assert credentials.get_credential().access_token == "token-1" + clock.now = NOW + 240 + assert credentials.get_credential().access_token == "token-2" + + bodies = [json.loads(request["body"]) for request in transport.requests] + assert bodies == [{"assertion": "assertion-1"}, {"assertion": "assertion-2"}] + + +def test_refresh_credential_always_exchanges_a_new_assertion() -> None: + transport = FakeTransport([issued("token-1"), issued("token-2")]) + source = Assertions() + credentials = provider(transport, source) + + current = credentials.get_credential() + refreshed = credentials.refresh_credential(current) + + assert refreshed.access_token == "token-2" + assert source.calls == 2 + assert json.loads(transport.requests[1]["body"]) == {"assertion": "assertion-2"} + + +def test_forbidden_is_not_retried() -> None: + transport = FakeTransport([FakeResponse(403, headers={"x-request-id": "req-1"})]) + source = Assertions() + credentials = provider(transport, source) + + with pytest.raises(DeixicError) as error: + credentials.get_credential() + + assert error.value.kind == "authorization" + assert error.value.status_code == 403 + assert error.value.code == "workload_federation_forbidden" + assert error.value.request_id == "req-1" + assert source.calls == 1 + assert len(transport.requests) == 1 + + +def test_replay_is_not_retried() -> None: + transport = FakeTransport([FakeResponse(409)]) + source = Assertions() + credentials = provider(transport, source) + + with pytest.raises(DeixicError) as error: + credentials.get_credential() + + assert error.value.kind == "conflict" + assert error.value.code == "workload_assertion_replayed" + assert len(transport.requests) == 1 + + +def test_unavailable_is_retried_with_a_new_assertion() -> None: + transport = FakeTransport([FakeResponse(503), FakeResponse(503), issued("token-1")]) + source = Assertions() + sleeps: list[float] = [] + credentials = provider(transport, source, sleeps=sleeps) + + assert credentials.get_credential().access_token == "token-1" + bodies = [ + json.loads(request["body"])["assertion"] for request in transport.requests + ] + assert bodies == ["assertion-1", "assertion-2", "assertion-3"] + assert sleeps == [0.5, 1.0] + + +def test_unavailable_stops_after_max_attempts() -> None: + transport = FakeTransport([FakeResponse(503)] * 3) + credentials = provider(transport, Assertions()) + + with pytest.raises(DeixicError) as error: + credentials.get_credential() + + assert error.value.kind == "unavailable" + assert error.value.code == "workload_federation_unavailable" + assert len(transport.requests) == 3 + + +def test_transport_failure_is_retried() -> None: + transport = FakeTransport([ConnectionError("reset"), issued("token-1")]) + credentials = provider(transport, Assertions()) + + assert credentials.get_credential().access_token == "token-1" + assert len(transport.requests) == 2 + + +def test_reused_assertion_is_never_sent() -> None: + clock = Clock() + transport = FakeTransport([issued("token-1")]) + credentials = provider(transport, lambda: "same-assertion", clock) + + credentials.get_credential() + clock.now = NOW + 301 + with pytest.raises(DeixicError) as error: + credentials.get_credential() + + assert error.value.code == "workload_assertion_reused" + assert "same-assertion" not in str(error.value) + assert len(transport.requests) == 1 + + +def test_early_refresh_keeps_unexpired_token_when_source_has_not_rotated() -> None: + clock = Clock() + transport = FakeTransport([issued("token-1")]) + credentials = provider(transport, lambda: "same-assertion", clock) + + credentials.get_credential() + clock.now = NOW + 250 + assert credentials.get_credential().access_token == "token-1" + assert len(transport.requests) == 1 + + +def test_errors_do_not_contain_assertion_or_token() -> None: + transport = FakeTransport([FakeResponse(200, b'{"access_token":"leak"}')]) + credentials = provider(transport, lambda: "secret-assertion") + + with pytest.raises(DeixicError) as error: + credentials.get_credential() + + assert error.value.kind == "protocol" + assert "secret-assertion" not in str(error.value) + assert "leak" not in str(error.value) + assert "secret-assertion" not in repr(credentials) + + +def test_github_actions_source_requests_audience_bound_token() -> None: + transport = FakeTransport( + [ + FakeResponse(200, b'{"value":"github-jwt-1"}'), + FakeResponse(200, b'{"value":"github-jwt-2"}'), + ] + ) + source = github_actions_assertion_source( + EXCHANGE, + transport=transport, + environ={ + "ACTIONS_ID_TOKEN_REQUEST_URL": "https://token.actions.test/id?api-version=2.0", + "ACTIONS_ID_TOKEN_REQUEST_TOKEN": "request-token", + }, + ) + + assert source() == "github-jwt-1" + assert source() == "github-jwt-2" + request = transport.requests[0] + assert request["method"] == "GET" + assert request["url"] == ( + "https://token.actions.test/id?api-version=2.0&audience=" + "https%3A%2F%2Fidentity.deixic.test%2Fv1%2Fworkload-federation%2Fexchange" + ) + assert request["headers"]["Authorization"] == "Bearer request-token" + assert len(transport.requests) == 2 + + +def test_github_actions_source_requires_id_token_permission() -> None: + source = github_actions_assertion_source( + "aud", transport=FakeTransport([]), environ={} + ) + + with pytest.raises(DeixicError) as error: + source() + + assert error.value.code == "workload_assertion_unavailable" + assert "id-token: write" in str(error.value) + + +def test_file_source_rereads_the_file(tmp_path: Path) -> None: + token_file = tmp_path / "token" + token_file.write_text("k8s-jwt-1\n") + source = file_assertion_source(token_file) + + assert source() == "k8s-jwt-1" + token_file.write_text("k8s-jwt-2\n") + assert source() == "k8s-jwt-2" + + +def test_environment_source_reads_current_value() -> None: + environ = {"DEIXIC_WORKLOAD_ASSERTION": "env-jwt-1"} + source = environment_assertion_source("DEIXIC_WORKLOAD_ASSERTION", environ=environ) + + assert source() == "env-jwt-1" + environ["DEIXIC_WORKLOAD_ASSERTION"] = "env-jwt-2" + assert source() == "env-jwt-2" + + +def test_client_replays_401_with_a_newly_exchanged_token() -> None: + identity = FakeTransport([issued("token-1"), issued("token-2")]) + credentials = provider(identity, Assertions()) + platform = FakeTransport( + [ + FakeResponse(401), + FakeResponse( + 200, + public_pb2.GetThreadResponse().SerializeToString(deterministic=True), + ), + ] + ) + client = Deixic( + credential_provider=credentials, + organization_id="org-a", + workspace_id="workspace-a", + base_url="https://api.deixic.test", + transport=platform, + ) + + client.threads.get(channel_id="company") + + assert [r["headers"]["Authorization"] for r in platform.requests] == [ + "Bearer token-1", + "Bearer token-2", + ]