summaryrefslogtreecommitdiffstats
path: root/src/client
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2023-12-20 07:49:05 +0800
committerGitHub <noreply@github.com>2023-12-20 07:49:05 +0800
commit64c4edf7c8af771d1d7a229507187eae2e3d03d3 (patch)
tree7b36820dedc67e7342c3634afffe705760f2c306 /src/client
parent6c9d7a679ea0097ab12f295caf16bec3e75ee267 (diff)
downloadaichat-64c4edf7c8af771d1d7a229507187eae2e3d03d3.tar.gz
feat: support ollama (#276)
Diffstat (limited to 'src/client')
-rw-r--r--src/client/mod.rs3
-rw-r--r--src/client/ollama.rs209
2 files changed, 211 insertions, 1 deletions
diff --git a/src/client/mod.rs b/src/client/mod.rs
index d2ada44..7c157ef 100644
--- a/src/client/mod.rs
+++ b/src/client/mod.rs
@@ -9,14 +9,15 @@ pub use model::*;
register_client!(
(openai, "openai", OpenAIConfig, OpenAIClient),
+ (gemini, "gemini", GeminiConfig, GeminiClient),
(localai, "localai", LocalAIConfig, LocalAIClient),
+ (ollama, "ollama", OllamaConfig, OllamaClient),
(
azure_openai,
"azure-openai",
AzureOpenAIConfig,
AzureOpenAIClient
),
- (gemini, "gemini", GeminiConfig, GeminiClient),
(ernie, "ernie", ErnieConfig, ErnieClient),
(qianwen, "qianwen", QianwenConfig, QianwenClient),
);
diff --git a/src/client/ollama.rs b/src/client/ollama.rs
new file mode 100644
index 0000000..91e7bb2
--- /dev/null
+++ b/src/client/ollama.rs
@@ -0,0 +1,209 @@
+use super::{
+ message::*, patch_system_message, Client, ExtraConfig, Model, OllamaClient, PromptType,
+ SendData, TokensCountFactors,
+};
+
+use crate::{config::GlobalConfig, render::ReplyHandler, utils::PromptKind};
+
+use anyhow::{anyhow, bail, Result};
+use async_trait::async_trait;
+use futures_util::StreamExt;
+use reqwest::{Client as ReqwestClient, RequestBuilder};
+use serde::Deserialize;
+use serde_json::{json, Value};
+
+const TOKENS_COUNT_FACTORS: TokensCountFactors = (5, 2);
+
+#[derive(Debug, Clone, Deserialize, Default)]
+pub struct OllamaConfig {
+ pub name: Option<String>,
+ pub api_base: String,
+ pub api_key: Option<String>,
+ pub chat_endpoint: Option<String>,
+ pub models: Vec<LocalAIModel>,
+ pub extra: Option<ExtraConfig>,
+}
+
+#[derive(Debug, Clone, Deserialize)]
+pub struct LocalAIModel {
+ name: String,
+ max_tokens: Option<usize>,
+}
+
+#[async_trait]
+impl Client for OllamaClient {
+ fn config(&self) -> (&GlobalConfig, &Option<ExtraConfig>) {
+ (&self.global_config, &self.config.extra)
+ }
+
+ async fn send_message_inner(&self, client: &ReqwestClient, data: SendData) -> Result<String> {
+ let builder = self.request_builder(client, data)?;
+ send_message(builder).await
+ }
+
+ async fn send_message_streaming_inner(
+ &self,
+ client: &ReqwestClient,
+ handler: &mut ReplyHandler,
+ data: SendData,
+ ) -> Result<()> {
+ let builder = self.request_builder(client, data)?;
+ send_message_streaming(builder, handler).await
+ }
+}
+
+impl OllamaClient {
+ config_get_fn!(api_key, get_api_key);
+
+ pub const PROMPTS: [PromptType<'static>; 4] = [
+ ("api_base", "API Base:", true, PromptKind::String),
+ ("api_key", "API Key:", false, PromptKind::String),
+ ("models[].name", "Model Name:", true, PromptKind::String),
+ (
+ "models[].max_tokens",
+ "Max Tokens:",
+ false,
+ PromptKind::Integer,
+ ),
+ ];
+
+ pub fn list_models(local_config: &OllamaConfig) -> Vec<Model> {
+ let client_name = Self::name(local_config);
+
+ local_config
+ .models
+ .iter()
+ .map(|v| {
+ Model::new(client_name, &v.name)
+ .set_max_tokens(v.max_tokens)
+ .set_tokens_count_factors(TOKENS_COUNT_FACTORS)
+ })
+ .collect()
+ }
+
+ fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result<RequestBuilder> {
+ let api_key = self.get_api_key().ok();
+
+ let body = build_body(data, self.model.name.clone())?;
+
+ let chat_endpoint = self.config.chat_endpoint.as_deref().unwrap_or("/chat");
+
+ let url = format!("{}{chat_endpoint}", self.config.api_base);
+
+ debug!("Ollama Request: {url} {body}");
+
+ let mut builder = client.post(url).json(&body);
+ if let Some(api_key) = api_key {
+ builder = builder.header("Authorization", api_key)
+ }
+
+ Ok(builder)
+ }
+}
+
+async fn send_message(builder: RequestBuilder) -> Result<String> {
+ let res = builder.send().await?;
+ let status = res.status();
+ if status != 200 {
+ let text = res.text().await?;
+ bail!("{status}, {text}");
+ }
+ let data: Value = res.json().await?;
+ let output = data["message"]["content"]
+ .as_str()
+ .ok_or_else(|| anyhow!("Invalid response data: {data}"))?;
+ Ok(output.to_string())
+}
+
+async fn send_message_streaming(builder: RequestBuilder, handler: &mut ReplyHandler) -> Result<()> {
+ let res = builder.send().await?;
+ let status = res.status();
+ if status != 200 {
+ let text = res.text().await?;
+ bail!("{status}, {text}");
+ } else {
+ let mut stream = res.bytes_stream();
+ while let Some(chunk) = stream.next().await {
+ let chunk = chunk?;
+ let data: Value = serde_json::from_slice(&chunk)?;
+ if data["done"].is_boolean() {
+ if let Some(text) = data["message"]["content"].as_str() {
+ handler.text(text)?;
+ }
+ } else {
+ bail!("Invalid response data: {data}")
+ }
+ }
+ }
+ Ok(())
+}
+
+fn build_body(data: SendData, model: String) -> Result<Value> {
+ let SendData {
+ mut messages,
+ temperature,
+ stream,
+ } = data;
+
+ patch_system_message(&mut messages);
+
+ let mut network_image_urls = vec![];
+ let messages: Vec<Value> = messages
+ .into_iter()
+ .map(|message| {
+ let role = message.role;
+ match message.content {
+ MessageContent::Text(text) => json!({
+ "role": role,
+ "content": text,
+ }),
+ MessageContent::Array(list) => {
+ let mut content = vec![];
+ let mut images = vec![];
+ for item in list {
+ match item {
+ MessageContentPart::Text { text } => {
+ content.push(text);
+ }
+ MessageContentPart::ImageUrl {
+ image_url: ImageUrl { url },
+ } => {
+ if let Some((_, data)) = url
+ .strip_prefix("data:")
+ .and_then(|v| v.split_once(";base64,"))
+ {
+ images.push(data.to_string());
+ } else {
+ network_image_urls.push(url.clone());
+ }
+ }
+ }
+ }
+ let content = content.join("\n\n");
+ json!({ "role": role, "content": content, "images": images })
+ }
+ }
+ })
+ .collect();
+
+ if !network_image_urls.is_empty() {
+ bail!(
+ "The model does not support network images: {:?}",
+ network_image_urls
+ );
+ }
+
+ let mut body = json!({
+ "model": model,
+ "messages": messages,
+ "stream": stream,
+ });
+
+ if let Some(temperature) = temperature {
+ body["options"] = json!({
+ "temperature": temperature,
+ });
+ }
+
+ Ok(body)
+}