diff --git a/apps/aether-gateway/src/ai_serving/planner/passthrough/provider/plans.rs b/apps/aether-gateway/src/ai_serving/planner/passthrough/provider/plans.rs index e4117cd37..ea311e00a 100644 --- a/apps/aether-gateway/src/ai_serving/planner/passthrough/provider/plans.rs +++ b/apps/aether-gateway/src/ai_serving/planner/passthrough/provider/plans.rs @@ -194,7 +194,7 @@ impl LocalExecutionAttemptSource for LocalSameFormatProviderSyncA self.input .routing_policy .as_ref() - .map(|policy| policy.execution_policy) + .map(|policy| policy.execution_policy.clone()) } async fn next_execution_attempt(&mut self) -> Result, GatewayError> { @@ -246,7 +246,7 @@ impl LocalExecutionAttemptSource self.input .routing_policy .as_ref() - .map(|policy| policy.execution_policy) + .map(|policy| policy.execution_policy.clone()) } async fn next_execution_attempt(&mut self) -> Result, GatewayError> { diff --git a/apps/aether-gateway/src/ai_serving/planner/report_context.rs b/apps/aether-gateway/src/ai_serving/planner/report_context.rs index 190ebd8fc..bd621e27b 100644 --- a/apps/aether-gateway/src/ai_serving/planner/report_context.rs +++ b/apps/aether-gateway/src/ai_serving/planner/report_context.rs @@ -118,7 +118,7 @@ pub(crate) fn build_local_execution_report_context( insert_pool_key_lease_report_context_fields(&mut extra_fields, parts.pool_key_lease); insert_scheduler_affinity_policy_report_context_field(&mut extra_fields, parts.routing_policy); if let Some(policy) = parts.routing_policy { - if let Ok(value) = serde_json::to_value(policy.execution_policy) { + if let Ok(value) = serde_json::to_value(&policy.execution_policy) { extra_fields.insert(ROUTING_EXECUTION_POLICY_REPORT_FIELD.to_string(), value); } } diff --git a/apps/aether-gateway/src/ai_serving/planner/specialized/files.rs b/apps/aether-gateway/src/ai_serving/planner/specialized/files.rs index 99524fda7..7093a121f 100644 --- a/apps/aether-gateway/src/ai_serving/planner/specialized/files.rs +++ b/apps/aether-gateway/src/ai_serving/planner/specialized/files.rs @@ -179,7 +179,7 @@ impl LocalExecutionAttemptSource for LocalGeminiFilesSyncAttemptS self.input .routing_policy .as_ref() - .map(|policy| policy.execution_policy) + .map(|policy| policy.execution_policy.clone()) } async fn next_execution_attempt(&mut self) -> Result, GatewayError> { @@ -224,7 +224,7 @@ impl LocalExecutionAttemptSource for LocalGeminiFilesStreamAtte self.input .routing_policy .as_ref() - .map(|policy| policy.execution_policy) + .map(|policy| policy.execution_policy.clone()) } async fn next_execution_attempt(&mut self) -> Result, GatewayError> { diff --git a/apps/aether-gateway/src/ai_serving/planner/specialized/image.rs b/apps/aether-gateway/src/ai_serving/planner/specialized/image.rs index c68a40fde..ea4ab7768 100644 --- a/apps/aether-gateway/src/ai_serving/planner/specialized/image.rs +++ b/apps/aether-gateway/src/ai_serving/planner/specialized/image.rs @@ -257,7 +257,7 @@ impl LocalExecutionAttemptSource for LocalOpenAiImageSyncAttemptS self.input .routing_policy .as_ref() - .map(|policy| policy.execution_policy) + .map(|policy| policy.execution_policy.clone()) } async fn next_execution_attempt(&mut self) -> Result, GatewayError> { @@ -302,7 +302,7 @@ impl LocalExecutionAttemptSource for LocalOpenAiImageStreamAtte self.input .routing_policy .as_ref() - .map(|policy| policy.execution_policy) + .map(|policy| policy.execution_policy.clone()) } async fn next_execution_attempt(&mut self) -> Result, GatewayError> { diff --git a/apps/aether-gateway/src/ai_serving/planner/specialized/video.rs b/apps/aether-gateway/src/ai_serving/planner/specialized/video.rs index daeee1e2b..7788ff225 100644 --- a/apps/aether-gateway/src/ai_serving/planner/specialized/video.rs +++ b/apps/aether-gateway/src/ai_serving/planner/specialized/video.rs @@ -109,7 +109,7 @@ impl LocalExecutionAttemptSource for LocalVideoCreateSyncAttemptS self.input .routing_policy .as_ref() - .map(|policy| policy.execution_policy) + .map(|policy| policy.execution_policy.clone()) } async fn next_execution_attempt(&mut self) -> Result, GatewayError> { diff --git a/apps/aether-gateway/src/ai_serving/planner/standard/family/build.rs b/apps/aether-gateway/src/ai_serving/planner/standard/family/build.rs index 384f8c928..e96cb4f19 100644 --- a/apps/aether-gateway/src/ai_serving/planner/standard/family/build.rs +++ b/apps/aether-gateway/src/ai_serving/planner/standard/family/build.rs @@ -182,7 +182,7 @@ impl LocalExecutionAttemptSource for LocalStandardSyncAttemptSour self.input .routing_policy .as_ref() - .map(|policy| policy.execution_policy) + .map(|policy| policy.execution_policy.clone()) } async fn next_execution_attempt(&mut self) -> Result, GatewayError> { @@ -232,7 +232,7 @@ impl LocalExecutionAttemptSource for LocalStandardStreamAttempt self.input .routing_policy .as_ref() - .map(|policy| policy.execution_policy) + .map(|policy| policy.execution_policy.clone()) } async fn next_execution_attempt(&mut self) -> Result, GatewayError> { diff --git a/apps/aether-gateway/src/ai_serving/planner/standard/openai/chat/plans/stream.rs b/apps/aether-gateway/src/ai_serving/planner/standard/openai/chat/plans/stream.rs index 9873bcf9d..0f6ae98c6 100644 --- a/apps/aether-gateway/src/ai_serving/planner/standard/openai/chat/plans/stream.rs +++ b/apps/aether-gateway/src/ai_serving/planner/standard/openai/chat/plans/stream.rs @@ -124,7 +124,7 @@ impl LocalExecutionAttemptSource for LocalOpenAiChatStreamAttem self.input .routing_policy .as_ref() - .map(|policy| policy.execution_policy) + .map(|policy| policy.execution_policy.clone()) } async fn next_execution_attempt(&mut self) -> Result, GatewayError> { diff --git a/apps/aether-gateway/src/ai_serving/planner/standard/openai/chat/plans/sync.rs b/apps/aether-gateway/src/ai_serving/planner/standard/openai/chat/plans/sync.rs index 6c54303d8..4badc4420 100644 --- a/apps/aether-gateway/src/ai_serving/planner/standard/openai/chat/plans/sync.rs +++ b/apps/aether-gateway/src/ai_serving/planner/standard/openai/chat/plans/sync.rs @@ -97,7 +97,7 @@ impl LocalExecutionAttemptSource for LocalOpenAiChatSyncAttemptSo self.input .routing_policy .as_ref() - .map(|policy| policy.execution_policy) + .map(|policy| policy.execution_policy.clone()) } async fn next_execution_attempt(&mut self) -> Result, GatewayError> { diff --git a/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/plans.rs b/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/plans.rs index 53cfd5d77..2e49eea0e 100644 --- a/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/plans.rs +++ b/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/plans.rs @@ -166,7 +166,7 @@ impl LocalExecutionAttemptSource for LocalOpenAiResponsesSyncAtte self.input .routing_policy .as_ref() - .map(|policy| policy.execution_policy) + .map(|policy| policy.execution_policy.clone()) } async fn next_execution_attempt(&mut self) -> Result, GatewayError> { @@ -216,7 +216,7 @@ impl LocalExecutionAttemptSource for LocalOpenAiResponsesStream self.input .routing_policy .as_ref() - .map(|policy| policy.execution_policy) + .map(|policy| policy.execution_policy.clone()) } async fn next_execution_attempt(&mut self) -> Result, GatewayError> { diff --git a/apps/aether-gateway/src/execution_runtime/fallback.rs b/apps/aether-gateway/src/execution_runtime/fallback.rs index 907932abc..7eb990a77 100644 --- a/apps/aether-gateway/src/execution_runtime/fallback.rs +++ b/apps/aether-gateway/src/execution_runtime/fallback.rs @@ -1036,6 +1036,7 @@ mod tests { policy, LocalFailoverPolicy { max_retries: Some(1), + routing_rules: Default::default(), max_transfer_count: 0, max_transfer_timeout_seconds: 0, stop_status_codes: [503].into_iter().collect(), diff --git a/apps/aether-gateway/src/execution_runtime/stream/commit_policy.rs b/apps/aether-gateway/src/execution_runtime/stream/commit_policy.rs index 9cfd0cae1..a5c1e4dcf 100644 --- a/apps/aether-gateway/src/execution_runtime/stream/commit_policy.rs +++ b/apps/aether-gateway/src/execution_runtime/stream/commit_policy.rs @@ -11,6 +11,10 @@ const GEMINI_PRECOMMIT_MAX_WAIT: Duration = Duration::from_millis(750); pub(super) enum StreamCommitPolicy { ResponseHeaders, FirstClassifiedBody, + FirstSseSemanticEvent { + max_bytes: usize, + max_wait: Duration, + }, FirstAnthropicSemanticEvent { max_bytes: usize, max_wait: Duration, @@ -36,16 +40,21 @@ impl StreamCommitPolicy { return Self::FirstClassifiedBody; } - if force_prefetch { - return Self::FirstClassifiedBody; - } - let content_type = content_type .map(str::trim) .filter(|value| !value.is_empty()) .unwrap_or_default() .to_ascii_lowercase(); if content_type.contains("text/event-stream") { + if provider_api_format.eq_ignore_ascii_case("openai:image") + || client_api_format.eq_ignore_ascii_case("openai:image") + { + return if force_prefetch { + Self::FirstClassifiedBody + } else { + Self::ResponseHeaders + }; + } if provider_api_format.eq_ignore_ascii_case("claude:messages") && provider_api_format.eq_ignore_ascii_case(client_api_format) && !has_private_stream_normalizer @@ -62,7 +71,14 @@ impl StreamCommitPolicy { max_wait: GEMINI_PRECOMMIT_MAX_WAIT, }; } - return Self::ResponseHeaders; + return Self::FirstSseSemanticEvent { + max_bytes: MAX_STREAM_PREFETCH_BYTES, + max_wait: Duration::from_secs(30), + }; + } + + if force_prefetch { + return Self::FirstClassifiedBody; } if has_private_stream_normalizer || has_local_stream_rewriter { @@ -91,14 +107,17 @@ impl StreamCommitPolicy { pub(super) const fn requires_bounded_frame_wait(self) -> bool { matches!( self, - Self::FirstAnthropicSemanticEvent { .. } | Self::FirstGeminiSemanticEvent { .. } + Self::FirstAnthropicSemanticEvent { .. } + | Self::FirstGeminiSemanticEvent { .. } + | Self::FirstSseSemanticEvent { .. } ) } pub(super) const fn max_precommit_wait(self) -> Option { match self { Self::FirstAnthropicSemanticEvent { max_wait, .. } - | Self::FirstGeminiSemanticEvent { max_wait, .. } => Some(max_wait), + | Self::FirstGeminiSemanticEvent { max_wait, .. } + | Self::FirstSseSemanticEvent { max_wait, .. } => Some(max_wait), Self::ResponseHeaders | Self::FirstClassifiedBody => None, } } @@ -110,6 +129,16 @@ impl StreamCommitPolicy { pub(super) const fn is_gemini(self) -> bool { matches!(self, Self::FirstGeminiSemanticEvent { .. }) } + + pub(super) fn with_precommit_wait(mut self, wait: Duration) -> Self { + match &mut self { + Self::FirstAnthropicSemanticEvent { max_wait, .. } + | Self::FirstGeminiSemanticEvent { max_wait, .. } + | Self::FirstSseSemanticEvent { max_wait, .. } => *max_wait = wait, + _ => {} + } + self + } } #[derive(Debug, Clone, Copy, PartialEq, Eq)] @@ -133,6 +162,7 @@ pub(super) struct StreamCommitGate { observed_bytes: usize, anthropic: AnthropicSsePrecommitInspector, gemini: GeminiSsePrecommitInspector, + generic: GenericSsePrecommitInspector, } impl StreamCommitGate { @@ -148,6 +178,7 @@ impl StreamCommitGate { observed_bytes: 0, anthropic: AnthropicSsePrecommitInspector::default(), gemini: GeminiSsePrecommitInspector::default(), + generic: GenericSsePrecommitInspector::default(), } } @@ -171,6 +202,9 @@ impl StreamCommitGate { StreamCommitPolicy::FirstGeminiSemanticEvent { max_bytes, .. } => { (max_bytes, self.gemini.observe(chunk, max_bytes)) } + StreamCommitPolicy::FirstSseSemanticEvent { max_bytes, .. } => { + (max_bytes, self.generic.observe(chunk, max_bytes)) + } StreamCommitPolicy::ResponseHeaders | StreamCommitPolicy::FirstClassifiedBody => { return StreamPrecommitObservation::Pending; } @@ -217,6 +251,152 @@ enum SemanticSseObservation { Error { status_code: u16, body_json: Value }, } +#[derive(Debug, Default)] +struct GenericSsePrecommitInspector { + buffered: Vec, +} + +impl GenericSsePrecommitInspector { + fn observe(&mut self, chunk: &[u8], max_bytes: usize) -> SemanticSseObservation { + let remaining = max_bytes.saturating_sub(self.buffered.len()); + self.buffered + .extend_from_slice(&chunk[..chunk.len().min(remaining)]); + while let Some((record_end, separator_len)) = find_sse_record_boundary(&self.buffered) { + let record = self.buffered[..record_end].to_vec(); + self.buffered.drain(..record_end + separator_len); + match classify_generic_sse_record(&record) { + SemanticSseObservation::Pending => {} + observation => return observation, + } + } + if chunk.len() > remaining { + SemanticSseObservation::SemanticEvent + } else { + SemanticSseObservation::Pending + } + } +} + +fn classify_generic_sse_record(record: &[u8]) -> SemanticSseObservation { + let Ok(record) = std::str::from_utf8(record) else { + return SemanticSseObservation::SemanticEvent; + }; + let normalized = record.replace("\r\n", "\n").replace('\r', "\n"); + let event_type = normalized + .lines() + .find_map(|line| line.strip_prefix("event:").map(str::trim)); + let data = normalized + .lines() + .filter_map(|line| line.strip_prefix("data:").map(str::trim_start)) + .collect::>() + .join("\n"); + if data.trim().is_empty() || matches!(event_type, Some("ping" | "heartbeat" | "keepalive")) { + return SemanticSseObservation::Pending; + } + if data.trim() == "[DONE]" { + return SemanticSseObservation::SemanticEvent; + } + let Ok(body_json) = serde_json::from_str::(data.trim()) else { + return SemanticSseObservation::SemanticEvent; + }; + let payload_type = body_json.get("type").and_then(Value::as_str).or(event_type); + if payload_type.is_some_and(is_anthropic_semantic_event_type) { + return classify_anthropic_sse_record(record.as_bytes()); + } + let error = body_json + .get("error") + .filter(|value| !value.is_null()) + .or_else(|| { + body_json + .pointer("/response/error") + .filter(|value| !value.is_null()) + }); + if error.is_some() + || matches!(payload_type, Some("error" | "response.failed")) + || body_json.get("status").and_then(Value::as_str) == Some("failed") + { + let failure = error + .map(|error| serde_json::json!({ "error": error })) + .unwrap_or_else(|| body_json.clone()); + return SemanticSseObservation::Error { + status_code: crate::execution_runtime::submission::resolve_local_sync_error_status_code( + 200, &failure, + ), + body_json: failure, + }; + } + if matches!( + payload_type, + Some("ping" | "response.created" | "response.in_progress" | "response.queued") + ) { + return SemanticSseObservation::Pending; + } + if payload_type == Some("response.output_item.added") + && matches!( + body_json.pointer("/item/type").and_then(Value::as_str), + Some("message" | "reasoning") + ) + && body_json + .pointer("/item/content") + .and_then(Value::as_array) + .is_none_or(Vec::is_empty) + && body_json + .pointer("/item/summary") + .and_then(Value::as_array) + .is_none_or(Vec::is_empty) + { + return SemanticSseObservation::Pending; + } + if matches!( + payload_type, + Some("response.content_part.added" | "response.reasoning_summary_part.added") + ) && matches!( + body_json.pointer("/part/type").and_then(Value::as_str), + Some("output_text" | "summary_text" | "refusal") + ) && !body_json + .pointer("/part/text") + .is_some_and(value_has_semantic_content) + && !body_json + .pointer("/part/refusal") + .is_some_and(value_has_semantic_content) + { + return SemanticSseObservation::Pending; + } + if let Some(choices) = body_json.get("choices").and_then(Value::as_array) { + let semantic = choices.iter().any(|choice| { + choice + .get("finish_reason") + .is_some_and(|value| !value.is_null()) + || choice.get("text").is_some_and(value_has_semantic_content) + || choice + .get("delta") + .or_else(|| choice.get("message")) + .and_then(Value::as_object) + .is_some_and(|delta| { + delta.iter().any(|(name, value)| { + name != "role" && value_has_semantic_content(value) + }) + }) + }); + return if semantic { + SemanticSseObservation::SemanticEvent + } else { + SemanticSseObservation::Pending + }; + } + SemanticSseObservation::SemanticEvent +} + +fn value_has_semantic_content(value: &Value) -> bool { + match value { + Value::Null => false, + Value::String(text) => !text.is_empty(), + Value::Array(values) => !values.is_empty(), + Value::Object(values) => !values.is_empty(), + _ => true, + } +} + #[derive(Debug, Default)] struct AnthropicSsePrecommitInspector { buffered: Vec, @@ -355,7 +535,30 @@ fn classify_anthropic_sse_record(record: &[u8]) -> SemanticSseObservation { (None, Some(payload_type)) => Some(payload_type), _ => None, }; - if semantic_type.is_some_and(is_anthropic_semantic_event_type) { + let setup_only = match semantic_type { + Some("message_start") => body_json + .pointer("/message/content") + .and_then(Value::as_array) + .is_none_or(Vec::is_empty), + Some("content_block_start") => { + let block_type = body_json + .pointer("/content_block/type") + .and_then(Value::as_str); + matches!(block_type, Some("text" | "thinking")) + && !body_json + .pointer("/content_block/text") + .is_some_and(value_has_semantic_content) + && !body_json + .pointer("/content_block/thinking") + .is_some_and(value_has_semantic_content) + } + Some("content_block_stop") => true, + Some("message_delta") => body_json + .pointer("/delta/stop_reason") + .is_none_or(Value::is_null), + _ => false, + }; + if !setup_only && semantic_type.is_some_and(is_anthropic_semantic_event_type) { SemanticSseObservation::SemanticEvent } else { SemanticSseObservation::Pending @@ -507,6 +710,89 @@ pub(super) fn anthropic_error_status_code(body_json: &Value) -> u16 { #[cfg(test)] mod tests { + #[test] + fn image_streams_only_prefetch_when_explicitly_requested() { + for force_prefetch in [false, true] { + let policy = super::StreamCommitPolicy::for_response( + true, + Some("text/event-stream"), + "openai:image", + "openai:image", + false, + false, + force_prefetch, + ); + assert_eq!(policy.commits_on_response_headers(), !force_prefetch); + assert!(!policy.requires_bounded_frame_wait()); + } + } + + #[test] + fn generic_sse_waits_through_setup_and_classifies_fragmented_errors() { + let setup = b"event: response.created\ndata: {\"type\":\"response.created\"}\n\n"; + let failure = b"event: response.failed\ndata: {\"type\":\"response.failed\",\"response\":{\"error\":{\"type\":\"server_error\",\"message\":\"capacity exhausted\"}}}\n\n"; + for split in 1..failure.len() { + let policy = super::StreamCommitPolicy::FirstSseSemanticEvent { + max_bytes: 4096, + max_wait: std::time::Duration::from_secs(1), + }; + let mut gate = super::StreamCommitGate::new(policy); + assert_eq!( + gate.observe_provider_bytes(setup), + super::StreamPrecommitObservation::Pending + ); + for control in [ + b"event: ping\ndata: keepalive\n\n".as_slice(), + b"data: {\"type\":\"response.output_item.added\",\"item\":{\"type\":\"reasoning\",\"summary\":[]}}\n\n".as_slice(), + b"data: {\"type\":\"response.reasoning_summary_part.added\",\"part\":{\"type\":\"summary_text\",\"text\":\"\"}}\n\n".as_slice(), + ] { + assert_eq!(gate.observe_provider_bytes(control), super::StreamPrecommitObservation::Pending); + } + assert_eq!( + gate.observe_provider_bytes(&failure[..split]), + super::StreamPrecommitObservation::Pending + ); + assert!(matches!( + gate.observe_provider_bytes(&failure[split..]), + super::StreamPrecommitObservation::UpstreamError { .. } + )); + } + } + + #[test] + fn generic_sse_commits_on_content_or_tool_call_but_not_role() { + for output in [ + "data: {\"choices\":[{\"delta\":{\"content\":\"hello\"}}]}\n\n", + "data: {\"choices\":[{\"delta\":{\"tool_calls\":[{\"id\":\"call-1\"}]}}]}\n\n", + ] { + let mut gate = + super::StreamCommitGate::new(super::StreamCommitPolicy::FirstSseSemanticEvent { + max_bytes: 4096, + max_wait: std::time::Duration::from_secs(1), + }); + assert_eq!(gate.observe_provider_bytes(b"data: {\"choices\":[{\"delta\":{\"role\":\"assistant\",\"content\":\"\"}}]}\n\n"), super::StreamPrecommitObservation::Pending); + assert_eq!( + gate.observe_provider_bytes(output.as_bytes()), + super::StreamPrecommitObservation::Commit + ); + assert_eq!( + gate.observe_provider_bytes(b"data: {\"error\":{\"message\":\"late error\"}}\n\n"), + super::StreamPrecommitObservation::Commit + ); + } + } + + #[test] + fn native_anthropic_setup_does_not_hide_an_early_error() { + let mut gate = + super::StreamCommitGate::new(super::StreamCommitPolicy::FirstAnthropicSemanticEvent { + max_bytes: 4096, + max_wait: std::time::Duration::from_secs(1), + }); + assert_eq!(gate.observe_provider_bytes(b"event: message_start\ndata: {\"type\":\"message_start\",\"message\":{\"content\":[]}}\n\n"), super::StreamPrecommitObservation::Pending); + assert_eq!(gate.observe_provider_bytes(b"event: content_block_start\ndata: {\"type\":\"content_block_start\",\"content_block\":{\"type\":\"text\",\"text\":\"\"}}\n\n"), super::StreamPrecommitObservation::Pending); + assert!(matches!(gate.observe_provider_bytes(b"event: error\ndata: {\"type\":\"error\",\"error\":{\"type\":\"overloaded_error\"}}\n\n"), super::StreamPrecommitObservation::UpstreamError { status_code: 529, .. })); + } use std::time::Duration; use super::{ @@ -553,7 +839,7 @@ mod tests { false, false, ) - .commits_on_response_headers()); + .requires_bounded_frame_wait()); assert!(StreamCommitPolicy::for_response( true, Some("text/event-stream"), @@ -563,7 +849,7 @@ mod tests { true, false, ) - .commits_on_response_headers()); + .requires_bounded_frame_wait()); } #[test] @@ -745,8 +1031,8 @@ mod tests { let mut gate = StreamCommitGate::new(native_anthropic_policy()); let observation = gate.observe_provider_bytes( concat!( - "event: message_start\n", - "data: {\"type\":\"message_start\",\"message\":{}}\n\n", + "event: content_block_delta\n", + "data: {\"type\":\"content_block_delta\",\"delta\":{\"type\":\"text_delta\",\"text\":\"hello\"}}\n\n", "event: error\n", "data: {\"type\":\"error\",\"error\":{\"type\":\"overloaded_error\"}}\n\n", ) diff --git a/apps/aether-gateway/src/execution_runtime/stream/execution.rs b/apps/aether-gateway/src/execution_runtime/stream/execution.rs index da0d04f08..31e5b22b3 100644 --- a/apps/aether-gateway/src/execution_runtime/stream/execution.rs +++ b/apps/aether-gateway/src/execution_runtime/stream/execution.rs @@ -1414,6 +1414,7 @@ async fn prefetch_direct_anthropic_stream_failure( } let mut gate = StreamCommitGate::new(policy); + let mut semantic_commit_observed = false; let precommit_started_at = Instant::now(); let max_wait = policy.max_precommit_wait()?; let mut observed_first_body = execution @@ -1464,7 +1465,10 @@ async fn prefetch_direct_anthropic_stream_failure( execution.prefetched_body.push_back(Ok(chunk.clone())); match gate.observe_provider_bytes(&chunk) { StreamPrecommitObservation::Pending => {} - StreamPrecommitObservation::Commit => break, + StreamPrecommitObservation::Commit => { + semantic_commit_observed = true; + break; + } StreamPrecommitObservation::UpstreamError { status_code, body_json, @@ -1478,7 +1482,7 @@ async fn prefetch_direct_anthropic_stream_failure( } } } - execution.stream_precommit_committed = !gate.is_uncommitted(); + execution.stream_precommit_committed = semantic_commit_observed; None } @@ -5696,6 +5700,30 @@ fn should_probe_success_failover_before_stream(headers: &BTreeMap, + elapsed_ms: u64, +) { + let finished_at = current_request_candidate_unix_ms(); + record_local_request_candidate_status( + state, + plan, + report_context, + SchedulerRequestCandidateStatusUpdate { + status: RequestCandidateStatus::Failed, + status_code: Some(200), + error_type: Some("success_failover_pattern".to_string()), + error_message: Some("HTTP 200 response matched a precommit failover rule".to_string()), + latency_ms: Some(elapsed_ms), + started_at_unix_ms: None, + finished_at_unix_ms: Some(finished_at), + }, + ) + .await; +} + async fn probe_local_stream_success_failover_text( buffered_frames: &mut VecDeque, lines: &mut FramedRead, @@ -6316,6 +6344,23 @@ async fn execute_stream_from_frame_stream_with_retry_scope( report_context.as_ref(), ) .is_some_and(|policy| policy.cyber_continue_failover); + let prefetch_failover_policy = + crate::orchestration::resolve_local_failover_policy(state, &plan, report_context.as_ref()) + .await; + let prefetch_success_patterns = prefetch_failover_policy + .routing_rules + .success_failover_patterns + .iter() + .map(|rule| (&rule.pattern, &rule.status_codes)) + .chain( + prefetch_failover_policy + .success_failover_patterns + .iter() + .map(|rule| (&rule.pattern, &rule.status_codes)), + ) + .filter(|(_, status_codes)| status_codes.is_empty() || status_codes.contains(&200)) + .filter_map(|(pattern, _)| regex::Regex::new(pattern.trim()).ok()) + .collect::>(); let stream_commit_policy = StreamCommitPolicy::for_response( direct_stream_finalize_kind.is_some(), upstream_content_type, @@ -6323,15 +6368,24 @@ async fn execute_stream_from_frame_stream_with_retry_scope( plan.client_api_format.as_str(), private_stream_normalizer.is_some(), local_stream_rewriter.is_some(), - prefetch_for_cyber_failover, - ); - let reuse_committed_precommit = - stream_precommit_committed && stream_commit_policy.is_native_anthropic(); + prefetch_for_cyber_failover || !prefetch_success_patterns.is_empty(), + ) + .with_precommit_wait(Duration::from_millis( + plan.timeouts + .as_ref() + .and_then(|timeouts| timeouts.first_byte_ms) + .unwrap_or(30_000) + .max(1), + )); + let reuse_committed_precommit = stream_precommit_committed + && stream_commit_policy.is_native_anthropic() + && prefetch_success_patterns.is_empty(); let skip_direct_finalize_prefetch = stream_commit_policy.commits_on_response_headers() || reuse_committed_precommit; let limit_direct_finalize_prefetch = should_limit_direct_finalize_prefetch(plan_kind, local_stream_rewriter.is_some()) - || stream_commit_policy.requires_bounded_frame_wait(); + || stream_commit_policy.requires_bounded_frame_wait() + || !prefetch_success_patterns.is_empty(); let mut stream_commit_gate = StreamCommitGate::new(stream_commit_policy); let mut prefetch_client_completion_tracker = ClientVisibleStreamCompletionTracker::default(); let mut prefetched_client_visible_stream_completed = false; @@ -6383,7 +6437,8 @@ async fn execute_stream_from_frame_stream_with_retry_scope( .max_precommit_wait() .map(|max_wait| max_wait.saturating_sub(precommit_started_at.elapsed())) .unwrap_or(REWRITTEN_STREAM_PREFETCH_TIMEOUT); - if prefetch_timeout.is_zero() { + if prefetch_timeout.is_zero() && !stream_commit_policy.requires_bounded_frame_wait() + { stream_commit_gate.commit(); debug!( event_name = "execution_runtime_stream_prefetch_limited", @@ -6414,6 +6469,29 @@ async fn execute_stream_from_frame_stream_with_retry_scope( { Ok(result) => result, Err(_) => { + if stream_commit_policy.requires_bounded_frame_wait() { + let failure = build_stream_transport_failure_report( + "first_byte_timeout", "Upstream did not produce a semantic event before the first byte deadline", 504, + ); + return handle_prefetch_stream_failure( + state, + trace_id, + decision, + &plan, + report_context, + request_id, + candidate_id, + report_kind, + headers, + prefetched_usage_telemetry.clone(), + &provider_prefetched_body, + candidate_started_unix_secs, + stream_elapsed_ms_since(stream_started_at), + failure, + retry_scope_out.as_deref_mut(), + ) + .await; + } stream_commit_gate.commit(); debug!( event_name = "execution_runtime_stream_prefetch_limited", @@ -6468,10 +6546,11 @@ async fn execute_stream_from_frame_stream_with_retry_scope( .await; } }) else { - if stream_commit_policy.is_native_anthropic() && stream_commit_gate.is_uncommitted() + if stream_commit_policy.requires_bounded_frame_wait() + && stream_commit_gate.is_uncommitted() { let error_body_json = anthropic_premature_eof_error_body( - "upstream Anthropic stream ended before the first semantic event", + "upstream stream ended before the first semantic event", ); let error_status_code = anthropic_error_status_code(&error_body_json); return handle_prefetch_provider_private_stream_error( @@ -6511,6 +6590,12 @@ async fn execute_stream_from_frame_stream_with_retry_scope( "stream_first_data", stream_elapsed_ms_at(stream_started_at, frame_observed_at), ); + state.usage_runtime.record_stream_started( + state.usage_lifecycle_data_state().as_ref(), + &lifecycle_seed, + status_code, + prefetched_usage_telemetry.as_ref(), + ); } let mut chunk = match decode_stream_data_chunk(chunk_b64.as_deref(), text.as_deref()) { @@ -6576,7 +6661,29 @@ async fn execute_stream_from_frame_stream_with_retry_scope( &mut prefetched_inspection_body_truncated, ); - let anthropic_commit_ready = + if !prefetch_success_patterns.is_empty() + && crate::orchestration::attempt_identity_from_report_context( + report_context.as_ref(), + ) + .is_some() + { + let response_text = String::from_utf8_lossy(&prefetched_inspection_body); + if prefetch_success_patterns.iter().any(|pattern| pattern.is_match(&response_text)) + && crate::orchestration::classify_local_failover( + &prefetch_failover_policy, + crate::orchestration::LocalFailoverInput::new(status_code, Some(&response_text)), + ) == crate::orchestration::LocalFailoverClassification::RetrySuccessPattern + { + record_prefetch_success_failover(state, &plan, report_context.as_ref(), stream_elapsed_ms_since(stream_started_at)).await; + if let Some(retry_scope) = retry_scope_out.as_deref_mut() { + *retry_scope = AiAttemptRetryScope::Candidate; + } + warn!(event_name = "local_stream_candidate_retry_scheduled", log_type = "event", trace_id, request_id, status_code, "gateway retrying after a precommit success pattern match"); + return Ok(None); + } + } + + let semantic_commit_ready = match stream_commit_gate.observe_provider_bytes(&chunk) { StreamPrecommitObservation::Pending => false, StreamPrecommitObservation::Commit => true, @@ -6584,6 +6691,14 @@ async fn execute_stream_from_frame_stream_with_retry_scope( status_code: error_status_code, body_json: error_body_json, } => { + let error_status_code = if plan + .provider_api_format + .eq_ignore_ascii_case("claude:messages") + { + anthropic_error_status_code(&error_body_json) + } else { + error_status_code + }; return handle_prefetch_provider_private_stream_error( state, trace_id, @@ -6606,7 +6721,7 @@ async fn execute_stream_from_frame_stream_with_retry_scope( } }; - if !anthropic_commit_ready { + if !semantic_commit_ready || private_stream_normalizer.is_some() { if let Some(error_body_json) = extract_provider_private_stream_error_body( report_context.as_ref(), &prefetched_inspection_body, @@ -6638,9 +6753,7 @@ async fn execute_stream_from_frame_stream_with_retry_scope( } } - let inspection = if stream_commit_policy.is_native_anthropic() - || stream_commit_policy.is_gemini() - { + let inspection = if stream_commit_policy.requires_bounded_frame_wait() { StreamPrefetchInspection::NeedMore } else { inspect_prefetched_stream_body( @@ -6650,55 +6763,30 @@ async fn execute_stream_from_frame_stream_with_retry_scope( }; match inspection { StreamPrefetchInspection::EmbeddedError(body_json) => { - debug!( - event_name = "execution_runtime_stream_prefetch_embedded_error_detected", - log_type = "debug", - trace_id = %trace_id, - request_id = %request_id_for_log, - candidate_id = ?candidate_id, - plan_kind, - report_kind, - provider_name, - endpoint_id = %plan.endpoint_id, - key_id = %plan.key_id, - model_name, - candidate_index = candidate_index.as_str(), - provider_prefetched_body_bytes = provider_prefetched_body.len(), - "gateway detected embedded error while prefetching execution runtime stream" - ); - let request_diagnostics = current_request_diagnostics(); - let terminal_report_context = report_context_with_request_diagnostics( - report_context, - request_diagnostics.as_ref(), - stream_started_at, - prefetched_usage_telemetry.as_ref(), - ); - let payload = build_stream_sync_payload( - trace_id, - report_kind.clone(), - terminal_report_context, + let error_status_code = resolve_provider_stream_error_status_code( + plan.provider_api_format.as_str(), status_code, - headers, - Some(body_json), - None, - prefetched_usage_telemetry.clone(), + &body_json, ); - record_sync_terminal_usage_with_handoff( + return handle_prefetch_provider_private_stream_error( state, + trace_id, + decision, &plan, - payload.report_context.as_ref(), - &payload, + report_context, + request_id, + candidate_id, + report_kind, + headers, + prefetched_usage_telemetry.clone(), + &provider_prefetched_body, + status_code, + error_status_code, + body_json, + retry_scope_out.as_deref_mut(), + retry_fallback_out.as_deref_mut(), ) .await; - let response = submit_local_core_error_or_sync_finalize( - state, trace_id, decision, payload, - ) - .await?; - return Ok(Some(attach_control_metadata_headers( - response, - Some(request_id), - candidate_id, - )?)); } StreamPrefetchInspection::NeedMore => {} StreamPrefetchInspection::NonError => {} @@ -6842,8 +6930,12 @@ async fn execute_stream_from_frame_stream_with_retry_scope( prefetched_chunks.push(Bytes::from(rewritten_chunk)); } - if anthropic_commit_ready + if semantic_commit_ready || (matches!(inspection, StreamPrefetchInspection::NonError) + && (prefetch_success_patterns.is_empty() + || response_headers_indicate_sse(&upstream_headers) + || parse_prefetched_sync_json_body(&prefetched_inspection_body) + .is_some()) && (!prefetch_for_cyber_failover || prefetched_openai_responses_body_has_output_boundary( &prefetched_inspection_body, @@ -6858,11 +6950,11 @@ async fn execute_stream_from_frame_stream_with_retry_scope( prefetched_telemetry = Some(frame_telemetry); } StreamFramePayload::Eof { summary } => { - if stream_commit_policy.is_native_anthropic() + if stream_commit_policy.requires_bounded_frame_wait() && stream_commit_gate.is_uncommitted() { let error_body_json = anthropic_premature_eof_error_body( - "upstream Anthropic stream ended before the first semantic event", + "upstream stream ended before the first semantic event", ); let error_status_code = anthropic_error_status_code(&error_body_json); return handle_prefetch_provider_private_stream_error( @@ -7003,7 +7095,7 @@ async fn execute_stream_from_frame_stream_with_retry_scope( let native_anthropic_stream_for_report = stream_commit_policy.is_native_anthropic(); let plan_for_report = plan; let emit_passthrough_sse_terminal_error = (skip_direct_finalize_prefetch - || stream_commit_policy.is_native_anthropic() + || stream_commit_policy.requires_bounded_frame_wait() || normalized_declared_stream_headers) && (response_headers_indicate_sse(&upstream_headers) || normalized_declared_stream_headers) && !is_openai_image_stream_for_report; @@ -8921,6 +9013,182 @@ mod tests { } } + async fn execute_generic_sse_precommit( + chunks: Vec<&str>, + routing_policy: Value, + provider_config: Option, + stall: bool, + ) -> Option> { + execute_generic_stream_precommit( + chunks, + routing_policy, + provider_config, + stall, + "text/event-stream", + ) + .await + } + + async fn execute_generic_stream_precommit( + chunks: Vec<&str>, + routing_policy: Value, + provider_config: Option, + stall: bool, + content_type: &str, + ) -> Option> { + let request_id = format!("generic-precommit-{}", uuid::Uuid::new_v4()); + let mut plan = native_anthropic_stream_plan(&request_id); + plan.provider_api_format = "openai:responses".to_string(); + plan.client_api_format = "openai:responses".to_string(); + plan.timeouts = Some(ExecutionTimeouts { + first_byte_ms: Some(20), + ..Default::default() + }); + let provider_catalog = provider_catalog_for_plan(&plan, provider_config); + let data_state = crate::data::GatewayDataState::with_provider_transport_reader_for_tests( + Arc::new(provider_catalog), + DEVELOPMENT_ENCRYPTION_KEY, + ); + let state = AppState::new() + .unwrap() + .with_data_state_for_tests(data_state); + let chunks = chunks.into_iter().map(str::to_string).collect::>(); + let content_type = content_type.to_string(); + let frames = stream! { + yield Ok::(ndjson_frame(StreamFrame { + frame_type: StreamFrameType::Headers, + payload: StreamFramePayload::Headers { + status_code: 200, + headers: BTreeMap::from([("content-type".to_string(), content_type)]), + response_observation: None, + }, + })); + for chunk in chunks { + yield Ok(ndjson_frame(StreamFrame { + frame_type: StreamFrameType::Data, + payload: StreamFramePayload::Data { text: Some(chunk), chunk_b64: None }, + })); + } + if stall { std::future::pending::<()>().await; } + yield Ok(ndjson_frame(StreamFrame { + frame_type: StreamFrameType::Eof, + payload: StreamFramePayload::Eof { summary: None }, + })); + } + .boxed(); + let mut scope = AiAttemptRetryScope::Provider; + execute_stream_from_frame_stream_with_retry_scope( + &state, + plan, + "trace-generic-precommit", + &test_decision(), + "openai_responses_stream", + Some("openai_responses_stream_success".to_string()), + Some(json!({ + "request_id": request_id, "candidate_id": format!("candidate-{request_id}"), + "candidate_index": 0, "retry_index": 0, + "provider_api_format": "openai:responses", "client_api_format": "openai:responses", + "routing_execution_policy": routing_policy, + })), + crate::clock::current_unix_ms(), + Instant::now(), + RequestStageTrace::from_env(), + true, + frames, + false, + None, + Some(&mut scope), + None, + None, + ) + .await + .unwrap() + } + + #[tokio::test] + async fn generic_stream_success_regex_matches_fragmented_plain_body() { + assert!(execute_generic_stream_precommit( + vec!["upstream CAPACITY ", "exhausted"], + json!({"failover_rules": {"success_failover_patterns": [{"pattern": "(?i)capacity.*exhausted"}]}}), + None, + false, + "text/plain", + ).await.is_none()); + } + + #[tokio::test] + async fn generic_stream_200_json_error_obeys_global_stop_rules() { + for stop in [false, true] { + let response = execute_generic_stream_precommit( + vec![r#"{"error":{"type":"server_error","message":"do not retry"}}"#], + if stop { json!({"failover_rules": {"error_stop_patterns": [{"pattern": "do not retry"}]}}) } else { json!({}) }, + None, + false, + "application/json", + ).await; + assert_eq!(response.is_some(), stop); + if let Some(response) = response { + assert!(response.status().is_server_error()); + } + } + } + + #[tokio::test] + async fn generic_sse_200_setup_then_error_retries_before_client_output() { + let response = execute_generic_sse_precommit(vec![ + "event: response.created\ndata: {\"type\":\"response.created\"}\n\n", + "event: response.failed\ndata: {\"type\":\"response.failed\",\"response\":{\"error\":{\"type\":\"server_error\",\"message\":\"capacity exhausted\"}}}\n\n", + ], json!({}), None, false).await; + assert!(response.is_none()); + } + + #[tokio::test] + async fn generic_sse_global_stop_rule_overrides_retryable_embedded_error() { + let response = execute_generic_sse_precommit(vec![ + "event: response.failed\ndata: {\"type\":\"response.failed\",\"response\":{\"error\":{\"type\":\"server_error\",\"message\":\"capacity exhausted\"}}}\n\n", + ], json!({ "failover_rules": { "error_stop_patterns": [{ "pattern": "capacity" }] } }), None, false).await; + let response = response.expect("global stop must return a terminal response"); + assert!(response.status().is_server_error()); + to_bytes(response.into_body(), usize::MAX).await.unwrap(); + } + + #[tokio::test] + async fn generic_sse_success_regex_applies_to_global_and_provider_rules() { + let rule = json!({ "success_failover_patterns": [{ "pattern": "(?i)CAPACITY" }] }); + for (global, provider) in [ + (json!({ "failover_rules": rule.clone() }), None), + (json!({}), Some(json!({ "failover_rules": rule }))), + ] { + let response = execute_generic_sse_precommit(vec![ + "event: response.output_text.delta\ndata: {\"type\":\"response.output_text.delta\",\"delta\":\"capacity exhausted\"}\n\n", + ], global, provider, false).await; + assert!(response.is_none()); + } + } + + #[tokio::test] + async fn generic_sse_late_error_does_not_replay_committed_content() { + let response = execute_generic_sse_precommit(vec![ + "event: response.output_text.delta\ndata: {\"type\":\"response.output_text.delta\",\"delta\":\"hello\"}\n\n", + "event: response.failed\ndata: {\"type\":\"response.failed\",\"response\":{\"error\":{\"type\":\"server_error\",\"message\":\"late failure\"}}}\n\n", + ], json!({}), None, false).await.expect("committed stream must not retry"); + assert_eq!(response.status(), axum::http::StatusCode::OK); + let body = to_bytes(response.into_body(), usize::MAX).await.unwrap(); + assert!(String::from_utf8_lossy(&body).contains("hello")); + } + + #[tokio::test] + async fn generic_sse_setup_timeout_always_retries() { + let response = execute_generic_sse_precommit( + vec!["event: response.created\ndata: {\"type\":\"response.created\"}\n\n"], + json!({}), + None, + true, + ) + .await; + assert!(response.is_none()); + } + fn native_anthropic_stream_plan(request_id: &str) -> ExecutionPlan { ExecutionPlan { request_id: request_id.to_string(), @@ -10752,7 +11020,7 @@ mod tests { .await .expect("default Codex cyber policy handling should return the provider error"); - assert_eq!(response.status().as_u16(), 200); + assert_eq!(response.status(), axum::http::StatusCode::BAD_REQUEST); } #[tokio::test] @@ -12071,6 +12339,8 @@ mod tests { let raw = concat!( "event: message_start\n", "data: {\"type\":\"message_start\",\"message\":{}}\n\n", + "event: content_block_delta\n", + "data: {\"type\":\"content_block_delta\",\"index\":0,\"delta\":{\"type\":\"text_delta\",\"text\":\"hello\"}}\n\n", "event: error\n", "data: {\"type\":\"error\",\"error\":{\"type\":\"overloaded_error\",\"message\":\"late\"}}\n\n", ); @@ -12094,6 +12364,8 @@ mod tests { let message_start = concat!( "event: message_start\n", "data: {\"type\":\"message_start\",\"message\":{}}\n\n", + "event: content_block_delta\n", + "data: {\"type\":\"content_block_delta\",\"index\":0,\"delta\":{\"type\":\"text_delta\",\"text\":\"hello\"}}\n\n", ); let original_error = "upstream disconnected after message_start"; let outcome = execute_native_anthropic_prefetch_stream_with_terminal_error( @@ -12123,6 +12395,8 @@ mod tests { let message_start = concat!( "event: message_start\n", "data: {\"type\":\"message_start\",\"message\":{}}\n\n", + "event: content_block_delta\n", + "data: {\"type\":\"content_block_delta\",\"index\":0,\"delta\":{\"type\":\"text_delta\",\"text\":\"hello\"}}\n\n", ); let outcome = execute_native_anthropic_prefetch_stream( "req-anthropic-postcommit-eof", @@ -12148,6 +12422,8 @@ mod tests { let message_start = concat!( "event: message_start\n", "data: {\"type\":\"message_start\",\"message\":{}}\n\n", + "event: content_block_delta\n", + "data: {\"type\":\"content_block_delta\",\"index\":0,\"delta\":{\"type\":\"text_delta\",\"text\":\"hello\"}}\n\n", ); let done = "data: [DONE]\n\n"; let outcome = execute_native_anthropic_prefetch_stream( @@ -12156,7 +12432,7 @@ mod tests { ) .await; let AiAttemptExecutionOutcome::Responded(response) = outcome else { - panic!("message_start should commit the selected candidate") + panic!("text output should commit the selected candidate") }; let body = to_bytes(response.into_body(), usize::MAX) .await @@ -12748,8 +13024,8 @@ mod tests { } #[test] - fn skips_prefetch_for_event_streams_even_when_cross_format_or_rewritten() { - assert!(should_skip_direct_finalize_prefetch( + fn keeps_prefetch_for_event_streams_even_when_cross_format_or_rewritten() { + assert!(!should_skip_direct_finalize_prefetch( Some("claude_cli_sync_finalize"), Some("text/event-stream"), "openai:chat", @@ -14383,21 +14659,21 @@ mod tests { ) .with_execution_runtime_candidate(true); - let response = execute_execution_runtime_stream( - &state, - plan, - "trace-live-stream-first-event", - &decision, - "openai_chat_stream", - None, - Some(json!({ - "provider_api_format": "openai:chat", - "client_api_format": "openai:chat", - })), - ) - .await - .expect("execution should succeed") - .expect("execution should return a client response"); + let execution_task = tokio::spawn(async move { + execute_execution_runtime_stream( + &state, + plan, + "trace-live-stream-first-event", + &decision, + "openai_chat_stream", + None, + Some(json!({ + "provider_api_format": "openai:chat", + "client_api_format": "openai:chat", + })), + ) + .await + }); first_event_seen.notified().await; let deadline = tokio::time::Instant::now() + Duration::from_secs(15); @@ -14418,9 +14694,16 @@ mod tests { tokio::time::sleep(Duration::from_millis(10)).await; }; assert!(first_event_usage.first_byte_time_ms.is_some()); + assert!(!execution_task.is_finished()); release_text.notify_one(); text_seen.notified().await; + let response = tokio::time::timeout(Duration::from_secs(1), execution_task) + .await + .expect("semantic text should commit the response") + .expect("execution task should complete") + .expect("execution should succeed") + .expect("execution should return a client response"); release_terminal.notify_one(); let body = to_bytes(response.into_body(), usize::MAX) diff --git a/apps/aether-gateway/src/execution_runtime/submission.rs b/apps/aether-gateway/src/execution_runtime/submission.rs index 9e7186058..3ee844002 100644 --- a/apps/aether-gateway/src/execution_runtime/submission.rs +++ b/apps/aether-gateway/src/execution_runtime/submission.rs @@ -529,7 +529,13 @@ fn classify_local_sync_error_kind( { return LocalCoreSyncErrorKind::Overloaded; } - if (500..600).contains(&status_code) { + if (500..600).contains(&status_code) + || raw_type.is_some_and(|value| { + ["server_error", "internal_error", "api_error"] + .iter() + .any(|kind| value.trim().eq_ignore_ascii_case(kind)) + }) + { return LocalCoreSyncErrorKind::ServerError; } LocalCoreSyncErrorKind::InvalidRequest @@ -676,6 +682,13 @@ pub(crate) async fn submit_local_core_error_or_sync_finalize( #[cfg(test)] mod tests { + #[test] + fn success_http_status_does_not_misclassify_explicit_server_errors_as_bad_requests() { + for error_type in ["server_error", "internal_error", "api_error"] { + let body = serde_json::json!({ "error": { "type": error_type, "message": "failed" } }); + assert_eq!(super::resolve_local_sync_error_status_code(200, &body), 500); + } + } use axum::body::to_bytes; use serde_json::json; diff --git a/apps/aether-gateway/src/executor/candidate_loop.rs b/apps/aether-gateway/src/executor/candidate_loop.rs index 86495df4d..d0743ec87 100644 --- a/apps/aether-gateway/src/executor/candidate_loop.rs +++ b/apps/aether-gateway/src/executor/candidate_loop.rs @@ -655,6 +655,73 @@ struct ProviderTransferState { struct ProviderTransferStateTracker { by_provider: BTreeMap, exhausted_provider_ids: BTreeSet, + global: GlobalTransferState, +} + +#[derive(Debug, Default)] +struct GlobalTransferState { + first_attempt_started_at: Option, + last_candidate: Option<(String, String, String)>, + transfer_count: u64, + limits: Option, + exhausted: bool, +} + +impl GlobalTransferState { + fn load_policy(&mut self, report_context: Option<&serde_json::Value>) { + if self.limits.is_none() { + if let Some(policy) = + crate::orchestration::routing_execution_policy_from_report_context(report_context) + { + self.limits = Some(ProviderTransferLimits { + max_transfer_count: policy.max_transfer_count, + max_transfer_timeout_seconds: policy.max_transfer_timeout_seconds, + }); + } + } + } + + fn changes_candidate(&self, plan: &aether_contracts::ExecutionPlan) -> bool { + self.last_candidate + .as_ref() + .is_some_and(|(provider, endpoint, key)| { + provider != &plan.provider_id + || endpoint != &plan.endpoint_id + || key != &plan.key_id + }) + } + + fn record_attempt_started(&mut self, plan: &aether_contracts::ExecutionPlan, now: Instant) { + self.first_attempt_started_at.get_or_insert(now); + if self.changes_candidate(plan) { + self.transfer_count = self.transfer_count.saturating_add(1); + } + self.last_candidate = Some(( + plan.provider_id.clone(), + plan.endpoint_id.clone(), + plan.key_id.clone(), + )); + } + + fn check_before_attempt( + &mut self, + plan: &aether_contracts::ExecutionPlan, + now: Instant, + ) -> Option<(bool, bool)> { + let limits = self.limits?; + let started_at = self.first_attempt_started_at?; + let count_reached = self.changes_candidate(plan) + && limits.max_transfer_count > 0 + && self.transfer_count >= limits.max_transfer_count; + let timeout_reached = limits.max_transfer_timeout_seconds > 0 + && now.saturating_duration_since(started_at) + >= Duration::from_secs(limits.max_transfer_timeout_seconds); + if !count_reached && !timeout_reached { + return None; + } + self.exhausted = true; + Some((count_reached, timeout_reached)) + } } #[derive(Clone, Debug, Default)] @@ -717,6 +784,7 @@ struct ProviderTransferLimitReached { impl ProviderTransferStateTracker { fn record_attempt_started(&mut self, plan: &aether_contracts::ExecutionPlan, now: Instant) { + self.global.record_attempt_started(plan, now); match self.by_provider.entry(plan.provider_id.clone()) { std::collections::btree_map::Entry::Vacant(entry) => { entry.insert(ProviderTransferState { @@ -903,11 +971,42 @@ async fn should_skip_provider_transfer_attempt( where Attempt: AiExecutionAttempt + Send + Sync + 'static, { - let reached = tracker - .state - .lock() - .await - .check_before_attempt(attempt.execution_plan(), Instant::now()); + let owned_report_context = attempt + .report_context_ref() + .is_none() + .then(|| attempt.report_context()) + .flatten(); + let report_context = attempt + .report_context_ref() + .or(owned_report_context.as_ref()); + let mut tracker = tracker.state.lock().await; + tracker.global.load_policy(report_context); + if tracker.global.exhausted { + return true; + } + let now = Instant::now(); + if let Some((count_reached, timeout_reached)) = tracker + .global + .check_before_attempt(attempt.execution_plan(), now) + { + warn!( + event_name = "routing_transfer_limit_reached", + log_type = "event", + trace_id, + plan_kind, + transfer_count = tracker.global.transfer_count, + elapsed_ms = tracker + .global + .first_attempt_started_at + .map(|started| now.saturating_duration_since(started).as_millis() as u64) + .unwrap_or(0), + count_reached, + timeout_reached, + "gateway exhausted the routing strategy transfer budget" + ); + return true; + } + let reached = tracker.check_before_attempt(attempt.execution_plan(), now); let Some(reached) = reached else { return false; }; @@ -2465,6 +2564,93 @@ mod tests { assert_eq!(port.unused.lock().unwrap().as_slice(), ["a-key3-retry0"]); } + #[tokio::test] + async fn routing_transfer_budget_counts_switches_across_providers_not_same_key_retries() { + for (limit, succeeds) in [(1, false), (2, true)] { + let state = AppState::new().unwrap(); + let port = TransferTestPort::new(&state); + let mut attempts = transfer_test_attempts(); + for attempt in &mut attempts { + attempt.report_context["routing_execution_policy"] = + json!({ "max_transfer_count": limit }); + } + let outcome = run_ai_attempt_loop(&port, attempts).await.unwrap(); + assert_eq!( + matches!(outcome, AiAttemptLoopOutcome::Responded(_)), + succeeds + ); + { + let executed = port.executed.lock().unwrap(); + assert_eq!( + &executed[..3], + ["a-key1-retry0", "a-key2-retry0", "a-key2-retry1"] + ); + assert_eq!(executed.len(), if succeeds { 4 } else { 3 }); + } + assert_eq!(port.tracker.state.lock().await.global.transfer_count, limit); + } + } + + #[tokio::test] + async fn dynamic_loop_honors_global_transfer_budget_across_providers() { + let state = AppState::new().unwrap(); + let port = TransferTestPort::new(&state); + let mut attempts = transfer_test_attempts(); + for attempt in &mut attempts { + attempt.report_context["routing_execution_policy"] = json!({ "max_transfer_count": 1 }); + } + let mut source = TransferTestAttemptSource { + attempts: attempts.into(), + skipped_providers: Vec::new(), + }; + let outcome = run_dynamic_attempt_loop( + &port, + &mut source, + "global-budget", + "test", + Duration::from_secs(1), + ) + .await + .unwrap(); + assert!(matches!( + outcome, + LocalExecutionRequestOutcome::Exhausted(_) + )); + assert_eq!( + port.executed.lock().unwrap().as_slice(), + ["a-key1-retry0", "a-key2-retry0", "a-key2-retry1"] + ); + assert_eq!(source.skipped_providers, ["provider-a", "provider-b"]); + } + + #[test] + fn routing_time_budget_is_cumulative_and_zero_is_unlimited() { + let mut global = super::GlobalTransferState::default(); + global.load_policy(Some( + &json!({ "routing_execution_policy": { "max_transfer_timeout_seconds": 60 } }), + )); + let now = tokio::time::Instant::now(); + let plan = test_plan(None); + global.record_attempt_started(&plan, now); + global.record_attempt_started(&plan, now + Duration::from_secs(40)); + assert_eq!(global.transfer_count, 0); + assert_eq!( + global.check_before_attempt(&plan, now + Duration::from_secs(59)), + None + ); + assert_eq!( + global.check_before_attempt(&plan, now + Duration::from_secs(60)), + Some((false, true)) + ); + let mut unlimited = super::GlobalTransferState::default(); + unlimited.load_policy(Some(&json!({ "routing_execution_policy": {} }))); + unlimited.record_attempt_started(&plan, now); + assert_eq!( + unlimited.check_before_attempt(&plan, now + Duration::from_secs(86_400)), + None + ); + } + #[tokio::test] async fn cloned_tracker_preserves_transfer_budget_across_candidate_loops() { let state = AppState::new().expect("state should build"); diff --git a/apps/aether-gateway/src/orchestration/classifier.rs b/apps/aether-gateway/src/orchestration/classifier.rs index 416e8c28a..dd13dab2a 100644 --- a/apps/aether-gateway/src/orchestration/classifier.rs +++ b/apps/aether-gateway/src/orchestration/classifier.rs @@ -300,6 +300,34 @@ pub(crate) fn classify_local_failover( policy: &LocalFailoverPolicy, input: LocalFailoverInput<'_>, ) -> LocalFailoverClassification { + if input.status_code >= 400 + && policy.routing_rules.error_stop_patterns.iter().any(|rule| { + failover_pattern_matches( + &rule.pattern, + &rule.status_codes, + input.response_text, + input.status_code, + ) + }) + { + return LocalFailoverClassification::StopErrorPattern; + } + if input.status_code == 200 + && policy + .routing_rules + .success_failover_patterns + .iter() + .any(|rule| { + failover_pattern_matches( + &rule.pattern, + &rule.status_codes, + input.response_text, + input.status_code, + ) + }) + { + return LocalFailoverClassification::RetrySuccessPattern; + } if policy.stop_status_codes.contains(&input.status_code) { return LocalFailoverClassification::StopStatusCode; } @@ -487,13 +515,27 @@ fn local_failover_regex_rule_matches( response_text: Option<&str>, status_code: u16, ) -> bool { - if !rule.status_codes.is_empty() && !rule.status_codes.contains(&status_code) { + failover_pattern_matches( + &rule.pattern, + &rule.status_codes, + response_text, + status_code, + ) +} + +fn failover_pattern_matches( + pattern: &str, + status_codes: &std::collections::BTreeSet, + response_text: Option<&str>, + status_code: u16, +) -> bool { + if !status_codes.is_empty() && !status_codes.contains(&status_code) { return false; } - let pattern = rule.pattern.trim(); + let pattern = pattern.trim(); if pattern.is_empty() { - return !rule.status_codes.is_empty(); + return !status_codes.is_empty(); } let Some(response_text) = response_text else { @@ -507,6 +549,72 @@ fn local_failover_regex_rule_matches( #[cfg(test)] mod tests { + #[test] + fn routing_rules_precede_provider_rules_and_keep_provider_fallback() { + let policy = super::LocalFailoverPolicy { + routing_rules: aether_routing_core::RoutingFailoverRules { + success_failover_patterns: vec![aether_routing_core::RoutingFailoverRule { + pattern: "(?i)capacity.*exhausted".to_string(), + ..Default::default() + }], + error_stop_patterns: vec![aether_routing_core::RoutingFailoverRule { + pattern: "invalid.*parameter".to_string(), + status_codes: [400].into_iter().collect(), + }], + ..Default::default() + }, + stop_status_codes: [200, 403].into_iter().collect(), + continue_status_codes: [400].into_iter().collect(), + ..Default::default() + }; + for (status, body, expected) in [ + ( + 200, + "CAPACITY exhausted", + super::LocalFailoverClassification::RetrySuccessPattern, + ), + ( + 400, + "invalid request parameter", + super::LocalFailoverClassification::StopErrorPattern, + ), + ( + 400, + "capacity exhausted", + super::LocalFailoverClassification::RetryStatusCode, + ), + ( + 403, + "permission denied", + super::LocalFailoverClassification::StopStatusCode, + ), + ( + 429, + "rate limited", + super::LocalFailoverClassification::RetryUpstreamFailure, + ), + ] { + assert_eq!( + super::classify_local_failover( + &policy, + super::LocalFailoverInput::new(status, Some(body)) + ), + expected + ); + } + } + + #[test] + fn routing_transport_stop_cannot_be_overridden_by_provider() { + let policy = super::LocalFailoverPolicy { + stop_on_transport_errors: true, + ..Default::default() + }; + assert_eq!( + super::classify_local_transport_error(&policy), + super::LocalTransportFailoverClassification::StopTransportError + ); + } use std::collections::BTreeSet; use super::{ diff --git a/apps/aether-gateway/src/orchestration/policy.rs b/apps/aether-gateway/src/orchestration/policy.rs index 0722a11cd..397d9d96f 100644 --- a/apps/aether-gateway/src/orchestration/policy.rs +++ b/apps/aether-gateway/src/orchestration/policy.rs @@ -4,7 +4,7 @@ use aether_contracts::ExecutionPlan; use serde_json::{json, Value}; use tracing::debug; -use aether_routing_core::RoutingExecutionPolicy; +use aether_routing_core::{RoutingExecutionPolicy, RoutingFailoverRules}; use crate::provider_transport::GatewayProviderTransportSnapshot; use crate::AppState; @@ -14,6 +14,7 @@ pub(crate) const ROUTING_EXECUTION_POLICY_REPORT_FIELD: &str = "routing_executio #[derive(Debug, Clone, PartialEq, Eq)] pub(crate) struct LocalFailoverPolicy { + pub(crate) routing_rules: RoutingFailoverRules, pub(crate) max_retries: Option, pub(crate) max_transfer_count: u64, pub(crate) max_transfer_timeout_seconds: u64, @@ -29,6 +30,7 @@ pub(crate) struct LocalFailoverPolicy { impl Default for LocalFailoverPolicy { fn default() -> Self { Self { + routing_rules: RoutingFailoverRules::default(), max_retries: None, max_transfer_count: 0, max_transfer_timeout_seconds: 0, @@ -61,8 +63,10 @@ pub(crate) async fn resolve_local_failover_policy( Ok(Some(transport)) => local_failover_policy_from_transport(&transport), Ok(None) | Err(_) => LocalFailoverPolicy::default(), }; - let cyber_continue_failover = routing_execution_policy_from_report_context(report_context) - .is_some_and(|policy| policy.cyber_continue_failover); + let routing_policy = + routing_execution_policy_from_report_context(report_context).unwrap_or_default(); + let cyber_continue_failover = routing_policy.cyber_continue_failover; + policy.routing_rules = routing_policy.failover_rules; policy.stop_cyber_policy_errors = !cyber_continue_failover; debug!( event_name = "local_failover_policy_loaded", @@ -80,6 +84,8 @@ pub(crate) async fn resolve_local_failover_policy( stop_on_transport_errors = policy.stop_on_transport_errors, success_failover_pattern_count = policy.success_failover_patterns.len(), error_stop_pattern_count = policy.error_stop_patterns.len(), + global_success_pattern_count = policy.routing_rules.success_failover_patterns.len(), + global_stop_pattern_count = policy.routing_rules.error_stop_patterns.len(), cyber_continue_failover, "gateway loaded local failover policy from transport snapshot" ); @@ -122,6 +128,7 @@ pub(crate) fn local_failover_policy_from_transport( }); LocalFailoverPolicy { + routing_rules: RoutingFailoverRules::default(), max_retries, max_transfer_count: provider_config .and_then(|value| value.get("max_transfer_count")) @@ -184,6 +191,10 @@ pub(crate) fn local_failover_policy_from_report_context( .as_object()?; Some(LocalFailoverPolicy { + routing_rules: object + .get("routing_rules") + .and_then(|value| serde_json::from_value(value.clone()).ok()) + .unwrap_or_default(), max_retries: object.get("max_retries").and_then(parse_u64_value), max_transfer_count: object .get("max_transfer_count") @@ -267,6 +278,7 @@ fn parse_status_code_list(value: &Value) -> BTreeSet { fn local_failover_policy_to_value(policy: &LocalFailoverPolicy) -> Value { json!({ + "routing_rules": policy.routing_rules, "max_retries": policy.max_retries, "max_transfer_count": policy.max_transfer_count, "max_transfer_timeout_seconds": policy.max_transfer_timeout_seconds, @@ -525,6 +537,7 @@ mod tests { assert_eq!( local_failover_policy_from_report_context(Some(&report_context)), Some(LocalFailoverPolicy { + routing_rules: Default::default(), max_retries: Some(2), max_transfer_count: 10, max_transfer_timeout_seconds: 60, diff --git a/apps/aether-gateway/src/routing/resolver.rs b/apps/aether-gateway/src/routing/resolver.rs index 0c0904579..49e28bbf9 100644 --- a/apps/aether-gateway/src/routing/resolver.rs +++ b/apps/aether-gateway/src/routing/resolver.rs @@ -74,7 +74,7 @@ pub(crate) fn resolve_gateway_routing_policy( }, ) .map_err(routing_policy_error)?; - crate::request_lifecycle::configure_client_disconnect(policy.execution_policy); + crate::request_lifecycle::configure_client_disconnect(policy.execution_policy.clone()); Ok(policy) } @@ -84,7 +84,7 @@ pub(crate) fn resolve_gateway_static_default_routing_policy( let Some(default_policy) = static_default_policy_fields(input.group_config_json)? else { return Ok(None); }; - crate::request_lifecycle::configure_client_disconnect(default_policy.execution_policy); + crate::request_lifecycle::configure_client_disconnect(default_policy.execution_policy.clone()); Ok(Some(ResolvedRoutingPolicy { group_id: input.group_id.map(str::to_string), @@ -145,32 +145,11 @@ fn static_default_policy_fields( .ok_or_else(invalid_routing_group_config)?, None => DEFAULT_STICKY_KEY_ATTEMPTS, }; - let enable_cf_heartbeat = routing_bool_field( - default_policy.get("enable_cf_heartbeat"), - "enable_cf_heartbeat", - )?; - // Older strategies stored separate image/text heartbeat flags. Treat - // either legacy flag as enabling the unified CF heartbeat setting while - // allowing newly saved strategies to use only the canonical key. - let legacy_image_heartbeat = routing_bool_field( - default_policy.get("enable_openai_image_sync_heartbeat"), - "enable_openai_image_sync_heartbeat", - )?; - let legacy_text_heartbeat = routing_bool_field( - default_policy.get("enable_standard_text_sync_heartbeat"), - "enable_standard_text_sync_heartbeat", - )?; - let execution_policy = aether_routing_core::RoutingExecutionPolicy { - enable_cf_heartbeat: enable_cf_heartbeat || legacy_image_heartbeat || legacy_text_heartbeat, - cyber_continue_failover: routing_bool_field( - default_policy.get("cyber_continue_failover"), - "cyber_continue_failover", - )?, - cancel_on_client_disconnect: routing_bool_field( - default_policy.get("cancel_on_client_disconnect"), - "cancel_on_client_disconnect", - )?, - }; + let execution_policy: aether_routing_core::RoutingExecutionPolicy = + serde_json::from_value(Value::Object(default_policy.clone())) + .map_err(|_| invalid_routing_group_config())?; + aether_routing_core::validate_routing_failover_rules(&execution_policy.failover_rules) + .map_err(|_| invalid_routing_group_config())?; Ok(Some(RoutingDefaultPolicy { priority_mode, @@ -181,13 +160,6 @@ fn static_default_policy_fields( })) } -fn routing_bool_field(value: Option<&Value>, _field: &str) -> Result { - match value { - Some(value) => value.as_bool().ok_or_else(invalid_routing_group_config), - None => Ok(false), - } -} - fn routing_array_field_is_missing_or_empty( object: &serde_json::Map, key: &str, @@ -246,7 +218,13 @@ mod tests { "priority_mode": "global_key", "scheduling_mode": "load_balance", "keep_priority_on_conversion": true, - "cancel_on_client_disconnect": true + "cancel_on_client_disconnect": true, + "max_transfer_count": 3, + "max_transfer_timeout_seconds": 90, + "failover_rules": { + "success_failover_patterns": [{"pattern": "(?i)capacity.*exhausted"}], + "error_stop_patterns": [{"status_codes": [400]}] + } }, "allowed_models": ["legacy-model"], "model_policies": [], @@ -282,6 +260,19 @@ mod tests { .expect("full policy should resolve"); assert_eq!(static_policy, full_policy); + assert_eq!(static_policy.execution_policy.max_transfer_count, 3); + assert_eq!( + static_policy.execution_policy.max_transfer_timeout_seconds, + 90 + ); + assert_eq!( + static_policy + .execution_policy + .failover_rules + .error_stop_patterns + .len(), + 1 + ); assert_eq!( static_policy.priority_mode, RoutingSetPriorityMode::GlobalKey diff --git a/crates/aether-routing-core/src/failover.rs b/crates/aether-routing-core/src/failover.rs new file mode 100644 index 000000000..bb09cad89 --- /dev/null +++ b/crates/aether-routing-core/src/failover.rs @@ -0,0 +1,120 @@ +use std::collections::BTreeSet; + +use regex::Regex; +use serde::{Deserialize, Serialize}; + +pub const MAX_ROUTING_FAILOVER_RULES: usize = 64; +pub const MAX_ROUTING_FAILOVER_PATTERN_BYTES: usize = 4096; + +#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)] +#[serde(default)] +pub struct RoutingFailoverRule { + pub pattern: String, + pub status_codes: BTreeSet, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)] +#[serde(default)] +pub struct RoutingFailoverRules { + pub success_failover_patterns: Vec, + pub error_stop_patterns: Vec, +} + +pub fn validate_routing_failover_rules(rules: &RoutingFailoverRules) -> Result<(), String> { + for (name, entries, success) in [ + ( + "success_failover_patterns", + &rules.success_failover_patterns, + true, + ), + ("error_stop_patterns", &rules.error_stop_patterns, false), + ] { + if entries.len() > MAX_ROUTING_FAILOVER_RULES { + return Err(format!("{name} exceeds {MAX_ROUTING_FAILOVER_RULES} rules")); + } + for (index, rule) in entries.iter().enumerate() { + let pattern = rule.pattern.trim(); + if pattern.is_empty() && (success || rule.status_codes.is_empty()) { + return Err(format!( + "{name}[{index}] requires a pattern or error status codes" + )); + } + if pattern.len() > MAX_ROUTING_FAILOVER_PATTERN_BYTES { + return Err(format!( + "{name}[{index}] pattern exceeds {MAX_ROUTING_FAILOVER_PATTERN_BYTES} bytes" + )); + } + if !pattern.is_empty() { + Regex::new(pattern) + .map_err(|error| format!("{name}[{index}] invalid regex: {error}"))?; + } + if rule.status_codes.iter().any(|status| { + if success { + *status != 200 + } else { + !(400..=599).contains(status) + } + }) { + return Err(format!("{name}[{index}] contains invalid status codes")); + } + } + } + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn validates_regex_and_status_only_stop_rules() { + let rules = RoutingFailoverRules { + success_failover_patterns: vec![RoutingFailoverRule { + pattern: "(?i)capacity.*exhausted".to_string(), + ..Default::default() + }], + error_stop_patterns: vec![RoutingFailoverRule { + status_codes: [400, 413].into_iter().collect(), + ..Default::default() + }], + ..Default::default() + }; + assert!(validate_routing_failover_rules(&rules).is_ok()); + } + + #[test] + fn rejects_invalid_or_unbounded_rule_configuration() { + for rule in [ + RoutingFailoverRule::default(), + RoutingFailoverRule { + pattern: "[".to_string(), + ..Default::default() + }, + RoutingFailoverRule { + pattern: "error".to_string(), + status_codes: [429].into_iter().collect(), + }, + RoutingFailoverRule { + pattern: "a".repeat(MAX_ROUTING_FAILOVER_PATTERN_BYTES + 1), + ..Default::default() + }, + ] { + let rules = RoutingFailoverRules { + success_failover_patterns: vec![rule], + ..Default::default() + }; + assert!(validate_routing_failover_rules(&rules).is_err()); + } + let rules = RoutingFailoverRules { + error_stop_patterns: vec![ + RoutingFailoverRule { + status_codes: [400].into_iter().collect(), + ..Default::default() + }; + MAX_ROUTING_FAILOVER_RULES + 1 + ], + ..Default::default() + }; + assert!(validate_routing_failover_rules(&rules).is_err()); + } +} diff --git a/crates/aether-routing-core/src/lib.rs b/crates/aether-routing-core/src/lib.rs index a9a97742c..f5fa2bb5a 100644 --- a/crates/aether-routing-core/src/lib.rs +++ b/crates/aether-routing-core/src/lib.rs @@ -1,5 +1,6 @@ mod actions; mod conditions; +mod failover; mod model; mod mutations; mod policy; @@ -12,6 +13,10 @@ pub use actions::{ RoutingSchedulingMode, RoutingSetPriorityMode, }; pub use conditions::{RoutingCondition, RoutingConditionContext, RoutingConditionOp}; +pub use failover::{ + validate_routing_failover_rules, RoutingFailoverRule, RoutingFailoverRules, + MAX_ROUTING_FAILOVER_PATTERN_BYTES, MAX_ROUTING_FAILOVER_RULES, +}; pub use model::{ RoutingDefaultPolicy, RoutingExecutionPolicy, RoutingGroupBinding, RoutingGroupBindingSubject, RoutingGroupConfig, RoutingGroupRecord, RoutingGroupVersionRecord, RoutingModelPolicy, diff --git a/crates/aether-routing-core/src/model.rs b/crates/aether-routing-core/src/model.rs index 03b9412bf..fc48854e9 100644 --- a/crates/aether-routing-core/src/model.rs +++ b/crates/aether-routing-core/src/model.rs @@ -7,6 +7,7 @@ use crate::actions::{ RoutingAction, RoutingRulePhase, RoutingSchedulingMode, RoutingSetPriorityMode, }; use crate::conditions::RoutingCondition; +use crate::RoutingFailoverRules; #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] pub struct RoutingSchedulingPreset { @@ -33,7 +34,7 @@ pub const DEFAULT_STICKY_KEY_ATTEMPTS: u32 = 2; /// transport configuration. A resolved policy is snapshotted for the request /// and can therefore be consumed by execution without rereading mutable /// system settings. -#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Default)] +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Default)] pub struct RoutingExecutionPolicy { #[serde(default, skip_serializing_if = "is_false")] pub enable_cf_heartbeat: bool, @@ -41,6 +42,12 @@ pub struct RoutingExecutionPolicy { pub cyber_continue_failover: bool, #[serde(default, skip_serializing_if = "is_false")] pub cancel_on_client_disconnect: bool, + #[serde(default)] + pub max_transfer_count: u64, + #[serde(default)] + pub max_transfer_timeout_seconds: u64, + #[serde(default)] + pub failover_rules: RoutingFailoverRules, } impl<'de> Deserialize<'de> for RoutingExecutionPolicy { @@ -60,6 +67,12 @@ impl<'de> Deserialize<'de> for RoutingExecutionPolicy { cyber_continue_failover: bool, #[serde(default)] cancel_on_client_disconnect: bool, + #[serde(default)] + max_transfer_count: u64, + #[serde(default)] + max_transfer_timeout_seconds: u64, + #[serde(default)] + failover_rules: RoutingFailoverRules, } let value = LegacyCompatibleExecutionPolicy::deserialize(deserializer)?; @@ -69,6 +82,9 @@ impl<'de> Deserialize<'de> for RoutingExecutionPolicy { || value.enable_standard_text_sync_heartbeat, cyber_continue_failover: value.cyber_continue_failover, cancel_on_client_disconnect: value.cancel_on_client_disconnect, + max_transfer_count: value.max_transfer_count, + max_transfer_timeout_seconds: value.max_transfer_timeout_seconds, + failover_rules: value.failover_rules, }) } } @@ -116,6 +132,29 @@ fn is_false(value: &bool) -> bool { mod execution_policy_tests { use super::*; + #[test] + fn routing_failover_configuration_round_trips_and_validates() { + let config: RoutingGroupConfig = serde_json::from_value(serde_json::json!({ + "default_policy": { + "max_transfer_count": 3, + "max_transfer_timeout_seconds": 90, + "failover_rules": { + "success_failover_patterns": [{ "pattern": "(?i)capacity" }], + "error_stop_patterns": [{ "status_codes": [400, 413] }] + } + } + })) + .unwrap(); + crate::validate_routing_group_config(&config).unwrap(); + assert_eq!(config.default_policy.execution_policy.max_transfer_count, 3); + let value = serde_json::to_value(&config).unwrap(); + assert_eq!(value["default_policy"]["max_transfer_timeout_seconds"], 90); + assert_eq!( + serde_json::from_value::(value).unwrap(), + config + ); + } + #[test] fn cancellation_defaults_off_and_round_trips_with_legacy_heartbeat() { let default: RoutingDefaultPolicy = serde_json::from_str("{}").unwrap(); diff --git a/crates/aether-routing-core/src/policy.rs b/crates/aether-routing-core/src/policy.rs index 710b0bbbe..526128e7e 100644 --- a/crates/aether-routing-core/src/policy.rs +++ b/crates/aether-routing-core/src/policy.rs @@ -89,7 +89,7 @@ pub fn resolve_routing_policy( scheduling_mode: config.default_policy.scheduling_mode, keep_priority_on_conversion: config.default_policy.keep_priority_on_conversion, sticky_key_attempts: config.default_policy.sticky_key_attempts, - execution_policy: config.default_policy.execution_policy, + execution_policy: config.default_policy.execution_policy.clone(), ranking_overlay: RankingOverlay::default(), mutation_plan: MutationPlan::default(), pool_policy_overrides: BTreeMap::new(), diff --git a/crates/aether-routing-core/src/validation.rs b/crates/aether-routing-core/src/validation.rs index 8314cc8c8..a4c3fb8d7 100644 --- a/crates/aether-routing-core/src/validation.rs +++ b/crates/aether-routing-core/src/validation.rs @@ -28,6 +28,8 @@ const ROUTING_POOL_PRESETS: &[&str] = &[ #[derive(Debug, Error, Clone, PartialEq, Eq)] pub enum RoutingValidationError { + #[error("routing failover rules are invalid: {0}")] + InvalidFailoverRules(String), #[error("routing rule id is empty")] EmptyRuleId, #[error("duplicate routing rule id: {0}")] @@ -69,6 +71,8 @@ pub enum RoutingValidationError { pub fn validate_routing_group_config( config: &RoutingGroupConfig, ) -> Result<(), RoutingValidationError> { + crate::validate_routing_failover_rules(&config.default_policy.execution_policy.failover_rules) + .map_err(RoutingValidationError::InvalidFailoverRules)?; let mut rule_ids = BTreeSet::new(); for model_policy in &config.model_policies { if model_policy.model.trim().is_empty() { diff --git a/docs/operations/routing-failover.md b/docs/operations/routing-failover.md new file mode 100644 index 000000000..a16e89bea --- /dev/null +++ b/docs/operations/routing-failover.md @@ -0,0 +1,55 @@ +# 调度策略级故障转移 + +调度策略的 `default_policy` 支持跨提供商的转移预算与错误规则。它们跟随当前请求的 `routing_execution_policy` 快照进入执行器,不依赖运行中修改全局系统设置。已有策略缺省为不限次数、不限累计时间、无全局错误规则。 + +```json +{ + "default_policy": { + "sticky_key_attempts": 2, + "max_transfer_count": 3, + "max_transfer_timeout_seconds": 90, + "failover_rules": { + "success_failover_patterns": [ + { "pattern": "(?i)capacity.*exhausted" } + ], + "error_stop_patterns": [ + { "status_codes": [400, 413], "pattern": "invalid.*parameter" }, + { "status_codes": [422] } + ] + } + } +} +``` + +## 预算语义 + +- `sticky_key_attempts` 是首个粘性候选上的总尝试次数,`2` 表示首次请求加一次同 Key 重试。该行为保持不变。 +- 全局 `max_transfer_count` 统计切换候选的次数,首次尝试不计数。同一提供商、端点、Key 上的重试不计数;改变该组合计一次。`3` 最多允许首次候选之后再切换三次。 +- 全局 `max_transfer_timeout_seconds` 从首次候选开始执行时计时,覆盖后续重试与切换间的累计耗时。它在准备下一次尝试时检查,不会强制打断已经执行中的调用或已提交给客户端的流;单次连接、首字节、读取和非流式完整调用超时仍独立生效。 +- 两个全局预算的 `0` 都表示不限制。被筛除、禁用或被提供商级预算跳过而未执行的候选不计数。 +- 提供商自身的转移次数和时间预算继续生效。提供商预算耗尽只跳过该提供商,仍可尝试其他提供商;全局预算耗尽则不再执行任何提供商的新尝试。 +- 预算仅约束当前请求,不能重置或替代客户端取消策略、权限校验及本地执行异常的终止行为。 + +## 错误规则 + +先匹配调度策略的全局显式规则;未匹配时继续使用提供商的规则及协议默认行为。本地执行函数真正返回 `Err` 时仍然终止,不做兜底重放。 + +- **成功转移规则**:仅当 HTTP 200 响应匹配配置的正则时继续转移,不是对所有 200 进行重试。非流式请求匹配响应体;流式请求只匹配尚未交付业务输出的有界预读取内容。 +- **错误提前终止**:适用于 400–599 错误。状态码与正则都填写时要求同时满足;只填状态码表示该状态一律终止;只填正则表示在所有错误状态上匹配。流内错误使用解析后的错误状态,而不是外层 200。 +- **网络错误**:没有上游 HTTP 状态的连接、TLS、DNS、提交前超时等错误统一继续转移;提供商级别若单独配置了停止规则,仍按提供商规则处理。 +- 正则使用 Rust `regex` 语法,支持 `(?i)` 等内联标志。服务端拒绝无效正则、无意义的空规则以及错误状态范围。每组最多 64 条,每条表达式最多 4096 字节。 + +## 流式 200 的恢复窗口 + +上游 HTTP 200 响应头不再默认关闭标准文本 SSE 的恢复窗口。执行器先缓冲协议开场事件,例如 Responses 的 `response.created`、Chat 的 role-only 增量、Anthropic 的空 `message_start` / 文本块起始事件。图片专用流保留原来的响应头提交行为,除非显式配置了需要预读取的规则。 + +首个业务内容之前的结构化错误、过早 EOF、首字节超时以及 200 正则命中会进入统一故障转移判断。真实文本、思考、工具调用或正常结束事件确定后,缓冲内容按原顺序交付;之后发生的错误保持终止,不重新执行原请求。 + +预读取受单次首字节时间和既有字节上限约束。达到字节上限时保守提交,避免无限缓冲;这意味着不能承诺识别响应任意位置的错误或正则。原始 HTTP 状态与最终执行结果是不同观测值,不能因流内失败而伪改已经发送的 HTTP 状态码。 + +## 排查 + +- `routing_transfer_limit_reached`:当前调度策略的累计次数或时间预算耗尽。 +- `provider_transfer_limit_reached`:提供商自身的预算耗尽。 +- `local_stream_candidate_retry_scheduled`:输出前的流内错误或 200 规则触发了继续调度。 +- `local_stream_transport_retry_scheduled`:输出前传输错误触发了继续调度。 diff --git a/frontend/src/features/routing/__tests__/RoutingFailoverPolicyEditor.spec.ts b/frontend/src/features/routing/__tests__/RoutingFailoverPolicyEditor.spec.ts new file mode 100644 index 000000000..6be181185 --- /dev/null +++ b/frontend/src/features/routing/__tests__/RoutingFailoverPolicyEditor.spec.ts @@ -0,0 +1,111 @@ +import { afterEach, describe, expect, it } from 'vitest' +import { createApp, h, nextTick, ref, type App } from 'vue' +import RoutingFailoverPolicyEditor from '../components/RoutingFailoverPolicyEditor.vue' +import { normalizeRoutingFailoverPolicy, type RoutingFailoverPolicy } from '../utils/routingFailover' + +const mounted: Array<{ app: App, root: HTMLElement }> = [] + +function mountEditor() { + const policy = ref(normalizeRoutingFailoverPolicy()) + const root = document.createElement('div') + document.body.appendChild(root) + const app = createApp({ + setup: () => () => h(RoutingFailoverPolicyEditor, { + modelValue: policy.value, + 'onUpdate:modelValue': (value: RoutingFailoverPolicy) => { policy.value = value }, + }), + }) + app.mount(root) + mounted.push({ app, root }) + return { root, policy } +} + +function control(root: HTMLElement, label: string): T { + const element = root.querySelector(`[aria-label="${label}"]`) + if (!element) throw new Error(`Missing control: ${label}`) + return element +} + +afterEach(() => { + for (const { app, root } of mounted.splice(0)) { + app.unmount() + root.remove() + } +}) + +describe('RoutingFailoverPolicyEditor', () => { + it('edits independent global budgets and documents sticky retry exclusion', async () => { + const { root, policy } = mountEditor() + expect(root.textContent).toContain('首次尝试和粘性同 Key 重试不计入') + expect(root.textContent).toContain('不会中断已开始的调用') + const count = control(root, '全局最大转移次数') + count.value = '4' + count.dispatchEvent(new Event('input', { bubbles: true })) + await nextTick() + expect(policy.value.max_transfer_count).toBe(4) + expect(policy.value.max_transfer_timeout_seconds).toBe(0) + }) + + it('adds regex and status-only rules and reports invalid drafts', async () => { + const { root, policy } = mountEditor() + control(root, '添加成功转移规则').click() + await nextTick() + expect(root.querySelector('[role="alert"]')?.textContent).toContain('正则表达式') + const regex = control(root, '成功转移规则 1 正则') + regex.value = '(?i)capacity.*exhausted' + regex.dispatchEvent(new Event('input', { bubbles: true })) + await nextTick() + expect(policy.value.failover_rules.success_failover_patterns[0].pattern).toBe('(?i)capacity.*exhausted') + control(root, '添加错误提前终止规则').click() + await nextTick() + const statuses = control(root, '终止规则 1 状态码') + statuses.value = '400, 413' + statuses.dispatchEvent(new Event('input', { bubbles: true })) + await nextTick() + expect(policy.value.failover_rules.error_stop_patterns[0].status_codes).toEqual([400, 413]) + expect(root.querySelector('[role="alert"]')).toBeNull() + control(root, '删除成功转移规则 1').click() + await nextTick() + expect(policy.value.failover_rules.success_failover_patterns).toHaveLength(0) + }) + + it('edits and applies both rule groups through JSON mode', async () => { + const { root, policy } = mountEditor() + control(root, '切到成功转移规则 JSON').click() + await nextTick() + const successJson = root.querySelector('textarea') + if (!successJson) throw new Error('Missing success JSON editor') + successJson.value = '[{"pattern":"capacity"}]' + successJson.dispatchEvent(new Event('input', { bubbles: true })) + await nextTick() + control(root, '切回成功转移规则表单').click() + await nextTick() + expect(policy.value.failover_rules.success_failover_patterns).toEqual([{ pattern: 'capacity', status_codes: [] }]) + + control(root, '切到错误提前终止规则 JSON').click() + await nextTick() + const errorJson = root.querySelector('textarea') + if (!errorJson) throw new Error('Missing error JSON editor') + errorJson.value = '[{"status_codes":[429,500],"pattern":"rate"}]' + errorJson.dispatchEvent(new Event('input', { bubbles: true })) + await nextTick() + control(root, '切回错误提前终止规则表单').click() + await nextTick() + expect(policy.value.failover_rules.error_stop_patterns).toEqual([{ pattern: 'rate', status_codes: [429, 500] }]) + }) + + it('keeps invalid JSON visible until it is corrected', async () => { + const { root, policy } = mountEditor() + control(root, '切到错误提前终止规则 JSON').click() + await nextTick() + const editor = root.querySelector('textarea') + if (!editor) throw new Error('Missing JSON editor') + editor.value = '{' + editor.dispatchEvent(new Event('input', { bubbles: true })) + await nextTick() + control(root, '切回错误提前终止规则表单').click() + await nextTick() + expect(root.querySelector('[role="alert"]')?.textContent).toContain('JSON') + expect(policy.value.failover_rules.error_stop_patterns).toHaveLength(0) + }) +}) diff --git a/frontend/src/features/routing/__tests__/routingFailover.spec.ts b/frontend/src/features/routing/__tests__/routingFailover.spec.ts new file mode 100644 index 000000000..cf085bc27 --- /dev/null +++ b/frontend/src/features/routing/__tests__/routingFailover.spec.ts @@ -0,0 +1,49 @@ +import { describe, expect, it } from 'vitest' +import { normalizeRoutingFailoverPolicy, validateRoutingFailoverPolicy } from '../utils/routingFailover' +import { createEmptyRoutingGroupConfig, getModelScheduling, normalizeRoutingGroupConfig, upsertModelSchedulingRule } from '../utils/routingPolicy' + +describe('routing failover policy', () => { + it('keeps legacy strategies unlimited with empty global rules', () => { + const policy = normalizeRoutingGroupConfig({}).default_policy + expect(policy.max_transfer_count).toBe(0) + expect(policy.max_transfer_timeout_seconds).toBe(0) + expect(policy.failover_rules).toEqual({ success_failover_patterns: [], error_stop_patterns: [] }) + expect(validateRoutingFailoverPolicy(policy)).toBeNull() + }) + + it('preserves global limits and rules across model edits without sharing mutable arrays', () => { + const config = createEmptyRoutingGroupConfig() + Object.assign(config.default_policy, { + max_transfer_count: 3, + max_transfer_timeout_seconds: 90, + failover_rules: { + success_failover_patterns: [{ pattern: '(?i)capacity', status_codes: [] }], + error_stop_patterns: [{ pattern: '', status_codes: [400, 413] }], + }, + }) + const updated = upsertModelSchedulingRule(config, 'model-a', { priority_mode: 'provider', scheduling_mode: 'fixed_order' }) + const policy = getModelScheduling(updated, 'model-a') + expect(policy.max_transfer_count).toBe(3) + expect(policy.max_transfer_timeout_seconds).toBe(90) + expect(policy.failover_rules).toEqual(config.default_policy.failover_rules) + expect(validateRoutingFailoverPolicy(policy)).toBeNull() + policy.failover_rules.error_stop_patterns[0].status_codes.push(422) + expect(config.default_policy.failover_rules.error_stop_patterns[0].status_codes).toEqual([400, 413]) + }) + + it('rejects invalid budgets and ambiguous empty rules before saving', () => { + const policy = normalizeRoutingFailoverPolicy() + policy.max_transfer_count = -1 + expect(validateRoutingFailoverPolicy(policy)).toContain('非负整数') + policy.max_transfer_count = 0 + policy.max_transfer_timeout_seconds = 0.5 + expect(validateRoutingFailoverPolicy(policy)).toContain('非负整数') + policy.max_transfer_timeout_seconds = 0 + policy.failover_rules.error_stop_patterns.push({ pattern: '', status_codes: [] }) + expect(validateRoutingFailoverPolicy(policy)).toContain('状态码或正则') + policy.failover_rules.error_stop_patterns[0].status_codes = [200] + expect(validateRoutingFailoverPolicy(policy)).toContain('400–599') + policy.failover_rules.error_stop_patterns[0].status_codes = [400] + expect(validateRoutingFailoverPolicy(policy)).toBeNull() + }) +}) diff --git a/frontend/src/features/routing/__tests__/routingPolicy.spec.ts b/frontend/src/features/routing/__tests__/routingPolicy.spec.ts index ed3008f98..38cd38569 100644 --- a/frontend/src/features/routing/__tests__/routingPolicy.spec.ts +++ b/frontend/src/features/routing/__tests__/routingPolicy.spec.ts @@ -91,7 +91,7 @@ describe('routingPolicy', () => { expect(createEmptyRoutingGroupConfig().default_policy.sticky_key_attempts).toBe(2) expect(normalizeRoutingGroupConfig({}).default_policy.sticky_key_attempts).toBe(2) expect(normalizeRoutingGroupConfig({ - default_policy: { priority_mode: 'provider', scheduling_mode: 'cache_affinity', keep_priority_on_conversion: false, sticky_key_attempts: 3, enable_cf_heartbeat: false, cyber_continue_failover: false, cancel_on_client_disconnect: false }, + default_policy: { ...createEmptyRoutingGroupConfig().default_policy, priority_mode: 'provider', scheduling_mode: 'cache_affinity', keep_priority_on_conversion: false, sticky_key_attempts: 3, enable_cf_heartbeat: false, cyber_continue_failover: false, cancel_on_client_disconnect: false }, }).default_policy.sticky_key_attempts).toBe(3) expect(normalizeStickyKeyAttempts('5')).toBe(5) expect(normalizeStickyKeyAttempts(-1)).toBe(2) diff --git a/frontend/src/features/routing/components/RoutingFailoverPolicyEditor.vue b/frontend/src/features/routing/components/RoutingFailoverPolicyEditor.vue new file mode 100644 index 000000000..eef987bd6 --- /dev/null +++ b/frontend/src/features/routing/components/RoutingFailoverPolicyEditor.vue @@ -0,0 +1,373 @@ +