summaryrefslogtreecommitdiffstats
path: root/src/client/ollama.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-04-23 16:46:48 +0800
committerGitHub <noreply@github.com>2024-04-23 16:46:48 +0800
commitd1aafa11153ab689c21c2c57c47da52337d8e8d1 (patch)
treedc20dc033e9d376aab09941835a842b22fe32c02 /src/client/ollama.rs
parent1cc89eff514669d6459f62cc4eed8e0b34d8ef0c (diff)
downloadaichat-d1aafa11153ab689c21c2c57c47da52337d8e8d1.tar.gz
feat: customize model's max_output_tokens (#428)
Diffstat (limited to 'src/client/ollama.rs')
-rw-r--r--src/client/ollama.rs31
1 files changed, 12 insertions, 19 deletions
diff --git a/src/client/ollama.rs b/src/client/ollama.rs
index e652634..1692942 100644
--- a/src/client/ollama.rs
+++ b/src/client/ollama.rs
@@ -1,6 +1,6 @@
use super::{
- message::*, patch_system_message, Client, ExtraConfig, Model, ModelConfig, OllamaClient,
- PromptType, ReplyHandler, SendData,
+ convert_models, message::*, patch_system_message, Client, ExtraConfig, Model, ModelConfig,
+ OllamaClient, PromptType, ReplyHandler, SendData,
};
use crate::utils::PromptKind;
@@ -59,23 +59,13 @@ impl OllamaClient {
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_capabilities(v.capabilities)
- .set_max_input_tokens(v.max_input_tokens)
- .set_extra_fields(v.extra_fields.clone())
- })
- .collect()
+ convert_models(client_name, &local_config.models)
}
fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result<RequestBuilder> {
let api_key = self.get_api_key().ok();
- let mut body = build_body(data, self.model.name.clone())?;
+ let mut body = build_body(data, &self.model)?;
self.model.merge_extra_fields(&mut body);
let chat_endpoint = self.config.chat_endpoint.as_deref().unwrap_or("/api/chat");
@@ -133,7 +123,7 @@ async fn send_message_streaming(builder: RequestBuilder, handler: &mut ReplyHand
Ok(())
}
-fn build_body(data: SendData, model: String) -> Result<Value> {
+fn build_body(data: SendData, model: &Model) -> Result<Value> {
let SendData {
mut messages,
temperature,
@@ -189,15 +179,18 @@ fn build_body(data: SendData, model: String) -> Result<Value> {
}
let mut body = json!({
- "model": model,
+ "model": &model.name,
"messages": messages,
"stream": stream,
+ "options": {},
});
+ if let Some(num_predict) = model.max_output_tokens {
+ body["options"]["num_predict"] = num_predict.into();
+ }
+
if let Some(temperature) = temperature {
- body["options"] = json!({
- "temperature": temperature,
- });
+ body["options"]["temperature"] = temperature.into();
}
Ok(body)