use serde::{Deserialize, Serialize};
use serde_json::Value;
use crate::{Completion, Message, Tool, ToolChoice, ToolResult, Toolbox};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum Effort {
Off,
Low,
Medium,
High,
#[serde(rename = "xhigh")]
XHigh,
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum ResponseFormat {
#[default]
Text,
Json,
Schema(Schema),
}
impl ResponseFormat {
pub(crate) fn is_text(&self) -> bool {
matches!(self, Self::Text)
}
}
impl From<Schema> for ResponseFormat {
fn from(schema: Schema) -> Self {
Self::Schema(schema)
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct Schema {
pub name: String,
pub schema: Value,
#[serde(default, skip_serializing_if = "std::ops::Not::not")]
pub strict: bool,
}
impl Schema {
pub fn new(name: impl Into<String>, schema: Value) -> Self {
Self {
name: name.into(),
schema,
strict: false,
}
}
#[cfg(feature = "schemars")]
pub fn of<T: schemars::JsonSchema>() -> Self {
Self::new(T::schema_name(), crate::tool::schema_of::<T>())
}
pub fn strict(mut self, strict: bool) -> Self {
self.strict = strict;
self
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct Request {
pub model: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub system: Option<String>,
#[serde(default)]
pub messages: Vec<Message>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub tools: Vec<Tool>,
#[serde(default, skip_serializing_if = "ToolChoice::is_auto")]
pub tool_choice: ToolChoice,
#[serde(default, skip_serializing_if = "ResponseFormat::is_text")]
pub response_format: ResponseFormat,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub max_tokens: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub temperature: Option<f32>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub reasoning: Option<Effort>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub include_usage: Option<bool>,
#[serde(default, skip_serializing_if = "std::ops::Not::not")]
pub send_reasoning: bool,
}
impl Request {
pub fn new(model: impl Into<String>) -> Self {
Self {
model: model.into(),
system: None,
messages: Vec::new(),
tools: Vec::new(),
tool_choice: ToolChoice::Auto,
response_format: ResponseFormat::Text,
max_tokens: None,
temperature: None,
reasoning: None,
include_usage: None,
send_reasoning: false,
}
}
pub fn system(mut self, prompt: impl Into<String>) -> Self {
self.system = Some(prompt.into());
self
}
pub fn user(self, text: impl Into<String>) -> Self {
self.message(Message::user(text))
}
pub fn message(mut self, message: Message) -> Self {
self.messages.push(message);
self
}
pub fn assistant(self, done: Completion) -> Self {
self.message(Message::from(done))
}
pub fn tool_result(self, call_id: impl Into<String>, content: impl Into<String>) -> Self {
self.message(Message::tool_result(ToolResult::new(call_id, content)))
}
pub fn tool_results(mut self, results: impl IntoIterator<Item = ToolResult>) -> Self {
self.messages
.extend(results.into_iter().map(Message::tool_result));
self
}
pub fn tool(mut self, tool: Tool) -> Self {
self.tools.push(tool);
self
}
pub fn tools(mut self, toolbox: &(impl Toolbox + ?Sized)) -> Self {
self.tools.extend(toolbox.tools());
self
}
pub fn tool_choice(mut self, choice: ToolChoice) -> Self {
self.tool_choice = choice;
self
}
pub fn response_format(mut self, format: impl Into<ResponseFormat>) -> Self {
self.response_format = format.into();
self
}
pub fn reasoning(mut self, effort: Effort) -> Self {
self.reasoning = Some(effort);
self
}
pub fn max_tokens(mut self, tokens: u64) -> Self {
self.max_tokens = Some(tokens);
self
}
pub fn temperature(mut self, temperature: f32) -> Self {
self.temperature = Some(temperature);
self
}
pub fn include_usage(mut self, include: bool) -> Self {
self.include_usage = Some(include);
self
}
pub fn send_reasoning(mut self, send: bool) -> Self {
self.send_reasoning = send;
self
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{FinishReason, Part, Role, ToolCall};
#[test]
fn a_tool_round_trip_reads_as_the_conversation() {
let mut done = Completion::new(FinishReason::ToolCalls);
done.calls = vec![
ToolCall::new("call-a", "lookup", r#"{"value":1}"#),
ToolCall::new("call-b", "lookup", r#"{"value":2}"#),
];
let request = Request::new("m")
.system("Be brief.")
.tool(Tool::new("lookup", "Look up a value."))
.user("Use lookup")
.assistant(done)
.tool_results([
ToolResult::new("call-a", "42"),
ToolResult::new("call-b", "43"),
]);
let roles: Vec<_> = request.messages.iter().map(|m| m.role).collect();
assert_eq!(roles, [Role::User, Role::Assistant, Role::Tool, Role::Tool]);
assert_eq!(request.messages[1].parts.len(), 2);
assert_eq!(
request.messages[3].parts,
[Part::ToolResult(ToolResult::new("call-b", "43"))]
);
assert_eq!(request.system.as_deref(), Some("Be brief."));
}
#[test]
fn nothing_is_set_until_asked() {
let request = Request::new("m");
assert_eq!(request.max_tokens, None);
assert_eq!(request.temperature, None);
assert_eq!(request.reasoning, None);
assert_eq!(request.include_usage, None);
assert_eq!(request.tool_choice, ToolChoice::Auto);
assert_eq!(request.response_format, ResponseFormat::Text);
assert!(!request.send_reasoning);
let request = request
.max_tokens(16)
.temperature(0.5)
.reasoning(Effort::Low)
.include_usage(false)
.tool_choice(ToolChoice::tool("lookup"))
.response_format(ResponseFormat::Json)
.send_reasoning(true);
assert_eq!(request.max_tokens, Some(16));
assert_eq!(request.temperature, Some(0.5));
assert_eq!(request.reasoning, Some(Effort::Low));
assert_eq!(request.include_usage, Some(false));
assert_eq!(request.tool_choice, ToolChoice::Tool("lookup".into()));
assert_eq!(request.response_format, ResponseFormat::Json);
assert!(request.send_reasoning);
}
#[test]
fn a_schema_is_a_response_format_and_strict_only_when_asked() {
let schema = Schema::new("weather", serde_json::json!({"type": "object"}));
assert!(!schema.strict);
let request = Request::new("m").response_format(schema.strict(true));
let ResponseFormat::Schema(sent) = request.response_format else {
panic!("not a schema");
};
assert_eq!(sent.name, "weather");
assert!(sent.strict);
}
#[cfg(feature = "schemars")]
#[test]
fn a_schema_of_a_type_is_named_after_it() {
#[derive(schemars::JsonSchema)]
#[allow(dead_code)]
struct Weather {
city: String,
celsius: f64,
}
let schema = Schema::of::<Weather>();
assert_eq!(schema.name, "Weather");
assert_eq!(schema.schema["type"], "object");
assert_eq!(schema.schema["description"], "The weather in a city.");
assert_eq!(
schema.schema["properties"]["city"]["description"],
"The city."
);
assert_eq!(schema.schema.get("$schema"), None);
assert_eq!(schema.schema.get("title"), None);
}
}