1use std::collections::VecDeque;
5use std::sync::Mutex;
6
7use af_llm::{ChatMessage, CompletionRequest, CompletionResponse, LlmError};
8use async_trait::async_trait;
9use tokio::sync::mpsc::UnboundedSender;
10
11use crate::ChatModel;
12
13#[derive(Default)]
17pub struct ScriptedModel {
18 responses: Mutex<VecDeque<Result<CompletionResponse, LlmError>>>,
19 requests: Mutex<Vec<CompletionRequest>>,
20}
21
22impl ScriptedModel {
23 pub fn new(responses: impl IntoIterator<Item = Result<CompletionResponse, LlmError>>) -> Self {
25 Self {
26 responses: Mutex::new(responses.into_iter().collect()),
27 requests: Mutex::new(Vec::new()),
28 }
29 }
30
31 pub fn replies(messages: impl IntoIterator<Item = ChatMessage>) -> Self {
34 Self::new(
35 messages
36 .into_iter()
37 .map(|message| Ok(Self::response(message))),
38 )
39 }
40
41 pub fn response(message: ChatMessage) -> CompletionResponse {
43 let finish_reason = if message
44 .tool_calls
45 .as_ref()
46 .is_some_and(|calls| !calls.is_empty())
47 {
48 af_llm::FinishReason::ToolCalls
49 } else {
50 af_llm::FinishReason::Stop
51 };
52 CompletionResponse {
53 id: "scripted".into(),
54 choices: vec![af_llm::Choice {
55 index: 0,
56 message,
57 finish_reason: Some(finish_reason),
58 output_blocks: Vec::new(),
59 }],
60 usage: Some(af_llm::Usage {
61 prompt_tokens: 1,
62 completion_tokens: 1,
63 total_tokens: 2,
64 }),
65 }
66 }
67
68 pub fn remaining(&self) -> usize {
70 self.responses
71 .lock()
72 .unwrap_or_else(std::sync::PoisonError::into_inner)
73 .len()
74 }
75
76 pub fn requests(&self) -> Vec<CompletionRequest> {
78 self.requests
79 .lock()
80 .unwrap_or_else(std::sync::PoisonError::into_inner)
81 .clone()
82 }
83}
84
85#[async_trait]
86impl ChatModel for ScriptedModel {
87 async fn complete_streaming(
88 &self,
89 request: &CompletionRequest,
90 delta_tx: UnboundedSender<(String, bool)>,
91 ) -> Result<CompletionResponse, LlmError> {
92 self.requests
93 .lock()
94 .unwrap_or_else(std::sync::PoisonError::into_inner)
95 .push(request.clone());
96 let next = self
97 .responses
98 .lock()
99 .unwrap_or_else(std::sync::PoisonError::into_inner)
100 .pop_front();
101 let response = match next {
102 Some(response) => response?,
103 None => return Err(LlmError::StreamProtocol("scripted model exhausted".into())),
104 };
105 if let Some(content) = response.first_content() {
106 let _ = delta_tx.send((content.into(), response.first_tool_calls().is_some()));
107 }
108 Ok(response)
109 }
110}
111
112#[cfg(test)]
113mod tests {
114 use super::*;
115
116 #[tokio::test]
117 async fn scripted_model_replays_records_and_exhausts_loudly() {
118 let model = ScriptedModel::replies([ChatMessage::assistant("hi")]);
119 assert_eq!(model.remaining(), 1);
120 let (tx, mut rx) = tokio::sync::mpsc::unbounded_channel();
121 let request = CompletionRequest::new("m", vec![ChatMessage::user("hello")]);
122 let response = model.complete_streaming(&request, tx).await.unwrap();
123 assert_eq!(response.first_content(), Some("hi"));
124 assert_eq!(rx.recv().await.unwrap(), ("hi".to_string(), false));
125 assert_eq!(model.requests().len(), 1);
126 assert_eq!(model.remaining(), 0);
127 let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
128 assert!(matches!(
129 model.complete_streaming(&request, tx).await,
130 Err(LlmError::StreamProtocol(_))
131 ));
132 let failing = ScriptedModel::new([Err(LlmError::CircuitOpen)]);
133 let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
134 assert!(matches!(
135 failing.complete_streaming(&request, tx).await,
136 Err(LlmError::CircuitOpen)
137 ));
138 }
139}