use crate::error::{CredentialError, ProviderError, RunFailureReason};
use crate::ids::{ModelId, ProviderId};
use crate::message::{ContentBlock, FinishReason, ToolResultPayload, Usage};
use crate::tool::ToolSpec;
use async_trait::async_trait;
use futures_core::Stream;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use std::pin::Pin;
use tokio_util::sync::CancellationToken;
pub type ProviderStream = Pin<Box<dyn Stream<Item = Result<ProviderEvent, ProviderError>> + Send>>;
#[async_trait]
pub trait ModelProvider: Send + Sync {
fn id(&self) -> ProviderId;
async fn capabilities(
&self,
model: &ModelId,
credentials: &dyn CredentialProvider,
cancel: CancellationToken,
) -> Result<ModelCapabilities, ProviderError>;
async fn stream(
&self,
request: ModelRequest,
credentials: &dyn CredentialProvider,
cancel: CancellationToken,
) -> Result<ProviderStream, ProviderError>;
}
#[async_trait]
pub trait TokenCounter: Send + Sync {
async fn count_input(
&self,
request: &ModelRequest,
credentials: &dyn CredentialProvider,
cancel: CancellationToken,
) -> Result<TokenMeasurement, ProviderError>;
}
#[async_trait]
pub trait CredentialProvider: Send + Sync {
async fn resolve(&self, request: CredentialRequest) -> Result<Credentials, CredentialError>;
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct CredentialRequest {
pub provider_id: ProviderId,
pub model: ModelId,
}
#[derive(Clone, Debug, PartialEq, Eq, Default, Serialize, Deserialize)]
pub struct Credentials {
pub bearer: Option<String>,
pub headers: Vec<(String, String)>,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct ModelRequest {
pub model: ModelId,
pub system_prompt: String,
pub messages: Vec<ModelMessage>,
pub tools: Vec<ToolSpec>,
pub reasoning: ReasoningLevel,
pub generation: GenerationOptions,
pub provider_options: Value,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
#[serde(tag = "role")]
pub enum ModelMessage {
User {
blocks: Vec<ContentBlock>,
},
Assistant {
blocks: Vec<ContentBlock>,
},
ToolResults {
results: Vec<ToolResultPayload>,
},
Summary {
text: String,
},
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub enum ReasoningLevel {
Off,
Low,
Medium,
High,
}
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
pub struct GenerationOptions {
pub max_output_tokens: Option<u64>,
pub temperature: Option<f64>,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct ModelCapabilities {
pub context_tokens: u32,
pub max_output_tokens: u32,
pub supports_tools: bool,
pub supports_images: bool,
pub supports_reasoning: bool,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub enum TokenMeasurementSource {
ExactProvider,
ExactLocal,
UsageCalibrated,
Heuristic,
}
impl std::fmt::Display for TokenMeasurementSource {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{self:?}")
}
}
impl std::error::Error for TokenMeasurementSource {}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct TokenMeasurement {
pub input_tokens: u64,
pub source: TokenMeasurementSource,
pub safety_margin_tokens: u64,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub enum ProviderEvent {
ResponseStarted,
TextDelta {
block: u32,
text: String,
},
ReasoningDelta {
block: u32,
text: String,
},
ToolCallStarted {
block: u32,
id: String,
name: String,
},
ToolCallArgumentsDelta {
block: u32,
json: String,
},
UsageUpdated(Usage),
ResponseCompleted {
finish_reason: FinishReason,
},
}
pub fn validate_request_capabilities(
provider_id: &ProviderId,
request: &ModelRequest,
capabilities: &ModelCapabilities,
) -> Result<(), RunFailureReason> {
let mismatch = |requirement: &str| RunFailureReason::ModelCapabilityMismatch {
requirement: requirement.to_string(),
provider: provider_id.clone(),
model: request.model.clone(),
};
if !request.tools.is_empty() && !capabilities.supports_tools {
return Err(mismatch("tools"));
}
let has_image = request.messages.iter().any(|m| match m {
ModelMessage::User { blocks } | ModelMessage::Assistant { blocks } => blocks
.iter()
.any(|b| matches!(b, ContentBlock::Image { .. })),
ModelMessage::ToolResults { .. } | ModelMessage::Summary { .. } => false,
});
if has_image && !capabilities.supports_images {
return Err(mismatch("images"));
}
if request.reasoning != ReasoningLevel::Off && !capabilities.supports_reasoning {
return Err(mismatch("reasoning"));
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::message::BinaryRef;
use crate::tool::Concurrency;
fn caps(
supports_tools: bool,
supports_images: bool,
supports_reasoning: bool,
) -> ModelCapabilities {
ModelCapabilities {
context_tokens: 200_000,
max_output_tokens: 8_192,
supports_tools,
supports_images,
supports_reasoning,
}
}
fn request_with_tools(n: usize, reasoning: ReasoningLevel, image: bool) -> ModelRequest {
let mut blocks = vec![ContentBlock::Text { text: "q".into() }];
if image {
blocks.push(ContentBlock::Image {
media_type: "image/png".into(),
data: BinaryRef {
sha256: "h".into(),
byte_len: 1,
},
});
}
ModelRequest {
model: ModelId::from("m"),
system_prompt: String::new(),
messages: vec![ModelMessage::User { blocks }],
tools: (0..n)
.map(|i| ToolSpec {
name: format!("t{i}"),
description: String::new(),
parameters_schema: serde_json::json!({"type": "object"}),
concurrency: Concurrency::Sequential,
})
.collect(),
reasoning,
generation: GenerationOptions::default(),
provider_options: serde_json::Value::Null,
}
}
#[test]
fn rejects_tools_when_unsupported() {
let err = validate_request_capabilities(
&ProviderId::from(ProviderId::ANTHROPIC),
&request_with_tools(1, ReasoningLevel::Off, false),
&caps(false, true, true),
)
.unwrap_err();
assert!(
matches!(err, RunFailureReason::ModelCapabilityMismatch { requirement, .. } if requirement == "tools")
);
}
#[test]
fn rejects_images_when_unsupported() {
let err = validate_request_capabilities(
&ProviderId::from(ProviderId::ANTHROPIC),
&request_with_tools(0, ReasoningLevel::Off, true),
&caps(true, false, true),
)
.unwrap_err();
assert!(
matches!(err, RunFailureReason::ModelCapabilityMismatch { requirement, .. } if requirement == "images")
);
}
#[test]
fn rejects_reasoning_when_unsupported() {
let err = validate_request_capabilities(
&ProviderId::from(ProviderId::ANTHROPIC),
&request_with_tools(0, ReasoningLevel::High, false),
&caps(true, true, false),
)
.unwrap_err();
assert!(
matches!(err, RunFailureReason::ModelCapabilityMismatch { requirement, .. } if requirement == "reasoning")
);
}
#[test]
fn accepts_matching_request() {
assert!(validate_request_capabilities(
&ProviderId::from(ProviderId::ANTHROPIC),
&request_with_tools(2, ReasoningLevel::Medium, true),
&caps(true, true, true),
)
.is_ok());
}
#[test]
fn request_roundtrips_through_serde() {
let req = request_with_tools(1, ReasoningLevel::Off, false);
let json = serde_json::to_string(&req).unwrap();
let back: ModelRequest = serde_json::from_str(&json).unwrap();
assert_eq!(back, req);
}
#[test]
fn summary_message_roundtrips_through_serde() {
let message = ModelMessage::Summary {
text: "摘要".into(),
};
let json = serde_json::to_string(&message).unwrap();
assert!(json.contains(r#""role""#));
let back: ModelMessage = serde_json::from_str(&json).unwrap();
assert_eq!(back, message);
}
}