turnframe_tasks/
testing.rs1use 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
25enum Scripted {
27 Answer(serde_json::Value),
28 Failure(ProviderError),
29}
30
31pub struct ScriptedTasks {
33 profile: ModelProfile,
34 answers: Mutex<BTreeMap<String, VecDeque<Scripted>>>,
35 calls: Mutex<Vec<ModelRequest>>,
36}
37
38impl ScriptedTasks {
39 #[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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 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 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}