mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-08 18:37:46 +08:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
535ee098c3 | ||
|
|
b45df89ce4 | ||
|
|
c8118edf36 | ||
|
|
4a0775c4ea | ||
|
|
6fc02dad3e | ||
|
|
dbf2809bd6 | ||
|
|
1d3051cb89 | ||
|
|
59e27524da | ||
|
|
247e7105a2 | ||
|
|
dc3743aecf | ||
|
|
1c5ee5228c | ||
|
|
9d80281b53 | ||
|
|
3b036299d4 | ||
|
|
f70ae68273 | ||
|
|
1353d76e07 | ||
|
|
621a528083 | ||
|
|
a498875591 | ||
|
|
71b54070e8 | ||
|
|
9a0d346ff3 | ||
|
|
32944538e9 | ||
|
|
0b17026eab | ||
|
|
b13d9b9b40 | ||
|
|
b7fca851b8 | ||
|
|
810c3dfe2b | ||
|
|
a1d64e5239 | ||
|
|
fb33ea57b0 | ||
|
|
5b0c763086 | ||
|
|
f3a12c1008 | ||
|
|
ca35e09eaa | ||
|
|
8cf381b0c3 | ||
|
|
654c4f6978 | ||
|
|
edb8362adc | ||
|
|
29fa4aed19 | ||
|
|
41e93858e1 | ||
|
|
8d918d0459 | ||
|
|
3a759fae89 | ||
|
|
985ff3c36a | ||
|
|
908d4f2603 | ||
|
|
4d67569873 | ||
|
|
1aab31a148 | ||
|
|
aedff9a704 | ||
|
|
669f4bddc5 | ||
|
|
1a4eede34d | ||
|
|
0318808db9 |
+1
-1
@@ -66,7 +66,7 @@ ADMIN_USERNAME=admin123456
|
||||
# docker compose 下 app 启动前自动执行 pending migration/backfill(默认 true)
|
||||
# AETHER_GATEWAY_AUTO_PREPARE_DATABASE=true
|
||||
|
||||
# PostgreSQL 连接池配置(默认按 CPU 自动计算;正式高并发环境可显式预算)
|
||||
# PostgreSQL 连接池配置(默认每核 4 条、总池至少 32 条且最多 100 条;多实例部署应显式分配每实例预算)
|
||||
# AETHER_GATEWAY_DATA_POSTGRES_MIN_CONNECTIONS=12
|
||||
# AETHER_GATEWAY_DATA_POSTGRES_MAX_CONNECTIONS=80
|
||||
# AETHER_GATEWAY_MAX_IN_FLIGHT_REQUESTS=2048
|
||||
|
||||
@@ -200,13 +200,13 @@ jobs:
|
||||
RUSTFLAGS: "-C link-arg=-fuse-ld=mold"
|
||||
run: cargo nextest run -p aether-gateway --lib
|
||||
|
||||
- name: Test bin
|
||||
- name: Test bins
|
||||
env:
|
||||
RUSTC_WRAPPER: sccache
|
||||
SCCACHE_GHA_ENABLED: "true"
|
||||
RUST_MIN_STACK: "16777216"
|
||||
RUSTFLAGS: "-C link-arg=-fuse-ld=mold"
|
||||
run: cargo nextest run -p aether-gateway --bin aether-gateway
|
||||
run: cargo nextest run -p aether-gateway --bins
|
||||
|
||||
- name: Show sccache stats
|
||||
if: always()
|
||||
@@ -387,11 +387,11 @@ jobs:
|
||||
- name: Setup sccache
|
||||
uses: mozilla-actions/[email protected]
|
||||
|
||||
- name: Test scenario binaries
|
||||
- name: Test scenario binaries and end-to-end suites
|
||||
env:
|
||||
RUSTC_WRAPPER: sccache
|
||||
SCCACHE_GHA_ENABLED: "true"
|
||||
run: cargo test -p aether-integration-tests --bins
|
||||
run: cargo test -p aether-integration-tests --bins --tests
|
||||
|
||||
- name: Show sccache stats
|
||||
if: always()
|
||||
|
||||
Generated
+4
@@ -325,6 +325,7 @@ dependencies = [
|
||||
"axum",
|
||||
"base64 0.22.1",
|
||||
"bcrypt",
|
||||
"brotli",
|
||||
"bytes",
|
||||
"chrono",
|
||||
"chrono-tz",
|
||||
@@ -346,6 +347,7 @@ dependencies = [
|
||||
"reqwest",
|
||||
"rsa",
|
||||
"rustls 0.23.37",
|
||||
"semver",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"sha1",
|
||||
@@ -441,6 +443,7 @@ name = "aether-integration-tests"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"aether-contracts",
|
||||
"aether-crypto",
|
||||
"aether-data",
|
||||
"aether-data-contracts",
|
||||
"aether-gateway",
|
||||
@@ -457,6 +460,7 @@ dependencies = [
|
||||
"sqlx",
|
||||
"tokio",
|
||||
"tokio-tungstenite 0.28.0",
|
||||
"uuid",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
||||
+2
-1
@@ -106,6 +106,7 @@ async-trait = "0.1"
|
||||
axum = "0.8"
|
||||
base64 = "0.22"
|
||||
bcrypt = "0.16"
|
||||
brotli = "8"
|
||||
bytes = "1"
|
||||
cbc = "0.1"
|
||||
chrono = { version = "0.4", features = ["serde"] }
|
||||
@@ -135,7 +136,7 @@ tokio = { version = "1", features = ["macros", "net", "rt-multi-thread", "signal
|
||||
tokio-util = { version = "0.7", features = ["codec", "io-util"] }
|
||||
tracing = "0.1"
|
||||
tracing-subscriber = { version = "0.3", features = ["env-filter", "json"] }
|
||||
uuid = { version = "1", features = ["serde", "v4", "v5"] }
|
||||
uuid = { version = "1", features = ["serde", "v4", "v5", "v7"] }
|
||||
webpki-roots = "0.26"
|
||||
wreq = { version = "6.0.0-rc.28", default-features = false, features = ["json", "stream", "socks", "webpki-roots", "ws"] }
|
||||
wreq-util = "3.0.0-rc.10"
|
||||
|
||||
@@ -137,12 +137,14 @@ Aether Tunnel 是配套的正向代理节点,部署在海外 VPS 上,为墙
|
||||
|
||||
- Embeddings: [OpenAI compatible `POST /v1/embeddings`](docs/api/embeddings.md)
|
||||
- Rerank: [OpenAI/Jina compatible `POST /v1/rerank`](docs/api/rerank.md)
|
||||
- Responses WebSocket mode: [protocol and Aether behavior](docs/WebSocket-Mode.md)
|
||||
- WebSocket probes: [Codex](docs/operations/codex-responses-websocket-probe.md) · [OpenAI Responses](docs/operations/openai-responses-websocket-probe.md)
|
||||
|
||||
## 环境变量
|
||||
|
||||
- `APP_PORT`:`aether-gateway` 唯一监听端口,固定绑定 `0.0.0.0:${APP_PORT}`
|
||||
- `DATABASE_URL`:数据库连接串;SQLite 例如 `sqlite:///opt/aether/data/aether.db`,Postgres 例如 `postgresql://postgres:aether@postgres:5432/aether`
|
||||
- `AETHER_GATEWAY_DATA_POSTGRES_MIN_CONNECTIONS` / `AETHER_GATEWAY_DATA_POSTGRES_MAX_CONNECTIONS`:数据库连接池手动覆盖值;未配置时会自动推导,SQLite 固定 `1/1`,Postgres/MySQL 按 CPU 核心数计算并默认封顶 `100`
|
||||
- `AETHER_GATEWAY_DATA_POSTGRES_MIN_CONNECTIONS` / `AETHER_GATEWAY_DATA_POSTGRES_MAX_CONNECTIONS`:数据库连接池手动覆盖值;未配置时 SQLite 固定 `1/1`,Postgres/MySQL 按每核 `4` 条自动推导,总池范围为 `32-100`。该预算按进程计算,多实例部署应按数据库连接上限显式分配
|
||||
- `AETHER_GATEWAY_MAX_IN_FLIGHT_REQUESTS`:单实例请求并发上限;未配置时按 CPU 自动推导(基础范围 `512-65536`),低文件描述符预算时会进一步下调
|
||||
- `AETHER_GATEWAY_REQUEST_BODY_BUFFER_BUDGET_MB`:单实例同时读取和解压请求体的加权内存预算,默认 `256MB`
|
||||
- `AETHER_GATEWAY_REQUEST_BODY_READ_TIMEOUT_MS`:请求体完整读取超时,默认 `120000ms`
|
||||
|
||||
@@ -52,6 +52,7 @@ async-trait.workspace = true
|
||||
axum = { version = "0.8", features = ["ws"] }
|
||||
base64.workspace = true
|
||||
bcrypt.workspace = true
|
||||
brotli.workspace = true
|
||||
bytes.workspace = true
|
||||
chrono.workspace = true
|
||||
chrono-tz.workspace = true
|
||||
@@ -75,6 +76,7 @@ rsa = "0.9.10"
|
||||
rustls.workspace = true
|
||||
serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
semver.workspace = true
|
||||
sha1 = "0.10"
|
||||
sha2 = { workspace = true, features = ["oid"] }
|
||||
socket2.workspace = true
|
||||
|
||||
@@ -69,6 +69,9 @@ pub(crate) use aether_ai_formats::api::{
|
||||
};
|
||||
pub(crate) use aether_ai_formats::protocol::stream::CanonicalUsage as StreamingCanonicalUsage;
|
||||
pub(crate) use aether_ai_formats::CODEX_RESPONSES_LITE_HEADER;
|
||||
/// Codex client identity headers re-exported for out-of-crate probe binaries,
|
||||
/// which must reach `aether_ai_formats` through this seam.
|
||||
pub use aether_ai_formats::{CODEX_CLIENT_ORIGINATOR, CODEX_CLIENT_USER_AGENT};
|
||||
|
||||
pub(crate) fn parse_direct_request_body(
|
||||
parts: &http::request::Parts,
|
||||
|
||||
@@ -52,17 +52,19 @@ pub(crate) use self::planner::{
|
||||
build_standard_family_sync_plan_and_reports, build_standard_stream_plan_from_decision,
|
||||
build_standard_sync_plan_from_decision, candidate_auth_channel_skip_reason,
|
||||
codex_model_capabilities_for_transport, extract_pool_sticky_session_token,
|
||||
maybe_build_stream_decision_payload, maybe_build_stream_plan_payload,
|
||||
maybe_build_sync_decision_payload, maybe_build_sync_plan_payload,
|
||||
planner_is_matching_stream_request, provider_key_pool_score_id, provider_key_pool_score_scope,
|
||||
read_candidate_transport_snapshot, record_local_runtime_candidate_skip_reason,
|
||||
maybe_build_responses_websocket_decision, maybe_build_stream_decision_payload,
|
||||
maybe_build_stream_plan_payload, maybe_build_sync_decision_payload,
|
||||
maybe_build_sync_plan_payload, planner_is_matching_stream_request, provider_key_pool_score_id,
|
||||
provider_key_pool_score_scope, read_candidate_transport_snapshot,
|
||||
record_local_runtime_candidate_skip_reason, resolve_provider_chat_pii_redaction,
|
||||
resolve_tunnel_scheduler_affinity_context, resolve_upstream_is_stream_for_provider,
|
||||
set_local_openai_chat_execution_exhausted_diagnostic,
|
||||
set_local_openai_image_execution_exhausted_diagnostic, validate_final_openai_provider_request,
|
||||
CandidateFailureDiagnostic, CandidateFailureDiagnosticKind, EligibleLocalExecutionCandidate,
|
||||
GatewayAuthApiKeySnapshot, GatewayProviderTransportSnapshot, LocalExecutionAttemptSource,
|
||||
LocalExecutionCandidateKind, LocalResolvedOAuthRequestAuth, PlannerAppState,
|
||||
SkippedLocalExecutionCandidate,
|
||||
ResponsesWebSocketBodyNormalization, ResponsesWebSocketDecision,
|
||||
ResponsesWebSocketPinnedCandidate, SkippedLocalExecutionCandidate,
|
||||
};
|
||||
pub(crate) use self::pure::*;
|
||||
pub(crate) use self::response_history::{
|
||||
|
||||
@@ -188,9 +188,14 @@ pub(crate) async fn scheduler_ordering_config_for_routing_policy(
|
||||
state: PlannerAppState<'_>,
|
||||
routing_policy: Option<&ResolvedRoutingPolicy>,
|
||||
) -> SchedulerOrderingConfig {
|
||||
let system_config = read_scheduler_ordering_config_or_default(state).await;
|
||||
match routing_policy {
|
||||
Some(policy) => scheduler_ordering_config_from_routing_policy(policy),
|
||||
None => read_scheduler_ordering_config_or_default(state).await,
|
||||
Some(policy) => {
|
||||
let mut config = scheduler_ordering_config_from_routing_policy(policy);
|
||||
config.keep_priority_on_conversion |= system_config.keep_priority_on_conversion;
|
||||
config
|
||||
}
|
||||
None => system_config,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -390,6 +395,43 @@ mod tests {
|
||||
assert_eq!(overlaid.key_global_priority_for_format, Some(2));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn routing_policy_inherits_global_conversion_priority_override() {
|
||||
let data_state = GatewayDataState::default().with_system_config_values_for_tests([(
|
||||
"keep_priority_on_conversion".to_string(),
|
||||
json!(true),
|
||||
)]);
|
||||
let state = AppState::new()
|
||||
.expect("state should build")
|
||||
.with_data_state_for_tests(data_state);
|
||||
let policy = aether_routing_core::ResolvedRoutingPolicy {
|
||||
group_id: Some("group-1".to_string()),
|
||||
group_version: Some(1),
|
||||
selection_source: "system_default".to_string(),
|
||||
requested_model: "gpt-5.4-mini".to_string(),
|
||||
resolved_model: "gpt-5.4-mini".to_string(),
|
||||
priority_mode: aether_routing_core::RoutingSetPriorityMode::Provider,
|
||||
scheduling_mode: aether_routing_core::RoutingSchedulingMode::FixedOrder,
|
||||
keep_priority_on_conversion: false,
|
||||
ranking_overlay: Default::default(),
|
||||
mutation_plan: Default::default(),
|
||||
pool_policy_overrides: Default::default(),
|
||||
matched_rules: Vec::new(),
|
||||
};
|
||||
|
||||
let ordering = super::scheduler_ordering_config_for_routing_policy(
|
||||
PlannerAppState::new(&state),
|
||||
Some(&policy),
|
||||
)
|
||||
.await;
|
||||
|
||||
assert_eq!(
|
||||
ordering.scheduling_mode,
|
||||
crate::scheduler::config::SchedulerSchedulingMode::FixedOrder
|
||||
);
|
||||
assert!(ordering.keep_priority_on_conversion);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn routing_policy_uses_pool_priority_for_pool_group_global_key_slot() {
|
||||
let mut candidate = sample_candidate("endpoint-1", "representative-key");
|
||||
|
||||
@@ -2540,4 +2540,146 @@ mod tests {
|
||||
vec!["provider-openai-responses-regular"]
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn fixed_order_prefers_codex_responses_when_conversion_keeps_priority() {
|
||||
let mut codex = standard_candidate_row("provider-codex", "openai:responses", 0);
|
||||
codex.provider_type = "codex".to_string();
|
||||
codex.key_auth_type = "oauth".to_string();
|
||||
codex.global_model_name = "gpt-5.4-mini".to_string();
|
||||
codex.model_provider_model_name = "gpt-5.4-mini".to_string();
|
||||
let mut custom_chat = standard_candidate_row("provider-custom", "openai:chat", 10);
|
||||
custom_chat.global_model_name = "gpt-5.4-mini".to_string();
|
||||
custom_chat.model_provider_model_name = "gpt-5.4-mini".to_string();
|
||||
|
||||
let candidate_repository: Arc<dyn MinimalCandidateSelectionReadRepository> =
|
||||
Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed([
|
||||
codex.clone(),
|
||||
custom_chat.clone(),
|
||||
]));
|
||||
let catalog_items = [
|
||||
provider_catalog_for_standard_row(&codex, false),
|
||||
provider_catalog_for_standard_row(&custom_chat, false),
|
||||
];
|
||||
let provider_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
catalog_items
|
||||
.iter()
|
||||
.map(|(provider, _, _)| provider.clone())
|
||||
.collect(),
|
||||
catalog_items
|
||||
.iter()
|
||||
.map(|(_, endpoint, _)| endpoint.clone())
|
||||
.collect(),
|
||||
catalog_items
|
||||
.iter()
|
||||
.map(|(_, _, key)| key.clone())
|
||||
.collect(),
|
||||
));
|
||||
let data_state =
|
||||
GatewayDataState::with_provider_catalog_and_minimal_candidate_selection_for_tests(
|
||||
provider_repository,
|
||||
candidate_repository,
|
||||
)
|
||||
.with_encryption_key_for_tests("development-key")
|
||||
.with_system_config_values_for_tests([
|
||||
(
|
||||
"scheduling_mode".to_string(),
|
||||
serde_json::json!("fixed_order"),
|
||||
),
|
||||
(
|
||||
"keep_priority_on_conversion".to_string(),
|
||||
serde_json::json!(true),
|
||||
),
|
||||
]);
|
||||
let app = AppState::new()
|
||||
.expect("gateway state should build")
|
||||
.with_data_state_for_tests(data_state);
|
||||
let auth_snapshot = unrestricted_auth_snapshot();
|
||||
let model_directive_policy =
|
||||
crate::system_features::ModelDirectivePolicySnapshot::load(&app).await;
|
||||
let routing_policy = ResolvedRoutingPolicy {
|
||||
group_id: Some("routing-group-codex-first".to_string()),
|
||||
group_version: Some(1),
|
||||
selection_source: "test".to_string(),
|
||||
requested_model: "gpt-5.4-mini".to_string(),
|
||||
resolved_model: "gpt-5.4-mini".to_string(),
|
||||
priority_mode: aether_routing_core::RoutingSetPriorityMode::Provider,
|
||||
scheduling_mode: aether_routing_core::RoutingSchedulingMode::FixedOrder,
|
||||
keep_priority_on_conversion: false,
|
||||
ranking_overlay: Default::default(),
|
||||
mutation_plan: Default::default(),
|
||||
pool_policy_overrides: Default::default(),
|
||||
matched_rules: Vec::new(),
|
||||
};
|
||||
let mut cursor = LocalCandidatePreselectionPageCursor::new(
|
||||
PlannerAppState::new(&app),
|
||||
&model_directive_policy,
|
||||
"openai:chat",
|
||||
"gpt-5.4-mini",
|
||||
None,
|
||||
false,
|
||||
None,
|
||||
&auth_snapshot,
|
||||
Some(&routing_policy),
|
||||
None,
|
||||
None,
|
||||
true,
|
||||
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModelAndApiFormat,
|
||||
false,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
|
||||
let first_page = cursor
|
||||
.next_page()
|
||||
.await
|
||||
.expect("preselection should succeed")
|
||||
.expect("Codex and custom candidates should share the priority page");
|
||||
assert!(
|
||||
first_page.skipped_candidates.is_empty(),
|
||||
"priority page unexpectedly skipped candidates: {:?}",
|
||||
first_page
|
||||
.skipped_candidates
|
||||
.iter()
|
||||
.map(|candidate| {
|
||||
(
|
||||
candidate.candidate.provider_id.as_str(),
|
||||
candidate.skip_reason,
|
||||
)
|
||||
})
|
||||
.collect::<Vec<_>>()
|
||||
);
|
||||
assert_eq!(
|
||||
first_page
|
||||
.candidates
|
||||
.iter()
|
||||
.map(|candidate| candidate.provider_id.as_str())
|
||||
.collect::<Vec<_>>(),
|
||||
vec!["provider-custom", "provider-codex"]
|
||||
);
|
||||
let (ranked, skipped) =
|
||||
super::super::candidate_resolution::resolve_and_rank_logical_local_execution_candidates(
|
||||
PlannerAppState::new(&app),
|
||||
first_page.candidates,
|
||||
"openai:chat",
|
||||
Some("gpt-5.4-mini"),
|
||||
Some(&auth_snapshot),
|
||||
None,
|
||||
None,
|
||||
Some(&routing_policy),
|
||||
None,
|
||||
None,
|
||||
aether_ai_serving::AiCandidateResolutionMode::Standard,
|
||||
)
|
||||
.await;
|
||||
|
||||
assert!(skipped.is_empty());
|
||||
assert_eq!(
|
||||
ranked
|
||||
.iter()
|
||||
.map(|candidate| candidate.candidate.provider_id.as_str())
|
||||
.collect::<Vec<_>>(),
|
||||
vec!["provider-codex", "provider-custom"]
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -55,6 +55,7 @@ pub(crate) struct LocalRequestedModelDecisionInput {
|
||||
pub(crate) client_surface: Option<ClientSurface>,
|
||||
pub(crate) gateway_credential_carrier: Option<GatewayCredentialCarrier>,
|
||||
pub(crate) client_session_affinity: Option<ClientSessionAffinity>,
|
||||
pub(crate) original_client_session_id: Option<String>,
|
||||
pub(crate) routing_policy: Option<ResolvedRoutingPolicy>,
|
||||
pub(crate) routing_trace_seed: Option<RoutingDecisionTrace>,
|
||||
pub(crate) routing_context: Option<LocalRoutingRequestContext>,
|
||||
@@ -150,6 +151,12 @@ pub(crate) fn apply_provider_request_routing_policy_to_decision(
|
||||
provider_api_format.as_str(),
|
||||
);
|
||||
}
|
||||
apply_codex_oauth_fingerprint_convergence_to_decision(
|
||||
input,
|
||||
decision,
|
||||
transport,
|
||||
provider_api_format.as_str(),
|
||||
);
|
||||
return Ok(());
|
||||
};
|
||||
let provider_body_rules = decision
|
||||
@@ -207,6 +214,12 @@ pub(crate) fn apply_provider_request_routing_policy_to_decision(
|
||||
provider_api_format.as_str(),
|
||||
);
|
||||
}
|
||||
apply_codex_oauth_fingerprint_convergence_to_decision(
|
||||
input,
|
||||
decision,
|
||||
transport,
|
||||
provider_api_format.as_str(),
|
||||
);
|
||||
return Ok(());
|
||||
}
|
||||
if original_provider_request_body.is_none() && !policy.mutation_plan.body_patch.is_empty() {
|
||||
@@ -305,13 +318,40 @@ pub(crate) fn apply_provider_request_routing_policy_to_decision(
|
||||
if original_provider_request_body.is_some() {
|
||||
decision.provider_request_body = Some(provider_request_body);
|
||||
}
|
||||
apply_codex_oauth_fingerprint_convergence_to_decision(
|
||||
input,
|
||||
decision,
|
||||
transport,
|
||||
provider_api_format.as_str(),
|
||||
);
|
||||
update_report_context_provider_request_mutation(decision, &policy);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn apply_codex_oauth_fingerprint_convergence_to_decision(
|
||||
input: &LocalRequestedModelDecisionInput,
|
||||
decision: &mut AiExecutionDecision,
|
||||
transport: Option<&GatewayProviderTransportSnapshot>,
|
||||
provider_api_format: &str,
|
||||
) {
|
||||
let (Some(transport), Some(provider_request_body)) =
|
||||
(transport, decision.provider_request_body.as_mut())
|
||||
else {
|
||||
return;
|
||||
};
|
||||
crate::ai_serving::transport::apply_codex_oauth_fingerprint_convergence(
|
||||
transport,
|
||||
provider_api_format,
|
||||
input.original_client_session_id.as_deref(),
|
||||
&mut decision.provider_request_headers,
|
||||
provider_request_body,
|
||||
);
|
||||
}
|
||||
|
||||
struct GatewayAuthenticatedDecisionInputPort<'a> {
|
||||
state: PlannerAppState<'a>,
|
||||
now_unix_secs: u64,
|
||||
auth_snapshot_override: Option<GatewayAuthApiKeySnapshot>,
|
||||
model_directive_policy: &'a crate::system_features::ModelDirectivePolicySnapshot,
|
||||
model_directive_base_model: Option<String>,
|
||||
}
|
||||
@@ -328,6 +368,17 @@ impl AiAuthenticatedDecisionInputPort for GatewayAuthenticatedDecisionInputPort<
|
||||
&self,
|
||||
auth_context: &Self::AuthContext,
|
||||
) -> Result<Option<Self::AuthSnapshot>, Self::Error> {
|
||||
if let Some(snapshot) = self.auth_snapshot_override.as_ref() {
|
||||
if snapshot.user_id != auth_context.user_id
|
||||
|| snapshot.api_key_id != auth_context.api_key_id
|
||||
{
|
||||
return Err(GatewayError::Internal(
|
||||
"WebSocket auth snapshot identity does not match its control decision"
|
||||
.to_string(),
|
||||
));
|
||||
}
|
||||
return Ok(Some(snapshot.clone()));
|
||||
}
|
||||
self.state
|
||||
.read_auth_api_key_snapshot(
|
||||
&auth_context.user_id,
|
||||
@@ -383,6 +434,7 @@ pub(crate) fn build_local_requested_model_decision_input(
|
||||
client_surface: None,
|
||||
gateway_credential_carrier: None,
|
||||
client_session_affinity: None,
|
||||
original_client_session_id: None,
|
||||
routing_policy: None,
|
||||
routing_trace_seed: None,
|
||||
routing_context: None,
|
||||
@@ -397,6 +449,8 @@ pub(crate) async fn attach_routing_policy_to_local_requested_model_input(
|
||||
body_json: &Value,
|
||||
client_api_format: &str,
|
||||
) -> Result<(), GatewayError> {
|
||||
input.original_client_session_id = routing_header_value_str(&parts.headers, "session-id")
|
||||
.or_else(|| routing_header_value_str(&parts.headers, "session_id"));
|
||||
let explicit_group = routing_header_value_str(&parts.headers, ROUTING_GROUP_HEADER);
|
||||
let selected_group = match state.routing_group_read_repository() {
|
||||
Some(repository) => {
|
||||
@@ -708,6 +762,27 @@ pub(crate) async fn resolve_local_authenticated_decision_input(
|
||||
requested_model_api_format: Option<&str>,
|
||||
explicit_required_capabilities: Option<&serde_json::Value>,
|
||||
model_directive_policy: &crate::system_features::ModelDirectivePolicySnapshot,
|
||||
) -> Result<Option<ResolvedLocalDecisionAuthInput>, GatewayError> {
|
||||
resolve_local_authenticated_decision_input_with_snapshot(
|
||||
state,
|
||||
auth_context,
|
||||
None,
|
||||
requested_model,
|
||||
requested_model_api_format,
|
||||
explicit_required_capabilities,
|
||||
model_directive_policy,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn resolve_local_authenticated_decision_input_with_snapshot(
|
||||
state: &AppState,
|
||||
auth_context: ExecutionRuntimeAuthContext,
|
||||
auth_snapshot_override: Option<GatewayAuthApiKeySnapshot>,
|
||||
requested_model: Option<&str>,
|
||||
requested_model_api_format: Option<&str>,
|
||||
explicit_required_capabilities: Option<&serde_json::Value>,
|
||||
model_directive_policy: &crate::system_features::ModelDirectivePolicySnapshot,
|
||||
) -> Result<Option<ResolvedLocalDecisionAuthInput>, GatewayError> {
|
||||
let model_directive_base_model = match (requested_model, requested_model_api_format) {
|
||||
(Some(model), Some(api_format)) => model_directive_policy
|
||||
@@ -719,6 +794,7 @@ pub(crate) async fn resolve_local_authenticated_decision_input(
|
||||
let port = GatewayAuthenticatedDecisionInputPort {
|
||||
state: PlannerAppState::new(state),
|
||||
now_unix_secs: current_unix_secs(),
|
||||
auth_snapshot_override,
|
||||
model_directive_policy,
|
||||
model_directive_base_model,
|
||||
};
|
||||
@@ -1024,6 +1100,52 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn explicit_auth_snapshot_override_does_not_fall_back_to_the_planner_cache() {
|
||||
// AppState::new has no auth snapshot repository. Without the explicit
|
||||
// override this resolver returns None; a WebSocket strong snapshot must
|
||||
// therefore be the exact value used to build the planner input.
|
||||
let state = AppState::new().expect("test state should build");
|
||||
let mut strong_snapshot = sample_auth_snapshot();
|
||||
strong_snapshot.api_key_allowed_models = Some(vec!["gpt-live-only".to_string()]);
|
||||
|
||||
let resolved = resolve_local_authenticated_decision_input_with_snapshot(
|
||||
&state,
|
||||
sample_auth_context(),
|
||||
Some(strong_snapshot.clone()),
|
||||
Some("gpt-live-only"),
|
||||
Some("openai:responses"),
|
||||
None,
|
||||
&Default::default(),
|
||||
)
|
||||
.await
|
||||
.expect("snapshot override should resolve")
|
||||
.expect("the explicit snapshot should replace the missing cached value");
|
||||
|
||||
assert_eq!(resolved.auth_snapshot, strong_snapshot);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn explicit_auth_snapshot_override_rejects_an_identity_mismatch() {
|
||||
let state = AppState::new().expect("test state should build");
|
||||
let mut wrong_snapshot = sample_auth_snapshot();
|
||||
wrong_snapshot.api_key_id = "another-key".to_string();
|
||||
|
||||
let error = resolve_local_authenticated_decision_input_with_snapshot(
|
||||
&state,
|
||||
sample_auth_context(),
|
||||
Some(wrong_snapshot),
|
||||
Some("gpt-live-only"),
|
||||
Some("openai:responses"),
|
||||
None,
|
||||
&Default::default(),
|
||||
)
|
||||
.await
|
||||
.expect_err("a snapshot for another API key must never be injected");
|
||||
|
||||
assert!(matches!(error, GatewayError::Internal(message) if message.contains("identity")));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn explicit_routing_attachment_authorizes_and_caches_per_principal() {
|
||||
let repository = Arc::new(InMemoryRoutingGroupRepository::default());
|
||||
@@ -1154,6 +1276,7 @@ mod tests {
|
||||
client_surface: None,
|
||||
gateway_credential_carrier: None,
|
||||
client_session_affinity: None,
|
||||
original_client_session_id: None,
|
||||
routing_policy: None,
|
||||
routing_trace_seed: None,
|
||||
model_directive_policy: Default::default(),
|
||||
@@ -1312,6 +1435,57 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
fn sample_codex_fingerprint_transport() -> GatewayProviderTransportSnapshot {
|
||||
let mut transport = sample_codex_transport_with_card();
|
||||
transport.provider.config = Some(json!({
|
||||
"codex": {"fingerprint_convergence_enabled": true}
|
||||
}));
|
||||
transport.endpoint.api_format = "openai:responses".to_string();
|
||||
transport.endpoint.endpoint_kind = Some("responses".to_string());
|
||||
transport.key.api_formats = Some(vec!["openai:responses".to_string()]);
|
||||
transport.key.decrypted_auth_config =
|
||||
Some(json!({"account_id": "account-codex-1"}).to_string());
|
||||
transport
|
||||
}
|
||||
|
||||
fn sample_codex_fingerprint_decision() -> AiExecutionDecision {
|
||||
let prompt_cache_key = "172c39e6-c0a0-5a70-8b63-e0f8e0d185a3";
|
||||
let mut decision = sample_decision();
|
||||
decision.provider_type = Some("codex".to_string());
|
||||
decision.provider_api_format = Some("openai:responses".to_string());
|
||||
decision.client_api_format = Some("openai:responses".to_string());
|
||||
decision.provider_request_headers.extend([
|
||||
("session-id".to_string(), "spoofed-session".to_string()),
|
||||
("thread-id".to_string(), "spoofed-thread".to_string()),
|
||||
(
|
||||
"x-codex-turn-metadata".to_string(),
|
||||
json!({
|
||||
"installation_id": "spoofed-installation",
|
||||
"session_id": "spoofed-session",
|
||||
"thread_source": "cli"
|
||||
})
|
||||
.to_string(),
|
||||
),
|
||||
]);
|
||||
decision.provider_request_body = Some(json!({
|
||||
"model": "gpt-5",
|
||||
"input": [],
|
||||
"metadata": {},
|
||||
"prompt_cache_key": prompt_cache_key,
|
||||
"client_metadata": {
|
||||
"session_id": "spoofed-session",
|
||||
"thread_id": "spoofed-thread",
|
||||
"caller": "sdk",
|
||||
"x-codex-turn-metadata": json!({
|
||||
"installation_id": "spoofed-installation",
|
||||
"session_id": "spoofed-session",
|
||||
"sandbox": "workspace-write"
|
||||
}).to_string()
|
||||
}
|
||||
}));
|
||||
decision
|
||||
}
|
||||
|
||||
fn set_provider_request_rules(
|
||||
input: &mut LocalRequestedModelDecisionInput,
|
||||
allowed_models: &[&str],
|
||||
@@ -1351,6 +1525,7 @@ mod tests {
|
||||
client_surface: None,
|
||||
gateway_credential_carrier: None,
|
||||
client_session_affinity: None,
|
||||
original_client_session_id: None,
|
||||
routing_policy: None,
|
||||
routing_trace_seed: None,
|
||||
model_directive_policy: Default::default(),
|
||||
@@ -1420,6 +1595,7 @@ mod tests {
|
||||
client_surface: None,
|
||||
gateway_credential_carrier: None,
|
||||
client_session_affinity: None,
|
||||
original_client_session_id: None,
|
||||
routing_policy: None,
|
||||
routing_trace_seed: None,
|
||||
routing_context: None,
|
||||
@@ -1488,6 +1664,124 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn codex_fingerprint_convergence_runs_at_every_provider_routing_success_exit() {
|
||||
let transport = sample_codex_fingerprint_transport();
|
||||
let mut no_context = sample_decision_input();
|
||||
no_context.routing_context = None;
|
||||
let mut empty_mutation = sample_decision_input();
|
||||
empty_mutation
|
||||
.routing_context
|
||||
.as_mut()
|
||||
.expect("routing context")
|
||||
.group_config_json = json!({
|
||||
"allowed_models": ["gpt-5"],
|
||||
"rules": []
|
||||
});
|
||||
let mut with_mutation = sample_decision_input();
|
||||
for input in [&mut no_context, &mut empty_mutation, &mut with_mutation] {
|
||||
input.original_client_session_id = Some("client-session-1".to_string());
|
||||
}
|
||||
|
||||
let mut stable_identity = None;
|
||||
let mut turn_ids = std::collections::BTreeSet::new();
|
||||
for (exit_name, input) in [
|
||||
("no_context", no_context),
|
||||
("empty_mutation", empty_mutation),
|
||||
("with_mutation", with_mutation),
|
||||
] {
|
||||
let mut decision = sample_codex_fingerprint_decision();
|
||||
|
||||
apply_provider_request_routing_policy_to_decision(
|
||||
&input,
|
||||
&mut decision,
|
||||
Some(&transport),
|
||||
)
|
||||
.unwrap_or_else(|error| panic!("{exit_name} should converge: {error:?}"));
|
||||
|
||||
let session_id = decision.provider_request_headers["session-id"].clone();
|
||||
let thread_id = decision.provider_request_headers["thread-id"].clone();
|
||||
let installation_id =
|
||||
decision.provider_request_headers["x-codex-installation-id"].clone();
|
||||
let window_id = decision.provider_request_headers["x-codex-window-id"].clone();
|
||||
assert_eq!(decision.provider_request_headers["session_id"], session_id);
|
||||
assert_eq!(
|
||||
decision.provider_request_headers["x-client-request-id"],
|
||||
thread_id
|
||||
);
|
||||
assert_eq!(window_id, format!("{thread_id}:0"));
|
||||
assert_eq!(
|
||||
uuid::Uuid::parse_str(&session_id)
|
||||
.expect("session UUID")
|
||||
.get_version_num(),
|
||||
4
|
||||
);
|
||||
assert_eq!(
|
||||
uuid::Uuid::parse_str(&thread_id)
|
||||
.expect("thread UUID")
|
||||
.get_version_num(),
|
||||
4
|
||||
);
|
||||
|
||||
let body = decision
|
||||
.provider_request_body
|
||||
.as_ref()
|
||||
.expect("request body");
|
||||
assert_eq!(
|
||||
body["prompt_cache_key"],
|
||||
"172c39e6-c0a0-5a70-8b63-e0f8e0d185a3"
|
||||
);
|
||||
assert_eq!(body["client_metadata"]["session_id"], session_id);
|
||||
assert_eq!(body["client_metadata"]["thread_id"], thread_id);
|
||||
assert_eq!(body["client_metadata"]["caller"], "sdk");
|
||||
assert_eq!(
|
||||
body["client_metadata"]["x-codex-installation-id"],
|
||||
installation_id
|
||||
);
|
||||
assert_eq!(body["client_metadata"]["x-codex-window-id"], window_id);
|
||||
|
||||
let header_metadata: Value =
|
||||
serde_json::from_str(&decision.provider_request_headers["x-codex-turn-metadata"])
|
||||
.expect("header turn metadata");
|
||||
let body_metadata: Value = serde_json::from_str(
|
||||
body["client_metadata"]["x-codex-turn-metadata"]
|
||||
.as_str()
|
||||
.expect("embedded turn metadata"),
|
||||
)
|
||||
.expect("embedded turn metadata JSON");
|
||||
assert_eq!(header_metadata["thread_source"], "cli");
|
||||
assert_eq!(body_metadata["sandbox"], "workspace-write");
|
||||
assert_eq!(
|
||||
header_metadata["turn_id"],
|
||||
body["client_metadata"]["turn_id"]
|
||||
);
|
||||
assert_eq!(body_metadata["turn_id"], body["client_metadata"]["turn_id"]);
|
||||
assert_eq!(
|
||||
header_metadata["turn_started_at_unix_ms"],
|
||||
body_metadata["turn_started_at_unix_ms"]
|
||||
);
|
||||
let turn_id = body["client_metadata"]["turn_id"]
|
||||
.as_str()
|
||||
.expect("turn ID")
|
||||
.to_string();
|
||||
assert_eq!(
|
||||
uuid::Uuid::parse_str(&turn_id)
|
||||
.expect("turn UUID")
|
||||
.get_version_num(),
|
||||
7
|
||||
);
|
||||
turn_ids.insert(turn_id);
|
||||
|
||||
let identity = (installation_id, session_id, thread_id);
|
||||
if let Some(expected) = stable_identity.as_ref() {
|
||||
assert_eq!(&identity, expected, "identity changed at {exit_name}");
|
||||
} else {
|
||||
stable_identity = Some(identity);
|
||||
}
|
||||
}
|
||||
assert_eq!(turn_ids.len(), 3, "each request needs a fresh turn ID");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_request_routing_policy_cannot_restore_credentials_or_aether_internal_headers() {
|
||||
for header_name in [
|
||||
|
||||
@@ -49,6 +49,7 @@ pub(crate) use self::plan_builders::{
|
||||
pub(crate) use self::pool_scores::{
|
||||
build_provider_key_pool_score_upsert, provider_key_pool_score_id, provider_key_pool_score_scope,
|
||||
};
|
||||
pub(crate) use self::redaction::resolve_provider_chat_pii_redaction;
|
||||
pub(crate) use self::request_gzip::resolve_transport_request_encoding_policy;
|
||||
pub(crate) use self::route::is_matching_stream_request as planner_is_matching_stream_request;
|
||||
pub(crate) use self::runtime_miss::{
|
||||
@@ -80,8 +81,10 @@ pub(crate) use self::standard::{
|
||||
build_local_stream_plan_and_reports as build_standard_family_stream_plan_and_reports,
|
||||
build_local_sync_attempt_source as build_standard_family_sync_attempt_source,
|
||||
build_local_sync_plan_and_reports as build_standard_family_sync_plan_and_reports,
|
||||
codex_model_capabilities_for_transport, set_local_openai_chat_execution_exhausted_diagnostic,
|
||||
validate_final_openai_provider_request,
|
||||
codex_model_capabilities_for_transport, maybe_build_responses_websocket_decision,
|
||||
set_local_openai_chat_execution_exhausted_diagnostic, validate_final_openai_provider_request,
|
||||
ResponsesWebSocketBodyNormalization, ResponsesWebSocketDecision,
|
||||
ResponsesWebSocketPinnedCandidate,
|
||||
};
|
||||
pub(crate) use self::state::{
|
||||
GatewayAuthApiKeySnapshot, GatewayProviderTransportSnapshot, LocalResolvedOAuthRequestAuth,
|
||||
|
||||
+8
-6
@@ -167,20 +167,21 @@ pub(crate) async fn materialize_local_same_format_provider_candidate_attempts(
|
||||
.collect(),
|
||||
LocalCandidateResolutionMode::Standard,
|
||||
|eligible| {
|
||||
let provider_api_format = eligible.provider_api_format.clone();
|
||||
let (execution_strategy, conversion_mode) = ai_local_execution_contract_for_formats(
|
||||
spec_metadata.api_format,
|
||||
spec_metadata.api_format,
|
||||
&provider_api_format,
|
||||
);
|
||||
Some(build_local_execution_candidate_contract_metadata(
|
||||
LocalExecutionCandidateMetadataParts {
|
||||
eligible,
|
||||
provider_api_format: spec_metadata.api_format,
|
||||
provider_api_format: provider_api_format.as_str(),
|
||||
client_api_format: spec_metadata.api_format,
|
||||
extra_fields: serde_json::Map::new(),
|
||||
},
|
||||
execution_strategy,
|
||||
conversion_mode,
|
||||
spec_metadata.api_format,
|
||||
provider_api_format.as_str(),
|
||||
))
|
||||
},
|
||||
|mut skipped_candidate| {
|
||||
@@ -273,20 +274,21 @@ pub(crate) async fn build_local_same_format_provider_candidate_attempt_source<'a
|
||||
.collect(),
|
||||
LocalCandidateResolutionMode::Standard,
|
||||
|eligible| {
|
||||
let provider_api_format = eligible.provider_api_format.clone();
|
||||
let (execution_strategy, conversion_mode) = ai_local_execution_contract_for_formats(
|
||||
spec_metadata.api_format,
|
||||
spec_metadata.api_format,
|
||||
&provider_api_format,
|
||||
);
|
||||
Some(build_local_execution_candidate_contract_metadata(
|
||||
LocalExecutionCandidateMetadataParts {
|
||||
eligible,
|
||||
provider_api_format: spec_metadata.api_format,
|
||||
provider_api_format: provider_api_format.as_str(),
|
||||
client_api_format: spec_metadata.api_format,
|
||||
extra_fields: serde_json::Map::new(),
|
||||
},
|
||||
execution_strategy,
|
||||
conversion_mode,
|
||||
spec_metadata.api_format,
|
||||
provider_api_format.as_str(),
|
||||
))
|
||||
},
|
||||
|mut skipped_candidate| {
|
||||
|
||||
@@ -55,8 +55,6 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
|
||||
..
|
||||
} = &attempt;
|
||||
let candidate = &eligible.candidate;
|
||||
let (execution_strategy, conversion_mode) =
|
||||
ai_local_execution_contract_for_formats(spec_metadata.api_format, spec_metadata.api_format);
|
||||
let Some(resolved) = resolve_local_same_format_provider_candidate_payload_parts(
|
||||
state, parts, trace_id, body_json, input, &attempt, spec,
|
||||
)
|
||||
@@ -164,6 +162,10 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
|
||||
}
|
||||
}
|
||||
let provider_api_format = resolved.provider_api_format.clone();
|
||||
let (execution_strategy, conversion_mode) = ai_local_execution_contract_for_formats(
|
||||
spec_metadata.api_format,
|
||||
provider_api_format.as_str(),
|
||||
);
|
||||
let effective_headers = input.effective_headers(&parts.headers);
|
||||
let report_context = append_local_failover_policy_to_value(
|
||||
append_execution_contract_fields_to_value(
|
||||
@@ -207,7 +209,10 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
|
||||
.unwrap_or(false),
|
||||
upstream_is_stream: resolved.upstream_is_stream,
|
||||
has_envelope: resolved.is_kiro || resolved.is_antigravity || resolved.is_gemini_cli,
|
||||
needs_conversion: false,
|
||||
needs_conversion: matches!(
|
||||
conversion_mode,
|
||||
crate::ai_serving::ConversionMode::Bidirectional
|
||||
),
|
||||
extra_fields,
|
||||
}),
|
||||
execution_strategy,
|
||||
|
||||
@@ -377,6 +377,7 @@ mod tests {
|
||||
client_surface: None,
|
||||
gateway_credential_carrier: None,
|
||||
client_session_affinity: None,
|
||||
original_client_session_id: None,
|
||||
routing_policy: None,
|
||||
routing_trace_seed: None,
|
||||
routing_context: None,
|
||||
|
||||
@@ -1314,21 +1314,19 @@ async fn resolve_local_gemini_image_to_openai_image_candidate_payload_parts(
|
||||
.as_object_mut()?
|
||||
.insert("stream".to_string(), Value::Bool(true));
|
||||
}
|
||||
provider_request_body = project_openai_image_api_request_body(
|
||||
&provider_request_body,
|
||||
&prepared_candidate.mapped_model,
|
||||
converted.operation,
|
||||
crate::image_capabilities::openai_image_provider_max_generation_count_for_model(
|
||||
transport.provider.provider_type.as_str(),
|
||||
Some(prepared_candidate.mapped_model.as_str()),
|
||||
),
|
||||
)?;
|
||||
if is_codex {
|
||||
provider_request_body = project_codex_openai_image_api_request_body(
|
||||
provider_request_body = if is_codex {
|
||||
project_codex_openai_image_api_request_body(&provider_request_body, converted.operation)?
|
||||
} else {
|
||||
project_openai_image_api_request_body(
|
||||
&provider_request_body,
|
||||
&prepared_candidate.mapped_model,
|
||||
converted.operation,
|
||||
)?;
|
||||
}
|
||||
crate::image_capabilities::openai_image_provider_max_generation_count_for_model(
|
||||
transport.provider.provider_type.as_str(),
|
||||
Some(prepared_candidate.mapped_model.as_str()),
|
||||
),
|
||||
)?
|
||||
};
|
||||
let request_path = match converted.operation {
|
||||
OpenAiImageOperation::Generate => "/v1/images/generations",
|
||||
OpenAiImageOperation::Edit => "/v1/images/edits",
|
||||
|
||||
@@ -42,13 +42,15 @@ pub(crate) use self::openai::{
|
||||
build_local_openai_responses_sync_attempt_source_for_kind,
|
||||
build_local_openai_responses_sync_plan_and_reports_for_kind, copy_request_number_field,
|
||||
copy_request_number_field_as, map_openai_reasoning_effort_to_claude_output,
|
||||
map_openai_reasoning_effort_to_gemini_budget, maybe_build_stream_local_decision_payload,
|
||||
map_openai_reasoning_effort_to_gemini_budget, maybe_build_responses_websocket_decision,
|
||||
maybe_build_stream_local_decision_payload,
|
||||
maybe_build_stream_local_openai_responses_decision_payload,
|
||||
maybe_build_sync_local_decision_payload,
|
||||
maybe_build_sync_local_openai_embedding_decision_payload,
|
||||
maybe_build_sync_local_openai_responses_decision_payload, parse_openai_stop_sequences,
|
||||
resolve_openai_chat_max_tokens, set_local_openai_chat_execution_exhausted_diagnostic,
|
||||
value_as_u64,
|
||||
value_as_u64, ResponsesWebSocketBodyNormalization, ResponsesWebSocketDecision,
|
||||
ResponsesWebSocketPinnedCandidate,
|
||||
};
|
||||
pub(crate) use crate::ai_serving::normalize_standard_request_to_openai_chat_request;
|
||||
pub(crate) use crate::ai_serving::{
|
||||
|
||||
@@ -352,7 +352,7 @@ fn final_openai_provider_contract_uses_the_mapped_model_for_reasoning() {
|
||||
false,
|
||||
)
|
||||
.is_some());
|
||||
assert!(build_local_openai_responses_request_body(
|
||||
let remapped = build_local_openai_responses_request_body(
|
||||
&alias,
|
||||
"gpt-5.4",
|
||||
false,
|
||||
@@ -364,14 +364,15 @@ fn final_openai_provider_contract_uses_the_mapped_model_for_reasoning() {
|
||||
&http::HeaderMap::new(),
|
||||
false,
|
||||
)
|
||||
.is_none());
|
||||
.expect("explicit reasoning effort should pass through to the mapped model");
|
||||
assert_eq!(remapped["reasoning"]["effort"], "max");
|
||||
|
||||
let minimal = json!({
|
||||
"model": "deployment-alias",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"reasoning_effort": "minimal"
|
||||
});
|
||||
assert!(build_local_openai_chat_request_body(
|
||||
let minimal = build_local_openai_chat_request_body(
|
||||
&minimal,
|
||||
"gpt-5.6-terra",
|
||||
false,
|
||||
@@ -380,7 +381,8 @@ fn final_openai_provider_contract_uses_the_mapped_model_for_reasoning() {
|
||||
&http::HeaderMap::new(),
|
||||
false,
|
||||
)
|
||||
.is_none());
|
||||
.expect("explicit chat reasoning effort should be validated by the upstream");
|
||||
assert_eq!(minimal["reasoning_effort"], "minimal");
|
||||
|
||||
let opaque_mapping = json!({
|
||||
"model": "gpt-5.6-sol-max",
|
||||
@@ -426,7 +428,7 @@ fn final_openai_provider_contract_validates_body_rule_output() {
|
||||
let model_override = json!([
|
||||
{"action":"set","path":"model","value":"gpt-5.4"}
|
||||
]);
|
||||
assert!(build_local_openai_responses_request_body(
|
||||
let provider_request = build_local_openai_responses_request_body(
|
||||
&body,
|
||||
"gpt-5.6-sol",
|
||||
false,
|
||||
@@ -438,7 +440,9 @@ fn final_openai_provider_contract_validates_body_rule_output() {
|
||||
&http::HeaderMap::new(),
|
||||
false,
|
||||
)
|
||||
.is_none());
|
||||
.expect("body rule output should preserve explicit reasoning effort");
|
||||
assert_eq!(provider_request["model"], "gpt-5.4");
|
||||
assert_eq!(provider_request["reasoning"]["effort"], "max");
|
||||
|
||||
let cache_override = json!([
|
||||
{"action":"set","path":"prompt_cache_options.ttl","value":"1h"}
|
||||
|
||||
+12
-29
@@ -1417,38 +1417,20 @@ async fn resolve_openai_chat_to_openai_image_payload_parts(
|
||||
return Ok(None);
|
||||
};
|
||||
if !is_chatgpt_web {
|
||||
let Some(projected) = project_openai_image_api_request_body(
|
||||
&provider_request_body,
|
||||
&prepared_candidate.mapped_model,
|
||||
operation,
|
||||
crate::image_capabilities::openai_image_provider_max_generation_count_for_model(
|
||||
transport.provider.provider_type.as_str(),
|
||||
Some(prepared_candidate.mapped_model.as_str()),
|
||||
),
|
||||
) else {
|
||||
mark_skipped_local_openai_chat_candidate_with_extra_data(
|
||||
state,
|
||||
input,
|
||||
trace_id,
|
||||
candidate,
|
||||
candidate_index,
|
||||
candidate_id,
|
||||
"provider_request_body_build_failed",
|
||||
request_body_build_failure_extra_data(
|
||||
body_json,
|
||||
"openai:chat",
|
||||
provider_api_format,
|
||||
let projected = if is_codex {
|
||||
project_codex_openai_image_api_request_body(&provider_request_body, operation)
|
||||
} else {
|
||||
project_openai_image_api_request_body(
|
||||
&provider_request_body,
|
||||
&prepared_candidate.mapped_model,
|
||||
operation,
|
||||
crate::image_capabilities::openai_image_provider_max_generation_count_for_model(
|
||||
transport.provider.provider_type.as_str(),
|
||||
Some(prepared_candidate.mapped_model.as_str()),
|
||||
),
|
||||
)
|
||||
.await;
|
||||
return Ok(None);
|
||||
};
|
||||
provider_request_body = projected;
|
||||
}
|
||||
if is_codex {
|
||||
let Some(projected) =
|
||||
project_codex_openai_image_api_request_body(&provider_request_body, operation)
|
||||
else {
|
||||
let Some(projected) = projected else {
|
||||
mark_skipped_local_openai_chat_candidate_with_extra_data(
|
||||
state,
|
||||
input,
|
||||
@@ -2195,6 +2177,7 @@ mod tests {
|
||||
client_surface: None,
|
||||
gateway_credential_carrier: None,
|
||||
client_session_affinity: None,
|
||||
original_client_session_id: None,
|
||||
routing_policy: None,
|
||||
routing_trace_seed: None,
|
||||
routing_context: None,
|
||||
|
||||
@@ -23,6 +23,8 @@ pub(crate) use responses::{
|
||||
build_local_openai_responses_stream_plan_and_reports_for_kind,
|
||||
build_local_openai_responses_sync_attempt_source_for_kind,
|
||||
build_local_openai_responses_sync_plan_and_reports_for_kind,
|
||||
maybe_build_responses_websocket_decision,
|
||||
maybe_build_stream_local_openai_responses_decision_payload,
|
||||
maybe_build_sync_local_openai_responses_decision_payload,
|
||||
maybe_build_sync_local_openai_responses_decision_payload, ResponsesWebSocketBodyNormalization,
|
||||
ResponsesWebSocketDecision, ResponsesWebSocketPinnedCandidate,
|
||||
};
|
||||
|
||||
@@ -9,7 +9,9 @@ pub(super) use self::payload::maybe_build_local_openai_responses_decision_payloa
|
||||
pub(super) use self::support::{
|
||||
build_local_openai_responses_candidate_attempt_source,
|
||||
materialize_local_openai_responses_candidate_attempts,
|
||||
resolve_local_openai_responses_decision_input, LocalOpenAiResponsesCandidateAttempt,
|
||||
LocalOpenAiResponsesCandidateAttemptSource, LocalOpenAiResponsesDecisionInput,
|
||||
resolve_local_openai_responses_decision_input,
|
||||
resolve_local_openai_responses_decision_input_with_snapshot,
|
||||
LocalOpenAiResponsesCandidateAttempt, LocalOpenAiResponsesCandidateAttemptSource,
|
||||
LocalOpenAiResponsesDecisionInput,
|
||||
};
|
||||
pub(super) use crate::ai_serving::LocalOpenAiResponsesSpec;
|
||||
|
||||
+4
-5
@@ -1395,7 +1395,10 @@ async fn resolve_openai_responses_to_openai_image_payload_parts(
|
||||
return None;
|
||||
};
|
||||
let operation = openai_image_operation_from_summary(&image_request_summary)?;
|
||||
if !is_chatgpt_web {
|
||||
if is_codex {
|
||||
provider_request_body =
|
||||
project_codex_openai_image_api_request_body(&provider_request_body, operation)?;
|
||||
} else if !is_chatgpt_web {
|
||||
provider_request_body = project_openai_image_api_request_body(
|
||||
&provider_request_body,
|
||||
&prepared_candidate.mapped_model,
|
||||
@@ -1406,10 +1409,6 @@ async fn resolve_openai_responses_to_openai_image_payload_parts(
|
||||
),
|
||||
)?;
|
||||
}
|
||||
if is_codex {
|
||||
provider_request_body =
|
||||
project_codex_openai_image_api_request_body(&provider_request_body, operation)?;
|
||||
}
|
||||
|
||||
let upstream_url = if is_chatgpt_web {
|
||||
chatgpt_web_image_internal_url(&transport.endpoint.base_url)
|
||||
|
||||
+40
-11
@@ -22,6 +22,7 @@ use crate::ai_serving::planner::common::extract_standard_requested_model;
|
||||
use crate::ai_serving::planner::decision_input::{
|
||||
attach_routing_policy_to_local_requested_model_input,
|
||||
build_local_requested_model_decision_input, resolve_local_authenticated_decision_input,
|
||||
resolve_local_authenticated_decision_input_with_snapshot,
|
||||
};
|
||||
use crate::ai_serving::planner::materialization_policy::{
|
||||
build_local_candidate_persistence_policy, LocalCandidatePersistencePolicyKind,
|
||||
@@ -32,7 +33,8 @@ use crate::ai_serving::planner::CandidateFailureDiagnostic;
|
||||
use crate::ai_serving::{
|
||||
ai_local_execution_contract_for_formats, extract_pool_sticky_session_token,
|
||||
openai_responses_request_operation, resolve_local_decision_execution_runtime_auth_context,
|
||||
ExecutionRuntimeAuthContext, GatewayControlDecision, PlannerAppState,
|
||||
ExecutionRuntimeAuthContext, GatewayAuthApiKeySnapshot, GatewayControlDecision,
|
||||
PlannerAppState,
|
||||
};
|
||||
use crate::client_session_affinity::client_session_affinity_from_parts;
|
||||
use crate::{AppState, GatewayError};
|
||||
@@ -51,6 +53,21 @@ pub(crate) async fn resolve_local_openai_responses_decision_input(
|
||||
decision: &GatewayControlDecision,
|
||||
body_json: &serde_json::Value,
|
||||
plan_kind: &str,
|
||||
) -> Result<Option<LocalOpenAiResponsesDecisionInput>, GatewayError> {
|
||||
resolve_local_openai_responses_decision_input_with_snapshot(
|
||||
state, parts, trace_id, decision, body_json, plan_kind, None,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn resolve_local_openai_responses_decision_input_with_snapshot(
|
||||
state: &AppState,
|
||||
parts: &http::request::Parts,
|
||||
trace_id: &str,
|
||||
decision: &GatewayControlDecision,
|
||||
body_json: &serde_json::Value,
|
||||
plan_kind: &str,
|
||||
auth_snapshot_override: Option<&GatewayAuthApiKeySnapshot>,
|
||||
) -> Result<Option<LocalOpenAiResponsesDecisionInput>, GatewayError> {
|
||||
let Some(auth_context) = resolve_local_decision_execution_runtime_auth_context(decision) else {
|
||||
warn!(
|
||||
@@ -87,16 +104,28 @@ pub(crate) async fn resolve_local_openai_responses_decision_input(
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let resolved_input = match resolve_local_authenticated_decision_input(
|
||||
state,
|
||||
auth_context.clone(),
|
||||
Some(requested_model.as_str()),
|
||||
decision.auth_endpoint_signature.as_deref(),
|
||||
None,
|
||||
&decision.model_directive_policy,
|
||||
)
|
||||
.await
|
||||
{
|
||||
let resolved_input = match if let Some(auth_snapshot) = auth_snapshot_override {
|
||||
resolve_local_authenticated_decision_input_with_snapshot(
|
||||
state,
|
||||
auth_context.clone(),
|
||||
Some(auth_snapshot.clone()),
|
||||
Some(requested_model.as_str()),
|
||||
decision.auth_endpoint_signature.as_deref(),
|
||||
None,
|
||||
&decision.model_directive_policy,
|
||||
)
|
||||
.await
|
||||
} else {
|
||||
resolve_local_authenticated_decision_input(
|
||||
state,
|
||||
auth_context.clone(),
|
||||
Some(requested_model.as_str()),
|
||||
decision.auth_endpoint_signature.as_deref(),
|
||||
None,
|
||||
&decision.model_directive_policy,
|
||||
)
|
||||
.await
|
||||
} {
|
||||
Ok(Some(resolved_input)) => resolved_input,
|
||||
Ok(None) => {
|
||||
warn!(
|
||||
|
||||
@@ -1,6 +1,59 @@
|
||||
use crate::ai_serving::planner::common::endpoint_config_forces_body_stream_field;
|
||||
use crate::ai_serving::planner::plan_builders::{AiStreamAttempt, AiSyncAttempt};
|
||||
use crate::ai_serving::planner::spec_metadata::local_openai_responses_spec_metadata;
|
||||
use crate::ai_serving::planner::standard::codex::codex_model_capabilities_for_transport;
|
||||
use crate::ai_serving::planner::standard::normalize::build_local_openai_responses_request_body_with_codex_model_capabilities;
|
||||
use crate::ai_serving::GatewayControlDecision;
|
||||
use crate::orchestration::{
|
||||
codex_quota_breaker_blocks_candidate, log_codex_quota_breaker_check_failure,
|
||||
responses_websocket_adapter, ResponsesWebSocketAdapter,
|
||||
};
|
||||
use crate::{AiExecutionDecision, AppState, GatewayError};
|
||||
use aether_runtime_state::RuntimeLockLease;
|
||||
use std::collections::BTreeSet;
|
||||
|
||||
/// Releases a scheduler pool-key lease if WebSocket planning is cancelled
|
||||
/// after candidate selection but before ownership reaches the turn lifecycle.
|
||||
struct ResponsesWebSocketPlanningLeaseGuard {
|
||||
state: AppState,
|
||||
lease: Option<RuntimeLockLease>,
|
||||
}
|
||||
|
||||
impl ResponsesWebSocketPlanningLeaseGuard {
|
||||
fn new(state: &AppState, lease: Option<&RuntimeLockLease>) -> Self {
|
||||
Self {
|
||||
state: state.clone(),
|
||||
lease: lease.cloned(),
|
||||
}
|
||||
}
|
||||
|
||||
async fn release(mut self) {
|
||||
// Keep the lease armed across the await. If the owner task is aborted
|
||||
// or reaches its hard deadline while the runtime backend is stalled,
|
||||
// Drop can still hand cleanup to a detached owner.
|
||||
if release_responses_websocket_planning_lease(&self.state, self.lease.as_ref()).await {
|
||||
self.lease = None;
|
||||
}
|
||||
}
|
||||
|
||||
fn disarm(&mut self) {
|
||||
self.lease = None;
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for ResponsesWebSocketPlanningLeaseGuard {
|
||||
fn drop(&mut self) {
|
||||
let Some(lease) = self.lease.take() else {
|
||||
return;
|
||||
};
|
||||
let state = self.state.clone();
|
||||
if let Ok(handle) = tokio::runtime::Handle::try_current() {
|
||||
handle.spawn(async move {
|
||||
let _ = release_responses_websocket_planning_lease(&state, Some(&lease)).await;
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
mod decision;
|
||||
mod plans;
|
||||
@@ -9,6 +62,7 @@ use self::decision::{
|
||||
build_local_openai_responses_candidate_attempt_source,
|
||||
maybe_build_local_openai_responses_decision_payload_for_candidate,
|
||||
resolve_local_openai_responses_decision_input,
|
||||
resolve_local_openai_responses_decision_input_with_snapshot,
|
||||
};
|
||||
use self::plans::{
|
||||
build_local_stream_attempt_source, build_local_stream_plan_and_reports,
|
||||
@@ -165,3 +219,373 @@ pub(crate) async fn maybe_build_stream_local_openai_responses_decision_payload(
|
||||
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
/// One eligible upstream plus the adapter that is allowed to speak to it.
|
||||
///
|
||||
/// The adapter is selected from the provider-scoped capability before the
|
||||
/// decision leaves the planner. This prevents a public Responses socket from
|
||||
/// choosing an arbitrary provider protocol after scheduling has completed.
|
||||
pub(crate) struct ResponsesWebSocketDecision {
|
||||
pub(crate) execution: AiExecutionDecision,
|
||||
pub(crate) adapter: ResponsesWebSocketAdapter,
|
||||
pub(crate) normalization: ResponsesWebSocketBodyNormalization,
|
||||
}
|
||||
|
||||
/// The scheduler identity a continuation is allowed to reuse.
|
||||
///
|
||||
/// A `previous_response_id` chain cannot move to another provider connection,
|
||||
/// but it still has to pass the current scheduler runtime checks on every
|
||||
/// turn. The planner uses this identity as a filter rather than selecting an
|
||||
/// arbitrary eligible replacement.
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub(crate) struct ResponsesWebSocketPinnedCandidate {
|
||||
provider_id: String,
|
||||
endpoint_id: String,
|
||||
key_id: String,
|
||||
}
|
||||
|
||||
impl ResponsesWebSocketPinnedCandidate {
|
||||
pub(crate) fn from_decision(decision: &AiExecutionDecision) -> Option<Self> {
|
||||
Some(Self {
|
||||
provider_id: non_empty_decision_identity(decision.provider_id.as_deref())?,
|
||||
endpoint_id: non_empty_decision_identity(decision.endpoint_id.as_deref())?,
|
||||
key_id: non_empty_decision_identity(decision.key_id.as_deref())?,
|
||||
})
|
||||
}
|
||||
|
||||
fn matches(
|
||||
&self,
|
||||
candidate: &aether_scheduler_core::SchedulerMinimalCandidateSelectionCandidate,
|
||||
) -> bool {
|
||||
candidate.provider_id == self.provider_id
|
||||
&& candidate.endpoint_id == self.endpoint_id
|
||||
&& candidate.key_id == self.key_id
|
||||
}
|
||||
}
|
||||
|
||||
fn non_empty_decision_identity(value: Option<&str>) -> Option<String> {
|
||||
value
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(str::to_string)
|
||||
}
|
||||
|
||||
/// Everything needed to re-run provider-body normalization for the candidate a
|
||||
/// socket is already bound to.
|
||||
///
|
||||
/// A continuation turn (`previous_response_id` on the bound upstream) cannot
|
||||
/// re-enter the planner, because planning selects a candidate and a different
|
||||
/// key would break the response chain. Without this, such turns reached the
|
||||
/// provider with only their `model` rewritten — skipping model directives,
|
||||
/// endpoint body rules, and the Codex body contract that turn 1 received.
|
||||
///
|
||||
/// This value holds cloned scalars and JSON only: no candidate, no pool key
|
||||
/// lease, no `AppState`. It cannot influence selection.
|
||||
#[derive(Debug, Clone)]
|
||||
pub(crate) struct ResponsesWebSocketBodyNormalization {
|
||||
provider_type: String,
|
||||
provider_api_format: String,
|
||||
client_api_format: String,
|
||||
mapped_model: String,
|
||||
requested_model: String,
|
||||
upstream_is_stream: bool,
|
||||
force_body_stream_field: bool,
|
||||
body_rules: Option<serde_json::Value>,
|
||||
request_headers: http::HeaderMap,
|
||||
codex_model_capabilities: Option<crate::ai_serving::CodexResponsesModelCapabilities>,
|
||||
model_directive_patch: Option<serde_json::Value>,
|
||||
}
|
||||
|
||||
impl ResponsesWebSocketBodyNormalization {
|
||||
/// Builds a normalizer for a plain `openai:responses` upstream with no
|
||||
/// endpoint body rules, directives or Codex capabilities, so relay tests can
|
||||
/// construct a bound connection without standing up a provider snapshot.
|
||||
#[cfg(test)]
|
||||
pub(crate) fn for_tests(mapped_model: &str) -> Self {
|
||||
Self {
|
||||
provider_type: "openai".to_string(),
|
||||
provider_api_format: "openai:responses".to_string(),
|
||||
client_api_format: "openai:responses".to_string(),
|
||||
mapped_model: mapped_model.to_string(),
|
||||
requested_model: mapped_model.to_string(),
|
||||
upstream_is_stream: true,
|
||||
force_body_stream_field: false,
|
||||
body_rules: None,
|
||||
request_headers: http::HeaderMap::new(),
|
||||
codex_model_capabilities: None,
|
||||
model_directive_patch: None,
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub(crate) fn with_provider_type_for_tests(mut self, provider_type: &str) -> Self {
|
||||
self.provider_type = provider_type.to_string();
|
||||
self
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub(crate) fn with_model_directive_patch_for_tests(mut self, patch: serde_json::Value) -> Self {
|
||||
self.model_directive_patch = Some(patch);
|
||||
self
|
||||
}
|
||||
|
||||
/// Applies the same body transformations the planner applied on the turn
|
||||
/// that bound this upstream.
|
||||
///
|
||||
/// Mirrors the same-format branch of
|
||||
/// `resolve_local_openai_responses_candidate_payload_parts`. The
|
||||
/// cross-format, Kiro, Windsurf and Antigravity branches are unreachable
|
||||
/// here: the WebSocket planner only returns candidates whose provider API
|
||||
/// format is `openai:responses`.
|
||||
///
|
||||
/// Returns `None` when normalization fails, leaving the caller to fall back
|
||||
/// to the unnormalized event — a continuation cannot re-select a candidate,
|
||||
/// so failing the turn outright would be worse than sending it as-is.
|
||||
pub(crate) fn normalize_response_create(
|
||||
&self,
|
||||
client_event: &serde_json::Value,
|
||||
) -> Option<serde_json::Value> {
|
||||
use crate::ai_serving::planner::common::{
|
||||
enforce_provider_body_stream_policy, request_requires_body_stream_field,
|
||||
};
|
||||
|
||||
let source_model = client_event
|
||||
.get("model")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.unwrap_or(self.requested_model.as_str());
|
||||
let require_body_stream_field =
|
||||
request_requires_body_stream_field(client_event, self.force_body_stream_field);
|
||||
let mut body = build_local_openai_responses_request_body_with_codex_model_capabilities(
|
||||
client_event,
|
||||
&self.mapped_model,
|
||||
self.upstream_is_stream,
|
||||
self.force_body_stream_field,
|
||||
self.provider_type.as_str(),
|
||||
self.provider_api_format.as_str(),
|
||||
self.body_rules.as_ref(),
|
||||
&self.request_headers,
|
||||
self.codex_model_capabilities.as_ref(),
|
||||
false,
|
||||
)?;
|
||||
if let Some(patch) = self.model_directive_patch.as_ref() {
|
||||
crate::ai_serving::apply_model_directive_mapping_patch(&mut body, patch);
|
||||
// The patch is a deep merge and may reintroduce `stream`.
|
||||
enforce_provider_body_stream_policy(
|
||||
&mut body,
|
||||
self.provider_api_format.as_str(),
|
||||
self.upstream_is_stream,
|
||||
require_body_stream_field,
|
||||
);
|
||||
}
|
||||
crate::ai_serving::finalize_openai_provider_request_with_codex_model_capabilities(
|
||||
&mut body,
|
||||
crate::ai_serving::OpenAiProviderRequestFinalization {
|
||||
source_api_format: self.client_api_format.as_str(),
|
||||
provider_api_format: self.provider_api_format.as_str(),
|
||||
provider_type: self.provider_type.as_str(),
|
||||
provider_model: self.mapped_model.as_str(),
|
||||
source_model,
|
||||
body_rules: self.body_rules.as_ref(),
|
||||
upstream_is_stream: self.upstream_is_stream,
|
||||
require_body_stream_field,
|
||||
},
|
||||
self.codex_model_capabilities.as_ref(),
|
||||
)
|
||||
.ok()?;
|
||||
Some(body)
|
||||
}
|
||||
}
|
||||
|
||||
/// Builds one upstream decision for a Responses WebSocket turn. The session
|
||||
/// reuses this decision for same-model turns and invokes the planner again when
|
||||
/// a later `response.create` changes the public model.
|
||||
pub(crate) async fn maybe_build_responses_websocket_decision(
|
||||
state: &AppState,
|
||||
parts: &http::request::Parts,
|
||||
trace_id: &str,
|
||||
decision: &GatewayControlDecision,
|
||||
auth_snapshot: Option<&crate::ai_serving::GatewayAuthApiKeySnapshot>,
|
||||
body_json: &serde_json::Value,
|
||||
excluded_key_ids: Option<&BTreeSet<String>>,
|
||||
excluded_codex_account_ids: Option<&BTreeSet<String>>,
|
||||
pinned_candidate: Option<&ResponsesWebSocketPinnedCandidate>,
|
||||
) -> Result<Option<ResponsesWebSocketDecision>, GatewayError> {
|
||||
let Some(spec) = resolve_stream_spec(crate::ai_serving::OPENAI_RESPONSES_STREAM_PLAN_KIND)
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
let Some(input) = resolve_local_openai_responses_decision_input_with_snapshot(
|
||||
state,
|
||||
parts,
|
||||
trace_id,
|
||||
decision,
|
||||
body_json,
|
||||
spec.decision_kind,
|
||||
auth_snapshot,
|
||||
)
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
let body_json = input.effective_body_json(body_json);
|
||||
let (mut source, _) = build_local_openai_responses_candidate_attempt_source(
|
||||
state, trace_id, &input, body_json, spec,
|
||||
)
|
||||
.await?;
|
||||
|
||||
while let Some(attempt) = source.next_attempt().await? {
|
||||
// `next_attempt` may return with a distributed pool-key lease. Arm a
|
||||
// guard before the first await so owner-task timeout/cancellation
|
||||
// cannot strand that lease until its server-side TTL expires.
|
||||
let mut planning_lease = ResponsesWebSocketPlanningLeaseGuard::new(
|
||||
state,
|
||||
attempt.eligible.orchestration.pool_key_lease.as_ref(),
|
||||
);
|
||||
if pinned_candidate.is_some_and(|pinned| !pinned.matches(&attempt.eligible.candidate)) {
|
||||
planning_lease.release().await;
|
||||
continue;
|
||||
}
|
||||
if excluded_key_ids
|
||||
.is_some_and(|key_ids| key_ids.contains(attempt.eligible.candidate.key_id.as_str()))
|
||||
{
|
||||
planning_lease.release().await;
|
||||
continue;
|
||||
}
|
||||
let Some(adapter) = responses_websocket_adapter(
|
||||
&attempt.eligible.transport.provider.provider_type,
|
||||
attempt.eligible.transport.provider.config.as_ref(),
|
||||
) else {
|
||||
planning_lease.release().await;
|
||||
continue;
|
||||
};
|
||||
// Captured before `attempt` is consumed so a later continuation turn can
|
||||
// reproduce this candidate's body normalization without re-planning.
|
||||
let transport = std::sync::Arc::clone(&attempt.eligible.transport);
|
||||
let candidate_provider_api_format = attempt.eligible.provider_api_format.clone();
|
||||
let payload = match maybe_build_local_openai_responses_decision_payload_for_candidate(
|
||||
state, parts, trace_id, body_json, &input, attempt, spec,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(Some(payload)) => payload,
|
||||
Ok(None) => {
|
||||
planning_lease.release().await;
|
||||
continue;
|
||||
}
|
||||
Err(error) => {
|
||||
planning_lease.release().await;
|
||||
return Err(error);
|
||||
}
|
||||
};
|
||||
if payload
|
||||
.provider_type
|
||||
.as_deref()
|
||||
.is_some_and(|value| value.trim().eq_ignore_ascii_case("codex"))
|
||||
&& crate::orchestration::codex_account_id_from_headers(
|
||||
&payload.provider_request_headers,
|
||||
)
|
||||
.is_some_and(|account_id| {
|
||||
excluded_codex_account_ids
|
||||
.is_some_and(|account_ids| account_ids.contains(account_id))
|
||||
})
|
||||
{
|
||||
planning_lease.release().await;
|
||||
continue;
|
||||
}
|
||||
match codex_quota_breaker_blocks_candidate(
|
||||
state,
|
||||
payload.provider_type.as_deref(),
|
||||
payload.key_id.as_deref(),
|
||||
&payload.provider_request_headers,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(true) => {
|
||||
planning_lease.release().await;
|
||||
continue;
|
||||
}
|
||||
Ok(false) => {}
|
||||
Err(error) => log_codex_quota_breaker_check_failure(&error),
|
||||
}
|
||||
if payload
|
||||
.provider_type
|
||||
.as_deref()
|
||||
.is_some_and(|value| adapter.supports_provider_type(value))
|
||||
&& payload.provider_api_format.as_deref().is_some_and(|value| {
|
||||
crate::ai_serving::normalize_api_format_alias(value) == "openai:responses"
|
||||
})
|
||||
{
|
||||
let mapped_model = payload.mapped_model.clone().unwrap_or_default();
|
||||
let source_model = body_json
|
||||
.get("model")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.unwrap_or(input.requested_model.as_str());
|
||||
let normalization = ResponsesWebSocketBodyNormalization {
|
||||
provider_type: transport.provider.provider_type.clone(),
|
||||
provider_api_format: candidate_provider_api_format.clone(),
|
||||
client_api_format: local_openai_responses_spec_metadata(spec)
|
||||
.api_format
|
||||
.to_string(),
|
||||
requested_model: input.requested_model.clone(),
|
||||
upstream_is_stream: payload.upstream_is_stream,
|
||||
force_body_stream_field: endpoint_config_forces_body_stream_field(
|
||||
transport.endpoint.config.as_ref(),
|
||||
),
|
||||
body_rules: transport.endpoint.body_rules.clone(),
|
||||
request_headers: input.effective_headers(&parts.headers).clone(),
|
||||
codex_model_capabilities: codex_model_capabilities_for_transport(
|
||||
&transport,
|
||||
candidate_provider_api_format.as_str(),
|
||||
mapped_model.as_str(),
|
||||
source_model,
|
||||
),
|
||||
model_directive_patch: input
|
||||
.model_directive_policy
|
||||
.resolve_reasoning(
|
||||
candidate_provider_api_format.as_str(),
|
||||
Some(&input.requested_model),
|
||||
)
|
||||
.mapping_patch_for_mapped_model(mapped_model.as_str())
|
||||
.ok()
|
||||
.flatten(),
|
||||
mapped_model,
|
||||
};
|
||||
let decision = ResponsesWebSocketDecision {
|
||||
execution: payload,
|
||||
adapter,
|
||||
normalization,
|
||||
};
|
||||
// The decision report context now carries the lease identity. The
|
||||
// WebSocket ownership layer takes over before any further await.
|
||||
planning_lease.disarm();
|
||||
return Ok(Some(decision));
|
||||
}
|
||||
planning_lease.release().await;
|
||||
}
|
||||
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
async fn release_responses_websocket_planning_lease(
|
||||
state: &AppState,
|
||||
lease: Option<&RuntimeLockLease>,
|
||||
) -> bool {
|
||||
let Some(lease) = lease else {
|
||||
return true;
|
||||
};
|
||||
match crate::handlers::shared::provider_pool::release_admin_provider_pool_key_lease(
|
||||
state.runtime_state.as_ref(),
|
||||
lease,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(_) => true,
|
||||
Err(error) => {
|
||||
tracing::warn!(
|
||||
error = ?error,
|
||||
"gateway Responses WebSocket planner failed to release an unused pool key lease"
|
||||
);
|
||||
false
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -82,10 +82,10 @@ pub(crate) use aether_ai_formats::api::{
|
||||
parse_codex_auth_identity, parse_direct_request_body, parse_model_directive,
|
||||
parse_model_directive_with_suffixes, parse_openai_stop_sequences,
|
||||
parse_openai_tool_result_content, prepare_local_success_response_parts,
|
||||
prepare_local_success_response_parts_owned, project_codex_openai_image_api_request_body,
|
||||
project_openai_image_api_request_body, provider_adaptation_allows_sync_finalize_envelope,
|
||||
provider_adaptation_anchor_api_format, provider_adaptation_descriptor_for_envelope,
|
||||
provider_adaptation_descriptor_for_provider_type,
|
||||
prepare_local_success_response_parts_owned, project_codex_catalog_model_card,
|
||||
project_codex_openai_image_api_request_body, project_openai_image_api_request_body,
|
||||
provider_adaptation_allows_sync_finalize_envelope, provider_adaptation_anchor_api_format,
|
||||
provider_adaptation_descriptor_for_envelope, provider_adaptation_descriptor_for_provider_type,
|
||||
provider_adaptation_requires_eventstream_accept,
|
||||
provider_adaptation_should_unwrap_stream_envelope,
|
||||
provider_private_response_allows_sync_finalize, record_converted_response_history,
|
||||
@@ -175,6 +175,7 @@ pub(crate) use aether_ai_formats::{
|
||||
is_rerank_api_format, openai_responses_request_operation,
|
||||
openai_responses_synthetic_reasoning_item_id,
|
||||
strip_incompatible_openai_responses_reasoning_items, ApiOperation, ClientSurface,
|
||||
CODEX_CLIENT_VERSION,
|
||||
};
|
||||
|
||||
pub(crate) fn plan_kind_matches_api_operation(
|
||||
|
||||
@@ -59,8 +59,9 @@ pub(crate) mod windsurf {
|
||||
}
|
||||
|
||||
pub(crate) use aether_provider_transport::{
|
||||
append_transport_diagnostics_to_value, apply_local_auth_config_header_overrides,
|
||||
apply_local_body_rules, apply_local_body_rules_with_request_headers, apply_local_header_rules,
|
||||
append_transport_diagnostics_to_value, apply_codex_oauth_fingerprint_convergence,
|
||||
apply_local_auth_config_header_overrides, apply_local_body_rules,
|
||||
apply_local_body_rules_with_request_headers, apply_local_header_rules,
|
||||
apply_local_header_rules_with_request_headers, apply_standard_provider_request_body_rules,
|
||||
apply_standard_provider_request_body_rules_with_request_headers,
|
||||
apply_transport_request_body_semantics, body_rules_are_locally_supported,
|
||||
|
||||
@@ -1,13 +1,17 @@
|
||||
use axum::body::Body;
|
||||
use axum::extract::Request;
|
||||
use axum::http::{header, HeaderValue, Response, StatusCode};
|
||||
use axum::routing::{any, post};
|
||||
use axum::routing::{any, get, post};
|
||||
use axum::Router;
|
||||
|
||||
use super::{aliyun, claude, doubao, gemini, jina, openai};
|
||||
use crate::api::response::build_local_http_error_response_with_request_path;
|
||||
use crate::headers::extract_or_generate_trace_id;
|
||||
use crate::{handlers::proxy::proxy_request, state::AppState, GatewayError};
|
||||
use crate::{
|
||||
handlers::proxy::{proxy_request, responses_websocket},
|
||||
state::AppState,
|
||||
GatewayError,
|
||||
};
|
||||
|
||||
// Router registration patterns live here so AI public ingress has a single mount registry.
|
||||
// They intentionally stay separate from manifest-facing route inventories in constants.rs,
|
||||
@@ -51,7 +55,11 @@ const AI_ANY_ROUTE_PATTERNS: &[&str] = &[
|
||||
|
||||
pub(crate) fn mount_ai_routes(mut router: Router<AppState>) -> Router<AppState> {
|
||||
for path in AI_POST_ROUTE_PATTERNS {
|
||||
router = router.route(path, post(proxy_request));
|
||||
router = if *path == "/v1/responses" {
|
||||
router.route(path, get(responses_websocket).post(proxy_request))
|
||||
} else {
|
||||
router.route(path, post(proxy_request))
|
||||
};
|
||||
}
|
||||
for path in CLAUDE_POST_ROUTE_PATTERNS {
|
||||
router = router.route(
|
||||
|
||||
@@ -50,12 +50,40 @@ pub(crate) async fn health(State(state): State<AppState>) -> impl IntoResponse {
|
||||
"rejected": snapshot.rejected,
|
||||
})
|
||||
});
|
||||
let websocket_connection_concurrency =
|
||||
state
|
||||
.websocket_connection_concurrency_snapshot()
|
||||
.map(|snapshot| {
|
||||
json!({
|
||||
"limit": snapshot.limit,
|
||||
"in_flight": snapshot.in_flight,
|
||||
"available_permits": snapshot.available_permits,
|
||||
"high_watermark": snapshot.high_watermark,
|
||||
"rejected": snapshot.rejected,
|
||||
})
|
||||
});
|
||||
let distributed_websocket_connection_concurrency = state
|
||||
.distributed_websocket_connection_concurrency_snapshot()
|
||||
.await
|
||||
.ok()
|
||||
.flatten()
|
||||
.map(|snapshot| {
|
||||
json!({
|
||||
"limit": snapshot.limit,
|
||||
"in_flight": snapshot.in_flight,
|
||||
"available_permits": snapshot.available_permits,
|
||||
"high_watermark": snapshot.high_watermark,
|
||||
"rejected": snapshot.rejected,
|
||||
})
|
||||
});
|
||||
Json(json!({
|
||||
"status": "ok",
|
||||
"component": "aether-gateway",
|
||||
"control_api_enabled": true,
|
||||
"request_concurrency": request_concurrency,
|
||||
"distributed_request_concurrency": distributed_request_concurrency,
|
||||
"websocket_connection_concurrency": websocket_connection_concurrency,
|
||||
"distributed_websocket_connection_concurrency": distributed_websocket_connection_concurrency,
|
||||
}))
|
||||
}
|
||||
|
||||
@@ -113,6 +141,12 @@ pub(crate) async fn frontdoor_manifest(State(state): State<AppState>) -> impl In
|
||||
"execution_runtime_configured": state.execution_runtime_configured(),
|
||||
"request_concurrency_enabled": state.request_concurrency_snapshot().is_some(),
|
||||
"distributed_request_concurrency_enabled": state.distributed_request_gate.is_some(),
|
||||
"websocket_connection_concurrency_enabled": state
|
||||
.websocket_connection_concurrency_snapshot()
|
||||
.is_some(),
|
||||
"distributed_websocket_connection_concurrency_enabled": state
|
||||
.distributed_websocket_connection_gate
|
||||
.is_some(),
|
||||
"frontdoor_cors_enabled": cors_enabled,
|
||||
"frontdoor_cors_allow_credentials": cors_allow_credentials,
|
||||
"frontdoor_cors_allowed_origins": cors_allowed_origins,
|
||||
|
||||
@@ -85,29 +85,17 @@ async fn read_last_backup_slot(app: &AppState) -> Result<Option<String>, Gateway
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use crate::backup::schedule::{BackupSchedule, BackupScheduleUnit};
|
||||
use crate::task_runtime::{task_definition, TASK_KEY_SYSTEM_S3_BACKUP};
|
||||
|
||||
#[test]
|
||||
fn backup_worker_skips_already_recorded_slot() {
|
||||
let schedule = BackupSchedule {
|
||||
unit: BackupScheduleUnit::Days,
|
||||
interval: 1,
|
||||
minute: 0,
|
||||
hour: 3,
|
||||
weekday: 1,
|
||||
month_day: 1,
|
||||
};
|
||||
let now = chrono::DateTime::parse_from_rfc3339("2026-05-24T03:00:30+08:00")
|
||||
.unwrap()
|
||||
.with_timezone(&chrono::Utc);
|
||||
let slot = schedule.due_slot(now).expect("slot should be due");
|
||||
let slot = "days:2026-05-23T19:00:00Z";
|
||||
|
||||
assert!(super::should_start_scheduled_backup(
|
||||
Some("days:2026-05-22T19:00:00Z"),
|
||||
&slot
|
||||
slot
|
||||
));
|
||||
assert!(!super::should_start_scheduled_backup(Some(&slot), &slot));
|
||||
assert!(!super::should_start_scheduled_backup(Some(slot), slot));
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
@@ -0,0 +1,124 @@
|
||||
//! Credential-safe compatibility probe for the Codex Responses WebSocket path.
|
||||
//!
|
||||
//! This binary preserves the established Codex CLI and environment contract.
|
||||
//! The common Responses WebSocket flow lives in `support/responses_ws_probe`;
|
||||
//! this profile owns only Codex authentication and header requirements.
|
||||
|
||||
#[path = "support/responses_ws_probe.rs"]
|
||||
mod responses_ws_probe;
|
||||
|
||||
use aether_gateway::{CODEX_CLIENT_ORIGINATOR, CODEX_CLIENT_USER_AGENT};
|
||||
use clap::Parser;
|
||||
use http::header::{AUTHORIZATION, USER_AGENT};
|
||||
use http::{HeaderMap, HeaderName, HeaderValue};
|
||||
use responses_ws_probe::{
|
||||
bearer_authorization_value, required_env, resolve_probe_url, run_profile_probe, turn_timeout,
|
||||
ProbeArgs, ProbeConfig, ProbeFailure, ResponsesWebSocketProbeProfile,
|
||||
};
|
||||
|
||||
const ACCESS_TOKEN_ENV: &str = "AETHER_CODEX_WS_PROBE_ACCESS_TOKEN";
|
||||
const ACCOUNT_ID_ENV: &str = "AETHER_CODEX_WS_PROBE_ACCOUNT_ID";
|
||||
const MODEL_ENV: &str = "AETHER_CODEX_WS_PROBE_MODEL";
|
||||
const URL_ENV: &str = "AETHER_CODEX_WS_PROBE_URL";
|
||||
|
||||
#[derive(Parser)]
|
||||
#[command(
|
||||
name = "aether-codex-ws-probe",
|
||||
about = "Verify a Codex Responses WebSocket endpoint without exposing credentials"
|
||||
)]
|
||||
struct Args {
|
||||
/// WebSocket endpoint. If omitted, AETHER_CODEX_WS_PROBE_URL is used.
|
||||
#[arg(long)]
|
||||
url: Option<String>,
|
||||
|
||||
/// Per-turn receive timeout in seconds.
|
||||
#[arg(long, default_value_t = 20, value_parser = clap::value_parser!(u64).range(1..=120))]
|
||||
timeout_secs: u64,
|
||||
}
|
||||
|
||||
impl From<Args> for ProbeArgs {
|
||||
fn from(args: Args) -> Self {
|
||||
Self {
|
||||
url: args.url,
|
||||
timeout_secs: args.timeout_secs,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct CodexResponsesProbeProfile;
|
||||
|
||||
impl ResponsesWebSocketProbeProfile for CodexResponsesProbeProfile {
|
||||
fn build_config(args: &ProbeArgs) -> Result<ProbeConfig, ProbeFailure> {
|
||||
let url = resolve_probe_url(args, URL_ENV, None)?;
|
||||
let access_token = required_env(ACCESS_TOKEN_ENV)?;
|
||||
let account_id = required_env(ACCOUNT_ID_ENV)?;
|
||||
let model = required_env(MODEL_ENV)?;
|
||||
Ok(ProbeConfig::new(
|
||||
url,
|
||||
model,
|
||||
turn_timeout(args),
|
||||
handshake_headers(&access_token, &account_id)?,
|
||||
Self::sent_header_names(),
|
||||
))
|
||||
}
|
||||
|
||||
fn sent_header_names() -> Vec<&'static str> {
|
||||
vec![
|
||||
"authorization",
|
||||
"chatgpt-account-id",
|
||||
"user-agent",
|
||||
"originator",
|
||||
]
|
||||
}
|
||||
}
|
||||
|
||||
fn handshake_headers(access_token: &str, account_id: &str) -> Result<HeaderMap, ProbeFailure> {
|
||||
let account_id =
|
||||
HeaderValue::from_str(account_id).map_err(|_| ProbeFailure::MissingConfiguration)?;
|
||||
let mut headers = HeaderMap::new();
|
||||
headers.insert(AUTHORIZATION, bearer_authorization_value(access_token)?);
|
||||
headers.insert(HeaderName::from_static("chatgpt-account-id"), account_id);
|
||||
headers.insert(
|
||||
USER_AGENT,
|
||||
HeaderValue::from_static(CODEX_CLIENT_USER_AGENT),
|
||||
);
|
||||
headers.insert(
|
||||
HeaderName::from_static("originator"),
|
||||
HeaderValue::from_static(CODEX_CLIENT_ORIGINATOR),
|
||||
);
|
||||
Ok(headers)
|
||||
}
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() {
|
||||
let exit_code = run_profile_probe::<CodexResponsesProbeProfile>(Args::parse().into()).await;
|
||||
if exit_code != 0 {
|
||||
std::process::exit(exit_code);
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use http::header::{AUTHORIZATION, USER_AGENT};
|
||||
|
||||
use super::{handshake_headers, CodexResponsesProbeProfile, ResponsesWebSocketProbeProfile};
|
||||
|
||||
#[test]
|
||||
fn codex_profile_keeps_its_required_handshake_headers() {
|
||||
let headers =
|
||||
handshake_headers("test-token", "test-account").expect("headers should build");
|
||||
assert!(headers.contains_key(AUTHORIZATION));
|
||||
assert!(headers.contains_key("chatgpt-account-id"));
|
||||
assert!(headers.contains_key(USER_AGENT));
|
||||
assert!(headers.contains_key("originator"));
|
||||
assert_eq!(
|
||||
CodexResponsesProbeProfile::sent_header_names(),
|
||||
vec![
|
||||
"authorization",
|
||||
"chatgpt-account-id",
|
||||
"user-agent",
|
||||
"originator",
|
||||
]
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,104 @@
|
||||
//! Credential-safe compatibility probe for the official OpenAI Responses
|
||||
//! WebSocket endpoint.
|
||||
//!
|
||||
//! This profile uses standard API-key Bearer authentication and shares the
|
||||
//! protocol flow with the Codex probe without inheriting Codex-specific
|
||||
//! account headers or quota assumptions.
|
||||
|
||||
#[path = "support/responses_ws_probe.rs"]
|
||||
mod responses_ws_probe;
|
||||
|
||||
use clap::Parser;
|
||||
use http::header::AUTHORIZATION;
|
||||
use http::HeaderMap;
|
||||
use responses_ws_probe::{
|
||||
bearer_authorization_value, required_env, resolve_probe_url, run_profile_probe, turn_timeout,
|
||||
ProbeArgs, ProbeConfig, ProbeFailure, ResponsesWebSocketProbeProfile,
|
||||
};
|
||||
|
||||
const API_KEY_ENV: &str = "AETHER_OPENAI_WS_PROBE_API_KEY";
|
||||
const MODEL_ENV: &str = "AETHER_OPENAI_WS_PROBE_MODEL";
|
||||
const URL_ENV: &str = "AETHER_OPENAI_WS_PROBE_URL";
|
||||
const DEFAULT_URL: &str = "wss://api.openai.com/v1/responses";
|
||||
|
||||
#[derive(Parser)]
|
||||
#[command(
|
||||
name = "aether-openai-responses-ws-probe",
|
||||
about = "Verify an OpenAI Responses WebSocket endpoint without exposing credentials"
|
||||
)]
|
||||
struct Args {
|
||||
/// WebSocket endpoint. If omitted, AETHER_OPENAI_WS_PROBE_URL or the
|
||||
/// official OpenAI endpoint is used.
|
||||
#[arg(long)]
|
||||
url: Option<String>,
|
||||
|
||||
/// Per-turn receive timeout in seconds.
|
||||
#[arg(long, default_value_t = 20, value_parser = clap::value_parser!(u64).range(1..=120))]
|
||||
timeout_secs: u64,
|
||||
}
|
||||
|
||||
impl From<Args> for ProbeArgs {
|
||||
fn from(args: Args) -> Self {
|
||||
Self {
|
||||
url: args.url,
|
||||
timeout_secs: args.timeout_secs,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct OpenAiResponsesProbeProfile;
|
||||
|
||||
impl ResponsesWebSocketProbeProfile for OpenAiResponsesProbeProfile {
|
||||
fn build_config(args: &ProbeArgs) -> Result<ProbeConfig, ProbeFailure> {
|
||||
let url = resolve_probe_url(args, URL_ENV, Some(DEFAULT_URL))?;
|
||||
let api_key = required_env(API_KEY_ENV)?;
|
||||
let model = required_env(MODEL_ENV)?;
|
||||
let mut headers = HeaderMap::new();
|
||||
headers.insert(AUTHORIZATION, bearer_authorization_value(&api_key)?);
|
||||
Ok(ProbeConfig::new(
|
||||
url,
|
||||
model,
|
||||
turn_timeout(args),
|
||||
headers,
|
||||
Self::sent_header_names(),
|
||||
))
|
||||
}
|
||||
|
||||
fn sent_header_names() -> Vec<&'static str> {
|
||||
vec!["authorization"]
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() {
|
||||
let exit_code = run_profile_probe::<OpenAiResponsesProbeProfile>(Args::parse().into()).await;
|
||||
if exit_code != 0 {
|
||||
std::process::exit(exit_code);
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use http::header::AUTHORIZATION;
|
||||
|
||||
use super::{
|
||||
bearer_authorization_value, responses_ws_probe::parse_probe_url,
|
||||
OpenAiResponsesProbeProfile, ResponsesWebSocketProbeProfile, DEFAULT_URL,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn openai_profile_exposes_only_standard_bearer_authentication() {
|
||||
let authorization = bearer_authorization_value("test-key").expect("header should build");
|
||||
assert_eq!(authorization.to_str().ok(), Some("Bearer test-key"));
|
||||
assert_eq!(
|
||||
OpenAiResponsesProbeProfile::sent_header_names(),
|
||||
vec![AUTHORIZATION.as_str()]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn openai_profile_uses_the_official_responses_websocket_endpoint_by_default() {
|
||||
let url = parse_probe_url(DEFAULT_URL).expect("default OpenAI endpoint should be valid");
|
||||
assert_eq!(url.as_str(), DEFAULT_URL);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,567 @@
|
||||
//! Shared, credential-safe engine for Responses WebSocket compatibility probes.
|
||||
//!
|
||||
//! Provider profiles own their environment variables and handshake headers.
|
||||
//! This module owns the common Responses WebSocket contract: two sequential
|
||||
//! `response.create` warmups, continuation with `previous_response_id`, safe
|
||||
//! event observation, and a redacted JSON report.
|
||||
|
||||
use std::env;
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
use http::{HeaderMap, HeaderValue};
|
||||
use serde::Serialize;
|
||||
use serde_json::{json, Value};
|
||||
use url::Url;
|
||||
use wreq::ws::message::Message as WreqWsMessage;
|
||||
|
||||
const MAX_FRAME_SIZE: usize = 1 << 20;
|
||||
const MAX_EVENTS_PER_TURN: usize = 16;
|
||||
|
||||
pub(crate) struct ProbeArgs {
|
||||
pub(crate) url: Option<String>,
|
||||
pub(crate) timeout_secs: u64,
|
||||
}
|
||||
|
||||
pub(crate) struct ProbeConfig {
|
||||
url: Url,
|
||||
model: String,
|
||||
turn_timeout: Duration,
|
||||
handshake_headers: HeaderMap,
|
||||
sent_header_names: Vec<&'static str>,
|
||||
}
|
||||
|
||||
impl ProbeConfig {
|
||||
pub(crate) fn new(
|
||||
url: Url,
|
||||
model: String,
|
||||
turn_timeout: Duration,
|
||||
handshake_headers: HeaderMap,
|
||||
sent_header_names: Vec<&'static str>,
|
||||
) -> Self {
|
||||
Self {
|
||||
url,
|
||||
model,
|
||||
turn_timeout,
|
||||
handshake_headers,
|
||||
sent_header_names,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// A profile retains provider-specific authentication and configuration while
|
||||
/// reusing one Responses protocol probe engine.
|
||||
pub(crate) trait ResponsesWebSocketProbeProfile {
|
||||
fn build_config(args: &ProbeArgs) -> Result<ProbeConfig, ProbeFailure>;
|
||||
fn sent_header_names() -> Vec<&'static str>;
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
pub(crate) enum ProbeFailure {
|
||||
MissingConfiguration,
|
||||
InvalidEndpoint,
|
||||
ClientBuild,
|
||||
Handshake,
|
||||
Upgrade,
|
||||
Send,
|
||||
ReceiveTimeout,
|
||||
Receive,
|
||||
RemoteError,
|
||||
MissingResponseId,
|
||||
UnexpectedFrame,
|
||||
}
|
||||
|
||||
impl ProbeFailure {
|
||||
const fn code(self) -> &'static str {
|
||||
match self {
|
||||
Self::MissingConfiguration => "missing_configuration",
|
||||
Self::InvalidEndpoint => "invalid_endpoint",
|
||||
Self::ClientBuild => "client_build_failed",
|
||||
Self::Handshake => "handshake_failed",
|
||||
Self::Upgrade => "upgrade_failed",
|
||||
Self::Send => "send_failed",
|
||||
Self::ReceiveTimeout => "receive_timeout",
|
||||
Self::Receive => "receive_failed",
|
||||
Self::RemoteError => "upstream_error_event",
|
||||
Self::MissingResponseId => "response_id_not_observed",
|
||||
Self::UnexpectedFrame => "unexpected_frame",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Serialize)]
|
||||
struct ProbeReport {
|
||||
status: &'static str,
|
||||
target_host: Option<String>,
|
||||
handshake_status: Option<u16>,
|
||||
sent_header_names: Vec<&'static str>,
|
||||
received_header_names: Vec<String>,
|
||||
observed_event_types: Vec<String>,
|
||||
continuation_confirmed: bool,
|
||||
elapsed_ms: u64,
|
||||
error: Option<&'static str>,
|
||||
}
|
||||
|
||||
impl ProbeReport {
|
||||
fn failed(
|
||||
config: Option<&ProbeConfig>,
|
||||
sent_header_names: Vec<&'static str>,
|
||||
started_at: Instant,
|
||||
error: ProbeFailure,
|
||||
) -> Self {
|
||||
Self {
|
||||
status: "failed",
|
||||
target_host: config.and_then(target_host),
|
||||
handshake_status: None,
|
||||
sent_header_names,
|
||||
received_header_names: Vec::new(),
|
||||
observed_event_types: Vec::new(),
|
||||
continuation_confirmed: false,
|
||||
elapsed_ms: started_at.elapsed().as_millis() as u64,
|
||||
error: Some(error.code()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Runs a profile and returns the process exit code after emitting exactly one
|
||||
/// credential-safe JSON report.
|
||||
pub(crate) async fn run_profile_probe<P: ResponsesWebSocketProbeProfile>(args: ProbeArgs) -> i32 {
|
||||
let started_at = Instant::now();
|
||||
let config = match P::build_config(&args) {
|
||||
Ok(config) => config,
|
||||
Err(error) => {
|
||||
print_report(&ProbeReport::failed(
|
||||
None,
|
||||
P::sent_header_names(),
|
||||
started_at,
|
||||
error,
|
||||
));
|
||||
return 2;
|
||||
}
|
||||
};
|
||||
|
||||
match run_probe(&config, started_at).await {
|
||||
Ok(report) => {
|
||||
print_report(&report);
|
||||
0
|
||||
}
|
||||
Err(error) => {
|
||||
print_report(&ProbeReport::failed(
|
||||
Some(&config),
|
||||
config.sent_header_names.clone(),
|
||||
started_at,
|
||||
error,
|
||||
));
|
||||
1
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn required_env(name: &str) -> Result<String, ProbeFailure> {
|
||||
env::var(name)
|
||||
.ok()
|
||||
.map(|value| value.trim().to_string())
|
||||
.filter(|value| !value.is_empty())
|
||||
.ok_or(ProbeFailure::MissingConfiguration)
|
||||
}
|
||||
|
||||
pub(crate) fn resolve_probe_url(
|
||||
args: &ProbeArgs,
|
||||
url_env: &str,
|
||||
default_url: Option<&str>,
|
||||
) -> Result<Url, ProbeFailure> {
|
||||
let raw_url = args
|
||||
.url
|
||||
.as_deref()
|
||||
.map(str::to_owned)
|
||||
.or_else(|| env::var(url_env).ok())
|
||||
.or_else(|| default_url.map(str::to_owned));
|
||||
let Some(raw_url) = raw_url
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
else {
|
||||
return Err(ProbeFailure::MissingConfiguration);
|
||||
};
|
||||
parse_probe_url(raw_url)
|
||||
}
|
||||
|
||||
pub(crate) fn parse_probe_url(raw: &str) -> Result<Url, ProbeFailure> {
|
||||
let url = Url::parse(raw).map_err(|_| ProbeFailure::InvalidEndpoint)?;
|
||||
if !matches!(url.scheme(), "ws" | "wss")
|
||||
|| url.host_str().is_none()
|
||||
|| !url.username().is_empty()
|
||||
|| url.password().is_some()
|
||||
|| url.query().is_some()
|
||||
|| url.fragment().is_some()
|
||||
{
|
||||
return Err(ProbeFailure::InvalidEndpoint);
|
||||
}
|
||||
Ok(url)
|
||||
}
|
||||
|
||||
pub(crate) fn bearer_authorization_value(token: &str) -> Result<HeaderValue, ProbeFailure> {
|
||||
HeaderValue::from_str(format!("Bearer {token}").as_str())
|
||||
.map_err(|_| ProbeFailure::MissingConfiguration)
|
||||
}
|
||||
|
||||
pub(crate) const fn turn_timeout(args: &ProbeArgs) -> Duration {
|
||||
Duration::from_secs(args.timeout_secs)
|
||||
}
|
||||
|
||||
async fn run_probe(config: &ProbeConfig, started_at: Instant) -> Result<ProbeReport, ProbeFailure> {
|
||||
let client = wreq::Client::builder()
|
||||
.connect_timeout(config.turn_timeout)
|
||||
.timeout(config.turn_timeout)
|
||||
.build()
|
||||
.map_err(|_| ProbeFailure::ClientBuild)?;
|
||||
let response = client
|
||||
.websocket(config.url.as_str())
|
||||
.headers(config.handshake_headers.clone())
|
||||
.max_frame_size(MAX_FRAME_SIZE)
|
||||
.max_message_size(MAX_FRAME_SIZE)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|_| ProbeFailure::Handshake)?;
|
||||
let handshake_status = response.status().as_u16();
|
||||
let received_header_names = response
|
||||
.headers()
|
||||
.keys()
|
||||
.map(|name| name.as_str().to_string())
|
||||
.collect();
|
||||
let mut socket = response
|
||||
.into_websocket()
|
||||
.await
|
||||
.map_err(|_| ProbeFailure::Upgrade)?;
|
||||
let mut observed_event_types = Vec::new();
|
||||
|
||||
send_warmup(&mut socket, &config.model, None).await?;
|
||||
let first_response_id =
|
||||
receive_completed_response_id(&mut socket, config.turn_timeout, &mut observed_event_types)
|
||||
.await?;
|
||||
|
||||
send_warmup(&mut socket, &config.model, Some(&first_response_id)).await?;
|
||||
let _second_response_id =
|
||||
receive_completed_response_id(&mut socket, config.turn_timeout, &mut observed_event_types)
|
||||
.await?;
|
||||
|
||||
Ok(ProbeReport {
|
||||
status: "passed",
|
||||
target_host: target_host(config),
|
||||
handshake_status: Some(handshake_status),
|
||||
sent_header_names: config.sent_header_names.clone(),
|
||||
received_header_names,
|
||||
observed_event_types,
|
||||
continuation_confirmed: true,
|
||||
elapsed_ms: started_at.elapsed().as_millis() as u64,
|
||||
error: None,
|
||||
})
|
||||
}
|
||||
|
||||
fn target_host(config: &ProbeConfig) -> Option<String> {
|
||||
config.url.host_str().map(|host| match config.url.port() {
|
||||
Some(port) => format!("{host}:{port}"),
|
||||
None => host.to_string(),
|
||||
})
|
||||
}
|
||||
|
||||
async fn send_warmup(
|
||||
socket: &mut wreq::ws::WebSocket,
|
||||
model: &str,
|
||||
previous_response_id: Option<&str>,
|
||||
) -> Result<(), ProbeFailure> {
|
||||
let mut event = json!({
|
||||
"type": "response.create",
|
||||
"model": model,
|
||||
"store": false,
|
||||
"generate": false,
|
||||
"input": [],
|
||||
"tools": [],
|
||||
});
|
||||
if let Some(previous_response_id) = previous_response_id {
|
||||
event["previous_response_id"] = Value::String(previous_response_id.to_string());
|
||||
}
|
||||
socket
|
||||
.send(WreqWsMessage::text(event.to_string()))
|
||||
.await
|
||||
.map_err(|_| ProbeFailure::Send)
|
||||
}
|
||||
|
||||
async fn receive_completed_response_id(
|
||||
socket: &mut wreq::ws::WebSocket,
|
||||
timeout: Duration,
|
||||
observed_event_types: &mut Vec<String>,
|
||||
) -> Result<String, ProbeFailure> {
|
||||
let mut response_id = None;
|
||||
for _ in 0..MAX_EVENTS_PER_TURN {
|
||||
let message = tokio::time::timeout(timeout, socket.recv())
|
||||
.await
|
||||
.map_err(|_| ProbeFailure::ReceiveTimeout)?
|
||||
.ok_or(ProbeFailure::MissingResponseId)?
|
||||
.map_err(|_| ProbeFailure::Receive)?;
|
||||
match message {
|
||||
WreqWsMessage::Text(text) => {
|
||||
let event: Value = serde_json::from_str(text.as_str())
|
||||
.map_err(|_| ProbeFailure::UnexpectedFrame)?;
|
||||
let event_type = event
|
||||
.get("type")
|
||||
.and_then(Value::as_str)
|
||||
.map(safe_event_label)
|
||||
.unwrap_or_else(|| "unknown".to_string());
|
||||
let is_remote_error = event_type == "error";
|
||||
let is_completed = event_type == "response.completed";
|
||||
observed_event_types.push(event_type);
|
||||
if is_remote_error {
|
||||
return Err(ProbeFailure::RemoteError);
|
||||
}
|
||||
if let Some(observed_response_id) = event
|
||||
.pointer("/response/id")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
response_id = Some(observed_response_id.to_string());
|
||||
}
|
||||
if is_completed {
|
||||
return response_id.ok_or(ProbeFailure::MissingResponseId);
|
||||
}
|
||||
}
|
||||
WreqWsMessage::Ping(_) | WreqWsMessage::Pong(_) => continue,
|
||||
WreqWsMessage::Close(_) => return Err(ProbeFailure::MissingResponseId),
|
||||
_ => return Err(ProbeFailure::UnexpectedFrame),
|
||||
}
|
||||
}
|
||||
Err(ProbeFailure::MissingResponseId)
|
||||
}
|
||||
|
||||
fn safe_event_label(value: &str) -> String {
|
||||
let trimmed = value.trim();
|
||||
if trimmed.is_empty()
|
||||
|| trimmed.len() > 80
|
||||
|| !trimmed
|
||||
.bytes()
|
||||
.all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'.' | b'_' | b'-'))
|
||||
{
|
||||
return "unknown".to_string();
|
||||
}
|
||||
trimmed.to_string()
|
||||
}
|
||||
|
||||
fn print_report(report: &ProbeReport) {
|
||||
match serde_json::to_string(report) {
|
||||
Ok(json) => println!("{json}"),
|
||||
Err(_) => println!("{{\"status\":\"failed\",\"error\":\"report_serialization_failed\"}}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::Arc;
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade};
|
||||
use axum::extract::State;
|
||||
use axum::http::header::AUTHORIZATION;
|
||||
use axum::http::{HeaderMap, HeaderValue};
|
||||
use axum::response::IntoResponse;
|
||||
use axum::routing::get;
|
||||
use axum::Router;
|
||||
use futures_util::{SinkExt, StreamExt};
|
||||
use serde_json::Value;
|
||||
use tokio::sync::{oneshot, Mutex};
|
||||
|
||||
use super::{parse_probe_url, run_probe, ProbeConfig};
|
||||
|
||||
#[derive(Default)]
|
||||
struct MockState {
|
||||
observed: Mutex<Option<oneshot::Sender<ObservedClientMessages>>>,
|
||||
}
|
||||
|
||||
struct ObservedClientMessages {
|
||||
authorization_present: bool,
|
||||
profile_header_present: bool,
|
||||
second_before_first_completion: bool,
|
||||
first: Value,
|
||||
second: Value,
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn probe_confirms_sequential_response_continuation_without_exposing_values() {
|
||||
let (url, observed, server) = spawn_mock_server().await;
|
||||
let mut headers = HeaderMap::new();
|
||||
headers.insert(
|
||||
AUTHORIZATION,
|
||||
HeaderValue::from_static("Bearer test-token-that-must-not-be-reported"),
|
||||
);
|
||||
headers.insert(
|
||||
"x-aether-probe-profile",
|
||||
HeaderValue::from_static("test-profile-id"),
|
||||
);
|
||||
let config = ProbeConfig::new(
|
||||
parse_probe_url(url.as_str()).expect("mock URL should be valid"),
|
||||
"gpt-test".to_string(),
|
||||
Duration::from_secs(2),
|
||||
headers,
|
||||
vec!["authorization", "x-aether-probe-profile"],
|
||||
);
|
||||
|
||||
let report = run_probe(&config, Instant::now())
|
||||
.await
|
||||
.expect("probe should complete against mock server");
|
||||
let client_messages = observed.await.expect("mock should observe client messages");
|
||||
server.abort();
|
||||
|
||||
assert_eq!(report.status, "passed");
|
||||
assert!(report.continuation_confirmed);
|
||||
assert!(report
|
||||
.observed_event_types
|
||||
.contains(&"response.created".to_string()));
|
||||
assert!(report
|
||||
.observed_event_types
|
||||
.contains(&"response.completed".to_string()));
|
||||
assert!(client_messages.authorization_present);
|
||||
assert!(client_messages.profile_header_present);
|
||||
assert!(!client_messages.second_before_first_completion);
|
||||
assert_eq!(client_messages.first["type"], "response.create");
|
||||
assert_eq!(client_messages.first["generate"], false);
|
||||
assert_eq!(client_messages.first["store"], false);
|
||||
assert_eq!(client_messages.second["previous_response_id"], "resp-first");
|
||||
let report_json = serde_json::to_string(&report).expect("report should serialize");
|
||||
assert!(!report_json.contains("test-token-that-must-not-be-reported"));
|
||||
assert!(!report_json.contains("test-profile-id"));
|
||||
assert!(!report_json.contains("resp-first"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn probe_url_rejects_credentials_and_query_strings() {
|
||||
assert!(parse_probe_url("wss://example.test/v1/responses").is_ok());
|
||||
assert!(parse_probe_url("https://example.test/v1/responses").is_err());
|
||||
assert!(parse_probe_url("wss://[email protected]/v1/responses").is_err());
|
||||
assert!(parse_probe_url("wss://example.test/v1/responses?token=secret").is_err());
|
||||
}
|
||||
|
||||
async fn spawn_mock_server() -> (
|
||||
String,
|
||||
oneshot::Receiver<ObservedClientMessages>,
|
||||
tokio::task::JoinHandle<()>,
|
||||
) {
|
||||
let (observed_tx, observed_rx) = oneshot::channel();
|
||||
let state = Arc::new(MockState {
|
||||
observed: Mutex::new(Some(observed_tx)),
|
||||
});
|
||||
let app = Router::new()
|
||||
.route("/v1/responses", get(mock_websocket))
|
||||
.with_state(state);
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
|
||||
.await
|
||||
.expect("mock listener should bind");
|
||||
let address = listener
|
||||
.local_addr()
|
||||
.expect("mock listener should expose address");
|
||||
let server = tokio::spawn(async move {
|
||||
axum::serve(listener, app)
|
||||
.await
|
||||
.expect("mock server should run");
|
||||
});
|
||||
(format!("ws://{address}/v1/responses"), observed_rx, server)
|
||||
}
|
||||
|
||||
async fn mock_websocket(
|
||||
ws: WebSocketUpgrade,
|
||||
State(state): State<Arc<MockState>>,
|
||||
headers: HeaderMap,
|
||||
) -> impl IntoResponse {
|
||||
let authorization_present = headers
|
||||
.get(AUTHORIZATION)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.is_some_and(|value| value.starts_with("Bearer "));
|
||||
let profile_header_present = headers.contains_key("x-aether-probe-profile");
|
||||
ws.on_upgrade(move |socket| async move {
|
||||
serve_mock_socket(socket, state, authorization_present, profile_header_present).await;
|
||||
})
|
||||
}
|
||||
|
||||
async fn serve_mock_socket(
|
||||
socket: WebSocket,
|
||||
state: Arc<MockState>,
|
||||
authorization_present: bool,
|
||||
profile_header_present: bool,
|
||||
) {
|
||||
let (mut sender, mut receiver) = socket.split();
|
||||
let first = receive_json(&mut receiver).await;
|
||||
let _ = sender
|
||||
.send(Message::Text(
|
||||
serde_json::json!({
|
||||
"type": "response.created",
|
||||
"response": {"id": "resp-first"}
|
||||
})
|
||||
.to_string()
|
||||
.into(),
|
||||
))
|
||||
.await;
|
||||
let early_second = tokio::select! {
|
||||
message = receiver.next() => Some(message),
|
||||
_ = tokio::time::sleep(Duration::from_millis(50)) => None,
|
||||
};
|
||||
let second_before_first_completion = early_second.is_some();
|
||||
let _ = sender
|
||||
.send(Message::Text(
|
||||
serde_json::json!({
|
||||
"type": "response.completed",
|
||||
"response": {"id": "resp-first", "status": "completed"}
|
||||
})
|
||||
.to_string()
|
||||
.into(),
|
||||
))
|
||||
.await;
|
||||
let second = match early_second {
|
||||
Some(Some(Ok(Message::Text(text)))) => {
|
||||
serde_json::from_str(text.as_str()).expect("early client message should be JSON")
|
||||
}
|
||||
Some(Some(Ok(_))) => panic!("expected text continuation message"),
|
||||
Some(Some(Err(error))) => panic!("client message should be valid: {error}"),
|
||||
Some(None) => panic!("client closed before continuation"),
|
||||
None => receive_json(&mut receiver).await,
|
||||
};
|
||||
let _ = sender
|
||||
.send(Message::Text(
|
||||
serde_json::json!({
|
||||
"type": "response.created",
|
||||
"response": {"id": "resp-second"}
|
||||
})
|
||||
.to_string()
|
||||
.into(),
|
||||
))
|
||||
.await;
|
||||
let _ = sender
|
||||
.send(Message::Text(
|
||||
serde_json::json!({
|
||||
"type": "response.completed",
|
||||
"response": {"id": "resp-second", "status": "completed"}
|
||||
})
|
||||
.to_string()
|
||||
.into(),
|
||||
))
|
||||
.await;
|
||||
if let Some(observed) = state.observed.lock().await.take() {
|
||||
let _ = observed.send(ObservedClientMessages {
|
||||
authorization_present,
|
||||
profile_header_present,
|
||||
second_before_first_completion,
|
||||
first,
|
||||
second,
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
async fn receive_json(receiver: &mut futures_util::stream::SplitStream<WebSocket>) -> Value {
|
||||
let message = receiver
|
||||
.next()
|
||||
.await
|
||||
.expect("client should send a message")
|
||||
.expect("client message should be valid");
|
||||
let Message::Text(text) = message else {
|
||||
panic!("expected text message");
|
||||
};
|
||||
serde_json::from_str(text.as_str()).expect("client message should be JSON")
|
||||
}
|
||||
}
|
||||
-10
@@ -1274,7 +1274,6 @@ mod tests {
|
||||
let first_cache = Arc::clone(&cache);
|
||||
let first_key = key.clone();
|
||||
let first_calls = Arc::clone(&calls);
|
||||
let first_started = Instant::now();
|
||||
let first = tokio::spawn(async move {
|
||||
first_cache
|
||||
.get_or_load_once_stale_while_refreshing::<(), _, _>(
|
||||
@@ -1292,7 +1291,6 @@ mod tests {
|
||||
});
|
||||
|
||||
let follower_cache = Arc::clone(&cache);
|
||||
let follower_started = Instant::now();
|
||||
let follower_calls = Arc::clone(&calls);
|
||||
let follower = tokio::spawn(async move {
|
||||
follower_cache
|
||||
@@ -1310,15 +1308,7 @@ mod tests {
|
||||
});
|
||||
|
||||
assert_eq!(first.await.unwrap().unwrap(), Some(1));
|
||||
assert!(
|
||||
first_started.elapsed() < Duration::from_millis(80),
|
||||
"stale value should not wait for request-path refresh"
|
||||
);
|
||||
assert_eq!(follower.await.unwrap().unwrap(), Some(1));
|
||||
assert!(
|
||||
follower_started.elapsed() < Duration::from_millis(80),
|
||||
"follower should return stale value without waiting for refresh"
|
||||
);
|
||||
assert_eq!(calls.load(Ordering::Acquire), 0);
|
||||
}
|
||||
|
||||
|
||||
@@ -177,6 +177,19 @@ fn extract_trusted_auth_headers(headers: &http::HeaderMap) -> Option<GatewayTrus
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(not(test))]
|
||||
pub(super) fn extract_trusted_admin_headers(
|
||||
_headers: &http::HeaderMap,
|
||||
) -> Option<GatewayTrustedAdminHeaders> {
|
||||
// The public gateway has no authenticated upstream that is allowed to
|
||||
// assert an administrator principal. `x-aether-gateway` is also emitted
|
||||
// on public responses, so it cannot serve as proof that these headers were
|
||||
// produced by a trusted hop. Production requests must authenticate with a
|
||||
// real admin session or management bearer token instead.
|
||||
None
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub(super) fn extract_trusted_admin_headers(
|
||||
headers: &http::HeaderMap,
|
||||
) -> Option<GatewayTrustedAdminHeaders> {
|
||||
|
||||
@@ -11,8 +11,9 @@ pub(crate) use gate::{
|
||||
should_buffer_request_for_local_auth, trusted_auth_local_rejection, GatewayLocalAuthRejection,
|
||||
};
|
||||
pub(crate) use resolution::{
|
||||
refresh_execution_runtime_auth_context, resolve_execution_runtime_auth_context,
|
||||
GatewayAdminPrincipalContext, GatewayControlAuthContext,
|
||||
refresh_execution_runtime_auth_context, refresh_execution_runtime_auth_context_with_snapshot,
|
||||
resolve_execution_runtime_auth_context, GatewayAdminPrincipalContext,
|
||||
GatewayControlAuthContext,
|
||||
};
|
||||
pub(super) use resolution::{resolve_control_decision_auth, ControlDecisionAuthResolution};
|
||||
pub(crate) use types::GatewayCredentialCarrier;
|
||||
|
||||
@@ -725,20 +725,47 @@ pub(crate) async fn refresh_execution_runtime_auth_context(
|
||||
auth_context: GatewayControlAuthContext,
|
||||
auth_endpoint_signature: Option<&str>,
|
||||
) -> Result<GatewayControlAuthContext, GatewayError> {
|
||||
refresh_execution_runtime_auth_context_with_snapshot(
|
||||
state,
|
||||
auth_context,
|
||||
auth_endpoint_signature,
|
||||
)
|
||||
.await
|
||||
.map(|(auth_context, _)| auth_context)
|
||||
}
|
||||
|
||||
/// Strongly refreshes the long-lived execution authorization context and
|
||||
/// returns the exact API-key snapshot that produced it.
|
||||
///
|
||||
/// WebSocket turns need both values: using the refreshed context for RPM and
|
||||
/// balance checks while letting the planner independently read its normal
|
||||
/// cache can authorize a different provider/model snapshot for up to the cache
|
||||
/// TTL. Ordinary HTTP callers keep using [`refresh_execution_runtime_auth_context`].
|
||||
pub(crate) async fn refresh_execution_runtime_auth_context_with_snapshot(
|
||||
state: &AppState,
|
||||
auth_context: GatewayControlAuthContext,
|
||||
auth_endpoint_signature: Option<&str>,
|
||||
) -> Result<
|
||||
(
|
||||
GatewayControlAuthContext,
|
||||
Option<crate::ai_serving::GatewayAuthApiKeySnapshot>,
|
||||
),
|
||||
GatewayError,
|
||||
> {
|
||||
if auth_context.local_rejection.is_some() || !auth_context.access_allowed {
|
||||
return Ok(auth_context);
|
||||
return Ok((auth_context, None));
|
||||
}
|
||||
let Some(auth_endpoint_signature) = auth_endpoint_signature
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
else {
|
||||
return Ok(auth_context);
|
||||
return Ok((auth_context, None));
|
||||
};
|
||||
if !state.has_auth_api_key_reader()
|
||||
|| auth_context.user_id.trim().is_empty()
|
||||
|| auth_context.api_key_id.trim().is_empty()
|
||||
{
|
||||
return Ok(auth_context);
|
||||
return Ok((auth_context, None));
|
||||
}
|
||||
|
||||
let snapshot = {
|
||||
@@ -758,19 +785,20 @@ pub(crate) async fn refresh_execution_runtime_auth_context(
|
||||
denied.access_allowed = false;
|
||||
denied.local_rejection = Some(GatewayLocalAuthRejection::InvalidApiKey);
|
||||
denied.balance_remaining = None;
|
||||
return Ok(denied);
|
||||
return Ok((denied, None));
|
||||
};
|
||||
|
||||
let wallet_access = resolve_wallet_auth_gate_uncached(state, &snapshot).await?;
|
||||
Ok(build_data_backed_auth_context(
|
||||
let refreshed = build_data_backed_auth_context(
|
||||
state,
|
||||
snapshot,
|
||||
snapshot.clone(),
|
||||
auth_endpoint_signature,
|
||||
Some(true),
|
||||
auth_context.balance_remaining,
|
||||
wallet_access,
|
||||
)
|
||||
.await)
|
||||
.await;
|
||||
Ok((refreshed, Some(snapshot)))
|
||||
}
|
||||
|
||||
fn put_cached_auth_context(
|
||||
|
||||
@@ -9,10 +9,11 @@ mod route;
|
||||
|
||||
pub(crate) use auth::{
|
||||
execution_plan_balance_capacity_rejection, extract_requested_model,
|
||||
refresh_execution_runtime_auth_context, request_model_local_rejection,
|
||||
resolve_execution_runtime_auth_context, should_buffer_request_for_local_auth,
|
||||
trusted_auth_local_rejection, GatewayAdminPrincipalContext, GatewayControlAuthContext,
|
||||
GatewayCredentialCarrier, GatewayLocalAuthRejection,
|
||||
refresh_execution_runtime_auth_context, refresh_execution_runtime_auth_context_with_snapshot,
|
||||
request_model_local_rejection, resolve_execution_runtime_auth_context,
|
||||
should_buffer_request_for_local_auth, trusted_auth_local_rejection,
|
||||
GatewayAdminPrincipalContext, GatewayControlAuthContext, GatewayCredentialCarrier,
|
||||
GatewayLocalAuthRejection,
|
||||
};
|
||||
pub(crate) use execute::{allows_control_execute_emergency, maybe_execute_via_control};
|
||||
pub(crate) use management_token_permissions::{
|
||||
|
||||
@@ -35,7 +35,10 @@ pub(super) fn classify_ai_public_route(
|
||||
"openai:rerank",
|
||||
true,
|
||||
))
|
||||
} else if method == http::Method::POST
|
||||
} else if (method == http::Method::POST
|
||||
|| (method == http::Method::GET
|
||||
&& normalized_path == "/v1/responses"
|
||||
&& is_websocket_upgrade_request(headers)))
|
||||
&& matches!(normalized_path, "/v1/responses" | "/v1/responses/compact")
|
||||
{
|
||||
if normalized_path.ends_with("/compact") {
|
||||
@@ -199,6 +202,24 @@ fn claude_request_auth_channel(headers: &http::HeaderMap) -> &'static str {
|
||||
}
|
||||
}
|
||||
|
||||
fn is_websocket_upgrade_request(headers: &http::HeaderMap) -> bool {
|
||||
let has_upgrade_connection = headers
|
||||
.get(http::header::CONNECTION)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.is_some_and(|value| {
|
||||
value
|
||||
.split(',')
|
||||
.map(str::trim)
|
||||
.any(|value| value.eq_ignore_ascii_case("upgrade"))
|
||||
});
|
||||
let has_websocket_upgrade = headers
|
||||
.get(http::header::UPGRADE)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.is_some_and(|value| value.eq_ignore_ascii_case("websocket"));
|
||||
|
||||
has_upgrade_connection && has_websocket_upgrade
|
||||
}
|
||||
|
||||
fn is_gemini_operation_method(method: &http::Method, normalized_path: &str) -> bool {
|
||||
method == http::Method::GET
|
||||
|| (method == http::Method::POST && normalized_path.ends_with(":cancel"))
|
||||
@@ -242,3 +263,32 @@ fn classify_antigravity_v1internal_route(
|
||||
execution_runtime_candidate,
|
||||
))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use axum::http::header::{CONNECTION, UPGRADE};
|
||||
use axum::http::{HeaderMap, HeaderValue, Method};
|
||||
|
||||
use super::classify_ai_public_route;
|
||||
|
||||
#[test]
|
||||
fn classifies_websocket_upgrade_on_responses_route() {
|
||||
let mut headers = HeaderMap::new();
|
||||
headers.insert(CONNECTION, HeaderValue::from_static("keep-alive, Upgrade"));
|
||||
headers.insert(UPGRADE, HeaderValue::from_static("websocket"));
|
||||
|
||||
let route = classify_ai_public_route(&Method::GET, "/v1/responses", &headers)
|
||||
.expect("Responses WebSocket should be an AI public route");
|
||||
assert_eq!(route.route_class, "ai_public");
|
||||
assert_eq!(route.route_family, "openai");
|
||||
assert_eq!(route.route_kind, "responses");
|
||||
assert_eq!(route.auth_endpoint_signature, "openai:responses");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn does_not_classify_plain_get_as_responses_websocket() {
|
||||
assert!(
|
||||
classify_ai_public_route(&Method::GET, "/v1/responses", &HeaderMap::new()).is_none()
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -247,8 +247,7 @@ pub(super) fn detect_public_models_auth_signature(uri: &Uri, headers: &http::Hea
|
||||
|
||||
let has_codex_client_version = uri.path() == "/v1/models"
|
||||
&& uri.query().is_some_and(|query| {
|
||||
url::form_urlencoded::parse(query.as_bytes())
|
||||
.any(|(key, value)| key == "client_version" && !value.trim().is_empty())
|
||||
url::form_urlencoded::parse(query.as_bytes()).any(|(key, _)| key == "client_version")
|
||||
});
|
||||
if has_codex_client_version {
|
||||
return "openai:responses".to_string();
|
||||
|
||||
@@ -39,7 +39,7 @@ fn classifies_codex_models_list_with_responses_auth_signature() {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn empty_codex_client_version_keeps_standard_openai_models_signature() {
|
||||
fn empty_codex_client_version_uses_responses_signature_for_bounded_fallback() {
|
||||
let headers = headers(&[("authorization", "Bearer sk-test")]);
|
||||
let uri: Uri = "/v1/models?client_version="
|
||||
.parse()
|
||||
@@ -49,7 +49,7 @@ fn empty_codex_client_version_keeps_standard_openai_models_signature() {
|
||||
|
||||
assert_eq!(
|
||||
decision.auth_endpoint_signature.as_deref(),
|
||||
Some("openai:chat")
|
||||
Some("openai:responses")
|
||||
);
|
||||
}
|
||||
|
||||
|
||||
@@ -195,6 +195,35 @@ mod tests {
|
||||
assert_eq!(background.pool.min_connections, 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn runtime_pool_split_gives_default_small_server_more_foreground_capacity() {
|
||||
let config = GatewayDataConfig::from_database_config(
|
||||
SqlDatabaseConfig::new(
|
||||
DatabaseDriver::Postgres,
|
||||
"postgres://localhost/aether",
|
||||
SqlPoolConfig {
|
||||
min_connections: 4,
|
||||
max_connections: 32,
|
||||
..SqlPoolConfig::default()
|
||||
},
|
||||
)
|
||||
.expect("database config should be valid"),
|
||||
);
|
||||
|
||||
let (foreground, background) = config.split_runtime_pools_with_background_max(None);
|
||||
let foreground = foreground.database().expect("foreground database");
|
||||
let background = background
|
||||
.expect("background database config")
|
||||
.database()
|
||||
.expect("background database")
|
||||
.clone();
|
||||
|
||||
assert_eq!(foreground.pool.max_connections, 26);
|
||||
assert_eq!(background.pool.max_connections, 6);
|
||||
assert_eq!(foreground.pool.min_connections, 4);
|
||||
assert_eq!(background.pool.min_connections, 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn runtime_pool_split_can_be_disabled_or_degrade_for_single_connection() {
|
||||
let mut database = SqlDatabaseConfig::sqlite_default();
|
||||
|
||||
@@ -1,14 +1,14 @@
|
||||
use super::{
|
||||
ApiKeyLastUsedDelta, DataLayerError, GatewayDataState, GeminiFileMappingListQuery,
|
||||
GeminiFileMappingStats, ProviderCatalogKeyAdaptiveStateUpdate,
|
||||
ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyListQuery,
|
||||
ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthRuntimeStateCasUpdate,
|
||||
ProviderCatalogKeyRuntimeMetadataUpdate, ProviderCatalogKeyStatusSnapshotUpdate,
|
||||
PublicHealthStatusCount, PublicHealthTimelineBucket, StoredGeminiFileMapping,
|
||||
StoredGeminiFileMappingListPage, StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
|
||||
StoredProviderCatalogKeyMaintenanceSummary, StoredProviderCatalogKeyPage,
|
||||
StoredProviderCatalogKeyStats, StoredProviderCatalogProvider, StoredRequestCandidate,
|
||||
UpsertGeminiFileMappingRecord, UpsertRequestCandidateRecord,
|
||||
ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyHealthStateUpdate,
|
||||
ProviderCatalogKeyListQuery, ProviderCatalogKeyOAuthCredentialCasDelete,
|
||||
ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ProviderCatalogKeyRuntimeMetadataUpdate,
|
||||
ProviderCatalogKeyStatusSnapshotUpdate, PublicHealthStatusCount, PublicHealthTimelineBucket,
|
||||
StoredGeminiFileMapping, StoredGeminiFileMappingListPage, StoredProviderCatalogEndpoint,
|
||||
StoredProviderCatalogKey, StoredProviderCatalogKeyMaintenanceSummary,
|
||||
StoredProviderCatalogKeyPage, StoredProviderCatalogKeyStats, StoredProviderCatalogProvider,
|
||||
StoredRequestCandidate, UpsertGeminiFileMappingRecord, UpsertRequestCandidateRecord,
|
||||
};
|
||||
|
||||
impl GatewayDataState {
|
||||
@@ -282,6 +282,16 @@ impl GatewayDataState {
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn list_provider_catalog_keys_by_ids_strong(
|
||||
&self,
|
||||
key_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogKey>, DataLayerError> {
|
||||
match &self.provider_catalog_reader {
|
||||
Some(repository) => repository.list_keys_by_ids_strong(key_ids).await,
|
||||
None => Ok(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn list_provider_catalog_keys_by_provider_ids(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
@@ -561,6 +571,20 @@ impl GatewayDataState {
|
||||
Ok(updated)
|
||||
}
|
||||
|
||||
pub(crate) async fn compare_and_update_provider_catalog_key_admin_state(
|
||||
&self,
|
||||
update: &ProviderCatalogKeyAdminCasUpdate,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
let updated = match &self.provider_catalog_writer {
|
||||
Some(repository) => repository.compare_and_update_key_admin_state(update).await,
|
||||
None => Ok(false),
|
||||
}?;
|
||||
// Clear on both success and conflict so a retry cannot reuse the stale
|
||||
// credential snapshot that lost the CAS.
|
||||
self.clear_provider_catalog_cache();
|
||||
Ok(updated)
|
||||
}
|
||||
|
||||
pub(crate) async fn update_provider_catalog_keys(
|
||||
&self,
|
||||
keys: &[StoredProviderCatalogKey],
|
||||
|
||||
@@ -121,13 +121,13 @@ use aether_data_contracts::repository::pool_scores::{
|
||||
UpsertPoolMemberScore,
|
||||
};
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
ProviderCatalogKeyAdaptiveStateUpdate, ProviderCatalogKeyHealthStateUpdate,
|
||||
ProviderCatalogKeyListQuery, ProviderCatalogKeyOAuthCredentialCasDelete,
|
||||
ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ProviderCatalogKeyRuntimeMetadataUpdate,
|
||||
ProviderCatalogKeyStatusSnapshotUpdate, ProviderCatalogReadRepository,
|
||||
ProviderCatalogWriteRepository, StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
|
||||
StoredProviderCatalogKeyMaintenanceSummary, StoredProviderCatalogKeyPage,
|
||||
StoredProviderCatalogKeyStats, StoredProviderCatalogProvider,
|
||||
ProviderCatalogKeyAdaptiveStateUpdate, ProviderCatalogKeyAdminCasUpdate,
|
||||
ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyListQuery,
|
||||
ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthRuntimeStateCasUpdate,
|
||||
ProviderCatalogKeyRuntimeMetadataUpdate, ProviderCatalogKeyStatusSnapshotUpdate,
|
||||
ProviderCatalogReadRepository, ProviderCatalogWriteRepository, StoredProviderCatalogEndpoint,
|
||||
StoredProviderCatalogKey, StoredProviderCatalogKeyMaintenanceSummary,
|
||||
StoredProviderCatalogKeyPage, StoredProviderCatalogKeyStats, StoredProviderCatalogProvider,
|
||||
};
|
||||
use aether_data_contracts::repository::quota::{
|
||||
ProviderQuotaReadRepository, ProviderQuotaWriteRepository, StoredProviderQuotaSnapshot,
|
||||
|
||||
@@ -334,6 +334,13 @@ impl ProviderCatalogReadRepository for CachedProviderCatalogReadRepository {
|
||||
}
|
||||
}
|
||||
|
||||
async fn list_keys_by_ids_strong(
|
||||
&self,
|
||||
key_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogKey>, DataLayerError> {
|
||||
self.inner.list_keys_by_ids_strong(key_ids).await
|
||||
}
|
||||
|
||||
async fn list_keys_by_provider_ids(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
@@ -538,6 +545,7 @@ fn normalize_ids(ids: &[String]) -> Vec<String> {
|
||||
mod tests {
|
||||
use super::*;
|
||||
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
|
||||
use aether_data_contracts::repository::provider_catalog::ProviderCatalogWriteRepository;
|
||||
|
||||
fn cache() -> CachedProviderCatalogReadRepository {
|
||||
CachedProviderCatalogReadRepository::new(Arc::new(
|
||||
@@ -555,6 +563,55 @@ mod tests {
|
||||
.expect("provider should be valid")
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn provider_catalog_strong_key_read_bypasses_fresh_cached_generation() {
|
||||
let old_metadata = serde_json::json!({
|
||||
"codex": {"credential_generation": "old"}
|
||||
});
|
||||
let mut key = StoredProviderCatalogKey::new(
|
||||
"key-1".to_string(),
|
||||
"provider-1".to_string(),
|
||||
"key-1".to_string(),
|
||||
"oauth".to_string(),
|
||||
None,
|
||||
true,
|
||||
)
|
||||
.expect("key should be valid");
|
||||
key.upstream_metadata = Some(old_metadata.clone());
|
||||
let inner = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![provider("provider-1")],
|
||||
Vec::new(),
|
||||
vec![key],
|
||||
));
|
||||
let cache = CachedProviderCatalogReadRepository::new(inner.clone());
|
||||
let key_ids = vec!["key-1".to_string()];
|
||||
|
||||
let first = cache
|
||||
.list_keys_by_ids(&key_ids)
|
||||
.await
|
||||
.expect("initial key read should succeed");
|
||||
assert_eq!(first[0].upstream_metadata.as_ref(), Some(&old_metadata));
|
||||
|
||||
let new_metadata = serde_json::json!({
|
||||
"codex": {"credential_generation": "new"}
|
||||
});
|
||||
assert!(inner
|
||||
.upsert_key_upstream_metadata_namespace("key-1", "codex", &new_metadata["codex"], None,)
|
||||
.await
|
||||
.expect("inner metadata update should succeed"));
|
||||
|
||||
let cached = cache
|
||||
.list_keys_by_ids(&key_ids)
|
||||
.await
|
||||
.expect("cached key read should succeed");
|
||||
assert_eq!(cached[0].upstream_metadata.as_ref(), Some(&old_metadata));
|
||||
let strong = cache
|
||||
.list_keys_by_ids_strong(&key_ids)
|
||||
.await
|
||||
.expect("strong key read should succeed");
|
||||
assert_eq!(strong[0].upstream_metadata.as_ref(), Some(&new_metadata));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn provider_catalog_follower_observes_completion_before_first_poll() {
|
||||
let cache = cache();
|
||||
|
||||
@@ -0,0 +1,62 @@
|
||||
//! Shared admission helpers for local upstream execution.
|
||||
//!
|
||||
//! The stream candidate loop and long-lived WebSocket turns both need to
|
||||
//! participate in the same gateway-wide upstream execution gate. Keep the
|
||||
//! provider abstraction here so tests can supply an isolated gate while
|
||||
//! production callers use `AppState` directly.
|
||||
|
||||
use std::time::Duration;
|
||||
|
||||
use aether_runtime::{ConcurrencyGate, ConcurrencyPermit};
|
||||
use tokio::time::timeout;
|
||||
|
||||
use crate::stage_metrics::observe_gateway_stage_ms;
|
||||
use crate::{AppState, GatewayError};
|
||||
|
||||
pub(crate) const UPSTREAM_EXECUTION_GATE_NAME: &str = "gateway_upstream_execution";
|
||||
|
||||
pub(crate) trait UpstreamExecutionGateProvider {
|
||||
fn upstream_execution_gate(&self) -> Option<&ConcurrencyGate>;
|
||||
fn upstream_execution_gate_queue_budget(&self) -> Duration;
|
||||
}
|
||||
|
||||
impl UpstreamExecutionGateProvider for AppState {
|
||||
fn upstream_execution_gate(&self) -> Option<&ConcurrencyGate> {
|
||||
self.upstream_execution_gate.as_deref()
|
||||
}
|
||||
|
||||
fn upstream_execution_gate_queue_budget(&self) -> Duration {
|
||||
self.frontdoor_runtime_guards.internal_gate_queue_budget
|
||||
}
|
||||
}
|
||||
|
||||
/// Acquires the shared gateway-wide upstream execution permit.
|
||||
///
|
||||
/// A missing gate is an intentional configuration (unlimited), so callers
|
||||
/// receive `Ok(None)`. Saturation keeps the existing candidate-level
|
||||
/// `AdmissionTimeout` contract used by the HTTP stream path.
|
||||
pub(crate) async fn acquire_upstream_execution_gate(
|
||||
state: &(impl UpstreamExecutionGateProvider + ?Sized),
|
||||
trace_id: &str,
|
||||
) -> Result<Option<ConcurrencyPermit>, GatewayError> {
|
||||
let Some(gate) = state.upstream_execution_gate() else {
|
||||
return Ok(None);
|
||||
};
|
||||
let budget = state.upstream_execution_gate_queue_budget();
|
||||
let gate_wait_started_at = std::time::Instant::now();
|
||||
match timeout(budget, gate.acquire()).await {
|
||||
Ok(Ok(permit)) => {
|
||||
observe_gateway_stage_ms(
|
||||
"upstream_execution_gate_wait",
|
||||
gate_wait_started_at.elapsed().as_millis() as u64,
|
||||
);
|
||||
Ok(Some(permit))
|
||||
}
|
||||
Ok(Err(err)) => Err(GatewayError::Internal(err.to_string())),
|
||||
Err(_) => Err(GatewayError::AdmissionTimeout {
|
||||
trace_id: trace_id.to_string(),
|
||||
gate: UPSTREAM_EXECUTION_GATE_NAME,
|
||||
queue_budget_ms: budget.as_millis() as u64,
|
||||
}),
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -2573,6 +2573,7 @@ fn json_execution_result(
|
||||
candidate_id: plan.candidate_id.clone(),
|
||||
status_code,
|
||||
headers: BTreeMap::from([("content-type".to_string(), "application/json".to_string())]),
|
||||
response_observation: None,
|
||||
body: Some(ResponseBody {
|
||||
json_body: Some(body),
|
||||
body_bytes_b64: None,
|
||||
@@ -2615,6 +2616,7 @@ fn bytes_execution_result(
|
||||
candidate_id: plan.candidate_id.clone(),
|
||||
status_code,
|
||||
headers,
|
||||
response_observation: None,
|
||||
body: Some(ResponseBody {
|
||||
json_body: None,
|
||||
body_bytes_b64: Some(base64::engine::general_purpose::STANDARD.encode(body)),
|
||||
@@ -2637,6 +2639,7 @@ fn execution_result_frame_stream(
|
||||
payload: StreamFramePayload::Headers {
|
||||
status_code: result.status_code,
|
||||
headers: result.headers.clone(),
|
||||
response_observation: result.response_observation.clone(),
|
||||
},
|
||||
},
|
||||
StreamFrame {
|
||||
|
||||
@@ -480,6 +480,7 @@ mod tests {
|
||||
candidate_id: None,
|
||||
status_code: 502,
|
||||
headers: Default::default(),
|
||||
response_observation: None,
|
||||
body: None,
|
||||
telemetry: None,
|
||||
error: None,
|
||||
@@ -519,6 +520,7 @@ mod tests {
|
||||
candidate_id: None,
|
||||
status_code: 502,
|
||||
headers: Default::default(),
|
||||
response_observation: None,
|
||||
body: None,
|
||||
telemetry: None,
|
||||
error: None,
|
||||
@@ -583,6 +585,7 @@ mod tests {
|
||||
candidate_id: None,
|
||||
status_code: 429,
|
||||
headers: Default::default(),
|
||||
response_observation: None,
|
||||
body: None,
|
||||
telemetry: None,
|
||||
error: None,
|
||||
@@ -614,6 +617,7 @@ mod tests {
|
||||
candidate_id: None,
|
||||
status_code: 401,
|
||||
headers: Default::default(),
|
||||
response_observation: None,
|
||||
body: None,
|
||||
telemetry: None,
|
||||
error: None,
|
||||
@@ -645,6 +649,7 @@ mod tests {
|
||||
candidate_id: None,
|
||||
status_code: 502,
|
||||
headers: Default::default(),
|
||||
response_observation: None,
|
||||
body: None,
|
||||
telemetry: None,
|
||||
error: Some(ExecutionError {
|
||||
@@ -708,6 +713,7 @@ mod tests {
|
||||
candidate_id: None,
|
||||
status_code: 404,
|
||||
headers: Default::default(),
|
||||
response_observation: None,
|
||||
body: None,
|
||||
telemetry: None,
|
||||
error: None,
|
||||
@@ -900,6 +906,7 @@ mod tests {
|
||||
candidate_id: None,
|
||||
status_code: 200,
|
||||
headers: Default::default(),
|
||||
response_observation: None,
|
||||
body: None,
|
||||
telemetry: None,
|
||||
error: None,
|
||||
@@ -1022,6 +1029,7 @@ mod tests {
|
||||
candidate_id: None,
|
||||
status_code: 429,
|
||||
headers: Default::default(),
|
||||
response_observation: None,
|
||||
body: None,
|
||||
telemetry: None,
|
||||
error: None,
|
||||
@@ -1068,6 +1076,7 @@ mod tests {
|
||||
candidate_id: None,
|
||||
status_code: 429,
|
||||
headers: Default::default(),
|
||||
response_observation: None,
|
||||
body: None,
|
||||
telemetry: None,
|
||||
error: None,
|
||||
@@ -1174,6 +1183,7 @@ mod tests {
|
||||
candidate_id: None,
|
||||
status_code: 200,
|
||||
headers: Default::default(),
|
||||
response_observation: None,
|
||||
body: None,
|
||||
telemetry: None,
|
||||
error: None,
|
||||
@@ -1211,6 +1221,7 @@ mod tests {
|
||||
candidate_id: None,
|
||||
status_code: 400,
|
||||
headers: Default::default(),
|
||||
response_observation: None,
|
||||
body: None,
|
||||
telemetry: None,
|
||||
error: None,
|
||||
@@ -1259,6 +1270,7 @@ mod tests {
|
||||
candidate_id: None,
|
||||
status_code: 429,
|
||||
headers: Default::default(),
|
||||
response_observation: None,
|
||||
body: None,
|
||||
telemetry: None,
|
||||
error: None,
|
||||
|
||||
@@ -841,6 +841,7 @@ fn encode_grok_headers_frame(
|
||||
payload: StreamFramePayload::Headers {
|
||||
status_code,
|
||||
headers,
|
||||
response_observation: None,
|
||||
},
|
||||
})
|
||||
}
|
||||
@@ -2157,6 +2158,7 @@ fn grok_execution_result(
|
||||
candidate_id: plan.candidate_id.clone(),
|
||||
status_code,
|
||||
headers: BTreeMap::from([("content-type".to_string(), "application/json".to_string())]),
|
||||
response_observation: None,
|
||||
body: Some(ResponseBody {
|
||||
json_body: Some(body_json),
|
||||
body_bytes_b64: None,
|
||||
@@ -2220,6 +2222,7 @@ fn grok_collected_frame_stream(
|
||||
"application/json".to_string()
|
||||
},
|
||||
)]),
|
||||
response_observation: None,
|
||||
},
|
||||
},
|
||||
StreamFrame {
|
||||
|
||||
@@ -279,6 +279,7 @@ fn raw_response_frame_stream(
|
||||
payload: StreamFramePayload::Headers {
|
||||
status_code,
|
||||
headers,
|
||||
response_observation: None,
|
||||
},
|
||||
},
|
||||
StreamFrame {
|
||||
@@ -1449,6 +1450,7 @@ mod tests {
|
||||
candidate_id: None,
|
||||
status_code: 200,
|
||||
headers: BTreeMap::new(),
|
||||
response_observation: None,
|
||||
body: Some(aether_contracts::ResponseBody {
|
||||
json_body: Some(json!({
|
||||
"jsonrpc": "2.0",
|
||||
|
||||
@@ -3,6 +3,8 @@ use std::collections::BTreeMap;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
pub(crate) mod admission;
|
||||
pub(crate) mod attempt_lifecycle;
|
||||
mod chatgpt_web_image;
|
||||
mod constants;
|
||||
mod fallback;
|
||||
@@ -23,6 +25,9 @@ pub(crate) mod transport;
|
||||
mod transport_failure;
|
||||
mod windsurf;
|
||||
|
||||
pub(crate) use self::admission::{
|
||||
acquire_upstream_execution_gate, UpstreamExecutionGateProvider, UPSTREAM_EXECUTION_GATE_NAME,
|
||||
};
|
||||
pub(crate) use self::chatgpt_web_image::maybe_execute_chatgpt_web_image_sync;
|
||||
pub(crate) use self::constants::{
|
||||
MAX_ERROR_BODY_BYTES, MAX_STREAM_PREFETCH_BYTES, MAX_STREAM_PREFETCH_FRAMES,
|
||||
|
||||
@@ -2,10 +2,10 @@ use aether_contracts::ExecutionPlan;
|
||||
use tracing::warn;
|
||||
|
||||
use crate::orchestration::{
|
||||
oauth_status_may_be_invalid as status_may_be_oauth_invalid,
|
||||
local_failover_error_message, oauth_status_may_be_invalid as status_may_be_oauth_invalid,
|
||||
oauth_status_proves_access_token_invalid as status_proves_access_token_invalid,
|
||||
};
|
||||
use crate::state::AgentIdentityAuthConfigFence;
|
||||
use crate::state::{AgentIdentityAuthConfigFence, CodexRuntimeOAuthObservation};
|
||||
use crate::{provider_transport::LocalOAuthRefreshError, AppState};
|
||||
|
||||
pub(crate) async fn refresh_oauth_plan_auth_for_retry(
|
||||
@@ -14,6 +14,9 @@ pub(crate) async fn refresh_oauth_plan_auth_for_retry(
|
||||
status_code: u16,
|
||||
response_text: Option<&str>,
|
||||
trace_id: &str,
|
||||
report_context: Option<&serde_json::Value>,
|
||||
request_started_at_unix_ms: Option<u64>,
|
||||
request_order_id: Option<&str>,
|
||||
) -> bool {
|
||||
if !status_may_be_oauth_invalid(status_code, response_text) {
|
||||
return false;
|
||||
@@ -109,15 +112,49 @@ pub(crate) async fn refresh_oauth_plan_auth_for_retry(
|
||||
body_excerpt,
|
||||
..
|
||||
}) if matches!(refresh_status_code, 400 | 401 | 403) => {
|
||||
if let Err(err) = state
|
||||
.persist_local_oauth_refresh_failure_state(
|
||||
&transport,
|
||||
refresh_status_code,
|
||||
body_excerpt.as_str(),
|
||||
access_token_invalid_proven,
|
||||
)
|
||||
.await
|
||||
{
|
||||
let observed_credential_generation =
|
||||
report_context_string(report_context, "codex_credential_generation");
|
||||
let runtime_invalid_message = local_failover_error_message(response_text);
|
||||
let runtime_invalid_reason =
|
||||
aether_admin::provider::quota::codex_runtime_invalid_reason(
|
||||
status_code,
|
||||
runtime_invalid_message.as_deref(),
|
||||
);
|
||||
let persist_result = match (request_started_at_unix_ms, request_order_id) {
|
||||
(Some(request_started_at_unix_ms), Some(request_order_id))
|
||||
if transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("codex") =>
|
||||
{
|
||||
state
|
||||
.persist_local_oauth_refresh_failure_state_observed(
|
||||
&transport,
|
||||
refresh_status_code,
|
||||
body_excerpt.as_str(),
|
||||
access_token_invalid_proven,
|
||||
CodexRuntimeOAuthObservation {
|
||||
request_started_at_unix_ms,
|
||||
request_order_id,
|
||||
observed_credential_generation,
|
||||
runtime_invalid_reason: runtime_invalid_reason.as_deref(),
|
||||
},
|
||||
)
|
||||
.await
|
||||
}
|
||||
_ => {
|
||||
state
|
||||
.persist_local_oauth_refresh_failure_state(
|
||||
&transport,
|
||||
refresh_status_code,
|
||||
body_excerpt.as_str(),
|
||||
access_token_invalid_proven,
|
||||
)
|
||||
.await
|
||||
}
|
||||
};
|
||||
if let Err(err) = persist_result {
|
||||
warn!(
|
||||
event_name = "local_oauth_retry_refresh_failure_persist_failed",
|
||||
log_type = "ops",
|
||||
@@ -161,6 +198,17 @@ pub(crate) async fn refresh_oauth_plan_auth_for_retry(
|
||||
}
|
||||
}
|
||||
|
||||
fn report_context_string<'a>(
|
||||
report_context: Option<&'a serde_json::Value>,
|
||||
field: &str,
|
||||
) -> Option<&'a str> {
|
||||
report_context
|
||||
.and_then(|context| context.get(field))
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
}
|
||||
|
||||
fn execution_plan_authorization(plan: &ExecutionPlan) -> Option<&str> {
|
||||
plan.headers
|
||||
.iter()
|
||||
@@ -209,6 +257,7 @@ mod tests {
|
||||
use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY};
|
||||
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyOAuthCredentialFence,
|
||||
ProviderCatalogReadRepository, ProviderCatalogWriteRepository,
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||
};
|
||||
@@ -312,7 +361,7 @@ mod tests {
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn auto_removes_codex_key_after_request_proven_terminal_refresh_failure() {
|
||||
async fn retains_codex_key_after_request_proven_terminal_refresh_failure() {
|
||||
let token_hits = Arc::new(Mutex::new(0usize));
|
||||
let token_hits_clone = Arc::clone(&token_hits);
|
||||
let token_server = Router::new().route(
|
||||
@@ -458,16 +507,34 @@ mod tests {
|
||||
401,
|
||||
Some(r#"{"error":"oauth_token_invalid"}"#),
|
||||
"trace-oauth-retry",
|
||||
None,
|
||||
Some(1_000),
|
||||
Some("01900000-0000-7000-8000-000000000010"),
|
||||
)
|
||||
.await;
|
||||
|
||||
assert!(!retried);
|
||||
assert_eq!(*token_hits.lock().expect("mutex should lock"), 1);
|
||||
let keys = provider_catalog_repository
|
||||
let stored_key = provider_catalog_repository
|
||||
.list_keys_by_ids(&["key-codex-oauth-retry".to_string()])
|
||||
.await
|
||||
.expect("keys should read");
|
||||
assert!(keys.is_empty());
|
||||
.expect("keys should read")
|
||||
.into_iter()
|
||||
.next()
|
||||
.expect("request-scoped refresh failure should retain the key");
|
||||
let invalid_reason = stored_key
|
||||
.oauth_invalid_reason
|
||||
.as_deref()
|
||||
.expect("combined invalid reason should persist");
|
||||
assert!(invalid_reason.contains("[OAUTH_EXPIRED]"));
|
||||
assert!(invalid_reason.contains("[REFRESH_FAILED]"));
|
||||
assert_eq!(
|
||||
stored_key
|
||||
.upstream_metadata
|
||||
.as_ref()
|
||||
.and_then(|metadata| metadata.pointer("/codex/oauth_state_request_id")),
|
||||
Some(&json!("01900000-0000-7000-8000-000000000010"))
|
||||
);
|
||||
|
||||
token_handle.abort();
|
||||
}
|
||||
@@ -619,6 +686,9 @@ mod tests {
|
||||
401,
|
||||
Some(r#"{"error":"invalid_token"}"#),
|
||||
"trace-claude-oauth-fence-first",
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.await
|
||||
);
|
||||
@@ -647,6 +717,9 @@ mod tests {
|
||||
401,
|
||||
Some(r#"{"error":"invalid_token"}"#),
|
||||
"trace-claude-oauth-fence-stale",
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.await
|
||||
);
|
||||
@@ -665,6 +738,7 @@ mod tests {
|
||||
.expect("Claude key should load")
|
||||
.pop()
|
||||
.expect("Claude key should exist");
|
||||
let expected_admin_replacement = admin_replacement.clone();
|
||||
admin_replacement.encrypted_api_key = Some(
|
||||
encrypt_python_fernet_plaintext(
|
||||
DEVELOPMENT_ENCRYPTION_KEY,
|
||||
@@ -673,10 +747,23 @@ mod tests {
|
||||
.expect("admin access token should encrypt"),
|
||||
);
|
||||
admin_replacement.expires_at_unix_secs = Some(4_102_444_800);
|
||||
provider_catalog_repository
|
||||
.update_key(&admin_replacement)
|
||||
assert!(provider_catalog_repository
|
||||
.compare_and_update_key_admin_state(&ProviderCatalogKeyAdminCasUpdate {
|
||||
expected_encrypted_auth_config: expected_admin_replacement
|
||||
.encrypted_auth_config
|
||||
.clone(),
|
||||
expected_credential: ProviderCatalogKeyOAuthCredentialFence {
|
||||
encrypted_api_key: expected_admin_replacement.encrypted_api_key.clone(),
|
||||
auth_type: expected_admin_replacement.auth_type.clone(),
|
||||
provider_id: expected_admin_replacement.provider_id.clone(),
|
||||
provider_type: "claude_code".to_string(),
|
||||
},
|
||||
key: admin_replacement,
|
||||
codex_rotation: None,
|
||||
reset_oauth_runtime: true,
|
||||
})
|
||||
.await
|
||||
.expect("admin replacement should persist");
|
||||
.expect("admin replacement CAS should run"));
|
||||
|
||||
let admin_result = state
|
||||
.force_local_oauth_refresh_entry(&stale_transport)
|
||||
|
||||
@@ -10,6 +10,10 @@ use crate::{AppState, GatewayError};
|
||||
const RESPONSE_HEADER_RULES_KEY: &str = "response_header_rules";
|
||||
const RESPONSE_HEADER_RULES_CAMEL_KEY: &str = "responseHeaderRules";
|
||||
const PROVIDER_RESPONSE_HEADERS_CONTEXT_KEY: &str = "provider_response_headers";
|
||||
const PROVIDER_REQUEST_STARTED_AT_UNIX_MS_CONTEXT_KEY: &str = "provider_request_started_at_unix_ms";
|
||||
const PROVIDER_REQUEST_ORDER_ID_CONTEXT_KEY: &str = "provider_request_order_id";
|
||||
const PROVIDER_RESPONSE_HEADERS_OBSERVED_AT_UNIX_MS_CONTEXT_KEY: &str =
|
||||
"provider_response_headers_observed_at_unix_ms";
|
||||
const RESPONSE_HEADER_RULE_PROTECTED_KEYS: &[&str] = &["content-length"];
|
||||
const RESPONSE_HEADER_RULES_CACHE_TTL: Duration = Duration::from_secs(5);
|
||||
|
||||
@@ -98,6 +102,9 @@ pub(crate) async fn apply_endpoint_response_header_rules(
|
||||
pub(crate) fn attach_provider_response_headers_to_report_context(
|
||||
report_context: Option<Value>,
|
||||
provider_headers: &BTreeMap<String, String>,
|
||||
provider_request_started_at_unix_ms: u64,
|
||||
provider_response_headers_observed_at_unix_ms: u64,
|
||||
provider_request_order_id: &str,
|
||||
) -> Option<Value> {
|
||||
let provider_headers = serde_json::to_value(provider_headers).ok()?;
|
||||
let mut object = match report_context {
|
||||
@@ -105,9 +112,99 @@ pub(crate) fn attach_provider_response_headers_to_report_context(
|
||||
Some(other) => Map::from_iter([("seed".to_string(), other)]),
|
||||
None => Map::new(),
|
||||
};
|
||||
object.insert(
|
||||
PROVIDER_RESPONSE_HEADERS_CONTEXT_KEY.to_string(),
|
||||
provider_headers,
|
||||
);
|
||||
let observation_is_absent = !object.contains_key(PROVIDER_RESPONSE_HEADERS_CONTEXT_KEY)
|
||||
&& !object.contains_key(PROVIDER_REQUEST_STARTED_AT_UNIX_MS_CONTEXT_KEY)
|
||||
&& !object.contains_key(PROVIDER_RESPONSE_HEADERS_OBSERVED_AT_UNIX_MS_CONTEXT_KEY)
|
||||
&& !object.contains_key(PROVIDER_REQUEST_ORDER_ID_CONTEXT_KEY);
|
||||
if observation_is_absent {
|
||||
object.insert(
|
||||
PROVIDER_RESPONSE_HEADERS_CONTEXT_KEY.to_string(),
|
||||
provider_headers,
|
||||
);
|
||||
object.insert(
|
||||
PROVIDER_REQUEST_STARTED_AT_UNIX_MS_CONTEXT_KEY.to_string(),
|
||||
Value::from(provider_request_started_at_unix_ms),
|
||||
);
|
||||
object.insert(
|
||||
PROVIDER_RESPONSE_HEADERS_OBSERVED_AT_UNIX_MS_CONTEXT_KEY.to_string(),
|
||||
Value::from(provider_response_headers_observed_at_unix_ms),
|
||||
);
|
||||
object.insert(
|
||||
PROVIDER_REQUEST_ORDER_ID_CONTEXT_KEY.to_string(),
|
||||
Value::from(provider_request_order_id),
|
||||
);
|
||||
}
|
||||
Some(Value::Object(object))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use serde_json::json;
|
||||
|
||||
#[test]
|
||||
fn provider_response_observation_is_first_write_wins() {
|
||||
let first_headers =
|
||||
BTreeMap::from([("x-codex-primary-used-percent".to_string(), "10".to_string())]);
|
||||
let second_headers =
|
||||
BTreeMap::from([("x-codex-primary-used-percent".to_string(), "20".to_string())]);
|
||||
|
||||
let report_context = attach_provider_response_headers_to_report_context(
|
||||
Some(json!("seed-value")),
|
||||
&first_headers,
|
||||
100,
|
||||
200,
|
||||
"observation-1",
|
||||
);
|
||||
let report_context = attach_provider_response_headers_to_report_context(
|
||||
report_context,
|
||||
&second_headers,
|
||||
300,
|
||||
400,
|
||||
"observation-2",
|
||||
)
|
||||
.expect("report context should exist");
|
||||
|
||||
assert_eq!(report_context["seed"], json!("seed-value"));
|
||||
assert_eq!(
|
||||
report_context["provider_response_headers"]["x-codex-primary-used-percent"],
|
||||
json!("10")
|
||||
);
|
||||
assert_eq!(
|
||||
report_context["provider_request_started_at_unix_ms"],
|
||||
json!(100)
|
||||
);
|
||||
assert_eq!(
|
||||
report_context["provider_response_headers_observed_at_unix_ms"],
|
||||
json!(200)
|
||||
);
|
||||
assert_eq!(
|
||||
report_context["provider_request_order_id"],
|
||||
json!("observation-1")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_response_observation_does_not_complete_a_partial_triplet() {
|
||||
let report_context = attach_provider_response_headers_to_report_context(
|
||||
Some(json!({"provider_response_headers": {"x-existing": "1"}})),
|
||||
&BTreeMap::from([("x-new".to_string(), "2".to_string())]),
|
||||
300,
|
||||
400,
|
||||
"observation-2",
|
||||
)
|
||||
.expect("report context should exist");
|
||||
|
||||
assert_eq!(
|
||||
report_context["provider_response_headers"]["x-existing"],
|
||||
json!("1")
|
||||
);
|
||||
assert!(report_context
|
||||
.get("provider_request_started_at_unix_ms")
|
||||
.is_none());
|
||||
assert!(report_context
|
||||
.get("provider_response_headers_observed_at_unix_ms")
|
||||
.is_none());
|
||||
assert!(report_context.get("provider_request_order_id").is_none());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -11,8 +11,8 @@ use std::time::{Duration, Instant};
|
||||
|
||||
use aether_ai_serving::{AiAttemptExecutionOutcome, AiAttemptRetryScope};
|
||||
use aether_contracts::{
|
||||
ExecutionPlan, ExecutionStreamTerminalSummary, ExecutionTelemetry, StandardizedUsage,
|
||||
StreamFrame, StreamFramePayload,
|
||||
ExecutionPlan, ExecutionResponseObservation, ExecutionStreamTerminalSummary,
|
||||
ExecutionTelemetry, StandardizedUsage, StreamFrame, StreamFramePayload,
|
||||
};
|
||||
use aether_data_contracts::repository::candidates::{
|
||||
RequestCandidateStatus, UpsertRequestCandidateRecord,
|
||||
@@ -112,12 +112,13 @@ use crate::execution_runtime::{
|
||||
use crate::log_ids::short_request_id;
|
||||
use crate::orchestration::{
|
||||
apply_local_execution_effect, build_local_error_flow_metadata, classify_failure_disposition,
|
||||
cyber_continue_failover_enabled, trace_upstream_response_body, with_error_flow_report_context,
|
||||
cyber_continue_failover_enabled, spawn_local_oauth_success_effect,
|
||||
trace_upstream_response_body, with_error_flow_report_context,
|
||||
with_upstream_response_report_context, FailureDisposition, FailureTokenAction,
|
||||
LocalAdaptiveRateLimitEffect, LocalAdaptiveSuccessEffect, LocalAttemptFailureEffect,
|
||||
LocalExecutionEffect, LocalExecutionEffectContext, LocalFailoverAnalysis,
|
||||
LocalHealthFailureEffect, LocalHealthSuccessEffect, LocalOAuthInvalidationEffect,
|
||||
LocalPoolErrorEffect,
|
||||
LocalOAuthSuccessEffect, LocalPoolErrorEffect,
|
||||
};
|
||||
use crate::provider_pool_demand::{
|
||||
acquire_provider_pool_in_flight_guard, ProviderPoolInFlightGuard,
|
||||
@@ -1249,6 +1250,9 @@ async fn execute_in_process_stream_with_oauth_retry(
|
||||
retry_status_code,
|
||||
response_text.as_deref(),
|
||||
trace_id,
|
||||
report_context,
|
||||
Some(execution.response_observation.request_started_at_unix_ms),
|
||||
Some(&execution.response_observation.request_order_id),
|
||||
)
|
||||
.await
|
||||
{
|
||||
@@ -2818,6 +2822,7 @@ async fn execute_stream_from_direct_passthrough(
|
||||
stream_precommit_committed: _,
|
||||
response,
|
||||
started_at: upstream_started_at,
|
||||
response_observation,
|
||||
stream_first_byte_timeout,
|
||||
upstream_target_permit,
|
||||
} = execution;
|
||||
@@ -2834,8 +2839,23 @@ async fn execute_stream_from_direct_passthrough(
|
||||
let request_id = plan.request_id.clone();
|
||||
let candidate_id = plan.candidate_id.clone();
|
||||
let request_id_for_log = short_request_id(request_id.as_str());
|
||||
let mut report_context =
|
||||
attach_provider_response_headers_to_report_context(report_context, &headers);
|
||||
let mut report_context = attach_provider_response_headers_to_report_context(
|
||||
report_context,
|
||||
&headers,
|
||||
response_observation.request_started_at_unix_ms,
|
||||
response_observation.response_headers_observed_at_unix_ms,
|
||||
&response_observation.request_order_id,
|
||||
);
|
||||
spawn_local_oauth_success_effect(
|
||||
state.clone(),
|
||||
&plan,
|
||||
report_context.as_ref(),
|
||||
LocalOAuthSuccessEffect {
|
||||
status_code,
|
||||
request_started_at_unix_ms: Some(response_observation.request_started_at_unix_ms),
|
||||
request_order_id: Some(&response_observation.request_order_id),
|
||||
},
|
||||
);
|
||||
if status_code == 200 {
|
||||
seed_kiro_simulated_cache_enabled(state, &plan, &mut report_context).await;
|
||||
if kiro_simulated_cache_enabled_from_report_context(report_context.as_ref()) {
|
||||
@@ -3819,6 +3839,7 @@ async fn execute_execution_runtime_stream_inner(
|
||||
provider_pool_in_flight_guard.take(),
|
||||
retry_scope_out.as_deref_mut(),
|
||||
retry_fallback_out.as_deref_mut(),
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
@@ -3891,6 +3912,7 @@ async fn execute_execution_runtime_stream_inner(
|
||||
provider_pool_in_flight_guard.take(),
|
||||
retry_scope_out.as_deref_mut(),
|
||||
retry_fallback_out.as_deref_mut(),
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
@@ -3963,6 +3985,7 @@ async fn execute_execution_runtime_stream_inner(
|
||||
provider_pool_in_flight_guard.take(),
|
||||
retry_scope_out.as_deref_mut(),
|
||||
retry_fallback_out.as_deref_mut(),
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
@@ -4035,6 +4058,7 @@ async fn execute_execution_runtime_stream_inner(
|
||||
provider_pool_in_flight_guard.take(),
|
||||
retry_scope_out.as_deref_mut(),
|
||||
retry_fallback_out.as_deref_mut(),
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
@@ -4192,6 +4216,15 @@ async fn execute_execution_runtime_stream_inner(
|
||||
record_stream_pending_lifecycle(state, seed, &mut stage_trace).await;
|
||||
lifecycle_pending_recorded = true;
|
||||
}
|
||||
let report_context = attach_provider_response_headers_to_report_context(
|
||||
report_context,
|
||||
&execution.headers,
|
||||
execution.response_observation.request_started_at_unix_ms,
|
||||
execution
|
||||
.response_observation
|
||||
.response_headers_observed_at_unix_ms,
|
||||
&execution.response_observation.request_order_id,
|
||||
);
|
||||
let stream_precommit_committed = execution.stream_precommit_committed;
|
||||
let frame_stream = build_direct_execution_frame_stream(execution).boxed();
|
||||
return execute_stream_from_frame_stream_with_retry_scope(
|
||||
@@ -4211,6 +4244,7 @@ async fn execute_execution_runtime_stream_inner(
|
||||
provider_pool_in_flight_guard.take(),
|
||||
retry_scope_out,
|
||||
retry_fallback_out,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
@@ -4327,6 +4361,15 @@ async fn execute_execution_runtime_stream_inner(
|
||||
record_stream_pending_lifecycle(state, seed, &mut stage_trace).await;
|
||||
lifecycle_pending_recorded = true;
|
||||
}
|
||||
let report_context = attach_provider_response_headers_to_report_context(
|
||||
report_context,
|
||||
&execution.headers,
|
||||
execution.response_observation.request_started_at_unix_ms,
|
||||
execution
|
||||
.response_observation
|
||||
.response_headers_observed_at_unix_ms,
|
||||
&execution.response_observation.request_order_id,
|
||||
);
|
||||
let stream_precommit_committed = execution.stream_precommit_committed;
|
||||
let frame_stream = build_direct_execution_frame_stream(execution).boxed();
|
||||
return execute_stream_from_frame_stream_with_retry_scope(
|
||||
@@ -4346,10 +4389,13 @@ async fn execute_execution_runtime_stream_inner(
|
||||
provider_pool_in_flight_guard.take(),
|
||||
retry_scope_out.as_deref_mut(),
|
||||
retry_fallback_out.as_deref_mut(),
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
let remote_request_started_at_unix_ms = current_request_candidate_unix_ms();
|
||||
let remote_request_order_id = uuid::Uuid::now_v7().to_string();
|
||||
let response = match post_stream_plan_to_remote_execution_runtime(
|
||||
state,
|
||||
remote_execution_runtime_base_url,
|
||||
@@ -4431,6 +4477,12 @@ async fn execute_execution_runtime_stream_inner(
|
||||
)?));
|
||||
}
|
||||
|
||||
let remote_response_observed_at_unix_ms = current_request_candidate_unix_ms();
|
||||
let remote_fallback_observation = ExecutionResponseObservation {
|
||||
request_started_at_unix_ms: remote_request_started_at_unix_ms,
|
||||
response_headers_observed_at_unix_ms: remote_response_observed_at_unix_ms,
|
||||
request_order_id: remote_request_order_id,
|
||||
};
|
||||
let frame_stream = response
|
||||
.bytes_stream()
|
||||
.map_err(|err| IoError::other(err.to_string()))
|
||||
@@ -4452,6 +4504,7 @@ async fn execute_execution_runtime_stream_inner(
|
||||
provider_pool_in_flight_guard.take(),
|
||||
retry_scope_out.as_deref_mut(),
|
||||
retry_fallback_out.as_deref_mut(),
|
||||
Some(remote_fallback_observation),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
@@ -5481,6 +5534,7 @@ async fn execute_stream_from_frame_stream(
|
||||
in_flight_guard,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.await
|
||||
}
|
||||
@@ -5503,6 +5557,7 @@ async fn execute_stream_from_frame_stream_with_retry_scope(
|
||||
in_flight_guard: Option<ProviderPoolInFlightGuard>,
|
||||
mut retry_scope_out: Option<&mut AiAttemptRetryScope>,
|
||||
mut retry_fallback_out: Option<&mut Option<Response<Body>>>,
|
||||
fallback_response_observation: Option<ExecutionResponseObservation>,
|
||||
) -> Result<Option<Response<Body>>, GatewayError> {
|
||||
let request_id = plan.request_id.as_str();
|
||||
let request_id_for_log = short_request_id(request_id);
|
||||
@@ -5535,14 +5590,37 @@ async fn execute_stream_from_frame_stream_with_retry_scope(
|
||||
let StreamFramePayload::Headers {
|
||||
status_code,
|
||||
mut headers,
|
||||
response_observation,
|
||||
} = first_frame.payload
|
||||
else {
|
||||
return Err(GatewayError::Internal(
|
||||
"execution runtime stream must start with headers frame".to_string(),
|
||||
));
|
||||
};
|
||||
let mut report_context =
|
||||
attach_provider_response_headers_to_report_context(report_context, &headers);
|
||||
let response_observation = response_observation
|
||||
.or(fallback_response_observation)
|
||||
.unwrap_or(ExecutionResponseObservation {
|
||||
request_started_at_unix_ms: candidate_started_unix_secs,
|
||||
response_headers_observed_at_unix_ms: current_request_candidate_unix_ms(),
|
||||
request_order_id: uuid::Uuid::now_v7().to_string(),
|
||||
});
|
||||
let mut report_context = attach_provider_response_headers_to_report_context(
|
||||
report_context,
|
||||
&headers,
|
||||
response_observation.request_started_at_unix_ms,
|
||||
response_observation.response_headers_observed_at_unix_ms,
|
||||
&response_observation.request_order_id,
|
||||
);
|
||||
spawn_local_oauth_success_effect(
|
||||
state.clone(),
|
||||
&plan,
|
||||
report_context.as_ref(),
|
||||
LocalOAuthSuccessEffect {
|
||||
status_code,
|
||||
request_started_at_unix_ms: Some(response_observation.request_started_at_unix_ms),
|
||||
request_order_id: Some(&response_observation.request_order_id),
|
||||
},
|
||||
);
|
||||
if status_code == 200 {
|
||||
seed_kiro_simulated_cache_enabled(state, &plan, &mut report_context).await;
|
||||
if kiro_simulated_cache_enabled_from_report_context(report_context.as_ref()) {
|
||||
@@ -8310,6 +8388,7 @@ mod tests {
|
||||
"content-type".to_string(),
|
||||
"text/event-stream".to_string(),
|
||||
)]),
|
||||
response_observation: None,
|
||||
},
|
||||
}));
|
||||
yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame {
|
||||
@@ -8389,6 +8468,7 @@ mod tests {
|
||||
"content-type".to_string(),
|
||||
"text/event-stream".to_string(),
|
||||
)]),
|
||||
response_observation: None,
|
||||
},
|
||||
}));
|
||||
yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame {
|
||||
@@ -8431,6 +8511,7 @@ mod tests {
|
||||
None,
|
||||
Some(&mut retry_scope),
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.await
|
||||
.expect("prefetch transport execution should resolve");
|
||||
@@ -8480,6 +8561,7 @@ mod tests {
|
||||
"content-type".to_string(),
|
||||
"text/event-stream".to_string(),
|
||||
)]),
|
||||
response_observation: None,
|
||||
},
|
||||
}));
|
||||
yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame {
|
||||
@@ -8522,6 +8604,7 @@ mod tests {
|
||||
None,
|
||||
Some(&mut retry_scope),
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.await
|
||||
.expect("prefetch HTTP status execution should resolve");
|
||||
@@ -8680,6 +8763,7 @@ mod tests {
|
||||
"content-type".to_string(),
|
||||
"text/event-stream".to_string(),
|
||||
)]),
|
||||
response_observation: None,
|
||||
},
|
||||
}));
|
||||
for chunk in chunks {
|
||||
@@ -8725,6 +8809,7 @@ mod tests {
|
||||
None,
|
||||
Some(&mut retry_scope),
|
||||
Some(&mut fallback_response),
|
||||
None,
|
||||
)
|
||||
.await
|
||||
.expect("native Anthropic stream execution should succeed");
|
||||
@@ -9364,6 +9449,7 @@ mod tests {
|
||||
"content-type".to_string(),
|
||||
"text/event-stream".to_string(),
|
||||
)]),
|
||||
response_observation: None,
|
||||
},
|
||||
}));
|
||||
yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame {
|
||||
@@ -9413,6 +9499,7 @@ mod tests {
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
),
|
||||
)
|
||||
.await
|
||||
@@ -9850,6 +9937,7 @@ mod tests {
|
||||
"content-type".to_string(),
|
||||
"text/event-stream".to_string(),
|
||||
)]),
|
||||
response_observation: None,
|
||||
},
|
||||
}));
|
||||
}
|
||||
@@ -11532,6 +11620,7 @@ mod tests {
|
||||
"content-type".to_string(),
|
||||
"text/event-stream".to_string(),
|
||||
)]),
|
||||
response_observation: None,
|
||||
},
|
||||
}));
|
||||
yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame {
|
||||
@@ -11660,6 +11749,7 @@ mod tests {
|
||||
"content-type".to_string(),
|
||||
"text/event-stream".to_string(),
|
||||
)]),
|
||||
response_observation: None,
|
||||
},
|
||||
}));
|
||||
yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame {
|
||||
@@ -12384,6 +12474,7 @@ mod tests {
|
||||
"content-type".to_string(),
|
||||
"text/event-stream".to_string(),
|
||||
)]),
|
||||
response_observation: None,
|
||||
},
|
||||
}));
|
||||
yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame {
|
||||
|
||||
@@ -4,8 +4,9 @@ use std::io::Error as IoError;
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
use aether_contracts::{
|
||||
ExecutionError, ExecutionErrorKind, ExecutionPhase, ExecutionStreamTerminalSummary,
|
||||
ExecutionTelemetry, StreamFrame, StreamFramePayload, StreamFrameType,
|
||||
ExecutionError, ExecutionErrorKind, ExecutionPhase, ExecutionResponseObservation,
|
||||
ExecutionStreamTerminalSummary, ExecutionTelemetry, StreamFrame, StreamFramePayload,
|
||||
StreamFrameType,
|
||||
};
|
||||
use async_stream::stream;
|
||||
use axum::body::Bytes;
|
||||
@@ -44,6 +45,7 @@ pub(crate) fn build_direct_execution_frame_stream(
|
||||
stream_precommit_committed: _,
|
||||
response,
|
||||
started_at,
|
||||
response_observation,
|
||||
stream_first_byte_timeout,
|
||||
upstream_target_permit,
|
||||
} = execution;
|
||||
@@ -108,7 +110,11 @@ pub(crate) fn build_direct_execution_frame_stream(
|
||||
}
|
||||
}
|
||||
|
||||
match encode_headers_frame(status_code, response_headers) {
|
||||
match encode_headers_frame(
|
||||
status_code,
|
||||
response_headers,
|
||||
&response_observation,
|
||||
) {
|
||||
Ok(frame) => yield Ok(frame),
|
||||
Err(err) => {
|
||||
yield Err(err);
|
||||
@@ -153,7 +159,11 @@ pub(crate) fn build_direct_execution_frame_stream(
|
||||
upstream_bytes,
|
||||
first_byte_timeout,
|
||||
}) => {
|
||||
match encode_headers_frame(status_code, original_headers) {
|
||||
match encode_headers_frame(
|
||||
status_code,
|
||||
original_headers,
|
||||
&response_observation,
|
||||
) {
|
||||
Ok(frame) => yield Ok(frame),
|
||||
Err(err) => {
|
||||
yield Err(err);
|
||||
@@ -192,7 +202,11 @@ pub(crate) fn build_direct_execution_frame_stream(
|
||||
return;
|
||||
}
|
||||
|
||||
match encode_headers_frame(status_code, headers) {
|
||||
match encode_headers_frame(
|
||||
status_code,
|
||||
headers,
|
||||
&response_observation,
|
||||
) {
|
||||
Ok(frame) => yield Ok(frame),
|
||||
Err(err) => {
|
||||
yield Err(err);
|
||||
@@ -611,12 +625,14 @@ pub(crate) fn build_direct_execution_frame_stream(
|
||||
fn encode_headers_frame(
|
||||
status_code: u16,
|
||||
headers: BTreeMap<String, String>,
|
||||
response_observation: &ExecutionResponseObservation,
|
||||
) -> Result<Bytes, IoError> {
|
||||
encode_stream_frame_ndjson(&StreamFrame {
|
||||
frame_type: StreamFrameType::Headers,
|
||||
payload: StreamFramePayload::Headers {
|
||||
status_code,
|
||||
headers,
|
||||
response_observation: Some(response_observation.clone()),
|
||||
},
|
||||
})
|
||||
}
|
||||
@@ -1606,43 +1622,47 @@ mod tests {
|
||||
.expect("listener should bind");
|
||||
let addr = listener.local_addr().expect("local addr should resolve");
|
||||
let server = tokio::spawn(async move {
|
||||
let app = Router::new().route(
|
||||
"/responses",
|
||||
post(|| async {
|
||||
let body = serde_json::json!({
|
||||
"id": "resp_sync_bridge_123",
|
||||
"object": "response",
|
||||
"model": "gpt-5.4",
|
||||
"status": "completed",
|
||||
"output": [{
|
||||
"type": "message",
|
||||
"id": "msg_sync_bridge_123",
|
||||
"role": "assistant",
|
||||
"content": [{
|
||||
"type": "output_text",
|
||||
"text": "Hello from buffered JSON stream",
|
||||
"annotations": []
|
||||
}]
|
||||
}],
|
||||
"usage": {
|
||||
"input_tokens": 1,
|
||||
"output_tokens": 2,
|
||||
"total_tokens": 3
|
||||
}
|
||||
});
|
||||
let mut response = axum::http::Response::new(Body::from(
|
||||
serde_json::to_vec(&body).expect("json should encode"),
|
||||
));
|
||||
response.headers_mut().insert(
|
||||
header::CONTENT_TYPE,
|
||||
HeaderValue::from_static("application/json"),
|
||||
);
|
||||
response
|
||||
}),
|
||||
);
|
||||
axum::serve(listener, app)
|
||||
let (mut socket, _) = listener.accept().await.expect("client should connect");
|
||||
let mut request = [0_u8; 4096];
|
||||
let _ = socket
|
||||
.read(&mut request)
|
||||
.await
|
||||
.expect("server should start");
|
||||
.expect("request should read");
|
||||
let body = serde_json::to_vec(&serde_json::json!({
|
||||
"id": "resp_sync_bridge_123",
|
||||
"object": "response",
|
||||
"model": "gpt-5.4",
|
||||
"status": "completed",
|
||||
"output": [{
|
||||
"type": "message",
|
||||
"id": "msg_sync_bridge_123",
|
||||
"role": "assistant",
|
||||
"content": [{
|
||||
"type": "output_text",
|
||||
"text": "Hello from buffered JSON stream",
|
||||
"annotations": []
|
||||
}]
|
||||
}],
|
||||
"usage": {
|
||||
"input_tokens": 1,
|
||||
"output_tokens": 2,
|
||||
"total_tokens": 3
|
||||
}
|
||||
}))
|
||||
.expect("json should encode");
|
||||
socket
|
||||
.write_all(
|
||||
format!(
|
||||
"HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\n\r\n",
|
||||
body.len()
|
||||
)
|
||||
.as_bytes(),
|
||||
)
|
||||
.await
|
||||
.expect("headers should write");
|
||||
socket.flush().await.expect("headers should flush");
|
||||
tokio::time::sleep(Duration::from_millis(75)).await;
|
||||
socket.write_all(&body).await.expect("body should write");
|
||||
});
|
||||
|
||||
let runtime = DirectSyncExecutionRuntime::new();
|
||||
@@ -1678,6 +1698,12 @@ mod tests {
|
||||
})
|
||||
.await
|
||||
.expect("stream execution should succeed");
|
||||
let expected_observation = execution.response_observation.clone();
|
||||
assert!(
|
||||
expected_observation.response_headers_observed_at_unix_ms
|
||||
>= expected_observation.request_started_at_unix_ms
|
||||
);
|
||||
assert!(!expected_observation.request_order_id.is_empty());
|
||||
|
||||
let frames = build_direct_execution_frame_stream(execution)
|
||||
.map(|item| item.expect("frame should encode"))
|
||||
@@ -1691,6 +1717,10 @@ mod tests {
|
||||
|
||||
let header_frame: Value =
|
||||
serde_json::from_str(&frames[0]).expect("headers frame should parse");
|
||||
let encoded_observation: aether_contracts::ExecutionResponseObservation =
|
||||
serde_json::from_value(header_frame["payload"]["response_observation"].clone())
|
||||
.expect("headers frame should retain the response observation");
|
||||
assert_eq!(encoded_observation, expected_observation);
|
||||
assert_eq!(
|
||||
header_frame
|
||||
.get("payload")
|
||||
|
||||
@@ -5,8 +5,8 @@ use std::time::{Duration, Instant};
|
||||
|
||||
use aether_ai_serving::{AiAttemptExecutionOutcome, AiAttemptRetryScope, UPSTREAM_IS_STREAM_KEY};
|
||||
use aether_contracts::{
|
||||
ExecutionError, ExecutionErrorKind, ExecutionPhase, ExecutionPlan, ExecutionResult,
|
||||
ExecutionTelemetry,
|
||||
ExecutionError, ExecutionErrorKind, ExecutionPhase, ExecutionPlan,
|
||||
ExecutionResponseObservation, ExecutionResult, ExecutionTelemetry,
|
||||
};
|
||||
use aether_data_contracts::repository::candidates::RequestCandidateStatus;
|
||||
use aether_scheduler_core::{
|
||||
@@ -55,8 +55,9 @@ use crate::execution_runtime::submission::{
|
||||
resolve_local_sync_error_status_code, submit_local_core_error_or_sync_finalize,
|
||||
};
|
||||
use crate::execution_runtime::transport::{
|
||||
append_upstream_response_body_chunk, build_execution_response_body, build_request_body,
|
||||
collect_response_headers, decode_response_body_bytes, execution_response_body_mode,
|
||||
append_upstream_response_body_chunk_with_limit, build_execution_response_body,
|
||||
build_request_body, collect_response_headers, decode_response_body_bytes_with_limit,
|
||||
execution_plan_response_body_limit_bytes, execution_response_body_mode,
|
||||
format_hyper_error_chain, format_upstream_request_error, format_wreq_upstream_request_error,
|
||||
response_body_is_json, send_request, DirectHttpResponse, DirectSyncExecutionRuntime,
|
||||
ExecutionRuntimeTransportError,
|
||||
@@ -70,11 +71,12 @@ use crate::execution_runtime::{
|
||||
};
|
||||
use crate::log_ids::short_request_id;
|
||||
use crate::orchestration::{
|
||||
apply_local_execution_effect, build_local_error_flow_metadata, trace_upstream_response_body,
|
||||
with_error_flow_report_context, with_upstream_response_report_context,
|
||||
LocalAdaptiveRateLimitEffect, LocalAdaptiveSuccessEffect, LocalAttemptFailureEffect,
|
||||
LocalExecutionEffect, LocalExecutionEffectContext, LocalHealthFailureEffect,
|
||||
LocalHealthSuccessEffect, LocalOAuthInvalidationEffect, LocalPoolErrorEffect,
|
||||
apply_local_execution_effect, build_local_error_flow_metadata,
|
||||
spawn_local_oauth_success_effect, trace_upstream_response_body, with_error_flow_report_context,
|
||||
with_upstream_response_report_context, LocalAdaptiveRateLimitEffect,
|
||||
LocalAdaptiveSuccessEffect, LocalAttemptFailureEffect, LocalExecutionEffect,
|
||||
LocalExecutionEffectContext, LocalHealthFailureEffect, LocalHealthSuccessEffect,
|
||||
LocalOAuthInvalidationEffect, LocalOAuthSuccessEffect, LocalPoolErrorEffect,
|
||||
};
|
||||
use crate::provider_pool_demand::acquire_provider_pool_in_flight_guard;
|
||||
use crate::request_candidate_runtime::{
|
||||
@@ -1379,7 +1381,19 @@ async fn execute_direct_sync_runtime_candidate(
|
||||
candidate_started_unix_ms,
|
||||
event.status_code,
|
||||
event.ttfb_ms,
|
||||
)
|
||||
);
|
||||
spawn_local_oauth_success_effect(
|
||||
state_for_response_started.clone(),
|
||||
plan,
|
||||
report_context,
|
||||
LocalOAuthSuccessEffect {
|
||||
status_code: event.status_code,
|
||||
request_started_at_unix_ms: Some(
|
||||
event.response_observation.request_started_at_unix_ms,
|
||||
),
|
||||
request_order_id: Some(&event.response_observation.request_order_id),
|
||||
},
|
||||
);
|
||||
})
|
||||
.await
|
||||
.map_err(SyncExecutionFailure::from_transport);
|
||||
@@ -1478,17 +1492,31 @@ async fn execute_openai_image_sync_upstream_sse_candidate(
|
||||
progress_snapshot: Option<Arc<Mutex<OpenAiImageSyncProgressSnapshot>>>,
|
||||
) -> Result<ExecutionResult, SyncExecutionFailure> {
|
||||
let request_body = build_request_body(plan).map_err(SyncExecutionFailure::from_transport)?;
|
||||
let response_body_limit_bytes = execution_plan_response_body_limit_bytes(plan);
|
||||
let started_at = Instant::now();
|
||||
let mut progress =
|
||||
OpenAiImageSyncProgressRecorder::new(state, plan, report_context, progress_snapshot);
|
||||
progress.record_connecting().await;
|
||||
|
||||
let request_started_at_unix_ms = current_request_candidate_unix_ms();
|
||||
let request_order_id = uuid::Uuid::now_v7().to_string();
|
||||
let response = send_request(plan, request_body)
|
||||
.await
|
||||
.map_err(SyncExecutionFailure::from_transport)?;
|
||||
let ttfb_ms = started_at.elapsed().as_millis() as u64;
|
||||
let response_headers_observed_at_unix_ms = current_request_candidate_unix_ms();
|
||||
let status_code = response.status_code();
|
||||
let headers = response.headers();
|
||||
spawn_local_oauth_success_effect(
|
||||
state.clone(),
|
||||
plan,
|
||||
report_context,
|
||||
LocalOAuthSuccessEffect {
|
||||
status_code,
|
||||
request_started_at_unix_ms: Some(request_started_at_unix_ms),
|
||||
request_order_id: Some(&request_order_id),
|
||||
},
|
||||
);
|
||||
progress.record_response_started(status_code, ttfb_ms).await;
|
||||
|
||||
let mut body_bytes = Vec::new();
|
||||
@@ -1503,8 +1531,12 @@ async fn execute_openai_image_sync_upstream_sse_candidate(
|
||||
),
|
||||
)
|
||||
})?;
|
||||
append_upstream_response_body_chunk(&mut body_bytes, &chunk)
|
||||
.map_err(SyncExecutionFailure::from_transport)?;
|
||||
append_upstream_response_body_chunk_with_limit(
|
||||
&mut body_bytes,
|
||||
&chunk,
|
||||
response_body_limit_bytes,
|
||||
)
|
||||
.map_err(SyncExecutionFailure::from_transport)?;
|
||||
let elapsed_ms = started_at.elapsed().as_millis() as u64;
|
||||
progress
|
||||
.observe_chunk(&chunk, status_code, elapsed_ms)
|
||||
@@ -1521,8 +1553,12 @@ async fn execute_openai_image_sync_upstream_sse_candidate(
|
||||
)),
|
||||
)
|
||||
})?;
|
||||
append_upstream_response_body_chunk(&mut body_bytes, &chunk)
|
||||
.map_err(SyncExecutionFailure::from_transport)?;
|
||||
append_upstream_response_body_chunk_with_limit(
|
||||
&mut body_bytes,
|
||||
&chunk,
|
||||
response_body_limit_bytes,
|
||||
)
|
||||
.map_err(SyncExecutionFailure::from_transport)?;
|
||||
let elapsed_ms = started_at.elapsed().as_millis() as u64;
|
||||
progress
|
||||
.observe_chunk(&chunk, status_code, elapsed_ms)
|
||||
@@ -1539,8 +1575,12 @@ async fn execute_openai_image_sync_upstream_sse_candidate(
|
||||
),
|
||||
)
|
||||
})?;
|
||||
append_upstream_response_body_chunk(&mut body_bytes, &chunk)
|
||||
.map_err(SyncExecutionFailure::from_transport)?;
|
||||
append_upstream_response_body_chunk_with_limit(
|
||||
&mut body_bytes,
|
||||
&chunk,
|
||||
response_body_limit_bytes,
|
||||
)
|
||||
.map_err(SyncExecutionFailure::from_transport)?;
|
||||
let elapsed_ms = started_at.elapsed().as_millis() as u64;
|
||||
progress
|
||||
.observe_chunk(&chunk, status_code, elapsed_ms)
|
||||
@@ -1549,8 +1589,9 @@ async fn execute_openai_image_sync_upstream_sse_candidate(
|
||||
}
|
||||
}
|
||||
|
||||
let decoded_body_bytes = decode_response_body_bytes(&headers, &body_bytes)
|
||||
.map_err(SyncExecutionFailure::from_transport)?;
|
||||
let decoded_body_bytes =
|
||||
decode_response_body_bytes_with_limit(&headers, &body_bytes, response_body_limit_bytes)
|
||||
.map_err(SyncExecutionFailure::from_transport)?;
|
||||
let elapsed_ms = started_at.elapsed().as_millis() as u64;
|
||||
let upstream_bytes = body_bytes.len() as u64;
|
||||
progress.finish(status_code, elapsed_ms).await;
|
||||
@@ -1569,6 +1610,11 @@ async fn execute_openai_image_sync_upstream_sse_candidate(
|
||||
candidate_id: plan.candidate_id.clone(),
|
||||
status_code,
|
||||
headers,
|
||||
response_observation: Some(ExecutionResponseObservation {
|
||||
request_started_at_unix_ms,
|
||||
response_headers_observed_at_unix_ms,
|
||||
request_order_id,
|
||||
}),
|
||||
body,
|
||||
telemetry: Some(ExecutionTelemetry {
|
||||
ttfb_ms: Some(ttfb_ms),
|
||||
@@ -2461,6 +2507,16 @@ async fn execute_execution_runtime_sync_impl(
|
||||
};
|
||||
let mut candidate_first_byte_elapsed_ms =
|
||||
calibrated_sync_candidate_first_byte_elapsed_ms(candidate_started_at, &result);
|
||||
let initial_response_observed_at_unix_ms = current_request_candidate_unix_ms();
|
||||
let mut provider_response_observation =
|
||||
result
|
||||
.response_observation
|
||||
.clone()
|
||||
.unwrap_or(ExecutionResponseObservation {
|
||||
request_started_at_unix_ms: candidate_started_unix_secs,
|
||||
response_headers_observed_at_unix_ms: initial_response_observed_at_unix_ms,
|
||||
request_order_id: uuid::Uuid::now_v7().to_string(),
|
||||
});
|
||||
let mut oauth_retry_attempted = false;
|
||||
let (
|
||||
result_error_type,
|
||||
@@ -2473,6 +2529,18 @@ async fn execute_execution_runtime_sync_impl(
|
||||
local_failover_response_text,
|
||||
local_failover_analysis,
|
||||
) = loop {
|
||||
spawn_local_oauth_success_effect(
|
||||
state.clone(),
|
||||
&plan,
|
||||
report_context.as_ref(),
|
||||
LocalOAuthSuccessEffect {
|
||||
status_code: result.status_code,
|
||||
request_started_at_unix_ms: Some(
|
||||
provider_response_observation.request_started_at_unix_ms,
|
||||
),
|
||||
request_order_id: Some(&provider_response_observation.request_order_id),
|
||||
},
|
||||
);
|
||||
let result_latency_ms = result
|
||||
.telemetry
|
||||
.as_ref()
|
||||
@@ -2534,10 +2602,15 @@ async fn execute_execution_runtime_sync_impl(
|
||||
result.status_code,
|
||||
local_failover_response_text.as_deref(),
|
||||
trace_id,
|
||||
report_context.as_ref(),
|
||||
Some(provider_response_observation.request_started_at_unix_ms),
|
||||
Some(&provider_response_observation.request_order_id),
|
||||
)
|
||||
.await
|
||||
{
|
||||
oauth_retry_attempted = true;
|
||||
let retry_started_at_unix_ms = current_request_candidate_unix_ms();
|
||||
let retry_request_order_id = uuid::Uuid::now_v7().to_string();
|
||||
match crate::execution_runtime::execute_execution_runtime_sync_plan(
|
||||
state,
|
||||
Some(trace_id),
|
||||
@@ -2546,6 +2619,16 @@ async fn execute_execution_runtime_sync_impl(
|
||||
.await
|
||||
{
|
||||
Ok(retry_result) => {
|
||||
let retry_response_observed_at_unix_ms = current_request_candidate_unix_ms();
|
||||
provider_response_observation = retry_result
|
||||
.response_observation
|
||||
.clone()
|
||||
.unwrap_or(ExecutionResponseObservation {
|
||||
request_started_at_unix_ms: retry_started_at_unix_ms,
|
||||
response_headers_observed_at_unix_ms:
|
||||
retry_response_observed_at_unix_ms,
|
||||
request_order_id: retry_request_order_id,
|
||||
});
|
||||
candidate_first_byte_elapsed_ms =
|
||||
calibrated_sync_candidate_first_byte_elapsed_ms(
|
||||
candidate_started_at,
|
||||
@@ -2594,6 +2677,13 @@ async fn execute_execution_runtime_sync_impl(
|
||||
local_failover_analysis,
|
||||
);
|
||||
};
|
||||
let mut report_context = attach_provider_response_headers_to_report_context(
|
||||
report_context,
|
||||
&headers,
|
||||
provider_response_observation.request_started_at_unix_ms,
|
||||
provider_response_observation.response_headers_observed_at_unix_ms,
|
||||
&provider_response_observation.request_order_id,
|
||||
);
|
||||
if result.status_code >= 400 {
|
||||
apply_local_execution_effect(
|
||||
state,
|
||||
@@ -2739,8 +2829,6 @@ async fn execute_execution_runtime_sync_impl(
|
||||
}
|
||||
let status_code = result.status_code;
|
||||
let has_body_bytes = body_base64.is_some();
|
||||
let mut report_context =
|
||||
attach_provider_response_headers_to_report_context(report_context, &headers);
|
||||
if (200..300).contains(&status_code) {
|
||||
seed_kiro_sync_simulated_cache_enabled(state, &plan, &mut report_context).await;
|
||||
if kiro_simulated_cache_enabled_from_report_context(report_context.as_ref()) {
|
||||
@@ -3231,6 +3319,8 @@ async fn execute_sync_via_remote_execution_runtime(
|
||||
candidate_started_unix_secs: u64,
|
||||
candidate_started_at: Instant,
|
||||
) -> Result<RemoteSyncFallbackOutcome, GatewayError> {
|
||||
let remote_request_started_at_unix_ms = current_request_candidate_unix_ms();
|
||||
let remote_request_order_id = uuid::Uuid::now_v7().to_string();
|
||||
let response = match post_sync_plan_to_remote_execution_runtime(
|
||||
state,
|
||||
remote_execution_runtime_base_url,
|
||||
@@ -3299,11 +3389,19 @@ async fn execute_sync_via_remote_execution_runtime(
|
||||
));
|
||||
}
|
||||
|
||||
response
|
||||
.json()
|
||||
let remote_response_observed_at_unix_ms = current_request_candidate_unix_ms();
|
||||
let mut result = response
|
||||
.json::<ExecutionResult>()
|
||||
.await
|
||||
.map(RemoteSyncFallbackOutcome::Executed)
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
result
|
||||
.response_observation
|
||||
.get_or_insert(ExecutionResponseObservation {
|
||||
request_started_at_unix_ms: remote_request_started_at_unix_ms,
|
||||
response_headers_observed_at_unix_ms: remote_response_observed_at_unix_ms,
|
||||
request_order_id: remote_request_order_id,
|
||||
});
|
||||
Ok(RemoteSyncFallbackOutcome::Executed(result))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
|
||||
@@ -9,18 +9,19 @@ use std::sync::{Arc, LazyLock, Mutex as StdMutex, OnceLock, RwLock as StdRwLock}
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
use aether_contracts::{
|
||||
ExecutionPlan, ExecutionResponseBodyMode, ExecutionResult, ExecutionTelemetry, ProxySnapshot,
|
||||
ResolvedTransportProfile, ResponseBody, EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER,
|
||||
EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER, EXECUTION_REQUEST_HTTP1_ONLY_HEADER,
|
||||
EXECUTION_RESPONSE_BODY_MODE_HEADER, TRANSPORT_BACKEND_BROWSER_WREQ,
|
||||
TRANSPORT_BACKEND_REQWEST_RUSTLS, TRANSPORT_HTTP_MODE_H2C_PRIOR_KNOWLEDGE,
|
||||
TRANSPORT_HTTP_MODE_HTTP1_ONLY,
|
||||
ExecutionPlan, ExecutionResponseBodyMode, ExecutionResponseObservation, ExecutionResult,
|
||||
ExecutionTelemetry, ProxySnapshot, ResolvedTransportProfile, ResponseBody,
|
||||
EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER, EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER,
|
||||
EXECUTION_REQUEST_HTTP1_ONLY_HEADER, EXECUTION_RESPONSE_BODY_MODE_HEADER,
|
||||
TRANSPORT_BACKEND_BROWSER_WREQ, TRANSPORT_BACKEND_REQWEST_RUSTLS,
|
||||
TRANSPORT_HTTP_MODE_H2C_PRIOR_KNOWLEDGE, TRANSPORT_HTTP_MODE_HTTP1_ONLY,
|
||||
};
|
||||
use aether_data::repository::proxy_nodes::ProxyNodeTrafficMutation;
|
||||
use aether_http::{apply_http_client_config, HttpClientConfig};
|
||||
use aether_runtime::{MetricKind, MetricSample};
|
||||
use axum::body::Bytes;
|
||||
use base64::Engine as _;
|
||||
use brotli::Decompressor as BrotliDecoder;
|
||||
use flate2::read::{DeflateDecoder, GzDecoder};
|
||||
use flate2::write::GzEncoder;
|
||||
use flate2::Compression;
|
||||
@@ -62,6 +63,10 @@ const DEFAULT_STREAM_FIRST_BYTE_TIMEOUT_MS: u64 = 30_000;
|
||||
const DEFAULT_NON_STREAM_TOTAL_TIMEOUT_MS: u64 = 300_000;
|
||||
const DEFAULT_CODEX_COMPACT_TOTAL_TIMEOUT_MS: u64 = 1_200_000;
|
||||
const MIN_TUNNEL_TIMEOUT_SECS: u64 = 1;
|
||||
const EXECUTION_RESPONSE_BODY_LIMIT_HEADER: &str = "x-aether-execution-response-body-limit-bytes";
|
||||
const DEFAULT_SCOPED_RESPONSE_BODY_LIMIT_BYTES: usize = 8 * 1024 * 1024;
|
||||
const MIN_SCOPED_RESPONSE_BODY_LIMIT_BYTES: usize = 64 * 1024;
|
||||
const MAX_SCOPED_RESPONSE_BODY_LIMIT_BYTES: usize = 64 * 1024 * 1024;
|
||||
const DIRECT_REQWEST_H2_CLIENT_SHARDS_ENV: &str = "AETHER_GATEWAY_DIRECT_REQWEST_H2_CLIENT_SHARDS";
|
||||
const DIRECT_REQWEST_CLIENT_SHARDS_ENV: &str = "AETHER_GATEWAY_DIRECT_REQWEST_CLIENT_SHARDS";
|
||||
const DIRECT_REQWEST_H2_TARGET_STREAMS_PER_CLIENT_ENV: &str =
|
||||
@@ -607,6 +612,56 @@ impl std::fmt::Display for UpstreamResponseBodyPhase {
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn with_upstream_response_body_limit(
|
||||
plan: &ExecutionPlan,
|
||||
limit_bytes: usize,
|
||||
) -> ExecutionPlan {
|
||||
let mut bounded_plan = plan.clone();
|
||||
bounded_plan
|
||||
.headers
|
||||
.retain(|name, _| !name.eq_ignore_ascii_case(EXECUTION_RESPONSE_BODY_LIMIT_HEADER));
|
||||
bounded_plan.headers.insert(
|
||||
EXECUTION_RESPONSE_BODY_LIMIT_HEADER.to_string(),
|
||||
normalize_scoped_response_body_limit(limit_bytes)
|
||||
.unwrap_or(DEFAULT_SCOPED_RESPONSE_BODY_LIMIT_BYTES)
|
||||
.to_string(),
|
||||
);
|
||||
bounded_plan
|
||||
}
|
||||
|
||||
pub(crate) fn execution_plan_response_body_limit_bytes(plan: &ExecutionPlan) -> usize {
|
||||
effective_response_body_limit_bytes(
|
||||
execution_transport_header_value(&plan.headers, EXECUTION_RESPONSE_BODY_LIMIT_HEADER),
|
||||
crate::headers::max_internal_buffered_body_bytes(),
|
||||
)
|
||||
}
|
||||
|
||||
fn effective_response_body_limit_bytes(
|
||||
raw_scoped_limit: Option<&str>,
|
||||
global_limit: usize,
|
||||
) -> usize {
|
||||
let Some(raw_scoped_limit) = raw_scoped_limit else {
|
||||
return global_limit;
|
||||
};
|
||||
parse_scoped_response_body_limit(raw_scoped_limit)
|
||||
.unwrap_or(DEFAULT_SCOPED_RESPONSE_BODY_LIMIT_BYTES)
|
||||
.min(global_limit)
|
||||
}
|
||||
|
||||
fn parse_scoped_response_body_limit(value: &str) -> Option<usize> {
|
||||
let raw_limit = value.trim().parse::<u64>().ok()?;
|
||||
usize::try_from(raw_limit)
|
||||
.ok()
|
||||
.and_then(normalize_scoped_response_body_limit)
|
||||
}
|
||||
|
||||
fn normalize_scoped_response_body_limit(limit_bytes: usize) -> Option<usize> {
|
||||
(limit_bytes > 0).then_some(limit_bytes.clamp(
|
||||
MIN_SCOPED_RESPONSE_BODY_LIMIT_BYTES,
|
||||
MAX_SCOPED_RESPONSE_BODY_LIMIT_BYTES,
|
||||
))
|
||||
}
|
||||
|
||||
pub(crate) fn append_upstream_response_body_chunk(
|
||||
body: &mut Vec<u8>,
|
||||
chunk: &[u8],
|
||||
@@ -618,7 +673,7 @@ pub(crate) fn append_upstream_response_body_chunk(
|
||||
)
|
||||
}
|
||||
|
||||
fn append_upstream_response_body_chunk_with_limit(
|
||||
pub(crate) fn append_upstream_response_body_chunk_with_limit(
|
||||
body: &mut Vec<u8>,
|
||||
chunk: &[u8],
|
||||
limit_bytes: usize,
|
||||
@@ -691,6 +746,7 @@ pub(crate) struct DirectUpstreamStreamExecution {
|
||||
pub(crate) stream_precommit_committed: bool,
|
||||
pub(crate) response: DirectUpstreamResponse,
|
||||
pub(crate) started_at: Instant,
|
||||
pub(crate) response_observation: ExecutionResponseObservation,
|
||||
pub(crate) stream_first_byte_timeout: Option<Duration>,
|
||||
pub(crate) upstream_target_permit: Option<UpstreamTargetAdmissionPermit>,
|
||||
}
|
||||
@@ -699,6 +755,7 @@ pub(crate) struct DirectUpstreamStreamExecution {
|
||||
pub(crate) struct DirectSyncResponseStarted {
|
||||
pub(crate) status_code: u16,
|
||||
pub(crate) ttfb_ms: u64,
|
||||
pub(crate) response_observation: ExecutionResponseObservation,
|
||||
}
|
||||
|
||||
impl DirectSyncExecutionRuntime {
|
||||
@@ -722,20 +779,35 @@ impl DirectSyncExecutionRuntime {
|
||||
F: FnOnce(DirectSyncResponseStarted),
|
||||
{
|
||||
let body_bytes = build_request_body(plan)?;
|
||||
let response_body_limit_bytes = execution_plan_response_body_limit_bytes(plan);
|
||||
|
||||
let started_at = Instant::now();
|
||||
let request_started_at_unix_ms = crate::clock::current_unix_ms();
|
||||
let request_order_id = uuid::Uuid::now_v7().to_string();
|
||||
with_non_stream_total_timeout(plan, async move {
|
||||
let response = send_request_inner(plan, body_bytes, false).await?;
|
||||
let ttfb_ms = started_at.elapsed().as_millis() as u64;
|
||||
let response_headers_observed_at_unix_ms = crate::clock::current_unix_ms();
|
||||
let status_code = response.status_code();
|
||||
let headers = response.headers();
|
||||
let response_observation = ExecutionResponseObservation {
|
||||
request_started_at_unix_ms,
|
||||
response_headers_observed_at_unix_ms,
|
||||
request_order_id,
|
||||
};
|
||||
on_response_started(DirectSyncResponseStarted {
|
||||
status_code,
|
||||
ttfb_ms,
|
||||
response_observation: response_observation.clone(),
|
||||
});
|
||||
let (body_bytes, stream_ttfb_ms) =
|
||||
response.bytes_with_stream_timeout(plan, started_at).await?;
|
||||
let decoded_body_bytes = decode_response_body_bytes(&headers, &body_bytes)?;
|
||||
let (body_bytes, stream_ttfb_ms) = response
|
||||
.bytes_with_stream_timeout(plan, started_at, response_body_limit_bytes)
|
||||
.await?;
|
||||
let decoded_body_bytes = decode_response_body_bytes_with_limit(
|
||||
&headers,
|
||||
&body_bytes,
|
||||
response_body_limit_bytes,
|
||||
)?;
|
||||
let elapsed_ms = started_at.elapsed().as_millis() as u64;
|
||||
let upstream_bytes = body_bytes.len() as u64;
|
||||
|
||||
@@ -752,6 +824,7 @@ impl DirectSyncExecutionRuntime {
|
||||
candidate_id: plan.candidate_id.clone(),
|
||||
status_code,
|
||||
headers,
|
||||
response_observation: Some(response_observation),
|
||||
body,
|
||||
telemetry: Some(ExecutionTelemetry {
|
||||
ttfb_ms: stream_ttfb_ms.or(Some(ttfb_ms)),
|
||||
@@ -776,6 +849,8 @@ impl DirectSyncExecutionRuntime {
|
||||
);
|
||||
|
||||
let started_at = Instant::now();
|
||||
let request_started_at_unix_ms = crate::clock::current_unix_ms();
|
||||
let request_order_id = uuid::Uuid::now_v7().to_string();
|
||||
let response = send_request(plan, body_bytes).await?;
|
||||
observe_gateway_stage_ms(
|
||||
"direct_send_headers",
|
||||
@@ -783,6 +858,7 @@ impl DirectSyncExecutionRuntime {
|
||||
);
|
||||
let status_code = response.status_code();
|
||||
let headers = response.headers();
|
||||
let response_headers_observed_at_unix_ms = crate::clock::current_unix_ms();
|
||||
|
||||
let stream_summary_report_context = build_stream_summary_report_context(plan);
|
||||
|
||||
@@ -797,6 +873,11 @@ impl DirectSyncExecutionRuntime {
|
||||
stream_precommit_committed: false,
|
||||
response: response.into_direct_upstream_response(),
|
||||
started_at,
|
||||
response_observation: ExecutionResponseObservation {
|
||||
request_started_at_unix_ms,
|
||||
response_headers_observed_at_unix_ms,
|
||||
request_order_id,
|
||||
},
|
||||
stream_first_byte_timeout: resolve_stream_first_byte_timeout(plan),
|
||||
upstream_target_permit: None,
|
||||
})
|
||||
@@ -834,7 +915,7 @@ pub(crate) async fn execute_sync_plan_with_report_context(
|
||||
}
|
||||
|
||||
if resolve_local_tunnel_node_id(state, plan.proxy.as_ref()).is_some() {
|
||||
return execute_sync_plan_via_local_tunnel(state, plan)
|
||||
return execute_sync_plan_via_local_tunnel(state, plan, report_context)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()));
|
||||
}
|
||||
@@ -857,7 +938,24 @@ pub(crate) async fn execute_sync_plan_with_report_context(
|
||||
Ok(None) => {}
|
||||
Err(err) => return Err(GatewayError::Internal(err.to_string())),
|
||||
}
|
||||
match DirectSyncExecutionRuntime::new().execute_sync(plan).await {
|
||||
let state_for_response_started = state.clone();
|
||||
match DirectSyncExecutionRuntime::new()
|
||||
.execute_sync_with_response_started(plan, move |event| {
|
||||
crate::orchestration::spawn_local_oauth_success_effect(
|
||||
state_for_response_started,
|
||||
plan,
|
||||
report_context,
|
||||
crate::orchestration::LocalOAuthSuccessEffect {
|
||||
status_code: event.status_code,
|
||||
request_started_at_unix_ms: Some(
|
||||
event.response_observation.request_started_at_unix_ms,
|
||||
),
|
||||
request_order_id: Some(&event.response_observation.request_order_id),
|
||||
},
|
||||
);
|
||||
})
|
||||
.await
|
||||
{
|
||||
Ok(result) => {
|
||||
record_manual_proxy_request_outcome(state, plan, result.status_code).await;
|
||||
Ok(result)
|
||||
@@ -889,6 +987,8 @@ pub(crate) async fn execute_stream_plan_via_local_tunnel(
|
||||
plan.body.body_bytes_b64.is_some(),
|
||||
)?;
|
||||
let started_at = Instant::now();
|
||||
let request_started_at_unix_ms = crate::clock::current_unix_ms();
|
||||
let request_order_id = uuid::Uuid::now_v7().to_string();
|
||||
let response = state
|
||||
.tunnel
|
||||
.open_direct_relay_stream(
|
||||
@@ -900,6 +1000,7 @@ pub(crate) async fn execute_stream_plan_via_local_tunnel(
|
||||
.map_err(ExecutionRuntimeTransportError::RelayError)?;
|
||||
let status_code = response.status();
|
||||
let headers = collect_tunnel_response_headers(response.headers());
|
||||
let response_headers_observed_at_unix_ms = crate::clock::current_unix_ms();
|
||||
|
||||
Ok(Some(DirectUpstreamStreamExecution {
|
||||
request_id: plan.request_id.clone(),
|
||||
@@ -912,6 +1013,11 @@ pub(crate) async fn execute_stream_plan_via_local_tunnel(
|
||||
stream_precommit_committed: false,
|
||||
response: DirectUpstreamResponse::LocalTunnel(response),
|
||||
started_at,
|
||||
response_observation: ExecutionResponseObservation {
|
||||
request_started_at_unix_ms,
|
||||
response_headers_observed_at_unix_ms,
|
||||
request_order_id,
|
||||
},
|
||||
stream_first_byte_timeout: resolve_stream_first_byte_timeout(plan),
|
||||
upstream_target_permit: None,
|
||||
}))
|
||||
@@ -991,13 +1097,19 @@ fn manual_proxy_node_id(proxy: Option<&ProxySnapshot>) -> Option<String> {
|
||||
async fn execute_sync_plan_via_local_tunnel(
|
||||
state: &AppState,
|
||||
plan: &ExecutionPlan,
|
||||
report_context: Option<&serde_json::Value>,
|
||||
) -> Result<ExecutionResult, ExecutionRuntimeTransportError> {
|
||||
with_non_stream_total_timeout(plan, execute_sync_plan_via_local_tunnel_inner(state, plan)).await
|
||||
with_non_stream_total_timeout(
|
||||
plan,
|
||||
execute_sync_plan_via_local_tunnel_inner(state, plan, report_context),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn execute_sync_plan_via_local_tunnel_inner(
|
||||
state: &AppState,
|
||||
plan: &ExecutionPlan,
|
||||
report_context: Option<&serde_json::Value>,
|
||||
) -> Result<ExecutionResult, ExecutionRuntimeTransportError> {
|
||||
let node_id = resolve_local_tunnel_node_id(state, plan.proxy.as_ref()).ok_or_else(|| {
|
||||
ExecutionRuntimeTransportError::RelayError("local tunnel node unavailable".to_string())
|
||||
@@ -1007,6 +1119,7 @@ async fn execute_sync_plan_via_local_tunnel_inner(
|
||||
}
|
||||
|
||||
let body_bytes = build_request_body(plan)?;
|
||||
let response_body_limit_bytes = execution_plan_response_body_limit_bytes(plan);
|
||||
let transport_controls = resolve_execution_transport_controls(&plan.headers);
|
||||
let headers = build_request_headers(
|
||||
&plan.headers,
|
||||
@@ -1030,6 +1143,8 @@ async fn execute_sync_plan_via_local_tunnel_inner(
|
||||
"gateway execution runtime local tunnel request prepared"
|
||||
);
|
||||
let started_at = Instant::now();
|
||||
let request_started_at_unix_ms = crate::clock::current_unix_ms();
|
||||
let request_order_id = uuid::Uuid::now_v7().to_string();
|
||||
let mut response = state
|
||||
.tunnel
|
||||
.open_direct_relay_stream(
|
||||
@@ -1040,12 +1155,30 @@ async fn execute_sync_plan_via_local_tunnel_inner(
|
||||
.await
|
||||
.map_err(ExecutionRuntimeTransportError::RelayError)?;
|
||||
let ttfb_ms = started_at.elapsed().as_millis() as u64;
|
||||
let response_headers_observed_at_unix_ms = crate::clock::current_unix_ms();
|
||||
let status_code = response.status();
|
||||
let headers = collect_tunnel_response_headers(response.headers());
|
||||
let response_observation = ExecutionResponseObservation {
|
||||
request_started_at_unix_ms,
|
||||
response_headers_observed_at_unix_ms,
|
||||
request_order_id,
|
||||
};
|
||||
crate::orchestration::spawn_local_oauth_success_effect(
|
||||
state.clone(),
|
||||
plan,
|
||||
report_context,
|
||||
crate::orchestration::LocalOAuthSuccessEffect {
|
||||
status_code,
|
||||
request_started_at_unix_ms: Some(response_observation.request_started_at_unix_ms),
|
||||
request_order_id: Some(&response_observation.request_order_id),
|
||||
},
|
||||
);
|
||||
let proxy_timing = execution_header_for_log(&headers, "x-proxy-timing").unwrap_or("-");
|
||||
let (body_bytes, stream_ttfb_ms) =
|
||||
collect_local_tunnel_response_body(response, plan, started_at).await?;
|
||||
let decoded_body_bytes = decode_response_body_bytes(&headers, &body_bytes)?;
|
||||
collect_local_tunnel_response_body(response, plan, started_at, response_body_limit_bytes)
|
||||
.await?;
|
||||
let decoded_body_bytes =
|
||||
decode_response_body_bytes_with_limit(&headers, &body_bytes, response_body_limit_bytes)?;
|
||||
let elapsed_ms = started_at.elapsed().as_millis() as u64;
|
||||
let upstream_bytes = body_bytes.len() as u64;
|
||||
if status_code >= 400 {
|
||||
@@ -1095,6 +1228,7 @@ async fn execute_sync_plan_via_local_tunnel_inner(
|
||||
candidate_id: plan.candidate_id.clone(),
|
||||
status_code,
|
||||
headers,
|
||||
response_observation: Some(response_observation),
|
||||
body,
|
||||
telemetry: Some(ExecutionTelemetry {
|
||||
ttfb_ms: stream_ttfb_ms.or(Some(ttfb_ms)),
|
||||
@@ -1109,6 +1243,7 @@ async fn collect_local_tunnel_response_body(
|
||||
mut response: tunnel::DirectRelayResponse,
|
||||
plan: &ExecutionPlan,
|
||||
started_at: Instant,
|
||||
response_body_limit_bytes: usize,
|
||||
) -> Result<(Vec<u8>, Option<u64>), ExecutionRuntimeTransportError> {
|
||||
let mut body_bytes = Vec::new();
|
||||
let mut first_byte_ms = None;
|
||||
@@ -1131,7 +1266,11 @@ async fn collect_local_tunnel_response_body(
|
||||
if plan.stream && first_byte_ms.is_none() && !chunk.is_empty() {
|
||||
first_byte_ms = Some(started_at.elapsed().as_millis() as u64);
|
||||
}
|
||||
append_upstream_response_body_chunk(&mut body_bytes, &chunk)?;
|
||||
append_upstream_response_body_chunk_with_limit(
|
||||
&mut body_bytes,
|
||||
&chunk,
|
||||
response_body_limit_bytes,
|
||||
)?;
|
||||
}
|
||||
|
||||
Ok((body_bytes, first_byte_ms))
|
||||
@@ -1292,20 +1431,28 @@ impl DirectHttpResponse {
|
||||
}
|
||||
|
||||
pub(crate) async fn bytes(self) -> Result<Bytes, ExecutionRuntimeTransportError> {
|
||||
self.bytes_with_limit(crate::headers::max_internal_buffered_body_bytes())
|
||||
.await
|
||||
}
|
||||
|
||||
async fn bytes_with_limit(
|
||||
self,
|
||||
response_body_limit_bytes: usize,
|
||||
) -> Result<Bytes, ExecutionRuntimeTransportError> {
|
||||
let started_at = Instant::now();
|
||||
match self {
|
||||
DirectHttpResponse::Reqwest(response) => {
|
||||
collect_reqwest_stream_body(response, started_at, None)
|
||||
collect_reqwest_stream_body(response, started_at, None, response_body_limit_bytes)
|
||||
.await
|
||||
.map(|(body, _)| body)
|
||||
}
|
||||
DirectHttpResponse::HyperH2c(response) => {
|
||||
collect_hyper_stream_body(response, started_at, None)
|
||||
collect_hyper_stream_body(response, started_at, None, response_body_limit_bytes)
|
||||
.await
|
||||
.map(|(body, _)| body)
|
||||
}
|
||||
DirectHttpResponse::BrowserWreq(response) => {
|
||||
collect_wreq_stream_body(response, started_at, None)
|
||||
collect_wreq_stream_body(response, started_at, None, response_body_limit_bytes)
|
||||
.await
|
||||
.map(|(body, _)| body)
|
||||
}
|
||||
@@ -1316,21 +1463,43 @@ impl DirectHttpResponse {
|
||||
self,
|
||||
plan: &ExecutionPlan,
|
||||
started_at: Instant,
|
||||
response_body_limit_bytes: usize,
|
||||
) -> Result<(Bytes, Option<u64>), ExecutionRuntimeTransportError> {
|
||||
if !plan.stream {
|
||||
return self.bytes().await.map(|bytes| (bytes, None));
|
||||
return self
|
||||
.bytes_with_limit(response_body_limit_bytes)
|
||||
.await
|
||||
.map(|bytes| (bytes, None));
|
||||
}
|
||||
|
||||
let first_byte_timeout = resolve_stream_first_byte_timeout(plan);
|
||||
match self {
|
||||
DirectHttpResponse::Reqwest(response) => {
|
||||
collect_reqwest_stream_body(response, started_at, first_byte_timeout).await
|
||||
collect_reqwest_stream_body(
|
||||
response,
|
||||
started_at,
|
||||
first_byte_timeout,
|
||||
response_body_limit_bytes,
|
||||
)
|
||||
.await
|
||||
}
|
||||
DirectHttpResponse::HyperH2c(response) => {
|
||||
collect_hyper_stream_body(response, started_at, first_byte_timeout).await
|
||||
collect_hyper_stream_body(
|
||||
response,
|
||||
started_at,
|
||||
first_byte_timeout,
|
||||
response_body_limit_bytes,
|
||||
)
|
||||
.await
|
||||
}
|
||||
DirectHttpResponse::BrowserWreq(response) => {
|
||||
collect_wreq_stream_body(response, started_at, first_byte_timeout).await
|
||||
collect_wreq_stream_body(
|
||||
response,
|
||||
started_at,
|
||||
first_byte_timeout,
|
||||
response_body_limit_bytes,
|
||||
)
|
||||
.await
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -1376,6 +1545,7 @@ async fn collect_reqwest_stream_body(
|
||||
response: reqwest::Response,
|
||||
started_at: Instant,
|
||||
first_byte_timeout: Option<Duration>,
|
||||
response_body_limit_bytes: usize,
|
||||
) -> Result<(Bytes, Option<u64>), ExecutionRuntimeTransportError> {
|
||||
let mut stream = response.bytes_stream();
|
||||
let mut body_bytes = Vec::new();
|
||||
@@ -1396,7 +1566,11 @@ async fn collect_reqwest_stream_body(
|
||||
if first_byte_ms.is_none() && !chunk.is_empty() {
|
||||
first_byte_ms = Some(started_at.elapsed().as_millis() as u64);
|
||||
}
|
||||
append_upstream_response_body_chunk(&mut body_bytes, &chunk)?;
|
||||
append_upstream_response_body_chunk_with_limit(
|
||||
&mut body_bytes,
|
||||
&chunk,
|
||||
response_body_limit_bytes,
|
||||
)?;
|
||||
}
|
||||
|
||||
Ok((Bytes::from(body_bytes), first_byte_ms))
|
||||
@@ -1406,6 +1580,7 @@ async fn collect_hyper_stream_body(
|
||||
response: hyper::Response<HyperIncomingBody>,
|
||||
started_at: Instant,
|
||||
first_byte_timeout: Option<Duration>,
|
||||
response_body_limit_bytes: usize,
|
||||
) -> Result<(Bytes, Option<u64>), ExecutionRuntimeTransportError> {
|
||||
let mut stream = response.into_body().into_data_stream();
|
||||
let mut body_bytes = Vec::new();
|
||||
@@ -1426,7 +1601,11 @@ async fn collect_hyper_stream_body(
|
||||
if first_byte_ms.is_none() && !chunk.is_empty() {
|
||||
first_byte_ms = Some(started_at.elapsed().as_millis() as u64);
|
||||
}
|
||||
append_upstream_response_body_chunk(&mut body_bytes, &chunk)?;
|
||||
append_upstream_response_body_chunk_with_limit(
|
||||
&mut body_bytes,
|
||||
&chunk,
|
||||
response_body_limit_bytes,
|
||||
)?;
|
||||
}
|
||||
|
||||
Ok((Bytes::from(body_bytes), first_byte_ms))
|
||||
@@ -1436,6 +1615,7 @@ async fn collect_wreq_stream_body(
|
||||
response: wreq::Response,
|
||||
started_at: Instant,
|
||||
first_byte_timeout: Option<Duration>,
|
||||
response_body_limit_bytes: usize,
|
||||
) -> Result<(Bytes, Option<u64>), ExecutionRuntimeTransportError> {
|
||||
let mut stream = response.bytes_stream();
|
||||
let mut body_bytes = Vec::new();
|
||||
@@ -1456,7 +1636,11 @@ async fn collect_wreq_stream_body(
|
||||
if first_byte_ms.is_none() && !chunk.is_empty() {
|
||||
first_byte_ms = Some(started_at.elapsed().as_millis() as u64);
|
||||
}
|
||||
append_upstream_response_body_chunk(&mut body_bytes, &chunk)?;
|
||||
append_upstream_response_body_chunk_with_limit(
|
||||
&mut body_bytes,
|
||||
&chunk,
|
||||
response_body_limit_bytes,
|
||||
)?;
|
||||
}
|
||||
|
||||
Ok((Bytes::from(body_bytes), first_byte_ms))
|
||||
@@ -2308,10 +2492,31 @@ async fn send_via_tunnel_relay(
|
||||
error_kind = %kind,
|
||||
"gateway execution runtime tunnel relay returned relay error"
|
||||
);
|
||||
let message = response
|
||||
.text()
|
||||
.await
|
||||
.unwrap_or_else(|_| format!("hub relay error: {kind}"));
|
||||
let response_headers = collect_response_headers(response.headers());
|
||||
let response_body_limit_bytes = execution_plan_response_body_limit_bytes(plan);
|
||||
let (wire_body, _) =
|
||||
collect_reqwest_stream_body(response, Instant::now(), None, response_body_limit_bytes)
|
||||
.await
|
||||
.map_err(|error| {
|
||||
ExecutionRuntimeTransportError::RelayError(format!(
|
||||
"hub relay error: {kind}: bounded error body read failed: {error}"
|
||||
))
|
||||
})?;
|
||||
let decoded_body = decode_response_body_bytes_with_limit(
|
||||
&response_headers,
|
||||
&wire_body,
|
||||
response_body_limit_bytes,
|
||||
)
|
||||
.map_err(|error| {
|
||||
ExecutionRuntimeTransportError::RelayError(format!(
|
||||
"hub relay error: {kind}: bounded error body decode failed: {error}"
|
||||
))
|
||||
})?;
|
||||
let message = if decoded_body.is_empty() {
|
||||
format!("hub relay error: {kind}")
|
||||
} else {
|
||||
String::from_utf8_lossy(decoded_body.as_ref()).into_owned()
|
||||
};
|
||||
return Err(ExecutionRuntimeTransportError::RelayError(message));
|
||||
}
|
||||
|
||||
@@ -3833,6 +4038,7 @@ pub(crate) fn build_request_headers(
|
||||
|| normalized_key == EXECUTION_REQUEST_HTTP1_ONLY_HEADER
|
||||
|| normalized_key == EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER
|
||||
|| normalized_key == EXECUTION_RESPONSE_BODY_MODE_HEADER
|
||||
|| normalized_key == EXECUTION_RESPONSE_BODY_LIMIT_HEADER
|
||||
{
|
||||
continue;
|
||||
}
|
||||
@@ -3981,7 +4187,7 @@ pub(crate) fn decode_response_body_bytes<'a>(
|
||||
)
|
||||
}
|
||||
|
||||
fn decode_response_body_bytes_with_limit<'a>(
|
||||
pub(crate) fn decode_response_body_bytes_with_limit<'a>(
|
||||
headers: &BTreeMap<String, String>,
|
||||
body_bytes: &'a [u8],
|
||||
limit_bytes: usize,
|
||||
@@ -4003,6 +4209,11 @@ fn decode_response_body_bytes_with_limit<'a>(
|
||||
read_upstream_response_decoder_with_limit("deflate", &mut decoder, limit_bytes)
|
||||
.map(Cow::Owned)
|
||||
}
|
||||
Some("br") => {
|
||||
let mut decoder = BrotliDecoder::new(body_bytes, 4_096);
|
||||
read_upstream_response_decoder_with_limit("br", &mut decoder, limit_bytes)
|
||||
.map(Cow::Owned)
|
||||
}
|
||||
_ => Ok(Cow::Borrowed(body_bytes)),
|
||||
}
|
||||
}
|
||||
@@ -4122,12 +4333,16 @@ mod tests {
|
||||
use super::{
|
||||
append_upstream_response_body_chunk_with_limit, build_browser_wreq_client, build_client,
|
||||
build_direct_tunnel_request_meta, build_execution_response_body, build_request_headers,
|
||||
decode_response_body_bytes_with_limit, execute_sync_plan, execution_response_body_mode,
|
||||
decode_response_body_bytes_with_limit, effective_response_body_limit_bytes,
|
||||
execute_sync_plan, execution_plan_response_body_limit_bytes, execution_response_body_mode,
|
||||
record_manual_proxy_request_failure, record_manual_proxy_request_outcome,
|
||||
record_manual_proxy_request_success, record_manual_proxy_stream_error,
|
||||
resolve_execution_transport_controls, resolve_non_stream_total_timeout,
|
||||
resolve_stream_first_byte_timeout, response_body_is_json, DirectSyncExecutionRuntime,
|
||||
resolve_stream_first_byte_timeout, response_body_is_json,
|
||||
with_upstream_response_body_limit, DirectSyncExecutionRuntime,
|
||||
ExecutionRuntimeTransportError, ExecutionTransportControls, UpstreamResponseBodyPhase,
|
||||
DEFAULT_SCOPED_RESPONSE_BODY_LIMIT_BYTES, EXECUTION_RESPONSE_BODY_LIMIT_HEADER,
|
||||
MAX_SCOPED_RESPONSE_BODY_LIMIT_BYTES, MIN_SCOPED_RESPONSE_BODY_LIMIT_BYTES,
|
||||
};
|
||||
use crate::constants::{
|
||||
EXECUTION_RUNTIME_LOOP_GUARD_HEADER, EXECUTION_RUNTIME_LOOP_GUARD_VIA_TOKEN,
|
||||
@@ -4182,6 +4397,162 @@ mod tests {
|
||||
assert!(!materialized.contains_key("x-aether-future-control"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn scoped_response_body_limit_injection_preserves_transport_profile_and_extra() {
|
||||
let mut plan = tunnel_timeout_plan(false);
|
||||
let original_profile = ResolvedTransportProfile {
|
||||
profile_id: "existing-profile".into(),
|
||||
backend: TRANSPORT_BACKEND_BROWSER_WREQ.into(),
|
||||
http_mode: TRANSPORT_HTTP_MODE_HTTP1_ONLY.into(),
|
||||
pool_scope: "provider".into(),
|
||||
header_fingerprint: Some(json!({"user_agent": "existing"})),
|
||||
extra: Some(json!({"existing": {"nested": true}})),
|
||||
};
|
||||
plan.transport_profile = Some(original_profile.clone());
|
||||
|
||||
let bounded_plan =
|
||||
with_upstream_response_body_limit(&plan, DEFAULT_SCOPED_RESPONSE_BODY_LIMIT_BYTES);
|
||||
|
||||
assert_eq!(plan.transport_profile, Some(original_profile.clone()));
|
||||
assert_eq!(bounded_plan.transport_profile, Some(original_profile));
|
||||
assert_eq!(
|
||||
bounded_plan
|
||||
.headers
|
||||
.get(EXECUTION_RESPONSE_BODY_LIMIT_HEADER)
|
||||
.and_then(|value| value.parse::<usize>().ok()),
|
||||
Some(DEFAULT_SCOPED_RESPONSE_BODY_LIMIT_BYTES)
|
||||
);
|
||||
assert_eq!(
|
||||
execution_plan_response_body_limit_bytes(&bounded_plan),
|
||||
DEFAULT_SCOPED_RESPONSE_BODY_LIMIT_BYTES
|
||||
);
|
||||
|
||||
let unprofiled_plan = tunnel_timeout_plan(false);
|
||||
let bounded_unprofiled_plan = with_upstream_response_body_limit(
|
||||
&unprofiled_plan,
|
||||
DEFAULT_SCOPED_RESPONSE_BODY_LIMIT_BYTES,
|
||||
);
|
||||
assert!(unprofiled_plan.transport_profile.is_none());
|
||||
assert!(bounded_unprofiled_plan.transport_profile.is_none());
|
||||
assert_eq!(
|
||||
execution_plan_response_body_limit_bytes(&bounded_unprofiled_plan),
|
||||
DEFAULT_SCOPED_RESPONSE_BODY_LIMIT_BYTES
|
||||
);
|
||||
|
||||
let mut shadowed_plan = tunnel_timeout_plan(false);
|
||||
shadowed_plan.headers.insert(
|
||||
EXECUTION_RESPONSE_BODY_LIMIT_HEADER.to_ascii_uppercase(),
|
||||
"65536".to_string(),
|
||||
);
|
||||
let bounded_shadowed_plan = with_upstream_response_body_limit(
|
||||
&shadowed_plan,
|
||||
DEFAULT_SCOPED_RESPONSE_BODY_LIMIT_BYTES,
|
||||
);
|
||||
assert_eq!(
|
||||
bounded_shadowed_plan
|
||||
.headers
|
||||
.keys()
|
||||
.filter(|name| name.eq_ignore_ascii_case(EXECUTION_RESPONSE_BODY_LIMIT_HEADER))
|
||||
.count(),
|
||||
1
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn scoped_response_body_limit_parsing_rejects_invalid_values_and_clamps_bounds() {
|
||||
let scoped_plan = |raw_limit: &str| {
|
||||
let mut plan = tunnel_timeout_plan(false);
|
||||
plan.headers.insert(
|
||||
EXECUTION_RESPONSE_BODY_LIMIT_HEADER.to_string(),
|
||||
raw_limit.to_string(),
|
||||
);
|
||||
plan
|
||||
};
|
||||
|
||||
for invalid in ["0", "-1", "1.5", "", "invalid"] {
|
||||
assert_eq!(
|
||||
execution_plan_response_body_limit_bytes(&scoped_plan(invalid)),
|
||||
DEFAULT_SCOPED_RESPONSE_BODY_LIMIT_BYTES
|
||||
);
|
||||
}
|
||||
assert_eq!(
|
||||
execution_plan_response_body_limit_bytes(&scoped_plan("1")),
|
||||
MIN_SCOPED_RESPONSE_BODY_LIMIT_BYTES
|
||||
);
|
||||
assert_eq!(
|
||||
execution_plan_response_body_limit_bytes(&scoped_plan(
|
||||
&(MAX_SCOPED_RESPONSE_BODY_LIMIT_BYTES as u64 + 1).to_string()
|
||||
)),
|
||||
MAX_SCOPED_RESPONSE_BODY_LIMIT_BYTES
|
||||
);
|
||||
assert_eq!(
|
||||
execution_plan_response_body_limit_bytes(&scoped_plan("1048576")),
|
||||
1_048_576
|
||||
);
|
||||
assert_eq!(
|
||||
effective_response_body_limit_bytes(
|
||||
Some(&(DEFAULT_SCOPED_RESPONSE_BODY_LIMIT_BYTES * 2).to_string()),
|
||||
1024 * 1024,
|
||||
),
|
||||
1024 * 1024,
|
||||
"a scoped limit must never raise the operator's global cap"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn scoped_response_body_wire_limit_rejects_overflow() {
|
||||
let bounded_plan = with_upstream_response_body_limit(
|
||||
&tunnel_timeout_plan(false),
|
||||
MIN_SCOPED_RESPONSE_BODY_LIMIT_BYTES,
|
||||
);
|
||||
let limit_bytes = execution_plan_response_body_limit_bytes(&bounded_plan);
|
||||
let mut body = vec![b'x'; limit_bytes];
|
||||
|
||||
let error =
|
||||
append_upstream_response_body_chunk_with_limit(&mut body, b"overflow", limit_bytes)
|
||||
.expect_err("wire body above the plan-scoped limit should fail");
|
||||
|
||||
assert!(matches!(
|
||||
error,
|
||||
ExecutionRuntimeTransportError::UpstreamResponseTooLarge {
|
||||
phase: UpstreamResponseBodyPhase::Wire,
|
||||
limit_bytes: MIN_SCOPED_RESPONSE_BODY_LIMIT_BYTES,
|
||||
}
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn scoped_response_body_limit_rejects_gzip_bomb_after_wire_check() {
|
||||
let bounded_plan = with_upstream_response_body_limit(
|
||||
&tunnel_timeout_plan(false),
|
||||
MIN_SCOPED_RESPONSE_BODY_LIMIT_BYTES,
|
||||
);
|
||||
let limit_bytes = execution_plan_response_body_limit_bytes(&bounded_plan);
|
||||
let payload = vec![b'x'; limit_bytes + 1];
|
||||
let mut encoder = flate2::write::GzEncoder::new(Vec::new(), flate2::Compression::default());
|
||||
encoder
|
||||
.write_all(&payload)
|
||||
.expect("gzip payload should encode");
|
||||
let encoded = encoder.finish().expect("gzip payload should finish");
|
||||
assert!(encoded.len() < limit_bytes);
|
||||
|
||||
let mut wire_body = Vec::new();
|
||||
append_upstream_response_body_chunk_with_limit(&mut wire_body, &encoded, limit_bytes)
|
||||
.expect("compressed wire body should fit within the plan-scoped limit");
|
||||
let headers = BTreeMap::from([("content-encoding".to_string(), "gzip".to_string())]);
|
||||
|
||||
let error = decode_response_body_bytes_with_limit(&headers, &wire_body, limit_bytes)
|
||||
.expect_err("decoded body above the plan-scoped limit should fail");
|
||||
|
||||
assert!(matches!(
|
||||
error,
|
||||
ExecutionRuntimeTransportError::UpstreamResponseTooLarge {
|
||||
phase: UpstreamResponseBodyPhase::Decoded,
|
||||
limit_bytes: MIN_SCOPED_RESPONSE_BODY_LIMIT_BYTES,
|
||||
}
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn upstream_response_wire_limit_allows_exact_body_and_rejects_next_byte() {
|
||||
let mut body = Vec::new();
|
||||
@@ -5605,6 +5976,8 @@ mod tests {
|
||||
)
|
||||
.await
|
||||
.expect("headers should write");
|
||||
socket.flush().await.expect("headers should flush");
|
||||
tokio::time::sleep(std::time::Duration::from_millis(40)).await;
|
||||
socket
|
||||
.write_all(b"b\r\ndata: one\n\n\r\n")
|
||||
.await
|
||||
@@ -5634,12 +6007,34 @@ mod tests {
|
||||
|
||||
let body = result
|
||||
.body
|
||||
.clone()
|
||||
.and_then(|body| body.body_bytes_b64)
|
||||
.and_then(|body| base64::engine::general_purpose::STANDARD.decode(body).ok())
|
||||
.expect("stream body should be captured as bytes");
|
||||
let body = String::from_utf8(body).expect("stream body should be utf8");
|
||||
assert!(body.contains("data: one"));
|
||||
assert!(body.contains("data: two"));
|
||||
let observation = result
|
||||
.response_observation
|
||||
.expect("stream sync execution should preserve header observation");
|
||||
let telemetry = result
|
||||
.telemetry
|
||||
.expect("stream sync execution should include telemetry");
|
||||
let ttfb_ms = telemetry
|
||||
.ttfb_ms
|
||||
.expect("stream sync execution should measure the first body byte");
|
||||
assert!(
|
||||
observation.response_headers_observed_at_unix_ms
|
||||
>= observation.request_started_at_unix_ms
|
||||
);
|
||||
assert!(
|
||||
observation
|
||||
.response_headers_observed_at_unix_ms
|
||||
.saturating_sub(observation.request_started_at_unix_ms)
|
||||
< ttfb_ms,
|
||||
"header observation must not be derived from body-byte ttfb"
|
||||
);
|
||||
assert!(!observation.request_order_id.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
|
||||
@@ -266,6 +266,7 @@ pub(crate) async fn maybe_execute_windsurf_sync(
|
||||
candidate_id: prepared.candidate_id,
|
||||
status_code: 200,
|
||||
headers: BTreeMap::from([("content-type".to_string(), "application/json".to_string())]),
|
||||
response_observation: None,
|
||||
body: Some(ResponseBody {
|
||||
json_body: Some(body_json),
|
||||
body_bytes_b64: None,
|
||||
@@ -527,6 +528,7 @@ fn build_windsurf_stream_frame_stream(
|
||||
("cache-control".to_string(), "no-cache".to_string()),
|
||||
("content-type".to_string(), "text/event-stream".to_string()),
|
||||
]),
|
||||
response_observation: None,
|
||||
},
|
||||
});
|
||||
|
||||
|
||||
@@ -20,9 +20,11 @@ use crate::ai_serving::LocalExecutionAttemptSource;
|
||||
use crate::clock::current_unix_ms;
|
||||
use crate::control::GatewayControlDecision;
|
||||
use crate::execution_runtime::{
|
||||
build_transport_error_stop_response, execute_execution_runtime_stream_with_retry_scope,
|
||||
acquire_upstream_execution_gate, build_transport_error_stop_response,
|
||||
execute_execution_runtime_stream_with_retry_scope,
|
||||
execute_execution_runtime_sync_with_retry_scope,
|
||||
mark_stream_candidate_watchdog_terminal_started, StreamCandidateWatchdogProgress,
|
||||
UpstreamExecutionGateProvider, UPSTREAM_EXECUTION_GATE_NAME,
|
||||
};
|
||||
use crate::executor::{
|
||||
build_local_execution_exhaustion, mark_deferred_upstream_response, LocalExecutionRequestOutcome,
|
||||
@@ -43,7 +45,6 @@ use crate::stage_metrics::observe_gateway_stage_ms;
|
||||
use crate::{AppState, GatewayError};
|
||||
|
||||
const DEFAULT_STREAM_FIRST_BYTE_WATCHDOG_TIMEOUT_MS: u64 = 30_000;
|
||||
const UPSTREAM_EXECUTION_GATE_NAME: &str = "gateway_upstream_execution";
|
||||
const UPSTREAM_TARGET_GATE_NAME: &str = "gateway_upstream_target";
|
||||
const UPSTREAM_EXECUTION_GATE_HOLD_STREAM_RESPONSE_ENV: &str =
|
||||
"AETHER_GATEWAY_UPSTREAM_EXECUTION_GATE_HOLD_STREAM_RESPONSE";
|
||||
@@ -1612,47 +1613,6 @@ fn hold_response_upstream_execution_permit(
|
||||
Response::from_parts(parts, Body::from_stream(stream))
|
||||
}
|
||||
|
||||
trait UpstreamExecutionGateProvider {
|
||||
fn upstream_execution_gate(&self) -> Option<&aether_runtime::ConcurrencyGate>;
|
||||
fn upstream_execution_gate_queue_budget(&self) -> Duration;
|
||||
}
|
||||
|
||||
impl UpstreamExecutionGateProvider for AppState {
|
||||
fn upstream_execution_gate(&self) -> Option<&aether_runtime::ConcurrencyGate> {
|
||||
self.upstream_execution_gate.as_deref()
|
||||
}
|
||||
|
||||
fn upstream_execution_gate_queue_budget(&self) -> Duration {
|
||||
self.frontdoor_runtime_guards.internal_gate_queue_budget
|
||||
}
|
||||
}
|
||||
|
||||
async fn acquire_upstream_execution_gate(
|
||||
state: &(impl UpstreamExecutionGateProvider + ?Sized),
|
||||
trace_id: &str,
|
||||
) -> Result<Option<ConcurrencyPermit>, GatewayError> {
|
||||
let Some(gate) = state.upstream_execution_gate() else {
|
||||
return Ok(None);
|
||||
};
|
||||
let budget = state.upstream_execution_gate_queue_budget();
|
||||
let gate_wait_started_at = std::time::Instant::now();
|
||||
match timeout(budget, gate.acquire()).await {
|
||||
Ok(Ok(permit)) => {
|
||||
observe_gateway_stage_ms(
|
||||
"upstream_execution_gate_wait",
|
||||
gate_wait_started_at.elapsed().as_millis() as u64,
|
||||
);
|
||||
Ok(Some(permit))
|
||||
}
|
||||
Ok(Err(err)) => Err(GatewayError::Internal(err.to_string())),
|
||||
Err(_) => Err(GatewayError::AdmissionTimeout {
|
||||
trace_id: trace_id.to_string(),
|
||||
gate: UPSTREAM_EXECUTION_GATE_NAME,
|
||||
queue_budget_ms: budget.as_millis() as u64,
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn mark_unused_local_candidate_items<T, FPlan, FContext>(
|
||||
state: &AppState,
|
||||
remaining: Vec<T>,
|
||||
|
||||
@@ -1687,6 +1687,7 @@ mod tests {
|
||||
CONTENT_TYPE.as_str().to_string(),
|
||||
"application/json".to_string(),
|
||||
)]),
|
||||
response_observation: None,
|
||||
body: Some(aether_contracts::ResponseBody {
|
||||
json_body: Some(body_json),
|
||||
body_bytes_b64: None,
|
||||
|
||||
+36
-3
@@ -55,6 +55,33 @@ pub(super) async fn maybe_handle(
|
||||
if idempotency_key.is_empty() {
|
||||
return Ok(Some(bad_request_response("idempotency_key 不能为空")));
|
||||
}
|
||||
if idempotency_key.len() > 256 {
|
||||
return Ok(Some(bad_request_response(
|
||||
"idempotency_key 不能超过 256 个字节",
|
||||
)));
|
||||
}
|
||||
let expected_credential_generation = match payload.expected_credential_generation {
|
||||
serde_json::Value::Null => None,
|
||||
serde_json::Value::String(value) => {
|
||||
let value = value.trim().to_string();
|
||||
if value.is_empty() {
|
||||
return Ok(Some(bad_request_response(
|
||||
"expected_credential_generation 不能为空字符串",
|
||||
)));
|
||||
}
|
||||
if value.len() > 256 {
|
||||
return Ok(Some(bad_request_response(
|
||||
"expected_credential_generation 不能超过 256 个字节",
|
||||
)));
|
||||
}
|
||||
Some(value)
|
||||
}
|
||||
_ => {
|
||||
return Ok(Some(bad_request_response(
|
||||
"expected_credential_generation 必须是字符串或 null",
|
||||
)));
|
||||
}
|
||||
};
|
||||
|
||||
let Some(key) = state
|
||||
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
|
||||
@@ -93,9 +120,15 @@ pub(super) async fn maybe_handle(
|
||||
)));
|
||||
};
|
||||
|
||||
let (status, payload) =
|
||||
consume_codex_reset_credit_locally(state, &provider, &endpoint, key, &idempotency_key)
|
||||
.await?;
|
||||
let (status, payload) = consume_codex_reset_credit_locally(
|
||||
state,
|
||||
&provider,
|
||||
&endpoint,
|
||||
key,
|
||||
&idempotency_key,
|
||||
expected_credential_generation.as_deref(),
|
||||
)
|
||||
.await?;
|
||||
Ok(Some((status, Json(payload)).into_response()))
|
||||
}
|
||||
|
||||
|
||||
@@ -1,7 +1,10 @@
|
||||
use crate::handlers::admin::admin_provider_pool_config;
|
||||
use crate::handlers::admin::provider::shared::paths::admin_update_key_id;
|
||||
use crate::handlers::admin::provider::shared::payloads::AdminProviderKeyUpdatePatch;
|
||||
use crate::handlers::admin::provider::write::keys::admin_provider_key_update_requires_immediate_model_fetch;
|
||||
use crate::handlers::admin::provider::write::keys::{
|
||||
admin_provider_key_update_requires_immediate_model_fetch,
|
||||
build_provider_catalog_key_admin_cas_update,
|
||||
};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::maintenance::ensure_provider_key_pool_scores_for_keys;
|
||||
use crate::provider_key_auth::provider_key_effective_api_formats;
|
||||
@@ -82,7 +85,25 @@ pub(super) async fn maybe_handle(
|
||||
Ok(record) => record,
|
||||
Err(detail) => return Ok(Some(bad_request_response(detail))),
|
||||
};
|
||||
let Some(mut updated) = state.update_provider_catalog_key(&updated_record).await? else {
|
||||
let admin_update = build_provider_catalog_key_admin_cas_update(
|
||||
&existing_key,
|
||||
updated_record.clone(),
|
||||
&provider.provider_type,
|
||||
);
|
||||
if !state
|
||||
.compare_and_update_provider_catalog_key_admin_state(&admin_update)
|
||||
.await?
|
||||
{
|
||||
return Ok(Some(conflict_response(
|
||||
"Key 凭据或配置已被其他请求更新,请刷新后重试",
|
||||
)));
|
||||
}
|
||||
let Some(mut updated) = state
|
||||
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
|
||||
.await?
|
||||
.into_iter()
|
||||
.next()
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
if updated_record.learned_rpm_limit != existing_key.learned_rpm_limit {
|
||||
@@ -183,3 +204,11 @@ fn not_found_response(detail: impl Into<String>) -> Response<Body> {
|
||||
)
|
||||
.into_response()
|
||||
}
|
||||
|
||||
fn conflict_response(detail: impl Into<String>) -> Response<Body> {
|
||||
(
|
||||
http::StatusCode::CONFLICT,
|
||||
Json(json!({ "detail": detail.into() })),
|
||||
)
|
||||
.into_response()
|
||||
}
|
||||
|
||||
@@ -22,16 +22,60 @@ 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 aether_data_contracts::repository::provider_catalog::{
|
||||
ProviderCatalogKeyOAuthRuntimeStateCasUpdate,
|
||||
ProviderCatalogUpstreamMetadataNamespaceExpectation,
|
||||
};
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
http,
|
||||
response::{IntoResponse, Response},
|
||||
Json,
|
||||
};
|
||||
use serde_json::json;
|
||||
use serde_json::{json, Value};
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
const CODEX_OAUTH_COMPLETE_NAMESPACE_CAS_MAX_RETRIES: usize = 3;
|
||||
const CODEX_CREDENTIAL_GENERATION_KEY: &str = "credential_generation";
|
||||
|
||||
#[derive(Debug, PartialEq)]
|
||||
enum CodexOAuthCompleteCasMissAction {
|
||||
AlreadyCompleted,
|
||||
RetryNamespace(Option<Value>),
|
||||
Conflict,
|
||||
}
|
||||
|
||||
fn codex_oauth_complete_cas_miss_action(
|
||||
latest_encrypted_auth_config: Option<&str>,
|
||||
latest_upstream_metadata: Option<&Value>,
|
||||
latest_status_snapshot: Option<&Value>,
|
||||
expected_encrypted_auth_config: Option<&str>,
|
||||
persisted_encrypted_auth_config: &str,
|
||||
expected_codex_metadata_value: Option<&Value>,
|
||||
replacement_codex_metadata_value: &Value,
|
||||
) -> CodexOAuthCompleteCasMissAction {
|
||||
let latest_codex_metadata_value = latest_upstream_metadata
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|metadata| metadata.get("codex"))
|
||||
.cloned();
|
||||
let quota_is_cleared = latest_status_snapshot
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|snapshot| snapshot.get("quota"))
|
||||
== Some(&Value::Null);
|
||||
if latest_encrypted_auth_config == Some(persisted_encrypted_auth_config)
|
||||
&& latest_codex_metadata_value.as_ref() == Some(replacement_codex_metadata_value)
|
||||
&& quota_is_cleared
|
||||
{
|
||||
return CodexOAuthCompleteCasMissAction::AlreadyCompleted;
|
||||
}
|
||||
if latest_encrypted_auth_config != expected_encrypted_auth_config
|
||||
|| latest_codex_metadata_value.as_ref() == expected_codex_metadata_value
|
||||
{
|
||||
return CodexOAuthCompleteCasMissAction::Conflict;
|
||||
}
|
||||
CodexOAuthCompleteCasMissAction::RetryNamespace(latest_codex_metadata_value)
|
||||
}
|
||||
|
||||
pub(super) async fn handle_admin_provider_oauth_complete_key(
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
@@ -277,29 +321,98 @@ pub(super) async fn handle_admin_provider_oauth_complete_key(
|
||||
.and_then(|snapshot| snapshot.get("oauth"))
|
||||
.cloned()
|
||||
.unwrap_or(serde_json::Value::Null);
|
||||
let mut expected_codex_metadata_value = key
|
||||
.upstream_metadata
|
||||
.as_ref()
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.and_then(|metadata| metadata.get("codex"))
|
||||
.cloned();
|
||||
let mut status_snapshot_patch =
|
||||
serde_json::Map::from_iter([("oauth".to_string(), oauth_status)]);
|
||||
if provider_type == "codex" {
|
||||
status_snapshot_patch.insert("quota".to_string(), 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,
|
||||
expected_credential: None,
|
||||
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;
|
||||
let replacement_codex_metadata_value = json!({
|
||||
CODEX_CREDENTIAL_GENERATION_KEY: uuid::Uuid::now_v7().to_string()
|
||||
});
|
||||
let expected_encrypted_auth_config = state_data.expected_encrypted_auth_config.clone();
|
||||
let updated_result: Result<bool, GatewayError> = async {
|
||||
let max_namespace_retries = if provider_type == "codex" {
|
||||
CODEX_OAUTH_COMPLETE_NAMESPACE_CAS_MAX_RETRIES
|
||||
} else {
|
||||
0
|
||||
};
|
||||
for retry in 0..=max_namespace_retries {
|
||||
let updated = state
|
||||
.app()
|
||||
.compare_and_update_provider_catalog_key_oauth_runtime_state(
|
||||
&ProviderCatalogKeyOAuthRuntimeStateCasUpdate {
|
||||
key_id: key_id.clone(),
|
||||
expected_encrypted_auth_config: expected_encrypted_auth_config.clone(),
|
||||
expected_credential: None,
|
||||
expected_upstream_metadata_namespace: (provider_type == "codex").then(
|
||||
|| ProviderCatalogUpstreamMetadataNamespaceExpectation {
|
||||
namespace: "codex".to_string(),
|
||||
expected_value: expected_codex_metadata_value.clone(),
|
||||
},
|
||||
),
|
||||
encrypted_auth_config: persisted_encrypted_auth_config.clone(),
|
||||
encrypted_api_key_update: Some(encrypted_api_key.clone()),
|
||||
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: (provider_type == "codex")
|
||||
.then(|| json!({"codex": replacement_codex_metadata_value.clone()})),
|
||||
upstream_metadata_namespace_to_remove: None,
|
||||
status_snapshot_patch: serde_json::Value::Object(
|
||||
status_snapshot_patch.clone(),
|
||||
),
|
||||
updated_at_unix_secs: Some(now_unix_secs),
|
||||
},
|
||||
)
|
||||
.await?;
|
||||
if updated {
|
||||
return Ok(true);
|
||||
}
|
||||
if provider_type != "codex" {
|
||||
return Ok(false);
|
||||
}
|
||||
|
||||
let Some(latest_key) = state
|
||||
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
|
||||
.await?
|
||||
.into_iter()
|
||||
.next()
|
||||
else {
|
||||
return Ok(false);
|
||||
};
|
||||
match codex_oauth_complete_cas_miss_action(
|
||||
latest_key.encrypted_auth_config.as_deref(),
|
||||
latest_key.upstream_metadata.as_ref(),
|
||||
latest_key.status_snapshot.as_ref(),
|
||||
expected_encrypted_auth_config.as_deref(),
|
||||
&persisted_encrypted_auth_config,
|
||||
expected_codex_metadata_value.as_ref(),
|
||||
&replacement_codex_metadata_value,
|
||||
) {
|
||||
CodexOAuthCompleteCasMissAction::AlreadyCompleted => return Ok(true),
|
||||
CodexOAuthCompleteCasMissAction::Conflict => return Ok(false),
|
||||
CodexOAuthCompleteCasMissAction::RetryNamespace(latest_codex_metadata_value) => {
|
||||
if retry == max_namespace_retries {
|
||||
return Ok(false);
|
||||
}
|
||||
expected_codex_metadata_value = latest_codex_metadata_value;
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(false)
|
||||
}
|
||||
.await;
|
||||
let _ = state
|
||||
.app()
|
||||
.invalidate_local_oauth_refresh_entry(&key_id)
|
||||
@@ -397,3 +510,78 @@ pub(super) async fn handle_admin_provider_oauth_complete_key(
|
||||
}))
|
||||
.into_response())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{codex_oauth_complete_cas_miss_action, CodexOAuthCompleteCasMissAction};
|
||||
use serde_json::json;
|
||||
|
||||
#[test]
|
||||
fn codex_oauth_complete_retries_only_when_namespace_changed() {
|
||||
let expected_codex = json!({"request_id": "old"});
|
||||
let replacement_codex = json!({"credential_generation": "generation-new"});
|
||||
let latest_metadata = json!({
|
||||
"codex": {"request_id": "new"},
|
||||
"unrelated": {"preserved": true}
|
||||
});
|
||||
|
||||
assert_eq!(
|
||||
codex_oauth_complete_cas_miss_action(
|
||||
Some("old-auth"),
|
||||
Some(&latest_metadata),
|
||||
None,
|
||||
Some("old-auth"),
|
||||
"new-auth",
|
||||
Some(&expected_codex),
|
||||
&replacement_codex,
|
||||
),
|
||||
CodexOAuthCompleteCasMissAction::RetryNamespace(Some(json!({
|
||||
"request_id": "new"
|
||||
})))
|
||||
);
|
||||
assert_eq!(
|
||||
codex_oauth_complete_cas_miss_action(
|
||||
Some("old-auth"),
|
||||
Some(&json!({"codex": expected_codex.clone()})),
|
||||
None,
|
||||
Some("old-auth"),
|
||||
"new-auth",
|
||||
Some(&expected_codex),
|
||||
&replacement_codex,
|
||||
),
|
||||
CodexOAuthCompleteCasMissAction::Conflict
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn codex_oauth_complete_accepts_an_ambiguous_success_but_rejects_auth_rotation() {
|
||||
let replacement_codex = json!({"credential_generation": "generation-new"});
|
||||
assert_eq!(
|
||||
codex_oauth_complete_cas_miss_action(
|
||||
Some("new-auth"),
|
||||
Some(&json!({
|
||||
"codex": replacement_codex.clone(),
|
||||
"unrelated": {"preserved": true}
|
||||
})),
|
||||
Some(&json!({"quota": null})),
|
||||
Some("old-auth"),
|
||||
"new-auth",
|
||||
Some(&json!({"request_id": "old"})),
|
||||
&replacement_codex,
|
||||
),
|
||||
CodexOAuthCompleteCasMissAction::AlreadyCompleted
|
||||
);
|
||||
assert_eq!(
|
||||
codex_oauth_complete_cas_miss_action(
|
||||
Some("other-auth"),
|
||||
Some(&json!({"codex": {"request_id": "new"}})),
|
||||
Some(&json!({"quota": null})),
|
||||
Some("old-auth"),
|
||||
"new-auth",
|
||||
Some(&json!({"request_id": "old"})),
|
||||
&replacement_codex,
|
||||
),
|
||||
CodexOAuthCompleteCasMissAction::Conflict
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -12,6 +12,7 @@ use crate::ai_serving::{
|
||||
build_provider_key_pool_score_upsert, provider_key_pool_score_id, provider_key_pool_score_scope,
|
||||
};
|
||||
use crate::handlers::admin::admin_provider_pool_config;
|
||||
use crate::handlers::admin::provider::write::keys::build_provider_catalog_key_admin_cas_update;
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::provider_key_auth::provider_active_api_formats;
|
||||
use crate::GatewayError;
|
||||
@@ -286,6 +287,61 @@ fn grok_oauth_catalog_key_fingerprint(
|
||||
grok_browser_transport_fingerprint_from_auth_config(auth_config)
|
||||
}
|
||||
|
||||
pub(crate) fn rotate_codex_credential_generation(
|
||||
key: &mut StoredProviderCatalogKey,
|
||||
provider_type: &str,
|
||||
) {
|
||||
if !provider_type.trim().eq_ignore_ascii_case("codex") {
|
||||
return;
|
||||
}
|
||||
|
||||
let mut upstream_metadata = key
|
||||
.upstream_metadata
|
||||
.as_ref()
|
||||
.and_then(Value::as_object)
|
||||
.cloned()
|
||||
.unwrap_or_default();
|
||||
upstream_metadata.insert(
|
||||
"codex".to_string(),
|
||||
json!({
|
||||
aether_admin::provider::quota::CODEX_CREDENTIAL_GENERATION_KEY:
|
||||
Uuid::now_v7().to_string(),
|
||||
}),
|
||||
);
|
||||
key.upstream_metadata = Some(Value::Object(upstream_metadata));
|
||||
|
||||
if let Some(mut status_snapshot) = key
|
||||
.status_snapshot
|
||||
.as_ref()
|
||||
.and_then(Value::as_object)
|
||||
.cloned()
|
||||
{
|
||||
status_snapshot.insert("quota".to_string(), Value::Null);
|
||||
key.status_snapshot = Some(Value::Object(status_snapshot));
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn ensure_codex_credential_generation_rotated(
|
||||
key: &mut StoredProviderCatalogKey,
|
||||
provider_type: &str,
|
||||
previous_generation: Option<&str>,
|
||||
) {
|
||||
if !provider_type.trim().eq_ignore_ascii_case("codex") {
|
||||
return;
|
||||
}
|
||||
|
||||
let current_generation = key
|
||||
.upstream_metadata
|
||||
.as_ref()
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|metadata| metadata.get("codex"))
|
||||
.and_then(|codex| aether_admin::provider::quota::codex_credential_generation(Some(codex)));
|
||||
let already_rotated = current_generation.is_some() && current_generation != previous_generation;
|
||||
if !already_rotated {
|
||||
rotate_codex_credential_generation(key, provider_type);
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn create_provider_oauth_catalog_key(
|
||||
state: &AdminAppState<'_>,
|
||||
provider_id: &str,
|
||||
@@ -344,6 +400,7 @@ pub(crate) async fn create_provider_oauth_catalog_key(
|
||||
record.circuit_breaker_by_format = Some(json!({}));
|
||||
record.created_at_unix_ms = Some(now_unix_secs);
|
||||
record.updated_at_unix_secs = Some(now_unix_secs);
|
||||
rotate_codex_credential_generation(&mut record, provider_type);
|
||||
let created = state.create_provider_catalog_key(&record).await?;
|
||||
if let Some(key) = created.as_ref() {
|
||||
let _ = state
|
||||
@@ -395,17 +452,23 @@ pub(crate) async fn update_existing_provider_oauth_catalog_key(
|
||||
updated.proxy = Some(proxy);
|
||||
}
|
||||
updated.updated_at_unix_secs = Some(now_unix_secs);
|
||||
if state.update_provider_catalog_key(&updated).await?.is_none() {
|
||||
return Ok(None);
|
||||
}
|
||||
rotate_codex_credential_generation(&mut updated, provider_type);
|
||||
let admin_update =
|
||||
build_provider_catalog_key_admin_cas_update(existing_key, updated.clone(), provider_type);
|
||||
if !state
|
||||
.clear_provider_catalog_key_oauth_invalid_marker(&updated.id)
|
||||
.compare_and_update_provider_catalog_key_admin_state(&admin_update)
|
||||
.await?
|
||||
{
|
||||
return Ok(None);
|
||||
}
|
||||
let persisted = state
|
||||
.reset_provider_catalog_key_recovery_state(&updated.id)
|
||||
.reset_provider_catalog_key_recovery_state_fenced(
|
||||
&updated.id,
|
||||
updated
|
||||
.encrypted_auth_config
|
||||
.as_deref()
|
||||
.expect("OAuth update always supplies encrypted auth_config"),
|
||||
)
|
||||
.await?;
|
||||
if let Some(key) = persisted.as_ref() {
|
||||
let _ = state
|
||||
@@ -502,10 +565,12 @@ fn provider_oauth_catalog_key_api_formats(
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
grok_oauth_catalog_key_fingerprint, provider_oauth_token_payload_expires_at_unix_secs,
|
||||
ensure_codex_credential_generation_rotated, grok_oauth_catalog_key_fingerprint,
|
||||
provider_oauth_token_payload_expires_at_unix_secs, rotate_codex_credential_generation,
|
||||
};
|
||||
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
|
||||
use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
|
||||
use serde_json::json;
|
||||
use serde_json::{json, Value};
|
||||
|
||||
fn sample_unsigned_jwt(payload: serde_json::Value) -> String {
|
||||
let header = URL_SAFE_NO_PAD.encode(r#"{"alg":"none","typ":"JWT"}"#);
|
||||
@@ -609,4 +674,101 @@ mod tests {
|
||||
|
||||
assert!(grok_oauth_catalog_key_fingerprint("openai", auth_config).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn codex_credential_rotation_replaces_quota_namespace_and_preserves_unrelated_state() {
|
||||
let mut key = StoredProviderCatalogKey::new(
|
||||
"key".to_string(),
|
||||
"provider".to_string(),
|
||||
"Codex".to_string(),
|
||||
"oauth".to_string(),
|
||||
None,
|
||||
true,
|
||||
)
|
||||
.expect("key should build");
|
||||
key.upstream_metadata = Some(json!({
|
||||
"codex": {
|
||||
"credential_generation": "old-generation",
|
||||
"primary_used_percent": 75.0,
|
||||
},
|
||||
"unrelated": {"preserved": true},
|
||||
}));
|
||||
key.status_snapshot = Some(json!({
|
||||
"oauth": {"status": "valid"},
|
||||
"quota": {"used_ratio": 0.75},
|
||||
}));
|
||||
|
||||
rotate_codex_credential_generation(&mut key, "codex");
|
||||
|
||||
let codex = key
|
||||
.upstream_metadata
|
||||
.as_ref()
|
||||
.and_then(|metadata| metadata.get("codex"))
|
||||
.and_then(Value::as_object)
|
||||
.expect("codex namespace should exist");
|
||||
assert_eq!(codex.len(), 1);
|
||||
assert_ne!(
|
||||
codex
|
||||
.get(aether_admin::provider::quota::CODEX_CREDENTIAL_GENERATION_KEY)
|
||||
.and_then(Value::as_str),
|
||||
Some("old-generation")
|
||||
);
|
||||
assert_eq!(
|
||||
key.upstream_metadata
|
||||
.as_ref()
|
||||
.and_then(|metadata| metadata.pointer("/unrelated/preserved")),
|
||||
Some(&json!(true))
|
||||
);
|
||||
assert_eq!(
|
||||
key.status_snapshot
|
||||
.as_ref()
|
||||
.and_then(|snapshot| snapshot.get("quota")),
|
||||
Some(&Value::Null)
|
||||
);
|
||||
assert_eq!(
|
||||
key.status_snapshot
|
||||
.as_ref()
|
||||
.and_then(|snapshot| snapshot.pointer("/oauth/status")),
|
||||
Some(&json!("valid"))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn codex_credential_rotation_ensure_does_not_rotate_twice_in_one_write() {
|
||||
let mut key = StoredProviderCatalogKey::new(
|
||||
"key".to_string(),
|
||||
"provider".to_string(),
|
||||
"Codex".to_string(),
|
||||
"oauth".to_string(),
|
||||
None,
|
||||
true,
|
||||
)
|
||||
.expect("key should build");
|
||||
key.upstream_metadata = Some(json!({
|
||||
"codex": {"credential_generation": "generation-before-write"}
|
||||
}));
|
||||
|
||||
rotate_codex_credential_generation(&mut key, "codex");
|
||||
let builder_generation = key
|
||||
.upstream_metadata
|
||||
.as_ref()
|
||||
.and_then(|metadata| metadata.pointer("/codex/credential_generation"))
|
||||
.and_then(Value::as_str)
|
||||
.expect("builder should rotate the generation")
|
||||
.to_string();
|
||||
|
||||
ensure_codex_credential_generation_rotated(
|
||||
&mut key,
|
||||
"codex",
|
||||
Some("generation-before-write"),
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
key.upstream_metadata
|
||||
.as_ref()
|
||||
.and_then(|metadata| metadata.pointer("/codex/credential_generation"))
|
||||
.and_then(Value::as_str),
|
||||
Some(builder_generation.as_str())
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -470,6 +470,7 @@ mod tests {
|
||||
candidate_id: None,
|
||||
status_code: 403,
|
||||
headers: BTreeMap::new(),
|
||||
response_observation: None,
|
||||
body: Some(ResponseBody {
|
||||
json_body: None,
|
||||
body_bytes_b64: Some(base64::engine::general_purpose::STANDARD.encode(body)),
|
||||
|
||||
@@ -18,24 +18,229 @@ use self::plan::{
|
||||
execute_codex_reset_credit_plan,
|
||||
};
|
||||
use super::shared::{
|
||||
build_quota_snapshot_payload, extract_execution_error_message,
|
||||
oauth_refresh_auto_removed_result, persist_fenced_provider_quota_refresh_state,
|
||||
persist_provider_quota_refresh_state, provider_auto_remove_banned_keys,
|
||||
build_quota_snapshot_payload, complete_codex_account_reset, extract_execution_error_message,
|
||||
oauth_refresh_auto_removed_result, persist_codex_provider_quota_refresh_state,
|
||||
persist_fenced_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_oauth_invalid_key,
|
||||
ProviderQuotaExecutionOutcome,
|
||||
quota_refresh_success_invalid_state, reserve_codex_account_reset,
|
||||
should_auto_remove_oauth_invalid_key, CodexAccountResetCompleteResult,
|
||||
CodexAccountResetReserveResult, CodexAccountResetTerminal, ProviderQuotaExecutionOutcome,
|
||||
};
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
|
||||
use crate::provider_key_auth::provider_key_is_oauth_managed;
|
||||
use crate::state::ProviderTransportCredentialFence;
|
||||
use crate::GatewayError;
|
||||
use aether_contracts::ProxySnapshot;
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthCredentialFence,
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||
ProviderCatalogKeyOAuthCredentialCasDelete,
|
||||
ProviderCatalogUpstreamMetadataNamespaceExpectation, StoredProviderCatalogEndpoint,
|
||||
StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||
};
|
||||
use axum::http::StatusCode;
|
||||
use serde_json::{json, Map, Value};
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
const CODEX_OAUTH_CREDENTIAL_STABILIZATION_ATTEMPTS: usize = 3;
|
||||
const CODEX_RESET_QUOTA_RECONCILIATION_DELAYS_MS: [u64; 4] = [1_000, 2_000, 4_000, 8_000];
|
||||
|
||||
enum CodexOAuthRequestPreparation {
|
||||
Ready {
|
||||
transport: AdminGatewayProviderTransportSnapshot,
|
||||
auth: (String, String),
|
||||
credential_fence: ProviderTransportCredentialFence,
|
||||
},
|
||||
MissingAuth,
|
||||
Conflict,
|
||||
}
|
||||
|
||||
async fn prepare_codex_oauth_request(
|
||||
state: &AdminAppState<'_>,
|
||||
initial_transport: &AdminGatewayProviderTransportSnapshot,
|
||||
) -> Result<CodexOAuthRequestPreparation, GatewayError> {
|
||||
for _ in 0..CODEX_OAUTH_CREDENTIAL_STABILIZATION_ATTEMPTS {
|
||||
let Some(transport) = state
|
||||
.read_provider_transport_snapshot_uncached(
|
||||
&initial_transport.provider.id,
|
||||
&initial_transport.endpoint.id,
|
||||
&initial_transport.key.id,
|
||||
)
|
||||
.await?
|
||||
else {
|
||||
return Ok(CodexOAuthRequestPreparation::Conflict);
|
||||
};
|
||||
if !crate::state::provider_transport_context_allows_credential_rotation(
|
||||
initial_transport,
|
||||
&transport,
|
||||
) {
|
||||
return Ok(CodexOAuthRequestPreparation::Conflict);
|
||||
}
|
||||
let Some(before_fence) = state
|
||||
.app()
|
||||
.capture_provider_transport_credential_fence(&transport)
|
||||
.await?
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
|
||||
let resolved_auth = state.resolve_local_oauth_header_auth(&transport).await?;
|
||||
let Some(current_transport) = state
|
||||
.read_provider_transport_snapshot_uncached(
|
||||
&initial_transport.provider.id,
|
||||
&initial_transport.endpoint.id,
|
||||
&initial_transport.key.id,
|
||||
)
|
||||
.await?
|
||||
else {
|
||||
return Ok(CodexOAuthRequestPreparation::Conflict);
|
||||
};
|
||||
if !crate::state::provider_transport_context_allows_credential_rotation(
|
||||
initial_transport,
|
||||
¤t_transport,
|
||||
) {
|
||||
return Ok(CodexOAuthRequestPreparation::Conflict);
|
||||
}
|
||||
let Some(after_fence) = state
|
||||
.app()
|
||||
.capture_provider_transport_credential_fence(¤t_transport)
|
||||
.await?
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
if before_fence != after_fence {
|
||||
continue;
|
||||
}
|
||||
|
||||
return Ok(match resolved_auth {
|
||||
Some(auth) => CodexOAuthRequestPreparation::Ready {
|
||||
transport: current_transport,
|
||||
auth,
|
||||
credential_fence: after_fence,
|
||||
},
|
||||
None => CodexOAuthRequestPreparation::MissingAuth,
|
||||
});
|
||||
}
|
||||
|
||||
Ok(CodexOAuthRequestPreparation::Conflict)
|
||||
}
|
||||
|
||||
fn codex_reset_refresh_succeeded(payload: Option<&Value>, key_id: &str) -> bool {
|
||||
payload
|
||||
.and_then(|payload| payload.get("results"))
|
||||
.and_then(Value::as_array)
|
||||
.into_iter()
|
||||
.flatten()
|
||||
.filter_map(Value::as_object)
|
||||
.find(|item| item.get("key_id").and_then(Value::as_str) == Some(key_id))
|
||||
.and_then(|item| item.get("status"))
|
||||
.and_then(Value::as_str)
|
||||
.is_some_and(|status| status.eq_ignore_ascii_case("success"))
|
||||
}
|
||||
|
||||
async fn codex_reset_fence_is_still_pending(
|
||||
state: &AdminAppState<'_>,
|
||||
key_id: &str,
|
||||
expected_credential: &ProviderTransportCredentialFence,
|
||||
reset_fence: &super::shared::CodexAccountResetFence,
|
||||
) -> Result<bool, GatewayError> {
|
||||
let Some(key) = state
|
||||
.read_provider_catalog_keys_by_ids(&[key_id.to_string()])
|
||||
.await?
|
||||
.into_iter()
|
||||
.next()
|
||||
else {
|
||||
return Ok(false);
|
||||
};
|
||||
if key.encrypted_auth_config.as_deref()
|
||||
!= Some(expected_credential.encrypted_auth_config.as_str())
|
||||
|| key.encrypted_api_key != expected_credential.credential.encrypted_api_key
|
||||
|| key.auth_type != expected_credential.credential.auth_type
|
||||
|| key.provider_id != expected_credential.credential.provider_id
|
||||
{
|
||||
return Ok(false);
|
||||
}
|
||||
let provider_type_matches = state
|
||||
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&key.provider_id))
|
||||
.await?
|
||||
.into_iter()
|
||||
.next()
|
||||
.is_some_and(|provider| {
|
||||
provider.provider_type == expected_credential.credential.provider_type
|
||||
});
|
||||
if !provider_type_matches {
|
||||
return Ok(false);
|
||||
}
|
||||
|
||||
let codex = key
|
||||
.upstream_metadata
|
||||
.as_ref()
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|metadata| metadata.get("codex"))
|
||||
.and_then(Value::as_object);
|
||||
Ok(codex.is_some_and(|codex| {
|
||||
codex
|
||||
.get(aether_admin::provider::quota::CODEX_QUOTA_ACCOUNT_RESET_FENCE_ID_KEY)
|
||||
.and_then(Value::as_str)
|
||||
== Some(reset_fence.id.as_str())
|
||||
&& codex
|
||||
.get(aether_admin::provider::quota::CODEX_QUOTA_ACCOUNT_RESET_GENERATION_KEY)
|
||||
.and_then(aether_admin::provider::quota::coerce_json_u64)
|
||||
== Some(reset_fence.generation)
|
||||
&& codex
|
||||
.get(
|
||||
aether_admin::provider::quota::CODEX_QUOTA_ACCOUNT_RESET_PENDING_GENERATION_KEY,
|
||||
)
|
||||
.and_then(aether_admin::provider::quota::coerce_json_u64)
|
||||
== Some(reset_fence.generation)
|
||||
&& codex
|
||||
.get(aether_admin::provider::quota::CODEX_QUOTA_ACCOUNT_RESET_PENDING_KEY)
|
||||
.and_then(Value::as_bool)
|
||||
== Some(true)
|
||||
}))
|
||||
}
|
||||
|
||||
async fn refresh_codex_quota_after_reset_until_settled(
|
||||
state: &AdminAppState<'_>,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
endpoint: &StoredProviderCatalogEndpoint,
|
||||
key: &StoredProviderCatalogKey,
|
||||
reset_fence: &super::shared::CodexAccountResetFence,
|
||||
expected_credential: &ProviderTransportCredentialFence,
|
||||
) -> Result<Option<Value>, GatewayError> {
|
||||
let mut latest_payload = None;
|
||||
for attempt in 0..=CODEX_RESET_QUOTA_RECONCILIATION_DELAYS_MS.len() {
|
||||
if attempt > 0 {
|
||||
tokio::time::sleep(std::time::Duration::from_millis(
|
||||
CODEX_RESET_QUOTA_RECONCILIATION_DELAYS_MS[attempt - 1],
|
||||
))
|
||||
.await;
|
||||
}
|
||||
if !codex_reset_fence_is_still_pending(state, &key.id, expected_credential, reset_fence)
|
||||
.await?
|
||||
{
|
||||
break;
|
||||
}
|
||||
|
||||
let payload = refresh_codex_provider_quota_locally_with_reset_fence(
|
||||
state,
|
||||
provider,
|
||||
endpoint,
|
||||
vec![key.clone()],
|
||||
None,
|
||||
Some(reset_fence.id.as_str()),
|
||||
Some(reset_fence.generation),
|
||||
Some(expected_credential),
|
||||
)
|
||||
.await?;
|
||||
let refresh_succeeded = codex_reset_refresh_succeeded(payload.as_ref(), &key.id);
|
||||
latest_payload = payload;
|
||||
if !refresh_succeeded
|
||||
|| !codex_reset_fence_is_still_pending(state, &key.id, expected_credential, reset_fence)
|
||||
.await?
|
||||
{
|
||||
break;
|
||||
}
|
||||
}
|
||||
Ok(latest_payload)
|
||||
}
|
||||
|
||||
fn merge_codex_quota_metadata(
|
||||
header_metadata: Option<&serde_json::Value>,
|
||||
@@ -53,6 +258,26 @@ fn merge_codex_quota_metadata(
|
||||
serde_json::Value::Object(merged)
|
||||
}
|
||||
|
||||
fn codex_quota_window_coverage(
|
||||
body_json: Option<&Value>,
|
||||
) -> aether_admin::provider::quota::CodexQuotaWindowCoverage {
|
||||
let body = body_json.and_then(Value::as_object);
|
||||
let has_account_snapshot = body
|
||||
.and_then(|body| body.get("rate_limit"))
|
||||
.and_then(Value::as_object)
|
||||
.is_some();
|
||||
let has_spark_snapshot = body
|
||||
.and_then(|body| body.get("additional_rate_limits"))
|
||||
.and_then(Value::as_array)
|
||||
.is_some();
|
||||
|
||||
match (has_account_snapshot, has_spark_snapshot) {
|
||||
(true, true) => aether_admin::provider::quota::CodexQuotaWindowCoverage::FullSnapshot,
|
||||
(true, false) => aether_admin::provider::quota::CodexQuotaWindowCoverage::AccountSnapshot,
|
||||
_ => aether_admin::provider::quota::CodexQuotaWindowCoverage::Patch,
|
||||
}
|
||||
}
|
||||
|
||||
fn truncate_codex_reset_credit_detail_error(message: impl Into<String>) -> String {
|
||||
let message = message.into();
|
||||
let mut sanitized = message.replace('\n', " ");
|
||||
@@ -198,6 +423,10 @@ fn codex_consume_success_status(outcome: &str) -> &'static str {
|
||||
}
|
||||
}
|
||||
|
||||
fn codex_reset_credit_outcome_allows_usage_drop(outcome: &str) -> bool {
|
||||
matches!(outcome, "reset" | "already_redeemed")
|
||||
}
|
||||
|
||||
fn codex_extract_refresh_result_fields(
|
||||
refresh_payload: Option<&Value>,
|
||||
key_id: &str,
|
||||
@@ -247,12 +476,74 @@ fn codex_extract_refresh_result_fields(
|
||||
)
|
||||
}
|
||||
|
||||
async fn finish_codex_reset_replay(
|
||||
state: &AdminAppState<'_>,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
endpoint: &StoredProviderCatalogEndpoint,
|
||||
key: &StoredProviderCatalogKey,
|
||||
credential: &ProviderTransportCredentialFence,
|
||||
terminal: CodexAccountResetTerminal,
|
||||
) -> Result<(StatusCode, Value), GatewayError> {
|
||||
let mut refresh_status = "skipped".to_string();
|
||||
let mut refresh_error = None;
|
||||
let mut metadata = None;
|
||||
let mut quota_snapshot = None;
|
||||
if codex_reset_credit_outcome_allows_usage_drop(&terminal.outcome) {
|
||||
let fence = super::shared::CodexAccountResetFence {
|
||||
unix_ms: crate::clock::current_unix_ms(),
|
||||
id: format!("reset:{}", terminal.idempotency_key),
|
||||
generation: terminal.generation,
|
||||
};
|
||||
if codex_reset_fence_is_still_pending(state, &key.id, credential, &fence).await? {
|
||||
match refresh_codex_quota_after_reset_until_settled(
|
||||
state, provider, endpoint, key, &fence, credential,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(payload) => {
|
||||
(refresh_status, refresh_error, metadata, quota_snapshot) =
|
||||
codex_extract_refresh_result_fields(payload.as_ref(), &key.id);
|
||||
}
|
||||
Err(err) => {
|
||||
refresh_status = "failed".to_string();
|
||||
refresh_error =
|
||||
Some(truncate_codex_reset_credit_detail_error(err.into_message()));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
let mut payload = Map::new();
|
||||
payload.insert("key_id".to_string(), json!(key.id));
|
||||
payload.insert(
|
||||
"status".to_string(),
|
||||
json!(codex_consume_success_status(&terminal.outcome)),
|
||||
);
|
||||
payload.insert("outcome".to_string(), json!(terminal.outcome));
|
||||
payload.insert(
|
||||
"idempotency_key".to_string(),
|
||||
json!(terminal.idempotency_key),
|
||||
);
|
||||
payload.insert("replay".to_string(), json!(true));
|
||||
payload.insert("refresh_status".to_string(), json!(refresh_status));
|
||||
if let Some(refresh_error) = refresh_error {
|
||||
payload.insert("refresh_error".to_string(), json!(refresh_error));
|
||||
}
|
||||
if let Some(metadata) = metadata {
|
||||
payload.insert("metadata".to_string(), metadata);
|
||||
}
|
||||
if let Some(quota_snapshot) = quota_snapshot {
|
||||
payload.insert("quota_snapshot".to_string(), quota_snapshot);
|
||||
}
|
||||
Ok((StatusCode::OK, Value::Object(payload)))
|
||||
}
|
||||
|
||||
pub(crate) async fn consume_codex_reset_credit_locally(
|
||||
state: &AdminAppState<'_>,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
endpoint: &StoredProviderCatalogEndpoint,
|
||||
key: StoredProviderCatalogKey,
|
||||
idempotency_key: &str,
|
||||
expected_credential_generation: Option<&str>,
|
||||
) -> Result<(StatusCode, Value), GatewayError> {
|
||||
let transport = match state
|
||||
.read_provider_transport_snapshot(&provider.id, &endpoint.id, &key.id)
|
||||
@@ -273,22 +564,48 @@ pub(crate) async fn consume_codex_reset_credit_locally(
|
||||
};
|
||||
|
||||
let is_oauth_managed = provider_key_is_oauth_managed(&key, provider.provider_type.as_str());
|
||||
let resolved_oauth_auth = if is_oauth_managed {
|
||||
state.resolve_local_oauth_header_auth(&transport).await?
|
||||
} else {
|
||||
None
|
||||
};
|
||||
if is_oauth_managed && resolved_oauth_auth.is_none() {
|
||||
if !is_oauth_managed {
|
||||
return Ok((
|
||||
StatusCode::BAD_REQUEST,
|
||||
json!({
|
||||
"key_id": key.id,
|
||||
"status": "error",
|
||||
"outcome": "error",
|
||||
"message": "缺少 Codex OAuth 认证信息,请先重新授权/刷新 Token",
|
||||
"message": "Codex reset credit 仅支持 OAuth 托管账号",
|
||||
}),
|
||||
));
|
||||
}
|
||||
let (transport, resolved_oauth_auth, reset_credential_fence) =
|
||||
match prepare_codex_oauth_request(state, &transport).await? {
|
||||
CodexOAuthRequestPreparation::Ready {
|
||||
transport,
|
||||
auth,
|
||||
credential_fence,
|
||||
} => (transport, Some(auth), credential_fence),
|
||||
CodexOAuthRequestPreparation::MissingAuth => {
|
||||
return Ok((
|
||||
StatusCode::BAD_REQUEST,
|
||||
json!({
|
||||
"key_id": key.id,
|
||||
"status": "error",
|
||||
"outcome": "error",
|
||||
"message": "缺少 Codex OAuth 认证信息,请先重新授权/刷新 Token",
|
||||
}),
|
||||
));
|
||||
}
|
||||
CodexOAuthRequestPreparation::Conflict => {
|
||||
return Ok((
|
||||
StatusCode::CONFLICT,
|
||||
json!({
|
||||
"key_id": key.id,
|
||||
"status": "error",
|
||||
"outcome": "error",
|
||||
"idempotency_key": idempotency_key,
|
||||
"message": "Codex credential changed before reset credit could be consumed",
|
||||
}),
|
||||
));
|
||||
}
|
||||
};
|
||||
|
||||
let request_spec = match build_codex_reset_credit_consume_request_spec(
|
||||
&transport,
|
||||
@@ -309,6 +626,79 @@ pub(crate) async fn consume_codex_reset_credit_locally(
|
||||
}
|
||||
};
|
||||
|
||||
let reservation = match reserve_codex_account_reset(
|
||||
state,
|
||||
&key.id,
|
||||
reset_credential_fence.encrypted_auth_config.as_str(),
|
||||
&reset_credential_fence.credential,
|
||||
expected_credential_generation,
|
||||
idempotency_key,
|
||||
)
|
||||
.await?
|
||||
{
|
||||
Some(CodexAccountResetReserveResult::Reserved(reservation)) => reservation,
|
||||
Some(CodexAccountResetReserveResult::Replay(terminal)) => {
|
||||
return finish_codex_reset_replay(
|
||||
state,
|
||||
provider,
|
||||
endpoint,
|
||||
&key,
|
||||
&reset_credential_fence,
|
||||
terminal,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
Some(CodexAccountResetReserveResult::LegacyReplay) => {
|
||||
return Ok((
|
||||
StatusCode::OK,
|
||||
json!({
|
||||
"key_id": key.id,
|
||||
"status": "success",
|
||||
"outcome": "historical_replay",
|
||||
"idempotency_key": idempotency_key,
|
||||
"refresh_status": "skipped",
|
||||
}),
|
||||
));
|
||||
}
|
||||
Some(CodexAccountResetReserveResult::Busy(active)) => {
|
||||
return Ok((
|
||||
StatusCode::CONFLICT,
|
||||
json!({
|
||||
"key_id": key.id,
|
||||
"status": "error",
|
||||
"outcome": "busy",
|
||||
"idempotency_key": idempotency_key,
|
||||
"active_idempotency_key": active.idempotency_key,
|
||||
"message": "Another Codex reset credit operation is unresolved",
|
||||
}),
|
||||
));
|
||||
}
|
||||
Some(CodexAccountResetReserveResult::CredentialGenerationMismatch) => {
|
||||
return Ok((
|
||||
StatusCode::CONFLICT,
|
||||
json!({
|
||||
"key_id": key.id,
|
||||
"status": "error",
|
||||
"outcome": "credential_changed",
|
||||
"idempotency_key": idempotency_key,
|
||||
"message": "Codex credential changed since this reset request was prepared",
|
||||
}),
|
||||
));
|
||||
}
|
||||
None => {
|
||||
return Ok((
|
||||
StatusCode::CONFLICT,
|
||||
json!({
|
||||
"key_id": key.id,
|
||||
"status": "error",
|
||||
"outcome": "error",
|
||||
"idempotency_key": idempotency_key,
|
||||
"message": "Codex reset reservation could not be persisted",
|
||||
}),
|
||||
));
|
||||
}
|
||||
};
|
||||
|
||||
let result =
|
||||
match execute_codex_reset_credit_plan(state, &transport, request_spec, None).await? {
|
||||
ProviderQuotaExecutionOutcome::Response(result) => result,
|
||||
@@ -331,11 +721,11 @@ pub(crate) async fn consume_codex_reset_credit_locally(
|
||||
.and_then(|body| body.json_body.as_ref());
|
||||
let outcome = normalize_codex_reset_credit_consume_outcome(body_json)
|
||||
.unwrap_or_else(|| "unknown".to_string());
|
||||
let known_non_error_outcome = matches!(
|
||||
let known_terminal_outcome = matches!(
|
||||
outcome.as_str(),
|
||||
"reset" | "already_redeemed" | "nothing_to_reset" | "no_credit"
|
||||
);
|
||||
if result.status_code >= 400 && !known_non_error_outcome {
|
||||
if !known_terminal_outcome {
|
||||
let detail = extract_execution_error_message(&result)
|
||||
.unwrap_or_else(|| format!("HTTP {}", result.status_code));
|
||||
return Ok((
|
||||
@@ -345,19 +735,88 @@ pub(crate) async fn consume_codex_reset_credit_locally(
|
||||
"status": "error",
|
||||
"outcome": "error",
|
||||
"idempotency_key": idempotency_key,
|
||||
"message": format!("reset credit consume 返回状态码 {}: {detail}", result.status_code),
|
||||
"message": format!("reset credit consume outcome is ambiguous: {detail}"),
|
||||
"status_code": result.status_code,
|
||||
}),
|
||||
));
|
||||
}
|
||||
|
||||
let (refresh_status, refresh_error, metadata, quota_snapshot) =
|
||||
match refresh_codex_provider_quota_locally(
|
||||
let fence_unix_ms = result
|
||||
.response_observation
|
||||
.as_ref()
|
||||
.map(|observation| observation.response_headers_observed_at_unix_ms)
|
||||
.unwrap_or_else(crate::clock::current_unix_ms);
|
||||
let Some(completed) = complete_codex_account_reset(
|
||||
state,
|
||||
&key.id,
|
||||
reset_credential_fence.encrypted_auth_config.as_str(),
|
||||
&reset_credential_fence.credential,
|
||||
&reservation,
|
||||
&outcome,
|
||||
fence_unix_ms,
|
||||
)
|
||||
.await?
|
||||
else {
|
||||
return Ok((
|
||||
StatusCode::CONFLICT,
|
||||
json!({
|
||||
"key_id": key.id,
|
||||
"status": "error",
|
||||
"outcome": "error",
|
||||
"idempotency_key": idempotency_key,
|
||||
"message": "Codex reset completion could not be persisted",
|
||||
}),
|
||||
));
|
||||
};
|
||||
|
||||
let (effective_outcome, reset_fence) = match completed {
|
||||
CodexAccountResetCompleteResult::Activated(fence) => (outcome.clone(), Some(fence)),
|
||||
CodexAccountResetCompleteResult::Noop(terminal)
|
||||
| CodexAccountResetCompleteResult::Replay(terminal) => {
|
||||
let fence =
|
||||
codex_reset_credit_outcome_allows_usage_drop(&terminal.outcome).then(|| {
|
||||
super::shared::CodexAccountResetFence {
|
||||
unix_ms: fence_unix_ms,
|
||||
id: format!("reset:{}", terminal.idempotency_key),
|
||||
generation: terminal.generation,
|
||||
}
|
||||
});
|
||||
(terminal.outcome, fence)
|
||||
}
|
||||
};
|
||||
|
||||
let (refresh_status, refresh_error, metadata, quota_snapshot) = match reset_fence.as_ref() {
|
||||
Some(reset_fence) => {
|
||||
match refresh_codex_quota_after_reset_until_settled(
|
||||
state,
|
||||
provider,
|
||||
endpoint,
|
||||
&key,
|
||||
reset_fence,
|
||||
&reset_credential_fence,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(refresh_payload) => {
|
||||
codex_extract_refresh_result_fields(refresh_payload.as_ref(), &key.id)
|
||||
}
|
||||
Err(err) => (
|
||||
"failed".to_string(),
|
||||
Some(truncate_codex_reset_credit_detail_error(err.into_message())),
|
||||
None,
|
||||
None,
|
||||
),
|
||||
}
|
||||
}
|
||||
None => match refresh_codex_provider_quota_locally_with_reset_fence(
|
||||
state,
|
||||
provider,
|
||||
endpoint,
|
||||
vec![key.clone()],
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
Some(&reset_credential_fence),
|
||||
)
|
||||
.await
|
||||
{
|
||||
@@ -370,15 +829,16 @@ pub(crate) async fn consume_codex_reset_credit_locally(
|
||||
None,
|
||||
None,
|
||||
),
|
||||
};
|
||||
},
|
||||
};
|
||||
|
||||
let mut payload = Map::new();
|
||||
payload.insert("key_id".to_string(), json!(key.id));
|
||||
payload.insert(
|
||||
"status".to_string(),
|
||||
json!(codex_consume_success_status(&outcome)),
|
||||
json!(codex_consume_success_status(&effective_outcome)),
|
||||
);
|
||||
payload.insert("outcome".to_string(), json!(outcome));
|
||||
payload.insert("outcome".to_string(), json!(effective_outcome));
|
||||
payload.insert("idempotency_key".to_string(), json!(idempotency_key));
|
||||
payload.insert("refresh_status".to_string(), json!(refresh_status));
|
||||
if let Some(refresh_error) = refresh_error {
|
||||
@@ -400,6 +860,29 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
endpoint: &StoredProviderCatalogEndpoint,
|
||||
keys: Vec<StoredProviderCatalogKey>,
|
||||
proxy_override: Option<ProxySnapshot>,
|
||||
) -> Result<Option<serde_json::Value>, GatewayError> {
|
||||
refresh_codex_provider_quota_locally_with_reset_fence(
|
||||
state,
|
||||
provider,
|
||||
endpoint,
|
||||
keys,
|
||||
proxy_override,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn refresh_codex_provider_quota_locally_with_reset_fence(
|
||||
state: &AdminAppState<'_>,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
endpoint: &StoredProviderCatalogEndpoint,
|
||||
keys: Vec<StoredProviderCatalogKey>,
|
||||
proxy_override: Option<ProxySnapshot>,
|
||||
account_reset_fence_id: Option<&str>,
|
||||
authoritative_reset_generation: Option<u64>,
|
||||
expected_reset_credential: Option<&crate::state::ProviderTransportCredentialFence>,
|
||||
) -> Result<Option<serde_json::Value>, GatewayError> {
|
||||
let mut results = Vec::new();
|
||||
let mut success_count = 0usize;
|
||||
@@ -412,7 +895,7 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
for key in keys {
|
||||
let had_oauth_refresh_issue =
|
||||
codex_oauth_refresh_issue_reason(key.oauth_invalid_reason.as_deref());
|
||||
let transport = match state
|
||||
let initial_transport = match state
|
||||
.read_provider_transport_snapshot(&provider.id, &endpoint.id, &key.id)
|
||||
.await?
|
||||
{
|
||||
@@ -429,48 +912,70 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
}
|
||||
};
|
||||
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 => {
|
||||
let (transport, resolved_oauth_auth, quota_credential_fence) = if is_oauth_managed {
|
||||
match prepare_codex_oauth_request(state, &initial_transport).await? {
|
||||
CodexOAuthRequestPreparation::Ready {
|
||||
transport,
|
||||
auth,
|
||||
credential_fence,
|
||||
} => (transport, Some(auth), Some(credential_fence)),
|
||||
CodexOAuthRequestPreparation::MissingAuth => {
|
||||
failed_count += 1;
|
||||
results.push(json!({
|
||||
"key_id": key.id,
|
||||
"key_name": key.name,
|
||||
"status": "error",
|
||||
"message": "OAuth credential changed before quota refresh",
|
||||
"message": "缺少 Codex OAuth 认证信息,请先重新授权/刷新 Token",
|
||||
}));
|
||||
continue;
|
||||
}
|
||||
CodexOAuthRequestPreparation::Conflict => {
|
||||
if quota_key_auto_removed(state, &key.id).await? {
|
||||
auto_removed_count += 1;
|
||||
results.push(oauth_refresh_auto_removed_result(&key));
|
||||
} else {
|
||||
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
|
||||
(initial_transport, None, None)
|
||||
};
|
||||
|
||||
let resolved_oauth_auth = if is_oauth_managed {
|
||||
state.resolve_local_oauth_header_auth(&transport).await?
|
||||
} else {
|
||||
None
|
||||
};
|
||||
if is_oauth_managed && quota_key_auto_removed(state, &key.id).await? {
|
||||
auto_removed_count += 1;
|
||||
results.push(oauth_refresh_auto_removed_result(&key));
|
||||
continue;
|
||||
}
|
||||
if is_oauth_managed && resolved_oauth_auth.is_none() {
|
||||
failed_count += 1;
|
||||
results.push(json!({
|
||||
"key_id": key.id,
|
||||
"key_name": key.name,
|
||||
"status": "error",
|
||||
"message": "缺少 Codex OAuth 认证信息,请先重新授权/刷新 Token",
|
||||
}));
|
||||
continue;
|
||||
if let Some(expected_reset_credential) = expected_reset_credential {
|
||||
if quota_credential_fence.as_ref() != Some(expected_reset_credential) {
|
||||
failed_count += 1;
|
||||
results.push(json!({
|
||||
"key_id": key.id,
|
||||
"key_name": key.name,
|
||||
"status": "error",
|
||||
"message": "Codex credential changed after reset credit was consumed",
|
||||
}));
|
||||
continue;
|
||||
}
|
||||
}
|
||||
let transport_codex_metadata = transport
|
||||
.key
|
||||
.upstream_metadata
|
||||
.as_ref()
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|metadata| metadata.get("codex"));
|
||||
let observed_reset_generation = authoritative_reset_generation.or_else(|| {
|
||||
Some(
|
||||
aether_admin::provider::quota::codex_quota_account_reset_generation(
|
||||
transport_codex_metadata,
|
||||
),
|
||||
)
|
||||
});
|
||||
let observed_credential_generation =
|
||||
aether_admin::provider::quota::codex_credential_generation(transport_codex_metadata)
|
||||
.map(ToOwned::to_owned);
|
||||
|
||||
let request_spec =
|
||||
match build_codex_quota_request_spec(&transport, resolved_oauth_auth.clone()) {
|
||||
@@ -487,6 +992,8 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
}
|
||||
};
|
||||
|
||||
let quota_request_fallback_started_at_unix_ms = crate::clock::current_unix_ms();
|
||||
let quota_request_fallback_order_id = uuid::Uuid::now_v7().to_string();
|
||||
let result = match execute_codex_quota_plan(
|
||||
state,
|
||||
&transport,
|
||||
@@ -508,16 +1015,25 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
continue;
|
||||
}
|
||||
};
|
||||
let now_unix_secs = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(0);
|
||||
let quota_response_fallback_observed_at_unix_ms = crate::clock::current_unix_ms();
|
||||
let quota_response_observation = result.response_observation.as_ref();
|
||||
let quota_request_started_at_unix_ms = quota_response_observation
|
||||
.map(|observation| observation.request_started_at_unix_ms)
|
||||
.unwrap_or(quota_request_fallback_started_at_unix_ms);
|
||||
let quota_response_observed_at_unix_ms = quota_response_observation
|
||||
.map(|observation| observation.response_headers_observed_at_unix_ms)
|
||||
.unwrap_or(quota_response_fallback_observed_at_unix_ms);
|
||||
let quota_request_order_id = quota_response_observation
|
||||
.map(|observation| observation.request_order_id.as_str())
|
||||
.unwrap_or(quota_request_fallback_order_id.as_str());
|
||||
let now_unix_secs = quota_response_observed_at_unix_ms / 1_000;
|
||||
|
||||
let header_metadata = parse_codex_usage_headers(&result.headers, now_unix_secs);
|
||||
let mut metadata_update = header_metadata
|
||||
.as_ref()
|
||||
.map(|metadata| json!({ "codex": metadata }));
|
||||
let mut quota_window_coverage =
|
||||
aether_admin::provider::quota::CodexQuotaWindowCoverage::Patch;
|
||||
let (mut oauth_invalid_at_unix_secs, mut oauth_invalid_reason) = (None, None);
|
||||
let mut status = "error".to_string();
|
||||
let mut message = None::<String>;
|
||||
@@ -544,6 +1060,7 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
now_unix_secs,
|
||||
)
|
||||
.await?;
|
||||
quota_window_coverage = codex_quota_window_coverage(Some(body_json));
|
||||
metadata_update = Some(json!({
|
||||
"codex": codex_metadata
|
||||
}));
|
||||
@@ -586,6 +1103,8 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
}
|
||||
402 => {
|
||||
if codex_looks_like_workspace_deactivated(err_msg.as_deref()) {
|
||||
quota_window_coverage =
|
||||
aether_admin::provider::quota::CodexQuotaWindowCoverage::Patch;
|
||||
let mut codex_meta = metadata_update
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("codex"))
|
||||
@@ -627,6 +1146,8 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
oauth_invalid_reason = reason;
|
||||
status = "workspace_deactivated".to_string();
|
||||
} else {
|
||||
quota_window_coverage =
|
||||
aether_admin::provider::quota::CodexQuotaWindowCoverage::Patch;
|
||||
let plan_type = transport
|
||||
.key
|
||||
.decrypted_auth_config
|
||||
@@ -667,24 +1188,45 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
}
|
||||
}
|
||||
|
||||
let persisted = if let Some(expected_auth_config) = quota_auth_config_fence.as_deref() {
|
||||
let persisted = if let Some(expected_credential) = quota_credential_fence.as_ref() {
|
||||
persist_fenced_provider_quota_refresh_state(
|
||||
state,
|
||||
&key.id,
|
||||
expected_auth_config,
|
||||
expected_credential.encrypted_auth_config.as_str(),
|
||||
metadata_update.as_ref(),
|
||||
oauth_invalid_at_unix_secs,
|
||||
oauth_invalid_reason.clone(),
|
||||
aether_admin::provider::quota::CodexQuotaMergeContext {
|
||||
observed_at_unix_secs: now_unix_secs,
|
||||
request_started_at_unix_ms: Some(quota_request_started_at_unix_ms),
|
||||
request_order_id: Some(quota_request_order_id),
|
||||
observed_reset_generation,
|
||||
authoritative_reset_generation,
|
||||
observed_credential_generation: observed_credential_generation.as_deref(),
|
||||
account_reset_fence_id,
|
||||
coverage: quota_window_coverage,
|
||||
},
|
||||
Some(&expected_credential.credential),
|
||||
)
|
||||
.await?
|
||||
} else {
|
||||
persist_provider_quota_refresh_state(
|
||||
persist_codex_provider_quota_refresh_state(
|
||||
state,
|
||||
&key.id,
|
||||
metadata_update.as_ref(),
|
||||
oauth_invalid_at_unix_secs,
|
||||
oauth_invalid_reason.clone(),
|
||||
None,
|
||||
aether_admin::provider::quota::CodexQuotaMergeContext {
|
||||
observed_at_unix_secs: now_unix_secs,
|
||||
request_started_at_unix_ms: Some(quota_request_started_at_unix_ms),
|
||||
request_order_id: Some(quota_request_order_id),
|
||||
observed_reset_generation,
|
||||
authoritative_reset_generation,
|
||||
observed_credential_generation: observed_credential_generation.as_deref(),
|
||||
account_reset_fence_id,
|
||||
coverage: quota_window_coverage,
|
||||
},
|
||||
)
|
||||
.await?
|
||||
};
|
||||
@@ -698,26 +1240,57 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
}));
|
||||
continue;
|
||||
}
|
||||
let credential_cas_delete = quota_auth_config_fence.as_ref().map(|auth_config| {
|
||||
let persisted_key = state
|
||||
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key.id))
|
||||
.await?
|
||||
.into_iter()
|
||||
.next();
|
||||
let persisted_codex_metadata = persisted_key
|
||||
.as_ref()
|
||||
.and_then(|key| key.upstream_metadata.as_ref())
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.and_then(|metadata| metadata.get("codex"))
|
||||
.cloned();
|
||||
if let Some(codex_metadata) = persisted_codex_metadata.as_ref() {
|
||||
metadata_update = Some(json!({"codex": codex_metadata}));
|
||||
}
|
||||
let persisted_codex_object = persisted_codex_metadata
|
||||
.as_ref()
|
||||
.and_then(serde_json::Value::as_object);
|
||||
let request_owns_persisted_oauth_state = quota_credential_fence.is_none()
|
||||
|| (persisted_codex_object
|
||||
.and_then(|codex| codex.get("oauth_state_request_started_at_unix_ms"))
|
||||
.and_then(aether_admin::provider::quota::coerce_json_u64)
|
||||
== Some(quota_request_started_at_unix_ms)
|
||||
&& persisted_codex_object
|
||||
.and_then(|codex| codex.get("oauth_state_request_id"))
|
||||
.and_then(serde_json::Value::as_str)
|
||||
== Some(quota_request_order_id));
|
||||
let credential_cas_delete = quota_credential_fence.as_ref().map(|credential_fence| {
|
||||
ProviderCatalogKeyOAuthCredentialCasDelete {
|
||||
key_id: key.id.clone(),
|
||||
expected_encrypted_auth_config: Some(auth_config.clone()),
|
||||
expected_credential: ProviderCatalogKeyOAuthCredentialFence {
|
||||
encrypted_api_key: key.encrypted_api_key.clone(),
|
||||
auth_type: key.auth_type.clone(),
|
||||
provider_id: key.provider_id.clone(),
|
||||
provider_type: provider.provider_type.clone(),
|
||||
},
|
||||
expected_encrypted_auth_config: Some(
|
||||
credential_fence.encrypted_auth_config.clone(),
|
||||
),
|
||||
expected_credential: credential_fence.credential.clone(),
|
||||
expected_upstream_metadata_namespace: Some(
|
||||
ProviderCatalogUpstreamMetadataNamespaceExpectation {
|
||||
namespace: "codex".to_string(),
|
||||
expected_value: persisted_codex_metadata.clone(),
|
||||
},
|
||||
),
|
||||
}
|
||||
});
|
||||
let should_auto_remove_hard_banned =
|
||||
provider_auto_remove_banned_keys(provider.config.as_ref())
|
||||
&& should_auto_remove_oauth_invalid_key(
|
||||
&key,
|
||||
oauth_invalid_reason.as_deref(),
|
||||
matches!(status_code, Some(401 | 403)),
|
||||
now_unix_secs,
|
||||
);
|
||||
let should_auto_remove_hard_banned = request_owns_persisted_oauth_state
|
||||
&& provider_auto_remove_banned_keys(provider.config.as_ref())
|
||||
&& should_auto_remove_oauth_invalid_key(
|
||||
persisted_key.as_ref().unwrap_or(&key),
|
||||
persisted_key
|
||||
.as_ref()
|
||||
.and_then(|key| key.oauth_invalid_reason.as_deref()),
|
||||
matches!(status_code, Some(401 | 403)),
|
||||
now_unix_secs,
|
||||
);
|
||||
let auto_removed_hard_banned = if should_auto_remove_hard_banned {
|
||||
match credential_cas_delete.as_ref() {
|
||||
Some(delete) => {
|
||||
@@ -735,6 +1308,7 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
auto_removed_hard_banned_count += 1;
|
||||
}
|
||||
let auto_removed_quota_exhausted = if !auto_removed_hard_banned
|
||||
&& request_owns_persisted_oauth_state
|
||||
&& status == "quota_exhausted"
|
||||
&& provider_auto_remove_quota_exhausted_keys(provider.config.as_ref())
|
||||
{
|
||||
@@ -798,7 +1372,9 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
}
|
||||
if let Some(quota_snapshot) = build_quota_snapshot_payload(
|
||||
"codex",
|
||||
key.status_snapshot.as_ref(),
|
||||
persisted_key
|
||||
.as_ref()
|
||||
.and_then(|key| key.status_snapshot.as_ref()),
|
||||
metadata_update.as_ref(),
|
||||
) {
|
||||
payload.insert("quota_snapshot".to_string(), quota_snapshot);
|
||||
@@ -866,6 +1442,19 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn codex_reset_credit_only_allows_usage_drop_after_confirmed_redemption() {
|
||||
assert!(codex_reset_credit_outcome_allows_usage_drop("reset"));
|
||||
assert!(codex_reset_credit_outcome_allows_usage_drop(
|
||||
"already_redeemed"
|
||||
));
|
||||
assert!(!codex_reset_credit_outcome_allows_usage_drop(
|
||||
"nothing_to_reset"
|
||||
));
|
||||
assert!(!codex_reset_credit_outcome_allows_usage_drop("no_credit"));
|
||||
assert!(!codex_reset_credit_outcome_allows_usage_drop("unknown"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn codex_reset_credit_detail_failure_records_attempt_time() {
|
||||
let mut metadata = Map::new();
|
||||
@@ -880,4 +1469,29 @@ mod tests {
|
||||
Some(&json!(1_777_000_000u64))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn codex_quota_coverage_only_replaces_observed_window_families() {
|
||||
assert_eq!(
|
||||
codex_quota_window_coverage(Some(&json!({"credits":{"balance":5}}))),
|
||||
aether_admin::provider::quota::CodexQuotaWindowCoverage::Patch
|
||||
);
|
||||
assert_eq!(
|
||||
codex_quota_window_coverage(None),
|
||||
aether_admin::provider::quota::CodexQuotaWindowCoverage::Patch
|
||||
);
|
||||
assert_eq!(
|
||||
codex_quota_window_coverage(Some(&json!({
|
||||
"rate_limit":{"primary_window":{}}
|
||||
}))),
|
||||
aether_admin::provider::quota::CodexQuotaWindowCoverage::AccountSnapshot
|
||||
);
|
||||
assert_eq!(
|
||||
codex_quota_window_coverage(Some(&json!({
|
||||
"rate_limit":{"primary_window":{}},
|
||||
"additional_rate_limits":[]
|
||||
}))),
|
||||
aether_admin::provider::quota::CodexQuotaWindowCoverage::FullSnapshot
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -796,6 +796,7 @@ mod tests {
|
||||
candidate_id: None,
|
||||
status_code: 403,
|
||||
headers: BTreeMap::new(),
|
||||
response_observation: None,
|
||||
body: Some(ResponseBody {
|
||||
json_body: None,
|
||||
body_bytes_b64: Some(base64::engine::general_purpose::STANDARD.encode(body)),
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -399,7 +399,8 @@ async fn provider_query_read_cached_models(
|
||||
let cache_key = format!("upstream_models:{provider_id}:{key_id}");
|
||||
let raw = state.runtime_state().kv_get(&cache_key).await.ok()??;
|
||||
let parsed = serde_json::from_str::<Vec<Value>>(&raw).ok()?;
|
||||
Some(aggregate_models_for_cache(&parsed))
|
||||
let models = aggregate_models_for_cache(&parsed);
|
||||
(!models.is_empty()).then_some(models)
|
||||
}
|
||||
|
||||
async fn provider_query_read_provider_cached_models(
|
||||
@@ -409,7 +410,8 @@ async fn provider_query_read_provider_cached_models(
|
||||
let cache_key = format!("{ANTIGRAVITY_PROVIDER_CACHE_KEY_PREFIX}{provider_id}");
|
||||
let raw = state.runtime_state().kv_get(&cache_key).await.ok()??;
|
||||
let parsed = serde_json::from_str::<Vec<Value>>(&raw).ok()?;
|
||||
Some(aggregate_models_for_cache(&parsed))
|
||||
let models = aggregate_models_for_cache(&parsed);
|
||||
(!models.is_empty()).then_some(models)
|
||||
}
|
||||
|
||||
async fn provider_query_write_provider_cached_models(
|
||||
@@ -417,7 +419,11 @@ async fn provider_query_write_provider_cached_models(
|
||||
provider_id: &str,
|
||||
models: &[Value],
|
||||
) {
|
||||
let Ok(serialized) = serde_json::to_string(&aggregate_models_for_cache(models)) else {
|
||||
let models = aggregate_models_for_cache(models);
|
||||
if models.is_empty() {
|
||||
return;
|
||||
}
|
||||
let Ok(serialized) = serde_json::to_string(&models) else {
|
||||
return;
|
||||
};
|
||||
let cache_key = format!("{ANTIGRAVITY_PROVIDER_CACHE_KEY_PREFIX}{provider_id}");
|
||||
@@ -577,7 +583,7 @@ async fn provider_query_fetch_models_for_key(
|
||||
};
|
||||
|
||||
all_errors.extend(outcome.errors);
|
||||
let unique_models = aggregate_models_for_cache(&outcome.cached_models);
|
||||
let unique_models = outcome.legacy_models;
|
||||
if outcome.has_success && !unique_models.is_empty() {
|
||||
<AppState as ModelFetchRuntimeState>::write_upstream_models_cache(
|
||||
state.app(),
|
||||
|
||||
@@ -635,7 +635,6 @@ fn provider_query_build_test_request_body_for_api_format_with_search_session(
|
||||
"model": model,
|
||||
"input": message,
|
||||
"max_output_tokens": 30,
|
||||
"temperature": 0.7,
|
||||
"stream": true,
|
||||
}),
|
||||
"openai:search" => json!({
|
||||
@@ -654,7 +653,6 @@ fn provider_query_build_test_request_body_for_api_format_with_search_session(
|
||||
"content": message
|
||||
}],
|
||||
"max_tokens": 30,
|
||||
"temperature": 0.7,
|
||||
"stream": true,
|
||||
}),
|
||||
_ => json!({
|
||||
@@ -664,7 +662,6 @@ fn provider_query_build_test_request_body_for_api_format_with_search_session(
|
||||
"content": message
|
||||
}],
|
||||
"max_tokens": 30,
|
||||
"temperature": 0.7,
|
||||
"stream": true,
|
||||
}),
|
||||
}
|
||||
@@ -811,7 +808,6 @@ fn provider_query_build_test_request_body_with_model_policy(
|
||||
"content": provider_query_extract_message(payload)
|
||||
.unwrap_or_else(|| DEFAULT_PROVIDER_QUERY_TEST_MESSAGE.to_string())
|
||||
}],
|
||||
"temperature": 0.7,
|
||||
"stream": true,
|
||||
})
|
||||
}
|
||||
@@ -1304,6 +1300,10 @@ fn provider_query_pool_catalog_key_context(
|
||||
provider_type,
|
||||
quota_snapshot,
|
||||
),
|
||||
quota_hard_blocked: admin_provider_pool_pure::admin_pool_key_quota_hard_blocked(
|
||||
key,
|
||||
provider_type,
|
||||
),
|
||||
health_score,
|
||||
latency_avg_ms,
|
||||
catalog_lru_score: Some(key.last_used_at_unix_secs.unwrap_or(0) as f64),
|
||||
|
||||
@@ -178,6 +178,7 @@ fn provider_query_execution_json_body_decodes_stream_encoded_json_response() {
|
||||
"content-type".to_string(),
|
||||
"application/json".to_string(),
|
||||
)]),
|
||||
response_observation: None,
|
||||
body: Some(aether_contracts::ResponseBody {
|
||||
json_body: None,
|
||||
body_bytes_b64: Some(encoded_body),
|
||||
@@ -242,6 +243,26 @@ fn provider_query_default_test_request_body_does_not_set_max_tokens() {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_default_test_request_bodies_do_not_set_temperature() {
|
||||
let payload = json!({});
|
||||
let default_body = provider_query_build_test_request_body(&payload, "fallback-model");
|
||||
assert!(default_body.get("temperature").is_none());
|
||||
|
||||
for api_format in ["openai:chat", "openai:responses", "claude:messages"] {
|
||||
let body = provider_query_build_test_request_body_for_api_format(
|
||||
&payload,
|
||||
"fallback-model",
|
||||
"/api/admin/provider-query/test-model",
|
||||
api_format,
|
||||
);
|
||||
assert!(
|
||||
body.get("temperature").is_none(),
|
||||
"admin model test must not set temperature for {api_format}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_failover_request_body_overrides_custom_model() {
|
||||
let payload = json!({
|
||||
@@ -416,6 +437,7 @@ fn provider_query_standard_test_aggregates_responses_stream_body() {
|
||||
candidate_id: Some("candidate-0".to_string()),
|
||||
status_code: 200,
|
||||
headers: BTreeMap::new(),
|
||||
response_observation: None,
|
||||
body: Some(aether_contracts::ResponseBody {
|
||||
json_body: None,
|
||||
body_bytes_b64: Some(
|
||||
@@ -448,6 +470,7 @@ fn provider_query_standard_test_aggregates_responses_image_generation_call() {
|
||||
candidate_id: Some("candidate-0".to_string()),
|
||||
status_code: 200,
|
||||
headers: BTreeMap::new(),
|
||||
response_observation: None,
|
||||
body: Some(aether_contracts::ResponseBody {
|
||||
json_body: None,
|
||||
body_bytes_b64: Some(
|
||||
@@ -581,6 +604,7 @@ fn provider_query_search_success_requires_non_empty_output() {
|
||||
candidate_id: Some("candidate-0".to_string()),
|
||||
status_code: 200,
|
||||
headers: BTreeMap::new(),
|
||||
response_observation: None,
|
||||
body: Some(aether_contracts::ResponseBody {
|
||||
json_body: Some(body),
|
||||
body_bytes_b64: None,
|
||||
@@ -731,6 +755,7 @@ fn provider_query_standard_test_rejects_gemini_success_without_visible_output()
|
||||
candidate_id: Some("candidate-0".to_string()),
|
||||
status_code: 200,
|
||||
headers: BTreeMap::new(),
|
||||
response_observation: None,
|
||||
body: Some(aether_contracts::ResponseBody {
|
||||
json_body: Some(json!({
|
||||
"candidates": [{
|
||||
|
||||
@@ -122,6 +122,7 @@ pub(crate) struct AdminProviderQuotaRefreshRequest {
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub(crate) struct AdminCodexResetCreditConsumeRequest {
|
||||
pub(crate) idempotency_key: String,
|
||||
pub(crate) expected_credential_generation: serde_json::Value,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
@@ -151,6 +152,10 @@ pub(crate) struct AdminProviderCreateRequest {
|
||||
#[serde(default)]
|
||||
pub(crate) keep_priority_on_conversion: Option<bool>,
|
||||
#[serde(default)]
|
||||
pub(crate) codex_fingerprint_convergence_enabled: Option<bool>,
|
||||
#[serde(default)]
|
||||
pub(crate) responses_websocket_enabled: Option<bool>,
|
||||
#[serde(default)]
|
||||
pub(crate) is_active: Option<bool>,
|
||||
#[serde(default)]
|
||||
pub(crate) concurrent_limit: Option<i32>,
|
||||
@@ -210,6 +215,10 @@ pub(crate) struct AdminProviderUpdateRequest {
|
||||
#[serde(default)]
|
||||
pub(crate) keep_priority_on_conversion: Option<bool>,
|
||||
#[serde(default)]
|
||||
pub(crate) codex_fingerprint_convergence_enabled: Option<bool>,
|
||||
#[serde(default)]
|
||||
pub(crate) responses_websocket_enabled: Option<bool>,
|
||||
#[serde(default)]
|
||||
pub(crate) is_active: Option<bool>,
|
||||
#[serde(default)]
|
||||
pub(crate) concurrent_limit: Option<i32>,
|
||||
@@ -335,3 +344,37 @@ pub(crate) struct AdminImportProviderModelsRequest {
|
||||
)]
|
||||
pub(crate) price_per_request: Option<f64>,
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::AdminCodexResetCreditConsumeRequest;
|
||||
|
||||
#[test]
|
||||
fn codex_reset_credit_consume_requires_an_explicit_credential_generation() {
|
||||
assert!(
|
||||
serde_json::from_value::<AdminCodexResetCreditConsumeRequest>(
|
||||
serde_json::json!({"idempotency_key":"reset-old-client"}),
|
||||
)
|
||||
.is_err()
|
||||
);
|
||||
|
||||
let legacy_account =
|
||||
serde_json::from_value::<AdminCodexResetCreditConsumeRequest>(serde_json::json!({
|
||||
"idempotency_key":"reset-legacy-account",
|
||||
"expected_credential_generation":null,
|
||||
}))
|
||||
.expect("explicit null should fence an account without a generation");
|
||||
assert!(legacy_account.expected_credential_generation.is_null());
|
||||
|
||||
let generated_account =
|
||||
serde_json::from_value::<AdminCodexResetCreditConsumeRequest>(serde_json::json!({
|
||||
"idempotency_key":"reset-generated-account",
|
||||
"expected_credential_generation":"credential-v2",
|
||||
}))
|
||||
.expect("string generation should deserialize");
|
||||
assert_eq!(
|
||||
generated_account.expected_credential_generation,
|
||||
serde_json::json!("credential-v2")
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -4,7 +4,7 @@ use crate::handlers::admin::provider::shared::support::{
|
||||
};
|
||||
use crate::handlers::admin::shared::unix_secs_to_rfc3339;
|
||||
use crate::handlers::public::{request_candidate_event_unix_ms, request_candidate_status_label};
|
||||
use crate::orchestration::codex_cyber_flag_passthrough_enabled;
|
||||
use crate::orchestration::{codex_cyber_flag_passthrough_enabled, responses_websocket_adapter};
|
||||
use crate::provider_key_auth::provider_key_effective_api_formats;
|
||||
use aether_data_contracts::repository::candidates::{
|
||||
RequestCandidateStatus, StoredRequestCandidate,
|
||||
@@ -218,6 +218,11 @@ pub(crate) fn build_admin_provider_summary_value(
|
||||
"ops_architecture_id": ops_architecture_id,
|
||||
"kiro_simulated_cache_enabled": kiro_simulated_cache_enabled,
|
||||
"codex_cyber_flag_passthrough_enabled": codex_cyber_flag_passthrough_enabled(&provider.provider_type, provider.config.as_ref()),
|
||||
"codex_fingerprint_convergence_enabled": crate::provider_transport::codex_fingerprint_convergence_enabled(
|
||||
&provider.provider_type,
|
||||
provider.config.as_ref(),
|
||||
),
|
||||
"responses_websocket_enabled": responses_websocket_adapter(&provider.provider_type, provider.config.as_ref()).is_some(),
|
||||
"ops_quota_alert_enabled": ops_quota_alert_enabled,
|
||||
"created_at": endpoint_timestamp_or_now(provider.created_at_unix_ms, now_unix_secs),
|
||||
"updated_at": endpoint_timestamp_or_now(provider.updated_at_unix_secs, now_unix_secs),
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
use crate::handlers::admin::provider::oauth::provisioning::rotate_codex_credential_generation;
|
||||
use crate::handlers::admin::provider::shared::payloads::AdminProviderKeyCreateRequest;
|
||||
use crate::handlers::admin::provider::write::normalize::{
|
||||
normalize_allow_auth_channel_mismatch_formats, normalize_api_format_json_object_keys,
|
||||
@@ -216,6 +217,7 @@ pub(crate) async fn build_admin_create_provider_key_record(
|
||||
)?;
|
||||
key.created_at_unix_ms = Some(now_unix_secs);
|
||||
key.updated_at_unix_secs = Some(now_unix_secs);
|
||||
rotate_codex_credential_generation(&mut key, &provider.provider_type);
|
||||
Ok(key)
|
||||
}
|
||||
|
||||
|
||||
@@ -6,6 +6,7 @@ pub(crate) use self::update::build_admin_update_provider_key_record;
|
||||
pub(crate) use self::update::{
|
||||
admin_provider_key_update_requires_immediate_model_fetch,
|
||||
build_admin_update_provider_key_record_with_existing_keys,
|
||||
build_provider_catalog_key_admin_cas_update,
|
||||
};
|
||||
|
||||
mod batch;
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
use crate::handlers::admin::provider::oauth::provisioning::rotate_codex_credential_generation;
|
||||
use crate::handlers::admin::provider::shared::payloads::AdminProviderKeyUpdatePatch;
|
||||
use crate::handlers::admin::provider::write::normalize::{
|
||||
normalize_allow_auth_channel_mismatch_formats, normalize_api_format_json_object_keys,
|
||||
@@ -13,6 +14,7 @@ use crate::handlers::admin::shared::{
|
||||
use crate::handlers::shared::normalize_optional_api_key_concurrent_limit;
|
||||
use crate::provider_key_auth::provider_key_is_oauth_managed;
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyOAuthCredentialFence,
|
||||
StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||
};
|
||||
use aether_provider_transport::provider_types::provider_type_is_fixed;
|
||||
@@ -368,6 +370,12 @@ pub(crate) fn build_admin_update_provider_key_record_with_existing_keys(
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs());
|
||||
let credential_identity_changed = !updated.auth_type.eq_ignore_ascii_case(&existing.auth_type)
|
||||
|| updated.encrypted_api_key != existing.encrypted_api_key
|
||||
|| updated.encrypted_auth_config != existing.encrypted_auth_config;
|
||||
if credential_identity_changed {
|
||||
rotate_codex_credential_generation(&mut updated, &provider.provider_type);
|
||||
}
|
||||
Ok(updated)
|
||||
}
|
||||
|
||||
@@ -382,6 +390,49 @@ pub(crate) fn admin_provider_key_update_requires_immediate_model_fetch(
|
||||
&& (!existing.auto_fetch_models || filters_changed || locked_models_changed)
|
||||
}
|
||||
|
||||
pub(crate) fn build_provider_catalog_key_admin_cas_update(
|
||||
existing: &StoredProviderCatalogKey,
|
||||
updated: StoredProviderCatalogKey,
|
||||
provider_type: &str,
|
||||
) -> ProviderCatalogKeyAdminCasUpdate {
|
||||
let previous_generation = existing
|
||||
.upstream_metadata
|
||||
.as_ref()
|
||||
.and_then(|metadata| metadata.pointer("/codex/credential_generation"))
|
||||
.and_then(serde_json::Value::as_str);
|
||||
let next_generation = updated
|
||||
.upstream_metadata
|
||||
.as_ref()
|
||||
.and_then(|metadata| metadata.pointer("/codex/credential_generation"))
|
||||
.and_then(serde_json::Value::as_str);
|
||||
let credential_changed = existing.auth_type != updated.auth_type
|
||||
|| existing.encrypted_api_key != updated.encrypted_api_key
|
||||
|| existing.encrypted_auth_config != updated.encrypted_auth_config;
|
||||
let codex_rotation = provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("codex")
|
||||
.then(|| next_generation.filter(|next| Some(*next) != previous_generation))
|
||||
.flatten()
|
||||
.map(|generation| {
|
||||
json!({
|
||||
aether_admin::provider::quota::CODEX_CREDENTIAL_GENERATION_KEY: generation,
|
||||
})
|
||||
});
|
||||
|
||||
ProviderCatalogKeyAdminCasUpdate {
|
||||
expected_encrypted_auth_config: existing.encrypted_auth_config.clone(),
|
||||
expected_credential: ProviderCatalogKeyOAuthCredentialFence {
|
||||
encrypted_api_key: existing.encrypted_api_key.clone(),
|
||||
auth_type: existing.auth_type.clone(),
|
||||
provider_id: existing.provider_id.clone(),
|
||||
provider_type: provider_type.to_string(),
|
||||
},
|
||||
key: updated,
|
||||
codex_rotation,
|
||||
reset_oauth_runtime: credential_changed,
|
||||
}
|
||||
}
|
||||
|
||||
fn raw_secret_auth_type(value: &str) -> bool {
|
||||
matches!(
|
||||
value.trim().to_ascii_lowercase().as_str(),
|
||||
|
||||
@@ -208,6 +208,53 @@ pub(crate) fn normalize_chat_pii_redaction_config(
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn set_responses_websocket_enabled(
|
||||
config: &mut serde_json::Map<String, serde_json::Value>,
|
||||
enabled: bool,
|
||||
) -> Result<(), String> {
|
||||
let mut responses = match config.remove("responses_websocket") {
|
||||
None => serde_json::Map::new(),
|
||||
Some(serde_json::Value::Object(config)) => config,
|
||||
Some(_) => return Err("config.responses_websocket 必须是 JSON 对象".to_string()),
|
||||
};
|
||||
responses.insert("enabled".to_string(), serde_json::Value::Bool(enabled));
|
||||
config.insert(
|
||||
"responses_websocket".to_string(),
|
||||
serde_json::Value::Object(responses),
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub(crate) fn remove_responses_websocket_enabled(
|
||||
config: &mut serde_json::Map<String, serde_json::Value>,
|
||||
) {
|
||||
let Some(serde_json::Value::Object(responses)) = config.get_mut("responses_websocket") else {
|
||||
return;
|
||||
};
|
||||
responses.remove("enabled");
|
||||
if responses.is_empty() {
|
||||
config.remove("responses_websocket");
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn validate_responses_websocket_config(
|
||||
config: &serde_json::Map<String, serde_json::Value>,
|
||||
) -> Result<(), String> {
|
||||
if let Some(value) = config.get("responses_websocket") {
|
||||
let responses = value
|
||||
.as_object()
|
||||
.ok_or_else(|| "config.responses_websocket 必须是 JSON 对象".to_string())?;
|
||||
let enabled = responses
|
||||
.get("enabled")
|
||||
.ok_or_else(|| "config.responses_websocket.enabled 为必填布尔值".to_string())?;
|
||||
if !enabled.is_boolean() {
|
||||
return Err("config.responses_websocket.enabled 必须是布尔值".to_string());
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub(crate) fn validate_vertex_api_formats(
|
||||
provider_type: &str,
|
||||
auth_type: &str,
|
||||
@@ -258,7 +305,9 @@ mod tests {
|
||||
normalize_api_format_list, normalize_auth_type, normalize_auth_type_by_format,
|
||||
normalize_chat_pii_redaction_config, normalize_pool_advanced_config,
|
||||
normalize_provider_type_input, normalize_rate_multipliers,
|
||||
reconcile_allow_auth_channel_mismatch_formats, validate_vertex_api_formats,
|
||||
reconcile_allow_auth_channel_mismatch_formats, remove_responses_websocket_enabled,
|
||||
set_responses_websocket_enabled, validate_responses_websocket_config,
|
||||
validate_vertex_api_formats,
|
||||
};
|
||||
use serde_json::json;
|
||||
|
||||
@@ -317,6 +366,21 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn responses_websocket_setting_is_available_to_explicitly_enabled_providers() {
|
||||
let mut config = serde_json::Map::new();
|
||||
set_responses_websocket_enabled(&mut config, true)
|
||||
.expect("Responses setting should be accepted");
|
||||
assert_eq!(
|
||||
config.get("responses_websocket"),
|
||||
Some(&json!({"enabled": true}))
|
||||
);
|
||||
validate_responses_websocket_config(&config).expect("Responses setting should validate");
|
||||
|
||||
remove_responses_websocket_enabled(&mut config);
|
||||
assert!(config.get("responses_websocket").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn normalize_auth_type_supports_bearer() {
|
||||
assert_eq!(
|
||||
|
||||
@@ -7,6 +7,8 @@ use crate::handlers::admin::provider::shared::support::{
|
||||
use crate::handlers::admin::provider::write::normalize::normalize_chat_pii_redaction_config;
|
||||
use crate::handlers::admin::provider::write::normalize::normalize_pool_advanced_config;
|
||||
use crate::handlers::admin::provider::write::normalize::normalize_provider_type_input;
|
||||
use crate::handlers::admin::provider::write::normalize::set_responses_websocket_enabled;
|
||||
use crate::handlers::admin::provider::write::normalize::validate_responses_websocket_config;
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::handlers::admin::shared::normalize_json_object;
|
||||
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogProvider;
|
||||
@@ -140,6 +142,28 @@ pub(crate) async fn build_admin_create_provider_record(
|
||||
if let Some(value) = normalize_pool_advanced_config(payload.pool_advanced)? {
|
||||
config_map.insert("pool_advanced".to_string(), value);
|
||||
}
|
||||
if let Some(enabled) = payload.codex_fingerprint_convergence_enabled {
|
||||
if provider_type != "codex" && enabled {
|
||||
return Err(
|
||||
"codex_fingerprint_convergence_enabled 仅适用于 provider_type=codex".to_string(),
|
||||
);
|
||||
}
|
||||
if provider_type == "codex" {
|
||||
let codex_config = config_map
|
||||
.entry(crate::provider_transport::CODEX_FINGERPRINT_CONFIG_NAMESPACE.to_string())
|
||||
.or_insert_with(|| json!({}));
|
||||
let Some(codex_config) = codex_config.as_object_mut() else {
|
||||
return Err("config.codex 必须是 JSON 对象".to_string());
|
||||
};
|
||||
codex_config.insert(
|
||||
crate::provider_transport::CODEX_FINGERPRINT_ENABLED_CONFIG_KEY.to_string(),
|
||||
json!(enabled),
|
||||
);
|
||||
}
|
||||
}
|
||||
if provider_type != "codex" {
|
||||
remove_codex_fingerprint_config(&mut config_map);
|
||||
}
|
||||
if let Some(value) = normalize_json_object(payload.failover_rules, "failover_rules")? {
|
||||
config_map.insert("failover_rules".to_string(), value);
|
||||
}
|
||||
@@ -157,6 +181,10 @@ pub(crate) async fn build_admin_create_provider_record(
|
||||
config_map.insert("chat_pii_redaction".to_string(), value);
|
||||
}
|
||||
}
|
||||
if let Some(enabled) = payload.responses_websocket_enabled {
|
||||
set_responses_websocket_enabled(&mut config_map, enabled)?;
|
||||
}
|
||||
validate_responses_websocket_config(&config_map)?;
|
||||
let config = (!config_map.is_empty()).then_some(serde_json::Value::Object(config_map));
|
||||
crate::provider_transport::validate_anthropic_compatibility_profile_config(config.as_ref())
|
||||
.map_err(|_| "无效的 Anthropic compatibility profile".to_string())?;
|
||||
@@ -204,3 +232,46 @@ pub(crate) async fn build_admin_create_provider_record(
|
||||
|
||||
Ok((record, shift_existing_priorities_from))
|
||||
}
|
||||
|
||||
fn remove_codex_fingerprint_config(config_map: &mut serde_json::Map<String, serde_json::Value>) {
|
||||
let namespace = crate::provider_transport::CODEX_FINGERPRINT_CONFIG_NAMESPACE;
|
||||
let key = crate::provider_transport::CODEX_FINGERPRINT_ENABLED_CONFIG_KEY;
|
||||
let mut remove_namespace = false;
|
||||
if let Some(codex_config) = config_map
|
||||
.get_mut(namespace)
|
||||
.and_then(|value| value.as_object_mut())
|
||||
{
|
||||
codex_config.remove(key);
|
||||
remove_namespace = codex_config.is_empty();
|
||||
}
|
||||
if remove_namespace {
|
||||
config_map.remove(namespace);
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
#[test]
|
||||
fn removing_fingerprint_setting_preserves_other_codex_config() {
|
||||
let mut config = json!({
|
||||
"codex": {
|
||||
"fingerprint_convergence_enabled": true,
|
||||
"pass_through_cyber_flag_interrupt": true
|
||||
},
|
||||
"other": {"kept": true}
|
||||
})
|
||||
.as_object()
|
||||
.expect("config object")
|
||||
.clone();
|
||||
|
||||
super::remove_codex_fingerprint_config(&mut config);
|
||||
|
||||
assert_eq!(
|
||||
config["codex"],
|
||||
json!({"pass_through_cyber_flag_interrupt": true})
|
||||
);
|
||||
assert_eq!(config["other"], json!({"kept": true}));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -7,6 +7,8 @@ use crate::handlers::admin::provider::shared::support::{
|
||||
use crate::handlers::admin::provider::write::normalize::normalize_chat_pii_redaction_config;
|
||||
use crate::handlers::admin::provider::write::normalize::normalize_pool_advanced_config;
|
||||
use crate::handlers::admin::provider::write::normalize::normalize_provider_type_input;
|
||||
use crate::handlers::admin::provider::write::normalize::set_responses_websocket_enabled;
|
||||
use crate::handlers::admin::provider::write::normalize::validate_responses_websocket_config;
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::handlers::admin::shared::normalize_json_object;
|
||||
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogProvider;
|
||||
@@ -244,6 +246,32 @@ pub(crate) async fn build_admin_update_provider_record(
|
||||
}
|
||||
}
|
||||
|
||||
if fields.contains("codex_fingerprint_convergence_enabled") {
|
||||
let Some(enabled) = payload.codex_fingerprint_convergence_enabled else {
|
||||
return Err("codex_fingerprint_convergence_enabled 必须是布尔值".to_string());
|
||||
};
|
||||
if target_provider_type != "codex" && enabled {
|
||||
return Err(
|
||||
"codex_fingerprint_convergence_enabled 仅适用于 provider_type=codex".to_string(),
|
||||
);
|
||||
}
|
||||
if target_provider_type == "codex" {
|
||||
let codex_config = config_map
|
||||
.entry(crate::provider_transport::CODEX_FINGERPRINT_CONFIG_NAMESPACE.to_string())
|
||||
.or_insert_with(|| json!({}));
|
||||
let Some(codex_config) = codex_config.as_object_mut() else {
|
||||
return Err("config.codex 必须是 JSON 对象".to_string());
|
||||
};
|
||||
codex_config.insert(
|
||||
crate::provider_transport::CODEX_FINGERPRINT_ENABLED_CONFIG_KEY.to_string(),
|
||||
json!(enabled),
|
||||
);
|
||||
}
|
||||
}
|
||||
if target_provider_type != "codex" {
|
||||
remove_codex_fingerprint_config(&mut config_map);
|
||||
}
|
||||
|
||||
for (field_name, payload_value) in [
|
||||
(
|
||||
PROVIDER_MAX_TRANSFER_COUNT_CONFIG_KEY,
|
||||
@@ -311,6 +339,14 @@ pub(crate) async fn build_admin_update_provider_record(
|
||||
}
|
||||
}
|
||||
|
||||
if fields.contains("responses_websocket_enabled") {
|
||||
let enabled = payload
|
||||
.responses_websocket_enabled
|
||||
.ok_or_else(|| "responses_websocket_enabled 必须是布尔值".to_string())?;
|
||||
set_responses_websocket_enabled(&mut config_map, enabled)?;
|
||||
}
|
||||
validate_responses_websocket_config(&config_map)?;
|
||||
|
||||
updated.config = (!config_map.is_empty()).then_some(serde_json::Value::Object(config_map));
|
||||
crate::provider_transport::validate_anthropic_compatibility_profile_config(
|
||||
updated.config.as_ref(),
|
||||
@@ -322,3 +358,46 @@ pub(crate) async fn build_admin_update_provider_record(
|
||||
.map(|duration| duration.as_secs());
|
||||
Ok(updated)
|
||||
}
|
||||
|
||||
fn remove_codex_fingerprint_config(config_map: &mut serde_json::Map<String, serde_json::Value>) {
|
||||
let namespace = crate::provider_transport::CODEX_FINGERPRINT_CONFIG_NAMESPACE;
|
||||
let key = crate::provider_transport::CODEX_FINGERPRINT_ENABLED_CONFIG_KEY;
|
||||
let mut remove_namespace = false;
|
||||
if let Some(codex_config) = config_map
|
||||
.get_mut(namespace)
|
||||
.and_then(|value| value.as_object_mut())
|
||||
{
|
||||
codex_config.remove(key);
|
||||
remove_namespace = codex_config.is_empty();
|
||||
}
|
||||
if remove_namespace {
|
||||
config_map.remove(namespace);
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
#[test]
|
||||
fn removing_fingerprint_setting_preserves_other_codex_config() {
|
||||
let mut config = json!({
|
||||
"codex": {
|
||||
"fingerprint_convergence_enabled": true,
|
||||
"pass_through_cyber_flag_interrupt": true
|
||||
},
|
||||
"other": {"kept": true}
|
||||
})
|
||||
.as_object()
|
||||
.expect("config object")
|
||||
.clone();
|
||||
|
||||
super::remove_codex_fingerprint_config(&mut config);
|
||||
|
||||
assert_eq!(
|
||||
config["codex"],
|
||||
json!({"pass_through_cyber_flag_interrupt": true})
|
||||
);
|
||||
assert_eq!(config["other"], json!({"kept": true}));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -186,6 +186,15 @@ impl<'a> AdminAppState<'a> {
|
||||
self.app.update_provider_catalog_key(key).await
|
||||
}
|
||||
|
||||
pub(crate) async fn compare_and_update_provider_catalog_key_admin_state(
|
||||
&self,
|
||||
update: &aether_data_contracts::repository::provider_catalog::ProviderCatalogKeyAdminCasUpdate,
|
||||
) -> Result<bool, GatewayError> {
|
||||
self.app
|
||||
.compare_and_update_provider_catalog_key_admin_state(update)
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn compare_and_update_provider_catalog_key_adaptive_state(
|
||||
&self,
|
||||
update: &aether_data_contracts::repository::provider_catalog::ProviderCatalogKeyAdaptiveStateUpdate,
|
||||
|
||||
@@ -513,6 +513,7 @@ impl<'a> AdminAppState<'a> {
|
||||
provider_id: key.provider_id.clone(),
|
||||
provider_type: provider.provider_type.clone(),
|
||||
},
|
||||
expected_upstream_metadata_namespace: None,
|
||||
},
|
||||
)
|
||||
.await
|
||||
|
||||
@@ -4,10 +4,12 @@ use crate::api::ai::admin_endpoint_signature_parts;
|
||||
use crate::handlers::admin::admin_provider_pool_config;
|
||||
use crate::handlers::admin::model::ADMIN_EXTERNAL_MODELS_PROXY_NODE_CONFIG_KEY;
|
||||
use crate::handlers::admin::provider::endpoints_admin::payloads::AdminProviderEndpointUpdatePatch;
|
||||
use crate::handlers::admin::provider::oauth::provisioning::ensure_codex_credential_generation_rotated;
|
||||
use crate::handlers::admin::provider::shared::payloads::{
|
||||
AdminProviderCreateRequest, AdminProviderKeyCreateRequest, AdminProviderKeyUpdatePatch,
|
||||
AdminProviderUpdatePatch,
|
||||
};
|
||||
use crate::handlers::admin::provider::write::keys::build_provider_catalog_key_admin_cas_update;
|
||||
use crate::handlers::admin::shared::{
|
||||
normalize_json_array, normalize_json_object, normalize_string_list,
|
||||
};
|
||||
@@ -377,10 +379,14 @@ fn normalize_import_key_raw_payload(
|
||||
|
||||
fn apply_imported_oauth_key_credentials(
|
||||
state: &AdminAppState<'_>,
|
||||
provider_type: &str,
|
||||
previous_codex_credential_generation: Option<&str>,
|
||||
raw_key: &Map<String, Value>,
|
||||
normalized_auth_config: Option<&Value>,
|
||||
record: &mut aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey,
|
||||
) -> Result<bool, String> {
|
||||
let previous_encrypted_api_key = record.encrypted_api_key.clone();
|
||||
let previous_encrypted_auth_config = record.encrypted_auth_config.clone();
|
||||
let mut credentials_supplied = false;
|
||||
let mut api_key_supplied = false;
|
||||
if let Some(api_key_value) = raw_key.get("api_key") {
|
||||
@@ -424,10 +430,19 @@ fn apply_imported_oauth_key_credentials(
|
||||
api_key_supplied,
|
||||
);
|
||||
|
||||
let credential_material_changed = record.encrypted_api_key != previous_encrypted_api_key
|
||||
|| record.encrypted_auth_config != previous_encrypted_auth_config;
|
||||
if credentials_supplied {
|
||||
record.oauth_invalid_at_unix_secs = None;
|
||||
record.oauth_invalid_reason = None;
|
||||
}
|
||||
if credential_material_changed {
|
||||
ensure_codex_credential_generation_rotated(
|
||||
record,
|
||||
provider_type,
|
||||
previous_codex_credential_generation,
|
||||
);
|
||||
}
|
||||
|
||||
Ok(credentials_supplied)
|
||||
}
|
||||
@@ -1861,6 +1876,15 @@ impl<'a> AdminAppState<'a> {
|
||||
|
||||
if let Some(existing_index) = existing_key_index {
|
||||
let existing_key = existing_keys[existing_index].clone();
|
||||
let previous_codex_credential_generation = existing_key
|
||||
.upstream_metadata
|
||||
.as_ref()
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|metadata| metadata.get("codex"))
|
||||
.and_then(|codex| {
|
||||
aether_admin::provider::quota::codex_credential_generation(Some(codex))
|
||||
})
|
||||
.map(ToOwned::to_owned);
|
||||
match merge_mode {
|
||||
AdminImportMergeMode::Skip => {
|
||||
stats.keys.skipped += 1;
|
||||
@@ -1890,6 +1914,8 @@ impl<'a> AdminAppState<'a> {
|
||||
let oauth_credentials_supplied = if auth_type == "oauth" {
|
||||
invalid!(apply_imported_oauth_key_credentials(
|
||||
self,
|
||||
&provider.provider_type,
|
||||
previous_codex_credential_generation.as_deref(),
|
||||
&raw_key,
|
||||
normalized_auth_config.as_ref(),
|
||||
&mut updated,
|
||||
@@ -1903,8 +1929,31 @@ impl<'a> AdminAppState<'a> {
|
||||
imported_key.fingerprint.clone(),
|
||||
"fingerprint",
|
||||
));
|
||||
let Some(mut persisted) =
|
||||
self.update_provider_catalog_key(&updated).await?
|
||||
let admin_update = build_provider_catalog_key_admin_cas_update(
|
||||
&existing_key,
|
||||
updated.clone(),
|
||||
&provider.provider_type,
|
||||
);
|
||||
if !self
|
||||
.compare_and_update_provider_catalog_key_admin_state(&admin_update)
|
||||
.await?
|
||||
{
|
||||
return Ok(Err((
|
||||
http::StatusCode::CONFLICT,
|
||||
json!({
|
||||
"detail": format!(
|
||||
"Provider '{provider_name}' 的 Key 已被其他请求更新,请重试"
|
||||
)
|
||||
}),
|
||||
)));
|
||||
}
|
||||
let Some(mut persisted) = self
|
||||
.read_provider_catalog_keys_by_ids(std::slice::from_ref(
|
||||
&updated.id,
|
||||
))
|
||||
.await?
|
||||
.into_iter()
|
||||
.next()
|
||||
else {
|
||||
return Ok(Err(invalid_request(format!(
|
||||
"更新 Provider '{provider_name}' 的 Key 失败"
|
||||
@@ -1926,16 +1975,15 @@ impl<'a> AdminAppState<'a> {
|
||||
persisted = reloaded;
|
||||
}
|
||||
if oauth_credentials_supplied {
|
||||
if !self
|
||||
.clear_provider_catalog_key_oauth_invalid_marker(&updated.id)
|
||||
.await?
|
||||
{
|
||||
return Ok(Err(invalid_request(format!(
|
||||
"更新 Provider '{provider_name}' 的 Key 失败"
|
||||
))));
|
||||
}
|
||||
let Some(reloaded) = self
|
||||
.reset_provider_catalog_key_recovery_state(&updated.id)
|
||||
.reset_provider_catalog_key_recovery_state_fenced(
|
||||
&updated.id,
|
||||
updated.encrypted_auth_config.as_deref().ok_or_else(|| {
|
||||
GatewayError::Internal(format!(
|
||||
"OAuth Provider '{provider_name}' imported without auth_config"
|
||||
))
|
||||
})?,
|
||||
)
|
||||
.await?
|
||||
else {
|
||||
return Ok(Err(invalid_request(format!(
|
||||
@@ -1975,6 +2023,8 @@ impl<'a> AdminAppState<'a> {
|
||||
let oauth_credentials_supplied = if auth_type == "oauth" {
|
||||
invalid!(apply_imported_oauth_key_credentials(
|
||||
self,
|
||||
&provider.provider_type,
|
||||
None,
|
||||
&raw_key,
|
||||
normalized_auth_config.as_ref(),
|
||||
&mut record,
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
mod body_buffer;
|
||||
mod local;
|
||||
mod websocket;
|
||||
|
||||
use self::body_buffer::{
|
||||
buffer_and_normalize_request_body, build_request_body_buffer_error_response,
|
||||
@@ -8,6 +9,7 @@ use self::body_buffer::{
|
||||
use self::local::{
|
||||
maybe_build_local_admin_proxy_response, maybe_build_local_internal_proxy_response,
|
||||
};
|
||||
pub(crate) use self::websocket::responses::responses_websocket;
|
||||
use super::internal::resolve_local_proxy_execution_path;
|
||||
pub(crate) use super::public::matches_model_mapping_for_models;
|
||||
use crate::ai_serving::api::{
|
||||
|
||||
@@ -0,0 +1,527 @@
|
||||
//! Authenticated public WebSocket upgrade admission shared by AI adapters.
|
||||
|
||||
use std::future::Future;
|
||||
use std::net::{IpAddr, SocketAddr};
|
||||
|
||||
use axum::body::Body;
|
||||
use axum::extract::ws::{WebSocket, WebSocketUpgrade};
|
||||
use axum::http::header::{
|
||||
AUTHORIZATION, CONNECTION, COOKIE, HOST, PROXY_AUTHORIZATION, TE, TRAILER, TRANSFER_ENCODING,
|
||||
UPGRADE,
|
||||
};
|
||||
use axum::http::uri::PathAndQuery;
|
||||
use axum::http::{HeaderMap, HeaderName, Method, Response, StatusCode, Uri};
|
||||
use tracing::{info, warn};
|
||||
|
||||
use crate::api::response::{
|
||||
build_local_auth_rejection_response, build_local_http_error_response,
|
||||
build_local_overloaded_response,
|
||||
};
|
||||
use crate::control::{
|
||||
trusted_auth_local_rejection, GatewayControlDecision, GatewayCredentialCarrier,
|
||||
GatewayLocalAuthRejection,
|
||||
};
|
||||
use crate::handlers::proxy::websocket::session::{WebSocketSessionLimits, WEBSOCKET_LOG_TRANSPORT};
|
||||
use crate::handlers::shared::ip_rules_allow;
|
||||
use crate::headers::{effective_client_ip, extract_or_generate_trace_id};
|
||||
use crate::router::RequestAdmissionError;
|
||||
use crate::{AppState, GatewayError};
|
||||
|
||||
/// Request facts that survive the HTTP Upgrade and are needed by a protocol
|
||||
/// adapter for planning, rate limiting, and connection-scoped audit logs.
|
||||
pub(crate) struct WebSocketRequestContext {
|
||||
pub(crate) trace_id: String,
|
||||
pub(crate) headers: HeaderMap,
|
||||
pub(crate) uri: Uri,
|
||||
pub(crate) remote_addr: SocketAddr,
|
||||
/// Effective client IP resolved once from the authenticated Upgrade. Every
|
||||
/// turn re-checks live API-key/admin IP policy against this immutable fact.
|
||||
pub(crate) client_ip: IpAddr,
|
||||
pub(crate) decision: GatewayControlDecision,
|
||||
/// Held for the lifetime of the upgraded socket. The Responses session
|
||||
/// polls its health and closes the client when a distributed lease is
|
||||
/// revoked or expires.
|
||||
pub(crate) websocket_connection_permit: Option<aether_runtime::AdmissionPermit>,
|
||||
}
|
||||
|
||||
/// Adapter-specific wording and event identifiers for generic upgrade checks.
|
||||
#[derive(Clone, Copy)]
|
||||
pub(crate) struct WebSocketIngressSpec {
|
||||
pub(crate) route_unavailable_message: &'static str,
|
||||
}
|
||||
|
||||
/// Performs the HTTP-only part of an AI WebSocket request.
|
||||
///
|
||||
/// The ordinary request permit covers only the HTTP Upgrade window. A
|
||||
/// dedicated WebSocket connection permit is held for the socket lifetime so
|
||||
/// idle clients cannot consume capacity reserved for normal HTTP requests.
|
||||
pub(crate) async fn upgrade_authenticated_ai_websocket<F, Fut>(
|
||||
state: AppState,
|
||||
remote_addr: SocketAddr,
|
||||
ws: WebSocketUpgrade,
|
||||
headers: HeaderMap,
|
||||
uri: Uri,
|
||||
limits: WebSocketSessionLimits,
|
||||
spec: WebSocketIngressSpec,
|
||||
run_session: F,
|
||||
) -> Result<Response<Body>, GatewayError>
|
||||
where
|
||||
F: FnOnce(WebSocket, AppState, WebSocketRequestContext) -> Fut + Send + 'static,
|
||||
Fut: Future<Output = ()> + Send + 'static,
|
||||
{
|
||||
let trace_id = extract_or_generate_trace_id(&headers);
|
||||
let client_ip = effective_client_ip(&headers, &remote_addr);
|
||||
if state.admin_security_ip_blacklisted(client_ip).await? {
|
||||
return build_local_http_error_response(
|
||||
&trace_id,
|
||||
None,
|
||||
StatusCode::FORBIDDEN,
|
||||
"当前 IP 已被禁止访问",
|
||||
);
|
||||
}
|
||||
|
||||
let request_context = crate::control::resolve_public_request_context(
|
||||
&state,
|
||||
&Method::GET,
|
||||
&uri,
|
||||
&headers,
|
||||
&trace_id,
|
||||
)
|
||||
.await?;
|
||||
let Some(mut decision) = request_context.control_decision else {
|
||||
return build_local_http_error_response(
|
||||
&trace_id,
|
||||
None,
|
||||
StatusCode::NOT_FOUND,
|
||||
spec.route_unavailable_message,
|
||||
);
|
||||
};
|
||||
if let Some(rejection) = trusted_auth_local_rejection(Some(&decision), &headers) {
|
||||
return build_local_auth_rejection_response(&trace_id, Some(&decision), &rejection);
|
||||
}
|
||||
// Browsers attach cookies to WebSocket handshakes automatically and the
|
||||
// WebSocket API does not let callers add an Authorization header. A
|
||||
// cookie-only public upgrade would therefore be vulnerable to cross-site
|
||||
// WebSocket hijacking unless every deployment maintained an Origin
|
||||
// allowlist. Explicit API-key/bearer credentials (or trusted internal
|
||||
// auth resolved by the control plane) remain supported.
|
||||
if !websocket_credential_carrier_is_allowed(decision.gateway_credential_carrier) {
|
||||
warn!(
|
||||
event_name = "ai_websocket_cookie_only_auth_rejected",
|
||||
log_type = "security",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %trace_id,
|
||||
client_ip = %client_ip,
|
||||
"gateway rejected cookie-only public WebSocket authentication"
|
||||
);
|
||||
return build_local_auth_rejection_response(
|
||||
&trace_id,
|
||||
Some(&decision),
|
||||
&GatewayLocalAuthRejection::InvalidApiKey,
|
||||
);
|
||||
}
|
||||
let Some(auth_context) = decision.auth_context.as_ref() else {
|
||||
return build_local_auth_rejection_response(
|
||||
&trace_id,
|
||||
Some(&decision),
|
||||
&GatewayLocalAuthRejection::InvalidApiKey,
|
||||
);
|
||||
};
|
||||
if !auth_context.access_allowed
|
||||
|| auth_context.user_id.trim().is_empty()
|
||||
|| auth_context.api_key_id.trim().is_empty()
|
||||
{
|
||||
return build_local_auth_rejection_response(
|
||||
&trace_id,
|
||||
Some(&decision),
|
||||
&GatewayLocalAuthRejection::InvalidApiKey,
|
||||
);
|
||||
}
|
||||
if !ip_rules_allow(auth_context.ip_rules.as_deref(), client_ip) {
|
||||
return build_local_auth_rejection_response(
|
||||
&trace_id,
|
||||
Some(&decision),
|
||||
&GatewayLocalAuthRejection::IpNotAllowed {
|
||||
remote_ip: client_ip.to_string(),
|
||||
},
|
||||
);
|
||||
}
|
||||
|
||||
let request_permit = match state.try_acquire_request_permit().await {
|
||||
Ok(permit) => permit,
|
||||
Err(error) => {
|
||||
return websocket_admission_error_response(
|
||||
&trace_id,
|
||||
&decision,
|
||||
Some(uri.path()),
|
||||
error,
|
||||
)
|
||||
}
|
||||
};
|
||||
let websocket_connection_permit = match state.try_acquire_websocket_connection_permit().await {
|
||||
Ok(permit) => permit,
|
||||
Err(error) => {
|
||||
return websocket_admission_error_response(
|
||||
&trace_id,
|
||||
&decision,
|
||||
Some(uri.path()),
|
||||
error,
|
||||
)
|
||||
}
|
||||
};
|
||||
|
||||
// Authentication has consumed the downstream credentials. From this
|
||||
// point on the URI and headers become planner input, so retain neither an
|
||||
// API key from the query string nor client authentication/handshake
|
||||
// headers. Provider authentication is added independently by the
|
||||
// planner and is therefore unaffected by this boundary.
|
||||
let uri = websocket_planning_uri(&uri);
|
||||
decision.public_query_string = uri.query().map(ToOwned::to_owned);
|
||||
let headers = websocket_planning_headers(headers);
|
||||
let context = WebSocketRequestContext {
|
||||
trace_id,
|
||||
headers,
|
||||
uri,
|
||||
remote_addr,
|
||||
client_ip,
|
||||
decision,
|
||||
websocket_connection_permit,
|
||||
};
|
||||
Ok(ws
|
||||
.max_frame_size(limits.max_frame_size)
|
||||
.max_message_size(limits.max_message_size)
|
||||
.on_upgrade(move |socket| async move {
|
||||
drop(request_permit);
|
||||
run_session(socket, state, context).await;
|
||||
}))
|
||||
}
|
||||
|
||||
fn websocket_credential_carrier_is_allowed(carrier: Option<GatewayCredentialCarrier>) -> bool {
|
||||
carrier != Some(GatewayCredentialCarrier::CookieHeader)
|
||||
}
|
||||
|
||||
fn websocket_planning_uri(uri: &Uri) -> Uri {
|
||||
let Some(query) = uri.query() else {
|
||||
return uri.clone();
|
||||
};
|
||||
let mut retained = Vec::new();
|
||||
let mut removed_sensitive_value = false;
|
||||
for (name, value) in url::form_urlencoded::parse(query.as_bytes()) {
|
||||
if websocket_query_parameter_is_sensitive(name.as_ref()) {
|
||||
removed_sensitive_value = true;
|
||||
} else {
|
||||
retained.push((name.into_owned(), value.into_owned()));
|
||||
}
|
||||
}
|
||||
if !removed_sensitive_value {
|
||||
return uri.clone();
|
||||
}
|
||||
|
||||
let mut serializer = url::form_urlencoded::Serializer::new(String::new());
|
||||
serializer.extend_pairs(retained.iter().map(|(name, value)| (name, value)));
|
||||
let retained_query = serializer.finish();
|
||||
let path_and_query = if retained_query.is_empty() {
|
||||
uri.path().to_string()
|
||||
} else {
|
||||
format!("{}?{retained_query}", uri.path())
|
||||
};
|
||||
let path_and_query = path_and_query
|
||||
.parse::<PathAndQuery>()
|
||||
.expect("a valid URI path plus form-encoded query must remain valid");
|
||||
let mut parts = uri.clone().into_parts();
|
||||
parts.path_and_query = Some(path_and_query);
|
||||
Uri::from_parts(parts).expect("replacing only path-and-query must preserve a valid URI")
|
||||
}
|
||||
|
||||
fn websocket_query_parameter_is_sensitive(name: &str) -> bool {
|
||||
matches!(
|
||||
name.to_ascii_lowercase().as_str(),
|
||||
"key" | "api_key" | "api-key" | "access_token" | "authorization" | "token"
|
||||
)
|
||||
}
|
||||
|
||||
fn websocket_planning_headers(mut headers: HeaderMap) -> HeaderMap {
|
||||
// RFC 9110 permits Connection to name additional hop-by-hop fields. Read
|
||||
// those names before removing Connection itself.
|
||||
let connection_scoped_names = headers
|
||||
.get_all(CONNECTION)
|
||||
.iter()
|
||||
.filter_map(|value| value.to_str().ok())
|
||||
.flat_map(|value| value.split(','))
|
||||
.filter_map(|name| HeaderName::from_bytes(name.trim().as_bytes()).ok())
|
||||
.collect::<Vec<_>>();
|
||||
for name in connection_scoped_names {
|
||||
headers.remove(name);
|
||||
}
|
||||
|
||||
for name in [
|
||||
AUTHORIZATION,
|
||||
CONNECTION,
|
||||
COOKIE,
|
||||
HOST,
|
||||
PROXY_AUTHORIZATION,
|
||||
TE,
|
||||
TRAILER,
|
||||
TRANSFER_ENCODING,
|
||||
UPGRADE,
|
||||
] {
|
||||
headers.remove(name);
|
||||
}
|
||||
for name in [
|
||||
"api-key",
|
||||
"keep-alive",
|
||||
"proxy-connection",
|
||||
"x-api-key",
|
||||
"x-goog-api-key",
|
||||
crate::constants::GATEWAY_HEADER,
|
||||
crate::constants::TRUSTED_AUTH_USER_ID_HEADER,
|
||||
crate::constants::TRUSTED_AUTH_API_KEY_ID_HEADER,
|
||||
crate::constants::TRUSTED_AUTH_BALANCE_HEADER,
|
||||
crate::constants::TRUSTED_AUTH_ACCESS_ALLOWED_HEADER,
|
||||
crate::constants::TRUSTED_ADMIN_USER_ID_HEADER,
|
||||
crate::constants::TRUSTED_ADMIN_USER_ROLE_HEADER,
|
||||
crate::constants::TRUSTED_ADMIN_SESSION_ID_HEADER,
|
||||
crate::constants::TRUSTED_ADMIN_MANAGEMENT_TOKEN_ID_HEADER,
|
||||
] {
|
||||
headers.remove(name);
|
||||
}
|
||||
let websocket_managed_names = headers
|
||||
.keys()
|
||||
.filter(|name| name.as_str().starts_with("sec-websocket-"))
|
||||
.cloned()
|
||||
.collect::<Vec<_>>();
|
||||
for name in websocket_managed_names {
|
||||
headers.remove(name);
|
||||
}
|
||||
headers
|
||||
}
|
||||
|
||||
fn websocket_admission_error_response(
|
||||
trace_id: &str,
|
||||
decision: &GatewayControlDecision,
|
||||
request_path: Option<&str>,
|
||||
error: RequestAdmissionError,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
match error {
|
||||
RequestAdmissionError::Local(aether_runtime::ConcurrencyError::Saturated {
|
||||
gate,
|
||||
limit,
|
||||
})
|
||||
| RequestAdmissionError::Distributed(
|
||||
aether_runtime_state::RuntimeSemaphoreError::Saturated { gate, limit },
|
||||
)
|
||||
| RequestAdmissionError::Distributed(
|
||||
aether_runtime_state::RuntimeSemaphoreError::Unavailable { gate, limit, .. },
|
||||
) => build_local_overloaded_response(trace_id, Some(decision), request_path, gate, limit),
|
||||
RequestAdmissionError::Local(aether_runtime::ConcurrencyError::Closed { gate }) => Err(
|
||||
GatewayError::Internal(format!("gateway concurrency gate {gate} is closed")),
|
||||
),
|
||||
RequestAdmissionError::Distributed(
|
||||
aether_runtime_state::RuntimeSemaphoreError::InvalidConfiguration(message),
|
||||
) => Err(GatewayError::Internal(message)),
|
||||
}
|
||||
}
|
||||
|
||||
/// Connection-level access log fields which are independent of a protocol's
|
||||
/// per-turn usage lifecycle.
|
||||
#[derive(Clone, Copy)]
|
||||
pub(crate) struct WebSocketConnectionLogSpec {
|
||||
pub(crate) opened_event_name: &'static str,
|
||||
pub(crate) closed_event_name: &'static str,
|
||||
pub(crate) opened_message: &'static str,
|
||||
pub(crate) closed_message: &'static str,
|
||||
pub(crate) execution_path: &'static str,
|
||||
pub(crate) provider_type: &'static str,
|
||||
}
|
||||
|
||||
pub(crate) struct WebSocketConnectionLog {
|
||||
spec: WebSocketConnectionLogSpec,
|
||||
trace_id: String,
|
||||
remote_addr: SocketAddr,
|
||||
path: String,
|
||||
route_class: String,
|
||||
user_id: String,
|
||||
api_key_id: String,
|
||||
started_at: std::time::Instant,
|
||||
}
|
||||
|
||||
impl WebSocketConnectionLog {
|
||||
pub(crate) fn new(context: &WebSocketRequestContext, spec: WebSocketConnectionLogSpec) -> Self {
|
||||
let auth_context = context.decision.auth_context.as_ref();
|
||||
Self {
|
||||
spec,
|
||||
trace_id: context.trace_id.clone(),
|
||||
remote_addr: context.remote_addr,
|
||||
path: context.uri.path().to_string(),
|
||||
route_class: context
|
||||
.decision
|
||||
.route_class
|
||||
.as_deref()
|
||||
.unwrap_or("ai_public")
|
||||
.to_string(),
|
||||
user_id: auth_context
|
||||
.map(|auth_context| auth_context.user_id.clone())
|
||||
.unwrap_or_else(|| "-".to_string()),
|
||||
api_key_id: auth_context
|
||||
.map(|auth_context| auth_context.api_key_id.clone())
|
||||
.unwrap_or_else(|| "-".to_string()),
|
||||
started_at: std::time::Instant::now(),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn log_opened(&self) {
|
||||
info!(
|
||||
event_name = self.spec.opened_event_name,
|
||||
log_type = "access",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
status = "upgraded",
|
||||
status_code = 101u16,
|
||||
trace_id = %self.trace_id,
|
||||
remote_addr = %self.remote_addr,
|
||||
method = "GET",
|
||||
path = %self.path,
|
||||
user_id = %self.user_id,
|
||||
api_key_id = %self.api_key_id,
|
||||
route_class = %self.route_class,
|
||||
execution_path = self.spec.execution_path,
|
||||
provider_type = self.spec.provider_type,
|
||||
message = self.spec.opened_message,
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for WebSocketConnectionLog {
|
||||
fn drop(&mut self) {
|
||||
info!(
|
||||
event_name = self.spec.closed_event_name,
|
||||
log_type = "access",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
status = "closed",
|
||||
status_code = 101u16,
|
||||
trace_id = %self.trace_id,
|
||||
remote_addr = %self.remote_addr,
|
||||
method = "GET",
|
||||
path = %self.path,
|
||||
user_id = %self.user_id,
|
||||
api_key_id = %self.api_key_id,
|
||||
route_class = %self.route_class,
|
||||
execution_path = self.spec.execution_path,
|
||||
provider_type = self.spec.provider_type,
|
||||
elapsed_ms = self.started_at.elapsed().as_millis() as u64,
|
||||
message = self.spec.closed_message,
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use axum::http::header::{
|
||||
AUTHORIZATION, CONNECTION, COOKIE, HOST, ORIGIN, SEC_WEBSOCKET_KEY, UPGRADE, USER_AGENT,
|
||||
};
|
||||
use axum::http::{HeaderMap, HeaderValue, Uri};
|
||||
|
||||
use super::{
|
||||
websocket_credential_carrier_is_allowed, websocket_planning_headers, websocket_planning_uri,
|
||||
};
|
||||
use crate::control::GatewayCredentialCarrier;
|
||||
|
||||
#[test]
|
||||
fn planning_uri_removes_query_credentials_without_losing_safe_parameters() {
|
||||
let uri: Uri = "/v1/responses?key=downstream-secret&client_hint=a%20b&token=also-secret"
|
||||
.parse()
|
||||
.expect("request URI should parse");
|
||||
|
||||
let sanitized = websocket_planning_uri(&uri);
|
||||
|
||||
assert_eq!(sanitized.path(), "/v1/responses");
|
||||
assert_eq!(sanitized.query(), Some("client_hint=a+b"));
|
||||
assert!(!sanitized.to_string().contains("downstream-secret"));
|
||||
assert!(!sanitized.to_string().contains("also-secret"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn planning_uri_leaves_an_uncredentialed_query_byte_for_byte_unchanged() {
|
||||
let uri: Uri = "/v1/responses?client_hint=a%20b&empty="
|
||||
.parse()
|
||||
.expect("request URI should parse");
|
||||
|
||||
assert_eq!(websocket_planning_uri(&uri), uri);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn planning_headers_drop_client_auth_cookie_and_websocket_transport_state() {
|
||||
let mut headers = HeaderMap::new();
|
||||
headers.insert(
|
||||
AUTHORIZATION,
|
||||
HeaderValue::from_static("Bearer client-secret"),
|
||||
);
|
||||
headers.insert(COOKIE, HeaderValue::from_static("session=client-secret"));
|
||||
headers.insert("x-api-key", HeaderValue::from_static("client-secret"));
|
||||
headers.insert(HOST, HeaderValue::from_static("gateway.example"));
|
||||
headers.insert(
|
||||
CONNECTION,
|
||||
HeaderValue::from_static("keep-alive, Upgrade, x-connection-secret"),
|
||||
);
|
||||
headers.insert(UPGRADE, HeaderValue::from_static("websocket"));
|
||||
headers.insert(SEC_WEBSOCKET_KEY, HeaderValue::from_static("handshake-key"));
|
||||
headers.insert(
|
||||
"sec-websocket-future-field",
|
||||
HeaderValue::from_static("future-handshake-value"),
|
||||
);
|
||||
headers.insert(
|
||||
"x-connection-secret",
|
||||
HeaderValue::from_static("connection-secret"),
|
||||
);
|
||||
headers.insert(ORIGIN, HeaderValue::from_static("https://client.example"));
|
||||
headers.insert(USER_AGENT, HeaderValue::from_static("codex-cli/test"));
|
||||
headers.insert("x-client-hint", HeaderValue::from_static("safe"));
|
||||
|
||||
let sanitized = websocket_planning_headers(headers);
|
||||
|
||||
for name in [
|
||||
AUTHORIZATION.as_str(),
|
||||
COOKIE.as_str(),
|
||||
"x-api-key",
|
||||
HOST.as_str(),
|
||||
CONNECTION.as_str(),
|
||||
UPGRADE.as_str(),
|
||||
SEC_WEBSOCKET_KEY.as_str(),
|
||||
"sec-websocket-future-field",
|
||||
"x-connection-secret",
|
||||
] {
|
||||
assert!(sanitized.get(name).is_none(), "{name} must not survive");
|
||||
}
|
||||
assert_eq!(
|
||||
sanitized.get(ORIGIN),
|
||||
Some(&HeaderValue::from_static("https://client.example"))
|
||||
);
|
||||
assert_eq!(
|
||||
sanitized.get(USER_AGENT),
|
||||
Some(&HeaderValue::from_static("codex-cli/test"))
|
||||
);
|
||||
assert_eq!(
|
||||
sanitized.get("x-client-hint"),
|
||||
Some(&HeaderValue::from_static("safe"))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn websocket_auth_requires_an_explicit_credential_instead_of_cookie_only() {
|
||||
assert!(!websocket_credential_carrier_is_allowed(Some(
|
||||
GatewayCredentialCarrier::CookieHeader
|
||||
)));
|
||||
for carrier in [
|
||||
None,
|
||||
Some(GatewayCredentialCarrier::AuthorizationBearer),
|
||||
Some(GatewayCredentialCarrier::XApiKey),
|
||||
Some(GatewayCredentialCarrier::ApiKey),
|
||||
Some(GatewayCredentialCarrier::XGoogApiKey),
|
||||
Some(GatewayCredentialCarrier::QueryKey),
|
||||
] {
|
||||
assert!(websocket_credential_carrier_is_allowed(carrier));
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,12 @@
|
||||
//! Shared infrastructure for public AI WebSocket bridges.
|
||||
//!
|
||||
//! Protocol adapters live below [`responses`]. This layer deliberately owns
|
||||
//! only transport concerns that are common to future adapters: authenticated
|
||||
//! upgrade admission, connection limits, upstream handshakes, and frame
|
||||
//! conversion. It does not interpret provider events or make routing
|
||||
//! decisions.
|
||||
|
||||
pub(crate) mod ingress;
|
||||
pub(crate) mod responses;
|
||||
pub(crate) mod session;
|
||||
pub(crate) mod transport;
|
||||
@@ -0,0 +1,234 @@
|
||||
//! Provider-specific hooks for the standard Responses WebSocket session.
|
||||
|
||||
use async_trait::async_trait;
|
||||
use serde_json::Value;
|
||||
|
||||
use super::adapters::CODEX_RESPONSES_WEBSOCKET_ADAPTER;
|
||||
use crate::ai_serving::AiExecutionDecision;
|
||||
use crate::handlers::proxy::websocket::transport::UpstreamWebSocketErrorCodes;
|
||||
use crate::orchestration::ResponsesWebSocketAdapter;
|
||||
use crate::AppState;
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
pub(super) struct ResponsesWebSocketDrainDirective {
|
||||
pub(super) error_code: &'static str,
|
||||
/// The terminal upstream event may be replayed only when the session has
|
||||
/// not exposed any standard Responses event to the client.
|
||||
pub(super) retry_current_turn: bool,
|
||||
/// When present, the exhausted provider key remains excluded from later
|
||||
/// turns on this client socket until the upstream's reported reset time.
|
||||
pub(super) retry_exclusion_until_unix_secs: Option<u64>,
|
||||
}
|
||||
|
||||
/// Provider-specific observation produced while relaying an upstream frame.
|
||||
/// The session can make the retry/drain decision synchronously, while the
|
||||
/// optional persistence sink runs outside the frame-forwarding path.
|
||||
#[derive(Debug, Clone)]
|
||||
pub(super) struct ResponsesWebSocketAdapterObservation {
|
||||
pub(super) drain: Option<ResponsesWebSocketDrainDirective>,
|
||||
pub(super) quota_metadata: Option<Value>,
|
||||
}
|
||||
|
||||
/// Provider identity used by the shared session's temporary exclusion table.
|
||||
/// The session does not need to know how a provider derives its account id.
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub(super) struct ResponsesWebSocketExclusionIdentity {
|
||||
pub(super) account_id: Option<String>,
|
||||
}
|
||||
|
||||
/// Whether receiving an upstream event still leaves the active client turn
|
||||
/// safe to replay on a freshly bound upstream. The shared session keeps the
|
||||
/// conservative default; provider adapters may explicitly whitelist their
|
||||
/// documented, pre-response advisory events.
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub(super) enum ResponsesWebSocketRebindSafety {
|
||||
Safe,
|
||||
Unsafe { reason: &'static str },
|
||||
}
|
||||
|
||||
/// How an upstream text frame crosses the public Responses WebSocket boundary.
|
||||
///
|
||||
/// The normal path is deliberately byte-opaque: callers forward the parsed
|
||||
/// frame's original text without rebuilding it from a gateway-owned schema.
|
||||
/// Codex is the only adapter that may peel its documented private batch
|
||||
/// envelope. Even then, the retained events are borrowed whole so unknown
|
||||
/// `response.*` event types and unknown fields survive unchanged.
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub(super) enum ResponsesWebSocketRelayDirective<'a> {
|
||||
/// Forward the provider frame's original text exactly as received.
|
||||
ForwardOriginal,
|
||||
/// The provider frame was a private batch envelope. Forward each retained
|
||||
/// event in document order by serializing the complete borrowed value.
|
||||
ForwardEvents(Vec<&'a Value>),
|
||||
/// The entire frame was an explicitly recognized provider-private
|
||||
/// envelope and therefore has no public event to relay.
|
||||
SuppressProviderPrivate,
|
||||
}
|
||||
|
||||
/// Boundary between the standard Responses protocol engine and provider
|
||||
/// behavior. Adapters receive already-planned provider requests; they never
|
||||
/// own public WebSocket parsing, turn accounting, or model scheduling.
|
||||
#[async_trait]
|
||||
pub(super) trait ResponsesWebSocketProtocolAdapter: Send + Sync {
|
||||
fn kind(&self) -> ResponsesWebSocketAdapter;
|
||||
|
||||
fn upstream_errors(&self) -> UpstreamWebSocketErrorCodes;
|
||||
|
||||
/// Adds provider-specific metadata to an otherwise standard Responses
|
||||
/// stream report. The event payload is never rewritten for the client.
|
||||
fn decorate_turn_report_context(&self, report_context: &mut Option<Value>, event: &Value);
|
||||
|
||||
/// Whether this adapter needs the shared session to parse each upstream
|
||||
/// text event before normal turn accounting runs.
|
||||
fn observes_upstream_events(&self) -> bool;
|
||||
|
||||
/// Classifies whether a received upstream event can be followed by a
|
||||
/// transparent quota-driven rebind. An adapter must return `Safe` only
|
||||
/// for events that neither create public Responses state nor make a replay
|
||||
/// observably ambiguous to the client.
|
||||
fn rebind_safety_for_upstream_event(&self, event: &Value) -> ResponsesWebSocketRebindSafety;
|
||||
|
||||
/// Selects the public relay shape without projecting a provider event
|
||||
/// through an Aether-owned field or event-type allowlist.
|
||||
fn relay_directive_for_upstream_event<'a>(
|
||||
&self,
|
||||
_event: &'a Value,
|
||||
) -> ResponsesWebSocketRelayDirective<'a> {
|
||||
ResponsesWebSocketRelayDirective::ForwardOriginal
|
||||
}
|
||||
|
||||
/// Lets an adapter classify provider-only events. Returning a directive
|
||||
/// asks the shared session to drain after the active standard response.
|
||||
fn observe_upstream_event(&self, event: &Value)
|
||||
-> Option<ResponsesWebSocketAdapterObservation>;
|
||||
|
||||
fn exhaustion_exclusion_identity(
|
||||
&self,
|
||||
_decision: &AiExecutionDecision,
|
||||
) -> Option<ResponsesWebSocketExclusionIdentity> {
|
||||
None
|
||||
}
|
||||
|
||||
/// Persists an adapter observation outside the frame-forwarding path.
|
||||
async fn persist_upstream_observation(
|
||||
&self,
|
||||
state: &AppState,
|
||||
trace_id: &str,
|
||||
report_context: Option<&Value>,
|
||||
observation: ResponsesWebSocketAdapterObservation,
|
||||
);
|
||||
}
|
||||
|
||||
pub(super) fn resolve_responses_websocket_adapter(
|
||||
kind: ResponsesWebSocketAdapter,
|
||||
) -> &'static dyn ResponsesWebSocketProtocolAdapter {
|
||||
match kind {
|
||||
ResponsesWebSocketAdapter::Standard => &STANDARD_RESPONSES_WEBSOCKET_ADAPTER,
|
||||
ResponsesWebSocketAdapter::Codex => &CODEX_RESPONSES_WEBSOCKET_ADAPTER,
|
||||
}
|
||||
}
|
||||
|
||||
struct StandardResponsesWebSocketAdapter;
|
||||
|
||||
const STANDARD_UPSTREAM_WEBSOCKET_ERRORS: UpstreamWebSocketErrorCodes =
|
||||
UpstreamWebSocketErrorCodes {
|
||||
upstream_url_missing: "responses_upstream_url_missing",
|
||||
upstream_url_invalid: "responses_upstream_url_invalid",
|
||||
headers_invalid: "responses_websocket_headers_invalid",
|
||||
client_build_failed: "responses_websocket_client_build_failed",
|
||||
proxy_invalid: "responses_websocket_proxy_invalid",
|
||||
tunnel_proxy_unsupported: "responses_websocket_tunnel_proxy_unsupported",
|
||||
handshake_failed: "responses_websocket_handshake_failed",
|
||||
upgrade_rejected: "responses_websocket_upgrade_rejected",
|
||||
upgrade_failed: "responses_websocket_upgrade_failed",
|
||||
};
|
||||
|
||||
static STANDARD_RESPONSES_WEBSOCKET_ADAPTER: StandardResponsesWebSocketAdapter =
|
||||
StandardResponsesWebSocketAdapter;
|
||||
|
||||
#[async_trait]
|
||||
impl ResponsesWebSocketProtocolAdapter for StandardResponsesWebSocketAdapter {
|
||||
fn kind(&self) -> ResponsesWebSocketAdapter {
|
||||
ResponsesWebSocketAdapter::Standard
|
||||
}
|
||||
|
||||
fn upstream_errors(&self) -> UpstreamWebSocketErrorCodes {
|
||||
STANDARD_UPSTREAM_WEBSOCKET_ERRORS
|
||||
}
|
||||
|
||||
fn decorate_turn_report_context(&self, _report_context: &mut Option<Value>, _event: &Value) {}
|
||||
|
||||
fn observes_upstream_events(&self) -> bool {
|
||||
false
|
||||
}
|
||||
|
||||
fn rebind_safety_for_upstream_event(&self, event: &Value) -> ResponsesWebSocketRebindSafety {
|
||||
let reason = if is_standard_responses_event(event) {
|
||||
"standard_response_event"
|
||||
} else {
|
||||
"unrecognized_upstream_event"
|
||||
};
|
||||
ResponsesWebSocketRebindSafety::Unsafe { reason }
|
||||
}
|
||||
|
||||
fn observe_upstream_event(
|
||||
&self,
|
||||
_event: &Value,
|
||||
) -> Option<ResponsesWebSocketAdapterObservation> {
|
||||
None
|
||||
}
|
||||
|
||||
async fn persist_upstream_observation(
|
||||
&self,
|
||||
_state: &AppState,
|
||||
_trace_id: &str,
|
||||
_report_context: Option<&Value>,
|
||||
_observation: ResponsesWebSocketAdapterObservation,
|
||||
) {
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn is_standard_responses_event(event: &Value) -> bool {
|
||||
event
|
||||
.get("type")
|
||||
.and_then(Value::as_str)
|
||||
.is_some_and(|event_type| event_type.starts_with("response."))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::{
|
||||
resolve_responses_websocket_adapter, ResponsesWebSocketProtocolAdapter,
|
||||
ResponsesWebSocketRelayDirective,
|
||||
};
|
||||
use crate::orchestration::ResponsesWebSocketAdapter;
|
||||
|
||||
#[test]
|
||||
fn standard_adapter_has_no_codex_extensions() {
|
||||
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard);
|
||||
|
||||
assert_eq!(adapter.kind(), ResponsesWebSocketAdapter::Standard);
|
||||
assert!(!adapter.observes_upstream_events());
|
||||
assert_eq!(
|
||||
adapter.upstream_errors().handshake_failed,
|
||||
"responses_websocket_handshake_failed"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn standard_adapter_always_forwards_future_events_opaquely() {
|
||||
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard);
|
||||
let event = json!({
|
||||
"type": "response.future_capability.delta",
|
||||
"delta": {"future_shape": [1, {"nested": true}]},
|
||||
"unknown_top_level": {"must": "survive"},
|
||||
});
|
||||
|
||||
assert_eq!(
|
||||
adapter.relay_directive_for_upstream_event(&event),
|
||||
ResponsesWebSocketRelayDirective::ForwardOriginal
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,472 @@
|
||||
//! Codex-specific extensions for the standard Responses WebSocket session.
|
||||
|
||||
use async_trait::async_trait;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use super::super::adapter::{
|
||||
is_standard_responses_event, ResponsesWebSocketAdapterObservation,
|
||||
ResponsesWebSocketDrainDirective, ResponsesWebSocketExclusionIdentity,
|
||||
ResponsesWebSocketProtocolAdapter, ResponsesWebSocketRebindSafety,
|
||||
ResponsesWebSocketRelayDirective,
|
||||
};
|
||||
use crate::ai_serving::AiExecutionDecision;
|
||||
use crate::clock::current_unix_secs;
|
||||
use crate::handlers::proxy::websocket::transport::UpstreamWebSocketErrorCodes;
|
||||
use crate::orchestration::{
|
||||
codex_account_id_from_headers, codex_quota_exhaustion_reset_at,
|
||||
sync_codex_websocket_quota_metadata, ResponsesWebSocketAdapter,
|
||||
};
|
||||
use crate::AppState;
|
||||
|
||||
const CODEX_WEBSOCKET_LOG_TARGET: &str = "aether_gateway::handlers::proxy::codex_ws";
|
||||
const CODEX_WEBSOCKET_RATE_LIMITS_REPORT_CONTEXT_FIELD: &str = "codex_websocket_rate_limits";
|
||||
|
||||
const CODEX_UPSTREAM_WEBSOCKET_ERRORS: UpstreamWebSocketErrorCodes = UpstreamWebSocketErrorCodes {
|
||||
upstream_url_missing: "codex_upstream_url_missing",
|
||||
upstream_url_invalid: "codex_upstream_url_invalid",
|
||||
headers_invalid: "codex_websocket_headers_invalid",
|
||||
client_build_failed: "codex_websocket_client_build_failed",
|
||||
proxy_invalid: "codex_websocket_proxy_invalid",
|
||||
tunnel_proxy_unsupported: "codex_websocket_tunnel_proxy_unsupported",
|
||||
handshake_failed: "codex_websocket_handshake_failed",
|
||||
upgrade_rejected: "codex_websocket_upgrade_rejected",
|
||||
upgrade_failed: "codex_websocket_upgrade_failed",
|
||||
};
|
||||
|
||||
pub(crate) static CODEX_RESPONSES_WEBSOCKET_ADAPTER: CodexResponsesWebSocketAdapter =
|
||||
CodexResponsesWebSocketAdapter;
|
||||
|
||||
pub(crate) struct CodexResponsesWebSocketAdapter;
|
||||
|
||||
#[async_trait]
|
||||
impl ResponsesWebSocketProtocolAdapter for CodexResponsesWebSocketAdapter {
|
||||
fn kind(&self) -> ResponsesWebSocketAdapter {
|
||||
ResponsesWebSocketAdapter::Codex
|
||||
}
|
||||
|
||||
fn upstream_errors(&self) -> UpstreamWebSocketErrorCodes {
|
||||
CODEX_UPSTREAM_WEBSOCKET_ERRORS
|
||||
}
|
||||
|
||||
fn decorate_turn_report_context(&self, report_context: &mut Option<Value>, event: &Value) {
|
||||
let Some(rate_limits) = parse_codex_rate_limits(event) else {
|
||||
return;
|
||||
};
|
||||
let context = report_context.get_or_insert_with(|| Value::Object(Map::new()));
|
||||
let Some(context) = context.as_object_mut() else {
|
||||
return;
|
||||
};
|
||||
context.insert(
|
||||
CODEX_WEBSOCKET_RATE_LIMITS_REPORT_CONTEXT_FIELD.to_string(),
|
||||
rate_limits,
|
||||
);
|
||||
}
|
||||
|
||||
fn observes_upstream_events(&self) -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
fn rebind_safety_for_upstream_event(&self, event: &Value) -> ResponsesWebSocketRebindSafety {
|
||||
let mut saw_event = false;
|
||||
if event.get("type").and_then(Value::as_str).is_some() {
|
||||
saw_event = true;
|
||||
let safety = codex_direct_rebind_safety(event);
|
||||
if matches!(safety, ResponsesWebSocketRebindSafety::Unsafe { .. }) {
|
||||
return safety;
|
||||
}
|
||||
}
|
||||
match event.get("chunks") {
|
||||
Some(Value::Array(chunks)) => {
|
||||
for chunk in chunks {
|
||||
saw_event = true;
|
||||
let safety = codex_direct_rebind_safety(chunk);
|
||||
if matches!(safety, ResponsesWebSocketRebindSafety::Unsafe { .. }) {
|
||||
return safety;
|
||||
}
|
||||
}
|
||||
}
|
||||
Some(_) => {
|
||||
return ResponsesWebSocketRebindSafety::Unsafe {
|
||||
reason: "unrecognized_upstream_event",
|
||||
};
|
||||
}
|
||||
None => {}
|
||||
}
|
||||
if saw_event {
|
||||
ResponsesWebSocketRebindSafety::Safe
|
||||
} else {
|
||||
ResponsesWebSocketRebindSafety::Unsafe {
|
||||
reason: "unrecognized_upstream_event",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn relay_directive_for_upstream_event<'a>(
|
||||
&self,
|
||||
event: &'a Value,
|
||||
) -> ResponsesWebSocketRelayDirective<'a> {
|
||||
codex_relay_directive(event)
|
||||
}
|
||||
|
||||
fn observe_upstream_event(
|
||||
&self,
|
||||
event: &Value,
|
||||
) -> Option<ResponsesWebSocketAdapterObservation> {
|
||||
let rate_limits = parse_codex_rate_limits(event)?;
|
||||
let exhausted =
|
||||
aether_admin::provider::quota::codex_rate_limit_metadata_exhausted(&rate_limits);
|
||||
let retry_exclusion_until_unix_secs =
|
||||
codex_quota_exhaustion_reset_at(&rate_limits, current_unix_secs());
|
||||
Some(ResponsesWebSocketAdapterObservation {
|
||||
drain: exhausted.then_some(ResponsesWebSocketDrainDirective {
|
||||
error_code: "codex_account_quota_exhausted",
|
||||
retry_current_turn: true,
|
||||
retry_exclusion_until_unix_secs,
|
||||
}),
|
||||
quota_metadata: Some(rate_limits),
|
||||
})
|
||||
}
|
||||
|
||||
fn exhaustion_exclusion_identity(
|
||||
&self,
|
||||
decision: &AiExecutionDecision,
|
||||
) -> Option<ResponsesWebSocketExclusionIdentity> {
|
||||
Some(ResponsesWebSocketExclusionIdentity {
|
||||
account_id: codex_account_id_from_headers(&decision.provider_request_headers)
|
||||
.map(str::to_string),
|
||||
})
|
||||
}
|
||||
|
||||
async fn persist_upstream_observation(
|
||||
&self,
|
||||
state: &AppState,
|
||||
trace_id: &str,
|
||||
report_context: Option<&Value>,
|
||||
observation: ResponsesWebSocketAdapterObservation,
|
||||
) {
|
||||
let Some(rate_limits) = observation.quota_metadata else {
|
||||
return;
|
||||
};
|
||||
if let Err(error) =
|
||||
sync_codex_websocket_quota_metadata(state, report_context, rate_limits).await
|
||||
{
|
||||
tracing::warn!(
|
||||
target: CODEX_WEBSOCKET_LOG_TARGET,
|
||||
event_name = "codex_websocket_quota_sync_failed",
|
||||
log_type = "ops",
|
||||
transport = "websocket",
|
||||
websocket = true,
|
||||
trace_id = %trace_id,
|
||||
error = ?error,
|
||||
"gateway failed to persist Codex WebSocket quota metadata"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn codex_direct_rebind_safety(event: &Value) -> ResponsesWebSocketRebindSafety {
|
||||
let event_type = event
|
||||
.get("type")
|
||||
.and_then(Value::as_str)
|
||||
.unwrap_or_default();
|
||||
if matches!(event_type, "codex.rate_limits" | "codex.response.metadata") {
|
||||
// Codex emits these as pre-response advisory metadata. They do
|
||||
// not create a public `response.*` object, so a replacement
|
||||
// upstream can safely emit its own current snapshot.
|
||||
return ResponsesWebSocketRebindSafety::Safe;
|
||||
}
|
||||
if event_type == "error"
|
||||
&& event.pointer("/error/type").and_then(Value::as_str) == Some("usage_limit_reached")
|
||||
&& parse_codex_rate_limits(event).is_some()
|
||||
{
|
||||
// This terminal quota event has not been relayed yet. It can trigger
|
||||
// one transparent attempt on another key as long as no earlier public
|
||||
// response event made the logical turn unsafe. If replanning fails,
|
||||
// the connection layer forwards this exact upstream error instead of
|
||||
// manufacturing a gateway continuation error.
|
||||
return ResponsesWebSocketRebindSafety::Safe;
|
||||
}
|
||||
let reason = if is_standard_responses_event(event) {
|
||||
"standard_response_event"
|
||||
} else {
|
||||
"unrecognized_upstream_event"
|
||||
};
|
||||
ResponsesWebSocketRebindSafety::Unsafe { reason }
|
||||
}
|
||||
|
||||
fn codex_relay_directive(event: &Value) -> ResponsesWebSocketRelayDirective<'_> {
|
||||
match event.get("chunks") {
|
||||
Some(Value::Array(chunks)) if is_explicit_codex_batch_envelope(event) => {
|
||||
let public_events = chunks
|
||||
.iter()
|
||||
.filter(|chunk| !is_codex_private_leaf_event(chunk))
|
||||
.collect::<Vec<_>>();
|
||||
if public_events.is_empty() {
|
||||
ResponsesWebSocketRelayDirective::SuppressProviderPrivate
|
||||
} else {
|
||||
ResponsesWebSocketRelayDirective::ForwardEvents(public_events)
|
||||
}
|
||||
}
|
||||
// A malformed or future shape is not proven private. Preserve it
|
||||
// opaquely rather than guessing at a provider schema.
|
||||
Some(_) => ResponsesWebSocketRelayDirective::ForwardOriginal,
|
||||
None if is_codex_private_leaf_event(event) => {
|
||||
ResponsesWebSocketRelayDirective::SuppressProviderPrivate
|
||||
}
|
||||
None => ResponsesWebSocketRelayDirective::ForwardOriginal,
|
||||
}
|
||||
}
|
||||
|
||||
/// Recognizes only Codex's private batch container. A type-less object must
|
||||
/// contain exactly `chunks`; unknown siblings could be future public protocol
|
||||
/// data and therefore force opaque forwarding. A named Codex private root may
|
||||
/// carry provider metadata alongside its chunks and is safe to peel.
|
||||
fn is_explicit_codex_batch_envelope(event: &Value) -> bool {
|
||||
if is_codex_private_event_type(event) {
|
||||
return true;
|
||||
}
|
||||
event.as_object().is_some_and(|object| {
|
||||
object.len() == 1
|
||||
&& object.contains_key("chunks")
|
||||
&& event.get("type").and_then(Value::as_str).is_none()
|
||||
})
|
||||
}
|
||||
|
||||
fn is_codex_private_leaf_event(event: &Value) -> bool {
|
||||
is_codex_private_event_type(event) && event.get("chunks").is_none()
|
||||
}
|
||||
|
||||
fn is_codex_private_event_type(event: &Value) -> bool {
|
||||
matches!(
|
||||
event.get("type").and_then(Value::as_str),
|
||||
Some("codex.rate_limits" | "codex.response.metadata")
|
||||
)
|
||||
}
|
||||
|
||||
fn parse_codex_rate_limits(event: &Value) -> Option<Value> {
|
||||
aether_admin::provider::quota::parse_codex_websocket_rate_limits_response(
|
||||
event,
|
||||
current_unix_secs(),
|
||||
)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::{
|
||||
CodexResponsesWebSocketAdapter, ResponsesWebSocketProtocolAdapter,
|
||||
ResponsesWebSocketRebindSafety, ResponsesWebSocketRelayDirective,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn codex_rate_limit_chunk_is_kept_for_the_terminal_report() {
|
||||
let adapter = CodexResponsesWebSocketAdapter;
|
||||
assert!(adapter.observes_upstream_events());
|
||||
let mut context = Some(json!({"key_id": "codex-key"}));
|
||||
adapter.decorate_turn_report_context(
|
||||
&mut context,
|
||||
&json!({
|
||||
"chunks": [{
|
||||
"type": "codex.rate_limits",
|
||||
"plan_type": "free",
|
||||
"rate_limits": {
|
||||
"allowed": true,
|
||||
"limit_reached": false,
|
||||
"primary": {
|
||||
"used_percent": 91,
|
||||
"window_minutes": 43200,
|
||||
"reset_after_seconds": 2590791
|
||||
}
|
||||
}
|
||||
}]
|
||||
}),
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
context.as_ref().and_then(
|
||||
|context| context.pointer("/codex_websocket_rate_limits/primary_used_percent")
|
||||
),
|
||||
Some(&json!(91.0))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn usage_limit_error_is_kept_for_the_terminal_report() {
|
||||
let adapter = CodexResponsesWebSocketAdapter;
|
||||
let mut context = Some(json!({"key_id": "codex-key"}));
|
||||
adapter.decorate_turn_report_context(
|
||||
&mut context,
|
||||
&json!({
|
||||
"type": "error",
|
||||
"error": {
|
||||
"type": "usage_limit_reached",
|
||||
"plan_type": "free",
|
||||
"resets_at": 1_787_274_385u64,
|
||||
},
|
||||
"status_code": 429,
|
||||
"headers": {
|
||||
"X-Codex-Primary-Used-Percent": "100",
|
||||
"X-Codex-Primary-Reset-At": "1787274385",
|
||||
},
|
||||
}),
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
context
|
||||
.as_ref()
|
||||
.and_then(|context| context.pointer("/codex_websocket_rate_limits/allowed")),
|
||||
Some(&json!(false))
|
||||
);
|
||||
assert_eq!(
|
||||
context.as_ref().and_then(|context| {
|
||||
context.pointer("/codex_websocket_rate_limits/primary_used_percent")
|
||||
}),
|
||||
Some(&json!(100.0))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn only_known_codex_pre_response_signals_are_safe_to_rebind() {
|
||||
let adapter = CodexResponsesWebSocketAdapter;
|
||||
|
||||
assert_eq!(
|
||||
adapter.rebind_safety_for_upstream_event(&json!({
|
||||
"type": "codex.rate_limits",
|
||||
"rate_limits": {"allowed": true}
|
||||
})),
|
||||
ResponsesWebSocketRebindSafety::Safe
|
||||
);
|
||||
assert_eq!(
|
||||
adapter.rebind_safety_for_upstream_event(&json!({
|
||||
"type": "codex.response.metadata"
|
||||
})),
|
||||
ResponsesWebSocketRebindSafety::Safe
|
||||
);
|
||||
assert_eq!(
|
||||
adapter.rebind_safety_for_upstream_event(&json!({
|
||||
"chunks": [
|
||||
{"type": "codex.rate_limits", "rate_limits": {"allowed": true}},
|
||||
{"type": "codex.response.metadata"}
|
||||
]
|
||||
})),
|
||||
ResponsesWebSocketRebindSafety::Safe
|
||||
);
|
||||
assert_eq!(
|
||||
adapter.rebind_safety_for_upstream_event(&json!({
|
||||
"type": "response.created"
|
||||
})),
|
||||
ResponsesWebSocketRebindSafety::Unsafe {
|
||||
reason: "standard_response_event"
|
||||
}
|
||||
);
|
||||
assert_eq!(
|
||||
adapter.rebind_safety_for_upstream_event(&json!({
|
||||
"type": "codex.unknown"
|
||||
})),
|
||||
ResponsesWebSocketRebindSafety::Unsafe {
|
||||
reason: "unrecognized_upstream_event"
|
||||
}
|
||||
);
|
||||
assert_eq!(
|
||||
adapter.rebind_safety_for_upstream_event(&json!({
|
||||
"type": "error",
|
||||
"error": {
|
||||
"type": "usage_limit_reached",
|
||||
"plan_type": "plus",
|
||||
"resets_in_seconds": 3_600
|
||||
},
|
||||
"status_code": 429
|
||||
})),
|
||||
ResponsesWebSocketRebindSafety::Safe
|
||||
);
|
||||
assert_eq!(
|
||||
adapter.rebind_safety_for_upstream_event(&json!({
|
||||
"type": "error",
|
||||
"error": {"type": "usage_limit_reached"}
|
||||
})),
|
||||
ResponsesWebSocketRebindSafety::Unsafe {
|
||||
reason: "unrecognized_upstream_event"
|
||||
}
|
||||
);
|
||||
assert_eq!(
|
||||
adapter.rebind_safety_for_upstream_event(&json!({
|
||||
"type": "response.future_capability.delta",
|
||||
"chunks": [{"type": "codex.rate_limits"}]
|
||||
})),
|
||||
ResponsesWebSocketRebindSafety::Unsafe {
|
||||
reason: "standard_response_event"
|
||||
}
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn codex_suppresses_only_explicit_private_events_and_envelopes() {
|
||||
let adapter = CodexResponsesWebSocketAdapter;
|
||||
|
||||
for event in [
|
||||
json!({"type": "codex.rate_limits", "rate_limits": {"allowed": true}}),
|
||||
json!({"type": "codex.response.metadata", "account_hint": "private"}),
|
||||
json!({"chunks": [
|
||||
{"type": "codex.rate_limits"},
|
||||
{"type": "codex.response.metadata"}
|
||||
]}),
|
||||
] {
|
||||
assert_eq!(
|
||||
adapter.relay_directive_for_upstream_event(&event),
|
||||
ResponsesWebSocketRelayDirective::SuppressProviderPrivate
|
||||
);
|
||||
}
|
||||
|
||||
for event in [
|
||||
json!({"type": "error", "error": {"type": "usage_limit_reached"}}),
|
||||
json!({"type": "codex.future_private_maybe", "future": true}),
|
||||
json!({"chunks": [], "future_envelope_field": {"must": "survive"}}),
|
||||
json!({"type": "response.future.done", "future_capability": true}),
|
||||
] {
|
||||
assert_eq!(
|
||||
adapter.relay_directive_for_upstream_event(&event),
|
||||
ResponsesWebSocketRelayDirective::ForwardOriginal
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn mixed_codex_batch_forwards_whole_non_private_events_in_order() {
|
||||
let adapter = CodexResponsesWebSocketAdapter;
|
||||
let event = json!({
|
||||
"chunks": [
|
||||
{
|
||||
"type": "response.created",
|
||||
"response": {"id": "resp_future"},
|
||||
"future_created_field": {"opaque": true}
|
||||
},
|
||||
{"type": "codex.rate_limits", "account_hint": "private"},
|
||||
{
|
||||
"type": "response.future_capability.delta",
|
||||
"future_capability": {"nested": [1, 2, 3]},
|
||||
"sequence_number": 2
|
||||
},
|
||||
{"provider_future_event": {"unknown": "must be forwarded"}},
|
||||
{
|
||||
"type": "error",
|
||||
"error": {"type": "future_error", "future_detail": 7}
|
||||
}
|
||||
]
|
||||
});
|
||||
|
||||
let ResponsesWebSocketRelayDirective::ForwardEvents(events) =
|
||||
adapter.relay_directive_for_upstream_event(&event)
|
||||
else {
|
||||
panic!("a mixed private envelope must retain all non-private events");
|
||||
};
|
||||
assert_eq!(events.len(), 4);
|
||||
assert_eq!(events[0]["future_created_field"], json!({"opaque": true}));
|
||||
assert_eq!(events[1]["future_capability"], json!({"nested": [1, 2, 3]}));
|
||||
assert_eq!(
|
||||
events[2]["provider_future_event"],
|
||||
json!({"unknown": "must be forwarded"})
|
||||
);
|
||||
assert_eq!(events[3]["error"]["future_detail"], json!(7));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,5 @@
|
||||
//! Provider-specific Responses WebSocket adapters.
|
||||
|
||||
mod codex;
|
||||
|
||||
pub(super) use codex::CODEX_RESPONSES_WEBSOCKET_ADAPTER;
|
||||
@@ -0,0 +1,80 @@
|
||||
//! Per-turn resource admission for the Responses WebSocket bridge.
|
||||
//!
|
||||
//! A WebSocket connection may live for a long time, but each `response.create`
|
||||
//! is still one active upstream execution. Keep the resource leases attached
|
||||
//! to the turn instead of the socket so idle connections do not consume
|
||||
//! upstream capacity.
|
||||
|
||||
use std::time::Instant;
|
||||
|
||||
use aether_contracts::ExecutionPlan;
|
||||
|
||||
use crate::execution_runtime::acquire_upstream_execution_gate;
|
||||
use crate::provider_pool_demand::{
|
||||
acquire_provider_pool_in_flight_guard, ProviderPoolInFlightGuard,
|
||||
};
|
||||
use crate::upstream_admission::UpstreamTargetAdmissionPermit;
|
||||
use crate::{AppState, GatewayError};
|
||||
|
||||
pub(super) struct ResponsesWebSocketTurnAdmission {
|
||||
upstream_execution: Option<aether_runtime::ConcurrencyPermit>,
|
||||
upstream_target: Option<UpstreamTargetAdmissionPermit>,
|
||||
provider_pool: Option<ProviderPoolInFlightGuard>,
|
||||
acquired_at: Instant,
|
||||
}
|
||||
|
||||
impl ResponsesWebSocketTurnAdmission {
|
||||
pub(super) async fn acquire(
|
||||
state: &AppState,
|
||||
plan: &ExecutionPlan,
|
||||
trace_id: &str,
|
||||
) -> Result<Self, GatewayError> {
|
||||
let upstream_execution = acquire_upstream_execution_gate(state, trace_id).await?;
|
||||
let upstream_target = match state
|
||||
.upstream_target_admission
|
||||
.acquire(plan, trace_id)
|
||||
.await
|
||||
{
|
||||
Ok(permit) => permit,
|
||||
Err(error) => {
|
||||
drop(upstream_execution);
|
||||
return Err(error);
|
||||
}
|
||||
};
|
||||
let provider_pool = acquire_provider_pool_in_flight_guard(
|
||||
state.runtime_state.clone(),
|
||||
&plan.provider_id,
|
||||
&plan.request_id,
|
||||
plan.candidate_id.as_deref(),
|
||||
&plan.key_id,
|
||||
)
|
||||
.await;
|
||||
|
||||
Ok(Self {
|
||||
upstream_execution,
|
||||
upstream_target,
|
||||
provider_pool,
|
||||
acquired_at: Instant::now(),
|
||||
})
|
||||
}
|
||||
|
||||
/// Release the distributed provider token before the turn's persistence
|
||||
/// work. The remaining permits are local RAII guards and are dropped with
|
||||
/// this value.
|
||||
pub(super) async fn release(mut self) {
|
||||
if let Some(provider_pool) = self.provider_pool.take() {
|
||||
provider_pool.release().await;
|
||||
}
|
||||
drop(self.upstream_target.take());
|
||||
drop(self.upstream_execution.take());
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for ResponsesWebSocketTurnAdmission {
|
||||
fn drop(&mut self) {
|
||||
crate::stage_metrics::observe_gateway_stage_ms(
|
||||
"websocket_turn_admission_held",
|
||||
self.acquired_at.elapsed().as_millis() as u64,
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,560 @@
|
||||
//! Identity of the physical upstream connection backing a Responses session.
|
||||
//!
|
||||
//! A Responses continuation carries state that lives on one provider socket.
|
||||
//! Comparing only the selected key is therefore not sufficient: transport
|
||||
//! settings, stable account headers, credentials, and the protocol adapter can
|
||||
//! all change the connection that would receive the next event. Ordinary Codex
|
||||
//! OAuth access-token refreshes retain the credential generation and therefore
|
||||
//! do not unnecessarily replace an already-upgraded socket.
|
||||
|
||||
use std::collections::{BTreeMap, BTreeSet};
|
||||
use std::fmt;
|
||||
|
||||
use aether_contracts::{ProxySnapshot, ResolvedTransportProfile};
|
||||
use sha2::{Digest, Sha256};
|
||||
|
||||
use super::adapter::ResponsesWebSocketProtocolAdapter;
|
||||
use crate::ai_serving::AiExecutionDecision;
|
||||
use crate::handlers::proxy::websocket::transport::{
|
||||
websocket_handshake_headers, websocket_upstream_url,
|
||||
};
|
||||
use crate::orchestration::ResponsesWebSocketAdapter;
|
||||
|
||||
/// Stable, comparable identity for the actual WebSocket connection target.
|
||||
///
|
||||
/// The identity deliberately owns the normalized handshake values rather than
|
||||
/// retaining a reference to the planner decision. A later re-plan can then
|
||||
/// be compared without accidentally ignoring a field that changes the
|
||||
/// physical connection.
|
||||
#[derive(Clone, PartialEq)]
|
||||
pub(super) struct UpstreamBindingIdentity {
|
||||
adapter_kind: ResponsesWebSocketAdapter,
|
||||
provider_id: Option<String>,
|
||||
endpoint_id: Option<String>,
|
||||
key_id: Option<String>,
|
||||
upstream_url: String,
|
||||
handshake_headers: BTreeMap<String, String>,
|
||||
/// One-way identity for the credential generation used by this socket.
|
||||
///
|
||||
/// A provider key id identifies a catalog row, not the secret currently
|
||||
/// stored in that row. Codex decisions carry a server-owned credential
|
||||
/// generation which is stable across access-token refreshes but rotates
|
||||
/// when the account/static/refresh credential is replaced. Other
|
||||
/// decisions conservatively fingerprint the effective authentication
|
||||
/// handshake values.
|
||||
credential_fingerprint: [u8; 32],
|
||||
proxy: Option<ProxySnapshot>,
|
||||
transport_profile: Option<ResolvedTransportProfile>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub(super) enum UpstreamBindingIdentityError {
|
||||
MissingUpstreamUrl,
|
||||
InvalidUpstreamUrl,
|
||||
InvalidHandshakeHeaders,
|
||||
}
|
||||
|
||||
impl UpstreamBindingIdentity {
|
||||
/// Builds an identity from the same normalized URL and headers used by
|
||||
/// the WebSocket transport client.
|
||||
pub(super) fn from_decision(
|
||||
adapter: &'static dyn ResponsesWebSocketProtocolAdapter,
|
||||
decision: &AiExecutionDecision,
|
||||
) -> Result<Self, UpstreamBindingIdentityError> {
|
||||
let raw_url = decision
|
||||
.upstream_url
|
||||
.as_deref()
|
||||
.filter(|value| !value.trim().is_empty())
|
||||
.ok_or(UpstreamBindingIdentityError::MissingUpstreamUrl)?;
|
||||
let upstream_url = websocket_upstream_url(raw_url, "invalid")
|
||||
.map_err(|_| UpstreamBindingIdentityError::InvalidUpstreamUrl)?
|
||||
.to_string();
|
||||
|
||||
let headers = websocket_handshake_headers(&decision.provider_request_headers, "invalid")
|
||||
.map_err(|_| UpstreamBindingIdentityError::InvalidHandshakeHeaders)?;
|
||||
let authentication_header_names = authentication_header_names(decision);
|
||||
let mut handshake_headers = BTreeMap::new();
|
||||
let mut authentication_headers = BTreeMap::new();
|
||||
for (name, value) in &headers {
|
||||
let name = name.as_str().to_ascii_lowercase();
|
||||
let value = value
|
||||
.to_str()
|
||||
.map_err(|_| UpstreamBindingIdentityError::InvalidHandshakeHeaders)?;
|
||||
if authentication_header_names.contains(name.as_str()) {
|
||||
authentication_headers.insert(name, value.to_string());
|
||||
} else {
|
||||
handshake_headers.insert(name, value.to_string());
|
||||
}
|
||||
}
|
||||
let credential_fingerprint =
|
||||
credential_binding_fingerprint(decision, &authentication_headers);
|
||||
|
||||
Ok(Self {
|
||||
adapter_kind: adapter.kind(),
|
||||
provider_id: decision.provider_id.clone(),
|
||||
endpoint_id: decision.endpoint_id.clone(),
|
||||
key_id: decision.key_id.clone(),
|
||||
upstream_url,
|
||||
handshake_headers,
|
||||
credential_fingerprint,
|
||||
proxy: effective_proxy_snapshot(decision.proxy.as_ref()),
|
||||
transport_profile: decision.transport_profile.clone(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/// Header names that carry credentials in the provider handshake. The
|
||||
/// planner's explicit `auth_header` extends this list for provider-specific
|
||||
/// schemes; unknown headers remain part of the stable handshake identity.
|
||||
fn authentication_header_names(decision: &AiExecutionDecision) -> BTreeSet<String> {
|
||||
let mut names = BTreeSet::from([
|
||||
"authorization".to_string(),
|
||||
"proxy-authorization".to_string(),
|
||||
"x-api-key".to_string(),
|
||||
"api-key".to_string(),
|
||||
"x-goog-api-key".to_string(),
|
||||
"x-azure-api-key".to_string(),
|
||||
]);
|
||||
if let Some(name) = decision
|
||||
.auth_header
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|name| !name.is_empty())
|
||||
{
|
||||
names.insert(name.to_ascii_lowercase());
|
||||
}
|
||||
names
|
||||
}
|
||||
|
||||
fn fingerprint_headers(headers: &BTreeMap<String, String>) -> [u8; 32] {
|
||||
let mut hasher = Sha256::new();
|
||||
hasher.update(b"aether-responses-websocket-auth-headers-v1");
|
||||
for (name, value) in headers {
|
||||
hasher.update((name.len() as u64).to_be_bytes());
|
||||
hasher.update(name.as_bytes());
|
||||
hasher.update((value.len() as u64).to_be_bytes());
|
||||
hasher.update(value.as_bytes());
|
||||
}
|
||||
hasher.finalize().into()
|
||||
}
|
||||
|
||||
/// Returns the non-secret credential identity represented by a planner
|
||||
/// decision. The generation is emitted by Aether's trusted Codex planner from
|
||||
/// provider-key metadata; it is not sourced from the downstream request.
|
||||
fn credential_binding_fingerprint(
|
||||
decision: &AiExecutionDecision,
|
||||
authentication_headers: &BTreeMap<String, String>,
|
||||
) -> [u8; 32] {
|
||||
if decision
|
||||
.provider_type
|
||||
.as_deref()
|
||||
.is_some_and(|provider_type| provider_type.trim().eq_ignore_ascii_case("codex"))
|
||||
{
|
||||
if let Some(generation) = decision
|
||||
.report_context
|
||||
.as_ref()
|
||||
.and_then(|context| context.get("codex_credential_generation"))
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|generation| !generation.is_empty())
|
||||
{
|
||||
let mut hasher = Sha256::new();
|
||||
hasher.update(b"aether-responses-websocket-codex-credential-generation-v1");
|
||||
hasher.update((generation.len() as u64).to_be_bytes());
|
||||
hasher.update(generation.as_bytes());
|
||||
// Only a planner-owned Codex bearer access token is expected to
|
||||
// rotate without changing credential generation. Compare the
|
||||
// effective handshake value with the decision's original auth
|
||||
// value: auth-config/routing/header overrides change only the
|
||||
// former and therefore must force a rebind.
|
||||
let stable_authentication_headers = authentication_headers
|
||||
.iter()
|
||||
.filter(|(name, value)| {
|
||||
!is_planner_owned_codex_bearer(decision, name.as_str(), value.as_str())
|
||||
})
|
||||
.map(|(name, value)| (name.clone(), value.clone()))
|
||||
.collect::<BTreeMap<_, _>>();
|
||||
hasher.update(fingerprint_headers(&stable_authentication_headers));
|
||||
return hasher.finalize().into();
|
||||
}
|
||||
}
|
||||
|
||||
// Fail closed when no trusted generation is available. Rebinding after an
|
||||
// access-token change is preferable to sending a continuation over a
|
||||
// socket authenticated with a credential that may have been replaced.
|
||||
fingerprint_headers(authentication_headers)
|
||||
}
|
||||
|
||||
fn is_planner_owned_codex_bearer(
|
||||
decision: &AiExecutionDecision,
|
||||
name: &str,
|
||||
effective_value: &str,
|
||||
) -> bool {
|
||||
name.eq_ignore_ascii_case("authorization")
|
||||
&& decision
|
||||
.auth_header
|
||||
.as_deref()
|
||||
.is_some_and(|header| header.eq_ignore_ascii_case(name))
|
||||
&& decision.auth_value.as_deref() == Some(effective_value)
|
||||
&& effective_value
|
||||
.get(.."bearer ".len())
|
||||
.is_some_and(|scheme| scheme.eq_ignore_ascii_case("bearer "))
|
||||
}
|
||||
|
||||
/// Normalize only values that are provably direct transport. Keep node/tunnel
|
||||
/// fields even though the current WebSocket builder rejects those proxies: a
|
||||
/// re-plan must not accidentally reuse an already-bound direct socket for a
|
||||
/// decision that selected a different proxy topology.
|
||||
fn effective_proxy_snapshot(proxy: Option<&ProxySnapshot>) -> Option<ProxySnapshot> {
|
||||
let proxy = proxy?;
|
||||
if proxy.enabled == Some(false) {
|
||||
return None;
|
||||
}
|
||||
let mut normalized = proxy.clone();
|
||||
normalized.url = normalized
|
||||
.url
|
||||
.take()
|
||||
.map(|value| value.trim().to_string())
|
||||
.filter(|value| !value.is_empty());
|
||||
normalized.mode = normalized
|
||||
.mode
|
||||
.take()
|
||||
.map(|value| value.trim().to_ascii_lowercase())
|
||||
.filter(|value| !value.is_empty());
|
||||
normalized.node_id = normalized
|
||||
.node_id
|
||||
.take()
|
||||
.map(|value| value.trim().to_string())
|
||||
.filter(|value| !value.is_empty());
|
||||
normalized.label = normalized
|
||||
.label
|
||||
.take()
|
||||
.map(|value| value.trim().to_string())
|
||||
.filter(|value| !value.is_empty());
|
||||
|
||||
let has_effective_proxy = normalized.url.is_some()
|
||||
|| normalized.node_id.is_some()
|
||||
|| normalized.mode.is_some()
|
||||
|| normalized.extra.is_some();
|
||||
has_effective_proxy.then_some(normalized)
|
||||
}
|
||||
|
||||
impl fmt::Debug for UpstreamBindingIdentity {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("UpstreamBindingIdentity")
|
||||
.field("adapter_kind", &self.adapter_kind)
|
||||
.field("provider_id", &self.provider_id)
|
||||
.field("endpoint_id", &self.endpoint_id)
|
||||
.field("key_id", &self.key_id)
|
||||
.field("upstream_url", &self.upstream_url)
|
||||
.field(
|
||||
"handshake_header_names",
|
||||
&self.handshake_headers.keys().collect::<Vec<_>>(),
|
||||
)
|
||||
.field("proxy_configured", &self.proxy.is_some())
|
||||
.field(
|
||||
"transport_profile_id",
|
||||
&self
|
||||
.transport_profile
|
||||
.as_ref()
|
||||
.map(|profile| profile.profile_id.as_str()),
|
||||
)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use serde_json::json;
|
||||
|
||||
use super::{UpstreamBindingIdentity, UpstreamBindingIdentityError};
|
||||
use crate::ai_serving::AiExecutionDecision;
|
||||
use crate::handlers::proxy::websocket::responses::adapter::resolve_responses_websocket_adapter;
|
||||
use crate::orchestration::ResponsesWebSocketAdapter;
|
||||
|
||||
fn decision() -> AiExecutionDecision {
|
||||
AiExecutionDecision {
|
||||
action: "execute".to_string(),
|
||||
decision_kind: None,
|
||||
execution_strategy: None,
|
||||
conversion_mode: None,
|
||||
request_id: Some("request-1".to_string()),
|
||||
candidate_id: Some("candidate-1".to_string()),
|
||||
provider_name: Some("provider".to_string()),
|
||||
provider_type: Some("openai".to_string()),
|
||||
provider_id: Some("provider-1".to_string()),
|
||||
endpoint_id: Some("endpoint-1".to_string()),
|
||||
key_id: Some("key-1".to_string()),
|
||||
upstream_base_url: Some("https://api.example.test".to_string()),
|
||||
upstream_url: Some("https://api.example.test/v1/responses".to_string()),
|
||||
provider_request_method: Some("POST".to_string()),
|
||||
auth_header: Some("authorization".to_string()),
|
||||
auth_value: Some("Bearer secret".to_string()),
|
||||
provider_api_format: Some("openai:responses".to_string()),
|
||||
client_api_format: Some("openai:responses".to_string()),
|
||||
provider_contract: None,
|
||||
client_contract: None,
|
||||
model_name: Some("gpt-5.6-sol".to_string()),
|
||||
mapped_model: None,
|
||||
prompt_cache_key: None,
|
||||
extra_headers: BTreeMap::new(),
|
||||
provider_request_headers: BTreeMap::from([
|
||||
("Authorization".to_string(), "Bearer secret".to_string()),
|
||||
("X-Client".to_string(), "aether".to_string()),
|
||||
("Connection".to_string(), "keep-alive".to_string()),
|
||||
]),
|
||||
provider_request_body: Some(json!({"model": "gpt-5.6-sol"})),
|
||||
provider_request_body_base64: None,
|
||||
content_type: Some("application/json".to_string()),
|
||||
content_encoding: None,
|
||||
request_gzip: None,
|
||||
proxy: None,
|
||||
transport_profile: None,
|
||||
timeouts: None,
|
||||
upstream_is_stream: true,
|
||||
report_kind: None,
|
||||
report_context: None,
|
||||
auth_context: None,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn identity_normalizes_url_and_hop_by_hop_headers() {
|
||||
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard);
|
||||
let identity = UpstreamBindingIdentity::from_decision(adapter, &decision()).unwrap();
|
||||
|
||||
assert_eq!(identity.upstream_url, "wss://api.example.test/v1/responses");
|
||||
assert_eq!(
|
||||
identity.handshake_headers,
|
||||
BTreeMap::from([("x-client".to_string(), "aether".to_string())])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn identity_changes_when_physical_binding_changes() {
|
||||
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard);
|
||||
let base = decision();
|
||||
let identity = UpstreamBindingIdentity::from_decision(adapter, &base).unwrap();
|
||||
|
||||
let codex_adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Codex);
|
||||
assert_ne!(
|
||||
identity,
|
||||
UpstreamBindingIdentity::from_decision(codex_adapter, &base).unwrap()
|
||||
);
|
||||
|
||||
for mutate in [
|
||||
|decision: &mut AiExecutionDecision| {
|
||||
decision.key_id = Some("key-2".to_string());
|
||||
},
|
||||
|decision: &mut AiExecutionDecision| {
|
||||
decision.upstream_url = Some("https://other.example.test/v1/responses".to_string());
|
||||
},
|
||||
|decision: &mut AiExecutionDecision| {
|
||||
decision
|
||||
.provider_request_headers
|
||||
.insert("X-Client".to_string(), "other".to_string());
|
||||
},
|
||||
|decision: &mut AiExecutionDecision| {
|
||||
decision.proxy = Some(aether_contracts::ProxySnapshot {
|
||||
enabled: Some(true),
|
||||
url: Some("http://proxy.example.test:8080".to_string()),
|
||||
..Default::default()
|
||||
});
|
||||
},
|
||||
|decision: &mut AiExecutionDecision| {
|
||||
decision.transport_profile = Some(aether_contracts::ResolvedTransportProfile {
|
||||
profile_id: "chrome136".to_string(),
|
||||
..Default::default()
|
||||
});
|
||||
},
|
||||
] {
|
||||
let mut changed = base.clone();
|
||||
mutate(&mut changed);
|
||||
let changed_identity =
|
||||
UpstreamBindingIdentity::from_decision(adapter, &changed).unwrap();
|
||||
assert_ne!(identity, changed_identity);
|
||||
}
|
||||
|
||||
let mut static_secret_rotated = base.clone();
|
||||
static_secret_rotated
|
||||
.provider_request_headers
|
||||
.insert("Authorization".to_string(), "Bearer rotated".to_string());
|
||||
assert_ne!(
|
||||
identity,
|
||||
UpstreamBindingIdentity::from_decision(adapter, &static_secret_rotated).unwrap()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stable_key_identity_rejects_custom_static_auth_value_rotation() {
|
||||
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard);
|
||||
let mut base = decision();
|
||||
base.auth_header = Some("X-Provider-Token".to_string());
|
||||
base.provider_request_headers.remove("Authorization");
|
||||
base.provider_request_headers.insert(
|
||||
"X-Provider-Token".to_string(),
|
||||
"provider-token-1".to_string(),
|
||||
);
|
||||
let identity = UpstreamBindingIdentity::from_decision(adapter, &base).unwrap();
|
||||
assert!(!identity.handshake_headers.contains_key("x-provider-token"));
|
||||
|
||||
let mut rotated = base;
|
||||
rotated.provider_request_headers.insert(
|
||||
"X-Provider-Token".to_string(),
|
||||
"provider-token-2".to_string(),
|
||||
);
|
||||
assert_ne!(
|
||||
identity,
|
||||
UpstreamBindingIdentity::from_decision(adapter, &rotated).unwrap()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn codex_access_token_refresh_reuses_the_same_credential_generation() {
|
||||
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Codex);
|
||||
let mut first = decision();
|
||||
first.provider_type = Some("codex".to_string());
|
||||
first.report_context = Some(json!({
|
||||
"codex_credential_generation": "credential-generation-1"
|
||||
}));
|
||||
let first_identity = UpstreamBindingIdentity::from_decision(adapter, &first).unwrap();
|
||||
|
||||
let mut access_token_refreshed = first;
|
||||
access_token_refreshed.auth_value = Some("Bearer refreshed-access-token".to_string());
|
||||
access_token_refreshed.provider_request_headers.insert(
|
||||
"Authorization".to_string(),
|
||||
"Bearer refreshed-access-token".to_string(),
|
||||
);
|
||||
assert_eq!(
|
||||
first_identity,
|
||||
UpstreamBindingIdentity::from_decision(adapter, &access_token_refreshed).unwrap()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn codex_authorization_override_changes_binding_with_the_same_generation() {
|
||||
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Codex);
|
||||
let mut first = decision();
|
||||
first.provider_type = Some("codex".to_string());
|
||||
first.report_context = Some(json!({
|
||||
"codex_credential_generation": "credential-generation-1"
|
||||
}));
|
||||
let first_identity = UpstreamBindingIdentity::from_decision(adapter, &first).unwrap();
|
||||
|
||||
// The planner-owned auth value remains unchanged while an effective
|
||||
// auth-config/header override replaces the actual handshake value.
|
||||
first.provider_request_headers.insert(
|
||||
"Authorization".to_string(),
|
||||
"Bearer endpoint-override".to_string(),
|
||||
);
|
||||
assert_ne!(
|
||||
first_identity,
|
||||
UpstreamBindingIdentity::from_decision(adapter, &first).unwrap()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn codex_credential_replacement_changes_binding_for_the_same_key_id() {
|
||||
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Codex);
|
||||
let mut first = decision();
|
||||
first.provider_type = Some("codex".to_string());
|
||||
first.report_context = Some(json!({
|
||||
"codex_credential_generation": "credential-generation-1"
|
||||
}));
|
||||
let first_identity = UpstreamBindingIdentity::from_decision(adapter, &first).unwrap();
|
||||
|
||||
let mut replaced = first;
|
||||
replaced.provider_request_headers.insert(
|
||||
"Authorization".to_string(),
|
||||
"Bearer replacement-access-token".to_string(),
|
||||
);
|
||||
replaced.report_context = Some(json!({
|
||||
"codex_credential_generation": "credential-generation-2"
|
||||
}));
|
||||
assert_ne!(
|
||||
first_identity,
|
||||
UpstreamBindingIdentity::from_decision(adapter, &replaced).unwrap()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn codex_custom_auth_rotation_changes_binding_with_the_same_generation() {
|
||||
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Codex);
|
||||
let mut first = decision();
|
||||
first.provider_type = Some("codex".to_string());
|
||||
first.auth_header = Some("X-Provider-Token".to_string());
|
||||
first.provider_request_headers.insert(
|
||||
"X-Provider-Token".to_string(),
|
||||
"provider-token-1".to_string(),
|
||||
);
|
||||
first.report_context = Some(json!({
|
||||
"codex_credential_generation": "credential-generation-1"
|
||||
}));
|
||||
let first_identity = UpstreamBindingIdentity::from_decision(adapter, &first).unwrap();
|
||||
|
||||
first.provider_request_headers.insert(
|
||||
"X-Provider-Token".to_string(),
|
||||
"provider-token-2".to_string(),
|
||||
);
|
||||
assert_ne!(
|
||||
first_identity,
|
||||
UpstreamBindingIdentity::from_decision(adapter, &first).unwrap()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn missing_codex_credential_generation_fails_closed_on_auth_rotation() {
|
||||
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Codex);
|
||||
let mut first = decision();
|
||||
first.provider_type = Some("codex".to_string());
|
||||
let first_identity = UpstreamBindingIdentity::from_decision(adapter, &first).unwrap();
|
||||
|
||||
first.provider_request_headers.insert(
|
||||
"Authorization".to_string(),
|
||||
"Bearer possibly-replaced-credential".to_string(),
|
||||
);
|
||||
assert_ne!(
|
||||
first_identity,
|
||||
UpstreamBindingIdentity::from_decision(adapter, &first).unwrap()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn disabled_proxy_is_equivalent_to_direct_transport() {
|
||||
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard);
|
||||
let direct = decision();
|
||||
let direct_identity = UpstreamBindingIdentity::from_decision(adapter, &direct).unwrap();
|
||||
let mut explicitly_disabled = direct;
|
||||
explicitly_disabled.proxy = Some(aether_contracts::ProxySnapshot {
|
||||
enabled: Some(false),
|
||||
url: Some("http://ignored.example.test:8080".to_string()),
|
||||
..Default::default()
|
||||
});
|
||||
|
||||
assert_eq!(
|
||||
direct_identity,
|
||||
UpstreamBindingIdentity::from_decision(adapter, &explicitly_disabled).unwrap()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn identity_rejects_missing_or_invalid_connection_fields() {
|
||||
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard);
|
||||
let mut missing = decision();
|
||||
missing.upstream_url = None;
|
||||
assert_eq!(
|
||||
UpstreamBindingIdentity::from_decision(adapter, &missing),
|
||||
Err(UpstreamBindingIdentityError::MissingUpstreamUrl)
|
||||
);
|
||||
|
||||
let mut invalid = decision();
|
||||
invalid.upstream_url = Some("file:///tmp/responses".to_string());
|
||||
assert_eq!(
|
||||
UpstreamBindingIdentity::from_decision(adapter, &invalid),
|
||||
Err(UpstreamBindingIdentityError::InvalidUpstreamUrl)
|
||||
);
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,623 @@
|
||||
//! Connection-level Responses WebSocket FSM.
|
||||
|
||||
use std::time::Duration;
|
||||
|
||||
use axum::extract::ws::{Message as AxumWsMessage, WebSocket};
|
||||
use futures_util::{SinkExt, StreamExt};
|
||||
use serde_json::Value;
|
||||
use wreq::ws::message::Message as WreqWsMessage;
|
||||
|
||||
use super::adapter::ResponsesWebSocketRelayDirective;
|
||||
use super::client::{adapter_drain_ready, forward_client_message, RelayDisposition};
|
||||
use super::frame::{encode_opaque_websocket_event, ParsedResponsesWebSocketFrame};
|
||||
use super::lifecycle::{
|
||||
await_pending_adapter_observation, finalize_active_turn, queue_turn_finalization,
|
||||
settle_turn_finalization, spawn_bounded_adapter_observation, PreviousAttemptSettled,
|
||||
};
|
||||
use super::quota::{
|
||||
detach_exhausted_upstream, is_usage_limit_error_event, mark_active_response_retry_unsafe,
|
||||
observe_active_response_rebind_safety, retry_active_turn_after_quota_exhaustion,
|
||||
};
|
||||
use super::relay_policy::{
|
||||
classify_quota_relay, fatal_relay_policy, FatalRelaySignal, QuotaRelayAction, QuotaRelayFacts,
|
||||
};
|
||||
use super::settlement::settle_signal_for_client_delivery_failure;
|
||||
use super::state::BoundResponsesConnection;
|
||||
use super::turn::{
|
||||
ResponsesProviderAttempt, ResponsesWebSocketTurnObservation, ResponsesWebSocketTurnOutcome,
|
||||
};
|
||||
use super::upstream::{close_bound_upstream, receive_optional_upstream};
|
||||
use crate::handlers::proxy::websocket::ingress::WebSocketRequestContext;
|
||||
use crate::handlers::proxy::websocket::session::{
|
||||
wait_for_optional_deadline, CLOSE_INTERNAL_ERROR, CLOSE_TRY_AGAIN, WEBSOCKET_LOG_TRANSPORT,
|
||||
};
|
||||
use crate::handlers::proxy::websocket::transport::{
|
||||
close_client_socket, send_client_message, send_gateway_error_with_status,
|
||||
send_responses_websocket_error, upstream_message_to_client,
|
||||
};
|
||||
use crate::AppState;
|
||||
|
||||
const LOG_TARGET: &str = "aether_gateway::handlers::proxy::responses_ws";
|
||||
|
||||
/// 写客户端 socket 失败时记录的投递失败原因。刻意不说「客户端在终态前断开」:
|
||||
/// 供应商的终态可能已经到达,只是最后一跳没送出去。
|
||||
const CLIENT_DELIVERY_FAILED_REASON: &str =
|
||||
"gateway could not relay the provider event to the client";
|
||||
|
||||
macro_rules! debug {
|
||||
($($arg:tt)*) => {
|
||||
tracing::debug!(target: LOG_TARGET, $($arg)*)
|
||||
};
|
||||
}
|
||||
|
||||
macro_rules! warn {
|
||||
($($arg:tt)*) => {
|
||||
tracing::warn!(target: LOG_TARGET, $($arg)*)
|
||||
};
|
||||
}
|
||||
|
||||
pub(super) async fn relay_bound_connection(
|
||||
client_socket: &mut WebSocket,
|
||||
bound: &mut BoundResponsesConnection,
|
||||
state: &AppState,
|
||||
context: &WebSocketRequestContext,
|
||||
) {
|
||||
loop {
|
||||
let active_turn_deadline = bound.turn_state.attempt().map(|turn| turn.deadline());
|
||||
tokio::select! {
|
||||
_ = wait_for_optional_deadline(active_turn_deadline.map(|deadline| deadline.deadline)) => {
|
||||
let Some(turn_deadline) = active_turn_deadline else {
|
||||
continue;
|
||||
};
|
||||
warn!(
|
||||
event_name = "responses_websocket_turn_timeout",
|
||||
log_type = "ops",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %context.trace_id,
|
||||
timeout_phase = ?turn_deadline.phase,
|
||||
timeout_ms = turn_deadline.timeout.as_millis() as u64,
|
||||
"Responses WebSocket response did not reach its configured deadline"
|
||||
);
|
||||
finalize_active_turn(bound, state, turn_deadline.phase.outcome()).await;
|
||||
send_gateway_error_with_status(
|
||||
client_socket,
|
||||
504,
|
||||
turn_deadline.phase.error_code(),
|
||||
turn_deadline.phase.client_message(),
|
||||
).await;
|
||||
close_bound_upstream(bound).await;
|
||||
close_client_socket(
|
||||
client_socket,
|
||||
CLOSE_TRY_AGAIN,
|
||||
turn_deadline.phase.error_code(),
|
||||
).await;
|
||||
break;
|
||||
}
|
||||
client_message = client_socket.next() => {
|
||||
let Some(client_message) = client_message else {
|
||||
finalize_active_turn(
|
||||
bound,
|
||||
state,
|
||||
ResponsesWebSocketTurnOutcome::client_disconnected(),
|
||||
).await;
|
||||
close_bound_upstream(bound).await;
|
||||
break;
|
||||
};
|
||||
let Ok(client_message) = client_message else {
|
||||
warn!(
|
||||
event_name = "responses_websocket_client_receive_failed",
|
||||
log_type = "ops",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %context.trace_id,
|
||||
"client WebSocket receive failed"
|
||||
);
|
||||
finalize_active_turn(
|
||||
bound,
|
||||
state,
|
||||
ResponsesWebSocketTurnOutcome::client_disconnected(),
|
||||
).await;
|
||||
close_bound_upstream(bound).await;
|
||||
break;
|
||||
};
|
||||
match Box::pin(forward_client_message(
|
||||
client_message,
|
||||
bound,
|
||||
client_socket,
|
||||
state,
|
||||
context,
|
||||
))
|
||||
.await
|
||||
{
|
||||
RelayDisposition::Continue => {}
|
||||
RelayDisposition::Close => {
|
||||
finalize_active_turn(
|
||||
bound,
|
||||
state,
|
||||
ResponsesWebSocketTurnOutcome::client_disconnected(),
|
||||
).await;
|
||||
break;
|
||||
}
|
||||
RelayDisposition::UpstreamError(code) => {
|
||||
warn!(
|
||||
event_name = "responses_websocket_upstream_send_failed",
|
||||
log_type = "ops",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %context.trace_id,
|
||||
error_code = code,
|
||||
"Upstream WebSocket send failed"
|
||||
);
|
||||
finalize_active_turn(
|
||||
bound,
|
||||
state,
|
||||
ResponsesWebSocketTurnOutcome::upstream_send_failed(),
|
||||
).await;
|
||||
send_gateway_error_with_status(
|
||||
client_socket,
|
||||
502,
|
||||
code,
|
||||
"Gateway could not forward the WebSocket event upstream",
|
||||
).await;
|
||||
close_bound_upstream(bound).await;
|
||||
close_client_socket(client_socket, CLOSE_INTERNAL_ERROR, code).await;
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
upstream_message = receive_optional_upstream(&mut bound.upstream) => {
|
||||
let Some(upstream_message) = upstream_message else {
|
||||
finalize_active_turn(
|
||||
bound,
|
||||
state,
|
||||
ResponsesWebSocketTurnOutcome::upstream_closed(),
|
||||
).await;
|
||||
bound.upstream = None;
|
||||
close_client_socket(client_socket, 1000, "upstream_closed").await;
|
||||
break;
|
||||
};
|
||||
let Ok(upstream_message) = upstream_message else {
|
||||
warn!(
|
||||
event_name = "responses_websocket_upstream_receive_failed",
|
||||
log_type = "ops",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %context.trace_id,
|
||||
"Upstream WebSocket receive failed"
|
||||
);
|
||||
finalize_active_turn(
|
||||
bound,
|
||||
state,
|
||||
ResponsesWebSocketTurnOutcome::upstream_receive_failed(),
|
||||
).await;
|
||||
send_gateway_error_with_status(
|
||||
client_socket,
|
||||
502,
|
||||
"responses_websocket_receive_failed",
|
||||
"Provider connection closed unexpectedly",
|
||||
).await;
|
||||
bound.upstream = None;
|
||||
close_client_socket(client_socket, CLOSE_INTERNAL_ERROR, "upstream_receive_failed").await;
|
||||
break;
|
||||
};
|
||||
let parsed_upstream_frame = match &upstream_message {
|
||||
WreqWsMessage::Text(text) => {
|
||||
ParsedResponsesWebSocketFrame::parse(text.as_str()).ok()
|
||||
}
|
||||
_ => None,
|
||||
};
|
||||
let parsed_upstream_event = parsed_upstream_frame
|
||||
.as_ref()
|
||||
.map(ParsedResponsesWebSocketFrame::event);
|
||||
if let WreqWsMessage::Text(text) = &upstream_message {
|
||||
debug!(
|
||||
event_name = "responses_websocket_upstream_event",
|
||||
log_type = "event",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %context.trace_id,
|
||||
event_type = %parsed_upstream_frame
|
||||
.as_ref()
|
||||
.map(ParsedResponsesWebSocketFrame::event_type_for_log)
|
||||
.unwrap_or_else(|| "invalid_json".to_string()),
|
||||
frame_bytes = text.len(),
|
||||
chunked = parsed_upstream_frame
|
||||
.as_ref()
|
||||
.is_some_and(ParsedResponsesWebSocketFrame::is_chunked),
|
||||
active_turn = bound.turn_state.response_in_flight(),
|
||||
"gateway received Responses WebSocket event"
|
||||
);
|
||||
}
|
||||
if matches!(&upstream_message, WreqWsMessage::Binary(_)) {
|
||||
mark_active_response_retry_unsafe(bound, "upstream_binary_frame");
|
||||
} else if matches!(&upstream_message, WreqWsMessage::Text(_))
|
||||
&& parsed_upstream_event.is_none()
|
||||
{
|
||||
mark_active_response_retry_unsafe(bound, "invalid_upstream_event");
|
||||
}
|
||||
if let Some(event) = parsed_upstream_event {
|
||||
observe_active_response_rebind_safety(bound, event);
|
||||
if bound.pending_adapter_drain.is_none()
|
||||
&& bound.adapter.observes_upstream_events()
|
||||
{
|
||||
let adapter = bound.adapter;
|
||||
if let Some(observation) = adapter.observe_upstream_event(event) {
|
||||
let directive = observation.drain;
|
||||
await_pending_adapter_observation(bound).await;
|
||||
let state_for_observation = state.clone();
|
||||
let trace_id = context.trace_id.clone();
|
||||
let report_context = bound.decision_template.report_context.clone();
|
||||
bound.pending_adapter_observation = Some(spawn_bounded_adapter_observation(async move {
|
||||
adapter
|
||||
.persist_upstream_observation(
|
||||
&state_for_observation,
|
||||
&trace_id,
|
||||
report_context.as_ref(),
|
||||
observation,
|
||||
)
|
||||
.await;
|
||||
}));
|
||||
if let Some(directive) = directive {
|
||||
bound.pending_adapter_drain = Some(directive);
|
||||
// A definitive quota signal must be visible to
|
||||
// the next planner before a transparent retry.
|
||||
await_pending_adapter_observation(bound).await;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
let observation = match &upstream_message {
|
||||
WreqWsMessage::Text(text) => {
|
||||
let adapter = bound.adapter;
|
||||
match parsed_upstream_frame.as_ref() {
|
||||
Some(frame) => bound
|
||||
.turn_state
|
||||
.attempt_mut()
|
||||
.and_then(|turn| turn.observe_upstream_frame(frame, adapter)),
|
||||
None => {
|
||||
if let Some(turn) = bound.turn_state.attempt_mut() {
|
||||
turn.observe_invalid_upstream_text(text.as_str())
|
||||
}
|
||||
else {
|
||||
None
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
_ => None,
|
||||
};
|
||||
if matches!(
|
||||
observation,
|
||||
Some(ResponsesWebSocketTurnObservation::Started)
|
||||
| Some(ResponsesWebSocketTurnObservation::Terminal(_))
|
||||
) {
|
||||
if let Some(turn) = bound.turn_state.attempt_mut() {
|
||||
turn.mark_stream_started(state).await;
|
||||
}
|
||||
}
|
||||
let terminal_outcome = match observation {
|
||||
Some(ResponsesWebSocketTurnObservation::Terminal(outcome)) => Some(outcome),
|
||||
_ => None,
|
||||
};
|
||||
if matches!(&upstream_message, WreqWsMessage::Text(_))
|
||||
&& parsed_upstream_frame.is_none()
|
||||
{
|
||||
let policy = fatal_relay_policy(FatalRelaySignal::InvalidUpstreamText);
|
||||
finalize_active_turn(
|
||||
bound,
|
||||
state,
|
||||
terminal_outcome.unwrap_or_else(
|
||||
ResponsesWebSocketTurnOutcome::upstream_receive_failed,
|
||||
),
|
||||
)
|
||||
.await;
|
||||
send_responses_websocket_error(
|
||||
client_socket,
|
||||
policy.status_code,
|
||||
"server_error",
|
||||
policy.error_code,
|
||||
policy.client_message,
|
||||
)
|
||||
.await;
|
||||
close_bound_upstream(bound).await;
|
||||
close_client_socket(
|
||||
client_socket,
|
||||
policy.close_code,
|
||||
policy.close_reason,
|
||||
)
|
||||
.await;
|
||||
break;
|
||||
}
|
||||
let is_close = matches!(upstream_message, WreqWsMessage::Close(_));
|
||||
let drain_for_adapter = adapter_drain_ready(
|
||||
bound.pending_adapter_drain,
|
||||
bound.turn_state.response_in_flight(),
|
||||
observation,
|
||||
is_close,
|
||||
);
|
||||
let quota_facts = QuotaRelayFacts {
|
||||
drain_ready: drain_for_adapter,
|
||||
retry_current_turn: bound
|
||||
.pending_adapter_drain
|
||||
.is_some_and(|directive| directive.retry_current_turn)
|
||||
&& bound
|
||||
.turn_state
|
||||
.logical()
|
||||
.is_some_and(|turn| turn.quota_retry_block_reason().is_none()),
|
||||
transparent_retry_failed: false,
|
||||
usage_limit_error: parsed_upstream_event.is_some_and(is_usage_limit_error_event),
|
||||
upstream_closed: is_close,
|
||||
};
|
||||
let mut quota_relay_action = classify_quota_relay(quota_facts);
|
||||
if matches!(quota_relay_action, QuotaRelayAction::AttemptTransparentRetry) {
|
||||
// detach_attempt 保留 logical turn:重试是同一轮请求的下一个 attempt。
|
||||
let retry_turn = bound.turn_state.detach_attempt();
|
||||
// 先结算旧 attempt 并等它落地,再规划下一个 attempt。两个理由:
|
||||
//
|
||||
// 1. 规划要读 health / adaptive / pool 状态,而这些正是旧
|
||||
// attempt 结算时才投射的。普通的新 turn 早就在 client.rs 里
|
||||
// 用 await_pending_turn_finalization 挡住了「基于陈旧状态
|
||||
// 规划」,透明重试这条路径原先漏了这一步。
|
||||
// 2. 旧 attempt 还占着自己的 pool key lease。不先释放,重试就
|
||||
// 可能因为「这把 key 仍被占用」而挑不到本该可用的替代 key,
|
||||
// 或者干脆判成无可用供应商。
|
||||
let settled = match retry_turn {
|
||||
Some(mut turn) => {
|
||||
turn.release_admission().await;
|
||||
settle_turn_finalization(
|
||||
bound,
|
||||
state,
|
||||
turn,
|
||||
terminal_outcome.unwrap_or_else(
|
||||
ResponsesWebSocketTurnOutcome::upstream_closed,
|
||||
),
|
||||
)
|
||||
.await
|
||||
}
|
||||
None => PreviousAttemptSettled::nothing_to_settle(),
|
||||
};
|
||||
// Planning and binding a replacement carries the complete
|
||||
// scheduler/provider state machine. Keep that large future
|
||||
// off the relay task's stack; the default Tokio/test worker
|
||||
// stack is otherwise easy to exhaust on this rare branch.
|
||||
if Box::pin(retry_active_turn_after_quota_exhaustion(
|
||||
bound, state, context, settled,
|
||||
))
|
||||
.await
|
||||
{
|
||||
continue;
|
||||
}
|
||||
// 重试失败。旧 attempt 已经结算,logical turn 仍停在
|
||||
// Replanning,所以后面分支里的 end() / finalize_active_turn
|
||||
// 只会清掉 logical turn 而不会交出 attempt——不存在重复结算。
|
||||
quota_relay_action = classify_quota_relay(QuotaRelayFacts {
|
||||
retry_current_turn: false,
|
||||
transparent_retry_failed: true,
|
||||
..quota_facts
|
||||
});
|
||||
}
|
||||
let detach_after_forward =
|
||||
matches!(quota_relay_action, QuotaRelayAction::ForwardQuotaAndDetach);
|
||||
if detach_after_forward && is_close {
|
||||
let directive = bound
|
||||
.pending_adapter_drain
|
||||
.expect("adapter drain state should be present");
|
||||
finalize_active_turn(
|
||||
bound,
|
||||
state,
|
||||
terminal_outcome
|
||||
.unwrap_or_else(ResponsesWebSocketTurnOutcome::provider_quota_exhausted),
|
||||
)
|
||||
.await;
|
||||
send_gateway_error_with_status(
|
||||
client_socket,
|
||||
429,
|
||||
directive.error_code,
|
||||
"Provider connection closed after reporting exhausted quota; send a new response.create to select another Provider connection",
|
||||
)
|
||||
.await;
|
||||
detach_exhausted_upstream(bound, directive, &context.trace_id).await;
|
||||
continue;
|
||||
}
|
||||
// Standard Responses frames cross the gateway byte-for-byte unless PII
|
||||
// restoration has something to replace. Codex may wrap public events with
|
||||
// provider-private side-channel chunks; only that explicit envelope is
|
||||
// peeled, and each retained event is serialized as a complete opaque Value.
|
||||
// Observation and capture continue to consume the redacted event, while the
|
||||
// final client hop receives restored text.
|
||||
let relay_directive = parsed_upstream_frame
|
||||
.as_ref()
|
||||
.map(|frame| {
|
||||
bound
|
||||
.adapter
|
||||
.relay_directive_for_upstream_event(frame.event())
|
||||
});
|
||||
let mut relay_send_error = None;
|
||||
let mut relay_serialization_failed = false;
|
||||
match relay_directive {
|
||||
Some(ResponsesWebSocketRelayDirective::ForwardOriginal) => {
|
||||
let restored = parsed_upstream_frame
|
||||
.as_ref()
|
||||
.and_then(|frame| {
|
||||
bound
|
||||
.redaction_restorer
|
||||
.restore_provider_frame_text(frame.event())
|
||||
});
|
||||
let client_frame = match restored {
|
||||
Some(text) => AxumWsMessage::Text(text.into()),
|
||||
None => upstream_message_to_client(upstream_message.clone()),
|
||||
};
|
||||
match send_client_message(client_socket, client_frame).await {
|
||||
Ok(()) => {
|
||||
if let (Some(turn), Some(frame)) = (
|
||||
bound.turn_state.attempt_mut(),
|
||||
parsed_upstream_frame.as_ref(),
|
||||
) {
|
||||
turn.capture_client_frame(frame.event());
|
||||
}
|
||||
}
|
||||
Err(error) => relay_send_error = Some(error),
|
||||
}
|
||||
}
|
||||
Some(ResponsesWebSocketRelayDirective::ForwardEvents(events)) => {
|
||||
for event in events {
|
||||
let text = match bound
|
||||
.redaction_restorer
|
||||
.restore_provider_frame_text(event)
|
||||
{
|
||||
Some(restored) => restored,
|
||||
None => match encode_opaque_websocket_event(event) {
|
||||
Ok(encoded) => encoded,
|
||||
Err(_) => {
|
||||
relay_serialization_failed = true;
|
||||
break;
|
||||
}
|
||||
},
|
||||
};
|
||||
match send_client_message(
|
||||
client_socket,
|
||||
AxumWsMessage::Text(text.into()),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(()) => {
|
||||
if let Some(turn) = bound.turn_state.attempt_mut() {
|
||||
turn.capture_client_frame(event);
|
||||
}
|
||||
}
|
||||
Err(error) => {
|
||||
relay_send_error = Some(error);
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
Some(ResponsesWebSocketRelayDirective::SuppressProviderPrivate) => {}
|
||||
None => {
|
||||
if let Err(error) = send_client_message(
|
||||
client_socket,
|
||||
upstream_message_to_client(upstream_message.clone()),
|
||||
)
|
||||
.await
|
||||
{
|
||||
relay_send_error = Some(error);
|
||||
}
|
||||
}
|
||||
}
|
||||
if relay_serialization_failed {
|
||||
warn!(
|
||||
event_name = "responses_websocket_provider_event_serialization_failed",
|
||||
log_type = "ops",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %context.trace_id,
|
||||
provider_terminal_reached = terminal_outcome.is_some(),
|
||||
"gateway could not serialize an opaque provider event"
|
||||
);
|
||||
bound
|
||||
.turn_state
|
||||
.record_client_delivery_aborted(CLIENT_DELIVERY_FAILED_REASON);
|
||||
finalize_active_turn(
|
||||
bound,
|
||||
state,
|
||||
settle_signal_for_client_delivery_failure(terminal_outcome),
|
||||
)
|
||||
.await;
|
||||
send_gateway_error_with_status(
|
||||
client_socket,
|
||||
502,
|
||||
"responses_websocket_event_serialization_failed",
|
||||
"Gateway could not relay the provider event",
|
||||
)
|
||||
.await;
|
||||
close_bound_upstream(bound).await;
|
||||
close_client_socket(
|
||||
client_socket,
|
||||
CLOSE_INTERNAL_ERROR,
|
||||
"provider_event_serialization_failed",
|
||||
)
|
||||
.await;
|
||||
break;
|
||||
}
|
||||
if let Some(error) = relay_send_error {
|
||||
warn!(
|
||||
event_name = "responses_websocket_client_send_failed",
|
||||
log_type = "ops",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %context.trace_id,
|
||||
error_code = error.as_str(),
|
||||
provider_terminal_reached = terminal_outcome.is_some(),
|
||||
"gateway could not relay a provider event to the client"
|
||||
);
|
||||
// 投递失败是独立事实,不能覆盖已经到达的 provider 终态:
|
||||
// 供应商已经完成推理并消耗 token,账单按它的终态计。
|
||||
bound
|
||||
.turn_state
|
||||
.record_client_delivery_aborted(CLIENT_DELIVERY_FAILED_REASON);
|
||||
finalize_active_turn(
|
||||
bound,
|
||||
state,
|
||||
settle_signal_for_client_delivery_failure(terminal_outcome),
|
||||
).await;
|
||||
close_bound_upstream(bound).await;
|
||||
break;
|
||||
}
|
||||
if let Some(outcome) = terminal_outcome {
|
||||
finalize_active_turn(bound, state, outcome).await;
|
||||
} else if is_close {
|
||||
finalize_active_turn(
|
||||
bound,
|
||||
state,
|
||||
ResponsesWebSocketTurnOutcome::upstream_closed(),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
if detach_after_forward {
|
||||
let directive = bound
|
||||
.pending_adapter_drain
|
||||
.expect("adapter drain state should be present");
|
||||
if bound.turn_state.response_in_flight() {
|
||||
finalize_active_turn(
|
||||
bound,
|
||||
state,
|
||||
ResponsesWebSocketTurnOutcome::provider_quota_exhausted(),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
detach_exhausted_upstream(bound, directive, &context.trace_id).await;
|
||||
continue;
|
||||
}
|
||||
if drain_for_adapter {
|
||||
let directive = bound
|
||||
.pending_adapter_drain
|
||||
.expect("adapter drain state should be present");
|
||||
detach_exhausted_upstream(bound, directive, &context.trace_id).await;
|
||||
continue;
|
||||
}
|
||||
if is_close {
|
||||
bound.upstream = None;
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) async fn wait_for_connection_permit_loss(
|
||||
permit: Option<&aether_runtime::AdmissionPermit>,
|
||||
) {
|
||||
let Some(permit) = permit else {
|
||||
std::future::pending::<()>().await;
|
||||
return;
|
||||
};
|
||||
let mut health = tokio::time::interval(Duration::from_secs(1));
|
||||
health.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
|
||||
loop {
|
||||
health.tick().await;
|
||||
if !permit.is_healthy() {
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,173 @@
|
||||
//! Per-turn control-plane refresh for long-lived Responses WebSockets.
|
||||
//!
|
||||
//! An Upgrade authenticates the connection, but it must not freeze API-key,
|
||||
//! wallet, IP, model, or RPM policy for up to an hour. This module produces one
|
||||
//! live decision and its exact strong API-key snapshot for every
|
||||
//! `response.create`; the caller uses that pair consistently for rate limiting,
|
||||
//! redaction, model authorization, planning, admission, balance, and retries.
|
||||
|
||||
use axum::http::StatusCode;
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::ai_serving::GatewayAuthApiKeySnapshot;
|
||||
use crate::control::{
|
||||
refresh_execution_runtime_auth_context_with_snapshot, request_model_local_rejection,
|
||||
GatewayControlDecision, GatewayLocalAuthRejection,
|
||||
};
|
||||
use crate::handlers::proxy::websocket::ingress::WebSocketRequestContext;
|
||||
use crate::handlers::proxy::websocket::session::WEBSOCKET_LOG_TRANSPORT;
|
||||
use crate::handlers::shared::ip_rules_allow;
|
||||
use crate::{AppState, GatewayError};
|
||||
|
||||
const LOG_TARGET: &str = "aether_gateway::handlers::proxy::responses_ws";
|
||||
|
||||
macro_rules! warn {
|
||||
($($arg:tt)*) => {
|
||||
tracing::warn!(target: LOG_TARGET, $($arg)*)
|
||||
};
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub(super) struct ResponsesWebSocketTurnControl {
|
||||
pub(super) decision: GatewayControlDecision,
|
||||
pub(super) auth_snapshot: Option<GatewayAuthApiKeySnapshot>,
|
||||
pub(super) rpm_bypassed: bool,
|
||||
}
|
||||
|
||||
pub(super) async fn resolve_responses_websocket_turn_control(
|
||||
state: &AppState,
|
||||
context: &WebSocketRequestContext,
|
||||
parts: &http::request::Parts,
|
||||
client_event: &Value,
|
||||
) -> Result<ResponsesWebSocketTurnControl, GatewayError> {
|
||||
if state
|
||||
.admin_security_ip_blacklisted(context.client_ip)
|
||||
.await?
|
||||
{
|
||||
return Err(GatewayError::Client {
|
||||
status: StatusCode::FORBIDDEN,
|
||||
message: "The current IP is blocked".to_string(),
|
||||
});
|
||||
}
|
||||
|
||||
let mut decision = context.decision.clone();
|
||||
let auth_snapshot = if let Some(auth_context) = decision.auth_context.take() {
|
||||
let (refreshed, snapshot) = refresh_execution_runtime_auth_context_with_snapshot(
|
||||
state,
|
||||
auth_context,
|
||||
decision.auth_endpoint_signature.as_deref(),
|
||||
)
|
||||
.await?;
|
||||
decision.local_auth_rejection = refreshed.local_rejection.clone();
|
||||
decision.auth_context = Some(refreshed);
|
||||
snapshot
|
||||
} else {
|
||||
None
|
||||
};
|
||||
// Model-directive configuration is mutable policy too; do not retain the
|
||||
// Upgrade-time snapshot for the lifetime of the socket.
|
||||
decision.model_directive_policy =
|
||||
crate::system_features::ModelDirectivePolicySnapshot::load(state).await;
|
||||
|
||||
if let Some(rejection) = decision.local_auth_rejection.clone() {
|
||||
return Err(websocket_auth_rejection_error(rejection));
|
||||
}
|
||||
let Some(auth_context) = decision.auth_context.as_ref() else {
|
||||
return Err(websocket_auth_rejection_error(
|
||||
GatewayLocalAuthRejection::InvalidApiKey,
|
||||
));
|
||||
};
|
||||
if !auth_context.access_allowed
|
||||
|| auth_context.user_id.trim().is_empty()
|
||||
|| auth_context.api_key_id.trim().is_empty()
|
||||
{
|
||||
return Err(websocket_auth_rejection_error(
|
||||
GatewayLocalAuthRejection::InvalidApiKey,
|
||||
));
|
||||
}
|
||||
if !ip_rules_allow(auth_context.ip_rules.as_deref(), context.client_ip) {
|
||||
return Err(websocket_auth_rejection_error(
|
||||
GatewayLocalAuthRejection::IpNotAllowed {
|
||||
remote_ip: context.client_ip.to_string(),
|
||||
},
|
||||
));
|
||||
}
|
||||
|
||||
let body = serde_json::to_vec(client_event)
|
||||
.map(axum::body::Bytes::from)
|
||||
.map_err(|error| GatewayError::Internal(error.to_string()))?;
|
||||
if let Some(rejection) =
|
||||
request_model_local_rejection(state, Some(&decision), &parts.uri, &parts.headers, &body)
|
||||
.await?
|
||||
{
|
||||
return Err(websocket_auth_rejection_error(rejection));
|
||||
}
|
||||
|
||||
let rpm_bypassed = match state.admin_security_ip_whitelisted(context.client_ip).await {
|
||||
Ok(value) => value,
|
||||
Err(error) => {
|
||||
warn!(
|
||||
event_name = "responses_websocket_turn_ip_whitelist_check_failed",
|
||||
log_type = "ops",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %context.trace_id,
|
||||
client_ip = %context.client_ip,
|
||||
error = ?error,
|
||||
"gateway applied ordinary WebSocket RPM after the live IP whitelist check failed"
|
||||
);
|
||||
false
|
||||
}
|
||||
};
|
||||
|
||||
Ok(ResponsesWebSocketTurnControl {
|
||||
decision,
|
||||
auth_snapshot,
|
||||
rpm_bypassed,
|
||||
})
|
||||
}
|
||||
|
||||
fn websocket_auth_rejection_error(rejection: GatewayLocalAuthRejection) -> GatewayError {
|
||||
let (status, message) = match rejection {
|
||||
GatewayLocalAuthRejection::InvalidApiKey => {
|
||||
(StatusCode::UNAUTHORIZED, "The API key is invalid")
|
||||
}
|
||||
GatewayLocalAuthRejection::LockedApiKey => (
|
||||
StatusCode::FORBIDDEN,
|
||||
"The API key is locked and cannot be used",
|
||||
),
|
||||
GatewayLocalAuthRejection::WalletUnavailable => {
|
||||
(StatusCode::FORBIDDEN, "The account wallet is unavailable")
|
||||
}
|
||||
GatewayLocalAuthRejection::BalanceDenied { remaining } => {
|
||||
let message = match remaining {
|
||||
Some(remaining) => format!("Insufficient balance (remaining: ${remaining:.2})"),
|
||||
None => "Insufficient balance".to_string(),
|
||||
};
|
||||
return GatewayError::Client {
|
||||
status: StatusCode::TOO_MANY_REQUESTS,
|
||||
message,
|
||||
};
|
||||
}
|
||||
GatewayLocalAuthRejection::ProviderNotAllowed { .. } => (
|
||||
StatusCode::FORBIDDEN,
|
||||
"The provider is not allowed for this API key",
|
||||
),
|
||||
GatewayLocalAuthRejection::ApiFormatNotAllowed { .. } => (
|
||||
StatusCode::FORBIDDEN,
|
||||
"The API format is not allowed for this API key",
|
||||
),
|
||||
GatewayLocalAuthRejection::ModelNotAllowed { .. } => (
|
||||
StatusCode::FORBIDDEN,
|
||||
"The requested model is not allowed for this API key",
|
||||
),
|
||||
GatewayLocalAuthRejection::IpNotAllowed { .. } => (
|
||||
StatusCode::UNAUTHORIZED,
|
||||
"The current IP is not allowed for this API key",
|
||||
),
|
||||
};
|
||||
GatewayError::Client {
|
||||
status,
|
||||
message: message.to_string(),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,570 @@
|
||||
//! Parsed OpenAI Responses WebSocket text frames.
|
||||
//!
|
||||
//! A relay frame is parsed once and then shared by the protocol adapter, turn
|
||||
//! accounting, retry safety, and connection lifecycle code. Keeping the raw
|
||||
//! text as a borrow avoids copying the websocket payload while the relay is
|
||||
//! processing it.
|
||||
|
||||
use serde_json::Value;
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub(super) struct ResponsesWebSocketFrameTerminal {
|
||||
pub(super) status_code: u16,
|
||||
pub(super) cancelled: bool,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub(super) struct ParsedResponsesWebSocketFrame<'a> {
|
||||
raw_text: &'a str,
|
||||
event: Value,
|
||||
event_type: Option<String>,
|
||||
status: Option<u16>,
|
||||
started: bool,
|
||||
terminal: Option<ResponsesWebSocketFrameTerminal>,
|
||||
terminal_event: Option<Value>,
|
||||
chunked: bool,
|
||||
}
|
||||
|
||||
impl<'a> ParsedResponsesWebSocketFrame<'a> {
|
||||
pub(super) fn parse(raw_text: &'a str) -> serde_json::Result<Self> {
|
||||
let event = serde_json::from_str::<Value>(raw_text)?;
|
||||
let events = protocol_events_of(&event);
|
||||
let started = events.iter().copied().any(event_is_started);
|
||||
// A batch carries at most one terminal in practice. Taking the first
|
||||
// in document order keeps the outcome deterministic if that ever
|
||||
// stops being true.
|
||||
let terminal_entry = events
|
||||
.iter()
|
||||
.copied()
|
||||
.find_map(|candidate| terminal_for_event(candidate).map(|term| (candidate, term)));
|
||||
let terminal = terminal_entry.map(|(_, terminal)| terminal);
|
||||
// The terminal event describes the turn's outcome, so it is the one
|
||||
// worth naming in logs and recording as the terminal error body.
|
||||
let event_type = terminal_entry
|
||||
.map(|(candidate, _)| candidate)
|
||||
.or_else(|| events.last().copied())
|
||||
.and_then(event_type_of)
|
||||
.map(str::to_string);
|
||||
let terminal_event = terminal_entry.map(|(candidate, _)| candidate.clone());
|
||||
let chunked = event.get("chunks").and_then(Value::as_array).is_some();
|
||||
let status = terminal.map(|terminal| terminal.status_code);
|
||||
|
||||
Ok(Self {
|
||||
raw_text,
|
||||
event,
|
||||
event_type,
|
||||
status,
|
||||
started,
|
||||
terminal,
|
||||
terminal_event,
|
||||
chunked,
|
||||
})
|
||||
}
|
||||
|
||||
/// The protocol events this frame carries.
|
||||
///
|
||||
/// Codex batches standard `response.*` events into a `{"chunks":[...]}`
|
||||
/// envelope, so one frame can carry several events — and the terminal one
|
||||
/// may be buried inside the batch. Every consumer that interprets event
|
||||
/// semantics must walk this rather than the envelope, or a batched
|
||||
/// `response.completed` goes unnoticed and wedges the turn.
|
||||
pub(super) fn protocol_events(&self) -> Vec<&Value> {
|
||||
protocol_events_of(&self.event)
|
||||
}
|
||||
|
||||
/// The individual event that ended the turn, unwrapped from its batch.
|
||||
pub(super) fn terminal_event(&self) -> Option<&Value> {
|
||||
self.terminal_event.as_ref()
|
||||
}
|
||||
|
||||
pub(super) fn is_chunked(&self) -> bool {
|
||||
self.chunked
|
||||
}
|
||||
|
||||
pub(super) fn raw_text(&self) -> &'a str {
|
||||
self.raw_text
|
||||
}
|
||||
|
||||
pub(super) fn event(&self) -> &Value {
|
||||
&self.event
|
||||
}
|
||||
|
||||
pub(super) fn event_type(&self) -> Option<&str> {
|
||||
self.event_type.as_deref()
|
||||
}
|
||||
|
||||
pub(super) fn status(&self) -> Option<u16> {
|
||||
self.status
|
||||
}
|
||||
|
||||
pub(super) fn is_started(&self) -> bool {
|
||||
self.started
|
||||
}
|
||||
|
||||
pub(super) fn is_terminal(&self) -> bool {
|
||||
self.terminal.is_some()
|
||||
}
|
||||
|
||||
pub(super) fn terminal(&self) -> Option<ResponsesWebSocketFrameTerminal> {
|
||||
self.terminal
|
||||
}
|
||||
|
||||
/// Return a bounded label suitable for structured logs. Event payloads
|
||||
/// are never inserted directly into a log field.
|
||||
pub(super) fn event_type_for_log(&self) -> String {
|
||||
self.event_type
|
||||
.as_deref()
|
||||
.map(safe_websocket_event_label)
|
||||
.unwrap_or_else(|| "invalid_json".to_string())
|
||||
}
|
||||
}
|
||||
|
||||
/// Encodes one event peeled from a provider-private envelope without applying
|
||||
/// an event-type or field projection.
|
||||
///
|
||||
/// Direct provider events should use [`ParsedResponsesWebSocketFrame::raw_text`]
|
||||
/// so their bytes remain identical. This helper exists only for batch
|
||||
/// envelopes that cannot be relayed as a whole: serializing the complete
|
||||
/// [`Value`] preserves every known and future JSON member.
|
||||
pub(super) fn encode_opaque_websocket_event(event: &Value) -> serde_json::Result<String> {
|
||||
serde_json::to_string(event)
|
||||
}
|
||||
|
||||
/// Flattens a frame into the events it carries. An envelope may name its own
|
||||
/// `type` *and* batch further events under `chunks`; both are protocol events.
|
||||
fn protocol_events_of(event: &Value) -> Vec<&Value> {
|
||||
let mut events = Vec::new();
|
||||
if event_type_of(event).is_some() {
|
||||
events.push(event);
|
||||
}
|
||||
if let Some(chunks) = event.get("chunks").and_then(Value::as_array) {
|
||||
events.extend(chunks.iter().filter(|chunk| event_type_of(chunk).is_some()));
|
||||
}
|
||||
// An unrecognized shape is still relayed and still accounted for, so it
|
||||
// must not vanish from the observer's view of the stream.
|
||||
if events.is_empty() {
|
||||
events.push(event);
|
||||
}
|
||||
events
|
||||
}
|
||||
|
||||
fn event_type_of(event: &Value) -> Option<&str> {
|
||||
event.get("type").and_then(Value::as_str)
|
||||
}
|
||||
|
||||
fn event_is_started(event: &Value) -> bool {
|
||||
matches!(
|
||||
event_type_of(event).unwrap_or_default(),
|
||||
"response.created" | "response.in_progress" | "response.queued"
|
||||
)
|
||||
}
|
||||
|
||||
/// 读取 `response.incomplete` 携带的 `incomplete_details.reason`。
|
||||
///
|
||||
/// 标准位置是 `response.incomplete_details.reason`;批量封装偶尔把
|
||||
/// `incomplete_details` 直接放在事件顶层,两处都要看,否则合法终态会被漏判。
|
||||
fn responses_incomplete_reason(event: &Value) -> Option<&str> {
|
||||
[
|
||||
event.pointer("/response/incomplete_details/reason"),
|
||||
event.pointer("/incomplete_details/reason"),
|
||||
]
|
||||
.into_iter()
|
||||
.flatten()
|
||||
.filter_map(Value::as_str)
|
||||
.map(str::trim)
|
||||
.find(|reason| !reason.is_empty())
|
||||
}
|
||||
|
||||
fn responses_incomplete_has_explicit_error(event: &Value) -> bool {
|
||||
[event.get("error"), event.pointer("/response/error")]
|
||||
.into_iter()
|
||||
.flatten()
|
||||
.any(|error| !error.is_null())
|
||||
}
|
||||
|
||||
/// Derives only the fallback status for `response.incomplete`.
|
||||
///
|
||||
/// A non-empty reason is provider-owned protocol data. Treating it as a fixed
|
||||
/// allowlist would turn every future legitimate reason into a synthetic 502
|
||||
/// and incorrectly penalize provider health. Missing/malformed reasons and
|
||||
/// explicit error markers still fail closed; numeric status and recognized
|
||||
/// error codes continue to override this fallback in
|
||||
/// [`websocket_event_status_code`].
|
||||
fn responses_incomplete_default_status(event: &Value) -> u16 {
|
||||
match responses_incomplete_reason(event) {
|
||||
None => 502,
|
||||
Some(reason)
|
||||
if reason.eq_ignore_ascii_case("error")
|
||||
|| reason.eq_ignore_ascii_case("server_error") =>
|
||||
{
|
||||
502
|
||||
}
|
||||
Some(_) if responses_incomplete_has_explicit_error(event) => 502,
|
||||
Some(_) => 200,
|
||||
}
|
||||
}
|
||||
|
||||
fn terminal_for_event(event: &Value) -> Option<ResponsesWebSocketFrameTerminal> {
|
||||
match event_type_of(event).unwrap_or_default() {
|
||||
"response.completed" => Some(ResponsesWebSocketFrameTerminal {
|
||||
status_code: websocket_event_status_code(event, 200),
|
||||
cancelled: false,
|
||||
}),
|
||||
// A non-empty provider reason is a normal terminal by default, including
|
||||
// future reasons Aether does not yet know. Explicit status/error data
|
||||
// still wins, so quota and server failures retain their failure status.
|
||||
"response.incomplete" => Some(ResponsesWebSocketFrameTerminal {
|
||||
status_code: websocket_event_status_code(
|
||||
event,
|
||||
responses_incomplete_default_status(event),
|
||||
),
|
||||
cancelled: false,
|
||||
}),
|
||||
"response.cancelled" => Some(ResponsesWebSocketFrameTerminal {
|
||||
status_code: 499,
|
||||
cancelled: true,
|
||||
}),
|
||||
"response.failed" => Some(ResponsesWebSocketFrameTerminal {
|
||||
status_code: websocket_event_status_code(event, 502),
|
||||
cancelled: false,
|
||||
}),
|
||||
"error" => Some(ResponsesWebSocketFrameTerminal {
|
||||
status_code: websocket_event_status_code(event, 502),
|
||||
cancelled: false,
|
||||
}),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn websocket_event_status_code(event: &Value, default: u16) -> u16 {
|
||||
if let Some(status_code) = event
|
||||
.get("status_code")
|
||||
.or_else(|| event.get("status"))
|
||||
.or_else(|| {
|
||||
event
|
||||
.get("response")
|
||||
.and_then(|response| response.get("status_code"))
|
||||
})
|
||||
.and_then(Value::as_u64)
|
||||
.and_then(|value| u16::try_from(value).ok())
|
||||
.filter(|value| *value > 0)
|
||||
{
|
||||
return status_code;
|
||||
}
|
||||
|
||||
let error_code = [
|
||||
event.pointer("/error/type"),
|
||||
event.pointer("/error/code"),
|
||||
event.pointer("/response/error/type"),
|
||||
event.pointer("/response/error/code"),
|
||||
]
|
||||
.into_iter()
|
||||
.flatten()
|
||||
.filter_map(Value::as_str)
|
||||
.map(str::to_ascii_lowercase)
|
||||
.find(|value| !value.trim().is_empty());
|
||||
match error_code.as_deref() {
|
||||
Some(
|
||||
"usage_limit_reached" | "insufficient_quota" | "rate_limit_exceeded" | "quota_exceeded",
|
||||
) => 429,
|
||||
Some("invalid_api_key" | "authentication_error") => 401,
|
||||
Some("invalid_request_error" | "invalid_request" | "model_not_found") => 400,
|
||||
Some("overloaded" | "server_error" | "service_unavailable") => 503,
|
||||
_ => default,
|
||||
}
|
||||
}
|
||||
|
||||
fn safe_websocket_event_label(value: &str) -> String {
|
||||
let value = value.trim();
|
||||
if value.is_empty()
|
||||
|| value.len() > 80
|
||||
|| !value
|
||||
.bytes()
|
||||
.all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'.' | b'_' | b'-'))
|
||||
{
|
||||
return "unknown".to_string();
|
||||
}
|
||||
value.to_string()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::{encode_opaque_websocket_event, ParsedResponsesWebSocketFrame};
|
||||
|
||||
#[test]
|
||||
fn parses_started_frame_once_with_raw_text_and_event_metadata() {
|
||||
let raw = r#"{"type":"response.in_progress","response":{"status":200}}"#;
|
||||
let frame = ParsedResponsesWebSocketFrame::parse(raw).expect("valid frame");
|
||||
|
||||
assert_eq!(frame.raw_text(), raw);
|
||||
assert_eq!(frame.event_type(), Some("response.in_progress"));
|
||||
assert_eq!(frame.status(), None);
|
||||
assert!(frame.is_started());
|
||||
assert!(!frame.is_terminal());
|
||||
assert_eq!(frame.event()["response"]["status"], 200);
|
||||
assert_eq!(frame.event_type_for_log(), "response.in_progress");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn future_response_event_keeps_its_exact_original_text_and_unknown_fields() {
|
||||
let raw = "{ \n \"future_top_level\": {\"nested\": [1, true, null]}, \n \"type\": \"response.future_capability.delta\", \n \"delta\": {\"new_wire_shape\": \"opaque\"}\n}";
|
||||
let frame = ParsedResponsesWebSocketFrame::parse(raw).expect("valid future event");
|
||||
|
||||
assert_eq!(frame.raw_text(), raw);
|
||||
assert_eq!(
|
||||
frame.event()["future_top_level"],
|
||||
json!({"nested": [1, true, null]})
|
||||
);
|
||||
assert_eq!(frame.event()["delta"], json!({"new_wire_shape": "opaque"}));
|
||||
assert!(!frame.is_terminal());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn peeled_batch_event_encoding_preserves_the_complete_opaque_value() {
|
||||
let frame = ParsedResponsesWebSocketFrame::parse(
|
||||
r#"{"chunks":[{"type":"response.future.done","future_capability":{"mode":"new"},"response":{"id":"resp_future","future_usage":{"novel_tokens":7}}}]}"#,
|
||||
)
|
||||
.expect("valid private envelope");
|
||||
let events = frame.protocol_events();
|
||||
let event = events.first().expect("one future response event");
|
||||
let encoded = encode_opaque_websocket_event(event).expect("Value serialization succeeds");
|
||||
let round_trip: serde_json::Value =
|
||||
serde_json::from_str(&encoded).expect("encoded event stays valid JSON");
|
||||
|
||||
assert_eq!(round_trip, **event);
|
||||
assert_eq!(round_trip["future_capability"], json!({"mode": "new"}));
|
||||
assert_eq!(round_trip["response"]["future_usage"]["novel_tokens"], 7);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn classifies_terminal_status_and_cancellation() {
|
||||
let completed = ParsedResponsesWebSocketFrame::parse(
|
||||
r#"{"type":"response.completed","status_code":201}"#,
|
||||
)
|
||||
.expect("valid frame");
|
||||
assert_eq!(completed.status(), Some(201));
|
||||
assert_eq!(
|
||||
completed
|
||||
.terminal()
|
||||
.map(|terminal| (terminal.status_code, terminal.cancelled)),
|
||||
Some((201, false))
|
||||
);
|
||||
|
||||
let cancelled = ParsedResponsesWebSocketFrame::parse(r#"{"type":"response.cancelled"}"#)
|
||||
.expect("valid frame");
|
||||
assert_eq!(cancelled.status(), Some(499));
|
||||
assert_eq!(
|
||||
cancelled
|
||||
.terminal()
|
||||
.map(|terminal| (terminal.status_code, terminal.cancelled)),
|
||||
Some((499, true))
|
||||
);
|
||||
|
||||
let error = ParsedResponsesWebSocketFrame::parse(
|
||||
r#"{"type":"error","status_code":429,"error":{"type":"usage_limit_reached"}}"#,
|
||||
)
|
||||
.expect("valid frame");
|
||||
assert_eq!(error.status(), Some(429));
|
||||
assert!(error.is_terminal());
|
||||
|
||||
let failed = ParsedResponsesWebSocketFrame::parse(
|
||||
r#"{"type":"response.failed","response":{"error":{"code":"rate_limit_exceeded"}}}"#,
|
||||
)
|
||||
.expect("valid frame");
|
||||
assert_eq!(failed.status(), Some(429));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn a_legitimate_incomplete_is_a_terminal_but_not_a_provider_failure() {
|
||||
for reason in [
|
||||
"max_output_tokens",
|
||||
"max_tokens",
|
||||
"content_filter",
|
||||
"tool_calls",
|
||||
"function_call",
|
||||
"MAX_OUTPUT_TOKENS",
|
||||
] {
|
||||
let raw = format!(
|
||||
r#"{{"type":"response.incomplete","response":{{"status":"incomplete","incomplete_details":{{"reason":"{reason}"}}}}}}"#
|
||||
);
|
||||
let frame = ParsedResponsesWebSocketFrame::parse(&raw).expect("valid frame");
|
||||
|
||||
assert!(frame.is_terminal(), "{reason} should end the turn");
|
||||
assert_eq!(
|
||||
frame
|
||||
.terminal()
|
||||
.map(|terminal| (terminal.status_code, terminal.cancelled)),
|
||||
Some((200, false)),
|
||||
"{reason} is a legitimate terminal result, not a 502 provider failure"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn a_top_level_incomplete_details_reason_is_also_honored() {
|
||||
let frame = ParsedResponsesWebSocketFrame::parse(
|
||||
r#"{"type":"response.incomplete","incomplete_details":{"reason":"max_output_tokens"}}"#,
|
||||
)
|
||||
.expect("valid frame");
|
||||
|
||||
assert_eq!(frame.status(), Some(200));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn an_incomplete_without_a_reason_or_with_a_failure_reason_stays_a_provider_failure() {
|
||||
for raw in [
|
||||
r#"{"type":"response.incomplete"}"#,
|
||||
r#"{"type":"response.incomplete","response":{"incomplete_details":null}}"#,
|
||||
r#"{"type":"response.incomplete","response":{"incomplete_details":{"reason":""}}}"#,
|
||||
r#"{"type":"response.incomplete","response":{"incomplete_details":{"reason":"error"}}}"#,
|
||||
r#"{"type":"response.incomplete","response":{"incomplete_details":{"reason":"server_error"}}}"#,
|
||||
] {
|
||||
let frame = ParsedResponsesWebSocketFrame::parse(raw).expect("valid frame");
|
||||
|
||||
assert_eq!(
|
||||
frame.status(),
|
||||
Some(502),
|
||||
"an incomplete without a usable reason must stay a provider failure: {raw}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn a_future_incomplete_reason_is_forward_compatible_without_hiding_explicit_errors() {
|
||||
let future = ParsedResponsesWebSocketFrame::parse(
|
||||
r#"{"type":"response.incomplete","response":{"incomplete_details":{"reason":"future_context_boundary"}}}"#,
|
||||
)
|
||||
.expect("valid frame");
|
||||
assert_eq!(future.status(), Some(200));
|
||||
|
||||
let future_with_error = ParsedResponsesWebSocketFrame::parse(
|
||||
r#"{"type":"response.incomplete","response":{"error":{"code":"future_provider_error"},"incomplete_details":{"reason":"future_context_boundary"}}}"#,
|
||||
)
|
||||
.expect("valid frame");
|
||||
assert_eq!(future_with_error.status(), Some(502));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn a_legitimate_incomplete_still_respects_an_explicit_provider_status() {
|
||||
let explicit = ParsedResponsesWebSocketFrame::parse(
|
||||
r#"{"type":"response.incomplete","status_code":503,"response":{"incomplete_details":{"reason":"max_output_tokens"}}}"#,
|
||||
)
|
||||
.expect("valid frame");
|
||||
assert_eq!(explicit.status(), Some(503));
|
||||
|
||||
let quota = ParsedResponsesWebSocketFrame::parse(
|
||||
r#"{"type":"response.incomplete","response":{"error":{"code":"rate_limit_exceeded"},"incomplete_details":{"reason":"max_output_tokens"}}}"#,
|
||||
)
|
||||
.expect("valid frame");
|
||||
assert_eq!(quota.status(), Some(429));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn a_legitimate_incomplete_batched_inside_a_chunks_envelope_is_not_a_failure() {
|
||||
let frame = ParsedResponsesWebSocketFrame::parse(
|
||||
r#"{"chunks":[{"type":"response.output_text.delta","delta":"hi"},{"type":"response.incomplete","response":{"incomplete_details":{"reason":"max_output_tokens"},"usage":{"total_tokens":9}}}]}"#,
|
||||
)
|
||||
.expect("valid frame");
|
||||
|
||||
assert!(frame.is_chunked());
|
||||
assert!(frame.is_terminal());
|
||||
assert_eq!(frame.status(), Some(200));
|
||||
assert_eq!(frame.event_type(), Some("response.incomplete"));
|
||||
assert_eq!(
|
||||
frame.terminal_event().and_then(|event| event
|
||||
.pointer("/response/usage/total_tokens")
|
||||
.and_then(serde_json::Value::as_u64)),
|
||||
Some(9)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn detects_a_terminal_batched_inside_a_chunks_envelope() {
|
||||
let frame = ParsedResponsesWebSocketFrame::parse(
|
||||
r#"{"chunks":[{"type":"response.output_text.delta","delta":"hi"},{"type":"response.completed","response":{"usage":{"total_tokens":8}}}]}"#,
|
||||
)
|
||||
.expect("valid frame");
|
||||
|
||||
assert!(frame.is_chunked());
|
||||
assert!(frame.is_terminal());
|
||||
assert_eq!(frame.status(), Some(200));
|
||||
// The label and the recorded error body must name the event that ended
|
||||
// the turn, not the envelope.
|
||||
assert_eq!(frame.event_type(), Some("response.completed"));
|
||||
assert_eq!(
|
||||
frame.terminal_event().and_then(|event| event
|
||||
.pointer("/response/usage/total_tokens")
|
||||
.and_then(serde_json::Value::as_u64)),
|
||||
Some(8)
|
||||
);
|
||||
assert_eq!(frame.protocol_events().len(), 2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn detects_a_start_event_batched_inside_a_chunks_envelope() {
|
||||
let frame = ParsedResponsesWebSocketFrame::parse(
|
||||
r#"{"chunks":[{"type":"codex.rate_limits"},{"type":"response.created"}]}"#,
|
||||
)
|
||||
.expect("valid frame");
|
||||
|
||||
assert!(frame.is_started());
|
||||
assert!(!frame.is_terminal());
|
||||
assert_eq!(frame.protocol_events().len(), 2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn an_envelope_may_carry_its_own_type_alongside_batched_events() {
|
||||
let frame = ParsedResponsesWebSocketFrame::parse(
|
||||
r#"{"type":"codex.response.metadata","chunks":[{"type":"response.failed","response":{"error":{"code":"rate_limit_exceeded"}}}]}"#,
|
||||
)
|
||||
.expect("valid frame");
|
||||
|
||||
assert_eq!(frame.protocol_events().len(), 2);
|
||||
assert!(frame.is_terminal());
|
||||
assert_eq!(frame.status(), Some(429));
|
||||
assert_eq!(frame.event_type(), Some("response.failed"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn a_batch_without_a_terminal_does_not_end_the_turn() {
|
||||
let frame = ParsedResponsesWebSocketFrame::parse(
|
||||
r#"{"chunks":[{"type":"response.output_text.delta","delta":"a"},{"type":"response.output_text.delta","delta":"b"}]}"#,
|
||||
)
|
||||
.expect("valid frame");
|
||||
|
||||
assert!(!frame.is_terminal());
|
||||
assert!(!frame.is_started());
|
||||
assert!(frame.terminal_event().is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn an_unrecognized_shape_is_still_surfaced_as_one_event() {
|
||||
let frame =
|
||||
ParsedResponsesWebSocketFrame::parse(r#"{"unexpected":true}"#).expect("valid frame");
|
||||
|
||||
assert_eq!(frame.protocol_events().len(), 1);
|
||||
assert!(!frame.is_chunked());
|
||||
assert!(!frame.is_terminal());
|
||||
assert_eq!(frame.event_type(), None);
|
||||
assert_eq!(frame.event_type_for_log(), "invalid_json");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn preserves_safe_log_label_boundaries() {
|
||||
let unsafe_label =
|
||||
ParsedResponsesWebSocketFrame::parse(r#"{"type":"not safe / contains spaces"}"#)
|
||||
.expect("valid frame");
|
||||
assert_eq!(unsafe_label.event_type_for_log(), "unknown");
|
||||
|
||||
let missing_label =
|
||||
ParsedResponsesWebSocketFrame::parse(r#"{"message":"ok"}"#).expect("valid frame");
|
||||
assert_eq!(missing_label.event_type_for_log(), "invalid_json");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_invalid_json() {
|
||||
assert!(ParsedResponsesWebSocketFrame::parse("not-json").is_err());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,588 @@
|
||||
//! Turn finalization and terminal error mapping for a Responses WebSocket.
|
||||
//!
|
||||
//! A connection can outlive a turn, so persistence and adapter observation
|
||||
//! handles are joined in order before the next turn is planned.
|
||||
|
||||
use std::time::Duration;
|
||||
|
||||
use axum::extract::ws::WebSocket;
|
||||
use axum::http::StatusCode;
|
||||
use tokio::task::JoinHandle;
|
||||
use tokio::time::timeout;
|
||||
|
||||
use super::state::BoundResponsesConnection;
|
||||
use super::turn::{
|
||||
begin_unowned_responses_websocket_turn, ResponsesProviderAttempt, ResponsesWebSocketTurnOutcome,
|
||||
};
|
||||
use crate::handlers::proxy::websocket::session::{
|
||||
CLOSE_INTERNAL_ERROR, CLOSE_POLICY_VIOLATION, CLOSE_TRY_AGAIN, WEBSOCKET_LOG_TRANSPORT,
|
||||
};
|
||||
use crate::handlers::proxy::websocket::transport::send_responses_websocket_error;
|
||||
use crate::{AppState, GatewayError};
|
||||
|
||||
const RESPONSES_WEBSOCKET_ADAPTER_OBSERVATION_TIMEOUT: Duration = Duration::from_secs(5);
|
||||
const LOG_TARGET: &str = "aether_gateway::handlers::proxy::responses_ws";
|
||||
|
||||
macro_rules! warn {
|
||||
($($arg:tt)*) => {
|
||||
tracing::warn!(target: LOG_TARGET, $($arg)*)
|
||||
};
|
||||
}
|
||||
|
||||
/// Owns the in-flight turn so that losing the relay task still finalizes it.
|
||||
///
|
||||
/// Every ordinary exit path takes the turn out of here and finalizes it
|
||||
/// explicitly. This guard only covers the paths that are not exit paths at all
|
||||
/// — a panic in the relay loop, or the task being dropped — where the turn
|
||||
/// would otherwise be discarded with its usage row left `Pending`, its
|
||||
/// candidate row left `Streaming`, and its distributed pool key lease leaked
|
||||
/// until the lease expires. Mirrors the HTTP path's `DirectPassthroughFinalizer`.
|
||||
pub(super) struct ActiveProviderAttempt {
|
||||
turn: Option<ResponsesProviderAttempt>,
|
||||
state: AppState,
|
||||
}
|
||||
|
||||
impl ActiveProviderAttempt {
|
||||
pub(super) fn new(state: &AppState, turn: ResponsesProviderAttempt) -> Self {
|
||||
Self {
|
||||
turn: Some(turn),
|
||||
state: state.clone(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Hands the turn back to a caller that will finalize it explicitly.
|
||||
pub(super) fn disarm(mut self) -> ResponsesProviderAttempt {
|
||||
self.turn
|
||||
.take()
|
||||
.expect("an armed active turn always holds its turn")
|
||||
}
|
||||
}
|
||||
|
||||
impl std::ops::Deref for ActiveProviderAttempt {
|
||||
type Target = ResponsesProviderAttempt;
|
||||
|
||||
fn deref(&self) -> &Self::Target {
|
||||
self.turn
|
||||
.as_ref()
|
||||
.expect("an armed active turn always holds its turn")
|
||||
}
|
||||
}
|
||||
|
||||
impl std::ops::DerefMut for ActiveProviderAttempt {
|
||||
fn deref_mut(&mut self) -> &mut Self::Target {
|
||||
self.turn
|
||||
.as_mut()
|
||||
.expect("an armed active turn always holds its turn")
|
||||
}
|
||||
}
|
||||
|
||||
/// Starts a turn and arms its cancellation fallback before control returns to
|
||||
/// code that can await an upstream bind or socket write.
|
||||
pub(super) async fn begin_responses_websocket_turn(
|
||||
state: &AppState,
|
||||
trace_id: &str,
|
||||
parts: http::request::Parts,
|
||||
control_decision: &crate::control::GatewayControlDecision,
|
||||
decision: crate::ai_serving::AiExecutionDecision,
|
||||
client_event: &serde_json::Value,
|
||||
) -> Result<ActiveProviderAttempt, GatewayError> {
|
||||
let state = state.clone();
|
||||
let trace_id = trace_id.to_string();
|
||||
let owner_timeout = state
|
||||
.frontdoor_runtime_guards
|
||||
.local_execution_planning_timeout;
|
||||
let control_decision = control_decision.clone();
|
||||
let client_event = client_event.clone();
|
||||
|
||||
// Beginning an attempt performs several indispensable async writes before
|
||||
// an `ActiveProviderAttempt` can exist (balance/admission, Pending usage,
|
||||
// and candidate state). Run that whole transition in an owned task. If the
|
||||
// relay/session future is cancelled while awaiting it, Tokio detaches this
|
||||
// task; it still reaches either an explicitly cleaned-up error or an armed
|
||||
// guard whose dropped output finalizes the attempt.
|
||||
await_owned_turn_begin(
|
||||
async move {
|
||||
let turn = begin_unowned_responses_websocket_turn(
|
||||
&state,
|
||||
&parts,
|
||||
&control_decision,
|
||||
decision,
|
||||
&client_event,
|
||||
)
|
||||
.await?;
|
||||
Ok(ActiveProviderAttempt::new(&state, turn))
|
||||
},
|
||||
owner_timeout,
|
||||
trace_id,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn await_owned_turn_begin<T>(
|
||||
begin: impl std::future::Future<Output = Result<T, GatewayError>> + Send + 'static,
|
||||
owner_timeout: Duration,
|
||||
trace_id: String,
|
||||
) -> Result<T, GatewayError>
|
||||
where
|
||||
T: Send + 'static,
|
||||
{
|
||||
await_owned_turn_begin_with_timeout(begin, owner_timeout, trace_id).await
|
||||
}
|
||||
|
||||
async fn await_owned_turn_begin_with_timeout<T>(
|
||||
begin: impl std::future::Future<Output = Result<T, GatewayError>> + Send + 'static,
|
||||
owner_timeout: Duration,
|
||||
trace_id: String,
|
||||
) -> Result<T, GatewayError>
|
||||
where
|
||||
T: Send + 'static,
|
||||
{
|
||||
tokio::spawn(async move {
|
||||
tokio::time::timeout(owner_timeout, begin)
|
||||
.await
|
||||
.map_err(|_| GatewayError::LocalExecutionPlanningTimeout {
|
||||
trace_id,
|
||||
phase: "responses_websocket_turn_begin_owner",
|
||||
timeout_ms: owner_timeout.as_millis() as u64,
|
||||
})?
|
||||
})
|
||||
.await
|
||||
.map_err(|error| {
|
||||
GatewayError::Internal(format!(
|
||||
"Responses WebSocket turn begin task failed before ownership transfer: {error}"
|
||||
))
|
||||
})?
|
||||
}
|
||||
|
||||
impl Drop for ActiveProviderAttempt {
|
||||
fn drop(&mut self) {
|
||||
let Some(turn) = self.turn.take() else {
|
||||
return;
|
||||
};
|
||||
let outcome = turn.abandonment_outcome();
|
||||
let state = self.state.clone();
|
||||
// No runtime means the process is going down; the spawn could not
|
||||
// complete anyway.
|
||||
if let Ok(handle) = tokio::runtime::Handle::try_current() {
|
||||
warn!(
|
||||
event_name = "responses_websocket_turn_abandoned",
|
||||
log_type = "ops",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
"gateway finalized a Responses WebSocket turn whose relay task went away"
|
||||
);
|
||||
handle.spawn(async move {
|
||||
turn.finalize_detached(&state, outcome).await;
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 结束当前 logical turn 并结算它的 attempt。
|
||||
///
|
||||
/// `end()` 同时清掉 logical turn 和 attempt,取代原来「take active_turn +
|
||||
/// 在每个出口手写 `active_response_create = None`」的两步组合。
|
||||
pub(super) async fn finalize_active_turn(
|
||||
bound: &mut BoundResponsesConnection,
|
||||
state: &AppState,
|
||||
outcome: ResponsesWebSocketTurnOutcome,
|
||||
) {
|
||||
if let Some(turn) = bound.turn_state.end() {
|
||||
queue_turn_finalization(bound, state, turn, outcome).await;
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) async fn queue_turn_finalization(
|
||||
bound: &mut BoundResponsesConnection,
|
||||
state: &AppState,
|
||||
turn: ActiveProviderAttempt,
|
||||
outcome: ResponsesWebSocketTurnOutcome,
|
||||
) {
|
||||
await_pending_adapter_observation(bound).await;
|
||||
await_pending_turn_finalization(bound).await;
|
||||
bound.pending_turn_finalization = Some(spawn_guarded_turn_finalization(
|
||||
state.clone(),
|
||||
turn,
|
||||
outcome,
|
||||
));
|
||||
}
|
||||
|
||||
/// 「上一个 attempt 已经结算完毕」的凭证。
|
||||
///
|
||||
/// 只能由本模块颁发,且只有在结算真正落地之后。规划下一个 attempt 的入口
|
||||
/// ([`super::quota::retry_active_turn_after_quota_exhaustion`]) 要求这个参数,
|
||||
/// 于是「先结算、再规划」成为签名的一部分,而不是一句注释——顺序写反连编译都
|
||||
/// 过不了。
|
||||
pub(super) struct PreviousAttemptSettled(());
|
||||
|
||||
impl PreviousAttemptSettled {
|
||||
/// 没有 attempt 要结算(连接此刻不在 `Responding`)。
|
||||
pub(super) const fn nothing_to_settle() -> Self {
|
||||
Self(())
|
||||
}
|
||||
}
|
||||
|
||||
/// 结算一个 attempt 并等它落地。
|
||||
///
|
||||
/// 与 [`queue_turn_finalization`] 的区别只在于「等」:后者把 handle 挂在连接上
|
||||
/// 让 relay loop 继续跑,适用于结算之后不再需要读取共享状态的出口;这个用在
|
||||
/// 必须先看到结算结果才能继续的路径上——典型的就是透明重试,它紧接着要按
|
||||
/// health / adaptive / pool 状态规划下一个 attempt。
|
||||
pub(super) async fn settle_turn_finalization(
|
||||
bound: &mut BoundResponsesConnection,
|
||||
state: &AppState,
|
||||
turn: ActiveProviderAttempt,
|
||||
outcome: ResponsesWebSocketTurnOutcome,
|
||||
) -> PreviousAttemptSettled {
|
||||
queue_turn_finalization(bound, state, turn, outcome).await;
|
||||
await_pending_turn_finalization(bound).await;
|
||||
PreviousAttemptSettled(())
|
||||
}
|
||||
|
||||
pub(super) fn spawn_bounded_adapter_observation(
|
||||
observation: impl std::future::Future<Output = ()> + Send + 'static,
|
||||
) -> JoinHandle<()> {
|
||||
spawn_bounded_adapter_observation_with_timeout(
|
||||
observation,
|
||||
RESPONSES_WEBSOCKET_ADAPTER_OBSERVATION_TIMEOUT,
|
||||
)
|
||||
}
|
||||
|
||||
fn spawn_bounded_adapter_observation_with_timeout(
|
||||
observation: impl std::future::Future<Output = ()> + Send + 'static,
|
||||
owner_timeout: Duration,
|
||||
) -> JoinHandle<()> {
|
||||
tokio::spawn(async move {
|
||||
if timeout(owner_timeout, observation).await.is_err() {
|
||||
warn!(
|
||||
event_name = "responses_websocket_adapter_observation_timeout",
|
||||
log_type = "ops",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
timeout_ms = owner_timeout.as_millis() as u64,
|
||||
"gateway stopped a timed-out Responses WebSocket adapter observation"
|
||||
);
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) async fn await_pending_adapter_observation(bound: &mut BoundResponsesConnection) {
|
||||
if let Some(handle) = bound.pending_adapter_observation.take() {
|
||||
if let Err(error) = handle.await {
|
||||
warn!(
|
||||
event_name = "responses_websocket_adapter_observation_join_failed",
|
||||
log_type = "ops",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
error = ?error,
|
||||
"gateway Responses WebSocket adapter observation task failed"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn finalize_unbound_turn(
|
||||
state: AppState,
|
||||
turn: ActiveProviderAttempt,
|
||||
outcome: ResponsesWebSocketTurnOutcome,
|
||||
) -> JoinHandle<()> {
|
||||
spawn_guarded_turn_finalization(state, turn, outcome)
|
||||
}
|
||||
|
||||
fn spawn_guarded_turn_finalization(
|
||||
state: AppState,
|
||||
turn: ActiveProviderAttempt,
|
||||
outcome: ResponsesWebSocketTurnOutcome,
|
||||
) -> JoinHandle<()> {
|
||||
// Spawn synchronously while the armed guard is still owned here. Caller
|
||||
// cancellation cannot drop an unguarded attempt between cleanup awaits.
|
||||
tokio::spawn(async move {
|
||||
let mut turn = turn;
|
||||
turn.release_admission().await;
|
||||
turn.disarm().finalize_detached(&state, outcome).await;
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) async fn await_turn_finalization_handle(handle: JoinHandle<()>) {
|
||||
// Do not abort terminal persistence here. Each I/O stage inside the turn
|
||||
// finalizer is independently bounded, and aborting the owner would skip
|
||||
// pool-lease cleanup and leave usage/candidate state non-terminal.
|
||||
match handle.await {
|
||||
Ok(()) => {}
|
||||
Err(error) => {
|
||||
warn!(
|
||||
event_name = "responses_websocket_turn_finalization_join_failed",
|
||||
log_type = "ops",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
error = ?error,
|
||||
"gateway Responses WebSocket turn finalizer task failed"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) async fn await_pending_turn_finalization(bound: &mut BoundResponsesConnection) {
|
||||
if let Some(handle) = bound.pending_turn_finalization.take() {
|
||||
await_turn_finalization_handle(handle).await;
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) async fn send_responses_websocket_turn_start_error(
|
||||
client_socket: &mut WebSocket,
|
||||
error: &GatewayError,
|
||||
) {
|
||||
let status_code = responses_websocket_turn_start_http_status(error);
|
||||
match error {
|
||||
GatewayError::Client { status, message } => {
|
||||
let (error_type, code) = if status.as_u16() == 429 {
|
||||
("rate_limit_error", "gateway_request_capacity_exceeded")
|
||||
} else {
|
||||
("invalid_request_error", "gateway_request_not_allowed")
|
||||
};
|
||||
send_responses_websocket_error(client_socket, status_code, error_type, code, message)
|
||||
.await;
|
||||
}
|
||||
GatewayError::AdmissionTimeout { .. } => {
|
||||
send_responses_websocket_error(
|
||||
client_socket,
|
||||
status_code,
|
||||
"server_error",
|
||||
"gateway_admission_timeout",
|
||||
"Gateway capacity is busy; retry this response",
|
||||
)
|
||||
.await;
|
||||
}
|
||||
GatewayError::LocalExecutionPlanningTimeout { .. } => {
|
||||
send_responses_websocket_error(
|
||||
client_socket,
|
||||
status_code,
|
||||
"server_error",
|
||||
"gateway_planning_timeout",
|
||||
"Gateway planning timed out; retry this response",
|
||||
)
|
||||
.await;
|
||||
}
|
||||
_ => {
|
||||
send_responses_websocket_error(
|
||||
client_socket,
|
||||
status_code,
|
||||
"server_error",
|
||||
"responses_websocket_turn_start_failed",
|
||||
"Gateway could not start this response",
|
||||
)
|
||||
.await;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn responses_websocket_turn_start_http_status(error: &GatewayError) -> u16 {
|
||||
match error {
|
||||
GatewayError::Client { status, .. } => status.as_u16(),
|
||||
GatewayError::AdmissionTimeout { .. } => StatusCode::TOO_MANY_REQUESTS.as_u16(),
|
||||
GatewayError::LocalExecutionPlanningTimeout { .. } => StatusCode::GATEWAY_TIMEOUT.as_u16(),
|
||||
_ => StatusCode::INTERNAL_SERVER_ERROR.as_u16(),
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn responses_websocket_turn_start_close(error: &GatewayError) -> (u16, &'static str) {
|
||||
match error {
|
||||
GatewayError::Client { .. } => (CLOSE_POLICY_VIOLATION, "request_not_allowed"),
|
||||
GatewayError::AdmissionTimeout { .. }
|
||||
| GatewayError::LocalExecutionPlanningTimeout { .. } => (CLOSE_TRY_AGAIN, "gateway_busy"),
|
||||
_ => (CLOSE_INTERNAL_ERROR, "turn_start_failed"),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
|
||||
use super::{
|
||||
await_owned_turn_begin, await_owned_turn_begin_with_timeout,
|
||||
await_turn_finalization_handle, responses_websocket_turn_start_close,
|
||||
responses_websocket_turn_start_http_status, spawn_bounded_adapter_observation_with_timeout,
|
||||
};
|
||||
use crate::GatewayError;
|
||||
|
||||
#[test]
|
||||
fn admission_timeout_uses_http_429_and_keeps_the_retry_later_close_code() {
|
||||
let error = GatewayError::AdmissionTimeout {
|
||||
trace_id: "turn-admission".to_string(),
|
||||
gate: "gateway_upstream_execution",
|
||||
queue_budget_ms: 25,
|
||||
};
|
||||
|
||||
assert_eq!(responses_websocket_turn_start_http_status(&error), 429);
|
||||
assert_eq!(
|
||||
responses_websocket_turn_start_close(&error),
|
||||
(1013, "gateway_busy")
|
||||
);
|
||||
}
|
||||
|
||||
/// C6 依赖的性质:结算是「等到落地」而不是「排进队列」。
|
||||
///
|
||||
/// 透明重试在这之后立刻按 health / adaptive / pool 状态规划下一个 attempt,
|
||||
/// 所以结算任务必须已经跑完——只把 handle 挂起来是不够的。
|
||||
#[tokio::test]
|
||||
async fn awaiting_a_finalization_handle_runs_the_settlement_to_completion() {
|
||||
let settled = Arc::new(AtomicBool::new(false));
|
||||
let flag = Arc::clone(&settled);
|
||||
let handle = tokio::spawn(async move {
|
||||
tokio::time::sleep(Duration::from_millis(60)).await;
|
||||
flag.store(true, Ordering::SeqCst);
|
||||
});
|
||||
|
||||
assert!(
|
||||
!settled.load(Ordering::SeqCst),
|
||||
"the settlement has not finished yet"
|
||||
);
|
||||
await_turn_finalization_handle(handle).await;
|
||||
assert!(
|
||||
settled.load(Ordering::SeqCst),
|
||||
"the settlement must be complete before the caller proceeds"
|
||||
);
|
||||
}
|
||||
|
||||
/// 顺序型:结算的每一步都要排在规划之前。
|
||||
///
|
||||
/// 用计数器替身重放透明重试的两步——旧 attempt 结算完成写入 1,规划开始时
|
||||
/// 读到的必须已经是 1。旧实现在这里先规划、再把结算排进队列,规划读到的是 0。
|
||||
#[tokio::test]
|
||||
async fn transparent_retry_replans_only_after_the_previous_attempt_is_settled() {
|
||||
let steps = Arc::new(AtomicUsize::new(0));
|
||||
|
||||
// 第一步:结算旧 attempt(等到落地)。
|
||||
let recorder = Arc::clone(&steps);
|
||||
let settlement = tokio::spawn(async move {
|
||||
tokio::time::sleep(Duration::from_millis(40)).await;
|
||||
recorder.store(1, Ordering::SeqCst);
|
||||
});
|
||||
await_turn_finalization_handle(settlement).await;
|
||||
|
||||
// 第二步:规划下一个 attempt,它读到的状态必须是结算之后的。
|
||||
let observed_at_planning = steps.load(Ordering::SeqCst);
|
||||
assert_eq!(
|
||||
observed_at_planning, 1,
|
||||
"planning must observe the state projected by the settled attempt"
|
||||
);
|
||||
}
|
||||
|
||||
/// 结算任务失败(panic / cancel)也必须让调用方继续,不能把 relay loop 卡死。
|
||||
#[tokio::test]
|
||||
async fn a_failed_finalization_task_still_releases_the_caller() {
|
||||
let handle = tokio::spawn(async { panic!("settlement task exploded") });
|
||||
await_turn_finalization_handle(handle).await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn cancelling_the_caller_does_not_cancel_turn_begin_or_drop_an_unowned_result() {
|
||||
struct DropProbe(Arc<AtomicBool>);
|
||||
impl Drop for DropProbe {
|
||||
fn drop(&mut self) {
|
||||
self.0.store(true, Ordering::SeqCst);
|
||||
}
|
||||
}
|
||||
|
||||
let begin_finished = Arc::new(AtomicBool::new(false));
|
||||
let result_dropped = Arc::new(AtomicBool::new(false));
|
||||
let finished = Arc::clone(&begin_finished);
|
||||
let dropped = Arc::clone(&result_dropped);
|
||||
let caller = tokio::spawn(async move {
|
||||
await_owned_turn_begin(
|
||||
async move {
|
||||
tokio::time::sleep(Duration::from_millis(60)).await;
|
||||
finished.store(true, Ordering::SeqCst);
|
||||
Ok(DropProbe(dropped))
|
||||
},
|
||||
Duration::from_secs(1),
|
||||
"turn-begin-cancel".to_string(),
|
||||
)
|
||||
.await
|
||||
});
|
||||
|
||||
tokio::time::sleep(Duration::from_millis(10)).await;
|
||||
caller.abort();
|
||||
let _ = caller.await;
|
||||
tokio::time::sleep(Duration::from_millis(120)).await;
|
||||
|
||||
assert!(
|
||||
begin_finished.load(Ordering::SeqCst),
|
||||
"the owned begin task must outlive its cancelled relay caller"
|
||||
);
|
||||
assert!(
|
||||
result_dropped.load(Ordering::SeqCst),
|
||||
"an undeliverable armed result must be dropped so its cleanup guard runs"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn turn_begin_owner_deadline_drops_stalled_work_and_its_guards() {
|
||||
struct DropProbe(Arc<AtomicBool>);
|
||||
impl Drop for DropProbe {
|
||||
fn drop(&mut self) {
|
||||
self.0.store(true, Ordering::SeqCst);
|
||||
}
|
||||
}
|
||||
|
||||
let dropped = Arc::new(AtomicBool::new(false));
|
||||
let task_dropped = Arc::clone(&dropped);
|
||||
let result: Result<(), GatewayError> = await_owned_turn_begin_with_timeout(
|
||||
async move {
|
||||
let _probe = DropProbe(task_dropped);
|
||||
std::future::pending::<()>().await;
|
||||
Ok(())
|
||||
},
|
||||
Duration::from_millis(20),
|
||||
"turn-begin-deadline".to_string(),
|
||||
)
|
||||
.await;
|
||||
|
||||
assert!(matches!(
|
||||
result,
|
||||
Err(GatewayError::LocalExecutionPlanningTimeout {
|
||||
trace_id,
|
||||
phase: "responses_websocket_turn_begin_owner",
|
||||
timeout_ms: 20,
|
||||
}) if trace_id == "turn-begin-deadline"
|
||||
));
|
||||
assert!(
|
||||
dropped.load(Ordering::SeqCst),
|
||||
"owner timeout must drop the stalled begin future so RAII cleanup runs"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn cancelling_observation_waiter_cannot_bypass_the_owner_timeout() {
|
||||
struct DropProbe(Arc<AtomicBool>);
|
||||
impl Drop for DropProbe {
|
||||
fn drop(&mut self) {
|
||||
self.0.store(true, Ordering::SeqCst);
|
||||
}
|
||||
}
|
||||
|
||||
let dropped = Arc::new(AtomicBool::new(false));
|
||||
let task_dropped = Arc::clone(&dropped);
|
||||
let observation = async move {
|
||||
let _probe = DropProbe(task_dropped);
|
||||
std::future::pending::<()>().await;
|
||||
};
|
||||
let owner =
|
||||
spawn_bounded_adapter_observation_with_timeout(observation, Duration::from_millis(20));
|
||||
let waiter = tokio::spawn(async move {
|
||||
let _ = owner.await;
|
||||
});
|
||||
waiter.abort();
|
||||
let _ = waiter.await;
|
||||
|
||||
tokio::time::timeout(Duration::from_secs(1), async {
|
||||
while !dropped.load(Ordering::SeqCst) {
|
||||
tokio::task::yield_now().await;
|
||||
}
|
||||
})
|
||||
.await
|
||||
.expect("detached observation owner must enforce its own timeout");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,66 @@
|
||||
//! OpenAI Responses WebSocket protocol entry point, session engine, and adapters.
|
||||
//!
|
||||
//! The route is protocol-oriented. `session` bootstraps the authenticated
|
||||
//! connection, `connection` owns the socket FSM, `client` and `quota` own
|
||||
//! protocol/retry policy, and `lifecycle`/`turn` bridge each turn into the
|
||||
//! existing usage and audit runtime. Adapters contain only provider-specific
|
||||
//! connection and metadata behavior.
|
||||
|
||||
mod adapter;
|
||||
mod adapters;
|
||||
mod admission;
|
||||
mod binding;
|
||||
mod client;
|
||||
mod connection;
|
||||
mod control;
|
||||
mod frame;
|
||||
mod lifecycle;
|
||||
mod observation;
|
||||
mod ownership;
|
||||
mod quota;
|
||||
mod redaction;
|
||||
mod relay_policy;
|
||||
mod request;
|
||||
mod session;
|
||||
mod settlement;
|
||||
mod state;
|
||||
mod turn;
|
||||
mod turn_state;
|
||||
mod upstream;
|
||||
|
||||
use std::net::SocketAddr;
|
||||
|
||||
use axum::body::Body;
|
||||
use axum::extract::ws::WebSocketUpgrade;
|
||||
use axum::extract::{ConnectInfo, State};
|
||||
use axum::http::{HeaderMap, Response, Uri};
|
||||
|
||||
use crate::handlers::proxy::websocket::ingress::{
|
||||
upgrade_authenticated_ai_websocket, WebSocketIngressSpec,
|
||||
};
|
||||
use crate::handlers::proxy::websocket::session::RESPONSES_WEBSOCKET_SESSION_LIMITS;
|
||||
use crate::{AppState, GatewayError};
|
||||
|
||||
pub(crate) async fn responses_websocket(
|
||||
State(state): State<AppState>,
|
||||
ConnectInfo(remote_addr): ConnectInfo<SocketAddr>,
|
||||
ws: WebSocketUpgrade,
|
||||
headers: HeaderMap,
|
||||
uri: Uri,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
upgrade_authenticated_ai_websocket(
|
||||
state,
|
||||
remote_addr,
|
||||
ws,
|
||||
headers,
|
||||
uri,
|
||||
RESPONSES_WEBSOCKET_SESSION_LIMITS,
|
||||
RESPONSES_WEBSOCKET_INGRESS_SPEC,
|
||||
session::run_responses_websocket,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
const RESPONSES_WEBSOCKET_INGRESS_SPEC: WebSocketIngressSpec = WebSocketIngressSpec {
|
||||
route_unavailable_message: "WebSocket route is unavailable",
|
||||
};
|
||||
@@ -0,0 +1,220 @@
|
||||
//! Responses WebSocket 的终态观测入口。
|
||||
//!
|
||||
//! 这条传输收到的本来就是结构化的 Responses 协议事件。之前为了复用面向 SSE 的
|
||||
//! `push_line`,观测路径要先把每个事件序列化成 `data: {json}\n\n`,解析器再把它
|
||||
//! 解码回 `Value`——一次纯粹的往返,而且这个「伪 SSE」形状是随手拼的,一旦
|
||||
//! 上游事件里出现需要转义的内容,或者以后有人给拼装函数加了换行/分块逻辑,
|
||||
//! 观测结果就会和真实事件悄悄分叉。
|
||||
//!
|
||||
//! 现在观测走 [`StreamingStandardTerminalObserver::push_event`],直接吃
|
||||
//! `frame.protocol_events()` 借出的事件,不再序列化、不再解码。
|
||||
//!
|
||||
//! **body capture 不走这条路,仍然保持 SSE 形状**(`data: {json}\n\n`):
|
||||
//! `aether_usage_runtime::report` 用 `line.strip_prefix("data:")` 解析被捕获的
|
||||
//! body 来判定 `StreamCapturedTerminalState`,而它是 `stream_report_represents_failure`
|
||||
//! 的一个 OR 项。把捕获内容换成结构化 JSON 会让终态判定恒为 Missing。
|
||||
//! 也就是说这一层只换「观测」,不换「捕获」——见
|
||||
//! [`super::turn::ResponsesProviderAttempt::capture_client_frame`] 一侧仍在用
|
||||
//! SSE 编码。
|
||||
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::ai_serving::api::StreamingStandardTerminalObserver;
|
||||
use aether_contracts::ExecutionStreamTerminalSummary;
|
||||
|
||||
/// 包一层 [`StreamingStandardTerminalObserver`],只暴露结构化入口。
|
||||
///
|
||||
/// 存在的意义是让「WS 不再拼 SSE」成为类型层面的事实:这里没有任何接受字节的
|
||||
/// 方法,所以不可能有人不小心把观测路径改回 `push_line`。
|
||||
#[derive(Default)]
|
||||
pub(super) struct ResponsesStructuredTerminalObserver {
|
||||
inner: StreamingStandardTerminalObserver,
|
||||
}
|
||||
|
||||
impl ResponsesStructuredTerminalObserver {
|
||||
/// 观测一帧里的全部协议事件。
|
||||
///
|
||||
/// 第一个被拒绝的事件就停止推进并把摘要标成 parser_error:解析器的状态机是
|
||||
/// 有顺序的,跳过一个事件继续喂后面的只会得到更没意义的摘要。
|
||||
pub(super) fn observe_events(&mut self, report_context: &Value, events: &[&Value]) {
|
||||
for event in events
|
||||
.iter()
|
||||
.copied()
|
||||
.filter(|event| event_is_relevant_to_terminal_observation(event))
|
||||
{
|
||||
if let Err(error) = self.inner.push_event(report_context, event) {
|
||||
self.inner.disable_with_error(error.to_string());
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn disable_with_error(&mut self, parser_error: impl Into<String>) {
|
||||
self.inner.disable_with_error(parser_error);
|
||||
}
|
||||
|
||||
pub(super) fn finish(&mut self, report_context: &Value) -> ExecutionStreamTerminalSummary {
|
||||
match self.inner.finish(report_context) {
|
||||
Ok(Some(summary)) => summary,
|
||||
Ok(None) => ExecutionStreamTerminalSummary::default(),
|
||||
Err(error) => {
|
||||
self.inner.disable_with_error(error.to_string());
|
||||
self.inner.latest_summary().cloned().unwrap_or_default()
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// The WebSocket relay is not a Responses schema gateway. It forwards all
|
||||
/// events opaquely, while this observer consumes only identity/terminal
|
||||
/// snapshots needed for usage and settlement. In particular, a future
|
||||
/// `response.*` delta must not become an observation failure merely because
|
||||
/// Aether's canonical streaming parser does not know it yet.
|
||||
fn event_is_relevant_to_terminal_observation(event: &Value) -> bool {
|
||||
matches!(
|
||||
event.get("type").and_then(Value::as_str),
|
||||
Some(
|
||||
"response.created"
|
||||
| "response.in_progress"
|
||||
| "response.queued"
|
||||
| "response.completed"
|
||||
| "response.done"
|
||||
| "response.failed"
|
||||
| "response.incomplete"
|
||||
| "response.cancelled"
|
||||
| "error"
|
||||
)
|
||||
)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::ResponsesStructuredTerminalObserver;
|
||||
|
||||
fn report_context() -> serde_json::Value {
|
||||
json!({
|
||||
"provider_api_format": "openai:responses",
|
||||
"client_api_format": "openai:responses",
|
||||
})
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn structured_events_reach_the_terminal_summary_without_sse_text() {
|
||||
let context = report_context();
|
||||
let created = json!({"type": "response.created", "response": {"id": "resp_ws", "model": "gpt-5-codex"}});
|
||||
let completed = json!({
|
||||
"type": "response.completed",
|
||||
"response": {
|
||||
"id": "resp_ws",
|
||||
"model": "gpt-5-codex",
|
||||
"status": "completed",
|
||||
"usage": {"input_tokens": 9, "output_tokens": 4, "total_tokens": 13},
|
||||
},
|
||||
});
|
||||
|
||||
let mut observer = ResponsesStructuredTerminalObserver::default();
|
||||
observer.observe_events(&context, &[&created, &completed]);
|
||||
let summary = observer.finish(&context);
|
||||
|
||||
assert!(summary.observed_finish);
|
||||
assert_eq!(summary.response_id.as_deref(), Some("resp_ws"));
|
||||
let usage = summary
|
||||
.standardized_usage
|
||||
.as_ref()
|
||||
.expect("a completed response carries usage");
|
||||
assert_eq!(usage.input_tokens, 9);
|
||||
assert_eq!(usage.output_tokens, 4);
|
||||
assert!(summary.parser_error.is_none());
|
||||
}
|
||||
|
||||
/// 批量帧里的多个事件按顺序喂入,usage 不能因为批量而丢失。
|
||||
#[test]
|
||||
fn a_batched_frame_keeps_the_usage_of_its_last_event() {
|
||||
let context = report_context();
|
||||
let events = [
|
||||
json!({"type": "response.created", "response": {"id": "resp_ws", "model": "m"}}),
|
||||
json!({
|
||||
"type": "response.output_text.delta",
|
||||
"item_id": "msg",
|
||||
"output_index": 0,
|
||||
"content_index": 0,
|
||||
"delta": "hi",
|
||||
}),
|
||||
json!({
|
||||
"type": "response.completed",
|
||||
"response": {
|
||||
"id": "resp_ws",
|
||||
"model": "m",
|
||||
"status": "completed",
|
||||
"usage": {"input_tokens": 3, "output_tokens": 1, "total_tokens": 4},
|
||||
},
|
||||
}),
|
||||
];
|
||||
let borrowed: Vec<&serde_json::Value> = events.iter().collect();
|
||||
|
||||
let mut observer = ResponsesStructuredTerminalObserver::default();
|
||||
observer.observe_events(&context, &borrowed);
|
||||
let summary = observer.finish(&context);
|
||||
|
||||
let usage = summary
|
||||
.standardized_usage
|
||||
.as_ref()
|
||||
.expect("usage survives batching");
|
||||
assert_eq!(usage.input_tokens, 3);
|
||||
assert_eq!(usage.output_tokens, 1);
|
||||
assert_eq!(usage.dimensions.get("total_tokens"), Some(&json!(4)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn future_and_provider_private_events_are_ignored_only_by_the_side_observer() {
|
||||
let context = report_context();
|
||||
let private = json!({
|
||||
"type": "codex.response.metadata",
|
||||
"private_future_field": {"shape": "unknown"},
|
||||
});
|
||||
let future = json!({
|
||||
"type": "response.future_capability.delta",
|
||||
"future_capability": {"nested": [1, 2, 3]},
|
||||
});
|
||||
let completed = json!({
|
||||
"type": "response.completed",
|
||||
"response": {
|
||||
"id": "resp_future",
|
||||
"model": "future-model",
|
||||
"status": "completed",
|
||||
"future_response_field": {"also": "unknown"},
|
||||
"usage": {"input_tokens": 5, "output_tokens": 2, "total_tokens": 7},
|
||||
},
|
||||
});
|
||||
|
||||
let mut observer = ResponsesStructuredTerminalObserver::default();
|
||||
observer.observe_events(&context, &[&private, &future, &completed]);
|
||||
let summary = observer.finish(&context);
|
||||
|
||||
assert!(summary.observed_finish);
|
||||
assert_eq!(summary.response_id.as_deref(), Some("resp_future"));
|
||||
assert_eq!(summary.unknown_event_count, 0);
|
||||
assert_eq!(
|
||||
summary
|
||||
.standardized_usage
|
||||
.as_ref()
|
||||
.map(|usage| (usage.input_tokens, usage.output_tokens)),
|
||||
Some((5, 2))
|
||||
);
|
||||
assert!(summary.parser_error.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn a_disabled_observer_reports_the_parser_error() {
|
||||
let context = report_context();
|
||||
let mut observer = ResponsesStructuredTerminalObserver::default();
|
||||
observer.disable_with_error("upstream event was not valid JSON");
|
||||
let summary = observer.finish(&context);
|
||||
assert_eq!(
|
||||
summary.parser_error.as_deref(),
|
||||
Some("upstream event was not valid JSON")
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,245 @@
|
||||
//! Cancellation-safe ownership handoff for WebSocket planning leases.
|
||||
//!
|
||||
//! The relay races every turn against connection and response deadlines. A
|
||||
//! planner future therefore cannot directly own a distributed pool-key lease:
|
||||
//! losing the race would drop the future between scheduler selection and turn
|
||||
//! startup, leaving that key unavailable until the lease TTL elapsed.
|
||||
|
||||
use std::collections::BTreeSet;
|
||||
use std::time::Duration;
|
||||
|
||||
use serde_json::Value;
|
||||
use tokio::task::JoinHandle;
|
||||
|
||||
use super::lifecycle::{begin_responses_websocket_turn, ActiveProviderAttempt};
|
||||
use crate::ai_serving::{
|
||||
maybe_build_responses_websocket_decision, AiExecutionDecision, GatewayAuthApiKeySnapshot,
|
||||
ResponsesWebSocketDecision, ResponsesWebSocketPinnedCandidate,
|
||||
};
|
||||
use crate::control::GatewayControlDecision;
|
||||
use crate::orchestration::release_pool_key_lease_from_report_context;
|
||||
use crate::{AppState, GatewayError};
|
||||
|
||||
/// Owns a selected pool-key lease until the attempt lifecycle has taken over
|
||||
/// the decision report context.
|
||||
pub(super) struct PlannedPoolKeyLeaseGuard {
|
||||
state: AppState,
|
||||
report_context: Option<Value>,
|
||||
}
|
||||
|
||||
/// Planner output coupled to both its request parts and lease guard.
|
||||
pub(super) struct OwnedResponsesWebSocketDecision {
|
||||
pub(super) planned: ResponsesWebSocketDecision,
|
||||
pub(super) planning_parts: http::request::Parts,
|
||||
pub(super) planned_lease: PlannedPoolKeyLeaseGuard,
|
||||
}
|
||||
|
||||
/// Runs planning in an owner task. Dropping the caller's waiter detaches this
|
||||
/// task; an unobserved successful output drops its guard and releases the
|
||||
/// selected pool-key lease.
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub(super) fn spawn_owned_responses_websocket_plan(
|
||||
state: AppState,
|
||||
parts: http::request::Parts,
|
||||
trace_id: String,
|
||||
control_decision: GatewayControlDecision,
|
||||
auth_snapshot: Option<GatewayAuthApiKeySnapshot>,
|
||||
client_event: Value,
|
||||
excluded_key_ids: Option<BTreeSet<String>>,
|
||||
excluded_codex_account_ids: Option<BTreeSet<String>>,
|
||||
pinned_candidate: Option<ResponsesWebSocketPinnedCandidate>,
|
||||
) -> JoinHandle<Result<Option<OwnedResponsesWebSocketDecision>, GatewayError>> {
|
||||
let owner_timeout = state
|
||||
.frontdoor_runtime_guards
|
||||
.local_execution_planning_timeout;
|
||||
tokio::spawn(async move {
|
||||
let planned = await_owned_planning_deadline(
|
||||
maybe_build_responses_websocket_decision(
|
||||
&state,
|
||||
&parts,
|
||||
&trace_id,
|
||||
&control_decision,
|
||||
auth_snapshot.as_ref(),
|
||||
&client_event,
|
||||
excluded_key_ids.as_ref(),
|
||||
excluded_codex_account_ids.as_ref(),
|
||||
pinned_candidate.as_ref(),
|
||||
),
|
||||
owner_timeout,
|
||||
)
|
||||
.await
|
||||
.map_err(|_| GatewayError::LocalExecutionPlanningTimeout {
|
||||
trace_id: trace_id.clone(),
|
||||
phase: "responses_websocket_plan_owner",
|
||||
timeout_ms: owner_timeout.as_millis() as u64,
|
||||
})??;
|
||||
|
||||
Ok(planned.map(|planned| {
|
||||
let planned_lease =
|
||||
PlannedPoolKeyLeaseGuard::new(&state, planned.execution.report_context.as_ref());
|
||||
OwnedResponsesWebSocketDecision {
|
||||
planned,
|
||||
planning_parts: parts,
|
||||
planned_lease,
|
||||
}
|
||||
}))
|
||||
})
|
||||
}
|
||||
|
||||
async fn await_owned_planning_deadline<F, T>(
|
||||
planning: F,
|
||||
deadline: Duration,
|
||||
) -> Result<T, tokio::time::error::Elapsed>
|
||||
where
|
||||
F: std::future::Future<Output = T>,
|
||||
{
|
||||
tokio::time::timeout(deadline, planning).await
|
||||
}
|
||||
|
||||
pub(super) async fn await_owned_responses_websocket_plan(
|
||||
handle: JoinHandle<Result<Option<OwnedResponsesWebSocketDecision>, GatewayError>>,
|
||||
) -> Result<Option<OwnedResponsesWebSocketDecision>, GatewayError> {
|
||||
handle.await.map_err(|error| {
|
||||
GatewayError::Internal(format!(
|
||||
"Responses WebSocket planning task failed before ownership transfer: {error}"
|
||||
))
|
||||
})?
|
||||
}
|
||||
|
||||
impl PlannedPoolKeyLeaseGuard {
|
||||
fn new(state: &AppState, report_context: Option<&Value>) -> Self {
|
||||
Self {
|
||||
state: state.clone(),
|
||||
report_context: report_context.cloned(),
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) async fn release(mut self) {
|
||||
release_pool_key_lease_from_report_context(&self.state, self.report_context.as_ref()).await;
|
||||
self.report_context = None;
|
||||
}
|
||||
|
||||
fn disarm(&mut self) {
|
||||
self.report_context = None;
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for PlannedPoolKeyLeaseGuard {
|
||||
fn drop(&mut self) {
|
||||
let Some(report_context) = self.report_context.take() else {
|
||||
return;
|
||||
};
|
||||
let state = self.state.clone();
|
||||
if let Ok(handle) = tokio::runtime::Handle::try_current() {
|
||||
handle.spawn(async move {
|
||||
release_pool_key_lease_from_report_context(&state, Some(&report_context)).await;
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Keeps the planning guard in the same detached owner task as lifecycle
|
||||
/// startup. If the relay loses a deadline race while awaiting startup, the
|
||||
/// task completes the handoff (or releases the lease on failure) without a
|
||||
/// cancellation gap.
|
||||
pub(super) async fn begin_responses_websocket_turn_with_planned_lease(
|
||||
state: &AppState,
|
||||
trace_id: &str,
|
||||
parts: http::request::Parts,
|
||||
control_decision: &GatewayControlDecision,
|
||||
decision: AiExecutionDecision,
|
||||
client_event: &Value,
|
||||
mut planned_lease: PlannedPoolKeyLeaseGuard,
|
||||
) -> Result<ActiveProviderAttempt, GatewayError> {
|
||||
let state = state.clone();
|
||||
let trace_id = trace_id.to_string();
|
||||
let control_decision = control_decision.clone();
|
||||
let client_event = client_event.clone();
|
||||
tokio::spawn(async move {
|
||||
let turn = begin_responses_websocket_turn(
|
||||
&state,
|
||||
&trace_id,
|
||||
parts,
|
||||
&control_decision,
|
||||
decision,
|
||||
&client_event,
|
||||
)
|
||||
.await?;
|
||||
// ActiveProviderAttempt now owns the report context containing the
|
||||
// lease. No await occurs between that handoff and disarming the guard.
|
||||
planned_lease.disarm();
|
||||
Ok(turn)
|
||||
})
|
||||
.await
|
||||
.map_err(|error| {
|
||||
GatewayError::Internal(format!(
|
||||
"Responses WebSocket guarded turn startup task failed: {error}"
|
||||
))
|
||||
})?
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
|
||||
struct DropProbe(Arc<AtomicUsize>);
|
||||
|
||||
impl Drop for DropProbe {
|
||||
fn drop(&mut self) {
|
||||
self.0.fetch_add(1, Ordering::SeqCst);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn dropping_a_planning_waiter_detaches_the_owner_and_drops_its_output() {
|
||||
let started = Arc::new(tokio::sync::Notify::new());
|
||||
let release = Arc::new(tokio::sync::Notify::new());
|
||||
let dropped = Arc::new(AtomicUsize::new(0));
|
||||
let task_started = Arc::clone(&started);
|
||||
let task_release = Arc::clone(&release);
|
||||
let task_dropped = Arc::clone(&dropped);
|
||||
let owner = tokio::spawn(async move {
|
||||
task_started.notify_one();
|
||||
task_release.notified().await;
|
||||
DropProbe(task_dropped)
|
||||
});
|
||||
started.notified().await;
|
||||
|
||||
let waiter = tokio::spawn(async move {
|
||||
let _ = owner.await;
|
||||
});
|
||||
waiter.abort();
|
||||
let _ = waiter.await;
|
||||
release.notify_one();
|
||||
|
||||
tokio::time::timeout(Duration::from_secs(1), async {
|
||||
while dropped.load(Ordering::SeqCst) == 0 {
|
||||
tokio::task::yield_now().await;
|
||||
}
|
||||
})
|
||||
.await
|
||||
.expect("detached owner output should be dropped after it finishes");
|
||||
assert_eq!(dropped.load(Ordering::SeqCst), 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn planning_owner_deadline_drops_stalled_work_and_its_guards() {
|
||||
let dropped = Arc::new(AtomicUsize::new(0));
|
||||
let task_dropped = Arc::clone(&dropped);
|
||||
let planning = async move {
|
||||
let _probe = DropProbe(task_dropped);
|
||||
std::future::pending::<()>().await;
|
||||
};
|
||||
|
||||
let result =
|
||||
super::await_owned_planning_deadline(planning, Duration::from_millis(20)).await;
|
||||
|
||||
assert!(
|
||||
result.is_err(),
|
||||
"stalled planning must hit its owner deadline"
|
||||
);
|
||||
assert_eq!(dropped.load(Ordering::SeqCst), 1);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,377 @@
|
||||
//! Quota exhaustion, replay safety, and upstream replacement policy.
|
||||
|
||||
use serde_json::Value;
|
||||
use uuid::Uuid;
|
||||
use wreq::ws::message::Message as WreqWsMessage;
|
||||
|
||||
use super::adapter::{
|
||||
resolve_responses_websocket_adapter, ResponsesWebSocketDrainDirective,
|
||||
ResponsesWebSocketRebindSafety,
|
||||
};
|
||||
use super::lifecycle::{queue_turn_finalization, PreviousAttemptSettled};
|
||||
use super::ownership::{
|
||||
await_owned_responses_websocket_plan, begin_responses_websocket_turn_with_planned_lease,
|
||||
spawn_owned_responses_websocket_plan, OwnedResponsesWebSocketDecision,
|
||||
};
|
||||
use super::request::{build_planning_parts, planned_response_create_event};
|
||||
use super::state::BoundResponsesConnection;
|
||||
use super::turn::{prepare_responses_websocket_turn_decision, ResponsesWebSocketTurnOutcome};
|
||||
use super::upstream::{bind_responses_upstream, close_bound_upstream};
|
||||
use crate::clock::current_unix_secs;
|
||||
use crate::handlers::proxy::websocket::ingress::WebSocketRequestContext;
|
||||
use crate::handlers::proxy::websocket::session::WEBSOCKET_LOG_TRANSPORT;
|
||||
use crate::handlers::proxy::websocket::transport::close_upstream_socket;
|
||||
use crate::AppState;
|
||||
|
||||
const LOG_TARGET: &str = "aether_gateway::handlers::proxy::responses_ws";
|
||||
|
||||
macro_rules! debug {
|
||||
($($arg:tt)*) => {
|
||||
tracing::debug!(target: LOG_TARGET, $($arg)*)
|
||||
};
|
||||
}
|
||||
|
||||
macro_rules! warn {
|
||||
($($arg:tt)*) => {
|
||||
tracing::warn!(target: LOG_TARGET, $($arg)*)
|
||||
};
|
||||
}
|
||||
|
||||
pub(super) async fn detach_exhausted_upstream(
|
||||
bound: &mut BoundResponsesConnection,
|
||||
directive: ResponsesWebSocketDrainDirective,
|
||||
trace_id: &str,
|
||||
) {
|
||||
let exclusion = record_exhausted_bound_key(bound, directive.retry_exclusion_until_unix_secs);
|
||||
close_bound_upstream(bound).await;
|
||||
// 调用方必须先结束当前 logical turn 再 detach:拆掉上游后 attempt 已经不可能
|
||||
// 收到终态,留着它只会等 deadline 或 drop guard 兜底。
|
||||
debug_assert!(
|
||||
!bound.turn_state.response_in_flight(),
|
||||
"an exhausted upstream must be detached after its logical turn ended"
|
||||
);
|
||||
bound.pending_adapter_drain = None;
|
||||
let now_unix_secs = current_unix_secs();
|
||||
debug!(
|
||||
event_name = "responses_websocket_upstream_detached",
|
||||
log_type = "event",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %trace_id,
|
||||
reason = directive.error_code,
|
||||
exhausted_key_id = ?exclusion.as_ref().map(|(key_id, _)| key_id),
|
||||
retry_exclusion_until_unix_secs = ?exclusion.as_ref().map(|(_, until)| until),
|
||||
exhausted_exclusion_count = bound.exhausted_exclusions.len(now_unix_secs),
|
||||
"gateway detached an exhausted Responses WebSocket upstream while preserving the client socket"
|
||||
);
|
||||
}
|
||||
|
||||
pub(super) fn record_exhausted_bound_key(
|
||||
bound: &mut BoundResponsesConnection,
|
||||
reset_at_unix_secs: Option<u64>,
|
||||
) -> Option<(String, u64)> {
|
||||
let key_id = bound
|
||||
.decision_template
|
||||
.key_id
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|key_id| !key_id.is_empty())?
|
||||
.to_string();
|
||||
let provider_account_id = bound
|
||||
.adapter
|
||||
.exhaustion_exclusion_identity(&bound.decision_template)
|
||||
.and_then(|identity| identity.account_id);
|
||||
let exclusion_until = bound.exhausted_exclusions.exclude(
|
||||
key_id.clone(),
|
||||
provider_account_id,
|
||||
reset_at_unix_secs,
|
||||
current_unix_secs(),
|
||||
);
|
||||
Some((key_id, exclusion_until))
|
||||
}
|
||||
|
||||
/// 为同一个 logical turn 规划并绑定下一个 attempt。
|
||||
///
|
||||
/// `_previous_settled` 不被使用,它只是把「上一个 attempt 已经结算完毕」这个
|
||||
/// 前置条件写进签名:规划要读 health / adaptive / pool 状态,而这些是上一个
|
||||
/// attempt 结算时才投射的;它的 pool key lease 也要先释放,否则替代 key 的挑选
|
||||
/// 会看到一把仍被占用的 key。
|
||||
pub(super) async fn retry_active_turn_after_quota_exhaustion(
|
||||
bound: &mut BoundResponsesConnection,
|
||||
state: &AppState,
|
||||
context: &WebSocketRequestContext,
|
||||
_previous_settled: PreviousAttemptSettled,
|
||||
) -> bool {
|
||||
let Some(active) = bound.turn_state.logical_mut() else {
|
||||
return false;
|
||||
};
|
||||
if let Some(reason) = active.quota_retry_block_reason() {
|
||||
debug!(
|
||||
event_name = "responses_websocket_quota_retry_skipped",
|
||||
log_type = "event",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %context.trace_id,
|
||||
turn_index = active.turn_index,
|
||||
logical_turn_id = %active.logical_turn_id,
|
||||
turn_attempt = active.turn_attempt,
|
||||
reason,
|
||||
"gateway will not transparently replay an unsafe Responses WebSocket turn"
|
||||
);
|
||||
return false;
|
||||
}
|
||||
active.retry_attempted = true;
|
||||
active.turn_attempt = active.turn_attempt.saturating_add(1);
|
||||
let client_event = active.client_event.clone();
|
||||
let Some(turn_control) = active.turn_control.clone() else {
|
||||
warn!(
|
||||
event_name = "responses_websocket_quota_retry_control_missing",
|
||||
log_type = "ops",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %context.trace_id,
|
||||
"gateway refused to retry a WebSocket turn without its live authorization snapshot"
|
||||
);
|
||||
return false;
|
||||
};
|
||||
let turn_index = active.turn_index;
|
||||
let logical_turn_id = active.logical_turn_id.clone();
|
||||
let turn_attempt = active.turn_attempt;
|
||||
|
||||
let retry_exclusion_until_unix_secs = bound
|
||||
.pending_adapter_drain
|
||||
.and_then(|directive| directive.retry_exclusion_until_unix_secs);
|
||||
let exhausted_key = record_exhausted_bound_key(bound, retry_exclusion_until_unix_secs);
|
||||
let exhausted_key_id = exhausted_key.as_ref().map(|(key_id, _)| key_id.clone());
|
||||
|
||||
let planning_parts = build_planning_parts(context);
|
||||
let turn_request_id = Uuid::new_v4().to_string();
|
||||
let now_unix_secs = current_unix_secs();
|
||||
let excluded_key_ids = bound.exhausted_exclusions.key_ids(now_unix_secs);
|
||||
let excluded_codex_account_ids = bound.exhausted_exclusions.codex_account_ids(now_unix_secs);
|
||||
let excluded_key_ids = (!excluded_key_ids.is_empty()).then_some(excluded_key_ids);
|
||||
let excluded_codex_account_ids =
|
||||
(!excluded_codex_account_ids.is_empty()).then_some(excluded_codex_account_ids);
|
||||
let planned = match await_owned_responses_websocket_plan(spawn_owned_responses_websocket_plan(
|
||||
state.clone(),
|
||||
planning_parts,
|
||||
turn_request_id.clone(),
|
||||
turn_control.decision.clone(),
|
||||
turn_control.auth_snapshot.clone(),
|
||||
client_event.clone(),
|
||||
excluded_key_ids,
|
||||
excluded_codex_account_ids,
|
||||
None,
|
||||
))
|
||||
.await
|
||||
{
|
||||
Ok(Some(decision)) => decision,
|
||||
Ok(None) => {
|
||||
warn!(
|
||||
event_name = "responses_websocket_quota_retry_provider_unavailable",
|
||||
log_type = "ops",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %context.trace_id,
|
||||
exhausted_key_id = ?exhausted_key_id,
|
||||
"gateway could not find an alternate Responses WebSocket provider after quota exhaustion"
|
||||
);
|
||||
return false;
|
||||
}
|
||||
Err(error) => {
|
||||
warn!(
|
||||
event_name = "responses_websocket_quota_retry_planning_failed",
|
||||
log_type = "ops",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %context.trace_id,
|
||||
exhausted_key_id = ?exhausted_key_id,
|
||||
error = ?error,
|
||||
"gateway could not plan an alternate Responses WebSocket provider after quota exhaustion"
|
||||
);
|
||||
return false;
|
||||
}
|
||||
};
|
||||
let OwnedResponsesWebSocketDecision {
|
||||
planned,
|
||||
planning_parts,
|
||||
planned_lease,
|
||||
} = planned;
|
||||
let adapter = resolve_responses_websocket_adapter(planned.adapter);
|
||||
let normalization = planned.normalization;
|
||||
let decision = planned.execution;
|
||||
if exhausted_key_id.as_deref() == decision.key_id.as_deref() {
|
||||
planned_lease.release().await;
|
||||
warn!(
|
||||
event_name = "responses_websocket_quota_retry_selected_exhausted_key",
|
||||
log_type = "ops",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %context.trace_id,
|
||||
key_id = ?decision.key_id,
|
||||
"gateway rejected an alternate Responses WebSocket plan that reused the exhausted key"
|
||||
);
|
||||
return false;
|
||||
}
|
||||
let provider_event = match planned_response_create_event(&decision, &client_event).and_then(
|
||||
|event| {
|
||||
serde_json::from_str::<Value>(&event)
|
||||
.map_err(|_| "response_create_serialization_failed")
|
||||
},
|
||||
) {
|
||||
Ok(event) => event,
|
||||
Err(code) => {
|
||||
planned_lease.release().await;
|
||||
warn!(
|
||||
event_name = "responses_websocket_quota_retry_normalization_failed",
|
||||
log_type = "ops",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %context.trace_id,
|
||||
error_code = code,
|
||||
"gateway could not rebuild a Responses response.create for transparent quota retry"
|
||||
);
|
||||
return false;
|
||||
}
|
||||
};
|
||||
let turn_decision = prepare_responses_websocket_turn_decision(
|
||||
&decision,
|
||||
turn_request_id,
|
||||
true,
|
||||
&client_event,
|
||||
&provider_event,
|
||||
&context.trace_id,
|
||||
turn_index,
|
||||
&logical_turn_id,
|
||||
turn_attempt,
|
||||
);
|
||||
let mut turn = match begin_responses_websocket_turn_with_planned_lease(
|
||||
state,
|
||||
&context.trace_id,
|
||||
planning_parts,
|
||||
&turn_control.decision,
|
||||
turn_decision,
|
||||
&client_event,
|
||||
planned_lease,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(turn) => turn,
|
||||
Err(error) => {
|
||||
warn!(
|
||||
event_name = "responses_websocket_quota_retry_reporting_unavailable",
|
||||
log_type = "ops",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %context.trace_id,
|
||||
error = ?error,
|
||||
"gateway could not start usage and audit tracking for transparent quota retry"
|
||||
);
|
||||
return false;
|
||||
}
|
||||
};
|
||||
let mut replacement = match bind_responses_upstream(
|
||||
&decision,
|
||||
normalization,
|
||||
&client_event,
|
||||
adapter,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(connection) => connection,
|
||||
Err(code) => {
|
||||
queue_turn_finalization(
|
||||
bound,
|
||||
state,
|
||||
turn,
|
||||
ResponsesWebSocketTurnOutcome::upstream_connect_failed(code),
|
||||
)
|
||||
.await;
|
||||
warn!(
|
||||
event_name = "responses_websocket_quota_retry_rebind_failed",
|
||||
log_type = "ops",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %context.trace_id,
|
||||
error_code = code,
|
||||
"gateway could not bind an alternate Responses WebSocket provider after quota exhaustion"
|
||||
);
|
||||
return false;
|
||||
}
|
||||
};
|
||||
|
||||
turn.mark_upstream_request_sent();
|
||||
turn.set_provider_response_headers(replacement.upstream_response_headers.clone());
|
||||
let replacement_upstream = replacement
|
||||
.upstream
|
||||
.take()
|
||||
.expect("newly bound Responses upstream should be present");
|
||||
if let Some(mut previous_upstream) = bound.upstream.replace(replacement_upstream) {
|
||||
close_upstream_socket(&mut previous_upstream, None).await;
|
||||
}
|
||||
let previous_key_id = bound.decision_template.key_id.clone();
|
||||
bound.adapter = replacement.adapter;
|
||||
bound.client_model = replacement.client_model;
|
||||
bound.provider_model = replacement.provider_model;
|
||||
bound.decision_template = replacement.decision_template;
|
||||
bound.body_normalization = replacement.body_normalization;
|
||||
bound.binding_identity = replacement.binding_identity;
|
||||
// 同一个 logical turn 的下一个 attempt 就位。状态不符时把 attempt 交回
|
||||
// drop guard 结算并让调用方走「透明重试失败」分支,不静默丢弃一条已经写了
|
||||
// pending usage 行、占着 candidate 和 pool key lease 的 attempt。
|
||||
if let Err(orphan) = bound.turn_state.resume(turn) {
|
||||
drop(orphan);
|
||||
return false;
|
||||
}
|
||||
bound.upstream_response_headers = replacement.upstream_response_headers;
|
||||
bound.pending_adapter_drain = None;
|
||||
debug!(
|
||||
event_name = "responses_websocket_quota_retry_rebound",
|
||||
log_type = "event",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %context.trace_id,
|
||||
turn_index,
|
||||
logical_turn_id = %logical_turn_id,
|
||||
turn_attempt,
|
||||
previous_key_id = ?previous_key_id,
|
||||
key_id = ?bound.decision_template.key_id,
|
||||
"gateway transparently rebound a Responses WebSocket turn after quota exhaustion"
|
||||
);
|
||||
true
|
||||
}
|
||||
|
||||
pub(super) fn is_usage_limit_error_event(event: &Value) -> bool {
|
||||
let is_error = |value: &Value| {
|
||||
value.get("type").and_then(Value::as_str) == Some("error")
|
||||
&& value.pointer("/error/type").and_then(Value::as_str) == Some("usage_limit_reached")
|
||||
};
|
||||
is_error(event)
|
||||
|| event
|
||||
.get("chunks")
|
||||
.and_then(Value::as_array)
|
||||
.is_some_and(|chunks| chunks.iter().any(is_error))
|
||||
}
|
||||
|
||||
pub(super) fn observe_active_response_rebind_safety(
|
||||
bound: &mut BoundResponsesConnection,
|
||||
event: &Value,
|
||||
) {
|
||||
let ResponsesWebSocketRebindSafety::Unsafe { reason } =
|
||||
bound.adapter.rebind_safety_for_upstream_event(event)
|
||||
else {
|
||||
return;
|
||||
};
|
||||
if let Some(active) = bound.turn_state.logical_mut() {
|
||||
active.mark_retry_unsafe(reason);
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn mark_active_response_retry_unsafe(
|
||||
bound: &mut BoundResponsesConnection,
|
||||
reason: &'static str,
|
||||
) {
|
||||
if let Some(active) = bound.turn_state.logical_mut() {
|
||||
active.mark_retry_unsafe(reason);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,850 @@
|
||||
//! Responses WebSocket 两侧的 PII 脱敏:请求侧 mask + 响应侧 restore。
|
||||
//!
|
||||
//! HTTP 路径在前门建 `RedactionSessionSlot` 并塞进 `parts.extensions`,planner
|
||||
//! 只有拿到这个 slot 才会脱敏。WS 的 planning Parts 是合成的:四个规划入口
|
||||
//! (首轮、换模型 re-plan、独立轮、配额透明重试)靠 `build_planning_parts` 注入
|
||||
//! slot 就能复用 planner 的脱敏;但复用已绑定 upstream 的 continuation 根本不进
|
||||
//! planner,必须在这里先把客户端事件脱敏,再交给协议归一化、上游发送和审计。
|
||||
//!
|
||||
//! 因此约定:**进入任何下游用途之前,客户端 `response.create` 只在这里脱敏一次**,
|
||||
//! 之后所有路径都只看脱敏后的事件。
|
||||
//!
|
||||
//! # 响应侧
|
||||
//!
|
||||
//! 只 mask 不 restore 是半个实现:HTTP 在把响应交给客户端之前会把占位符换回真实值
|
||||
//! (`privacy::restore_sync_response_body` / `privacy::StreamingResponseRestorer`),
|
||||
//! WS 少了这一步,客户端就会直接看到 `<AETHER:EMAIL:...>`。
|
||||
//! [`ResponsesWebSocketRedactionRestorer`] 补上这一跳,语义与 HTTP 完全一致:
|
||||
//! 复用 `privacy::restore_json_strings`,只还原本连接自己 mask 出来的映射,
|
||||
//! 未映射的占位符原样透传。
|
||||
//!
|
||||
//! ## session 为什么活在连接上而不是活在这一轮里
|
||||
//!
|
||||
//! mask session 由 planner 写进 per-turn 的 slot,而 slot 随 planning Parts 在
|
||||
//! 规划结束时就被丢弃,响应帧到达时已经无处可取。可选的存活范围有两个:
|
||||
//!
|
||||
//! * 挂在 `LogicalTurn` 上:这一轮结束即释放,是 HTTP「一个请求一个 session」的
|
||||
//! 直译。但 WS 的会话历史留在上游:continuation 只发增量输入,第 1 轮的
|
||||
//! `input` 不会在第 3 轮重发。于是第 3 轮的响应里若回显了第 1 轮的占位符
|
||||
//! ("你刚才给我的邮箱是……"),本轮 session 里没有这条映射,占位符就漏给客户端。
|
||||
//! HTTP 不会漏,是因为它每次都重发整段历史,重新 mask 同一个值会派生出同一个
|
||||
//! sentinel(HMAC over 规则 + bucket + 值),所以映射天然齐备。
|
||||
//! * 挂在连接上(当前实现):每轮仍然各自 mask、各自持有独立 session
|
||||
//! (per-turn 语义不变),连接只是把最近若干轮的 session 留下来一起参与还原,
|
||||
//! 凑出的映射集合正好等于「等价 HTTP 请求会拥有的那一份」。
|
||||
//!
|
||||
//! 选后者。代价是每帧最多对 [`MAX_RETAINED_TURN_REDACTION_SESSIONS`] 个 session
|
||||
//! 各扫一遍,以及这些 session 的映射会驻留到连接结束;用有界 FIFO 兜住上限。
|
||||
//! 窗口不够用或每帧成本变高时,正确的下一步是在 `privacy` 侧提供跨 session 的
|
||||
//! 合并匹配器,而不是把这个窗口调大。
|
||||
|
||||
use std::collections::VecDeque;
|
||||
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::ai_serving::{
|
||||
resolve_local_decision_execution_runtime_auth_context, resolve_provider_chat_pii_redaction,
|
||||
};
|
||||
use crate::control::GatewayControlDecision;
|
||||
use crate::privacy::{restore_json_strings, RedactionSession, RedactionSessionSlot};
|
||||
use crate::{AppState, GatewayError};
|
||||
|
||||
/// Responses WebSocket 只承载 `openai:responses`,脱敏规则按这个客户端格式选取。
|
||||
const RESPONSES_WEBSOCKET_CLIENT_API_FORMAT: &str = "openai:responses";
|
||||
|
||||
/// WS 在选出候选之前就要脱敏,所以脱敏 session 先记在这个固定 key 下。
|
||||
///
|
||||
/// slot 是 per-turn 的(见 `build_planning_parts`),这一轮之后即随 slot 一起丢弃;
|
||||
/// planner 后续用真实 candidate_id 再取一次配置时,body 已是脱敏态、不会重复写入。
|
||||
const WEBSOCKET_TURN_REDACTION_CANDIDATE_ID: &str = "responses_websocket_turn";
|
||||
|
||||
/// 一条连接最多留几轮的 mask session 用于响应侧还原。
|
||||
///
|
||||
/// 取值权衡见模块文档:调大会线性增加每帧还原成本和常驻映射量,调小则更容易漏还原
|
||||
/// 上游历史里更早那几轮的占位符。8 覆盖的是「上游最可能回显的最近窗口」。
|
||||
const MAX_RETAINED_TURN_REDACTION_SESSIONS: usize = 8;
|
||||
|
||||
/// 一轮客户端 `response.create` 的请求侧脱敏结果。
|
||||
#[derive(Debug)]
|
||||
pub(super) struct ResponsesWebSocketTurnRedaction {
|
||||
/// 脱敏后的客户端事件;这一轮之后所有下游路径都只看它。
|
||||
pub(super) client_event: Value,
|
||||
/// 这一轮 mask 出来的映射表,响应侧还原只能靠它。
|
||||
pub(super) session: RedactionSession,
|
||||
}
|
||||
|
||||
/// 对一条客户端 `response.create` 做请求侧脱敏。
|
||||
///
|
||||
/// 返回 `Some(..)` 仅当脱敏真正命中;`None` 表示未启用或没有命中,调用方
|
||||
/// 继续用原事件即可(避免未开启脱敏时多一次整包 clone)。
|
||||
///
|
||||
/// 脱敏只改写 `instructions` / `input`(见 `privacy::mask_openai_responses_request_value`),
|
||||
/// `type` / `model` / `previous_response_id` / `generate` 等协议字段原样保留,所以脱敏后的
|
||||
/// 事件仍可直接用于协议归一化和上游发送。
|
||||
///
|
||||
/// 出错必须让这一轮失败:脱敏已启用却读不到配置或加密密钥时,把原文发上游就是
|
||||
/// 静默旁路,正是本次要修的问题。
|
||||
pub(super) async fn redact_responses_websocket_client_event(
|
||||
state: &AppState,
|
||||
parts: &http::request::Parts,
|
||||
control_decision: &GatewayControlDecision,
|
||||
client_event: &Value,
|
||||
) -> Result<Option<ResponsesWebSocketTurnRedaction>, GatewayError> {
|
||||
let Some(auth_context) =
|
||||
resolve_local_decision_execution_runtime_auth_context(control_decision)
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
let redaction = resolve_provider_chat_pii_redaction(
|
||||
state,
|
||||
parts,
|
||||
client_event,
|
||||
&auth_context,
|
||||
RESPONSES_WEBSOCKET_CLIENT_API_FORMAT,
|
||||
WEBSOCKET_TURN_REDACTION_CANDIDATE_ID,
|
||||
)
|
||||
.await?;
|
||||
if !redaction.redacted {
|
||||
return Ok(None);
|
||||
}
|
||||
// mask 命中时 `resolve_provider_chat_pii_redaction` 必定把 session 写进 slot。
|
||||
// 取不到就是内部契约被破坏了,此时继续下发意味着这一轮的响应无法还原、占位符
|
||||
// 会漏给客户端;按本模块既有的「脱敏链路出错就让这一轮失败」处理,不做降级。
|
||||
let Some(session) = parts
|
||||
.extensions
|
||||
.get::<RedactionSessionSlot>()
|
||||
.and_then(|slot| slot.take_for_candidate(Some(WEBSOCKET_TURN_REDACTION_CANDIDATE_ID)))
|
||||
else {
|
||||
return Err(GatewayError::Internal(
|
||||
"chat pii redaction masked a Responses WebSocket turn without retaining its session"
|
||||
.to_string(),
|
||||
));
|
||||
};
|
||||
Ok(Some(ResponsesWebSocketTurnRedaction {
|
||||
client_event: redaction.body_json.into_owned(),
|
||||
session,
|
||||
}))
|
||||
}
|
||||
|
||||
/// 一条连接上「我们 mask 过哪些映射」的留存集合,供响应侧还原使用。
|
||||
///
|
||||
/// 每轮一个独立 session(per-turn mask 语义不变),连接按 FIFO 留最近
|
||||
/// [`MAX_RETAINED_TURN_REDACTION_SESSIONS`] 轮。上游重绑不清空:客户端仍在同一段
|
||||
/// 对话里,旧占位符可能随重发的输入再次出现。
|
||||
#[derive(Default)]
|
||||
pub(super) struct ResponsesWebSocketRedactionRestorer {
|
||||
sessions: VecDeque<RedactionSession>,
|
||||
}
|
||||
|
||||
impl ResponsesWebSocketRedactionRestorer {
|
||||
/// 登记这一轮的 mask session。
|
||||
pub(super) fn register(&mut self, session: RedactionSession) {
|
||||
if session.mapping_count() == 0 {
|
||||
return;
|
||||
}
|
||||
self.sessions.push_back(session);
|
||||
while self.sessions.len() > MAX_RETAINED_TURN_REDACTION_SESSIONS {
|
||||
self.sessions.pop_front();
|
||||
}
|
||||
}
|
||||
|
||||
/// 把一帧 provider 事件里的占位符换回真实值,返回要发给客户端的帧文本。
|
||||
///
|
||||
/// `None` 表示这一帧没有任何东西要还原,调用方必须原样转发上游字节:未启用
|
||||
/// 脱敏(没有任何 session)时连 clone 都不做。
|
||||
///
|
||||
/// 入参只读:审计与终态观测继续消费脱敏态的事件,还原只作用于发往客户端的
|
||||
/// 那一份拷贝,和 HTTP 侧「审计存脱敏体、线上还原」保持一致。
|
||||
pub(super) fn restore_provider_frame_text(&self, event: &Value) -> Option<String> {
|
||||
if self.sessions.is_empty() {
|
||||
return None;
|
||||
}
|
||||
let mut restored_event = event.clone();
|
||||
let mut restored = false;
|
||||
for session in &self.sessions {
|
||||
// 逐 session 还原而不是合并映射:每个 session 只认自己 mask 过的
|
||||
// sentinel(`RedactionSession::restore_text`),跨 session 合并会绕开
|
||||
// 这条边界。同一个值在不同轮派生出的 sentinel 相同,所以顺序无关。
|
||||
restored |= restore_json_strings(&mut restored_event, session);
|
||||
}
|
||||
if !restored {
|
||||
return None;
|
||||
}
|
||||
// 刚从 JSON 解析出来的 Value 再序列化不会失败;真失败时宁可让客户端看到
|
||||
// 占位符,也不能丢掉这一帧——丢帧会让客户端的协议状态机卡死。
|
||||
serde_json::to_string(&restored_event).ok()
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::net::SocketAddr;
|
||||
use std::sync::Arc;
|
||||
|
||||
use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY;
|
||||
use aether_data::repository::auth::{
|
||||
InMemoryAuthApiKeySnapshotRepository, StoredAuthApiKeyExportRecord,
|
||||
};
|
||||
use axum::http::{HeaderMap, Uri};
|
||||
use serde_json::{json, Value};
|
||||
|
||||
use super::super::request::{
|
||||
build_planning_parts, normalize_followup_response_create, planned_response_create_event,
|
||||
};
|
||||
use super::super::turn::prepare_responses_websocket_turn_decision;
|
||||
use super::super::turn_state::LogicalTurn;
|
||||
use super::{
|
||||
redact_responses_websocket_client_event, ResponsesWebSocketRedactionRestorer,
|
||||
ResponsesWebSocketTurnRedaction, MAX_RETAINED_TURN_REDACTION_SESSIONS,
|
||||
};
|
||||
use crate::ai_serving::{AiExecutionDecision, ResponsesWebSocketBodyNormalization};
|
||||
use crate::control::{GatewayControlAuthContext, GatewayControlDecision};
|
||||
use crate::handlers::proxy::websocket::ingress::WebSocketRequestContext;
|
||||
use crate::AppState;
|
||||
|
||||
const TEST_USER_ID: &str = "user-responses-ws-redaction";
|
||||
const TEST_API_KEY_ID: &str = "api-key-responses-ws-redaction";
|
||||
const TEST_EMAIL: &str = "[email protected]";
|
||||
/// 另一轮用的 PII,用来证明连接级还原覆盖到更早的轮次。
|
||||
const OTHER_TEST_EMAIL: &str = "[email protected]";
|
||||
/// 不是本连接 mask 出来的占位符:格式合法(符合 sentinel 正则),但没有任何
|
||||
/// session 记过它,必须原样透传。
|
||||
const FOREIGN_SENTINEL: &str = "<AETHER:EMAIL:AAAAAAAAAAAAAAAAAAAA>";
|
||||
|
||||
fn auth_export_record() -> StoredAuthApiKeyExportRecord {
|
||||
StoredAuthApiKeyExportRecord::new(
|
||||
TEST_USER_ID.to_string(),
|
||||
TEST_API_KEY_ID.to_string(),
|
||||
"hash-responses-ws-redaction".to_string(),
|
||||
None,
|
||||
Some("ws".to_string()),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
true,
|
||||
None,
|
||||
false,
|
||||
0,
|
||||
0,
|
||||
0.0,
|
||||
false,
|
||||
)
|
||||
.expect("auth api key export record should build")
|
||||
.with_feature_settings(Some(json!({
|
||||
"chat_pii_redaction": {"enabled": true}
|
||||
})))
|
||||
}
|
||||
|
||||
/// 只装脱敏真正需要的东西:系统配置开关 + 规则、加密密钥、带 feature settings
|
||||
/// 的 API Key 导出记录。候选/上游都不需要,这条链路在 planner 之前。
|
||||
fn redaction_enabled_state() -> AppState {
|
||||
let auth_repository = Arc::new(
|
||||
InMemoryAuthApiKeySnapshotRepository::seed(vec![])
|
||||
.with_export_records(vec![auth_export_record()]),
|
||||
);
|
||||
let data_state =
|
||||
crate::data::GatewayDataState::with_auth_api_key_reader_for_tests(auth_repository)
|
||||
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY)
|
||||
.with_system_config_values_for_tests(vec![
|
||||
("module.chat_pii_redaction.enabled".to_string(), json!(true)),
|
||||
(
|
||||
"module.chat_pii_redaction.rules".to_string(),
|
||||
json!([{
|
||||
"id": "email",
|
||||
"name": "邮箱",
|
||||
"pattern": r"(?i)[A-Z0-9._%+-]{1,64}@[A-Z0-9.-]{1,253}\.[A-Z]{2,63}",
|
||||
"enabled": true,
|
||||
"features": {"validator": "email"},
|
||||
"system": true
|
||||
}]),
|
||||
),
|
||||
(
|
||||
"module.chat_pii_redaction.cache_ttl_seconds".to_string(),
|
||||
json!(300),
|
||||
),
|
||||
]);
|
||||
AppState::new()
|
||||
.expect("gateway state should build")
|
||||
.with_data_state_for_tests(data_state)
|
||||
}
|
||||
|
||||
fn control_decision() -> GatewayControlDecision {
|
||||
let mut decision = GatewayControlDecision::synthetic(
|
||||
"/v1/responses".to_string(),
|
||||
Some("ai_public".to_string()),
|
||||
Some("openai".to_string()),
|
||||
Some("responses_websocket".to_string()),
|
||||
Some("openai:responses".to_string()),
|
||||
);
|
||||
decision.auth_context = Some(GatewayControlAuthContext {
|
||||
user_id: TEST_USER_ID.to_string(),
|
||||
api_key_id: TEST_API_KEY_ID.to_string(),
|
||||
username: Some("ws".to_string()),
|
||||
api_key_name: Some("ws".to_string()),
|
||||
balance_remaining: None,
|
||||
access_allowed: true,
|
||||
user_rate_limit: None,
|
||||
api_key_rate_limit: None,
|
||||
api_key_is_standalone: false,
|
||||
admin_bypass_limits: false,
|
||||
local_rejection: None,
|
||||
allowed_models: None,
|
||||
ip_rules: None,
|
||||
});
|
||||
decision
|
||||
}
|
||||
|
||||
fn websocket_context(decision: GatewayControlDecision) -> WebSocketRequestContext {
|
||||
WebSocketRequestContext {
|
||||
trace_id: "trace-responses-ws-redaction".to_string(),
|
||||
headers: HeaderMap::new(),
|
||||
uri: Uri::from_static("/v1/responses"),
|
||||
remote_addr: "127.0.0.1:65000"
|
||||
.parse::<SocketAddr>()
|
||||
.expect("remote address should parse"),
|
||||
client_ip: "127.0.0.1".parse().expect("client IP should parse"),
|
||||
decision,
|
||||
websocket_connection_permit: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn client_event() -> Value {
|
||||
client_event_with_email(TEST_EMAIL)
|
||||
}
|
||||
|
||||
fn client_event_with_email(email: &str) -> Value {
|
||||
json!({
|
||||
"type": "response.create",
|
||||
"model": "public-model",
|
||||
"previous_response_id": "resp-previous",
|
||||
"generate": false,
|
||||
"input": [{
|
||||
"role": "user",
|
||||
"content": [{"type": "input_text", "text": format!("mail {email}")}]
|
||||
}]
|
||||
})
|
||||
}
|
||||
|
||||
/// 真跑一遍请求侧脱敏,拿到这一轮的生效事件和 mask session。
|
||||
async fn turn_redaction(
|
||||
state: &AppState,
|
||||
decision: &GatewayControlDecision,
|
||||
email: &str,
|
||||
) -> ResponsesWebSocketTurnRedaction {
|
||||
let context = websocket_context(decision.clone());
|
||||
let parts = build_planning_parts(&context);
|
||||
let event = client_event_with_email(email);
|
||||
redact_responses_websocket_client_event(state, &parts, &context.decision, &event)
|
||||
.await
|
||||
.expect("redaction should resolve")
|
||||
.expect("an email in the request should be redacted")
|
||||
}
|
||||
|
||||
/// 这一轮为 `email` 派生出的占位符。
|
||||
fn sentinel_for(redaction: &ResponsesWebSocketTurnRedaction, email: &str) -> String {
|
||||
redaction
|
||||
.session
|
||||
.sentinel_for_original(email)
|
||||
.expect("a masked email must have a sentinel")
|
||||
.to_string()
|
||||
}
|
||||
|
||||
/// 上游回显占位符的一帧 provider 事件。
|
||||
fn provider_delta_frame(text: &str) -> Value {
|
||||
json!({
|
||||
"type": "response.output_text.delta",
|
||||
"item_id": "msg_ws",
|
||||
"output_index": 0,
|
||||
"content_index": 0,
|
||||
"delta": text,
|
||||
})
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn websocket_client_event_is_redacted_without_losing_protocol_fields() {
|
||||
let state = redaction_enabled_state();
|
||||
let context = websocket_context(control_decision());
|
||||
let parts = build_planning_parts(&context);
|
||||
let event = client_event();
|
||||
|
||||
let redacted =
|
||||
redact_responses_websocket_client_event(&state, &parts, &context.decision, &event)
|
||||
.await
|
||||
.expect("redaction should resolve")
|
||||
.expect("an email in the request should be redacted")
|
||||
.client_event;
|
||||
|
||||
let serialized = serde_json::to_string(&redacted).expect("event should serialize");
|
||||
assert!(!serialized.contains(TEST_EMAIL), "{serialized}");
|
||||
assert!(serialized.contains("<AETHER:EMAIL:"), "{serialized}");
|
||||
// 协议字段必须原样保留,否则 continuation 链路会断。
|
||||
assert_eq!(redacted["type"], "response.create");
|
||||
assert_eq!(redacted["model"], "public-model");
|
||||
assert_eq!(redacted["previous_response_id"], "resp-previous");
|
||||
assert_eq!(redacted["generate"], false);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn redacting_an_already_redacted_event_is_a_no_op() {
|
||||
// re-plan 与配额重试路径会把已脱敏的事件再交给 planner,planner 内部会对
|
||||
// 同一个 body 再跑一遍 mask。占位符本身不该被任何规则命中,否则会被二次
|
||||
// 替换、破坏与上游已有 previous_response_id 链的一致性。
|
||||
let state = redaction_enabled_state();
|
||||
let context = websocket_context(control_decision());
|
||||
let parts = build_planning_parts(&context);
|
||||
let event = client_event();
|
||||
|
||||
let redacted =
|
||||
redact_responses_websocket_client_event(&state, &parts, &context.decision, &event)
|
||||
.await
|
||||
.expect("redaction should resolve")
|
||||
.expect("an email in the request should be redacted")
|
||||
.client_event;
|
||||
|
||||
// 复用同一个 parts/slot,和 re-plan 在同一 turn 内二次脱敏的情形一致。
|
||||
let second_pass =
|
||||
redact_responses_websocket_client_event(&state, &parts, &context.decision, &redacted)
|
||||
.await
|
||||
.expect("second redaction pass should resolve");
|
||||
|
||||
assert!(
|
||||
second_pass.is_none(),
|
||||
"already redacted event should stay byte-identical: {second_pass:?}"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn redaction_is_skipped_without_a_local_auth_context() {
|
||||
let state = redaction_enabled_state();
|
||||
let mut decision = control_decision();
|
||||
decision.auth_context = None;
|
||||
let context = websocket_context(decision);
|
||||
let parts = build_planning_parts(&context);
|
||||
let event = client_event();
|
||||
|
||||
let redacted =
|
||||
redact_responses_websocket_client_event(&state, &parts, &context.decision, &event)
|
||||
.await
|
||||
.expect("redaction should resolve");
|
||||
|
||||
assert!(redacted.is_none());
|
||||
}
|
||||
|
||||
/// 真跑一遍脱敏,拿到这一轮的「生效事件」。
|
||||
async fn redacted_client_event(state: &AppState, decision: &GatewayControlDecision) -> Value {
|
||||
turn_redaction(state, decision, TEST_EMAIL)
|
||||
.await
|
||||
.client_event
|
||||
}
|
||||
|
||||
/// 只有 `action` 没有 serde 默认值,其余字段都能省略。
|
||||
fn decision_template(
|
||||
provider_request_body: Value,
|
||||
report_context: Value,
|
||||
) -> AiExecutionDecision {
|
||||
serde_json::from_value(json!({
|
||||
"action": "local",
|
||||
"candidate_id": "candidate-responses-ws",
|
||||
"provider_request_body": provider_request_body,
|
||||
"report_context": report_context,
|
||||
}))
|
||||
.expect("decision template should deserialize")
|
||||
}
|
||||
|
||||
/// planner 在脱敏 body 上做模型映射后的 provider body。
|
||||
fn provider_body_from(effective_event: &Value) -> Value {
|
||||
let mut provider_body = effective_event.clone();
|
||||
provider_body["model"] = json!("provider-model");
|
||||
provider_body
|
||||
}
|
||||
|
||||
/// 绑定那一轮留下的 report_context seed:故意带上原始 PII,用来证明这一轮
|
||||
/// 会用脱敏后的 body 覆盖它,而不是把原文带进审计。
|
||||
fn seed_report_context_with_raw_pii() -> Value {
|
||||
json!({
|
||||
"request_id": "connection",
|
||||
"candidate_id": "candidate-responses-ws",
|
||||
"original_request_body": {
|
||||
"type": "response.create",
|
||||
"model": "public-model",
|
||||
"input": format!("mail {TEST_EMAIL}")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
fn assert_redacted_json(value: &Value, label: &str) {
|
||||
let serialized = serde_json::to_string(value).expect("value should serialize");
|
||||
assert!(
|
||||
!serialized.contains(TEST_EMAIL),
|
||||
"{label} must not carry raw PII: {serialized}"
|
||||
);
|
||||
assert!(
|
||||
serialized.contains("<AETHER:EMAIL:"),
|
||||
"{label} must carry the redaction sentinel: {serialized}"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn first_turn_upstream_and_audit_bodies_are_redacted() {
|
||||
let state = redaction_enabled_state();
|
||||
let decision = control_decision();
|
||||
let effective_event = redacted_client_event(&state, &decision).await;
|
||||
let template = decision_template(
|
||||
provider_body_from(&effective_event),
|
||||
seed_report_context_with_raw_pii(),
|
||||
);
|
||||
// 首轮实际发上游的事件由 decision.provider_request_body 派生。
|
||||
let provider_event: Value = serde_json::from_str(
|
||||
&planned_response_create_event(&template, &effective_event)
|
||||
.expect("first provider event should serialize"),
|
||||
)
|
||||
.expect("first provider event should parse");
|
||||
|
||||
let turn_decision = prepare_responses_websocket_turn_decision(
|
||||
&template,
|
||||
"turn-1".to_string(),
|
||||
true,
|
||||
&effective_event,
|
||||
&provider_event,
|
||||
"connection",
|
||||
1,
|
||||
"logical-turn-1",
|
||||
1,
|
||||
);
|
||||
|
||||
assert_redacted_json(&provider_event, "first turn upstream event");
|
||||
assert_redacted_json(
|
||||
turn_decision
|
||||
.provider_request_body
|
||||
.as_ref()
|
||||
.expect("turn decision should carry a provider body"),
|
||||
"first turn provider request body",
|
||||
);
|
||||
let report_context = turn_decision
|
||||
.report_context
|
||||
.as_ref()
|
||||
.expect("turn decision should carry a report context");
|
||||
assert_redacted_json(
|
||||
&report_context["original_request_body"],
|
||||
"first turn audit body",
|
||||
);
|
||||
// 整个 report_context 都不该残留原文(seed 里的原始 body 必须被覆盖)。
|
||||
assert_redacted_json(report_context, "first turn report context");
|
||||
assert_eq!(provider_event["type"], "response.create");
|
||||
assert_eq!(provider_event["model"], "provider-model");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn continuation_upstream_and_audit_bodies_are_redacted() {
|
||||
let state = redaction_enabled_state();
|
||||
let decision = control_decision();
|
||||
let effective_event = redacted_client_event(&state, &decision).await;
|
||||
// continuation 复用已绑定的 upstream:不再规划,直接重放归一化器。
|
||||
let outbound = normalize_followup_response_create(
|
||||
&effective_event,
|
||||
"provider-model",
|
||||
&ResponsesWebSocketBodyNormalization::for_tests("provider-model"),
|
||||
)
|
||||
.expect("continuation should normalize");
|
||||
let provider_event: Value =
|
||||
serde_json::from_str(&outbound).expect("continuation event should parse");
|
||||
|
||||
let template = decision_template(
|
||||
provider_body_from(&effective_event),
|
||||
seed_report_context_with_raw_pii(),
|
||||
);
|
||||
let turn_decision = prepare_responses_websocket_turn_decision(
|
||||
&template,
|
||||
"turn-2".to_string(),
|
||||
false,
|
||||
&effective_event,
|
||||
&provider_event,
|
||||
"connection",
|
||||
2,
|
||||
"logical-turn-2",
|
||||
1,
|
||||
);
|
||||
|
||||
assert!(
|
||||
!outbound.contains(TEST_EMAIL),
|
||||
"continuation upstream frame must not carry raw PII: {outbound}"
|
||||
);
|
||||
assert!(
|
||||
outbound.contains("<AETHER:EMAIL:"),
|
||||
"continuation upstream frame must carry the sentinel: {outbound}"
|
||||
);
|
||||
assert_eq!(provider_event["previous_response_id"], "resp-previous");
|
||||
let report_context = turn_decision
|
||||
.report_context
|
||||
.as_ref()
|
||||
.expect("turn decision should carry a report context");
|
||||
assert_redacted_json(
|
||||
&report_context["original_request_body"],
|
||||
"continuation audit body",
|
||||
);
|
||||
assert_redacted_json(report_context, "continuation report context");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn quota_retry_replays_the_redacted_event() {
|
||||
let state = redaction_enabled_state();
|
||||
let decision = control_decision();
|
||||
let effective_event = redacted_client_event(&state, &decision).await;
|
||||
// 配额透明重试重放 LogicalTurn 里保存的事件,所以保存的必须
|
||||
// 已经是脱敏版,否则重试会把原文发给新的上游账号。
|
||||
let active = LogicalTurn::new(effective_event.clone(), 2, "logical-turn-2".to_string());
|
||||
assert_redacted_json(&active.client_event, "quota retry replay event");
|
||||
|
||||
let template = decision_template(
|
||||
provider_body_from(&active.client_event),
|
||||
seed_report_context_with_raw_pii(),
|
||||
);
|
||||
let provider_event: Value = serde_json::from_str(
|
||||
&planned_response_create_event(&template, &active.client_event)
|
||||
.expect("retry provider event should serialize"),
|
||||
)
|
||||
.expect("retry provider event should parse");
|
||||
let turn_decision = prepare_responses_websocket_turn_decision(
|
||||
&template,
|
||||
"turn-2-retry".to_string(),
|
||||
true,
|
||||
&active.client_event,
|
||||
&provider_event,
|
||||
"connection",
|
||||
active.turn_index,
|
||||
"logical-turn-2",
|
||||
2,
|
||||
);
|
||||
|
||||
assert_redacted_json(&provider_event, "quota retry upstream event");
|
||||
let report_context = turn_decision
|
||||
.report_context
|
||||
.as_ref()
|
||||
.expect("turn decision should carry a report context");
|
||||
assert_eq!(report_context["websocket_turn_attempt"], 2);
|
||||
assert_redacted_json(
|
||||
&report_context["original_request_body"],
|
||||
"quota retry audit body",
|
||||
);
|
||||
assert_redacted_json(report_context, "quota retry report context");
|
||||
}
|
||||
|
||||
// -----------------------------------------------------------------------
|
||||
// 响应侧还原
|
||||
// -----------------------------------------------------------------------
|
||||
|
||||
/// 本次修复的核心:上游把占位符回显在事件里,客户端必须拿到真实值。
|
||||
#[tokio::test]
|
||||
async fn provider_frame_placeholders_are_restored_before_client_delivery() {
|
||||
let state = redaction_enabled_state();
|
||||
let redaction = turn_redaction(&state, &control_decision(), TEST_EMAIL).await;
|
||||
let sentinel = sentinel_for(&redaction, TEST_EMAIL);
|
||||
let mut restorer = ResponsesWebSocketRedactionRestorer::default();
|
||||
restorer.register(redaction.session);
|
||||
|
||||
let frame = provider_delta_frame(&format!("your mail is {sentinel}"));
|
||||
let restored = restorer
|
||||
.restore_provider_frame_text(&frame)
|
||||
.expect("a frame echoing this turn's sentinel must be restored");
|
||||
|
||||
assert!(
|
||||
restored.contains(TEST_EMAIL),
|
||||
"the client must receive the real value: {restored}"
|
||||
);
|
||||
assert!(
|
||||
!restored.contains(&sentinel),
|
||||
"no sentinel may survive to the client: {restored}"
|
||||
);
|
||||
// 协议字段不受影响,客户端的状态机照旧。
|
||||
let restored: Value = serde_json::from_str(&restored).expect("restored frame is JSON");
|
||||
assert_eq!(restored["type"], "response.output_text.delta");
|
||||
assert_eq!(restored["item_id"], "msg_ws");
|
||||
assert_eq!(restored["output_index"], 0);
|
||||
}
|
||||
|
||||
/// Codex 把多个事件批量塞进 `{"chunks":[...]}`,还原必须走进批量里。
|
||||
#[tokio::test]
|
||||
async fn placeholders_batched_inside_a_chunks_envelope_are_restored() {
|
||||
let state = redaction_enabled_state();
|
||||
let redaction = turn_redaction(&state, &control_decision(), TEST_EMAIL).await;
|
||||
let sentinel = sentinel_for(&redaction, TEST_EMAIL);
|
||||
let mut restorer = ResponsesWebSocketRedactionRestorer::default();
|
||||
restorer.register(redaction.session);
|
||||
|
||||
let frame = json!({
|
||||
"chunks": [
|
||||
provider_delta_frame("plain delta"),
|
||||
provider_delta_frame(&format!("mail {sentinel}")),
|
||||
]
|
||||
});
|
||||
let restored = restorer
|
||||
.restore_provider_frame_text(&frame)
|
||||
.expect("a batched sentinel must be restored");
|
||||
|
||||
assert!(restored.contains(TEST_EMAIL), "{restored}");
|
||||
assert!(!restored.contains(&sentinel), "{restored}");
|
||||
}
|
||||
|
||||
/// 只还原本连接 mask 过的映射,和 `RedactionSession::restore_text` 一致:
|
||||
/// 别处来的占位符(比如客户端自己发的、或上一条连接的)保持原样。
|
||||
#[tokio::test]
|
||||
async fn an_unmapped_placeholder_is_left_untouched() {
|
||||
let state = redaction_enabled_state();
|
||||
let redaction = turn_redaction(&state, &control_decision(), TEST_EMAIL).await;
|
||||
let sentinel = sentinel_for(&redaction, TEST_EMAIL);
|
||||
let mut restorer = ResponsesWebSocketRedactionRestorer::default();
|
||||
restorer.register(redaction.session);
|
||||
|
||||
let frame = provider_delta_frame(&format!("{FOREIGN_SENTINEL} and {sentinel}"));
|
||||
let restored = restorer
|
||||
.restore_provider_frame_text(&frame)
|
||||
.expect("the mapped sentinel is still restored");
|
||||
|
||||
assert!(restored.contains(TEST_EMAIL), "{restored}");
|
||||
assert!(
|
||||
restored.contains(FOREIGN_SENTINEL),
|
||||
"an unmapped placeholder must survive verbatim: {restored}"
|
||||
);
|
||||
}
|
||||
|
||||
/// 没有命中还原时必须让调用方原样转发上游字节。
|
||||
#[tokio::test]
|
||||
async fn a_frame_without_known_placeholders_is_not_rewritten() {
|
||||
let state = redaction_enabled_state();
|
||||
let redaction = turn_redaction(&state, &control_decision(), TEST_EMAIL).await;
|
||||
let mut restorer = ResponsesWebSocketRedactionRestorer::default();
|
||||
restorer.register(redaction.session);
|
||||
|
||||
assert!(
|
||||
restorer
|
||||
.restore_provider_frame_text(&provider_delta_frame("nothing to restore"))
|
||||
.is_none(),
|
||||
"a frame with no mapped sentinel must be relayed byte-for-byte"
|
||||
);
|
||||
assert!(
|
||||
restorer
|
||||
.restore_provider_frame_text(&provider_delta_frame(FOREIGN_SENTINEL))
|
||||
.is_none(),
|
||||
"a frame that only carries unmapped placeholders must not be rewritten"
|
||||
);
|
||||
}
|
||||
|
||||
/// 未启用脱敏(或这条连接从没 mask 到东西)时,还原器必须完全不介入:
|
||||
/// 连 clone 都不做,输出就是上游原字节。
|
||||
#[tokio::test]
|
||||
async fn a_restorer_without_sessions_never_rewrites_a_frame() {
|
||||
let restorer = ResponsesWebSocketRedactionRestorer::default();
|
||||
|
||||
assert!(restorer
|
||||
.restore_provider_frame_text(&provider_delta_frame(FOREIGN_SENTINEL))
|
||||
.is_none());
|
||||
assert!(restorer
|
||||
.restore_provider_frame_text(&provider_delta_frame(TEST_EMAIL))
|
||||
.is_none());
|
||||
}
|
||||
|
||||
/// 空 session(启用了脱敏但这一轮没命中任何规则)不该被留下来白扫每一帧。
|
||||
#[tokio::test]
|
||||
async fn a_session_without_mappings_is_not_retained() {
|
||||
let state = redaction_enabled_state();
|
||||
let hmac_key = state
|
||||
.encryption_key()
|
||||
.expect("the test state carries an encryption key")
|
||||
.as_bytes()
|
||||
.to_vec();
|
||||
let empty_session = crate::privacy::RedactionSession::new(
|
||||
crate::privacy::RedactionSessionConfig::default_ttl(hmac_key, 0),
|
||||
);
|
||||
let mut restorer = ResponsesWebSocketRedactionRestorer::default();
|
||||
restorer.register(empty_session);
|
||||
|
||||
assert!(restorer
|
||||
.restore_provider_frame_text(&provider_delta_frame(FOREIGN_SENTINEL))
|
||||
.is_none());
|
||||
}
|
||||
|
||||
/// 还原只作用于发给客户端的那一份拷贝:审计和终态观测消费的事件必须保持脱敏态。
|
||||
#[tokio::test]
|
||||
async fn restoring_does_not_mutate_the_event_the_audit_path_keeps() {
|
||||
let state = redaction_enabled_state();
|
||||
let redaction = turn_redaction(&state, &control_decision(), TEST_EMAIL).await;
|
||||
let sentinel = sentinel_for(&redaction, TEST_EMAIL);
|
||||
let mut restorer = ResponsesWebSocketRedactionRestorer::default();
|
||||
restorer.register(redaction.session);
|
||||
|
||||
let frame = provider_delta_frame(&format!("mail {sentinel}"));
|
||||
let before = frame.clone();
|
||||
let _ = restorer
|
||||
.restore_provider_frame_text(&frame)
|
||||
.expect("the frame is restored for the client");
|
||||
|
||||
assert_eq!(
|
||||
frame, before,
|
||||
"capture_client_frame / 终态观测拿到的事件必须仍是脱敏态"
|
||||
);
|
||||
}
|
||||
|
||||
/// 连接级持有的意义:WS 的会话历史留在上游,continuation 只发增量输入,
|
||||
/// 所以第 2 轮的响应可能回显第 1 轮的占位符。per-turn 持有会漏掉这一条。
|
||||
#[tokio::test]
|
||||
async fn a_later_turn_restores_a_placeholder_first_masked_by_an_earlier_turn() {
|
||||
let state = redaction_enabled_state();
|
||||
let decision = control_decision();
|
||||
let first = turn_redaction(&state, &decision, TEST_EMAIL).await;
|
||||
let second = turn_redaction(&state, &decision, OTHER_TEST_EMAIL).await;
|
||||
let first_sentinel = sentinel_for(&first, TEST_EMAIL);
|
||||
let second_sentinel = sentinel_for(&second, OTHER_TEST_EMAIL);
|
||||
assert_ne!(first_sentinel, second_sentinel);
|
||||
|
||||
let mut restorer = ResponsesWebSocketRedactionRestorer::default();
|
||||
restorer.register(first.session);
|
||||
restorer.register(second.session);
|
||||
|
||||
let frame = provider_delta_frame(&format!("{first_sentinel} then {second_sentinel}"));
|
||||
let restored = restorer
|
||||
.restore_provider_frame_text(&frame)
|
||||
.expect("both turns' sentinels are restorable on this connection");
|
||||
|
||||
assert!(restored.contains(TEST_EMAIL), "{restored}");
|
||||
assert!(restored.contains(OTHER_TEST_EMAIL), "{restored}");
|
||||
assert!(!restored.contains(&first_sentinel), "{restored}");
|
||||
assert!(!restored.contains(&second_sentinel), "{restored}");
|
||||
}
|
||||
|
||||
/// 留存窗口是有界的:长连接不能无限累积映射,代价是更早的轮次会退回
|
||||
/// 「占位符原样透传」而不是被错误还原成别的值。
|
||||
#[tokio::test]
|
||||
async fn the_retained_session_window_is_bounded() {
|
||||
let state = redaction_enabled_state();
|
||||
let decision = control_decision();
|
||||
let oldest = turn_redaction(&state, &decision, TEST_EMAIL).await;
|
||||
let oldest_sentinel = sentinel_for(&oldest, TEST_EMAIL);
|
||||
let mut restorer = ResponsesWebSocketRedactionRestorer::default();
|
||||
restorer.register(oldest.session);
|
||||
|
||||
// 再灌满整个窗口,最老的那一轮必须被挤出去。
|
||||
let mut newest_sentinel = String::new();
|
||||
for index in 0..MAX_RETAINED_TURN_REDACTION_SESSIONS {
|
||||
let email = format!("ws.turn{index}@example.com");
|
||||
let redaction = turn_redaction(&state, &decision, &email).await;
|
||||
newest_sentinel = sentinel_for(&redaction, &email);
|
||||
restorer.register(redaction.session);
|
||||
}
|
||||
|
||||
assert!(
|
||||
restorer
|
||||
.restore_provider_frame_text(&provider_delta_frame(&oldest_sentinel))
|
||||
.is_none(),
|
||||
"the evicted turn's sentinel is relayed verbatim, never mis-restored"
|
||||
);
|
||||
assert!(
|
||||
restorer
|
||||
.restore_provider_frame_text(&provider_delta_frame(&newest_sentinel))
|
||||
.is_some(),
|
||||
"the most recent turns stay restorable"
|
||||
);
|
||||
}
|
||||
}
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user