mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-08 10:27:46 +08:00
feat: improve failover rules and request timeline
This commit is contained in:
@@ -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,
|
||||
|
||||
-9
@@ -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 {
|
||||
|
||||
@@ -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,
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user