use std::fmt;
use serde::{Deserialize, Serialize};
use crate::ids::{ModelKey, ModelRef, ProviderKey};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum StructuredOutputCapability {
NativeJsonSchema,
NativeFunctionSchema,
GrammarConstrained,
JsonObject,
PromptOnly,
None,
}
impl StructuredOutputCapability {
pub const ALL: [Self; 6] = [
Self::NativeJsonSchema,
Self::NativeFunctionSchema,
Self::GrammarConstrained,
Self::JsonObject,
Self::PromptOnly,
Self::None,
];
#[must_use]
pub const fn as_str(self) -> &'static str {
match self {
Self::NativeJsonSchema => "native_json_schema",
Self::NativeFunctionSchema => "native_function_schema",
Self::GrammarConstrained => "grammar_constrained",
Self::JsonObject => "json_object",
Self::PromptOnly => "prompt_only",
Self::None => "none",
}
}
#[must_use]
pub const fn enforces_schema(self) -> bool {
matches!(
self,
Self::NativeJsonSchema | Self::NativeFunctionSchema | Self::GrammarConstrained
)
}
}
impl fmt::Display for StructuredOutputCapability {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(self.as_str())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum ToolCallingCapability {
None,
Sequential,
Parallel,
}
impl ToolCallingCapability {
#[must_use]
pub const fn as_str(self) -> &'static str {
match self {
Self::None => "none",
Self::Sequential => "sequential",
Self::Parallel => "parallel",
}
}
}
impl fmt::Display for ToolCallingCapability {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(self.as_str())
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub struct ProviderCapabilities {
pub structured_output: StructuredOutputCapability,
pub tool_calling: ToolCallingCapability,
pub parallel_tool_calls: bool,
pub vision: bool,
#[serde(default)]
pub documents: bool,
pub audio_input: bool,
pub audio_output: bool,
pub streaming: bool,
pub prompt_caching: bool,
pub reasoning_controls: bool,
#[serde(default)]
pub temperature: bool,
#[serde(default)]
pub seed: bool,
pub max_context_tokens: Option<u64>,
pub preserves_call_ids: bool,
}
impl ProviderCapabilities {
#[must_use]
pub const fn minimal() -> Self {
Self {
structured_output: StructuredOutputCapability::None,
tool_calling: ToolCallingCapability::None,
parallel_tool_calls: false,
vision: false,
documents: false,
audio_input: false,
audio_output: false,
streaming: false,
prompt_caching: false,
reasoning_controls: false,
temperature: false,
seed: false,
max_context_tokens: None,
preserves_call_ids: false,
}
}
#[must_use]
pub const fn with_structured_output(mut self, capability: StructuredOutputCapability) -> Self {
self.structured_output = capability;
self
}
#[must_use]
pub const fn with_tool_calling(mut self, capability: ToolCallingCapability) -> Self {
self.tool_calling = capability;
self.parallel_tool_calls = matches!(capability, ToolCallingCapability::Parallel);
self
}
#[must_use]
pub const fn with_streaming(mut self, streaming: bool) -> Self {
self.streaming = streaming;
self
}
#[must_use]
pub const fn with_vision(mut self, vision: bool) -> Self {
self.vision = vision;
self
}
#[must_use]
pub const fn with_documents(mut self, documents: bool) -> Self {
self.documents = documents;
self
}
#[must_use]
pub const fn with_max_context_tokens(mut self, tokens: u64) -> Self {
self.max_context_tokens = Some(tokens);
self
}
#[must_use]
pub const fn with_preserves_call_ids(mut self, preserves: bool) -> Self {
self.preserves_call_ids = preserves;
self
}
#[must_use]
pub const fn with_prompt_caching(mut self, caching: bool) -> Self {
self.prompt_caching = caching;
self
}
#[must_use]
pub const fn with_reasoning_controls(mut self, controls: bool) -> Self {
self.reasoning_controls = controls;
self
}
#[must_use]
pub const fn with_temperature(mut self, temperature: bool) -> Self {
self.temperature = temperature;
self
}
#[must_use]
pub const fn with_seed(mut self, seed: bool) -> Self {
self.seed = seed;
self
}
#[must_use]
pub const fn supports_tools(&self) -> bool {
!matches!(self.tool_calling, ToolCallingCapability::None)
}
}
impl Default for ProviderCapabilities {
fn default() -> Self {
Self::minimal()
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(tag = "kind", rename_all = "snake_case")]
#[non_exhaustive]
pub enum MissingCapability {
StructuredOutput {
required: Vec<StructuredOutputCapability>,
declared: StructuredOutputCapability,
},
ToolCalling,
Streaming,
Vision,
Documents,
ContextWindow {
required: u64,
declared: Option<u64>,
},
}
impl fmt::Display for MissingCapability {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::StructuredOutput { required, declared } => {
write!(f, "structured_output(declared {declared}, required one of ")?;
for (index, capability) in required.iter().enumerate() {
if index > 0 {
f.write_str("|")?;
}
write!(f, "{capability}")?;
}
f.write_str(")")
}
Self::ToolCalling => f.write_str("tool_calling"),
Self::Streaming => f.write_str("streaming"),
Self::Vision => f.write_str("vision"),
Self::Documents => f.write_str("documents"),
Self::ContextWindow { required, declared } => match declared {
Some(declared) => write!(
f,
"context_window(declared {declared}, required {required})"
),
None => write!(f, "context_window(undeclared, required {required})"),
},
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize, thiserror::Error)]
pub struct CapabilityMismatch {
pub missing: Vec<MissingCapability>,
}
impl CapabilityMismatch {
#[must_use]
pub fn structured_output_unmet(&self) -> bool {
self.missing
.iter()
.any(|missing| matches!(missing, MissingCapability::StructuredOutput { .. }))
}
}
impl fmt::Display for CapabilityMismatch {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("capability mismatch: ")?;
for (index, missing) in self.missing.iter().enumerate() {
if index > 0 {
f.write_str(", ")?;
}
write!(f, "{missing}")?;
}
Ok(())
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash, Default, Serialize, Deserialize)]
pub struct CapabilityRequirements {
pub structured_output: Vec<StructuredOutputCapability>,
pub needs_tools: bool,
pub needs_streaming: bool,
pub min_context_tokens: Option<u64>,
pub needs_vision: bool,
#[serde(default)]
pub needs_documents: bool,
}
impl CapabilityRequirements {
#[must_use]
pub const fn none() -> Self {
Self {
structured_output: Vec::new(),
needs_tools: false,
needs_streaming: false,
min_context_tokens: None,
needs_vision: false,
needs_documents: false,
}
}
#[must_use]
pub fn with_tools(mut self) -> Self {
self.needs_tools = true;
self
}
#[must_use]
pub fn with_streaming(mut self) -> Self {
self.needs_streaming = true;
self
}
#[must_use]
pub fn with_vision(mut self) -> Self {
self.needs_vision = true;
self
}
#[must_use]
pub fn with_documents(mut self) -> Self {
self.needs_documents = true;
self
}
#[must_use]
pub fn with_min_context_tokens(mut self, tokens: u64) -> Self {
self.min_context_tokens = Some(tokens);
self
}
pub fn satisfied_by(&self, caps: &ProviderCapabilities) -> Result<(), CapabilityMismatch> {
let mut missing = Vec::new();
if !self.structured_output.is_empty()
&& !self.structured_output.contains(&caps.structured_output)
{
missing.push(MissingCapability::StructuredOutput {
required: self.structured_output.clone(),
declared: caps.structured_output,
});
}
if self.needs_tools && !caps.supports_tools() {
missing.push(MissingCapability::ToolCalling);
}
if self.needs_streaming && !caps.streaming {
missing.push(MissingCapability::Streaming);
}
if self.needs_vision && !caps.vision {
missing.push(MissingCapability::Vision);
}
if self.needs_documents && !caps.documents {
missing.push(MissingCapability::Documents);
}
if let Some(required) = self.min_context_tokens {
let ok = caps
.max_context_tokens
.is_some_and(|declared| declared >= required);
if !ok {
missing.push(MissingCapability::ContextWindow {
required,
declared: caps.max_context_tokens,
});
}
}
if missing.is_empty() {
Ok(())
} else {
Err(CapabilityMismatch { missing })
}
}
}
#[derive(
Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord, Default, Serialize, Deserialize,
)]
#[serde(transparent)]
pub struct MicroCents(pub u64);
impl MicroCents {
pub const PER_CENT: u64 = 1_000_000;
#[must_use]
pub const fn from_cents(cents: u64) -> Self {
Self(cents.saturating_mul(Self::PER_CENT))
}
#[must_use]
pub const fn from_dollars(dollars: u64) -> Self {
Self::from_cents(dollars.saturating_mul(100))
}
#[must_use]
pub const fn value(self) -> u64 {
self.0
}
#[must_use]
pub const fn saturating_add(self, other: Self) -> Self {
Self(self.0.saturating_add(other.0))
}
}
impl fmt::Display for MicroCents {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}µ¢", self.0)
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct ModelProfile {
pub provider: ProviderKey,
pub model: ModelKey,
pub capabilities: ProviderCapabilities,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub cost_per_million_input: Option<MicroCents>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub cost_per_million_output: Option<MicroCents>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub region: Option<String>,
#[serde(default)]
pub tags: Vec<String>,
}
impl ModelProfile {
#[must_use]
pub fn new(
provider: impl Into<ProviderKey>,
model: impl Into<ModelKey>,
capabilities: ProviderCapabilities,
) -> Self {
Self {
provider: provider.into(),
model: model.into(),
capabilities,
cost_per_million_input: None,
cost_per_million_output: None,
region: None,
tags: Vec::new(),
}
}
#[must_use]
pub fn with_cost(mut self, input: MicroCents, output: MicroCents) -> Self {
self.cost_per_million_input = Some(input);
self.cost_per_million_output = Some(output);
self
}
#[must_use]
pub fn with_region(mut self, region: impl Into<String>) -> Self {
self.region = Some(region.into());
self
}
#[must_use]
pub fn with_tag(mut self, tag: impl Into<String>) -> Self {
self.tags.push(tag.into());
self
}
#[must_use]
pub fn reference(&self) -> ModelRef {
ModelRef {
provider: self.provider.clone(),
model: self.model.clone(),
}
}
#[must_use]
pub fn max_cost_per_million(&self) -> Option<MicroCents> {
match (self.cost_per_million_input, self.cost_per_million_output) {
(Some(input), Some(output)) => Some(input.max(output)),
_ => None,
}
}
#[must_use]
pub fn estimate_cost(&self, usage: &crate::response::TokenUsage) -> Option<MicroCents> {
let input = self.cost_per_million_input?;
let output = self.cost_per_million_output?;
let per_token = |rate: MicroCents, tokens: u64| -> u64 {
rate.0.saturating_mul(tokens) / 1_000_000
};
Some(MicroCents(
per_token(input, usage.input).saturating_add(per_token(output, usage.output)),
))
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn satisfied_by_reports_every_gap_in_order() {
let requirements = CapabilityRequirements {
structured_output: vec![StructuredOutputCapability::NativeJsonSchema],
needs_tools: true,
needs_streaming: true,
min_context_tokens: Some(100_000),
needs_vision: true,
needs_documents: true,
};
let caps = ProviderCapabilities::minimal();
let mismatch = requirements.satisfied_by(&caps).unwrap_err();
assert_eq!(mismatch.missing.len(), 6);
assert!(mismatch.structured_output_unmet());
assert!(matches!(
mismatch.missing[0],
MissingCapability::StructuredOutput { .. }
));
assert!(matches!(mismatch.missing[4], MissingCapability::Documents));
assert!(matches!(
mismatch.missing[5],
MissingCapability::ContextWindow {
required: 100_000,
declared: None
}
));
let text = mismatch.to_string();
assert!(text.contains("structured_output(declared none"));
assert!(text.contains("documents"));
assert!(text.contains("context_window(undeclared"));
}
#[test]
fn documents_are_declared_apart_from_images() {
let sighted = ProviderCapabilities::minimal().with_vision(true);
let requirements = CapabilityRequirements::none().with_documents();
let mismatch = requirements.satisfied_by(&sighted).unwrap_err();
assert_eq!(mismatch.missing, vec![MissingCapability::Documents]);
assert_eq!(mismatch.to_string(), "capability mismatch: documents");
let reader = sighted.with_documents(true);
assert!(requirements.satisfied_by(&reader).is_ok());
let paper_only = ProviderCapabilities::minimal().with_documents(true);
assert!(
CapabilityRequirements::none()
.with_vision()
.satisfied_by(&paper_only)
.is_err()
);
}
#[test]
fn a_declaration_written_before_the_documents_flag_reads_as_false() {
let stored = serde_json::to_value(ProviderCapabilities::minimal().with_vision(true))
.expect("serializes");
let mut object = stored.as_object().expect("an object").clone();
object.remove("documents");
let older: ProviderCapabilities =
serde_json::from_value(serde_json::Value::Object(object)).expect("still decodes");
assert!(!older.documents);
assert!(older.vision);
}
#[test]
fn context_window_is_fail_closed_and_compared_numerically() {
let requirements = CapabilityRequirements::none().with_min_context_tokens(8_000);
assert!(
requirements
.satisfied_by(&ProviderCapabilities::minimal())
.is_err()
);
let small = ProviderCapabilities::minimal().with_max_context_tokens(4_000);
assert!(requirements.satisfied_by(&small).is_err());
let large = ProviderCapabilities::minimal().with_max_context_tokens(8_000);
assert!(requirements.satisfied_by(&large).is_ok());
}
#[test]
fn empty_structured_set_means_no_requirement() {
let requirements = CapabilityRequirements::none();
assert!(
requirements
.satisfied_by(&ProviderCapabilities::minimal())
.is_ok()
);
}
#[test]
fn profile_cost_helpers() {
let profile = ModelProfile::new("p", "m", ProviderCapabilities::minimal())
.with_cost(MicroCents::from_cents(250), MicroCents::from_dollars(10))
.with_region("eu")
.with_tag("cheap");
assert_eq!(
profile.max_cost_per_million(),
Some(MicroCents::from_dollars(10))
);
let usage = crate::response::TokenUsage::new(1_000_000, 500_000);
assert_eq!(
profile.estimate_cost(&usage),
Some(MicroCents::from_cents(250).saturating_add(MicroCents::from_dollars(5)))
);
assert_eq!(profile.reference().to_string(), "p/m");
let unknown = ModelProfile::new("p", "m", ProviderCapabilities::minimal());
assert_eq!(unknown.max_cost_per_million(), None);
assert_eq!(unknown.estimate_cost(&usage), None);
}
#[test]
fn labels_serialize_snake_case() {
assert_eq!(
serde_json::to_string(&StructuredOutputCapability::NativeJsonSchema).unwrap(),
"\"native_json_schema\""
);
assert_eq!(
serde_json::to_string(&ToolCallingCapability::Parallel).unwrap(),
"\"parallel\""
);
let caps =
ProviderCapabilities::minimal().with_tool_calling(ToolCallingCapability::Parallel);
assert!(caps.parallel_tool_calls && caps.supports_tools());
}
}