mirror of
https://github.com/euzu/tuliprox.git
synced 2026-10-01 13:32:04 +02:00
Merge pull request #383 from euzu/feature/shared_stream_burst_buffer
Added Shared Stream Burst buffer
This commit is contained in:
@@ -6,10 +6,12 @@ use std::{
|
||||
use std::cmp::{max};
|
||||
use tokio::sync::mpsc::{channel, Sender};
|
||||
use tokio_stream::wrappers::ReceiverStream;
|
||||
use crate::api::model::BoxedProviderStream;
|
||||
use crate::api::model::{BoxedProviderStream};
|
||||
use crate::api::model::StreamError;
|
||||
use crate::tools::atomic_once_flag::AtomicOnceFlag;
|
||||
|
||||
pub const CHANNEL_SIZE: usize = 1024;
|
||||
|
||||
pub(in crate::api::model) struct BufferedStream {
|
||||
stream: ReceiverStream<Result<bytes::Bytes, StreamError>>,
|
||||
close_signal: Arc<AtomicOnceFlag>
|
||||
@@ -17,7 +19,8 @@ pub(in crate::api::model) struct BufferedStream {
|
||||
|
||||
impl BufferedStream {
|
||||
pub fn new(stream: BoxedProviderStream, buffer_size: usize, client_close_signal: Arc<AtomicOnceFlag>, _url: &str) -> Self {
|
||||
let (tx, rx) = channel(max(buffer_size, 4096));
|
||||
// TODO make channel_size based on bytes not entries
|
||||
let (tx, rx) = channel(max(buffer_size, CHANNEL_SIZE));
|
||||
tokio::spawn(Self::buffer_stream(tx, stream, Arc::clone(&client_close_signal)));
|
||||
Self {
|
||||
stream: ReceiverStream::new(rx),
|
||||
@@ -39,9 +42,6 @@ impl BufferedStream {
|
||||
}
|
||||
}
|
||||
Some(Err(err)) => {
|
||||
//trace!("Buffered Stream Error: {err:?}");
|
||||
// tokio::time::sleep(sleep_duration).await;
|
||||
// Attempt to send error to client
|
||||
if tx.send(Err(err)).await.is_err() {
|
||||
client_close_signal.notify();
|
||||
}
|
||||
|
||||
@@ -10,7 +10,6 @@ mod provider_stream_factory;
|
||||
mod shared_stream_manager;
|
||||
mod active_client_stream;
|
||||
mod throttled_stream;
|
||||
|
||||
pub(in crate) use self::transport_stream_buffer::*;
|
||||
pub(in crate::api) use self::provider_stream::*;
|
||||
pub(in crate::api) use self::persist_pipe_stream::*;
|
||||
|
||||
@@ -23,8 +23,6 @@ use url::Url;
|
||||
use crate::api::model::streams::buffered_stream::BufferedStream;
|
||||
use crate::api::model::streams::client_stream::ClientStream;
|
||||
|
||||
// TODO make this configurable
|
||||
pub const STREAM_QUEUE_SIZE: usize = 4096; // mpsc channel holding messages. with possible 8192byte chunks
|
||||
const RETRY_SECONDS: u64 = 5;
|
||||
const ERR_MAX_RETRY_COUNT: u32 = 5;
|
||||
|
||||
@@ -56,7 +54,7 @@ impl ProviderStreamFactoryOptions {
|
||||
let buffer_size = if stream_options.buffer_enabled {
|
||||
stream_options.buffer_size
|
||||
} else {
|
||||
STREAM_QUEUE_SIZE
|
||||
0
|
||||
};
|
||||
let filter_header = get_header_filter_for_item_type(item_type);
|
||||
let mut req_headers = get_headers_from_request(req_headers, &filter_header);
|
||||
|
||||
@@ -1,11 +1,10 @@
|
||||
use crate::api::model::AppState;
|
||||
use crate::api::model::StreamError;
|
||||
use crate::api::model::STREAM_QUEUE_SIZE;
|
||||
use crate::utils::debug_if_enabled;
|
||||
use bytes::Bytes;
|
||||
use bytes::{Bytes};
|
||||
use futures::stream::BoxStream;
|
||||
use futures::{Stream, StreamExt};
|
||||
use std::collections::HashMap;
|
||||
use std::collections::{HashMap, VecDeque};
|
||||
use std::sync::Arc;
|
||||
|
||||
use crate::api::model::BoxedProviderStream;
|
||||
@@ -15,9 +14,16 @@ use crate::utils::{trace_if_enabled};
|
||||
use shared::utils::sanitize_sensitive_info;
|
||||
use std::pin::Pin;
|
||||
use std::task::{Context, Poll};
|
||||
use tokio::sync::{mpsc, RwLock};
|
||||
use tokio::sync::{mpsc, Mutex, RwLock};
|
||||
use tokio::sync::mpsc::Sender;
|
||||
use tokio_stream::wrappers::ReceiverStream;
|
||||
use tokio_util::sync::CancellationToken;
|
||||
use crate::api::model::streams::buffered_stream::CHANNEL_SIZE;
|
||||
|
||||
// TODO make this configurable
|
||||
const MIN_SHARED_BUFFER_SIZE: usize = 1024 * 1024 * 12; // 12 MB
|
||||
|
||||
const YIELD_COUNTER:usize = 200;
|
||||
|
||||
///
|
||||
/// Wraps a `ReceiverStream` as Stream<Item = Result<Bytes, `StreamError`>>
|
||||
@@ -47,6 +53,40 @@ fn convert_stream(stream: BoxStream<Bytes>) -> BoxStream<Result<Bytes, StreamErr
|
||||
|
||||
type SubscriberId = String;
|
||||
|
||||
|
||||
struct BurstBuffer {
|
||||
buffer: VecDeque<Bytes>,
|
||||
buffer_size: usize,
|
||||
current_bytes: usize,
|
||||
}
|
||||
|
||||
impl BurstBuffer {
|
||||
pub fn new(buf_size: usize) -> Self {
|
||||
Self {
|
||||
buffer: VecDeque::with_capacity(buf_size),
|
||||
buffer_size: buf_size,
|
||||
current_bytes: 0,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn snapshot(&self) -> VecDeque<Bytes> {
|
||||
self.buffer.iter().cloned().collect::<VecDeque<Bytes>>()
|
||||
}
|
||||
|
||||
pub fn push(&mut self, packet: &Bytes) {
|
||||
while self.current_bytes > self.buffer_size {
|
||||
if let Some(popped) = self.buffer.pop_front() {
|
||||
self.current_bytes -= popped.len();
|
||||
} else {
|
||||
self.current_bytes = 0;
|
||||
break;
|
||||
}
|
||||
}
|
||||
self.current_bytes += packet.len();
|
||||
self.buffer.push_back(packet.clone());
|
||||
}
|
||||
}
|
||||
|
||||
/// Represents the state of a shared provider URL.
|
||||
///
|
||||
/// - `headers`: The initial connection headers used during the setup of the shared stream.
|
||||
@@ -57,6 +97,7 @@ pub struct SharedStreamState {
|
||||
subscribers: RwLock<HashMap<SubscriberId, CancellationToken>>,
|
||||
broadcaster: tokio::sync::broadcast::Sender<Bytes>,
|
||||
stop_token: CancellationToken,
|
||||
burst_buffer: Arc<Mutex<BurstBuffer>>,
|
||||
}
|
||||
|
||||
impl Drop for SharedStreamState {
|
||||
@@ -74,13 +115,16 @@ impl SharedStreamState {
|
||||
guard.disable_release();
|
||||
}
|
||||
let (broadcaster, _) = tokio::sync::broadcast::channel(buf_size);
|
||||
// TODO channel size versus byte size, channels are chunk sized, burst_buffer byte sized
|
||||
let burst_buffer_size_in_bytes = MIN_SHARED_BUFFER_SIZE.max(buf_size * 1024 * 12);
|
||||
Self {
|
||||
headers,
|
||||
buf_size,
|
||||
provider_guard,
|
||||
subscribers: RwLock::new(HashMap::new()), //Arc::new(RwLock::new(Vec::new())),
|
||||
subscribers: RwLock::new(HashMap::new()),
|
||||
broadcaster,
|
||||
stop_token: CancellationToken::new(),
|
||||
burst_buffer : Arc::new(Mutex::new(BurstBuffer::new(burst_buffer_size_in_bytes))),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -91,9 +135,20 @@ impl SharedStreamState {
|
||||
self.subscribers.write().await.insert(addr.to_string(), cancel_token.clone());
|
||||
|
||||
let address = addr.to_string();
|
||||
let client_tx_clone = client_tx.clone();
|
||||
let burst_buffer = self.burst_buffer.clone();
|
||||
|
||||
tokio::spawn(async move {
|
||||
let snapshot = {
|
||||
let buffer = burst_buffer.lock().await;
|
||||
buffer.snapshot()
|
||||
};
|
||||
send_burst_buffer(&snapshot, &client_tx_clone, &cancel_token).await;
|
||||
|
||||
let mut loop_cnt = 0;
|
||||
loop {
|
||||
tokio::select! {
|
||||
loop_cnt += 1;
|
||||
tokio::select! {
|
||||
biased;
|
||||
|
||||
() = cancel_token.cancelled() => {
|
||||
@@ -107,10 +162,17 @@ impl SharedStreamState {
|
||||
debug!("Shared stream client send error: {address} {err}");
|
||||
break;
|
||||
}
|
||||
tokio::task::yield_now().await;
|
||||
if loop_cnt > YIELD_COUNTER {
|
||||
tokio::task::yield_now().await;
|
||||
loop_cnt = 0;
|
||||
}
|
||||
}
|
||||
Err(tokio::sync::broadcast::error::RecvError::Lagged(skipped)) => {
|
||||
trace!("Client lagged behind. Skipped {skipped} messages. {address}");
|
||||
if loop_cnt > YIELD_COUNTER {
|
||||
tokio::task::yield_now().await;
|
||||
loop_cnt = 0;
|
||||
}
|
||||
}
|
||||
Err(_) => break,
|
||||
}
|
||||
@@ -136,9 +198,10 @@ impl SharedStreamState {
|
||||
let streaming_url = stream_url.to_string();
|
||||
let sender = self.broadcaster.clone();
|
||||
let stop_token = self.stop_token.clone();
|
||||
let burst_buffer = self.burst_buffer.clone();
|
||||
|
||||
tokio::spawn(async move {
|
||||
let mut counter = 0u32;
|
||||
let mut counter = 0usize;
|
||||
loop {
|
||||
tokio::select! {
|
||||
biased;
|
||||
@@ -151,6 +214,11 @@ impl SharedStreamState {
|
||||
item = source_stream.next() => {
|
||||
match item {
|
||||
Some(Ok(data)) => {
|
||||
{
|
||||
let mut buffer = burst_buffer.lock().await;
|
||||
buffer.push(&data);
|
||||
}
|
||||
|
||||
match sender.send(data) {
|
||||
Ok(clients) => {
|
||||
if clients == 0 {
|
||||
@@ -158,7 +226,7 @@ impl SharedStreamState {
|
||||
break;
|
||||
}
|
||||
counter += 1;
|
||||
if counter >= 100 {
|
||||
if counter >= YIELD_COUNTER {
|
||||
tokio::task::yield_now().await;
|
||||
counter = 0;
|
||||
}
|
||||
@@ -171,6 +239,8 @@ impl SharedStreamState {
|
||||
}
|
||||
Some(Err(e)) => {
|
||||
trace!("Shared stream received error: {e:?}");
|
||||
tokio::task::yield_now().await;
|
||||
|
||||
}
|
||||
None => {
|
||||
debug_if_enabled!("Source stream ended. Closing shared provider stream {}", sanitize_sensitive_info(&streaming_url));
|
||||
@@ -299,7 +369,7 @@ impl SharedStreamManager {
|
||||
S: Stream<Item=Result<Bytes, E>> + Unpin + 'static + Send,
|
||||
E: std::fmt::Debug + Send,
|
||||
{
|
||||
let buf_size = STREAM_QUEUE_SIZE.max(buffer_size);
|
||||
let buf_size = CHANNEL_SIZE.max(buffer_size);
|
||||
let shared_state = Arc::new(SharedStreamState::new(headers, buf_size, provider_guard));
|
||||
app_state.shared_stream_manager.register(stream_url, Arc::clone(&shared_state)).await;
|
||||
debug_if_enabled!("Created shared provider stream {}", sanitize_sensitive_info(stream_url));
|
||||
@@ -317,4 +387,18 @@ impl SharedStreamManager {
|
||||
let manager = Arc::clone(&app_state.shared_stream_manager);
|
||||
app_state.shared_stream_manager.subscribe_stream(stream_url, addr, manager).await
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
async fn send_burst_buffer(
|
||||
start_buffer: &VecDeque<Bytes>,
|
||||
client_tx: &Sender<Bytes>,
|
||||
cancellation_token: &CancellationToken) {
|
||||
for buf in start_buffer {
|
||||
if cancellation_token.is_cancelled() { return; }
|
||||
if let Err(err) = client_tx.send(buf.clone()).await {
|
||||
debug!("Error sending current chunk: {err}");
|
||||
return; // stop on send error
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user