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