summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-10-19 20:10:22 +0800
committerGitHub <noreply@github.com>2024-10-19 20:10:22 +0800
commit7b20ab064f30b7eddd18cecb9afb93e2b2c2965a (patch)
tree45030db5b840684ad9ee404564870807e78ed06c /src
parent7e29f642a99682ea7bb074eb4e7aa7d34d523515 (diff)
downloadaichat-7b20ab064f30b7eddd18cecb9afb93e2b2c2965a.tar.gz
refactor: improve handling of no_stream/no_system_message (#936)
Diffstat (limited to 'src')
-rw-r--r--src/client/common.rs2
-rw-r--r--src/client/openai.rs6
-rw-r--r--src/client/stream.rs4
-rw-r--r--src/config/input.rs9
-rw-r--r--src/serve.rs55
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();
}