summaryrefslogtreecommitdiffstats
path: root/src/serve.rs
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/serve.rs
parent7e29f642a99682ea7bb074eb4e7aa7d34d523515 (diff)
downloadaichat-7b20ab064f30b7eddd18cecb9afb93e2b2c2965a.tar.gz
refactor: improve handling of no_stream/no_system_message (#936)
Diffstat (limited to 'src/serve.rs')
-rw-r--r--src/serve.rs55
1 files changed, 39 insertions, 16 deletions
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();
}