summaryrefslogtreecommitdiffstats
path: root/src/client/cohere.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/cohere.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/cohere.rs')
-rw-r--r--src/client/cohere.rs121
1 files changed, 103 insertions, 18 deletions
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)
}