Add workload identity token exchange support (#37610)

## What changed

- Add the `codex-workload-identity` crate for exchanging a file-backed JWT assertion and federation rule ID for short-lived ChatGPT credentials.
- Cache valid access tokens, refresh them before expiry or after rejection, and coalesce concurrent exchanges. Continue using a still-valid cached token when a proactive refresh fails transiently.
- Validate assertion files, token endpoints, and exchange responses; honor outbound proxy policy for HTTPS endpoints and redact access tokens from debug output.

## Testing

- Cover request encoding, assertion rotation, caching, concurrent refreshes, transient-failure fallback, configuration validation, and malformed inputs and responses.

GitOrigin-RevId: 5496851683c2dcf6aaad6840053b97f7c0be076e
This commit is contained in:
cooper-oai
2026-08-08 17:12:54 +00:00
committed by copyberry
parent c4513cb982
commit 936f5eb3ee
8 changed files with 832 additions and 0 deletions

View File

@@ -0,0 +1,6 @@
load("//:defs.bzl", "codex_rust_crate")
codex_rust_crate(
name = "workload-identity",
crate_name = "codex_workload_identity",
)

View File

@@ -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 }

View File

@@ -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<String, WorkloadIdentityError> {
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())
}

View File

@@ -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<CacheState>,
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<Self, WorkloadIdentityError> {
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<WorkloadIdentityToken, WorkloadIdentityError> {
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<WorkloadIdentityToken, WorkloadIdentityError> {
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<WorkloadIdentityToken, WorkloadIdentityError> {
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::<TokenExchangeResponse>(&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<bool, WorkloadIdentityError> {
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<CachedToken>,
last_attempt_error: Option<WorkloadIdentityError>,
token_generation: u64,
}
impl CacheState {
fn store(
&mut self,
mut token: WorkloadIdentityToken,
valid_from: Instant,
now: Instant,
) -> Result<WorkloadIdentityToken, WorkloadIdentityError> {
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<WorkloadIdentityToken> {
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<String>,
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<String>,
expires_in: u64,
issued_token_type: String,
scope: String,
token_type: String,
user_id: String,
}
impl TokenExchangeResponse {
fn into_token(self) -> Result<WorkloadIdentityToken, WorkloadIdentityError> {
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;

View File

@@ -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<Self, WorkloadIdentityError> {
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<std::io::Error>,
},
#[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,
}

View File

@@ -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<_>>(),
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::<Vec<_>>();
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::<Vec<_>>();
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<_>>(),
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)
));
}