use ferrin_spec::BoxFuture;
use ferrin_spec::CallOptions;
use ferrin_spec::Headers;
use ferrin_spec::JsonObject;
use ferrin_spec::JsonValue;
use ferrin_spec::ProviderOptions;
use ferrin_spec::ToolChoice;
use ferrin_spec::ToolDefinition;
use ferrin_spec::error::ProviderError;
use ferrin_spec::language_model::ResponseFormat;
use crate::middleware::LanguageModelMiddleware;
use crate::middleware::MiddlewareContext;
#[derive(Debug, Clone, Default)]
pub struct CallDefaults {
pub max_output_tokens: Option<u32>,
pub temperature: Option<f64>,
pub stop_sequences: Option<Vec<String>>,
pub top_p: Option<f64>,
pub top_k: Option<u32>,
pub presence_penalty: Option<f64>,
pub frequency_penalty: Option<f64>,
pub response_format: Option<ResponseFormat>,
pub seed: Option<u64>,
pub tools: Vec<ToolDefinition>,
pub tool_choice: Option<ToolChoice>,
pub headers: Headers,
pub provider_options: ProviderOptions,
}
#[derive(Debug, Clone)]
pub struct DefaultSettings {
defaults: CallDefaults,
}
#[must_use]
pub fn default_settings(defaults: CallDefaults) -> DefaultSettings {
DefaultSettings { defaults }
}
impl DefaultSettings {
#[must_use]
pub fn apply(&self, mut options: CallOptions) -> CallOptions {
let defaults = &self.defaults;
macro_rules! fill {
($($field:ident),* $(,)?) => {
$(
if options.$field.is_none() {
options.$field = defaults.$field.clone();
}
)*
};
}
fill!(
max_output_tokens,
temperature,
stop_sequences,
top_p,
top_k,
presence_penalty,
frequency_penalty,
response_format,
seed,
tool_choice,
);
if options.tools.is_empty() && !defaults.tools.is_empty() {
options.tools = defaults.tools.clone();
}
if !defaults.headers.is_empty() {
let mut headers = defaults.headers.clone();
headers.merge(&options.headers);
options.headers = headers;
}
if !defaults.provider_options.is_empty() {
options.provider_options =
merge_provider_options(&defaults.provider_options, options.provider_options);
}
options
}
}
impl LanguageModelMiddleware for DefaultSettings {
fn transform_params<'a>(
&'a self,
options: CallOptions,
_ctx: MiddlewareContext<'a>,
) -> BoxFuture<'a, Result<CallOptions, ProviderError>> {
let options = self.apply(options);
Box::pin(async move { Ok(options) })
}
}
fn merge_provider_options(base: &ProviderOptions, overrides: ProviderOptions) -> ProviderOptions {
let mut merged = base.clone();
for (provider, options) in overrides {
match merged.remove(&provider) {
Some(existing) => {
merged.insert(provider, merge_json_objects(&existing, options));
}
None => {
merged.insert(provider, options);
}
}
}
merged
}
#[must_use]
pub fn merge_json_objects(base: &JsonObject, overrides: JsonObject) -> JsonObject {
let mut merged = base.clone();
for (key, value) in overrides {
match (merged.remove(&key), value) {
(Some(JsonValue::Object(existing)), JsonValue::Object(incoming)) => {
merged.insert(
key,
JsonValue::Object(merge_json_objects(&existing, incoming)),
);
}
(_, value) => {
merged.insert(key, value);
}
}
}
merged
}