1use std::net::{IpAddr, SocketAddr};
8use std::sync::Arc;
9use std::time::Duration;
10
11use bytes::{Bytes, BytesMut};
12use quinn::{Connection, Endpoint, RecvStream, SendStream};
13use sipx_sip::{Limits, Message, parse_datagram};
14use tokio::sync::mpsc;
15use tokio::task::JoinSet;
16
17use crate::error::Result;
18use crate::target::{ConnectionKey, TransportKind};
19use crate::tcp::Event;
20
21#[derive(Debug, thiserror::Error)]
23#[non_exhaustive]
24pub enum QuicError {
25 #[error("certificate for {peer} has the wrong host: {detail}")]
27 WrongHost {
28 peer: String,
30 detail: String,
32 },
33 #[error("certificate for {peer} has an unknown issuer: {detail}")]
35 UnknownIssuer {
36 peer: String,
38 detail: String,
40 },
41 #[error("QUIC peer {peer} negotiated the wrong ALPN: {detail}")]
43 WrongAlpn {
44 peer: String,
46 detail: String,
48 },
49 #[error("QUIC connection to {peer} closed: {detail}")]
51 ConnectionClosed {
52 peer: String,
54 detail: String,
56 },
57 #[error("QUIC handshake with {peer}: {detail}")]
59 Handshake {
60 peer: String,
62 detail: String,
64 },
65}
66
67impl QuicError {
68 pub(crate) fn handshake(peer: String, detail: String) -> Self {
69 let folded = detail.to_ascii_lowercase();
70 if folded.contains("notvalidforname") || folded.contains("not valid for") {
71 Self::WrongHost { peer, detail }
72 } else if folded.contains("unknownissuer") || folded.contains("unknown issuer") {
73 Self::UnknownIssuer { peer, detail }
74 } else if folded.contains("application protocol")
75 || folded.contains("known protocol")
76 || folded.contains("alpn")
77 {
78 Self::WrongAlpn { peer, detail }
79 } else {
80 Self::Handshake { peer, detail }
81 }
82 }
83}
84
85pub(crate) type Reply = mpsc::Sender<Bytes>;
87
88const KEEPALIVE: Duration = Duration::from_secs(25);
89
90pub(crate) fn transport_config() -> Arc<quinn::TransportConfig> {
91 let mut transport = quinn::TransportConfig::default();
92 transport.keep_alive_interval(Some(KEEPALIVE));
93 transport.max_concurrent_bidi_streams(64_u8.into());
94 transport.max_concurrent_uni_streams(0_u8.into());
95 Arc::new(transport)
96}
97
98pub(crate) fn endpoint(
100 ip: IpAddr,
101 port: u16,
102 client: Option<&crate::tls::ClientTls>,
103 server: Option<&crate::tls::ServerTls>,
104) -> Result<Endpoint> {
105 let bind = SocketAddr::new(ip, port);
106 let mut endpoint = match server {
107 Some(server) => {
108 let mut config = server.quic_config()?;
109 config.transport_config(transport_config());
110 Endpoint::server(config, bind)?
111 }
112 None => Endpoint::client(bind)?,
113 };
114 if let Some(client) = client {
115 let mut config = client.quic_config()?;
116 config.transport_config(transport_config());
117 endpoint.set_default_client_config(config);
118 }
119 Ok(endpoint)
120}
121
122pub(crate) async fn pump(
124 connection: Connection,
125 key: ConnectionKey,
126 id: u64,
127 mut outgoing: mpsc::Receiver<Bytes>,
128 events: mpsc::Sender<Event>,
129 limits: Limits,
130) {
131 let mut streams = JoinSet::new();
132 let detail = loop {
133 tokio::select! {
134 incoming = connection.accept_bi() => match incoming {
135 Ok((send, recv)) => {
136 let events = events.clone();
137 let key = key.clone();
138 let connection = connection.clone();
139 streams.spawn(async move {
140 receive_request(connection, send, recv, key, id, events, limits).await;
141 });
142 }
143 Err(error) => break error.to_string(),
144 },
145 Some(bytes) = outgoing.recv() => {
146 let connection = connection.clone();
147 let events = events.clone();
148 let key = key.clone();
149 streams.spawn(async move {
150 send_request(connection, bytes, key, id, events, limits).await;
151 });
152 }
153 completed = streams.join_next(), if !streams.is_empty() => {
154 if let Some(Err(error)) = completed {
155 tracing::debug!(%error, peer = %key.peer, "QUIC stream task failed");
156 }
157 }
158 error = connection.closed() => break error.to_string(),
159 else => break "connection task ended".to_owned(),
160 }
161 };
162 streams.abort_all();
163 while streams.join_next().await.is_some() {}
164 let _ = events.send(Event::QuicClosed { key, id, detail }).await;
166}
167
168async fn receive_request(
169 connection: Connection,
170 send: SendStream,
171 mut recv: RecvStream,
172 key: ConnectionKey,
173 id: u64,
174 events: mpsc::Sender<Event>,
175 limits: Limits,
176) {
177 let peer = key.peer;
178 let bytes = match recv.read_to_end(limits.max_message_bytes).await {
179 Ok(bytes) => Bytes::from(bytes),
180 Err(quinn::ReadToEndError::Read(quinn::ReadError::ConnectionLost(error))) => {
181 tracing::debug!(%error, %peer, "QUIC connection closed while reading a request");
182 return;
183 }
184 Err(error) => {
185 tracing::debug!(%error, %peer, "QUIC request stream could not be read");
186 let _ = events.send(Event::FramingFailed { key }).await;
188 connection.close(1_u8.into(), b"malformed SIP stream");
189 return;
190 }
191 };
192 let message = match parse_one(bytes, &limits) {
193 Ok(message @ Message::Request(_)) => message,
194 Ok(Message::Response(_)) => {
195 tracing::debug!(%peer, "QUIC peer opened a stream with a response");
196 let _ = events.send(Event::FramingFailed { key }).await;
198 connection.close(1_u8.into(), b"response opened a request stream");
199 return;
200 }
201 Err(error) => {
202 tracing::debug!(%error, %peer, "malformed QUIC request stream");
204 let _ = events.send(Event::FramingFailed { key }).await;
206 connection.close(1_u8.into(), b"malformed SIP stream");
207 return;
208 }
209 };
210 let (reply, replies) = mpsc::channel(16);
211 if events
212 .send(Event::Message {
213 message: Box::new(message),
214 source: peer,
215 transport: TransportKind::Quic,
216 id,
217 quic_reply: Some(reply),
218 })
219 .await
220 .is_ok()
221 {
222 write_responses(send, replies).await;
223 }
224}
225
226async fn write_responses(mut send: SendStream, mut replies: mpsc::Receiver<Bytes>) {
227 while let Some(bytes) = replies.recv().await {
228 let final_response = parse_datagram(bytes.clone(), &Limits::stream())
229 .ok()
230 .and_then(|message| match message {
231 Message::Response(response) => Some(response.status.is_final()),
232 Message::Request(_) => None,
233 })
234 .unwrap_or(true);
235 if send.write_all(&bytes).await.is_err() {
236 return;
237 }
238 if final_response {
239 let _ = send.finish();
241 return;
242 }
243 }
244 let _ = send.finish();
246}
247
248async fn send_request(
249 connection: Connection,
250 bytes: Bytes,
251 key: ConnectionKey,
252 id: u64,
253 events: mpsc::Sender<Event>,
254 limits: Limits,
255) {
256 let peer = key.peer;
257 let Some(mut recv) = write_request(&connection, &events, &key, bytes).await else {
258 return;
259 };
260
261 let mut pending = BytesMut::new();
262 let mut final_response = None;
263 loop {
264 let chunk = match recv.read_chunk(usize::MAX, true).await {
265 Ok(Some(chunk)) => chunk.bytes,
266 Ok(None) => break,
267 Err(quinn::ReadError::ConnectionLost(error)) => {
268 tracing::debug!(%error, %peer, "QUIC connection closed while reading a response");
269 return;
270 }
271 Err(error) => {
272 reject_response_stream(&connection, &events, &key, peer, error.to_string()).await;
273 return;
274 }
275 };
276 if final_response.is_some() {
277 reject_response_stream(
278 &connection,
279 &events,
280 &key,
281 peer,
282 "bytes followed the final response".to_owned(),
283 )
284 .await;
285 return;
286 }
287 pending.extend_from_slice(&chunk);
288 loop {
289 let message = match take_response(&mut pending, false, &limits) {
290 Ok(Some(message)) => message,
291 Ok(None) => break,
292 Err(error) => {
293 reject_response_stream(&connection, &events, &key, peer, error).await;
294 return;
295 }
296 };
297 let is_final = match &message {
298 Message::Response(response) => response.status.is_final(),
299 Message::Request(_) => {
300 reject_response_stream(
301 &connection,
302 &events,
303 &key,
304 peer,
305 "request on response stream".to_owned(),
306 )
307 .await;
308 return;
309 }
310 };
311 if is_final {
312 final_response = Some(message);
313 if !pending.is_empty() {
314 reject_response_stream(
315 &connection,
316 &events,
317 &key,
318 peer,
319 "bytes followed the final response".to_owned(),
320 )
321 .await;
322 return;
323 }
324 break;
325 }
326 if send_response_event(&events, message, peer, id)
327 .await
328 .is_err()
329 {
330 return;
331 }
332 }
333 }
334
335 let message = match complete_response(final_response, &mut pending, &limits) {
336 Ok(message) => message,
337 Err(error) => {
338 reject_response_stream(&connection, &events, &key, peer, error).await;
339 return;
340 }
341 };
342 let _ = send_response_event(&events, message, peer, id).await;
344}
345
346async fn write_request(
347 connection: &Connection,
348 events: &mpsc::Sender<Event>,
349 key: &ConnectionKey,
350 bytes: Bytes,
351) -> Option<RecvStream> {
352 let peer = key.peer;
353 let (mut send, recv) = match connection.open_bi().await {
354 Ok(stream) => stream,
355 Err(error) => {
356 tracing::debug!(%error, %peer, "QUIC stream could not be opened");
357 return None;
358 }
359 };
360 if let Err(error) = send.write_all(&bytes).await {
361 match error {
362 quinn::WriteError::ConnectionLost(error) => {
363 tracing::debug!(%error, %peer, "QUIC connection closed while writing a request");
364 }
365 other => {
366 reject_response_stream(connection, events, key, peer, other.to_string()).await;
367 }
368 }
369 return None;
370 }
371 if let Err(error) = send.finish() {
372 tracing::debug!(%error, %peer, "QUIC request stream closed before it could finish");
373 return None;
374 }
375 Some(recv)
376}
377
378fn complete_response(
379 final_response: Option<Message>,
380 pending: &mut BytesMut,
381 limits: &Limits,
382) -> std::result::Result<Message, String> {
383 if let Some(message) = final_response {
384 return pending
385 .is_empty()
386 .then_some(message)
387 .ok_or_else(|| "bytes followed the final response".to_owned());
388 }
389 match take_response(pending, true, limits)? {
390 Some(message)
391 if matches!(&message, Message::Response(response) if response.status.is_final())
392 && pending.is_empty() =>
393 {
394 Ok(message)
395 }
396 Some(_) | None => Err("response stream ended without one final response".to_owned()),
397 }
398}
399
400async fn send_response_event(
401 events: &mpsc::Sender<Event>,
402 message: Message,
403 peer: SocketAddr,
404 id: u64,
405) -> std::result::Result<(), mpsc::error::SendError<Event>> {
406 events
407 .send(Event::Message {
408 message: Box::new(message),
409 source: peer,
410 transport: TransportKind::Quic,
411 id,
412 quic_reply: None,
413 })
414 .await
415}
416
417async fn reject_response_stream(
418 connection: &Connection,
419 events: &mpsc::Sender<Event>,
420 key: &ConnectionKey,
421 peer: SocketAddr,
422 detail: String,
423) {
424 tracing::debug!(%detail, %peer, "malformed QUIC response stream");
427 let _ = events.send(Event::FramingFailed { key: key.clone() }).await;
428 connection.close(1_u8.into(), b"malformed SIP stream");
429}
430
431fn take_response(
436 pending: &mut BytesMut,
437 end: bool,
438 limits: &Limits,
439) -> std::result::Result<Option<Message>, String> {
440 let Some(head_end) = pending.windows(4).position(|window| window == b"\r\n\r\n") else {
441 if pending.len() > limits.max_message_bytes || end && !pending.is_empty() {
442 return Err("response has no complete header section".to_owned());
443 }
444 return Ok(None);
445 };
446 let body_start = head_end
447 .checked_add(4)
448 .ok_or_else(|| "response length overflow".to_owned())?;
449 let header = pending
450 .get(..head_end)
451 .ok_or_else(|| "response header length overflow".to_owned())?;
452 let frame_len = match declared_content_length(header)? {
453 Some(declared) => body_start
454 .checked_add(declared)
455 .ok_or_else(|| "response length overflow".to_owned())?,
456 None if end => pending.len(),
457 None => {
458 if pending.len() > limits.max_message_bytes {
459 return Err("response exceeds the configured message limit".to_owned());
460 }
461 return Ok(None);
462 }
463 };
464 if frame_len > limits.max_message_bytes {
465 return Err("response exceeds the configured message limit".to_owned());
466 }
467 if pending.len() < frame_len {
468 if end {
469 return Err("response body ended before Content-Length".to_owned());
470 }
471 return Ok(None);
472 }
473 parse_one(pending.split_to(frame_len).freeze(), limits).map(Some)
474}
475
476fn declared_content_length(head: &[u8]) -> std::result::Result<Option<usize>, String> {
477 let mut lines = head.split(|byte| *byte == b'\n');
478 let _start_line = lines.next();
479 let mut value: Option<Vec<u8>> = None;
480 let mut continuing_length = false;
481 for raw_line in lines {
482 let line = raw_line.strip_suffix(b"\r").unwrap_or(raw_line);
483 if line.first().is_some_and(u8::is_ascii_whitespace) {
484 if continuing_length {
485 let stored = value
486 .as_mut()
487 .ok_or_else(|| "missing Content-Length value".to_owned())?;
488 stored.push(b' ');
489 stored.extend_from_slice(trim_ascii(line));
490 }
491 continue;
492 }
493 continuing_length = false;
494 let mut parts = line.splitn(2, |byte| *byte == b':');
495 let name = parts.next().unwrap_or_default();
496 let Some(raw_value) = parts.next() else {
497 continue;
498 };
499 let name = trim_ascii(name);
500 if name.eq_ignore_ascii_case(b"content-length") || name.eq_ignore_ascii_case(b"l") {
501 if value.is_some() {
502 return Err("repeated Content-Length".to_owned());
503 }
504 value = Some(trim_ascii(raw_value).to_vec());
505 continuing_length = true;
506 }
507 }
508 value
509 .map(|value| {
510 std::str::from_utf8(&value)
511 .map_err(|error| error.to_string())?
512 .parse::<usize>()
513 .map_err(|error| error.to_string())
514 })
515 .transpose()
516}
517
518fn trim_ascii(mut bytes: &[u8]) -> &[u8] {
519 while bytes.first().is_some_and(u8::is_ascii_whitespace) {
520 bytes = bytes.get(1..).unwrap_or_default();
521 }
522 while bytes.last().is_some_and(u8::is_ascii_whitespace) {
523 bytes = bytes
524 .get(..bytes.len().saturating_sub(1))
525 .unwrap_or_default();
526 }
527 bytes
528}
529
530#[allow(
532 clippy::needless_pass_by_value,
533 reason = "the parsed message retains views into this owned Bytes allocation"
534)]
535fn parse_one(bytes: Bytes, limits: &Limits) -> std::result::Result<Message, String> {
536 let head_end = bytes
537 .windows(4)
538 .position(|window| window == b"\r\n\r\n")
539 .ok_or_else(|| "no header terminator".to_owned())?;
540 let message = parse_datagram(bytes.clone(), limits).map_err(|error| error.to_string())?;
541 if let Some(value) = message
542 .headers()
543 .value(&sipx_sip::HeaderName::ContentLength)
544 {
545 let declared = std::str::from_utf8(&value)
546 .map_err(|error| error.to_string())?
547 .trim()
548 .parse::<usize>()
549 .map_err(|error| error.to_string())?;
550 let actual = bytes.len().saturating_sub(head_end + 4);
551 if declared != actual {
552 return Err(format!(
553 "Content-Length {declared} does not match {actual} stream bytes"
554 ));
555 }
556 }
557 Ok(message)
558}