summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-06-22 14:07:58 +0800
committerGitHub <noreply@github.com>2024-06-22 14:07:58 +0800
commit0ec370b48c951a69459f9232b56df882765ccb52 (patch)
treec12bab3bd72acc97adcd835b8d5aea82d1b96cd3 /src
parent250e0eb7fee0d180b08b7d38d10c530e9a6a120d (diff)
downloadaichat-0ec370b48c951a69459f9232b56df882765ccb52.tar.gz
refactor: improve clients (#632)
Diffstat (limited to 'src')
-rw-r--r--src/client/claude.rs28
-rw-r--r--src/client/cohere.rs12
-rw-r--r--src/client/common.rs4
-rw-r--r--src/client/openai.rs20
-rw-r--r--src/client/vertexai.rs20
-rw-r--r--src/config/input.rs10
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<ChatCompletionsOu
.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
+ if let Some(calls) = data["choices"][0]["message"]["tool_calls"].as_array() {
+ tool_calls = calls
.iter()
.filter_map(|call| {
if let (Some(name), Some(arguments), Some(id)) = (
diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs
index 69d1bc4..16910b6 100644
--- a/src/client/vertexai.rs
+++ b/src/client/vertexai.rs
@@ -295,29 +295,29 @@ pub fn gemini_build_chat_completions_body(
.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| {
+ MessageContent::ToolResults((tool_results, _)) => {
+ let model_parts: Vec<Value> = 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<Value> = tool_call_results.into_iter().map(|tool_call_result| {
+ let function_parts: Vec<Value> = 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<ToolResult>) -> Self {
+ pub fn merge_tool_call(mut self, output: String, tool_results: Vec<ToolResult>) -> 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
}