summaryrefslogtreecommitdiffstats
path: root/src/client/qianwen.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/qianwen.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/qianwen.rs')
-rw-r--r--src/client/qianwen.rs86
1 files changed, 75 insertions, 11 deletions
diff --git a/src/client/qianwen.rs b/src/client/qianwen.rs
index 0230e21..c34e409 100644
--- a/src/client/qianwen.rs
+++ b/src/client/qianwen.rs
@@ -1,8 +1,4 @@
-use super::{
- maybe_catch_error, message::*, sse_stream, ChatCompletionsData, ChatCompletionsOutput, Client,
- ExtraConfig, Model, ModelData, ModelPatches, PromptAction, PromptKind, QianwenClient,
- SseHandler, SseMmessage,
-};
+use super::*;
use crate::utils::{base64_decode, sha256};
@@ -16,12 +12,15 @@ use serde::Deserialize;
use serde_json::{json, Value};
use std::borrow::BorrowMut;
-const API_URL: &str =
+const CHAT_COMPLETIONS_API_URL: &str =
"https://dashscope.aliyuncs.com/api/v1/services/aigc/text-generation/generation";
-const API_URL_VL: &str =
+const CHAT_COMPLETIONS_API_URL_VL: &str =
"https://dashscope.aliyuncs.com/api/v1/services/aigc/multimodal-generation/generation";
+const EMBEDDINGS_API_URL: &str =
+ "https://dashscope.aliyuncs.com/api/v1/services/embeddings/text-embedding/text-embedding";
+
#[derive(Debug, Clone, Deserialize, Default)]
pub struct QianwenConfig {
pub name: Option<String>,
@@ -48,13 +47,13 @@ impl QianwenClient {
let stream = data.stream;
let url = match self.model.supports_vision() {
- true => API_URL_VL,
- false => API_URL,
+ true => CHAT_COMPLETIONS_API_URL_VL,
+ false => CHAT_COMPLETIONS_API_URL,
};
let (mut body, has_upload) = build_chat_completions_body(data, &self.model)?;
- self.patch_request_body(&mut body);
+ self.patch_chat_completions_body(&mut body);
- debug!("Qianwen Request: {url} {body}");
+ debug!("Qianwen Chat Completions Request: {url} {body}");
let mut builder = client.post(url).bearer_auth(api_key).json(&body);
if stream {
@@ -66,6 +65,37 @@ impl QianwenClient {
Ok(builder)
}
+
+ fn embeddings_builder(
+ &self,
+ client: &ReqwestClient,
+ data: EmbeddingsData,
+ ) -> Result<RequestBuilder> {
+ let api_key = self.get_api_key()?;
+
+ let text_type = match data.query {
+ true => "query",
+ false => "document",
+ };
+
+ let body = json!({
+ "model": self.model.name(),
+ "input": {
+ "texts": data.texts,
+ },
+ "parameters": {
+ "text_type": text_type,
+ }
+ });
+
+ let url = EMBEDDINGS_API_URL;
+
+ debug!("Qianwen Embeddings Request: {url} {body}");
+
+ let builder = client.post(url).bearer_auth(api_key).json(&body);
+
+ Ok(builder)
+ }
}
#[async_trait]
@@ -94,6 +124,15 @@ impl Client for QianwenClient {
let builder = self.chat_completions_builder(client, data)?;
chat_completions_streaming(builder, handler, &self.model).await
}
+
+ async fn embeddings_inner(
+ &self,
+ client: &ReqwestClient,
+ data: EmbeddingsData,
+ ) -> Result<Vec<Vec<f32>>> {
+ let builder = self.embeddings_builder(client, data)?;
+ embeddings(builder).await
+ }
}
async fn chat_completions(builder: RequestBuilder, model: &Model) -> Result<ChatCompletionsOutput> {
@@ -210,6 +249,31 @@ fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Resu
Ok((body, has_upload))
}
+async fn embeddings(
+ builder: RequestBuilder,
+) -> Result<EmbeddingsOutput> {
+ let data: Value = builder.send().await?.json().await?;
+ maybe_catch_error(&data)?;
+ let res_body: EmbeddingsResBody = serde_json::from_value(data).context("Invalid request data")?;
+ let output = res_body.output.embeddings.into_iter().map(|v| v.embedding).collect();
+ Ok(output)
+}
+
+#[derive(Deserialize)]
+struct EmbeddingsResBody {
+ output: EmbeddingsResBodyOutput,
+}
+
+#[derive(Deserialize)]
+struct EmbeddingsResBodyOutput {
+ embeddings: Vec<EmbeddingsResBodyOutputEmbedding>,
+}
+
+#[derive(Deserialize)]
+struct EmbeddingsResBodyOutputEmbedding {
+ embedding: Vec<f32>,
+}
+
fn extract_chat_completions_text(data: &Value, model: &Model) -> Result<ChatCompletionsOutput> {
let err = || anyhow!("Invalid response data: {data}");
let text = if model.name() == "qwen-long" {