use std::sync::Mutex;
use async_trait::async_trait;
use crate::inference::adapter::InferenceAdapter;
use crate::inference::error::InferenceError;
use crate::inference::registry::ProviderCapabilities;
use crate::inference::types::{
AssistantMessage, ChatChoice, ChatRequest, ChatResponse, UsageBlock,
};
pub struct ScriptedAdapter {
name: String,
capabilities: ProviderCapabilities,
echo: bool,
queue: Mutex<Vec<Result<ChatResponse, InferenceError>>>,
}
impl ScriptedAdapter {
pub fn new(name: impl Into<String>, capabilities: &ProviderCapabilities) -> Self {
Self {
name: name.into(),
capabilities: *capabilities,
echo: false,
queue: Mutex::new(Vec::new()),
}
}
pub fn echo(name: impl Into<String>, capabilities: &ProviderCapabilities) -> Self {
Self {
name: name.into(),
capabilities: *capabilities,
echo: true,
queue: Mutex::new(Vec::new()),
}
}
pub fn with_response(self, response: ChatResponse) -> Self {
self.push(Ok(response));
self
}
pub fn with_error(self, error: InferenceError) -> Self {
self.push(Err(error));
self
}
fn push(&self, outcome: Result<ChatResponse, InferenceError>) {
let mut q = self.queue.lock().unwrap_or_else(|p| p.into_inner());
q.push(outcome);
}
fn echo_response(request: &ChatRequest) -> ChatResponse {
let content = request
.messages
.iter()
.rev()
.find(|m| m.role == "user")
.and_then(|m| m.content.clone())
.unwrap_or_default();
let prompt_tokens = request
.messages
.iter()
.filter_map(|m| m.content.as_ref())
.map(|c| c.split_whitespace().count() as u32)
.sum();
let completion_tokens = content.split_whitespace().count() as u32;
ChatResponse {
id: "scripted-echo".into(),
model: request.model.clone(),
choices: vec![ChatChoice {
message: AssistantMessage {
content: Some(content),
tool_calls: Vec::new(),
},
finish_reason: Some("stop".into()),
}],
usage: UsageBlock {
prompt_tokens,
completion_tokens,
total_tokens: prompt_tokens + completion_tokens,
..Default::default()
},
}
}
}
#[async_trait]
impl InferenceAdapter for ScriptedAdapter {
fn name(&self) -> &str {
&self.name
}
fn capabilities(&self) -> &ProviderCapabilities {
&self.capabilities
}
async fn chat(&self, request: &ChatRequest) -> Result<ChatResponse, InferenceError> {
let mut q = self.queue.lock().unwrap_or_else(|p| p.into_inner());
if q.is_empty() {
drop(q);
return if self.echo {
Ok(Self::echo_response(request))
} else {
Err(InferenceError::Provider(
"scripted adapter queue exhausted".into(),
))
};
}
q.remove(0)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::inference::registry::{ProviderId, capabilities};
use crate::inference::types::ChatMessage;
fn caps() -> &'static ProviderCapabilities {
capabilities(ProviderId::OpenRouter)
}
#[tokio::test]
async fn echo_reflects_last_user_message() {
let adapter = ScriptedAdapter::echo("scripted", caps());
let req = ChatRequest::new(
"x/y",
vec![ChatMessage::system("sys"), ChatMessage::user("hello world")],
);
let resp = adapter.chat(&req).await.expect("echo");
assert_eq!(resp.first_text().as_deref(), Some("hello world"));
assert!(resp.usage().total_tokens() > 0);
}
#[tokio::test]
async fn queue_is_fifo() {
let first = ScriptedAdapter::echo_response(&ChatRequest::new(
"m",
vec![ChatMessage::user("first")],
));
let adapter = ScriptedAdapter::echo("scripted", caps()).with_response(first);
let req = ChatRequest::new("m", vec![ChatMessage::user("second")]);
assert_eq!(
adapter.chat(&req).await.expect("q").first_text().as_deref(),
Some("first")
);
assert_eq!(
adapter
.chat(&req)
.await
.expect("echo")
.first_text()
.as_deref(),
Some("second")
);
}
#[tokio::test]
async fn strict_mode_exhaustion_errors() {
let adapter = ScriptedAdapter::new("scripted", caps());
let req = ChatRequest::new("m", vec![ChatMessage::user("x")]);
let err = adapter.chat(&req).await.expect_err("exhausted");
assert!(matches!(err, InferenceError::Provider(_)));
}
}