summaryrefslogtreecommitdiffstats
path: root/src/client/ernie.rs
diff options
context:
space:
mode:
Diffstat (limited to 'src/client/ernie.rs')
-rw-r--r--src/client/ernie.rs32
1 files changed, 21 insertions, 11 deletions
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,