summaryrefslogtreecommitdiffstats
path: root/src/client/model.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-06-21 06:00:26 +0800
committerGitHub <noreply@github.com>2024-06-21 06:00:26 +0800
commitabc588daac6053ec2edbdcde3f5a2dc5eb7d50b8 (patch)
tree9f1cd8a40bd959420dfdf4ac3ae8b52f823aa1e9 /src/client/model.rs
parent2eab71a641827e503b14952373aec82661192ba2 (diff)
downloadaichat-abc588daac6053ec2edbdcde3f5a2dc5eb7d50b8.tar.gz
feat: support rerank (#620)
Diffstat (limited to 'src/client/model.rs')
-rw-r--r--src/client/model.rs13
1 files changed, 10 insertions, 3 deletions
diff --git a/src/client/model.rs b/src/client/model.rs
index 56421bf..ebf1264 100644
--- a/src/client/model.rs
+++ b/src/client/model.rs
@@ -1,5 +1,5 @@
use super::{
- list_chat_models, list_embedding_models,
+ list_chat_models, list_embedding_models, list_rerank_models,
message::{Message, MessageContent},
EmbeddingsData,
};
@@ -46,14 +46,21 @@ impl Model {
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}'"),
+ None => bail!("Invalid chat 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}'"),
+ None => bail!("Invalid embedding model '{model_id}'"),
+ }
+ }
+
+ pub fn retrieve_rerank(config: &Config, model_id: &str) -> Result<Self> {
+ match Self::find(&list_rerank_models(config), model_id) {
+ Some(v) => Ok(v),
+ None => bail!("Invalid rerank model '{model_id}'"),
}
}