1use alloc::vec;
33use alloc::vec::Vec;
34use core::ops::Range;
35
36use stun_proto::agent::Transmit;
37use stun_types::message::{Message, MessageHeader};
38use tracing::{debug, trace};
39
40use crate::channel::ChannelData;
41
42#[derive(Debug)]
47pub enum IncomingTcp<T: AsRef<[u8]> + core::fmt::Debug> {
48 CompleteMessage(Transmit<T>, Range<usize>),
52 CompleteChannel(Transmit<T>, Range<usize>),
56 StoredMessage(Vec<u8>, Transmit<T>),
58 StoredChannel(Vec<u8>, Transmit<T>),
60}
61
62impl<T: AsRef<[u8]> + core::fmt::Debug> IncomingTcp<T> {
63 pub fn data(&self) -> &[u8] {
65 match self {
66 Self::CompleteMessage(transmit, range) => {
67 &transmit.data.as_ref()[range.start..range.end]
68 }
69 Self::CompleteChannel(transmit, range) => {
70 &transmit.data.as_ref()[range.start..range.end]
71 }
72 Self::StoredMessage(data, _transmit) => data,
73 Self::StoredChannel(data, _transmit) => data,
74 }
75 }
76
77 pub fn message(&self) -> Option<Message<'_>> {
79 if !matches!(
80 self,
81 Self::CompleteMessage(_, _) | Self::StoredMessage(_, _)
82 ) {
83 return None;
84 }
85 Message::from_bytes(self.data()).ok()
86 }
87
88 pub fn channel(&self) -> Option<ChannelData<'_>> {
90 if !matches!(
91 self,
92 Self::CompleteChannel(_, _) | Self::StoredChannel(_, _)
93 ) {
94 return None;
95 }
96 ChannelData::parse(self.data()).ok()
97 }
98}
99
100impl<T: AsRef<[u8]> + core::fmt::Debug> AsRef<[u8]> for IncomingTcp<T> {
101 fn as_ref(&self) -> &[u8] {
102 self.data()
103 }
104}
105
106#[derive(Debug)]
108pub enum StoredTcp {
109 Message(Vec<u8>),
111 Channel(Vec<u8>),
113}
114
115impl StoredTcp {
116 pub fn data(&self) -> &[u8] {
118 match self {
119 Self::Message(data) => data,
120 Self::Channel(data) => data,
121 }
122 }
123
124 fn into_incoming<T: AsRef<[u8]> + core::fmt::Debug>(
125 self,
126 transmit: Transmit<T>,
127 ) -> IncomingTcp<T> {
128 match self {
129 Self::Message(msg) => IncomingTcp::StoredMessage(msg, transmit),
130 Self::Channel(channel) => IncomingTcp::StoredChannel(channel, transmit),
131 }
132 }
133}
134
135impl AsRef<[u8]> for StoredTcp {
136 fn as_ref(&self) -> &[u8] {
137 self.data()
138 }
139}
140
141#[derive(Debug, Default)]
143pub struct TurnTcpBuffer {
144 tcp_buffer: Vec<u8>,
145}
146
147impl TurnTcpBuffer {
148 pub fn new() -> Self {
150 Self { tcp_buffer: vec![] }
151 }
152
153 #[tracing::instrument(
158 level = "trace",
159 skip(self, transmit),
160 fields(
161 transmit.data_len = transmit.data.as_ref().len(),
162 from = ?transmit.from
163 )
164 )]
165 pub fn incoming_tcp<T: AsRef<[u8]> + core::fmt::Debug>(
166 &mut self,
167 transmit: Transmit<T>,
168 ) -> Option<IncomingTcp<T>> {
169 if self.tcp_buffer.is_empty() {
170 let data = transmit.data.as_ref();
171 trace!("Trying to parse incoming data as a complete message/channel");
172 let Ok(hdr) = MessageHeader::from_bytes(data) else {
173 let Ok(channel) = ChannelData::parse(data) else {
174 self.tcp_buffer.extend_from_slice(data);
175 return None;
176 };
177 let channel_len = 4 + channel.data().len();
178 debug!(
179 channel.id = channel.id(),
180 channel.len = channel_len - 4,
181 "Incoming data contains a channel",
182 );
183 if channel_len < data.len() {
184 self.tcp_buffer.extend_from_slice(&data[channel_len..]);
185 }
186 return Some(IncomingTcp::CompleteChannel(transmit, 0..channel_len));
187 };
188 let msg_len = MessageHeader::LENGTH + hdr.data_length() as usize;
189 debug!(
190 msg.transaction = %hdr.transaction_id(),
191 msg.len = msg_len,
192 "Incoming data contains a message",
193 );
194 if data.len() < msg_len {
195 self.tcp_buffer.extend_from_slice(data);
196 return None;
197 }
198 if msg_len < data.len() {
199 self.tcp_buffer.extend_from_slice(&data[msg_len..]);
200 }
201 return Some(IncomingTcp::CompleteMessage(transmit, 0..msg_len));
202 }
203
204 self.tcp_buffer.extend_from_slice(transmit.data.as_ref());
205 self.poll_recv().map(|recv| recv.into_incoming(transmit))
206 }
207
208 #[tracing::instrument(
210 level = "trace",
211 skip(self),
212 fields(
213 buffered_len = self.tcp_buffer.len(),
214 )
215 )]
216 pub fn poll_recv(&mut self) -> Option<StoredTcp> {
217 let Ok(hdr) = MessageHeader::from_bytes(&self.tcp_buffer) else {
218 let Ok((id, channel_data_len)) = ChannelData::parse_header(&self.tcp_buffer) else {
219 trace!(
220 buffered.len = self.tcp_buffer.len(),
221 "cannot parse stored data"
222 );
223 return None;
224 };
225 let channel_len = 4 + channel_data_len;
226 if self.tcp_buffer.len() < channel_len {
227 trace!(
228 buffered.len = self.tcp_buffer.len(),
229 required = channel_len,
230 "need more bytes to complete channel data"
231 );
232 return None;
233 }
234 let (data, remaining) = self.tcp_buffer.split_at(channel_len);
235 let data_binding = data.to_vec();
236 debug!(
237 channel.id = id,
238 channel.len = channel_data_len,
239 remaining = remaining.len(),
240 "buffered data contains a channel",
241 );
242 self.tcp_buffer = remaining.to_vec();
243 return Some(StoredTcp::Channel(data_binding));
244 };
245 let msg_len = MessageHeader::LENGTH + hdr.data_length() as usize;
246 if self.tcp_buffer.len() < msg_len {
247 trace!(
248 buffered.len = self.tcp_buffer.len(),
249 required = msg_len,
250 "need more bytes to complete STUN message"
251 );
252 return None;
253 }
254 let (data, remaining) = self.tcp_buffer.split_at(msg_len);
255 let data_binding = data.to_vec();
256 debug!(
257 msg.transaction = %hdr.transaction_id(),
258 msg.len = msg_len,
259 remaining = remaining.len(),
260 "stored data contains a message",
261 );
262 self.tcp_buffer = remaining.to_vec();
263 Some(StoredTcp::Message(data_binding))
264 }
265
266 pub fn into_inner(self) -> Vec<u8> {
268 self.tcp_buffer
269 }
270
271 pub fn len(&self) -> usize {
273 self.tcp_buffer.len()
274 }
275
276 pub fn is_empty(&self) -> bool {
278 self.tcp_buffer.is_empty()
279 }
280}
281
282#[cfg(test)]
283mod tests {
284 use core::net::SocketAddr;
285
286 use stun_types::{
287 attribute::Software,
288 message::{Message, MessageWriteVec},
289 prelude::{MessageWrite, MessageWriteExt},
290 TransportType,
291 };
292 use tracing::info;
293
294 use crate::message::ALLOCATE;
295
296 use super::*;
297
298 fn generate_addresses() -> (SocketAddr, SocketAddr) {
299 (
300 "192.168.0.1:1000".parse().unwrap(),
301 "10.0.0.2:2000".parse().unwrap(),
302 )
303 }
304
305 fn generate_message() -> Vec<u8> {
306 let mut msg = Message::builder_request(ALLOCATE, MessageWriteVec::new());
307 msg.add_attribute(&Software::new("turn-types").unwrap())
308 .unwrap();
309 msg.add_fingerprint().unwrap();
310 msg.finish()
311 }
312
313 fn generate_message_in_channel() -> Vec<u8> {
314 let msg = generate_message();
315 let channel = ChannelData::new(0x4000, &msg);
316 let mut out = vec![0; msg.len() + 4];
317 channel.write_into_unchecked(&mut out);
318 out
319 }
320
321 #[test]
322 fn test_incoming_tcp_complete_message() {
323 let _init = crate::tests::test_init_log();
324 let (local_addr, remote_addr) = generate_addresses();
325 let mut tcp = TurnTcpBuffer::new();
326 let msg = generate_message();
327 let ret = tcp
328 .incoming_tcp(Transmit::new(
329 msg.clone(),
330 TransportType::Tcp,
331 remote_addr,
332 local_addr,
333 ))
334 .unwrap();
335 assert!(matches!(ret, IncomingTcp::CompleteMessage(_, _)));
336 assert_eq!(ret.data(), &msg);
337 assert_eq!(ret.as_ref(), &msg);
338 assert!(ret.message().is_some());
339 assert!(tcp.is_empty());
340 assert_eq!(tcp.len(), 0);
341 assert!(tcp.into_inner().is_empty());
342 }
343
344 #[test]
345 fn test_incoming_tcp_complete_message_in_channel() {
346 let _init = crate::tests::test_init_log();
347 let (local_addr, remote_addr) = generate_addresses();
348 let mut tcp = TurnTcpBuffer::new();
349 let msg = generate_message_in_channel();
350 let ret = tcp
351 .incoming_tcp(Transmit::new(
352 msg.clone(),
353 TransportType::Tcp,
354 remote_addr,
355 local_addr,
356 ))
357 .unwrap();
358 assert!(matches!(ret, IncomingTcp::CompleteChannel(_, _)));
359 assert_eq!(ret.data(), &msg);
360 assert_eq!(ret.as_ref(), &msg);
361 assert!(ret.channel().is_some());
362 assert!(tcp.is_empty());
363 assert_eq!(tcp.len(), 0);
364 assert!(tcp.into_inner().is_empty());
365 }
366
367 #[test]
368 fn test_incoming_tcp_partial_message() {
369 let _init = crate::tests::test_init_log();
370 let (local_addr, remote_addr) = generate_addresses();
371 let mut tcp = TurnTcpBuffer::new();
372 let msg = generate_message();
373 info!("message: {msg:x?}");
374 for i in 1..msg.len() {
375 let ret = tcp.incoming_tcp(Transmit::new(
376 &msg[i - 1..i],
377 TransportType::Tcp,
378 remote_addr,
379 local_addr,
380 ));
381 assert!(ret.is_none());
382
383 let data = tcp.into_inner();
384 assert_eq!(&data, &msg[..i]);
385 tcp = TurnTcpBuffer::new();
386 let ret = tcp.incoming_tcp(Transmit::new(
387 &data,
388 TransportType::Tcp,
389 remote_addr,
390 local_addr,
391 ));
392 assert!(ret.is_none());
393 assert!(!tcp.is_empty());
394 assert_eq!(tcp.len(), i);
395 }
396 let ret = tcp
397 .incoming_tcp(Transmit::new(
398 &msg[msg.len() - 1..],
399 TransportType::Tcp,
400 remote_addr,
401 local_addr,
402 ))
403 .unwrap();
404 assert_eq!(ret.data(), &msg);
405 assert_eq!(ret.as_ref(), &msg);
406 assert!(ret.message().is_some());
407 let IncomingTcp::StoredMessage(produced, _) = ret else {
408 unreachable!();
409 };
410 assert_eq!(produced, msg);
411 assert!(tcp.is_empty());
412 assert_eq!(tcp.len(), 0);
413 assert!(tcp.into_inner().is_empty());
414 }
415
416 #[test]
417 fn test_incoming_tcp_partial_channel() {
418 let _init = crate::tests::test_init_log();
419 let (local_addr, remote_addr) = generate_addresses();
420 let mut tcp = TurnTcpBuffer::new();
421 let channel = generate_message_in_channel();
422 info!("message: {channel:x?}");
423 for i in 1..channel.len() {
424 let ret = tcp.incoming_tcp(Transmit::new(
425 &channel[i - 1..i],
426 TransportType::Tcp,
427 remote_addr,
428 local_addr,
429 ));
430 assert!(ret.is_none());
431
432 let data = tcp.into_inner();
433 assert_eq!(&data, &channel[..i]);
434 tcp = TurnTcpBuffer::new();
435 let ret = tcp.incoming_tcp(Transmit::new(
436 &data,
437 TransportType::Tcp,
438 remote_addr,
439 local_addr,
440 ));
441 assert!(ret.is_none());
442 assert!(!tcp.is_empty());
443 assert_eq!(tcp.len(), i);
444 }
445 let ret = tcp
446 .incoming_tcp(Transmit::new(
447 &channel[channel.len() - 1..],
448 TransportType::Tcp,
449 remote_addr,
450 local_addr,
451 ))
452 .unwrap();
453 assert_eq!(ret.data(), &channel);
454 assert_eq!(ret.as_ref(), &channel);
455 assert!(ret.channel().is_some());
456 let IncomingTcp::StoredChannel(produced, _) = ret else {
457 unreachable!()
458 };
459 assert_eq!(produced, channel);
460 assert!(tcp.into_inner().is_empty());
461 }
462
463 #[test]
464 fn test_incoming_tcp_message_and_channel() {
465 let _init = crate::tests::test_init_log();
466 let (local_addr, remote_addr) = generate_addresses();
467 let mut tcp = TurnTcpBuffer::new();
468 let msg = generate_message();
469 let channel = generate_message_in_channel();
470 let mut input = msg.clone();
471 input.extend_from_slice(&channel);
472 let ret = tcp
473 .incoming_tcp(Transmit::new(
474 input.clone(),
475 TransportType::Tcp,
476 remote_addr,
477 local_addr,
478 ))
479 .unwrap();
480 assert_eq!(ret.data(), &msg);
481 assert_eq!(ret.as_ref(), &msg);
482 assert!(ret.message().is_some());
483 let IncomingTcp::CompleteMessage(transmit, msg_range) = ret else {
484 unreachable!();
485 };
486 assert_eq!(msg_range, 0..msg.len());
487 assert_eq!(transmit.data, input);
488 let ret = tcp.poll_recv().unwrap();
489 assert_eq!(ret.data(), &channel);
490 assert_eq!(ret.as_ref(), &channel);
491 let StoredTcp::Channel(produced) = ret else {
492 unreachable!()
493 };
494 assert_eq!(produced, channel);
495 }
496
497 #[test]
498 fn test_incoming_tcp_channel_and_message() {
499 let _init = crate::tests::test_init_log();
500 let (local_addr, remote_addr) = generate_addresses();
501 let mut tcp = TurnTcpBuffer::new();
502 let msg = generate_message();
503 let channel = generate_message_in_channel();
504 let mut input = channel.clone();
505 input.extend_from_slice(&msg);
506 let ret = tcp
507 .incoming_tcp(Transmit::new(
508 input.clone(),
509 TransportType::Tcp,
510 remote_addr,
511 local_addr,
512 ))
513 .unwrap();
514 assert_eq!(ret.data(), &channel);
515 assert_eq!(ret.as_ref(), &channel);
516 assert!(ret.channel().is_some());
517 let IncomingTcp::CompleteChannel(transmit, channel_range) = ret else {
518 unreachable!()
519 };
520 assert_eq!(channel_range, 0..channel.len());
521 assert_eq!(transmit.data, input);
522 let ret = tcp.poll_recv().unwrap();
523 assert_eq!(ret.data(), &msg);
524 assert_eq!(ret.as_ref(), &msg);
525 let StoredTcp::Message(produced) = ret else {
526 unreachable!()
527 };
528 assert_eq!(produced, msg);
529 }
530}