summaryrefslogtreecommitdiffstats
path: root/src/client
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
parentb33e2da75efab4246f9f1d74a0be56aba1d8bd07 (diff)
downloadaichat-865be2bf75bb62b6aeee059f684400b4b9938a15.tar.gz
feat: non-streaming returns completion stats (#456)
Diffstat (limited to 'src/client')
-rw-r--r--src/client/bedrock.rs51
-rw-r--r--src/client/claude.rs26
-rw-r--r--src/client/cohere.rs38
-rw-r--r--src/client/common.rs21
-rw-r--r--src/client/ernie.rs32
-rw-r--r--src/client/ollama.rs10
-rw-r--r--src/client/openai.rs23
-rw-r--r--src/client/qianwen.rs90
-rw-r--r--src/client/vertexai.rs25
9 files changed, 203 insertions, 113 deletions
diff --git a/src/client/bedrock.rs b/src/client/bedrock.rs
index 5f0a385..dd94a41 100644
--- a/src/client/bedrock.rs
+++ b/src/client/bedrock.rs
@@ -1,7 +1,8 @@
-use super::claude::claude_build_body;
+use super::claude::{claude_build_body, claude_extract_completion};
use super::{
- catch_error, generate_prompt, BedrockClient, Client, ExtraConfig, Model, ModelConfig,
- PromptFormat, PromptType, ReplyHandler, SendData, LLAMA2_PROMPT_FORMAT, LLAMA3_PROMPT_FORMAT,
+ catch_error, generate_prompt, BedrockClient, Client, CompletionStats, ExtraConfig, Model,
+ ModelConfig, PromptFormat, PromptType, ReplyHandler, SendData, LLAMA2_PROMPT_FORMAT,
+ LLAMA3_PROMPT_FORMAT,
};
use crate::utils::PromptKind;
@@ -40,7 +41,11 @@ pub struct BedrockConfig {
impl Client for BedrockClient {
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)> {
let model_category = ModelCategory::from_str(&self.model.name)?;
let builder = self.request_builder(client, data, &model_category)?;
send_message(builder, &model_category).await
@@ -124,7 +129,10 @@ impl BedrockClient {
}
}
-async fn send_message(builder: RequestBuilder, model_category: &ModelCategory) -> Result<String> {
+async fn send_message(
+ builder: RequestBuilder,
+ model_category: &ModelCategory,
+) -> Result<(String, CompletionStats)> {
let res = builder.send().await?;
let status = res.status();
let data: Value = res.json().await?;
@@ -133,15 +141,11 @@ async fn send_message(builder: RequestBuilder, model_category: &ModelCategory) -
catch_error(&data, status.as_u16())?;
}
- let output = match model_category {
- ModelCategory::Anthropic => data["content"][0]["text"].as_str(),
- ModelCategory::MetaLlama2 | ModelCategory::MetaLlama3 => data["generation"].as_str(),
- ModelCategory::Mistral => data["outputs"][0]["text"].as_str(),
- };
-
- let output = output.ok_or_else(|| anyhow!("Invalid response data: {data}"))?;
-
- Ok(output.to_string())
+ match model_category {
+ ModelCategory::Anthropic => claude_extract_completion(&data),
+ ModelCategory::MetaLlama2 | ModelCategory::MetaLlama3 => llama_extract_completion(&data),
+ ModelCategory::Mistral => mistral_extrat_completion(&data),
+ }
}
async fn send_message_streaming(
@@ -271,6 +275,25 @@ fn mistral_build_body(data: SendData, model: &Model) -> Result<Value> {
Ok(body)
}
+fn llama_extract_completion(data: &Value) -> Result<(String, CompletionStats)> {
+ let text = data["generation"]
+ .as_str()
+ .ok_or_else(|| anyhow!("Invalid response data: {data}"))?;
+ let stats = CompletionStats {
+ id: None,
+ input_tokens: data["prompt_token_count"].as_u64(),
+ output_tokens: data["generation_token_count"].as_u64(),
+ };
+ Ok((text.to_string(), stats))
+}
+
+fn mistral_extrat_completion(data: &Value) -> Result<(String, CompletionStats)> {
+ let text = data["outputs"][0]["text"]
+ .as_str()
+ .ok_or_else(|| anyhow!("Invalid response data: {data}"))?;
+ Ok((text.to_string(), CompletionStats::default()))
+}
+
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum ModelCategory {
Anthropic,
diff --git a/src/client/claude.rs b/src/client/claude.rs
index e4081e7..8bd87ee 100644
--- a/src/client/claude.rs
+++ b/src/client/claude.rs
@@ -1,6 +1,6 @@
use super::{
- catch_error, extract_system_message, ClaudeClient, ExtraConfig, ImageUrl, MessageContent,
- MessageContentPart, Model, ModelConfig, PromptType, ReplyHandler, SendData,
+ catch_error, extract_system_message, ClaudeClient, CompletionStats, ExtraConfig, ImageUrl,
+ MessageContent, MessageContentPart, Model, ModelConfig, PromptType, ReplyHandler, SendData,
};
use crate::utils::PromptKind;
@@ -54,19 +54,14 @@ impl_client_trait!(
claude_send_message_streaming
);
-pub async fn claude_send_message(builder: RequestBuilder) -> Result<String> {
+pub async fn claude_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 = data["content"][0]["text"]
- .as_str()
- .ok_or_else(|| anyhow!("Invalid response data: {data}"))?;
-
- Ok(output.to_string())
+ claude_extract_completion(&data)
}
pub async fn claude_send_message_streaming(
@@ -195,3 +190,16 @@ pub fn claude_build_body(data: SendData, model: &Model) -> Result<Value> {
}
Ok(body)
}
+
+pub fn claude_extract_completion(data: &Value) -> Result<(String, CompletionStats)> {
+ let text = data["content"][0]["text"]
+ .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"]["input_tokens"].as_u64(),
+ output_tokens: data["usage"]["output_tokens"].as_u64(),
+ };
+ Ok((text.to_string(), stats))
+}
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))
}
diff --git a/src/client/common.rs b/src/client/common.rs
index 695245a..a58c962 100644
--- a/src/client/common.rs
+++ b/src/client/common.rs
@@ -126,7 +126,7 @@ macro_rules! register_client {
client.set_model(model);
} else {
anyhow::bail!(
- "The current model lacks the corresponding capability."
+ "The current model is incapable of doing that."
);
}
}
@@ -260,7 +260,7 @@ macro_rules! impl_client_trait {
&self,
client: &reqwest::Client,
data: $crate::client::SendData,
- ) -> anyhow::Result<String> {
+ ) -> anyhow::Result<(String, $crate::client::CompletionStats)> {
let builder = self.request_builder(client, data)?;
$send_message(builder).await
}
@@ -330,11 +330,11 @@ pub trait Client: Sync + Send {
Ok(client)
}
- async fn send_message(&self, input: Input) -> Result<String> {
+ async fn send_message(&self, input: Input) -> Result<(String, CompletionStats)> {
let global_config = self.config().0;
if global_config.read().dry_run {
let content = global_config.read().echo_messages(&input);
- return Ok(content);
+ return Ok((content, CompletionStats::default()));
}
let client = self.build_client()?;
let data = global_config.read().prepare_send_data(&input, false)?;
@@ -384,7 +384,11 @@ pub trait Client: Sync + Send {
}
}
- 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)>;
async fn send_message_streaming_inner(
&self,
@@ -414,6 +418,13 @@ pub struct SendData {
pub stream: bool,
}
+#[derive(Debug, Clone, Default)]
+pub struct CompletionStats {
+ pub id: Option<String>,
+ pub input_tokens: Option<u64>,
+ pub output_tokens: Option<u64>,
+}
+
pub type PromptType<'a> = (&'a str, &'a str, bool, PromptKind);
pub fn create_config(list: &[PromptType], client: &str) -> Result<(String, Value)> {
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,
diff --git a/src/client/ollama.rs b/src/client/ollama.rs
index 5ec50a1..ec83cbd 100644
--- a/src/client/ollama.rs
+++ b/src/client/ollama.rs
@@ -1,6 +1,6 @@
use super::{
- catch_error, message::*, ExtraConfig, Model, ModelConfig, OllamaClient, PromptType,
- ReplyHandler, SendData,
+ catch_error, message::*, CompletionStats, ExtraConfig, Model, ModelConfig, OllamaClient,
+ PromptType, ReplyHandler, SendData,
};
use crate::utils::PromptKind;
@@ -59,17 +59,17 @@ impl OllamaClient {
impl_client_trait!(OllamaClient, 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 = res.json().await?;
if status != 200 {
catch_error(&data, status.as_u16())?;
}
- let output = data["message"]["content"]
+ let text = data["message"]["content"]
.as_str()
.ok_or_else(|| anyhow!("Invalid response data: {data}"))?;
- Ok(output.to_string())
+ Ok((text.to_string(), CompletionStats::default()))
}
async fn send_message_streaming(builder: RequestBuilder, handler: &mut ReplyHandler) -> Result<()> {
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,
diff --git a/src/client/qianwen.rs b/src/client/qianwen.rs
index 64d27d4..8a5a6d5 100644
--- a/src/client/qianwen.rs
+++ b/src/client/qianwen.rs
@@ -1,6 +1,6 @@
use super::{
- maybe_catch_error, message::*, Client, ExtraConfig, Model, ModelConfig, PromptType,
- QianwenClient, ReplyHandler, SendData,
+ maybe_catch_error, message::*, Client, CompletionStats, ExtraConfig, Model, ModelConfig,
+ PromptType, QianwenClient, ReplyHandler, SendData,
};
use crate::utils::{sha256sum, PromptKind};
@@ -69,19 +69,39 @@ impl QianwenClient {
}
}
-async fn send_message(builder: RequestBuilder, is_vl: bool) -> Result<String> {
- let data: Value = builder.send().await?.json().await?;
- maybe_catch_error(&data)?;
+#[async_trait]
+impl Client for QianwenClient {
+ client_common_fns!();
- let output = if is_vl {
- data["output"]["choices"][0]["message"]["content"][0]["text"].as_str()
- } else {
- data["output"]["text"].as_str()
- };
+ async fn send_message_inner(
+ &self,
+ client: &ReqwestClient,
+ mut data: SendData,
+ ) -> Result<(String, CompletionStats)> {
+ let api_key = self.get_api_key()?;
+ patch_messages(&self.model.name, &api_key, &mut data.messages).await?;
+ let builder = self.request_builder(client, data)?;
+ send_message(builder, self.is_vl()).await
+ }
- let output = output.ok_or_else(|| anyhow!("Unexpected response {data}"))?;
+ async fn send_message_streaming_inner(
+ &self,
+ client: &ReqwestClient,
+ handler: &mut ReplyHandler,
+ mut data: SendData,
+ ) -> Result<()> {
+ let api_key = self.get_api_key()?;
+ patch_messages(&self.model.name, &api_key, &mut data.messages).await?;
+ let builder = self.request_builder(client, data)?;
+ send_message_streaming(builder, handler, self.is_vl()).await
+ }
+}
- Ok(output.to_string())
+async fn send_message(builder: RequestBuilder, is_vl: bool) -> Result<(String, CompletionStats)> {
+ let data: Value = builder.send().await?.json().await?;
+ maybe_catch_error(&data)?;
+
+ extract_completion_text(&data, is_vl)
}
async fn send_message_streaming(
@@ -190,6 +210,24 @@ fn build_body(data: SendData, model: &Model, is_vl: bool) -> Result<(Value, bool
Ok((body, has_upload))
}
+fn extract_completion_text(data: &Value, is_vl: bool) -> Result<(String, CompletionStats)> {
+ let err = || anyhow!("Invalid response data: {data}");
+ let text = if is_vl {
+ data["output"]["choices"][0]["message"]["content"][0]["text"]
+ .as_str()
+ .ok_or_else(err)?
+ } else {
+ data["output"]["text"].as_str().ok_or_else(err)?
+ };
+ let stats = CompletionStats {
+ id: data["request_id"].as_str().map(|v| v.to_string()),
+ input_tokens: data["usage"]["input_tokens"].as_u64(),
+ output_tokens: data["usage"]["output_tokens"].as_u64(),
+ };
+
+ Ok((text.to_string(), stats))
+}
+
/// Patch messsages, upload embedded images to oss
async fn patch_messages(model: &str, api_key: &str, messages: &mut Vec<Message>) -> Result<()> {
for message in messages {
@@ -283,31 +321,3 @@ async fn upload(model: &str, api_key: &str, url: &str) -> Result<String> {
}
Ok(format!("oss://{key}"))
}
-
-#[async_trait]
-impl Client for QianwenClient {
- client_common_fns!();
-
- async fn send_message_inner(
- &self,
- client: &ReqwestClient,
- mut data: SendData,
- ) -> Result<String> {
- let api_key = self.get_api_key()?;
- patch_messages(&self.model.name, &api_key, &mut data.messages).await?;
- let builder = self.request_builder(client, data)?;
- send_message(builder, self.is_vl()).await
- }
-
- async fn send_message_streaming_inner(
- &self,
- client: &ReqwestClient,
- handler: &mut ReplyHandler,
- mut data: SendData,
- ) -> Result<()> {
- let api_key = self.get_api_key()?;
- patch_messages(&self.model.name, &api_key, &mut data.messages).await?;
- let builder = self.request_builder(client, data)?;
- send_message_streaming(builder, handler, self.is_vl()).await
- }
-}
diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs
index 317ed2a..f4079f0 100644
--- a/src/client/vertexai.rs
+++ b/src/client/vertexai.rs
@@ -1,7 +1,7 @@
use super::claude::{claude_build_body, claude_send_message, claude_send_message_streaming};
use super::{
- catch_error, json_stream, message::*, patch_system_message, Client, ExtraConfig, Model,
- ModelConfig, PromptType, ReplyHandler, SendData, VertexAIClient,
+ catch_error, json_stream, message::*, patch_system_message, Client, CompletionStats,
+ ExtraConfig, Model, ModelConfig, PromptType, ReplyHandler, SendData, VertexAIClient,
};
use crate::utils::PromptKind;
@@ -81,7 +81,11 @@ impl VertexAIClient {
impl Client for VertexAIClient {
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)> {
let model_category = ModelCategory::from_str(&self.model.name)?;
self.prepare_access_token().await?;
let builder = self.request_builder(client, data, &model_category)?;
@@ -107,15 +111,14 @@ impl Client for VertexAIClient {
}
}
-pub async fn gemini_send_message(builder: RequestBuilder) -> Result<String> {
+pub async fn gemini_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 = gemini_extract_text(&data)?;
- Ok(output.to_string())
+ gemini_extract_completion_text(&data)
}
pub async fn gemini_send_message_streaming(
@@ -138,6 +141,16 @@ pub async fn gemini_send_message_streaming(
Ok(())
}
+fn gemini_extract_completion_text(data: &Value) -> Result<(String, CompletionStats)> {
+ let text = gemini_extract_text(data)?;
+ let stats = CompletionStats {
+ id: None,
+ input_tokens: data["usageMetadata"]["promptTokenCount"].as_u64(),
+ output_tokens: data["usageMetadata"]["candidatesTokenCount"].as_u64(),
+ };
+ Ok((text.to_string(), stats))
+}
+
fn gemini_extract_text(data: &Value) -> Result<&str> {
match data["candidates"][0]["content"]["parts"][0]["text"].as_str() {
Some(text) => Ok(text),