From 1c29583187031902eaf81a698cd820ecdacea274 Mon Sep 17 00:00:00 2001 From: Francisco Javier Arceo Date: Tue, 28 Jul 2026 06:44:51 -0400 Subject: [PATCH] feat: add optional OIDC bearer authentication Signed-off-by: Francisco Javier Arceo --- Cargo.lock | 222 ++++- Cargo.toml | 1 + README.md | 36 + crates/agentic-server/Cargo.toml | 4 + crates/agentic-server/src/app.rs | 32 +- crates/agentic-server/src/auth.rs | 837 ++++++++++++++++ crates/agentic-server/src/handler/mod.rs | 1 + .../src/handler/websocket/error.rs | 5 + .../src/handler/websocket/mod.rs | 1 + .../src/handler/websocket/responses.rs | 57 +- crates/agentic-server/src/lib.rs | 1 + crates/agentic-server/src/main.rs | 48 +- crates/agentic-server/src/server.rs | 91 +- crates/agentic-server/tests/oidc_auth_test.rs | 897 ++++++++++++++++++ docs/api/index.md | 40 + docs/deploying/container.md | 22 +- docs/design/oidc-bearer-authentication.md | 67 ++ 17 files changed, 2316 insertions(+), 46 deletions(-) create mode 100644 crates/agentic-server/src/auth.rs create mode 100644 crates/agentic-server/tests/oidc_auth_test.rs create mode 100644 docs/design/oidc-bearer-authentication.md diff --git a/Cargo.lock b/Cargo.lock index 7a7d66e6..fbfcbd32 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -21,7 +21,10 @@ dependencies = [ "either", "futures", "http", + "jsonwebtoken", + "rand 0.8.6", "reqwest 0.12.28", + "rsa", "serde", "serde_json", "thiserror", @@ -31,6 +34,7 @@ dependencies = [ "tower-http", "tracing", "tracing-subscriber", + "url", "uuid", ] @@ -280,6 +284,12 @@ dependencies = [ "tracing", ] +[[package]] +name = "base16ct" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4c7f02d4ea65f2c1853089ffd8d2787bdbc63de2f0d29dedbcf8ccdfa0ccd4cf" + [[package]] name = "base64" version = "0.22.1" @@ -597,6 +607,18 @@ version = "0.2.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "460fbee9c2c2f33933d720630a6a0bac33ba7053db5344fac858d4b8952d77d5" +[[package]] +name = "crypto-bigint" +version = "0.5.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0dc92fb57ca44df6db8059111ab3af99a63d5d0f8375d9972e319a379c6bab76" +dependencies = [ + "generic-array", + "rand_core 0.6.4", + "subtle", + "zeroize", +] + [[package]] name = "crypto-common" version = "0.1.7" @@ -607,6 +629,33 @@ dependencies = [ "typenum", ] +[[package]] +name = "curve25519-dalek" +version = "4.1.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "97fb8b7c4503de7d6ae7b42ab72a5a59857b4c937ec27a3d4539dba95b5ab2be" +dependencies = [ + "cfg-if", + "cpufeatures", + "curve25519-dalek-derive", + "digest", + "fiat-crypto", + "rustc_version", + "subtle", + "zeroize", +] + +[[package]] +name = "curve25519-dalek-derive" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f46882e17999c6cc590af592290432be3bce0428cb0d5f8b6715e4dc7b383eb3" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "data-encoding" version = "2.11.0" @@ -659,6 +708,44 @@ version = "1.0.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "92773504d58c093f6de2459af4af33faa518c13451eb8f2b5698ed3d36e7c813" +[[package]] +name = "ecdsa" +version = "0.16.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ee27f32b5c5292967d2d4a9d7f1e0b0aed2c15daded5a60300e4abb9d8020bca" +dependencies = [ + "der", + "digest", + "elliptic-curve", + "rfc6979", + "signature", + "spki", +] + +[[package]] +name = "ed25519" +version = "2.2.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "115531babc129696a58c64a4fef0a8bf9e9698629fb97e9e40767d235cfbcd53" +dependencies = [ + "pkcs8", + "signature", +] + +[[package]] +name = "ed25519-dalek" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "70e796c081cee67dc755e1a36a0a172b897fab85fc3f6bc48307991f64e4eca9" +dependencies = [ + "curve25519-dalek", + "ed25519", + "serde", + "sha2", + "subtle", + "zeroize", +] + [[package]] name = "either" version = "1.16.0" @@ -668,6 +755,27 @@ dependencies = [ "serde", ] +[[package]] +name = "elliptic-curve" +version = "0.13.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b5e6043086bf7973472e0c7dff2142ea0b680d30e18d9cc40f267efbf222bd47" +dependencies = [ + "base16ct", + "crypto-bigint", + "digest", + "ff", + "generic-array", + "group", + "hkdf", + "pem-rfc7468", + "pkcs8", + "rand_core 0.6.4", + "sec1", + "subtle", + "zeroize", +] + [[package]] name = "equivalent" version = "1.0.2" @@ -712,6 +820,22 @@ version = "2.4.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9f1f227452a390804cdb637b74a86990f2a7d7ba4b7d5693aac9b4dd6defd8d6" +[[package]] +name = "ff" +version = "0.13.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c0b50bfb653653f9ca9095b427bed08ab8d75a137839d9ad64eb11810d5b6393" +dependencies = [ + "rand_core 0.6.4", + "subtle", +] + +[[package]] +name = "fiat-crypto" +version = "0.2.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "28dea519a9695b9977216879a3ebfddf92f1c08c05d984f8996aecd6ecdc811d" + [[package]] name = "find-msvc-tools" version = "0.1.9" @@ -872,6 +996,7 @@ checksum = "85649ca51fd72272d7821adaf274ad91c288277713d9c18820d8499a7ff69e9a" dependencies = [ "typenum", "version_check", + "zeroize", ] [[package]] @@ -914,6 +1039,17 @@ dependencies = [ "wasip3", ] +[[package]] +name = "group" +version = "0.13.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f0f9ef7462f7c099f518d754361858f86d8a07af53ba9af0fe635bbccb151a63" +dependencies = [ + "ff", + "rand_core 0.6.4", + "subtle", +] + [[package]] name = "half" version = "2.7.1" @@ -1371,6 +1507,27 @@ dependencies = [ "wasm-bindgen", ] +[[package]] +name = "jsonwebtoken" +version = "10.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0529410abe238729a60b108898784df8984c87f6054c9c4fcacc47e4803c1ce1" +dependencies = [ + "base64", + "ed25519-dalek", + "getrandom 0.2.17", + "hmac", + "js-sys", + "p256", + "p384", + "rand 0.8.6", + "rsa", + "serde", + "serde_json", + "sha2", + "signature", +] + [[package]] name = "lazy_static" version = "1.5.0" @@ -1657,6 +1814,30 @@ dependencies = [ "vcpkg", ] +[[package]] +name = "p256" +version = "0.13.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c9863ad85fa8f4460f9c48cb909d38a0d689dba1f6f6988a5e3e0d31071bcd4b" +dependencies = [ + "ecdsa", + "elliptic-curve", + "primeorder", + "sha2", +] + +[[package]] +name = "p384" +version = "0.13.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fe42f1670a52a47d448f14b6a5c61dd78fce51856e68edaa38f7ae3a46b8d6b6" +dependencies = [ + "ecdsa", + "elliptic-curve", + "primeorder", + "sha2", +] + [[package]] name = "parking" version = "2.2.1" @@ -1796,6 +1977,15 @@ dependencies = [ "syn", ] +[[package]] +name = "primeorder" +version = "0.13.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "353e1ca18966c16d9deb1c69278edbc5f194139612772bd9537af60ac231e1e6" +dependencies = [ + "elliptic-curve", +] + [[package]] name = "proc-macro2" version = "1.0.106" @@ -2106,6 +2296,16 @@ dependencies = [ "web-sys", ] +[[package]] +name = "rfc6979" +version = "0.4.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f8dd2a808d456c4a54e300a23e9f5a67e122c3024119acbfd73e3bf664491cb2" +dependencies = [ + "hmac", + "subtle", +] + [[package]] name = "ring" version = "0.17.14" @@ -2188,7 +2388,7 @@ dependencies = [ "errno", "libc", "linux-raw-sys", - "windows-sys 0.52.0", + "windows-sys 0.61.2", ] [[package]] @@ -2246,7 +2446,7 @@ dependencies = [ "security-framework", "security-framework-sys", "webpki-root-certs", - "windows-sys 0.52.0", + "windows-sys 0.61.2", ] [[package]] @@ -2303,6 +2503,20 @@ version = "1.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "94143f37725109f92c262ed2cf5e59bce7498c01bcc1502d7b9afe439a4e9f49" +[[package]] +name = "sec1" +version = "0.7.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d3e97a565f76233a6003f9f5c54be1d9c5bdfa3eccfb189469f11ec4901c47dc" +dependencies = [ + "base16ct", + "der", + "generic-array", + "pkcs8", + "subtle", + "zeroize", +] + [[package]] name = "security-framework" version = "3.7.0" @@ -2813,10 +3027,10 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "32497e9a4c7b38532efcdebeef879707aa9f794296a4f0244f6f69e9bc8574bd" dependencies = [ "fastrand", - "getrandom 0.3.4", + "getrandom 0.4.2", "once_cell", "rustix", - "windows-sys 0.52.0", + "windows-sys 0.61.2", ] [[package]] diff --git a/Cargo.toml b/Cargo.toml index 82f3388e..b09fc8f1 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -26,6 +26,7 @@ criterion = { version = "0.5", features = ["async_tokio"] } futures = "0.3" indexmap = "2" http = "1" +jsonwebtoken = { version = "=10.3.0", default-features = false, features = ["rust_crypto"] } reqwest = { version = "0.12", default-features = false } rmcp-reqwest = { package = "reqwest", version = "0.13.2", default-features = false, features = ["json", "stream", "rustls"] } rmcp = { version = "1.8", default-features = false } diff --git a/README.md b/README.md index e5bd27b8..e3baf2b0 100644 --- a/README.md +++ b/README.md @@ -124,6 +124,28 @@ Then launch Codex: codex --disable image_generation -c model_provider=agentic-api -m Qwen/Qwen3-30B-A3B-FP8 ``` +If the gateway enables OIDC, configure Codex's supported command-backed bearer authentication instead of +`requires_openai_auth = false`: + +```toml +[model_providers.agentic-api] +name = "agentic-api" +base_url = "http://localhost:9000/v1" +wire_api = "responses" +supports_websockets = true + +[model_providers.agentic-api.auth] +command = "/absolute/path/to/print-oidc-token" +args = ["--audience", "agentic-api"] +refresh_interval_ms = 300000 +``` + +The command must print only a current OIDC token to stdout. Codex refreshes it before expiry and sends it as the +provider bearer token. See the +[Codex custom-provider authentication reference](https://developers.openai.com/codex/config-advanced#custom-model-providers). +Keep the inference credential in the gateway's `OPENAI_API_KEY`; do not print that service credential from the token +command. + ## 🧑‍💻 Claude Code on your own GPUs Agentic API serves the Anthropic Messages protocol at `/v1/messages`, so Claude Code (CLI or Agent SDK) runs against open models. Point it at the gateway: @@ -136,6 +158,20 @@ export ANTHROPIC_MODEL="Qwen/Qwen3-30B-A3B-FP8" # match the served model claude -p "summarize the files in this directory" ``` +With OIDC enabled, use Claude Code's bearer-token variable and leave its API-key variable unset so the identity token +is not also sent as an upstream `x-api-key`: + +```bash +export ANTHROPIC_BASE_URL="http://localhost:9000" +export ANTHROPIC_AUTH_TOKEN="$(/absolute/path/to/print-oidc-token --audience agentic-api)" +unset ANTHROPIC_API_KEY + +claude -p "summarize the files in this directory" +``` + +Refresh `ANTHROPIC_AUTH_TOKEN` before it expires. For supported dynamic credential helpers, see Anthropic's +[LLM gateway authentication guide](https://docs.anthropic.com/en/docs/claude-code/llm-gateway). + Claude Code's own tools (Bash, Edit, Read, …) stay **client-owned** — Claude Code runs them, as usual. ### Running Claude Code's web search on the gateway diff --git a/crates/agentic-server/Cargo.toml b/crates/agentic-server/Cargo.toml index 2d0f3564..62cfd1dc 100644 --- a/crates/agentic-server/Cargo.toml +++ b/crates/agentic-server/Cargo.toml @@ -14,6 +14,7 @@ clap.workspace = true either.workspace = true futures.workspace = true http.workspace = true +jsonwebtoken.workspace = true reqwest = { workspace = true, default-features = false, features = ["rustls-tls"] } serde.workspace = true serde_json.workspace = true @@ -23,12 +24,15 @@ tokio-util.workspace = true tower-http.workspace = true tracing.workspace = true tracing-subscriber.workspace = true +url.workspace = true [dev-dependencies] bytes.workspace = true criterion.workspace = true futures.workspace = true reqwest = { workspace = true, features = ["json"] } +rand = "0.8" +rsa = "0.9" serde_json.workspace = true tokio = { workspace = true, features = ["test-util"] } tokio-tungstenite.workspace = true diff --git a/crates/agentic-server/src/app.rs b/crates/agentic-server/src/app.rs index c2771c57..3d967a35 100644 --- a/crates/agentic-server/src/app.rs +++ b/crates/agentic-server/src/app.rs @@ -2,6 +2,7 @@ use std::sync::Arc; use std::sync::atomic::{AtomicUsize, Ordering}; use axum::Router; +use axum::middleware; use axum::routing::{get, post}; use http::HeaderValue; use tokio::sync::Notify; @@ -11,7 +12,8 @@ use tower_http::cors::{AllowOrigin, Any, CorsLayer}; use agentic_core::executor::ExecutionContext; use agentic_core::proxy::ProxyState; -use crate::handler::{conversations, count_tokens, health, messages, models, ready, responses, responses_ws}; +use crate::auth::{ANTHROPIC_COUNT_TOKENS_PATH, ANTHROPIC_MESSAGES_PATH, OidcAuthenticator, require_oidc}; +use crate::handler::{conversations, count_tokens, health, messages, models, ready, responses, responses_ws_with_auth}; #[derive(Clone, Default)] pub struct WebSocketTracker { @@ -119,14 +121,30 @@ pub struct AppState { } pub fn build_router(state: AppState, server_config: &ServerConfig) -> Router { - Router::new() - .route("/health", get(health)) - .route("/ready", get(ready)) + build_router_with_auth(state, server_config, None) +} + +pub fn build_router_with_auth( + state: AppState, + server_config: &ServerConfig, + authenticator: Option, +) -> Router { + let public_routes = Router::new().route("/health", get(health)).route("/ready", get(ready)); + let protected_routes = Router::new() .route("/v1/conversations", post(conversations)) .route("/v1/models", get(models)) - .route("/v1/messages", post(messages)) - .route("/v1/messages/count_tokens", post(count_tokens)) - .route("/v1/responses", post(responses).get(responses_ws)) + .route(ANTHROPIC_MESSAGES_PATH, post(messages)) + .route(ANTHROPIC_COUNT_TOKENS_PATH, post(count_tokens)) + .route("/v1/responses", post(responses).get(responses_ws_with_auth)); + let protected_routes = match authenticator { + Some(authenticator) => { + protected_routes.route_layer(middleware::from_fn_with_state(authenticator, require_oidc)) + } + None => protected_routes, + }; + + public_routes + .merge(protected_routes) .layer(server_config.cors_layer()) .with_state(state) } diff --git a/crates/agentic-server/src/auth.rs b/crates/agentic-server/src/auth.rs new file mode 100644 index 00000000..48366d34 --- /dev/null +++ b/crates/agentic-server/src/auth.rs @@ -0,0 +1,837 @@ +use std::collections::HashMap; +use std::sync::Arc; +use std::time::Duration; +use std::time::Instant; + +use axum::body::Body; +use axum::extract::State; +use axum::http::{HeaderMap, Request, StatusCode, header}; +use axum::middleware::Next; +use axum::response::Response; +use jsonwebtoken::jwk::{Jwk, JwkSet, KeyAlgorithm, KeyOperations, PublicKeyUse}; +use jsonwebtoken::{Algorithm, DecodingKey, Validation, decode, decode_header}; +use serde::Deserialize; +use serde_json::json; +use tokio::sync::RwLock; +use tracing::{debug, warn}; +use url::{Host, Url}; + +const OIDC_HTTP_TIMEOUT: Duration = Duration::from_secs(10); +const JWKS_REFRESH_COOLDOWN: Duration = Duration::from_secs(30); +const JWKS_REFRESH_COALESCE_WINDOW: Duration = Duration::from_secs(1); +const DEFAULT_JWKS_TTL: Duration = Duration::from_secs(300); +const MAX_JWKS_TTL: Duration = Duration::from_secs(3600); +const MAX_PROVIDER_RESPONSE_BYTES: usize = 1024 * 1024; +const MAX_JWKS_KEYS: usize = 100; +const JWT_CLOCK_SKEW_SECONDS: u64 = 60; + +pub(crate) const ANTHROPIC_MESSAGES_PATH: &str = "/v1/messages"; +pub(crate) const ANTHROPIC_COUNT_TOKENS_PATH: &str = "/v1/messages/count_tokens"; + +#[derive(Clone)] +pub struct OidcConfig { + issuer: Url, + issuer_value: String, + audience: String, + allows_loopback_http: bool, +} + +impl OidcConfig { + /// Create an OIDC bearer-token configuration. + /// + /// # Errors + /// + /// Returns an error when the issuer is not an absolute HTTPS URL (except + /// for loopback test/development issuers) or the audience is empty. + pub fn new(issuer: &str, audience: &str) -> Result { + let issuer_value = issuer.trim().to_owned(); + let issuer = Url::parse(&issuer_value).map_err(OidcAuthError::InvalidIssuer)?; + let allows_loopback_http = is_loopback_http(&issuer); + if issuer.scheme() != "https" && !allows_loopback_http { + return Err(OidcAuthError::InsecureIssuer); + } + if issuer.query().is_some() || issuer.fragment().is_some() { + return Err(OidcAuthError::InvalidIssuerComponents); + } + + let audience = audience.trim().to_owned(); + if audience.is_empty() { + return Err(OidcAuthError::EmptyAudience); + } + + Ok(Self { + issuer, + issuer_value, + audience, + allows_loopback_http, + }) + } + + fn discovery_url(&self) -> Result { + Url::parse(&format!( + "{}/.well-known/openid-configuration", + self.issuer.as_str().trim_end_matches('/') + )) + .map_err(OidcAuthError::InvalidIssuer) + } +} + +fn is_loopback_http(url: &Url) -> bool { + url.scheme() == "http" + && match url.host() { + Some(Host::Ipv4(address)) => address.is_loopback(), + Some(Host::Ipv6(address)) => address.is_loopback(), + Some(Host::Domain(_)) | None => false, + } +} + +struct CachedKey { + decoding_key: DecodingKey, + algorithm: Option, +} + +struct CachedJwks { + keys: HashMap>, + expires_at: Instant, +} + +struct RefreshState { + last_completed: Instant, + retry_after: Option, + coalesce_until: Option, +} + +#[derive(Clone)] +pub struct OidcAuthenticator { + audience: String, + jwks_uri: Url, + keys: Arc>, + refresh_state: Arc>, + refresh_gate: Arc, + validations: Arc>, + client: reqwest::Client, +} + +impl OidcAuthenticator { + /// Discover an OIDC provider and cache its initial verification keys. + /// + /// # Errors + /// + /// Returns an error when provider discovery or the JSON Web Key Set + /// (JWKS) request fails, or when the discovered metadata is inconsistent. + pub async fn discover(config: OidcConfig) -> Result { + let client = reqwest::Client::builder() + .redirect(reqwest::redirect::Policy::none()) + .timeout(OIDC_HTTP_TIMEOUT) + .build() + .map_err(OidcAuthError::HttpClient)?; + let (metadata, _) = + fetch_json::(&client, config.discovery_url()?, ProviderRequest::Metadata).await?; + let discovered_issuer = Url::parse(&metadata.issuer).map_err(OidcAuthError::InvalidDiscoveredIssuer)?; + if discovered_issuer != config.issuer { + return Err(OidcAuthError::IssuerMismatch { + expected: config.issuer_value, + discovered: metadata.issuer, + }); + } + + let jwks_uri = Url::parse(&metadata.jwks_uri).map_err(OidcAuthError::InvalidJwksUri)?; + if jwks_uri.scheme() != "https" && !(config.allows_loopback_http && is_loopback_http(&jwks_uri)) { + return Err(OidcAuthError::InsecureJwksUri); + } + + let keys = fetch_jwks(&client, jwks_uri.clone()).await?; + let refresh_completed = Instant::now(); + let validations = build_validations(&metadata.issuer, &config.audience); + Ok(Self { + audience: config.audience, + jwks_uri, + keys: Arc::new(RwLock::new(keys)), + refresh_state: Arc::new(std::sync::Mutex::new(RefreshState { + last_completed: refresh_completed, + retry_after: None, + coalesce_until: None, + })), + refresh_gate: Arc::new(tokio::sync::Semaphore::new(1)), + validations: Arc::new(validations), + client, + }) + } + + async fn authenticate(&self, token: &str) -> Result { + let token_header = decode_header(token).map_err(OidcAuthError::InvalidToken)?; + if matches!(token_header.alg, Algorithm::HS256 | Algorithm::HS384 | Algorithm::HS512) { + return Err(OidcAuthError::UnsupportedTokenAlgorithm); + } + let kid = token_header.kid.ok_or(OidcAuthError::MissingKeyId)?; + let key = self.verification_key(&kid).await?; + if key.algorithm.is_some_and(|algorithm| algorithm != token_header.alg) { + return Err(OidcAuthError::AlgorithmMismatch); + } + let validation = self + .validations + .iter() + .find_map(|(algorithm, validation)| (*algorithm == token_header.alg).then_some(validation)) + .ok_or(OidcAuthError::UnsupportedTokenAlgorithm)?; + + let claims = decode::(token, &key.decoding_key, validation) + .map_err(OidcAuthError::InvalidToken)? + .claims; + if claims.sub.is_empty() { + return Err(OidcAuthError::EmptySubject); + } + if !claims.audience_allows(&self.audience) { + return Err(OidcAuthError::InvalidAuthorizedParty); + } + Ok(AuthenticatedPrincipal { + issuer: claims.iss, + subject: claims.sub, + expires_at: claims.exp, + }) + } + + async fn verification_key(&self, kid: &str) -> Result, OidcAuthError> { + { + let keys = self.keys.read().await; + if Instant::now() < keys.expires_at { + if let Some(key) = keys.keys.get(kid) { + return Ok(Arc::clone(key)); + } + } + } + + let _refresh_permit = self + .refresh_gate + .acquire() + .await + .map_err(|_| OidcAuthError::RefreshStateUnavailable)?; + { + let keys = self.keys.read().await; + let refresh_state = self + .refresh_state + .lock() + .map_err(|_| OidcAuthError::RefreshStateUnavailable)?; + let now = Instant::now(); + if now < keys.expires_at { + if let Some(key) = keys.keys.get(kid) { + return Ok(Arc::clone(key)); + } + if now.duration_since(refresh_state.last_completed) < JWKS_REFRESH_COOLDOWN { + return Err(OidcAuthError::UnknownKeyId); + } + } + if refresh_state.coalesce_until.is_some_and(|deadline| now < deadline) { + return keys.keys.get(kid).cloned().ok_or(OidcAuthError::UnknownKeyId); + } + if refresh_state.retry_after.is_some_and(|deadline| now < deadline) { + return Err(OidcAuthError::JwksRefreshBackoff); + } + } + + self.refresh_state + .lock() + .map_err(|_| OidcAuthError::RefreshStateUnavailable)? + .retry_after = Some(Instant::now() + JWKS_REFRESH_COOLDOWN); + let refreshed = fetch_jwks(&self.client, self.jwks_uri.clone()).await?; + let key = refreshed.keys.get(kid).cloned(); + *self.keys.write().await = refreshed; + let refresh_completed = Instant::now(); + let mut refresh_state = self + .refresh_state + .lock() + .map_err(|_| OidcAuthError::RefreshStateUnavailable)?; + refresh_state.last_completed = refresh_completed; + refresh_state.retry_after = None; + refresh_state.coalesce_until = Some(refresh_completed + JWKS_REFRESH_COALESCE_WINDOW); + key.ok_or(OidcAuthError::UnknownKeyId) + } +} + +#[derive(Clone, Debug, Eq, PartialEq)] +pub struct AuthenticatedPrincipal { + issuer: String, + subject: String, + expires_at: u64, +} + +impl AuthenticatedPrincipal { + #[must_use] + pub fn issuer(&self) -> &str { + &self.issuer + } + + #[must_use] + pub fn subject(&self) -> &str { + &self.subject + } + + #[must_use] + pub fn expires_at(&self) -> u64 { + self.expires_at + } + + pub(crate) fn is_expired(&self) -> bool { + self.is_expired_at(jsonwebtoken::get_current_timestamp()) + } + + fn is_expired_at(&self, timestamp: u64) -> bool { + timestamp > self.expires_at.saturating_add(JWT_CLOCK_SKEW_SECONDS) + } + + #[cfg(test)] + pub(crate) fn expired_for_test() -> Self { + Self { + issuer: "https://issuer.example".to_owned(), + subject: "subject".to_owned(), + expires_at: 0, + } + } +} + +pub async fn require_oidc( + State(authenticator): State, + mut request: Request, + next: Next, +) -> Response { + let error_format = AuthErrorFormat::for_path(request.uri().path()); + let Some(token) = bearer_token(request.headers()) else { + return authentication_error(error_format, "missing_bearer_token", "missing bearer token"); + }; + let duplicate_identity_api_key = request + .headers() + .get_all("x-api-key") + .iter() + .any(|value| value.to_str().ok().is_some_and(|value| value.trim() == token)); + + match authenticator.authenticate(token).await { + Ok(principal) => { + request.headers_mut().remove(header::AUTHORIZATION); + if duplicate_identity_api_key { + request.headers_mut().remove("x-api-key"); + } + request.extensions_mut().insert(principal); + next.run(request).await + } + Err(error) => { + if error.is_dependency_failure() { + warn!(error = %error, "OIDC token verification dependency failed"); + authentication_service_unavailable(error_format) + } else { + debug!(error = %error, "OIDC bearer token rejected"); + authentication_error(error_format, "invalid_token", "invalid bearer token") + } + } + } +} + +fn bearer_token(headers: &axum::http::HeaderMap) -> Option<&str> { + headers + .get(header::AUTHORIZATION)? + .to_str() + .ok()? + .split_once(' ') + .and_then(|(scheme, token)| { + let token = token.trim(); + (scheme.eq_ignore_ascii_case("bearer") && !token.is_empty()).then_some(token) + }) +} + +#[derive(Clone, Copy)] +enum AuthErrorFormat { + OpenAi, + Anthropic, +} + +impl AuthErrorFormat { + fn for_path(path: &str) -> Self { + if matches!(path, ANTHROPIC_MESSAGES_PATH | ANTHROPIC_COUNT_TOKENS_PATH) { + Self::Anthropic + } else { + Self::OpenAi + } + } +} + +fn authentication_error(format: AuthErrorFormat, code: &'static str, message: &'static str) -> Response { + protocol_error( + format, + StatusCode::UNAUTHORIZED, + "authentication_error", + "authentication_error", + code, + message, + true, + ) +} + +fn authentication_service_unavailable(format: AuthErrorFormat) -> Response { + protocol_error( + format, + StatusCode::SERVICE_UNAVAILABLE, + "server_error", + "api_error", + "authentication_service_unavailable", + "authentication service temporarily unavailable", + false, + ) +} + +fn protocol_error( + format: AuthErrorFormat, + status: StatusCode, + openai_error_type: &'static str, + anthropic_error_type: &'static str, + code: &'static str, + message: &'static str, + challenge: bool, +) -> Response { + let body = match format { + AuthErrorFormat::OpenAi => json!({ + "error": { + "message": message, + "type": openai_error_type, + "param": null, + "code": code + } + }), + AuthErrorFormat::Anthropic => json!({ + "type": "error", + "error": { + "type": anthropic_error_type, + "message": message + } + }), + }; + let mut builder = Response::builder() + .status(status) + .header(header::CONTENT_TYPE, "application/json"); + if challenge { + builder = builder.header(header::WWW_AUTHENTICATE, "Bearer"); + } + builder + .body(Body::from(body.to_string())) + .expect("valid authentication protocol error response") +} + +#[derive(Clone, Copy)] +enum ProviderRequest { + Metadata, + Jwks, +} + +impl ProviderRequest { + fn error(self, error: reqwest::Error) -> OidcAuthError { + match self { + Self::Metadata => OidcAuthError::ProviderMetadataRequest(error), + Self::Jwks => OidcAuthError::JwksRequest(error), + } + } +} + +async fn fetch_jwks(client: &reqwest::Client, uri: Url) -> Result { + let (keys, headers) = fetch_json::(client, uri, ProviderRequest::Jwks).await?; + compile_jwks(keys, jwks_ttl(&headers)) +} + +async fn fetch_json( + client: &reqwest::Client, + uri: Url, + request_kind: ProviderRequest, +) -> Result<(T, HeaderMap), OidcAuthError> +where + T: serde::de::DeserializeOwned, +{ + let mut response = client + .get(uri) + .send() + .await + .map_err(|error| request_kind.error(error))? + .error_for_status() + .map_err(|error| request_kind.error(error))?; + if response + .content_length() + .is_some_and(|length| length > MAX_PROVIDER_RESPONSE_BYTES as u64) + { + return Err(OidcAuthError::ProviderResponseTooLarge); + } + + let headers = response.headers().clone(); + let mut body = Vec::with_capacity( + response + .content_length() + .and_then(|length| usize::try_from(length).ok()) + .unwrap_or_default() + .min(MAX_PROVIDER_RESPONSE_BYTES), + ); + while let Some(chunk) = response.chunk().await.map_err(|error| request_kind.error(error))? { + if body.len().saturating_add(chunk.len()) > MAX_PROVIDER_RESPONSE_BYTES { + return Err(OidcAuthError::ProviderResponseTooLarge); + } + body.extend_from_slice(&chunk); + } + + let value = serde_json::from_slice(&body).map_err(OidcAuthError::InvalidProviderJson)?; + Ok((value, headers)) +} + +fn compile_jwks(keys: JwkSet, ttl: Duration) -> Result { + if keys.keys.is_empty() { + return Err(OidcAuthError::EmptyJwks); + } + if keys.keys.len() > MAX_JWKS_KEYS { + return Err(OidcAuthError::TooManyJwksKeys); + } + + let mut compiled = HashMap::with_capacity(keys.keys.len()); + for key in keys.keys { + let Some(kid) = key.common.key_id.clone() else { + continue; + }; + let algorithm = match verification_algorithm(&key) { + VerificationAlgorithm::Skip => continue, + VerificationAlgorithm::AnyAsymmetric => None, + VerificationAlgorithm::Exact(algorithm) => Some(algorithm), + }; + let Ok(decoding_key) = DecodingKey::from_jwk(&key) else { + continue; + }; + if compiled + .insert( + kid.clone(), + Arc::new(CachedKey { + decoding_key, + algorithm, + }), + ) + .is_some() + { + return Err(OidcAuthError::DuplicateKeyId(kid)); + } + } + if compiled.is_empty() { + return Err(OidcAuthError::EmptyJwks); + } + + Ok(CachedJwks { + keys: compiled, + expires_at: Instant::now() + ttl, + }) +} + +enum VerificationAlgorithm { + Skip, + AnyAsymmetric, + Exact(Algorithm), +} + +fn verification_algorithm(key: &Jwk) -> VerificationAlgorithm { + if key + .common + .public_key_use + .as_ref() + .is_some_and(|key_use| key_use != &PublicKeyUse::Signature) + || key + .common + .key_operations + .as_ref() + .is_some_and(|operations| !operations.contains(&KeyOperations::Verify)) + { + return VerificationAlgorithm::Skip; + } + + match key.common.key_algorithm { + Some(key_algorithm) => VerificationAlgorithm::Exact(match key_algorithm { + KeyAlgorithm::ES256 => Algorithm::ES256, + KeyAlgorithm::ES384 => Algorithm::ES384, + KeyAlgorithm::RS256 => Algorithm::RS256, + KeyAlgorithm::RS384 => Algorithm::RS384, + KeyAlgorithm::RS512 => Algorithm::RS512, + KeyAlgorithm::PS256 => Algorithm::PS256, + KeyAlgorithm::PS384 => Algorithm::PS384, + KeyAlgorithm::PS512 => Algorithm::PS512, + KeyAlgorithm::EdDSA => Algorithm::EdDSA, + KeyAlgorithm::HS256 + | KeyAlgorithm::HS384 + | KeyAlgorithm::HS512 + | KeyAlgorithm::RSA1_5 + | KeyAlgorithm::RSA_OAEP + | KeyAlgorithm::RSA_OAEP_256 + | KeyAlgorithm::UNKNOWN_ALGORITHM => return VerificationAlgorithm::Skip, + }), + None => VerificationAlgorithm::AnyAsymmetric, + } +} + +fn build_validations(issuer: &str, audience: &str) -> Vec<(Algorithm, Validation)> { + [ + Algorithm::ES256, + Algorithm::ES384, + Algorithm::RS256, + Algorithm::RS384, + Algorithm::RS512, + Algorithm::PS256, + Algorithm::PS384, + Algorithm::PS512, + Algorithm::EdDSA, + ] + .into_iter() + .map(|algorithm| { + let mut validation = Validation::new(algorithm); + validation.leeway = JWT_CLOCK_SKEW_SECONDS; + validation.set_audience(&[audience]); + validation.set_issuer(&[issuer]); + validation.set_required_spec_claims(&["exp", "iss", "aud", "sub"]); + validation.validate_nbf = true; + (algorithm, validation) + }) + .collect() +} + +fn jwks_ttl(headers: &HeaderMap) -> Duration { + let mut max_age: Option = None; + for value in headers.get_all(header::CACHE_CONTROL) { + let Ok(value) = value.to_str() else { + continue; + }; + for directive in value.split(',').map(str::trim) { + if directive.eq_ignore_ascii_case("no-cache") || directive.eq_ignore_ascii_case("no-store") { + return Duration::ZERO; + } + if let Some((name, seconds)) = directive.split_once('=') { + if name.trim().eq_ignore_ascii_case("max-age") { + let parsed = seconds.trim().trim_matches('"').parse::().ok(); + max_age = match (max_age, parsed) { + (Some(existing), Some(parsed)) => Some(existing.min(parsed)), + (None, parsed) => parsed, + (existing, None) => existing, + }; + } + } + } + } + let age = headers + .get_all(header::AGE) + .iter() + .filter_map(|value| value.to_str().ok()?.trim().parse::().ok()) + .max() + .map_or(Duration::ZERO, Duration::from_secs); + max_age.map_or_else( + || DEFAULT_JWKS_TTL.saturating_sub(age), + |seconds| Duration::from_secs(seconds).saturating_sub(age).min(MAX_JWKS_TTL), + ) +} + +#[derive(Deserialize)] +struct ProviderMetadata { + issuer: String, + jwks_uri: String, +} + +#[derive(Deserialize)] +struct IdentityClaims { + iss: String, + sub: String, + aud: AudienceClaim, + exp: u64, + #[serde(default)] + azp: Option, +} + +impl IdentityClaims { + fn audience_allows(&self, expected: &str) -> bool { + match &self.aud { + AudienceClaim::One(audience) => audience == expected, + AudienceClaim::Many(audiences) if audiences.len() == 1 => audiences[0] == expected, + AudienceClaim::Many(audiences) => { + audiences.iter().any(|audience| audience == expected) && self.azp.as_deref() == Some(expected) + } + } + } +} + +#[derive(Deserialize)] +#[serde(untagged)] +enum AudienceClaim { + One(String), + Many(Vec), +} + +#[derive(Debug, thiserror::Error)] +#[non_exhaustive] +pub enum OidcAuthError { + #[error("OIDC audience must not be empty")] + EmptyAudience, + #[error("OIDC issuer must not include a query or fragment")] + InvalidIssuerComponents, + #[error("OIDC issuer must use HTTPS; HTTP is only permitted for loopback IP addresses (127.0.0.1 or ::1)")] + InsecureIssuer, + #[error("OIDC JWKS URI must use HTTPS; HTTP is only permitted for loopback IP addresses (127.0.0.1 or ::1)")] + InsecureJwksUri, + #[error("invalid OIDC issuer URL")] + InvalidIssuer(#[source] url::ParseError), + #[error("invalid OIDC JWKS URI")] + InvalidJwksUri(#[source] url::ParseError), + #[error("OIDC discovery returned an invalid issuer URL")] + InvalidDiscoveredIssuer(#[source] url::ParseError), + #[error("failed to build OIDC HTTP client")] + HttpClient(#[source] reqwest::Error), + #[error("OIDC provider metadata request failed")] + ProviderMetadataRequest(#[source] reqwest::Error), + #[error("OIDC JWKS request failed")] + JwksRequest(#[source] reqwest::Error), + #[error("OIDC provider response exceeded the size limit")] + ProviderResponseTooLarge, + #[error("OIDC provider returned invalid JSON")] + InvalidProviderJson(#[source] serde_json::Error), + #[error("OIDC discovery returned issuer {discovered}, expected {expected}")] + IssuerMismatch { expected: String, discovered: String }, + #[error("OIDC provider returned an empty JWKS")] + EmptyJwks, + #[error("OIDC provider returned too many JWKs")] + TooManyJwksKeys, + #[error("OIDC provider returned duplicate JWK key ID {0}")] + DuplicateKeyId(String), + #[error("OIDC JWKS refresh is temporarily backed off after a provider failure")] + JwksRefreshBackoff, + #[error("OIDC JWKS refresh coordination is unavailable")] + RefreshStateUnavailable, + #[error("bearer token is missing a key ID")] + MissingKeyId, + #[error("bearer token references an unknown key ID")] + UnknownKeyId, + #[error("bearer token subject must not be empty")] + EmptySubject, + #[error("bearer token authorized party does not match the configured audience")] + InvalidAuthorizedParty, + #[error("bearer token uses an unsupported algorithm")] + UnsupportedTokenAlgorithm, + #[error("bearer token and JWK algorithms do not match")] + AlgorithmMismatch, + #[error("bearer token validation failed")] + InvalidToken(#[source] jsonwebtoken::errors::Error), +} + +impl OidcAuthError { + fn is_dependency_failure(&self) -> bool { + matches!( + self, + Self::JwksRequest(_) + | Self::ProviderResponseTooLarge + | Self::InvalidProviderJson(_) + | Self::EmptyJwks + | Self::TooManyJwksKeys + | Self::DuplicateKeyId(_) + | Self::JwksRefreshBackoff + | Self::RefreshStateUnavailable + ) + } +} + +#[cfg(test)] +mod tests { + use super::{ + AuthenticatedPrincipal, MAX_JWKS_KEYS, MAX_JWKS_TTL, OidcAuthError, VerificationAlgorithm, compile_jwks, + jwks_ttl, verification_algorithm, + }; + use axum::http::{HeaderMap, HeaderValue, header}; + use jsonwebtoken::jwk::{Jwk, JwkSet, KeyAlgorithm, KeyOperations, PublicKeyUse}; + use jsonwebtoken::{Algorithm, EncodingKey}; + use rand::rngs::OsRng; + use rsa::RsaPrivateKey; + use rsa::pkcs1::EncodeRsaPrivateKey; + use std::time::Duration; + + fn test_jwk() -> Jwk { + let private_key = RsaPrivateKey::new(&mut OsRng, 2048).expect("generate test RSA key"); + let private_key = private_key.to_pkcs1_der().expect("encode test RSA key"); + let mut jwk = Jwk::from_encoding_key(&EncodingKey::from_rsa_der(private_key.as_bytes()), Algorithm::RS256) + .expect("test JWK"); + jwk.common.key_id = Some("test-key".to_owned()); + jwk.common.key_algorithm = Some(KeyAlgorithm::RS256); + jwk.common.public_key_use = Some(PublicKeyUse::Signature); + jwk + } + + #[test] + fn jwks_cache_lifetime_uses_provider_max_age_with_a_cap() { + let mut headers = HeaderMap::new(); + headers.insert(header::CACHE_CONTROL, HeaderValue::from_static("public, max-age=60")); + assert_eq!(jwks_ttl(&headers), Duration::from_secs(60)); + + headers.insert(header::CACHE_CONTROL, HeaderValue::from_static("max-age=86400")); + assert_eq!(jwks_ttl(&headers), MAX_JWKS_TTL); + + headers.insert( + header::CACHE_CONTROL, + HeaderValue::from_static("private, Max-Age=\"30\""), + ); + assert_eq!(jwks_ttl(&headers), Duration::from_secs(30)); + + headers.insert(header::CACHE_CONTROL, HeaderValue::from_static("no-store")); + assert_eq!(jwks_ttl(&headers), Duration::ZERO); + + headers.insert(header::CACHE_CONTROL, HeaderValue::from_static("max-age=60")); + headers.insert(header::AGE, HeaderValue::from_static("55")); + assert_eq!(jwks_ttl(&headers), Duration::from_secs(5)); + + headers.insert(header::CACHE_CONTROL, HeaderValue::from_static("max-age=86400")); + headers.insert(header::AGE, HeaderValue::from_static("4000")); + assert_eq!(jwks_ttl(&headers), MAX_JWKS_TTL); + + headers.remove(header::AGE); + headers.append(header::CACHE_CONTROL, HeaderValue::from_static("no-cache")); + assert_eq!(jwks_ttl(&headers), Duration::ZERO); + } + + #[test] + fn authenticated_principal_expiration_includes_clock_skew() { + let principal = AuthenticatedPrincipal { + issuer: "https://issuer.example".to_owned(), + subject: "subject".to_owned(), + expires_at: 100, + }; + + assert!(!principal.is_expired_at(160)); + assert!(principal.is_expired_at(161)); + } + + #[test] + fn jwks_limits_and_signature_metadata_are_enforced() { + let mut encryption_key = test_jwk(); + encryption_key.common.public_key_use = Some(PublicKeyUse::Encryption); + assert!(matches!( + compile_jwks( + JwkSet { + keys: vec![encryption_key] + }, + Duration::from_secs(60) + ), + Err(OidcAuthError::EmptyJwks) + )); + + let mut non_verifying_key = test_jwk(); + non_verifying_key.common.key_operations = Some(vec![KeyOperations::Encrypt]); + assert!(matches!( + compile_jwks( + JwkSet { + keys: vec![non_verifying_key] + }, + Duration::from_secs(60) + ), + Err(OidcAuthError::EmptyJwks) + )); + + let mut mismatched_algorithm = test_jwk(); + mismatched_algorithm.common.key_algorithm = Some(KeyAlgorithm::RS512); + assert!(matches!( + verification_algorithm(&mismatched_algorithm), + VerificationAlgorithm::Exact(Algorithm::RS512) + )); + + let too_many_keys = vec![test_jwk(); MAX_JWKS_KEYS + 1]; + assert!(matches!( + compile_jwks(JwkSet { keys: too_many_keys }, Duration::from_secs(60)), + Err(OidcAuthError::TooManyJwksKeys) + )); + } +} diff --git a/crates/agentic-server/src/handler/mod.rs b/crates/agentic-server/src/handler/mod.rs index 157a6366..9d379d1f 100644 --- a/crates/agentic-server/src/handler/mod.rs +++ b/crates/agentic-server/src/handler/mod.rs @@ -5,3 +5,4 @@ pub mod websocket; pub use common::{convert_response, executor_error_response}; pub use http::{conversations, count_tokens, health, messages, models, ready, responses}; pub use websocket::responses_ws; +pub(crate) use websocket::responses_ws_with_auth; diff --git a/crates/agentic-server/src/handler/websocket/error.rs b/crates/agentic-server/src/handler/websocket/error.rs index 43b013d6..0ffe4652 100644 --- a/crates/agentic-server/src/handler/websocket/error.rs +++ b/crates/agentic-server/src/handler/websocket/error.rs @@ -21,6 +21,9 @@ pub(super) enum WsError { #[error("websocket messages must be JSON text frames")] BinaryFrame, + #[error("OIDC bearer token expired")] + AuthenticationExpired, + #[error("websocket send failed")] SendFailed, @@ -36,6 +39,7 @@ impl WsError { match self { Self::Executor(err) => err.http_status(), Self::InvalidJson(_) | Self::UnexpectedType | Self::BinaryFrame => StatusCode::BAD_REQUEST, + Self::AuthenticationExpired => StatusCode::UNAUTHORIZED, Self::SerializeJson(_) | Self::SendFailed | Self::ClientDisconnected | Self::Receive(_) => { StatusCode::INTERNAL_SERVER_ERROR } @@ -47,6 +51,7 @@ impl WsError { Self::Executor(err) => err.error_code(), Self::InvalidJson(_) => "invalid_json", Self::UnexpectedType | Self::BinaryFrame => "invalid_request_error", + Self::AuthenticationExpired => "invalid_token", Self::SerializeJson(_) | Self::SendFailed | Self::ClientDisconnected | Self::Receive(_) => "server_error", } } diff --git a/crates/agentic-server/src/handler/websocket/mod.rs b/crates/agentic-server/src/handler/websocket/mod.rs index 14c75011..e4e248c8 100644 --- a/crates/agentic-server/src/handler/websocket/mod.rs +++ b/crates/agentic-server/src/handler/websocket/mod.rs @@ -2,3 +2,4 @@ mod error; mod responses; pub use responses::responses_ws; +pub(crate) use responses::responses_ws_with_auth; diff --git a/crates/agentic-server/src/handler/websocket/responses.rs b/crates/agentic-server/src/handler/websocket/responses.rs index b5264a1b..97fd472f 100644 --- a/crates/agentic-server/src/handler/websocket/responses.rs +++ b/crates/agentic-server/src/handler/websocket/responses.rs @@ -1,8 +1,8 @@ use std::collections::VecDeque; use std::sync::Arc; -use axum::extract::State; use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade}; +use axum::extract::{Extension, State}; use axum::http::HeaderMap; use axum::response::Response; use either::Either; @@ -20,21 +20,45 @@ use agentic_core::utils::common::utcnow_str; use super::super::common::{MAX_BODY_SIZE, extract_bearer}; use super::error::WsError; use crate::app::AppState; +use crate::auth::AuthenticatedPrincipal; type WsSender = SplitSink; type WsReceiver = SplitStream; pub async fn responses_ws(State(state): State, headers: HeaderMap, ws: WebSocketUpgrade) -> Response { + upgrade_responses_ws(state, headers, ws, None) +} + +pub(crate) async fn responses_ws_with_auth( + State(state): State, + principal: Option>, + headers: HeaderMap, + ws: WebSocketUpgrade, +) -> Response { + upgrade_responses_ws(state, headers, ws, principal.map(|Extension(principal)| principal)) +} + +fn upgrade_responses_ws( + state: AppState, + headers: HeaderMap, + ws: WebSocketUpgrade, + principal: Option, +) -> Response { let websocket_guard = state.websocket_tracker.track(); ws.max_message_size(MAX_BODY_SIZE) .max_frame_size(MAX_BODY_SIZE) .on_upgrade(move |socket| async move { let _websocket_guard = websocket_guard; - responses_ws_loop(socket, state, headers).await; + responses_ws_loop(socket, state, headers, principal).await; }) } -async fn responses_ws_loop(socket: WebSocket, state: AppState, headers: HeaderMap) { +async fn responses_ws_loop( + socket: WebSocket, + state: AppState, + headers: HeaderMap, + principal: Option, +) { debug!("responses websocket session opened"); let shutdown_token = state.shutdown_token.clone(); let (mut sender, mut receiver) = socket.split(); @@ -78,6 +102,11 @@ async fn responses_ws_loop(socket: WebSocket, state: AppState, headers: HeaderMa } }; + if let Some(error) = websocket_identity_error(principal.as_ref()) { + let _ = send_ws_error(&mut sender, &error).await; + break; + } + match handle_ws_text( &mut sender, &mut receiver, @@ -101,6 +130,12 @@ async fn responses_ws_loop(socket: WebSocket, state: AppState, headers: HeaderMa debug!("responses websocket session closed"); } +fn websocket_identity_error(principal: Option<&AuthenticatedPrincipal>) -> Option { + principal + .is_some_and(AuthenticatedPrincipal::is_expired) + .then_some(WsError::AuthenticationExpired) +} + async fn next_ws_message( shutdown_token: &CancellationToken, receiver: &mut Receiver, @@ -415,7 +450,10 @@ mod tests { use futures::{Sink, Stream, StreamExt, sink, stream}; use tokio_util::sync::CancellationToken; - use super::{ShutdownInput, close_ws, keep_if_running, next_shutdown_input, next_ws_message}; + use super::{ + ShutdownInput, close_ws, keep_if_running, next_shutdown_input, next_ws_message, websocket_identity_error, + }; + use crate::auth::AuthenticatedPrincipal; struct CloseErrorSink; @@ -484,6 +522,17 @@ mod tests { assert_eq!(keep_if_running(&shutdown_token, "unpolled stream"), None); } + #[test] + fn websocket_identity_expiry_selects_the_unauthorized_error_event() { + assert!(websocket_identity_error(None).is_none()); + let error = + websocket_identity_error(Some(&AuthenticatedPrincipal::expired_for_test())).expect("expired-token error"); + let frame = error.to_ws_frame().expect("client-visible error frame"); + + assert_eq!(frame["status"], 401); + assert_eq!(frame["error"]["code"], "invalid_token"); + } + #[tokio::test] async fn close_ws_ignores_late_frames_until_peer_close() { let mut sender = sink::drain(); diff --git a/crates/agentic-server/src/lib.rs b/crates/agentic-server/src/lib.rs index 55f41a01..36a9b14f 100644 --- a/crates/agentic-server/src/lib.rs +++ b/crates/agentic-server/src/lib.rs @@ -1,2 +1,3 @@ pub mod app; +pub mod auth; pub mod handler; diff --git a/crates/agentic-server/src/main.rs b/crates/agentic-server/src/main.rs index 890fa4d2..df8ab323 100644 --- a/crates/agentic-server/src/main.rs +++ b/crates/agentic-server/src/main.rs @@ -5,6 +5,7 @@ use agentic_core::config::{ SqliteConfig, SqliteTempStore, normalize_base_url, }; use agentic_core::error::Error; +use agentic_server::auth::OidcConfig; mod server; @@ -13,6 +14,14 @@ struct CommonArgs { #[arg(long, env = "OPENAI_API_KEY", hide_env_values = true, global = true)] openai_api_key: Option, + /// OIDC issuer for optional inbound bearer-token authentication. + #[arg(long, env = "OIDC_ISSUER", global = true)] + oidc_issuer: Option, + + /// Required bearer-token audience when `OIDC_ISSUER` is configured. + #[arg(long, env = "OIDC_AUDIENCE", global = true)] + oidc_audience: Option, + #[arg(long, env = "GATEWAY_HOST", default_value = "0.0.0.0", global = true)] gateway_host: String, @@ -40,6 +49,17 @@ struct CommonArgs { db_url: String, } +fn oidc_config_from_values( + issuer: Option<&str>, + audience: Option<&str>, +) -> Result, server::ServerError> { + match (issuer, audience) { + (None, None) => Ok(None), + (Some(issuer), Some(audience)) => Ok(Some(OidcConfig::new(issuer, audience)?)), + _ => Err(Error::Config("OIDC_ISSUER and OIDC_AUDIENCE must be configured together".to_owned()).into()), + } +} + #[derive(Parser)] #[command(name = "agentic-server", about = "Stateful API gateway for vLLM Responses API")] struct Cli { @@ -148,7 +168,7 @@ fn build_config(llm_api_base: String, common: &CommonArgs) -> Result Result<(), Error> { +async fn main() -> Result<(), server::ServerError> { tracing_subscriber::fmt() .with_env_filter( tracing_subscriber::EnvFilter::try_from_default_env() @@ -161,6 +181,7 @@ async fn main() -> Result<(), Error> { llm_api_base, common, } = Cli::parse(); + let oidc_config = oidc_config_from_values(common.oidc_issuer.as_deref(), common.oidc_audience.as_deref())?; match command { None => { @@ -171,20 +192,21 @@ async fn main() -> Result<(), Error> { ) })?; let config = build_config(normalize_base_url(&base), &common)?; - server::run(config, &common.gateway_host, common.gateway_port).await + server::run(config, &common.gateway_host, common.gateway_port, oidc_config).await } Some(Commands::Serve { model, port, llm_args }) => { if llm_api_base.is_some() { return Err(Error::Config( "--llm-api-base is only valid in standalone mode; remove it when using `serve`".to_owned(), - )); + ) + .into()); } let config = build_config(normalize_base_url(&format!("http://127.0.0.1:{port}")), &common)?; let mut args = vec!["--model".to_owned(), model]; args.push("--port".to_owned()); args.push(port.to_string()); args.extend(llm_args); - server::run_with_llm(config, &common.gateway_host, common.gateway_port, args).await + server::run_with_llm(config, &common.gateway_host, common.gateway_port, args, oidc_config).await } } } @@ -193,7 +215,9 @@ async fn main() -> Result<(), Error> { mod tests { use clap::{CommandFactory, Parser}; - use super::{Cli, Commands, parse_env_temp_store_value, parse_env_u32_value, parse_env_u64_value}; + use super::{ + Cli, Commands, oidc_config_from_values, parse_env_temp_store_value, parse_env_u32_value, parse_env_u64_value, + }; use agentic_core::config::{DEFAULT_SQLITE_MAX_CONNECTIONS, SqliteTempStore}; #[test] @@ -221,6 +245,18 @@ mod tests { assert!(cli.common.skip_llm_ready_check); } + #[test] + fn oidc_configuration_requires_issuer_and_audience_together() { + assert!(oidc_config_from_values(None, None).expect("disabled OIDC").is_none()); + assert!(oidc_config_from_values(Some("https://issuer.example"), None).is_err()); + assert!(oidc_config_from_values(None, Some("agentic-api")).is_err()); + assert!( + oidc_config_from_values(Some("https://issuer.example"), Some("agentic-api")) + .expect("complete OIDC configuration") + .is_some() + ); + } + #[test] fn container_runtime_options_are_bound_to_environment_variables() { let command = Cli::command(); @@ -229,6 +265,8 @@ mod tests { ("llm_api_base", "LLM_API_BASE"), ("gateway_host", "GATEWAY_HOST"), ("gateway_port", "GATEWAY_PORT"), + ("oidc_issuer", "OIDC_ISSUER"), + ("oidc_audience", "OIDC_AUDIENCE"), ] { let env = command .get_arguments() diff --git a/crates/agentic-server/src/server.rs b/crates/agentic-server/src/server.rs index c9cd19f5..a976ec10 100644 --- a/crates/agentic-server/src/server.rs +++ b/crates/agentic-server/src/server.rs @@ -4,18 +4,35 @@ use std::sync::Arc; use std::time::Duration; use agentic_core::config::Config; -use agentic_core::error::Error; +use agentic_core::error::Error as CoreError; use agentic_core::executor::ExecutionContext; use agentic_core::proxy::ProxyState; use agentic_core::readiness::wait_llm_ready; -use agentic_server::app::{AppState, ServerConfig, WebSocketTracker, build_router}; +use agentic_server::app::{AppState, ServerConfig, WebSocketTracker, build_router_with_auth}; +use agentic_server::auth::{OidcAuthError, OidcAuthenticator, OidcConfig}; use tokio::net::TcpListener; use tokio_util::sync::CancellationToken; use tracing::{info, warn}; const GATEWAY_DRAIN_TIMEOUT: Duration = Duration::from_secs(8); -async fn build_state(config: &Config, shutdown_token: CancellationToken) -> Result { +#[derive(Debug, thiserror::Error)] +pub enum ServerError { + #[error(transparent)] + Core(#[from] CoreError), + #[error(transparent)] + Io(#[from] std::io::Error), + #[error("failed to initialize OIDC authentication: {0}")] + Oidc(#[source] OidcAuthError), +} + +impl From for ServerError { + fn from(error: OidcAuthError) -> Self { + Self::Oidc(error) + } +} + +async fn build_state(config: &Config, shutdown_token: CancellationToken) -> Result { let proxy_state = ProxyState::new(config.clone())?; let exec_ctx = Arc::new(ExecutionContext::from_config(config).await?); @@ -29,12 +46,17 @@ async fn build_state(config: &Config, shutdown_token: CancellationToken) -> Resu }) } -async fn serve_gateway(state: AppState, host: &str, port: u16) -> Result<(), Error> { +async fn serve_gateway( + state: AppState, + host: &str, + port: u16, + authenticator: Option, +) -> Result<(), ServerError> { let addr = format!("{host}:{port}"); let server_config = ServerConfig::from_env(); let shutdown_token = state.shutdown_token.clone(); let websocket_tracker = state.websocket_tracker.clone(); - let router = build_router(state, &server_config); + let router = build_router_with_auth(state, &server_config, authenticator); let listener = TcpListener::bind(&addr).await?; info!("gateway listening on {addr}"); axum::serve(listener, router) @@ -46,9 +68,14 @@ async fn serve_gateway(state: AppState, host: &str, port: u16) -> Result<(), Err Ok(()) } -async fn serve_gateway_until_signal(state: AppState, host: &str, port: u16) -> Result<(), Error> { +async fn serve_gateway_until_signal( + state: AppState, + host: &str, + port: u16, + authenticator: Option, +) -> Result<(), ServerError> { let shutdown_token = state.shutdown_token.clone(); - let gateway = serve_gateway(state, host, port); + let gateway = serve_gateway(state, host, port, authenticator); tokio::pin!(gateway); tokio::select! { @@ -62,9 +89,9 @@ async fn serve_gateway_until_signal(state: AppState, host: &str, port: u16) -> R } } -async fn drain_gateway(gateway: Pin<&mut F>) -> Result<(), Error> +async fn drain_gateway(gateway: Pin<&mut F>) -> Result<(), ServerError> where - F: Future>, + F: Future>, { if let Ok(result) = tokio::time::timeout(GATEWAY_DRAIN_TIMEOUT, gateway).await { result @@ -92,7 +119,7 @@ async fn shutdown_signal() -> Result<(), std::io::Error> { tokio::signal::ctrl_c().await } -async fn wait_until_llm_ready(config: &Config) -> Result<(), Error> { +async fn wait_until_llm_ready(config: &Config) -> Result<(), ServerError> { if config.skip_llm_ready_check { info!("skipping LLM readiness check: {}", config.llm_api_base); return Ok(()); @@ -107,21 +134,35 @@ async fn wait_until_llm_ready(config: &Config) -> Result<(), Error> { /// /// # Errors /// -/// Returns an error if DB initialisation, LLM readiness polling, or the -/// server binding fails. -pub async fn run(config: Config, host: &str, port: u16) -> Result<(), Error> { +/// Returns an error if OIDC discovery or verification-key loading, DB +/// initialisation, LLM readiness polling, or the server binding fails. +pub async fn run(config: Config, host: &str, port: u16, oidc_config: Option) -> Result<(), ServerError> { + let authenticator = match oidc_config { + Some(config) => Some(OidcAuthenticator::discover(config).await?), + None => None, + }; wait_until_llm_ready(&config).await?; let state = build_state(&config, CancellationToken::new()).await?; - serve_gateway_until_signal(state, host, port).await + serve_gateway_until_signal(state, host, port, authenticator).await } /// Spawn vLLM as a subprocess and run the gateway in the foreground. /// /// # Errors /// -/// Returns an error if vLLM fails to start, DB init fails, or the gateway -/// errors. -pub async fn run_with_llm(config: Config, host: &str, port: u16, llm_args: Vec) -> Result<(), Error> { +/// Returns an error if OIDC discovery or verification-key loading fails, vLLM +/// fails to start, DB initialisation fails, or the gateway errors. +pub async fn run_with_llm( + config: Config, + host: &str, + port: u16, + llm_args: Vec, + oidc_config: Option, +) -> Result<(), ServerError> { + let authenticator = match oidc_config { + Some(config) => Some(OidcAuthenticator::discover(config).await?), + None => None, + }; let mut cmd = tokio::process::Command::new("python"); cmd.arg("-m").arg("vllm.entrypoints.openai.api_server"); cmd.args(&llm_args); @@ -134,10 +175,10 @@ pub async fn run_with_llm(config: Config, host: &str, port: u16, llm_args: Vec ready.map(|()| true), + ready = wait_llm_ready(&config) => ready.map(|()| true).map_err(ServerError::from), status = child.wait() => { let status = status?; - Err(Error::LlmProcessExited { status: status.to_string() }) + Err(ServerError::from(CoreError::LlmProcessExited { status: status.to_string() })) } } }; @@ -162,7 +203,7 @@ pub async fn run_with_llm(config: Config, host: &str, port: u16, llm_args: Vec { shutdown_token.cancel(); let status = status?; - Err(Error::LlmProcessExited { status: status.to_string() }) + Err(ServerError::from(CoreError::LlmProcessExited { status: status.to_string() })) }, signal = shutdown_signal() => { match signal { @@ -191,12 +232,12 @@ pub async fn run_with_llm(config: Config, host: &str, port: u16, llm_args: Vec>(); + let gateway = std::future::pending::>(); tokio::pin!(gateway); drain_gateway(gateway.as_mut()).await.unwrap(); @@ -204,7 +245,7 @@ mod tests { #[tokio::test] async fn gateway_drain_preserves_server_errors() { - let gateway = std::future::ready(Err(Error::Config("gateway failed".to_owned()))); + let gateway = std::future::ready(Err(ServerError::from(CoreError::Config("gateway failed".to_owned())))); tokio::pin!(gateway); let error = drain_gateway(gateway.as_mut()).await.unwrap_err(); diff --git a/crates/agentic-server/tests/oidc_auth_test.rs b/crates/agentic-server/tests/oidc_auth_test.rs new file mode 100644 index 00000000..c5396fb9 --- /dev/null +++ b/crates/agentic-server/tests/oidc_auth_test.rs @@ -0,0 +1,897 @@ +#[allow(dead_code)] +mod common; + +use axum::body::{Body, Bytes}; +use axum::http::{HeaderMap, HeaderValue, Response, StatusCode}; +use axum::response::IntoResponse; +use axum::routing::get; +use axum::{Extension, Json, Router, middleware}; +use common::{test_config, test_state}; +use futures::stream; +use jsonwebtoken::jwk::{Jwk, KeyAlgorithm, PublicKeyUse}; +use jsonwebtoken::{Algorithm, EncodingKey, Header, encode}; +use rand::rngs::OsRng; +use rsa::RsaPrivateKey; +use rsa::pkcs1::EncodeRsaPrivateKey; +use serde::Serialize; +use serde_json::{Value, json}; +use std::sync::atomic::{AtomicUsize, Ordering}; +use tokio::net::TcpListener; +use tokio::task::JoinHandle; + +use agentic_server::app::{ServerConfig, build_router_with_auth}; +use agentic_server::auth::{AuthenticatedPrincipal, OidcAuthError, OidcAuthenticator, OidcConfig, require_oidc}; + +const TEST_AUDIENCE: &str = "agentic-api"; + +struct TestGateway { + address: std::net::SocketAddr, + handle: JoinHandle<()>, +} + +impl Drop for TestGateway { + fn drop(&mut self) { + self.handle.abort(); + } +} + +async fn spawn_gateway(authenticator: OidcAuthenticator, upstream_url: &str) -> TestGateway { + let config = test_config(upstream_url); + let router = build_router_with_auth( + test_state(&config), + &ServerConfig { + cors_allowed_origins: Vec::new(), + }, + Some(authenticator), + ); + spawn_router(router).await +} + +async fn spawn_router(router: Router) -> TestGateway { + let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind test server"); + let address = listener.local_addr().expect("gateway address"); + let handle = tokio::spawn(async move { + axum::serve(listener, router).await.expect("serve gateway"); + }); + TestGateway { address, handle } +} + +async fn discover_test_authenticator(issuer: &str) -> OidcAuthenticator { + OidcAuthenticator::discover(OidcConfig::new(issuer, TEST_AUDIENCE).expect("OIDC config")) + .await + .expect("OIDC discovery") +} + +fn test_key() -> (Vec, Value) { + test_key_with_id("test-key") +} + +fn test_key_with_id(kid: &str) -> (Vec, Value) { + let private_key = RsaPrivateKey::new(&mut OsRng, 2048).expect("generate test RSA key"); + let private_key_der = private_key.to_pkcs1_der().expect("encode test RSA key"); + let private_key_der = private_key_der.as_bytes().to_vec(); + let encoding_key = EncodingKey::from_rsa_der(&private_key_der); + let mut jwk = Jwk::from_encoding_key(&encoding_key, Algorithm::RS256).expect("test JWK"); + jwk.common.key_id = Some(kid.to_owned()); + jwk.common.key_algorithm = Some(KeyAlgorithm::RS256); + jwk.common.public_key_use = Some(PublicKeyUse::Signature); + (private_key_der, serde_json::to_value(jwk).expect("serialize test JWK")) +} + +async fn spawn_rotating_oidc_provider() -> (String, Vec, Vec, std::sync::Arc, JoinHandle<()>) { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("bind rotating OIDC provider"); + let issuer = format!("http://{}", listener.local_addr().expect("OIDC provider address")); + let discovery_issuer = issuer.clone(); + let discovery_jwks_uri = format!("{issuer}/jwks"); + let (old_private_key, old_jwk) = test_key_with_id("old-key"); + let (new_private_key, new_jwk) = test_key_with_id("new-key"); + let jwks_requests = std::sync::Arc::new(AtomicUsize::new(0)); + let observed_jwks_requests = std::sync::Arc::clone(&jwks_requests); + + let provider = Router::new() + .route( + "/.well-known/openid-configuration", + get(move || { + let issuer = discovery_issuer.clone(); + let jwks_uri = discovery_jwks_uri.clone(); + async move { Json(json!({"issuer": issuer, "jwks_uri": jwks_uri})) } + }), + ) + .route( + "/jwks", + get(move || { + let old_jwk = old_jwk.clone(); + let new_jwk = new_jwk.clone(); + let observed_jwks_requests = std::sync::Arc::clone(&observed_jwks_requests); + async move { + let request = observed_jwks_requests.fetch_add(1, Ordering::Relaxed); + let (cache_control, jwk) = if request == 0 { + ("max-age=0", old_jwk) + } else { + ("max-age=0", new_jwk) + }; + ( + [(reqwest::header::CACHE_CONTROL, cache_control)], + Json(json!({"keys": [jwk]})), + ) + } + }), + ); + let handle = tokio::spawn(async move { + axum::serve(listener, provider) + .await + .expect("serve rotating OIDC provider"); + }); + + (issuer, old_private_key, new_private_key, jwks_requests, handle) +} + +async fn spawn_failing_refresh_provider() -> (String, Vec, std::sync::Arc, JoinHandle<()>) { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("bind failing OIDC provider"); + let issuer = format!("http://{}", listener.local_addr().expect("OIDC provider address")); + let discovery_issuer = issuer.clone(); + let discovery_jwks_uri = format!("{issuer}/jwks"); + let (private_key, jwk) = test_key(); + let jwks_requests = std::sync::Arc::new(AtomicUsize::new(0)); + let observed_jwks_requests = std::sync::Arc::clone(&jwks_requests); + + let provider = Router::new() + .route( + "/.well-known/openid-configuration", + get(move || { + let issuer = discovery_issuer.clone(); + let jwks_uri = discovery_jwks_uri.clone(); + async move { Json(json!({"issuer": issuer, "jwks_uri": jwks_uri})) } + }), + ) + .route( + "/jwks", + get(move || { + let jwk = jwk.clone(); + let observed_jwks_requests = std::sync::Arc::clone(&observed_jwks_requests); + async move { + if observed_jwks_requests.fetch_add(1, Ordering::Relaxed) == 0 { + ( + [(reqwest::header::CACHE_CONTROL, "max-age=0")], + Json(json!({"keys": [jwk]})), + ) + .into_response() + } else { + StatusCode::SERVICE_UNAVAILABLE.into_response() + } + } + }), + ); + let handle = tokio::spawn(async move { + axum::serve(listener, provider) + .await + .expect("serve failing OIDC provider"); + }); + + (issuer, private_key, jwks_requests, handle) +} + +async fn spawn_metadata_provider(build_body: impl FnOnce(&str) -> String) -> (String, JoinHandle<()>) { + let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind metadata provider"); + let issuer = format!("http://{}", listener.local_addr().expect("metadata provider address")); + let body = std::sync::Arc::new(build_body(&issuer)); + let provider = Router::new().route( + "/.well-known/openid-configuration", + get(move || { + let body = std::sync::Arc::clone(&body); + async move { ([(reqwest::header::CONTENT_TYPE, "application/json")], body.to_string()) } + }), + ); + let handle = tokio::spawn(async move { + axum::serve(listener, provider).await.expect("serve metadata provider"); + }); + (issuer, handle) +} + +async fn spawn_chunked_metadata_provider() -> (String, JoinHandle<()>) { + let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind metadata provider"); + let issuer = format!("http://{}", listener.local_addr().expect("metadata provider address")); + let provider = Router::new().route( + "/.well-known/openid-configuration", + get(|| async { + let chunks = stream::iter([ + Ok::<_, std::convert::Infallible>(Bytes::from(vec![b' '; 768 * 1024])), + Ok(Bytes::from(vec![b' '; 768 * 1024])), + ]); + Response::builder() + .header(reqwest::header::CONTENT_TYPE, "application/json") + .body(Body::from_stream(chunks)) + .expect("chunked metadata response") + }), + ); + let handle = tokio::spawn(async move { + axum::serve(listener, provider).await.expect("serve metadata provider"); + }); + (issuer, handle) +} + +async fn spawn_oidc_provider() -> ( + String, + Vec, + Vec, + std::sync::Arc, + tokio::task::JoinHandle<()>, +) { + spawn_oidc_provider_with_algorithm(KeyAlgorithm::RS256).await +} + +async fn spawn_oidc_provider_with_algorithm( + key_algorithm: KeyAlgorithm, +) -> ( + String, + Vec, + Vec, + std::sync::Arc, + tokio::task::JoinHandle<()>, +) { + let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind OIDC provider"); + let issuer = format!("http://{}", listener.local_addr().expect("OIDC provider address")); + let discovery_issuer = issuer.clone(); + let discovery_jwks_uri = format!("{issuer}/jwks"); + let (private_key_der, mut jwk) = test_key(); + jwk["alg"] = Value::String( + match key_algorithm { + KeyAlgorithm::RS256 => "RS256", + KeyAlgorithm::RS512 => "RS512", + _ => panic!("test provider only supports RS256 and RS512 metadata"), + } + .to_owned(), + ); + let public_jwk = serde_json::to_vec(&jwk).expect("serialize public test JWK"); + let jwks_requests = std::sync::Arc::new(AtomicUsize::new(0)); + let observed_jwks_requests = std::sync::Arc::clone(&jwks_requests); + + let provider = Router::new() + .route( + "/.well-known/openid-configuration", + get(move || { + let issuer = discovery_issuer.clone(); + let jwks_uri = discovery_jwks_uri.clone(); + async move { + Json(json!({ + "issuer": issuer, + "jwks_uri": jwks_uri + })) + } + }), + ) + .route( + "/jwks", + get(move || { + let jwk = jwk.clone(); + let observed_jwks_requests = std::sync::Arc::clone(&observed_jwks_requests); + async move { + observed_jwks_requests.fetch_add(1, Ordering::Relaxed); + Json(json!({ + "keys": [jwk] + })) + } + }), + ); + let handle = tokio::spawn(async move { + axum::serve(listener, provider).await.expect("serve OIDC provider"); + }); + + (issuer, private_key_der, public_jwk, jwks_requests, handle) +} + +fn identity_token(issuer: &str, audience: &str, expires_at: u64, kid: &str, private_key_der: &[u8]) -> String { + #[derive(Serialize)] + struct Claims<'a> { + iss: &'a str, + sub: &'a str, + aud: &'a str, + exp: u64, + } + + let mut header = Header::new(Algorithm::RS256); + header.kid = Some(kid.to_owned()); + encode( + &header, + &Claims { + iss: issuer, + sub: "github-user-123", + aud: audience, + exp: expires_at, + }, + &EncodingKey::from_rsa_der(private_key_der), + ) + .expect("encode test identity token") +} + +fn identity_token_with_audiences( + issuer: &str, + audiences: &[&str], + authorized_party: Option<&str>, + private_key_der: &[u8], +) -> String { + let mut header = Header::new(Algorithm::RS256); + header.kid = Some("test-key".to_owned()); + encode( + &header, + &json!({ + "iss": issuer, + "sub": "github-user-123", + "aud": audiences, + "azp": authorized_party, + "exp": jsonwebtoken::get_current_timestamp() + 300 + }), + &EncodingKey::from_rsa_der(private_key_der), + ) + .expect("encode multi-audience identity token") +} + +fn custom_identity_token(header: &Header, claims: &Value, private_key_der: &[u8]) -> String { + encode(header, claims, &EncodingKey::from_rsa_der(private_key_der)).expect("encode custom identity token") +} + +fn hmac_identity_token(issuer: &str, audience: &str, secret: &[u8]) -> String { + #[derive(Serialize)] + struct Claims<'a> { + iss: &'a str, + sub: &'a str, + aud: &'a str, + exp: u64, + } + + let mut header = Header::new(Algorithm::HS256); + header.kid = Some("test-key".to_owned()); + encode( + &header, + &Claims { + iss: issuer, + sub: "github-user-123", + aud: audience, + exp: jsonwebtoken::get_current_timestamp() + 300, + }, + &EncodingKey::from_secret(secret), + ) + .expect("encode test HMAC identity token") +} + +async fn spawn_models_upstream() -> ( + String, + std::sync::Arc>>, + tokio::task::JoinHandle<()>, +) { + let observed_headers = std::sync::Arc::new(std::sync::Mutex::new(None)); + let captured_headers = std::sync::Arc::clone(&observed_headers); + let upstream = Router::new().route( + "/v1/models", + get(move |headers: HeaderMap| { + let captured_headers = std::sync::Arc::clone(&captured_headers); + async move { + *captured_headers.lock().expect("capture headers") = Some(headers); + Json(json!({"object": "list", "data": []})) + } + }), + ); + let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind upstream"); + let address = listener.local_addr().expect("upstream address"); + let handle = tokio::spawn(async move { + axum::serve(listener, upstream).await.expect("serve upstream"); + }); + (format!("http://{address}"), observed_headers, handle) +} + +#[tokio::test] +async fn configured_oidc_rejects_missing_bearer_before_upstream() { + let (issuer, _private_key, _public_jwk, _jwks_requests, _provider) = spawn_oidc_provider().await; + let authenticator = discover_test_authenticator(&issuer).await; + let gateway = spawn_gateway(authenticator, "http://127.0.0.1:9").await; + + let response = reqwest::get(format!("http://{}/v1/models", gateway.address)) + .await + .expect("request gateway"); + + assert_eq!(response.status(), reqwest::StatusCode::UNAUTHORIZED); + let body = response.json::().await.expect("JSON error body"); + assert_eq!(body["error"]["code"], "missing_bearer_token"); + + let health = reqwest::get(format!("http://{}/health", gateway.address)) + .await + .expect("request health"); + assert_eq!(health.status(), reqwest::StatusCode::OK); +} + +#[tokio::test] +async fn configured_oidc_rejects_invalid_bearer_before_upstream() { + let (issuer, _private_key, _public_jwk, _jwks_requests, _provider) = spawn_oidc_provider().await; + let authenticator = discover_test_authenticator(&issuer).await; + let gateway = spawn_gateway(authenticator, "http://127.0.0.1:9").await; + + let response = reqwest::Client::new() + .get(format!("http://{}/v1/models", gateway.address)) + .bearer_auth("not-a-jwt") + .send() + .await + .expect("request gateway"); + + assert_eq!(response.status(), reqwest::StatusCode::UNAUTHORIZED); + let body = response.json::().await.expect("JSON error body"); + assert_eq!(body["error"]["code"], "invalid_token"); + + let messages_response = reqwest::Client::new() + .post(format!("http://{}/v1/messages", gateway.address)) + .bearer_auth("not-a-jwt") + .send() + .await + .expect("request Messages API"); + assert_eq!(messages_response.status(), reqwest::StatusCode::UNAUTHORIZED); + let messages_body = messages_response.json::().await.expect("JSON error body"); + assert_eq!(messages_body["type"], "error"); + assert_eq!(messages_body["error"]["type"], "authentication_error"); +} + +#[tokio::test] +async fn configured_oidc_rejects_hmac_tokens_signed_with_public_key_material() { + let (issuer, _private_key, public_jwk, _jwks_requests, _provider) = spawn_oidc_provider().await; + let authenticator = discover_test_authenticator(&issuer).await; + let gateway = spawn_gateway(authenticator, "http://127.0.0.1:9").await; + + let response = reqwest::Client::new() + .get(format!("http://{}/v1/models", gateway.address)) + .bearer_auth(hmac_identity_token(&issuer, "agentic-api", &public_jwk)) + .send() + .await + .expect("request gateway"); + + assert_eq!(response.status(), reqwest::StatusCode::UNAUTHORIZED); + let body = response.json::().await.expect("JSON error body"); + assert_eq!(body["error"]["code"], "invalid_token"); +} + +#[tokio::test] +async fn configured_oidc_rejects_token_and_jwk_algorithm_mismatch() { + let (issuer, private_key, _public_jwk, _jwks_requests, _provider) = + spawn_oidc_provider_with_algorithm(KeyAlgorithm::RS512).await; + let authenticator = discover_test_authenticator(&issuer).await; + let gateway = spawn_gateway(authenticator, "http://127.0.0.1:9").await; + + let response = reqwest::Client::new() + .get(format!("http://{}/v1/models", gateway.address)) + .bearer_auth(identity_token( + &issuer, + TEST_AUDIENCE, + jsonwebtoken::get_current_timestamp() + 300, + "test-key", + &private_key, + )) + .send() + .await + .expect("algorithm-mismatch request"); + + assert_eq!(response.status(), reqwest::StatusCode::UNAUTHORIZED); +} + +#[tokio::test] +async fn configured_oidc_accepts_valid_identity_and_uses_service_upstream_credential() { + let (issuer, private_key_der, _public_jwk, _jwks_requests, _provider) = spawn_oidc_provider().await; + let authenticator = discover_test_authenticator(&issuer).await; + let (upstream_url, observed_headers, _upstream) = spawn_models_upstream().await; + let gateway = spawn_gateway(authenticator, &upstream_url).await; + let identity_token = identity_token( + &issuer, + "agentic-api", + jsonwebtoken::get_current_timestamp() + 300, + "test-key", + &private_key_der, + ); + let mut identity_headers = HeaderMap::new(); + identity_headers.append("x-api-key", HeaderValue::from_static("distinct-upstream-key")); + identity_headers.append( + "x-api-key", + HeaderValue::from_str(&identity_token).expect("identity header value"), + ); + + let response = reqwest::Client::new() + .get(format!("http://{}/v1/models", gateway.address)) + .bearer_auth(&identity_token) + .headers(identity_headers) + .send() + .await + .expect("request gateway"); + + assert_eq!(response.status(), reqwest::StatusCode::OK); + let headers = observed_headers + .lock() + .expect("read captured headers") + .clone() + .expect("upstream request"); + assert_eq!( + headers + .get(reqwest::header::AUTHORIZATION) + .expect("service credential") + .to_str() + .expect("valid authorization"), + "Bearer test-key" + ); + assert!( + headers.get("x-api-key").is_none(), + "duplicate identity credential must not reach the upstream" + ); +} + +#[tokio::test] +async fn configured_oidc_rejects_wrong_issuer_audience_and_expired_tokens() { + let (issuer, private_key_der, _public_jwk, _jwks_requests, _provider) = spawn_oidc_provider().await; + let authenticator = discover_test_authenticator(&issuer).await; + let gateway = spawn_gateway(authenticator, "http://127.0.0.1:9").await; + let now = jsonwebtoken::get_current_timestamp(); + let invalid_claims = [ + ("https://other-issuer.example", "agentic-api", now + 300), + (issuer.as_str(), "other-audience", now + 300), + (issuer.as_str(), "agentic-api", now - 120), + ]; + + for (token_issuer, audience, expires_at) in invalid_claims { + let response = reqwest::Client::new() + .get(format!("http://{}/v1/models", gateway.address)) + .bearer_auth(identity_token( + token_issuer, + audience, + expires_at, + "test-key", + &private_key_der, + )) + .send() + .await + .expect("request gateway"); + + assert_eq!(response.status(), reqwest::StatusCode::UNAUTHORIZED); + let body = response.json::().await.expect("JSON error body"); + assert_eq!(body["error"]["code"], "invalid_token"); + } +} + +#[tokio::test] +async fn configured_oidc_rejects_missing_claims_empty_subject_future_nbf_and_bad_signature() { + let (issuer, private_key, _public_jwk, _jwks_requests, _provider) = spawn_oidc_provider().await; + let (other_private_key, _) = test_key_with_id("test-key"); + let authenticator = discover_test_authenticator(&issuer).await; + let gateway = spawn_gateway(authenticator, "http://127.0.0.1:9").await; + let now = jsonwebtoken::get_current_timestamp(); + let mut header = Header::new(Algorithm::RS256); + header.kid = Some("test-key".to_owned()); + let valid_claims = json!({ + "iss": issuer, + "sub": "github-user-123", + "aud": "agentic-api", + "exp": now + 300 + }); + let mut missing_kid_header = Header::new(Algorithm::RS256); + missing_kid_header.kid = None; + let cases = [ + custom_identity_token(&missing_kid_header, &valid_claims, &private_key), + custom_identity_token( + &header, + &json!({"iss": issuer, "sub": "", "aud": "agentic-api", "exp": now + 300}), + &private_key, + ), + custom_identity_token( + &header, + &json!({ + "iss": issuer, + "sub": "github-user-123", + "aud": "agentic-api", + "exp": now + 300, + "nbf": now + 300 + }), + &private_key, + ), + custom_identity_token( + &header, + &json!({"iss": issuer, "sub": "github-user-123", "aud": "agentic-api"}), + &private_key, + ), + custom_identity_token(&header, &valid_claims, &other_private_key), + ]; + + for token in cases { + let response = reqwest::Client::new() + .get(format!("http://{}/v1/models", gateway.address)) + .bearer_auth(token) + .send() + .await + .expect("invalid token request"); + assert_eq!(response.status(), reqwest::StatusCode::UNAUTHORIZED); + } +} + +#[tokio::test] +async fn unknown_key_ids_do_not_refresh_jwks_during_cooldown() { + let (issuer, private_key_der, _public_jwk, jwks_requests, _provider) = spawn_oidc_provider().await; + let authenticator = discover_test_authenticator(&issuer).await; + assert_eq!(jwks_requests.load(Ordering::Relaxed), 1); + let gateway = spawn_gateway(authenticator, "http://127.0.0.1:9").await; + let expires_at = jsonwebtoken::get_current_timestamp() + 300; + + for kid in ["unknown-key-1", "unknown-key-2"] { + let response = reqwest::Client::new() + .get(format!("http://{}/v1/models", gateway.address)) + .bearer_auth(identity_token( + &issuer, + "agentic-api", + expires_at, + kid, + &private_key_der, + )) + .send() + .await + .expect("request gateway"); + assert_eq!(response.status(), reqwest::StatusCode::UNAUTHORIZED); + } + + assert_eq!(jwks_requests.load(Ordering::Relaxed), 1); +} + +#[test] +fn oidc_configuration_enforces_secure_endpoints_and_nonempty_audience() { + for issuer in ["https://issuer.example", "http://127.0.0.1:8080", "http://[::1]:8080"] { + OidcConfig::new(issuer, "agentic-api").expect("accepted issuer"); + } + + for issuer in [ + "http://localhost:8080", + "http://192.0.2.1:8080", + "https://issuer.example?tenant=one", + "https://issuer.example#fragment", + ] { + assert!( + OidcConfig::new(issuer, "agentic-api").is_err(), + "{issuer} must be rejected" + ); + } + assert!(OidcConfig::new("https://issuer.example", " \t").is_err()); +} + +#[tokio::test] +async fn discovery_rejects_mismatched_issuer_insecure_jwks_and_oversized_metadata() { + let (issuer, _provider) = spawn_metadata_provider(|_| { + json!({ + "issuer": "https://other-issuer.example", + "jwks_uri": "https://other-issuer.example/jwks" + }) + .to_string() + }) + .await; + assert!(matches!( + OidcAuthenticator::discover(OidcConfig::new(&issuer, TEST_AUDIENCE).expect("OIDC config")).await, + Err(OidcAuthError::IssuerMismatch { .. }) + )); + + let (issuer, _provider) = spawn_metadata_provider(|issuer| { + json!({ + "issuer": issuer, + "jwks_uri": "http://192.0.2.1/jwks" + }) + .to_string() + }) + .await; + assert!(matches!( + OidcAuthenticator::discover(OidcConfig::new(&issuer, TEST_AUDIENCE).expect("OIDC config")).await, + Err(OidcAuthError::InsecureJwksUri) + )); + + let (issuer, _provider) = spawn_metadata_provider(|_| " ".repeat(1024 * 1024 + 1)).await; + assert!(matches!( + OidcAuthenticator::discover(OidcConfig::new(&issuer, TEST_AUDIENCE).expect("OIDC config")).await, + Err(OidcAuthError::ProviderResponseTooLarge) + )); + + let (issuer, _provider) = spawn_chunked_metadata_provider().await; + assert!(matches!( + OidcAuthenticator::discover(OidcConfig::new(&issuer, TEST_AUDIENCE).expect("OIDC config")).await, + Err(OidcAuthError::ProviderResponseTooLarge) + )); +} + +#[tokio::test] +async fn zero_ttl_rotated_jwks_is_coalesced_and_revokes_the_old_key() { + let (issuer, old_private_key, new_private_key, jwks_requests, _provider) = spawn_rotating_oidc_provider().await; + let authenticator = discover_test_authenticator(&issuer).await; + assert_eq!(jwks_requests.load(Ordering::Relaxed), 1); + let (upstream_url, _observed_headers, _upstream) = spawn_models_upstream().await; + let gateway = spawn_gateway(authenticator, &upstream_url).await; + let expires_at = jsonwebtoken::get_current_timestamp() + 300; + let new_token = identity_token(&issuer, "agentic-api", expires_at, "new-key", &new_private_key); + let client = reqwest::Client::new(); + + let first = client + .get(format!("http://{}/v1/models", gateway.address)) + .bearer_auth(&new_token) + .send(); + let second = client + .get(format!("http://{}/v1/models", gateway.address)) + .bearer_auth(&new_token) + .send(); + let (first, second) = tokio::join!(first, second); + assert_eq!( + first.expect("first rotated-key request").status(), + reqwest::StatusCode::OK + ); + assert_eq!( + second.expect("second rotated-key request").status(), + reqwest::StatusCode::OK + ); + assert_eq!(jwks_requests.load(Ordering::Relaxed), 2); + + let cached = client + .get(format!("http://{}/v1/models", gateway.address)) + .bearer_auth(&new_token) + .send() + .await + .expect("cached rotated-key request"); + assert_eq!(cached.status(), reqwest::StatusCode::OK); + assert_eq!(jwks_requests.load(Ordering::Relaxed), 2); + + let revoked = client + .get(format!("http://{}/v1/models", gateway.address)) + .bearer_auth(identity_token( + &issuer, + "agentic-api", + expires_at, + "old-key", + &old_private_key, + )) + .send() + .await + .expect("revoked-key request"); + assert_eq!(revoked.status(), reqwest::StatusCode::UNAUTHORIZED); +} + +#[tokio::test] +async fn jwks_refresh_failure_returns_protocol_specific_service_errors() { + let (issuer, private_key, jwks_requests, _provider) = spawn_failing_refresh_provider().await; + let authenticator = discover_test_authenticator(&issuer).await; + let gateway = spawn_gateway(authenticator, "http://127.0.0.1:9").await; + let token = identity_token( + &issuer, + "agentic-api", + jsonwebtoken::get_current_timestamp() + 300, + "test-key", + &private_key, + ); + let client = reqwest::Client::new(); + + let openai_response = client + .get(format!("http://{}/v1/models", gateway.address)) + .bearer_auth(&token) + .send(); + let anthropic_response = client + .post(format!("http://{}/v1/messages", gateway.address)) + .bearer_auth(&token) + .send(); + let (response, anthropic_response) = tokio::join!(openai_response, anthropic_response); + let response = response.expect("OpenAI-style dependency failure"); + assert_eq!(response.status(), reqwest::StatusCode::SERVICE_UNAVAILABLE); + assert!(response.headers().get(reqwest::header::WWW_AUTHENTICATE).is_none()); + let body = response.json::().await.expect("OpenAI service error"); + assert_eq!(body["error"]["code"], "authentication_service_unavailable"); + + let response = anthropic_response.expect("Anthropic-style dependency failure"); + assert_eq!(response.status(), reqwest::StatusCode::SERVICE_UNAVAILABLE); + let body = response.json::().await.expect("Anthropic service error"); + assert_eq!(body["error"]["type"], "api_error"); + assert_eq!(jwks_requests.load(Ordering::Relaxed), 2); +} + +#[tokio::test] +async fn multi_audience_tokens_require_the_expected_authorized_party() { + let (issuer, private_key, _public_jwk, _jwks_requests, _provider) = spawn_oidc_provider().await; + let authenticator = discover_test_authenticator(&issuer).await; + let (upstream_url, _observed_headers, _upstream) = spawn_models_upstream().await; + let gateway = spawn_gateway(authenticator, &upstream_url).await; + + for (authorized_party, expected_status) in [ + (Some("other-client"), reqwest::StatusCode::UNAUTHORIZED), + (None, reqwest::StatusCode::UNAUTHORIZED), + (Some("agentic-api"), reqwest::StatusCode::OK), + ] { + let response = reqwest::Client::new() + .get(format!("http://{}/v1/models", gateway.address)) + .bearer_auth(identity_token_with_audiences( + &issuer, + &["agentic-api", "other-client"], + authorized_party, + &private_key, + )) + .send() + .await + .expect("multi-audience request"); + assert_eq!(response.status(), expected_status); + } +} + +#[tokio::test] +async fn authenticated_principal_is_inserted_and_identity_header_is_removed() { + let (issuer, private_key, _public_jwk, _jwks_requests, _provider) = spawn_oidc_provider().await; + let authenticator = discover_test_authenticator(&issuer).await; + let router = Router::new() + .route( + "/v1/principal", + get( + |Extension(principal): Extension, headers: HeaderMap| async move { + Json(json!({ + "issuer": principal.issuer(), + "subject": principal.subject(), + "authorization_present": headers.contains_key(reqwest::header::AUTHORIZATION) + })) + }, + ), + ) + .route_layer(middleware::from_fn_with_state(authenticator, require_oidc)); + let gateway = spawn_router(router).await; + + let response = reqwest::Client::new() + .get(format!("http://{}/v1/principal", gateway.address)) + .bearer_auth(identity_token( + &issuer, + "agentic-api", + jsonwebtoken::get_current_timestamp() + 300, + "test-key", + &private_key, + )) + .send() + .await + .expect("principal request"); + assert_eq!(response.status(), reqwest::StatusCode::OK); + let body = response.json::().await.expect("principal JSON"); + assert_eq!(body["issuer"], issuer); + assert_eq!(body["subject"], "github-user-123"); + assert_eq!(body["authorization_present"], false); +} + +#[tokio::test] +async fn every_v1_route_rejects_missing_credentials() { + let (issuer, _private_key, _public_jwk, _jwks_requests, _provider) = spawn_oidc_provider().await; + let authenticator = discover_test_authenticator(&issuer).await; + let gateway = spawn_gateway(authenticator, "http://127.0.0.1:9").await; + let client = reqwest::Client::new(); + + for path in [ + "/v1/conversations", + "/v1/messages", + "/v1/messages/count_tokens", + "/v1/responses", + ] { + let response = client + .post(format!("http://{}{path}", gateway.address)) + .send() + .await + .expect("protected POST"); + assert_eq!(response.status(), reqwest::StatusCode::UNAUTHORIZED, "{path}"); + } + let models = client + .get(format!("http://{}/v1/models", gateway.address)) + .send() + .await + .expect("protected models request"); + assert_eq!(models.status(), reqwest::StatusCode::UNAUTHORIZED); + + let websocket_error = tokio_tungstenite::connect_async(format!("ws://{}/v1/responses", gateway.address)) + .await + .expect_err("missing bearer must reject WebSocket upgrade"); + match websocket_error { + tokio_tungstenite::tungstenite::Error::Http(response) => { + assert_eq!(response.status(), reqwest::StatusCode::UNAUTHORIZED); + } + error => panic!("unexpected WebSocket error: {error}"), + } + + let ready = client + .get(format!("http://{}/ready", gateway.address)) + .send() + .await + .expect("public readiness request"); + assert_ne!(ready.status(), reqwest::StatusCode::UNAUTHORIZED); +} diff --git a/docs/api/index.md b/docs/api/index.md index f1b8d370..cf43816c 100644 --- a/docs/api/index.md +++ b/docs/api/index.md @@ -1,5 +1,45 @@ # API Reference +## Authentication + +Inbound authentication is optional. When the gateway starts with both `OIDC_ISSUER` and `OIDC_AUDIENCE`, every +`/v1/*` HTTP route and the `/v1/responses` WebSocket upgrade require an OIDC `Authorization: Bearer `. +`/health` and `/ready` remain public. Supplying only one OIDC setting is a startup error. + +The gateway validates the token signature, issuer, audience, authorized party for multi-audience tokens, subject, +expiration, and not-before time. It consumes the identity token at the gateway boundary instead of forwarding it to +the inference service. WebSocket sessions reject new `response.create` messages after the validated token expires. + +Missing or rejected credentials return `401 Unauthorized` with `WWW-Authenticate: Bearer`. OpenAI-compatible routes +use this envelope: + +```json +{ + "error": { + "message": "invalid bearer token", + "type": "authentication_error", + "param": null, + "code": "invalid_token" + } +} +``` + +`/v1/messages` and `/v1/messages/count_tokens` use the Anthropic-compatible envelope: + +```json +{ + "type": "error", + "error": { + "type": "authentication_error", + "message": "invalid bearer token" + } +} +``` + +A JWKS refresh failure returns `503 Service Unavailable`, without `WWW-Authenticate`, so clients can distinguish an +identity-provider dependency failure from rejected credentials. See +[OIDC bearer authentication](../design/oidc-bearer-authentication.md) for configuration and key-cache behavior. + ## Responses ### `POST /v1/responses` diff --git a/docs/deploying/container.md b/docs/deploying/container.md index ecb736df..16ec1e5b 100644 --- a/docs/deploying/container.md +++ b/docs/deploying/container.md @@ -30,6 +30,8 @@ The image starts `agentic-server` in standalone mode. At minimum, set `LLM_API_B | `GATEWAY_PORT` | `9000` | Listen port | | `DATABASE_URL` | `sqlite://./agentic_api.db` | SQLite or PostgreSQL persistence URL | | `OPENAI_API_KEY` | none | Credential sent to the upstream service when the client does not supply one | +| `OIDC_ISSUER` | none | Optional OIDC issuer for inbound bearer-token authentication | +| `OIDC_AUDIENCE` | none | Required token audience when `OIDC_ISSUER` is set | | `SKIP_LLM_READY_CHECK` | `false` | Skip the startup probe for hosted providers without `/health` | | `CORS_ALLOWED_ORIGINS` | none | Comma-separated browser origins | @@ -46,7 +48,25 @@ docker run --rm --name agentic-api \ agentic-api:dev ``` -The gateway does not provide inbound client authentication. `OPENAI_API_KEY` is an upstream credential, not a password for callers, so keep the port bound to loopback unless an authenticated ingress or proxy protects it. +Inbound authentication is disabled by default. `OPENAI_API_KEY` is an upstream credential, not a password for callers, +so keep the port bound to loopback unless an authenticated ingress protects it or OIDC is enabled. + +To enable OIDC bearer authentication, configure both variables: + +```console +docker run --rm --name agentic-api \ + --publish 127.0.0.1:9000:9000 \ + --env LLM_API_BASE=https://vllm.example.com \ + --env OPENAI_API_KEY \ + --env OIDC_ISSUER=https://identity.example.com \ + --env OIDC_AUDIENCE=agentic-api \ + agentic-api:dev +``` + +The gateway discovers the provider and its JSON Web Key Set before listening. `/health` and `/ready` remain public; +all `/v1/*` routes then require an OIDC `Authorization: Bearer` token. The identity token is consumed by the gateway, +and `OPENAI_API_KEY` supplies the upstream inference credential. See +[OIDC bearer authentication](../design/oidc-bearer-authentication.md) for the validation and key-rotation contract. If the upstream is running on the Docker host, use `http://host.docker.internal:` on Docker Desktop. On Linux, add `--add-host host.docker.internal:host-gateway`. diff --git a/docs/design/oidc-bearer-authentication.md b/docs/design/oidc-bearer-authentication.md new file mode 100644 index 00000000..fedf599d --- /dev/null +++ b/docs/design/oidc-bearer-authentication.md @@ -0,0 +1,67 @@ +# OIDC bearer authentication + +## Scope + +vLLM Agentic API can optionally authenticate API callers with JSON Web Tokens issued by an OpenID Connect (OIDC) +provider. This is the first authentication slice for [issue #104](https://github.com/vllm-project/agentic-api/issues/104): +it establishes a verified principal at the HTTP boundary without adding a browser login, callback, or server-side +session. + +The bearer-token model works with API clients such as Codex and Claude Code. GitHub login can be supplied by an OIDC +provider configured to federate GitHub identities; the gateway does not add GitHub-specific authorization logic. + +## Configuration and startup + +Authentication is disabled unless both `OIDC_ISSUER` and `OIDC_AUDIENCE` are set. Supplying only one is a startup +error. When enabled, the gateway: + +1. fetches the issuer's `/.well-known/openid-configuration` document without following redirects; +2. requires the discovered issuer to match the configured issuer; +3. fetches and caches the JSON Web Key Set (JWKS); +4. refuses to listen if discovery or the initial JWKS request fails. + +Issuer and JWKS URLs must use HTTPS. HTTP is accepted only for literal loopback IP addresses (`127.0.0.1` or `::1`) +in local tests and development, and an HTTPS issuer cannot redirect JWKS retrieval to loopback HTTP. Provider +responses are limited to 1 MiB and JWKS documents to 100 keys. + +Verification keys are cached for the provider's `Cache-Control: max-age` duration, capped at one hour, or five minutes +when no cache lifetime is supplied. A stale cache is refreshed before a cached key is accepted, so a provider can +remove a compromised key without requiring a gateway restart. Unknown key IDs can trigger at most one refresh per +30-second cooldown after a completed fetch. Refreshes are single-flight, and every successfully fetched key set is +installed even when it does not contain the key requested by the triggering token. +Concurrent refresh waiters reuse the completed result. After a refresh failure, another provider request is suppressed +for 30 seconds and callers receive `503 Service Unavailable`; a one-second coalescing window also prevents a +provider-supplied zero-second cache lifetime from causing one fetch per concurrent request. + +## Request boundary + +`/health` and `/ready` remain public so orchestrators can probe the process. Every `/v1/*` route requires +`Authorization: Bearer ` when OIDC is enabled, including HTTP streaming and the Responses WebSocket upgrade. +`/ready` continues to report inference-service readiness; it does not treat a temporary identity-provider refresh +failure as a reason to remove an otherwise healthy gateway from service. Those request-time dependency failures +return `503 Service Unavailable` as described above. + +The gateway verifies: + +- an asymmetric token signing algorithm and a signature from the provider JWKS; +- a signing key whose `kid`, `alg`, `use`, and `key_ops` permit verification; +- required `iss`, `aud`, `sub`, and `exp` claims; +- issuer and audience equality, plus `azp` equality when a token has multiple audiences; +- expiration and, when present, the not-before time. + +Successful authentication inserts the stable issuer and subject pair into request extensions as the authenticated +principal. Tenant and persisted-state authorization remain follow-up work under +[issue #107](https://github.com/vllm-project/agentic-api/issues/107). + +For WebSockets, authentication occurs during the HTTP upgrade. The validated expiration is retained with the +principal, and the gateway rejects new `response.create` messages after the token expires (including clock skew). + +## Credential separation + +The verified identity token is consumed at the gateway and is never forwarded to the inference service. OpenAI-style +upstream requests use `OPENAI_API_KEY` after authentication removes the inbound `Authorization` header. +Anthropic-compatible requests may continue to supply an upstream `x-api-key`; otherwise they also fall back to +`OPENAI_API_KEY`. + +This separation prevents an OIDC identity token from being mistaken for an inference-provider credential. Deployments +that do not enable OIDC retain the existing pass-through behavior for client credentials.