use std::{collections::BTreeMap, pin::Pin};
use async_trait::async_trait;
use futures::Stream;
use crate::{
capabilities::{
CapabilityError, MediaKind, MediaSupport, ModelCapabilities, ReasoningCapability,
},
config::{LanguageModelConfig, ResponseFormat},
error::LanguageModelError,
identifiers::ModelId,
message::Message,
response::{LanguageModelResponse, StreamDelta},
};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum ResponseFormatKind {
Text,
JsonObject,
JsonSchema,
}
impl ResponseFormatKind {
#[must_use]
pub const fn from_format(format: &ResponseFormat) -> Self {
match format {
ResponseFormat::Text => Self::Text,
ResponseFormat::JsonObject => Self::JsonObject,
ResponseFormat::JsonSchema { .. } => Self::JsonSchema,
}
}
}
pub struct GenerateRequest<'a> {
pub model: &'a ModelId,
pub messages: &'a [Message],
pub config: &'a LanguageModelConfig,
}
#[derive(Debug, Clone, PartialEq)]
pub struct ChatModelInfo {
pub id: ModelId,
pub display_name: Option<String>,
pub context_window: Option<u32>,
pub supports_streaming: bool,
pub supported_response_formats: Vec<ResponseFormatKind>,
pub media_support: BTreeMap<MediaKind, MediaSupport>,
pub reasoning: Option<ReasoningCapability>,
}
#[async_trait]
pub trait LanguageModelProvider: Send + Sync {
fn name(&self) -> &'static str;
async fn list_models(&self) -> Result<Vec<ChatModelInfo>, LanguageModelError>;
async fn generate(
&self,
request: GenerateRequest<'_>,
) -> Result<LanguageModelResponse, LanguageModelError>;
async fn generate_stream(
&self,
request: GenerateRequest<'_>,
) -> Result<
Pin<Box<dyn Stream<Item = Result<StreamDelta, LanguageModelError>> + Send>>,
LanguageModelError,
>;
fn capabilities(&self, model: &ModelId) -> Option<ModelCapabilities>;
fn validate_request(&self, request: &GenerateRequest<'_>) -> Result<(), CapabilityError> {
let has_media = request
.messages
.iter()
.any(|m| m.content.iter().any(|p| p.media_kind().is_some()));
if !has_media {
return Ok(());
}
let caps = self.capabilities(request.model);
let model_label = request.model.as_str();
for msg in request.messages {
for part in &msg.content {
let Some(kind) = part.media_kind() else {
continue;
};
let Some(source) = part.media_source() else {
continue;
};
let Some(caps_ref) = caps.as_ref() else {
return Err(CapabilityError::UnknownModel {
model: model_label.to_owned(),
kind,
});
};
caps_ref.validate(kind, source)?;
}
}
Ok(())
}
}