summaryrefslogtreecommitdiffstats
path: root/src/client/ollama.rs
diff options
context:
space:
mode:
Diffstat (limited to 'src/client/ollama.rs')
-rw-r--r--src/client/ollama.rs67
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,