Skip to main content

turnframe_test/providers/
script.rs

1//! A [`ModelProvider`] whose whole behaviour is a script the test writes.
2//!
3//! See the [module documentation](crate::providers) for why the kit ships this
4//! next to [`StaticProvider`](turnframe_provider::testing::StaticProvider).
5
6use std::collections::VecDeque;
7use std::fmt;
8use std::sync::{Arc, Mutex, MutexGuard, PoisonError};
9use std::time::Duration;
10
11use async_trait::async_trait;
12use turnframe_provider::capabilities::{
13    MicroCents, ModelProfile, ProviderCapabilities, StructuredOutputCapability,
14    ToolCallingCapability,
15};
16use turnframe_provider::error::ProviderError;
17use turnframe_provider::ids::{CallId, ModelKey, ProviderKey};
18use turnframe_provider::provider::ModelProvider;
19use turnframe_provider::purpose::ModelPurpose;
20use turnframe_provider::request::{ContentPart, ModelRequest, Role, ToolCall};
21use turnframe_provider::response::{FinishReason, ModelResponse, TokenUsage};
22use turnframe_provider::router::ProviderCandidate;
23use turnframe_provider::stream::{ModelStream, StreamAccumulator, StreamEvent};
24
25/// Error code of the failure returned for a call the script did not
26/// anticipate.
27pub const UNEXPECTED_CALL_CODE: &str = "turnframe.test.unexpected_call";
28
29/// Error code of the failure returned when a call arrives for another purpose
30/// than the next step expected.
31pub const WRONG_PURPOSE_CODE: &str = "turnframe.test.wrong_purpose";
32
33/// Error code of the failure returned when a scripted value cannot be
34/// serialized. Only reachable with a custom `Serialize` that fails.
35pub const UNSERIALIZABLE_REPLY_CODE: &str = "turnframe.test.unserializable_reply";
36
37/// What one scripted call answers with.
38///
39/// The variants cover the answers a turn has to survive: the plan you meant,
40/// the plan you did not mean, a body that is not JSON at all, a refusal, and
41/// the three transport failures whose handling differs (timeout, rate limit,
42/// anything else).
43#[derive(Debug, Clone)]
44#[non_exhaustive]
45pub enum ScriptedReply {
46    /// An arbitrary JSON document, for the shapes no typed value can express โ€”
47    /// an act with a field the schema forbids, a plan missing a required field.
48    Json(serde_json::Value),
49    /// Prose.
50    Text(String),
51    /// A single tool call, the shape a native function-schema transport
52    /// produces.
53    ToolCall {
54        /// Call id echoed back to the runtime.
55        id: CallId,
56        /// Tool name.
57        name: String,
58        /// Arguments, still untrusted.
59        arguments: serde_json::Value,
60    },
61    /// A body that is not JSON. Sent verbatim, so the test decides *how*
62    /// broken it is.
63    MalformedJson(String),
64    /// The model declined. A semantic outcome
65    /// ([`FinishReason::Refusal`]), not a transport failure โ€” the difference a
66    /// runtime must not lose.
67    Refusal(String),
68    /// The call did not answer in time.
69    Timeout,
70    /// The provider refused the call and asked the caller to wait.
71    RateLimited {
72        /// The `Retry-After` hint, when the provider gave one.
73        retry_after: Option<Duration>,
74    },
75    /// Any other normalized failure.
76    Fail(ProviderError),
77    /// Exactly these stream events, in this order.
78    Stream(Vec<StreamEvent>),
79    /// This text, delivered in these chunks. Equivalent to a [`Self::Stream`]
80    /// of one text delta per chunk followed by a stop, and shorter to write.
81    Chunks(Vec<String>),
82    /// A document built from the request's output schema, for a task whose schema
83    /// depends on its input.
84    FromSchema(fn(&serde_json::Value) -> serde_json::Value),
85}
86
87impl ScriptedReply {
88    /// Prose.
89    #[must_use]
90    pub fn text(text: impl Into<String>) -> Self {
91        Self::Text(text.into())
92    }
93
94    /// A refusal.
95    #[must_use]
96    pub fn refusal(text: impl Into<String>) -> Self {
97        Self::Refusal(text.into())
98    }
99
100    /// What the acknowledge and answer tasks write: a document holding `text`.
101    #[must_use]
102    pub fn written(text: impl Into<String>) -> Self {
103        Self::Json(serde_json::json!({ "text": text.into() }))
104    }
105
106    /// An answer the answer task writes.
107    #[must_use]
108    pub fn answer(text: impl Into<String>) -> Self {
109        Self::Json(serde_json::json!({ "kind": "answered", "text": text.into() }))
110    }
111
112    /// The answer task saying the facts do not answer the question.
113    #[must_use]
114    pub fn cannot_answer(reason: impl Into<String>) -> Self {
115        Self::Json(serde_json::json!({ "kind": "cannot_answer", "text": reason.into() }))
116    }
117
118    /// A review that passes the reply: each check it was asked answered in its favour.
119    #[must_use]
120    pub fn review_passes() -> Self {
121        Self::FromSchema(|schema| review(schema, false))
122    }
123
124    /// A review that finds the reply claims what its material does not hold.
125    #[must_use]
126    pub fn review_fails() -> Self {
127        Self::FromSchema(|schema| review(schema, true))
128    }
129
130    /// A rate limit with a `Retry-After` hint in seconds.
131    #[must_use]
132    pub fn rate_limited_after(seconds: u64) -> Self {
133        Self::RateLimited {
134            retry_after: Some(Duration::from_secs(seconds)),
135        }
136    }
137
138    /// A text answer split into chunks.
139    #[must_use]
140    pub fn chunks<I, S>(chunks: I) -> Self
141    where
142        I: IntoIterator<Item = S>,
143        S: Into<String>,
144    {
145        Self::Chunks(chunks.into_iter().map(Into::into).collect())
146    }
147
148    /// Stable snake-case label of the variant, safe in a failure message: it
149    /// never carries the scripted value.
150    #[must_use]
151    pub const fn label(&self) -> &'static str {
152        match self {
153            Self::Json(_) => "json",
154            Self::Text(_) => "text",
155            Self::ToolCall { .. } => "tool_call",
156            Self::MalformedJson(_) => "malformed_json",
157            Self::Refusal(_) => "refusal",
158            Self::Timeout => "timeout",
159            Self::RateLimited { .. } => "rate_limited",
160            Self::Fail(_) => "fail",
161            Self::Stream(_) => "stream",
162            Self::Chunks(_) => "chunks",
163            Self::FromSchema(_) => "from_schema",
164        }
165    }
166}
167
168/// A review answering every check its schema asks: the reply asks its ask, and claims
169/// beyond its material only when `claims`.
170fn review(schema: &serde_json::Value, claims: bool) -> serde_json::Value {
171    let mut document = serde_json::Map::new();
172    let checks = schema["properties"]
173        .as_object()
174        .into_iter()
175        .flat_map(|p| p.keys());
176    for check in checks {
177        let answer = match check.as_str() {
178            "reasoning" => serde_json::Value::from("Judged against the material."),
179            "asks_the_ask" => serde_json::Value::from(true),
180            "claims_beyond_material" => serde_json::Value::from(claims),
181            _ => serde_json::Value::from(false),
182        };
183        document.insert(check.clone(), answer);
184    }
185    serde_json::Value::Object(document)
186}
187
188/// One entry of a script: what to answer, and what the call is expected to be.
189#[derive(Debug, Clone)]
190pub struct ScriptStep {
191    /// The answer.
192    pub reply: ScriptedReply,
193    /// Purpose the call must carry. `None` accepts any purpose.
194    pub expected_purpose: Option<ModelPurpose>,
195}
196
197impl ScriptStep {
198    /// A step that accepts any call.
199    #[must_use]
200    pub fn new(reply: ScriptedReply) -> Self {
201        Self {
202            reply,
203            expected_purpose: None,
204        }
205    }
206
207    /// Restricts the step to one purpose.
208    #[must_use]
209    pub fn expecting(mut self, purpose: ModelPurpose) -> Self {
210        self.expected_purpose = Some(purpose);
211        self
212    }
213}
214
215/// One request the provider received.
216///
217/// The point of recording is to assert what the runtime *sent*, not only what
218/// it did with the answer: the catalog reaches a model as the output schema
219/// and the declared tools, so [`schema_mentions`](Self::schema_mentions) and
220/// [`tool_names`](Self::tool_names) are how a test checks that a workflow's
221/// operations were actually offered.
222#[derive(Debug, Clone, PartialEq)]
223pub struct RecordedCall {
224    /// Zero-based position in the call sequence.
225    pub index: usize,
226    /// `true` when the call arrived through
227    /// [`ModelProvider::stream`], `false` through
228    /// [`generate`](ModelProvider::generate).
229    pub streamed: bool,
230    /// The request, verbatim.
231    pub request: ModelRequest,
232}
233
234impl RecordedCall {
235    /// Why the call was made.
236    #[must_use]
237    pub fn purpose(&self) -> ModelPurpose {
238        self.request.purpose
239    }
240
241    /// The JSON Schema the call demanded, when it demanded one.
242    #[must_use]
243    pub fn schema(&self) -> Option<&serde_json::Value> {
244        self.request.output.schema()
245    }
246
247    /// Name the schema was labelled with.
248    #[must_use]
249    pub fn schema_name(&self) -> Option<&str> {
250        match &self.request.output {
251            turnframe_provider::request::OutputSpec::Json { name, .. } => Some(name),
252            _ => None,
253        }
254    }
255
256    /// Returns `true` when `needle` occurs anywhere in the serialized schema.
257    ///
258    /// This is how a test asks "was `trip.set_name` in the catalog this
259    /// call offered?" without depending on how the schema is shaped.
260    #[must_use]
261    pub fn schema_mentions(&self, needle: &str) -> bool {
262        self.schema()
263            .is_some_and(|schema| schema.to_string().contains(needle))
264    }
265
266    /// Names of the read-only tools the call declared, in order.
267    #[must_use]
268    pub fn tool_names(&self) -> Vec<&str> {
269        self.request
270            .tools
271            .iter()
272            .map(|tool| tool.name.as_str())
273            .collect()
274    }
275
276    /// Every message as `(role, flattened text)`, in order.
277    #[must_use]
278    pub fn messages(&self) -> Vec<(Role, String)> {
279        self.request
280            .messages
281            .iter()
282            .map(|message| (message.role, message.text()))
283            .collect()
284    }
285
286    /// The concatenated text of every user message.
287    #[must_use]
288    pub fn user_text(&self) -> String {
289        self.request
290            .messages
291            .iter()
292            .filter(|message| message.role == Role::User)
293            .map(turnframe_provider::request::Message::text)
294            .collect::<Vec<_>>()
295            .join("\n")
296    }
297
298    /// Returns `true` when `needle` occurs in the system prompt or in any
299    /// message of the call.
300    #[must_use]
301    pub fn prompt_mentions(&self, needle: &str) -> bool {
302        if self
303            .request
304            .system
305            .as_deref()
306            .is_some_and(|system| system.contains(needle))
307        {
308            return true;
309        }
310        self.request
311            .messages
312            .iter()
313            .any(|message| message.text().contains(needle))
314    }
315}
316
317/// A call the script did not anticipate, or a step it never reached.
318///
319/// Every variant names positions, purposes and labels โ€” never a prompt, a
320/// scripted value or anything a user typed.
321#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
322#[non_exhaustive]
323pub enum ScriptViolation {
324    /// A call arrived after the script ran out of steps.
325    #[error("call {call_index} ({purpose}) arrived after the script ran out of steps")]
326    UnexpectedCall {
327        /// Position of the call.
328        call_index: usize,
329        /// Purpose the call carried.
330        purpose: &'static str,
331    },
332    /// A call arrived for another purpose than the next step expected.
333    #[error("call {call_index} carried purpose {found}, but the next step expected {expected}")]
334    WrongPurpose {
335        /// Position of the call.
336        call_index: usize,
337        /// Purpose the step expected.
338        expected: &'static str,
339        /// Purpose the call carried.
340        found: &'static str,
341    },
342    /// The script still had steps when it was verified.
343    #[error("{remaining} scripted step(s) were never reached; the next one would answer {next}")]
344    StepsUnused {
345        /// How many steps are left.
346        remaining: usize,
347        /// Label of the first unused reply.
348        next: &'static str,
349    },
350}
351
352/// A [`ModelProvider`] that answers from a script and refuses to improvise.
353///
354/// Steps are consumed in order. When the script runs out, the provider does
355/// **not** fall back to a default answer: it records a
356/// [`ScriptViolation::UnexpectedCall`] and fails the call with
357/// [`UNEXPECTED_CALL_CODE`], so a runtime that calls a model one more time than
358/// the test believed cannot pass silently.
359///
360/// ```
361/// use turnframe_provider::prelude::*;
362/// use turnframe_test::providers::{ScriptedProvider, ScriptedReply};
363///
364/// # futures::executor::block_on(async {
365/// let provider = ScriptedProvider::builder("fake", "m")
366///     .reply(ScriptedReply::text("una frase"))
367///     .build();
368///
369/// let request = ModelRequest::new(ModelPurpose::Acknowledge).with_message(Message::user("ciao"));
370/// assert_eq!(provider.generate(request.clone()).await.unwrap().text(), "una frase");
371///
372/// // The second call was never scripted.
373/// assert!(provider.generate(request).await.is_err());
374/// assert!(provider.verify().is_err());
375/// # });
376/// ```
377pub struct ScriptedProvider {
378    profile: ModelProfile,
379    steps: Mutex<VecDeque<ScriptStep>>,
380    calls: Mutex<Vec<RecordedCall>>,
381    violations: Mutex<Vec<ScriptViolation>>,
382    usage: TokenUsage,
383    latency: Duration,
384}
385
386/// Collects steps and capabilities for a [`ScriptedProvider`].
387#[derive(Debug)]
388pub struct ScriptedProviderBuilder {
389    profile: ModelProfile,
390    steps: VecDeque<ScriptStep>,
391    usage: TokenUsage,
392    latency: Duration,
393}
394
395impl ScriptedProviderBuilder {
396    /// Declares capabilities, replacing the defaults.
397    ///
398    /// They may be dishonest on purpose: a test of the no-silent-downgrade rule
399    /// needs a profile that claims more than it delivers.
400    #[must_use]
401    pub fn capabilities(mut self, capabilities: ProviderCapabilities) -> Self {
402        self.profile.capabilities = capabilities;
403        self
404    }
405
406    /// Declares streaming support, which [`ModelProvider::stream`] refuses
407    /// without.
408    #[must_use]
409    pub fn streaming(mut self) -> Self {
410        self.profile.capabilities = self.profile.capabilities.with_streaming(true);
411        self
412    }
413
414    /// Declares a structured-output transport.
415    #[must_use]
416    pub fn structured_output(mut self, capability: StructuredOutputCapability) -> Self {
417        self.profile.capabilities = self.profile.capabilities.with_structured_output(capability);
418        self
419    }
420
421    /// Declares a region, so residency routing can be exercised.
422    #[must_use]
423    pub fn region(mut self, region: impl Into<String>) -> Self {
424        self.profile.region = Some(region.into());
425        self
426    }
427
428    /// Declares per-million prices, so cost ceilings can be exercised.
429    #[must_use]
430    pub fn cost(mut self, input: MicroCents, output: MicroCents) -> Self {
431        self.profile.cost_per_million_input = Some(input);
432        self.profile.cost_per_million_output = Some(output);
433        self
434    }
435
436    /// Reports this usage on every successful answer.
437    #[must_use]
438    pub fn usage(mut self, usage: TokenUsage) -> Self {
439        self.usage = usage;
440        self
441    }
442
443    /// Reports this latency on every successful answer.
444    #[must_use]
445    pub fn latency(mut self, latency: Duration) -> Self {
446        self.latency = latency;
447        self
448    }
449
450    /// Appends one step.
451    #[must_use]
452    pub fn step(mut self, step: ScriptStep) -> Self {
453        self.steps.push_back(step);
454        self
455    }
456
457    /// Appends a reply that accepts any purpose.
458    #[must_use]
459    pub fn reply(self, reply: ScriptedReply) -> Self {
460        self.step(ScriptStep::new(reply))
461    }
462
463    /// Appends a reply that only accepts calls made for `purpose`.
464    #[must_use]
465    pub fn reply_to(self, purpose: ModelPurpose, reply: ScriptedReply) -> Self {
466        self.step(ScriptStep::new(reply).expecting(purpose))
467    }
468
469    /// Appends prose.
470    #[must_use]
471    pub fn text(self, text: impl Into<String>) -> Self {
472        self.reply(ScriptedReply::text(text))
473    }
474
475    /// Appends the acknowledgement a turn writes, `text`, and the review that passes it.
476    #[must_use]
477    pub fn acknowledging(self, text: impl Into<String>) -> Self {
478        self.reply_to(ModelPurpose::Acknowledge, ScriptedReply::written(text))
479            .reply_to(ModelPurpose::Review, ScriptedReply::review_passes())
480    }
481
482    /// Appends one answer, `text`.
483    #[must_use]
484    pub fn answering(self, text: impl Into<String>) -> Self {
485        self.reply_to(ModelPurpose::Answer, ScriptedReply::answer(text))
486    }
487
488    /// Appends a body that is not JSON.
489    #[must_use]
490    pub fn malformed_json(self, body: impl Into<String>) -> Self {
491        self.reply(ScriptedReply::MalformedJson(body.into()))
492    }
493
494    /// Appends a refusal.
495    #[must_use]
496    pub fn refusing(self, text: impl Into<String>) -> Self {
497        self.reply(ScriptedReply::refusal(text))
498    }
499
500    /// Appends a timeout.
501    #[must_use]
502    pub fn timing_out(self) -> Self {
503        self.reply(ScriptedReply::Timeout)
504    }
505
506    /// Appends a rate limit.
507    #[must_use]
508    pub fn rate_limited(self, retry_after: Option<Duration>) -> Self {
509        self.reply(ScriptedReply::RateLimited { retry_after })
510    }
511
512    /// Appends a normalized failure.
513    #[must_use]
514    pub fn failing(self, error: ProviderError) -> Self {
515        self.reply(ScriptedReply::Fail(error))
516    }
517
518    /// Appends a stream delivered in these chunks.
519    #[must_use]
520    pub fn streaming_chunks<I, S>(self, chunks: I) -> Self
521    where
522        I: IntoIterator<Item = S>,
523        S: Into<String>,
524    {
525        self.streaming().reply(ScriptedReply::chunks(chunks))
526    }
527
528    /// Builds the provider.
529    #[must_use]
530    pub fn build(self) -> ScriptedProvider {
531        ScriptedProvider {
532            profile: self.profile,
533            steps: Mutex::new(self.steps),
534            calls: Mutex::new(Vec::new()),
535            violations: Mutex::new(Vec::new()),
536            usage: self.usage,
537            latency: self.latency,
538        }
539    }
540
541    /// Builds the provider behind an [`Arc`], the shape routing and fallback
542    /// take.
543    #[must_use]
544    pub fn build_shared(self) -> Arc<ScriptedProvider> {
545        Arc::new(self.build())
546    }
547}
548
549impl ScriptedProvider {
550    /// A builder for a provider that declares native JSON Schema output,
551    /// parallel tool calling and no streaming.
552    ///
553    /// Those are the capabilities a mutation-capable interpretation stage
554    /// requires (spec ยง20.4), which is what most tests need; call
555    /// [`capabilities`](ScriptedProviderBuilder::capabilities) to say something
556    /// else.
557    #[must_use]
558    pub fn builder(
559        provider: impl Into<ProviderKey>,
560        model: impl Into<ModelKey>,
561    ) -> ScriptedProviderBuilder {
562        ScriptedProviderBuilder {
563            profile: ModelProfile::new(
564                provider,
565                model,
566                ProviderCapabilities::minimal()
567                    .with_structured_output(StructuredOutputCapability::NativeJsonSchema)
568                    .with_tool_calling(ToolCallingCapability::Parallel)
569                    .with_preserves_call_ids(true),
570            ),
571            steps: VecDeque::new(),
572            usage: TokenUsage::none(),
573            latency: Duration::ZERO,
574        }
575    }
576
577    /// The one-line common case: a provider that acknowledges one turn with `text`,
578    /// passes its review, and refuses everything else.
579    #[must_use]
580    pub fn narrating(text: impl Into<String>) -> Self {
581        Self::builder("scripted", "narrator-1")
582            .acknowledging(text)
583            .build()
584    }
585
586    /// The routing profile, as configured.
587    #[must_use]
588    pub fn profile_ref(&self) -> &ModelProfile {
589        &self.profile
590    }
591
592    /// Wraps a shared provider as a healthy routing candidate.
593    #[must_use]
594    pub fn candidate(provider: Arc<Self>) -> ProviderCandidate {
595        let profile = provider.profile.clone();
596        ProviderCandidate {
597            provider,
598            profile,
599            healthy: true,
600        }
601    }
602
603    /// Every request received, in order.
604    #[must_use]
605    pub fn calls(&self) -> Vec<RecordedCall> {
606        self.lock(&self.calls).clone()
607    }
608
609    /// How many requests were received.
610    #[must_use]
611    pub fn call_count(&self) -> usize {
612        self.lock(&self.calls).len()
613    }
614
615    /// The `index`-th request received.
616    #[must_use]
617    pub fn nth_call(&self, index: usize) -> Option<RecordedCall> {
618        self.lock(&self.calls).get(index).cloned()
619    }
620
621    /// The most recent request received.
622    #[must_use]
623    pub fn last_call(&self) -> Option<RecordedCall> {
624        self.lock(&self.calls).last().cloned()
625    }
626
627    /// Every request made for one purpose, in order.
628    #[must_use]
629    pub fn calls_for(&self, purpose: ModelPurpose) -> Vec<RecordedCall> {
630        self.lock(&self.calls)
631            .iter()
632            .filter(|call| call.request.purpose == purpose)
633            .cloned()
634            .collect()
635    }
636
637    /// How many steps are still unused.
638    #[must_use]
639    pub fn remaining_steps(&self) -> usize {
640        self.lock(&self.steps).len()
641    }
642
643    /// Every violation the script observed, in order.
644    #[must_use]
645    pub fn violations(&self) -> Vec<ScriptViolation> {
646        self.lock(&self.violations).clone()
647    }
648
649    /// Checks that the script was followed exactly: no unanticipated call and
650    /// no unused step.
651    ///
652    /// A call the script did not anticipate already failed at the time it was
653    /// made; this is how the *test* learns about it even when the code under
654    /// test swallowed the error.
655    ///
656    /// # Errors
657    ///
658    /// The first [`ScriptViolation`] observed, or
659    /// [`ScriptViolation::StepsUnused`] when steps remain.
660    pub fn verify(&self) -> Result<(), ScriptViolation> {
661        if let Some(violation) = self.lock(&self.violations).first() {
662            return Err(violation.clone());
663        }
664        let steps = self.lock(&self.steps);
665        match steps.front() {
666            None => Ok(()),
667            Some(next) => Err(ScriptViolation::StepsUnused {
668                remaining: steps.len(),
669                next: next.reply.label(),
670            }),
671        }
672    }
673
674    /// Appends a step to a provider already in use, for a turn whose script
675    /// depends on identifiers an earlier turn produced.
676    pub fn push(&self, step: ScriptStep) {
677        self.lock(&self.steps).push_back(step);
678    }
679
680    /// Forgets the recorded calls and violations, keeping the remaining script.
681    pub fn clear_calls(&self) {
682        self.lock(&self.calls).clear();
683        self.lock(&self.violations).clear();
684    }
685
686    /// Locks, recovering from poisoning: a test that already failed must not
687    /// cascade into unrelated failures.
688    fn lock<'a, T>(&self, target: &'a Mutex<T>) -> MutexGuard<'a, T> {
689        target.lock().unwrap_or_else(PoisonError::into_inner)
690    }
691
692    fn record_violation(&self, violation: ScriptViolation) -> ProviderError {
693        let error = match &violation {
694            ScriptViolation::WrongPurpose { .. } => ProviderError::other(WRONG_PURPOSE_CODE),
695            _ => ProviderError::other(UNEXPECTED_CALL_CODE),
696        };
697        self.lock(&self.violations).push(violation);
698        error.with_model(&self.profile.reference())
699    }
700
701    /// Records the call and takes the next step, or explains why there is none.
702    fn take_step(
703        &self,
704        request: &ModelRequest,
705        streamed: bool,
706    ) -> Result<ScriptStep, ProviderError> {
707        let index = {
708            let mut calls = self.lock(&self.calls);
709            let index = calls.len();
710            calls.push(RecordedCall {
711                index,
712                streamed,
713                request: request.clone(),
714            });
715            index
716        };
717        let Some(step) = self.lock(&self.steps).pop_front() else {
718            return Err(self.record_violation(ScriptViolation::UnexpectedCall {
719                call_index: index,
720                purpose: request.purpose.as_str(),
721            }));
722        };
723        if let Some(expected) = step.expected_purpose
724            && expected != request.purpose
725        {
726            return Err(self.record_violation(ScriptViolation::WrongPurpose {
727                call_index: index,
728                expected: expected.as_str(),
729                found: request.purpose.as_str(),
730            }));
731        }
732        Ok(step)
733    }
734
735    /// Turns a reply into a response.
736    fn respond(
737        &self,
738        request: &ModelRequest,
739        reply: &ScriptedReply,
740    ) -> Result<ModelResponse, ProviderError> {
741        let base = ModelResponse::new(
742            request.request_id,
743            self.profile.provider.clone(),
744            self.profile.model.clone(),
745        )
746        .with_usage(self.usage)
747        .with_latency(self.latency);
748        match reply {
749            ScriptedReply::Json(value) => Ok(base.with_text(value.to_string())),
750            ScriptedReply::FromSchema(build) => {
751                let schema = request.output.schema().cloned().unwrap_or_default();
752                Ok(base.with_text(build(&schema).to_string()))
753            }
754            ScriptedReply::Text(text) => Ok(base.with_text(text.clone())),
755            ScriptedReply::ToolCall {
756                id,
757                name,
758                arguments,
759            } => Ok(base
760                .with_tool_call(ToolCall::new(id.clone(), name.clone(), arguments.clone()))
761                .with_finish(FinishReason::ToolCalls)),
762            ScriptedReply::MalformedJson(body) => Ok(base.with_text(body.clone())),
763            ScriptedReply::Refusal(text) => Ok(base
764                .with_text(text.clone())
765                .with_finish(FinishReason::Refusal)),
766            ScriptedReply::Timeout => Err(self.label(ProviderError::timeout())),
767            ScriptedReply::RateLimited { retry_after } => {
768                Err(self.label(ProviderError::rate_limited(*retry_after)))
769            }
770            ScriptedReply::Fail(error) => Err(self.label(error.clone())),
771            ScriptedReply::Stream(events) => self.reassemble(request, events.clone()),
772            ScriptedReply::Chunks(chunks) => self.reassemble(request, chunk_events(chunks)),
773        }
774    }
775
776    /// Renders a reply as the events a stream would deliver.
777    fn events_for(&self, reply: &ScriptedReply, response: &ModelResponse) -> Vec<StreamEvent> {
778        match reply {
779            ScriptedReply::Stream(events) => return events.clone(),
780            ScriptedReply::Chunks(chunks) => return chunk_events(chunks),
781            _ => {}
782        }
783        let mut events = Vec::new();
784        let text = response.text();
785        if !text.is_empty() {
786            events.push(StreamEvent::text(text));
787        }
788        for call in response.tool_calls() {
789            events.push(StreamEvent::tool_call_start(
790                call.id.clone(),
791                call.name.clone(),
792            ));
793            events.push(StreamEvent::tool_call_delta(
794                call.id.clone(),
795                call.arguments.to_string(),
796            ));
797            events.push(StreamEvent::tool_call_end(call.id.clone()));
798        }
799        if !response.usage.is_unreported() {
800            events.push(StreamEvent::Usage {
801                usage: response.usage,
802            });
803        }
804        events.push(StreamEvent::Finish {
805            reason: response.finish,
806        });
807        events
808    }
809
810    /// Rebuilds the whole answer from scripted events, so `generate` and
811    /// `stream` never disagree about the same step.
812    fn reassemble(
813        &self,
814        request: &ModelRequest,
815        events: Vec<StreamEvent>,
816    ) -> Result<ModelResponse, ProviderError> {
817        let mut accumulator = StreamAccumulator::new(
818            request.request_id,
819            self.profile.provider.clone(),
820            self.profile.model.clone(),
821        )
822        .with_latency(self.latency);
823        for event in events {
824            accumulator.push(event)?;
825        }
826        accumulator.finish()
827    }
828
829    fn label(&self, error: ProviderError) -> ProviderError {
830        error.with_model(&self.profile.reference())
831    }
832}
833
834/// One text delta per chunk, then a stop.
835fn chunk_events(chunks: &[String]) -> Vec<StreamEvent> {
836    let mut events: Vec<StreamEvent> = chunks.iter().map(StreamEvent::text).collect();
837    events.push(StreamEvent::Finish {
838        reason: FinishReason::Stop,
839    });
840    events
841}
842
843#[async_trait]
844impl ModelProvider for ScriptedProvider {
845    fn provider_key(&self) -> ProviderKey {
846        self.profile.provider.clone()
847    }
848
849    fn model_key(&self) -> ModelKey {
850        self.profile.model.clone()
851    }
852
853    fn capabilities(&self) -> ProviderCapabilities {
854        self.profile.capabilities.clone()
855    }
856
857    fn profile(&self) -> ModelProfile {
858        self.profile.clone()
859    }
860
861    async fn generate(&self, request: ModelRequest) -> Result<ModelResponse, ProviderError> {
862        let step = self.take_step(&request, false)?;
863        self.respond(&request, &step.reply)
864    }
865
866    async fn stream(&self, request: ModelRequest) -> Result<ModelStream, ProviderError> {
867        if !self.profile.capabilities.streaming {
868            // The double stays honest: a profile that declares no streaming
869            // does not get a stream synthesized for it.
870            return Err(self.label(ProviderError::unsupported("streaming")));
871        }
872        let step = self.take_step(&request, true)?;
873        let response = self.respond(&request, &step.reply)?;
874        Ok(ModelStream::from_events(
875            self.events_for(&step.reply, &response),
876        ))
877    }
878}
879
880impl fmt::Debug for ScriptedProvider {
881    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
882        f.debug_struct("ScriptedProvider")
883            .field("model", &self.profile.reference().to_string())
884            .field("remaining_steps", &self.remaining_steps())
885            .field("calls", &self.call_count())
886            .field("violations", &self.lock(&self.violations).len())
887            .finish_non_exhaustive()
888    }
889}
890
891/// Concatenated text of every part of a response, for a test that only cares
892/// about the prose.
893#[must_use]
894pub fn response_text(response: &ModelResponse) -> String {
895    response
896        .content
897        .iter()
898        .filter_map(ContentPart::as_text)
899        .collect()
900}
901
902#[cfg(test)]
903mod tests {
904    use super::*;
905    use turnframe_provider::ids::RequestId;
906    use turnframe_provider::request::{Message, ModelRequest, OutputSpec, ToolSpec};
907    use turnframe_provider::stream::reconstruct;
908
909    fn interpret() -> ModelRequest {
910        ModelRequest::new(ModelPurpose::Extract)
911            .with_request_id(RequestId::nil())
912            .with_system("Propose acts for trip.set_name.")
913            .with_message(Message::user("cambia l'oggetto"))
914            .with_output(OutputSpec::json(
915                "user_turn_plan",
916                serde_json::json!({"operations": ["trip.set_name"]}),
917            ))
918            .with_tools(vec![ToolSpec::new(
919                "case.get",
920                "load a case",
921                serde_json::json!({}),
922            )])
923    }
924
925    #[tokio::test]
926    async fn the_script_is_consumed_in_order() {
927        let provider = ScriptedProvider::builder("fake", "m")
928            .text("first")
929            .reply(ScriptedReply::Json(serde_json::json!({"acts": []})))
930            .build();
931        let first = provider.generate(interpret()).await.unwrap();
932        assert_eq!(first.text(), "first");
933        let second = provider.generate(interpret()).await.unwrap();
934        assert_eq!(second.text(), r#"{"acts":[]}"#);
935        assert_eq!(provider.remaining_steps(), 0);
936        assert!(provider.verify().is_ok());
937    }
938
939    #[tokio::test]
940    async fn a_call_the_script_did_not_anticipate_fails_loudly() {
941        let provider = ScriptedProvider::builder("fake", "m").text("once").build();
942        assert!(provider.generate(interpret()).await.is_ok());
943
944        let error = provider.generate(interpret()).await.unwrap_err();
945        assert_eq!(
946            error
947                .code()
948                .map(turnframe_provider::error::ErrorCode::as_str),
949            Some(UNEXPECTED_CALL_CODE)
950        );
951        assert_eq!(
952            provider.verify().unwrap_err(),
953            ScriptViolation::UnexpectedCall {
954                call_index: 1,
955                purpose: "extract",
956            }
957        );
958    }
959
960    #[tokio::test]
961    async fn a_step_bound_to_a_purpose_refuses_another_one() {
962        let provider = ScriptedProvider::builder("fake", "m")
963            .reply_to(ModelPurpose::Extract, ScriptedReply::text("{}"))
964            .build();
965        let error = provider
966            .generate(ModelRequest::new(ModelPurpose::Acknowledge))
967            .await
968            .unwrap_err();
969        assert_eq!(
970            error
971                .code()
972                .map(turnframe_provider::error::ErrorCode::as_str),
973            Some(WRONG_PURPOSE_CODE)
974        );
975        assert!(matches!(
976            provider.verify().unwrap_err(),
977            ScriptViolation::WrongPurpose { .. }
978        ));
979    }
980
981    #[test]
982    fn an_unused_step_is_a_violation() {
983        let provider = ScriptedProvider::builder("fake", "m")
984            .text("never asked for")
985            .build();
986        assert_eq!(
987            provider.verify().unwrap_err(),
988            ScriptViolation::StepsUnused {
989                remaining: 1,
990                next: "text",
991            }
992        );
993    }
994
995    #[tokio::test]
996    async fn every_request_is_recorded_with_what_was_sent() {
997        let provider = ScriptedProvider::builder("fake", "m").text("{}").build();
998        provider.generate(interpret()).await.unwrap();
999
1000        let call = provider.last_call().unwrap();
1001        assert_eq!(call.index, 0);
1002        assert!(!call.streamed);
1003        assert_eq!(call.purpose(), ModelPurpose::Extract);
1004        assert_eq!(call.schema_name(), Some("user_turn_plan"));
1005        assert!(call.schema_mentions("trip.set_name"));
1006        assert!(!call.schema_mentions("trip.rebook"));
1007        assert_eq!(call.tool_names(), vec!["case.get"]);
1008        assert_eq!(call.user_text(), "cambia l'oggetto");
1009        assert_eq!(
1010            call.messages(),
1011            vec![(Role::User, "cambia l'oggetto".to_owned())]
1012        );
1013        assert!(call.prompt_mentions("trip.set_name"));
1014        assert_eq!(provider.calls_for(ModelPurpose::Extract).len(), 1);
1015        assert_eq!(provider.nth_call(0), Some(call));
1016    }
1017
1018    #[tokio::test]
1019    async fn the_transport_failures_keep_their_families() {
1020        let provider = ScriptedProvider::builder("fake", "m")
1021            .timing_out()
1022            .rate_limited(Some(Duration::from_secs(3)))
1023            .refusing("non posso")
1024            .build();
1025        let timeout = provider
1026            .generate(ModelRequest::new(ModelPurpose::Acknowledge))
1027            .await
1028            .unwrap_err();
1029        assert_eq!(
1030            timeout.retry_class(),
1031            turnframe_provider::error::RetryClass::Retry
1032        );
1033        let limited = provider
1034            .generate(ModelRequest::new(ModelPurpose::Acknowledge))
1035            .await
1036            .unwrap_err();
1037        assert_eq!(limited.retry_after(), Some(Duration::from_secs(3)));
1038        assert_eq!(limited.model().map(ModelKey::as_str), Some("m"));
1039
1040        // A refusal is a response, not a transport failure.
1041        let refusal = provider
1042            .generate(ModelRequest::new(ModelPurpose::Acknowledge))
1043            .await
1044            .unwrap();
1045        assert_eq!(refusal.finish, FinishReason::Refusal);
1046        assert!(!refusal.finish.is_complete());
1047        assert_eq!(response_text(&refusal), "non posso");
1048    }
1049
1050    #[tokio::test]
1051    async fn malformed_json_is_delivered_verbatim() {
1052        let provider = ScriptedProvider::builder("fake", "m")
1053            .malformed_json("{\"acts\": [")
1054            .build();
1055        let response = provider
1056            .generate(ModelRequest::new(ModelPurpose::Extract))
1057            .await
1058            .unwrap();
1059        assert!(serde_json::from_str::<serde_json::Value>(&response.text()).is_err());
1060    }
1061
1062    #[tokio::test]
1063    async fn chunks_stream_and_reassemble_to_the_same_answer() {
1064        let provider = ScriptedProvider::builder("fake", "m")
1065            .streaming_chunks(["Ho preparato ", "la modifica."])
1066            .streaming_chunks(["Ho preparato ", "la modifica."])
1067            .build();
1068        let request =
1069            ModelRequest::new(ModelPurpose::Acknowledge).with_request_id(RequestId::nil());
1070
1071        let items = provider
1072            .stream(request.clone())
1073            .await
1074            .unwrap()
1075            .collect_items()
1076            .await;
1077        assert_eq!(items.len(), 3, "two deltas and a finish");
1078
1079        let whole = provider.generate(request).await.unwrap();
1080        assert_eq!(whole.text(), "Ho preparato la modifica.");
1081        assert!(provider.calls()[0].streamed);
1082        assert!(!provider.calls()[1].streamed);
1083    }
1084
1085    #[tokio::test]
1086    async fn streaming_is_refused_unless_declared() {
1087        let provider = ScriptedProvider::builder("fake", "m").text("x").build();
1088        let error = provider
1089            .stream(ModelRequest::new(ModelPurpose::Acknowledge))
1090            .await
1091            .unwrap_err();
1092        assert!(matches!(
1093            error.kind(),
1094            turnframe_provider::error::ProviderErrorKind::Unsupported { .. }
1095        ));
1096        // The refused call consumed no step.
1097        assert_eq!(provider.remaining_steps(), 1);
1098    }
1099
1100    #[tokio::test]
1101    async fn a_scripted_stream_reconstructs_into_the_generated_answer() {
1102        let provider = ScriptedProvider::builder("fake", "m")
1103            .streaming()
1104            .reply(ScriptedReply::Stream(vec![
1105                StreamEvent::text("ciao "),
1106                StreamEvent::text("mondo"),
1107                StreamEvent::Finish {
1108                    reason: FinishReason::Stop,
1109                },
1110            ]))
1111            .build();
1112        let stream = provider
1113            .stream(ModelRequest::new(ModelPurpose::Acknowledge).with_request_id(RequestId::nil()))
1114            .await
1115            .unwrap();
1116        let rebuilt = reconstruct(
1117            stream,
1118            StreamAccumulator::new(RequestId::nil(), "fake", "m"),
1119        )
1120        .await
1121        .unwrap();
1122        assert_eq!(rebuilt.text(), "ciao mondo");
1123    }
1124
1125    #[test]
1126    fn the_double_renders_its_state_without_the_script() {
1127        let provider = ScriptedProvider::builder("fake", "m")
1128            .text("secret")
1129            .build();
1130        let rendered = format!("{provider:?}");
1131        assert!(rendered.contains("fake/m"), "{rendered}");
1132        assert!(!rendered.contains("secret"), "{rendered}");
1133    }
1134
1135    #[test]
1136    fn the_one_line_constructors_declare_what_they_need() {
1137        let narrating = ScriptedProvider::narrating("x");
1138        assert!(narrating.capabilities().structured_output.enforces_schema());
1139        let candidate = ScriptedProvider::candidate(Arc::new(narrating));
1140        assert!(candidate.healthy);
1141        assert_eq!(candidate.reference().to_string(), "scripted/narrator-1");
1142        assert_eq!(
1143            ScriptedProvider::narrating("x")
1144                .profile_ref()
1145                .model
1146                .as_str(),
1147            "narrator-1"
1148        );
1149    }
1150
1151    #[tokio::test]
1152    async fn recorded_calls_survive_a_clear() {
1153        let provider = ScriptedProvider::builder("fake", "m")
1154            .text("a")
1155            .text("b")
1156            .build();
1157        provider
1158            .generate(ModelRequest::new(ModelPurpose::Acknowledge))
1159            .await
1160            .unwrap();
1161        assert_eq!(provider.call_count(), 1);
1162        provider.clear_calls();
1163        assert_eq!(provider.call_count(), 0);
1164        assert!(provider.violations().is_empty());
1165        assert_eq!(provider.remaining_steps(), 1);
1166    }
1167}