summaryrefslogtreecommitdiffstats
path: root/src/client/ollama.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-05-18 19:06:21 +0800
committerGitHub <noreply@github.com>2024-05-18 19:06:21 +0800
commitb4a40e3fedb438570770a224b890ea24f6e660a9 (patch)
tree344b96102da7cbedf1034d023aa82599940388b1 /src/client/ollama.rs
parent1348a62e5f8bc140a7218fbfe1b73f990ab16101 (diff)
downloadaichat-b4a40e3fedb438570770a224b890ea24f6e660a9.tar.gz
feat: support function calling (#514)
* feat: support function calling * fix on Windows OS * implement multi-steps function calling * fix on Windows OS * add error for client not support function calling * refactor message data structure and make claude client supporting function calling * support reuse previous call results * improve error handling for function calling * use prefix `may_` as indicator for `execute` type fucntions
Diffstat (limited to 'src/client/ollama.rs')
-rw-r--r--src/client/ollama.rs23
1 files changed, 18 insertions, 5 deletions
diff --git a/src/client/ollama.rs b/src/client/ollama.rs
index 6408d2e..24d2dfe 100644
--- a/src/client/ollama.rs
+++ b/src/client/ollama.rs
@@ -1,5 +1,5 @@
use super::{
- catch_error, message::*, CompletionDetails, ExtraConfig, Model, ModelConfig, OllamaClient,
+ catch_error, message::*, CompletionOutput, ExtraConfig, Model, ModelData, OllamaClient,
PromptAction, PromptKind, SendData, SseHandler,
};
@@ -15,7 +15,7 @@ pub struct OllamaConfig {
pub api_base: Option<String>,
pub api_auth: Option<String>,
pub chat_endpoint: Option<String>,
- pub models: Vec<ModelConfig>,
+ pub models: Vec<ModelData>,
pub extra: Option<ExtraConfig>,
}
@@ -59,17 +59,18 @@ impl OllamaClient {
impl_client_trait!(OllamaClient, send_message, send_message_streaming);
-async fn send_message(builder: RequestBuilder) -> Result<(String, CompletionDetails)> {
+async fn send_message(builder: RequestBuilder) -> Result<CompletionOutput> {
let res = builder.send().await?;
let status = res.status();
let data = res.json().await?;
if !status.is_success() {
catch_error(&data, status.as_u16())?;
}
+ debug!("non-stream-data: {data}");
let text = data["message"]["content"]
.as_str()
.ok_or_else(|| anyhow!("Invalid response data: {data}"))?;
- Ok((text.to_string(), CompletionDetails::default()))
+ Ok(CompletionOutput::new(text))
}
async fn send_message_streaming(builder: RequestBuilder, handler: &mut SseHandler) -> Result<()> {
@@ -86,6 +87,7 @@ async fn send_message_streaming(builder: RequestBuilder, handler: &mut SseHandle
continue;
}
let data: Value = serde_json::from_slice(&chunk)?;
+ debug!("stream-data: {data}");
if data["done"].is_boolean() {
if let Some(text) = data["message"]["content"].as_str() {
handler.text(text)?;
@@ -103,10 +105,13 @@ fn build_body(data: SendData, model: &Model) -> Result<Value> {
messages,
temperature,
top_p,
+ functions: _,
stream,
} = data;
+ let mut is_tool_call = false;
let mut network_image_urls = vec![];
+
let messages: Vec<Value> = messages
.into_iter()
.map(|message| {
@@ -141,10 +146,18 @@ fn build_body(data: SendData, model: &Model) -> Result<Value> {
let content = content.join("\n\n");
json!({ "role": role, "content": content, "images": images })
}
+ MessageContent::ToolResults(_) => {
+ is_tool_call = true;
+ json!({ "role": role })
+ }
}
})
.collect();
+ if is_tool_call {
+ bail!("The client does not support function calling",);
+ }
+
if !network_image_urls.is_empty() {
bail!(
"The model does not support network images: {:?}",
@@ -153,7 +166,7 @@ fn build_body(data: SendData, model: &Model) -> Result<Value> {
}
let mut body = json!({
- "model": &model.name,
+ "model": &model.name(),
"messages": messages,
"stream": stream,
"options": {},