summaryrefslogtreecommitdiffstats
path: root/src/client/openai.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/openai.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/openai.rs')
-rw-r--r--src/client/openai.rs137
1 files changed, 127 insertions, 10 deletions
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!(