diff options
| author | sigoden <sigoden@gmail.com> | 2024-04-29 06:51:03 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-04-29 06:51:03 +0800 |
| commit | 865be2bf75bb62b6aeee059f684400b4b9938a15 (patch) | |
| tree | b6296ccdde2af49251b02087b9f51c5239cc4d2b /src/serve.rs | |
| parent | b33e2da75efab4246f9f1d74a0be56aba1d8bd07 (diff) | |
| download | aichat-865be2bf75bb62b6aeee059f684400b4b9938a15.tar.gz | |
feat: non-streaming returns completion stats (#456)
Diffstat (limited to 'src/serve.rs')
| -rw-r--r-- | src/serve.rs | 39 |
1 files changed, 31 insertions, 8 deletions
diff --git a/src/serve.rs b/src/serve.rs index 394a0e4..3d8b1a3 100644 --- a/src/serve.rs +++ b/src/serve.rs @@ -1,5 +1,8 @@ use crate::{ - client::{init_client, ClientConfig, Message, Model, ReplyEvent, ReplyHandler, SendData}, + client::{ + init_client, ClientConfig, CompletionStats, Message, Model, ReplyEvent, ReplyHandler, + SendData, + }, config::{Config, GlobalConfig}, utils::create_abort_signal, }; @@ -248,10 +251,19 @@ impl Server { .body(BodyExt::boxed(StreamBody::new(stream)))?; Ok(res) } else { - let content = client.send_message_inner(&http_client, send_data).await?; + let (content, stats) = client.send_message_inner(&http_client, send_data).await?; let res = Response::builder() .header("Content-Type", "application/json") - .body(Full::new(ret_non_stream(&completion_id, created, &content)).boxed())?; + .body( + Full::new(ret_non_stream( + &completion_id, + &model_name, + created, + &content, + &stats, + )) + .boxed(), + )?; Ok(res) } } @@ -340,12 +352,22 @@ fn create_frame(id: &str, model: &str, created: i64, content: &str, done: bool) Frame::data(Bytes::from(output)) } -fn ret_non_stream(id: &str, created: i64, content: &str) -> Bytes { +fn ret_non_stream( + id: &str, + model: &str, + created: i64, + content: &str, + stats: &CompletionStats, +) -> Bytes { + let id = stats.id.as_deref().unwrap_or(id); + let input_tokens = stats.input_tokens.unwrap_or_default(); + let output_tokens = stats.output_tokens.unwrap_or_default(); + let total_tokens = input_tokens + output_tokens; let res_body = json!({ "id": id, "object": "chat.completion", "created": created, - "model": "gpt-3.5-turbo", + "model": model, "choices": [ { "index": 0, @@ -353,13 +375,14 @@ fn ret_non_stream(id: &str, created: i64, content: &str) -> Bytes { "role": "assistant", "content": content, }, + "logprobs": null, "finish_reason": "stop", }, ], "usage": { - "prompt_tokens": 0, - "completion_tokens": 0, - "total_tokens": 0, + "prompt_tokens": input_tokens, + "completion_tokens": output_tokens, + "total_tokens": total_tokens, }, }); Bytes::from(res_body.to_string()) |
