Skip to main content

vv_agent/
model_settings.rs

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}