use super::entry::StreamEntry;
use super::runner::Runner;
use super::transport::Transport;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
#[derive(Clone, Debug)]
pub(crate) struct StreamPoolOptions {
pub(crate) max_streams: usize,
pub(crate) max_outstanding_requests: Option<u64>,
pub(crate) max_outstanding_bytes: Option<u64>,
pub(crate) load_threshold: f64,
}
impl Default for StreamPoolOptions {
fn default() -> Self {
Self {
max_streams: 8,
max_outstanding_requests: Some(1000),
max_outstanding_bytes: None,
load_threshold: 0.2,
}
}
}
#[derive(Debug)]
pub(crate) struct StreamPool {
inner: Arc<Transport>,
next_stream_id: AtomicU64,
streams: Mutex<Vec<StreamEntry>>,
options: StreamPoolOptions,
}
impl StreamPool {
pub(crate) fn new(inner: Arc<Transport>, options: StreamPoolOptions) -> Self {
Self {
inner,
next_stream_id: AtomicU64::new(1),
streams: Mutex::new(Vec::new()),
options,
}
}
pub(crate) fn get(&self) -> StreamEntry {
let mut streams = self.streams.lock().unwrap();
self.get_impl(&mut streams)
}
pub(crate) fn evict_and_replace(&self, failed_id: u64) -> StreamEntry {
let mut streams = self.streams.lock().unwrap();
if let Some(pos) = streams.iter().position(|entry| entry.id == failed_id) {
let stream = self.new_stream_entry();
streams[pos] = stream.clone();
return stream;
}
self.get_impl(&mut streams)
}
fn get_impl(&self, streams: &mut Vec<StreamEntry>) -> StreamEntry {
let least_loaded = streams.iter().min_by(|a, b| {
let load_a = self.normalize_load(a);
let load_b = self.normalize_load(b);
load_a.total_cmp(&load_b)
});
let should_grow = least_loaded.is_none_or(|s| self.is_loaded(s));
if streams.len() < self.options.max_streams && should_grow {
let stream = self.new_stream_entry();
streams.push(stream.clone());
return stream;
}
match least_loaded {
Some(s) => s.clone(),
None => unreachable!("this can only happen when `max_streams == 0`"),
}
}
fn new_stream_entry(&self) -> StreamEntry {
let id = self.next_stream_id.fetch_add(1, Ordering::Relaxed);
let runner = Runner::new(self.inner.clone());
StreamEntry {
id,
req_tx: runner.req_tx,
outstanding_requests: Arc::new(AtomicU64::new(0)),
outstanding_bytes: Arc::new(AtomicU64::new(0)),
}
}
fn normalize_load(&self, entry: &StreamEntry) -> f64 {
let r = self
.options
.max_outstanding_requests
.map(|m| entry.outstanding_requests.load(Ordering::Relaxed) as f64 / m as f64)
.unwrap_or_default();
let b = self
.options
.max_outstanding_bytes
.map(|m| entry.outstanding_bytes.load(Ordering::Relaxed) as f64 / m as f64)
.unwrap_or_default();
f64::max(r, b)
}
fn is_loaded(&self, entry: &StreamEntry) -> bool {
self.normalize_load(entry) > self.options.load_threshold
}
}
#[cfg(test)]
mod tests {
use super::super::runner::WriteRequest;
use super::*;
use crate::write::test::*;
use bigquery_grpc_mock::{MockBigQueryWrite, start};
use gaxi::grpc::tonic::Response as TonicResponse;
use std::sync::MutexGuard;
use test_case::test_case;
use tokio::sync::{mpsc, oneshot};
use tokio::task::JoinSet;
#[test_case(10, Some(100), 10_000, Some(100_000), 0.1, false)]
#[test_case(90, Some(100), 10_000, Some(100_000), 0.9, true)]
#[test_case(10, Some(100), 90_000, Some(100_000), 0.9, true)]
#[test_case(90, Some(100), 90_000, Some(100_000), 0.9, true)]
#[test_case(10, None, 10_000, Some(100_000), 0.1, false)]
#[test_case(10, Some(100), 10_000, None, 0.1, false)]
#[test_case(10, None, 10_000, None, 0.0, false)]
#[tokio::test]
async fn load_math(
requests: u64,
max_outstanding_requests: Option<u64>,
bytes: u64,
max_outstanding_bytes: Option<u64>,
expected_load: f64,
expected_is_loaded: bool,
) -> anyhow::Result<()> {
let transport = Arc::new(test_transport("ignored").await?);
let options = StreamPoolOptions {
max_streams: 10,
max_outstanding_requests,
max_outstanding_bytes,
load_threshold: 0.2,
};
let pool = StreamPool::new(transport, options);
let s = pool.new_stream_entry();
s.outstanding_requests.store(requests, Ordering::Relaxed);
s.outstanding_bytes.store(bytes, Ordering::Relaxed);
assert_eq!(pool.normalize_load(&s), expected_load);
assert_eq!(pool.is_loaded(&s), expected_is_loaded);
Ok(())
}
#[tokio::test]
async fn empty_pool_get_basic() -> anyhow::Result<()> {
let (response_tx, response_rx) = mpsc::channel(10);
let mut mock = MockBigQueryWrite::new();
mock.expect_append_rows()
.return_once(|_| Ok(TonicResponse::from(response_rx)));
let (endpoint, _server) = start("0.0.0.0:0", mock).await?;
let transport = Arc::new(test_transport(endpoint).await?);
let pool = StreamPool::new(transport, StreamPoolOptions::default());
let s1 = pool.get();
assert_eq!(s1.id, 1);
assert_eq!(pool.stream_ids(), [1]);
let s2 = pool.get();
assert_eq!(s2.id, 1);
assert_eq!(pool.stream_ids(), [1]);
let (resp_tx1, resp_rx1) = oneshot::channel();
let write1 = WriteRequest {
req: test_request(1),
resp_tx: resp_tx1,
};
s1.req_tx.send(write1)?;
let (resp_tx2, resp_rx2) = oneshot::channel();
let write2 = WriteRequest {
req: test_request(2),
resp_tx: resp_tx2,
};
s2.req_tx.send(write2)?;
response_tx.send(Ok(convert(&test_response(1)))).await?;
let resp1 = resp_rx1.await??;
assert_eq!(resp1, test_response(1));
response_tx.send(Ok(convert(&test_response(2)))).await?;
let resp2 = resp_rx2.await??;
assert_eq!(resp2, test_response(2));
Ok(())
}
#[tokio::test]
async fn empty_pool_get_lock_contention() -> anyhow::Result<()> {
let transport = Arc::new(test_transport("ignored").await?);
let options = StreamPoolOptions {
max_streams: 10,
max_outstanding_requests: None,
max_outstanding_bytes: None,
load_threshold: 0.2,
};
let pool = Arc::new(StreamPool::new(transport, options));
let mut streams = JoinSet::new();
for _ in 0..1000 {
let p = pool.clone();
streams.spawn(async move { p.get() });
}
while let Some(s) = streams.join_next().await {
assert_eq!(s?.id, 1);
}
assert_eq!(pool.stream_ids(), [1]);
Ok(())
}
#[tokio::test]
async fn get_least_loaded() -> anyhow::Result<()> {
let transport = Arc::new(test_transport("ignored").await?);
let pool = StreamPool::new(transport, StreamPoolOptions::default());
pool.seed([8, 2, 2, 3, 1, 9]);
let s = pool.get();
assert_eq!(s.id, 5);
Ok(())
}
#[tokio::test]
async fn get_should_grow() -> anyhow::Result<()> {
let transport = Arc::new(test_transport("ignored").await?);
let options = StreamPoolOptions {
max_streams: 10,
max_outstanding_requests: Some(3),
max_outstanding_bytes: None,
load_threshold: 0.5,
};
let pool = Arc::new(StreamPool::new(transport, options));
let s = pool.get();
assert_eq!(s.id, 1);
assert_eq!(pool.stream_ids(), [1]);
s.outstanding_requests.fetch_add(1, Ordering::Relaxed);
let s = pool.get();
assert_eq!(s.id, 1);
assert_eq!(pool.stream_ids(), [1]);
s.outstanding_requests.fetch_add(1, Ordering::Relaxed);
let s = pool.get();
assert_eq!(s.id, 2);
assert_eq!(pool.stream_ids(), [1, 2]);
s.outstanding_requests.fetch_add(1, Ordering::Relaxed);
let s = pool.get();
assert_eq!(s.id, 2);
assert_eq!(pool.stream_ids(), [1, 2]);
s.outstanding_requests.fetch_add(1, Ordering::Relaxed);
let s = pool.get();
assert_eq!(s.id, 3);
assert_eq!(pool.stream_ids(), [1, 2, 3]);
Ok(())
}
#[tokio::test]
async fn fully_loaded_get() -> anyhow::Result<()> {
let transport = Arc::new(test_transport("ignored").await?);
let options = StreamPoolOptions {
max_streams: 6,
max_outstanding_requests: Some(10),
max_outstanding_bytes: None,
load_threshold: 0.2,
};
let pool = Arc::new(StreamPool::new(transport, options));
pool.seed([8, 5, 3, 8, 9, 11]);
assert_eq!(pool.stream_ids().len(), 6);
let s = pool.get();
assert_eq!(s.id, 3);
assert_eq!(pool.stream_ids().len(), 6);
Ok(())
}
#[tokio::test]
async fn evict_basic() -> anyhow::Result<()> {
let transport = Arc::new(test_transport("ignored").await?);
let pool = StreamPool::new(transport, StreamPoolOptions::default());
pool.seed([1, 2, 3, 4, 5, 6]);
assert_eq!(pool.stream_ids(), [1, 2, 3, 4, 5, 6]);
let s = pool.evict_and_replace(3);
assert_eq!(s.id, 7);
assert_eq!(pool.normalize_load(&s), 0.0);
assert_eq!(pool.stream_ids(), [1, 2, 4, 5, 6, 7]);
let s = pool.evict_and_replace(6);
assert_eq!(s.id, 8);
assert_eq!(pool.normalize_load(&s), 0.0);
assert_eq!(pool.stream_ids(), [1, 2, 4, 5, 7, 8]);
Ok(())
}
#[tokio::test]
async fn evict_lock_contention() -> anyhow::Result<()> {
let transport = Arc::new(test_transport("ignored").await?);
let pool = Arc::new(StreamPool::new(transport, StreamPoolOptions::default()));
pool.seed([1]);
let mut streams = JoinSet::new();
for _ in 0..1000 {
let p = pool.clone();
streams.spawn(async move { p.evict_and_replace(1) });
}
while let Some(s) = streams.join_next().await {
assert_eq!(s?.id, 2);
}
assert_eq!(pool.stream_ids(), [2]);
Ok(())
}
impl StreamPool {
pub(crate) fn seed(&self, loads: impl IntoIterator<Item = u64>) {
for load in loads.into_iter() {
let s = self.new_stream_entry();
s.outstanding_requests.store(load, Ordering::Relaxed);
self.streams.lock().unwrap().push(s);
}
}
pub(crate) fn stream_ids(&self) -> Vec<u64> {
let mut ids: Vec<_> = self.streams.lock().unwrap().iter().map(|s| s.id).collect();
ids.sort();
ids
}
pub(crate) fn lock(&self) -> MutexGuard<'_, Vec<StreamEntry>> {
self.streams.lock().unwrap()
}
}
}