diff --git a/codex-rs/Cargo.lock b/codex-rs/Cargo.lock index f3a2f5d684..0255bcf0bb 100644 --- a/codex-rs/Cargo.lock +++ b/codex-rs/Cargo.lock @@ -4700,6 +4700,21 @@ dependencies = [ "windows-sys 0.52.0", ] +[[package]] +name = "codex-workload-identity" +version = "0.0.0" +dependencies = [ + "codex-http-client", + "pretty_assertions", + "serde", + "serde_json", + "tempfile", + "thiserror 2.0.18", + "tokio", + "url", + "wiremock", +] + [[package]] name = "color-eyre" version = "0.6.5" diff --git a/codex-rs/Cargo.toml b/codex-rs/Cargo.toml index 8ce32222e2..e776322fa1 100644 --- a/codex-rs/Cargo.toml +++ b/codex-rs/Cargo.toml @@ -94,6 +94,7 @@ members = [ "tools", "v8-poc", "websocket-client", + "workload-identity", "utils/absolute-path", "utils/audio", "utils/path-uri", diff --git a/codex-rs/workload-identity/BUILD.bazel b/codex-rs/workload-identity/BUILD.bazel new file mode 100644 index 0000000000..ac046d02d4 --- /dev/null +++ b/codex-rs/workload-identity/BUILD.bazel @@ -0,0 +1,6 @@ +load("//:defs.bzl", "codex_rust_crate") + +codex_rust_crate( + name = "workload-identity", + crate_name = "codex_workload_identity", +) diff --git a/codex-rs/workload-identity/Cargo.toml b/codex-rs/workload-identity/Cargo.toml new file mode 100644 index 0000000000..b64177282a --- /dev/null +++ b/codex-rs/workload-identity/Cargo.toml @@ -0,0 +1,27 @@ +[package] +edition.workspace = true +license.workspace = true +name = "codex-workload-identity" +version.workspace = true + +[lib] +doctest = false +name = "codex_workload_identity" +path = "src/lib.rs" + +[lints] +workspace = true + +[dependencies] +codex-http-client = { workspace = true } +serde = { workspace = true, features = ["derive"] } +serde_json = { workspace = true } +thiserror = { workspace = true } +tokio = { workspace = true, features = ["fs", "io-util", "sync"] } +url = { workspace = true } + +[dev-dependencies] +pretty_assertions = { workspace = true } +tempfile = { workspace = true } +tokio = { workspace = true, features = ["macros", "rt-multi-thread"] } +wiremock = { workspace = true } diff --git a/codex-rs/workload-identity/src/assertion.rs b/codex-rs/workload-identity/src/assertion.rs new file mode 100644 index 0000000000..0c9d755e06 --- /dev/null +++ b/codex-rs/workload-identity/src/assertion.rs @@ -0,0 +1,35 @@ +use std::path::Path; + +use tokio::io::AsyncReadExt; + +use crate::WorkloadIdentityError; + +const MAX_ASSERTION_BYTES: u64 = 16 * 1024; + +/// Reopens the assertion file for each exchange so its owner can rotate the credential. +pub(crate) async fn read_assertion(path: &Path) -> Result { + let file = tokio::fs::File::open(path).await.map_err(|source| { + WorkloadIdentityError::AssertionFile { + path: path.to_path_buf(), + source: source.into(), + } + })?; + let mut bytes = Vec::new(); + file.take(MAX_ASSERTION_BYTES + 1) + .read_to_end(&mut bytes) + .await + .map_err(|source| WorkloadIdentityError::AssertionFile { + path: path.to_path_buf(), + source: source.into(), + })?; + if bytes.len() as u64 > MAX_ASSERTION_BYTES { + return Err(WorkloadIdentityError::AssertionTooLarge); + } + let assertion = + String::from_utf8(bytes).map_err(|_| WorkloadIdentityError::InvalidAssertion)?; + let assertion = assertion.trim(); + if assertion.is_empty() || assertion.as_bytes().contains(&0) { + return Err(WorkloadIdentityError::InvalidAssertion); + } + Ok(assertion.to_string()) +} diff --git a/codex-rs/workload-identity/src/exchange.rs b/codex-rs/workload-identity/src/exchange.rs new file mode 100644 index 0000000000..7bbd73014c --- /dev/null +++ b/codex-rs/workload-identity/src/exchange.rs @@ -0,0 +1,372 @@ +use std::fmt; +use std::sync::atomic::AtomicU64; +use std::sync::atomic::Ordering; +use std::time::Duration; +use std::time::Instant; + +use codex_http_client::ClientRouteClass; +use codex_http_client::HttpClient; +use codex_http_client::HttpClientBuilder; +use codex_http_client::HttpClientFactory; +use serde::Deserialize; +use tokio::sync::Mutex; +use url::Host; +use url::Url; + +use crate::WorkloadIdentityConfig; +use crate::WorkloadIdentityError; +use crate::assertion::read_assertion; + +const ACCESS_TOKEN_TYPE: &str = "urn:ietf:params:oauth:token-type:access_token"; +const JWT_BEARER_GRANT_TYPE: &str = "urn:ietf:params:oauth:grant-type:jwt-bearer"; +const MAX_ACCESS_TOKEN_LIFETIME: Duration = Duration::from_secs(60 * 60); +const MAX_RESPONSE_BYTES: usize = 1024 * 1024; +const REQUEST_TIMEOUT: Duration = Duration::from_secs(30); +const TRANSIENT_FAILURE_RETRY_DELAY: Duration = Duration::from_secs(30); + +/// Exchanges assertions and retains only the current short-lived access token in memory. +pub struct WorkloadIdentityExchange { + client: HttpClient, + completed_attempts: AtomicU64, + config: WorkloadIdentityConfig, + state: Mutex, + token_url: Url, +} + +impl WorkloadIdentityExchange { + /// Creates an exchange against a token URL selected by the trusted caller. + /// + /// HTTP is accepted only for loopback development servers. The supplied factory governs + /// production proxy and custom-CA policy; loopback requests connect directly. + pub fn new( + config: WorkloadIdentityConfig, + token_url: Url, + http_client_factory: HttpClientFactory, + ) -> Result { + let is_loopback = validate_token_url(&token_url)?; + let builder = HttpClientBuilder::new() + .without_redirects() + .without_request_logging(); + let client = if is_loopback { + builder + .build_direct() + .map_err(|_| WorkloadIdentityError::HttpClientConfiguration) + } else { + builder + .build_respecting_outbound_proxy_policy( + &http_client_factory, + token_url.as_str(), + ClientRouteClass::Auth, + ) + .map_err(|_| WorkloadIdentityError::HttpClientConfiguration) + }?; + Ok(Self { + client, + completed_attempts: AtomicU64::new(0), + config, + state: Mutex::new(CacheState::default()), + token_url, + }) + } + + /// Returns a cached token when possible and otherwise performs one shared exchange. + #[expect( + clippy::await_holding_invalid_type, + reason = "the mutex intentionally provides single-flight exchange ownership" + )] + pub async fn resolve(&self) -> Result { + let observed_attempts = self.completed_attempts.load(Ordering::Acquire); + let mut state = self.state.lock().await; + let now = Instant::now(); + if let Some(cached) = &state.cached + && cached.refresh_at > now + && let Some(token) = cached.token_at(now) + { + return Ok(token); + } + if self.completed_attempts.load(Ordering::Acquire) != observed_attempts { + if let Some(error) = state.last_attempt_error.clone() { + return Err(error); + } + if let Some(token) = state + .cached + .as_ref() + .and_then(|cached| cached.token_at(now)) + { + return Ok(token); + } + } + + let valid_from = Instant::now(); + let result = match self.exchange_uncached().await { + Ok(token) => state.store(token, valid_from, Instant::now()), + Err(error) if error.allows_cached_fallback() => { + let now = Instant::now(); + match state.cached.as_mut() { + Some(cached) if cached.expires_at > now => { + cached.refresh_at = + std::cmp::min(now + TRANSIENT_FAILURE_RETRY_DELAY, cached.expires_at); + cached + .token_at(now) + .ok_or(WorkloadIdentityError::InvalidExchangeResponse) + } + _ => Err(error), + } + } + Err(error) => Err(error), + }; + self.complete_attempt(&mut state, result.as_ref().err()); + result + } + + /// Exchanges after a downstream service rejects `observed_token_version`. + /// + /// Concurrent callers that rejected the same token share the first caller's result. + #[expect( + clippy::await_holding_invalid_type, + reason = "the mutex intentionally provides single-flight exchange ownership" + )] + pub async fn refresh( + &self, + observed_token_version: u64, + ) -> Result { + let observed_attempts = self.completed_attempts.load(Ordering::Acquire); + let mut state = self.state.lock().await; + if state.token_generation != observed_token_version + && let Some(token) = state + .cached + .as_ref() + .and_then(|cached| cached.token_at(Instant::now())) + { + return Ok(token); + } + if self.completed_attempts.load(Ordering::Acquire) != observed_attempts + && let Some(error) = state.last_attempt_error.clone() + { + return Err(error); + } + + state.cached = None; + let valid_from = Instant::now(); + let result = self + .exchange_uncached() + .await + .and_then(|token| state.store(token, valid_from, Instant::now())); + self.complete_attempt(&mut state, result.as_ref().err()); + result + } + + fn complete_attempt(&self, state: &mut CacheState, error: Option<&WorkloadIdentityError>) { + state.last_attempt_error = error.cloned(); + self.completed_attempts.fetch_add(1, Ordering::Release); + } + + async fn exchange_uncached(&self) -> Result { + let assertion = read_assertion(&self.config.assertion_file).await?; + let body = url::form_urlencoded::Serializer::new(String::new()) + .append_pair("grant_type", JWT_BEARER_GRANT_TYPE) + .append_pair("assertion", &assertion) + .append_pair("federation_rule_id", &self.config.federation_rule_id) + .finish(); + let response = self + .client + .post(self.token_url.as_str()) + .header("content-type", "application/x-www-form-urlencoded") + .body(body) + .timeout(REQUEST_TIMEOUT) + .send() + .await + .map_err(|_| WorkloadIdentityError::ExchangeUnavailable)?; + if !response.status().is_success() { + return Err(WorkloadIdentityError::ExchangeRejected( + response.status().as_u16(), + )); + } + if response + .content_length() + .is_some_and(|length| length > MAX_RESPONSE_BYTES as u64) + { + return Err(WorkloadIdentityError::InvalidExchangeResponse); + } + + let mut response = response; + let mut bytes = Vec::new(); + while let Some(chunk) = response + .chunk() + .await + .map_err(|_| WorkloadIdentityError::ExchangeUnavailable)? + { + if chunk.len() > MAX_RESPONSE_BYTES.saturating_sub(bytes.len()) { + return Err(WorkloadIdentityError::InvalidExchangeResponse); + } + bytes.extend_from_slice(&chunk); + } + serde_json::from_slice::(&bytes) + .map_err(|_| WorkloadIdentityError::InvalidExchangeResponse)? + .into_token() + } +} + +impl WorkloadIdentityError { + fn allows_cached_fallback(&self) -> bool { + matches!( + self, + Self::AssertionFile { .. } + | Self::ExchangeUnavailable + | Self::ExchangeRejected(408 | 429 | 500..=599) + ) + } +} + +fn validate_token_url(url: &Url) -> Result { + if !url.username().is_empty() + || url.password().is_some() + || url.query().is_some() + || url.fragment().is_some() + { + return Err(WorkloadIdentityError::InvalidTokenUrl); + } + let is_loopback = match url.host().ok_or(WorkloadIdentityError::InvalidTokenUrl)? { + Host::Domain(domain) => domain.eq_ignore_ascii_case("localhost"), + Host::Ipv4(address) => address.is_loopback(), + Host::Ipv6(address) => address.is_loopback(), + }; + if url.scheme() != "https" && !(url.scheme() == "http" && is_loopback) { + return Err(WorkloadIdentityError::InvalidTokenUrl); + } + Ok(is_loopback) +} + +#[derive(Default)] +struct CacheState { + cached: Option, + last_attempt_error: Option, + token_generation: u64, +} + +impl CacheState { + fn store( + &mut self, + mut token: WorkloadIdentityToken, + valid_from: Instant, + now: Instant, + ) -> Result { + let token_generation = self.token_generation.saturating_add(1); + token.version = token_generation; + let cached = CachedToken::new(token, valid_from); + let token = cached + .token_at(now) + .ok_or(WorkloadIdentityError::InvalidExchangeResponse)?; + self.cached = Some(cached); + self.token_generation = token_generation; + Ok(token) + } +} + +struct CachedToken { + expires_at: Instant, + refresh_at: Instant, + token: WorkloadIdentityToken, +} + +impl CachedToken { + fn new(token: WorkloadIdentityToken, valid_from: Instant) -> Self { + let lifetime = Duration::from_secs(token.expires_in); + let refresh_margin = std::cmp::min(Duration::from_secs(120), lifetime / 2); + Self { + expires_at: valid_from + lifetime, + refresh_at: valid_from + lifetime.saturating_sub(refresh_margin), + token, + } + } + + fn token_at(&self, now: Instant) -> Option { + let remaining = self.expires_at.checked_duration_since(now)?; + if remaining.is_zero() { + return None; + } + let mut token = self.token.clone(); + token.expires_in = remaining + .as_secs() + .saturating_add(u64::from(remaining.subsec_nanos() != 0)); + Some(token) + } +} + +#[derive(Clone, PartialEq, Eq)] +pub struct WorkloadIdentityToken { + pub access_token: String, + pub chatgpt_account_id: String, + pub chatgpt_account_user_id: String, + pub chatgpt_plan_type: Option, + pub expires_in: u64, + pub scope: String, + pub user_id: String, + version: u64, +} + +impl WorkloadIdentityToken { + pub fn version(&self) -> u64 { + self.version + } +} + +impl fmt::Debug for WorkloadIdentityToken { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("WorkloadIdentityToken") + .field("access_token", &"[redacted]") + .field("expires_in", &self.expires_in) + .field("version", &self.version) + .finish_non_exhaustive() + } +} + +#[derive(Deserialize)] +struct TokenExchangeResponse { + access_token: String, + chatgpt_account_id: String, + chatgpt_account_user_id: String, + chatgpt_plan_type: Option, + expires_in: u64, + issued_token_type: String, + scope: String, + token_type: String, + user_id: String, +} + +impl TokenExchangeResponse { + fn into_token(self) -> Result { + let lifetime = Duration::from_secs(self.expires_in); + if self.access_token.trim().is_empty() + || self.issued_token_type != ACCESS_TOKEN_TYPE + || !self.token_type.eq_ignore_ascii_case("bearer") + || lifetime.is_zero() + || lifetime > MAX_ACCESS_TOKEN_LIFETIME + || self.scope.trim().is_empty() + || self.chatgpt_account_id.trim().is_empty() + || self.chatgpt_account_user_id.trim().is_empty() + || self.user_id.trim().is_empty() + || self + .chatgpt_plan_type + .as_deref() + .is_some_and(|plan_type| plan_type.trim().is_empty()) + { + return Err(WorkloadIdentityError::InvalidExchangeResponse); + } + Ok(WorkloadIdentityToken { + access_token: self.access_token, + chatgpt_account_id: self.chatgpt_account_id, + chatgpt_account_user_id: self.chatgpt_account_user_id, + chatgpt_plan_type: self.chatgpt_plan_type, + expires_in: self.expires_in, + scope: self.scope, + user_id: self.user_id, + version: 0, + }) + } +} + +#[cfg(test)] +#[path = "workload_identity_tests.rs"] +mod tests; diff --git a/codex-rs/workload-identity/src/lib.rs b/codex-rs/workload-identity/src/lib.rs new file mode 100644 index 0000000000..570ffec28c --- /dev/null +++ b/codex-rs/workload-identity/src/lib.rs @@ -0,0 +1,63 @@ +mod assertion; +mod exchange; + +use std::path::PathBuf; +use std::sync::Arc; + +pub use exchange::WorkloadIdentityExchange; +pub use exchange::WorkloadIdentityToken; +use thiserror::Error; + +/// The inputs needed to exchange a file-backed assertion for ChatGPT auth. +#[derive(Clone)] +pub struct WorkloadIdentityConfig { + pub(crate) assertion_file: PathBuf, + pub(crate) federation_rule_id: String, +} + +impl WorkloadIdentityConfig { + pub fn new( + federation_rule_id: String, + assertion_file: PathBuf, + ) -> Result { + let federation_rule_id = federation_rule_id.trim(); + if federation_rule_id.is_empty() { + return Err(WorkloadIdentityError::InvalidFederationRuleId); + } + if !assertion_file.is_absolute() { + return Err(WorkloadIdentityError::AssertionFileMustBeAbsolute); + } + Ok(Self { + assertion_file, + federation_rule_id: federation_rule_id.to_string(), + }) + } +} + +#[derive(Clone, Debug, Error)] +pub enum WorkloadIdentityError { + #[error("the workload identity federation rule ID must not be empty")] + InvalidFederationRuleId, + #[error("the workload identity assertion file path must be absolute")] + AssertionFileMustBeAbsolute, + #[error("the workload identity assertion is invalid")] + InvalidAssertion, + #[error("the workload identity assertion exceeds 16 KiB")] + AssertionTooLarge, + #[error("could not read workload identity assertion file {path}")] + AssertionFile { + path: PathBuf, + #[source] + source: Arc, + }, + #[error("could not configure the workload identity HTTP client")] + HttpClientConfiguration, + #[error("the workload identity token URL must use HTTPS or loopback HTTP")] + InvalidTokenUrl, + #[error("the workload identity token exchange is unavailable")] + ExchangeUnavailable, + #[error("the workload identity token exchange was rejected with HTTP {0}")] + ExchangeRejected(u16), + #[error("the workload identity token exchange returned an invalid response")] + InvalidExchangeResponse, +} diff --git a/codex-rs/workload-identity/src/workload_identity_tests.rs b/codex-rs/workload-identity/src/workload_identity_tests.rs new file mode 100644 index 0000000000..da587e2706 --- /dev/null +++ b/codex-rs/workload-identity/src/workload_identity_tests.rs @@ -0,0 +1,313 @@ +use std::path::PathBuf; +use std::sync::Arc; +use std::sync::atomic::AtomicUsize; +use std::sync::atomic::Ordering; +use std::time::Duration; + +use codex_http_client::HttpClientFactory; +use codex_http_client::OutboundProxyPolicy; +use pretty_assertions::assert_eq; +use tempfile::TempDir; +use url::Url; +use wiremock::Mock; +use wiremock::MockServer; +use wiremock::ResponseTemplate; +use wiremock::matchers::header; +use wiremock::matchers::method; +use wiremock::matchers::path; + +use super::ACCESS_TOKEN_TYPE; +use super::JWT_BEARER_GRANT_TYPE; +use super::WorkloadIdentityExchange; +use super::WorkloadIdentityToken; +use crate::WorkloadIdentityConfig; +use crate::WorkloadIdentityError; + +fn assertion_file(assertion: &str) -> (TempDir, PathBuf) { + let temp_dir = TempDir::new().expect("tempdir"); + let path = temp_dir.path().join("identity-token"); + std::fs::write(&path, assertion).expect("write assertion"); + (temp_dir, path) +} + +fn make_exchange(path: PathBuf, server: &MockServer) -> WorkloadIdentityExchange { + WorkloadIdentityExchange::new( + WorkloadIdentityConfig::new("idpm_rule_one".to_string(), path).expect("valid config"), + Url::parse(&format!("{}/oauth/token", server.uri())).expect("valid token URL"), + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + ) + .expect("valid exchange") +} + +fn success(access_token: &str, expires_in: u64) -> ResponseTemplate { + ResponseTemplate::new(200).set_body_json(serde_json::json!({ + "access_token": access_token, + "issued_token_type": ACCESS_TOKEN_TYPE, + "token_type": "Bearer", + "expires_in": expires_in, + "scope": "openid profile email chatgpt.workspace.feature.allow-codex-local-access.access", + "chatgpt_account_id": "workspace-one", + "chatgpt_account_user_id": "membership-one", + "user_id": "user-one", + "chatgpt_plan_type": "enterprise" + })) +} + +#[tokio::test] +async fn exchange_sends_three_field_contract_and_caches_valid_response() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/oauth/token")) + .and(header("content-type", "application/x-www-form-urlencoded")) + .respond_with(success("sensitive-access-token", /*expires_in*/ 600)) + .mount(&server) + .await; + let (_temp_dir, assertion_path) = assertion_file("assertion-one\n"); + let exchange = make_exchange(assertion_path, &server); + + let expected = WorkloadIdentityToken { + access_token: "sensitive-access-token".to_string(), + chatgpt_account_id: "workspace-one".to_string(), + chatgpt_account_user_id: "membership-one".to_string(), + chatgpt_plan_type: Some("enterprise".to_string()), + expires_in: 600, + scope: "openid profile email chatgpt.workspace.feature.allow-codex-local-access.access" + .to_string(), + user_id: "user-one".to_string(), + version: 1, + }; + assert_eq!(exchange.resolve().await.expect("exchange"), expected); + assert_eq!(exchange.resolve().await.expect("cached token"), expected); + assert!(!format!("{expected:?}").contains("sensitive-access-token")); + + let requests = server.received_requests().await.expect("requests"); + assert_eq!(requests.len(), 1); + assert_eq!( + url::form_urlencoded::parse(&requests[0].body) + .into_owned() + .collect::>(), + vec![ + ("grant_type".to_string(), JWT_BEARER_GRANT_TYPE.to_string()), + ("assertion".to_string(), "assertion-one".to_string()), + ( + "federation_rule_id".to_string(), + "idpm_rule_one".to_string() + ), + ] + ); +} + +#[tokio::test] +async fn concurrent_resolve_and_rejected_token_refresh_are_single_flight() { + let server = MockServer::start().await; + let calls = Arc::new(AtomicUsize::new(0)); + Mock::given(method("POST")) + .respond_with({ + let calls = Arc::clone(&calls); + move |_request: &wiremock::Request| { + let call = calls.fetch_add(1, Ordering::SeqCst) + 1; + success(&format!("access-{call}"), /*expires_in*/ 600) + .set_delay(Duration::from_millis(50)) + } + }) + .mount(&server) + .await; + let (_temp_dir, assertion_path) = assertion_file("assertion-one"); + let exchange = Arc::new(make_exchange(assertion_path.clone(), &server)); + + let resolves = (0..8) + .map(|_| { + let exchange = Arc::clone(&exchange); + tokio::spawn(async move { exchange.resolve().await }) + }) + .collect::>(); + let mut initial = None; + for resolve in resolves { + let token = resolve.await.expect("join resolve").expect("resolve"); + assert_eq!(token.access_token, "access-1"); + initial = Some(token); + } + tokio::fs::write(&assertion_path, "assertion-two\n") + .await + .expect("rotate assertion"); + let version = initial.expect("initial token").version(); + let refreshes = (0..8) + .map(|_| { + let exchange = Arc::clone(&exchange); + tokio::spawn(async move { exchange.refresh(version).await }) + }) + .collect::>(); + for refresh in refreshes { + assert_eq!( + refresh + .await + .expect("join refresh") + .expect("refresh") + .access_token, + "access-2" + ); + } + + assert_eq!(calls.load(Ordering::SeqCst), 2); + assert_eq!( + server + .received_requests() + .await + .expect("requests") + .iter() + .map(|request| { + url::form_urlencoded::parse(&request.body) + .find(|(name, _)| name == "assertion") + .map(|(_, value)| value.into_owned()) + .expect("assertion field") + }) + .collect::>(), + vec!["assertion-one", "assertion-two"] + ); +} + +#[tokio::test] +async fn rejected_token_waiting_on_proactive_fallback_still_forces_refresh() { + let server = MockServer::start().await; + let calls = Arc::new(AtomicUsize::new(0)); + Mock::given(method("POST")) + .respond_with({ + let calls = Arc::clone(&calls); + move |_request: &wiremock::Request| match calls.fetch_add(1, Ordering::SeqCst) { + 0 => success("access-one", /*expires_in*/ 600), + 1 => ResponseTemplate::new(503).set_delay(Duration::from_millis(200)), + 2.. => success("access-three", /*expires_in*/ 600), + } + }) + .mount(&server) + .await; + let (_temp_dir, assertion_path) = assertion_file("assertion-one"); + let exchange = Arc::new(make_exchange(assertion_path, &server)); + let initial = exchange.resolve().await.expect("initial exchange"); + exchange + .state + .lock() + .await + .cached + .as_mut() + .expect("cached token") + .refresh_at = std::time::Instant::now(); + + let proactive = tokio::spawn({ + let exchange = Arc::clone(&exchange); + async move { exchange.resolve().await } + }); + tokio::time::timeout(Duration::from_secs(1), async { + while calls.load(Ordering::SeqCst) < 2 { + tokio::task::yield_now().await; + } + }) + .await + .expect("proactive exchange started"); + + let forced = exchange.refresh(initial.version()); + tokio::pin!(forced); + assert!( + tokio::time::timeout(Duration::from_millis(20), forced.as_mut()) + .await + .is_err(), + "forced refresh should wait for the proactive exchange" + ); + let fallback = proactive + .await + .expect("join proactive refresh") + .expect("cached fallback"); + assert_eq!(fallback.access_token, initial.access_token); + + let refreshed = forced.await.expect("forced refresh"); + assert_eq!(refreshed.access_token, "access-three"); + assert_ne!(refreshed.version(), initial.version()); + assert_eq!(calls.load(Ordering::SeqCst), 3); +} + +#[tokio::test] +async fn transient_proactive_refresh_failure_uses_still_valid_token() { + let server = MockServer::start().await; + let calls = Arc::new(AtomicUsize::new(0)); + Mock::given(method("POST")) + .respond_with({ + let calls = Arc::clone(&calls); + move |_request: &wiremock::Request| { + if calls.fetch_add(1, Ordering::SeqCst) == 0 { + success("access-one", /*expires_in*/ 4) + } else { + ResponseTemplate::new(503).set_body_string("sensitive server detail") + } + } + }) + .mount(&server) + .await; + let (_temp_dir, assertion_path) = assertion_file("assertion-one"); + let exchange = make_exchange(assertion_path, &server); + let initial = exchange.resolve().await.expect("initial exchange"); + exchange + .state + .lock() + .await + .cached + .as_mut() + .expect("cached token") + .refresh_at = std::time::Instant::now(); + + let fallback = exchange.resolve().await.expect("cached fallback"); + assert_eq!(fallback.access_token, initial.access_token); + assert_eq!(fallback.version(), initial.version()); + assert_eq!(exchange.resolve().await.expect("delayed retry"), fallback); + assert_eq!(calls.load(Ordering::SeqCst), 2); +} + +#[test] +fn configuration_requires_an_absolute_file_and_secure_token_url() { + assert!(matches!( + WorkloadIdentityConfig::new("idpm_rule_one".to_string(), PathBuf::from("relative.jwt")), + Err(WorkloadIdentityError::AssertionFileMustBeAbsolute) + )); + + let (_temp_dir, assertion_path) = assertion_file("assertion-one"); + let config = WorkloadIdentityConfig::new("idpm_rule_one".to_string(), assertion_path) + .expect("valid config"); + assert!(matches!( + WorkloadIdentityExchange::new( + config, + Url::parse("http://auth.example.com/oauth/token").expect("parse URL"), + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + ), + Err(WorkloadIdentityError::InvalidTokenUrl) + )); +} + +#[tokio::test] +async fn exchange_rejects_oversized_assertions_and_incomplete_responses() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({ + "access_token": "access-one", + "issued_token_type": ACCESS_TOKEN_TYPE, + "token_type": "Bearer", + "expires_in": 600, + "scope": "openid", + "chatgpt_account_id": "workspace-one", + "user_id": "user-one" + }))) + .mount(&server) + .await; + let (_temp_dir, assertion_path) = assertion_file(&"x".repeat(16 * 1024 + 1)); + let exchange = make_exchange(assertion_path.clone(), &server); + assert!(matches!( + exchange.resolve().await, + Err(WorkloadIdentityError::AssertionTooLarge) + )); + + tokio::fs::write(&assertion_path, "valid-assertion") + .await + .expect("replace assertion"); + assert!(matches!( + exchange.resolve().await, + Err(WorkloadIdentityError::InvalidExchangeResponse) + )); +}