1use serde::{Deserialize, Serialize};
6
7pub const DEFAULT_OPENROUTER_MODEL: &str = "mistralai/mistral-small-2603";
9pub const DEFAULT_GEMINI_MODEL: &str = "gemini-3.1-flash-lite";
11
12#[derive(Debug, Clone, Copy, PartialEq, Eq)]
14pub enum TaskType {
15 Triage,
17 Review,
19 Create,
21}
22
23#[derive(Debug, Deserialize, Serialize, Default, Clone)]
25#[serde(default)]
26pub struct TaskOverride {
27 pub provider: Option<String>,
29 pub model: Option<String>,
31 #[serde(default)]
33 pub small_model: Option<String>,
34 #[serde(default)]
36 pub large_model: Option<String>,
37 #[serde(default)]
39 pub routing_threshold_chars: Option<usize>,
40}
41
42#[derive(Debug, Deserialize, Serialize, Default, Clone)]
44#[serde(default)]
45pub struct TasksConfig {
46 pub triage: Option<TaskOverride>,
48 pub review: Option<TaskOverride>,
50 pub create: Option<TaskOverride>,
52}
53
54#[derive(Debug, Clone, Serialize)]
56pub struct FallbackEntry {
57 pub provider: String,
59 pub model: Option<String>,
61}
62
63impl<'de> Deserialize<'de> for FallbackEntry {
64 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
65 where
66 D: serde::Deserializer<'de>,
67 {
68 #[derive(Deserialize)]
69 #[serde(untagged)]
70 enum EntryVariant {
71 String(String),
72 Struct {
73 provider: String,
74 model: Option<String>,
75 },
76 }
77
78 match EntryVariant::deserialize(deserializer)? {
79 EntryVariant::String(provider) => Ok(FallbackEntry {
80 provider,
81 model: None,
82 }),
83 EntryVariant::Struct { provider, model } => Ok(FallbackEntry { provider, model }),
84 }
85 }
86}
87
88#[derive(Debug, Deserialize, Serialize, Clone, Default)]
90#[serde(default)]
91pub struct FallbackConfig {
92 pub chain: Vec<FallbackEntry>,
94}
95
96fn default_retry_max_attempts() -> u32 {
98 3
99}
100
101#[derive(Debug, Deserialize, Serialize, Clone)]
103#[serde(default)]
104pub struct AiConfig {
105 pub provider: String,
107 pub model: String,
109 pub timeout_seconds: u64,
111 pub allow_paid_models: bool,
113 pub max_tokens: u32,
115 pub temperature: f32,
117 pub circuit_breaker_threshold: u32,
119 pub circuit_breaker_reset_seconds: u64,
121 #[serde(default = "default_retry_max_attempts")]
123 pub retry_max_attempts: u32,
124 pub tasks: Option<TasksConfig>,
126 pub fallback: Option<FallbackConfig>,
128 pub custom_guidance: Option<String>,
134 pub validation_enabled: bool,
140}
141
142impl Default for AiConfig {
143 fn default() -> Self {
144 Self {
145 provider: "openrouter".to_string(),
146 model: DEFAULT_OPENROUTER_MODEL.to_string(),
147 timeout_seconds: 30,
148 allow_paid_models: true,
149 max_tokens: 4096,
150 temperature: 0.3,
151 circuit_breaker_threshold: 3,
152 circuit_breaker_reset_seconds: 60,
153 retry_max_attempts: default_retry_max_attempts(),
154 tasks: None,
155 fallback: None,
156 custom_guidance: None,
157 validation_enabled: true,
158 }
159 }
160}
161
162impl AiConfig {
163 #[must_use]
179 pub fn resolve_for_task(
180 &self,
181 task: TaskType,
182 estimated_size: Option<usize>,
183 ) -> (String, String) {
184 let task_override = match task {
185 TaskType::Triage => self.tasks.as_ref().and_then(|t| t.triage.as_ref()),
186 TaskType::Review => self.tasks.as_ref().and_then(|t| t.review.as_ref()),
187 TaskType::Create => self.tasks.as_ref().and_then(|t| t.create.as_ref()),
188 };
189
190 let provider = task_override
191 .and_then(|o| o.provider.clone())
192 .unwrap_or_else(|| self.provider.clone());
193
194 if let Some(model) = task_override.and_then(|o| o.model.clone()) {
196 return (provider, model);
197 }
198
199 if let (Some(small), Some(large), Some(size)) = (
201 task_override.and_then(|o| o.small_model.clone()),
202 task_override.and_then(|o| o.large_model.clone()),
203 estimated_size,
204 ) {
205 let default_threshold = match task {
206 TaskType::Review => 60_000,
207 TaskType::Triage | TaskType::Create => 8_192,
208 };
209 let threshold = task_override
210 .and_then(|o| o.routing_threshold_chars)
211 .unwrap_or(default_threshold);
212 if size < threshold {
213 return (provider.clone(), small);
214 }
215 return (provider, large);
216 }
217
218 let has_small = task_override.and_then(|o| o.small_model.as_ref()).is_some();
220 let has_large = task_override.and_then(|o| o.large_model.as_ref()).is_some();
221 if has_small != has_large {
222 let missing = if has_small {
223 "large_model"
224 } else {
225 "small_model"
226 };
227 tracing::warn!(
228 "TaskOverride has only one routing field set (missing {missing}); falling back to default model"
229 );
230 }
231
232 let model = self.model.clone();
234 (provider, model)
235 }
236}