Skip to main content

lc_testkit/
replay.rs

1//! `ReplayProvider`:从录制文件按 FIFO 顺序回放,零网络。
2//!
3//! 回放**不做消息匹配**(LLMChain 渲染出的 prompt 逐次可变),只按顺序弹出录播。
4//! 队列耗尽返回 [`TestkitError::ReplayExhausted`]。
5
6use 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
21/// 从录制文件回放的零网络 `BaseChatModel`。
22pub struct ReplayProvider {
23    queue: Mutex<VecDeque<RecordedExchange>>,
24    model_name: String,
25}
26
27impl ReplayProvider {
28    /// 读 JSONL 录制文件(缺文件/坏行 → `Err`)。
29    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    /// 内存构造(手写录播 = MockProvider 的等价物)。
53    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    /// 单一固定响应:任意请求都返回同一 `response`(最简 mock)。
61    pub fn single(response: LLMResult) -> Self {
62        Self::from_exchanges(vec![RecordedExchange {
63            messages: Vec::new(),
64            response,
65        }])
66    }
67
68    /// 剩余录播条数。
69    pub fn len(&self) -> usize {
70        self.queue.lock().unwrap_or_else(|e| e.into_inner()).len()
71    }
72
73    /// 是否已无剩余录播。
74    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        // 估算:约 4 字符/ token。
99        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}