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        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}