summaryrefslogtreecommitdiffstats
path: root/src/client/model.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-04-23 16:46:48 +0800
committerGitHub <noreply@github.com>2024-04-23 16:46:48 +0800
commitd1aafa11153ab689c21c2c57c47da52337d8e8d1 (patch)
treedc20dc033e9d376aab09941835a842b22fe32c02 /src/client/model.rs
parent1cc89eff514669d6459f62cc4eed8e0b34d8ef0c (diff)
downloadaichat-d1aafa11153ab689c21c2c57c47da52337d8e8d1.tar.gz
feat: customize model's max_output_tokens (#428)
Diffstat (limited to 'src/client/model.rs')
-rw-r--r--src/client/model.rs40
1 files changed, 37 insertions, 3 deletions
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<usize>,
+ pub max_output_tokens: Option<isize>,
pub extra_fields: Option<serde_json::Map<String, serde_json::Value>>,
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<isize>) -> 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<Model> {
+ 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<usize>,
+ pub max_output_tokens: Option<isize>,
pub extra_fields: Option<serde_json::Map<String, serde_json::Value>>,
#[serde(deserialize_with = "deserialize_capabilities")]
#[serde(default = "default_capabilities")]