1use 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
25pub const UNEXPECTED_CALL_CODE: &str = "turnframe.test.unexpected_call";
28
29pub const WRONG_PURPOSE_CODE: &str = "turnframe.test.wrong_purpose";
32
33pub const UNSERIALIZABLE_REPLY_CODE: &str = "turnframe.test.unserializable_reply";
36
37#[derive(Debug, Clone)]
44#[non_exhaustive]
45pub enum ScriptedReply {
46 Json(serde_json::Value),
49 Text(String),
51 ToolCall {
54 id: CallId,
56 name: String,
58 arguments: serde_json::Value,
60 },
61 MalformedJson(String),
64 Refusal(String),
68 Timeout,
70 RateLimited {
72 retry_after: Option<Duration>,
74 },
75 Fail(ProviderError),
77 Stream(Vec<StreamEvent>),
79 Chunks(Vec<String>),
82 FromSchema(fn(&serde_json::Value) -> serde_json::Value),
85}
86
87impl ScriptedReply {
88 #[must_use]
90 pub fn text(text: impl Into<String>) -> Self {
91 Self::Text(text.into())
92 }
93
94 #[must_use]
96 pub fn refusal(text: impl Into<String>) -> Self {
97 Self::Refusal(text.into())
98 }
99
100 #[must_use]
102 pub fn written(text: impl Into<String>) -> Self {
103 Self::Json(serde_json::json!({ "text": text.into() }))
104 }
105
106 #[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 #[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 #[must_use]
120 pub fn review_passes() -> Self {
121 Self::FromSchema(|schema| review(schema, false))
122 }
123
124 #[must_use]
126 pub fn review_fails() -> Self {
127 Self::FromSchema(|schema| review(schema, true))
128 }
129
130 #[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 #[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 #[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
168fn 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#[derive(Debug, Clone)]
192pub struct ScriptStep {
193 pub reply: ScriptedReply,
195 pub expected_purpose: Option<ModelPurpose>,
197}
198
199impl ScriptStep {
200 #[must_use]
202 pub fn new(reply: ScriptedReply) -> Self {
203 Self {
204 reply,
205 expected_purpose: None,
206 }
207 }
208
209 #[must_use]
211 pub fn expecting(mut self, purpose: ModelPurpose) -> Self {
212 self.expected_purpose = Some(purpose);
213 self
214 }
215}
216
217#[derive(Debug, Clone, PartialEq)]
225pub struct RecordedCall {
226 pub index: usize,
228 pub streamed: bool,
232 pub request: ModelRequest,
234}
235
236impl RecordedCall {
237 #[must_use]
239 pub fn purpose(&self) -> ModelPurpose {
240 self.request.purpose
241 }
242
243 #[must_use]
245 pub fn schema(&self) -> Option<&serde_json::Value> {
246 self.request.output.schema()
247 }
248
249 #[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 #[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 #[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 #[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 #[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 #[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#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
324#[non_exhaustive]
325pub enum ScriptViolation {
326 #[error("call {call_index} ({purpose}) arrived after the script ran out of steps")]
328 UnexpectedCall {
329 call_index: usize,
331 purpose: &'static str,
333 },
334 #[error("call {call_index} carried purpose {found}, but the next step expected {expected}")]
336 WrongPurpose {
337 call_index: usize,
339 expected: &'static str,
341 found: &'static str,
343 },
344 #[error("{remaining} scripted step(s) were never reached; the next one would answer {next}")]
346 StepsUnused {
347 remaining: usize,
349 next: &'static str,
351 },
352}
353
354pub 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#[derive(Debug)]
390pub struct ScriptedProviderBuilder {
391 profile: ModelProfile,
392 steps: VecDeque<ScriptStep>,
393 usage: TokenUsage,
394 latency: Duration,
395}
396
397impl ScriptedProviderBuilder {
398 #[must_use]
403 pub fn capabilities(mut self, capabilities: ProviderCapabilities) -> Self {
404 self.profile.capabilities = capabilities;
405 self
406 }
407
408 #[must_use]
411 pub fn streaming(mut self) -> Self {
412 self.profile.capabilities = self.profile.capabilities.with_streaming(true);
413 self
414 }
415
416 #[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 #[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 #[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 #[must_use]
440 pub fn usage(mut self, usage: TokenUsage) -> Self {
441 self.usage = usage;
442 self
443 }
444
445 #[must_use]
447 pub fn latency(mut self, latency: Duration) -> Self {
448 self.latency = latency;
449 self
450 }
451
452 #[must_use]
454 pub fn step(mut self, step: ScriptStep) -> Self {
455 self.steps.push_back(step);
456 self
457 }
458
459 #[must_use]
461 pub fn reply(self, reply: ScriptedReply) -> Self {
462 self.step(ScriptStep::new(reply))
463 }
464
465 #[must_use]
467 pub fn reply_to(self, purpose: ModelPurpose, reply: ScriptedReply) -> Self {
468 self.step(ScriptStep::new(reply).expecting(purpose))
469 }
470
471 #[must_use]
473 pub fn text(self, text: impl Into<String>) -> Self {
474 self.reply(ScriptedReply::text(text))
475 }
476
477 #[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 #[must_use]
486 pub fn answering(self, text: impl Into<String>) -> Self {
487 self.reply_to(ModelPurpose::Answer, ScriptedReply::answer(text))
488 }
489
490 #[must_use]
492 pub fn malformed_json(self, body: impl Into<String>) -> Self {
493 self.reply(ScriptedReply::MalformedJson(body.into()))
494 }
495
496 #[must_use]
498 pub fn refusing(self, text: impl Into<String>) -> Self {
499 self.reply(ScriptedReply::refusal(text))
500 }
501
502 #[must_use]
504 pub fn timing_out(self) -> Self {
505 self.reply(ScriptedReply::Timeout)
506 }
507
508 #[must_use]
510 pub fn rate_limited(self, retry_after: Option<Duration>) -> Self {
511 self.reply(ScriptedReply::RateLimited { retry_after })
512 }
513
514 #[must_use]
516 pub fn failing(self, error: ProviderError) -> Self {
517 self.reply(ScriptedReply::Fail(error))
518 }
519
520 #[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 #[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 #[must_use]
546 pub fn build_shared(self) -> Arc<ScriptedProvider> {
547 Arc::new(self.build())
548 }
549}
550
551impl ScriptedProvider {
552 #[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 #[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 #[must_use]
590 pub fn profile_ref(&self) -> &ModelProfile {
591 &self.profile
592 }
593
594 #[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 #[must_use]
607 pub fn calls(&self) -> Vec<RecordedCall> {
608 self.lock(&self.calls).clone()
609 }
610
611 #[must_use]
613 pub fn call_count(&self) -> usize {
614 self.lock(&self.calls).len()
615 }
616
617 #[must_use]
619 pub fn nth_call(&self, index: usize) -> Option<RecordedCall> {
620 self.lock(&self.calls).get(index).cloned()
621 }
622
623 #[must_use]
625 pub fn last_call(&self) -> Option<RecordedCall> {
626 self.lock(&self.calls).last().cloned()
627 }
628
629 #[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 #[must_use]
641 pub fn remaining_steps(&self) -> usize {
642 self.lock(&self.steps).len()
643 }
644
645 #[must_use]
647 pub fn violations(&self) -> Vec<ScriptViolation> {
648 self.lock(&self.violations).clone()
649 }
650
651 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 pub fn push(&self, step: ScriptStep) {
679 self.lock(&self.steps).push_back(step);
680 }
681
682 pub fn clear_calls(&self) {
684 self.lock(&self.calls).clear();
685 self.lock(&self.violations).clear();
686 }
687
688 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 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 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 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 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
836fn 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 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#[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 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 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}