summaryrefslogtreecommitdiffstats
path: root/src/client/openai.rs
diff options
context:
space:
mode:
Diffstat (limited to 'src/client/openai.rs')
-rw-r--r--src/client/openai.rs50
1 files changed, 28 insertions, 22 deletions
diff --git a/src/client/openai.rs b/src/client/openai.rs
index 902a215..4876ed3 100644
--- a/src/client/openai.rs
+++ b/src/client/openai.rs
@@ -108,9 +108,12 @@ pub async fn openai_chat_completions_streaming(
let handle = |message: SseMmessage| -> Result<bool> {
if message.data == "[DONE]" {
if !function_name.is_empty() {
+ let arguments: Value = function_arguments.parse().with_context(|| {
+ format!("Tool call '{function_name}' is invalid: arguments must be in valid JSON format")
+ })?;
handler.tool_call(ToolCall::new(
function_name.clone(),
- json!(function_arguments),
+ arguments,
normalize_function_id(&function_id),
))?;
}
@@ -128,9 +131,12 @@ pub async fn openai_chat_completions_streaming(
let index = index.unwrap_or_default();
if index != function_index {
if !function_name.is_empty() {
+ let arguments: Value = function_arguments.parse().with_context(|| {
+ format!("Tool call '{function_name}' is invalid: arguments must be in valid JSON format")
+ })?;
handler.tool_call(ToolCall::new(
function_name.clone(),
- json!(function_arguments),
+ arguments,
normalize_function_id(&function_id),
))?;
}
@@ -207,7 +213,7 @@ pub fn openai_build_chat_completions_body(data: ChatCompletionsData, model: &Mod
"type": "function",
"function": {
"name": tool_result.call.name,
- "arguments": tool_result.call.arguments,
+ "arguments": tool_result.call.arguments.to_string(),
},
})
}).collect();
@@ -237,7 +243,7 @@ pub fn openai_build_chat_completions_body(data: ChatCompletionsData, model: &Mod
"type": "function",
"function": {
"name": tool_result.call.name,
- "arguments": tool_result.call.arguments,
+ "arguments": tool_result.call.arguments.to_string(),
},
}
]
@@ -302,24 +308,24 @@ pub fn openai_extract_chat_completions(data: &Value) -> Result<ChatCompletionsOu
let mut tool_calls = vec![];
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)) = (
- 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()
+ for call in calls {
+ if let (Some(name), Some(arguments), Some(id)) = (
+ call["function"]["name"].as_str(),
+ call["function"]["arguments"].as_str(),
+ call["id"].as_str(),
+ ) {
+ let arguments: Value = arguments.parse().with_context(|| {
+ format!(
+ "Tool call '{name}' is invalid: arguments must be in valid JSON format"
+ )
+ })?;
+ tool_calls.push(ToolCall::new(
+ name.to_string(),
+ arguments,
+ Some(id.to_string()),
+ ));
+ }
+ }
};
if text.is_empty() && tool_calls.is_empty() {