diff options
| author | sigoden <sigoden@gmail.com> | 2024-04-29 15:34:24 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-04-29 15:34:24 +0800 |
| commit | c752ba9b272ddf8186f7786ac32a88a410e6eee8 (patch) | |
| tree | cf3058037e13eb2226e5041ddcf3a7f0c5f44a42 /src/client/model.rs | |
| parent | 3a00fb283e4f503901ce198ead65aebc46ae1875 (diff) | |
| download | aichat-c752ba9b272ddf8186f7786ac32a88a410e6eee8.tar.gz | |
feat: `.model` repl completions show max tokens and price (#462)
Diffstat (limited to 'src/client/model.rs')
| -rw-r--r-- | src/client/model.rs | 62 |
1 files changed, 57 insertions, 5 deletions
diff --git a/src/client/model.rs b/src/client/model.rs index 459d94e..aface38 100644 --- a/src/client/model.rs +++ b/src/client/model.rs @@ -1,6 +1,6 @@ use super::message::{Message, MessageContent}; -use crate::utils::count_tokens; +use crate::utils::{count_tokens, format_option_value}; use anyhow::{bail, Result}; use serde::Deserialize; @@ -14,6 +14,9 @@ pub struct Model { pub name: String, pub max_input_tokens: Option<usize>, pub max_output_tokens: Option<isize>, + pub ref_max_output_tokens: Option<isize>, + pub input_price: Option<f64>, + pub output_price: Option<f64>, pub extra_fields: Option<serde_json::Map<String, serde_json::Value>>, pub capabilities: ModelCapabilities, } @@ -32,6 +35,9 @@ impl Model { extra_fields: None, max_input_tokens: None, max_output_tokens: None, + ref_max_output_tokens: None, + input_price: None, + output_price: None, capabilities: ModelCapabilities::Text, } } @@ -43,13 +49,16 @@ impl Model { Model::new(client_name, &v.name) .set_max_input_tokens(v.max_input_tokens) .set_max_output_tokens(v.max_output_tokens) + .set_ref_max_output_tokens(v.ref_max_output_tokens) + .set_input_price(v.input_price) + .set_output_price(v.output_price) .set_supports_vision(v.supports_vision) .set_extra_fields(&v.extra_fields) }) .collect() } - pub fn find(models: &[Self], value: &str) -> Option<Self> { + pub fn find(models: &[&Self], value: &str) -> Option<Self> { let mut model = None; let (client_name, model_name) = match value.split_once(':') { Some((client_name, model_name)) => { @@ -64,16 +73,16 @@ impl Model { match model_name { Some(model_name) => { if let Some(found) = models.iter().find(|v| v.id() == value) { - model = Some(found.clone()); + model = Some((*found).clone()); } else if let Some(found) = models.iter().find(|v| v.client_name == client_name) { - let mut found = found.clone(); + let mut found = (*found).clone(); found.name = model_name.to_string(); model = Some(found) } } None => { if let Some(found) = models.iter().find(|v| v.client_name == client_name) { - model = Some(found.clone()); + model = Some((*found).clone()); } } } @@ -84,6 +93,23 @@ impl Model { format!("{}:{}", self.client_name, self.name) } + pub fn description(&self) -> String { + let max_input_tokens = format_option_value(&self.max_input_tokens); + let max_output_tokens = + format_option_value(&self.max_output_tokens.or(self.ref_max_output_tokens)); + let input_price = format_option_value(&self.input_price); + let output_price = format_option_value(&self.output_price); + let vision = if self.capabilities.contains(ModelCapabilities::Vision) { + "👁" + } else { + "" + }; + format!( + "{:>8} / {:>8} | {:>6} / {:>6} {}", + max_input_tokens, max_output_tokens, input_price, output_price, vision + ) + } + pub fn set_max_input_tokens(mut self, max_input_tokens: Option<usize>) -> Self { match max_input_tokens { None | Some(0) => self.max_input_tokens = None, @@ -100,6 +126,30 @@ impl Model { self } + pub fn set_ref_max_output_tokens(mut self, ref_max_output_tokens: Option<isize>) -> Self { + match ref_max_output_tokens { + None | Some(0) => self.ref_max_output_tokens = None, + _ => self.ref_max_output_tokens = ref_max_output_tokens, + } + self + } + + pub fn set_input_price(mut self, input_price: Option<f64>) -> Self { + match input_price { + None => self.input_price = None, + _ => self.input_price = input_price, + } + self + } + + pub fn set_output_price(mut self, output_price: Option<f64>) -> Self { + match output_price { + None => self.output_price = None, + _ => self.output_price = output_price, + } + self + } + pub fn set_supports_vision(mut self, supports_vision: bool) -> Self { if supports_vision { self.capabilities |= ModelCapabilities::Vision; @@ -178,6 +228,8 @@ pub struct ModelConfig { pub name: String, pub max_input_tokens: Option<usize>, pub max_output_tokens: Option<isize>, + #[serde(rename = "max_output_tokens?")] + pub ref_max_output_tokens: Option<isize>, pub input_price: Option<f64>, pub output_price: Option<f64>, #[serde(default)] |
