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/qianwen.rs | 90 ++++++++++++++++++++++++++++----------------------- 1 file changed, 50 insertions(+), 40 deletions(-) (limited to 'src/client/qianwen.rs') diff --git a/src/client/qianwen.rs b/src/client/qianwen.rs index 64d27d4..8a5a6d5 100644 --- a/src/client/qianwen.rs +++ b/src/client/qianwen.rs @@ -1,6 +1,6 @@ use super::{ - maybe_catch_error, message::*, Client, ExtraConfig, Model, ModelConfig, PromptType, - QianwenClient, ReplyHandler, SendData, + maybe_catch_error, message::*, Client, CompletionStats, ExtraConfig, Model, ModelConfig, + PromptType, QianwenClient, ReplyHandler, SendData, }; use crate::utils::{sha256sum, PromptKind}; @@ -69,19 +69,39 @@ impl QianwenClient { } } -async fn send_message(builder: RequestBuilder, is_vl: bool) -> Result { - let data: Value = builder.send().await?.json().await?; - maybe_catch_error(&data)?; +#[async_trait] +impl Client for QianwenClient { + client_common_fns!(); - let output = if is_vl { - data["output"]["choices"][0]["message"]["content"][0]["text"].as_str() - } else { - data["output"]["text"].as_str() - }; + async fn send_message_inner( + &self, + client: &ReqwestClient, + mut data: SendData, + ) -> Result<(String, CompletionStats)> { + let api_key = self.get_api_key()?; + patch_messages(&self.model.name, &api_key, &mut data.messages).await?; + let builder = self.request_builder(client, data)?; + send_message(builder, self.is_vl()).await + } - let output = output.ok_or_else(|| anyhow!("Unexpected response {data}"))?; + async fn send_message_streaming_inner( + &self, + client: &ReqwestClient, + handler: &mut ReplyHandler, + mut data: SendData, + ) -> Result<()> { + let api_key = self.get_api_key()?; + patch_messages(&self.model.name, &api_key, &mut data.messages).await?; + let builder = self.request_builder(client, data)?; + send_message_streaming(builder, handler, self.is_vl()).await + } +} - Ok(output.to_string()) +async fn send_message(builder: RequestBuilder, is_vl: bool) -> Result<(String, CompletionStats)> { + let data: Value = builder.send().await?.json().await?; + maybe_catch_error(&data)?; + + extract_completion_text(&data, is_vl) } async fn send_message_streaming( @@ -190,6 +210,24 @@ fn build_body(data: SendData, model: &Model, is_vl: bool) -> Result<(Value, bool Ok((body, has_upload)) } +fn extract_completion_text(data: &Value, is_vl: bool) -> Result<(String, CompletionStats)> { + let err = || anyhow!("Invalid response data: {data}"); + let text = if is_vl { + data["output"]["choices"][0]["message"]["content"][0]["text"] + .as_str() + .ok_or_else(err)? + } else { + data["output"]["text"].as_str().ok_or_else(err)? + }; + let stats = CompletionStats { + id: data["request_id"].as_str().map(|v| v.to_string()), + input_tokens: data["usage"]["input_tokens"].as_u64(), + output_tokens: data["usage"]["output_tokens"].as_u64(), + }; + + Ok((text.to_string(), stats)) +} + /// Patch messsages, upload embedded images to oss async fn patch_messages(model: &str, api_key: &str, messages: &mut Vec) -> Result<()> { for message in messages { @@ -283,31 +321,3 @@ async fn upload(model: &str, api_key: &str, url: &str) -> Result { } Ok(format!("oss://{key}")) } - -#[async_trait] -impl Client for QianwenClient { - client_common_fns!(); - - async fn send_message_inner( - &self, - client: &ReqwestClient, - mut data: SendData, - ) -> Result { - let api_key = self.get_api_key()?; - patch_messages(&self.model.name, &api_key, &mut data.messages).await?; - let builder = self.request_builder(client, data)?; - send_message(builder, self.is_vl()).await - } - - async fn send_message_streaming_inner( - &self, - client: &ReqwestClient, - handler: &mut ReplyHandler, - mut data: SendData, - ) -> Result<()> { - let api_key = self.get_api_key()?; - patch_messages(&self.model.name, &api_key, &mut data.messages).await?; - let builder = self.request_builder(client, data)?; - send_message_streaming(builder, handler, self.is_vl()).await - } -} -- cgit v1.2.3