summaryrefslogtreecommitdiffstats
path: root/src/client/openai.rs
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 /src/client/openai.rs
parentb33e2da75efab4246f9f1d74a0be56aba1d8bd07 (diff)
downloadaichat-865be2bf75bb62b6aeee059f684400b4b9938a15.tar.gz
feat: non-streaming returns completion stats (#456)
Diffstat (limited to 'src/client/openai.rs')
-rw-r--r--src/client/openai.rs23
1 files changed, 16 insertions, 7 deletions
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,