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 | |
| parent | b33e2da75efab4246f9f1d74a0be56aba1d8bd07 (diff) | |
| download | aichat-865be2bf75bb62b6aeee059f684400b4b9938a15.tar.gz | |
feat: non-streaming returns completion stats (#456)
| -rwxr-xr-x | Argcfile.sh | 54 | ||||
| -rw-r--r-- | models.yaml | 10 | ||||
| -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 | ||||
| -rw-r--r-- | src/main.rs | 4 | ||||
| -rw-r--r-- | src/repl/mod.rs | 2 | ||||
| -rw-r--r-- | src/serve.rs | 39 |
14 files changed, 271 insertions, 154 deletions
diff --git a/Argcfile.sh b/Argcfile.sh index 1e5229a..b830e28 100755 --- a/Argcfile.sh +++ b/Argcfile.sh @@ -36,7 +36,7 @@ test-clients() { } # @cmd Test proxy server -# @option -m --model=default +# @option -m --model[`_choice_model`] # @flag -S --no-stream # @arg text~ test-server() { @@ -44,9 +44,9 @@ test-server() { if [[ -n "$argc_no_stream" ]]; then args+=("-S") fi - argc generic-chat "${args[@]}" \ + argc chat-llm "${args[@]}" \ --api-base http://localhost:8000/v1 \ - --model $argc_model \ + --model "${argc_model:-default}" \ "$@" } @@ -56,7 +56,7 @@ test-server() { # @option -m --model! $$ # @flag -S --no-stream # @arg text~ -generic-chat() { +chat-llm() { curl_args="$CURL_ARGS" _openai_chat "$@" } @@ -64,7 +64,7 @@ generic-chat() { # @cmd List models by openai-comptabile api # @option --api-base! $$ # @option --api-key! $$ -generic-models() { +models-llm() { curl_args="$CURL_ARGS" _openai_models } @@ -75,7 +75,7 @@ generic-models() { # @option -m --model=gpt-3.5-turbo $OPENAI_MODEL # @flag -S --no-stream # @arg text~ -openai-chat() { +chat-openai() { api_base=https://api.openai.com/v1 api_key=$OPENAI_API_KEY curl_args="-i $OPENAI_CURL_ARGS" @@ -84,7 +84,7 @@ openai-chat() { # @cmd List openai models # @env OPENAI_API_KEY! -openai-models() { +models-openai() { api_base=https://api.openai.com/v1 api_key=$OPENAI_API_KEY curl_args="$OPENAI_CURL_ARGS" @@ -96,7 +96,7 @@ openai-models() { # @option -m --model=gemini-1.0-pro-latest $GEMINI_MODEL # @flag -S --no-stream # @arg text~ -gemini-chat() { +chat-gemini() { method="streamGenerateContent" if [[ -n "$argc_no_stream" ]]; then method="generateContent" @@ -112,7 +112,7 @@ gemini-chat() { # @cmd List gemini models # @env GEMINI_API_KEY! -gemini-models() { +models-gemini() { _wrapper curl $GEMINI_CURL_ARGS "https://generativelanguage.googleapis.com/v1beta/models?key=${GEMINI_API_KEY}" \ -H 'Content-Type: application/json' \ @@ -123,7 +123,7 @@ gemini-models() { # @option -m --model=claude-3-haiku-20240307 $CLAUDE_MODEL # @flag -S --no-stream # @arg text~ -claude-chat() { +chat-claude() { _wrapper curl -i $CLAUDE_CURL_ARGS https://api.anthropic.com/v1/messages \ -X POST \ -H 'content-type: application/json' \ @@ -143,7 +143,7 @@ claude-chat() { # @option -m --model=mistral-small-latest $MISTRAL_MODEL # @flag -S --no-stream # @arg text~ -mistral-chat() { +chat-mistral() { api_base=https://api.mistral.ai/v1 api_key=$MISTRAL_API_KEY curl_args="$MISTRAL_CURL_ARGS" @@ -152,7 +152,7 @@ mistral-chat() { # @cmd List mistral models # @env MISTRAL_API_KEY! -mistral-models() { +models-mistral() { api_base=https://api.mistral.ai/v1 api_key=$MISTRAL_API_KEY curl_args="$MISTRAL_CURL_ARGS" @@ -164,7 +164,7 @@ mistral-models() { # @option -m --model=command-r $COHERE_MODEL # @flag -S --no-stream # @arg text~ -cohere-chat() { +chat-cohere() { _wrapper curl -i $COHERE_CURL_ARGS https://api.cohere.ai/v1/chat \ -X POST \ -H 'Content-Type: application/json' \ @@ -179,7 +179,7 @@ cohere-chat() { # @cmd List cohere models # @env COHERE_API_KEY! -cohere-models() { +models-cohere() { _wrapper curl $COHERE_CURL_ARGS https://api.cohere.ai/v1/models \ -H "Authorization: Bearer $COHERE_API_KEY" \ @@ -190,7 +190,7 @@ cohere-models() { # @option -m --model=sonar-small-chat $PERPLEXITY_MODEL # @flag -S --no-stream # @arg text~ -perplexity-chat() { +chat-perplexity() { api_base=https://api.perplexity.ai api_key=$PERPLEXITY_API_KEY curl_args="$PERPLEXITY_CURL_ARGS" @@ -202,7 +202,7 @@ perplexity-chat() { # @option -m --model=llama3-70b-8192 $GROQ_MODEL # @flag -S --no-stream # @arg text~ -groq-chat() { +chat-groq() { api_base=https://api.groq.com/openai/v1 api_key=$GROQ_API_KEY curl_args="$GROQ_CURL_ARGS" @@ -211,7 +211,7 @@ groq-chat() { # @cmd List groq models # @env GROQ_API_KEY! -groq-models() { +models-groq() { api_base=https://api.groq.com/openai/v1 api_key=$GROQ_API_KEY curl_args="$GROQ_CURL_ARGS" @@ -222,7 +222,7 @@ groq-models() { # @option -m --model=codegemma $OLLAMA_MODEL # @flag -S --no-stream # @arg text~ -ollama-chat() { +chat-ollama() { _wrapper curl -i $OLLAMA_CURL_ARGS http://localhost:11434/api/chat \ -X POST \ -H 'Content-Type: application/json' \ @@ -240,7 +240,7 @@ ollama-chat() { # @option -m --model=gemini-1.0-pro $VERTEXAI_GEMINI_MODEL # @flag -S --no-stream # @arg text~ -vertexai-gemini-chat() { +chat-vertexai-gemini() { api_key="$(gcloud auth print-access-token)" func="streamGenerateContent" if [[ -n "$argc_no_stream" ]]; then @@ -264,7 +264,7 @@ vertexai-gemini-chat() { # @option -m --model=claude-3-haiku@20240307 $VERTEXAI_CLAUDE_MODEL # @flag -S --no-stream # @arg text~ -vertexai-claude-chat() { +chat-vertexai-claude() { api_key="$(gcloud auth print-access-token)" url=https://$VERTEXAI_LOCATION-aiplatform.googleapis.com/v1/projects/$VERTEXAI_PROJECT_ID/locations/$VERTEXAI_LOCATION/publishers/anthropic/models/$argc_model:streamRawPredict _wrapper curl -i $VERTEXAI_CURL_ARGS $url \ @@ -283,7 +283,7 @@ vertexai-claude-chat() { # @meta require-tools aws # @option -m --model=mistral.mistral-7b-instruct-v0:2 $BEDROCK_MODEL # @env AWS_REGION=us-east-1 -bedrock-chat() { +chat-bedrock() { file="$(mktemp)" case "$argc_model" in mistral.* | meta.*) @@ -315,7 +315,7 @@ bedrock-chat() { # @option -m --model=ernie-tiny-8k $ERNIE_MODEL # @flag -S --no-stream # @arg text~ -ernie-chat() { +chat-ernie() { ACCESS_TOKEN="$(curl -fsSL "https://aip.baidubce.com/oauth/2.0/token?grant_type=client_credentials&client_id=$ERNIE_API_KEY&client_secret=$ERNIE_SECRET_KEY" | jq -r '.access_token')" _wrapper curl -i $ERNIE_CURL_ARGS "https://aip.baidubce.com/rpc/2.0/ai_custom/v1/wenxinworkshop/chat/$argc_model?access_token=$ACCESS_TOKEN" \ -X POST \ @@ -331,7 +331,7 @@ ernie-chat() { # @option -m --model=qwen-turbo $QIANWEN_MODEL # @flag -S --no-stream # @arg text~ -qianwen-chat() { +chat-qianwen() { stream_args="-H X-DashScope-SSE:enable" parameters_args='{"incremental_output": true}' if [[ -n "$argc_no_stream" ]]; then @@ -358,7 +358,7 @@ qianwen-chat() { # @option -m --model=moonshot-v1-8k @MOONSHOT_MODEL # @flag -S --no-stream # @arg text~ -moonshot-chat() { +chat-moonshot() { api_base=https://api.moonshot.cn/v1 api_key=$MOONSHOT_API_KEY curl_args="$MOONSHOT_CURL_ARGS" @@ -367,13 +367,17 @@ moonshot-chat() { # @cmd List moonshot models # @env MOONSHOT_API_KEY! -moonshot-models() { +models-moonshot() { api_base=https://api.moonshot.cn/v1 api_key=$MOONSHOT_API_KEY curl_args="$MOONSHOT_CURL_ARGS" _openai_models } +_choice_model() { + aichat --list-models +} + _argc_before() { stream="true" if [[ -n "$argc_no_stream" ]]; then diff --git a/models.yaml b/models.yaml index bfa0e4a..5ef72dd 100644 --- a/models.yaml +++ b/models.yaml @@ -203,19 +203,19 @@ models: - name: llama3-8b-8192 max_input_tokens: 8192 - max_output_tokens: 8192 + max_output_tokens?: 8192 - name: llama3-70b-8192 max_input_tokens: 8192 - max_output_tokens: 8192 + max_output_tokens?: 8192 - name: llama2-70b-4096 max_input_tokens: 4096 - max_output_tokens: 4096 + max_output_tokens?: 4096 - name: mixtral-8x7b-32768 max_input_tokens: 32768 - max_output_tokens: 32768 + max_output_tokens?: 32768 - name: gemma-7b-it max_input_tokens: 8192 - max_output_tokens: 8192 + max_output_tokens?: 8192 - type: vertexai # docs: 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), diff --git a/src/main.rs b/src/main.rs index f691b6f..01f4bfb 100644 --- a/src/main.rs +++ b/src/main.rs @@ -143,7 +143,7 @@ async fn start_directive( let is_terminal_stdout = stdout().is_terminal(); let extract_code = !is_terminal_stdout && code_mode; let output = if no_stream || extract_code { - let output = client.send_message(input.clone()).await?; + let (output, _) = client.send_message(input.clone()).await?; let output = if extract_code && output.trim_start().starts_with("```") { extract_block(&output) } else { @@ -181,7 +181,7 @@ async fn execute(config: &GlobalConfig, mut input: Input) -> Result<()> { tokio::spawn(run_spinner(" Generating", spinner_rx)); let ret = client.send_message(input.clone()).await; let _ = spinner_tx.send(()); - let mut eval_str = ret?; + let (mut eval_str, _) = ret?; if let Ok(true) = CODE_BLOCK_RE.is_match(&eval_str) { eval_str = extract_block(&eval_str); } diff --git a/src/repl/mod.rs b/src/repl/mod.rs index 6306f0b..21b3ae9 100644 --- a/src/repl/mod.rs +++ b/src/repl/mod.rs @@ -447,7 +447,7 @@ async fn compress_session(config: &GlobalConfig) -> Result<()> { ); let mut client = init_client(config)?; ensure_model_capabilities(client.as_mut(), input.required_capabilities())?; - let summary = client.send_message(input).await?; + let (summary, _) = client.send_message(input).await?; config.write().compress_session(&summary); Ok(()) } 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()) |
