ferrin_spec/language_model/
call_options.rs1use serde::Deserialize;
4use serde::Serialize;
5use tokio_util::sync::CancellationToken;
6
7use super::prompt::Prompt;
8use super::tool::ToolDefinition;
9use crate::json::JsonValue;
10use crate::shared::Headers;
11use crate::shared::ProviderOptions;
12use crate::shared::ToolName;
13
14#[derive(Debug, Clone, Default)]
21pub struct CallOptions {
22 pub prompt: Prompt,
24 pub max_output_tokens: Option<u32>,
26 pub temperature: Option<f64>,
28 pub top_p: Option<f64>,
30 pub top_k: Option<u32>,
32 pub presence_penalty: Option<f64>,
34 pub frequency_penalty: Option<f64>,
36 pub stop_sequences: Option<Vec<String>>,
38 pub seed: Option<u64>,
40 pub response_format: Option<ResponseFormat>,
42 pub tools: Vec<ToolDefinition>,
44 pub tool_choice: Option<ToolChoice>,
46 pub include_raw_chunks: bool,
48 pub reasoning: ReasoningEffort,
50 pub headers: Headers,
52 pub provider_options: ProviderOptions,
54 pub cancellation: CancellationToken,
56}
57
58impl CallOptions {
59 #[must_use]
61 pub fn new(prompt: Prompt) -> Self {
62 Self {
63 prompt,
64 ..Self::default()
65 }
66 }
67
68 #[must_use]
73 pub fn to_recordable(&self) -> CallOptionsRecord {
74 CallOptionsRecord {
75 prompt: self.prompt.clone(),
76 max_output_tokens: self.max_output_tokens,
77 temperature: self.temperature,
78 top_p: self.top_p,
79 top_k: self.top_k,
80 presence_penalty: self.presence_penalty,
81 frequency_penalty: self.frequency_penalty,
82 stop_sequences: self.stop_sequences.clone(),
83 seed: self.seed,
84 response_format: self.response_format.clone(),
85 tools: self.tools.clone(),
86 tool_choice: self.tool_choice.clone(),
87 include_raw_chunks: self.include_raw_chunks,
88 reasoning: self.reasoning,
89 headers: self.headers.clone(),
90 provider_options: self.provider_options.clone(),
91 }
92 }
93
94 #[must_use]
96 pub fn provider_options_for(&self, provider_key: &str) -> Option<&crate::json::JsonObject> {
97 self.provider_options.get(provider_key)
98 }
99}
100
101#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
103pub struct CallOptionsRecord {
104 pub prompt: Prompt,
106 #[serde(default, skip_serializing_if = "Option::is_none")]
108 pub max_output_tokens: Option<u32>,
109 #[serde(default, skip_serializing_if = "Option::is_none")]
111 pub temperature: Option<f64>,
112 #[serde(default, skip_serializing_if = "Option::is_none")]
114 pub top_p: Option<f64>,
115 #[serde(default, skip_serializing_if = "Option::is_none")]
117 pub top_k: Option<u32>,
118 #[serde(default, skip_serializing_if = "Option::is_none")]
120 pub presence_penalty: Option<f64>,
121 #[serde(default, skip_serializing_if = "Option::is_none")]
123 pub frequency_penalty: Option<f64>,
124 #[serde(default, skip_serializing_if = "Option::is_none")]
126 pub stop_sequences: Option<Vec<String>>,
127 #[serde(default, skip_serializing_if = "Option::is_none")]
129 pub seed: Option<u64>,
130 #[serde(default, skip_serializing_if = "Option::is_none")]
132 pub response_format: Option<ResponseFormat>,
133 #[serde(default, skip_serializing_if = "Vec::is_empty")]
135 pub tools: Vec<ToolDefinition>,
136 #[serde(default, skip_serializing_if = "Option::is_none")]
138 pub tool_choice: Option<ToolChoice>,
139 #[serde(default, skip_serializing_if = "std::ops::Not::not")]
141 pub include_raw_chunks: bool,
142 #[serde(default)]
144 pub reasoning: ReasoningEffort,
145 #[serde(default, skip_serializing_if = "Headers::is_empty")]
147 pub headers: Headers,
148 #[serde(default, skip_serializing_if = "ProviderOptions::is_empty")]
150 pub provider_options: ProviderOptions,
151}
152
153impl From<CallOptionsRecord> for CallOptions {
154 fn from(record: CallOptionsRecord) -> Self {
155 Self {
156 prompt: record.prompt,
157 max_output_tokens: record.max_output_tokens,
158 temperature: record.temperature,
159 top_p: record.top_p,
160 top_k: record.top_k,
161 presence_penalty: record.presence_penalty,
162 frequency_penalty: record.frequency_penalty,
163 stop_sequences: record.stop_sequences,
164 seed: record.seed,
165 response_format: record.response_format,
166 tools: record.tools,
167 tool_choice: record.tool_choice,
168 include_raw_chunks: record.include_raw_chunks,
169 reasoning: record.reasoning,
170 headers: record.headers,
171 provider_options: record.provider_options,
172 cancellation: CancellationToken::new(),
173 }
174 }
175}
176
177#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
179#[serde(tag = "type", rename_all = "lowercase")]
180pub enum ResponseFormat {
181 Text,
183 Json {
185 #[serde(default, skip_serializing_if = "Option::is_none")]
187 schema: Option<JsonValue>,
188 #[serde(default, skip_serializing_if = "Option::is_none")]
190 name: Option<String>,
191 #[serde(default, skip_serializing_if = "Option::is_none")]
193 description: Option<String>,
194 },
195}
196
197impl ResponseFormat {
198 #[must_use]
200 pub fn json(schema: JsonValue) -> Self {
201 Self::Json {
202 schema: Some(schema),
203 name: None,
204 description: None,
205 }
206 }
207
208 #[must_use]
210 pub fn json_unconstrained() -> Self {
211 Self::Json {
212 schema: None,
213 name: None,
214 description: None,
215 }
216 }
217}
218
219#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
221#[serde(tag = "type", rename_all = "lowercase")]
222pub enum ToolChoice {
223 Auto,
225 None,
227 Required,
229 Tool {
231 tool_name: ToolName,
233 },
234}
235
236impl ToolChoice {
237 #[must_use]
239 pub fn tool(name: impl Into<ToolName>) -> Self {
240 Self::Tool {
241 tool_name: name.into(),
242 }
243 }
244}
245
246#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Hash, Serialize, Deserialize)]
248#[serde(rename_all = "kebab-case")]
249pub enum ReasoningEffort {
250 #[default]
252 ProviderDefault,
253 None,
255 Minimal,
257 Low,
259 Medium,
261 High,
263 #[serde(rename = "xhigh")]
265 XHigh,
266}
267
268impl ReasoningEffort {
269 #[must_use]
271 pub fn as_str(self) -> &'static str {
272 match self {
273 Self::ProviderDefault => "provider-default",
274 Self::None => "none",
275 Self::Minimal => "minimal",
276 Self::Low => "low",
277 Self::Medium => "medium",
278 Self::High => "high",
279 Self::XHigh => "xhigh",
280 }
281 }
282}
283
284impl std::fmt::Display for ReasoningEffort {
285 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
286 f.write_str(self.as_str())
287 }
288}