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/openai.rs | |
| parent | b33e2da75efab4246f9f1d74a0be56aba1d8bd07 (diff) | |
| download | aichat-865be2bf75bb62b6aeee059f684400b4b9938a15.tar.gz | |
feat: non-streaming returns completion stats (#456)
Diffstat (limited to 'src/client/openai.rs')
| -rw-r--r-- | src/client/openai.rs | 23 |
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, |
