From 7b20ab064f30b7eddd18cecb9afb93e2b2c2965a Mon Sep 17 00:00:00 2001 From: sigoden Date: Sat, 19 Oct 2024 20:10:22 +0800 Subject: refactor: improve handling of no_stream/no_system_message (#936) --- src/serve.rs | 55 +++++++++++++++++++++++++++++++++++++++---------------- 1 file changed, 39 insertions(+), 16 deletions(-) (limited to 'src/serve.rs') diff --git a/src/serve.rs b/src/serve.rs index 51cc010..ebb79af 100644 --- a/src/serve.rs +++ b/src/serve.rs @@ -281,7 +281,7 @@ impl Server { tools, } = req_body; - let messages = + let mut messages = parse_messages(messages).map_err(|err| anyhow!("Invalid request body, {err}"))?; let functions = parse_tools(tools).map_err(|err| anyhow!("Invalid request body, {err}"))?; @@ -313,7 +313,9 @@ impl Server { let completion_id = generate_completion_id(); let created = Utc::now().timestamp(); - + if client.model().no_system_message() { + patch_system_message(&mut messages); + } let data: ChatCompletionsData = ChatCompletionsData { messages, temperature, @@ -357,20 +359,41 @@ impl Server { tx: &UnboundedSender, is_first: Arc, ) { - let ret = client - .chat_completions_streaming_inner(http_client, handler, data) - .await; - let first = match ret { - Ok(()) => None, - Err(err) => Some(format!("{err:?}")), - }; - if is_first.load(Ordering::SeqCst) { - let _ = tx.send(ResEvent::First(first)); - is_first.store(false, Ordering::SeqCst) - } - let tool_calls = handler.get_tool_calls(); - if !tool_calls.is_empty() { - let _ = tx.send(ResEvent::ToolCalls(tool_calls.to_vec())); + if client.model().no_stream() { + let ret = client.chat_completions_inner(http_client, data).await; + match ret { + Ok(output) => { + let ChatCompletionsOutput { + text, tool_calls, .. + } = output; + let _ = tx.send(ResEvent::First(None)); + is_first.store(false, Ordering::SeqCst); + let _ = tx.send(ResEvent::Text(text)); + if !tool_calls.is_empty() { + let _ = tx.send(ResEvent::ToolCalls(tool_calls)); + } + } + Err(err) => { + let _ = tx.send(ResEvent::First(Some(format!("{err:?}")))); + is_first.store(false, Ordering::SeqCst) + } + }; + } else { + let ret = client + .chat_completions_streaming_inner(http_client, handler, data) + .await; + let first = match ret { + Ok(()) => None, + Err(err) => Some(format!("{err:?}")), + }; + if is_first.load(Ordering::SeqCst) { + let _ = tx.send(ResEvent::First(first)); + is_first.store(false, Ordering::SeqCst) + } + let tool_calls = handler.tool_calls().to_vec(); + if !tool_calls.is_empty() { + let _ = tx.send(ResEvent::ToolCalls(tool_calls)); + } } handler.done(); } -- cgit v1.2.3