summaryrefslogtreecommitdiffstats
path: root/src/client/common.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-06-25 07:39:35 +0800
committerGitHub <noreply@github.com>2024-06-25 07:39:35 +0800
commit2fbb5271af27eee7ba3ae6d2b3c1aa0818648c8f (patch)
tree5e72e72d592025c00a46cebd0eb7012ccd765702 /src/client/common.rs
parented71901611247d41daed8112a5106b42eb12395b (diff)
downloadaichat-2fbb5271af27eee7ba3ae6d2b3c1aa0818648c8f.tar.gz
feat: support rag-dedicated clients (jina and voyageai) (#645)
Diffstat (limited to 'src/client/common.rs')
-rw-r--r--src/client/common.rs10
1 files changed, 6 insertions, 4 deletions
diff --git a/src/client/common.rs b/src/client/common.rs
index 1b1e723..bf2f336 100644
--- a/src/client/common.rs
+++ b/src/client/common.rs
@@ -83,7 +83,9 @@ macro_rules! register_client {
let client_name = Self::name(local_config);
if local_config.models.is_empty() {
if let Some(models) = $crate::client::ALL_MODELS.iter().find(|v| {
- v.platform == $name || ($name == "openai-compatible" && local_config.name.as_deref() == Some(&v.platform))
+ v.platform == $name ||
+ ($name == OpenAICompatibleClient::NAME && local_config.name.as_deref() == Some(&v.platform)) ||
+ ($name == RagDedicatedClient::NAME && local_config.name.as_deref() == Some(&v.platform))
}) {
return Model::from_config(client_name, &models.models);
}
@@ -432,7 +434,7 @@ pub trait Client: Sync + Send {
_client: &ReqwestClient,
_data: EmbeddingsData,
) -> Result<EmbeddingsOutput> {
- bail!("No embeddings api")
+ bail!("The client doesn't support embeddings api")
}
async fn rerank_inner(
@@ -440,7 +442,7 @@ pub trait Client: Sync + Send {
_client: &ReqwestClient,
_data: RerankData,
) -> Result<RerankOutput> {
- bail!("No rerank api")
+ bail!("The client doesn't support rerank api")
}
}
@@ -566,7 +568,7 @@ pub fn create_openai_compatible_client_config(client: &str) -> Result<Option<(St
None => Ok(None),
Some((name, api_base)) => {
let mut config = json!({
- "type": "openai-compatible",
+ "type": OpenAICompatibleClient::NAME,
"name": name,
"api_base": api_base,
});