mirror of
https://github.com/openai/codex.git
synced 2026-08-23 13:09:46 +00:00
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:
6
codex-rs/workload-identity/BUILD.bazel
Normal file
6
codex-rs/workload-identity/BUILD.bazel
Normal file
@@ -0,0 +1,6 @@
|
||||
load("//:defs.bzl", "codex_rust_crate")
|
||||
|
||||
codex_rust_crate(
|
||||
name = "workload-identity",
|
||||
crate_name = "codex_workload_identity",
|
||||
)
|
||||
27
codex-rs/workload-identity/Cargo.toml
Normal file
27
codex-rs/workload-identity/Cargo.toml
Normal 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 }
|
||||
35
codex-rs/workload-identity/src/assertion.rs
Normal file
35
codex-rs/workload-identity/src/assertion.rs
Normal 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())
|
||||
}
|
||||
372
codex-rs/workload-identity/src/exchange.rs
Normal file
372
codex-rs/workload-identity/src/exchange.rs
Normal 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;
|
||||
63
codex-rs/workload-identity/src/lib.rs
Normal file
63
codex-rs/workload-identity/src/lib.rs
Normal 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,
|
||||
}
|
||||
313
codex-rs/workload-identity/src/workload_identity_tests.rs
Normal file
313
codex-rs/workload-identity/src/workload_identity_tests.rs
Normal 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)
|
||||
));
|
||||
}
|
||||
Reference in New Issue
Block a user