diff options
| author | sigoden <sigoden@gmail.com> | 2024-06-05 09:02:23 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-06-05 09:02:23 +0800 |
| commit | 1ec6abfaee2fdc189b348b7e3a8145bd9a84da74 (patch) | |
| tree | 6f68a860a39b0fbe87784de925b5a31e00e74e33 /src/client | |
| parent | 71f2e94579511d7524f5534377001ab3f02a9597 (diff) | |
| download | aichat-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.rs | 31 | ||||
| -rw-r--r-- | src/client/bedrock.rs | 10 | ||||
| -rw-r--r-- | src/client/claude.rs | 9 | ||||
| -rw-r--r-- | src/client/cloudflare.rs | 7 | ||||
| -rw-r--r-- | src/client/cohere.rs | 69 | ||||
| -rw-r--r-- | src/client/common.rs | 140 | ||||
| -rw-r--r-- | src/client/ernie.rs | 9 | ||||
| -rw-r--r-- | src/client/gemini.rs | 73 | ||||
| -rw-r--r-- | src/client/model.rs | 44 | ||||
| -rw-r--r-- | src/client/ollama.rs | 67 | ||||
| -rw-r--r-- | src/client/openai.rs | 67 | ||||
| -rw-r--r-- | src/client/openai_compatible.rs | 73 | ||||
| -rw-r--r-- | src/client/qianwen.rs | 86 | ||||
| -rw-r--r-- | src/client/replicate.rs | 9 | ||||
| -rw-r--r-- | src/client/vertexai.rs | 79 | ||||
| -rw-r--r-- | src/client/vertexai_claude.rs | 11 |
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"); } |
