1use 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#[derive(Debug, Clone, Serialize, Deserialize)]
23pub struct RecordedExchange {
24 pub messages: Vec<Message>,
26 pub response: LLMResult,
28}
29
30pub struct Recorder {
32 file: Mutex<std::fs::File>,
33}
34
35impl Recorder {
36 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 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
66fn to_testkit<E: Into<ProviderError>>(e: E) -> TestkitError {
68 TestkitError::Inner(e.into())
69}
70
71pub 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 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 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}