From c752ba9b272ddf8186f7786ac32a88a410e6eee8 Mon Sep 17 00:00:00 2001 From: sigoden Date: Mon, 29 Apr 2024 15:34:24 +0800 Subject: feat: `.model` repl completions show max tokens and price (#462) --- src/client/common.rs | 25 ++++++++++++--------- src/client/model.rs | 62 +++++++++++++++++++++++++++++++++++++++++++++++----- 2 files changed, 72 insertions(+), 15 deletions(-) (limited to 'src/client') diff --git a/src/client/common.rs b/src/client/common.rs index 842741c..85255ee 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -147,17 +147,22 @@ macro_rules! register_client { anyhow::bail!("Unknown client '{}'", client) } - pub fn list_models(config: &$crate::config::Config) -> Vec<$crate::client::Model> { - config - .clients - .iter() - .flat_map(|v| match v { - $(ClientConfig::$config(c) => $client::list_models(c),)+ - ClientConfig::Unknown => vec![], - }) - .collect() + static mut ALL_CLIENTS: Option> = None; + + pub fn list_models(config: &$crate::config::Config) -> Vec<&$crate::client::Model> { + if unsafe { ALL_CLIENTS.is_none() } { + let models: Vec<_> = config + .clients + .iter() + .flat_map(|v| match v { + $(ClientConfig::$config(c) => $client::list_models(c),)+ + ClientConfig::Unknown => vec![], + }) + .collect(); + unsafe { ALL_CLIENTS = Some(models) }; + } + unsafe { ALL_CLIENTS.as_ref().unwrap().iter().collect() } } - }; } 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, pub max_output_tokens: Option, + pub ref_max_output_tokens: Option, + pub input_price: Option, + pub output_price: Option, pub extra_fields: Option>, 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 { + pub fn find(models: &[&Self], value: &str) -> Option { 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) -> 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) -> 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) -> 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) -> 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, pub max_output_tokens: Option, + #[serde(rename = "max_output_tokens?")] + pub ref_max_output_tokens: Option, pub input_price: Option, pub output_price: Option, #[serde(default)] -- cgit v1.2.3