summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
-rw-r--r--src/client/cohere.rs42
1 files changed, 23 insertions, 19 deletions
diff --git a/src/client/cohere.rs b/src/client/cohere.rs
index 9954a61..8745347 100644
--- a/src/client/cohere.rs
+++ b/src/client/cohere.rs
@@ -207,25 +207,6 @@ fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Resu
"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();
}
@@ -247,6 +228,29 @@ fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Resu
body["stream"] = true.into();
}
+ 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(object) = body.as_object_mut() {
+ object.remove("chat_history");
+ object.remove("message");
+ }
+ }
+
if let Some(functions) = functions {
body["tools"] = functions
.iter()