Skip to main content

turnframe_tasks/
profile.rs

1//! How each task kind is run: which model, which settings, how many votes and repairs.
2//!
3//! Every field has a default per kind (see [`TaskProfile::default_for`]), and a
4//! deployment overrides only what it names, in code or in TOML:
5//!
6//! ```toml
7//! [route]
8//! votes = 3
9//! on_disagreement = "escalate"
10//! escalate_to = "large"
11//! ```
12
13use std::collections::BTreeMap;
14
15use serde::{Deserialize, Serialize};
16use turnframe_provider::request::ReasoningEffort;
17
18use crate::task::TaskKind;
19
20/// What a vote without a strict majority leads to.
21#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
22#[serde(rename_all = "snake_case")]
23#[non_exhaustive]
24pub enum Disagreement {
25    /// Run once more on the escalation model; without one, ask.
26    Escalate,
27    /// Hand the answers back so the caller can ask the user.
28    #[default]
29    Clarify,
30    /// Treat the task as failed.
31    Fail,
32    /// Run once more, shown the answers that disagreed; that answer stands.
33    Reread,
34}
35
36/// How one task kind is run.
37#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
38#[serde(default, deny_unknown_fields)]
39#[non_exhaustive]
40pub struct TaskProfile {
41    /// Whether the task runs at all; optional tasks are switched off here.
42    pub enabled: bool,
43    /// The pool tag of the model that answers. `None` takes the pool's order.
44    pub model: Option<String>,
45    /// The pool tag of the model an escalation runs on. `None` never escalates.
46    pub escalate_to: Option<String>,
47    /// Sampling temperature. `None` keeps the provider's default.
48    pub temperature: Option<f32>,
49    /// Temperature for votes when more than one is cast.
50    pub vote_temperature: f32,
51    /// Output cap in tokens.
52    pub max_output_tokens: Option<u32>,
53    /// How much a reasoning model may think.
54    pub reasoning_effort: Option<ReasoningEffort>,
55    /// Deadline of one call in seconds; the turn's budget may shorten it.
56    pub timeout_secs: Option<u64>,
57    /// Answers sampled and compared; `1` means no vote.
58    pub votes: u8,
59    /// What a vote without a strict majority leads to.
60    pub on_disagreement: Disagreement,
61    /// Rounds that send a structurally wrong answer back with the error.
62    pub repairs: u8,
63    /// Times a call is sent again as it was, after a provider failure another try
64    /// can change: a refusal, a filter, a broken answer, a timeout.
65    pub retries: u8,
66    /// Whether a written block is reviewed before it is published.
67    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    /// The shipped profile of `kind`.
92    #[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            // The one judge of a reading: at minimal reasoning it misjudged multi-part
116            // messages; reasoning counts against the output cap, so it gets room.
117            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            // A progress line is a preview: one call, and nothing waits for it.
153            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    /// A copy casting `votes` answers.
166    #[must_use]
167    pub fn with_votes(mut self, votes: u8) -> Self {
168        self.votes = votes.max(1);
169        self
170    }
171
172    /// A copy answered by the model tagged `tag`.
173    #[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    /// A copy escalating to the model tagged `tag`.
180    #[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    /// A copy deciding a vote without a majority as `disagreement` says.
187    #[must_use]
188    pub fn on_disagreement(mut self, disagreement: Disagreement) -> Self {
189        self.on_disagreement = disagreement;
190        self
191    }
192
193    /// A copy with `repairs` repair rounds.
194    #[must_use]
195    pub fn with_repairs(mut self, repairs: u8) -> Self {
196        self.repairs = repairs;
197        self
198    }
199
200    /// A copy with the task switched on or off.
201    #[must_use]
202    pub fn enabled(mut self, enabled: bool) -> Self {
203        self.enabled = enabled;
204        self
205    }
206
207    /// A copy thinking as much as `reasoning` allows.
208    #[must_use]
209    pub fn with_reasoning(mut self, reasoning: Option<ReasoningEffort>) -> Self {
210        self.reasoning_effort = reasoning;
211        self
212    }
213
214    /// A copy whose written block is reviewed, or not.
215    #[must_use]
216    pub fn with_review(mut self, review: bool) -> Self {
217        self.review = review;
218        self
219    }
220}
221
222/// The profiles of every task kind: the shipped ones, and a deployment's overrides.
223///
224/// Keyed by the kind's name (`segment`, `route`, …), which is also how TOML names them.
225#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
226#[serde(transparent)]
227pub struct TaskProfiles {
228    overrides: BTreeMap<String, TaskProfile>,
229}
230
231impl TaskProfiles {
232    /// The shipped profiles, with no override.
233    #[must_use]
234    pub const fn new() -> Self {
235        Self {
236            overrides: BTreeMap::new(),
237        }
238    }
239
240    /// The profile `kind` runs under.
241    #[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    /// Replaces the profile of `kind`.
250    #[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    /// Adjusts the profile of `kind`, starting from what it is now.
257    #[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    /// Every overridden kind name, so a configuration can be checked for typos.
264    pub fn overridden(&self) -> impl Iterator<Item = &str> {
265        self.overrides.keys().map(String::as_str)
266    }
267
268    /// Every pool tag the profiles name, so the pool can be checked for them.
269    #[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/// The fields of a [`TaskProfile`] to change; a field left out keeps its value.
284#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
285#[serde(default, deny_unknown_fields)]
286#[non_exhaustive]
287pub struct ProfileChange {
288    /// [`TaskProfile::enabled`].
289    pub enabled: Option<bool>,
290    /// [`TaskProfile::model`].
291    pub model: Option<String>,
292    /// [`TaskProfile::escalate_to`].
293    pub escalate_to: Option<String>,
294    /// [`TaskProfile::max_output_tokens`].
295    pub max_output_tokens: Option<u32>,
296    /// [`TaskProfile::reasoning_effort`].
297    pub reasoning_effort: Option<ReasoningEffort>,
298    /// [`TaskProfile::timeout_secs`].
299    pub timeout_secs: Option<u64>,
300    /// [`TaskProfile::votes`].
301    pub votes: Option<u8>,
302    /// [`TaskProfile::on_disagreement`].
303    pub on_disagreement: Option<Disagreement>,
304    /// [`TaskProfile::repairs`].
305    pub repairs: Option<u8>,
306    /// [`TaskProfile::retries`].
307    pub retries: Option<u8>,
308    /// [`TaskProfile::review`].
309    pub review: Option<bool>,
310}
311
312impl ProfileChange {
313    /// `profile` with every field this change names.
314    #[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/// Changes to several task kinds, keyed by the kind's name as TOML writes it.
355#[derive(Debug, Clone, Default, PartialEq, Serialize)]
356#[serde(transparent)]
357pub struct ProfileChanges {
358    changes: BTreeMap<String, ProfileChange>,
359}
360
361impl ProfileChanges {
362    /// No change.
363    #[must_use]
364    pub const fn new() -> Self {
365        Self {
366            changes: BTreeMap::new(),
367        }
368    }
369
370    /// Changes `kind` as `change` says.
371    #[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    /// `profiles` with every change applied.
378    #[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    /// Every pool tag the changes name, so the pool can be checked for them.
389    #[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}