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/client | |
| parent | b33e2da75efab4246f9f1d74a0be56aba1d8bd07 (diff) | |
| download | aichat-865be2bf75bb62b6aeee059f684400b4b9938a15.tar.gz | |
feat: non-streaming returns completion stats (#456)
Diffstat (limited to 'src/client')
| -rw-r--r-- | src/client/bedrock.rs | 51 | ||||
| -rw-r--r-- | src/client/claude.rs | 26 | ||||
| -rw-r--r-- | src/client/cohere.rs | 38 | ||||
| -rw-r--r-- | src/client/common.rs | 21 | ||||
| -rw-r--r-- | src/client/ernie.rs | 32 | ||||
| -rw-r--r-- | src/client/ollama.rs | 10 | ||||
| -rw-r--r-- | src/client/openai.rs | 23 | ||||
| -rw-r--r-- | src/client/qianwen.rs | 90 | ||||
| -rw-r--r-- | src/client/vertexai.rs | 25 |
9 files changed, 203 insertions, 113 deletions
diff --git a/src/client/bedrock.rs b/src/client/bedrock.rs index 5f0a385..dd94a41 100644 --- a/src/client/bedrock.rs +++ b/src/client/bedrock.rs @@ -1,7 +1,8 @@ -use super::claude::claude_build_body; +use super::claude::{claude_build_body, claude_extract_completion}; use super::{ - catch_error, generate_prompt, BedrockClient, Client, ExtraConfig, Model, ModelConfig, - PromptFormat, PromptType, ReplyHandler, SendData, LLAMA2_PROMPT_FORMAT, LLAMA3_PROMPT_FORMAT, + catch_error, generate_prompt, BedrockClient, Client, CompletionStats, ExtraConfig, Model, + ModelConfig, PromptFormat, PromptType, ReplyHandler, SendData, LLAMA2_PROMPT_FORMAT, + LLAMA3_PROMPT_FORMAT, }; use crate::utils::PromptKind; @@ -40,7 +41,11 @@ pub struct BedrockConfig { impl Client for BedrockClient { client_common_fns!(); - async fn send_message_inner(&self, client: &ReqwestClient, data: SendData) -> Result<String> { + async fn send_message_inner( + &self, + client: &ReqwestClient, + data: SendData, + ) -> Result<(String, CompletionStats)> { let model_category = ModelCategory::from_str(&self.model.name)?; let builder = self.request_builder(client, data, &model_category)?; send_message(builder, &model_category).await @@ -124,7 +129,10 @@ impl BedrockClient { } } -async fn send_message(builder: RequestBuilder, model_category: &ModelCategory) -> Result<String> { +async fn send_message( + builder: RequestBuilder, + model_category: &ModelCategory, +) -> Result<(String, CompletionStats)> { let res = builder.send().await?; let status = res.status(); let data: Value = res.json().await?; @@ -133,15 +141,11 @@ async fn send_message(builder: RequestBuilder, model_category: &ModelCategory) - catch_error(&data, status.as_u16())?; } - let output = match model_category { - ModelCategory::Anthropic => data["content"][0]["text"].as_str(), - ModelCategory::MetaLlama2 | ModelCategory::MetaLlama3 => data["generation"].as_str(), - ModelCategory::Mistral => data["outputs"][0]["text"].as_str(), - }; - - let output = output.ok_or_else(|| anyhow!("Invalid response data: {data}"))?; - - Ok(output.to_string()) + match model_category { + ModelCategory::Anthropic => claude_extract_completion(&data), + ModelCategory::MetaLlama2 | ModelCategory::MetaLlama3 => llama_extract_completion(&data), + ModelCategory::Mistral => mistral_extrat_completion(&data), + } } async fn send_message_streaming( @@ -271,6 +275,25 @@ fn mistral_build_body(data: SendData, model: &Model) -> Result<Value> { Ok(body) } +fn llama_extract_completion(data: &Value) -> Result<(String, CompletionStats)> { + let text = data["generation"] + .as_str() + .ok_or_else(|| anyhow!("Invalid response data: {data}"))?; + let stats = CompletionStats { + id: None, + input_tokens: data["prompt_token_count"].as_u64(), + output_tokens: data["generation_token_count"].as_u64(), + }; + Ok((text.to_string(), stats)) +} + +fn mistral_extrat_completion(data: &Value) -> Result<(String, CompletionStats)> { + let text = data["outputs"][0]["text"] + .as_str() + .ok_or_else(|| anyhow!("Invalid response data: {data}"))?; + Ok((text.to_string(), CompletionStats::default())) +} + #[derive(Debug, Clone, Copy, PartialEq, Eq)] enum ModelCategory { Anthropic, diff --git a/src/client/claude.rs b/src/client/claude.rs index e4081e7..8bd87ee 100644 --- a/src/client/claude.rs +++ b/src/client/claude.rs @@ -1,6 +1,6 @@ use super::{ - catch_error, extract_system_message, ClaudeClient, ExtraConfig, ImageUrl, MessageContent, - MessageContentPart, Model, ModelConfig, PromptType, ReplyHandler, SendData, + catch_error, extract_system_message, ClaudeClient, CompletionStats, ExtraConfig, ImageUrl, + MessageContent, MessageContentPart, Model, ModelConfig, PromptType, ReplyHandler, SendData, }; use crate::utils::PromptKind; @@ -54,19 +54,14 @@ impl_client_trait!( claude_send_message_streaming ); -pub async fn claude_send_message(builder: RequestBuilder) -> Result<String> { +pub async fn claude_send_message(builder: RequestBuilder) -> Result<(String, CompletionStats)> { let res = builder.send().await?; let status = res.status(); let data: Value = res.json().await?; if status != 200 { catch_error(&data, status.as_u16())?; } - - let output = data["content"][0]["text"] - .as_str() - .ok_or_else(|| anyhow!("Invalid response data: {data}"))?; - - Ok(output.to_string()) + claude_extract_completion(&data) } pub async fn claude_send_message_streaming( @@ -195,3 +190,16 @@ pub fn claude_build_body(data: SendData, model: &Model) -> Result<Value> { } Ok(body) } + +pub fn claude_extract_completion(data: &Value) -> Result<(String, CompletionStats)> { + let text = data["content"][0]["text"] + .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"]["input_tokens"].as_u64(), + output_tokens: data["usage"]["output_tokens"].as_u64(), + }; + Ok((text.to_string(), stats)) +} diff --git a/src/client/cohere.rs b/src/client/cohere.rs index d99276c..6186718 100644 --- a/src/client/cohere.rs +++ b/src/client/cohere.rs @@ -1,11 +1,11 @@ use super::{ - catch_error, extract_system_message, json_stream, message::*, CohereClient, + catch_error, extract_system_message, json_stream, message::*, CohereClient, CompletionStats, ExtraConfig, Model, ModelConfig, PromptType, ReplyHandler, SendData, }; use crate::utils::PromptKind; -use anyhow::{bail, Result}; +use anyhow::{anyhow, bail, Result}; use reqwest::{Client as ReqwestClient, RequestBuilder}; use serde::Deserialize; use serde_json::{json, Value}; @@ -47,15 +47,15 @@ impl CohereClient { impl_client_trait!(CohereClient, send_message, send_message_streaming); -async fn send_message(builder: RequestBuilder) -> Result<String> { +async fn send_message(builder: RequestBuilder) -> Result<(String, CompletionStats)> { let res = builder.send().await?; let status = res.status(); let data: Value = res.json().await?; if status != 200 { catch_error(&data, status.as_u16())?; } - let output = extract_text(&data)?; - Ok(output.to_string()) + + cohere_extract_completion(&data) } async fn send_message_streaming(builder: RequestBuilder, handler: &mut ReplyHandler) -> Result<()> { @@ -65,10 +65,12 @@ async fn send_message_streaming(builder: RequestBuilder, handler: &mut ReplyHand let data: Value = res.json().await?; catch_error(&data, status.as_u16())?; } else { - let handle = |value: &str| -> Result<()> { - let value: Value = serde_json::from_str(value)?; - if let Some("text-generation") = value["event_type"].as_str() { - handler.text(extract_text(&value)?)?; + let handle = |data: &str| -> Result<()> { + let data: Value = serde_json::from_str(data)?; + if let Some("text-generation") = data["event_type"].as_str() { + if let Some(text) = data["text"].as_str() { + handler.text(text)?; + } } Ok(()) }; @@ -154,11 +156,15 @@ fn build_body(data: SendData, model: &Model) -> Result<Value> { Ok(body) } -fn extract_text(data: &Value) -> Result<&str> { - match data["text"].as_str() { - Some(text) => Ok(text), - None => { - bail!("Invalid response data: {data}") - } - } +fn cohere_extract_completion(data: &Value) -> Result<(String, CompletionStats)> { + let text = data["text"] + .as_str() + .ok_or_else(|| anyhow!("Invalid response data: {data}"))?; + + let stats = CompletionStats { + id: data["generation_id"].as_str().map(|v| v.to_string()), + input_tokens: data["meta"]["billed_units"]["input_tokens"].as_u64(), + output_tokens: data["meta"]["billed_units"]["output_tokens"].as_u64(), + }; + Ok((text.to_string(), stats)) } diff --git a/src/client/common.rs b/src/client/common.rs index 695245a..a58c962 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -126,7 +126,7 @@ macro_rules! register_client { client.set_model(model); } else { anyhow::bail!( - "The current model lacks the corresponding capability." + "The current model is incapable of doing that." ); } } @@ -260,7 +260,7 @@ macro_rules! impl_client_trait { &self, client: &reqwest::Client, data: $crate::client::SendData, - ) -> anyhow::Result<String> { + ) -> anyhow::Result<(String, $crate::client::CompletionStats)> { let builder = self.request_builder(client, data)?; $send_message(builder).await } @@ -330,11 +330,11 @@ pub trait Client: Sync + Send { Ok(client) } - async fn send_message(&self, input: Input) -> Result<String> { + async fn send_message(&self, input: Input) -> Result<(String, CompletionStats)> { let global_config = self.config().0; if global_config.read().dry_run { let content = global_config.read().echo_messages(&input); - return Ok(content); + return Ok((content, CompletionStats::default())); } let client = self.build_client()?; let data = global_config.read().prepare_send_data(&input, false)?; @@ -384,7 +384,11 @@ pub trait Client: Sync + Send { } } - async fn send_message_inner(&self, client: &ReqwestClient, data: SendData) -> Result<String>; + async fn send_message_inner( + &self, + client: &ReqwestClient, + data: SendData, + ) -> Result<(String, CompletionStats)>; async fn send_message_streaming_inner( &self, @@ -414,6 +418,13 @@ pub struct SendData { pub stream: bool, } +#[derive(Debug, Clone, Default)] +pub struct CompletionStats { + pub id: Option<String>, + pub input_tokens: Option<u64>, + pub output_tokens: Option<u64>, +} + pub type PromptType<'a> = (&'a str, &'a str, bool, PromptKind); pub fn create_config(list: &[PromptType], client: &str) -> Result<(String, Value)> { diff --git a/src/client/ernie.rs b/src/client/ernie.rs index 7695061..dc00f2f 100644 --- a/src/client/ernie.rs +++ b/src/client/ernie.rs @@ -1,6 +1,6 @@ use super::{ - maybe_catch_error, patch_system_message, Client, ErnieClient, ExtraConfig, Model, ModelConfig, - PromptType, ReplyHandler, SendData, + maybe_catch_error, patch_system_message, Client, CompletionStats, ErnieClient, ExtraConfig, + Model, ModelConfig, PromptType, ReplyHandler, SendData, }; use crate::utils::PromptKind; @@ -31,7 +31,6 @@ pub struct ErnieConfig { } impl ErnieClient { - pub const PROMPTS: [PromptType<'static>; 2] = [ ("api_key", "API Key:", true, PromptKind::String), ("secret_key", "Secret Key:", true, PromptKind::String), @@ -80,7 +79,11 @@ impl ErnieClient { impl Client for ErnieClient { client_common_fns!(); - async fn send_message_inner(&self, client: &ReqwestClient, data: SendData) -> Result<String> { + async fn send_message_inner( + &self, + client: &ReqwestClient, + data: SendData, + ) -> Result<(String, CompletionStats)> { self.prepare_access_token().await?; let builder = self.request_builder(client, data)?; send_message(builder).await @@ -98,15 +101,10 @@ impl Client for ErnieClient { } } -async fn send_message(builder: RequestBuilder) -> Result<String> { +async fn send_message(builder: RequestBuilder) -> Result<(String, CompletionStats)> { let data: Value = builder.send().await?.json().await?; maybe_catch_error(&data)?; - - let output = data["result"] - .as_str() - .ok_or_else(|| anyhow!("Unexpected response {data}"))?; - - Ok(output.to_string()) + extract_completion_text(&data) } async fn send_message_streaming(builder: RequestBuilder, handler: &mut ReplyHandler) -> Result<()> { @@ -186,6 +184,18 @@ fn build_body(data: SendData, model: &Model) -> Value { body } +fn extract_completion_text(data: &Value) -> Result<(String, CompletionStats)> { + let text = data["result"] + .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)) +} + async fn fetch_access_token( client: &reqwest::Client, api_key: &str, diff --git a/src/client/ollama.rs b/src/client/ollama.rs index 5ec50a1..ec83cbd 100644 --- a/src/client/ollama.rs +++ b/src/client/ollama.rs @@ -1,6 +1,6 @@ use super::{ - catch_error, message::*, ExtraConfig, Model, ModelConfig, OllamaClient, PromptType, - ReplyHandler, SendData, + catch_error, message::*, CompletionStats, ExtraConfig, Model, ModelConfig, OllamaClient, + PromptType, ReplyHandler, SendData, }; use crate::utils::PromptKind; @@ -59,17 +59,17 @@ impl OllamaClient { impl_client_trait!(OllamaClient, send_message, send_message_streaming); -async fn send_message(builder: RequestBuilder) -> Result<String> { +async fn send_message(builder: RequestBuilder) -> Result<(String, CompletionStats)> { let res = builder.send().await?; let status = res.status(); let data = res.json().await?; if status != 200 { catch_error(&data, status.as_u16())?; } - let output = data["message"]["content"] + let text = data["message"]["content"] .as_str() .ok_or_else(|| anyhow!("Invalid response data: {data}"))?; - Ok(output.to_string()) + Ok((text.to_string(), CompletionStats::default())) } async fn send_message_streaming(builder: RequestBuilder, handler: &mut ReplyHandler) -> Result<()> { 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<String> { +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<String> { 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, 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<String> { - 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<Message>) -> Result<()> { for message in messages { @@ -283,31 +321,3 @@ async fn upload(model: &str, api_key: &str, url: &str) -> Result<String> { } 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<String> { - 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 - } -} diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs index 317ed2a..f4079f0 100644 --- a/src/client/vertexai.rs +++ b/src/client/vertexai.rs @@ -1,7 +1,7 @@ use super::claude::{claude_build_body, claude_send_message, claude_send_message_streaming}; use super::{ - catch_error, json_stream, message::*, patch_system_message, Client, ExtraConfig, Model, - ModelConfig, PromptType, ReplyHandler, SendData, VertexAIClient, + catch_error, json_stream, message::*, patch_system_message, Client, CompletionStats, + ExtraConfig, Model, ModelConfig, PromptType, ReplyHandler, SendData, VertexAIClient, }; use crate::utils::PromptKind; @@ -81,7 +81,11 @@ impl VertexAIClient { impl Client for VertexAIClient { client_common_fns!(); - async fn send_message_inner(&self, client: &ReqwestClient, data: SendData) -> Result<String> { + async fn send_message_inner( + &self, + client: &ReqwestClient, + data: SendData, + ) -> Result<(String, CompletionStats)> { let model_category = ModelCategory::from_str(&self.model.name)?; self.prepare_access_token().await?; let builder = self.request_builder(client, data, &model_category)?; @@ -107,15 +111,14 @@ impl Client for VertexAIClient { } } -pub async fn gemini_send_message(builder: RequestBuilder) -> Result<String> { +pub async fn gemini_send_message(builder: RequestBuilder) -> Result<(String, CompletionStats)> { let res = builder.send().await?; let status = res.status(); let data: Value = res.json().await?; if status != 200 { catch_error(&data, status.as_u16())?; } - let output = gemini_extract_text(&data)?; - Ok(output.to_string()) + gemini_extract_completion_text(&data) } pub async fn gemini_send_message_streaming( @@ -138,6 +141,16 @@ pub async fn gemini_send_message_streaming( Ok(()) } +fn gemini_extract_completion_text(data: &Value) -> Result<(String, CompletionStats)> { + let text = gemini_extract_text(data)?; + let stats = CompletionStats { + id: None, + input_tokens: data["usageMetadata"]["promptTokenCount"].as_u64(), + output_tokens: data["usageMetadata"]["candidatesTokenCount"].as_u64(), + }; + Ok((text.to_string(), stats)) +} + fn gemini_extract_text(data: &Value) -> Result<&str> { match data["candidates"][0]["content"]["parts"][0]["text"].as_str() { Some(text) => Ok(text), |
