summaryrefslogtreecommitdiffstats
path: root/src/client
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
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')
-rw-r--r--src/client/azure_openai.rs9
-rw-r--r--src/client/bedrock.rs35
-rw-r--r--src/client/claude.rs233
-rw-r--r--src/client/cloudflare.rs17
-rw-r--r--src/client/cohere.rs121
-rw-r--r--src/client/common.rs68
-rw-r--r--src/client/ernie.rs23
-rw-r--r--src/client/gemini.rs6
-rw-r--r--src/client/message.rs23
-rw-r--r--src/client/mod.rs1
-rw-r--r--src/client/model.rs196
-rw-r--r--src/client/ollama.rs23
-rw-r--r--src/client/openai.rs137
-rw-r--r--src/client/openai_compatible.rs6
-rw-r--r--src/client/prompt_format.rs1
-rw-r--r--src/client/qianwen.rs37
-rw-r--r--src/client/replicate.rs26
-rw-r--r--src/client/sse_handler.rs23
-rw-r--r--src/client/vertexai.rs172
-rw-r--r--src/client/vertexai_claude.rs8
20 files changed, 811 insertions, 354 deletions
diff --git a/src/client/azure_openai.rs b/src/client/azure_openai.rs
index d586921..95a2960 100644
--- a/src/client/azure_openai.rs
+++ b/src/client/azure_openai.rs
@@ -1,7 +1,5 @@
use super::openai::openai_build_body;
-use super::{
- AzureOpenAIClient, ExtraConfig, Model, ModelConfig, PromptAction, PromptKind, SendData,
-};
+use super::{AzureOpenAIClient, ExtraConfig, Model, ModelData, PromptAction, PromptKind, SendData};
use anyhow::Result;
use reqwest::{Client as ReqwestClient, RequestBuilder};
@@ -12,7 +10,7 @@ pub struct AzureOpenAIConfig {
pub name: Option<String>,
pub api_base: Option<String>,
pub api_key: Option<String>,
- pub models: Vec<ModelConfig>,
+ pub models: Vec<ModelData>,
pub extra: Option<ExtraConfig>,
}
@@ -41,7 +39,8 @@ impl AzureOpenAIClient {
let url = format!(
"{}/openai/deployments/{}/chat/completions?api-version=2024-02-01",
- &api_base, self.model.name
+ &api_base,
+ self.model.name()
);
debug!("AzureOpenAI Request: {url} {body}");
diff --git a/src/client/bedrock.rs b/src/client/bedrock.rs
index b07152b..8bbb3eb 100644
--- a/src/client/bedrock.rs
+++ b/src/client/bedrock.rs
@@ -1,8 +1,8 @@
use super::claude::{claude_build_body, claude_extract_completion};
use super::{
- catch_error, generate_prompt, BedrockClient, Client, CompletionDetails, ExtraConfig, Model,
- ModelConfig, PromptAction, PromptFormat, PromptKind, SendData, SseHandler,
- LLAMA3_PROMPT_FORMAT, MISTRAL_PROMPT_FORMAT,
+ catch_error, generate_prompt, BedrockClient, Client, CompletionOutput, ExtraConfig, Model,
+ ModelData, PromptAction, PromptFormat, PromptKind, SendData, SseHandler, LLAMA3_PROMPT_FORMAT,
+ MISTRAL_PROMPT_FORMAT,
};
use crate::utils::{base64_decode, encode_uri, hex_encode, hmac_sha256, sha256};
@@ -30,7 +30,7 @@ pub struct BedrockConfig {
pub secret_access_key: Option<String>,
pub region: Option<String>,
#[serde(default)]
- pub models: Vec<ModelConfig>,
+ pub models: Vec<ModelData>,
pub extra: Option<ExtraConfig>,
}
@@ -42,8 +42,8 @@ impl Client for BedrockClient {
&self,
client: &ReqwestClient,
data: SendData,
- ) -> Result<(String, CompletionDetails)> {
- let model_category = ModelCategory::from_str(&self.model.name)?;
+ ) -> Result<CompletionOutput> {
+ let model_category = ModelCategory::from_str(self.model.name())?;
let builder = self.request_builder(client, data, &model_category)?;
send_message(builder, &model_category).await
}
@@ -54,7 +54,7 @@ impl Client for BedrockClient {
handler: &mut SseHandler,
data: SendData,
) -> Result<()> {
- let model_category = ModelCategory::from_str(&self.model.name)?;
+ let model_category = ModelCategory::from_str(self.model.name())?;
let builder = self.request_builder(client, data, &model_category)?;
send_message_streaming(builder, handler, &model_category).await
}
@@ -91,7 +91,7 @@ impl BedrockClient {
let secret_access_key = self.get_secret_access_key()?;
let region = self.get_region()?;
- let model_name = &self.model.name;
+ let model_name = &self.model.name();
let uri = if data.stream {
format!("/model/{model_name}/invoke-with-response-stream")
} else {
@@ -129,7 +129,7 @@ impl BedrockClient {
async fn send_message(
builder: RequestBuilder,
model_category: &ModelCategory,
-) -> Result<(String, CompletionDetails)> {
+) -> Result<CompletionOutput> {
let res = builder.send().await?;
let status = res.status();
let data: Value = res.json().await?;
@@ -138,6 +138,7 @@ async fn send_message(
catch_error(&data, status.as_u16())?;
}
+ debug!("non-stream-data: {data}");
match model_category {
ModelCategory::Anthropic => claude_extract_completion(&data),
ModelCategory::MetaLlama3 => llama_extract_completion(&data),
@@ -172,7 +173,7 @@ async fn send_message_streaming(
let data: Value = decode_chunk(message.payload()).ok_or_else(|| {
anyhow!("Invalid chunk data: {}", hex_encode(message.payload()))
})?;
- // debug!("bedrock chunk: {data}");
+ debug!("stream-data: {data}");
match model_category {
ModelCategory::Anthropic => {
if let Some(typ) = data["type"].as_str() {
@@ -230,6 +231,7 @@ fn meta_llama_build_body(data: SendData, model: &Model, pt: PromptFormat) -> Res
messages,
temperature,
top_p,
+ functions: _,
stream: _,
} = data;
let prompt = generate_prompt(&messages, pt)?;
@@ -253,6 +255,7 @@ fn mistral_build_body(data: SendData, model: &Model) -> Result<Value> {
messages,
temperature,
top_p,
+ functions: _,
stream: _,
} = data;
let prompt = generate_prompt(&messages, MISTRAL_PROMPT_FORMAT)?;
@@ -271,23 +274,25 @@ fn mistral_build_body(data: SendData, model: &Model) -> Result<Value> {
Ok(body)
}
-fn llama_extract_completion(data: &Value) -> Result<(String, CompletionDetails)> {
+fn llama_extract_completion(data: &Value) -> Result<CompletionOutput> {
let text = data["generation"]
.as_str()
.ok_or_else(|| anyhow!("Invalid response data: {data}"))?;
- let details = CompletionDetails {
+ let output = CompletionOutput {
+ text: text.to_string(),
+ tool_calls: vec![],
id: None,
input_tokens: data["prompt_token_count"].as_u64(),
output_tokens: data["generation_token_count"].as_u64(),
};
- Ok((text.to_string(), details))
+ Ok(output)
}
-fn mistral_extract_completion(data: &Value) -> Result<(String, CompletionDetails)> {
+fn mistral_extract_completion(data: &Value) -> Result<CompletionOutput> {
let text = data["outputs"][0]["text"]
.as_str()
.ok_or_else(|| anyhow!("Invalid response data: {data}"))?;
- Ok((text.to_string(), CompletionDetails::default()))
+ Ok(CompletionOutput::new(text))
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
diff --git a/src/client/claude.rs b/src/client/claude.rs
index 89742f3..8dc4dce 100644
--- a/src/client/claude.rs
+++ b/src/client/claude.rs
@@ -1,10 +1,10 @@
use super::{
- catch_error, extract_system_message, sse_stream, ClaudeClient, CompletionDetails, ExtraConfig,
- ImageUrl, MessageContent, MessageContentPart, Model, ModelConfig, PromptAction, PromptKind,
- SendData, SsMmessage, SseHandler,
+ catch_error, extract_system_message, message::*, sse_stream, ClaudeClient, CompletionOutput,
+ ExtraConfig, ImageUrl, MessageContent, MessageContentPart, Model, ModelData, PromptAction,
+ PromptKind, SendData, SsMmessage, SseHandler, ToolCall,
};
-use anyhow::{anyhow, bail, Result};
+use anyhow::{bail, Context, Result};
use reqwest::{Client as ReqwestClient, RequestBuilder};
use serde::Deserialize;
use serde_json::{json, Value};
@@ -16,7 +16,7 @@ pub struct ClaudeConfig {
pub name: Option<String>,
pub api_key: Option<String>,
#[serde(default)]
- pub models: Vec<ModelConfig>,
+ pub models: Vec<ModelData>,
pub extra: Option<ExtraConfig>,
}
@@ -36,7 +36,9 @@ impl ClaudeClient {
debug!("Claude Request: {url} {body}");
let mut builder = client.post(url).json(&body);
- builder = builder.header("anthropic-version", "2023-06-01");
+ builder = builder
+ .header("anthropic-version", "2023-06-01")
+ .header("anthropic-beta", "tools-2024-05-16");
if let Some(api_key) = api_key {
builder = builder.header("x-api-key", api_key)
}
@@ -51,13 +53,14 @@ impl_client_trait!(
claude_send_message_streaming
);
-pub async fn claude_send_message(builder: RequestBuilder) -> Result<(String, CompletionDetails)> {
+pub async fn claude_send_message(builder: RequestBuilder) -> Result<CompletionOutput> {
let res = builder.send().await?;
let status = res.status();
let data: Value = res.json().await?;
if !status.is_success() {
catch_error(&data, status.as_u16())?;
}
+ debug!("non-stream-data: {data}");
claude_extract_completion(&data)
}
@@ -65,13 +68,59 @@ pub async fn claude_send_message_streaming(
builder: RequestBuilder,
handler: &mut SseHandler,
) -> Result<()> {
+ let mut function_name = String::new();
+ let mut function_arguments = String::new();
+ let mut function_id = String::new();
let handle = |message: SsMmessage| -> Result<bool> {
let data: Value = serde_json::from_str(&message.data)?;
+ debug!("stream-data: {data}");
if let Some(typ) = data["type"].as_str() {
- if typ == "content_block_delta" {
- if let Some(text) = data["delta"]["text"].as_str() {
- handler.text(text)?;
+ match typ {
+ "content_block_start" => {
+ if let (Some("tool_use"), Some(name), Some(id)) = (
+ data["content_block"]["type"].as_str(),
+ data["content_block"]["name"].as_str(),
+ data["content_block"]["id"].as_str(),
+ ) {
+ if !function_name.is_empty() {
+ let arguments: Value =
+ function_arguments.parse().with_context(|| {
+ format!("Tool call '{function_name}' is invalid: arguments must be in valid JSON format")
+ })?;
+ handler.tool_call(ToolCall::new(
+ function_name.clone(),
+ arguments,
+ Some(function_id.clone()),
+ ))?;
+ }
+ function_name = name.into();
+ function_arguments.clear();
+ function_id = id.into();
+ }
+ }
+ "content_block_delta" => {
+ if let Some(text) = data["delta"]["text"].as_str() {
+ handler.text(text)?;
+ } else if let (true, Some(partial_json)) = (
+ !function_name.is_empty(),
+ data["delta"]["partial_json"].as_str(),
+ ) {
+ function_arguments.push_str(partial_json);
+ }
}
+ "content_block_stop" => {
+ if !function_name.is_empty() {
+ let arguments: Value = function_arguments.parse().with_context(|| {
+ format!("Tool call '{function_name}' is invalid: arguments must be in valid JSON format")
+ })?;
+ handler.tool_call(ToolCall::new(
+ function_name.clone(),
+ arguments,
+ Some(function_id.clone()),
+ ))?;
+ }
+ }
+ _ => {}
}
}
Ok(false)
@@ -85,46 +134,91 @@ pub fn claude_build_body(data: SendData, model: &Model) -> Result<Value> {
mut messages,
temperature,
top_p,
+ functions,
stream,
} = data;
let system_message = extract_system_message(&mut messages);
let mut network_image_urls = vec![];
+
let messages: Vec<Value> = messages
.into_iter()
- .map(|message| {
- let role = message.role;
- let content = match message.content {
- MessageContent::Text(text) => vec![json!({"type": "text", "text": text})],
- MessageContent::Array(list) => list
- .into_iter()
- .map(|item| match item {
- MessageContentPart::Text { text } => json!({"type": "text", "text": text}),
- MessageContentPart::ImageUrl {
- image_url: ImageUrl { url },
- } => {
- if let Some((mime_type, data)) = url
- .strip_prefix("data:")
- .and_then(|v| v.split_once(";base64,"))
- {
- json!({
- "type": "image",
- "source": {
- "type": "base64",
- "media_type": mime_type,
- "data": data,
- }
- })
- } else {
- network_image_urls.push(url.clone());
- json!({ "url": url })
+ .flat_map(|message| {
+ let Message { role, content } = message;
+ match content {
+ MessageContent::Text(text) => vec![json!({
+ "role": role,
+ "content": text,
+ })],
+ MessageContent::Array(list) => {
+ let content: Vec<_> = list
+ .into_iter()
+ .map(|item| match item {
+ MessageContentPart::Text { text } => {
+ json!({"type": "text", "text": text})
}
- }
- })
- .collect(),
- };
- json!({ "role": role, "content": content })
+ MessageContentPart::ImageUrl {
+ image_url: ImageUrl { url },
+ } => {
+ if let Some((mime_type, data)) = url
+ .strip_prefix("data:")
+ .and_then(|v| v.split_once(";base64,"))
+ {
+ json!({
+ "type": "image",
+ "source": {
+ "type": "base64",
+ "media_type": mime_type,
+ "data": data,
+ }
+ })
+ } else {
+ network_image_urls.push(url.clone());
+ json!({ "url": url })
+ }
+ }
+ })
+ .collect();
+ vec![json!({
+ "role": role,
+ "content": content,
+ })]
+ }
+ MessageContent::ToolResults((tool_call_results, text)) => {
+ let mut tool_call = vec![];
+ let mut tool_result = vec![];
+ if !text.is_empty() {
+ tool_call.push(json!({
+ "type": "text",
+ "text": text,
+ }))
+ }
+ for tool_call_result in tool_call_results {
+ tool_call.push(json!({
+ "type": "tool_use",
+ "id": tool_call_result.call.id,
+ "name": tool_call_result.call.name,
+ "input": tool_call_result.call.arguments,
+ }));
+ tool_result.push(json!({
+ "type": "tool_result",
+ "tool_use_id": tool_call_result.call.id,
+ "content": tool_call_result.output.to_string(),
+ }));
+ }
+ vec![
+ json!({
+ "role": "assistant",
+ "content": tool_call,
+ }),
+ json!({
+ "role": "user",
+ "content": tool_result,
+ }),
+ ]
+ }
+ }
})
.collect();
@@ -136,7 +230,7 @@ pub fn claude_build_body(data: SendData, model: &Model) -> Result<Value> {
}
let mut body = json!({
- "model": &model.name,
+ "model": model.name(),
"messages": messages,
});
if let Some(v) = system_message {
@@ -154,18 +248,61 @@ pub fn claude_build_body(data: SendData, model: &Model) -> Result<Value> {
if stream {
body["stream"] = true.into();
}
+ if let Some(functions) = functions {
+ body["tools"] = functions
+ .iter()
+ .map(|v| {
+ json!({
+ "name": v.name,
+ "description": v.description,
+ "input_schema": v.parameters,
+ })
+ })
+ .collect();
+ }
Ok(body)
}
-pub fn claude_extract_completion(data: &Value) -> Result<(String, CompletionDetails)> {
- let text = data["content"][0]["text"]
- .as_str()
- .ok_or_else(|| anyhow!("Invalid response data: {data}"))?;
+pub fn claude_extract_completion(data: &Value) -> Result<CompletionOutput> {
+ let text = data["content"][0]["text"].as_str().unwrap_or_default();
+
+ let mut tool_calls = vec![];
+ if let Some(calls) = data["content"].as_array().map(|content| {
+ content
+ .iter()
+ .filter(|content| matches!(content["type"].as_str(), Some("tool_use")))
+ .collect::<Vec<&Value>>()
+ }) {
+ tool_calls = calls
+ .into_iter()
+ .filter_map(|call| {
+ if let (Some(name), Some(input), Some(id)) = (
+ call["name"].as_str(),
+ call.get("input"),
+ call["id"].as_str(),
+ ) {
+ Some(ToolCall::new(
+ name.to_string(),
+ input.clone(),
+ Some(id.to_string()),
+ ))
+ } else {
+ None
+ }
+ })
+ .collect();
+ };
+
+ if text.is_empty() && tool_calls.is_empty() {
+ bail!("Invalid response data: {data}");
+ }
- let details = CompletionDetails {
+ let output = CompletionOutput {
+ text: text.to_string(),
+ tool_calls,
id: data["id"].as_str().map(|v| v.to_string()),
input_tokens: data["usage"]["input_tokens"].as_u64(),
output_tokens: data["usage"]["output_tokens"].as_u64(),
};
- Ok((text.to_string(), details))
+ Ok(output)
}
diff --git a/src/client/cloudflare.rs b/src/client/cloudflare.rs
index 5a4bf8c..dfff009 100644
--- a/src/client/cloudflare.rs
+++ b/src/client/cloudflare.rs
@@ -1,5 +1,5 @@
use super::{
- catch_error, sse_stream, CloudflareClient, CompletionDetails, ExtraConfig, Model, ModelConfig,
+ catch_error, sse_stream, CloudflareClient, CompletionOutput, ExtraConfig, Model, ModelData,
PromptAction, PromptKind, SendData, SsMmessage, SseHandler,
};
@@ -16,7 +16,7 @@ pub struct CloudflareConfig {
pub account_id: Option<String>,
pub api_key: Option<String>,
#[serde(default)]
- pub models: Vec<ModelConfig>,
+ pub models: Vec<ModelData>,
pub extra: Option<ExtraConfig>,
}
@@ -37,7 +37,7 @@ impl CloudflareClient {
let url = format!(
"{API_BASE}/accounts/{account_id}/ai/run/{}",
- self.model.name
+ self.model.name()
);
debug!("Cloudflare Request: {url} {body}");
@@ -50,7 +50,7 @@ impl CloudflareClient {
impl_client_trait!(CloudflareClient, 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: Value = res.json().await?;
@@ -58,6 +58,7 @@ async fn send_message(builder: RequestBuilder) -> Result<(String, CompletionDeta
catch_error(&data, status.as_u16())?;
}
+ debug!("non-stream-data: {data}");
extract_completion(&data)
}
@@ -67,6 +68,7 @@ async fn send_message_streaming(builder: RequestBuilder, handler: &mut SseHandle
return Ok(true);
}
let data: Value = serde_json::from_str(&message.data)?;
+ debug!("stream-data: {data}");
if let Some(text) = data["response"].as_str() {
handler.text(text)?;
}
@@ -80,11 +82,12 @@ fn build_body(data: SendData, model: &Model) -> Result<Value> {
messages,
temperature,
top_p,
+ functions: _,
stream,
} = data;
let mut body = json!({
- "model": &model.name,
+ "model": &model.name(),
"messages": messages,
});
@@ -104,10 +107,10 @@ fn build_body(data: SendData, model: &Model) -> Result<Value> {
Ok(body)
}
-fn extract_completion(data: &Value) -> Result<(String, CompletionDetails)> {
+fn extract_completion(data: &Value) -> Result<CompletionOutput> {
let text = data["result"]["response"]
.as_str()
.ok_or_else(|| anyhow!("Invalid response data: {data}"))?;
- Ok((text.to_string(), CompletionDetails::default()))
+ Ok(CompletionOutput::new(text))
}
diff --git a/src/client/cohere.rs b/src/client/cohere.rs
index b5d6647..41b0e4b 100644
--- a/src/client/cohere.rs
+++ b/src/client/cohere.rs
@@ -1,9 +1,9 @@
use super::{
- catch_error, extract_system_message, json_stream, message::*, CohereClient, CompletionDetails,
- ExtraConfig, Model, ModelConfig, PromptAction, PromptKind, SendData, SseHandler,
+ catch_error, extract_system_message, json_stream, message::*, CohereClient, CompletionOutput,
+ ExtraConfig, Model, ModelData, PromptAction, PromptKind, SendData, SseHandler, ToolCall,
};
-use anyhow::{anyhow, bail, Result};
+use anyhow::{bail, Result};
use reqwest::{Client as ReqwestClient, RequestBuilder};
use serde::Deserialize;
use serde_json::{json, Value};
@@ -15,7 +15,7 @@ pub struct CohereConfig {
pub name: Option<String>,
pub api_key: Option<String>,
#[serde(default)]
- pub models: Vec<ModelConfig>,
+ pub models: Vec<ModelData>,
pub extra: Option<ExtraConfig>,
}
@@ -42,7 +42,7 @@ impl CohereClient {
impl_client_trait!(CohereClient, 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: Value = res.json().await?;
@@ -50,6 +50,7 @@ async fn send_message(builder: RequestBuilder) -> Result<(String, CompletionDeta
catch_error(&data, status.as_u16())?;
}
+ debug!("non-stream-data: {data}");
extract_completion(&data)
}
@@ -62,10 +63,25 @@ async fn send_message_streaming(builder: RequestBuilder, handler: &mut SseHandle
} else {
let handle = |data: &str| -> Result<()> {
let data: Value = serde_json::from_str(data)?;
+ debug!("stream-data: {data}");
if let Some("text-generation") = data["event_type"].as_str() {
if let Some(text) = data["text"].as_str() {
handler.text(text)?;
}
+ } else if let Some("tool-calls-generation") = data["event_type"].as_str() {
+ if let Some(tool_calls) = data["tool_calls"].as_array() {
+ for call in tool_calls {
+ if let (Some(name), Some(args)) =
+ (call["name"].as_str(), call["parameters"].as_object())
+ {
+ handler.tool_call(ToolCall::new(
+ name.to_string(),
+ json!(args),
+ None,
+ ))?;
+ }
+ }
+ }
}
Ok(())
};
@@ -79,24 +95,28 @@ fn build_body(data: SendData, model: &Model) -> Result<Value> {
mut messages,
temperature,
top_p,
+ functions,
stream,
} = data;
let system_message = extract_system_message(&mut messages);
let mut image_urls = vec![];
+ let mut tool_results = None;
+
let mut messages: Vec<Value> = messages
.into_iter()
- .map(|message| {
- let role = match message.role {
+ .filter_map(|message| {
+ let Message { role, content } = message;
+ let role = match role {
MessageRole::User => "USER",
_ => "CHATBOT",
};
- match message.content {
- MessageContent::Text(text) => json!({
+ match content {
+ MessageContent::Text(text) => Some(json!({
"role": role,
"message": text,
- }),
+ })),
MessageContent::Array(list) => {
let list: Vec<String> = list
.into_iter()
@@ -110,7 +130,11 @@ fn build_body(data: SendData, model: &Model) -> Result<Value> {
}
})
.collect();
- json!({ "role": role, "message": list.join("\n\n") })
+ Some(json!({ "role": role, "message": list.join("\n\n") }))
+ }
+ MessageContent::ToolResults((tool_call_results, _)) => {
+ tool_results = Some(tool_call_results);
+ None
}
}
})
@@ -123,10 +147,29 @@ fn build_body(data: SendData, model: &Model) -> Result<Value> {
let message = message["message"].as_str().unwrap_or_default();
let mut body = json!({
- "model": &model.name,
+ "model": &model.name(),
"message": message,
});
+ if let Some(tool_results) = tool_results {
+ let tool_results: Vec<_> = tool_results
+ .into_iter()
+ .map(|tool_call_result| {
+ json!({
+ "call": {
+ "name": tool_call_result.call.name,
+ "parameters": tool_call_result.call.arguments,
+ },
+ "outputs": [
+ tool_call_result.output,
+ ]
+
+ })
+ })
+ .collect();
+ body["tool_results"] = json!(tool_results);
+ }
+
if let Some(v) = system_message {
body["preamble"] = v.into();
}
@@ -148,18 +191,60 @@ fn build_body(data: SendData, model: &Model) -> Result<Value> {
body["stream"] = true.into();
}
+ if let Some(functions) = functions {
+ body["tools"] = functions
+ .iter()
+ .map(|v| {
+ let required = v.parameters.required.clone().unwrap_or_default();
+ let mut parameter_definitions = json!({});
+ if let Some(properties) = &v.parameters.properties {
+ for (key, value) in properties {
+ let mut value: Value = json!(value);
+ if value.is_object() && required.iter().any(|x| x == key) {
+ value["required"] = true.into();
+ }
+ parameter_definitions[key] = value;
+ }
+ }
+ json!({
+ "name": v.name,
+ "description": v.description,
+ "parameter_definitions": parameter_definitions,
+ })
+ })
+ .collect();
+ }
Ok(body)
}
-fn extract_completion(data: &Value) -> Result<(String, CompletionDetails)> {
- let text = data["text"]
- .as_str()
- .ok_or_else(|| anyhow!("Invalid response data: {data}"))?;
+fn extract_completion(data: &Value) -> Result<CompletionOutput> {
+ let text = data["text"].as_str().unwrap_or_default();
+
+ let mut tool_calls = vec![];
+ if let Some(calls) = data["tool_calls"].as_array() {
+ tool_calls = calls
+ .iter()
+ .filter_map(|call| {
+ if let (Some(name), Some(parameters)) =
+ (call["name"].as_str(), call["parameters"].as_object())
+ {
+ Some(ToolCall::new(name.to_string(), json!(parameters), None))
+ } else {
+ None
+ }
+ })
+ .collect()
+ }
- let details = CompletionDetails {
+ if text.is_empty() && tool_calls.is_empty() {
+ bail!("Invalid response data: {data}");
+ }
+ let output = CompletionOutput {
+ text: text.to_string(),
+ tool_calls,
id: data["generation_id"].as_str().map(|v| v.to_string()),
input_tokens: data["meta"]["billed_units"]["input_tokens"].as_u64(),
output_tokens: data["meta"]["billed_units"]["output_tokens"].as_u64(),
};
- Ok((text.to_string(), details))
+ Ok(output)
}
diff --git a/src/client/common.rs b/src/client/common.rs
index d42315b..4ef276d 100644
--- a/src/client/common.rs
+++ b/src/client/common.rs
@@ -2,6 +2,7 @@ use super::{openai::OpenAIConfig, BuiltinModels, ClientConfig, Message, Model, S
use crate::{
config::{GlobalConfig, Input},
+ function::{eval_tool_calls, FunctionDeclaration, ToolCall, ToolCallResult},
render::{render_error, render_stream},
utils::{prompt_input_integer, prompt_input_string, tokenize, AbortSignal, PromptKind},
};
@@ -52,7 +53,7 @@ macro_rules! register_client {
pub enum ClientModel {
$(
#[serde(rename = $name)]
- $config { models: Vec<ModelConfig> },
+ $config { models: Vec<ModelData> },
)+
#[serde(other)]
Unknown,
@@ -73,7 +74,7 @@ macro_rules! register_client {
pub fn init(global_config: &$crate::config::GlobalConfig, model: &$crate::client::Model) -> Option<Box<dyn Client>> {
let config = global_config.read().clients.iter().find_map(|client_config| {
if let ClientConfig::$config(c) = client_config {
- if Self::name(c) == &model.client_name {
+ if Self::name(c) == model.client_name() {
return Some(c.clone())
}
}
@@ -113,24 +114,10 @@ macro_rules! register_client {
None
$(.or_else(|| $client::init(config, &model)))+
.ok_or_else(|| {
- anyhow::anyhow!("Unknown client '{}'", model.client_name)
+ anyhow::anyhow!("Unknown client '{}'", model.client_name())
})
}
- pub fn ensure_model_capabilities(client: &mut dyn Client, capabilities: $crate::client::ModelCapabilities) -> anyhow::Result<()> {
- if !client.model().capabilities.contains(capabilities) {
- let models = client.list_models();
- if let Some(model) = models.into_iter().find(|v| v.capabilities.contains(capabilities)) {
- client.set_model(model);
- } else {
- anyhow::bail!(
- "The current model is incapable of doing that."
- );
- }
- }
- Ok(())
- }
-
pub fn list_client_types() -> Vec<&'static str> {
let mut client_types: Vec<_> = vec![$($client::NAME,)+];
client_types.extend($crate::client::OPENAI_COMPATIBLE_PLATFORMS.iter().map(|(name, _)| *name));
@@ -213,7 +200,7 @@ macro_rules! impl_client_trait {
&self,
client: &reqwest::Client,
data: $crate::client::SendData,
- ) -> anyhow::Result<(String, $crate::client::CompletionDetails)> {
+ ) -> anyhow::Result<$crate::client::CompletionOutput> {
let builder = self.request_builder(client, data)?;
$send_message(builder).await
}
@@ -261,14 +248,16 @@ macro_rules! unsupported_model {
pub trait Client: Sync + Send {
fn config(&self) -> (&GlobalConfig, &Option<ExtraConfig>);
- fn list_models(&self) -> Vec<Model>;
-
fn name(&self) -> &str;
+ #[allow(unused)]
+ fn list_models(&self) -> Vec<Model>;
+
fn model(&self) -> &Model;
fn model_mut(&mut self) -> &mut Model;
+ #[allow(unused)]
fn set_model(&mut self, model: Model);
fn build_client(&self) -> Result<ReqwestClient> {
@@ -287,14 +276,15 @@ pub trait Client: Sync + Send {
Ok(client)
}
- async fn send_message(&self, input: Input) -> Result<(String, CompletionDetails)> {
+ async fn send_message(&self, input: Input) -> Result<CompletionOutput> {
let global_config = self.config().0;
if global_config.read().dry_run {
let content = input.echo_messages();
- return Ok((content, CompletionDetails::default()));
+ return Ok(CompletionOutput::new(&content));
}
let client = self.build_client()?;
- let data = input.prepare_send_data(false)?;
+
+ let data = input.prepare_send_data(self.model(), false)?;
self.send_message_inner(&client, data)
.await
.with_context(|| "Failed to get answer")
@@ -324,7 +314,7 @@ pub trait Client: Sync + Send {
return Ok(());
}
let client = self.build_client()?;
- let data = input.prepare_send_data(true)?;
+ let data = input.prepare_send_data(self.model(), true)?;
self.send_message_streaming_inner(&client, handler, data).await
} => {
handler.done()?;
@@ -341,7 +331,7 @@ pub trait Client: Sync + Send {
&self,
client: &ReqwestClient,
data: SendData,
- ) -> Result<(String, CompletionDetails)>;
+ ) -> Result<CompletionOutput>;
async fn send_message_streaming_inner(
&self,
@@ -368,16 +358,28 @@ pub struct SendData {
pub messages: Vec<Message>,
pub temperature: Option<f64>,
pub top_p: Option<f64>,
+ pub functions: Option<Vec<FunctionDeclaration>>,
pub stream: bool,
}
#[derive(Debug, Clone, Default)]
-pub struct CompletionDetails {
+pub struct CompletionOutput {
+ pub text: String,
+ pub tool_calls: Vec<ToolCall>,
pub id: Option<String>,
pub input_tokens: Option<u64>,
pub output_tokens: Option<u64>,
}
+impl CompletionOutput {
+ pub fn new(text: &str) -> Self {
+ Self {
+ text: text.to_string(),
+ ..Default::default()
+ }
+ }
+}
+
pub type PromptAction<'a> = (&'a str, &'a str, bool, PromptKind);
pub fn create_config(prompts: &[PromptAction], client: &str) -> Result<(String, Value)> {
@@ -429,22 +431,24 @@ pub async fn send_stream(
client: &dyn Client,
config: &GlobalConfig,
abort: AbortSignal,
-) -> Result<String> {
+) -> Result<(String, Vec<ToolCallResult>)> {
let (tx, rx) = unbounded_channel();
- let mut stream_handler = SseHandler::new(tx, abort.clone());
+ let mut handler = SseHandler::new(tx, abort.clone());
let (send_ret, rend_ret) = tokio::join!(
- client.send_message_streaming(input, &mut stream_handler),
+ client.send_message_streaming(input, &mut handler),
render_stream(rx, config, abort.clone()),
);
if let Err(err) = rend_ret {
render_error(err, config.read().highlight);
}
- let output = stream_handler.get_buffer().to_string();
+ let (output, calls) = handler.take();
match send_ret {
Ok(_) => {
- println!();
- Ok(output)
+ if !output.is_empty() && !output.ends_with('\n') {
+ println!();
+ }
+ Ok((output, eval_tool_calls(config, calls)?))
}
Err(err) => {
if !output.is_empty() {
diff --git a/src/client/ernie.rs b/src/client/ernie.rs
index 28cb857..1d79138 100644
--- a/src/client/ernie.rs
+++ b/src/client/ernie.rs
@@ -1,7 +1,7 @@
use super::access_token::*;
use super::{
- maybe_catch_error, patch_system_message, sse_stream, Client, CompletionDetails, ErnieClient,
- ExtraConfig, Model, ModelConfig, PromptAction, PromptKind, SendData, SsMmessage, SseHandler,
+ maybe_catch_error, patch_system_message, sse_stream, Client, CompletionOutput, ErnieClient,
+ ExtraConfig, Model, ModelData, PromptAction, PromptKind, SendData, SsMmessage, SseHandler,
};
use anyhow::{anyhow, Context, Result};
@@ -20,7 +20,7 @@ pub struct ErnieConfig {
pub api_key: Option<String>,
pub secret_key: Option<String>,
#[serde(default)]
- pub models: Vec<ModelConfig>,
+ pub models: Vec<ModelData>,
pub extra: Option<ExtraConfig>,
}
@@ -36,7 +36,7 @@ impl ErnieClient {
let url = format!(
"{API_BASE}/wenxinworkshop/chat/{}?access_token={access_token}",
- &self.model.name,
+ &self.model.name(),
);
debug!("Ernie Request: {url} {body}");
@@ -78,7 +78,7 @@ impl Client for ErnieClient {
&self,
client: &ReqwestClient,
data: SendData,
- ) -> Result<(String, CompletionDetails)> {
+ ) -> Result<CompletionOutput> {
self.prepare_access_token().await?;
let builder = self.request_builder(client, data)?;
send_message(builder).await
@@ -96,15 +96,17 @@ impl Client for ErnieClient {
}
}
-async fn send_message(builder: RequestBuilder) -> Result<(String, CompletionDetails)> {
+async fn send_message(builder: RequestBuilder) -> Result<CompletionOutput> {
let data: Value = builder.send().await?.json().await?;
maybe_catch_error(&data)?;
+ debug!("non-stream-data: {data}");
extract_completion_text(&data)
}
async fn send_message_streaming(builder: RequestBuilder, handler: &mut SseHandler) -> Result<()> {
let handle = |message: SsMmessage| -> Result<bool> {
let data: Value = serde_json::from_str(&message.data)?;
+ debug!("stream-data: {data}");
if let Some(text) = data["result"].as_str() {
handler.text(text)?;
}
@@ -119,6 +121,7 @@ fn build_body(data: SendData, model: &Model) -> Value {
mut messages,
temperature,
top_p,
+ functions: _,
stream,
} = data;
@@ -145,16 +148,18 @@ fn build_body(data: SendData, model: &Model) -> Value {
body
}
-fn extract_completion_text(data: &Value) -> Result<(String, CompletionDetails)> {
+fn extract_completion_text(data: &Value) -> Result<CompletionOutput> {
let text = data["result"]
.as_str()
.ok_or_else(|| anyhow!("Invalid response data: {data}"))?;
- let details = CompletionDetails {
+ let output = CompletionOutput {
+ text: text.to_string(),
+ tool_calls: vec![],
id: data["id"].as_str().map(|v| v.to_string()),
input_tokens: data["usage"]["prompt_tokens"].as_u64(),
output_tokens: data["usage"]["completion_tokens"].as_u64(),
};
- Ok((text.to_string(), details))
+ Ok(output)
}
async fn fetch_access_token(
diff --git a/src/client/gemini.rs b/src/client/gemini.rs
index 28fc93e..1ffb9a9 100644
--- a/src/client/gemini.rs
+++ b/src/client/gemini.rs
@@ -1,5 +1,5 @@
use super::vertexai::gemini_build_body;
-use super::{ExtraConfig, GeminiClient, Model, ModelConfig, PromptAction, PromptKind, SendData};
+use super::{ExtraConfig, GeminiClient, Model, ModelData, PromptAction, PromptKind, SendData};
use anyhow::Result;
use reqwest::{Client as ReqwestClient, RequestBuilder};
@@ -14,7 +14,7 @@ pub struct GeminiConfig {
#[serde(rename = "safetySettings")]
pub safety_settings: Option<serde_json::Value>,
#[serde(default)]
- pub models: Vec<ModelConfig>,
+ pub models: Vec<ModelData>,
pub extra: Option<ExtraConfig>,
}
@@ -34,7 +34,7 @@ impl GeminiClient {
let body = gemini_build_body(data, &self.model, self.config.safety_settings.clone())?;
- let model = &self.model.name;
+ let model = &self.model.name();
let url = format!("{API_BASE}{}:{}?key={}", model, func, api_key);
diff --git a/src/client/message.rs b/src/client/message.rs
index 9621811..d7ba698 100644
--- a/src/client/message.rs
+++ b/src/client/message.rs
@@ -1,4 +1,4 @@
-use crate::config::Input;
+use super::ToolResults;
use serde::{Deserialize, Serialize};
@@ -8,15 +8,21 @@ pub struct Message {
pub content: MessageContent,
}
-impl Message {
- pub fn new(input: &Input) -> Self {
+impl Default for Message {
+ fn default() -> Self {
Self {
role: MessageRole::User,
- content: input.to_message_content(),
+ content: MessageContent::Text(String::new()),
}
}
}
+impl Message {
+ pub fn new(role: MessageRole, content: MessageContent) -> Self {
+ Self { role, content }
+ }
+}
+
#[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize, Serialize)]
#[serde(rename_all = "snake_case")]
pub enum MessageRole {
@@ -34,10 +40,6 @@ impl MessageRole {
pub fn is_user(&self) -> bool {
matches!(self, MessageRole::User)
}
-
- pub fn is_assistant(&self) -> bool {
- matches!(self, MessageRole::Assistant)
- }
}
#[derive(Debug, Clone, Deserialize, Serialize)]
@@ -45,6 +47,8 @@ impl MessageRole {
pub enum MessageContent {
Text(String),
Array(Vec<MessageContentPart>),
+ // Note: This type is primarily for convenience and does not exist in OpenAI's API.
+ ToolResults(ToolResults),
}
impl MessageContent {
@@ -68,6 +72,7 @@ impl MessageContent {
}
format!(".file {}{}", files.join(" "), concated_text)
}
+ MessageContent::ToolResults(_) => String::new(),
}
}
@@ -83,6 +88,7 @@ impl MessageContent {
*text = replace_fn(text)
}
}
+ MessageContent::ToolResults(_) => {}
}
}
@@ -98,6 +104,7 @@ impl MessageContent {
}
parts.join("\n\n")
}
+ MessageContent::ToolResults(_) => String::new(),
}
}
}
diff --git a/src/client/mod.rs b/src/client/mod.rs
index 7f902a1..acfc39c 100644
--- a/src/client/mod.rs
+++ b/src/client/mod.rs
@@ -6,6 +6,7 @@ mod model;
mod prompt_format;
mod sse_handler;
+pub use crate::function::{ToolCall, ToolResults};
pub use crate::utils::PromptKind;
pub use common::*;
pub use message::*;
diff --git a/src/client/model.rs b/src/client/model.rs
index 1af3e59..4b16ffd 100644
--- a/src/client/model.rs
+++ b/src/client/model.rs
@@ -10,15 +10,8 @@ const BASIS_TOKENS: usize = 2;
#[derive(Debug, Clone)]
pub struct Model {
- pub client_name: String,
- pub name: String,
- pub max_input_tokens: Option<usize>,
- pub max_output_tokens: Option<isize>,
- pub pass_max_tokens: bool,
- pub input_price: Option<f64>,
- pub output_price: Option<f64>,
- pub capabilities: ModelCapabilities,
- pub extra_fields: Option<serde_json::Map<String, serde_json::Value>>,
+ client_name: String,
+ data: ModelData,
}
impl Default for Model {
@@ -31,30 +24,16 @@ impl Model {
pub fn new(client_name: &str, name: &str) -> Self {
Self {
client_name: client_name.into(),
- name: name.into(),
- max_input_tokens: None,
- max_output_tokens: None,
- pass_max_tokens: false,
- input_price: None,
- output_price: None,
- capabilities: ModelCapabilities::Text,
- extra_fields: None,
+ data: ModelData::new(name),
}
}
- pub fn from_config(client_name: &str, models: &[ModelConfig]) -> Vec<Self> {
+ pub fn from_config(client_name: &str, models: &[ModelData]) -> Vec<Self> {
models
.iter()
- .map(|v| {
- let mut model = Model::new(client_name, &v.name);
- model
- .set_max_input_tokens(v.max_input_tokens)
- .set_max_tokens(v.max_output_tokens, v.pass_max_tokens)
- .set_input_price(v.input_price)
- .set_output_price(v.output_price)
- .set_supports_vision(v.supports_vision)
- .set_extra_fields(&v.extra_fields);
- model
+ .map(|v| Model {
+ client_name: client_name.to_string(),
+ data: v.clone(),
})
.collect()
}
@@ -77,7 +56,7 @@ impl Model {
model = Some((*found).clone());
} else if let Some(found) = models.iter().find(|v| v.client_name == client_name) {
let mut found = (*found).clone();
- found.name = model_name.to_string();
+ found.data.name = model_name.to_string();
model = Some(found)
}
}
@@ -91,99 +70,101 @@ impl Model {
}
pub fn id(&self) -> String {
- format!("{}:{}", self.client_name, self.name)
+ format!("{}:{}", self.client_name, self.data.name)
+ }
+
+ pub fn client_name(&self) -> &str {
+ &self.client_name
+ }
+
+ pub fn name(&self) -> &str {
+ &self.data.name
+ }
+
+ pub fn data(&self) -> &ModelData {
+ &self.data
+ }
+
+ pub fn data_mut(&mut self) -> &mut ModelData {
+ &mut self.data
}
pub fn description(&self) -> String {
- let max_input_tokens = format_option_value(&self.max_input_tokens);
- let max_output_tokens = format_option_value(&self.max_output_tokens);
- let input_price = format_option_value(&self.input_price);
- let output_price = format_option_value(&self.output_price);
- let vision = if self.capabilities.contains(ModelCapabilities::Vision) {
- "👁"
- } else {
- ""
+ let ModelData {
+ max_input_tokens,
+ max_output_tokens,
+ input_price,
+ output_price,
+ supports_vision,
+ supports_function_calling,
+ ..
+ } = &self.data;
+ let max_input_tokens = format_option_value(max_input_tokens);
+ let max_output_tokens = format_option_value(max_output_tokens);
+ let input_price = format_option_value(input_price);
+ let output_price = format_option_value(output_price);
+ let mut capabilities = vec![];
+ if *supports_vision {
+ capabilities.push('👁');
+ };
+ if *supports_function_calling {
+ capabilities.push('⚒');
};
+ let capabilities: String = capabilities
+ .into_iter()
+ .map(|v| format!("{v} "))
+ .collect::<Vec<String>>()
+ .join("");
format!(
- "{:>8} / {:>8} | {:>6} / {:>6} {}",
- max_input_tokens, max_output_tokens, input_price, output_price, vision
+ "{:>8} / {:>8} | {:>6} / {:>6} {:>6}",
+ max_input_tokens, max_output_tokens, input_price, output_price, capabilities
)
}
+ pub fn max_input_tokens(&self) -> Option<usize> {
+ self.data.max_input_tokens
+ }
+
+ pub fn max_output_tokens(&self) -> Option<isize> {
+ self.data.max_output_tokens
+ }
+
pub fn supports_vision(&self) -> bool {
- self.capabilities.contains(ModelCapabilities::Vision)
+ self.data.supports_vision
+ }
+
+ pub fn supports_function_calling(&self) -> bool {
+ self.data.supports_function_calling
}
pub fn max_tokens_param(&self) -> Option<isize> {
- if self.pass_max_tokens {
- self.max_output_tokens
+ if self.data.pass_max_tokens {
+ self.data.max_output_tokens
} else {
None
}
}
- pub fn set_max_input_tokens(&mut self, max_input_tokens: Option<usize>) -> &mut Self {
- match max_input_tokens {
- None | Some(0) => self.max_input_tokens = None,
- _ => self.max_input_tokens = max_input_tokens,
- }
- self
- }
-
pub fn set_max_tokens(
&mut self,
max_output_tokens: Option<isize>,
pass_max_tokens: bool,
) -> &mut Self {
match max_output_tokens {
- None | Some(0) => self.max_output_tokens = None,
- _ => self.max_output_tokens = max_output_tokens,
- }
- self.pass_max_tokens = pass_max_tokens;
- self
- }
-
- pub fn set_input_price(&mut self, input_price: Option<f64>) -> &mut Self {
- match input_price {
- None => self.input_price = None,
- _ => self.input_price = input_price,
- }
- self
- }
-
- pub fn set_output_price(&mut self, output_price: Option<f64>) -> &mut Self {
- match output_price {
- None => self.output_price = None,
- _ => self.output_price = output_price,
+ None | Some(0) => self.data.max_output_tokens = None,
+ _ => self.data.max_output_tokens = max_output_tokens,
}
- self
- }
-
- pub fn set_supports_vision(&mut self, supports_vision: bool) -> &mut Self {
- if supports_vision {
- self.capabilities |= ModelCapabilities::Vision;
- } else {
- self.capabilities &= !ModelCapabilities::Vision;
- }
- self
- }
-
- pub fn set_extra_fields(
- &mut self,
- extra_fields: &Option<serde_json::Map<String, serde_json::Value>>,
- ) -> &mut Self {
- self.extra_fields.clone_from(extra_fields);
+ self.data.pass_max_tokens = pass_max_tokens;
self
}
pub fn messages_tokens(&self, messages: &[Message]) -> usize {
messages
.iter()
- .map(|v| {
- match &v.content {
- MessageContent::Text(text) => estimate_token_length(text),
- MessageContent::Array(_) => 0, // TODO
- }
+ .map(|v| match &v.content {
+ MessageContent::Text(text) => estimate_token_length(text),
+ MessageContent::Array(_) => 0,
+ MessageContent::ToolResults(_) => 0,
})
.sum()
}
@@ -203,7 +184,7 @@ impl Model {
pub fn max_input_tokens_limit(&self, messages: &[Message]) -> Result<()> {
let total_tokens = self.total_tokens(messages) + BASIS_TOKENS;
- if let Some(max_input_tokens) = self.max_input_tokens {
+ if let Some(max_input_tokens) = self.data.max_input_tokens {
if total_tokens >= max_input_tokens {
bail!("Exceed max input tokens limit")
}
@@ -212,7 +193,7 @@ impl Model {
}
pub fn merge_extra_fields(&self, body: &mut serde_json::Value) {
- if let (Some(body), Some(extra_fields)) = (body.as_object_mut(), &self.extra_fields) {
+ if let (Some(body), Some(extra_fields)) = (body.as_object_mut(), &self.data.extra_fields) {
for (key, extra_field) in extra_fields {
if body.contains_key(key) {
if let (Some(sub_body), Some(extra_field)) =
@@ -232,30 +213,33 @@ impl Model {
}
}
-#[derive(Debug, Clone, Deserialize)]
-pub struct ModelConfig {
+#[derive(Debug, Clone, Default, Deserialize)]
+pub struct ModelData {
pub name: String,
pub max_input_tokens: Option<usize>,
pub max_output_tokens: Option<isize>,
+ #[serde(default)]
+ pub pass_max_tokens: bool,
pub input_price: Option<f64>,
pub output_price: Option<f64>,
#[serde(default)]
pub supports_vision: bool,
#[serde(default)]
- pub pass_max_tokens: bool,
+ pub supports_function_calling: bool,
pub extra_fields: Option<serde_json::Map<String, serde_json::Value>>,
}
+impl ModelData {
+ pub fn new(name: &str) -> Self {
+ Self {
+ name: name.to_string(),
+ ..Default::default()
+ }
+ }
+}
+
#[derive(Debug, Clone, Deserialize)]
pub struct BuiltinModels {
pub platform: String,
- pub models: Vec<ModelConfig>,
-}
-
-bitflags::bitflags! {
- #[derive(Debug, Clone, Copy, PartialEq)]
- pub struct ModelCapabilities: u32 {
- const Text = 0b00000001;
- const Vision = 0b00000010;
- }
+ pub models: Vec<ModelData>,
}
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": {},
diff --git a/src/client/openai.rs b/src/client/openai.rs
index 0b111db..f2fd222 100644
--- a/src/client/openai.rs
+++ b/src/client/openai.rs
@@ -1,9 +1,9 @@
use super::{
- catch_error, sse_stream, CompletionDetails, ExtraConfig, Model, ModelConfig, OpenAIClient,
- PromptAction, PromptKind, SendData, SsMmessage, SseHandler,
+ catch_error, message::*, sse_stream, CompletionOutput, ExtraConfig, Model, ModelData,
+ OpenAIClient, PromptAction, PromptKind, SendData, SsMmessage, SseHandler, ToolCall,
};
-use anyhow::{anyhow, Result};
+use anyhow::{bail, Result};
use reqwest::{Client as ReqwestClient, RequestBuilder};
use serde::Deserialize;
use serde_json::{json, Value};
@@ -17,7 +17,7 @@ pub struct OpenAIConfig {
pub api_base: Option<String>,
pub organization_id: Option<String>,
#[serde(default)]
- pub models: Vec<ModelConfig>,
+ pub models: Vec<ModelData>,
pub extra: Option<ExtraConfig>,
}
@@ -48,7 +48,7 @@ impl OpenAIClient {
}
}
-pub async fn openai_send_message(builder: RequestBuilder) -> Result<(String, CompletionDetails)> {
+pub async fn openai_send_message(builder: RequestBuilder) -> Result<CompletionOutput> {
let res = builder.send().await?;
let status = res.status();
let data: Value = res.json().await?;
@@ -56,6 +56,7 @@ pub async fn openai_send_message(builder: RequestBuilder) -> Result<(String, Com
catch_error(&data, status.as_u16())?;
}
+ debug!("non-stream-data: {data}");
openai_extract_completion(&data)
}
@@ -63,13 +64,53 @@ pub async fn openai_send_message_streaming(
builder: RequestBuilder,
handler: &mut SseHandler,
) -> Result<()> {
+ let mut function_index = 0;
+ let mut function_name = String::new();
+ let mut function_arguments = String::new();
+ let mut function_id = String::new();
let handle = |message: SsMmessage| -> Result<bool> {
if message.data == "[DONE]" {
+ if !function_name.is_empty() {
+ handler.tool_call(ToolCall::new(
+ function_name.clone(),
+ json!(function_arguments),
+ Some(function_id.clone()),
+ ))?;
+ }
return Ok(true);
}
let data: Value = serde_json::from_str(&message.data)?;
+ debug!("stream-data: {data}");
if let Some(text) = data["choices"][0]["delta"]["content"].as_str() {
handler.text(text)?;
+ } else if let (Some(function), index, id) = (
+ data["choices"][0]["delta"]["tool_calls"][0]["function"].as_object(),
+ data["choices"][0]["delta"]["tool_calls"][0]["index"].as_u64(),
+ data["choices"][0]["delta"]["tool_calls"][0]["id"].as_str(),
+ ) {
+ let index = index.unwrap_or_default();
+ if index != function_index {
+ if !function_name.is_empty() {
+ handler.tool_call(ToolCall::new(
+ function_name.clone(),
+ json!(function_arguments),
+ Some(function_id.clone()),
+ ))?;
+ }
+ function_name.clear();
+ function_arguments.clear();
+ function_id.clear();
+ function_index = index;
+ }
+ if let Some(name) = function.get("name").and_then(|v| v.as_str()) {
+ function_name = name.to_string();
+ }
+ if let Some(arguments) = function.get("arguments").and_then(|v| v.as_str()) {
+ function_arguments.push_str(arguments);
+ }
+ if let Some(id) = id {
+ function_id = id.to_string();
+ }
}
Ok(false)
};
@@ -82,11 +123,47 @@ pub fn openai_build_body(data: SendData, model: &Model) -> Value {
messages,
temperature,
top_p,
+ functions,
stream,
} = data;
+ let messages: Vec<Value> = messages
+ .into_iter()
+ .flat_map(|message| {
+ let Message { role, content } = message;
+ match content {
+ MessageContent::ToolResults((tool_call_results, text)) => {
+ let tool_calls: Vec<_> = tool_call_results.iter().map(|tool_call_result| {
+ json!({
+ "id": tool_call_result.call.id,
+ "type": "function",
+ "function": {
+ "name": tool_call_result.call.name,
+ "arguments": tool_call_result.call.arguments,
+ },
+ })
+ }).collect();
+ let mut messages = vec![
+ json!({ "role": MessageRole::Assistant, "content": text, "tool_calls": tool_calls })
+ ];
+ for tool_call_result in tool_call_results {
+ messages.push(
+ json!({
+ "role": "tool",
+ "content": tool_call_result.output.to_string(),
+ "tool_call_id": tool_call_result.call.id,
+ })
+ );
+ }
+ messages
+ },
+ _ => vec![json!({ "role": role, "content": content })]
+ }
+ })
+ .collect();
+
let mut body = json!({
- "model": &model.name,
+ "model": &model.name(),
"messages": messages,
});
@@ -102,19 +179,59 @@ pub fn openai_build_body(data: SendData, model: &Model) -> Value {
if stream {
body["stream"] = true.into();
}
+ if let Some(functions) = functions {
+ body["tools"] = functions
+ .iter()
+ .map(|v| {
+ json!({
+ "type": "function",
+ "function": v,
+ })
+ })
+ .collect();
+ body["tool_choice"] = "auto".into();
+ }
body
}
-pub fn openai_extract_completion(data: &Value) -> Result<(String, CompletionDetails)> {
+pub fn openai_extract_completion(data: &Value) -> Result<CompletionOutput> {
let text = data["choices"][0]["message"]["content"]
.as_str()
- .ok_or_else(|| anyhow!("Invalid response data: {data}"))?;
- let details = CompletionDetails {
+ .unwrap_or_default();
+
+ let mut tool_calls = vec![];
+ if let Some(tools_call) = data["choices"][0]["message"]["tool_calls"].as_array() {
+ tool_calls = tools_call
+ .iter()
+ .filter_map(|call| {
+ if let (Some(name), Some(arguments), Some(id)) = (
+ call["function"]["name"].as_str(),
+ call["function"]["arguments"].as_str(),
+ call["id"].as_str(),
+ ) {
+ Some(ToolCall::new(
+ name.to_string(),
+ json!(arguments),
+ Some(id.to_string()),
+ ))
+ } else {
+ None
+ }
+ })
+ .collect()
+ };
+
+ if text.is_empty() && tool_calls.is_empty() {
+ bail!("Invalid response data: {data}");
+ }
+ let output = CompletionOutput {
+ text: text.to_string(),
+ tool_calls,
id: data["id"].as_str().map(|v| v.to_string()),
input_tokens: data["usage"]["prompt_tokens"].as_u64(),
output_tokens: data["usage"]["completion_tokens"].as_u64(),
};
- Ok((text.to_string(), details))
+ Ok(output)
}
impl_client_trait!(
diff --git a/src/client/openai_compatible.rs b/src/client/openai_compatible.rs
index 6eae77b..38b291c 100644
--- a/src/client/openai_compatible.rs
+++ b/src/client/openai_compatible.rs
@@ -2,7 +2,7 @@ use crate::client::OPENAI_COMPATIBLE_PLATFORMS;
use super::openai::openai_build_body;
use super::{
- ExtraConfig, Model, ModelConfig, OpenAICompatibleClient, PromptAction, PromptKind, SendData,
+ ExtraConfig, Model, ModelData, OpenAICompatibleClient, PromptAction, PromptKind, SendData,
};
use anyhow::Result;
@@ -16,7 +16,7 @@ pub struct OpenAICompatibleConfig {
pub api_key: Option<String>,
pub chat_endpoint: Option<String>,
#[serde(default)]
- pub models: Vec<ModelConfig>,
+ pub models: Vec<ModelData>,
pub extra: Option<ExtraConfig>,
}
@@ -44,7 +44,7 @@ impl OpenAICompatibleClient {
match OPENAI_COMPATIBLE_PLATFORMS
.into_iter()
.find_map(|(name, api_base)| {
- if name == self.model.client_name {
+ if name == self.model.client_name() {
Some(api_base.to_string())
} else {
None
diff --git a/src/client/prompt_format.rs b/src/client/prompt_format.rs
index 8a90924..320f105 100644
--- a/src/client/prompt_format.rs
+++ b/src/client/prompt_format.rs
@@ -108,6 +108,7 @@ pub fn generate_prompt(messages: &[Message], format: PromptFormat) -> anyhow::Re
}
parts.join("\n\n")
}
+ MessageContent::ToolResults(_) => String::new(),
};
match role {
MessageRole::System => prompt.push_str(&format!(
diff --git a/src/client/qianwen.rs b/src/client/qianwen.rs
index 7391a38..b33d13e 100644
--- a/src/client/qianwen.rs
+++ b/src/client/qianwen.rs
@@ -1,6 +1,6 @@
use super::{
- maybe_catch_error, message::*, sse_stream, Client, CompletionDetails, ExtraConfig, Model,
- ModelConfig, PromptAction, PromptKind, QianwenClient, SendData, SsMmessage, SseHandler,
+ maybe_catch_error, message::*, sse_stream, Client, CompletionOutput, ExtraConfig, Model,
+ ModelData, PromptAction, PromptKind, QianwenClient, SendData, SsMmessage, SseHandler,
};
use crate::utils::{base64_decode, sha256};
@@ -26,7 +26,7 @@ pub struct QianwenConfig {
pub name: Option<String>,
pub api_key: Option<String>,
#[serde(default)]
- pub models: Vec<ModelConfig>,
+ pub models: Vec<ModelData>,
pub extra: Option<ExtraConfig>,
}
@@ -62,7 +62,7 @@ impl QianwenClient {
}
fn is_vl(&self) -> bool {
- self.model.name.starts_with("qwen-vl")
+ self.model.name().starts_with("qwen-vl")
}
}
@@ -74,9 +74,9 @@ impl Client for QianwenClient {
&self,
client: &ReqwestClient,
mut data: SendData,
- ) -> Result<(String, CompletionDetails)> {
+ ) -> Result<CompletionOutput> {
let api_key = self.get_api_key()?;
- patch_messages(&self.model.name, &api_key, &mut data.messages).await?;
+ patch_messages(self.model.name(), &api_key, &mut data.messages).await?;
let builder = self.request_builder(client, data)?;
send_message(builder, self.is_vl()).await
}
@@ -88,16 +88,17 @@ impl Client for QianwenClient {
mut data: SendData,
) -> Result<()> {
let api_key = self.get_api_key()?;
- patch_messages(&self.model.name, &api_key, &mut data.messages).await?;
+ patch_messages(self.model.name(), &api_key, &mut data.messages).await?;
let builder = self.request_builder(client, data)?;
send_message_streaming(builder, handler, self.is_vl()).await
}
}
-async fn send_message(builder: RequestBuilder, is_vl: bool) -> Result<(String, CompletionDetails)> {
+async fn send_message(builder: RequestBuilder, is_vl: bool) -> Result<CompletionOutput> {
let data: Value = builder.send().await?.json().await?;
maybe_catch_error(&data)?;
+ debug!("non-stream-data: {data}");
extract_completion_text(&data, is_vl)
}
@@ -109,6 +110,7 @@ async fn send_message_streaming(
let handle = |message: SsMmessage| -> Result<bool> {
let data: Value = serde_json::from_str(&message.data)?;
maybe_catch_error(&data)?;
+ debug!("stream-data: {data}");
if is_vl {
if let Some(text) =
data["output"]["choices"][0]["message"]["content"][0]["text"].as_str()
@@ -129,10 +131,12 @@ fn build_body(data: SendData, model: &Model, is_vl: bool) -> Result<(Value, bool
messages,
temperature,
top_p,
+ functions: _,
stream,
} = data;
let mut has_upload = false;
+ let mut is_tool_call = false;
let input = if is_vl {
let messages: Vec<Value> = messages
.into_iter()
@@ -154,6 +158,10 @@ fn build_body(data: SendData, model: &Model, is_vl: bool) -> Result<(Value, bool
}
})
.collect(),
+ MessageContent::ToolResults(_) => {
+ is_tool_call = true;
+ vec![]
+ }
};
json!({ "role": role, "content": content })
})
@@ -167,6 +175,9 @@ fn build_body(data: SendData, model: &Model, is_vl: bool) -> Result<(Value, bool
"messages": messages,
})
};
+ if is_tool_call {
+ bail!("The client does not support function calling",);
+ }
let mut parameters = json!({});
if stream {
@@ -184,7 +195,7 @@ fn build_body(data: SendData, model: &Model, is_vl: bool) -> Result<(Value, bool
}
let body = json!({
- "model": &model.name,
+ "model": &model.name(),
"input": input,
"parameters": parameters
});
@@ -192,7 +203,7 @@ fn build_body(data: SendData, model: &Model, is_vl: bool) -> Result<(Value, bool
Ok((body, has_upload))
}
-fn extract_completion_text(data: &Value, is_vl: bool) -> Result<(String, CompletionDetails)> {
+fn extract_completion_text(data: &Value, is_vl: bool) -> Result<CompletionOutput> {
let err = || anyhow!("Invalid response data: {data}");
let text = if is_vl {
data["output"]["choices"][0]["message"]["content"][0]["text"]
@@ -201,13 +212,15 @@ fn extract_completion_text(data: &Value, is_vl: bool) -> Result<(String, Complet
} else {
data["output"]["text"].as_str().ok_or_else(err)?
};
- let details = CompletionDetails {
+ let output = CompletionOutput {
+ text: text.to_string(),
+ tool_calls: vec![],
id: data["request_id"].as_str().map(|v| v.to_string()),
input_tokens: data["usage"]["input_tokens"].as_u64(),
output_tokens: data["usage"]["output_tokens"].as_u64(),
};
- Ok((text.to_string(), details))
+ Ok(output)
}
/// Patch messages, upload embedded images to oss
diff --git a/src/client/replicate.rs b/src/client/replicate.rs
index 34cfd94..c0a77c2 100644
--- a/src/client/replicate.rs
+++ b/src/client/replicate.rs
@@ -1,9 +1,9 @@
use std::time::Duration;
use super::{
- catch_error, generate_prompt, smart_prompt_format, sse_stream, Client, CompletionDetails,
- ExtraConfig, Model, ModelConfig, PromptAction, PromptKind, ReplicateClient, SendData,
- SsMmessage, SseHandler,
+ catch_error, generate_prompt, smart_prompt_format, sse_stream, Client, CompletionOutput,
+ ExtraConfig, Model, ModelData, PromptAction, PromptKind, ReplicateClient, SendData, SsMmessage,
+ SseHandler,
};
use anyhow::{anyhow, Result};
@@ -19,7 +19,7 @@ pub struct ReplicateConfig {
pub name: Option<String>,
pub api_key: Option<String>,
#[serde(default)]
- pub models: Vec<ModelConfig>,
+ pub models: Vec<ModelData>,
pub extra: Option<ExtraConfig>,
}
@@ -37,7 +37,7 @@ impl ReplicateClient {
) -> Result<RequestBuilder> {
let body = build_body(data, &self.model)?;
- let url = format!("{API_BASE}/models/{}/predictions", self.model.name);
+ let url = format!("{API_BASE}/models/{}/predictions", self.model.name());
debug!("Replicate Request: {url} {body}");
@@ -55,7 +55,7 @@ impl Client for ReplicateClient {
&self,
client: &ReqwestClient,
data: SendData,
- ) -> Result<(String, CompletionDetails)> {
+ ) -> Result<CompletionOutput> {
let api_key = self.get_api_key()?;
let builder = self.request_builder(client, data, &api_key)?;
send_message(client, builder, &api_key).await
@@ -77,7 +77,7 @@ async fn send_message(
client: &ReqwestClient,
builder: RequestBuilder,
api_key: &str,
-) -> Result<(String, CompletionDetails)> {
+) -> Result<CompletionOutput> {
let res = builder.send().await?;
let status = res.status();
let data: Value = res.json().await?;
@@ -96,6 +96,7 @@ async fn send_message(
.await?
.json()
.await?;
+ debug!("non-stream-data: {prediction_data}");
let err = || anyhow!("Invalid response data: {prediction_data}");
let status = prediction_data["status"].as_str().ok_or_else(err)?;
if status == "succeeded" {
@@ -138,10 +139,11 @@ fn build_body(data: SendData, model: &Model) -> Result<Value> {
messages,
temperature,
top_p,
+ functions: _,
stream,
} = data;
- let prompt = generate_prompt(&messages, smart_prompt_format(&model.name))?;
+ let prompt = generate_prompt(&messages, smart_prompt_format(model.name()))?;
let mut input = json!({
"prompt": prompt,
@@ -170,7 +172,7 @@ fn build_body(data: SendData, model: &Model) -> Result<Value> {
Ok(body)
}
-fn extract_completion(data: &Value) -> Result<(String, CompletionDetails)> {
+fn extract_completion(data: &Value) -> Result<CompletionOutput> {
let text = data["output"]
.as_array()
.map(|parts| {
@@ -182,11 +184,13 @@ fn extract_completion(data: &Value) -> Result<(String, CompletionDetails)> {
})
.ok_or_else(|| anyhow!("Invalid response data: {data}"))?;
- let details = CompletionDetails {
+ let output = CompletionOutput {
+ text: text.to_string(),
+ tool_calls: vec![],
id: data["id"].as_str().map(|v| v.to_string()),
input_tokens: data["metrics"]["input_token_count"].as_u64(),
output_tokens: data["metrics"]["output_token_count"].as_u64(),
};
- Ok((text.to_string(), details))
+ Ok(output)
}
diff --git a/src/client/sse_handler.rs b/src/client/sse_handler.rs
index dbf78c2..ddbdbcd 100644
--- a/src/client/sse_handler.rs
+++ b/src/client/sse_handler.rs
@@ -3,10 +3,13 @@ use crate::utils::AbortSignal;
use anyhow::{Context, Result};
use tokio::sync::mpsc::UnboundedSender;
+use super::ToolCall;
+
pub struct SseHandler {
sender: UnboundedSender<SseEvent>,
- buffer: String,
abort: AbortSignal,
+ buffer: String,
+ tool_calls: Vec<ToolCall>,
}
impl SseHandler {
@@ -15,11 +18,12 @@ impl SseHandler {
sender,
abort,
buffer: String::new(),
+ tool_calls: Vec::new(),
}
}
pub fn text(&mut self, text: &str) -> Result<()> {
- // debug!("ReplyText: {}", text);
+ // debug!("HandleText: {}", text);
if text.is_empty() {
return Ok(());
}
@@ -33,7 +37,7 @@ impl SseHandler {
}
pub fn done(&mut self) -> Result<()> {
- // debug!("ReplyDone");
+ // debug!("HandleDone");
let ret = self
.sender
.send(SseEvent::Done)
@@ -42,14 +46,23 @@ impl SseHandler {
Ok(())
}
- pub fn get_buffer(&self) -> &str {
- &self.buffer
+ pub fn tool_call(&mut self, call: ToolCall) -> Result<()> {
+ // debug!("HandleCall: {:?}", call);
+ self.tool_calls.push(call);
+ Ok(())
}
pub fn get_abort(&self) -> AbortSignal {
self.abort.clone()
}
+ pub fn take(self) -> (String, Vec<ToolCall>) {
+ let Self {
+ buffer, tool_calls, ..
+ } = self;
+ (buffer, tool_calls)
+ }
+
fn safe_ret(&self, ret: Result<()>) -> Result<()> {
if ret.is_err() && self.abort.aborted() {
return Ok(());
diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs
index 4d06934..21fc154 100644
--- a/src/client/vertexai.rs
+++ b/src/client/vertexai.rs
@@ -1,8 +1,7 @@
-use super::access_token::*;
use super::{
- catch_error, json_stream, message::*, patch_system_message, Client, CompletionDetails,
- ExtraConfig, Model, ModelConfig, PromptAction, PromptKind, SendData, SseHandler,
- VertexAIClient,
+ access_token::*, catch_error, json_stream, message::*, patch_system_message, Client,
+ CompletionOutput, ExtraConfig, Model, ModelData, PromptAction, PromptKind, SendData,
+ SseHandler, ToolCall, VertexAIClient,
};
use anyhow::{anyhow, bail, Context, Result};
@@ -22,7 +21,7 @@ pub struct VertexAIConfig {
#[serde(rename = "safetySettings")]
pub safety_settings: Option<Value>,
#[serde(default)]
- pub models: Vec<ModelConfig>,
+ pub models: Vec<ModelData>,
pub extra: Option<ExtraConfig>,
}
@@ -46,7 +45,7 @@ impl VertexAIClient {
true => "streamGenerateContent",
false => "generateContent",
};
- let url = format!("{base_url}/google/models/{}:{func}", self.model.name);
+ let url = format!("{base_url}/google/models/{}:{func}", self.model.name());
let body = gemini_build_body(data, &self.model, self.config.safety_settings.clone())?;
@@ -66,7 +65,7 @@ impl Client for VertexAIClient {
&self,
client: &ReqwestClient,
data: SendData,
- ) -> Result<(String, CompletionDetails)> {
+ ) -> Result<CompletionOutput> {
prepare_gcloud_access_token(client, self.name(), &self.config.adc_file).await?;
let builder = self.request_builder(client, data)?;
gemini_send_message(builder).await
@@ -84,13 +83,14 @@ impl Client for VertexAIClient {
}
}
-pub async fn gemini_send_message(builder: RequestBuilder) -> Result<(String, CompletionDetails)> {
+pub async fn gemini_send_message(builder: RequestBuilder) -> Result<CompletionOutput> {
let res = builder.send().await?;
let status = res.status();
let data: Value = res.json().await?;
if !status.is_success() {
catch_error(&data, status.as_u16())?;
}
+ debug!("non-stream-data: {data}");
gemini_extract_completion_text(&data)
}
@@ -105,8 +105,28 @@ pub async fn gemini_send_message_streaming(
catch_error(&data, status.as_u16())?;
} else {
let handle = |value: &str| -> Result<()> {
- let value: Value = serde_json::from_str(value)?;
- handler.text(gemini_extract_text(&value)?)?;
+ let data: Value = serde_json::from_str(value)?;
+ debug!("stream-data: {data}");
+ if let Some(text) = data["candidates"][0]["content"]["parts"][0]["text"].as_str() {
+ if !text.is_empty() {
+ handler.text(text)?;
+ }
+ } else if let Some("SAFETY") = data["promptFeedback"]["blockReason"]
+ .as_str()
+ .or_else(|| data["candidates"][0]["finishReason"].as_str())
+ {
+ bail!("Content Blocked")
+ } else if let Some(parts) = data["candidates"][0]["content"]["parts"].as_array() {
+ for part in parts {
+ if let (Some(name), Some(args)) = (
+ part["functionCall"]["name"].as_str(),
+ part["functionCall"]["args"].as_object(),
+ ) {
+ handler.tool_call(ToolCall::new(name.to_string(), json!(args), None))?;
+ }
+ }
+ }
+
Ok(())
};
json_stream(res.bytes_stream(), handle).await?;
@@ -114,30 +134,45 @@ pub async fn gemini_send_message_streaming(
Ok(())
}
-fn gemini_extract_completion_text(data: &Value) -> Result<(String, CompletionDetails)> {
- let text = gemini_extract_text(data)?;
- let details = CompletionDetails {
+fn gemini_extract_completion_text(data: &Value) -> Result<CompletionOutput> {
+ let text = data["candidates"][0]["content"]["parts"][0]["text"]
+ .as_str()
+ .unwrap_or_default();
+
+ let mut tool_calls = vec![];
+ if let Some(parts) = data["candidates"][0]["content"]["parts"].as_array() {
+ tool_calls = parts
+ .iter()
+ .filter_map(|part| {
+ if let (Some(name), Some(args)) = (
+ part["functionCall"]["name"].as_str(),
+ part["functionCall"]["args"].as_object(),
+ ) {
+ Some(ToolCall::new(name.to_string(), json!(args), None))
+ } else {
+ None
+ }
+ })
+ .collect()
+ }
+ if text.is_empty() && tool_calls.is_empty() {
+ if let Some("SAFETY") = data["promptFeedback"]["blockReason"]
+ .as_str()
+ .or_else(|| data["candidates"][0]["finishReason"].as_str())
+ {
+ bail!("Content Blocked")
+ } else {
+ bail!("Invalid response data: {data}");
+ }
+ }
+ let output = CompletionOutput {
+ text: text.to_string(),
+ tool_calls,
id: None,
input_tokens: data["usageMetadata"]["promptTokenCount"].as_u64(),
output_tokens: data["usageMetadata"]["candidatesTokenCount"].as_u64(),
};
- Ok((text.to_string(), details))
-}
-
-fn gemini_extract_text(data: &Value) -> Result<&str> {
- match data["candidates"][0]["content"]["parts"][0]["text"].as_str() {
- Some(text) => Ok(text),
- None => {
- if let Some("SAFETY") = data["promptFeedback"]["blockReason"]
- .as_str()
- .or_else(|| data["candidates"][0]["finishReason"].as_str())
- {
- bail!("Blocked by safety settings,consider adjusting `safetySettings` in the client configuration")
- } else {
- bail!("Invalid response data: {data}")
- }
- }
- }
+ Ok(output)
}
pub(crate) fn gemini_build_body(
@@ -149,6 +184,7 @@ pub(crate) fn gemini_build_body(
mut messages,
temperature,
top_p,
+ functions,
stream: _,
} = data;
@@ -157,34 +193,60 @@ pub(crate) fn gemini_build_body(
let mut network_image_urls = vec![];
let contents: Vec<Value> = messages
.into_iter()
- .map(|message| {
- let role = match message.role {
+ .flat_map(|message| {
+ let Message { role, content } = message;
+ let role = match role {
MessageRole::User => "user",
_ => "model",
};
- match message.content {
- MessageContent::Text(text) => json!({
- "role": role,
- "parts": [{ "text": text }]
- }),
- MessageContent::Array(list) => {
- let list: Vec<Value> = list
- .into_iter()
- .map(|item| match item {
- MessageContentPart::Text { text } => json!({"text": text}),
- MessageContentPart::ImageUrl { image_url: ImageUrl { url } } => {
- if let Some((mime_type, data)) = url.strip_prefix("data:").and_then(|v| v.split_once(";base64,")) {
- json!({ "inline_data": { "mime_type": mime_type, "data": data } })
- } else {
- network_image_urls.push(url.clone());
- json!({ "url": url })
+ match content {
+ MessageContent::Text(text) => vec![json!({
+ "role": role,
+ "parts": [{ "text": text }]
+ })],
+ MessageContent::Array(list) => {
+ let parts: Vec<Value> = list
+ .into_iter()
+ .map(|item| match item {
+ MessageContentPart::Text { text } => json!({"text": text}),
+ MessageContentPart::ImageUrl { image_url: ImageUrl { url } } => {
+ if let Some((mime_type, data)) = url.strip_prefix("data:").and_then(|v| v.split_once(";base64,")) {
+ json!({ "inline_data": { "mime_type": mime_type, "data": data } })
+ } else {
+ network_image_urls.push(url.clone());
+ json!({ "url": url })
+ }
+ },
+ })
+ .collect();
+ vec![json!({ "role": role, "parts": parts })]
+ },
+ MessageContent::ToolResults((tool_call_results, _)) => {
+ let function_call_parts: Vec<Value> = tool_call_results.iter().map(|tool_call_result| {
+ json!({
+ "functionCall": {
+ "name": tool_call_result.call.name,
+ "args": tool_call_result.call.arguments,
+ }
+ })
+ }).collect();
+ let function_response_parts: Vec<Value> = tool_call_results.into_iter().map(|tool_call_result| {
+ json!({
+ "functionResponse": {
+ "name": tool_call_result.call.name,
+ "response": {
+ "name": tool_call_result.call.name,
+ "content": tool_call_result.output,
+ }
}
- },
- })
- .collect();
- json!({ "role": role, "parts": list })
+ })
+ }).collect();
+ vec![
+ json!({ "role": "model", "parts": function_call_parts }),
+ json!({ "role": "function", "parts": function_response_parts }),
+ ]
+ }
}
- }
})
.collect();
@@ -211,6 +273,10 @@ pub(crate) fn gemini_build_body(
body["generationConfig"]["topP"] = v.into();
}
+ if let Some(functions) = functions {
+ body["tools"] = json!([{ "functionDeclarations": *functions }]);
+ }
+
Ok(body)
}
diff --git a/src/client/vertexai_claude.rs b/src/client/vertexai_claude.rs
index 2b5f763..2c51d59 100644
--- a/src/client/vertexai_claude.rs
+++ b/src/client/vertexai_claude.rs
@@ -2,7 +2,7 @@ use super::access_token::*;
use super::claude::{claude_build_body, claude_send_message, claude_send_message_streaming};
use super::vertexai::prepare_gcloud_access_token;
use super::{
- Client, CompletionDetails, ExtraConfig, Model, ModelConfig, PromptAction, PromptKind, SendData,
+ Client, CompletionOutput, ExtraConfig, Model, ModelData, PromptAction, PromptKind, SendData,
SseHandler, VertexAIClaudeClient,
};
@@ -18,7 +18,7 @@ pub struct VertexAIClaudeConfig {
pub location: Option<String>,
pub adc_file: Option<String>,
#[serde(default)]
- pub models: Vec<ModelConfig>,
+ pub models: Vec<ModelData>,
pub extra: Option<ExtraConfig>,
}
@@ -39,7 +39,7 @@ impl VertexAIClaudeClient {
let base_url = format!("https://{location}-aiplatform.googleapis.com/v1/projects/{project_id}/locations/{location}/publishers");
let url = format!(
"{base_url}/anthropic/models/{}:streamRawPredict",
- self.model.name
+ self.model.name()
);
let mut body = claude_build_body(data, &self.model)?;
@@ -64,7 +64,7 @@ impl Client for VertexAIClaudeClient {
&self,
client: &ReqwestClient,
data: SendData,
- ) -> Result<(String, CompletionDetails)> {
+ ) -> Result<CompletionOutput> {
prepare_gcloud_access_token(client, self.name(), &self.config.adc_file).await?;
let builder = self.request_builder(client, data)?;
claude_send_message(builder).await