1use std::collections::VecDeque;
7use std::io::BufRead;
8use std::path::Path;
9use std::pin::Pin;
10use std::sync::Mutex;
11
12use async_trait::async_trait;
13use futures_util::Stream;
14use lc_core::language_models::{BaseChatModel, BaseLanguageModel, LLMResult};
15use lc_core::runnables::{Runnable, RunnableConfig};
16use lc_schema::Message;
17
18use crate::error::TestkitError;
19use crate::recording::RecordedExchange;
20
21pub struct ReplayProvider {
23 queue: Mutex<VecDeque<RecordedExchange>>,
24 model_name: String,
25}
26
27impl ReplayProvider {
28 pub fn from_file(path: impl AsRef<Path>) -> Result<Self, TestkitError> {
30 let file = std::fs::File::open(path)?;
31 let reader = std::io::BufReader::new(file);
32 let mut queue = VecDeque::new();
33 for line in reader.lines() {
34 let line = line?.trim().to_string();
35 if line.is_empty() {
36 continue;
37 }
38 let exchange: RecordedExchange = serde_json::from_str(&line).map_err(|e| {
39 TestkitError::Io(std::io::Error::new(
40 std::io::ErrorKind::InvalidData,
41 format!("invalid recording line: {e}"),
42 ))
43 })?;
44 queue.push_back(exchange);
45 }
46 Ok(Self {
47 queue: Mutex::new(queue),
48 model_name: "replay".to_string(),
49 })
50 }
51
52 pub fn from_exchanges(exchanges: Vec<RecordedExchange>) -> Self {
54 Self {
55 queue: Mutex::new(exchanges.into()),
56 model_name: "replay".to_string(),
57 }
58 }
59
60 pub fn single(response: LLMResult) -> Self {
62 Self::from_exchanges(vec![RecordedExchange {
63 messages: Vec::new(),
64 response,
65 }])
66 }
67
68 pub fn len(&self) -> usize {
70 self.queue.lock().unwrap_or_else(|e| e.into_inner()).len()
71 }
72
73 pub fn is_empty(&self) -> bool {
75 self.len() == 0
76 }
77}
78
79#[async_trait]
80impl Runnable<Vec<Message>, LLMResult> for ReplayProvider {
81 type Error = TestkitError;
82
83 async fn invoke(
84 &self,
85 input: Vec<Message>,
86 config: Option<RunnableConfig>,
87 ) -> Result<LLMResult, Self::Error> {
88 self.chat(input, config).await
89 }
90}
91
92impl BaseLanguageModel<Vec<Message>, LLMResult> for ReplayProvider {
93 fn model_name(&self) -> &str {
94 &self.model_name
95 }
96
97 fn get_num_tokens(&self, text: &str) -> usize {
98 text.chars().count() / 4 + 1
100 }
101
102 fn temperature(&self) -> Option<f32> {
103 None
104 }
105
106 fn max_tokens(&self) -> Option<usize> {
107 None
108 }
109
110 fn with_temperature(self, _temp: f32) -> Self {
111 self
112 }
113
114 fn with_max_tokens(self, _max: usize) -> Self {
115 self
116 }
117}
118
119#[async_trait]
120impl BaseChatModel for ReplayProvider {
121 async fn chat(
122 &self,
123 messages: Vec<Message>,
124 _config: Option<RunnableConfig>,
125 ) -> Result<LLMResult, Self::Error> {
126 let mut queue = self.queue.lock().unwrap_or_else(|e| e.into_inner());
127 let Some(exchange) = queue.pop_front() else {
128 return Err(TestkitError::ReplayExhausted {
129 requested: messages.len(),
130 });
131 };
132 Ok(exchange.response)
133 }
134
135 async fn stream_chat(
136 &self,
137 messages: Vec<Message>,
138 config: Option<RunnableConfig>,
139 ) -> Result<Pin<Box<dyn Stream<Item = Result<String, Self::Error>> + Send>>, Self::Error> {
140 let response = self.chat(messages, config).await?;
141 let stream = futures_util::stream::iter(vec![Ok(response.content)]);
142 Ok(Box::pin(stream))
143 }
144}
145
146#[cfg(test)]
147mod tests {
148 use super::*;
149 use lc_core::language_models::TokenUsage;
150
151 fn exchange(content: &str) -> RecordedExchange {
152 RecordedExchange {
153 messages: vec![Message::system("ping")],
154 response: LLMResult {
155 content: content.to_string(),
156 model: "replay".to_string(),
157 token_usage: Some(TokenUsage {
158 prompt_tokens: 1,
159 completion_tokens: 2,
160 total_tokens: 3,
161 }),
162 ..Default::default()
163 },
164 }
165 }
166
167 #[tokio::test]
168 async fn single_returns_fixed_response_for_any_request() {
169 let provider = ReplayProvider::single(exchange("hello").response);
170 let result = provider
171 .chat(vec![Message::system("any")], None)
172 .await
173 .unwrap();
174 assert_eq!(result.content, "hello");
175 }
176
177 #[tokio::test]
178 async fn replay_is_fifo_ordered() {
179 let provider = ReplayProvider::from_exchanges(vec![exchange("first"), exchange("second")]);
180 let first = provider
181 .chat(vec![Message::system("a")], None)
182 .await
183 .unwrap();
184 let second = provider
185 .chat(vec![Message::system("b")], None)
186 .await
187 .unwrap();
188 assert_eq!(first.content, "first");
189 assert_eq!(second.content, "second");
190 }
191
192 #[tokio::test]
193 async fn replay_exhausted_returns_error() {
194 let provider = ReplayProvider::from_exchanges(vec![exchange("only")]);
195 provider
196 .chat(vec![Message::system("a")], None)
197 .await
198 .unwrap();
199 let err = provider
200 .chat(vec![Message::system("b")], None)
201 .await
202 .unwrap_err();
203 assert!(matches!(
204 err,
205 TestkitError::ReplayExhausted { requested: 1 }
206 ));
207 }
208}