use async_trait::async_trait;
use serde_json::Value;
use crate::inference::error::InferenceError;
use crate::inference::registry::{ProviderCapabilities, context_window};
use crate::inference::streaming::{ChatStream, buffered_stream};
use crate::inference::types::{ChatRequest, ChatResponse, ToolChoice, openai_tool_choice};
#[async_trait]
pub trait InferenceAdapter: Send + Sync {
fn name(&self) -> &str;
fn capabilities(&self) -> &ProviderCapabilities;
fn capabilities_for(&self, model: &str) -> &ProviderCapabilities {
let _ = model;
self.capabilities()
}
async fn chat(&self, request: &ChatRequest) -> Result<ChatResponse, InferenceError>;
async fn chat_stream(&self, request: &ChatRequest) -> Result<ChatStream, InferenceError> {
let response = self.chat(request).await?;
Ok(buffered_stream(response))
}
fn map_tool_choice(&self, choice: ToolChoice) -> Value {
openai_tool_choice(choice)
}
fn supports_native_tools(&self) -> bool {
self.capabilities().native_tool_calling
}
fn supports_prompt_caching(&self) -> bool {
self.capabilities().prompt_caching
}
fn supports_structured_output(&self) -> bool {
self.capabilities().structured_output
}
fn wants_detailed_usage(&self) -> bool {
self.capabilities().detailed_usage_accounting
}
fn context_window(&self, model: &str) -> usize {
context_window(model, Some(self.capabilities_for(model)))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::inference::registry::{ProviderId, capabilities};
struct SingleProvider;
#[async_trait]
impl InferenceAdapter for SingleProvider {
fn name(&self) -> &str {
"single"
}
fn capabilities(&self) -> &ProviderCapabilities {
capabilities(ProviderId::OpenRouter)
}
async fn chat(&self, _request: &ChatRequest) -> Result<ChatResponse, InferenceError> {
Err(InferenceError::Unsupported("test double".into()))
}
}
struct RoutingAdapter;
#[async_trait]
impl InferenceAdapter for RoutingAdapter {
fn name(&self) -> &str {
"routing"
}
fn capabilities(&self) -> &ProviderCapabilities {
capabilities(ProviderId::OpenRouter)
}
fn capabilities_for(&self, model: &str) -> &ProviderCapabilities {
if model.starts_with("fireworks/") {
capabilities(ProviderId::Fireworks)
} else {
capabilities(ProviderId::OpenRouter)
}
}
async fn chat(&self, _request: &ChatRequest) -> Result<ChatResponse, InferenceError> {
Err(InferenceError::Unsupported("test double".into()))
}
}
#[test]
fn adapter_capabilities_for_defaults_to_capabilities() {
let adapter = SingleProvider;
for slug in ["openai/gpt-4o-mini", "fireworks/accounts/x/models/y"] {
assert_eq!(
adapter.capabilities_for(slug).id,
adapter.capabilities().id,
"slug {slug} must fall back to the adapter's own capabilities"
);
}
}
#[test]
fn context_window_default_follows_capabilities_for() {
let adapter = RoutingAdapter;
let fireworks_slug = "fireworks/accounts/fireworks/models/llama-v3p1-70b-instruct";
assert_eq!(
adapter.capabilities_for(fireworks_slug).id,
ProviderId::Fireworks
);
assert_eq!(adapter.context_window(fireworks_slug), 128_000);
assert_eq!(
adapter
.capabilities_for("qwen/qwen-2.5-coder-32b-instruct")
.id,
ProviderId::OpenRouter
);
assert_eq!(
adapter.context_window("qwen/qwen-2.5-coder-32b-instruct"),
200_000
);
}
}