summaryrefslogtreecommitdiffstats
path: root/src/client/model.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-06-14 19:12:18 +0800
committerGitHub <noreply@github.com>2024-06-14 19:12:18 +0800
commit746b087111fabc10ec3f3f3e9ef3628d1eb47fd8 (patch)
tree19f072b9b0f3b7befd04e128dbc01eb468c9b9f1 /src/client/model.rs
parentc1d39e4621373d232cb7520b2a425b536c7f2dd9 (diff)
downloadaichat-746b087111fabc10ec3f3f3e9ef3628d1eb47fd8.tar.gz
refactor: add/modify rag-related config (#599)
Diffstat (limited to 'src/client/model.rs')
-rw-r--r--src/client/model.rs11
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(':') {