use std::collections::BTreeSet;
use serde_json::Value;
use tea_protocol::{
CanonicalMessage, ContentBlock, ModelId, ProtocolMetadata, ReasoningEffort, TokenCount,
};
use thiserror::Error;
use crate::{HostedToolKind, HostedToolOptions, ModelSpec};
pub const MAX_SYSTEM_PROMPT_BYTES: usize = 1024 * 1024;
pub const MAX_REQUEST_MESSAGES: usize = 4096;
pub const MAX_MODEL_TOOLS: usize = 256;
pub const MAX_TOOL_DESCRIPTION_BYTES: usize = 16 * 1024;
pub const MAX_TOOL_SCHEMA_BYTES: usize = 256 * 1024;
pub const MAX_TOOL_SCHEMA_DEPTH: usize = 32;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ReasoningOptions {
effort: ReasoningEffort,
budget_tokens: Option<TokenCount>,
}
impl ReasoningOptions {
#[must_use]
pub const fn new(effort: ReasoningEffort) -> Self {
Self {
effort,
budget_tokens: None,
}
}
#[must_use]
pub const fn with_budget(mut self, budget_tokens: TokenCount) -> Self {
self.budget_tokens = Some(budget_tokens);
self
}
#[must_use]
pub const fn effort(self) -> ReasoningEffort {
self.effort
}
#[must_use]
pub const fn budget_tokens(self) -> Option<TokenCount> {
self.budget_tokens
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct FunctionToolDefinition {
name: String,
description: String,
input_schema: Value,
}
impl FunctionToolDefinition {
pub fn new(
name: impl Into<String>,
description: impl Into<String>,
input_schema: Value,
) -> Result<Self, ModelRequestError> {
let name = name.into();
let description = description.into();
validate_tool_contract(&name, &description, &input_schema)?;
Ok(Self {
name,
description,
input_schema,
})
}
#[must_use]
pub fn name(&self) -> &str {
&self.name
}
#[must_use]
pub fn description(&self) -> &str {
&self.description
}
#[must_use]
pub const fn input_schema(&self) -> &Value {
&self.input_schema
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct HostedToolDefinition {
name: String,
description: String,
input_schema: Value,
options: HostedToolOptions,
}
impl HostedToolDefinition {
fn new(
description: impl Into<String>,
input_schema: Value,
options: HostedToolOptions,
) -> Result<Self, ModelRequestError> {
let name = options.kind().name().to_owned();
let description = description.into();
validate_tool_contract(&name, &description, &input_schema)?;
Ok(Self {
name,
description,
input_schema,
options,
})
}
#[must_use]
pub fn name(&self) -> &str {
&self.name
}
#[must_use]
pub fn description(&self) -> &str {
&self.description
}
#[must_use]
pub const fn input_schema(&self) -> &Value {
&self.input_schema
}
#[must_use]
pub const fn kind(&self) -> HostedToolKind {
self.options.kind()
}
#[must_use]
pub const fn options(&self) -> &HostedToolOptions {
&self.options
}
}
#[derive(Debug, Clone, PartialEq)]
pub enum ModelToolDefinition {
Function(FunctionToolDefinition),
Hosted(HostedToolDefinition),
}
impl ModelToolDefinition {
pub fn new(
name: impl Into<String>,
description: impl Into<String>,
input_schema: Value,
) -> Result<Self, ModelRequestError> {
FunctionToolDefinition::new(name, description, input_schema).map(Self::Function)
}
pub fn hosted(
description: impl Into<String>,
input_schema: Value,
options: HostedToolOptions,
) -> Result<Self, ModelRequestError> {
HostedToolDefinition::new(description, input_schema, options).map(Self::Hosted)
}
#[must_use]
pub fn name(&self) -> &str {
match self {
Self::Function(tool) => &tool.name,
Self::Hosted(tool) => &tool.name,
}
}
#[must_use]
pub fn description(&self) -> &str {
match self {
Self::Function(tool) => &tool.description,
Self::Hosted(tool) => &tool.description,
}
}
#[must_use]
pub const fn input_schema(&self) -> &Value {
match self {
Self::Function(tool) => &tool.input_schema,
Self::Hosted(tool) => &tool.input_schema,
}
}
#[must_use]
pub const fn as_function(&self) -> Option<&FunctionToolDefinition> {
match self {
Self::Function(tool) => Some(tool),
Self::Hosted(_) => None,
}
}
#[must_use]
pub const fn as_hosted(&self) -> Option<&HostedToolDefinition> {
match self {
Self::Function(_) => None,
Self::Hosted(tool) => Some(tool),
}
}
#[must_use]
pub const fn hosted_kind(&self) -> Option<HostedToolKind> {
match self {
Self::Function(_) => None,
Self::Hosted(tool) => Some(tool.kind()),
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct ModelRequest {
model_id: ModelId,
system_prompt: Option<String>,
messages: Vec<CanonicalMessage>,
tools: Vec<ModelToolDefinition>,
allow_parallel_tool_calls: bool,
reasoning: Option<ReasoningOptions>,
max_output_tokens: Option<TokenCount>,
metadata: ProtocolMetadata,
}
impl ModelRequest {
pub fn new(
model_id: ModelId,
messages: Vec<CanonicalMessage>,
) -> Result<Self, ModelRequestError> {
validate_messages(&messages)?;
Ok(Self {
model_id,
system_prompt: None,
messages,
tools: Vec::new(),
allow_parallel_tool_calls: false,
reasoning: None,
max_output_tokens: None,
metadata: ProtocolMetadata::default(),
})
}
pub fn with_system_prompt(
mut self,
system_prompt: impl Into<String>,
) -> Result<Self, ModelRequestError> {
let system_prompt = system_prompt.into();
if system_prompt.is_empty()
|| system_prompt.len() > MAX_SYSTEM_PROMPT_BYTES
|| system_prompt.contains('\0')
{
return Err(ModelRequestError::InvalidSystemPrompt);
}
self.system_prompt = Some(system_prompt);
Ok(self)
}
pub fn with_tools(
mut self,
tools: Vec<ModelToolDefinition>,
allow_parallel: bool,
) -> Result<Self, ModelRequestError> {
if tools.len() > MAX_MODEL_TOOLS {
return Err(ModelRequestError::TooManyTools);
}
let mut names = BTreeSet::new();
if tools.iter().any(|tool| !names.insert(tool.name())) {
return Err(ModelRequestError::DuplicateToolName);
}
self.tools = tools;
self.allow_parallel_tool_calls = allow_parallel;
Ok(self)
}
#[must_use]
pub const fn with_reasoning(mut self, reasoning: ReasoningOptions) -> Self {
self.reasoning = Some(reasoning);
self
}
#[must_use]
pub const fn with_max_output_tokens(mut self, max_output_tokens: TokenCount) -> Self {
self.max_output_tokens = Some(max_output_tokens);
self
}
#[must_use]
pub fn with_metadata(mut self, metadata: ProtocolMetadata) -> Self {
self.metadata = metadata;
self
}
pub fn validate_for(&self, model: &ModelSpec) -> Result<(), ModelRequestError> {
if self.model_id != *model.model_id() {
return Err(ModelRequestError::ModelMismatch);
}
validate_messages(&self.messages)?;
let capabilities = model.capabilities();
if request_contains_image(&self.messages) && !capabilities.accepts_images() {
return Err(ModelRequestError::ImageInputUnsupported);
}
if let Some(reasoning) = self.reasoning {
let Some(profile) = model.reasoning_profile() else {
return Err(ModelRequestError::ReasoningUnsupported);
};
if !profile.supported_efforts().contains(&reasoning.effort()) {
return Err(ModelRequestError::ReasoningEffortUnsupported);
}
}
if self.tools.iter().any(|tool| tool.as_function().is_some())
&& !capabilities.supports_tools()
{
return Err(ModelRequestError::ToolsUnsupported);
}
if self.tools.iter().any(|tool| {
tool.hosted_kind()
.is_some_and(|kind| !capabilities.supports_hosted_tool(kind))
}) {
return Err(ModelRequestError::HostedToolUnsupported);
}
if self.allow_parallel_tool_calls
&& self.tools.iter().any(|tool| tool.as_function().is_some())
&& !capabilities.supports_parallel_tool_calls()
{
return Err(ModelRequestError::ParallelToolsUnsupported);
}
let output_limit = self
.max_output_tokens
.unwrap_or_else(|| model.max_output_tokens());
if output_limit.get() == 0 || output_limit > model.max_output_tokens() {
return Err(ModelRequestError::OutputLimitUnsupported);
}
if self
.reasoning
.and_then(ReasoningOptions::budget_tokens)
.is_some_and(|budget| budget.get() == 0 || budget > output_limit)
{
return Err(ModelRequestError::ReasoningBudgetUnsupported);
}
Ok(())
}
#[must_use]
pub const fn model_id(&self) -> &ModelId {
&self.model_id
}
#[must_use]
pub fn system_prompt(&self) -> Option<&str> {
self.system_prompt.as_deref()
}
#[must_use]
pub fn messages(&self) -> &[CanonicalMessage] {
&self.messages
}
#[must_use]
pub fn tools(&self) -> &[ModelToolDefinition] {
&self.tools
}
#[must_use]
pub const fn allow_parallel_tool_calls(&self) -> bool {
self.allow_parallel_tool_calls
}
#[must_use]
pub const fn reasoning(&self) -> Option<ReasoningOptions> {
self.reasoning
}
#[must_use]
pub const fn max_output_tokens(&self) -> Option<TokenCount> {
self.max_output_tokens
}
#[must_use]
pub const fn metadata(&self) -> &ProtocolMetadata {
&self.metadata
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Error)]
pub enum ModelRequestError {
#[error("model request requires at least one message")]
EmptyMessages,
#[error("model request contains too many messages")]
TooManyMessages,
#[error("model request contains an invalid canonical message")]
InvalidMessage,
#[error("system prompt is invalid")]
InvalidSystemPrompt,
#[error("tool name is invalid")]
InvalidToolName,
#[error("tool description is invalid")]
InvalidToolDescription,
#[error("web-search domain must be a canonical lowercase hostname")]
InvalidWebSearchDomain,
#[error("web-search domain policy contains too many domains")]
TooManyWebSearchDomains,
#[error("web-search allowed and blocked domains are mutually exclusive")]
ConflictingWebSearchDomainFilters,
#[error("web-search location is invalid")]
InvalidWebSearchLocation,
#[error("tool input schema must be a JSON object")]
ToolSchemaMustBeObject,
#[error("tool input schema must declare object type")]
ToolSchemaMustDeclareObject,
#[error("tool input schema exceeds supported bounds")]
ToolSchemaOutOfBounds,
#[error("model request contains too many tools")]
TooManyTools,
#[error("model request contains a duplicate tool name")]
DuplicateToolName,
#[error("model request does not match model specification")]
ModelMismatch,
#[error("model does not support image input")]
ImageInputUnsupported,
#[error("model does not support reasoning")]
ReasoningUnsupported,
#[error("model does not support the requested reasoning effort")]
ReasoningEffortUnsupported,
#[error("model does not support tools")]
ToolsUnsupported,
#[error("model does not support parallel tool calls")]
ParallelToolsUnsupported,
#[error("model does not support a requested hosted tool")]
HostedToolUnsupported,
#[error("requested output limit is unsupported")]
OutputLimitUnsupported,
#[error("requested reasoning budget is unsupported")]
ReasoningBudgetUnsupported,
}
fn validate_messages(messages: &[CanonicalMessage]) -> Result<(), ModelRequestError> {
if messages.is_empty() {
return Err(ModelRequestError::EmptyMessages);
}
if messages.len() > MAX_REQUEST_MESSAGES {
return Err(ModelRequestError::TooManyMessages);
}
if messages
.iter()
.any(|message| serde_json::to_value(message).is_err())
{
return Err(ModelRequestError::InvalidMessage);
}
Ok(())
}
fn request_contains_image(messages: &[CanonicalMessage]) -> bool {
messages.iter().any(|message| {
let content = match message {
CanonicalMessage::User { content, .. }
| CanonicalMessage::Assistant { content, .. }
| CanonicalMessage::ToolResult { content, .. } => content,
};
content
.iter()
.any(|block| matches!(block, ContentBlock::Image { .. }))
})
}
fn validate_tool_name(value: &str) -> Result<(), ModelRequestError> {
let mut bytes = value.bytes();
if value.len() > 128
|| !bytes.next().is_some_and(|byte| byte.is_ascii_lowercase())
|| !bytes.all(|byte| {
byte.is_ascii_lowercase() || byte.is_ascii_digit() || matches!(byte, b'_' | b'-' | b'.')
})
{
return Err(ModelRequestError::InvalidToolName);
}
Ok(())
}
fn validate_tool_contract(
name: &str,
description: &str,
input_schema: &Value,
) -> Result<(), ModelRequestError> {
validate_tool_name(name)?;
if description.is_empty()
|| description.len() > MAX_TOOL_DESCRIPTION_BYTES
|| description.contains('\0')
{
return Err(ModelRequestError::InvalidToolDescription);
}
let object = input_schema
.as_object()
.ok_or(ModelRequestError::ToolSchemaMustBeObject)?;
if object.get("type").and_then(Value::as_str) != Some("object") {
return Err(ModelRequestError::ToolSchemaMustDeclareObject);
}
validate_schema_bounds(input_schema)
}
fn validate_schema_bounds(value: &Value) -> Result<(), ModelRequestError> {
if serde_json::to_vec(value)
.map_err(|_| ModelRequestError::ToolSchemaOutOfBounds)?
.len()
> MAX_TOOL_SCHEMA_BYTES
|| json_depth(value) > MAX_TOOL_SCHEMA_DEPTH
{
return Err(ModelRequestError::ToolSchemaOutOfBounds);
}
Ok(())
}
fn json_depth(value: &Value) -> usize {
match value {
Value::Array(values) => 1 + values.iter().map(json_depth).max().unwrap_or(0),
Value::Object(values) => 1 + values.values().map(json_depth).max().unwrap_or(0),
_ => 1,
}
}