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 pub openrouter_data_collection: String,
142 pub openrouter_zdr: bool,
144}
145
146impl Default for AiConfig {
147 fn default() -> Self {
148 Self {
149 provider: "openrouter".to_string(),
150 model: DEFAULT_OPENROUTER_MODEL.to_string(),
151 timeout_seconds: 30,
152 allow_paid_models: true,
153 max_tokens: 4096,
154 temperature: 0.3,
155 circuit_breaker_threshold: 3,
156 circuit_breaker_reset_seconds: 60,
157 retry_max_attempts: default_retry_max_attempts(),
158 tasks: None,
159 fallback: None,
160 custom_guidance: None,
161 validation_enabled: true,
162 openrouter_data_collection: "deny".to_string(),
163 openrouter_zdr: true,
164 }
165 }
166}
167
168impl AiConfig {
169 #[must_use]
185 pub fn resolve_for_task(
186 &self,
187 task: TaskType,
188 estimated_size: Option<usize>,
189 ) -> (String, String) {
190 let task_override = match task {
191 TaskType::Triage => self.tasks.as_ref().and_then(|t| t.triage.as_ref()),
192 TaskType::Review => self.tasks.as_ref().and_then(|t| t.review.as_ref()),
193 TaskType::Create => self.tasks.as_ref().and_then(|t| t.create.as_ref()),
194 };
195
196 let provider = task_override
197 .and_then(|o| o.provider.clone())
198 .unwrap_or_else(|| self.provider.clone());
199
200 if let Some(model) = task_override.and_then(|o| o.model.clone()) {
202 return (provider, model);
203 }
204
205 if let (Some(small), Some(large), Some(size)) = (
207 task_override.and_then(|o| o.small_model.clone()),
208 task_override.and_then(|o| o.large_model.clone()),
209 estimated_size,
210 ) {
211 let default_threshold = match task {
212 TaskType::Review => 60_000,
213 TaskType::Triage | TaskType::Create => 8_192,
214 };
215 let threshold = task_override
216 .and_then(|o| o.routing_threshold_chars)
217 .unwrap_or(default_threshold);
218 if size < threshold {
219 return (provider.clone(), small);
220 }
221 return (provider, large);
222 }
223
224 let has_small = task_override.and_then(|o| o.small_model.as_ref()).is_some();
226 let has_large = task_override.and_then(|o| o.large_model.as_ref()).is_some();
227 if has_small != has_large {
228 let missing = if has_small {
229 "large_model"
230 } else {
231 "small_model"
232 };
233 tracing::warn!(
234 "TaskOverride has only one routing field set (missing {missing}); falling back to default model"
235 );
236 }
237
238 let model = self.model.clone();
240 (provider, model)
241 }
242}