From 0ec370b48c951a69459f9232b56df882765ccb52 Mon Sep 17 00:00:00 2001 From: sigoden Date: Sat, 22 Jun 2024 14:07:58 +0800 Subject: refactor: improve clients (#632) --- src/client/claude.rs | 28 ++++++++++++++-------------- src/client/cohere.rs | 12 ++++++------ src/client/common.rs | 4 ++-- src/client/openai.rs | 20 ++++++++++---------- src/client/vertexai.rs | 20 ++++++++++---------- src/config/input.rs | 10 +++++----- 6 files changed, 47 insertions(+), 47 deletions(-) diff --git a/src/client/claude.rs b/src/client/claude.rs index 2b836ee..e239fc4 100644 --- a/src/client/claude.rs +++ b/src/client/claude.rs @@ -188,36 +188,36 @@ pub fn claude_build_chat_completions_body( "content": content, })] } - MessageContent::ToolResults((tool_call_results, text)) => { - let mut tool_call = vec![]; - let mut tool_result = vec![]; + MessageContent::ToolResults((tool_results, text)) => { + let mut assistant_parts = vec![]; + let mut user_parts = vec![]; if !text.is_empty() { - tool_call.push(json!({ + assistant_parts.push(json!({ "type": "text", "text": text, })) } - for tool_call_result in tool_call_results { - tool_call.push(json!({ + for tool_result in tool_results { + assistant_parts.push(json!({ "type": "tool_use", - "id": tool_call_result.call.id, - "name": tool_call_result.call.name, - "input": tool_call_result.call.arguments, + "id": tool_result.call.id, + "name": tool_result.call.name, + "input": tool_result.call.arguments, })); - tool_result.push(json!({ + user_parts.push(json!({ "type": "tool_result", - "tool_use_id": tool_call_result.call.id, - "content": tool_call_result.output.to_string(), + "tool_use_id": tool_result.call.id, + "content": tool_result.output.to_string(), })); } vec![ json!({ "role": "assistant", - "content": tool_call, + "content": assistant_parts, }), json!({ "role": "user", - "content": tool_result, + "content": user_parts, }), ] } diff --git a/src/client/cohere.rs b/src/client/cohere.rs index 698c2a6..3e1cb2b 100644 --- a/src/client/cohere.rs +++ b/src/client/cohere.rs @@ -220,8 +220,8 @@ fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Resu .collect(); Some(json!({ "role": role, "message": list.join("\n\n") })) } - MessageContent::ToolResults((tool_call_results, _)) => { - tool_results = Some(tool_call_results); + MessageContent::ToolResults((results, _)) => { + tool_results = Some(results); None } } @@ -263,14 +263,14 @@ fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Resu if let Some(tool_results) = tool_results { let tool_results: Vec<_> = tool_results .into_iter() - .map(|tool_call_result| { + .map(|tool_result| { json!({ "call": { - "name": tool_call_result.call.name, - "parameters": tool_call_result.call.arguments, + "name": tool_result.call.name, + "parameters": tool_result.call.arguments, }, "outputs": [ - tool_call_result.output, + tool_result.output, ] }) diff --git a/src/client/common.rs b/src/client/common.rs index 13fd220..1b1e723 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -355,7 +355,7 @@ pub trait Client: Sync + Send { let data = input.prepare_completion_data(self.model(), false)?; self.chat_completions_inner(&client, data) .await - .with_context(|| "Failed to call chat completions api") + .with_context(|| "Failed to call chat-completions api") } async fn chat_completions_streaming( @@ -381,7 +381,7 @@ pub trait Client: Sync + Send { self.chat_completions_streaming_inner(&client, handler, data).await } => { handler.done()?; - ret.with_context(|| "Failed to call chat completions api") + ret.with_context(|| "Failed to call chat-completions api") } _ = watch_abort_signal(abort_signal) => { handler.done()?; diff --git a/src/client/openai.rs b/src/client/openai.rs index 40058e6..bfe896f 100644 --- a/src/client/openai.rs +++ b/src/client/openai.rs @@ -177,26 +177,26 @@ pub fn openai_build_chat_completions_body(data: ChatCompletionsData, model: &Mod .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| { + MessageContent::ToolResults((tool_results, text)) => { + let tool_calls: Vec<_> = tool_results.iter().map(|tool_result| { json!({ - "id": tool_call_result.call.id, + "id": tool_result.call.id, "type": "function", "function": { - "name": tool_call_result.call.name, - "arguments": tool_call_result.call.arguments, + "name": tool_result.call.name, + "arguments": tool_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 { + for tool_result in tool_results { messages.push( json!({ "role": "tool", - "content": tool_call_result.output.to_string(), - "tool_call_id": tool_call_result.call.id, + "content": tool_result.output.to_string(), + "tool_call_id": tool_result.call.id, }) ); } @@ -251,8 +251,8 @@ pub fn openai_extract_chat_completions(data: &Value) -> Result { - let function_call_parts: Vec = tool_call_results.iter().map(|tool_call_result| { + MessageContent::ToolResults((tool_results, _)) => { + let model_parts: Vec = tool_results.iter().map(|tool_result| { json!({ "functionCall": { - "name": tool_call_result.call.name, - "args": tool_call_result.call.arguments, + "name": tool_result.call.name, + "args": tool_result.call.arguments, } }) }).collect(); - let function_response_parts: Vec = tool_call_results.into_iter().map(|tool_call_result| { + let function_parts: Vec = tool_results.into_iter().map(|tool_result| { json!({ "functionResponse": { - "name": tool_call_result.call.name, + "name": tool_result.call.name, "response": { - "name": tool_call_result.call.name, - "content": tool_call_result.output, + "name": tool_result.call.name, + "content": tool_result.output, } } }) }).collect(); vec![ - json!({ "role": "model", "parts": function_call_parts }), - json!({ "role": "function", "parts": function_response_parts }), + json!({ "role": "model", "parts": model_parts }), + json!({ "role": "function", "parts": function_parts }), ] } } diff --git a/src/config/input.rs b/src/config/input.rs index b0b0bd9..3a258b9 100644 --- a/src/config/input.rs +++ b/src/config/input.rs @@ -209,13 +209,13 @@ impl Input { self.rag_name.as_deref() } - pub fn merge_tool_call(mut self, output: String, tool_call_results: Vec) -> Self { + pub fn merge_tool_call(mut self, output: String, tool_results: Vec) -> Self { match self.tool_call.as_mut() { - Some(exist_tool_call_results) => { - exist_tool_call_results.0.extend(tool_call_results); - exist_tool_call_results.1 = output; + Some(exist_tool_results) => { + exist_tool_results.0.extend(tool_results); + exist_tool_results.1 = output; } - None => self.tool_call = Some((tool_call_results, output)), + None => self.tool_call = Some((tool_results, output)), } self } -- cgit v1.2.3