summaryrefslogtreecommitdiffstats
path: root/src/client/ollama.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-01-13 19:52:07 +0800
committerGitHub <noreply@github.com>2024-01-13 19:52:07 +0800
commitfe35cfd9419302f01baf9672493c0b0a4b41d889 (patch)
tree94e763745fb7989c97af39cc1dfb44250440a5eb /src/client/ollama.rs
parent4e99df4c1bd4028a77251bdb00ff23c664372b5f (diff)
downloadaichat-fe35cfd9419302f01baf9672493c0b0a4b41d889.tar.gz
feat: supports model capabilities (#297)
1. automatically switch to the model that has the necessary capabilities. 2. throw an error if the client does not have a model with the necessary capabilities
Diffstat (limited to 'src/client/ollama.rs')
-rw-r--r--src/client/ollama.rs19
1 files changed, 6 insertions, 13 deletions
diff --git a/src/client/ollama.rs b/src/client/ollama.rs
index fc6e148..0705f2e 100644
--- a/src/client/ollama.rs
+++ b/src/client/ollama.rs
@@ -1,9 +1,9 @@
use super::{
- message::*, patch_system_message, Client, ExtraConfig, Model, OllamaClient, PromptType,
- SendData, TokensCountFactors,
+ message::*, patch_system_message, Client, ExtraConfig, Model, ModelConfig, OllamaClient,
+ PromptType, SendData, TokensCountFactors,
};
-use crate::{config::GlobalConfig, render::ReplyHandler, utils::PromptKind};
+use crate::{render::ReplyHandler, utils::PromptKind};
use anyhow::{anyhow, bail, Result};
use async_trait::async_trait;
@@ -20,21 +20,13 @@ pub struct OllamaConfig {
pub api_base: String,
pub api_key: Option<String>,
pub chat_endpoint: Option<String>,
- pub models: Vec<LocalAIModel>,
+ pub models: Vec<ModelConfig>,
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)
- }
+ client_common_fns!();
async fn send_message_inner(&self, client: &ReqwestClient, data: SendData) -> Result<String> {
let builder = self.request_builder(client, data)?;
@@ -75,6 +67,7 @@ impl OllamaClient {
.iter()
.map(|v| {
Model::new(client_name, &v.name)
+ .set_capabilities(v.capabilities)
.set_max_tokens(v.max_tokens)
.set_tokens_count_factors(TOKENS_COUNT_FACTORS)
})