use std::{collections::BTreeMap, ops::RangeInclusive};
use enumset::{EnumSet, EnumSetType};
use serde::{Deserialize, Serialize};
use thiserror::Error;
use crate::{
config::{ReasoningConfig, ReasoningEffort},
media::{MediaSource, MediaType, SourceKind},
};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Ord, PartialOrd, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum MediaKind {
Image,
Document,
Audio,
Video,
}
impl MediaKind {
#[must_use]
pub fn from_media_type(media_type: &MediaType) -> Option<Self> {
match media_type.top_level() {
"image" => Some(Self::Image),
"audio" => Some(Self::Audio),
"video" => Some(Self::Video),
"application" | "text" => Some(Self::Document),
_ => None,
}
}
#[must_use]
pub const fn label(self) -> &'static str {
match self {
Self::Image => "image",
Self::Document => "document",
Self::Audio => "audio",
Self::Video => "video",
}
}
}
impl std::fmt::Display for MediaKind {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.label())
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct MediaSupport {
pub sources: EnumSet<SourceKind>,
pub formats: &'static [&'static str],
pub max_bytes: Option<u64>,
pub max_count_per_message: Option<u8>,
}
impl MediaSupport {
pub fn validate(
&self,
model: &str,
kind: MediaKind,
source: &MediaSource,
) -> Result<(), CapabilityError> {
let source_kind = source.kind();
if !self.sources.contains(source_kind) {
return Err(CapabilityError::SourceKindUnsupported {
model: model.to_owned(),
kind,
attempted: source_kind,
accepted: self.sources,
});
}
if let MediaSource::InlineBytes { mime, data } = source {
if !self.formats.is_empty() {
let subtype = mime.subtype();
if !self.formats.contains(&subtype) {
return Err(CapabilityError::FormatUnsupported {
model: model.to_owned(),
kind,
format: subtype.to_owned(),
accepted: self.formats.to_vec(),
});
}
}
if let Some(max) = self.max_bytes {
let bytes = u64::try_from(data.len()).unwrap_or(u64::MAX);
if bytes > max {
return Err(CapabilityError::SizeExceeded {
model: model.to_owned(),
kind,
bytes,
max,
});
}
}
}
Ok(())
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct ModelCapabilities {
pub model_id: String,
pub media_support: BTreeMap<MediaKind, MediaSupport>,
pub reasoning: Option<ReasoningCapability>,
pub latency_optimized_supported: bool,
pub extended_cache_ttl_supported: bool,
}
impl ModelCapabilities {
pub fn validate(&self, kind: MediaKind, source: &MediaSource) -> Result<(), CapabilityError> {
let support =
self.media_support
.get(&kind)
.ok_or_else(|| CapabilityError::ModalityUnsupported {
model: self.model_id.clone(),
kind,
})?;
support.validate(&self.model_id, kind, source)
}
}
#[derive(EnumSetType, Debug, Hash)]
pub enum ReasoningMode {
Adaptive,
Manual,
}
impl ReasoningMode {
#[must_use]
pub const fn label(self) -> &'static str {
match self {
Self::Adaptive => "adaptive",
Self::Manual => "manual",
}
}
}
impl std::fmt::Display for ReasoningMode {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.label())
}
}
#[derive(Debug, Clone, PartialEq, Default)]
pub struct ReasoningParamConflicts {
pub temperature_forbidden: bool,
pub top_k_forbidden: bool,
pub top_p_allowed_range: Option<RangeInclusive<f64>>,
}
#[derive(Debug, Clone, PartialEq)]
pub struct ReasoningCapability {
pub supported_modes: EnumSet<ReasoningMode>,
pub supported_efforts: EnumSet<ReasoningEffort>,
pub manual_budget_range: Option<RangeInclusive<u32>>,
pub conflicts: ReasoningParamConflicts,
pub sampling_params_removed: bool,
}
impl ReasoningCapability {
pub fn validate(
&self,
model_id: &str,
config: &ReasoningConfig,
) -> Result<(), ReasoningValidationError> {
match config {
ReasoningConfig::Off => Ok(()),
ReasoningConfig::Adaptive { effort } => {
if !self.supported_modes.contains(ReasoningMode::Adaptive) {
return Err(ReasoningValidationError::ModeUnsupported {
model: model_id.to_owned(),
requested: ReasoningMode::Adaptive,
supported: self.supported_modes,
});
}
if !self.supported_efforts.contains(*effort) {
return Err(ReasoningValidationError::EffortUnsupported {
model: model_id.to_owned(),
requested: *effort,
supported: self.supported_efforts,
});
}
Ok(())
}
ReasoningConfig::Manual { budget_tokens } => {
if !self.supported_modes.contains(ReasoningMode::Manual) {
return Err(ReasoningValidationError::ModeUnsupported {
model: model_id.to_owned(),
requested: ReasoningMode::Manual,
supported: self.supported_modes,
});
}
let range = self.manual_budget_range.as_ref().ok_or_else(|| {
ReasoningValidationError::ModeUnsupported {
model: model_id.to_owned(),
requested: ReasoningMode::Manual,
supported: self.supported_modes,
}
})?;
if !range.contains(budget_tokens) {
return Err(ReasoningValidationError::BudgetOutOfRange {
model: model_id.to_owned(),
requested: *budget_tokens,
min: *range.start(),
max: *range.end(),
});
}
Ok(())
}
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Error)]
pub enum ReasoningValidationError {
#[error("model `{model}` does not support reasoning")]
Unsupported { model: String },
#[error(
"model `{model}` does not support {requested} reasoning mode (supported: {supported:?})"
)]
ModeUnsupported {
model: String,
requested: ReasoningMode,
supported: EnumSet<ReasoningMode>,
},
#[error(
"model `{model}` does not support reasoning effort `{requested}` (supported: {supported:?})"
)]
EffortUnsupported {
model: String,
requested: ReasoningEffort,
supported: EnumSet<ReasoningEffort>,
},
#[error(
"model `{model}` rejects manual budget_tokens={requested} \
(supported range: {min}..={max})"
)]
BudgetOutOfRange {
model: String,
requested: u32,
min: u32,
max: u32,
},
}
#[derive(Debug, Clone, PartialEq, Eq, Error)]
pub enum CapabilityError {
#[error("model `{model}` does not accept {kind} content")]
ModalityUnsupported { model: String, kind: MediaKind },
#[error(
"model `{model}` accepts {kind} content but not via {attempted:?} (accepted: {accepted:?})"
)]
SourceKindUnsupported {
model: String,
kind: MediaKind,
attempted: SourceKind,
accepted: EnumSet<SourceKind>,
},
#[error("model `{model}` rejects {kind} format `{format}` (accepted: {accepted:?})")]
FormatUnsupported {
model: String,
kind: MediaKind,
format: String,
accepted: Vec<&'static str>,
},
#[error("{kind} payload {bytes} exceeds model `{model}` cap {max}")]
SizeExceeded {
model: String,
kind: MediaKind,
bytes: u64,
max: u64,
},
#[error("model `{model}` has no capability metadata; cannot validate {kind} content")]
UnknownModel { model: String, kind: MediaKind },
}
pub trait AcceptsImageUrl {}
pub trait AcceptsImageBytes {}
pub trait AcceptsImageS3 {}
pub trait AcceptsAudioBytes {}
pub trait AcceptsDocumentBytes {}
pub trait AcceptsVideoBytes {}
pub trait AcceptsVideoS3 {}
#[cfg(test)]
mod tests {
use enumset::enum_set;
use super::*;
use crate::media::{HttpsUrl, MediaType};
fn anthropic_image_support() -> MediaSupport {
MediaSupport {
sources: enum_set!(
SourceKind::Url | SourceKind::InlineBytes | SourceKind::ProviderFile
),
formats: &["png", "jpeg", "gif", "webp"],
max_bytes: Some(5 * 1024 * 1024),
max_count_per_message: None,
}
}
fn bedrock_image_support() -> MediaSupport {
MediaSupport {
sources: enum_set!(SourceKind::InlineBytes | SourceKind::S3),
formats: &["png", "jpeg", "gif", "webp"],
max_bytes: Some(3_932_160), max_count_per_message: None,
}
}
#[test]
fn media_kind_from_media_type_buckets_correctly() {
let png = MediaType::parse("image/png").unwrap();
assert_eq!(MediaKind::from_media_type(&png), Some(MediaKind::Image));
let mp3 = MediaType::parse("audio/mpeg").unwrap();
assert_eq!(MediaKind::from_media_type(&mp3), Some(MediaKind::Audio));
let mp4 = MediaType::parse("video/mp4").unwrap();
assert_eq!(MediaKind::from_media_type(&mp4), Some(MediaKind::Video));
let pdf = MediaType::parse("application/pdf").unwrap();
assert_eq!(MediaKind::from_media_type(&pdf), Some(MediaKind::Document));
let txt = MediaType::parse("text/plain").unwrap();
assert_eq!(MediaKind::from_media_type(&txt), Some(MediaKind::Document));
let multipart = MediaType::parse("multipart/form-data").unwrap();
assert_eq!(MediaKind::from_media_type(&multipart), None);
}
#[test]
fn support_accepts_url_when_listed() {
let support = anthropic_image_support();
let src = MediaSource::Url {
url: HttpsUrl::parse("https://x/y.png").unwrap(),
};
assert!(support.validate("claude", MediaKind::Image, &src).is_ok());
}
#[test]
fn support_rejects_url_when_not_listed() {
let support = bedrock_image_support();
let src = MediaSource::Url {
url: HttpsUrl::parse("https://x/y.png").unwrap(),
};
let err = support
.validate("bedrock-claude", MediaKind::Image, &src)
.unwrap_err();
assert!(matches!(
err,
CapabilityError::SourceKindUnsupported {
attempted: SourceKind::Url,
..
}
));
}
#[test]
fn support_rejects_inline_bytes_with_wrong_subtype() {
let support = anthropic_image_support();
let src = MediaSource::InlineBytes {
mime: MediaType::parse("image/bmp").unwrap(),
data: vec![0, 1, 2, 3],
};
let err = support
.validate("claude", MediaKind::Image, &src)
.unwrap_err();
match err {
CapabilityError::FormatUnsupported {
format, accepted, ..
} => {
assert_eq!(format, "bmp");
assert_eq!(accepted, vec!["png", "jpeg", "gif", "webp"]);
}
other => panic!("expected FormatUnsupported, got {other:?}"),
}
}
#[test]
fn support_rejects_oversize_inline_bytes() {
let support = MediaSupport {
sources: enum_set!(SourceKind::InlineBytes),
formats: &["png"],
max_bytes: Some(8),
max_count_per_message: None,
};
let src = MediaSource::InlineBytes {
mime: MediaType::parse("image/png").unwrap(),
data: vec![0; 16],
};
let err = support.validate("m", MediaKind::Image, &src).unwrap_err();
match err {
CapabilityError::SizeExceeded { bytes, max, .. } => {
assert_eq!(bytes, 16);
assert_eq!(max, 8);
}
other => panic!("expected SizeExceeded, got {other:?}"),
}
}
#[test]
fn support_skips_format_check_for_non_inline_sources() {
let support = bedrock_image_support();
let src = MediaSource::S3 {
uri: crate::media::S3Uri::parse("s3://bucket/key").unwrap(),
bucket_owner: None,
};
assert!(
support
.validate("bedrock-claude", MediaKind::Image, &src)
.is_ok()
);
}
#[test]
fn capabilities_rejects_unknown_modality() {
let caps = ModelCapabilities {
model_id: "claude".to_owned(),
media_support: BTreeMap::from([(MediaKind::Image, anthropic_image_support())]),
reasoning: None,
latency_optimized_supported: false,
extended_cache_ttl_supported: false,
};
let src = MediaSource::InlineBytes {
mime: MediaType::parse("audio/mpeg").unwrap(),
data: vec![0],
};
let err = caps.validate(MediaKind::Audio, &src).unwrap_err();
assert!(matches!(
err,
CapabilityError::ModalityUnsupported {
kind: MediaKind::Audio,
..
}
));
}
#[test]
fn capabilities_routes_through_to_support_validate() {
let caps = ModelCapabilities {
model_id: "claude".to_owned(),
media_support: BTreeMap::from([(MediaKind::Image, anthropic_image_support())]),
reasoning: None,
latency_optimized_supported: false,
extended_cache_ttl_supported: false,
};
let src = MediaSource::Url {
url: HttpsUrl::parse("https://x/y.png").unwrap(),
};
assert!(caps.validate(MediaKind::Image, &src).is_ok());
}
}