use std::{sync::Arc, time::Duration};
use async_dropper_simple::AsyncDropper;
use recloser::AsyncRecloser;
use reqwest::header::{HeaderMap, HeaderValue};
use reqwest_middleware::ClientBuilder;
use reqwest_retry::RetryTransientMiddleware;
use std::sync::LazyLock;
use crate::agent::buffer::Buffer;
use crate::agent::usage_agent::{non_empty_string, AgentError, UsageAgent, UsageAgentInner};
use crate::agent::utils::OperationProcessor;
use crate::circuit_breaker;
use crate::expressions::CompileExpression;
use retry_policies::policies::ExponentialBackoff;
pub struct UsageAgentBuilder {
token: Option<String>,
endpoint: String,
target_id: Option<String>,
buffer_size: usize,
connect_timeout: Duration,
request_timeout: Duration,
accept_invalid_certs: bool,
flush_interval: Duration,
retry_policy: ExponentialBackoff,
user_agent: Option<String>,
circuit_breaker: Option<AsyncRecloser>,
exclude_expression: Option<String>,
}
pub static DEFAULT_HIVE_USAGE_ENDPOINT: &str = "https://app.graphql-hive.com/usage";
impl Default for UsageAgentBuilder {
fn default() -> Self {
Self {
endpoint: DEFAULT_HIVE_USAGE_ENDPOINT.to_string(),
token: None,
target_id: None,
buffer_size: 1000,
connect_timeout: Duration::from_secs(5),
request_timeout: Duration::from_secs(15),
accept_invalid_certs: false,
flush_interval: Duration::from_secs(5),
retry_policy: ExponentialBackoff::builder().build_with_max_retries(3),
user_agent: None,
circuit_breaker: None,
exclude_expression: None,
}
}
}
fn is_legacy_token(token: &str) -> bool {
!token.starts_with("hvo1/") && !token.starts_with("hvu1/") && !token.starts_with("hvp1/")
}
impl UsageAgentBuilder {
pub fn token(mut self, token: String) -> Self {
if let Some(token) = non_empty_string(Some(token)) {
self.token = Some(token);
}
self
}
pub fn endpoint(mut self, endpoint: String) -> Self {
if let Some(endpoint) = non_empty_string(Some(endpoint)) {
self.endpoint = endpoint;
}
self
}
pub fn target_id(mut self, target_id: String) -> Self {
if let Some(target_id) = non_empty_string(Some(target_id)) {
self.target_id = Some(target_id);
}
self
}
pub fn buffer_size(mut self, buffer_size: usize) -> Self {
self.buffer_size = buffer_size;
self
}
pub fn connect_timeout(mut self, connect_timeout: Duration) -> Self {
self.connect_timeout = connect_timeout;
self
}
pub fn request_timeout(mut self, request_timeout: Duration) -> Self {
self.request_timeout = request_timeout;
self
}
pub fn accept_invalid_certs(mut self, accept_invalid_certs: bool) -> Self {
self.accept_invalid_certs = accept_invalid_certs;
self
}
pub fn flush_interval(mut self, flush_interval: Duration) -> Self {
self.flush_interval = flush_interval;
self
}
pub fn user_agent(mut self, user_agent: String) -> Self {
if let Some(user_agent) = non_empty_string(Some(user_agent)) {
self.user_agent = Some(user_agent);
}
self
}
pub fn retry_policy(mut self, retry_policy: ExponentialBackoff) -> Self {
self.retry_policy = retry_policy;
self
}
pub fn max_retries(mut self, max_retries: u32) -> Self {
self.retry_policy = ExponentialBackoff::builder().build_with_max_retries(max_retries);
self
}
pub(crate) fn build_agent(self) -> Result<UsageAgentInner, AgentError> {
let mut default_headers = HeaderMap::new();
default_headers.insert("X-Usage-API-Version", HeaderValue::from_static("2"));
let token = match self.token {
Some(token) => token,
None => return Err(AgentError::MissingToken),
};
let mut authorization_header = HeaderValue::from_str(&format!("Bearer {}", token))
.map_err(|_| AgentError::InvalidToken)?;
authorization_header.set_sensitive(true);
default_headers.insert(reqwest::header::AUTHORIZATION, authorization_header);
default_headers.insert(
reqwest::header::CONTENT_TYPE,
HeaderValue::from_static("application/json"),
);
let mut reqwest_agent = reqwest::Client::builder()
.danger_accept_invalid_certs(self.accept_invalid_certs)
.connect_timeout(self.connect_timeout)
.timeout(self.request_timeout)
.default_headers(default_headers);
if let Some(user_agent) = &self.user_agent {
reqwest_agent = reqwest_agent.user_agent(user_agent);
}
let reqwest_agent = reqwest_agent
.build()
.map_err(AgentError::HTTPClientCreationError)?;
let client = ClientBuilder::new(reqwest_agent)
.with(RetryTransientMiddleware::new_with_policy(self.retry_policy))
.build();
let mut endpoint = self.endpoint;
match self.target_id {
Some(_) if is_legacy_token(&token) => return Err(AgentError::TargetIdWithLegacyToken),
Some(target_id) if !is_legacy_token(&token) => {
let target_id = validate_target_id(&target_id)?;
endpoint.push_str(&format!("/{}", target_id));
}
None if !is_legacy_token(&token) => return Err(AgentError::MissingTargetId),
_ => {}
}
let circuit_breaker = if let Some(cb) = self.circuit_breaker {
cb
} else {
circuit_breaker::CircuitBreakerBuilder::default()
.build_async()
.map_err(AgentError::CircuitBreakerCreationError)?
};
let buffer = Buffer::new(self.buffer_size);
let exclude = if let Some(expr) = self.exclude_expression {
expr.compile_expression(None)?.into()
} else {
None
};
Ok(UsageAgentInner {
endpoint,
buffer,
processor: OperationProcessor::new(),
client,
flush_interval: self.flush_interval,
circuit_breaker,
exclude_expression: exclude,
})
}
pub fn exclude_expression(mut self, expression: String) -> Self {
if let Some(expression) = non_empty_string(Some(expression)) {
self.exclude_expression = Some(expression);
}
self
}
pub fn build(self) -> Result<UsageAgent, AgentError> {
let agent = self.build_agent()?;
Ok(Arc::new(AsyncDropper::new(agent)))
}
}
static SLUG_REGEX: LazyLock<regex_automata::meta::Regex> = LazyLock::new(|| {
regex_automata::meta::Regex::new(r"^[a-zA-Z0-9-_]+\/[a-zA-Z0-9-_]+\/[a-zA-Z0-9-_]+$").unwrap()
});
static UUID_REGEX: LazyLock<regex_automata::meta::Regex> = LazyLock::new(|| {
regex_automata::meta::Regex::new(
r"^[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}$",
)
.unwrap()
});
fn validate_target_id(target_id: &str) -> Result<&str, AgentError> {
let trimmed_s = target_id.trim();
if trimmed_s.is_empty() {
Err(AgentError::InvalidTargetId("<empty>".to_string()))
} else {
if SLUG_REGEX.is_match(trimmed_s) {
return Ok(trimmed_s);
}
if UUID_REGEX.is_match(trimmed_s) {
return Ok(trimmed_s);
}
Err(AgentError::InvalidTargetId(format!(
"Invalid target_id format: '{}'. It must be either in slug format '$organizationSlug/$projectSlug/$targetSlug' or UUID format 'a0f4c605-6541-4350-8cfe-b31f21a4bf80'",
trimmed_s
)))
}
}