diff options
Diffstat (limited to 'src')
| -rw-r--r-- | src/client/model.rs | 80 | ||||
| -rw-r--r-- | src/rag/mod.rs | 9 | ||||
| -rw-r--r-- | src/utils/prompt_input.rs | 18 |
3 files changed, 75 insertions, 32 deletions
diff --git a/src/client/model.rs b/src/client/model.rs index 4c69d78..d661d85 100644 --- a/src/client/model.rs +++ b/src/client/model.rs @@ -124,35 +124,56 @@ impl Model { } pub fn description(&self) -> String { - let ModelData { - max_input_tokens, - max_output_tokens, - input_price, - output_price, - supports_vision, - supports_function_calling, - .. - } = &self.data; - let max_input_tokens = format_option_value(max_input_tokens); - let max_output_tokens = format_option_value(max_output_tokens); - let input_price = format_option_value(input_price); - let output_price = format_option_value(output_price); - let mut capabilities = vec![]; - if *supports_vision { - capabilities.push('👁'); - }; - if *supports_function_calling { - capabilities.push('⚒'); - }; - let capabilities: String = capabilities - .into_iter() - .map(|v| format!("{v} ")) - .collect::<Vec<String>>() - .join(""); - format!( - "{:>8} / {:>8} | {:>6} / {:>6} {:>6}", - max_input_tokens, max_output_tokens, input_price, output_price, capabilities - ) + match self.model_type() { + "chat" => { + let ModelData { + max_input_tokens, + max_output_tokens, + input_price, + output_price, + supports_vision, + supports_function_calling, + .. + } = &self.data; + let max_input_tokens = format_option_value(max_input_tokens); + let max_output_tokens = format_option_value(max_output_tokens); + let input_price = format_option_value(input_price); + let output_price = format_option_value(output_price); + let mut capabilities = vec![]; + if *supports_vision { + capabilities.push('👁'); + }; + if *supports_function_calling { + capabilities.push('⚒'); + }; + let capabilities: String = capabilities + .into_iter() + .map(|v| format!("{v} ")) + .collect::<Vec<String>>() + .join(""); + format!( + "{:>8} / {:>8} | {:>6} / {:>6} {:>6}", + max_input_tokens, max_output_tokens, input_price, output_price, capabilities + ) + } + "embedding" => { + let ModelData { + max_input_tokens, + input_price, + output_vector_size, + max_batch_size, + .. + } = &self.data; + let dimension = format_option_value(output_vector_size); + let max_tokens = format_option_value(max_input_tokens); + let price = format_option_value(input_price); + let batch = format_option_value(max_batch_size); + format!( + "dimension:{dimension}; max-tokens:{max_tokens}; price:{price}; batch:{batch}" + ) + } + _ => String::new(), + } } pub fn max_input_tokens(&self) -> Option<usize> { @@ -261,6 +282,7 @@ pub struct ModelData { pub supports_function_calling: bool, // embedding-only properties + pub output_vector_size: Option<usize>, pub default_chunk_size: Option<usize>, pub max_batch_size: Option<usize>, } diff --git a/src/rag/mod.rs b/src/rag/mod.rs index 116beea..30ca92f 100644 --- a/src/rag/mod.rs +++ b/src/rag/mod.rs @@ -548,9 +548,12 @@ pub fn split_document_id(value: DocumentId) -> (usize, usize) { } fn select_embedding_model(models: &[&Model]) -> Result<String> { - let model_ids: Vec<_> = models.iter().map(|v| v.id()).collect(); - let model_id = Select::new("Select embedding model:", model_ids).prompt()?; - Ok(model_id) + let models: Vec<_> = models + .iter() + .map(|v| SelectOption::new(v.id(), v.description())) + .collect(); + let result = Select::new("Select embedding model:", models).prompt()?; + Ok(result.value) } fn set_chunk_size(model: &Model) -> Result<usize> { diff --git a/src/utils/prompt_input.rs b/src/utils/prompt_input.rs index e887b9f..26343a2 100644 --- a/src/utils/prompt_input.rs +++ b/src/utils/prompt_input.rs @@ -54,3 +54,21 @@ fn validate_integer(text: &str) -> Validation { Validation::Valid } } + +#[derive(Debug)] +pub struct SelectOption { + pub value: String, + pub description: String, +} + +impl SelectOption { + pub fn new(value: String, description: String) -> Self { + Self { value, description } + } +} + +impl std::fmt::Display for SelectOption { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "{} ({})", self.value, self.description) + } +} |
