1use std::collections::BTreeMap;
2use std::time::Duration;
3
4use serde::{Deserialize, Serialize};
5use serde_json::{Map, Value};
6
7#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
8#[serde(deny_unknown_fields)]
9pub struct ModelSettings {
10 #[serde(
11 default,
12 with = "temperature_option",
13 skip_serializing_if = "Option::is_none"
14 )]
15 pub temperature: Option<f64>,
16 #[serde(
17 default,
18 with = "top_p_option",
19 skip_serializing_if = "Option::is_none"
20 )]
21 pub top_p: Option<f64>,
22 #[serde(
23 default,
24 alias = "max_output_tokens",
25 with = "positive_u32_option",
26 skip_serializing_if = "Option::is_none"
27 )]
28 pub max_tokens: Option<u32>,
29 #[serde(default, skip_serializing_if = "Option::is_none")]
30 pub tool_choice: Option<ToolChoice>,
31 #[serde(default, skip_serializing_if = "Option::is_none")]
32 pub parallel_tool_calls: Option<bool>,
33 #[serde(
34 default,
35 with = "reasoning_option",
36 skip_serializing_if = "reasoning_option::is_none_or_empty"
37 )]
38 pub reasoning: Option<Value>,
39 #[serde(default, skip_serializing_if = "Option::is_none")]
40 pub response_format: Option<ResponseFormat>,
41 #[serde(
42 default,
43 rename = "timeout_seconds",
44 with = "duration_seconds_option",
45 skip_serializing_if = "Option::is_none"
46 )]
47 pub timeout: Option<Duration>,
48 #[serde(default, skip_serializing_if = "Option::is_none")]
49 pub retry: Option<RetrySettings>,
50 #[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
51 pub extra_headers: BTreeMap<String, String>,
52 #[serde(default, skip_serializing_if = "Map::is_empty")]
53 pub extra_body: Map<String, Value>,
54 #[serde(default, skip_serializing_if = "Map::is_empty")]
55 pub extra_args: Map<String, Value>,
56}
57
58impl ModelSettings {
59 pub fn builder() -> ModelSettingsBuilder {
60 ModelSettingsBuilder::default()
61 }
62
63 pub fn merge(&self, override_settings: &ModelSettings) -> ModelSettings {
64 let mut merged = self.clone();
65 if override_settings.temperature.is_some() {
66 merged.temperature = override_settings.temperature;
67 }
68 if override_settings.top_p.is_some() {
69 merged.top_p = override_settings.top_p;
70 }
71 if override_settings.max_tokens.is_some() {
72 merged.max_tokens = override_settings.max_tokens;
73 }
74 if override_settings.tool_choice.is_some() {
75 merged.tool_choice = override_settings.tool_choice.clone();
76 }
77 if override_settings.parallel_tool_calls.is_some() {
78 merged.parallel_tool_calls = override_settings.parallel_tool_calls;
79 }
80 if override_settings.reasoning.is_some() {
81 merged.reasoning = override_settings.reasoning.clone();
82 }
83 if override_settings.response_format.is_some() {
84 merged.response_format = override_settings.response_format.clone();
85 }
86 if override_settings.timeout.is_some() {
87 merged.timeout = override_settings.timeout;
88 }
89 if override_settings.retry.is_some() {
90 merged.retry = override_settings.retry.clone();
91 }
92 merged
93 .extra_headers
94 .extend(override_settings.extra_headers.clone());
95 merged
96 .extra_body
97 .extend(override_settings.extra_body.clone());
98 merged
99 .extra_args
100 .extend(override_settings.extra_args.clone());
101 merged
102 }
103
104 pub fn to_value(&self) -> Value {
105 serde_json::to_value(self).unwrap_or(Value::Null)
106 }
107
108 pub fn validate(&self) -> Result<(), String> {
109 validate_finite_min("temperature", self.temperature, 0.0, false)?;
110 validate_finite_range("top_p", self.top_p, 0.0, 1.0)?;
111 if self.max_tokens == Some(0) {
112 return Err("max_tokens must be greater than zero".to_string());
113 }
114 if self.timeout.is_some_and(|timeout| timeout.is_zero()) {
115 return Err("timeout_seconds must be greater than zero".to_string());
116 }
117 if let Some(retry) = self.retry.as_ref() {
118 retry.validate()?;
119 }
120 if self.tool_choice.as_ref().is_some_and(
121 |choice| matches!(choice, ToolChoice::Tool(name) if name.trim().is_empty()),
122 ) {
123 return Err("named tool_choice requires a non-empty function name".to_string());
124 }
125 if self
126 .reasoning
127 .as_ref()
128 .is_some_and(|value| !value.is_object())
129 {
130 return Err("reasoning must be an object".to_string());
131 }
132 Ok(())
133 }
134}
135
136#[derive(Debug, Clone, Default)]
137pub struct ModelSettingsBuilder {
138 settings: ModelSettings,
139}
140
141impl ModelSettingsBuilder {
142 pub fn temperature(mut self, temperature: f64) -> Self {
143 self.settings.temperature = Some(temperature);
144 self
145 }
146
147 pub fn top_p(mut self, top_p: f64) -> Self {
148 self.settings.top_p = Some(top_p);
149 self
150 }
151
152 pub fn max_tokens(mut self, max_tokens: u32) -> Self {
153 self.settings.max_tokens = Some(max_tokens);
154 self
155 }
156
157 pub fn max_output_tokens(self, max_output_tokens: u32) -> Self {
158 self.max_tokens(max_output_tokens)
159 }
160
161 pub fn tool_choice(mut self, tool_choice: ToolChoice) -> Self {
162 self.settings.tool_choice = Some(tool_choice);
163 self
164 }
165
166 pub fn parallel_tool_calls(mut self, parallel_tool_calls: bool) -> Self {
167 self.settings.parallel_tool_calls = Some(parallel_tool_calls);
168 self
169 }
170
171 pub fn reasoning(mut self, reasoning: Value) -> Self {
172 self.settings.reasoning =
173 (!reasoning.as_object().is_some_and(Map::is_empty)).then_some(reasoning);
174 self
175 }
176
177 pub fn response_format(mut self, response_format: ResponseFormat) -> Self {
178 self.settings.response_format = Some(response_format);
179 self
180 }
181
182 pub fn timeout(mut self, timeout: Duration) -> Self {
183 self.settings.timeout = Some(timeout);
184 self
185 }
186
187 pub fn retry(mut self, retry: RetrySettings) -> Self {
188 self.settings.retry = Some(retry);
189 self
190 }
191
192 pub fn extra_header(mut self, key: impl Into<String>, value: impl Into<String>) -> Self {
193 self.settings.extra_headers.insert(key.into(), value.into());
194 self
195 }
196
197 pub fn extra_body(mut self, key: impl Into<String>, value: Value) -> Self {
198 self.settings.extra_body.insert(key.into(), value);
199 self
200 }
201
202 pub fn extra_arg(mut self, key: impl Into<String>, value: Value) -> Self {
203 self.settings.extra_args.insert(key.into(), value);
204 self
205 }
206
207 pub fn build(self) -> ModelSettings {
208 self.settings
209 }
210}
211
212#[derive(Debug, Clone, PartialEq, Eq)]
213pub enum ToolChoice {
214 Auto,
215 None,
216 Required,
217 Tool(String),
218}
219
220impl Serialize for ToolChoice {
221 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
222 where
223 S: serde::Serializer,
224 {
225 match self {
226 Self::Auto => serializer.serialize_str("auto"),
227 Self::None => serializer.serialize_str("none"),
228 Self::Required => serializer.serialize_str("required"),
229 Self::Tool(name) => {
230 if name.trim().is_empty() {
231 return Err(serde::ser::Error::custom(
232 "named tool_choice requires a non-empty function name",
233 ));
234 }
235 serde_json::json!({
236 "type": "function",
237 "function": {"name": name},
238 })
239 .serialize(serializer)
240 }
241 }
242 }
243}
244
245impl<'de> Deserialize<'de> for ToolChoice {
246 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
247 where
248 D: serde::Deserializer<'de>,
249 {
250 let value = Value::deserialize(deserializer)?;
251 if let Some(mode) = value.as_str() {
252 return match mode {
253 "auto" => Ok(Self::Auto),
254 "none" => Ok(Self::None),
255 "required" => Ok(Self::Required),
256 _ => Err(serde::de::Error::custom(format!(
257 "unknown tool_choice mode: {mode}"
258 ))),
259 };
260 }
261 let object = value.as_object().ok_or_else(|| {
262 serde::de::Error::custom("tool_choice must be a mode or function object")
263 })?;
264 if object.len() != 2 || object.get("type") != Some(&Value::String("function".to_string())) {
265 return Err(serde::de::Error::custom(
266 "named tool_choice must use the standard function object",
267 ));
268 }
269 let function = object
270 .get("function")
271 .and_then(Value::as_object)
272 .filter(|function| function.len() == 1)
273 .ok_or_else(|| {
274 serde::de::Error::custom("tool_choice function must contain only name")
275 })?;
276 let name = function
277 .get("name")
278 .and_then(Value::as_str)
279 .filter(|name| !name.trim().is_empty())
280 .ok_or_else(|| serde::de::Error::custom("tool_choice function name cannot be empty"))?;
281 Ok(Self::Tool(name.to_string()))
282 }
283}
284
285#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
286#[serde(tag = "type", rename_all = "snake_case")]
287pub enum ResponseFormat {
288 Text,
289 JsonObject,
290 JsonSchema { json_schema: Map<String, Value> },
291}
292
293impl<'de> Deserialize<'de> for ResponseFormat {
294 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
295 where
296 D: serde::Deserializer<'de>,
297 {
298 let value = Value::deserialize(deserializer)?;
299 let object = value
300 .as_object()
301 .ok_or_else(|| serde::de::Error::custom("response_format must be an object"))?;
302 let format_type = object
303 .get("type")
304 .and_then(Value::as_str)
305 .ok_or_else(|| serde::de::Error::custom("response_format.type must be a string"))?;
306 match format_type {
307 "text" if object.len() == 1 => Ok(Self::Text),
308 "json_object" if object.len() == 1 => Ok(Self::JsonObject),
309 "json_schema" if object.len() == 2 => {
310 let json_schema = object
311 .get("json_schema")
312 .and_then(Value::as_object)
313 .cloned()
314 .ok_or_else(|| {
315 serde::de::Error::custom("json_schema response_format requires an object")
316 })?;
317 Ok(Self::JsonSchema { json_schema })
318 }
319 _ => Err(serde::de::Error::custom(
320 "invalid or unsupported response_format wire shape",
321 )),
322 }
323 }
324}
325
326#[derive(Debug, Clone, PartialEq, Serialize)]
327pub struct RetrySettings {
328 pub max_attempts: u32,
329 pub backoff_seconds: f64,
330}
331
332impl RetrySettings {
333 pub fn new(max_attempts: u32) -> Self {
334 Self {
335 max_attempts,
336 backoff_seconds: 2.0,
337 }
338 }
339
340 pub fn with_backoff_seconds(mut self, backoff_seconds: f64) -> Self {
341 self.backoff_seconds = backoff_seconds;
342 self
343 }
344
345 pub fn validate(&self) -> Result<(), String> {
346 if self.max_attempts == 0 {
347 return Err("retry.max_attempts must be greater than zero".to_string());
348 }
349 if !self.backoff_seconds.is_finite() || self.backoff_seconds < 0.0 {
350 return Err("retry.backoff_seconds must be a finite non-negative number".to_string());
351 }
352 Ok(())
353 }
354}
355
356impl Default for RetrySettings {
357 fn default() -> Self {
358 Self {
359 max_attempts: 3,
360 backoff_seconds: 2.0,
361 }
362 }
363}
364
365#[derive(Deserialize)]
366#[serde(default, deny_unknown_fields)]
367struct RetrySettingsWire {
368 max_attempts: u32,
369 backoff_seconds: f64,
370}
371
372impl Default for RetrySettingsWire {
373 fn default() -> Self {
374 let defaults = RetrySettings::default();
375 Self {
376 max_attempts: defaults.max_attempts,
377 backoff_seconds: defaults.backoff_seconds,
378 }
379 }
380}
381
382impl<'de> Deserialize<'de> for RetrySettings {
383 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
384 where
385 D: serde::Deserializer<'de>,
386 {
387 let wire = RetrySettingsWire::deserialize(deserializer)?;
388 let settings = Self {
389 max_attempts: wire.max_attempts,
390 backoff_seconds: wire.backoff_seconds,
391 };
392 settings.validate().map_err(serde::de::Error::custom)?;
393 Ok(settings)
394 }
395}
396
397pub type RetryPolicy = RetrySettings;
398
399mod duration_seconds_option {
400 use std::time::Duration;
401
402 use serde::{Deserialize, Deserializer, Serializer};
403
404 pub fn serialize<S>(value: &Option<Duration>, serializer: S) -> Result<S::Ok, S::Error>
405 where
406 S: Serializer,
407 {
408 match value {
409 Some(value) => serializer.serialize_some(&value.as_secs_f64()),
410 None => serializer.serialize_none(),
411 }
412 }
413
414 pub fn deserialize<'de, D>(deserializer: D) -> Result<Option<Duration>, D::Error>
415 where
416 D: Deserializer<'de>,
417 {
418 let seconds = Option::<f64>::deserialize(deserializer)?;
419 seconds
420 .map(|seconds| {
421 if seconds.is_finite() && seconds > 0.0 {
422 Ok(Duration::from_secs_f64(seconds))
423 } else {
424 Err(serde::de::Error::custom(
425 "timeout_seconds must be a finite positive number",
426 ))
427 }
428 })
429 .transpose()
430 }
431}
432
433macro_rules! finite_option_module {
434 ($module:ident, $validator:expr) => {
435 mod $module {
436 use serde::{Deserialize, Deserializer, Serialize, Serializer};
437
438 pub fn serialize<S>(value: &Option<f64>, serializer: S) -> Result<S::Ok, S::Error>
439 where
440 S: Serializer,
441 {
442 value.serialize(serializer)
443 }
444
445 pub fn deserialize<'de, D>(deserializer: D) -> Result<Option<f64>, D::Error>
446 where
447 D: Deserializer<'de>,
448 {
449 let value = Option::<f64>::deserialize(deserializer)?;
450 if value.is_some_and(|value| !($validator)(value)) {
451 return Err(serde::de::Error::custom(concat!(
452 stringify!($module),
453 " is outside its valid range"
454 )));
455 }
456 Ok(value)
457 }
458 }
459 };
460}
461
462finite_option_module!(temperature_option, |value: f64| value.is_finite()
463 && value >= 0.0);
464finite_option_module!(top_p_option, |value: f64| value.is_finite()
465 && (0.0..=1.0).contains(&value));
466
467mod positive_u32_option {
468 use serde::{Deserialize, Deserializer, Serialize, Serializer};
469
470 pub fn serialize<S>(value: &Option<u32>, serializer: S) -> Result<S::Ok, S::Error>
471 where
472 S: Serializer,
473 {
474 value.serialize(serializer)
475 }
476
477 pub fn deserialize<'de, D>(deserializer: D) -> Result<Option<u32>, D::Error>
478 where
479 D: Deserializer<'de>,
480 {
481 let value = Option::<u32>::deserialize(deserializer)?;
482 if value == Some(0) {
483 return Err(serde::de::Error::custom(
484 "max_tokens must be greater than zero",
485 ));
486 }
487 Ok(value)
488 }
489}
490
491mod reasoning_option {
492 use serde::{Deserialize, Deserializer, Serialize, Serializer};
493 use serde_json::{Map, Value};
494
495 pub fn is_none_or_empty(value: &Option<Value>) -> bool {
496 value
497 .as_ref()
498 .is_none_or(|value| value.as_object().is_some_and(Map::is_empty))
499 }
500
501 pub fn serialize<S>(value: &Option<Value>, serializer: S) -> Result<S::Ok, S::Error>
502 where
503 S: Serializer,
504 {
505 value.serialize(serializer)
506 }
507
508 pub fn deserialize<'de, D>(deserializer: D) -> Result<Option<Value>, D::Error>
509 where
510 D: Deserializer<'de>,
511 {
512 let value = Option::<Value>::deserialize(deserializer)?;
513 match value {
514 Some(Value::Object(object)) if object.is_empty() => Ok(None),
515 Some(value @ Value::Object(_)) => Ok(Some(value)),
516 Some(_) => Err(serde::de::Error::custom("reasoning must be an object")),
517 None => Ok(None),
518 }
519 }
520}
521
522fn validate_finite_min(
523 name: &str,
524 value: Option<f64>,
525 minimum: f64,
526 exclusive: bool,
527) -> Result<(), String> {
528 let Some(value) = value else {
529 return Ok(());
530 };
531 if !value.is_finite() || (exclusive && value <= minimum) || (!exclusive && value < minimum) {
532 return Err(format!("{name} is outside its valid range"));
533 }
534 Ok(())
535}
536
537fn validate_finite_range(
538 name: &str,
539 value: Option<f64>,
540 minimum: f64,
541 maximum: f64,
542) -> Result<(), String> {
543 validate_finite_min(name, value, minimum, false)?;
544 if value.is_some_and(|value| value > maximum) {
545 return Err(format!("{name} is outside its valid range"));
546 }
547 Ok(())
548}