use std::sync::Mutex;
use async_trait::async_trait;
use trusty_common::inference::{
ChatRequest, ChatResponse, InferenceAdapter, InferenceError, ProviderCapabilities,
};
pub struct RecordingAdapter {
capabilities: &'static ProviderCapabilities,
seen: Mutex<Vec<ChatRequest>>,
response: ChatResponse,
}
impl RecordingAdapter {
pub fn new(capabilities: &'static ProviderCapabilities, response: ChatResponse) -> Self {
Self {
capabilities,
seen: Mutex::new(Vec::new()),
response,
}
}
pub fn only_request(&self) -> ChatRequest {
let seen = self.seen.lock().unwrap_or_else(|p| p.into_inner());
assert_eq!(
seen.len(),
1,
"expected exactly one call, got {}",
seen.len()
);
seen[0].clone()
}
pub fn only_system_turn(&self) -> String {
self.only_request().messages[0]
.content
.clone()
.unwrap_or_default()
}
}
#[async_trait]
impl InferenceAdapter for RecordingAdapter {
fn name(&self) -> &str {
"recording"
}
fn capabilities(&self) -> &ProviderCapabilities {
self.capabilities
}
async fn chat(&self, request: &ChatRequest) -> Result<ChatResponse, InferenceError> {
self.seen
.lock()
.unwrap_or_else(|p| p.into_inner())
.push(request.clone());
Ok(self.response.clone())
}
}