1use std::collections::VecDeque;
7use std::io::BufRead;
8use std::path::Path;
9use std::pin::Pin;
10use std::sync::{Arc, Mutex};
11
12use async_trait::async_trait;
13use futures_util::Stream;
14use lc_core::language_models::{BaseChatModel, BaseLanguageModel, LLMResult, StreamChunk};
15use lc_core::runnables::{Runnable, RunnableConfig};
16use lc_core::tools::ToolDefinition;
17use lc_schema::Message;
18
19use crate::error::TestkitError;
20use crate::recording::RecordedExchange;
21
22#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
24pub enum ReplayStrategy {
25 #[default]
32 Fifo,
33 ByToolName,
41 Exact,
52}
53
54#[derive(Clone)]
60pub struct ReplayProvider {
61 queue: Arc<Mutex<VecDeque<RecordedExchange>>>,
62 model_name: String,
63 strategy: ReplayStrategy,
64 tools: Option<Vec<ToolDefinition>>,
66}
67
68impl ReplayProvider {
69 pub fn from_file(path: impl AsRef<Path>) -> Result<Self, TestkitError> {
71 let file = std::fs::File::open(path)?;
72 let reader = std::io::BufReader::new(file);
73 let mut queue = VecDeque::new();
74 for line in reader.lines() {
75 let line = line?.trim().to_string();
76 if line.is_empty() {
77 continue;
78 }
79 let exchange: RecordedExchange = serde_json::from_str(&line).map_err(|e| {
80 TestkitError::Io(std::io::Error::new(
81 std::io::ErrorKind::InvalidData,
82 format!("invalid recording line: {e}"),
83 ))
84 })?;
85 queue.push_back(exchange);
86 }
87 Ok(Self {
88 queue: Arc::new(Mutex::new(queue)),
89 model_name: "replay".to_string(),
90 strategy: ReplayStrategy::Fifo,
91 tools: None,
92 })
93 }
94
95 pub fn from_exchanges(exchanges: Vec<RecordedExchange>) -> Self {
97 Self {
98 queue: Arc::new(Mutex::new(exchanges.into())),
99 model_name: "replay".to_string(),
100 strategy: ReplayStrategy::Fifo,
101 tools: None,
102 }
103 }
104
105 pub fn single(response: LLMResult) -> Self {
107 Self::from_exchanges(vec![RecordedExchange {
108 messages: Vec::new(),
109 response,
110 tools: None,
111 }])
112 }
113
114 pub fn with_strategy(mut self, strategy: ReplayStrategy) -> Self {
116 self.strategy = strategy;
117 self
118 }
119
120 pub fn bind_tools(&self, tools: Vec<ToolDefinition>) -> Self {
126 Self {
127 queue: self.queue.clone(),
128 model_name: self.model_name.clone(),
129 strategy: self.strategy,
130 tools: Some(tools),
131 }
132 }
133
134 pub fn len(&self) -> usize {
136 self.queue.lock().unwrap_or_else(|e| e.into_inner()).len()
137 }
138
139 pub fn is_empty(&self) -> bool {
141 self.len() == 0
142 }
143}
144
145fn exchange_matches(exchange: &RecordedExchange, tool_name: &str) -> bool {
148 if let Some(tools) = &exchange.tools {
149 if tools.iter().any(|t| t.function.name == tool_name) {
150 return true;
151 }
152 }
153 if let Some(calls) = &exchange.response.tool_calls {
154 if calls.iter().any(|c| c.name() == tool_name) {
155 return true;
156 }
157 }
158 false
159}
160
161fn messages_match(request: &[Message], recorded: &[Message]) -> bool {
168 request == recorded
169}
170
171#[async_trait]
172impl Runnable<Vec<Message>, LLMResult> for ReplayProvider {
173 type Error = TestkitError;
174
175 async fn invoke(
176 &self,
177 input: Vec<Message>,
178 config: Option<RunnableConfig>,
179 ) -> Result<LLMResult, Self::Error> {
180 self.chat(input, config).await
181 }
182}
183
184impl BaseLanguageModel<Vec<Message>, LLMResult> for ReplayProvider {
185 fn model_name(&self) -> &str {
186 &self.model_name
187 }
188
189 fn get_num_tokens(&self, text: &str) -> usize {
190 text.chars().count() / 4 + 1
192 }
193
194 fn temperature(&self) -> Option<f32> {
195 None
196 }
197
198 fn max_tokens(&self) -> Option<usize> {
199 None
200 }
201
202 fn with_temperature(self, _temp: f32) -> Self {
203 self
204 }
205
206 fn with_max_tokens(self, _max: usize) -> Self {
207 self
208 }
209}
210
211#[async_trait]
212impl BaseChatModel for ReplayProvider {
213 async fn chat(
214 &self,
215 messages: Vec<Message>,
216 _config: Option<RunnableConfig>,
217 ) -> Result<LLMResult, Self::Error> {
218 let mut queue = self.queue.lock().unwrap_or_else(|e| e.into_inner());
219 let exchange = match self.strategy {
220 ReplayStrategy::Fifo => queue.pop_front(),
221 ReplayStrategy::ByToolName => {
222 let want = self
223 .tools
224 .as_ref()
225 .and_then(|tools| tools.first().map(|t| t.function.name.clone()));
226 match want {
227 Some(name) => queue
229 .iter()
230 .position(|ex| exchange_matches(ex, &name))
231 .map(|i| queue.remove(i).expect("position 必有元素")),
232 None => queue.pop_front(),
233 }
234 }
235 ReplayStrategy::Exact => {
236 match queue
239 .iter()
240 .position(|ex| messages_match(&messages, &ex.messages))
241 {
242 Some(i) => Some(queue.remove(i).expect("position 必有元素")),
243 None => {
244 return Err(TestkitError::ReplayNoMatch { left: queue.len() });
245 }
246 }
247 }
248 };
249 let Some(exchange) = exchange else {
250 return Err(TestkitError::ReplayExhausted {
251 requested: messages.len(),
252 });
253 };
254 Ok(exchange.response)
255 }
256
257 async fn stream_chat(
258 &self,
259 messages: Vec<Message>,
260 config: Option<RunnableConfig>,
261 ) -> Result<Pin<Box<dyn Stream<Item = Result<StreamChunk, Self::Error>> + Send>>, Self::Error>
262 {
263 let response = self.chat(messages, config).await?;
264 let stream = futures_util::stream::iter(vec![Ok(StreamChunk {
265 text: response.content,
266 token_usage: response.token_usage,
267 })]);
268 Ok(Box::pin(stream))
269 }
270
271 fn bind_tools(
272 &self,
273 tools: Vec<ToolDefinition>,
274 ) -> Option<Box<dyn BaseChatModel<Error = Self::Error> + Send + Sync>> {
275 Some(Box::new(self.bind_tools(tools)))
277 }
278}
279
280#[cfg(test)]
281mod tests {
282 use super::*;
283 use lc_core::language_models::TokenUsage;
284 use lc_core::tools::{ToolCall, ToolDefinition};
285
286 fn exchange(content: &str) -> RecordedExchange {
287 RecordedExchange {
288 messages: vec![Message::system("ping")],
289 response: LLMResult {
290 content: content.to_string(),
291 model: "replay".to_string(),
292 token_usage: Some(TokenUsage {
293 prompt_tokens: 1,
294 completion_tokens: 2,
295 total_tokens: 3,
296 }),
297 ..Default::default()
298 },
299 tools: None,
300 }
301 }
302
303 fn exchange_with_tool_call(tool_name: &str, content: &str) -> RecordedExchange {
305 let mut response = exchange(content).response;
306 response.tool_calls = Some(vec![ToolCall::builder("call_1")
307 .name(tool_name)
308 .arguments("{}".to_string())
309 .build()]);
310 RecordedExchange {
311 response,
312 ..exchange(content)
313 }
314 }
315
316 fn exchange_with_bound_tool(tool_name: &str, content: &str) -> RecordedExchange {
318 RecordedExchange {
319 tools: Some(vec![ToolDefinition::new(tool_name, "a tool")]),
320 ..exchange(content)
321 }
322 }
323
324 #[tokio::test]
325 async fn single_returns_fixed_response_for_any_request() {
326 let provider = ReplayProvider::single(exchange("hello").response);
327 let result = provider
328 .chat(vec![Message::system("any")], None)
329 .await
330 .unwrap();
331 assert_eq!(result.content, "hello");
332 }
333
334 #[tokio::test]
335 async fn replay_is_fifo_ordered() {
336 let provider = ReplayProvider::from_exchanges(vec![exchange("first"), exchange("second")]);
337 let first = provider
338 .chat(vec![Message::system("a")], None)
339 .await
340 .unwrap();
341 let second = provider
342 .chat(vec![Message::system("b")], None)
343 .await
344 .unwrap();
345 assert_eq!(first.content, "first");
346 assert_eq!(second.content, "second");
347 }
348
349 #[tokio::test]
350 async fn replay_exhausted_returns_error() {
351 let provider = ReplayProvider::from_exchanges(vec![exchange("only")]);
352 provider
353 .chat(vec![Message::system("a")], None)
354 .await
355 .unwrap();
356 let err = provider
357 .chat(vec![Message::system("b")], None)
358 .await
359 .unwrap_err();
360 assert!(matches!(
361 err,
362 TestkitError::ReplayExhausted { requested: 1 }
363 ));
364 }
365
366 #[test]
367 fn bind_tools_returns_some_and_carries_tools() {
368 let provider = ReplayProvider::from_exchanges(vec![exchange("x")]);
369 let bound = provider.bind_tools(vec![ToolDefinition::new("calculator", "calc")]);
371 assert!(bound.tools.is_some());
372 assert_eq!(bound.tools.as_ref().unwrap()[0].function.name, "calculator");
373 assert_eq!(provider.len(), 1);
374 assert_eq!(bound.len(), 1);
375 let trait_bound = BaseChatModel::bind_tools(&provider, vec![ToolDefinition::new("x", "y")]);
377 assert!(trait_bound.is_some());
378 }
379
380 #[tokio::test(flavor = "multi_thread", worker_threads = 4)]
381 async fn parallel_replay_fifo_is_order_independent() {
382 let provider = ReplayProvider::from_exchanges(vec![
386 exchange("first"),
387 exchange("second"),
388 exchange("third"),
389 ]);
390 let provider = std::sync::Arc::new(provider);
391
392 let mut handles = Vec::new();
393 for _ in 0..3 {
394 let p = provider.clone();
395 handles.push(tokio::spawn(async move {
396 p.chat(vec![Message::system("parallel")], None)
397 .await
398 .expect("并发回放不应失败")
399 }));
400 }
401 let mut contents: Vec<String> = Vec::new();
402 for handle in handles {
403 contents.push(handle.await.unwrap().content);
404 }
405 contents.sort();
406 assert_eq!(
407 contents,
408 vec![
409 "first".to_string(),
410 "second".to_string(),
411 "third".to_string()
412 ]
413 );
414 assert!(provider.is_empty(), "并发回放应恰好耗尽全部录播");
415 }
416
417 #[tokio::test]
418 async fn by_tool_name_routes_to_matching_exchange() {
419 let provider = ReplayProvider::from_exchanges(vec![
421 exchange_with_tool_call("search", "search result"),
422 exchange_with_tool_call("calc", "calc result"),
423 ])
424 .with_strategy(ReplayStrategy::ByToolName);
425
426 let search =
427 BaseChatModel::bind_tools(&provider, vec![ToolDefinition::new("search", "s")]).unwrap();
428 let calc =
429 BaseChatModel::bind_tools(&provider, vec![ToolDefinition::new("calc", "c")]).unwrap();
430
431 let calc_res = calc.chat(vec![Message::system("q")], None).await.unwrap();
432 let search_res = search.chat(vec![Message::system("q")], None).await.unwrap();
433
434 assert_eq!(search_res.content, "search result");
435 assert_eq!(calc_res.content, "calc result");
436 assert!(provider.is_empty());
437 }
438
439 #[tokio::test]
440 async fn by_tool_name_matches_request_side_tools() {
441 let provider = ReplayProvider::from_exchanges(vec![
443 exchange_with_bound_tool("weather", "sunny"),
444 exchange_with_bound_tool("news", "headlines"),
445 ])
446 .with_strategy(ReplayStrategy::ByToolName);
447
448 let weather =
449 BaseChatModel::bind_tools(&provider, vec![ToolDefinition::new("weather", "w")])
450 .unwrap();
451 let res = weather
452 .chat(vec![Message::system("q")], None)
453 .await
454 .unwrap();
455 assert_eq!(res.content, "sunny");
456 }
457
458 fn exchange_with_messages(messages: Vec<Message>, content: &str) -> RecordedExchange {
460 RecordedExchange {
461 messages,
462 ..exchange(content)
463 }
464 }
465
466 #[tokio::test]
467 async fn exact_strategy_matches_by_full_signature_out_of_order() {
468 let provider = ReplayProvider::from_exchanges(vec![
470 exchange_with_messages(vec![Message::system("ping")], "pong"),
471 exchange_with_messages(vec![Message::human("hello")], "hi"),
472 ])
473 .with_strategy(ReplayStrategy::Exact);
474
475 let hello_res = provider
476 .chat(vec![Message::human("hello")], None)
477 .await
478 .unwrap();
479 let ping_res = provider
480 .chat(vec![Message::system("ping")], None)
481 .await
482 .unwrap();
483
484 assert_eq!(hello_res.content, "hi");
485 assert_eq!(ping_res.content, "pong");
486 assert!(provider.is_empty(), "两条请求应精确消耗两条录播");
487 }
488
489 #[tokio::test]
490 async fn exact_strategy_matches_full_message_sequence() {
491 let msgs = vec![
493 Message::system("You are a calculator."),
494 Message::human("2 + 2"),
495 Message::ai("I'll compute that."),
496 ];
497 let provider = ReplayProvider::from_exchanges(vec![
498 exchange_with_messages(vec![Message::human("other")], "wrong"),
499 exchange_with_messages(msgs.clone(), "42"),
500 ])
501 .with_strategy(ReplayStrategy::Exact);
502
503 let res = provider.chat(msgs, None).await.unwrap();
504 assert_eq!(res.content, "42");
505 }
506
507 #[tokio::test]
508 async fn exact_strategy_no_match_returns_explicit_error() {
509 let provider = ReplayProvider::from_exchanges(vec![exchange_with_messages(
511 vec![Message::system("ping")],
512 "pong",
513 )])
514 .with_strategy(ReplayStrategy::Exact);
515
516 let err = provider
517 .chat(vec![Message::system("different")], None)
518 .await
519 .unwrap_err();
520 assert!(
521 matches!(err, TestkitError::ReplayNoMatch { left: 1 }),
522 "无匹配应返回 ReplayNoMatch,剩余录播保留"
523 );
524 }
525
526 #[tokio::test]
527 async fn exact_strategy_distinguishes_message_type() {
528 let provider = ReplayProvider::from_exchanges(vec![
530 exchange_with_messages(vec![Message::human("q")], "human response"),
531 exchange_with_messages(vec![Message::system("q")], "system response"),
532 ])
533 .with_strategy(ReplayStrategy::Exact);
534
535 let human = provider
536 .chat(vec![Message::human("q")], None)
537 .await
538 .unwrap();
539 assert_eq!(human.content, "human response");
540
541 let system = provider
542 .chat(vec![Message::system("q")], None)
543 .await
544 .unwrap();
545 assert_eq!(system.content, "system response");
546 }
547
548 #[tokio::test]
549 async fn exact_strategy_distinguishes_tool_calls() {
550 let call = ToolCall::builder("call_1")
552 .name("weather")
553 .arguments("{}".to_string())
554 .build();
555 let provider = ReplayProvider::from_exchanges(vec![
556 exchange_with_messages(vec![Message::ai("q")], "plain"),
557 exchange_with_messages(
558 vec![Message::ai_with_tool_calls("q", vec![call.clone()])],
559 "with tool",
560 ),
561 ])
562 .with_strategy(ReplayStrategy::Exact);
563
564 let plain = provider.chat(vec![Message::ai("q")], None).await.unwrap();
565 assert_eq!(plain.content, "plain");
566
567 let with_tool = provider
568 .chat(vec![Message::ai_with_tool_calls("q", vec![call])], None)
569 .await
570 .unwrap();
571 assert_eq!(with_tool.content, "with tool");
572 }
573}