From 865be2bf75bb62b6aeee059f684400b4b9938a15 Mon Sep 17 00:00:00 2001 From: sigoden Date: Mon, 29 Apr 2024 06:51:03 +0800 Subject: feat: non-streaming returns completion stats (#456) --- src/client/openai.rs | 23 ++++++++++++++++------- 1 file changed, 16 insertions(+), 7 deletions(-) (limited to 'src/client/openai.rs') diff --git a/src/client/openai.rs b/src/client/openai.rs index 54197ac..eba0992 100644 --- a/src/client/openai.rs +++ b/src/client/openai.rs @@ -1,5 +1,6 @@ use super::{ - catch_error, ExtraConfig, Model, ModelConfig, OpenAIClient, PromptType, ReplyHandler, SendData, + catch_error, CompletionStats, ExtraConfig, Model, ModelConfig, OpenAIClient, PromptType, + ReplyHandler, SendData, }; use crate::utils::PromptKind; @@ -51,7 +52,7 @@ impl OpenAIClient { } } -pub async fn openai_send_message(builder: RequestBuilder) -> Result { +pub async fn openai_send_message(builder: RequestBuilder) -> Result<(String, CompletionStats)> { let res = builder.send().await?; let status = res.status(); let data: Value = res.json().await?; @@ -59,11 +60,7 @@ pub async fn openai_send_message(builder: RequestBuilder) -> Result { catch_error(&data, status.as_u16())?; } - let output = data["choices"][0]["message"]["content"] - .as_str() - .ok_or_else(|| anyhow!("Invalid response data: {data}"))?; - - Ok(output.to_string()) + openai_extract_completion(&data) } pub async fn openai_send_message_streaming( @@ -143,6 +140,18 @@ pub fn openai_build_body(data: SendData, model: &Model) -> Value { body } +pub fn openai_extract_completion(data: &Value) -> Result<(String, CompletionStats)> { + let text = data["choices"][0]["message"]["content"] + .as_str() + .ok_or_else(|| anyhow!("Invalid response data: {data}"))?; + let stats = CompletionStats { + id: data["id"].as_str().map(|v| v.to_string()), + input_tokens: data["usage"]["prompt_tokens"].as_u64(), + output_tokens: data["usage"]["completion_tokens"].as_u64(), + }; + Ok((text.to_string(), stats)) +} + impl_client_trait!( OpenAIClient, openai_send_message, -- cgit v1.2.3