use serde::Deserialize;
use serde::Serialize;
use tokio_util::sync::CancellationToken;
use super::prompt::Prompt;
use super::tool::ToolDefinition;
use crate::json::JsonValue;
use crate::shared::Headers;
use crate::shared::ProviderOptions;
use crate::shared::ToolName;
#[derive(Debug, Clone, Default)]
pub struct CallOptions {
pub prompt: Prompt,
pub max_output_tokens: Option<u32>,
pub temperature: Option<f64>,
pub top_p: Option<f64>,
pub top_k: Option<u32>,
pub presence_penalty: Option<f64>,
pub frequency_penalty: Option<f64>,
pub stop_sequences: Option<Vec<String>>,
pub seed: Option<u64>,
pub response_format: Option<ResponseFormat>,
pub tools: Vec<ToolDefinition>,
pub tool_choice: Option<ToolChoice>,
pub include_raw_chunks: bool,
pub reasoning: ReasoningEffort,
pub headers: Headers,
pub provider_options: ProviderOptions,
pub cancellation: CancellationToken,
}
impl CallOptions {
#[must_use]
pub fn new(prompt: Prompt) -> Self {
Self {
prompt,
..Self::default()
}
}
#[must_use]
pub fn to_recordable(&self) -> CallOptionsRecord {
CallOptionsRecord {
prompt: self.prompt.clone(),
max_output_tokens: self.max_output_tokens,
temperature: self.temperature,
top_p: self.top_p,
top_k: self.top_k,
presence_penalty: self.presence_penalty,
frequency_penalty: self.frequency_penalty,
stop_sequences: self.stop_sequences.clone(),
seed: self.seed,
response_format: self.response_format.clone(),
tools: self.tools.clone(),
tool_choice: self.tool_choice.clone(),
include_raw_chunks: self.include_raw_chunks,
reasoning: self.reasoning,
headers: self.headers.clone(),
provider_options: self.provider_options.clone(),
}
}
#[must_use]
pub fn provider_options_for(&self, provider_key: &str) -> Option<&crate::json::JsonObject> {
self.provider_options.get(provider_key)
}
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
pub struct CallOptionsRecord {
pub prompt: Prompt,
#[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<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub top_p: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub top_k: Option<u32>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub presence_penalty: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub frequency_penalty: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub stop_sequences: Option<Vec<String>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub seed: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub response_format: Option<ResponseFormat>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub tools: Vec<ToolDefinition>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub tool_choice: Option<ToolChoice>,
#[serde(default, skip_serializing_if = "std::ops::Not::not")]
pub include_raw_chunks: bool,
#[serde(default)]
pub reasoning: ReasoningEffort,
#[serde(default, skip_serializing_if = "Headers::is_empty")]
pub headers: Headers,
#[serde(default, skip_serializing_if = "ProviderOptions::is_empty")]
pub provider_options: ProviderOptions,
}
impl From<CallOptionsRecord> for CallOptions {
fn from(record: CallOptionsRecord) -> Self {
Self {
prompt: record.prompt,
max_output_tokens: record.max_output_tokens,
temperature: record.temperature,
top_p: record.top_p,
top_k: record.top_k,
presence_penalty: record.presence_penalty,
frequency_penalty: record.frequency_penalty,
stop_sequences: record.stop_sequences,
seed: record.seed,
response_format: record.response_format,
tools: record.tools,
tool_choice: record.tool_choice,
include_raw_chunks: record.include_raw_chunks,
reasoning: record.reasoning,
headers: record.headers,
provider_options: record.provider_options,
cancellation: CancellationToken::new(),
}
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "lowercase")]
pub enum ResponseFormat {
Text,
Json {
#[serde(default, skip_serializing_if = "Option::is_none")]
schema: Option<JsonValue>,
#[serde(default, skip_serializing_if = "Option::is_none")]
name: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
description: Option<String>,
},
}
impl ResponseFormat {
#[must_use]
pub fn json(schema: JsonValue) -> Self {
Self::Json {
schema: Some(schema),
name: None,
description: None,
}
}
#[must_use]
pub fn json_unconstrained() -> Self {
Self::Json {
schema: None,
name: None,
description: None,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "lowercase")]
pub enum ToolChoice {
Auto,
None,
Required,
Tool {
tool_name: ToolName,
},
}
impl ToolChoice {
#[must_use]
pub fn tool(name: impl Into<ToolName>) -> Self {
Self::Tool {
tool_name: name.into(),
}
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "kebab-case")]
pub enum ReasoningEffort {
#[default]
ProviderDefault,
None,
Minimal,
Low,
Medium,
High,
#[serde(rename = "xhigh")]
XHigh,
}
impl ReasoningEffort {
#[must_use]
pub fn as_str(self) -> &'static str {
match self {
Self::ProviderDefault => "provider-default",
Self::None => "none",
Self::Minimal => "minimal",
Self::Low => "low",
Self::Medium => "medium",
Self::High => "high",
Self::XHigh => "xhigh",
}
}
}
impl std::fmt::Display for ReasoningEffort {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.as_str())
}
}