Skip to main content

lc_testkit/
recording.rs

1//! `RecordingProvider`:真实调用一次,把请求/响应对追加到 JSONL 录制文件。
2//!
3//! 录制是**旁路**:真实调用失败就返回失败、不写录播;真实调用成功但写盘失败
4//! 仅 `log::warn!`,不阻断真实结果。
5
6use std::io::Write;
7use std::path::Path;
8use std::pin::Pin;
9use std::sync::{Arc, Mutex};
10
11use async_trait::async_trait;
12use futures_util::{Stream, StreamExt};
13use lc_core::language_models::{BaseChatModel, BaseLanguageModel, LLMResult};
14use lc_core::runnables::{Runnable, RunnableConfig};
15use lc_providers::ProviderError;
16use lc_schema::Message;
17use serde::{Deserialize, Serialize};
18
19use crate::error::TestkitError;
20
21/// 一次录制的请求/响应对,序列化为 JSONL 一行。
22#[derive(Debug, Clone, Serialize, Deserialize)]
23pub struct RecordedExchange {
24    /// 请求(含 system/user/assistant/tool 历史)。
25    pub messages: Vec<Message>,
26    /// 完整响应。
27    pub response: LLMResult,
28}
29
30/// 追加写录制文件的共享句柄(append 模式,std Mutex 保护)。
31pub struct Recorder {
32    file: Mutex<std::fs::File>,
33}
34
35impl Recorder {
36    /// 打开/创建录制文件。打不开 → 构造期直接 `Err`(fail fast)。
37    pub fn new(path: impl AsRef<Path>) -> std::io::Result<Self> {
38        let file = std::fs::OpenOptions::new()
39            .create(true)
40            .append(true)
41            .open(path)?;
42        Ok(Self {
43            file: Mutex::new(file),
44        })
45    }
46
47    /// Best-effort 追加一条录播:失败只 `log::warn!`,绝不向上传播。
48    pub fn record(&self, exchange: &RecordedExchange) {
49        let line = match serde_json::to_string(exchange) {
50            Ok(line) => line,
51            Err(e) => {
52                log::warn!("lc-testkit: failed to serialize recording: {e}");
53                return;
54            }
55        };
56        let Ok(mut file) = self.file.lock() else {
57            log::warn!("lc-testkit: recording lock poisoned");
58            return;
59        };
60        if let Err(e) = writeln!(file, "{line}") {
61            log::warn!("lc-testkit: failed to append recording: {e}");
62        }
63    }
64}
65
66/// 把内层模型错误映射为 `TestkitError`(经 `ProviderError` 无损透传)。
67fn to_testkit<E: Into<ProviderError>>(e: E) -> TestkitError {
68    TestkitError::Inner(e.into())
69}
70
71/// 包裹任意 `BaseChatModel`:成功响应后把请求/响应对追加到 JSONL。
72pub struct RecordingProvider<M> {
73    inner: M,
74    recorder: Arc<Recorder>,
75    model_name: String,
76}
77
78impl<M> RecordingProvider<M>
79where
80    M: BaseChatModel + Send + Sync + 'static,
81    M::Error: Into<ProviderError>,
82{
83    /// 用内层模型 + 录制文件构造。文件打不开 → `Err`。
84    pub fn new(inner: M, path: impl AsRef<Path>) -> std::io::Result<Self> {
85        let model_name = format!("{}-recorded", inner.model_name());
86        let recorder = Arc::new(Recorder::new(path)?);
87        Ok(Self {
88            inner,
89            recorder,
90            model_name,
91        })
92    }
93
94    /// 访问内层模型。
95    pub fn inner(&self) -> &M {
96        &self.inner
97    }
98}
99
100#[async_trait]
101impl<M> Runnable<Vec<Message>, LLMResult> for RecordingProvider<M>
102where
103    M: BaseChatModel + Send + Sync + 'static,
104    M::Error: Into<ProviderError>,
105{
106    type Error = TestkitError;
107
108    async fn invoke(
109        &self,
110        input: Vec<Message>,
111        config: Option<RunnableConfig>,
112    ) -> Result<LLMResult, Self::Error> {
113        self.chat(input, config).await
114    }
115}
116
117impl<M> BaseLanguageModel<Vec<Message>, LLMResult> for RecordingProvider<M>
118where
119    M: BaseChatModel + Send + Sync + 'static,
120    M::Error: Into<ProviderError>,
121{
122    fn model_name(&self) -> &str {
123        &self.model_name
124    }
125
126    fn get_num_tokens(&self, text: &str) -> usize {
127        self.inner.get_num_tokens(text)
128    }
129
130    fn temperature(&self) -> Option<f32> {
131        self.inner.temperature()
132    }
133
134    fn max_tokens(&self) -> Option<usize> {
135        self.inner.max_tokens()
136    }
137
138    fn with_temperature(mut self, temp: f32) -> Self {
139        self.inner = self.inner.with_temperature(temp);
140        self
141    }
142
143    fn with_max_tokens(mut self, max: usize) -> Self {
144        self.inner = self.inner.with_max_tokens(max);
145        self
146    }
147}
148
149#[async_trait]
150impl<M> BaseChatModel for RecordingProvider<M>
151where
152    M: BaseChatModel + Send + Sync + 'static,
153    M::Error: Into<ProviderError>,
154{
155    async fn chat(
156        &self,
157        messages: Vec<Message>,
158        config: Option<RunnableConfig>,
159    ) -> Result<LLMResult, Self::Error> {
160        let response = self
161            .inner
162            .chat(messages.clone(), config)
163            .await
164            .map_err(to_testkit)?;
165        self.recorder.record(&RecordedExchange {
166            messages,
167            response: response.clone(),
168        });
169        Ok(response)
170    }
171
172    async fn stream_chat(
173        &self,
174        messages: Vec<Message>,
175        config: Option<RunnableConfig>,
176    ) -> Result<Pin<Box<dyn Stream<Item = Result<String, Self::Error>> + Send>>, Self::Error> {
177        let mut stream = self
178            .inner
179            .stream_chat(messages.clone(), config)
180            .await
181            .map_err(to_testkit)?;
182        let mut chunks = Vec::new();
183        while let Some(chunk) = stream.next().await {
184            chunks.push(chunk.map_err(to_testkit)?);
185        }
186        let full = chunks.concat();
187        let response = LLMResult {
188            content: full.clone(),
189            model: self.model_name.clone(),
190            ..Default::default()
191        };
192        self.recorder
193            .record(&RecordedExchange { messages, response });
194        let stream = futures_util::stream::iter(vec![Ok(full)]);
195        Ok(Box::pin(stream))
196    }
197}