use std::collections::HashMap;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, RwLock};
use std::time::Duration;
use prost::Message;
use tokio::sync::mpsc;
use tokio::time::interval;
use crate::error::PolicyError;
use crate::policy::Policy;
use crate::proto::tero::policy::v1::{ClientMetadata, SyncRequest, SyncResponse};
use super::sync::{collect_policy_statuses, collect_volume};
use super::{PolicyCallback, PolicyProvider, StatsCollector};
use crate::volume::VolumeTracker;
#[derive(Debug, Clone)]
pub struct HttpProviderConfig {
pub url: String,
pub headers: HashMap<String, String>,
pub poll_interval_ns: u64,
pub client_metadata: Option<ClientMetadata>,
pub content_type: ContentType,
}
#[derive(Debug, Clone, Copy, Default)]
pub enum ContentType {
#[default]
Protobuf,
Json,
}
impl HttpProviderConfig {
pub fn new(url: impl Into<String>) -> Self {
Self {
url: url.into(),
headers: HashMap::new(),
poll_interval_ns: Duration::from_secs(60).as_nanos() as u64,
client_metadata: None,
content_type: ContentType::default(),
}
}
pub fn header(mut self, key: impl Into<String>, value: impl Into<String>) -> Self {
self.headers.insert(key.into(), value.into());
self
}
pub fn headers(mut self, headers: HashMap<String, String>) -> Self {
self.headers.extend(headers);
self
}
pub fn poll_interval(mut self, interval: Duration) -> Self {
self.poll_interval_ns = interval.as_nanos() as u64;
self
}
pub fn poll_interval_ns(mut self, ns: u64) -> Self {
self.poll_interval_ns = ns;
self
}
pub fn client_metadata(mut self, metadata: ClientMetadata) -> Self {
self.client_metadata = Some(metadata);
self
}
pub fn content_type(mut self, content_type: ContentType) -> Self {
self.content_type = content_type;
self
}
}
pub struct HttpProvider {
config: HttpProviderConfig,
client: reqwest::Client,
last_hash: RwLock<Option<String>>,
last_sync_timestamp: RwLock<u64>,
running: AtomicBool,
stats_collector: RwLock<Option<StatsCollector>>,
volume_tracker: RwLock<Option<Arc<VolumeTracker>>>,
initial_policies: RwLock<Option<Vec<Policy>>>,
}
impl HttpProvider {
pub fn new(config: HttpProviderConfig) -> Self {
let client = reqwest::Client::new();
Self {
config,
client,
last_hash: RwLock::new(None),
last_sync_timestamp: RwLock::new(0),
running: AtomicBool::new(false),
stats_collector: RwLock::new(None),
volume_tracker: RwLock::new(None),
initial_policies: RwLock::new(None),
}
}
pub async fn new_with_initial_fetch(config: HttpProviderConfig) -> Result<Self, PolicyError> {
let provider = Self::new(config);
let policies = provider.sync(true).await?;
*provider.initial_policies.write().unwrap() = Some(policies);
Ok(provider)
}
pub async fn load(&self) -> Result<Vec<Policy>, PolicyError> {
self.sync(true).await
}
fn build_sync_request(&self, full_sync: bool) -> SyncRequest {
let last_hash = self.last_hash.read().unwrap().clone().unwrap_or_default();
let last_timestamp = *self.last_sync_timestamp.read().unwrap();
let policy_statuses = collect_policy_statuses(&self.stats_collector.read().unwrap());
let volume = collect_volume(&self.volume_tracker.read().unwrap());
SyncRequest {
client_metadata: self.config.client_metadata.clone(),
full_sync,
last_sync_timestamp_unix_nano: last_timestamp,
last_successful_hash: last_hash,
policy_statuses,
volume,
}
}
async fn sync(&self, full_sync: bool) -> Result<Vec<Policy>, PolicyError> {
let request = self.build_sync_request(full_sync);
let mut http_request = self.client.post(&self.config.url);
for (key, value) in &self.config.headers {
http_request = http_request.header(key, value);
}
let response = match self.config.content_type {
ContentType::Protobuf => {
let body = request.encode_to_vec();
http_request
.header("Content-Type", "application/x-protobuf")
.header("Accept", "application/x-protobuf")
.body(body)
.send()
.await
.map_err(|e| PolicyError::HttpError(e.to_string()))?
}
ContentType::Json => {
http_request
.header("Content-Type", "application/json")
.header("Accept", "application/json")
.json(&request)
.send()
.await
.map_err(|e| PolicyError::HttpError(e.to_string()))?
}
};
if !response.status().is_success() {
return Err(PolicyError::HttpError(format!(
"HTTP error: {} - {}",
response.status(),
response
.text()
.await
.unwrap_or_else(|_| "unknown".to_string())
)));
}
let sync_response: SyncResponse = match self.config.content_type {
ContentType::Protobuf => {
let bytes = response
.bytes()
.await
.map_err(|e| PolicyError::HttpError(e.to_string()))?;
SyncResponse::decode(bytes).map_err(|e| PolicyError::HttpError(e.to_string()))?
}
ContentType::Json => {
let text = response
.text()
.await
.map_err(|e| PolicyError::HttpError(e.to_string()))?;
serde_json::from_str(&text).map_err(|e| {
PolicyError::HttpError(format!(
"JSON decode error: {} - response: {}",
e,
&text[..text.len().min(500)]
))
})?
}
};
if !sync_response.error_message.is_empty() {
return Err(PolicyError::HttpError(format!(
"Sync error: {}",
sync_response.error_message
)));
}
if !sync_response.hash.is_empty() {
*self.last_hash.write().unwrap() = Some(sync_response.hash);
}
if sync_response.sync_timestamp_unix_nano > 0 {
*self.last_sync_timestamp.write().unwrap() = sync_response.sync_timestamp_unix_nano;
}
let policies = sync_response
.policies
.into_iter()
.map(Policy::new)
.collect();
Ok(policies)
}
pub fn start_polling(
&self,
) -> mpsc::Receiver<Result<(Option<String>, Vec<Policy>), PolicyError>> {
let (tx, rx) = mpsc::channel(16);
self.running.store(true, Ordering::SeqCst);
let config = self.config.clone();
let client = self.client.clone();
let last_hash = Arc::new(RwLock::new(None::<String>));
let last_sync_timestamp = Arc::new(RwLock::new(0u64));
let stats_collector = self.stats_collector.read().unwrap().clone();
let volume_tracker = self.volume_tracker.read().unwrap().clone();
let running = Arc::new(AtomicBool::new(true));
let running_clone = running.clone();
let last_hash_clone = last_hash.clone();
let last_sync_timestamp_clone = last_sync_timestamp.clone();
tokio::spawn(async move {
let poll_duration = Duration::from_nanos(config.poll_interval_ns);
let mut interval = interval(poll_duration);
let mut first = true;
while running_clone.load(Ordering::SeqCst) {
interval.tick().await;
let request = {
let last_hash = last_hash_clone.read().unwrap().clone().unwrap_or_default();
let last_timestamp = *last_sync_timestamp_clone.read().unwrap();
let policy_statuses = collect_policy_statuses(&stats_collector);
let volume = collect_volume(&volume_tracker);
SyncRequest {
client_metadata: config.client_metadata.clone(),
full_sync: first,
last_sync_timestamp_unix_nano: last_timestamp,
last_successful_hash: last_hash,
policy_statuses,
volume,
}
};
first = false;
let mut http_request = client.post(&config.url);
for (key, value) in &config.headers {
http_request = http_request.header(key, value);
}
let result = async {
let response = match config.content_type {
ContentType::Protobuf => {
let body = request.encode_to_vec();
http_request
.header("Content-Type", "application/x-protobuf")
.header("Accept", "application/x-protobuf")
.body(body)
.send()
.await
.map_err(|e| PolicyError::HttpError(e.to_string()))?
}
ContentType::Json => http_request
.header("Content-Type", "application/json")
.header("Accept", "application/json")
.json(&request)
.send()
.await
.map_err(|e| PolicyError::HttpError(e.to_string()))?,
};
if !response.status().is_success() {
return Err(PolicyError::HttpError(format!(
"HTTP error: {} - {}",
response.status(),
response
.text()
.await
.unwrap_or_else(|_| "unknown".to_string())
)));
}
let sync_response: SyncResponse = match config.content_type {
ContentType::Protobuf => {
let bytes = response
.bytes()
.await
.map_err(|e| PolicyError::HttpError(e.to_string()))?;
SyncResponse::decode(bytes)
.map_err(|e| PolicyError::HttpError(e.to_string()))?
}
ContentType::Json => response
.json()
.await
.map_err(|e| PolicyError::HttpError(e.to_string()))?,
};
if !sync_response.error_message.is_empty() {
return Err(PolicyError::HttpError(format!(
"Sync error: {}",
sync_response.error_message
)));
}
let new_hash = if !sync_response.hash.is_empty() {
let hash = Some(sync_response.hash);
*last_hash_clone.write().unwrap() = hash.clone();
hash
} else {
None
};
if sync_response.sync_timestamp_unix_nano > 0 {
*last_sync_timestamp_clone.write().unwrap() =
sync_response.sync_timestamp_unix_nano;
}
let policies: Vec<Policy> = sync_response
.policies
.into_iter()
.map(Policy::new)
.collect();
Ok((new_hash, policies))
}
.await;
if tx.send(result).await.is_err() {
break; }
}
});
rx
}
pub fn stop(&self) {
self.running.store(false, Ordering::SeqCst);
}
}
impl PolicyProvider for HttpProvider {
fn set_stats_collector(&self, collector: StatsCollector) {
*self.stats_collector.write().unwrap() = Some(collector);
}
fn set_volume_tracker(&self, tracker: Arc<VolumeTracker>) {
*self.volume_tracker.write().unwrap() = Some(tracker);
}
fn subscribe(&self, callback: PolicyCallback) -> Result<(), PolicyError> {
let policies = self
.initial_policies
.write()
.unwrap()
.take()
.expect("HttpProvider::subscribe() requires new_with_initial_fetch()");
callback(policies);
let initial_hash = self.last_hash.read().unwrap().clone();
let mut rx = self.start_polling();
let callback = callback.clone();
tokio::spawn(async move {
let mut last_known_hash = initial_hash;
while let Some(result) = rx.recv().await {
match result {
Ok((new_hash, policies)) => {
if new_hash != last_known_hash {
last_known_hash = new_hash;
callback(policies);
}
}
Err(e) => {
eprintln!("HTTP provider sync error: {}", e);
}
}
}
});
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
fn unreachable_provider() -> (HttpProvider, Arc<VolumeTracker>) {
let provider = HttpProvider::new(HttpProviderConfig::new("http://127.0.0.1:1"));
let tracker = Arc::new(VolumeTracker::new());
provider.set_volume_tracker(Arc::clone(&tracker));
(provider, tracker)
}
#[test]
fn sync_request_omits_untracked_volume() {
let (provider, _tracker) = unreachable_provider();
assert!(provider.build_sync_request(true).volume.is_none());
}
#[test]
fn sync_request_drains_observed_volume() {
let (provider, tracker) = unreachable_provider();
tracker.record_log();
tracker.add_log_bytes(400);
let volume = provider.build_sync_request(true).volume.unwrap();
assert_eq!(volume.log_records, 1);
assert_eq!(volume.log_bytes, 400);
assert!(provider.build_sync_request(true).volume.is_none());
}
#[tokio::test]
async fn failed_sync_does_not_replay_volume() {
let (provider, tracker) = unreachable_provider();
tracker.record_log();
tracker.add_log_bytes(400);
assert!(provider.sync(true).await.is_err());
assert!(tracker.collect().is_none());
}
}