mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-12 04:09:48 +08:00
fix(provider): harden Agent Identity OAuth lifecycle
This commit is contained in:
@@ -13,11 +13,11 @@ use crypto_box::{
|
||||
};
|
||||
use ed25519_dalek::{
|
||||
pkcs8::{DecodePrivateKey, EncodePrivateKey},
|
||||
Signer, SigningKey,
|
||||
Signature, Signer, SigningKey, Verifier,
|
||||
};
|
||||
use serde::Deserialize;
|
||||
use serde_json::{json, Map, Value};
|
||||
use sha2::{Digest, Sha512};
|
||||
use sha2::{Digest, Sha256, Sha512};
|
||||
use thiserror::Error;
|
||||
use url::Url;
|
||||
|
||||
@@ -40,9 +40,17 @@ const ASSERTION_PREFIX: &str = "AgentAssertion ";
|
||||
const CODEX_AGENT_IDENTITY_AGENT_HARNESS_ID: &str = "codex-cli";
|
||||
const CODEX_AGENT_IDENTITY_RUNNING_LOCATION: &str = "local";
|
||||
|
||||
/// The AgentAssertion scheme is generated internally after an Agent Identity
|
||||
/// task has been registered. Keep the scheme check separate from envelope
|
||||
/// validation so a malformed in-flight assertion is still treated as an
|
||||
/// Agent-originated request by defensive runtime state writers.
|
||||
pub fn is_codex_agent_identity_authorization(value: &str) -> bool {
|
||||
encoded_agent_identity_assertion(value).is_some()
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Error)]
|
||||
pub enum CodexAgentIdentityEnrollmentError {
|
||||
#[error("ChatGPT Session Token 不能为空")]
|
||||
#[error("ChatGPT Access Token 不能为空")]
|
||||
MissingSessionToken,
|
||||
#[error("Agent Identity 注册请求失败")]
|
||||
RegistrationRequestFailed,
|
||||
@@ -67,6 +75,14 @@ struct AgentIdentityCredentials {
|
||||
task_id: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct AgentIdentityAssertionEnvelope {
|
||||
agent_runtime_id: String,
|
||||
task_id: String,
|
||||
timestamp: String,
|
||||
signature: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct AgentTaskRegistrationResponse {
|
||||
#[serde(default)]
|
||||
@@ -143,6 +159,90 @@ impl CodexAgentIdentityRefreshAdapter {
|
||||
}
|
||||
}
|
||||
|
||||
/// Returns a stable, non-secret fingerprint for the Agent Identity key pair
|
||||
/// and runtime. The task id is deliberately excluded so task rotation does not
|
||||
/// look like a credential replacement.
|
||||
pub fn codex_agent_identity_credential_fingerprint(config: &Value) -> Option<String> {
|
||||
let credentials = agent_identity_credentials(config).ok()?;
|
||||
Some(agent_identity_credential_fingerprint_from_credentials(
|
||||
&credentials,
|
||||
))
|
||||
}
|
||||
|
||||
pub fn codex_agent_identity_transport_credential_fingerprint(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<String> {
|
||||
CodexAgentIdentityRefreshAdapter::config_from_transport(transport)
|
||||
.and_then(|config| codex_agent_identity_credential_fingerprint(&config))
|
||||
}
|
||||
|
||||
/// Returns the fencing generation for a transport/config entry, including the
|
||||
/// current task id. It is safe to log/compare but must never be used as a key.
|
||||
pub fn codex_agent_identity_refresh_fingerprint(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: Option<&CachedOAuthEntry>,
|
||||
) -> Option<String> {
|
||||
let transport_config = CodexAgentIdentityRefreshAdapter::config_from_transport(transport)?;
|
||||
let transport_credentials = agent_identity_credentials(&transport_config).ok()?;
|
||||
let transport_credential_fingerprint =
|
||||
agent_identity_credential_fingerprint_from_credentials(&transport_credentials);
|
||||
let config = entry
|
||||
.filter(|entry| {
|
||||
entry.source_fingerprint.as_deref() == Some(transport_credential_fingerprint.as_str())
|
||||
})
|
||||
.and_then(CodexAgentIdentityRefreshAdapter::config_from_entry)
|
||||
.filter(|entry_config| {
|
||||
agent_identity_credentials(entry_config)
|
||||
.ok()
|
||||
.is_some_and(|entry_credentials| {
|
||||
entry_credentials.task_id == transport_credentials.task_id
|
||||
})
|
||||
})
|
||||
.unwrap_or(transport_config);
|
||||
codex_agent_identity_config_refresh_fingerprint(&config)
|
||||
}
|
||||
|
||||
pub fn codex_agent_identity_config_refresh_fingerprint(config: &Value) -> Option<String> {
|
||||
let credentials = agent_identity_credentials(config).ok()?;
|
||||
let credential_fingerprint =
|
||||
agent_identity_credential_fingerprint_from_credentials(&credentials);
|
||||
let task_id = credentials.task_id.as_deref().unwrap_or_default();
|
||||
let mut digest = Sha256::new();
|
||||
digest.update(credential_fingerprint.as_bytes());
|
||||
digest.update([0]);
|
||||
digest.update(task_id.as_bytes());
|
||||
Some(URL_SAFE_NO_PAD.encode(digest.finalize()))
|
||||
}
|
||||
|
||||
pub fn codex_agent_identity_cached_entry_from_transport(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<CachedOAuthEntry> {
|
||||
let config = CodexAgentIdentityRefreshAdapter::config_from_transport(transport)?;
|
||||
let credentials = agent_identity_credentials(&config).ok()?;
|
||||
let task_id = credentials.task_id.as_deref()?;
|
||||
let auth_header_value = build_agent_assertion(&credentials, task_id, Utc::now()).ok()?;
|
||||
Some(CachedOAuthEntry {
|
||||
provider_type: CODEX_AGENT_IDENTITY_CACHED_ENTRY_PROVIDER_TYPE.to_string(),
|
||||
auth_header_name: AUTHORIZATION_HEADER.to_string(),
|
||||
auth_header_value,
|
||||
expires_at_unix_secs: None,
|
||||
metadata: Some(config),
|
||||
source_fingerprint: Some(agent_identity_credential_fingerprint_from_credentials(
|
||||
&credentials,
|
||||
)),
|
||||
})
|
||||
}
|
||||
|
||||
fn agent_identity_credential_fingerprint_from_credentials(
|
||||
credentials: &AgentIdentityCredentials,
|
||||
) -> String {
|
||||
let mut digest = Sha256::new();
|
||||
digest.update(credentials.runtime_id.as_bytes());
|
||||
digest.update([0]);
|
||||
digest.update(credentials.signing_key.to_bytes());
|
||||
URL_SAFE_NO_PAD.encode(digest.finalize())
|
||||
}
|
||||
|
||||
pub fn is_codex_agent_identity_auth_config_value(config: &Value) -> bool {
|
||||
let Some(root) = config.as_object() else {
|
||||
return false;
|
||||
@@ -168,6 +268,90 @@ pub fn is_codex_agent_identity_transport(transport: &GatewayProviderTransportSna
|
||||
.is_some_and(is_codex_agent_identity_auth_config_value)
|
||||
}
|
||||
|
||||
/// Verifies that an in-flight AgentAssertion was signed by the exact Agent
|
||||
/// Identity credential and task represented by the current transport.
|
||||
pub fn codex_agent_identity_authorization_matches_transport(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
authorization: &str,
|
||||
) -> bool {
|
||||
if !is_codex_agent_identity_transport(transport) {
|
||||
return false;
|
||||
}
|
||||
let Some(encoded) = encoded_agent_identity_assertion(authorization) else {
|
||||
return false;
|
||||
};
|
||||
let Ok(envelope_bytes) = URL_SAFE_NO_PAD.decode(encoded) else {
|
||||
return false;
|
||||
};
|
||||
let Ok(envelope) = serde_json::from_slice::<AgentIdentityAssertionEnvelope>(&envelope_bytes)
|
||||
else {
|
||||
return false;
|
||||
};
|
||||
let Some(config) = CodexAgentIdentityRefreshAdapter::config_from_transport(transport) else {
|
||||
return false;
|
||||
};
|
||||
let Ok(credentials) = agent_identity_credentials(&config) else {
|
||||
return false;
|
||||
};
|
||||
let Some(current_task_id) = credentials.task_id.as_deref() else {
|
||||
return false;
|
||||
};
|
||||
let runtime_id = envelope.agent_runtime_id.trim();
|
||||
let task_id = envelope.task_id.trim();
|
||||
let timestamp = envelope.timestamp.trim();
|
||||
if runtime_id != credentials.runtime_id || task_id != current_task_id || timestamp.is_empty() {
|
||||
return false;
|
||||
}
|
||||
let Ok(signature_bytes) = STANDARD.decode(envelope.signature.trim()) else {
|
||||
return false;
|
||||
};
|
||||
let Ok(signature) = Signature::from_slice(&signature_bytes) else {
|
||||
return false;
|
||||
};
|
||||
let payload = format!("{runtime_id}:{task_id}:{timestamp}");
|
||||
credentials
|
||||
.signing_key
|
||||
.verifying_key()
|
||||
.verify(payload.as_bytes(), &signature)
|
||||
.is_ok()
|
||||
}
|
||||
|
||||
/// Agent task rotation may change only task_id. Any other auth_config change
|
||||
/// means the caller must rebuild the whole request from a fresh transport.
|
||||
pub fn codex_agent_identity_transport_allows_task_rotation_from(
|
||||
initial: &GatewayProviderTransportSnapshot,
|
||||
current: &GatewayProviderTransportSnapshot,
|
||||
) -> bool {
|
||||
let Some(initial_config) = CodexAgentIdentityRefreshAdapter::config_from_transport(initial)
|
||||
.and_then(agent_identity_config_without_task)
|
||||
else {
|
||||
return false;
|
||||
};
|
||||
let Some(current_config) = CodexAgentIdentityRefreshAdapter::config_from_transport(current)
|
||||
.and_then(agent_identity_config_without_task)
|
||||
else {
|
||||
return false;
|
||||
};
|
||||
initial_config == current_config
|
||||
}
|
||||
|
||||
pub fn codex_agent_identity_entry_allows_task_rotation_from(
|
||||
initial: &GatewayProviderTransportSnapshot,
|
||||
entry: &CachedOAuthEntry,
|
||||
) -> bool {
|
||||
let Some(initial_config) = CodexAgentIdentityRefreshAdapter::config_from_transport(initial)
|
||||
.and_then(agent_identity_config_without_task)
|
||||
else {
|
||||
return false;
|
||||
};
|
||||
let Some(entry_config) = CodexAgentIdentityRefreshAdapter::config_from_entry(entry)
|
||||
.and_then(agent_identity_config_without_task)
|
||||
else {
|
||||
return false;
|
||||
};
|
||||
initial_config == entry_config
|
||||
}
|
||||
|
||||
pub fn is_codex_agent_identity_cached_entry(entry: &CachedOAuthEntry) -> bool {
|
||||
entry
|
||||
.provider_type
|
||||
@@ -179,6 +363,18 @@ pub fn validate_codex_agent_identity_auth_config(config: &Value) -> Result<(), S
|
||||
agent_identity_credentials(config).map(|_| ())
|
||||
}
|
||||
|
||||
pub fn codex_agent_identity_auth_config_has_task_id(config: &Value) -> bool {
|
||||
let Some(root) = config.as_object() else {
|
||||
return false;
|
||||
};
|
||||
string_from_maps(
|
||||
root,
|
||||
agent_identity_nested_object(root),
|
||||
&["task_id", "taskId"],
|
||||
)
|
||||
.is_some()
|
||||
}
|
||||
|
||||
/// Returns whether an upstream response proves that the registered Agent Identity task is no
|
||||
/// longer usable. Only this condition should trigger task registration again; an arbitrary 401
|
||||
/// can instead mean that the account itself has lost access.
|
||||
@@ -228,6 +424,22 @@ fn agent_identity_nested_object(root: &Map<String, Value>) -> Option<&Map<String
|
||||
.and_then(Value::as_object)
|
||||
}
|
||||
|
||||
fn agent_identity_config_without_task(mut config: Value) -> Option<Value> {
|
||||
if !is_codex_agent_identity_auth_config_value(&config) {
|
||||
return None;
|
||||
}
|
||||
let root = config.as_object_mut()?;
|
||||
root.remove("task_id");
|
||||
root.remove("taskId");
|
||||
for nested_key in ["agent_identity", "agentIdentity"] {
|
||||
if let Some(nested) = root.get_mut(nested_key).and_then(Value::as_object_mut) {
|
||||
nested.remove("task_id");
|
||||
nested.remove("taskId");
|
||||
}
|
||||
}
|
||||
Some(config)
|
||||
}
|
||||
|
||||
fn string_from_map(map: &Map<String, Value>, keys: &[&str]) -> Option<String> {
|
||||
keys.iter().find_map(|key| {
|
||||
map.get(*key)
|
||||
@@ -279,6 +491,19 @@ fn agent_identity_timestamp(now: DateTime<Utc>) -> String {
|
||||
now.to_rfc3339_opts(SecondsFormat::Secs, true)
|
||||
}
|
||||
|
||||
fn encoded_agent_identity_assertion(value: &str) -> Option<&str> {
|
||||
let mut parts = value.split_ascii_whitespace();
|
||||
let scheme = parts.next()?;
|
||||
let encoded = parts.next()?;
|
||||
if !scheme.eq_ignore_ascii_case(ASSERTION_PREFIX.trim())
|
||||
|| encoded.is_empty()
|
||||
|| parts.next().is_some()
|
||||
{
|
||||
return None;
|
||||
}
|
||||
Some(encoded)
|
||||
}
|
||||
|
||||
fn build_agent_assertion(
|
||||
credentials: &AgentIdentityCredentials,
|
||||
task_id: &str,
|
||||
@@ -373,20 +598,44 @@ fn agent_runtime_id_from_registration_response(body: &str) -> Result<String, ()>
|
||||
.ok_or(())
|
||||
}
|
||||
|
||||
/// Uses a ChatGPT session token once to register a fresh Agent Identity. The returned config
|
||||
/// contains only the generated signing credentials and is deliberately free of the session token.
|
||||
/// Uses a ChatGPT access token once to register a fresh Agent Identity. The returned config
|
||||
/// contains only the generated signing credentials and is deliberately free of the access token.
|
||||
pub async fn register_codex_agent_identity_from_access_token(
|
||||
executor: &dyn OAuthHttpExecutor,
|
||||
access_token: &str,
|
||||
network: OAuthNetworkContext,
|
||||
) -> Result<Map<String, Value>, CodexAgentIdentityEnrollmentError> {
|
||||
register_codex_agent_identity_from_access_token_with_auth_api_base_url(
|
||||
executor,
|
||||
access_token,
|
||||
network,
|
||||
CODEX_AGENT_IDENTITY_AUTH_API_BASE_URL,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
/// Registers an Agent Identity and its initial task for backwards compatibility.
|
||||
pub async fn create_codex_agent_identity_from_access_token(
|
||||
executor: &dyn OAuthHttpExecutor,
|
||||
access_token: &str,
|
||||
network: OAuthNetworkContext,
|
||||
) -> Result<Map<String, Value>, CodexAgentIdentityEnrollmentError> {
|
||||
create_codex_agent_identity_from_session_token_with_auth_api_base_url(
|
||||
executor,
|
||||
access_token,
|
||||
network,
|
||||
CODEX_AGENT_IDENTITY_AUTH_API_BASE_URL,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
/// Backwards-compatible alias for callers using the original, inaccurate name.
|
||||
pub async fn create_codex_agent_identity_from_session_token(
|
||||
executor: &dyn OAuthHttpExecutor,
|
||||
session_token: &str,
|
||||
network: OAuthNetworkContext,
|
||||
) -> Result<Map<String, Value>, CodexAgentIdentityEnrollmentError> {
|
||||
create_codex_agent_identity_from_session_token_with_auth_api_base_url(
|
||||
executor,
|
||||
session_token,
|
||||
network,
|
||||
CODEX_AGENT_IDENTITY_AUTH_API_BASE_URL,
|
||||
)
|
||||
.await
|
||||
create_codex_agent_identity_from_access_token(executor, session_token, network).await
|
||||
}
|
||||
|
||||
async fn create_codex_agent_identity_from_session_token_with_auth_api_base_url(
|
||||
@@ -395,75 +644,13 @@ async fn create_codex_agent_identity_from_session_token_with_auth_api_base_url(
|
||||
network: OAuthNetworkContext,
|
||||
auth_api_base_url: &str,
|
||||
) -> Result<Map<String, Value>, CodexAgentIdentityEnrollmentError> {
|
||||
let session_token = session_token.trim();
|
||||
if session_token.is_empty() {
|
||||
return Err(CodexAgentIdentityEnrollmentError::MissingSessionToken);
|
||||
}
|
||||
|
||||
let signing_key = generate_agent_identity_signing_key();
|
||||
let private_key_der = signing_key
|
||||
.to_pkcs8_der()
|
||||
.map_err(|_| CodexAgentIdentityEnrollmentError::KeyGenerationFailed)?;
|
||||
let agent_private_key = STANDARD.encode(private_key_der.as_bytes());
|
||||
let agent_public_key = agent_identity_ssh_public_key(&signing_key);
|
||||
let registration_url = agent_registration_url(auth_api_base_url)
|
||||
.map_err(|_| CodexAgentIdentityEnrollmentError::RegistrationRequestFailed)?;
|
||||
let registration_response = executor
|
||||
.execute(OAuthHttpRequest {
|
||||
request_id: CODEX_AGENT_IDENTITY_AGENT_REGISTRATION_REQUEST_ID.to_string(),
|
||||
method: reqwest::Method::POST,
|
||||
url: registration_url,
|
||||
headers: BTreeMap::from([
|
||||
("accept".to_string(), "application/json".to_string()),
|
||||
("content-type".to_string(), "application/json".to_string()),
|
||||
(
|
||||
"authorization".to_string(),
|
||||
format!("Bearer {session_token}"),
|
||||
),
|
||||
(
|
||||
"user-agent".to_string(),
|
||||
aether_ai_formats::CODEX_CLIENT_USER_AGENT.to_string(),
|
||||
),
|
||||
(
|
||||
"originator".to_string(),
|
||||
aether_ai_formats::CODEX_CLIENT_ORIGINATOR.to_string(),
|
||||
),
|
||||
]),
|
||||
content_type: Some("application/json".to_string()),
|
||||
json_body: Some(json!({
|
||||
"abom": {
|
||||
"agent_version": aether_ai_formats::CODEX_CLIENT_VERSION,
|
||||
"agent_harness_id": CODEX_AGENT_IDENTITY_AGENT_HARNESS_ID,
|
||||
"running_location": CODEX_AGENT_IDENTITY_RUNNING_LOCATION,
|
||||
},
|
||||
"agent_public_key": agent_public_key,
|
||||
})),
|
||||
body_bytes: None,
|
||||
network: network.clone(),
|
||||
})
|
||||
.await
|
||||
.map_err(|_| CodexAgentIdentityEnrollmentError::RegistrationRequestFailed)?;
|
||||
if !(200..300).contains(®istration_response.status_code) {
|
||||
return Err(CodexAgentIdentityEnrollmentError::RegistrationRejected {
|
||||
status_code: registration_response.status_code,
|
||||
});
|
||||
}
|
||||
let agent_runtime_id =
|
||||
agent_runtime_id_from_registration_response(registration_response.body_text.as_str())
|
||||
.map_err(|_| CodexAgentIdentityEnrollmentError::InvalidRegistrationResponse)?;
|
||||
|
||||
let mut auth_config = Map::from_iter([
|
||||
(
|
||||
"provider_type".to_string(),
|
||||
json!(CODEX_AGENT_IDENTITY_PROVIDER_TYPE),
|
||||
),
|
||||
(
|
||||
"auth_mode".to_string(),
|
||||
json!(CODEX_AGENT_IDENTITY_AUTH_MODE),
|
||||
),
|
||||
("agent_runtime_id".to_string(), json!(agent_runtime_id)),
|
||||
("agent_private_key".to_string(), json!(agent_private_key)),
|
||||
]);
|
||||
let mut auth_config = register_codex_agent_identity_from_access_token_with_auth_api_base_url(
|
||||
executor,
|
||||
session_token,
|
||||
network.clone(),
|
||||
auth_api_base_url,
|
||||
)
|
||||
.await?;
|
||||
let config_value = Value::Object(auth_config.clone());
|
||||
let credentials = agent_identity_credentials(&config_value)
|
||||
.map_err(|_| CodexAgentIdentityEnrollmentError::KeyGenerationFailed)?;
|
||||
@@ -505,6 +692,84 @@ async fn create_codex_agent_identity_from_session_token_with_auth_api_base_url(
|
||||
Ok(auth_config)
|
||||
}
|
||||
|
||||
async fn register_codex_agent_identity_from_access_token_with_auth_api_base_url(
|
||||
executor: &dyn OAuthHttpExecutor,
|
||||
access_token: &str,
|
||||
network: OAuthNetworkContext,
|
||||
auth_api_base_url: &str,
|
||||
) -> Result<Map<String, Value>, CodexAgentIdentityEnrollmentError> {
|
||||
let access_token = access_token.trim();
|
||||
if access_token.is_empty() {
|
||||
return Err(CodexAgentIdentityEnrollmentError::MissingSessionToken);
|
||||
}
|
||||
|
||||
let signing_key = generate_agent_identity_signing_key();
|
||||
let private_key_der = signing_key
|
||||
.to_pkcs8_der()
|
||||
.map_err(|_| CodexAgentIdentityEnrollmentError::KeyGenerationFailed)?;
|
||||
let agent_private_key = STANDARD.encode(private_key_der.as_bytes());
|
||||
let agent_public_key = agent_identity_ssh_public_key(&signing_key);
|
||||
let registration_url = agent_registration_url(auth_api_base_url)
|
||||
.map_err(|_| CodexAgentIdentityEnrollmentError::RegistrationRequestFailed)?;
|
||||
let response = executor
|
||||
.execute(OAuthHttpRequest {
|
||||
request_id: CODEX_AGENT_IDENTITY_AGENT_REGISTRATION_REQUEST_ID.to_string(),
|
||||
method: reqwest::Method::POST,
|
||||
url: registration_url,
|
||||
headers: BTreeMap::from([
|
||||
("accept".to_string(), "application/json".to_string()),
|
||||
("content-type".to_string(), "application/json".to_string()),
|
||||
(
|
||||
"authorization".to_string(),
|
||||
format!("Bearer {access_token}"),
|
||||
),
|
||||
(
|
||||
"user-agent".to_string(),
|
||||
aether_ai_formats::CODEX_CLIENT_USER_AGENT.to_string(),
|
||||
),
|
||||
(
|
||||
"originator".to_string(),
|
||||
aether_ai_formats::CODEX_CLIENT_ORIGINATOR.to_string(),
|
||||
),
|
||||
]),
|
||||
content_type: Some("application/json".to_string()),
|
||||
json_body: Some(json!({
|
||||
"abom": {
|
||||
"agent_version": aether_ai_formats::CODEX_CLIENT_VERSION,
|
||||
"agent_harness_id": CODEX_AGENT_IDENTITY_AGENT_HARNESS_ID,
|
||||
"running_location": CODEX_AGENT_IDENTITY_RUNNING_LOCATION,
|
||||
},
|
||||
"agent_public_key": agent_public_key,
|
||||
})),
|
||||
body_bytes: None,
|
||||
network,
|
||||
})
|
||||
.await
|
||||
.map_err(|_| CodexAgentIdentityEnrollmentError::RegistrationRequestFailed)?;
|
||||
if !(200..300).contains(&response.status_code) {
|
||||
return Err(CodexAgentIdentityEnrollmentError::RegistrationRejected {
|
||||
status_code: response.status_code,
|
||||
});
|
||||
}
|
||||
let agent_runtime_id = agent_runtime_id_from_registration_response(response.body_text.as_str())
|
||||
.map_err(|_| CodexAgentIdentityEnrollmentError::InvalidRegistrationResponse)?;
|
||||
let auth_config = Map::from_iter([
|
||||
(
|
||||
"provider_type".to_string(),
|
||||
json!(CODEX_AGENT_IDENTITY_PROVIDER_TYPE),
|
||||
),
|
||||
(
|
||||
"auth_mode".to_string(),
|
||||
json!(CODEX_AGENT_IDENTITY_AUTH_MODE),
|
||||
),
|
||||
("agent_runtime_id".to_string(), json!(agent_runtime_id)),
|
||||
("agent_private_key".to_string(), json!(agent_private_key)),
|
||||
]);
|
||||
validate_codex_agent_identity_auth_config(&Value::Object(auth_config.clone()))
|
||||
.map_err(|_| CodexAgentIdentityEnrollmentError::KeyGenerationFailed)?;
|
||||
Ok(auth_config)
|
||||
}
|
||||
|
||||
fn task_id_from_registration_response(
|
||||
credentials: &AgentIdentityCredentials,
|
||||
body: &str,
|
||||
@@ -592,9 +857,40 @@ impl LocalOAuthRefreshAdapter for CodexAgentIdentityRefreshAdapter {
|
||||
|
||||
fn resolve_cached(
|
||||
&self,
|
||||
_transport: &GatewayProviderTransportSnapshot,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: &CachedOAuthEntry,
|
||||
) -> Option<LocalResolvedOAuthRequestAuth> {
|
||||
let current_fingerprint = codex_agent_identity_transport_credential_fingerprint(transport)?;
|
||||
if entry.source_fingerprint.as_deref() != Some(current_fingerprint.as_str()) {
|
||||
return None;
|
||||
}
|
||||
let transport_config = Self::config_from_transport(transport)?;
|
||||
let entry_config = Self::config_from_entry(entry)?;
|
||||
let transport_task = agent_identity_credentials(&transport_config).ok()?.task_id;
|
||||
let entry_task = agent_identity_credentials(&entry_config).ok()?.task_id;
|
||||
if transport_task != entry_task {
|
||||
return None;
|
||||
}
|
||||
Self::resolve_from_config(&entry_config)
|
||||
}
|
||||
|
||||
fn resolve_fenced_cached(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: &CachedOAuthEntry,
|
||||
) -> Option<LocalResolvedOAuthRequestAuth> {
|
||||
self.resolve_cached(transport, entry)
|
||||
}
|
||||
|
||||
fn resolve_refreshed(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: &CachedOAuthEntry,
|
||||
) -> Option<LocalResolvedOAuthRequestAuth> {
|
||||
let current_fingerprint = codex_agent_identity_transport_credential_fingerprint(transport)?;
|
||||
if entry.source_fingerprint.as_deref() != Some(current_fingerprint.as_str()) {
|
||||
return None;
|
||||
}
|
||||
Self::config_from_entry(entry).and_then(|config| Self::resolve_from_config(&config))
|
||||
}
|
||||
|
||||
@@ -618,6 +914,36 @@ impl LocalOAuthRefreshAdapter for CodexAgentIdentityRefreshAdapter {
|
||||
.is_some_and(|credentials| credentials.task_id.is_none())
|
||||
}
|
||||
|
||||
fn refresh_fingerprint(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: Option<&CachedOAuthEntry>,
|
||||
) -> Option<String> {
|
||||
codex_agent_identity_refresh_fingerprint(transport, entry)
|
||||
}
|
||||
|
||||
fn cached_entry_from_transport(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<CachedOAuthEntry> {
|
||||
codex_agent_identity_cached_entry_from_transport(transport)
|
||||
}
|
||||
|
||||
fn should_backoff_after_error(&self, error: &LocalOAuthRefreshError) -> bool {
|
||||
match error {
|
||||
LocalOAuthRefreshError::HttpStatus { status_code, .. } => {
|
||||
*status_code == 429 || *status_code >= 500
|
||||
}
|
||||
LocalOAuthRefreshError::Transport { .. }
|
||||
| LocalOAuthRefreshError::TransportMessage { .. }
|
||||
| LocalOAuthRefreshError::InvalidResponse { .. } => true,
|
||||
}
|
||||
}
|
||||
|
||||
fn requires_distributed_refresh_lock(&self) -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
async fn refresh(
|
||||
&self,
|
||||
executor: &dyn LocalOAuthHttpExecutor,
|
||||
@@ -676,6 +1002,9 @@ impl LocalOAuthRefreshAdapter for CodexAgentIdentityRefreshAdapter {
|
||||
auth_header_value,
|
||||
expires_at_unix_secs: None,
|
||||
metadata: Some(config),
|
||||
source_fingerprint: Some(agent_identity_credential_fingerprint_from_credentials(
|
||||
&updated_credentials,
|
||||
)),
|
||||
}))
|
||||
}
|
||||
}
|
||||
@@ -693,16 +1022,22 @@ mod tests {
|
||||
|
||||
use super::{
|
||||
agent_identity_credentials, build_agent_assertion,
|
||||
codex_agent_identity_authorization_matches_transport,
|
||||
codex_agent_identity_cached_entry_from_transport,
|
||||
codex_agent_identity_config_refresh_fingerprint, codex_agent_identity_refresh_fingerprint,
|
||||
codex_agent_identity_transport_allows_task_rotation_from,
|
||||
create_codex_agent_identity_from_session_token_with_auth_api_base_url,
|
||||
decrypt_agent_task_id, is_codex_agent_identity_auth_config_value,
|
||||
is_codex_agent_identity_invalid_task_response, task_id_from_registration_response,
|
||||
validate_codex_agent_identity_auth_config, with_agent_identity_task_id,
|
||||
CodexAgentIdentityEnrollmentError, CodexAgentIdentityRefreshAdapter,
|
||||
CODEX_AGENT_IDENTITY_CACHED_ENTRY_PROVIDER_TYPE,
|
||||
is_codex_agent_identity_authorization, is_codex_agent_identity_invalid_task_response,
|
||||
register_codex_agent_identity_from_access_token_with_auth_api_base_url,
|
||||
task_id_from_registration_response, validate_codex_agent_identity_auth_config,
|
||||
with_agent_identity_task_id, CodexAgentIdentityEnrollmentError,
|
||||
CodexAgentIdentityRefreshAdapter, CODEX_AGENT_IDENTITY_CACHED_ENTRY_PROVIDER_TYPE,
|
||||
};
|
||||
use crate::oauth_refresh::{
|
||||
LocalOAuthHttpExecutor, LocalOAuthHttpRequest, LocalOAuthHttpResponse,
|
||||
LocalOAuthRefreshAdapter, LocalOAuthRefreshError, LocalResolvedOAuthRequestAuth,
|
||||
LocalOAuthRefreshAdapter, LocalOAuthRefreshCoordinator, LocalOAuthRefreshError,
|
||||
LocalResolvedOAuthRequestAuth,
|
||||
};
|
||||
use crate::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
@@ -826,6 +1161,79 @@ mod tests {
|
||||
.expect("assertion signature should verify");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn recognizes_agent_assertion_authorization_scheme_without_parsing_secrets() {
|
||||
assert!(is_codex_agent_identity_authorization(
|
||||
"AgentAssertion assertion-envelope"
|
||||
));
|
||||
assert!(is_codex_agent_identity_authorization(
|
||||
"agentassertion\tassertion-envelope"
|
||||
));
|
||||
assert!(!is_codex_agent_identity_authorization(
|
||||
"Bearer access-token"
|
||||
));
|
||||
assert!(!is_codex_agent_identity_authorization("AgentAssertion"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn assertion_match_binds_runtime_task_and_signing_key_generation() {
|
||||
let config = test_auth_config(Some("task-test"));
|
||||
let credentials = agent_identity_credentials(&config).expect("credentials should parse");
|
||||
let assertion = build_agent_assertion(
|
||||
&credentials,
|
||||
"task-test",
|
||||
Utc.with_ymd_and_hms(2030, 1, 2, 3, 4, 5).unwrap(),
|
||||
)
|
||||
.expect("assertion should build");
|
||||
|
||||
assert!(codex_agent_identity_authorization_matches_transport(
|
||||
&sample_transport(config.clone()),
|
||||
&assertion,
|
||||
));
|
||||
assert!(!codex_agent_identity_authorization_matches_transport(
|
||||
&sample_transport(test_auth_config(Some("task-replaced"))),
|
||||
&assertion,
|
||||
));
|
||||
|
||||
let replacement_key = SigningKey::from_bytes(&[8u8; 32])
|
||||
.to_pkcs8_der()
|
||||
.expect("replacement key should encode");
|
||||
let mut replacement_config = config;
|
||||
replacement_config["agent_private_key"] =
|
||||
json!(STANDARD.encode(replacement_key.as_bytes()));
|
||||
assert!(!codex_agent_identity_authorization_matches_transport(
|
||||
&sample_transport(replacement_config),
|
||||
&assertion,
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn task_rotation_context_rejects_metadata_and_credential_replacement() {
|
||||
let initial = sample_transport(test_auth_config(Some("task-old")));
|
||||
let rotated = sample_transport(test_auth_config(Some("task-new")));
|
||||
assert!(codex_agent_identity_transport_allows_task_rotation_from(
|
||||
&initial, &rotated
|
||||
));
|
||||
|
||||
let mut metadata_rewrite = test_auth_config(Some("task-new"));
|
||||
metadata_rewrite["account_id"] = json!("account-replaced");
|
||||
assert!(!codex_agent_identity_transport_allows_task_rotation_from(
|
||||
&initial,
|
||||
&sample_transport(metadata_rewrite),
|
||||
));
|
||||
|
||||
let replacement_key = SigningKey::from_bytes(&[9u8; 32])
|
||||
.to_pkcs8_der()
|
||||
.expect("replacement key should encode");
|
||||
let mut credential_rewrite = test_auth_config(Some("task-new"));
|
||||
credential_rewrite["agent_private_key"] =
|
||||
json!(STANDARD.encode(replacement_key.as_bytes()));
|
||||
assert!(!codex_agent_identity_transport_allows_task_rotation_from(
|
||||
&initial,
|
||||
&sample_transport(credential_rewrite),
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn decrypts_sealed_task_registration_response() {
|
||||
let config = test_auth_config(None);
|
||||
@@ -886,6 +1294,16 @@ mod tests {
|
||||
assert!(updated["agent_identity"].get("taskId").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn refresh_fingerprint_fences_task_generation_not_only_keypair() {
|
||||
let pending = test_auth_config(None);
|
||||
let registered = test_auth_config(Some("task-winner"));
|
||||
assert_ne!(
|
||||
codex_agent_identity_config_refresh_fingerprint(&pending),
|
||||
codex_agent_identity_config_refresh_fingerprint(®istered)
|
||||
);
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct RecordingExecutor {
|
||||
requests: Arc<Mutex<Vec<LocalOAuthHttpRequest>>>,
|
||||
@@ -1009,6 +1427,37 @@ mod tests {
|
||||
assert!(!requests[1].headers.contains_key("authorization"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn register_only_returns_pending_config_without_registering_task() {
|
||||
let requests = Arc::new(Mutex::new(Vec::new()));
|
||||
let executor = RecordingEnrollmentExecutor {
|
||||
requests: Arc::clone(&requests),
|
||||
responses: Arc::new(Mutex::new(vec![OAuthHttpResponse {
|
||||
status_code: 200,
|
||||
body_text: r#"{"agent_runtime_id":"runtime-pending"}"#.to_string(),
|
||||
json_body: None,
|
||||
}])),
|
||||
};
|
||||
let config = register_codex_agent_identity_from_access_token_with_auth_api_base_url(
|
||||
&executor,
|
||||
"access-token-for-test-only",
|
||||
OAuthNetworkContext::direct_identity(),
|
||||
"https://auth.test/api/accounts",
|
||||
)
|
||||
.await
|
||||
.expect("register-only enrollment should succeed");
|
||||
validate_codex_agent_identity_auth_config(&serde_json::Value::Object(config.clone()))
|
||||
.expect("pending config should be valid");
|
||||
assert_eq!(
|
||||
config.get("agent_runtime_id"),
|
||||
Some(&json!("runtime-pending"))
|
||||
);
|
||||
assert!(!config.contains_key("task_id"));
|
||||
let requests = requests.lock().expect("recording lock should hold");
|
||||
assert_eq!(requests.len(), 1);
|
||||
assert!(requests[0].url.ends_with("/v1/agent/register"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn enrollment_error_does_not_echo_session_token_or_response_body() {
|
||||
let executor = RecordingEnrollmentExecutor {
|
||||
@@ -1066,9 +1515,10 @@ mod tests {
|
||||
.and_then(|value| value.get("task_id")),
|
||||
Some(&json!("task-registered"))
|
||||
);
|
||||
assert!(adapter.resolve_cached(&transport, &entry).is_none());
|
||||
let cached_auth = adapter
|
||||
.resolve_cached(&transport, &entry)
|
||||
.expect("cached task should create a new assertion");
|
||||
.resolve_refreshed(&transport, &entry)
|
||||
.expect("new refresh result should create an assertion");
|
||||
assert!(matches!(
|
||||
cached_auth,
|
||||
LocalResolvedOAuthRequestAuth::Header { ref name, ref value }
|
||||
@@ -1093,6 +1543,120 @@ mod tests {
|
||||
.is_some());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn rejects_cached_task_after_agent_credential_rotation() {
|
||||
let transport = sample_transport(test_auth_config(None));
|
||||
let requests = Arc::new(Mutex::new(Vec::new()));
|
||||
let adapter = CodexAgentIdentityRefreshAdapter::default()
|
||||
.with_auth_api_base_url_for_tests("https://auth.test/api/accounts");
|
||||
let entry = adapter
|
||||
.refresh(
|
||||
&RecordingExecutor {
|
||||
requests: Arc::clone(&requests),
|
||||
},
|
||||
&transport,
|
||||
None,
|
||||
)
|
||||
.await
|
||||
.expect("registration should succeed")
|
||||
.expect("registration should return an entry");
|
||||
let mut rotated_config = test_auth_config(Some("task-rotated"));
|
||||
rotated_config["agent_runtime_id"] = json!("runtime-rotated");
|
||||
let rotated_transport = sample_transport(rotated_config);
|
||||
|
||||
assert!(adapter.resolve_cached(&rotated_transport, &entry).is_none());
|
||||
assert!(adapter
|
||||
.resolve_without_refresh(&rotated_transport)
|
||||
.is_some());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn rejects_cached_task_after_remote_task_rotation() {
|
||||
let transport = sample_transport(test_auth_config(None));
|
||||
let adapter = CodexAgentIdentityRefreshAdapter::default()
|
||||
.with_auth_api_base_url_for_tests("https://auth.test/api/accounts");
|
||||
let entry = adapter
|
||||
.refresh(
|
||||
&RecordingExecutor {
|
||||
requests: Arc::new(Mutex::new(Vec::new())),
|
||||
},
|
||||
&transport,
|
||||
None,
|
||||
)
|
||||
.await
|
||||
.expect("registration should succeed")
|
||||
.expect("registration should return an entry");
|
||||
let rotated_transport = sample_transport(test_auth_config(Some("task-new-winner")));
|
||||
|
||||
assert!(adapter.resolve_cached(&rotated_transport, &entry).is_none());
|
||||
assert!(adapter
|
||||
.resolve_without_refresh(&rotated_transport)
|
||||
.is_some());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn newer_transport_task_is_authoritative_over_stale_local_cache() {
|
||||
let stale_transport = sample_transport(test_auth_config(Some("task-stale")));
|
||||
let stale_entry = codex_agent_identity_cached_entry_from_transport(&stale_transport)
|
||||
.expect("stale transport should produce a cache entry");
|
||||
let current_config = test_auth_config(Some("task-current"));
|
||||
let current_transport = sample_transport(current_config.clone());
|
||||
|
||||
assert_eq!(
|
||||
codex_agent_identity_refresh_fingerprint(¤t_transport, Some(&stale_entry)),
|
||||
codex_agent_identity_config_refresh_fingerprint(¤t_config)
|
||||
);
|
||||
assert!(CodexAgentIdentityRefreshAdapter::default()
|
||||
.resolve_fenced_cached(¤t_transport, &stale_entry)
|
||||
.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn distributed_waiter_reuses_reloaded_transport_not_stale_cache() {
|
||||
let stale_transport = sample_transport(test_auth_config(Some("task-stale")));
|
||||
let stale_entry = codex_agent_identity_cached_entry_from_transport(&stale_transport)
|
||||
.expect("stale transport should produce a cache entry");
|
||||
let expected = codex_agent_identity_refresh_fingerprint(&stale_transport, None)
|
||||
.expect("stale transport should have a generation");
|
||||
let current_transport = sample_transport(test_auth_config(Some("task-current")));
|
||||
let requests = Arc::new(Mutex::new(Vec::new()));
|
||||
let coordinator = LocalOAuthRefreshCoordinator::with_adapters_for_tests(vec![Arc::new(
|
||||
CodexAgentIdentityRefreshAdapter::default()
|
||||
.with_auth_api_base_url_for_tests("https://auth.test/api/accounts"),
|
||||
)]);
|
||||
coordinator
|
||||
.store_cached_entry(current_transport.key.id.as_str(), stale_entry)
|
||||
.await;
|
||||
|
||||
let resolution = coordinator
|
||||
.force_refresh_with_result_fenced(
|
||||
&RecordingExecutor {
|
||||
requests: Arc::clone(&requests),
|
||||
},
|
||||
¤t_transport,
|
||||
None,
|
||||
None,
|
||||
Some(expected.as_str()),
|
||||
)
|
||||
.await
|
||||
.expect("waiter resolution should succeed")
|
||||
.expect("waiter should reuse the DB winner");
|
||||
|
||||
assert!(resolution.reused_refresh);
|
||||
assert_eq!(
|
||||
resolution
|
||||
.refreshed_entry
|
||||
.as_ref()
|
||||
.and_then(|entry| entry.metadata.as_ref())
|
||||
.and_then(|config| config.get("task_id")),
|
||||
Some(&json!("task-current"))
|
||||
);
|
||||
assert!(requests
|
||||
.lock()
|
||||
.expect("recording lock should hold")
|
||||
.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn registration_response_accepts_encrypted_task_aliases() {
|
||||
let config = test_auth_config(Some("task-original"));
|
||||
|
||||
@@ -6,6 +6,7 @@ use aether_oauth::provider::providers::{
|
||||
use aether_oauth::provider::{ProviderOAuthAccount, ProviderOAuthAdapter, ProviderOAuthTokenSet};
|
||||
use async_trait::async_trait;
|
||||
use serde_json::Value;
|
||||
use sha2::{Digest, Sha256};
|
||||
|
||||
use super::oauth_refresh::{
|
||||
oauth_error_to_local_refresh_error, provider_oauth_transport_context_from_snapshot,
|
||||
@@ -73,6 +74,7 @@ impl GenericOAuthRefreshAdapter {
|
||||
entry
|
||||
.provider_type
|
||||
.eq_ignore_ascii_case(transport.provider.provider_type.as_str())
|
||||
&& generic_oauth_cached_entry_matches_transport(transport, entry)
|
||||
})
|
||||
.cloned()
|
||||
}
|
||||
@@ -147,6 +149,7 @@ impl GenericOAuthRefreshAdapter {
|
||||
|
||||
fn build_cached_entry(
|
||||
provider_type: &'static str,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
refreshed: ProviderOAuthTokenSet,
|
||||
) -> CachedOAuthEntry {
|
||||
CachedOAuthEntry {
|
||||
@@ -155,6 +158,7 @@ impl GenericOAuthRefreshAdapter {
|
||||
auth_header_value: refreshed.token_set.bearer_header_value(),
|
||||
expires_at_unix_secs: refreshed.token_set.expires_at_unix_secs,
|
||||
metadata: Some(refreshed.auth_config),
|
||||
source_fingerprint: Some(generic_oauth_transport_source_fingerprint(transport)),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -188,6 +192,9 @@ impl LocalOAuthRefreshAdapter for GenericOAuthRefreshAdapter {
|
||||
value,
|
||||
});
|
||||
}
|
||||
if !generic_oauth_cached_entry_matches_transport(transport, entry) {
|
||||
return None;
|
||||
}
|
||||
if expires_at_requires_refresh(entry.expires_at_unix_secs) {
|
||||
return None;
|
||||
}
|
||||
@@ -293,10 +300,46 @@ impl LocalOAuthRefreshAdapter for GenericOAuthRefreshAdapter {
|
||||
"gateway generic oauth refresh succeeded"
|
||||
);
|
||||
|
||||
Ok(Some(Self::build_cached_entry(provider_type, refreshed)))
|
||||
Ok(Some(Self::build_cached_entry(
|
||||
provider_type,
|
||||
transport,
|
||||
refreshed,
|
||||
)))
|
||||
}
|
||||
}
|
||||
|
||||
fn generic_oauth_transport_source_fingerprint(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> String {
|
||||
let provider_type = transport.provider.provider_type.trim().to_ascii_lowercase();
|
||||
let auth_type = transport.key.auth_type.trim().to_ascii_lowercase();
|
||||
let auth_config = transport
|
||||
.key
|
||||
.decrypted_auth_config
|
||||
.as_deref()
|
||||
.unwrap_or_default();
|
||||
let api_key = transport.key.decrypted_api_key.as_str();
|
||||
let mut digest = Sha256::new();
|
||||
for field in [
|
||||
provider_type.as_bytes(),
|
||||
auth_type.as_bytes(),
|
||||
auth_config.as_bytes(),
|
||||
api_key.as_bytes(),
|
||||
] {
|
||||
digest.update((field.len() as u64).to_be_bytes());
|
||||
digest.update(field);
|
||||
}
|
||||
format!("{:x}", digest.finalize())
|
||||
}
|
||||
|
||||
fn generic_oauth_cached_entry_matches_transport(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: &CachedOAuthEntry,
|
||||
) -> bool {
|
||||
let transport_fingerprint = generic_oauth_transport_source_fingerprint(transport);
|
||||
entry.source_fingerprint.as_deref() == Some(transport_fingerprint.as_str())
|
||||
}
|
||||
|
||||
fn generic_provider_type(provider_type: &str) -> Option<&'static str> {
|
||||
let normalized = provider_type.trim();
|
||||
GENERIC_PROVIDER_OAUTH_TEMPLATES
|
||||
@@ -362,6 +405,7 @@ fn current_access_token(
|
||||
entry: Option<&CachedOAuthEntry>,
|
||||
) -> Option<String> {
|
||||
entry
|
||||
.filter(|entry| generic_oauth_cached_entry_matches_transport(transport, entry))
|
||||
.and_then(|entry| {
|
||||
entry
|
||||
.auth_header_value
|
||||
@@ -388,7 +432,10 @@ mod tests {
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||
};
|
||||
use super::GenericOAuthRefreshAdapter;
|
||||
use super::{
|
||||
current_access_token, generic_oauth_transport_source_fingerprint,
|
||||
GenericOAuthRefreshAdapter,
|
||||
};
|
||||
|
||||
fn sample_transport() -> GatewayProviderTransportSnapshot {
|
||||
GatewayProviderTransportSnapshot {
|
||||
@@ -481,6 +528,7 @@ mod tests {
|
||||
auth_header_value: "Bearer refreshed-access-token".to_string(),
|
||||
expires_at_unix_secs: Some(u64::MAX),
|
||||
metadata: None,
|
||||
source_fingerprint: None,
|
||||
};
|
||||
let auth = adapter
|
||||
.resolve_cached(&sample_transport(), &entry)
|
||||
@@ -494,4 +542,98 @@ mod tests {
|
||||
}
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_cached_bearer_and_metadata_from_replaced_credential_generation() {
|
||||
let adapter = GenericOAuthRefreshAdapter::default();
|
||||
let mut original = sample_transport();
|
||||
original.key.decrypted_api_key = "access-a".to_string();
|
||||
original.key.decrypted_auth_config = Some(
|
||||
json!({
|
||||
"provider_type": "codex",
|
||||
"refresh_token": "refresh-a",
|
||||
"expires_at": 1,
|
||||
"updated_at": 100,
|
||||
})
|
||||
.to_string(),
|
||||
);
|
||||
let entry = CachedOAuthEntry {
|
||||
provider_type: "codex".to_string(),
|
||||
auth_header_name: "authorization".to_string(),
|
||||
auth_header_value: "Bearer cached-access-a".to_string(),
|
||||
expires_at_unix_secs: Some(u64::MAX),
|
||||
metadata: Some(json!({
|
||||
"provider_type": "codex",
|
||||
"refresh_token": "rotated-refresh-a",
|
||||
"expires_at": u64::MAX,
|
||||
"updated_at": 200,
|
||||
})),
|
||||
source_fingerprint: Some(generic_oauth_transport_source_fingerprint(&original)),
|
||||
};
|
||||
|
||||
let mut replacement = original.clone();
|
||||
replacement.key.decrypted_api_key = "access-b".to_string();
|
||||
replacement.key.decrypted_auth_config = Some(
|
||||
json!({
|
||||
"provider_type": "codex",
|
||||
"refresh_token": "refresh-b",
|
||||
"expires_at": 1,
|
||||
"updated_at": 300,
|
||||
})
|
||||
.to_string(),
|
||||
);
|
||||
|
||||
assert!(adapter.resolve_cached(&replacement, &entry).is_none());
|
||||
assert_eq!(
|
||||
adapter.base_auth_config(&replacement, Some(&entry)),
|
||||
replacement
|
||||
.key
|
||||
.decrypted_auth_config
|
||||
.as_deref()
|
||||
.and_then(|value| serde_json::from_str(value).ok())
|
||||
);
|
||||
assert_eq!(
|
||||
current_access_token(&replacement, Some(&entry)).as_deref(),
|
||||
Some("access-b")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn reuses_cached_bearer_from_matching_credential_generation() {
|
||||
let adapter = GenericOAuthRefreshAdapter::default();
|
||||
let mut transport = sample_transport();
|
||||
transport.key.decrypted_api_key = "access-a".to_string();
|
||||
transport.key.decrypted_auth_config = Some(
|
||||
json!({
|
||||
"provider_type": "codex",
|
||||
"refresh_token": "refresh-a",
|
||||
"expires_at": 1,
|
||||
})
|
||||
.to_string(),
|
||||
);
|
||||
let entry = CachedOAuthEntry {
|
||||
provider_type: "codex".to_string(),
|
||||
auth_header_name: "authorization".to_string(),
|
||||
auth_header_value: "Bearer refreshed-access-a".to_string(),
|
||||
expires_at_unix_secs: Some(u64::MAX),
|
||||
metadata: Some(json!({
|
||||
"provider_type": "codex",
|
||||
"refresh_token": "rotated-refresh-a",
|
||||
"expires_at": u64::MAX,
|
||||
})),
|
||||
source_fingerprint: Some(generic_oauth_transport_source_fingerprint(&transport)),
|
||||
};
|
||||
|
||||
assert_eq!(
|
||||
adapter.resolve_cached(&transport, &entry),
|
||||
Some(LocalResolvedOAuthRequestAuth::Header {
|
||||
name: "authorization".to_string(),
|
||||
value: "Bearer refreshed-access-a".to_string(),
|
||||
})
|
||||
);
|
||||
assert_eq!(
|
||||
current_access_token(&transport, Some(&entry)).as_deref(),
|
||||
Some("refreshed-access-a")
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -77,6 +77,7 @@ impl KiroOAuthRefreshAdapter {
|
||||
auth_header_value: request_auth.value,
|
||||
expires_at_unix_secs: auth_config.expires_at,
|
||||
metadata: Some(auth_config.to_json_value()),
|
||||
source_fingerprint: None,
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
@@ -30,13 +30,21 @@ pub mod windsurf;
|
||||
|
||||
pub use aether_oauth as oauth;
|
||||
pub use agent_identity::{
|
||||
create_codex_agent_identity_from_session_token, is_codex_agent_identity_auth_config_value,
|
||||
codex_agent_identity_auth_config_has_task_id,
|
||||
codex_agent_identity_authorization_matches_transport,
|
||||
codex_agent_identity_cached_entry_from_transport,
|
||||
codex_agent_identity_config_refresh_fingerprint, codex_agent_identity_credential_fingerprint,
|
||||
codex_agent_identity_entry_allows_task_rotation_from, codex_agent_identity_refresh_fingerprint,
|
||||
codex_agent_identity_transport_allows_task_rotation_from,
|
||||
codex_agent_identity_transport_credential_fingerprint,
|
||||
create_codex_agent_identity_from_access_token, create_codex_agent_identity_from_session_token,
|
||||
is_codex_agent_identity_auth_config_value, is_codex_agent_identity_authorization,
|
||||
is_codex_agent_identity_cached_entry, is_codex_agent_identity_invalid_task_response,
|
||||
is_codex_agent_identity_transport, validate_codex_agent_identity_auth_config,
|
||||
CodexAgentIdentityEnrollmentError, CodexAgentIdentityRefreshAdapter,
|
||||
CODEX_AGENT_IDENTITY_AGENT_REGISTRATION_REQUEST_ID, CODEX_AGENT_IDENTITY_AUTH_MODE,
|
||||
CODEX_AGENT_IDENTITY_CACHED_ENTRY_PROVIDER_TYPE, CODEX_AGENT_IDENTITY_PROVIDER_TYPE,
|
||||
CODEX_AGENT_IDENTITY_TASK_REGISTRATION_REQUEST_ID,
|
||||
is_codex_agent_identity_transport, register_codex_agent_identity_from_access_token,
|
||||
validate_codex_agent_identity_auth_config, CodexAgentIdentityEnrollmentError,
|
||||
CodexAgentIdentityRefreshAdapter, CODEX_AGENT_IDENTITY_AGENT_REGISTRATION_REQUEST_ID,
|
||||
CODEX_AGENT_IDENTITY_AUTH_MODE, CODEX_AGENT_IDENTITY_CACHED_ENTRY_PROVIDER_TYPE,
|
||||
CODEX_AGENT_IDENTITY_PROVIDER_TYPE, CODEX_AGENT_IDENTITY_TASK_REGISTRATION_REQUEST_ID,
|
||||
};
|
||||
pub use auth::{build_passthrough_headers, ensure_upstream_auth_header};
|
||||
pub use auth_config::apply_local_auth_config_header_overrides;
|
||||
@@ -91,7 +99,8 @@ pub use network::{
|
||||
pub use oauth_refresh::{
|
||||
supports_local_oauth_request_auth_resolution, CachedOAuthEntry, LocalOAuthHttpExecutor,
|
||||
LocalOAuthHttpRequest, LocalOAuthHttpResponse, LocalOAuthRefreshCoordinator,
|
||||
LocalOAuthRefreshError, LocalResolvedOAuthRequestAuth, ReqwestLocalOAuthHttpExecutor,
|
||||
LocalOAuthRefreshError, LocalOAuthResolution, LocalResolvedOAuthRequestAuth,
|
||||
ReqwestLocalOAuthHttpExecutor,
|
||||
};
|
||||
pub use openai_image::{
|
||||
build_openai_image_headers, build_openai_image_upstream_url,
|
||||
|
||||
@@ -1,13 +1,14 @@
|
||||
use std::collections::BTreeMap;
|
||||
use std::fmt;
|
||||
use std::sync::Arc;
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
use aether_oauth::core::OAuthError;
|
||||
use aether_oauth::network::{
|
||||
OAuthHttpExecutor, OAuthHttpRequest, OAuthHttpResponse, OAuthNetworkContext,
|
||||
};
|
||||
use aether_oauth::provider::ProviderOAuthTransportContext;
|
||||
use aether_runtime_state::RuntimeState;
|
||||
use aether_runtime_state::{RuntimeLockLease, RuntimeState};
|
||||
use async_trait::async_trait;
|
||||
use serde_json::Value;
|
||||
use thiserror::Error;
|
||||
@@ -40,6 +41,12 @@ pub struct LocalOAuthResolution {
|
||||
pub auth: Option<LocalResolvedOAuthRequestAuth>,
|
||||
pub refreshed_entry: Option<CachedOAuthEntry>,
|
||||
pub refresh_in_flight: bool,
|
||||
/// Indicates that a forced caller reused a newer completed refresh rather
|
||||
/// than producing a new entry that needs persistence.
|
||||
pub reused_refresh: bool,
|
||||
/// Held until the caller persists `refreshed_entry`. The lease TTL remains
|
||||
/// the cancellation fallback if the caller is dropped.
|
||||
pub distributed_lease: Option<RuntimeLockLease>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
@@ -49,6 +56,8 @@ pub struct CachedOAuthEntry {
|
||||
pub auth_header_value: String,
|
||||
pub expires_at_unix_secs: Option<u64>,
|
||||
pub metadata: Option<Value>,
|
||||
/// Non-secret fingerprint of the credential/configuration that produced it.
|
||||
pub source_fingerprint: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
@@ -297,6 +306,28 @@ pub trait LocalOAuthRefreshAdapter: Send + Sync {
|
||||
entry: &CachedOAuthEntry,
|
||||
) -> Option<LocalResolvedOAuthRequestAuth>;
|
||||
|
||||
/// Resolves a cache entry that is known to have advanced the caller's
|
||||
/// refresh fence. Agent task rotation can safely use the winner even while
|
||||
/// the caller still holds the pre-refresh transport snapshot.
|
||||
fn resolve_fenced_cached(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: &CachedOAuthEntry,
|
||||
) -> Option<LocalResolvedOAuthRequestAuth> {
|
||||
self.resolve_cached(transport, entry)
|
||||
}
|
||||
|
||||
/// Resolves the entry returned by this adapter's immediately preceding
|
||||
/// refresh. Unlike a reusable cache entry, this entry is expected to have
|
||||
/// advanced the transport generation.
|
||||
fn resolve_refreshed(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: &CachedOAuthEntry,
|
||||
) -> Option<LocalResolvedOAuthRequestAuth> {
|
||||
self.resolve_cached(transport, entry)
|
||||
}
|
||||
|
||||
fn resolve_without_refresh(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
@@ -308,6 +339,36 @@ pub trait LocalOAuthRefreshAdapter: Send + Sync {
|
||||
entry: Option<&CachedOAuthEntry>,
|
||||
) -> bool;
|
||||
|
||||
/// Identifies the credential/configuration generation used by a refresh.
|
||||
/// Adapters that support fencing override this method.
|
||||
fn refresh_fingerprint(
|
||||
&self,
|
||||
_transport: &GatewayProviderTransportSnapshot,
|
||||
_entry: Option<&CachedOAuthEntry>,
|
||||
) -> Option<String> {
|
||||
None
|
||||
}
|
||||
|
||||
/// Reconstructs a cache entry from an already-persisted transport after a
|
||||
/// distributed refresh waiter reloads the winner.
|
||||
fn cached_entry_from_transport(
|
||||
&self,
|
||||
_transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<CachedOAuthEntry> {
|
||||
None
|
||||
}
|
||||
|
||||
/// Enables bounded negative backoff for transient refresh failures.
|
||||
fn should_backoff_after_error(&self, _error: &LocalOAuthRefreshError) -> bool {
|
||||
false
|
||||
}
|
||||
|
||||
/// Agent task registration is a non-idempotent external mutation and must
|
||||
/// not continue unlocked when a configured distributed lock is unavailable.
|
||||
fn requires_distributed_refresh_lock(&self) -> bool {
|
||||
false
|
||||
}
|
||||
|
||||
async fn refresh(
|
||||
&self,
|
||||
executor: &dyn LocalOAuthHttpExecutor,
|
||||
@@ -320,6 +381,13 @@ pub struct LocalOAuthRefreshCoordinator {
|
||||
adapters: Vec<Arc<dyn LocalOAuthRefreshAdapter>>,
|
||||
cache: Mutex<BTreeMap<String, CachedOAuthEntry>>,
|
||||
key_locks: Mutex<BTreeMap<String, Arc<Mutex<()>>>>,
|
||||
refresh_backoff: Mutex<BTreeMap<String, RefreshBackoffState>>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
struct RefreshBackoffState {
|
||||
failures: u32,
|
||||
retry_after: Instant,
|
||||
}
|
||||
|
||||
impl fmt::Debug for LocalOAuthRefreshCoordinator {
|
||||
@@ -337,7 +405,10 @@ impl Default for LocalOAuthRefreshCoordinator {
|
||||
}
|
||||
|
||||
impl LocalOAuthRefreshCoordinator {
|
||||
const DISTRIBUTED_REFRESH_LOCK_TTL_MS: u64 = 30_000;
|
||||
// Keep the lease alive through the 30s upstream HTTP timeout and the
|
||||
// subsequent encrypted DB CAS/persistence step. Cancellation still relies
|
||||
// on expiry as the last-resort release path.
|
||||
const DISTRIBUTED_REFRESH_LOCK_TTL_MS: u64 = 120_000;
|
||||
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
@@ -349,6 +420,7 @@ impl LocalOAuthRefreshCoordinator {
|
||||
],
|
||||
cache: Mutex::new(BTreeMap::new()),
|
||||
key_locks: Mutex::new(BTreeMap::new()),
|
||||
refresh_backoff: Mutex::new(BTreeMap::new()),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -373,7 +445,9 @@ impl LocalOAuthRefreshCoordinator {
|
||||
}
|
||||
|
||||
pub async fn invalidate_cached_entry(&self, key_id: &str) -> bool {
|
||||
self.cache.lock().await.remove(key_id).is_some()
|
||||
let removed = self.cache.lock().await.remove(key_id).is_some();
|
||||
self.clear_refresh_backoff(key_id).await;
|
||||
removed
|
||||
}
|
||||
|
||||
pub async fn resolve_with_result(
|
||||
@@ -389,6 +463,7 @@ impl LocalOAuthRefreshCoordinator {
|
||||
distributed_lock,
|
||||
distributed_owner,
|
||||
false,
|
||||
None,
|
||||
)
|
||||
.await
|
||||
}
|
||||
@@ -399,6 +474,27 @@ impl LocalOAuthRefreshCoordinator {
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
distributed_lock: Option<&RuntimeState>,
|
||||
distributed_owner: Option<&str>,
|
||||
) -> Result<Option<LocalOAuthResolution>, LocalOAuthRefreshError> {
|
||||
self.force_refresh_with_result_fenced(
|
||||
executor,
|
||||
transport,
|
||||
distributed_lock,
|
||||
distributed_owner,
|
||||
None,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
/// Force a refresh unless another request has already advanced the supplied
|
||||
/// refresh fence. This prevents a distributed waiter from re-registering a
|
||||
/// task after the winner has persisted it.
|
||||
pub async fn force_refresh_with_result_fenced(
|
||||
&self,
|
||||
executor: &dyn LocalOAuthHttpExecutor,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
distributed_lock: Option<&RuntimeState>,
|
||||
distributed_owner: Option<&str>,
|
||||
expected_refresh_fingerprint: Option<&str>,
|
||||
) -> Result<Option<LocalOAuthResolution>, LocalOAuthRefreshError> {
|
||||
self.resolve_with_result_mode(
|
||||
executor,
|
||||
@@ -406,6 +502,7 @@ impl LocalOAuthRefreshCoordinator {
|
||||
distributed_lock,
|
||||
distributed_owner,
|
||||
true,
|
||||
expected_refresh_fingerprint,
|
||||
)
|
||||
.await
|
||||
}
|
||||
@@ -417,6 +514,7 @@ impl LocalOAuthRefreshCoordinator {
|
||||
distributed_lock: Option<&RuntimeState>,
|
||||
distributed_owner: Option<&str>,
|
||||
force_refresh: bool,
|
||||
expected_refresh_fingerprint: Option<&str>,
|
||||
) -> Result<Option<LocalOAuthResolution>, LocalOAuthRefreshError> {
|
||||
let Some(adapter) = self
|
||||
.adapters
|
||||
@@ -450,10 +548,37 @@ impl LocalOAuthRefreshCoordinator {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
if force_refresh {
|
||||
if let Some(resolution) = Self::resolve_if_refresh_fence_advanced(
|
||||
adapter.as_ref(),
|
||||
transport,
|
||||
cached_entry.as_ref(),
|
||||
expected_refresh_fingerprint,
|
||||
) {
|
||||
return Ok(Some(resolution));
|
||||
}
|
||||
}
|
||||
if let Some(error) = self.backoff_error(key_id, adapter.provider_type()).await {
|
||||
return Err(error);
|
||||
}
|
||||
|
||||
let key_lock = self.lock_for_key(key_id).await;
|
||||
let _key_guard = key_lock.lock().await;
|
||||
|
||||
let cached_entry = self.cached_entry(key_id).await;
|
||||
if force_refresh {
|
||||
if let Some(resolution) = Self::resolve_if_refresh_fence_advanced(
|
||||
adapter.as_ref(),
|
||||
transport,
|
||||
cached_entry.as_ref(),
|
||||
expected_refresh_fingerprint,
|
||||
) {
|
||||
return Ok(Some(resolution));
|
||||
}
|
||||
}
|
||||
if let Some(error) = self.backoff_error(key_id, adapter.provider_type()).await {
|
||||
return Err(error);
|
||||
}
|
||||
if !force_refresh {
|
||||
if let Some(auth) = cached_entry
|
||||
.as_ref()
|
||||
@@ -488,6 +613,16 @@ impl LocalOAuthRefreshCoordinator {
|
||||
error = ?err,
|
||||
"gateway local oauth refresh distributed lock unavailable"
|
||||
);
|
||||
if adapter.requires_distributed_refresh_lock() {
|
||||
let error = LocalOAuthRefreshError::TransportMessage {
|
||||
provider_type: adapter.provider_type(),
|
||||
message: "distributed refresh lock is unavailable".to_string(),
|
||||
};
|
||||
if adapter.should_backoff_after_error(&error) {
|
||||
self.record_refresh_failure(key_id).await;
|
||||
}
|
||||
return Err(error);
|
||||
}
|
||||
None
|
||||
}
|
||||
}
|
||||
@@ -501,22 +636,147 @@ impl LocalOAuthRefreshCoordinator {
|
||||
// came from the original transport snapshot.
|
||||
let refresh_entry = cached_entry.as_ref();
|
||||
let refresh_result = adapter.refresh(executor, transport, refresh_entry).await;
|
||||
if let (Some(lock), Some(lease)) = (distributed_lock, distributed_lease.as_ref()) {
|
||||
if let Err(err) = lock.lock_release(lease).await {
|
||||
tracing::warn!(
|
||||
key_id = %key_id,
|
||||
provider_type = adapter.provider_type(),
|
||||
error = ?err,
|
||||
"gateway local oauth refresh distributed lock release failed"
|
||||
);
|
||||
let refreshed_entry = match refresh_result {
|
||||
Ok(Some(entry)) => {
|
||||
self.clear_refresh_backoff(key_id).await;
|
||||
entry
|
||||
}
|
||||
Ok(None) => {
|
||||
Self::release_distributed_lease(
|
||||
distributed_lock,
|
||||
distributed_lease.as_ref(),
|
||||
key_id,
|
||||
adapter.provider_type(),
|
||||
)
|
||||
.await;
|
||||
return Ok(None);
|
||||
}
|
||||
Err(error) => {
|
||||
if adapter.should_backoff_after_error(&error) {
|
||||
self.record_refresh_failure(key_id).await;
|
||||
}
|
||||
Self::release_distributed_lease(
|
||||
distributed_lock,
|
||||
distributed_lease.as_ref(),
|
||||
key_id,
|
||||
adapter.provider_type(),
|
||||
)
|
||||
.await;
|
||||
return Err(error);
|
||||
}
|
||||
};
|
||||
// In production the distributed lease is held through the gateway's
|
||||
// DB CAS. Do not publish a provisional task before that CAS succeeds;
|
||||
// otherwise a waiter could consume an assertion that loses the CAS.
|
||||
// Lock-free/test callers retain the historical in-memory behavior.
|
||||
if distributed_lease.is_none() {
|
||||
self.insert_cached_entry(key_id, refreshed_entry.clone())
|
||||
.await;
|
||||
}
|
||||
let Some(refreshed_entry) = refresh_result? else {
|
||||
let Some(auth) = adapter.resolve_refreshed(transport, &refreshed_entry) else {
|
||||
Self::release_distributed_lease(
|
||||
distributed_lock,
|
||||
distributed_lease.as_ref(),
|
||||
key_id,
|
||||
adapter.provider_type(),
|
||||
)
|
||||
.await;
|
||||
return Ok(None);
|
||||
};
|
||||
Ok(adapter
|
||||
.resolve_cached(transport, &refreshed_entry)
|
||||
.map(|auth| LocalOAuthResolution::resolved(auth, Some(refreshed_entry))))
|
||||
Ok(Some(LocalOAuthResolution::refreshed(
|
||||
auth,
|
||||
refreshed_entry,
|
||||
distributed_lease,
|
||||
)))
|
||||
}
|
||||
|
||||
fn resolve_if_refresh_fence_advanced(
|
||||
adapter: &dyn LocalOAuthRefreshAdapter,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: Option<&CachedOAuthEntry>,
|
||||
expected_refresh_fingerprint: Option<&str>,
|
||||
) -> Option<LocalOAuthResolution> {
|
||||
let expected = expected_refresh_fingerprint?;
|
||||
if adapter.refresh_fingerprint(transport, entry).as_deref() == Some(expected) {
|
||||
return None;
|
||||
}
|
||||
entry
|
||||
.and_then(|entry| adapter.resolve_fenced_cached(transport, entry))
|
||||
.map(|auth| {
|
||||
LocalOAuthResolution::reused(
|
||||
auth,
|
||||
entry.expect("cached auth is required when a refresh fence advanced"),
|
||||
)
|
||||
})
|
||||
.or_else(|| {
|
||||
adapter
|
||||
.cached_entry_from_transport(transport)
|
||||
.and_then(|entry| {
|
||||
adapter
|
||||
.resolve_cached(transport, &entry)
|
||||
.map(|auth| LocalOAuthResolution::reused(auth, &entry))
|
||||
})
|
||||
})
|
||||
.or_else(|| {
|
||||
adapter
|
||||
.resolve_without_refresh(transport)
|
||||
.map(|auth| LocalOAuthResolution::resolved(auth, None))
|
||||
})
|
||||
}
|
||||
|
||||
async fn backoff_error(
|
||||
&self,
|
||||
key_id: &str,
|
||||
provider_type: &'static str,
|
||||
) -> Option<LocalOAuthRefreshError> {
|
||||
let backoff = self.refresh_backoff.lock().await;
|
||||
let state = backoff.get(key_id)?;
|
||||
let remaining = state.retry_after.checked_duration_since(Instant::now())?;
|
||||
Some(LocalOAuthRefreshError::InvalidResponse {
|
||||
provider_type,
|
||||
message: format!(
|
||||
"refresh temporarily backed off after {} failed attempts (retry in {}ms)",
|
||||
state.failures,
|
||||
remaining.as_millis()
|
||||
),
|
||||
})
|
||||
}
|
||||
|
||||
async fn record_refresh_failure(&self, key_id: &str) {
|
||||
let mut backoff = self.refresh_backoff.lock().await;
|
||||
let state = backoff
|
||||
.entry(key_id.to_string())
|
||||
.or_insert(RefreshBackoffState {
|
||||
failures: 0,
|
||||
retry_after: Instant::now(),
|
||||
});
|
||||
state.failures = state.failures.saturating_add(1);
|
||||
let exponent = state.failures.saturating_sub(1).min(4);
|
||||
let delay = Duration::from_millis(500u64.saturating_mul(1u64 << exponent));
|
||||
state.retry_after = Instant::now() + delay.min(Duration::from_secs(8));
|
||||
}
|
||||
|
||||
async fn clear_refresh_backoff(&self, key_id: &str) {
|
||||
self.refresh_backoff.lock().await.remove(key_id);
|
||||
}
|
||||
|
||||
async fn release_distributed_lease(
|
||||
distributed_lock: Option<&RuntimeState>,
|
||||
lease: Option<&RuntimeLockLease>,
|
||||
key_id: &str,
|
||||
provider_type: &'static str,
|
||||
) {
|
||||
let (Some(lock), Some(lease)) = (distributed_lock, lease) else {
|
||||
return;
|
||||
};
|
||||
if let Err(err) = lock.lock_release(lease).await {
|
||||
tracing::warn!(
|
||||
key_id = %key_id,
|
||||
provider_type,
|
||||
error = ?err,
|
||||
"gateway local oauth refresh distributed lock release failed"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_adapters_for_tests(adapters: Vec<Arc<dyn LocalOAuthRefreshAdapter>>) -> Self {
|
||||
@@ -524,6 +784,7 @@ impl LocalOAuthRefreshCoordinator {
|
||||
adapters,
|
||||
cache: Mutex::new(BTreeMap::new()),
|
||||
key_locks: Mutex::new(BTreeMap::new()),
|
||||
refresh_backoff: Mutex::new(BTreeMap::new()),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -537,6 +798,37 @@ impl LocalOAuthResolution {
|
||||
auth: Some(auth),
|
||||
refreshed_entry,
|
||||
refresh_in_flight: false,
|
||||
reused_refresh: false,
|
||||
distributed_lease: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn refreshed(
|
||||
auth: LocalResolvedOAuthRequestAuth,
|
||||
refreshed_entry: CachedOAuthEntry,
|
||||
distributed_lease: Option<RuntimeLockLease>,
|
||||
) -> Self {
|
||||
Self {
|
||||
auth: Some(auth),
|
||||
refreshed_entry: Some(refreshed_entry),
|
||||
refresh_in_flight: false,
|
||||
reused_refresh: false,
|
||||
distributed_lease,
|
||||
}
|
||||
}
|
||||
|
||||
fn reused(auth: LocalResolvedOAuthRequestAuth, entry: &CachedOAuthEntry) -> Self {
|
||||
let mut refreshed_entry = entry.clone();
|
||||
if let LocalResolvedOAuthRequestAuth::Header { name, value } = &auth {
|
||||
refreshed_entry.auth_header_name = name.clone();
|
||||
refreshed_entry.auth_header_value = value.clone();
|
||||
}
|
||||
Self {
|
||||
auth: Some(auth),
|
||||
refreshed_entry: Some(refreshed_entry),
|
||||
refresh_in_flight: false,
|
||||
reused_refresh: true,
|
||||
distributed_lease: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -545,6 +837,8 @@ impl LocalOAuthResolution {
|
||||
auth: None,
|
||||
refreshed_entry: None,
|
||||
refresh_in_flight: true,
|
||||
reused_refresh: false,
|
||||
distributed_lease: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -560,7 +854,7 @@ pub fn supports_local_oauth_request_auth_resolution(
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
|
||||
|
||||
use super::super::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
@@ -580,6 +874,12 @@ mod tests {
|
||||
refresh_with_entry_hits: Arc<AtomicUsize>,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
struct FencedTestAdapter {
|
||||
refresh_hits: Arc<AtomicUsize>,
|
||||
fail_refresh: Arc<AtomicBool>,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LocalOAuthRefreshAdapter for TestAdapter {
|
||||
fn provider_type(&self) -> &'static str {
|
||||
@@ -634,6 +934,77 @@ mod tests {
|
||||
auth_header_value: "Bearer refreshed-token".to_string(),
|
||||
expires_at_unix_secs: Some(4_102_444_800),
|
||||
metadata: None,
|
||||
source_fingerprint: None,
|
||||
}))
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LocalOAuthRefreshAdapter for FencedTestAdapter {
|
||||
fn provider_type(&self) -> &'static str {
|
||||
"test-oauth"
|
||||
}
|
||||
|
||||
fn resolve_cached(
|
||||
&self,
|
||||
_transport: &GatewayProviderTransportSnapshot,
|
||||
entry: &CachedOAuthEntry,
|
||||
) -> Option<LocalResolvedOAuthRequestAuth> {
|
||||
Some(LocalResolvedOAuthRequestAuth::Header {
|
||||
name: entry.auth_header_name.clone(),
|
||||
value: "fresh-winner-assertion".to_string(),
|
||||
})
|
||||
}
|
||||
|
||||
fn resolve_without_refresh(
|
||||
&self,
|
||||
_transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<LocalResolvedOAuthRequestAuth> {
|
||||
None
|
||||
}
|
||||
|
||||
fn should_refresh(
|
||||
&self,
|
||||
_transport: &GatewayProviderTransportSnapshot,
|
||||
_entry: Option<&CachedOAuthEntry>,
|
||||
) -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
fn refresh_fingerprint(
|
||||
&self,
|
||||
_transport: &GatewayProviderTransportSnapshot,
|
||||
entry: Option<&CachedOAuthEntry>,
|
||||
) -> Option<String> {
|
||||
entry
|
||||
.and_then(|entry| entry.source_fingerprint.clone())
|
||||
.or_else(|| Some("generation-1".to_string()))
|
||||
}
|
||||
|
||||
fn should_backoff_after_error(&self, _error: &LocalOAuthRefreshError) -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
async fn refresh(
|
||||
&self,
|
||||
_executor: &dyn LocalOAuthHttpExecutor,
|
||||
_transport: &GatewayProviderTransportSnapshot,
|
||||
_entry: Option<&CachedOAuthEntry>,
|
||||
) -> Result<Option<CachedOAuthEntry>, LocalOAuthRefreshError> {
|
||||
self.refresh_hits.fetch_add(1, Ordering::SeqCst);
|
||||
if self.fail_refresh.load(Ordering::SeqCst) {
|
||||
return Err(LocalOAuthRefreshError::TransportMessage {
|
||||
provider_type: "test-oauth",
|
||||
message: "temporary failure".to_string(),
|
||||
});
|
||||
}
|
||||
Ok(Some(CachedOAuthEntry {
|
||||
provider_type: "test-oauth".to_string(),
|
||||
auth_header_name: "authorization".to_string(),
|
||||
auth_header_value: "stale-winner-cache-value".to_string(),
|
||||
expires_at_unix_secs: None,
|
||||
metadata: None,
|
||||
source_fingerprint: Some("generation-2".to_string()),
|
||||
}))
|
||||
}
|
||||
}
|
||||
@@ -739,8 +1110,11 @@ mod tests {
|
||||
auth_header_value: "Bearer refreshed-token".to_string(),
|
||||
expires_at_unix_secs: Some(4_102_444_800),
|
||||
metadata: None,
|
||||
source_fingerprint: None,
|
||||
}),
|
||||
refresh_in_flight: false,
|
||||
reused_refresh: false,
|
||||
distributed_lease: None,
|
||||
})
|
||||
);
|
||||
assert_eq!(
|
||||
@@ -752,6 +1126,8 @@ mod tests {
|
||||
}),
|
||||
refreshed_entry: None,
|
||||
refresh_in_flight: false,
|
||||
reused_refresh: false,
|
||||
distributed_lease: None,
|
||||
})
|
||||
);
|
||||
}
|
||||
@@ -791,4 +1167,86 @@ mod tests {
|
||||
assert_eq!(refresh_hits.load(Ordering::SeqCst), 2);
|
||||
assert_eq!(refresh_with_entry_hits.load(Ordering::SeqCst), 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn fenced_force_refresh_reuses_the_winner() {
|
||||
let refresh_hits = Arc::new(AtomicUsize::new(0));
|
||||
let coordinator = LocalOAuthRefreshCoordinator::with_adapters_for_tests(vec![Arc::new(
|
||||
FencedTestAdapter {
|
||||
refresh_hits: Arc::clone(&refresh_hits),
|
||||
fail_refresh: Arc::new(AtomicBool::new(false)),
|
||||
},
|
||||
)]);
|
||||
let transport = sample_transport();
|
||||
let executor = ReqwestLocalOAuthHttpExecutor::new(reqwest::Client::new());
|
||||
|
||||
let first = coordinator
|
||||
.force_refresh_with_result_fenced(
|
||||
&executor,
|
||||
&transport,
|
||||
None,
|
||||
None,
|
||||
Some("generation-1"),
|
||||
)
|
||||
.await
|
||||
.expect("first refresh should succeed")
|
||||
.expect("first refresh should resolve");
|
||||
let waiter = coordinator
|
||||
.force_refresh_with_result_fenced(
|
||||
&executor,
|
||||
&transport,
|
||||
None,
|
||||
None,
|
||||
Some("generation-1"),
|
||||
)
|
||||
.await
|
||||
.expect("waiter should reuse winner")
|
||||
.expect("waiter should resolve");
|
||||
|
||||
assert!(first.refreshed_entry.is_some());
|
||||
assert!(waiter.refreshed_entry.is_some());
|
||||
assert!(waiter.reused_refresh);
|
||||
assert_eq!(
|
||||
waiter
|
||||
.refreshed_entry
|
||||
.as_ref()
|
||||
.expect("reused entry")
|
||||
.auth_header_value,
|
||||
"fresh-winner-assertion"
|
||||
);
|
||||
assert_eq!(refresh_hits.load(Ordering::SeqCst), 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn refresh_failure_enters_bounded_negative_backoff() {
|
||||
let refresh_hits = Arc::new(AtomicUsize::new(0));
|
||||
let fail_refresh = Arc::new(AtomicBool::new(true));
|
||||
let coordinator = LocalOAuthRefreshCoordinator::with_adapters_for_tests(vec![Arc::new(
|
||||
FencedTestAdapter {
|
||||
refresh_hits: Arc::clone(&refresh_hits),
|
||||
fail_refresh: Arc::clone(&fail_refresh),
|
||||
},
|
||||
)]);
|
||||
let transport = sample_transport();
|
||||
let executor = ReqwestLocalOAuthHttpExecutor::new(reqwest::Client::new());
|
||||
|
||||
assert!(coordinator
|
||||
.force_refresh_with_result(&executor, &transport, None, None)
|
||||
.await
|
||||
.is_err());
|
||||
let second = coordinator
|
||||
.force_refresh_with_result(&executor, &transport, None, None)
|
||||
.await
|
||||
.expect_err("second refresh should be backed off");
|
||||
assert!(second.to_string().contains("temporarily backed off"));
|
||||
assert_eq!(refresh_hits.load(Ordering::SeqCst), 1);
|
||||
fail_refresh.store(false, Ordering::SeqCst);
|
||||
coordinator.invalidate_cached_entry("key-1").await;
|
||||
assert!(coordinator
|
||||
.force_refresh_with_result(&executor, &transport, None, None)
|
||||
.await
|
||||
.expect("replacement should refresh immediately")
|
||||
.is_some());
|
||||
assert_eq!(refresh_hits.load(Ordering::SeqCst), 2);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -259,6 +259,7 @@ impl LocalOAuthRefreshAdapter for VertexServiceAccountRefreshAdapter {
|
||||
"project_id": auth_config.project_id,
|
||||
"client_email": auth_config.client_email,
|
||||
})),
|
||||
source_fingerprint: None,
|
||||
}))
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user