diff options
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, |
