diff options
| -rw-r--r-- | src/client/common.rs | 2 | ||||
| -rw-r--r-- | src/client/openai.rs | 6 | ||||
| -rw-r--r-- | src/client/stream.rs | 4 | ||||
| -rw-r--r-- | src/config/input.rs | 9 | ||||
| -rw-r--r-- | src/serve.rs | 55 |
5 files changed, 49 insertions, 27 deletions
diff --git a/src/client/common.rs b/src/client/common.rs index 0cf900f..5fd4f24 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -67,7 +67,7 @@ pub trait Client: Sync + Send { input: &Input, handler: &mut SseHandler, ) -> Result<()> { - let abort_signal = handler.get_abort(); + let abort_signal = handler.abort(); let input = input.clone(); tokio::select! { ret = async { diff --git a/src/client/openai.rs b/src/client/openai.rs index 831b64f..c4c2b0c 100644 --- a/src/client/openai.rs +++ b/src/client/openai.rs @@ -193,17 +193,13 @@ struct EmbeddingsResBodyEmbedding { pub fn openai_build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Value { let ChatCompletionsData { - mut messages, + messages, temperature, top_p, functions, stream, } = data; - if model.no_system_message() { - patch_system_message(&mut messages); - } - let messages: Vec<Value> = messages .into_iter() .flat_map(|message| { diff --git a/src/client/stream.rs b/src/client/stream.rs index 7fa1abd..735ec93 100644 --- a/src/client/stream.rs +++ b/src/client/stream.rs @@ -61,11 +61,11 @@ impl SseHandler { Ok(()) } - pub fn get_abort(&self) -> AbortSignal { + pub fn abort(&self) -> AbortSignal { self.abort.clone() } - pub fn get_tool_calls(&self) -> &[ToolCall] { + pub fn tool_calls(&self) -> &[ToolCall] { &self.tool_calls } diff --git a/src/config/input.rs b/src/config/input.rs index 0325e1c..bf092de 100644 --- a/src/config/input.rs +++ b/src/config/input.rs @@ -1,8 +1,8 @@ use super::*; use crate::client::{ - init_client, ChatCompletionsData, Client, ImageUrl, Message, MessageContent, - MessageContentPart, MessageRole, Model, + init_client, patch_system_message, ChatCompletionsData, Client, ImageUrl, Message, + MessageContent, MessageContentPart, MessageRole, Model, }; use crate::function::{ToolResult, ToolResults}; use crate::utils::{base64_encode, sha256, AbortSignal}; @@ -206,7 +206,10 @@ impl Input { if !self.medias.is_empty() && !model.supports_vision() { bail!("The current model does not support vision. Is the model configured with `supports_vision: true`?"); } - let messages = self.build_messages()?; + let mut messages = self.build_messages()?; + if model.no_system_message() { + patch_system_message(&mut messages); + } model.guard_max_input_tokens(&messages)?; let temperature = self.role().temperature(); let top_p = self.role().top_p(); 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<ResEvent>, is_first: Arc<AtomicBool>, ) { - 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(); } |
