Skip to main content

turnframe_tasks/
testing.rs

1//! A provider that answers each task by its id, for testing anything built on tasks.
2//!
3//! Tasks run concurrently, so a queue answered in call order is a race. [`ScriptedTasks`]
4//! reads the task id every request carries under [`TASK_LABEL`] and answers from that
5//! task's own queue. An exact call id (`u1/extract#repair1`) is looked up before its task
6//! (`u1/extract`), so a repair or a vote can be scripted apart from the first call.
7
8use std::collections::{BTreeMap, VecDeque};
9use std::fmt;
10use std::sync::{Arc, Mutex, PoisonError};
11
12use async_trait::async_trait;
13use turnframe_provider::capabilities::{
14    ModelProfile, ProviderCapabilities, StructuredOutputCapability,
15};
16use turnframe_provider::error::ProviderError;
17use turnframe_provider::ids::{ModelKey, ProviderKey};
18use turnframe_provider::provider::ModelProvider;
19use turnframe_provider::request::ModelRequest;
20use turnframe_provider::response::{ModelResponse, TokenUsage};
21use turnframe_provider::router::{PolicyRouter, ProviderPool};
22
23use crate::engine::TASK_LABEL;
24
25/// One queued reply: an answer, or a provider failure.
26enum Scripted {
27    Answer(serde_json::Value),
28    Failure(ProviderError),
29}
30
31/// Answers each task from its own queue of JSON documents.
32pub struct ScriptedTasks {
33    profile: ModelProfile,
34    answers: Mutex<BTreeMap<String, VecDeque<Scripted>>>,
35    calls: Mutex<Vec<ModelRequest>>,
36}
37
38impl ScriptedTasks {
39    /// A provider with native schema support and no answers yet.
40    #[must_use]
41    pub fn new(provider: impl Into<ProviderKey>, model: impl Into<ModelKey>) -> Self {
42        let capabilities = ProviderCapabilities::minimal()
43            .with_structured_output(StructuredOutputCapability::NativeJsonSchema)
44            .with_temperature(true);
45        Self {
46            profile: ModelProfile::new(provider, model, capabilities),
47            answers: Mutex::new(BTreeMap::new()),
48            calls: Mutex::new(Vec::new()),
49        }
50    }
51
52    /// Queues `answer` for the task or call `task`.
53    #[must_use]
54    pub fn answer(self, task: &str, answer: serde_json::Value) -> Self {
55        self.lock_answers()
56            .entry(task.to_owned())
57            .or_default()
58            .push_back(Scripted::Answer(answer));
59        self
60    }
61
62    /// Queues a provider failure for the task or call `task`.
63    #[must_use]
64    pub fn failing(self, task: &str, error: ProviderError) -> Self {
65        self.lock_answers()
66            .entry(task.to_owned())
67            .or_default()
68            .push_back(Scripted::Failure(error));
69        self
70    }
71
72    /// Adds a tag to the profile, for routing by tag.
73    #[must_use]
74    pub fn tagged(mut self, tag: impl Into<String>) -> Self {
75        self.profile.tags.push(tag.into());
76        self
77    }
78
79    /// Every request received, in order.
80    #[must_use]
81    pub fn calls(&self) -> Vec<ModelRequest> {
82        self.calls
83            .lock()
84            .unwrap_or_else(PoisonError::into_inner)
85            .clone()
86    }
87
88    /// The task ids called, in order.
89    #[must_use]
90    pub fn called(&self) -> Vec<String> {
91        self.calls()
92            .iter()
93            .map(|request| {
94                request
95                    .metadata
96                    .get(TASK_LABEL)
97                    .unwrap_or_default()
98                    .to_owned()
99            })
100            .collect()
101    }
102
103    /// Tasks with answers never asked for.
104    #[must_use]
105    pub fn unanswered(&self) -> Vec<String> {
106        self.lock_answers()
107            .iter()
108            .filter(|(_, queue)| !queue.is_empty())
109            .map(|(task, _)| task.clone())
110            .collect()
111    }
112
113    /// A router over this provider alone.
114    #[must_use]
115    pub fn router(self: &Arc<Self>) -> Arc<PolicyRouter> {
116        let provider: Arc<dyn ModelProvider> = self.clone();
117        let pool = ProviderPool::builder()
118            .provider(provider)
119            .build()
120            .unwrap_or_else(|error| unreachable!("a pool of one provider builds: {error}"));
121        Arc::new(PolicyRouter::new(Arc::new(pool)))
122    }
123
124    fn lock_answers(&self) -> std::sync::MutexGuard<'_, BTreeMap<String, VecDeque<Scripted>>> {
125        self.answers.lock().unwrap_or_else(PoisonError::into_inner)
126    }
127
128    fn take(&self, call: &str) -> Option<Scripted> {
129        let mut answers = self.lock_answers();
130        let task = call.split('#').next().unwrap_or(call);
131        for key in [call, task] {
132            if let Some(answer) = answers.get_mut(key).and_then(VecDeque::pop_front) {
133                return Some(answer);
134            }
135        }
136        None
137    }
138}
139
140#[async_trait]
141impl ModelProvider for ScriptedTasks {
142    fn provider_key(&self) -> ProviderKey {
143        self.profile.provider.clone()
144    }
145
146    fn model_key(&self) -> ModelKey {
147        self.profile.model.clone()
148    }
149
150    fn capabilities(&self) -> ProviderCapabilities {
151        self.profile.capabilities.clone()
152    }
153
154    fn profile(&self) -> ModelProfile {
155        self.profile.clone()
156    }
157
158    async fn generate(&self, request: ModelRequest) -> Result<ModelResponse, ProviderError> {
159        self.calls
160            .lock()
161            .unwrap_or_else(PoisonError::into_inner)
162            .push(request.clone());
163        let Some(call) = request.metadata.get(TASK_LABEL) else {
164            // Not a task: another provider in the pool may answer it.
165            return Err(ProviderError::unsupported("untasked_request"));
166        };
167        let answer = match self.take(call) {
168            Some(Scripted::Answer(answer)) => answer,
169            Some(Scripted::Failure(error)) => return Err(error),
170            None => {
171                return Err(ProviderError::invalid_request("unscripted_task").with_detail(call));
172            }
173        };
174        // Usage a budget can be tested against: about four characters a token.
175        let text = answer.to_string();
176        let prompt: usize = request
177            .messages
178            .iter()
179            .map(|message| message.text().len())
180            .sum();
181        let prompt = prompt + request.system.as_ref().map_or(0, String::len);
182        let tokens = |characters: usize| u64::try_from(characters / 4).unwrap_or(u64::MAX);
183        let usage = TokenUsage::new(tokens(prompt), tokens(text.len()));
184        Ok(ModelResponse::new(
185            request.request_id,
186            self.profile.provider.clone(),
187            self.profile.model.clone(),
188        )
189        .with_text(text)
190        .with_usage(usage))
191    }
192}
193
194impl fmt::Debug for ScriptedTasks {
195    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
196        f.debug_struct("ScriptedTasks")
197            .field("model", &self.profile.reference().to_string())
198            .field("unanswered", &self.unanswered())
199            .finish_non_exhaustive()
200    }
201}