feat(routing): add strategy failover controls

This commit is contained in:
elky
2026-09-09 09:12:09 +08:00
parent e58570d79d
commit f2839ae6a7
31 changed files with 1881 additions and 160 deletions
@@ -194,7 +194,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> 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<Option<AiSyncAttempt>, GatewayError> {
@@ -246,7 +246,7 @@ impl LocalExecutionAttemptSource<AiStreamAttempt>
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<Option<AiStreamAttempt>, GatewayError> {
@@ -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);
}
}
@@ -179,7 +179,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> 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<Option<AiSyncAttempt>, GatewayError> {
@@ -224,7 +224,7 @@ impl LocalExecutionAttemptSource<AiStreamAttempt> 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<Option<AiStreamAttempt>, GatewayError> {
@@ -257,7 +257,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> 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<Option<AiSyncAttempt>, GatewayError> {
@@ -302,7 +302,7 @@ impl LocalExecutionAttemptSource<AiStreamAttempt> 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<Option<AiStreamAttempt>, GatewayError> {
@@ -109,7 +109,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> 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<Option<AiSyncAttempt>, GatewayError> {
@@ -182,7 +182,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> 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<Option<AiSyncAttempt>, GatewayError> {
@@ -232,7 +232,7 @@ impl LocalExecutionAttemptSource<AiStreamAttempt> 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<Option<AiStreamAttempt>, GatewayError> {
@@ -124,7 +124,7 @@ impl LocalExecutionAttemptSource<AiStreamAttempt> 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<Option<AiStreamAttempt>, GatewayError> {
@@ -97,7 +97,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> 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<Option<AiSyncAttempt>, GatewayError> {
@@ -166,7 +166,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> 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<Option<AiSyncAttempt>, GatewayError> {
@@ -216,7 +216,7 @@ impl LocalExecutionAttemptSource<AiStreamAttempt> 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<Option<AiStreamAttempt>, GatewayError> {
@@ -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(),
@@ -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<Duration> {
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<u8>,
}
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::<Vec<_>>()
.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::<Value>(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<u8>,
@@ -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",
)
@@ -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<String, String
content_type.contains("json") || content_type.ends_with("+json")
}
async fn record_prefetch_success_failover(
state: &AppState,
plan: &ExecutionPlan,
report_context: Option<&Value>,
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<R>(
buffered_frames: &mut VecDeque<ObservedStreamFrame>,
lines: &mut FramedRead<R, LinesCodec>,
@@ -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::<Vec<_>>();
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<Value>,
stall: bool,
) -> Option<axum::http::Response<Body>> {
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<Value>,
stall: bool,
content_type: &str,
) -> Option<axum::http::Response<Body>> {
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::<Vec<_>>();
let content_type = content_type.to_string();
let frames = stream! {
yield Ok::<Bytes, std::io::Error>(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)
@@ -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;
@@ -655,6 +655,73 @@ struct ProviderTransferState {
struct ProviderTransferStateTracker {
by_provider: BTreeMap<String, ProviderTransferState>,
exhausted_provider_ids: BTreeSet<String>,
global: GlobalTransferState,
}
#[derive(Debug, Default)]
struct GlobalTransferState {
first_attempt_started_at: Option<Instant>,
last_candidate: Option<(String, String, String)>,
transfer_count: u64,
limits: Option<ProviderTransferLimits>,
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<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");
@@ -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<u16>,
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::{
@@ -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<u64>,
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<u16> {
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,
+27 -36
View File
@@ -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<bool, GatewayError> {
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<String, Value>,
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