1use crate::poll::{PollEvents, POLLERR, POLLHUP, POLLIN, POLLOUT};
2use crate::vfs::normalize_path;
3use std::collections::{BTreeMap, BTreeSet, VecDeque};
4use std::error::Error;
5use std::fmt;
6use std::net::{Ipv4Addr, Ipv6Addr};
7use std::sync::{Arc, Mutex, MutexGuard};
8
9pub type SocketId = u64;
10pub type SocketResult<T> = Result<T, SocketTableError>;
11
12#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord)]
13pub struct InetSocketAddress {
14 host: String,
15 port: u16,
16}
17
18impl InetSocketAddress {
19 pub fn new(host: impl Into<String>, port: u16) -> Self {
20 Self {
21 host: host.into(),
22 port,
23 }
24 }
25
26 pub fn host(&self) -> &str {
27 &self.host
28 }
29
30 pub const fn port(&self) -> u16 {
31 self.port
32 }
33}
34
35#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
36pub enum SocketDomain {
37 Inet,
38 Inet6,
39 Unix,
40}
41
42#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
43pub enum SocketType {
44 Stream,
45 Datagram,
46}
47
48#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
49pub enum SocketState {
50 Created,
51 Bound,
52 Listening,
53 Connected,
54}
55
56impl SocketState {
57 pub const fn counts_as_listener(self) -> bool {
58 matches!(self, Self::Listening)
59 }
60
61 pub const fn counts_as_connection(self) -> bool {
62 matches!(self, Self::Connected)
63 }
64}
65
66#[derive(Debug, Clone, Copy, PartialEq, Eq)]
67pub enum SocketShutdown {
68 Read,
69 Write,
70 Both,
71}
72
73#[derive(Debug, Clone, Copy, PartialEq, Eq)]
74pub enum DatagramSocketOption {
75 ReuseAddr,
76 ReusePort,
77 Broadcast,
78}
79
80#[derive(Debug, Clone, Copy, PartialEq, Eq)]
81pub struct SocketSpec {
82 pub domain: SocketDomain,
83 pub socket_type: SocketType,
84}
85
86impl SocketSpec {
87 pub const fn new(domain: SocketDomain, socket_type: SocketType) -> Self {
88 Self {
89 domain,
90 socket_type,
91 }
92 }
93
94 pub const fn tcp() -> Self {
95 Self::new(SocketDomain::Inet, SocketType::Stream)
96 }
97
98 pub const fn udp() -> Self {
99 Self::new(SocketDomain::Inet, SocketType::Datagram)
100 }
101
102 pub const fn unix_stream() -> Self {
103 Self::new(SocketDomain::Unix, SocketType::Stream)
104 }
105
106 pub const fn unix_datagram() -> Self {
107 Self::new(SocketDomain::Unix, SocketType::Datagram)
108 }
109}
110
111#[derive(Debug, Clone, PartialEq, Eq)]
112pub struct SocketRecord {
113 id: SocketId,
114 owner_pid: u32,
115 spec: SocketSpec,
116 state: SocketState,
117 local_address: Option<InetSocketAddress>,
118 peer_address: Option<InetSocketAddress>,
119 local_unix_path: Option<String>,
120 peer_unix_path: Option<String>,
121 listener_state: Option<ListenerState>,
122 connection_state: Option<ConnectionState>,
123 datagram_state: Option<DatagramState>,
124}
125
126impl SocketRecord {
127 pub const fn id(&self) -> SocketId {
128 self.id
129 }
130
131 pub const fn owner_pid(&self) -> u32 {
132 self.owner_pid
133 }
134
135 pub const fn spec(&self) -> SocketSpec {
136 self.spec
137 }
138
139 pub const fn state(&self) -> SocketState {
140 self.state
141 }
142
143 pub fn local_address(&self) -> Option<&InetSocketAddress> {
144 self.local_address.as_ref()
145 }
146
147 pub fn peer_address(&self) -> Option<&InetSocketAddress> {
148 self.peer_address.as_ref()
149 }
150
151 pub fn local_unix_path(&self) -> Option<&str> {
152 self.local_unix_path.as_deref()
153 }
154
155 pub fn peer_unix_path(&self) -> Option<&str> {
156 self.peer_unix_path.as_deref()
157 }
158
159 pub fn listen_backlog(&self) -> Option<usize> {
160 self.listener_state.as_ref().map(|state| state.backlog)
161 }
162
163 pub fn pending_accept_count(&self) -> usize {
164 self.listener_state
165 .as_ref()
166 .map(|state| state.pending_accepts.len())
167 .unwrap_or(0)
168 }
169
170 pub fn peer_socket_id(&self) -> Option<SocketId> {
171 self.connection_state
172 .as_ref()
173 .and_then(|state| state.peer_socket_id)
174 }
175
176 pub fn buffered_read_bytes(&self) -> usize {
177 self.connection_state
178 .as_ref()
179 .map(|state| state.recv_buffer.len())
180 .unwrap_or(0)
181 }
182
183 pub fn read_shutdown(&self) -> bool {
184 self.connection_state
185 .as_ref()
186 .map(|state| state.read_shutdown)
187 .unwrap_or(false)
188 }
189
190 pub fn write_shutdown(&self) -> bool {
191 self.connection_state
192 .as_ref()
193 .map(|state| state.write_shutdown)
194 .unwrap_or(false)
195 }
196
197 pub fn peer_write_shutdown(&self) -> bool {
198 self.connection_state
199 .as_ref()
200 .map(|state| state.peer_write_shutdown)
201 .unwrap_or(false)
202 }
203
204 pub fn queued_datagrams(&self) -> usize {
205 self.datagram_state
206 .as_ref()
207 .map(|state| state.recv_queue.len())
208 .unwrap_or(0)
209 }
210
211 pub fn reuse_address(&self) -> bool {
212 self.datagram_state
213 .as_ref()
214 .map(|state| state.reuse_addr)
215 .unwrap_or(false)
216 }
217
218 pub fn reuse_port(&self) -> bool {
219 self.datagram_state
220 .as_ref()
221 .map(|state| state.reuse_port)
222 .unwrap_or(false)
223 }
224
225 pub fn broadcast_enabled(&self) -> bool {
226 self.datagram_state
227 .as_ref()
228 .map(|state| state.broadcast)
229 .unwrap_or(false)
230 }
231
232 pub fn multicast_membership_count(&self) -> usize {
233 self.datagram_state
234 .as_ref()
235 .map(|state| state.multicast_memberships.len())
236 .unwrap_or(0)
237 }
238
239 pub fn has_multicast_membership(&self, membership: &SocketMulticastMembership) -> bool {
240 self.datagram_state
241 .as_ref()
242 .map(|state| state.multicast_memberships.contains(membership))
243 .unwrap_or(false)
244 }
245}
246
247#[derive(Debug, Clone, PartialEq, Eq)]
248pub struct ReceivedDatagram {
249 source_address: Option<InetSocketAddress>,
250 payload: Vec<u8>,
251}
252
253impl ReceivedDatagram {
254 pub fn source_address(&self) -> Option<&InetSocketAddress> {
255 self.source_address.as_ref()
256 }
257
258 pub fn payload(&self) -> &[u8] {
259 &self.payload
260 }
261
262 pub fn into_parts(self) -> (Option<InetSocketAddress>, Vec<u8>) {
263 (self.source_address, self.payload)
264 }
265}
266
267#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
268pub struct SocketTableSnapshot {
269 pub sockets: usize,
270 pub listeners: usize,
271 pub connections: usize,
272}
273
274#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord)]
275pub struct SocketMulticastMembership {
276 group_address: String,
277 interface_address: Option<String>,
278}
279
280impl SocketMulticastMembership {
281 pub fn new(group_address: impl Into<String>, interface_address: Option<String>) -> Self {
282 Self {
283 group_address: group_address.into(),
284 interface_address,
285 }
286 }
287
288 pub fn group_address(&self) -> &str {
289 &self.group_address
290 }
291
292 pub fn interface_address(&self) -> Option<&str> {
293 self.interface_address.as_deref()
294 }
295}
296
297#[derive(Debug, Clone, PartialEq, Eq)]
298pub struct SocketTableError {
299 code: &'static str,
300 message: String,
301}
302
303impl SocketTableError {
304 pub fn code(&self) -> &'static str {
305 self.code
306 }
307
308 fn not_found(socket_id: SocketId) -> Self {
309 Self {
310 code: "ENOENT",
311 message: format!("no such socket {socket_id}"),
312 }
313 }
314
315 fn invalid_argument(message: impl Into<String>) -> Self {
316 Self {
317 code: "EINVAL",
318 message: message.into(),
319 }
320 }
321
322 fn address_in_use(message: impl Into<String>) -> Self {
323 Self {
324 code: "EADDRINUSE",
325 message: message.into(),
326 }
327 }
328
329 fn address_not_available(message: impl Into<String>) -> Self {
330 Self {
331 code: "EADDRNOTAVAIL",
332 message: message.into(),
333 }
334 }
335
336 fn not_found_address(message: impl Into<String>) -> Self {
337 Self {
338 code: "ECONNREFUSED",
339 message: message.into(),
340 }
341 }
342
343 fn would_block(message: impl Into<String>) -> Self {
344 Self {
345 code: "EAGAIN",
346 message: message.into(),
347 }
348 }
349
350 fn not_connected(message: impl Into<String>) -> Self {
351 Self {
352 code: "ENOTCONN",
353 message: message.into(),
354 }
355 }
356
357 fn broken_pipe(message: impl Into<String>) -> Self {
358 Self {
359 code: "EPIPE",
360 message: message.into(),
361 }
362 }
363}
364
365impl fmt::Display for SocketTableError {
366 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
367 write!(f, "{}: {}", self.code, self.message)
368 }
369}
370
371impl Error for SocketTableError {}
372
373#[derive(Debug, Default)]
374struct SocketTableState {
375 sockets: BTreeMap<SocketId, SocketRecord>,
376 by_owner: BTreeMap<u32, BTreeSet<SocketId>>,
377 bound_inet_streams: BTreeMap<InetSocketAddress, SocketId>,
378 bound_inet_datagrams: BTreeMap<InetSocketAddress, BTreeSet<SocketId>>,
379 bound_unix_streams: BTreeMap<String, SocketId>,
380 multicast_groups: BTreeMap<SocketMulticastMembership, BTreeSet<SocketId>>,
381 next_socket_id: SocketId,
382}
383
384#[derive(Debug, Clone, PartialEq, Eq)]
385struct ListenerState {
386 backlog: usize,
387 pending_accepts: VecDeque<PendingConnection>,
388}
389
390#[derive(Debug, Clone, PartialEq, Eq, Default)]
391struct ConnectionState {
392 peer_socket_id: Option<SocketId>,
393 recv_buffer: VecDeque<u8>,
394 read_shutdown: bool,
395 write_shutdown: bool,
396 peer_write_shutdown: bool,
397}
398
399#[derive(Debug, Clone, PartialEq, Eq)]
400struct PendingConnection {
401 peer_address: Option<InetSocketAddress>,
402 peer_unix_path: Option<String>,
403 accepted_socket_id: Option<SocketId>,
404}
405
406#[derive(Debug, Clone, PartialEq, Eq, Default)]
407struct DatagramState {
408 recv_queue: VecDeque<QueuedDatagram>,
409 reuse_addr: bool,
410 reuse_port: bool,
411 broadcast: bool,
412 multicast_memberships: BTreeSet<SocketMulticastMembership>,
413}
414
415#[derive(Debug, Clone, PartialEq, Eq)]
416struct QueuedDatagram {
417 source_address: Option<InetSocketAddress>,
418 payload: Vec<u8>,
419}
420
421#[derive(Debug, Default)]
422struct SocketTableInner {
423 state: Mutex<SocketTableState>,
424}
425
426#[derive(Debug, Clone, Default)]
427pub struct SocketTable {
428 inner: Arc<SocketTableInner>,
429}
430
431impl SocketTable {
432 pub fn new() -> Self {
433 Self::default()
434 }
435
436 pub fn allocate(&self, owner_pid: u32, spec: SocketSpec) -> SocketRecord {
437 self.allocate_with_state(owner_pid, spec, SocketState::Created)
438 }
439
440 pub fn allocate_with_state(
441 &self,
442 owner_pid: u32,
443 spec: SocketSpec,
444 state: SocketState,
445 ) -> SocketRecord {
446 let mut table = lock_or_recover(&self.inner.state);
447 let socket_id = next_socket_id(&mut table);
448 let record = SocketRecord {
449 id: socket_id,
450 owner_pid,
451 spec,
452 state,
453 local_address: None,
454 peer_address: None,
455 local_unix_path: None,
456 peer_unix_path: None,
457 listener_state: None,
458 connection_state: default_connection_state(spec, state),
459 datagram_state: default_datagram_state(spec),
460 };
461 table.sockets.insert(socket_id, record.clone());
462 table
463 .by_owner
464 .entry(owner_pid)
465 .or_default()
466 .insert(socket_id);
467 record
468 }
469
470 pub fn get(&self, socket_id: SocketId) -> Option<SocketRecord> {
471 lock_or_recover(&self.inner.state)
472 .sockets
473 .get(&socket_id)
474 .cloned()
475 }
476
477 pub fn update_state(
478 &self,
479 socket_id: SocketId,
480 new_state: SocketState,
481 ) -> SocketResult<SocketRecord> {
482 let mut table = lock_or_recover(&self.inner.state);
483 let record = table
484 .sockets
485 .get_mut(&socket_id)
486 .ok_or_else(|| SocketTableError::not_found(socket_id))?;
487 validate_state_transition(record.state, new_state)?;
488 record.state = new_state;
489 if new_state != SocketState::Listening {
490 record.listener_state = None;
491 }
492 if new_state == SocketState::Connected && supports_connection_lifecycle(record.spec) {
493 record
494 .connection_state
495 .get_or_insert_with(ConnectionState::default);
496 } else if new_state != SocketState::Connected {
497 record.connection_state = None;
498 }
499 Ok(record.clone())
500 }
501
502 pub fn bind_inet(
503 &self,
504 socket_id: SocketId,
505 address: InetSocketAddress,
506 ) -> SocketResult<SocketRecord> {
507 let address = normalize_inet_address(address);
508 let mut table = lock_or_recover(&self.inner.state);
509 let existing = table
510 .sockets
511 .get(&socket_id)
512 .cloned()
513 .ok_or_else(|| SocketTableError::not_found(socket_id))?;
514 if !supports_inet_bind(existing.spec) {
515 return Err(SocketTableError::invalid_argument(format!(
516 "socket {socket_id} is not an INET socket"
517 )));
518 }
519 let conflicting_ids =
520 lookup_conflicting_bound_inet_socket_ids(&table, existing.spec, &address);
521 if has_incompatible_inet_bind_conflict(&table, &existing, &conflicting_ids) {
522 return Err(SocketTableError::address_in_use(format!(
523 "address {}:{} is already bound",
524 address.host(),
525 address.port()
526 )));
527 }
528 let cloned = {
529 let record = table
530 .sockets
531 .get_mut(&socket_id)
532 .ok_or_else(|| SocketTableError::not_found(socket_id))?;
533
534 match record.state {
535 SocketState::Created => {}
536 SocketState::Bound if record.local_address.as_ref() == Some(&address) => {
537 return Ok(record.clone());
538 }
539 SocketState::Bound | SocketState::Listening | SocketState::Connected => {
540 return Err(SocketTableError::invalid_argument(format!(
541 "socket {socket_id} cannot bind in state {:?}",
542 record.state
543 )));
544 }
545 }
546
547 record.local_address = Some(address.clone());
548 record.peer_address = None;
549 record.local_unix_path = None;
550 record.peer_unix_path = None;
551 record.listener_state = None;
552 record.connection_state = None;
553 record.state = SocketState::Bound;
554 record.clone()
555 };
556 register_bound_inet_socket(&mut table, cloned.spec, address, socket_id);
557 Ok(cloned)
558 }
559
560 pub fn set_datagram_socket_option(
561 &self,
562 socket_id: SocketId,
563 option: DatagramSocketOption,
564 enabled: bool,
565 ) -> SocketResult<SocketRecord> {
566 let mut table = lock_or_recover(&self.inner.state);
567 let record = table
568 .sockets
569 .get_mut(&socket_id)
570 .ok_or_else(|| SocketTableError::not_found(socket_id))?;
571 let datagram_state = datagram_state_mut(record)?;
572
573 match option {
574 DatagramSocketOption::ReuseAddr => datagram_state.reuse_addr = enabled,
575 DatagramSocketOption::ReusePort => datagram_state.reuse_port = enabled,
576 DatagramSocketOption::Broadcast => datagram_state.broadcast = enabled,
577 }
578
579 Ok(record.clone())
580 }
581
582 pub fn add_multicast_membership(
583 &self,
584 socket_id: SocketId,
585 membership: SocketMulticastMembership,
586 ) -> SocketResult<SocketRecord> {
587 let mut table = lock_or_recover(&self.inner.state);
588 let normalized_membership = {
589 let record = table
590 .sockets
591 .get(&socket_id)
592 .ok_or_else(|| SocketTableError::not_found(socket_id))?;
593 validate_multicast_socket(record)?;
594 normalize_multicast_membership(record.spec, membership)?
595 };
596
597 let cloned = {
598 let record = table
599 .sockets
600 .get_mut(&socket_id)
601 .ok_or_else(|| SocketTableError::not_found(socket_id))?;
602 let datagram_state = datagram_state_mut(record)?;
603 datagram_state
604 .multicast_memberships
605 .insert(normalized_membership.clone());
606 record.clone()
607 };
608
609 table
610 .multicast_groups
611 .entry(normalized_membership)
612 .or_default()
613 .insert(socket_id);
614 Ok(cloned)
615 }
616
617 pub fn drop_multicast_membership(
618 &self,
619 socket_id: SocketId,
620 membership: SocketMulticastMembership,
621 ) -> SocketResult<SocketRecord> {
622 let mut table = lock_or_recover(&self.inner.state);
623 let normalized_membership = {
624 let record = table
625 .sockets
626 .get(&socket_id)
627 .ok_or_else(|| SocketTableError::not_found(socket_id))?;
628 validate_multicast_socket(record)?;
629 normalize_multicast_membership(record.spec, membership)?
630 };
631
632 let cloned = {
633 let record = table
634 .sockets
635 .get_mut(&socket_id)
636 .ok_or_else(|| SocketTableError::not_found(socket_id))?;
637 let datagram_state = datagram_state_mut(record)?;
638 if !datagram_state
639 .multicast_memberships
640 .remove(&normalized_membership)
641 {
642 return Err(SocketTableError::address_not_available(format!(
643 "socket {socket_id} has not joined multicast group {}",
644 normalized_membership.group_address()
645 )));
646 }
647 record.clone()
648 };
649
650 if let Some(members) = table.multicast_groups.get_mut(&normalized_membership) {
651 members.remove(&socket_id);
652 if members.is_empty() {
653 table.multicast_groups.remove(&normalized_membership);
654 }
655 }
656
657 Ok(cloned)
658 }
659
660 pub fn bind_unix(
661 &self,
662 socket_id: SocketId,
663 path: impl Into<String>,
664 ) -> SocketResult<SocketRecord> {
665 let path = normalize_unix_socket_path(path.into())?;
666 let mut table = lock_or_recover(&self.inner.state);
667 let existing = table
668 .sockets
669 .get(&socket_id)
670 .cloned()
671 .ok_or_else(|| SocketTableError::not_found(socket_id))?;
672 if !supports_unix_stream_lifecycle(existing.spec) {
673 return Err(SocketTableError::invalid_argument(format!(
674 "socket {socket_id} is not a Unix stream socket"
675 )));
676 }
677 let existing_id = table.bound_unix_streams.get(&path).copied();
678 let cloned = {
679 let record = table
680 .sockets
681 .get_mut(&socket_id)
682 .ok_or_else(|| SocketTableError::not_found(socket_id))?;
683
684 if let Some(bound_socket_id) = existing_id {
685 if bound_socket_id != socket_id {
686 return Err(SocketTableError::address_in_use(format!(
687 "path {path} is already bound"
688 )));
689 }
690 }
691
692 match record.state {
693 SocketState::Created => {}
694 SocketState::Bound if record.local_unix_path.as_deref() == Some(path.as_str()) => {
695 return Ok(record.clone());
696 }
697 SocketState::Bound | SocketState::Listening | SocketState::Connected => {
698 return Err(SocketTableError::invalid_argument(format!(
699 "socket {socket_id} cannot bind in state {:?}",
700 record.state
701 )));
702 }
703 }
704
705 record.local_address = None;
706 record.peer_address = None;
707 record.local_unix_path = Some(path.clone());
708 record.peer_unix_path = None;
709 record.listener_state = None;
710 record.connection_state = None;
711 record.state = SocketState::Bound;
712 record.clone()
713 };
714 table.bound_unix_streams.insert(path, socket_id);
715 Ok(cloned)
716 }
717
718 pub fn listen(&self, socket_id: SocketId, backlog: usize) -> SocketResult<SocketRecord> {
719 if backlog == 0 {
720 return Err(SocketTableError::invalid_argument(
721 "listener backlog must be greater than zero",
722 ));
723 }
724
725 let mut table = lock_or_recover(&self.inner.state);
726 let record = table
727 .sockets
728 .get_mut(&socket_id)
729 .ok_or_else(|| SocketTableError::not_found(socket_id))?;
730
731 if !supports_listener_lifecycle(record.spec) {
732 return Err(SocketTableError::invalid_argument(format!(
733 "socket {socket_id} is not a stream socket"
734 )));
735 }
736 if record.state != SocketState::Bound || !has_bound_endpoint(record) {
737 return Err(SocketTableError::invalid_argument(format!(
738 "socket {socket_id} must be bound before listen"
739 )));
740 }
741
742 record.state = SocketState::Listening;
743 record.listener_state = Some(ListenerState {
744 backlog,
745 pending_accepts: VecDeque::new(),
746 });
747 Ok(record.clone())
748 }
749
750 pub fn enqueue_incoming_tcp_connection(
751 &self,
752 listener_socket_id: SocketId,
753 peer_address: InetSocketAddress,
754 ) -> SocketResult<()> {
755 let mut table = lock_or_recover(&self.inner.state);
756 let record = table
757 .sockets
758 .get_mut(&listener_socket_id)
759 .ok_or_else(|| SocketTableError::not_found(listener_socket_id))?;
760
761 if record.state != SocketState::Listening {
762 return Err(SocketTableError::invalid_argument(format!(
763 "socket {listener_socket_id} is not listening"
764 )));
765 }
766
767 let listener_state = record.listener_state.as_mut().ok_or_else(|| {
768 SocketTableError::invalid_argument(format!(
769 "socket {listener_socket_id} has no listener state"
770 ))
771 })?;
772
773 if listener_state.pending_accepts.len() >= listener_state.backlog {
774 return Err(SocketTableError::would_block(format!(
775 "listener {listener_socket_id} backlog is full"
776 )));
777 }
778
779 listener_state.pending_accepts.push_back(PendingConnection {
780 peer_address: Some(peer_address),
781 peer_unix_path: None,
782 accepted_socket_id: None,
783 });
784 Ok(())
785 }
786
787 pub fn accept(&self, listener_socket_id: SocketId) -> SocketResult<SocketRecord> {
788 let mut table = lock_or_recover(&self.inner.state);
789 let (owner_pid, spec, local_address, local_unix_path, pending) = {
790 let record = table
791 .sockets
792 .get_mut(&listener_socket_id)
793 .ok_or_else(|| SocketTableError::not_found(listener_socket_id))?;
794
795 if record.state != SocketState::Listening {
796 return Err(SocketTableError::invalid_argument(format!(
797 "socket {listener_socket_id} is not listening"
798 )));
799 }
800
801 let listener_state = record.listener_state.as_mut().ok_or_else(|| {
802 SocketTableError::invalid_argument(format!(
803 "socket {listener_socket_id} has no listener state"
804 ))
805 })?;
806 let pending = listener_state.pending_accepts.pop_front().ok_or_else(|| {
807 SocketTableError::would_block(format!(
808 "listener {listener_socket_id} has no pending connections"
809 ))
810 })?;
811
812 (
813 record.owner_pid,
814 record.spec,
815 record.local_address.clone(),
816 record.local_unix_path.clone(),
817 pending,
818 )
819 };
820
821 if let Some(accepted_socket_id) = pending.accepted_socket_id {
822 return table
823 .sockets
824 .get(&accepted_socket_id)
825 .cloned()
826 .ok_or_else(|| SocketTableError::not_found(accepted_socket_id));
827 }
828
829 let socket_id = next_socket_id(&mut table);
830 let record = SocketRecord {
831 id: socket_id,
832 owner_pid,
833 spec,
834 state: SocketState::Connected,
835 local_address,
836 peer_address: pending.peer_address,
837 local_unix_path,
838 peer_unix_path: pending.peer_unix_path,
839 listener_state: None,
840 connection_state: default_connection_state(spec, SocketState::Connected),
841 datagram_state: default_datagram_state(spec),
842 };
843 table.sockets.insert(socket_id, record.clone());
844 table
845 .by_owner
846 .entry(owner_pid)
847 .or_default()
848 .insert(socket_id);
849 Ok(record)
850 }
851
852 pub fn connect_pair(
853 &self,
854 socket_id: SocketId,
855 peer_socket_id: SocketId,
856 ) -> SocketResult<(SocketRecord, SocketRecord)> {
857 if socket_id == peer_socket_id {
858 return Err(SocketTableError::invalid_argument(
859 "socket cannot connect to itself",
860 ));
861 }
862
863 let mut table = lock_or_recover(&self.inner.state);
864 let mut socket = table
865 .sockets
866 .remove(&socket_id)
867 .ok_or_else(|| SocketTableError::not_found(socket_id))?;
868 let Some(mut peer) = table.sockets.remove(&peer_socket_id) else {
869 table.sockets.insert(socket_id, socket);
870 return Err(SocketTableError::not_found(peer_socket_id));
871 };
872
873 if let Err(error) = validate_connect_pair(&socket, &peer) {
874 table.sockets.insert(socket_id, socket);
875 table.sockets.insert(peer_socket_id, peer);
876 return Err(error);
877 }
878
879 socket.state = SocketState::Connected;
880 socket.peer_address = peer.local_address.clone();
881 socket.peer_unix_path = peer.local_unix_path.clone();
882 socket.listener_state = None;
883 socket.connection_state = Some(ConnectionState {
884 peer_socket_id: Some(peer_socket_id),
885 ..ConnectionState::default()
886 });
887
888 peer.state = SocketState::Connected;
889 peer.peer_address = socket.local_address.clone();
890 peer.peer_unix_path = socket.local_unix_path.clone();
891 peer.listener_state = None;
892 peer.connection_state = Some(ConnectionState {
893 peer_socket_id: Some(socket_id),
894 ..ConnectionState::default()
895 });
896
897 let socket_clone = socket.clone();
898 let peer_clone = peer.clone();
899 table.sockets.insert(socket_id, socket);
900 table.sockets.insert(peer_socket_id, peer);
901 Ok((socket_clone, peer_clone))
902 }
903
904 pub fn find_bound_inet_socket(
905 &self,
906 spec: SocketSpec,
907 address: &InetSocketAddress,
908 ) -> Option<SocketRecord> {
909 let address = normalize_inet_address(address.clone());
910 let table = lock_or_recover(&self.inner.state);
911 let socket_id = lookup_bound_inet_socket(&table, spec, &address)?;
912 table.sockets.get(&socket_id).cloned()
913 }
914
915 pub fn connect_to_bound_inet_stream(
916 &self,
917 socket_id: SocketId,
918 target_address: InetSocketAddress,
919 ) -> SocketResult<()> {
920 let target_address = normalize_inet_address(target_address);
921 let mut table = lock_or_recover(&self.inner.state);
922 let listener_socket_id =
923 lookup_bound_inet_socket_in_table(&table.bound_inet_streams, &target_address)
924 .ok_or_else(|| {
925 SocketTableError::not_found_address(format!(
926 "no listening socket bound at {}:{}",
927 target_address.host(),
928 target_address.port()
929 ))
930 })?;
931
932 if socket_id == listener_socket_id {
933 return Err(SocketTableError::invalid_argument(
934 "socket cannot connect to its own listening endpoint",
935 ));
936 }
937
938 let mut client = table
939 .sockets
940 .remove(&socket_id)
941 .ok_or_else(|| SocketTableError::not_found(socket_id))?;
942 let accepted_socket_id = next_socket_id(&mut table);
943
944 let result = (|| {
945 let listener = table
946 .sockets
947 .get_mut(&listener_socket_id)
948 .ok_or_else(|| SocketTableError::not_found(listener_socket_id))?;
949 validate_connect_to_listener(&client, listener)?;
950
951 let listener_state = listener.listener_state.as_mut().ok_or_else(|| {
952 SocketTableError::invalid_argument(format!(
953 "socket {listener_socket_id} has no listener state"
954 ))
955 })?;
956 if listener_state.pending_accepts.len() >= listener_state.backlog {
957 return Err(SocketTableError::would_block(format!(
958 "listener {listener_socket_id} backlog is full"
959 )));
960 }
961
962 let accepted = SocketRecord {
963 id: accepted_socket_id,
964 owner_pid: listener.owner_pid,
965 spec: listener.spec,
966 state: SocketState::Connected,
967 local_address: listener.local_address.clone(),
968 peer_address: client.local_address.clone(),
969 local_unix_path: None,
970 peer_unix_path: None,
971 listener_state: None,
972 connection_state: Some(ConnectionState {
973 peer_socket_id: Some(socket_id),
974 ..ConnectionState::default()
975 }),
976 datagram_state: default_datagram_state(listener.spec),
977 };
978
979 listener_state.pending_accepts.push_back(PendingConnection {
980 peer_address: client.local_address.clone(),
981 peer_unix_path: None,
982 accepted_socket_id: Some(accepted_socket_id),
983 });
984
985 client.state = SocketState::Connected;
986 client.peer_address = listener.local_address.clone();
987 client.peer_unix_path = None;
988 client.listener_state = None;
989 client.connection_state = Some(ConnectionState {
990 peer_socket_id: Some(accepted_socket_id),
991 ..ConnectionState::default()
992 });
993
994 Ok(accepted)
995 })();
996
997 match result {
998 Ok(accepted) => {
999 table.sockets.insert(socket_id, client);
1000 table.sockets.insert(accepted_socket_id, accepted.clone());
1001 table
1002 .by_owner
1003 .entry(accepted.owner_pid)
1004 .or_default()
1005 .insert(accepted_socket_id);
1006 Ok(())
1007 }
1008 Err(error) => {
1009 table.sockets.insert(socket_id, client);
1010 Err(error)
1011 }
1012 }
1013 }
1014
1015 pub fn find_bound_unix_socket(&self, path: &str) -> Option<SocketRecord> {
1016 let path = normalize_unix_socket_path(path).ok()?;
1017 let table = lock_or_recover(&self.inner.state);
1018 let socket_id = table.bound_unix_streams.get(&path).copied()?;
1019 table.sockets.get(&socket_id).cloned()
1020 }
1021
1022 pub fn connect_to_bound_unix_stream(
1023 &self,
1024 socket_id: SocketId,
1025 target_path: impl Into<String>,
1026 ) -> SocketResult<()> {
1027 let target_path = normalize_unix_socket_path(target_path.into())?;
1028 let mut table = lock_or_recover(&self.inner.state);
1029 let listener_socket_id = table
1030 .bound_unix_streams
1031 .get(&target_path)
1032 .copied()
1033 .ok_or_else(|| {
1034 SocketTableError::not_found_address(format!(
1035 "no listening socket bound at path {target_path}"
1036 ))
1037 })?;
1038
1039 if socket_id == listener_socket_id {
1040 return Err(SocketTableError::invalid_argument(
1041 "socket cannot connect to its own listening endpoint",
1042 ));
1043 }
1044
1045 let mut client = table
1046 .sockets
1047 .remove(&socket_id)
1048 .ok_or_else(|| SocketTableError::not_found(socket_id))?;
1049 let accepted_socket_id = next_socket_id(&mut table);
1050
1051 let result = (|| {
1052 let listener = table
1053 .sockets
1054 .get_mut(&listener_socket_id)
1055 .ok_or_else(|| SocketTableError::not_found(listener_socket_id))?;
1056 validate_connect_to_listener(&client, listener)?;
1057
1058 let listener_state = listener.listener_state.as_mut().ok_or_else(|| {
1059 SocketTableError::invalid_argument(format!(
1060 "socket {listener_socket_id} has no listener state"
1061 ))
1062 })?;
1063 if listener_state.pending_accepts.len() >= listener_state.backlog {
1064 return Err(SocketTableError::would_block(format!(
1065 "listener {listener_socket_id} backlog is full"
1066 )));
1067 }
1068
1069 let accepted = SocketRecord {
1070 id: accepted_socket_id,
1071 owner_pid: listener.owner_pid,
1072 spec: listener.spec,
1073 state: SocketState::Connected,
1074 local_address: None,
1075 peer_address: None,
1076 local_unix_path: listener.local_unix_path.clone(),
1077 peer_unix_path: client.local_unix_path.clone(),
1078 listener_state: None,
1079 connection_state: Some(ConnectionState {
1080 peer_socket_id: Some(socket_id),
1081 ..ConnectionState::default()
1082 }),
1083 datagram_state: default_datagram_state(listener.spec),
1084 };
1085
1086 listener_state.pending_accepts.push_back(PendingConnection {
1087 peer_address: None,
1088 peer_unix_path: client.local_unix_path.clone(),
1089 accepted_socket_id: Some(accepted_socket_id),
1090 });
1091
1092 client.state = SocketState::Connected;
1093 client.peer_address = None;
1094 client.peer_unix_path = listener.local_unix_path.clone();
1095 client.listener_state = None;
1096 client.connection_state = Some(ConnectionState {
1097 peer_socket_id: Some(accepted_socket_id),
1098 ..ConnectionState::default()
1099 });
1100
1101 Ok(accepted)
1102 })();
1103
1104 match result {
1105 Ok(accepted) => {
1106 table.sockets.insert(socket_id, client);
1107 table.sockets.insert(accepted_socket_id, accepted.clone());
1108 table
1109 .by_owner
1110 .entry(accepted.owner_pid)
1111 .or_default()
1112 .insert(accepted_socket_id);
1113 Ok(())
1114 }
1115 Err(error) => {
1116 table.sockets.insert(socket_id, client);
1117 Err(error)
1118 }
1119 }
1120 }
1121
1122 pub fn send_to_bound_udp_socket(
1123 &self,
1124 socket_id: SocketId,
1125 target_address: InetSocketAddress,
1126 data: &[u8],
1127 ) -> SocketResult<usize> {
1128 let target_address = normalize_inet_address(target_address);
1129 let mut table = lock_or_recover(&self.inner.state);
1130 let sender = table
1131 .sockets
1132 .get(&socket_id)
1133 .cloned()
1134 .ok_or_else(|| SocketTableError::not_found(socket_id))?;
1135 validate_bound_udp_sender(&sender)?;
1136
1137 let receiver_socket_id = table
1138 .bound_inet_datagrams
1139 .get(&target_address)
1140 .and_then(|socket_ids| socket_ids.first().copied())
1141 .ok_or_else(|| {
1142 SocketTableError::not_found_address(format!(
1143 "no UDP socket bound at {}:{}",
1144 target_address.host(),
1145 target_address.port()
1146 ))
1147 })?;
1148 let receiver = table
1149 .sockets
1150 .get_mut(&receiver_socket_id)
1151 .ok_or_else(|| SocketTableError::not_found(receiver_socket_id))?;
1152 validate_bound_udp_receiver(receiver)?;
1153
1154 let datagram_state = receiver.datagram_state.as_mut().ok_or_else(|| {
1155 SocketTableError::invalid_argument(format!(
1156 "socket {receiver_socket_id} does not support datagrams"
1157 ))
1158 })?;
1159 datagram_state.recv_queue.push_back(QueuedDatagram {
1160 source_address: sender.local_address.clone(),
1161 payload: data.to_vec(),
1162 });
1163 Ok(data.len())
1164 }
1165
1166 pub fn recv_datagram(
1167 &self,
1168 socket_id: SocketId,
1169 max_bytes: usize,
1170 ) -> SocketResult<Option<ReceivedDatagram>> {
1171 let mut table = lock_or_recover(&self.inner.state);
1172 let record = table
1173 .sockets
1174 .get_mut(&socket_id)
1175 .ok_or_else(|| SocketTableError::not_found(socket_id))?;
1176 validate_bound_udp_receiver(record)?;
1177
1178 let datagram_state = record.datagram_state.as_mut().ok_or_else(|| {
1179 SocketTableError::invalid_argument(format!(
1180 "socket {socket_id} does not support datagrams"
1181 ))
1182 })?;
1183 let Some(datagram) = datagram_state.recv_queue.pop_front() else {
1184 return Err(SocketTableError::would_block(format!(
1185 "socket {socket_id} has no queued datagrams"
1186 )));
1187 };
1188
1189 let payload = if datagram.payload.len() > max_bytes {
1190 datagram.payload[..max_bytes].to_vec()
1191 } else {
1192 datagram.payload
1193 };
1194 Ok(Some(ReceivedDatagram {
1195 source_address: datagram.source_address,
1196 payload,
1197 }))
1198 }
1199
1200 pub fn poll(&self, socket_id: SocketId, requested: PollEvents) -> SocketResult<PollEvents> {
1201 let table = lock_or_recover(&self.inner.state);
1202 let record = table
1203 .sockets
1204 .get(&socket_id)
1205 .ok_or_else(|| SocketTableError::not_found(socket_id))?;
1206
1207 let mut events = PollEvents::empty();
1208 match record.state {
1209 SocketState::Listening => {
1210 if requested.intersects(POLLIN) && record.pending_accept_count() > 0 {
1211 events |= POLLIN;
1212 }
1213 }
1214 SocketState::Connected => {
1215 let connection = record.connection_state.as_ref().ok_or_else(|| {
1216 SocketTableError::not_connected(format!("socket {socket_id} is not connected"))
1217 })?;
1218 let peer = connection
1219 .peer_socket_id
1220 .and_then(|peer_socket_id| table.sockets.get(&peer_socket_id));
1221
1222 if requested.intersects(POLLIN) && !connection.recv_buffer.is_empty() {
1223 events |= POLLIN;
1224 }
1225 if connection.peer_write_shutdown || peer.is_none() {
1226 events |= POLLHUP;
1227 }
1228
1229 if requested.intersects(POLLOUT) && !connection.write_shutdown {
1230 if peer
1231 .and_then(|peer| peer.connection_state.as_ref())
1232 .map(|peer_connection| peer_connection.read_shutdown)
1233 .unwrap_or(true)
1234 {
1235 events |= POLLERR;
1236 } else {
1237 events |= POLLOUT;
1238 }
1239 }
1240 }
1241 SocketState::Bound if supports_inet_datagram_lifecycle(record.spec) => {
1242 let datagram_state = record.datagram_state.as_ref().ok_or_else(|| {
1243 SocketTableError::invalid_argument(format!(
1244 "socket {socket_id} does not support datagrams"
1245 ))
1246 })?;
1247 if requested.intersects(POLLIN) && !datagram_state.recv_queue.is_empty() {
1248 events |= POLLIN;
1249 }
1250 if requested.intersects(POLLOUT) {
1251 events |= POLLOUT;
1252 }
1253 }
1254 SocketState::Created | SocketState::Bound => {}
1255 }
1256
1257 Ok(events)
1258 }
1259
1260 pub fn write(&self, socket_id: SocketId, data: &[u8]) -> SocketResult<usize> {
1261 let mut table = lock_or_recover(&self.inner.state);
1262 let record = table
1263 .sockets
1264 .get(&socket_id)
1265 .cloned()
1266 .ok_or_else(|| SocketTableError::not_found(socket_id))?;
1267 let connection = record.connection_state.as_ref().ok_or_else(|| {
1268 SocketTableError::not_connected(format!("socket {socket_id} is not connected"))
1269 })?;
1270 if record.state != SocketState::Connected {
1271 return Err(SocketTableError::not_connected(format!(
1272 "socket {socket_id} is not connected"
1273 )));
1274 }
1275 if connection.write_shutdown {
1276 return Err(SocketTableError::broken_pipe(format!(
1277 "socket {socket_id} write side is shut down"
1278 )));
1279 }
1280
1281 let peer_socket_id = connection.peer_socket_id.ok_or_else(|| {
1282 SocketTableError::broken_pipe(format!("socket {socket_id} peer is closed"))
1283 })?;
1284 let peer = table.sockets.get_mut(&peer_socket_id).ok_or_else(|| {
1285 SocketTableError::broken_pipe(format!("socket {socket_id} peer is closed"))
1286 })?;
1287 let peer_connection = peer.connection_state.as_mut().ok_or_else(|| {
1288 SocketTableError::broken_pipe(format!("socket {socket_id} peer is closed"))
1289 })?;
1290 if peer_connection.read_shutdown {
1291 return Err(SocketTableError::broken_pipe(format!(
1292 "socket {peer_socket_id} read side is shut down"
1293 )));
1294 }
1295
1296 peer_connection.recv_buffer.extend(data.iter().copied());
1297 Ok(data.len())
1298 }
1299
1300 pub fn read(&self, socket_id: SocketId, max_bytes: usize) -> SocketResult<Option<Vec<u8>>> {
1301 if max_bytes == 0 {
1302 return Ok(Some(Vec::new()));
1303 }
1304
1305 let mut table = lock_or_recover(&self.inner.state);
1306 let record = table
1307 .sockets
1308 .get(&socket_id)
1309 .cloned()
1310 .ok_or_else(|| SocketTableError::not_found(socket_id))?;
1311 if record.state != SocketState::Connected {
1312 return Err(SocketTableError::not_connected(format!(
1313 "socket {socket_id} is not connected"
1314 )));
1315 }
1316
1317 let connection = record.connection_state.as_ref().ok_or_else(|| {
1318 SocketTableError::not_connected(format!("socket {socket_id} is not connected"))
1319 })?;
1320 if connection.read_shutdown {
1321 return Ok(None);
1322 }
1323 if !connection.recv_buffer.is_empty() {
1324 let record = table
1325 .sockets
1326 .get_mut(&socket_id)
1327 .ok_or_else(|| SocketTableError::not_found(socket_id))?;
1328 let connection = record.connection_state.as_mut().ok_or_else(|| {
1329 SocketTableError::not_connected(format!("socket {socket_id} is not connected"))
1330 })?;
1331 let read_len = connection.recv_buffer.len().min(max_bytes);
1332 let bytes = connection.recv_buffer.drain(..read_len).collect::<Vec<_>>();
1333 return Ok(Some(bytes));
1334 }
1335
1336 let peer_open = connection
1337 .peer_socket_id
1338 .map(|peer_socket_id| table.sockets.contains_key(&peer_socket_id))
1339 .unwrap_or(false);
1340 if connection.peer_write_shutdown || !peer_open {
1341 return Ok(None);
1342 }
1343
1344 Err(SocketTableError::would_block(format!(
1345 "socket {socket_id} has no readable data"
1346 )))
1347 }
1348
1349 pub fn shutdown(&self, socket_id: SocketId, how: SocketShutdown) -> SocketResult<SocketRecord> {
1350 let mut table = lock_or_recover(&self.inner.state);
1351 let record = table
1352 .sockets
1353 .remove(&socket_id)
1354 .ok_or_else(|| SocketTableError::not_found(socket_id))?;
1355
1356 if record.state != SocketState::Connected {
1357 table.sockets.insert(socket_id, record);
1358 return Err(SocketTableError::not_connected(format!(
1359 "socket {socket_id} is not connected"
1360 )));
1361 }
1362
1363 let Some(mut connection) = record.connection_state.clone() else {
1364 table.sockets.insert(socket_id, record);
1365 return Err(SocketTableError::not_connected(format!(
1366 "socket {socket_id} is not connected"
1367 )));
1368 };
1369
1370 if matches!(how, SocketShutdown::Read | SocketShutdown::Both) {
1371 connection.recv_buffer.clear();
1372 connection.read_shutdown = true;
1373 }
1374 if matches!(how, SocketShutdown::Write | SocketShutdown::Both) {
1375 connection.write_shutdown = true;
1376 if let Some(peer_socket_id) = connection.peer_socket_id {
1377 if let Some(peer) = table.sockets.get_mut(&peer_socket_id) {
1378 if let Some(peer_connection) = peer.connection_state.as_mut() {
1379 peer_connection.peer_write_shutdown = true;
1380 }
1381 }
1382 }
1383 }
1384
1385 let mut record = record;
1386 record.connection_state = Some(connection);
1387 let cloned = record.clone();
1388 table.sockets.insert(socket_id, record);
1389 Ok(cloned)
1390 }
1391
1392 pub fn remove(&self, socket_id: SocketId) -> SocketResult<SocketRecord> {
1393 let mut table = lock_or_recover(&self.inner.state);
1394 remove_socket(&mut table, socket_id).ok_or_else(|| SocketTableError::not_found(socket_id))
1395 }
1396
1397 pub fn remove_all_for_pid(&self, owner_pid: u32) -> Vec<SocketRecord> {
1398 let mut table = lock_or_recover(&self.inner.state);
1399 let Some(socket_ids) = table.by_owner.remove(&owner_pid) else {
1400 return Vec::new();
1401 };
1402
1403 socket_ids
1404 .into_iter()
1405 .filter_map(|socket_id| remove_socket(&mut table, socket_id))
1406 .collect()
1407 }
1408
1409 pub fn snapshot(&self) -> SocketTableSnapshot {
1410 let table = lock_or_recover(&self.inner.state);
1411 let mut snapshot = SocketTableSnapshot {
1412 sockets: table.sockets.len(),
1413 ..SocketTableSnapshot::default()
1414 };
1415 for record in table.sockets.values() {
1416 if record.state.counts_as_listener() {
1417 snapshot.listeners += 1;
1418 }
1419 if record.state.counts_as_connection() {
1420 snapshot.connections += 1;
1421 }
1422 }
1423 snapshot
1424 }
1425}
1426
1427fn next_socket_id(table: &mut SocketTableState) -> SocketId {
1428 if table.next_socket_id == 0 {
1429 table.next_socket_id = 1;
1430 }
1431 let socket_id = table.next_socket_id;
1432 table.next_socket_id = table.next_socket_id.saturating_add(1);
1433 socket_id
1434}
1435
1436fn validate_state_transition(current: SocketState, next: SocketState) -> SocketResult<()> {
1437 if current == SocketState::Connected && next != SocketState::Connected {
1438 return Err(SocketTableError::invalid_argument(format!(
1439 "invalid socket state transition from {current:?} to {next:?}"
1440 )));
1441 }
1442 Ok(())
1443}
1444
1445fn validate_connect_pair(socket: &SocketRecord, peer: &SocketRecord) -> SocketResult<()> {
1446 if !supports_connection_lifecycle(socket.spec) {
1447 return Err(SocketTableError::invalid_argument(format!(
1448 "socket {} does not support stream connections",
1449 socket.id
1450 )));
1451 }
1452 if !supports_connection_lifecycle(peer.spec) {
1453 return Err(SocketTableError::invalid_argument(format!(
1454 "socket {} does not support stream connections",
1455 peer.id
1456 )));
1457 }
1458 if !matches!(socket.state, SocketState::Created | SocketState::Bound) {
1459 return Err(SocketTableError::invalid_argument(format!(
1460 "socket {} cannot connect in state {:?}",
1461 socket.id, socket.state
1462 )));
1463 }
1464 if !matches!(peer.state, SocketState::Created | SocketState::Bound) {
1465 return Err(SocketTableError::invalid_argument(format!(
1466 "socket {} cannot connect in state {:?}",
1467 peer.id, peer.state
1468 )));
1469 }
1470 Ok(())
1471}
1472
1473fn default_connection_state(spec: SocketSpec, state: SocketState) -> Option<ConnectionState> {
1474 if state == SocketState::Connected && supports_connection_lifecycle(spec) {
1475 Some(ConnectionState::default())
1476 } else {
1477 None
1478 }
1479}
1480
1481fn default_datagram_state(spec: SocketSpec) -> Option<DatagramState> {
1482 if supports_inet_datagram_lifecycle(spec) {
1483 Some(DatagramState::default())
1484 } else {
1485 None
1486 }
1487}
1488
1489fn supports_connection_lifecycle(spec: SocketSpec) -> bool {
1490 matches!(spec.socket_type, SocketType::Stream)
1491}
1492
1493fn supports_listener_lifecycle(spec: SocketSpec) -> bool {
1494 matches!(spec.socket_type, SocketType::Stream)
1495 && matches!(
1496 spec.domain,
1497 SocketDomain::Inet | SocketDomain::Inet6 | SocketDomain::Unix
1498 )
1499}
1500
1501fn supports_inet_bind(spec: SocketSpec) -> bool {
1502 matches!(spec.domain, SocketDomain::Inet | SocketDomain::Inet6)
1503 && matches!(spec.socket_type, SocketType::Stream | SocketType::Datagram)
1504}
1505
1506fn supports_unix_stream_lifecycle(spec: SocketSpec) -> bool {
1507 matches!(spec.socket_type, SocketType::Stream) && matches!(spec.domain, SocketDomain::Unix)
1508}
1509
1510fn supports_inet_stream_lifecycle(spec: SocketSpec) -> bool {
1511 matches!(spec.socket_type, SocketType::Stream)
1512 && matches!(spec.domain, SocketDomain::Inet | SocketDomain::Inet6)
1513}
1514
1515fn supports_inet_datagram_lifecycle(spec: SocketSpec) -> bool {
1516 matches!(spec.socket_type, SocketType::Datagram)
1517 && matches!(spec.domain, SocketDomain::Inet | SocketDomain::Inet6)
1518}
1519
1520fn lookup_conflicting_bound_inet_socket_ids(
1521 table: &SocketTableState,
1522 spec: SocketSpec,
1523 address: &InetSocketAddress,
1524) -> Vec<SocketId> {
1525 if supports_inet_stream_lifecycle(spec) {
1526 table
1527 .bound_inet_streams
1528 .iter()
1529 .find_map(|(bound_address, socket_id)| {
1530 inet_stream_bind_addresses_overlap(bound_address, address).then_some(*socket_id)
1531 })
1532 .into_iter()
1533 .collect()
1534 } else if supports_inet_datagram_lifecycle(spec) {
1535 table
1536 .bound_inet_datagrams
1537 .iter()
1538 .filter(|(bound_address, _)| inet_stream_bind_addresses_overlap(bound_address, address))
1539 .flat_map(|(_, socket_ids)| socket_ids.iter().copied())
1540 .collect()
1541 } else {
1542 Vec::new()
1543 }
1544}
1545
1546fn lookup_bound_inet_socket(
1547 table: &SocketTableState,
1548 spec: SocketSpec,
1549 address: &InetSocketAddress,
1550) -> Option<SocketId> {
1551 if supports_inet_stream_lifecycle(spec) {
1552 lookup_bound_inet_socket_in_table(&table.bound_inet_streams, address)
1553 } else if supports_inet_datagram_lifecycle(spec) {
1554 lookup_bound_inet_datagram_socket_in_table(&table.bound_inet_datagrams, address)
1555 } else {
1556 None
1557 }
1558}
1559
1560fn inet_stream_bind_addresses_overlap(
1561 existing: &InetSocketAddress,
1562 requested: &InetSocketAddress,
1563) -> bool {
1564 if existing == requested {
1565 return true;
1566 }
1567
1568 wildcard_inet_address(existing).as_ref() == Some(requested)
1569 || wildcard_inet_address(requested).as_ref() == Some(existing)
1570}
1571
1572fn lookup_bound_inet_socket_in_table(
1573 sockets: &BTreeMap<InetSocketAddress, SocketId>,
1574 address: &InetSocketAddress,
1575) -> Option<SocketId> {
1576 sockets.get(address).copied().or_else(|| {
1577 wildcard_inet_address(address).and_then(|wildcard| sockets.get(&wildcard).copied())
1578 })
1579}
1580
1581fn lookup_bound_inet_datagram_socket_in_table(
1582 sockets: &BTreeMap<InetSocketAddress, BTreeSet<SocketId>>,
1583 address: &InetSocketAddress,
1584) -> Option<SocketId> {
1585 sockets
1586 .get(address)
1587 .and_then(|socket_ids| socket_ids.first().copied())
1588 .or_else(|| {
1589 wildcard_inet_address(address).and_then(|wildcard| {
1590 sockets
1591 .get(&wildcard)
1592 .and_then(|socket_ids| socket_ids.first().copied())
1593 })
1594 })
1595}
1596
1597fn register_bound_inet_socket(
1598 table: &mut SocketTableState,
1599 spec: SocketSpec,
1600 address: InetSocketAddress,
1601 socket_id: SocketId,
1602) {
1603 if supports_inet_stream_lifecycle(spec) {
1604 table.bound_inet_streams.insert(address, socket_id);
1605 } else if supports_inet_datagram_lifecycle(spec) {
1606 table
1607 .bound_inet_datagrams
1608 .entry(address)
1609 .or_default()
1610 .insert(socket_id);
1611 }
1612}
1613
1614fn validate_connect_to_listener(
1615 client: &SocketRecord,
1616 listener: &SocketRecord,
1617) -> SocketResult<()> {
1618 if !supports_connection_lifecycle(client.spec) {
1619 return Err(SocketTableError::invalid_argument(format!(
1620 "socket {} does not support stream connections",
1621 client.id
1622 )));
1623 }
1624 if !supports_listener_lifecycle(listener.spec) {
1625 return Err(SocketTableError::invalid_argument(format!(
1626 "socket {} is not a stream listener",
1627 listener.id
1628 )));
1629 }
1630 if !matches!(client.state, SocketState::Created | SocketState::Bound) {
1631 return Err(SocketTableError::invalid_argument(format!(
1632 "socket {} cannot connect in state {:?}",
1633 client.id, client.state
1634 )));
1635 }
1636 if listener.state != SocketState::Listening {
1637 return Err(SocketTableError::invalid_argument(format!(
1638 "socket {} is not listening",
1639 listener.id
1640 )));
1641 }
1642 Ok(())
1643}
1644
1645fn has_bound_endpoint(record: &SocketRecord) -> bool {
1646 record.local_address.is_some() || record.local_unix_path.is_some()
1647}
1648
1649fn validate_bound_udp_sender(sender: &SocketRecord) -> SocketResult<()> {
1650 if !supports_inet_datagram_lifecycle(sender.spec) {
1651 return Err(SocketTableError::invalid_argument(format!(
1652 "socket {} is not an INET datagram socket",
1653 sender.id
1654 )));
1655 }
1656 if sender.state != SocketState::Bound || sender.local_address.is_none() {
1657 return Err(SocketTableError::invalid_argument(format!(
1658 "socket {} must be bound before sending datagrams",
1659 sender.id
1660 )));
1661 }
1662 Ok(())
1663}
1664
1665fn validate_bound_udp_receiver(receiver: &SocketRecord) -> SocketResult<()> {
1666 if !supports_inet_datagram_lifecycle(receiver.spec) {
1667 return Err(SocketTableError::invalid_argument(format!(
1668 "socket {} is not an INET datagram socket",
1669 receiver.id
1670 )));
1671 }
1672 if receiver.state != SocketState::Bound || receiver.local_address.is_none() {
1673 return Err(SocketTableError::invalid_argument(format!(
1674 "socket {} must be bound to receive datagrams",
1675 receiver.id
1676 )));
1677 }
1678 Ok(())
1679}
1680
1681fn datagram_state_mut(record: &mut SocketRecord) -> SocketResult<&mut DatagramState> {
1682 if !supports_inet_datagram_lifecycle(record.spec) {
1683 return Err(SocketTableError::invalid_argument(format!(
1684 "socket {} is not an INET datagram socket",
1685 record.id
1686 )));
1687 }
1688 record.datagram_state.as_mut().ok_or_else(|| {
1689 SocketTableError::invalid_argument(format!(
1690 "socket {} does not support datagrams",
1691 record.id
1692 ))
1693 })
1694}
1695
1696fn validate_multicast_socket(record: &SocketRecord) -> SocketResult<()> {
1697 validate_bound_udp_receiver(record)?;
1698 if record.spec.domain != SocketDomain::Inet {
1699 return Err(SocketTableError::invalid_argument(format!(
1700 "socket {} multicast membership is only implemented for IPv4 datagrams",
1701 record.id
1702 )));
1703 }
1704 Ok(())
1705}
1706
1707fn normalize_multicast_membership(
1708 spec: SocketSpec,
1709 membership: SocketMulticastMembership,
1710) -> SocketResult<SocketMulticastMembership> {
1711 let group_address = membership.group_address.trim().to_ascii_lowercase();
1712 let interface_address = membership
1713 .interface_address
1714 .map(|value| value.trim().to_ascii_lowercase())
1715 .filter(|value| !value.is_empty());
1716
1717 match spec.domain {
1718 SocketDomain::Inet => {
1719 let parsed = group_address.parse::<Ipv4Addr>().map_err(|_| {
1720 SocketTableError::invalid_argument(format!(
1721 "invalid IPv4 multicast address {group_address}"
1722 ))
1723 })?;
1724 if !parsed.is_multicast() {
1725 return Err(SocketTableError::invalid_argument(format!(
1726 "address {group_address} is not an IPv4 multicast group"
1727 )));
1728 }
1729 }
1730 SocketDomain::Inet6 => {
1731 let parsed = group_address.parse::<Ipv6Addr>().map_err(|_| {
1732 SocketTableError::invalid_argument(format!(
1733 "invalid IPv6 multicast address {group_address}"
1734 ))
1735 })?;
1736 if !parsed.is_multicast() {
1737 return Err(SocketTableError::invalid_argument(format!(
1738 "address {group_address} is not an IPv6 multicast group"
1739 )));
1740 }
1741 }
1742 SocketDomain::Unix => {
1743 return Err(SocketTableError::invalid_argument(
1744 "unix sockets do not support multicast membership",
1745 ));
1746 }
1747 }
1748
1749 Ok(SocketMulticastMembership::new(
1750 group_address,
1751 interface_address,
1752 ))
1753}
1754
1755fn has_incompatible_inet_bind_conflict(
1756 table: &SocketTableState,
1757 record: &SocketRecord,
1758 conflicting_ids: &[SocketId],
1759) -> bool {
1760 conflicting_ids.iter().any(|conflicting_id| {
1761 if *conflicting_id == record.id {
1762 return false;
1763 }
1764
1765 let Some(existing) = table.sockets.get(conflicting_id) else {
1766 return false;
1767 };
1768
1769 if supports_inet_datagram_lifecycle(record.spec) {
1770 !inet_datagram_bind_shares_port(record, existing)
1771 } else {
1772 true
1773 }
1774 })
1775}
1776
1777fn inet_datagram_bind_shares_port(requested: &SocketRecord, existing: &SocketRecord) -> bool {
1778 (requested.reuse_port() && existing.reuse_port())
1779 || (requested.reuse_address() && existing.reuse_address())
1780}
1781
1782fn remove_socket(table: &mut SocketTableState, socket_id: SocketId) -> Option<SocketRecord> {
1783 let record = table.sockets.remove(&socket_id)?;
1784 unregister_bound_socket(table, &record);
1785 unregister_multicast_memberships(table, &record);
1786 if let Some(listener_state) = record.listener_state.as_ref() {
1787 let pending_socket_ids = listener_state
1788 .pending_accepts
1789 .iter()
1790 .filter_map(|pending| pending.accepted_socket_id)
1791 .collect::<Vec<_>>();
1792 for pending_socket_id in pending_socket_ids {
1793 let _ = remove_socket(table, pending_socket_id);
1794 }
1795 }
1796 if let Some(connection) = record.connection_state.as_ref() {
1797 if let Some(peer_socket_id) = connection.peer_socket_id {
1798 if let Some(peer) = table.sockets.get_mut(&peer_socket_id) {
1799 if let Some(peer_connection) = peer.connection_state.as_mut() {
1800 if peer_connection.peer_socket_id == Some(socket_id) {
1801 peer_connection.peer_socket_id = None;
1802 }
1803 peer_connection.peer_write_shutdown = true;
1804 }
1805 }
1806 }
1807 }
1808 if let Some(owner_sockets) = table.by_owner.get_mut(&record.owner_pid) {
1809 owner_sockets.remove(&socket_id);
1810 if owner_sockets.is_empty() {
1811 table.by_owner.remove(&record.owner_pid);
1812 }
1813 }
1814 Some(record)
1815}
1816
1817fn unregister_bound_socket(table: &mut SocketTableState, record: &SocketRecord) {
1818 let Some(address) = record.local_address.as_ref() else {
1819 if supports_unix_stream_lifecycle(record.spec) {
1820 if let Some(path) = record.local_unix_path.as_ref() {
1821 if table.bound_unix_streams.get(path).copied() == Some(record.id) {
1822 table.bound_unix_streams.remove(path);
1823 }
1824 }
1825 }
1826 return;
1827 };
1828 if supports_inet_stream_lifecycle(record.spec)
1829 && table.bound_inet_streams.get(address).copied() == Some(record.id)
1830 {
1831 table.bound_inet_streams.remove(address);
1832 }
1833 if supports_inet_datagram_lifecycle(record.spec) {
1834 if let Some(socket_ids) = table.bound_inet_datagrams.get_mut(address) {
1835 socket_ids.remove(&record.id);
1836 if socket_ids.is_empty() {
1837 table.bound_inet_datagrams.remove(address);
1838 }
1839 }
1840 }
1841}
1842
1843fn unregister_multicast_memberships(table: &mut SocketTableState, record: &SocketRecord) {
1844 let Some(datagram_state) = record.datagram_state.as_ref() else {
1845 return;
1846 };
1847
1848 for membership in &datagram_state.multicast_memberships {
1849 if let Some(socket_ids) = table.multicast_groups.get_mut(membership) {
1850 socket_ids.remove(&record.id);
1851 if socket_ids.is_empty() {
1852 table.multicast_groups.remove(membership);
1853 }
1854 }
1855 }
1856}
1857
1858fn normalize_inet_address(address: InetSocketAddress) -> InetSocketAddress {
1859 match address.host().to_ascii_lowercase().as_str() {
1860 "localhost" => InetSocketAddress::new("127.0.0.1", address.port()),
1861 _ => address,
1862 }
1863}
1864
1865fn wildcard_inet_address(address: &InetSocketAddress) -> Option<InetSocketAddress> {
1866 match address.host() {
1867 "127.0.0.1" => Some(InetSocketAddress::new("0.0.0.0", address.port())),
1868 "::1" => Some(InetSocketAddress::new("::", address.port())),
1869 _ => None,
1870 }
1871}
1872
1873fn normalize_unix_socket_path(path: impl AsRef<str>) -> SocketResult<String> {
1874 let normalized = normalize_path(path.as_ref());
1875 if normalized == "/" {
1876 return Err(SocketTableError::invalid_argument(
1877 "unix socket path must not be empty or root",
1878 ));
1879 }
1880 Ok(normalized)
1881}
1882
1883fn lock_or_recover<'a, T>(mutex: &'a Mutex<T>) -> MutexGuard<'a, T> {
1884 match mutex.lock() {
1885 Ok(guard) => guard,
1886 Err(poisoned) => poisoned.into_inner(),
1887 }
1888}