mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-07 09:57:47 +08:00
feat(gateway): configure cyber policy failover
This commit is contained in:
@@ -135,10 +135,7 @@ fn json_value_has_cyber_policy_code(value: &Value, depth: usize) -> bool {
|
||||
}
|
||||
match value {
|
||||
Value::Object(object) => object.iter().any(|(key, value)| {
|
||||
(key == "code"
|
||||
&& value
|
||||
.as_str()
|
||||
.is_some_and(|code| code.eq_ignore_ascii_case("cyber_policy")))
|
||||
(key.eq_ignore_ascii_case("code") && value.as_str().is_some_and(is_cyber_policy_code))
|
||||
|| json_value_has_cyber_policy_code(value, depth + 1)
|
||||
}),
|
||||
Value::Array(values) => values
|
||||
@@ -157,6 +154,11 @@ fn json_value_has_cyber_policy_code(value: &Value, depth: usize) -> bool {
|
||||
}
|
||||
}
|
||||
|
||||
fn is_cyber_policy_code(code: &str) -> bool {
|
||||
let code = code.trim();
|
||||
code.eq_ignore_ascii_case("cyber_policy") || code.eq_ignore_ascii_case("cyber_policy_violation")
|
||||
}
|
||||
|
||||
fn parse_local_error_response(response_text: Option<&str>) -> ParsedLocalErrorResponse {
|
||||
let raw = response_text
|
||||
.map(str::trim)
|
||||
@@ -382,6 +384,16 @@ mod tests {
|
||||
),
|
||||
LocalFailoverClassification::StopCyberPolicy
|
||||
);
|
||||
assert_eq!(
|
||||
classify_local_failover(
|
||||
&policy,
|
||||
LocalFailoverInput::new(
|
||||
400,
|
||||
Some(r#"{"error":{"code":"cyber_policy_violation"}}"#)
|
||||
)
|
||||
),
|
||||
LocalFailoverClassification::StopCyberPolicy
|
||||
);
|
||||
assert_eq!(
|
||||
classify_local_failover(
|
||||
&policy,
|
||||
@@ -403,9 +415,13 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn classifier_retries_cyber_policy_when_policy_disabled() {
|
||||
let policy = LocalFailoverPolicy {
|
||||
stop_cyber_policy_errors: false,
|
||||
..LocalFailoverPolicy::default()
|
||||
};
|
||||
assert_eq!(
|
||||
classify_local_failover(
|
||||
&LocalFailoverPolicy::default(),
|
||||
&policy,
|
||||
LocalFailoverInput::new(
|
||||
400,
|
||||
Some(r#"{"error":{"code":"cyber_policy","message":"flagged"}}"#)
|
||||
|
||||
@@ -39,8 +39,9 @@ pub(crate) use self::health::{
|
||||
};
|
||||
pub(crate) use self::policy::{
|
||||
append_local_failover_policy_to_value, codex_cyber_flag_passthrough_enabled,
|
||||
local_failover_policy_from_report_context, local_failover_policy_from_transport,
|
||||
resolve_local_failover_policy, LocalFailoverPolicy, LocalFailoverRegexRule,
|
||||
cyber_continue_failover_enabled, local_failover_policy_from_report_context,
|
||||
local_failover_policy_from_transport, resolve_local_failover_policy, LocalFailoverPolicy,
|
||||
LocalFailoverRegexRule, CYBER_CONTINUE_FAILOVER_CONFIG_KEY,
|
||||
};
|
||||
pub(crate) use self::recovery::{
|
||||
analyze_local_failover, recover_local_failover_decision, LocalFailoverAnalysis,
|
||||
|
||||
@@ -7,6 +7,8 @@ use tracing::debug;
|
||||
use crate::provider_transport::GatewayProviderTransportSnapshot;
|
||||
use crate::AppState;
|
||||
|
||||
pub(crate) const CYBER_CONTINUE_FAILOVER_CONFIG_KEY: &str = "cyber_continue_failover";
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub(crate) struct LocalFailoverPolicy {
|
||||
pub(crate) max_retries: Option<u64>,
|
||||
@@ -26,7 +28,7 @@ impl Default for LocalFailoverPolicy {
|
||||
continue_status_codes: BTreeSet::new(),
|
||||
success_failover_patterns: Vec::new(),
|
||||
error_stop_patterns: Vec::new(),
|
||||
stop_cyber_policy_errors: false,
|
||||
stop_cyber_policy_errors: true,
|
||||
retry_client_errors_by_default: true,
|
||||
}
|
||||
}
|
||||
@@ -43,14 +45,15 @@ pub(crate) async fn resolve_local_failover_policy(
|
||||
plan: &ExecutionPlan,
|
||||
_report_context: Option<&serde_json::Value>,
|
||||
) -> LocalFailoverPolicy {
|
||||
let transport = match state
|
||||
let mut policy = match state
|
||||
.read_provider_transport_snapshot(&plan.provider_id, &plan.endpoint_id, &plan.key_id)
|
||||
.await
|
||||
{
|
||||
Ok(Some(transport)) => transport,
|
||||
Ok(None) | Err(_) => return LocalFailoverPolicy::default(),
|
||||
Ok(Some(transport)) => local_failover_policy_from_transport(&transport),
|
||||
Ok(None) | Err(_) => LocalFailoverPolicy::default(),
|
||||
};
|
||||
let policy = local_failover_policy_from_transport(&transport);
|
||||
let cyber_continue_failover = cyber_continue_failover_enabled(state).await;
|
||||
policy.stop_cyber_policy_errors = !cyber_continue_failover;
|
||||
debug!(
|
||||
event_name = "local_failover_policy_loaded",
|
||||
log_type = "debug",
|
||||
@@ -64,11 +67,23 @@ pub(crate) async fn resolve_local_failover_policy(
|
||||
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(),
|
||||
cyber_continue_failover,
|
||||
"gateway loaded local failover policy from transport snapshot"
|
||||
);
|
||||
policy
|
||||
}
|
||||
|
||||
pub(crate) async fn cyber_continue_failover_enabled(state: &AppState) -> bool {
|
||||
state
|
||||
.read_system_config_json_value(CYBER_CONTINUE_FAILOVER_CONFIG_KEY)
|
||||
.await
|
||||
.ok()
|
||||
.flatten()
|
||||
.as_ref()
|
||||
.and_then(Value::as_bool)
|
||||
.unwrap_or(false)
|
||||
}
|
||||
|
||||
pub(crate) fn local_failover_policy_from_transport(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> LocalFailoverPolicy {
|
||||
@@ -100,10 +115,7 @@ pub(crate) fn local_failover_policy_from_transport(
|
||||
crate::ai_serving::api_format_defaults_to_client_error_failover(
|
||||
&transport.endpoint.api_format,
|
||||
),
|
||||
stop_cyber_policy_errors: codex_cyber_flag_passthrough_enabled(
|
||||
&transport.provider.provider_type,
|
||||
transport.provider.config.as_ref(),
|
||||
),
|
||||
stop_cyber_policy_errors: true,
|
||||
stop_status_codes: rules
|
||||
.map(|value| {
|
||||
parse_status_code_set(
|
||||
@@ -162,7 +174,7 @@ pub(crate) fn local_failover_policy_from_report_context(
|
||||
stop_cyber_policy_errors: object
|
||||
.get("stop_cyber_policy_errors")
|
||||
.and_then(Value::as_bool)
|
||||
.unwrap_or(false),
|
||||
.unwrap_or(true),
|
||||
retry_client_errors_by_default: object
|
||||
.get("retry_client_errors_by_default")
|
||||
.and_then(Value::as_bool)
|
||||
@@ -398,7 +410,7 @@ mod tests {
|
||||
pattern: "validation".to_string(),
|
||||
status_codes: [422].into_iter().collect(),
|
||||
}],
|
||||
stop_cyber_policy_errors: false,
|
||||
stop_cyber_policy_errors: true,
|
||||
retry_client_errors_by_default: true,
|
||||
})
|
||||
);
|
||||
@@ -420,7 +432,7 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn codex_cyber_policy_passthrough_defaults_on_and_can_be_disabled() {
|
||||
fn transport_policy_defaults_to_stopping_cyber_policy() {
|
||||
let mut transport = sample_transport(None, None, None);
|
||||
transport.provider.provider_type = "codex".to_string();
|
||||
assert!(local_failover_policy_from_transport(&transport).stop_cyber_policy_errors);
|
||||
@@ -428,7 +440,7 @@ mod tests {
|
||||
transport.provider.config = Some(json!({
|
||||
"codex": {"pass_through_cyber_flag_interrupt": false}
|
||||
}));
|
||||
assert!(!local_failover_policy_from_transport(&transport).stop_cyber_policy_errors);
|
||||
assert!(local_failover_policy_from_transport(&transport).stop_cyber_policy_errors);
|
||||
|
||||
transport.provider.config = Some(json!({
|
||||
"codex": {"passthrough_cyber_flag_interrupt": true}
|
||||
@@ -436,6 +448,6 @@ mod tests {
|
||||
assert!(local_failover_policy_from_transport(&transport).stop_cyber_policy_errors);
|
||||
|
||||
transport.provider.provider_type = "llm".to_string();
|
||||
assert!(!local_failover_policy_from_transport(&transport).stop_cyber_policy_errors);
|
||||
assert!(local_failover_policy_from_transport(&transport).stop_cyber_policy_errors);
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user