1use std::collections::BTreeMap;
14
15use serde::{Deserialize, Serialize};
16use turnframe_provider::request::ReasoningEffort;
17
18use crate::task::TaskKind;
19
20#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
22#[serde(rename_all = "snake_case")]
23#[non_exhaustive]
24pub enum Disagreement {
25 Escalate,
27 #[default]
29 Clarify,
30 Fail,
32 Reread,
34}
35
36#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
38#[serde(default, deny_unknown_fields)]
39#[non_exhaustive]
40pub struct TaskProfile {
41 pub enabled: bool,
43 pub model: Option<String>,
45 pub escalate_to: Option<String>,
47 pub temperature: Option<f32>,
49 pub vote_temperature: f32,
51 pub max_output_tokens: Option<u32>,
53 pub reasoning_effort: Option<ReasoningEffort>,
55 pub timeout_secs: Option<u64>,
57 pub votes: u8,
59 pub on_disagreement: Disagreement,
61 pub repairs: u8,
63 pub retries: u8,
66 pub review: bool,
68}
69
70impl Default for TaskProfile {
71 fn default() -> Self {
72 Self {
73 enabled: true,
74 model: None,
75 escalate_to: None,
76 temperature: Some(0.0),
77 vote_temperature: 0.7,
78 max_output_tokens: None,
79 reasoning_effort: Some(ReasoningEffort::Minimal),
80 timeout_secs: None,
81 votes: 1,
82 on_disagreement: Disagreement::Clarify,
83 repairs: 1,
84 retries: 1,
85 review: false,
86 }
87 }
88}
89
90impl TaskProfile {
91 #[must_use]
93 pub fn default_for(kind: TaskKind) -> Self {
94 let base = Self::default();
95 match kind {
96 TaskKind::Segment => Self {
97 max_output_tokens: Some(800),
98 ..base
99 },
100 TaskKind::Coverage => Self {
101 max_output_tokens: Some(200),
102 repairs: 0,
103 ..base
104 },
105 TaskKind::TakeUp | TaskKind::Route | TaskKind::Locate | TaskKind::QuestionFrame => {
106 Self {
107 max_output_tokens: Some(150),
108 ..base
109 }
110 }
111 TaskKind::Extract => Self {
112 max_output_tokens: Some(600),
113 ..base
114 },
115 TaskKind::Verify => Self {
118 max_output_tokens: Some(2_000),
119 reasoning_effort: Some(ReasoningEffort::Low),
120 repairs: 0,
121 ..base
122 },
123 TaskKind::CrossCheck => Self {
124 max_output_tokens: Some(400),
125 ..base
126 },
127 TaskKind::Respects => Self {
128 max_output_tokens: Some(200),
129 ..base
130 },
131 TaskKind::Investigate => Self {
132 enabled: false,
133 max_output_tokens: Some(400),
134 ..base
135 },
136 TaskKind::Acknowledge => Self {
137 temperature: None,
138 reasoning_effort: None,
139 max_output_tokens: Some(300),
140 review: true,
141 ..base
142 },
143 TaskKind::Answer => Self {
144 temperature: None,
145 reasoning_effort: None,
146 ..base
147 },
148 TaskKind::Review => Self {
149 max_output_tokens: Some(300),
150 ..base
151 },
152 TaskKind::Progress => Self {
154 temperature: None,
155 reasoning_effort: None,
156 max_output_tokens: Some(80),
157 repairs: 0,
158 retries: 0,
159 ..base
160 },
161 _ => base,
162 }
163 }
164
165 #[must_use]
167 pub fn with_votes(mut self, votes: u8) -> Self {
168 self.votes = votes.max(1);
169 self
170 }
171
172 #[must_use]
174 pub fn on_model(mut self, tag: impl Into<String>) -> Self {
175 self.model = Some(tag.into());
176 self
177 }
178
179 #[must_use]
181 pub fn escalating_to(mut self, tag: impl Into<String>) -> Self {
182 self.escalate_to = Some(tag.into());
183 self
184 }
185
186 #[must_use]
188 pub fn on_disagreement(mut self, disagreement: Disagreement) -> Self {
189 self.on_disagreement = disagreement;
190 self
191 }
192
193 #[must_use]
195 pub fn with_repairs(mut self, repairs: u8) -> Self {
196 self.repairs = repairs;
197 self
198 }
199
200 #[must_use]
202 pub fn enabled(mut self, enabled: bool) -> Self {
203 self.enabled = enabled;
204 self
205 }
206
207 #[must_use]
209 pub fn with_reasoning(mut self, reasoning: Option<ReasoningEffort>) -> Self {
210 self.reasoning_effort = reasoning;
211 self
212 }
213
214 #[must_use]
216 pub fn with_review(mut self, review: bool) -> Self {
217 self.review = review;
218 self
219 }
220}
221
222#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
226#[serde(transparent)]
227pub struct TaskProfiles {
228 overrides: BTreeMap<String, TaskProfile>,
229}
230
231impl TaskProfiles {
232 #[must_use]
234 pub const fn new() -> Self {
235 Self {
236 overrides: BTreeMap::new(),
237 }
238 }
239
240 #[must_use]
242 pub fn get(&self, kind: TaskKind) -> TaskProfile {
243 self.overrides
244 .get(kind.as_str())
245 .cloned()
246 .unwrap_or_else(|| TaskProfile::default_for(kind))
247 }
248
249 #[must_use]
251 pub fn with(mut self, kind: TaskKind, profile: TaskProfile) -> Self {
252 self.overrides.insert(kind.as_str().to_owned(), profile);
253 self
254 }
255
256 #[must_use]
258 pub fn adjust(self, kind: TaskKind, change: impl FnOnce(TaskProfile) -> TaskProfile) -> Self {
259 let current = self.get(kind);
260 self.with(kind, change(current))
261 }
262
263 pub fn overridden(&self) -> impl Iterator<Item = &str> {
265 self.overrides.keys().map(String::as_str)
266 }
267
268 #[must_use]
270 pub fn tags(&self) -> Vec<String> {
271 let mut tags: Vec<String> = TaskKind::ALL
272 .iter()
273 .map(|kind| self.get(*kind))
274 .flat_map(|profile| [profile.model, profile.escalate_to])
275 .flatten()
276 .collect();
277 tags.sort();
278 tags.dedup();
279 tags
280 }
281}
282
283#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
285#[serde(default, deny_unknown_fields)]
286#[non_exhaustive]
287pub struct ProfileChange {
288 pub enabled: Option<bool>,
290 pub model: Option<String>,
292 pub escalate_to: Option<String>,
294 pub max_output_tokens: Option<u32>,
296 pub reasoning_effort: Option<ReasoningEffort>,
298 pub timeout_secs: Option<u64>,
300 pub votes: Option<u8>,
302 pub on_disagreement: Option<Disagreement>,
304 pub repairs: Option<u8>,
306 pub retries: Option<u8>,
308 pub review: Option<bool>,
310}
311
312impl ProfileChange {
313 #[must_use]
315 pub fn apply(&self, mut profile: TaskProfile) -> TaskProfile {
316 let change = self.clone();
317 if let Some(value) = change.enabled {
318 profile.enabled = value;
319 }
320 if let Some(value) = change.model {
321 profile.model = Some(value);
322 }
323 if let Some(value) = change.escalate_to {
324 profile.escalate_to = Some(value);
325 }
326 if let Some(value) = change.max_output_tokens {
327 profile.max_output_tokens = Some(value);
328 }
329 if let Some(value) = change.reasoning_effort {
330 profile.reasoning_effort = Some(value);
331 }
332 if let Some(value) = change.timeout_secs {
333 profile.timeout_secs = Some(value);
334 }
335 if let Some(value) = change.votes {
336 profile.votes = value.max(1);
337 }
338 if let Some(value) = change.on_disagreement {
339 profile.on_disagreement = value;
340 }
341 if let Some(value) = change.repairs {
342 profile.repairs = value;
343 }
344 if let Some(value) = change.retries {
345 profile.retries = value;
346 }
347 if let Some(value) = change.review {
348 profile.review = value;
349 }
350 profile
351 }
352}
353
354#[derive(Debug, Clone, Default, PartialEq, Serialize)]
356#[serde(transparent)]
357pub struct ProfileChanges {
358 changes: BTreeMap<String, ProfileChange>,
359}
360
361impl ProfileChanges {
362 #[must_use]
364 pub const fn new() -> Self {
365 Self {
366 changes: BTreeMap::new(),
367 }
368 }
369
370 #[must_use]
372 pub fn with(mut self, kind: TaskKind, change: ProfileChange) -> Self {
373 self.changes.insert(kind.as_str().to_owned(), change);
374 self
375 }
376
377 #[must_use]
379 pub fn apply(&self, mut profiles: TaskProfiles) -> TaskProfiles {
380 for kind in TaskKind::ALL {
381 if let Some(change) = self.changes.get(kind.as_str()) {
382 profiles = profiles.adjust(kind, |profile| change.apply(profile));
383 }
384 }
385 profiles
386 }
387
388 #[must_use]
390 pub fn tags(&self) -> Vec<String> {
391 self.changes
392 .values()
393 .flat_map(|change| [change.model.clone(), change.escalate_to.clone()])
394 .flatten()
395 .collect()
396 }
397}
398
399impl<'de> Deserialize<'de> for ProfileChanges {
400 fn deserialize<D: serde::Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
401 let changes = BTreeMap::<String, ProfileChange>::deserialize(deserializer)?;
402 if let Some(unknown) = changes.keys().find(|name| {
403 !TaskKind::ALL
404 .iter()
405 .any(|kind| kind.as_str() == name.as_str())
406 }) {
407 return Err(serde::de::Error::custom(format!(
408 "`{unknown}` is not a task kind"
409 )));
410 }
411 Ok(Self { changes })
412 }
413}
414
415#[cfg(test)]
416mod tests {
417 use super::*;
418
419 #[test]
420 fn the_verifier_thinks_before_it_judges() {
421 let verify = TaskProfile::default_for(TaskKind::Verify);
422 assert_eq!(verify.reasoning_effort, Some(ReasoningEffort::Low));
423 assert!(
424 verify.max_output_tokens >= Some(2_000),
425 "reasoning counts against the cap"
426 );
427 assert_eq!(
428 TaskProfile::default_for(TaskKind::Extract).reasoning_effort,
429 Some(ReasoningEffort::Minimal)
430 );
431 }
432
433 #[test]
434 fn an_override_names_only_what_changes() {
435 let profiles: TaskProfiles = toml::from_str(
436 r#"
437 [route]
438 votes = 3
439 on_disagreement = "escalate"
440 escalate_to = "large"
441 "#,
442 )
443 .expect("parses");
444 let route = profiles.get(TaskKind::Route);
445 assert_eq!(route.votes, 3);
446 assert_eq!(route.on_disagreement, Disagreement::Escalate);
447 assert_eq!(route.escalate_to.as_deref(), Some("large"));
448 assert_eq!(profiles.get(TaskKind::Extract).max_output_tokens, Some(600));
449 assert_eq!(profiles.tags(), vec!["large".to_owned()]);
450 }
451
452 #[test]
453 fn a_change_keeps_what_it_does_not_name() {
454 let changes: ProfileChanges = toml::from_str(
455 r#"
456 [route]
457 votes = 3
458 "#,
459 )
460 .expect("parses");
461 let profiles = changes.apply(TaskProfiles::new());
462 let route = profiles.get(TaskKind::Route);
463 assert_eq!(route.votes, 3);
464 assert_eq!(
465 route.max_output_tokens,
466 Some(150),
467 "route's own cap survives"
468 );
469 }
470
471 #[test]
472 fn a_change_to_a_task_that_does_not_exist_is_refused_by_name() {
473 let refused = toml::from_str::<ProfileChanges>("[extrakt]\nvotes = 3\n").unwrap_err();
474 assert!(refused.to_string().contains("extrakt"), "{refused}");
475 }
476
477 #[test]
478 fn optional_tasks_ship_as_the_spec_says() {
479 let profiles = TaskProfiles::new();
480 assert!(profiles.get(TaskKind::Coverage).enabled);
481 assert!(!profiles.get(TaskKind::Investigate).enabled);
482 assert!(profiles.get(TaskKind::Acknowledge).review);
483 assert!(!profiles.get(TaskKind::Answer).review);
484 assert_eq!(profiles.get(TaskKind::Verify).repairs, 0);
485 }
486}