From 7d33a2c660a4a675d1ec9d3b10a107ce953d76ae Mon Sep 17 00:00:00 2001 From: sigoden Date: Fri, 29 Nov 2024 09:33:21 +0800 Subject: fix: stream function calling for some openai-compatible clients (#1017) --- src/client/openai.rs | 25 +++++++++++++++---------- 1 file changed, 15 insertions(+), 10 deletions(-) (limited to 'src/client/openai.rs') diff --git a/src/client/openai.rs b/src/client/openai.rs index 7f449fb..253bf21 100644 --- a/src/client/openai.rs +++ b/src/client/openai.rs @@ -101,7 +101,7 @@ pub async fn openai_chat_completions_streaming( handler: &mut SseHandler, _model: &Model, ) -> Result<()> { - let mut function_index = 0; + let mut call_id = String::new(); let mut function_name = String::new(); let mut function_arguments = String::new(); let mut function_id = String::new(); @@ -109,7 +109,7 @@ pub async fn openai_chat_completions_streaming( 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") + format!("Tool call '{function_name}' have non-JSON arguments '{function_arguments}'") })?; handler.tool_call(ToolCall::new( function_name.clone(), @@ -121,18 +121,23 @@ pub async fn openai_chat_completions_streaming( } let data: Value = serde_json::from_str(&message.data)?; debug!("stream-data: {data}"); - if let Some(text) = data["choices"][0]["delta"]["content"].as_str() { + if let Some(text) = data["choices"][0]["delta"]["content"] + .as_str() + .filter(|v| !v.is_empty()) + { handler.text(text)?; } else if let (Some(function), index, id) = ( data["choices"][0]["delta"]["tool_calls"][0]["function"].as_object(), data["choices"][0]["delta"]["tool_calls"][0]["index"].as_u64(), - data["choices"][0]["delta"]["tool_calls"][0]["id"].as_str(), + data["choices"][0]["delta"]["tool_calls"][0]["id"] + .as_str() + .filter(|v| !v.is_empty()), ) { - let index = index.unwrap_or_default(); - if index != function_index { + let maybe_call_id = format!("{}/{}", id.unwrap_or_default(), index.unwrap_or_default()); + if maybe_call_id != call_id && maybe_call_id.len() >= call_id.len() { 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") + format!("Tool call '{function_name}' have non-JSON arguments '{function_arguments}'") })?; handler.tool_call(ToolCall::new( function_name.clone(), @@ -143,7 +148,7 @@ pub async fn openai_chat_completions_streaming( function_name.clear(); function_arguments.clear(); function_id.clear(); - function_index = index; + call_id = maybe_call_id; } if let Some(name) = function.get("name").and_then(|v| v.as_str()) { if name.starts_with(&function_name) { @@ -240,7 +245,7 @@ pub fn openai_build_chat_completions_body(data: ChatCompletionsData, model: &Mod vec![ json!({ "role": MessageRole::Assistant, - "content": null, + "content": "", "tool_calls": [ { "id": tool_result.call.id, @@ -319,7 +324,7 @@ pub fn openai_extract_chat_completions(data: &Value) -> Result