summaryrefslogtreecommitdiffstats
path: root/src/client/azure_openai.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-06-05 09:02:23 +0800
committerGitHub <noreply@github.com>2024-06-05 09:02:23 +0800
commit1ec6abfaee2fdc189b348b7e3a8145bd9a84da74 (patch)
tree6f68a860a39b0fbe87784de925b5a31e00e74e33 /src/client/azure_openai.rs
parent71f2e94579511d7524f5534377001ab3f02a9597 (diff)
downloadaichat-1ec6abfaee2fdc189b348b7e3a8145bd9a84da74.tar.gz
feat: support RAG (#560)
* feat: support RAG * support more embeddings models and implement concurrent embedding api * show the progress of addings paths * ignore embedding context when saving message * embedding model max_chunk_size => default_chunk_size * support pdf and pandoc formats (docx, epub, ipynb)
Diffstat (limited to 'src/client/azure_openai.rs')
-rw-r--r--src/client/azure_openai.rs31
1 files changed, 24 insertions, 7 deletions
diff --git a/src/client/azure_openai.rs b/src/client/azure_openai.rs
index 52c8a34..19d234a 100644
--- a/src/client/azure_openai.rs
+++ b/src/client/azure_openai.rs
@@ -1,7 +1,5 @@
-use super::{
- openai::*, AzureOpenAIClient, ChatCompletionsData, Client, ExtraConfig, Model, ModelData,
- ModelPatches, PromptAction, PromptKind,
-};
+use super::*;
+use super::openai::*;
use anyhow::Result;
use reqwest::{Client as ReqwestClient, RequestBuilder};
@@ -12,6 +10,7 @@ pub struct AzureOpenAIConfig {
pub name: Option<String>,
pub api_base: Option<String>,
pub api_key: Option<String>,
+ #[serde(default)]
pub models: Vec<ModelData>,
pub patches: Option<ModelPatches>,
pub extra: Option<ExtraConfig>,
@@ -42,7 +41,7 @@ impl AzureOpenAIClient {
let api_key = self.get_api_key()?;
let mut body = openai_build_chat_completions_body(data, &self.model);
- self.patch_request_body(&mut body);
+ self.patch_chat_completions_body(&mut body);
let url = format!(
"{}/openai/deployments/{}/chat/completions?api-version=2024-02-01",
@@ -56,10 +55,28 @@ impl AzureOpenAIClient {
Ok(builder)
}
+
+ fn embeddings_builder(
+ &self,
+ client: &ReqwestClient,
+ data: EmbeddingsData,
+ ) -> Result<RequestBuilder> {
+ let api_base = self.get_api_base()?;
+ let api_key = self.get_api_key()?;
+
+ let body = openai_build_embeddings_body(data, &self.model);
+
+ let url = format!("{api_base}/embeddings");
+
+ let builder = client.post(url).bearer_auth(api_key).json(&body);
+
+ Ok(builder)
+ }
}
impl_client_trait!(
AzureOpenAIClient,
- crate::client::openai::openai_chat_completions,
- crate::client::openai::openai_chat_completions_streaming
+ openai_chat_completions,
+ openai_chat_completions_streaming,
+ openai_embeddings
);