diff options
| author | sigoden <sigoden@gmail.com> | 2024-06-14 19:12:18 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-06-14 19:12:18 +0800 |
| commit | 746b087111fabc10ec3f3f3e9ef3628d1eb47fd8 (patch) | |
| tree | 19f072b9b0f3b7befd04e128dbc01eb468c9b9f1 /src/client/model.rs | |
| parent | c1d39e4621373d232cb7520b2a425b536c7f2dd9 (diff) | |
| download | aichat-746b087111fabc10ec3f3f3e9ef3628d1eb47fd8.tar.gz | |
refactor: add/modify rag-related config (#599)
Diffstat (limited to 'src/client/model.rs')
| -rw-r--r-- | src/client/model.rs | 11 |
1 files changed, 9 insertions, 2 deletions
diff --git a/src/client/model.rs b/src/client/model.rs index d555232..56421bf 100644 --- a/src/client/model.rs +++ b/src/client/model.rs @@ -1,5 +1,5 @@ use super::{ - list_chat_models, + list_chat_models, list_embedding_models, message::{Message, MessageContent}, EmbeddingsData, }; @@ -43,13 +43,20 @@ impl Model { .collect() } - pub fn retrieve(config: &Config, model_id: &str) -> Result<Self> { + pub fn retrieve_chat(config: &Config, model_id: &str) -> Result<Self> { match Self::find(&list_chat_models(config), model_id) { Some(v) => Ok(v), None => bail!("Invalid model '{model_id}'"), } } + pub fn retrieve_embedding(config: &Config, model_id: &str) -> Result<Self> { + match Self::find(&list_embedding_models(config), model_id) { + Some(v) => Ok(v), + None => bail!("Invalid model '{model_id}'"), + } + } + pub fn find(models: &[&Self], model_id: &str) -> Option<Self> { let mut model = None; let (client_name, model_name) = match model_id.split_once(':') { |
