fix: harden routing failover, model testing, and wallet queries

This commit is contained in:
elky
2026-09-09 10:38:25 +08:00
parent a893bd0557
commit 6630856061
14 changed files with 1530 additions and 84 deletions
@@ -14,7 +14,7 @@ fn sync_plan_kind_disables_local_candidate_failover(plan_kind: &str) -> bool {
)
}
fn openai_image_success_disables_local_success_failover(
pub(super) fn openai_image_success_disables_local_success_failover(
plan: &ExecutionPlan,
status_code: u16,
) -> bool {
@@ -4759,6 +4759,31 @@ fn parse_prefetched_sync_json_body(body: &[u8]) -> Option<Value> {
serde_json::from_slice::<Value>(stripped).ok()
}
fn success_failover_matchable_body<'body>(
headers: &BTreeMap<String, String>,
body: &'body [u8],
) -> Option<&'body [u8]> {
if body.is_empty() {
return None;
}
if response_headers_indicate_sse(headers) {
let mut complete_end = 0;
while let Some((record_end, separator_len)) =
find_sse_record_boundary(&body[complete_end..])
{
complete_end += record_end + separator_len;
}
return (complete_end > 0).then_some(&body[..complete_end]);
}
let stripped = strip_utf8_bom_and_ws(body);
if stripped.starts_with(b"{") || stripped.starts_with(b"[") {
if serde_json::from_slice::<Value>(stripped).is_err_and(|error| error.is_eof()) {
return None;
}
}
Some(body)
}
fn resolve_provider_stream_error_status_code(
provider_api_format: &str,
upstream_status_code: u16,
@@ -5881,7 +5906,11 @@ async fn execute_stream_from_frame_stream_with_retry_scope(
}
let mut buffered_frames = VecDeque::new();
let mut stream_terminal_summary: Option<ExecutionStreamTerminalSummary> = None;
if status_code == 200 && should_probe_success_failover_before_stream(&headers) {
let direct_stream_finalize_kind = resolve_core_stream_direct_finalize_report_kind(plan_kind);
if status_code == 200
&& direct_stream_finalize_kind.is_none()
&& should_probe_success_failover_before_stream(&headers)
{
let success_probe_text =
probe_local_stream_success_failover_text(&mut buffered_frames, &mut lines).await?;
if should_retry_next_local_candidate_stream(
@@ -6298,7 +6327,6 @@ async fn execute_stream_from_frame_stream_with_retry_scope(
)?));
}
let direct_stream_finalize_kind = resolve_core_stream_direct_finalize_report_kind(plan_kind);
let normalized_stream_report_context =
normalize_provider_private_report_context(report_context.as_ref());
let upstream_headers = headers.clone();
@@ -6358,6 +6386,12 @@ async fn execute_stream_from_frame_stream_with_retry_scope(
.iter()
.map(|rule| (&rule.pattern, &rule.status_codes)),
)
.filter(|_| {
!crate::execution_runtime::fallback::openai_image_success_disables_local_success_failover(
&plan,
status_code,
)
})
.filter(|(_, status_codes)| status_codes.is_empty() || status_codes.contains(&200))
.filter_map(|(pattern, _)| regex::Regex::new(pattern.trim()).ok())
.collect::<Vec<_>>();
@@ -6661,28 +6695,6 @@ async fn execute_stream_from_frame_stream_with_retry_scope(
&mut prefetched_inspection_body_truncated,
);
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,
@@ -6792,6 +6804,33 @@ async fn execute_stream_from_frame_stream_with_retry_scope(
StreamPrefetchInspection::NonError => {}
}
if !prefetch_success_patterns.is_empty()
&& crate::orchestration::attempt_identity_from_report_context(
report_context.as_ref(),
)
.is_some()
{
if let Some(matchable_body) = success_failover_matchable_body(
&upstream_headers,
&prefetched_inspection_body,
) {
let response_text = String::from_utf8_lossy(matchable_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);
}
}
}
if !response_headers_indicate_sse(&upstream_headers)
&& (200..300).contains(&status_code)
{
@@ -9035,11 +9074,30 @@ mod tests {
provider_config: Option<Value>,
stall: bool,
content_type: &str,
) -> Option<axum::http::Response<Body>> {
execute_stream_precommit_for_format(
chunks,
routing_policy,
provider_config,
stall,
content_type,
"openai:responses",
)
.await
}
async fn execute_stream_precommit_for_format(
chunks: Vec<&str>,
routing_policy: Value,
provider_config: Option<Value>,
stall: bool,
content_type: &str,
api_format: &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.provider_api_format = api_format.to_string();
plan.client_api_format = api_format.to_string();
plan.timeouts = Some(ExecutionTimeouts {
first_byte_ms: Some(20),
..Default::default()
@@ -9077,17 +9135,22 @@ mod tests {
}
.boxed();
let mut scope = AiAttemptRetryScope::Provider;
let plan_kind = if api_format == "openai:image" {
"openai_image_stream"
} else {
"openai_responses_stream"
};
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()),
plan_kind,
Some(format!("{plan_kind}_success")),
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",
"provider_api_format": api_format, "client_api_format": api_format,
"routing_execution_policy": routing_policy,
})),
crate::clock::current_unix_ms(),
@@ -9107,13 +9170,18 @@ mod tests {
#[tokio::test]
async fn generic_stream_success_regex_matches_fragmented_plain_body() {
assert!(execute_generic_stream_precommit(
for chunks in [
vec!["upstream CAPACITY ", "exhausted"],
json!({"failover_rules": {"success_failover_patterns": [{"pattern": "(?i)capacity.*exhausted"}]}}),
None,
false,
"text/plain",
).await.is_none());
vec!["[upstream] CAPACITY ", "exhausted"],
] {
assert!(execute_generic_stream_precommit(
chunks,
json!({"failover_rules": {"success_failover_patterns": [{"pattern": "(?i)capacity.*exhausted"}]}}),
None,
false,
"text/plain",
).await.is_none());
}
}
#[tokio::test]
@@ -9152,6 +9220,128 @@ mod tests {
to_bytes(response.into_body(), usize::MAX).await.unwrap();
}
#[tokio::test]
async fn generic_image_success_is_not_replayed_by_global_or_provider_success_regex() {
let rule = json!({ "success_failover_patterns": [{ "pattern": "b64_json" }] });
for (routing_policy, provider_config) in [
(json!({ "failover_rules": rule }), None),
(json!({}), Some(json!({ "failover_rules": rule }))),
] {
let response = execute_stream_precommit_for_format(
vec![r#"{"created":1,"data":[{"b64_json":"aGVsbG8="}]}"#],
routing_policy,
provider_config,
false,
"application/json",
"openai:image",
)
.await
.expect("successful image responses must retain their no-replay protection");
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("aGVsbG8="));
}
}
#[tokio::test]
async fn generic_complete_setup_events_and_json_bodies_still_match_success_regex() {
for (content_type, chunks) in [
(
"text/event-stream",
vec![
"event: response.created\ndata: {\"type\":\"response.created\",\"response\":{\"metadata\":{\"warning\":\"capacity",
" exhausted\"}}}\n\n",
],
),
(
"application/json",
vec!["{\"warning\":\"capacity", " exhausted\"}"],
),
] {
assert!(execute_generic_stream_precommit(
chunks,
json!({ "failover_rules": {
"success_failover_patterns": [{ "pattern": "capacity.*exhausted" }],
} }),
None,
false,
content_type,
)
.await
.is_none());
}
}
#[tokio::test]
async fn generic_fragmented_errors_apply_stop_rules_before_success_regex() {
for (content_type, chunks) in [
(
"text/event-stream",
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",
],
),
(
"application/json",
vec![
"{\"error\":{\"type\":\"server_error\",\"message\":\"capacity",
" exhausted\"}}",
],
),
(
"application/json",
vec![r#"{"error":{"type":"server_error","message":"capacity exhausted"}}"#],
),
] {
let response = execute_generic_stream_precommit(
chunks,
json!({ "failover_rules": {
"success_failover_patterns": [{ "pattern": "capacity" }],
"error_stop_patterns": [{ "status_codes": [500], "pattern": "capacity" }],
} }),
None,
false,
content_type,
)
.await
.unwrap_or_else(|| panic!("partial errors must be parsed before applying success regex rules ({content_type})"));
assert!(response.status().is_server_error());
to_bytes(response.into_body(), usize::MAX).await.unwrap();
}
}
#[tokio::test]
async fn generic_sse_global_error_stop_precedes_global_or_provider_success_regex() {
let success_rules = json!({ "success_failover_patterns": [{ "pattern": "capacity" }] });
let stop_rule = json!([{ "status_codes": [500], "pattern": "capacity" }]);
for (routing_policy, provider_config) in [
(
json!({ "failover_rules": {
"success_failover_patterns": success_rules["success_failover_patterns"],
"error_stop_patterns": stop_rule,
} }),
None,
),
(
json!({ "failover_rules": { "error_stop_patterns": stop_rule } }),
Some(json!({ "failover_rules": success_rules })),
),
] {
let response = execute_generic_sse_precommit(
vec!["event: response.failed\ndata: {\"type\":\"response.failed\",\"response\":{\"error\":{\"type\":\"server_error\",\"message\":\"capacity exhausted\"}}}\n\n"],
routing_policy,
provider_config,
false,
)
.await
.expect("a matching global stop rule must win over a 200 success regex");
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" }] });
@@ -9177,6 +9367,18 @@ mod tests {
assert!(String::from_utf8_lossy(&body).contains("hello"));
}
#[tokio::test]
async fn generic_sse_setup_timeout_ignores_removed_global_transport_stop_flag() {
let response = execute_generic_sse_precommit(
vec!["event: response.created\ndata: {\"type\":\"response.created\"}\n\n"],
json!({ "failover_rules": { "stop_on_transport_errors": true } }),
None,
true,
)
.await;
assert!(response.is_none());
}
#[tokio::test]
async fn generic_sse_setup_timeout_always_retries() {
let response = execute_generic_sse_precommit(
@@ -2651,6 +2651,43 @@ mod tests {
);
}
#[tokio::test]
async fn cloned_tracker_preserves_global_budget_across_candidate_loops() {
let state = AppState::new().unwrap();
let tracker = ProviderTransferTracker::default();
let mut attempts = transfer_test_attempts();
for attempt in &mut attempts {
attempt.report_context["routing_execution_policy"] = json!({ "max_transfer_count": 1 });
}
let remaining = attempts.split_off(3);
let first_port = TransferTestPort::with_tracker(&state, tracker.clone());
let first_outcome = run_ai_attempt_loop(&first_port, attempts).await.unwrap();
assert!(matches!(first_outcome, AiAttemptLoopOutcome::Exhausted(_)));
assert_eq!(tracker.state.lock().await.global.transfer_count, 1);
let second_port = TransferTestPort::with_tracker(&state, tracker.clone());
let mut source = TransferTestAttemptSource {
attempts: remaining.into(),
skipped_providers: Vec::new(),
};
let second_outcome = run_dynamic_attempt_loop(
&second_port,
&mut source,
"global-budget-across-loops",
"test",
Duration::from_secs(1),
)
.await
.unwrap();
assert!(matches!(
second_outcome,
LocalExecutionRequestOutcome::NoPath
));
assert!(second_port.executed.lock().unwrap().is_empty());
assert_eq!(source.skipped_providers, ["provider-a", "provider-b"]);
assert!(tracker.state.lock().await.global.exhausted);
}
#[tokio::test]
async fn cloned_tracker_preserves_transfer_budget_across_candidate_loops() {
let state = AppState::new().expect("state should build");
@@ -605,7 +605,7 @@ mod tests {
}
#[test]
fn routing_transport_stop_cannot_be_overridden_by_provider() {
fn provider_transport_stop_rule_is_respected() {
let policy = super::LocalFailoverPolicy {
stop_on_transport_errors: true,
..Default::default()