Skip to main content

ferrin_core/middleware/builtin/
default_settings.rs

1//! Default call settings.
2
3use ferrin_spec::BoxFuture;
4use ferrin_spec::CallOptions;
5use ferrin_spec::Headers;
6use ferrin_spec::JsonObject;
7use ferrin_spec::JsonValue;
8use ferrin_spec::ProviderOptions;
9use ferrin_spec::ToolChoice;
10use ferrin_spec::ToolDefinition;
11use ferrin_spec::error::ProviderError;
12use ferrin_spec::language_model::ResponseFormat;
13
14use crate::middleware::LanguageModelMiddleware;
15use crate::middleware::MiddlewareContext;
16
17/// Settings applied when the call does not set them.
18///
19/// Scalars fill `None` values. `tools` applies when the call sends no
20/// tools. `headers` and `provider_options` are merged with the call's
21/// values taking precedence (provider options merge recursively).
22#[derive(Debug, Clone, Default)]
23pub struct CallDefaults {
24    /// See [`CallOptions::max_output_tokens`].
25    pub max_output_tokens: Option<u32>,
26    /// See [`CallOptions::temperature`].
27    pub temperature: Option<f64>,
28    /// See [`CallOptions::stop_sequences`].
29    pub stop_sequences: Option<Vec<String>>,
30    /// See [`CallOptions::top_p`].
31    pub top_p: Option<f64>,
32    /// See [`CallOptions::top_k`].
33    pub top_k: Option<u32>,
34    /// See [`CallOptions::presence_penalty`].
35    pub presence_penalty: Option<f64>,
36    /// See [`CallOptions::frequency_penalty`].
37    pub frequency_penalty: Option<f64>,
38    /// See [`CallOptions::response_format`].
39    pub response_format: Option<ResponseFormat>,
40    /// See [`CallOptions::seed`].
41    pub seed: Option<u64>,
42    /// See [`CallOptions::tools`].
43    pub tools: Vec<ToolDefinition>,
44    /// See [`CallOptions::tool_choice`].
45    pub tool_choice: Option<ToolChoice>,
46    /// See [`CallOptions::headers`].
47    pub headers: Headers,
48    /// See [`CallOptions::provider_options`].
49    pub provider_options: ProviderOptions,
50}
51
52/// Middleware created by [`default_settings`].
53#[derive(Debug, Clone)]
54pub struct DefaultSettings {
55    defaults: CallDefaults,
56}
57
58/// Fills unset call options from `defaults`.
59#[must_use]
60pub fn default_settings(defaults: CallDefaults) -> DefaultSettings {
61    DefaultSettings { defaults }
62}
63
64impl DefaultSettings {
65    /// Applies the defaults to `options`.
66    #[must_use]
67    pub fn apply(&self, mut options: CallOptions) -> CallOptions {
68        let defaults = &self.defaults;
69        macro_rules! fill {
70            ($($field:ident),* $(,)?) => {
71                $(
72                    if options.$field.is_none() {
73                        options.$field = defaults.$field.clone();
74                    }
75                )*
76            };
77        }
78        fill!(
79            max_output_tokens,
80            temperature,
81            stop_sequences,
82            top_p,
83            top_k,
84            presence_penalty,
85            frequency_penalty,
86            response_format,
87            seed,
88            tool_choice,
89        );
90        if options.tools.is_empty() && !defaults.tools.is_empty() {
91            options.tools = defaults.tools.clone();
92        }
93        if !defaults.headers.is_empty() {
94            let mut headers = defaults.headers.clone();
95            headers.merge(&options.headers);
96            options.headers = headers;
97        }
98        if !defaults.provider_options.is_empty() {
99            options.provider_options =
100                merge_provider_options(&defaults.provider_options, options.provider_options);
101        }
102        options
103    }
104}
105
106impl LanguageModelMiddleware for DefaultSettings {
107    fn transform_params<'a>(
108        &'a self,
109        options: CallOptions,
110        _ctx: MiddlewareContext<'a>,
111    ) -> BoxFuture<'a, Result<CallOptions, ProviderError>> {
112        let options = self.apply(options);
113        Box::pin(async move { Ok(options) })
114    }
115}
116
117fn merge_provider_options(base: &ProviderOptions, overrides: ProviderOptions) -> ProviderOptions {
118    let mut merged = base.clone();
119    for (provider, options) in overrides {
120        match merged.remove(&provider) {
121            Some(existing) => {
122                merged.insert(provider, merge_json_objects(&existing, options));
123            }
124            None => {
125                merged.insert(provider, options);
126            }
127        }
128    }
129    merged
130}
131
132/// Deeply merges two JSON objects: keys of `overrides` win, except that
133/// nested objects on both sides are merged recursively. Arrays and scalars
134/// (including `null`) override.
135#[must_use]
136pub fn merge_json_objects(base: &JsonObject, overrides: JsonObject) -> JsonObject {
137    let mut merged = base.clone();
138    for (key, value) in overrides {
139        match (merged.remove(&key), value) {
140            (Some(JsonValue::Object(existing)), JsonValue::Object(incoming)) => {
141                merged.insert(
142                    key,
143                    JsonValue::Object(merge_json_objects(&existing, incoming)),
144                );
145            }
146            (_, value) => {
147                merged.insert(key, value);
148            }
149        }
150    }
151    merged
152}