Skip to main content

agent_os_kernel/
socket_table.rs

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}