From b73496e703f1b7a0d058cd4ae8b0a75f3d566017 Mon Sep 17 00:00:00 2001 From: Luke Curley Date: Mon, 28 Sep 2026 19:08:10 -0700 Subject: [PATCH 1/8] quest: claim quest/m1/rs2ts/varint-codec Co-Authored-By: Claude Opus 5.5 From a508c2dd078971b043eaeace674cc71383360318 Mon Sep 17 00:00:00 2001 From: Luke Curley Date: Mon, 28 Sep 2026 19:16:11 -0700 Subject: [PATCH 2/8] bench(net): measure the wire codec on its own Co-Authored-By: Claude Opus 5.5 --- rs/moq-net/Cargo.toml | 6 ++ rs/moq-net/benches/codec.rs | 72 +++++++++++++++++ rs/moq-net/src/fuzz.rs | 152 ++++++++++++++++++++++++++++++++++++ 3 files changed, 230 insertions(+) create mode 100644 rs/moq-net/benches/codec.rs diff --git a/rs/moq-net/Cargo.toml b/rs/moq-net/Cargo.toml index 967afb259d..9d30aee10d 100644 --- a/rs/moq-net/Cargo.toml +++ b/rs/moq-net/Cargo.toml @@ -75,6 +75,12 @@ name = "announce" harness = false required-features = ["fuzz"] +# Reaches the private message codecs through the hidden `fuzz` module. +[[bench]] +name = "codec" +harness = false +required-features = ["fuzz"] + [[bench]] name = "group" harness = false diff --git a/rs/moq-net/benches/codec.rs b/rs/moq-net/benches/codec.rs new file mode 100644 index 0000000000..39f6dad934 --- /dev/null +++ b/rs/moq-net/benches/codec.rs @@ -0,0 +1,72 @@ +//! The wire codec on its own: a fixed mix of moq-lite and moq-transport messages, and raw +//! varints in each wire form. Every message and frame header pays this cost. +//! +//! Run with `cargo bench -p moq-net --features fuzz --bench codec`. + +use criterion::{BenchmarkId, Criterion, Throughput, criterion_group, criterion_main}; +use moq_net::fuzz::{Messages, decode_varints, encode_varints}; + +type Encode = fn(&Messages, &mut Vec); +type Decode = fn(&Messages, &[u8]); + +/// Varints spanning the length classes of both wire forms. +fn varints() -> Vec { + (0..1_024u64) + .map(|n| match n % 4 { + 0 => n % 60, + 1 => 1_000 + n, + 2 => 1_000_000 + n, + _ => (1 << 40) + n, + }) + .collect() +} + +fn bench(c: &mut Criterion) { + let messages = Messages::default(); + + let mut group = c.benchmark_group("codec_messages"); + let protocols: [(&str, Encode, Decode); 2] = [ + ("lite", Messages::encode_lite, Messages::decode_lite), + ("ietf", Messages::encode_ietf, Messages::decode_ietf), + ]; + for (name, encode, decode) in protocols { + let mut encoded = Vec::new(); + encode(&messages, &mut encoded); + group.throughput(Throughput::Bytes(encoded.len() as u64)); + + group.bench_function(BenchmarkId::new("encode", name), |b| { + let mut out = Vec::with_capacity(encoded.len()); + b.iter(|| { + out.clear(); + encode(&messages, &mut out); + }); + }); + group.bench_function(BenchmarkId::new("decode", name), |b| { + b.iter(|| decode(&messages, &encoded)) + }); + } + group.finish(); + + let values = varints(); + let mut group = c.benchmark_group("codec_varint"); + group.throughput(Throughput::Elements(values.len() as u64)); + for (name, ietf) in [("quic", false), ("leading_ones", true)] { + let mut encoded = Vec::new(); + encode_varints(&values, ietf, &mut encoded); + + group.bench_function(BenchmarkId::new("encode", name), |b| { + let mut out = Vec::with_capacity(encoded.len()); + b.iter(|| { + out.clear(); + encode_varints(&values, ietf, &mut out); + }); + }); + group.bench_function(BenchmarkId::new("decode", name), |b| { + b.iter(|| decode_varints(&encoded, ietf)) + }); + } + group.finish(); +} + +criterion_group!(benches, bench); +criterion_main!(benches); diff --git a/rs/moq-net/src/fuzz.rs b/rs/moq-net/src/fuzz.rs index 5fc9f793b6..cfb59f1bad 100644 --- a/rs/moq-net/src/fuzz.rs +++ b/rs/moq-net/src/fuzz.rs @@ -370,6 +370,158 @@ pub fn varint(data: &[u8]) -> bool { true } +/// A fixed mix of control and data-stream messages, for the codec benchmark. +/// +/// Encoding appends every message to a buffer; decoding reads them back in the same +/// order. The mix leans on the messages every subscription pays for. +pub struct Messages { + lite_subscribe: lite::Subscribe<'static>, + lite_update: lite::SubscribeUpdate, + lite_start: lite::SubscribeResponse, + lite_info: lite::TrackInfo, + lite_group: lite::Group, + ietf_subscribe: ietf::Subscribe<'static>, + ietf_ok: ietf::SubscribeOk, + ietf_group: ietf::GroupHeader, +} + +impl Default for Messages { + fn default() -> Self { + use std::time::Duration; + + Self { + lite_subscribe: lite::Subscribe { + id: 7, + broadcast: Path::new("room/alice"), + track: "video".into(), + priority: 3, + max_age: Duration::from_millis(500), + start_group: Some(1_000), + end_group: None, + start_frame: 0, + end_frame: None, + }, + lite_update: lite::SubscribeUpdate { + priority: 4, + max_age: Duration::from_millis(500), + start_group: Some(1_000), + end_group: Some(2_000), + start_frame: 0, + end_frame: None, + }, + lite_start: lite::SubscribeResponse::Start(lite::SubscribeStart { group: 1_234 }), + lite_info: lite::TrackInfo { + priority: 1, + max_age: Some(Duration::from_secs(10)), + timescale: crate::Timescale::MICRO, + }, + lite_group: lite::Group { + subscribe: 7, + sequence: 123_456, + frame_start: 0, + }, + ietf_subscribe: ietf::Subscribe { + request_id: ietf::RequestId(2), + track_namespace: Path::new("room/alice"), + track_name: "video".into(), + subscriber_priority: 128, + group_order: ietf::GroupOrder::Descending, + filter: ietf::Filter::NextObject, + fill: None, + properties_wanted: false, + }, + // Draft-17+ carries the request id in the control message framing instead. + ietf_ok: ietf::SubscribeOk { + request_id: None, + track_alias: 5, + largest: Some(ietf::Location { + group: 1_000, + object: 3, + }), + properties: Default::default(), + }, + ietf_group: ietf::GroupHeader { + track_alias: 5, + group_id: 1_000, + sub_group_id: 0, + publisher_priority: 128, + flags: Default::default(), + }, + } + } +} + +/// The moq-lite version the [`Messages`] mix is encoded at. +const BENCH_LITE: lite::Version = lite::Version::Lite06; + +/// The moq-transport draft the [`Messages`] mix is encoded at: a leading-ones varint draft. +const BENCH_IETF: ietf::Version = ietf::Version::Draft20; + +impl Messages { + /// Append the moq-lite messages to `out`. + pub fn encode_lite(&self, out: &mut Vec) { + let v = BENCH_LITE; + self.lite_subscribe.encode(out, v).unwrap(); + self.lite_update.encode(out, v).unwrap(); + self.lite_start.encode(out, v).unwrap(); + self.lite_info.encode(out, v).unwrap(); + self.lite_group.encode(out, v).unwrap(); + } + + /// Decode what [`Self::encode_lite`] wrote. + pub fn decode_lite(&self, mut data: &[u8]) { + let v = BENCH_LITE; + lite::Subscribe::decode(&mut data, v).unwrap(); + lite::SubscribeUpdate::decode(&mut data, v).unwrap(); + lite::SubscribeResponse::decode(&mut data, v).unwrap(); + lite::TrackInfo::decode(&mut data, v).unwrap(); + lite::Group::decode(&mut data, v).unwrap(); + assert!(data.is_empty()); + } + + /// Append the moq-transport messages to `out`. + pub fn encode_ietf(&self, out: &mut Vec) { + let v = BENCH_IETF; + self.ietf_subscribe.encode(out, v).unwrap(); + self.ietf_ok.encode(out, v).unwrap(); + self.ietf_group.encode(out, v).unwrap(); + } + + /// Decode what [`Self::encode_ietf`] wrote. + pub fn decode_ietf(&self, mut data: &[u8]) { + let v = BENCH_IETF; + ietf::Subscribe::decode(&mut data, v).unwrap(); + ietf::SubscribeOk::decode(&mut data, v).unwrap(); + ietf::GroupHeader::decode(&mut data, v).unwrap(); + assert!(data.is_empty()); + } +} + +/// Append `values` as varints in the wire form of moq-lite (`ietf == false`) or of a +/// leading-ones moq-transport draft. +pub fn encode_varints(values: &[u64], ietf: bool, out: &mut Vec) { + for value in values { + let value = VarInt::try_from(*value).unwrap(); + match ietf { + false => value.encode(out, BENCH_LITE).unwrap(), + true => value.encode(out, BENCH_IETF).unwrap(), + } + } +} + +/// Sum the varints [`encode_varints`] wrote. +pub fn decode_varints(mut data: &[u8], ietf: bool) -> u64 { + let mut sum = 0u64; + while !data.is_empty() { + let value = match ietf { + false => VarInt::decode(&mut data, BENCH_LITE).unwrap(), + true => VarInt::decode(&mut data, BENCH_IETF).unwrap(), + }; + sum = sum.wrapping_add(value.into_inner()); + } + sum +} + /// Exercise the [`Path`] invariants against arbitrary text. /// /// The input is UTF-8, split at the first newline into a target path and a base. The From ee9e41d0d4ff95e1161549db279d3e42c47396f2 Mon Sep 17 00:00:00 2001 From: Luke Curley Date: Mon, 28 Sep 2026 19:59:01 -0700 Subject: [PATCH 3/8] refactor(net)!: encode through a concrete VarInt codec VarInt is now the only integer with Encode/Decode, and it holds the full u64 range. Messages read from a slice-based Decoder and write to a Vec-backed Encoder instead of generic Buf/BufMut traits implemented on u64, usize, bool, String, Option and Vec. Parameters are Vec-backed. Co-Authored-By: Claude Opus 5.5 --- rs/hang/src/container/frame.rs | 5 +- rs/moq-archive/src/segment.rs | 3 +- rs/moq-loc/src/lib.rs | 8 +- rs/moq-net/src/client.rs | 19 +- rs/moq-net/src/coding/decode.rs | 200 +++--- rs/moq-net/src/coding/encode.rs | 219 +++---- rs/moq-net/src/coding/mod.rs | 23 +- rs/moq-net/src/coding/reader.rs | 28 +- rs/moq-net/src/coding/size.rs | 212 ------ rs/moq-net/src/coding/varint.rs | 722 ++++++++------------- rs/moq-net/src/coding/version.rs | 38 +- rs/moq-net/src/coding/writer.rs | 35 +- rs/moq-net/src/fuzz.rs | 288 ++++---- rs/moq-net/src/ietf/adapter.rs | 113 ++-- rs/moq-net/src/ietf/cluster.rs | 96 +-- rs/moq-net/src/ietf/fetch.rs | 136 ++-- rs/moq-net/src/ietf/filter.rs | 215 +++--- rs/moq-net/src/ietf/goaway.rs | 74 ++- rs/moq-net/src/ietf/group.rs | 193 +++--- rs/moq-net/src/ietf/location.rs | 14 +- rs/moq-net/src/ietf/message.rs | 90 ++- rs/moq-net/src/ietf/mod.rs | 2 +- rs/moq-net/src/ietf/namespace.rs | 36 +- rs/moq-net/src/ietf/parameters.rs | 448 ++++++------- rs/moq-net/src/ietf/properties.rs | 180 +++-- rs/moq-net/src/ietf/publish.rs | 121 ++-- rs/moq-net/src/ietf/publish_namespace.rs | 106 +-- rs/moq-net/src/ietf/publisher.rs | 308 +++++---- rs/moq-net/src/ietf/request.rs | 46 +- rs/moq-net/src/ietf/session.rs | 74 +-- rs/moq-net/src/ietf/subscribe.rs | 159 ++--- rs/moq-net/src/ietf/subscribe_namespace.rs | 125 ++-- rs/moq-net/src/ietf/subscriber.rs | 254 ++++---- rs/moq-net/src/ietf/token.rs | 48 +- rs/moq-net/src/ietf/track.rs | 44 +- rs/moq-net/src/ietf/version.rs | 7 +- rs/moq-net/src/lite/announce.rs | 293 +++++---- rs/moq-net/src/lite/compress.rs | 39 +- rs/moq-net/src/lite/datagram.rs | 57 +- rs/moq-net/src/lite/fetch.rs | 45 +- rs/moq-net/src/lite/goaway.rs | 67 +- rs/moq-net/src/lite/group.rs | 32 +- rs/moq-net/src/lite/info.rs | 8 +- rs/moq-net/src/lite/message.rs | 101 +-- rs/moq-net/src/lite/parameters.rs | 55 +- rs/moq-net/src/lite/probe.rs | 19 +- rs/moq-net/src/lite/publisher.rs | 11 +- rs/moq-net/src/lite/setup.rs | 69 +- rs/moq-net/src/lite/stream.rs | 16 +- rs/moq-net/src/lite/subscribe.rs | 263 ++++---- rs/moq-net/src/lite/subscriber.rs | 42 +- rs/moq-net/src/lite/track.rs | 87 +-- rs/moq-net/src/model/origin.rs | 44 +- rs/moq-net/src/model/time.rs | 18 +- rs/moq-net/src/model/track.rs | 5 +- rs/moq-net/src/path/mod.rs | 39 +- rs/moq-net/src/server.rs | 27 +- rs/moq-net/src/setup.rs | 136 ++-- rs/moq-net/src/version.rs | 4 +- 59 files changed, 2970 insertions(+), 3196 deletions(-) delete mode 100644 rs/moq-net/src/coding/size.rs diff --git a/rs/hang/src/container/frame.rs b/rs/hang/src/container/frame.rs index 77c7ae8955..b971af1fc5 100644 --- a/rs/hang/src/container/frame.rs +++ b/rs/hang/src/container/frame.rs @@ -120,8 +120,9 @@ impl Frame { /// Write the VarInt timestamp prefix, normalized to [`TIMESCALE`]. fn encode_header(&self, buf: &mut impl BufMut) -> Result<(), Error> { let timestamp = self.timestamp.convert(TIMESCALE)?; - let value = VarInt::try_from(timestamp.value()).map_err(moq_net::Error::from)?; - value.encode_quic(buf).map_err(moq_net::Error::from)?; + VarInt::from(timestamp.value()) + .encode_quic(buf) + .map_err(moq_net::Error::from)?; Ok(()) } diff --git a/rs/moq-archive/src/segment.rs b/rs/moq-archive/src/segment.rs index c2fd68c5aa..ceacc54238 100644 --- a/rs/moq-archive/src/segment.rs +++ b/rs/moq-archive/src/segment.rs @@ -203,8 +203,7 @@ fn validate(groups: &[Group]) -> Result<()> { } fn write_varint(buf: &mut impl BufMut, value: u64) -> Result<()> { - let value = VarInt::try_from(value).map_err(|_| Error::Overflow)?; - value.encode_quic(buf).map_err(|_| Error::Overflow) + VarInt::from(value).encode_quic(buf).map_err(|_| Error::Overflow) } fn read_varint(buf: &mut impl Buf) -> Result { diff --git a/rs/moq-loc/src/lib.rs b/rs/moq-loc/src/lib.rs index af241919a5..145eeca3be 100644 --- a/rs/moq-loc/src/lib.rs +++ b/rs/moq-loc/src/lib.rs @@ -167,11 +167,11 @@ pub fn decode(mut buf: Bytes) -> Result { /// catalog timescale to interpret `timestamp`. pub fn encode(timestamp: u64, payload: &[u8]) -> Result { let mut props = BytesMut::with_capacity(16); - VarInt::try_from(PROP_TIMESTAMP)?.encode_quic(&mut props)?; - VarInt::try_from(timestamp)?.encode_quic(&mut props)?; + VarInt::from(PROP_TIMESTAMP).encode_quic(&mut props)?; + VarInt::from(timestamp).encode_quic(&mut props)?; let mut out = BytesMut::with_capacity(props.len() + payload.len() + 8); - VarInt::try_from(props.len() as u64)?.encode_quic(&mut out)?; + VarInt::from(props.len()).encode_quic(&mut out)?; out.extend_from_slice(&props); out.extend_from_slice(payload); @@ -184,7 +184,7 @@ mod tests { /// Test helper: write a u64 as a QUIC varint into `buf`. fn write_varint(buf: &mut BytesMut, value: u64) { - VarInt::try_from(value).unwrap().encode_quic(buf).unwrap(); + VarInt::from(value).encode_quic(buf).unwrap(); } #[test] diff --git a/rs/moq-net/src/client.rs b/rs/moq-net/src/client.rs index 47d1aecd88..31ddad3ab6 100644 --- a/rs/moq-net/src/client.rs +++ b/rs/moq-net/src/client.rs @@ -353,7 +353,7 @@ impl Client { stream.writer.encode(&client).await?; - let mut server: setup::Server = stream.reader.decode().await?; + let server: setup::Server = stream.reader.decode().await?; let version = supported .iter() @@ -387,7 +387,7 @@ impl Client { Version::Ietf(v) => { // Decode the parameters to get the initial request ID and what the server // requires of us. - let parameters = ietf::Parameters::decode(&mut server.parameters, v)?; + let (parameters, _) = ietf::Parameters::decode_slice(&server.parameters, v)?; let request_id_max = parameters .get_varint(ietf::ParameterVarInt::MaxRequestId) .map(ietf::RequestId); @@ -652,13 +652,17 @@ mod tests { parameters: Bytes::new(), }; server - .encode(&mut encoded, Version::Ietf(ietf::Version::Draft14)) + .encode( + &mut crate::coding::Encoder::new(&mut encoded, (Version::Ietf(ietf::Version::Draft14)).into()), + Version::Ietf(ietf::Version::Draft14), + ) .unwrap(); // Add a setup-stream SessionInfo frame using the negotiated Lite version. let info = lite::SessionInfo { bitrate: Some(1) }; let lite_v = lite::Version::try_from(negotiated).unwrap(); - info.encode(&mut encoded, lite_v).unwrap(); + info.encode(&mut crate::coding::Encoder::new(&mut encoded, lite_v.into()), lite_v) + .unwrap(); encoded } @@ -684,7 +688,12 @@ mod tests { // Verify the client setup was encoded using Draft14 framing (ALPN_LITE fallback path). let mut setup_bytes = Bytes::from(fake.control_writes()); - let setup = setup::Client::decode(&mut setup_bytes, Version::Ietf(ietf::Version::Draft14)).unwrap(); + let setup = crate::coding::decode_buf( + &mut setup_bytes, + Version::Ietf(ietf::Version::Draft14), + setup::Client::decode, + ) + .unwrap(); let advertised: Vec = setup.versions.iter().map(|v| Version::try_from(*v).unwrap()).collect(); assert_eq!( advertised, diff --git a/rs/moq-net/src/coding/decode.rs b/rs/moq-net/src/coding/decode.rs index 13579e7330..9d95ececaf 100644 --- a/rs/moq-net/src/coding/decode.rs +++ b/rs/moq-net/src/coding/decode.rs @@ -1,12 +1,24 @@ -use std::{borrow::Cow, string::FromUtf8Error}; +use std::string::FromUtf8Error; use thiserror::Error; -/// Read the from the buffer using the given version. +use super::{BoundsExceeded, Form, VarInt}; + +/// Read the value from a [`Decoder`] using the given version. /// /// If [DecodeError::Short] is returned, the caller should try again with more data. pub trait Decode: Sized { - /// Decode the value from the given buffer. - fn decode(buf: &mut B, version: V) -> Result; + /// Decode the value from the front of the decoder. + fn decode(r: &mut Decoder<'_>, version: V) -> Result; + + /// Decode the value from the front of `buf`, returning it and the bytes it took. + fn decode_slice(buf: &[u8], version: V) -> Result<(Self, usize), DecodeError> + where + V: Into
+ Copy, + { + let mut r = Decoder::new(buf, version.into()); + let value = Self::decode(&mut r, version)?; + Ok((value, buf.len() - r.remaining())) + } } /// A decode error. @@ -41,7 +53,7 @@ pub enum DecodeError { #[error("too many")] TooMany, - /// An integer was too large for the QUIC varint range. + /// An integer was too large for the field it was read into. #[error("bounds exceeded")] BoundsExceeded, @@ -83,119 +95,121 @@ pub enum DecodeError { Version, } -impl Decode for bool { - fn decode(r: &mut R, version: V) -> Result { - match u8::decode(r, version)? { - 0 => Ok(false), - 1 => Ok(true), - _ => Err(DecodeError::InvalidValue), - } +impl From for DecodeError { + fn from(_: BoundsExceeded) -> Self { + Self::BoundsExceeded } } -impl Decode for u8 { - fn decode(r: &mut R, _: V) -> Result { - match r.has_remaining() { - true => Ok(r.get_u8()), - false => Err(DecodeError::Short), - } - } +/// Reads wire primitives from the front of a byte slice. +/// +/// A read either consumes exactly what it returns or fails and consumes nothing, so a +/// [`DecodeError::Short`] can be retried once more bytes arrive. +#[derive(Debug, Clone)] +pub struct Decoder<'a> { + buf: &'a [u8], + form: Form, } -impl Decode for u16 { - fn decode(r: &mut R, _: V) -> Result { - match r.remaining() >= 2 { - true => Ok(r.get_u16()), - false => Err(DecodeError::Short), - } +impl<'a> Decoder<'a> { + /// Read `buf`, with varints in the given form. + pub fn new(buf: &'a [u8], form: Form) -> Self { + Self { buf, form } } -} -impl Decode for String -where - usize: Decode, -{ - /// Decode a string with a varint length prefix. - fn decode(r: &mut R, version: V) -> Result { - let v = Vec::::decode(r, version)?; - let str = String::from_utf8(v)?; + /// The varint form this decoder reads. + pub fn form(&self) -> Form { + self.form + } - Ok(str) + /// The number of unread bytes. + pub fn remaining(&self) -> usize { + self.buf.len() } -} -impl Decode for Vec -where - usize: Decode, -{ - fn decode(buf: &mut B, version: V) -> Result { - let size = usize::decode(buf, version)?; + /// Whether every byte has been read. + pub fn is_empty(&self) -> bool { + self.buf.is_empty() + } - if buf.remaining() < size { + /// Read `len` raw bytes. + pub fn slice(&mut self, len: usize) -> Result<&'a [u8], DecodeError> { + let Some((head, rest)) = self.buf.split_at_checked(len) else { return Err(DecodeError::Short); - } + }; + self.buf = rest; + Ok(head) + } - let bytes = buf.copy_to_bytes(size); - Ok(bytes.to_vec()) + /// Read every remaining byte. + pub fn rest(&mut self) -> &'a [u8] { + std::mem::take(&mut self.buf) } -} -impl Decode for i8 { - fn decode(r: &mut R, _: V) -> Result { - if !r.has_remaining() { - return Err(DecodeError::Short); - } + /// Split off the next `len` bytes as their own decoder, e.g. a size-prefixed body. + pub fn sub(&mut self, len: usize) -> Result { + Ok(Self::new(self.slice(len)?, self.form)) + } - // This is not the usual way of encoding negative numbers. - // i8 doesn't exist in the draft, but we use it instead of u8 for priority. - // A default of 0 is more ergonomic for the user than a default of 128. - Ok(((r.get_u8() as i16) - 128) as i8) + /// Read a single byte. + pub fn u8(&mut self) -> Result { + Ok(self.slice(1)?[0]) } -} -impl Decode for bytes::Bytes -where - usize: Decode, -{ - fn decode(r: &mut R, version: V) -> Result { - let len = usize::decode(r, version)?; - if r.remaining() < len { - return Err(DecodeError::Short); + /// Read a big-endian `u16`. + pub fn u16(&mut self) -> Result { + let b = self.slice(2)?; + Ok(u16::from_be_bytes([b[0], b[1]])) + } + + /// Read a byte that must be 0 or 1. + pub fn bool(&mut self) -> Result { + match self.u8()? { + 0 => Ok(false), + 1 => Ok(true), + _ => Err(DecodeError::InvalidValue), } - let bytes = r.copy_to_bytes(len); - Ok(bytes) } -} -// TODO Support borrowed strings. -impl Decode for Cow<'_, str> -where - usize: Decode, -{ - fn decode(r: &mut R, version: V) -> Result { - let s = String::decode(r, version)?; - Ok(Cow::Owned(s)) + /// Read a varint. + pub fn varint(&mut self) -> Result { + let (value, len) = VarInt::decode_form(self.buf, self.form)?; + self.buf = &self.buf[len..]; + Ok(value) } -} -impl Decode for Option -where - u64: Decode, -{ - fn decode(r: &mut R, version: V) -> Result { - match u64::decode(r, version)? { - 0 => Ok(None), - value => Ok(Some(value - 1)), - } + /// Read an optional varint: 0 is `None`, and `n + 1` is `Some(n)`. + pub fn varint_opt(&mut self) -> Result, DecodeError> { + Ok(self.varint()?.into_inner().checked_sub(1)) + } + + /// Read a varint length, then that many raw bytes. + pub fn bytes(&mut self) -> Result<&'a [u8], DecodeError> { + let start = self.buf; + let len = usize::try_from(self.varint()?)?; + self.slice(len).inspect_err(|_| self.buf = start) + } + + /// Read a varint length, then that many bytes of UTF-8. + pub fn string(&mut self) -> Result { + Ok(String::from_utf8(self.bytes()?.to_vec())?) } } -impl Decode for std::time::Duration -where - u64: Decode, -{ - fn decode(r: &mut R, version: V) -> Result { - let value = u64::decode(r, version)?; - Ok(Self::from_millis(value)) +#[cfg(test)] +mod tests { + use super::*; + + /// A short read must leave the decoder where it was, or a retry with more bytes + /// would start mid-value. + #[test] + fn short_consumes_nothing() { + let mut r = Decoder::new(&[0x05, b'a', b'b'], Form::Quic); + assert!(matches!(r.bytes(), Err(DecodeError::Short))); + assert_eq!(r.remaining(), 3); + + let mut r = Decoder::new(&[0x40], Form::Quic); + assert!(matches!(r.varint(), Err(DecodeError::Short))); + assert_eq!(r.remaining(), 1); } } diff --git a/rs/moq-net/src/coding/encode.rs b/rs/moq-net/src/coding/encode.rs index 04e7560765..6045fc835c 100644 --- a/rs/moq-net/src/coding/encode.rs +++ b/rs/moq-net/src/coding/encode.rs @@ -1,8 +1,6 @@ -use std::{borrow::Cow, sync::Arc}; +use bytes::Bytes; -use bytes::{Bytes, BytesMut}; - -use super::BoundsExceeded; +use super::{BoundsExceeded, Form, VarInt}; /// An error that occurs during encoding. #[derive(thiserror::Error, Debug, Clone)] @@ -37,161 +35,130 @@ impl From for EncodeError { } } -/// Check that the writer has enough remaining capacity. -fn check_remaining(w: &impl bytes::BufMut, needed: usize) -> Result<(), EncodeError> { - if w.remaining_mut() < needed { - return Err(EncodeError::Short); +/// Write the value to an [`Encoder`] using the given version. +pub trait Encode { + /// Encode the value to the given encoder. + fn encode(&self, w: &mut Encoder<'_>, version: V) -> Result<(), EncodeError>; + + /// Encode the value into a fresh [Bytes] buffer. + fn encode_bytes(&self, version: V) -> Result + where + V: Into + Copy, + { + let mut buf = Vec::new(); + self.encode(&mut Encoder::new(&mut buf, version.into()), version)?; + Ok(buf.into()) } - Ok(()) } -/// Write the value to the buffer using the given version. -pub trait Encode: Sized { - /// Encode the value to the given writer. - fn encode(&self, w: &mut W, version: V) -> Result<(), EncodeError>; +/// Appends wire primitives to a byte buffer. +/// +/// The buffer grows as needed, so only a value the wire cannot express fails. +#[derive(Debug)] +pub struct Encoder<'a> { + buf: &'a mut Vec, + form: Form, +} - /// Encode the value into a [Bytes] buffer. - /// - /// NOTE: This will allocate. - fn encode_bytes(&self, v: V) -> Result { - let mut buf = BytesMut::new(); - self.encode(&mut buf, v)?; - Ok(buf.freeze()) +impl<'a> Encoder<'a> { + /// Append to `buf`, with varints in the given form. + pub fn new(buf: &'a mut Vec, form: Form) -> Self { + Self { buf, form } } -} -impl Encode for bool { - fn encode(&self, w: &mut W, _: V) -> Result<(), EncodeError> { - check_remaining(&*w, 1)?; - w.put_u8(*self as u8); - Ok(()) + /// The varint form this encoder writes. + pub fn form(&self) -> Form { + self.form } -} -impl Encode for u8 { - fn encode(&self, w: &mut W, _: V) -> Result<(), EncodeError> { - check_remaining(&*w, 1)?; - w.put_u8(*self); - Ok(()) + /// Where the next byte goes: the buffer's length, including any written before this + /// encoder. Mark a size-prefixed body's start with it. + pub fn position(&self) -> usize { + self.buf.len() } -} -impl Encode for u16 { - fn encode(&self, w: &mut W, _: V) -> Result<(), EncodeError> { - check_remaining(&*w, 2)?; - w.put_u16(*self); - Ok(()) + /// Write raw bytes. + pub fn slice(&mut self, v: &[u8]) { + self.buf.extend_from_slice(v); } -} -impl Encode for String -where - usize: Encode, -{ - fn encode(&self, w: &mut W, version: V) -> Result<(), EncodeError> { - self.as_str().encode(w, version) + /// Write a single byte. + pub fn u8(&mut self, v: u8) { + self.buf.push(v); } -} -impl Encode for &str -where - usize: Encode, -{ - fn encode(&self, w: &mut W, version: V) -> Result<(), EncodeError> { - self.len().encode(w, version)?; - check_remaining(&*w, self.len())?; - w.put(self.as_bytes()); - Ok(()) + /// Write a big-endian `u16`. + pub fn u16(&mut self, v: u16) { + self.buf.extend_from_slice(&v.to_be_bytes()); } -} -impl Encode for i8 { - fn encode(&self, w: &mut W, _: V) -> Result<(), EncodeError> { - // This is not the usual way of encoding negative numbers. - // i8 doesn't exist in the draft, but we use it instead of u8 for priority. - // A default of 0 is more ergonomic for the user than a default of 128. - check_remaining(&*w, 1)?; - w.put_u8(((*self as i16) + 128) as u8); - Ok(()) + /// Write a bool as a 0 or 1 byte. + pub fn bool(&mut self, v: bool) { + self.buf.push(v as u8); } -} -impl> Encode for &[T] -where - usize: Encode, -{ - fn encode(&self, w: &mut W, version: V) -> Result<(), EncodeError> { - self.len().encode(w, version)?; - for item in self.iter() { - item.encode(w, version)?; - } + /// Write a varint, or fail with [`EncodeError::BoundsExceeded`] if the form cannot carry it. + pub fn varint(&mut self, v: VarInt) -> Result<(), EncodeError> { + let (buf, len) = v.encode_form(self.form)?; + self.buf.extend_from_slice(&buf[..len]); Ok(()) } -} -impl Encode for Vec -where - usize: Encode, -{ - fn encode(&self, w: &mut W, version: V) -> Result<(), EncodeError> { - self.len().encode(w, version)?; - check_remaining(&*w, self.len())?; - w.put_slice(self); - Ok(()) + /// Write an optional varint: `None` as 0, and `Some(n)` as `n + 1`. + pub fn varint_opt(&mut self, v: Option) -> Result<(), EncodeError> { + let v = match v { + Some(v) => v.checked_add(1).ok_or(EncodeError::TooLarge)?, + None => 0, + }; + self.varint(v.into()) } -} -impl Encode for bytes::Bytes -where - usize: Encode, -{ - fn encode(&self, w: &mut W, version: V) -> Result<(), EncodeError> { - self.len().encode(w, version)?; - check_remaining(&*w, self.len())?; - w.put_slice(self); + /// Write a varint length, then the raw bytes. + pub fn bytes(&mut self, v: &[u8]) -> Result<(), EncodeError> { + self.varint(v.len().into())?; + self.slice(v); Ok(()) } -} -impl, V> Encode for Arc { - fn encode(&self, w: &mut W, version: V) -> Result<(), EncodeError> { - (**self).encode(w, version) + /// Write a varint length, then the UTF-8 bytes. + pub fn string(&mut self, v: &str) -> Result<(), EncodeError> { + self.bytes(v.as_bytes()) } -} -impl Encode for Cow<'_, str> -where - usize: Encode, -{ - fn encode(&self, w: &mut W, version: V) -> Result<(), EncodeError> { - self.len().encode(w, version)?; - check_remaining(&*w, self.len())?; - w.put(self.as_bytes()); + /// Prefix everything written since [`Self::position`] was `start` with its varint length. + pub fn prefix_varint(&mut self, start: usize) -> Result<(), EncodeError> { + let size = VarInt::from(self.buf.len() - start); + let (prefix, len) = size.encode_form(self.form)?; + self.buf.splice(start..start, prefix[..len].iter().copied()); Ok(()) } -} -impl Encode for Option -where - u64: Encode, -{ - fn encode(&self, w: &mut W, version: V) -> Result<(), EncodeError> { - match self { - Some(value) => value.checked_add(1).ok_or(EncodeError::TooLarge)?.encode(w, version), - None => 0u64.encode(w, version), - } + /// Prefix everything written since [`Self::position`] was `start` with its `u16` length. + pub fn prefix_u16(&mut self, start: usize) -> Result<(), EncodeError> { + let size = u16::try_from(self.buf.len() - start).map_err(|_| EncodeError::TooLarge)?; + self.buf.splice(start..start, size.to_be_bytes()); + Ok(()) } } -impl Encode for std::time::Duration -where - super::VarInt: Encode, -{ - /// Milliseconds as a QUIC varint. Sub-millisecond precision is truncated - /// (`Duration::as_millis`). Fails with [`EncodeError::BoundsExceeded`] if the - /// millisecond count exceeds `2^62 - 1`. - fn encode(&self, w: &mut W, version: V) -> Result<(), EncodeError> { - let ms = super::VarInt::try_from(self.as_millis())?; - ms.encode(w, version) +#[cfg(test)] +mod tests { + use super::*; + + /// The prefix lands before the body, sized to it, even when it takes more than a byte. + #[test] + fn prefix_varint_sizes_the_body() { + for size in [0usize, 63, 64, 20_000] { + let mut buf = vec![0xaa]; + let mut w = Encoder::new(&mut buf, Form::Quic); + let start = w.position(); + w.slice(&vec![0x55; size]); + w.prefix_varint(start).unwrap(); + + let mut r = super::super::Decoder::new(&buf[1..], Form::Quic); + assert_eq!(buf[0], 0xaa); + assert_eq!(r.varint().unwrap().into_inner(), size as u64); + assert_eq!(r.rest(), vec![0x55; size]); + } } } diff --git a/rs/moq-net/src/coding/mod.rs b/rs/moq-net/src/coding/mod.rs index 5f427d0974..c44cc6f23a 100644 --- a/rs/moq-net/src/coding/mod.rs +++ b/rs/moq-net/src/coding/mod.rs @@ -4,7 +4,6 @@ mod codes; mod decode; mod encode; mod reader; -mod size; mod stream; mod varint; mod version; @@ -14,8 +13,28 @@ pub use codes::*; pub use decode::*; pub use encode::*; pub use reader::*; -pub use size::*; pub use stream::*; pub use varint::*; pub use version::*; pub use writer::*; + +/// Decode from the front of a test buffer with `decode`, advancing it past what was read. +#[cfg(test)] +pub(crate) fn decode_buf + Copy, T>( + buf: &mut B, + version: V, + decode: impl FnOnce(&mut Decoder<'_>, V) -> Result, +) -> Result { + let chunk = buf.chunk(); + let mut r = Decoder::new(chunk, version.into()); + let value = decode(&mut r, version)?; + let used = chunk.len() - r.remaining(); + buf.advance(used); + Ok(value) +} + +/// Decode one varint from the front of a test buffer, advancing it past what was read. +#[cfg(test)] +pub(crate) fn decode_varint + Copy>(buf: &mut B, version: V) -> Result { + decode_buf(buf, version, |r, _| Ok(r.varint()?.into_inner())) +} diff --git a/rs/moq-net/src/coding/reader.rs b/rs/moq-net/src/coding/reader.rs index 434e962551..dbc0968207 100644 --- a/rs/moq-net/src/coding/reader.rs +++ b/rs/moq-net/src/coding/reader.rs @@ -1,7 +1,6 @@ use std::{ cmp, fmt::Debug, - io, task::{Context, Poll, ready}, }; @@ -44,13 +43,14 @@ impl Reader { /// Poll for the next message on the stream. pub fn poll_decode + Debug>(&mut self, cx: &mut Context<'_>) -> Poll> where - V: Clone, + V: Into + Copy, { loop { - let mut cursor = io::Cursor::new(&self.buffer); - match T::decode(&mut cursor, self.version.clone()) { + let mut r = Decoder::new(&self.buffer, self.version.into()); + match T::decode(&mut r, self.version) { Ok(msg) => { - self.buffer.advance(cursor.position() as usize); + let used = self.buffer.len() - r.remaining(); + self.buffer.advance(used); return Poll::Ready(Ok(msg)); } // Stream closed while we still need more data. @@ -66,7 +66,7 @@ impl Reader { /// Decode the next message from the stream. pub async fn decode + Debug>(&mut self) -> Result where - V: Clone, + V: Into + Copy, { std::future::poll_fn(|cx| self.poll_decode(cx)).await } @@ -74,7 +74,7 @@ impl Reader { /// Poll for the next message unless the stream is closed cleanly first. pub fn poll_decode_maybe + Debug>(&mut self, cx: &mut Context<'_>) -> Poll, Error>> where - V: Clone, + V: Into + Copy, { if !ready!(self.poll_has_more(cx))? { return Poll::Ready(Ok(None)); @@ -86,7 +86,7 @@ impl Reader { /// Decode the next message unless the stream is closed. pub async fn decode_maybe + Debug>(&mut self) -> Result, Error> where - V: Clone, + V: Into + Copy, { std::future::poll_fn(|cx| self.poll_decode_maybe(cx)).await } @@ -94,11 +94,11 @@ impl Reader { /// Poll for the next message without consuming it. pub fn poll_decode_peek + Debug>(&mut self, cx: &mut Context<'_>) -> Poll> where - V: Clone, + V: Into + Copy, { loop { - let mut cursor = io::Cursor::new(&self.buffer); - match T::decode(&mut cursor, self.version.clone()) { + let mut r = Decoder::new(&self.buffer, self.version.into()); + match T::decode(&mut r, self.version) { Ok(msg) => return Poll::Ready(Ok(msg)), Err(DecodeError::Short) if !ready!(self.poll_read_more(cx))? => { return Poll::Ready(Err(DecodeError::Short.into())); @@ -112,7 +112,7 @@ impl Reader { /// Decode the next message from the stream without consuming it. pub async fn decode_peek + Debug>(&mut self) -> Result where - V: Clone, + V: Into + Copy, { std::future::poll_fn(|cx| self.poll_decode_peek(cx)).await } @@ -123,7 +123,7 @@ impl Reader { cx: &mut Context<'_>, ) -> Poll, Error>> where - V: Clone, + V: Into + Copy, { if !ready!(self.poll_has_more(cx))? { return Poll::Ready(Ok(None)); @@ -135,7 +135,7 @@ impl Reader { /// Peek the next message unless the stream is closed. pub async fn decode_peek_maybe + Debug>(&mut self) -> Result, Error> where - V: Clone, + V: Into + Copy, { std::future::poll_fn(|cx| self.poll_decode_peek_maybe(cx)).await } diff --git a/rs/moq-net/src/coding/size.rs b/rs/moq-net/src/coding/size.rs deleted file mode 100644 index 7642e02fc3..0000000000 --- a/rs/moq-net/src/coding/size.rs +++ /dev/null @@ -1,212 +0,0 @@ -use std::mem::MaybeUninit; - -use bytes::{Buf, BufMut, buf::UninitSlice}; - -/// A [BufMut] implementation that only counts the size of the buffer. -/// -/// Used to calculate the size of a message before encoding it. -#[derive(Default)] -pub struct Sizer { - pub size: usize, -} - -unsafe impl BufMut for Sizer { - unsafe fn advance_mut(&mut self, cnt: usize) { - self.size += cnt; - } - - fn chunk_mut(&mut self) -> &mut UninitSlice { - // We need to return a valid slice, but it won't actually be written to - // Use a thread-local static buffer to avoid safety issues - thread_local! { - static BUFFER: std::cell::UnsafeCell<[MaybeUninit; 8192]> = - const { std::cell::UnsafeCell::new([MaybeUninit::uninit(); 8192]) }; - } - - BUFFER.with(|buf| { - let ptr = buf.get(); - unsafe { - let slice = (*ptr).as_mut_ptr(); - bytes::buf::UninitSlice::from_raw_parts_mut(slice as *mut u8, 8192) - } - }) - } - - fn remaining_mut(&self) -> usize { - usize::MAX - } - - fn has_remaining_mut(&self) -> bool { - true - } - - fn put(&mut self, mut src: T) { - self.size += src.remaining(); - src.advance(src.remaining()); - } - - fn put_bytes(&mut self, _val: u8, cnt: usize) { - self.size += cnt; - } - - fn put_f32(&mut self, _val: f32) { - self.size += 4; - } - - fn put_f32_le(&mut self, _: f32) { - self.size += 4 - } - - fn put_f32_ne(&mut self, _: f32) { - self.size += 4 - } - - fn put_f64(&mut self, _: f64) { - self.size += 8 - } - - fn put_f64_le(&mut self, _: f64) { - self.size += 8 - } - - fn put_f64_ne(&mut self, _: f64) { - self.size += 8 - } - - fn put_i128(&mut self, _: i128) { - self.size += 16 - } - - fn put_i128_le(&mut self, _: i128) { - self.size += 16 - } - - fn put_i128_ne(&mut self, _: i128) { - self.size += 16 - } - - fn put_i16(&mut self, _: i16) { - self.size += 2 - } - - fn put_i16_le(&mut self, _: i16) { - self.size += 2 - } - - fn put_i16_ne(&mut self, _: i16) { - self.size += 2 - } - - fn put_i32(&mut self, _: i32) { - self.size += 4 - } - - fn put_i32_le(&mut self, _: i32) { - self.size += 4 - } - - fn put_i32_ne(&mut self, _: i32) { - self.size += 4 - } - - fn put_i64(&mut self, _: i64) { - self.size += 8 - } - - fn put_i64_le(&mut self, _: i64) { - self.size += 8 - } - - fn put_i64_ne(&mut self, _: i64) { - self.size += 8 - } - - fn put_i8(&mut self, _: i8) { - self.size += 1 - } - - fn put_int(&mut self, _: i64, nbytes: usize) { - self.size += nbytes - } - - fn put_int_le(&mut self, _: i64, nbytes: usize) { - self.size += nbytes - } - - fn put_int_ne(&mut self, _: i64, nbytes: usize) { - self.size += nbytes - } - - fn put_slice(&mut self, src: &[u8]) { - self.size += src.len(); - } - - fn put_u128(&mut self, _: u128) { - self.size += 16 - } - - fn put_u128_le(&mut self, _: u128) { - self.size += 16 - } - - fn put_u128_ne(&mut self, _: u128) { - self.size += 16 - } - - fn put_u16(&mut self, _: u16) { - self.size += 2 - } - - fn put_u16_le(&mut self, _: u16) { - self.size += 2 - } - - fn put_u16_ne(&mut self, _: u16) { - self.size += 2 - } - - fn put_u32(&mut self, _: u32) { - self.size += 4 - } - - fn put_u32_le(&mut self, _: u32) { - self.size += 4 - } - - fn put_u32_ne(&mut self, _: u32) { - self.size += 4 - } - - fn put_u64(&mut self, _: u64) { - self.size += 8 - } - - fn put_u64_le(&mut self, _: u64) { - self.size += 8 - } - - fn put_u64_ne(&mut self, _: u64) { - self.size += 8 - } - - fn put_u8(&mut self, _: u8) { - self.size += 1 - } - - fn put_uint(&mut self, _: u64, nbytes: usize) { - self.size += nbytes - } - - fn put_uint_le(&mut self, _: u64, nbytes: usize) { - self.size += nbytes - } - - fn put_uint_ne(&mut self, _: u64, nbytes: usize) { - self.size += nbytes - } - - // TODO - // fn writer(self) -> bytes::buf::Writer { - // fn chain_mut(self, next: U) -> bytes::buf::Chain - // fn limit(self, limit: usize) -> bytes::buf::Limit -} diff --git a/rs/moq-net/src/coding/varint.rs b/rs/moq-net/src/coding/varint.rs index a933977174..5a55bb8547 100644 --- a/rs/moq-net/src/coding/varint.rs +++ b/rs/moq-net/src/coding/varint.rs @@ -2,47 +2,49 @@ // https://github.com/quinn-rs/quinn/blob/main/quinn-proto/src/varint.rs // Licensed via Apache 2.0 and MIT -use std::convert::{TryFrom, TryInto}; use std::fmt; use thiserror::Error; -use super::{Decode, DecodeError, Encode, EncodeError}; +use super::{Decode, DecodeError, Decoder, Encode, EncodeError, Encoder}; +use crate::{Version, ietf, lite}; -/// The number is too large to fit in a VarInt (62 bits). +/// The number does not fit the target: a varint wire form or a narrower integer. #[derive(Debug, Copy, Clone, Eq, PartialEq, Error)] #[error("value out of range")] pub struct BoundsExceeded; -/// An integer less than 2^62 +/// An integer destined for the wire as a variable-length integer. /// -/// Values of this type are suitable for encoding as QUIC variable-length integer. -/// It would be neat if we could express to Rust that the top two bits are available for use as enum -/// discriminants +/// It holds the full `u64` range, which the leading-ones form of moq-transport draft-17+ +/// can carry. The QUIC form (moq-lite, drafts 14-16) tops out at `2^62 - 1`, so encoding +/// a larger value there fails with [`BoundsExceeded`] rather than truncating. #[derive(Debug, Default, Copy, Clone, Eq, PartialEq, Ord, PartialOrd, Hash)] pub struct VarInt(u64); impl VarInt { /// The largest possible value. - pub const MAX: Self = Self((1 << 62) - 1); + pub const MAX: Self = Self(u64::MAX); + + /// The largest value the QUIC form can carry: `2^62 - 1`. + pub const MAX_QUIC: Self = Self((1 << 62) - 1); /// The smallest possible value. pub const ZERO: Self = Self(0); - /// Construct a `VarInt` infallibly using the largest available type. - /// Larger values need to use `try_from` instead. + /// Construct from a `u32`, usable in `const` contexts. pub const fn from_u32(x: u32) -> Self { Self(x as u64) } - /// Construct from a `u64`, or `None` if it exceeds [`Self::MAX`]. - pub const fn from_u64(x: u64) -> Option { - if x <= Self::MAX.0 { Some(Self(x)) } else { None } + /// Construct from a `u64`, usable in `const` contexts. + pub const fn from_u64(x: u64) -> Self { + Self(x) } /// Construct from a `u128`, or `None` if it exceeds [`Self::MAX`]. pub const fn from_u128(x: u128) -> Option { - if x <= Self::MAX.0 as u128 { + if x <= u64::MAX as u128 { Some(Self(x as u64)) } else { None @@ -54,17 +56,12 @@ impl VarInt { self.0 } - /// Encode a signed `i64` as a zigzag-then-unsigned varint: `(n << 1) ^ (n >> 63)`. + /// Map a signed `i64` onto the unsigned range with zigzag: `(n << 1) ^ (n >> 63)`. /// - /// Small negative numbers map to small unsigneds (-1 -> 1, 1 -> 2, -2 -> 3, ...). - /// Returns [`BoundsExceeded`] if `signed` is outside `[-2^61, 2^61 - 1]`, since the - /// zigzag-encoded result must fit in a 62-bit varint. - pub const fn from_zigzag(signed: i64) -> Result { - const RANGE: i64 = 1 << 61; - if signed < -RANGE || signed >= RANGE { - return Err(BoundsExceeded); - } - Ok(Self(((signed << 1) ^ (signed >> 63)) as u64)) + /// Small negative numbers map to small unsigneds (-1 -> 1, 1 -> 2, -2 -> 3, ...), and + /// the whole `i64` range fits. + pub const fn from_zigzag(signed: i64) -> Self { + Self(((signed << 1) ^ (signed >> 63)) as u64) } /// Decode this varint as a signed `i64` via the inverse zigzag transform. @@ -80,12 +77,6 @@ impl From for u64 { } } -impl From for usize { - fn from(x: VarInt) -> Self { - x.0 as usize - } -} - impl From for u128 { fn from(x: VarInt) -> Self { x.0 as u128 @@ -110,35 +101,34 @@ impl From for VarInt { } } -impl TryFrom for VarInt { - type Error = BoundsExceeded; +impl From for VarInt { + fn from(x: u64) -> Self { + Self(x) + } +} - /// Succeeds iff `x` < 2^62 - fn try_from(x: u64) -> Result { - let x = Self(x); - if x <= Self::MAX { Ok(x) } else { Err(BoundsExceeded) } +impl From for VarInt { + fn from(x: usize) -> Self { + // usize is at most 64 bits on every target Rust supports. + Self(x as u64) } } impl TryFrom for VarInt { type Error = BoundsExceeded; - /// Succeeds iff `x` < 2^62 + /// Succeeds iff `x` < 2^64 fn try_from(x: u128) -> Result { - if x <= Self::MAX.into() { - Ok(Self(x as u64)) - } else { - Err(BoundsExceeded) - } + Self::from_u128(x).ok_or(BoundsExceeded) } } -impl TryFrom for VarInt { +impl TryFrom for usize { type Error = BoundsExceeded; - /// Succeeds iff `x` < 2^62 - fn try_from(x: usize) -> Result { - Self::try_from(x as u64) + /// Succeeds iff `x` fits the target's pointer width. + fn try_from(x: VarInt) -> Result { + usize::try_from(x.0).map_err(|_| BoundsExceeded) } } @@ -147,11 +137,7 @@ impl TryFrom for u32 { /// Succeeds iff `x` < 2^32 fn try_from(x: VarInt) -> Result { - if x.0 <= u32::MAX.into() { - Ok(x.0 as u32) - } else { - Err(BoundsExceeded) - } + u32::try_from(x.0).map_err(|_| BoundsExceeded) } } @@ -160,11 +146,7 @@ impl TryFrom for u16 { /// Succeeds iff `x` < 2^16 fn try_from(x: VarInt) -> Result { - if x.0 <= u16::MAX.into() { - Ok(x.0 as u16) - } else { - Err(BoundsExceeded) - } + u16::try_from(x.0).map_err(|_| BoundsExceeded) } } @@ -173,11 +155,7 @@ impl TryFrom for u8 { /// Succeeds iff `x` < 2^8 fn try_from(x: VarInt) -> Result { - if x.0 <= u8::MAX.into() { - Ok(x.0 as u8) - } else { - Err(BoundsExceeded) - } + u8::try_from(x.0).map_err(|_| BoundsExceeded) } } @@ -187,384 +165,183 @@ impl fmt::Display for VarInt { } } -impl VarInt { - /// Decode a QUIC-style varint (2-bit length tag in top bits). - pub fn decode_quic(r: &mut R) -> Result { - if !r.has_remaining() { - return Err(DecodeError::Short); - } - - let b = r.get_u8(); - let tag = b >> 6; - - let mut buf = [0u8; 8]; - buf[0] = b & 0b0011_1111; - - let x = match tag { - 0b00 => u64::from(buf[0]), - 0b01 => { - if !r.has_remaining() { - return Err(DecodeError::Short); - } - r.copy_to_slice(buf[1..2].as_mut()); - u64::from(u16::from_be_bytes(buf[..2].try_into().unwrap())) - } - 0b10 => { - if r.remaining() < 3 { - return Err(DecodeError::Short); - } - r.copy_to_slice(buf[1..4].as_mut()); - u64::from(u32::from_be_bytes(buf[..4].try_into().unwrap())) - } - 0b11 => { - if r.remaining() < 7 { - return Err(DecodeError::Short); - } - r.copy_to_slice(buf[1..8].as_mut()); - u64::from_be_bytes(buf) - } - _ => unreachable!(), - }; - - Ok(Self(x)) - } +/// How a protocol version lays out a [`VarInt`] on the wire. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum Form { + /// QUIC's two-bit length tag, up to [`VarInt::MAX_QUIC`]. + Quic, + /// Leading one bits count the length, up to [`VarInt::MAX`]. + LeadingOnes { + /// Whether the 7-byte form (`1111110x`) is accepted on decode, which draft-17 forbids. + seven: bool, + }, +} - /// Encode a QUIC-style varint (2-bit length tag in top bits). - pub fn encode_quic(&self, w: &mut W) -> Result<(), EncodeError> { - let remaining = w.remaining_mut(); - if self.0 < (1u64 << 6) { - if remaining < 1 { - return Err(EncodeError::Short); - } - w.put_u8(self.0 as u8); - } else if self.0 < (1u64 << 14) { - if remaining < 2 { - return Err(EncodeError::Short); - } - w.put_u16((0b01 << 14) | self.0 as u16); - } else if self.0 < (1u64 << 30) { - if remaining < 4 { - return Err(EncodeError::Short); - } - w.put_u32((0b10 << 30) | self.0 as u32); - } else if self.0 < (1u64 << 62) { - if remaining < 8 { - return Err(EncodeError::Short); - } - w.put_u64((0b11 << 62) | self.0); - } else { - return Err(BoundsExceeded.into()); +impl From for Form { + fn from(version: lite::Version) -> Self { + match version { + lite::Version::Lite01 + | lite::Version::Lite02 + | lite::Version::Lite03 + | lite::Version::Lite04 + | lite::Version::Lite05 + | lite::Version::Lite06 + | lite::Version::Lite07 => Self::Quic, } - Ok(()) } +} - /// Decode a leading-1-bits varint (draft-17+ Section 1.4.1). - /// - /// The number of leading 1-bits determines the byte length: - /// - `0xxxxxxx` → 1 byte, 7 usable bits - /// - `10xxxxxx` → 2 bytes, 14 usable bits - /// - `110xxxxx` → 3 bytes, 21 usable bits - /// - `1110xxxx` → 4 bytes, 28 usable bits - /// - `11110xxx` → 5 bytes, 35 usable bits - /// - `111110xx` → 6 bytes, 42 usable bits - /// - `1111110x` → 7 bytes, 49 usable bits (draft-18+, INVALID in draft-17 per #1595) - /// - `11111110` → 8 bytes, 56 usable bits - /// - `11111111` → 9 bytes, 64 usable bits - fn decode_leading_ones(r: &mut R, version: ietf::Version) -> Result { - if !r.has_remaining() { - return Err(DecodeError::Short); - } - - let b = r.get_u8(); - let ones = b.leading_ones() as usize; - - match ones { - 0 => { - // 0xxxxxxx: 7 bits - Ok(Self(u64::from(b))) - } - 1 => { - // 10xxxxxx + 1 byte: 14 bits - if !r.has_remaining() { - return Err(DecodeError::Short); - } - let hi = u64::from(b & 0x3F); - let lo = u64::from(r.get_u8()); - Ok(Self((hi << 8) | lo)) - } - 2 => { - // 110xxxxx + 2 bytes: 21 bits - if r.remaining() < 2 { - return Err(DecodeError::Short); - } - let hi = u64::from(b & 0x1F); - let mut buf = [0u8; 2]; - r.copy_to_slice(&mut buf); - Ok(Self((hi << 16) | u64::from(u16::from_be_bytes(buf)))) - } - 3 => { - // 1110xxxx + 3 bytes: 28 bits - if r.remaining() < 3 { - return Err(DecodeError::Short); - } - let hi = u64::from(b & 0x0F); - let mut buf = [0u8; 3]; - r.copy_to_slice(&mut buf); - Ok(Self( - (hi << 24) | u64::from(buf[0]) << 16 | u64::from(buf[1]) << 8 | u64::from(buf[2]), - )) - } - 4 => { - // 11110xxx + 4 bytes: 35 bits - if r.remaining() < 4 { - return Err(DecodeError::Short); - } - let hi = u64::from(b & 0x07); - let mut buf = [0u8; 4]; - r.copy_to_slice(&mut buf); - Ok(Self((hi << 32) | u64::from(u32::from_be_bytes(buf)))) - } - 5 => { - // 111110xx + 5 bytes: 42 bits - if r.remaining() < 5 { - return Err(DecodeError::Short); - } - let hi = u64::from(b & 0x03); - let mut buf = [0u8; 5]; - r.copy_to_slice(&mut buf); - let lo = u64::from(buf[0]) << 32 - | u64::from(buf[1]) << 24 - | u64::from(buf[2]) << 16 - | u64::from(buf[3]) << 8 - | u64::from(buf[4]); - Ok(Self((hi << 40) | lo)) - } - 6 => { - // 1111110x + 6 bytes, 49 bits (draft-18+, INVALID in draft-17 per #1595) - if matches!(version, ietf::Version::Draft17) { - return Err(DecodeError::InvalidValue); - } - if r.remaining() < 6 { - return Err(DecodeError::Short); - } - let hi = u64::from(b & 0x01); - let mut buf = [0u8; 8]; - r.copy_to_slice(&mut buf[2..]); - Ok(Self((hi << 48) | u64::from_be_bytes(buf))) - } - 7 => { - // 11111110 + 7 bytes: 56 bits - if r.remaining() < 7 { - return Err(DecodeError::Short); - } - let mut buf = [0u8; 8]; - buf[0] = 0; - r.copy_to_slice(&mut buf[1..]); - Ok(Self(u64::from_be_bytes(buf))) - } - 8 => { - // 11111111 + 8 bytes: 64 bits - if r.remaining() < 8 { - return Err(DecodeError::Short); - } - let mut buf = [0u8; 8]; - r.copy_to_slice(&mut buf); - Ok(Self(u64::from_be_bytes(buf))) - } - _ => unreachable!(), +impl From for Form { + fn from(version: ietf::Version) -> Self { + match version { + ietf::Version::Draft14 | ietf::Version::Draft15 | ietf::Version::Draft16 => Self::Quic, + ietf::Version::Draft17 => Self::LeadingOnes { seven: false }, + _ => Self::LeadingOnes { seven: true }, } } +} - /// Encode a leading-1-bits varint (draft-17+ Section 1.4.1). - /// - /// Always emits the minimal canonical form. Draft-18 also accepts 7-byte form - /// (`1111110x`) on decode but we never emit it because the 8-byte form is one byte - /// larger but simpler and is universally valid. - fn encode_leading_ones(&self, w: &mut W, _version: ietf::Version) -> Result<(), EncodeError> { - let x = self.0; - let remaining = w.remaining_mut(); - - if x < (1 << 7) { - // 0xxxxxxx: 1 byte - if remaining < 1 { - return Err(EncodeError::Short); - } - w.put_u8(x as u8); - } else if x < (1 << 14) { - // 10xxxxxx: 2 bytes - if remaining < 2 { - return Err(EncodeError::Short); - } - w.put_u8(0x80 | (x >> 8) as u8); - w.put_u8(x as u8); - } else if x < (1 << 21) { - // 110xxxxx: 3 bytes - if remaining < 3 { - return Err(EncodeError::Short); - } - w.put_u8(0xC0 | (x >> 16) as u8); - w.put_u16(x as u16); - } else if x < (1 << 28) { - // 1110xxxx: 4 bytes - if remaining < 4 { - return Err(EncodeError::Short); - } - w.put_u8(0xE0 | (x >> 24) as u8); - w.put_u8((x >> 16) as u8); - w.put_u16(x as u16); - } else if x < (1 << 35) { - // 11110xxx: 5 bytes - if remaining < 5 { - return Err(EncodeError::Short); - } - w.put_u8(0xF0 | (x >> 32) as u8); - w.put_u32(x as u32); - } else if x < (1 << 42) { - // 111110xx: 6 bytes - if remaining < 6 { - return Err(EncodeError::Short); - } - w.put_u8(0xF8 | (x >> 40) as u8); - w.put_u8((x >> 32) as u8); - w.put_u32(x as u32); - } else if x < (1 << 56) { - // 11111110: 8 bytes (skips 7) - if remaining < 8 { - return Err(EncodeError::Short); - } - w.put_u8(0xFE); - // Write 7 bytes: high byte then low 6 bytes - w.put_u8((x >> 48) as u8); - w.put_u16((x >> 32) as u16); - w.put_u32(x as u32); - } else { - // 11111111: 9 bytes - if remaining < 9 { - return Err(EncodeError::Short); - } - w.put_u8(0xFF); - w.put_u64(x); +impl From for Form { + fn from(version: Version) -> Self { + match version { + Version::Lite(v) => v.into(), + Version::Ietf(v) => v.into(), } - - Ok(()) } } -use crate::{Version, ietf, lite}; - -// All lite versions use QUIC-style varint encoding. -impl Encode for VarInt { - fn encode(&self, w: &mut W, _: lite::Version) -> Result<(), EncodeError> { - self.encode_quic(w) +impl VarInt { + /// The high and low 32 bits. + /// + /// The codec works on the halves so it never needs 64-bit bitwise math, which a + /// JavaScript `number` cannot do. + const fn to_halves(self) -> (u32, u32) { + ((self.0 >> 32) as u32, self.0 as u32) } -} -impl Decode for VarInt { - fn decode(r: &mut R, _: lite::Version) -> Result { - Self::decode_quic(r) + /// The inverse of [`Self::to_halves`]. + const fn from_halves(hi: u32, lo: u32) -> Self { + Self(((hi as u64) << 32) | lo as u64) } -} -// Draft14-16 use QUIC-style varints; draft-17+ uses leading-ones. -impl Encode for VarInt { - fn encode(&self, w: &mut W, version: ietf::Version) -> Result<(), EncodeError> { - match version { - ietf::Version::Draft14 | ietf::Version::Draft15 | ietf::Version::Draft16 => self.encode_quic(w), - _ => self.encode_leading_ones(w, version), - } - } -} + /// Decode from the front of `buf`, returning the value and the bytes it took. + pub(super) fn decode_form(buf: &[u8], form: Form) -> Result<(Self, usize), DecodeError> { + let Some(&first) = buf.first() else { + return Err(DecodeError::Short); + }; -impl Decode for VarInt { - fn decode(r: &mut R, version: ietf::Version) -> Result { - match version { - ietf::Version::Draft14 | ietf::Version::Draft15 | ietf::Version::Draft16 => Self::decode_quic(r), - _ => Self::decode_leading_ones(r, version), - } - } -} + let (len, head) = match form { + Form::Quic => (1usize << (first >> 6), first & 0x3f), + Form::LeadingOnes { seven } => { + let ones = first.leading_ones(); + if ones == 6 && !seven { + return Err(DecodeError::InvalidValue); + } + // `0x7f >> 8` would overflow, and there are no value bits left anyway. + let head = if ones >= 7 { 0 } else { first & (0x7f >> ones) }; + (ones as usize + 1, head) + } + }; -// The top-level Version delegates to the sub-version impls. -impl Encode for VarInt { - fn encode(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { - match version { - Version::Lite(v) => self.encode(w, v), - Version::Ietf(v) => self.encode(w, v), - } - } -} + let Some(rest) = buf.get(1..len) else { + return Err(DecodeError::Short); + }; -impl Decode for VarInt { - fn decode(r: &mut R, version: Version) -> Result { - match version { - Version::Lite(v) => Self::decode(r, v), - Version::Ietf(v) => Self::decode(r, v), + let mut hi = 0u32; + let mut lo = head as u32; + for &byte in rest { + hi = (hi << 8) | (lo >> 24); + lo = (lo << 8) | byte as u32; } - } -} -// Blanket impls for integer types that delegate to VarInt. -impl Encode for u64 -where - VarInt: Encode, -{ - fn encode(&self, w: &mut W, version: V) -> Result<(), EncodeError> { - VarInt::try_from(*self)?.encode(w, version) + Ok((Self::from_halves(hi, lo), len)) } -} -impl Decode for u64 -where - VarInt: Decode, -{ - fn decode(r: &mut R, version: V) -> Result { - VarInt::decode(r, version).map(|v| v.into_inner()) + /// The minimal encoding in the given form: the bytes, and how many of them are used. + /// + /// Fails past [`Self::MAX_QUIC`] in the QUIC form. + pub(super) fn encode_form(self, form: Form) -> Result<([u8; 9], usize), BoundsExceeded> { + let (hi, lo) = self.to_halves(); + let [a, b, c, d] = lo.to_be_bytes(); + let [e, f, g, h] = hi.to_be_bytes(); + + Ok(match form { + Form::Quic if hi == 0 && lo < 1 << 6 => ([d, 0, 0, 0, 0, 0, 0, 0, 0], 1), + Form::Quic if hi == 0 && lo < 1 << 14 => ([0x40 | c, d, 0, 0, 0, 0, 0, 0, 0], 2), + Form::Quic if hi == 0 && lo < 1 << 30 => ([0x80 | a, b, c, d, 0, 0, 0, 0, 0], 4), + Form::Quic if hi < 1 << 30 => ([0xc0 | e, f, g, h, a, b, c, d, 0], 8), + Form::Quic => return Err(BoundsExceeded), + Form::LeadingOnes { .. } if hi == 0 && lo < 1 << 7 => ([d, 0, 0, 0, 0, 0, 0, 0, 0], 1), + Form::LeadingOnes { .. } if hi == 0 && lo < 1 << 14 => ([0x80 | c, d, 0, 0, 0, 0, 0, 0, 0], 2), + Form::LeadingOnes { .. } if hi == 0 && lo < 1 << 21 => ([0xc0 | b, c, d, 0, 0, 0, 0, 0, 0], 3), + Form::LeadingOnes { .. } if hi == 0 && lo < 1 << 28 => ([0xe0 | a, b, c, d, 0, 0, 0, 0, 0], 4), + Form::LeadingOnes { .. } if hi < 1 << 3 => ([0xf0 | h, a, b, c, d, 0, 0, 0, 0], 5), + Form::LeadingOnes { .. } if hi < 1 << 10 => ([0xf8 | g, h, a, b, c, d, 0, 0, 0], 6), + // The 7-byte form is skipped: one byte longer, but legal on every draft. + Form::LeadingOnes { .. } if hi < 1 << 24 => ([0xfe, f, g, h, a, b, c, d, 0], 8), + Form::LeadingOnes { .. } => ([0xff, e, f, g, h, a, b, c, d], 9), + }) + } + + /// The bytes this takes on the wire in the given form, or [`BoundsExceeded`] if it + /// does not fit. + pub(crate) fn size(self, form: Form) -> Result { + Ok(self.encode_form(form)?.1) } -} -impl Encode for usize -where - VarInt: Encode, -{ - fn encode(&self, w: &mut W, version: V) -> Result<(), EncodeError> { - VarInt::try_from(*self)?.encode(w, version) + /// Decode a QUIC-style varint (2-bit length tag in top bits). + pub fn decode_quic(r: &mut R) -> Result { + let Some(&first) = r.chunk().first() else { + return Err(DecodeError::Short); + }; + + // Copy out so a varint split across chunks still decodes. + let len = 1usize << (first >> 6); + if r.remaining() < len { + return Err(DecodeError::Short); + } + let mut buf = [0u8; 8]; + r.copy_to_slice(&mut buf[..len]); + + Ok(Self::decode_form(&buf[..len], Form::Quic)?.0) } -} -impl Decode for usize -where - VarInt: Decode, -{ - fn decode(r: &mut R, version: V) -> Result { - VarInt::decode(r, version).map(|v| v.into_inner() as usize) + /// Encode a QUIC-style varint (2-bit length tag in top bits). + /// + /// Fails with [`EncodeError::BoundsExceeded`] past [`Self::MAX_QUIC`]. + pub fn encode_quic(&self, w: &mut W) -> Result<(), EncodeError> { + let (buf, len) = self.encode_form(Form::Quic)?; + if w.remaining_mut() < len { + return Err(EncodeError::Short); + } + w.put_slice(&buf[..len]); + Ok(()) } } -impl Encode for u32 -where - VarInt: Encode, -{ - fn encode(&self, w: &mut W, version: V) -> Result<(), EncodeError> { - VarInt::from(*self).encode(w, version) +impl Encode for VarInt { + fn encode(&self, w: &mut Encoder<'_>, _: V) -> Result<(), EncodeError> { + w.varint(*self) } } -impl Decode for u32 -where - VarInt: Decode, -{ - fn decode(r: &mut R, version: V) -> Result { - let v = VarInt::decode(r, version)?; - let v = v.try_into().map_err(|_| DecodeError::BoundsExceeded)?; - Ok(v) +impl Decode for VarInt { + fn decode(r: &mut Decoder<'_>, _: V) -> Result { + r.varint() } } #[cfg(test)] mod tests { use super::*; - use crate::{ietf, lite}; - use bytes::Bytes; + + const DRAFT17: Form = Form::LeadingOnes { seven: false }; + const DRAFT18: Form = Form::LeadingOnes { seven: true }; + + fn encode(value: VarInt, form: Form) -> Result, BoundsExceeded> { + let (buf, len) = value.encode_form(form)?; + Ok(buf[..len].to_vec()) + } /// Test vectors from the draft-17 spec (Table 2: Example Integer Encodings), /// excluding the known-buggy example 4 (0xdd7f3e7d). @@ -588,26 +365,17 @@ mod tests { ]; for (bytes, expected) in cases { - // Test decoding - let mut buf = Bytes::from(bytes.to_vec()); - let decoded = VarInt::decode_leading_ones(&mut buf, ietf::Version::Draft17).expect("decode should succeed"); + let (decoded, len) = VarInt::decode_form(bytes, DRAFT17).expect("decode should succeed"); assert_eq!( decoded.into_inner(), *expected, "decode mismatch for bytes {bytes:02x?}" ); - assert_eq!(buf.len(), 0, "all bytes should be consumed for {bytes:02x?}"); - - // Test round-trip encode: - // - Skip non-minimal encoding (0x8025 for 37) - // - Skip u64::MAX which exceeds VarInt::MAX (2^62-1) but is decodable - if let Some(varint) = VarInt::from_u64(*expected) - && (bytes.len() == 1 || *expected != 37) - { - let mut encoded = Vec::new(); - varint - .encode_leading_ones(&mut encoded, ietf::Version::Draft17) - .expect("encode should succeed"); + assert_eq!(len, bytes.len(), "all bytes should be consumed for {bytes:02x?}"); + + // Skip the non-minimal encoding (0x8025 for 37); we only emit the minimal one. + if bytes.len() == 1 || *expected != 37 { + let encoded = encode(VarInt::from(*expected), DRAFT17).unwrap(); assert_eq!(&encoded, bytes, "encode mismatch for value {expected}"); } } @@ -616,12 +384,8 @@ mod tests { /// 11111100 (0xFC) is an invalid code point on draft-17 (allowed as 7-byte form on draft-18+). #[test] fn leading_ones_invalid_0xfc() { - let mut buf = Bytes::from_static(&[0xFC]); assert!( - matches!( - VarInt::decode_leading_ones(&mut buf, ietf::Version::Draft17), - Err(DecodeError::InvalidValue) - ), + matches!(VarInt::decode_form(&[0xFC], DRAFT17), Err(DecodeError::InvalidValue)), "0xFC should be rejected as invalid on draft-17" ); } @@ -638,37 +402,72 @@ mod tests { ]; for (value, expected_len) in cases { - let varint = VarInt::from_u64(value).expect("value should be representable as VarInt"); - let mut encoded = Vec::new(); - varint - .encode_leading_ones(&mut encoded, ietf::Version::Draft17) - .expect("leading-ones encode should succeed"); + let encoded = encode(VarInt::from(value), DRAFT17).unwrap(); assert_eq!( encoded.len(), expected_len, "unexpected encoded length for value {value}" ); - let mut bytes = Bytes::from(encoded); - let decoded = VarInt::decode_leading_ones(&mut bytes, ietf::Version::Draft17) - .expect("leading-ones decode should succeed"); + let (decoded, _) = VarInt::decode_form(&encoded, DRAFT17).expect("leading-ones decode should succeed"); assert_eq!(decoded.into_inner(), value, "round-trip mismatch for value {value}"); } } + /// Every length class of both forms survives a round trip, including the boundaries + /// of each class. + #[test] + fn every_length_round_trips() { + for bits in 0..64 { + for value in [1u64 << bits, (1u64 << bits) - 1, (1u64 << bits) + 1] { + let value = VarInt::from(value); + for form in [Form::Quic, DRAFT17, DRAFT18] { + let Ok(encoded) = encode(value, form) else { + assert_eq!(form, Form::Quic); + assert!(value > VarInt::MAX_QUIC); + continue; + }; + let (decoded, len) = VarInt::decode_form(&encoded, form).unwrap(); + assert_eq!((decoded, len), (value, encoded.len()), "{form:?} {value}"); + } + } + } + } + + /// The QUIC form stops at 2^62 - 1 and refuses anything past it rather than + /// truncating, while the leading-ones form carries the whole u64. + #[test] + fn quic_refuses_past_62_bits() { + let max = VarInt::MAX_QUIC; + assert_eq!(max.into_inner(), (1 << 62) - 1); + assert_eq!(encode(max, Form::Quic).unwrap(), [0xff; 8]); + + for value in [1u64 << 62, u64::MAX] { + assert_eq!(encode(VarInt::from(value), Form::Quic), Err(BoundsExceeded)); + assert!(matches!( + VarInt::from(value).encode_quic(&mut Vec::new()), + Err(EncodeError::BoundsExceeded) + )); + } + + for value in [(1u64 << 62) - 1, 1u64 << 62, u64::MAX] { + let encoded = encode(VarInt::from(value), DRAFT18).unwrap(); + assert_eq!(encoded.len(), 9); + assert_eq!(VarInt::decode_form(&encoded, DRAFT18).unwrap().0.into_inner(), value); + } + } + #[test] fn draft17_rejects_7_byte_varint() { // 1111110x prefix: invalid on draft-17. - let bytes = Bytes::from(vec![0xFC, 0, 0, 0, 0, 0, 0]); - let mut buf = bytes.clone(); - let err = VarInt::decode_leading_ones(&mut buf, ietf::Version::Draft17).unwrap_err(); + let err = VarInt::decode_form(&[0xFC, 0, 0, 0, 0, 0, 0], DRAFT17).unwrap_err(); assert!(matches!(err, DecodeError::InvalidValue)); } #[test] fn zigzag_roundtrip_small() { for n in [-3i64, -2, -1, 0, 1, 2, 3, 100, -100] { - let v = VarInt::from_zigzag(n).unwrap(); + let v = VarInt::from_zigzag(n); assert_eq!(v.to_zigzag(), n, "roundtrip failed for {}", n); } } @@ -676,45 +475,35 @@ mod tests { #[test] fn zigzag_small_values_compact() { // First few values should fit in 1 byte (varint range 0..=63 = top-2-bits tag 00). - assert_eq!(VarInt::from_zigzag(0).unwrap().into_inner(), 0); - assert_eq!(VarInt::from_zigzag(-1).unwrap().into_inner(), 1); - assert_eq!(VarInt::from_zigzag(1).unwrap().into_inner(), 2); - assert_eq!(VarInt::from_zigzag(-2).unwrap().into_inner(), 3); - assert_eq!(VarInt::from_zigzag(2).unwrap().into_inner(), 4); + assert_eq!(VarInt::from_zigzag(0).into_inner(), 0); + assert_eq!(VarInt::from_zigzag(-1).into_inner(), 1); + assert_eq!(VarInt::from_zigzag(1).into_inner(), 2); + assert_eq!(VarInt::from_zigzag(-2).into_inner(), 3); + assert_eq!(VarInt::from_zigzag(2).into_inner(), 4); } + /// Zigzag covers the whole i64 range; only the QUIC form bounds what goes on the wire. #[test] fn zigzag_roundtrip_boundary() { - // Boundary values in the valid input range [-2^61, 2^61 - 1]. - let max = (1i64 << 61) - 1; - let min = -(1i64 << 61); let mid = (1i64 << 30) + 17; - for n in [max, min, mid, -mid] { - let v = VarInt::from_zigzag(n).unwrap(); + for n in [i64::MAX, i64::MIN, (1i64 << 61) - 1, -(1i64 << 61), mid, -mid] { + let v = VarInt::from_zigzag(n); assert_eq!(v.to_zigzag(), n); } - } - #[test] - fn zigzag_out_of_range_rejected() { - // Values past the i61 boundary are out of varint range. - assert!(VarInt::from_zigzag(1i64 << 61).is_err()); - assert!(VarInt::from_zigzag(-(1i64 << 61) - 1).is_err()); - assert!(VarInt::from_zigzag(i64::MAX).is_err()); - assert!(VarInt::from_zigzag(i64::MIN).is_err()); + assert_eq!(VarInt::from_zigzag(i64::MIN), VarInt::MAX); + assert!(VarInt::from_zigzag(1i64 << 61) > VarInt::MAX_QUIC); + assert_eq!(VarInt::from_zigzag(-(1i64 << 61)), VarInt::MAX_QUIC); } #[test] fn zigzag_quic_varint_roundtrip() { // Encode a zigzag value through the QUIC varint wire format. for n in [-5000i64, 0, 100, -1, 1_000_000, -1_000_000] { - let v = VarInt::from_zigzag(n).unwrap(); - - let mut buf = bytes::BytesMut::new(); - v.encode(&mut buf, lite::Version::Lite01).unwrap(); - let mut bytes = buf.freeze(); - let decoded = VarInt::decode(&mut bytes, lite::Version::Lite01).unwrap(); + let v = VarInt::from_zigzag(n); + let bytes = v.encode_bytes(lite::Version::Lite01).unwrap(); + let (decoded, _) = VarInt::decode_slice(&bytes, lite::Version::Lite01).unwrap(); assert_eq!(decoded.to_zigzag(), n); } } @@ -731,8 +520,23 @@ mod tests { for shift in (0..48).step_by(8).rev() { bytes.push(((value >> shift) & 0xFF) as u8); } - let mut buf = Bytes::from(bytes); - let decoded = VarInt::decode_leading_ones(&mut buf, ietf::Version::Draft18).unwrap(); + let (decoded, _) = VarInt::decode_form(&bytes, DRAFT18).unwrap(); assert_eq!(decoded.into_inner(), value); } + + /// The Buf-based helpers other crates use read and write the same bytes, even when + /// the varint straddles two chunks. + #[test] + fn quic_helpers_match_the_codec() { + use bytes::Buf; + + let value = VarInt::from(0x1234_5678u64); + let mut out = Vec::new(); + value.encode_quic(&mut out).unwrap(); + assert_eq!(out, encode(value, Form::Quic).unwrap()); + + let mut split = (&out[..2]).chain(&out[2..]); + assert_eq!(VarInt::decode_quic(&mut split).unwrap(), value); + assert!(!split.has_remaining()); + } } diff --git a/rs/moq-net/src/coding/version.rs b/rs/moq-net/src/coding/version.rs index fc2e134c65..11000271ab 100644 --- a/rs/moq-net/src/coding/version.rs +++ b/rs/moq-net/src/coding/version.rs @@ -18,23 +18,17 @@ impl From for u64 { } } -impl Decode for Version -where - u64: Decode, -{ +impl Decode for Version { /// Decode the version number. - fn decode(r: &mut R, version: V) -> Result { - let v = u64::decode(r, version)?; - Ok(Self(v)) + fn decode(r: &mut Decoder<'_>, _: V) -> Result { + Ok(Self(r.varint()?.into_inner())) } } -impl Encode for Version -where - u64: Encode, -{ - fn encode(&self, w: &mut W, version: V) -> Result<(), EncodeError> { - self.0.encode(w, version) +impl Encode for Version { + fn encode(&self, w: &mut Encoder<'_>, _: V) -> Result<(), EncodeError> { + w.varint(self.0.into())?; + Ok(()) } } @@ -48,13 +42,10 @@ impl fmt::Debug for Version { #[derive(Clone, PartialEq, Eq, PartialOrd, Ord, Hash, Default)] pub struct Versions(Vec); -impl Decode for Versions -where - u64: Decode, -{ +impl Decode for Versions { /// Decode the version list. - fn decode(r: &mut R, version: V) -> Result { - let count = u64::decode(r, version)?; + fn decode(r: &mut Decoder<'_>, version: V) -> Result { + let count = r.varint()?.into_inner(); let mut vs = Vec::new(); for _ in 0..count { @@ -66,13 +57,10 @@ where } } -impl Encode for Versions -where - u64: Encode, -{ +impl Encode for Versions { /// Encode the version list. - fn encode(&self, w: &mut W, version: V) -> Result<(), EncodeError> { - (self.0.len() as u64).encode(w, version)?; + fn encode(&self, w: &mut Encoder<'_>, version: V) -> Result<(), EncodeError> { + w.varint(self.0.len().into())?; for v in &self.0 { v.encode(w, version)?; diff --git a/rs/moq-net/src/coding/writer.rs b/rs/moq-net/src/coding/writer.rs index 33ec964ae3..328cd21777 100644 --- a/rs/moq-net/src/coding/writer.rs +++ b/rs/moq-net/src/coding/writer.rs @@ -14,7 +14,9 @@ use crate::{Error, StreamError, coding::*, ietf}; /// cancelled wrapper) resumes mid-message instead of desynchronizing the stream. pub struct Writer { stream: Option, - buffer: bytes::BytesMut, + buffer: Vec, + /// How much of `buffer` has already hit the stream. + flushed: usize, version: V, } @@ -24,6 +26,7 @@ impl Writer { Self { stream: Some(stream), buffer: Default::default(), + flushed: 0, version, } } @@ -32,10 +35,11 @@ impl Writer { /// [`Self::poll_flush`]. An encode error leaves the buffer untouched. pub fn buffer + Debug>(&mut self, msg: &T) -> Result<(), Error> where - V: Clone, + V: Into + Copy, { let start = self.buffer.len(); - if let Err(err) = msg.encode(&mut self.buffer, self.version.clone()) { + let mut w = Encoder::new(&mut self.buffer, self.version.into()); + if let Err(err) = msg.encode(&mut w, self.version) { // Drop the partial encode: flushing it would corrupt the stream. self.buffer.truncate(start); return Err(err.into()); @@ -51,10 +55,14 @@ impl Writer { /// Poll until the write buffer has fully hit the stream. pub fn poll_flush(&mut self, cx: &mut Context<'_>) -> Poll> { - while !self.buffer.is_empty() { - ready!(self.stream.as_mut().unwrap().poll_write_buf(cx, &mut self.buffer)) + while self.flushed < self.buffer.len() { + let stream = self.stream.as_mut().unwrap(); + let n = ready!(stream.poll_write(cx, &self.buffer[self.flushed..])) .map_err(|err| self.version.transport_error(err))?; + self.flushed += n; } + self.buffer.clear(); + self.flushed = 0; Poll::Ready(Ok(())) } @@ -67,7 +75,7 @@ impl Writer { /// write) instead. pub async fn encode + Debug>(&mut self, msg: &T) -> Result<(), Error> where - V: Clone, + V: Into + Copy, { self.buffer(msg)?; std::future::poll_fn(|cx| self.poll_flush(cx)).await @@ -183,6 +191,7 @@ impl Writer { // We need to use an Option so Drop doesn't reset the stream. stream: self.stream.take(), buffer: std::mem::take(&mut self.buffer), + flushed: self.flushed, version, } } @@ -191,7 +200,7 @@ impl Writer { impl Writer { /// Encode an IETF `Message` to the stream, writing `[type_id][size][body]`. pub async fn encode_message(&mut self, msg: &T) -> Result<(), Error> { - self.buffer(&T::ID)?; + self.buffer(&VarInt::from(T::ID))?; self.encode(msg).await } } @@ -299,8 +308,8 @@ mod tests { struct Poison; impl Encode for Poison { - fn encode(&self, w: &mut W, _: crate::lite::Version) -> Result<(), EncodeError> { - w.put_slice(b"junk"); + fn encode(&self, w: &mut Encoder<'_>, _: crate::lite::Version) -> Result<(), EncodeError> { + w.slice(b"junk"); Err(EncodeError::BoundsExceeded) } } @@ -311,9 +320,9 @@ mod tests { fn a_failed_encode_leaves_no_partial_bytes() { let mut writer = Writer::new(SinkSend::new(Log::default()), crate::lite::Version::Lite05); - writer.buffer(&5u8).unwrap(); + writer.buffer(&VarInt::from_u32(5)).unwrap(); writer.buffer(&Poison).unwrap_err(); - writer.buffer(&7u8).unwrap(); + writer.buffer(&VarInt::from_u32(7)).unwrap(); let log = writer.stream.as_ref().unwrap().log.clone(); let mut cx = std::task::Context::from_waker(Waker::noop()); @@ -332,13 +341,13 @@ mod tests { ); let log = writer.stream.as_ref().unwrap().log.clone(); - writer.buffer(&5u8).unwrap(); + writer.buffer(&VarInt::from_u32(5)).unwrap(); let mut cx = std::task::Context::from_waker(Waker::noop()); assert!(writer.poll_flush(&mut cx).is_pending()); assert!(log.writes.lock().unwrap().is_empty()); // The bytes survive the Pending (and a second message queued behind them). - writer.buffer(&7u8).unwrap(); + writer.buffer(&VarInt::from_u32(7)).unwrap(); let Ok(mut open) = gate.write() else { panic!("gate closed") }; diff --git a/rs/moq-net/src/fuzz.rs b/rs/moq-net/src/fuzz.rs index cfb59f1bad..6b3c991066 100644 --- a/rs/moq-net/src/fuzz.rs +++ b/rs/moq-net/src/fuzz.rs @@ -10,11 +10,11 @@ //! Compiled only under `cfg(test)` or the `fuzz` feature, so none of this is part of //! the published API. See `fuzz/README.md` for the workflow. -use bytes::Buf; +use bytes::Bytes; use crate::{ Hops, Path, PathOwned, Pattern, - coding::{Decode, Encode, VarInt}, + coding::{Decode, Decoder, Encode, Encoder, Form, VarInt}, ietf, lite, path::Relative, }; @@ -78,17 +78,13 @@ fn select(data: &[u8], versions: usize) -> Option<(usize, u8, &[u8])> { /// at a version that cannot express it again (a duration past the varint range, a /// message a later draft dropped). What is never legitimate is emitting bytes we then /// refuse, or refuse to consume in full, since the peer's decoder is this same code. -/// -/// `stable` asks for the stronger check that the second encoding matches the first, -/// byte for byte. It is off wherever a parameter map reaches the wire, because those -/// encoders walk a `HashMap`, whose iteration order differs per instance. -fn roundtrip(data: &[u8], version: V, stable: bool) -> bool +/// The second encoding must also match the first byte for byte. +fn roundtrip(data: &[u8], version: V) -> bool where T: Decode + Encode, - V: Copy, + V: Copy + Into, { - let mut buf = data; - let Ok(decoded) = T::decode(&mut buf, version) else { + let Ok((decoded, _)) = T::decode_slice(data, version) else { return false; }; @@ -96,21 +92,29 @@ where return true; }; - let mut echo = first.clone(); - let decoded = T::decode(&mut echo, version).expect("could not decode our own encoding"); - assert!( - !echo.has_remaining(), - "our own encoding left {} bytes", - echo.remaining() - ); + let (decoded, used) = T::decode_slice(&first, version).expect("could not decode our own encoding"); + assert_eq!(used, first.len(), "our own encoding left {} bytes", first.len() - used); let Ok(second) = decoded.encode_bytes(version) else { panic!("could not re-encode what we just encoded"); }; + assert_eq!(first, second, "encoding is not stable"); - if stable { - assert_eq!(first, second, "encoding is not stable"); - } + true +} + +/// [`roundtrip`] for a datagram body, which runs to the end of its buffer. +fn datagram(data: &[u8], version: lite::Version) -> bool { + let Ok(decoded) = lite::Datagram::decode(Bytes::copy_from_slice(data), version) else { + return false; + }; + + let Ok(first) = decoded.encode_bytes(version) else { + return true; + }; + + let echo = lite::Datagram::decode(first, version).expect("could not decode our own encoding"); + assert_eq!(echo, decoded, "datagram did not survive a round trip"); true } @@ -126,31 +130,28 @@ pub fn lite_wire(data: &[u8]) -> bool { let version = LITE_VERSIONS[version]; let kind = kind % LITE_KINDS; - // SETUP is a parameter map, and so is the map itself. - let stable = !matches!(kind, 0 | 20); - match kind { - 0 => roundtrip::(rest, version, stable), - 1 => roundtrip::(rest, version, stable), - 2 => roundtrip::, _>(rest, version, stable), - 3 => roundtrip::, _>(rest, version, stable), - 4 => roundtrip::(rest, version, stable), - 5 => roundtrip::, _>(rest, version, stable), - 6 => roundtrip::, _>(rest, version, stable), - 7 => roundtrip::(rest, version, stable), - 8 => roundtrip::(rest, version, stable), - 9 => roundtrip::(rest, version, stable), - 10 => roundtrip::(rest, version, stable), - 11 => roundtrip::(rest, version, stable), - 12 => roundtrip::(rest, version, stable), - 13 => roundtrip::, _>(rest, version, stable), - 14 => roundtrip::(rest, version, stable), - 15 => roundtrip::, _>(rest, version, stable), - 16 => roundtrip::, _>(rest, version, stable), - 17 => roundtrip::(rest, version, stable), - 18 => roundtrip::(rest, version, stable), - 19 => roundtrip::(rest, version, stable), - 20 => roundtrip::(rest, version, stable), + 0 => roundtrip::(rest, version), + 1 => roundtrip::(rest, version), + 2 => roundtrip::, _>(rest, version), + 3 => roundtrip::, _>(rest, version), + 4 => roundtrip::(rest, version), + 5 => roundtrip::, _>(rest, version), + 6 => roundtrip::, _>(rest, version), + 7 => roundtrip::(rest, version), + 8 => roundtrip::(rest, version), + 9 => roundtrip::(rest, version), + 10 => roundtrip::(rest, version), + 11 => roundtrip::(rest, version), + 12 => roundtrip::(rest, version), + 13 => roundtrip::, _>(rest, version), + 14 => roundtrip::(rest, version), + 15 => roundtrip::, _>(rest, version), + 16 => roundtrip::, _>(rest, version), + 17 => roundtrip::(rest, version), + 18 => roundtrip::(rest, version), + 19 => datagram(rest, version), + 20 => roundtrip::(rest, version), _ => unreachable!("kind is taken modulo LITE_KINDS"), } } @@ -211,7 +212,7 @@ impl AnnounceWriter { lite::AnnounceBroadcast::EndedId { id: *id } } }; - msg.encode(data, self.version) + msg.encode(&mut Encoder::new(data, self.version.into()), self.version) .expect("could not encode an announcement"); } } @@ -229,8 +230,9 @@ pub fn encode_announces(announced: &[Announced], compress: bool) -> Vec { /// Decode and resolve an announce stream, as [`encode_announces`] writes it, stopping /// at the first message that fails to decode or resolve: a subscriber closes the /// session there. -pub fn decode_announces(mut data: &[u8], compress: bool) -> Vec { +pub fn decode_announces(data: &[u8], compress: bool) -> Vec { let version = announce_version(compress); + let mut data = Decoder::new(data, version.into()); let mut decoder = lite::AnnounceDecoder::default(); let mut resolved = Vec::new(); while let Ok(msg) = lite::AnnounceBroadcast::decode(&mut data, version) { @@ -274,50 +276,46 @@ pub fn ietf_wire(data: &[u8]) -> bool { let version = IETF_VERSIONS[version]; let kind = kind % IETF_KINDS; - // Draft-14 and draft-15 write a parameter map straight out of a `HashMap`; every - // later draft sorts by key first, so only these two are order-dependent. - let stable = !matches!(version, ietf::Version::Draft14 | ietf::Version::Draft15); - match kind { - 0 => roundtrip::, _>(rest, version, stable), - 1 => roundtrip::, _>(rest, version, stable), - 2 => roundtrip::(rest, version, stable), - 3 => roundtrip::, _>(rest, version, stable), - 4 => roundtrip::(rest, version, stable), - 5 => roundtrip::(rest, version, stable), - 6 => roundtrip::, _>(rest, version, stable), - 7 => roundtrip::, _>(rest, version, stable), - 8 => roundtrip::(rest, version, stable), - 9 => roundtrip::, _>(rest, version, stable), - 10 => roundtrip::, _>(rest, version, stable), - 11 => roundtrip::, _>(rest, version, stable), - 12 => roundtrip::, _>(rest, version, stable), - 13 => roundtrip::, _>(rest, version, stable), - 14 => roundtrip::(rest, version, stable), - 15 => roundtrip::, _>(rest, version, stable), - 16 => roundtrip::, _>(rest, version, stable), - 17 => roundtrip::, _>(rest, version, stable), - 18 => roundtrip::, _>(rest, version, stable), - 19 => roundtrip::(rest, version, stable), - 20 => roundtrip::, _>(rest, version, stable), - 21 => roundtrip::(rest, version, stable), - 22 => roundtrip::(rest, version, stable), - 23 => roundtrip::, _>(rest, version, stable), - 24 => roundtrip::, _>(rest, version, stable), - 25 => roundtrip::(rest, version, stable), - 26 => roundtrip::, _>(rest, version, stable), - 27 => roundtrip::(rest, version, stable), - 28 => roundtrip::, _>(rest, version, stable), - 29 => roundtrip::, _>(rest, version, stable), - 30 => roundtrip::(rest, version, stable), - 31 => roundtrip::(rest, version, stable), - 32 => roundtrip::(rest, version, stable), - 33 => roundtrip::, _>(rest, version, stable), - 34 => roundtrip::(rest, version, stable), - 35 => roundtrip::(rest, version, stable), - 36 => roundtrip::(rest, version, stable), - 37 => roundtrip::(rest, version, stable), - 38 => roundtrip::(rest, version, stable), + 0 => roundtrip::, _>(rest, version), + 1 => roundtrip::, _>(rest, version), + 2 => roundtrip::(rest, version), + 3 => roundtrip::, _>(rest, version), + 4 => roundtrip::(rest, version), + 5 => roundtrip::(rest, version), + 6 => roundtrip::, _>(rest, version), + 7 => roundtrip::, _>(rest, version), + 8 => roundtrip::(rest, version), + 9 => roundtrip::, _>(rest, version), + 10 => roundtrip::, _>(rest, version), + 11 => roundtrip::, _>(rest, version), + 12 => roundtrip::, _>(rest, version), + 13 => roundtrip::, _>(rest, version), + 14 => roundtrip::(rest, version), + 15 => roundtrip::, _>(rest, version), + 16 => roundtrip::, _>(rest, version), + 17 => roundtrip::, _>(rest, version), + 18 => roundtrip::, _>(rest, version), + 19 => roundtrip::(rest, version), + 20 => roundtrip::, _>(rest, version), + 21 => roundtrip::(rest, version), + 22 => roundtrip::(rest, version), + 23 => roundtrip::, _>(rest, version), + 24 => roundtrip::, _>(rest, version), + 25 => roundtrip::(rest, version), + 26 => roundtrip::, _>(rest, version), + 27 => roundtrip::(rest, version), + 28 => roundtrip::, _>(rest, version), + 29 => roundtrip::, _>(rest, version), + 30 => roundtrip::(rest, version), + 31 => roundtrip::(rest, version), + 32 => roundtrip::(rest, version), + 33 => roundtrip::, _>(rest, version), + 34 => roundtrip::(rest, version), + 35 => roundtrip::(rest, version), + 36 => roundtrip::(rest, version), + 37 => roundtrip::(rest, version), + 38 => roundtrip::(rest, version), _ => unreachable!("kind is taken modulo IETF_KINDS"), } } @@ -328,9 +326,8 @@ pub fn ietf_wire(data: &[u8]) -> bool { /// the QUIC two-bit length tag, while draft-17+ counts leading ones, and the two /// disagree about which byte sequences are even legal. /// -/// The decoded value is deliberately not asserted to be within [`VarInt::MAX`]: the -/// leading-ones form spans the full `u64` by design, so a 9-byte encoding decodes -/// above the 62-bit ceiling and only fails when re-encoded for a QUIC-form version. +/// The leading-ones form spans the full `u64`, while the QUIC form stops at +/// [`VarInt::MAX_QUIC`]; a value always re-encodes in the form it was read in. pub fn varint(data: &[u8]) -> bool { let Some((&selector, rest)) = data.split_first() else { return false; @@ -342,28 +339,23 @@ pub fn varint(data: &[u8]) -> bool { _ => IETF_VERSIONS[(selector as usize / 2) % IETF_VERSIONS.len()].into(), }; - let mut buf = rest; - let Ok(value) = VarInt::decode(&mut buf, version) else { + let Ok((value, _)) = VarInt::decode_slice(rest, version) else { return false; }; - // Zigzag is a pure mapping on top of the wire value, so it must round-trip - // whenever the signed value is back in range. + // Zigzag is a bijection on top of the wire value, so it must round-trip. let signed = value.to_zigzag(); - if let Ok(mapped) = VarInt::from_zigzag(signed) { - assert_eq!(mapped.to_zigzag(), signed, "zigzag is not its own inverse"); - } + assert_eq!(VarInt::from_zigzag(signed), value, "zigzag is not its own inverse"); - let Ok(encoded) = value.encode_bytes(version) else { - return true; - }; - - let mut echo = encoded.clone(); - let again = VarInt::decode(&mut echo, version).expect("could not decode our own encoding"); - assert!( - !echo.has_remaining(), + let encoded = value + .encode_bytes(version) + .expect("a varint re-encodes in the form it was read in"); + let (again, used) = VarInt::decode_slice(&encoded, version).expect("could not decode our own encoding"); + assert_eq!( + used, + encoded.len(), "our own encoding left {} bytes", - echo.remaining() + encoded.len() - used ); assert_eq!(value, again, "varint did not survive a round trip"); @@ -461,63 +453,69 @@ impl Messages { /// Append the moq-lite messages to `out`. pub fn encode_lite(&self, out: &mut Vec) { let v = BENCH_LITE; - self.lite_subscribe.encode(out, v).unwrap(); - self.lite_update.encode(out, v).unwrap(); - self.lite_start.encode(out, v).unwrap(); - self.lite_info.encode(out, v).unwrap(); - self.lite_group.encode(out, v).unwrap(); + let w = &mut Encoder::new(out, v.into()); + self.lite_subscribe.encode(w, v).unwrap(); + self.lite_update.encode(w, v).unwrap(); + self.lite_start.encode(w, v).unwrap(); + self.lite_info.encode(w, v).unwrap(); + self.lite_group.encode(w, v).unwrap(); } /// Decode what [`Self::encode_lite`] wrote. - pub fn decode_lite(&self, mut data: &[u8]) { + pub fn decode_lite(&self, data: &[u8]) { let v = BENCH_LITE; - lite::Subscribe::decode(&mut data, v).unwrap(); - lite::SubscribeUpdate::decode(&mut data, v).unwrap(); - lite::SubscribeResponse::decode(&mut data, v).unwrap(); - lite::TrackInfo::decode(&mut data, v).unwrap(); - lite::Group::decode(&mut data, v).unwrap(); - assert!(data.is_empty()); + let r = &mut Decoder::new(data, v.into()); + lite::Subscribe::decode(r, v).unwrap(); + lite::SubscribeUpdate::decode(r, v).unwrap(); + lite::SubscribeResponse::decode(r, v).unwrap(); + lite::TrackInfo::decode(r, v).unwrap(); + lite::Group::decode(r, v).unwrap(); + assert!(r.is_empty()); } /// Append the moq-transport messages to `out`. pub fn encode_ietf(&self, out: &mut Vec) { let v = BENCH_IETF; - self.ietf_subscribe.encode(out, v).unwrap(); - self.ietf_ok.encode(out, v).unwrap(); - self.ietf_group.encode(out, v).unwrap(); + let w = &mut Encoder::new(out, v.into()); + self.ietf_subscribe.encode(w, v).unwrap(); + self.ietf_ok.encode(w, v).unwrap(); + self.ietf_group.encode(w, v).unwrap(); } /// Decode what [`Self::encode_ietf`] wrote. - pub fn decode_ietf(&self, mut data: &[u8]) { + pub fn decode_ietf(&self, data: &[u8]) { let v = BENCH_IETF; - ietf::Subscribe::decode(&mut data, v).unwrap(); - ietf::SubscribeOk::decode(&mut data, v).unwrap(); - ietf::GroupHeader::decode(&mut data, v).unwrap(); - assert!(data.is_empty()); + let r = &mut Decoder::new(data, v.into()); + ietf::Subscribe::decode(r, v).unwrap(); + ietf::SubscribeOk::decode(r, v).unwrap(); + ietf::GroupHeader::decode(r, v).unwrap(); + assert!(r.is_empty()); + } +} + +/// The varint form of moq-lite (`ietf == false`) or of a leading-ones moq-transport draft. +fn bench_form(ietf: bool) -> Form { + match ietf { + false => BENCH_LITE.into(), + true => BENCH_IETF.into(), } } /// Append `values` as varints in the wire form of moq-lite (`ietf == false`) or of a /// leading-ones moq-transport draft. pub fn encode_varints(values: &[u64], ietf: bool, out: &mut Vec) { + let mut w = Encoder::new(out, bench_form(ietf)); for value in values { - let value = VarInt::try_from(*value).unwrap(); - match ietf { - false => value.encode(out, BENCH_LITE).unwrap(), - true => value.encode(out, BENCH_IETF).unwrap(), - } + w.varint((*value).into()).unwrap(); } } /// Sum the varints [`encode_varints`] wrote. -pub fn decode_varints(mut data: &[u8], ietf: bool) -> u64 { +pub fn decode_varints(data: &[u8], ietf: bool) -> u64 { + let mut r = Decoder::new(data, bench_form(ietf)); let mut sum = 0u64; - while !data.is_empty() { - let value = match ietf { - false => VarInt::decode(&mut data, BENCH_LITE).unwrap(), - true => VarInt::decode(&mut data, BENCH_IETF).unwrap(), - }; - sum = sum.wrapping_add(value.into_inner()); + while !r.is_empty() { + sum = sum.wrapping_add(r.varint().unwrap().into_inner()); } sum } @@ -579,12 +577,12 @@ pub fn path(data: &[u8]) -> bool { // The wire form is the same path back, whenever the path is expressible at all. let version = lite::Version::Lite05; if let Ok(encoded) = target.encode_bytes(version) { - let mut echo = encoded; - let decoded = Path::decode(&mut echo, version).expect("could not decode our own encoding"); - assert!( - !echo.has_remaining(), + let (decoded, used) = Path::decode_slice(&encoded, version).expect("could not decode our own encoding"); + assert_eq!( + used, + encoded.len(), "our own encoding left {} bytes", - echo.remaining() + encoded.len() - used ); assert_eq!(decoded, target, "path did not survive a round trip"); } diff --git a/rs/moq-net/src/ietf/adapter.rs b/rs/moq-net/src/ietf/adapter.rs index 2ec6597457..578b805d3e 100644 --- a/rs/moq-net/src/ietf/adapter.rs +++ b/rs/moq-net/src/ietf/adapter.rs @@ -8,7 +8,7 @@ use bytes::{Buf, BufMut, Bytes, BytesMut}; use crate::{ Error, PathOwned, - coding::{Decode, Encode, Reader, Writer}, + coding::{Decode, Decoder, Encoder, Reader, VarInt, Writer}, ietf::{self, RequestId}, }; @@ -274,31 +274,29 @@ impl OutgoingRegistration { /// Try to parse the request_id (and optionally namespace) from the accumulated bytes. /// Returns Ok(None) if not enough data yet, Err if the message is malformed. fn try_parse(&self) -> Result, crate::Error> { - let mut cursor = std::io::Cursor::new(&self.buf); - let Ok(type_id) = u64::decode(&mut cursor, self.version) else { + let mut r = Decoder::new(&self.buf, self.version.into()); + let Ok(type_id) = r.varint() else { return Ok(None); }; - let Ok(size) = u16::decode(&mut cursor, self.version) else { + let Ok(size) = r.u16() else { return Ok(None); }; // We know the full message size now: header bytes + body. - let header_len = cursor.position() as usize; - let message_len = header_len + size as usize; - if self.buf.len() < message_len { + let Ok(mut body) = r.sub(size as usize) else { return Ok(None); - } + }; // We have enough bytes for the full message; decoding must succeed. - let request_id = RequestId::decode(&mut cursor, self.version)?; + let request_id = RequestId::decode(&mut body, self.version)?; // For PublishNamespace, also extract the namespace for reverse lookup. - if type_id == ietf::PublishNamespace::ID { + if type_id.into_inner() == ietf::PublishNamespace::ID { if self.version == Version::Draft17 { // v17 has required_request_id_delta after request_id - let _ = u64::decode(&mut cursor, self.version); + let _ = body.varint(); } - if let Ok(ns) = crate::ietf::namespace::decode_namespace(&mut cursor, self.version) { + if let Ok(ns) = crate::ietf::namespace::decode_namespace(&mut body) { self.shared.namespaces.register(Direction::Outgoing, ns, request_id); } } @@ -781,26 +779,19 @@ impl ControlStreamAdapter { timeout: timeout_ms, }; - let mut body = BytesMut::new(); - if let Err(err) = msg.encode_msg(&mut body, version) { + // The size prefix is a u16 on a stream shared with every other request, so the + // encode refuses a body that would wrap it and desynchronize the framing for all. + let mut raw = Vec::new(); + let mut w = Encoder::new(&mut raw, version.into()); + if let Err(err) = w + .varint(crate::ietf::GoAway::ID.into()) + .and_then(|()| msg.encode(&mut w, version)) + { tracing::warn!(%err, "failed to encode goaway"); return; } - // The size prefix is a u16 on a stream shared with every other request, so a - // wrapping cast here would desynchronize the framing for all of them. - let Ok(size) = u16::try_from(body.len()) else { - tracing::warn!(len = body.len(), "goaway too large for the control stream"); - return; - }; - - let mut raw = BytesMut::new(); - if crate::ietf::GoAway::ID.encode(&mut raw, version).is_err() || size.encode(&mut raw, version).is_err() { - return; - } - raw.extend_from_slice(&body); - - if !self.shared.control.push(raw.freeze()) { + if !self.shared.control.push(raw.into()) { tracing::debug!("control stream closed; goaway not sent"); } } @@ -821,17 +812,15 @@ impl ControlStreamAdapter { goaway: crate::goaway::Protocol, ) -> Result<(), Error> { loop { - let type_id: u64 = match reader.decode_maybe().await? { - Some(id) => id, + let type_id = match reader.decode_maybe::().await? { + Some(id) => id.into_inner(), None => return Ok(()), }; - let size: u16 = reader.decode::().await?; - - let body = reader.read_exact(size as usize).await?; + let body = reader.decode::().await?.0; // Reconstruct raw message bytes: [type_id][size][body] - let raw = encode_raw(type_id, size, &body, self.version); + let raw = encode_raw(type_id, &body, self.version); // Classify and route match classify(type_id, &body, self.version, &self.shared.namespaces)? { @@ -841,7 +830,7 @@ impl ControlStreamAdapter { Route::MaxRequestId(max) => self.control.max_request_id(max), Route::Ignore => {} Route::GoAway => { - let mut data = body; + let mut data = Decoder::new(&body, self.version.into()); let msg = crate::ietf::GoAway::decode_msg(&mut data, self.version)?; tracing::info!(message = ?msg, "received GOAWAY"); @@ -1064,8 +1053,8 @@ fn lookup_namespace_request_id( namespaces: &Namespaces, direction: Direction, ) -> Result, Error> { - let mut cursor = std::io::Cursor::new(body); - let ns = crate::ietf::namespace::decode_namespace(&mut cursor, version)?; + let mut r = Decoder::new(body, version.into()); + let ns = crate::ietf::namespace::decode_namespace(&mut r)?; Ok(namespaces.get(direction, &ns)) } @@ -1167,18 +1156,18 @@ enum Route { } /// Encode raw message bytes as [type_id varint][size u16][body]. -fn encode_raw(type_id: u64, size: u16, body: &Bytes, version: Version) -> Bytes { - let mut buf = BytesMut::new(); - type_id.encode(&mut buf, version).expect("encode type_id"); - size.encode(&mut buf, version).expect("encode size"); - buf.extend_from_slice(body); - buf.freeze() +fn encode_raw(type_id: u64, body: &Bytes, version: Version) -> Bytes { + let mut buf = Vec::new(); + let mut w = Encoder::new(&mut buf, version.into()); + w.varint(type_id.into()).expect("type_id was read from the same wire"); + w.u16(u16::try_from(body.len()).expect("body was read with a u16 size")); + w.slice(body); + buf.into() } /// Decode just the request_id from the beginning of a message body. fn decode_request_id(body: &Bytes, version: Version) -> Result { - let mut cursor = std::io::Cursor::new(body); - let request_id = RequestId::decode(&mut cursor, version)?; + let (request_id, _) = RequestId::decode_slice(body, version)?; Ok(request_id) } @@ -1190,28 +1179,26 @@ fn decode_response_request_id(body: &Bytes, version: Version) -> Result Result { - let mut cursor = std::io::Cursor::new(body); + let mut r = Decoder::new(body, version.into()); // Skip request_id - let _request_id = RequestId::decode(&mut cursor, version)?; + let _request_id = RequestId::decode(&mut r, version)?; // v17 has required_request_id_delta if version == Version::Draft17 { - let _ = u64::decode(&mut cursor, version)?; + r.varint()?; } - let ns = crate::ietf::namespace::decode_namespace(&mut cursor, version)?; + let ns = crate::ietf::namespace::decode_namespace(&mut r)?; Ok(ns.into_owned()) } #[cfg(test)] mod tests { use super::*; + use crate::coding::Encode; use crate::transport::poll::{RecvStream as _, SendStream as _}; - use bytes::BytesMut; use futures::FutureExt as _; fn make_body_with_request_id(id: u64, version: Version) -> Bytes { - let mut buf = BytesMut::new(); - RequestId(id).encode(&mut buf, version).unwrap(); - buf.freeze() + RequestId(id).encode_bytes(version).unwrap() } /// Classify against an empty namespace map, for the messages that don't use it. @@ -1319,14 +1306,13 @@ mod tests { fn test_encode_raw_roundtrip() { let version = Version::Draft15; let body = Bytes::from_static(b"hello"); - let raw = encode_raw(0x03, 5, &body, version); + let raw = encode_raw(0x03, &body, version); // Decode the raw bytes - let mut cursor = std::io::Cursor::new(&raw[..]); - let type_id = u64::decode(&mut cursor, version).unwrap(); - let size = u16::decode(&mut cursor, version).unwrap(); - assert_eq!(type_id, 0x03); - assert_eq!(size, 5); + let mut r = Decoder::new(&raw, version.into()); + assert_eq!(r.varint().unwrap().into_inner(), 0x03); + assert_eq!(r.u16().unwrap(), 5); + assert_eq!(r.rest(), b"hello"); } #[tokio::test] @@ -1404,15 +1390,16 @@ mod tests { /// Encode a message body (no type_id/size header). fn encode_body(msg: &M, version: Version) -> Bytes { - let mut buf = BytesMut::new(); - msg.encode_msg(&mut buf, version).unwrap(); - buf.freeze() + let mut buf = Vec::new(); + msg.encode_msg(&mut Encoder::new(&mut buf, version.into()), version) + .unwrap(); + bytes::Bytes::from(buf) } /// Encode a full control message: [type_id][size][body]. fn encode_msg(msg: &M, version: Version) -> Bytes { let body = encode_body(msg, version); - encode_raw(M::ID, body.len() as u16, &body, version) + encode_raw(M::ID, &body, version) } fn publish_namespace(request_id: RequestId, namespace: &str) -> ietf::PublishNamespace<'_> { diff --git a/rs/moq-net/src/ietf/cluster.rs b/rs/moq-net/src/ietf/cluster.rs index e5e35a3793..3182cfe33b 100644 --- a/rs/moq-net/src/ietf/cluster.rs +++ b/rs/moq-net/src/ietf/cluster.rs @@ -15,9 +15,7 @@ //! [`crate::origin::Route`]); this module is only the moq-transport binding. //! Negotiated on draft-17+ only, where SETUP is a Key-Value-Pair block. -use bytes::Buf; - -use crate::coding::{Decode, DecodeError, Encode, EncodeError}; +use crate::coding::{Decode, DecodeError, Decoder, Encode, EncodeError, Encoder}; use crate::{Hop, Hops}; use super::{Param, Version}; @@ -94,22 +92,22 @@ impl HopPath { } impl Param for HopPath { - fn param_encode(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn param_encode(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { // The entries fill the value, so they are written with no count and framed by // the parameter's own length prefix. let mut buf = Vec::new(); + let mut inner = Encoder::new(&mut buf, w.form()); for hop in &self.0 { - hop.encode(&mut buf, version)?; + hop.encode(&mut inner, version)?; } - buf.encode(w, version) + w.bytes(&buf) } - fn param_decode(r: &mut R, version: Version) -> Result { - let value = Vec::::decode(r, version)?; - let mut buf = bytes::Bytes::from(value); + fn param_decode(r: &mut Decoder<'_>, version: Version) -> Result { + let mut buf = Decoder::new(r.bytes()?, r.form()); let mut hops = Hops::new(); - while buf.has_remaining() { + while !buf.is_empty() { // A short read here means the entries did not exactly fill the length. hops.push(Hop::decode(&mut buf, version)?)?; } @@ -264,7 +262,6 @@ pub fn peer_into_setup(params: &mut super::Parameters, self_hop: Hop, cost: Opti #[cfg(test)] mod tests { use super::*; - use bytes::BytesMut; const VERSION: Version = Version::Draft19; @@ -284,11 +281,12 @@ mod tests { } fn round_trip(path: &HopPath) -> Result { - let mut buf = BytesMut::new(); - path.param_encode(&mut buf, VERSION).unwrap(); - let mut bytes = buf.freeze(); - let decoded = HopPath::param_decode(&mut bytes, VERSION)?; - assert!(!bytes.has_remaining(), "trailing bytes after decode"); + let mut buf = Vec::new(); + path.param_encode(&mut Encoder::new(&mut buf, VERSION.into()), VERSION) + .unwrap(); + let mut bytes = bytes::Bytes::from(buf); + let decoded = crate::coding::decode_buf(&mut bytes, VERSION, HopPath::param_decode)?; + assert!(bytes.is_empty(), "trailing bytes after decode"); Ok(decoded) } @@ -317,8 +315,10 @@ mod tests { fn hop_path_has_no_inner_count() { // The entries fill the parameter length; a count would make us unreadable to // every other implementation. One byte of length plus one byte per small varint. - let mut buf = BytesMut::new(); - hop_path(&[1, 2, 3]).param_encode(&mut buf, VERSION).unwrap(); + let mut buf = Vec::new(); + hop_path(&[1, 2, 3]) + .param_encode(&mut Encoder::new(&mut buf, VERSION.into()), VERSION) + .unwrap(); assert_eq!(buf.to_vec(), vec![0x03, 0x01, 0x02, 0x03]); } @@ -336,15 +336,17 @@ mod tests { // no longer lets one be built: only a non-conforming sender produces this. let mut value = Vec::new(); for id in [4u64, 8, 4] { - hop(id).encode(&mut value, VERSION).unwrap(); + hop(id) + .encode(&mut Encoder::new(&mut value, VERSION.into()), VERSION) + .unwrap(); } - let mut buf = BytesMut::new(); - value.encode(&mut buf, VERSION).unwrap(); + let mut buf = Vec::new(); + Encoder::new(&mut buf, VERSION.into()).bytes(&value).unwrap(); - let mut bytes = buf.freeze(); + let mut bytes = bytes::Bytes::from(buf); assert!(matches!( - HopPath::param_decode(&mut bytes, VERSION), + crate::coding::decode_buf(&mut bytes, VERSION, HopPath::param_decode), Err(DecodeError::InvalidValue) )); } @@ -379,11 +381,13 @@ mod tests { #[test] fn hop_path_rejects_empty() { // The list always has at least one entry: the original publisher. - let mut buf = BytesMut::new(); - HopPath::default().param_encode(&mut buf, VERSION).unwrap(); - let mut bytes = buf.freeze(); + let mut buf = Vec::new(); + HopPath::default() + .param_encode(&mut Encoder::new(&mut buf, VERSION.into()), VERSION) + .unwrap(); + let mut bytes = bytes::Bytes::from(buf); assert!(matches!( - HopPath::param_decode(&mut bytes, VERSION), + crate::coding::decode_buf(&mut bytes, VERSION, HopPath::param_decode), Err(DecodeError::InvalidValue) )); } @@ -393,15 +397,17 @@ mod tests { // Entries must exactly fill Length. Chop the last byte off a multi-byte hop id // and shrink the length to match, so the value ends mid-varint. let mut value = Vec::new(); - hop(300).encode(&mut value, VERSION).unwrap(); + hop(300) + .encode(&mut Encoder::new(&mut value, VERSION.into()), VERSION) + .unwrap(); assert!(value.len() > 1, "300 should not fit in one byte"); value.pop(); - let mut buf = BytesMut::new(); - value.encode(&mut buf, VERSION).unwrap(); + let mut buf = Vec::new(); + Encoder::new(&mut buf, VERSION.into()).bytes(&value).unwrap(); - let mut bytes = buf.freeze(); - assert!(HopPath::param_decode(&mut bytes, VERSION).is_err()); + let mut bytes = bytes::Bytes::from(buf); + assert!(crate::coding::decode_buf(&mut bytes, VERSION, HopPath::param_decode).is_err()); } #[test] @@ -501,10 +507,12 @@ mod tests { let mut params = super::super::Parameters::default(); peer_into_setup(&mut params, self_hop, cost, VERSION); - let mut buf = BytesMut::new(); - params.encode(&mut buf, VERSION).unwrap(); - let mut bytes = buf.freeze(); - let decoded = super::super::Parameters::decode(&mut bytes, VERSION).unwrap(); + let mut buf = Vec::new(); + params + .encode(&mut Encoder::new(&mut buf, VERSION.into()), VERSION) + .unwrap(); + let mut bytes = bytes::Bytes::from(buf); + let decoded = crate::coding::decode_buf(&mut bytes, VERSION, super::super::Parameters::decode).unwrap(); let peer = peer_from_setup(&decoded, VERSION).unwrap(); assert_eq!(peer.hop, Some(self_hop)); @@ -519,8 +527,10 @@ mod tests { let mut params = super::super::Parameters::default(); peer_into_setup(&mut params, hop(42), None, VERSION); - let mut buf = BytesMut::new(); - params.encode(&mut buf, VERSION).unwrap(); + let mut buf = Vec::new(); + params + .encode(&mut Encoder::new(&mut buf, VERSION.into()), VERSION) + .unwrap(); assert_eq!(buf.to_vec(), vec![0xC4, 0x0B, 0x54, 0x2A]); } @@ -539,10 +549,12 @@ mod tests { let mut params = super::super::Parameters::default(); params.set_bytes(super::super::ParameterBytes::Unknown(0x40B55), vec![0x2A]); - let mut buf = BytesMut::new(); - params.encode(&mut buf, VERSION).unwrap(); - let mut bytes = buf.freeze(); - let decoded = super::super::Parameters::decode(&mut bytes, VERSION).unwrap(); + let mut buf = Vec::new(); + params + .encode(&mut Encoder::new(&mut buf, VERSION.into()), VERSION) + .unwrap(); + let mut bytes = bytes::Bytes::from(buf); + let decoded = crate::coding::decode_buf(&mut bytes, VERSION, super::super::Parameters::decode).unwrap(); let peer = peer_from_setup(&decoded, VERSION).unwrap(); assert!(!peer.negotiated()); diff --git a/rs/moq-net/src/ietf/fetch.rs b/rs/moq-net/src/ietf/fetch.rs index 5987d8efbc..5a126ee80f 100644 --- a/rs/moq-net/src/ietf/fetch.rs +++ b/rs/moq-net/src/ietf/fetch.rs @@ -2,7 +2,7 @@ use std::borrow::Cow; use crate::{ Path, - coding::{Decode, DecodeError, Encode, EncodeError}, + coding::{Decode, DecodeError, Decoder, Encode, EncodeError, Encoder, VarInt}, ietf::{ GroupOrder, Location, Parameters, RequestId, namespace::{decode_namespace, encode_namespace}, @@ -33,7 +33,7 @@ pub enum FetchType<'a> { } impl Encode for FetchType<'_> { - fn encode(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { match self { FetchType::Standalone { namespace, @@ -41,9 +41,9 @@ impl Encode for FetchType<'_> { start, end, } => { - 1u8.encode(w, version)?; - encode_namespace(w, namespace, version)?; - track.encode(w, version)?; + w.u8(1u8); + encode_namespace(w, namespace)?; + w.string(track)?; start.encode(w, version)?; end.encode(w, version)?; } @@ -51,17 +51,17 @@ impl Encode for FetchType<'_> { subscriber_request_id, group_offset, } => { - 2u8.encode(w, version)?; + w.u8(2u8); subscriber_request_id.encode(w, version)?; - group_offset.encode(w, version)?; + w.varint(VarInt::from(*group_offset))?; } FetchType::AbsoluteJoining { subscriber_request_id, group_id, } => { - 3u8.encode(w, version)?; + w.u8(3u8); subscriber_request_id.encode(w, version)?; - group_id.encode(w, version)?; + w.varint(VarInt::from(*group_id))?; } } Ok(()) @@ -69,12 +69,12 @@ impl Encode for FetchType<'_> { } impl Decode for FetchType<'_> { - fn decode(buf: &mut B, version: Version) -> Result { - let fetch_type = u64::decode(buf, version)?; + fn decode(buf: &mut Decoder<'_>, version: Version) -> Result { + let fetch_type = buf.varint()?.into_inner(); Ok(match fetch_type { 0x1 => { - let namespace = decode_namespace(buf, version)?; - let track = Cow::::decode(buf, version)?; + let namespace = decode_namespace(buf)?; + let track = Cow::Owned(buf.string()?); let start = Location::decode(buf, version)?; let end = Location::decode(buf, version)?; FetchType::Standalone { @@ -86,7 +86,7 @@ impl Decode for FetchType<'_> { } 0x2 => { let subscriber_request_id = RequestId::decode(buf, version)?; - let group_offset = u64::decode(buf, version)?; + let group_offset = buf.varint()?.into_inner(); FetchType::RelativeJoining { subscriber_request_id, group_offset, @@ -94,7 +94,7 @@ impl Decode for FetchType<'_> { } 0x3 => { let subscriber_request_id = RequestId::decode(buf, version)?; - let group_id = u64::decode(buf, version)?; + let group_id = buf.varint()?.into_inner(); FetchType::AbsoluteJoining { subscriber_request_id, group_id, @@ -116,18 +116,18 @@ pub struct Fetch<'a> { impl Message for Fetch<'_> { const ID: u64 = 0x16; - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.request_id.encode(w, version)?; if version == Version::Draft17 { - 0u64.encode(w, version)?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) + w.varint(VarInt::from(0u64))?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) } match version { Version::Draft14 => { - self.subscriber_priority.encode(w, version)?; + w.u8(self.subscriber_priority); self.group_order.encode(w, version)?; self.fetch_type.encode(w, version)?; - 0u8.encode(w, version)?; // no parameters + w.u8(0u8); // no parameters } _ => { self.fetch_type.encode(w, version)?; @@ -140,15 +140,15 @@ impl Message for Fetch<'_> { Ok(()) } - fn decode_msg(buf: &mut B, version: Version) -> Result { + fn decode_msg(buf: &mut Decoder<'_>, version: Version) -> Result { let request_id = RequestId::decode(buf, version)?; if version == Version::Draft17 { - let _required_request_id_delta = u64::decode(buf, version)?; + let _required_request_id_delta = buf.varint()?.into_inner(); } match version { Version::Draft14 => { - let subscriber_priority = u8::decode(buf, version)?; + let subscriber_priority = buf.u8()?; let group_order = GroupOrder::decode(buf, version)?; let fetch_type = FetchType::decode(buf, version)?; let _params = Parameters::decode(buf, version)?; @@ -190,7 +190,7 @@ pub struct FetchOk { impl Message for FetchOk { const ID: u64 = 0x18; - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { if matches!(version, Version::Draft14 | Version::Draft15 | Version::Draft16) { self.request_id .expect("request_id required for draft14-16") @@ -202,14 +202,14 @@ impl Message for FetchOk { match version { Version::Draft14 => { self.group_order.encode(w, version)?; - self.end_of_track.encode(w, version)?; + w.bool(self.end_of_track); self.end_location.encode(w, version)?; - 0u8.encode(w, version)?; // no parameters + w.u8(0u8); // no parameters } _ => { // GROUP_ORDER is not a legal FETCH_OK parameter in any draft after 14; the order // of the response is whatever the FETCH asked for. - self.end_of_track.encode(w, version)?; + w.bool(self.end_of_track); self.end_location.encode(w, version)?; encode_params!(w, version,); } @@ -217,7 +217,7 @@ impl Message for FetchOk { Ok(()) } - fn decode_msg(buf: &mut B, version: Version) -> Result { + fn decode_msg(buf: &mut Decoder<'_>, version: Version) -> Result { let request_id = if matches!(version, Version::Draft14 | Version::Draft15 | Version::Draft16) { Some(RequestId::decode(buf, version)?) } else { @@ -227,7 +227,7 @@ impl Message for FetchOk { match version { Version::Draft14 => { let group_order = GroupOrder::decode(buf, version)?; - let end_of_track = bool::decode(buf, version)?; + let end_of_track = buf.bool()?; let end_location = Location::decode(buf, version)?; let _params = Parameters::decode(buf, version)?; Ok(Self { @@ -238,7 +238,7 @@ impl Message for FetchOk { }) } _ => { - let end_of_track = bool::decode(buf, version)?; + let end_of_track = buf.bool()?; let end_location = Location::decode(buf, version)?; // GROUP_ORDER isn't legal here, but keep accepting it so a peer that still sends // it doesn't have its session torn down over a hint. @@ -272,17 +272,17 @@ pub struct FetchError<'a> { impl Message for FetchError<'_> { const ID: u64 = 0x19; - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.request_id.encode(w, version)?; - self.error_code.encode(w, version)?; - self.reason_phrase.encode(w, version)?; + w.varint(VarInt::from(self.error_code))?; + w.string(&self.reason_phrase)?; Ok(()) } - fn decode_msg(buf: &mut B, version: Version) -> Result { + fn decode_msg(buf: &mut Decoder<'_>, version: Version) -> Result { let request_id = RequestId::decode(buf, version)?; - let error_code = u64::decode(buf, version)?; - let reason_phrase = Cow::::decode(buf, version)?; + let error_code = buf.varint()?.into_inner(); + let reason_phrase = Cow::Owned(buf.string()?); Ok(Self { request_id, error_code, @@ -298,12 +298,12 @@ pub struct FetchCancel { impl Message for FetchCancel { const ID: u64 = 0x17; - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.request_id.encode(w, version)?; Ok(()) } - fn decode_msg(buf: &mut B, version: Version) -> Result { + fn decode_msg(buf: &mut Decoder<'_>, version: Version) -> Result { let request_id = RequestId::decode(buf, version)?; Ok(Self { request_id }) } @@ -319,14 +319,14 @@ impl FetchHeader { } impl Encode for FetchHeader { - fn encode(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.request_id.encode(w, version)?; Ok(()) } } impl Decode for FetchHeader { - fn decode(buf: &mut B, version: Version) -> Result { + fn decode(buf: &mut Decoder<'_>, version: Version) -> Result { let request_id = RequestId::decode(buf, version)?; Ok(Self { request_id }) } @@ -411,15 +411,15 @@ impl FetchObject { } impl Encode for FetchObject { - fn encode(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode(&self, w: &mut Encoder<'_>, _: Version) -> Result<(), EncodeError> { match self { Self::EndOfRange { reason, group, object } => { if !Self::END_OF_RANGE.contains(reason) { return Err(EncodeError::InvalidState); } - reason.encode(w, version)?; - group.encode(w, version)?; - object.encode(w, version)?; + w.varint(VarInt::from(*reason))?; + w.varint(VarInt::from(*group))?; + w.varint(VarInt::from(*object))?; } Self::Object { subgroup, @@ -447,22 +447,22 @@ impl Encode for FetchObject { if properties.is_some() { flags |= flag::PROPERTIES; } - flags.encode(w, version)?; + w.varint(VarInt::from(flags))?; if let Some(group) = group { - group.encode(w, version)?; + w.varint(VarInt::from(*group))?; } if let FetchSubgroup::Explicit(subgroup) = subgroup { - subgroup.encode(w, version)?; + w.varint(VarInt::from(*subgroup))?; } if let Some(object) = object { - object.encode(w, version)?; + w.varint(VarInt::from(*object))?; } if let Some(priority) = priority { - priority.encode(w, version)?; + w.u8(*priority); } if let Some(properties) = properties { - properties.encode(w, version)?; + w.bytes(properties)?; } } } @@ -471,8 +471,8 @@ impl Encode for FetchObject { } impl Decode for FetchObject { - fn decode(buf: &mut B, version: Version) -> Result { - let flags = u64::decode(buf, version)?; + fn decode(buf: &mut Decoder<'_>, _: Version) -> Result { + let flags = buf.varint()?.into_inner(); // Anything at or above 128 is a named value rather than a set of flags, and only // the three End of Range markers are defined. @@ -482,14 +482,14 @@ impl Decode for FetchObject { } return Ok(Self::EndOfRange { reason: flags, - group: u64::decode(buf, version)?, - object: u64::decode(buf, version)?, + group: buf.varint()?.into_inner(), + object: buf.varint()?.into_inner(), }); } // Wire order: Group ID Delta, Subgroup ID, Object ID Delta, Priority, Properties. let group = match flags & flag::GROUP_ID != 0 { - true => Some(u64::decode(buf, version)?), + true => Some(buf.varint()?.into_inner()), false => None, }; @@ -499,22 +499,22 @@ impl Decode for FetchObject { 0 => FetchSubgroup::Zero, 1 => FetchSubgroup::Prior, 2 => FetchSubgroup::PriorPlusOne, - _ => FetchSubgroup::Explicit(u64::decode(buf, version)?), + _ => FetchSubgroup::Explicit(buf.varint()?.into_inner()), }, }; let object = match flags & flag::OBJECT_ID != 0 { - true => Some(u64::decode(buf, version)?), + true => Some(buf.varint()?.into_inner()), false => None, }; let priority = match flags & flag::PRIORITY != 0 { - true => Some(u8::decode(buf, version)?), + true => Some(buf.u8()?), false => None, }; let properties = match flags & flag::PROPERTIES != 0 { - true => Some(Vec::::decode(buf, version)?), + true => Some(buf.bytes()?.to_vec()), false => None, }; @@ -531,17 +531,17 @@ impl Decode for FetchObject { #[cfg(test)] mod tests { use super::*; - use bytes::BytesMut; fn encode_message(msg: &M, version: Version) -> Vec { - let mut buf = BytesMut::new(); - msg.encode_msg(&mut buf, version).unwrap(); + let mut buf = Vec::new(); + msg.encode_msg(&mut Encoder::new(&mut buf, version.into()), version) + .unwrap(); buf.to_vec() } fn decode_message(bytes: &[u8], version: Version) -> Result { let mut buf = bytes::Bytes::from(bytes.to_vec()); - M::decode_msg(&mut buf, version) + crate::coding::decode_buf(&mut buf, version, M::decode_msg) } #[test] @@ -762,16 +762,18 @@ mod tests { #[cfg(test)] mod object_tests { use super::*; - use bytes::{Buf as _, BytesMut}; + use bytes::Buf as _; const VERSION: Version = Version::Draft20; fn round_trip(object: &FetchObject) -> (Vec, FetchObject) { - let mut buf = BytesMut::new(); - object.encode(&mut buf, VERSION).expect("encode"); + let mut buf = Vec::new(); + object + .encode(&mut Encoder::new(&mut buf, VERSION.into()), VERSION) + .expect("encode"); let mut bytes = bytes::Bytes::from(buf.to_vec()); - let decoded = FetchObject::decode(&mut bytes, VERSION).expect("decode"); + let decoded = crate::coding::decode_buf(&mut bytes, VERSION, FetchObject::decode).expect("decode"); assert!(!bytes.has_remaining(), "the object header is fully consumed"); (buf.to_vec(), decoded) @@ -863,6 +865,6 @@ mod object_tests { fn an_undefined_value_is_refused() { // 0x8D, one past End of Non-Existent Range, in the draft-17+ leading-ones form. let mut bytes = bytes::Bytes::from_static(&[0x80, 0x8D, 0x00, 0x00]); - assert!(FetchObject::decode(&mut bytes, VERSION).is_err()); + assert!(crate::coding::decode_buf(&mut bytes, VERSION, FetchObject::decode).is_err()); } } diff --git a/rs/moq-net/src/ietf/filter.rs b/rs/moq-net/src/ietf/filter.rs index a909fb4458..5aef4ef0f6 100644 --- a/rs/moq-net/src/ietf/filter.rs +++ b/rs/moq-net/src/ietf/filter.rs @@ -1,8 +1,6 @@ //! The Location Filter carried by SUBSCRIBE, PUBLISH and REQUEST_UPDATE. -use bytes::Buf as _; - -use crate::coding::{Decode, DecodeError, Encode, EncodeError}; +use crate::coding::{Decode, DecodeError, Decoder, Encode, EncodeError, Encoder, VarInt}; use super::{Location, Param, Version}; @@ -82,27 +80,27 @@ impl Filter { } /// Encode the draft-20 field list, without the enclosing length prefix. - fn encode_fields(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_fields(&self, w: &mut Encoder<'_>) -> Result<(), EncodeError> { match *self { // Zero fields. The draft defines an open ended absolute {0, 0} as equivalent, so // that spelling normalizes to this one rather than colliding with NextObject. Self::Unfiltered => {} Self::NextObject => { - 0u64.encode(w, version)?; - 0u64.encode(w, version)?; + w.varint(VarInt::from(0u64))?; + w.varint(VarInt::from(0u64))?; } - Self::Relative(groups) => groups.encode(w, version)?, + Self::Relative(groups) => w.varint(VarInt::from(groups))?, Self::Absolute { start: Location { group: 0, object: 0 }, end: None, } => {} Self::Absolute { start, end } => { - start.group.encode(w, version)?; - start.object.encode(w, version)?; + w.varint(VarInt::from(start.group))?; + w.varint(VarInt::from(start.object))?; if let Some(end) = end { - Self::end_delta(start.group, end.group)?.encode(w, version)?; + w.varint(VarInt::from(Self::end_delta(start.group, end.group)?))?; if let Some(object) = end.object { - object.encode(w, version)?; + w.varint(VarInt::from(object))?; } } } @@ -111,13 +109,13 @@ impl Filter { } /// Decode the draft-20 field list, which the caller has already delimited. - fn decode_fields(r: &mut R, version: Version) -> Result { + fn decode_fields(r: &mut Decoder<'_>) -> Result { let mut fields = Vec::with_capacity(4); - while r.has_remaining() { + while !r.is_empty() { if fields.len() == 4 { return Err(DecodeError::TrailingBytes); } - fields.push(u64::decode(r, version)?); + fields.push(r.varint()?.into_inner()); } Ok(match fields[..] { @@ -148,20 +146,20 @@ impl Filter { } /// Encode the draft-19 and earlier tag form, without the enclosing length prefix. - fn encode_tag(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_tag(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { match *self { // No tag means "everything", which only the absolute spelling can say. Self::Unfiltered => { - tag::ABSOLUTE_START.encode(w, version)?; + w.varint(VarInt::from(tag::ABSOLUTE_START))?; Location::default().encode(w, version)?; } - Self::NextObject => tag::LARGEST_OBJECT.encode(w, version)?, - Self::Relative(0) => tag::NEXT_GROUP.encode(w, version)?, + Self::NextObject => w.varint(VarInt::from(tag::LARGEST_OBJECT))?, + Self::Relative(0) => w.varint(VarInt::from(tag::NEXT_GROUP))?, // Only draft-20 can name a start further back than the next group without // knowing Largest Object, so there is no honest tag to fall back to. Self::Relative(_) => return Err(EncodeError::Unsupported), Self::Absolute { start, end: None } => { - tag::ABSOLUTE_START.encode(w, version)?; + w.varint(VarInt::from(tag::ABSOLUTE_START))?; start.encode(w, version)?; } // Draft-19's AbsoluteRange ends on a group, so an object-bounded range has no @@ -171,17 +169,17 @@ impl Filter { .. } => return Err(EncodeError::Unsupported), Self::Absolute { start, end: Some(end) } => { - tag::ABSOLUTE_RANGE.encode(w, version)?; + w.varint(VarInt::from(tag::ABSOLUTE_RANGE))?; start.encode(w, version)?; - Self::end_delta(start.group, end.group)?.encode(w, version)?; + w.varint(VarInt::from(Self::end_delta(start.group, end.group)?))?; } } Ok(()) } /// Decode the draft-19 and earlier tag form. - fn decode_tag(r: &mut R, version: Version) -> Result { - Ok(match u64::decode(r, version)? { + fn decode_tag(r: &mut Decoder<'_>, version: Version) -> Result { + Ok(match r.varint()?.into_inner() { tag::NEXT_GROUP => Self::Relative(0), tag::LARGEST_OBJECT => Self::NextObject, tag::ABSOLUTE_START => Self::Absolute { @@ -190,7 +188,7 @@ impl Filter { }, tag::ABSOLUTE_RANGE => { let start = Location::decode(r, version)?; - let delta = u64::decode(r, version)?; + let delta = r.varint()?.into_inner(); Self::Absolute { start, end: Some(EndLocation { @@ -207,13 +205,13 @@ impl Filter { /// The inline form, used by draft-14 where the filter is message fields rather than a /// parameter. Later drafts go through [`Param`] instead. impl Encode for Filter { - fn encode(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.encode_tag(w, version) } } impl Decode for Filter { - fn decode(r: &mut R, version: Version) -> Result { + fn decode(r: &mut Decoder<'_>, version: Version) -> Result { Self::decode_tag(r, version) } } @@ -225,7 +223,7 @@ impl Param for Filter { !matches!(self, Self::Unfiltered) } - fn param_encode(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn param_encode(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { let mut buf = Vec::new(); // Inner varints use the draft-15 leading-ones encoding on the drafts that predate @@ -235,30 +233,30 @@ impl Param for Filter { _ => version, }; + let mut inner = Encoder::new(&mut buf, sv.into()); if Self::is_draft20(version) { - self.encode_fields(&mut buf, sv)?; + self.encode_fields(&mut inner)?; } else { - self.encode_tag(&mut buf, sv)?; + self.encode_tag(&mut inner, sv)?; } - buf.encode(w, version) + w.bytes(&buf) } - fn param_decode(r: &mut R, version: Version) -> Result { - let data = Vec::::decode(r, version)?; - let mut buf = bytes::Bytes::from(data); + fn param_decode(r: &mut Decoder<'_>, version: Version) -> Result { let sv = match version { Version::Draft14 | Version::Draft15 | Version::Draft16 => Version::Draft15, _ => version, }; + let mut buf = Decoder::new(r.bytes()?, sv.into()); if Self::is_draft20(version) { // The field count is what carries the meaning, so the value is consumed whole. - return Self::decode_fields(&mut buf, sv); + return Self::decode_fields(&mut buf); } let filter = Self::decode_tag(&mut buf, sv)?; - if buf.has_remaining() { + if !buf.is_empty() { return Err(DecodeError::TrailingBytes); } Ok(filter) @@ -274,19 +272,23 @@ mod tests { fn round_trip(filter: Filter, version: Version) -> Filter { let mut buf = Vec::new(); - filter.param_encode(&mut buf, version).expect("encode"); + filter + .param_encode(&mut Encoder::new(&mut buf, version.into()), version) + .expect("encode"); let mut bytes = bytes::Bytes::from(buf); - let decoded = Filter::param_decode(&mut bytes, version).expect("decode"); - assert!(!bytes.has_remaining(), "parameter left trailing bytes"); + let decoded = crate::coding::decode_buf(&mut bytes, version, Filter::param_decode).expect("decode"); + assert!(bytes.is_empty(), "parameter left trailing bytes"); decoded } /// The value bytes, without the length prefix a parameter carries. fn value(filter: Filter, version: Version) -> Vec { let mut buf = Vec::new(); - filter.param_encode(&mut buf, version).expect("encode"); + filter + .param_encode(&mut Encoder::new(&mut buf, version.into()), version) + .expect("encode"); let mut bytes = bytes::Bytes::from(buf); - Vec::::decode(&mut bytes, version).expect("length prefix") + crate::coding::decode_buf(&mut bytes, version, |r, _| Ok(r.bytes()?.to_vec())).expect("length prefix") } #[test] @@ -435,7 +437,11 @@ mod tests { object: Some(4), }), }; - assert!(bounded.param_encode(&mut buf, OLD).is_err()); + assert!( + bounded + .param_encode(&mut Encoder::new(&mut buf, OLD.into()), OLD) + .is_err() + ); } /// Draft-19 has a tag per case and none of them mean "two groups back", so refuse @@ -443,7 +449,11 @@ mod tests { #[test] fn draft19_cannot_name_a_relative_group() { let mut buf = Vec::new(); - assert!(Filter::Relative(2).param_encode(&mut buf, OLD).is_err()); + assert!( + Filter::Relative(2) + .param_encode(&mut Encoder::new(&mut buf, OLD.into()), OLD) + .is_err() + ); } #[test] @@ -454,16 +464,21 @@ mod tests { }; for version in [OLD, NEW] { let mut buf = Vec::new(); - assert!(backwards.param_encode(&mut buf, version).is_err(), "{version}"); + assert!( + backwards + .param_encode(&mut Encoder::new(&mut buf, version.into()), version) + .is_err(), + "{version}" + ); } } #[test] fn rejects_too_many_fields() { let mut buf = Vec::new(); - vec![0u8; 5].encode(&mut buf, NEW).expect("encode"); + Encoder::new(&mut buf, NEW.into()).bytes(&[0u8; 5]).unwrap(); let mut bytes = bytes::Bytes::from(buf); - assert!(Filter::param_decode(&mut bytes, NEW).is_err()); + assert!(crate::coding::decode_buf(&mut bytes, NEW, Filter::param_decode).is_err()); } } @@ -532,30 +547,30 @@ impl Fill { } impl Param for Fill { - fn param_encode(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn param_encode(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { let mut buf = Vec::new(); + let mut inner = Encoder::new(&mut buf, w.form()); // A nested scope is encoded like a message's parameters: a count, then the KVPs. // An omitted filter inherits the subscription's, so the scope is empty. An explicit // Unfiltered still encodes, as a zero-length filter meaning the whole track. match self.filter { - None => 0u64.encode(&mut buf, version)?, + None => inner.varint(VarInt::from(0u64))?, Some(filter) => { - 1u64.encode(&mut buf, version)?; + inner.varint(VarInt::from(1u64))?; // The first type in a scope is not delta encoded, so this is the raw id. - Self::LOCATION_FILTER.encode(&mut buf, version)?; - filter.param_encode(&mut buf, version)?; + inner.varint(VarInt::from(Self::LOCATION_FILTER))?; + filter.param_encode(&mut inner, version)?; } } - buf.encode(w, version) + w.bytes(&buf) } - fn param_decode(r: &mut R, version: Version) -> Result { - let data = Vec::::decode(r, version)?; - let mut buf = bytes::Bytes::from(data); + fn param_decode(r: &mut Decoder<'_>, version: Version) -> Result { + let mut buf = Decoder::new(r.bytes()?, r.form()); - let count = u64::decode(&mut buf, version)?; + let count = buf.varint()?.into_inner(); if count > 64 { return Err(DecodeError::TooMany); } @@ -564,7 +579,7 @@ impl Param for Fill { let mut range_filters = false; let mut prev = 0u64; for i in 0..count { - let delta = u64::decode(&mut buf, version)?; + let delta = buf.varint()?.into_inner(); let key = if i == 0 { delta } else { @@ -592,18 +607,18 @@ impl Param for Fill { // consumed or the remaining keys desync. match framing { Framing::Bytes => { - Vec::::decode(&mut buf, version)?; + buf.bytes()?; } Framing::Varint => { - u64::decode(&mut buf, version)?; + buf.varint()?; } Framing::Byte => { - u8::decode(&mut buf, version)?; + buf.u8()?; } } } - if buf.has_remaining() { + if !buf.is_empty() { return Err(DecodeError::TrailingBytes); } @@ -619,10 +634,11 @@ mod fill_tests { fn round_trip(fill: Fill) -> Fill { let mut buf = Vec::new(); - fill.param_encode(&mut buf, NEW).expect("encode"); + fill.param_encode(&mut Encoder::new(&mut buf, NEW.into()), NEW) + .expect("encode"); let mut bytes = bytes::Bytes::from(buf); - let decoded = Fill::param_decode(&mut bytes, NEW).expect("decode"); - assert!(!bytes.has_remaining()); + let decoded = crate::coding::decode_buf(&mut bytes, NEW, Fill::param_decode).expect("decode"); + assert!(bytes.is_empty()); decoded } @@ -657,10 +673,11 @@ mod fill_tests { range_filters: false, }; let mut buf = Vec::new(); - fill.param_encode(&mut buf, NEW).expect("encode"); + fill.param_encode(&mut Encoder::new(&mut buf, NEW.into()), NEW) + .expect("encode"); // count=1, type=0x21, len=1, StartGroup=1 let mut bytes = bytes::Bytes::from(buf); - let value = Vec::::decode(&mut bytes, NEW).expect("length prefix"); + let value = crate::coding::decode_buf(&mut bytes, NEW, |r, _| Ok(r.bytes()?.to_vec())).expect("length prefix"); assert_eq!(value, vec![0x01, 0x21, 0x01, 0x01]); } @@ -669,14 +686,16 @@ mod fill_tests { #[test] fn rejects_a_disallowed_parameter() { let mut value = Vec::new(); - 1u64.encode(&mut value, NEW).unwrap(); - 0x10u64.encode(&mut value, NEW).unwrap(); // FORWARD, not allowed in a fill - 0u64.encode(&mut value, NEW).unwrap(); + Encoder::new(&mut value, NEW.into()).varint(VarInt::from(1u64)).unwrap(); + Encoder::new(&mut value, NEW.into()) + .varint(VarInt::from(0x10u64)) + .unwrap(); // FORWARD, not allowed in a fill + Encoder::new(&mut value, NEW.into()).varint(VarInt::from(0u64)).unwrap(); let mut buf = Vec::new(); - value.encode(&mut buf, NEW).unwrap(); + Encoder::new(&mut buf, NEW.into()).bytes(&value).unwrap(); let mut bytes = bytes::Bytes::from(buf); - assert!(Fill::param_decode(&mut bytes, NEW).is_err()); + assert!(crate::coding::decode_buf(&mut bytes, NEW, Fill::param_decode).is_err()); } /// A uint8 parameter is one raw byte, so a priority of 128 or more would be read as a @@ -684,16 +703,20 @@ mod fill_tests { #[test] fn skips_a_uint8_whose_value_has_a_leading_one() { let mut value = Vec::new(); - 2u64.encode(&mut value, NEW).unwrap(); - 0x20u64.encode(&mut value, NEW).unwrap(); // SUBSCRIBER_PRIORITY - 0x80u8.encode(&mut value, NEW).unwrap(); // a raw byte, not a varint - 1u64.encode(&mut value, NEW).unwrap(); // delta to 0x21 - Filter::Relative(1).param_encode(&mut value, NEW).unwrap(); + Encoder::new(&mut value, NEW.into()).varint(VarInt::from(2u64)).unwrap(); + Encoder::new(&mut value, NEW.into()) + .varint(VarInt::from(0x20u64)) + .unwrap(); // SUBSCRIBER_PRIORITY + Encoder::new(&mut value, NEW.into()).u8(0x80u8); // a raw byte, not a varint + Encoder::new(&mut value, NEW.into()).varint(VarInt::from(1u64)).unwrap(); // delta to 0x21 + Filter::Relative(1) + .param_encode(&mut Encoder::new(&mut value, NEW.into()), NEW) + .unwrap(); let mut buf = Vec::new(); - value.encode(&mut buf, NEW).unwrap(); + Encoder::new(&mut buf, NEW.into()).bytes(&value).unwrap(); let mut bytes = bytes::Bytes::from(buf); - let fill = Fill::param_decode(&mut bytes, NEW).expect("decode"); + let fill = crate::coding::decode_buf(&mut bytes, NEW, Fill::param_decode).expect("decode"); assert_eq!(fill.filter, Some(Filter::Relative(1))); } @@ -704,16 +727,20 @@ mod fill_tests { #[test] fn skips_a_length_prefixed_range_filter() { let mut value = Vec::new(); - 2u64.encode(&mut value, NEW).unwrap(); - 0x26u64.encode(&mut value, NEW).unwrap(); // OBJECTID_FILTER, length prefixed - vec![0xAAu8, 0xBB, 0xCC].encode(&mut value, NEW).unwrap(); - 1u64.encode(&mut value, NEW).unwrap(); // delta to 0x27 - vec![0xDDu8].encode(&mut value, NEW).unwrap(); // PRIORITY_FILTER + Encoder::new(&mut value, NEW.into()).varint(VarInt::from(2u64)).unwrap(); + Encoder::new(&mut value, NEW.into()) + .varint(VarInt::from(0x26u64)) + .unwrap(); // OBJECTID_FILTER, length prefixed + Encoder::new(&mut value, NEW.into()) + .bytes(&[0xAAu8, 0xBB, 0xCC]) + .unwrap(); + Encoder::new(&mut value, NEW.into()).varint(VarInt::from(1u64)).unwrap(); // delta to 0x27 + Encoder::new(&mut value, NEW.into()).bytes(&[0xDDu8]).unwrap(); // PRIORITY_FILTER let mut buf = Vec::new(); - value.encode(&mut buf, NEW).unwrap(); + Encoder::new(&mut buf, NEW.into()).bytes(&value).unwrap(); let mut bytes = bytes::Bytes::from(buf); - let fill = Fill::param_decode(&mut bytes, NEW).expect("decode"); + let fill = crate::coding::decode_buf(&mut bytes, NEW, Fill::param_decode).expect("decode"); assert_eq!(fill.filter, None, "no Location Filter in the scope means inherit"); assert!(fill.range_filters, "a Range Filter's presence must be recorded"); } @@ -721,16 +748,22 @@ mod fill_tests { #[test] fn skips_allowed_parameters_it_ignores() { let mut value = Vec::new(); - 2u64.encode(&mut value, NEW).unwrap(); - 0x20u64.encode(&mut value, NEW).unwrap(); // SUBSCRIBER_PRIORITY, even: one varint - 42u64.encode(&mut value, NEW).unwrap(); - 1u64.encode(&mut value, NEW).unwrap(); // delta to 0x21, odd: length prefixed - Filter::Relative(2).param_encode(&mut value, NEW).unwrap(); + Encoder::new(&mut value, NEW.into()).varint(VarInt::from(2u64)).unwrap(); + Encoder::new(&mut value, NEW.into()) + .varint(VarInt::from(0x20u64)) + .unwrap(); // SUBSCRIBER_PRIORITY, even: one varint + Encoder::new(&mut value, NEW.into()) + .varint(VarInt::from(42u64)) + .unwrap(); + Encoder::new(&mut value, NEW.into()).varint(VarInt::from(1u64)).unwrap(); // delta to 0x21, odd: length prefixed + Filter::Relative(2) + .param_encode(&mut Encoder::new(&mut value, NEW.into()), NEW) + .unwrap(); let mut buf = Vec::new(); - value.encode(&mut buf, NEW).unwrap(); + Encoder::new(&mut buf, NEW.into()).bytes(&value).unwrap(); let mut bytes = bytes::Bytes::from(buf); - let fill = Fill::param_decode(&mut bytes, NEW).expect("decode"); + let fill = crate::coding::decode_buf(&mut bytes, NEW, Fill::param_decode).expect("decode"); assert_eq!(fill.filter, Some(Filter::Relative(2))); } } diff --git a/rs/moq-net/src/ietf/goaway.rs b/rs/moq-net/src/ietf/goaway.rs index 5df6c3cdf9..e19d0fe0ae 100644 --- a/rs/moq-net/src/ietf/goaway.rs +++ b/rs/moq-net/src/ietf/goaway.rs @@ -19,11 +19,11 @@ pub struct GoAway<'a> { impl Message for GoAway<'_> { const ID: u64 = 0x10; - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { - self.new_session_uri.encode(w, version)?; + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { + w.string(&self.new_session_uri)?; // Draft-17+ adds a timeout field. if !matches!(version, Version::Draft14 | Version::Draft15 | Version::Draft16) { - self.timeout.encode(w, version)?; + w.varint(VarInt::from(self.timeout))?; } // Draft-18 (#1559) requires a Request ID when GOAWAY is sent on the // control stream, which is the only place we send it. We don't track @@ -32,13 +32,13 @@ impl Message for GoAway<'_> { // conformant peer must treat as a PROTOCOL_VIOLATION. Draft-19 // removed the field again (#1623). if matches!(version, Version::Draft18) { - 0u64.encode(w, version)?; + w.varint(VarInt::from(0u64))?; } Ok(()) } - fn decode_msg(r: &mut R, version: Version) -> Result { - let new_session_uri = Cow::::decode(r, version)?; + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { + let new_session_uri = Cow::::Owned(r.string()?); // All drafts cap the New Session URI at 8,192 bytes; a longer one is a // protocol violation. if new_session_uri.len() > 8192 { @@ -47,17 +47,17 @@ impl Message for GoAway<'_> { let timeout = match version { Version::Draft14 | Version::Draft15 | Version::Draft16 => 0, Version::Draft18 => { - let timeout = u64::decode(r, version)?; + let timeout = r.varint()?.into_inner(); // Draft-18 trailing Request ID (#1559): required on the control // stream, but tolerate its absence from lenient peers. We don't // act on per-request GOAWAY so the value is discarded. Draft-19 // removed this field again (#1623). - if r.has_remaining() { - let _ = u64::decode(r, version)?; + if !r.is_empty() { + let _ = r.varint()?.into_inner(); } timeout } - _ => u64::decode(r, version)?, + _ => r.varint()?.into_inner(), }; Ok(Self { new_session_uri, @@ -69,17 +69,17 @@ impl Message for GoAway<'_> { #[cfg(test)] mod tests { use super::*; - use bytes::BytesMut; fn encode_message(msg: &M) -> Vec { - let mut buf = BytesMut::new(); - msg.encode_msg(&mut buf, Version::Draft14).unwrap(); + let mut buf = Vec::new(); + msg.encode_msg(&mut Encoder::new(&mut buf, Version::Draft14.into()), Version::Draft14) + .unwrap(); buf.to_vec() } fn decode_message(bytes: &[u8]) -> Result { let mut buf = bytes::Bytes::from(bytes.to_vec()); - M::decode_msg(&mut buf, Version::Draft14) + crate::coding::decode_buf(&mut buf, Version::Draft14, M::decode_msg) } #[test] @@ -115,11 +115,12 @@ mod tests { timeout: 5000, }; - let mut buf = BytesMut::new(); - msg.encode_msg(&mut buf, Version::Draft17).unwrap(); + let mut buf = Vec::new(); + msg.encode_msg(&mut Encoder::new(&mut buf, Version::Draft17.into()), Version::Draft17) + .unwrap(); let mut bytes = bytes::Bytes::from(buf.to_vec()); - let decoded: GoAway = GoAway::decode_msg(&mut bytes, Version::Draft17).unwrap(); + let decoded: GoAway = crate::coding::decode_buf(&mut bytes, Version::Draft17, GoAway::decode_msg).unwrap(); assert_eq!(decoded.new_session_uri, "https://example.com/new"); assert_eq!(decoded.timeout, 5000); @@ -132,17 +133,19 @@ mod tests { timeout: 5000, }; - let mut buf = BytesMut::new(); - msg.encode_msg(&mut buf, Version::Draft18).unwrap(); + let mut buf = Vec::new(); + msg.encode_msg(&mut Encoder::new(&mut buf, Version::Draft18.into()), Version::Draft18) + .unwrap(); // Draft-18 requires a trailing Request ID on the control stream, so the // v18 body must be exactly one varint longer than the v17 body. - let mut buf17 = BytesMut::new(); - msg.encode_msg(&mut buf17, Version::Draft17).unwrap(); + let mut buf17 = Vec::new(); + msg.encode_msg(&mut Encoder::new(&mut buf17, Version::Draft17.into()), Version::Draft17) + .unwrap(); assert_eq!(buf.len(), buf17.len() + 1, "v18 must append the Request ID varint"); let mut bytes = bytes::Bytes::from(buf.to_vec()); - let decoded: GoAway = GoAway::decode_msg(&mut bytes, Version::Draft18).unwrap(); + let decoded: GoAway = crate::coding::decode_buf(&mut bytes, Version::Draft18, GoAway::decode_msg).unwrap(); assert_eq!(decoded.new_session_uri, "moqt://relay.example/"); assert_eq!(decoded.timeout, 5000); @@ -157,11 +160,12 @@ mod tests { timeout: 5000, }; - let mut buf = BytesMut::new(); - msg.encode_msg(&mut buf, Version::Draft19).unwrap(); + let mut buf = Vec::new(); + msg.encode_msg(&mut Encoder::new(&mut buf, Version::Draft19.into()), Version::Draft19) + .unwrap(); let mut bytes = bytes::Bytes::from(buf.to_vec()); - let decoded: GoAway = GoAway::decode_msg(&mut bytes, Version::Draft19).unwrap(); + let decoded: GoAway = crate::coding::decode_buf(&mut bytes, Version::Draft19, GoAway::decode_msg).unwrap(); assert_eq!(decoded.new_session_uri, "moqt://relay.example/"); assert_eq!(decoded.timeout, 5000); @@ -171,20 +175,24 @@ mod tests { /// one must not break our decoder; we drain and discard it. #[test] fn test_goaway_v18_drains_optional_request_id() { - use bytes::Buf; - // Hand-construct a draft-18 GOAWAY body that includes the optional Request ID. - let mut buf = BytesMut::new(); - "moqt://relay.example/".encode(&mut buf, Version::Draft18).unwrap(); - 5000u64.encode(&mut buf, Version::Draft18).unwrap(); + let mut buf = Vec::new(); + Encoder::new(&mut buf, Version::Draft18.into()) + .string("moqt://relay.example/") + .unwrap(); + Encoder::new(&mut buf, Version::Draft18.into()) + .varint(VarInt::from(5000u64)) + .unwrap(); // Optional trailing Request ID: - 42u64.encode(&mut buf, Version::Draft18).unwrap(); + Encoder::new(&mut buf, Version::Draft18.into()) + .varint(VarInt::from(42u64)) + .unwrap(); let mut bytes = bytes::Bytes::from(buf.to_vec()); - let decoded: GoAway = GoAway::decode_msg(&mut bytes, Version::Draft18).unwrap(); + let decoded: GoAway = crate::coding::decode_buf(&mut bytes, Version::Draft18, GoAway::decode_msg).unwrap(); assert_eq!(decoded.new_session_uri, "moqt://relay.example/"); assert_eq!(decoded.timeout, 5000); - assert!(!bytes.has_remaining(), "trailing Request ID should be consumed"); + assert!(!!bytes.is_empty(), "trailing Request ID should be consumed"); } } diff --git a/rs/moq-net/src/ietf/group.rs b/rs/moq-net/src/ietf/group.rs index 0196865a1a..c577b77715 100644 --- a/rs/moq-net/src/ietf/group.rs +++ b/rs/moq-net/src/ietf/group.rs @@ -1,4 +1,4 @@ -use crate::coding::{Decode, DecodeError, Encode, EncodeError}; +use crate::coding::{Decode, DecodeError, Decoder, Encode, EncodeError, Encoder, VarInt}; use crate::{Timescale, Timestamp}; use num_enum::{IntoPrimitive, TryFromPrimitive}; @@ -29,29 +29,24 @@ const PROP_TIMESTAMP_DRAFT03: u64 = 0x06; /// timescale on every object would cost bytes per frame to restate something fixed for /// the track's lifetime. `timescale` is what the track advertised, and the timestamp is /// converted into it so the value on the wire matches the declared units. -pub fn encode_object_time( - w: &mut W, +pub fn encode_object_time( + w: &mut Encoder<'_>, timestamp: Timestamp, timescale: Timescale, version: Version, ) -> Result<(), EncodeError> { let timestamp = timestamp.convert(timescale).map_err(|_| EncodeError::BoundsExceeded)?; encode_object_property_type(w, PROP_TIMESTAMP, 0, version)?; - timestamp.value().encode(w, version)?; + w.varint(VarInt::from(timestamp.value()))?; Ok(()) } -fn encode_object_property_type( - w: &mut W, - kind: u64, - prev: u64, - version: Version, -) -> Result<(), EncodeError> { +fn encode_object_property_type(w: &mut Encoder<'_>, kind: u64, prev: u64, version: Version) -> Result<(), EncodeError> { let encoded = match version { Version::Draft14 | Version::Draft15 => kind, _ => kind.checked_sub(prev).ok_or(EncodeError::BoundsExceeded)?, }; - encoded.encode(w, version) + w.varint(VarInt::from(encoded)) } /// Decode the Timestamp (0x10) Object Property from an object's extension block, @@ -60,8 +55,8 @@ fn encode_object_property_type( /// `timescale` is the track's declared units. An object-scope Timescale (0x08) overrides /// it for that object alone, which draft-ietf-moq-loc-04 permits and we still honor on /// decode even though we no longer write one. -pub fn decode_object_time( - r: &mut R, +pub fn decode_object_time( + r: &mut Decoder<'_>, timescale: Timescale, version: Version, ) -> Result, DecodeError> { @@ -70,8 +65,8 @@ pub fn decode_object_time( let mut prev_type: u64 = 0; let mut first = true; - while r.has_remaining() { - let step = u64::decode(r, version)?; + while !r.is_empty() { + let step = r.varint()?.into_inner(); let abs = match version { Version::Draft14 | Version::Draft15 => step, _ if first => step, @@ -82,7 +77,7 @@ pub fn decode_object_time( if abs % 2 == 0 { // Even type: a single varint value. - let value = u64::decode(r, version)?; + let value = r.varint()?.into_inner(); match abs { PROP_TIMESTAMP | PROP_TIMESTAMP_DRAFT03 => timestamp = Some(value), PROP_TIMESCALE => override_scale = Some(value), @@ -90,11 +85,8 @@ pub fn decode_object_time( } } else { // Odd type: length-prefixed bytes we don't care about. - let len = u64::decode(r, version)? as usize; - if r.remaining() < len { - return Err(DecodeError::Short); - } - r.advance(len); + let len = usize::try_from(r.varint()?)?; + r.slice(len)?; } } @@ -129,24 +121,24 @@ impl GroupOrder { } impl Encode for GroupOrder { - fn encode(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { - u8::from(*self).encode(w, version)?; + fn encode(&self, w: &mut Encoder<'_>, _: Version) -> Result<(), EncodeError> { + w.u8(u8::from(*self)); Ok(()) } } impl Decode for GroupOrder { - fn decode(r: &mut R, version: Version) -> Result { - Self::try_from(u8::decode(r, version)?).map_err(|_| DecodeError::InvalidValue) + fn decode(r: &mut Decoder<'_>, _: Version) -> Result { + Self::try_from(r.u8()?).map_err(|_| DecodeError::InvalidValue) } } impl Param for GroupOrder { - fn param_encode(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn param_encode(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { u8::from(*self).param_encode(w, version) } - fn param_decode(r: &mut R, version: Version) -> Result { + fn param_decode(r: &mut Decoder<'_>, version: Version) -> Result { let v = u8::param_decode(r, version)?; Ok(GroupOrder::try_from(v) .unwrap_or(GroupOrder::Descending) @@ -293,42 +285,42 @@ pub struct GroupHeader { } impl Encode for GroupHeader { - fn encode(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { tracing::trace!(?self, "encoding group header"); - self.flags.encode(version)?.encode(w, version)?; - self.track_alias.encode(w, version)?; - self.group_id.encode(w, version)?; + w.varint(VarInt::from(self.flags.encode(version)?))?; + w.varint(VarInt::from(self.track_alias))?; + w.varint(VarInt::from(self.group_id))?; if !self.flags.has_subgroup && self.sub_group_id != 0 { return Err(EncodeError::InvalidState); } if self.flags.has_subgroup { - self.sub_group_id.encode(w, version)?; + w.varint(VarInt::from(self.sub_group_id))?; } // Publisher priority (only if has_priority flag is set) if self.flags.has_priority { - self.publisher_priority.encode(w, version)?; + w.u8(self.publisher_priority); } Ok(()) } } impl Decode for GroupHeader { - fn decode(r: &mut R, version: Version) -> Result { - let flags = GroupFlags::decode(u64::decode(r, version)?, version)?; - let track_alias = u64::decode(r, version)?; - let group_id = u64::decode(r, version)?; + fn decode(r: &mut Decoder<'_>, version: Version) -> Result { + let flags = GroupFlags::decode(r.varint()?.into_inner(), version)?; + let track_alias = r.varint()?.into_inner(); + let group_id = r.varint()?.into_inner(); let sub_group_id = match flags.has_subgroup { - true => u64::decode(r, version)?, + true => r.varint()?.into_inner(), false => 0, }; // Priority present only if has_priority flag is set let publisher_priority = if flags.has_priority { - u8::decode(r, version)? + r.u8()? } else { 128 // Default priority when absent }; @@ -349,22 +341,45 @@ impl Decode for GroupHeader { mod tests { use super::*; - use bytes::Buf; + /// Encode an object's properties at `version`. + fn encode_time(ts: Timestamp, timescale: Timescale, version: Version) -> Vec { + let mut buf = Vec::new(); + encode_object_time(&mut Encoder::new(&mut buf, version.into()), ts, timescale, version).unwrap(); + buf + } + + /// Read `buf` back as a flat list of varints. + fn varints(buf: &[u8], version: Version) -> Vec { + let mut r = Decoder::new(buf, version.into()); + std::iter::from_fn(|| (!r.is_empty()).then(|| r.varint().unwrap().into_inner())).collect() + } + + /// Write `values` as a flat list of varints. + fn from_varints(values: &[u64], version: Version) -> Vec { + let mut buf = Vec::new(); + let mut w = Encoder::new(&mut buf, version.into()); + for value in values { + w.varint((*value).into()).unwrap(); + } + buf + } + + fn decode_time(buf: &[u8], timescale: Timescale, version: Version) -> Option { + let mut r = Decoder::new(buf, version.into()); + let decoded = decode_object_time(&mut r, timescale, version).unwrap(); + assert!(r.is_empty()); + decoded + } /// An object Timestamp round-trips through encode/decode at the track's scale. #[test] fn test_object_time_roundtrip() { let ts = Timestamp::new(96_000, Timescale::MICRO).unwrap(); - let mut buf = bytes::BytesMut::new(); - encode_object_time(&mut buf, ts, Timescale::MICRO, Version::Draft18).unwrap(); + let buf = encode_time(ts, Timescale::MICRO, Version::Draft18); - let mut bytes = buf.freeze(); - let decoded = decode_object_time(&mut bytes, Timescale::MICRO, Version::Draft18) - .unwrap() - .unwrap(); + let decoded = decode_time(&buf, Timescale::MICRO, Version::Draft18).unwrap(); assert_eq!(decoded.value(), 96_000); assert_eq!(decoded.scale(), Timescale::MICRO); - assert!(!bytes.has_remaining()); } /// The value on the wire is in the track's units, not the frame's. @@ -372,18 +387,10 @@ mod tests { fn test_object_time_converts_into_the_track_scale() { // 2 seconds, expressed in milliseconds by the frame. let ts = Timestamp::new(2_000, Timescale::MILLI).unwrap(); - let mut buf = bytes::BytesMut::new(); - encode_object_time(&mut buf, ts, Timescale::MICRO, Version::Draft18).unwrap(); + let buf = encode_time(ts, Timescale::MICRO, Version::Draft18); + assert_eq!(varints(&buf, Version::Draft18), [PROP_TIMESTAMP, 2_000_000]); - let mut bytes = buf.clone().freeze(); - assert_eq!(u64::decode(&mut bytes, Version::Draft18).unwrap(), PROP_TIMESTAMP); - assert_eq!(u64::decode(&mut bytes, Version::Draft18).unwrap(), 2_000_000); - assert!(!bytes.has_remaining()); - - let mut bytes = buf.freeze(); - let decoded = decode_object_time(&mut bytes, Timescale::MICRO, Version::Draft18) - .unwrap() - .unwrap(); + let decoded = decode_time(&buf, Timescale::MICRO, Version::Draft18).unwrap(); assert_eq!(decoded.value(), 2_000_000); assert_eq!(decoded.scale(), Timescale::MICRO); } @@ -392,43 +399,32 @@ mod tests { #[test] fn test_object_time_omits_the_timescale() { let ts = Timestamp::new(96_000, Timescale::MILLI).unwrap(); - let mut buf = bytes::BytesMut::new(); - encode_object_time(&mut buf, ts, Timescale::MILLI, Version::Draft16).unwrap(); - - let mut bytes = buf.freeze(); - assert_eq!(u64::decode(&mut bytes, Version::Draft16).unwrap(), PROP_TIMESTAMP); - assert_eq!(u64::decode(&mut bytes, Version::Draft16).unwrap(), 96_000); - assert!(!bytes.has_remaining()); + let buf = encode_time(ts, Timescale::MILLI, Version::Draft16); + assert_eq!(varints(&buf, Version::Draft16), [PROP_TIMESTAMP, 96_000]); } /// Draft-14/15 write absolute property types rather than deltas. #[test] fn test_object_time_legacy_uses_absolute_types() { let ts = Timestamp::new(96_000, Timescale::MILLI).unwrap(); - let mut buf = bytes::BytesMut::new(); - encode_object_time(&mut buf, ts, Timescale::MILLI, Version::Draft15).unwrap(); - - let mut bytes = buf.freeze(); - assert_eq!(u64::decode(&mut bytes, Version::Draft15).unwrap(), PROP_TIMESTAMP); - assert_eq!(u64::decode(&mut bytes, Version::Draft15).unwrap(), 96_000); - assert!(!bytes.has_remaining()); + let buf = encode_time(ts, Timescale::MILLI, Version::Draft15); + assert_eq!(varints(&buf, Version::Draft15), [PROP_TIMESTAMP, 96_000]); } /// An object-scope Timescale (which LOC permits) overrides the track's for that object. #[test] fn test_object_time_honors_an_object_scope_timescale() { - let mut buf = bytes::BytesMut::new(); - PROP_TIMESCALE.encode(&mut buf, Version::Draft18).unwrap(); - u64::from(Timescale::MILLI).encode(&mut buf, Version::Draft18).unwrap(); - (PROP_TIMESTAMP - PROP_TIMESCALE) - .encode(&mut buf, Version::Draft18) - .unwrap(); - 42u64.encode(&mut buf, Version::Draft18).unwrap(); + let buf = from_varints( + &[ + PROP_TIMESCALE, + u64::from(Timescale::MILLI), + PROP_TIMESTAMP - PROP_TIMESCALE, + 42, + ], + Version::Draft18, + ); - let mut bytes = buf.freeze(); - let decoded = decode_object_time(&mut bytes, Timescale::MICRO, Version::Draft18) - .unwrap() - .unwrap(); + let decoded = decode_time(&buf, Timescale::MICRO, Version::Draft18).unwrap(); assert_eq!(decoded.value(), 42); assert_eq!(decoded.scale(), Timescale::MILLI); } @@ -436,14 +432,9 @@ mod tests { /// Without an object-scope override, the track's timescale supplies the units. #[test] fn test_object_time_defaults_to_the_track_scale() { - let mut buf = bytes::BytesMut::new(); - PROP_TIMESTAMP.encode(&mut buf, Version::Draft18).unwrap(); - 1234u64.encode(&mut buf, Version::Draft18).unwrap(); + let buf = from_varints(&[PROP_TIMESTAMP, 1234], Version::Draft18); - let mut bytes = buf.freeze(); - let decoded = decode_object_time(&mut bytes, Timescale::MILLI, Version::Draft18) - .unwrap() - .unwrap(); + let decoded = decode_time(&buf, Timescale::MILLI, Version::Draft18).unwrap(); assert_eq!(decoded.value(), 1234); assert_eq!(decoded.scale(), Timescale::MILLI); } @@ -451,14 +442,9 @@ mod tests { /// A peer on draft-ietf-moq-loc-03 wrote the Timestamp at 0x06; still decode it. #[test] fn test_object_time_decodes_draft03_timestamp() { - let mut buf = bytes::BytesMut::new(); - PROP_TIMESTAMP_DRAFT03.encode(&mut buf, Version::Draft18).unwrap(); - 777u64.encode(&mut buf, Version::Draft18).unwrap(); + let buf = from_varints(&[PROP_TIMESTAMP_DRAFT03, 777], Version::Draft18); - let mut bytes = buf.freeze(); - let decoded = decode_object_time(&mut bytes, Timescale::MICRO, Version::Draft18) - .unwrap() - .unwrap(); + let decoded = decode_time(&buf, Timescale::MICRO, Version::Draft18).unwrap(); assert_eq!(decoded.value(), 777); assert_eq!(decoded.scale(), Timescale::MICRO); } @@ -466,12 +452,7 @@ mod tests { /// No Timestamp property at all yields None (the caller wall-clock-stamps). #[test] fn test_object_time_absent() { - let mut empty = bytes::Bytes::new(); - assert!( - decode_object_time(&mut empty, Timescale::MICRO, Version::Draft18) - .unwrap() - .is_none() - ); + assert!(decode_time(&[], Timescale::MICRO, Version::Draft18).is_none()); } // Test table from draft-ietf-moq-transport-14 Section 10.4.2 Table 7 @@ -715,8 +696,10 @@ mod tests { flags: GroupFlags::default(), }; - let mut buf = bytes::BytesMut::new(); - header.encode(&mut buf, Version::Draft18).unwrap(); + let mut buf = Vec::new(); + header + .encode(&mut Encoder::new(&mut buf, Version::Draft18.into()), Version::Draft18) + .unwrap(); let type_byte = buf[0] as u64; // The check in session.rs::run_uni_group. diff --git a/rs/moq-net/src/ietf/location.rs b/rs/moq-net/src/ietf/location.rs index 71bb2bd64f..5a3c4f174a 100644 --- a/rs/moq-net/src/ietf/location.rs +++ b/rs/moq-net/src/ietf/location.rs @@ -1,4 +1,4 @@ -use crate::coding::{Decode, DecodeError, Encode, EncodeError}; +use crate::coding::{Decode, DecodeError, Decoder, Encode, EncodeError, Encoder, VarInt}; use super::Version; @@ -9,17 +9,17 @@ pub struct Location { } impl Encode for Location { - fn encode(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { - self.group.encode(w, version)?; - self.object.encode(w, version)?; + fn encode(&self, w: &mut Encoder<'_>, _: Version) -> Result<(), EncodeError> { + w.varint(VarInt::from(self.group))?; + w.varint(VarInt::from(self.object))?; Ok(()) } } impl Decode for Location { - fn decode(buf: &mut B, version: Version) -> Result { - let group = u64::decode(buf, version)?; - let object = u64::decode(buf, version)?; + fn decode(buf: &mut Decoder<'_>, _: Version) -> Result { + let group = buf.varint()?.into_inner(); + let object = buf.varint()?.into_inner(); Ok(Self { group, object }) } } diff --git a/rs/moq-net/src/ietf/message.rs b/rs/moq-net/src/ietf/message.rs index 4bf95248df..a2c4439288 100644 --- a/rs/moq-net/src/ietf/message.rs +++ b/rs/moq-net/src/ietf/message.rs @@ -1,6 +1,4 @@ -use bytes::{Buf, BufMut}; - -use crate::coding::{Decode, DecodeError, Encode, EncodeError, Sizer}; +use crate::coding::{Decode, DecodeError, Decoder, Encode, EncodeError, Encoder}; use super::Version; @@ -11,63 +9,55 @@ pub trait Message: Sized + std::fmt::Debug { const ID: u64; /// Encode this message body (without size prefix). - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError>; + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError>; /// Decode a message body (without size prefix). - fn decode_msg(buf: &mut B, version: Version) -> Result; + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result; } impl Encode for T { - fn encode(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { tracing::trace!(?self, "encoding"); - let mut sizer = Sizer::default(); - self.encode_msg(&mut sizer, version)?; - let size: u16 = sizer.size.try_into().map_err(|_| EncodeError::TooLarge)?; - size.encode(w, version)?; - self.encode_msg(w, version) + let start = w.position(); + self.encode_msg(w, version)?; + w.prefix_u16(start) } } -impl Decode for T { - fn decode(buf: &mut B, version: Version) -> Result { - let size = u16::decode(buf, version)? as usize; +/// A control message body not decoded yet: a `u16` length, then that many bytes. +/// +/// Read after the message type, when the type decides how to decode the rest. +#[derive(Debug)] +pub struct Body(pub bytes::Bytes); + +impl Decode for Body { + fn decode(r: &mut Decoder<'_>, _: Version) -> Result { + let size = r.u16()? as usize; + Ok(Self(bytes::Bytes::copy_from_slice(r.slice(size)?))) + } +} - if tracing::enabled!(tracing::Level::TRACE) { - if buf.remaining() < size { - return Err(DecodeError::Short); - } - let raw = buf.copy_to_bytes(size); - let mut slice = &raw[..]; - match Self::decode_msg(&mut slice, version) { - Ok(result) => { - if slice.remaining() > 0 { - return Err(DecodeError::Long); - } - tracing::trace!(?result, "decoded"); - Ok(result) - } - Err(e) => { - tracing::warn!(%e, ?raw, "decode failed"); - Err(e) - } - } - } else { - if buf.remaining() < size { - return Err(DecodeError::Short); - } - let mut limited = buf.take(size); - match Self::decode_msg(&mut limited, version) { - Ok(result) => { - if limited.remaining() > 0 { - return Err(DecodeError::Long); - } - Ok(result) - } - Err(e) => { - tracing::warn!(%e, "decode failed"); - Err(e) - } - } +impl Body { + /// A decoder over the body, with `version`'s varints. + pub fn decoder(&self, version: Version) -> Decoder<'_> { + Decoder::new(&self.0, version.into()) + } +} + +impl Decode for T { + fn decode(r: &mut Decoder<'_>, version: Version) -> Result { + let size = r.u16()? as usize; + let mut body = r.sub(size)?; + + let result = Self::decode_msg(&mut body, version).and_then(|msg| match body.is_empty() { + true => Ok(msg), + false => Err(DecodeError::Long), + }); + + match &result { + Ok(msg) => tracing::trace!(?msg, "decoded"), + Err(err) => tracing::warn!(%err, "decode failed"), } + result } } diff --git a/rs/moq-net/src/ietf/mod.rs b/rs/moq-net/src/ietf/mod.rs index f97355e820..9e890d70a3 100644 --- a/rs/moq-net/src/ietf/mod.rs +++ b/rs/moq-net/src/ietf/mod.rs @@ -41,7 +41,7 @@ pub use filter::*; pub use goaway::*; pub use group::*; pub use location::*; -pub use message::Message; +pub use message::{Body, Message}; pub use parameters::*; pub use properties::Properties; pub use publish::*; diff --git a/rs/moq-net/src/ietf/namespace.rs b/rs/moq-net/src/ietf/namespace.rs index 70df7144b0..e82eaad8a0 100644 --- a/rs/moq-net/src/ietf/namespace.rs +++ b/rs/moq-net/src/ietf/namespace.rs @@ -1,7 +1,5 @@ use crate::{Path, coding::*}; -use super::Version; - fn to_tuple(namespace: &Path) -> Vec { let path = namespace.as_str(); if path.is_empty() { @@ -49,7 +47,7 @@ fn from_tuple(parts: &[String]) -> Path<'static> { } /// Helper function to encode namespace as tuple of strings -pub fn encode_namespace(w: &mut W, namespace: &Path, version: Version) -> Result<(), EncodeError> { +pub fn encode_namespace(w: &mut Encoder<'_>, namespace: &Path) -> Result<(), EncodeError> { let parts = to_tuple(namespace); // The IETF draft limits namespaces to 32 parts. @@ -57,16 +55,16 @@ pub fn encode_namespace(w: &mut W, namespace: &Path, version: return Err(BoundsExceeded.into()); } - (parts.len() as u64).encode(w, version)?; + w.varint(VarInt::from(parts.len()))?; for part in parts { - part.encode(w, version)?; + w.string(&part)?; } Ok(()) } /// Helper function to decode namespace from tuple of strings -pub fn decode_namespace(r: &mut R, version: Version) -> Result, DecodeError> { - let count = u64::decode(r, version)?; +pub fn decode_namespace(r: &mut Decoder<'_>) -> Result, DecodeError> { + let count = r.varint()?.into_inner(); if count == 0 { return Ok(Path::from(String::new())); @@ -80,7 +78,7 @@ pub fn decode_namespace(r: &mut R, version: Version) -> Result(r: &mut R, version: Version) -> Result Vec { encode_path(&Path::from(path.to_string())) } fn encode_path(path: &Path<'_>) -> Vec { - let mut buf = BytesMut::new(); - encode_namespace(&mut buf, path, Version::Draft17).unwrap(); - buf.to_vec() + let mut buf = Vec::new(); + encode_namespace(&mut Encoder::new(&mut buf, FORM), path).unwrap(); + buf } fn decode_ns(bytes: &[u8]) -> Path<'static> { - let mut buf = bytes::Bytes::from(bytes.to_vec()); - decode_namespace(&mut buf, Version::Draft17).unwrap() + decode_namespace(&mut Decoder::new(bytes, FORM)).unwrap() } fn encode_tuple(parts: &[&str]) -> Vec { - let mut buf = BytesMut::new(); - (parts.len() as u64).encode(&mut buf, Version::Draft17).unwrap(); + let mut buf = Vec::new(); + let mut w = Encoder::new(&mut buf, FORM); + w.varint(parts.len().into()).unwrap(); for part in parts { - (*part).encode(&mut buf, Version::Draft17).unwrap(); + w.string(part).unwrap(); } - buf.to_vec() + buf } #[test] diff --git a/rs/moq-net/src/ietf/parameters.rs b/rs/moq-net/src/ietf/parameters.rs index 5fb8ad044a..9a28fd417c 100644 --- a/rs/moq-net/src/ietf/parameters.rs +++ b/rs/moq-net/src/ietf/parameters.rs @@ -1,6 +1,3 @@ -use std::collections::{HashMap, hash_map}; - -use bytes::Buf; use num_enum::{FromPrimitive, IntoPrimitive}; use crate::coding::*; @@ -45,164 +42,117 @@ pub enum ParameterBytes { Unknown(u64), } +/// SETUP parameters, in the order they were set or decoded. +/// +/// A handful at most, so a linear scan beats hashing, and the encoding is deterministic. #[derive(Default, Debug, Clone)] pub struct Parameters { - vars: HashMap, - bytes: HashMap>, + vars: Vec<(ParameterVarInt, u64)>, + bytes: Vec<(ParameterBytes, Vec)>, } impl Decode for Parameters { - fn decode(mut r: &mut R, version: Version) -> Result { - let mut vars = HashMap::new(); - let mut bytes = HashMap::new(); + fn decode(r: &mut Decoder<'_>, version: Version) -> Result { + let mut params = Self::default(); - match version { - Version::Draft14 | Version::Draft15 | Version::Draft16 => { - let count = u64::decode(r, version)?; + // Draft-14/15/16 count the pairs; draft-17+ reads them until the buffer is empty. + let count = match version { + Version::Draft14 | Version::Draft15 | Version::Draft16 => Some(r.varint()?.into_inner()), + _ => None, + }; + if count.is_some_and(|count| count > MAX_PARAMS) { + return Err(DecodeError::TooMany); + } - if count > MAX_PARAMS { - return Err(DecodeError::TooMany); - } + // Draft-16+ delta-encodes the types; even is a varint value, odd is length-prefixed bytes. + let delta = !matches!(version, Version::Draft14 | Version::Draft15); + let mut prev = 0u64; + let mut i = 0u64; - let mut prev_type: u64 = 0; - - for i in 0..count { - let kind = match version { - Version::Draft16 => { - let delta = u64::decode(r, version)?; - let abs = if i == 0 { - delta - } else { - prev_type.checked_add(delta).ok_or(DecodeError::BoundsExceeded)? - }; - prev_type = abs; - abs - } - Version::Draft14 | Version::Draft15 => u64::decode(r, version)?, - _ => unreachable!("handled above"), - }; - - if kind % 2 == 0 { - let kind = ParameterVarInt::from(kind); - match vars.entry(kind) { - hash_map::Entry::Occupied(_) => return Err(DecodeError::Duplicate), - hash_map::Entry::Vacant(entry) => entry.insert(u64::decode(&mut r, version)?), - }; - } else { - let kind = ParameterBytes::from(kind); - let val = Vec::::decode(&mut r, version)?; - if val.len() > MAX_KVP_VALUE_LEN { - return Err(DecodeError::BoundsExceeded); - } - match bytes.entry(kind) { - hash_map::Entry::Occupied(_) => return Err(DecodeError::Duplicate), - hash_map::Entry::Vacant(entry) => entry.insert(val), - }; - } - } + while count.map_or(!r.is_empty(), |count| i < count) { + if i >= MAX_PARAMS { + return Err(DecodeError::TooMany); } - _ => { - // Draft17+: no count prefix, read Key-Value-Pairs until buffer empty. - // Delta-encoded types, even = varint value, odd = length-prefixed bytes. - let mut prev_type: u64 = 0; - let mut i = 0u64; - while r.has_remaining() { - if i >= MAX_PARAMS { - return Err(DecodeError::TooMany); - } - let delta = u64::decode(&mut r, version)?; - let abs = if i == 0 { - delta - } else { - prev_type.checked_add(delta).ok_or(DecodeError::BoundsExceeded)? - }; - prev_type = abs; - i += 1; - - if abs % 2 == 0 { - let kind = ParameterVarInt::from(abs); - match vars.entry(kind) { - hash_map::Entry::Occupied(_) => return Err(DecodeError::Duplicate), - hash_map::Entry::Vacant(entry) => entry.insert(u64::decode(&mut r, version)?), - }; - } else { - let kind = ParameterBytes::from(abs); - let val = Vec::::decode(&mut r, version)?; - if val.len() > MAX_KVP_VALUE_LEN { - return Err(DecodeError::BoundsExceeded); - } - match bytes.entry(kind) { - hash_map::Entry::Occupied(_) => return Err(DecodeError::Duplicate), - hash_map::Entry::Vacant(entry) => entry.insert(val), - }; - } + + let kind = r.varint()?.into_inner(); + let kind = match delta && i > 0 { + true => prev.checked_add(kind).ok_or(DecodeError::BoundsExceeded)?, + false => kind, + }; + prev = kind; + i += 1; + + if kind % 2 == 0 { + let kind = ParameterVarInt::from(kind); + if params.get_varint(kind).is_some() { + return Err(DecodeError::Duplicate); } + params.vars.push((kind, r.varint()?.into_inner())); + } else { + let kind = ParameterBytes::from(kind); + let value = r.bytes()?; + if value.len() > MAX_KVP_VALUE_LEN { + return Err(DecodeError::BoundsExceeded); + } + if params.get_bytes(kind).is_some() { + return Err(DecodeError::Duplicate); + } + params.bytes.push((kind, value.to_vec())); } } - Ok(Parameters { vars, bytes }) + Ok(params) } } impl Encode for Parameters { - fn encode(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { let count = self.vars.len() + self.bytes.len(); if count as u64 > MAX_PARAMS { return Err(EncodeError::TooMany); } + if self.bytes.iter().any(|(_, value)| value.len() > MAX_KVP_VALUE_LEN) { + return Err(EncodeError::BoundsExceeded); + } match version { Version::Draft14 | Version::Draft15 => { - count.encode(w, version)?; + w.varint(count.into())?; - for (kind, value) in self.vars.iter() { - u64::from(*kind).encode(w, version)?; - value.encode(w, version)?; + for (kind, value) in &self.vars { + w.varint(u64::from(*kind).into())?; + w.varint((*value).into())?; } - for (kind, value) in self.bytes.iter() { - if value.len() > MAX_KVP_VALUE_LEN { - return Err(EncodeError::BoundsExceeded); - } - u64::from(*kind).encode(w, version)?; - value.encode(w, version)?; + for (kind, value) in &self.bytes { + w.varint(u64::from(*kind).into())?; + w.bytes(value)?; } } _ => { // Draft16: count prefix + delta encoding // Draft17+: NO count prefix + delta encoding if matches!(version, Version::Draft16) { - count.encode(w, version)?; + w.varint(count.into())?; } - // Collect all keys, sort, encode deltas enum ParamRef<'a> { - Var(&'a u64), - Bytes(&'a Vec), - } - let mut all: Vec<(u64, ParamRef)> = Vec::new(); - for (k, v) in self.vars.iter() { - all.push((u64::from(*k), ParamRef::Var(v))); - } - for (k, v) in self.bytes.iter() { - all.push((u64::from(*k), ParamRef::Bytes(v))); + Var(u64), + Bytes(&'a [u8]), } + let mut all: Vec<(u64, ParamRef)> = Vec::with_capacity(count); + all.extend(self.vars.iter().map(|(k, v)| (u64::from(*k), ParamRef::Var(*v)))); + all.extend(self.bytes.iter().map(|(k, v)| (u64::from(*k), ParamRef::Bytes(v)))); all.sort_by_key(|(k, _)| *k); - let mut prev_type: u64 = 0; - for (idx, (kind, val)) in all.iter().enumerate() { - let delta = if idx == 0 { *kind } else { kind - prev_type }; - prev_type = *kind; - delta.encode(w, version)?; - - match val { - ParamRef::Var(v) => v.encode(w, version)?, - ParamRef::Bytes(v) => { - if v.len() > MAX_KVP_VALUE_LEN { - return Err(EncodeError::BoundsExceeded); - } - v.encode(w, version)?; - } + let mut prev = 0u64; + for (kind, value) in all { + w.varint((kind - prev).into())?; + prev = kind; + + match value { + ParamRef::Var(v) => w.varint(v.into())?, + ParamRef::Bytes(v) => w.bytes(v)?, } } } @@ -214,19 +164,25 @@ impl Encode for Parameters { impl Parameters { pub fn get_varint(&self, kind: ParameterVarInt) -> Option { - self.vars.get(&kind).copied() + self.vars.iter().find(|(k, _)| *k == kind).map(|(_, v)| *v) } pub fn set_varint(&mut self, kind: ParameterVarInt, value: u64) { - self.vars.insert(kind, value); + match self.vars.iter_mut().find(|(k, _)| *k == kind) { + Some((_, v)) => *v = value, + None => self.vars.push((kind, value)), + } } pub fn get_bytes(&self, kind: ParameterBytes) -> Option<&[u8]> { - self.bytes.get(&kind).map(|v| v.as_slice()) + self.bytes.iter().find(|(k, _)| *k == kind).map(|(_, v)| v.as_slice()) } pub fn set_bytes(&mut self, kind: ParameterBytes, value: Vec) { - self.bytes.insert(kind, value); + match self.bytes.iter_mut().find(|(k, _)| *k == kind) { + Some((_, v)) => *v = value, + None => self.bytes.push((kind, value)), + } } } @@ -240,8 +196,8 @@ impl Parameters { /// /// Use `_ =>` for the newest draft behavior so future versions default forward. pub trait Param: Sized { - fn param_encode(&self, w: &mut W, version: Version) -> Result<(), EncodeError>; - fn param_decode(r: &mut R, version: Version) -> Result; + fn param_encode(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError>; + fn param_decode(r: &mut Decoder<'_>, version: Version) -> Result; /// Whether this parameter should be encoded. Returns false to skip. fn param_present(&self) -> bool { @@ -250,56 +206,55 @@ pub trait Param: Sized { } impl Param for u8 { - fn param_encode(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn param_encode(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { match version { - // Draft-14/15/16: u8 encoded as varint (cast to u64) - Version::Draft14 | Version::Draft15 | Version::Draft16 => (*self as u64).encode(w, version), - _ => Encode::encode(self, w, version), + // Draft-14/15/16: u8 encoded as varint + Version::Draft14 | Version::Draft15 | Version::Draft16 => w.varint((*self).into())?, + _ => w.u8(*self), } + Ok(()) } - fn param_decode(r: &mut R, version: Version) -> Result { + fn param_decode(r: &mut Decoder<'_>, version: Version) -> Result { match version { Version::Draft14 | Version::Draft15 | Version::Draft16 => { - let v = u64::decode(r, version)?; - u8::try_from(v).map_err(|_| DecodeError::InvalidValue) + u8::try_from(r.varint()?).map_err(|_| DecodeError::InvalidValue) } - _ => u8::decode(r, version), + _ => r.u8(), } } } impl Param for bool { - fn param_encode(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn param_encode(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { match version { - // Draft-14/15/16: bool encoded as varint (cast to u64) - Version::Draft14 | Version::Draft15 | Version::Draft16 => (*self as u64).encode(w, version), - _ => Encode::encode(self, w, version), + // Draft-14/15/16: bool encoded as varint + Version::Draft14 | Version::Draft15 | Version::Draft16 => w.varint(u8::from(*self).into())?, + _ => w.bool(*self), } + Ok(()) } - fn param_decode(r: &mut R, version: Version) -> Result { + fn param_decode(r: &mut Decoder<'_>, version: Version) -> Result { match version { - Version::Draft14 | Version::Draft15 | Version::Draft16 => { - let v = u64::decode(r, version)?; - match v { - 0 => Ok(false), - 1 => Ok(true), - _ => Err(DecodeError::InvalidValue), - } - } - _ => bool::decode(r, version), + Version::Draft14 | Version::Draft15 | Version::Draft16 => match r.varint()?.into_inner() { + 0 => Ok(false), + 1 => Ok(true), + _ => Err(DecodeError::InvalidValue), + }, + _ => r.bool(), } } } impl Param for u64 { - fn param_encode(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { - self.encode(w, version) + fn param_encode(&self, w: &mut Encoder<'_>, _: Version) -> Result<(), EncodeError> { + w.varint((*self).into())?; + Ok(()) } - fn param_decode(r: &mut R, version: Version) -> Result { - u64::decode(r, version) + fn param_decode(r: &mut Decoder<'_>, _: Version) -> Result { + Ok(r.varint()?.into_inner()) } } @@ -312,38 +267,39 @@ impl Param for u64 { /// all, and lists "Location: Two consecutive varints (Group, Object)" as its own value /// encoding beside Length-prefixed. So from draft-17 the two varints are written bare. impl Param for Location { - fn param_encode(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn param_encode(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { match version { Version::Draft14 | Version::Draft15 | Version::Draft16 => { // The drafts before 17 pin the inner varints to the draft-15 encoding, // matching the other length-prefixed parameters. let mut buf = Vec::new(); - self.group.encode(&mut buf, Version::Draft15)?; - self.object.encode(&mut buf, Version::Draft15)?; - buf.encode(w, version) + let mut inner = Encoder::new(&mut buf, Version::Draft15.into()); + inner.varint(self.group.into())?; + inner.varint(self.object.into())?; + w.bytes(&buf) } _ => { - self.group.encode(w, version)?; - self.object.encode(w, version) + w.varint(self.group.into())?; + w.varint(self.object.into())?; + Ok(()) } } } - fn param_decode(r: &mut R, version: Version) -> Result { + fn param_decode(r: &mut Decoder<'_>, version: Version) -> Result { match version { Version::Draft14 | Version::Draft15 | Version::Draft16 => { - let data = Vec::::decode(r, version)?; - let mut buf = bytes::Bytes::from(data); - let group = u64::decode(&mut buf, Version::Draft15)?; - let object = u64::decode(&mut buf, Version::Draft15)?; - if buf.has_remaining() { + let mut inner = Decoder::new(r.bytes()?, Version::Draft15.into()); + let group = inner.varint()?.into_inner(); + let object = inner.varint()?.into_inner(); + if !inner.is_empty() { return Err(DecodeError::TrailingBytes); } Ok(Location { group, object }) } _ => { - let group = u64::decode(r, version)?; - let object = u64::decode(r, version)?; + let group = r.varint()?.into_inner(); + let object = r.varint()?.into_inner(); Ok(Location { group, object }) } } @@ -355,14 +311,14 @@ impl Param for Option { self.is_some() } - fn param_encode(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn param_encode(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { match self { Some(v) => v.param_encode(w, version), None => Ok(()), } } - fn param_decode(r: &mut R, version: Version) -> Result { + fn param_decode(r: &mut Decoder<'_>, version: Version) -> Result { Ok(Some(T::param_decode(r, version)?)) } } @@ -380,9 +336,6 @@ impl Param for Option { /// ``` macro_rules! encode_params { ($w:expr, $version:expr, $($key:expr => $val:expr),* $(,)?) => {{ - #[allow(unused_imports)] - use $crate::coding::Encode as _; - #[allow(unused)] const _: () = { let _keys: &[u64] = &[$($key),*]; @@ -398,7 +351,7 @@ macro_rules! encode_params { #[allow(unused_mut)] let mut _count: usize = 0; $(_count += if $crate::ietf::Param::param_present(&$val) { 1 } else { 0 };)* - _count.encode($w, _version)?; + $w.varint($crate::coding::VarInt::from(_count))?; #[allow(unused_mut, unused_assignments)] let mut _prev_key: u64 = 0; @@ -407,15 +360,12 @@ macro_rules! encode_params { $( if $crate::ietf::Param::param_present(&$val) { let _key: u64 = $key; - match _version { - $crate::ietf::Version::Draft14 | $crate::ietf::Version::Draft15 => { - _key.encode($w, _version)?; - } - _ => { - let _delta = if _first { _key } else { _key - _prev_key }; - _delta.encode($w, _version)?; - } - } + let _wire = match _version { + $crate::ietf::Version::Draft14 | $crate::ietf::Version::Draft15 => _key, + _ if _first => _key, + _ => _key - _prev_key, + }; + $w.varint($crate::coding::VarInt::from(_wire))?; _prev_key = _key; _first = false; $crate::ietf::Param::param_encode(&$val, $w, _version)?; @@ -457,11 +407,8 @@ macro_rules! decode_params { $(#[allow(unused_mut, non_snake_case)] let mut $name: Option<$ty> = None;)* { - #[allow(unused_imports)] - use $crate::coding::Decode as _; - let _version: $crate::ietf::Version = $version; - let _count = >::decode($r, _version)?; + let _count = $r.varint()?.into_inner(); if _count > 64 { return Err($crate::coding::DecodeError::TooMany); } @@ -469,21 +416,13 @@ macro_rules! decode_params { #[allow(unused_mut, unused_assignments)] let mut _prev_key: u64 = 0; for _i in 0.._count { + let _wire = $r.varint()?.into_inner(); let _key: u64 = match _version { - $crate::ietf::Version::Draft14 | $crate::ietf::Version::Draft15 => { - >::decode($r, _version)? - } - _ => { - let _delta = >::decode($r, _version)?; - let _abs = if _i == 0 { - _delta - } else { - _prev_key.checked_add(_delta).ok_or($crate::coding::DecodeError::BoundsExceeded)? - }; - _prev_key = _abs; - _abs - } + $crate::ietf::Version::Draft14 | $crate::ietf::Version::Draft15 => _wire, + _ if _i == 0 => _wire, + _ => _prev_key.checked_add(_wire).ok_or($crate::coding::DecodeError::BoundsExceeded)?, }; + _prev_key = _key; // An if-chain rather than a `match`, so a key can be a named constant: // the macro captures it as an expression, which is not a legal pattern. @@ -509,7 +448,6 @@ macro_rules! decode_params { mod tests { use super::super::Filter; use super::*; - use bytes::{Buf, Bytes, BytesMut}; // ---- Setup Parameters tests (unchanged) ---- @@ -520,11 +458,13 @@ mod tests { params.set_varint(ParameterVarInt::MaxRequestId, 100); params.set_bytes(ParameterBytes::Implementation, b"test-impl".to_vec()); - let mut buf = BytesMut::new(); - params.encode(&mut buf, Version::Draft16).unwrap(); + let mut buf = Vec::new(); + params + .encode(&mut Encoder::new(&mut buf, Version::Draft16.into()), Version::Draft16) + .unwrap(); - let mut bytes = buf.freeze(); - let decoded = Parameters::decode(&mut bytes, Version::Draft16).unwrap(); + let mut bytes = bytes::Bytes::from(buf); + let decoded = crate::coding::decode_buf(&mut bytes, Version::Draft16, Parameters::decode).unwrap(); assert_eq!(decoded.get_bytes(ParameterBytes::Path), Some(b"/test".as_ref())); assert_eq!(decoded.get_varint(ParameterVarInt::MaxRequestId), Some(100)); @@ -540,11 +480,13 @@ mod tests { params.set_bytes(ParameterBytes::Path, b"/test".to_vec()); params.set_varint(ParameterVarInt::MaxRequestId, 100); - let mut buf = BytesMut::new(); - params.encode(&mut buf, Version::Draft15).unwrap(); + let mut buf = Vec::new(); + params + .encode(&mut Encoder::new(&mut buf, Version::Draft15.into()), Version::Draft15) + .unwrap(); - let mut bytes = buf.freeze(); - let decoded = Parameters::decode(&mut bytes, Version::Draft15).unwrap(); + let mut bytes = bytes::Bytes::from(buf); + let decoded = crate::coding::decode_buf(&mut bytes, Version::Draft15, Parameters::decode).unwrap(); assert_eq!(decoded.get_bytes(ParameterBytes::Path), Some(b"/test".as_ref())); assert_eq!(decoded.get_varint(ParameterVarInt::MaxRequestId), Some(100)); @@ -557,11 +499,13 @@ mod tests { params.set_varint(ParameterVarInt::MaxAuthTokenCacheSize, 4096); params.set_bytes(ParameterBytes::Implementation, b"test-impl".to_vec()); - let mut buf = BytesMut::new(); - params.encode(&mut buf, Version::Draft17).unwrap(); + let mut buf = Vec::new(); + params + .encode(&mut Encoder::new(&mut buf, Version::Draft17.into()), Version::Draft17) + .unwrap(); - let mut bytes = buf.freeze(); - let decoded = Parameters::decode(&mut bytes, Version::Draft17).unwrap(); + let mut bytes = bytes::Bytes::from(buf); + let decoded = crate::coding::decode_buf(&mut bytes, Version::Draft17, Parameters::decode).unwrap(); assert_eq!(decoded.get_bytes(ParameterBytes::Path), Some(b"/test".as_ref())); assert_eq!(decoded.get_varint(ParameterVarInt::MaxAuthTokenCacheSize), Some(4096)); @@ -569,7 +513,7 @@ mod tests { decoded.get_bytes(ParameterBytes::Implementation), Some(b"test-impl".as_ref()) ); - assert!(!bytes.has_remaining()); + assert!(bytes.is_empty()); } #[test] @@ -577,11 +521,15 @@ mod tests { let mut params = Parameters::default(); params.set_bytes(ParameterBytes::Path, b"/x".to_vec()); - let mut buf15 = BytesMut::new(); - params.encode(&mut buf15, Version::Draft15).unwrap(); + let mut buf15 = Vec::new(); + params + .encode(&mut Encoder::new(&mut buf15, Version::Draft15.into()), Version::Draft15) + .unwrap(); - let mut buf17 = BytesMut::new(); - params.encode(&mut buf17, Version::Draft17).unwrap(); + let mut buf17 = Vec::new(); + params + .encode(&mut Encoder::new(&mut buf17, Version::Draft17.into()), Version::Draft17) + .unwrap(); assert!(buf17.len() < buf15.len()); } @@ -590,14 +538,14 @@ mod tests { fn round_trip_params( version: Version, - encode_fn: impl FnOnce(&mut BytesMut, Version) -> Result<(), EncodeError>, - decode_fn: impl FnOnce(&mut bytes::Bytes, Version) -> Result<(), DecodeError>, + encode_fn: impl FnOnce(&mut Encoder<'_>, Version) -> Result<(), EncodeError>, + decode_fn: impl FnOnce(&mut Decoder<'_>, Version) -> Result<(), DecodeError>, ) { - let mut buf = BytesMut::new(); - encode_fn(&mut buf, version).unwrap(); - let mut bytes = buf.freeze(); - decode_fn(&mut bytes, version).unwrap(); - assert!(!bytes.has_remaining(), "buffer not fully consumed for {version}"); + let mut buf = Vec::new(); + encode_fn(&mut Encoder::new(&mut buf, version.into()), version).unwrap(); + let mut r = Decoder::new(&buf, version.into()); + decode_fn(&mut r, version).unwrap(); + assert!(r.is_empty(), "buffer not fully consumed for {version}"); } #[test] @@ -632,18 +580,18 @@ mod tests { (Version::Draft18, &[0x01, 0x20, 0xff][..]), (Version::Draft19, &[0x01, 0x20, 0xff][..]), ] { - let mut buf = BytesMut::new(); - encode_params!(&mut buf, version, 0x20 => u8::MAX); + let mut buf = Vec::new(); + encode_params!(&mut Encoder::new(&mut buf, version.into()), version, 0x20 => u8::MAX); assert_eq!(&buf[..], expected, "{version}"); - let mut encoded = Bytes::copy_from_slice(expected); + let mut encoded = Decoder::new(expected, version.into()); let decoded = (|| -> Result, DecodeError> { decode_params!(&mut encoded, version, 0x20 => value: Option); Ok(value) })() .expect("fixed uint8 vector should decode"); assert_eq!(decoded, Some(u8::MAX), "{version}"); - assert!(!encoded.has_remaining(), "{version}"); + assert!(encoded.is_empty(), "{version}"); } Ok(()) } @@ -713,18 +661,18 @@ mod tests { (Version::Draft19, &[0x01, 0x09, 0x80, 0xff, 0x80, 0x80][..]), (Version::Draft20, &[0x01, 0x09, 0x80, 0xff, 0x80, 0x80][..]), ] { - let mut buf = BytesMut::new(); - encode_params!(&mut buf, version, 0x09 => location.clone()); + let mut buf = Vec::new(); + encode_params!(&mut Encoder::new(&mut buf, version.into()), version, 0x09 => location.clone()); assert_eq!(&buf[..], expected, "{version}"); - let mut encoded = Bytes::copy_from_slice(expected); + let mut encoded = Decoder::new(expected, version.into()); let decoded = (|| -> Result, DecodeError> { decode_params!(&mut encoded, version, 0x09 => value: Option); Ok(value) })() .expect("fixed Location vector should decode"); assert_eq!(decoded, Some(location), "{version}"); - assert!(!encoded.has_remaining(), "{version}"); + assert!(encoded.is_empty(), "{version}"); } Ok(()) } @@ -917,12 +865,13 @@ mod tests { Version::Draft17, Version::Draft18, ] { - let mut buf = BytesMut::new(); - 1usize.encode(&mut buf, version).unwrap(); - 0x10u64.encode(&mut buf, version).unwrap(); - true.param_encode(&mut buf, version).unwrap(); + let mut buf = Vec::new(); + let mut w = Encoder::new(&mut buf, version.into()); + w.varint(VarInt::from(1usize)).unwrap(); + w.varint(VarInt::from(0x10u64)).unwrap(); + true.param_encode(&mut w, version).unwrap(); - let mut bytes = buf.freeze(); + let mut bytes = Decoder::new(&buf, version.into()); let result: Result<(), DecodeError> = (|| { decode_params!(&mut bytes, version, 0x20 => val: Option); let _ = val; @@ -945,27 +894,28 @@ mod tests { Version::Draft17, Version::Draft18, ] { - let mut buf = BytesMut::new(); + let mut buf = Vec::new(); + let mut w = Encoder::new(&mut buf, version.into()); // Encode count = 2 - 2usize.encode(&mut buf, version).unwrap(); + w.varint(VarInt::from(2usize)).unwrap(); match version { Version::Draft14 | Version::Draft15 => { // Plain (non-delta) keys: first key=0x20, second key=0x20 - 0x20u64.encode(&mut buf, version).unwrap(); - 100u8.param_encode(&mut buf, version).unwrap(); - 0x20u64.encode(&mut buf, version).unwrap(); - 200u8.param_encode(&mut buf, version).unwrap(); + w.varint(VarInt::from(0x20u64)).unwrap(); + 100u8.param_encode(&mut w, version).unwrap(); + w.varint(VarInt::from(0x20u64)).unwrap(); + 200u8.param_encode(&mut w, version).unwrap(); } _ => { // Delta-encoded: first delta=0x20 (abs=0x20), second delta=0 (abs=0x20) - 0x20u64.encode(&mut buf, version).unwrap(); - 100u8.param_encode(&mut buf, version).unwrap(); - 0u64.encode(&mut buf, version).unwrap(); - 200u8.param_encode(&mut buf, version).unwrap(); + w.varint(VarInt::from(0x20u64)).unwrap(); + 100u8.param_encode(&mut w, version).unwrap(); + w.varint(VarInt::from(0u64)).unwrap(); + 200u8.param_encode(&mut w, version).unwrap(); } } - let mut bytes = buf.freeze(); + let mut bytes = Decoder::new(&buf, version.into()); let result: Result<(), DecodeError> = (|| { decode_params!(&mut bytes, version, 0x20 => val: Option); let _ = val; diff --git a/rs/moq-net/src/ietf/properties.rs b/rs/moq-net/src/ietf/properties.rs index cf17a9b861..11ca7512dc 100644 --- a/rs/moq-net/src/ietf/properties.rs +++ b/rs/moq-net/src/ietf/properties.rs @@ -9,11 +9,10 @@ /// /// MAX_CACHE_DURATION, TIMESCALE, DEFAULT_PUBLISHER_PRIORITY, and DEFAULT_PUBLISHER_GROUP_ORDER are understood; /// the rest are parsed and discarded. -use bytes::Buf; use std::time::Duration; use crate::Timescale; -use crate::coding::{Decode, DecodeError, Encode, EncodeError}; +use crate::coding::{DecodeError, Decoder, EncodeError, Encoder, VarInt}; use super::{GroupOrder, Version}; @@ -61,7 +60,7 @@ impl Properties { /// so the caller must not append anything after it. /// /// Properties are serialized in ascending order by type, delta-encoded. - pub fn encode(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + pub fn encode(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { // Draft-16 carries the same block under the name Track Extensions, but we only write it // from draft-17 on: a draft-16 peer running an older build of this crate rejects any // trailing bytes it doesn't parse, and draft-16 never registered TIMESCALE (0x08). We @@ -74,26 +73,26 @@ impl Properties { let mut prev_type = 0; if let Some(age) = self.max_cache_duration { - 4u64.encode(w, version)?; - age.encode(w, version)?; + w.varint(VarInt::from(4u64))?; + w.varint(VarInt::try_from(age.as_millis())?)?; prev_type = 4; } if let Some(timescale) = self.timescale { - (TIMESCALE - prev_type).encode(w, version)?; - u64::from(timescale).encode(w, version)?; + w.varint(VarInt::from(TIMESCALE - prev_type))?; + w.varint(VarInt::from(u64::from(timescale)))?; prev_type = TIMESCALE; } if let Some(priority) = self.priority { - (DEFAULT_PUBLISHER_PRIORITY - prev_type).encode(w, version)?; - u64::from(priority).encode(w, version)?; + w.varint(VarInt::from(DEFAULT_PUBLISHER_PRIORITY - prev_type))?; + w.varint(VarInt::from(u64::from(priority)))?; prev_type = DEFAULT_PUBLISHER_PRIORITY; } if let Some(group_order) = self.group_order { - (DEFAULT_PUBLISHER_GROUP_ORDER - prev_type).encode(w, version)?; - u64::from(u8::from(group_order)).encode(w, version)?; + w.varint(VarInt::from(DEFAULT_PUBLISHER_GROUP_ORDER - prev_type))?; + w.varint(VarInt::from(u64::from(u8::from(group_order))))?; } Ok(()) @@ -109,7 +108,7 @@ impl Properties { /// fatal, which is what lets a relay forward properties it does not implement. /// /// Drafts before 16 have no such block, so this reads nothing and leaves the buffer alone. - pub fn decode(r: &mut R, version: Version) -> Result { + pub fn decode(r: &mut Decoder<'_>, version: Version) -> Result { let mut properties = Self::default(); // Draft-16 calls the block Track Extensions, draft-17+ Track Properties. Same encoding, @@ -122,12 +121,12 @@ impl Properties { let mut prev_type: u64 = 0; let mut i: u64 = 0; - while r.has_remaining() { + while !r.is_empty() { if i >= MAX_PROPERTIES { return Err(DecodeError::TooMany); } - let delta = u64::decode(r, version)?; + let delta = r.varint()?.into_inner(); let abs = if i == 0 { delta } else { @@ -138,7 +137,7 @@ impl Properties { if abs % 2 == 0 { // Even type: single varint value - let value = u64::decode(r, version)?; + let value = r.varint()?.into_inner(); match abs { 4 => properties.max_cache_duration = Some(Duration::from_millis(value)), TIMESCALE => { @@ -162,14 +161,11 @@ impl Properties { } } else { // Odd type: length-prefixed bytes - let len = u64::decode(r, version)? as usize; + let len = usize::try_from(r.varint()?)?; if len > MAX_KVP_VALUE_LEN { return Err(DecodeError::BoundsExceeded); } - if r.remaining() < len { - return Err(DecodeError::Short); - } - r.advance(len); + r.slice(len)?; } } @@ -180,14 +176,12 @@ impl Properties { #[cfg(test)] mod tests { use super::*; - use crate::coding::Encode; - use bytes::BytesMut; #[test] fn test_skip_empty_properties() { let mut buf = bytes::Bytes::new(); assert_eq!( - Properties::decode(&mut buf, Version::Draft17).unwrap(), + crate::coding::decode_buf(&mut buf, Version::Draft17, Properties::decode).unwrap(), Properties::default() ); } @@ -195,43 +189,63 @@ mod tests { #[test] fn test_skip_varint_property() { // Even type (0x02 = DELIVERY_TIMEOUT), varint value - let mut buf = BytesMut::new(); - 0x02u64.encode(&mut buf, Version::Draft17).unwrap(); // delta type - 5000u64.encode(&mut buf, Version::Draft17).unwrap(); // value - let mut bytes = buf.freeze(); - Properties::decode(&mut bytes, Version::Draft17).unwrap(); - assert!(!bytes.has_remaining()); + let mut buf = Vec::new(); + Encoder::new(&mut buf, Version::Draft17.into()) + .varint(VarInt::from(0x02u64)) + .unwrap(); // delta type + Encoder::new(&mut buf, Version::Draft17.into()) + .varint(VarInt::from(5000u64)) + .unwrap(); // value + let mut bytes = bytes::Bytes::from(buf); + crate::coding::decode_buf(&mut bytes, Version::Draft17, Properties::decode).unwrap(); + assert!(!!bytes.is_empty()); } #[test] fn test_skip_bytes_property() { // Odd type (0x0B = IMMUTABLE_PROPERTIES), length-prefixed - let mut buf = BytesMut::new(); - 0x0Bu64.encode(&mut buf, Version::Draft17).unwrap(); // delta type - 3u64.encode(&mut buf, Version::Draft17).unwrap(); // length + let mut buf = Vec::new(); + Encoder::new(&mut buf, Version::Draft17.into()) + .varint(VarInt::from(0x0Bu64)) + .unwrap(); // delta type + Encoder::new(&mut buf, Version::Draft17.into()) + .varint(VarInt::from(3u64)) + .unwrap(); // length buf.extend_from_slice(&[0x01, 0x02, 0x03]); // value bytes - let mut bytes = buf.freeze(); - Properties::decode(&mut bytes, Version::Draft17).unwrap(); - assert!(!bytes.has_remaining()); + let mut bytes = bytes::Bytes::from(buf); + crate::coding::decode_buf(&mut bytes, Version::Draft17, Properties::decode).unwrap(); + assert!(!!bytes.is_empty()); } #[test] fn test_skip_multiple_properties() { - let mut buf = BytesMut::new(); + let mut buf = Vec::new(); // First: type 0x02 (even), varint value - 0x02u64.encode(&mut buf, Version::Draft17).unwrap(); - 1000u64.encode(&mut buf, Version::Draft17).unwrap(); + Encoder::new(&mut buf, Version::Draft17.into()) + .varint(VarInt::from(0x02u64)) + .unwrap(); + Encoder::new(&mut buf, Version::Draft17.into()) + .varint(VarInt::from(1000u64)) + .unwrap(); // Second: delta = 0x02 → abs type 0x04 (even), varint value - 0x02u64.encode(&mut buf, Version::Draft17).unwrap(); - 2000u64.encode(&mut buf, Version::Draft17).unwrap(); + Encoder::new(&mut buf, Version::Draft17.into()) + .varint(VarInt::from(0x02u64)) + .unwrap(); + Encoder::new(&mut buf, Version::Draft17.into()) + .varint(VarInt::from(2000u64)) + .unwrap(); // Third: delta = 0x07 → abs type 0x0B (odd), length-prefixed - 0x07u64.encode(&mut buf, Version::Draft17).unwrap(); - 2u64.encode(&mut buf, Version::Draft17).unwrap(); + Encoder::new(&mut buf, Version::Draft17.into()) + .varint(VarInt::from(0x07u64)) + .unwrap(); + Encoder::new(&mut buf, Version::Draft17.into()) + .varint(VarInt::from(2u64)) + .unwrap(); buf.extend_from_slice(&[0xAA, 0xBB]); - let mut bytes = buf.freeze(); - Properties::decode(&mut bytes, Version::Draft17).unwrap(); - assert!(!bytes.has_remaining()); + let mut bytes = bytes::Bytes::from(buf); + crate::coding::decode_buf(&mut bytes, Version::Draft17, Properties::decode).unwrap(); + assert!(!!bytes.is_empty()); } #[test] @@ -243,12 +257,17 @@ mod tests { group_order: Some(GroupOrder::Descending), }; - let mut buf = BytesMut::new(); - properties.encode(&mut buf, Version::Draft18).unwrap(); + let mut buf = Vec::new(); + properties + .encode(&mut Encoder::new(&mut buf, Version::Draft18.into()), Version::Draft18) + .unwrap(); - let mut bytes = buf.freeze(); - assert_eq!(Properties::decode(&mut bytes, Version::Draft18).unwrap(), properties); - assert!(!bytes.has_remaining()); + let mut bytes = bytes::Bytes::from(buf); + assert_eq!( + crate::coding::decode_buf(&mut bytes, Version::Draft18, Properties::decode).unwrap(), + properties + ); + assert!(!!bytes.is_empty()); } #[test] @@ -262,11 +281,15 @@ mod tests { Version::Draft21, Version::Draft22, ] { - let mut buf = BytesMut::new(); - 0x0eu64.encode(&mut buf, version).unwrap(); - 256u64.encode(&mut buf, version).unwrap(); + let mut buf = Vec::new(); + Encoder::new(&mut buf, version.into()) + .varint(VarInt::from(0x0eu64)) + .unwrap(); + Encoder::new(&mut buf, version.into()) + .varint(VarInt::from(256u64)) + .unwrap(); assert!(matches!( - Properties::decode(&mut buf.freeze(), version), + crate::coding::decode_buf(&mut bytes::Bytes::from(buf), version, Properties::decode), Err(DecodeError::InvalidValue) )); } @@ -276,26 +299,34 @@ mod tests { /// the "publisher decides" it means in the draft-14 fields. #[test] fn test_rejects_zero_group_order() { - let mut buf = BytesMut::new(); - 0x22u64.encode(&mut buf, Version::Draft18).unwrap(); - 0u64.encode(&mut buf, Version::Draft18).unwrap(); - - let mut bytes = buf.freeze(); - assert!(Properties::decode(&mut bytes, Version::Draft18).is_err()); + let mut buf = Vec::new(); + Encoder::new(&mut buf, Version::Draft18.into()) + .varint(VarInt::from(0x22u64)) + .unwrap(); + Encoder::new(&mut buf, Version::Draft18.into()) + .varint(VarInt::from(0u64)) + .unwrap(); + + let mut bytes = bytes::Bytes::from(buf); + assert!(crate::coding::decode_buf(&mut bytes, Version::Draft18, Properties::decode).is_err()); } /// Draft-16 carries the same block under the name Track Extensions. We don't write one /// there, but a peer that does must be understood rather than faulted. #[test] fn test_decodes_draft16_track_extensions() { - let mut buf = BytesMut::new(); - 0x22u64.encode(&mut buf, Version::Draft16).unwrap(); - 2u64.encode(&mut buf, Version::Draft16).unwrap(); - - let mut bytes = buf.freeze(); - let properties = Properties::decode(&mut bytes, Version::Draft16).unwrap(); + let mut buf = Vec::new(); + Encoder::new(&mut buf, Version::Draft16.into()) + .varint(VarInt::from(0x22u64)) + .unwrap(); + Encoder::new(&mut buf, Version::Draft16.into()) + .varint(VarInt::from(2u64)) + .unwrap(); + + let mut bytes = bytes::Bytes::from(buf); + let properties = crate::coding::decode_buf(&mut bytes, Version::Draft16, Properties::decode).unwrap(); assert_eq!(properties.group_order, Some(GroupOrder::Descending)); - assert!(!bytes.has_remaining()); + assert!(!!bytes.is_empty()); } /// The group order property is delta-encoded against the timescale that precedes it, @@ -309,11 +340,16 @@ mod tests { group_order: Some(GroupOrder::Descending), }; - let mut buf = BytesMut::new(); - properties.encode(&mut buf, Version::Draft18).unwrap(); + let mut buf = Vec::new(); + properties + .encode(&mut Encoder::new(&mut buf, Version::Draft18.into()), Version::Draft18) + .unwrap(); - let mut bytes = buf.freeze(); - assert_eq!(Properties::decode(&mut bytes, Version::Draft18).unwrap(), properties); - assert!(!bytes.has_remaining()); + let mut bytes = bytes::Bytes::from(buf); + assert_eq!( + crate::coding::decode_buf(&mut bytes, Version::Draft18, Properties::decode).unwrap(), + properties + ); + assert!(!!bytes.is_empty()); } } diff --git a/rs/moq-net/src/ietf/publish.rs b/rs/moq-net/src/ietf/publish.rs index 7a5b24164b..a71ffd615b 100644 --- a/rs/moq-net/src/ietf/publish.rs +++ b/rs/moq-net/src/ietf/publish.rs @@ -108,7 +108,7 @@ use std::borrow::Cow; use crate::{ Path, - coding::{Decode, DecodeError, Encode, EncodeError}, + coding::{Decode, DecodeError, Decoder, Encode, EncodeError, Encoder, VarInt}, ietf::{ Filter, GroupOrder, Location, Parameters, Properties, RequestId, namespace::{decode_namespace, encode_namespace}, @@ -190,7 +190,7 @@ impl PublishDone<'_> { impl Message for PublishDone<'_> { const ID: u64 = 0x0b; - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { if matches!(version, Version::Draft14 | Version::Draft15 | Version::Draft16) { self.request_id .expect("request_id required for draft14-16") @@ -198,21 +198,21 @@ impl Message for PublishDone<'_> { } else { assert!(self.request_id.is_none(), "request_id must be None for draft17+"); } - self.status_code.encode(w, version)?; - self.stream_count.encode(w, version)?; - self.reason_phrase.encode(w, version)?; + w.varint(VarInt::from(self.status_code))?; + w.varint(VarInt::from(self.stream_count))?; + w.string(&self.reason_phrase)?; Ok(()) } - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { let request_id = if matches!(version, Version::Draft14 | Version::Draft15 | Version::Draft16) { Some(RequestId::decode(r, version)?) } else { None }; - let status_code = u64::decode(r, version)?; - let stream_count = u64::decode(r, version)?; - let reason_phrase = Cow::::decode(r, version)?; + let status_code = r.varint()?.into_inner(); + let stream_count = r.varint()?.into_inner(); + let reason_phrase = Cow::Owned(r.string()?); Ok(Self { request_id, @@ -240,14 +240,14 @@ pub struct Publish<'a> { impl Message for Publish<'_> { const ID: u64 = 0x1D; - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.request_id.encode(w, version)?; if version == Version::Draft17 { - 0u64.encode(w, version)?; // required_request_id_delta = 0 + w.varint(VarInt::from(0u64))?; // required_request_id_delta = 0 } - encode_namespace(w, &self.track_namespace, version)?; - self.track_name.encode(w, version)?; - self.track_alias.encode(w, version)?; + encode_namespace(w, &self.track_namespace)?; + w.string(&self.track_name)?; + w.varint(VarInt::from(self.track_alias))?; match version { Version::Draft14 => { @@ -256,15 +256,15 @@ impl Message for Publish<'_> { .unwrap_or(GroupOrder::Ascending) .encode(w, version)?; if let Some(location) = &self.largest_location { - true.encode(w, version)?; + w.bool(true); location.encode(w, version)?; } else { - false.encode(w, version)?; + w.bool(false); } - self.forward.encode(w, version)?; + w.bool(self.forward); // parameters - 0u8.encode(w, version)?; + w.u8(0u8); } _ => { // GROUP_ORDER is a legal PUBLISH parameter only through draft-15; a later peer @@ -290,24 +290,24 @@ impl Message for Publish<'_> { Ok(()) } - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { let request_id = RequestId::decode(r, version)?; if version == Version::Draft17 { - let _required_request_id_delta = u64::decode(r, version)?; + let _required_request_id_delta = r.varint()?.into_inner(); } - let track_namespace = decode_namespace(r, version)?; - let track_name = Cow::::decode(r, version)?; - let track_alias = u64::decode(r, version)?; + let track_namespace = decode_namespace(r)?; + let track_name = Cow::Owned(r.string()?); + let track_alias = r.varint()?.into_inner(); match version { Version::Draft14 => { let group_order = GroupOrder::decode(r, version)?.any_to_descending(); - let content_exists = bool::decode(r, version)?; + let content_exists = r.bool()?; let largest_location = match content_exists { true => Some(Location::decode(r, version)?), false => None, }; - let forward = bool::decode(r, version)?; + let forward = r.bool()?; // parameters let _params = Parameters::decode(r, version)?; @@ -386,7 +386,7 @@ pub struct PublishOk { impl Message for PublishOk { const ID: u64 = 0x1E; - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { if matches!(version, Version::Draft14 | Version::Draft15 | Version::Draft16) { self.request_id .expect("request_id required for draft14-16") @@ -397,14 +397,14 @@ impl Message for PublishOk { match version { Version::Draft14 => { - self.forward.encode(w, version)?; - self.subscriber_priority.encode(w, version)?; + w.bool(self.forward); + w.u8(self.subscriber_priority); self.group_order.encode(w, version)?; // Same as SUBSCRIBE: the Location an absolute filter carries is dropped on // decode, so encoding one would truncate the message. self.filter.encode(w, version)?; // no parameters - 0u8.encode(w, version)?; + w.u8(0u8); } // Draft-20 moved the subscription parameters out of PUBLISH_OK; they belong to // PUBLISH and REQUEST_UPDATE now, so a PUBLISH_OK carries none of them. @@ -422,7 +422,7 @@ impl Message for PublishOk { Ok(()) } - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { let request_id = if matches!(version, Version::Draft14 | Version::Draft15 | Version::Draft16) { Some(RequestId::decode(r, version)?) } else { @@ -431,8 +431,8 @@ impl Message for PublishOk { match version { Version::Draft14 => { - let forward = bool::decode(r, version)?; - let subscriber_priority = u8::decode(r, version)?; + let forward = r.bool()?; + let subscriber_priority = r.u8()?; let group_order = GroupOrder::decode(r, version)?; let filter = Filter::decode(r, version)?; @@ -483,17 +483,17 @@ pub struct PublishError<'a> { impl Message for PublishError<'_> { const ID: u64 = 0x1F; - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.request_id.encode(w, version)?; - self.error_code.encode(w, version)?; - self.reason_phrase.encode(w, version)?; + w.varint(VarInt::from(self.error_code))?; + w.string(&self.reason_phrase)?; Ok(()) } - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { let request_id = RequestId::decode(r, version)?; - let error_code = u64::decode(r, version)?; - let reason_phrase = Cow::::decode(r, version)?; + let error_code = r.varint()?.into_inner(); + let reason_phrase = Cow::Owned(r.string()?); Ok(Self { request_id, error_code, @@ -511,50 +511,53 @@ mod tests { /// decode, which kills the session instead of letting the peer get its refusal. #[test] fn publish_accepts_the_relocated_subscription_parameters() -> Result<(), EncodeError> { + let version = Version::Draft20; let mut body = Vec::new(); - RequestId(1).encode(&mut body, Version::Draft20).unwrap(); - super::super::namespace::encode_namespace(&mut body, &crate::Path::new("broadcast"), Version::Draft20).unwrap(); - "video".encode(&mut body, Version::Draft20).unwrap(); - 1u64.encode(&mut body, Version::Draft20).unwrap(); // track alias + let w = &mut Encoder::new(&mut body, version.into()); + RequestId(1).encode(w, version)?; + super::super::namespace::encode_namespace(w, &crate::Path::new("broadcast"))?; + w.string("video")?; + w.varint(1u64.into())?; // track alias // SUBSCRIBER_PRIORITY then LOCATION_FILTER, delta encoded from 0. - encode_params!(&mut body, Version::Draft20, + encode_params!(w, version, 0x20 => 128u8, 0x21 => Filter::NextObject, ); - Properties::default().encode(&mut body, Version::Draft20).unwrap(); + Properties::default().encode(w, version)?; - let mut buf = bytes::Bytes::from(body); - Publish::decode_msg(&mut buf, Version::Draft20).expect("draft-20 PUBLISH parameters must parse"); + Publish::decode_msg(&mut Decoder::new(&body, version.into()), version) + .expect("draft-20 PUBLISH parameters must parse"); Ok(()) } /// They arrived in draft-20, so an earlier peer sending one is still a violation. #[test] fn older_drafts_reject_the_relocated_parameters() -> Result<(), EncodeError> { + let version = Version::Draft19; let mut body = Vec::new(); - RequestId(1).encode(&mut body, Version::Draft19).unwrap(); - super::super::namespace::encode_namespace(&mut body, &crate::Path::new("broadcast"), Version::Draft19).unwrap(); - "video".encode(&mut body, Version::Draft19).unwrap(); - 1u64.encode(&mut body, Version::Draft19).unwrap(); - encode_params!(&mut body, Version::Draft19, 0x20 => 128u8); - Properties::default().encode(&mut body, Version::Draft19).unwrap(); - - let mut buf = bytes::Bytes::from(body); - assert!(Publish::decode_msg(&mut buf, Version::Draft19).is_err()); + let w = &mut Encoder::new(&mut body, version.into()); + RequestId(1).encode(w, version)?; + super::super::namespace::encode_namespace(w, &crate::Path::new("broadcast"))?; + w.string("video")?; + w.varint(1u64.into())?; + encode_params!(w, version, 0x20 => 128u8); + Properties::default().encode(w, version)?; + + assert!(Publish::decode_msg(&mut Decoder::new(&body, version.into()), version).is_err()); Ok(()) } - use bytes::BytesMut; fn encode_message(msg: &M, version: Version) -> Vec { - let mut buf = BytesMut::new(); - msg.encode_msg(&mut buf, version).unwrap(); + let mut buf = Vec::new(); + msg.encode_msg(&mut Encoder::new(&mut buf, version.into()), version) + .unwrap(); buf.to_vec() } fn decode_message(bytes: &[u8], version: Version) -> Result { let mut buf = bytes::Bytes::from(bytes.to_vec()); - M::decode_msg(&mut buf, version) + crate::coding::decode_buf(&mut buf, version, M::decode_msg) } #[test] diff --git a/rs/moq-net/src/ietf/publish_namespace.rs b/rs/moq-net/src/ietf/publish_namespace.rs index 863826dfeb..3ec5793692 100644 --- a/rs/moq-net/src/ietf/publish_namespace.rs +++ b/rs/moq-net/src/ietf/publish_namespace.rs @@ -30,12 +30,12 @@ impl PublishNamespace<'_> { /// The negotiation is session state rather than anything in the message, so the /// caller supplies it. A negotiated session that omits HOP_PATH is a protocol /// violation, which surfaces here as [`DecodeError::InvalidValue`]. - pub fn decode_body(r: &mut R, version: Version, negotiated: bool) -> Result { + pub fn decode_body(r: &mut Decoder<'_>, version: Version, negotiated: bool) -> Result { let request_id = RequestId::decode(r, version)?; if version == Version::Draft17 { - let _required_request_id_delta = u64::decode(r, version)?; + let _required_request_id_delta = r.varint()?.into_inner(); } - let track_namespace = decode_namespace(r, version)?; + let track_namespace = decode_namespace(r)?; let cluster = decode_cluster_params(r, version, negotiated)?; Ok(Self { @@ -49,16 +49,16 @@ impl PublishNamespace<'_> { impl Message for PublishNamespace<'_> { const ID: u64 = 0x06; - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.request_id.encode(w, version)?; if version == Version::Draft17 { - 0u64.encode(w, version)?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) + w.varint(VarInt::from(0u64))?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) } - encode_namespace(w, &self.track_namespace, version)?; + encode_namespace(w, &self.track_namespace)?; encode_cluster_params(w, version, self.cluster.as_ref()) } - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { Self::decode_body(r, version, false) } } @@ -101,12 +101,12 @@ impl PublishNamespaceUpdate { impl Message for PublishNamespaceUpdate { const ID: u64 = 0x02; - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { match version { Version::Draft14 | Version::Draft15 | Version::Draft16 => return Err(EncodeError::Version), Version::Draft17 => { self.request_id.encode(w, version)?; - 0u64.encode(w, version)?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) + w.varint(VarInt::from(0u64))?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) } _ => self.request_id.encode(w, version)?, } @@ -117,12 +117,12 @@ impl Message for PublishNamespaceUpdate { Ok(()) } - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { let request_id = match version { Version::Draft14 | Version::Draft15 | Version::Draft16 => return Err(DecodeError::Version), Version::Draft17 => { let request_id = RequestId::decode(r, version)?; - let _required_request_id_delta = u64::decode(r, version)?; + let _required_request_id_delta = r.varint()?.into_inner(); request_id } _ => RequestId::decode(r, version)?, @@ -139,8 +139,8 @@ impl Message for PublishNamespaceUpdate { /// /// On a session that negotiated the MoQ Cluster extension every advertisement carries /// HOP_PATH; ROUTE_COST is optional and absent means 0, so a free path sends nothing. -pub(super) fn encode_cluster_params( - w: &mut W, +pub(super) fn encode_cluster_params( + w: &mut Encoder<'_>, version: Version, advert: Option<&cluster::Advert>, ) -> Result<(), EncodeError> { @@ -158,8 +158,8 @@ pub(super) fn encode_cluster_params( } /// Read the Parameters field of an advertisement. See [`encode_cluster_params`]. -pub(super) fn decode_cluster_params( - r: &mut R, +pub(super) fn decode_cluster_params( + r: &mut Decoder<'_>, version: Version, negotiated: bool, ) -> Result, DecodeError> { @@ -190,12 +190,12 @@ pub struct PublishNamespaceOk { impl Message for PublishNamespaceOk { const ID: u64 = 0x07; - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.request_id.encode(w, version)?; Ok(()) } - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { let request_id = RequestId::decode(r, version)?; Ok(Self { request_id }) } @@ -212,17 +212,17 @@ pub struct PublishNamespaceError<'a> { impl Message for PublishNamespaceError<'_> { const ID: u64 = 0x08; - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.request_id.encode(w, version)?; - self.error_code.encode(w, version)?; - self.reason_phrase.encode(w, version)?; + w.varint(VarInt::from(self.error_code))?; + w.string(&self.reason_phrase)?; Ok(()) } - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { let request_id = RequestId::decode(r, version)?; - let error_code = u64::decode(r, version)?; - let reason_phrase = Cow::::decode(r, version)?; + let error_code = r.varint()?.into_inner(); + let reason_phrase = Cow::Owned(r.string()?); Ok(Self { request_id, @@ -245,10 +245,10 @@ pub struct PublishNamespaceDone<'a> { impl Message for PublishNamespaceDone<'_> { const ID: u64 = 0x09; - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { match version { Version::Draft14 | Version::Draft15 => { - encode_namespace(w, &self.track_namespace, version)?; + encode_namespace(w, &self.track_namespace)?; } Version::Draft16 => { self.request_id.encode(w, version)?; @@ -258,10 +258,10 @@ impl Message for PublishNamespaceDone<'_> { Ok(()) } - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { match version { Version::Draft14 | Version::Draft15 => { - let track_namespace = decode_namespace(r, version)?; + let track_namespace = decode_namespace(r)?; Ok(Self { track_namespace, request_id: RequestId(0), @@ -294,10 +294,10 @@ pub struct PublishNamespaceCancel<'a> { impl Message for PublishNamespaceCancel<'_> { const ID: u64 = 0x0c; - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { match version { Version::Draft14 | Version::Draft15 => { - encode_namespace(w, &self.track_namespace, version)?; + encode_namespace(w, &self.track_namespace)?; } Version::Draft16 => { self.request_id.encode(w, version)?; @@ -306,15 +306,15 @@ impl Message for PublishNamespaceCancel<'_> { return Err(EncodeError::Version); } } - self.error_code.encode(w, version)?; - self.reason_phrase.encode(w, version)?; + w.varint(VarInt::from(self.error_code))?; + w.string(&self.reason_phrase)?; Ok(()) } - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { let (track_namespace, request_id) = match version { Version::Draft14 | Version::Draft15 => { - let track_namespace = decode_namespace(r, version)?; + let track_namespace = decode_namespace(r)?; (track_namespace, RequestId(0)) } Version::Draft16 => { @@ -325,8 +325,8 @@ impl Message for PublishNamespaceCancel<'_> { return Err(DecodeError::Version); } }; - let error_code = u64::decode(r, version)?; - let reason_phrase = Cow::::decode(r, version)?; + let error_code = r.varint()?.into_inner(); + let reason_phrase = Cow::Owned(r.string()?); Ok(Self { track_namespace, request_id, @@ -339,17 +339,17 @@ impl Message for PublishNamespaceCancel<'_> { #[cfg(test)] mod tests { use super::*; - use bytes::BytesMut; fn encode_message(msg: &M, version: Version) -> Vec { - let mut buf = BytesMut::new(); - msg.encode_msg(&mut buf, version).unwrap(); + let mut buf = Vec::new(); + msg.encode_msg(&mut Encoder::new(&mut buf, version.into()), version) + .unwrap(); buf.to_vec() } fn decode_message(bytes: &[u8], version: Version) -> Result { let mut buf = bytes::Bytes::from(bytes.to_vec()); - M::decode_msg(&mut buf, version) + crate::coding::decode_buf(&mut buf, version, M::decode_msg) } #[test] @@ -478,8 +478,11 @@ mod tests { request_id: RequestId(42), }; - let mut buf = BytesMut::new(); - assert!(msg.encode_msg(&mut buf, Version::Draft18).is_err()); + let mut buf = Vec::new(); + assert!( + msg.encode_msg(&mut Encoder::new(&mut buf, Version::Draft18.into()), Version::Draft18) + .is_err() + ); } #[test] @@ -489,8 +492,11 @@ mod tests { request_id: RequestId(42), }; - let mut buf = BytesMut::new(); - assert!(msg.encode_msg(&mut buf, Version::Draft17).is_err()); + let mut buf = Vec::new(); + assert!( + msg.encode_msg(&mut Encoder::new(&mut buf, Version::Draft17.into()), Version::Draft17) + .is_err() + ); } #[test] @@ -502,8 +508,11 @@ mod tests { reason_phrase: "Shutdown".into(), }; - let mut buf = BytesMut::new(); - assert!(msg.encode_msg(&mut buf, Version::Draft17).is_err()); + let mut buf = Vec::new(); + assert!( + msg.encode_msg(&mut Encoder::new(&mut buf, Version::Draft17.into()), Version::Draft17) + .is_err() + ); } fn hop_path(ids: &[u64]) -> cluster::HopPath { @@ -555,8 +564,11 @@ mod tests { hops: None, cost: Some(0), }; - let mut buf = BytesMut::new(); - assert!(msg.encode_msg(&mut buf, Version::Draft16).is_err()); + let mut buf = Vec::new(); + assert!( + msg.encode_msg(&mut Encoder::new(&mut buf, Version::Draft16.into()), Version::Draft16) + .is_err() + ); assert!(decode_message::(&[0x02, 0x00], Version::Draft16).is_err()); } diff --git a/rs/moq-net/src/ietf/publisher.rs b/rs/moq-net/src/ietf/publisher.rs index 66e18e069a..75c49ed782 100644 --- a/rs/moq-net/src/ietf/publisher.rs +++ b/rs/moq-net/src/ietf/publisher.rs @@ -15,7 +15,7 @@ use web_transport_trait::poll::SendStream as _; use crate::{ AsPath, Error, Timescale, Timestamp, - coding::{Stream, Writer}, + coding::{Encoder, Stream, VarInt, Writer}, ietf::{self, Control, EndLocation, FetchHeader, FetchType, Filter, GroupOrder, Location, RequestId}, track::Subscription, util::{MaybeBoxedExt, MaybeSendBox}, @@ -382,10 +382,11 @@ where pub fn handle_stream( &self, id: u64, - mut data: bytes::Bytes, + body: ietf::Body, stream: Stream, ) -> Result, Error> { let this = self.clone(); + let mut data = body.decoder(this.version); let task = match id { ietf::Subscribe::ID => { let msg = ietf::Subscribe::decode_msg(&mut data, this.version)?; @@ -562,7 +563,7 @@ where .map(|fill| (fill_range(fill, msg.filter, edge.largest), cache, timescale)); // Send SubscribeOk on the stream - stream.writer.encode(&ietf::SubscribeOk::ID).await?; + stream.writer.encode(&VarInt::from(ietf::SubscribeOk::ID)).await?; stream .writer .encode(&ietf::SubscribeOk { @@ -647,7 +648,7 @@ where Ok(()) => (ietf::PublishDoneStatus::TrackEnded, "track ended"), Err(_) => (ietf::PublishDoneStatus::InternalError, "internal error"), }; - let _ = stream.writer.encode(&ietf::PublishDone::ID).await; + let _ = stream.writer.encode(&VarInt::from(ietf::PublishDone::ID)).await; let _ = stream .writer .encode(&ietf::PublishDone { @@ -701,7 +702,7 @@ where match self.version { Version::Draft14 => { - writer.encode(&ietf::SubscribeError::ID).await?; + writer.encode(&VarInt::from(ietf::SubscribeError::ID)).await?; writer .encode(&ietf::SubscribeError { request_id, @@ -711,7 +712,7 @@ where .await?; } Version::Draft15 | Version::Draft16 => { - writer.encode(&ietf::RequestError::ID).await?; + writer.encode(&VarInt::from(ietf::RequestError::ID)).await?; writer .encode(&ietf::RequestError { request_id: Some(request_id), @@ -722,7 +723,7 @@ where .await?; } _ => { - writer.encode(&ietf::RequestError::ID).await?; + writer.encode(&VarInt::from(ietf::RequestError::ID)).await?; writer .encode(&ietf::RequestError { request_id: None, @@ -768,7 +769,7 @@ where stream.set_priority(priority); let res = async { - stream.encode(&FetchHeader::TYPE).await?; + stream.encode(&VarInt::from(FetchHeader::TYPE)).await?; stream.encode(&FetchHeader { request_id }).await?; let FillServe::Group { sequence, skip, until } = fill else { @@ -871,9 +872,9 @@ where version, ) .await?; - stream.encode(&(frame.payload.len() as u64)).await?; + stream.encode(&VarInt::from(frame.payload.len())).await?; if frame.payload.is_empty() && matches!(version, Version::Draft14 | Version::Draft15) { - stream.encode(&0u64).await?; + stream.encode(&VarInt::from(0u64)).await?; } if !frame.payload.is_empty() { let mut payload = frame.payload; @@ -918,9 +919,9 @@ where .await?; index += 1; - stream.encode(&frame.size).await?; + stream.encode(&VarInt::from(frame.size)).await?; if frame.size == 0 && matches!(version, Version::Draft14 | Version::Draft15) { - stream.encode(&0u64).await?; + stream.encode(&VarInt::from(0u64)).await?; } loop { let chunk = { @@ -965,19 +966,28 @@ where // than written empty: the track declared no units to read a timestamp in. let properties = match timescale { Some(timescale) => { - let mut properties = bytes::BytesMut::new(); - ietf::encode_object_time(&mut properties, timestamp, timescale, version)?; - Some(properties.to_vec()) + let mut properties = Vec::new(); + ietf::encode_object_time( + &mut Encoder::new(&mut properties, version.into()), + timestamp, + timescale, + version, + )?; + Some(properties) } None => None, }; if version == Version::Draft14 { - stream.encode(&sequence).await?; - stream.encode(&0u64).await?; - stream.encode(&object).await?; - stream.encode(&0u8).await?; - stream.encode(&properties.unwrap_or_default()).await?; + let properties = properties.unwrap_or_default(); + stream.buffer(&VarInt::from(sequence))?; + stream.buffer(&VarInt::from(0u64))?; + stream.buffer(&VarInt::from(object))?; + // Publisher priority, a raw byte. + stream.buffer_raw(&[0]); + stream.buffer(&VarInt::from(properties.len()))?; + stream.buffer_raw(&properties); + std::future::poll_fn(|cx| stream.poll_flush(cx)).await?; return Ok(()); } @@ -1178,7 +1188,7 @@ where // FETCH_OK on every draft, never REQUEST_OK: section 5.2 allows exactly one FETCH_OK or // REQUEST_ERROR in answer to a FETCH, and REQUEST_OK's own definition lists the other // requests it answers without ever naming this one. - stream.writer.encode(&ietf::FetchOk::ID).await?; + stream.writer.encode(&VarInt::from(ietf::FetchOk::ID)).await?; stream .writer .encode(&ietf::FetchOk { @@ -1197,7 +1207,7 @@ where let uni = self.session.open_uni().await.map_err(Error::from_transport)?; let mut writer = Writer::new(uni, self.version); writer.set_priority(priority); - writer.encode(&FetchHeader::TYPE).await?; + writer.encode(&VarInt::from(FetchHeader::TYPE)).await?; writer .encode(&FetchHeader { request_id: msg.request_id, @@ -1214,9 +1224,9 @@ where self.version, ) .await?; - writer.encode(&(frame.payload.len() as u64)).await?; + writer.encode(&VarInt::from(frame.payload.len())).await?; if frame.payload.is_empty() && matches!(self.version, Version::Draft14 | Version::Draft15) { - writer.encode(&0u64).await?; + writer.encode(&VarInt::from(0u64)).await?; } if !frame.payload.is_empty() { let mut payload = frame.payload; @@ -1236,7 +1246,7 @@ where async fn reject_track_status(&self, mut stream: Stream, request_id: RequestId) -> Result<(), Error> { let error_code = request::to_code(&Error::Unsupported, request::Kind::TrackStatus, self.version); if self.version == Version::Draft14 { - stream.writer.encode(&0x0fu64).await?; // TRACK_STATUS_ERROR has the SUBSCRIBE_ERROR body. + stream.writer.encode(&VarInt::from(0x0fu64)).await?; // TRACK_STATUS_ERROR has the SUBSCRIBE_ERROR body. stream .writer .encode(&ietf::SubscribeError { @@ -1246,7 +1256,7 @@ where }) .await?; } else { - stream.writer.encode(&ietf::RequestError::ID).await?; + stream.writer.encode(&VarInt::from(ietf::RequestError::ID)).await?; stream .writer .encode(&ietf::RequestError { @@ -1288,7 +1298,7 @@ where match self.version { Version::Draft14 => { - writer.encode(&ietf::FetchError::ID).await?; + writer.encode(&VarInt::from(ietf::FetchError::ID)).await?; writer .encode(&ietf::FetchError { request_id, @@ -1298,7 +1308,7 @@ where .await?; } Version::Draft15 | Version::Draft16 => { - writer.encode(&ietf::RequestError::ID).await?; + writer.encode(&VarInt::from(ietf::RequestError::ID)).await?; writer .encode(&ietf::RequestError { request_id: Some(request_id), @@ -1309,7 +1319,7 @@ where .await?; } _ => { - writer.encode(&ietf::RequestError::ID).await?; + writer.encode(&VarInt::from(ietf::RequestError::ID)).await?; writer .encode(&ietf::RequestError { request_id: None, @@ -1407,7 +1417,7 @@ where match (advert.wanted(), held) { (true, _) => { tracing::debug!(broadcast = %absolute, "namespace"); - stream.writer.encode(&ietf::Namespace::ID).await?; + stream.writer.encode(&VarInt::from(ietf::Namespace::ID)).await?; stream .writer .encode(&ietf::Namespace { @@ -1418,7 +1428,7 @@ where } (false, true) => { tracing::debug!(broadcast = %absolute, "namespace_done"); - stream.writer.encode(&ietf::NamespaceDone::ID).await?; + stream.writer.encode(&VarInt::from(ietf::NamespaceDone::ID)).await?; stream .writer .encode(&ietf::NamespaceDone { @@ -1466,7 +1476,7 @@ where return Ok(Refused::No); }; - request.writer.encode(&ietf::PublishNamespace::ID).await?; + request.writer.encode(&VarInt::from(ietf::PublishNamespace::ID)).await?; request .writer .encode(&ietf::PublishNamespace { @@ -1478,10 +1488,11 @@ where // Bounded for the same reason the open is: a peer that takes the stream and answers // nothing would park this loop forever, and every withdrawal queued behind it. - let Some((type_id, mut data)) = self.read_response(&mut request).await? else { + let Some((type_id, body)) = self.read_response(&mut request).await? else { tracing::debug!(broadcast = %self.origin.absolute(path), "no answer to the advertisement"); return Ok(Refused::No); }; + let mut data = body.decoder(self.version); match (self.version, type_id) { (Version::Draft14, ietf::PublishNamespaceOk::ID) => { @@ -1570,16 +1581,21 @@ where let request_id = self.control.next_request_id(&self.runtime).await?; let update = ietf::PublishNamespaceUpdate::between(request_id, &held, &next); - request.stream.writer.encode(&ietf::PublishNamespaceUpdate::ID).await?; + request + .stream + .writer + .encode(&VarInt::from(ietf::PublishNamespaceUpdate::ID)) + .await?; request.stream.writer.encode(&update).await?; let absolute = self.origin.absolute(&request.path).to_owned(); - let Some((type_id, mut data)) = self.read_response(&mut request.stream).await? else { + let Some((type_id, body)) = self.read_response(&mut request.stream).await? else { tracing::debug!(broadcast = %absolute, "no answer to the update"); // Abrupt: a peer that never answers is not owed the FIN handshake. requests.remove(suffix); return Ok(Refused::No); }; + let mut data = body.decoder(self.version); match type_id { ietf::RequestOk::ID => { @@ -1650,12 +1666,11 @@ where /// behavior seen a step later: a stream the peer accepts and never answers on holds the /// loop just as effectively as one it never grants. Giving up records nothing, so the /// namespace stays outstanding and the retry re-offers it. - async fn read_response(&self, request: &mut Stream) -> Result, Error> { + async fn read_response(&self, request: &mut Stream) -> Result, Error> { let mut read = std::pin::pin!(async { - let type_id: u64 = request.reader.decode().await?; - let size: u16 = request.reader.decode().await?; - let data = request.reader.read_exact(size as usize).await?; - Ok::<_, Error>((type_id, data)) + let type_id = request.reader.decode::().await?.into_inner(); + let body: ietf::Body = request.reader.decode().await?; + Ok::<_, Error>((type_id, body)) }); let mut timeout = crate::runtime::Deadline::after(&self.runtime, ADVERTISE_TIMEOUT); @@ -1705,7 +1720,7 @@ where } } Target::Inline(stream) => { - stream.writer.encode(&ietf::NamespaceDone::ID).await?; + stream.writer.encode(&VarInt::from(ietf::NamespaceDone::ID)).await?; stream .writer .encode(&ietf::NamespaceDone { @@ -1837,7 +1852,10 @@ where // Send OK response match self.version { Version::Draft14 => { - stream.writer.encode(&ietf::SubscribeNamespaceOk::ID).await?; + stream + .writer + .encode(&VarInt::from(ietf::SubscribeNamespaceOk::ID)) + .await?; stream .writer .encode(&ietf::SubscribeNamespaceOk { @@ -1846,7 +1864,7 @@ where .await?; } Version::Draft15 | Version::Draft16 => { - stream.writer.encode(&ietf::RequestOk::ID).await?; + stream.writer.encode(&VarInt::from(ietf::RequestOk::ID)).await?; stream .writer .encode(&ietf::RequestOk { @@ -1856,7 +1874,7 @@ where .await?; } _ => { - stream.writer.encode(&ietf::RequestOk::ID).await?; + stream.writer.encode(&VarInt::from(ietf::RequestOk::ID)).await?; stream .writer .encode(&ietf::RequestOk { @@ -2124,9 +2142,9 @@ impl TrackServe { flags: ietf::GroupFlags::default(), })?; // Object ID delta 0, then an empty object whose status is END_OF_TRACK. - writer.buffer(&0u64)?; - writer.buffer(&0u64)?; - writer.encode(&END_OF_TRACK).await?; + writer.buffer(&VarInt::from(0u64))?; + writer.buffer(&VarInt::from(0u64))?; + writer.encode(&VarInt::from(END_OF_TRACK)).await?; // PUBLISH_DONE follows once this closes, like every other data stream. writer.close().await } @@ -2492,20 +2510,25 @@ fn buffer_object_info( timescale: Option, version: Version, ) -> Result<(), Error> { - writer.buffer(&delta)?; + writer.buffer(&VarInt::from(delta))?; if let Some(timescale) = timescale.filter(|_| has_extensions) { // Per-object extension headers carry the frame's presentation timestamp. - let mut ext = bytes::BytesMut::new(); - ietf::encode_object_time(&mut ext, timestamp, timescale, version)?; - writer.buffer(&(ext.len() as u64))?; + let mut ext = Vec::new(); + ietf::encode_object_time( + &mut Encoder::new(&mut ext, version.into()), + timestamp, + timescale, + version, + )?; + writer.buffer(&VarInt::from(ext.len()))?; writer.buffer_raw(&ext); } - writer.buffer(&size)?; + writer.buffer(&VarInt::from(size))?; if size == 0 { - // Have to write the object status too. - writer.buffer(&0u8)?; + // Have to write the object status too: Normal (0). + writer.buffer(&VarInt::ZERO)?; } Ok(()) } @@ -2598,7 +2621,8 @@ mod group_priority_test { let written = log.writes.lock().unwrap().clone(); let mut buf = bytes::Bytes::from(written); - let header = ietf::GroupHeader::decode(&mut buf, Version::Draft14).expect("a group header"); + let header = + crate::coding::decode_buf(&mut buf, Version::Draft14, ietf::GroupHeader::decode).expect("a group header"); assert_eq!( header.publisher_priority, priority::to_wire(hang_audio_priority()), @@ -2904,7 +2928,7 @@ mod serve_tests { match version { Version::Draft14 => { - writer.encode(&ietf::SubscribeError::ID).await.unwrap(); + writer.encode(&VarInt::from(ietf::SubscribeError::ID)).await.unwrap(); writer .encode(&ietf::SubscribeError { request_id: RequestId(REQUEST_ID), @@ -2915,7 +2939,7 @@ mod serve_tests { .unwrap(); } _ => { - writer.encode(&ietf::RequestError::ID).await.unwrap(); + writer.encode(&VarInt::from(ietf::RequestError::ID)).await.unwrap(); writer .encode(&ietf::RequestError { request_id: match version { @@ -2963,12 +2987,13 @@ mod serve_tests { track_namespace: crate::Path::new("live"), track_name: "video".into(), }; - let mut body = bytes::BytesMut::new(); - msg.encode_msg(&mut body, version).unwrap(); + let mut body = Vec::new(); + msg.encode_msg(&mut Encoder::new(&mut body, version.into()), version) + .unwrap(); let mark = h.log.writes.lock().unwrap().len(); h.publisher .clone() - .handle_stream(ietf::TrackStatus::ID, body.freeze(), stream) + .handle_stream(ietf::TrackStatus::ID, ietf::Body(bytes::Bytes::from(body)), stream) .unwrap() .await; let actual = h.log.writes.lock().unwrap()[mark..].to_vec(); @@ -2977,7 +3002,7 @@ mod serve_tests { let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(log.clone()), version); if version == Version::Draft14 { - writer.encode(&0x0fu64).await.unwrap(); + writer.encode(&VarInt::from(0x0fu64)).await.unwrap(); writer .encode(&ietf::SubscribeError { request_id: RequestId(REQUEST_ID), @@ -2987,7 +3012,7 @@ mod serve_tests { .await .unwrap(); } else { - writer.encode(&ietf::RequestError::ID).await.unwrap(); + writer.encode(&VarInt::from(ietf::RequestError::ID)).await.unwrap(); writer .encode(&ietf::RequestError { request_id: matches!(version, Version::Draft15 | Version::Draft16) @@ -3234,10 +3259,10 @@ mod serve_tests { let response = joining_fetch(&h, mark).await.unwrap(); let mut buf = bytes::Bytes::from(response); - let id = u64::decode(&mut buf, version).unwrap(); + let id = crate::coding::decode_varint(&mut buf, version).unwrap(); assert_eq!(id, ietf::FetchOk::ID, "{version}: not a FETCH_OK"); - let ok = ietf::FetchOk::decode(&mut buf, version).unwrap(); + let ok = crate::coding::decode_buf(&mut buf, version, ietf::FetchOk::decode).unwrap(); assert_eq!( ok.request_id, match version { @@ -3259,30 +3284,38 @@ mod serve_tests { "{version}: wrong saved end boundary" ); - assert_eq!(u64::decode(&mut buf, version).unwrap(), FetchHeader::TYPE); - assert_eq!(FetchHeader::decode(&mut buf, version).unwrap().request_id, FETCH_ID); + assert_eq!( + crate::coding::decode_varint(&mut buf, version).unwrap(), + FetchHeader::TYPE + ); + assert_eq!( + crate::coding::decode_buf(&mut buf, version, FetchHeader::decode) + .unwrap() + .request_id, + FETCH_ID + ); for (index, payload) in payloads.iter().enumerate() { if version == Version::Draft14 { - assert_eq!(u64::decode(&mut buf, version).unwrap(), LATEST); - assert_eq!(u64::decode(&mut buf, version).unwrap(), 0); - assert_eq!(u64::decode(&mut buf, version).unwrap(), index as u64); - assert_eq!(u8::decode(&mut buf, version).unwrap(), 0); + assert_eq!(crate::coding::decode_varint(&mut buf, version).unwrap(), LATEST); + assert_eq!(crate::coding::decode_varint(&mut buf, version).unwrap(), 0); + assert_eq!(crate::coding::decode_varint(&mut buf, version).unwrap(), index as u64); + assert_eq!(bytes::Buf::try_get_u8(&mut buf).unwrap(), 0); } else { assert_eq!( - u64::decode(&mut buf, version).unwrap(), + crate::coding::decode_varint(&mut buf, version).unwrap(), if index == 0 { 0x3c } else { 0x20 } ); if index == 0 { - assert_eq!(u64::decode(&mut buf, version).unwrap(), LATEST); - assert_eq!(u64::decode(&mut buf, version).unwrap(), 0); - assert_eq!(u8::decode(&mut buf, version).unwrap(), 0); + assert_eq!(crate::coding::decode_varint(&mut buf, version).unwrap(), LATEST); + assert_eq!(crate::coding::decode_varint(&mut buf, version).unwrap(), 0); + assert_eq!(bytes::Buf::try_get_u8(&mut buf).unwrap(), 0); } } - let _properties = Vec::::decode(&mut buf, version).unwrap(); - let size = u64::decode(&mut buf, version).unwrap() as usize; + let _properties = crate::coding::decode_buf(&mut buf, version, |r, _| Ok(r.bytes()?.to_vec())).unwrap(); + let size = crate::coding::decode_varint(&mut buf, version).unwrap() as usize; assert_eq!(size, payload.len()); if size == 0 && matches!(version, Version::Draft14 | Version::Draft15) { - assert_eq!(u64::decode(&mut buf, version).unwrap(), 0); + assert_eq!(crate::coding::decode_varint(&mut buf, version).unwrap(), 0); } assert_eq!(buf.split_to(size).as_ref(), *payload); } @@ -3290,7 +3323,10 @@ mod serve_tests { let mark = h.log.writes.lock().unwrap().len(); assert!(futures::poll!(serve.as_mut()).is_pending()); let mut tail = bytes::Bytes::from(h.log.writes.lock().unwrap()[mark..].to_vec()); - assert_eq!(u64::decode(&mut tail, version).unwrap(), payloads.len() as u64); + assert_eq!( + crate::coding::decode_varint(&mut tail, version).unwrap(), + payloads.len() as u64 + ); assert_eq!(occurrences(&h.log, b"new-object"), 1); assert!(h.log.resets().is_empty(), "{version}: an answered fetch must not reset"); } @@ -3304,13 +3340,13 @@ mod serve_tests { publish_groups(&mut h, 5); settle().await; let stream = Stream::open(&mut h.session.clone(), version).await.unwrap(); - let mut data = bytes::BytesMut::new(); + let mut data = Vec::new(); subscribe(Filter::NextObject, None) - .encode_msg(&mut data, version) + .encode_msg(&mut Encoder::new(&mut data, version.into()), version) .unwrap(); let mut serving = h .publisher - .handle_stream(ietf::Subscribe::ID, data.freeze(), stream) + .handle_stream(ietf::Subscribe::ID, ietf::Body(bytes::Bytes::from(data)), stream) .unwrap(); let mut fetching = Box::pin(joining_fetch(&h, 0)); assert!( @@ -3337,13 +3373,13 @@ mod serve_tests { "allow SUBSCRIBE to arrive on its stream" ); let stream = Stream::open(&mut h.session.clone(), version).await.unwrap(); - let mut data = bytes::BytesMut::new(); + let mut data = Vec::new(); subscribe(Filter::NextObject, None) - .encode_msg(&mut data, version) + .encode_msg(&mut Encoder::new(&mut data, version.into()), version) .unwrap(); let mut serving = h .publisher - .handle_stream(ietf::Subscribe::ID, data.freeze(), stream) + .handle_stream(ietf::Subscribe::ID, ietf::Body(bytes::Bytes::from(data)), stream) .unwrap(); tokio::select! { response = &mut fetching => { response.unwrap(); } @@ -3358,29 +3394,33 @@ mod serve_tests { for version in JOINING_DRAFTS { let h = serve(version); let stream = Stream::open(&mut h.session.clone(), version).await.unwrap(); - let mut data = bytes::BytesMut::new(); + let mut data = Vec::new(); subscribe(Filter::NextObject, None) - .encode_msg(&mut data, version) + .encode_msg(&mut Encoder::new(&mut data, version.into()), version) .unwrap(); let serving = h .publisher - .handle_stream(ietf::Subscribe::ID, data.freeze(), stream) + .handle_stream(ietf::Subscribe::ID, ietf::Body(bytes::Bytes::from(data)), stream) .unwrap(); let mut fetching = Box::pin(joining_fetch(&h, 0)); assert!(futures::poll!(fetching.as_mut()).is_pending()); drop(serving); let mut response = bytes::Bytes::from(fetching.await.unwrap()); - let id = u64::decode(&mut response, version).unwrap(); + let id = crate::coding::decode_varint(&mut response, version).unwrap(); if version == Version::Draft14 { assert_eq!(id, ietf::FetchError::ID); assert_eq!( - ietf::FetchError::decode(&mut response, version).unwrap().error_code, + crate::coding::decode_buf(&mut response, version, ietf::FetchError::decode) + .unwrap() + .error_code, invalid_joining_request_id(version) ); } else { assert_eq!(id, ietf::RequestError::ID); assert_eq!( - ietf::RequestError::decode(&mut response, version).unwrap().error_code, + crate::coding::decode_buf(&mut response, version, ietf::RequestError::decode) + .unwrap() + .error_code, invalid_joining_request_id(version) ); } @@ -3399,9 +3439,14 @@ mod serve_tests { .clone() .run_subscribe_stream(stream, subscribe(Filter::NextObject, None)); let mut response = bytes::Bytes::from(joining_fetch(&h, 0).await.unwrap()); - assert_eq!(u64::decode(&mut response, version).unwrap(), ietf::RequestError::ID); assert_eq!( - ietf::RequestError::decode(&mut response, version).unwrap().error_code, + crate::coding::decode_varint(&mut response, version).unwrap(), + ietf::RequestError::ID + ); + assert_eq!( + crate::coding::decode_buf(&mut response, version, ietf::RequestError::decode) + .unwrap() + .error_code, 0x2 ); assert!(response.is_empty()); @@ -3428,9 +3473,14 @@ mod serve_tests { registered(&h, serving.as_mut()).await; let mark = h.log.writes.lock().unwrap().len(); let mut response = bytes::Bytes::from(joining_fetch(&h, mark).await.expect("refuse before opening data")); - assert_eq!(u64::decode(&mut response, version).unwrap(), ietf::RequestError::ID); assert_eq!( - ietf::RequestError::decode(&mut response, version).unwrap().error_code, + crate::coding::decode_varint(&mut response, version).unwrap(), + ietf::RequestError::ID + ); + assert_eq!( + crate::coding::decode_buf(&mut response, version, ietf::RequestError::decode) + .unwrap() + .error_code, does_not_exist(version) ); assert!(response.is_empty()); @@ -3447,17 +3497,21 @@ mod serve_tests { run_live(&mut h, subscribe(Filter::NextObject, None)).await; let mark = h.log.writes.lock().unwrap().len(); let mut buf = bytes::Bytes::from(joining_fetch(&h, mark).await.unwrap()); - let id = u64::decode(&mut buf, version).unwrap(); + let id = crate::coding::decode_varint(&mut buf, version).unwrap(); if version == Version::Draft14 { assert_eq!(id, ietf::FetchError::ID); assert_eq!( - ietf::FetchError::decode(&mut buf, version).unwrap().error_code, + crate::coding::decode_buf(&mut buf, version, ietf::FetchError::decode) + .unwrap() + .error_code, invalid_joining_request_id(version) ); } else { assert_eq!(id, ietf::RequestError::ID); assert_eq!( - ietf::RequestError::decode(&mut buf, version).unwrap().error_code, + crate::coding::decode_buf(&mut buf, version, ietf::RequestError::decode) + .unwrap() + .error_code, invalid_joining_request_id(version) ); } @@ -3487,8 +3541,16 @@ mod serve_tests { assert_eq!(h.log.closes()[0].0, 0x3); } else { let mut buf = bytes::Bytes::from(result.unwrap()); - assert_eq!(u64::decode(&mut buf, version).unwrap(), ietf::RequestError::ID); - assert_eq!(ietf::RequestError::decode(&mut buf, version).unwrap().error_code, 0x3); + assert_eq!( + crate::coding::decode_varint(&mut buf, version).unwrap(), + ietf::RequestError::ID + ); + assert_eq!( + crate::coding::decode_buf(&mut buf, version, ietf::RequestError::decode) + .unwrap() + .error_code, + 0x3 + ); } } } @@ -3510,17 +3572,21 @@ mod serve_tests { registered(&h, serving.as_mut()).await; let mark = h.log.writes.lock().unwrap().len(); let mut buf = bytes::Bytes::from(joining_fetch(&h, mark).await.unwrap()); - let id = u64::decode(&mut buf, version).unwrap(); + let id = crate::coding::decode_varint(&mut buf, version).unwrap(); if version == Version::Draft14 { assert_eq!(id, ietf::FetchError::ID); assert_eq!( - ietf::FetchError::decode(&mut buf, version).unwrap().error_code, + crate::coding::decode_buf(&mut buf, version, ietf::FetchError::decode) + .unwrap() + .error_code, invalid_range(version) ); } else { assert_eq!(id, ietf::RequestError::ID); assert_eq!( - ietf::RequestError::decode(&mut buf, version).unwrap().error_code, + crate::coding::decode_buf(&mut buf, version, ietf::RequestError::decode) + .unwrap() + .error_code, invalid_range(version) ); } @@ -3533,9 +3599,12 @@ mod serve_tests { async fn a_joining_fetch_is_not_supported_on_draft20() { let h = serve(Version::Draft20); let mut buf = bytes::Bytes::from(joining_fetch(&h, 0).await.unwrap()); - assert_eq!(u64::decode(&mut buf, Version::Draft20).unwrap(), ietf::RequestError::ID); assert_eq!( - ietf::RequestError::decode(&mut buf, Version::Draft20) + crate::coding::decode_varint(&mut buf, Version::Draft20).unwrap(), + ietf::RequestError::ID + ); + assert_eq!( + crate::coding::decode_buf(&mut buf, Version::Draft20, ietf::RequestError::decode) .unwrap() .error_code, 0x3 @@ -3580,8 +3649,13 @@ mod serve_tests { // SUBSCRIBE_OK is the first thing written, before any group stream opens. let writes = h.log.writes.lock().unwrap().clone(); let mut buf = writes.as_slice(); - assert_eq!(u64::decode(&mut buf, version).unwrap(), ietf::SubscribeOk::ID); - ietf::SubscribeOk::decode(&mut buf, version).unwrap().largest + assert_eq!( + crate::coding::decode_varint(&mut buf, version).unwrap(), + ietf::SubscribeOk::ID + ); + crate::coding::decode_buf(&mut buf, version, ietf::SubscribeOk::decode) + .unwrap() + .largest } #[tokio::test] @@ -3972,7 +4046,7 @@ mod tests { crate::lite::test_transport::SinkSend::new(expected.clone()), Version::Draft17, ); - writer.encode(&ietf::RequestOk::ID).await.unwrap(); + writer.encode(&VarInt::from(ietf::RequestOk::ID)).await.unwrap(); writer .encode(&ietf::RequestOk { request_id: None, @@ -4072,7 +4146,7 @@ mod tests { let log = crate::lite::test_transport::Log::default(); let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(log.clone()), version); - writer.encode(&ietf::RequestError::ID).await.unwrap(); + writer.encode(&VarInt::from(ietf::RequestError::ID)).await.unwrap(); writer .encode(&ietf::RequestError { request_id: matches!(version, Version::Draft15 | Version::Draft16).then_some(RequestId(1)), @@ -4144,7 +4218,10 @@ mod tests { match version { Version::Draft14 => { - writer.encode(&ietf::PublishNamespaceOk::ID).await.unwrap(); + writer + .encode(&VarInt::from(ietf::PublishNamespaceOk::ID)) + .await + .unwrap(); writer .encode(&ietf::PublishNamespaceOk { request_id: RequestId(1), @@ -4153,7 +4230,7 @@ mod tests { .unwrap(); } Version::Draft15 | Version::Draft16 => { - writer.encode(&ietf::RequestOk::ID).await.unwrap(); + writer.encode(&VarInt::from(ietf::RequestOk::ID)).await.unwrap(); writer .encode(&ietf::RequestOk { request_id: Some(RequestId(1)), @@ -4164,7 +4241,7 @@ mod tests { } // Draft-17+ dropped the request id: the response rides the request's stream. _ => { - writer.encode(&ietf::RequestOk::ID).await.unwrap(); + writer.encode(&VarInt::from(ietf::RequestOk::ID)).await.unwrap(); writer .encode(&ietf::RequestOk { request_id: None, @@ -4325,7 +4402,7 @@ mod tests { let ok = crate::lite::test_transport::Log::default(); let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(ok.clone()), VERSION); - writer.encode(&ietf::RequestOk::ID).await.unwrap(); + writer.encode(&VarInt::from(ietf::RequestOk::ID)).await.unwrap(); writer .encode(&ietf::RequestOk { request_id: None, @@ -4666,7 +4743,10 @@ mod tests { async fn request_update(version: Version, msg: &ietf::PublishNamespaceUpdate) -> Vec { let log = crate::lite::test_transport::Log::default(); let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(log.clone()), version); - writer.encode(&ietf::PublishNamespaceUpdate::ID).await.unwrap(); + writer + .encode(&VarInt::from(ietf::PublishNamespaceUpdate::ID)) + .await + .unwrap(); writer.encode(msg).await.unwrap(); log.writes.lock().unwrap().clone() } @@ -4755,7 +4835,7 @@ mod tests { let counted = crate::lite::test_transport::Log::default(); let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(counted.clone()), VERSION); - writer.encode(&ietf::RequestOk::ID).await.unwrap(); + writer.encode(&VarInt::from(ietf::RequestOk::ID)).await.unwrap(); writer .encode(&ietf::RequestOk { request_id: None, diff --git a/rs/moq-net/src/ietf/request.rs b/rs/moq-net/src/ietf/request.rs index dc0cfb2889..e1e0aaf198 100644 --- a/rs/moq-net/src/ietf/request.rs +++ b/rs/moq-net/src/ietf/request.rs @@ -1,6 +1,6 @@ use std::borrow::Cow; -use crate::coding::{Decode, DecodeError, Encode, EncodeError}; +use crate::coding::{Decode, DecodeError, Decoder, Encode, EncodeError, Encoder, VarInt}; use super::Message; use super::active_count::ACTIVE_COUNT_PARAM; @@ -29,15 +29,15 @@ impl std::fmt::Display for RequestId { } impl Encode for RequestId { - fn encode(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { - self.0.encode(w, version)?; + fn encode(&self, w: &mut Encoder<'_>, _: Version) -> Result<(), EncodeError> { + w.varint(VarInt::from(self.0))?; Ok(()) } } impl Decode for RequestId { - fn decode(r: &mut R, version: Version) -> Result { - let request_id = u64::decode(r, version)?; + fn decode(r: &mut Decoder<'_>, _: Version) -> Result { + let request_id = r.varint()?.into_inner(); Ok(Self(request_id)) } } @@ -50,12 +50,12 @@ pub struct MaxRequestId { impl Message for MaxRequestId { const ID: u64 = 0x15; - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.request_id.encode(w, version)?; Ok(()) } - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { let request_id = RequestId::decode(r, version)?; Ok(Self { request_id }) } @@ -69,12 +69,12 @@ pub struct RequestsBlocked { impl Message for RequestsBlocked { const ID: u64 = 0x1a; - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.request_id.encode(w, version)?; Ok(()) } - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { let request_id = RequestId::decode(r, version)?; Ok(Self { request_id }) } @@ -94,7 +94,7 @@ pub struct RequestOk { impl Message for RequestOk { const ID: u64 = 0x07; - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { if matches!(version, Version::Draft14 | Version::Draft15 | Version::Draft16) { self.request_id .expect("request_id required for draft14-16") @@ -106,7 +106,7 @@ impl Message for RequestOk { Ok(()) } - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { let request_id = if matches!(version, Version::Draft14 | Version::Draft15 | Version::Draft16) { Some(RequestId::decode(r, version)?) } else { @@ -136,7 +136,7 @@ pub struct RequestError<'a> { impl Message for RequestError<'_> { const ID: u64 = 0x05; - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { if matches!(version, Version::Draft14 | Version::Draft15 | Version::Draft16) { self.request_id .expect("request_id required for draft14-16") @@ -144,26 +144,26 @@ impl Message for RequestError<'_> { } else { assert!(self.request_id.is_none(), "request_id must be None for draft17+"); } - self.error_code.encode(w, version)?; + w.varint(VarInt::from(self.error_code))?; if !matches!(version, Version::Draft14 | Version::Draft15) { - self.retry_interval.encode(w, version)?; + w.varint(VarInt::from(self.retry_interval))?; } - self.reason_phrase.encode(w, version)?; + w.string(&self.reason_phrase)?; Ok(()) } - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { let request_id = if matches!(version, Version::Draft14 | Version::Draft15 | Version::Draft16) { Some(RequestId::decode(r, version)?) } else { None }; - let error_code = u64::decode(r, version)?; + let error_code = r.varint()?.into_inner(); let retry_interval = match version { Version::Draft14 | Version::Draft15 => 0, - _ => u64::decode(r, version)?, + _ => r.varint()?.into_inner(), }; - let reason_phrase = Cow::::decode(r, version)?; + let reason_phrase = Cow::Owned(r.string()?); Ok(Self { request_id, error_code, @@ -176,17 +176,17 @@ impl Message for RequestError<'_> { #[cfg(test)] mod tests { use super::*; - use bytes::BytesMut; fn encode_message(msg: &M, version: Version) -> Vec { - let mut buf = BytesMut::new(); - msg.encode_msg(&mut buf, version).unwrap(); + let mut buf = Vec::new(); + msg.encode_msg(&mut Encoder::new(&mut buf, version.into()), version) + .unwrap(); buf.to_vec() } fn decode_message(bytes: &[u8], version: Version) -> Result { let mut buf = bytes::Bytes::from(bytes.to_vec()); - M::decode_msg(&mut buf, version) + crate::coding::decode_buf(&mut buf, version, M::decode_msg) } #[test] diff --git a/rs/moq-net/src/ietf/session.rs b/rs/moq-net/src/ietf/session.rs index 6243fc339f..3d8ace7137 100644 --- a/rs/moq-net/src/ietf/session.rs +++ b/rs/moq-net/src/ietf/session.rs @@ -1,7 +1,7 @@ use crate::origin; use crate::{ Error, Hop, SessionError, StreamError, - coding::{Decode, DecodeError, Encode, Reader, Stream, Writer}, + coding::{Decode, DecodeError, Encode, Reader, Stream, VarInt, Writer}, ietf::{self, FetchHeader, RequestId}, setup, util::{MaybeBoxedExt, MaybeSendBox, TaskSet, err_only}, @@ -474,15 +474,14 @@ pub async fn accept_setup( let recv = session.accept_uni().await.map_err(Error::from_transport)?; let mut reader: Reader = Reader::new(recv, outer_version); - if reader.decode_peek::().await? != setup::SETUP_V17 { + if reader.decode_peek::().await?.into_inner() != setup::SETUP_V17 { // Not the SETUP (group data this early is unexpected). Reject and keep waiting. reader.abort(&Error::UnexpectedStream); continue; } let setup: setup::Setup = reader.decode().await?; - let mut bytes = setup.parameters.clone(); - let params = ietf::Parameters::decode(&mut bytes, version)?; + let (params, _) = ietf::Parameters::decode_slice(&setup.parameters, version)?; let path = match params.get_bytes(ietf::ParameterBytes::Path) { Some(bytes) => Some( std::str::from_utf8(bytes) @@ -505,8 +504,7 @@ pub async fn accept_setup( /// Parse the Setup Options we act on out of a raw SETUP parameter block. fn decode_peer_setup(parameters: bytes::Bytes, version: Version) -> Result { - let mut bytes = parameters; - let params = ietf::Parameters::decode(&mut bytes, version)?; + let (params, _) = ietf::Parameters::decode_slice(¶meters, version)?; peer_from_params(¶ms, version) } @@ -577,17 +575,8 @@ async fn run_setup( // Frame as [type_id varint][size u16][body], the same shape as the // control-stream messages this channel otherwise carries. - let mut body = bytes::BytesMut::new(); - msg.encode_msg(&mut body, version)?; - let size: u16 = body - .len() - .try_into() - .map_err(|_| Error::BoundsExceeded(crate::coding::BoundsExceeded))?; - let mut writer = writer.with_version(version); - writer.encode(&ietf::GoAway::ID).await?; - writer.encode(&size).await?; - writer.write_all(&mut std::io::Cursor::new(body)).await?; + writer.encode_message(&msg).await?; crate::goaway::enforce(&runtime, &mut session, payload.timeout).await; session.closed().await; @@ -638,14 +627,14 @@ where // the peer wrote to. Failing the loop here would tear down the whole session // over a single stream the peer had already given up on. Only death is // tolerated: bytes that arrive and do not parse stay session-fatal. - let kind: u64 = match tasks + let kind = match tasks .drive(|waiter| { let mut cx = waiter.context(); - reader.poll_decode_peek(&mut cx) + reader.poll_decode_peek::(&mut cx) }) .await { - Ok(kind) => kind, + Ok(kind) => kind.into_inner(), Err(err @ (Error::Cancel | Error::Stream(_) | Error::Remote(_) | Error::Decode(DecodeError::Short))) => { tracing::debug!(%err, "dropping uni stream that died before its type"); continue; @@ -726,7 +715,7 @@ async fn run_uni_group( where S: crate::transport::poll::Boxable, { - let kind: u64 = stream.decode_peek().await?; + let kind = stream.decode_peek::().await?.into_inner(); // SUBGROUP_HEADER type bytes match the form 0b0XX1XXXX (spec §11.4.2): // draft-14-17 use 0x10-0x1D and 0x30-0x3D, draft-18 adds 0x40 (FIRST_OBJECT) @@ -778,20 +767,17 @@ where // The intermediate results live outside the poll closure, so a Pending // mid-header resumes where it left off. let mut hdr_id: Option = None; - let mut hdr_size: Option = None; let header = tasks .drive(|waiter| { let mut cx = waiter.context(); let id = match hdr_id { Some(id) => id, - None => *hdr_id.insert(std::task::ready!(stream.reader.poll_decode(&mut cx))?), - }; - let size = match hdr_size { - Some(size) => size, - None => *hdr_size.insert(std::task::ready!(stream.reader.poll_decode(&mut cx))?), + None => { + *hdr_id.insert(std::task::ready!(stream.reader.poll_decode::(&mut cx))?.into_inner()) + } }; - let data = std::task::ready!(stream.reader.poll_read_exact(&mut cx, size as usize))?; - std::task::Poll::Ready(Ok::<_, Error>((id, data))) + let body = std::task::ready!(stream.reader.poll_decode::(&mut cx))?; + std::task::Poll::Ready(Ok::<_, Error>((id, body))) }) .await; // Same tolerance as `run_unis`: a request stream that dies before its header @@ -835,13 +821,13 @@ async fn run_goaway( version: Version, goaway: crate::goaway::Protocol, ) -> Result<(), Error> { - let id: u64 = match reader.decode_maybe().await? { - Some(id) => id, + let id = match reader.decode_maybe::().await? { + Some(id) => id.into_inner(), None => return Ok(()), }; - let size: u16 = reader.decode::().await?; - let mut data = reader.read_exact(size as usize).await?; + let body: ietf::Body = reader.decode().await?; + let mut data = body.decoder(version); if id != ietf::GoAway::ID { return Err(Error::UnexpectedMessage); @@ -868,12 +854,12 @@ async fn run_goaway( // control stream enforces, so close over it here too rather than logging; // anything else is merely unexpected and discarded. loop { - let id: u64 = match reader.decode_maybe().await? { - Some(id) => id, + let id = match reader.decode_maybe::().await? { + Some(id) => id.into_inner(), None => return Ok(()), }; - let size: u16 = reader.decode::().await?; - let mut data = reader.read_exact(size as usize).await?; + let body: ietf::Body = reader.decode().await?; + let mut data = body.decoder(version); if id == ietf::GoAway::ID { let msg = ietf::GoAway::decode_msg(&mut data, version)?; @@ -906,7 +892,7 @@ mod tests { let log = crate::lite::test_transport::Log::default(); let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(log.clone()), version); - writer.encode(&ietf::RequestOk::ID).await.unwrap(); + writer.encode(&VarInt::from(ietf::RequestOk::ID)).await.unwrap(); writer .encode(&ietf::RequestOk { request_id: None, @@ -914,7 +900,7 @@ mod tests { }) .await .unwrap(); - writer.encode(&ietf::Namespace::ID).await.unwrap(); + writer.encode(&VarInt::from(ietf::Namespace::ID)).await.unwrap(); writer .encode(&ietf::Namespace { suffix: crate::Path::new("cam"), @@ -1292,7 +1278,7 @@ mod tests { let log = crate::lite::test_transport::Log::default(); let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(log.clone()), version); - writer.encode(&ietf::PublishNamespace::ID).await.unwrap(); + writer.encode(&VarInt::from(ietf::PublishNamespace::ID)).await.unwrap(); writer .encode(&ietf::PublishNamespace { request_id: RequestId(1), @@ -1303,7 +1289,10 @@ mod tests { .unwrap(); for _ in 0..2 { - writer.encode(&ietf::PublishNamespaceDone::ID).await.unwrap(); + writer + .encode(&VarInt::from(ietf::PublishNamespaceDone::ID)) + .await + .unwrap(); writer .encode(&ietf::PublishNamespaceDone { track_namespace: crate::Path::new("room/host"), @@ -1326,7 +1315,10 @@ mod tests { Version::Draft14, ); - writer.encode(&ietf::PublishNamespaceOk::ID).await.unwrap(); + writer + .encode(&VarInt::from(ietf::PublishNamespaceOk::ID)) + .await + .unwrap(); writer.encode(&ietf::PublishNamespaceOk { request_id }).await.unwrap(); let writes = log.writes.lock().unwrap(); diff --git a/rs/moq-net/src/ietf/subscribe.rs b/rs/moq-net/src/ietf/subscribe.rs index dd9c1e8b82..dcfd43310f 100644 --- a/rs/moq-net/src/ietf/subscribe.rs +++ b/rs/moq-net/src/ietf/subscribe.rs @@ -24,12 +24,12 @@ use super::Version; pub struct IncludeProperties(pub bool); impl Param for IncludeProperties { - fn param_encode(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { - vec![u8::from(self.0)].encode(w, version) + fn param_encode(&self, w: &mut Encoder<'_>, _: Version) -> Result<(), EncodeError> { + w.bytes(&[u8::from(self.0)]) } - fn param_decode(r: &mut R, version: Version) -> Result { - match Vec::::decode(r, version)?[..] { + fn param_decode(r: &mut Decoder<'_>, _: Version) -> Result { + match r.bytes()?.to_vec()[..] { // The draft allows exactly 0 or 1; anything else is a protocol violation. [0] => Ok(Self(false)), [1] => Ok(Self(true)), @@ -58,20 +58,20 @@ pub struct Subscribe<'a> { impl Message for Subscribe<'_> { const ID: u64 = 0x03; - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { let request_id = RequestId::decode(r, version)?; if version == Version::Draft17 { - let _required_request_id_delta = u64::decode(r, version)?; + let _required_request_id_delta = r.varint()?.into_inner(); } - let track_namespace = decode_namespace(r, version)?; - let track_name = Cow::::decode(r, version)?; + let track_namespace = decode_namespace(r)?; + let track_name = Cow::Owned(r.string()?); match version { Version::Draft14 => { - let subscriber_priority = u8::decode(r, version)?; + let subscriber_priority = r.u8()?; let group_order = GroupOrder::decode(r, version)?; - let forward = bool::decode(r, version)?; + let forward = r.bool()?; if !forward { return Err(DecodeError::Unsupported); } @@ -149,22 +149,22 @@ impl Message for Subscribe<'_> { } } - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.request_id.encode(w, version)?; if version == Version::Draft17 { - 0u64.encode(w, version)?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) + w.varint(VarInt::from(0u64))?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) } - encode_namespace(w, &self.track_namespace, version)?; - self.track_name.encode(w, version)?; + encode_namespace(w, &self.track_namespace)?; + w.string(&self.track_name)?; match version { Version::Draft14 => { - self.subscriber_priority.encode(w, version)?; + w.u8(self.subscriber_priority); self.group_order.encode(w, version)?; - true.encode(w, version)?; // forward + w.bool(true); // forward self.filter.encode(w, version)?; - 0u8.encode(w, version)?; // no parameters + w.u8(0u8); // no parameters } _ => { // FILL_PARAMETERS arrived in draft-20. Sending it to an older peer would be an @@ -210,7 +210,7 @@ pub struct SubscribeOk { impl Message for SubscribeOk { const ID: u64 = 0x04; - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { if matches!(version, Version::Draft14 | Version::Draft15 | Version::Draft16) { self.request_id .expect("request_id required for draft14-16") @@ -218,20 +218,20 @@ impl Message for SubscribeOk { } else { assert!(self.request_id.is_none(), "request_id must be None for draft17+"); } - self.track_alias.encode(w, version)?; + w.varint(VarInt::from(self.track_alias))?; match version { Version::Draft14 => { - 0u64.encode(w, version)?; // expires = 0 + w.varint(VarInt::from(0u64))?; // expires = 0 self.properties .group_order .unwrap_or(GroupOrder::Ascending) .encode(w, version)?; - self.largest.is_some().encode(w, version)?; + w.bool(self.largest.is_some()); if let Some(largest) = self.largest { largest.encode(w, version)?; } - 0u8.encode(w, version)?; // no parameters + w.u8(0u8); // no parameters } _ => { // GROUP_ORDER is a legal SUBSCRIBE_OK parameter only through draft-15; a later @@ -257,13 +257,13 @@ impl Message for SubscribeOk { Ok(()) } - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { let request_id = if matches!(version, Version::Draft14 | Version::Draft15 | Version::Draft16) { Some(RequestId::decode(r, version)?) } else { None }; - let track_alias = u64::decode(r, version)?; + let track_alias = r.varint()?.into_inner(); let mut properties = Properties::default(); let mut largest = None; @@ -271,11 +271,11 @@ impl Message for SubscribeOk { Version::Draft14 => { // EXPIRES is when the publisher expects to end the subscription. That end // arrives as PUBLISH_DONE regardless, so there is nothing to act on. - let _expires = u64::decode(r, version)?; + let _expires = r.varint()?.into_inner(); properties.group_order = Some(GroupOrder::decode(r, version)?.any_to_descending()); - if bool::decode(r, version)? { + if r.bool()? { largest = Some(Location::decode(r, version)?); } @@ -331,17 +331,17 @@ pub struct SubscribeError<'a> { impl Message for SubscribeError<'_> { const ID: u64 = 0x05; - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.request_id.encode(w, version)?; - self.error_code.encode(w, version)?; - self.reason_phrase.encode(w, version)?; + w.varint(VarInt::from(self.error_code))?; + w.string(&self.reason_phrase)?; Ok(()) } - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { let request_id = RequestId::decode(r, version)?; - let error_code = u64::decode(r, version)?; - let reason_phrase = Cow::::decode(r, version)?; + let error_code = r.varint()?.into_inner(); + let reason_phrase = Cow::Owned(r.string()?); Ok(Self { request_id, @@ -360,12 +360,12 @@ pub struct Unsubscribe { impl Message for Unsubscribe { const ID: u64 = 0x0a; - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.request_id.encode(w, version)?; Ok(()) } - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { let request_id = RequestId::decode(r, version)?; Ok(Self { request_id }) } @@ -385,7 +385,7 @@ pub struct SubscribeUpdate { impl Message for SubscribeUpdate { const ID: u64 = 0x02; - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { match version { Version::Draft14 => { self.request_id.encode(w, version)?; @@ -393,10 +393,10 @@ impl Message for SubscribeUpdate { .expect("subscription_request_id required for draft14") .encode(w, version)?; self.start_location.encode(w, version)?; - self.end_group.encode(w, version)?; - self.subscriber_priority.encode(w, version)?; - self.forward.encode(w, version)?; - 0u8.encode(w, version)?; // no parameters + w.varint(VarInt::from(self.end_group))?; + w.u8(self.subscriber_priority); + w.bool(self.forward); + w.u8(0u8); // no parameters } Version::Draft15 | Version::Draft16 => { self.request_id.encode(w, version)?; @@ -417,7 +417,7 @@ impl Message for SubscribeUpdate { // REQUEST_UPDATE self.request_id.encode(w, version)?; if matches!(version, Version::Draft17) { - 0u64.encode(w, version)?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) + w.varint(VarInt::from(0u64))?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) } encode_params!(w, version, 0x10 => self.forward, @@ -430,15 +430,15 @@ impl Message for SubscribeUpdate { Ok(()) } - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { match version { Version::Draft14 => { let request_id = RequestId::decode(r, version)?; let subscription_request_id = Some(RequestId::decode(r, version)?); let start_location = Location::decode(r, version)?; - let end_group = u64::decode(r, version)?; - let subscriber_priority = u8::decode(r, version)?; - let forward = bool::decode(r, version)?; + let end_group = r.varint()?.into_inner(); + let subscriber_priority = r.u8()?; + let forward = r.bool()?; let _parameters = Parameters::decode(r, version)?; Ok(Self { @@ -477,7 +477,7 @@ impl Message for SubscribeUpdate { // REQUEST_UPDATE let request_id = RequestId::decode(r, version)?; if matches!(version, Version::Draft17) { - let _required_request_id_delta = u64::decode(r, version)?; + let _required_request_id_delta = r.varint()?.into_inner(); } decode_params!(r, version, 0x02 => _object_delivery_timeout: Option, @@ -513,17 +513,17 @@ impl Message for SubscribeUpdate { #[cfg(test)] mod tests { use super::*; - use bytes::BytesMut; fn encode_message(msg: &M, version: Version) -> Vec { - let mut buf = BytesMut::new(); - msg.encode_msg(&mut buf, version).unwrap(); + let mut buf = Vec::new(); + msg.encode_msg(&mut Encoder::new(&mut buf, version.into()), version) + .unwrap(); buf.to_vec() } fn decode_message(bytes: &[u8], version: Version) -> Result { let mut buf = bytes::Bytes::from(bytes.to_vec()); - M::decode_msg(&mut buf, version) + crate::coding::decode_buf(&mut buf, version, M::decode_msg) } #[test] @@ -573,17 +573,18 @@ mod tests { /// Build a SUBSCRIBE body carrying a single RENDEZVOUS_TIMEOUT parameter. fn subscribe_with_rendezvous(millis: u64, version: Version) -> Vec { fn build(millis: u64, version: Version) -> Result, EncodeError> { - let mut buf = BytesMut::new(); - RequestId(1).encode(&mut buf, version)?; + let mut buf = Vec::new(); + let w = &mut Encoder::new(&mut buf, version.into()); + RequestId(1).encode(w, version)?; if version == Version::Draft17 { - 0u64.encode(&mut buf, version)?; // required_request_id_delta + w.varint(VarInt::ZERO)?; // required_request_id_delta } - encode_namespace(&mut buf, &Path::new("test"), version)?; - Cow::Borrowed("video").encode(&mut buf, version)?; - encode_params!(&mut buf, version, + encode_namespace(w, &Path::new("test"))?; + w.string("video")?; + encode_params!(w, version, 0x04 => millis, ); - Ok(buf.to_vec()) + Ok(buf) } build(millis, version).unwrap() @@ -624,17 +625,17 @@ mod tests { /// parameter values are typed per key with no generic skip rule, which is exactly why the /// draft makes an unknown one a protocol violation. fn first_param_key(encoded: &[u8], version: Version) -> Option { - let mut buf = bytes::Bytes::copy_from_slice(encoded); - RequestId::decode(&mut buf, version).unwrap(); + let r = &mut Decoder::new(encoded, version.into()); + RequestId::decode(r, version).unwrap(); if version == Version::Draft17 { - u64::decode(&mut buf, version).unwrap(); + r.varint().unwrap(); } - decode_namespace(&mut buf, version).unwrap(); - Cow::::decode(&mut buf, version).unwrap(); + decode_namespace(r).unwrap(); + r.string().unwrap(); // draft-14/15 write absolute keys, draft-16+ deltas, but the first is absolute either way. - let count = u64::decode(&mut buf, version).unwrap(); - (count > 0).then(|| u64::decode(&mut buf, version).unwrap()) + let count = r.varint().unwrap().into_inner(); + (count > 0).then(|| r.varint().unwrap().into_inner()) } /// We never ask a peer to hold a subscription open, so the parameter stays off our wire. @@ -827,16 +828,17 @@ mod tests { fn subscribe_accepts_the_delivery_timeouts() -> Result<(), EncodeError> { for version in [Version::Draft19, Version::Draft20] { let mut body = Vec::new(); - RequestId(1).encode(&mut body, version)?; - encode_namespace(&mut body, &crate::Path::new("broadcast"), version)?; - "video".encode(&mut body, version)?; - encode_params!(&mut body, version, + let w = &mut Encoder::new(&mut body, version.into()); + RequestId(1).encode(w, version)?; + encode_namespace(w, &crate::Path::new("broadcast"))?; + w.string("video")?; + encode_params!(w, version, 0x02 => 5000u64, 0x06 => 9000u64, ); - let mut buf = bytes::Bytes::from(body); - Subscribe::decode_msg(&mut buf, version).unwrap_or_else(|e| panic!("{version}: {e}")); + Subscribe::decode_msg(&mut Decoder::new(&body, version.into()), version) + .unwrap_or_else(|e| panic!("{version}: {e}")); } Ok(()) } @@ -1183,12 +1185,12 @@ mod tests { ]; // Go through the size-prefixed path: that's what rejects unread trailing bytes. - let mut buf = BytesMut::new(); - (body.len() as u16).encode(&mut buf, Version::Draft16).unwrap(); + let mut buf = Vec::new(); + Encoder::new(&mut buf, Version::Draft16.into()).u16(body.len() as u16); buf.extend_from_slice(&body); - let mut bytes = buf.freeze(); - let decoded = SubscribeOk::decode(&mut bytes, Version::Draft16).unwrap(); + let mut bytes = bytes::Bytes::from(buf); + let decoded = crate::coding::decode_buf(&mut bytes, Version::Draft16, SubscribeOk::decode).unwrap(); assert_eq!(decoded.track_alias, 42); assert_eq!(decoded.properties.group_order, Some(GroupOrder::Descending)); } @@ -1278,8 +1280,10 @@ mod cache_duration_tests { Version::Draft16 => vec![0, 0, 0, 4], _ => vec![0, 0, 4], }; - age.encode(&mut payload, version).unwrap(); - let got = SubscribeOk::decode_msg(&mut payload.as_slice(), version).unwrap(); + Encoder::new(&mut payload, version.into()) + .varint(VarInt::from(age)) + .unwrap(); + let got = crate::coding::decode_buf(&mut payload.as_slice(), version, SubscribeOk::decode_msg).unwrap(); assert_eq!( got.properties.max_cache_duration, Some(Duration::from_millis(age)), @@ -1314,8 +1318,9 @@ mod cache_duration_tests { }, }; let mut payload = Vec::new(); - ok.encode_msg(&mut payload, version).unwrap(); - let got = SubscribeOk::decode_msg(&mut payload.as_slice(), version).unwrap(); + ok.encode_msg(&mut Encoder::new(&mut payload, version.into()), version) + .unwrap(); + let got = crate::coding::decode_buf(&mut payload.as_slice(), version, SubscribeOk::decode_msg).unwrap(); assert_eq!( got.properties.max_cache_duration, if legacy { None } else { age }, diff --git a/rs/moq-net/src/ietf/subscribe_namespace.rs b/rs/moq-net/src/ietf/subscribe_namespace.rs index bf57ee0bb7..a93ad85b6b 100644 --- a/rs/moq-net/src/ietf/subscribe_namespace.rs +++ b/rs/moq-net/src/ietf/subscribe_namespace.rs @@ -60,22 +60,22 @@ fn hidden_from_param(value: Option) -> Result { impl Message for SubscribeNamespace<'_> { const ID: u64 = 0x50; - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { if is_legacy_version(version) { return Err(EncodeError::Version); } self.request_id.encode(w, version)?; - encode_namespace(w, &self.namespace, version)?; + encode_namespace(w, &self.namespace)?; encode_params!(w, version, HIDDEN_PARAM => hidden_param(self.hidden)); Ok(()) } - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { if is_legacy_version(version) { return Err(DecodeError::Version); } let request_id = RequestId::decode(r, version)?; - let namespace = decode_namespace(r, version)?; + let namespace = decode_namespace(r)?; decode_params!(r, version, HIDDEN_PARAM => hidden: Option); Ok(Self { @@ -105,33 +105,33 @@ pub struct SubscribeNamespaceLegacy<'a> { impl Message for SubscribeNamespaceLegacy<'_> { const ID: u64 = 0x11; - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { if !is_legacy_version(version) { return Err(EncodeError::Version); } self.request_id.encode(w, version)?; if version == Version::Draft17 { - 0u64.encode(w, version)?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) + w.varint(VarInt::from(0u64))?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) } - encode_namespace(w, &self.namespace, version)?; + encode_namespace(w, &self.namespace)?; if matches!(version, Version::Draft16 | Version::Draft17) { - self.subscribe_options.encode(w, version)?; + w.varint(VarInt::from(self.subscribe_options))?; } encode_params!(w, version, HIDDEN_PARAM => hidden_param(self.hidden)); Ok(()) } - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { if !is_legacy_version(version) { return Err(DecodeError::Version); } let request_id = RequestId::decode(r, version)?; if version == Version::Draft17 { - let _required_request_id_delta = u64::decode(r, version)?; + let _required_request_id_delta = r.varint()?.into_inner(); } - let namespace = decode_namespace(r, version)?; + let namespace = decode_namespace(r)?; let subscribe_options = match version { - Version::Draft16 | Version::Draft17 => u64::decode(r, version)?, + Version::Draft16 | Version::Draft17 => r.varint()?.into_inner(), _ => 0x01, }; @@ -155,12 +155,12 @@ pub struct SubscribeNamespaceOk { impl Message for SubscribeNamespaceOk { const ID: u64 = 0x12; - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.request_id.encode(w, version)?; Ok(()) } - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { let request_id = RequestId::decode(r, version)?; Ok(Self { request_id }) } @@ -177,17 +177,17 @@ pub struct SubscribeNamespaceError<'a> { impl Message for SubscribeNamespaceError<'_> { const ID: u64 = 0x13; - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.request_id.encode(w, version)?; - self.error_code.encode(w, version)?; - self.reason_phrase.encode(w, version)?; + w.varint(VarInt::from(self.error_code))?; + w.string(&self.reason_phrase)?; Ok(()) } - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { let request_id = RequestId::decode(r, version)?; - let error_code = u64::decode(r, version)?; - let reason_phrase = Cow::::decode(r, version)?; + let error_code = r.varint()?.into_inner(); + let reason_phrase = Cow::Owned(r.string()?); Ok(Self { request_id, @@ -206,12 +206,12 @@ pub struct UnsubscribeNamespace { impl Message for UnsubscribeNamespace { const ID: u64 = 0x14; - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.request_id.encode(w, version)?; Ok(()) } - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { let request_id = RequestId::decode(r, version)?; Ok(Self { request_id }) } @@ -235,8 +235,8 @@ pub struct Namespace<'a> { impl Namespace<'_> { /// Decode the message body, expecting the extended form when the session negotiated /// the MoQ Cluster extension. See [`super::PublishNamespace::decode_body`]. - pub fn decode_body(r: &mut R, version: Version, negotiated: bool) -> Result { - let suffix = decode_namespace(r, version)?; + pub fn decode_body(r: &mut Decoder<'_>, version: Version, negotiated: bool) -> Result { + let suffix = decode_namespace(r)?; // The base form has no Parameters field at all, so there is nothing to read // (and nothing to reject) unless the extension is on. @@ -252,15 +252,15 @@ impl Namespace<'_> { impl Message for Namespace<'_> { const ID: u64 = 0x08; - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { - encode_namespace(w, &self.suffix, version)?; + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { + encode_namespace(w, &self.suffix)?; match &self.cluster { Some(advert) => super::publish_namespace::encode_cluster_params(w, version, Some(advert)), None => Ok(()), } } - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { Self::decode_body(r, version, false) } } @@ -277,19 +277,19 @@ pub struct PublishBlocked<'a> { impl Message for PublishBlocked<'_> { const ID: u64 = 0x0F; - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { assert!(version == Version::Draft17, "PublishBlocked is draft17 only"); - encode_namespace(w, &self.suffix, version)?; - self.track_name.encode(w, version)?; + encode_namespace(w, &self.suffix)?; + w.string(&self.track_name)?; Ok(()) } - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { if version != Version::Draft17 { return Err(DecodeError::Unsupported); } - let suffix = decode_namespace(r, version)?; - let track_name = Cow::::decode(r, version)?; + let suffix = decode_namespace(r)?; + let track_name = Cow::Owned(r.string()?); Ok(Self { suffix, track_name }) } } @@ -304,13 +304,13 @@ pub struct NamespaceDone<'a> { impl Message for NamespaceDone<'_> { const ID: u64 = 0x0E; - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { - encode_namespace(w, &self.suffix, version)?; + fn encode_msg(&self, w: &mut Encoder<'_>, _: Version) -> Result<(), EncodeError> { + encode_namespace(w, &self.suffix)?; Ok(()) } - fn decode_msg(r: &mut R, version: Version) -> Result { - let suffix = decode_namespace(r, version)?; + fn decode_msg(r: &mut Decoder<'_>, _: Version) -> Result { + let suffix = decode_namespace(r)?; Ok(Self { suffix }) } } @@ -318,11 +318,11 @@ impl Message for NamespaceDone<'_> { #[cfg(test)] mod tests { use super::*; - use bytes::BytesMut; fn body(msg: &M, version: Version) -> Vec { - let mut buf = BytesMut::new(); - msg.encode_msg(&mut buf, version).unwrap(); + let mut buf = Vec::new(); + msg.encode_msg(&mut Encoder::new(&mut buf, version.into()), version) + .unwrap(); buf.to_vec() } @@ -346,7 +346,7 @@ mod tests { let encoded = body(&msg, version); let mut buf = bytes::Bytes::from(encoded.clone()); - let decoded = Namespace::decode_body(&mut buf, version, true).unwrap(); + let decoded = crate::coding::decode_buf(&mut buf, version, |r, v| Namespace::decode_body(r, v, true)).unwrap(); assert!(buf.is_empty()); assert_eq!(decoded.suffix.as_str(), "alice.hang"); assert_eq!(decoded.cluster, msg.cluster); @@ -361,7 +361,7 @@ mod tests { assert!(base_encoded.len() < encoded.len()); let mut buf = bytes::Bytes::from(encoded); - let decoded = Namespace::decode_body(&mut buf, version, false).unwrap(); + let decoded = crate::coding::decode_buf(&mut buf, version, |r, v| Namespace::decode_body(r, v, false)).unwrap(); assert!(decoded.cluster.is_none()); assert!(!buf.is_empty(), "the parameters were not consumed"); } @@ -387,7 +387,7 @@ mod tests { assert!(body(&free, version).len() < body(&priced, version).len()); let mut buf = bytes::Bytes::from(body(&free, version)); - let decoded = Namespace::decode_body(&mut buf, version, true).unwrap(); + let decoded = crate::coding::decode_buf(&mut buf, version, |r, v| Namespace::decode_body(r, v, true)).unwrap(); assert_eq!(decoded.cluster.unwrap().cost, 0); } @@ -395,14 +395,16 @@ mod tests { #[test] fn cluster_namespace_requires_hop_path() { let version = Version::Draft19; - let mut buf = BytesMut::new(); - encode_namespace(&mut buf, &Path::new("a"), version).unwrap(); + let mut buf = Vec::new(); + encode_namespace(&mut Encoder::new(&mut buf, version.into()), &Path::new("a")).unwrap(); // Number of Parameters = 0. - 0u64.encode(&mut buf, version).unwrap(); + Encoder::new(&mut buf, version.into()) + .varint(VarInt::from(0u64)) + .unwrap(); - let mut bytes = buf.freeze(); + let mut bytes = bytes::Bytes::from(buf); assert!(matches!( - Namespace::decode_body(&mut bytes, version, true), + crate::coding::decode_buf(&mut bytes, version, |r, v| Namespace::decode_body(r, v, true)), Err(DecodeError::InvalidValue) )); } @@ -421,10 +423,16 @@ mod tests { }), }; - let mut buf = BytesMut::new(); - msg.encode_msg(&mut buf, version).unwrap(); - let mut bytes = buf.freeze(); - assert!(super::super::PublishNamespace::decode_body(&mut bytes, version, false).is_err()); + let mut buf = Vec::new(); + msg.encode_msg(&mut Encoder::new(&mut buf, version.into()), version) + .unwrap(); + let mut bytes = bytes::Bytes::from(buf); + assert!( + crate::coding::decode_buf(&mut bytes, version, |r, v| super::super::PublishNamespace::decode_body( + r, v, false + )) + .is_err() + ); } #[test] @@ -463,7 +471,7 @@ mod tests { hidden: false, }; let mut buf = bytes::Bytes::from(body(&msg, Version::Draft18)); - let decoded = SubscribeNamespace::decode_msg(&mut buf, Version::Draft18).unwrap(); + let decoded = crate::coding::decode_buf(&mut buf, Version::Draft18, SubscribeNamespace::decode_msg).unwrap(); assert!(buf.is_empty()); assert_eq!(decoded.request_id, RequestId(4)); assert_eq!(decoded.namespace.as_str(), "example/meeting"); @@ -480,7 +488,7 @@ mod tests { }; let mut buf = bytes::Bytes::from(body(&msg, Version::Draft18)); assert_eq!( - SubscribeNamespace::decode_msg(&mut buf, Version::Draft18) + crate::coding::decode_buf(&mut buf, Version::Draft18, SubscribeNamespace::decode_msg) .unwrap() .hidden, hidden @@ -494,7 +502,8 @@ mod tests { hidden, }; let mut buf = bytes::Bytes::from(body(&msg, version)); - let decoded = SubscribeNamespaceLegacy::decode_msg(&mut buf, version).unwrap(); + let decoded = + crate::coding::decode_buf(&mut buf, version, SubscribeNamespaceLegacy::decode_msg).unwrap(); assert_eq!(decoded.hidden, hidden, "{version:?}"); } } @@ -510,7 +519,7 @@ mod tests { hidden: false, }; let mut buf = bytes::Bytes::from(body(&msg, version)); - let decoded = SubscribeNamespaceLegacy::decode_msg(&mut buf, version).unwrap(); + let decoded = crate::coding::decode_buf(&mut buf, version, SubscribeNamespaceLegacy::decode_msg).unwrap(); assert!(buf.is_empty(), "trailing bytes for {version:?}"); assert_eq!(decoded.request_id, RequestId(4)); assert_eq!(decoded.namespace.as_str(), "example/meeting"); @@ -524,7 +533,7 @@ mod tests { for version in [Version::Draft14, Version::Draft16, Version::Draft17] { let mut buf = bytes::Bytes::from(vec![0x00, 0x00, 0x00]); assert!(matches!( - SubscribeNamespace::decode_msg(&mut buf, version), + crate::coding::decode_buf(&mut buf, version, SubscribeNamespace::decode_msg), Err(DecodeError::Version) )); } @@ -532,7 +541,7 @@ mod tests { // The legacy 0x11 message only exists in draft-14..17. let mut buf = bytes::Bytes::from(vec![0x00, 0x00, 0x00]); assert!(matches!( - SubscribeNamespaceLegacy::decode_msg(&mut buf, Version::Draft18), + crate::coding::decode_buf(&mut buf, Version::Draft18, SubscribeNamespaceLegacy::decode_msg), Err(DecodeError::Version) )); } diff --git a/rs/moq-net/src/ietf/subscriber.rs b/rs/moq-net/src/ietf/subscriber.rs index 876e25770b..5584423b41 100644 --- a/rs/moq-net/src/ietf/subscriber.rs +++ b/rs/moq-net/src/ietf/subscriber.rs @@ -7,7 +7,7 @@ use std::{ use crate::{ Error, Path, PathOwned, SessionError, Timescale, broadcast, - coding::{Decode, DecodeError, Reader, Stream}, + coding::{Decode, DecodeError, Decoder, Reader, Stream, VarInt}, frame, group, ietf::{self, Control, FetchType, Filter, GroupOrder, RequestId}, origin, track, @@ -745,7 +745,10 @@ where subscribe_options: 0x01, // NAMESPACE only hidden, }; - stream.writer.encode(&ietf::SubscribeNamespaceLegacy::ID).await?; + stream + .writer + .encode(&VarInt::from(ietf::SubscribeNamespaceLegacy::ID)) + .await?; stream.writer.encode(&msg).await?; } _ => { @@ -754,7 +757,10 @@ where namespace: prefix.clone(), hidden, }; - stream.writer.encode(&ietf::SubscribeNamespace::ID).await?; + stream + .writer + .encode(&VarInt::from(ietf::SubscribeNamespace::ID)) + .await?; stream.writer.encode(&msg).await?; } } @@ -762,9 +768,9 @@ where tracing::debug!(%prefix, "subscribe_namespace sent"); // Read response - let type_id: u64 = stream.reader.decode().await?; - let size: u16 = stream.reader.decode().await?; - let mut data = stream.reader.read_exact(size as usize).await?; + let type_id = stream.reader.decode::().await?.into_inner(); + let body: ietf::Body = stream.reader.decode().await?; + let mut data = body.decoder(self.version); let count = match type_id { ietf::SubscribeNamespaceOk::ID if self.version == Version::Draft14 => { @@ -839,7 +845,7 @@ where ) -> Result<(), Error> { loop { let next = { - let mut decode = std::pin::pin!(stream.reader.decode_maybe::()); + let mut decode = std::pin::pin!(stream.reader.decode_maybe::()); kio::wait(|waiter| { // Land before decoding past the boundary, so no live update enters the // origin ahead of the marker. @@ -852,15 +858,15 @@ where }) .await }; - let type_id: u64 = match next? { - Some(id) => id, + let type_id = match next? { + Some(id) => id.into_inner(), None => break, // Stream closed }; if let Some((_, Landing::Quiet(quiet))) = landing { quiet.heard(); } - let size: u16 = stream.reader.decode().await?; - let mut data = stream.reader.read_exact(size as usize).await?; + let body: ietf::Body = stream.reader.decode().await?; + let mut data = body.decoder(self.version); match type_id { // The suffix is relative to the prefix we subscribed, which is itself @@ -939,12 +945,13 @@ where pub fn handle_stream( &mut self, id: u64, - mut data: bytes::Bytes, + body: ietf::Body, stream: Stream, peer: cluster::Peer, declared: Option, ) -> Result, Error> { let mut this = self.clone(); + let mut data = body.decoder(this.version); let task = match id { ietf::Publish::ID => { let msg = ietf::Publish::decode_msg(&mut data, this.version)?; @@ -1114,7 +1121,7 @@ where attached: &mut bool, ) -> Result<(), Error> { loop { - let type_id: u64 = match stream.reader.decode_maybe().await? { + let type_id = match stream.reader.decode_maybe::().await?.map(VarInt::into_inner) { Some(id) => id, None => return Ok(()), }; @@ -1126,8 +1133,8 @@ where return Err(Error::UnexpectedMessage); } - let size: u16 = stream.reader.decode().await?; - let mut data = stream.reader.read_exact(size as usize).await?; + let body: ietf::Body = stream.reader.decode().await?; + let mut data = body.decoder(self.version); if terminal { ietf::PublishNamespaceDone::decode_msg(&mut data, self.version)?; @@ -1259,11 +1266,14 @@ where async fn write_ok(&self, stream: &mut Stream, request_id: RequestId) -> Result<(), Error> { match self.version { Version::Draft14 => { - stream.writer.encode(&ietf::PublishNamespaceOk::ID).await?; + stream + .writer + .encode(&VarInt::from(ietf::PublishNamespaceOk::ID)) + .await?; stream.writer.encode(&ietf::PublishNamespaceOk { request_id }).await?; } Version::Draft15 | Version::Draft16 => { - stream.writer.encode(&ietf::RequestOk::ID).await?; + stream.writer.encode(&VarInt::from(ietf::RequestOk::ID)).await?; stream .writer .encode(&ietf::RequestOk { @@ -1273,7 +1283,7 @@ where .await?; } _ => { - stream.writer.encode(&ietf::RequestOk::ID).await?; + stream.writer.encode(&VarInt::from(ietf::RequestOk::ID)).await?; stream .writer .encode(&ietf::RequestOk { @@ -1298,7 +1308,10 @@ where match self.version { Version::Draft14 => { - stream.writer.encode(&ietf::PublishNamespaceError::ID).await?; + stream + .writer + .encode(&VarInt::from(ietf::PublishNamespaceError::ID)) + .await?; stream .writer .encode(&ietf::PublishNamespaceError { @@ -1309,7 +1322,7 @@ where .await?; } Version::Draft15 | Version::Draft16 => { - stream.writer.encode(&ietf::RequestError::ID).await?; + stream.writer.encode(&VarInt::from(ietf::RequestError::ID)).await?; stream .writer .encode(&ietf::RequestError { @@ -1321,7 +1334,7 @@ where .await?; } _ => { - stream.writer.encode(&ietf::RequestError::ID).await?; + stream.writer.encode(&VarInt::from(ietf::RequestError::ID)).await?; stream .writer .encode(&ietf::RequestError { @@ -1348,7 +1361,7 @@ where match self.version { Version::Draft14 => { - stream.writer.encode(&ietf::PublishError::ID).await?; + stream.writer.encode(&VarInt::from(ietf::PublishError::ID)).await?; stream .writer .encode(&ietf::PublishError { @@ -1359,7 +1372,7 @@ where .await?; } Version::Draft15 | Version::Draft16 => { - stream.writer.encode(&ietf::RequestError::ID).await?; + stream.writer.encode(&VarInt::from(ietf::RequestError::ID)).await?; stream .writer .encode(&ietf::RequestError { @@ -1371,7 +1384,7 @@ where .await?; } _ => { - stream.writer.encode(&ietf::RequestError::ID).await?; + stream.writer.encode(&VarInt::from(ietf::RequestError::ID)).await?; stream .writer .encode(&ietf::RequestError { @@ -1970,7 +1983,7 @@ where /// The publisher must send it before its FIN (draft-19 section 3.3.2), so a FIN /// without one is a failed request, not a clean end. async fn read_publish_done(reader: &mut Reader, version: Version) -> Result { - match reader.decode_maybe::().await? { + match reader.decode_maybe::().await?.map(VarInt::into_inner) { Some(ietf::PublishDone::ID) => {} Some(_) => return Err(Error::UnexpectedMessage), None => return Err(Error::ProtocolViolation), @@ -2029,7 +2042,7 @@ where writer: &mut crate::coding::Writer, request_id: RequestId, ) -> Result<(), Error> { - writer.encode(&ietf::Unsubscribe::ID).await?; + writer.encode(&VarInt::from(ietf::Unsubscribe::ID)).await?; writer.encode(&ietf::Unsubscribe { request_id }).await?; Ok(()) } @@ -2045,7 +2058,7 @@ where // Read the aggregate now: a subscriber can join while the request ID and stream // were awaited, and nothing updates the priority after SUBSCRIBE. let priority = request.subscription().map(|s| s.priority).unwrap_or(0); - stream.writer.encode(&ietf::Subscribe::ID).await?; + stream.writer.encode(&VarInt::from(ietf::Subscribe::ID)).await?; stream .writer .encode(&ietf::Subscribe { @@ -2115,7 +2128,7 @@ where }; if let Err(err) = async { - stream.writer.encode(&ietf::Fetch::ID).await?; + stream.writer.encode(&VarInt::from(ietf::Fetch::ID)).await?; stream .writer .encode(&ietf::Fetch { @@ -2161,9 +2174,9 @@ where /// `true` when the publisher answered FETCH_OK. A FETCH_ERROR / REQUEST_ERROR is a /// refusal, not a session error: the live subscription continues. async fn read_fetch_response(&self, stream: &mut Stream) -> Result { - let type_id: u64 = stream.reader.decode().await?; - let size: u16 = stream.reader.decode().await?; - let mut data = stream.reader.read_exact(size as usize).await?; + let type_id = stream.reader.decode::().await?.into_inner(); + let body: ietf::Body = stream.reader.decode().await?; + let mut data = body.decoder(self.version); match type_id { ietf::FetchOk::ID => { @@ -2184,9 +2197,9 @@ where async fn read_subscribe_response(&self, stream: &mut Stream) -> Result, Error> { // Read type_id + size + body from the stream - let type_id: u64 = stream.reader.decode().await?; - let size: u16 = stream.reader.decode().await?; - let mut data = stream.reader.read_exact(size as usize).await?; + let type_id = stream.reader.decode::().await?.into_inner(); + let body: ietf::Body = stream.reader.decode().await?; + let mut data = body.decoder(self.version); match type_id { ietf::SubscribeOk::ID => { @@ -2408,17 +2421,13 @@ struct FirstObject { struct PeekFirst(FirstObject); impl Decode for PeekFirst { - fn decode(buf: &mut B, version: Version) -> Result { - let id = u64::decode(buf, version)?; + fn decode(buf: &mut Decoder<'_>, _: Version) -> Result { + let id = buf.varint()?.into_inner(); if EXTENSIONS { - let size = usize::decode(buf, version)?; - if buf.remaining() < size { - return Err(DecodeError::Short); - } - buf.advance(size); + buf.bytes()?; } - let size = u64::decode(buf, version)?; - let end_of_track = size == 0 && u64::decode(buf, version)? == END_OF_TRACK; + let size = buf.varint()?.into_inner(); + let end_of_track = size == 0 && buf.varint()?.into_inner() == END_OF_TRACK; Ok(Self(FirstObject { id, end_of_track })) } } @@ -2600,7 +2609,7 @@ where /// signal, and arrives here as a read error, which drops the head and the join with it. pub async fn recv_fill(&mut self, stream: &mut Reader) -> Result<(), Error> { // The dispatcher peeked the stream type to get here. - let _: u64 = stream.decode().await?; + let _ = stream.decode::().await?.into_inner(); let header: ietf::FetchHeader = stream.decode().await?; let (subscribe_id, fill, joining, largest, _counted) = { @@ -2823,7 +2832,7 @@ where // out, so its frames are stamped on arrival instead. let timestamp = match (object.properties, timescale) { (Some(properties), Some(timescale)) => { - let mut properties = bytes::Bytes::from(properties); + let mut properties = Decoder::new(&properties, self.version.into()); ietf::decode_object_time(&mut properties, timescale, self.version)? } _ => None, @@ -2832,9 +2841,9 @@ where // A fetch object has no status field from draft-16 on; a zero length is simply // an empty object. Draft-14 and 15 still encode Normal (0) after a zero length. - let size: u64 = stream.decode().await?; + let size = stream.decode::().await?.into_inner(); if size == 0 && matches!(self.version, Version::Draft14 | Version::Draft15) { - let status: u64 = stream.decode().await?; + let status = stream.decode::().await?.into_inner(); if status != 0 { return Err(Error::Unsupported); } @@ -2872,13 +2881,14 @@ async fn decode_fetch_object( version: Version, ) -> Result, Error> { if version == Version::Draft14 { - let Some(group) = stream.decode_maybe::().await? else { + let Some(group) = stream.decode_maybe::().await?.map(VarInt::into_inner) else { return Ok(None); }; - let subgroup: u64 = stream.decode().await?; - let object: u64 = stream.decode().await?; - let _priority: u8 = stream.decode().await?; - let properties: Vec = stream.decode().await?; + let subgroup = stream.decode::().await?.into_inner(); + let object = stream.decode::().await?.into_inner(); + let _priority = stream.read_exact(1).await?; + let size = usize::try_from(stream.decode::().await?)?; + let properties = stream.read_exact(size).await?.to_vec(); return Ok(Some(FetchedObject { group: Some(group), object: Some(object), @@ -2995,7 +3005,8 @@ impl GroupIngest { loop { match &mut self.phase { IngestPhase::Delta => { - let Some(id_delta) = ready!(reader.poll_decode_maybe::(&mut cx))? else { + let Some(id_delta) = ready!(reader.poll_decode_maybe::(&mut cx))?.map(VarInt::into_inner) + else { return Poll::Ready(Ok(Ended::Group)); }; self.prior_object = Some(next_object_id(self.prior_object, id_delta, self.start)?); @@ -3005,7 +3016,9 @@ impl GroupIngest { }; } IngestPhase::ExtSize => { - let size: usize = ready!(reader.poll_decode(&mut cx))?; + let size = ready!(reader.poll_decode::(&mut cx))?; + let size = + usize::try_from(size).map_err(|_| Error::BoundsExceeded(crate::coding::BoundsExceeded))?; self.phase = IngestPhase::ExtBytes { size }; } IngestPhase::ExtBytes { size } => { @@ -3014,15 +3027,18 @@ impl GroupIngest { // declared. A track that declared no timescale opted out, so its // objects are stamped on arrival even if one carries a Timestamp we // could not interpret. - let mut ext = ready!(reader.poll_read_exact(&mut cx, *size))?; + let ext = ready!(reader.poll_read_exact(&mut cx, *size))?; let timestamp = match self.timescale { - Some(timescale) => ietf::decode_object_time(&mut ext, timescale, self.version)?, + Some(timescale) => { + let mut ext = Decoder::new(&ext, self.version.into()); + ietf::decode_object_time(&mut ext, timescale, self.version)? + } None => None, }; self.phase = IngestPhase::Size { timestamp }; } IngestPhase::Size { timestamp } => { - let size: u64 = ready!(reader.poll_decode(&mut cx))?; + let size = ready!(reader.poll_decode::(&mut cx))?.into_inner(); if size == 0 { self.phase = IngestPhase::Status { timestamp: *timestamp }; continue; @@ -3034,7 +3050,7 @@ impl GroupIngest { self.phase = IngestPhase::Payload { frame }; } IngestPhase::Status { timestamp } => { - let status: u64 = ready!(reader.poll_decode(&mut cx))?; + let status = ready!(reader.poll_decode::(&mut cx))?.into_inner(); if status == 0 { let timestamp = timestamp.unwrap_or_else(|| crate::Timestamp::from(self.runtime.now())); let frame = group.create_frame_owned(frame::Info { size: 0, timestamp })?; @@ -3100,14 +3116,19 @@ mod tests { use crate::coding::Encode; let mut responses = Vec::new(); if clean { - ietf::PublishDone::ID.encode(&mut responses, Version::Draft19).unwrap(); + crate::coding::Encoder::new(&mut responses, Version::Draft19.into()) + .varint(VarInt::from(ietf::PublishDone::ID)) + .unwrap(); ietf::PublishDone { request_id: None, status_code: ietf::PublishDoneStatus::TrackEnded.code(Version::Draft19), stream_count: 0, reason_phrase: "done".into(), } - .encode(&mut responses, Version::Draft19) + .encode( + &mut crate::coding::Encoder::new(&mut responses, Version::Draft19.into()), + Version::Draft19, + ) .unwrap(); } responses @@ -3325,7 +3346,7 @@ mod tests { let log = crate::lite::test_transport::Log::default(); let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(log.clone()), version); - writer.encode(&ietf::RequestOk::ID).await.unwrap(); + writer.encode(&VarInt::from(ietf::RequestOk::ID)).await.unwrap(); writer .encode(&ietf::RequestOk { request_id: None, @@ -3333,7 +3354,7 @@ mod tests { }) .await .unwrap(); - writer.encode(&ietf::Namespace::ID).await.unwrap(); + writer.encode(&VarInt::from(ietf::Namespace::ID)).await.unwrap(); writer .encode(&ietf::Namespace { suffix: crate::Path::new(suffix), @@ -3417,7 +3438,7 @@ mod tests { let log = crate::lite::test_transport::Log::default(); let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(log.clone()), VERSION); - writer.encode(&ietf::RequestOk::ID).await.unwrap(); + writer.encode(&VarInt::from(ietf::RequestOk::ID)).await.unwrap(); writer .encode(&ietf::RequestOk { request_id: None, @@ -3685,7 +3706,7 @@ mod tests { let log = crate::lite::test_transport::Log::default(); let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(log.clone()), VERSION); - writer.encode(&ietf::RequestError::ID).await.unwrap(); + writer.encode(&VarInt::from(ietf::RequestError::ID)).await.unwrap(); writer .encode(&ietf::RequestError { request_id: Some(RequestId(1)), @@ -3879,23 +3900,20 @@ mod tests { /// Decoding the framing rather than scanning for a byte: a type id is one varint among /// many, and a substring match would happily find one inside a length or a payload. fn control_message_types(log: &crate::lite::test_transport::Log, version: Version) -> Vec { - use crate::coding::Decode; - let writes = log.writes.lock().unwrap().clone(); - let mut buf = writes.as_slice(); + let mut buf = Decoder::new(&writes, version.into()); let mut types = Vec::new(); while !buf.is_empty() { - let Ok(type_id) = u64::decode(&mut buf, version) else { + let Ok(type_id) = buf.varint().map(VarInt::into_inner) else { break; }; - let Ok(size) = u16::decode(&mut buf, version) else { + let Ok(size) = buf.u16() else { break; }; - if buf.len() < size as usize { + if buf.slice(size as usize).is_err() { break; } - buf = &buf[size as usize..]; types.push(type_id); } @@ -4024,7 +4042,7 @@ mod tests { let log = crate::lite::test_transport::Log::default(); let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(log.clone()), version); - writer.encode(&ietf::SubscribeOk::ID).await.unwrap(); + writer.encode(&VarInt::from(ietf::SubscribeOk::ID)).await.unwrap(); writer .encode(&ietf::SubscribeOk { request_id: match version { @@ -4127,7 +4145,7 @@ mod tests { let log = crate::lite::test_transport::Log::default(); let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(log.clone()), version); - writer.encode(&ietf::SubscribeOk::ID).await.unwrap(); + writer.encode(&VarInt::from(ietf::SubscribeOk::ID)).await.unwrap(); writer .encode(&ietf::SubscribeOk { request_id: None, @@ -4570,7 +4588,7 @@ mod tests { let log = crate::lite::test_transport::Log::default(); let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(log.clone()), VERSION); - writer.encode(&ietf::RequestOk::ID).await.unwrap(); + writer.encode(&VarInt::from(ietf::RequestOk::ID)).await.unwrap(); writer .encode(&ietf::RequestOk { request_id: None, @@ -4579,7 +4597,7 @@ mod tests { .await .unwrap(); for cost in [4, 0] { - writer.encode(&ietf::Namespace::ID).await.unwrap(); + writer.encode(&VarInt::from(ietf::Namespace::ID)).await.unwrap(); writer .encode(&ietf::Namespace { suffix: crate::Path::new("x.hang"), @@ -4731,7 +4749,7 @@ mod tests { // error while the advertisement is still live. let log = crate::lite::test_transport::Log::default(); let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(log.clone()), VERSION); - writer.encode(&ietf::NamespaceDone::ID).await.unwrap(); + writer.encode(&VarInt::from(ietf::NamespaceDone::ID)).await.unwrap(); let script = log.writes.lock().unwrap().clone(); let session = crate::lite::test_transport::ScriptedSession::eof(script); @@ -5031,7 +5049,10 @@ mod tests { let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(log.clone()), VERSION); for (i, advert) in updates.iter().enumerate() { - writer.encode(&ietf::PublishNamespaceUpdate::ID).await.unwrap(); + writer + .encode(&VarInt::from(ietf::PublishNamespaceUpdate::ID)) + .await + .unwrap(); writer .encode(&ietf::PublishNamespaceUpdate { // Each update consumes a request id of the peer's parity. @@ -5261,7 +5282,10 @@ mod tests { let log = crate::lite::test_transport::Log::default(); let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(log.clone()), VERSION); - writer.encode(&ietf::PublishNamespaceUpdate::ID).await.unwrap(); + writer + .encode(&VarInt::from(ietf::PublishNamespaceUpdate::ID)) + .await + .unwrap(); writer .encode(&ietf::PublishNamespaceUpdate { request_id: RequestId(3), @@ -5415,7 +5439,7 @@ mod tests { let log = crate::lite::test_transport::Log::default(); let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(log.clone()), VERSION); - writer.encode(&ietf::PublishNamespace::ID).await.unwrap(); + writer.encode(&VarInt::from(ietf::PublishNamespace::ID)).await.unwrap(); writer .encode(&ietf::PublishNamespace { request_id: RequestId(1), @@ -5547,7 +5571,7 @@ mod tests { match version { Version::Draft14 => { - writer.encode(&ietf::PublishError::ID).await.unwrap(); + writer.encode(&VarInt::from(ietf::PublishError::ID)).await.unwrap(); writer .encode(&ietf::PublishError { request_id: RequestId(1), @@ -5558,7 +5582,7 @@ mod tests { .unwrap(); } _ => { - writer.encode(&ietf::RequestError::ID).await.unwrap(); + writer.encode(&VarInt::from(ietf::RequestError::ID)).await.unwrap(); writer .encode(&ietf::RequestError { request_id: None, @@ -5897,7 +5921,7 @@ mod stitch_tests { use super::*; use crate::{ Timestamp, - coding::Encode as _, + coding::{Encode as _, Encoder}, lite::test_transport::ScriptedSession, model::ProduceTest, transport::poll::Session as _, @@ -5926,9 +5950,13 @@ mod stitch_tests { /// stamped on arrival share one epoch with the live tail; a 1000µs presentation time /// against a wall-clock tail would convict every earlier group as stale. fn fill_stream_for>(request_id: RequestId, groups: &[(u64, &[B])], timed: bool) -> Vec { - let mut buf = bytes::BytesMut::new(); - ietf::FetchHeader::TYPE.encode(&mut buf, VERSION).unwrap(); - ietf::FetchHeader { request_id }.encode(&mut buf, VERSION).unwrap(); + let mut buf = Vec::new(); + crate::coding::Encoder::new(&mut buf, VERSION.into()) + .varint(VarInt::from(ietf::FetchHeader::TYPE)) + .unwrap(); + ietf::FetchHeader { request_id } + .encode(&mut crate::coding::Encoder::new(&mut buf, VERSION.into()), VERSION) + .unwrap(); let mut object_index = 0usize; let mut prev_group = None; @@ -5936,10 +5964,10 @@ mod stitch_tests { for (index, payload) in payloads.iter().enumerate() { let payload = payload.as_ref(); let properties = timed.then(|| { - let mut properties = bytes::BytesMut::new(); - ietf::encode_object_time(&mut properties, timestamp(object_index), Timescale::MICRO, VERSION) - .unwrap(); - properties.to_vec() + let mut properties = Vec::new(); + let w = &mut Encoder::new(&mut properties, VERSION.into()); + ietf::encode_object_time(w, timestamp(object_index), Timescale::MICRO, VERSION).unwrap(); + properties }); // The first object of the stream carries the absolute Group ID. From @@ -5958,10 +5986,12 @@ mod stitch_tests { priority: first.then_some(0), properties, } - .encode(&mut buf, VERSION) + .encode(&mut crate::coding::Encoder::new(&mut buf, VERSION.into()), VERSION) .unwrap(); - (payload.len() as u64).encode(&mut buf, VERSION).unwrap(); + crate::coding::Encoder::new(&mut buf, VERSION.into()) + .varint(VarInt::from(payload.len())) + .unwrap(); buf.put_slice(payload); object_index += 1; } @@ -5974,7 +6004,7 @@ mod stitch_tests { /// The subscription's own subgroup stream, starting at `start` because a strict /// publisher delivers nothing before it: that head is the fill's job. fn tail_stream(sequence: u64, start: u64, payloads: &[&[u8]]) -> Vec { - let mut buf = bytes::BytesMut::new(); + let mut buf = Vec::new(); ietf::GroupHeader { track_alias: ALIAS, group_id: sequence, @@ -5985,7 +6015,7 @@ mod stitch_tests { ..Default::default() }, } - .encode(&mut buf, VERSION) + .encode(&mut crate::coding::Encoder::new(&mut buf, VERSION.into()), VERSION) .unwrap(); for (index, payload) in payloads.iter().enumerate() { @@ -5995,8 +6025,12 @@ mod stitch_tests { 0 => start, _ => 0, }; - delta.encode(&mut buf, VERSION).unwrap(); - (payload.len() as u64).encode(&mut buf, VERSION).unwrap(); + crate::coding::Encoder::new(&mut buf, VERSION.into()) + .varint(VarInt::from(delta)) + .unwrap(); + crate::coding::Encoder::new(&mut buf, VERSION.into()) + .varint(VarInt::from(payload.len())) + .unwrap(); buf.put_slice(payload); } @@ -6146,7 +6180,9 @@ mod stitch_tests { /// Append an END_OF_TRACK object: delta 0, an empty payload, then its status. fn end_of_track(mut stream: Vec) -> Vec { for value in [0u64, 0, END_OF_TRACK] { - value.encode(&mut stream, VERSION).unwrap(); + crate::coding::Encoder::new(&mut stream, VERSION.into()) + .varint(VarInt::from(value)) + .unwrap(); } stream } @@ -6610,8 +6646,11 @@ mod joining_fetch_tests { fn message_bytes(id: u64, msg: &M, version: Version) -> Vec { let mut buf = Vec::new(); - id.encode(&mut buf, version).unwrap(); - msg.encode(&mut buf, version).unwrap(); + crate::coding::Encoder::new(&mut buf, version.into()) + .varint(VarInt::from(id)) + .unwrap(); + msg.encode(&mut crate::coding::Encoder::new(&mut buf, version.into()), version) + .unwrap(); buf } @@ -6685,21 +6724,16 @@ mod joining_fetch_tests { use crate::coding::Decode; let writes = log.writes.lock().unwrap().clone(); - let mut buf = writes.as_slice(); + let mut buf = Decoder::new(&writes, version.into()); let mut messages = Vec::new(); while !buf.is_empty() { - let Ok(type_id) = u64::decode(&mut buf, version) else { + let Ok(type_id) = buf.varint() else { break; }; - let Ok(size) = u16::decode(&mut buf, version) else { + let Ok(body) = ietf::Body::decode(&mut buf, version) else { break; }; - if buf.len() < size as usize { - break; - } - let (body, rest) = buf.split_at(size as usize); - messages.push((type_id, bytes::Bytes::copy_from_slice(body))); - buf = rest; + messages.push((type_id.into_inner(), body.0)); } messages } @@ -6793,13 +6827,13 @@ mod joining_fetch_tests { .find(|(id, _)| *id == ietf::Subscribe::ID) .expect("SUBSCRIBE"); let mut body = subscribe.1.clone(); - let msg = ietf::Subscribe::decode_msg(&mut body, version).unwrap(); + let msg = crate::coding::decode_buf(&mut body, version, ietf::Subscribe::decode_msg).unwrap(); assert_eq!(msg.filter, Filter::NextObject, "{version}"); assert!(msg.fill.is_none(), "{version}"); let fetch = messages.iter().find(|(id, _)| *id == ietf::Fetch::ID).expect("FETCH"); let mut body = fetch.1.clone(); - let msg = ietf::Fetch::decode_msg(&mut body, version).unwrap(); + let msg = crate::coding::decode_buf(&mut body, version, ietf::Fetch::decode_msg).unwrap(); assert_eq!( msg.fetch_type, FetchType::RelativeJoining { @@ -6828,7 +6862,7 @@ mod joining_fetch_tests { let messages = decode_messages(&run.session.log, version); let fetch = messages.iter().find(|(id, _)| *id == ietf::Fetch::ID).expect("FETCH"); let mut body = fetch.1.clone(); - let msg = ietf::Fetch::decode_msg(&mut body, version).unwrap(); + let msg = crate::coding::decode_buf(&mut body, version, ietf::Fetch::decode_msg).unwrap(); assert_eq!( msg.fetch_type, FetchType::AbsoluteJoining { diff --git a/rs/moq-net/src/ietf/token.rs b/rs/moq-net/src/ietf/token.rs index 129e9e17e4..2dad2a6fbb 100644 --- a/rs/moq-net/src/ietf/token.rs +++ b/rs/moq-net/src/ietf/token.rs @@ -6,7 +6,7 @@ use crate::{ Error, SessionError, - coding::{Decode, Encode, EncodeError}, + coding::{Decoder, EncodeError, Encoder}, setup::Token, }; @@ -36,39 +36,42 @@ pub fn from_setup(params: &Parameters, version: Version) -> Result #[cfg_attr(not(test), expect(dead_code))] pub fn into_setup(params: &mut Parameters, token: &Token, version: Version) -> Result<(), EncodeError> { let mut value = Vec::new(); - USE_VALUE.encode(&mut value, version)?; - token.kind.encode(&mut value, version)?; - value.extend_from_slice(&token.value); + let mut w = Encoder::new(&mut value, version.into()); + w.varint(USE_VALUE.into())?; + w.varint(token.kind.into())?; + w.slice(&token.value); params.set_bytes(ParameterBytes::AuthorizationToken, value); Ok(()) } /// Decode a Token structure, refusing what a SETUP cannot carry. -fn decode(mut buf: &[u8], version: Version) -> Result { +fn decode(buf: &[u8], version: Version) -> Result { // Section 8.9: a structure that cannot be decoded closes with KEY_VALUE_FORMATTING_ERROR. let malformed = |_| Error::Session(SessionError::KeyValueFormatting); + let mut r = Decoder::new(buf, version.into()); - match u64::decode(&mut buf, version).map_err(malformed)? { + match r.varint().map_err(malformed)?.into_inner() { USE_VALUE => {} // With no cache, section 9.1.4 treats a registration as a value; the alias is unused. REGISTER => { - u64::decode(&mut buf, version).map_err(malformed)?; + r.varint().map_err(malformed)?; } // Section 9.1.4: nothing can have been registered before SETUP. DELETE | USE_ALIAS => return Err(Error::ProtocolViolation), _ => return Err(Error::Session(SessionError::KeyValueFormatting)), } - let kind = u64::decode(&mut buf, version).map_err(malformed)?; + let kind = r.varint().map_err(malformed)?.into_inner(); Ok(Token { kind, - value: buf.to_vec(), + value: r.rest().to_vec(), }) } #[cfg(test)] mod tests { use super::*; + use crate::coding::{Decode, Encode}; const VERSIONS: [Version; 9] = [ Version::Draft14, @@ -93,15 +96,16 @@ mod tests { /// The option as it arrives, after a trip through the SETUP parameter block. fn received(params: &Parameters, version: Version) -> Parameters { - let mut bytes = params.encode_bytes(version).unwrap(); - Parameters::decode(&mut bytes, version).unwrap() + let bytes = params.encode_bytes(version).unwrap(); + Parameters::decode_slice(&bytes, version).unwrap().0 } /// A raw Token structure: the alias type then its varint fields then a value. fn structure(version: Version, fields: &[u64], value: &[u8]) -> Parameters { let mut raw = Vec::new(); + let mut w = Encoder::new(&mut raw, version.into()); for field in fields { - field.encode(&mut raw, version).unwrap(); + w.varint((*field).into()).unwrap(); } raw.extend_from_slice(value); let mut params = Parameters::default(); @@ -173,8 +177,12 @@ mod tests { fn two_tokens_are_refused() { for version in VERSIONS { let mut value = Vec::new(); - USE_VALUE.encode(&mut value, version).unwrap(); - Token::OUT_OF_BAND.encode(&mut value, version).unwrap(); + Encoder::new(&mut value, version.into()) + .varint(crate::coding::VarInt::from(USE_VALUE)) + .unwrap(); + Encoder::new(&mut value, version.into()) + .varint(crate::coding::VarInt::from(Token::OUT_OF_BAND)) + .unwrap(); let key = u64::from(ParameterBytes::AuthorizationToken); let (count, keys): (Option, [u64; 2]) = match version { @@ -185,14 +193,18 @@ mod tests { }; let mut raw = Vec::new(); if let Some(count) = count { - count.encode(&mut raw, version).unwrap(); + Encoder::new(&mut raw, version.into()) + .varint(crate::coding::VarInt::from(count)) + .unwrap(); } for key in keys { - key.encode(&mut raw, version).unwrap(); - value.encode(&mut raw, version).unwrap(); + Encoder::new(&mut raw, version.into()) + .varint(crate::coding::VarInt::from(key)) + .unwrap(); + Encoder::new(&mut raw, version.into()).bytes(&value).unwrap(); } - let err = Parameters::decode(&mut raw.as_slice(), version).unwrap_err(); + let err = Parameters::decode_slice(&raw, version).unwrap_err(); assert!(matches!(err, crate::DecodeError::Duplicate), "{version:?}: {err:?}"); } } diff --git a/rs/moq-net/src/ietf/track.rs b/rs/moq-net/src/ietf/track.rs index 2c600c331c..9aabea1f0a 100644 --- a/rs/moq-net/src/ietf/track.rs +++ b/rs/moq-net/src/ietf/track.rs @@ -28,21 +28,21 @@ pub struct TrackStatus<'a> { impl Message for TrackStatus<'_> { const ID: u64 = 0x0d; - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.request_id.encode(w, version)?; if version == Version::Draft17 { - 0u64.encode(w, version)?; // required_request_id_delta = 0 + w.varint(VarInt::from(0u64))?; // required_request_id_delta = 0 } - encode_namespace(w, &self.track_namespace, version)?; - self.track_name.encode(w, version)?; + encode_namespace(w, &self.track_namespace)?; + w.string(&self.track_name)?; match version { Version::Draft14 => { - 0u8.encode(w, version)?; // subscriber priority + w.u8(0u8); // subscriber priority GroupOrder::Descending.encode(w, version)?; - false.encode(w, version)?; // forward + w.bool(false); // forward Filter::NextObject.encode(w, version)?; // filter - 0u8.encode(w, version)?; // no parameters + w.u8(0u8); // no parameters } _ => { encode_params!(w, version,); @@ -51,20 +51,20 @@ impl Message for TrackStatus<'_> { Ok(()) } - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { let request_id = RequestId::decode(r, version)?; if version == Version::Draft17 { - let _required_request_id_delta = u64::decode(r, version)?; + let _required_request_id_delta = r.varint()?.into_inner(); } - let track_namespace = decode_namespace(r, version)?; - let track_name = Cow::::decode(r, version)?; + let track_namespace = decode_namespace(r)?; + let track_name = Cow::Owned(r.string()?); match version { Version::Draft14 => { - let _subscriber_priority = u8::decode(r, version)?; + let _subscriber_priority = r.u8()?; let _group_order = GroupOrder::decode(r, version)?; - let _forward = bool::decode(r, version)?; - let _filter_type = u64::decode(r, version)?; + let _forward = r.bool()?; + let _filter_type = r.varint()?.into_inner(); let _params = Parameters::decode(r, version)?; } _ => { @@ -90,32 +90,32 @@ pub enum TrackStatusCode { } impl Encode for TrackStatusCode { - fn encode(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { - u64::from(*self).encode(w, version)?; + fn encode(&self, w: &mut Encoder<'_>, _: Version) -> Result<(), EncodeError> { + w.varint(VarInt::from(u64::from(*self)))?; Ok(()) } } impl Decode for TrackStatusCode { - fn decode(r: &mut R, version: Version) -> Result { - Self::try_from(u64::decode(r, version)?).map_err(|_| DecodeError::InvalidValue) + fn decode(r: &mut Decoder<'_>, _: Version) -> Result { + Self::try_from(r.varint()?.into_inner()).map_err(|_| DecodeError::InvalidValue) } } #[cfg(test)] mod tests { use super::*; - use bytes::BytesMut; fn encode_message(msg: &M, version: Version) -> Vec { - let mut buf = BytesMut::new(); - msg.encode_msg(&mut buf, version).unwrap(); + let mut buf = Vec::new(); + msg.encode_msg(&mut Encoder::new(&mut buf, version.into()), version) + .unwrap(); buf.to_vec() } fn decode_message(bytes: &[u8], version: Version) -> Result { let mut buf = bytes::Bytes::from(bytes.to_vec()); - M::decode_msg(&mut buf, version) + crate::coding::decode_buf(&mut buf, version, M::decode_msg) } #[test] diff --git a/rs/moq-net/src/ietf/version.rs b/rs/moq-net/src/ietf/version.rs index 34f68362af..11d7db02e6 100644 --- a/rs/moq-net/src/ietf/version.rs +++ b/rs/moq-net/src/ietf/version.rs @@ -60,13 +60,16 @@ mod tests { fn message(msg: &M, version: Version) -> Vec { let mut buf = Vec::new(); - msg.encode_msg(&mut buf, version).expect("encode"); + msg.encode_msg(&mut crate::coding::Encoder::new(&mut buf, version.into()), version) + .expect("encode"); buf } fn field>(value: &E, version: Version) -> Vec { let mut buf = Vec::new(); - value.encode(&mut buf, version).expect("encode"); + value + .encode(&mut crate::coding::Encoder::new(&mut buf, version.into()), version) + .expect("encode"); buf } diff --git a/rs/moq-net/src/lite/announce.rs b/rs/moq-net/src/lite/announce.rs index e54733e4f7..33de5dab76 100644 --- a/rs/moq-net/src/lite/announce.rs +++ b/rs/moq-net/src/lite/announce.rs @@ -1,4 +1,3 @@ -use bytes::{Buf, BufMut}; use num_enum::{IntoPrimitive, TryFromPrimitive}; use crate::{Hop, Hops, Path, coding::*, origin::Cost}; @@ -83,10 +82,10 @@ impl<'a> PathRef<'a> { } impl Encode for PathRef<'_> { - fn encode(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { if version.has_announce_compression() { - self.base.encode(w, version)?; - self.keep.encode(w, version)?; + w.varint(VarInt::from(self.base))?; + w.varint(VarInt::from(self.keep))?; } else if self.base != 0 || self.keep != 0 { return Err(EncodeError::Version); } @@ -95,12 +94,12 @@ impl Encode for PathRef<'_> { } impl Decode for PathRef<'_> { - fn decode(buf: &mut B, version: Version) -> Result { + fn decode(buf: &mut Decoder<'_>, version: Version) -> Result { if !version.has_announce_compression() { return Ok(Self::literal(Path::decode(buf, version)?)); } - let base = u64::decode(buf, version)?; - let keep = u64::decode(buf, version)?; + let base = buf.varint()?.into_inner(); + let keep = buf.varint()?.into_inner(); if base == 0 && keep != 0 { return Err(DecodeError::InvalidValue); } @@ -133,27 +132,27 @@ impl HopsRef { } impl Encode for HopsRef { - fn encode(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { if !version.has_announce_compression() { if self.base != 0 || self.keep != 0 { return Err(EncodeError::Version); } return self.literal.encode(w, version); } - self.base.encode(w, version)?; + w.varint(VarInt::from(self.base))?; self.literal.encode(w, version)?; - self.keep.encode(w, version) + w.varint(VarInt::from(self.keep)) } } impl Decode for HopsRef { - fn decode(buf: &mut B, version: Version) -> Result { + fn decode(buf: &mut Decoder<'_>, version: Version) -> Result { if !version.has_announce_compression() { return Ok(Self::literal(Hops::decode(buf, version)?)); } - let base = u64::decode(buf, version)?; + let base = buf.varint()?.into_inner(); let literal = Hops::decode(buf, version)?; - let keep = u64::decode(buf, version)?; + let keep = buf.varint()?.into_inner(); if base == 0 && keep != 0 { return Err(DecodeError::InvalidValue); } @@ -187,65 +186,64 @@ impl AnnounceBroadcast<'_> { } impl Encode for Cost { - fn encode(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { if !version.has_route_cost() { return Ok(()); } - self.warm.encode(w, version)?; - self.cold.encode(w, version) + w.varint(VarInt::from(self.warm))?; + w.varint(VarInt::from(self.cold)) } } impl Decode for Cost { - fn decode(buf: &mut B, version: Version) -> Result { + fn decode(buf: &mut Decoder<'_>, version: Version) -> Result { if !version.has_route_cost() { return Ok(Cost::UNKNOWN); } Ok(Cost { - warm: u64::decode(buf, version)?, - cold: u64::decode(buf, version)?, + warm: buf.varint()?.into_inner(), + cold: buf.varint()?.into_inner(), }) } } impl Encode for AnnounceBroadcast<'_> { - fn encode(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { if version.has_announce_id() { // Lite06+: outer type discriminator, then a size-prefixed body (like the - // subscribe stream). The body varies by type. Announce messages are small and - // infrequent, so the scratch buffer is cheap. - let mut body = Vec::new(); + // subscribe stream). The body varies by type. let typ = match self { - Self::Active { suffix, hops, cost } => { - suffix.encode(&mut body, version)?; - hops.encode(&mut body, version)?; - cost.encode(&mut body, version)?; - ANNOUNCE_START - } - Self::EndedId { id } => { - id.encode(&mut body, version)?; - ANNOUNCE_END - } - Self::Restart { id, hops, cost } => { - id.encode(&mut body, version)?; - hops.encode(&mut body, version)?; - cost.encode(&mut body, version)?; - ANNOUNCE_RESTART - } + Self::Active { .. } => ANNOUNCE_START, + Self::EndedId { .. } => ANNOUNCE_END, + Self::Restart { .. } => ANNOUNCE_RESTART, // The pre-lite-06 path-form retraction has no place on lite-06. Self::Ended { .. } => return Err(EncodeError::Version), // Decode-only: an unknown type is never sent. Self::Skipped => return Err(EncodeError::Unsupported), }; - typ.encode(w, version)?; - (body.len() as u64).encode(w, version)?; - w.put_slice(&body); - return Ok(()); + w.varint(VarInt::from(typ))?; + + let start = w.position(); + match self { + Self::Active { suffix, hops, cost } => { + suffix.encode(w, version)?; + hops.encode(w, version)?; + cost.encode(w, version)?; + } + Self::EndedId { id } => w.varint(VarInt::from(*id))?, + Self::Restart { id, hops, cost } => { + w.varint(VarInt::from(*id))?; + hops.encode(w, version)?; + cost.encode(w, version)?; + } + Self::Ended { .. } | Self::Skipped => unreachable!("refused above"), + } + return w.prefix_varint(start); } // Older versions: a single ANNOUNCE_BROADCAST message, size-prefixed, with the // status carried inside the body. - let mut body = Vec::new(); + let start = w.position(); match self { // The cost is a lite-06 addition, so it is simply not on the wire here. Self::Active { suffix, hops, .. } => { @@ -253,36 +251,31 @@ impl Encode for AnnounceBroadcast<'_> { if suffix.base != 0 || hops.base != 0 { return Err(EncodeError::Version); } - AnnounceStatus::Active.encode(&mut body, version)?; - suffix.rest.encode(&mut body, version)?; - encode_hops(&mut body, version, &hops.literal)?; + AnnounceStatus::Active.encode(w, version)?; + suffix.rest.encode(w, version)?; + encode_hops(w, version, &hops.literal)?; } Self::Ended { suffix, hops } => { - AnnounceStatus::Ended.encode(&mut body, version)?; - suffix.encode(&mut body, version)?; - encode_hops(&mut body, version, hops)?; + AnnounceStatus::Ended.encode(w, version)?; + suffix.encode(w, version)?; + encode_hops(w, version, hops)?; } // The id-referencing forms only exist on lite-06+. Self::EndedId { .. } | Self::Restart { .. } | Self::Skipped => { return Err(EncodeError::Version); } } - (body.len() as u64).encode(w, version)?; - w.put_slice(&body); - Ok(()) + w.prefix_varint(start) } } impl Decode for AnnounceBroadcast<'_> { - fn decode(buf: &mut B, version: Version) -> Result { + fn decode(buf: &mut Decoder<'_>, version: Version) -> Result { if version.has_announce_id() { // Lite06+: outer type, then a size-prefixed body decoded within its bounds. - let typ = u64::decode(buf, version)?; - let size = decode_size(buf, version)?; - if buf.remaining() < size { - return Err(DecodeError::Short); - } - let mut body = buf.take(size); + let typ = buf.varint()?.into_inner(); + let size = decode_size(buf)?; + let mut body = buf.sub(size)?; let msg = match typ { ANNOUNCE_START => Self::Active { suffix: PathRef::decode(&mut body, version)?, @@ -290,35 +283,31 @@ impl Decode for AnnounceBroadcast<'_> { cost: Cost::decode(&mut body, version)?, }, ANNOUNCE_END => Self::EndedId { - id: u64::decode(&mut body, version)?, + id: body.varint()?.into_inner(), }, ANNOUNCE_RESTART => Self::Restart { - id: u64::decode(&mut body, version)?, + id: body.varint()?.into_inner(), hops: HopsRef::decode(&mut body, version)?, cost: Cost::decode(&mut body, version)?, }, // Unknown types are skipped by length so an earlier Lite06 build // negotiating the same ALPN does not kill the announce stream. _ => { - let remaining = body.remaining(); - bytes::Buf::advance(&mut body, remaining); + body.rest(); Self::Skipped } }; - if body.remaining() > 0 { + if !body.is_empty() { return Err(DecodeError::Long); } return Ok(msg); } // Older versions: a single size-prefixed ANNOUNCE_BROADCAST with an inner status. - let size = decode_size(buf, version)?; - if buf.remaining() < size { - return Err(DecodeError::Short); - } - let mut body = buf.take(size); + let size = decode_size(buf)?; + let mut body = buf.sub(size)?; let msg = Self::decode_legacy(&mut body, version)?; - if body.remaining() > 0 { + if !body.is_empty() { return Err(DecodeError::Long); } Ok(msg) @@ -327,7 +316,7 @@ impl Decode for AnnounceBroadcast<'_> { impl AnnounceBroadcast<'_> { /// Decode the body of a pre-lite-06 ANNOUNCE_BROADCAST (inner status + path + hops). - fn decode_legacy(r: &mut R, version: Version) -> Result { + fn decode_legacy(r: &mut Decoder<'_>, version: Version) -> Result { let status = AnnounceStatus::decode(r, version)?; let suffix = Path::decode(r, version)?; @@ -336,7 +325,7 @@ impl AnnounceBroadcast<'_> { Version::Lite03 => { // Lite03 sends only a hop count, not individual ids. Fill with UNKNOWN placeholders. // push() enforces MAX_HOPS and `?` lifts the overflow to DecodeError::BoundsExceeded. - let count = u64::decode(r, version)? as usize; + let count = r.varint()?.into_inner() as usize; let mut list = Hops::new(); for _ in 0..count { list.push(Hop::UNKNOWN)?; @@ -368,10 +357,13 @@ impl AnnounceBroadcast<'_> { } } -fn encode_hops(w: &mut W, version: Version, hops: &Hops) -> Result<(), EncodeError> { +fn encode_hops(w: &mut Encoder<'_>, version: Version, hops: &Hops) -> Result<(), EncodeError> { match version { Version::Lite01 | Version::Lite02 => Ok(()), - Version::Lite03 => (hops.len() as u64).encode(w, version), + Version::Lite03 => { + w.varint(VarInt::from(hops.len()))?; + Ok(()) + } _ => hops.encode(w, version), } } @@ -394,14 +386,14 @@ pub struct AnnounceRequest<'a> { } impl Message for AnnounceRequest<'_> { - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { let prefix = Path::decode(r, version)?; let exclude_hop = match version.has_exclude_hop() { - true => u64::decode(r, version)?, + true => r.varint()?.into_inner(), false => 0, }; let hidden = match version.has_hidden() { - true => bool::decode(r, version)?, + true => r.bool()?, false => false, }; Ok(Self { @@ -411,13 +403,13 @@ impl Message for AnnounceRequest<'_> { }) } - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.prefix.encode(w, version)?; if version.has_exclude_hop() { - self.exclude_hop.encode(w, version)?; + w.varint(VarInt::from(self.exclude_hop))?; } if version.has_hidden() { - self.hidden.encode(w, version)?; + w.bool(self.hidden); } Ok(()) @@ -436,15 +428,16 @@ enum AnnounceStatus { } impl Decode for AnnounceStatus { - fn decode(r: &mut R, version: Version) -> Result { - let status = u8::decode(r, version)?; + fn decode(r: &mut Decoder<'_>, _: Version) -> Result { + let status = r.u8()?; status.try_into().map_err(|_| DecodeError::InvalidValue) } } impl Encode for AnnounceStatus { - fn encode(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { - (*self as u8).encode(w, version) + fn encode(&self, w: &mut Encoder<'_>, _: Version) -> Result<(), EncodeError> { + w.u8(*self as u8); + Ok(()) } } @@ -458,7 +451,7 @@ pub struct AnnounceInit<'a> { } impl Message for AnnounceInit<'_> { - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { match version { Version::Lite01 | Version::Lite02 => {} _ => { @@ -466,7 +459,7 @@ impl Message for AnnounceInit<'_> { } } - let count = u64::decode(r, version)?; + let count = r.varint()?.into_inner(); // Don't allocate more than 1024 elements upfront let mut paths = Vec::with_capacity(count.min(1024) as usize); @@ -478,7 +471,7 @@ impl Message for AnnounceInit<'_> { Ok(Self { suffixes: paths }) } - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { match version { Version::Lite01 | Version::Lite02 => {} _ => { @@ -486,7 +479,7 @@ impl Message for AnnounceInit<'_> { } } - (self.suffixes.len() as u64).encode(w, version)?; + w.varint(VarInt::from(self.suffixes.len()))?; for path in &self.suffixes { path.encode(w, version)?; } @@ -510,23 +503,23 @@ pub struct AnnounceOk { } impl Message for AnnounceOk { - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { if !version.has_announce_ok() { return Err(DecodeError::Version); } let origin = Hop::decode(r, version)?; - let active = u64::decode(r, version)?; + let active = r.varint()?.into_inner(); Ok(Self { origin, active }) } - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { if !version.has_announce_ok() { return Err(EncodeError::Version); } self.origin.encode(w, version)?; - self.active.encode(w, version) + w.varint(VarInt::from(self.active)) } } @@ -538,13 +531,13 @@ mod tests { // Forge an ANNOUNCE_BROADCAST with the draft's explicit `restart` status (2) for the given version. fn encode_forged_restart(version: Version) -> bytes::Bytes { // Encode a normal Active, then flip its status byte (1 -> 2). - let mut buf = bytes::BytesMut::new(); + let mut buf = Vec::new(); AnnounceBroadcast::Active { suffix: PathRef::literal(Path::new("foo/bar")), hops: HopsRef::default(), cost: Cost::default(), } - .encode(&mut buf, version) + .encode(&mut Encoder::new(&mut buf, version.into()), version) .expect("encode"); // Layout: <...>. The message is small, so the size is one byte and @@ -555,7 +548,7 @@ mod tests { "expected an Active status byte" ); buf[1] = u8::from(AnnounceStatus::Restart); - buf.freeze() + bytes::Bytes::from(buf) } // On lite-05+ the explicit `restart` status is accepted and surfaced as an `Active` (the @@ -564,7 +557,8 @@ mod tests { fn decodes_explicit_restart_status_as_active_on_lite05() { let version = Version::Lite05; let mut slice = encode_forged_restart(version); - let decoded = AnnounceBroadcast::decode(&mut slice, version).expect("explicit restart must decode"); + let decoded = crate::coding::decode_buf(&mut slice, version, AnnounceBroadcast::decode) + .expect("explicit restart must decode"); assert!(!slice.has_remaining(), "trailing bytes after decode"); assert!( matches!(decoded, AnnounceBroadcast::Active { .. }), @@ -579,7 +573,7 @@ mod tests { let mut slice = encode_forged_restart(version); assert!( matches!( - AnnounceBroadcast::decode(&mut slice, version), + crate::coding::decode_buf(&mut slice, version, AnnounceBroadcast::decode), Err(DecodeError::InvalidValue) ), "restart status must be rejected before lite-05" @@ -587,10 +581,11 @@ mod tests { } fn round_trip(msg: &AnnounceOk) -> AnnounceOk { - let mut buf = bytes::BytesMut::new(); - msg.encode(&mut buf, Version::Lite05).unwrap(); + let mut buf = Vec::new(); + msg.encode(&mut Encoder::new(&mut buf, Version::Lite05.into()), Version::Lite05) + .unwrap(); let mut slice = &buf[..]; - let got = AnnounceOk::decode(&mut slice, Version::Lite05).unwrap(); + let got = crate::coding::decode_buf(&mut slice, Version::Lite05, AnnounceOk::decode).unwrap(); assert!(slice.is_empty(), "trailing bytes after decode"); got } @@ -614,10 +609,11 @@ mod tests { } fn broadcast_round_trip(msg: &AnnounceBroadcast, version: Version) -> AnnounceBroadcast<'static> { - let mut buf = bytes::BytesMut::new(); - msg.encode(&mut buf, version).unwrap(); + let mut buf = Vec::new(); + msg.encode(&mut Encoder::new(&mut buf, version.into()), version) + .unwrap(); let mut slice = &buf[..]; - let got = AnnounceBroadcast::decode(&mut slice, version).unwrap(); + let got = crate::coding::decode_buf(&mut slice, version, AnnounceBroadcast::decode).unwrap(); assert!(slice.is_empty(), "trailing bytes after decode"); got.into_owned() } @@ -710,7 +706,7 @@ mod tests { let mut buf = vec![ANNOUNCE_START as u8, body.len() as u8]; buf.extend_from_slice(&body); assert!(matches!( - AnnounceBroadcast::decode(&mut &buf[..], Version::Lite07), + crate::coding::decode_buf(&mut &buf[..], Version::Lite07, AnnounceBroadcast::decode), Err(DecodeError::InvalidValue) )); } @@ -729,17 +725,21 @@ mod tests { cost: Cost::default(), }; for version in [Version::Lite05, Version::Lite06] { - let mut buf = bytes::BytesMut::new(); - assert!(matches!(msg.encode(&mut buf, version), Err(EncodeError::Version))); + let mut buf = Vec::new(); + assert!(matches!( + msg.encode(&mut Encoder::new(&mut buf, version.into()), version), + Err(EncodeError::Version) + )); } } // The id-referencing forms don't exist before lite-06, and the path form is gone on lite-06. #[test] fn announce_broadcast_rejects_cross_version_forms() { - let mut buf = bytes::BytesMut::new(); + let mut buf = Vec::new(); assert!(matches!( - AnnounceBroadcast::EndedId { id: 1 }.encode(&mut buf, Version::Lite05), + AnnounceBroadcast::EndedId { id: 1 } + .encode(&mut Encoder::new(&mut buf, Version::Lite05.into()), Version::Lite05), Err(EncodeError::Version) )); assert!(matches!( @@ -748,7 +748,7 @@ mod tests { hops: HopsRef::default(), cost: Cost::default() } - .encode(&mut buf, Version::Lite05), + .encode(&mut Encoder::new(&mut buf, Version::Lite05.into()), Version::Lite05), Err(EncodeError::Version) )); assert!(matches!( @@ -756,7 +756,7 @@ mod tests { suffix: Path::new("room/cam"), hops: Hops::new() } - .encode(&mut buf, Version::Lite06), + .encode(&mut Encoder::new(&mut buf, Version::Lite06.into()), Version::Lite06), Err(EncodeError::Version) )); } @@ -790,25 +790,35 @@ mod tests { let mut buf = Vec::new(); crate::origin::Cost::MAX .charged(1) - .encode(&mut buf, Version::Lite06) + .encode(&mut Encoder::new(&mut buf, Version::Lite06.into()), Version::Lite06) .expect("a charged cost must stay encodable"); } #[test] fn unknown_announce_type_is_skipped() { let mut body = Vec::new(); - Path::new("room/cam").encode(&mut body, Version::Lite06).unwrap(); - Hops::new().encode(&mut body, Version::Lite06).unwrap(); - Cost::default().encode(&mut body, Version::Lite06).unwrap(); + Path::new("room/cam") + .encode(&mut Encoder::new(&mut body, Version::Lite06.into()), Version::Lite06) + .unwrap(); + Hops::new() + .encode(&mut Encoder::new(&mut body, Version::Lite06.into()), Version::Lite06) + .unwrap(); + Cost::default() + .encode(&mut Encoder::new(&mut body, Version::Lite06.into()), Version::Lite06) + .unwrap(); - let mut buf = bytes::BytesMut::new(); - 4u64.encode(&mut buf, Version::Lite06).unwrap(); - (body.len() as u64).encode(&mut buf, Version::Lite06).unwrap(); + let mut buf = Vec::new(); + Encoder::new(&mut buf, Version::Lite06.into()) + .varint(VarInt::from(4u64)) + .unwrap(); + Encoder::new(&mut buf, Version::Lite06.into()) + .varint(VarInt::from(body.len())) + .unwrap(); buf.extend_from_slice(&body); let mut slice = &buf[..]; - let got = - AnnounceBroadcast::decode(&mut slice, Version::Lite06).expect("unknown type must not kill the stream"); + let got = crate::coding::decode_buf(&mut slice, Version::Lite06, AnnounceBroadcast::decode) + .expect("unknown type must not kill the stream"); assert!(slice.is_empty()); assert_eq!(got, AnnounceBroadcast::Skipped); } @@ -816,18 +826,19 @@ mod tests { // An ANNOUNCE_END message on lite-06 is tiny: type byte, size prefix, id varint. #[test] fn ended_by_id_is_three_bytes() { - let mut buf = bytes::BytesMut::new(); + let mut buf = Vec::new(); AnnounceBroadcast::EndedId { id: 42 } - .encode(&mut buf, Version::Lite06) + .encode(&mut Encoder::new(&mut buf, Version::Lite06.into()), Version::Lite06) .unwrap(); assert_eq!(buf.len(), 3); } fn request_round_trip(msg: &AnnounceRequest, version: Version) -> AnnounceRequest<'static> { - let mut buf = bytes::BytesMut::new(); - msg.encode(&mut buf, version).unwrap(); + let mut buf = Vec::new(); + msg.encode(&mut Encoder::new(&mut buf, version.into()), version) + .unwrap(); let mut slice = &buf[..]; - let got = AnnounceRequest::decode(&mut slice, version).unwrap(); + let got = crate::coding::decode_buf(&mut slice, version, AnnounceRequest::decode).unwrap(); assert!(slice.is_empty(), "trailing bytes after decode"); AnnounceRequest { prefix: got.prefix.to_owned(), @@ -853,13 +864,17 @@ mod tests { // A flag byte other than 0 or 1 is malformed, not a future extension. #[test] fn announce_request_rejects_a_bad_hidden_flag() { - let mut buf = bytes::BytesMut::new(); + let mut buf = Vec::new(); let mut body = Vec::new(); - Path::new("room").encode(&mut body, Version::Lite07).unwrap(); + Path::new("room") + .encode(&mut Encoder::new(&mut body, Version::Lite07.into()), Version::Lite07) + .unwrap(); body.push(2); - (body.len() as u64).encode(&mut buf, Version::Lite07).unwrap(); + Encoder::new(&mut buf, Version::Lite07.into()) + .varint(VarInt::from(body.len())) + .unwrap(); buf.extend_from_slice(&body); - assert!(AnnounceRequest::decode(&mut &buf[..], Version::Lite07).is_err()); + assert!(crate::coding::decode_buf(&mut &buf[..], Version::Lite07, AnnounceRequest::decode).is_err()); } // Lite04/05 carry the subscriber's origin id so the publisher can skip reflected @@ -886,10 +901,12 @@ mod tests { assert_eq!(request_round_trip(&msg, Version::Lite06).exclude_hop, 0); // And it costs nothing on the wire: the body is just the prefix. - let mut with = bytes::BytesMut::new(); - msg.encode(&mut with, Version::Lite05).unwrap(); - let mut without = bytes::BytesMut::new(); - msg.encode(&mut without, Version::Lite06).unwrap(); + let mut with = Vec::new(); + msg.encode(&mut Encoder::new(&mut with, Version::Lite05.into()), Version::Lite05) + .unwrap(); + let mut without = Vec::new(); + msg.encode(&mut Encoder::new(&mut without, Version::Lite06.into()), Version::Lite06) + .unwrap(); assert!( without.len() < with.len(), "lite06 must not encode the exclude_hop varint" @@ -902,9 +919,9 @@ mod tests { origin: Hop::new(1).unwrap(), active: 0, }; - let mut buf = bytes::BytesMut::new(); + let mut buf = Vec::new(); assert!(matches!( - msg.encode(&mut buf, Version::Lite04), + msg.encode(&mut Encoder::new(&mut buf, Version::Lite04.into()), Version::Lite04), Err(EncodeError::Version) )); } @@ -912,12 +929,12 @@ mod tests { #[test] fn announce_ok_accepts_zero_origin() { // Encode a well-formed message then patch the origin to 0 on the wire. - let mut buf = bytes::BytesMut::new(); + let mut buf = Vec::new(); AnnounceOk { origin: Hop::new(1).unwrap(), active: 0, } - .encode(&mut buf, Version::Lite05) + .encode(&mut Encoder::new(&mut buf, Version::Lite05.into()), Version::Lite05) .unwrap(); // origin id 1 sits right after the size prefix; rewrite it to 0. let bytes = &buf[..]; @@ -925,7 +942,7 @@ mod tests { // size(1 byte) | origin varint(1 byte = 0x01) | active varint(1 byte) patched[1] = 0x00; let mut slice = &patched[..]; - let got = AnnounceOk::decode(&mut slice, Version::Lite05).unwrap(); + let got = crate::coding::decode_buf(&mut slice, Version::Lite05, AnnounceOk::decode).unwrap(); assert_eq!(got.origin.id(), 0); assert_eq!(got.active, 0); } diff --git a/rs/moq-net/src/lite/compress.rs b/rs/moq-net/src/lite/compress.rs index e908e8cbec..8f35f992b3 100644 --- a/rs/moq-net/src/lite/compress.rs +++ b/rs/moq-net/src/lite/compress.rs @@ -13,7 +13,10 @@ use std::{ hash::Hash, }; -use crate::{Error, Hop, Hops, Path, PathOwned, coding::Encode, coding::Sizer}; +use crate::{ + Error, Hop, Hops, Path, PathOwned, + coding::{Form, VarInt}, +}; use super::{HopsRef, PathRef, Version}; @@ -287,14 +290,6 @@ impl AnnounceEncoder { self.tails.remove(reversed(&entry.hops), id); } - fn size>(&self, value: &T) -> usize { - let mut sizer = Sizer::default(); - value - .encode(&mut sizer, self.version) - .expect("sizing an encodable value"); - sizer.size - } - /// The wire form of `suffix`, sized from its parts so only the winner is built. fn path_ref(&self, suffix: &PathOwned, next: u64) -> PathRef<'static> { let literal = || PathRef::literal(suffix.clone()); @@ -308,8 +303,8 @@ impl AnnounceEncoder { let base = next - id; let keep = keep as u64; - let size = self.size(&base) + self.size(&keep) + self.size(&rest); - match size < self.size(&0u64) * 2 + self.size(suffix) { + let size = varint_size(base) + varint_size(keep) + string_size(&rest); + match size < varint_size(0) * 2 + string_size(suffix) { true => PathRef { base, keep, rest }, false => literal(), } @@ -326,9 +321,9 @@ impl AnnounceEncoder { let base = next - id; let keep = keep as u64; let chain = - |hops: &[Hop]| self.size(&(hops.len() as u64)) + hops.iter().map(|hop| self.size(hop)).sum::(); - let size = self.size(&base) + chain(head) + self.size(&keep); - match size < self.size(&0u64) * 2 + chain(hops.as_slice()) { + |hops: &[Hop]| varint_size(hops.len() as u64) + hops.iter().map(|hop| varint_size(hop.id())).sum::(); + let size = varint_size(base) + chain(head) + varint_size(keep); + match size < varint_size(0) * 2 + chain(hops.as_slice()) { true => HopsRef { base, literal: Hops::try_from(head.to_vec()).expect("a prefix of a valid chain is valid"), @@ -339,9 +334,23 @@ impl AnnounceEncoder { } } +/// The bytes `value` takes as a moq-lite varint. +fn varint_size(value: u64) -> usize { + VarInt::from(value) + .size(Form::Quic) + .expect("sizing a value in varint range") +} + +/// The bytes `path` takes on the wire: a varint length, then the string. +fn string_size(path: &Path<'_>) -> usize { + let len = path.as_str().len(); + varint_size(len as u64) + len +} + #[cfg(test)] mod tests { use super::*; + use crate::coding::Encode; const VERSION: Version = Version::Lite07; @@ -498,7 +507,7 @@ mod tests { /// resolved values match the input. fn start(encoder: &mut AnnounceEncoder, decoder: &mut AnnounceDecoder, suffix: &str, ids: &[u64]) -> (u64, usize) { let (id, wire_path, wire_hops) = encoder.start(Path::new(suffix).to_owned(), hops(ids)); - let size = encoder.size(&wire_path) + encoder.size(&wire_hops); + let size = wire_path.encode_bytes(VERSION).unwrap().len() + wire_hops.encode_bytes(VERSION).unwrap().len(); let (got_path, got_hops) = decoder.start(wire_path, wire_hops).unwrap(); assert_eq!(got_path.as_str(), suffix); assert_eq!(got_hops, hops(ids)); diff --git a/rs/moq-net/src/lite/datagram.rs b/rs/moq-net/src/lite/datagram.rs index 6e81d09b1e..e02c44cf2b 100644 --- a/rs/moq-net/src/lite/datagram.rs +++ b/rs/moq-net/src/lite/datagram.rs @@ -5,9 +5,9 @@ //! boundary, so unlike a [`super::Message`] there is no inner length prefix. The model counterpart //! is [`crate::Datagram`]. -use bytes::{Buf, BufMut, Bytes}; +use bytes::Bytes; -use crate::coding::{Decode, DecodeError, Encode, EncodeError}; +use crate::coding::{DecodeError, Decoder, Encode, EncodeError, Encoder, VarInt}; use super::Version; @@ -25,35 +25,35 @@ pub struct Datagram { } impl Encode for Datagram { - fn encode(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { if !version.has_datagrams() { return Err(EncodeError::Version); } - self.subscribe.encode(w, version)?; - self.sequence.encode(w, version)?; - self.timestamp.encode(w, version)?; + w.varint(VarInt::from(self.subscribe))?; + w.varint(VarInt::from(self.sequence))?; + w.varint(VarInt::from(self.timestamp))?; // Payload runs to the datagram boundary: written raw, no length prefix. - if w.remaining_mut() < self.payload.len() { - return Err(EncodeError::Short); - } - w.put_slice(&self.payload); + w.slice(&self.payload); Ok(()) } } -impl Decode for Datagram { - fn decode(r: &mut R, version: Version) -> Result { +impl Datagram { + /// Decode a whole datagram body. The payload shares `buf` rather than copying it. + pub fn decode(buf: Bytes, version: Version) -> Result { if !version.has_datagrams() { return Err(DecodeError::Version); } - let subscribe = u64::decode(r, version)?; - let sequence = u64::decode(r, version)?; - let timestamp = u64::decode(r, version)?; + let mut r = Decoder::new(&buf, version.into()); + let subscribe = r.varint()?.into_inner(); + let sequence = r.varint()?.into_inner(); + let timestamp = r.varint()?.into_inner(); + // Everything remaining is the payload (the datagram boundary delimits it). - let payload = r.copy_to_bytes(r.remaining()); + let payload = buf.slice(buf.len() - r.remaining()..); Ok(Self { subscribe, @@ -67,7 +67,6 @@ impl Decode for Datagram { #[cfg(test)] mod test { use super::*; - use bytes::BytesMut; #[test] fn roundtrip() { @@ -77,12 +76,9 @@ mod test { timestamp: 1_000, payload: Bytes::from_static(b"hello"), }; - let mut buf = BytesMut::new(); - original.encode(&mut buf, Version::Lite05).unwrap(); - let mut slice = &buf[..]; - let decoded = Datagram::decode(&mut slice, Version::Lite05).unwrap(); - assert_eq!(decoded, original); - assert!(!slice.has_remaining(), "payload has no trailing length prefix"); + let buf = original.encode_bytes(Version::Lite05).unwrap(); + let decoded = Datagram::decode(buf, Version::Lite05).unwrap(); + assert_eq!(decoded, original, "payload has no trailing length prefix"); } #[test] @@ -93,10 +89,8 @@ mod test { timestamp: 0, payload: Bytes::new(), }; - let mut buf = BytesMut::new(); - original.encode(&mut buf, Version::Lite05).unwrap(); - let mut slice = &buf[..]; - let decoded = Datagram::decode(&mut slice, Version::Lite05).unwrap(); + let buf = original.encode_bytes(Version::Lite05).unwrap(); + let decoded = Datagram::decode(buf, Version::Lite05).unwrap(); assert_eq!(decoded, original); } @@ -124,15 +118,10 @@ mod test { timestamp: 3, payload: Bytes::from_static(b"x"), }; - let mut buf = BytesMut::new(); - assert!(matches!( - dg.encode(&mut buf, Version::Lite04), - Err(EncodeError::Version) - )); + assert!(matches!(dg.encode_bytes(Version::Lite04), Err(EncodeError::Version))); - let mut slice = &b"\x01\x02\x03x"[..]; assert!(matches!( - Datagram::decode(&mut slice, Version::Lite04), + Datagram::decode(Bytes::from_static(b"\x01\x02\x03x"), Version::Lite04), Err(DecodeError::Version) )); } diff --git a/rs/moq-net/src/lite/fetch.rs b/rs/moq-net/src/lite/fetch.rs index c098c5990a..a234fbab8d 100644 --- a/rs/moq-net/src/lite/fetch.rs +++ b/rs/moq-net/src/lite/fetch.rs @@ -2,7 +2,7 @@ use std::borrow::Cow; use crate::{ Path, - coding::{Decode, DecodeError, Encode, EncodeError}, + coding::{Decode, DecodeError, Decoder, Encode, EncodeError, Encoder, VarInt}, }; use super::{Message, Version}; @@ -24,7 +24,7 @@ pub struct Fetch<'a> { } impl Message for Fetch<'_> { - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { match version { Version::Lite01 | Version::Lite02 => { return Err(DecodeError::Version); @@ -33,12 +33,12 @@ impl Message for Fetch<'_> { } let broadcast = Path::decode(r, version)?; - let track = Cow::::decode(r, version)?; - let priority = u8::decode(r, version)?; - let group = u64::decode(r, version)?; + let track = Cow::Owned(r.string()?); + let priority = r.u8()?; + let group = r.varint()?.into_inner(); let (start_frame, end_frame) = match version.has_frame_bounds() { - true => (u64::decode(r, version)?, Option::::decode(r, version)?), + true => (r.varint()?.into_inner(), r.varint_opt()?), false => (0, None), }; // A range that ends before it starts can never be served. @@ -56,7 +56,7 @@ impl Message for Fetch<'_> { }) } - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { match version { Version::Lite01 | Version::Lite02 => { return Err(EncodeError::Version); @@ -65,13 +65,13 @@ impl Message for Fetch<'_> { } self.broadcast.encode(w, version)?; - self.track.encode(w, version)?; - self.priority.encode(w, version)?; - self.group.encode(w, version)?; + w.string(&self.track)?; + w.u8(self.priority); + w.varint(VarInt::from(self.group))?; if version.has_frame_bounds() { - self.start_frame.encode(w, version)?; - self.end_frame.encode(w, version)?; + w.varint(VarInt::from(self.start_frame))?; + w.varint_opt(self.end_frame)?; } else if self.start_frame != 0 || self.end_frame.is_some() { // The peer would serve the whole group, including frames the caller excluded. return Err(EncodeError::Version); @@ -98,9 +98,10 @@ mod test { fn fetch_roundtrip(version: Version, msg: &Fetch<'_>) -> Fetch<'static> { let mut buf = Vec::new(); - msg.encode_msg(&mut buf, version).unwrap(); + msg.encode_msg(&mut Encoder::new(&mut buf, version.into()), version) + .unwrap(); let mut slice = buf.as_slice(); - Fetch::decode_msg(&mut slice, version).unwrap() + crate::coding::decode_buf(&mut slice, version, Fetch::decode_msg).unwrap() } #[test] @@ -132,9 +133,10 @@ mod test { msg.end_frame = Some(2); let mut buf = Vec::new(); - msg.encode_msg(&mut buf, Version::Lite06).unwrap(); + msg.encode_msg(&mut Encoder::new(&mut buf, Version::Lite06.into()), Version::Lite06) + .unwrap(); assert!(matches!( - Fetch::decode_msg(&mut buf.as_slice(), Version::Lite06), + crate::coding::decode_buf(&mut buf.as_slice(), Version::Lite06, Fetch::decode_msg), Err(DecodeError::InvalidSubscribeLocation) )); } @@ -146,12 +148,19 @@ mod test { msg.start_frame = 2; let mut buf = Vec::new(); - assert!(msg.encode_msg(&mut buf, Version::Lite05).is_err()); + assert!( + msg.encode_msg(&mut Encoder::new(&mut buf, Version::Lite05.into()), Version::Lite05) + .is_err() + ); } #[test] fn fetch_rejected_before_lite03() { let mut buf = Vec::new(); - assert!(fetch_sample().encode_msg(&mut buf, Version::Lite02).is_err()); + assert!( + fetch_sample() + .encode_msg(&mut Encoder::new(&mut buf, Version::Lite02.into()), Version::Lite02) + .is_err() + ); } } diff --git a/rs/moq-net/src/lite/goaway.rs b/rs/moq-net/src/lite/goaway.rs index 0152c2d4bd..a7b7e7fb10 100644 --- a/rs/moq-net/src/lite/goaway.rs +++ b/rs/moq-net/src/lite/goaway.rs @@ -13,7 +13,7 @@ pub struct Goaway<'a> { } impl Message for Goaway<'_> { - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { match version { Version::Lite01 | Version::Lite02 | Version::Lite03 => { return Err(DecodeError::Version); @@ -25,18 +25,15 @@ impl Message for Goaway<'_> { // cap. Rejected from the string's length prefix alone, before allocating // or validating the payload. (Buffering is bounded separately by the // outer message-size prefix that frames every lite control message.) - let len = usize::decode(r, version)?; + let len = r.varint()?.into_inner(); if len > 8192 { return Err(DecodeError::InvalidValue); } - if r.remaining() < len { - return Err(DecodeError::Short); - } - let uri = String::from_utf8(r.copy_to_bytes(len).to_vec())?; + let uri = String::from_utf8(r.slice(len as usize)?.to_vec())?; Ok(Self { uri: Cow::Owned(uri) }) } - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { match version { Version::Lite01 | Version::Lite02 | Version::Lite03 => { return Err(EncodeError::Version); @@ -44,7 +41,7 @@ impl Message for Goaway<'_> { _ => {} } - self.uri.encode(w, version)?; + w.string(&self.uri)?; Ok(()) } } @@ -52,27 +49,30 @@ impl Message for Goaway<'_> { #[cfg(test)] mod tests { use super::*; - use bytes::BytesMut; #[test] fn roundtrip_with_uri() { let msg = Goaway { uri: Cow::Borrowed("https://relay.example/new"), }; - let mut buf = BytesMut::new(); - msg.encode_msg(&mut buf, Version::Lite04).unwrap(); + let mut buf = Vec::new(); + msg.encode_msg(&mut Encoder::new(&mut buf, Version::Lite04.into()), Version::Lite04) + .unwrap(); - let decoded = Goaway::decode_msg(&mut buf.freeze(), Version::Lite04).unwrap(); + let decoded = + crate::coding::decode_buf(&mut bytes::Bytes::from(buf), Version::Lite04, Goaway::decode_msg).unwrap(); assert_eq!(decoded.uri, "https://relay.example/new"); } #[test] fn roundtrip_empty() { let msg = Goaway { uri: Cow::Borrowed("") }; - let mut buf = BytesMut::new(); - msg.encode_msg(&mut buf, Version::Lite04).unwrap(); + let mut buf = Vec::new(); + msg.encode_msg(&mut Encoder::new(&mut buf, Version::Lite04.into()), Version::Lite04) + .unwrap(); - let decoded = Goaway::decode_msg(&mut buf.freeze(), Version::Lite04).unwrap(); + let decoded = + crate::coding::decode_buf(&mut bytes::Bytes::from(buf), Version::Lite04, Goaway::decode_msg).unwrap(); assert_eq!(decoded.uri, ""); } @@ -81,15 +81,25 @@ mod tests { let msg = Goaway { uri: Cow::Borrowed("https://relay.example/new"), }; - let mut buf = BytesMut::new(); + let mut buf = Vec::new(); // Encoding should fail on Lite03. - assert!(msg.encode_msg(&mut buf, Version::Lite03).is_err()); + assert!( + msg.encode_msg(&mut Encoder::new(&mut buf, Version::Lite03.into()), Version::Lite03) + .is_err() + ); // Even if we have valid bytes, decoding on Lite03 should fail. - let mut encode_buf = BytesMut::new(); - msg.encode_msg(&mut encode_buf, Version::Lite04).unwrap(); - assert!(Goaway::decode_msg(&mut encode_buf.freeze(), Version::Lite03).is_err()); + let mut encode_buf = Vec::new(); + msg.encode_msg( + &mut Encoder::new(&mut encode_buf, Version::Lite04.into()), + Version::Lite04, + ) + .unwrap(); + assert!( + crate::coding::decode_buf(&mut bytes::Bytes::from(encode_buf), Version::Lite03, Goaway::decode_msg) + .is_err() + ); } /// The URI is capped at 8,192 bytes (matching the IETF wire), rejected from @@ -101,9 +111,11 @@ mod tests { let msg = Goaway { uri: Cow::Borrowed(&at_cap), }; - let mut buf = BytesMut::new(); - msg.encode_msg(&mut buf, Version::Lite04).unwrap(); - let decoded = Goaway::decode_msg(&mut buf.freeze(), Version::Lite04).unwrap(); + let mut buf = Vec::new(); + msg.encode_msg(&mut Encoder::new(&mut buf, Version::Lite04.into()), Version::Lite04) + .unwrap(); + let decoded = + crate::coding::decode_buf(&mut bytes::Bytes::from(buf), Version::Lite04, Goaway::decode_msg).unwrap(); assert_eq!(decoded.uri.len(), 8192); // One byte over: rejected as InvalidValue, without needing the payload @@ -112,13 +124,14 @@ mod tests { let msg = Goaway { uri: Cow::Borrowed(&over_cap), }; - let mut buf = BytesMut::new(); - msg.encode_msg(&mut buf, Version::Lite04).unwrap(); - let mut truncated = buf.freeze(); + let mut buf = Vec::new(); + msg.encode_msg(&mut Encoder::new(&mut buf, Version::Lite04.into()), Version::Lite04) + .unwrap(); + let mut truncated = bytes::Bytes::from(buf); // Keep only the length prefix plus a little payload. let mut short = truncated.split_to(16); assert!(matches!( - Goaway::decode_msg(&mut short, Version::Lite04), + crate::coding::decode_buf(&mut short, Version::Lite04, Goaway::decode_msg), Err(DecodeError::InvalidValue) )); } diff --git a/rs/moq-net/src/lite/group.rs b/rs/moq-net/src/lite/group.rs index 2fb42358fe..2cf52a84d9 100644 --- a/rs/moq-net/src/lite/group.rs +++ b/rs/moq-net/src/lite/group.rs @@ -18,11 +18,11 @@ pub struct Group { } impl Message for Group { - fn decode_msg(r: &mut R, version: Version) -> Result { - let subscribe = u64::decode(r, version)?; - let sequence = u64::decode(r, version)?; + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { + let subscribe = r.varint()?.into_inner(); + let sequence = r.varint()?.into_inner(); let frame_start = match version.has_frame_bounds() { - true => u64::decode(r, version)?, + true => r.varint()?.into_inner(), false => 0, }; @@ -33,12 +33,12 @@ impl Message for Group { }) } - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { - self.subscribe.encode(w, version)?; - self.sequence.encode(w, version)?; + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { + w.varint(VarInt::from(self.subscribe))?; + w.varint(VarInt::from(self.sequence))?; if version.has_frame_bounds() { - self.frame_start.encode(w, version)?; + w.varint(VarInt::from(self.frame_start))?; } else if self.frame_start != 0 { // The peer would number the frames from 0 and silently misalign the group. return Err(EncodeError::Version); @@ -61,8 +61,9 @@ mod test { frame_start: 4, }; let mut buf = Vec::new(); - msg.encode_msg(&mut buf, Version::Lite06).unwrap(); - let got = Group::decode_msg(&mut buf.as_slice(), Version::Lite06).unwrap(); + msg.encode_msg(&mut Encoder::new(&mut buf, Version::Lite06.into()), Version::Lite06) + .unwrap(); + let got = crate::coding::decode_buf(&mut buf.as_slice(), Version::Lite06, Group::decode_msg).unwrap(); assert_eq!((got.sequence, got.frame_start), (7, 4)); } @@ -75,12 +76,17 @@ mod test { frame_start: 4, }; let mut buf = Vec::new(); - assert!(msg.encode_msg(&mut buf, Version::Lite05).is_err()); + assert!( + msg.encode_msg(&mut Encoder::new(&mut buf, Version::Lite05.into()), Version::Lite05) + .is_err() + ); let whole = Group { frame_start: 0, ..msg }; let mut buf = Vec::new(); - whole.encode_msg(&mut buf, Version::Lite05).unwrap(); - let got = Group::decode_msg(&mut buf.as_slice(), Version::Lite05).unwrap(); + whole + .encode_msg(&mut Encoder::new(&mut buf, Version::Lite05.into()), Version::Lite05) + .unwrap(); + let got = crate::coding::decode_buf(&mut buf.as_slice(), Version::Lite05, Group::decode_msg).unwrap(); assert_eq!(got.frame_start, 0); } } diff --git a/rs/moq-net/src/lite/info.rs b/rs/moq-net/src/lite/info.rs index d85827e8fa..102c3c3c3f 100644 --- a/rs/moq-net/src/lite/info.rs +++ b/rs/moq-net/src/lite/info.rs @@ -8,7 +8,7 @@ pub struct SessionInfo { } impl Message for SessionInfo { - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { match version { Version::Lite01 | Version::Lite02 => {} _ => { @@ -16,7 +16,7 @@ impl Message for SessionInfo { } } - let bitrate = match u64::decode(r, version)? { + let bitrate = match r.varint()?.into_inner() { 0 => None, bitrate => Some(bitrate), }; @@ -24,7 +24,7 @@ impl Message for SessionInfo { Ok(Self { bitrate }) } - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { match version { Version::Lite01 | Version::Lite02 => {} _ => { @@ -32,7 +32,7 @@ impl Message for SessionInfo { } } - self.bitrate.unwrap_or(0).encode(w, version)?; + w.varint(VarInt::from(self.bitrate.unwrap_or(0)))?; Ok(()) } } diff --git a/rs/moq-net/src/lite/message.rs b/rs/moq-net/src/lite/message.rs index 0d28eb735f..3133665543 100644 --- a/rs/moq-net/src/lite/message.rs +++ b/rs/moq-net/src/lite/message.rs @@ -1,6 +1,4 @@ -use bytes::{Buf, BufMut}; - -use crate::coding::{Decode, DecodeError, Encode, EncodeError, Sizer}; +use crate::coding::{Decode, DecodeError, Decoder, Encode, EncodeError, Encoder}; use super::Version; @@ -8,15 +6,16 @@ use super::Version; // decoding, so the limit must be checked as soon as their length prefix arrives. pub(super) const MAX_MESSAGE_SIZE: usize = 64 * 1024 * 1024; -pub(super) fn decode_size(buf: &mut B, version: Version) -> Result { - let size = usize::decode(buf, version)?; - if size > MAX_MESSAGE_SIZE { - return Err(DecodeError::MessageTooLarge { - size, +/// Read a lite message's varint size prefix, refusing one past [`MAX_MESSAGE_SIZE`]. +pub(super) fn decode_size(r: &mut Decoder<'_>) -> Result { + let size = r.varint()?.into_inner(); + match usize::try_from(size) { + Ok(size) if size <= MAX_MESSAGE_SIZE => Ok(size), + _ => Err(DecodeError::MessageTooLarge { + size: usize::try_from(size).unwrap_or(usize::MAX), max: MAX_MESSAGE_SIZE, - }); + }), } - Ok(size) } /// A trait for lite messages that are automatically size-prefixed during encoding/decoding. @@ -24,91 +23,64 @@ pub(super) fn decode_size(buf: &mut B, version: Version) -> Result(&self, w: &mut W, version: Version) -> Result<(), EncodeError>; + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError>; /// Decode a message body (without size prefix). - fn decode_msg(buf: &mut B, version: Version) -> Result; + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result; } impl Encode for T { - fn encode(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { tracing::trace!(?self, "encoding"); - let mut sizer = Sizer::default(); - self.encode_msg(&mut sizer, version)?; - sizer.size.encode(w, version)?; - self.encode_msg(w, version) + let start = w.position(); + self.encode_msg(w, version)?; + w.prefix_varint(start) } } impl Decode for T { - fn decode(buf: &mut B, version: Version) -> Result { - let size = decode_size(buf, version)?; - - if tracing::enabled!(tracing::Level::TRACE) { - if buf.remaining() < size { - return Err(DecodeError::Short); - } - let raw = buf.copy_to_bytes(size); - let mut slice = &raw[..]; - match Self::decode_msg(&mut slice, version) { - Ok(result) => { - if slice.remaining() > 0 { - return Err(DecodeError::Long); - } - tracing::trace!(?result, "decoded"); - Ok(result) - } - Err(e) => { - tracing::warn!(%e, ?raw, "decode failed"); - Err(e) - } - } - } else { - if buf.remaining() < size { - return Err(DecodeError::Short); - } - let mut limited = buf.take(size); - match Self::decode_msg(&mut limited, version) { - Ok(result) => { - if limited.remaining() > 0 { - return Err(DecodeError::Long); - } - Ok(result) - } - Err(e) => { - tracing::warn!(%e, "decode failed"); - Err(e) - } - } + fn decode(r: &mut Decoder<'_>, version: Version) -> Result { + let size = decode_size(r)?; + let mut body = r.sub(size)?; + + let result = Self::decode_msg(&mut body, version).and_then(|msg| match body.is_empty() { + true => Ok(msg), + false => Err(DecodeError::Long), + }); + + match &result { + Ok(msg) => tracing::trace!(?msg, "decoded"), + Err(err) => tracing::warn!(%err, "decode failed"), } + result } } #[cfg(test)] mod tests { use super::*; + use crate::coding::VarInt; #[derive(Debug)] struct Empty; impl Message for Empty { - fn encode_msg(&self, _: &mut W, _: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, _: &mut Encoder<'_>, _: Version) -> Result<(), EncodeError> { Ok(()) } - fn decode_msg(_: &mut B, _: Version) -> Result { + fn decode_msg(_: &mut Decoder<'_>, _: Version) -> Result { Ok(Self) } } #[test] fn rejects_oversized_message_before_reading_the_body() { - let mut wire = Vec::new(); - ((MAX_MESSAGE_SIZE + 1) as u64) - .encode(&mut wire, Version::Lite06) + let wire = VarInt::from(MAX_MESSAGE_SIZE + 1) + .encode_bytes(Version::Lite06) .unwrap(); - let err = Empty::decode(&mut wire.as_slice(), Version::Lite06).unwrap_err(); + let err = Empty::decode_slice(&wire, Version::Lite06).unwrap_err(); assert!(matches!( err, DecodeError::MessageTooLarge { @@ -120,10 +92,9 @@ mod tests { #[test] fn accepts_message_at_the_limit() { - let mut wire = Vec::new(); - (MAX_MESSAGE_SIZE as u64).encode(&mut wire, Version::Lite06).unwrap(); + let wire = VarInt::from(MAX_MESSAGE_SIZE).encode_bytes(Version::Lite06).unwrap(); - let err = Empty::decode(&mut wire.as_slice(), Version::Lite06).unwrap_err(); + let err = Empty::decode_slice(&wire, Version::Lite06).unwrap_err(); assert!(matches!(err, DecodeError::Short)); } } diff --git a/rs/moq-net/src/lite/parameters.rs b/rs/moq-net/src/lite/parameters.rs index b0148a6aae..e10cd1909b 100644 --- a/rs/moq-net/src/lite/parameters.rs +++ b/rs/moq-net/src/lite/parameters.rs @@ -1,5 +1,3 @@ -use std::collections::HashMap; - use crate::coding::*; use super::Version; @@ -9,35 +7,45 @@ const MAX_PARAMS: u64 = 64; /// A bag of `id -> raw bytes` parameters, the body shared by SETUP (and any other /// parameterized message). Encoded as a varint count followed by `id, length, value` /// triples; duplicate ids are rejected on decode. +/// +/// A handful at most, so a linear scan beats hashing, and the encoding keeps the +/// order the parameters were set or decoded in. #[derive(Default, Debug, Clone)] -pub struct Parameters(HashMap>); +pub struct Parameters(Vec<(u64, Vec)>); impl Parameters { /// Set a parameter to a raw byte value, replacing any existing entry. pub fn set_bytes(&mut self, id: u64, value: Vec) { - self.0.insert(id, value); + match self.0.iter_mut().find(|(k, _)| *k == id) { + Some((_, v)) => *v = value, + None => self.0.push((id, value)), + } } /// Borrow a parameter's raw byte value, if present. pub fn get_bytes(&self, id: u64) -> Option<&[u8]> { - self.0.get(&id).map(Vec::as_slice) + self.0.iter().find(|(k, _)| *k == id).map(|(_, v)| v.as_slice()) } /// Set a parameter to a varint value, replacing any existing entry. + /// + /// Panics past [`VarInt::MAX_QUIC`], which no parameter we set comes near. pub fn set_varint(&mut self, id: u64, value: u64) { let mut buf = Vec::new(); - // Infallible: writing into a Vec never runs short. - value.encode(&mut buf, Version::Lite05).expect("varint encode into Vec"); - self.0.insert(id, buf); + Encoder::new(&mut buf, Form::Quic) + .varint(value.into()) + .expect("parameter varint in range"); + self.set_bytes(id, buf); } /// Decode a parameter as a single varint, if present. Errors if trailing bytes remain. pub fn get_varint(&self, id: u64) -> Result, DecodeError> { - let Some(mut bytes) = self.0.get(&id).map(Vec::as_slice) else { + let Some(bytes) = self.get_bytes(id) else { return Ok(None); }; - let value = u64::decode(&mut bytes, Version::Lite05)?; - if !bytes.is_empty() { + let mut r = Decoder::new(bytes, Form::Quic); + let value = r.varint()?.into_inner(); + if !r.is_empty() { return Err(DecodeError::Long); } Ok(Some(value)) @@ -45,40 +53,39 @@ impl Parameters { } impl Decode for Parameters { - fn decode(mut r: &mut R, version: Version) -> Result { - let mut map = HashMap::new(); + fn decode(r: &mut Decoder<'_>, _: Version) -> Result { + let mut params = Self::default(); // I hate this encoding so much; let me encode my role and get on with my life. - let count = u64::decode(r, version)?; + let count = r.varint()?.into_inner(); if count > MAX_PARAMS { return Err(DecodeError::TooMany); } for _ in 0..count { - let kind = u64::decode(r, version)?; - if map.contains_key(&kind) { + let kind = r.varint()?.into_inner(); + if params.get_bytes(kind).is_some() { return Err(DecodeError::Duplicate); } - let data = Vec::::decode(&mut r, version)?; - map.insert(kind, data); + params.0.push((kind, r.bytes()?.to_vec())); } - Ok(Parameters(map)) + Ok(params) } } impl Encode for Parameters { - fn encode(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode(&self, w: &mut Encoder<'_>, _: Version) -> Result<(), EncodeError> { if self.0.len() as u64 > MAX_PARAMS { return Err(EncodeError::TooMany); } - self.0.len().encode(w, version)?; + w.varint(VarInt::from(self.0.len()))?; - for (kind, value) in self.0.iter() { - kind.encode(w, version)?; - value.encode(w, version)?; + for (kind, value) in &self.0 { + w.varint(VarInt::from(*kind))?; + w.bytes(value)?; } Ok(()) diff --git a/rs/moq-net/src/lite/probe.rs b/rs/moq-net/src/lite/probe.rs index dc64fb5b64..59347600dc 100644 --- a/rs/moq-net/src/lite/probe.rs +++ b/rs/moq-net/src/lite/probe.rs @@ -13,7 +13,7 @@ pub struct Probe { } impl Message for Probe { - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { match version { Version::Lite01 | Version::Lite02 => { return Err(DecodeError::Version); @@ -23,13 +23,13 @@ impl Message for Probe { // 0 means unknown, the same as RTT below. A publisher whose transport // exposes no congestion controller reports the RTT half alone. - let bitrate = match u64::decode(r, version)? { + let bitrate = match r.varint()?.into_inner() { 0 => None, v => Some(v), }; let rtt = match version.has_probe_rtt() { false => None, - true => match u64::decode(r, version)? { + true => match r.varint()?.into_inner() { 0 => None, v => Some(v), }, @@ -38,7 +38,7 @@ impl Message for Probe { Ok(Self { bitrate, rtt }) } - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { match version { Version::Lite01 | Version::Lite02 => { return Err(EncodeError::Version); @@ -48,10 +48,10 @@ impl Message for Probe { // 0 means unknown; round Some(0) up to 1. let wire = self.bitrate.map(|v| v.max(1)).unwrap_or(0); - wire.encode(w, version)?; + w.varint(VarInt::from(wire))?; if version.has_probe_rtt() { let wire = self.rtt.map(|v| v.max(1)).unwrap_or(0); - wire.encode(w, version)?; + w.varint(VarInt::from(wire))?; } Ok(()) } @@ -62,10 +62,11 @@ mod tests { use super::*; fn round_trip(msg: &Probe, version: Version) -> Probe { - let mut buf = bytes::BytesMut::new(); - msg.encode(&mut buf, version).unwrap(); + let mut buf = Vec::new(); + msg.encode(&mut Encoder::new(&mut buf, version.into()), version) + .unwrap(); let mut slice = &buf[..]; - let got = Probe::decode(&mut slice, version).unwrap(); + let got = crate::coding::decode_buf(&mut slice, version, Probe::decode).unwrap(); assert!(bytes::Buf::remaining(&slice) == 0, "trailing bytes after decode"); got } diff --git a/rs/moq-net/src/lite/publisher.rs b/rs/moq-net/src/lite/publisher.rs index 5ce1eaefc8..bd8c428237 100644 --- a/rs/moq-net/src/lite/publisher.rs +++ b/rs/moq-net/src/lite/publisher.rs @@ -1693,7 +1693,7 @@ mod announce_test { fn take_ok(&mut self) -> lite::AnnounceOk { let buf = self.pending(); let mut slice = &buf[..]; - let ok = lite::AnnounceOk::decode(&mut slice, VERSION).expect("announce ok"); + let ok = crate::coding::decode_buf(&mut slice, VERSION, lite::AnnounceOk::decode).expect("announce ok"); self.cursor += buf.len() - slice.len(); ok } @@ -1705,7 +1705,7 @@ mod announce_test { let mut msgs = Vec::new(); while !slice.is_empty() { msgs.push( - lite::AnnounceBroadcast::decode(&mut slice, VERSION) + crate::coding::decode_buf(&mut slice, VERSION, lite::AnnounceBroadcast::decode) .expect("announce message") .into_owned(), ); @@ -1938,7 +1938,7 @@ fn buffer_frame_info( if timescale.is_some() { buffer_zigzag_delta(writer, timestamp.value(), prev_ts)?; } - writer.buffer(&size)?; + writer.buffer(&crate::coding::VarInt::from(size))?; Ok(()) } @@ -1952,8 +1952,7 @@ fn buffer_zigzag_delta( let delta: i64 = (curr as i128 - *prev as i128) .try_into() .map_err(|_| Error::BoundsExceeded(crate::coding::BoundsExceeded))?; - let zz = crate::coding::VarInt::from_zigzag(delta).map_err(crate::coding::EncodeError::from)?; - writer.buffer(&zz)?; + writer.buffer(&crate::coding::VarInt::from_zigzag(delta))?; *prev = curr; Ok(()) } @@ -3504,7 +3503,7 @@ mod tests { let mut slice = bytes; let mut out = Vec::new(); while bytes::Buf::remaining(&slice) > 0 { - out.push(lite::Probe::decode(&mut slice, version).unwrap()); + out.push(crate::coding::decode_buf(&mut slice, version, lite::Probe::decode).unwrap()); } out } diff --git a/rs/moq-net/src/lite/setup.rs b/rs/moq-net/src/lite/setup.rs index 51225882cd..b20631ea5f 100644 --- a/rs/moq-net/src/lite/setup.rs +++ b/rs/moq-net/src/lite/setup.rs @@ -185,7 +185,7 @@ pub struct Setup { } impl Message for Setup { - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { if !version.has_setup_stream() { return Err(DecodeError::Version); } @@ -218,7 +218,7 @@ impl Message for Setup { }) } - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { if !version.has_setup_stream() { return Err(EncodeError::Version); } @@ -300,10 +300,11 @@ mod tests { use super::*; fn round_trip(msg: &Setup) -> Setup { - let mut buf = bytes::BytesMut::new(); - msg.encode(&mut buf, Version::Lite05).unwrap(); + let mut buf = Vec::new(); + msg.encode(&mut Encoder::new(&mut buf, Version::Lite05.into()), Version::Lite05) + .unwrap(); let mut slice = &buf[..]; - let got = Setup::decode(&mut slice, Version::Lite05).unwrap(); + let got = crate::coding::decode_buf(&mut slice, Version::Lite05, Setup::decode).unwrap(); assert!(bytes::Buf::remaining(&slice) == 0, "trailing bytes after decode"); got } @@ -391,14 +392,18 @@ mod tests { let version = Version::Lite05; let mut params = Parameters::default(); params.set_varint(super::PARAM_HOP, 0); - let mut body = bytes::BytesMut::new(); - params.encode(&mut body, version).unwrap(); + let mut body = Vec::new(); + params + .encode(&mut Encoder::new(&mut body, version.into()), version) + .unwrap(); // Frame the body with the Message Length prefix `Setup::decode` expects. - let mut buf = bytes::BytesMut::new(); - (body.len() as u64).encode(&mut buf, version).unwrap(); + let mut buf = Vec::new(); + Encoder::new(&mut buf, version.into()) + .varint(VarInt::from(body.len())) + .unwrap(); buf.extend_from_slice(&body); let mut slice = &buf[..]; - let got = Setup::decode(&mut slice, version).unwrap(); + let got = crate::coding::decode_buf(&mut slice, version, Setup::decode).unwrap(); assert_eq!(got.hop, None); } @@ -432,14 +437,18 @@ mod tests { let mut params = Parameters::default(); params.set_varint(PARAM_PROBE, 99); let mut body = Vec::new(); - params.encode(&mut body, Version::Lite05).unwrap(); - - let mut buf = bytes::BytesMut::new(); - body.len().encode(&mut buf, Version::Lite05).unwrap(); + params + .encode(&mut Encoder::new(&mut body, Version::Lite05.into()), Version::Lite05) + .unwrap(); + + let mut buf = Vec::new(); + Encoder::new(&mut buf, Version::Lite05.into()) + .varint(VarInt::from(body.len())) + .unwrap(); buf.extend_from_slice(&body); let mut slice = &buf[..]; - let got = Setup::decode(&mut slice, Version::Lite05).unwrap(); + let got = crate::coding::decode_buf(&mut slice, Version::Lite05, Setup::decode).unwrap(); assert_eq!(got.probe, ProbeLevel::Increase); } @@ -462,14 +471,18 @@ mod tests { let mut params = Parameters::default(); params.set_varint(PARAM_ROLE, code); let mut body = Vec::new(); - params.encode(&mut body, Version::Lite05).unwrap(); - - let mut buf = bytes::BytesMut::new(); - body.len().encode(&mut buf, Version::Lite05).unwrap(); + params + .encode(&mut Encoder::new(&mut body, Version::Lite05.into()), Version::Lite05) + .unwrap(); + + let mut buf = Vec::new(); + Encoder::new(&mut buf, Version::Lite05.into()) + .varint(VarInt::from(body.len())) + .unwrap(); buf.extend_from_slice(&body); let mut slice = &buf[..]; - let got = Setup::decode(&mut slice, Version::Lite05).unwrap(); + let got = crate::coding::decode_buf(&mut slice, Version::Lite05, Setup::decode).unwrap(); assert_eq!(got.role, None, "role code {code} should decode as bidirectional"); } } @@ -477,9 +490,9 @@ mod tests { #[test] fn rejects_before_lite05() { let msg = Setup::default(); - let mut buf = bytes::BytesMut::new(); + let mut buf = Vec::new(); assert!(matches!( - msg.encode(&mut buf, Version::Lite04), + msg.encode(&mut Encoder::new(&mut buf, Version::Lite04.into()), Version::Lite04), Err(EncodeError::Version) )); } @@ -492,15 +505,19 @@ mod tests { params.set_bytes(0xbeef, b"whatever".to_vec()); let mut body = Vec::new(); - params.encode(&mut body, Version::Lite05).unwrap(); + params + .encode(&mut Encoder::new(&mut body, Version::Lite05.into()), Version::Lite05) + .unwrap(); // Wrap with the message size prefix the Message impl expects. - let mut buf = bytes::BytesMut::new(); - body.len().encode(&mut buf, Version::Lite05).unwrap(); + let mut buf = Vec::new(); + Encoder::new(&mut buf, Version::Lite05.into()) + .varint(VarInt::from(body.len())) + .unwrap(); buf.extend_from_slice(&body); let mut slice = &buf[..]; - let got = Setup::decode(&mut slice, Version::Lite05).unwrap(); + let got = crate::coding::decode_buf(&mut slice, Version::Lite05, Setup::decode).unwrap(); assert_eq!(got.path.as_deref(), Some("/foo")); } } diff --git a/rs/moq-net/src/lite/stream.rs b/rs/moq-net/src/lite/stream.rs index 519ea3ef9b..8c0f89ce53 100644 --- a/rs/moq-net/src/lite/stream.rs +++ b/rs/moq-net/src/lite/stream.rs @@ -17,16 +17,16 @@ pub enum ControlType { } impl Decode for ControlType { - fn decode(r: &mut R, version: Version) -> Result { - let t = u64::decode(r, version)?; + fn decode(r: &mut Decoder<'_>, _: Version) -> Result { + let t = r.varint()?.into_inner(); t.try_into().map_err(|_| DecodeError::InvalidValue) } } impl Encode for ControlType { - fn encode(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode(&self, w: &mut Encoder<'_>, _: Version) -> Result<(), EncodeError> { let v: u64 = (*self).into(); - v.encode(w, version)?; + w.varint(VarInt::from(v))?; Ok(()) } } @@ -41,16 +41,16 @@ pub enum DataType { } impl Decode for DataType { - fn decode(r: &mut R, version: Version) -> Result { - let t = u64::decode(r, version)?; + fn decode(r: &mut Decoder<'_>, _: Version) -> Result { + let t = r.varint()?.into_inner(); t.try_into().map_err(|_| DecodeError::InvalidValue) } } impl Encode for DataType { - fn encode(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode(&self, w: &mut Encoder<'_>, _: Version) -> Result<(), EncodeError> { let v: u64 = (*self).into(); - v.encode(w, version)?; + w.varint(VarInt::from(v))?; Ok(()) } } diff --git a/rs/moq-net/src/lite/subscribe.rs b/rs/moq-net/src/lite/subscribe.rs index 4656e94521..0b8c730b66 100644 --- a/rs/moq-net/src/lite/subscribe.rs +++ b/rs/moq-net/src/lite/subscribe.rs @@ -2,7 +2,7 @@ use std::borrow::Cow; use crate::{ Path, - coding::{Decode, DecodeError, Encode, EncodeError, Sizer}, + coding::{Decode, DecodeError, Decoder, Encode, EncodeError, Encoder, VarInt}, }; use super::{Message, Version}; @@ -44,19 +44,19 @@ impl Version { } impl Message for Subscribe<'_> { - fn decode_msg(r: &mut R, version: Version) -> Result { - let id = u64::decode(r, version)?; + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { + let id = r.varint()?.into_inner(); let broadcast = Path::decode(r, version)?; - let track = Cow::::decode(r, version)?; - let priority = u8::decode(r, version)?; + let track = Cow::Owned(r.string()?); + let priority = r.u8()?; let (max_age, start_group, end_group) = match version { Version::Lite01 | Version::Lite02 => (std::time::Duration::ZERO, None, None), _ => { skip_group_order(r, version)?; - let max_age = std::time::Duration::decode(r, version)?; + let max_age = std::time::Duration::from_millis(r.varint()?.into_inner()); let start_group = decode_start_group(r, version)?; - let end_group = Option::::decode(r, version)?; + let end_group = r.varint_opt()?; (max_age, start_group, end_group) } }; @@ -77,19 +77,19 @@ impl Message for Subscribe<'_> { }) } - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { - self.id.encode(w, version)?; + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { + w.varint(VarInt::from(self.id))?; self.broadcast.encode(w, version)?; - self.track.encode(w, version)?; - self.priority.encode(w, version)?; + w.string(&self.track)?; + w.u8(self.priority); match version { Version::Lite01 | Version::Lite02 => {} _ => { pad_group_order(w, version)?; - self.max_age.encode(w, version)?; + w.varint(VarInt::try_from(self.max_age.as_millis())?)?; encode_start_group(w, version, self.start_group)?; - self.end_group.encode(w, version)?; + w.varint_opt(self.end_group)?; } } @@ -110,17 +110,17 @@ impl Message for Subscribe<'_> { /// /// The value is ignored: group order is fixed, so a peer that still sets it gets the /// same newest-first delivery as one that doesn't. -pub(super) fn skip_group_order(r: &mut R, version: Version) -> Result<(), DecodeError> { +pub(super) fn skip_group_order(r: &mut Decoder<'_>, version: Version) -> Result<(), DecodeError> { if version.has_group_order() { - u8::decode(r, version)?; + r.u8()?; } Ok(()) } /// Write the retired `Ordered` byte as 0, keeping a deployed version's field offsets. -pub(super) fn pad_group_order(w: &mut W, version: Version) -> Result<(), EncodeError> { +pub(super) fn pad_group_order(w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { if version.has_group_order() { - 0u8.encode(w, version)?; + w.u8(0u8); } Ok(()) } @@ -131,11 +131,11 @@ pub(super) fn pad_group_order(w: &mut W, version: Version) -> /// `Some(0)`. Pre-06 wires encode the sequence + 1, with 0 meaning the latest group /// (`None`). Callers canonicalize with [`canonical_start_group`] once the frame bounds /// are known. -fn decode_start_group(r: &mut R, version: Version) -> Result, DecodeError> { +fn decode_start_group(r: &mut Decoder<'_>, version: Version) -> Result, DecodeError> { if version.resolves_start() { - return Ok(Some(u64::decode(r, version)?)); + return Ok(Some(r.varint()?.into_inner())); } - Option::::decode(r, version) + r.varint_opt() } /// Canonicalize a decoded floor: a lite-06 `Group Start` of 0 with no frame offset is the @@ -155,15 +155,11 @@ fn canonical_start_group(version: Version, start_group: Option, start_frame /// `Some(0)` are the same absence of a constraint), while a pre-06 wire gets `Some(0)` /// folded back to absent. On those wires an explicit group 0 means "replay from the /// beginning", which is not what a vacuous floor asks for. -fn encode_start_group( - w: &mut W, - version: Version, - start_group: Option, -) -> Result<(), EncodeError> { +fn encode_start_group(w: &mut Encoder<'_>, version: Version, start_group: Option) -> Result<(), EncodeError> { if version.resolves_start() { - return start_group.unwrap_or(0).encode(w, version); + return w.varint(VarInt::from(start_group.unwrap_or(0))); } - start_group.filter(|&group| group > 0).encode(w, version) + w.varint_opt(start_group.filter(|&group| group > 0)) } /// Decode the trailing `Frame Start` / `Frame End` pair shared by SUBSCRIBE, @@ -172,8 +168,8 @@ fn encode_start_group( /// Older versions carry no such fields, so they decode as the whole group. A frame bound /// without the group bound it qualifies is a protocol violation: frames are numbered per /// group, so there is nothing to count from. -fn decode_frame_bounds( - r: &mut R, +fn decode_frame_bounds( + r: &mut Decoder<'_>, version: Version, start_group: Option, end_group: Option, @@ -182,8 +178,8 @@ fn decode_frame_bounds( return Ok((0, None)); } - let start_frame = u64::decode(r, version)?; - let end_frame = Option::::decode(r, version)?; + let start_frame = r.varint()?.into_inner(); + let end_frame = r.varint_opt()?; if (start_frame != 0 && start_group.is_none()) || (end_frame.is_some() && end_group.is_none()) { return Err(DecodeError::InvalidSubscribeLocation); @@ -193,8 +189,8 @@ fn decode_frame_bounds( } /// Encode the trailing `Frame Start` / `Frame End` pair, a no-op before lite-06. -fn encode_frame_bounds( - w: &mut W, +fn encode_frame_bounds( + w: &mut Encoder<'_>, version: Version, start_group: Option, start_frame: u64, @@ -214,8 +210,8 @@ fn encode_frame_bounds( return Ok(()); } - start_frame.encode(w, version)?; - end_frame.encode(w, version) + w.varint(VarInt::from(start_frame))?; + w.varint_opt(end_frame) } /// Publisher's acknowledgement on the Subscribe Stream for drafts 01-04. @@ -232,30 +228,30 @@ pub struct SubscribeOk { } impl Message for SubscribeOk { - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { match version { Version::Lite01 => { - self.priority.encode(w, version)?; + w.u8(self.priority); } Version::Lite02 => {} // Lite05+ never sends SUBSCRIBE_OK, but keep the field layout matching // Lite03/04 so a stray future use stays well-formed. _ => { - self.priority.encode(w, version)?; + w.u8(self.priority); pad_group_order(w, version)?; - self.max_age.encode(w, version)?; - self.start_group.encode(w, version)?; - self.end_group.encode(w, version)?; + w.varint(VarInt::try_from(self.max_age.as_millis())?)?; + w.varint_opt(self.start_group)?; + w.varint_opt(self.end_group)?; } } Ok(()) } - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { match version { Version::Lite01 => Ok(Self { - priority: u8::decode(r, version)?, + priority: r.u8()?, max_age: std::time::Duration::ZERO, start_group: None, end_group: None, @@ -267,11 +263,11 @@ impl Message for SubscribeOk { end_group: None, }), _ => { - let priority = u8::decode(r, version)?; + let priority = r.u8()?; skip_group_order(r, version)?; - let max_age = std::time::Duration::decode(r, version)?; - let start_group = Option::::decode(r, version)?; - let end_group = Option::::decode(r, version)?; + let max_age = std::time::Duration::from_millis(r.varint()?.into_inner()); + let start_group = r.varint_opt()?; + let end_group = r.varint_opt()?; Ok(Self { priority, @@ -298,20 +294,20 @@ pub struct SubscribeStart { } impl Message for SubscribeStart { - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { if !version.has_track_stream() { return Err(DecodeError::Version); } Ok(Self { - group: u64::decode(r, version)?, + group: r.varint()?.into_inner(), }) } - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { if !version.has_track_stream() { return Err(EncodeError::Version); } - self.group.encode(w, version) + w.varint(VarInt::from(self.group)) } } @@ -328,25 +324,25 @@ pub struct SubscribeEnd { } impl Message for SubscribeEnd { - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { if !version.has_track_stream() { return Err(DecodeError::Version); } - let group = u64::decode(r, version)?; + let group = r.varint()?.into_inner(); let streams = match version.has_stream_count() { - true => u64::decode(r, version)?, + true => r.varint()?.into_inner(), false => 0, }; Ok(Self { group, streams }) } - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { if !version.has_track_stream() { return Err(EncodeError::Version); } - self.group.encode(w, version)?; + w.varint(VarInt::from(self.group))?; if version.has_stream_count() { - self.streams.encode(w, version)?; + w.varint(VarInt::from(self.streams))?; } Ok(()) } @@ -369,7 +365,7 @@ pub struct SubscribeUpdate { } impl Message for SubscribeUpdate { - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { match version { Version::Lite01 | Version::Lite02 => { return Err(DecodeError::Version); @@ -377,11 +373,11 @@ impl Message for SubscribeUpdate { _ => {} } - let priority = u8::decode(r, version)?; + let priority = r.u8()?; skip_group_order(r, version)?; - let max_age = std::time::Duration::decode(r, version)?; + let max_age = std::time::Duration::from_millis(r.varint()?.into_inner()); let start_group = decode_start_group(r, version)?; - let end_group = match u64::decode(r, version)? { + let end_group = match r.varint()?.into_inner() { 0 => None, group => Some(group - 1), }; @@ -399,7 +395,7 @@ impl Message for SubscribeUpdate { }) } - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { match version { Version::Lite01 | Version::Lite02 => { return Err(EncodeError::Version); @@ -407,19 +403,13 @@ impl Message for SubscribeUpdate { _ => {} } - self.priority.encode(w, version)?; + w.u8(self.priority); pad_group_order(w, version)?; - self.max_age.encode(w, version)?; + w.varint(VarInt::try_from(self.max_age.as_millis())?)?; encode_start_group(w, version, self.start_group)?; - match self.end_group { - Some(end_group) => end_group - .checked_add(1) - .ok_or(EncodeError::TooLarge)? - .encode(w, version)?, - None => 0u64.encode(w, version)?, - } + w.varint_opt(self.end_group)?; encode_frame_bounds( w, @@ -454,7 +444,7 @@ pub struct SubscribeDrop { } impl Message for SubscribeDrop { - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { match version { Version::Lite01 | Version::Lite02 => { return Err(DecodeError::Version); @@ -464,13 +454,13 @@ impl Message for SubscribeDrop { } Ok(Self { - start: u64::decode(r, version)?, - end: u64::decode(r, version)?, - error: u64::decode(r, version)?, + start: r.varint()?.into_inner(), + end: r.varint()?.into_inner(), + error: r.varint()?.into_inner(), }) } - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { match version { Version::Lite01 | Version::Lite02 => { return Err(EncodeError::Version); @@ -479,9 +469,9 @@ impl Message for SubscribeDrop { _ => {} } - self.start.encode(w, version)?; - self.end.encode(w, version)?; - self.error.encode(w, version)?; + w.varint(VarInt::from(self.start))?; + w.varint(VarInt::from(self.end))?; + w.varint(VarInt::from(self.error))?; Ok(()) } @@ -504,29 +494,16 @@ pub enum SubscribeResponse { } /// Write a `type` varint followed by the size-prefixed message body. -fn encode_typed( - w: &mut W, - typ: u64, - msg: &M, - version: Version, -) -> Result<(), EncodeError> { - typ.encode(w, version)?; - let mut sizer = Sizer::default(); - msg.encode_msg(&mut sizer, version)?; - sizer.size.encode(w, version)?; - msg.encode_msg(w, version) +fn encode_typed(w: &mut Encoder<'_>, typ: u64, msg: &M, version: Version) -> Result<(), EncodeError> { + w.varint(VarInt::from(typ))?; + msg.encode(w, version) } impl Encode for SubscribeResponse { - fn encode(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { match version { Version::Lite01 | Version::Lite02 => match self { - Self::Ok(ok) => { - let mut sizer = Sizer::default(); - Message::encode_msg(ok, &mut sizer, version)?; - sizer.size.encode(w, version)?; - Message::encode_msg(ok, w, version)?; - } + Self::Ok(ok) => ok.encode(w, version)?, _ => return Err(EncodeError::Version), }, Version::Lite03 | Version::Lite04 => match self { @@ -548,11 +525,11 @@ impl Encode for SubscribeResponse { } impl Decode for SubscribeResponse { - fn decode(buf: &mut B, version: Version) -> Result { + fn decode(buf: &mut Decoder<'_>, version: Version) -> Result { match version { Version::Lite01 | Version::Lite02 => Ok(Self::Ok(SubscribeOk::decode(buf, version)?)), Version::Lite03 | Version::Lite04 => { - let typ = u64::decode(buf, version)?; + let typ = buf.varint()?.into_inner(); match typ { 0 => Ok(Self::Ok(SubscribeOk::decode(buf, version)?)), 1 => Ok(Self::Drop(SubscribeDrop::decode(buf, version)?)), @@ -560,7 +537,7 @@ impl Decode for SubscribeResponse { } } _ => { - let typ = u64::decode(buf, version)?; + let typ = buf.varint()?.into_inner(); match typ { 0 => Ok(Self::Start(SubscribeStart::decode(buf, version)?)), 1 => Ok(Self::End(SubscribeEnd::decode(buf, version)?)), @@ -580,9 +557,10 @@ mod test { fn subscribe_start_roundtrips_on_lite05() { let resp = SubscribeResponse::Start(SubscribeStart { group: 42 }); let mut buf = Vec::new(); - resp.encode(&mut buf, Version::Lite05).unwrap(); + resp.encode(&mut Encoder::new(&mut buf, Version::Lite05.into()), Version::Lite05) + .unwrap(); let mut slice = buf.as_slice(); - match SubscribeResponse::decode(&mut slice, Version::Lite05).unwrap() { + match crate::coding::decode_buf(&mut slice, Version::Lite05, SubscribeResponse::decode).unwrap() { SubscribeResponse::Start(start) => assert_eq!(start.group, 42), other => panic!("expected Start, got {other:?}"), } @@ -592,11 +570,12 @@ mod test { fn subscribe_end_roundtrips_on_lite05() { let resp = SubscribeResponse::End(SubscribeEnd { group: 7, streams: 3 }); let mut buf = Vec::new(); - resp.encode(&mut buf, Version::Lite05).unwrap(); + resp.encode(&mut Encoder::new(&mut buf, Version::Lite05.into()), Version::Lite05) + .unwrap(); // Type, length, group: no stream count before lite-07. assert_eq!(buf, [1, 1, 7]); let mut slice = buf.as_slice(); - match SubscribeResponse::decode(&mut slice, Version::Lite05).unwrap() { + match crate::coding::decode_buf(&mut slice, Version::Lite05, SubscribeResponse::decode).unwrap() { SubscribeResponse::End(end) => assert_eq!((end.group, end.streams), (7, 0)), other => panic!("expected End, got {other:?}"), } @@ -606,10 +585,11 @@ mod test { fn subscribe_end_carries_the_stream_count_on_lite07() { let resp = SubscribeResponse::End(SubscribeEnd { group: 7, streams: 3 }); let mut buf = Vec::new(); - resp.encode(&mut buf, Version::Lite07).unwrap(); + resp.encode(&mut Encoder::new(&mut buf, Version::Lite07.into()), Version::Lite07) + .unwrap(); assert_eq!(buf, [1, 2, 7, 3]); let mut slice = buf.as_slice(); - match SubscribeResponse::decode(&mut slice, Version::Lite07).unwrap() { + match crate::coding::decode_buf(&mut slice, Version::Lite07, SubscribeResponse::decode).unwrap() { SubscribeResponse::End(end) => assert_eq!((end.group, end.streams), (7, 3)), other => panic!("expected End, got {other:?}"), } @@ -624,15 +604,16 @@ mod test { }); let mut buf = Vec::new(); assert!(matches!( - resp.encode(&mut buf, Version::Lite07), + resp.encode(&mut Encoder::new(&mut buf, Version::Lite07.into()), Version::Lite07), Err(EncodeError::Version) )); // A lite-06 DROP is an unknown response type on lite-07. let mut buf = Vec::new(); - resp.encode(&mut buf, Version::Lite06).unwrap(); + resp.encode(&mut Encoder::new(&mut buf, Version::Lite06.into()), Version::Lite06) + .unwrap(); assert!(matches!( - SubscribeResponse::decode(&mut buf.as_slice(), Version::Lite07), + crate::coding::decode_buf(&mut buf.as_slice(), Version::Lite07, SubscribeResponse::decode), Err(DecodeError::InvalidMessage(2)) )); } @@ -645,12 +626,13 @@ mod test { error: 0, }); let mut buf = Vec::new(); - resp.encode(&mut buf, Version::Lite05).unwrap(); + resp.encode(&mut Encoder::new(&mut buf, Version::Lite05.into()), Version::Lite05) + .unwrap(); // Type discriminator is the first varint; on Lite05 DROP is 0x2. assert_eq!(buf[0], 2); let mut slice = buf.as_slice(); - match SubscribeResponse::decode(&mut slice, Version::Lite05).unwrap() { + match crate::coding::decode_buf(&mut slice, Version::Lite05, SubscribeResponse::decode).unwrap() { SubscribeResponse::Drop(drop) => assert_eq!((drop.start, drop.end), (1, 3)), other => panic!("expected Drop, got {other:?}"), } @@ -664,7 +646,8 @@ mod test { error: 0, }); let mut buf = Vec::new(); - resp.encode(&mut buf, Version::Lite04).unwrap(); + resp.encode(&mut Encoder::new(&mut buf, Version::Lite04.into()), Version::Lite04) + .unwrap(); assert_eq!(buf[0], 1); } @@ -686,8 +669,9 @@ mod test { fn subscribe_frame_bounds_roundtrip() { let msg = subscribe_sample(); let mut buf = Vec::new(); - msg.encode_msg(&mut buf, Version::Lite06).unwrap(); - let got = Subscribe::decode_msg(&mut buf.as_slice(), Version::Lite06).unwrap(); + msg.encode_msg(&mut Encoder::new(&mut buf, Version::Lite06.into()), Version::Lite06) + .unwrap(); + let got = crate::coding::decode_buf(&mut buf.as_slice(), Version::Lite06, Subscribe::decode_msg).unwrap(); assert_eq!((got.start_group, got.start_frame), (Some(7), 4)); assert_eq!((got.end_group, got.end_frame), (Some(9), Some(2))); } @@ -703,9 +687,11 @@ mod test { msg.end_frame = None; let mut lite05 = Vec::new(); - msg.encode_msg(&mut lite05, Version::Lite05).unwrap(); + msg.encode_msg(&mut Encoder::new(&mut lite05, Version::Lite05.into()), Version::Lite05) + .unwrap(); let mut lite06 = Vec::new(); - msg.encode_msg(&mut lite06, Version::Lite06).unwrap(); + msg.encode_msg(&mut Encoder::new(&mut lite06, Version::Lite06.into()), Version::Lite06) + .unwrap(); // The two layouts diverge in exactly one place: the retired byte lite-05 still // reserves. A deployed peer's field offsets depend on it being there and zero. @@ -722,7 +708,7 @@ mod test { assert_eq!(&lite06[..spliced.len()], &spliced[..]); assert_eq!(&lite06[spliced.len()..], &[0, 0]); - let got = Subscribe::decode_msg(&mut lite05.as_slice(), Version::Lite05).unwrap(); + let got = crate::coding::decode_buf(&mut lite05.as_slice(), Version::Lite05, Subscribe::decode_msg).unwrap(); assert_eq!((got.start_frame, got.end_frame), (0, None)); } @@ -736,12 +722,14 @@ mod test { msg.end_frame = None; let mut lite05 = Vec::new(); - msg.encode_msg(&mut lite05, Version::Lite05).unwrap(); + msg.encode_msg(&mut Encoder::new(&mut lite05, Version::Lite05.into()), Version::Lite05) + .unwrap(); let mut lite06 = Vec::new(); - msg.encode_msg(&mut lite06, Version::Lite06).unwrap(); + msg.encode_msg(&mut Encoder::new(&mut lite06, Version::Lite06.into()), Version::Lite06) + .unwrap(); - let on05 = Subscribe::decode_msg(&mut lite05.as_slice(), Version::Lite05).unwrap(); - let on06 = Subscribe::decode_msg(&mut lite06.as_slice(), Version::Lite06).unwrap(); + let on05 = crate::coding::decode_buf(&mut lite05.as_slice(), Version::Lite05, Subscribe::decode_msg).unwrap(); + let on06 = crate::coding::decode_buf(&mut lite06.as_slice(), Version::Lite06, Subscribe::decode_msg).unwrap(); assert_eq!(on05.start_group, Some(7)); assert_eq!(on06.start_group, Some(7)); // The raw byte differs: 7 on the wire, not 7 + 1. @@ -751,18 +739,21 @@ mod test { // byte-identical on the wire, and canonicalized to absent on decode. msg.start_group = None; let mut absent = Vec::new(); - msg.encode_msg(&mut absent, Version::Lite06).unwrap(); + msg.encode_msg(&mut Encoder::new(&mut absent, Version::Lite06.into()), Version::Lite06) + .unwrap(); msg.start_group = Some(0); let mut zero = Vec::new(); - msg.encode_msg(&mut zero, Version::Lite06).unwrap(); + msg.encode_msg(&mut Encoder::new(&mut zero, Version::Lite06.into()), Version::Lite06) + .unwrap(); assert_eq!(absent, zero); - let got = Subscribe::decode_msg(&mut zero.as_slice(), Version::Lite06).unwrap(); + let got = crate::coding::decode_buf(&mut zero.as_slice(), Version::Lite06, Subscribe::decode_msg).unwrap(); assert_eq!(got.start_group, None); // On the pre-06 wire the vacuous floor folds to absent (the latest group). let mut folded = Vec::new(); - msg.encode_msg(&mut folded, Version::Lite05).unwrap(); - let got = Subscribe::decode_msg(&mut folded.as_slice(), Version::Lite05).unwrap(); + msg.encode_msg(&mut Encoder::new(&mut folded, Version::Lite05.into()), Version::Lite05) + .unwrap(); + let got = crate::coding::decode_buf(&mut folded.as_slice(), Version::Lite05, Subscribe::decode_msg).unwrap(); assert_eq!(got.start_group, None); } @@ -775,8 +766,9 @@ mod test { msg.start_frame = 4; let mut buf = Vec::new(); - msg.encode_msg(&mut buf, Version::Lite06).unwrap(); - let got = Subscribe::decode_msg(&mut buf.as_slice(), Version::Lite06).unwrap(); + msg.encode_msg(&mut Encoder::new(&mut buf, Version::Lite06.into()), Version::Lite06) + .unwrap(); + let got = crate::coding::decode_buf(&mut buf.as_slice(), Version::Lite06, Subscribe::decode_msg).unwrap(); assert_eq!((got.start_group, got.start_frame), (Some(0), 4)); } @@ -784,7 +776,11 @@ mod test { #[test] fn subscribe_frame_bounds_rejected_before_lite06() { let mut buf = Vec::new(); - assert!(subscribe_sample().encode_msg(&mut buf, Version::Lite05).is_err()); + assert!( + subscribe_sample() + .encode_msg(&mut Encoder::new(&mut buf, Version::Lite05.into()), Version::Lite05) + .is_err() + ); } /// Frames are numbered per group, so a frame bound without its group bound has @@ -799,14 +795,14 @@ mod test { let mut buf = Vec::new(); assert!(matches!( - msg.encode_msg(&mut buf, Version::Lite06), + msg.encode_msg(&mut Encoder::new(&mut buf, Version::Lite06.into()), Version::Lite06), Err(EncodeError::InvalidState) )); msg.start_frame = 0; msg.end_frame = Some(7); assert!(matches!( - msg.encode_msg(&mut buf, Version::Lite06), + msg.encode_msg(&mut Encoder::new(&mut buf, Version::Lite06.into()), Version::Lite06), Err(EncodeError::InvalidState) )); } @@ -820,6 +816,9 @@ mod test { end_group: None, }); let mut buf = Vec::new(); - assert!(resp.encode(&mut buf, Version::Lite05).is_err()); + assert!( + resp.encode(&mut Encoder::new(&mut buf, Version::Lite05.into()), Version::Lite05) + .is_err() + ); } } diff --git a/rs/moq-net/src/lite/subscriber.rs b/rs/moq-net/src/lite/subscriber.rs index fa9033b837..964a2e38dd 100644 --- a/rs/moq-net/src/lite/subscriber.rs +++ b/rs/moq-net/src/lite/subscriber.rs @@ -9,7 +9,7 @@ use std::{ use crate::{ AsPath, Error, Path, PathOwned, Timescale, Timestamp, bandwidth, - coding::{Decode, Reader, Stream}, + coding::{Reader, Stream}, lite, track::{Position, Subscription}, }; @@ -440,8 +440,7 @@ impl Subscriber { /// Decode one datagram body and hand it to the matching subscription's producer. fn route_datagram(&self, payload: bytes::Bytes) -> Result<(), Error> { - let mut buf = payload; - let dg = lite::Datagram::decode(&mut buf, self.version)?; + let dg = lite::Datagram::decode(payload, self.version)?; // Write through the map rather than cloning the entry out: a `TrackEntry` clone // is a handful of atomic bumps on every datagram, and a producer held past its @@ -904,9 +903,10 @@ impl FrameIngest { }; } IngestPhase::Size { timestamp } => { - let Some(size) = ready!(reader.poll_decode_maybe::(&mut cx))? else { + let Some(size) = ready!(reader.poll_decode_maybe::(&mut cx))? else { return Poll::Ready(Ok(())); }; + let size = size.into_inner(); // `create_frame_owned` is the allocation chokepoint and rejects an // oversized `size` before allocating, so no pre-check is needed. No // wire timestamp (pre-lite-05) means local receive time. @@ -1484,12 +1484,18 @@ mod tests { let mut responses = Vec::new(); if started { lite::SubscribeResponse::Start(lite::SubscribeStart { group: 0 }) - .encode(&mut responses, version) + .encode( + &mut crate::coding::Encoder::new(&mut responses, version.into()), + version, + ) .unwrap(); } if clean { lite::SubscribeResponse::End(lite::SubscribeEnd { group: 0, streams: 0 }) - .encode(&mut responses, version) + .encode( + &mut crate::coding::Encoder::new(&mut responses, version.into()), + version, + ) .unwrap(); } responses @@ -1689,10 +1695,10 @@ mod tests { let writes = session.log.writes.lock().unwrap().clone(); let mut wire = writes.as_slice(); assert_eq!( - lite::ControlType::decode(&mut wire, VERSION).unwrap(), + crate::coding::decode_buf(&mut wire, VERSION, lite::ControlType::decode).unwrap(), lite::ControlType::Subscribe ); - let msg = lite::Subscribe::decode(&mut wire, VERSION).unwrap(); + let msg = crate::coding::decode_buf(&mut wire, VERSION, lite::Subscribe::decode).unwrap(); assert_eq!(msg.id, 0); assert_eq!(msg.track, "catalog.json"); assert!(wire.is_empty(), "a second SUBSCRIBE trailed the first"); @@ -1783,10 +1789,10 @@ mod tests { let wire = h.wire(); let mut wire = wire.as_slice(); assert_eq!( - lite::ControlType::decode(&mut wire, Version::Lite05).unwrap(), + crate::coding::decode_buf(&mut wire, Version::Lite05, lite::ControlType::decode).unwrap(), lite::ControlType::Subscribe ); - let msg = lite::Subscribe::decode(&mut wire, Version::Lite05).unwrap(); + let msg = crate::coding::decode_buf(&mut wire, Version::Lite05, lite::Subscribe::decode).unwrap(); // The group bounds survive; only the frame offsets are widened away. assert_eq!((msg.start_group, msg.end_group), (Some(5), Some(5))); assert_eq!((msg.start_frame, msg.end_frame), (0, None)); @@ -1813,10 +1819,10 @@ mod tests { let wire = h.wire(); let mut wire = wire.as_slice(); assert_eq!( - lite::ControlType::decode(&mut wire, Version::Lite06).unwrap(), + crate::coding::decode_buf(&mut wire, Version::Lite06, lite::ControlType::decode).unwrap(), lite::ControlType::Subscribe ); - let msg = lite::Subscribe::decode(&mut wire, Version::Lite06).unwrap(); + let msg = crate::coding::decode_buf(&mut wire, Version::Lite06, lite::Subscribe::decode).unwrap(); assert_eq!((msg.start_frame, msg.end_frame), (3, Some(7))); } @@ -2025,7 +2031,7 @@ mod tests { // SUBSCRIBE_UPDATE rides the subscribe stream with no control type ahead of it. let wire = h.wire(); let mut wire = &wire[established..]; - let msg = lite::SubscribeUpdate::decode(&mut wire, Version::Lite05).unwrap(); + let msg = crate::coding::decode_buf(&mut wire, Version::Lite05, lite::SubscribeUpdate::decode).unwrap(); assert_eq!((msg.start_group, msg.start_frame), (Some(5), 0)); assert_eq!((msg.end_group, msg.end_frame), (Some(5), None)); } @@ -2722,12 +2728,16 @@ mod tests { origin: crate::Hop::new(9).unwrap(), active: 1, } - .encode(&mut script, VERSION) + .encode(&mut crate::coding::Encoder::new(&mut script, VERSION.into()), VERSION) .unwrap(); // An unknown announce type with an empty body, which decodes as `Skipped`. script.extend([0x3f, 0x00]); - start("a").encode(&mut script, VERSION).unwrap(); - start("b").encode(&mut script, VERSION).unwrap(); + start("a") + .encode(&mut crate::coding::Encoder::new(&mut script, VERSION.into()), VERSION) + .unwrap(); + start("b") + .encode(&mut crate::coding::Encoder::new(&mut script, VERSION.into()), VERSION) + .unwrap(); let origin = origin::Config::new(crate::Hop::new(1).unwrap()).produce(); let consumer = origin.consume(); diff --git a/rs/moq-net/src/lite/track.rs b/rs/moq-net/src/lite/track.rs index 3e4fca45ca..f764dd7c27 100644 --- a/rs/moq-net/src/lite/track.rs +++ b/rs/moq-net/src/lite/track.rs @@ -2,7 +2,7 @@ use std::{borrow::Cow, time::Duration}; use crate::{ Path, Timescale, - coding::{Decode, DecodeError, Encode, EncodeError}, + coding::{Decode, DecodeError, Decoder, Encode, EncodeError, Encoder, VarInt}, }; use super::{Message, Version}; @@ -21,24 +21,24 @@ pub struct Track<'a> { } impl Message for Track<'_> { - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { if !version.has_track_stream() { return Err(DecodeError::Version); } let broadcast = Path::decode(r, version)?; - let track = Cow::::decode(r, version)?; + let track = Cow::Owned(r.string()?); Ok(Self { broadcast, track }) } - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { if !version.has_track_stream() { return Err(EncodeError::Version); } self.broadcast.encode(w, version)?; - self.track.encode(w, version)?; + w.string(&self.track)?; Ok(()) } } @@ -61,19 +61,19 @@ pub struct TrackInfo { } impl Message for TrackInfo { - fn decode_msg(r: &mut R, version: Version) -> Result { + fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { if !version.has_track_stream() { return Err(DecodeError::Version); } - let priority = u8::decode(r, version)?; + let priority = r.u8()?; super::subscribe::skip_group_order(r, version)?; - let encoded = u64::decode(r, version)?; + let encoded = r.varint()?.into_inner(); let max_age = match version { Version::Lite05 | Version::Lite06 => (encoded < LEGACY_UNLIMITED).then(|| Duration::from_millis(encoded)), _ => encoded.checked_sub(1).map(Duration::from_millis), }; - let timescale = Timescale::new(u64::decode(r, version)?).map_err(|_| DecodeError::InvalidValue)?; + let timescale = Timescale::new(r.varint()?.into_inner()).map_err(|_| DecodeError::InvalidValue)?; Ok(Self { priority, @@ -82,12 +82,12 @@ impl Message for TrackInfo { }) } - fn encode_msg(&self, w: &mut W, version: Version) -> Result<(), EncodeError> { + fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { if !version.has_track_stream() { return Err(EncodeError::Version); } - self.priority.encode(w, version)?; + w.u8(self.priority); super::subscribe::pad_group_order(w, version)?; let encoded = match (version, self.max_age) { (Version::Lite05 | Version::Lite06, None) => LEGACY_UNLIMITED, @@ -95,8 +95,8 @@ impl Message for TrackInfo { (_, None) => 0, (_, Some(age)) => u64::try_from(age.as_millis() + 1).map_err(|_| EncodeError::BoundsExceeded)?, }; - encoded.encode(w, version)?; - u64::from(self.timescale).encode(w, version)?; + w.varint(VarInt::from(encoded))?; + w.varint(VarInt::from(u64::from(self.timescale)))?; Ok(()) } } @@ -115,9 +115,10 @@ mod test { fn info_roundtrip(version: Version, info: &TrackInfo) -> TrackInfo { let mut buf = Vec::new(); - info.encode_msg(&mut buf, version).unwrap(); + info.encode_msg(&mut Encoder::new(&mut buf, version.into()), version) + .unwrap(); let mut slice = buf.as_slice(); - TrackInfo::decode_msg(&mut slice, version).unwrap() + crate::coding::decode_buf(&mut slice, version, TrackInfo::decode_msg).unwrap() } #[test] @@ -145,7 +146,7 @@ mod test { max_age: age, ..info_sample() } - .encode_msg(&mut buf, Version::Lite07) + .encode_msg(&mut Encoder::new(&mut buf, Version::Lite07.into()), Version::Lite07) .unwrap(); assert_eq!(buf[1], encoded); } @@ -156,11 +157,12 @@ mod test { for version in [Version::Lite05, Version::Lite06] { for millis in [LEGACY_UNLIMITED - 1, LEGACY_UNLIMITED, 1 << 53, 1 << 60, (1 << 62) - 1] { let mut raw = Vec::new(); - 0u8.encode(&mut raw, version).unwrap(); - super::super::subscribe::pad_group_order(&mut raw, version).unwrap(); - millis.encode(&mut raw, version).unwrap(); - 1000u64.encode(&mut raw, version).unwrap(); - let decoded = TrackInfo::decode_msg(&mut raw.as_slice(), version).unwrap(); + let w = &mut Encoder::new(&mut raw, version.into()); + w.u8(0); + super::super::subscribe::pad_group_order(w, version).unwrap(); + w.varint(millis.into()).unwrap(); + w.varint(1000u64.into()).unwrap(); + let decoded = TrackInfo::decode_msg(&mut Decoder::new(&raw, version.into()), version).unwrap(); assert_eq!( decoded.max_age, (millis < LEGACY_UNLIMITED).then(|| Duration::from_millis(millis)) @@ -170,14 +172,12 @@ mod test { ..info_sample() }; let mut encoded = Vec::new(); - info.encode_msg(&mut encoded, version).unwrap(); - let mut old_reader = encoded.as_slice(); - u8::decode(&mut old_reader, version).unwrap(); - super::super::subscribe::skip_group_order(&mut old_reader, version).unwrap(); - assert_eq!( - u64::decode(&mut old_reader, version).unwrap(), - millis.min(LEGACY_UNLIMITED) - ); + info.encode_msg(&mut Encoder::new(&mut encoded, version.into()), version) + .unwrap(); + let old_reader = &mut Decoder::new(&encoded, version.into()); + old_reader.u8().unwrap(); + super::super::subscribe::skip_group_order(old_reader, version).unwrap(); + assert_eq!(old_reader.varint().unwrap().into_inner(), millis.min(LEGACY_UNLIMITED)); } } } @@ -206,7 +206,8 @@ mod test { timescale: info.timescale, }; let mut buf = Vec::new(); - info.encode(&mut buf, Version::Lite05).unwrap(); + info.encode(&mut Encoder::new(&mut buf, Version::Lite05.into()), Version::Lite05) + .unwrap(); assert_eq!( buf, @@ -219,7 +220,11 @@ mod test { #[test] fn track_info_errors_before_lite05() { let mut buf = Vec::new(); - assert!(info_sample().encode_msg(&mut buf, Version::Lite04).is_err()); + assert!( + info_sample() + .encode_msg(&mut Encoder::new(&mut buf, Version::Lite04.into()), Version::Lite04) + .is_err() + ); } #[test] @@ -247,15 +252,17 @@ mod test { } #[test] - fn track_info_encode_rejects_max_age_past_varint_without_writing() { + fn track_info_encode_rejects_max_age_past_varint() { let info = TrackInfo { priority: 7, max_age: Some(Duration::from_millis(1u64 << 62)), timescale: Timescale::MILLI, }; - let mut buf = Vec::new(); - assert!(info.encode(&mut buf, Version::Lite07).is_err()); - assert!(buf.is_empty()); + // The partial bytes are the Writer's to drop; see `a_failed_encode_leaves_no_partial_bytes`. + assert!(matches!( + info.encode_bytes(Version::Lite07), + Err(EncodeError::BoundsExceeded) + )); } #[test] @@ -265,9 +272,10 @@ mod test { track: Cow::Borrowed("video"), }; let mut buf = Vec::new(); - msg.encode_msg(&mut buf, Version::Lite05).unwrap(); + msg.encode_msg(&mut Encoder::new(&mut buf, Version::Lite05.into()), Version::Lite05) + .unwrap(); let mut slice = buf.as_slice(); - let got = Track::decode_msg(&mut slice, Version::Lite05).unwrap(); + let got = crate::coding::decode_buf(&mut slice, Version::Lite05, Track::decode_msg).unwrap(); assert_eq!(got.broadcast, Path::new("room")); assert_eq!(got.track, "video"); } @@ -279,6 +287,9 @@ mod test { track: Cow::Borrowed("video"), }; let mut buf = Vec::new(); - assert!(msg.encode_msg(&mut buf, Version::Lite04).is_err()); + assert!( + msg.encode_msg(&mut Encoder::new(&mut buf, Version::Lite04.into()), Version::Lite04) + .is_err() + ); } } diff --git a/rs/moq-net/src/model/origin.rs b/rs/moq-net/src/model/origin.rs index a992112d9a..a45f8dbcf8 100644 --- a/rs/moq-net/src/model/origin.rs +++ b/rs/moq-net/src/model/origin.rs @@ -18,7 +18,7 @@ use super::{ }; use crate::{ AsPath, Error, InvalidPattern, Path, PathOwned, Pattern, Patterns, - coding::{BoundsExceeded, Decode, DecodeError, Encode, EncodeError}, + coding::{BoundsExceeded, Decode, DecodeError, Decoder, Encode, EncodeError, Encoder, VarInt}, path::Segment, runtime::{Instant, Timers}, time::Clock, @@ -156,21 +156,16 @@ impl fmt::Display for Hop { } } -impl Encode for Hop -where - u64: Encode, -{ - fn encode(&self, w: &mut W, version: V) -> Result<(), EncodeError> { - self.id.encode(w, version) +impl Encode for Hop { + fn encode(&self, w: &mut Encoder<'_>, _: V) -> Result<(), EncodeError> { + w.varint(VarInt::from(self.id))?; + Ok(()) } } -impl Decode for Hop -where - u64: Decode, -{ - fn decode(r: &mut R, version: V) -> Result { - Self::from_wire(u64::decode(r, version)?) +impl Decode for Hop { + fn decode(r: &mut Decoder<'_>, _: V) -> Result { + Self::from_wire(r.varint()?.into_inner()) } } @@ -304,13 +299,9 @@ impl<'a> IntoIterator for &'a Hops { } } -impl Encode for Hops -where - u64: Encode, - Hop: Encode, -{ - fn encode(&self, w: &mut W, version: V) -> Result<(), EncodeError> { - (self.0.len() as u64).encode(w, version)?; +impl Encode for Hops { + fn encode(&self, w: &mut Encoder<'_>, version: V) -> Result<(), EncodeError> { + w.varint(VarInt::from(self.0.len()))?; for origin in &self.0 { origin.encode(w, version)?; } @@ -318,13 +309,9 @@ where } } -impl Decode for Hops -where - u64: Decode, - Hop: Decode, -{ - fn decode(r: &mut R, version: V) -> Result { - let count = u64::decode(r, version)? as usize; +impl Decode for Hops { + fn decode(r: &mut Decoder<'_>, version: V) -> Result { + let count = r.varint()?.into_inner() as usize; if count > MAX_HOPS { return Err(DecodeError::BoundsExceeded); } @@ -8009,9 +7996,8 @@ mod tests { fn drain_cost_is_encodable() { use crate::coding::Encode; - let mut buf = Vec::new(); Cost::DRAIN - .encode(&mut buf, crate::lite::Version::Lite06) + .encode_bytes(crate::lite::Version::Lite06) .expect("a draining route is still forwarded, so its cost must encode"); } } diff --git a/rs/moq-net/src/model/time.rs b/rs/moq-net/src/model/time.rs index e1a7b5a118..764ff26fb3 100644 --- a/rs/moq-net/src/model/time.rs +++ b/rs/moq-net/src/model/time.rs @@ -2,6 +2,16 @@ use std::num::NonZero; use crate::coding::VarInt; +/// `value` as a [`VarInt`] the QUIC form can carry, or `None` past `2^62 - 1`, so every +/// timestamp stays encodable on moq-lite. +const fn quic(value: u128) -> Option { + if value <= VarInt::MAX_QUIC.into_inner() as u128 { + Some(VarInt::from_u64(value as u64)) + } else { + None + } +} + /// Returned when a [`Timestamp`] operation would exceed the QUIC VarInt range /// (`2^62 - 1`), overflow during scale conversion or arithmetic, or attempt /// arithmetic between timestamps with mismatched scales. @@ -51,7 +61,7 @@ impl Timescale { pub const fn new(units_per_second: u64) -> Result { // Reject values that wouldn't fit in a QUIC varint, keeping the constraint // symmetric with Timestamp's raw value. - if VarInt::from_u64(units_per_second).is_none() { + if quic(units_per_second as u128).is_none() { return Err(TimeOverflow); } match NonZero::new(units_per_second) { @@ -174,7 +184,7 @@ impl Timestamp { /// Construct a timestamp directly from a raw value at the given scale. /// Returns [`TimeOverflow`] if `value` exceeds `2^62 - 1`. pub const fn new(value: u64, scale: Timescale) -> Result { - match VarInt::from_u64(value) { + match quic(value as u128) { Some(value) => Ok(Self { value, scale }), None => Err(TimeOverflow), } @@ -236,7 +246,7 @@ impl Timestamp { return Ok(self); } match (self.value.into_inner() as u128).checked_mul(new_scale.0.get() as u128) { - Some(scaled) => match VarInt::from_u128(scaled / self.scale.0.get() as u128) { + Some(scaled) => match quic(scaled / self.scale.0.get() as u128) { Some(value) => Ok(Self { value, scale: new_scale, @@ -314,7 +324,7 @@ impl TryFrom for Timestamp { /// Convert a [`std::time::Duration`] into a nanosecond-scale timestamp. fn try_from(duration: std::time::Duration) -> Result { - match VarInt::from_u128(duration.as_nanos()) { + match quic(duration.as_nanos()) { Some(value) => Ok(Self { value, scale: Timescale::NANO, diff --git a/rs/moq-net/src/model/track.rs b/rs/moq-net/src/model/track.rs index 9fc80e0115..6ae6baccdc 100644 --- a/rs/moq-net/src/model/track.rs +++ b/rs/moq-net/src/model/track.rs @@ -4884,7 +4884,7 @@ mod test { /// the wire codec and the model. #[tokio::test] async fn datagram_wire_roundtrip_between_tracks() { - use crate::coding::{Decode, Encode}; + use crate::coding::Encode; use crate::lite; let version = lite::Version::Lite05; @@ -4906,8 +4906,7 @@ mod test { .unwrap(); // Subscriber decodes the body and writes it downstream, preserving the sequence. - let mut slice = &body[..]; - let wire = lite::Datagram::decode(&mut slice, version).unwrap(); + let wire = lite::Datagram::decode(body, version).unwrap(); let mut downstream = track_producer("test", None); let mut downstream_dg = downstream.subscribe(None); downstream diff --git a/rs/moq-net/src/path/mod.rs b/rs/moq-net/src/path/mod.rs index 21acabb44e..70e534c8be 100644 --- a/rs/moq-net/src/path/mod.rs +++ b/rs/moq-net/src/path/mod.rs @@ -14,7 +14,7 @@ use std::borrow::Cow; use std::fmt::{self, Display}; use std::sync::Arc; -use crate::coding::{Decode, DecodeError, Encode, EncodeError}; +use crate::coding::{Decode, DecodeError, Decoder, Encode, EncodeError, Encoder}; /// An owned version of [`Path`] with a `'static` lifetime. pub type PathOwned = Path<'static>; @@ -579,12 +579,9 @@ impl Display for Path<'_> { } } -impl Decode for Path<'_> -where - String: Decode, -{ - fn decode(r: &mut R, version: V) -> Result { - let path: Path = String::decode(r, version)?.into(); +impl Decode for Path<'_> { + fn decode(r: &mut Decoder<'_>, _: V) -> Result { + let path: Path = r.string()?.into(); if path.parts().count() > Path::MAX_PARTS { return Err(DecodeError::BoundsExceeded); } @@ -592,16 +589,12 @@ where } } -impl Encode for Path<'_> -where - for<'a> &'a str: Encode, -{ - fn encode(&self, w: &mut W, version: V) -> Result<(), EncodeError> { +impl Encode for Path<'_> { + fn encode(&self, w: &mut Encoder<'_>, _: V) -> Result<(), EncodeError> { if self.parts().count() > Path::MAX_PARTS { return Err(EncodeError::BoundsExceeded); } - self.as_str().encode(w, version)?; - Ok(()) + w.string(self.as_str()) } } @@ -1251,22 +1244,26 @@ mod tests { let too_deep = format!("{ok}/extra"); // Encode enforces the limit. - let mut buf = bytes::BytesMut::new(); - Path::new(&ok).encode(&mut buf, Version::Lite04).unwrap(); + let mut buf = Vec::new(); + Path::new(&ok) + .encode(&mut Encoder::new(&mut buf, Version::Lite04.into()), Version::Lite04) + .unwrap(); assert!(matches!( - Path::new(&too_deep).encode(&mut bytes::BytesMut::new(), Version::Lite04), + Path::new(&too_deep).encode_bytes(Version::Lite04), Err(EncodeError::BoundsExceeded) )); // Decode round-trips at the limit. - let decoded = Path::decode(&mut buf.freeze(), Version::Lite04).unwrap(); + let decoded = crate::coding::decode_buf(&mut bytes::Bytes::from(buf), Version::Lite04, Path::decode).unwrap(); assert_eq!(decoded.as_str(), ok); // Decode enforces the limit on a raw string that encode would have refused. - let mut buf = bytes::BytesMut::new(); - too_deep.as_str().encode(&mut buf, Version::Lite04).unwrap(); + let mut buf = Vec::new(); + Encoder::new(&mut buf, Version::Lite04.into()) + .string(too_deep.as_str()) + .unwrap(); assert!(matches!( - Path::decode(&mut buf.freeze(), Version::Lite04), + crate::coding::decode_buf(&mut bytes::Bytes::from(buf), Version::Lite04, Path::decode), Err(DecodeError::BoundsExceeded) )); } diff --git a/rs/moq-net/src/server.rs b/rs/moq-net/src/server.rs index 6a35f32394..265b7c009c 100644 --- a/rs/moq-net/src/server.rs +++ b/rs/moq-net/src/server.rs @@ -334,7 +334,7 @@ impl Server { // Legacy bidi SETUP exchange (IETF 14-16, lite 01/02). Read the client's // SETUP to choose the version; `ok()` sends the server SETUP and starts. let mut stream = Stream::accept(&mut session, encoding).await?; - let mut client: setup::Client = stream.reader.decode().await?; + let client: setup::Client = stream.reader.decode().await?; let version = client .versions @@ -348,7 +348,7 @@ impl Server { // in its SETUP just like lite-05. let (path, token, request_id_max, peer_declared) = match version { Version::Ietf(v) => { - let params = ietf::Parameters::decode(&mut client.parameters, v)?; + let (params, _) = ietf::Parameters::decode_slice(&client.parameters, v)?; let path = match params.get_bytes(ietf::ParameterBytes::Path) { Some(bytes) => Some( std::str::from_utf8(bytes) @@ -979,7 +979,9 @@ mod tests { fn lite05_setup(path: Option<&str>, role: Option, hop: Option) -> Vec { let v = lite::Version::Lite05; let mut buf = Vec::new(); - lite::DataType::Setup.encode(&mut buf, v).unwrap(); + lite::DataType::Setup + .encode(&mut crate::coding::Encoder::new(&mut buf, v.into()), v) + .unwrap(); lite::Setup { probe: lite::ProbeLevel::None, path: path.map(str::to_string), @@ -987,7 +989,7 @@ mod tests { cost: None, hop, } - .encode(&mut buf, v) + .encode(&mut crate::coding::Encoder::new(&mut buf, v.into()), v) .unwrap(); buf } @@ -1007,7 +1009,10 @@ mod tests { let mut buf = Vec::new(); setup::Setup { parameters } - .encode(&mut buf, crate::Version::Ietf(version)) + .encode( + &mut crate::coding::Encoder::new(&mut buf, (crate::Version::Ietf(version)).into()), + crate::Version::Ietf(version), + ) .unwrap(); buf } @@ -1019,7 +1024,10 @@ mod tests { versions: crate::coding::Versions::from([crate::Version::Ietf(version).into()]), parameters: params.encode_bytes(version).unwrap(), } - .encode(&mut buf, crate::Version::Ietf(version)) + .encode( + &mut crate::coding::Encoder::new(&mut buf, (crate::Version::Ietf(version)).into()), + crate::Version::Ietf(version), + ) .unwrap(); buf } @@ -1138,7 +1146,12 @@ mod tests { /// Encode a lite-05 GROUP uni stream header (just the `DataType::Group` tag). fn lite05_group() -> Vec { let mut buf = Vec::new(); - lite::DataType::Group.encode(&mut buf, lite::Version::Lite05).unwrap(); + lite::DataType::Group + .encode( + &mut crate::coding::Encoder::new(&mut buf, lite::Version::Lite05.into()), + lite::Version::Lite05, + ) + .unwrap(); buf } diff --git a/rs/moq-net/src/setup.rs b/rs/moq-net/src/setup.rs index 3a4a96cba0..529a26c652 100644 --- a/rs/moq-net/src/setup.rs +++ b/rs/moq-net/src/setup.rs @@ -6,7 +6,7 @@ use bytes::Bytes; use crate::{ Version, - coding::{self, Decode, DecodeError, Encode, EncodeError, Sizer}, + coding::{self, Decode, DecodeError, Decoder, Encode, EncodeError, Encoder, VarInt}, ietf, lite, }; @@ -51,33 +51,25 @@ impl Setup { } impl Encode for Setup { - fn encode(&self, w: &mut W, v: Version) -> Result<(), EncodeError> { + fn encode(&self, w: &mut Encoder<'_>, v: Version) -> Result<(), EncodeError> { Self::check_version(v); - SETUP_V17.encode(w, v)?; - u16::try_from(self.parameters.len()) - .map_err(|_| EncodeError::TooLarge)? - .encode(w, v)?; - if w.remaining_mut() < self.parameters.len() { - return Err(EncodeError::Short); - } - w.put_slice(&self.parameters); - Ok(()) + w.varint(VarInt::from(SETUP_V17))?; + let start = w.position(); + w.slice(&self.parameters); + w.prefix_u16(start) } } impl Decode for Setup { - fn decode(r: &mut R, v: Version) -> Result { + fn decode(r: &mut Decoder<'_>, v: Version) -> Result { Self::check_version(v); - let kind = u64::decode(r, v)?; + let kind = r.varint()?.into_inner(); if kind != SETUP_V17 { return Err(DecodeError::InvalidValue); } - let size = u16::decode(r, v)? as usize; - if r.remaining() < size { - return Err(DecodeError::Short); - } - let msg = r.copy_to_bytes(size); - Ok(Self { parameters: msg }) + let size = r.u16()? as usize; + let parameters = Bytes::copy_from_slice(r.slice(size)?); + Ok(Self { parameters }) } } @@ -119,7 +111,7 @@ pub(crate) struct Client { } impl Client { - fn encode_inner(&self, w: &mut W, v: Version) -> Result<(), EncodeError> { + fn encode_inner(&self, w: &mut Encoder<'_>, v: Version) -> Result<(), EncodeError> { match SetupVersion::from_version(v) { SetupVersion::Draft15Plus => { // Draft15+: no versions list, parameters only. @@ -127,33 +119,20 @@ impl Client { SetupVersion::Draft14 | SetupVersion::LiteLegacy => self.versions.encode(w, v)?, SetupVersion::Modern | SetupVersion::Unsupported => return Err(EncodeError::Version), }; - if w.remaining_mut() < self.parameters.len() { - return Err(EncodeError::Short); - } - w.put_slice(&self.parameters); + w.slice(&self.parameters); Ok(()) } } impl Decode for Client { /// Decode a client setup message (draft-14 through draft-16 only). - fn decode(r: &mut R, v: Version) -> Result { - let kind = u8::decode(r, v)?; + fn decode(r: &mut Decoder<'_>, v: Version) -> Result { + let kind = r.u8()?; if kind != CLIENT_SETUP { return Err(DecodeError::InvalidValue); } - let size = match SetupVersion::from_version(v) { - SetupVersion::Draft14 | SetupVersion::Draft15Plus => u16::decode(r, v)? as usize, - SetupVersion::LiteLegacy => u64::decode(r, v)? as usize, - SetupVersion::Modern | SetupVersion::Unsupported => return Err(DecodeError::Version), - }; - - if r.remaining() < size { - return Err(DecodeError::Short); - } - - let mut msg = r.copy_to_bytes(size); + let mut msg = decode_body(r, v)?; let versions = match SetupVersion::from_version(v) { SetupVersion::Draft15Plus => { @@ -166,28 +145,37 @@ impl Decode for Client { Ok(Self { versions, - parameters: msg, + parameters: Bytes::copy_from_slice(msg.rest()), }) } } impl Encode for Client { /// Encode a client setup message (draft-14 through draft-16 only). - fn encode(&self, w: &mut W, v: Version) -> Result<(), EncodeError> { - CLIENT_SETUP.encode(w, v)?; + fn encode(&self, w: &mut Encoder<'_>, v: Version) -> Result<(), EncodeError> { + w.u8(CLIENT_SETUP); + let start = w.position(); + self.encode_inner(w, v)?; + prefix_body(w, v, start) + } +} - let mut sizer = Sizer::default(); - self.encode_inner(&mut sizer, v)?; - let size = sizer.size; +/// Read a pre-draft-17 SETUP body: its size, then that many bytes. +fn decode_body<'a>(r: &mut Decoder<'a>, v: Version) -> Result, DecodeError> { + let size = match SetupVersion::from_version(v) { + SetupVersion::Draft14 | SetupVersion::Draft15Plus => r.u16()? as usize, + SetupVersion::LiteLegacy => usize::try_from(r.varint()?)?, + SetupVersion::Modern | SetupVersion::Unsupported => return Err(DecodeError::Version), + }; + r.sub(size) +} - match SetupVersion::from_version(v) { - SetupVersion::Draft14 | SetupVersion::Draft15Plus => { - u16::try_from(size).map_err(|_| EncodeError::TooLarge)?.encode(w, v)?; - } - SetupVersion::LiteLegacy => (size as u64).encode(w, v)?, - SetupVersion::Modern | SetupVersion::Unsupported => return Err(EncodeError::Version), - } - self.encode_inner(w, v) +/// Size-prefix a pre-draft-17 SETUP body written since `start`. +fn prefix_body(w: &mut Encoder<'_>, v: Version, start: usize) -> Result<(), EncodeError> { + match SetupVersion::from_version(v) { + SetupVersion::Draft14 | SetupVersion::Draft15Plus => w.prefix_u16(start), + SetupVersion::LiteLegacy => w.prefix_varint(start), + SetupVersion::Modern | SetupVersion::Unsupported => Err(EncodeError::Version), } } @@ -202,7 +190,7 @@ pub(crate) struct Server { } impl Server { - fn encode_inner(&self, w: &mut W, v: Version) -> Result<(), EncodeError> { + fn encode_inner(&self, w: &mut Encoder<'_>, v: Version) -> Result<(), EncodeError> { match SetupVersion::from_version(v) { SetupVersion::Draft15Plus => { // Draft15+: No version field, parameters only. @@ -210,54 +198,30 @@ impl Server { SetupVersion::Draft14 | SetupVersion::LiteLegacy => self.version.encode(w, v)?, SetupVersion::Modern | SetupVersion::Unsupported => return Err(EncodeError::Version), }; - if w.remaining_mut() < self.parameters.len() { - return Err(EncodeError::Short); - } - w.put_slice(&self.parameters); + w.slice(&self.parameters); Ok(()) } } impl Encode for Server { /// Encode a server setup message (draft-14 through draft-16 only). - fn encode(&self, w: &mut W, v: Version) -> Result<(), EncodeError> { - SERVER_SETUP.encode(w, v)?; - - let mut sizer = Sizer::default(); - self.encode_inner(&mut sizer, v)?; - let size = sizer.size; - - match SetupVersion::from_version(v) { - SetupVersion::Draft14 | SetupVersion::Draft15Plus => { - u16::try_from(size).map_err(|_| EncodeError::TooLarge)?.encode(w, v)?; - } - SetupVersion::LiteLegacy => (size as u64).encode(w, v)?, - SetupVersion::Modern | SetupVersion::Unsupported => return Err(EncodeError::Version), - } - - self.encode_inner(w, v) + fn encode(&self, w: &mut Encoder<'_>, v: Version) -> Result<(), EncodeError> { + w.u8(SERVER_SETUP); + let start = w.position(); + self.encode_inner(w, v)?; + prefix_body(w, v, start) } } impl Decode for Server { /// Decode a server setup message (draft-14 through draft-16 only). - fn decode(r: &mut R, v: Version) -> Result { - let kind = u8::decode(r, v)?; + fn decode(r: &mut Decoder<'_>, v: Version) -> Result { + let kind = r.u8()?; if kind != SERVER_SETUP { return Err(DecodeError::InvalidValue); } - let size = match SetupVersion::from_version(v) { - SetupVersion::Draft14 | SetupVersion::Draft15Plus => u16::decode(r, v)? as usize, - SetupVersion::LiteLegacy => u64::decode(r, v)? as usize, - SetupVersion::Modern | SetupVersion::Unsupported => return Err(DecodeError::Version), - }; - - if r.remaining() < size { - return Err(DecodeError::Short); - } - - let mut msg = r.copy_to_bytes(size); + let mut msg = decode_body(r, v)?; let version = match SetupVersion::from_version(v) { SetupVersion::Draft15Plus => v.into(), SetupVersion::Draft14 | SetupVersion::LiteLegacy => coding::Version::decode(&mut msg, v)?, @@ -266,7 +230,7 @@ impl Decode for Server { Ok(Self { version, - parameters: msg, + parameters: Bytes::copy_from_slice(msg.rest()), }) } } diff --git a/rs/moq-net/src/version.rs b/rs/moq-net/src/version.rs index e84ab91b38..e1cccfb7b3 100644 --- a/rs/moq-net/src/version.rs +++ b/rs/moq-net/src/version.rs @@ -267,13 +267,13 @@ impl TryFrom for Version { } impl coding::Decode for Version { - fn decode(r: &mut R, version: Version) -> Result { + fn decode(r: &mut coding::Decoder<'_>, version: Version) -> Result { coding::Version::decode(r, version).and_then(|v| v.try_into().map_err(|_| coding::DecodeError::InvalidValue)) } } impl coding::Encode for Version { - fn encode(&self, w: &mut W, v: Version) -> Result<(), coding::EncodeError> { + fn encode(&self, w: &mut coding::Encoder<'_>, v: Version) -> Result<(), coding::EncodeError> { coding::Version::from(*self).encode(w, v) } } From 526b82d890e00bea7c4349746fb1ea4edb5b669a Mon Sep 17 00:00:00 2001 From: Luke Curley Date: Mon, 28 Sep 2026 21:13:00 -0700 Subject: [PATCH 4/8] perf(net): fixed-size varint reads and reserved size prefixes Split the varint codec per form with fixed-size reads and writes, and reserve a message's size prefix ahead of its body so a small body never moves. Delete the finished varint-codec quest and move the questline's varint text to the full 64-bit range. Co-Authored-By: Claude Opus 5.5 --- quest/m1/rs2ts/README.md | 11 +- quest/m1/rs2ts/js-varint.md | 7 +- quest/m1/rs2ts/translator.md | 2 +- quest/m1/rs2ts/varint-codec.md | 31 --- rs/moq-net/src/coding/decode.rs | 5 +- rs/moq-net/src/coding/encode.rs | 109 ++++++++-- rs/moq-net/src/coding/varint.rs | 238 +++++++++++++++------ rs/moq-net/src/ietf/fetch.rs | 12 +- rs/moq-net/src/ietf/filter.rs | 8 +- rs/moq-net/src/ietf/goaway.rs | 2 +- rs/moq-net/src/ietf/message.rs | 4 +- rs/moq-net/src/ietf/parameters.rs | 2 +- rs/moq-net/src/ietf/properties.rs | 2 +- rs/moq-net/src/ietf/publish.rs | 6 +- rs/moq-net/src/ietf/publish_namespace.rs | 4 +- rs/moq-net/src/ietf/publisher.rs | 12 +- rs/moq-net/src/ietf/subscribe.rs | 12 +- rs/moq-net/src/ietf/subscribe_namespace.rs | 6 +- rs/moq-net/src/ietf/track.rs | 6 +- rs/moq-net/src/lite/announce.rs | 8 +- rs/moq-net/src/lite/message.rs | 4 +- rs/moq-net/src/lite/subscribe.rs | 2 +- rs/moq-net/src/setup.rs | 20 +- 23 files changed, 328 insertions(+), 185 deletions(-) delete mode 100644 quest/m1/rs2ts/varint-codec.md diff --git a/quest/m1/rs2ts/README.md b/quest/m1/rs2ts/README.md index 9fffcb51c3..534fa7d855 100644 --- a/quest/m1/rs2ts/README.md +++ b/quest/m1/rs2ts/README.md @@ -30,9 +30,11 @@ Decided in planning (2026-09-27), with the spike data in and bytes out, no runtime. The async helper methods move behind an `async` cargo feature; rs2ts reads the crate without it and JS reimplements the helpers with Promises. No second crate. -- Varints stay 62-bit on the wire; the spec is not bounded to 2^53. Rust's - `VarInt` newtype carries Encode/Decode and JS gets a matching `VarInt` type - with checked conversion to and from `number`. +- `VarInt` holds the full 64-bit range; the spec is not bounded to 2^53. The + leading-ones form (moq-transport draft-17+) carries all of it, and the QUIC + form (moq-lite, drafts 14-16) refuses anything past 2^62 - 1 rather than + truncating. Rust's `VarInt` newtype carries Encode/Decode and JS gets a + matching `VarInt` type with checked conversion to and from `number`. - The generated TypeScript is committed and a CI lane regenerates it and fails on drift, so JS contributors and npm publishing never need the nightly toolchain Charon pins. It lives inside js/net and `@moq/net` stays @@ -54,8 +56,7 @@ js/net it replaces, measured with the [browser benchmarks](/quest/m1/browser-ben ## Required -- [VarInt codec](/quest/m1/rs2ts/varint-codec.md) - moq-net encodes through a `VarInt` newtype and a concrete slice-based codec, not generic traits on primitives -- [JS VarInt](/quest/m1/rs2ts/js-varint.md) - js/net has a 62-bit `VarInt` type with checked `number` conversion and no BigInt on the hot path +- [JS VarInt](/quest/m1/rs2ts/js-varint.md) - js/net has a 64-bit `VarInt` type with checked `number` conversion and no BigInt on the hot path - [rs2ts](/quest/m1/rs2ts/translator.md) - a Charon-based translator emits readable TypeScript for moq-net's lite codec, committed and checked for drift in CI - [Sans-IO moq-net](/quest/m1/rs2ts/sans-io/README.md) - moq-net builds and runs without a runtime; async helpers sit behind an `async` feature - [Mock-clock tests](/quest/m1/rs2ts/mock-clock.md) - moq-net's tests run on the sans-IO clock instead of tokio, so they translate with the code diff --git a/quest/m1/rs2ts/js-varint.md b/quest/m1/rs2ts/js-varint.md index 0c17bc4432..15a197b55e 100644 --- a/quest/m1/rs2ts/js-varint.md +++ b/quest/m1/rs2ts/js-varint.md @@ -2,7 +2,7 @@ ## Goal -js/net has a `VarInt` type that holds the full 62-bit range, converts to and +js/net has a `VarInt` type that holds the full 64-bit range, converts to and from `number` with a loud error outside the safe range, and encodes and decodes without BigInt on the hot path. It is the TypeScript type rs2ts maps Rust's `VarInt` to. @@ -21,7 +21,8 @@ Guidance: and comparison and increment methods so sequence logic never converts. - Move js/net's varint reading and writing onto it, dropping the BigInt round trip for QUIC and leading-ones varints. -- Unit-test the boundaries (2^30, 2^53, 2^62 - 1) against Rust's encoder in - `just test interop`. +- Unit-test the boundaries (2^30, 2^53, 2^62 - 1, 2^62, 2^64 - 1) against + Rust's encoder in `just test interop`. The QUIC form refuses anything past + 2^62 - 1, as Rust's does. Public API: additive to `@moq/net`; lands on `main`. Wire: none. diff --git a/quest/m1/rs2ts/translator.md b/quest/m1/rs2ts/translator.md index 8459752e54..8941968c97 100644 --- a/quest/m1/rs2ts/translator.md +++ b/quest/m1/rs2ts/translator.md @@ -55,5 +55,5 @@ translates. Wire: none. ## Required -- [VarInt codec](/quest/m1/rs2ts/varint-codec.md) - the codec shape the translator targets + - [JS VarInt](/quest/m1/rs2ts/js-varint.md) - the TypeScript type `VarInt` maps to diff --git a/quest/m1/rs2ts/varint-codec.md b/quest/m1/rs2ts/varint-codec.md deleted file mode 100644 index 39557cb31b..0000000000 --- a/quest/m1/rs2ts/varint-codec.md +++ /dev/null @@ -1,31 +0,0 @@ -# [M] VarInt codec - -## Goal - -Every varint moq-net puts on the wire is a `VarInt`, and `VarInt` is the only -integer type with Encode/Decode. Messages encode and decode through a -concrete, slice-based codec instead of generic traits implemented on `u64`, -`usize`, `bool`, `String`, `Option`, and `Vec`. - -## Plan - -Two reasons, one refactor. Not every `u64` is a valid varint, so the type -should say which fields are. And the rs2ts translator cannot map generic -traits on primitives without dictionary passing, its most expensive feature; -the same generics are most of why moq-net's lite codec built to 47 KB gzip in -WASM against 6 KB for a hand-carved one. - -Guidance: - -- Keep the 62-bit range; the wire is not bounded to 2^53. -- Message types keep a local trait; the primitives become inherent methods - on concrete reader and writer types (`varint`, `string`, `bytes`, ...). - Make the version a concrete type rather than a generic `V` where possible. -- Avoid bit operations on `u64` in the codec: write the 8-byte form as two - `u32` halves, so the generated TypeScript never needs 64-bit bitwise math. -- `Parameters` becomes Vec-backed, and decode paths stop branching on - `tracing::enabled!` (log after decoding instead). -- Benchmark the codec before and after (Criterion); it is on every message. - -Public API: breaks moq-net's `coding` module (Encode/Decode on primitives -go away), so this retargets to `dev`. Wire: none. diff --git a/rs/moq-net/src/coding/decode.rs b/rs/moq-net/src/coding/decode.rs index 9d95ececaf..c1fbfc0dab 100644 --- a/rs/moq-net/src/coding/decode.rs +++ b/rs/moq-net/src/coding/decode.rs @@ -172,9 +172,10 @@ impl<'a> Decoder<'a> { } /// Read a varint. + #[inline] pub fn varint(&mut self) -> Result { - let (value, len) = VarInt::decode_form(self.buf, self.form)?; - self.buf = &self.buf[len..]; + let (value, rest) = VarInt::read(self.buf, self.form)?; + self.buf = rest; Ok(value) } diff --git a/rs/moq-net/src/coding/encode.rs b/rs/moq-net/src/coding/encode.rs index 6045fc835c..7076a8a08d 100644 --- a/rs/moq-net/src/coding/encode.rs +++ b/rs/moq-net/src/coding/encode.rs @@ -71,12 +71,6 @@ impl<'a> Encoder<'a> { self.form } - /// Where the next byte goes: the buffer's length, including any written before this - /// encoder. Mark a size-prefixed body's start with it. - pub fn position(&self) -> usize { - self.buf.len() - } - /// Write raw bytes. pub fn slice(&mut self, v: &[u8]) { self.buf.extend_from_slice(v); @@ -98,10 +92,9 @@ impl<'a> Encoder<'a> { } /// Write a varint, or fail with [`EncodeError::BoundsExceeded`] if the form cannot carry it. + #[inline] pub fn varint(&mut self, v: VarInt) -> Result<(), EncodeError> { - let (buf, len) = v.encode_form(self.form)?; - self.buf.extend_from_slice(&buf[..len]); - Ok(()) + Ok(v.write(self.form, self.buf)?) } /// Write an optional varint: `None` as 0, and `Some(n)` as `n + 1`. @@ -125,40 +118,110 @@ impl<'a> Encoder<'a> { self.bytes(v.as_bytes()) } - /// Prefix everything written since [`Self::position`] was `start` with its varint length. - pub fn prefix_varint(&mut self, start: usize) -> Result<(), EncodeError> { - let size = VarInt::from(self.buf.len() - start); - let (prefix, len) = size.encode_form(self.form)?; - self.buf.splice(start..start, prefix[..len].iter().copied()); - Ok(()) + /// Reserve a varint size prefix for the body written next; [`Self::fill`] sizes it. + /// + /// One byte is reserved, which most bodies fit, so they never move. + pub fn prefix_varint(&mut self) -> Prefix { + self.buf.push(0); + Prefix { + body: self.buf.len(), + kind: PrefixKind::VarInt, + } } - /// Prefix everything written since [`Self::position`] was `start` with its `u16` length. - pub fn prefix_u16(&mut self, start: usize) -> Result<(), EncodeError> { - let size = u16::try_from(self.buf.len() - start).map_err(|_| EncodeError::TooLarge)?; - self.buf.splice(start..start, size.to_be_bytes()); + /// Reserve a big-endian `u16` size prefix for the body written next; [`Self::fill`] sizes it. + pub fn prefix_u16(&mut self) -> Prefix { + self.buf.extend_from_slice(&[0, 0]); + Prefix { + body: self.buf.len(), + kind: PrefixKind::U16, + } + } + + /// Size a reserved prefix to everything written since. + pub fn fill(&mut self, prefix: Prefix) -> Result<(), EncodeError> { + let body = prefix.body; + let end = self.buf.len(); + + match prefix.kind { + PrefixKind::U16 => { + let size = u16::try_from(end - body).map_err(|_| EncodeError::TooLarge)?; + self.buf[body - 2..body].copy_from_slice(&size.to_be_bytes()); + } + PrefixKind::VarInt => { + // Encode the size past the body, then move it into the reserved byte, + // shifting the body up when the size needs more than that one byte. + VarInt::from(end - body).write(self.form, self.buf)?; + let len = self.buf.len() - end; + if len == 1 { + self.buf[body - 1] = self.buf[end]; + self.buf.truncate(end); + return Ok(()); + } + + let mut size = [0u8; 9]; + size[..len].copy_from_slice(&self.buf[end..]); + self.buf.truncate(end + len - 1); + self.buf.copy_within(body..end, body + len - 1); + self.buf[body - 1..body - 1 + len].copy_from_slice(&size[..len]); + } + } Ok(()) } } +/// A size prefix reserved ahead of a body, sized by [`Encoder::fill`] once it is written. +#[must_use = "an unfilled prefix leaves a zero size on the wire"] +#[derive(Debug)] +pub struct Prefix { + /// Where the body starts, just past the reserved bytes. + body: usize, + kind: PrefixKind, +} + +#[derive(Debug)] +enum PrefixKind { + VarInt, + U16, +} + #[cfg(test)] mod tests { use super::*; - /// The prefix lands before the body, sized to it, even when it takes more than a byte. + /// The prefix lands before the body, sized to it, even when it outgrows the one byte + /// reserved for it. #[test] fn prefix_varint_sizes_the_body() { for size in [0usize, 63, 64, 20_000] { let mut buf = vec![0xaa]; let mut w = Encoder::new(&mut buf, Form::Quic); - let start = w.position(); + let prefix = w.prefix_varint(); w.slice(&vec![0x55; size]); - w.prefix_varint(start).unwrap(); + w.fill(prefix).unwrap(); + w.u8(0xbb); let mut r = super::super::Decoder::new(&buf[1..], Form::Quic); assert_eq!(buf[0], 0xaa); assert_eq!(r.varint().unwrap().into_inner(), size as u64); - assert_eq!(r.rest(), vec![0x55; size]); + assert_eq!(r.slice(size).unwrap(), vec![0x55; size]); + assert_eq!(r.rest(), [0xbb]); } } + + #[test] + fn prefix_u16_sizes_the_body() { + let mut buf = Vec::new(); + let mut w = Encoder::new(&mut buf, Form::Quic); + let prefix = w.prefix_u16(); + w.slice(&[0x55; 300]); + w.fill(prefix).unwrap(); + assert_eq!(buf[..2], 300u16.to_be_bytes()); + assert_eq!(buf.len(), 302); + + let mut w = Encoder::new(&mut buf, Form::Quic); + let prefix = w.prefix_u16(); + w.slice(&vec![0; 1 << 16]); + assert!(matches!(w.fill(prefix), Err(EncodeError::TooLarge))); + } } diff --git a/rs/moq-net/src/coding/varint.rs b/rs/moq-net/src/coding/varint.rs index 5a55bb8547..9f5e5a3a59 100644 --- a/rs/moq-net/src/coding/varint.rs +++ b/rs/moq-net/src/coding/varint.rs @@ -224,69 +224,167 @@ impl VarInt { Self(((hi as u64) << 32) | lo as u64) } - /// Decode from the front of `buf`, returning the value and the bytes it took. - pub(super) fn decode_form(buf: &[u8], form: Form) -> Result<(Self, usize), DecodeError> { - let Some(&first) = buf.first() else { - return Err(DecodeError::Short); - }; + /// The bytes this takes on the wire in the given form, or [`BoundsExceeded`] if it + /// does not fit. + pub(crate) fn size(self, form: Form) -> Result { + let (hi, lo) = self.to_halves(); + Ok(match form { + Form::Quic if hi == 0 && lo < 1 << 6 => 1, + Form::Quic if hi == 0 && lo < 1 << 14 => 2, + Form::Quic if hi == 0 && lo < 1 << 30 => 4, + Form::Quic if hi < 1 << 30 => 8, + Form::Quic => return Err(BoundsExceeded), + Form::LeadingOnes { .. } if hi == 0 && lo < 1 << 7 => 1, + Form::LeadingOnes { .. } if hi == 0 && lo < 1 << 14 => 2, + Form::LeadingOnes { .. } if hi == 0 && lo < 1 << 21 => 3, + Form::LeadingOnes { .. } if hi == 0 && lo < 1 << 28 => 4, + Form::LeadingOnes { .. } if hi < 1 << 3 => 5, + Form::LeadingOnes { .. } if hi < 1 << 10 => 6, + // The 7-byte form is skipped: one byte longer, but legal on every draft. + Form::LeadingOnes { .. } if hi < 1 << 24 => 8, + Form::LeadingOnes { .. } => 9, + }) + } - let (len, head) = match form { - Form::Quic => (1usize << (first >> 6), first & 0x3f), - Form::LeadingOnes { seven } => { - let ones = first.leading_ones(); - if ones == 6 && !seven { - return Err(DecodeError::InvalidValue); - } - // `0x7f >> 8` would overflow, and there are no value bits left anyway. - let head = if ones >= 7 { 0 } else { first & (0x7f >> ones) }; - (ones as usize + 1, head) + /// Append the minimal encoding in the given form. + /// + /// Fails past [`Self::MAX_QUIC`] in the QUIC form, writing nothing. + #[inline] + pub(super) fn write(self, form: Form, out: &mut Vec) -> Result<(), BoundsExceeded> { + match form { + Form::Quic => self.write_quic(out), + Form::LeadingOnes { .. } => { + self.write_leading_ones(out); + Ok(()) } - }; + } + } - let Some(rest) = buf.get(1..len) else { - return Err(DecodeError::Short); - }; + // Each arm below is a fixed-size write or read, which is what keeps the codec as fast + // as a hand-rolled `put_u16`/`get_u32`. - let mut hi = 0u32; - let mut lo = head as u32; - for &byte in rest { - hi = (hi << 8) | (lo >> 24); - lo = (lo << 8) | byte as u32; + #[inline] + fn write_quic(self, out: &mut Vec) -> Result<(), BoundsExceeded> { + let (hi, lo) = self.to_halves(); + if hi == 0 && lo < 1 << 6 { + out.push(lo as u8); + } else if hi == 0 && lo < 1 << 14 { + out.extend_from_slice(&(0x4000 | lo as u16).to_be_bytes()); + } else if hi == 0 && lo < 1 << 30 { + out.extend_from_slice(&(0x8000_0000 | lo).to_be_bytes()); + } else if hi < 1 << 30 { + let [a, b, c, d] = (0xc000_0000 | hi).to_be_bytes(); + let [e, f, g, h] = lo.to_be_bytes(); + out.extend_from_slice(&[a, b, c, d, e, f, g, h]); + } else { + return Err(BoundsExceeded); } - - Ok((Self::from_halves(hi, lo), len)) + Ok(()) } - /// The minimal encoding in the given form: the bytes, and how many of them are used. - /// - /// Fails past [`Self::MAX_QUIC`] in the QUIC form. - pub(super) fn encode_form(self, form: Form) -> Result<([u8; 9], usize), BoundsExceeded> { + #[inline] + fn write_leading_ones(self, out: &mut Vec) { let (hi, lo) = self.to_halves(); let [a, b, c, d] = lo.to_be_bytes(); - let [e, f, g, h] = hi.to_be_bytes(); - - Ok(match form { - Form::Quic if hi == 0 && lo < 1 << 6 => ([d, 0, 0, 0, 0, 0, 0, 0, 0], 1), - Form::Quic if hi == 0 && lo < 1 << 14 => ([0x40 | c, d, 0, 0, 0, 0, 0, 0, 0], 2), - Form::Quic if hi == 0 && lo < 1 << 30 => ([0x80 | a, b, c, d, 0, 0, 0, 0, 0], 4), - Form::Quic if hi < 1 << 30 => ([0xc0 | e, f, g, h, a, b, c, d, 0], 8), - Form::Quic => return Err(BoundsExceeded), - Form::LeadingOnes { .. } if hi == 0 && lo < 1 << 7 => ([d, 0, 0, 0, 0, 0, 0, 0, 0], 1), - Form::LeadingOnes { .. } if hi == 0 && lo < 1 << 14 => ([0x80 | c, d, 0, 0, 0, 0, 0, 0, 0], 2), - Form::LeadingOnes { .. } if hi == 0 && lo < 1 << 21 => ([0xc0 | b, c, d, 0, 0, 0, 0, 0, 0], 3), - Form::LeadingOnes { .. } if hi == 0 && lo < 1 << 28 => ([0xe0 | a, b, c, d, 0, 0, 0, 0, 0], 4), - Form::LeadingOnes { .. } if hi < 1 << 3 => ([0xf0 | h, a, b, c, d, 0, 0, 0, 0], 5), - Form::LeadingOnes { .. } if hi < 1 << 10 => ([0xf8 | g, h, a, b, c, d, 0, 0, 0], 6), + if hi == 0 && lo < 1 << 7 { + out.push(d); + } else if hi == 0 && lo < 1 << 14 { + out.extend_from_slice(&[0x80 | c, d]); + } else if hi == 0 && lo < 1 << 21 { + out.extend_from_slice(&[0xc0 | b, c, d]); + } else if hi == 0 && lo < 1 << 28 { + out.extend_from_slice(&[0xe0 | a, b, c, d]); + } else if hi < 1 << 3 { + out.extend_from_slice(&[0xf0 | hi as u8, a, b, c, d]); + } else if hi < 1 << 10 { + out.extend_from_slice(&[0xf8 | (hi >> 8) as u8, hi as u8, a, b, c, d]); + } else if hi < 1 << 24 { // The 7-byte form is skipped: one byte longer, but legal on every draft. - Form::LeadingOnes { .. } if hi < 1 << 24 => ([0xfe, f, g, h, a, b, c, d, 0], 8), - Form::LeadingOnes { .. } => ([0xff, e, f, g, h, a, b, c, d], 9), + let [_, f, g, h] = hi.to_be_bytes(); + out.extend_from_slice(&[0xfe, f, g, h, a, b, c, d]); + } else { + let [e, f, g, h] = hi.to_be_bytes(); + out.extend_from_slice(&[0xff, e, f, g, h, a, b, c, d]); + } + } + + /// Decode from the front of `buf`, returning the value and the rest of `buf`. + #[inline] + pub(super) fn read(buf: &[u8], form: Form) -> Result<(Self, &[u8]), DecodeError> { + match form { + Form::Quic => Self::read_quic(buf), + Form::LeadingOnes { seven } => Self::read_leading_ones(buf, seven), + } + } + + #[inline] + fn read_quic(buf: &[u8]) -> Result<(Self, &[u8]), DecodeError> { + let Some((&first, rest)) = buf.split_first() else { + return Err(DecodeError::Short); + }; + + let be = u32::from_be_bytes; + Ok(match first >> 6 { + 0 => (Self::from_u32(first as u32), rest), + 1 => { + let ([a, b], rest) = buf.split_first_chunk().ok_or(DecodeError::Short)?; + (Self::from_u32(be([0, 0, a & 0x3f, *b])), rest) + } + 2 => { + let ([a, b, c, d], rest) = buf.split_first_chunk().ok_or(DecodeError::Short)?; + (Self::from_u32(be([a & 0x3f, *b, *c, *d])), rest) + } + _ => { + let ([a, b, c, d, lo @ ..], rest) = buf.split_first_chunk::<8>().ok_or(DecodeError::Short)?; + (Self::from_halves(be([a & 0x3f, *b, *c, *d]), be(*lo)), rest) + } }) } - /// The bytes this takes on the wire in the given form, or [`BoundsExceeded`] if it - /// does not fit. - pub(crate) fn size(self, form: Form) -> Result { - Ok(self.encode_form(form)?.1) + #[inline] + fn read_leading_ones(buf: &[u8], seven: bool) -> Result<(Self, &[u8]), DecodeError> { + let Some((&first, rest)) = buf.split_first() else { + return Err(DecodeError::Short); + }; + + let be = u32::from_be_bytes; + Ok(match first.leading_ones() { + 0 => (Self::from_u32(first as u32), rest), + 1 => { + let ([a, b], rest) = buf.split_first_chunk().ok_or(DecodeError::Short)?; + (Self::from_u32(be([0, 0, a & 0x3f, *b])), rest) + } + 2 => { + let ([a, b, c], rest) = buf.split_first_chunk().ok_or(DecodeError::Short)?; + (Self::from_u32(be([0, a & 0x1f, *b, *c])), rest) + } + 3 => { + let ([a, b, c, d], rest) = buf.split_first_chunk().ok_or(DecodeError::Short)?; + (Self::from_u32(be([a & 0x0f, *b, *c, *d])), rest) + } + 4 => { + let ([a, lo @ ..], rest) = buf.split_first_chunk::<5>().ok_or(DecodeError::Short)?; + (Self::from_halves((a & 0x07) as u32, be(*lo)), rest) + } + 5 => { + let ([a, b, lo @ ..], rest) = buf.split_first_chunk::<6>().ok_or(DecodeError::Short)?; + (Self::from_halves(be([0, 0, a & 0x03, *b]), be(*lo)), rest) + } + // 1111110x: the 7-byte form, which draft-17 forbids. + 6 if !seven => return Err(DecodeError::InvalidValue), + 6 => { + let ([a, b, c, lo @ ..], rest) = buf.split_first_chunk::<7>().ok_or(DecodeError::Short)?; + (Self::from_halves(be([0, a & 0x01, *b, *c]), be(*lo)), rest) + } + 7 => { + let ([_, b, c, d, lo @ ..], rest) = buf.split_first_chunk::<8>().ok_or(DecodeError::Short)?; + (Self::from_halves(be([0, *b, *c, *d]), be(*lo)), rest) + } + _ => { + let ([_, hi @ .., e, f, g, h], rest) = buf.split_first_chunk::<9>().ok_or(DecodeError::Short)?; + (Self::from_halves(be(*hi), be([*e, *f, *g, *h])), rest) + } + }) } /// Decode a QUIC-style varint (2-bit length tag in top bits). @@ -303,18 +401,28 @@ impl VarInt { let mut buf = [0u8; 8]; r.copy_to_slice(&mut buf[..len]); - Ok(Self::decode_form(&buf[..len], Form::Quic)?.0) + Ok(Self::read(&buf[..len], Form::Quic)?.0) } /// Encode a QUIC-style varint (2-bit length tag in top bits). /// /// Fails with [`EncodeError::BoundsExceeded`] past [`Self::MAX_QUIC`]. pub fn encode_quic(&self, w: &mut W) -> Result<(), EncodeError> { - let (buf, len) = self.encode_form(Form::Quic)?; + let len = self.size(Form::Quic)?; if w.remaining_mut() < len { return Err(EncodeError::Short); } - w.put_slice(&buf[..len]); + + let (hi, lo) = self.to_halves(); + match len { + 1 => w.put_u8(lo as u8), + 2 => w.put_u16(0x4000 | lo as u16), + 4 => w.put_u32(0x8000_0000 | lo), + _ => { + w.put_u32(0xc000_0000 | hi); + w.put_u32(lo); + } + } Ok(()) } } @@ -339,8 +447,10 @@ mod tests { const DRAFT18: Form = Form::LeadingOnes { seven: true }; fn encode(value: VarInt, form: Form) -> Result, BoundsExceeded> { - let (buf, len) = value.encode_form(form)?; - Ok(buf[..len].to_vec()) + let mut buf = Vec::new(); + value.write(form, &mut buf)?; + assert_eq!(buf.len(), value.size(form)?, "size disagrees with the encoding"); + Ok(buf) } /// Test vectors from the draft-17 spec (Table 2: Example Integer Encodings), @@ -365,13 +475,13 @@ mod tests { ]; for (bytes, expected) in cases { - let (decoded, len) = VarInt::decode_form(bytes, DRAFT17).expect("decode should succeed"); + let (decoded, rest) = VarInt::read(bytes, DRAFT17).expect("decode should succeed"); assert_eq!( decoded.into_inner(), *expected, "decode mismatch for bytes {bytes:02x?}" ); - assert_eq!(len, bytes.len(), "all bytes should be consumed for {bytes:02x?}"); + assert!(rest.is_empty(), "all bytes should be consumed for {bytes:02x?}"); // Skip the non-minimal encoding (0x8025 for 37); we only emit the minimal one. if bytes.len() == 1 || *expected != 37 { @@ -385,7 +495,7 @@ mod tests { #[test] fn leading_ones_invalid_0xfc() { assert!( - matches!(VarInt::decode_form(&[0xFC], DRAFT17), Err(DecodeError::InvalidValue)), + matches!(VarInt::read(&[0xFC], DRAFT17), Err(DecodeError::InvalidValue)), "0xFC should be rejected as invalid on draft-17" ); } @@ -409,7 +519,7 @@ mod tests { "unexpected encoded length for value {value}" ); - let (decoded, _) = VarInt::decode_form(&encoded, DRAFT17).expect("leading-ones decode should succeed"); + let (decoded, _) = VarInt::read(&encoded, DRAFT17).expect("leading-ones decode should succeed"); assert_eq!(decoded.into_inner(), value, "round-trip mismatch for value {value}"); } } @@ -427,8 +537,8 @@ mod tests { assert!(value > VarInt::MAX_QUIC); continue; }; - let (decoded, len) = VarInt::decode_form(&encoded, form).unwrap(); - assert_eq!((decoded, len), (value, encoded.len()), "{form:?} {value}"); + let (decoded, rest) = VarInt::read(&encoded, form).unwrap(); + assert_eq!((decoded, rest), (value, &[][..]), "{form:?} {value}"); } } } @@ -453,14 +563,14 @@ mod tests { for value in [(1u64 << 62) - 1, 1u64 << 62, u64::MAX] { let encoded = encode(VarInt::from(value), DRAFT18).unwrap(); assert_eq!(encoded.len(), 9); - assert_eq!(VarInt::decode_form(&encoded, DRAFT18).unwrap().0.into_inner(), value); + assert_eq!(VarInt::read(&encoded, DRAFT18).unwrap().0.into_inner(), value); } } #[test] fn draft17_rejects_7_byte_varint() { // 1111110x prefix: invalid on draft-17. - let err = VarInt::decode_form(&[0xFC, 0, 0, 0, 0, 0, 0], DRAFT17).unwrap_err(); + let err = VarInt::read(&[0xFC, 0, 0, 0, 0, 0, 0], DRAFT17).unwrap_err(); assert!(matches!(err, DecodeError::InvalidValue)); } @@ -520,7 +630,7 @@ mod tests { for shift in (0..48).step_by(8).rev() { bytes.push(((value >> shift) & 0xFF) as u8); } - let (decoded, _) = VarInt::decode_form(&bytes, DRAFT18).unwrap(); + let (decoded, _) = VarInt::read(&bytes, DRAFT18).unwrap(); assert_eq!(decoded.into_inner(), value); } diff --git a/rs/moq-net/src/ietf/fetch.rs b/rs/moq-net/src/ietf/fetch.rs index 5a126ee80f..e9d1db939d 100644 --- a/rs/moq-net/src/ietf/fetch.rs +++ b/rs/moq-net/src/ietf/fetch.rs @@ -41,7 +41,7 @@ impl Encode for FetchType<'_> { start, end, } => { - w.u8(1u8); + w.u8(1); encode_namespace(w, namespace)?; w.string(track)?; start.encode(w, version)?; @@ -51,7 +51,7 @@ impl Encode for FetchType<'_> { subscriber_request_id, group_offset, } => { - w.u8(2u8); + w.u8(2); subscriber_request_id.encode(w, version)?; w.varint(VarInt::from(*group_offset))?; } @@ -59,7 +59,7 @@ impl Encode for FetchType<'_> { subscriber_request_id, group_id, } => { - w.u8(3u8); + w.u8(3); subscriber_request_id.encode(w, version)?; w.varint(VarInt::from(*group_id))?; } @@ -119,7 +119,7 @@ impl Message for Fetch<'_> { fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.request_id.encode(w, version)?; if version == Version::Draft17 { - w.varint(VarInt::from(0u64))?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) + w.varint(VarInt::ZERO)?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) } match version { @@ -127,7 +127,7 @@ impl Message for Fetch<'_> { w.u8(self.subscriber_priority); self.group_order.encode(w, version)?; self.fetch_type.encode(w, version)?; - w.u8(0u8); // no parameters + w.u8(0); // no parameters } _ => { self.fetch_type.encode(w, version)?; @@ -204,7 +204,7 @@ impl Message for FetchOk { self.group_order.encode(w, version)?; w.bool(self.end_of_track); self.end_location.encode(w, version)?; - w.u8(0u8); // no parameters + w.u8(0); // no parameters } _ => { // GROUP_ORDER is not a legal FETCH_OK parameter in any draft after 14; the order diff --git a/rs/moq-net/src/ietf/filter.rs b/rs/moq-net/src/ietf/filter.rs index 5aef4ef0f6..3ea0d99a2a 100644 --- a/rs/moq-net/src/ietf/filter.rs +++ b/rs/moq-net/src/ietf/filter.rs @@ -86,8 +86,8 @@ impl Filter { // that spelling normalizes to this one rather than colliding with NextObject. Self::Unfiltered => {} Self::NextObject => { - w.varint(VarInt::from(0u64))?; - w.varint(VarInt::from(0u64))?; + w.varint(VarInt::ZERO)?; + w.varint(VarInt::ZERO)?; } Self::Relative(groups) => w.varint(VarInt::from(groups))?, Self::Absolute { @@ -555,7 +555,7 @@ impl Param for Fill { // An omitted filter inherits the subscription's, so the scope is empty. An explicit // Unfiltered still encodes, as a zero-length filter meaning the whole track. match self.filter { - None => inner.varint(VarInt::from(0u64))?, + None => inner.varint(VarInt::ZERO)?, Some(filter) => { inner.varint(VarInt::from(1u64))?; // The first type in a scope is not delta encoded, so this is the raw id. @@ -690,7 +690,7 @@ mod fill_tests { Encoder::new(&mut value, NEW.into()) .varint(VarInt::from(0x10u64)) .unwrap(); // FORWARD, not allowed in a fill - Encoder::new(&mut value, NEW.into()).varint(VarInt::from(0u64)).unwrap(); + Encoder::new(&mut value, NEW.into()).varint(VarInt::ZERO).unwrap(); let mut buf = Vec::new(); Encoder::new(&mut buf, NEW.into()).bytes(&value).unwrap(); diff --git a/rs/moq-net/src/ietf/goaway.rs b/rs/moq-net/src/ietf/goaway.rs index e19d0fe0ae..0812742d9f 100644 --- a/rs/moq-net/src/ietf/goaway.rs +++ b/rs/moq-net/src/ietf/goaway.rs @@ -32,7 +32,7 @@ impl Message for GoAway<'_> { // conformant peer must treat as a PROTOCOL_VIOLATION. Draft-19 // removed the field again (#1623). if matches!(version, Version::Draft18) { - w.varint(VarInt::from(0u64))?; + w.varint(VarInt::ZERO)?; } Ok(()) } diff --git a/rs/moq-net/src/ietf/message.rs b/rs/moq-net/src/ietf/message.rs index a2c4439288..05526e2530 100644 --- a/rs/moq-net/src/ietf/message.rs +++ b/rs/moq-net/src/ietf/message.rs @@ -18,9 +18,9 @@ pub trait Message: Sized + std::fmt::Debug { impl Encode for T { fn encode(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { tracing::trace!(?self, "encoding"); - let start = w.position(); + let prefix = w.prefix_u16(); self.encode_msg(w, version)?; - w.prefix_u16(start) + w.fill(prefix) } } diff --git a/rs/moq-net/src/ietf/parameters.rs b/rs/moq-net/src/ietf/parameters.rs index 9a28fd417c..322366215b 100644 --- a/rs/moq-net/src/ietf/parameters.rs +++ b/rs/moq-net/src/ietf/parameters.rs @@ -910,7 +910,7 @@ mod tests { // Delta-encoded: first delta=0x20 (abs=0x20), second delta=0 (abs=0x20) w.varint(VarInt::from(0x20u64)).unwrap(); 100u8.param_encode(&mut w, version).unwrap(); - w.varint(VarInt::from(0u64)).unwrap(); + w.varint(VarInt::ZERO).unwrap(); 200u8.param_encode(&mut w, version).unwrap(); } } diff --git a/rs/moq-net/src/ietf/properties.rs b/rs/moq-net/src/ietf/properties.rs index 11ca7512dc..3f8dc33d0f 100644 --- a/rs/moq-net/src/ietf/properties.rs +++ b/rs/moq-net/src/ietf/properties.rs @@ -304,7 +304,7 @@ mod tests { .varint(VarInt::from(0x22u64)) .unwrap(); Encoder::new(&mut buf, Version::Draft18.into()) - .varint(VarInt::from(0u64)) + .varint(VarInt::ZERO) .unwrap(); let mut bytes = bytes::Bytes::from(buf); diff --git a/rs/moq-net/src/ietf/publish.rs b/rs/moq-net/src/ietf/publish.rs index a71ffd615b..54fd642457 100644 --- a/rs/moq-net/src/ietf/publish.rs +++ b/rs/moq-net/src/ietf/publish.rs @@ -243,7 +243,7 @@ impl Message for Publish<'_> { fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.request_id.encode(w, version)?; if version == Version::Draft17 { - w.varint(VarInt::from(0u64))?; // required_request_id_delta = 0 + w.varint(VarInt::ZERO)?; // required_request_id_delta = 0 } encode_namespace(w, &self.track_namespace)?; w.string(&self.track_name)?; @@ -264,7 +264,7 @@ impl Message for Publish<'_> { w.bool(self.forward); // parameters - w.u8(0u8); + w.u8(0); } _ => { // GROUP_ORDER is a legal PUBLISH parameter only through draft-15; a later peer @@ -404,7 +404,7 @@ impl Message for PublishOk { // decode, so encoding one would truncate the message. self.filter.encode(w, version)?; // no parameters - w.u8(0u8); + w.u8(0); } // Draft-20 moved the subscription parameters out of PUBLISH_OK; they belong to // PUBLISH and REQUEST_UPDATE now, so a PUBLISH_OK carries none of them. diff --git a/rs/moq-net/src/ietf/publish_namespace.rs b/rs/moq-net/src/ietf/publish_namespace.rs index 3ec5793692..5239f8bda5 100644 --- a/rs/moq-net/src/ietf/publish_namespace.rs +++ b/rs/moq-net/src/ietf/publish_namespace.rs @@ -52,7 +52,7 @@ impl Message for PublishNamespace<'_> { fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.request_id.encode(w, version)?; if version == Version::Draft17 { - w.varint(VarInt::from(0u64))?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) + w.varint(VarInt::ZERO)?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) } encode_namespace(w, &self.track_namespace)?; encode_cluster_params(w, version, self.cluster.as_ref()) @@ -106,7 +106,7 @@ impl Message for PublishNamespaceUpdate { Version::Draft14 | Version::Draft15 | Version::Draft16 => return Err(EncodeError::Version), Version::Draft17 => { self.request_id.encode(w, version)?; - w.varint(VarInt::from(0u64))?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) + w.varint(VarInt::ZERO)?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) } _ => self.request_id.encode(w, version)?, } diff --git a/rs/moq-net/src/ietf/publisher.rs b/rs/moq-net/src/ietf/publisher.rs index 75c49ed782..661a371f3a 100644 --- a/rs/moq-net/src/ietf/publisher.rs +++ b/rs/moq-net/src/ietf/publisher.rs @@ -874,7 +874,7 @@ where .await?; stream.encode(&VarInt::from(frame.payload.len())).await?; if frame.payload.is_empty() && matches!(version, Version::Draft14 | Version::Draft15) { - stream.encode(&VarInt::from(0u64)).await?; + stream.encode(&VarInt::ZERO).await?; } if !frame.payload.is_empty() { let mut payload = frame.payload; @@ -921,7 +921,7 @@ where stream.encode(&VarInt::from(frame.size)).await?; if frame.size == 0 && matches!(version, Version::Draft14 | Version::Draft15) { - stream.encode(&VarInt::from(0u64)).await?; + stream.encode(&VarInt::ZERO).await?; } loop { let chunk = { @@ -981,7 +981,7 @@ where if version == Version::Draft14 { let properties = properties.unwrap_or_default(); stream.buffer(&VarInt::from(sequence))?; - stream.buffer(&VarInt::from(0u64))?; + stream.buffer(&VarInt::ZERO)?; stream.buffer(&VarInt::from(object))?; // Publisher priority, a raw byte. stream.buffer_raw(&[0]); @@ -1226,7 +1226,7 @@ where .await?; writer.encode(&VarInt::from(frame.payload.len())).await?; if frame.payload.is_empty() && matches!(self.version, Version::Draft14 | Version::Draft15) { - writer.encode(&VarInt::from(0u64)).await?; + writer.encode(&VarInt::ZERO).await?; } if !frame.payload.is_empty() { let mut payload = frame.payload; @@ -2142,8 +2142,8 @@ impl TrackServe { flags: ietf::GroupFlags::default(), })?; // Object ID delta 0, then an empty object whose status is END_OF_TRACK. - writer.buffer(&VarInt::from(0u64))?; - writer.buffer(&VarInt::from(0u64))?; + writer.buffer(&VarInt::ZERO)?; + writer.buffer(&VarInt::ZERO)?; writer.encode(&VarInt::from(END_OF_TRACK)).await?; // PUBLISH_DONE follows once this closes, like every other data stream. writer.close().await diff --git a/rs/moq-net/src/ietf/subscribe.rs b/rs/moq-net/src/ietf/subscribe.rs index dcfd43310f..2241d82b72 100644 --- a/rs/moq-net/src/ietf/subscribe.rs +++ b/rs/moq-net/src/ietf/subscribe.rs @@ -152,7 +152,7 @@ impl Message for Subscribe<'_> { fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.request_id.encode(w, version)?; if version == Version::Draft17 { - w.varint(VarInt::from(0u64))?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) + w.varint(VarInt::ZERO)?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) } encode_namespace(w, &self.track_namespace)?; w.string(&self.track_name)?; @@ -164,7 +164,7 @@ impl Message for Subscribe<'_> { w.bool(true); // forward self.filter.encode(w, version)?; - w.u8(0u8); // no parameters + w.u8(0); // no parameters } _ => { // FILL_PARAMETERS arrived in draft-20. Sending it to an older peer would be an @@ -222,7 +222,7 @@ impl Message for SubscribeOk { match version { Version::Draft14 => { - w.varint(VarInt::from(0u64))?; // expires = 0 + w.varint(VarInt::ZERO)?; // expires = 0 self.properties .group_order .unwrap_or(GroupOrder::Ascending) @@ -231,7 +231,7 @@ impl Message for SubscribeOk { if let Some(largest) = self.largest { largest.encode(w, version)?; } - w.u8(0u8); // no parameters + w.u8(0); // no parameters } _ => { // GROUP_ORDER is a legal SUBSCRIBE_OK parameter only through draft-15; a later @@ -396,7 +396,7 @@ impl Message for SubscribeUpdate { w.varint(VarInt::from(self.end_group))?; w.u8(self.subscriber_priority); w.bool(self.forward); - w.u8(0u8); // no parameters + w.u8(0); // no parameters } Version::Draft15 | Version::Draft16 => { self.request_id.encode(w, version)?; @@ -417,7 +417,7 @@ impl Message for SubscribeUpdate { // REQUEST_UPDATE self.request_id.encode(w, version)?; if matches!(version, Version::Draft17) { - w.varint(VarInt::from(0u64))?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) + w.varint(VarInt::ZERO)?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) } encode_params!(w, version, 0x10 => self.forward, diff --git a/rs/moq-net/src/ietf/subscribe_namespace.rs b/rs/moq-net/src/ietf/subscribe_namespace.rs index a93ad85b6b..35ec81655b 100644 --- a/rs/moq-net/src/ietf/subscribe_namespace.rs +++ b/rs/moq-net/src/ietf/subscribe_namespace.rs @@ -111,7 +111,7 @@ impl Message for SubscribeNamespaceLegacy<'_> { } self.request_id.encode(w, version)?; if version == Version::Draft17 { - w.varint(VarInt::from(0u64))?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) + w.varint(VarInt::ZERO)?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) } encode_namespace(w, &self.namespace)?; if matches!(version, Version::Draft16 | Version::Draft17) { @@ -398,9 +398,7 @@ mod tests { let mut buf = Vec::new(); encode_namespace(&mut Encoder::new(&mut buf, version.into()), &Path::new("a")).unwrap(); // Number of Parameters = 0. - Encoder::new(&mut buf, version.into()) - .varint(VarInt::from(0u64)) - .unwrap(); + Encoder::new(&mut buf, version.into()).varint(VarInt::ZERO).unwrap(); let mut bytes = bytes::Bytes::from(buf); assert!(matches!( diff --git a/rs/moq-net/src/ietf/track.rs b/rs/moq-net/src/ietf/track.rs index 9aabea1f0a..a066d04d96 100644 --- a/rs/moq-net/src/ietf/track.rs +++ b/rs/moq-net/src/ietf/track.rs @@ -31,18 +31,18 @@ impl Message for TrackStatus<'_> { fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.request_id.encode(w, version)?; if version == Version::Draft17 { - w.varint(VarInt::from(0u64))?; // required_request_id_delta = 0 + w.varint(VarInt::ZERO)?; // required_request_id_delta = 0 } encode_namespace(w, &self.track_namespace)?; w.string(&self.track_name)?; match version { Version::Draft14 => { - w.u8(0u8); // subscriber priority + w.u8(0); // subscriber priority GroupOrder::Descending.encode(w, version)?; w.bool(false); // forward Filter::NextObject.encode(w, version)?; // filter - w.u8(0u8); // no parameters + w.u8(0); // no parameters } _ => { encode_params!(w, version,); diff --git a/rs/moq-net/src/lite/announce.rs b/rs/moq-net/src/lite/announce.rs index 33de5dab76..3cf12efd2c 100644 --- a/rs/moq-net/src/lite/announce.rs +++ b/rs/moq-net/src/lite/announce.rs @@ -223,7 +223,7 @@ impl Encode for AnnounceBroadcast<'_> { }; w.varint(VarInt::from(typ))?; - let start = w.position(); + let prefix = w.prefix_varint(); match self { Self::Active { suffix, hops, cost } => { suffix.encode(w, version)?; @@ -238,12 +238,12 @@ impl Encode for AnnounceBroadcast<'_> { } Self::Ended { .. } | Self::Skipped => unreachable!("refused above"), } - return w.prefix_varint(start); + return w.fill(prefix); } // Older versions: a single ANNOUNCE_BROADCAST message, size-prefixed, with the // status carried inside the body. - let start = w.position(); + let prefix = w.prefix_varint(); match self { // The cost is a lite-06 addition, so it is simply not on the wire here. Self::Active { suffix, hops, .. } => { @@ -265,7 +265,7 @@ impl Encode for AnnounceBroadcast<'_> { return Err(EncodeError::Version); } } - w.prefix_varint(start) + w.fill(prefix) } } diff --git a/rs/moq-net/src/lite/message.rs b/rs/moq-net/src/lite/message.rs index 3133665543..574c6d5487 100644 --- a/rs/moq-net/src/lite/message.rs +++ b/rs/moq-net/src/lite/message.rs @@ -32,9 +32,9 @@ pub trait Message: Sized + std::fmt::Debug { impl Encode for T { fn encode(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { tracing::trace!(?self, "encoding"); - let start = w.position(); + let prefix = w.prefix_varint(); self.encode_msg(w, version)?; - w.prefix_varint(start) + w.fill(prefix) } } diff --git a/rs/moq-net/src/lite/subscribe.rs b/rs/moq-net/src/lite/subscribe.rs index 0b8c730b66..2ef7c983cf 100644 --- a/rs/moq-net/src/lite/subscribe.rs +++ b/rs/moq-net/src/lite/subscribe.rs @@ -120,7 +120,7 @@ pub(super) fn skip_group_order(r: &mut Decoder<'_>, version: Version) -> Result< /// Write the retired `Ordered` byte as 0, keeping a deployed version's field offsets. pub(super) fn pad_group_order(w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { if version.has_group_order() { - w.u8(0u8); + w.u8(0); } Ok(()) } diff --git a/rs/moq-net/src/setup.rs b/rs/moq-net/src/setup.rs index 529a26c652..f5be4f531c 100644 --- a/rs/moq-net/src/setup.rs +++ b/rs/moq-net/src/setup.rs @@ -54,9 +54,9 @@ impl Encode for Setup { fn encode(&self, w: &mut Encoder<'_>, v: Version) -> Result<(), EncodeError> { Self::check_version(v); w.varint(VarInt::from(SETUP_V17))?; - let start = w.position(); + let prefix = w.prefix_u16(); w.slice(&self.parameters); - w.prefix_u16(start) + w.fill(prefix) } } @@ -154,9 +154,9 @@ impl Encode for Client { /// Encode a client setup message (draft-14 through draft-16 only). fn encode(&self, w: &mut Encoder<'_>, v: Version) -> Result<(), EncodeError> { w.u8(CLIENT_SETUP); - let start = w.position(); + let prefix = prefix_body(w, v)?; self.encode_inner(w, v)?; - prefix_body(w, v, start) + w.fill(prefix) } } @@ -170,11 +170,11 @@ fn decode_body<'a>(r: &mut Decoder<'a>, v: Version) -> Result, Decod r.sub(size) } -/// Size-prefix a pre-draft-17 SETUP body written since `start`. -fn prefix_body(w: &mut Encoder<'_>, v: Version, start: usize) -> Result<(), EncodeError> { +/// Reserve the size prefix of a pre-draft-17 SETUP body. +fn prefix_body(w: &mut Encoder<'_>, v: Version) -> Result { match SetupVersion::from_version(v) { - SetupVersion::Draft14 | SetupVersion::Draft15Plus => w.prefix_u16(start), - SetupVersion::LiteLegacy => w.prefix_varint(start), + SetupVersion::Draft14 | SetupVersion::Draft15Plus => Ok(w.prefix_u16()), + SetupVersion::LiteLegacy => Ok(w.prefix_varint()), SetupVersion::Modern | SetupVersion::Unsupported => Err(EncodeError::Version), } } @@ -207,9 +207,9 @@ impl Encode for Server { /// Encode a server setup message (draft-14 through draft-16 only). fn encode(&self, w: &mut Encoder<'_>, v: Version) -> Result<(), EncodeError> { w.u8(SERVER_SETUP); - let start = w.position(); + let prefix = prefix_body(w, v)?; self.encode_inner(w, v)?; - prefix_body(w, v, start) + w.fill(prefix) } } From 0684eb3fa2639c04ad0872e5765325f2a8706925 Mon Sep 17 00:00:00 2001 From: Luke Curley Date: Mon, 28 Sep 2026 21:56:22 -0700 Subject: [PATCH 5/8] refactor(net)!: drop the VarInt newtype for plain u64 Varint is a wire encoding, not a type. Encoder::varint and Decoder::varint take and return u64, the stream Reader/Writer gain varint methods, and moq_net::varint exposes MAX_QUIC plus the QUIC Buf helpers other crates use. Co-Authored-By: Claude Opus 5.5 --- quest/m1/rs2ts/README.md | 13 +- quest/m1/rs2ts/ietf.md | 2 +- quest/m1/rs2ts/lite.md | 2 +- quest/m1/rs2ts/translator.md | 16 +- rs/hang/src/catalog/container.rs | 2 +- rs/hang/src/container/frame.rs | 15 +- rs/moq-archive/src/segment.rs | 6 +- rs/moq-c/src/api.rs | 2 +- rs/moq-c/src/video.rs | 2 +- rs/moq-ffi/src/video.rs | 2 +- rs/moq-loc/src/lib.rs | 20 +- rs/moq-mux/src/catalog/hang/container.rs | 2 +- rs/moq-mux/src/catalog/msf/consumer.rs | 2 +- rs/moq-mux/src/container/legacy/mod.rs | 2 +- rs/moq-net/src/coding/decode.rs | 10 +- rs/moq-net/src/coding/encode.rs | 20 +- rs/moq-net/src/coding/mod.rs | 7 +- rs/moq-net/src/coding/reader.rs | 130 ++-- rs/moq-net/src/coding/varint.rs | 654 ++++++++------------- rs/moq-net/src/coding/version.rs | 8 +- rs/moq-net/src/coding/writer.rs | 38 +- rs/moq-net/src/fuzz.rs | 29 +- rs/moq-net/src/ietf/adapter.rs | 14 +- rs/moq-net/src/ietf/fetch.rs | 46 +- rs/moq-net/src/ietf/filter.rs | 80 ++- rs/moq-net/src/ietf/goaway.rs | 18 +- rs/moq-net/src/ietf/group.rs | 32 +- rs/moq-net/src/ietf/location.rs | 10 +- rs/moq-net/src/ietf/namespace.rs | 6 +- rs/moq-net/src/ietf/parameters.rs | 66 +-- rs/moq-net/src/ietf/properties.rs | 88 +-- rs/moq-net/src/ietf/publish.rs | 26 +- rs/moq-net/src/ietf/publish_namespace.rs | 16 +- rs/moq-net/src/ietf/publisher.rs | 117 ++-- rs/moq-net/src/ietf/request.rs | 14 +- rs/moq-net/src/ietf/session.rs | 38 +- rs/moq-net/src/ietf/subscribe.rs | 34 +- rs/moq-net/src/ietf/subscribe_namespace.rs | 14 +- rs/moq-net/src/ietf/subscriber.rs | 138 ++--- rs/moq-net/src/ietf/token.rs | 24 +- rs/moq-net/src/ietf/track.rs | 10 +- rs/moq-net/src/lib.rs | 2 +- rs/moq-net/src/lite/announce.rs | 62 +- rs/moq-net/src/lite/compress.rs | 6 +- rs/moq-net/src/lite/datagram.rs | 14 +- rs/moq-net/src/lite/fetch.rs | 10 +- rs/moq-net/src/lite/goaway.rs | 2 +- rs/moq-net/src/lite/group.rs | 12 +- rs/moq-net/src/lite/info.rs | 4 +- rs/moq-net/src/lite/message.rs | 18 +- rs/moq-net/src/lite/parameters.rs | 14 +- rs/moq-net/src/lite/probe.rs | 8 +- rs/moq-net/src/lite/publisher.rs | 4 +- rs/moq-net/src/lite/setup.rs | 8 +- rs/moq-net/src/lite/stream.rs | 8 +- rs/moq-net/src/lite/subscribe.rs | 60 +- rs/moq-net/src/lite/subscriber.rs | 8 +- rs/moq-net/src/lite/track.rs | 16 +- rs/moq-net/src/model/origin.rs | 10 +- rs/moq-net/src/model/time.rs | 36 +- rs/moq-net/src/setup.rs | 8 +- 61 files changed, 947 insertions(+), 1138 deletions(-) diff --git a/quest/m1/rs2ts/README.md b/quest/m1/rs2ts/README.md index 3f873a4f59..111d37e856 100644 --- a/quest/m1/rs2ts/README.md +++ b/quest/m1/rs2ts/README.md @@ -30,16 +30,17 @@ Decided in planning (2026-09-27), with the spike data in and bytes out, no runtime. The async helper methods move behind an `async` cargo feature; rs2ts reads the crate without it and JS reimplements the helpers with Promises. No second crate. -- `VarInt` holds the full 64-bit range; the spec is not bounded to 2^53. The - leading-ones form (moq-transport draft-17+) carries all of it, and the QUIC - form (moq-lite, drafts 14-16) refuses anything past 2^62 - 1 rather than - truncating. Rust's `VarInt` newtype carries Encode/Decode and JS gets a - matching `VarInt` type with checked conversion to and from `number`. +- Values are plain `u64` in Rust, and varint is a wire encoding in the codec, + not a type. The spec is not bounded to 2^53: the leading-ones form + (moq-transport draft-17+) carries all 64 bits, and the QUIC form (moq-lite, + drafts 14-16) refuses anything past 2^62 - 1 rather than truncating. Rust + `u64` maps to a TypeScript `U64` (two `u32` halves), generically, with + checked conversion to and from `number`. - The generated TypeScript is committed and a CI lane regenerates it and fails on drift, so JS contributors and npm publishing never need the nightly toolchain Charon pins. It lives inside js/net and `@moq/net` stays the package. -- The `@moq/net` API may change (disposable handles, `VarInt`) as long as it +- The `@moq/net` API may change (disposable handles, `U64`) as long as it is no worse to use; watch, publish, hang, and the demos update in the same change. - Parity: `just test interop --all`, plus moq-net's own tests translated with diff --git a/quest/m1/rs2ts/ietf.md b/quest/m1/rs2ts/ietf.md index a037e66a5f..3de2f1ec70 100644 --- a/quest/m1/rs2ts/ietf.md +++ b/quest/m1/rs2ts/ietf.md @@ -9,7 +9,7 @@ the hand-written js/net IETF code (about 8.7k lines) is deleted, with ## Plan Values above 2^53 are legal on the IETF wire (request ids, track aliases); -they stay exact as `VarInt` and only fail where code converts them to +they stay exact as `U64` and only fail where code converts them to `number`. Public API: breaks `@moq/net`; retargets to `dev`. Wire: none. diff --git a/quest/m1/rs2ts/lite.md b/quest/m1/rs2ts/lite.md index 256dd08783..f42bccd742 100644 --- a/quest/m1/rs2ts/lite.md +++ b/quest/m1/rs2ts/lite.md @@ -12,7 +12,7 @@ first-frame latency are no worse than the hand-written js/net. ## Plan - The `@moq/net` API may change where the generated shape is no worse to - use: disposable handles (`using`), `VarInt` for sequences and ids. Update + use: disposable handles (`using`), `U64` for sequences and ids. Update watch, publish, hang, room, and the demos in the same change, and the `doc/` pages for anything user-facing. - A forgotten `drop()` leaves a track open forever: add a debug-only diff --git a/quest/m1/rs2ts/translator.md b/quest/m1/rs2ts/translator.md index 071409d9ba..c0ba400843 100644 --- a/quest/m1/rs2ts/translator.md +++ b/quest/m1/rs2ts/translator.md @@ -29,12 +29,12 @@ Mapping decided in planning: `[Symbol.dispose]`; `Arc`/`Rc` of a type with drop glue become an explicit refcount. JS is single-threaded, so `Mutex` and atomics become plain access. -- Rust `VarInt` maps to the [JS VarInt](/quest/m1/rs2ts/js-varint.md) type. - Integers up to 32 bits and `usize` map to `number` with checked arithmetic - that throws on overflow; never wrap silently. A `u64` or `i64` never maps - to a lossy `number`: the model accepts `u64::MAX` (e.g. - `model/subscription.rs`), so each one either becomes `VarInt` or an - `Option` in the source, or maps to a full-width 64-bit TypeScript type. +- Rust `u64` maps to a TypeScript `U64` (two `u32` halves), generically; see + [JS U64](/quest/m1/rs2ts/js-varint.md). Integers up to 32 bits and `usize` + map to `number` with checked arithmetic that throws on overflow; never wrap + silently. A `u64` or `i64` never maps to a lossy `number`: the model accepts + `u64::MAX` (e.g. `model/subscription.rs`). Varint is a wire encoding in the + codec, not a type, so nothing maps by that name. Guidance: @@ -42,7 +42,7 @@ Guidance: maps by Rust type and construct (e.g. `u64` to one 64-bit TypeScript type), never by moq-net names; anything project-specific lives in moq-net's source or a small config, so another crate could use rs2ts unchanged. -- `VarInt::from_zigzag` and `to_zigzag` (lite per-frame timestamps) still use +- `varint::zigzag` and `unzigzag` (lite per-frame timestamps) still use 64-bit bit math, which the subset forbids. Rewrite them on two `u32` halves, or give the 64-bit TypeScript type the operations they need. - Pin Charon and its nightly in the nix shell for the regeneration lane only. @@ -63,4 +63,4 @@ translates. Wire: none. ## Required -- [JS VarInt](/quest/m1/rs2ts/js-varint.md) - the TypeScript type `VarInt` maps to +- [JS U64](/quest/m1/rs2ts/js-varint.md) - the TypeScript `U64` that Rust `u64` maps to diff --git a/rs/hang/src/catalog/container.rs b/rs/hang/src/catalog/container.rs index 126f42601a..2466f3dbd2 100644 --- a/rs/hang/src/catalog/container.rs +++ b/rs/hang/src/catalog/container.rs @@ -15,7 +15,7 @@ use serde_with::{base64::Base64, serde_as}; /// rendition must be ignored by consumers. #[derive(Debug, Clone, PartialEq, Default)] pub enum Container { - /// A QUIC VarInt timestamp prefix followed by the raw codec payload. + /// A QUIC varint timestamp prefix followed by the raw codec payload. /// Timestamps are in microseconds. #[default] Legacy, diff --git a/rs/hang/src/container/frame.rs b/rs/hang/src/container/frame.rs index b971af1fc5..384f2ed2a3 100644 --- a/rs/hang/src/container/frame.rs +++ b/rs/hang/src/container/frame.rs @@ -1,7 +1,6 @@ use super::MAX_AGE; use bytes::{Buf, BufMut, Bytes, BytesMut}; use derive_more::Debug; -use moq_net::VarInt; use crate::Error; @@ -9,7 +8,7 @@ pub use moq_net::{Timescale, Timestamp}; /// Canonical timescale for the hang legacy wire format: microseconds. /// -/// The legacy container's on-wire timestamp is a single VarInt with no scale tag, +/// The legacy container's on-wire timestamp is a single varint with no scale tag, /// so encoders normalize to this scale and decoders attach it. pub const TIMESCALE: Timescale = Timescale::MICRO; @@ -67,7 +66,7 @@ pub struct Frame { } impl Frame { - /// Encode the frame: VarInt timestamp prefix followed by the raw codec payload. + /// Encode the frame: varint timestamp prefix followed by the raw codec payload. /// /// The timestamp is normalized to [`TIMESCALE`] (microseconds) so peers using a /// different source scale (e.g. nanoseconds from MKV) can decode without knowing @@ -78,12 +77,12 @@ impl Frame { Ok(()) } - /// Decode a frame from raw bytes (VarInt timestamp prefix + payload). + /// Decode a frame from raw bytes (varint timestamp prefix + payload). /// /// Attaches [`TIMESCALE`] (microseconds) to the decoded timestamp, matching what /// [`Self::encode`] writes. Inverse of [`Self::encode`]. pub fn decode(mut buf: impl Buf) -> Result { - let value: u64 = VarInt::decode_quic(&mut buf).map_err(moq_net::Error::from)?.into(); + let value: u64 = moq_net::varint::decode_quic(&mut buf).map_err(moq_net::Error::from)?; let timestamp = Timestamp::new(value, TIMESCALE)?; let payload = buf.copy_to_bytes(buf.remaining()); @@ -117,12 +116,10 @@ impl Frame { Ok(()) } - /// Write the VarInt timestamp prefix, normalized to [`TIMESCALE`]. + /// Write the varint timestamp prefix, normalized to [`TIMESCALE`]. fn encode_header(&self, buf: &mut impl BufMut) -> Result<(), Error> { let timestamp = self.timestamp.convert(TIMESCALE)?; - VarInt::from(timestamp.value()) - .encode_quic(buf) - .map_err(moq_net::Error::from)?; + moq_net::varint::encode_quic(timestamp.value(), buf).map_err(moq_net::Error::from)?; Ok(()) } diff --git a/rs/moq-archive/src/segment.rs b/rs/moq-archive/src/segment.rs index ceacc54238..9b1cac5a89 100644 --- a/rs/moq-archive/src/segment.rs +++ b/rs/moq-archive/src/segment.rs @@ -1,7 +1,7 @@ use std::ops::RangeInclusive; use bytes::{Buf, BufMut, Bytes, BytesMut}; -use moq_net::VarInt; +use moq_net::varint; use crate::path::{check_id, check_range}; use crate::{Error, Result, VERSION}; @@ -203,11 +203,11 @@ fn validate(groups: &[Group]) -> Result<()> { } fn write_varint(buf: &mut impl BufMut, value: u64) -> Result<()> { - VarInt::from(value).encode_quic(buf).map_err(|_| Error::Overflow) + varint::encode_quic(value, buf).map_err(|_| Error::Overflow) } fn read_varint(buf: &mut impl Buf) -> Result { - Ok(VarInt::decode_quic(buf).map_err(|_| Error::Table)?.into_inner()) + varint::decode_quic(buf).map_err(|_| Error::Table) } fn read_count(buf: &mut impl Buf, min_entry: usize) -> Result { diff --git a/rs/moq-c/src/api.rs b/rs/moq-c/src/api.rs index 7639b60b43..f397dc41ef 100644 --- a/rs/moq-c/src/api.rs +++ b/rs/moq-c/src/api.rs @@ -15,7 +15,7 @@ use tracing::Level; #[allow(non_camel_case_types)] #[derive(Clone, Copy, Debug)] pub enum moq_container_kind { - /// A QUIC VarInt timestamp prefix followed by the raw codec payload. + /// A QUIC varint timestamp prefix followed by the raw codec payload. /// Timestamps are in microseconds. MOQ_CONTAINER_KIND_LEGACY = 0, /// Fragmented MP4: each frame is a complete moof+mdat fragment, described by diff --git a/rs/moq-c/src/video.rs b/rs/moq-c/src/video.rs index f180d124d6..f37df0d75c 100644 --- a/rs/moq-c/src/video.rs +++ b/rs/moq-c/src/video.rs @@ -968,7 +968,7 @@ pub unsafe extern "C" fn moq_decode_video_frame(id: u32, dst: *mut moq_video_fra let frame = State::lock().video.frame(id)?; let pixels = frame.pixels()?; *dst = moq_video_frame { - // The decoded Timestamp is bounded by a QUIC VarInt, so its microseconds fit. + // The decoded Timestamp is bounded by a QUIC varint, so its microseconds fit. timestamp_us: frame.frame.timestamp.as_micros() as u64, width: pixels.width, height: pixels.height, diff --git a/rs/moq-ffi/src/video.rs b/rs/moq-ffi/src/video.rs index 347b59b502..a7b2b9ca84 100644 --- a/rs/moq-ffi/src/video.rs +++ b/rs/moq-ffi/src/video.rs @@ -621,7 +621,7 @@ pub struct MoqVideoDecodedFrame { impl MoqVideoDecodedFrame { /// Presentation timestamp, in microseconds. pub fn timestamp_us(&self) -> u64 { - // A decoded Timestamp is bounded by a QUIC VarInt, so its microseconds fit. + // A decoded Timestamp is bounded by a QUIC varint, so its microseconds fit. self.frame.timestamp.as_micros() as u64 } diff --git a/rs/moq-loc/src/lib.rs b/rs/moq-loc/src/lib.rs index 145eeca3be..1d19e7ab37 100644 --- a/rs/moq-loc/src/lib.rs +++ b/rs/moq-loc/src/lib.rs @@ -23,10 +23,10 @@ //! encode. Public properties are not handled here. They belong in the MoQ //! object header and are stripped by the transport layer. //! -//! Varint encoding is QUIC-style throughout via [`moq_net::VarInt`]. +//! Varint encoding is QUIC-style throughout via [`moq_net::varint`]. use bytes::{Buf, Bytes, BytesMut}; -use moq_net::{BoundsExceeded, DecodeError, EncodeError, VarInt}; +use moq_net::{BoundsExceeded, DecodeError, EncodeError, varint}; /// Property IDs recognized by this implementation. const PROP_TIMESCALE: u64 = 0x08; @@ -104,7 +104,7 @@ impl From for Error { /// Consumes the properties_length prefix, walks the bounded property block, /// and returns the remainder as `payload`. pub fn decode(mut buf: Bytes) -> Result { - let properties_length: u64 = VarInt::decode_quic(&mut buf)?.into(); + let properties_length = varint::decode_quic(&mut buf)?; let properties_length: usize = properties_length.try_into().map_err(|_| Error::MalformedProperties)?; if properties_length > buf.remaining() { @@ -119,7 +119,7 @@ pub fn decode(mut buf: Bytes) -> Result { let mut first = true; while props.has_remaining() { - let delta: u64 = VarInt::decode_quic(&mut props)?.into(); + let delta = varint::decode_quic(&mut props)?; let abs = if first { first = false; delta @@ -129,7 +129,7 @@ pub fn decode(mut buf: Bytes) -> Result { prev_type = abs; if abs % 2 == 0 { - let value: u64 = VarInt::decode_quic(&mut props)?.into(); + let value = varint::decode_quic(&mut props)?; match abs { PROP_TIMESTAMP | PROP_TIMESTAMP_DRAFT03 => timestamp = Some(value), PROP_TIMESCALE => { @@ -141,7 +141,7 @@ pub fn decode(mut buf: Bytes) -> Result { _ => {} } } else { - let len: u64 = VarInt::decode_quic(&mut props)?.into(); + let len = varint::decode_quic(&mut props)?; let len: usize = len.try_into().map_err(|_| Error::MalformedProperties)?; if len > props.remaining() { return Err(Error::MalformedProperties); @@ -167,11 +167,11 @@ pub fn decode(mut buf: Bytes) -> Result { /// catalog timescale to interpret `timestamp`. pub fn encode(timestamp: u64, payload: &[u8]) -> Result { let mut props = BytesMut::with_capacity(16); - VarInt::from(PROP_TIMESTAMP).encode_quic(&mut props)?; - VarInt::from(timestamp).encode_quic(&mut props)?; + varint::encode_quic(PROP_TIMESTAMP, &mut props)?; + varint::encode_quic(timestamp, &mut props)?; let mut out = BytesMut::with_capacity(props.len() + payload.len() + 8); - VarInt::from(props.len()).encode_quic(&mut out)?; + varint::encode_quic(props.len() as u64, &mut out)?; out.extend_from_slice(&props); out.extend_from_slice(payload); @@ -184,7 +184,7 @@ mod tests { /// Test helper: write a u64 as a QUIC varint into `buf`. fn write_varint(buf: &mut BytesMut, value: u64) { - VarInt::from(value).encode_quic(buf).unwrap(); + varint::encode_quic(value, buf).unwrap(); } #[test] diff --git a/rs/moq-mux/src/catalog/hang/container.rs b/rs/moq-mux/src/catalog/hang/container.rs index be57cfc8e2..f780a415e3 100644 --- a/rs/moq-mux/src/catalog/hang/container.rs +++ b/rs/moq-mux/src/catalog/hang/container.rs @@ -6,7 +6,7 @@ use crate::container::{Container as ContainerTrait, Frame, Kind, fmp4, legacy, l /// /// Built from a track's audio or video configuration, including its container. pub enum Container { - /// VarInt timestamp + raw codec bitstream. The original hang wire format. + /// varint timestamp + raw codec bitstream. The original hang wire format. Legacy(Kind), /// ISO-BMFF moof+mdat fragments. The wrapped [`fmp4::Wire`] holds /// the track's `trak` box so per-frame writes and reads have the diff --git a/rs/moq-mux/src/catalog/msf/consumer.rs b/rs/moq-mux/src/catalog/msf/consumer.rs index c8e2a256a0..1b488bb5d7 100644 --- a/rs/moq-mux/src/catalog/msf/consumer.rs +++ b/rs/moq-mux/src/catalog/msf/consumer.rs @@ -166,7 +166,7 @@ pub(crate) fn from_msf(msf: &moq_msf::Catalog) -> Result Result> { match &track.packaging { - // Neither is ISO-BMFF boxed, but they frame differently: a LOC property block against a VarInt + // Neither is ISO-BMFF boxed, but they frame differently: a LOC property block against a varint // timestamp prefix. Reading one as the other misparses the head of every frame. moq_msf::Packaging::Loc => Ok(Some(Container::Loc)), moq_msf::Packaging::Legacy => Ok(Some(Container::Legacy)), diff --git a/rs/moq-mux/src/container/legacy/mod.rs b/rs/moq-mux/src/container/legacy/mod.rs index 2307148f3b..c63ead824d 100644 --- a/rs/moq-mux/src/container/legacy/mod.rs +++ b/rs/moq-mux/src/container/legacy/mod.rs @@ -1,6 +1,6 @@ //! The original hang wire format. //! -//! Each moq frame holds one media frame: a VarInt-encoded timestamp +//! Each moq frame holds one media frame: a varint-encoded timestamp //! followed by the raw codec bitstream. Simple but ad-hoc; new //! broadcasts should use [`crate::container::loc`] instead. diff --git a/rs/moq-net/src/coding/decode.rs b/rs/moq-net/src/coding/decode.rs index c1fbfc0dab..c5844e202f 100644 --- a/rs/moq-net/src/coding/decode.rs +++ b/rs/moq-net/src/coding/decode.rs @@ -1,7 +1,7 @@ use std::string::FromUtf8Error; use thiserror::Error; -use super::{BoundsExceeded, Form, VarInt}; +use super::{BoundsExceeded, Form, varint}; /// Read the value from a [`Decoder`] using the given version. /// @@ -173,21 +173,21 @@ impl<'a> Decoder<'a> { /// Read a varint. #[inline] - pub fn varint(&mut self) -> Result { - let (value, rest) = VarInt::read(self.buf, self.form)?; + pub fn varint(&mut self) -> Result { + let (value, rest) = varint::read(self.buf, self.form)?; self.buf = rest; Ok(value) } /// Read an optional varint: 0 is `None`, and `n + 1` is `Some(n)`. pub fn varint_opt(&mut self) -> Result, DecodeError> { - Ok(self.varint()?.into_inner().checked_sub(1)) + Ok(self.varint()?.checked_sub(1)) } /// Read a varint length, then that many raw bytes. pub fn bytes(&mut self) -> Result<&'a [u8], DecodeError> { let start = self.buf; - let len = usize::try_from(self.varint()?)?; + let len = usize::try_from(self.varint()?).map_err(|_| DecodeError::BoundsExceeded)?; self.slice(len).inspect_err(|_| self.buf = start) } diff --git a/rs/moq-net/src/coding/encode.rs b/rs/moq-net/src/coding/encode.rs index 7076a8a08d..dd70fb4451 100644 --- a/rs/moq-net/src/coding/encode.rs +++ b/rs/moq-net/src/coding/encode.rs @@ -1,6 +1,6 @@ use bytes::Bytes; -use super::{BoundsExceeded, Form, VarInt}; +use super::{BoundsExceeded, Form, varint}; /// An error that occurs during encoding. #[derive(thiserror::Error, Debug, Clone)] @@ -93,8 +93,8 @@ impl<'a> Encoder<'a> { /// Write a varint, or fail with [`EncodeError::BoundsExceeded`] if the form cannot carry it. #[inline] - pub fn varint(&mut self, v: VarInt) -> Result<(), EncodeError> { - Ok(v.write(self.form, self.buf)?) + pub fn varint(&mut self, v: u64) -> Result<(), EncodeError> { + Ok(varint::write(v, self.form, self.buf)?) } /// Write an optional varint: `None` as 0, and `Some(n)` as `n + 1`. @@ -103,12 +103,12 @@ impl<'a> Encoder<'a> { Some(v) => v.checked_add(1).ok_or(EncodeError::TooLarge)?, None => 0, }; - self.varint(v.into()) + self.varint(v) } /// Write a varint length, then the raw bytes. pub fn bytes(&mut self, v: &[u8]) -> Result<(), EncodeError> { - self.varint(v.len().into())?; + self.varint(v.len() as u64)?; self.slice(v); Ok(()) } @@ -125,7 +125,7 @@ impl<'a> Encoder<'a> { self.buf.push(0); Prefix { body: self.buf.len(), - kind: PrefixKind::VarInt, + kind: PrefixKind::Varint, } } @@ -148,10 +148,10 @@ impl<'a> Encoder<'a> { let size = u16::try_from(end - body).map_err(|_| EncodeError::TooLarge)?; self.buf[body - 2..body].copy_from_slice(&size.to_be_bytes()); } - PrefixKind::VarInt => { + PrefixKind::Varint => { // Encode the size past the body, then move it into the reserved byte, // shifting the body up when the size needs more than that one byte. - VarInt::from(end - body).write(self.form, self.buf)?; + varint::write((end - body) as u64, self.form, self.buf)?; let len = self.buf.len() - end; if len == 1 { self.buf[body - 1] = self.buf[end]; @@ -181,7 +181,7 @@ pub struct Prefix { #[derive(Debug)] enum PrefixKind { - VarInt, + Varint, U16, } @@ -203,7 +203,7 @@ mod tests { let mut r = super::super::Decoder::new(&buf[1..], Form::Quic); assert_eq!(buf[0], 0xaa); - assert_eq!(r.varint().unwrap().into_inner(), size as u64); + assert_eq!(r.varint().unwrap(), size as u64); assert_eq!(r.slice(size).unwrap(), vec![0x55; size]); assert_eq!(r.rest(), [0xbb]); } diff --git a/rs/moq-net/src/coding/mod.rs b/rs/moq-net/src/coding/mod.rs index c44cc6f23a..263cc221d3 100644 --- a/rs/moq-net/src/coding/mod.rs +++ b/rs/moq-net/src/coding/mod.rs @@ -5,7 +5,7 @@ mod decode; mod encode; mod reader; mod stream; -mod varint; +pub mod varint; mod version; mod writer; @@ -14,7 +14,8 @@ pub use decode::*; pub use encode::*; pub use reader::*; pub use stream::*; -pub use varint::*; +pub use varint::BoundsExceeded; +pub(crate) use varint::Form; pub use version::*; pub use writer::*; @@ -36,5 +37,5 @@ pub(crate) fn decode_buf + Copy, T>( /// Decode one varint from the front of a test buffer, advancing it past what was read. #[cfg(test)] pub(crate) fn decode_varint + Copy>(buf: &mut B, version: V) -> Result { - decode_buf(buf, version, |r, _| Ok(r.varint()?.into_inner())) + decode_buf(buf, version, |r, _| r.varint()) } diff --git a/rs/moq-net/src/coding/reader.rs b/rs/moq-net/src/coding/reader.rs index dbc0968207..21ced73bd6 100644 --- a/rs/moq-net/src/coding/reader.rs +++ b/rs/moq-net/src/coding/reader.rs @@ -40,18 +40,26 @@ impl Reader { } } - /// Poll for the next message on the stream. - pub fn poll_decode + Debug>(&mut self, cx: &mut Context<'_>) -> Poll> + /// Poll `decode` against the buffered bytes, reading more until it stops coming up + /// short. `consume` drops what it read; otherwise the bytes stay for the next read. + fn poll_with( + &mut self, + cx: &mut Context<'_>, + consume: bool, + decode: impl Fn(&mut Decoder<'_>) -> Result, + ) -> Poll> where V: Into + Copy, { loop { let mut r = Decoder::new(&self.buffer, self.version.into()); - match T::decode(&mut r, self.version) { - Ok(msg) => { - let used = self.buffer.len() - r.remaining(); - self.buffer.advance(used); - return Poll::Ready(Ok(msg)); + match decode(&mut r) { + Ok(value) => { + if consume { + let used = self.buffer.len() - r.remaining(); + self.buffer.advance(used); + } + return Poll::Ready(Ok(value)); } // Stream closed while we still need more data. Err(DecodeError::Short) if !ready!(self.poll_read_more(cx))? => { @@ -63,6 +71,32 @@ impl Reader { } } + /// [`Self::poll_with`], or `None` if the stream closes cleanly before any byte. + fn poll_with_maybe( + &mut self, + cx: &mut Context<'_>, + consume: bool, + decode: impl Fn(&mut Decoder<'_>) -> Result, + ) -> Poll, Error>> + where + V: Into + Copy, + { + if !ready!(self.poll_has_more(cx))? { + return Poll::Ready(Ok(None)); + } + + self.poll_with(cx, consume, decode).map_ok(Some) + } + + /// Poll for the next message on the stream. + pub fn poll_decode + Debug>(&mut self, cx: &mut Context<'_>) -> Poll> + where + V: Into + Copy, + { + let version = self.version; + self.poll_with(cx, true, |r| T::decode(r, version)) + } + /// Decode the next message from the stream. pub async fn decode + Debug>(&mut self) -> Result where @@ -76,11 +110,8 @@ impl Reader { where V: Into + Copy, { - if !ready!(self.poll_has_more(cx))? { - return Poll::Ready(Ok(None)); - } - - self.poll_decode(cx).map_ok(Some) + let version = self.version; + self.poll_with_maybe(cx, true, |r| T::decode(r, version)) } /// Decode the next message unless the stream is closed. @@ -91,53 +122,72 @@ impl Reader { std::future::poll_fn(|cx| self.poll_decode_maybe(cx)).await } - /// Poll for the next message without consuming it. - pub fn poll_decode_peek + Debug>(&mut self, cx: &mut Context<'_>) -> Poll> + /// Poll for the next message without consuming it unless the stream closes cleanly first. + pub fn poll_decode_peek_maybe + Debug>( + &mut self, + cx: &mut Context<'_>, + ) -> Poll, Error>> where V: Into + Copy, { - loop { - let mut r = Decoder::new(&self.buffer, self.version.into()); - match T::decode(&mut r, self.version) { - Ok(msg) => return Poll::Ready(Ok(msg)), - Err(DecodeError::Short) if !ready!(self.poll_read_more(cx))? => { - return Poll::Ready(Err(DecodeError::Short.into())); - } - Err(DecodeError::Short) => {} - Err(e) => return Poll::Ready(Err(e.into())), - } - } + let version = self.version; + self.poll_with_maybe(cx, false, |r| T::decode(r, version)) } - /// Decode the next message from the stream without consuming it. - pub async fn decode_peek + Debug>(&mut self) -> Result + /// Peek the next message unless the stream is closed. + pub async fn decode_peek_maybe + Debug>(&mut self) -> Result, Error> where V: Into + Copy, { - std::future::poll_fn(|cx| self.poll_decode_peek(cx)).await + std::future::poll_fn(|cx| self.poll_decode_peek_maybe(cx)).await } - /// Poll for the next message without consuming it unless the stream closes cleanly first. - pub fn poll_decode_peek_maybe + Debug>( - &mut self, - cx: &mut Context<'_>, - ) -> Poll, Error>> + /// Poll for the next varint on the stream. + pub fn poll_varint(&mut self, cx: &mut Context<'_>) -> Poll> where V: Into + Copy, { - if !ready!(self.poll_has_more(cx))? { - return Poll::Ready(Ok(None)); - } + self.poll_with(cx, true, |r| r.varint()) + } - self.poll_decode_peek(cx).map_ok(Some) + /// Read the next varint from the stream. + pub async fn varint(&mut self) -> Result + where + V: Into + Copy, + { + std::future::poll_fn(|cx| self.poll_varint(cx)).await } - /// Peek the next message unless the stream is closed. - pub async fn decode_peek_maybe + Debug>(&mut self) -> Result, Error> + /// Poll for the next varint unless the stream is closed cleanly first. + pub fn poll_varint_maybe(&mut self, cx: &mut Context<'_>) -> Poll, Error>> where V: Into + Copy, { - std::future::poll_fn(|cx| self.poll_decode_peek_maybe(cx)).await + self.poll_with_maybe(cx, true, |r| r.varint()) + } + + /// Read the next varint unless the stream is closed. + pub async fn varint_maybe(&mut self) -> Result, Error> + where + V: Into + Copy, + { + std::future::poll_fn(|cx| self.poll_varint_maybe(cx)).await + } + + /// Poll for the next varint without consuming it. + pub fn poll_varint_peek(&mut self, cx: &mut Context<'_>) -> Poll> + where + V: Into + Copy, + { + self.poll_with(cx, false, |r| r.varint()) + } + + /// Read the next varint without consuming it. + pub async fn varint_peek(&mut self) -> Result + where + V: Into + Copy, + { + std::future::poll_fn(|cx| self.poll_varint_peek(cx)).await } /// Poll for the next chunk, draining the reader's internal buffer first. diff --git a/rs/moq-net/src/coding/varint.rs b/rs/moq-net/src/coding/varint.rs index 9f5e5a3a59..45ef4b7fd8 100644 --- a/rs/moq-net/src/coding/varint.rs +++ b/rs/moq-net/src/coding/varint.rs @@ -1,12 +1,12 @@ -// Based on quinn-proto -// https://github.com/quinn-rs/quinn/blob/main/quinn-proto/src/varint.rs -// Licensed via Apache 2.0 and MIT - -use std::fmt; +//! Variable-length integers: QUIC's two-bit length tag, and moq-transport's leading ones. +//! +//! A varint is a wire encoding of a plain `u64`, not a type. The codec works on the two +//! `u32` halves so it never needs 64-bit bitwise math, which a JavaScript `number` +//! cannot do. use thiserror::Error; -use super::{Decode, DecodeError, Decoder, Encode, EncodeError, Encoder}; +use super::{DecodeError, EncodeError}; use crate::{Version, ietf, lite}; /// The number does not fit the target: a varint wire form or a narrower integer. @@ -14,163 +14,15 @@ use crate::{Version, ietf, lite}; #[error("value out of range")] pub struct BoundsExceeded; -/// An integer destined for the wire as a variable-length integer. -/// -/// It holds the full `u64` range, which the leading-ones form of moq-transport draft-17+ -/// can carry. The QUIC form (moq-lite, drafts 14-16) tops out at `2^62 - 1`, so encoding -/// a larger value there fails with [`BoundsExceeded`] rather than truncating. -#[derive(Debug, Default, Copy, Clone, Eq, PartialEq, Ord, PartialOrd, Hash)] -pub struct VarInt(u64); - -impl VarInt { - /// The largest possible value. - pub const MAX: Self = Self(u64::MAX); - - /// The largest value the QUIC form can carry: `2^62 - 1`. - pub const MAX_QUIC: Self = Self((1 << 62) - 1); - - /// The smallest possible value. - pub const ZERO: Self = Self(0); - - /// Construct from a `u32`, usable in `const` contexts. - pub const fn from_u32(x: u32) -> Self { - Self(x as u64) - } - - /// Construct from a `u64`, usable in `const` contexts. - pub const fn from_u64(x: u64) -> Self { - Self(x) - } - - /// Construct from a `u128`, or `None` if it exceeds [`Self::MAX`]. - pub const fn from_u128(x: u128) -> Option { - if x <= u64::MAX as u128 { - Some(Self(x as u64)) - } else { - None - } - } - - /// Extract the integer value - pub const fn into_inner(self) -> u64 { - self.0 - } - - /// Map a signed `i64` onto the unsigned range with zigzag: `(n << 1) ^ (n >> 63)`. - /// - /// Small negative numbers map to small unsigneds (-1 -> 1, 1 -> 2, -2 -> 3, ...), and - /// the whole `i64` range fits. - pub const fn from_zigzag(signed: i64) -> Self { - Self(((signed << 1) ^ (signed >> 63)) as u64) - } - - /// Decode this varint as a signed `i64` via the inverse zigzag transform. - pub const fn to_zigzag(self) -> i64 { - let v = self.0; - ((v >> 1) as i64) ^ -((v & 1) as i64) - } -} - -impl From for u64 { - fn from(x: VarInt) -> Self { - x.0 - } -} - -impl From for u128 { - fn from(x: VarInt) -> Self { - x.0 as u128 - } -} - -impl From for VarInt { - fn from(x: u8) -> Self { - Self(x.into()) - } -} - -impl From for VarInt { - fn from(x: u16) -> Self { - Self(x.into()) - } -} - -impl From for VarInt { - fn from(x: u32) -> Self { - Self(x.into()) - } -} - -impl From for VarInt { - fn from(x: u64) -> Self { - Self(x) - } -} - -impl From for VarInt { - fn from(x: usize) -> Self { - // usize is at most 64 bits on every target Rust supports. - Self(x as u64) - } -} - -impl TryFrom for VarInt { - type Error = BoundsExceeded; - - /// Succeeds iff `x` < 2^64 - fn try_from(x: u128) -> Result { - Self::from_u128(x).ok_or(BoundsExceeded) - } -} - -impl TryFrom for usize { - type Error = BoundsExceeded; - - /// Succeeds iff `x` fits the target's pointer width. - fn try_from(x: VarInt) -> Result { - usize::try_from(x.0).map_err(|_| BoundsExceeded) - } -} - -impl TryFrom for u32 { - type Error = BoundsExceeded; - - /// Succeeds iff `x` < 2^32 - fn try_from(x: VarInt) -> Result { - u32::try_from(x.0).map_err(|_| BoundsExceeded) - } -} - -impl TryFrom for u16 { - type Error = BoundsExceeded; - - /// Succeeds iff `x` < 2^16 - fn try_from(x: VarInt) -> Result { - u16::try_from(x.0).map_err(|_| BoundsExceeded) - } -} - -impl TryFrom for u8 { - type Error = BoundsExceeded; - - /// Succeeds iff `x` < 2^8 - fn try_from(x: VarInt) -> Result { - u8::try_from(x.0).map_err(|_| BoundsExceeded) - } -} - -impl fmt::Display for VarInt { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - self.0.fmt(f) - } -} +/// The largest value the QUIC form can carry: `2^62 - 1`. +pub const MAX_QUIC: u64 = (1 << 62) - 1; -/// How a protocol version lays out a [`VarInt`] on the wire. +/// How a protocol version lays out a varint on the wire. #[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub enum Form { - /// QUIC's two-bit length tag, up to [`VarInt::MAX_QUIC`]. +pub(crate) enum Form { + /// QUIC's two-bit length tag, up to [`MAX_QUIC`]. Quic, - /// Leading one bits count the length, up to [`VarInt::MAX`]. + /// Leading one bits count the length, up to `u64::MAX`. LeadingOnes { /// Whether the 7-byte form (`1111110x`) is accepted on decode, which draft-17 forbids. seven: bool, @@ -210,233 +62,227 @@ impl From for Form { } } -impl VarInt { - /// The high and low 32 bits. - /// - /// The codec works on the halves so it never needs 64-bit bitwise math, which a - /// JavaScript `number` cannot do. - const fn to_halves(self) -> (u32, u32) { - ((self.0 >> 32) as u32, self.0 as u32) - } - - /// The inverse of [`Self::to_halves`]. - const fn from_halves(hi: u32, lo: u32) -> Self { - Self(((hi as u64) << 32) | lo as u64) - } - - /// The bytes this takes on the wire in the given form, or [`BoundsExceeded`] if it - /// does not fit. - pub(crate) fn size(self, form: Form) -> Result { - let (hi, lo) = self.to_halves(); - Ok(match form { - Form::Quic if hi == 0 && lo < 1 << 6 => 1, - Form::Quic if hi == 0 && lo < 1 << 14 => 2, - Form::Quic if hi == 0 && lo < 1 << 30 => 4, - Form::Quic if hi < 1 << 30 => 8, - Form::Quic => return Err(BoundsExceeded), - Form::LeadingOnes { .. } if hi == 0 && lo < 1 << 7 => 1, - Form::LeadingOnes { .. } if hi == 0 && lo < 1 << 14 => 2, - Form::LeadingOnes { .. } if hi == 0 && lo < 1 << 21 => 3, - Form::LeadingOnes { .. } if hi == 0 && lo < 1 << 28 => 4, - Form::LeadingOnes { .. } if hi < 1 << 3 => 5, - Form::LeadingOnes { .. } if hi < 1 << 10 => 6, - // The 7-byte form is skipped: one byte longer, but legal on every draft. - Form::LeadingOnes { .. } if hi < 1 << 24 => 8, - Form::LeadingOnes { .. } => 9, - }) - } +/// The high and low 32 bits: the only place the codec splits a `u64`. +const fn to_halves(value: u64) -> (u32, u32) { + ((value >> 32) as u32, value as u32) +} - /// Append the minimal encoding in the given form. - /// - /// Fails past [`Self::MAX_QUIC`] in the QUIC form, writing nothing. - #[inline] - pub(super) fn write(self, form: Form, out: &mut Vec) -> Result<(), BoundsExceeded> { - match form { - Form::Quic => self.write_quic(out), - Form::LeadingOnes { .. } => { - self.write_leading_ones(out); - Ok(()) - } - } - } +/// The inverse of [`to_halves`]. +const fn from_halves(hi: u32, lo: u32) -> u64 { + ((hi as u64) << 32) | lo as u64 +} - // Each arm below is a fixed-size write or read, which is what keeps the codec as fast - // as a hand-rolled `put_u16`/`get_u32`. - - #[inline] - fn write_quic(self, out: &mut Vec) -> Result<(), BoundsExceeded> { - let (hi, lo) = self.to_halves(); - if hi == 0 && lo < 1 << 6 { - out.push(lo as u8); - } else if hi == 0 && lo < 1 << 14 { - out.extend_from_slice(&(0x4000 | lo as u16).to_be_bytes()); - } else if hi == 0 && lo < 1 << 30 { - out.extend_from_slice(&(0x8000_0000 | lo).to_be_bytes()); - } else if hi < 1 << 30 { - let [a, b, c, d] = (0xc000_0000 | hi).to_be_bytes(); - let [e, f, g, h] = lo.to_be_bytes(); - out.extend_from_slice(&[a, b, c, d, e, f, g, h]); - } else { - return Err(BoundsExceeded); - } - Ok(()) - } +/// The bytes `value` takes on the wire in `form`, or [`BoundsExceeded`] if it does not fit. +pub(crate) fn size(value: u64, form: Form) -> Result { + let (hi, lo) = to_halves(value); + Ok(match form { + Form::Quic if hi == 0 && lo < 1 << 6 => 1, + Form::Quic if hi == 0 && lo < 1 << 14 => 2, + Form::Quic if hi == 0 && lo < 1 << 30 => 4, + Form::Quic if hi < 1 << 30 => 8, + Form::Quic => return Err(BoundsExceeded), + Form::LeadingOnes { .. } if hi == 0 && lo < 1 << 7 => 1, + Form::LeadingOnes { .. } if hi == 0 && lo < 1 << 14 => 2, + Form::LeadingOnes { .. } if hi == 0 && lo < 1 << 21 => 3, + Form::LeadingOnes { .. } if hi == 0 && lo < 1 << 28 => 4, + Form::LeadingOnes { .. } if hi < 1 << 3 => 5, + Form::LeadingOnes { .. } if hi < 1 << 10 => 6, + // The 7-byte form is skipped: one byte longer, but legal on every draft. + Form::LeadingOnes { .. } if hi < 1 << 24 => 8, + Form::LeadingOnes { .. } => 9, + }) +} - #[inline] - fn write_leading_ones(self, out: &mut Vec) { - let (hi, lo) = self.to_halves(); - let [a, b, c, d] = lo.to_be_bytes(); - if hi == 0 && lo < 1 << 7 { - out.push(d); - } else if hi == 0 && lo < 1 << 14 { - out.extend_from_slice(&[0x80 | c, d]); - } else if hi == 0 && lo < 1 << 21 { - out.extend_from_slice(&[0xc0 | b, c, d]); - } else if hi == 0 && lo < 1 << 28 { - out.extend_from_slice(&[0xe0 | a, b, c, d]); - } else if hi < 1 << 3 { - out.extend_from_slice(&[0xf0 | hi as u8, a, b, c, d]); - } else if hi < 1 << 10 { - out.extend_from_slice(&[0xf8 | (hi >> 8) as u8, hi as u8, a, b, c, d]); - } else if hi < 1 << 24 { - // The 7-byte form is skipped: one byte longer, but legal on every draft. - let [_, f, g, h] = hi.to_be_bytes(); - out.extend_from_slice(&[0xfe, f, g, h, a, b, c, d]); - } else { - let [e, f, g, h] = hi.to_be_bytes(); - out.extend_from_slice(&[0xff, e, f, g, h, a, b, c, d]); +/// Append the minimal encoding of `value` in `form`. +/// +/// Fails past [`MAX_QUIC`] in the QUIC form, writing nothing. +#[inline] +pub(super) fn write(value: u64, form: Form, out: &mut Vec) -> Result<(), BoundsExceeded> { + match form { + Form::Quic => write_quic(value, out), + Form::LeadingOnes { .. } => { + write_leading_ones(value, out); + Ok(()) } } +} - /// Decode from the front of `buf`, returning the value and the rest of `buf`. - #[inline] - pub(super) fn read(buf: &[u8], form: Form) -> Result<(Self, &[u8]), DecodeError> { - match form { - Form::Quic => Self::read_quic(buf), - Form::LeadingOnes { seven } => Self::read_leading_ones(buf, seven), - } - } +// Each arm below is a fixed-size write or read, which is what keeps the codec as fast as +// a hand-rolled `put_u16`/`get_u32`. + +#[inline] +fn write_quic(value: u64, out: &mut Vec) -> Result<(), BoundsExceeded> { + let (hi, lo) = to_halves(value); + if hi == 0 && lo < 1 << 6 { + out.push(lo as u8); + } else if hi == 0 && lo < 1 << 14 { + out.extend_from_slice(&(0x4000 | lo as u16).to_be_bytes()); + } else if hi == 0 && lo < 1 << 30 { + out.extend_from_slice(&(0x8000_0000 | lo).to_be_bytes()); + } else if hi < 1 << 30 { + let [a, b, c, d] = (0xc000_0000 | hi).to_be_bytes(); + let [e, f, g, h] = lo.to_be_bytes(); + out.extend_from_slice(&[a, b, c, d, e, f, g, h]); + } else { + return Err(BoundsExceeded); + } + Ok(()) +} - #[inline] - fn read_quic(buf: &[u8]) -> Result<(Self, &[u8]), DecodeError> { - let Some((&first, rest)) = buf.split_first() else { - return Err(DecodeError::Short); - }; - - let be = u32::from_be_bytes; - Ok(match first >> 6 { - 0 => (Self::from_u32(first as u32), rest), - 1 => { - let ([a, b], rest) = buf.split_first_chunk().ok_or(DecodeError::Short)?; - (Self::from_u32(be([0, 0, a & 0x3f, *b])), rest) - } - 2 => { - let ([a, b, c, d], rest) = buf.split_first_chunk().ok_or(DecodeError::Short)?; - (Self::from_u32(be([a & 0x3f, *b, *c, *d])), rest) - } - _ => { - let ([a, b, c, d, lo @ ..], rest) = buf.split_first_chunk::<8>().ok_or(DecodeError::Short)?; - (Self::from_halves(be([a & 0x3f, *b, *c, *d]), be(*lo)), rest) - } - }) +#[inline] +fn write_leading_ones(value: u64, out: &mut Vec) { + let (hi, lo) = to_halves(value); + let [a, b, c, d] = lo.to_be_bytes(); + if hi == 0 && lo < 1 << 7 { + out.push(d); + } else if hi == 0 && lo < 1 << 14 { + out.extend_from_slice(&[0x80 | c, d]); + } else if hi == 0 && lo < 1 << 21 { + out.extend_from_slice(&[0xc0 | b, c, d]); + } else if hi == 0 && lo < 1 << 28 { + out.extend_from_slice(&[0xe0 | a, b, c, d]); + } else if hi < 1 << 3 { + out.extend_from_slice(&[0xf0 | hi as u8, a, b, c, d]); + } else if hi < 1 << 10 { + out.extend_from_slice(&[0xf8 | (hi >> 8) as u8, hi as u8, a, b, c, d]); + } else if hi < 1 << 24 { + // The 7-byte form is skipped: one byte longer, but legal on every draft. + let [_, f, g, h] = hi.to_be_bytes(); + out.extend_from_slice(&[0xfe, f, g, h, a, b, c, d]); + } else { + let [e, f, g, h] = hi.to_be_bytes(); + out.extend_from_slice(&[0xff, e, f, g, h, a, b, c, d]); } +} - #[inline] - fn read_leading_ones(buf: &[u8], seven: bool) -> Result<(Self, &[u8]), DecodeError> { - let Some((&first, rest)) = buf.split_first() else { - return Err(DecodeError::Short); - }; - - let be = u32::from_be_bytes; - Ok(match first.leading_ones() { - 0 => (Self::from_u32(first as u32), rest), - 1 => { - let ([a, b], rest) = buf.split_first_chunk().ok_or(DecodeError::Short)?; - (Self::from_u32(be([0, 0, a & 0x3f, *b])), rest) - } - 2 => { - let ([a, b, c], rest) = buf.split_first_chunk().ok_or(DecodeError::Short)?; - (Self::from_u32(be([0, a & 0x1f, *b, *c])), rest) - } - 3 => { - let ([a, b, c, d], rest) = buf.split_first_chunk().ok_or(DecodeError::Short)?; - (Self::from_u32(be([a & 0x0f, *b, *c, *d])), rest) - } - 4 => { - let ([a, lo @ ..], rest) = buf.split_first_chunk::<5>().ok_or(DecodeError::Short)?; - (Self::from_halves((a & 0x07) as u32, be(*lo)), rest) - } - 5 => { - let ([a, b, lo @ ..], rest) = buf.split_first_chunk::<6>().ok_or(DecodeError::Short)?; - (Self::from_halves(be([0, 0, a & 0x03, *b]), be(*lo)), rest) - } - // 1111110x: the 7-byte form, which draft-17 forbids. - 6 if !seven => return Err(DecodeError::InvalidValue), - 6 => { - let ([a, b, c, lo @ ..], rest) = buf.split_first_chunk::<7>().ok_or(DecodeError::Short)?; - (Self::from_halves(be([0, a & 0x01, *b, *c]), be(*lo)), rest) - } - 7 => { - let ([_, b, c, d, lo @ ..], rest) = buf.split_first_chunk::<8>().ok_or(DecodeError::Short)?; - (Self::from_halves(be([0, *b, *c, *d]), be(*lo)), rest) - } - _ => { - let ([_, hi @ .., e, f, g, h], rest) = buf.split_first_chunk::<9>().ok_or(DecodeError::Short)?; - (Self::from_halves(be(*hi), be([*e, *f, *g, *h])), rest) - } - }) +/// Decode a varint in `form` from the front of `buf`, returning it and the rest of `buf`. +#[inline] +pub(super) fn read(buf: &[u8], form: Form) -> Result<(u64, &[u8]), DecodeError> { + match form { + Form::Quic => read_quic(buf), + Form::LeadingOnes { seven } => read_leading_ones(buf, seven), } +} - /// Decode a QUIC-style varint (2-bit length tag in top bits). - pub fn decode_quic(r: &mut R) -> Result { - let Some(&first) = r.chunk().first() else { - return Err(DecodeError::Short); - }; +#[inline] +fn read_quic(buf: &[u8]) -> Result<(u64, &[u8]), DecodeError> { + let Some((&first, rest)) = buf.split_first() else { + return Err(DecodeError::Short); + }; + + let be = u32::from_be_bytes; + Ok(match first >> 6 { + 0 => (first as u64, rest), + 1 => { + let ([a, b], rest) = buf.split_first_chunk().ok_or(DecodeError::Short)?; + (be([0, 0, a & 0x3f, *b]) as u64, rest) + } + 2 => { + let ([a, b, c, d], rest) = buf.split_first_chunk().ok_or(DecodeError::Short)?; + (be([a & 0x3f, *b, *c, *d]) as u64, rest) + } + _ => { + let ([a, b, c, d, lo @ ..], rest) = buf.split_first_chunk::<8>().ok_or(DecodeError::Short)?; + (from_halves(be([a & 0x3f, *b, *c, *d]), be(*lo)), rest) + } + }) +} - // Copy out so a varint split across chunks still decodes. - let len = 1usize << (first >> 6); - if r.remaining() < len { - return Err(DecodeError::Short); +#[inline] +fn read_leading_ones(buf: &[u8], seven: bool) -> Result<(u64, &[u8]), DecodeError> { + let Some((&first, rest)) = buf.split_first() else { + return Err(DecodeError::Short); + }; + + let be = u32::from_be_bytes; + Ok(match first.leading_ones() { + 0 => (first as u64, rest), + 1 => { + let ([a, b], rest) = buf.split_first_chunk().ok_or(DecodeError::Short)?; + (be([0, 0, a & 0x3f, *b]) as u64, rest) + } + 2 => { + let ([a, b, c], rest) = buf.split_first_chunk().ok_or(DecodeError::Short)?; + (be([0, a & 0x1f, *b, *c]) as u64, rest) + } + 3 => { + let ([a, b, c, d], rest) = buf.split_first_chunk().ok_or(DecodeError::Short)?; + (be([a & 0x0f, *b, *c, *d]) as u64, rest) + } + 4 => { + let ([a, lo @ ..], rest) = buf.split_first_chunk::<5>().ok_or(DecodeError::Short)?; + (from_halves((a & 0x07) as u32, be(*lo)), rest) + } + 5 => { + let ([a, b, lo @ ..], rest) = buf.split_first_chunk::<6>().ok_or(DecodeError::Short)?; + (from_halves(be([0, 0, a & 0x03, *b]), be(*lo)), rest) } - let mut buf = [0u8; 8]; - r.copy_to_slice(&mut buf[..len]); + // 1111110x: the 7-byte form, which draft-17 forbids. + 6 if !seven => return Err(DecodeError::InvalidValue), + 6 => { + let ([a, b, c, lo @ ..], rest) = buf.split_first_chunk::<7>().ok_or(DecodeError::Short)?; + (from_halves(be([0, a & 0x01, *b, *c]), be(*lo)), rest) + } + 7 => { + let ([_, b, c, d, lo @ ..], rest) = buf.split_first_chunk::<8>().ok_or(DecodeError::Short)?; + (from_halves(be([0, *b, *c, *d]), be(*lo)), rest) + } + _ => { + let ([_, hi @ .., e, f, g, h], rest) = buf.split_first_chunk::<9>().ok_or(DecodeError::Short)?; + (from_halves(be(*hi), be([*e, *f, *g, *h])), rest) + } + }) +} - Ok(Self::read(&buf[..len], Form::Quic)?.0) +/// Decode a QUIC varint (two-bit length tag) from the front of `r`. +pub fn decode_quic(r: &mut R) -> Result { + let Some(&first) = r.chunk().first() else { + return Err(DecodeError::Short); + }; + + // Copy out so a varint split across chunks still decodes. + let len = 1usize << (first >> 6); + if r.remaining() < len { + return Err(DecodeError::Short); } + let mut buf = [0u8; 8]; + r.copy_to_slice(&mut buf[..len]); - /// Encode a QUIC-style varint (2-bit length tag in top bits). - /// - /// Fails with [`EncodeError::BoundsExceeded`] past [`Self::MAX_QUIC`]. - pub fn encode_quic(&self, w: &mut W) -> Result<(), EncodeError> { - let len = self.size(Form::Quic)?; - if w.remaining_mut() < len { - return Err(EncodeError::Short); - } + Ok(read_quic(&buf[..len])?.0) +} - let (hi, lo) = self.to_halves(); - match len { - 1 => w.put_u8(lo as u8), - 2 => w.put_u16(0x4000 | lo as u16), - 4 => w.put_u32(0x8000_0000 | lo), - _ => { - w.put_u32(0xc000_0000 | hi); - w.put_u32(lo); - } +/// Encode `value` as a QUIC varint (two-bit length tag). +/// +/// Fails with [`EncodeError::BoundsExceeded`] past [`MAX_QUIC`], writing nothing. +pub fn encode_quic(value: u64, w: &mut W) -> Result<(), EncodeError> { + let len = size(value, Form::Quic)?; + if w.remaining_mut() < len { + return Err(EncodeError::Short); + } + + let (hi, lo) = to_halves(value); + match len { + 1 => w.put_u8(lo as u8), + 2 => w.put_u16(0x4000 | lo as u16), + 4 => w.put_u32(0x8000_0000 | lo), + _ => { + w.put_u32(0xc000_0000 | hi); + w.put_u32(lo); } - Ok(()) } + Ok(()) } -impl Encode for VarInt { - fn encode(&self, w: &mut Encoder<'_>, _: V) -> Result<(), EncodeError> { - w.varint(*self) - } +/// Map a signed value onto the unsigned range: `(n << 1) ^ (n >> 63)`. +/// +/// Small magnitudes stay small (-1 -> 1, 1 -> 2, -2 -> 3, ...), and all of `i64` fits. +pub(crate) const fn zigzag(signed: i64) -> u64 { + ((signed << 1) ^ (signed >> 63)) as u64 } -impl Decode for VarInt { - fn decode(r: &mut Decoder<'_>, _: V) -> Result { - r.varint() - } +/// The inverse of [`zigzag`]. +pub(crate) const fn unzigzag(value: u64) -> i64 { + ((value >> 1) as i64) ^ -((value & 1) as i64) } #[cfg(test)] @@ -446,10 +292,10 @@ mod tests { const DRAFT17: Form = Form::LeadingOnes { seven: false }; const DRAFT18: Form = Form::LeadingOnes { seven: true }; - fn encode(value: VarInt, form: Form) -> Result, BoundsExceeded> { + fn encode(value: u64, form: Form) -> Result, BoundsExceeded> { let mut buf = Vec::new(); - value.write(form, &mut buf)?; - assert_eq!(buf.len(), value.size(form)?, "size disagrees with the encoding"); + write(value, form, &mut buf)?; + assert_eq!(buf.len(), size(value, form)?, "size disagrees with the encoding"); Ok(buf) } @@ -475,17 +321,13 @@ mod tests { ]; for (bytes, expected) in cases { - let (decoded, rest) = VarInt::read(bytes, DRAFT17).expect("decode should succeed"); - assert_eq!( - decoded.into_inner(), - *expected, - "decode mismatch for bytes {bytes:02x?}" - ); + let (decoded, rest) = read(bytes, DRAFT17).expect("decode should succeed"); + assert_eq!(decoded, *expected, "decode mismatch for bytes {bytes:02x?}"); assert!(rest.is_empty(), "all bytes should be consumed for {bytes:02x?}"); // Skip the non-minimal encoding (0x8025 for 37); we only emit the minimal one. if bytes.len() == 1 || *expected != 37 { - let encoded = encode(VarInt::from(*expected), DRAFT17).unwrap(); + let encoded = encode(*expected, DRAFT17).unwrap(); assert_eq!(&encoded, bytes, "encode mismatch for value {expected}"); } } @@ -495,7 +337,7 @@ mod tests { #[test] fn leading_ones_invalid_0xfc() { assert!( - matches!(VarInt::read(&[0xFC], DRAFT17), Err(DecodeError::InvalidValue)), + matches!(read(&[0xFC], DRAFT17), Err(DecodeError::InvalidValue)), "0xFC should be rejected as invalid on draft-17" ); } @@ -512,15 +354,15 @@ mod tests { ]; for (value, expected_len) in cases { - let encoded = encode(VarInt::from(value), DRAFT17).unwrap(); + let encoded = encode(value, DRAFT17).unwrap(); assert_eq!( encoded.len(), expected_len, "unexpected encoded length for value {value}" ); - let (decoded, _) = VarInt::read(&encoded, DRAFT17).expect("leading-ones decode should succeed"); - assert_eq!(decoded.into_inner(), value, "round-trip mismatch for value {value}"); + let (decoded, _) = read(&encoded, DRAFT17).expect("leading-ones decode should succeed"); + assert_eq!(decoded, value, "round-trip mismatch for value {value}"); } } @@ -530,14 +372,13 @@ mod tests { fn every_length_round_trips() { for bits in 0..64 { for value in [1u64 << bits, (1u64 << bits) - 1, (1u64 << bits) + 1] { - let value = VarInt::from(value); for form in [Form::Quic, DRAFT17, DRAFT18] { let Ok(encoded) = encode(value, form) else { assert_eq!(form, Form::Quic); - assert!(value > VarInt::MAX_QUIC); + assert!(value > MAX_QUIC); continue; }; - let (decoded, rest) = VarInt::read(&encoded, form).unwrap(); + let (decoded, rest) = read(&encoded, form).unwrap(); assert_eq!((decoded, rest), (value, &[][..]), "{form:?} {value}"); } } @@ -548,48 +389,41 @@ mod tests { /// truncating, while the leading-ones form carries the whole u64. #[test] fn quic_refuses_past_62_bits() { - let max = VarInt::MAX_QUIC; - assert_eq!(max.into_inner(), (1 << 62) - 1); - assert_eq!(encode(max, Form::Quic).unwrap(), [0xff; 8]); + assert_eq!(encode(MAX_QUIC, Form::Quic).unwrap(), [0xff; 8]); for value in [1u64 << 62, u64::MAX] { - assert_eq!(encode(VarInt::from(value), Form::Quic), Err(BoundsExceeded)); + assert_eq!(encode(value, Form::Quic), Err(BoundsExceeded)); assert!(matches!( - VarInt::from(value).encode_quic(&mut Vec::new()), + encode_quic(value, &mut Vec::new()), Err(EncodeError::BoundsExceeded) )); } - for value in [(1u64 << 62) - 1, 1u64 << 62, u64::MAX] { - let encoded = encode(VarInt::from(value), DRAFT18).unwrap(); + for value in [MAX_QUIC, 1u64 << 62, u64::MAX] { + let encoded = encode(value, DRAFT18).unwrap(); assert_eq!(encoded.len(), 9); - assert_eq!(VarInt::read(&encoded, DRAFT18).unwrap().0.into_inner(), value); + assert_eq!(read(&encoded, DRAFT18).unwrap().0, value); } } #[test] fn draft17_rejects_7_byte_varint() { // 1111110x prefix: invalid on draft-17. - let err = VarInt::read(&[0xFC, 0, 0, 0, 0, 0, 0], DRAFT17).unwrap_err(); + let err = read(&[0xFC, 0, 0, 0, 0, 0, 0], DRAFT17).unwrap_err(); assert!(matches!(err, DecodeError::InvalidValue)); } #[test] fn zigzag_roundtrip_small() { for n in [-3i64, -2, -1, 0, 1, 2, 3, 100, -100] { - let v = VarInt::from_zigzag(n); - assert_eq!(v.to_zigzag(), n, "roundtrip failed for {}", n); + assert_eq!(unzigzag(zigzag(n)), n, "roundtrip failed for {}", n); } } #[test] fn zigzag_small_values_compact() { // First few values should fit in 1 byte (varint range 0..=63 = top-2-bits tag 00). - assert_eq!(VarInt::from_zigzag(0).into_inner(), 0); - assert_eq!(VarInt::from_zigzag(-1).into_inner(), 1); - assert_eq!(VarInt::from_zigzag(1).into_inner(), 2); - assert_eq!(VarInt::from_zigzag(-2).into_inner(), 3); - assert_eq!(VarInt::from_zigzag(2).into_inner(), 4); + assert_eq!([0, -1, 1, -2, 2].map(zigzag), [0, 1, 2, 3, 4]); } /// Zigzag covers the whole i64 range; only the QUIC form bounds what goes on the wire. @@ -598,23 +432,21 @@ mod tests { let mid = (1i64 << 30) + 17; for n in [i64::MAX, i64::MIN, (1i64 << 61) - 1, -(1i64 << 61), mid, -mid] { - let v = VarInt::from_zigzag(n); - assert_eq!(v.to_zigzag(), n); + assert_eq!(unzigzag(zigzag(n)), n); } - assert_eq!(VarInt::from_zigzag(i64::MIN), VarInt::MAX); - assert!(VarInt::from_zigzag(1i64 << 61) > VarInt::MAX_QUIC); - assert_eq!(VarInt::from_zigzag(-(1i64 << 61)), VarInt::MAX_QUIC); + assert_eq!(zigzag(i64::MIN), u64::MAX); + assert!(zigzag(1i64 << 61) > MAX_QUIC); + assert_eq!(zigzag(-(1i64 << 61)), MAX_QUIC); } #[test] fn zigzag_quic_varint_roundtrip() { // Encode a zigzag value through the QUIC varint wire format. for n in [-5000i64, 0, 100, -1, 1_000_000, -1_000_000] { - let v = VarInt::from_zigzag(n); - let bytes = v.encode_bytes(lite::Version::Lite01).unwrap(); - let (decoded, _) = VarInt::decode_slice(&bytes, lite::Version::Lite01).unwrap(); - assert_eq!(decoded.to_zigzag(), n); + let bytes = encode(zigzag(n), Form::Quic).unwrap(); + let (decoded, _) = read(&bytes, Form::Quic).unwrap(); + assert_eq!(unzigzag(decoded), n); } } @@ -630,8 +462,8 @@ mod tests { for shift in (0..48).step_by(8).rev() { bytes.push(((value >> shift) & 0xFF) as u8); } - let (decoded, _) = VarInt::read(&bytes, DRAFT18).unwrap(); - assert_eq!(decoded.into_inner(), value); + let (decoded, _) = read(&bytes, DRAFT18).unwrap(); + assert_eq!(decoded, value); } /// The Buf-based helpers other crates use read and write the same bytes, even when @@ -640,13 +472,13 @@ mod tests { fn quic_helpers_match_the_codec() { use bytes::Buf; - let value = VarInt::from(0x1234_5678u64); + let value = 0x1234_5678u64; let mut out = Vec::new(); - value.encode_quic(&mut out).unwrap(); + encode_quic(value, &mut out).unwrap(); assert_eq!(out, encode(value, Form::Quic).unwrap()); let mut split = (&out[..2]).chain(&out[2..]); - assert_eq!(VarInt::decode_quic(&mut split).unwrap(), value); + assert_eq!(decode_quic(&mut split).unwrap(), value); assert!(!split.has_remaining()); } } diff --git a/rs/moq-net/src/coding/version.rs b/rs/moq-net/src/coding/version.rs index 11000271ab..705912ef01 100644 --- a/rs/moq-net/src/coding/version.rs +++ b/rs/moq-net/src/coding/version.rs @@ -21,13 +21,13 @@ impl From for u64 { impl Decode for Version { /// Decode the version number. fn decode(r: &mut Decoder<'_>, _: V) -> Result { - Ok(Self(r.varint()?.into_inner())) + Ok(Self(r.varint()?)) } } impl Encode for Version { fn encode(&self, w: &mut Encoder<'_>, _: V) -> Result<(), EncodeError> { - w.varint(self.0.into())?; + w.varint(self.0)?; Ok(()) } } @@ -45,7 +45,7 @@ pub struct Versions(Vec); impl Decode for Versions { /// Decode the version list. fn decode(r: &mut Decoder<'_>, version: V) -> Result { - let count = r.varint()?.into_inner(); + let count = r.varint()?; let mut vs = Vec::new(); for _ in 0..count { @@ -60,7 +60,7 @@ impl Decode for Versions { impl Encode for Versions { /// Encode the version list. fn encode(&self, w: &mut Encoder<'_>, version: V) -> Result<(), EncodeError> { - w.varint(self.0.len().into())?; + w.varint(self.0.len() as u64)?; for v in &self.0 { v.encode(w, version)?; diff --git a/rs/moq-net/src/coding/writer.rs b/rs/moq-net/src/coding/writer.rs index 328cd21777..e14a9dcf45 100644 --- a/rs/moq-net/src/coding/writer.rs +++ b/rs/moq-net/src/coding/writer.rs @@ -34,12 +34,29 @@ impl Writer { /// Encode the given message into the write buffer, to be sent by /// [`Self::poll_flush`]. An encode error leaves the buffer untouched. pub fn buffer + Debug>(&mut self, msg: &T) -> Result<(), Error> + where + V: Into + Copy, + { + let version = self.version; + self.buffer_with(|w| msg.encode(w, version)) + } + + /// Encode a varint into the write buffer, to be sent by [`Self::poll_flush`]. + pub fn buffer_varint(&mut self, value: u64) -> Result<(), Error> + where + V: Into + Copy, + { + self.buffer_with(|w| w.varint(value)) + } + + /// Run `encode` into the write buffer, dropping its partial bytes if it fails. + fn buffer_with(&mut self, encode: impl FnOnce(&mut Encoder<'_>) -> Result<(), EncodeError>) -> Result<(), Error> where V: Into + Copy, { let start = self.buffer.len(); let mut w = Encoder::new(&mut self.buffer, self.version.into()); - if let Err(err) = msg.encode(&mut w, self.version) { + if let Err(err) = encode(&mut w) { // Drop the partial encode: flushing it would corrupt the stream. self.buffer.truncate(start); return Err(err.into()); @@ -81,6 +98,15 @@ impl Writer { std::future::poll_fn(|cx| self.poll_flush(cx)).await } + /// Encode a varint to the stream, with the same cancellation rules as [`Self::encode`]. + pub async fn varint(&mut self, value: u64) -> Result<(), Error> + where + V: Into + Copy, + { + self.buffer_varint(value)?; + std::future::poll_fn(|cx| self.poll_flush(cx)).await + } + /// Poll a write of `buf`, flushing any buffered message bytes first so the /// stream never reorders around the buffer. pub fn poll_write(&mut self, cx: &mut Context<'_>, buf: &mut Buf) -> Poll> { @@ -200,7 +226,7 @@ impl Writer { impl Writer { /// Encode an IETF `Message` to the stream, writing `[type_id][size][body]`. pub async fn encode_message(&mut self, msg: &T) -> Result<(), Error> { - self.buffer(&VarInt::from(T::ID))?; + self.buffer_varint(T::ID)?; self.encode(msg).await } } @@ -320,9 +346,9 @@ mod tests { fn a_failed_encode_leaves_no_partial_bytes() { let mut writer = Writer::new(SinkSend::new(Log::default()), crate::lite::Version::Lite05); - writer.buffer(&VarInt::from_u32(5)).unwrap(); + writer.buffer_varint(5).unwrap(); writer.buffer(&Poison).unwrap_err(); - writer.buffer(&VarInt::from_u32(7)).unwrap(); + writer.buffer_varint(7).unwrap(); let log = writer.stream.as_ref().unwrap().log.clone(); let mut cx = std::task::Context::from_waker(Waker::noop()); @@ -341,13 +367,13 @@ mod tests { ); let log = writer.stream.as_ref().unwrap().log.clone(); - writer.buffer(&VarInt::from_u32(5)).unwrap(); + writer.buffer_varint(5).unwrap(); let mut cx = std::task::Context::from_waker(Waker::noop()); assert!(writer.poll_flush(&mut cx).is_pending()); assert!(log.writes.lock().unwrap().is_empty()); // The bytes survive the Pending (and a second message queued behind them). - writer.buffer(&VarInt::from_u32(7)).unwrap(); + writer.buffer_varint(7).unwrap(); let Ok(mut open) = gate.write() else { panic!("gate closed") }; diff --git a/rs/moq-net/src/fuzz.rs b/rs/moq-net/src/fuzz.rs index 6b3c991066..a28535d810 100644 --- a/rs/moq-net/src/fuzz.rs +++ b/rs/moq-net/src/fuzz.rs @@ -14,7 +14,7 @@ use bytes::Bytes; use crate::{ Hops, Path, PathOwned, Pattern, - coding::{Decode, Decoder, Encode, Encoder, Form, VarInt}, + coding::{Decode, Decoder, Encode, Encoder, Form, varint}, ietf, lite, path::Relative, }; @@ -327,7 +327,7 @@ pub fn ietf_wire(data: &[u8]) -> bool { /// disagree about which byte sequences are even legal. /// /// The leading-ones form spans the full `u64`, while the QUIC form stops at -/// [`VarInt::MAX_QUIC`]; a value always re-encodes in the form it was read in. +/// [`crate::coding::varint::MAX_QUIC`]; a value always re-encodes in the form it was read in. pub fn varint(data: &[u8]) -> bool { let Some((&selector, rest)) = data.split_first() else { return false; @@ -339,24 +339,21 @@ pub fn varint(data: &[u8]) -> bool { _ => IETF_VERSIONS[(selector as usize / 2) % IETF_VERSIONS.len()].into(), }; - let Ok((value, _)) = VarInt::decode_slice(rest, version) else { + let Ok(value) = Decoder::new(rest, version.into()).varint() else { return false; }; // Zigzag is a bijection on top of the wire value, so it must round-trip. - let signed = value.to_zigzag(); - assert_eq!(VarInt::from_zigzag(signed), value, "zigzag is not its own inverse"); + let signed = varint::unzigzag(value); + assert_eq!(varint::zigzag(signed), value, "zigzag is not its own inverse"); - let encoded = value - .encode_bytes(version) + let mut encoded = Vec::new(); + Encoder::new(&mut encoded, version.into()) + .varint(value) .expect("a varint re-encodes in the form it was read in"); - let (again, used) = VarInt::decode_slice(&encoded, version).expect("could not decode our own encoding"); - assert_eq!( - used, - encoded.len(), - "our own encoding left {} bytes", - encoded.len() - used - ); + let mut echo = Decoder::new(&encoded, version.into()); + let again = echo.varint().expect("could not decode our own encoding"); + assert!(echo.is_empty(), "our own encoding left {} bytes", echo.remaining()); assert_eq!(value, again, "varint did not survive a round trip"); true @@ -506,7 +503,7 @@ fn bench_form(ietf: bool) -> Form { pub fn encode_varints(values: &[u64], ietf: bool, out: &mut Vec) { let mut w = Encoder::new(out, bench_form(ietf)); for value in values { - w.varint((*value).into()).unwrap(); + w.varint(*value).unwrap(); } } @@ -515,7 +512,7 @@ pub fn decode_varints(data: &[u8], ietf: bool) -> u64 { let mut r = Decoder::new(data, bench_form(ietf)); let mut sum = 0u64; while !r.is_empty() { - sum = sum.wrapping_add(r.varint().unwrap().into_inner()); + sum = sum.wrapping_add(r.varint().unwrap()); } sum } diff --git a/rs/moq-net/src/ietf/adapter.rs b/rs/moq-net/src/ietf/adapter.rs index 578b805d3e..a08324663a 100644 --- a/rs/moq-net/src/ietf/adapter.rs +++ b/rs/moq-net/src/ietf/adapter.rs @@ -8,7 +8,7 @@ use bytes::{Buf, BufMut, Bytes, BytesMut}; use crate::{ Error, PathOwned, - coding::{Decode, Decoder, Encoder, Reader, VarInt, Writer}, + coding::{Decode, Decoder, Encoder, Reader, Writer}, ietf::{self, RequestId}, }; @@ -291,7 +291,7 @@ impl OutgoingRegistration { let request_id = RequestId::decode(&mut body, self.version)?; // For PublishNamespace, also extract the namespace for reverse lookup. - if type_id.into_inner() == ietf::PublishNamespace::ID { + if type_id == ietf::PublishNamespace::ID { if self.version == Version::Draft17 { // v17 has required_request_id_delta after request_id let _ = body.varint(); @@ -784,7 +784,7 @@ impl ControlStreamAdapter { let mut raw = Vec::new(); let mut w = Encoder::new(&mut raw, version.into()); if let Err(err) = w - .varint(crate::ietf::GoAway::ID.into()) + .varint(crate::ietf::GoAway::ID) .and_then(|()| msg.encode(&mut w, version)) { tracing::warn!(%err, "failed to encode goaway"); @@ -812,8 +812,8 @@ impl ControlStreamAdapter { goaway: crate::goaway::Protocol, ) -> Result<(), Error> { loop { - let type_id = match reader.decode_maybe::().await? { - Some(id) => id.into_inner(), + let type_id = match reader.varint_maybe().await? { + Some(id) => id, None => return Ok(()), }; @@ -1159,7 +1159,7 @@ enum Route { fn encode_raw(type_id: u64, body: &Bytes, version: Version) -> Bytes { let mut buf = Vec::new(); let mut w = Encoder::new(&mut buf, version.into()); - w.varint(type_id.into()).expect("type_id was read from the same wire"); + w.varint(type_id).expect("type_id was read from the same wire"); w.u16(u16::try_from(body.len()).expect("body was read with a u16 size")); w.slice(body); buf.into() @@ -1310,7 +1310,7 @@ mod tests { // Decode the raw bytes let mut r = Decoder::new(&raw, version.into()); - assert_eq!(r.varint().unwrap().into_inner(), 0x03); + assert_eq!(r.varint().unwrap(), 0x03); assert_eq!(r.u16().unwrap(), 5); assert_eq!(r.rest(), b"hello"); } diff --git a/rs/moq-net/src/ietf/fetch.rs b/rs/moq-net/src/ietf/fetch.rs index e9d1db939d..3281bf50c2 100644 --- a/rs/moq-net/src/ietf/fetch.rs +++ b/rs/moq-net/src/ietf/fetch.rs @@ -2,7 +2,7 @@ use std::borrow::Cow; use crate::{ Path, - coding::{Decode, DecodeError, Decoder, Encode, EncodeError, Encoder, VarInt}, + coding::{Decode, DecodeError, Decoder, Encode, EncodeError, Encoder}, ietf::{ GroupOrder, Location, Parameters, RequestId, namespace::{decode_namespace, encode_namespace}, @@ -53,7 +53,7 @@ impl Encode for FetchType<'_> { } => { w.u8(2); subscriber_request_id.encode(w, version)?; - w.varint(VarInt::from(*group_offset))?; + w.varint(*group_offset)?; } FetchType::AbsoluteJoining { subscriber_request_id, @@ -61,7 +61,7 @@ impl Encode for FetchType<'_> { } => { w.u8(3); subscriber_request_id.encode(w, version)?; - w.varint(VarInt::from(*group_id))?; + w.varint(*group_id)?; } } Ok(()) @@ -70,7 +70,7 @@ impl Encode for FetchType<'_> { impl Decode for FetchType<'_> { fn decode(buf: &mut Decoder<'_>, version: Version) -> Result { - let fetch_type = buf.varint()?.into_inner(); + let fetch_type = buf.varint()?; Ok(match fetch_type { 0x1 => { let namespace = decode_namespace(buf)?; @@ -86,7 +86,7 @@ impl Decode for FetchType<'_> { } 0x2 => { let subscriber_request_id = RequestId::decode(buf, version)?; - let group_offset = buf.varint()?.into_inner(); + let group_offset = buf.varint()?; FetchType::RelativeJoining { subscriber_request_id, group_offset, @@ -94,7 +94,7 @@ impl Decode for FetchType<'_> { } 0x3 => { let subscriber_request_id = RequestId::decode(buf, version)?; - let group_id = buf.varint()?.into_inner(); + let group_id = buf.varint()?; FetchType::AbsoluteJoining { subscriber_request_id, group_id, @@ -119,7 +119,7 @@ impl Message for Fetch<'_> { fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.request_id.encode(w, version)?; if version == Version::Draft17 { - w.varint(VarInt::ZERO)?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) + w.varint(0)?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) } match version { @@ -143,7 +143,7 @@ impl Message for Fetch<'_> { fn decode_msg(buf: &mut Decoder<'_>, version: Version) -> Result { let request_id = RequestId::decode(buf, version)?; if version == Version::Draft17 { - let _required_request_id_delta = buf.varint()?.into_inner(); + let _required_request_id_delta = buf.varint()?; } match version { @@ -274,14 +274,14 @@ impl Message for FetchError<'_> { fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.request_id.encode(w, version)?; - w.varint(VarInt::from(self.error_code))?; + w.varint(self.error_code)?; w.string(&self.reason_phrase)?; Ok(()) } fn decode_msg(buf: &mut Decoder<'_>, version: Version) -> Result { let request_id = RequestId::decode(buf, version)?; - let error_code = buf.varint()?.into_inner(); + let error_code = buf.varint()?; let reason_phrase = Cow::Owned(buf.string()?); Ok(Self { request_id, @@ -417,9 +417,9 @@ impl Encode for FetchObject { if !Self::END_OF_RANGE.contains(reason) { return Err(EncodeError::InvalidState); } - w.varint(VarInt::from(*reason))?; - w.varint(VarInt::from(*group))?; - w.varint(VarInt::from(*object))?; + w.varint(*reason)?; + w.varint(*group)?; + w.varint(*object)?; } Self::Object { subgroup, @@ -447,16 +447,16 @@ impl Encode for FetchObject { if properties.is_some() { flags |= flag::PROPERTIES; } - w.varint(VarInt::from(flags))?; + w.varint(flags)?; if let Some(group) = group { - w.varint(VarInt::from(*group))?; + w.varint(*group)?; } if let FetchSubgroup::Explicit(subgroup) = subgroup { - w.varint(VarInt::from(*subgroup))?; + w.varint(*subgroup)?; } if let Some(object) = object { - w.varint(VarInt::from(*object))?; + w.varint(*object)?; } if let Some(priority) = priority { w.u8(*priority); @@ -472,7 +472,7 @@ impl Encode for FetchObject { impl Decode for FetchObject { fn decode(buf: &mut Decoder<'_>, _: Version) -> Result { - let flags = buf.varint()?.into_inner(); + let flags = buf.varint()?; // Anything at or above 128 is a named value rather than a set of flags, and only // the three End of Range markers are defined. @@ -482,14 +482,14 @@ impl Decode for FetchObject { } return Ok(Self::EndOfRange { reason: flags, - group: buf.varint()?.into_inner(), - object: buf.varint()?.into_inner(), + group: buf.varint()?, + object: buf.varint()?, }); } // Wire order: Group ID Delta, Subgroup ID, Object ID Delta, Priority, Properties. let group = match flags & flag::GROUP_ID != 0 { - true => Some(buf.varint()?.into_inner()), + true => Some(buf.varint()?), false => None, }; @@ -499,12 +499,12 @@ impl Decode for FetchObject { 0 => FetchSubgroup::Zero, 1 => FetchSubgroup::Prior, 2 => FetchSubgroup::PriorPlusOne, - _ => FetchSubgroup::Explicit(buf.varint()?.into_inner()), + _ => FetchSubgroup::Explicit(buf.varint()?), }, }; let object = match flags & flag::OBJECT_ID != 0 { - true => Some(buf.varint()?.into_inner()), + true => Some(buf.varint()?), false => None, }; diff --git a/rs/moq-net/src/ietf/filter.rs b/rs/moq-net/src/ietf/filter.rs index 3ea0d99a2a..98e1d3f5b1 100644 --- a/rs/moq-net/src/ietf/filter.rs +++ b/rs/moq-net/src/ietf/filter.rs @@ -1,6 +1,6 @@ //! The Location Filter carried by SUBSCRIBE, PUBLISH and REQUEST_UPDATE. -use crate::coding::{Decode, DecodeError, Decoder, Encode, EncodeError, Encoder, VarInt}; +use crate::coding::{Decode, DecodeError, Decoder, Encode, EncodeError, Encoder}; use super::{Location, Param, Version}; @@ -86,21 +86,21 @@ impl Filter { // that spelling normalizes to this one rather than colliding with NextObject. Self::Unfiltered => {} Self::NextObject => { - w.varint(VarInt::ZERO)?; - w.varint(VarInt::ZERO)?; + w.varint(0)?; + w.varint(0)?; } - Self::Relative(groups) => w.varint(VarInt::from(groups))?, + Self::Relative(groups) => w.varint(groups)?, Self::Absolute { start: Location { group: 0, object: 0 }, end: None, } => {} Self::Absolute { start, end } => { - w.varint(VarInt::from(start.group))?; - w.varint(VarInt::from(start.object))?; + w.varint(start.group)?; + w.varint(start.object)?; if let Some(end) = end { - w.varint(VarInt::from(Self::end_delta(start.group, end.group)?))?; + w.varint(Self::end_delta(start.group, end.group)?)?; if let Some(object) = end.object { - w.varint(VarInt::from(object))?; + w.varint(object)?; } } } @@ -115,7 +115,7 @@ impl Filter { if fields.len() == 4 { return Err(DecodeError::TrailingBytes); } - fields.push(r.varint()?.into_inner()); + fields.push(r.varint()?); } Ok(match fields[..] { @@ -150,16 +150,16 @@ impl Filter { match *self { // No tag means "everything", which only the absolute spelling can say. Self::Unfiltered => { - w.varint(VarInt::from(tag::ABSOLUTE_START))?; + w.varint(tag::ABSOLUTE_START)?; Location::default().encode(w, version)?; } - Self::NextObject => w.varint(VarInt::from(tag::LARGEST_OBJECT))?, - Self::Relative(0) => w.varint(VarInt::from(tag::NEXT_GROUP))?, + Self::NextObject => w.varint(tag::LARGEST_OBJECT)?, + Self::Relative(0) => w.varint(tag::NEXT_GROUP)?, // Only draft-20 can name a start further back than the next group without // knowing Largest Object, so there is no honest tag to fall back to. Self::Relative(_) => return Err(EncodeError::Unsupported), Self::Absolute { start, end: None } => { - w.varint(VarInt::from(tag::ABSOLUTE_START))?; + w.varint(tag::ABSOLUTE_START)?; start.encode(w, version)?; } // Draft-19's AbsoluteRange ends on a group, so an object-bounded range has no @@ -169,9 +169,9 @@ impl Filter { .. } => return Err(EncodeError::Unsupported), Self::Absolute { start, end: Some(end) } => { - w.varint(VarInt::from(tag::ABSOLUTE_RANGE))?; + w.varint(tag::ABSOLUTE_RANGE)?; start.encode(w, version)?; - w.varint(VarInt::from(Self::end_delta(start.group, end.group)?))?; + w.varint(Self::end_delta(start.group, end.group)?)?; } } Ok(()) @@ -179,7 +179,7 @@ impl Filter { /// Decode the draft-19 and earlier tag form. fn decode_tag(r: &mut Decoder<'_>, version: Version) -> Result { - Ok(match r.varint()?.into_inner() { + Ok(match r.varint()? { tag::NEXT_GROUP => Self::Relative(0), tag::LARGEST_OBJECT => Self::NextObject, tag::ABSOLUTE_START => Self::Absolute { @@ -188,7 +188,7 @@ impl Filter { }, tag::ABSOLUTE_RANGE => { let start = Location::decode(r, version)?; - let delta = r.varint()?.into_inner(); + let delta = r.varint()?; Self::Absolute { start, end: Some(EndLocation { @@ -555,11 +555,11 @@ impl Param for Fill { // An omitted filter inherits the subscription's, so the scope is empty. An explicit // Unfiltered still encodes, as a zero-length filter meaning the whole track. match self.filter { - None => inner.varint(VarInt::ZERO)?, + None => inner.varint(0)?, Some(filter) => { - inner.varint(VarInt::from(1u64))?; + inner.varint(1u64)?; // The first type in a scope is not delta encoded, so this is the raw id. - inner.varint(VarInt::from(Self::LOCATION_FILTER))?; + inner.varint(Self::LOCATION_FILTER)?; filter.param_encode(&mut inner, version)?; } } @@ -570,7 +570,7 @@ impl Param for Fill { fn param_decode(r: &mut Decoder<'_>, version: Version) -> Result { let mut buf = Decoder::new(r.bytes()?, r.form()); - let count = buf.varint()?.into_inner(); + let count = buf.varint()?; if count > 64 { return Err(DecodeError::TooMany); } @@ -579,7 +579,7 @@ impl Param for Fill { let mut range_filters = false; let mut prev = 0u64; for i in 0..count { - let delta = buf.varint()?.into_inner(); + let delta = buf.varint()?; let key = if i == 0 { delta } else { @@ -686,11 +686,9 @@ mod fill_tests { #[test] fn rejects_a_disallowed_parameter() { let mut value = Vec::new(); - Encoder::new(&mut value, NEW.into()).varint(VarInt::from(1u64)).unwrap(); - Encoder::new(&mut value, NEW.into()) - .varint(VarInt::from(0x10u64)) - .unwrap(); // FORWARD, not allowed in a fill - Encoder::new(&mut value, NEW.into()).varint(VarInt::ZERO).unwrap(); + Encoder::new(&mut value, NEW.into()).varint(1u64).unwrap(); + Encoder::new(&mut value, NEW.into()).varint(0x10u64).unwrap(); // FORWARD, not allowed in a fill + Encoder::new(&mut value, NEW.into()).varint(0).unwrap(); let mut buf = Vec::new(); Encoder::new(&mut buf, NEW.into()).bytes(&value).unwrap(); @@ -703,12 +701,10 @@ mod fill_tests { #[test] fn skips_a_uint8_whose_value_has_a_leading_one() { let mut value = Vec::new(); - Encoder::new(&mut value, NEW.into()).varint(VarInt::from(2u64)).unwrap(); - Encoder::new(&mut value, NEW.into()) - .varint(VarInt::from(0x20u64)) - .unwrap(); // SUBSCRIBER_PRIORITY + Encoder::new(&mut value, NEW.into()).varint(2u64).unwrap(); + Encoder::new(&mut value, NEW.into()).varint(0x20u64).unwrap(); // SUBSCRIBER_PRIORITY Encoder::new(&mut value, NEW.into()).u8(0x80u8); // a raw byte, not a varint - Encoder::new(&mut value, NEW.into()).varint(VarInt::from(1u64)).unwrap(); // delta to 0x21 + Encoder::new(&mut value, NEW.into()).varint(1u64).unwrap(); // delta to 0x21 Filter::Relative(1) .param_encode(&mut Encoder::new(&mut value, NEW.into()), NEW) .unwrap(); @@ -727,14 +723,12 @@ mod fill_tests { #[test] fn skips_a_length_prefixed_range_filter() { let mut value = Vec::new(); - Encoder::new(&mut value, NEW.into()).varint(VarInt::from(2u64)).unwrap(); - Encoder::new(&mut value, NEW.into()) - .varint(VarInt::from(0x26u64)) - .unwrap(); // OBJECTID_FILTER, length prefixed + Encoder::new(&mut value, NEW.into()).varint(2u64).unwrap(); + Encoder::new(&mut value, NEW.into()).varint(0x26u64).unwrap(); // OBJECTID_FILTER, length prefixed Encoder::new(&mut value, NEW.into()) .bytes(&[0xAAu8, 0xBB, 0xCC]) .unwrap(); - Encoder::new(&mut value, NEW.into()).varint(VarInt::from(1u64)).unwrap(); // delta to 0x27 + Encoder::new(&mut value, NEW.into()).varint(1u64).unwrap(); // delta to 0x27 Encoder::new(&mut value, NEW.into()).bytes(&[0xDDu8]).unwrap(); // PRIORITY_FILTER let mut buf = Vec::new(); @@ -748,14 +742,10 @@ mod fill_tests { #[test] fn skips_allowed_parameters_it_ignores() { let mut value = Vec::new(); - Encoder::new(&mut value, NEW.into()).varint(VarInt::from(2u64)).unwrap(); - Encoder::new(&mut value, NEW.into()) - .varint(VarInt::from(0x20u64)) - .unwrap(); // SUBSCRIBER_PRIORITY, even: one varint - Encoder::new(&mut value, NEW.into()) - .varint(VarInt::from(42u64)) - .unwrap(); - Encoder::new(&mut value, NEW.into()).varint(VarInt::from(1u64)).unwrap(); // delta to 0x21, odd: length prefixed + Encoder::new(&mut value, NEW.into()).varint(2u64).unwrap(); + Encoder::new(&mut value, NEW.into()).varint(0x20u64).unwrap(); // SUBSCRIBER_PRIORITY, even: one varint + Encoder::new(&mut value, NEW.into()).varint(42u64).unwrap(); + Encoder::new(&mut value, NEW.into()).varint(1u64).unwrap(); // delta to 0x21, odd: length prefixed Filter::Relative(2) .param_encode(&mut Encoder::new(&mut value, NEW.into()), NEW) .unwrap(); diff --git a/rs/moq-net/src/ietf/goaway.rs b/rs/moq-net/src/ietf/goaway.rs index 0812742d9f..f132dc2b86 100644 --- a/rs/moq-net/src/ietf/goaway.rs +++ b/rs/moq-net/src/ietf/goaway.rs @@ -23,7 +23,7 @@ impl Message for GoAway<'_> { w.string(&self.new_session_uri)?; // Draft-17+ adds a timeout field. if !matches!(version, Version::Draft14 | Version::Draft15 | Version::Draft16) { - w.varint(VarInt::from(self.timeout))?; + w.varint(self.timeout)?; } // Draft-18 (#1559) requires a Request ID when GOAWAY is sent on the // control stream, which is the only place we send it. We don't track @@ -32,7 +32,7 @@ impl Message for GoAway<'_> { // conformant peer must treat as a PROTOCOL_VIOLATION. Draft-19 // removed the field again (#1623). if matches!(version, Version::Draft18) { - w.varint(VarInt::ZERO)?; + w.varint(0)?; } Ok(()) } @@ -47,17 +47,17 @@ impl Message for GoAway<'_> { let timeout = match version { Version::Draft14 | Version::Draft15 | Version::Draft16 => 0, Version::Draft18 => { - let timeout = r.varint()?.into_inner(); + let timeout = r.varint()?; // Draft-18 trailing Request ID (#1559): required on the control // stream, but tolerate its absence from lenient peers. We don't // act on per-request GOAWAY so the value is discarded. Draft-19 // removed this field again (#1623). if !r.is_empty() { - let _ = r.varint()?.into_inner(); + let _ = r.varint()?; } timeout } - _ => r.varint()?.into_inner(), + _ => r.varint()?, }; Ok(Self { new_session_uri, @@ -180,13 +180,9 @@ mod tests { Encoder::new(&mut buf, Version::Draft18.into()) .string("moqt://relay.example/") .unwrap(); - Encoder::new(&mut buf, Version::Draft18.into()) - .varint(VarInt::from(5000u64)) - .unwrap(); + Encoder::new(&mut buf, Version::Draft18.into()).varint(5000u64).unwrap(); // Optional trailing Request ID: - Encoder::new(&mut buf, Version::Draft18.into()) - .varint(VarInt::from(42u64)) - .unwrap(); + Encoder::new(&mut buf, Version::Draft18.into()).varint(42u64).unwrap(); let mut bytes = bytes::Bytes::from(buf.to_vec()); let decoded: GoAway = crate::coding::decode_buf(&mut bytes, Version::Draft18, GoAway::decode_msg).unwrap(); diff --git a/rs/moq-net/src/ietf/group.rs b/rs/moq-net/src/ietf/group.rs index c577b77715..fa2ae4fdf0 100644 --- a/rs/moq-net/src/ietf/group.rs +++ b/rs/moq-net/src/ietf/group.rs @@ -1,4 +1,4 @@ -use crate::coding::{Decode, DecodeError, Decoder, Encode, EncodeError, Encoder, VarInt}; +use crate::coding::{Decode, DecodeError, Decoder, Encode, EncodeError, Encoder}; use crate::{Timescale, Timestamp}; use num_enum::{IntoPrimitive, TryFromPrimitive}; @@ -37,7 +37,7 @@ pub fn encode_object_time( ) -> Result<(), EncodeError> { let timestamp = timestamp.convert(timescale).map_err(|_| EncodeError::BoundsExceeded)?; encode_object_property_type(w, PROP_TIMESTAMP, 0, version)?; - w.varint(VarInt::from(timestamp.value()))?; + w.varint(timestamp.value())?; Ok(()) } @@ -46,7 +46,7 @@ fn encode_object_property_type(w: &mut Encoder<'_>, kind: u64, prev: u64, versio Version::Draft14 | Version::Draft15 => kind, _ => kind.checked_sub(prev).ok_or(EncodeError::BoundsExceeded)?, }; - w.varint(VarInt::from(encoded)) + w.varint(encoded) } /// Decode the Timestamp (0x10) Object Property from an object's extension block, @@ -66,7 +66,7 @@ pub fn decode_object_time( let mut first = true; while !r.is_empty() { - let step = r.varint()?.into_inner(); + let step = r.varint()?; let abs = match version { Version::Draft14 | Version::Draft15 => step, _ if first => step, @@ -77,7 +77,7 @@ pub fn decode_object_time( if abs % 2 == 0 { // Even type: a single varint value. - let value = r.varint()?.into_inner(); + let value = r.varint()?; match abs { PROP_TIMESTAMP | PROP_TIMESTAMP_DRAFT03 => timestamp = Some(value), PROP_TIMESCALE => override_scale = Some(value), @@ -85,7 +85,7 @@ pub fn decode_object_time( } } else { // Odd type: length-prefixed bytes we don't care about. - let len = usize::try_from(r.varint()?)?; + let len = usize::try_from(r.varint()?).map_err(|_| DecodeError::BoundsExceeded)?; r.slice(len)?; } } @@ -287,16 +287,16 @@ pub struct GroupHeader { impl Encode for GroupHeader { fn encode(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { tracing::trace!(?self, "encoding group header"); - w.varint(VarInt::from(self.flags.encode(version)?))?; - w.varint(VarInt::from(self.track_alias))?; - w.varint(VarInt::from(self.group_id))?; + w.varint(self.flags.encode(version)?)?; + w.varint(self.track_alias)?; + w.varint(self.group_id)?; if !self.flags.has_subgroup && self.sub_group_id != 0 { return Err(EncodeError::InvalidState); } if self.flags.has_subgroup { - w.varint(VarInt::from(self.sub_group_id))?; + w.varint(self.sub_group_id)?; } // Publisher priority (only if has_priority flag is set) @@ -309,12 +309,12 @@ impl Encode for GroupHeader { impl Decode for GroupHeader { fn decode(r: &mut Decoder<'_>, version: Version) -> Result { - let flags = GroupFlags::decode(r.varint()?.into_inner(), version)?; - let track_alias = r.varint()?.into_inner(); - let group_id = r.varint()?.into_inner(); + let flags = GroupFlags::decode(r.varint()?, version)?; + let track_alias = r.varint()?; + let group_id = r.varint()?; let sub_group_id = match flags.has_subgroup { - true => r.varint()?.into_inner(), + true => r.varint()?, false => 0, }; @@ -351,7 +351,7 @@ mod tests { /// Read `buf` back as a flat list of varints. fn varints(buf: &[u8], version: Version) -> Vec { let mut r = Decoder::new(buf, version.into()); - std::iter::from_fn(|| (!r.is_empty()).then(|| r.varint().unwrap().into_inner())).collect() + std::iter::from_fn(|| (!r.is_empty()).then(|| r.varint().unwrap())).collect() } /// Write `values` as a flat list of varints. @@ -359,7 +359,7 @@ mod tests { let mut buf = Vec::new(); let mut w = Encoder::new(&mut buf, version.into()); for value in values { - w.varint((*value).into()).unwrap(); + w.varint(*value).unwrap(); } buf } diff --git a/rs/moq-net/src/ietf/location.rs b/rs/moq-net/src/ietf/location.rs index 5a3c4f174a..f6e95479e7 100644 --- a/rs/moq-net/src/ietf/location.rs +++ b/rs/moq-net/src/ietf/location.rs @@ -1,4 +1,4 @@ -use crate::coding::{Decode, DecodeError, Decoder, Encode, EncodeError, Encoder, VarInt}; +use crate::coding::{Decode, DecodeError, Decoder, Encode, EncodeError, Encoder}; use super::Version; @@ -10,16 +10,16 @@ pub struct Location { impl Encode for Location { fn encode(&self, w: &mut Encoder<'_>, _: Version) -> Result<(), EncodeError> { - w.varint(VarInt::from(self.group))?; - w.varint(VarInt::from(self.object))?; + w.varint(self.group)?; + w.varint(self.object)?; Ok(()) } } impl Decode for Location { fn decode(buf: &mut Decoder<'_>, _: Version) -> Result { - let group = buf.varint()?.into_inner(); - let object = buf.varint()?.into_inner(); + let group = buf.varint()?; + let object = buf.varint()?; Ok(Self { group, object }) } } diff --git a/rs/moq-net/src/ietf/namespace.rs b/rs/moq-net/src/ietf/namespace.rs index e82eaad8a0..53d630c8ee 100644 --- a/rs/moq-net/src/ietf/namespace.rs +++ b/rs/moq-net/src/ietf/namespace.rs @@ -55,7 +55,7 @@ pub fn encode_namespace(w: &mut Encoder<'_>, namespace: &Path) -> Result<(), Enc return Err(BoundsExceeded.into()); } - w.varint(VarInt::from(parts.len()))?; + w.varint(parts.len() as u64)?; for part in parts { w.string(&part)?; } @@ -64,7 +64,7 @@ pub fn encode_namespace(w: &mut Encoder<'_>, namespace: &Path) -> Result<(), Enc /// Helper function to decode namespace from tuple of strings pub fn decode_namespace(r: &mut Decoder<'_>) -> Result, DecodeError> { - let count = r.varint()?.into_inner(); + let count = r.varint()?; if count == 0 { return Ok(Path::from(String::new())); @@ -109,7 +109,7 @@ mod tests { fn encode_tuple(parts: &[&str]) -> Vec { let mut buf = Vec::new(); let mut w = Encoder::new(&mut buf, FORM); - w.varint(parts.len().into()).unwrap(); + w.varint(parts.len() as u64).unwrap(); for part in parts { w.string(part).unwrap(); } diff --git a/rs/moq-net/src/ietf/parameters.rs b/rs/moq-net/src/ietf/parameters.rs index 322366215b..5625c51976 100644 --- a/rs/moq-net/src/ietf/parameters.rs +++ b/rs/moq-net/src/ietf/parameters.rs @@ -57,7 +57,7 @@ impl Decode for Parameters { // Draft-14/15/16 count the pairs; draft-17+ reads them until the buffer is empty. let count = match version { - Version::Draft14 | Version::Draft15 | Version::Draft16 => Some(r.varint()?.into_inner()), + Version::Draft14 | Version::Draft15 | Version::Draft16 => Some(r.varint()?), _ => None, }; if count.is_some_and(|count| count > MAX_PARAMS) { @@ -74,7 +74,7 @@ impl Decode for Parameters { return Err(DecodeError::TooMany); } - let kind = r.varint()?.into_inner(); + let kind = r.varint()?; let kind = match delta && i > 0 { true => prev.checked_add(kind).ok_or(DecodeError::BoundsExceeded)?, false => kind, @@ -87,7 +87,7 @@ impl Decode for Parameters { if params.get_varint(kind).is_some() { return Err(DecodeError::Duplicate); } - params.vars.push((kind, r.varint()?.into_inner())); + params.vars.push((kind, r.varint()?)); } else { let kind = ParameterBytes::from(kind); let value = r.bytes()?; @@ -117,15 +117,15 @@ impl Encode for Parameters { match version { Version::Draft14 | Version::Draft15 => { - w.varint(count.into())?; + w.varint(count as u64)?; for (kind, value) in &self.vars { - w.varint(u64::from(*kind).into())?; - w.varint((*value).into())?; + w.varint(u64::from(*kind))?; + w.varint(*value)?; } for (kind, value) in &self.bytes { - w.varint(u64::from(*kind).into())?; + w.varint(u64::from(*kind))?; w.bytes(value)?; } } @@ -133,7 +133,7 @@ impl Encode for Parameters { // Draft16: count prefix + delta encoding // Draft17+: NO count prefix + delta encoding if matches!(version, Version::Draft16) { - w.varint(count.into())?; + w.varint(count as u64)?; } enum ParamRef<'a> { @@ -147,11 +147,11 @@ impl Encode for Parameters { let mut prev = 0u64; for (kind, value) in all { - w.varint((kind - prev).into())?; + w.varint(kind - prev)?; prev = kind; match value { - ParamRef::Var(v) => w.varint(v.into())?, + ParamRef::Var(v) => w.varint(v)?, ParamRef::Bytes(v) => w.bytes(v)?, } } @@ -209,7 +209,7 @@ impl Param for u8 { fn param_encode(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { match version { // Draft-14/15/16: u8 encoded as varint - Version::Draft14 | Version::Draft15 | Version::Draft16 => w.varint((*self).into())?, + Version::Draft14 | Version::Draft15 | Version::Draft16 => w.varint(u64::from(*self))?, _ => w.u8(*self), } Ok(()) @@ -237,7 +237,7 @@ impl Param for bool { fn param_decode(r: &mut Decoder<'_>, version: Version) -> Result { match version { - Version::Draft14 | Version::Draft15 | Version::Draft16 => match r.varint()?.into_inner() { + Version::Draft14 | Version::Draft15 | Version::Draft16 => match r.varint()? { 0 => Ok(false), 1 => Ok(true), _ => Err(DecodeError::InvalidValue), @@ -249,12 +249,12 @@ impl Param for bool { impl Param for u64 { fn param_encode(&self, w: &mut Encoder<'_>, _: Version) -> Result<(), EncodeError> { - w.varint((*self).into())?; + w.varint(*self)?; Ok(()) } fn param_decode(r: &mut Decoder<'_>, _: Version) -> Result { - Ok(r.varint()?.into_inner()) + r.varint() } } @@ -274,13 +274,13 @@ impl Param for Location { // matching the other length-prefixed parameters. let mut buf = Vec::new(); let mut inner = Encoder::new(&mut buf, Version::Draft15.into()); - inner.varint(self.group.into())?; - inner.varint(self.object.into())?; + inner.varint(self.group)?; + inner.varint(self.object)?; w.bytes(&buf) } _ => { - w.varint(self.group.into())?; - w.varint(self.object.into())?; + w.varint(self.group)?; + w.varint(self.object)?; Ok(()) } } @@ -290,16 +290,16 @@ impl Param for Location { match version { Version::Draft14 | Version::Draft15 | Version::Draft16 => { let mut inner = Decoder::new(r.bytes()?, Version::Draft15.into()); - let group = inner.varint()?.into_inner(); - let object = inner.varint()?.into_inner(); + let group = inner.varint()?; + let object = inner.varint()?; if !inner.is_empty() { return Err(DecodeError::TrailingBytes); } Ok(Location { group, object }) } _ => { - let group = r.varint()?.into_inner(); - let object = r.varint()?.into_inner(); + let group = r.varint()?; + let object = r.varint()?; Ok(Location { group, object }) } } @@ -351,7 +351,7 @@ macro_rules! encode_params { #[allow(unused_mut)] let mut _count: usize = 0; $(_count += if $crate::ietf::Param::param_present(&$val) { 1 } else { 0 };)* - $w.varint($crate::coding::VarInt::from(_count))?; + $w.varint(_count as u64)?; #[allow(unused_mut, unused_assignments)] let mut _prev_key: u64 = 0; @@ -365,7 +365,7 @@ macro_rules! encode_params { _ if _first => _key, _ => _key - _prev_key, }; - $w.varint($crate::coding::VarInt::from(_wire))?; + $w.varint(_wire)?; _prev_key = _key; _first = false; $crate::ietf::Param::param_encode(&$val, $w, _version)?; @@ -408,7 +408,7 @@ macro_rules! decode_params { { let _version: $crate::ietf::Version = $version; - let _count = $r.varint()?.into_inner(); + let _count = $r.varint()?; if _count > 64 { return Err($crate::coding::DecodeError::TooMany); } @@ -416,7 +416,7 @@ macro_rules! decode_params { #[allow(unused_mut, unused_assignments)] let mut _prev_key: u64 = 0; for _i in 0.._count { - let _wire = $r.varint()?.into_inner(); + let _wire = $r.varint()?; let _key: u64 = match _version { $crate::ietf::Version::Draft14 | $crate::ietf::Version::Draft15 => _wire, _ if _i == 0 => _wire, @@ -867,8 +867,8 @@ mod tests { ] { let mut buf = Vec::new(); let mut w = Encoder::new(&mut buf, version.into()); - w.varint(VarInt::from(1usize)).unwrap(); - w.varint(VarInt::from(0x10u64)).unwrap(); + w.varint(1).unwrap(); + w.varint(0x10u64).unwrap(); true.param_encode(&mut w, version).unwrap(); let mut bytes = Decoder::new(&buf, version.into()); @@ -897,20 +897,20 @@ mod tests { let mut buf = Vec::new(); let mut w = Encoder::new(&mut buf, version.into()); // Encode count = 2 - w.varint(VarInt::from(2usize)).unwrap(); + w.varint(2).unwrap(); match version { Version::Draft14 | Version::Draft15 => { // Plain (non-delta) keys: first key=0x20, second key=0x20 - w.varint(VarInt::from(0x20u64)).unwrap(); + w.varint(0x20u64).unwrap(); 100u8.param_encode(&mut w, version).unwrap(); - w.varint(VarInt::from(0x20u64)).unwrap(); + w.varint(0x20u64).unwrap(); 200u8.param_encode(&mut w, version).unwrap(); } _ => { // Delta-encoded: first delta=0x20 (abs=0x20), second delta=0 (abs=0x20) - w.varint(VarInt::from(0x20u64)).unwrap(); + w.varint(0x20u64).unwrap(); 100u8.param_encode(&mut w, version).unwrap(); - w.varint(VarInt::ZERO).unwrap(); + w.varint(0).unwrap(); 200u8.param_encode(&mut w, version).unwrap(); } } diff --git a/rs/moq-net/src/ietf/properties.rs b/rs/moq-net/src/ietf/properties.rs index 3f8dc33d0f..a4ee99f29a 100644 --- a/rs/moq-net/src/ietf/properties.rs +++ b/rs/moq-net/src/ietf/properties.rs @@ -12,7 +12,7 @@ use std::time::Duration; use crate::Timescale; -use crate::coding::{DecodeError, Decoder, EncodeError, Encoder, VarInt}; +use crate::coding::{DecodeError, Decoder, EncodeError, Encoder}; use super::{GroupOrder, Version}; @@ -73,26 +73,26 @@ impl Properties { let mut prev_type = 0; if let Some(age) = self.max_cache_duration { - w.varint(VarInt::from(4u64))?; - w.varint(VarInt::try_from(age.as_millis())?)?; + w.varint(4u64)?; + w.varint(u64::try_from(age.as_millis()).map_err(|_| EncodeError::BoundsExceeded)?)?; prev_type = 4; } if let Some(timescale) = self.timescale { - w.varint(VarInt::from(TIMESCALE - prev_type))?; - w.varint(VarInt::from(u64::from(timescale)))?; + w.varint(TIMESCALE - prev_type)?; + w.varint(u64::from(timescale))?; prev_type = TIMESCALE; } if let Some(priority) = self.priority { - w.varint(VarInt::from(DEFAULT_PUBLISHER_PRIORITY - prev_type))?; - w.varint(VarInt::from(u64::from(priority)))?; + w.varint(DEFAULT_PUBLISHER_PRIORITY - prev_type)?; + w.varint(u64::from(priority))?; prev_type = DEFAULT_PUBLISHER_PRIORITY; } if let Some(group_order) = self.group_order { - w.varint(VarInt::from(DEFAULT_PUBLISHER_GROUP_ORDER - prev_type))?; - w.varint(VarInt::from(u64::from(u8::from(group_order))))?; + w.varint(DEFAULT_PUBLISHER_GROUP_ORDER - prev_type)?; + w.varint(u64::from(u8::from(group_order)))?; } Ok(()) @@ -126,7 +126,7 @@ impl Properties { return Err(DecodeError::TooMany); } - let delta = r.varint()?.into_inner(); + let delta = r.varint()?; let abs = if i == 0 { delta } else { @@ -137,7 +137,7 @@ impl Properties { if abs % 2 == 0 { // Even type: single varint value - let value = r.varint()?.into_inner(); + let value = r.varint()?; match abs { 4 => properties.max_cache_duration = Some(Duration::from_millis(value)), TIMESCALE => { @@ -161,7 +161,7 @@ impl Properties { } } else { // Odd type: length-prefixed bytes - let len = usize::try_from(r.varint()?)?; + let len = usize::try_from(r.varint()?).map_err(|_| DecodeError::BoundsExceeded)?; if len > MAX_KVP_VALUE_LEN { return Err(DecodeError::BoundsExceeded); } @@ -190,12 +190,8 @@ mod tests { fn test_skip_varint_property() { // Even type (0x02 = DELIVERY_TIMEOUT), varint value let mut buf = Vec::new(); - Encoder::new(&mut buf, Version::Draft17.into()) - .varint(VarInt::from(0x02u64)) - .unwrap(); // delta type - Encoder::new(&mut buf, Version::Draft17.into()) - .varint(VarInt::from(5000u64)) - .unwrap(); // value + Encoder::new(&mut buf, Version::Draft17.into()).varint(0x02u64).unwrap(); // delta type + Encoder::new(&mut buf, Version::Draft17.into()).varint(5000u64).unwrap(); // value let mut bytes = bytes::Bytes::from(buf); crate::coding::decode_buf(&mut bytes, Version::Draft17, Properties::decode).unwrap(); assert!(!!bytes.is_empty()); @@ -205,12 +201,8 @@ mod tests { fn test_skip_bytes_property() { // Odd type (0x0B = IMMUTABLE_PROPERTIES), length-prefixed let mut buf = Vec::new(); - Encoder::new(&mut buf, Version::Draft17.into()) - .varint(VarInt::from(0x0Bu64)) - .unwrap(); // delta type - Encoder::new(&mut buf, Version::Draft17.into()) - .varint(VarInt::from(3u64)) - .unwrap(); // length + Encoder::new(&mut buf, Version::Draft17.into()).varint(0x0Bu64).unwrap(); // delta type + Encoder::new(&mut buf, Version::Draft17.into()).varint(3u64).unwrap(); // length buf.extend_from_slice(&[0x01, 0x02, 0x03]); // value bytes let mut bytes = bytes::Bytes::from(buf); crate::coding::decode_buf(&mut bytes, Version::Draft17, Properties::decode).unwrap(); @@ -221,26 +213,14 @@ mod tests { fn test_skip_multiple_properties() { let mut buf = Vec::new(); // First: type 0x02 (even), varint value - Encoder::new(&mut buf, Version::Draft17.into()) - .varint(VarInt::from(0x02u64)) - .unwrap(); - Encoder::new(&mut buf, Version::Draft17.into()) - .varint(VarInt::from(1000u64)) - .unwrap(); + Encoder::new(&mut buf, Version::Draft17.into()).varint(0x02u64).unwrap(); + Encoder::new(&mut buf, Version::Draft17.into()).varint(1000u64).unwrap(); // Second: delta = 0x02 → abs type 0x04 (even), varint value - Encoder::new(&mut buf, Version::Draft17.into()) - .varint(VarInt::from(0x02u64)) - .unwrap(); - Encoder::new(&mut buf, Version::Draft17.into()) - .varint(VarInt::from(2000u64)) - .unwrap(); + Encoder::new(&mut buf, Version::Draft17.into()).varint(0x02u64).unwrap(); + Encoder::new(&mut buf, Version::Draft17.into()).varint(2000u64).unwrap(); // Third: delta = 0x07 → abs type 0x0B (odd), length-prefixed - Encoder::new(&mut buf, Version::Draft17.into()) - .varint(VarInt::from(0x07u64)) - .unwrap(); - Encoder::new(&mut buf, Version::Draft17.into()) - .varint(VarInt::from(2u64)) - .unwrap(); + Encoder::new(&mut buf, Version::Draft17.into()).varint(0x07u64).unwrap(); + Encoder::new(&mut buf, Version::Draft17.into()).varint(2u64).unwrap(); buf.extend_from_slice(&[0xAA, 0xBB]); let mut bytes = bytes::Bytes::from(buf); @@ -282,12 +262,8 @@ mod tests { Version::Draft22, ] { let mut buf = Vec::new(); - Encoder::new(&mut buf, version.into()) - .varint(VarInt::from(0x0eu64)) - .unwrap(); - Encoder::new(&mut buf, version.into()) - .varint(VarInt::from(256u64)) - .unwrap(); + Encoder::new(&mut buf, version.into()).varint(0x0eu64).unwrap(); + Encoder::new(&mut buf, version.into()).varint(256u64).unwrap(); assert!(matches!( crate::coding::decode_buf(&mut bytes::Bytes::from(buf), version, Properties::decode), Err(DecodeError::InvalidValue) @@ -300,12 +276,8 @@ mod tests { #[test] fn test_rejects_zero_group_order() { let mut buf = Vec::new(); - Encoder::new(&mut buf, Version::Draft18.into()) - .varint(VarInt::from(0x22u64)) - .unwrap(); - Encoder::new(&mut buf, Version::Draft18.into()) - .varint(VarInt::ZERO) - .unwrap(); + Encoder::new(&mut buf, Version::Draft18.into()).varint(0x22u64).unwrap(); + Encoder::new(&mut buf, Version::Draft18.into()).varint(0).unwrap(); let mut bytes = bytes::Bytes::from(buf); assert!(crate::coding::decode_buf(&mut bytes, Version::Draft18, Properties::decode).is_err()); @@ -316,12 +288,8 @@ mod tests { #[test] fn test_decodes_draft16_track_extensions() { let mut buf = Vec::new(); - Encoder::new(&mut buf, Version::Draft16.into()) - .varint(VarInt::from(0x22u64)) - .unwrap(); - Encoder::new(&mut buf, Version::Draft16.into()) - .varint(VarInt::from(2u64)) - .unwrap(); + Encoder::new(&mut buf, Version::Draft16.into()).varint(0x22u64).unwrap(); + Encoder::new(&mut buf, Version::Draft16.into()).varint(2u64).unwrap(); let mut bytes = bytes::Bytes::from(buf); let properties = crate::coding::decode_buf(&mut bytes, Version::Draft16, Properties::decode).unwrap(); diff --git a/rs/moq-net/src/ietf/publish.rs b/rs/moq-net/src/ietf/publish.rs index 54fd642457..45e9ecd0dc 100644 --- a/rs/moq-net/src/ietf/publish.rs +++ b/rs/moq-net/src/ietf/publish.rs @@ -108,7 +108,7 @@ use std::borrow::Cow; use crate::{ Path, - coding::{Decode, DecodeError, Decoder, Encode, EncodeError, Encoder, VarInt}, + coding::{Decode, DecodeError, Decoder, Encode, EncodeError, Encoder}, ietf::{ Filter, GroupOrder, Location, Parameters, Properties, RequestId, namespace::{decode_namespace, encode_namespace}, @@ -198,8 +198,8 @@ impl Message for PublishDone<'_> { } else { assert!(self.request_id.is_none(), "request_id must be None for draft17+"); } - w.varint(VarInt::from(self.status_code))?; - w.varint(VarInt::from(self.stream_count))?; + w.varint(self.status_code)?; + w.varint(self.stream_count)?; w.string(&self.reason_phrase)?; Ok(()) } @@ -210,8 +210,8 @@ impl Message for PublishDone<'_> { } else { None }; - let status_code = r.varint()?.into_inner(); - let stream_count = r.varint()?.into_inner(); + let status_code = r.varint()?; + let stream_count = r.varint()?; let reason_phrase = Cow::Owned(r.string()?); Ok(Self { @@ -243,11 +243,11 @@ impl Message for Publish<'_> { fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.request_id.encode(w, version)?; if version == Version::Draft17 { - w.varint(VarInt::ZERO)?; // required_request_id_delta = 0 + w.varint(0)?; // required_request_id_delta = 0 } encode_namespace(w, &self.track_namespace)?; w.string(&self.track_name)?; - w.varint(VarInt::from(self.track_alias))?; + w.varint(self.track_alias)?; match version { Version::Draft14 => { @@ -293,11 +293,11 @@ impl Message for Publish<'_> { fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { let request_id = RequestId::decode(r, version)?; if version == Version::Draft17 { - let _required_request_id_delta = r.varint()?.into_inner(); + let _required_request_id_delta = r.varint()?; } let track_namespace = decode_namespace(r)?; let track_name = Cow::Owned(r.string()?); - let track_alias = r.varint()?.into_inner(); + let track_alias = r.varint()?; match version { Version::Draft14 => { @@ -485,14 +485,14 @@ impl Message for PublishError<'_> { fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.request_id.encode(w, version)?; - w.varint(VarInt::from(self.error_code))?; + w.varint(self.error_code)?; w.string(&self.reason_phrase)?; Ok(()) } fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { let request_id = RequestId::decode(r, version)?; - let error_code = r.varint()?.into_inner(); + let error_code = r.varint()?; let reason_phrase = Cow::Owned(r.string()?); Ok(Self { request_id, @@ -517,7 +517,7 @@ mod tests { RequestId(1).encode(w, version)?; super::super::namespace::encode_namespace(w, &crate::Path::new("broadcast"))?; w.string("video")?; - w.varint(1u64.into())?; // track alias + w.varint(1u64)?; // track alias // SUBSCRIBER_PRIORITY then LOCATION_FILTER, delta encoded from 0. encode_params!(w, version, @@ -540,7 +540,7 @@ mod tests { RequestId(1).encode(w, version)?; super::super::namespace::encode_namespace(w, &crate::Path::new("broadcast"))?; w.string("video")?; - w.varint(1u64.into())?; + w.varint(1u64)?; encode_params!(w, version, 0x20 => 128u8); Properties::default().encode(w, version)?; diff --git a/rs/moq-net/src/ietf/publish_namespace.rs b/rs/moq-net/src/ietf/publish_namespace.rs index 5239f8bda5..3449d23cb3 100644 --- a/rs/moq-net/src/ietf/publish_namespace.rs +++ b/rs/moq-net/src/ietf/publish_namespace.rs @@ -33,7 +33,7 @@ impl PublishNamespace<'_> { pub fn decode_body(r: &mut Decoder<'_>, version: Version, negotiated: bool) -> Result { let request_id = RequestId::decode(r, version)?; if version == Version::Draft17 { - let _required_request_id_delta = r.varint()?.into_inner(); + let _required_request_id_delta = r.varint()?; } let track_namespace = decode_namespace(r)?; let cluster = decode_cluster_params(r, version, negotiated)?; @@ -52,7 +52,7 @@ impl Message for PublishNamespace<'_> { fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.request_id.encode(w, version)?; if version == Version::Draft17 { - w.varint(VarInt::ZERO)?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) + w.varint(0)?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) } encode_namespace(w, &self.track_namespace)?; encode_cluster_params(w, version, self.cluster.as_ref()) @@ -106,7 +106,7 @@ impl Message for PublishNamespaceUpdate { Version::Draft14 | Version::Draft15 | Version::Draft16 => return Err(EncodeError::Version), Version::Draft17 => { self.request_id.encode(w, version)?; - w.varint(VarInt::ZERO)?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) + w.varint(0)?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) } _ => self.request_id.encode(w, version)?, } @@ -122,7 +122,7 @@ impl Message for PublishNamespaceUpdate { Version::Draft14 | Version::Draft15 | Version::Draft16 => return Err(DecodeError::Version), Version::Draft17 => { let request_id = RequestId::decode(r, version)?; - let _required_request_id_delta = r.varint()?.into_inner(); + let _required_request_id_delta = r.varint()?; request_id } _ => RequestId::decode(r, version)?, @@ -214,14 +214,14 @@ impl Message for PublishNamespaceError<'_> { fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.request_id.encode(w, version)?; - w.varint(VarInt::from(self.error_code))?; + w.varint(self.error_code)?; w.string(&self.reason_phrase)?; Ok(()) } fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { let request_id = RequestId::decode(r, version)?; - let error_code = r.varint()?.into_inner(); + let error_code = r.varint()?; let reason_phrase = Cow::Owned(r.string()?); Ok(Self { @@ -306,7 +306,7 @@ impl Message for PublishNamespaceCancel<'_> { return Err(EncodeError::Version); } } - w.varint(VarInt::from(self.error_code))?; + w.varint(self.error_code)?; w.string(&self.reason_phrase)?; Ok(()) } @@ -325,7 +325,7 @@ impl Message for PublishNamespaceCancel<'_> { return Err(DecodeError::Version); } }; - let error_code = r.varint()?.into_inner(); + let error_code = r.varint()?; let reason_phrase = Cow::Owned(r.string()?); Ok(Self { track_namespace, diff --git a/rs/moq-net/src/ietf/publisher.rs b/rs/moq-net/src/ietf/publisher.rs index 661a371f3a..242aebed99 100644 --- a/rs/moq-net/src/ietf/publisher.rs +++ b/rs/moq-net/src/ietf/publisher.rs @@ -15,7 +15,7 @@ use web_transport_trait::poll::SendStream as _; use crate::{ AsPath, Error, Timescale, Timestamp, - coding::{Encoder, Stream, VarInt, Writer}, + coding::{Encoder, Stream, Writer}, ietf::{self, Control, EndLocation, FetchHeader, FetchType, Filter, GroupOrder, Location, RequestId}, track::Subscription, util::{MaybeBoxedExt, MaybeSendBox}, @@ -563,7 +563,7 @@ where .map(|fill| (fill_range(fill, msg.filter, edge.largest), cache, timescale)); // Send SubscribeOk on the stream - stream.writer.encode(&VarInt::from(ietf::SubscribeOk::ID)).await?; + stream.writer.varint(ietf::SubscribeOk::ID).await?; stream .writer .encode(&ietf::SubscribeOk { @@ -648,7 +648,7 @@ where Ok(()) => (ietf::PublishDoneStatus::TrackEnded, "track ended"), Err(_) => (ietf::PublishDoneStatus::InternalError, "internal error"), }; - let _ = stream.writer.encode(&VarInt::from(ietf::PublishDone::ID)).await; + let _ = stream.writer.varint(ietf::PublishDone::ID).await; let _ = stream .writer .encode(&ietf::PublishDone { @@ -702,7 +702,7 @@ where match self.version { Version::Draft14 => { - writer.encode(&VarInt::from(ietf::SubscribeError::ID)).await?; + writer.varint(ietf::SubscribeError::ID).await?; writer .encode(&ietf::SubscribeError { request_id, @@ -712,7 +712,7 @@ where .await?; } Version::Draft15 | Version::Draft16 => { - writer.encode(&VarInt::from(ietf::RequestError::ID)).await?; + writer.varint(ietf::RequestError::ID).await?; writer .encode(&ietf::RequestError { request_id: Some(request_id), @@ -723,7 +723,7 @@ where .await?; } _ => { - writer.encode(&VarInt::from(ietf::RequestError::ID)).await?; + writer.varint(ietf::RequestError::ID).await?; writer .encode(&ietf::RequestError { request_id: None, @@ -769,7 +769,7 @@ where stream.set_priority(priority); let res = async { - stream.encode(&VarInt::from(FetchHeader::TYPE)).await?; + stream.varint(FetchHeader::TYPE).await?; stream.encode(&FetchHeader { request_id }).await?; let FillServe::Group { sequence, skip, until } = fill else { @@ -872,9 +872,9 @@ where version, ) .await?; - stream.encode(&VarInt::from(frame.payload.len())).await?; + stream.varint(frame.payload.len() as u64).await?; if frame.payload.is_empty() && matches!(version, Version::Draft14 | Version::Draft15) { - stream.encode(&VarInt::ZERO).await?; + stream.varint(0).await?; } if !frame.payload.is_empty() { let mut payload = frame.payload; @@ -919,9 +919,9 @@ where .await?; index += 1; - stream.encode(&VarInt::from(frame.size)).await?; + stream.varint(frame.size).await?; if frame.size == 0 && matches!(version, Version::Draft14 | Version::Draft15) { - stream.encode(&VarInt::ZERO).await?; + stream.varint(0).await?; } loop { let chunk = { @@ -980,12 +980,12 @@ where if version == Version::Draft14 { let properties = properties.unwrap_or_default(); - stream.buffer(&VarInt::from(sequence))?; - stream.buffer(&VarInt::ZERO)?; - stream.buffer(&VarInt::from(object))?; + stream.buffer_varint(sequence)?; + stream.buffer_varint(0)?; + stream.buffer_varint(object)?; // Publisher priority, a raw byte. stream.buffer_raw(&[0]); - stream.buffer(&VarInt::from(properties.len()))?; + stream.buffer_varint(properties.len() as u64)?; stream.buffer_raw(&properties); std::future::poll_fn(|cx| stream.poll_flush(cx)).await?; return Ok(()); @@ -1188,7 +1188,7 @@ where // FETCH_OK on every draft, never REQUEST_OK: section 5.2 allows exactly one FETCH_OK or // REQUEST_ERROR in answer to a FETCH, and REQUEST_OK's own definition lists the other // requests it answers without ever naming this one. - stream.writer.encode(&VarInt::from(ietf::FetchOk::ID)).await?; + stream.writer.varint(ietf::FetchOk::ID).await?; stream .writer .encode(&ietf::FetchOk { @@ -1207,7 +1207,7 @@ where let uni = self.session.open_uni().await.map_err(Error::from_transport)?; let mut writer = Writer::new(uni, self.version); writer.set_priority(priority); - writer.encode(&VarInt::from(FetchHeader::TYPE)).await?; + writer.varint(FetchHeader::TYPE).await?; writer .encode(&FetchHeader { request_id: msg.request_id, @@ -1224,9 +1224,9 @@ where self.version, ) .await?; - writer.encode(&VarInt::from(frame.payload.len())).await?; + writer.varint(frame.payload.len() as u64).await?; if frame.payload.is_empty() && matches!(self.version, Version::Draft14 | Version::Draft15) { - writer.encode(&VarInt::ZERO).await?; + writer.varint(0).await?; } if !frame.payload.is_empty() { let mut payload = frame.payload; @@ -1246,7 +1246,7 @@ where async fn reject_track_status(&self, mut stream: Stream, request_id: RequestId) -> Result<(), Error> { let error_code = request::to_code(&Error::Unsupported, request::Kind::TrackStatus, self.version); if self.version == Version::Draft14 { - stream.writer.encode(&VarInt::from(0x0fu64)).await?; // TRACK_STATUS_ERROR has the SUBSCRIBE_ERROR body. + stream.writer.varint(0x0fu64).await?; // TRACK_STATUS_ERROR has the SUBSCRIBE_ERROR body. stream .writer .encode(&ietf::SubscribeError { @@ -1256,7 +1256,7 @@ where }) .await?; } else { - stream.writer.encode(&VarInt::from(ietf::RequestError::ID)).await?; + stream.writer.varint(ietf::RequestError::ID).await?; stream .writer .encode(&ietf::RequestError { @@ -1298,7 +1298,7 @@ where match self.version { Version::Draft14 => { - writer.encode(&VarInt::from(ietf::FetchError::ID)).await?; + writer.varint(ietf::FetchError::ID).await?; writer .encode(&ietf::FetchError { request_id, @@ -1308,7 +1308,7 @@ where .await?; } Version::Draft15 | Version::Draft16 => { - writer.encode(&VarInt::from(ietf::RequestError::ID)).await?; + writer.varint(ietf::RequestError::ID).await?; writer .encode(&ietf::RequestError { request_id: Some(request_id), @@ -1319,7 +1319,7 @@ where .await?; } _ => { - writer.encode(&VarInt::from(ietf::RequestError::ID)).await?; + writer.varint(ietf::RequestError::ID).await?; writer .encode(&ietf::RequestError { request_id: None, @@ -1417,7 +1417,7 @@ where match (advert.wanted(), held) { (true, _) => { tracing::debug!(broadcast = %absolute, "namespace"); - stream.writer.encode(&VarInt::from(ietf::Namespace::ID)).await?; + stream.writer.varint(ietf::Namespace::ID).await?; stream .writer .encode(&ietf::Namespace { @@ -1428,7 +1428,7 @@ where } (false, true) => { tracing::debug!(broadcast = %absolute, "namespace_done"); - stream.writer.encode(&VarInt::from(ietf::NamespaceDone::ID)).await?; + stream.writer.varint(ietf::NamespaceDone::ID).await?; stream .writer .encode(&ietf::NamespaceDone { @@ -1476,7 +1476,7 @@ where return Ok(Refused::No); }; - request.writer.encode(&VarInt::from(ietf::PublishNamespace::ID)).await?; + request.writer.varint(ietf::PublishNamespace::ID).await?; request .writer .encode(&ietf::PublishNamespace { @@ -1581,11 +1581,7 @@ where let request_id = self.control.next_request_id(&self.runtime).await?; let update = ietf::PublishNamespaceUpdate::between(request_id, &held, &next); - request - .stream - .writer - .encode(&VarInt::from(ietf::PublishNamespaceUpdate::ID)) - .await?; + request.stream.writer.varint(ietf::PublishNamespaceUpdate::ID).await?; request.stream.writer.encode(&update).await?; let absolute = self.origin.absolute(&request.path).to_owned(); @@ -1668,7 +1664,7 @@ where /// namespace stays outstanding and the retry re-offers it. async fn read_response(&self, request: &mut Stream) -> Result, Error> { let mut read = std::pin::pin!(async { - let type_id = request.reader.decode::().await?.into_inner(); + let type_id = request.reader.varint().await?; let body: ietf::Body = request.reader.decode().await?; Ok::<_, Error>((type_id, body)) }); @@ -1720,7 +1716,7 @@ where } } Target::Inline(stream) => { - stream.writer.encode(&VarInt::from(ietf::NamespaceDone::ID)).await?; + stream.writer.varint(ietf::NamespaceDone::ID).await?; stream .writer .encode(&ietf::NamespaceDone { @@ -1852,10 +1848,7 @@ where // Send OK response match self.version { Version::Draft14 => { - stream - .writer - .encode(&VarInt::from(ietf::SubscribeNamespaceOk::ID)) - .await?; + stream.writer.varint(ietf::SubscribeNamespaceOk::ID).await?; stream .writer .encode(&ietf::SubscribeNamespaceOk { @@ -1864,7 +1857,7 @@ where .await?; } Version::Draft15 | Version::Draft16 => { - stream.writer.encode(&VarInt::from(ietf::RequestOk::ID)).await?; + stream.writer.varint(ietf::RequestOk::ID).await?; stream .writer .encode(&ietf::RequestOk { @@ -1874,7 +1867,7 @@ where .await?; } _ => { - stream.writer.encode(&VarInt::from(ietf::RequestOk::ID)).await?; + stream.writer.varint(ietf::RequestOk::ID).await?; stream .writer .encode(&ietf::RequestOk { @@ -2142,9 +2135,9 @@ impl TrackServe { flags: ietf::GroupFlags::default(), })?; // Object ID delta 0, then an empty object whose status is END_OF_TRACK. - writer.buffer(&VarInt::ZERO)?; - writer.buffer(&VarInt::ZERO)?; - writer.encode(&VarInt::from(END_OF_TRACK)).await?; + writer.buffer_varint(0)?; + writer.buffer_varint(0)?; + writer.varint(END_OF_TRACK).await?; // PUBLISH_DONE follows once this closes, like every other data stream. writer.close().await } @@ -2510,7 +2503,7 @@ fn buffer_object_info( timescale: Option, version: Version, ) -> Result<(), Error> { - writer.buffer(&VarInt::from(delta))?; + writer.buffer_varint(delta)?; if let Some(timescale) = timescale.filter(|_| has_extensions) { // Per-object extension headers carry the frame's presentation timestamp. @@ -2521,14 +2514,14 @@ fn buffer_object_info( timescale, version, )?; - writer.buffer(&VarInt::from(ext.len()))?; + writer.buffer_varint(ext.len() as u64)?; writer.buffer_raw(&ext); } - writer.buffer(&VarInt::from(size))?; + writer.buffer_varint(size)?; if size == 0 { // Have to write the object status too: Normal (0). - writer.buffer(&VarInt::ZERO)?; + writer.buffer_varint(0)?; } Ok(()) } @@ -2928,7 +2921,7 @@ mod serve_tests { match version { Version::Draft14 => { - writer.encode(&VarInt::from(ietf::SubscribeError::ID)).await.unwrap(); + writer.varint(ietf::SubscribeError::ID).await.unwrap(); writer .encode(&ietf::SubscribeError { request_id: RequestId(REQUEST_ID), @@ -2939,7 +2932,7 @@ mod serve_tests { .unwrap(); } _ => { - writer.encode(&VarInt::from(ietf::RequestError::ID)).await.unwrap(); + writer.varint(ietf::RequestError::ID).await.unwrap(); writer .encode(&ietf::RequestError { request_id: match version { @@ -3002,7 +2995,7 @@ mod serve_tests { let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(log.clone()), version); if version == Version::Draft14 { - writer.encode(&VarInt::from(0x0fu64)).await.unwrap(); + writer.varint(0x0fu64).await.unwrap(); writer .encode(&ietf::SubscribeError { request_id: RequestId(REQUEST_ID), @@ -3012,7 +3005,7 @@ mod serve_tests { .await .unwrap(); } else { - writer.encode(&VarInt::from(ietf::RequestError::ID)).await.unwrap(); + writer.varint(ietf::RequestError::ID).await.unwrap(); writer .encode(&ietf::RequestError { request_id: matches!(version, Version::Draft15 | Version::Draft16) @@ -4046,7 +4039,7 @@ mod tests { crate::lite::test_transport::SinkSend::new(expected.clone()), Version::Draft17, ); - writer.encode(&VarInt::from(ietf::RequestOk::ID)).await.unwrap(); + writer.varint(ietf::RequestOk::ID).await.unwrap(); writer .encode(&ietf::RequestOk { request_id: None, @@ -4146,7 +4139,7 @@ mod tests { let log = crate::lite::test_transport::Log::default(); let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(log.clone()), version); - writer.encode(&VarInt::from(ietf::RequestError::ID)).await.unwrap(); + writer.varint(ietf::RequestError::ID).await.unwrap(); writer .encode(&ietf::RequestError { request_id: matches!(version, Version::Draft15 | Version::Draft16).then_some(RequestId(1)), @@ -4218,10 +4211,7 @@ mod tests { match version { Version::Draft14 => { - writer - .encode(&VarInt::from(ietf::PublishNamespaceOk::ID)) - .await - .unwrap(); + writer.varint(ietf::PublishNamespaceOk::ID).await.unwrap(); writer .encode(&ietf::PublishNamespaceOk { request_id: RequestId(1), @@ -4230,7 +4220,7 @@ mod tests { .unwrap(); } Version::Draft15 | Version::Draft16 => { - writer.encode(&VarInt::from(ietf::RequestOk::ID)).await.unwrap(); + writer.varint(ietf::RequestOk::ID).await.unwrap(); writer .encode(&ietf::RequestOk { request_id: Some(RequestId(1)), @@ -4241,7 +4231,7 @@ mod tests { } // Draft-17+ dropped the request id: the response rides the request's stream. _ => { - writer.encode(&VarInt::from(ietf::RequestOk::ID)).await.unwrap(); + writer.varint(ietf::RequestOk::ID).await.unwrap(); writer .encode(&ietf::RequestOk { request_id: None, @@ -4402,7 +4392,7 @@ mod tests { let ok = crate::lite::test_transport::Log::default(); let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(ok.clone()), VERSION); - writer.encode(&VarInt::from(ietf::RequestOk::ID)).await.unwrap(); + writer.varint(ietf::RequestOk::ID).await.unwrap(); writer .encode(&ietf::RequestOk { request_id: None, @@ -4743,10 +4733,7 @@ mod tests { async fn request_update(version: Version, msg: &ietf::PublishNamespaceUpdate) -> Vec { let log = crate::lite::test_transport::Log::default(); let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(log.clone()), version); - writer - .encode(&VarInt::from(ietf::PublishNamespaceUpdate::ID)) - .await - .unwrap(); + writer.varint(ietf::PublishNamespaceUpdate::ID).await.unwrap(); writer.encode(msg).await.unwrap(); log.writes.lock().unwrap().clone() } @@ -4835,7 +4822,7 @@ mod tests { let counted = crate::lite::test_transport::Log::default(); let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(counted.clone()), VERSION); - writer.encode(&VarInt::from(ietf::RequestOk::ID)).await.unwrap(); + writer.varint(ietf::RequestOk::ID).await.unwrap(); writer .encode(&ietf::RequestOk { request_id: None, diff --git a/rs/moq-net/src/ietf/request.rs b/rs/moq-net/src/ietf/request.rs index e1e0aaf198..248c585926 100644 --- a/rs/moq-net/src/ietf/request.rs +++ b/rs/moq-net/src/ietf/request.rs @@ -1,6 +1,6 @@ use std::borrow::Cow; -use crate::coding::{Decode, DecodeError, Decoder, Encode, EncodeError, Encoder, VarInt}; +use crate::coding::{Decode, DecodeError, Decoder, Encode, EncodeError, Encoder}; use super::Message; use super::active_count::ACTIVE_COUNT_PARAM; @@ -30,14 +30,14 @@ impl std::fmt::Display for RequestId { impl Encode for RequestId { fn encode(&self, w: &mut Encoder<'_>, _: Version) -> Result<(), EncodeError> { - w.varint(VarInt::from(self.0))?; + w.varint(self.0)?; Ok(()) } } impl Decode for RequestId { fn decode(r: &mut Decoder<'_>, _: Version) -> Result { - let request_id = r.varint()?.into_inner(); + let request_id = r.varint()?; Ok(Self(request_id)) } } @@ -144,9 +144,9 @@ impl Message for RequestError<'_> { } else { assert!(self.request_id.is_none(), "request_id must be None for draft17+"); } - w.varint(VarInt::from(self.error_code))?; + w.varint(self.error_code)?; if !matches!(version, Version::Draft14 | Version::Draft15) { - w.varint(VarInt::from(self.retry_interval))?; + w.varint(self.retry_interval)?; } w.string(&self.reason_phrase)?; Ok(()) @@ -158,10 +158,10 @@ impl Message for RequestError<'_> { } else { None }; - let error_code = r.varint()?.into_inner(); + let error_code = r.varint()?; let retry_interval = match version { Version::Draft14 | Version::Draft15 => 0, - _ => r.varint()?.into_inner(), + _ => r.varint()?, }; let reason_phrase = Cow::Owned(r.string()?); Ok(Self { diff --git a/rs/moq-net/src/ietf/session.rs b/rs/moq-net/src/ietf/session.rs index 3d8ace7137..dbe3bd5d73 100644 --- a/rs/moq-net/src/ietf/session.rs +++ b/rs/moq-net/src/ietf/session.rs @@ -1,7 +1,7 @@ use crate::origin; use crate::{ Error, Hop, SessionError, StreamError, - coding::{Decode, DecodeError, Encode, Reader, Stream, VarInt, Writer}, + coding::{Decode, DecodeError, Encode, Reader, Stream, Writer}, ietf::{self, FetchHeader, RequestId}, setup, util::{MaybeBoxedExt, MaybeSendBox, TaskSet, err_only}, @@ -474,7 +474,7 @@ pub async fn accept_setup( let recv = session.accept_uni().await.map_err(Error::from_transport)?; let mut reader: Reader = Reader::new(recv, outer_version); - if reader.decode_peek::().await?.into_inner() != setup::SETUP_V17 { + if reader.varint_peek().await? != setup::SETUP_V17 { // Not the SETUP (group data this early is unexpected). Reject and keep waiting. reader.abort(&Error::UnexpectedStream); continue; @@ -630,11 +630,11 @@ where let kind = match tasks .drive(|waiter| { let mut cx = waiter.context(); - reader.poll_decode_peek::(&mut cx) + reader.poll_varint_peek(&mut cx) }) .await { - Ok(kind) => kind.into_inner(), + Ok(kind) => kind, Err(err @ (Error::Cancel | Error::Stream(_) | Error::Remote(_) | Error::Decode(DecodeError::Short))) => { tracing::debug!(%err, "dropping uni stream that died before its type"); continue; @@ -715,7 +715,7 @@ async fn run_uni_group( where S: crate::transport::poll::Boxable, { - let kind = stream.decode_peek::().await?.into_inner(); + let kind = stream.varint_peek().await?; // SUBGROUP_HEADER type bytes match the form 0b0XX1XXXX (spec §11.4.2): // draft-14-17 use 0x10-0x1D and 0x30-0x3D, draft-18 adds 0x40 (FIRST_OBJECT) @@ -772,9 +772,7 @@ where let mut cx = waiter.context(); let id = match hdr_id { Some(id) => id, - None => { - *hdr_id.insert(std::task::ready!(stream.reader.poll_decode::(&mut cx))?.into_inner()) - } + None => *hdr_id.insert(std::task::ready!(stream.reader.poll_varint(&mut cx))?), }; let body = std::task::ready!(stream.reader.poll_decode::(&mut cx))?; std::task::Poll::Ready(Ok::<_, Error>((id, body))) @@ -821,8 +819,8 @@ async fn run_goaway( version: Version, goaway: crate::goaway::Protocol, ) -> Result<(), Error> { - let id = match reader.decode_maybe::().await? { - Some(id) => id.into_inner(), + let id = match reader.varint_maybe().await? { + Some(id) => id, None => return Ok(()), }; @@ -854,8 +852,8 @@ async fn run_goaway( // control stream enforces, so close over it here too rather than logging; // anything else is merely unexpected and discarded. loop { - let id = match reader.decode_maybe::().await? { - Some(id) => id.into_inner(), + let id = match reader.varint_maybe().await? { + Some(id) => id, None => return Ok(()), }; let body: ietf::Body = reader.decode().await?; @@ -892,7 +890,7 @@ mod tests { let log = crate::lite::test_transport::Log::default(); let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(log.clone()), version); - writer.encode(&VarInt::from(ietf::RequestOk::ID)).await.unwrap(); + writer.varint(ietf::RequestOk::ID).await.unwrap(); writer .encode(&ietf::RequestOk { request_id: None, @@ -900,7 +898,7 @@ mod tests { }) .await .unwrap(); - writer.encode(&VarInt::from(ietf::Namespace::ID)).await.unwrap(); + writer.varint(ietf::Namespace::ID).await.unwrap(); writer .encode(&ietf::Namespace { suffix: crate::Path::new("cam"), @@ -1278,7 +1276,7 @@ mod tests { let log = crate::lite::test_transport::Log::default(); let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(log.clone()), version); - writer.encode(&VarInt::from(ietf::PublishNamespace::ID)).await.unwrap(); + writer.varint(ietf::PublishNamespace::ID).await.unwrap(); writer .encode(&ietf::PublishNamespace { request_id: RequestId(1), @@ -1289,10 +1287,7 @@ mod tests { .unwrap(); for _ in 0..2 { - writer - .encode(&VarInt::from(ietf::PublishNamespaceDone::ID)) - .await - .unwrap(); + writer.varint(ietf::PublishNamespaceDone::ID).await.unwrap(); writer .encode(&ietf::PublishNamespaceDone { track_namespace: crate::Path::new("room/host"), @@ -1315,10 +1310,7 @@ mod tests { Version::Draft14, ); - writer - .encode(&VarInt::from(ietf::PublishNamespaceOk::ID)) - .await - .unwrap(); + writer.varint(ietf::PublishNamespaceOk::ID).await.unwrap(); writer.encode(&ietf::PublishNamespaceOk { request_id }).await.unwrap(); let writes = log.writes.lock().unwrap(); diff --git a/rs/moq-net/src/ietf/subscribe.rs b/rs/moq-net/src/ietf/subscribe.rs index 2241d82b72..9f7cf3b97e 100644 --- a/rs/moq-net/src/ietf/subscribe.rs +++ b/rs/moq-net/src/ietf/subscribe.rs @@ -61,7 +61,7 @@ impl Message for Subscribe<'_> { fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { let request_id = RequestId::decode(r, version)?; if version == Version::Draft17 { - let _required_request_id_delta = r.varint()?.into_inner(); + let _required_request_id_delta = r.varint()?; } let track_namespace = decode_namespace(r)?; let track_name = Cow::Owned(r.string()?); @@ -152,7 +152,7 @@ impl Message for Subscribe<'_> { fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.request_id.encode(w, version)?; if version == Version::Draft17 { - w.varint(VarInt::ZERO)?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) + w.varint(0)?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) } encode_namespace(w, &self.track_namespace)?; w.string(&self.track_name)?; @@ -218,11 +218,11 @@ impl Message for SubscribeOk { } else { assert!(self.request_id.is_none(), "request_id must be None for draft17+"); } - w.varint(VarInt::from(self.track_alias))?; + w.varint(self.track_alias)?; match version { Version::Draft14 => { - w.varint(VarInt::ZERO)?; // expires = 0 + w.varint(0)?; // expires = 0 self.properties .group_order .unwrap_or(GroupOrder::Ascending) @@ -263,7 +263,7 @@ impl Message for SubscribeOk { } else { None }; - let track_alias = r.varint()?.into_inner(); + let track_alias = r.varint()?; let mut properties = Properties::default(); let mut largest = None; @@ -271,7 +271,7 @@ impl Message for SubscribeOk { Version::Draft14 => { // EXPIRES is when the publisher expects to end the subscription. That end // arrives as PUBLISH_DONE regardless, so there is nothing to act on. - let _expires = r.varint()?.into_inner(); + let _expires = r.varint()?; properties.group_order = Some(GroupOrder::decode(r, version)?.any_to_descending()); @@ -333,14 +333,14 @@ impl Message for SubscribeError<'_> { fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.request_id.encode(w, version)?; - w.varint(VarInt::from(self.error_code))?; + w.varint(self.error_code)?; w.string(&self.reason_phrase)?; Ok(()) } fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { let request_id = RequestId::decode(r, version)?; - let error_code = r.varint()?.into_inner(); + let error_code = r.varint()?; let reason_phrase = Cow::Owned(r.string()?); Ok(Self { @@ -393,7 +393,7 @@ impl Message for SubscribeUpdate { .expect("subscription_request_id required for draft14") .encode(w, version)?; self.start_location.encode(w, version)?; - w.varint(VarInt::from(self.end_group))?; + w.varint(self.end_group)?; w.u8(self.subscriber_priority); w.bool(self.forward); w.u8(0); // no parameters @@ -417,7 +417,7 @@ impl Message for SubscribeUpdate { // REQUEST_UPDATE self.request_id.encode(w, version)?; if matches!(version, Version::Draft17) { - w.varint(VarInt::ZERO)?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) + w.varint(0)?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) } encode_params!(w, version, 0x10 => self.forward, @@ -436,7 +436,7 @@ impl Message for SubscribeUpdate { let request_id = RequestId::decode(r, version)?; let subscription_request_id = Some(RequestId::decode(r, version)?); let start_location = Location::decode(r, version)?; - let end_group = r.varint()?.into_inner(); + let end_group = r.varint()?; let subscriber_priority = r.u8()?; let forward = r.bool()?; let _parameters = Parameters::decode(r, version)?; @@ -477,7 +477,7 @@ impl Message for SubscribeUpdate { // REQUEST_UPDATE let request_id = RequestId::decode(r, version)?; if matches!(version, Version::Draft17) { - let _required_request_id_delta = r.varint()?.into_inner(); + let _required_request_id_delta = r.varint()?; } decode_params!(r, version, 0x02 => _object_delivery_timeout: Option, @@ -577,7 +577,7 @@ mod tests { let w = &mut Encoder::new(&mut buf, version.into()); RequestId(1).encode(w, version)?; if version == Version::Draft17 { - w.varint(VarInt::ZERO)?; // required_request_id_delta + w.varint(0)?; // required_request_id_delta } encode_namespace(w, &Path::new("test"))?; w.string("video")?; @@ -634,8 +634,8 @@ mod tests { r.string().unwrap(); // draft-14/15 write absolute keys, draft-16+ deltas, but the first is absolute either way. - let count = r.varint().unwrap().into_inner(); - (count > 0).then(|| r.varint().unwrap().into_inner()) + let count = r.varint().unwrap(); + (count > 0).then(|| r.varint().unwrap()) } /// We never ask a peer to hold a subscription open, so the parameter stays off our wire. @@ -1280,9 +1280,7 @@ mod cache_duration_tests { Version::Draft16 => vec![0, 0, 0, 4], _ => vec![0, 0, 4], }; - Encoder::new(&mut payload, version.into()) - .varint(VarInt::from(age)) - .unwrap(); + Encoder::new(&mut payload, version.into()).varint(age).unwrap(); let got = crate::coding::decode_buf(&mut payload.as_slice(), version, SubscribeOk::decode_msg).unwrap(); assert_eq!( got.properties.max_cache_duration, diff --git a/rs/moq-net/src/ietf/subscribe_namespace.rs b/rs/moq-net/src/ietf/subscribe_namespace.rs index 35ec81655b..7f8909389a 100644 --- a/rs/moq-net/src/ietf/subscribe_namespace.rs +++ b/rs/moq-net/src/ietf/subscribe_namespace.rs @@ -111,11 +111,11 @@ impl Message for SubscribeNamespaceLegacy<'_> { } self.request_id.encode(w, version)?; if version == Version::Draft17 { - w.varint(VarInt::ZERO)?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) + w.varint(0)?; // required_request_id_delta = 0 (draft-17 only, removed in draft-18 per #1615) } encode_namespace(w, &self.namespace)?; if matches!(version, Version::Draft16 | Version::Draft17) { - w.varint(VarInt::from(self.subscribe_options))?; + w.varint(self.subscribe_options)?; } encode_params!(w, version, HIDDEN_PARAM => hidden_param(self.hidden)); Ok(()) @@ -127,11 +127,11 @@ impl Message for SubscribeNamespaceLegacy<'_> { } let request_id = RequestId::decode(r, version)?; if version == Version::Draft17 { - let _required_request_id_delta = r.varint()?.into_inner(); + let _required_request_id_delta = r.varint()?; } let namespace = decode_namespace(r)?; let subscribe_options = match version { - Version::Draft16 | Version::Draft17 => r.varint()?.into_inner(), + Version::Draft16 | Version::Draft17 => r.varint()?, _ => 0x01, }; @@ -179,14 +179,14 @@ impl Message for SubscribeNamespaceError<'_> { fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.request_id.encode(w, version)?; - w.varint(VarInt::from(self.error_code))?; + w.varint(self.error_code)?; w.string(&self.reason_phrase)?; Ok(()) } fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { let request_id = RequestId::decode(r, version)?; - let error_code = r.varint()?.into_inner(); + let error_code = r.varint()?; let reason_phrase = Cow::Owned(r.string()?); Ok(Self { @@ -398,7 +398,7 @@ mod tests { let mut buf = Vec::new(); encode_namespace(&mut Encoder::new(&mut buf, version.into()), &Path::new("a")).unwrap(); // Number of Parameters = 0. - Encoder::new(&mut buf, version.into()).varint(VarInt::ZERO).unwrap(); + Encoder::new(&mut buf, version.into()).varint(0).unwrap(); let mut bytes = bytes::Bytes::from(buf); assert!(matches!( diff --git a/rs/moq-net/src/ietf/subscriber.rs b/rs/moq-net/src/ietf/subscriber.rs index 5584423b41..27642b0947 100644 --- a/rs/moq-net/src/ietf/subscriber.rs +++ b/rs/moq-net/src/ietf/subscriber.rs @@ -7,7 +7,7 @@ use std::{ use crate::{ Error, Path, PathOwned, SessionError, Timescale, broadcast, - coding::{Decode, DecodeError, Decoder, Reader, Stream, VarInt}, + coding::{Decode, DecodeError, Decoder, Reader, Stream}, frame, group, ietf::{self, Control, FetchType, Filter, GroupOrder, RequestId}, origin, track, @@ -745,10 +745,7 @@ where subscribe_options: 0x01, // NAMESPACE only hidden, }; - stream - .writer - .encode(&VarInt::from(ietf::SubscribeNamespaceLegacy::ID)) - .await?; + stream.writer.varint(ietf::SubscribeNamespaceLegacy::ID).await?; stream.writer.encode(&msg).await?; } _ => { @@ -757,10 +754,7 @@ where namespace: prefix.clone(), hidden, }; - stream - .writer - .encode(&VarInt::from(ietf::SubscribeNamespace::ID)) - .await?; + stream.writer.varint(ietf::SubscribeNamespace::ID).await?; stream.writer.encode(&msg).await?; } } @@ -768,7 +762,7 @@ where tracing::debug!(%prefix, "subscribe_namespace sent"); // Read response - let type_id = stream.reader.decode::().await?.into_inner(); + let type_id = stream.reader.varint().await?; let body: ietf::Body = stream.reader.decode().await?; let mut data = body.decoder(self.version); @@ -845,7 +839,7 @@ where ) -> Result<(), Error> { loop { let next = { - let mut decode = std::pin::pin!(stream.reader.decode_maybe::()); + let mut decode = std::pin::pin!(stream.reader.varint_maybe()); kio::wait(|waiter| { // Land before decoding past the boundary, so no live update enters the // origin ahead of the marker. @@ -859,7 +853,7 @@ where .await }; let type_id = match next? { - Some(id) => id.into_inner(), + Some(id) => id, None => break, // Stream closed }; if let Some((_, Landing::Quiet(quiet))) = landing { @@ -1121,7 +1115,7 @@ where attached: &mut bool, ) -> Result<(), Error> { loop { - let type_id = match stream.reader.decode_maybe::().await?.map(VarInt::into_inner) { + let type_id = match stream.reader.varint_maybe().await? { Some(id) => id, None => return Ok(()), }; @@ -1266,14 +1260,11 @@ where async fn write_ok(&self, stream: &mut Stream, request_id: RequestId) -> Result<(), Error> { match self.version { Version::Draft14 => { - stream - .writer - .encode(&VarInt::from(ietf::PublishNamespaceOk::ID)) - .await?; + stream.writer.varint(ietf::PublishNamespaceOk::ID).await?; stream.writer.encode(&ietf::PublishNamespaceOk { request_id }).await?; } Version::Draft15 | Version::Draft16 => { - stream.writer.encode(&VarInt::from(ietf::RequestOk::ID)).await?; + stream.writer.varint(ietf::RequestOk::ID).await?; stream .writer .encode(&ietf::RequestOk { @@ -1283,7 +1274,7 @@ where .await?; } _ => { - stream.writer.encode(&VarInt::from(ietf::RequestOk::ID)).await?; + stream.writer.varint(ietf::RequestOk::ID).await?; stream .writer .encode(&ietf::RequestOk { @@ -1308,10 +1299,7 @@ where match self.version { Version::Draft14 => { - stream - .writer - .encode(&VarInt::from(ietf::PublishNamespaceError::ID)) - .await?; + stream.writer.varint(ietf::PublishNamespaceError::ID).await?; stream .writer .encode(&ietf::PublishNamespaceError { @@ -1322,7 +1310,7 @@ where .await?; } Version::Draft15 | Version::Draft16 => { - stream.writer.encode(&VarInt::from(ietf::RequestError::ID)).await?; + stream.writer.varint(ietf::RequestError::ID).await?; stream .writer .encode(&ietf::RequestError { @@ -1334,7 +1322,7 @@ where .await?; } _ => { - stream.writer.encode(&VarInt::from(ietf::RequestError::ID)).await?; + stream.writer.varint(ietf::RequestError::ID).await?; stream .writer .encode(&ietf::RequestError { @@ -1361,7 +1349,7 @@ where match self.version { Version::Draft14 => { - stream.writer.encode(&VarInt::from(ietf::PublishError::ID)).await?; + stream.writer.varint(ietf::PublishError::ID).await?; stream .writer .encode(&ietf::PublishError { @@ -1372,7 +1360,7 @@ where .await?; } Version::Draft15 | Version::Draft16 => { - stream.writer.encode(&VarInt::from(ietf::RequestError::ID)).await?; + stream.writer.varint(ietf::RequestError::ID).await?; stream .writer .encode(&ietf::RequestError { @@ -1384,7 +1372,7 @@ where .await?; } _ => { - stream.writer.encode(&VarInt::from(ietf::RequestError::ID)).await?; + stream.writer.varint(ietf::RequestError::ID).await?; stream .writer .encode(&ietf::RequestError { @@ -1983,7 +1971,7 @@ where /// The publisher must send it before its FIN (draft-19 section 3.3.2), so a FIN /// without one is a failed request, not a clean end. async fn read_publish_done(reader: &mut Reader, version: Version) -> Result { - match reader.decode_maybe::().await?.map(VarInt::into_inner) { + match reader.varint_maybe().await? { Some(ietf::PublishDone::ID) => {} Some(_) => return Err(Error::UnexpectedMessage), None => return Err(Error::ProtocolViolation), @@ -2042,7 +2030,7 @@ where writer: &mut crate::coding::Writer, request_id: RequestId, ) -> Result<(), Error> { - writer.encode(&VarInt::from(ietf::Unsubscribe::ID)).await?; + writer.varint(ietf::Unsubscribe::ID).await?; writer.encode(&ietf::Unsubscribe { request_id }).await?; Ok(()) } @@ -2058,7 +2046,7 @@ where // Read the aggregate now: a subscriber can join while the request ID and stream // were awaited, and nothing updates the priority after SUBSCRIBE. let priority = request.subscription().map(|s| s.priority).unwrap_or(0); - stream.writer.encode(&VarInt::from(ietf::Subscribe::ID)).await?; + stream.writer.varint(ietf::Subscribe::ID).await?; stream .writer .encode(&ietf::Subscribe { @@ -2128,7 +2116,7 @@ where }; if let Err(err) = async { - stream.writer.encode(&VarInt::from(ietf::Fetch::ID)).await?; + stream.writer.varint(ietf::Fetch::ID).await?; stream .writer .encode(&ietf::Fetch { @@ -2174,7 +2162,7 @@ where /// `true` when the publisher answered FETCH_OK. A FETCH_ERROR / REQUEST_ERROR is a /// refusal, not a session error: the live subscription continues. async fn read_fetch_response(&self, stream: &mut Stream) -> Result { - let type_id = stream.reader.decode::().await?.into_inner(); + let type_id = stream.reader.varint().await?; let body: ietf::Body = stream.reader.decode().await?; let mut data = body.decoder(self.version); @@ -2197,7 +2185,7 @@ where async fn read_subscribe_response(&self, stream: &mut Stream) -> Result, Error> { // Read type_id + size + body from the stream - let type_id = stream.reader.decode::().await?.into_inner(); + let type_id = stream.reader.varint().await?; let body: ietf::Body = stream.reader.decode().await?; let mut data = body.decoder(self.version); @@ -2422,12 +2410,12 @@ struct PeekFirst(FirstObject); impl Decode for PeekFirst { fn decode(buf: &mut Decoder<'_>, _: Version) -> Result { - let id = buf.varint()?.into_inner(); + let id = buf.varint()?; if EXTENSIONS { buf.bytes()?; } - let size = buf.varint()?.into_inner(); - let end_of_track = size == 0 && buf.varint()?.into_inner() == END_OF_TRACK; + let size = buf.varint()?; + let end_of_track = size == 0 && buf.varint()? == END_OF_TRACK; Ok(Self(FirstObject { id, end_of_track })) } } @@ -2609,7 +2597,7 @@ where /// signal, and arrives here as a read error, which drops the head and the join with it. pub async fn recv_fill(&mut self, stream: &mut Reader) -> Result<(), Error> { // The dispatcher peeked the stream type to get here. - let _ = stream.decode::().await?.into_inner(); + let _ = stream.varint().await?; let header: ietf::FetchHeader = stream.decode().await?; let (subscribe_id, fill, joining, largest, _counted) = { @@ -2841,9 +2829,9 @@ where // A fetch object has no status field from draft-16 on; a zero length is simply // an empty object. Draft-14 and 15 still encode Normal (0) after a zero length. - let size = stream.decode::().await?.into_inner(); + let size = stream.varint().await?; if size == 0 && matches!(self.version, Version::Draft14 | Version::Draft15) { - let status = stream.decode::().await?.into_inner(); + let status = stream.varint().await?; if status != 0 { return Err(Error::Unsupported); } @@ -2881,13 +2869,14 @@ async fn decode_fetch_object( version: Version, ) -> Result, Error> { if version == Version::Draft14 { - let Some(group) = stream.decode_maybe::().await?.map(VarInt::into_inner) else { + let Some(group) = stream.varint_maybe().await? else { return Ok(None); }; - let subgroup = stream.decode::().await?.into_inner(); - let object = stream.decode::().await?.into_inner(); + let subgroup = stream.varint().await?; + let object = stream.varint().await?; let _priority = stream.read_exact(1).await?; - let size = usize::try_from(stream.decode::().await?)?; + let size = usize::try_from(stream.varint().await?) + .map_err(|_| Error::BoundsExceeded(crate::coding::BoundsExceeded))?; let properties = stream.read_exact(size).await?.to_vec(); return Ok(Some(FetchedObject { group: Some(group), @@ -3005,8 +2994,7 @@ impl GroupIngest { loop { match &mut self.phase { IngestPhase::Delta => { - let Some(id_delta) = ready!(reader.poll_decode_maybe::(&mut cx))?.map(VarInt::into_inner) - else { + let Some(id_delta) = ready!(reader.poll_varint_maybe(&mut cx))? else { return Poll::Ready(Ok(Ended::Group)); }; self.prior_object = Some(next_object_id(self.prior_object, id_delta, self.start)?); @@ -3016,7 +3004,7 @@ impl GroupIngest { }; } IngestPhase::ExtSize => { - let size = ready!(reader.poll_decode::(&mut cx))?; + let size = ready!(reader.poll_varint(&mut cx))?; let size = usize::try_from(size).map_err(|_| Error::BoundsExceeded(crate::coding::BoundsExceeded))?; self.phase = IngestPhase::ExtBytes { size }; @@ -3038,7 +3026,7 @@ impl GroupIngest { self.phase = IngestPhase::Size { timestamp }; } IngestPhase::Size { timestamp } => { - let size = ready!(reader.poll_decode::(&mut cx))?.into_inner(); + let size = ready!(reader.poll_varint(&mut cx))?; if size == 0 { self.phase = IngestPhase::Status { timestamp: *timestamp }; continue; @@ -3050,7 +3038,7 @@ impl GroupIngest { self.phase = IngestPhase::Payload { frame }; } IngestPhase::Status { timestamp } => { - let status = ready!(reader.poll_decode::(&mut cx))?.into_inner(); + let status = ready!(reader.poll_varint(&mut cx))?; if status == 0 { let timestamp = timestamp.unwrap_or_else(|| crate::Timestamp::from(self.runtime.now())); let frame = group.create_frame_owned(frame::Info { size: 0, timestamp })?; @@ -3117,7 +3105,7 @@ mod tests { let mut responses = Vec::new(); if clean { crate::coding::Encoder::new(&mut responses, Version::Draft19.into()) - .varint(VarInt::from(ietf::PublishDone::ID)) + .varint(ietf::PublishDone::ID) .unwrap(); ietf::PublishDone { request_id: None, @@ -3346,7 +3334,7 @@ mod tests { let log = crate::lite::test_transport::Log::default(); let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(log.clone()), version); - writer.encode(&VarInt::from(ietf::RequestOk::ID)).await.unwrap(); + writer.varint(ietf::RequestOk::ID).await.unwrap(); writer .encode(&ietf::RequestOk { request_id: None, @@ -3354,7 +3342,7 @@ mod tests { }) .await .unwrap(); - writer.encode(&VarInt::from(ietf::Namespace::ID)).await.unwrap(); + writer.varint(ietf::Namespace::ID).await.unwrap(); writer .encode(&ietf::Namespace { suffix: crate::Path::new(suffix), @@ -3438,7 +3426,7 @@ mod tests { let log = crate::lite::test_transport::Log::default(); let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(log.clone()), VERSION); - writer.encode(&VarInt::from(ietf::RequestOk::ID)).await.unwrap(); + writer.varint(ietf::RequestOk::ID).await.unwrap(); writer .encode(&ietf::RequestOk { request_id: None, @@ -3706,7 +3694,7 @@ mod tests { let log = crate::lite::test_transport::Log::default(); let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(log.clone()), VERSION); - writer.encode(&VarInt::from(ietf::RequestError::ID)).await.unwrap(); + writer.varint(ietf::RequestError::ID).await.unwrap(); writer .encode(&ietf::RequestError { request_id: Some(RequestId(1)), @@ -3905,7 +3893,7 @@ mod tests { let mut types = Vec::new(); while !buf.is_empty() { - let Ok(type_id) = buf.varint().map(VarInt::into_inner) else { + let Ok(type_id) = buf.varint() else { break; }; let Ok(size) = buf.u16() else { @@ -4042,7 +4030,7 @@ mod tests { let log = crate::lite::test_transport::Log::default(); let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(log.clone()), version); - writer.encode(&VarInt::from(ietf::SubscribeOk::ID)).await.unwrap(); + writer.varint(ietf::SubscribeOk::ID).await.unwrap(); writer .encode(&ietf::SubscribeOk { request_id: match version { @@ -4145,7 +4133,7 @@ mod tests { let log = crate::lite::test_transport::Log::default(); let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(log.clone()), version); - writer.encode(&VarInt::from(ietf::SubscribeOk::ID)).await.unwrap(); + writer.varint(ietf::SubscribeOk::ID).await.unwrap(); writer .encode(&ietf::SubscribeOk { request_id: None, @@ -4588,7 +4576,7 @@ mod tests { let log = crate::lite::test_transport::Log::default(); let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(log.clone()), VERSION); - writer.encode(&VarInt::from(ietf::RequestOk::ID)).await.unwrap(); + writer.varint(ietf::RequestOk::ID).await.unwrap(); writer .encode(&ietf::RequestOk { request_id: None, @@ -4597,7 +4585,7 @@ mod tests { .await .unwrap(); for cost in [4, 0] { - writer.encode(&VarInt::from(ietf::Namespace::ID)).await.unwrap(); + writer.varint(ietf::Namespace::ID).await.unwrap(); writer .encode(&ietf::Namespace { suffix: crate::Path::new("x.hang"), @@ -4749,7 +4737,7 @@ mod tests { // error while the advertisement is still live. let log = crate::lite::test_transport::Log::default(); let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(log.clone()), VERSION); - writer.encode(&VarInt::from(ietf::NamespaceDone::ID)).await.unwrap(); + writer.varint(ietf::NamespaceDone::ID).await.unwrap(); let script = log.writes.lock().unwrap().clone(); let session = crate::lite::test_transport::ScriptedSession::eof(script); @@ -5049,10 +5037,7 @@ mod tests { let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(log.clone()), VERSION); for (i, advert) in updates.iter().enumerate() { - writer - .encode(&VarInt::from(ietf::PublishNamespaceUpdate::ID)) - .await - .unwrap(); + writer.varint(ietf::PublishNamespaceUpdate::ID).await.unwrap(); writer .encode(&ietf::PublishNamespaceUpdate { // Each update consumes a request id of the peer's parity. @@ -5282,10 +5267,7 @@ mod tests { let log = crate::lite::test_transport::Log::default(); let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(log.clone()), VERSION); - writer - .encode(&VarInt::from(ietf::PublishNamespaceUpdate::ID)) - .await - .unwrap(); + writer.varint(ietf::PublishNamespaceUpdate::ID).await.unwrap(); writer .encode(&ietf::PublishNamespaceUpdate { request_id: RequestId(3), @@ -5439,7 +5421,7 @@ mod tests { let log = crate::lite::test_transport::Log::default(); let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(log.clone()), VERSION); - writer.encode(&VarInt::from(ietf::PublishNamespace::ID)).await.unwrap(); + writer.varint(ietf::PublishNamespace::ID).await.unwrap(); writer .encode(&ietf::PublishNamespace { request_id: RequestId(1), @@ -5571,7 +5553,7 @@ mod tests { match version { Version::Draft14 => { - writer.encode(&VarInt::from(ietf::PublishError::ID)).await.unwrap(); + writer.varint(ietf::PublishError::ID).await.unwrap(); writer .encode(&ietf::PublishError { request_id: RequestId(1), @@ -5582,7 +5564,7 @@ mod tests { .unwrap(); } _ => { - writer.encode(&VarInt::from(ietf::RequestError::ID)).await.unwrap(); + writer.varint(ietf::RequestError::ID).await.unwrap(); writer .encode(&ietf::RequestError { request_id: None, @@ -5952,7 +5934,7 @@ mod stitch_tests { fn fill_stream_for>(request_id: RequestId, groups: &[(u64, &[B])], timed: bool) -> Vec { let mut buf = Vec::new(); crate::coding::Encoder::new(&mut buf, VERSION.into()) - .varint(VarInt::from(ietf::FetchHeader::TYPE)) + .varint(ietf::FetchHeader::TYPE) .unwrap(); ietf::FetchHeader { request_id } .encode(&mut crate::coding::Encoder::new(&mut buf, VERSION.into()), VERSION) @@ -5990,7 +5972,7 @@ mod stitch_tests { .unwrap(); crate::coding::Encoder::new(&mut buf, VERSION.into()) - .varint(VarInt::from(payload.len())) + .varint(payload.len() as u64) .unwrap(); buf.put_slice(payload); object_index += 1; @@ -6026,10 +6008,10 @@ mod stitch_tests { _ => 0, }; crate::coding::Encoder::new(&mut buf, VERSION.into()) - .varint(VarInt::from(delta)) + .varint(delta) .unwrap(); crate::coding::Encoder::new(&mut buf, VERSION.into()) - .varint(VarInt::from(payload.len())) + .varint(payload.len() as u64) .unwrap(); buf.put_slice(payload); } @@ -6181,7 +6163,7 @@ mod stitch_tests { fn end_of_track(mut stream: Vec) -> Vec { for value in [0u64, 0, END_OF_TRACK] { crate::coding::Encoder::new(&mut stream, VERSION.into()) - .varint(VarInt::from(value)) + .varint(value) .unwrap(); } stream @@ -6647,7 +6629,7 @@ mod joining_fetch_tests { fn message_bytes(id: u64, msg: &M, version: Version) -> Vec { let mut buf = Vec::new(); crate::coding::Encoder::new(&mut buf, version.into()) - .varint(VarInt::from(id)) + .varint(id) .unwrap(); msg.encode(&mut crate::coding::Encoder::new(&mut buf, version.into()), version) .unwrap(); @@ -6733,7 +6715,7 @@ mod joining_fetch_tests { let Ok(body) = ietf::Body::decode(&mut buf, version) else { break; }; - messages.push((type_id.into_inner(), body.0)); + messages.push((type_id, body.0)); } messages } diff --git a/rs/moq-net/src/ietf/token.rs b/rs/moq-net/src/ietf/token.rs index 2dad2a6fbb..e5e1c3558d 100644 --- a/rs/moq-net/src/ietf/token.rs +++ b/rs/moq-net/src/ietf/token.rs @@ -37,8 +37,8 @@ pub fn from_setup(params: &Parameters, version: Version) -> Result pub fn into_setup(params: &mut Parameters, token: &Token, version: Version) -> Result<(), EncodeError> { let mut value = Vec::new(); let mut w = Encoder::new(&mut value, version.into()); - w.varint(USE_VALUE.into())?; - w.varint(token.kind.into())?; + w.varint(USE_VALUE)?; + w.varint(token.kind)?; w.slice(&token.value); params.set_bytes(ParameterBytes::AuthorizationToken, value); Ok(()) @@ -50,7 +50,7 @@ fn decode(buf: &[u8], version: Version) -> Result { let malformed = |_| Error::Session(SessionError::KeyValueFormatting); let mut r = Decoder::new(buf, version.into()); - match r.varint().map_err(malformed)?.into_inner() { + match r.varint().map_err(malformed)? { USE_VALUE => {} // With no cache, section 9.1.4 treats a registration as a value; the alias is unused. REGISTER => { @@ -61,7 +61,7 @@ fn decode(buf: &[u8], version: Version) -> Result { _ => return Err(Error::Session(SessionError::KeyValueFormatting)), } - let kind = r.varint().map_err(malformed)?.into_inner(); + let kind = r.varint().map_err(malformed)?; Ok(Token { kind, value: r.rest().to_vec(), @@ -105,7 +105,7 @@ mod tests { let mut raw = Vec::new(); let mut w = Encoder::new(&mut raw, version.into()); for field in fields { - w.varint((*field).into()).unwrap(); + w.varint(*field).unwrap(); } raw.extend_from_slice(value); let mut params = Parameters::default(); @@ -177,11 +177,9 @@ mod tests { fn two_tokens_are_refused() { for version in VERSIONS { let mut value = Vec::new(); + Encoder::new(&mut value, version.into()).varint(USE_VALUE).unwrap(); Encoder::new(&mut value, version.into()) - .varint(crate::coding::VarInt::from(USE_VALUE)) - .unwrap(); - Encoder::new(&mut value, version.into()) - .varint(crate::coding::VarInt::from(Token::OUT_OF_BAND)) + .varint(Token::OUT_OF_BAND) .unwrap(); let key = u64::from(ParameterBytes::AuthorizationToken); @@ -193,14 +191,10 @@ mod tests { }; let mut raw = Vec::new(); if let Some(count) = count { - Encoder::new(&mut raw, version.into()) - .varint(crate::coding::VarInt::from(count)) - .unwrap(); + Encoder::new(&mut raw, version.into()).varint(count).unwrap(); } for key in keys { - Encoder::new(&mut raw, version.into()) - .varint(crate::coding::VarInt::from(key)) - .unwrap(); + Encoder::new(&mut raw, version.into()).varint(key).unwrap(); Encoder::new(&mut raw, version.into()).bytes(&value).unwrap(); } diff --git a/rs/moq-net/src/ietf/track.rs b/rs/moq-net/src/ietf/track.rs index a066d04d96..6dca94d0f4 100644 --- a/rs/moq-net/src/ietf/track.rs +++ b/rs/moq-net/src/ietf/track.rs @@ -31,7 +31,7 @@ impl Message for TrackStatus<'_> { fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.request_id.encode(w, version)?; if version == Version::Draft17 { - w.varint(VarInt::ZERO)?; // required_request_id_delta = 0 + w.varint(0)?; // required_request_id_delta = 0 } encode_namespace(w, &self.track_namespace)?; w.string(&self.track_name)?; @@ -54,7 +54,7 @@ impl Message for TrackStatus<'_> { fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { let request_id = RequestId::decode(r, version)?; if version == Version::Draft17 { - let _required_request_id_delta = r.varint()?.into_inner(); + let _required_request_id_delta = r.varint()?; } let track_namespace = decode_namespace(r)?; let track_name = Cow::Owned(r.string()?); @@ -64,7 +64,7 @@ impl Message for TrackStatus<'_> { let _subscriber_priority = r.u8()?; let _group_order = GroupOrder::decode(r, version)?; let _forward = r.bool()?; - let _filter_type = r.varint()?.into_inner(); + let _filter_type = r.varint()?; let _params = Parameters::decode(r, version)?; } _ => { @@ -91,14 +91,14 @@ pub enum TrackStatusCode { impl Encode for TrackStatusCode { fn encode(&self, w: &mut Encoder<'_>, _: Version) -> Result<(), EncodeError> { - w.varint(VarInt::from(u64::from(*self)))?; + w.varint(u64::from(*self))?; Ok(()) } } impl Decode for TrackStatusCode { fn decode(r: &mut Decoder<'_>, _: Version) -> Result { - Self::try_from(r.varint()?.into_inner()).map_err(|_| DecodeError::InvalidValue) + Self::try_from(r.varint()?).map_err(|_| DecodeError::InvalidValue) } } diff --git a/rs/moq-net/src/lib.rs b/rs/moq-net/src/lib.rs index ffb5a32415..6d8797b294 100644 --- a/rs/moq-net/src/lib.rs +++ b/rs/moq-net/src/lib.rs @@ -100,7 +100,7 @@ pub mod time; pub mod transport; pub use client::*; -pub use coding::{BoundsExceeded, DecodeError, EncodeError, VarInt}; +pub use coding::{BoundsExceeded, DecodeError, EncodeError, varint}; pub use driver::Driver; pub use error::*; /// The session direction a client advertises in its SETUP (moq-lite-05+). diff --git a/rs/moq-net/src/lite/announce.rs b/rs/moq-net/src/lite/announce.rs index 3cf12efd2c..a2e9bfbc87 100644 --- a/rs/moq-net/src/lite/announce.rs +++ b/rs/moq-net/src/lite/announce.rs @@ -84,8 +84,8 @@ impl<'a> PathRef<'a> { impl Encode for PathRef<'_> { fn encode(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { if version.has_announce_compression() { - w.varint(VarInt::from(self.base))?; - w.varint(VarInt::from(self.keep))?; + w.varint(self.base)?; + w.varint(self.keep)?; } else if self.base != 0 || self.keep != 0 { return Err(EncodeError::Version); } @@ -98,8 +98,8 @@ impl Decode for PathRef<'_> { if !version.has_announce_compression() { return Ok(Self::literal(Path::decode(buf, version)?)); } - let base = buf.varint()?.into_inner(); - let keep = buf.varint()?.into_inner(); + let base = buf.varint()?; + let keep = buf.varint()?; if base == 0 && keep != 0 { return Err(DecodeError::InvalidValue); } @@ -139,9 +139,9 @@ impl Encode for HopsRef { } return self.literal.encode(w, version); } - w.varint(VarInt::from(self.base))?; + w.varint(self.base)?; self.literal.encode(w, version)?; - w.varint(VarInt::from(self.keep)) + w.varint(self.keep) } } @@ -150,9 +150,9 @@ impl Decode for HopsRef { if !version.has_announce_compression() { return Ok(Self::literal(Hops::decode(buf, version)?)); } - let base = buf.varint()?.into_inner(); + let base = buf.varint()?; let literal = Hops::decode(buf, version)?; - let keep = buf.varint()?.into_inner(); + let keep = buf.varint()?; if base == 0 && keep != 0 { return Err(DecodeError::InvalidValue); } @@ -190,8 +190,8 @@ impl Encode for Cost { if !version.has_route_cost() { return Ok(()); } - w.varint(VarInt::from(self.warm))?; - w.varint(VarInt::from(self.cold)) + w.varint(self.warm)?; + w.varint(self.cold) } } @@ -201,8 +201,8 @@ impl Decode for Cost { return Ok(Cost::UNKNOWN); } Ok(Cost { - warm: buf.varint()?.into_inner(), - cold: buf.varint()?.into_inner(), + warm: buf.varint()?, + cold: buf.varint()?, }) } } @@ -221,7 +221,7 @@ impl Encode for AnnounceBroadcast<'_> { // Decode-only: an unknown type is never sent. Self::Skipped => return Err(EncodeError::Unsupported), }; - w.varint(VarInt::from(typ))?; + w.varint(typ)?; let prefix = w.prefix_varint(); match self { @@ -230,9 +230,9 @@ impl Encode for AnnounceBroadcast<'_> { hops.encode(w, version)?; cost.encode(w, version)?; } - Self::EndedId { id } => w.varint(VarInt::from(*id))?, + Self::EndedId { id } => w.varint(*id)?, Self::Restart { id, hops, cost } => { - w.varint(VarInt::from(*id))?; + w.varint(*id)?; hops.encode(w, version)?; cost.encode(w, version)?; } @@ -273,7 +273,7 @@ impl Decode for AnnounceBroadcast<'_> { fn decode(buf: &mut Decoder<'_>, version: Version) -> Result { if version.has_announce_id() { // Lite06+: outer type, then a size-prefixed body decoded within its bounds. - let typ = buf.varint()?.into_inner(); + let typ = buf.varint()?; let size = decode_size(buf)?; let mut body = buf.sub(size)?; let msg = match typ { @@ -282,11 +282,9 @@ impl Decode for AnnounceBroadcast<'_> { hops: HopsRef::decode(&mut body, version)?, cost: Cost::decode(&mut body, version)?, }, - ANNOUNCE_END => Self::EndedId { - id: body.varint()?.into_inner(), - }, + ANNOUNCE_END => Self::EndedId { id: body.varint()? }, ANNOUNCE_RESTART => Self::Restart { - id: body.varint()?.into_inner(), + id: body.varint()?, hops: HopsRef::decode(&mut body, version)?, cost: Cost::decode(&mut body, version)?, }, @@ -325,7 +323,7 @@ impl AnnounceBroadcast<'_> { Version::Lite03 => { // Lite03 sends only a hop count, not individual ids. Fill with UNKNOWN placeholders. // push() enforces MAX_HOPS and `?` lifts the overflow to DecodeError::BoundsExceeded. - let count = r.varint()?.into_inner() as usize; + let count = r.varint()? as usize; let mut list = Hops::new(); for _ in 0..count { list.push(Hop::UNKNOWN)?; @@ -361,7 +359,7 @@ fn encode_hops(w: &mut Encoder<'_>, version: Version, hops: &Hops) -> Result<(), match version { Version::Lite01 | Version::Lite02 => Ok(()), Version::Lite03 => { - w.varint(VarInt::from(hops.len()))?; + w.varint(hops.len() as u64)?; Ok(()) } _ => hops.encode(w, version), @@ -389,7 +387,7 @@ impl Message for AnnounceRequest<'_> { fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { let prefix = Path::decode(r, version)?; let exclude_hop = match version.has_exclude_hop() { - true => r.varint()?.into_inner(), + true => r.varint()?, false => 0, }; let hidden = match version.has_hidden() { @@ -406,7 +404,7 @@ impl Message for AnnounceRequest<'_> { fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { self.prefix.encode(w, version)?; if version.has_exclude_hop() { - w.varint(VarInt::from(self.exclude_hop))?; + w.varint(self.exclude_hop)?; } if version.has_hidden() { w.bool(self.hidden); @@ -459,7 +457,7 @@ impl Message for AnnounceInit<'_> { } } - let count = r.varint()?.into_inner(); + let count = r.varint()?; // Don't allocate more than 1024 elements upfront let mut paths = Vec::with_capacity(count.min(1024) as usize); @@ -479,7 +477,7 @@ impl Message for AnnounceInit<'_> { } } - w.varint(VarInt::from(self.suffixes.len()))?; + w.varint(self.suffixes.len() as u64)?; for path in &self.suffixes { path.encode(w, version)?; } @@ -509,7 +507,7 @@ impl Message for AnnounceOk { } let origin = Hop::decode(r, version)?; - let active = r.varint()?.into_inner(); + let active = r.varint()?; Ok(Self { origin, active }) } @@ -519,7 +517,7 @@ impl Message for AnnounceOk { } self.origin.encode(w, version)?; - w.varint(VarInt::from(self.active)) + w.varint(self.active) } } @@ -808,11 +806,9 @@ mod tests { .unwrap(); let mut buf = Vec::new(); + Encoder::new(&mut buf, Version::Lite06.into()).varint(4u64).unwrap(); Encoder::new(&mut buf, Version::Lite06.into()) - .varint(VarInt::from(4u64)) - .unwrap(); - Encoder::new(&mut buf, Version::Lite06.into()) - .varint(VarInt::from(body.len())) + .varint(body.len() as u64) .unwrap(); buf.extend_from_slice(&body); @@ -871,7 +867,7 @@ mod tests { .unwrap(); body.push(2); Encoder::new(&mut buf, Version::Lite07.into()) - .varint(VarInt::from(body.len())) + .varint(body.len() as u64) .unwrap(); buf.extend_from_slice(&body); assert!(crate::coding::decode_buf(&mut &buf[..], Version::Lite07, AnnounceRequest::decode).is_err()); diff --git a/rs/moq-net/src/lite/compress.rs b/rs/moq-net/src/lite/compress.rs index 8f35f992b3..b13edc5ea5 100644 --- a/rs/moq-net/src/lite/compress.rs +++ b/rs/moq-net/src/lite/compress.rs @@ -15,7 +15,7 @@ use std::{ use crate::{ Error, Hop, Hops, Path, PathOwned, - coding::{Form, VarInt}, + coding::{Form, varint}, }; use super::{HopsRef, PathRef, Version}; @@ -336,9 +336,7 @@ impl AnnounceEncoder { /// The bytes `value` takes as a moq-lite varint. fn varint_size(value: u64) -> usize { - VarInt::from(value) - .size(Form::Quic) - .expect("sizing a value in varint range") + varint::size(value, Form::Quic).expect("sizing a value in varint range") } /// The bytes `path` takes on the wire: a varint length, then the string. diff --git a/rs/moq-net/src/lite/datagram.rs b/rs/moq-net/src/lite/datagram.rs index e02c44cf2b..90055d2fcc 100644 --- a/rs/moq-net/src/lite/datagram.rs +++ b/rs/moq-net/src/lite/datagram.rs @@ -7,7 +7,7 @@ use bytes::Bytes; -use crate::coding::{DecodeError, Decoder, Encode, EncodeError, Encoder, VarInt}; +use crate::coding::{DecodeError, Decoder, Encode, EncodeError, Encoder}; use super::Version; @@ -30,9 +30,9 @@ impl Encode for Datagram { return Err(EncodeError::Version); } - w.varint(VarInt::from(self.subscribe))?; - w.varint(VarInt::from(self.sequence))?; - w.varint(VarInt::from(self.timestamp))?; + w.varint(self.subscribe)?; + w.varint(self.sequence)?; + w.varint(self.timestamp)?; // Payload runs to the datagram boundary: written raw, no length prefix. w.slice(&self.payload); @@ -48,9 +48,9 @@ impl Datagram { } let mut r = Decoder::new(&buf, version.into()); - let subscribe = r.varint()?.into_inner(); - let sequence = r.varint()?.into_inner(); - let timestamp = r.varint()?.into_inner(); + let subscribe = r.varint()?; + let sequence = r.varint()?; + let timestamp = r.varint()?; // Everything remaining is the payload (the datagram boundary delimits it). let payload = buf.slice(buf.len() - r.remaining()..); diff --git a/rs/moq-net/src/lite/fetch.rs b/rs/moq-net/src/lite/fetch.rs index a234fbab8d..7bb16636c6 100644 --- a/rs/moq-net/src/lite/fetch.rs +++ b/rs/moq-net/src/lite/fetch.rs @@ -2,7 +2,7 @@ use std::borrow::Cow; use crate::{ Path, - coding::{Decode, DecodeError, Decoder, Encode, EncodeError, Encoder, VarInt}, + coding::{Decode, DecodeError, Decoder, Encode, EncodeError, Encoder}, }; use super::{Message, Version}; @@ -35,10 +35,10 @@ impl Message for Fetch<'_> { let broadcast = Path::decode(r, version)?; let track = Cow::Owned(r.string()?); let priority = r.u8()?; - let group = r.varint()?.into_inner(); + let group = r.varint()?; let (start_frame, end_frame) = match version.has_frame_bounds() { - true => (r.varint()?.into_inner(), r.varint_opt()?), + true => (r.varint()?, r.varint_opt()?), false => (0, None), }; // A range that ends before it starts can never be served. @@ -67,10 +67,10 @@ impl Message for Fetch<'_> { self.broadcast.encode(w, version)?; w.string(&self.track)?; w.u8(self.priority); - w.varint(VarInt::from(self.group))?; + w.varint(self.group)?; if version.has_frame_bounds() { - w.varint(VarInt::from(self.start_frame))?; + w.varint(self.start_frame)?; w.varint_opt(self.end_frame)?; } else if self.start_frame != 0 || self.end_frame.is_some() { // The peer would serve the whole group, including frames the caller excluded. diff --git a/rs/moq-net/src/lite/goaway.rs b/rs/moq-net/src/lite/goaway.rs index a7b7e7fb10..f142a50ec7 100644 --- a/rs/moq-net/src/lite/goaway.rs +++ b/rs/moq-net/src/lite/goaway.rs @@ -25,7 +25,7 @@ impl Message for Goaway<'_> { // cap. Rejected from the string's length prefix alone, before allocating // or validating the payload. (Buffering is bounded separately by the // outer message-size prefix that frames every lite control message.) - let len = r.varint()?.into_inner(); + let len = r.varint()?; if len > 8192 { return Err(DecodeError::InvalidValue); } diff --git a/rs/moq-net/src/lite/group.rs b/rs/moq-net/src/lite/group.rs index 2cf52a84d9..3e18a108a0 100644 --- a/rs/moq-net/src/lite/group.rs +++ b/rs/moq-net/src/lite/group.rs @@ -19,10 +19,10 @@ pub struct Group { impl Message for Group { fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { - let subscribe = r.varint()?.into_inner(); - let sequence = r.varint()?.into_inner(); + let subscribe = r.varint()?; + let sequence = r.varint()?; let frame_start = match version.has_frame_bounds() { - true => r.varint()?.into_inner(), + true => r.varint()?, false => 0, }; @@ -34,11 +34,11 @@ impl Message for Group { } fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { - w.varint(VarInt::from(self.subscribe))?; - w.varint(VarInt::from(self.sequence))?; + w.varint(self.subscribe)?; + w.varint(self.sequence)?; if version.has_frame_bounds() { - w.varint(VarInt::from(self.frame_start))?; + w.varint(self.frame_start)?; } else if self.frame_start != 0 { // The peer would number the frames from 0 and silently misalign the group. return Err(EncodeError::Version); diff --git a/rs/moq-net/src/lite/info.rs b/rs/moq-net/src/lite/info.rs index 102c3c3c3f..301764fdca 100644 --- a/rs/moq-net/src/lite/info.rs +++ b/rs/moq-net/src/lite/info.rs @@ -16,7 +16,7 @@ impl Message for SessionInfo { } } - let bitrate = match r.varint()?.into_inner() { + let bitrate = match r.varint()? { 0 => None, bitrate => Some(bitrate), }; @@ -32,7 +32,7 @@ impl Message for SessionInfo { } } - w.varint(VarInt::from(self.bitrate.unwrap_or(0)))?; + w.varint(self.bitrate.unwrap_or(0))?; Ok(()) } } diff --git a/rs/moq-net/src/lite/message.rs b/rs/moq-net/src/lite/message.rs index 574c6d5487..f601f97521 100644 --- a/rs/moq-net/src/lite/message.rs +++ b/rs/moq-net/src/lite/message.rs @@ -8,7 +8,7 @@ pub(super) const MAX_MESSAGE_SIZE: usize = 64 * 1024 * 1024; /// Read a lite message's varint size prefix, refusing one past [`MAX_MESSAGE_SIZE`]. pub(super) fn decode_size(r: &mut Decoder<'_>) -> Result { - let size = r.varint()?.into_inner(); + let size = r.varint()?; match usize::try_from(size) { Ok(size) if size <= MAX_MESSAGE_SIZE => Ok(size), _ => Err(DecodeError::MessageTooLarge { @@ -59,7 +59,6 @@ impl Decode for T { #[cfg(test)] mod tests { use super::*; - use crate::coding::VarInt; #[derive(Debug)] struct Empty; @@ -74,11 +73,18 @@ mod tests { } } + /// A lite size prefix announcing `size` bytes, with no body behind it. + fn prefix(size: usize) -> Vec { + let mut wire = Vec::new(); + Encoder::new(&mut wire, Version::Lite06.into()) + .varint(size as u64) + .unwrap(); + wire + } + #[test] fn rejects_oversized_message_before_reading_the_body() { - let wire = VarInt::from(MAX_MESSAGE_SIZE + 1) - .encode_bytes(Version::Lite06) - .unwrap(); + let wire = prefix(MAX_MESSAGE_SIZE + 1); let err = Empty::decode_slice(&wire, Version::Lite06).unwrap_err(); assert!(matches!( @@ -92,7 +98,7 @@ mod tests { #[test] fn accepts_message_at_the_limit() { - let wire = VarInt::from(MAX_MESSAGE_SIZE).encode_bytes(Version::Lite06).unwrap(); + let wire = prefix(MAX_MESSAGE_SIZE); let err = Empty::decode_slice(&wire, Version::Lite06).unwrap_err(); assert!(matches!(err, DecodeError::Short)); diff --git a/rs/moq-net/src/lite/parameters.rs b/rs/moq-net/src/lite/parameters.rs index e10cd1909b..10927d2e47 100644 --- a/rs/moq-net/src/lite/parameters.rs +++ b/rs/moq-net/src/lite/parameters.rs @@ -29,11 +29,11 @@ impl Parameters { /// Set a parameter to a varint value, replacing any existing entry. /// - /// Panics past [`VarInt::MAX_QUIC`], which no parameter we set comes near. + /// Panics past [`crate::coding::varint::MAX_QUIC`], which no parameter we set comes near. pub fn set_varint(&mut self, id: u64, value: u64) { let mut buf = Vec::new(); Encoder::new(&mut buf, Form::Quic) - .varint(value.into()) + .varint(value) .expect("parameter varint in range"); self.set_bytes(id, buf); } @@ -44,7 +44,7 @@ impl Parameters { return Ok(None); }; let mut r = Decoder::new(bytes, Form::Quic); - let value = r.varint()?.into_inner(); + let value = r.varint()?; if !r.is_empty() { return Err(DecodeError::Long); } @@ -57,13 +57,13 @@ impl Decode for Parameters { let mut params = Self::default(); // I hate this encoding so much; let me encode my role and get on with my life. - let count = r.varint()?.into_inner(); + let count = r.varint()?; if count > MAX_PARAMS { return Err(DecodeError::TooMany); } for _ in 0..count { - let kind = r.varint()?.into_inner(); + let kind = r.varint()?; if params.get_bytes(kind).is_some() { return Err(DecodeError::Duplicate); } @@ -81,10 +81,10 @@ impl Encode for Parameters { return Err(EncodeError::TooMany); } - w.varint(VarInt::from(self.0.len()))?; + w.varint(self.0.len() as u64)?; for (kind, value) in &self.0 { - w.varint(VarInt::from(*kind))?; + w.varint(*kind)?; w.bytes(value)?; } diff --git a/rs/moq-net/src/lite/probe.rs b/rs/moq-net/src/lite/probe.rs index 59347600dc..812d1f8362 100644 --- a/rs/moq-net/src/lite/probe.rs +++ b/rs/moq-net/src/lite/probe.rs @@ -23,13 +23,13 @@ impl Message for Probe { // 0 means unknown, the same as RTT below. A publisher whose transport // exposes no congestion controller reports the RTT half alone. - let bitrate = match r.varint()?.into_inner() { + let bitrate = match r.varint()? { 0 => None, v => Some(v), }; let rtt = match version.has_probe_rtt() { false => None, - true => match r.varint()?.into_inner() { + true => match r.varint()? { 0 => None, v => Some(v), }, @@ -48,10 +48,10 @@ impl Message for Probe { // 0 means unknown; round Some(0) up to 1. let wire = self.bitrate.map(|v| v.max(1)).unwrap_or(0); - w.varint(VarInt::from(wire))?; + w.varint(wire)?; if version.has_probe_rtt() { let wire = self.rtt.map(|v| v.max(1)).unwrap_or(0); - w.varint(VarInt::from(wire))?; + w.varint(wire)?; } Ok(()) } diff --git a/rs/moq-net/src/lite/publisher.rs b/rs/moq-net/src/lite/publisher.rs index bd8c428237..80336ec56a 100644 --- a/rs/moq-net/src/lite/publisher.rs +++ b/rs/moq-net/src/lite/publisher.rs @@ -1938,7 +1938,7 @@ fn buffer_frame_info( if timescale.is_some() { buffer_zigzag_delta(writer, timestamp.value(), prev_ts)?; } - writer.buffer(&crate::coding::VarInt::from(size))?; + writer.buffer_varint(size)?; Ok(()) } @@ -1952,7 +1952,7 @@ fn buffer_zigzag_delta( let delta: i64 = (curr as i128 - *prev as i128) .try_into() .map_err(|_| Error::BoundsExceeded(crate::coding::BoundsExceeded))?; - writer.buffer(&crate::coding::VarInt::from_zigzag(delta))?; + writer.buffer_varint(crate::coding::varint::zigzag(delta))?; *prev = curr; Ok(()) } diff --git a/rs/moq-net/src/lite/setup.rs b/rs/moq-net/src/lite/setup.rs index b20631ea5f..77014db0e5 100644 --- a/rs/moq-net/src/lite/setup.rs +++ b/rs/moq-net/src/lite/setup.rs @@ -399,7 +399,7 @@ mod tests { // Frame the body with the Message Length prefix `Setup::decode` expects. let mut buf = Vec::new(); Encoder::new(&mut buf, version.into()) - .varint(VarInt::from(body.len())) + .varint(body.len() as u64) .unwrap(); buf.extend_from_slice(&body); let mut slice = &buf[..]; @@ -443,7 +443,7 @@ mod tests { let mut buf = Vec::new(); Encoder::new(&mut buf, Version::Lite05.into()) - .varint(VarInt::from(body.len())) + .varint(body.len() as u64) .unwrap(); buf.extend_from_slice(&body); @@ -477,7 +477,7 @@ mod tests { let mut buf = Vec::new(); Encoder::new(&mut buf, Version::Lite05.into()) - .varint(VarInt::from(body.len())) + .varint(body.len() as u64) .unwrap(); buf.extend_from_slice(&body); @@ -512,7 +512,7 @@ mod tests { // Wrap with the message size prefix the Message impl expects. let mut buf = Vec::new(); Encoder::new(&mut buf, Version::Lite05.into()) - .varint(VarInt::from(body.len())) + .varint(body.len() as u64) .unwrap(); buf.extend_from_slice(&body); diff --git a/rs/moq-net/src/lite/stream.rs b/rs/moq-net/src/lite/stream.rs index 8c0f89ce53..9de36c9594 100644 --- a/rs/moq-net/src/lite/stream.rs +++ b/rs/moq-net/src/lite/stream.rs @@ -18,7 +18,7 @@ pub enum ControlType { impl Decode for ControlType { fn decode(r: &mut Decoder<'_>, _: Version) -> Result { - let t = r.varint()?.into_inner(); + let t = r.varint()?; t.try_into().map_err(|_| DecodeError::InvalidValue) } } @@ -26,7 +26,7 @@ impl Decode for ControlType { impl Encode for ControlType { fn encode(&self, w: &mut Encoder<'_>, _: Version) -> Result<(), EncodeError> { let v: u64 = (*self).into(); - w.varint(VarInt::from(v))?; + w.varint(v)?; Ok(()) } } @@ -42,7 +42,7 @@ pub enum DataType { impl Decode for DataType { fn decode(r: &mut Decoder<'_>, _: Version) -> Result { - let t = r.varint()?.into_inner(); + let t = r.varint()?; t.try_into().map_err(|_| DecodeError::InvalidValue) } } @@ -50,7 +50,7 @@ impl Decode for DataType { impl Encode for DataType { fn encode(&self, w: &mut Encoder<'_>, _: Version) -> Result<(), EncodeError> { let v: u64 = (*self).into(); - w.varint(VarInt::from(v))?; + w.varint(v)?; Ok(()) } } diff --git a/rs/moq-net/src/lite/subscribe.rs b/rs/moq-net/src/lite/subscribe.rs index 2ef7c983cf..cb6de29c5a 100644 --- a/rs/moq-net/src/lite/subscribe.rs +++ b/rs/moq-net/src/lite/subscribe.rs @@ -2,7 +2,7 @@ use std::borrow::Cow; use crate::{ Path, - coding::{Decode, DecodeError, Decoder, Encode, EncodeError, Encoder, VarInt}, + coding::{Decode, DecodeError, Decoder, Encode, EncodeError, Encoder}, }; use super::{Message, Version}; @@ -45,7 +45,7 @@ impl Version { impl Message for Subscribe<'_> { fn decode_msg(r: &mut Decoder<'_>, version: Version) -> Result { - let id = r.varint()?.into_inner(); + let id = r.varint()?; let broadcast = Path::decode(r, version)?; let track = Cow::Owned(r.string()?); let priority = r.u8()?; @@ -54,7 +54,7 @@ impl Message for Subscribe<'_> { Version::Lite01 | Version::Lite02 => (std::time::Duration::ZERO, None, None), _ => { skip_group_order(r, version)?; - let max_age = std::time::Duration::from_millis(r.varint()?.into_inner()); + let max_age = std::time::Duration::from_millis(r.varint()?); let start_group = decode_start_group(r, version)?; let end_group = r.varint_opt()?; (max_age, start_group, end_group) @@ -78,7 +78,7 @@ impl Message for Subscribe<'_> { } fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { - w.varint(VarInt::from(self.id))?; + w.varint(self.id)?; self.broadcast.encode(w, version)?; w.string(&self.track)?; w.u8(self.priority); @@ -87,7 +87,7 @@ impl Message for Subscribe<'_> { Version::Lite01 | Version::Lite02 => {} _ => { pad_group_order(w, version)?; - w.varint(VarInt::try_from(self.max_age.as_millis())?)?; + w.varint(u64::try_from(self.max_age.as_millis()).map_err(|_| EncodeError::BoundsExceeded)?)?; encode_start_group(w, version, self.start_group)?; w.varint_opt(self.end_group)?; } @@ -133,7 +133,7 @@ pub(super) fn pad_group_order(w: &mut Encoder<'_>, version: Version) -> Result<( /// are known. fn decode_start_group(r: &mut Decoder<'_>, version: Version) -> Result, DecodeError> { if version.resolves_start() { - return Ok(Some(r.varint()?.into_inner())); + return Ok(Some(r.varint()?)); } r.varint_opt() } @@ -157,7 +157,7 @@ fn canonical_start_group(version: Version, start_group: Option, start_frame /// beginning", which is not what a vacuous floor asks for. fn encode_start_group(w: &mut Encoder<'_>, version: Version, start_group: Option) -> Result<(), EncodeError> { if version.resolves_start() { - return w.varint(VarInt::from(start_group.unwrap_or(0))); + return w.varint(start_group.unwrap_or(0)); } w.varint_opt(start_group.filter(|&group| group > 0)) } @@ -178,7 +178,7 @@ fn decode_frame_bounds( return Ok((0, None)); } - let start_frame = r.varint()?.into_inner(); + let start_frame = r.varint()?; let end_frame = r.varint_opt()?; if (start_frame != 0 && start_group.is_none()) || (end_frame.is_some() && end_group.is_none()) { @@ -210,7 +210,7 @@ fn encode_frame_bounds( return Ok(()); } - w.varint(VarInt::from(start_frame))?; + w.varint(start_frame)?; w.varint_opt(end_frame) } @@ -239,7 +239,7 @@ impl Message for SubscribeOk { _ => { w.u8(self.priority); pad_group_order(w, version)?; - w.varint(VarInt::try_from(self.max_age.as_millis())?)?; + w.varint(u64::try_from(self.max_age.as_millis()).map_err(|_| EncodeError::BoundsExceeded)?)?; w.varint_opt(self.start_group)?; w.varint_opt(self.end_group)?; } @@ -265,7 +265,7 @@ impl Message for SubscribeOk { _ => { let priority = r.u8()?; skip_group_order(r, version)?; - let max_age = std::time::Duration::from_millis(r.varint()?.into_inner()); + let max_age = std::time::Duration::from_millis(r.varint()?); let start_group = r.varint_opt()?; let end_group = r.varint_opt()?; @@ -298,16 +298,14 @@ impl Message for SubscribeStart { if !version.has_track_stream() { return Err(DecodeError::Version); } - Ok(Self { - group: r.varint()?.into_inner(), - }) + Ok(Self { group: r.varint()? }) } fn encode_msg(&self, w: &mut Encoder<'_>, version: Version) -> Result<(), EncodeError> { if !version.has_track_stream() { return Err(EncodeError::Version); } - w.varint(VarInt::from(self.group)) + w.varint(self.group) } } @@ -328,9 +326,9 @@ impl Message for SubscribeEnd { if !version.has_track_stream() { return Err(DecodeError::Version); } - let group = r.varint()?.into_inner(); + let group = r.varint()?; let streams = match version.has_stream_count() { - true => r.varint()?.into_inner(), + true => r.varint()?, false => 0, }; Ok(Self { group, streams }) @@ -340,9 +338,9 @@ impl Message for SubscribeEnd { if !version.has_track_stream() { return Err(EncodeError::Version); } - w.varint(VarInt::from(self.group))?; + w.varint(self.group)?; if version.has_stream_count() { - w.varint(VarInt::from(self.streams))?; + w.varint(self.streams)?; } Ok(()) } @@ -375,9 +373,9 @@ impl Message for SubscribeUpdate { let priority = r.u8()?; skip_group_order(r, version)?; - let max_age = std::time::Duration::from_millis(r.varint()?.into_inner()); + let max_age = std::time::Duration::from_millis(r.varint()?); let start_group = decode_start_group(r, version)?; - let end_group = match r.varint()?.into_inner() { + let end_group = match r.varint()? { 0 => None, group => Some(group - 1), }; @@ -405,7 +403,7 @@ impl Message for SubscribeUpdate { w.u8(self.priority); pad_group_order(w, version)?; - w.varint(VarInt::try_from(self.max_age.as_millis())?)?; + w.varint(u64::try_from(self.max_age.as_millis()).map_err(|_| EncodeError::BoundsExceeded)?)?; encode_start_group(w, version, self.start_group)?; @@ -454,9 +452,9 @@ impl Message for SubscribeDrop { } Ok(Self { - start: r.varint()?.into_inner(), - end: r.varint()?.into_inner(), - error: r.varint()?.into_inner(), + start: r.varint()?, + end: r.varint()?, + error: r.varint()?, }) } @@ -469,9 +467,9 @@ impl Message for SubscribeDrop { _ => {} } - w.varint(VarInt::from(self.start))?; - w.varint(VarInt::from(self.end))?; - w.varint(VarInt::from(self.error))?; + w.varint(self.start)?; + w.varint(self.end)?; + w.varint(self.error)?; Ok(()) } @@ -495,7 +493,7 @@ pub enum SubscribeResponse { /// Write a `type` varint followed by the size-prefixed message body. fn encode_typed(w: &mut Encoder<'_>, typ: u64, msg: &M, version: Version) -> Result<(), EncodeError> { - w.varint(VarInt::from(typ))?; + w.varint(typ)?; msg.encode(w, version) } @@ -529,7 +527,7 @@ impl Decode for SubscribeResponse { match version { Version::Lite01 | Version::Lite02 => Ok(Self::Ok(SubscribeOk::decode(buf, version)?)), Version::Lite03 | Version::Lite04 => { - let typ = buf.varint()?.into_inner(); + let typ = buf.varint()?; match typ { 0 => Ok(Self::Ok(SubscribeOk::decode(buf, version)?)), 1 => Ok(Self::Drop(SubscribeDrop::decode(buf, version)?)), @@ -537,7 +535,7 @@ impl Decode for SubscribeResponse { } } _ => { - let typ = buf.varint()?.into_inner(); + let typ = buf.varint()?; match typ { 0 => Ok(Self::Start(SubscribeStart::decode(buf, version)?)), 1 => Ok(Self::End(SubscribeEnd::decode(buf, version)?)), diff --git a/rs/moq-net/src/lite/subscriber.rs b/rs/moq-net/src/lite/subscriber.rs index 964a2e38dd..33091356b3 100644 --- a/rs/moq-net/src/lite/subscriber.rs +++ b/rs/moq-net/src/lite/subscriber.rs @@ -889,10 +889,10 @@ impl FrameIngest { continue; }; // The timestamp delta doubles as the per-frame sentinel. - let Some(zz) = ready!(reader.poll_decode_maybe::(&mut cx))? else { + let Some(zz) = ready!(reader.poll_varint_maybe(&mut cx))? else { return Poll::Ready(Ok(())); }; - let next: u64 = (self.prev_ts as i128 + zz.to_zigzag() as i128) + let next: u64 = (self.prev_ts as i128 + crate::coding::varint::unzigzag(zz) as i128) .try_into() .map_err(|_| Error::BoundsExceeded(crate::coding::BoundsExceeded))?; self.prev_ts = next; @@ -903,10 +903,10 @@ impl FrameIngest { }; } IngestPhase::Size { timestamp } => { - let Some(size) = ready!(reader.poll_decode_maybe::(&mut cx))? else { + let Some(size) = ready!(reader.poll_varint_maybe(&mut cx))? else { return Poll::Ready(Ok(())); }; - let size = size.into_inner(); + // `create_frame_owned` is the allocation chokepoint and rejects an // oversized `size` before allocating, so no pre-check is needed. No // wire timestamp (pre-lite-05) means local receive time. diff --git a/rs/moq-net/src/lite/track.rs b/rs/moq-net/src/lite/track.rs index f764dd7c27..ffd21f1a81 100644 --- a/rs/moq-net/src/lite/track.rs +++ b/rs/moq-net/src/lite/track.rs @@ -2,7 +2,7 @@ use std::{borrow::Cow, time::Duration}; use crate::{ Path, Timescale, - coding::{Decode, DecodeError, Decoder, Encode, EncodeError, Encoder, VarInt}, + coding::{Decode, DecodeError, Decoder, Encode, EncodeError, Encoder}, }; use super::{Message, Version}; @@ -68,12 +68,12 @@ impl Message for TrackInfo { let priority = r.u8()?; super::subscribe::skip_group_order(r, version)?; - let encoded = r.varint()?.into_inner(); + let encoded = r.varint()?; let max_age = match version { Version::Lite05 | Version::Lite06 => (encoded < LEGACY_UNLIMITED).then(|| Duration::from_millis(encoded)), _ => encoded.checked_sub(1).map(Duration::from_millis), }; - let timescale = Timescale::new(r.varint()?.into_inner()).map_err(|_| DecodeError::InvalidValue)?; + let timescale = Timescale::new(r.varint()?).map_err(|_| DecodeError::InvalidValue)?; Ok(Self { priority, @@ -95,8 +95,8 @@ impl Message for TrackInfo { (_, None) => 0, (_, Some(age)) => u64::try_from(age.as_millis() + 1).map_err(|_| EncodeError::BoundsExceeded)?, }; - w.varint(VarInt::from(encoded))?; - w.varint(VarInt::from(u64::from(self.timescale)))?; + w.varint(encoded)?; + w.varint(u64::from(self.timescale))?; Ok(()) } } @@ -160,8 +160,8 @@ mod test { let w = &mut Encoder::new(&mut raw, version.into()); w.u8(0); super::super::subscribe::pad_group_order(w, version).unwrap(); - w.varint(millis.into()).unwrap(); - w.varint(1000u64.into()).unwrap(); + w.varint(millis).unwrap(); + w.varint(1000u64).unwrap(); let decoded = TrackInfo::decode_msg(&mut Decoder::new(&raw, version.into()), version).unwrap(); assert_eq!( decoded.max_age, @@ -177,7 +177,7 @@ mod test { let old_reader = &mut Decoder::new(&encoded, version.into()); old_reader.u8().unwrap(); super::super::subscribe::skip_group_order(old_reader, version).unwrap(); - assert_eq!(old_reader.varint().unwrap().into_inner(), millis.min(LEGACY_UNLIMITED)); + assert_eq!(old_reader.varint().unwrap(), millis.min(LEGACY_UNLIMITED)); } } } diff --git a/rs/moq-net/src/model/origin.rs b/rs/moq-net/src/model/origin.rs index a45f8dbcf8..45f14575ba 100644 --- a/rs/moq-net/src/model/origin.rs +++ b/rs/moq-net/src/model/origin.rs @@ -18,7 +18,7 @@ use super::{ }; use crate::{ AsPath, Error, InvalidPattern, Path, PathOwned, Pattern, Patterns, - coding::{BoundsExceeded, Decode, DecodeError, Decoder, Encode, EncodeError, Encoder, VarInt}, + coding::{BoundsExceeded, Decode, DecodeError, Decoder, Encode, EncodeError, Encoder}, path::Segment, runtime::{Instant, Timers}, time::Clock, @@ -158,14 +158,14 @@ impl fmt::Display for Hop { impl Encode for Hop { fn encode(&self, w: &mut Encoder<'_>, _: V) -> Result<(), EncodeError> { - w.varint(VarInt::from(self.id))?; + w.varint(self.id)?; Ok(()) } } impl Decode for Hop { fn decode(r: &mut Decoder<'_>, _: V) -> Result { - Self::from_wire(r.varint()?.into_inner()) + Self::from_wire(r.varint()?) } } @@ -301,7 +301,7 @@ impl<'a> IntoIterator for &'a Hops { impl Encode for Hops { fn encode(&self, w: &mut Encoder<'_>, version: V) -> Result<(), EncodeError> { - w.varint(VarInt::from(self.0.len()))?; + w.varint(self.0.len() as u64)?; for origin in &self.0 { origin.encode(w, version)?; } @@ -311,7 +311,7 @@ impl Encode for Hops { impl Decode for Hops { fn decode(r: &mut Decoder<'_>, version: V) -> Result { - let count = r.varint()?.into_inner() as usize; + let count = r.varint()? as usize; if count > MAX_HOPS { return Err(DecodeError::BoundsExceeded); } diff --git a/rs/moq-net/src/model/time.rs b/rs/moq-net/src/model/time.rs index 764ff26fb3..aa6284fbc4 100644 --- a/rs/moq-net/src/model/time.rs +++ b/rs/moq-net/src/model/time.rs @@ -1,18 +1,18 @@ use std::num::NonZero; -use crate::coding::VarInt; +use crate::coding::varint::MAX_QUIC; -/// `value` as a [`VarInt`] the QUIC form can carry, or `None` past `2^62 - 1`, so every -/// timestamp stays encodable on moq-lite. -const fn quic(value: u128) -> Option { - if value <= VarInt::MAX_QUIC.into_inner() as u128 { - Some(VarInt::from_u64(value as u64)) +/// `value`, or `None` past the QUIC varint limit (`2^62 - 1`), so every timestamp stays +/// encodable on moq-lite. +const fn quic(value: u128) -> Option { + if value <= MAX_QUIC as u128 { + Some(value as u64) } else { None } } -/// Returned when a [`Timestamp`] operation would exceed the QUIC VarInt range +/// Returned when a [`Timestamp`] operation would exceed the QUIC varint range /// (`2^62 - 1`), overflow during scale conversion or arithmetic, or attempt /// arithmetic between timestamps with mismatched scales. #[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)] @@ -136,7 +136,7 @@ impl std::fmt::Display for Timescale { /// A timestamp in a track's timescale (units per second). /// /// All timestamps within a track are relative, so zero for one track is not zero for another. -/// The underlying value is constrained to fit within a QUIC VarInt (`2^62 - 1`) so it can be +/// The underlying value is constrained to fit within a QUIC varint (`2^62 - 1`) so it can be /// encoded and decoded easily; the scale is carried alongside so frames from different /// sources can be compared and converted without lossy detours through a single fixed scale. /// @@ -167,7 +167,7 @@ impl std::fmt::Display for Timescale { /// want "same instant regardless of encoding", compare after a [`Self::convert`] to a common scale. #[derive(Clone, Copy, PartialEq, Eq, Hash)] pub struct Timestamp { - value: VarInt, + value: u64, scale: Timescale, } @@ -226,7 +226,7 @@ impl Timestamp { /// The raw value in the timestamp's own scale. pub const fn value(self) -> u64 { - self.value.into_inner() + self.value } /// The scale (units per second) attached to this timestamp. @@ -236,7 +236,7 @@ impl Timestamp { /// Whether the raw value is zero. Does not consider scale. pub const fn is_zero(self) -> bool { - self.value.into_inner() == 0 + self.value == 0 } /// Re-express this timestamp at a new scale. Returns [`TimeOverflow`] if the new @@ -245,7 +245,7 @@ impl Timestamp { if self.scale.0.get() == new_scale.0.get() { return Ok(self); } - match (self.value.into_inner() as u128).checked_mul(new_scale.0.get() as u128) { + match (self.value as u128).checked_mul(new_scale.0.get() as u128) { Some(scaled) => match quic(scaled / self.scale.0.get() as u128) { Some(value) => Ok(Self { value, @@ -259,12 +259,12 @@ impl Timestamp { /// The value re-expressed at `target` as a `u128`. pub const fn as_scale(self, target: Timescale) -> u128 { - self.value.into_inner() as u128 * target.0.get() as u128 / self.scale.0.get() as u128 + self.value as u128 * target.0.get() as u128 / self.scale.0.get() as u128 } /// The value re-expressed in seconds. pub const fn as_secs(self) -> u64 { - self.value.into_inner() / self.scale.0.get() + self.value / self.scale.0.get() } /// The value re-expressed in milliseconds. @@ -288,7 +288,7 @@ impl Timestamp { if self.scale.0.get() != rhs.scale.0.get() { return Err(TimeOverflow); } - match self.value.into_inner().checked_add(rhs.value.into_inner()) { + match self.value.checked_add(rhs.value) { Some(result) => Self::new(result, self.scale), None => Err(TimeOverflow), } @@ -300,7 +300,7 @@ impl Timestamp { if self.scale.0.get() != rhs.scale.0.get() { return Err(TimeOverflow); } - match self.value.into_inner().checked_sub(rhs.value.into_inner()) { + match self.value.checked_sub(rhs.value) { Some(result) => Self::new(result, self.scale), None => Err(TimeOverflow), } @@ -377,8 +377,8 @@ impl Ord for Timestamp { if self.scale.0.get() == other.scale.0.get() { return self.value.cmp(&other.value); } - let lhs = self.value.into_inner() as u128 * other.scale.0.get() as u128; - let rhs = other.value.into_inner() as u128 * self.scale.0.get() as u128; + let lhs = self.value as u128 * other.scale.0.get() as u128; + let rhs = other.value as u128 * self.scale.0.get() as u128; lhs.cmp(&rhs) .then_with(|| self.scale.0.get().cmp(&other.scale.0.get())) .then_with(|| self.value.cmp(&other.value)) diff --git a/rs/moq-net/src/setup.rs b/rs/moq-net/src/setup.rs index f5be4f531c..4d12549e74 100644 --- a/rs/moq-net/src/setup.rs +++ b/rs/moq-net/src/setup.rs @@ -6,7 +6,7 @@ use bytes::Bytes; use crate::{ Version, - coding::{self, Decode, DecodeError, Decoder, Encode, EncodeError, Encoder, VarInt}, + coding::{self, Decode, DecodeError, Decoder, Encode, EncodeError, Encoder}, ietf, lite, }; @@ -53,7 +53,7 @@ impl Setup { impl Encode for Setup { fn encode(&self, w: &mut Encoder<'_>, v: Version) -> Result<(), EncodeError> { Self::check_version(v); - w.varint(VarInt::from(SETUP_V17))?; + w.varint(SETUP_V17)?; let prefix = w.prefix_u16(); w.slice(&self.parameters); w.fill(prefix) @@ -63,7 +63,7 @@ impl Encode for Setup { impl Decode for Setup { fn decode(r: &mut Decoder<'_>, v: Version) -> Result { Self::check_version(v); - let kind = r.varint()?.into_inner(); + let kind = r.varint()?; if kind != SETUP_V17 { return Err(DecodeError::InvalidValue); } @@ -164,7 +164,7 @@ impl Encode for Client { fn decode_body<'a>(r: &mut Decoder<'a>, v: Version) -> Result, DecodeError> { let size = match SetupVersion::from_version(v) { SetupVersion::Draft14 | SetupVersion::Draft15Plus => r.u16()? as usize, - SetupVersion::LiteLegacy => usize::try_from(r.varint()?)?, + SetupVersion::LiteLegacy => usize::try_from(r.varint()?).map_err(|_| DecodeError::BoundsExceeded)?, SetupVersion::Modern | SetupVersion::Unsupported => return Err(DecodeError::Version), }; r.sub(size) From d1a681bdede2590027c4316a4ff3395522c9b8c5 Mon Sep 17 00:00:00 2001 From: Luke Curley Date: Tue, 29 Sep 2026 07:09:58 -0700 Subject: [PATCH 6/8] perf(net): inline the varint path natively, QUIC read as an if-chain Dropping the VarInt newtype flipped LLVM's inlining of Decoder::varint: the call stayed out of line with its Result going through memory, doubling a varint in a tight loop. Force it inline off wasm32 (size-optimized there), and read the QUIC tag with an if-chain instead of a jump table. Co-Authored-By: Claude Opus 5.5 --- rs/moq-net/src/coding/decode.rs | 3 +- rs/moq-net/src/coding/encode.rs | 3 +- rs/moq-net/src/coding/varint.rs | 54 +++++++++++++++++++-------------- 3 files changed, 35 insertions(+), 25 deletions(-) diff --git a/rs/moq-net/src/coding/decode.rs b/rs/moq-net/src/coding/decode.rs index c5844e202f..d628371748 100644 --- a/rs/moq-net/src/coding/decode.rs +++ b/rs/moq-net/src/coding/decode.rs @@ -172,7 +172,8 @@ impl<'a> Decoder<'a> { } /// Read a varint. - #[inline] + #[cfg_attr(target_arch = "wasm32", inline)] + #[cfg_attr(not(target_arch = "wasm32"), inline(always))] pub fn varint(&mut self) -> Result { let (value, rest) = varint::read(self.buf, self.form)?; self.buf = rest; diff --git a/rs/moq-net/src/coding/encode.rs b/rs/moq-net/src/coding/encode.rs index dd70fb4451..87f323c43d 100644 --- a/rs/moq-net/src/coding/encode.rs +++ b/rs/moq-net/src/coding/encode.rs @@ -92,7 +92,8 @@ impl<'a> Encoder<'a> { } /// Write a varint, or fail with [`EncodeError::BoundsExceeded`] if the form cannot carry it. - #[inline] + #[cfg_attr(target_arch = "wasm32", inline)] + #[cfg_attr(not(target_arch = "wasm32"), inline(always))] pub fn varint(&mut self, v: u64) -> Result<(), EncodeError> { Ok(varint::write(v, self.form, self.buf)?) } diff --git a/rs/moq-net/src/coding/varint.rs b/rs/moq-net/src/coding/varint.rs index 45ef4b7fd8..10368f3a68 100644 --- a/rs/moq-net/src/coding/varint.rs +++ b/rs/moq-net/src/coding/varint.rs @@ -96,7 +96,8 @@ pub(crate) fn size(value: u64, form: Form) -> Result { /// Append the minimal encoding of `value` in `form`. /// /// Fails past [`MAX_QUIC`] in the QUIC form, writing nothing. -#[inline] +#[cfg_attr(target_arch = "wasm32", inline)] +#[cfg_attr(not(target_arch = "wasm32"), inline(always))] pub(super) fn write(value: u64, form: Form, out: &mut Vec) -> Result<(), BoundsExceeded> { match form { Form::Quic => write_quic(value, out), @@ -108,9 +109,15 @@ pub(super) fn write(value: u64, form: Form, out: &mut Vec) -> Result<(), Bou } // Each arm below is a fixed-size write or read, which is what keeps the codec as fast as -// a hand-rolled `put_u16`/`get_u32`. - -#[inline] +// a hand-rolled `put_u16`/`get_u32`. Natively the varint path is `inline(always)` from +// `Encoder::varint`/`Decoder::varint` down: left to LLVM's heuristics it stays a call +// whose `Result` goes through memory, which doubles the cost of a varint +// in a tight loop. wasm32 builds optimize for size, where forcing it grew moq-wasm ~10% +// gzipped, so they keep the heuristics. The QUIC read is an if-chain rather than a `match` +// on the tag: the jump table measured ~35% slower on a mixed-length stream. + +#[cfg_attr(target_arch = "wasm32", inline)] +#[cfg_attr(not(target_arch = "wasm32"), inline(always))] fn write_quic(value: u64, out: &mut Vec) -> Result<(), BoundsExceeded> { let (hi, lo) = to_halves(value); if hi == 0 && lo < 1 << 6 { @@ -129,7 +136,8 @@ fn write_quic(value: u64, out: &mut Vec) -> Result<(), BoundsExceeded> { Ok(()) } -#[inline] +#[cfg_attr(target_arch = "wasm32", inline)] +#[cfg_attr(not(target_arch = "wasm32"), inline(always))] fn write_leading_ones(value: u64, out: &mut Vec) { let (hi, lo) = to_halves(value); let [a, b, c, d] = lo.to_be_bytes(); @@ -156,7 +164,8 @@ fn write_leading_ones(value: u64, out: &mut Vec) { } /// Decode a varint in `form` from the front of `buf`, returning it and the rest of `buf`. -#[inline] +#[cfg_attr(target_arch = "wasm32", inline)] +#[cfg_attr(not(target_arch = "wasm32"), inline(always))] pub(super) fn read(buf: &[u8], form: Form) -> Result<(u64, &[u8]), DecodeError> { match form { Form::Quic => read_quic(buf), @@ -164,31 +173,30 @@ pub(super) fn read(buf: &[u8], form: Form) -> Result<(u64, &[u8]), DecodeError> } } -#[inline] +#[cfg_attr(target_arch = "wasm32", inline)] +#[cfg_attr(not(target_arch = "wasm32"), inline(always))] fn read_quic(buf: &[u8]) -> Result<(u64, &[u8]), DecodeError> { let Some((&first, rest)) = buf.split_first() else { return Err(DecodeError::Short); }; let be = u32::from_be_bytes; - Ok(match first >> 6 { - 0 => (first as u64, rest), - 1 => { - let ([a, b], rest) = buf.split_first_chunk().ok_or(DecodeError::Short)?; - (be([0, 0, a & 0x3f, *b]) as u64, rest) - } - 2 => { - let ([a, b, c, d], rest) = buf.split_first_chunk().ok_or(DecodeError::Short)?; - (be([a & 0x3f, *b, *c, *d]) as u64, rest) - } - _ => { - let ([a, b, c, d, lo @ ..], rest) = buf.split_first_chunk::<8>().ok_or(DecodeError::Short)?; - (from_halves(be([a & 0x3f, *b, *c, *d]), be(*lo)), rest) - } - }) + if first < 0x40 { + Ok((first as u64, rest)) + } else if first < 0x80 { + let ([a, b], rest) = buf.split_first_chunk().ok_or(DecodeError::Short)?; + Ok((be([0, 0, a & 0x3f, *b]) as u64, rest)) + } else if first < 0xc0 { + let ([a, b, c, d], rest) = buf.split_first_chunk().ok_or(DecodeError::Short)?; + Ok((be([a & 0x3f, *b, *c, *d]) as u64, rest)) + } else { + let ([a, b, c, d, lo @ ..], rest) = buf.split_first_chunk::<8>().ok_or(DecodeError::Short)?; + Ok((from_halves(be([a & 0x3f, *b, *c, *d]), be(*lo)), rest)) + } } -#[inline] +#[cfg_attr(target_arch = "wasm32", inline)] +#[cfg_attr(not(target_arch = "wasm32"), inline(always))] fn read_leading_ones(buf: &[u8], seven: bool) -> Result<(u64, &[u8]), DecodeError> { let Some((&first, rest)) = buf.split_first() else { return Err(DecodeError::Short); From 3b70086970293b2d00f9c4b12d1b33e9295009a2 Mon Sep 17 00:00:00 2001 From: Luke Curley Date: Tue, 29 Sep 2026 10:22:13 -0700 Subject: [PATCH 7/8] test(net): drop leftover double negations in IETF codec tests Also removes a stray blank line under translator.md's Required heading. Co-Authored-By: Claude Opus 5.5 --- quest/m1/rs2ts/translator.md | 1 - rs/moq-net/src/ietf/goaway.rs | 2 +- rs/moq-net/src/ietf/properties.rs | 12 ++++++------ 3 files changed, 7 insertions(+), 8 deletions(-) diff --git a/quest/m1/rs2ts/translator.md b/quest/m1/rs2ts/translator.md index c0ba400843..608c28b985 100644 --- a/quest/m1/rs2ts/translator.md +++ b/quest/m1/rs2ts/translator.md @@ -62,5 +62,4 @@ translates. Wire: none. ## Required - - [JS U64](/quest/m1/rs2ts/js-varint.md) - the TypeScript `U64` that Rust `u64` maps to diff --git a/rs/moq-net/src/ietf/goaway.rs b/rs/moq-net/src/ietf/goaway.rs index f132dc2b86..e0df6db16d 100644 --- a/rs/moq-net/src/ietf/goaway.rs +++ b/rs/moq-net/src/ietf/goaway.rs @@ -189,6 +189,6 @@ mod tests { assert_eq!(decoded.new_session_uri, "moqt://relay.example/"); assert_eq!(decoded.timeout, 5000); - assert!(!!bytes.is_empty(), "trailing Request ID should be consumed"); + assert!(bytes.is_empty(), "trailing Request ID should be consumed"); } } diff --git a/rs/moq-net/src/ietf/properties.rs b/rs/moq-net/src/ietf/properties.rs index a4ee99f29a..7514f286f7 100644 --- a/rs/moq-net/src/ietf/properties.rs +++ b/rs/moq-net/src/ietf/properties.rs @@ -194,7 +194,7 @@ mod tests { Encoder::new(&mut buf, Version::Draft17.into()).varint(5000u64).unwrap(); // value let mut bytes = bytes::Bytes::from(buf); crate::coding::decode_buf(&mut bytes, Version::Draft17, Properties::decode).unwrap(); - assert!(!!bytes.is_empty()); + assert!(bytes.is_empty()); } #[test] @@ -206,7 +206,7 @@ mod tests { buf.extend_from_slice(&[0x01, 0x02, 0x03]); // value bytes let mut bytes = bytes::Bytes::from(buf); crate::coding::decode_buf(&mut bytes, Version::Draft17, Properties::decode).unwrap(); - assert!(!!bytes.is_empty()); + assert!(bytes.is_empty()); } #[test] @@ -225,7 +225,7 @@ mod tests { let mut bytes = bytes::Bytes::from(buf); crate::coding::decode_buf(&mut bytes, Version::Draft17, Properties::decode).unwrap(); - assert!(!!bytes.is_empty()); + assert!(bytes.is_empty()); } #[test] @@ -247,7 +247,7 @@ mod tests { crate::coding::decode_buf(&mut bytes, Version::Draft18, Properties::decode).unwrap(), properties ); - assert!(!!bytes.is_empty()); + assert!(bytes.is_empty()); } #[test] @@ -294,7 +294,7 @@ mod tests { let mut bytes = bytes::Bytes::from(buf); let properties = crate::coding::decode_buf(&mut bytes, Version::Draft16, Properties::decode).unwrap(); assert_eq!(properties.group_order, Some(GroupOrder::Descending)); - assert!(!!bytes.is_empty()); + assert!(bytes.is_empty()); } /// The group order property is delta-encoded against the timescale that precedes it, @@ -318,6 +318,6 @@ mod tests { crate::coding::decode_buf(&mut bytes, Version::Draft18, Properties::decode).unwrap(), properties ); - assert!(!!bytes.is_empty()); + assert!(bytes.is_empty()); } } From 0aaab8f30b16c7e982afff348623637d9a077056 Mon Sep 17 00:00:00 2001 From: Luke Curley Date: Tue, 29 Sep 2026 11:23:46 -0700 Subject: [PATCH 8/8] docs(js): point zigzag comments at moq_net::varint Co-Authored-By: Claude Opus 5.5 --- js/net/src/lite/group.ts | 2 +- js/net/src/lite/publisher.ts | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/js/net/src/lite/group.ts b/js/net/src/lite/group.ts index 2bb058051f..dd85e2de6e 100644 --- a/js/net/src/lite/group.ts +++ b/js/net/src/lite/group.ts @@ -64,7 +64,7 @@ export class Group { } } -/** Decode an unsigned zigzag varint back to a signed delta (mirrors Rust `VarInt::to_zigzag`). */ +/** Decode an unsigned zigzag varint back to a signed delta (mirrors Rust `varint::unzigzag`). */ function unzigzag(v: bigint): bigint { return (v >> 1n) ^ -(v & 1n); } diff --git a/js/net/src/lite/publisher.ts b/js/net/src/lite/publisher.ts index ed2a854dd8..26c5821e7a 100644 --- a/js/net/src/lite/publisher.ts +++ b/js/net/src/lite/publisher.ts @@ -42,7 +42,7 @@ const PROBE_MAX_AGE = 10_000; // ms const PROBE_MAX_DELTA = 0.25; const PROBE_RTT_DELTA = 0.25; -/** Map a signed delta to an unsigned zigzag varint value (mirrors Rust `VarInt::from_zigzag`). */ +/** Map a signed delta to an unsigned zigzag varint value (mirrors Rust `varint::zigzag`). */ function zigzag(delta: bigint): bigint { return delta >= 0n ? delta << 1n : (-delta << 1n) - 1n; }