Skip to main content

turn_types/
tcp.rs

1// Copyright (C) 2025 Matthew Waters <matthew@centricular.com>
2//
3// Licensed under the Apache License, Version 2.0 <LICENSE-APACHE or
4// http://www.apache.org/licenses/LICENSE-2.0> or the MIT license
5// <LICENSE-MIT or http://opensource.org/licenses/MIT>, at your
6// option. This file may not be copied, modified, or distributed
7// except according to those terms.
8//
9// SPDX-License-Identifier: MIT OR Apache-2.0
10
11//! Handle TURN over TCP.
12//!
13//! A TURN connection between a client and a server can have two types of data:
14//! - STUN [`Message`]s, and
15//! - [`ChannelData`]s
16//!
17//! Unlike a UDP connection which inherently contains a size for every message, TCP is a
18//! stream-based protocol and the size of a message must be infered from the contained data. This
19//! module performs the relevant buffering of incoming data over a TCP connection and produces
20//! [`Message`]s or [`ChannelData`] as they are completely received.
21//!
22//! The buffering performed by [`TurnTcpBuffer`] is only applicable for the TCP connection between
23//! the TURN client and the TURN server when using a UDP allocation. Use of TURN-TCP ([RFC6062])
24//! requires the TURN client to connect to the TURN server using a TCP connection (optionally with
25//! TLS) and data on the separate data TCP connection is forwarded as-is. The control connection for
26//! TURN-TCP requires buffering of only STUN Messages without any framing and can also be performed
27//! by [`TurnTcpBuffer`] if any [`ChannelData`] messages received are considered fatal TURN protocol
28//! errors.
29//!
30//! [RFC6062]: https://tools.ietf.org/html/rfc6062
31
32use 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/// Reply to [`TurnTcpBuffer::incoming_tcp()`].
43///
44/// The `Transmit<T>` in each value  is always the original value passed to
45/// [`TurnTcpBuffer::incoming_tcp()`].
46#[derive(Debug)]
47pub enum IncomingTcp<T: AsRef<[u8]> + core::fmt::Debug> {
48    /// Input data (with the provided range) contains a complete STUN Message.
49    ///
50    /// Any extra data after the range is stored for later processing.
51    CompleteMessage(Transmit<T>, Range<usize>),
52    /// Input data (with the provided range) contains a complete Channel data message.
53    ///
54    /// Any extra data after the range is stored for later processing.
55    CompleteChannel(Transmit<T>, Range<usize>),
56    /// A STUN message has been produced from the buffered data.
57    StoredMessage(Vec<u8>, Transmit<T>),
58    /// A Channel data message has been produced from the buffered data.
59    StoredChannel(Vec<u8>, Transmit<T>),
60}
61
62impl<T: AsRef<[u8]> + core::fmt::Debug> IncomingTcp<T> {
63    /// The byte slice for this incoming, or stored message, or channel.
64    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    /// The [`Message`] contained in this incoming or stored data.
78    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    /// The [`ChannelData`] contained in this incoming or stored data.
89    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/// A stored [`Message`] or [`ChannelData`]
107#[derive(Debug)]
108pub enum StoredTcp {
109    /// A STUN [`Message`] has been received.
110    Message(Vec<u8>),
111    /// A [`ChannelData`] has been received.
112    Channel(Vec<u8>),
113}
114
115impl StoredTcp {
116    /// The byte slice for this stored data.
117    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/// A TCP buffer for TURN messages.
142#[derive(Debug, Default)]
143pub struct TurnTcpBuffer {
144    tcp_buffer: Vec<u8>,
145}
146
147impl TurnTcpBuffer {
148    /// Construct a new [`TurnTcpBuffer`].
149    pub fn new() -> Self {
150        Self { tcp_buffer: vec![] }
151    }
152
153    /// Provide incoming TCP data to parse.
154    ///
155    /// A return value of `None` indicates that the more data is required to provide a complete
156    /// STUN [`Message`] or a [`ChannelData`].
157    #[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    /// Return the next complete message (if any).
209    #[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    /// Returns the underlying buffer.
267    pub fn into_inner(self) -> Vec<u8> {
268        self.tcp_buffer
269    }
270
271    /// The number of bytes contained in this buffer.
272    pub fn len(&self) -> usize {
273        self.tcp_buffer.len()
274    }
275
276    /// Whether the buffer currently contains 0 bytes of data.
277    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}