Skip to main content

relay_knowledge/model_provider/
fallback.rs

1//! Owns model fallback defaults and policy validation.
2
3use std::collections::BTreeSet;
4
5use serde::{Deserialize, Serialize};
6use tokio::fs;
7
8use super::{
9    ModelProviderConfigService, ModelProviderError, persistence::write_json,
10    profile_config::validate_profile_name,
11};
12
13/// Built-in fallback policy strategy.
14#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
15#[serde(rename_all = "snake_case")]
16pub enum ModelFallbackStrategy {
17    SameProviderThenOtherProvider,
18    OtherProviderOnly,
19}
20
21/// Model fallback policy used after retryable provider failures.
22#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
23pub struct ModelFallbackPolicy {
24    pub policy_id: String,
25    pub name: String,
26    pub description: String,
27    pub enabled: bool,
28    pub strategy: ModelFallbackStrategy,
29    pub max_hops: u32,
30    pub cooldown_seconds: u32,
31}
32
33/// Fallback config returned by Settings.
34#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
35pub struct ModelFallbackConfig {
36    pub policies: Vec<ModelFallbackPolicy>,
37}
38
39fn default_fallback() -> ModelFallbackConfig {
40    ModelFallbackConfig {
41        policies: vec![
42            ModelFallbackPolicy {
43                policy_id: "same_provider_then_other_provider".to_owned(),
44                name: "Same Provider Then Other Provider".to_owned(),
45                description: "Retry same-provider alternatives before switching providers."
46                    .to_owned(),
47                enabled: true,
48                strategy: ModelFallbackStrategy::SameProviderThenOtherProvider,
49                max_hops: 3,
50                cooldown_seconds: 60,
51            },
52            ModelFallbackPolicy {
53                policy_id: "other_provider_only".to_owned(),
54                name: "Other Provider Only".to_owned(),
55                description: "Fail over directly to profiles from other providers.".to_owned(),
56                enabled: true,
57                strategy: ModelFallbackStrategy::OtherProviderOnly,
58                max_hops: 3,
59                cooldown_seconds: 60,
60            },
61        ],
62    }
63}
64
65fn validate_fallback_config(config: &ModelFallbackConfig) -> Result<(), ModelProviderError> {
66    let mut ids = BTreeSet::new();
67    for policy in &config.policies {
68        let id = validate_profile_name(&policy.policy_id)?;
69        if !ids.insert(id.clone()) {
70            return Err(ModelProviderError::InvalidInput(format!(
71                "duplicate fallback policy id '{id}'"
72            )));
73        }
74        if policy.max_hops == 0 || policy.cooldown_seconds > 3600 {
75            return Err(ModelProviderError::InvalidInput(
76                "fallback policy max_hops must be positive and cooldown_seconds <= 3600".to_owned(),
77            ));
78        }
79    }
80    Ok(())
81}
82
83impl ModelProviderConfigService {
84    pub async fn fallback_config(&self) -> Result<ModelFallbackConfig, ModelProviderError> {
85        match fs::read_to_string(self.paths.model_fallback_file()).await {
86            Ok(raw) => serde_json::from_str(&raw).map_err(ModelProviderError::from),
87            Err(error) if error.kind() == std::io::ErrorKind::NotFound => Ok(default_fallback()),
88            Err(error) => Err(ModelProviderError::from(error)),
89        }
90    }
91
92    pub async fn save_fallback_config(
93        &self,
94        config: ModelFallbackConfig,
95    ) -> Result<ModelFallbackConfig, ModelProviderError> {
96        validate_fallback_config(&config)?;
97        write_json(self.paths.model_fallback_file(), &config).await?;
98        Ok(config)
99    }
100}
101
102#[cfg(test)]
103#[path = "fallback_tests.rs"]
104mod tests;