summaryrefslogtreecommitdiffstats
path: root/src/client/openai_compatible.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/openai_compatible.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/openai_compatible.rs')
-rw-r--r--src/client/openai_compatible.rs73
1 files changed, 48 insertions, 25 deletions
diff --git a/src/client/openai_compatible.rs b/src/client/openai_compatible.rs
index 74cd954..af7cd0e 100644
--- a/src/client/openai_compatible.rs
+++ b/src/client/openai_compatible.rs
@@ -1,7 +1,5 @@
-use super::{
- openai::*, ChatCompletionsData, Client, ExtraConfig, Model, ModelData, ModelPatches,
- OpenAICompatibleClient, PromptAction, PromptKind, OPENAI_COMPATIBLE_PLATFORMS,
-};
+use super::*;
+use super::openai::*;
use anyhow::Result;
use reqwest::{Client as ReqwestClient, RequestBuilder};
@@ -41,27 +39,11 @@ impl OpenAICompatibleClient {
client: &ReqwestClient,
data: ChatCompletionsData,
) -> Result<RequestBuilder> {
- let api_base = match self.get_api_base() {
- Ok(v) => v,
- Err(err) => {
- match OPENAI_COMPATIBLE_PLATFORMS
- .into_iter()
- .find_map(|(name, api_base)| {
- if name == self.model.client_name() {
- Some(api_base.to_string())
- } else {
- None
- }
- }) {
- Some(v) => v,
- None => return Err(err),
- }
- }
- };
let api_key = self.get_api_key().ok();
+ let api_base = self.get_api_base_ext()?;
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 chat_endpoint = self
.config
@@ -71,7 +53,7 @@ impl OpenAICompatibleClient {
let url = format!("{api_base}{chat_endpoint}");
- debug!("OpenAICompatible Request: {url} {body}");
+ debug!("OpenAICompatible Chat Completions Request: {url} {body}");
let mut builder = client.post(url).json(&body);
if let Some(api_key) = api_key {
@@ -80,10 +62,51 @@ impl OpenAICompatibleClient {
Ok(builder)
}
+
+ fn embeddings_builder(
+ &self,
+ client: &ReqwestClient,
+ data: EmbeddingsData,
+ ) -> Result<RequestBuilder> {
+ let api_key = self.get_api_key()?;
+ let api_base = self.get_api_base_ext()?;
+
+ let body = openai_build_embeddings_body(data, &self.model);
+
+ let url = format!("{api_base}/embeddings");
+
+ debug!("OpenAICompatible Embeddings Request: {url} {body}");
+
+ let builder = client.post(url).bearer_auth(api_key).json(&body);
+
+ Ok(builder)
+ }
+
+ fn get_api_base_ext(&self) -> Result<String> {
+ let api_base = match self.get_api_base() {
+ Ok(v) => v,
+ Err(err) => {
+ match OPENAI_COMPATIBLE_PLATFORMS
+ .into_iter()
+ .find_map(|(name, api_base)| {
+ if name == self.model.client_name() {
+ Some(api_base.to_string())
+ } else {
+ None
+ }
+ }) {
+ Some(v) => v,
+ None => return Err(err),
+ }
+ }
+ };
+ Ok(api_base)
+ }
}
impl_client_trait!(
OpenAICompatibleClient,
- crate::client::openai::openai_chat_completions,
- crate::client::openai::openai_chat_completions_streaming
+ openai_chat_completions,
+ openai_chat_completions_streaming,
+ openai_embeddings
);