use crate::config::Config;
use crate::errors::Error;
use crate::types::{Content, GenerateContentRequest, GenerateContentResponse, Part, Role};
use backon::{ExponentialBuilder, Retryable};
use reqwest::Client;
use serde::{de::DeserializeOwned, Deserialize, Serialize};
const API_BASE_URL: &str = "https://generativelanguage.googleapis.com/v1beta";
pub struct GeminiClient {
client: Client,
config: Config,
base_url: String,
}
impl GeminiClient {
pub fn new() -> Result<Self, Error> {
let config = Config::new()?;
let client = Client::new();
Ok(Self {
client,
config,
base_url: API_BASE_URL.to_string(),
})
}
#[cfg(test)]
pub(crate) fn with_config(config: Config) -> Self {
Self {
client: Client::new(),
config,
base_url: API_BASE_URL.to_string(),
}
}
#[cfg(test)]
pub(crate) fn with_base_url(mut self, base_url: &str) -> Self {
self.base_url = base_url.to_string();
self
}
pub async fn list_models(&self) -> Result<ListModelsResponse, Error> {
let url = format!("{}/models", self.base_url);
self.get_with_retry(&url).await
}
pub async fn generate_text(&self, model: &str, prompt: &str) -> Result<String, Error> {
let url = format!("{}/models/{}:generateContent", self.base_url, model);
let request = GenerateContentRequest {
contents: vec![Content {
role: Role::User,
parts: vec![Part {
text: prompt.to_string(),
}],
}],
};
let response: GenerateContentResponse = self.post_with_retry(&url, &request).await?;
if let Some(candidate) = response.candidates.first() {
if let Some(part) = candidate.content.parts.first() {
return Ok(part.text.clone());
}
}
Err(Error::Api("No content found in response".to_string()))
}
async fn get_with_retry<R>(&self, url: &str) -> Result<R, Error>
where
R: DeserializeOwned,
{
let url = url.to_string();
let api_key = self.config.api_key().to_string();
let client = self.client.clone();
(|| async {
let response = client
.get(&url)
.header("x-goog-api-key", &api_key)
.send()
.await
.map_err(Error::Network)?;
if response.status().is_server_error() {
let status_code = response.status().as_u16();
let error_text = response.text().await.unwrap_or_default();
return Err(Error::Api(format!("HTTP {}: {}", status_code, error_text)));
}
if !response.status().is_success() {
return Err(Error::Api(response.text().await.unwrap_or_default()));
}
response.json::<R>().await.map_err(Error::Json)
})
.retry(ExponentialBuilder::default())
.when(|e| match e {
Error::Network(_) => true,
Error::Api(msg) => {
msg.contains("503")
|| msg.contains("502")
|| msg.contains("500")
|| msg.contains("504")
}
_ => false,
})
.await
}
async fn post_with_retry<T, R>(&self, url: &str, body: &T) -> Result<R, Error>
where
T: Serialize + Clone,
R: DeserializeOwned,
{
let url = url.to_string();
let api_key = self.config.api_key().to_string();
let client = self.client.clone();
let body = body.clone();
(|| async {
let response = client
.post(&url)
.header("x-goog-api-key", &api_key)
.json(&body)
.send()
.await
.map_err(Error::Network)?;
if response.status().is_server_error() {
let status_code = response.status().as_u16();
let error_text = response.text().await.unwrap_or_default();
return Err(Error::Api(format!("HTTP {}: {}", status_code, error_text)));
}
if !response.status().is_success() {
return Err(Error::Api(response.text().await.unwrap_or_default()));
}
response.json::<R>().await.map_err(Error::Json)
})
.retry(ExponentialBuilder::default())
.when(|e| match e {
Error::Network(_) => true,
Error::Api(msg) => {
msg.contains("503")
|| msg.contains("502")
|| msg.contains("500")
|| msg.contains("504")
}
_ => false,
})
.await
}
}
#[derive(Deserialize, Debug)]
pub struct ListModelsResponse {
pub models: Vec<Model>,
}
#[derive(Deserialize, Debug)]
#[serde(rename_all = "camelCase")]
pub struct Model {
pub name: String,
#[serde(default)]
pub version: String,
#[serde(default)]
pub display_name: String,
#[serde(default)]
pub description: String,
#[serde(default)]
pub input_token_limit: u32,
#[serde(default)]
pub output_token_limit: u32,
#[serde(default)]
pub supported_generation_methods: Vec<String>,
}