diff options
| author | sigoden <sigoden@gmail.com> | 2024-09-09 18:39:14 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-09-09 18:39:14 +0800 |
| commit | 69965466e6510db730f0eaf2384addecdcc66a84 (patch) | |
| tree | 891b0169cc6d972da545873b4832d8faacf444f9 /src/client/openai.rs | |
| parent | a56d5f2ddfcfb290c308f2b1a3464c3bebc340de (diff) | |
| download | aichat-69965466e6510db730f0eaf2384addecdcc66a84.tar.gz | |
feat: proxy chat-completions api with tools support (#850)
Diffstat (limited to 'src/client/openai.rs')
| -rw-r--r-- | src/client/openai.rs | 50 |
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() { |
