use std::collections::{BTreeMap, VecDeque};
use std::fmt;
use std::sync::{Arc, Mutex, PoisonError};
use async_trait::async_trait;
use turnframe_provider::capabilities::{
ModelProfile, ProviderCapabilities, StructuredOutputCapability,
};
use turnframe_provider::error::ProviderError;
use turnframe_provider::ids::{ModelKey, ProviderKey};
use turnframe_provider::provider::ModelProvider;
use turnframe_provider::request::ModelRequest;
use turnframe_provider::response::{ModelResponse, TokenUsage};
use turnframe_provider::router::{PolicyRouter, ProviderPool};
use crate::engine::TASK_LABEL;
enum Scripted {
Answer(serde_json::Value),
Failure(ProviderError),
}
pub struct ScriptedTasks {
profile: ModelProfile,
answers: Mutex<BTreeMap<String, VecDeque<Scripted>>>,
calls: Mutex<Vec<ModelRequest>>,
}
impl ScriptedTasks {
#[must_use]
pub fn new(provider: impl Into<ProviderKey>, model: impl Into<ModelKey>) -> Self {
let capabilities = ProviderCapabilities::minimal()
.with_structured_output(StructuredOutputCapability::NativeJsonSchema)
.with_temperature(true);
Self {
profile: ModelProfile::new(provider, model, capabilities),
answers: Mutex::new(BTreeMap::new()),
calls: Mutex::new(Vec::new()),
}
}
#[must_use]
pub fn answer(self, task: &str, answer: serde_json::Value) -> Self {
self.lock_answers()
.entry(task.to_owned())
.or_default()
.push_back(Scripted::Answer(answer));
self
}
#[must_use]
pub fn failing(self, task: &str, error: ProviderError) -> Self {
self.lock_answers()
.entry(task.to_owned())
.or_default()
.push_back(Scripted::Failure(error));
self
}
#[must_use]
pub fn tagged(mut self, tag: impl Into<String>) -> Self {
self.profile.tags.push(tag.into());
self
}
#[must_use]
pub fn calls(&self) -> Vec<ModelRequest> {
self.calls
.lock()
.unwrap_or_else(PoisonError::into_inner)
.clone()
}
#[must_use]
pub fn called(&self) -> Vec<String> {
self.calls()
.iter()
.map(|request| {
request
.metadata
.get(TASK_LABEL)
.unwrap_or_default()
.to_owned()
})
.collect()
}
#[must_use]
pub fn unanswered(&self) -> Vec<String> {
self.lock_answers()
.iter()
.filter(|(_, queue)| !queue.is_empty())
.map(|(task, _)| task.clone())
.collect()
}
#[must_use]
pub fn router(self: &Arc<Self>) -> Arc<PolicyRouter> {
let provider: Arc<dyn ModelProvider> = self.clone();
let pool = ProviderPool::builder()
.provider(provider)
.build()
.unwrap_or_else(|error| unreachable!("a pool of one provider builds: {error}"));
Arc::new(PolicyRouter::new(Arc::new(pool)))
}
fn lock_answers(&self) -> std::sync::MutexGuard<'_, BTreeMap<String, VecDeque<Scripted>>> {
self.answers.lock().unwrap_or_else(PoisonError::into_inner)
}
fn take(&self, call: &str) -> Option<Scripted> {
let mut answers = self.lock_answers();
let task = call.split('#').next().unwrap_or(call);
for key in [call, task] {
if let Some(answer) = answers.get_mut(key).and_then(VecDeque::pop_front) {
return Some(answer);
}
}
None
}
}
#[async_trait]
impl ModelProvider for ScriptedTasks {
fn provider_key(&self) -> ProviderKey {
self.profile.provider.clone()
}
fn model_key(&self) -> ModelKey {
self.profile.model.clone()
}
fn capabilities(&self) -> ProviderCapabilities {
self.profile.capabilities.clone()
}
fn profile(&self) -> ModelProfile {
self.profile.clone()
}
async fn generate(&self, request: ModelRequest) -> Result<ModelResponse, ProviderError> {
self.calls
.lock()
.unwrap_or_else(PoisonError::into_inner)
.push(request.clone());
let Some(call) = request.metadata.get(TASK_LABEL) else {
return Err(ProviderError::unsupported("untasked_request"));
};
let answer = match self.take(call) {
Some(Scripted::Answer(answer)) => answer,
Some(Scripted::Failure(error)) => return Err(error),
None => {
return Err(ProviderError::invalid_request("unscripted_task").with_detail(call));
}
};
let text = answer.to_string();
let prompt: usize = request
.messages
.iter()
.map(|message| message.text().len())
.sum();
let prompt = prompt + request.system.as_ref().map_or(0, String::len);
let tokens = |characters: usize| u64::try_from(characters / 4).unwrap_or(u64::MAX);
let usage = TokenUsage::new(tokens(prompt), tokens(text.len()));
Ok(ModelResponse::new(
request.request_id,
self.profile.provider.clone(),
self.profile.model.clone(),
)
.with_text(text)
.with_usage(usage))
}
}
impl fmt::Debug for ScriptedTasks {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ScriptedTasks")
.field("model", &self.profile.reference().to_string())
.field("unanswered", &self.unanswered())
.finish_non_exhaustive()
}
}