use std::collections::BTreeMap;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use crate::{
ContentPart, FeaturePolicy, ModelError, ModelErrorKind, ModelRequest, ModelWarning,
OutputFormat, SupportLevel::Emulated, SupportLevel::Native, SupportLevel::Unknown,
SupportLevel::Unsupported, ToolChoice,
};
#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum SupportLevel {
Native,
Emulated,
Unsupported,
Unknown,
}
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
pub struct FeatureSupport {
pub level: SupportLevel,
pub constraints: BTreeMap<String, Value>,
}
impl FeatureSupport {
pub fn new(level: SupportLevel) -> Self {
Self {
level,
constraints: BTreeMap::new(),
}
}
}
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
pub struct ModelCapabilities {
pub streaming: FeatureSupport,
pub tools: FeatureSupport,
pub parallel_tools: FeatureSupport,
pub structured_output: FeatureSupport,
pub reasoning: FeatureSupport,
pub image_input: FeatureSupport,
pub audio_input: FeatureSupport,
pub document_input: FeatureSupport,
pub max_context_tokens: Option<u64>,
pub extensions: BTreeMap<String, FeatureSupport>,
}
impl Default for ModelCapabilities {
fn default() -> Self {
let unknown = || FeatureSupport::new(SupportLevel::Unknown);
Self {
streaming: unknown(),
tools: unknown(),
parallel_tools: unknown(),
structured_output: unknown(),
reasoning: unknown(),
image_input: unknown(),
audio_input: unknown(),
document_input: unknown(),
max_context_tokens: None,
extensions: BTreeMap::new(),
}
}
}
impl ModelCapabilities {
pub fn validate_request(
&self,
request: &ModelRequest,
streaming: bool,
) -> Result<Vec<ModelWarning>, ModelError> {
let requires_tools = !request.tools.is_empty()
|| !request.provider_tools().is_empty()
|| matches!(
request.tool_choice,
ToolChoice::Required | ToolChoice::Named { .. }
);
let requires_structured_output = !matches!(request.output_format, OutputFormat::Text);
let has_capability_sensitive_content = request
.messages
.iter()
.flat_map(|message| &message.content)
.any(|part| {
matches!(
part,
ContentPart::Image { .. }
| ContentPart::Audio { .. }
| ContentPart::Document { .. }
| ContentPart::Reasoning(_)
)
});
if !requires_tools && !requires_structured_output && !has_capability_sensitive_content {
let mut warnings = Vec::new();
if streaming {
assess_feature(
"streaming",
&self.streaming,
request.feature_policy,
&mut warnings,
)?;
}
return Ok(warnings);
}
let mut required = Vec::new();
if streaming {
required.push(("streaming", &self.streaming));
}
if requires_tools {
required.push(("tools", &self.tools));
}
if requires_structured_output {
required.push(("structured_output", &self.structured_output));
}
for message in &request.messages {
for part in &message.content {
match part {
ContentPart::Image { .. } => required.push(("image_input", &self.image_input)),
ContentPart::Audio { .. } => required.push(("audio_input", &self.audio_input)),
ContentPart::Document { .. } => {
required.push(("document_input", &self.document_input));
}
ContentPart::Reasoning(_) => required.push(("reasoning", &self.reasoning)),
_ => {}
}
}
}
required.sort_by_key(|(name, _)| *name);
required.dedup_by_key(|(name, _)| *name);
let mut warnings = Vec::new();
for (name, support) in required {
assess_feature(name, support, request.feature_policy, &mut warnings)?;
}
Ok(warnings)
}
}
fn assess_feature(
name: &str,
support: &FeatureSupport,
policy: FeaturePolicy,
warnings: &mut Vec<ModelWarning>,
) -> Result<(), ModelError> {
match (support.level, policy) {
(Native, _) => Ok(()),
(Emulated, FeaturePolicy::AllowEmulation | FeaturePolicy::BestEffort) => {
warnings.push(compatibility_warning(
"runifold.feature_emulated",
name,
support,
));
Ok(())
}
(Unknown, FeaturePolicy::BestEffort) => {
warnings.push(compatibility_warning(
"runifold.feature_support_unknown",
name,
support,
));
Ok(())
}
(Unsupported, _) => Err(unsupported_error(name, "unsupported", support)),
(Emulated, FeaturePolicy::Strict) => Err(unsupported_error(name, "emulated", support)),
(Unknown, FeaturePolicy::Strict | FeaturePolicy::AllowEmulation) => {
Err(unsupported_error(name, "unknown", support))
}
}
}
fn compatibility_warning(code: &str, name: &str, support: &FeatureSupport) -> ModelWarning {
let mut metadata = support.constraints.clone();
metadata.insert("feature".into(), Value::String(name.into()));
ModelWarning {
code: code.into(),
message: format!("feature `{name}` is not natively supported"),
metadata,
}
}
fn unsupported_error(name: &str, level: &str, support: &FeatureSupport) -> ModelError {
let mut error = ModelError::local(
ModelErrorKind::UnsupportedFeature,
format!("required feature `{name}` has support level `{level}`"),
);
error
.metadata
.insert("feature".into(), Value::String(name.into()));
error.metadata.insert(
"constraints".into(),
Value::Object(support.constraints.clone().into_iter().collect()),
);
error
}
#[cfg(test)]
mod tests {
use crate::{
FeaturePolicy, FeatureSupport, Message, ModelCapabilities, ModelErrorKind, ModelRef,
ModelRequest, OutputFormat, SupportLevel,
};
#[test]
fn strict_policy_rejects_unknown_required_support() {
let request = ModelRequest::new(ModelRef::new("test", "model"), Message::user("hello"));
let error = ModelCapabilities::default()
.validate_request(&request, true)
.unwrap_err();
assert_eq!(error.kind, ModelErrorKind::UnsupportedFeature);
}
#[test]
fn best_effort_makes_unknown_support_visible() {
let request = ModelRequest::new(ModelRef::new("test", "model"), Message::user("hello"))
.feature_policy(FeaturePolicy::BestEffort);
let warnings = ModelCapabilities::default()
.validate_request(&request, true)
.unwrap();
assert_eq!(warnings.len(), 1);
assert_eq!(warnings[0].code, "runifold.feature_support_unknown");
}
#[test]
fn unsupported_features_fail_even_in_best_effort_mode() {
let capabilities = ModelCapabilities {
structured_output: FeatureSupport::new(SupportLevel::Unsupported),
streaming: FeatureSupport::new(SupportLevel::Native),
..ModelCapabilities::default()
};
let request = ModelRequest::new(ModelRef::new("test", "model"), Message::user("hello"))
.output_format(OutputFormat::Json)
.feature_policy(FeaturePolicy::BestEffort);
let error = capabilities.validate_request(&request, true).unwrap_err();
assert_eq!(error.kind, ModelErrorKind::UnsupportedFeature);
assert_eq!(error.metadata["feature"], "structured_output");
}
}