1use std::fmt;
11
12use serde::{Deserialize, Serialize};
13
14use crate::ids::{ModelKey, ModelRef, ProviderKey};
15
16#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize, Deserialize)]
20#[serde(rename_all = "snake_case")]
21pub enum StructuredOutputCapability {
22 NativeJsonSchema,
24 NativeFunctionSchema,
27 GrammarConstrained,
29 JsonObject,
31 PromptOnly,
33 None,
35}
36
37impl StructuredOutputCapability {
38 pub const ALL: [Self; 6] = [
40 Self::NativeJsonSchema,
41 Self::NativeFunctionSchema,
42 Self::GrammarConstrained,
43 Self::JsonObject,
44 Self::PromptOnly,
45 Self::None,
46 ];
47
48 #[must_use]
50 pub const fn as_str(self) -> &'static str {
51 match self {
52 Self::NativeJsonSchema => "native_json_schema",
53 Self::NativeFunctionSchema => "native_function_schema",
54 Self::GrammarConstrained => "grammar_constrained",
55 Self::JsonObject => "json_object",
56 Self::PromptOnly => "prompt_only",
57 Self::None => "none",
58 }
59 }
60
61 #[must_use]
64 pub const fn enforces_schema(self) -> bool {
65 matches!(
66 self,
67 Self::NativeJsonSchema | Self::NativeFunctionSchema | Self::GrammarConstrained
68 )
69 }
70}
71
72impl fmt::Display for StructuredOutputCapability {
73 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
74 f.write_str(self.as_str())
75 }
76}
77
78#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize, Deserialize)]
80#[serde(rename_all = "snake_case")]
81pub enum ToolCallingCapability {
82 None,
84 Sequential,
86 Parallel,
88}
89
90impl ToolCallingCapability {
91 #[must_use]
93 pub const fn as_str(self) -> &'static str {
94 match self {
95 Self::None => "none",
96 Self::Sequential => "sequential",
97 Self::Parallel => "parallel",
98 }
99 }
100}
101
102impl fmt::Display for ToolCallingCapability {
103 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
104 f.write_str(self.as_str())
105 }
106}
107
108#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)]
110pub struct ProviderCapabilities {
111 pub structured_output: StructuredOutputCapability,
113 pub tool_calling: ToolCallingCapability,
115 pub parallel_tool_calls: bool,
117 pub vision: bool,
119 #[serde(default)]
126 pub documents: bool,
127 pub audio_input: bool,
129 pub audio_output: bool,
131 pub streaming: bool,
133 pub prompt_caching: bool,
135 pub reasoning_controls: bool,
137 #[serde(default)]
140 pub temperature: bool,
141 #[serde(default)]
143 pub seed: bool,
144 pub max_context_tokens: Option<u64>,
146 pub preserves_call_ids: bool,
148}
149
150impl ProviderCapabilities {
151 #[must_use]
153 pub const fn minimal() -> Self {
154 Self {
155 structured_output: StructuredOutputCapability::None,
156 tool_calling: ToolCallingCapability::None,
157 parallel_tool_calls: false,
158 vision: false,
159 documents: false,
160 audio_input: false,
161 audio_output: false,
162 streaming: false,
163 prompt_caching: false,
164 reasoning_controls: false,
165 temperature: false,
166 seed: false,
167 max_context_tokens: None,
168 preserves_call_ids: false,
169 }
170 }
171
172 #[must_use]
174 pub const fn with_structured_output(mut self, capability: StructuredOutputCapability) -> Self {
175 self.structured_output = capability;
176 self
177 }
178
179 #[must_use]
181 pub const fn with_tool_calling(mut self, capability: ToolCallingCapability) -> Self {
182 self.tool_calling = capability;
183 self.parallel_tool_calls = matches!(capability, ToolCallingCapability::Parallel);
184 self
185 }
186
187 #[must_use]
189 pub const fn with_streaming(mut self, streaming: bool) -> Self {
190 self.streaming = streaming;
191 self
192 }
193
194 #[must_use]
196 pub const fn with_vision(mut self, vision: bool) -> Self {
197 self.vision = vision;
198 self
199 }
200
201 #[must_use]
203 pub const fn with_documents(mut self, documents: bool) -> Self {
204 self.documents = documents;
205 self
206 }
207
208 #[must_use]
210 pub const fn with_max_context_tokens(mut self, tokens: u64) -> Self {
211 self.max_context_tokens = Some(tokens);
212 self
213 }
214
215 #[must_use]
217 pub const fn with_preserves_call_ids(mut self, preserves: bool) -> Self {
218 self.preserves_call_ids = preserves;
219 self
220 }
221
222 #[must_use]
224 pub const fn with_prompt_caching(mut self, caching: bool) -> Self {
225 self.prompt_caching = caching;
226 self
227 }
228
229 #[must_use]
231 pub const fn with_reasoning_controls(mut self, controls: bool) -> Self {
232 self.reasoning_controls = controls;
233 self
234 }
235
236 #[must_use]
238 pub const fn with_temperature(mut self, temperature: bool) -> Self {
239 self.temperature = temperature;
240 self
241 }
242
243 #[must_use]
245 pub const fn with_seed(mut self, seed: bool) -> Self {
246 self.seed = seed;
247 self
248 }
249
250 #[must_use]
252 pub const fn supports_tools(&self) -> bool {
253 !matches!(self.tool_calling, ToolCallingCapability::None)
254 }
255}
256
257impl Default for ProviderCapabilities {
258 fn default() -> Self {
259 Self::minimal()
260 }
261}
262
263#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)]
265#[serde(tag = "kind", rename_all = "snake_case")]
266#[non_exhaustive]
267pub enum MissingCapability {
268 StructuredOutput {
270 required: Vec<StructuredOutputCapability>,
272 declared: StructuredOutputCapability,
274 },
275 ToolCalling,
277 Streaming,
279 Vision,
281 Documents,
283 ContextWindow {
285 required: u64,
287 declared: Option<u64>,
289 },
290}
291
292impl fmt::Display for MissingCapability {
293 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
294 match self {
295 Self::StructuredOutput { required, declared } => {
296 write!(f, "structured_output(declared {declared}, required one of ")?;
297 for (index, capability) in required.iter().enumerate() {
298 if index > 0 {
299 f.write_str("|")?;
300 }
301 write!(f, "{capability}")?;
302 }
303 f.write_str(")")
304 }
305 Self::ToolCalling => f.write_str("tool_calling"),
306 Self::Streaming => f.write_str("streaming"),
307 Self::Vision => f.write_str("vision"),
308 Self::Documents => f.write_str("documents"),
309 Self::ContextWindow { required, declared } => match declared {
310 Some(declared) => write!(
311 f,
312 "context_window(declared {declared}, required {required})"
313 ),
314 None => write!(f, "context_window(undeclared, required {required})"),
315 },
316 }
317 }
318}
319
320#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize, thiserror::Error)]
324pub struct CapabilityMismatch {
325 pub missing: Vec<MissingCapability>,
327}
328
329impl CapabilityMismatch {
330 #[must_use]
333 pub fn structured_output_unmet(&self) -> bool {
334 self.missing
335 .iter()
336 .any(|missing| matches!(missing, MissingCapability::StructuredOutput { .. }))
337 }
338}
339
340impl fmt::Display for CapabilityMismatch {
341 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
342 f.write_str("capability mismatch: ")?;
343 for (index, missing) in self.missing.iter().enumerate() {
344 if index > 0 {
345 f.write_str(", ")?;
346 }
347 write!(f, "{missing}")?;
348 }
349 Ok(())
350 }
351}
352
353#[derive(Debug, Clone, PartialEq, Eq, Hash, Default, Serialize, Deserialize)]
357pub struct CapabilityRequirements {
358 pub structured_output: Vec<StructuredOutputCapability>,
360 pub needs_tools: bool,
362 pub needs_streaming: bool,
364 pub min_context_tokens: Option<u64>,
366 pub needs_vision: bool,
368 #[serde(default)]
370 pub needs_documents: bool,
371}
372
373impl CapabilityRequirements {
374 #[must_use]
376 pub const fn none() -> Self {
377 Self {
378 structured_output: Vec::new(),
379 needs_tools: false,
380 needs_streaming: false,
381 min_context_tokens: None,
382 needs_vision: false,
383 needs_documents: false,
384 }
385 }
386
387 #[must_use]
389 pub fn with_tools(mut self) -> Self {
390 self.needs_tools = true;
391 self
392 }
393
394 #[must_use]
396 pub fn with_streaming(mut self) -> Self {
397 self.needs_streaming = true;
398 self
399 }
400
401 #[must_use]
403 pub fn with_vision(mut self) -> Self {
404 self.needs_vision = true;
405 self
406 }
407
408 #[must_use]
410 pub fn with_documents(mut self) -> Self {
411 self.needs_documents = true;
412 self
413 }
414
415 #[must_use]
417 pub fn with_min_context_tokens(mut self, tokens: u64) -> Self {
418 self.min_context_tokens = Some(tokens);
419 self
420 }
421
422 pub fn satisfied_by(&self, caps: &ProviderCapabilities) -> Result<(), CapabilityMismatch> {
428 let mut missing = Vec::new();
429 if !self.structured_output.is_empty()
430 && !self.structured_output.contains(&caps.structured_output)
431 {
432 missing.push(MissingCapability::StructuredOutput {
433 required: self.structured_output.clone(),
434 declared: caps.structured_output,
435 });
436 }
437 if self.needs_tools && !caps.supports_tools() {
438 missing.push(MissingCapability::ToolCalling);
439 }
440 if self.needs_streaming && !caps.streaming {
441 missing.push(MissingCapability::Streaming);
442 }
443 if self.needs_vision && !caps.vision {
444 missing.push(MissingCapability::Vision);
445 }
446 if self.needs_documents && !caps.documents {
447 missing.push(MissingCapability::Documents);
448 }
449 if let Some(required) = self.min_context_tokens {
450 let ok = caps
451 .max_context_tokens
452 .is_some_and(|declared| declared >= required);
453 if !ok {
454 missing.push(MissingCapability::ContextWindow {
455 required,
456 declared: caps.max_context_tokens,
457 });
458 }
459 }
460 if missing.is_empty() {
461 Ok(())
462 } else {
463 Err(CapabilityMismatch { missing })
464 }
465 }
466}
467
468#[derive(
471 Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord, Default, Serialize, Deserialize,
472)]
473#[serde(transparent)]
474pub struct MicroCents(pub u64);
475
476impl MicroCents {
477 pub const PER_CENT: u64 = 1_000_000;
479
480 #[must_use]
482 pub const fn from_cents(cents: u64) -> Self {
483 Self(cents.saturating_mul(Self::PER_CENT))
484 }
485
486 #[must_use]
488 pub const fn from_dollars(dollars: u64) -> Self {
489 Self::from_cents(dollars.saturating_mul(100))
490 }
491
492 #[must_use]
494 pub const fn value(self) -> u64 {
495 self.0
496 }
497
498 #[must_use]
500 pub const fn saturating_add(self, other: Self) -> Self {
501 Self(self.0.saturating_add(other.0))
502 }
503}
504
505impl fmt::Display for MicroCents {
506 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
507 write!(f, "{}µ¢", self.0)
508 }
509}
510
511#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
517pub struct ModelProfile {
518 pub provider: ProviderKey,
520 pub model: ModelKey,
522 pub capabilities: ProviderCapabilities,
524 #[serde(default, skip_serializing_if = "Option::is_none")]
526 pub cost_per_million_input: Option<MicroCents>,
527 #[serde(default, skip_serializing_if = "Option::is_none")]
529 pub cost_per_million_output: Option<MicroCents>,
530 #[serde(default, skip_serializing_if = "Option::is_none")]
532 pub region: Option<String>,
533 #[serde(default)]
535 pub tags: Vec<String>,
536}
537
538impl ModelProfile {
539 #[must_use]
541 pub fn new(
542 provider: impl Into<ProviderKey>,
543 model: impl Into<ModelKey>,
544 capabilities: ProviderCapabilities,
545 ) -> Self {
546 Self {
547 provider: provider.into(),
548 model: model.into(),
549 capabilities,
550 cost_per_million_input: None,
551 cost_per_million_output: None,
552 region: None,
553 tags: Vec::new(),
554 }
555 }
556
557 #[must_use]
559 pub fn with_cost(mut self, input: MicroCents, output: MicroCents) -> Self {
560 self.cost_per_million_input = Some(input);
561 self.cost_per_million_output = Some(output);
562 self
563 }
564
565 #[must_use]
567 pub fn with_region(mut self, region: impl Into<String>) -> Self {
568 self.region = Some(region.into());
569 self
570 }
571
572 #[must_use]
574 pub fn with_tag(mut self, tag: impl Into<String>) -> Self {
575 self.tags.push(tag.into());
576 self
577 }
578
579 #[must_use]
581 pub fn reference(&self) -> ModelRef {
582 ModelRef {
583 provider: self.provider.clone(),
584 model: self.model.clone(),
585 }
586 }
587
588 #[must_use]
593 pub fn max_cost_per_million(&self) -> Option<MicroCents> {
594 match (self.cost_per_million_input, self.cost_per_million_output) {
595 (Some(input), Some(output)) => Some(input.max(output)),
596 _ => None,
597 }
598 }
599
600 #[must_use]
602 pub fn estimate_cost(&self, usage: &crate::response::TokenUsage) -> Option<MicroCents> {
603 let input = self.cost_per_million_input?;
604 let output = self.cost_per_million_output?;
605 let per_token = |rate: MicroCents, tokens: u64| -> u64 {
606 rate.0.saturating_mul(tokens) / 1_000_000
608 };
609 Some(MicroCents(
610 per_token(input, usage.input).saturating_add(per_token(output, usage.output)),
611 ))
612 }
613}
614
615#[cfg(test)]
616mod tests {
617 use super::*;
618
619 #[test]
620 fn satisfied_by_reports_every_gap_in_order() {
621 let requirements = CapabilityRequirements {
622 structured_output: vec![StructuredOutputCapability::NativeJsonSchema],
623 needs_tools: true,
624 needs_streaming: true,
625 min_context_tokens: Some(100_000),
626 needs_vision: true,
627 needs_documents: true,
628 };
629 let caps = ProviderCapabilities::minimal();
630 let mismatch = requirements.satisfied_by(&caps).unwrap_err();
631 assert_eq!(mismatch.missing.len(), 6);
632 assert!(mismatch.structured_output_unmet());
633 assert!(matches!(
634 mismatch.missing[0],
635 MissingCapability::StructuredOutput { .. }
636 ));
637 assert!(matches!(mismatch.missing[4], MissingCapability::Documents));
638 assert!(matches!(
639 mismatch.missing[5],
640 MissingCapability::ContextWindow {
641 required: 100_000,
642 declared: None
643 }
644 ));
645 let text = mismatch.to_string();
646 assert!(text.contains("structured_output(declared none"));
647 assert!(text.contains("documents"));
648 assert!(text.contains("context_window(undeclared"));
649 }
650
651 #[test]
652 fn documents_are_declared_apart_from_images() {
653 let sighted = ProviderCapabilities::minimal().with_vision(true);
656 let requirements = CapabilityRequirements::none().with_documents();
657 let mismatch = requirements.satisfied_by(&sighted).unwrap_err();
658 assert_eq!(mismatch.missing, vec![MissingCapability::Documents]);
659 assert_eq!(mismatch.to_string(), "capability mismatch: documents");
660
661 let reader = sighted.with_documents(true);
662 assert!(requirements.satisfied_by(&reader).is_ok());
663 let paper_only = ProviderCapabilities::minimal().with_documents(true);
665 assert!(
666 CapabilityRequirements::none()
667 .with_vision()
668 .satisfied_by(&paper_only)
669 .is_err()
670 );
671 }
672
673 #[test]
674 fn a_declaration_written_before_the_documents_flag_reads_as_false() {
675 let stored = serde_json::to_value(ProviderCapabilities::minimal().with_vision(true))
678 .expect("serializes");
679 let mut object = stored.as_object().expect("an object").clone();
680 object.remove("documents");
681 let older: ProviderCapabilities =
682 serde_json::from_value(serde_json::Value::Object(object)).expect("still decodes");
683 assert!(!older.documents);
684 assert!(older.vision);
685 }
686
687 #[test]
688 fn context_window_is_fail_closed_and_compared_numerically() {
689 let requirements = CapabilityRequirements::none().with_min_context_tokens(8_000);
690 assert!(
691 requirements
692 .satisfied_by(&ProviderCapabilities::minimal())
693 .is_err()
694 );
695 let small = ProviderCapabilities::minimal().with_max_context_tokens(4_000);
696 assert!(requirements.satisfied_by(&small).is_err());
697 let large = ProviderCapabilities::minimal().with_max_context_tokens(8_000);
698 assert!(requirements.satisfied_by(&large).is_ok());
699 }
700
701 #[test]
702 fn empty_structured_set_means_no_requirement() {
703 let requirements = CapabilityRequirements::none();
704 assert!(
705 requirements
706 .satisfied_by(&ProviderCapabilities::minimal())
707 .is_ok()
708 );
709 }
710
711 #[test]
712 fn profile_cost_helpers() {
713 let profile = ModelProfile::new("p", "m", ProviderCapabilities::minimal())
714 .with_cost(MicroCents::from_cents(250), MicroCents::from_dollars(10))
715 .with_region("eu")
716 .with_tag("cheap");
717 assert_eq!(
718 profile.max_cost_per_million(),
719 Some(MicroCents::from_dollars(10))
720 );
721 let usage = crate::response::TokenUsage::new(1_000_000, 500_000);
722 assert_eq!(
723 profile.estimate_cost(&usage),
724 Some(MicroCents::from_cents(250).saturating_add(MicroCents::from_dollars(5)))
725 );
726 assert_eq!(profile.reference().to_string(), "p/m");
727 let unknown = ModelProfile::new("p", "m", ProviderCapabilities::minimal());
728 assert_eq!(unknown.max_cost_per_million(), None);
729 assert_eq!(unknown.estimate_cost(&usage), None);
730 }
731
732 #[test]
733 fn labels_serialize_snake_case() {
734 assert_eq!(
735 serde_json::to_string(&StructuredOutputCapability::NativeJsonSchema).unwrap(),
736 "\"native_json_schema\""
737 );
738 assert_eq!(
739 serde_json::to_string(&ToolCallingCapability::Parallel).unwrap(),
740 "\"parallel\""
741 );
742 let caps =
743 ProviderCapabilities::minimal().with_tool_calling(ToolCallingCapability::Parallel);
744 assert!(caps.parallel_tool_calls && caps.supports_tools());
745 }
746}