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:
@@ -0,0 +1,120 @@
|
||||
use std::collections::BTreeSet;
|
||||
|
||||
use regex::Regex;
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
pub const MAX_ROUTING_FAILOVER_RULES: usize = 64;
|
||||
pub const MAX_ROUTING_FAILOVER_PATTERN_BYTES: usize = 4096;
|
||||
|
||||
#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(default)]
|
||||
pub struct RoutingFailoverRule {
|
||||
pub pattern: String,
|
||||
pub status_codes: BTreeSet<u16>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(default)]
|
||||
pub struct RoutingFailoverRules {
|
||||
pub success_failover_patterns: Vec<RoutingFailoverRule>,
|
||||
pub error_stop_patterns: Vec<RoutingFailoverRule>,
|
||||
}
|
||||
|
||||
pub fn validate_routing_failover_rules(rules: &RoutingFailoverRules) -> Result<(), String> {
|
||||
for (name, entries, success) in [
|
||||
(
|
||||
"success_failover_patterns",
|
||||
&rules.success_failover_patterns,
|
||||
true,
|
||||
),
|
||||
("error_stop_patterns", &rules.error_stop_patterns, false),
|
||||
] {
|
||||
if entries.len() > MAX_ROUTING_FAILOVER_RULES {
|
||||
return Err(format!("{name} exceeds {MAX_ROUTING_FAILOVER_RULES} rules"));
|
||||
}
|
||||
for (index, rule) in entries.iter().enumerate() {
|
||||
let pattern = rule.pattern.trim();
|
||||
if pattern.is_empty() && (success || rule.status_codes.is_empty()) {
|
||||
return Err(format!(
|
||||
"{name}[{index}] requires a pattern or error status codes"
|
||||
));
|
||||
}
|
||||
if pattern.len() > MAX_ROUTING_FAILOVER_PATTERN_BYTES {
|
||||
return Err(format!(
|
||||
"{name}[{index}] pattern exceeds {MAX_ROUTING_FAILOVER_PATTERN_BYTES} bytes"
|
||||
));
|
||||
}
|
||||
if !pattern.is_empty() {
|
||||
Regex::new(pattern)
|
||||
.map_err(|error| format!("{name}[{index}] invalid regex: {error}"))?;
|
||||
}
|
||||
if rule.status_codes.iter().any(|status| {
|
||||
if success {
|
||||
*status != 200
|
||||
} else {
|
||||
!(400..=599).contains(status)
|
||||
}
|
||||
}) {
|
||||
return Err(format!("{name}[{index}] contains invalid status codes"));
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn validates_regex_and_status_only_stop_rules() {
|
||||
let rules = RoutingFailoverRules {
|
||||
success_failover_patterns: vec![RoutingFailoverRule {
|
||||
pattern: "(?i)capacity.*exhausted".to_string(),
|
||||
..Default::default()
|
||||
}],
|
||||
error_stop_patterns: vec![RoutingFailoverRule {
|
||||
status_codes: [400, 413].into_iter().collect(),
|
||||
..Default::default()
|
||||
}],
|
||||
..Default::default()
|
||||
};
|
||||
assert!(validate_routing_failover_rules(&rules).is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_invalid_or_unbounded_rule_configuration() {
|
||||
for rule in [
|
||||
RoutingFailoverRule::default(),
|
||||
RoutingFailoverRule {
|
||||
pattern: "[".to_string(),
|
||||
..Default::default()
|
||||
},
|
||||
RoutingFailoverRule {
|
||||
pattern: "error".to_string(),
|
||||
status_codes: [429].into_iter().collect(),
|
||||
},
|
||||
RoutingFailoverRule {
|
||||
pattern: "a".repeat(MAX_ROUTING_FAILOVER_PATTERN_BYTES + 1),
|
||||
..Default::default()
|
||||
},
|
||||
] {
|
||||
let rules = RoutingFailoverRules {
|
||||
success_failover_patterns: vec![rule],
|
||||
..Default::default()
|
||||
};
|
||||
assert!(validate_routing_failover_rules(&rules).is_err());
|
||||
}
|
||||
let rules = RoutingFailoverRules {
|
||||
error_stop_patterns: vec![
|
||||
RoutingFailoverRule {
|
||||
status_codes: [400].into_iter().collect(),
|
||||
..Default::default()
|
||||
};
|
||||
MAX_ROUTING_FAILOVER_RULES + 1
|
||||
],
|
||||
..Default::default()
|
||||
};
|
||||
assert!(validate_routing_failover_rules(&rules).is_err());
|
||||
}
|
||||
}
|
||||
@@ -1,5 +1,6 @@
|
||||
mod actions;
|
||||
mod conditions;
|
||||
mod failover;
|
||||
mod model;
|
||||
mod mutations;
|
||||
mod policy;
|
||||
@@ -12,6 +13,10 @@ pub use actions::{
|
||||
RoutingSchedulingMode, RoutingSetPriorityMode,
|
||||
};
|
||||
pub use conditions::{RoutingCondition, RoutingConditionContext, RoutingConditionOp};
|
||||
pub use failover::{
|
||||
validate_routing_failover_rules, RoutingFailoverRule, RoutingFailoverRules,
|
||||
MAX_ROUTING_FAILOVER_PATTERN_BYTES, MAX_ROUTING_FAILOVER_RULES,
|
||||
};
|
||||
pub use model::{
|
||||
RoutingDefaultPolicy, RoutingExecutionPolicy, RoutingGroupBinding, RoutingGroupBindingSubject,
|
||||
RoutingGroupConfig, RoutingGroupRecord, RoutingGroupVersionRecord, RoutingModelPolicy,
|
||||
|
||||
@@ -7,6 +7,7 @@ use crate::actions::{
|
||||
RoutingAction, RoutingRulePhase, RoutingSchedulingMode, RoutingSetPriorityMode,
|
||||
};
|
||||
use crate::conditions::RoutingCondition;
|
||||
use crate::RoutingFailoverRules;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct RoutingSchedulingPreset {
|
||||
@@ -33,7 +34,7 @@ pub const DEFAULT_STICKY_KEY_ATTEMPTS: u32 = 2;
|
||||
/// transport configuration. A resolved policy is snapshotted for the request
|
||||
/// and can therefore be consumed by execution without rereading mutable
|
||||
/// system settings.
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Default)]
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Default)]
|
||||
pub struct RoutingExecutionPolicy {
|
||||
#[serde(default, skip_serializing_if = "is_false")]
|
||||
pub enable_cf_heartbeat: bool,
|
||||
@@ -41,6 +42,12 @@ pub struct RoutingExecutionPolicy {
|
||||
pub cyber_continue_failover: bool,
|
||||
#[serde(default, skip_serializing_if = "is_false")]
|
||||
pub cancel_on_client_disconnect: bool,
|
||||
#[serde(default)]
|
||||
pub max_transfer_count: u64,
|
||||
#[serde(default)]
|
||||
pub max_transfer_timeout_seconds: u64,
|
||||
#[serde(default)]
|
||||
pub failover_rules: RoutingFailoverRules,
|
||||
}
|
||||
|
||||
impl<'de> Deserialize<'de> for RoutingExecutionPolicy {
|
||||
@@ -60,6 +67,12 @@ impl<'de> Deserialize<'de> for RoutingExecutionPolicy {
|
||||
cyber_continue_failover: bool,
|
||||
#[serde(default)]
|
||||
cancel_on_client_disconnect: bool,
|
||||
#[serde(default)]
|
||||
max_transfer_count: u64,
|
||||
#[serde(default)]
|
||||
max_transfer_timeout_seconds: u64,
|
||||
#[serde(default)]
|
||||
failover_rules: RoutingFailoverRules,
|
||||
}
|
||||
|
||||
let value = LegacyCompatibleExecutionPolicy::deserialize(deserializer)?;
|
||||
@@ -69,6 +82,9 @@ impl<'de> Deserialize<'de> for RoutingExecutionPolicy {
|
||||
|| value.enable_standard_text_sync_heartbeat,
|
||||
cyber_continue_failover: value.cyber_continue_failover,
|
||||
cancel_on_client_disconnect: value.cancel_on_client_disconnect,
|
||||
max_transfer_count: value.max_transfer_count,
|
||||
max_transfer_timeout_seconds: value.max_transfer_timeout_seconds,
|
||||
failover_rules: value.failover_rules,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -116,6 +132,29 @@ fn is_false(value: &bool) -> bool {
|
||||
mod execution_policy_tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn routing_failover_configuration_round_trips_and_validates() {
|
||||
let config: RoutingGroupConfig = serde_json::from_value(serde_json::json!({
|
||||
"default_policy": {
|
||||
"max_transfer_count": 3,
|
||||
"max_transfer_timeout_seconds": 90,
|
||||
"failover_rules": {
|
||||
"success_failover_patterns": [{ "pattern": "(?i)capacity" }],
|
||||
"error_stop_patterns": [{ "status_codes": [400, 413] }]
|
||||
}
|
||||
}
|
||||
}))
|
||||
.unwrap();
|
||||
crate::validate_routing_group_config(&config).unwrap();
|
||||
assert_eq!(config.default_policy.execution_policy.max_transfer_count, 3);
|
||||
let value = serde_json::to_value(&config).unwrap();
|
||||
assert_eq!(value["default_policy"]["max_transfer_timeout_seconds"], 90);
|
||||
assert_eq!(
|
||||
serde_json::from_value::<RoutingGroupConfig>(value).unwrap(),
|
||||
config
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cancellation_defaults_off_and_round_trips_with_legacy_heartbeat() {
|
||||
let default: RoutingDefaultPolicy = serde_json::from_str("{}").unwrap();
|
||||
|
||||
@@ -89,7 +89,7 @@ pub fn resolve_routing_policy(
|
||||
scheduling_mode: config.default_policy.scheduling_mode,
|
||||
keep_priority_on_conversion: config.default_policy.keep_priority_on_conversion,
|
||||
sticky_key_attempts: config.default_policy.sticky_key_attempts,
|
||||
execution_policy: config.default_policy.execution_policy,
|
||||
execution_policy: config.default_policy.execution_policy.clone(),
|
||||
ranking_overlay: RankingOverlay::default(),
|
||||
mutation_plan: MutationPlan::default(),
|
||||
pool_policy_overrides: BTreeMap::new(),
|
||||
|
||||
@@ -28,6 +28,8 @@ const ROUTING_POOL_PRESETS: &[&str] = &[
|
||||
|
||||
#[derive(Debug, Error, Clone, PartialEq, Eq)]
|
||||
pub enum RoutingValidationError {
|
||||
#[error("routing failover rules are invalid: {0}")]
|
||||
InvalidFailoverRules(String),
|
||||
#[error("routing rule id is empty")]
|
||||
EmptyRuleId,
|
||||
#[error("duplicate routing rule id: {0}")]
|
||||
@@ -69,6 +71,8 @@ pub enum RoutingValidationError {
|
||||
pub fn validate_routing_group_config(
|
||||
config: &RoutingGroupConfig,
|
||||
) -> Result<(), RoutingValidationError> {
|
||||
crate::validate_routing_failover_rules(&config.default_policy.execution_policy.failover_rules)
|
||||
.map_err(RoutingValidationError::InvalidFailoverRules)?;
|
||||
let mut rule_ids = BTreeSet::new();
|
||||
for model_policy in &config.model_policies {
|
||||
if model_policy.model.trim().is_empty() {
|
||||
|
||||
Reference in New Issue
Block a user