summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
Diffstat (limited to 'src')
-rw-r--r--src/client/model.rs80
-rw-r--r--src/rag/mod.rs9
-rw-r--r--src/utils/prompt_input.rs18
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)
+ }
+}