use std::sync::{Arc, Mutex, PoisonError};
use std::time::{Duration, Instant};
use futures::future::join_all;
use turnframe_core::effort::Effort;
use turnframe_core::hash::Digest;
use turnframe_core::locale::Locale;
use turnframe_core::observe::{NoopObserver, Observer, Signal, SignalLabels};
use turnframe_core::prompt::{PromptSelector, PromptSource};
use turnframe_core::replay::{BudgetReport, TaskParams, TaskRecord, TaskVerdict};
use turnframe_provider::fallback::{FallbackOptions, FallbackStage, execute_with_fallback};
use turnframe_provider::ids::ModelRef;
use turnframe_provider::request::{CacheHint, Message, ModelRequest, OutputSpec};
use turnframe_provider::response::ModelResponse;
use turnframe_provider::router::{ProviderRouter, RoutingPolicy};
use turnframe_provider::structured::{CompiledSchema, SchemaCache, parse_structured};
use crate::budget::{Budget, BudgetBound, BudgetTracker};
use crate::instructions::{self, Instructions};
use crate::profile::{Disagreement, TaskProfile, TaskProfiles};
use crate::task::{ModelTask, TaskId, TaskKind};
pub const TASK_LABEL: &str = "task";
pub const TURN_LABEL: &str = "turn";
const BUILT_IN_REPAIR: &str = "Your previous answer was not accepted. Answer again with a \
document that satisfies the schema and fixes this:";
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
#[non_exhaustive]
pub struct RecordPolicy {
pub keep_prompts: bool,
pub keep_raw_output: bool,
}
#[derive(Debug)]
pub struct TaskScope {
budget: BudgetTracker,
records: Mutex<Vec<TaskRecord>>,
locale: Locale,
turn: Option<String>,
profiles: Option<TaskProfiles>,
effort: Option<Effort>,
}
impl TaskScope {
#[must_use]
pub fn new(budget: Budget, locale: Locale) -> Self {
Self {
budget: BudgetTracker::new(budget),
records: Mutex::new(Vec::new()),
locale,
turn: None,
profiles: None,
effort: None,
}
}
#[must_use]
pub fn with_profiles(mut self, profiles: TaskProfiles) -> Self {
self.profiles = Some(profiles);
self
}
#[must_use]
pub const fn with_effort(mut self, effort: Effort) -> Self {
self.effort = Some(effort);
self
}
#[must_use]
pub const fn effort(&self) -> Option<Effort> {
self.effort
}
fn labels(&self) -> SignalLabels {
let labels = SignalLabels::default();
match self.effort {
Some(effort) => labels.with_effort(effort),
None => labels,
}
}
#[must_use]
pub fn for_turn(mut self, turn: impl Into<String>) -> Self {
self.turn = Some(turn.into());
self
}
#[must_use]
pub fn records(&self) -> Vec<TaskRecord> {
self.lock().clone()
}
#[must_use]
pub fn budget_report(&self) -> BudgetReport {
self.budget.report()
}
#[must_use]
pub fn exhausted(&self) -> Option<BudgetBound> {
self.budget.exhausted()
}
#[must_use]
pub const fn locale(&self) -> &Locale {
&self.locale
}
fn push(&self, record: TaskRecord) {
self.lock().push(record);
}
fn mark(&self, task_id: &str, verdict: &TaskVerdict) {
if let Some(record) = self.lock().iter_mut().find(|r| r.task_id == task_id) {
record.verdict = verdict.clone();
}
}
fn lock(&self) -> std::sync::MutexGuard<'_, Vec<TaskRecord>> {
self.records.lock().unwrap_or_else(PoisonError::into_inner)
}
}
#[derive(Debug, Clone, Copy)]
pub struct TaskCall<'a> {
pub id: &'a TaskId,
pub parent: Option<&'a TaskId>,
pub depth: u8,
}
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub enum TaskFailure {
Disabled,
Budget(BudgetBound),
Routing(String),
Provider(String),
Invalid {
code: String,
reason: String,
},
Disagreement,
}
impl TaskFailure {
#[must_use]
pub fn code(&self) -> String {
match self {
Self::Disabled => "disabled".to_owned(),
Self::Budget(bound) => format!("budget_{}", bound.as_str()),
Self::Routing(_) => "routing".to_owned(),
Self::Provider(code) => format!("provider_{code}"),
Self::Invalid { code, .. } => code.clone(),
Self::Disagreement => "vote_disagreement".to_owned(),
}
}
}
#[derive(Debug, Clone)]
pub enum TaskOutcome<O> {
Accepted {
output: O,
depth: u8,
},
Disagreed {
answers: Vec<O>,
depth: u8,
},
Failed {
failure: TaskFailure,
depth: u8,
},
}
impl<O> TaskOutcome<O> {
#[must_use]
pub fn accepted(self) -> Option<O> {
match self {
Self::Accepted { output, .. } => Some(output),
_ => None,
}
}
#[must_use]
pub const fn depth(&self) -> u8 {
match self {
Self::Accepted { depth, .. }
| Self::Disagreed { depth, .. }
| Self::Failed { depth, .. } => *depth,
}
}
fn label(&self) -> String {
match self {
Self::Accepted { .. } => "accepted".to_owned(),
Self::Disagreed { .. } => "disagreed".to_owned(),
Self::Failed { failure, .. } => failure.code(),
}
}
}
#[derive(Clone)]
pub struct TaskEngine {
router: Arc<dyn ProviderRouter>,
routing: RoutingPolicy,
fallback: Arc<FallbackOptions>,
schemas: SchemaCache,
profiles: TaskProfiles,
prompts: Option<Arc<dyn PromptSource>>,
selector: PromptSelector,
records: RecordPolicy,
observer: Arc<dyn Observer>,
}
impl std::fmt::Debug for TaskEngine {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("TaskEngine")
.field("profiles", &self.profiles)
.field("records", &self.records)
.finish_non_exhaustive()
}
}
#[derive(Debug)]
pub struct TaskEngineBuilder {
engine: TaskEngine,
}
impl TaskEngineBuilder {
#[must_use]
pub fn routing(mut self, routing: RoutingPolicy) -> Self {
self.engine.routing = routing;
self
}
#[must_use]
pub fn fallback(mut self, options: FallbackOptions) -> Self {
self.engine.fallback = Arc::new(options);
self
}
#[must_use]
pub fn profiles(mut self, profiles: TaskProfiles) -> Self {
self.engine.profiles = profiles;
self
}
#[must_use]
pub fn prompts(mut self, source: Arc<dyn PromptSource>, selector: PromptSelector) -> Self {
self.engine.prompts = Some(source);
self.engine.selector = selector;
self
}
#[must_use]
pub fn records(mut self, policy: RecordPolicy) -> Self {
self.engine.records = policy;
self
}
#[must_use]
pub fn observer(mut self, observer: Arc<dyn Observer>) -> Self {
self.engine.observer = observer;
self
}
#[must_use]
pub fn build(self) -> TaskEngine {
self.engine
}
}
#[derive(Clone)]
struct Prepared {
kind: TaskKind,
parent: Option<String>,
profile: TaskProfile,
instructions: Instructions,
repair: Instructions,
schema: CompiledSchema,
messages: Vec<Message>,
}
enum Chain<O> {
Answered {
output: O,
depth: u8,
record: String,
},
Unusable {
failure: TaskFailure,
depth: u8,
},
}
impl TaskEngine {
#[must_use]
pub fn builder(router: Arc<dyn ProviderRouter>) -> TaskEngineBuilder {
TaskEngineBuilder {
engine: Self {
router,
routing: RoutingPolicy::new(),
fallback: Arc::new(FallbackOptions::new()),
schemas: SchemaCache::new(),
profiles: TaskProfiles::new(),
prompts: None,
selector: PromptSelector::Latest,
records: RecordPolicy::default(),
observer: Arc::new(NoopObserver),
},
}
}
#[must_use]
pub const fn profiles(&self) -> &TaskProfiles {
&self.profiles
}
#[must_use]
pub fn profile(&self, scope: &TaskScope, kind: TaskKind) -> TaskProfile {
scope.profiles.as_ref().unwrap_or(&self.profiles).get(kind)
}
pub async fn run<T: ModelTask>(
&self,
scope: &TaskScope,
call: TaskCall<'_>,
task: &T,
input: &T::Input,
) -> TaskOutcome<T::Output> {
self.run_inner(scope, call, task, input, None).await
}
pub async fn run_with_feedback<T: ModelTask>(
&self,
scope: &TaskScope,
call: TaskCall<'_>,
task: &T,
input: &T::Input,
previous: &T::Output,
feedback: &str,
) -> TaskOutcome<T::Output> {
self.run_inner(scope, call, task, input, Some((previous, feedback)))
.await
}
async fn run_inner<T: ModelTask>(
&self,
scope: &TaskScope,
call: TaskCall<'_>,
task: &T,
input: &T::Input,
feedback: Option<(&T::Output, &str)>,
) -> TaskOutcome<T::Output> {
let outcome = match self.prepare(scope, call, task, input, feedback).await {
Ok(prepared) => self.run_prepared(scope, call, task, input, &prepared).await,
Err(failure) => TaskOutcome::Failed {
failure,
depth: call.depth,
},
};
let labels = scope
.labels()
.with_purpose(task.kind().as_str())
.with_error_code(outcome.label());
self.observer
.observe_labeled(&Signal::TaskCompleted, &labels);
outcome
}
async fn prepare<T: ModelTask>(
&self,
scope: &TaskScope,
call: TaskCall<'_>,
task: &T,
input: &T::Input,
feedback: Option<(&T::Output, &str)>,
) -> Result<Prepared, TaskFailure> {
let kind = task.kind();
let profile = self.profile(scope, kind);
if !profile.enabled {
return Err(TaskFailure::Disabled);
}
let schema =
self.schemas
.compile(&task.schema(input))
.map_err(|error| TaskFailure::Invalid {
code: "schema_compile".to_owned(),
reason: error.to_string(),
})?;
let name = task.prompt_name();
let instructions = instructions::resolve(
self.prompts.as_ref(),
&self.selector,
name,
scope.locale(),
task.instructions(),
)
.await;
let repair = instructions::resolve(
self.prompts.as_ref(),
&self.selector,
&format!("{name}.repair"),
scope.locale(),
BUILT_IN_REPAIR,
)
.await;
let mut messages = task.render(input);
if let Some((previous, note)) = feedback {
messages.push(Message::assistant(
serde_json::to_string(previous).unwrap_or_default(),
));
messages.push(Message::user(format!("{}\n\n{note}", repair.text)));
}
Ok(Prepared {
kind,
parent: call.parent.map(|parent| parent.as_str().to_owned()),
profile,
instructions,
repair,
schema,
messages,
})
}
async fn run_prepared<T: ModelTask>(
&self,
scope: &TaskScope,
call: TaskCall<'_>,
task: &T,
input: &T::Input,
prepared: &Prepared,
) -> TaskOutcome<T::Output> {
let tag = prepared.profile.model.as_deref();
let votes = prepared.profile.votes.max(1);
if votes == 1 {
let id = call.id.as_str().to_owned();
return match self
.chain(scope, prepared, task, input, tag, call.depth, id, None)
.await
{
Chain::Answered { output, depth, .. } => TaskOutcome::Accepted { output, depth },
Chain::Unusable { failure, depth } => {
self.escalate(scope, call, task, input, prepared, failure, depth)
.await
}
};
}
let temperature = Some(prepared.profile.vote_temperature);
let chains = join_all((1..=votes).map(|vote| {
let id = call.id.call(format!("vote{vote}"));
self.chain(
scope,
prepared,
task,
input,
tag,
call.depth,
id,
temperature,
)
}))
.await;
let depth = chains
.iter()
.map(|chain| match chain {
Chain::Answered { depth, .. } | Chain::Unusable { depth, .. } => *depth,
})
.max()
.unwrap_or(call.depth);
let answered: Vec<(&T::Output, &str)> = chains
.iter()
.filter_map(|chain| match chain {
Chain::Answered { output, record, .. } => Some((output, record.as_str())),
Chain::Unusable { .. } => None,
})
.collect();
if let Some(winner) = majority(task, &answered, usize::from(votes)) {
for (index, (_, record)) in answered.iter().enumerate() {
if !winner.contains(&index) {
scope.mark(record, &TaskVerdict::Outvoted);
}
}
let output = answered[winner[0]].0.clone();
return TaskOutcome::Accepted { output, depth };
}
self.observer.observe_labeled(
&Signal::TaskVoteDisagreement,
&scope.labels().with_purpose(prepared.kind.as_str()),
);
match prepared.profile.on_disagreement {
Disagreement::Reread if !answered.is_empty() => {
let shown: Vec<String> = answered
.iter()
.map(|(output, _)| serde_json::to_string(output).unwrap_or_default())
.collect();
let mut again = prepared.clone();
again.messages.push(Message::user(format!(
"Readings of this that disagreed:\n{}\n\nRead it again and give the answer \
the message supports.",
shown.join("\n")
)));
let id = call.id.call("reread");
match self
.chain(scope, &again, task, input, tag, depth, id, None)
.await
{
Chain::Answered { output, depth, .. } => {
TaskOutcome::Accepted { output, depth }
}
Chain::Unusable { failure, depth } => TaskOutcome::Failed { failure, depth },
}
}
Disagreement::Escalate => {
self.escalate(
scope,
call,
task,
input,
prepared,
TaskFailure::Disagreement,
depth,
)
.await
}
Disagreement::Clarify => TaskOutcome::Disagreed {
answers: answered
.into_iter()
.map(|(output, _)| output.clone())
.collect(),
depth,
},
_ => TaskOutcome::Failed {
failure: TaskFailure::Disagreement,
depth,
},
}
}
#[allow(clippy::too_many_arguments)]
async fn escalate<T: ModelTask>(
&self,
scope: &TaskScope,
call: TaskCall<'_>,
task: &T,
input: &T::Input,
prepared: &Prepared,
failure: TaskFailure,
depth: u8,
) -> TaskOutcome<T::Output> {
let escalates = matches!(
failure,
TaskFailure::Invalid { .. }
| TaskFailure::Disagreement
| TaskFailure::Provider(_)
| TaskFailure::Routing(_)
);
let Some(tag) = prepared
.profile
.escalate_to
.as_deref()
.filter(|_| escalates)
else {
return TaskOutcome::Failed { failure, depth };
};
self.observer.observe_labeled(
&Signal::TaskEscalated,
&scope
.labels()
.with_purpose(prepared.kind.as_str())
.with_error_code(failure.code()),
);
let id = call.id.call("escalation");
match self
.chain(
scope,
prepared,
task,
input,
Some(tag),
depth.saturating_add(1),
id,
None,
)
.await
{
Chain::Answered { output, depth, .. } => TaskOutcome::Accepted { output, depth },
Chain::Unusable { failure, depth } => TaskOutcome::Failed { failure, depth },
}
}
#[allow(clippy::too_many_arguments)]
async fn chain<T: ModelTask>(
&self,
scope: &TaskScope,
prepared: &Prepared,
task: &T,
input: &T::Input,
tag: Option<&str>,
depth: u8,
id: String,
temperature: Option<f32>,
) -> Chain<T::Output> {
let mut messages = prepared.messages.clone();
let mut depth = depth;
let mut failure = TaskFailure::Invalid {
code: "no_answer".to_owned(),
reason: String::new(),
};
let rounds = prepared.profile.repairs.saturating_add(1);
for round in 0..rounds {
let record_id = if round == 0 {
id.clone()
} else {
format!("{id}#repair{round}")
};
if let Err(bound) = scope.budget.reserve(depth) {
self.observer.observe_labeled(
&Signal::BudgetExhausted,
&scope.labels().with_error_code(bound.as_str()),
);
let mut record = self.record(prepared, &record_id, depth, None);
record.verdict = TaskVerdict::Failed {
code: format!("budget_{}", bound.as_str()),
};
scope.push(record);
return Chain::Unusable {
failure: TaskFailure::Budget(bound),
depth,
};
}
let mut retries = prepared.profile.retries;
let mut call_id = record_id.clone();
let (request, answer) = loop {
let request = self.request(scope, prepared, &messages, temperature, &call_id);
let answer = self.send(scope, prepared, tag, &request).await;
match &answer {
Err(TaskFailure::Provider(kind)) if retries > 0 && retried_in_place(kind) => {
let mut record = self.record(prepared, &call_id, depth, Some(&request));
record.verdict = TaskVerdict::Failed {
code: format!("provider_{kind}"),
};
scope.push(record);
if let Err(bound) = scope.budget.reserve(depth) {
return Chain::Unusable {
failure: TaskFailure::Budget(bound),
depth,
};
}
retries -= 1;
call_id =
format!("{record_id}#retry{}", prepared.profile.retries - retries);
}
_ => break (request, answer),
}
};
let mut record = self.record(prepared, &call_id, depth, Some(&request));
let response = match answer {
Ok((response, served_by, latency)) => {
record.provider_key = Some(served_by.provider.clone());
record.model_key = Some(served_by.model.clone());
record.input_tokens = Some(response.usage.input);
record.output_tokens = Some(response.usage.output);
record.latency_ms = u64::try_from(latency.as_millis()).ok();
response
}
Err(failed) => {
record.verdict = TaskVerdict::Failed {
code: failed.code(),
};
scope.push(record);
return Chain::Unusable {
failure: failed,
depth,
};
}
};
let raw = response.text();
if self.records.keep_raw_output {
record.raw_output = Some(raw.clone());
}
match judge(task, input, &prepared.schema, &response) {
Ok(output) => {
record.parsed = serde_json::to_value(&output).ok();
record.verdict = TaskVerdict::Accepted;
scope.push(record);
return Chain::Answered {
output,
depth,
record: record_id,
};
}
Err((code, reason)) => {
record.verdict = TaskVerdict::Rejected {
code: code.clone(),
reason: reason.clone(),
};
scope.push(record);
if round + 1 < rounds {
self.observer.observe_labeled(
&Signal::TaskRepaired,
&scope
.labels()
.with_purpose(prepared.kind.as_str())
.with_error_code(code.clone()),
);
messages.push(Message::assistant(raw));
messages.push(Message::user(format!(
"{}\n\n{reason}",
prepared.repair.text
)));
depth = depth.saturating_add(1);
}
failure = TaskFailure::Invalid { code, reason };
}
}
}
Chain::Unusable { failure, depth }
}
fn request(
&self,
scope: &TaskScope,
prepared: &Prepared,
messages: &[Message],
temperature: Option<f32>,
task_id: &str,
) -> ModelRequest {
let profile = &prepared.profile;
let timeout = scope
.budget
.call_timeout(profile.timeout_secs.map(Duration::from_secs));
let mut request = ModelRequest::new(prepared.kind)
.with_system(prepared.instructions.text.clone())
.with_output(OutputSpec::json(
prepared.kind.as_str(),
prepared.schema.schema().clone(),
))
.with_timeout(timeout)
.with_cache_hint(CacheHint::System);
for message in messages {
request = request.with_message(message.clone());
}
if let Some(temperature) = temperature.or(profile.temperature) {
request = request.with_temperature(temperature);
}
if let Some(tokens) = profile.max_output_tokens {
request = request.with_max_output_tokens(tokens);
}
if let Some(effort) = profile.reasoning_effort {
request = request.with_reasoning_effort(effort);
}
let _ = request.metadata.insert(TASK_LABEL, task_id);
if let Some(turn) = &scope.turn {
let _ = request.metadata.insert(TURN_LABEL, turn.clone());
}
request
}
async fn send(
&self,
scope: &TaskScope,
prepared: &Prepared,
tag: Option<&str>,
request: &ModelRequest,
) -> Result<(ModelResponse, ModelRef, Duration), TaskFailure> {
let routing = match tag {
Some(tag) => self.routing.clone().with_required_tag(tag),
None => self.routing.clone(),
};
let candidates = self
.router
.select(prepared.kind, &request.requirements(), &routing)
.map_err(|error| {
crate::signals::observe_routing_error(self.observer.as_ref(), &error);
TaskFailure::Routing(error.to_string())
})?;
let stage = if prepared.kind.is_critical() {
FallbackStage::PreCommit
} else {
FallbackStage::PostCommitNarration
};
let _permit = scope.budget.permit().await;
let started = Instant::now();
let outcome = execute_with_fallback(&candidates, request, stage, &self.fallback)
.await
.map_err(|failure| {
crate::signals::observe_attempts(self.observer.as_ref(), &failure.attempts);
tracing::warn!(
target: "turnframe.tasks",
task = request.metadata.get(TASK_LABEL).unwrap_or_default(),
kind = failure.error.kind().as_str(),
detail = failure.error.detail().map_or("", |detail| detail.as_str()),
"a task call failed at the provider"
);
TaskFailure::Provider(failure.error.kind().as_str().to_owned())
})?;
let latency = started.elapsed();
crate::signals::observe_attempts(self.observer.as_ref(), &outcome.attempts);
scope.budget.record_tokens(outcome.response.usage.input);
let served_by = outcome.served_by();
self.observer.observe_duration(
&Signal::TaskLatency,
latency,
&scope
.labels()
.with_purpose(prepared.kind.as_str())
.with_provider(served_by.provider.clone())
.with_model(served_by.model.clone()),
);
Ok((outcome.response, served_by, latency))
}
fn record(
&self,
prepared: &Prepared,
task_id: &str,
depth: u8,
request: Option<&ModelRequest>,
) -> TaskRecord {
let mut record = TaskRecord::new(task_id, prepared.kind.as_str(), TaskVerdict::Accepted);
record.parent.clone_from(&prepared.parent);
record.depth = depth;
record.prompt_ref = Some(prepared.instructions.reference.clone());
if let Some(request) = request {
record.params = params_of(request);
record.input_digest = Some(digest_of(request));
if self.records.keep_prompts {
record.rendered = serde_json::to_value(request).ok();
}
}
record
}
}
fn judge<T: ModelTask>(
task: &T,
input: &T::Input,
schema: &CompiledSchema,
response: &ModelResponse,
) -> Result<T::Output, (String, String)> {
let output: T::Output = parse_structured(response, schema)
.map_err(|error| ("schema".to_owned(), error.to_string()))?;
task.check(input, &output)
.map_err(|error| (error.code.to_owned(), error.message))?;
Ok(output)
}
fn majority<T: ModelTask>(
task: &T,
answered: &[(&T::Output, &str)],
votes: usize,
) -> Option<Vec<usize>> {
let mut groups: Vec<Vec<usize>> = Vec::new();
for (index, (output, _)) in answered.iter().enumerate() {
match groups
.iter_mut()
.find(|group| task.agree(answered[group[0]].0, output))
{
Some(group) => group.push(index),
None => groups.push(vec![index]),
}
}
let largest = groups
.into_iter()
.max_by_key(|group| (group.len(), usize::MAX - group[0]))?;
(largest.len() * 2 > votes).then_some(largest)
}
fn params_of(request: &ModelRequest) -> TaskParams {
let mut params = TaskParams::default();
params.temperature = request.temperature;
params.max_output_tokens = request.max_output_tokens;
params.reasoning_effort = request
.reasoning_effort
.map(|effort| effort.as_str().to_owned());
params.seed = request.seed;
params.timeout_ms = u64::try_from(request.timeout.as_millis()).unwrap_or(u64::MAX);
params
}
fn digest_of(request: &ModelRequest) -> Digest {
let shown = serde_json::json!({
"system": request.system,
"messages": request.messages,
"output": request.output,
});
Digest::of_bytes(&serde_json::to_vec(&shown).unwrap_or_default())
}
fn retried_in_place(kind: &str) -> bool {
matches!(
kind,
"refusal" | "content_filter" | "malformed" | "transport" | "server" | "timeout" | "other"
)
}