use async_trait::async_trait;
use reqwest::Client;
use std::time::Duration;
use tokio::sync::Mutex;
use tokio::time::Instant;
use url::Url;
use uuid::Uuid;
use crate::{
transport::{Transport, TransportError},
EventData,
};
#[derive(Debug, Clone)]
pub struct BatchConfig {
pub max_batch_size: usize,
pub max_batch_age: Duration,
}
impl Default for BatchConfig {
fn default() -> Self {
Self {
max_batch_size: 500,
max_batch_age: Duration::from_secs(1),
}
}
}
impl BatchConfig {
pub fn new(max_batch_size: usize, max_batch_age: Duration) -> Self {
Self {
max_batch_size,
max_batch_age,
}
}
}
const MAX_BATCH_BYTES: usize = 4 * 1024 * 1024;
pub struct BatchingHttpTransport {
client: Client,
batch_url: Url,
config: BatchConfig,
auth_token: Option<String>,
buffer: Mutex<BatchBuffer>,
}
struct BatchBuffer {
events: Vec<EventData>,
oldest_event_time: Option<Instant>,
serialized_bytes: usize,
}
impl BatchBuffer {
fn new() -> Self {
Self {
events: Vec::new(),
oldest_event_time: None,
serialized_bytes: 2,
}
}
fn push(&mut self, event: EventData) {
if self.events.is_empty() {
self.oldest_event_time = Some(Instant::now());
}
self.serialized_bytes += Self::event_bytes(&event);
self.events.push(event);
}
fn take(&mut self) -> Vec<EventData> {
self.oldest_event_time = None;
self.serialized_bytes = 2;
std::mem::take(&mut self.events)
}
fn event_bytes(event: &EventData) -> usize {
serde_json::to_vec(event)
.expect("event payload is JSON")
.len()
+ 1
}
fn is_empty(&self) -> bool {
self.events.is_empty()
}
fn len(&self) -> usize {
self.events.len()
}
fn age(&self) -> Option<Duration> {
self.oldest_event_time.map(|t| t.elapsed())
}
}
impl BatchingHttpTransport {
pub fn new(
base_url: Url,
org_id: Uuid,
app_id: Uuid,
config: BatchConfig,
auth_token: Option<String>,
) -> Result<Self, TransportError> {
let batch_url = base_url
.join(&format!(
"/api/orgs/{}/apps/{}/events/batch",
org_id, app_id
))
.map_err(|e| TransportError::Configuration(format!("Invalid URL: {}", e)))?;
Ok(Self {
client: Client::builder()
.timeout(Duration::from_secs(10))
.build()
.map_err(|error| TransportError::Configuration(error.to_string()))?,
batch_url,
config,
auth_token,
buffer: Mutex::new(BatchBuffer::new()),
})
}
pub fn with_default_config(
base_url: Url,
org_id: Uuid,
app_id: Uuid,
auth_token: Option<String>,
) -> Result<Self, TransportError> {
Self::new(base_url, org_id, app_id, BatchConfig::default(), auth_token)
}
async fn flush_buffer(&self) -> Result<(), TransportError> {
let mut buffer = self.buffer.lock().await;
if buffer.is_empty() {
return Ok(());
}
self.send_batch(&buffer.events).await?;
buffer.take();
Ok(())
}
async fn send_batch(&self, events: &[EventData]) -> Result<(), TransportError> {
if events.is_empty() {
return Ok(());
}
let mut request = self.client.post(self.batch_url.clone()).json(&events);
if let Some(token) = &self.auth_token {
request = request.bearer_auth(token);
}
let response = request
.send()
.await
.map_err(|e| TransportError::Send(format!("HTTP batch request failed: {}", e)))?;
if !response.status().is_success() {
return Err(TransportError::Send(format!(
"Server returned error status for batch: {}",
response.status()
)));
}
Ok(())
}
async fn should_flush(&self) -> bool {
let buffer = self.buffer.lock().await;
if buffer.len() >= self.config.max_batch_size {
return true;
}
if let Some(age) = buffer.age() {
if age >= self.config.max_batch_age {
return true;
}
}
false
}
}
#[async_trait]
impl Transport for BatchingHttpTransport {
async fn connect(&mut self) -> Result<(), TransportError> {
Ok(())
}
async fn send(&mut self, event: EventData) -> Result<(), TransportError> {
let exceeds_bytes = {
let buffer = self.buffer.lock().await;
!buffer.is_empty()
&& buffer
.serialized_bytes
.saturating_add(BatchBuffer::event_bytes(&event))
> MAX_BATCH_BYTES
};
if self.should_flush().await || exceeds_bytes {
self.flush_buffer().await?;
}
self.buffer.lock().await.push(event);
Ok(())
}
async fn flush(&mut self) -> Result<(), TransportError> {
self.flush_buffer().await
}
fn flush_interval(&self) -> Option<Duration> {
Some(self.config.max_batch_age)
}
async fn close(&mut self) -> Result<(), TransportError> {
self.flush_buffer().await
}
}
#[cfg(test)]
mod tests {
use super::*;
use chrono::Utc;
#[test]
fn test_batch_config_default() {
let config = BatchConfig::default();
assert_eq!(config.max_batch_size, 500);
assert_eq!(config.max_batch_age, Duration::from_secs(1));
}
#[test]
fn test_batch_config_custom() {
let config = BatchConfig::new(50, Duration::from_millis(500));
assert_eq!(config.max_batch_size, 50);
assert_eq!(config.max_batch_age, Duration::from_millis(500));
}
#[test]
fn test_batch_buffer_operations() {
let mut buffer = BatchBuffer::new();
assert!(buffer.is_empty());
assert_eq!(buffer.len(), 0);
assert!(buffer.age().is_none());
let event = EventData {
event_type: "test".to_string(),
event_data: serde_json::json!({}),
event_timestamp: Utc::now(),
process_instance_id: None,
};
buffer.push(event);
assert!(!buffer.is_empty());
assert_eq!(buffer.len(), 1);
assert!(buffer.age().is_some());
let events = buffer.take();
assert_eq!(events.len(), 1);
assert!(buffer.is_empty());
assert!(buffer.age().is_none());
}
#[test]
fn test_batching_transport_creation() {
let base_url = Url::parse("http://localhost:4318").unwrap();
let org_id = Uuid::new_v4();
let app_id = Uuid::new_v4();
let transport =
BatchingHttpTransport::with_default_config(base_url.clone(), org_id, app_id, None);
assert!(transport.is_ok());
let custom_config = BatchConfig::new(50, Duration::from_millis(100));
let transport = BatchingHttpTransport::new(base_url, org_id, app_id, custom_config, None);
assert!(transport.is_ok());
}
#[tokio::test]
async fn test_batching_transport_url() {
let base_url = Url::parse("http://localhost:4318").unwrap();
let org_id = Uuid::parse_str("12345678-1234-1234-1234-123456789012").unwrap();
let app_id = Uuid::parse_str("87654321-4321-4321-4321-210987654321").unwrap();
let transport = BatchingHttpTransport::with_default_config(base_url, org_id, app_id, None)
.expect("should create transport");
assert_eq!(
transport.batch_url.as_str(),
"http://localhost:4318/api/orgs/12345678-1234-1234-1234-123456789012/apps/87654321-4321-4321-4321-210987654321/events/batch"
);
}
#[tokio::test]
async fn test_should_flush_by_size() {
let base_url = Url::parse("http://localhost:4318").unwrap();
let org_id = Uuid::new_v4();
let app_id = Uuid::new_v4();
let config = BatchConfig::new(2, Duration::from_secs(60));
let transport = BatchingHttpTransport::new(base_url, org_id, app_id, config, None)
.expect("should create transport");
{
let mut buffer = transport.buffer.lock().await;
buffer.push(EventData {
event_type: "test".to_string(),
event_data: serde_json::json!({}),
event_timestamp: Utc::now(),
process_instance_id: None,
});
}
assert!(!transport.should_flush().await);
{
let mut buffer = transport.buffer.lock().await;
buffer.push(EventData {
event_type: "test".to_string(),
event_data: serde_json::json!({}),
event_timestamp: Utc::now(),
process_instance_id: None,
});
}
assert!(transport.should_flush().await);
}
#[tokio::test]
async fn test_should_flush_by_age() {
let base_url = Url::parse("http://localhost:4318").unwrap();
let org_id = Uuid::new_v4();
let app_id = Uuid::new_v4();
let config = BatchConfig::new(1000, Duration::from_millis(10));
let transport = BatchingHttpTransport::new(base_url, org_id, app_id, config, None)
.expect("should create transport");
{
let mut buffer = transport.buffer.lock().await;
buffer.push(EventData {
event_type: "test".to_string(),
event_data: serde_json::json!({}),
event_timestamp: Utc::now(),
process_instance_id: None,
});
}
assert!(!transport.should_flush().await);
tokio::time::sleep(Duration::from_millis(15)).await;
assert!(transport.should_flush().await);
}
fn test_event(name: &str) -> EventData {
EventData {
event_type: name.into(),
event_data: serde_json::json!({}),
event_timestamp: Utc::now(),
process_instance_id: None,
}
}
fn server(
replies: Vec<(u16, Duration)>,
) -> (Url, tokio::sync::oneshot::Receiver<Vec<Vec<EventData>>>) {
use std::io::{BufRead, BufReader, Read, Write};
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
listener.set_nonblocking(true).unwrap();
let address = listener.local_addr().unwrap();
let (sent, received) = tokio::sync::oneshot::channel();
std::thread::spawn(move || {
let mut requests = Vec::new();
for (status, delay) in replies {
let deadline = std::time::Instant::now() + Duration::from_secs(3);
let mut stream = loop {
match listener.accept() {
Ok((stream, _)) => break stream,
Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => {
assert!(std::time::Instant::now() < deadline, "batch never arrived");
std::thread::sleep(Duration::from_millis(1));
}
Err(error) => panic!("{error}"),
}
};
stream
.set_read_timeout(Some(Duration::from_secs(2)))
.unwrap();
let mut reader = BufReader::new(&mut stream);
let mut length = 0;
loop {
let mut line = String::new();
assert!(reader.read_line(&mut line).unwrap() > 0);
if line == "\r\n" {
break;
}
if let Some(value) = line.to_ascii_lowercase().strip_prefix("content-length:") {
length = value.trim().parse::<usize>().unwrap();
}
}
let mut body = vec![0; length];
reader.read_exact(&mut body).unwrap();
requests.push(serde_json::from_slice(&body).unwrap());
std::thread::sleep(delay);
let _ = write!(stream, "HTTP/1.1 {status} Result\r\nContent-Length: 2\r\nConnection: close\r\n\r\n{{}}");
}
let _ = sent.send(requests);
});
(Url::parse(&format!("http://{address}")).unwrap(), received)
}
#[tokio::test]
async fn failed_batch_is_retained_and_retry_does_not_duplicate_the_next_event() {
let (url, received) = server(vec![
(500, Duration::ZERO),
(201, Duration::ZERO),
(201, Duration::ZERO),
]);
let mut transport = BatchingHttpTransport::new(
url,
Uuid::new_v4(),
Uuid::new_v4(),
BatchConfig::new(2, Duration::from_secs(60)),
None,
)
.unwrap();
transport.send(test_event("first")).await.unwrap();
transport.send(test_event("second")).await.unwrap();
assert!(transport.send(test_event("third")).await.is_err());
assert_eq!(transport.buffer.lock().await.len(), 2);
transport.send(test_event("third")).await.unwrap();
transport.close().await.unwrap();
let requests = received.await.unwrap();
let names: Vec<Vec<_>> = requests
.iter()
.map(|batch| {
batch
.iter()
.map(|event| event.event_type.as_str())
.collect()
})
.collect();
assert_eq!(
names,
vec![
vec!["first", "second"],
vec!["first", "second"],
vec!["third"]
]
);
}
#[tokio::test]
async fn cancelling_a_flush_keeps_every_accepted_event_for_retry() {
let (url, received) = server(vec![
(201, Duration::from_millis(100)),
(201, Duration::ZERO),
]);
let mut transport =
BatchingHttpTransport::with_default_config(url, Uuid::new_v4(), Uuid::new_v4(), None)
.unwrap();
transport.send(test_event("retained")).await.unwrap();
assert!(
tokio::time::timeout(Duration::from_millis(30), transport.flush_buffer())
.await
.is_err()
);
assert_eq!(transport.buffer.lock().await.len(), 1);
transport.close().await.unwrap();
let requests = received.await.unwrap();
assert_eq!(requests.len(), 2);
assert!(requests
.iter()
.all(|batch| batch.len() == 1 && batch[0].event_type == "retained"));
}
#[tokio::test]
async fn idle_batch_flushes_without_another_event_or_shutdown() {
use std::sync::{atomic::AtomicU64, Arc};
let (url, received) = server(vec![(201, Duration::ZERO)]);
let transport = BatchingHttpTransport::new(
url,
Uuid::new_v4(),
Uuid::new_v4(),
BatchConfig::new(100, Duration::from_millis(25)),
None,
)
.unwrap();
let (sender, receiver) = tokio::sync::mpsc::channel(10);
let (shutdown, shutdown_rx) = tokio::sync::oneshot::channel();
let (completed, completion) = tokio::sync::oneshot::channel();
tokio::spawn(crate::transport::run_transport_loop(
Box::new(transport),
receiver,
shutdown_rx,
completed,
Arc::new(AtomicU64::new(0)),
crate::transport::TransportLoopConfig::default(),
));
sender.send(test_event("quiet tail")).await.unwrap();
let requests = tokio::time::timeout(Duration::from_secs(1), received)
.await
.unwrap()
.unwrap();
assert_eq!(requests[0][0].event_type, "quiet tail");
shutdown.send(()).unwrap();
completion.await.unwrap();
}
#[tokio::test]
async fn large_events_split_batches_before_the_http_body_limit() {
let (url, received) = server(vec![(201, Duration::ZERO), (201, Duration::ZERO)]);
let mut transport =
BatchingHttpTransport::with_default_config(url, Uuid::new_v4(), Uuid::new_v4(), None)
.unwrap();
let mut event = test_event("large");
event.event_data = serde_json::json!({"message":"x".repeat(3 * 1024 * 1024)});
transport.send(event.clone()).await.unwrap();
transport.send(event).await.unwrap();
transport.close().await.unwrap();
let requests = received.await.unwrap();
assert_eq!(requests.len(), 2);
assert!(requests.iter().all(|batch| batch.len() == 1));
}
}