mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-08 02:17:46 +08:00
fix(provider): harden Agent Identity OAuth lifecycle
This commit is contained in:
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user