use std::collections::BTreeMap;
use serde::{Deserialize, Serialize};
use turnframe_provider::request::ReasoningEffort;
use crate::task::TaskKind;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum Disagreement {
Escalate,
#[default]
Clarify,
Fail,
Reread,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(default, deny_unknown_fields)]
#[non_exhaustive]
pub struct TaskProfile {
pub enabled: bool,
pub model: Option<String>,
pub escalate_to: Option<String>,
pub temperature: Option<f32>,
pub vote_temperature: f32,
pub max_output_tokens: Option<u32>,
pub reasoning_effort: Option<ReasoningEffort>,
pub timeout_secs: Option<u64>,
pub votes: u8,
pub on_disagreement: Disagreement,
pub repairs: u8,
pub retries: u8,
pub review: bool,
}
impl Default for TaskProfile {
fn default() -> Self {
Self {
enabled: true,
model: None,
escalate_to: None,
temperature: Some(0.0),
vote_temperature: 0.7,
max_output_tokens: None,
reasoning_effort: Some(ReasoningEffort::Minimal),
timeout_secs: None,
votes: 1,
on_disagreement: Disagreement::Clarify,
repairs: 1,
retries: 1,
review: false,
}
}
}
impl TaskProfile {
#[must_use]
pub fn default_for(kind: TaskKind) -> Self {
let base = Self::default();
match kind {
TaskKind::Segment => Self {
max_output_tokens: Some(800),
..base
},
TaskKind::Coverage => Self {
max_output_tokens: Some(200),
repairs: 0,
..base
},
TaskKind::Route | TaskKind::Locate | TaskKind::QuestionFrame => Self {
max_output_tokens: Some(150),
..base
},
TaskKind::Extract => Self {
max_output_tokens: Some(600),
..base
},
TaskKind::Verify => Self {
max_output_tokens: Some(2_000),
reasoning_effort: Some(ReasoningEffort::Low),
repairs: 0,
..base
},
TaskKind::CrossCheck => Self {
max_output_tokens: Some(400),
..base
},
TaskKind::Respects => Self {
max_output_tokens: Some(200),
..base
},
TaskKind::Investigate => Self {
enabled: false,
max_output_tokens: Some(400),
..base
},
TaskKind::Acknowledge => Self {
temperature: None,
reasoning_effort: None,
max_output_tokens: Some(300),
review: true,
..base
},
TaskKind::Answer => Self {
temperature: None,
reasoning_effort: None,
..base
},
TaskKind::Review => Self {
max_output_tokens: Some(300),
..base
},
TaskKind::Progress => Self {
temperature: None,
reasoning_effort: None,
max_output_tokens: Some(80),
repairs: 0,
retries: 0,
..base
},
_ => base,
}
}
#[must_use]
pub fn with_votes(mut self, votes: u8) -> Self {
self.votes = votes.max(1);
self
}
#[must_use]
pub fn on_model(mut self, tag: impl Into<String>) -> Self {
self.model = Some(tag.into());
self
}
#[must_use]
pub fn escalating_to(mut self, tag: impl Into<String>) -> Self {
self.escalate_to = Some(tag.into());
self
}
#[must_use]
pub fn on_disagreement(mut self, disagreement: Disagreement) -> Self {
self.on_disagreement = disagreement;
self
}
#[must_use]
pub fn with_repairs(mut self, repairs: u8) -> Self {
self.repairs = repairs;
self
}
#[must_use]
pub fn enabled(mut self, enabled: bool) -> Self {
self.enabled = enabled;
self
}
#[must_use]
pub fn with_reasoning(mut self, reasoning: Option<ReasoningEffort>) -> Self {
self.reasoning_effort = reasoning;
self
}
#[must_use]
pub fn with_review(mut self, review: bool) -> Self {
self.review = review;
self
}
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
#[serde(transparent)]
pub struct TaskProfiles {
overrides: BTreeMap<String, TaskProfile>,
}
impl TaskProfiles {
#[must_use]
pub const fn new() -> Self {
Self {
overrides: BTreeMap::new(),
}
}
#[must_use]
pub fn get(&self, kind: TaskKind) -> TaskProfile {
self.overrides
.get(kind.as_str())
.cloned()
.unwrap_or_else(|| TaskProfile::default_for(kind))
}
#[must_use]
pub fn with(mut self, kind: TaskKind, profile: TaskProfile) -> Self {
self.overrides.insert(kind.as_str().to_owned(), profile);
self
}
#[must_use]
pub fn adjust(self, kind: TaskKind, change: impl FnOnce(TaskProfile) -> TaskProfile) -> Self {
let current = self.get(kind);
self.with(kind, change(current))
}
pub fn overridden(&self) -> impl Iterator<Item = &str> {
self.overrides.keys().map(String::as_str)
}
#[must_use]
pub fn tags(&self) -> Vec<String> {
let mut tags: Vec<String> = TaskKind::ALL
.iter()
.map(|kind| self.get(*kind))
.flat_map(|profile| [profile.model, profile.escalate_to])
.flatten()
.collect();
tags.sort();
tags.dedup();
tags
}
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
#[serde(default, deny_unknown_fields)]
#[non_exhaustive]
pub struct ProfileChange {
pub enabled: Option<bool>,
pub model: Option<String>,
pub escalate_to: Option<String>,
pub max_output_tokens: Option<u32>,
pub reasoning_effort: Option<ReasoningEffort>,
pub timeout_secs: Option<u64>,
pub votes: Option<u8>,
pub on_disagreement: Option<Disagreement>,
pub repairs: Option<u8>,
pub retries: Option<u8>,
pub review: Option<bool>,
}
impl ProfileChange {
#[must_use]
pub fn apply(&self, mut profile: TaskProfile) -> TaskProfile {
let change = self.clone();
if let Some(value) = change.enabled {
profile.enabled = value;
}
if let Some(value) = change.model {
profile.model = Some(value);
}
if let Some(value) = change.escalate_to {
profile.escalate_to = Some(value);
}
if let Some(value) = change.max_output_tokens {
profile.max_output_tokens = Some(value);
}
if let Some(value) = change.reasoning_effort {
profile.reasoning_effort = Some(value);
}
if let Some(value) = change.timeout_secs {
profile.timeout_secs = Some(value);
}
if let Some(value) = change.votes {
profile.votes = value.max(1);
}
if let Some(value) = change.on_disagreement {
profile.on_disagreement = value;
}
if let Some(value) = change.repairs {
profile.repairs = value;
}
if let Some(value) = change.retries {
profile.retries = value;
}
if let Some(value) = change.review {
profile.review = value;
}
profile
}
}
#[derive(Debug, Clone, Default, PartialEq, Serialize)]
#[serde(transparent)]
pub struct ProfileChanges {
changes: BTreeMap<String, ProfileChange>,
}
impl ProfileChanges {
#[must_use]
pub const fn new() -> Self {
Self {
changes: BTreeMap::new(),
}
}
#[must_use]
pub fn with(mut self, kind: TaskKind, change: ProfileChange) -> Self {
self.changes.insert(kind.as_str().to_owned(), change);
self
}
#[must_use]
pub fn apply(&self, mut profiles: TaskProfiles) -> TaskProfiles {
for kind in TaskKind::ALL {
if let Some(change) = self.changes.get(kind.as_str()) {
profiles = profiles.adjust(kind, |profile| change.apply(profile));
}
}
profiles
}
#[must_use]
pub fn tags(&self) -> Vec<String> {
self.changes
.values()
.flat_map(|change| [change.model.clone(), change.escalate_to.clone()])
.flatten()
.collect()
}
}
impl<'de> Deserialize<'de> for ProfileChanges {
fn deserialize<D: serde::Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
let changes = BTreeMap::<String, ProfileChange>::deserialize(deserializer)?;
if let Some(unknown) = changes.keys().find(|name| {
!TaskKind::ALL
.iter()
.any(|kind| kind.as_str() == name.as_str())
}) {
return Err(serde::de::Error::custom(format!(
"`{unknown}` is not a task kind"
)));
}
Ok(Self { changes })
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn the_verifier_thinks_before_it_judges() {
let verify = TaskProfile::default_for(TaskKind::Verify);
assert_eq!(verify.reasoning_effort, Some(ReasoningEffort::Low));
assert!(
verify.max_output_tokens >= Some(2_000),
"reasoning counts against the cap"
);
assert_eq!(
TaskProfile::default_for(TaskKind::Extract).reasoning_effort,
Some(ReasoningEffort::Minimal)
);
}
#[test]
fn an_override_names_only_what_changes() {
let profiles: TaskProfiles = toml::from_str(
r#"
[route]
votes = 3
on_disagreement = "escalate"
escalate_to = "large"
"#,
)
.expect("parses");
let route = profiles.get(TaskKind::Route);
assert_eq!(route.votes, 3);
assert_eq!(route.on_disagreement, Disagreement::Escalate);
assert_eq!(route.escalate_to.as_deref(), Some("large"));
assert_eq!(profiles.get(TaskKind::Extract).max_output_tokens, Some(600));
assert_eq!(profiles.tags(), vec!["large".to_owned()]);
}
#[test]
fn a_change_keeps_what_it_does_not_name() {
let changes: ProfileChanges = toml::from_str(
r#"
[route]
votes = 3
"#,
)
.expect("parses");
let profiles = changes.apply(TaskProfiles::new());
let route = profiles.get(TaskKind::Route);
assert_eq!(route.votes, 3);
assert_eq!(
route.max_output_tokens,
Some(150),
"route's own cap survives"
);
}
#[test]
fn a_change_to_a_task_that_does_not_exist_is_refused_by_name() {
let refused = toml::from_str::<ProfileChanges>("[extrakt]\nvotes = 3\n").unwrap_err();
assert!(refused.to_string().contains("extrakt"), "{refused}");
}
#[test]
fn optional_tasks_ship_as_the_spec_says() {
let profiles = TaskProfiles::new();
assert!(profiles.get(TaskKind::Coverage).enabled);
assert!(!profiles.get(TaskKind::Investigate).enabled);
assert!(profiles.get(TaskKind::Acknowledge).review);
assert!(!profiles.get(TaskKind::Answer).review);
assert_eq!(profiles.get(TaskKind::Verify).repairs, 0);
}
}