mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-04 16:37:46 +08:00
feat(routing): add strategy failover controls
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user