From d1aafa11153ab689c21c2c57c47da52337d8e8d1 Mon Sep 17 00:00:00 2001 From: sigoden Date: Tue, 23 Apr 2024 16:46:48 +0800 Subject: feat: customize model's max_output_tokens (#428) --- src/client/model.rs | 40 +++++++++++++++++++++++++++++++++++++--- 1 file changed, 37 insertions(+), 3 deletions(-) (limited to 'src/client/model.rs') diff --git a/src/client/model.rs b/src/client/model.rs index 03877b9..53d1834 100644 --- a/src/client/model.rs +++ b/src/client/model.rs @@ -13,6 +13,7 @@ pub struct Model { pub client_name: String, pub name: String, pub max_input_tokens: Option, + pub max_output_tokens: Option, pub extra_fields: Option>, pub capabilities: ModelCapabilities, } @@ -30,6 +31,7 @@ impl Model { name: name.into(), extra_fields: None, max_input_tokens: None, + max_output_tokens: None, capabilities: ModelCapabilities::Text, } } @@ -90,6 +92,14 @@ impl Model { self } + pub fn set_max_output_tokens(mut self, max_output_tokens: Option) -> Self { + match max_output_tokens { + None | Some(0) => self.max_output_tokens = None, + _ => self.max_output_tokens = max_output_tokens, + } + self + } + pub fn messages_tokens(&self, messages: &[Message]) -> usize { messages .iter() @@ -127,19 +137,43 @@ impl Model { pub fn merge_extra_fields(&self, body: &mut serde_json::Value) { if let (Some(body), Some(extra_fields)) = (body.as_object_mut(), &self.extra_fields) { - for (k, v) in extra_fields { - if !body.contains_key(k) { - body.insert(k.clone(), v.clone()); + for (key, extra_field) in extra_fields { + if body.contains_key(key) { + if let (Some(sub_body), Some(extra_field)) = + (body[key].as_object_mut(), extra_field.as_object()) + { + for (subkey, sub_field) in extra_field { + if !sub_body.contains_key(subkey) { + sub_body.insert(subkey.clone(), sub_field.clone()); + } + } + } + } else { + body.insert(key.clone(), extra_field.clone()); } } } } } +pub fn convert_models(client_name: &str, models: &[ModelConfig]) -> Vec { + models + .iter() + .map(|v| { + Model::new(client_name, &v.name) + .set_capabilities(v.capabilities) + .set_max_input_tokens(v.max_input_tokens) + .set_max_output_tokens(v.max_output_tokens) + .set_extra_fields(v.extra_fields.clone()) + }) + .collect() +} + #[derive(Debug, Clone, Deserialize)] pub struct ModelConfig { pub name: String, pub max_input_tokens: Option, + pub max_output_tokens: Option, pub extra_fields: Option>, #[serde(deserialize_with = "deserialize_capabilities")] #[serde(default = "default_capabilities")] -- cgit v1.2.3