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