use std::fmt;
use std::sync::Arc;
use ferrin_provider_util::IdGenerator;
use ferrin_provider_util::PrefixedIdGenerator;
use ferrin_provider_util::SharedTransport;
use ferrin_provider_util::base_url::join_path;
use ferrin_provider_util::secure_url::UrlPolicy;
use ferrin_provider_util::settings::ApiKeyConfig;
use ferrin_provider_util::settings::load_api_key;
use ferrin_spec::Headers;
use ferrin_spec::ProviderId;
use ferrin_spec::error::InvalidArgumentError;
use ferrin_spec::error::ProviderError;
use secrecy::ExposeSecret;
use secrecy::SecretString;
use url::Url;
pub const USER_AGENT: &str = concat!("ferrin-google/", env!("CARGO_PKG_VERSION"));
pub const API_KEY_ENV: &str = "GOOGLE_GENERATIVE_AI_API_KEY";
pub const DEFAULT_BASE_URL: &str = "https://generativelanguage.googleapis.com/v1beta";
pub const API_KEY_HEADER: &str = "x-goog-api-key";
pub const CANONICAL_OPTIONS_KEY: &str = "google";
pub const DEFAULT_NAME: &str = "google";
pub const UPLOAD_PATH: &str = "/upload/v1beta/files";
pub const DOWNLOAD_PATH_PREFIX: &str = "/download/v1beta/";
pub const AUTH_TOKENS_PATH: &str = "/v1alpha/auth_tokens";
pub struct GoogleConfig {
pub name: String,
pub base_url: Url,
pub api_key: Option<SecretString>,
pub headers: Headers,
pub url_policy: UrlPolicy,
pub transport: SharedTransport,
pub id_generator: Arc<dyn IdGenerator>,
}
impl fmt::Debug for GoogleConfig {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("GoogleConfig")
.field("name", &self.name)
.field("base_url", &self.base_url)
.field("api_key", &self.api_key.as_ref().map(|_| "***"))
.field("headers", &self.headers)
.finish_non_exhaustive()
}
}
impl GoogleConfig {
pub fn new(name: impl Into<String>, base_url: Url) -> Result<Self, ProviderError> {
let transport = ferrin_provider_util::default_transport().map_err(ProviderError::other)?;
Ok(Self::with_transport(name, base_url, transport))
}
#[must_use]
pub fn with_transport(
name: impl Into<String>,
base_url: Url,
transport: SharedTransport,
) -> Self {
Self {
name: name.into(),
base_url,
api_key: None,
headers: Headers::new(),
url_policy: UrlPolicy::default(),
transport,
id_generator: Arc::new(PrefixedIdGenerator::default()),
}
}
#[must_use]
pub fn provider_id(&self, family: &str) -> ProviderId {
ProviderId::new(format!("{}.{family}", self.name))
}
#[must_use]
pub fn options_key(&self) -> &str {
&self.name
}
#[must_use]
pub fn url(&self, path: &str) -> Url {
join_path(&self.base_url, path)
}
#[must_use]
pub fn model_path(model_id: &str) -> String {
if model_id.contains('/') {
model_id.to_owned()
} else {
format!("models/{model_id}")
}
}
#[must_use]
pub fn model_url(&self, model_id: &str, action: &str) -> Url {
self.url(&format!("{}:{action}", Self::model_path(model_id)))
}
#[must_use]
pub fn origin_url(&self, path: &str) -> Url {
let mut url = self.base_url.clone();
url.set_path(path);
url.set_query(None);
url.set_fragment(None);
url
}
#[must_use]
pub fn websocket_url(&self, service_path: &str) -> Url {
let mut url = self.base_url.clone();
let mut segments: Vec<&str> = url
.path()
.split('/')
.filter(|segment| !segment.is_empty())
.collect();
if matches!(segments.last(), Some(&"v1beta" | &"v1alpha")) {
segments.pop();
}
let mut path = segments.join("/");
if !path.is_empty() {
path.insert(0, '/');
}
path.push_str("/ws/");
path.push_str(service_path);
url.set_path(&path);
url.set_query(None);
url.set_fragment(None);
let scheme = if url.scheme() == "http" { "ws" } else { "wss" };
let _ = url.set_scheme(scheme);
url
}
pub fn api_key(&self) -> Result<SecretString, ProviderError> {
Ok(load_api_key(ApiKeyConfig {
api_key: self.api_key.clone(),
environment_variable: API_KEY_ENV,
parameter_name: "api_key",
description: "Google Generative AI",
})?)
}
pub fn headers(&self, call_headers: &Headers) -> Result<Headers, ProviderError> {
let mut headers = Headers::new();
let key = self.api_key()?;
headers
.insert(API_KEY_HEADER, key.expose_secret())
.map_err(|_| {
ProviderError::InvalidArgument(InvalidArgumentError::new(
"api_key",
"api_key is not a valid header value",
))
})?;
headers.merge(&self.headers);
headers.merge(call_headers);
Ok(headers.with_user_agent_suffix([USER_AGENT]))
}
#[must_use]
pub fn unauthenticated_headers(&self, call_headers: &Headers) -> Headers {
let mut headers = self.headers.clone();
headers.merge(call_headers);
headers.with_user_agent_suffix([USER_AGENT])
}
#[must_use]
pub fn generate_id(&self) -> String {
self.id_generator.generate()
}
}
pub type SharedConfig = Arc<GoogleConfig>;