mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-09 18:59:50 +08:00
fix(provider): harden Agent Identity OAuth lifecycle
This commit is contained in:
@@ -222,6 +222,28 @@ pub(super) fn classify_oauth_route(
|
||||
"admin:provider_oauth",
|
||||
false,
|
||||
))
|
||||
} else if method == http::Method::POST
|
||||
&& normalized_path.starts_with("/api/admin/provider-oauth/providers/")
|
||||
&& normalized_path.ends_with("/agent-identity-import/tasks")
|
||||
{
|
||||
Some(classified(
|
||||
"admin_proxy",
|
||||
"provider_oauth_manage",
|
||||
"start_agent_identity_import_task",
|
||||
"admin:provider_oauth",
|
||||
false,
|
||||
))
|
||||
} else if method == http::Method::GET
|
||||
&& normalized_path.starts_with("/api/admin/provider-oauth/providers/")
|
||||
&& normalized_path.contains("/agent-identity-import/tasks/")
|
||||
{
|
||||
Some(classified(
|
||||
"admin_proxy",
|
||||
"provider_oauth_manage",
|
||||
"get_agent_identity_import_task_status",
|
||||
"admin:provider_oauth",
|
||||
false,
|
||||
))
|
||||
} else if method == http::Method::POST
|
||||
&& normalized_path.starts_with("/api/admin/provider-oauth/providers/")
|
||||
&& normalized_path.ends_with("/batch-import")
|
||||
|
||||
@@ -109,6 +109,20 @@ fn classifies_admin_provider_oauth_maintenance_routes_as_admin_proxy_route() {
|
||||
"admin:provider_oauth",
|
||||
"admin:provider_oauth:write",
|
||||
),
|
||||
(
|
||||
http::Method::POST,
|
||||
"/api/admin/provider-oauth/providers/provider-123/agent-identity-import/tasks",
|
||||
"start_agent_identity_import_task",
|
||||
"admin:provider_oauth",
|
||||
"admin:provider_oauth:write",
|
||||
),
|
||||
(
|
||||
http::Method::GET,
|
||||
"/api/admin/provider-oauth/providers/provider-123/agent-identity-import/tasks/agent-identity-task-123",
|
||||
"get_agent_identity_import_task_status",
|
||||
"admin:provider_oauth",
|
||||
"admin:provider_oauth:read",
|
||||
),
|
||||
(
|
||||
http::Method::POST,
|
||||
"/api/admin/provider-oauth/providers/provider-123/batch-import",
|
||||
|
||||
@@ -2,12 +2,12 @@ use super::{
|
||||
ApiKeyLastUsedDelta, DataLayerError, GatewayDataState, GeminiFileMappingListQuery,
|
||||
GeminiFileMappingStats, ProviderCatalogKeyAdaptiveStateUpdate,
|
||||
ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyListQuery,
|
||||
ProviderCatalogKeyRuntimeMetadataUpdate, ProviderCatalogKeyStatusSnapshotUpdate,
|
||||
PublicHealthStatusCount, PublicHealthTimelineBucket, StoredGeminiFileMapping,
|
||||
StoredGeminiFileMappingListPage, StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
|
||||
StoredProviderCatalogKeyMaintenanceSummary, StoredProviderCatalogKeyPage,
|
||||
StoredProviderCatalogKeyStats, StoredProviderCatalogProvider, StoredRequestCandidate,
|
||||
UpsertGeminiFileMappingRecord, UpsertRequestCandidateRecord,
|
||||
ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ProviderCatalogKeyRuntimeMetadataUpdate,
|
||||
ProviderCatalogKeyStatusSnapshotUpdate, PublicHealthStatusCount, PublicHealthTimelineBucket,
|
||||
StoredGeminiFileMapping, StoredGeminiFileMappingListPage, StoredProviderCatalogEndpoint,
|
||||
StoredProviderCatalogKey, StoredProviderCatalogKeyMaintenanceSummary,
|
||||
StoredProviderCatalogKeyPage, StoredProviderCatalogKeyStats, StoredProviderCatalogProvider,
|
||||
StoredRequestCandidate, UpsertGeminiFileMappingRecord, UpsertRequestCandidateRecord,
|
||||
};
|
||||
|
||||
impl GatewayDataState {
|
||||
@@ -400,6 +400,24 @@ impl GatewayDataState {
|
||||
Ok(updated)
|
||||
}
|
||||
|
||||
pub(crate) async fn compare_and_update_provider_catalog_key_oauth_runtime_state(
|
||||
&self,
|
||||
update: &ProviderCatalogKeyOAuthRuntimeStateCasUpdate,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
let updated = match &self.provider_catalog_writer {
|
||||
Some(repository) => {
|
||||
repository
|
||||
.compare_and_update_key_oauth_runtime_state(update)
|
||||
.await
|
||||
}
|
||||
None => Ok(false),
|
||||
}?;
|
||||
// A false result is a credential CAS conflict. Clear cached snapshots
|
||||
// either way so the next read observes the authoritative row.
|
||||
self.clear_provider_catalog_cache();
|
||||
Ok(updated)
|
||||
}
|
||||
|
||||
pub(crate) async fn create_provider_catalog_key(
|
||||
&self,
|
||||
key: &StoredProviderCatalogKey,
|
||||
|
||||
@@ -122,11 +122,11 @@ use aether_data_contracts::repository::pool_scores::{
|
||||
};
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
ProviderCatalogKeyAdaptiveStateUpdate, ProviderCatalogKeyHealthStateUpdate,
|
||||
ProviderCatalogKeyListQuery, ProviderCatalogKeyRuntimeMetadataUpdate,
|
||||
ProviderCatalogKeyStatusSnapshotUpdate, ProviderCatalogReadRepository,
|
||||
ProviderCatalogWriteRepository, StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
|
||||
StoredProviderCatalogKeyMaintenanceSummary, StoredProviderCatalogKeyPage,
|
||||
StoredProviderCatalogKeyStats, StoredProviderCatalogProvider,
|
||||
ProviderCatalogKeyListQuery, ProviderCatalogKeyOAuthRuntimeStateCasUpdate,
|
||||
ProviderCatalogKeyRuntimeMetadataUpdate, ProviderCatalogKeyStatusSnapshotUpdate,
|
||||
ProviderCatalogReadRepository, ProviderCatalogWriteRepository, StoredProviderCatalogEndpoint,
|
||||
StoredProviderCatalogKey, StoredProviderCatalogKeyMaintenanceSummary,
|
||||
StoredProviderCatalogKeyPage, StoredProviderCatalogKeyStats, StoredProviderCatalogProvider,
|
||||
};
|
||||
use aether_data_contracts::repository::quota::{
|
||||
ProviderQuotaReadRepository, ProviderQuotaWriteRepository, StoredProviderQuotaSnapshot,
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
use aether_contracts::ExecutionPlan;
|
||||
use tracing::warn;
|
||||
|
||||
use crate::state::AgentIdentityAuthConfigFence;
|
||||
use crate::{provider_transport::LocalOAuthRefreshError, AppState};
|
||||
|
||||
pub(crate) async fn refresh_oauth_plan_auth_for_retry(
|
||||
@@ -13,8 +14,11 @@ pub(crate) async fn refresh_oauth_plan_auth_for_retry(
|
||||
if !status_may_be_oauth_invalid(status_code, response_text) {
|
||||
return false;
|
||||
}
|
||||
let access_token_invalid_proven =
|
||||
status_proves_access_token_invalid(status_code, response_text);
|
||||
let request_authorization = execution_plan_authorization(plan);
|
||||
let request_uses_agent_identity = request_authorization
|
||||
.is_some_and(aether_provider_transport::is_codex_agent_identity_authorization);
|
||||
let access_token_invalid_proven = !request_uses_agent_identity
|
||||
&& status_proves_access_token_invalid(status_code, response_text);
|
||||
|
||||
let transport = match state
|
||||
.read_provider_transport_snapshot(&plan.provider_id, &plan.endpoint_id, &plan.key_id)
|
||||
@@ -37,11 +41,44 @@ pub(crate) async fn refresh_oauth_plan_auth_for_retry(
|
||||
}
|
||||
};
|
||||
|
||||
if aether_provider_transport::is_codex_agent_identity_transport(&transport)
|
||||
&& !aether_provider_transport::is_codex_agent_identity_invalid_task_response(
|
||||
status_code,
|
||||
response_text,
|
||||
)
|
||||
let current_uses_agent_identity =
|
||||
aether_provider_transport::is_codex_agent_identity_transport(&transport);
|
||||
if request_uses_agent_identity {
|
||||
if !current_uses_agent_identity
|
||||
|| !aether_provider_transport::is_codex_agent_identity_invalid_task_response(
|
||||
status_code,
|
||||
response_text,
|
||||
)
|
||||
|| !request_authorization.is_some_and(|authorization| {
|
||||
aether_provider_transport::codex_agent_identity_authorization_matches_transport(
|
||||
&transport,
|
||||
authorization,
|
||||
)
|
||||
})
|
||||
{
|
||||
return false;
|
||||
}
|
||||
if !matches!(
|
||||
state
|
||||
.capture_agent_identity_auth_config_fence(&transport)
|
||||
.await,
|
||||
Ok(AgentIdentityAuthConfigFence::Current(_))
|
||||
) {
|
||||
return false;
|
||||
}
|
||||
} else if current_uses_agent_identity {
|
||||
// A bearer-token response cannot authorize refreshing an Agent Identity
|
||||
// installed under the same key id while the request was in flight.
|
||||
return false;
|
||||
} else if transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("codex")
|
||||
&& transport.key.auth_type.trim().eq_ignore_ascii_case("oauth")
|
||||
&& !request_authorization.is_some_and(|authorization| {
|
||||
bearer_authorization_matches_transport(authorization, &transport)
|
||||
})
|
||||
{
|
||||
return false;
|
||||
}
|
||||
@@ -120,6 +157,26 @@ pub(crate) async fn refresh_oauth_plan_auth_for_retry(
|
||||
}
|
||||
}
|
||||
|
||||
fn execution_plan_authorization(plan: &ExecutionPlan) -> Option<&str> {
|
||||
plan.headers
|
||||
.iter()
|
||||
.find(|(name, _)| name.eq_ignore_ascii_case("authorization"))
|
||||
.map(|(_, value)| value.as_str())
|
||||
}
|
||||
|
||||
fn bearer_authorization_matches_transport(
|
||||
authorization: &str,
|
||||
transport: &aether_provider_transport::GatewayProviderTransportSnapshot,
|
||||
) -> bool {
|
||||
let current_token = transport.key.decrypted_api_key.trim();
|
||||
!current_token.is_empty()
|
||||
&& authorization
|
||||
.trim()
|
||||
.strip_prefix("Bearer ")
|
||||
.map(str::trim)
|
||||
.is_some_and(|token| token == current_token)
|
||||
}
|
||||
|
||||
fn status_may_be_oauth_invalid(status_code: u16, response_text: Option<&str>) -> bool {
|
||||
if status_code == 401 {
|
||||
return true;
|
||||
|
||||
@@ -98,7 +98,9 @@ use crate::execution_runtime::{
|
||||
resolve_local_candidate_failover_analysis_stream, should_fallback_to_control_stream,
|
||||
should_retry_next_local_candidate_stream, LocalFailoverDecision,
|
||||
};
|
||||
use crate::execution_runtime::{MAX_STREAM_PREFETCH_BYTES, MAX_STREAM_PREFETCH_FRAMES};
|
||||
use crate::execution_runtime::{
|
||||
MAX_ERROR_BODY_BYTES, MAX_STREAM_PREFETCH_BYTES, MAX_STREAM_PREFETCH_FRAMES,
|
||||
};
|
||||
use crate::log_ids::short_request_id;
|
||||
use crate::orchestration::{
|
||||
apply_local_execution_effect, build_local_error_flow_metadata, cyber_continue_failover_enabled,
|
||||
@@ -1120,9 +1122,22 @@ async fn execute_in_process_stream_with_oauth_retry(
|
||||
) -> Result<DirectUpstreamStreamExecution, InProcessStreamExecutionError> {
|
||||
let mut execution = execute_in_process_stream(state, plan, trace_id).await?;
|
||||
apply_stream_summary_report_context(&mut execution, report_context);
|
||||
let response_text = if execution.status_code == 401
|
||||
&& stream_plan_uses_codex_agent_identity(state, plan).await
|
||||
{
|
||||
prefetch_direct_stream_error_body(&mut execution).await
|
||||
} else {
|
||||
None
|
||||
};
|
||||
if execution.status_code >= 400
|
||||
&& refresh_oauth_plan_auth_for_retry(state, plan, execution.status_code, None, trace_id)
|
||||
.await
|
||||
&& refresh_oauth_plan_auth_for_retry(
|
||||
state,
|
||||
plan,
|
||||
execution.status_code,
|
||||
response_text.as_deref(),
|
||||
trace_id,
|
||||
)
|
||||
.await
|
||||
{
|
||||
drop(execution);
|
||||
execution = execute_in_process_stream(state, plan, trace_id).await?;
|
||||
@@ -1131,6 +1146,103 @@ async fn execute_in_process_stream_with_oauth_retry(
|
||||
Ok(execution)
|
||||
}
|
||||
|
||||
async fn stream_plan_uses_codex_agent_identity(state: &AppState, plan: &ExecutionPlan) -> bool {
|
||||
state
|
||||
.read_provider_transport_snapshot(&plan.provider_id, &plan.endpoint_id, &plan.key_id)
|
||||
.await
|
||||
.ok()
|
||||
.flatten()
|
||||
.as_ref()
|
||||
.is_some_and(aether_provider_transport::is_codex_agent_identity_transport)
|
||||
}
|
||||
|
||||
async fn next_direct_upstream_response_chunk(
|
||||
response: &mut DirectUpstreamResponse,
|
||||
) -> Result<Option<Bytes>, String> {
|
||||
match response {
|
||||
DirectUpstreamResponse::Reqwest(response) => response
|
||||
.chunk()
|
||||
.await
|
||||
.map_err(|err| format_upstream_request_error(&err)),
|
||||
DirectUpstreamResponse::HyperH2c(response) => loop {
|
||||
let Some(frame) = response.body_mut().frame().await else {
|
||||
return Ok(None);
|
||||
};
|
||||
let frame = frame.map_err(|err| format_hyper_error_chain(&err))?;
|
||||
if let Ok(chunk) = frame.into_data() {
|
||||
return Ok(Some(chunk));
|
||||
}
|
||||
},
|
||||
DirectUpstreamResponse::BrowserWreq(response) => response
|
||||
.chunk()
|
||||
.await
|
||||
.map_err(|err| format_wreq_upstream_request_error(&err)),
|
||||
DirectUpstreamResponse::LocalTunnel(response) => response.next_chunk().await,
|
||||
}
|
||||
}
|
||||
|
||||
async fn prefetch_direct_stream_error_body(
|
||||
execution: &mut DirectUpstreamStreamExecution,
|
||||
) -> Option<String> {
|
||||
let mut inspected = Vec::with_capacity(MAX_ERROR_BODY_BYTES);
|
||||
let mut fully_buffered = false;
|
||||
while inspected.len() < MAX_ERROR_BODY_BYTES {
|
||||
let next_chunk = if execution.prefetched_body.is_empty() {
|
||||
match await_direct_passthrough_first_item(
|
||||
next_direct_upstream_response_chunk(&mut execution.response),
|
||||
execution.started_at,
|
||||
execution.stream_first_byte_timeout,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(item) => item,
|
||||
Err(_) => break,
|
||||
}
|
||||
} else {
|
||||
next_direct_upstream_response_chunk(&mut execution.response).await
|
||||
};
|
||||
let chunk = match next_chunk {
|
||||
Ok(Some(chunk)) => chunk,
|
||||
Ok(None) => {
|
||||
fully_buffered = true;
|
||||
break;
|
||||
}
|
||||
Err(error) => {
|
||||
execution.prefetched_body.push_back(Err(error));
|
||||
break;
|
||||
}
|
||||
};
|
||||
if chunk.is_empty() {
|
||||
continue;
|
||||
}
|
||||
|
||||
let remaining = MAX_ERROR_BODY_BYTES.saturating_sub(inspected.len());
|
||||
inspected.extend_from_slice(&chunk[..chunk.len().min(remaining)]);
|
||||
execution.prefetched_body.push_back(Ok(chunk));
|
||||
|
||||
let response_text = String::from_utf8_lossy(&inspected);
|
||||
if aether_provider_transport::is_codex_agent_identity_invalid_task_response(
|
||||
execution.status_code,
|
||||
Some(response_text.as_ref()),
|
||||
) {
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
if inspected.is_empty() {
|
||||
return None;
|
||||
}
|
||||
if fully_buffered {
|
||||
let (body_json, _) = decode_stream_error_body(&execution.headers, &inspected);
|
||||
if let Some(body_json) = body_json {
|
||||
if let Ok(response_text) = serde_json::to_string(&body_json) {
|
||||
return Some(response_text);
|
||||
}
|
||||
}
|
||||
}
|
||||
Some(String::from_utf8_lossy(&inspected).into_owned())
|
||||
}
|
||||
|
||||
fn should_use_direct_sse_passthrough(
|
||||
plan: &ExecutionPlan,
|
||||
plan_kind: &str,
|
||||
@@ -1184,9 +1296,10 @@ fn should_use_direct_sse_passthrough(
|
||||
type DirectUpstreamByteStream = BoxStream<'static, Result<Bytes, String>>;
|
||||
|
||||
fn direct_upstream_response_byte_stream(
|
||||
prefetched_body: VecDeque<Result<Bytes, String>>,
|
||||
response: DirectUpstreamResponse,
|
||||
) -> DirectUpstreamByteStream {
|
||||
match response {
|
||||
let response_stream = match response {
|
||||
DirectUpstreamResponse::Reqwest(response) => response
|
||||
.bytes_stream()
|
||||
.map(|item| item.map_err(|err| format_upstream_request_error(&err)))
|
||||
@@ -1213,7 +1326,10 @@ fn direct_upstream_response_byte_stream(
|
||||
}
|
||||
}
|
||||
.boxed(),
|
||||
}
|
||||
};
|
||||
futures_stream::iter(prefetched_body)
|
||||
.chain(response_stream)
|
||||
.boxed()
|
||||
}
|
||||
|
||||
async fn await_direct_passthrough_first_item<T, F>(
|
||||
@@ -1952,12 +2068,14 @@ impl DirectPassthroughFinalizerCore {
|
||||
|
||||
fn build_direct_passthrough_inline_body_stream(
|
||||
finalizer: DirectPassthroughFinalizer,
|
||||
prefetched_body: VecDeque<Result<Bytes, String>>,
|
||||
response: DirectUpstreamResponse,
|
||||
upstream_started_at: Instant,
|
||||
stream_first_byte_timeout: Option<Duration>,
|
||||
) -> impl futures_util::Stream<Item = Result<Bytes, IoError>> + Send + 'static {
|
||||
let state = DirectPassthroughInlineBodyState::new(
|
||||
finalizer,
|
||||
prefetched_body,
|
||||
response,
|
||||
upstream_started_at,
|
||||
stream_first_byte_timeout,
|
||||
@@ -1982,13 +2100,17 @@ struct DirectPassthroughInlineBodyState {
|
||||
impl DirectPassthroughInlineBodyState {
|
||||
fn new(
|
||||
finalizer: DirectPassthroughFinalizer,
|
||||
prefetched_body: VecDeque<Result<Bytes, String>>,
|
||||
response: DirectUpstreamResponse,
|
||||
upstream_started_at: Instant,
|
||||
stream_first_byte_timeout: Option<Duration>,
|
||||
) -> Self {
|
||||
Self {
|
||||
finalizer: Some(finalizer),
|
||||
upstream: Some(direct_upstream_response_byte_stream(response)),
|
||||
upstream: Some(direct_upstream_response_byte_stream(
|
||||
prefetched_body,
|
||||
response,
|
||||
)),
|
||||
upstream_control_filter: Some(SseControlBlockFilter::default()),
|
||||
upstream_started_at,
|
||||
stream_first_byte_timeout,
|
||||
@@ -2271,6 +2393,7 @@ async fn execute_stream_from_direct_passthrough(
|
||||
mut headers,
|
||||
provider_api_format: _,
|
||||
stream_summary_report_context: _,
|
||||
prefetched_body,
|
||||
response,
|
||||
started_at: upstream_started_at,
|
||||
stream_first_byte_timeout,
|
||||
@@ -2405,6 +2528,7 @@ async fn execute_stream_from_direct_passthrough(
|
||||
});
|
||||
let body_stream = build_direct_passthrough_inline_body_stream(
|
||||
finalizer,
|
||||
prefetched_body,
|
||||
response,
|
||||
upstream_started_at,
|
||||
stream_first_byte_timeout,
|
||||
@@ -2480,7 +2604,7 @@ async fn execute_stream_from_direct_passthrough(
|
||||
let mut last_client_chunk_elapsed_ms = 0u64;
|
||||
let mut downstream_dropped = false;
|
||||
let mut terminal_failure: Option<StreamFailureReport> = None;
|
||||
let mut upstream = direct_upstream_response_byte_stream(response);
|
||||
let mut upstream = direct_upstream_response_byte_stream(prefetched_body, response);
|
||||
let mut observed_first_upstream_body = false;
|
||||
let mut observed_first_client_send = false;
|
||||
|
||||
@@ -6333,8 +6457,8 @@ mod tests {
|
||||
use std::collections::BTreeMap;
|
||||
use std::convert::Infallible;
|
||||
use std::sync::{
|
||||
atomic::{AtomicBool, Ordering},
|
||||
Arc,
|
||||
atomic::{AtomicBool, AtomicUsize, Ordering},
|
||||
Arc, Mutex,
|
||||
};
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
@@ -6371,7 +6495,9 @@ mod tests {
|
||||
use axum::extract::ws::Message;
|
||||
use axum::extract::Request;
|
||||
use axum::routing::any;
|
||||
use axum::{http::header, http::HeaderValue, Router};
|
||||
use axum::{
|
||||
http::header, http::HeaderValue, http::StatusCode, response::IntoResponse, Json, Router,
|
||||
};
|
||||
use base64::Engine as _;
|
||||
use futures_util::StreamExt as _;
|
||||
use serde_json::{json, Value};
|
||||
@@ -6380,10 +6506,12 @@ mod tests {
|
||||
use super::{
|
||||
build_sse_body_stream, build_stream_sync_payload,
|
||||
client_format_allows_proxy_generated_sse_control_blocks,
|
||||
direct_upstream_response_byte_stream,
|
||||
ensure_stream_terminal_summary_for_missing_observed_finish,
|
||||
execute_execution_runtime_stream, execute_stream_from_frame_stream,
|
||||
maybe_apply_kiro_prompt_cache_usage_to_stream_summary, merge_stream_terminal_summary,
|
||||
parse_direct_passthrough_mode, prefetched_openai_responses_body_has_output_boundary,
|
||||
execute_execution_runtime_stream, execute_in_process_stream_with_oauth_retry,
|
||||
execute_stream_from_frame_stream, maybe_apply_kiro_prompt_cache_usage_to_stream_summary,
|
||||
merge_stream_terminal_summary, parse_direct_passthrough_mode,
|
||||
prefetch_direct_stream_error_body, prefetched_openai_responses_body_has_output_boundary,
|
||||
record_sync_terminal_usage_with_handoff,
|
||||
record_sync_terminal_usage_with_handoff_after_spawn, should_limit_direct_finalize_prefetch,
|
||||
should_probe_success_failover_before_stream, should_skip_direct_finalize_prefetch,
|
||||
@@ -6399,6 +6527,7 @@ mod tests {
|
||||
use crate::stage_metrics::RequestStageTrace;
|
||||
use crate::tunnel::{tunnel_protocol, TunnelProxyConn};
|
||||
use crate::AppState;
|
||||
use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY};
|
||||
|
||||
fn provider_catalog_stop_429_for_plan(
|
||||
plan: &ExecutionPlan,
|
||||
@@ -6481,6 +6610,124 @@ mod tests {
|
||||
InMemoryProviderCatalogReadRepository::seed(vec![provider], vec![endpoint], vec![key])
|
||||
}
|
||||
|
||||
fn provider_catalog_for_stream_auth_plan(
|
||||
plan: &ExecutionPlan,
|
||||
provider_type: &str,
|
||||
auth_type: &str,
|
||||
auth_config: Option<Value>,
|
||||
) -> InMemoryProviderCatalogReadRepository {
|
||||
let provider = StoredProviderCatalogProvider::new(
|
||||
plan.provider_id.clone(),
|
||||
plan.provider_id.clone(),
|
||||
Some("https://provider.example".to_string()),
|
||||
provider_type.to_string(),
|
||||
)
|
||||
.expect("provider should build");
|
||||
let endpoint = StoredProviderCatalogEndpoint::new(
|
||||
plan.endpoint_id.clone(),
|
||||
plan.provider_id.clone(),
|
||||
plan.provider_api_format.clone(),
|
||||
None,
|
||||
None,
|
||||
true,
|
||||
)
|
||||
.expect("endpoint should build")
|
||||
.with_transport_fields(
|
||||
plan.url.clone(),
|
||||
None,
|
||||
None,
|
||||
Some(2),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.expect("endpoint transport should build");
|
||||
let encrypted_auth_config = auth_config.map(|config| {
|
||||
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, &config.to_string())
|
||||
.expect("auth config should encrypt")
|
||||
});
|
||||
let key = StoredProviderCatalogKey::new(
|
||||
plan.key_id.clone(),
|
||||
plan.provider_id.clone(),
|
||||
plan.key_id.clone(),
|
||||
auth_type.to_string(),
|
||||
None,
|
||||
true,
|
||||
)
|
||||
.expect("key should build")
|
||||
.with_transport_fields(
|
||||
Some(json!([plan.provider_api_format.clone()])),
|
||||
None,
|
||||
encrypted_auth_config,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.expect("key transport should build");
|
||||
|
||||
InMemoryProviderCatalogReadRepository::seed(vec![provider], vec![endpoint], vec![key])
|
||||
}
|
||||
|
||||
fn direct_stream_test_plan(request_id: &str, url: String) -> ExecutionPlan {
|
||||
ExecutionPlan {
|
||||
request_id: request_id.to_string(),
|
||||
candidate_id: Some(format!("candidate-{request_id}")),
|
||||
provider_name: Some("codex".to_string()),
|
||||
provider_id: format!("provider-{request_id}"),
|
||||
endpoint_id: format!("endpoint-{request_id}"),
|
||||
key_id: format!("key-{request_id}"),
|
||||
method: "POST".to_string(),
|
||||
url,
|
||||
headers: BTreeMap::from([
|
||||
("content-type".to_string(), "application/json".to_string()),
|
||||
(
|
||||
"authorization".to_string(),
|
||||
"AgentAssertion stale-task".to_string(),
|
||||
),
|
||||
]),
|
||||
content_type: Some("application/json".to_string()),
|
||||
content_encoding: None,
|
||||
body: RequestBody::from_json(json!({"stream": true})),
|
||||
stream: true,
|
||||
client_api_format: "openai:responses".to_string(),
|
||||
provider_api_format: "openai:responses".to_string(),
|
||||
model_name: Some("gpt-5".to_string()),
|
||||
proxy: None,
|
||||
transport_profile: None,
|
||||
timeouts: Some(ExecutionTimeouts {
|
||||
connect_ms: Some(30_000),
|
||||
first_byte_ms: Some(30_000),
|
||||
..ExecutionTimeouts::default()
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
fn agent_identity_test_auth_config(task_id: &str) -> Value {
|
||||
json!({
|
||||
"provider_type": "codex",
|
||||
"auth_mode": "agentIdentity",
|
||||
"agent_runtime_id": "runtime-test",
|
||||
"agent_private_key": "MC4CAQAwBQYDK2VwBCIEIAcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcH",
|
||||
"task_id": task_id
|
||||
})
|
||||
}
|
||||
|
||||
async fn collect_direct_execution_body(
|
||||
mut execution: crate::execution_runtime::DirectUpstreamStreamExecution,
|
||||
) -> Result<Vec<u8>, String> {
|
||||
let prefetched_body = std::mem::take(&mut execution.prefetched_body);
|
||||
let mut stream = direct_upstream_response_byte_stream(prefetched_body, execution.response);
|
||||
let mut body = Vec::new();
|
||||
while let Some(item) = stream.next().await {
|
||||
body.extend_from_slice(&item?);
|
||||
}
|
||||
Ok(body)
|
||||
}
|
||||
|
||||
fn codex_cyber_policy_plan(request_id: &str) -> ExecutionPlan {
|
||||
ExecutionPlan {
|
||||
request_id: request_id.to_string(),
|
||||
@@ -6609,6 +6856,326 @@ mod tests {
|
||||
AppState::new().expect("gateway state should build")
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn agent_identity_stream_error_prefetch_is_bounded_and_replayed() {
|
||||
let upstream_body = format!(
|
||||
"{}{}",
|
||||
"x".repeat(crate::execution_runtime::MAX_ERROR_BODY_BYTES),
|
||||
"body-after-inspection-limit"
|
||||
);
|
||||
let expected_body = upstream_body.clone().into_bytes();
|
||||
let listener = crate::test_support::bind_loopback_listener()
|
||||
.await
|
||||
.expect("listener should bind");
|
||||
let addr = listener.local_addr().expect("address should resolve");
|
||||
let server = tokio::spawn(async move {
|
||||
let app = Router::new().route(
|
||||
"/responses",
|
||||
any(move || {
|
||||
let body = upstream_body.clone();
|
||||
async move { (StatusCode::UNAUTHORIZED, body) }
|
||||
}),
|
||||
);
|
||||
axum::serve(listener, app)
|
||||
.await
|
||||
.expect("server should start");
|
||||
});
|
||||
let plan =
|
||||
direct_stream_test_plan("bounded-agent-error", format!("http://{addr}/responses"));
|
||||
let mut execution = crate::execution_runtime::DirectSyncExecutionRuntime::new()
|
||||
.execute_stream(&plan)
|
||||
.await
|
||||
.expect("stream headers should execute");
|
||||
|
||||
let inspected = prefetch_direct_stream_error_body(&mut execution)
|
||||
.await
|
||||
.expect("error body should be inspected");
|
||||
|
||||
assert_eq!(
|
||||
inspected.len(),
|
||||
crate::execution_runtime::MAX_ERROR_BODY_BYTES
|
||||
);
|
||||
let replayed = collect_direct_execution_body(execution)
|
||||
.await
|
||||
.expect("prefetched response should replay");
|
||||
assert_eq!(replayed, expected_body);
|
||||
server.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn non_agent_stream_401_is_not_prefetched_and_body_passes_through() {
|
||||
let upstream_body = br#"{"error":{"code":"ordinary_unauthorized","message":"sign in"}}"#;
|
||||
let listener = crate::test_support::bind_loopback_listener()
|
||||
.await
|
||||
.expect("listener should bind");
|
||||
let addr = listener.local_addr().expect("address should resolve");
|
||||
let server = tokio::spawn(async move {
|
||||
let app = Router::new().route(
|
||||
"/responses",
|
||||
any(|| async {
|
||||
(
|
||||
StatusCode::UNAUTHORIZED,
|
||||
[(header::CONTENT_TYPE, "application/json")],
|
||||
upstream_body.as_slice(),
|
||||
)
|
||||
}),
|
||||
);
|
||||
axum::serve(listener, app)
|
||||
.await
|
||||
.expect("server should start");
|
||||
});
|
||||
let mut plan = direct_stream_test_plan("non-agent-401", format!("http://{addr}/responses"));
|
||||
plan.provider_name = Some("openai".to_string());
|
||||
let repository = provider_catalog_for_stream_auth_plan(&plan, "openai", "api_key", None);
|
||||
let state = AppState::new()
|
||||
.expect("state should build")
|
||||
.with_data_state_for_tests(
|
||||
crate::data::GatewayDataState::with_provider_transport_reader_for_tests(
|
||||
Arc::new(repository),
|
||||
DEVELOPMENT_ENCRYPTION_KEY,
|
||||
),
|
||||
);
|
||||
|
||||
let execution = execute_in_process_stream_with_oauth_retry(
|
||||
&state,
|
||||
&mut plan,
|
||||
"trace-non-agent-401",
|
||||
None,
|
||||
)
|
||||
.await
|
||||
.expect("stream request should execute");
|
||||
|
||||
assert!(execution.prefetched_body.is_empty());
|
||||
let replayed = collect_direct_execution_body(execution)
|
||||
.await
|
||||
.expect("response body should pass through");
|
||||
assert_eq!(replayed, upstream_body);
|
||||
server.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn agent_identity_stream_non_task_401_replays_original_body_without_refresh() {
|
||||
let upstream_body =
|
||||
br#"{"error":{"code":"account_disabled","message":"account unavailable"}}"#;
|
||||
let listener = crate::test_support::bind_loopback_listener()
|
||||
.await
|
||||
.expect("listener should bind");
|
||||
let addr = listener.local_addr().expect("address should resolve");
|
||||
let task_registration_hits = Arc::new(AtomicUsize::new(0));
|
||||
let task_registration_hits_for_server = Arc::clone(&task_registration_hits);
|
||||
let server = tokio::spawn(async move {
|
||||
let app = Router::new()
|
||||
.route(
|
||||
"/responses",
|
||||
any(|| async {
|
||||
(
|
||||
StatusCode::UNAUTHORIZED,
|
||||
[(header::CONTENT_TYPE, "application/json")],
|
||||
upstream_body.as_slice(),
|
||||
)
|
||||
}),
|
||||
)
|
||||
.route(
|
||||
"/api/accounts/v1/agent/runtime-test/task/register",
|
||||
any(move || {
|
||||
let hits = Arc::clone(&task_registration_hits_for_server);
|
||||
async move {
|
||||
hits.fetch_add(1, Ordering::SeqCst);
|
||||
Json(json!({"task_id": "unexpected-task"}))
|
||||
}
|
||||
}),
|
||||
);
|
||||
axum::serve(listener, app)
|
||||
.await
|
||||
.expect("server should start");
|
||||
});
|
||||
let mut plan =
|
||||
direct_stream_test_plan("agent-non-task-401", format!("http://{addr}/responses"));
|
||||
let repository = Arc::new(provider_catalog_for_stream_auth_plan(
|
||||
&plan,
|
||||
"codex",
|
||||
"oauth",
|
||||
Some(agent_identity_test_auth_config("task-old")),
|
||||
));
|
||||
let oauth_refresh =
|
||||
aether_provider_transport::LocalOAuthRefreshCoordinator::with_adapters_for_tests(vec![
|
||||
Arc::new(
|
||||
aether_provider_transport::CodexAgentIdentityRefreshAdapter::default()
|
||||
.with_auth_api_base_url_for_tests(format!("http://{addr}/api/accounts")),
|
||||
),
|
||||
]);
|
||||
let state = AppState::new()
|
||||
.expect("state should build")
|
||||
.with_data_state_for_tests(
|
||||
crate::data::GatewayDataState::with_provider_catalog_repository_for_tests(
|
||||
repository,
|
||||
)
|
||||
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
|
||||
)
|
||||
.with_oauth_refresh_coordinator_for_tests(oauth_refresh);
|
||||
|
||||
let execution = execute_in_process_stream_with_oauth_retry(
|
||||
&state,
|
||||
&mut plan,
|
||||
"trace-agent-non-task-401",
|
||||
None,
|
||||
)
|
||||
.await
|
||||
.expect("stream request should execute");
|
||||
|
||||
assert!(!execution.prefetched_body.is_empty());
|
||||
assert_eq!(task_registration_hits.load(Ordering::SeqCst), 0);
|
||||
let replayed = collect_direct_execution_body(execution)
|
||||
.await
|
||||
.expect("response body should replay");
|
||||
assert_eq!(replayed, upstream_body);
|
||||
server.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn agent_identity_stream_invalid_task_refreshes_and_retries_once() {
|
||||
let upstream_hits = Arc::new(AtomicUsize::new(0));
|
||||
let upstream_hits_for_server = Arc::clone(&upstream_hits);
|
||||
let task_registration_hits = Arc::new(AtomicUsize::new(0));
|
||||
let task_registration_hits_for_server = Arc::clone(&task_registration_hits);
|
||||
let observed_authorization = Arc::new(Mutex::new(Vec::<String>::new()));
|
||||
let observed_authorization_for_server = Arc::clone(&observed_authorization);
|
||||
let listener = crate::test_support::bind_loopback_listener()
|
||||
.await
|
||||
.expect("listener should bind");
|
||||
let addr = listener.local_addr().expect("address should resolve");
|
||||
let server = tokio::spawn(async move {
|
||||
let app = Router::new()
|
||||
.route(
|
||||
"/responses",
|
||||
any(move |request: Request| {
|
||||
let hits = Arc::clone(&upstream_hits_for_server);
|
||||
let authorizations = Arc::clone(&observed_authorization_for_server);
|
||||
async move {
|
||||
let authorization = request
|
||||
.headers()
|
||||
.get(header::AUTHORIZATION)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.unwrap_or_default()
|
||||
.to_string();
|
||||
authorizations
|
||||
.lock()
|
||||
.expect("authorization mutex should lock")
|
||||
.push(authorization);
|
||||
if hits.fetch_add(1, Ordering::SeqCst) == 0 {
|
||||
(
|
||||
StatusCode::UNAUTHORIZED,
|
||||
Json(json!({
|
||||
"error": {
|
||||
"code": "invalid_task_id",
|
||||
"message": "registered task expired"
|
||||
}
|
||||
})),
|
||||
)
|
||||
.into_response()
|
||||
} else {
|
||||
(StatusCode::OK, Json(json!({"ok": true}))).into_response()
|
||||
}
|
||||
}
|
||||
}),
|
||||
)
|
||||
.route(
|
||||
"/api/accounts/v1/agent/runtime-test/task/register",
|
||||
any(move || {
|
||||
let hits = Arc::clone(&task_registration_hits_for_server);
|
||||
async move {
|
||||
hits.fetch_add(1, Ordering::SeqCst);
|
||||
Json(json!({"task_id": "task-new"}))
|
||||
}
|
||||
}),
|
||||
);
|
||||
axum::serve(listener, app)
|
||||
.await
|
||||
.expect("server should start");
|
||||
});
|
||||
let mut plan =
|
||||
direct_stream_test_plan("agent-invalid-task", format!("http://{addr}/responses"));
|
||||
let repository = Arc::new(provider_catalog_for_stream_auth_plan(
|
||||
&plan,
|
||||
"codex",
|
||||
"oauth",
|
||||
Some(agent_identity_test_auth_config("task-old")),
|
||||
));
|
||||
let oauth_refresh =
|
||||
aether_provider_transport::LocalOAuthRefreshCoordinator::with_adapters_for_tests(vec![
|
||||
Arc::new(
|
||||
aether_provider_transport::CodexAgentIdentityRefreshAdapter::default()
|
||||
.with_auth_api_base_url_for_tests(format!("http://{addr}/api/accounts")),
|
||||
),
|
||||
]);
|
||||
let state = AppState::new()
|
||||
.expect("state should build")
|
||||
.with_data_state_for_tests(
|
||||
crate::data::GatewayDataState::with_provider_catalog_repository_for_tests(
|
||||
repository,
|
||||
)
|
||||
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
|
||||
)
|
||||
.with_oauth_refresh_coordinator_for_tests(oauth_refresh);
|
||||
let transport = state
|
||||
.read_provider_transport_snapshot(&plan.provider_id, &plan.endpoint_id, &plan.key_id)
|
||||
.await
|
||||
.expect("transport should load")
|
||||
.expect("transport should exist");
|
||||
let initial_authorization = match state
|
||||
.resolve_local_oauth_request_auth(&transport)
|
||||
.await
|
||||
.expect("initial Agent Identity auth should resolve")
|
||||
.expect("initial Agent Identity auth should exist")
|
||||
{
|
||||
aether_provider_transport::LocalResolvedOAuthRequestAuth::Header { name, value } => {
|
||||
assert_eq!(name, "authorization");
|
||||
value
|
||||
}
|
||||
aether_provider_transport::LocalResolvedOAuthRequestAuth::Kiro(_) => {
|
||||
panic!("Agent Identity should resolve to header auth")
|
||||
}
|
||||
};
|
||||
assert!(
|
||||
aether_provider_transport::codex_agent_identity_authorization_matches_transport(
|
||||
&transport,
|
||||
&initial_authorization,
|
||||
)
|
||||
);
|
||||
plan.headers
|
||||
.insert("authorization".to_string(), initial_authorization.clone());
|
||||
|
||||
let execution = execute_in_process_stream_with_oauth_retry(
|
||||
&state,
|
||||
&mut plan,
|
||||
"trace-agent-invalid-task",
|
||||
None,
|
||||
)
|
||||
.await
|
||||
.expect("stream request should recover");
|
||||
|
||||
assert_eq!(execution.status_code, 200);
|
||||
assert!(execution.prefetched_body.is_empty());
|
||||
assert_eq!(upstream_hits.load(Ordering::SeqCst), 2);
|
||||
assert_eq!(task_registration_hits.load(Ordering::SeqCst), 1);
|
||||
let authorizations = observed_authorization
|
||||
.lock()
|
||||
.expect("authorization mutex should lock");
|
||||
assert_eq!(authorizations.len(), 2);
|
||||
assert_eq!(authorizations[0], initial_authorization);
|
||||
assert!(authorizations[1].starts_with("AgentAssertion "));
|
||||
assert_ne!(authorizations[1], authorizations[0]);
|
||||
drop(authorizations);
|
||||
let replayed = collect_direct_execution_body(execution)
|
||||
.await
|
||||
.expect("retried response body should read");
|
||||
assert_eq!(
|
||||
serde_json::from_slice::<Value>(&replayed).expect("response should be JSON"),
|
||||
json!({"ok": true})
|
||||
);
|
||||
server.abort();
|
||||
}
|
||||
|
||||
struct BlockingStreamingRequestCandidateRepository {
|
||||
inner: InMemoryRequestCandidateRepository,
|
||||
block_streaming: AtomicBool,
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
use std::collections::BTreeMap;
|
||||
use std::collections::{BTreeMap, VecDeque};
|
||||
use std::future::Future;
|
||||
use std::io::Error as IoError;
|
||||
use std::time::{Duration, Instant};
|
||||
@@ -38,6 +38,7 @@ pub(crate) fn build_direct_execution_frame_stream(
|
||||
headers,
|
||||
provider_api_format,
|
||||
stream_summary_report_context,
|
||||
prefetched_body,
|
||||
response,
|
||||
started_at,
|
||||
stream_first_byte_timeout,
|
||||
@@ -70,7 +71,14 @@ pub(crate) fn build_direct_execution_frame_stream(
|
||||
|
||||
if should_buffer_non_stream_response(&headers, &observer_context) {
|
||||
let original_headers = headers.clone();
|
||||
match buffer_non_sse_upstream_body(response, started_at, stream_first_byte_timeout).await {
|
||||
match buffer_non_sse_upstream_body(
|
||||
prefetched_body,
|
||||
response,
|
||||
started_at,
|
||||
stream_first_byte_timeout,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(buffered) => {
|
||||
let mut response_headers = original_headers;
|
||||
let mut response_body = Bytes::from(buffered.body_bytes);
|
||||
@@ -192,6 +200,61 @@ pub(crate) fn build_direct_execution_frame_stream(
|
||||
let mut upstream_bytes = 0u64;
|
||||
let mut ttfb_ms = None;
|
||||
let mut first_chunk_telemetry_emitted = false;
|
||||
let mut prefetched_body_failed = false;
|
||||
for item in prefetched_body {
|
||||
match item {
|
||||
Ok(chunk) => {
|
||||
if ttfb_ms.is_none() {
|
||||
ttfb_ms = Some(started_at.elapsed().as_millis() as u64);
|
||||
}
|
||||
if !first_chunk_telemetry_emitted {
|
||||
match encode_telemetry_frame(ttfb_ms, ttfb_ms, upstream_bytes) {
|
||||
Ok(frame) => yield Ok(frame),
|
||||
Err(err) => {
|
||||
yield Err(err);
|
||||
return;
|
||||
}
|
||||
}
|
||||
first_chunk_telemetry_emitted = true;
|
||||
}
|
||||
upstream_bytes += chunk.len() as u64;
|
||||
observe_stream_chunk(
|
||||
&mut stream_terminal_observer,
|
||||
&normalized_observer_context,
|
||||
private_stream_normalizer.as_mut(),
|
||||
&mut observer_buffered,
|
||||
chunk.as_ref(),
|
||||
);
|
||||
match encode_data_frame(&chunk) {
|
||||
Ok(frame) => yield Ok(frame),
|
||||
Err(err) => {
|
||||
yield Err(err);
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(message) => {
|
||||
warn!(
|
||||
event_name = "stream_pump_body_read_error",
|
||||
log_type = "ops",
|
||||
status_code,
|
||||
upstream_bytes,
|
||||
error = %message,
|
||||
"upstream body stream read error"
|
||||
);
|
||||
match encode_error_frame(status_code, message) {
|
||||
Ok(frame) => yield Ok(frame),
|
||||
Err(encode_err) => {
|
||||
yield Err(encode_err);
|
||||
return;
|
||||
}
|
||||
}
|
||||
prefetched_body_failed = true;
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
if !prefetched_body_failed {
|
||||
match response {
|
||||
DirectUpstreamResponse::Reqwest(response) => {
|
||||
let mut bytes_stream = response.bytes_stream();
|
||||
@@ -516,6 +579,7 @@ pub(crate) fn build_direct_execution_frame_stream(
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
let summary = finalize_stream_terminal_summary(
|
||||
&mut stream_terminal_observer,
|
||||
&normalized_observer_context,
|
||||
@@ -691,6 +755,7 @@ fn should_buffer_non_stream_response(
|
||||
}
|
||||
|
||||
async fn buffer_non_sse_upstream_body(
|
||||
mut prefetched_body: VecDeque<Result<Bytes, String>>,
|
||||
response: DirectUpstreamResponse,
|
||||
started_at: Instant,
|
||||
stream_first_byte_timeout: Option<Duration>,
|
||||
@@ -699,6 +764,26 @@ async fn buffer_non_sse_upstream_body(
|
||||
let mut upstream_bytes = 0u64;
|
||||
let mut ttfb_ms = None;
|
||||
|
||||
while let Some(item) = prefetched_body.pop_front() {
|
||||
match item {
|
||||
Ok(chunk) => {
|
||||
if ttfb_ms.is_none() {
|
||||
ttfb_ms = Some(started_at.elapsed().as_millis() as u64);
|
||||
}
|
||||
upstream_bytes += chunk.len() as u64;
|
||||
body_bytes.extend_from_slice(&chunk);
|
||||
}
|
||||
Err(message) => {
|
||||
return Err(BufferedUpstreamBodyError {
|
||||
message,
|
||||
ttfb_ms,
|
||||
upstream_bytes,
|
||||
first_byte_timeout: None,
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
match response {
|
||||
DirectUpstreamResponse::Reqwest(response) => {
|
||||
let mut bytes_stream = response.bytes_stream();
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
use std::collections::{BTreeMap, HashMap, HashSet};
|
||||
use std::collections::{BTreeMap, HashMap, HashSet, VecDeque};
|
||||
use std::error::Error as _;
|
||||
use std::future::Future;
|
||||
use std::io::Read;
|
||||
@@ -607,6 +607,7 @@ pub(crate) struct DirectUpstreamStreamExecution {
|
||||
pub(crate) headers: BTreeMap<String, String>,
|
||||
pub(crate) provider_api_format: String,
|
||||
pub(crate) stream_summary_report_context: Value,
|
||||
pub(crate) prefetched_body: VecDeque<Result<Bytes, String>>,
|
||||
pub(crate) response: DirectUpstreamResponse,
|
||||
pub(crate) started_at: Instant,
|
||||
pub(crate) stream_first_byte_timeout: Option<Duration>,
|
||||
@@ -711,6 +712,7 @@ impl DirectSyncExecutionRuntime {
|
||||
headers,
|
||||
provider_api_format: plan.provider_api_format.clone(),
|
||||
stream_summary_report_context,
|
||||
prefetched_body: VecDeque::new(),
|
||||
response: response.into_direct_upstream_response(),
|
||||
started_at,
|
||||
stream_first_byte_timeout: resolve_stream_first_byte_timeout(plan),
|
||||
@@ -824,6 +826,7 @@ pub(crate) async fn execute_stream_plan_via_local_tunnel(
|
||||
headers,
|
||||
provider_api_format: plan.provider_api_format.clone(),
|
||||
stream_summary_report_context: build_stream_summary_report_context(plan),
|
||||
prefetched_body: VecDeque::new(),
|
||||
response: DirectUpstreamResponse::LocalTunnel(response),
|
||||
started_at,
|
||||
stream_first_byte_timeout: resolve_stream_first_byte_timeout(plan),
|
||||
|
||||
@@ -190,7 +190,10 @@ pub(crate) async fn build_admin_global_model_routing_payload(
|
||||
let payload = json!({
|
||||
"id": key.id,
|
||||
"name": key.name,
|
||||
"masked_key": state.masked_catalog_api_key(key),
|
||||
"masked_key": state.masked_catalog_api_key_for_provider(
|
||||
key,
|
||||
&provider.provider_type,
|
||||
),
|
||||
"is_active": key.is_active,
|
||||
"is_adaptive": is_adaptive,
|
||||
"effective_rpm": effective_rpm,
|
||||
@@ -252,9 +255,9 @@ pub(crate) async fn build_admin_global_model_routing_payload(
|
||||
.iter()
|
||||
.map(|provider| provider.id.clone())
|
||||
.collect::<Vec<_>>();
|
||||
let active_provider_name_by_id = active_providers
|
||||
let active_provider_metadata_by_id = active_providers
|
||||
.into_iter()
|
||||
.map(|provider| (provider.id, provider.name))
|
||||
.map(|provider| (provider.id, (provider.name, provider.provider_type)))
|
||||
.collect::<BTreeMap<_, _>>();
|
||||
let active_keys = if active_provider_ids.is_empty() {
|
||||
Vec::new()
|
||||
@@ -273,14 +276,14 @@ pub(crate) async fn build_admin_global_model_routing_payload(
|
||||
if allowed_models.is_empty() {
|
||||
continue;
|
||||
}
|
||||
let provider_name = active_provider_name_by_id
|
||||
let (provider_name, provider_type) = active_provider_metadata_by_id
|
||||
.get(&key.provider_id)
|
||||
.cloned()
|
||||
.unwrap_or_default();
|
||||
all_keys_whitelist.push(json!({
|
||||
"key_id": key.id,
|
||||
"key_name": key.name,
|
||||
"masked_key": state.masked_catalog_api_key(&key),
|
||||
"masked_key": state.masked_catalog_api_key_for_provider(&key, &provider_type),
|
||||
"provider_id": key.provider_id,
|
||||
"provider_name": provider_name,
|
||||
"allowed_models": allowed_models,
|
||||
|
||||
+9
-2
@@ -149,8 +149,15 @@ pub(super) async fn build_admin_monitoring_cache_affinities_response(
|
||||
.map(|item| item.base_url.clone())
|
||||
.filter(|value| !value.trim().is_empty());
|
||||
let key_name = key.map(|item| item.name.clone());
|
||||
let key_prefix =
|
||||
key.and_then(|item| admin_monitoring_masked_provider_key_prefix(state, item));
|
||||
let key_prefix = key.and_then(|item| {
|
||||
admin_monitoring_masked_provider_key_prefix(
|
||||
state,
|
||||
item,
|
||||
provider
|
||||
.map(|provider| provider.provider_type.as_str())
|
||||
.unwrap_or(""),
|
||||
)
|
||||
});
|
||||
let user_id_text = user_id.clone();
|
||||
let username = user.map(|item| item.username.clone());
|
||||
let email = user.and_then(|item| item.email.clone());
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::provider_key_auth::provider_key_auth_config_uses_header_authorization;
|
||||
use crate::provider_key_auth::{
|
||||
provider_key_auth_config_is_agent_identity, provider_key_auth_config_uses_header_authorization,
|
||||
};
|
||||
use aether_crypto::decrypt_python_fernet_ciphertext;
|
||||
#[cfg(test)]
|
||||
use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY;
|
||||
@@ -26,13 +28,15 @@ pub(super) fn admin_monitoring_masked_user_api_key_prefix(
|
||||
pub(super) fn admin_monitoring_masked_provider_key_prefix(
|
||||
state: &AdminAppState<'_>,
|
||||
key: &StoredProviderCatalogKey,
|
||||
provider_type: &str,
|
||||
) -> Option<String> {
|
||||
match key.auth_type.trim() {
|
||||
"service_account" | "vertex_ai" => Some("[Service Account]".to_string()),
|
||||
"oauth" => {
|
||||
if provider_key_auth_config_uses_header_authorization(
|
||||
state.parse_catalog_auth_config_json(key).as_ref(),
|
||||
) {
|
||||
let auth_config = state.parse_catalog_auth_config_json(key);
|
||||
if provider_key_auth_config_is_agent_identity(provider_type, auth_config.as_ref()) {
|
||||
Some("[Agent Identity]".to_string())
|
||||
} else if provider_key_auth_config_uses_header_authorization(auth_config.as_ref()) {
|
||||
Some("[OAuth Header]".to_string())
|
||||
} else {
|
||||
Some("[OAuth Token]".to_string())
|
||||
@@ -115,3 +119,52 @@ pub(super) fn admin_monitoring_cache_affinity_sort_value(value: Option<&serde_js
|
||||
}
|
||||
0.0
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::admin_monitoring_masked_provider_key_prefix;
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::AppState;
|
||||
use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY};
|
||||
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
|
||||
|
||||
#[test]
|
||||
fn monitoring_labels_agent_identity_instead_of_oauth_token() {
|
||||
let app = AppState::new().expect("gateway should build");
|
||||
let state = AdminAppState::new(&app);
|
||||
let encrypted_placeholder =
|
||||
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "__placeholder__")
|
||||
.expect("placeholder should encrypt");
|
||||
let encrypted_auth_config = encrypt_python_fernet_plaintext(
|
||||
DEVELOPMENT_ENCRYPTION_KEY,
|
||||
r#"{"auth_mode":"agentIdentity","agent_runtime_id":"runtime-1","agent_private_key":"base64-private-key","task_id":"task-1"}"#,
|
||||
)
|
||||
.expect("auth config should encrypt");
|
||||
let key = StoredProviderCatalogKey::new(
|
||||
"key-agent".to_string(),
|
||||
"provider-codex".to_string(),
|
||||
"agent".to_string(),
|
||||
"oauth".to_string(),
|
||||
None,
|
||||
true,
|
||||
)
|
||||
.expect("key should build")
|
||||
.with_transport_fields(
|
||||
None,
|
||||
encrypted_placeholder,
|
||||
Some(encrypted_auth_config),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.expect("transport should build");
|
||||
|
||||
assert_eq!(
|
||||
admin_monitoring_masked_provider_key_prefix(&state, &key, "codex").as_deref(),
|
||||
Some("[Agent Identity]")
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -4,9 +4,11 @@ use crate::handlers::admin::provider::shared::support::{
|
||||
};
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::handlers::admin::shared::{
|
||||
decrypt_catalog_secret_with_fallbacks, json_string_list, take_secret_prefix, take_secret_suffix,
|
||||
decrypt_catalog_secret_with_fallbacks, json_string_list, parse_catalog_auth_config_json,
|
||||
take_secret_prefix, take_secret_suffix,
|
||||
};
|
||||
use crate::handlers::public::matches_model_mapping_for_models;
|
||||
use crate::provider_key_auth::provider_key_auth_config_is_agent_identity;
|
||||
use crate::{GatewayError, LocalProviderDeleteTaskState};
|
||||
use aether_data_contracts::repository::global_models::{
|
||||
AdminGlobalModelListQuery, AdminProviderModelListQuery, PublicGlobalModelQuery,
|
||||
@@ -169,7 +171,12 @@ pub(crate) fn global_model_mapping_patterns_from_config(
|
||||
pub(crate) fn mapping_preview_masked_catalog_api_key(
|
||||
state: &AdminAppState<'_>,
|
||||
key: &StoredProviderCatalogKey,
|
||||
provider_type: &str,
|
||||
) -> String {
|
||||
let auth_config = parse_catalog_auth_config_json(state.as_ref(), key);
|
||||
if provider_key_auth_config_is_agent_identity(provider_type, auth_config.as_ref()) {
|
||||
return "[Agent Identity]".to_string();
|
||||
}
|
||||
let ciphertext = key.encrypted_api_key.as_deref().unwrap_or("").trim();
|
||||
if ciphertext.is_empty() {
|
||||
return "***".to_string();
|
||||
@@ -345,7 +352,11 @@ pub(crate) async fn build_admin_provider_mapping_preview_payload(
|
||||
key_payloads.push(json!({
|
||||
"key_id": key.id,
|
||||
"key_name": key.name,
|
||||
"masked_key": mapping_preview_masked_catalog_api_key(state, &key),
|
||||
"masked_key": mapping_preview_masked_catalog_api_key(
|
||||
state,
|
||||
&key,
|
||||
&provider.provider_type,
|
||||
),
|
||||
"is_active": key.is_active,
|
||||
"allowed_models": allowed_models,
|
||||
"matching_global_models": matching_global_models,
|
||||
@@ -363,3 +374,52 @@ pub(crate) async fn build_admin_provider_mapping_preview_payload(
|
||||
"truncated_models": truncated_models,
|
||||
}))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::mapping_preview_masked_catalog_api_key;
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::AppState;
|
||||
use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY};
|
||||
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
|
||||
|
||||
#[test]
|
||||
fn delete_preview_never_masks_internal_agent_identity_placeholder() {
|
||||
let app = AppState::new().expect("gateway should build");
|
||||
let state = AdminAppState::new(&app);
|
||||
let encrypted_placeholder =
|
||||
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "__placeholder__")
|
||||
.expect("placeholder should encrypt");
|
||||
let encrypted_auth_config = encrypt_python_fernet_plaintext(
|
||||
DEVELOPMENT_ENCRYPTION_KEY,
|
||||
r#"{"auth_mode":"agentIdentity","agent_runtime_id":"runtime-1","agent_private_key":"base64-private-key","task_id":"task-1"}"#,
|
||||
)
|
||||
.expect("auth config should encrypt");
|
||||
let key = StoredProviderCatalogKey::new(
|
||||
"key-agent".to_string(),
|
||||
"provider-codex".to_string(),
|
||||
"agent".to_string(),
|
||||
"oauth".to_string(),
|
||||
None,
|
||||
true,
|
||||
)
|
||||
.expect("key should build")
|
||||
.with_transport_fields(
|
||||
None,
|
||||
encrypted_placeholder,
|
||||
Some(encrypted_auth_config),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.expect("transport should build");
|
||||
|
||||
assert_eq!(
|
||||
mapping_preview_masked_catalog_api_key(&state, &key, "codex"),
|
||||
"[Agent Identity]"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
+197
-106
@@ -13,7 +13,10 @@ use super::progress::{
|
||||
maybe_report_admin_provider_oauth_batch_import_progress,
|
||||
AdminProviderOAuthBatchProgressReporter,
|
||||
};
|
||||
use crate::handlers::admin::provider::oauth::duplicates::find_duplicate_provider_oauth_key;
|
||||
use crate::handlers::admin::provider::oauth::duplicates::{
|
||||
acquire_codex_oauth_account_locks, find_duplicate_provider_oauth_key,
|
||||
release_codex_oauth_account_locks,
|
||||
};
|
||||
use crate::handlers::admin::provider::oauth::provisioning::build_provider_oauth_auth_config_from_token_payload;
|
||||
use crate::handlers::admin::provider::oauth::provisioning::{
|
||||
create_provider_oauth_catalog_key, provider_oauth_active_api_formats,
|
||||
@@ -53,38 +56,92 @@ fn sanitize_windsurf_batch_import_error(error: &OAuthError) -> String {
|
||||
}
|
||||
}
|
||||
|
||||
fn copy_codex_agent_identity_field(
|
||||
const CODEX_AGENT_IDENTITY_SAFE_FIELDS: &[(&str, &[&str])] = &[
|
||||
("agent_runtime_id", &["agent_runtime_id", "agentRuntimeId"]),
|
||||
(
|
||||
"agent_private_key",
|
||||
&["agent_private_key", "agentPrivateKey"],
|
||||
),
|
||||
("task_id", &["task_id", "taskId"]),
|
||||
(
|
||||
"account_id",
|
||||
&[
|
||||
"account_id",
|
||||
"accountId",
|
||||
"chatgpt_account_id",
|
||||
"chatgptAccountId",
|
||||
],
|
||||
),
|
||||
(
|
||||
"account_user_id",
|
||||
&[
|
||||
"account_user_id",
|
||||
"accountUserId",
|
||||
"chatgpt_account_user_id",
|
||||
"chatgptAccountUserId",
|
||||
],
|
||||
),
|
||||
(
|
||||
"user_id",
|
||||
&["user_id", "userId", "chatgpt_user_id", "chatgptUserId"],
|
||||
),
|
||||
("email", &["email"]),
|
||||
(
|
||||
"plan_type",
|
||||
&[
|
||||
"plan_type",
|
||||
"planType",
|
||||
"chatgpt_plan_type",
|
||||
"chatgptPlanType",
|
||||
],
|
||||
),
|
||||
("account_name", &["account_name", "accountName"]),
|
||||
(
|
||||
"is_fedramp",
|
||||
&[
|
||||
"is_fedramp",
|
||||
"chatgpt_account_is_fedramp",
|
||||
"chatgptAccountIsFedramp",
|
||||
],
|
||||
),
|
||||
("workspace_id", &["workspace_id", "workspaceId"]),
|
||||
];
|
||||
|
||||
fn copy_codex_agent_identity_safe_fields(
|
||||
auth_config: &mut Map<String, Value>,
|
||||
nested: &Map<String, Value>,
|
||||
canonical_key: &str,
|
||||
aliases: &[&str],
|
||||
preferred: Option<&Map<String, Value>>,
|
||||
fallback: &Map<String, Value>,
|
||||
) {
|
||||
if auth_config.contains_key(canonical_key) {
|
||||
return;
|
||||
}
|
||||
if let Some(value) = aliases.iter().find_map(|key| nested.get(*key)).cloned() {
|
||||
auth_config.insert(canonical_key.to_string(), value);
|
||||
for (canonical_key, aliases) in CODEX_AGENT_IDENTITY_SAFE_FIELDS {
|
||||
if auth_config.contains_key(*canonical_key) {
|
||||
continue;
|
||||
}
|
||||
let value = preferred
|
||||
.and_then(|map| aliases.iter().find_map(|key| map.get(*key)))
|
||||
.or_else(|| aliases.iter().find_map(|key| fallback.get(*key)))
|
||||
.cloned();
|
||||
let Some(value) = value else {
|
||||
continue;
|
||||
};
|
||||
let type_is_allowed = if *canonical_key == "is_fedramp" {
|
||||
value.is_boolean()
|
||||
} else {
|
||||
value.as_str().is_some_and(|text| !text.trim().is_empty())
|
||||
};
|
||||
if !type_is_allowed {
|
||||
continue;
|
||||
}
|
||||
auth_config.insert((*canonical_key).to_string(), value);
|
||||
}
|
||||
}
|
||||
|
||||
fn remove_codex_agent_identity_oauth_tokens(auth_config: &mut Map<String, Value>) {
|
||||
for key in [
|
||||
"access_token",
|
||||
"accessToken",
|
||||
"refresh_token",
|
||||
"refreshToken",
|
||||
"id_token",
|
||||
"idToken",
|
||||
"expires_at",
|
||||
"expiresAt",
|
||||
"expires_in",
|
||||
"expiresIn",
|
||||
] {
|
||||
auth_config.remove(key);
|
||||
}
|
||||
fn sanitize_codex_agent_identity_nested_fields(nested: &Map<String, Value>) -> Map<String, Value> {
|
||||
let mut sanitized = Map::new();
|
||||
copy_codex_agent_identity_safe_fields(&mut sanitized, None, nested);
|
||||
sanitized
|
||||
}
|
||||
|
||||
fn codex_agent_identity_auth_config_from_import(
|
||||
pub(super) fn codex_agent_identity_auth_config_from_import(
|
||||
entry: &AdminProviderOAuthBatchImportEntry,
|
||||
) -> Result<Option<Map<String, Value>>, String> {
|
||||
let Some(raw_credentials) = entry.raw_credentials.as_ref() else {
|
||||
@@ -93,81 +150,24 @@ fn codex_agent_identity_auth_config_from_import(
|
||||
if !aether_provider_transport::is_codex_agent_identity_auth_config_value(raw_credentials) {
|
||||
return Ok(None);
|
||||
}
|
||||
let mut auth_config = raw_credentials
|
||||
let root = raw_credentials
|
||||
.as_object()
|
||||
.cloned()
|
||||
.ok_or_else(|| "Agent Identity 凭据必须是 JSON 对象".to_string())?;
|
||||
remove_codex_agent_identity_oauth_tokens(&mut auth_config);
|
||||
for nested_key in ["agent_identity", "agentIdentity"] {
|
||||
if let Some(nested) = auth_config
|
||||
.get_mut(nested_key)
|
||||
.and_then(Value::as_object_mut)
|
||||
{
|
||||
remove_codex_agent_identity_oauth_tokens(nested);
|
||||
}
|
||||
}
|
||||
let nested = auth_config
|
||||
let nested = root
|
||||
.get("agent_identity")
|
||||
.or_else(|| auth_config.get("agentIdentity"))
|
||||
.or_else(|| root.get("agentIdentity"))
|
||||
.and_then(Value::as_object)
|
||||
.cloned();
|
||||
let root = auth_config.clone();
|
||||
for (canonical_key, aliases) in [
|
||||
(
|
||||
"agent_runtime_id",
|
||||
&["agent_runtime_id", "agentRuntimeId"][..],
|
||||
),
|
||||
(
|
||||
"agent_private_key",
|
||||
&["agent_private_key", "agentPrivateKey"][..],
|
||||
),
|
||||
("task_id", &["task_id", "taskId"][..]),
|
||||
(
|
||||
"account_id",
|
||||
&[
|
||||
"account_id",
|
||||
"accountId",
|
||||
"chatgpt_account_id",
|
||||
"chatgptAccountId",
|
||||
][..],
|
||||
),
|
||||
(
|
||||
"account_user_id",
|
||||
&[
|
||||
"account_user_id",
|
||||
"accountUserId",
|
||||
"chatgpt_account_user_id",
|
||||
"chatgptAccountUserId",
|
||||
][..],
|
||||
),
|
||||
(
|
||||
"user_id",
|
||||
&["user_id", "userId", "chatgpt_user_id", "chatgptUserId"][..],
|
||||
),
|
||||
("email", &["email"][..]),
|
||||
(
|
||||
"plan_type",
|
||||
&[
|
||||
"plan_type",
|
||||
"planType",
|
||||
"chatgpt_plan_type",
|
||||
"chatgptPlanType",
|
||||
][..],
|
||||
),
|
||||
("account_name", &["account_name", "accountName"][..]),
|
||||
(
|
||||
"is_fedramp",
|
||||
&[
|
||||
"is_fedramp",
|
||||
"chatgpt_account_is_fedramp",
|
||||
"chatgptAccountIsFedramp",
|
||||
][..],
|
||||
),
|
||||
] {
|
||||
if let Some(nested) = nested.as_ref() {
|
||||
copy_codex_agent_identity_field(&mut auth_config, nested, canonical_key, aliases);
|
||||
let mut auth_config = Map::new();
|
||||
copy_codex_agent_identity_safe_fields(&mut auth_config, nested.as_ref(), root);
|
||||
if let Some(nested) = nested.as_ref() {
|
||||
let sanitized_nested = sanitize_codex_agent_identity_nested_fields(nested);
|
||||
if !sanitized_nested.is_empty() {
|
||||
auth_config.insert(
|
||||
"agent_identity".to_string(),
|
||||
Value::Object(sanitized_nested),
|
||||
);
|
||||
}
|
||||
copy_codex_agent_identity_field(&mut auth_config, &root, canonical_key, aliases);
|
||||
}
|
||||
auth_config.insert("provider_type".to_string(), json!("codex"));
|
||||
auth_config.insert("auth_mode".to_string(), json!("agentIdentity"));
|
||||
@@ -548,10 +548,49 @@ pub(super) async fn execute_admin_provider_oauth_batch_import(
|
||||
}
|
||||
}
|
||||
|
||||
let is_agent_identity = provider_type.eq_ignore_ascii_case("codex")
|
||||
&& aether_provider_transport::is_codex_agent_identity_auth_config_value(
|
||||
&Value::Object(auth_config.clone()),
|
||||
);
|
||||
let codex_oauth_account_leases =
|
||||
if provider_type.eq_ignore_ascii_case("codex") && !is_agent_identity {
|
||||
match acquire_codex_oauth_account_locks(
|
||||
state,
|
||||
provider_id,
|
||||
&auth_config,
|
||||
"batch-import",
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(leases) => leases,
|
||||
Err(error) => {
|
||||
failed += 1;
|
||||
results.push(json!({
|
||||
"index": index,
|
||||
"status": "error",
|
||||
"error": error.detail(),
|
||||
"replaced": false,
|
||||
}));
|
||||
maybe_report_admin_provider_oauth_batch_import_progress(
|
||||
&mut progress,
|
||||
entries.len(),
|
||||
success,
|
||||
failed,
|
||||
&results,
|
||||
)
|
||||
.await;
|
||||
continue;
|
||||
}
|
||||
}
|
||||
} else {
|
||||
Vec::new()
|
||||
};
|
||||
|
||||
let duplicate =
|
||||
match find_duplicate_provider_oauth_key(state, provider_id, &auth_config, None).await {
|
||||
Ok(value) => value,
|
||||
Err(detail) => {
|
||||
release_codex_oauth_account_locks(state, codex_oauth_account_leases).await;
|
||||
failed += 1;
|
||||
results.push(json!({
|
||||
"index": index,
|
||||
@@ -573,7 +612,7 @@ pub(super) async fn execute_admin_provider_oauth_batch_import(
|
||||
|
||||
let replaced = duplicate.is_some();
|
||||
let (persisted_key, key_name) = if let Some(existing_key) = duplicate {
|
||||
match update_existing_provider_oauth_catalog_key(
|
||||
let update_result = update_existing_provider_oauth_catalog_key(
|
||||
state,
|
||||
&existing_key,
|
||||
provider_type,
|
||||
@@ -583,10 +622,15 @@ pub(super) async fn execute_admin_provider_oauth_batch_import(
|
||||
key_proxy.clone(),
|
||||
expires_at,
|
||||
)
|
||||
.await?
|
||||
{
|
||||
Some(key) => (key, existing_key.name.clone()),
|
||||
None => {
|
||||
.await;
|
||||
match update_result {
|
||||
Err(error) => {
|
||||
release_codex_oauth_account_locks(state, codex_oauth_account_leases).await;
|
||||
return Err(error);
|
||||
}
|
||||
Ok(Some(key)) => (key, existing_key.name.clone()),
|
||||
Ok(None) => {
|
||||
release_codex_oauth_account_locks(state, codex_oauth_account_leases).await;
|
||||
failed += 1;
|
||||
results.push(json!({
|
||||
"index": index,
|
||||
@@ -611,7 +655,7 @@ pub(super) async fn execute_admin_provider_oauth_batch_import(
|
||||
&auth_config,
|
||||
Some(index),
|
||||
);
|
||||
match create_provider_oauth_catalog_key(
|
||||
let create_result = create_provider_oauth_catalog_key(
|
||||
state,
|
||||
provider_id,
|
||||
provider_type,
|
||||
@@ -622,10 +666,15 @@ pub(super) async fn execute_admin_provider_oauth_batch_import(
|
||||
key_proxy.clone(),
|
||||
expires_at,
|
||||
)
|
||||
.await?
|
||||
{
|
||||
Some(key) => (key, key_name),
|
||||
None => {
|
||||
.await;
|
||||
match create_result {
|
||||
Err(error) => {
|
||||
release_codex_oauth_account_locks(state, codex_oauth_account_leases).await;
|
||||
return Err(error);
|
||||
}
|
||||
Ok(Some(key)) => (key, key_name),
|
||||
Ok(None) => {
|
||||
release_codex_oauth_account_locks(state, codex_oauth_account_leases).await;
|
||||
failed += 1;
|
||||
results.push(json!({
|
||||
"index": index,
|
||||
@@ -645,6 +694,7 @@ pub(super) async fn execute_admin_provider_oauth_batch_import(
|
||||
}
|
||||
}
|
||||
};
|
||||
release_codex_oauth_account_locks(state, codex_oauth_account_leases).await;
|
||||
|
||||
spawn_provider_oauth_account_state_refresh_after_update(
|
||||
state.cloned_app(),
|
||||
@@ -726,13 +776,24 @@ mod tests {
|
||||
"credentials":{
|
||||
"auth_mode":"agentIdentity",
|
||||
"id_token":"stale-id-token",
|
||||
"sessionToken":"stale-session-token",
|
||||
"apiKey":"stale-api-key",
|
||||
"cookie":"stale-cookie",
|
||||
"headers":{
|
||||
"authorization":"Bearer stale-bearer-token",
|
||||
"cookie":"stale-header-cookie",
|
||||
"x-api-key":"stale-header-api-key"
|
||||
},
|
||||
"profile":{"token":"stale-deep-token"},
|
||||
"agent_identity":{
|
||||
"agent_runtime_id":"runtime-1",
|
||||
"agent_private_key":"MC4CAQAwBQYDK2VwBCIEIAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA",
|
||||
"accountId":"account-1",
|
||||
"chatgptUserId":"user-1",
|
||||
"chatgptAccountIsFedramp":true,
|
||||
"access_token":"stale-access-token"
|
||||
"access_token":"stale-access-token",
|
||||
"refreshToken":"stale-refresh-token",
|
||||
"headers":{"authorization":"Bearer stale-nested-bearer"}
|
||||
}
|
||||
}
|
||||
}]
|
||||
@@ -761,5 +822,35 @@ mod tests {
|
||||
.get("agent_identity")
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.is_some_and(|nested| !nested.contains_key("access_token")));
|
||||
let serialized = serde_json::to_string(&auth_config).expect("auth config should serialize");
|
||||
for secret in [
|
||||
"stale-id-token",
|
||||
"stale-session-token",
|
||||
"stale-api-key",
|
||||
"stale-cookie",
|
||||
"stale-bearer-token",
|
||||
"stale-header-cookie",
|
||||
"stale-header-api-key",
|
||||
"stale-deep-token",
|
||||
"stale-access-token",
|
||||
"stale-refresh-token",
|
||||
"stale-nested-bearer",
|
||||
] {
|
||||
assert!(!serialized.contains(secret), "secret leaked: {secret}");
|
||||
}
|
||||
for forbidden_key in [
|
||||
"access_token",
|
||||
"refreshToken",
|
||||
"sessionToken",
|
||||
"apiKey",
|
||||
"cookie",
|
||||
"headers",
|
||||
"profile",
|
||||
] {
|
||||
assert!(
|
||||
!serialized.contains(forbidden_key),
|
||||
"forbidden key persisted: {forbidden_key}"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -6,4 +6,7 @@ mod progress;
|
||||
mod task;
|
||||
|
||||
pub(super) use orchestration::handle_admin_provider_oauth_batch_import;
|
||||
pub(super) use task::handle_admin_provider_oauth_start_batch_import_task;
|
||||
pub(super) use task::{
|
||||
handle_admin_provider_oauth_start_agent_identity_import_task,
|
||||
handle_admin_provider_oauth_start_batch_import_task,
|
||||
};
|
||||
|
||||
@@ -3,6 +3,7 @@ use super::execution::{
|
||||
execute_admin_provider_oauth_batch_import_for_provider_type,
|
||||
};
|
||||
use super::parse::{
|
||||
admin_provider_oauth_batch_contains_agent_identity,
|
||||
build_admin_provider_oauth_batch_import_response,
|
||||
parse_admin_provider_oauth_batch_import_request, AdminProviderOAuthBatchImportRequest,
|
||||
};
|
||||
@@ -41,6 +42,12 @@ pub(in super::super) async fn handle_admin_provider_oauth_batch_import(
|
||||
Ok(payload) => payload,
|
||||
Err(response) => return Ok(response),
|
||||
};
|
||||
if admin_provider_oauth_batch_contains_agent_identity(&payload.credentials) {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"Agent Identity JSON 必须使用专属导入接口",
|
||||
));
|
||||
}
|
||||
|
||||
let Some(provider) = state
|
||||
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
|
||||
|
||||
@@ -78,6 +78,36 @@ pub(super) fn parse_admin_provider_oauth_batch_import_request(
|
||||
}
|
||||
}
|
||||
|
||||
fn json_value_contains_agent_identity(value: &serde_json::Value) -> bool {
|
||||
if value.as_object().is_some_and(|_| {
|
||||
aether_provider_transport::is_codex_agent_identity_auth_config_value(value)
|
||||
}) {
|
||||
return true;
|
||||
}
|
||||
match value {
|
||||
serde_json::Value::Array(items) => items.iter().any(json_value_contains_agent_identity),
|
||||
serde_json::Value::Object(object) => {
|
||||
object.values().any(json_value_contains_agent_identity)
|
||||
}
|
||||
_ => false,
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn admin_provider_oauth_batch_contains_agent_identity(raw_credentials: &str) -> bool {
|
||||
let raw = raw_credentials.trim();
|
||||
if raw.is_empty() {
|
||||
return false;
|
||||
}
|
||||
if let Ok(value) = serde_json::from_str::<serde_json::Value>(raw) {
|
||||
return json_value_contains_agent_identity(&value);
|
||||
}
|
||||
raw.lines()
|
||||
.map(str::trim)
|
||||
.filter(|line| !line.is_empty() && !line.starts_with('#'))
|
||||
.filter_map(|line| serde_json::from_str::<serde_json::Value>(line).ok())
|
||||
.any(|value| json_value_contains_agent_identity(&value))
|
||||
}
|
||||
|
||||
fn coerce_admin_provider_oauth_import_str(value: Option<&serde_json::Value>) -> Option<String> {
|
||||
value
|
||||
.and_then(serde_json::Value::as_str)
|
||||
@@ -623,6 +653,47 @@ pub(super) fn parse_admin_provider_oauth_batch_import_entries(
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub(super) fn parse_admin_provider_oauth_agent_identity_import_entries(
|
||||
raw_credentials: &str,
|
||||
) -> Result<Vec<AdminProviderOAuthBatchImportEntry>, String> {
|
||||
let raw = raw_credentials.trim();
|
||||
if raw.is_empty() {
|
||||
return Err("Agent Identity 凭据不能为空".to_string());
|
||||
}
|
||||
let value = serde_json::from_str::<serde_json::Value>(raw)
|
||||
.map_err(|error| format!("Agent Identity JSON 解析失败: {error}"))?;
|
||||
let entries = match &value {
|
||||
serde_json::Value::Array(items) => items
|
||||
.iter()
|
||||
.enumerate()
|
||||
.map(|(index, item)| {
|
||||
extract_admin_provider_oauth_batch_import_entry("codex", item).unwrap_or_else(
|
||||
|| {
|
||||
parse_error_entry(format!(
|
||||
"第 {} 个条目没有可导入的 Agent Identity 凭据",
|
||||
index + 1
|
||||
))
|
||||
},
|
||||
)
|
||||
})
|
||||
.collect(),
|
||||
serde_json::Value::Object(object) => parse_sub2api_export_accounts("codex", object)
|
||||
.unwrap_or_else(|| {
|
||||
vec![
|
||||
extract_admin_provider_oauth_batch_import_entry("codex", &value)
|
||||
.unwrap_or_else(|| {
|
||||
parse_error_entry("没有可导入的 Agent Identity 凭据".to_string())
|
||||
}),
|
||||
]
|
||||
}),
|
||||
_ => return Err("Agent Identity 凭据必须是 JSON 对象、数组或 sub2api 导出".to_string()),
|
||||
};
|
||||
if entries.is_empty() {
|
||||
return Err("Agent Identity 凭据不能为空".to_string());
|
||||
}
|
||||
Ok(entries)
|
||||
}
|
||||
|
||||
fn parse_error_entry(error: String) -> AdminProviderOAuthBatchImportEntry {
|
||||
AdminProviderOAuthBatchImportEntry {
|
||||
parse_error: Some(error),
|
||||
@@ -811,6 +882,7 @@ pub(super) fn build_admin_provider_oauth_batch_task_state(
|
||||
task_id: &str,
|
||||
provider_id: &str,
|
||||
provider_type: &str,
|
||||
import_kind: &str,
|
||||
status: &str,
|
||||
total: usize,
|
||||
processed: usize,
|
||||
@@ -839,6 +911,7 @@ pub(super) fn build_admin_provider_oauth_batch_task_state(
|
||||
"task_id": task_id,
|
||||
"provider_id": provider_id,
|
||||
"provider_type": provider_type,
|
||||
"import_kind": import_kind,
|
||||
"status": status,
|
||||
"total": total,
|
||||
"processed": processed,
|
||||
@@ -860,6 +933,7 @@ pub(super) fn build_admin_provider_oauth_batch_task_state(
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
admin_provider_oauth_batch_contains_agent_identity,
|
||||
apply_admin_provider_oauth_batch_import_hints,
|
||||
parse_admin_provider_oauth_batch_import_entries,
|
||||
};
|
||||
@@ -874,6 +948,37 @@ mod tests {
|
||||
format!("{}.{}.signature", encode(header), encode(payload))
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn ordinary_batch_guard_detects_agent_identity_in_all_json_shapes() {
|
||||
let single = json!({
|
||||
"auth_mode": "agentIdentity",
|
||||
"agent_runtime_id": "runtime-1",
|
||||
"agent_private_key": "private-key"
|
||||
});
|
||||
assert!(admin_provider_oauth_batch_contains_agent_identity(
|
||||
&single.to_string()
|
||||
));
|
||||
assert!(admin_provider_oauth_batch_contains_agent_identity(
|
||||
&json!([{"refresh_token":"ordinary"}, single.clone()]).to_string()
|
||||
));
|
||||
assert!(admin_provider_oauth_batch_contains_agent_identity(
|
||||
&format!("ordinary-token\n{}", single)
|
||||
));
|
||||
assert!(admin_provider_oauth_batch_contains_agent_identity(
|
||||
&json!({
|
||||
"type": "sub2api-data",
|
||||
"accounts": [{"credentials": single.clone()}]
|
||||
})
|
||||
.to_string()
|
||||
));
|
||||
assert!(!admin_provider_oauth_batch_contains_agent_identity(
|
||||
&json!([{"refresh_token":"ordinary"}]).to_string()
|
||||
));
|
||||
assert!(!admin_provider_oauth_batch_contains_agent_identity(
|
||||
"ordinary-token"
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_access_token_only_entry() {
|
||||
let entries = parse_admin_provider_oauth_batch_import_entries(
|
||||
|
||||
@@ -1,19 +1,26 @@
|
||||
use super::execution::{
|
||||
estimate_admin_provider_oauth_batch_import_total,
|
||||
codex_agent_identity_auth_config_from_import, estimate_admin_provider_oauth_batch_import_total,
|
||||
execute_admin_provider_oauth_batch_import_for_provider_type,
|
||||
};
|
||||
use super::parse::{
|
||||
build_admin_provider_oauth_batch_task_state, parse_admin_provider_oauth_batch_import_request,
|
||||
admin_provider_oauth_batch_contains_agent_identity,
|
||||
build_admin_provider_oauth_batch_task_state,
|
||||
parse_admin_provider_oauth_agent_identity_import_entries,
|
||||
parse_admin_provider_oauth_batch_import_request,
|
||||
};
|
||||
use super::progress::{
|
||||
AdminProviderOAuthBatchImportProgress, AdminProviderOAuthBatchProgressReporter,
|
||||
};
|
||||
use crate::handlers::admin::provider::oauth::duplicates::codex_agent_identity_account_lock_keys;
|
||||
use crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response;
|
||||
use crate::handlers::admin::provider::oauth::state::{
|
||||
admin_provider_oauth_template, build_admin_provider_oauth_backend_unavailable_response,
|
||||
is_fixed_provider_type_for_provider_oauth,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_batch_import_task_provider_id;
|
||||
use crate::handlers::admin::provider::shared::paths::{
|
||||
admin_provider_oauth_agent_identity_import_task_provider_id,
|
||||
admin_provider_oauth_batch_import_task_provider_id,
|
||||
};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::task_runtime::{
|
||||
append_event_with_logging, now_unix_secs, task_definition, update_run_status,
|
||||
@@ -23,24 +30,177 @@ use crate::GatewayError;
|
||||
use aether_data_contracts::repository::background_tasks::{
|
||||
BackgroundTaskKind, BackgroundTaskStatus, UpsertBackgroundTaskRun,
|
||||
};
|
||||
use aether_runtime_state::RuntimeLockLease;
|
||||
use axum::{
|
||||
body::Bytes,
|
||||
http,
|
||||
response::{IntoResponse, Response},
|
||||
Json,
|
||||
};
|
||||
use serde_json::json;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
use serde_json::{json, Map, Value};
|
||||
use sha2::{Digest, Sha256};
|
||||
use std::time::{Duration, SystemTime, UNIX_EPOCH};
|
||||
use tokio::task;
|
||||
use uuid::Uuid;
|
||||
|
||||
const PROVIDER_OAUTH_BATCH_TASK_MAX_ERROR_SAMPLES: usize = 20;
|
||||
const PROVIDER_OAUTH_BATCH_IMPORT_KIND: &str = "oauth_batch";
|
||||
const PROVIDER_AGENT_IDENTITY_IMPORT_KIND: &str = "agent_identity";
|
||||
const PROVIDER_AGENT_IDENTITY_IMPORT_LOCK_TTL: Duration = Duration::from_secs(180);
|
||||
const PROVIDER_AGENT_IDENTITY_IMPORT_LOCK_RENEW_INTERVAL: Duration = Duration::from_secs(60);
|
||||
|
||||
fn provider_oauth_import_kind(agent_identity_only: bool) -> &'static str {
|
||||
if agent_identity_only {
|
||||
PROVIDER_AGENT_IDENTITY_IMPORT_KIND
|
||||
} else {
|
||||
PROVIDER_OAUTH_BATCH_IMPORT_KIND
|
||||
}
|
||||
}
|
||||
|
||||
fn codex_agent_identity_import_auth_configs(
|
||||
credentials: &str,
|
||||
) -> Result<Vec<Map<String, Value>>, String> {
|
||||
let entries = parse_admin_provider_oauth_agent_identity_import_entries(credentials)?;
|
||||
entries
|
||||
.iter()
|
||||
.enumerate()
|
||||
.map(|(index, entry)| {
|
||||
if let Some(error) = entry.parse_error.as_deref() {
|
||||
return Err(format!("第 {} 个条目无效: {error}", index + 1));
|
||||
}
|
||||
match codex_agent_identity_auth_config_from_import(entry) {
|
||||
Ok(Some(auth_config)) => Ok(auth_config),
|
||||
Ok(None) => Err(format!("第 {} 个条目不是 Agent Identity", index + 1)),
|
||||
Err(error) => Err(format!("第 {} 个条目无效: {error}", index + 1)),
|
||||
}
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn provider_agent_identity_import_lock_key(provider_id: &str, agent_runtime_id: &str) -> String {
|
||||
let digest = Sha256::digest(format!("{provider_id}\0{agent_runtime_id}").as_bytes());
|
||||
format!("provider_oauth_agent_identity_import:{digest:x}")
|
||||
}
|
||||
|
||||
fn provider_agent_identity_import_lock_keys(
|
||||
provider_id: &str,
|
||||
auth_configs: &[Map<String, Value>],
|
||||
) -> Vec<String> {
|
||||
let mut lock_keys = Vec::new();
|
||||
for auth_config in auth_configs {
|
||||
if let Some(agent_runtime_id) = auth_config
|
||||
.get("agent_runtime_id")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
lock_keys.push(provider_agent_identity_import_lock_key(
|
||||
provider_id,
|
||||
agent_runtime_id,
|
||||
));
|
||||
}
|
||||
lock_keys.extend(codex_agent_identity_account_lock_keys(
|
||||
provider_id,
|
||||
auth_config,
|
||||
));
|
||||
}
|
||||
lock_keys.sort_unstable();
|
||||
lock_keys.dedup();
|
||||
lock_keys
|
||||
}
|
||||
|
||||
async fn acquire_provider_agent_identity_import_locks(
|
||||
state: &AdminAppState<'_>,
|
||||
provider_id: &str,
|
||||
lock_keys: &[String],
|
||||
task_id: &str,
|
||||
) -> Result<Vec<RuntimeLockLease>, Response> {
|
||||
let owner = format!("aether-gateway-agent-identity-import-{task_id}");
|
||||
let mut leases = Vec::with_capacity(lock_keys.len());
|
||||
for lock_key in lock_keys {
|
||||
match state
|
||||
.runtime_state()
|
||||
.lock_try_acquire(
|
||||
lock_key.as_str(),
|
||||
owner.as_str(),
|
||||
PROVIDER_AGENT_IDENTITY_IMPORT_LOCK_TTL,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(Some(lease)) => leases.push(lease),
|
||||
Ok(None) => {
|
||||
release_provider_agent_identity_import_locks(state, leases).await;
|
||||
return Err(build_internal_control_error_response(
|
||||
http::StatusCode::CONFLICT,
|
||||
"其中一个 Agent Identity 正在导入或创建,请稍后重试",
|
||||
));
|
||||
}
|
||||
Err(error) => {
|
||||
tracing::warn!(
|
||||
provider_id = %provider_id,
|
||||
lock_key = %lock_key,
|
||||
error = ?error,
|
||||
"gateway Agent Identity import lock unavailable"
|
||||
);
|
||||
release_provider_agent_identity_import_locks(state, leases).await;
|
||||
return Err(build_internal_control_error_response(
|
||||
http::StatusCode::SERVICE_UNAVAILABLE,
|
||||
"Agent Identity 导入锁暂不可用,请稍后重试",
|
||||
));
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(leases)
|
||||
}
|
||||
|
||||
async fn release_provider_agent_identity_import_locks(
|
||||
state: &AdminAppState<'_>,
|
||||
leases: Vec<RuntimeLockLease>,
|
||||
) {
|
||||
for lease in leases {
|
||||
match state.runtime_state().lock_release(&lease).await {
|
||||
Ok(true) => {}
|
||||
Ok(false) => tracing::warn!(
|
||||
lock_key = %lease.key,
|
||||
"gateway Agent Identity import lock was not owned during release"
|
||||
),
|
||||
Err(error) => tracing::warn!(
|
||||
lock_key = %lease.key,
|
||||
error = ?error,
|
||||
"gateway Agent Identity import lock release failed"
|
||||
),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn renew_provider_agent_identity_import_locks(
|
||||
state: &AdminAppState<'_>,
|
||||
leases: &[RuntimeLockLease],
|
||||
ttl: Duration,
|
||||
) -> Result<(), String> {
|
||||
for lease in leases {
|
||||
match state.runtime_state().lock_renew(lease, ttl).await {
|
||||
Ok(true) => {}
|
||||
Ok(false) => {
|
||||
return Err(format!("Agent Identity 导入锁已失效: {}", lease.key));
|
||||
}
|
||||
Err(error) => {
|
||||
return Err(format!(
|
||||
"Agent Identity 导入锁续租失败 ({}): {error:?}",
|
||||
lease.key
|
||||
));
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
struct BatchTaskProgressReporter {
|
||||
app: crate::AppState,
|
||||
task_id: String,
|
||||
provider_id: String,
|
||||
provider_type: String,
|
||||
import_kind: &'static str,
|
||||
created_at: u64,
|
||||
started_at: u64,
|
||||
error_samples: Vec<serde_json::Value>,
|
||||
@@ -65,6 +225,7 @@ impl AdminProviderOAuthBatchProgressReporter for BatchTaskProgressReporter {
|
||||
&self.task_id,
|
||||
&self.provider_id,
|
||||
&self.provider_type,
|
||||
self.import_kind,
|
||||
"processing",
|
||||
progress.total,
|
||||
progress.processed,
|
||||
@@ -89,13 +250,33 @@ pub(in super::super) async fn handle_admin_provider_oauth_start_batch_import_tas
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
request_body: Option<&Bytes>,
|
||||
) -> Result<Response, GatewayError> {
|
||||
handle_admin_provider_oauth_start_import_task(state, request_context, request_body, false).await
|
||||
}
|
||||
|
||||
pub(in super::super) async fn handle_admin_provider_oauth_start_agent_identity_import_task(
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
request_body: Option<&Bytes>,
|
||||
) -> Result<Response, GatewayError> {
|
||||
handle_admin_provider_oauth_start_import_task(state, request_context, request_body, true).await
|
||||
}
|
||||
|
||||
async fn handle_admin_provider_oauth_start_import_task(
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
request_body: Option<&Bytes>,
|
||||
agent_identity_only: bool,
|
||||
) -> Result<Response, GatewayError> {
|
||||
if !state.has_provider_catalog_data_reader() {
|
||||
return Ok(build_admin_provider_oauth_backend_unavailable_response());
|
||||
}
|
||||
let Some(provider_id) =
|
||||
let provider_id = if agent_identity_only {
|
||||
admin_provider_oauth_agent_identity_import_task_provider_id(request_context.path())
|
||||
} else {
|
||||
admin_provider_oauth_batch_import_task_provider_id(request_context.path())
|
||||
else {
|
||||
};
|
||||
let Some(provider_id) = provider_id else {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
"Provider 不存在",
|
||||
@@ -105,6 +286,14 @@ pub(in super::super) async fn handle_admin_provider_oauth_start_batch_import_tas
|
||||
Ok(payload) => payload,
|
||||
Err(response) => return Ok(response),
|
||||
};
|
||||
if !agent_identity_only
|
||||
&& admin_provider_oauth_batch_contains_agent_identity(&payload.credentials)
|
||||
{
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"Agent Identity JSON 必须使用专属导入接口",
|
||||
));
|
||||
}
|
||||
|
||||
let Some(provider) = state
|
||||
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
|
||||
@@ -118,6 +307,12 @@ pub(in super::super) async fn handle_admin_provider_oauth_start_batch_import_tas
|
||||
));
|
||||
};
|
||||
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
|
||||
if agent_identity_only && provider_type != "codex" {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"仅 Codex Provider 支持导入 Agent Identity",
|
||||
));
|
||||
}
|
||||
if !is_fixed_provider_type_for_provider_oauth(&provider_type) {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
@@ -131,9 +326,27 @@ pub(in super::super) async fn handle_admin_provider_oauth_start_batch_import_tas
|
||||
return Ok(build_admin_provider_oauth_backend_unavailable_response());
|
||||
}
|
||||
|
||||
let total = estimate_admin_provider_oauth_batch_import_total(
|
||||
&provider_type,
|
||||
payload.credentials.as_str(),
|
||||
let agent_identity_auth_configs = if agent_identity_only {
|
||||
match codex_agent_identity_import_auth_configs(&payload.credentials) {
|
||||
Ok(auth_configs) => Some(auth_configs),
|
||||
Err(detail) => {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
format!("该接口仅接受有效的 Agent Identity JSON: {detail}"),
|
||||
));
|
||||
}
|
||||
}
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let total = agent_identity_auth_configs.as_ref().map_or_else(
|
||||
|| {
|
||||
estimate_admin_provider_oauth_batch_import_total(
|
||||
&provider_type,
|
||||
payload.credentials.as_str(),
|
||||
)
|
||||
},
|
||||
Vec::len,
|
||||
);
|
||||
if total == 0 {
|
||||
return Ok(build_internal_control_error_response(
|
||||
@@ -142,12 +355,35 @@ pub(in super::super) async fn handle_admin_provider_oauth_start_batch_import_tas
|
||||
));
|
||||
}
|
||||
|
||||
let task_id = Uuid::new_v4().to_string();
|
||||
let task_id = if agent_identity_only {
|
||||
format!("agent-identity-{}", Uuid::new_v4())
|
||||
} else {
|
||||
Uuid::new_v4().to_string()
|
||||
};
|
||||
let mut agent_identity_import_leases =
|
||||
if let Some(auth_configs) = agent_identity_auth_configs.as_deref() {
|
||||
let lock_keys = provider_agent_identity_import_lock_keys(&provider_id, auth_configs);
|
||||
match acquire_provider_agent_identity_import_locks(
|
||||
state,
|
||||
&provider_id,
|
||||
&lock_keys,
|
||||
&task_id,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(leases) => leases,
|
||||
Err(response) => return Ok(response),
|
||||
}
|
||||
} else {
|
||||
Vec::new()
|
||||
};
|
||||
let import_kind = provider_oauth_import_kind(agent_identity_only);
|
||||
let created_at = now_unix_secs();
|
||||
let submitted_state = build_admin_provider_oauth_batch_task_state(
|
||||
&task_id,
|
||||
&provider_id,
|
||||
&provider_type,
|
||||
import_kind,
|
||||
"submitted",
|
||||
total,
|
||||
0,
|
||||
@@ -167,6 +403,11 @@ pub(in super::super) async fn handle_admin_provider_oauth_start_batch_import_tas
|
||||
.await
|
||||
.is_err()
|
||||
{
|
||||
release_provider_agent_identity_import_locks(
|
||||
state,
|
||||
std::mem::take(&mut agent_identity_import_leases),
|
||||
)
|
||||
.await;
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::SERVICE_UNAVAILABLE,
|
||||
"provider oauth batch task redis unavailable",
|
||||
@@ -191,6 +432,7 @@ pub(in super::super) async fn handle_admin_provider_oauth_start_batch_import_tas
|
||||
payload_json: Some(json!({
|
||||
"provider_id": provider_id.clone(),
|
||||
"provider_type": provider_type.clone(),
|
||||
"import_kind": import_kind,
|
||||
"total": total,
|
||||
})),
|
||||
result_json: None,
|
||||
@@ -211,6 +453,7 @@ pub(in super::super) async fn handle_admin_provider_oauth_start_batch_import_tas
|
||||
Some(json!({
|
||||
"provider_id": provider_id.clone(),
|
||||
"provider_type": provider_type.clone(),
|
||||
"import_kind": import_kind,
|
||||
"total": total,
|
||||
})),
|
||||
)
|
||||
@@ -223,6 +466,7 @@ pub(in super::super) async fn handle_admin_provider_oauth_start_batch_import_tas
|
||||
let provider_type_for_worker = provider_type.clone();
|
||||
let proxy_node_id = payload.proxy_node_id.clone();
|
||||
let raw_credentials = payload.credentials.clone();
|
||||
let agent_identity_import_leases_for_worker = std::mem::take(&mut agent_identity_import_leases);
|
||||
task::spawn(async move {
|
||||
let started_at = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
@@ -233,6 +477,7 @@ pub(in super::super) async fn handle_admin_provider_oauth_start_batch_import_tas
|
||||
&task_id_for_worker,
|
||||
&provider_id_for_worker,
|
||||
&provider_type_for_worker,
|
||||
import_kind,
|
||||
"processing",
|
||||
total,
|
||||
0,
|
||||
@@ -277,20 +522,55 @@ pub(in super::super) async fn handle_admin_provider_oauth_start_batch_import_tas
|
||||
task_id: task_id_for_worker.clone(),
|
||||
provider_id: provider_id_for_worker.clone(),
|
||||
provider_type: provider_type_for_worker.clone(),
|
||||
import_kind,
|
||||
created_at,
|
||||
started_at,
|
||||
error_samples: Vec::new(),
|
||||
};
|
||||
match execute_admin_provider_oauth_batch_import_for_provider_type(
|
||||
&AdminAppState::new(&task_state),
|
||||
let task_admin_state = AdminAppState::new(&task_state);
|
||||
let execution = execute_admin_provider_oauth_batch_import_for_provider_type(
|
||||
&task_admin_state,
|
||||
&provider_id_for_worker,
|
||||
&provider_type_for_worker,
|
||||
raw_credentials.as_str(),
|
||||
proxy_node_id.as_deref(),
|
||||
Some(&mut progress_reporter),
|
||||
);
|
||||
tokio::pin!(execution);
|
||||
let execution_result = if agent_identity_import_leases_for_worker.is_empty() {
|
||||
execution.await
|
||||
} else {
|
||||
let mut renew_timer =
|
||||
tokio::time::interval(PROVIDER_AGENT_IDENTITY_IMPORT_LOCK_RENEW_INTERVAL);
|
||||
renew_timer.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
|
||||
renew_timer.tick().await;
|
||||
loop {
|
||||
tokio::select! {
|
||||
result = &mut execution => break result,
|
||||
_ = renew_timer.tick() => {
|
||||
if let Err(detail) = renew_provider_agent_identity_import_locks(
|
||||
&task_admin_state,
|
||||
&agent_identity_import_leases_for_worker,
|
||||
PROVIDER_AGENT_IDENTITY_IMPORT_LOCK_TTL,
|
||||
).await {
|
||||
tracing::error!(
|
||||
provider_id = %provider_id_for_worker,
|
||||
task_id = %task_id_for_worker,
|
||||
detail = %detail,
|
||||
"gateway Agent Identity import lease lost"
|
||||
);
|
||||
break Err(GatewayError::Internal(detail));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
};
|
||||
release_provider_agent_identity_import_locks(
|
||||
&task_admin_state,
|
||||
agent_identity_import_leases_for_worker,
|
||||
)
|
||||
.await
|
||||
{
|
||||
.await;
|
||||
match execution_result {
|
||||
Ok(outcome) => {
|
||||
let finished_at = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
@@ -324,6 +604,7 @@ pub(in super::super) async fn handle_admin_provider_oauth_start_batch_import_tas
|
||||
&task_id_for_worker,
|
||||
&provider_id_for_worker,
|
||||
&provider_type_for_worker,
|
||||
import_kind,
|
||||
"completed",
|
||||
outcome.total,
|
||||
outcome.total,
|
||||
@@ -350,6 +631,7 @@ pub(in super::super) async fn handle_admin_provider_oauth_start_batch_import_tas
|
||||
Some(json!({
|
||||
"provider_id": provider_id_for_worker,
|
||||
"provider_type": provider_type_for_worker,
|
||||
"import_kind": import_kind,
|
||||
"total": outcome.total,
|
||||
"success": outcome.success,
|
||||
"failed": outcome.failed,
|
||||
@@ -381,6 +663,7 @@ pub(in super::super) async fn handle_admin_provider_oauth_start_batch_import_tas
|
||||
&task_id_for_worker,
|
||||
&provider_id_for_worker,
|
||||
&provider_type_for_worker,
|
||||
import_kind,
|
||||
"failed",
|
||||
total,
|
||||
0,
|
||||
@@ -432,6 +715,7 @@ pub(in super::super) async fn handle_admin_provider_oauth_start_batch_import_tas
|
||||
&task_id,
|
||||
&provider_id,
|
||||
&provider_type,
|
||||
import_kind,
|
||||
"submitted",
|
||||
total,
|
||||
0,
|
||||
@@ -448,3 +732,251 @@ pub(in super::super) async fn handle_admin_provider_oauth_start_batch_import_tas
|
||||
);
|
||||
Ok(Json(submitted_response).into_response())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
acquire_provider_agent_identity_import_locks, codex_agent_identity_import_auth_configs,
|
||||
provider_agent_identity_import_lock_key, provider_agent_identity_import_lock_keys,
|
||||
release_provider_agent_identity_import_locks, renew_provider_agent_identity_import_locks,
|
||||
};
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::AppState;
|
||||
use serde_json::json;
|
||||
|
||||
fn agent_identity(runtime_id: &str) -> serde_json::Value {
|
||||
json!({
|
||||
"auth_mode": "agentIdentity",
|
||||
"agent_runtime_id": runtime_id,
|
||||
"agent_private_key": "MC4CAQAwBQYDK2VwBCIEIAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA",
|
||||
"task_id": format!("task-{runtime_id}"),
|
||||
"account_id": "account-1",
|
||||
"user_id": "user-1",
|
||||
"email": "[email protected]"
|
||||
})
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn dedicated_import_accepts_root_array_and_sub2api_agent_identities() {
|
||||
let single = agent_identity("runtime-1");
|
||||
assert_eq!(
|
||||
codex_agent_identity_import_auth_configs(&single.to_string())
|
||||
.expect("single Agent Identity should parse")
|
||||
.len(),
|
||||
1
|
||||
);
|
||||
|
||||
let array = json!([agent_identity("runtime-1"), agent_identity("runtime-2")]);
|
||||
assert_eq!(
|
||||
codex_agent_identity_import_auth_configs(&array.to_string())
|
||||
.expect("Agent Identity array should parse")
|
||||
.len(),
|
||||
2
|
||||
);
|
||||
|
||||
let sub2api = json!({
|
||||
"type": "sub2api-data",
|
||||
"accounts": [
|
||||
{
|
||||
"name": "[email protected]",
|
||||
"platform": "openai",
|
||||
"credentials": agent_identity("runtime-1")
|
||||
},
|
||||
{
|
||||
"name": "[email protected]",
|
||||
"platform": "openai",
|
||||
"credentials": agent_identity("runtime-2")
|
||||
},
|
||||
{
|
||||
"name": "[email protected]",
|
||||
"platform": "anthropic",
|
||||
"credentials": { "access_token": "ignored" }
|
||||
}
|
||||
]
|
||||
});
|
||||
assert_eq!(
|
||||
codex_agent_identity_import_auth_configs(&sub2api.to_string())
|
||||
.expect("sub2api Agent Identity export should parse")
|
||||
.len(),
|
||||
2
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn dedicated_import_rejects_mixed_and_invalid_entries() {
|
||||
let mixed = json!([
|
||||
agent_identity("runtime-1"),
|
||||
{ "refresh_token": "refresh-token" }
|
||||
]);
|
||||
assert!(codex_agent_identity_import_auth_configs(&mixed.to_string()).is_err());
|
||||
|
||||
let invalid = json!({
|
||||
"auth_mode": "agentIdentity",
|
||||
"agent_runtime_id": "runtime-invalid",
|
||||
"agent_private_key": "not-a-pkcs8-key"
|
||||
});
|
||||
assert!(codex_agent_identity_import_auth_configs(&invalid.to_string()).is_err());
|
||||
assert!(codex_agent_identity_import_auth_configs("[]").is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn agent_identity_import_lock_key_is_scoped_to_provider_and_runtime() {
|
||||
let first = provider_agent_identity_import_lock_key("provider-a", "runtime-1");
|
||||
assert_eq!(
|
||||
first,
|
||||
provider_agent_identity_import_lock_key("provider-a", "runtime-1")
|
||||
);
|
||||
assert_ne!(
|
||||
first,
|
||||
provider_agent_identity_import_lock_key("provider-b", "runtime-1")
|
||||
);
|
||||
assert_ne!(
|
||||
first,
|
||||
provider_agent_identity_import_lock_key("provider-a", "runtime-2")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn agent_identity_imports_with_distinct_runtimes_share_account_lock_keys() {
|
||||
let first =
|
||||
codex_agent_identity_import_auth_configs(&agent_identity("runtime-1").to_string())
|
||||
.expect("first Agent Identity should parse");
|
||||
let second =
|
||||
codex_agent_identity_import_auth_configs(&agent_identity("runtime-2").to_string())
|
||||
.expect("second Agent Identity should parse");
|
||||
let first_keys = provider_agent_identity_import_lock_keys("provider-a", &first);
|
||||
let second_keys = provider_agent_identity_import_lock_keys("provider-a", &second);
|
||||
|
||||
assert!(first_keys.iter().any(|key| second_keys.contains(key)));
|
||||
assert!(
|
||||
first_keys.contains(&provider_agent_identity_import_lock_key(
|
||||
"provider-a",
|
||||
"runtime-1"
|
||||
))
|
||||
);
|
||||
assert!(
|
||||
second_keys.contains(&provider_agent_identity_import_lock_key(
|
||||
"provider-a",
|
||||
"runtime-2"
|
||||
))
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn agent_identity_import_runtime_lock_releases_after_contention() {
|
||||
let app = AppState::new().expect("app state should build");
|
||||
let state = AdminAppState::new(&app);
|
||||
let lock_keys = vec![provider_agent_identity_import_lock_key(
|
||||
"provider-a",
|
||||
"runtime-1",
|
||||
)];
|
||||
let first = acquire_provider_agent_identity_import_locks(
|
||||
&state,
|
||||
"provider-a",
|
||||
&lock_keys,
|
||||
"agent-identity-task-1",
|
||||
)
|
||||
.await
|
||||
.expect("first lock should acquire");
|
||||
let second = acquire_provider_agent_identity_import_locks(
|
||||
&state,
|
||||
"provider-a",
|
||||
&lock_keys,
|
||||
"agent-identity-task-2",
|
||||
)
|
||||
.await
|
||||
.expect_err("second lock should be rejected");
|
||||
assert_eq!(second.status(), axum::http::StatusCode::CONFLICT);
|
||||
release_provider_agent_identity_import_locks(&state, first).await;
|
||||
let third = acquire_provider_agent_identity_import_locks(
|
||||
&state,
|
||||
"provider-a",
|
||||
&lock_keys,
|
||||
"agent-identity-task-3",
|
||||
)
|
||||
.await
|
||||
.expect("lock should be reusable after release");
|
||||
release_provider_agent_identity_import_locks(&state, third).await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn agent_identity_import_partial_lock_failure_releases_acquired_leases() {
|
||||
let app = AppState::new().expect("app state should build");
|
||||
let state = AdminAppState::new(&app);
|
||||
let held = state
|
||||
.runtime_state()
|
||||
.lock_try_acquire("z-held", "other-task", std::time::Duration::from_secs(30))
|
||||
.await
|
||||
.expect("runtime lock should be available")
|
||||
.expect("held lock should acquire");
|
||||
let lock_keys = vec!["a-free".to_string(), "z-held".to_string()];
|
||||
|
||||
let response = acquire_provider_agent_identity_import_locks(
|
||||
&state,
|
||||
"provider-a",
|
||||
&lock_keys,
|
||||
"agent-identity-task-partial",
|
||||
)
|
||||
.await
|
||||
.expect_err("second lock should cause contention");
|
||||
assert_eq!(response.status(), axum::http::StatusCode::CONFLICT);
|
||||
|
||||
let free = state
|
||||
.runtime_state()
|
||||
.lock_try_acquire(
|
||||
"a-free",
|
||||
"verification-task",
|
||||
std::time::Duration::from_secs(30),
|
||||
)
|
||||
.await
|
||||
.expect("runtime lock should be available")
|
||||
.expect("partially acquired lock should have been released");
|
||||
assert!(state
|
||||
.runtime_state()
|
||||
.lock_release(&free)
|
||||
.await
|
||||
.expect("free lock should release"));
|
||||
assert!(state
|
||||
.runtime_state()
|
||||
.lock_release(&held)
|
||||
.await
|
||||
.expect("held lock should release"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn agent_identity_import_lock_renewal_extends_all_leases() {
|
||||
let app = AppState::new().expect("app state should build");
|
||||
let state = AdminAppState::new(&app);
|
||||
let lease = state
|
||||
.runtime_state()
|
||||
.lock_try_acquire(
|
||||
"renewed-agent-lock",
|
||||
"agent-identity-task-renew",
|
||||
std::time::Duration::from_secs(1),
|
||||
)
|
||||
.await
|
||||
.expect("runtime lock should be available")
|
||||
.expect("lock should acquire");
|
||||
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
|
||||
renew_provider_agent_identity_import_locks(
|
||||
&state,
|
||||
std::slice::from_ref(&lease),
|
||||
std::time::Duration::from_secs(2),
|
||||
)
|
||||
.await
|
||||
.expect("lock should renew");
|
||||
tokio::time::sleep(std::time::Duration::from_millis(1_000)).await;
|
||||
|
||||
let contender = state
|
||||
.runtime_state()
|
||||
.lock_try_acquire(
|
||||
"renewed-agent-lock",
|
||||
"contending-task",
|
||||
std::time::Duration::from_secs(1),
|
||||
)
|
||||
.await
|
||||
.expect("runtime lock should be available");
|
||||
assert!(contender.is_none(), "renewed lease should still be held");
|
||||
release_provider_agent_identity_import_locks(&state, vec![lease]).await;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,3 +1,7 @@
|
||||
use super::super::super::duplicates::{
|
||||
acquire_codex_oauth_account_locks, find_duplicate_provider_oauth_key,
|
||||
release_codex_oauth_account_locks,
|
||||
};
|
||||
use super::super::super::errors::build_internal_control_error_response;
|
||||
use super::super::super::provisioning::{
|
||||
provider_oauth_token_payload_expires_at_unix_secs, seed_provider_oauth_pool_score,
|
||||
@@ -15,8 +19,10 @@ use super::shared::{
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_complete_key_id;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::handlers::shared::sync_provider_key_oauth_status_snapshot;
|
||||
use crate::provider_key_auth::provider_key_is_oauth_managed;
|
||||
use crate::GatewayError;
|
||||
use aether_data_contracts::repository::provider_catalog::ProviderCatalogKeyOAuthRuntimeStateCasUpdate;
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
http,
|
||||
@@ -88,6 +94,12 @@ pub(super) async fn handle_admin_provider_oauth_complete_key(
|
||||
"state 无效或已过期",
|
||||
));
|
||||
}
|
||||
if state_data.expected_encrypted_auth_config != key.encrypted_auth_config {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::CONFLICT,
|
||||
"授权期间 Key 认证信息已变更,请重新获取授权",
|
||||
));
|
||||
}
|
||||
|
||||
let provider_id = key.provider_id.clone();
|
||||
let provider = state
|
||||
@@ -208,39 +220,164 @@ pub(super) async fn handle_admin_provider_oauth_complete_key(
|
||||
"provider oauth encryption unavailable",
|
||||
));
|
||||
};
|
||||
let updated = state
|
||||
.update_provider_catalog_key_oauth_credentials(
|
||||
&key_id,
|
||||
&encrypted_api_key,
|
||||
Some(&encrypted_auth_config),
|
||||
expires_at,
|
||||
let codex_oauth_account_leases = if provider_type == "codex" {
|
||||
match acquire_codex_oauth_account_locks(state, &provider_id, &auth_config, "key-complete")
|
||||
.await
|
||||
{
|
||||
Ok(leases) => leases,
|
||||
Err(error) => {
|
||||
return Ok(build_internal_control_error_response(
|
||||
error.status_code(),
|
||||
error.detail(),
|
||||
));
|
||||
}
|
||||
}
|
||||
} else {
|
||||
Vec::new()
|
||||
};
|
||||
if provider_type == "codex" {
|
||||
let duplicate = match state
|
||||
.find_duplicate_provider_oauth_key(&provider_id, &auth_config, Some(&key_id))
|
||||
.await
|
||||
{
|
||||
Ok(duplicate) => duplicate,
|
||||
Err(detail) => {
|
||||
release_codex_oauth_account_locks(state, codex_oauth_account_leases).await;
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::CONFLICT,
|
||||
detail,
|
||||
));
|
||||
}
|
||||
};
|
||||
if let Some(duplicate) = duplicate {
|
||||
release_codex_oauth_account_locks(state, codex_oauth_account_leases).await;
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::CONFLICT,
|
||||
format!(
|
||||
"该 ChatGPT 账号已存在于其他 Key(名称: {})",
|
||||
duplicate.name
|
||||
),
|
||||
));
|
||||
}
|
||||
}
|
||||
let mut recovered_key = key.clone();
|
||||
recovered_key.encrypted_api_key = Some(encrypted_api_key.clone());
|
||||
recovered_key.encrypted_auth_config = Some(encrypted_auth_config.clone());
|
||||
recovered_key.expires_at_unix_secs = expires_at;
|
||||
recovered_key.oauth_invalid_at_unix_secs = None;
|
||||
recovered_key.oauth_invalid_reason = None;
|
||||
recovered_key.updated_at_unix_secs = Some(now_unix_secs);
|
||||
recovered_key.status_snapshot = sync_provider_key_oauth_status_snapshot(
|
||||
recovered_key.status_snapshot.as_ref(),
|
||||
&recovered_key,
|
||||
);
|
||||
let oauth_status = recovered_key
|
||||
.status_snapshot
|
||||
.as_ref()
|
||||
.and_then(|snapshot| snapshot.get("oauth"))
|
||||
.cloned()
|
||||
.unwrap_or(serde_json::Value::Null);
|
||||
let persisted_encrypted_auth_config = recovered_key
|
||||
.encrypted_auth_config
|
||||
.clone()
|
||||
.expect("recovered auth config should be present");
|
||||
let updated_result = state
|
||||
.app()
|
||||
.compare_and_update_provider_catalog_key_oauth_runtime_state(
|
||||
&ProviderCatalogKeyOAuthRuntimeStateCasUpdate {
|
||||
key_id: key_id.clone(),
|
||||
expected_encrypted_auth_config: state_data.expected_encrypted_auth_config,
|
||||
encrypted_auth_config: persisted_encrypted_auth_config.clone(),
|
||||
encrypted_api_key_update: Some(encrypted_api_key),
|
||||
expires_at_unix_secs_update: Some(expires_at),
|
||||
oauth_invalid_at_unix_secs: None,
|
||||
oauth_invalid_reason: None,
|
||||
reset_error_count: true,
|
||||
upstream_metadata_patch: None,
|
||||
status_snapshot_patch: json!({ "oauth": oauth_status }),
|
||||
updated_at_unix_secs: Some(now_unix_secs),
|
||||
},
|
||||
)
|
||||
.await?;
|
||||
.await;
|
||||
let _ = state
|
||||
.app()
|
||||
.invalidate_local_oauth_refresh_entry(&key_id)
|
||||
.await;
|
||||
let updated = match updated_result {
|
||||
Ok(updated) => updated,
|
||||
Err(error) => {
|
||||
release_codex_oauth_account_locks(state, codex_oauth_account_leases).await;
|
||||
return Err(error);
|
||||
}
|
||||
};
|
||||
if !updated {
|
||||
release_codex_oauth_account_locks(state, codex_oauth_account_leases).await;
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
"Key 不存在",
|
||||
http::StatusCode::CONFLICT,
|
||||
"授权期间 Key 认证信息已变更,请重新获取授权",
|
||||
));
|
||||
}
|
||||
if !state
|
||||
.clear_provider_catalog_key_oauth_invalid_marker(&key_id)
|
||||
.await?
|
||||
let current_after_cas = match state
|
||||
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
|
||||
.await
|
||||
{
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
"Key 不存在",
|
||||
));
|
||||
}
|
||||
let Some(recovered_key) = state
|
||||
.reset_provider_catalog_key_recovery_state(&key_id)
|
||||
.await?
|
||||
else {
|
||||
Ok(keys) => keys.into_iter().next(),
|
||||
Err(error) => {
|
||||
release_codex_oauth_account_locks(state, codex_oauth_account_leases).await;
|
||||
return Err(error);
|
||||
}
|
||||
};
|
||||
let Some(current_after_cas) = current_after_cas else {
|
||||
release_codex_oauth_account_locks(state, codex_oauth_account_leases).await;
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
"Key 不存在",
|
||||
));
|
||||
};
|
||||
if current_after_cas.encrypted_auth_config.as_deref()
|
||||
!= Some(persisted_encrypted_auth_config.as_str())
|
||||
{
|
||||
release_codex_oauth_account_locks(state, codex_oauth_account_leases).await;
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::CONFLICT,
|
||||
"授权期间 Key 认证信息已变更,请重新获取授权",
|
||||
));
|
||||
}
|
||||
let recovered_key = match state
|
||||
.reset_provider_catalog_key_recovery_state_fenced(&key_id, &persisted_encrypted_auth_config)
|
||||
.await
|
||||
{
|
||||
Ok(recovered_key) => recovered_key,
|
||||
Err(error) => {
|
||||
release_codex_oauth_account_locks(state, codex_oauth_account_leases).await;
|
||||
return Err(error);
|
||||
}
|
||||
};
|
||||
let Some(recovered_key) = recovered_key else {
|
||||
let key_exists = match state
|
||||
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
|
||||
.await
|
||||
{
|
||||
Ok(keys) => !keys.is_empty(),
|
||||
Err(error) => {
|
||||
release_codex_oauth_account_locks(state, codex_oauth_account_leases).await;
|
||||
return Err(error);
|
||||
}
|
||||
};
|
||||
release_codex_oauth_account_locks(state, codex_oauth_account_leases).await;
|
||||
if !key_exists {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
"Key 不存在",
|
||||
));
|
||||
}
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::CONFLICT,
|
||||
"授权期间 Key 认证信息已变更,请重新获取授权",
|
||||
));
|
||||
};
|
||||
seed_provider_oauth_pool_score(state, &provider.id, &recovered_key, now_unix_secs).await;
|
||||
release_codex_oauth_account_locks(state, codex_oauth_account_leases).await;
|
||||
|
||||
spawn_provider_oauth_account_state_refresh_after_update(
|
||||
state.cloned_app(),
|
||||
|
||||
+51
-12
@@ -1,4 +1,7 @@
|
||||
use super::super::super::duplicates::find_duplicate_provider_oauth_key;
|
||||
use super::super::super::duplicates::{
|
||||
acquire_codex_oauth_account_locks, find_duplicate_provider_oauth_key,
|
||||
release_codex_oauth_account_locks,
|
||||
};
|
||||
use super::super::super::errors::build_internal_control_error_response;
|
||||
use super::super::super::provisioning::{
|
||||
build_provider_oauth_auth_config_from_token_payload, create_provider_oauth_catalog_key,
|
||||
@@ -158,14 +161,39 @@ pub(super) async fn handle_admin_provider_oauth_complete_provider(
|
||||
};
|
||||
|
||||
let api_formats = provider_oauth_active_api_formats(&endpoints);
|
||||
let codex_oauth_account_leases = if provider_type == "codex" {
|
||||
match acquire_codex_oauth_account_locks(
|
||||
state,
|
||||
&provider_id,
|
||||
&auth_config,
|
||||
"provider-complete",
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(leases) => leases,
|
||||
Err(error) => {
|
||||
return Ok(build_internal_control_error_response(
|
||||
error.status_code(),
|
||||
error.detail(),
|
||||
));
|
||||
}
|
||||
}
|
||||
} else {
|
||||
Vec::new()
|
||||
};
|
||||
let duplicate = match state
|
||||
.find_duplicate_provider_oauth_key(&provider_id, &auth_config, None)
|
||||
.await
|
||||
{
|
||||
Ok(duplicate) => duplicate,
|
||||
Err(detail) => {
|
||||
release_codex_oauth_account_locks(state, codex_oauth_account_leases).await;
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
if provider_type == "codex" {
|
||||
http::StatusCode::CONFLICT
|
||||
} else {
|
||||
http::StatusCode::BAD_REQUEST
|
||||
},
|
||||
detail,
|
||||
));
|
||||
}
|
||||
@@ -173,7 +201,7 @@ pub(super) async fn handle_admin_provider_oauth_complete_provider(
|
||||
|
||||
let replaced = duplicate.is_some();
|
||||
let persisted_key = if let Some(existing_key) = duplicate {
|
||||
match state
|
||||
let update_result = state
|
||||
.update_existing_provider_oauth_catalog_key(
|
||||
&existing_key,
|
||||
&provider_type,
|
||||
@@ -183,10 +211,15 @@ pub(super) async fn handle_admin_provider_oauth_complete_provider(
|
||||
key_proxy.clone(),
|
||||
expires_at,
|
||||
)
|
||||
.await?
|
||||
{
|
||||
Some(key) => key,
|
||||
None => {
|
||||
.await;
|
||||
match update_result {
|
||||
Err(error) => {
|
||||
release_codex_oauth_account_locks(state, codex_oauth_account_leases).await;
|
||||
return Err(error);
|
||||
}
|
||||
Ok(Some(key)) => key,
|
||||
Ok(None) => {
|
||||
release_codex_oauth_account_locks(state, codex_oauth_account_leases).await;
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::SERVICE_UNAVAILABLE,
|
||||
"provider oauth write unavailable",
|
||||
@@ -214,7 +247,7 @@ pub(super) async fn handle_admin_provider_oauth_complete_provider(
|
||||
.unwrap_or(0)
|
||||
)
|
||||
});
|
||||
match state
|
||||
let create_result = state
|
||||
.create_provider_oauth_catalog_key(
|
||||
&provider_id,
|
||||
&provider_type,
|
||||
@@ -225,10 +258,15 @@ pub(super) async fn handle_admin_provider_oauth_complete_provider(
|
||||
key_proxy.clone(),
|
||||
expires_at,
|
||||
)
|
||||
.await?
|
||||
{
|
||||
Some(key) => key,
|
||||
None => {
|
||||
.await;
|
||||
match create_result {
|
||||
Err(error) => {
|
||||
release_codex_oauth_account_locks(state, codex_oauth_account_leases).await;
|
||||
return Err(error);
|
||||
}
|
||||
Ok(Some(key)) => key,
|
||||
Ok(None) => {
|
||||
release_codex_oauth_account_locks(state, codex_oauth_account_leases).await;
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::SERVICE_UNAVAILABLE,
|
||||
"provider oauth write unavailable",
|
||||
@@ -236,6 +274,7 @@ pub(super) async fn handle_admin_provider_oauth_complete_provider(
|
||||
}
|
||||
}
|
||||
};
|
||||
release_codex_oauth_account_locks(state, codex_oauth_account_leases).await;
|
||||
|
||||
spawn_provider_oauth_account_state_refresh_after_update(
|
||||
state.cloned_app(),
|
||||
|
||||
@@ -6,6 +6,31 @@ use axum::{
|
||||
use serde_json::{Map, Value};
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
pub(super) fn admin_provider_oauth_single_import_audit_taxonomy(
|
||||
request_body: Option<&axum::body::Bytes>,
|
||||
) -> (&'static str, &'static str) {
|
||||
let creates_agent_identity = request_body
|
||||
.and_then(|body| serde_json::from_slice::<Value>(body).ok())
|
||||
.and_then(|value| value.as_object().cloned())
|
||||
.is_some_and(|payload| {
|
||||
payload
|
||||
.get("create_agent_identity")
|
||||
.and_then(Value::as_bool)
|
||||
== Some(true)
|
||||
});
|
||||
if creates_agent_identity {
|
||||
(
|
||||
"admin_provider_oauth_agent_identity_created",
|
||||
"create_provider_agent_identity",
|
||||
)
|
||||
} else {
|
||||
(
|
||||
"admin_provider_oauth_refresh_token_imported",
|
||||
"import_provider_oauth_refresh_token",
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn attach_admin_provider_oauth_audit_response(
|
||||
response: Response<Body>,
|
||||
event_name: &'static str,
|
||||
@@ -96,4 +121,53 @@ mod tests {
|
||||
assert!(name.starts_with("codex_"));
|
||||
assert!(name.ends_with("_3"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn single_import_audit_distinguishes_agent_identity_creation_without_exposing_input() {
|
||||
let body = axum::body::Bytes::from(
|
||||
json!({
|
||||
"create_agent_identity": true,
|
||||
"access_token": "secret-access-token"
|
||||
})
|
||||
.to_string(),
|
||||
);
|
||||
assert_eq!(
|
||||
admin_provider_oauth_single_import_audit_taxonomy(Some(&body)),
|
||||
(
|
||||
"admin_provider_oauth_agent_identity_created",
|
||||
"create_provider_agent_identity",
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn single_import_audit_rejects_removed_session_token_creation_alias() {
|
||||
let body = axum::body::Bytes::from(
|
||||
json!({
|
||||
"create_agent_identity_from_session_token": true,
|
||||
"access_token": "secret-access-token"
|
||||
})
|
||||
.to_string(),
|
||||
);
|
||||
assert_eq!(
|
||||
admin_provider_oauth_single_import_audit_taxonomy(Some(&body)),
|
||||
(
|
||||
"admin_provider_oauth_refresh_token_imported",
|
||||
"import_provider_oauth_refresh_token",
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn single_import_audit_keeps_standard_import_taxonomy() {
|
||||
let body =
|
||||
axum::body::Bytes::from(json!({ "refresh_token": "secret-refresh-token" }).to_string());
|
||||
assert_eq!(
|
||||
admin_provider_oauth_single_import_audit_taxonomy(Some(&body)),
|
||||
(
|
||||
"admin_provider_oauth_refresh_token_imported",
|
||||
"import_provider_oauth_refresh_token",
|
||||
)
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,4 +1,7 @@
|
||||
use super::super::duplicates::find_duplicate_provider_oauth_key;
|
||||
use super::super::duplicates::{
|
||||
acquire_codex_oauth_account_locks, find_duplicate_provider_oauth_key,
|
||||
release_codex_oauth_account_locks,
|
||||
};
|
||||
use super::super::errors::build_internal_control_error_response;
|
||||
use super::super::provisioning::{
|
||||
build_provider_oauth_auth_config_from_token_payload, create_provider_oauth_catalog_key,
|
||||
@@ -27,10 +30,12 @@ use crate::handlers::admin::request::{
|
||||
};
|
||||
use crate::GatewayError;
|
||||
use aether_contracts::ProxySnapshot;
|
||||
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
|
||||
use aether_oauth::provider::{
|
||||
ProviderOAuthImportInput, ProviderOAuthService, ProviderOAuthTransportContext,
|
||||
};
|
||||
use aether_oauth::{core::OAuthError, network::OAuthNetworkContext};
|
||||
use aether_runtime_state::RuntimeLockLease;
|
||||
use axum::{
|
||||
body::Body,
|
||||
http,
|
||||
@@ -38,6 +43,15 @@ use axum::{
|
||||
Json,
|
||||
};
|
||||
use serde_json::json;
|
||||
use std::sync::{
|
||||
atomic::{AtomicBool, Ordering},
|
||||
Arc,
|
||||
};
|
||||
use std::time::Duration;
|
||||
use uuid::Uuid;
|
||||
|
||||
const CODEX_AGENT_IDENTITY_ENROLLMENT_LOCK_TTL: Duration = Duration::from_secs(180);
|
||||
const CODEX_AGENT_IDENTITY_ENROLLMENT_LOCK_RENEW_INTERVAL: Duration = Duration::from_secs(60);
|
||||
|
||||
struct AdminProviderOAuthSingleImportTokens {
|
||||
access_token: String,
|
||||
@@ -45,6 +59,19 @@ struct AdminProviderOAuthSingleImportTokens {
|
||||
expires_at: Option<u64>,
|
||||
}
|
||||
|
||||
struct CodexAgentIdentityEnrollment {
|
||||
leases: Vec<RuntimeLockLease>,
|
||||
duplicate: Option<StoredProviderCatalogKey>,
|
||||
lease_lost: Arc<AtomicBool>,
|
||||
heartbeat: tokio::task::JoinHandle<()>,
|
||||
}
|
||||
|
||||
impl Drop for CodexAgentIdentityEnrollment {
|
||||
fn drop(&mut self) {
|
||||
self.heartbeat.abort();
|
||||
}
|
||||
}
|
||||
|
||||
fn sanitize_windsurf_import_error(error: &OAuthError) -> String {
|
||||
match error {
|
||||
OAuthError::InvalidRequest(_) => "Windsurf 凭据验证失败: 请求参数无效".to_string(),
|
||||
@@ -115,14 +142,32 @@ fn import_payload_bool(payload: &serde_json::Map<String, serde_json::Value>, key
|
||||
.unwrap_or(false)
|
||||
}
|
||||
|
||||
fn codex_session_token_identity_hints(
|
||||
session_token: &str,
|
||||
fn codex_agent_identity_access_token_input(
|
||||
payload: &serde_json::Map<String, serde_json::Value>,
|
||||
) -> Option<String> {
|
||||
import_payload_string_any(payload, &["access_token", "accessToken"])
|
||||
}
|
||||
|
||||
fn import_payload_requests_agent_identity(
|
||||
payload: &serde_json::Map<String, serde_json::Value>,
|
||||
) -> bool {
|
||||
import_payload_bool(payload, "create_agent_identity")
|
||||
}
|
||||
|
||||
fn import_payload_requests_legacy_agent_identity(
|
||||
payload: &serde_json::Map<String, serde_json::Value>,
|
||||
) -> bool {
|
||||
import_payload_bool(payload, "create_agent_identity_from_session_token")
|
||||
}
|
||||
|
||||
fn codex_access_token_identity_hints(
|
||||
access_token: &str,
|
||||
) -> Result<serde_json::Map<String, serde_json::Value>, &'static str> {
|
||||
let mut hints = serde_json::Map::new();
|
||||
enrich_admin_provider_oauth_auth_config(
|
||||
"codex",
|
||||
&mut hints,
|
||||
&json!({ "access_token": session_token }),
|
||||
&json!({ "access_token": access_token }),
|
||||
);
|
||||
let account_id = hints
|
||||
.get("account_id")
|
||||
@@ -135,22 +180,173 @@ fn codex_session_token_identity_hints(
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty());
|
||||
if account_id.is_none() || user_id.is_none() {
|
||||
return Err("ChatGPT Session Token 缺少账号身份字段");
|
||||
return Err("ChatGPT Access Token 缺少账号身份字段");
|
||||
}
|
||||
Ok(hints)
|
||||
}
|
||||
|
||||
async fn resolve_admin_provider_oauth_codex_session_agent_identity_import(
|
||||
async fn prepare_codex_agent_identity_enrollment(
|
||||
state: &AdminAppState<'_>,
|
||||
session_token: &str,
|
||||
provider_id: &str,
|
||||
identity_hints: &serde_json::Map<String, serde_json::Value>,
|
||||
) -> Result<CodexAgentIdentityEnrollment, Response<Body>> {
|
||||
let lock_keys =
|
||||
crate::handlers::admin::provider::oauth::duplicates::codex_agent_identity_account_lock_keys(
|
||||
provider_id,
|
||||
identity_hints,
|
||||
);
|
||||
if lock_keys.is_empty() {
|
||||
return Err(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"ChatGPT Access Token 缺少账号身份字段",
|
||||
));
|
||||
}
|
||||
let owner = format!(
|
||||
"aether-gateway-agent-identity-enrollment-{}",
|
||||
Uuid::new_v4()
|
||||
);
|
||||
let mut leases = Vec::with_capacity(lock_keys.len());
|
||||
for lock_key in lock_keys {
|
||||
match state
|
||||
.runtime_state()
|
||||
.lock_try_acquire(
|
||||
lock_key.as_str(),
|
||||
owner.as_str(),
|
||||
CODEX_AGENT_IDENTITY_ENROLLMENT_LOCK_TTL,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(Some(lease)) => leases.push(lease),
|
||||
Ok(None) => {
|
||||
release_codex_agent_identity_leases(state, leases).await;
|
||||
return Err(build_internal_control_error_response(
|
||||
http::StatusCode::CONFLICT,
|
||||
"该 ChatGPT 账号正在创建 Agent Identity,请稍后重试",
|
||||
));
|
||||
}
|
||||
Err(error) => {
|
||||
tracing::warn!(
|
||||
provider_id = %provider_id,
|
||||
error = ?error,
|
||||
"gateway Agent Identity enrollment lock unavailable"
|
||||
);
|
||||
release_codex_agent_identity_leases(state, leases).await;
|
||||
return Err(build_internal_control_error_response(
|
||||
http::StatusCode::SERVICE_UNAVAILABLE,
|
||||
"Agent Identity 创建锁暂不可用,请稍后重试",
|
||||
));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let lease_lost = Arc::new(AtomicBool::new(false));
|
||||
let heartbeat = spawn_codex_agent_identity_enrollment_heartbeat(
|
||||
state.cloned_app(),
|
||||
leases.clone(),
|
||||
Arc::clone(&lease_lost),
|
||||
);
|
||||
|
||||
let duplicate = match state
|
||||
.find_duplicate_provider_oauth_key(provider_id, identity_hints, None)
|
||||
.await
|
||||
{
|
||||
Ok(duplicate) => duplicate,
|
||||
Err(detail) => {
|
||||
heartbeat.abort();
|
||||
release_codex_agent_identity_leases(state, leases).await;
|
||||
return Err(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
detail,
|
||||
));
|
||||
}
|
||||
};
|
||||
Ok(CodexAgentIdentityEnrollment {
|
||||
leases,
|
||||
duplicate,
|
||||
lease_lost,
|
||||
heartbeat,
|
||||
})
|
||||
}
|
||||
|
||||
fn spawn_codex_agent_identity_enrollment_heartbeat(
|
||||
app: crate::AppState,
|
||||
leases: Vec<RuntimeLockLease>,
|
||||
lease_lost: Arc<AtomicBool>,
|
||||
) -> tokio::task::JoinHandle<()> {
|
||||
tokio::spawn(async move {
|
||||
let mut renew_timer =
|
||||
tokio::time::interval(CODEX_AGENT_IDENTITY_ENROLLMENT_LOCK_RENEW_INTERVAL);
|
||||
renew_timer.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
|
||||
renew_timer.tick().await;
|
||||
loop {
|
||||
renew_timer.tick().await;
|
||||
for lease in &leases {
|
||||
match app
|
||||
.runtime_state()
|
||||
.lock_renew(lease, CODEX_AGENT_IDENTITY_ENROLLMENT_LOCK_TTL)
|
||||
.await
|
||||
{
|
||||
Ok(true) => {}
|
||||
Ok(false) => {
|
||||
lease_lost.store(true, Ordering::Release);
|
||||
tracing::error!(
|
||||
lock_key = %lease.key,
|
||||
"gateway Agent Identity enrollment lock was lost"
|
||||
);
|
||||
return;
|
||||
}
|
||||
Err(error) => {
|
||||
lease_lost.store(true, Ordering::Release);
|
||||
tracing::error!(
|
||||
lock_key = %lease.key,
|
||||
error = ?error,
|
||||
"gateway Agent Identity enrollment lock renewal failed"
|
||||
);
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
async fn release_codex_agent_identity_leases(
|
||||
state: &AdminAppState<'_>,
|
||||
leases: Vec<RuntimeLockLease>,
|
||||
) {
|
||||
for lease in leases {
|
||||
if let Err(error) = state.runtime_state().lock_release(&lease).await {
|
||||
tracing::warn!(
|
||||
lock_key = %lease.key,
|
||||
error = ?error,
|
||||
"gateway Agent Identity enrollment lock release failed"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn release_codex_agent_identity_enrollment(
|
||||
state: &AdminAppState<'_>,
|
||||
enrollment: Option<CodexAgentIdentityEnrollment>,
|
||||
) {
|
||||
let Some(mut enrollment) = enrollment else {
|
||||
return;
|
||||
};
|
||||
enrollment.heartbeat.abort();
|
||||
release_codex_agent_identity_leases(state, std::mem::take(&mut enrollment.leases)).await;
|
||||
}
|
||||
|
||||
async fn resolve_admin_provider_oauth_codex_access_token_agent_identity_import(
|
||||
state: &AdminAppState<'_>,
|
||||
access_token: &str,
|
||||
identity_hints: serde_json::Map<String, serde_json::Value>,
|
||||
request_proxy: Option<ProxySnapshot>,
|
||||
) -> Result<AdminProviderOAuthSingleImportTokens, Response<Body>> {
|
||||
let executor = crate::oauth::GatewayOAuthHttpExecutor::new(*state);
|
||||
let mut auth_config =
|
||||
aether_provider_transport::create_codex_agent_identity_from_session_token(
|
||||
aether_provider_transport::register_codex_agent_identity_from_access_token(
|
||||
&executor,
|
||||
session_token,
|
||||
access_token,
|
||||
OAuthNetworkContext::provider_operation(request_proxy),
|
||||
)
|
||||
.await
|
||||
@@ -584,18 +780,15 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
||||
));
|
||||
}
|
||||
};
|
||||
let create_agent_identity_from_session_token =
|
||||
import_payload_bool(&raw_payload, "create_agent_identity_from_session_token");
|
||||
let session_token_agent_identity_input = if create_agent_identity_from_session_token {
|
||||
import_payload_string_any(
|
||||
&raw_payload,
|
||||
&[
|
||||
"session_token",
|
||||
"sessionToken",
|
||||
"access_token",
|
||||
"accessToken",
|
||||
],
|
||||
)
|
||||
if import_payload_requests_legacy_agent_identity(&raw_payload) {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"旧版 Agent Identity 创建参数已停用,请使用 create_agent_identity 和 access_token",
|
||||
));
|
||||
}
|
||||
let create_agent_identity = import_payload_requests_agent_identity(&raw_payload);
|
||||
let agent_identity_access_token_input = if create_agent_identity {
|
||||
codex_agent_identity_access_token_input(&raw_payload)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
@@ -644,10 +837,7 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
||||
refresh_token_input.as_deref(),
|
||||
access_token_input.as_deref(),
|
||||
);
|
||||
if !create_agent_identity_from_session_token
|
||||
&& refresh_token_input.is_none()
|
||||
&& access_token_input.is_none()
|
||||
{
|
||||
if !create_agent_identity && refresh_token_input.is_none() && access_token_input.is_none() {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"Refresh Token、Access Token 或 sso_token 不能为空",
|
||||
@@ -665,16 +855,16 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
||||
"Kiro 不支持单条 Refresh Token 导入,请使用批量导入或设备授权。",
|
||||
));
|
||||
}
|
||||
if create_agent_identity_from_session_token && provider_type != "codex" {
|
||||
if create_agent_identity && provider_type != "codex" {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"仅 Codex Provider 支持使用 Session Token 创建 Agent Identity",
|
||||
"仅 Codex Provider 支持使用 Access Token 创建 Agent Identity",
|
||||
));
|
||||
}
|
||||
if create_agent_identity_from_session_token && refresh_token_input.is_some() {
|
||||
if create_agent_identity && refresh_token_input.is_some() {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"使用 Session Token 创建 Agent Identity 时不能同时提交 Refresh Token",
|
||||
"使用 Access Token 创建 Agent Identity 时不能同时提交 Refresh Token",
|
||||
));
|
||||
}
|
||||
let template = admin_provider_oauth_template(&provider_type);
|
||||
@@ -697,15 +887,16 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
||||
)
|
||||
.await;
|
||||
let key_proxy = provider_oauth_key_proxy_value(proxy_node_id.as_deref());
|
||||
let mut agent_identity_enrollment = None;
|
||||
|
||||
let resolved_import = if create_agent_identity_from_session_token {
|
||||
let Some(session_token) = session_token_agent_identity_input.as_deref() else {
|
||||
let resolved_import = if create_agent_identity {
|
||||
let Some(access_token) = agent_identity_access_token_input.as_deref() else {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"ChatGPT Session Token(JWT)不能为空",
|
||||
"ChatGPT Access Token(JWT)不能为空",
|
||||
));
|
||||
};
|
||||
let identity_hints = match codex_session_token_identity_hints(session_token) {
|
||||
let mut identity_hints = match codex_access_token_identity_hints(access_token) {
|
||||
Ok(hints) => hints,
|
||||
Err(detail) => {
|
||||
return Ok(build_internal_control_error_response(
|
||||
@@ -714,16 +905,30 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
||||
));
|
||||
}
|
||||
};
|
||||
match resolve_admin_provider_oauth_codex_session_agent_identity_import(
|
||||
identity_hints.insert("provider_type".to_string(), json!("codex"));
|
||||
let enrollment =
|
||||
match prepare_codex_agent_identity_enrollment(state, &provider_id, &identity_hints)
|
||||
.await
|
||||
{
|
||||
Ok(enrollment) => enrollment,
|
||||
Err(response) => return Ok(response),
|
||||
};
|
||||
match resolve_admin_provider_oauth_codex_access_token_agent_identity_import(
|
||||
state,
|
||||
session_token,
|
||||
access_token,
|
||||
identity_hints,
|
||||
request_proxy.clone(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(value) => value,
|
||||
Err(response) => return Ok(response),
|
||||
Ok(value) => {
|
||||
agent_identity_enrollment = Some(enrollment);
|
||||
value
|
||||
}
|
||||
Err(response) => {
|
||||
release_codex_agent_identity_enrollment(state, Some(enrollment)).await;
|
||||
return Ok(response);
|
||||
}
|
||||
}
|
||||
} else if provider_type == "windsurf" {
|
||||
if !import_payload_has_windsurf_credentials(&raw_payload) {
|
||||
@@ -772,7 +977,7 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
||||
mut auth_config,
|
||||
mut expires_at,
|
||||
} = resolved_import;
|
||||
if !create_agent_identity_from_session_token {
|
||||
if !create_agent_identity {
|
||||
apply_single_import_hints(&provider_type, &raw_payload, &mut auth_config);
|
||||
if let Some(header_access_token) =
|
||||
provider_oauth_import_authorization_bearer_token_from_object(&raw_payload)
|
||||
@@ -794,23 +999,85 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
||||
.map(str::trim)
|
||||
.is_some_and(|value| !value.is_empty());
|
||||
|
||||
let codex_oauth_account_leases = if !create_agent_identity && provider_type == "codex" {
|
||||
match acquire_codex_oauth_account_locks(state, &provider_id, &auth_config, "single-import")
|
||||
.await
|
||||
{
|
||||
Ok(leases) => leases,
|
||||
Err(error) => {
|
||||
return Ok(build_internal_control_error_response(
|
||||
error.status_code(),
|
||||
error.detail(),
|
||||
));
|
||||
}
|
||||
}
|
||||
} else {
|
||||
Vec::new()
|
||||
};
|
||||
|
||||
let api_formats = provider_oauth_active_api_formats(&endpoints);
|
||||
let duplicate = match state
|
||||
.find_duplicate_provider_oauth_key(&provider_id, &auth_config, None)
|
||||
.await
|
||||
{
|
||||
Ok(duplicate) => duplicate,
|
||||
Err(detail) => {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
detail,
|
||||
));
|
||||
let duplicate = if create_agent_identity {
|
||||
let initial_duplicate_id = agent_identity_enrollment
|
||||
.as_ref()
|
||||
.and_then(|enrollment| enrollment.duplicate.as_ref())
|
||||
.map(|key| key.id.clone());
|
||||
match state
|
||||
.find_duplicate_provider_oauth_key(&provider_id, &auth_config, None)
|
||||
.await
|
||||
{
|
||||
Ok(duplicate) => {
|
||||
let current_duplicate_id = duplicate.as_ref().map(|key| key.id.as_str());
|
||||
if initial_duplicate_id.as_deref() != current_duplicate_id {
|
||||
tracing::info!(
|
||||
provider_id = %provider_id,
|
||||
initial_duplicate_id = ?initial_duplicate_id,
|
||||
current_duplicate_id,
|
||||
"gateway Agent Identity duplicate changed during enrollment"
|
||||
);
|
||||
}
|
||||
duplicate
|
||||
}
|
||||
Err(detail) => {
|
||||
release_codex_agent_identity_enrollment(state, agent_identity_enrollment).await;
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::CONFLICT,
|
||||
detail,
|
||||
));
|
||||
}
|
||||
}
|
||||
} else {
|
||||
match state
|
||||
.find_duplicate_provider_oauth_key(&provider_id, &auth_config, None)
|
||||
.await
|
||||
{
|
||||
Ok(duplicate) => duplicate,
|
||||
Err(detail) => {
|
||||
release_codex_oauth_account_locks(state, codex_oauth_account_leases).await;
|
||||
return Ok(build_internal_control_error_response(
|
||||
if provider_type == "codex" {
|
||||
http::StatusCode::CONFLICT
|
||||
} else {
|
||||
http::StatusCode::BAD_REQUEST
|
||||
},
|
||||
detail,
|
||||
));
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
let replaced = duplicate.is_some();
|
||||
let persisted_key = if let Some(existing_key) = duplicate {
|
||||
match state
|
||||
if agent_identity_enrollment
|
||||
.as_ref()
|
||||
.is_some_and(|enrollment| enrollment.lease_lost.load(Ordering::Acquire))
|
||||
{
|
||||
release_codex_agent_identity_enrollment(state, agent_identity_enrollment).await;
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::SERVICE_UNAVAILABLE,
|
||||
"Agent Identity 创建锁已失效,请稍后重试",
|
||||
));
|
||||
}
|
||||
let persisted_key_result = if let Some(existing_key) = duplicate {
|
||||
state
|
||||
.update_existing_provider_oauth_catalog_key(
|
||||
&existing_key,
|
||||
&provider_type,
|
||||
@@ -820,21 +1087,12 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
||||
key_proxy.clone(),
|
||||
expires_at,
|
||||
)
|
||||
.await?
|
||||
{
|
||||
Some(key) => key,
|
||||
None => {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::SERVICE_UNAVAILABLE,
|
||||
"provider oauth write unavailable",
|
||||
));
|
||||
}
|
||||
}
|
||||
.await
|
||||
} else {
|
||||
let name = name.unwrap_or_else(|| {
|
||||
admin_provider_oauth_key_name_from_auth_config(&provider_type, &auth_config, None)
|
||||
});
|
||||
match state
|
||||
state
|
||||
.create_provider_oauth_catalog_key(
|
||||
&provider_id,
|
||||
&provider_type,
|
||||
@@ -845,17 +1103,68 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
||||
key_proxy.clone(),
|
||||
expires_at,
|
||||
)
|
||||
.await?
|
||||
{
|
||||
Some(key) => key,
|
||||
None => {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::SERVICE_UNAVAILABLE,
|
||||
"provider oauth write unavailable",
|
||||
));
|
||||
}
|
||||
.await
|
||||
};
|
||||
let persisted_key = match persisted_key_result {
|
||||
Ok(Some(key)) => key,
|
||||
Ok(None) => {
|
||||
release_codex_oauth_account_locks(state, codex_oauth_account_leases).await;
|
||||
release_codex_agent_identity_enrollment(state, agent_identity_enrollment).await;
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::SERVICE_UNAVAILABLE,
|
||||
"provider oauth write unavailable",
|
||||
));
|
||||
}
|
||||
Err(error) => {
|
||||
release_codex_oauth_account_locks(state, codex_oauth_account_leases).await;
|
||||
release_codex_agent_identity_enrollment(state, agent_identity_enrollment).await;
|
||||
return Err(error);
|
||||
}
|
||||
};
|
||||
release_codex_oauth_account_locks(state, codex_oauth_account_leases).await;
|
||||
|
||||
let agent_identity_task_ready = if create_agent_identity {
|
||||
match runtime_endpoint.as_ref() {
|
||||
Some(endpoint) => match state
|
||||
.read_provider_transport_snapshot_uncached(
|
||||
&provider_id,
|
||||
&endpoint.id,
|
||||
&persisted_key.id,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(Some(transport)) => {
|
||||
match state.resolve_local_oauth_request_auth(&transport).await {
|
||||
Ok(Some(_)) => true,
|
||||
Ok(None) => false,
|
||||
Err(error) => {
|
||||
tracing::warn!(
|
||||
provider_id = %provider_id,
|
||||
key_id = %persisted_key.id,
|
||||
error = ?error,
|
||||
"gateway Agent Identity initial task registration failed"
|
||||
);
|
||||
false
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(None) => false,
|
||||
Err(error) => {
|
||||
tracing::warn!(
|
||||
provider_id = %provider_id,
|
||||
key_id = %persisted_key.id,
|
||||
error = ?error,
|
||||
"gateway Agent Identity pending transport reload failed"
|
||||
);
|
||||
false
|
||||
}
|
||||
},
|
||||
None => false,
|
||||
}
|
||||
} else {
|
||||
true
|
||||
};
|
||||
release_codex_agent_identity_enrollment(state, agent_identity_enrollment).await;
|
||||
|
||||
spawn_provider_oauth_account_state_refresh_after_update(
|
||||
state.cloned_app(),
|
||||
@@ -864,6 +1173,28 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
||||
request_proxy.clone(),
|
||||
);
|
||||
|
||||
if !agent_identity_task_ready {
|
||||
return Ok((
|
||||
http::StatusCode::ACCEPTED,
|
||||
Json(json!({
|
||||
"detail": "Agent Identity 已安全保存,但 task 初始化暂未完成,系统将自动重试",
|
||||
"key_id": persisted_key.id,
|
||||
"provider_type": provider_type,
|
||||
"expires_at": serde_json::Value::Null,
|
||||
"has_refresh_token": false,
|
||||
"temporary": false,
|
||||
"email": auth_config
|
||||
.get("email")
|
||||
.cloned()
|
||||
.unwrap_or(serde_json::Value::Null),
|
||||
"replaced": replaced,
|
||||
"task_ready": false,
|
||||
"recoverable": true,
|
||||
})),
|
||||
)
|
||||
.into_response());
|
||||
}
|
||||
|
||||
Ok(Json(json!({
|
||||
"key_id": persisted_key.id,
|
||||
"provider_type": provider_type,
|
||||
@@ -882,9 +1213,12 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
apply_single_import_hints, codex_session_token_identity_hints, import_payload_bool,
|
||||
import_payload_string_any, import_payload_u64_any, sanitize_windsurf_import_error,
|
||||
apply_single_import_hints, codex_access_token_identity_hints,
|
||||
codex_agent_identity_access_token_input, import_payload_requests_agent_identity,
|
||||
import_payload_requests_legacy_agent_identity, import_payload_string_any,
|
||||
import_payload_u64_any, sanitize_windsurf_import_error,
|
||||
};
|
||||
use crate::handlers::admin::provider::oauth::duplicates::codex_agent_identity_account_lock_keys;
|
||||
use aether_oauth::core::OAuthError;
|
||||
use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
|
||||
use serde_json::json;
|
||||
@@ -911,8 +1245,36 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn session_token_agent_identity_hints_require_and_extract_chatgpt_identity() {
|
||||
let session_token = unsigned_jwt(json!({
|
||||
fn agent_identity_reads_access_token_and_ignores_session_token_alias() {
|
||||
let payload = json!({
|
||||
"accessToken": "access-token",
|
||||
"sessionToken": "session-token-must-not-win",
|
||||
})
|
||||
.as_object()
|
||||
.cloned()
|
||||
.expect("payload should be an object");
|
||||
|
||||
assert_eq!(
|
||||
codex_agent_identity_access_token_input(&payload).as_deref(),
|
||||
Some("access-token")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn new_agent_identity_flag_does_not_treat_session_token_as_access_token() {
|
||||
let payload = json!({
|
||||
"sessionToken": "session-token-must-not-be-used",
|
||||
})
|
||||
.as_object()
|
||||
.cloned()
|
||||
.expect("payload should be an object");
|
||||
|
||||
assert_eq!(codex_agent_identity_access_token_input(&payload), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn access_token_agent_identity_hints_require_and_extract_chatgpt_identity() {
|
||||
let access_token = unsigned_jwt(json!({
|
||||
"https://api.openai.com/auth": {
|
||||
"chatgpt_account_id": "account-1",
|
||||
"chatgpt_user_id": "user-1",
|
||||
@@ -923,8 +1285,8 @@ mod tests {
|
||||
}
|
||||
}));
|
||||
|
||||
let hints = codex_session_token_identity_hints(&session_token)
|
||||
.expect("session token identity hints should parse");
|
||||
let hints = codex_access_token_identity_hints(&access_token)
|
||||
.expect("access token identity hints should parse");
|
||||
|
||||
assert_eq!(hints.get("account_id"), Some(&json!("account-1")));
|
||||
assert_eq!(hints.get("user_id"), Some(&json!("user-1")));
|
||||
@@ -935,41 +1297,81 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn session_token_agent_identity_hints_reject_missing_identity() {
|
||||
let session_token = unsigned_jwt(json!({
|
||||
fn access_token_agent_identity_hints_reject_missing_identity() {
|
||||
let access_token = unsigned_jwt(json!({
|
||||
"https://api.openai.com/auth": {
|
||||
"chatgpt_account_id": "account-1"
|
||||
}
|
||||
}));
|
||||
|
||||
assert_eq!(
|
||||
codex_session_token_identity_hints(&session_token),
|
||||
Err("ChatGPT Session Token 缺少账号身份字段")
|
||||
codex_access_token_identity_hints(&access_token),
|
||||
Err("ChatGPT Access Token 缺少账号身份字段")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn session_token_agent_identity_flag_is_explicit_boolean_only() {
|
||||
fn agent_identity_enrollment_lock_is_stable_and_account_scoped() {
|
||||
let first = json!({
|
||||
"account_id": "account-1",
|
||||
"user_id": "user-1",
|
||||
})
|
||||
.as_object()
|
||||
.cloned()
|
||||
.expect("identity hints should be an object");
|
||||
let second = json!({
|
||||
"account_id": "account-2",
|
||||
"user_id": "user-1",
|
||||
})
|
||||
.as_object()
|
||||
.cloned()
|
||||
.expect("identity hints should be an object");
|
||||
|
||||
let first_keys = codex_agent_identity_account_lock_keys("provider-1", &first);
|
||||
assert_eq!(
|
||||
first_keys,
|
||||
codex_agent_identity_account_lock_keys("provider-1", &first)
|
||||
);
|
||||
assert_ne!(
|
||||
first_keys,
|
||||
codex_agent_identity_account_lock_keys("provider-1", &second)
|
||||
);
|
||||
assert_ne!(
|
||||
first_keys,
|
||||
codex_agent_identity_account_lock_keys("provider-2", &first)
|
||||
);
|
||||
assert!(first_keys
|
||||
.iter()
|
||||
.all(|key| !key.contains("account-1") && !key.contains("user-1")));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn agent_identity_flag_is_explicit_boolean_and_rejects_legacy_alias() {
|
||||
let payload = json!({
|
||||
"create_agent_identity": true,
|
||||
})
|
||||
.as_object()
|
||||
.cloned()
|
||||
.expect("payload should be an object");
|
||||
assert!(import_payload_requests_agent_identity(&payload));
|
||||
|
||||
let string_payload = json!({
|
||||
"create_agent_identity": "true",
|
||||
})
|
||||
.as_object()
|
||||
.cloned()
|
||||
.expect("payload should be an object");
|
||||
assert!(!import_payload_requests_agent_identity(&string_payload));
|
||||
|
||||
let legacy_payload = json!({
|
||||
"create_agent_identity_from_session_token": true,
|
||||
})
|
||||
.as_object()
|
||||
.cloned()
|
||||
.expect("payload should be an object");
|
||||
assert!(import_payload_bool(
|
||||
&payload,
|
||||
"create_agent_identity_from_session_token"
|
||||
));
|
||||
|
||||
let string_payload = json!({
|
||||
"create_agent_identity_from_session_token": "true",
|
||||
})
|
||||
.as_object()
|
||||
.cloned()
|
||||
.expect("payload should be an object");
|
||||
assert!(!import_payload_bool(
|
||||
&string_payload,
|
||||
"create_agent_identity_from_session_token"
|
||||
assert!(!import_payload_requests_agent_identity(&legacy_payload));
|
||||
assert!(import_payload_requests_legacy_agent_identity(
|
||||
&legacy_payload
|
||||
));
|
||||
}
|
||||
|
||||
|
||||
@@ -3,6 +3,7 @@ use super::state::{
|
||||
build_admin_provider_oauth_supported_types_payload,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::paths::{
|
||||
admin_provider_oauth_agent_identity_import_task_provider_id,
|
||||
admin_provider_oauth_batch_import_provider_id,
|
||||
admin_provider_oauth_batch_import_task_provider_id, admin_provider_oauth_complete_key_id,
|
||||
admin_provider_oauth_complete_provider_id, admin_provider_oauth_device_authorize_provider_id,
|
||||
@@ -83,6 +84,16 @@ pub(crate) async fn maybe_build_local_admin_provider_oauth_response(
|
||||
));
|
||||
}
|
||||
|
||||
if route_kind == Some("get_agent_identity_import_task_status") && *method == http::Method::GET {
|
||||
return Ok(Some(
|
||||
tasks::handle_admin_provider_oauth_agent_identity_import_task_status(
|
||||
state,
|
||||
request_context,
|
||||
)
|
||||
.await?,
|
||||
));
|
||||
}
|
||||
|
||||
if route_kind == Some("complete_key_oauth") && *method == http::Method::POST {
|
||||
let response = complete::handle_admin_provider_oauth_complete_key(
|
||||
state,
|
||||
@@ -128,6 +139,8 @@ pub(crate) async fn maybe_build_local_admin_provider_oauth_response(
|
||||
}
|
||||
|
||||
if route_kind == Some("import_refresh_token") && *method == http::Method::POST {
|
||||
let (event_name, action) =
|
||||
helpers::admin_provider_oauth_single_import_audit_taxonomy(request_body);
|
||||
let response = import::handle_admin_provider_oauth_import_refresh_token(
|
||||
state,
|
||||
request_context,
|
||||
@@ -136,8 +149,8 @@ pub(crate) async fn maybe_build_local_admin_provider_oauth_response(
|
||||
.await?;
|
||||
return Ok(Some(helpers::attach_admin_provider_oauth_audit_response(
|
||||
response,
|
||||
"admin_provider_oauth_refresh_token_imported",
|
||||
"import_provider_oauth_refresh_token",
|
||||
event_name,
|
||||
action,
|
||||
"provider",
|
||||
admin_provider_oauth_import_provider_id(request_context.path()),
|
||||
)));
|
||||
@@ -172,6 +185,22 @@ pub(crate) async fn maybe_build_local_admin_provider_oauth_response(
|
||||
)));
|
||||
}
|
||||
|
||||
if route_kind == Some("start_agent_identity_import_task") && *method == http::Method::POST {
|
||||
let response = batch::handle_admin_provider_oauth_start_agent_identity_import_task(
|
||||
state,
|
||||
request_context,
|
||||
request_body,
|
||||
)
|
||||
.await?;
|
||||
return Ok(Some(helpers::attach_admin_provider_oauth_audit_response(
|
||||
response,
|
||||
"admin_provider_oauth_agent_identity_import_started",
|
||||
"start_provider_agent_identity_import",
|
||||
"provider",
|
||||
admin_provider_oauth_agent_identity_import_task_provider_id(request_context.path()),
|
||||
)));
|
||||
}
|
||||
|
||||
if route_kind == Some("device_authorize") && *method == http::Method::POST {
|
||||
let response = device::handle_admin_provider_oauth_device_authorize(
|
||||
state,
|
||||
|
||||
+35
-79
@@ -1,16 +1,7 @@
|
||||
use super::super::super::errors::{
|
||||
merge_provider_oauth_refresh_failure_reason, normalize_provider_oauth_refresh_error_message,
|
||||
};
|
||||
use super::super::super::quota::shared::{
|
||||
persist_provider_quota_refresh_state, provider_auto_remove_banned_keys,
|
||||
should_auto_remove_oauth_invalid_key,
|
||||
};
|
||||
use super::super::super::errors::normalize_provider_oauth_refresh_error_message;
|
||||
use super::super::super::runtime::refresh_provider_oauth_account_state_after_update;
|
||||
use super::helpers::{self, RefreshDispatch, RefreshRequestContext, RefreshSuccessContext};
|
||||
use super::response;
|
||||
use crate::handlers::admin::provider::shared::payloads::{
|
||||
OAUTH_ACCOUNT_BLOCK_PREFIX, OAUTH_REFRESH_FAILED_PREFIX,
|
||||
};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminLocalOAuthRefreshError};
|
||||
use crate::GatewayError;
|
||||
use axum::http;
|
||||
@@ -62,62 +53,36 @@ pub(super) async fn execute_admin_provider_oauth_refresh(
|
||||
"gateway manual provider oauth refresh failed"
|
||||
);
|
||||
if matches!(status_code, 400 | 401 | 403) {
|
||||
let failure_reason = format!(
|
||||
"{OAUTH_REFRESH_FAILED_PREFIX}Token 续期失败 ({status_code}): {error_reason}"
|
||||
);
|
||||
let merged_reason = merge_provider_oauth_refresh_failure_reason(
|
||||
key.oauth_invalid_reason.as_deref(),
|
||||
&failure_reason,
|
||||
);
|
||||
if let Some(merged_reason) = merged_reason {
|
||||
let _ = persist_provider_quota_refresh_state(
|
||||
state,
|
||||
&key_id,
|
||||
None,
|
||||
Some(helpers::unix_now_secs()),
|
||||
Some(merged_reason),
|
||||
None,
|
||||
let auto_removed = state
|
||||
.app()
|
||||
.persist_local_oauth_refresh_failure_state(
|
||||
&transport,
|
||||
status_code,
|
||||
body_excerpt.as_str(),
|
||||
false,
|
||||
)
|
||||
.await?;
|
||||
if provider_auto_remove_banned_keys(provider.config.as_ref()) {
|
||||
let now_unix_secs = helpers::unix_now_secs();
|
||||
let auto_removed = state
|
||||
.cleanup_provider_catalog_key_if_current(
|
||||
&provider,
|
||||
&key_id,
|
||||
|latest_key| {
|
||||
should_auto_remove_oauth_invalid_key(
|
||||
latest_key,
|
||||
Some(&failure_reason),
|
||||
false,
|
||||
now_unix_secs,
|
||||
)
|
||||
},
|
||||
)
|
||||
.await?;
|
||||
if auto_removed {
|
||||
tracing::info!(
|
||||
trace_id = %trace_id,
|
||||
key_id = %key_id,
|
||||
provider_id = %provider.id,
|
||||
provider_type = %provider_type,
|
||||
event_name = "auto_removed_oauth_refresh_failed",
|
||||
"gateway manual provider oauth refresh auto-removed unusable key"
|
||||
);
|
||||
return Ok(RefreshDispatch::Respond(
|
||||
response::oauth_refresh_auto_removed_response(&error_reason),
|
||||
));
|
||||
}
|
||||
}
|
||||
if auto_removed {
|
||||
tracing::info!(
|
||||
trace_id = %trace_id,
|
||||
key_id = %key_id,
|
||||
provider_id = %provider.id,
|
||||
provider_type = %provider_type,
|
||||
event_name = "refresh_failed_retained",
|
||||
"gateway manual provider oauth refresh failure retained key"
|
||||
event_name = "auto_removed_oauth_refresh_failed",
|
||||
"gateway manual provider oauth refresh auto-removed unusable key"
|
||||
);
|
||||
return Ok(RefreshDispatch::Respond(
|
||||
response::oauth_refresh_auto_removed_response(&error_reason),
|
||||
));
|
||||
}
|
||||
tracing::info!(
|
||||
trace_id = %trace_id,
|
||||
key_id = %key_id,
|
||||
provider_id = %provider.id,
|
||||
provider_type = %provider_type,
|
||||
event_name = "refresh_failed_retained",
|
||||
"gateway manual provider oauth refresh failure retained key"
|
||||
);
|
||||
}
|
||||
return Ok(RefreshDispatch::Respond(
|
||||
response::oauth_refresh_failed_bad_request_response(&error_reason),
|
||||
@@ -164,28 +129,6 @@ pub(super) async fn execute_admin_provider_oauth_refresh(
|
||||
}
|
||||
};
|
||||
|
||||
if !helpers::key_is_account_blocked(&key, OAUTH_ACCOUNT_BLOCK_PREFIX) {
|
||||
let previous_oauth_refresh_issue =
|
||||
key.oauth_invalid_reason.as_deref().is_some_and(|reason| {
|
||||
reason.lines().map(str::trim).any(|line| {
|
||||
line.starts_with("[OAUTH_EXPIRED]") || line.starts_with("[REFRESH_FAILED]")
|
||||
})
|
||||
});
|
||||
let cleared = state
|
||||
.clear_provider_catalog_key_oauth_invalid_marker(&key_id)
|
||||
.await?;
|
||||
if cleared && previous_oauth_refresh_issue {
|
||||
tracing::info!(
|
||||
trace_id = %trace_id,
|
||||
key_id = %key_id,
|
||||
provider_id = %provider.id,
|
||||
provider_type = %provider_type,
|
||||
event_name = "refresh_fixed",
|
||||
"gateway manual provider oauth refresh cleared oauth invalid marker"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
let refreshed_key = state
|
||||
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
|
||||
.await?
|
||||
@@ -223,3 +166,16 @@ pub(super) async fn execute_admin_provider_oauth_refresh(
|
||||
account_state_recheck_error,
|
||||
}))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
#[test]
|
||||
fn manual_refresh_uses_fenced_state_persistence_without_redundant_clear() {
|
||||
let source = include_str!("execution.rs");
|
||||
assert!(source.contains("persist_local_oauth_refresh_failure_state"));
|
||||
let redundant_clear = ["clear_provider_catalog_key_", "oauth_invalid_marker"].concat();
|
||||
let unfenced_persistence = ["persist_provider_quota_", "refresh_state"].concat();
|
||||
assert!(!source.contains(&redundant_clear));
|
||||
assert!(!source.contains(&unfenced_persistence));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -43,12 +43,6 @@ pub(super) async fn parse_admin_provider_oauth_refresh_request(
|
||||
)));
|
||||
};
|
||||
let parsed_auth_config = helpers::parse_auth_config_object(&decrypted_auth_config);
|
||||
if !helpers::auth_config_has_refresh_token(&parsed_auth_config) {
|
||||
return Ok(RefreshDispatch::Respond(response::control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"缺少 refresh_token,需要重新授权",
|
||||
)));
|
||||
}
|
||||
|
||||
let provider_id = key.provider_id.clone();
|
||||
let Some(provider) = state
|
||||
@@ -63,6 +57,16 @@ pub(super) async fn parse_admin_provider_oauth_refresh_request(
|
||||
)));
|
||||
};
|
||||
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
|
||||
let is_agent_identity = provider_type == "codex"
|
||||
&& crate::provider_transport::is_codex_agent_identity_auth_config_value(
|
||||
&serde_json::Value::Object(parsed_auth_config.clone()),
|
||||
);
|
||||
if !is_agent_identity && !helpers::auth_config_has_refresh_token(&parsed_auth_config) {
|
||||
return Ok(RefreshDispatch::Respond(response::control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"缺少 refresh_token,需要重新授权",
|
||||
)));
|
||||
}
|
||||
if !provider_key_is_oauth_managed(&key, provider_type.as_str()) {
|
||||
return Ok(RefreshDispatch::Respond(response::control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
|
||||
@@ -86,6 +86,7 @@ pub(super) async fn handle_admin_provider_oauth_start_key(
|
||||
&provider_id,
|
||||
&provider_type,
|
||||
pkce_verifier.as_deref(),
|
||||
key.encrypted_auth_config.as_deref(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
@@ -158,7 +159,13 @@ pub(super) async fn handle_admin_provider_oauth_start_provider(
|
||||
.then(generate_provider_oauth_pkce_verifier);
|
||||
let code_challenge = pkce_verifier.as_deref().map(provider_oauth_pkce_s256);
|
||||
let nonce = match state
|
||||
.save_provider_oauth_state("", &provider_id, &provider_type, pkce_verifier.as_deref())
|
||||
.save_provider_oauth_state(
|
||||
"",
|
||||
&provider_id,
|
||||
&provider_type,
|
||||
pkce_verifier.as_deref(),
|
||||
None,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(nonce) => nonce,
|
||||
|
||||
@@ -1,5 +1,8 @@
|
||||
use super::super::errors::build_internal_control_error_response;
|
||||
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_batch_import_task_path;
|
||||
use crate::handlers::admin::provider::shared::paths::{
|
||||
admin_provider_oauth_agent_identity_import_task_path,
|
||||
admin_provider_oauth_batch_import_task_path,
|
||||
};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::handlers::admin::shared::attach_admin_audit_response;
|
||||
use crate::GatewayError;
|
||||
@@ -10,13 +13,49 @@ use axum::{
|
||||
Json,
|
||||
};
|
||||
|
||||
const PROVIDER_AGENT_IDENTITY_IMPORT_KIND: &str = "agent_identity";
|
||||
|
||||
fn provider_oauth_import_task_matches_route(
|
||||
task_id: &str,
|
||||
payload: &serde_json::Value,
|
||||
agent_identity_only: bool,
|
||||
) -> bool {
|
||||
let has_agent_prefix = task_id.starts_with("agent-identity-");
|
||||
let import_kind = payload
|
||||
.get("import_kind")
|
||||
.and_then(serde_json::Value::as_str);
|
||||
if agent_identity_only {
|
||||
has_agent_prefix && import_kind == Some(PROVIDER_AGENT_IDENTITY_IMPORT_KIND)
|
||||
} else {
|
||||
!has_agent_prefix && import_kind != Some(PROVIDER_AGENT_IDENTITY_IMPORT_KIND)
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) async fn handle_admin_provider_oauth_batch_import_task_status(
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
let Some((provider_id, task_id)) =
|
||||
handle_admin_provider_oauth_import_task_status(state, request_context, false).await
|
||||
}
|
||||
|
||||
pub(super) async fn handle_admin_provider_oauth_agent_identity_import_task_status(
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
handle_admin_provider_oauth_import_task_status(state, request_context, true).await
|
||||
}
|
||||
|
||||
async fn handle_admin_provider_oauth_import_task_status(
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
agent_identity_only: bool,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
let task_path = if agent_identity_only {
|
||||
admin_provider_oauth_agent_identity_import_task_path(request_context.path())
|
||||
} else {
|
||||
admin_provider_oauth_batch_import_task_path(request_context.path())
|
||||
else {
|
||||
};
|
||||
let Some((provider_id, task_id)) = task_path else {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
"批量导入任务不存在",
|
||||
@@ -40,27 +79,86 @@ pub(super) async fn handle_admin_provider_oauth_batch_import_task_status(
|
||||
));
|
||||
}
|
||||
};
|
||||
if !provider_oauth_import_task_matches_route(&task_id, &payload, agent_identity_only) {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
"导入任务不存在或已过期",
|
||||
));
|
||||
}
|
||||
let status = payload
|
||||
.get("status")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(ToOwned::to_owned)
|
||||
.unwrap_or_default();
|
||||
let response = Json(payload).into_response();
|
||||
let (completed_event, failed_event, action, target_type) = if agent_identity_only {
|
||||
(
|
||||
"admin_provider_oauth_agent_identity_import_completed_viewed",
|
||||
"admin_provider_oauth_agent_identity_import_failed_viewed",
|
||||
"view_provider_agent_identity_import_terminal_state",
|
||||
"provider_agent_identity_import_task",
|
||||
)
|
||||
} else {
|
||||
(
|
||||
"admin_provider_oauth_batch_task_completed_viewed",
|
||||
"admin_provider_oauth_batch_task_failed_viewed",
|
||||
"view_provider_oauth_batch_task_terminal_state",
|
||||
"provider_oauth_batch_task",
|
||||
)
|
||||
};
|
||||
Ok(match status.as_str() {
|
||||
"completed" => attach_admin_audit_response(
|
||||
response,
|
||||
"admin_provider_oauth_batch_task_completed_viewed",
|
||||
"view_provider_oauth_batch_task_terminal_state",
|
||||
"provider_oauth_batch_task",
|
||||
completed_event,
|
||||
action,
|
||||
target_type,
|
||||
&format!("{provider_id}:{task_id}"),
|
||||
),
|
||||
"failed" => attach_admin_audit_response(
|
||||
response,
|
||||
"admin_provider_oauth_batch_task_failed_viewed",
|
||||
"view_provider_oauth_batch_task_terminal_state",
|
||||
"provider_oauth_batch_task",
|
||||
failed_event,
|
||||
action,
|
||||
target_type,
|
||||
&format!("{provider_id}:{task_id}"),
|
||||
),
|
||||
_ => response,
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::provider_oauth_import_task_matches_route;
|
||||
use serde_json::json;
|
||||
|
||||
#[test]
|
||||
fn import_task_status_routes_are_bidirectionally_isolated() {
|
||||
let agent_payload = json!({ "import_kind": "agent_identity" });
|
||||
let batch_payload = json!({ "import_kind": "oauth_batch" });
|
||||
|
||||
assert!(provider_oauth_import_task_matches_route(
|
||||
"agent-identity-task-1",
|
||||
&agent_payload,
|
||||
true,
|
||||
));
|
||||
assert!(!provider_oauth_import_task_matches_route(
|
||||
"agent-identity-task-1",
|
||||
&agent_payload,
|
||||
false,
|
||||
));
|
||||
assert!(provider_oauth_import_task_matches_route(
|
||||
"batch-task-1",
|
||||
&batch_payload,
|
||||
false,
|
||||
));
|
||||
assert!(!provider_oauth_import_task_matches_route(
|
||||
"batch-task-1",
|
||||
&batch_payload,
|
||||
true,
|
||||
));
|
||||
assert!(provider_oauth_import_task_matches_route(
|
||||
"legacy-batch-task",
|
||||
&json!({}),
|
||||
false,
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,7 +1,38 @@
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::provider_key_auth::provider_key_is_oauth_managed;
|
||||
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
use aether_runtime_state::RuntimeLockLease;
|
||||
use axum::http;
|
||||
use sha2::{Digest, Sha256};
|
||||
use std::time::{Duration, SystemTime, UNIX_EPOCH};
|
||||
use uuid::Uuid;
|
||||
|
||||
const CODEX_OAUTH_ACCOUNT_LOCK_TTL: Duration = Duration::from_secs(180);
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub(crate) enum CodexOAuthAccountLockError {
|
||||
MissingIdentity,
|
||||
Contended,
|
||||
Unavailable,
|
||||
}
|
||||
|
||||
impl CodexOAuthAccountLockError {
|
||||
pub(crate) const fn status_code(self) -> http::StatusCode {
|
||||
match self {
|
||||
Self::MissingIdentity => http::StatusCode::BAD_REQUEST,
|
||||
Self::Contended => http::StatusCode::CONFLICT,
|
||||
Self::Unavailable => http::StatusCode::SERVICE_UNAVAILABLE,
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) const fn detail(self) -> &'static str {
|
||||
match self {
|
||||
Self::MissingIdentity => "Codex 账号身份字段缺失,无法安全写入授权",
|
||||
Self::Contended => "该 ChatGPT 账号正在更新授权,请稍后重试",
|
||||
Self::Unavailable => "Codex 账号授权锁暂不可用,请稍后重试",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn normalize_codex_plan_group_for_provider_oauth(
|
||||
plan_type: Option<&serde_json::Value>,
|
||||
@@ -26,6 +57,174 @@ fn normalize_provider_oauth_identity_value(value: Option<&serde_json::Value>) ->
|
||||
.map(ToOwned::to_owned)
|
||||
}
|
||||
|
||||
fn normalize_provider_oauth_identity_value_from_keys(
|
||||
auth_config: &serde_json::Map<String, serde_json::Value>,
|
||||
keys: &[&str],
|
||||
) -> Option<String> {
|
||||
keys.iter()
|
||||
.find_map(|key| normalize_provider_oauth_identity_value(auth_config.get(*key)))
|
||||
}
|
||||
|
||||
fn codex_agent_identity_account_lock_key(
|
||||
provider_id: &str,
|
||||
identity_kind: &str,
|
||||
identity_parts: &[&str],
|
||||
) -> String {
|
||||
let mut digest = Sha256::new();
|
||||
digest.update(provider_id.trim().as_bytes());
|
||||
digest.update([0]);
|
||||
digest.update(identity_kind.as_bytes());
|
||||
for part in identity_parts {
|
||||
digest.update([0]);
|
||||
digest.update(part.as_bytes());
|
||||
}
|
||||
format!(
|
||||
"provider_oauth_agent_identity_account:{:x}",
|
||||
digest.finalize()
|
||||
)
|
||||
}
|
||||
|
||||
pub(crate) fn codex_agent_identity_account_lock_keys(
|
||||
provider_id: &str,
|
||||
auth_config: &serde_json::Map<String, serde_json::Value>,
|
||||
) -> Vec<String> {
|
||||
let account_user_id = normalize_provider_oauth_identity_value_from_keys(
|
||||
auth_config,
|
||||
&[
|
||||
"account_user_id",
|
||||
"accountUserId",
|
||||
"chatgpt_account_user_id",
|
||||
"chatgptAccountUserId",
|
||||
],
|
||||
);
|
||||
let account_id = normalize_provider_oauth_identity_value_from_keys(
|
||||
auth_config,
|
||||
&[
|
||||
"account_id",
|
||||
"accountId",
|
||||
"chatgpt_account_id",
|
||||
"chatgptAccountId",
|
||||
],
|
||||
);
|
||||
let user_id = normalize_provider_oauth_identity_value_from_keys(
|
||||
auth_config,
|
||||
&["user_id", "userId", "chatgpt_user_id", "chatgptUserId"],
|
||||
);
|
||||
let email = normalize_provider_oauth_identity_value_from_keys(auth_config, &["email"]);
|
||||
|
||||
let mut keys = Vec::with_capacity(5);
|
||||
if let Some(account_user_id) = account_user_id.as_deref() {
|
||||
keys.push(codex_agent_identity_account_lock_key(
|
||||
provider_id,
|
||||
"account_user_id",
|
||||
&[account_user_id],
|
||||
));
|
||||
}
|
||||
if let (Some(account_id), Some(user_id)) = (account_id.as_deref(), user_id.as_deref()) {
|
||||
keys.push(codex_agent_identity_account_lock_key(
|
||||
provider_id,
|
||||
"account_id_user_id",
|
||||
&[account_id, user_id],
|
||||
));
|
||||
}
|
||||
if let (Some(account_id), Some(email)) = (account_id.as_deref(), email.as_deref()) {
|
||||
keys.push(codex_agent_identity_account_lock_key(
|
||||
provider_id,
|
||||
"account_id_email",
|
||||
&[account_id, email],
|
||||
));
|
||||
}
|
||||
if let Some(user_id) = user_id.as_deref() {
|
||||
keys.push(codex_agent_identity_account_lock_key(
|
||||
provider_id,
|
||||
"user_id",
|
||||
&[user_id],
|
||||
));
|
||||
}
|
||||
if let Some(email) = email.as_deref() {
|
||||
keys.push(codex_agent_identity_account_lock_key(
|
||||
provider_id,
|
||||
"email",
|
||||
&[email],
|
||||
));
|
||||
}
|
||||
keys.sort_unstable();
|
||||
keys.dedup();
|
||||
keys
|
||||
}
|
||||
|
||||
pub(crate) async fn acquire_codex_oauth_account_locks(
|
||||
state: &AdminAppState<'_>,
|
||||
provider_id: &str,
|
||||
auth_config: &serde_json::Map<String, serde_json::Value>,
|
||||
operation: &str,
|
||||
) -> Result<Vec<RuntimeLockLease>, CodexOAuthAccountLockError> {
|
||||
let lock_keys = codex_agent_identity_account_lock_keys(provider_id, auth_config);
|
||||
if lock_keys.is_empty() {
|
||||
return Err(CodexOAuthAccountLockError::MissingIdentity);
|
||||
}
|
||||
|
||||
let owner = format!(
|
||||
"aether-gateway-codex-oauth-{}-{}",
|
||||
operation.trim(),
|
||||
Uuid::new_v4()
|
||||
);
|
||||
let mut leases = Vec::with_capacity(lock_keys.len());
|
||||
for lock_key in lock_keys {
|
||||
match state
|
||||
.runtime_state()
|
||||
.lock_try_acquire(
|
||||
lock_key.as_str(),
|
||||
owner.as_str(),
|
||||
CODEX_OAUTH_ACCOUNT_LOCK_TTL,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(Some(lease)) => leases.push(lease),
|
||||
Ok(None) => {
|
||||
release_codex_oauth_account_locks(state, leases).await;
|
||||
return Err(CodexOAuthAccountLockError::Contended);
|
||||
}
|
||||
Err(error) => {
|
||||
tracing::warn!(
|
||||
provider_id = %provider_id,
|
||||
lock_key = %lock_key,
|
||||
operation,
|
||||
error = ?error,
|
||||
"gateway Codex OAuth account lock unavailable"
|
||||
);
|
||||
release_codex_oauth_account_locks(state, leases).await;
|
||||
return Err(CodexOAuthAccountLockError::Unavailable);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// The lock is distributed, while the catalog cache is process-local. A
|
||||
// fresh read inside the lease is required to observe the previous holder.
|
||||
state.app().data.clear_provider_catalog_cache();
|
||||
Ok(leases)
|
||||
}
|
||||
|
||||
pub(crate) async fn release_codex_oauth_account_locks(
|
||||
state: &AdminAppState<'_>,
|
||||
leases: Vec<RuntimeLockLease>,
|
||||
) {
|
||||
for lease in leases.into_iter().rev() {
|
||||
match state.runtime_state().lock_release(&lease).await {
|
||||
Ok(true) => {}
|
||||
Ok(false) => tracing::warn!(
|
||||
lock_key = %lease.key,
|
||||
"gateway Codex OAuth account lock was not owned during release"
|
||||
),
|
||||
Err(error) => tracing::warn!(
|
||||
lock_key = %lease.key,
|
||||
error = ?error,
|
||||
"gateway Codex OAuth account lock release failed"
|
||||
),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn is_openai_provider_oauth_provider_type(value: Option<&serde_json::Value>) -> bool {
|
||||
value
|
||||
.and_then(serde_json::Value::as_str)
|
||||
@@ -55,6 +254,24 @@ fn match_codex_provider_oauth_identity(
|
||||
return None;
|
||||
}
|
||||
|
||||
let new_agent_runtime_id = normalize_provider_oauth_identity_value(
|
||||
new_auth_config
|
||||
.get("agent_runtime_id")
|
||||
.or_else(|| new_auth_config.get("agentRuntimeId")),
|
||||
);
|
||||
let existing_agent_runtime_id = normalize_provider_oauth_identity_value(
|
||||
existing_auth_config
|
||||
.get("agent_runtime_id")
|
||||
.or_else(|| existing_auth_config.get("agentRuntimeId")),
|
||||
);
|
||||
if new_agent_runtime_id
|
||||
.as_deref()
|
||||
.zip(existing_agent_runtime_id.as_deref())
|
||||
.is_some_and(|(left, right)| left == right)
|
||||
{
|
||||
return Some(true);
|
||||
}
|
||||
|
||||
let new_account_user_id =
|
||||
normalize_provider_oauth_identity_value(new_auth_config.get("account_user_id"));
|
||||
let existing_account_user_id =
|
||||
@@ -212,6 +429,11 @@ pub(crate) async fn find_duplicate_provider_oauth_key(
|
||||
let new_email = normalize_provider_oauth_identity_value(auth_config.get("email"));
|
||||
let new_user_id = normalize_provider_oauth_identity_value(auth_config.get("user_id"));
|
||||
let new_account_id = normalize_provider_oauth_identity_value(auth_config.get("account_id"));
|
||||
let new_agent_runtime_id = normalize_provider_oauth_identity_value(
|
||||
auth_config
|
||||
.get("agent_runtime_id")
|
||||
.or_else(|| auth_config.get("agentRuntimeId")),
|
||||
);
|
||||
let new_credential_fingerprint =
|
||||
normalize_provider_oauth_identity_value(auth_config.get("credential_fingerprint"));
|
||||
let new_auth_method = normalize_provider_oauth_identity_value(auth_config.get("auth_method"));
|
||||
@@ -220,11 +442,15 @@ pub(crate) async fn find_duplicate_provider_oauth_key(
|
||||
if new_email.is_none()
|
||||
&& new_user_id.is_none()
|
||||
&& new_account_id.is_none()
|
||||
&& new_agent_runtime_id.is_none()
|
||||
&& new_credential_fingerprint.is_none()
|
||||
{
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
// Duplicate checks are write admission checks. Never let a process-local
|
||||
// read-through cache hide a row committed by the previous lock holder.
|
||||
state.app().data.clear_provider_catalog_cache();
|
||||
let existing_keys = state
|
||||
.list_provider_catalog_keys_by_provider_ids(&[provider_id.to_string()])
|
||||
.await
|
||||
@@ -325,6 +551,7 @@ pub(crate) async fn find_duplicate_provider_oauth_key(
|
||||
let identifier =
|
||||
normalize_provider_oauth_identity_value(auth_config.get("account_user_id"))
|
||||
.or_else(|| normalize_provider_oauth_identity_value(auth_config.get("account_id")))
|
||||
.or_else(|| new_agent_runtime_id.clone())
|
||||
.or_else(|| {
|
||||
normalize_provider_oauth_identity_value(
|
||||
auth_config.get("credential_fingerprint"),
|
||||
@@ -345,7 +572,13 @@ pub(crate) async fn find_duplicate_provider_oauth_key(
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::match_windsurf_provider_oauth_identity;
|
||||
use super::{
|
||||
acquire_codex_oauth_account_locks, codex_agent_identity_account_lock_keys,
|
||||
match_codex_provider_oauth_identity, match_windsurf_provider_oauth_identity,
|
||||
release_codex_oauth_account_locks, CodexOAuthAccountLockError,
|
||||
};
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::AppState;
|
||||
use serde_json::{json, Map, Value};
|
||||
|
||||
fn auth_config(value: Value) -> Map<String, Value> {
|
||||
@@ -371,6 +604,189 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn codex_agent_identity_matches_runtime_without_account_metadata() {
|
||||
let new_auth_config = auth_config(json!({
|
||||
"provider_type": "codex",
|
||||
"auth_mode": "agentIdentity",
|
||||
"agent_runtime_id": "runtime-1",
|
||||
"agent_private_key": "new-private-key"
|
||||
}));
|
||||
let existing_auth_config = auth_config(json!({
|
||||
"provider_type": "codex",
|
||||
"auth_mode": "agentIdentity",
|
||||
"agentRuntimeId": "runtime-1",
|
||||
"agent_private_key": "existing-private-key"
|
||||
}));
|
||||
|
||||
assert_eq!(
|
||||
match_codex_provider_oauth_identity(&new_auth_config, &existing_auth_config),
|
||||
Some(true)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn direct_and_json_agent_identity_imports_share_account_lock_keys() {
|
||||
let direct_identity_hints = auth_config(json!({
|
||||
"provider_type": "codex",
|
||||
"account_id": "account-1",
|
||||
"account_user_id": "account-user-1",
|
||||
"user_id": "user-1",
|
||||
"email": "[email protected]"
|
||||
}));
|
||||
let imported_auth_config = auth_config(json!({
|
||||
"provider_type": "codex",
|
||||
"auth_mode": "agentIdentity",
|
||||
"agent_runtime_id": "runtime-1",
|
||||
"accountId": "account-1",
|
||||
"chatgptAccountUserId": "account-user-1",
|
||||
"chatgptUserId": "user-1",
|
||||
"email": "[email protected]"
|
||||
}));
|
||||
|
||||
let direct_keys =
|
||||
codex_agent_identity_account_lock_keys("provider-codex", &direct_identity_hints);
|
||||
let imported_keys =
|
||||
codex_agent_identity_account_lock_keys("provider-codex", &imported_auth_config);
|
||||
let shared_keys = direct_keys
|
||||
.iter()
|
||||
.filter(|key| imported_keys.contains(key))
|
||||
.collect::<Vec<_>>();
|
||||
|
||||
assert_eq!(shared_keys.len(), 5);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn ordinary_codex_oauth_and_agent_identity_share_runtime_account_locks() {
|
||||
let app = AppState::new().expect("app state should build");
|
||||
let state = AdminAppState::new(&app);
|
||||
let ordinary = auth_config(json!({
|
||||
"provider_type": "codex",
|
||||
"account_id": "account-1",
|
||||
"account_user_id": "account-user-1",
|
||||
"user_id": "user-1",
|
||||
"email": "[email protected]"
|
||||
}));
|
||||
let agent = auth_config(json!({
|
||||
"provider_type": "codex",
|
||||
"auth_mode": "agentIdentity",
|
||||
"agent_runtime_id": "runtime-1",
|
||||
"account_id": "account-1",
|
||||
"account_user_id": "account-user-1",
|
||||
"user_id": "user-1",
|
||||
"email": "[email protected]"
|
||||
}));
|
||||
|
||||
let first =
|
||||
acquire_codex_oauth_account_locks(&state, "provider-codex", &ordinary, "ordinary-test")
|
||||
.await
|
||||
.expect("ordinary OAuth lock should acquire");
|
||||
let second =
|
||||
acquire_codex_oauth_account_locks(&state, "provider-codex", &agent, "agent-test")
|
||||
.await
|
||||
.expect_err("Agent Identity must contend on the same account locks");
|
||||
assert_eq!(second, CodexOAuthAccountLockError::Contended);
|
||||
|
||||
release_codex_oauth_account_locks(&state, first).await;
|
||||
let third =
|
||||
acquire_codex_oauth_account_locks(&state, "provider-codex", &agent, "agent-retry-test")
|
||||
.await
|
||||
.expect("account locks should be reusable after release");
|
||||
release_codex_oauth_account_locks(&state, third).await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn codex_oauth_account_lock_rejects_identity_free_config() {
|
||||
let app = AppState::new().expect("app state should build");
|
||||
let state = AdminAppState::new(&app);
|
||||
let config = auth_config(json!({"provider_type": "codex"}));
|
||||
|
||||
let error = acquire_codex_oauth_account_locks(
|
||||
&state,
|
||||
"provider-codex",
|
||||
&config,
|
||||
"missing-identity-test",
|
||||
)
|
||||
.await
|
||||
.expect_err("identity-free Codex writes must not proceed unlocked");
|
||||
|
||||
assert_eq!(error, CodexOAuthAccountLockError::MissingIdentity);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn codex_oauth_account_lock_releases_partial_acquisition() {
|
||||
let app = AppState::new().expect("app state should build");
|
||||
let state = AdminAppState::new(&app);
|
||||
let config = auth_config(json!({
|
||||
"provider_type": "codex",
|
||||
"account_id": "account-partial",
|
||||
"account_user_id": "account-user-partial",
|
||||
"user_id": "user-partial",
|
||||
"email": "[email protected]"
|
||||
}));
|
||||
let keys = codex_agent_identity_account_lock_keys("provider-codex", &config);
|
||||
let held_key = keys.last().expect("account locks should not be empty");
|
||||
let held = state
|
||||
.runtime_state()
|
||||
.lock_try_acquire(held_key, "other-owner", std::time::Duration::from_secs(30))
|
||||
.await
|
||||
.expect("runtime lock should be available")
|
||||
.expect("last account lock should acquire");
|
||||
|
||||
let error =
|
||||
acquire_codex_oauth_account_locks(&state, "provider-codex", &config, "partial-test")
|
||||
.await
|
||||
.expect_err("held final lock should cause contention");
|
||||
assert_eq!(error, CodexOAuthAccountLockError::Contended);
|
||||
|
||||
let first_key = keys.first().expect("account locks should not be empty");
|
||||
let first = state
|
||||
.runtime_state()
|
||||
.lock_try_acquire(
|
||||
first_key,
|
||||
"verification-owner",
|
||||
std::time::Duration::from_secs(30),
|
||||
)
|
||||
.await
|
||||
.expect("runtime lock should be available")
|
||||
.expect("partially acquired account lock should have been released");
|
||||
assert!(state
|
||||
.runtime_state()
|
||||
.lock_release(&first)
|
||||
.await
|
||||
.expect("verification lock should release"));
|
||||
assert!(state
|
||||
.runtime_state()
|
||||
.lock_release(&held)
|
||||
.await
|
||||
.expect("held lock should release"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn agent_identity_account_locks_cover_generic_user_and_email_deduplication() {
|
||||
let first = auth_config(json!({
|
||||
"provider_type": "codex",
|
||||
"agent_runtime_id": "runtime-1",
|
||||
"user_id": "user-1",
|
||||
"email": "[email protected]"
|
||||
}));
|
||||
let second = auth_config(json!({
|
||||
"provider_type": "codex",
|
||||
"agent_runtime_id": "runtime-2",
|
||||
"user_id": "user-1",
|
||||
"email": "[email protected]"
|
||||
}));
|
||||
|
||||
let first_keys = codex_agent_identity_account_lock_keys("provider-codex", &first);
|
||||
let second_keys = codex_agent_identity_account_lock_keys("provider-codex", &second);
|
||||
let shared_keys = first_keys
|
||||
.iter()
|
||||
.filter(|key| second_keys.contains(key))
|
||||
.collect::<Vec<_>>();
|
||||
|
||||
assert_eq!(shared_keys.len(), 2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn windsurf_identity_rejects_different_account_id() {
|
||||
let new_auth_config = auth_config(json!({
|
||||
|
||||
@@ -19,10 +19,9 @@ use self::plan::{
|
||||
};
|
||||
use super::shared::{
|
||||
build_quota_snapshot_payload, extract_execution_error_message,
|
||||
oauth_refresh_auto_removed_result, persist_provider_quota_refresh_state,
|
||||
provider_auto_remove_banned_keys, provider_auto_remove_quota_exhausted_keys,
|
||||
quota_key_auto_removed, quota_refresh_success_invalid_state,
|
||||
should_auto_remove_structured_reason, ProviderQuotaExecutionOutcome,
|
||||
oauth_refresh_auto_removed_result, persist_fenced_provider_quota_refresh_state,
|
||||
persist_provider_quota_refresh_state, quota_key_auto_removed,
|
||||
quota_refresh_success_invalid_state, ProviderQuotaExecutionOutcome,
|
||||
};
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::provider_key_auth::provider_key_is_oauth_managed;
|
||||
@@ -399,9 +398,6 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
keys: Vec<StoredProviderCatalogKey>,
|
||||
proxy_override: Option<ProxySnapshot>,
|
||||
) -> Result<Option<serde_json::Value>, GatewayError> {
|
||||
let auto_remove_abnormal_keys = provider_auto_remove_banned_keys(provider.config.as_ref());
|
||||
let auto_remove_quota_exhausted_keys =
|
||||
provider_auto_remove_quota_exhausted_keys(provider.config.as_ref());
|
||||
let mut results = Vec::new();
|
||||
let mut success_count = 0usize;
|
||||
let mut failed_count = 0usize;
|
||||
@@ -429,8 +425,29 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
continue;
|
||||
}
|
||||
};
|
||||
|
||||
let is_oauth_managed = provider_key_is_oauth_managed(&key, provider.provider_type.as_str());
|
||||
let quota_auth_config_fence = if is_oauth_managed {
|
||||
match state
|
||||
.app()
|
||||
.capture_provider_transport_auth_config_fence(&transport)
|
||||
.await?
|
||||
{
|
||||
Some(ciphertext) => Some(ciphertext),
|
||||
None => {
|
||||
failed_count += 1;
|
||||
results.push(json!({
|
||||
"key_id": key.id,
|
||||
"key_name": key.name,
|
||||
"status": "error",
|
||||
"message": "OAuth credential changed before quota refresh",
|
||||
}));
|
||||
continue;
|
||||
}
|
||||
}
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
let resolved_oauth_auth = if is_oauth_managed {
|
||||
state.resolve_local_oauth_header_auth(&transport).await?
|
||||
} else {
|
||||
@@ -647,17 +664,27 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
}
|
||||
}
|
||||
|
||||
let auto_remove_candidate = auto_remove_abnormal_keys
|
||||
&& should_auto_remove_structured_reason(oauth_invalid_reason.as_deref());
|
||||
let persisted = persist_provider_quota_refresh_state(
|
||||
state,
|
||||
&key.id,
|
||||
metadata_update.as_ref(),
|
||||
oauth_invalid_at_unix_secs,
|
||||
oauth_invalid_reason.clone(),
|
||||
None,
|
||||
)
|
||||
.await?;
|
||||
let persisted = if let Some(expected_auth_config) = quota_auth_config_fence.as_deref() {
|
||||
persist_fenced_provider_quota_refresh_state(
|
||||
state,
|
||||
&key.id,
|
||||
expected_auth_config,
|
||||
metadata_update.as_ref(),
|
||||
oauth_invalid_at_unix_secs,
|
||||
oauth_invalid_reason.clone(),
|
||||
)
|
||||
.await?
|
||||
} else {
|
||||
persist_provider_quota_refresh_state(
|
||||
state,
|
||||
&key.id,
|
||||
metadata_update.as_ref(),
|
||||
oauth_invalid_at_unix_secs,
|
||||
oauth_invalid_reason.clone(),
|
||||
None,
|
||||
)
|
||||
.await?
|
||||
};
|
||||
if !persisted {
|
||||
failed_count += 1;
|
||||
results.push(json!({
|
||||
@@ -668,32 +695,15 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
}));
|
||||
continue;
|
||||
}
|
||||
let auto_removed_hard_banned = if auto_remove_candidate {
|
||||
state
|
||||
.cleanup_provider_catalog_key_if_current(provider, &key.id, |latest_key| {
|
||||
should_auto_remove_structured_reason(latest_key.oauth_invalid_reason.as_deref())
|
||||
})
|
||||
.await?
|
||||
} else {
|
||||
false
|
||||
};
|
||||
// Codex quota responses never auto-delete keys. Without a repository
|
||||
// conditional delete, any read-then-delete sequence could remove a
|
||||
// replacement Agent Identity installed while the response was in flight.
|
||||
let auto_removed_hard_banned = false;
|
||||
if auto_removed_hard_banned {
|
||||
auto_removed_count += 1;
|
||||
auto_removed_hard_banned_count += 1;
|
||||
}
|
||||
let auto_removed_quota_exhausted =
|
||||
if !auto_removed_hard_banned && auto_remove_quota_exhausted_keys {
|
||||
state
|
||||
.cleanup_provider_catalog_key_if_current(provider, &key.id, |latest_key| {
|
||||
aether_admin::provider::pool::admin_pool_key_account_quota_exhausted(
|
||||
latest_key,
|
||||
provider.provider_type.as_str(),
|
||||
)
|
||||
})
|
||||
.await?
|
||||
} else {
|
||||
false
|
||||
};
|
||||
let auto_removed_quota_exhausted = false;
|
||||
if auto_removed_quota_exhausted {
|
||||
auto_removed_count += 1;
|
||||
status = "quota_exhausted".to_string();
|
||||
|
||||
@@ -14,8 +14,9 @@ use aether_contracts::{
|
||||
ResolvedTransportProfile, EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER,
|
||||
};
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
ProviderCatalogKeyRuntimeMetadataUpdate, ProviderCatalogKeyStatusSnapshotUpdate,
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
|
||||
ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ProviderCatalogKeyRuntimeMetadataUpdate,
|
||||
ProviderCatalogKeyStatusSnapshotUpdate, StoredProviderCatalogEndpoint,
|
||||
StoredProviderCatalogKey,
|
||||
};
|
||||
use aether_provider_pool::{ProviderPoolQuotaRequestSpec, ProviderPoolService};
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
@@ -200,6 +201,10 @@ pub(super) fn extract_execution_error_message(result: &ExecutionResult) -> Optio
|
||||
admin_provider_quota_pure::extract_execution_error_message(result)
|
||||
}
|
||||
|
||||
fn extract_execution_error_detail(result: &ExecutionResult) -> Option<String> {
|
||||
admin_provider_quota_pure::extract_execution_error_detail(result)
|
||||
}
|
||||
|
||||
pub(super) fn quota_refresh_success_invalid_state(
|
||||
key: &StoredProviderCatalogKey,
|
||||
) -> (Option<u64>, Option<String>) {
|
||||
@@ -301,6 +306,86 @@ pub(crate) async fn persist_provider_quota_refresh_state(
|
||||
.await
|
||||
}
|
||||
|
||||
/// Persist a Codex Agent Identity quota response only when the exact encrypted
|
||||
/// auth_config used for the request is still installed. Metadata, OAuth state,
|
||||
/// and their status projection share one repository CAS so a replacement cannot
|
||||
/// receive any portion of an older response.
|
||||
pub(crate) async fn persist_fenced_provider_quota_refresh_state(
|
||||
state: &AdminAppState<'_>,
|
||||
key_id: &str,
|
||||
expected_encrypted_auth_config: &str,
|
||||
metadata_update: Option<&serde_json::Value>,
|
||||
oauth_invalid_at_unix_secs: Option<u64>,
|
||||
oauth_invalid_reason: Option<String>,
|
||||
) -> Result<bool, GatewayError> {
|
||||
let expected_encrypted_auth_config = expected_encrypted_auth_config.trim();
|
||||
if expected_encrypted_auth_config.is_empty() {
|
||||
return Ok(false);
|
||||
}
|
||||
if metadata_update.is_some_and(|value| !value.is_object()) {
|
||||
return Err(GatewayError::Internal(
|
||||
"fenced quota metadata update must be an object".to_string(),
|
||||
));
|
||||
}
|
||||
let Some(mut latest_key) = state
|
||||
.read_provider_catalog_keys_by_ids(&[key_id.to_string()])
|
||||
.await?
|
||||
.into_iter()
|
||||
.next()
|
||||
else {
|
||||
return Ok(false);
|
||||
};
|
||||
if latest_key.encrypted_auth_config.as_deref() != Some(expected_encrypted_auth_config) {
|
||||
return Ok(false);
|
||||
}
|
||||
|
||||
let quota_snapshot_provider_type =
|
||||
metadata_update.and_then(aether_provider_pool::provider_pool_quota_metadata_provider_type);
|
||||
if let Some(metadata_update) = metadata_update {
|
||||
latest_key.upstream_metadata = Some(merge_upstream_metadata(
|
||||
latest_key.upstream_metadata.as_ref(),
|
||||
metadata_update,
|
||||
));
|
||||
}
|
||||
latest_key.oauth_invalid_at_unix_secs = oauth_invalid_at_unix_secs;
|
||||
latest_key.oauth_invalid_reason = oauth_invalid_reason;
|
||||
if let Some(provider_type) = quota_snapshot_provider_type.as_deref() {
|
||||
latest_key.status_snapshot = sync_provider_key_quota_status_snapshot(
|
||||
latest_key.status_snapshot.as_ref(),
|
||||
provider_type,
|
||||
latest_key.upstream_metadata.as_ref(),
|
||||
"refresh_api",
|
||||
);
|
||||
}
|
||||
latest_key.status_snapshot =
|
||||
sync_provider_key_oauth_status_snapshot(latest_key.status_snapshot.as_ref(), &latest_key);
|
||||
latest_key.updated_at_unix_secs = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs());
|
||||
|
||||
state
|
||||
.app()
|
||||
.compare_and_update_provider_catalog_key_oauth_runtime_state(
|
||||
&ProviderCatalogKeyOAuthRuntimeStateCasUpdate {
|
||||
key_id: key_id.to_string(),
|
||||
expected_encrypted_auth_config: Some(expected_encrypted_auth_config.to_string()),
|
||||
encrypted_auth_config: expected_encrypted_auth_config.to_string(),
|
||||
encrypted_api_key_update: None,
|
||||
expires_at_unix_secs_update: None,
|
||||
oauth_invalid_at_unix_secs: latest_key.oauth_invalid_at_unix_secs,
|
||||
oauth_invalid_reason: latest_key.oauth_invalid_reason.clone(),
|
||||
upstream_metadata_patch: metadata_update.cloned(),
|
||||
status_snapshot_patch: provider_quota_refresh_status_patch(
|
||||
latest_key.status_snapshot.as_ref(),
|
||||
),
|
||||
reset_error_count: false,
|
||||
updated_at_unix_secs: latest_key.updated_at_unix_secs,
|
||||
},
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn persist_provider_quota_refresh_state_after_read<F>(
|
||||
state: &AdminAppState<'_>,
|
||||
key_id: &str,
|
||||
@@ -452,7 +537,7 @@ pub(super) async fn execute_provider_quota_plan(
|
||||
if !crate::provider_transport::is_codex_agent_identity_transport(transport)
|
||||
|| !crate::provider_transport::is_codex_agent_identity_invalid_task_response(
|
||||
result.status_code,
|
||||
extract_execution_error_message(&result).as_deref(),
|
||||
extract_execution_error_detail(&result).as_deref(),
|
||||
)
|
||||
{
|
||||
return Ok(ProviderQuotaExecutionOutcome::Response(result));
|
||||
|
||||
@@ -5,7 +5,7 @@ use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::handlers::admin::shared::{provider_key_status_snapshot_payload, unix_secs_to_rfc3339};
|
||||
use crate::provider_key_auth::{
|
||||
provider_key_auth_config_is_agent_identity, provider_key_auth_config_uses_header_authorization,
|
||||
provider_key_auth_semantics, provider_key_can_refresh_oauth,
|
||||
provider_key_auth_semantics, provider_key_can_export_oauth, provider_key_can_refresh_oauth,
|
||||
provider_key_effective_api_formats,
|
||||
};
|
||||
use aether_admin::provider::pool as admin_provider_pool_pure;
|
||||
@@ -1228,12 +1228,17 @@ pub(super) fn build_admin_pool_key_payload(
|
||||
"can_refresh_oauth".to_string(),
|
||||
json!(provider_key_can_refresh_oauth(
|
||||
auth_semantics,
|
||||
provider_type,
|
||||
auth_config.as_ref()
|
||||
)),
|
||||
);
|
||||
payload.insert(
|
||||
"can_export_oauth".to_string(),
|
||||
json!(auth_semantics.can_export_oauth()),
|
||||
json!(provider_key_can_export_oauth(
|
||||
auth_semantics,
|
||||
provider_type,
|
||||
auth_config.as_ref()
|
||||
)),
|
||||
);
|
||||
payload.insert(
|
||||
"can_edit_oauth".to_string(),
|
||||
|
||||
+11
-3
@@ -8,7 +8,7 @@ use super::{
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::provider_key_auth::{
|
||||
provider_key_auth_config_is_agent_identity, provider_key_auth_config_uses_header_authorization,
|
||||
provider_key_auth_semantics, provider_key_can_refresh_oauth,
|
||||
provider_key_auth_semantics, provider_key_can_export_oauth, provider_key_can_refresh_oauth,
|
||||
};
|
||||
use crate::GatewayError;
|
||||
use aether_admin::provider::pool as admin_provider_pool_pure;
|
||||
@@ -156,8 +156,16 @@ pub(super) async fn build_admin_pool_resolve_selection_response(
|
||||
&provider_type,
|
||||
auth_config.as_ref(),
|
||||
),
|
||||
"can_refresh_oauth": provider_key_can_refresh_oauth(auth_semantics, auth_config.as_ref()),
|
||||
"can_export_oauth": auth_semantics.can_export_oauth(),
|
||||
"can_refresh_oauth": provider_key_can_refresh_oauth(
|
||||
auth_semantics,
|
||||
&provider_type,
|
||||
auth_config.as_ref(),
|
||||
),
|
||||
"can_export_oauth": provider_key_can_export_oauth(
|
||||
auth_semantics,
|
||||
&provider_type,
|
||||
auth_config.as_ref(),
|
||||
),
|
||||
"can_edit_oauth": auth_semantics.can_edit_oauth(),
|
||||
"oauth_header_auth": auth_semantics.oauth_managed()
|
||||
&& provider_key_auth_config_uses_header_authorization(auth_config.as_ref()),
|
||||
|
||||
@@ -20,6 +20,8 @@ pub(crate) use self::endpoint_keys::{
|
||||
admin_reset_cycle_stats_key_id, admin_reveal_key_id, admin_update_key_id,
|
||||
};
|
||||
pub(crate) use self::oauth::{
|
||||
admin_provider_oauth_agent_identity_import_task_path,
|
||||
admin_provider_oauth_agent_identity_import_task_provider_id,
|
||||
admin_provider_oauth_batch_import_provider_id, admin_provider_oauth_batch_import_task_path,
|
||||
admin_provider_oauth_batch_import_task_provider_id, admin_provider_oauth_complete_key_id,
|
||||
admin_provider_oauth_complete_provider_id, admin_provider_oauth_device_authorize_provider_id,
|
||||
|
||||
@@ -44,6 +44,12 @@ pub(crate) fn admin_provider_oauth_batch_import_task_provider_id(
|
||||
provider_oauth_provider_id_for_suffix(request_path, "/batch-import/tasks")
|
||||
}
|
||||
|
||||
pub(crate) fn admin_provider_oauth_agent_identity_import_task_provider_id(
|
||||
request_path: &str,
|
||||
) -> Option<String> {
|
||||
provider_oauth_provider_id_for_suffix(request_path, "/agent-identity-import/tasks")
|
||||
}
|
||||
|
||||
pub(crate) fn admin_provider_oauth_batch_import_task_path(
|
||||
request_path: &str,
|
||||
) -> Option<(String, String)> {
|
||||
@@ -62,6 +68,25 @@ pub(crate) fn admin_provider_oauth_batch_import_task_path(
|
||||
Some((provider_id.to_string(), task_path.to_string()))
|
||||
}
|
||||
|
||||
pub(crate) fn admin_provider_oauth_agent_identity_import_task_path(
|
||||
request_path: &str,
|
||||
) -> Option<(String, String)> {
|
||||
let suffix = request_path
|
||||
.strip_prefix("/api/admin/provider-oauth/providers/")?
|
||||
.strip_suffix("/")
|
||||
.unwrap_or(request_path.strip_prefix("/api/admin/provider-oauth/providers/")?);
|
||||
let (provider_id, task_path) = suffix.split_once("/agent-identity-import/tasks/")?;
|
||||
if provider_id.is_empty()
|
||||
|| provider_id.contains('/')
|
||||
|| task_path.is_empty()
|
||||
|| task_path.contains('/')
|
||||
|| !task_path.starts_with("agent-identity-")
|
||||
{
|
||||
return None;
|
||||
}
|
||||
Some((provider_id.to_string(), task_path.to_string()))
|
||||
}
|
||||
|
||||
pub(crate) fn admin_provider_oauth_device_authorize_provider_id(
|
||||
request_path: &str,
|
||||
) -> Option<String> {
|
||||
@@ -79,3 +104,39 @@ fn provider_oauth_provider_id_for_suffix(request_path: &str, suffix: &str) -> Op
|
||||
.filter(|provider_id| !provider_id.is_empty() && !provider_id.contains('/'))
|
||||
.map(ToOwned::to_owned)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
admin_provider_oauth_agent_identity_import_task_path,
|
||||
admin_provider_oauth_agent_identity_import_task_provider_id,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn parses_dedicated_agent_identity_import_task_paths() {
|
||||
assert_eq!(
|
||||
admin_provider_oauth_agent_identity_import_task_provider_id(
|
||||
"/api/admin/provider-oauth/providers/provider-codex/agent-identity-import/tasks",
|
||||
)
|
||||
.as_deref(),
|
||||
Some("provider-codex")
|
||||
);
|
||||
assert_eq!(
|
||||
admin_provider_oauth_agent_identity_import_task_path(
|
||||
"/api/admin/provider-oauth/providers/provider-codex/agent-identity-import/tasks/agent-identity-task-1",
|
||||
),
|
||||
Some((
|
||||
"provider-codex".to_string(),
|
||||
"agent-identity-task-1".to_string(),
|
||||
))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn dedicated_status_path_rejects_generic_batch_task_ids() {
|
||||
assert!(admin_provider_oauth_agent_identity_import_task_path(
|
||||
"/api/admin/provider-oauth/providers/provider-codex/agent-identity-import/tasks/task-1",
|
||||
)
|
||||
.is_none());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -52,17 +52,14 @@ pub(crate) async fn build_admin_create_provider_key_record(
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.cloned();
|
||||
|
||||
if auth_type == "oauth"
|
||||
&& provider.provider_type.trim().eq_ignore_ascii_case("codex")
|
||||
&& auth_config
|
||||
.as_ref()
|
||||
.is_some_and(aether_provider_transport::is_codex_agent_identity_auth_config_value)
|
||||
if auth_config
|
||||
.as_ref()
|
||||
.is_some_and(aether_provider_transport::is_codex_agent_identity_auth_config_value)
|
||||
{
|
||||
aether_provider_transport::validate_codex_agent_identity_auth_config(
|
||||
auth_config
|
||||
.as_ref()
|
||||
.expect("Agent Identity auth_config was checked"),
|
||||
)?;
|
||||
return Err(
|
||||
"Agent Identity 凭据必须通过专属创建或导入接口管理,不能通过通用 Key 接口写入"
|
||||
.to_string(),
|
||||
);
|
||||
}
|
||||
|
||||
match auth_type.as_str() {
|
||||
|
||||
@@ -78,17 +78,14 @@ pub(crate) fn build_admin_update_provider_key_record_with_existing_keys(
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.cloned();
|
||||
|
||||
if target_auth_type == "oauth"
|
||||
&& provider.provider_type.trim().eq_ignore_ascii_case("codex")
|
||||
&& auth_config
|
||||
.as_ref()
|
||||
.is_some_and(aether_provider_transport::is_codex_agent_identity_auth_config_value)
|
||||
if auth_config
|
||||
.as_ref()
|
||||
.is_some_and(aether_provider_transport::is_codex_agent_identity_auth_config_value)
|
||||
{
|
||||
aether_provider_transport::validate_codex_agent_identity_auth_config(
|
||||
auth_config
|
||||
.as_ref()
|
||||
.expect("Agent Identity auth_config was checked"),
|
||||
)?;
|
||||
return Err(
|
||||
"Agent Identity 凭据必须通过专属创建或导入接口管理,不能通过通用 Key 接口写入"
|
||||
.to_string(),
|
||||
);
|
||||
}
|
||||
|
||||
match target_auth_type.as_str() {
|
||||
|
||||
@@ -31,6 +31,16 @@ pub(crate) fn build_admin_reveal_key_payload(
|
||||
key: &StoredProviderCatalogKey,
|
||||
) -> Result<serde_json::Value, String> {
|
||||
let parsed_auth_config = state.parse_catalog_auth_config_json(key);
|
||||
if parsed_auth_config.as_ref().is_some_and(|auth_config| {
|
||||
aether_provider_transport::is_codex_agent_identity_auth_config_value(
|
||||
&serde_json::Value::Object(auth_config.clone()),
|
||||
)
|
||||
}) {
|
||||
return Err(
|
||||
"Agent Identity 凭据不能通过通用 Key 查看接口读取,请使用专属 provider-oauth 管理面"
|
||||
.to_string(),
|
||||
);
|
||||
}
|
||||
let provider_type = reveal_provider_type_from_auth_config(parsed_auth_config.as_ref());
|
||||
let auth_semantics = provider_key_auth_semantics(key, provider_type.as_str());
|
||||
let auth_type = if auth_semantics.oauth_managed() {
|
||||
@@ -183,6 +193,15 @@ pub(crate) async fn build_admin_export_key_payload(
|
||||
.and_then(|value| value.as_object().cloned())
|
||||
.ok_or_else(|| "无法解密认证配置".to_string())?;
|
||||
|
||||
if aether_provider_transport::is_codex_agent_identity_auth_config_value(
|
||||
&serde_json::Value::Object(auth_config.clone()),
|
||||
) {
|
||||
return Err(
|
||||
"Agent Identity 凭据不能通过通用 Key 导出接口导出,请使用专属 provider-oauth 管理面"
|
||||
.to_string(),
|
||||
);
|
||||
}
|
||||
|
||||
let provider_type_from_config = auth_config
|
||||
.get("provider_type")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
|
||||
@@ -69,11 +69,16 @@ impl<'a> AdminAppState<'a> {
|
||||
)
|
||||
}
|
||||
|
||||
pub(crate) fn masked_catalog_api_key(
|
||||
pub(crate) fn masked_catalog_api_key_for_provider(
|
||||
&self,
|
||||
key: &aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey,
|
||||
provider_type: &str,
|
||||
) -> String {
|
||||
crate::handlers::admin::shared::masked_catalog_api_key(self.app, key)
|
||||
crate::handlers::admin::shared::masked_catalog_api_key_for_provider(
|
||||
self.app,
|
||||
key,
|
||||
provider_type,
|
||||
)
|
||||
}
|
||||
|
||||
pub(crate) async fn build_admin_provider_keys_payload(
|
||||
|
||||
@@ -227,6 +227,7 @@ impl<'a> AdminAppState<'a> {
|
||||
.compare_and_update_provider_catalog_key_adaptive_state(
|
||||
&ProviderCatalogKeyAdaptiveStateUpdate {
|
||||
key_id: key_id.to_string(),
|
||||
expected_encrypted_auth_config: None,
|
||||
expected,
|
||||
next,
|
||||
status_snapshot_patch: serde_json::json!({
|
||||
@@ -257,6 +258,33 @@ impl<'a> AdminAppState<'a> {
|
||||
) -> Result<
|
||||
Option<aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey>,
|
||||
GatewayError,
|
||||
> {
|
||||
self.reset_provider_catalog_key_recovery_state_inner(key_id, None)
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn reset_provider_catalog_key_recovery_state_fenced(
|
||||
&self,
|
||||
key_id: &str,
|
||||
expected_encrypted_auth_config: &str,
|
||||
) -> Result<
|
||||
Option<aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey>,
|
||||
GatewayError,
|
||||
> {
|
||||
self.reset_provider_catalog_key_recovery_state_inner(
|
||||
key_id,
|
||||
Some(expected_encrypted_auth_config),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn reset_provider_catalog_key_recovery_state_inner(
|
||||
&self,
|
||||
key_id: &str,
|
||||
expected_auth_config: Option<&str>,
|
||||
) -> Result<
|
||||
Option<aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey>,
|
||||
GatewayError,
|
||||
> {
|
||||
use aether_data_contracts::repository::provider_catalog::ProviderCatalogKeyHealthStateUpdate;
|
||||
|
||||
@@ -271,6 +299,11 @@ impl<'a> AdminAppState<'a> {
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
if expected_auth_config
|
||||
.is_some_and(|expected| current.encrypted_auth_config.as_deref() != Some(expected))
|
||||
{
|
||||
return Ok(None);
|
||||
}
|
||||
if current.health_by_format.as_ref() == Some(&empty)
|
||||
&& current.circuit_breaker_by_format.as_ref() == Some(&empty)
|
||||
{
|
||||
@@ -282,6 +315,7 @@ impl<'a> AdminAppState<'a> {
|
||||
.compare_and_update_provider_catalog_key_health_state(
|
||||
&ProviderCatalogKeyHealthStateUpdate {
|
||||
key_id: key_id.to_string(),
|
||||
expected_encrypted_auth_config: expected_auth_config.map(ToOwned::to_owned),
|
||||
expected_health_by_format: current.health_by_format,
|
||||
expected_circuit_breaker_by_format: current.circuit_breaker_by_format,
|
||||
health_by_format: Some(empty.clone()),
|
||||
@@ -299,15 +333,24 @@ impl<'a> AdminAppState<'a> {
|
||||
"provider key {key_id} health state changed repeatedly while resetting OAuth recovery state"
|
||||
)));
|
||||
}
|
||||
if !self.reset_provider_catalog_key_error_count(key_id).await? {
|
||||
if expected_auth_config.is_none()
|
||||
&& !self.reset_provider_catalog_key_error_count(key_id).await?
|
||||
{
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
Ok(self
|
||||
let current = self
|
||||
.read_provider_catalog_keys_by_ids(&[key_id.to_string()])
|
||||
.await?
|
||||
.into_iter()
|
||||
.next())
|
||||
.next();
|
||||
if current.as_ref().is_some_and(|key| {
|
||||
expected_auth_config
|
||||
.is_some_and(|expected| key.encrypted_auth_config.as_deref() != Some(expected))
|
||||
}) {
|
||||
return Ok(None);
|
||||
}
|
||||
Ok(current)
|
||||
}
|
||||
|
||||
pub(crate) async fn update_provider_catalog_key_status_snapshot(
|
||||
|
||||
@@ -90,6 +90,7 @@ impl<'a> AdminAppState<'a> {
|
||||
provider_id: &str,
|
||||
provider_type: &str,
|
||||
pkce_verifier: Option<&str>,
|
||||
expected_encrypted_auth_config: Option<&str>,
|
||||
) -> Result<String, GatewayError> {
|
||||
let nonce = aether_admin::provider::state::generate_provider_oauth_nonce();
|
||||
let payload = json!({
|
||||
@@ -98,6 +99,7 @@ impl<'a> AdminAppState<'a> {
|
||||
"provider_id": provider_id,
|
||||
"provider_type": provider_type,
|
||||
"pkce_verifier": pkce_verifier,
|
||||
"expected_encrypted_auth_config": expected_encrypted_auth_config,
|
||||
"created_at": aether_admin::provider::state::current_unix_secs(),
|
||||
});
|
||||
let key = provider_oauth_state_storage_key(&nonce);
|
||||
|
||||
@@ -16,6 +16,10 @@ use serde_json::{json, Map, Value};
|
||||
use std::collections::{BTreeMap, BTreeSet};
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
fn provider_skips_automatic_key_cleanup(provider: &StoredProviderCatalogProvider) -> bool {
|
||||
provider.provider_type.trim().eq_ignore_ascii_case("codex")
|
||||
}
|
||||
|
||||
impl<'a> AdminAppState<'a> {
|
||||
pub(crate) async fn clear_admin_provider_pool_cooldown(&self, provider_id: &str, key_id: &str) {
|
||||
crate::handlers::admin::provider::pool::runtime::clear_admin_provider_pool_cooldown(
|
||||
@@ -368,6 +372,13 @@ impl<'a> AdminAppState<'a> {
|
||||
) -> Result<usize, GatewayError> {
|
||||
use aether_admin::provider::pool as admin_provider_pool_pure;
|
||||
|
||||
// Codex OAuth credentials can be replaced by a long-lived Agent
|
||||
// Identity under the same key id. Until deletes support an auth_config
|
||||
// CAS, automatic cleanup must retain every Codex key.
|
||||
if provider_skips_automatic_key_cleanup(provider) {
|
||||
return Ok(0);
|
||||
}
|
||||
|
||||
let banned_keys = self
|
||||
.list_provider_catalog_keys_by_provider_ids(std::slice::from_ref(&provider.id))
|
||||
.await?
|
||||
@@ -408,6 +419,10 @@ impl<'a> AdminAppState<'a> {
|
||||
) -> Result<usize, GatewayError> {
|
||||
use aether_admin::provider::pool as admin_provider_pool_pure;
|
||||
|
||||
if provider_skips_automatic_key_cleanup(provider) {
|
||||
return Ok(0);
|
||||
}
|
||||
|
||||
let keys = self
|
||||
.list_provider_catalog_keys_by_provider_ids(std::slice::from_ref(&provider.id))
|
||||
.await?;
|
||||
@@ -849,3 +864,26 @@ impl<'a> AdminAppState<'a> {
|
||||
.into_response())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod automatic_cleanup_tests {
|
||||
use super::provider_skips_automatic_key_cleanup;
|
||||
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogProvider;
|
||||
|
||||
fn provider(provider_type: &str) -> StoredProviderCatalogProvider {
|
||||
StoredProviderCatalogProvider::new(
|
||||
format!("provider-{provider_type}"),
|
||||
provider_type.to_string(),
|
||||
None,
|
||||
provider_type.to_string(),
|
||||
)
|
||||
.expect("provider should build")
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn codex_automatic_cleanup_is_disabled_for_replaceable_agent_credentials() {
|
||||
assert!(provider_skips_automatic_key_cleanup(&provider("codex")));
|
||||
assert!(provider_skips_automatic_key_cleanup(&provider("CoDeX")));
|
||||
assert!(!provider_skips_automatic_key_cleanup(&provider("kiro")));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -153,6 +153,7 @@ impl<'a> AdminAppState<'a> {
|
||||
.compare_and_update_provider_catalog_key_adaptive_state(
|
||||
&ProviderCatalogKeyAdaptiveStateUpdate {
|
||||
key_id: key.id.clone(),
|
||||
expected_encrypted_auth_config: None,
|
||||
expected,
|
||||
next,
|
||||
status_snapshot_patch: json!({
|
||||
|
||||
@@ -11,10 +11,10 @@ pub(crate) use crate::handlers::shared::{
|
||||
attach_admin_audit_response, build_admin_provider_key_response,
|
||||
decrypt_catalog_secret_with_fallbacks, default_provider_key_status_snapshot,
|
||||
effective_catalog_encryption_key, encrypt_catalog_secret_with_fallbacks, json_string_list,
|
||||
masked_catalog_api_key, normalize_json_array, normalize_json_object, normalize_string_list,
|
||||
parse_catalog_auth_config_json, provider_catalog_key_supports_format,
|
||||
provider_key_health_summary, provider_key_health_summary_at,
|
||||
provider_key_status_snapshot_payload, query_param_bool, query_param_optional_bool,
|
||||
query_param_value, take_secret_prefix, take_secret_suffix, unix_secs_to_rfc3339,
|
||||
OFFICIAL_EXTERNAL_MODEL_PROVIDERS,
|
||||
masked_catalog_api_key, masked_catalog_api_key_for_provider, normalize_json_array,
|
||||
normalize_json_object, normalize_string_list, parse_catalog_auth_config_json,
|
||||
provider_catalog_key_supports_format, provider_key_health_summary,
|
||||
provider_key_health_summary_at, provider_key_status_snapshot_payload, query_param_bool,
|
||||
query_param_optional_bool, query_param_value, take_secret_prefix, take_secret_suffix,
|
||||
unix_secs_to_rfc3339, OFFICIAL_EXTERNAL_MODEL_PROVIDERS,
|
||||
};
|
||||
|
||||
@@ -1,7 +1,8 @@
|
||||
use super::enabled_key_capability_short_names;
|
||||
use crate::handlers::shared::{parse_catalog_auth_config_json, unix_secs_to_rfc3339};
|
||||
use crate::provider_key_auth::{
|
||||
provider_key_auth_config_uses_header_authorization, provider_key_effective_api_formats,
|
||||
provider_key_auth_config_is_agent_identity, provider_key_auth_config_uses_header_authorization,
|
||||
provider_key_effective_api_formats,
|
||||
};
|
||||
use crate::AppState;
|
||||
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
|
||||
@@ -10,13 +11,18 @@ use serde_json::json;
|
||||
use std::collections::{BTreeMap, HashMap};
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
fn grouped_key_masked_label(state: &AppState, key: &StoredProviderCatalogKey) -> &'static str {
|
||||
fn grouped_key_masked_label(
|
||||
state: &AppState,
|
||||
key: &StoredProviderCatalogKey,
|
||||
provider_type: &str,
|
||||
) -> &'static str {
|
||||
match key.auth_type.trim() {
|
||||
"service_account" | "vertex_ai" => "[Service Account]",
|
||||
"oauth" => {
|
||||
if provider_key_auth_config_uses_header_authorization(
|
||||
parse_catalog_auth_config_json(state, key).as_ref(),
|
||||
) {
|
||||
let auth_config = parse_catalog_auth_config_json(state, key);
|
||||
if provider_key_auth_config_is_agent_identity(provider_type, auth_config.as_ref()) {
|
||||
"[Agent Identity]"
|
||||
} else if provider_key_auth_config_uses_header_authorization(auth_config.as_ref()) {
|
||||
"[OAuth Header]"
|
||||
} else {
|
||||
"[OAuth Token]"
|
||||
@@ -161,7 +167,7 @@ pub(crate) async fn build_admin_keys_grouped_by_format_payload(
|
||||
"provider_id": key.provider_id,
|
||||
"name": key.name,
|
||||
"auth_type": key.auth_type,
|
||||
"api_key_masked": grouped_key_masked_label(state, &key),
|
||||
"api_key_masked": grouped_key_masked_label(state, &key, provider_type),
|
||||
"internal_priority": key.internal_priority,
|
||||
"global_priority_by_format": key.global_priority_by_format,
|
||||
"rate_multipliers": key.rate_multipliers,
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
use crate::handlers::shared::{json_string_list, unix_secs_to_rfc3339};
|
||||
use crate::provider_key_auth::{
|
||||
provider_key_auth_config_is_agent_identity, provider_key_auth_config_uses_header_authorization,
|
||||
provider_key_auth_semantics, provider_key_can_refresh_oauth,
|
||||
provider_key_auth_semantics, provider_key_can_export_oauth, provider_key_can_refresh_oauth,
|
||||
provider_key_configured_api_formats, provider_key_inherits_provider_api_formats,
|
||||
};
|
||||
use crate::AppState;
|
||||
@@ -168,6 +168,19 @@ pub(crate) fn masked_catalog_api_key(state: &AppState, key: &StoredProviderCatal
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn masked_catalog_api_key_for_provider(
|
||||
state: &AppState,
|
||||
key: &StoredProviderCatalogKey,
|
||||
provider_type: &str,
|
||||
) -> String {
|
||||
let auth_config = parse_catalog_auth_config_json(state, key);
|
||||
if provider_key_auth_config_is_agent_identity(provider_type, auth_config.as_ref()) {
|
||||
"[Agent Identity]".to_string()
|
||||
} else {
|
||||
masked_catalog_api_key(state, key)
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn parse_catalog_auth_config_json(
|
||||
state: &AppState,
|
||||
key: &StoredProviderCatalogKey,
|
||||
@@ -2496,11 +2509,11 @@ pub(crate) fn build_admin_provider_key_response(
|
||||
);
|
||||
payload.insert(
|
||||
"api_key_masked".to_string(),
|
||||
json!(if agent_identity {
|
||||
"[Agent Identity]".to_string()
|
||||
} else {
|
||||
masked_catalog_api_key(state, key)
|
||||
}),
|
||||
json!(masked_catalog_api_key_for_provider(
|
||||
state,
|
||||
key,
|
||||
provider_type,
|
||||
)),
|
||||
);
|
||||
payload.insert("api_key_plain".to_string(), serde_json::Value::Null);
|
||||
payload.insert("auth_type".to_string(), json!(key.auth_type));
|
||||
@@ -2529,12 +2542,17 @@ pub(crate) fn build_admin_provider_key_response(
|
||||
"can_refresh_oauth".to_string(),
|
||||
json!(provider_key_can_refresh_oauth(
|
||||
auth_semantics,
|
||||
provider_type,
|
||||
auth_config.as_ref()
|
||||
)),
|
||||
);
|
||||
payload.insert(
|
||||
"can_export_oauth".to_string(),
|
||||
json!(auth_semantics.can_export_oauth()),
|
||||
json!(provider_key_can_export_oauth(
|
||||
auth_semantics,
|
||||
provider_type,
|
||||
auth_config.as_ref()
|
||||
)),
|
||||
);
|
||||
payload.insert(
|
||||
"can_edit_oauth".to_string(),
|
||||
@@ -2874,6 +2892,46 @@ mod tests {
|
||||
assert_ne!(masked, "***ERROR***");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_aware_mask_labels_agent_identity_without_exposing_placeholder() {
|
||||
let state = AppState::new().expect("gateway should build");
|
||||
let encrypted_placeholder =
|
||||
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "__placeholder__")
|
||||
.expect("placeholder ciphertext should build");
|
||||
let encrypted_auth_config = encrypt_python_fernet_plaintext(
|
||||
DEVELOPMENT_ENCRYPTION_KEY,
|
||||
r#"{"provider_type":"codex","auth_mode":"agentIdentity","agent_runtime_id":"runtime-1","agent_private_key":"base64-private-key","task_id":"task-1"}"#,
|
||||
)
|
||||
.expect("auth config ciphertext should build");
|
||||
let key = StoredProviderCatalogKey::new(
|
||||
"key-agent".to_string(),
|
||||
"provider-codex".to_string(),
|
||||
"agent".to_string(),
|
||||
"oauth".to_string(),
|
||||
None,
|
||||
true,
|
||||
)
|
||||
.expect("key should build")
|
||||
.with_transport_fields(
|
||||
Some(json!(["openai:responses"])),
|
||||
encrypted_placeholder,
|
||||
Some(encrypted_auth_config),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.expect("key transport should build");
|
||||
|
||||
assert_eq!(
|
||||
masked_catalog_api_key_for_provider(&state, &key, "codex"),
|
||||
"[Agent Identity]"
|
||||
);
|
||||
assert!(!masked_catalog_api_key_for_provider(&state, &key, "codex").contains("placeholder"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_key_status_snapshot_payload_backfills_missing_quota_from_upstream_metadata() {
|
||||
let mut key = sample_catalog_key();
|
||||
|
||||
@@ -24,7 +24,8 @@ pub(crate) use self::api_keys::{
|
||||
pub(crate) use self::catalog::{
|
||||
build_admin_provider_key_response, decrypt_catalog_secret_with_fallbacks,
|
||||
default_provider_key_status_snapshot, effective_catalog_encryption_key,
|
||||
encrypt_catalog_secret_with_fallbacks, masked_catalog_api_key, parse_catalog_auth_config_json,
|
||||
encrypt_catalog_secret_with_fallbacks, masked_catalog_api_key,
|
||||
masked_catalog_api_key_for_provider, parse_catalog_auth_config_json,
|
||||
provider_catalog_key_supports_format, provider_key_health_summary,
|
||||
provider_key_health_summary_at, provider_key_status_snapshot_payload,
|
||||
sync_provider_key_oauth_status_snapshot, sync_provider_key_quota_status_snapshot,
|
||||
|
||||
@@ -14,6 +14,7 @@ use crate::{AppState, GatewayError};
|
||||
use super::system_config_bool;
|
||||
|
||||
const OAUTH_TOKEN_REFRESH_LOOKAHEAD_SECS: u64 = 120;
|
||||
const OAUTH_REFRESH_FAILED_PREFIX: &str = "[REFRESH_FAILED] ";
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, serde::Serialize)]
|
||||
pub(crate) struct OAuthTokenRefreshRunSummary {
|
||||
@@ -91,13 +92,34 @@ pub(crate) async fn perform_oauth_token_refresh_once(
|
||||
summary.skipped = summary.skipped.saturating_add(1);
|
||||
continue;
|
||||
};
|
||||
if !auth_config_has_refresh_token(transport.key.decrypted_auth_config.as_deref()) {
|
||||
let is_agent_identity =
|
||||
crate::provider_transport::is_codex_agent_identity_transport(&transport);
|
||||
let needs_agent_task_recovery = is_agent_identity
|
||||
&& agent_identity_needs_task_recovery(
|
||||
transport.key.decrypted_auth_config.as_deref(),
|
||||
key.oauth_invalid_reason.as_deref(),
|
||||
);
|
||||
if !needs_agent_task_recovery
|
||||
&& !auth_config_has_refresh_token(transport.key.decrypted_auth_config.as_deref())
|
||||
{
|
||||
summary.skipped = summary.skipped.saturating_add(1);
|
||||
continue;
|
||||
}
|
||||
|
||||
match state.resolve_local_oauth_request_auth(&transport).await {
|
||||
Ok(Some(_auth)) => {
|
||||
let refresh_result = if needs_agent_task_recovery {
|
||||
state
|
||||
.force_local_oauth_refresh_entry(&transport)
|
||||
.await
|
||||
.map(|entry| entry.map(|_| ()))
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))
|
||||
} else {
|
||||
state
|
||||
.resolve_local_oauth_request_auth(&transport)
|
||||
.await
|
||||
.map(|auth| auth.map(|_| ()))
|
||||
};
|
||||
match refresh_result {
|
||||
Ok(Some(())) => {
|
||||
summary.resolved = summary.resolved.saturating_add(1);
|
||||
if provider_key_credentials_changed(state, key).await? {
|
||||
summary.refreshed = summary.refreshed.saturating_add(1);
|
||||
@@ -171,19 +193,45 @@ fn oauth_refresh_candidate(
|
||||
key: &StoredProviderCatalogKey,
|
||||
refresh_cutoff_unix_secs: u64,
|
||||
) -> bool {
|
||||
key.is_active
|
||||
&& key.oauth_invalid_at_unix_secs.is_none()
|
||||
&& key
|
||||
.encrypted_auth_config
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.is_some_and(|value| !value.is_empty())
|
||||
let has_auth_config = key
|
||||
.encrypted_auth_config
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.is_some_and(|value| !value.is_empty());
|
||||
let regular_oauth_candidate = key.oauth_invalid_at_unix_secs.is_none()
|
||||
&& key
|
||||
.expires_at_unix_secs
|
||||
.is_some_and(|expires_at| expires_at <= refresh_cutoff_unix_secs)
|
||||
.is_some_and(|expires_at| expires_at <= refresh_cutoff_unix_secs);
|
||||
// The catalog row is encrypted here, so exact Agent Identity validation is
|
||||
// deferred until the transport snapshot has decrypted auth_config.
|
||||
let possible_agent_candidate = provider.provider_type.trim().eq_ignore_ascii_case("codex")
|
||||
&& key.auth_type.trim().eq_ignore_ascii_case("oauth")
|
||||
&& (key.expires_at_unix_secs.is_none()
|
||||
|| key
|
||||
.oauth_invalid_reason
|
||||
.as_deref()
|
||||
.is_some_and(|reason| reason.contains(OAUTH_REFRESH_FAILED_PREFIX)));
|
||||
key.is_active
|
||||
&& has_auth_config
|
||||
&& (regular_oauth_candidate || possible_agent_candidate)
|
||||
&& provider_key_is_oauth_managed(key, provider.provider_type.as_str())
|
||||
}
|
||||
|
||||
fn agent_identity_needs_task_recovery(
|
||||
auth_config: Option<&str>,
|
||||
oauth_invalid_reason: Option<&str>,
|
||||
) -> bool {
|
||||
if oauth_invalid_reason.is_some_and(|reason| reason.contains(OAUTH_REFRESH_FAILED_PREFIX)) {
|
||||
return true;
|
||||
}
|
||||
auth_config
|
||||
.and_then(|value| serde_json::from_str::<Value>(value).ok())
|
||||
.is_some_and(|config| {
|
||||
crate::provider_transport::is_codex_agent_identity_auth_config_value(&config)
|
||||
&& !crate::provider_transport::codex_agent_identity_auth_config_has_task_id(&config)
|
||||
})
|
||||
}
|
||||
|
||||
async fn provider_key_credentials_changed(
|
||||
state: &AppState,
|
||||
before: &StoredProviderCatalogKey,
|
||||
@@ -222,3 +270,29 @@ fn now_unix_secs() -> u64 {
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or_default()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::agent_identity_needs_task_recovery;
|
||||
|
||||
#[test]
|
||||
fn pending_agent_identity_without_task_is_recoverable() {
|
||||
let config = serde_json::json!({
|
||||
"auth_mode": "agentIdentity",
|
||||
"agent_runtime_id": "runtime-1",
|
||||
"agent_private_key": "private-key-present",
|
||||
});
|
||||
assert!(agent_identity_needs_task_recovery(
|
||||
Some(&config.to_string()),
|
||||
None,
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn refresh_failure_marker_forces_agent_task_recovery() {
|
||||
assert!(agent_identity_needs_task_recovery(
|
||||
Some("{}"),
|
||||
Some("[REFRESH_FAILED] temporary"),
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -44,9 +44,7 @@ use crate::handlers::shared::provider_pool::{
|
||||
use crate::orchestration::local_execution_candidate_metadata_from_report_context;
|
||||
use crate::scheduler::affinity::SCHEDULER_AFFINITY_TTL;
|
||||
use crate::scheduler::config::{read_scheduler_ordering_config, SchedulerSchedulingMode};
|
||||
use crate::{
|
||||
provider_transport::snapshot::GatewayProviderTransportProvider, AppState, GatewayError,
|
||||
};
|
||||
use crate::AppState;
|
||||
|
||||
const POOL_SCORE_FEEDBACK_GATE_MAX_ENTRIES: usize = 50_000;
|
||||
const HEALTH_SUCCESS_PERSIST_GATE_MAX_ENTRIES: usize = 50_000;
|
||||
@@ -253,6 +251,21 @@ struct PoolFeedbackContext {
|
||||
sticky_session_token: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
enum LocalExecutionAuthConfigFence {
|
||||
Unfenced,
|
||||
Fenced(String),
|
||||
}
|
||||
|
||||
impl LocalExecutionAuthConfigFence {
|
||||
fn encrypted_auth_config(&self) -> Option<&str> {
|
||||
match self {
|
||||
Self::Unfenced => None,
|
||||
Self::Fenced(ciphertext) => Some(ciphertext),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
const ADAPTIVE_RPM_RECENT_CANDIDATE_LIMIT: usize = 512;
|
||||
const LOCAL_EXECUTION_SCHEDULER_AFFINITY_MAX_ENTRIES: usize = 10_000;
|
||||
|
||||
@@ -405,6 +418,73 @@ async fn local_execution_plan_uses_pool(state: &AppState, plan: &ExecutionPlan)
|
||||
admin_provider_pool_config_from_config_value(transport.provider.config.as_ref()).is_some()
|
||||
}
|
||||
|
||||
async fn capture_local_execution_auth_config_fence(
|
||||
state: &AppState,
|
||||
plan: &ExecutionPlan,
|
||||
) -> Option<LocalExecutionAuthConfigFence> {
|
||||
let transport = match state
|
||||
.read_provider_transport_snapshot(&plan.provider_id, &plan.endpoint_id, &plan.key_id)
|
||||
.await
|
||||
{
|
||||
Ok(Some(transport)) => transport,
|
||||
Ok(None) => return None,
|
||||
Err(err) => {
|
||||
warn!(
|
||||
provider_id = %plan.provider_id,
|
||||
endpoint_id = %plan.endpoint_id,
|
||||
key_id = %plan.key_id,
|
||||
error = ?err,
|
||||
"gateway orchestration effects: failed to read transport for credential fencing"
|
||||
);
|
||||
return None;
|
||||
}
|
||||
};
|
||||
if !transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("codex")
|
||||
|| !transport.key.auth_type.trim().eq_ignore_ascii_case("oauth")
|
||||
{
|
||||
return Some(LocalExecutionAuthConfigFence::Unfenced);
|
||||
}
|
||||
|
||||
let authorization = execution_plan_authorization(plan)?;
|
||||
let current_uses_agent_identity =
|
||||
crate::provider_transport::is_codex_agent_identity_transport(&transport);
|
||||
let authorization_matches = if current_uses_agent_identity {
|
||||
crate::provider_transport::codex_agent_identity_authorization_matches_transport(
|
||||
&transport,
|
||||
authorization,
|
||||
)
|
||||
} else if crate::provider_transport::is_codex_agent_identity_authorization(authorization) {
|
||||
false
|
||||
} else {
|
||||
execution_plan_bearer_matches_transport(plan, &transport)
|
||||
};
|
||||
if !authorization_matches {
|
||||
return None;
|
||||
}
|
||||
|
||||
match state
|
||||
.capture_provider_transport_auth_config_fence(&transport)
|
||||
.await
|
||||
{
|
||||
Ok(Some(ciphertext)) => Some(LocalExecutionAuthConfigFence::Fenced(ciphertext)),
|
||||
Ok(None) => None,
|
||||
Err(err) => {
|
||||
warn!(
|
||||
provider_id = %plan.provider_id,
|
||||
endpoint_id = %plan.endpoint_id,
|
||||
key_id = %plan.key_id,
|
||||
error = ?err,
|
||||
"gateway orchestration effects: failed to capture credential fence"
|
||||
);
|
||||
None
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn local_scheduler_affinity_matches_failed_target(
|
||||
state: &AppState,
|
||||
plan: &ExecutionPlan,
|
||||
@@ -480,6 +560,7 @@ async fn resolve_pool_feedback_context(
|
||||
context: LocalExecutionEffectContext<'_>,
|
||||
) -> Option<PoolFeedbackContext> {
|
||||
let plan = context.plan;
|
||||
capture_local_execution_auth_config_fence(state, plan).await?;
|
||||
let transport = match state
|
||||
.read_provider_transport_snapshot(&plan.provider_id, &plan.endpoint_id, &plan.key_id)
|
||||
.await
|
||||
@@ -602,6 +683,11 @@ async fn record_adaptive_rate_limit_effect(
|
||||
context: LocalExecutionEffectContext<'_>,
|
||||
effect: LocalAdaptiveRateLimitEffect<'_>,
|
||||
) {
|
||||
let Some(auth_config_fence) =
|
||||
capture_local_execution_auth_config_fence(state, context.plan).await
|
||||
else {
|
||||
return;
|
||||
};
|
||||
let effect_lock = PROVIDER_KEY_EFFECT_LOCKS.lock_for(&context.plan.key_id);
|
||||
let _effect_guard = effect_lock.lock().await;
|
||||
let observed_at_unix_secs = current_unix_secs();
|
||||
@@ -626,6 +712,12 @@ async fn record_adaptive_rate_limit_effect(
|
||||
else {
|
||||
return;
|
||||
};
|
||||
if auth_config_fence
|
||||
.encrypted_auth_config()
|
||||
.is_some_and(|expected| current_key.encrypted_auth_config.as_deref() != Some(expected))
|
||||
{
|
||||
return;
|
||||
}
|
||||
let Some(projection) = project_local_adaptive_rate_limit(
|
||||
¤t_key,
|
||||
effect.classification,
|
||||
@@ -648,6 +740,9 @@ async fn record_adaptive_rate_limit_effect(
|
||||
next.last_rpm_peak = projection.last_rpm_peak;
|
||||
let update = ProviderCatalogKeyAdaptiveStateUpdate {
|
||||
key_id: context.plan.key_id.clone(),
|
||||
expected_encrypted_auth_config: auth_config_fence
|
||||
.encrypted_auth_config()
|
||||
.map(ToOwned::to_owned),
|
||||
expected,
|
||||
next,
|
||||
status_snapshot_patch: adaptive_status_snapshot_patch(&projection.status_snapshot),
|
||||
@@ -705,6 +800,11 @@ async fn record_adaptive_success_effect(
|
||||
context: LocalExecutionEffectContext<'_>,
|
||||
_effect: LocalAdaptiveSuccessEffect,
|
||||
) {
|
||||
let Some(auth_config_fence) =
|
||||
capture_local_execution_auth_config_fence(state, context.plan).await
|
||||
else {
|
||||
return;
|
||||
};
|
||||
let observed_at_unix_secs = current_unix_secs();
|
||||
let Some(current_key) = state
|
||||
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&context.plan.key_id))
|
||||
@@ -714,6 +814,12 @@ async fn record_adaptive_success_effect(
|
||||
else {
|
||||
return;
|
||||
};
|
||||
if auth_config_fence
|
||||
.encrypted_auth_config()
|
||||
.is_some_and(|expected| current_key.encrypted_auth_config.as_deref() != Some(expected))
|
||||
{
|
||||
return;
|
||||
}
|
||||
if current_key.rpm_limit.is_some()
|
||||
|| current_key
|
||||
.learned_rpm_limit
|
||||
@@ -757,6 +863,12 @@ async fn record_adaptive_success_effect(
|
||||
else {
|
||||
return;
|
||||
};
|
||||
if auth_config_fence
|
||||
.encrypted_auth_config()
|
||||
.is_some_and(|expected| current_key.encrypted_auth_config.as_deref() != Some(expected))
|
||||
{
|
||||
return;
|
||||
}
|
||||
if current_key.rpm_limit.is_some()
|
||||
|| current_key
|
||||
.learned_rpm_limit
|
||||
@@ -778,6 +890,9 @@ async fn record_adaptive_success_effect(
|
||||
next.last_probe_increase_at_unix_secs = projection.last_probe_increase_at_unix_secs;
|
||||
let update = ProviderCatalogKeyAdaptiveStateUpdate {
|
||||
key_id: context.plan.key_id.clone(),
|
||||
expected_encrypted_auth_config: auth_config_fence
|
||||
.encrypted_auth_config()
|
||||
.map(ToOwned::to_owned),
|
||||
expected,
|
||||
next,
|
||||
status_snapshot_patch: adaptive_status_snapshot_patch(&projection.status_snapshot),
|
||||
@@ -853,6 +968,11 @@ async fn record_health_failure_effect(
|
||||
if api_format.is_empty() {
|
||||
return;
|
||||
}
|
||||
let Some(auth_config_fence) =
|
||||
capture_local_execution_auth_config_fence(state, context.plan).await
|
||||
else {
|
||||
return;
|
||||
};
|
||||
|
||||
let effect_lock = PROVIDER_KEY_EFFECT_LOCKS.lock_for(&context.plan.key_id);
|
||||
let _effect_guard = effect_lock.lock().await;
|
||||
@@ -869,6 +989,12 @@ async fn record_health_failure_effect(
|
||||
else {
|
||||
return;
|
||||
};
|
||||
if auth_config_fence
|
||||
.encrypted_auth_config()
|
||||
.is_some_and(|expected| current_key.encrypted_auth_config.as_deref() != Some(expected))
|
||||
{
|
||||
return;
|
||||
}
|
||||
let Some(health_by_format) = project_local_failure_health(
|
||||
current_key.health_by_format.as_ref(),
|
||||
api_format,
|
||||
@@ -897,6 +1023,9 @@ async fn record_health_failure_effect(
|
||||
};
|
||||
let update = ProviderCatalogKeyHealthStateUpdate {
|
||||
key_id: context.plan.key_id.clone(),
|
||||
expected_encrypted_auth_config: auth_config_fence
|
||||
.encrypted_auth_config()
|
||||
.map(ToOwned::to_owned),
|
||||
expected_health_by_format: current_key.health_by_format,
|
||||
expected_circuit_breaker_by_format: current_key.circuit_breaker_by_format,
|
||||
health_by_format: Some(health_by_format),
|
||||
@@ -934,6 +1063,11 @@ async fn record_health_success_effect(
|
||||
if api_format.is_empty() {
|
||||
return;
|
||||
}
|
||||
let Some(auth_config_fence) =
|
||||
capture_local_execution_auth_config_fence(state, context.plan).await
|
||||
else {
|
||||
return;
|
||||
};
|
||||
|
||||
// Health updates replace both JSON snapshots in one write. Serialize the success
|
||||
// read/project/write with failure and circuit-clear effects for this provider key so a
|
||||
@@ -953,6 +1087,12 @@ async fn record_health_success_effect(
|
||||
else {
|
||||
return;
|
||||
};
|
||||
if auth_config_fence
|
||||
.encrypted_auth_config()
|
||||
.is_some_and(|expected| current_key.encrypted_auth_config.as_deref() != Some(expected))
|
||||
{
|
||||
return;
|
||||
}
|
||||
let Some(health_by_format) =
|
||||
project_local_success_health(current_key.health_by_format.as_ref(), api_format)
|
||||
else {
|
||||
@@ -991,6 +1131,9 @@ async fn record_health_success_effect(
|
||||
};
|
||||
let update = ProviderCatalogKeyHealthStateUpdate {
|
||||
key_id: context.plan.key_id.clone(),
|
||||
expected_encrypted_auth_config: auth_config_fence
|
||||
.encrypted_auth_config()
|
||||
.map(ToOwned::to_owned),
|
||||
expected_health_by_format: current_key.health_by_format,
|
||||
expected_circuit_breaker_by_format: current_key.circuit_breaker_by_format,
|
||||
health_by_format: Some(health_by_format),
|
||||
@@ -1104,6 +1247,12 @@ async fn record_pool_error_effect(
|
||||
};
|
||||
|
||||
clear_pool_key_circuit_breaker(state, context).await;
|
||||
if capture_local_execution_auth_config_fence(state, context.plan)
|
||||
.await
|
||||
.is_none()
|
||||
{
|
||||
return;
|
||||
}
|
||||
record_admin_provider_pool_error(
|
||||
state.runtime_state.as_ref(),
|
||||
&context.plan.provider_id,
|
||||
@@ -1135,6 +1284,11 @@ async fn clear_pool_key_circuit_breaker(
|
||||
state: &AppState,
|
||||
context: LocalExecutionEffectContext<'_>,
|
||||
) {
|
||||
let Some(auth_config_fence) =
|
||||
capture_local_execution_auth_config_fence(state, context.plan).await
|
||||
else {
|
||||
return;
|
||||
};
|
||||
let effect_lock = PROVIDER_KEY_EFFECT_LOCKS.lock_for(&context.plan.key_id);
|
||||
let _effect_guard = effect_lock.lock().await;
|
||||
|
||||
@@ -1147,11 +1301,20 @@ async fn clear_pool_key_circuit_breaker(
|
||||
else {
|
||||
return;
|
||||
};
|
||||
if auth_config_fence
|
||||
.encrypted_auth_config()
|
||||
.is_some_and(|expected| current_key.encrypted_auth_config.as_deref() != Some(expected))
|
||||
{
|
||||
return;
|
||||
}
|
||||
if current_key.circuit_breaker_by_format.is_none() {
|
||||
return;
|
||||
}
|
||||
let update = ProviderCatalogKeyHealthStateUpdate {
|
||||
key_id: context.plan.key_id.clone(),
|
||||
expected_encrypted_auth_config: auth_config_fence
|
||||
.encrypted_auth_config()
|
||||
.map(ToOwned::to_owned),
|
||||
expected_health_by_format: current_key.health_by_format.clone(),
|
||||
expected_circuit_breaker_by_format: current_key.circuit_breaker_by_format,
|
||||
health_by_format: current_key.health_by_format,
|
||||
@@ -1188,6 +1351,12 @@ async fn record_oauth_invalidation_effect(
|
||||
}
|
||||
|
||||
let plan = context.plan;
|
||||
// Agent assertions are long-lived credential requests whose task can rotate
|
||||
// while the response is in flight. Runtime 401/403 handling must not project
|
||||
// that response onto whichever credential generation is stored later.
|
||||
if execution_plan_uses_codex_agent_identity(plan) {
|
||||
return;
|
||||
}
|
||||
let transport = match state
|
||||
.read_provider_transport_snapshot(&plan.provider_id, &plan.endpoint_id, &plan.key_id)
|
||||
.await
|
||||
@@ -1202,9 +1371,23 @@ async fn record_oauth_invalidation_effect(
|
||||
return;
|
||||
}
|
||||
};
|
||||
// The inverse replacement is equally unsafe: a response sent with an old
|
||||
// bearer token must not invalidate a newly installed Agent Identity.
|
||||
if crate::provider_transport::is_codex_agent_identity_transport(&transport) {
|
||||
return;
|
||||
}
|
||||
if !transport.key.auth_type.trim().eq_ignore_ascii_case("oauth") {
|
||||
return;
|
||||
}
|
||||
if transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("codex")
|
||||
&& !execution_plan_bearer_matches_transport(plan, &transport)
|
||||
{
|
||||
return;
|
||||
}
|
||||
|
||||
let Some(invalid_reason) = resolve_local_oauth_invalid_reason(
|
||||
transport.provider.provider_type.as_str(),
|
||||
@@ -1214,11 +1397,27 @@ async fn record_oauth_invalidation_effect(
|
||||
return;
|
||||
};
|
||||
|
||||
let expected_auth_config = match state
|
||||
.capture_provider_transport_auth_config_fence(&transport)
|
||||
.await
|
||||
{
|
||||
Ok(Some(ciphertext)) => ciphertext,
|
||||
Ok(None) => return,
|
||||
Err(err) => {
|
||||
warn!(
|
||||
"gateway orchestration effects: failed to capture oauth invalidation fence for provider {} endpoint {} key {}: {:?}",
|
||||
plan.provider_id, plan.endpoint_id, plan.key_id, err
|
||||
);
|
||||
return;
|
||||
}
|
||||
};
|
||||
|
||||
if let Err(err) = state
|
||||
.mark_provider_catalog_key_oauth_invalid(
|
||||
.mark_provider_catalog_key_oauth_invalid_fenced(
|
||||
&plan.key_id,
|
||||
transport.provider.provider_type.as_str(),
|
||||
invalid_reason.as_str(),
|
||||
expected_auth_config.as_str(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
@@ -1227,98 +1426,34 @@ async fn record_oauth_invalidation_effect(
|
||||
plan.provider_id, plan.endpoint_id, plan.key_id, err
|
||||
);
|
||||
}
|
||||
record_pool_score_schedule_feedback(
|
||||
state,
|
||||
context,
|
||||
Some(false),
|
||||
Some(PoolMemberHardState::AuthInvalid),
|
||||
Some(-2_000),
|
||||
serde_json::json!({
|
||||
"last_request_feedback": {
|
||||
"source": "oauth_invalidation",
|
||||
"status_code": effect.status_code,
|
||||
"reason": invalid_reason.as_str()
|
||||
}
|
||||
}),
|
||||
)
|
||||
.await;
|
||||
|
||||
match auto_remove_runtime_oauth_invalid_key(
|
||||
state,
|
||||
&transport.provider,
|
||||
&plan.key_id,
|
||||
invalid_reason.as_str(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(true) => {
|
||||
tracing::info!(
|
||||
provider_id = %plan.provider_id,
|
||||
endpoint_id = %plan.endpoint_id,
|
||||
key_id = %plan.key_id,
|
||||
provider_type = %transport.provider.provider_type,
|
||||
event_name = "auto_removed_oauth_runtime_invalid",
|
||||
"gateway auto-removed runtime invalid oauth key"
|
||||
);
|
||||
}
|
||||
Ok(false) => {}
|
||||
Err(err) => {
|
||||
warn!(
|
||||
"gateway orchestration effects: failed to auto-remove oauth invalid key for provider {} endpoint {} key {}: {:?}",
|
||||
plan.provider_id, plan.endpoint_id, plan.key_id, err
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn auto_remove_runtime_oauth_invalid_key(
|
||||
state: &AppState,
|
||||
provider: &GatewayProviderTransportProvider,
|
||||
key_id: &str,
|
||||
invalid_reason: &str,
|
||||
) -> Result<bool, GatewayError> {
|
||||
if !admin_provider_quota_pure::provider_auto_remove_banned_keys(provider.config.as_ref()) {
|
||||
return Ok(false);
|
||||
}
|
||||
fn execution_plan_uses_codex_agent_identity(plan: &ExecutionPlan) -> bool {
|
||||
execution_plan_authorization(plan)
|
||||
.is_some_and(crate::provider_transport::is_codex_agent_identity_authorization)
|
||||
}
|
||||
|
||||
let key_ids = [key_id.to_string()];
|
||||
let Some(key) = state
|
||||
.read_provider_catalog_keys_by_ids(&key_ids)
|
||||
.await?
|
||||
.into_iter()
|
||||
.next()
|
||||
else {
|
||||
return Ok(false);
|
||||
};
|
||||
if key.provider_id != provider.id {
|
||||
return Ok(false);
|
||||
}
|
||||
fn execution_plan_authorization(plan: &ExecutionPlan) -> Option<&str> {
|
||||
plan.headers
|
||||
.iter()
|
||||
.find(|(name, _)| name.eq_ignore_ascii_case("authorization"))
|
||||
.map(|(_, value)| value.as_str())
|
||||
}
|
||||
|
||||
if !admin_provider_quota_pure::should_auto_remove_oauth_invalid_key(
|
||||
&key,
|
||||
Some(invalid_reason),
|
||||
true,
|
||||
current_unix_secs(),
|
||||
) {
|
||||
return Ok(false);
|
||||
}
|
||||
|
||||
let deleted_key_id = key.id.clone();
|
||||
if !state.delete_provider_catalog_key(&deleted_key_id).await? {
|
||||
return Ok(false);
|
||||
}
|
||||
state
|
||||
.cleanup_deleted_provider_catalog_refs(
|
||||
&provider.id,
|
||||
false,
|
||||
&[],
|
||||
std::slice::from_ref(&deleted_key_id),
|
||||
)
|
||||
.await?;
|
||||
let _ = state
|
||||
.invalidate_local_oauth_refresh_entry(&deleted_key_id)
|
||||
.await;
|
||||
Ok(true)
|
||||
fn execution_plan_bearer_matches_transport(
|
||||
plan: &ExecutionPlan,
|
||||
transport: &crate::provider_transport::GatewayProviderTransportSnapshot,
|
||||
) -> bool {
|
||||
let current_token = transport.key.decrypted_api_key.trim();
|
||||
!current_token.is_empty()
|
||||
&& plan.headers.iter().any(|(name, value)| {
|
||||
name.eq_ignore_ascii_case("authorization")
|
||||
&& value
|
||||
.trim()
|
||||
.strip_prefix("Bearer ")
|
||||
.map(str::trim)
|
||||
.is_some_and(|token| token == current_token)
|
||||
})
|
||||
}
|
||||
|
||||
fn resolve_local_oauth_invalid_reason(
|
||||
@@ -1409,6 +1544,12 @@ async fn record_pool_score_schedule_feedback(
|
||||
if context.plan.provider_id.trim().is_empty() || context.plan.key_id.trim().is_empty() {
|
||||
return;
|
||||
}
|
||||
if capture_local_execution_auth_config_fence(state, context.plan)
|
||||
.await
|
||||
.is_none()
|
||||
{
|
||||
return;
|
||||
}
|
||||
if !pool_score_feedback_gate_allows(context.plan, succeeded, hard_state, score_delta) {
|
||||
return;
|
||||
}
|
||||
@@ -1560,7 +1701,8 @@ mod tests {
|
||||
};
|
||||
use aether_data_contracts::repository::pool_scores::PoolMemberHardState;
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||
ProviderCatalogKeyAdaptiveState, StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
|
||||
StoredProviderCatalogProvider,
|
||||
};
|
||||
use aether_test_support::ManagedRedisServer;
|
||||
use serde_json::{json, Value};
|
||||
@@ -1720,7 +1862,10 @@ mod tests {
|
||||
key_id: "key-codex-cli-local-1".to_string(),
|
||||
method: "POST".to_string(),
|
||||
url: "https://chatgpt.com/backend-api/codex".to_string(),
|
||||
headers: BTreeMap::new(),
|
||||
headers: BTreeMap::from([(
|
||||
"authorization".to_string(),
|
||||
"Bearer __placeholder__".to_string(),
|
||||
)]),
|
||||
content_type: Some("application/json".to_string()),
|
||||
content_encoding: None,
|
||||
body: RequestBody::from_json(json!({"model":"gpt-5.4"})),
|
||||
@@ -1734,6 +1879,15 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
fn sample_codex_agent_identity_plan() -> ExecutionPlan {
|
||||
let mut plan = sample_codex_plan();
|
||||
plan.headers.insert(
|
||||
"Authorization".to_string(),
|
||||
"AgentAssertion in-flight-assertion".to_string(),
|
||||
);
|
||||
plan
|
||||
}
|
||||
|
||||
fn sample_codex_provider() -> StoredProviderCatalogProvider {
|
||||
StoredProviderCatalogProvider::new(
|
||||
"provider-codex-cli-local-1".to_string(),
|
||||
@@ -1818,6 +1972,19 @@ mod tests {
|
||||
.expect("key transport should build")
|
||||
}
|
||||
|
||||
fn sample_codex_agent_identity_key() -> StoredProviderCatalogKey {
|
||||
let mut key = sample_codex_key();
|
||||
key.name = "Agent Identity".to_string();
|
||||
key.encrypted_auth_config = Some(
|
||||
encrypt_python_fernet_plaintext(
|
||||
DEVELOPMENT_ENCRYPTION_KEY,
|
||||
r#"{"provider_type":"codex","auth_mode":"agentIdentity","agent_runtime_id":"runtime-current","agent_private_key":"MC4CAQAwBQYDK2VwBCIEIAcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcH","task_id":"task-current"}"#,
|
||||
)
|
||||
.expect("Agent Identity auth config should encrypt"),
|
||||
);
|
||||
key
|
||||
}
|
||||
|
||||
fn codex_state() -> AppState {
|
||||
codex_state_with_provider(sample_codex_provider())
|
||||
}
|
||||
@@ -1827,10 +1994,17 @@ mod tests {
|
||||
}
|
||||
|
||||
fn codex_state_with_provider(provider: StoredProviderCatalogProvider) -> AppState {
|
||||
codex_state_with_provider_and_key(provider, sample_codex_key())
|
||||
}
|
||||
|
||||
fn codex_state_with_provider_and_key(
|
||||
provider: StoredProviderCatalogProvider,
|
||||
key: StoredProviderCatalogKey,
|
||||
) -> AppState {
|
||||
let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![provider],
|
||||
vec![sample_codex_endpoint()],
|
||||
vec![sample_codex_key()],
|
||||
vec![key],
|
||||
));
|
||||
AppState::new()
|
||||
.expect("gateway state should build")
|
||||
@@ -2744,7 +2918,7 @@ mod tests {
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn oauth_invalidation_auto_removes_inactive_pat_owner_when_enabled() {
|
||||
async fn oauth_invalidation_retains_inactive_pat_owner_when_auto_remove_is_enabled() {
|
||||
let state = codex_state_with_auto_remove();
|
||||
let plan = sample_codex_plan();
|
||||
|
||||
@@ -2763,13 +2937,16 @@ mod tests {
|
||||
)
|
||||
.await;
|
||||
|
||||
let keys = state
|
||||
let stored_key = state
|
||||
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&plan.key_id))
|
||||
.await
|
||||
.expect("provider catalog keys should load");
|
||||
assert!(
|
||||
keys.is_empty(),
|
||||
"hard-invalid PAT owner should be auto removed"
|
||||
.expect("provider catalog keys should load")
|
||||
.into_iter()
|
||||
.next()
|
||||
.expect("runtime invalidation must not race-delete a replacement key");
|
||||
assert_eq!(
|
||||
stored_key.oauth_invalid_reason.as_deref(),
|
||||
Some("[OAUTH_EXPIRED] Personal access token owner is inactive.")
|
||||
);
|
||||
}
|
||||
|
||||
@@ -2806,6 +2983,253 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn oauth_invalidation_does_not_mutate_replacement_after_agent_request() {
|
||||
let state = codex_state_with_auto_remove();
|
||||
let plan = sample_codex_agent_identity_plan();
|
||||
|
||||
apply_local_execution_effect(
|
||||
&state,
|
||||
LocalExecutionEffectContext {
|
||||
plan: &plan,
|
||||
report_context: None,
|
||||
},
|
||||
LocalExecutionEffect::OauthInvalidation(LocalOAuthInvalidationEffect {
|
||||
status_code: 403,
|
||||
response_text: Some(
|
||||
r#"{"error":{"code":"biscuit_baker_service_auth_credential_error_status","message":"Personal access token owner is inactive."},"status":403}"#,
|
||||
),
|
||||
}),
|
||||
)
|
||||
.await;
|
||||
|
||||
let stored_key = state
|
||||
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&plan.key_id))
|
||||
.await
|
||||
.expect("provider catalog keys should load")
|
||||
.into_iter()
|
||||
.next()
|
||||
.expect("replacement OAuth key should not be removed");
|
||||
assert_eq!(stored_key.oauth_invalid_at_unix_secs, None);
|
||||
assert_eq!(stored_key.oauth_invalid_reason, None);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn oauth_invalidation_does_not_mutate_agent_replacement_after_bearer_request() {
|
||||
let state = codex_state_with_provider_and_key(
|
||||
sample_codex_provider_with_auto_remove(),
|
||||
sample_codex_agent_identity_key(),
|
||||
);
|
||||
let mut plan = sample_codex_plan();
|
||||
plan.headers.insert(
|
||||
"authorization".to_string(),
|
||||
"Bearer old-access-token".to_string(),
|
||||
);
|
||||
|
||||
apply_local_execution_effect(
|
||||
&state,
|
||||
LocalExecutionEffectContext {
|
||||
plan: &plan,
|
||||
report_context: None,
|
||||
},
|
||||
LocalExecutionEffect::OauthInvalidation(LocalOAuthInvalidationEffect {
|
||||
status_code: 403,
|
||||
response_text: Some(
|
||||
r#"{"error":{"code":"biscuit_baker_service_auth_credential_error_status","message":"Personal access token owner is inactive."},"status":403}"#,
|
||||
),
|
||||
}),
|
||||
)
|
||||
.await;
|
||||
|
||||
let stored_key = state
|
||||
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&plan.key_id))
|
||||
.await
|
||||
.expect("provider catalog keys should load")
|
||||
.into_iter()
|
||||
.next()
|
||||
.expect("Agent Identity replacement should not be removed");
|
||||
assert_eq!(stored_key.oauth_invalid_at_unix_secs, None);
|
||||
assert_eq!(stored_key.oauth_invalid_reason, None);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn health_failure_updates_codex_key_for_current_bearer_request() {
|
||||
let state = codex_state();
|
||||
let plan = sample_codex_plan();
|
||||
|
||||
apply_local_execution_effect(
|
||||
&state,
|
||||
LocalExecutionEffectContext {
|
||||
plan: &plan,
|
||||
report_context: None,
|
||||
},
|
||||
LocalExecutionEffect::HealthFailure(LocalHealthFailureEffect {
|
||||
status_code: 503,
|
||||
classification: LocalFailoverClassification::RetryUpstreamFailure,
|
||||
}),
|
||||
)
|
||||
.await;
|
||||
|
||||
let stored_key = state
|
||||
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&plan.key_id))
|
||||
.await
|
||||
.expect("provider catalog keys should load")
|
||||
.into_iter()
|
||||
.next()
|
||||
.expect("stored key should exist");
|
||||
assert_eq!(
|
||||
stored_key
|
||||
.health_by_format
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("openai:responses"))
|
||||
.and_then(|value| value.get("consecutive_failures"))
|
||||
.and_then(Value::as_u64),
|
||||
Some(1)
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn health_failure_does_not_mutate_codex_bearer_replacement() {
|
||||
let mut replacement = sample_codex_key();
|
||||
replacement.encrypted_api_key = Some(
|
||||
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "replacement-token")
|
||||
.expect("replacement token should encrypt"),
|
||||
);
|
||||
replacement.encrypted_auth_config = Some(
|
||||
encrypt_python_fernet_plaintext(
|
||||
DEVELOPMENT_ENCRYPTION_KEY,
|
||||
r#"{"provider_type":"codex","refresh_token":"replacement-refresh-token"}"#,
|
||||
)
|
||||
.expect("replacement auth config should encrypt"),
|
||||
);
|
||||
let expected_health = replacement.health_by_format.clone();
|
||||
let expected_circuit = replacement.circuit_breaker_by_format.clone();
|
||||
let state = codex_state_with_provider_and_key(sample_codex_provider(), replacement);
|
||||
let plan = sample_codex_plan();
|
||||
|
||||
apply_local_execution_effect(
|
||||
&state,
|
||||
LocalExecutionEffectContext {
|
||||
plan: &plan,
|
||||
report_context: None,
|
||||
},
|
||||
LocalExecutionEffect::HealthFailure(LocalHealthFailureEffect {
|
||||
status_code: 503,
|
||||
classification: LocalFailoverClassification::RetryUpstreamFailure,
|
||||
}),
|
||||
)
|
||||
.await;
|
||||
|
||||
let stored_key = state
|
||||
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&plan.key_id))
|
||||
.await
|
||||
.expect("provider catalog keys should load")
|
||||
.into_iter()
|
||||
.next()
|
||||
.expect("replacement key should exist");
|
||||
assert_eq!(stored_key.health_by_format, expected_health);
|
||||
assert_eq!(stored_key.circuit_breaker_by_format, expected_circuit);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn adaptive_rate_limit_does_not_mutate_codex_bearer_replacement() {
|
||||
let mut replacement = sample_codex_key();
|
||||
replacement.encrypted_api_key = Some(
|
||||
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "replacement-token")
|
||||
.expect("replacement token should encrypt"),
|
||||
);
|
||||
replacement.encrypted_auth_config = Some(
|
||||
encrypt_python_fernet_plaintext(
|
||||
DEVELOPMENT_ENCRYPTION_KEY,
|
||||
r#"{"provider_type":"codex","refresh_token":"replacement-refresh-token"}"#,
|
||||
)
|
||||
.expect("replacement auth config should encrypt"),
|
||||
);
|
||||
replacement.learned_rpm_limit = Some(12);
|
||||
replacement.rpm_429_count = Some(1);
|
||||
let expected_adaptive_state = ProviderCatalogKeyAdaptiveState::from(&replacement);
|
||||
let state = codex_state_with_provider_and_key(sample_codex_provider(), replacement);
|
||||
let plan = sample_codex_plan();
|
||||
|
||||
apply_local_execution_effect(
|
||||
&state,
|
||||
LocalExecutionEffectContext {
|
||||
plan: &plan,
|
||||
report_context: None,
|
||||
},
|
||||
LocalExecutionEffect::AdaptiveRateLimit(LocalAdaptiveRateLimitEffect {
|
||||
status_code: 429,
|
||||
classification: LocalFailoverClassification::RetryUpstreamFailure,
|
||||
headers: Some(&BTreeMap::from([(
|
||||
"x-ratelimit-limit-requests".to_string(),
|
||||
"42".to_string(),
|
||||
)])),
|
||||
}),
|
||||
)
|
||||
.await;
|
||||
|
||||
let stored_key = state
|
||||
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&plan.key_id))
|
||||
.await
|
||||
.expect("provider catalog keys should load")
|
||||
.into_iter()
|
||||
.next()
|
||||
.expect("replacement key should exist");
|
||||
assert_eq!(
|
||||
ProviderCatalogKeyAdaptiveState::from(&stored_key),
|
||||
expected_adaptive_state
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn pool_error_does_not_clear_codex_bearer_replacement_circuit() {
|
||||
let legacy_circuit = json!({
|
||||
"openai:responses": {
|
||||
"open": true,
|
||||
"reason": "replacement-state"
|
||||
}
|
||||
});
|
||||
let mut replacement = sample_codex_key();
|
||||
replacement.encrypted_api_key = Some(
|
||||
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "replacement-token")
|
||||
.expect("replacement token should encrypt"),
|
||||
);
|
||||
replacement.encrypted_auth_config = Some(
|
||||
encrypt_python_fernet_plaintext(
|
||||
DEVELOPMENT_ENCRYPTION_KEY,
|
||||
r#"{"provider_type":"codex","refresh_token":"replacement-refresh-token"}"#,
|
||||
)
|
||||
.expect("replacement auth config should encrypt"),
|
||||
);
|
||||
replacement.circuit_breaker_by_format = Some(legacy_circuit.clone());
|
||||
let state = codex_state_with_provider_and_key(sample_codex_provider(), replacement);
|
||||
let plan = sample_codex_plan();
|
||||
|
||||
apply_local_execution_effect(
|
||||
&state,
|
||||
LocalExecutionEffectContext {
|
||||
plan: &plan,
|
||||
report_context: None,
|
||||
},
|
||||
LocalExecutionEffect::PoolError(LocalPoolErrorEffect {
|
||||
status_code: 401,
|
||||
classification: LocalFailoverClassification::StopErrorPattern,
|
||||
headers: &BTreeMap::new(),
|
||||
error_body: Some(r#"{"error":{"message":"account has been deactivated"}}"#),
|
||||
}),
|
||||
)
|
||||
.await;
|
||||
|
||||
let stored_key = state
|
||||
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&plan.key_id))
|
||||
.await
|
||||
.expect("provider catalog keys should load")
|
||||
.into_iter()
|
||||
.next()
|
||||
.expect("replacement key should exist");
|
||||
assert_eq!(stored_key.circuit_breaker_by_format, Some(legacy_circuit));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn health_failure_projection_updates_key_health_for_format() {
|
||||
let state = health_state();
|
||||
|
||||
@@ -85,14 +85,25 @@ impl ProviderKeyAuthSemantics {
|
||||
|
||||
pub(crate) fn provider_key_can_refresh_oauth(
|
||||
auth_semantics: ProviderKeyAuthSemantics,
|
||||
provider_type: &str,
|
||||
auth_config: Option<&Map<String, Value>>,
|
||||
) -> bool {
|
||||
auth_semantics.can_refresh_oauth()
|
||||
&& auth_config
|
||||
.and_then(|config| config.get("refresh_token"))
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.is_some_and(|value| !value.is_empty())
|
||||
&& (provider_key_auth_config_is_agent_identity(provider_type, auth_config)
|
||||
|| auth_config
|
||||
.and_then(|config| config.get("refresh_token"))
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.is_some_and(|value| !value.is_empty()))
|
||||
}
|
||||
|
||||
pub(crate) fn provider_key_can_export_oauth(
|
||||
auth_semantics: ProviderKeyAuthSemantics,
|
||||
provider_type: &str,
|
||||
auth_config: Option<&Map<String, Value>>,
|
||||
) -> bool {
|
||||
auth_semantics.can_export_oauth()
|
||||
&& !provider_key_auth_config_is_agent_identity(provider_type, auth_config)
|
||||
}
|
||||
|
||||
pub(crate) fn provider_key_auth_config_uses_header_authorization(
|
||||
@@ -291,9 +302,10 @@ mod tests {
|
||||
use super::{
|
||||
provider_active_api_formats, provider_key_auth_config_is_agent_identity,
|
||||
provider_key_auth_config_uses_header_authorization, provider_key_auth_semantics,
|
||||
provider_key_can_refresh_oauth, provider_key_configured_api_formats,
|
||||
provider_key_effective_api_formats, provider_key_inherits_provider_api_formats,
|
||||
ProviderKeyCredentialKind, ProviderKeyRuntimeAuthKind,
|
||||
provider_key_can_export_oauth, provider_key_can_refresh_oauth,
|
||||
provider_key_configured_api_formats, provider_key_effective_api_formats,
|
||||
provider_key_inherits_provider_api_formats, ProviderKeyCredentialKind,
|
||||
ProviderKeyRuntimeAuthKind,
|
||||
};
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
|
||||
@@ -396,6 +408,7 @@ mod tests {
|
||||
|
||||
assert!(!provider_key_can_refresh_oauth(
|
||||
semantics,
|
||||
"codex",
|
||||
json!({
|
||||
"access_token": "access-token",
|
||||
"access_token_import_temporary": true
|
||||
@@ -404,12 +417,24 @@ mod tests {
|
||||
));
|
||||
assert!(!provider_key_can_refresh_oauth(
|
||||
semantics,
|
||||
"codex",
|
||||
json!({ "refresh_token": " " }).as_object()
|
||||
));
|
||||
assert!(provider_key_can_refresh_oauth(
|
||||
semantics,
|
||||
"codex",
|
||||
json!({ "refresh_token": "refresh-token" }).as_object()
|
||||
));
|
||||
assert!(provider_key_can_refresh_oauth(
|
||||
semantics,
|
||||
"codex",
|
||||
json!({
|
||||
"auth_mode": "agentIdentity",
|
||||
"agent_runtime_id": "runtime-1",
|
||||
"agent_private_key": "private-key-present"
|
||||
})
|
||||
.as_object()
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -455,6 +480,28 @@ mod tests {
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn agent_identity_is_not_exportable_through_generic_oauth_export() {
|
||||
let semantics = provider_key_auth_semantics(&sample_key("oauth"), "codex");
|
||||
let agent_identity = json!({
|
||||
"auth_mode": "agentIdentity",
|
||||
"agent_runtime_id": "runtime-1",
|
||||
"agent_private_key": "base64-private-key",
|
||||
"task_id": "task-1"
|
||||
});
|
||||
|
||||
assert!(!provider_key_can_export_oauth(
|
||||
semantics,
|
||||
"codex",
|
||||
agent_identity.as_object()
|
||||
));
|
||||
assert!(provider_key_can_export_oauth(
|
||||
semantics,
|
||||
"codex",
|
||||
json!({ "refresh_token": "refresh-token" }).as_object()
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn recognizes_legacy_kiro_bearer_key_with_auth_config_as_oauth_managed() {
|
||||
let mut key = sample_key("bearer");
|
||||
|
||||
@@ -1465,6 +1465,7 @@ mod tests {
|
||||
.compare_and_update_provider_catalog_key_health_state(
|
||||
&aether_data_contracts::repository::provider_catalog::ProviderCatalogKeyHealthStateUpdate {
|
||||
key_id: "key-1".to_string(),
|
||||
expected_encrypted_auth_config: None,
|
||||
expected_health_by_format: None,
|
||||
expected_circuit_breaker_by_format: None,
|
||||
health_by_format: Some(health_by_format),
|
||||
|
||||
@@ -38,6 +38,7 @@ pub(crate) use self::cache::{
|
||||
PROVIDER_TRANSPORT_SNAPSHOT_CACHE_STALE_TTL, PROVIDER_TRANSPORT_SNAPSHOT_CACHE_TTL,
|
||||
};
|
||||
pub use self::cors::FrontdoorCorsConfig;
|
||||
pub(crate) use self::oauth::AgentIdentityAuthConfigFence;
|
||||
pub(crate) use self::types::{
|
||||
AdminWalletMutationOutcome, GatewayAdminPaymentCallbackView, GatewayUserPreferenceView,
|
||||
GatewayUserSessionView, LocalExecutionRuntimeMissDiagnostic, LocalMutationOutcome,
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -677,6 +677,103 @@ async fn gateway_creates_admin_provider_key_locally_with_trusted_admin_principal
|
||||
upstream_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn generic_key_routes_reject_agent_identity_credential_writes() {
|
||||
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![sample_provider("provider-codex", "codex", 10)],
|
||||
vec![],
|
||||
vec![sample_key(
|
||||
"key-codex-existing",
|
||||
"provider-codex",
|
||||
"openai:responses",
|
||||
"existing-secret",
|
||||
)],
|
||||
));
|
||||
let gateway = build_router_with_state(
|
||||
AppState::new()
|
||||
.expect("gateway should build")
|
||||
.with_data_state_for_tests(
|
||||
GatewayDataState::with_provider_catalog_repository_for_tests(
|
||||
provider_catalog_repository.clone(),
|
||||
)
|
||||
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
|
||||
),
|
||||
);
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
let agent_identity = json!({
|
||||
"provider_type": "codex",
|
||||
"auth_mode": "agentIdentity",
|
||||
"agent_runtime_id": "runtime-bypass",
|
||||
"agent_private_key": "private-key-must-use-dedicated-import"
|
||||
});
|
||||
let client = reqwest::Client::new();
|
||||
|
||||
let create_response = client
|
||||
.post(format!(
|
||||
"{gateway_url}/api/admin/endpoints/providers/provider-codex/keys"
|
||||
))
|
||||
.header(crate::constants::GATEWAY_HEADER, "rust-phase3b")
|
||||
.header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123")
|
||||
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
|
||||
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
|
||||
.json(&json!({
|
||||
"api_formats": ["openai:responses"],
|
||||
"auth_type": "oauth",
|
||||
"auth_config": agent_identity,
|
||||
"name": "bypass create"
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
.expect("create request should complete");
|
||||
assert_eq!(create_response.status(), StatusCode::BAD_REQUEST);
|
||||
let create_payload: serde_json::Value = create_response
|
||||
.json()
|
||||
.await
|
||||
.expect("create error should be JSON");
|
||||
assert!(create_payload["detail"]
|
||||
.as_str()
|
||||
.is_some_and(|detail| detail.contains("专属创建或导入接口")));
|
||||
|
||||
let update_response = client
|
||||
.put(format!(
|
||||
"{gateway_url}/api/admin/endpoints/keys/key-codex-existing"
|
||||
))
|
||||
.header(crate::constants::GATEWAY_HEADER, "rust-phase3b")
|
||||
.header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123")
|
||||
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
|
||||
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
|
||||
.json(&json!({
|
||||
"auth_type": "oauth",
|
||||
"auth_config": {
|
||||
"provider_type": "codex",
|
||||
"auth_mode": "agentIdentity",
|
||||
"agent_runtime_id": "runtime-bypass-update",
|
||||
"agent_private_key": "private-key-must-use-dedicated-import"
|
||||
}
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
.expect("update request should complete");
|
||||
assert_eq!(update_response.status(), StatusCode::BAD_REQUEST);
|
||||
let update_payload: serde_json::Value = update_response
|
||||
.json()
|
||||
.await
|
||||
.expect("update error should be JSON");
|
||||
assert!(update_payload["detail"]
|
||||
.as_str()
|
||||
.is_some_and(|detail| detail.contains("专属创建或导入接口")));
|
||||
|
||||
let keys = provider_catalog_repository
|
||||
.list_keys_by_provider_ids(&["provider-codex".to_string()])
|
||||
.await
|
||||
.expect("keys should read");
|
||||
assert_eq!(keys.len(), 1);
|
||||
assert_eq!(keys[0].id, "key-codex-existing");
|
||||
assert_eq!(keys[0].auth_type, "api_key");
|
||||
|
||||
gateway_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn provider_key_concurrent_limit_create_and_list_responses() {
|
||||
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
@@ -1380,6 +1477,121 @@ async fn gateway_export_preserves_distinct_imported_access_token_with_authorizat
|
||||
upstream_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_generic_export_rejects_agent_identity_without_exposing_private_key() {
|
||||
let private_key = "agent-private-key-must-not-leak";
|
||||
let mut provider = sample_provider("provider-codex", "codex", 10);
|
||||
provider.provider_type = "codex".to_string();
|
||||
let mut key = sample_key(
|
||||
"key-codex-agent",
|
||||
"provider-codex",
|
||||
"openai:responses",
|
||||
"__placeholder__",
|
||||
);
|
||||
key.auth_type = "oauth".to_string();
|
||||
key.encrypted_auth_config = Some(
|
||||
encrypt_python_fernet_plaintext(
|
||||
DEVELOPMENT_ENCRYPTION_KEY,
|
||||
&json!({
|
||||
"provider_type": "codex",
|
||||
"auth_mode": "agentIdentity",
|
||||
"agent_runtime_id": "runtime-must-not-leak",
|
||||
"agent_private_key": private_key,
|
||||
"task_id": "task-must-not-leak"
|
||||
})
|
||||
.to_string(),
|
||||
)
|
||||
.expect("Agent Identity auth config should encrypt"),
|
||||
);
|
||||
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![provider],
|
||||
vec![],
|
||||
vec![key],
|
||||
));
|
||||
let gateway = build_router_with_state(
|
||||
AppState::new()
|
||||
.expect("gateway should build")
|
||||
.with_data_state_for_tests(
|
||||
GatewayDataState::with_provider_catalog_reader_for_tests(
|
||||
provider_catalog_repository,
|
||||
)
|
||||
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
|
||||
),
|
||||
);
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
let client = reqwest::Client::new();
|
||||
|
||||
let reveal_response = client
|
||||
.get(format!(
|
||||
"{gateway_url}/api/admin/endpoints/keys/key-codex-agent/reveal"
|
||||
))
|
||||
.header(GATEWAY_HEADER, "rust-phase3b")
|
||||
.header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123")
|
||||
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
|
||||
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
|
||||
.send()
|
||||
.await
|
||||
.expect("reveal request should complete");
|
||||
|
||||
assert_eq!(reveal_response.status(), StatusCode::BAD_REQUEST);
|
||||
let reveal_body = reveal_response
|
||||
.text()
|
||||
.await
|
||||
.expect("reveal error body should read");
|
||||
assert!(reveal_body.contains("专属 provider-oauth 管理面"));
|
||||
assert!(!reveal_body.contains(private_key));
|
||||
assert!(!reveal_body.contains("runtime-must-not-leak"));
|
||||
assert!(!reveal_body.contains("task-must-not-leak"));
|
||||
|
||||
let export_response = client
|
||||
.get(format!(
|
||||
"{gateway_url}/api/admin/endpoints/keys/key-codex-agent/export"
|
||||
))
|
||||
.header(GATEWAY_HEADER, "rust-phase3b")
|
||||
.header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123")
|
||||
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
|
||||
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
|
||||
.send()
|
||||
.await
|
||||
.expect("export request should complete");
|
||||
|
||||
assert_eq!(export_response.status(), StatusCode::BAD_REQUEST);
|
||||
let export_body = export_response
|
||||
.text()
|
||||
.await
|
||||
.expect("export error body should read");
|
||||
assert!(export_body.contains("专属 provider-oauth 管理面"));
|
||||
assert!(!export_body.contains(private_key));
|
||||
assert!(!export_body.contains("runtime-must-not-leak"));
|
||||
assert!(!export_body.contains("task-must-not-leak"));
|
||||
|
||||
let list_response = client
|
||||
.get(format!(
|
||||
"{gateway_url}/api/admin/endpoints/providers/provider-codex/keys"
|
||||
))
|
||||
.header(GATEWAY_HEADER, "rust-phase3b")
|
||||
.header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123")
|
||||
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
|
||||
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
|
||||
.send()
|
||||
.await
|
||||
.expect("list request should complete");
|
||||
|
||||
assert_eq!(list_response.status(), StatusCode::OK);
|
||||
let list_payload: serde_json::Value = list_response
|
||||
.json()
|
||||
.await
|
||||
.expect("list body should be JSON");
|
||||
let agent = list_payload
|
||||
.as_array()
|
||||
.and_then(|items| items.iter().find(|item| item["id"] == "key-codex-agent"))
|
||||
.expect("Agent Identity key should be listed");
|
||||
assert_eq!(agent["agent_identity"], true);
|
||||
assert_eq!(agent["can_export_oauth"], false);
|
||||
|
||||
gateway_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_clears_admin_provider_key_oauth_invalid_locally_with_trusted_admin_principal() {
|
||||
let upstream_hits = Arc::new(Mutex::new(0usize));
|
||||
@@ -2601,11 +2813,27 @@ async fn gateway_handles_admin_keys_grouped_by_format_locally_with_trusted_admin
|
||||
key_b.created_at_unix_ms = Some(1_711_100_000);
|
||||
key_b.updated_at_unix_secs = Some(1_711_100_100);
|
||||
|
||||
let mut key_agent = sample_key(
|
||||
"key-codex-agent",
|
||||
"provider-codex",
|
||||
"openai:responses",
|
||||
"__placeholder__",
|
||||
);
|
||||
key_agent.auth_type = "oauth".to_string();
|
||||
key_agent.encrypted_auth_config = Some(
|
||||
encrypt_python_fernet_plaintext(
|
||||
DEVELOPMENT_ENCRYPTION_KEY,
|
||||
r#"{"provider_type":"codex","auth_mode":"agentIdentity","agent_runtime_id":"runtime-1","agent_private_key":"base64-private-key","task_id":"task-1"}"#,
|
||||
)
|
||||
.expect("Agent Identity auth config should encrypt"),
|
||||
);
|
||||
|
||||
let provider_catalog_repository = Arc::new(SummaryNullingProviderCatalogReadRepository::seed(
|
||||
vec![
|
||||
sample_provider("provider-openai", "openai", 10),
|
||||
sample_provider("provider-claude", "claude", 20)
|
||||
.with_transport_fields(false, false, true, None, None, None, None, None, None),
|
||||
sample_provider("provider-codex", "codex", 30),
|
||||
],
|
||||
vec![
|
||||
sample_endpoint(
|
||||
@@ -2620,8 +2848,14 @@ async fn gateway_handles_admin_keys_grouped_by_format_locally_with_trusted_admin
|
||||
"claude:messages",
|
||||
"https://api.claude.example",
|
||||
),
|
||||
sample_endpoint(
|
||||
"endpoint-codex-responses",
|
||||
"provider-codex",
|
||||
"openai:responses",
|
||||
"https://api.codex.example",
|
||||
),
|
||||
],
|
||||
vec![key_a, key_b],
|
||||
vec![key_a, key_b, key_agent],
|
||||
));
|
||||
|
||||
let (upstream_url, upstream_handle) = start_server(upstream).await;
|
||||
@@ -2665,6 +2899,11 @@ async fn gateway_handles_admin_keys_grouped_by_format_locally_with_trusted_admin
|
||||
);
|
||||
assert_eq!(payload["openai:chat"][0]["internal_priority"], 10);
|
||||
assert_eq!(payload["claude:messages"][0]["provider_active"], false);
|
||||
let agent_item = payload["openai:responses"]
|
||||
.as_array()
|
||||
.and_then(|items| items.iter().find(|item| item["id"] == "key-codex-agent"))
|
||||
.expect("Agent Identity key should be grouped");
|
||||
assert_eq!(agent_item["api_key_masked"], "[Agent Identity]");
|
||||
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
|
||||
|
||||
gateway_handle.abort();
|
||||
|
||||
@@ -8968,6 +8968,137 @@ async fn gateway_allows_management_token_with_pool_write_for_provider_oauth_batc
|
||||
gateway_handle.abort();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn gateway_prevents_pool_write_token_from_importing_agent_identity_via_batch_routes() {
|
||||
run_admin_oauth_test(
|
||||
"gateway_prevents_pool_write_token_from_importing_agent_identity_via_batch_routes",
|
||||
gateway_prevents_pool_write_token_from_importing_agent_identity_via_batch_routes_impl,
|
||||
);
|
||||
}
|
||||
|
||||
async fn gateway_prevents_pool_write_token_from_importing_agent_identity_via_batch_routes_impl() {
|
||||
let raw_token = "ae-provider-oauth-agent-identity-pool-write";
|
||||
let state = AppState::new().expect("gateway should build");
|
||||
let admin_user = state
|
||||
.create_local_auth_user_with_settings(
|
||||
Some("[email protected]".to_string()),
|
||||
true,
|
||||
"admin".to_string(),
|
||||
"hash".to_string(),
|
||||
"admin".to_string(),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.await
|
||||
.expect("admin user should be created")
|
||||
.expect("admin user should exist");
|
||||
let mut management_token = sample_management_token(
|
||||
"token-provider-oauth-agent-identity-pool",
|
||||
&admin_user.id,
|
||||
"provider-oauth-agent-identity-pool",
|
||||
true,
|
||||
);
|
||||
management_token.token.allowed_ips = None;
|
||||
management_token.token.permissions = Some(json!(["admin:pool:read", "admin:pool:write"]));
|
||||
let management_token_repository =
|
||||
Arc::new(InMemoryManagementTokenRepository::seed_with_hashes(
|
||||
vec![management_token],
|
||||
vec![(
|
||||
hash_management_token(raw_token),
|
||||
"token-provider-oauth-agent-identity-pool".to_string(),
|
||||
)],
|
||||
));
|
||||
|
||||
let mut provider = sample_provider("provider-codex", "codex", 10);
|
||||
provider.provider_type = "codex".to_string();
|
||||
let endpoint = sample_endpoint(
|
||||
"endpoint-codex-agent-identity",
|
||||
"provider-codex",
|
||||
"openai:chat",
|
||||
"https://chatgpt.com/backend-api/codex",
|
||||
);
|
||||
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![provider],
|
||||
vec![endpoint],
|
||||
vec![],
|
||||
));
|
||||
let data_state =
|
||||
GatewayDataState::with_management_token_repository_for_tests(management_token_repository)
|
||||
.attach_provider_catalog_repository_for_tests(provider_catalog_repository.clone())
|
||||
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY);
|
||||
let gateway = build_router_with_state(state.with_data_state_for_tests(data_state));
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
|
||||
let credentials = json!({
|
||||
"auth_mode": "agentIdentity",
|
||||
"agent_runtime_id": "runtime-rbac-guard",
|
||||
"agent_private_key": "MC4CAQAwBQYDK2VwBCIEIAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA",
|
||||
"task_id": "task-rbac-guard"
|
||||
})
|
||||
.to_string();
|
||||
let client = reqwest::Client::new();
|
||||
for path in [
|
||||
"/api/admin/provider-oauth/providers/provider-codex/batch-import",
|
||||
"/api/admin/provider-oauth/providers/provider-codex/batch-import/tasks",
|
||||
] {
|
||||
let response = client
|
||||
.post(format!("{gateway_url}{path}"))
|
||||
.header(crate::constants::GATEWAY_HEADER, "rust-phase3b")
|
||||
.bearer_auth(raw_token)
|
||||
.json(&json!({ "credentials": credentials }))
|
||||
.send()
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
let status = response.status();
|
||||
let payload: serde_json::Value = response.json().await.expect("json body should parse");
|
||||
assert_eq!(
|
||||
status,
|
||||
StatusCode::BAD_REQUEST,
|
||||
"path={path} payload={payload}"
|
||||
);
|
||||
assert_eq!(
|
||||
payload["detail"],
|
||||
"Agent Identity JSON 必须使用专属导入接口"
|
||||
);
|
||||
}
|
||||
|
||||
let dedicated_response = client
|
||||
.post(format!(
|
||||
"{gateway_url}/api/admin/provider-oauth/providers/provider-codex/agent-identity-import/tasks"
|
||||
))
|
||||
.header(crate::constants::GATEWAY_HEADER, "rust-phase3b")
|
||||
.bearer_auth(raw_token)
|
||||
.json(&json!({ "credentials": credentials }))
|
||||
.send()
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
let dedicated_status = dedicated_response.status();
|
||||
let dedicated_payload: serde_json::Value = dedicated_response
|
||||
.json()
|
||||
.await
|
||||
.expect("json body should parse");
|
||||
assert_eq!(
|
||||
dedicated_status,
|
||||
StatusCode::FORBIDDEN,
|
||||
"payload={dedicated_payload}"
|
||||
);
|
||||
assert_eq!(
|
||||
dedicated_payload["required_permission"],
|
||||
"admin:provider_oauth:write"
|
||||
);
|
||||
|
||||
let keys = provider_catalog_repository
|
||||
.list_keys_by_provider_ids(&["provider-codex".to_string()])
|
||||
.await
|
||||
.expect("keys should load");
|
||||
assert!(keys.is_empty());
|
||||
|
||||
gateway_handle.abort();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn gateway_rejects_management_token_without_pool_write_for_provider_oauth_batch_import() {
|
||||
run_admin_oauth_test(
|
||||
|
||||
@@ -3147,6 +3147,8 @@ async fn gateway_pool_keys_classify_oauth_credentials() {
|
||||
.expect("Agent Identity key should exist");
|
||||
assert_eq!(agent_identity_key["oauth_header_auth"], false);
|
||||
assert_eq!(agent_identity_key["agent_identity"], true);
|
||||
assert_eq!(agent_identity_key["can_refresh_oauth"], true);
|
||||
assert_eq!(agent_identity_key["can_export_oauth"], false);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
|
||||
Reference in New Issue
Block a user