1use std::collections::BTreeMap;
2
3use serde::{Deserialize, Serialize};
4use serde_json::Value;
5
6use crate::{
7 ContentPart, FeaturePolicy, ModelError, ModelErrorKind, ModelRequest, ModelWarning,
8 OutputFormat, SupportLevel::Emulated, SupportLevel::Native, SupportLevel::Unknown,
9 SupportLevel::Unsupported, ToolChoice,
10};
11
12#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)]
14#[serde(rename_all = "snake_case")]
15#[non_exhaustive]
16pub enum SupportLevel {
17 Native,
19 Emulated,
21 Unsupported,
23 Unknown,
25}
26
27#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
29pub struct FeatureSupport {
30 pub level: SupportLevel,
32 pub constraints: BTreeMap<String, Value>,
34}
35
36impl FeatureSupport {
37 pub fn new(level: SupportLevel) -> Self {
39 Self {
40 level,
41 constraints: BTreeMap::new(),
42 }
43 }
44}
45
46#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
48pub struct ModelCapabilities {
49 pub streaming: FeatureSupport,
51 pub tools: FeatureSupport,
53 pub parallel_tools: FeatureSupport,
55 pub structured_output: FeatureSupport,
57 pub reasoning: FeatureSupport,
59 pub image_input: FeatureSupport,
61 pub audio_input: FeatureSupport,
63 pub document_input: FeatureSupport,
65 pub max_context_tokens: Option<u64>,
67 pub extensions: BTreeMap<String, FeatureSupport>,
69}
70
71impl Default for ModelCapabilities {
72 fn default() -> Self {
73 let unknown = || FeatureSupport::new(SupportLevel::Unknown);
74 Self {
75 streaming: unknown(),
76 tools: unknown(),
77 parallel_tools: unknown(),
78 structured_output: unknown(),
79 reasoning: unknown(),
80 image_input: unknown(),
81 audio_input: unknown(),
82 document_input: unknown(),
83 max_context_tokens: None,
84 extensions: BTreeMap::new(),
85 }
86 }
87}
88
89impl ModelCapabilities {
90 pub fn validate_request(
103 &self,
104 request: &ModelRequest,
105 streaming: bool,
106 ) -> Result<Vec<ModelWarning>, ModelError> {
107 let requires_tools = !request.tools.is_empty()
108 || !request.provider_tools().is_empty()
109 || matches!(
110 request.tool_choice,
111 ToolChoice::Required | ToolChoice::Named { .. }
112 );
113 let requires_structured_output = !matches!(request.output_format, OutputFormat::Text);
114 let has_capability_sensitive_content = request
115 .messages
116 .iter()
117 .flat_map(|message| &message.content)
118 .any(|part| {
119 matches!(
120 part,
121 ContentPart::Image { .. }
122 | ContentPart::Audio { .. }
123 | ContentPart::Document { .. }
124 | ContentPart::Reasoning(_)
125 )
126 });
127 if !requires_tools && !requires_structured_output && !has_capability_sensitive_content {
128 let mut warnings = Vec::new();
129 if streaming {
130 assess_feature(
131 "streaming",
132 &self.streaming,
133 request.feature_policy,
134 &mut warnings,
135 )?;
136 }
137 return Ok(warnings);
138 }
139
140 let mut required = Vec::new();
141 if streaming {
142 required.push(("streaming", &self.streaming));
143 }
144 if requires_tools {
145 required.push(("tools", &self.tools));
146 }
147 if requires_structured_output {
148 required.push(("structured_output", &self.structured_output));
149 }
150 for message in &request.messages {
151 for part in &message.content {
152 match part {
153 ContentPart::Image { .. } => required.push(("image_input", &self.image_input)),
154 ContentPart::Audio { .. } => required.push(("audio_input", &self.audio_input)),
155 ContentPart::Document { .. } => {
156 required.push(("document_input", &self.document_input));
157 }
158 ContentPart::Reasoning(_) => required.push(("reasoning", &self.reasoning)),
159 _ => {}
160 }
161 }
162 }
163
164 required.sort_by_key(|(name, _)| *name);
165 required.dedup_by_key(|(name, _)| *name);
166 let mut warnings = Vec::new();
167 for (name, support) in required {
168 assess_feature(name, support, request.feature_policy, &mut warnings)?;
169 }
170 Ok(warnings)
171 }
172}
173
174fn assess_feature(
175 name: &str,
176 support: &FeatureSupport,
177 policy: FeaturePolicy,
178 warnings: &mut Vec<ModelWarning>,
179) -> Result<(), ModelError> {
180 match (support.level, policy) {
181 (Native, _) => Ok(()),
182 (Emulated, FeaturePolicy::AllowEmulation | FeaturePolicy::BestEffort) => {
183 warnings.push(compatibility_warning(
184 "runifold.feature_emulated",
185 name,
186 support,
187 ));
188 Ok(())
189 }
190 (Unknown, FeaturePolicy::BestEffort) => {
191 warnings.push(compatibility_warning(
192 "runifold.feature_support_unknown",
193 name,
194 support,
195 ));
196 Ok(())
197 }
198 (Unsupported, _) => Err(unsupported_error(name, "unsupported", support)),
199 (Emulated, FeaturePolicy::Strict) => Err(unsupported_error(name, "emulated", support)),
200 (Unknown, FeaturePolicy::Strict | FeaturePolicy::AllowEmulation) => {
201 Err(unsupported_error(name, "unknown", support))
202 }
203 }
204}
205
206fn compatibility_warning(code: &str, name: &str, support: &FeatureSupport) -> ModelWarning {
207 let mut metadata = support.constraints.clone();
208 metadata.insert("feature".into(), Value::String(name.into()));
209 ModelWarning {
210 code: code.into(),
211 message: format!("feature `{name}` is not natively supported"),
212 metadata,
213 }
214}
215
216fn unsupported_error(name: &str, level: &str, support: &FeatureSupport) -> ModelError {
217 let mut error = ModelError::local(
218 ModelErrorKind::UnsupportedFeature,
219 format!("required feature `{name}` has support level `{level}`"),
220 );
221 error
222 .metadata
223 .insert("feature".into(), Value::String(name.into()));
224 error.metadata.insert(
225 "constraints".into(),
226 Value::Object(support.constraints.clone().into_iter().collect()),
227 );
228 error
229}
230
231#[cfg(test)]
232mod tests {
233 use crate::{
234 FeaturePolicy, FeatureSupport, Message, ModelCapabilities, ModelErrorKind, ModelRef,
235 ModelRequest, OutputFormat, SupportLevel,
236 };
237
238 #[test]
239 fn strict_policy_rejects_unknown_required_support() {
240 let request = ModelRequest::new(ModelRef::new("test", "model"), Message::user("hello"));
241
242 let error = ModelCapabilities::default()
243 .validate_request(&request, true)
244 .unwrap_err();
245
246 assert_eq!(error.kind, ModelErrorKind::UnsupportedFeature);
247 }
248
249 #[test]
250 fn best_effort_makes_unknown_support_visible() {
251 let request = ModelRequest::new(ModelRef::new("test", "model"), Message::user("hello"))
252 .feature_policy(FeaturePolicy::BestEffort);
253
254 let warnings = ModelCapabilities::default()
255 .validate_request(&request, true)
256 .unwrap();
257
258 assert_eq!(warnings.len(), 1);
259 assert_eq!(warnings[0].code, "runifold.feature_support_unknown");
260 }
261
262 #[test]
263 fn unsupported_features_fail_even_in_best_effort_mode() {
264 let capabilities = ModelCapabilities {
265 structured_output: FeatureSupport::new(SupportLevel::Unsupported),
266 streaming: FeatureSupport::new(SupportLevel::Native),
267 ..ModelCapabilities::default()
268 };
269 let request = ModelRequest::new(ModelRef::new("test", "model"), Message::user("hello"))
270 .output_format(OutputFormat::Json)
271 .feature_policy(FeaturePolicy::BestEffort);
272
273 let error = capabilities.validate_request(&request, true).unwrap_err();
274
275 assert_eq!(error.kind, ModelErrorKind::UnsupportedFeature);
276 assert_eq!(error.metadata["feature"], "structured_output");
277 }
278}