use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use crate::transport::{RecvHalf, SendHalf};
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
use weida_core::{Error, ErrorCode, PeerIdentity, TraceContext};
use weida_protocol::header::limits::MAX_REPORT_LEVELS;
use weida_protocol::header::{Acknowledgement, CursorLevel, ReportMode};
use weida_protocol::{DataHeader, ErrorHeader, FrameKind, codes};
use crate::conn::{ConnHandle, Ctl, read_frame, write_error_frame};
use crate::cursor::{Cursors, Reporter, order_report};
use crate::drain::Receipt;
use crate::ordering::Gap;
#[derive(Clone, Debug, Default)]
pub struct TransferMeta {
pub content_type: Option<String>,
pub content_len: Option<u64>,
pub trace: Option<TraceContext>,
pub topic: Option<String>,
pub achieved: Option<Acknowledgement>,
pub report: Vec<CursorLevel>,
pub report_mode: ReportMode,
}
impl TransferMeta {
pub fn with_content_type(mut self, content_type: impl Into<String>) -> TransferMeta {
self.content_type = Some(content_type.into());
self
}
pub fn with_content_len(mut self, len: u64) -> TransferMeta {
self.content_len = Some(len);
self
}
pub fn with_trace(mut self, trace: TraceContext) -> TransferMeta {
self.trace = Some(trace);
self
}
pub fn with_topic(mut self, topic: impl Into<String>) -> TransferMeta {
self.topic = Some(topic.into());
self
}
pub fn with_achieved(mut self, achieved: Acknowledgement) -> TransferMeta {
self.achieved = Some(achieved);
self
}
pub fn with_report(mut self, levels: impl IntoIterator<Item = CursorLevel>) -> TransferMeta {
let mut levels: Vec<CursorLevel> = levels.into_iter().collect();
levels.sort_unstable();
levels.dedup();
self.report = levels;
self
}
pub fn with_report_mode(mut self, mode: ReportMode) -> TransferMeta {
self.report_mode = mode;
self
}
}
#[derive(Clone, Debug)]
pub struct IncomingMeta {
pub endpoint: Option<String>,
pub content_len: Option<u64>,
pub content_type: Option<String>,
pub trace: Option<TraceContext>,
pub tracestate: Option<String>,
pub topic: Option<String>,
pub peer: Option<PeerIdentity>,
pub sequence: Option<u64>,
pub gap: Option<Gap>,
pub achieved: Option<Acknowledgement>,
pub report: Vec<CursorLevel>,
pub report_mode: ReportMode,
pub report_id: Option<u64>,
}
impl IncomingMeta {
pub(crate) fn from_header(header: &DataHeader, peer: Option<PeerIdentity>) -> IncomingMeta {
IncomingMeta {
endpoint: header.endpoint.clone(),
content_len: header.content_len,
content_type: header.content_type.clone(),
trace: header
.traceparent
.as_deref()
.and_then(|v| TraceContext::parse_traceparent(v).ok()),
tracestate: header.tracestate.clone(),
topic: header.topic.clone(),
peer,
sequence: header.sequence,
gap: None,
achieved: header.achieved,
report: header.report.clone(),
report_mode: header.report_mode,
report_id: header.report_id,
}
}
pub(crate) fn with_gap(mut self, gap: Option<Gap>) -> IncomingMeta {
self.gap = gap;
self
}
}
pub fn new_trace() -> TraceContext {
use rand::Rng;
let mut rng = rand::rng();
loop {
let trace_id: [u8; 16] = rng.random();
let span_id: [u8; 8] = rng.random();
if let Some(ctx) = TraceContext::new(trace_id, span_id, 0x01) {
return ctx;
}
}
}
pub(crate) fn data_header(
endpoint: Option<&str>,
meta: &TransferMeta,
tracestate: Option<String>,
report_id: Option<u64>,
) -> Result<(DataHeader, Option<TraceContext>), Error> {
if meta.report.len() > MAX_REPORT_LEVELS {
return Err(Error::LimitExceeded);
}
let trace = meta.trace;
let header = DataHeader {
endpoint: endpoint.map(str::to_owned),
content_len: meta.content_len,
content_type: meta.content_type.clone(),
traceparent: trace.map(|t| t.to_traceparent()),
tracestate,
topic: meta.topic.clone(),
sequence: None,
producer: None,
achieved: meta.achieved,
report_id,
report: meta.report.clone(),
report_mode: meta.report_mode,
};
Ok((header, trace))
}
pub(crate) fn outgoing_header(
conn: &ConnHandle,
endpoint: Option<&str>,
meta: &TransferMeta,
tracestate: Option<String>,
) -> Result<(DataHeader, Option<TraceContext>, Option<Cursors>), Error> {
if meta.report.is_empty() {
let (header, trace) = data_header(endpoint, meta, tracestate, None)?;
return Ok((header, trace, None));
}
if meta.report.len() > MAX_REPORT_LEVELS {
return Err(Error::LimitExceeded);
}
let (report_id, cursors) = order_report(conn);
let (header, trace) = data_header(endpoint, meta, tracestate, Some(report_id))?;
Ok((header, trace, Some(cursors)))
}
pub struct OutgoingTransfer {
stream: SendHalf,
trace: Option<TraceContext>,
settled: bool,
conn: ConnHandle,
cursors: Option<Cursors>,
}
impl OutgoingTransfer {
pub(crate) fn new(
stream: SendHalf,
trace: Option<TraceContext>,
conn: ConnHandle,
cursors: Option<Cursors>,
) -> OutgoingTransfer {
OutgoingTransfer {
stream,
trace,
settled: false,
conn,
cursors,
}
}
pub fn trace(&self) -> Option<TraceContext> {
self.trace
}
pub fn cursors(&mut self) -> Option<Cursors> {
self.cursors.take()
}
pub async fn write_all(&mut self, buf: &[u8]) -> Result<(), Error> {
self.stream.write_all(buf).await
}
pub fn finish(mut self) -> Result<Delivery, Error> {
self.settled = true;
self.stream.finish()?;
Ok(Delivery {
stopped: Some(self.stream.stopped()),
conn: Arc::clone(&self.conn),
})
}
pub fn cancel(mut self) {
self.settled = true;
self.stream.reset(codes::CANCELED);
}
}
impl Drop for OutgoingTransfer {
fn drop(&mut self) {
if !self.settled {
self.stream.reset(codes::CANCELED);
}
}
}
impl AsyncWrite for OutgoingTransfer {
fn poll_write(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<std::io::Result<usize>> {
AsyncWrite::poll_write(Pin::new(&mut self.stream), cx, buf)
}
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
AsyncWrite::poll_flush(Pin::new(&mut self.stream), cx)
}
fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
AsyncWrite::poll_shutdown(Pin::new(&mut self.stream), cx)
}
}
impl futures_io::AsyncWrite for OutgoingTransfer {
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<std::io::Result<usize>> {
<Self as AsyncWrite>::poll_write(self, cx, buf)
}
fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
<Self as AsyncWrite>::poll_flush(self, cx)
}
fn poll_close(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
<Self as AsyncWrite>::poll_shutdown(self, cx)
}
}
impl std::fmt::Debug for OutgoingTransfer {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("OutgoingTransfer")
.field("settled", &self.settled)
.finish_non_exhaustive()
}
}
pub struct Delivery {
stopped: Option<Receipt>,
conn: ConnHandle,
}
impl Delivery {
pub async fn delivered(mut self) -> Result<(), Error> {
let stopped = self.stopped.take().expect("receipt taken only here");
match stopped.await {
Ok(None) => Ok(()),
Ok(Some(code)) => Err(codes::stop_reason(code).into()),
Err(e) => Err(e),
}
}
}
impl Drop for Delivery {
fn drop(&mut self) {
if let Some(receipt) = self.stopped.take()
&& self.conn.parked.park(receipt)
{
self.conn.shared.drain.evict();
}
}
}
impl std::fmt::Debug for Delivery {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Delivery").finish_non_exhaustive()
}
}
pub struct IncomingTransfer {
stream: RecvHalf,
meta: Arc<IncomingMeta>,
conn: ConnHandle,
done: bool,
}
impl IncomingTransfer {
pub(crate) fn new(
stream: RecvHalf,
meta: Arc<IncomingMeta>,
conn: ConnHandle,
) -> IncomingTransfer {
IncomingTransfer {
stream,
meta,
conn,
done: false,
}
}
pub fn meta(&self) -> &IncomingMeta {
&self.meta
}
pub fn reporter(&self) -> Option<Reporter> {
let report_id = self.meta.report_id?;
if self.meta.report.is_empty() {
return None;
}
Some(Reporter::new(
ConnHandle::clone(&self.conn),
report_id,
self.meta.report.clone(),
self.meta.report_mode,
))
}
pub(crate) fn refuse(mut self, code: u64) {
self.done = true;
self.stream.stop(code);
}
pub async fn read_capped(&mut self, max_bytes: usize) -> Result<Vec<u8>, Error> {
let mut out = Vec::new();
let mut chunk = vec![0u8; 64 * 1024];
loop {
match self.stream.read(&mut chunk).await {
Ok(Some(0)) => continue,
Ok(Some(n)) => {
if out.len() + n > max_bytes {
self.done = true;
self.stream.stop(codes::REJECTED);
return Err(Error::LimitExceeded);
}
out.extend_from_slice(&chunk[..n]);
}
Ok(None) => {
self.done = true;
return Ok(out);
}
Err(e) => {
self.done = true;
return Err(e);
}
}
}
}
pub async fn collect(mut self, max_bytes: usize) -> Result<Vec<u8>, Error> {
self.read_capped(max_bytes).await
}
}
impl Drop for IncomingTransfer {
fn drop(&mut self) {
if !self.done {
self.stream.stop(codes::REJECTED);
}
}
}
impl std::fmt::Debug for IncomingTransfer {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("IncomingTransfer")
.field("meta", &self.meta)
.field("done", &self.done)
.finish_non_exhaustive()
}
}
impl AsyncRead for IncomingTransfer {
fn poll_read(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<std::io::Result<()>> {
let before = buf.filled().len();
match AsyncRead::poll_read(Pin::new(&mut self.stream), cx, buf) {
Poll::Ready(Ok(())) => {
if buf.filled().len() == before {
self.done = true;
}
Poll::Ready(Ok(()))
}
Poll::Ready(Err(e)) => {
self.done = true;
Poll::Ready(Err(e))
}
pending => pending,
}
}
}
impl futures_io::AsyncRead for IncomingTransfer {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut [u8],
) -> Poll<std::io::Result<usize>> {
let mut read_buf = ReadBuf::new(buf);
match <Self as AsyncRead>::poll_read(self, cx, &mut read_buf) {
Poll::Ready(Ok(())) => Poll::Ready(Ok(read_buf.filled().len())),
Poll::Ready(Err(e)) => Poll::Ready(Err(e)),
Poll::Pending => Poll::Pending,
}
}
}
struct ReplyHalf {
send: Option<SendHalf>,
conn: ConnHandle,
}
impl Drop for ReplyHalf {
fn drop(&mut self) {
if let Some(send) = self.send.take() {
self.conn.notify(Ctl::ReplyError {
send,
code: ErrorCode::NoReply,
});
}
}
}
pub struct IncomingRequest {
body: Option<IncomingTransfer>,
meta: Arc<IncomingMeta>,
reply: ReplyHalf,
}
impl IncomingRequest {
pub(crate) fn new(body: IncomingTransfer, send: SendHalf, conn: ConnHandle) -> IncomingRequest {
IncomingRequest {
meta: Arc::clone(&body.meta),
body: Some(body),
reply: ReplyHalf {
send: Some(send),
conn,
},
}
}
pub fn meta(&self) -> &IncomingMeta {
&self.meta
}
pub fn body(&mut self) -> &mut IncomingTransfer {
self.body
.as_mut()
.expect("the request body was detached by take_body")
}
pub fn take_body(&mut self) -> IncomingTransfer {
self.body
.take()
.expect("the request body is detached at most once")
}
pub fn canceled(&self) -> impl Future<Output = ()> + Send + use<> {
let stopped = self.reply.send.as_ref().map(|s| s.stopped());
async move {
match stopped {
Some(fut) => {
let _ = fut.await;
}
None => std::future::pending().await,
}
}
}
pub async fn reply(mut self, meta: TransferMeta) -> Result<OutgoingTransfer, Error> {
drop(self.body.take());
let mut send = self
.reply
.send
.take()
.expect("the send half is taken exactly once, by this method");
let meta = match (meta.trace, self.meta.trace) {
(None, Some(inherited)) => meta.with_trace(inherited),
_ => meta,
};
let (header, trace, cursors) =
outgoing_header(&self.reply.conn, None, &meta, self.meta.tracestate.clone())?;
write_data_preamble(&mut send, &header).await?;
Ok(OutgoingTransfer::new(
send,
trace,
Arc::clone(&self.reply.conn),
cursors,
))
}
pub async fn refuse(self, code: ErrorCode) {
self.refuse_coded(code, codes::REJECTED).await;
}
pub(crate) async fn refuse_coded(mut self, code: ErrorCode, stop: u64) {
if let Some(body) = self.body.take() {
body.refuse(stop);
}
let Some(mut send) = self.reply.send.take() else {
return;
};
if let Err(e) = write_error_frame(&mut send, code).await {
tracing::debug!(error = %e, "failed to refuse an exchange");
}
}
}
impl std::fmt::Debug for IncomingRequest {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("IncomingRequest")
.field("meta", self.meta())
.finish_non_exhaustive()
}
}
fn data_frame(header: &DataHeader) -> (Vec<u8>, usize) {
const HEADER_HINT: usize = 160;
let pre = weida_protocol::MAX_PREAMBLE_LEN;
let mut buf = Vec::with_capacity(pre + HEADER_HINT);
buf.resize(pre, 0);
header.encode_into(&mut buf);
let header_len = buf.len() - pre;
let (bytes, len) = weida_protocol::preamble_bytes(FrameKind::Data, header_len as u64);
let start = pre - len;
buf[start..pre].copy_from_slice(&bytes[..len]);
(buf, start)
}
pub(crate) async fn write_data_preamble(
stream: &mut SendHalf,
header: &DataHeader,
) -> Result<(), Error> {
let (buf, start) = data_frame(header);
stream.write_all(&buf[start..]).await
}
pub struct ReplyStream {
recv: Option<RecvHalf>,
conn: ConnHandle,
}
impl ReplyStream {
pub(crate) fn new(recv: RecvHalf, conn: ConnHandle) -> ReplyStream {
ReplyStream {
recv: Some(recv),
conn,
}
}
pub async fn recv(mut self) -> Result<IncomingTransfer, Error> {
let mut recv = self.recv.take().expect("the receiver is taken once");
let (preamble, header) = read_frame(&mut recv, self.conn.limits.max_header_bytes)
.await
.map_err(indeterminate_on_loss)?;
match preamble.kind {
FrameKind::Data => {
let header = DataHeader::decode(&header)?;
let meta = Arc::new(IncomingMeta::from_header(&header, self.conn.peer.clone()));
Ok(IncomingTransfer::new(
recv,
meta,
ConnHandle::clone(&self.conn),
))
}
FrameKind::Error => {
let header = ErrorHeader::decode(&header)?;
Err(header.error_code().map_or_else(
|| Error::Transport(format!("peer reported error code {}", header.code)),
Error::from,
))
}
other => Err(Error::Protocol(format!(
"{other} is not legal on the reply half of an exchange"
))),
}
}
}
fn indeterminate_on_loss(e: Error) -> Error {
match e {
Error::ConnectionLost(_) => Error::Indeterminate,
other => other,
}
}
impl Drop for ReplyStream {
fn drop(&mut self) {
if let Some(mut recv) = self.recv.take() {
recv.stop(codes::CANCELED);
}
}
}
impl std::fmt::Debug for ReplyStream {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ReplyStream")
.field("awaiting", &self.recv.is_some())
.finish_non_exhaustive()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn transfer_meta_builders_compose() {
let trace = new_trace();
let meta = TransferMeta::default()
.with_content_type("text/plain")
.with_content_len(7)
.with_trace(trace);
assert_eq!(meta.content_type.as_deref(), Some("text/plain"));
assert_eq!(meta.content_len, Some(7));
assert_eq!(meta.trace, Some(trace));
}
#[test]
fn generated_trace_contexts_are_valid_and_distinct() {
let a = new_trace();
let b = new_trace();
assert_ne!(a.trace_id, b.trace_id);
assert!(a.is_sampled());
assert_eq!(
TraceContext::parse_traceparent(&a.to_traceparent()).unwrap(),
a
);
}
#[test]
fn a_header_carries_no_trace_context_unless_one_was_supplied() {
let meta = TransferMeta::default().with_content_len(3);
let (header, trace) = data_header(Some("/transform"), &meta, None, None).unwrap();
assert_eq!(header.endpoint.as_deref(), Some("/transform"));
assert_eq!(header.content_len, Some(3));
assert_eq!(trace, None, "nothing mints a context");
assert_eq!(header.traceparent, None, "and nothing writes one");
assert_eq!(DataHeader::decode(&header.encode()).unwrap(), header);
let supplied = new_trace();
let (header, trace) = data_header(
Some("/transform"),
&meta.clone().with_trace(supplied),
None,
None,
)
.unwrap();
assert_eq!(trace, Some(supplied));
assert_eq!(
header.traceparent.as_deref(),
Some(&*supplied.to_traceparent()),
"a supplied context is propagated verbatim"
);
assert_eq!(DataHeader::decode(&header.encode()).unwrap(), header);
}
#[test]
fn a_trace_context_costs_fifty_eight_bytes_of_frame() {
let meta = TransferMeta::default().with_content_len(64);
let (bare, _) = data_header(Some("/t"), &meta, None, None).unwrap();
let (traced, _) = data_header(
Some("/t"),
&meta.clone().with_trace(new_trace()),
None,
None,
)
.unwrap();
assert_eq!(
traced.encode().len() - bare.encode().len(),
58,
"1 byte of key, 2 of the tstr prefix, 55 of the value"
);
}
#[test]
fn reply_headers_carry_no_endpoint_but_keep_tracestate() {
let (header, _) = data_header(
None,
&TransferMeta::default(),
Some("vendor=x".into()),
None,
)
.unwrap();
assert_eq!(header.endpoint, None);
assert_eq!(header.tracestate.as_deref(), Some("vendor=x"));
assert_eq!(DataHeader::decode(&header.encode()).unwrap(), header);
}
#[test]
fn a_data_frame_is_built_in_one_buffer_and_is_byte_identical() {
use weida_protocol::{encode_frame, header::limits::MAX_ENDPOINT_BYTES};
for path_len in [1usize, 60, 61, 62, 200, MAX_ENDPOINT_BYTES] {
let path = format!("/{}", "a".repeat(path_len - 1));
let meta = TransferMeta::default()
.with_content_len(1 << 20)
.with_trace(new_trace());
let (header, _) = data_header(Some(&path), &meta, None, None).unwrap();
let (buf, start) = data_frame(&header);
let want = encode_frame(FrameKind::Data, &header.encode());
assert_eq!(
&buf[start..],
want.as_slice(),
"a {path_len}-byte path frames differently in one buffer"
);
let (preamble, used) =
weida_protocol::parse_preamble(&buf[start..], 64 * 1024).expect("a preamble");
assert_eq!(preamble.kind, FrameKind::Data);
assert_eq!(
DataHeader::decode(&buf[start + used..]).unwrap(),
header,
"the header the length field points at"
);
}
}
#[test]
fn a_report_order_is_sorted_deduplicated_and_capped() {
let accepted = CursorLevel::Known(Acknowledgement::Accepted);
let stored = CursorLevel::Known(Acknowledgement::Stored);
let app = CursorLevel::Application(17);
let meta = TransferMeta::default().with_report([app, stored, accepted, stored]);
assert_eq!(meta.report, vec![accepted, stored, app]);
let (header, _) = data_header(Some("/t"), &meta, None, Some(1)).unwrap();
assert_eq!(header.report_id, Some(1));
assert_eq!(DataHeader::decode(&header.encode()).unwrap(), header);
let over = TransferMeta::default().with_report(
(0..=MAX_REPORT_LEVELS as u64)
.map(|i| CursorLevel::Application(CursorLevel::APPLICATION_FLOOR + i)),
);
assert!(matches!(
data_header(Some("/t"), &over, None, Some(1)),
Err(Error::LimitExceeded)
));
}
#[test]
fn incoming_meta_ignores_a_malformed_traceparent() {
let mut header = DataHeader::addressed("/x");
header.traceparent = Some("not-a-traceparent".into());
let meta = IncomingMeta::from_header(&header, None);
assert!(meta.trace.is_none());
let good = new_trace();
header.traceparent = Some(good.to_traceparent());
assert_eq!(IncomingMeta::from_header(&header, None).trace, Some(good));
}
}