feat: improve failover rules and request timeline

This commit is contained in:
elky
2026-06-11 00:49:29 +08:00
parent 31fade82f6
commit 30b545785f
15 changed files with 1576 additions and 199 deletions
@@ -83,15 +83,6 @@ pub(super) async fn maybe_build_local_openai_image_decision_payload_for_candidat
"chatgpt_web_image".to_string(),
serde_json::Value::Bool(true),
);
extra_fields.insert(
"local_failover_policy".to_string(),
serde_json::json!({
"stop_status_codes": [400, 401, 403, 429, 500, 502, 503, 504],
"error_stop_patterns": [
{ "pattern": ".*" }
]
}),
);
}
let upstream_is_stream = resolved
.provider_request_body
@@ -105,15 +105,6 @@ pub(crate) async fn maybe_build_local_openai_chat_decision_payload_for_candidate
.eq_ignore_ascii_case("chatgpt_web")
{
extra_fields.insert("chatgpt_web_image".to_string(), serde_json::json!(true));
extra_fields.insert(
"local_failover_policy".to_string(),
serde_json::json!({
"stop_status_codes": [400, 401, 403, 429, 500, 502, 503, 504],
"error_stop_patterns": [
{ "pattern": ".*" }
]
}),
);
}
let super::request::LocalOpenAiChatCandidatePayloadParts {
client_api_format,
@@ -102,15 +102,6 @@ pub(crate) async fn maybe_build_local_openai_responses_decision_payload_for_cand
.eq_ignore_ascii_case("chatgpt_web")
{
extra_fields.insert("chatgpt_web_image".to_string(), json!(true));
extra_fields.insert(
"local_failover_policy".to_string(),
json!({
"stop_status_codes": [400, 401, 403, 429, 500, 502, 503, 504],
"error_stop_patterns": [
{ "pattern": ".*" }
]
}),
);
}
insert_provider_stream_event_api_format(
&mut extra_fields,
@@ -1007,6 +1007,100 @@ mod tests {
);
}
#[tokio::test]
async fn provider_failover_rules_can_stop_rate_limit_status() {
let result = ExecutionResult {
request_id: "req-1".to_string(),
candidate_id: None,
status_code: 429,
headers: Default::default(),
body: None,
telemetry: None,
error: None,
};
let local_report_context = serde_json::json!({
"candidate_index": 0,
"retry_index": 0,
});
let state = build_state_with_provider_config(Some(serde_json::json!({
"failover_rules": {
"stop_on_status_codes": [429]
}
})));
let plan = sample_plan();
assert!(
should_stop_local_candidate_failover_sync(
&state,
&plan,
"openai_chat_sync",
Some(&local_report_context),
&result,
Some("{\"error\":{\"message\":\"rate limited\"}}"),
)
.await
);
assert!(
!should_retry_next_local_candidate_sync(
&state,
&plan,
"openai_chat_sync",
Some(&local_report_context),
&result,
Some("{\"error\":{\"message\":\"rate limited\"}}"),
)
.await
);
}
#[tokio::test]
async fn status_only_error_stop_rule_can_stop_rate_limit_status() {
let result = ExecutionResult {
request_id: "req-1".to_string(),
candidate_id: None,
status_code: 429,
headers: Default::default(),
body: None,
telemetry: None,
error: None,
};
let local_report_context = serde_json::json!({
"candidate_index": 0,
"retry_index": 0,
});
let state = build_state_with_provider_config(Some(serde_json::json!({
"failover_rules": {
"error_stop_patterns": [
{"status_codes": [429]}
]
}
})));
let plan = sample_plan();
assert!(
should_stop_local_candidate_failover_sync(
&state,
&plan,
"openai_chat_sync",
Some(&local_report_context),
&result,
None,
)
.await
);
assert!(
!should_retry_next_local_candidate_sync(
&state,
&plan,
"openai_chat_sync",
Some(&local_report_context),
&result,
None,
)
.await
);
}
#[test]
fn resolve_local_failover_policy_reads_regex_rules() {
let state = build_state_with_provider_config(Some(serde_json::json!({
@@ -1039,6 +1133,32 @@ mod tests {
);
}
#[test]
fn resolve_local_failover_policy_reads_status_only_error_stop_rules() {
let state = build_state_with_provider_config(Some(serde_json::json!({
"failover_rules": {
"success_failover_patterns": [
{"status_codes": [200]}
],
"error_stop_patterns": [
{"status_codes": [429]}
]
}
})));
let plan = sample_plan();
let runtime = tokio::runtime::Runtime::new().expect("runtime should build");
let policy = runtime.block_on(resolve_local_failover_policy(&state, &plan, None));
assert!(policy.success_failover_patterns.is_empty());
assert_eq!(
policy.error_stop_patterns,
vec![LocalFailoverRegexRule {
pattern: String::new(),
status_codes: [429].into_iter().collect(),
}]
);
}
#[tokio::test]
async fn success_failover_pattern_can_retry_sync_candidate() {
let result = ExecutionResult {
@@ -1125,11 +1245,11 @@ mod tests {
}
#[tokio::test]
async fn chatgpt_web_report_context_stops_local_sync_failover_on_transport_errors() {
async fn report_context_failover_policy_does_not_override_provider_config() {
let result = ExecutionResult {
request_id: "req-1".to_string(),
candidate_id: None,
status_code: 503,
status_code: 429,
headers: Default::default(),
body: None,
telemetry: None,
@@ -1150,7 +1270,7 @@ mod tests {
let plan = sample_plan();
assert!(
should_stop_local_candidate_failover_sync(
should_retry_next_local_candidate_sync(
&state,
&plan,
"openai_image_sync",
@@ -1161,7 +1281,7 @@ mod tests {
.await
);
assert!(
!should_retry_next_local_candidate_sync(
!should_stop_local_candidate_failover_sync(
&state,
&plan,
"openai_image_sync",
@@ -3921,10 +3921,14 @@ mod tests {
StreamFrame, StreamFramePayload, StreamFrameType,
};
use aether_data::repository::candidates::InMemoryRequestCandidateRepository;
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
use aether_data::repository::usage::InMemoryUsageReadRepository;
use aether_data_contracts::repository::candidates::{
RequestCandidateReadRepository, RequestCandidateStatus,
};
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use aether_data_contracts::repository::usage::UsageReadRepository;
use aether_usage_runtime::UsageRuntimeConfig;
use async_stream::stream;
@@ -3954,6 +3958,77 @@ mod tests {
use crate::tunnel::{tunnel_protocol, TunnelProxyConn};
use crate::AppState;
fn provider_catalog_stop_429_for_plan(
plan: &ExecutionPlan,
) -> InMemoryProviderCatalogReadRepository {
let provider_type = plan.provider_name.as_deref().unwrap_or("custom");
let provider = StoredProviderCatalogProvider::new(
plan.provider_id.clone(),
plan.provider_id.clone(),
Some("https://provider.example".to_string()),
provider_type.to_string(),
)
.expect("provider should build")
.with_transport_fields(
true,
false,
false,
None,
Some(3),
None,
None,
None,
Some(json!({
"failover_rules": {
"stop_status_codes": [429]
}
})),
);
let endpoint = StoredProviderCatalogEndpoint::new(
plan.endpoint_id.clone(),
plan.provider_id.clone(),
plan.provider_api_format.clone(),
None,
None,
true,
)
.expect("endpoint should build")
.with_transport_fields(
"https://provider.example".to_string(),
None,
None,
Some(2),
None,
None,
None,
None,
)
.expect("endpoint transport should build");
let key = StoredProviderCatalogKey::new(
plan.key_id.clone(),
plan.provider_id.clone(),
plan.key_id.clone(),
"api_key".to_string(),
None,
true,
)
.expect("key should build")
.with_transport_fields(
Some(json!([plan.provider_api_format.clone()])),
"plain-upstream-key".to_string(),
None,
None,
Some(json!({ "openai:chat": 1 })),
None,
None,
None,
None,
)
.expect("key transport should build");
InMemoryProviderCatalogReadRepository::seed(vec![provider], vec![endpoint], vec![key])
}
fn test_decision() -> GatewayControlDecision {
GatewayControlDecision::synthetic(
"/v1/chat/completions",
@@ -5458,6 +5533,14 @@ mod tests {
transport_profile: None,
timeouts: None,
};
let state = state.with_data_state_for_tests(
crate::data::GatewayDataState::with_request_candidate_and_usage_repository_for_tests(
Arc::clone(&request_candidate_repository),
Arc::clone(&usage_repository),
)
.with_provider_catalog_reader(Arc::new(provider_catalog_stop_429_for_plan(&plan)))
.with_encryption_key_for_tests("development-key"),
);
let trailer_error = connect_json_frame(
2,
br#"{"error":{"code":"resource_exhausted","message":"an internal error occurred"}}"#,
@@ -5493,10 +5576,7 @@ mod tests {
"client_api_format": "claude:messages",
"needs_conversion": true,
"has_envelope": true,
"envelope_name": "windsurf:GetChatMessage",
"local_failover_policy": {
"stop_status_codes": [429]
}
"envelope_name": "windsurf:GetChatMessage"
})),
crate::clock::current_unix_ms(),
Instant::now(),
@@ -5590,6 +5670,14 @@ mod tests {
transport_profile: None,
timeouts: None,
};
let state = state.with_data_state_for_tests(
crate::data::GatewayDataState::with_request_candidate_and_usage_repository_for_tests(
Arc::clone(&request_candidate_repository),
Arc::clone(&usage_repository),
)
.with_provider_catalog_reader(Arc::new(provider_catalog_stop_429_for_plan(&plan)))
.with_encryption_key_for_tests("development-key"),
);
let connect_error = connect_json_frame(
2,
br#"{"error":{"code":"resource_exhausted","message":"quota exhausted"}}"#,
@@ -5624,10 +5712,7 @@ mod tests {
"client_api_format": "claude:messages",
"needs_conversion": true,
"has_envelope": true,
"envelope_name": "windsurf:GetChatMessage",
"local_failover_policy": {
"stop_status_codes": [429]
}
"envelope_name": "windsurf:GetChatMessage"
})),
crate::clock::current_unix_ms(),
Instant::now(),
@@ -61,11 +61,8 @@ pub(crate) fn classify_local_failover(
}
if input.status_code >= 400
&& input.response_text.is_some_and(|text| {
policy
.error_stop_patterns
.iter()
.any(|rule| local_failover_regex_rule_matches(rule, text, input.status_code))
&& policy.error_stop_patterns.iter().any(|rule| {
local_failover_regex_rule_matches(rule, input.response_text, input.status_code)
})
{
return LocalFailoverClassification::StopErrorPattern;
@@ -76,7 +73,7 @@ pub(crate) fn classify_local_failover(
policy
.success_failover_patterns
.iter()
.any(|rule| local_failover_regex_rule_matches(rule, text, input.status_code))
.any(|rule| local_failover_regex_rule_matches(rule, Some(text), input.status_code))
})
{
return LocalFailoverClassification::RetrySuccessPattern;
@@ -190,14 +187,23 @@ fn first_non_empty_json_text(
fn local_failover_regex_rule_matches(
rule: &LocalFailoverRegexRule,
response_text: &str,
response_text: Option<&str>,
status_code: u16,
) -> bool {
if !rule.status_codes.is_empty() && !rule.status_codes.contains(&status_code) {
return false;
}
Regex::new(&rule.pattern)
let pattern = rule.pattern.trim();
if pattern.is_empty() {
return !rule.status_codes.is_empty();
}
let Some(response_text) = response_text else {
return false;
};
Regex::new(pattern)
.ok()
.is_some_and(|regex| regex.is_match(response_text))
}
@@ -260,6 +266,50 @@ mod tests {
);
}
#[test]
fn classifier_detects_error_stop_pattern_without_status_codes_on_any_error_status() {
let policy = LocalFailoverPolicy {
error_stop_patterns: vec![LocalFailoverRegexRule {
pattern: "content_policy_violation".to_string(),
status_codes: BTreeSet::new(),
}],
..LocalFailoverPolicy::default()
};
for status_code in [400, 429, 503] {
assert_eq!(
classify_local_failover(
&policy,
LocalFailoverInput::new(
status_code,
Some("{\"error\":\"content_policy_violation\"}")
)
),
LocalFailoverClassification::StopErrorPattern
);
}
}
#[test]
fn classifier_detects_status_only_error_stop_rule_without_response_text() {
let policy = LocalFailoverPolicy {
error_stop_patterns: vec![LocalFailoverRegexRule {
pattern: String::new(),
status_codes: [429].into_iter().collect(),
}],
..LocalFailoverPolicy::default()
};
assert_eq!(
classify_local_failover(&policy, LocalFailoverInput::new(429, None)),
LocalFailoverClassification::StopErrorPattern
);
assert_eq!(
classify_local_failover(&policy, LocalFailoverInput::new(503, None)),
LocalFailoverClassification::RetryUpstreamFailure
);
}
#[test]
fn classifier_detects_success_continue_status_code() {
let policy = LocalFailoverPolicy {
+19 -30
View File
@@ -25,27 +25,8 @@ pub(crate) struct LocalFailoverRegexRule {
pub(crate) async fn resolve_local_failover_policy(
state: &AppState,
plan: &ExecutionPlan,
report_context: Option<&serde_json::Value>,
_report_context: Option<&serde_json::Value>,
) -> LocalFailoverPolicy {
if let Some(policy) = local_failover_policy_from_report_context(report_context) {
debug!(
event_name = "local_failover_policy_loaded",
log_type = "debug",
request_id = %plan.request_id,
provider_id = %plan.provider_id,
endpoint_id = %plan.endpoint_id,
key_id = %plan.key_id,
source = "report_context",
max_retries = ?policy.max_retries,
stop_status_code_count = policy.stop_status_codes.len(),
continue_status_code_count = policy.continue_status_codes.len(),
success_failover_pattern_count = policy.success_failover_patterns.len(),
error_stop_pattern_count = policy.error_stop_patterns.len(),
"gateway loaded local failover policy from report context"
);
return policy;
}
let transport = match state
.read_provider_transport_snapshot(&plan.provider_id, &plan.endpoint_id, &plan.key_id)
.await
@@ -201,31 +182,39 @@ fn parse_regex_rules(
rules: &serde_json::Map<String, serde_json::Value>,
key: &str,
) -> Vec<LocalFailoverRegexRule> {
let allow_status_only = key == "error_stop_patterns";
rules
.get(key)
.and_then(Value::as_array)
.into_iter()
.flat_map(|items| items.iter())
.filter_map(parse_regex_rule)
.filter_map(|value| parse_regex_rule(value, allow_status_only))
.collect()
}
fn parse_regex_rule(value: &serde_json::Value) -> Option<LocalFailoverRegexRule> {
fn parse_regex_rule(
value: &serde_json::Value,
allow_status_only: bool,
) -> Option<LocalFailoverRegexRule> {
let object = value.as_object()?;
let pattern = object
.get("pattern")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())?;
.unwrap_or_default();
let status_codes: BTreeSet<u16> = object
.get("status_codes")
.and_then(Value::as_array)
.into_iter()
.flat_map(|values| values.iter())
.filter_map(|value| parse_u64_value(value).and_then(|value| u16::try_from(value).ok()))
.collect();
if pattern.is_empty() && (!allow_status_only || status_codes.is_empty()) {
return None;
}
Some(LocalFailoverRegexRule {
pattern: pattern.to_string(),
status_codes: object
.get("status_codes")
.and_then(Value::as_array)
.into_iter()
.flat_map(|values| values.iter())
.filter_map(|value| parse_u64_value(value).and_then(|value| u16::try_from(value).ok()))
.collect(),
status_codes,
})
}