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#[derive(Debug, Clone, Copy)]
23pub struct AddressCodec;
24
25#[derive(Debug, Clone, PartialEq)]
27pub enum Address {
28 None,
30 Domain(String, u16),
32 IPv4(Ipv4Addr, u16),
34 IPv6(Ipv6Addr, u16),
36}
37
38#[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
50impl 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#[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 if src.is_empty() {
78 return Ok(None);
79 }
80
81 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 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 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 if src.len() < 1 + 1 {
118 return Ok(None);
119 }
120 let domain_len = src[1] as usize;
121
122 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 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 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 if domain.len() > u8::MAX as usize {
177 return DomainTooLongSnafu { domain }.fail();
178 }
179
180 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#[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 #[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 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 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 #[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 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 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 {
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 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 #[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 FramedWrite::new(&mut buffer, AddressCodec).send(vars[1].clone()).await?;
294 println!("{}", hex::encode(buffer));
295 Ok(())
296 }
297}