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::Route | TaskKind::Locate | TaskKind::QuestionFrame => Self {
106                max_output_tokens: Some(150),
107                ..base
108            },
109            TaskKind::Extract => Self {
110                max_output_tokens: Some(600),
111                ..base
112            },
113            // The one judge of a reading: at minimal reasoning it misjudged multi-part
114            // messages; reasoning counts against the output cap, so it gets room.
115            TaskKind::Verify => Self {
116                max_output_tokens: Some(2_000),
117                reasoning_effort: Some(ReasoningEffort::Low),
118                repairs: 0,
119                ..base
120            },
121            TaskKind::CrossCheck => Self {
122                max_output_tokens: Some(400),
123                ..base
124            },
125            TaskKind::Respects => Self {
126                max_output_tokens: Some(200),
127                ..base
128            },
129            TaskKind::Investigate => Self {
130                enabled: false,
131                max_output_tokens: Some(400),
132                ..base
133            },
134            TaskKind::Acknowledge => Self {
135                temperature: None,
136                reasoning_effort: None,
137                max_output_tokens: Some(300),
138                review: true,
139                ..base
140            },
141            TaskKind::Answer => Self {
142                temperature: None,
143                reasoning_effort: None,
144                ..base
145            },
146            TaskKind::Review => Self {
147                max_output_tokens: Some(300),
148                ..base
149            },
150            // A progress line is a preview: one call, and nothing waits for it.
151            TaskKind::Progress => Self {
152                temperature: None,
153                reasoning_effort: None,
154                max_output_tokens: Some(80),
155                repairs: 0,
156                retries: 0,
157                ..base
158            },
159            _ => base,
160        }
161    }
162
163    /// A copy casting `votes` answers.
164    #[must_use]
165    pub fn with_votes(mut self, votes: u8) -> Self {
166        self.votes = votes.max(1);
167        self
168    }
169
170    /// A copy answered by the model tagged `tag`.
171    #[must_use]
172    pub fn on_model(mut self, tag: impl Into<String>) -> Self {
173        self.model = Some(tag.into());
174        self
175    }
176
177    /// A copy escalating to the model tagged `tag`.
178    #[must_use]
179    pub fn escalating_to(mut self, tag: impl Into<String>) -> Self {
180        self.escalate_to = Some(tag.into());
181        self
182    }
183
184    /// A copy deciding a vote without a majority as `disagreement` says.
185    #[must_use]
186    pub fn on_disagreement(mut self, disagreement: Disagreement) -> Self {
187        self.on_disagreement = disagreement;
188        self
189    }
190
191    /// A copy with `repairs` repair rounds.
192    #[must_use]
193    pub fn with_repairs(mut self, repairs: u8) -> Self {
194        self.repairs = repairs;
195        self
196    }
197
198    /// A copy with the task switched on or off.
199    #[must_use]
200    pub fn enabled(mut self, enabled: bool) -> Self {
201        self.enabled = enabled;
202        self
203    }
204
205    /// A copy thinking as much as `reasoning` allows.
206    #[must_use]
207    pub fn with_reasoning(mut self, reasoning: Option<ReasoningEffort>) -> Self {
208        self.reasoning_effort = reasoning;
209        self
210    }
211
212    /// A copy whose written block is reviewed, or not.
213    #[must_use]
214    pub fn with_review(mut self, review: bool) -> Self {
215        self.review = review;
216        self
217    }
218}
219
220/// The profiles of every task kind: the shipped ones, and a deployment's overrides.
221///
222/// Keyed by the kind's name (`segment`, `route`, …), which is also how TOML names them.
223#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
224#[serde(transparent)]
225pub struct TaskProfiles {
226    overrides: BTreeMap<String, TaskProfile>,
227}
228
229impl TaskProfiles {
230    /// The shipped profiles, with no override.
231    #[must_use]
232    pub const fn new() -> Self {
233        Self {
234            overrides: BTreeMap::new(),
235        }
236    }
237
238    /// The profile `kind` runs under.
239    #[must_use]
240    pub fn get(&self, kind: TaskKind) -> TaskProfile {
241        self.overrides
242            .get(kind.as_str())
243            .cloned()
244            .unwrap_or_else(|| TaskProfile::default_for(kind))
245    }
246
247    /// Replaces the profile of `kind`.
248    #[must_use]
249    pub fn with(mut self, kind: TaskKind, profile: TaskProfile) -> Self {
250        self.overrides.insert(kind.as_str().to_owned(), profile);
251        self
252    }
253
254    /// Adjusts the profile of `kind`, starting from what it is now.
255    #[must_use]
256    pub fn adjust(self, kind: TaskKind, change: impl FnOnce(TaskProfile) -> TaskProfile) -> Self {
257        let current = self.get(kind);
258        self.with(kind, change(current))
259    }
260
261    /// Every overridden kind name, so a configuration can be checked for typos.
262    pub fn overridden(&self) -> impl Iterator<Item = &str> {
263        self.overrides.keys().map(String::as_str)
264    }
265
266    /// Every pool tag the profiles name, so the pool can be checked for them.
267    #[must_use]
268    pub fn tags(&self) -> Vec<String> {
269        let mut tags: Vec<String> = TaskKind::ALL
270            .iter()
271            .map(|kind| self.get(*kind))
272            .flat_map(|profile| [profile.model, profile.escalate_to])
273            .flatten()
274            .collect();
275        tags.sort();
276        tags.dedup();
277        tags
278    }
279}
280
281/// The fields of a [`TaskProfile`] to change; a field left out keeps its value.
282#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
283#[serde(default, deny_unknown_fields)]
284#[non_exhaustive]
285pub struct ProfileChange {
286    /// [`TaskProfile::enabled`].
287    pub enabled: Option<bool>,
288    /// [`TaskProfile::model`].
289    pub model: Option<String>,
290    /// [`TaskProfile::escalate_to`].
291    pub escalate_to: Option<String>,
292    /// [`TaskProfile::max_output_tokens`].
293    pub max_output_tokens: Option<u32>,
294    /// [`TaskProfile::reasoning_effort`].
295    pub reasoning_effort: Option<ReasoningEffort>,
296    /// [`TaskProfile::timeout_secs`].
297    pub timeout_secs: Option<u64>,
298    /// [`TaskProfile::votes`].
299    pub votes: Option<u8>,
300    /// [`TaskProfile::on_disagreement`].
301    pub on_disagreement: Option<Disagreement>,
302    /// [`TaskProfile::repairs`].
303    pub repairs: Option<u8>,
304    /// [`TaskProfile::retries`].
305    pub retries: Option<u8>,
306    /// [`TaskProfile::review`].
307    pub review: Option<bool>,
308}
309
310impl ProfileChange {
311    /// `profile` with every field this change names.
312    #[must_use]
313    pub fn apply(&self, mut profile: TaskProfile) -> TaskProfile {
314        let change = self.clone();
315        if let Some(value) = change.enabled {
316            profile.enabled = value;
317        }
318        if let Some(value) = change.model {
319            profile.model = Some(value);
320        }
321        if let Some(value) = change.escalate_to {
322            profile.escalate_to = Some(value);
323        }
324        if let Some(value) = change.max_output_tokens {
325            profile.max_output_tokens = Some(value);
326        }
327        if let Some(value) = change.reasoning_effort {
328            profile.reasoning_effort = Some(value);
329        }
330        if let Some(value) = change.timeout_secs {
331            profile.timeout_secs = Some(value);
332        }
333        if let Some(value) = change.votes {
334            profile.votes = value.max(1);
335        }
336        if let Some(value) = change.on_disagreement {
337            profile.on_disagreement = value;
338        }
339        if let Some(value) = change.repairs {
340            profile.repairs = value;
341        }
342        if let Some(value) = change.retries {
343            profile.retries = value;
344        }
345        if let Some(value) = change.review {
346            profile.review = value;
347        }
348        profile
349    }
350}
351
352/// Changes to several task kinds, keyed by the kind's name as TOML writes it.
353#[derive(Debug, Clone, Default, PartialEq, Serialize)]
354#[serde(transparent)]
355pub struct ProfileChanges {
356    changes: BTreeMap<String, ProfileChange>,
357}
358
359impl ProfileChanges {
360    /// No change.
361    #[must_use]
362    pub const fn new() -> Self {
363        Self {
364            changes: BTreeMap::new(),
365        }
366    }
367
368    /// Changes `kind` as `change` says.
369    #[must_use]
370    pub fn with(mut self, kind: TaskKind, change: ProfileChange) -> Self {
371        self.changes.insert(kind.as_str().to_owned(), change);
372        self
373    }
374
375    /// `profiles` with every change applied.
376    #[must_use]
377    pub fn apply(&self, mut profiles: TaskProfiles) -> TaskProfiles {
378        for kind in TaskKind::ALL {
379            if let Some(change) = self.changes.get(kind.as_str()) {
380                profiles = profiles.adjust(kind, |profile| change.apply(profile));
381            }
382        }
383        profiles
384    }
385
386    /// Every pool tag the changes name, so the pool can be checked for them.
387    #[must_use]
388    pub fn tags(&self) -> Vec<String> {
389        self.changes
390            .values()
391            .flat_map(|change| [change.model.clone(), change.escalate_to.clone()])
392            .flatten()
393            .collect()
394    }
395}
396
397impl<'de> Deserialize<'de> for ProfileChanges {
398    fn deserialize<D: serde::Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
399        let changes = BTreeMap::<String, ProfileChange>::deserialize(deserializer)?;
400        if let Some(unknown) = changes.keys().find(|name| {
401            !TaskKind::ALL
402                .iter()
403                .any(|kind| kind.as_str() == name.as_str())
404        }) {
405            return Err(serde::de::Error::custom(format!(
406                "`{unknown}` is not a task kind"
407            )));
408        }
409        Ok(Self { changes })
410    }
411}
412
413#[cfg(test)]
414mod tests {
415    use super::*;
416
417    #[test]
418    fn the_verifier_thinks_before_it_judges() {
419        let verify = TaskProfile::default_for(TaskKind::Verify);
420        assert_eq!(verify.reasoning_effort, Some(ReasoningEffort::Low));
421        assert!(
422            verify.max_output_tokens >= Some(2_000),
423            "reasoning counts against the cap"
424        );
425        assert_eq!(
426            TaskProfile::default_for(TaskKind::Extract).reasoning_effort,
427            Some(ReasoningEffort::Minimal)
428        );
429    }
430
431    #[test]
432    fn an_override_names_only_what_changes() {
433        let profiles: TaskProfiles = toml::from_str(
434            r#"
435            [route]
436            votes = 3
437            on_disagreement = "escalate"
438            escalate_to = "large"
439            "#,
440        )
441        .expect("parses");
442        let route = profiles.get(TaskKind::Route);
443        assert_eq!(route.votes, 3);
444        assert_eq!(route.on_disagreement, Disagreement::Escalate);
445        assert_eq!(route.escalate_to.as_deref(), Some("large"));
446        assert_eq!(profiles.get(TaskKind::Extract).max_output_tokens, Some(600));
447        assert_eq!(profiles.tags(), vec!["large".to_owned()]);
448    }
449
450    #[test]
451    fn a_change_keeps_what_it_does_not_name() {
452        let changes: ProfileChanges = toml::from_str(
453            r#"
454            [route]
455            votes = 3
456            "#,
457        )
458        .expect("parses");
459        let profiles = changes.apply(TaskProfiles::new());
460        let route = profiles.get(TaskKind::Route);
461        assert_eq!(route.votes, 3);
462        assert_eq!(
463            route.max_output_tokens,
464            Some(150),
465            "route's own cap survives"
466        );
467    }
468
469    #[test]
470    fn a_change_to_a_task_that_does_not_exist_is_refused_by_name() {
471        let refused = toml::from_str::<ProfileChanges>("[extrakt]\nvotes = 3\n").unwrap_err();
472        assert!(refused.to_string().contains("extrakt"), "{refused}");
473    }
474
475    #[test]
476    fn optional_tasks_ship_as_the_spec_says() {
477        let profiles = TaskProfiles::new();
478        assert!(profiles.get(TaskKind::Coverage).enabled);
479        assert!(!profiles.get(TaskKind::Investigate).enabled);
480        assert!(profiles.get(TaskKind::Acknowledge).review);
481        assert!(!profiles.get(TaskKind::Answer).review);
482        assert_eq!(profiles.get(TaskKind::Verify).repairs, 0);
483    }
484}