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