ferrin_core/middleware/builtin/
default_settings.rs1use 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#[derive(Debug, Clone, Default)]
23pub struct CallDefaults {
24 pub max_output_tokens: Option<u32>,
26 pub temperature: Option<f64>,
28 pub stop_sequences: Option<Vec<String>>,
30 pub top_p: Option<f64>,
32 pub top_k: Option<u32>,
34 pub presence_penalty: Option<f64>,
36 pub frequency_penalty: Option<f64>,
38 pub response_format: Option<ResponseFormat>,
40 pub seed: Option<u64>,
42 pub tools: Vec<ToolDefinition>,
44 pub tool_choice: Option<ToolChoice>,
46 pub headers: Headers,
48 pub provider_options: ProviderOptions,
50}
51
52#[derive(Debug, Clone)]
54pub struct DefaultSettings {
55 defaults: CallDefaults,
56}
57
58#[must_use]
60pub fn default_settings(defaults: CallDefaults) -> DefaultSettings {
61 DefaultSettings { defaults }
62}
63
64impl DefaultSettings {
65 #[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#[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}