use std::io::Write;
use std::path::Path;
use std::pin::Pin;
use std::sync::{Arc, Mutex};
use async_trait::async_trait;
use futures_util::{Stream, StreamExt};
use lc_core::language_models::{
BaseChatModel, BaseLanguageModel, LLMResult, StreamChunk, TokenUsage,
};
use lc_core::runnables::{Runnable, RunnableConfig};
use lc_core::tools::ToolDefinition;
use lc_providers::ProviderError;
use lc_schema::Message;
use serde::{Deserialize, Serialize};
use crate::error::TestkitError;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RecordedExchange {
pub messages: Vec<Message>,
pub response: LLMResult,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub tools: Option<Vec<ToolDefinition>>,
}
pub struct Recorder {
file: Mutex<std::fs::File>,
}
impl Recorder {
pub fn new(path: impl AsRef<Path>) -> std::io::Result<Self> {
let file = std::fs::OpenOptions::new()
.create(true)
.append(true)
.open(path)?;
Ok(Self {
file: Mutex::new(file),
})
}
pub fn record(&self, exchange: &RecordedExchange) {
let line = match serde_json::to_string(exchange) {
Ok(line) => line,
Err(e) => {
log::warn!("lc-testkit: failed to serialize recording: {e}");
return;
}
};
let Ok(mut file) = self.file.lock() else {
log::warn!("lc-testkit: recording lock poisoned");
return;
};
if let Err(e) = writeln!(file, "{line}") {
log::warn!("lc-testkit: failed to append recording: {e}");
}
}
}
fn to_testkit<E: Into<ProviderError>>(e: E) -> TestkitError {
TestkitError::Inner(e.into())
}
pub struct RecordingProvider<M> {
inner: M,
recorder: Arc<Recorder>,
model_name: String,
tools: Option<Vec<ToolDefinition>>,
}
impl<M> RecordingProvider<M>
where
M: BaseChatModel + Send + Sync + 'static,
M::Error: Into<ProviderError>,
{
pub fn new(inner: M, path: impl AsRef<Path>) -> std::io::Result<Self> {
let model_name = format!("{}-recorded", inner.model_name());
let recorder = Arc::new(Recorder::new(path)?);
Ok(Self {
inner,
recorder,
model_name,
tools: None,
})
}
pub fn inner(&self) -> &M {
&self.inner
}
pub fn bind_tools(&self, tools: Vec<ToolDefinition>) -> Self
where
M: Clone,
{
Self {
inner: self.inner.clone(),
recorder: self.recorder.clone(),
model_name: self.model_name.clone(),
tools: Some(tools),
}
}
}
#[async_trait]
impl<M> Runnable<Vec<Message>, LLMResult> for RecordingProvider<M>
where
M: BaseChatModel + Clone + Send + Sync + 'static,
M::Error: Into<ProviderError>,
{
type Error = TestkitError;
async fn invoke(
&self,
input: Vec<Message>,
config: Option<RunnableConfig>,
) -> Result<LLMResult, Self::Error> {
self.chat(input, config).await
}
}
impl<M> BaseLanguageModel<Vec<Message>, LLMResult> for RecordingProvider<M>
where
M: BaseChatModel + Clone + Send + Sync + 'static,
M::Error: Into<ProviderError>,
{
fn model_name(&self) -> &str {
&self.model_name
}
fn get_num_tokens(&self, text: &str) -> usize {
self.inner.get_num_tokens(text)
}
fn temperature(&self) -> Option<f32> {
self.inner.temperature()
}
fn max_tokens(&self) -> Option<usize> {
self.inner.max_tokens()
}
fn with_temperature(mut self, temp: f32) -> Self {
self.inner = self.inner.with_temperature(temp);
self
}
fn with_max_tokens(mut self, max: usize) -> Self {
self.inner = self.inner.with_max_tokens(max);
self
}
}
#[async_trait]
impl<M> BaseChatModel for RecordingProvider<M>
where
M: BaseChatModel + Clone + Send + Sync + 'static,
M::Error: Into<ProviderError>,
{
async fn chat(
&self,
messages: Vec<Message>,
config: Option<RunnableConfig>,
) -> Result<LLMResult, Self::Error> {
let response = self
.inner
.chat(messages.clone(), config)
.await
.map_err(to_testkit)?;
self.recorder.record(&RecordedExchange {
messages,
response: response.clone(),
tools: self.tools.clone(),
});
Ok(response)
}
async fn stream_chat(
&self,
messages: Vec<Message>,
config: Option<RunnableConfig>,
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamChunk, Self::Error>> + Send>>, Self::Error>
{
let mut stream = self
.inner
.stream_chat(messages.clone(), config)
.await
.map_err(to_testkit)?;
let mut full = String::new();
let mut usage: Option<TokenUsage> = None;
while let Some(chunk) = stream.next().await {
let chunk = chunk.map_err(to_testkit)?;
full.push_str(&chunk.text);
if chunk.token_usage.is_some() {
usage = chunk.token_usage;
}
}
let response = LLMResult {
content: full.clone(),
model: self.model_name.clone(),
token_usage: usage.clone(),
..Default::default()
};
self.recorder.record(&RecordedExchange {
messages,
response,
tools: self.tools.clone(),
});
let stream = futures_util::stream::iter(vec![Ok(StreamChunk {
text: full,
token_usage: usage,
tool_calls: None,
})]);
Ok(Box::pin(stream))
}
fn bind_tools(
&self,
tools: Vec<ToolDefinition>,
) -> Option<Box<dyn BaseChatModel<Error = Self::Error> + Send + Sync>> {
Some(Box::new(self.bind_tools(tools)))
}
}