summaryrefslogtreecommitdiffstats
path: root/src/client
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
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')
-rw-r--r--src/client/azure_openai.rs31
-rw-r--r--src/client/bedrock.rs10
-rw-r--r--src/client/claude.rs9
-rw-r--r--src/client/cloudflare.rs7
-rw-r--r--src/client/cohere.rs69
-rw-r--r--src/client/common.rs140
-rw-r--r--src/client/ernie.rs9
-rw-r--r--src/client/gemini.rs73
-rw-r--r--src/client/model.rs44
-rw-r--r--src/client/ollama.rs67
-rw-r--r--src/client/openai.rs67
-rw-r--r--src/client/openai_compatible.rs73
-rw-r--r--src/client/qianwen.rs86
-rw-r--r--src/client/replicate.rs9
-rw-r--r--src/client/vertexai.rs79
-rw-r--r--src/client/vertexai_claude.rs11
16 files changed, 604 insertions, 180 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
);
diff --git a/src/client/bedrock.rs b/src/client/bedrock.rs
index 3dfa977..981d1cb 100644
--- a/src/client/bedrock.rs
+++ b/src/client/bedrock.rs
@@ -1,8 +1,6 @@
-use super::{
- catch_error, claude::*, prompt_format::*, BedrockClient, ChatCompletionsData,
- ChatCompletionsOutput, Client, ExtraConfig, Model, ModelData, ModelPatches, PromptAction,
- PromptKind, SseHandler,
-};
+use super::*;
+use super::claude::*;
+use super::prompt_format::*;
use crate::utils::{base64_decode, encode_uri, hex_encode, hmac_sha256, sha256};
@@ -102,7 +100,7 @@ impl BedrockClient {
let headers = IndexMap::new();
let mut body = build_chat_completions_body(data, &self.model, model_category)?;
- self.patch_request_body(&mut body);
+ self.patch_chat_completions_body(&mut body);
let builder = aws_fetch(
client,
diff --git a/src/client/claude.rs b/src/client/claude.rs
index 6533a9a..a16722b 100644
--- a/src/client/claude.rs
+++ b/src/client/claude.rs
@@ -1,9 +1,4 @@
-use super::{
- catch_error, extract_system_message, message::*, sse_stream, ChatCompletionsData,
- ChatCompletionsOutput, ClaudeClient, Client, ExtraConfig, ImageUrl, MessageContent,
- MessageContentPart, Model, ModelData, ModelPatches, PromptAction, PromptKind, SseHandler,
- SseMmessage, ToolCall,
-};
+use super::*;
use anyhow::{bail, Context, Result};
use reqwest::{Client as ReqwestClient, RequestBuilder};
@@ -36,7 +31,7 @@ impl ClaudeClient {
let api_key = self.get_api_key().ok();
let mut body = claude_build_chat_completions_body(data, &self.model)?;
- self.patch_request_body(&mut body);
+ self.patch_chat_completions_body(&mut body);
let url = API_BASE;
diff --git a/src/client/cloudflare.rs b/src/client/cloudflare.rs
index 966cee4..965f20a 100644
--- a/src/client/cloudflare.rs
+++ b/src/client/cloudflare.rs
@@ -1,7 +1,4 @@
-use super::{
- catch_error, sse_stream, ChatCompletionsData, ChatCompletionsOutput, Client, CloudflareClient,
- ExtraConfig, Model, ModelData, ModelPatches, PromptAction, PromptKind, SseHandler, SseMmessage,
-};
+use super::*;
use anyhow::{anyhow, Result};
use reqwest::{Client as ReqwestClient, RequestBuilder};
@@ -39,7 +36,7 @@ impl CloudflareClient {
let api_key = self.get_api_key()?;
let mut body = build_chat_completions_body(data, &self.model)?;
- self.patch_request_body(&mut body);
+ self.patch_chat_completions_body(&mut body);
let url = format!(
"{API_BASE}/accounts/{account_id}/ai/run/{}",
diff --git a/src/client/cohere.rs b/src/client/cohere.rs
index e0a5eec..69c343b 100644
--- a/src/client/cohere.rs
+++ b/src/client/cohere.rs
@@ -1,15 +1,12 @@
-use super::{
- catch_error, extract_system_message, json_stream, message::*, ChatCompletionsData,
- ChatCompletionsOutput, Client, CohereClient, ExtraConfig, Model, ModelData, ModelPatches,
- PromptAction, PromptKind, SseHandler, ToolCall,
-};
+use super::*;
-use anyhow::{bail, Result};
+use anyhow::{bail, Context, Result};
use reqwest::{Client as ReqwestClient, RequestBuilder};
use serde::Deserialize;
use serde_json::{json, Value};
-const API_URL: &str = "https://api.cohere.ai/v1/chat";
+const CHAT_COMPLETIONS_API_URL: &str = "https://api.cohere.ai/v1/chat";
+const EMBEDDINGS_API_URL: &str = "https://api.cohere.ai/v1/embed";
#[derive(Debug, Clone, Deserialize, Default)]
pub struct CohereConfig {
@@ -35,11 +32,38 @@ impl CohereClient {
let api_key = self.get_api_key()?;
let mut body = build_chat_completions_body(data, &self.model)?;
- self.patch_request_body(&mut body);
+ self.patch_chat_completions_body(&mut body);
- let url = API_URL;
+ let url = CHAT_COMPLETIONS_API_URL;
- debug!("Cohere Request: {url} {body}");
+ debug!("Cohere Chat Completions Request: {url} {body}");
+
+ let builder = client.post(url).bearer_auth(api_key).json(&body);
+
+ Ok(builder)
+ }
+
+ fn embeddings_builder(
+ &self,
+ client: &ReqwestClient,
+ data: EmbeddingsData,
+ ) -> Result<RequestBuilder> {
+ let api_key = self.get_api_key()?;
+
+ let input_type = match data.query {
+ true => "search_query",
+ false => "search_document",
+ };
+
+ let body = json!({
+ "model": self.model.name(),
+ "texts": data.texts,
+ "input_type": input_type,
+ });
+
+ let url = EMBEDDINGS_API_URL;
+
+ debug!("Cohere Embeddings Request: {url} {body}");
let builder = client.post(url).bearer_auth(api_key).json(&body);
@@ -47,7 +71,12 @@ impl CohereClient {
}
}
-impl_client_trait!(CohereClient, chat_completions, chat_completions_streaming);
+impl_client_trait!(
+ CohereClient,
+ chat_completions,
+ chat_completions_streaming,
+ embeddings
+);
async fn chat_completions(builder: RequestBuilder) -> Result<ChatCompletionsOutput> {
let res = builder.send().await?;
@@ -100,6 +129,24 @@ async fn chat_completions_streaming(
Ok(())
}
+async fn embeddings(
+ builder: RequestBuilder,
+) -> Result<EmbeddingsOutput> {
+ let res = builder.send().await?;
+ let status = res.status();
+ let data: Value = res.json().await?;
+ if !status.is_success() {
+ catch_error(&data, status.as_u16())?;
+ }
+ let res_body: EmbeddingsResBody = serde_json::from_value(data).context("Invalid request data")?;
+ Ok(res_body.embeddings)
+}
+
+#[derive(Deserialize)]
+struct EmbeddingsResBody {
+ embeddings: Vec<Vec<f32>>,
+}
+
fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Result<Value> {
let ChatCompletionsData {
mut messages,
diff --git a/src/client/common.rs b/src/client/common.rs
index 96ec90b..1055b84 100644
--- a/src/client/common.rs
+++ b/src/client/common.rs
@@ -1,10 +1,13 @@
-use super::{openai::OpenAIConfig, BuiltinModels, ClientConfig, Message, Model, SseHandler};
+use super::*;
use crate::{
config::{GlobalConfig, Input},
function::{eval_tool_calls, FunctionDeclaration, ToolCall, ToolCallResult},
render::{render_error, render_stream},
- utils::{prompt_input_integer, prompt_input_string, tokenize, AbortSignal, PromptKind},
+ utils::{
+ prompt_input_integer, prompt_input_string, tokenize, watch_abort_signal, AbortSignal,
+ PromptKind,
+ },
};
use anyhow::{bail, Context, Result};
@@ -16,13 +19,12 @@ use reqwest::{Client as ReqwestClient, ClientBuilder, Proxy, RequestBuilder};
use serde::Deserialize;
use serde_json::{json, Value};
use std::{env, future::Future, time::Duration};
-use tokio::{sync::mpsc::unbounded_channel, time::sleep};
+use tokio::sync::mpsc::unbounded_channel;
const MODELS_YAML: &str = include_str!("../../models.yaml");
lazy_static! {
- pub static ref ALL_CLIENT_MODELS: Vec<BuiltinModels> =
- serde_yaml::from_str(MODELS_YAML).unwrap();
+ pub static ref ALL_MODELS: Vec<BuiltinModels> = serde_yaml::from_str(MODELS_YAML).unwrap();
static ref ESCAPE_SLASH_RE: Regex = Regex::new(r"(?<!\\)/").unwrap();
}
@@ -92,10 +94,10 @@ macro_rules! register_client {
pub fn list_models(local_config: &$config) -> Vec<Model> {
let client_name = Self::name(local_config);
if local_config.models.is_empty() {
- if let Some(client_models) = $crate::client::ALL_CLIENT_MODELS.iter().find(|v| {
+ 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))
}) {
- return Model::from_config(client_name, &client_models.models);
+ return Model::from_config(client_name, &models.models);
}
vec![]
} else {
@@ -137,10 +139,10 @@ macro_rules! register_client {
anyhow::bail!("Unknown client '{}'", client)
}
- static mut ALL_CLIENTS: Option<Vec<$crate::client::Model>> = None;
+ static mut ALL_CLIENT_MODELS: Option<Vec<$crate::client::Model>> = None;
- pub fn list_models(config: &$crate::config::Config) -> Vec<&$crate::client::Model> {
- if unsafe { ALL_CLIENTS.is_none() } {
+ pub fn list_models(config: &$crate::config::Config) -> Vec<&'static $crate::client::Model> {
+ if unsafe { ALL_CLIENT_MODELS.is_none() } {
let models: Vec<_> = config
.clients
.iter()
@@ -149,9 +151,17 @@ macro_rules! register_client {
ClientConfig::Unknown => vec![],
})
.collect();
- unsafe { ALL_CLIENTS = Some(models) };
+ unsafe { ALL_CLIENT_MODELS = Some(models) };
}
- unsafe { ALL_CLIENTS.as_ref().unwrap().iter().collect() }
+ unsafe { ALL_CLIENT_MODELS.as_ref().unwrap().iter().collect() }
+ }
+
+ pub fn list_chat_models(config: &$crate::config::Config) -> Vec<&'static $crate::client::Model> {
+ list_models(config).into_iter().filter(|v| v.mode() == "chat").collect()
+ }
+
+ pub fn list_embedding_models(config: &$crate::config::Config) -> Vec<&'static $crate::client::Model> {
+ list_models(config).into_iter().filter(|v| v.mode() == "embedding").collect()
}
};
}
@@ -171,10 +181,6 @@ macro_rules! client_common_fns {
self.config.patches.as_ref()
}
- fn list_models(&self) -> Vec<Model> {
- Self::list_models(&self.config)
- }
-
fn name(&self) -> &str {
Self::name(&self.config)
}
@@ -186,10 +192,6 @@ macro_rules! client_common_fns {
fn model_mut(&mut self) -> &mut Model {
&mut self.model
}
-
- fn set_model(&mut self, model: Model) {
- self.model = model;
- }
};
}
@@ -220,6 +222,40 @@ macro_rules! impl_client_trait {
}
}
};
+ ($client:ident, $chat_completions:path, $chat_completions_streaming:path, $embeddings:path) => {
+ #[async_trait::async_trait]
+ impl $crate::client::Client for $crate::client::$client {
+ client_common_fns!();
+
+ async fn chat_completions_inner(
+ &self,
+ client: &reqwest::Client,
+ data: $crate::client::ChatCompletionsData,
+ ) -> anyhow::Result<$crate::client::ChatCompletionsOutput> {
+ let builder = self.chat_completions_builder(client, data)?;
+ $chat_completions(builder).await
+ }
+
+ async fn chat_completions_streaming_inner(
+ &self,
+ client: &reqwest::Client,
+ handler: &mut $crate::client::SseHandler,
+ data: $crate::client::ChatCompletionsData,
+ ) -> Result<()> {
+ let builder = self.chat_completions_builder(client, data)?;
+ $chat_completions_streaming(builder, handler).await
+ }
+
+ async fn embeddings_inner(
+ &self,
+ client: &ReqwestClient,
+ data: EmbeddingsData,
+ ) -> Result<Vec<Vec<f32>>> {
+ let builder = self.embeddings_builder(client, data)?;
+ $embeddings(builder).await
+ }
+ }
+ };
}
#[macro_export]
@@ -256,19 +292,12 @@ pub trait Client: Sync + Send {
fn patches_config(&self) -> Option<&ModelPatches>;
- #[allow(unused)]
fn name(&self) -> &str;
- #[allow(unused)]
- fn list_models(&self) -> Vec<Model>;
-
fn model(&self) -> &Model;
fn model_mut(&mut self) -> &mut Model;
- #[allow(unused)]
- fn set_model(&mut self, model: Model);
-
fn build_client(&self) -> Result<ReqwestClient> {
let mut builder = ReqwestClient::builder();
let extra = self.extra_config();
@@ -288,11 +317,10 @@ pub trait Client: Sync + Send {
return Ok(ChatCompletionsOutput::new(&content));
}
let client = self.build_client()?;
-
let data = input.prepare_completion_data(self.model(), false)?;
self.chat_completions_inner(&client, data)
.await
- .with_context(|| "Failed to get answer")
+ .with_context(|| "Failed to get chat completions")
}
async fn chat_completions_streaming(
@@ -300,15 +328,7 @@ pub trait Client: Sync + Send {
input: &Input,
handler: &mut SseHandler,
) -> Result<()> {
- async fn watch_abort(abort: AbortSignal) {
- loop {
- if abort.aborted() {
- break;
- }
- sleep(Duration::from_millis(100)).await;
- }
- }
- let abort = handler.get_abort();
+ let abort_signal = handler.get_abort();
let input = input.clone();
tokio::select! {
ret = async {
@@ -326,20 +346,28 @@ pub trait Client: Sync + Send {
self.chat_completions_streaming_inner(&client, handler, data).await
} => {
handler.done()?;
- ret.with_context(|| "Failed to get answer")
+ ret.with_context(|| "Failed to get chat completions")
}
- _ = watch_abort(abort.clone()) => {
+ _ = watch_abort_signal(abort_signal) => {
handler.done()?;
Ok(())
},
}
}
- fn patch_request_body(&self, body: &mut Value) {
+ async fn embeddings(&self, data: EmbeddingsData) -> Result<Vec<Vec<f32>>> {
+ let client = self.build_client()?;
+ self.model().guard_max_concurrent_chunks(&data)?;
+ self.embeddings_inner(&client, data)
+ .await
+ .with_context(|| "Failed to get embeddings")
+ }
+
+ fn patch_chat_completions_body(&self, body: &mut Value) {
let model_name = self.model().name();
if let Some(patch_data) = select_model_patch(self.patches_config(), model_name) {
- if body.is_object() && patch_data.request_body.is_object() {
- json_patch::merge(body, &patch_data.request_body)
+ if body.is_object() && patch_data.chat_completions_body.is_object() {
+ json_patch::merge(body, &patch_data.chat_completions_body)
}
}
}
@@ -356,6 +384,14 @@ pub trait Client: Sync + Send {
handler: &mut SseHandler,
data: ChatCompletionsData,
) -> Result<()>;
+
+ async fn embeddings_inner(
+ &self,
+ _client: &ReqwestClient,
+ _data: EmbeddingsData,
+ ) -> Result<Vec<Vec<f32>>> {
+ bail!("No embeddings api")
+ }
}
impl Default for ClientConfig {
@@ -375,7 +411,7 @@ pub type ModelPatches = IndexMap<String, ModelPatch>;
#[derive(Debug, Clone, Deserialize)]
pub struct ModelPatch {
#[serde(default)]
- pub request_body: Value,
+ pub chat_completions_body: Value,
}
pub fn select_model_patch<'a>(
@@ -421,6 +457,20 @@ impl ChatCompletionsOutput {
}
}
+#[derive(Debug)]
+pub struct EmbeddingsData {
+ pub texts: Vec<String>,
+ pub query: bool,
+}
+
+impl EmbeddingsData {
+ pub fn new(texts: Vec<String>, query: bool) -> Self {
+ Self { texts, query }
+ }
+}
+
+pub type EmbeddingsOutput = Vec<Vec<f32>>;
+
pub type PromptAction<'a> = (&'a str, &'a str, bool, PromptKind);
pub fn create_config(prompts: &[PromptAction], client: &str) -> Result<(String, Value)> {
@@ -445,7 +495,7 @@ pub fn create_openai_compatible_client_config(client: &str) -> Result<Option<(St
"name": name,
"api_base": api_base,
});
- let prompts = if ALL_CLIENT_MODELS.iter().any(|v| &v.platform == name) {
+ let prompts = if ALL_MODELS.iter().any(|v| &v.platform == name) {
vec![("api_key", "API Key:", false, PromptKind::String)]
} else {
vec![
diff --git a/src/client/ernie.rs b/src/client/ernie.rs
index 097ee68..77f4741 100644
--- a/src/client/ernie.rs
+++ b/src/client/ernie.rs
@@ -1,8 +1,5 @@
-use super::{
- access_token::*, maybe_catch_error, patch_system_message, sse_stream, ChatCompletionsData,
- ChatCompletionsOutput, Client, ErnieClient, ExtraConfig, Model, ModelData, ModelPatches,
- PromptAction, PromptKind, SseHandler, SseMmessage,
-};
+use super::*;
+use super::access_token::*;
use anyhow::{anyhow, Context, Result};
use async_trait::async_trait;
@@ -37,7 +34,7 @@ impl ErnieClient {
data: ChatCompletionsData,
) -> Result<RequestBuilder> {
let mut body = build_chat_completions_body(data, &self.model);
- self.patch_request_body(&mut body);
+ self.patch_chat_completions_body(&mut body);
let access_token = get_access_token(self.name())?;
diff --git a/src/client/gemini.rs b/src/client/gemini.rs
index 5cc45c5..03eef7a 100644
--- a/src/client/gemini.rs
+++ b/src/client/gemini.rs
@@ -1,11 +1,10 @@
-use super::{
- vertexai::*, ChatCompletionsData, Client, ExtraConfig, GeminiClient, Model, ModelData,
- ModelPatches, PromptAction, PromptKind,
-};
+use super::vertexai::*;
+use super::*;
-use anyhow::Result;
+use anyhow::{Context, Result};
use reqwest::{Client as ReqwestClient, RequestBuilder};
use serde::Deserialize;
+use serde_json::{json, Value};
const API_BASE: &str = "https://generativelanguage.googleapis.com/v1beta/models/";
@@ -38,13 +37,41 @@ impl GeminiClient {
};
let mut body = gemini_build_chat_completions_body(data, &self.model)?;
- self.patch_request_body(&mut body);
+ self.patch_chat_completions_body(&mut body);
- let model = &self.model.name();
+ let url = format!("{API_BASE}{}:{}?key={}", &self.model.name(), func, api_key);
- let url = format!("{API_BASE}{}:{}?key={}", model, func, api_key);
+ debug!("Gemini Chat Completions Request: {url} {body}");
- debug!("Gemini Request: {url} {body}");
+ let builder = client.post(url).json(&body);
+
+ Ok(builder)
+ }
+
+ fn embeddings_builder(
+ &self,
+ client: &ReqwestClient,
+ data: EmbeddingsData,
+ ) -> Result<RequestBuilder> {
+ let api_key = self.get_api_key()?;
+
+ let body = json!({
+ "content": {
+ "parts": [
+ {
+ "text": data.texts[0],
+ }
+ ]
+ }
+ });
+
+ let url = format!(
+ "{API_BASE}{}:embedContent?key={}",
+ &self.model.name(),
+ api_key
+ );
+
+ debug!("Gemini Embeddings Request: {url} {body}");
let builder = client.post(url).json(&body);
@@ -54,6 +81,30 @@ impl GeminiClient {
impl_client_trait!(
GeminiClient,
- crate::client::vertexai::gemini_chat_completions,
- crate::client::vertexai::gemini_chat_completions_streaming
+ gemini_chat_completions,
+ gemini_chat_completions_streaming,
+ gemini_embeddings
);
+
+async fn gemini_embeddings(builder: RequestBuilder) -> Result<EmbeddingsOutput> {
+ let res = builder.send().await?;
+ let status = res.status();
+ let data: Value = res.json().await?;
+ if !status.is_success() {
+ catch_error(&data, status.as_u16())?;
+ }
+ let res_body: EmbeddingsResBody =
+ serde_json::from_value(data).context("Invalid request data")?;
+ let output = vec![res_body.embedding.values];
+ Ok(output)
+}
+
+#[derive(Deserialize)]
+struct EmbeddingsResBody {
+ embedding: EmbeddingsResBodyEmbedding,
+}
+
+#[derive(Deserialize)]
+struct EmbeddingsResBodyEmbedding {
+ values: Vec<f32>,
+}
diff --git a/src/client/model.rs b/src/client/model.rs
index 65e4143..e16cb4e 100644
--- a/src/client/model.rs
+++ b/src/client/model.rs
@@ -1,4 +1,7 @@
-use super::message::{Message, MessageContent};
+use super::{
+ message::{Message, MessageContent},
+ EmbeddingsData,
+};
use crate::utils::{estimate_token_length, format_option_value};
@@ -81,6 +84,10 @@ impl Model {
&self.data.name
}
+ pub fn mode(&self) -> &str {
+ &self.data.mode
+ }
+
pub fn data(&self) -> &ModelData {
&self.data
}
@@ -137,6 +144,14 @@ impl Model {
self.data.supports_function_calling
}
+ pub fn default_chunk_size(&self) -> usize {
+ self.data.default_chunk_size.unwrap_or(1000)
+ }
+
+ pub fn max_concurrent_chunks(&self) -> usize {
+ self.data.max_concurrent_chunks.unwrap_or(1)
+ }
+
pub fn max_tokens_param(&self) -> Option<isize> {
if self.data.pass_max_tokens {
self.data.max_output_tokens
@@ -182,30 +197,45 @@ impl Model {
}
}
- pub fn max_input_tokens_limit(&self, messages: &[Message]) -> Result<()> {
+ pub fn guard_max_input_tokens(&self, messages: &[Message]) -> Result<()> {
let total_tokens = self.total_tokens(messages) + BASIS_TOKENS;
if let Some(max_input_tokens) = self.data.max_input_tokens {
if total_tokens >= max_input_tokens {
- bail!("Exceed max input tokens limit")
+ bail!("Exceed max_input_tokens limit")
}
}
Ok(())
}
+
+ pub fn guard_max_concurrent_chunks(&self, data: &EmbeddingsData) -> Result<()> {
+ if data.texts.len() > self.max_concurrent_chunks() {
+ bail!("Exceed max_concurrent_chunks limit");
+ }
+ Ok(())
+ }
}
#[derive(Debug, Clone, Default, Deserialize)]
pub struct ModelData {
pub name: String,
+ #[serde(default = "default_model_mode")]
+ pub mode: String,
pub max_input_tokens: Option<usize>,
+ pub input_price: Option<f64>,
+ pub output_price: Option<f64>,
+
+ // chat-only properties
pub max_output_tokens: Option<isize>,
#[serde(default)]
pub pass_max_tokens: bool,
- pub input_price: Option<f64>,
- pub output_price: Option<f64>,
#[serde(default)]
pub supports_vision: bool,
#[serde(default)]
pub supports_function_calling: bool,
+
+ // embedding-only properties
+ pub default_chunk_size: Option<usize>,
+ pub max_concurrent_chunks: Option<usize>,
}
impl ModelData {
@@ -222,3 +252,7 @@ pub struct BuiltinModels {
pub platform: String,
pub models: Vec<ModelData>,
}
+
+fn default_model_mode() -> String {
+ "chat".into()
+}
diff --git a/src/client/ollama.rs b/src/client/ollama.rs
index beba8a1..f9bf8d5 100644
--- a/src/client/ollama.rs
+++ b/src/client/ollama.rs
@@ -1,10 +1,6 @@
-use super::{
- catch_error, json_stream, message::*, ChatCompletionsData, ChatCompletionsOutput, Client,
- ExtraConfig, Model, ModelData, ModelPatches, OllamaClient, PromptAction, PromptKind,
- SseHandler,
-};
+use super::*;
-use anyhow::{anyhow, bail, Result};
+use anyhow::{anyhow, bail, Context, Result};
use reqwest::{Client as ReqwestClient, RequestBuilder};
use serde::Deserialize;
use serde_json::{json, Value};
@@ -14,7 +10,7 @@ pub struct OllamaConfig {
pub name: Option<String>,
pub api_base: Option<String>,
pub api_auth: Option<String>,
- pub chat_endpoint: Option<String>,
+ #[serde(default)]
pub models: Vec<ModelData>,
pub patches: Option<ModelPatches>,
pub extra: Option<ExtraConfig>,
@@ -45,13 +41,36 @@ impl OllamaClient {
let api_auth = self.get_api_auth().ok();
let mut body = 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.chat_endpoint.as_deref().unwrap_or("/api/chat");
+ let url = format!("{api_base}/api/chat");
- let url = format!("{api_base}{chat_endpoint}");
+ debug!("Ollama Chat Completions Request: {url} {body}");
- debug!("Ollama Request: {url} {body}");
+ let mut builder = client.post(url).json(&body);
+ if let Some(api_auth) = api_auth {
+ builder = builder.header("Authorization", api_auth)
+ }
+
+ Ok(builder)
+ }
+
+ fn embeddings_builder(
+ &self,
+ client: &ReqwestClient,
+ data: EmbeddingsData,
+ ) -> Result<RequestBuilder> {
+ let api_base = self.get_api_base()?;
+ let api_auth = self.get_api_auth().ok();
+
+ let body = json!({
+ "model": self.model.name(),
+ "prompt": data.texts[0],
+ });
+
+ let url = format!("{api_base}/api/embeddings");
+
+ debug!("Ollama Embeddings Request: {url} {body}");
let mut builder = client.post(url).json(&body);
if let Some(api_auth) = api_auth {
@@ -62,7 +81,12 @@ impl OllamaClient {
}
}
-impl_client_trait!(OllamaClient, chat_completions, chat_completions_streaming);
+impl_client_trait!(
+ OllamaClient,
+ chat_completions,
+ chat_completions_streaming,
+ embeddings
+);
async fn chat_completions(builder: RequestBuilder) -> Result<ChatCompletionsOutput> {
let res = builder.send().await?;
@@ -109,6 +133,25 @@ async fn chat_completions_streaming(
Ok(())
}
+async fn embeddings(
+ builder: RequestBuilder,
+) -> Result<EmbeddingsOutput> {
+ let res = builder.send().await?;
+ let status = res.status();
+ let data = res.json().await?;
+ if !status.is_success() {
+ catch_error(&data, status.as_u16())?;
+ }
+ let res_body: EmbeddingsResBody = serde_json::from_value(data).context("Invalid request data")?;
+ let output = vec![res_body.embedding];
+ Ok(output)
+}
+
+#[derive(Deserialize)]
+struct EmbeddingsResBody {
+ embedding: Vec<f32>,
+}
+
fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Result<Value> {
let ChatCompletionsData {
messages,
diff --git a/src/client/openai.rs b/src/client/openai.rs
index 3cdea24..0da8166 100644
--- a/src/client/openai.rs
+++ b/src/client/openai.rs
@@ -1,10 +1,6 @@
-use super::{
- catch_error, message::*, sse_stream, ChatCompletionsData, ChatCompletionsOutput, Client,
- ExtraConfig, Model, ModelData, ModelPatches, OpenAIClient, PromptAction, PromptKind,
- SseHandler, SseMmessage, ToolCall,
-};
+use super::*;
-use anyhow::{bail, Result};
+use anyhow::{bail, Context, Result};
use reqwest::{Client as ReqwestClient, RequestBuilder};
use serde::Deserialize;
use serde_json::{json, Value};
@@ -39,11 +35,11 @@ impl OpenAIClient {
let api_base = self.get_api_base().unwrap_or_else(|_| API_BASE.to_string());
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!("{api_base}/chat/completions");
- debug!("OpenAI Request: {url} {body}");
+ debug!("OpenAI Chat Completions Request: {url} {body}");
let mut builder = client.post(url).bearer_auth(api_key).json(&body);
@@ -53,6 +49,25 @@ impl OpenAIClient {
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().unwrap_or_else(|_| API_BASE.to_string());
+
+ let body = openai_build_embeddings_body(data, &self.model);
+
+ let url = format!("{api_base}/embeddings");
+
+ debug!("OpenAI Embeddings Request: {url} {body}");
+
+ let builder = client.post(url).bearer_auth(api_key).json(&body);
+
+ Ok(builder)
+ }
}
pub async fn openai_chat_completions(builder: RequestBuilder) -> Result<ChatCompletionsOutput> {
@@ -125,6 +140,30 @@ pub async fn openai_chat_completions_streaming(
sse_stream(builder, handle).await
}
+pub async fn openai_embeddings(
+ builder: RequestBuilder,
+) -> Result<EmbeddingsOutput> {
+ let res = builder.send().await?;
+ let status = res.status();
+ let data: Value = res.json().await?;
+ if !status.is_success() {
+ catch_error(&data, status.as_u16())?;
+ }
+ let res_body: EmbeddingsResBody = serde_json::from_value(data).context("Invalid request data")?;
+ let output = res_body.data.into_iter().map(|v| v.embedding).collect();
+ Ok(output)
+}
+
+#[derive(Deserialize)]
+struct EmbeddingsResBody {
+ data: Vec<EmbeddingsResBodyEmbedding>,
+}
+
+#[derive(Deserialize)]
+struct EmbeddingsResBodyEmbedding {
+ embedding: Vec<f32>,
+}
+
pub fn openai_build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Value {
let ChatCompletionsData {
messages,
@@ -201,6 +240,15 @@ pub fn openai_build_chat_completions_body(data: ChatCompletionsData, model: &Mod
body
}
+
+pub fn openai_build_embeddings_body(data: EmbeddingsData, model: &Model) -> Value {
+ json!({
+ "input": data.texts,
+ "model": model.name(),
+ "encoding_format": "float",
+ })
+}
+
pub fn openai_extract_chat_completions(data: &Value) -> Result<ChatCompletionsOutput> {
let text = data["choices"][0]["message"]["content"]
.as_str()
@@ -244,5 +292,6 @@ pub fn openai_extract_chat_completions(data: &Value) -> Result<ChatCompletionsOu
impl_client_trait!(
OpenAIClient,
openai_chat_completions,
- openai_chat_completions_streaming
+ openai_chat_completions_streaming,
+ openai_embeddings
);
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
);
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" {
diff --git a/src/client/replicate.rs b/src/client/replicate.rs
index 92c7e18..e96ed64 100644
--- a/src/client/replicate.rs
+++ b/src/client/replicate.rs
@@ -1,8 +1,5 @@
-use super::{
- catch_error, prompt_format::*, sse_stream, ChatCompletionsData, ChatCompletionsOutput, Client,
- ExtraConfig, Model, ModelData, ModelPatches, PromptAction, PromptKind, ReplicateClient,
- SseHandler, SseMmessage,
-};
+use super::*;
+use super::prompt_format::*;
use anyhow::{anyhow, Result};
use async_trait::async_trait;
@@ -36,7 +33,7 @@ impl ReplicateClient {
api_key: &str,
) -> Result<RequestBuilder> {
let mut body = build_chat_completions_body(data, &self.model)?;
- self.patch_request_body(&mut body);
+ self.patch_chat_completions_body(&mut body);
let url = format!("{API_BASE}/models/{}/predictions", self.model.name());
diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs
index b40247d..a9f84f8 100644
--- a/src/client/vertexai.rs
+++ b/src/client/vertexai.rs
@@ -1,8 +1,5 @@
-use super::{
- access_token::*, catch_error, json_stream, message::*, patch_system_message,
- ChatCompletionsData, ChatCompletionsOutput, Client, ExtraConfig, Model, ModelData,
- ModelPatches, PromptAction, PromptKind, SseHandler, ToolCall, VertexAIClient,
-};
+use super::*;
+use super::access_token::*;
use anyhow::{anyhow, bail, Context, Result};
use async_trait::async_trait;
@@ -51,9 +48,37 @@ impl VertexAIClient {
let url = format!("{base_url}/google/models/{}:{func}", self.model.name());
let mut body = gemini_build_chat_completions_body(data, &self.model)?;
- self.patch_request_body(&mut body);
+ self.patch_chat_completions_body(&mut body);
- debug!("VertexAI Request: {url} {body}");
+ debug!("VertexAI Chat Completions Request: {url} {body}");
+
+ let builder = client.post(url).bearer_auth(access_token).json(&body);
+
+ Ok(builder)
+ }
+
+ fn embeddings_builder(
+ &self,
+ client: &ReqwestClient,
+ data: EmbeddingsData,
+ ) -> Result<RequestBuilder> {
+ let project_id = self.get_project_id()?;
+ let location = self.get_location()?;
+ let access_token = get_access_token(self.name())?;
+
+ let base_url = format!("https://{location}-aiplatform.googleapis.com/v1/projects/{project_id}/locations/{location}/publishers");
+ let url = format!("{base_url}/google/models/{}:predict", self.model.name());
+
+ let task_type = match data.query {
+ true => "RETRIEVAL_DOCUMENT",
+ false => "QUESTION_ANSWERING",
+ };
+ let instances: Vec<_> = data.texts.into_iter().map(|v| json!({"task_type": task_type, "content": v})).collect();
+ let body = json!({
+ "instances": instances,
+ });
+
+ debug!("VertexAI Embeddings Request: {url} {body}");
let builder = client.post(url).bearer_auth(access_token).json(&body);
@@ -85,6 +110,16 @@ impl Client for VertexAIClient {
let builder = self.chat_completions_builder(client, data)?;
gemini_chat_completions_streaming(builder, handler).await
}
+
+ async fn embeddings_inner(
+ &self,
+ client: &ReqwestClient,
+ data: EmbeddingsData,
+ ) -> Result<Vec<Vec<f32>>> {
+ prepare_gcloud_access_token(client, self.name(), &self.config.adc_file).await?;
+ let builder = self.embeddings_builder(client, data)?;
+ embeddings(builder).await
+ }
}
pub async fn gemini_chat_completions(builder: RequestBuilder) -> Result<ChatCompletionsOutput> {
@@ -138,6 +173,34 @@ pub async fn gemini_chat_completions_streaming(
Ok(())
}
+async fn embeddings(builder: RequestBuilder) -> Result<EmbeddingsOutput> {
+ let res = builder.send().await?;
+ let status = res.status();
+ let data: Value = res.json().await?;
+ if !status.is_success() {
+ catch_error(&data, status.as_u16())?;
+ }
+ let res_body: EmbeddingsResBody =
+ serde_json::from_value(data).context("Invalid request data")?;
+ let output = res_body.predictions.into_iter().map(|v| v.embeddings.values).collect();
+ Ok(output)
+}
+
+#[derive(Deserialize)]
+struct EmbeddingsResBody {
+ predictions: Vec<EmbeddingsResBodyPrediction>,
+}
+
+#[derive(Deserialize)]
+struct EmbeddingsResBodyPrediction {
+ embeddings: EmbeddingsResBodyPredictionEmbeddings,
+}
+
+#[derive(Deserialize)]
+struct EmbeddingsResBodyPredictionEmbeddings {
+ values: Vec<f32>
+}
+
fn gemini_extract_chat_completions_text(data: &Value) -> Result<ChatCompletionsOutput> {
let text = data["candidates"][0]["content"]["parts"][0]["text"]
.as_str()
@@ -179,7 +242,7 @@ fn gemini_extract_chat_completions_text(data: &Value) -> Result<ChatCompletionsO
Ok(output)
}
-pub(crate) fn gemini_build_chat_completions_body(
+pub fn gemini_build_chat_completions_body(
data: ChatCompletionsData,
model: &Model,
) -> Result<Value> {
diff --git a/src/client/vertexai_claude.rs b/src/client/vertexai_claude.rs
index bdce7d8..3993078 100644
--- a/src/client/vertexai_claude.rs
+++ b/src/client/vertexai_claude.rs
@@ -1,8 +1,7 @@
-use super::{
- access_token::*, claude::*, vertexai::*, ChatCompletionsData, ChatCompletionsOutput, Client,
- ExtraConfig, Model, ModelData, ModelPatches, PromptAction, PromptKind, SseHandler,
- VertexAIClaudeClient,
-};
+use super::*;
+use super::access_token::*;
+use super::claude::*;
+use super::vertexai::*;
use anyhow::Result;
use async_trait::async_trait;
@@ -46,7 +45,7 @@ impl VertexAIClaudeClient {
);
let mut body = claude_build_chat_completions_body(data, &self.model)?;
- self.patch_request_body(&mut body);
+ self.patch_chat_completions_body(&mut body);
if let Some(body_obj) = body.as_object_mut() {
body_obj.remove("model");
}