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" => 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#[derive(Debug, Clone)]
190pub struct ScriptStep {
191 pub reply: ScriptedReply,
193 pub expected_purpose: Option<ModelPurpose>,
195}
196
197impl ScriptStep {
198 #[must_use]
200 pub fn new(reply: ScriptedReply) -> Self {
201 Self {
202 reply,
203 expected_purpose: None,
204 }
205 }
206
207 #[must_use]
209 pub fn expecting(mut self, purpose: ModelPurpose) -> Self {
210 self.expected_purpose = Some(purpose);
211 self
212 }
213}
214
215#[derive(Debug, Clone, PartialEq)]
223pub struct RecordedCall {
224 pub index: usize,
226 pub streamed: bool,
230 pub request: ModelRequest,
232}
233
234impl RecordedCall {
235 #[must_use]
237 pub fn purpose(&self) -> ModelPurpose {
238 self.request.purpose
239 }
240
241 #[must_use]
243 pub fn schema(&self) -> Option<&serde_json::Value> {
244 self.request.output.schema()
245 }
246
247 #[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 #[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 #[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 #[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 #[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 #[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#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
322#[non_exhaustive]
323pub enum ScriptViolation {
324 #[error("call {call_index} ({purpose}) arrived after the script ran out of steps")]
326 UnexpectedCall {
327 call_index: usize,
329 purpose: &'static str,
331 },
332 #[error("call {call_index} carried purpose {found}, but the next step expected {expected}")]
334 WrongPurpose {
335 call_index: usize,
337 expected: &'static str,
339 found: &'static str,
341 },
342 #[error("{remaining} scripted step(s) were never reached; the next one would answer {next}")]
344 StepsUnused {
345 remaining: usize,
347 next: &'static str,
349 },
350}
351
352pub 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#[derive(Debug)]
388pub struct ScriptedProviderBuilder {
389 profile: ModelProfile,
390 steps: VecDeque<ScriptStep>,
391 usage: TokenUsage,
392 latency: Duration,
393}
394
395impl ScriptedProviderBuilder {
396 #[must_use]
401 pub fn capabilities(mut self, capabilities: ProviderCapabilities) -> Self {
402 self.profile.capabilities = capabilities;
403 self
404 }
405
406 #[must_use]
409 pub fn streaming(mut self) -> Self {
410 self.profile.capabilities = self.profile.capabilities.with_streaming(true);
411 self
412 }
413
414 #[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 #[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 #[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 #[must_use]
438 pub fn usage(mut self, usage: TokenUsage) -> Self {
439 self.usage = usage;
440 self
441 }
442
443 #[must_use]
445 pub fn latency(mut self, latency: Duration) -> Self {
446 self.latency = latency;
447 self
448 }
449
450 #[must_use]
452 pub fn step(mut self, step: ScriptStep) -> Self {
453 self.steps.push_back(step);
454 self
455 }
456
457 #[must_use]
459 pub fn reply(self, reply: ScriptedReply) -> Self {
460 self.step(ScriptStep::new(reply))
461 }
462
463 #[must_use]
465 pub fn reply_to(self, purpose: ModelPurpose, reply: ScriptedReply) -> Self {
466 self.step(ScriptStep::new(reply).expecting(purpose))
467 }
468
469 #[must_use]
471 pub fn text(self, text: impl Into<String>) -> Self {
472 self.reply(ScriptedReply::text(text))
473 }
474
475 #[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 #[must_use]
484 pub fn answering(self, text: impl Into<String>) -> Self {
485 self.reply_to(ModelPurpose::Answer, ScriptedReply::answer(text))
486 }
487
488 #[must_use]
490 pub fn malformed_json(self, body: impl Into<String>) -> Self {
491 self.reply(ScriptedReply::MalformedJson(body.into()))
492 }
493
494 #[must_use]
496 pub fn refusing(self, text: impl Into<String>) -> Self {
497 self.reply(ScriptedReply::refusal(text))
498 }
499
500 #[must_use]
502 pub fn timing_out(self) -> Self {
503 self.reply(ScriptedReply::Timeout)
504 }
505
506 #[must_use]
508 pub fn rate_limited(self, retry_after: Option<Duration>) -> Self {
509 self.reply(ScriptedReply::RateLimited { retry_after })
510 }
511
512 #[must_use]
514 pub fn failing(self, error: ProviderError) -> Self {
515 self.reply(ScriptedReply::Fail(error))
516 }
517
518 #[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 #[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 #[must_use]
544 pub fn build_shared(self) -> Arc<ScriptedProvider> {
545 Arc::new(self.build())
546 }
547}
548
549impl ScriptedProvider {
550 #[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 #[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 #[must_use]
588 pub fn profile_ref(&self) -> &ModelProfile {
589 &self.profile
590 }
591
592 #[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 #[must_use]
605 pub fn calls(&self) -> Vec<RecordedCall> {
606 self.lock(&self.calls).clone()
607 }
608
609 #[must_use]
611 pub fn call_count(&self) -> usize {
612 self.lock(&self.calls).len()
613 }
614
615 #[must_use]
617 pub fn nth_call(&self, index: usize) -> Option<RecordedCall> {
618 self.lock(&self.calls).get(index).cloned()
619 }
620
621 #[must_use]
623 pub fn last_call(&self) -> Option<RecordedCall> {
624 self.lock(&self.calls).last().cloned()
625 }
626
627 #[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 #[must_use]
639 pub fn remaining_steps(&self) -> usize {
640 self.lock(&self.steps).len()
641 }
642
643 #[must_use]
645 pub fn violations(&self) -> Vec<ScriptViolation> {
646 self.lock(&self.violations).clone()
647 }
648
649 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 pub fn push(&self, step: ScriptStep) {
677 self.lock(&self.steps).push_back(step);
678 }
679
680 pub fn clear_calls(&self) {
682 self.lock(&self.calls).clear();
683 self.lock(&self.violations).clear();
684 }
685
686 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 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 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 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 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
834fn 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 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#[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 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 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}