relay_knowledge/model_provider/
fallback.rs1use 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#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
15#[serde(rename_all = "snake_case")]
16pub enum ModelFallbackStrategy {
17 SameProviderThenOtherProvider,
18 OtherProviderOnly,
19}
20
21#[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#[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;