use std::fmt;
use std::str::FromStr;
use tea_protocol::{ModelId, ModelRef, ProviderId, ReasoningEffort, TokenCount};
use thiserror::Error;
use crate::HostedToolKind;
const MAX_DISPLAY_NAME_BYTES: usize = 256;
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct ModelDisplayName(String);
impl ModelDisplayName {
#[must_use]
pub fn as_str(&self) -> &str {
&self.0
}
}
impl FromStr for ModelDisplayName {
type Err = ModelTextParseError;
fn from_str(value: &str) -> Result<Self, Self::Err> {
if value.is_empty()
|| value.len() > MAX_DISPLAY_NAME_BYTES
|| value.chars().any(char::is_control)
{
return Err(ModelTextParseError::InvalidDisplayName);
}
Ok(Self(value.to_owned()))
}
}
impl fmt::Display for ModelDisplayName {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter.write_str(&self.0)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Error)]
pub enum ModelTextParseError {
#[error("model display name is invalid")]
InvalidDisplayName,
}
const CAP_IMAGE_INPUT: u16 = 1 << 0;
const CAP_REASONING: u16 = 1 << 1;
const CAP_TOOLS: u16 = 1 << 2;
const CAP_PARALLEL_TOOLS: u16 = 1 << 3;
const CAP_USAGE: u16 = 1 << 4;
const CAP_HOSTED_WEB_SEARCH: u16 = 1 << 5;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ModelCapabilities(u16);
impl ModelCapabilities {
#[must_use]
pub const fn text() -> Self {
Self(0)
}
#[must_use]
pub const fn with_image_input(mut self) -> Self {
self.0 |= CAP_IMAGE_INPUT;
self
}
#[must_use]
pub const fn with_reasoning(mut self) -> Self {
self.0 |= CAP_REASONING;
self
}
#[must_use]
pub const fn with_tools(mut self, parallel: bool) -> Self {
self.0 |= CAP_TOOLS;
if parallel {
self.0 |= CAP_PARALLEL_TOOLS;
} else {
self.0 &= !CAP_PARALLEL_TOOLS;
}
self
}
#[must_use]
pub const fn with_usage_reporting(mut self) -> Self {
self.0 |= CAP_USAGE;
self
}
#[must_use]
pub const fn with_hosted_tool(mut self, kind: HostedToolKind) -> Self {
match kind {
HostedToolKind::WebSearch => self.0 |= CAP_HOSTED_WEB_SEARCH,
}
self
}
#[must_use]
pub const fn accepts_text(self) -> bool {
true
}
#[must_use]
pub const fn accepts_images(self) -> bool {
self.0 & CAP_IMAGE_INPUT != 0
}
#[must_use]
pub const fn supports_reasoning(self) -> bool {
self.0 & CAP_REASONING != 0
}
#[must_use]
pub const fn supports_tools(self) -> bool {
self.0 & CAP_TOOLS != 0
}
#[must_use]
pub const fn supports_parallel_tool_calls(self) -> bool {
self.0 & CAP_PARALLEL_TOOLS != 0
}
#[must_use]
pub const fn reports_usage(self) -> bool {
self.0 & CAP_USAGE != 0
}
#[must_use]
pub const fn supports_hosted_tool(self, kind: HostedToolKind) -> bool {
match kind {
HostedToolKind::WebSearch => self.0 & CAP_HOSTED_WEB_SEARCH != 0,
}
}
}
impl Default for ModelCapabilities {
fn default() -> Self {
Self::text()
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ReasoningProfile {
default_effort: ReasoningEffort,
supported_efforts: Vec<ReasoningEffort>,
}
impl ReasoningProfile {
pub fn new(
default_effort: ReasoningEffort,
supported_efforts: impl IntoIterator<Item = ReasoningEffort>,
) -> Result<Self, ModelSpecError> {
let mut supported_efforts = supported_efforts.into_iter().collect::<Vec<_>>();
if supported_efforts.is_empty() {
return Err(ModelSpecError::EmptyReasoningEfforts);
}
supported_efforts.sort_unstable();
if supported_efforts
.windows(2)
.any(|levels| levels[0] == levels[1])
{
return Err(ModelSpecError::DuplicateReasoningEffort);
}
if !supported_efforts.contains(&default_effort) {
return Err(ModelSpecError::ReasoningDefaultUnsupported);
}
Ok(Self {
default_effort,
supported_efforts,
})
}
pub(crate) fn compatible_default() -> Self {
Self {
default_effort: ReasoningEffort::Medium,
supported_efforts: ReasoningEffort::SHORTCUT_LEVELS.to_vec(),
}
}
#[must_use]
pub const fn default_effort(&self) -> ReasoningEffort {
self.default_effort
}
#[must_use]
pub fn supported_efforts(&self) -> &[ReasoningEffort] {
&self.supported_efforts
}
#[must_use]
pub fn resolve(&self, requested: ReasoningEffort) -> ReasoningResolution {
let effective = if self.supported_efforts.contains(&requested) {
requested
} else {
self.supported_efforts
.iter()
.copied()
.find(|candidate| *candidate > requested)
.or_else(|| {
self.supported_efforts
.iter()
.rev()
.copied()
.find(|candidate| *candidate < requested)
})
.unwrap_or(self.default_effort)
};
ReasoningResolution {
requested,
effective,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ReasoningResolution {
requested: ReasoningEffort,
effective: ReasoningEffort,
}
impl ReasoningResolution {
#[must_use]
pub const fn requested(self) -> ReasoningEffort {
self.requested
}
#[must_use]
pub const fn effective(self) -> ReasoningEffort {
self.effective
}
#[must_use]
pub fn was_clamped(self) -> bool {
self.requested != self.effective
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ModelSpec {
model_ref: ModelRef,
display_name: ModelDisplayName,
context_window_tokens: TokenCount,
max_output_tokens: TokenCount,
capabilities: ModelCapabilities,
reasoning_profile: Option<ReasoningProfile>,
}
impl ModelSpec {
pub fn new(
model_id: ModelId,
provider_id: ProviderId,
display_name: ModelDisplayName,
context_window_tokens: TokenCount,
max_output_tokens: TokenCount,
capabilities: ModelCapabilities,
) -> Result<Self, ModelSpecError> {
if context_window_tokens.get() == 0 {
return Err(ModelSpecError::EmptyContextWindow);
}
if max_output_tokens.get() == 0 {
return Err(ModelSpecError::EmptyOutputLimit);
}
if max_output_tokens > context_window_tokens {
return Err(ModelSpecError::OutputExceedsContext);
}
let reasoning_profile = capabilities
.supports_reasoning()
.then(ReasoningProfile::compatible_default);
Ok(Self {
model_ref: ModelRef::new(provider_id, model_id),
display_name,
context_window_tokens,
max_output_tokens,
capabilities,
reasoning_profile,
})
}
#[must_use]
pub fn with_reasoning_profile(mut self, profile: ReasoningProfile) -> Self {
self.capabilities = self.capabilities.with_reasoning();
self.reasoning_profile = Some(profile);
self
}
#[must_use]
pub const fn model_id(&self) -> &ModelId {
self.model_ref.model_id()
}
#[must_use]
pub const fn provider_id(&self) -> &ProviderId {
self.model_ref.provider_id()
}
#[must_use]
pub const fn model_ref(&self) -> &ModelRef {
&self.model_ref
}
#[must_use]
pub const fn display_name(&self) -> &ModelDisplayName {
&self.display_name
}
#[must_use]
pub const fn context_window_tokens(&self) -> TokenCount {
self.context_window_tokens
}
#[must_use]
pub const fn max_output_tokens(&self) -> TokenCount {
self.max_output_tokens
}
#[must_use]
pub const fn capabilities(&self) -> ModelCapabilities {
self.capabilities
}
#[must_use]
pub const fn reasoning_profile(&self) -> Option<&ReasoningProfile> {
self.reasoning_profile.as_ref()
}
#[must_use]
pub fn resolve_reasoning(
&self,
requested: Option<ReasoningEffort>,
) -> Option<ReasoningResolution> {
match (&self.reasoning_profile, requested) {
(Some(profile), Some(requested)) => Some(profile.resolve(requested)),
(Some(profile), None) => Some(profile.resolve(profile.default_effort())),
(None, Some(requested)) => Some(ReasoningResolution {
requested,
effective: ReasoningEffort::Off,
}),
(None, None) => None,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Error)]
pub enum ModelSpecError {
#[error("model context window must be non-zero")]
EmptyContextWindow,
#[error("model output limit must be non-zero")]
EmptyOutputLimit,
#[error("model output limit exceeds context window")]
OutputExceedsContext,
#[error("model reasoning profile is empty")]
EmptyReasoningEfforts,
#[error("model reasoning profile contains a duplicate effort")]
DuplicateReasoningEffort,
#[error("model reasoning default is unsupported")]
ReasoningDefaultUnsupported,
}