Skip to main content

wind_tuic/proto/v5/
addr.rs

1use std::{
2	net::{Ipv4Addr, Ipv6Addr},
3	str,
4};
5
6use bytes::{Buf, BufMut};
7use num_enum::{FromPrimitive, IntoPrimitive};
8use snafu::{ResultExt, ensure};
9use tokio_util::codec::{Decoder, Encoder};
10use wind_core::types::TargetAddr;
11
12use crate::{
13	ProtoError,
14	proto::{BytesRemainingSnafu, DomainTooLongSnafu, FailParseDomainSnafu, UnknownAddressTypeSnafu},
15};
16
17//-----------------------------------------------------------------------------
18// Type Definitions
19//-----------------------------------------------------------------------------
20
21/// Codec for TUIC address encoding and decoding
22#[derive(Debug, Clone, Copy)]
23pub struct AddressCodec;
24
25/// TUIC Address representation
26#[derive(Debug, Clone, PartialEq)]
27pub enum Address {
28	/// No address
29	None,
30	/// Domain name and port
31	Domain(String, u16),
32	/// IPv4 address and port
33	IPv4(Ipv4Addr, u16),
34	/// IPv6 address and port
35	IPv6(Ipv6Addr, u16),
36}
37
38/// Address type indicators as defined in TUIC protocol
39#[derive(IntoPrimitive, FromPrimitive, Copy, Clone, Debug, PartialEq)]
40#[repr(u8)]
41pub enum AddressType {
42	None   = u8::MAX,
43	Domain = 0,
44	IPv4   = 1,
45	IPv6   = 2,
46	#[num_enum(catch_all)]
47	Other(u8),
48}
49
50//-----------------------------------------------------------------------------
51// Implementations
52//-----------------------------------------------------------------------------
53
54impl From<TargetAddr> for Address {
55	fn from(value: TargetAddr) -> Self {
56		match value {
57			TargetAddr::Domain(s, port) => Self::Domain(s, port),
58			TargetAddr::IPv4(addr, port) => Self::IPv4(addr, port),
59			TargetAddr::IPv6(addr, port) => Self::IPv6(addr, port),
60		}
61	}
62}
63
64//-----------------------------------------------------------------------------
65// Codec Implementation
66//-----------------------------------------------------------------------------
67
68/// Implementation according to TUIC specification:
69/// https://github.com/proxy-rs/wind/blob/main/crates/wind-tuic/SPEC.md#6-address-encoding
70#[cfg(feature = "decode")]
71impl Decoder for AddressCodec {
72	type Error = ProtoError;
73	type Item = Address;
74
75	fn decode(&mut self, src: &mut bytes::BytesMut) -> Result<Option<Self::Item>, Self::Error> {
76		// Return None if buffer is empty
77		if src.is_empty() {
78			return Ok(None);
79		}
80
81		// Parse address type from first byte
82		let addr_type = AddressType::from(src[0]);
83
84		ensure!(!matches!(addr_type, AddressType::Other(_)), UnknownAddressTypeSnafu { value: u8::from(addr_type) });
85
86		match addr_type {
87			AddressType::None => {
88				src.advance(1);
89				Ok(Some(Address::None))
90			}
91			AddressType::IPv4 => {
92				// Type (1) + IPv4 (4) + Port (2)
93				if src.len() < 1 + 4 + 2 {
94					return Ok(None);
95				}
96				src.advance(1);
97				let mut octets = [0; 4];
98				src.copy_to_slice(&mut octets);
99				let ip = Ipv4Addr::from(octets);
100				let port = src.get_u16();
101				Ok(Some(Address::IPv4(ip, port)))
102			}
103			AddressType::IPv6 => {
104				// Type (1) + IPv6 (16) + Port (2)
105				if src.len() < 1 + 16 + 2 {
106					return Ok(None);
107				}
108				src.advance(1);
109				let mut octets = [0; 16];
110				src.copy_to_slice(&mut octets);
111				let ip = Ipv6Addr::from(octets);
112				let port = src.get_u16();
113				Ok(Some(Address::IPv6(ip, port)))
114			}
115			AddressType::Domain => {
116				// Need at least type byte and length byte
117				if src.len() < 1 + 1 {
118					return Ok(None);
119				}
120				let domain_len = src[1] as usize;
121
122				// Type (1) + Length (1) + Domain + Port (2)
123				if src.len() < 1 + 1 + domain_len + 2 {
124					return Ok(None);
125				}
126				src.advance(2);
127
128				let domain = &src[..domain_len];
129				let domain = str::from_utf8(domain)
130					.context(FailParseDomainSnafu {
131						raw: hex::encode(domain),
132					})?
133					.to_string();
134				src.advance(domain_len);
135				let port = src.get_u16();
136				Ok(Some(Address::Domain(domain, port)))
137			}
138			_ => unreachable!(),
139		}
140	}
141
142	fn decode_eof(&mut self, buf: &mut bytes::BytesMut) -> Result<Option<Self::Item>, Self::Error> {
143		match self.decode(buf) {
144			Ok(None) => BytesRemainingSnafu.fail(),
145			v => v,
146		}
147	}
148}
149
150#[cfg(feature = "encode")]
151impl Encoder<Address> for AddressCodec {
152	type Error = ProtoError;
153
154	fn encode(&mut self, item: Address, dst: &mut bytes::BytesMut) -> Result<(), Self::Error> {
155		match item {
156			Address::None => {
157				dst.reserve(1);
158				dst.put_u8(AddressType::None.into());
159			}
160			Address::IPv4(ip, port) => {
161				// Type (1) + IPv4 (4) + Port (2)
162				dst.reserve(1 + 4 + 2);
163				dst.put_u8(AddressType::IPv4.into());
164				dst.put_slice(&ip.octets());
165				dst.put_u16(port);
166			}
167			Address::IPv6(ip, port) => {
168				// Type (1) + IPv6 (16) + Port (2)
169				dst.reserve(1 + 16 + 2);
170				dst.put_u8(AddressType::IPv6.into());
171				dst.put_slice(&ip.octets());
172				dst.put_u16(port);
173			}
174			Address::Domain(domain, port) => {
175				// Validate domain length
176				if domain.len() > u8::MAX as usize {
177					return DomainTooLongSnafu { domain }.fail();
178				}
179
180				// Type (1) + Length (1) + Domain + Port (2)
181				dst.reserve(1 + 1 + domain.len() + 2);
182				dst.put_u8(AddressType::Domain.into());
183				dst.put_u8(domain.len() as u8);
184				dst.put_slice(domain.as_bytes());
185				dst.put_u16(port);
186			}
187		}
188		Ok(())
189	}
190}
191
192//-----------------------------------------------------------------------------
193// Tests
194//-----------------------------------------------------------------------------
195
196#[cfg(test)]
197mod test {
198	use std::net::{Ipv4Addr, Ipv6Addr};
199
200	use futures_util::SinkExt as _;
201	use tokio_stream::StreamExt as _;
202	use tokio_util::codec::{FramedRead, FramedWrite};
203
204	use super::{Address, AddressCodec};
205	use crate::proto::ProtoError;
206
207	/// Test complete encoding and decoding cycle for all address types
208	#[tokio::test]
209	async fn test_addr_1() -> eyre::Result<()> {
210		let buffer = Vec::with_capacity(128);
211		let vars = vec![
212			Address::None,
213			Address::IPv4(Ipv4Addr::LOCALHOST, 80),
214			Address::IPv6(Ipv6Addr::UNSPECIFIED, 12),
215			Address::Domain(String::from("www.google.com"), 443),
216		];
217
218		// Test encoding
219		let mut writer = FramedWrite::new(buffer, AddressCodec);
220		let mut expect_len = 0;
221		for var in &vars {
222			match var {
223				Address::None => expect_len = expect_len + 1,
224				Address::Domain(domain, _) => expect_len = expect_len + 1 + 1 + domain.len() + 2,
225				Address::IPv4(..) => expect_len = expect_len + 1 + 4 + 2,
226				Address::IPv6(..) => expect_len = expect_len + 1 + 16 + 2,
227			}
228			writer.send(var.clone()).await?;
229			assert_eq!(writer.get_ref().len(), expect_len);
230		}
231
232		// Test decoding
233		let buffer = writer.get_ref();
234		let mut reader = FramedRead::new(buffer.as_slice(), AddressCodec);
235		for var in vars {
236			let frame = reader.next().await.unwrap()?;
237			assert_eq!(var, frame);
238		}
239		Ok(())
240	}
241
242	/// Test behavior with partial data (simulating streaming data arrival)
243	#[tokio::test]
244	async fn test_addr_2() -> eyre::Result<()> {
245		let vars = vec![
246			Address::IPv4(Ipv4Addr::LOCALHOST, 80),
247			Address::IPv6(Ipv6Addr::UNSPECIFIED, 12),
248			Address::Domain(String::from("www.google.com"), 443),
249		];
250
251		for addr in vars {
252			// Encode the address
253			let buffer = Vec::with_capacity(128);
254			let mut writer = FramedWrite::new(buffer, AddressCodec);
255			writer.send(addr.clone()).await?;
256			let mut buffer = writer.into_inner();
257
258			// Split the encoded data in half to simulate partial data arrival
259			let full_len = buffer.len();
260			let mut half_b = buffer.split_off(full_len / 2 as usize);
261			let mut half_a = buffer;
262
263			// First half should result in BytesRemaining error
264			{
265				let mut reader = FramedRead::new(half_a.as_slice(), AddressCodec);
266				assert!(matches!(
267					reader.next().await.unwrap().unwrap_err(),
268					ProtoError::BytesRemaining
269				));
270			}
271
272			// Recombined buffer should decode properly
273			half_a.append(&mut half_b);
274			let mut reader = FramedRead::new(half_a.as_slice(), AddressCodec);
275			assert_eq!(reader.next().await.unwrap()?, addr);
276		}
277
278		Ok(())
279	}
280
281	/// Test to generate and inspect hex encoding (useful for debugging)
282	#[tokio::test]
283	async fn hex_check() -> eyre::Result<()> {
284		let mut buffer = Vec::new();
285		let vars = vec![
286			Address::None,
287			Address::IPv4(Ipv4Addr::LOCALHOST, 80),
288			Address::IPv6(Ipv6Addr::LOCALHOST, 12),
289			Address::Domain(String::from("www.google.com"), 443),
290		];
291
292		// Encode the second address and print its hex representation
293		FramedWrite::new(&mut buffer, AddressCodec).send(vars[1].clone()).await?;
294		println!("{}", hex::encode(buffer));
295		Ok(())
296	}
297}