use std::collections::{BTreeMap, VecDeque};
use std::time::Duration;
use bytes::{Buf, BufMut, BytesMut};
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use tokio::io::AsyncWriteExt;
use tokio::sync::{Mutex, Notify, RwLock};
use tokio::task::yield_now;
use tokio::time::{Instant, sleep_until};
use wtransport::error::StreamReadError;
use wtransport::{RecvStream, SendStream};
use crate::model::control::fetch::Fetch;
use crate::model::control::subscribe::Subscribe;
use crate::model::data::constant::FetchHeaderType;
use crate::model::data::fetch_header::FetchHeader;
use crate::model::data::fetch_object::{FetchObject, FetchObjectContext};
use crate::model::data::object::Object;
use crate::model::data::subgroup_header::SubgroupHeader;
use crate::model::data::subgroup_object::SubgroupObject;
use crate::model::error::ParseError;
use tracing::{debug, error, info, trace};
type ObjectParseResult =
Result<(Option<u64>, Option<FetchObjectContext>, Option<Object>), ParseError>;
const DATA_STREAM_TIMEOUT: Duration = Duration::from_secs(15);
const MTU_SIZE: usize = 1500;
#[derive(Debug, Clone)]
pub enum HeaderInfo {
Fetch {
header: FetchHeader,
fetch_request: Fetch, },
Subgroup {
header: SubgroupHeader,
},
}
#[derive(Debug, Clone)]
pub struct FetchRequest {
pub original_request_id: u64,
pub requested_by: usize, pub fetch_request: Fetch,
pub track_alias: u64, }
#[derive(Debug, Clone)]
pub struct SubscribeRequest {
pub original_request_id: u64,
pub requested_by: usize, pub original_subscribe_request: Subscribe, pub subscribe_request_to_publisher: Option<Subscribe>, }
impl FetchRequest {
pub fn new(
original_request_id: u64,
requested_by: usize,
fetch_request: Fetch,
track_alias: u64,
) -> Self {
Self {
original_request_id,
requested_by,
fetch_request,
track_alias,
}
}
}
impl SubscribeRequest {
pub fn new(
original_request_id: u64,
requested_by: usize,
original_subscribe_request: Subscribe,
subscribe_request_to_publisher: Option<Subscribe>,
) -> Self {
Self {
original_request_id,
requested_by,
original_subscribe_request,
subscribe_request_to_publisher,
}
}
}
pub struct SendDataStream {
send_stream: Arc<Mutex<SendStream>>,
header_info: HeaderInfo,
fetch_prev_ctx: Option<FetchObjectContext>,
}
impl SendDataStream {
pub async fn new(
send_stream: Arc<Mutex<SendStream>>,
header_info: HeaderInfo,
) -> Result<Self, ParseError> {
let mut buf = BytesMut::new();
match &header_info {
HeaderInfo::Fetch { header, .. } => {
buf.extend_from_slice(&header.serialize()?);
}
HeaderInfo::Subgroup { header, .. } => {
buf.extend_from_slice(&header.serialize(None)?);
}
}
send_stream
.lock()
.await
.write_all(&buf)
.await
.map_err(|e| ParseError::Other {
context: "SendDataStream::new(header write)",
msg: e.to_string(),
})?;
Ok(Self {
send_stream,
header_info,
fetch_prev_ctx: None,
})
}
pub async fn send_object(
&mut self,
object: &Object,
previous_object_id: Option<u64>,
) -> Result<(), ParseError> {
let mut buf = BytesMut::new();
let object = object.clone();
match &self.header_info {
HeaderInfo::Fetch { .. } => {
let payload = object.try_into_fetch()?;
let fetch_obj = FetchObject::Object(payload);
buf.extend_from_slice(&fetch_obj.serialize(self.fetch_prev_ctx.as_ref())?);
self.fetch_prev_ctx = fetch_obj.context();
}
HeaderInfo::Subgroup { header, .. } => {
let has_extensions = header.header_type.has_extensions();
let subgroup_obj = object.try_into_subgroup()?;
buf.extend_from_slice(&subgroup_obj.serialize(previous_object_id, has_extensions)?);
}
}
self
.send_stream
.lock()
.await
.write_all(&buf)
.await
.map_err(|e| ParseError::Other {
context: "SendDataStream::send_object",
msg: e.to_string(),
})?;
Ok(())
}
pub async fn flush(&mut self) -> Result<(), ParseError> {
debug!("SendDataStream::flush() called");
self
.send_stream
.lock()
.await
.flush()
.await
.map_err(|e| ParseError::Other {
context: "SendDataStream::flush",
msg: e.to_string(),
})
}
pub async fn finish(&mut self) -> Result<(), ParseError> {
debug!("SendDataStream::finish() called");
self
.send_stream
.lock()
.await
.finish()
.await
.map_err(|e| ParseError::Other {
context: "SendDataStream::finish",
msg: e.to_string(),
})
}
}
#[derive(Debug)]
pub enum RecvDataStreamReadError {
ParseError(ParseError),
StreamClosed,
}
pub struct RecvDataStream {
recv_stream: Arc<Mutex<RecvStream>>,
header_info: Arc<Mutex<Option<HeaderInfo>>>,
pending_fetches: Arc<RwLock<BTreeMap<u64, FetchRequest>>>, objects: Arc<RwLock<VecDeque<Object>>>, is_closed: Arc<AtomicBool>,
started_read_task: Arc<AtomicBool>,
notify: Arc<Notify>,
}
impl RecvDataStream {
pub fn new(
recv_stream: RecvStream,
pending_fetches: Arc<RwLock<BTreeMap<u64, FetchRequest>>>, ) -> Self {
Self {
recv_stream: Arc::new(Mutex::new(recv_stream)),
header_info: Arc::new(Mutex::new(None)), pending_fetches,
objects: Arc::new(RwLock::new(VecDeque::new())), is_closed: Arc::new(AtomicBool::new(false)),
started_read_task: Arc::new(AtomicBool::new(false)),
notify: Arc::new(Notify::new()),
}
}
pub async fn get_header_info(&self) -> Option<HeaderInfo> {
debug!("RecvDataStream::get_header_info() called");
let header_info = self.header_info.lock().await;
header_info.clone()
}
async fn read(
recv_stream: Arc<Mutex<RecvStream>>,
is_closed: Arc<AtomicBool>,
the_header_info: Arc<Mutex<Option<HeaderInfo>>>,
pending_fetches: Arc<RwLock<BTreeMap<u64, FetchRequest>>>,
objects: Arc<RwLock<VecDeque<Object>>>,
notify: Arc<Notify>,
) -> Result<(), RecvDataStreamReadError> {
let mut header_info = None;
let mut recv_buf = Box::new([0u8; MTU_SIZE]);
let mut recv_bytes = BytesMut::new();
let mut timeout_at = Instant::now() + DATA_STREAM_TIMEOUT;
let mut previous_object_id: Option<u64> = None;
let mut fetch_prev_ctx: Option<FetchObjectContext> = None;
loop {
let bytes_cursor = recv_bytes.clone().freeze();
if !recv_bytes.is_empty() && header_info.is_none() {
let is_fetch = recv_bytes[0] == FetchHeaderType::Type0x05 as u8;
header_info = Self::read_header(
bytes_cursor,
is_fetch,
is_closed.clone(),
pending_fetches.clone(),
)
.await
.map_err(|e| {
error!("Failed to parse header: {:?}", e);
RecvDataStreamReadError::ParseError(e)
})?;
let consumed = if let Some((consumed, _)) = header_info.clone() {
consumed
} else {
0
};
recv_bytes.advance(consumed);
*the_header_info.lock().await = Some(header_info.clone().unwrap().1.clone());
}
loop {
let bytes_cursor = recv_bytes.clone().freeze();
let mut consumed = 0;
if !recv_bytes.is_empty() {
let header_info = header_info.clone().unwrap().1;
let (c, object_id, new_fetch_ctx) = Self::read_object(
bytes_cursor,
&header_info,
is_closed.clone(),
objects.clone(),
&previous_object_id,
&fetch_prev_ctx,
)
.await
.map_err(|e| {
error!("Failed to parse object: {:?}", e);
RecvDataStreamReadError::ParseError(e)
})?;
if c > 0 {
previous_object_id = object_id;
if new_fetch_ctx.is_some() {
fetch_prev_ctx = new_fetch_ctx;
}
notify.notify_waiters();
recv_bytes.advance(c);
}
consumed = c;
trace!(
"previous_object_id: {:?} object_id: {:?} consumed: {}",
previous_object_id, &object_id, consumed
);
}
if consumed > 0 {
continue;
} else {
break;
}
}
if is_closed.load(Ordering::Relaxed) {
notify.notify_waiters();
return Ok(());
}
let stream = recv_stream.clone();
let mut stream = stream.lock().await;
tokio::select! {
biased;
_ = sleep_until(timeout_at) => {
info!("Timeout while waiting for data");
is_closed.store(true, Ordering::Relaxed);
return Err(RecvDataStreamReadError::ParseError(ParseError::Timeout { context: "RecvDataStream::new(header_read)" }));
}
read_result = stream.read(&mut recv_buf[..]) => {
match read_result {
Ok(Some(n)) => {
if n > 0 {
recv_bytes.put_slice(&recv_buf[..n]);
timeout_at = Instant::now() + DATA_STREAM_TIMEOUT;
} else {
}
}
Ok(None) => {
is_closed.store(true, Ordering::Relaxed);
}
Err(e) => {
debug!("RecvDataStream::read() Read error: {:?}", e);
is_closed.store(true, Ordering::Relaxed);
if e == StreamReadError::NotConnected {
return Err(RecvDataStreamReadError::StreamClosed);
}
return Err(RecvDataStreamReadError::ParseError(ParseError::Other { context: "RecvDataStream::new(header_read)", msg:e.to_string() }));
}
}
}
}
}
}
async fn read_header(
mut bytes_cursor: bytes::Bytes,
is_fetch: bool,
is_closed: Arc<AtomicBool>,
pending_fetches: Arc<RwLock<BTreeMap<u64, FetchRequest>>>,
) -> Result<Option<(usize, HeaderInfo)>, ParseError> {
debug!("RecvDataStream::read_header() called");
let original_remaining = bytes_cursor.remaining();
if is_fetch {
match FetchHeader::deserialize(&mut bytes_cursor) {
Ok(fetch_header) => {
let pending_fetches = pending_fetches.read().await;
if let Some(fetch_request) = pending_fetches.get(&fetch_header.request_id) {
let consumed = original_remaining - bytes_cursor.remaining();
let header_info = HeaderInfo::Fetch {
header: fetch_header,
fetch_request: fetch_request.fetch_request.clone(),
};
debug!(
"RecvDataStream::read_header() Parsed FetchHeader: {:?}",
header_info
);
Ok(Some((consumed, header_info)))
} else {
drop(pending_fetches);
is_closed.store(true, Ordering::Relaxed);
Err(ParseError::ProtocolViolation {
context: "RecvDataStream::new(FetchHeader validation)",
details: format!(
"Received FetchHeader for unknown request_id: {}",
fetch_header.request_id
),
})
}
}
Err(ParseError::NotEnoughBytes { .. }) => {
Ok(None) }
Err(e) => {
is_closed.store(true, Ordering::Relaxed);
Err(ParseError::ProtocolViolation {
context: "RecvDataStream::new(FetchHeader validation)",
details: e.to_string(),
})
}
}
} else {
match SubgroupHeader::deserialize(&mut bytes_cursor) {
Ok(subgroup_header) => {
let consumed = original_remaining - bytes_cursor.remaining();
let header_info = HeaderInfo::Subgroup {
header: subgroup_header,
};
Ok(Some((consumed, header_info)))
}
Err(ParseError::NotEnoughBytes { .. }) => {
Ok(None) }
Err(e) => {
is_closed.store(true, Ordering::Relaxed);
Err(ParseError::ProtocolViolation {
context: "RecvDataStream::new(SubgroupHeader validation)",
details: e.to_string(),
})
}
}
}
}
async fn read_object(
mut bytes_cursor: bytes::Bytes,
header_info: &HeaderInfo,
is_closed: Arc<AtomicBool>,
objects: Arc<RwLock<VecDeque<Object>>>,
previous_object_id: &Option<u64>,
fetch_prev_ctx: &Option<FetchObjectContext>,
) -> Result<(usize, Option<u64>, Option<FetchObjectContext>), ParseError> {
if !bytes_cursor.is_empty() {
let original_remaining = bytes_cursor.remaining();
let parse_result: ObjectParseResult = match header_info {
HeaderInfo::Fetch { .. } => {
FetchObject::deserialize(&mut bytes_cursor, fetch_prev_ctx.as_ref()).and_then(
|fetch_obj| {
let new_ctx = fetch_obj.context();
match fetch_obj {
FetchObject::Object(payload) => {
let object = Object::try_from_fetch(payload, 0)?;
Ok((None, new_ctx, Some(object)))
}
FetchObject::EndOfRange { .. } => {
debug!("FetchObject::EndOfRange received, skipping");
Ok((None, None, None))
}
}
},
)
}
HeaderInfo::Subgroup { header, .. } => {
let has_extensions = header.header_type.has_extensions();
SubgroupObject::deserialize(&mut bytes_cursor, previous_object_id, has_extensions)
.and_then(|subgroup_obj| {
let object_id = subgroup_obj.object_id;
let object = Object::try_from_subgroup(
subgroup_obj,
header.track_alias,
header.group_id,
header.subgroup_id,
header.publisher_priority,
)?;
Ok((Some(object_id), None, Some(object)))
})
}
};
match parse_result {
Ok((object_id, new_ctx, maybe_object)) => {
let consumed = original_remaining - bytes_cursor.remaining();
debug!(
"consumed: {} Parsed payload object: {:?}",
consumed, maybe_object
);
if let Some(object) = maybe_object {
let mut objects = objects.write().await;
objects.push_back(object);
}
Ok((consumed, object_id, new_ctx))
}
Err(ParseError::NotEnoughBytes { .. }) => {
trace!("Not enough bytes to parse the object, continuing to read...");
Ok((0, None, None))
}
Err(e) => {
is_closed.store(true, Ordering::Relaxed);
Err(ParseError::ProtocolViolation {
context: "RecvDataStream::next_object(parse_result)",
details: e.to_string(),
})
}
}
} else {
debug!("No bytes available to parse an object");
Ok((0, None, None))
}
}
pub async fn next_object(&self) -> (&Self, Option<Object>) {
if self
.started_read_task
.compare_exchange(false, true, Ordering::Relaxed, Ordering::Relaxed)
.is_ok()
{
let recv_stream = self.recv_stream.clone();
let is_closed = self.is_closed.clone();
let pending_fetches = self.pending_fetches.clone();
let objects = self.objects.clone();
let header_info = self.header_info.clone();
let notify = self.notify.clone();
tokio::spawn(async move {
match Self::read(
recv_stream,
is_closed,
header_info,
pending_fetches,
objects,
notify,
)
.await
{
Ok(_) => debug!("RecvDataStream read task completed successfully"),
Err(e) => {
error!("RecvDataStream read task encountered an error: {:?}", e);
if matches!(e, RecvDataStreamReadError::StreamClosed) {
debug!("Stream is closed, returning EOF");
} else {
error!("RecvDataStream read task encountered an error: {:?}", e);
}
}
}
});
}
loop {
{
let mut objects = self.objects.write().await;
if let Some(object) = objects.pop_front() {
return (self, Some(object));
}
}
if self.is_closed.load(Ordering::Relaxed) {
if self.objects.read().await.is_empty() {
debug!("Stream is closed, returning EOF");
return (self, None);
} else {
debug!("Stream is closed, but has objects still");
yield_now().await
}
} else {
self.notify.notified().await;
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::model::common::location::Location;
use crate::model::common::pair::KeyValuePair;
use crate::model::common::tuple::{Tuple, TupleField};
use crate::model::control::constant::{FetchType, GroupOrder};
use crate::model::control::fetch::JoiningFetchProps;
use crate::model::control::{fetch::Fetch, subscribe::Subscribe};
use crate::model::data::constant::SubgroupHeaderType;
use crate::model::extension_header::object_extension::ObjectExtension;
use crate::model::parameter::authorization_token::AuthorizationToken;
use crate::model::parameter::message_parameter::MessageParameter;
use bytes::Bytes;
use std::error::Error;
use std::sync::Arc;
use tokio::sync::Mutex;
use tokio::time::{Duration, sleep};
use wtransport::endpoint::IntoConnectOptions;
use wtransport::{ClientConfig, Connection, Endpoint, Identity, RecvStream, SendStream};
fn make_fetch_header_and_request() -> (FetchHeader, Fetch) {
let fetch = Fetch {
request_id: 161803,
fetch_type: FetchType::AbsoluteFetch,
standalone_fetch_props: None,
joining_fetch_props: Some(JoiningFetchProps {
joining_request_id: 119,
joining_start: 73,
}),
parameters: vec![
MessageParameter::new_authorization_token(AuthorizationToken::new_use_value(
0,
Bytes::from_static(b"test-token"),
)),
MessageParameter::new_subscriber_priority(42),
MessageParameter::new_group_order(GroupOrder::Ascending),
],
};
(FetchHeader { request_id: 161803 }, fetch)
}
fn make_fetch_object() -> crate::model::data::fetch_object::FetchObjectPayload {
use crate::model::data::constant::ObjectForwardingPreference;
use crate::model::data::fetch_object::FetchObjectPayload;
FetchObjectPayload {
group_id: 9,
subgroup_id: 144,
object_id: 10,
publisher_priority: 255,
forwarding_preference: ObjectForwardingPreference::Subgroup,
extension_headers: Some(vec![
ObjectExtension::Unknown {
kvp: KeyValuePair::try_new_varint(0, 10).unwrap(),
},
ObjectExtension::Unknown {
kvp: KeyValuePair::try_new_bytes(1, Bytes::from_static(b"wololoo")).unwrap(),
},
]),
payload: Bytes::from_static(
b"01239gjawkk92837aljwdnjwandjnanwdjnajwndkjawndjkanwdkjnawkjddmi",
),
}
}
fn make_object_from_fetch(
fetch_obj: &crate::model::data::fetch_object::FetchObjectPayload,
) -> Object {
Object::try_from_fetch(fetch_obj.clone(), 0).unwrap()
}
#[allow(dead_code)]
fn make_subgroup_header_and_request() -> (SubgroupHeader, Subscribe) {
let request_id = 128242;
let track_namespace = Tuple::from_utf8_path("nein/nein/nein");
let track_name = TupleField::from_utf8("track_42");
let start_location = Location {
group: 81,
object: 81,
};
let subscribe = Subscribe::new_absolute_range(
request_id,
track_namespace,
track_name,
start_location,
25,
vec![
MessageParameter::new_subscriber_priority(31),
MessageParameter::new_group_order(GroupOrder::Original),
MessageParameter::new_forward(true),
],
);
let header_type = SubgroupHeaderType::try_new(0x15).unwrap();
let track_alias = 999;
let group_id = 9;
let subgroup_id = Some(11);
let publisher_priority = Some(255);
let subgroup_header = SubgroupHeader {
header_type,
track_alias,
group_id,
subgroup_id,
publisher_priority,
};
(subgroup_header, subscribe)
}
#[allow(dead_code)]
fn make_subgroup_object() -> SubgroupObject {
let object_id: u64 = 10;
let extension_headers = Some(vec![
ObjectExtension::Unknown {
kvp: KeyValuePair::try_new_varint(0, 10).unwrap(),
},
ObjectExtension::Unknown {
kvp: KeyValuePair::try_new_bytes(1, Bytes::from_static(b"wololoo")).unwrap(),
},
]);
let object_status = None;
let payload = Some(Bytes::from_static(b"01239gjawkk92837aldmi"));
SubgroupObject {
object_id,
extension_headers,
payload,
object_status,
}
}
#[allow(dead_code)]
fn make_object_from_subgroup(subgroup_obj: &SubgroupObject, header: &SubgroupHeader) -> Object {
Object::try_from_subgroup(
subgroup_obj.clone(),
header.track_alias,
header.group_id,
header.subgroup_id,
header.publisher_priority,
)
.unwrap()
}
struct TestSetup {
client: Connection,
server: Connection,
}
impl TestSetup {
async fn new() -> Result<Self, Box<dyn Error>> {
let server_identity = Identity::self_signed(std::iter::once("localhost"))
.map_err(|e| format!("Failed to create server identity: {e}"))?;
let server_cert_hash = server_identity.certificate_chain().as_slice()[0].hash();
let server_config = wtransport::ServerConfig::builder()
.with_bind_address(
"127.0.0.1:0"
.parse()
.map_err(|e| format!("Failed to parse bind address: {e}"))?,
)
.with_identity(server_identity)
.build();
let server_endpoint = Endpoint::server(server_config)
.map_err(|e| format!("Failed to create server endpoint: {e}"))?;
let server_addr = server_endpoint
.local_addr()
.map_err(|e| format!("Failed to get server local address: {e}"))?;
let (tx, rx) = tokio::sync::oneshot::channel();
tokio::spawn(async move {
let result = async {
let incoming = server_endpoint.accept().await;
let session_request = incoming
.await
.map_err(|e| format!("Failed to await session request: {e}"))
.unwrap();
let server = session_request.accept().await.unwrap();
Ok::<_, Box<dyn Error + Send>>(server)
}
.await;
if tx.send(result).is_err() {
eprintln!("Failed to send server connection result back through the channel");
}
});
let client_config = ClientConfig::builder()
.with_bind_default()
.with_server_certificate_hashes(vec![server_cert_hash])
.build();
let client_endpoint = Endpoint::client(client_config)
.map_err(|e| format!("Failed to create client endpoint: {e}"))?;
let client = client_endpoint
.connect(
format!("https://{}:{}", server_addr.ip(), server_addr.port())
.as_str()
.into_options(),
)
.await
.map_err(|e| format!("Client connection failed: {e}"))?;
let server = rx
.await
.map_err(|_| "Server task failed to send connection back")?
.map_err(|e| format!("Server connection error: {e}"))?;
Ok(Self { client, server })
}
async fn create_data_stream_pair(&self) -> Result<(SendStream, RecvStream), Box<dyn Error>> {
let client = self.client.clone();
let send_fut = tokio::spawn(async move {
let send_uni = client.open_uni().await.unwrap();
send_uni.await
});
let server = self.server.clone();
let recv_fut = tokio::spawn(async move {
server
.accept_uni()
.await
.map_err(|e| format!("Failed to accept server uni stream: {e}"))
});
let (send_res, recv_res) = tokio::try_join!(send_fut, recv_fut)?;
let send = send_res.map_err(|e| format!("Failed to open client uni stream: {e}"))?;
let recv = recv_res.map_err(|e| format!("Failed to accept server uni stream: {e}"))?;
Ok((send, recv))
}
}
async fn setup_stream_pair() -> (SendStream, RecvStream) {
let setup = TestSetup::new()
.await
.expect("Failed to setup test transport");
setup
.create_data_stream_pair()
.await
.expect("Failed to create data stream pair")
}
#[tokio::test]
async fn test_partial_object_completion() {
let (send, recv) = setup_stream_pair().await;
let (fetch_header, fetch_req) = make_fetch_header_and_request();
let mut pending_fetches = BTreeMap::new();
pending_fetches.insert(
fetch_req.request_id,
FetchRequest {
original_request_id: fetch_req.request_id,
requested_by: 1,
fetch_request: fetch_req.clone(),
track_alias: 1,
},
);
let sender = SendDataStream::new(
Arc::new(Mutex::new(send)),
HeaderInfo::Fetch {
header: fetch_header,
fetch_request: fetch_req.clone(),
},
)
.await
.unwrap();
let fetch_obj = make_fetch_object();
let object = make_object_from_fetch(&fetch_obj);
let pending_fetches = Arc::new(RwLock::new(pending_fetches));
let receiver = RecvDataStream::new(recv, pending_fetches);
let bytes = FetchObject::Object(fetch_obj.clone())
.serialize(None)
.unwrap();
let half = bytes.len() / 2;
let first_half = &bytes[..half];
let second_half = &bytes[half..];
sender
.send_stream
.lock()
.await
.write_all(first_half)
.await
.unwrap();
let second_half = second_half.to_vec();
tokio::spawn({
let send_stream = sender.send_stream.clone();
async move {
sleep(Duration::from_millis(100)).await;
let mut s = send_stream.lock().await;
s.write_all(&second_half).await.unwrap();
}
});
let received = receiver.next_object().await.1.unwrap();
assert_eq!(object, received);
}
}