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/ollama.rs | |
| 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/ollama.rs')
| -rw-r--r-- | src/client/ollama.rs | 67 |
1 files changed, 55 insertions, 12 deletions
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, |
