use std::collections::BTreeMap;
use std::fmt;
use std::time::Duration;
use serde::{Deserialize, Serialize};
use crate::ids::{CallId, RequestId};
use crate::purpose::ModelPurpose;
pub const DEFAULT_TIMEOUT: Duration = Duration::from_secs(60);
pub const MAX_METADATA_VALUE_LEN: usize = 64;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum Role {
System,
User,
Assistant,
Tool,
}
impl Role {
#[must_use]
pub const fn as_str(self) -> &'static str {
match self {
Self::System => "system",
Self::User => "user",
Self::Assistant => "assistant",
Self::Tool => "tool",
}
}
}
impl fmt::Display for Role {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(self.as_str())
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(tag = "kind", rename_all = "snake_case")]
#[non_exhaustive]
pub enum ImageSource {
Url {
url: String,
},
Base64 {
media_type: String,
data: String,
},
}
#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(tag = "kind", rename_all = "snake_case")]
#[non_exhaustive]
pub enum DocumentSource {
Url {
url: String,
media_type: String,
},
Base64 {
media_type: String,
data: String,
},
}
impl DocumentSource {
#[must_use]
pub fn media_type(&self) -> &str {
match self {
Self::Url { media_type, .. } | Self::Base64 { media_type, .. } => media_type,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct ToolCall {
pub id: CallId,
pub name: String,
pub arguments: serde_json::Value,
}
impl ToolCall {
#[must_use]
pub fn new(
id: impl Into<CallId>,
name: impl Into<String>,
arguments: serde_json::Value,
) -> Self {
Self {
id: id.into(),
name: name.into(),
arguments,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct ToolResult {
pub call_id: CallId,
pub content: String,
#[serde(default)]
pub is_error: bool,
}
impl ToolResult {
#[must_use]
pub fn ok(call_id: impl Into<CallId>, content: impl Into<String>) -> Self {
Self {
call_id: call_id.into(),
content: content.into(),
is_error: false,
}
}
#[must_use]
pub fn error(call_id: impl Into<CallId>, content: impl Into<String>) -> Self {
Self {
call_id: call_id.into(),
content: content.into(),
is_error: true,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
#[non_exhaustive]
pub enum ContentPart {
Text {
text: String,
},
Image {
source: ImageSource,
},
Document {
source: DocumentSource,
},
ToolCall(ToolCall),
ToolResult(ToolResult),
}
impl ContentPart {
#[must_use]
pub fn text(text: impl Into<String>) -> Self {
Self::Text { text: text.into() }
}
#[must_use]
pub fn image_url(url: impl Into<String>) -> Self {
Self::Image {
source: ImageSource::Url { url: url.into() },
}
}
#[must_use]
pub fn image_base64(media_type: impl Into<String>, data: impl Into<String>) -> Self {
Self::Image {
source: ImageSource::Base64 {
media_type: media_type.into(),
data: data.into(),
},
}
}
#[must_use]
pub fn inline_bytes(media_type: impl Into<String>, bytes: &[u8]) -> Self {
use base64::Engine as _;
let media_type = media_type.into();
let data = base64::engine::general_purpose::STANDARD.encode(bytes);
if media_type.starts_with("image/") {
Self::Image {
source: ImageSource::Base64 { media_type, data },
}
} else {
Self::Document {
source: DocumentSource::Base64 { media_type, data },
}
}
}
#[must_use]
pub fn document_url(url: impl Into<String>, media_type: impl Into<String>) -> Self {
Self::Document {
source: DocumentSource::Url {
url: url.into(),
media_type: media_type.into(),
},
}
}
#[must_use]
pub fn document_base64(media_type: impl Into<String>, data: impl Into<String>) -> Self {
Self::Document {
source: DocumentSource::Base64 {
media_type: media_type.into(),
data: data.into(),
},
}
}
#[must_use]
pub fn as_text(&self) -> Option<&str> {
match self {
Self::Text { text } => Some(text),
_ => None,
}
}
#[must_use]
pub fn as_tool_call(&self) -> Option<&ToolCall> {
match self {
Self::ToolCall(call) => Some(call),
_ => None,
}
}
#[must_use]
pub fn as_document(&self) -> Option<&DocumentSource> {
match self {
Self::Document { source } => Some(source),
_ => None,
}
}
#[must_use]
pub const fn kind(&self) -> &'static str {
match self {
Self::Text { .. } => "text",
Self::Image { .. } => "image",
Self::Document { .. } => "document",
Self::ToolCall(_) => "tool_call",
Self::ToolResult(_) => "tool_result",
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct Message {
pub role: Role,
pub content: Vec<ContentPart>,
}
impl Message {
#[must_use]
pub fn new(role: Role, content: Vec<ContentPart>) -> Self {
Self { role, content }
}
#[must_use]
pub fn user(text: impl Into<String>) -> Self {
Self::new(Role::User, vec![ContentPart::text(text)])
}
#[must_use]
pub fn assistant(text: impl Into<String>) -> Self {
Self::new(Role::Assistant, vec![ContentPart::text(text)])
}
#[must_use]
pub fn system(text: impl Into<String>) -> Self {
Self::new(Role::System, vec![ContentPart::text(text)])
}
#[must_use]
pub fn tool_result(result: ToolResult) -> Self {
Self::new(Role::Tool, vec![ContentPart::ToolResult(result)])
}
#[must_use]
pub fn with_part(mut self, part: ContentPart) -> Self {
self.content.push(part);
self
}
#[must_use]
pub fn text(&self) -> String {
let mut out = String::new();
for part in &self.content {
if let Some(text) = part.as_text() {
out.push_str(text);
}
}
out
}
#[must_use]
pub fn has_image(&self) -> bool {
self.content
.iter()
.any(|part| matches!(part, ContentPart::Image { .. }))
}
#[must_use]
pub fn has_document(&self) -> bool {
self.content
.iter()
.any(|part| matches!(part, ContentPart::Document { .. }))
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(tag = "kind", rename_all = "snake_case")]
#[non_exhaustive]
pub enum OutputSpec {
FreeText,
Json {
schema: serde_json::Value,
name: String,
strict: bool,
},
ToolCalls,
}
impl OutputSpec {
#[must_use]
pub fn json(name: impl Into<String>, schema: serde_json::Value) -> Self {
Self::Json {
schema,
name: name.into(),
strict: true,
}
}
#[must_use]
pub fn schema(&self) -> Option<&serde_json::Value> {
match self {
Self::Json { schema, .. } => Some(schema),
_ => None,
}
}
#[must_use]
pub const fn is_structured(&self) -> bool {
matches!(self, Self::Json { .. })
}
#[must_use]
pub const fn as_str(&self) -> &'static str {
match self {
Self::FreeText => "free_text",
Self::Json { .. } => "json",
Self::ToolCalls => "tool_calls",
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct ToolSpec {
pub name: String,
pub description: String,
pub parameters: serde_json::Value,
}
impl ToolSpec {
#[must_use]
pub fn new(
name: impl Into<String>,
description: impl Into<String>,
parameters: serde_json::Value,
) -> Self {
Self {
name: name.into(),
description: description.into(),
parameters,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash, Default, Serialize, Deserialize)]
#[serde(tag = "kind", rename_all = "snake_case")]
#[non_exhaustive]
pub enum ToolChoice {
#[default]
Auto,
None,
Required,
Named {
name: String,
},
}
impl ToolChoice {
#[must_use]
pub fn named(name: impl Into<String>) -> Self {
Self::Named { name: name.into() }
}
#[must_use]
pub const fn as_str(&self) -> &'static str {
match self {
Self::Auto => "auto",
Self::None => "none",
Self::Required => "required",
Self::Named { .. } => "named",
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum ReasoningEffort {
Minimal,
Low,
Medium,
High,
}
impl ReasoningEffort {
#[must_use]
pub const fn as_str(self) -> &'static str {
match self {
Self::Minimal => "minimal",
Self::Low => "low",
Self::Medium => "medium",
Self::High => "high",
}
}
}
#[derive(Debug, Clone, PartialEq)]
#[non_exhaustive]
pub struct Sampling {
pub temperature: Option<f32>,
pub seed: Option<u64>,
pub reasoning_effort: Option<ReasoningEffort>,
pub dropped: Vec<&'static str>,
}
#[derive(Debug, Clone, PartialEq, Eq, Hash, Default, Serialize, Deserialize)]
#[serde(tag = "kind", rename_all = "snake_case")]
#[non_exhaustive]
pub enum CacheHint {
#[default]
None,
System,
Prefix {
messages: usize,
},
}
#[derive(Debug, Clone, PartialEq, Eq, Default, Serialize)]
#[serde(transparent)]
pub struct RequestMetadata(BTreeMap<String, String>);
impl<'de> Deserialize<'de> for RequestMetadata {
fn deserialize<D: serde::Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
let raw = BTreeMap::<String, String>::deserialize(deserializer)?;
let mut metadata = Self::new();
for (key, value) in raw {
metadata
.insert(key, value)
.map_err(serde::de::Error::custom)?;
}
Ok(metadata)
}
}
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
#[non_exhaustive]
pub enum MetadataError {
#[error("metadata key is empty")]
EmptyKey,
#[error("metadata value for key {key} is {len} bytes, over the label limit")]
ValueTooLong {
key: String,
len: usize,
},
#[error("metadata value for key {key} is not a label")]
ValueNotALabel {
key: String,
},
}
impl RequestMetadata {
#[must_use]
pub fn new() -> Self {
Self::default()
}
pub fn insert(
&mut self,
key: impl Into<String>,
value: impl Into<String>,
) -> Result<(), MetadataError> {
let key = key.into();
let value = value.into();
if key.is_empty() {
return Err(MetadataError::EmptyKey);
}
if value.len() > MAX_METADATA_VALUE_LEN {
return Err(MetadataError::ValueTooLong {
key,
len: value.len(),
});
}
if value
.chars()
.any(|ch| ch.is_whitespace() || ch.is_control())
{
return Err(MetadataError::ValueNotALabel { key });
}
self.0.insert(key, value);
Ok(())
}
#[must_use]
pub fn get(&self, key: &str) -> Option<&str> {
self.0.get(key).map(String::as_str)
}
pub fn iter(&self) -> impl Iterator<Item = (&str, &str)> {
self.0.iter().map(|(k, v)| (k.as_str(), v.as_str()))
}
#[must_use]
pub fn len(&self) -> usize {
self.0.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.0.is_empty()
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct ModelRequest {
pub request_id: RequestId,
pub purpose: ModelPurpose,
pub messages: Vec<Message>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub system: Option<String>,
pub output: OutputSpec,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub tools: Vec<ToolSpec>,
#[serde(default)]
pub tool_choice: ToolChoice,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub max_output_tokens: Option<u32>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub temperature: Option<f32>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub reasoning_effort: Option<ReasoningEffort>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub seed: Option<u64>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub stop: Vec<String>,
#[serde(default, skip_serializing_if = "RequestMetadata::is_empty")]
pub metadata: RequestMetadata,
pub timeout: Duration,
#[serde(default)]
pub cache_hint: CacheHint,
}
impl ModelRequest {
#[must_use]
pub fn new(purpose: ModelPurpose) -> Self {
Self {
request_id: RequestId::new(),
purpose,
messages: Vec::new(),
system: None,
output: OutputSpec::FreeText,
tools: Vec::new(),
tool_choice: ToolChoice::Auto,
max_output_tokens: None,
temperature: None,
reasoning_effort: None,
seed: None,
stop: Vec::new(),
metadata: RequestMetadata::new(),
timeout: DEFAULT_TIMEOUT,
cache_hint: CacheHint::None,
}
}
#[must_use]
pub fn with_request_id(mut self, request_id: RequestId) -> Self {
self.request_id = request_id;
self
}
#[must_use]
pub fn with_system(mut self, system: impl Into<String>) -> Self {
self.system = Some(system.into());
self
}
#[must_use]
pub fn with_message(mut self, message: Message) -> Self {
self.messages.push(message);
self
}
#[must_use]
pub fn with_messages(mut self, messages: Vec<Message>) -> Self {
self.messages = messages;
self
}
#[must_use]
pub fn with_output(mut self, output: OutputSpec) -> Self {
self.output = output;
self
}
#[must_use]
pub fn with_tools(mut self, tools: Vec<ToolSpec>) -> Self {
self.tools = tools;
self
}
#[must_use]
pub fn with_tool_choice(mut self, tool_choice: ToolChoice) -> Self {
self.tool_choice = tool_choice;
self
}
#[must_use]
pub fn with_max_output_tokens(mut self, tokens: u32) -> Self {
self.max_output_tokens = Some(tokens);
self
}
#[must_use]
pub fn with_temperature(mut self, temperature: f32) -> Self {
self.temperature = Some(temperature);
self
}
#[must_use]
pub fn sampling_for(
&self,
capabilities: &crate::capabilities::ProviderCapabilities,
) -> Sampling {
let mut dropped = Vec::new();
let mut accept = |accepted: bool, name: &'static str| {
if !accepted {
dropped.push(name);
}
accepted
};
let temperature = self
.temperature
.filter(|_| accept(capabilities.temperature, "temperature"));
let seed = self.seed.filter(|_| accept(capabilities.seed, "seed"));
let reasoning_effort = self
.reasoning_effort
.filter(|_| accept(capabilities.reasoning_controls, "reasoning_effort"));
Sampling {
temperature,
seed,
reasoning_effort,
dropped,
}
}
#[must_use]
pub fn with_reasoning_effort(mut self, effort: ReasoningEffort) -> Self {
self.reasoning_effort = Some(effort);
self
}
#[must_use]
pub fn with_seed(mut self, seed: u64) -> Self {
self.seed = Some(seed);
self
}
#[must_use]
pub fn with_stop(mut self, stop: Vec<String>) -> Self {
self.stop = stop;
self
}
#[must_use]
pub fn with_timeout(mut self, timeout: Duration) -> Self {
self.timeout = timeout;
self
}
#[must_use]
pub fn with_cache_hint(mut self, cache_hint: CacheHint) -> Self {
self.cache_hint = cache_hint;
self
}
pub fn with_metadata(
mut self,
key: impl Into<String>,
value: impl Into<String>,
) -> Result<Self, MetadataError> {
self.metadata.insert(key, value)?;
Ok(self)
}
#[must_use]
pub fn requirements(&self) -> crate::capabilities::CapabilityRequirements {
let mut requirements = self.purpose.requirements();
if !self.tools.is_empty() || matches!(self.output, OutputSpec::ToolCalls) {
requirements.needs_tools = true;
}
if self.messages.iter().any(Message::has_image) {
requirements.needs_vision = true;
}
if self.messages.iter().any(Message::has_document) {
requirements.needs_documents = true;
}
requirements
}
}
#[cfg(test)]
mod tests {
#[test]
fn sampling_keeps_what_the_profile_takes_and_names_the_rest() {
use crate::capabilities::ProviderCapabilities;
let request = ModelRequest::new(ModelPurpose::Extract)
.with_temperature(0.0)
.with_seed(7)
.with_reasoning_effort(ReasoningEffort::Minimal);
let reasoning_model = ProviderCapabilities::minimal().with_reasoning_controls(true);
let sampling = request.sampling_for(&reasoning_model);
assert_eq!(sampling.temperature, None);
assert_eq!(sampling.seed, None);
assert_eq!(sampling.reasoning_effort, Some(ReasoningEffort::Minimal));
assert_eq!(sampling.dropped, vec!["temperature", "seed"]);
let chat_model = ProviderCapabilities::minimal()
.with_temperature(true)
.with_seed(true);
let sampling = request.sampling_for(&chat_model);
assert_eq!(sampling.temperature, Some(0.0));
assert_eq!(sampling.seed, Some(7));
assert_eq!(sampling.dropped, vec!["reasoning_effort"]);
let nothing_set = ModelRequest::new(ModelPurpose::Extract).sampling_for(&chat_model);
assert!(
nothing_set.dropped.is_empty(),
"an unset parameter is never dropped"
);
}
use super::*;
use serde_json::json;
#[test]
fn builder_defaults_are_conservative() {
let request = ModelRequest::new(ModelPurpose::Extract);
assert_eq!(request.timeout, DEFAULT_TIMEOUT);
assert_eq!(request.tool_choice, ToolChoice::Auto);
assert_eq!(request.cache_hint, CacheHint::None);
assert!(request.temperature.is_none());
assert!(request.metadata.is_empty());
assert_eq!(request.output.as_str(), "free_text");
}
#[test]
fn request_id_is_stable_when_pinned() {
let id = RequestId::nil();
let request = ModelRequest::new(ModelPurpose::Acknowledge).with_request_id(id);
let retried = request.clone().with_temperature(0.7);
assert_eq!(request.request_id, retried.request_id);
}
#[test]
fn metadata_rejects_anything_that_is_not_a_label() {
let mut metadata = RequestMetadata::new();
metadata.insert("turn", "0192f0aa-1b2c-7def").unwrap();
assert_eq!(metadata.get("turn"), Some("0192f0aa-1b2c-7def"));
assert_eq!(metadata.len(), 1);
assert!(matches!(
metadata.insert("", "x"),
Err(MetadataError::EmptyKey)
));
assert!(matches!(
metadata.insert("note", "hello world"),
Err(MetadataError::ValueNotALabel { .. })
));
assert!(matches!(
metadata.insert("note", "x".repeat(MAX_METADATA_VALUE_LEN + 1)),
Err(MetadataError::ValueTooLong { .. })
));
assert!(matches!(
metadata.insert("note", "line\nbreak"),
Err(MetadataError::ValueNotALabel { .. })
));
assert_eq!(metadata.len(), 1);
let collected: Vec<_> = metadata.iter().collect();
assert_eq!(collected, vec![("turn", "0192f0aa-1b2c-7def")]);
}
#[test]
fn requirements_follow_the_content_of_the_request() {
let plain = ModelRequest::new(ModelPurpose::Acknowledge);
let requirements = plain.requirements();
assert!(!requirements.needs_tools);
assert!(!requirements.needs_vision);
let with_tools = ModelRequest::new(ModelPurpose::Investigate)
.with_tools(vec![ToolSpec::new("case.get", "load a case", json!({}))])
.with_message(
Message::user("look at this").with_part(ContentPart::image_url("https://x.test/a")),
);
let requirements = with_tools.requirements();
assert!(requirements.needs_tools);
assert!(requirements.needs_vision);
assert!(!requirements.structured_output.is_empty());
}
#[test]
fn messages_flatten_their_text_parts() {
let message = Message::user("hello ").with_part(ContentPart::text("world"));
assert_eq!(message.text(), "hello world");
assert!(!message.has_image());
assert_eq!(message.role, Role::User);
assert_eq!(message.content[0].kind(), "text");
}
#[test]
fn a_document_states_its_media_type_on_both_forms() {
let inline = ContentPart::document_base64("application/pdf", "JVBERi0=");
let referenced = ContentPart::document_url("https://x.test/a", "application/pdf");
for part in [&inline, &referenced] {
assert_eq!(part.kind(), "document");
assert_eq!(
part.as_document().map(DocumentSource::media_type),
Some("application/pdf"),
"a document never leaves its media type to be guessed"
);
assert!(part.as_text().is_none());
}
let message = Message::new(Role::User, vec![inline]);
assert!(message.has_document());
assert!(!message.has_image(), "a document is not an image");
let request = ModelRequest::new(ModelPurpose::Acknowledge).with_message(message);
let requirements = request.requirements();
assert!(requirements.needs_documents);
assert!(!requirements.needs_vision);
}
#[test]
fn output_spec_exposes_its_schema() {
let spec = OutputSpec::json("user_turn_plan", json!({"type": "object"}));
assert!(spec.is_structured());
assert_eq!(spec.schema(), Some(&json!({"type": "object"})));
assert!(matches!(spec, OutputSpec::Json { strict: true, .. }));
assert_eq!(OutputSpec::FreeText.schema(), None);
}
#[test]
fn requests_round_trip_through_serde() {
let request = ModelRequest::new(ModelPurpose::Extract)
.with_request_id(RequestId::nil())
.with_system("be precise")
.with_message(Message::user("ciao"))
.with_message(Message::tool_result(ToolResult::ok("call_1", "{}")))
.with_output(OutputSpec::json("plan", json!({"type": "object"})))
.with_tools(vec![ToolSpec::new("case.get", "load", json!({}))])
.with_tool_choice(ToolChoice::named("case.get"))
.with_cache_hint(CacheHint::Prefix { messages: 1 })
.with_metadata("workflow", "trip")
.unwrap();
let json = serde_json::to_string(&request).unwrap();
let back: ModelRequest = serde_json::from_str(&json).unwrap();
assert_eq!(back, request);
assert!(!json.contains("\"stop\""), "empty vectors are skipped");
}
#[test]
fn unknown_fields_are_rejected() {
let json = json!({
"request_id": "00000000-0000-0000-0000-000000000000",
"purpose": "acknowledge",
"messages": [],
"output": {"kind": "free_text"},
"timeout": {"secs": 1, "nanos": 0},
"surprise": true
});
assert!(serde_json::from_value::<ModelRequest>(json).is_err());
}
#[test]
fn tool_calls_and_results_pair_by_id() {
let call = ToolCall::new("call_7", "case.get", json!({"id": "c1"}));
let result = ToolResult::ok(call.id.clone(), "{\"ok\":true}");
assert_eq!(result.call_id, call.id);
assert!(!result.is_error);
assert!(ToolResult::error("call_7", "boom").is_error);
let part = ContentPart::ToolCall(call);
assert_eq!(part.kind(), "tool_call");
assert!(part.as_tool_call().is_some());
assert!(part.as_text().is_none());
}
}