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 let (Some(default), Some(response_format)) =
91            (&defaults.response_format, &mut options.response_format)
92        {
93            merge_response_format(default, response_format);
94        }
95        if options.tools.is_empty() && !defaults.tools.is_empty() {
96            options.tools = defaults.tools.clone();
97        }
98        if !defaults.headers.is_empty() {
99            let mut headers = defaults.headers.clone();
100            headers.merge(&options.headers);
101            options.headers = headers;
102        }
103        if !defaults.provider_options.is_empty() {
104            options.provider_options =
105                merge_provider_options(&defaults.provider_options, options.provider_options);
106        }
107        options
108    }
109}
110
111fn merge_response_format(default: &ResponseFormat, response_format: &mut ResponseFormat) {
112    if let (
113        ResponseFormat::Json {
114            schema: base_schema,
115            name: base_name,
116            description: base_description,
117        },
118        ResponseFormat::Json {
119            schema,
120            name,
121            description,
122        },
123    ) = (default, response_format)
124    {
125        if schema.is_none() {
126            *schema = base_schema.clone();
127        } else if let (Some(JsonValue::Object(base)), Some(JsonValue::Object(overrides))) =
128            (base_schema, schema.as_mut())
129        {
130            *overrides = merge_json_objects(base, std::mem::take(overrides));
131        }
132        if name.is_none() {
133            *name = base_name.clone();
134        }
135        if description.is_none() {
136            *description = base_description.clone();
137        }
138    }
139}
140
141impl LanguageModelMiddleware for DefaultSettings {
142    fn transform_params<'a>(
143        &'a self,
144        options: CallOptions,
145        _ctx: MiddlewareContext<'a>,
146    ) -> BoxFuture<'a, Result<CallOptions, ProviderError>> {
147        let options = self.apply(options);
148        Box::pin(async move { Ok(options) })
149    }
150}
151
152pub(crate) fn merge_provider_options(
153    base: &ProviderOptions,
154    overrides: ProviderOptions,
155) -> ProviderOptions {
156    let mut merged = base.clone();
157    for (provider, options) in overrides {
158        if matches!(provider.as_str(), "__proto__" | "constructor" | "prototype") {
159            continue;
160        }
161        match merged.remove(&provider) {
162            Some(existing) => {
163                merged.insert(provider, merge_json_objects(&existing, options));
164            }
165            None => {
166                merged.insert(provider, options);
167            }
168        }
169    }
170    merged
171}
172
173/// Deeply merges two JSON objects: keys of `overrides` win, except that
174/// nested objects on both sides are merged recursively. Arrays and scalars
175/// (including `null`) override.
176#[must_use]
177pub fn merge_json_objects(base: &JsonObject, overrides: JsonObject) -> JsonObject {
178    let mut merged = base.clone();
179    for (key, value) in overrides {
180        if matches!(key.as_str(), "__proto__" | "constructor" | "prototype") {
181            continue;
182        }
183        match (merged.remove(&key), value) {
184            (Some(JsonValue::Object(existing)), JsonValue::Object(incoming)) => {
185                merged.insert(
186                    key,
187                    JsonValue::Object(merge_json_objects(&existing, incoming)),
188                );
189            }
190            (_, value) => {
191                merged.insert(key, value);
192            }
193        }
194    }
195    merged
196}