summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-04-29 06:51:03 +0800
committerGitHub <noreply@github.com>2024-04-29 06:51:03 +0800
commit865be2bf75bb62b6aeee059f684400b4b9938a15 (patch)
treeb6296ccdde2af49251b02087b9f51c5239cc4d2b
parentb33e2da75efab4246f9f1d74a0be56aba1d8bd07 (diff)
downloadaichat-865be2bf75bb62b6aeee059f684400b4b9938a15.tar.gz
feat: non-streaming returns completion stats (#456)
-rwxr-xr-xArgcfile.sh54
-rw-r--r--models.yaml10
-rw-r--r--src/client/bedrock.rs51
-rw-r--r--src/client/claude.rs26
-rw-r--r--src/client/cohere.rs38
-rw-r--r--src/client/common.rs21
-rw-r--r--src/client/ernie.rs32
-rw-r--r--src/client/ollama.rs10
-rw-r--r--src/client/openai.rs23
-rw-r--r--src/client/qianwen.rs90
-rw-r--r--src/client/vertexai.rs25
-rw-r--r--src/main.rs4
-rw-r--r--src/repl/mod.rs2
-rw-r--r--src/serve.rs39
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())