fix(provider): harden Agent Identity OAuth lifecycle

This commit is contained in:
elky
2026-07-23 09:33:00 +08:00
parent e49024d33b
commit 3606290ac8
84 changed files with 8063 additions and 1104 deletions
@@ -8,8 +8,8 @@ use serde_json::{json, Map, Value};
use super::{
ProviderCatalogKeyAdaptiveState, ProviderCatalogKeyAdaptiveStateUpdate,
ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyListQuery,
ProviderCatalogKeyRuntimeMetadataUpdate, ProviderCatalogKeyStatusSnapshotUpdate,
ProviderCatalogReadRepository, ProviderCatalogSnapshot,
ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ProviderCatalogKeyRuntimeMetadataUpdate,
ProviderCatalogKeyStatusSnapshotUpdate, ProviderCatalogReadRepository, ProviderCatalogSnapshot,
ProviderCatalogUpstreamMetadataNamespaceUpdate, ProviderCatalogWriteRepository,
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
StoredProviderCatalogKeyMaintenanceSummary, StoredProviderCatalogKeyPage,
@@ -770,6 +770,82 @@ impl ProviderCatalogWriteRepository for InMemoryProviderCatalogReadRepository {
Ok(true)
}
async fn compare_and_update_key_oauth_runtime_state(
&self,
update: &ProviderCatalogKeyOAuthRuntimeStateCasUpdate,
) -> Result<bool, DataLayerError> {
if update.encrypted_auth_config.trim().is_empty()
|| update
.encrypted_api_key_update
.as_deref()
.is_some_and(|value| value.trim().is_empty())
|| !update.status_snapshot_patch.is_object()
|| update
.upstream_metadata_patch
.as_ref()
.is_some_and(|patch| !patch.is_object())
{
return Err(DataLayerError::InvalidInput(
"provider catalog OAuth runtime CAS requires auth_config and object status patch"
.to_string(),
));
}
let patch = update
.status_snapshot_patch
.as_object()
.cloned()
.expect("status patch object was validated");
let mut index = self
.index
.write()
.expect("provider catalog repository lock");
let Some(key) = index.keys.get_mut(&update.key_id) else {
return Ok(false);
};
if key.encrypted_auth_config.as_deref() != update.expected_encrypted_auth_config.as_deref()
{
return Ok(false);
}
if let Some(encrypted_api_key) = update.encrypted_api_key_update.as_ref() {
key.encrypted_api_key = Some(encrypted_api_key.clone());
}
key.encrypted_auth_config = Some(update.encrypted_auth_config.clone());
if let Some(expires_at_unix_secs) = update.expires_at_unix_secs_update {
key.expires_at_unix_secs = expires_at_unix_secs;
}
key.oauth_invalid_at_unix_secs = update.oauth_invalid_at_unix_secs;
key.oauth_invalid_reason = update.oauth_invalid_reason.clone();
if update.reset_error_count {
key.error_count = Some(0);
}
if let Some(metadata_patch) = update
.upstream_metadata_patch
.as_ref()
.and_then(Value::as_object)
.cloned()
{
let upstream_metadata = json_object_for_merge(
key.upstream_metadata.as_ref(),
"provider catalog upstream metadata",
)?;
key.upstream_metadata = Some(Value::Object(merge_json_objects(
upstream_metadata,
metadata_patch,
)));
}
let status_snapshot = json_object_for_merge(
key.status_snapshot.as_ref(),
"provider catalog status snapshot",
)?;
key.status_snapshot = Some(Value::Object(merge_json_objects(status_snapshot, patch)));
key.updated_at_unix_secs = Some(
update
.updated_at_unix_secs
.unwrap_or_else(current_unix_secs),
);
Ok(true)
}
async fn update_key_health_state(
&self,
key_id: &str,
@@ -825,7 +901,12 @@ impl ProviderCatalogWriteRepository for InMemoryProviderCatalogReadRepository {
};
let expected = update.expected.canonicalized();
let next = update.next.canonicalized();
if ProviderCatalogKeyAdaptiveState::from(&*key) != expected {
if update
.expected_encrypted_auth_config
.as_deref()
.is_some_and(|expected| key.encrypted_auth_config.as_deref() != Some(expected))
|| ProviderCatalogKeyAdaptiveState::from(&*key) != expected
{
return Ok(false);
}
let status_snapshot = json_object_for_merge(
@@ -964,7 +1045,11 @@ impl ProviderCatalogWriteRepository for InMemoryProviderCatalogReadRepository {
let Some(key) = index.keys.get_mut(&update.key_id) else {
return Ok(false);
};
if key.health_by_format != update.expected_health_by_format
if update
.expected_encrypted_auth_config
.as_deref()
.is_some_and(|expected| key.encrypted_auth_config.as_deref() != Some(expected))
|| key.health_by_format != update.expected_health_by_format
|| key.circuit_breaker_by_format != update.expected_circuit_breaker_by_format
{
return Ok(false);
@@ -1085,9 +1170,10 @@ mod tests {
use crate::repository::provider_catalog::{
ProviderCatalogKeyAdaptiveState, ProviderCatalogKeyAdaptiveStateUpdate,
ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyListOrder,
ProviderCatalogKeyListQuery, ProviderCatalogKeyRuntimeMetadataUpdate,
ProviderCatalogReadRepository, ProviderCatalogWriteRepository,
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
ProviderCatalogKeyListQuery, ProviderCatalogKeyOAuthRuntimeStateCasUpdate,
ProviderCatalogKeyRuntimeMetadataUpdate, ProviderCatalogReadRepository,
ProviderCatalogWriteRepository, StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
StoredProviderCatalogProvider,
};
use crate::repository::usage::ProviderApiKeyUsageDelta;
use serde_json::{json, Value};
@@ -1236,6 +1322,75 @@ mod tests {
assert_eq!(stored[0].expires_at_unix_secs, Some(4_102_444_800));
}
#[tokio::test]
async fn oauth_runtime_cas_preserves_admin_fields_and_rejects_stale_config() {
let mut key = sample_key("key-1", "provider-1")
.with_transport_fields(
None,
"ciphertext-placeholder".to_string(),
Some("ciphertext-auth-1".to_string()),
None,
None,
None,
None,
None,
None,
)
.expect("key transport should build");
key.note = Some("admin-owned-note".to_string());
key.status_snapshot = Some(json!({"quota":{"remaining":7}}));
let repository = InMemoryProviderCatalogReadRepository::seed(
vec![sample_provider("provider-1")],
vec![],
vec![key],
);
let update = ProviderCatalogKeyOAuthRuntimeStateCasUpdate {
key_id: "key-1".to_string(),
expected_encrypted_auth_config: Some("ciphertext-auth-1".to_string()),
encrypted_auth_config: "ciphertext-auth-2".to_string(),
encrypted_api_key_update: Some("ciphertext-api-2".to_string()),
expires_at_unix_secs_update: Some(Some(456)),
oauth_invalid_at_unix_secs: None,
oauth_invalid_reason: None,
upstream_metadata_patch: Some(json!({"codex":{"remaining":3}})),
status_snapshot_patch: json!({"oauth":{"code":"none"}}),
reset_error_count: false,
updated_at_unix_secs: Some(123),
};
assert!(repository
.compare_and_update_key_oauth_runtime_state(&update)
.await
.expect("CAS should succeed"));
assert!(!repository
.compare_and_update_key_oauth_runtime_state(&update)
.await
.expect("stale CAS should not fail"));
let stored = repository
.list_keys_by_ids(&["key-1".to_string()])
.await
.expect("key should read")
.pop()
.expect("key should exist");
assert_eq!(stored.note.as_deref(), Some("admin-owned-note"));
assert_eq!(
stored.encrypted_api_key.as_deref(),
Some("ciphertext-api-2")
);
assert_eq!(stored.expires_at_unix_secs, Some(456));
assert_eq!(
stored.status_snapshot.as_ref().unwrap()["quota"]["remaining"],
7
);
assert_eq!(
stored.status_snapshot.as_ref().unwrap()["oauth"]["code"],
"none"
);
assert_eq!(
stored.upstream_metadata.as_ref().unwrap()["codex"]["remaining"],
3
);
}
#[tokio::test]
async fn materializes_codex_window_usage_stats_delta_in_memory() {
let mut key = sample_key("key-1", "provider-1");
@@ -1586,6 +1741,7 @@ mod tests {
);
let update = ProviderCatalogKeyHealthStateUpdate {
key_id: "key-1".to_string(),
expected_encrypted_auth_config: None,
expected_health_by_format: Some(json!({"openai:chat":{"consecutive_failures":1}})),
expected_circuit_breaker_by_format: None,
health_by_format: Some(json!({"openai:chat":{"consecutive_failures":2}})),
@@ -1617,6 +1773,7 @@ mod tests {
#[tokio::test]
async fn adaptive_cas_detects_conflicts_and_merges_only_owned_status_fields() {
let mut key = sample_key("key-1", "provider-1");
key.encrypted_auth_config = Some("auth-current".to_string());
key.learned_rpm_limit = Some(10);
key.rpm_429_count = Some(1);
key.status_snapshot = Some(json!({
@@ -1635,6 +1792,7 @@ mod tests {
);
let update = ProviderCatalogKeyAdaptiveStateUpdate {
key_id: "key-1".to_string(),
expected_encrypted_auth_config: Some("auth-current".to_string()),
expected,
next,
status_snapshot_patch: json!({
@@ -1645,6 +1803,14 @@ mod tests {
updated_at_unix_secs: Some(10),
};
let stale_generation_update = ProviderCatalogKeyAdaptiveStateUpdate {
expected_encrypted_auth_config: Some("auth-stale".to_string()),
..update.clone()
};
assert!(!repository
.compare_and_update_key_adaptive_state(&stale_generation_update)
.await
.expect("stale auth generation should conflict"));
assert!(repository
.compare_and_update_key_adaptive_state(&update)
.await
@@ -1887,6 +2053,7 @@ mod tests {
repository
.compare_and_update_key_health_state(&ProviderCatalogKeyHealthStateUpdate {
key_id: key.id.clone(),
expected_encrypted_auth_config: None,
expected_health_by_format: key.health_by_format.clone(),
expected_circuit_breaker_by_format: None,
health_by_format: Some(json!({"openai:chat":{"consecutive_failures":2}})),
@@ -1901,6 +2068,7 @@ mod tests {
repository
.compare_and_update_key_adaptive_state(&ProviderCatalogKeyAdaptiveStateUpdate {
key_id: key.id.clone(),
expected_encrypted_auth_config: None,
expected,
next,
status_snapshot_patch: json!({"observation_count":2}),
@@ -4,8 +4,8 @@ mod memory;
pub(crate) use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogKeyAdaptiveState, ProviderCatalogKeyAdaptiveStateUpdate,
ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery,
ProviderCatalogKeyRuntimeMetadataUpdate, ProviderCatalogKeyStatusSnapshotUpdate,
ProviderCatalogReadRepository, ProviderCatalogSnapshot,
ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ProviderCatalogKeyRuntimeMetadataUpdate,
ProviderCatalogKeyStatusSnapshotUpdate, ProviderCatalogReadRepository, ProviderCatalogSnapshot,
ProviderCatalogUpstreamMetadataNamespaceUpdate, ProviderCatalogWriteRepository,
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
StoredProviderCatalogKeyMaintenanceSummary, StoredProviderCatalogKeyPage,
@@ -40,6 +40,8 @@ pub struct StoredAdminProviderOAuthState {
pub provider_id: String,
pub provider_type: String,
pub pkce_verifier: Option<String>,
#[serde(default)]
pub expected_encrypted_auth_config: Option<String>,
}
pub fn provider_oauth_device_session_storage_key(session_id: &str) -> String {