From 0e740d81e94505bd57036755abaaecb12c3b26e3 Mon Sep 17 00:00:00 2001 From: sigoden Date: Sun, 28 Jul 2024 06:04:36 +0800 Subject: feat: abandon rag_dedicated client and improve (#757) --- src/client/ollama.rs | 75 ++++++++++++++++++++++++++++++---------------------- 1 file changed, 43 insertions(+), 32 deletions(-) (limited to 'src/client/ollama.rs') diff --git a/src/client/ollama.rs b/src/client/ollama.rs index 5c26b99..4c2f344 100644 --- a/src/client/ollama.rs +++ b/src/client/ollama.rs @@ -31,53 +31,63 @@ impl OllamaClient { PromptKind::Integer, ), ]; +} - fn prepare_chat_completions(&self, data: ChatCompletionsData) -> Result { - let api_base = self.get_api_base()?; - let api_auth = self.get_api_auth().ok(); +impl_client_trait!( + OllamaClient, + ( + prepare_chat_completions, + chat_completions, + chat_completions_streaming + ), + (prepare_embeddings, embeddings), + (noop_prepare_rerank, noop_rerank), +); - let url = format!("{api_base}/api/chat"); +fn prepare_chat_completions( + self_: &OllamaClient, + data: ChatCompletionsData, +) -> Result { + let api_base = self_.get_api_base()?; + let api_auth = self_.get_api_auth().ok(); - let body = build_chat_completions_body(data, &self.model)?; + let url = format!("{api_base}/api/chat"); - let mut request_data = RequestData::new(url, body); + let body = build_chat_completions_body(data, &self_.model)?; - if let Some(api_auth) = api_auth { - request_data.header("Authorization", api_auth) - } + let mut request_data = RequestData::new(url, body); - Ok(request_data) + if let Some(api_auth) = api_auth { + request_data.header("Authorization", api_auth) } - fn prepare_embeddings(&self, data: EmbeddingsData) -> Result { - let api_base = self.get_api_base()?; - let api_auth = self.get_api_auth().ok(); + Ok(request_data) +} - let url = format!("{api_base}/api/embed"); +fn prepare_embeddings(self_: &OllamaClient, data: EmbeddingsData) -> Result { + let api_base = self_.get_api_base()?; + let api_auth = self_.get_api_auth().ok(); - let body = json!({ - "model": self.model.name(), - "input": data.texts, - }); + let url = format!("{api_base}/api/embed"); - let mut request_data = RequestData::new(url, body); + let body = json!({ + "model": self_.model.name(), + "input": data.texts, + }); - if let Some(api_auth) = api_auth { - request_data.header("Authorization", api_auth) - } + let mut request_data = RequestData::new(url, body); - Ok(request_data) + if let Some(api_auth) = api_auth { + request_data.header("Authorization", api_auth) } -} -impl_client_trait!( - OllamaClient, - chat_completions, - chat_completions_streaming, - embeddings -); + Ok(request_data) +} -async fn chat_completions(builder: RequestBuilder) -> Result { +async fn chat_completions( + builder: RequestBuilder, + _model: &Model, +) -> Result { let res = builder.send().await?; let status = res.status(); let data = res.json().await?; @@ -92,6 +102,7 @@ async fn chat_completions(builder: RequestBuilder) -> Result Result<()> { let res = builder.send().await?; let status = res.status(); @@ -120,7 +131,7 @@ async fn chat_completions_streaming( Ok(()) } -async fn embeddings(builder: RequestBuilder) -> Result { +async fn embeddings(builder: RequestBuilder, _model: &Model) -> Result { let res = builder.send().await?; let status = res.status(); let data = res.json().await?; -- cgit v1.2.3