use async_trait::async_trait;
use crate::capabilities::{
CapabilityMismatch, CapabilityRequirements, ModelProfile, ProviderCapabilities,
};
use crate::error::ProviderError;
use crate::ids::{ModelKey, ModelRef, ProviderKey};
use crate::request::ModelRequest;
use crate::response::ModelResponse;
use crate::stream::ModelStream;
#[async_trait]
pub trait ModelProvider: Send + Sync {
fn provider_key(&self) -> ProviderKey;
fn model_key(&self) -> ModelKey;
fn capabilities(&self) -> ProviderCapabilities;
fn profile(&self) -> ModelProfile {
ModelProfile::new(self.provider_key(), self.model_key(), self.capabilities())
}
fn reference(&self) -> ModelRef {
ModelRef {
provider: self.provider_key(),
model: self.model_key(),
}
}
fn supports(&self, requirements: &CapabilityRequirements) -> Result<(), CapabilityMismatch> {
requirements.satisfied_by(&self.capabilities())
}
async fn generate(&self, request: ModelRequest) -> Result<ModelResponse, ProviderError>;
async fn stream(&self, request: ModelRequest) -> Result<ModelStream, ProviderError> {
let _ = request;
Err(ProviderError::unsupported("streaming").with_model(&self.reference()))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::capabilities::StructuredOutputCapability;
use crate::ids::RequestId;
use crate::purpose::ModelPurpose;
use crate::request::Message;
struct Fixed {
capabilities: ProviderCapabilities,
}
#[async_trait]
impl ModelProvider for Fixed {
fn provider_key(&self) -> ProviderKey {
ProviderKey::from("fixed")
}
fn model_key(&self) -> ModelKey {
ModelKey::from("fixed-1")
}
fn capabilities(&self) -> ProviderCapabilities {
self.capabilities.clone()
}
async fn generate(&self, request: ModelRequest) -> Result<ModelResponse, ProviderError> {
Ok(
ModelResponse::new(request.request_id, self.provider_key(), self.model_key())
.with_text("ok"),
)
}
}
fn provider() -> Fixed {
Fixed {
capabilities: ProviderCapabilities::minimal()
.with_structured_output(StructuredOutputCapability::NativeJsonSchema),
}
}
#[tokio::test]
async fn the_default_profile_mirrors_the_declared_capabilities() {
let provider = provider();
let profile = provider.profile();
assert_eq!(profile.provider, provider.provider_key());
assert_eq!(profile.model, provider.model_key());
assert_eq!(profile.capabilities, provider.capabilities());
assert_eq!(provider.reference().to_string(), "fixed/fixed-1");
assert!(profile.max_cost_per_million().is_none());
}
#[tokio::test]
async fn supports_answers_with_the_mismatch() {
let provider = provider();
let ok = ModelPurpose::Extract.requirements();
assert!(provider.supports(&ok).is_ok());
let needs_streaming = CapabilityRequirements::none().with_streaming();
let mismatch = provider.supports(&needs_streaming).unwrap_err();
assert_eq!(mismatch.missing.len(), 1);
assert!(!mismatch.structured_output_unmet());
}
#[tokio::test]
async fn streaming_defaults_to_an_honest_refusal() {
let provider = provider();
let request =
ModelRequest::new(ModelPurpose::Acknowledge).with_message(Message::user("hi"));
let error = provider.stream(request).await.unwrap_err();
assert!(matches!(
error.kind(),
crate::error::ProviderErrorKind::Unsupported { .. }
));
assert_eq!(error.retry_class(), crate::error::RetryClass::Fallback);
assert_eq!(error.provider().map(ProviderKey::as_str), Some("fixed"));
}
#[tokio::test]
async fn generate_echoes_the_request_id() {
let provider = provider();
let request =
ModelRequest::new(ModelPurpose::Acknowledge).with_request_id(RequestId::nil());
let response = provider.generate(request).await.unwrap();
assert_eq!(response.request_id, RequestId::nil());
}
#[tokio::test]
async fn the_trait_is_object_safe() {
let provider: Box<dyn ModelProvider> = Box::new(provider());
assert_eq!(provider.provider_key().as_str(), "fixed");
let response = provider
.generate(ModelRequest::new(ModelPurpose::Acknowledge))
.await
.unwrap();
assert_eq!(response.text(), "ok");
}
}