Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
107 changes: 66 additions & 41 deletions kernel/src/net/socket/inet/common/port.rs
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@ use system_error::SystemError;
use crate::{
arch::rand::rand,
libs::mutex::Mutex,
process::{ProcessManager, RawPid},
process::ProcessManager,
};

use super::Types::{self, *};
Expand All @@ -16,8 +16,8 @@ use super::Types::{self, *};
/// 如果 TCP/UDP 的 socket 绑定了某个端口,它会在对应的表中记录,以检测端口冲突。
#[derive(Debug)]
pub struct PortManager {
// TCP 端口记录表
tcp_port_table: Mutex<HashMap<u16, RawPid>>,
// TCP 端口记录表。一个端口可以有多条绑定记录(SO_REUSEPORT/SO_REUSEADDR 共享)。
tcp_port_table: Mutex<HashMap<u16, Vec<TcpPortBinding>>>,
// UDP 端口记录表
udp_port_table: Mutex<HashMap<u16, Vec<UdpPortBinding>>>,
}
Expand Down Expand Up @@ -95,16 +95,23 @@ impl PortManager {
}

#[inline]
pub fn bind_ephemeral_port(&self, socket_type: Types) -> Result<u16, SystemError> {
pub fn bind_tcp_ephemeral_port(
&self,
addr: IpAddress,
reuseaddr: bool,
reuseport: bool,
iface_nic_id: usize,
handle: smoltcp::iface::SocketHandle,
) -> Result<u16, SystemError> {
let (min, max) = Self::local_port_range();
let range = (max - min) as u32 + 1;
if range == 0 {
return Err(SystemError::EINVAL);
}
let mut remaining = range;
while remaining > 0 {
let port = self.get_ephemeral_port(socket_type)?;
match self.bind_port(socket_type, port) {
let port = self.get_ephemeral_port(Types::Tcp)?;
match self.bind_tcp_port(port, addr, reuseaddr, reuseport, iface_nic_id, handle) {
Ok(()) => return Ok(port),
Err(SystemError::EADDRINUSE) => {
// Race: another thread grabbed the port after we checked.
Expand Down Expand Up @@ -146,44 +153,52 @@ impl PortManager {
Err(SystemError::EADDRINUSE)
}

/// @brief 检测给定端口是否已被占用,如果未被占用则在 TCP 对应的表中记录
/// TCP: 绑定端口,支持 SO_REUSEADDR/SO_REUSEPORT。
///
/// UDP 复用逻辑请使用 `bind_udp_port`
pub fn bind_port(&self, socket_type: Types, port: u16) -> Result<(), SystemError> {
if port > 0 {
match socket_type {
Udp => {
let mut guard = self.udp_port_table.lock();
if guard.get(&port).is_some() {
return Err(SystemError::EADDRINUSE);
}
guard.insert(port, Vec::new());
}
Tcp => {
let mut guard = self.tcp_port_table.lock();
if guard.get(&port).is_some() {
return Err(SystemError::EADDRINUSE);
}
guard.insert(port, ProcessManager::current_pid());
}
_ => {}
};
/// 一条绑定记录以 `(iface_nic_id, handle)` 唯一标识(BoundInner 身份),
/// 因此多个进程/多个 socket 可以共享同一端口而不需要调用方保存额外 id。
pub fn bind_tcp_port(
&self,
port: u16,
addr: IpAddress,
reuseaddr: bool,
reuseport: bool,
iface_nic_id: usize,
handle: smoltcp::iface::SocketHandle,
) -> Result<(), SystemError> {
if port == 0 {
return Err(SystemError::EINVAL);
}
let mut guard = self.tcp_port_table.lock();
let bindings = guard.entry(port).or_default();
for binding in bindings.iter() {
if !addrs_conflict(addr, binding.addr) {
continue;
}
let share_ok = (reuseport && binding.reuseport) || (reuseaddr && binding.reuseaddr);
if !share_ok {
return Err(SystemError::EADDRINUSE);
}
}
return Ok(());
bindings.push(TcpPortBinding {
addr,
reuseaddr,
reuseport,
iface_nic_id,
handle,
});
Ok(())
}

/// @brief 在对应的端口记录表中将端口和 socket 解绑
/// should call this function when socket is closed or aborted
pub fn unbind_port(&self, socket_type: Types, port: u16) {
match socket_type {
Udp => {
self.udp_port_table.lock().remove(&port);
}
Tcp => {
self.tcp_port_table.lock().remove(&port);
/// TCP: 解绑端口(按 BoundInner 身份)
pub fn unbind_tcp_port(&self, port: u16, iface_nic_id: usize, handle: smoltcp::iface::SocketHandle) {
let mut guard = self.tcp_port_table.lock();
if let Some(list) = guard.get_mut(&port) {
list.retain(|b| b.iface_nic_id != iface_nic_id || b.handle != handle);
if list.is_empty() {
guard.remove(&port);
}
_ => {}
};
}
}

/// UDP: 绑定端口,支持 SO_REUSEADDR/SO_REUSEPORT
Expand All @@ -201,7 +216,7 @@ impl PortManager {
let mut guard = self.udp_port_table.lock();
let bindings = guard.entry(port).or_default();
for binding in bindings.iter() {
if !udp_addrs_conflict(addr, binding.addr) {
if !addrs_conflict(addr, binding.addr) {
continue;
}
let share_ok = (reuseport && binding.reuseport) || (reuseaddr && binding.reuseaddr);
Expand Down Expand Up @@ -238,8 +253,18 @@ struct UdpPortBinding {
bind_id: usize,
}

/// TCP 端口绑定记录。`(iface_nic_id, handle)` 是绑定的 BoundInner 身份。
#[derive(Debug, Clone)]
struct TcpPortBinding {
addr: IpAddress,
reuseaddr: bool,
reuseport: bool,
iface_nic_id: usize,
handle: smoltcp::iface::SocketHandle,
}

#[inline]
fn udp_addrs_conflict(a: IpAddress, b: IpAddress) -> bool {
fn addrs_conflict(a: IpAddress, b: IpAddress) -> bool {
if a.version() != b.version() {
return false;
}
Expand Down
2 changes: 1 addition & 1 deletion kernel/src/net/socket/inet/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@ pub mod stream;
pub mod syscall;

pub use common::BoundInner;
pub use common::Types;

pub use datagram::UdpSocket;
pub use raw::RawSocket;

Expand Down
71 changes: 51 additions & 20 deletions kernel/src/net/socket/inet/stream/inner.rs
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@ use core::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use crate::filesystem::epoll::EPollEventType;
use crate::libs::mutex::Mutex;
use crate::libs::rwsem::RwSem;
use crate::net::socket::{self, inet::Types};
use crate::net::socket::{self};
use crate::process::namespace::net_namespace::NetNamespace;
use crate::syscall::user_buffer::UserBuffer;
use alloc::boxed::Box;
Expand Down Expand Up @@ -51,11 +51,13 @@ fn new_smoltcp_socket() -> smoltcp::socket::tcp::Socket<'static> {

fn new_listen_smoltcp_socket<T>(
local_endpoint: T,
reuseport: bool,
) -> Result<smoltcp::socket::tcp::Socket<'static>, SystemError>
where
T: Into<smoltcp::wire::IpListenEndpoint>,
{
let mut socket = new_smoltcp_socket();
socket.set_reuseport(reuseport);
socket.listen(local_endpoint).map_err(|e| match e {
tcp::ListenError::InvalidState => SystemError::EINVAL, // TODO: Check is right impl
tcp::ListenError::Unaddressable => SystemError::EADDRINUSE,
Expand Down Expand Up @@ -112,6 +114,8 @@ impl Init {
self,
local_endpoint: smoltcp::wire::IpEndpoint,
netns: Arc<NetNamespace>,
reuseaddr: bool,
reuseport: bool,
) -> Result<Self, (Self, SystemError)> {
match self {
Init::Unbound((socket, ver)) => {
Expand All @@ -128,7 +132,13 @@ impl Init {

// Handle ephemeral port assignment (port 0)
let bind_port = if local_endpoint.port == 0 {
match bound.port_manager().bind_ephemeral_port(Types::Tcp) {
match bound.port_manager().bind_tcp_ephemeral_port(
local_endpoint.addr,
reuseaddr,
reuseport,
bound.iface().nic_id(),
bound.handle(),
) {
Ok(port) => port,
Err(err) => {
let smoltcp::socket::Socket::Tcp(socket) = bound.into_socket() else {
Expand All @@ -138,10 +148,14 @@ impl Init {
}
}
} else {
if let Err(err) = bound
.port_manager()
.bind_port(Types::Tcp, local_endpoint.port)
{
if let Err(err) = bound.port_manager().bind_tcp_port(
local_endpoint.port,
local_endpoint.addr,
reuseaddr,
reuseport,
bound.iface().nic_id(),
bound.handle(),
) {
let smoltcp::socket::Socket::Tcp(socket) = bound.into_socket() else {
unreachable!("TCP BoundInner should contain a TCP socket");
};
Expand Down Expand Up @@ -178,7 +192,13 @@ impl Init {
return Err((Self::Unbound((Box::new(socket), ver)), err))
}
};
let bound_port = match bound.port_manager().bind_ephemeral_port(Types::Tcp) {
let bound_port = match bound.port_manager().bind_tcp_ephemeral_port(
address,
false,
false,
bound.iface().nic_id(),
bound.handle(),
) {
Ok(port) => port,
Err(err) => {
let smoltcp::socket::Socket::Tcp(socket) = bound.into_socket() else {
Expand Down Expand Up @@ -230,6 +250,8 @@ impl Init {
self,
backlog: usize,
netns: Arc<NetNamespace>,
reuseaddr: bool,
reuseport: bool,
) -> Result<Listening, (Self, SystemError)> {
// If unbound, auto-bind to INADDR_ANY:ephemeral (Linux compat).
let bound_self = if matches!(self, Init::Unbound(_)) {
Expand All @@ -246,7 +268,7 @@ impl Init {
}
};
let auto_bind_ep = smoltcp::wire::IpEndpoint::new(unspec_addr, 0);
match self.bind(auto_bind_ep, netns.clone()) {
match self.bind(auto_bind_ep, netns.clone(), reuseaddr, reuseport) {
Ok(bound) => bound,
Err((init, err)) => return Err((init, err)),
}
Expand Down Expand Up @@ -298,7 +320,7 @@ impl Init {
continue; // primary inner already covers this iface
}
let new_listen = socket::inet::BoundInner::bind_on_iface(
new_listen_smoltcp_socket(listen_addr)?,
new_listen_smoltcp_socket(listen_addr, reuseport)?,
iface.clone(),
inner.netns(),
)?;
Expand All @@ -308,7 +330,7 @@ impl Init {
let remaining = backlog.saturating_sub(1 + inners.len());
for _ in 0..remaining {
let new_listen = socket::inet::BoundInner::bind_on_iface(
new_listen_smoltcp_socket(listen_addr)?,
new_listen_smoltcp_socket(listen_addr, reuseport)?,
inner.iface().clone(),
inner.netns(),
)?;
Expand All @@ -319,7 +341,7 @@ impl Init {
let additional_sockets = backlog.saturating_sub(1);
for _ in 0..additional_sockets {
let new_listen = socket::inet::BoundInner::bind(
new_listen_smoltcp_socket(listen_addr)?,
new_listen_smoltcp_socket(listen_addr, reuseport)?,
listen_addr
.addr
.as_ref()
Expand All @@ -337,6 +359,7 @@ impl Init {
}

if let Err(err) = inner.with_mut::<smoltcp::socket::tcp::Socket, _, _>(|socket| {
socket.set_reuseport(reuseport);
socket.listen(listen_addr).map_err(|err| match err {
tcp::ListenError::InvalidState => SystemError::EINVAL,
tcp::ListenError::Unaddressable => SystemError::EINVAL,
Expand All @@ -350,14 +373,17 @@ impl Init {
inners,
connect: AtomicUsize::new(0),
listen_addr,
reuseport,
});
}

pub(super) fn close(&self) {
match self {
Init::Unbound(_) => {}
Init::Bound((inner, endpoint)) => {
inner.port_manager().unbind_port(Types::Tcp, endpoint.port);
inner
.port_manager()
.unbind_tcp_port(endpoint.port, inner.iface().nic_id(), inner.handle());
inner.with_mut::<smoltcp::socket::tcp::Socket, _, _>(|socket| socket.close());
}
}
Expand Down Expand Up @@ -434,9 +460,11 @@ impl Connecting {
| ConnectResult::ShutdownReset
| ConnectResult::ShutdownResetConsumed => {
// unbind port
self.inner
.port_manager()
.unbind_port(Types::Tcp, self.local.port);
self.inner.port_manager().unbind_tcp_port(
self.local.port,
self.inner.iface().nic_id(),
self.inner.handle(),
);
let socket = self.inner.into_socket();
let socket = match socket {
smoltcp::socket::Socket::Tcp(s) => s,
Expand Down Expand Up @@ -690,6 +718,7 @@ pub struct Listening {
pub inners: Vec<socket::inet::BoundInner>,
connect: AtomicUsize,
listen_addr: smoltcp::wire::IpListenEndpoint,
reuseport: bool,
}

impl Listening {
Expand All @@ -716,13 +745,13 @@ impl Listening {
// where each interface has its own listen socket in the smoltcp SocketSet.
let mut new_listen = if self.listen_addr.addr.is_none() {
socket::inet::BoundInner::bind_on_iface(
new_listen_smoltcp_socket(self.listen_addr)?,
new_listen_smoltcp_socket(self.listen_addr, self.reuseport)?,
connected.iface().clone(),
connected.netns(),
)?
} else {
socket::inet::BoundInner::bind(
new_listen_smoltcp_socket(self.listen_addr)?,
new_listen_smoltcp_socket(self.listen_addr, self.reuseport)?,
self.listen_addr
.addr
.as_ref()
Expand Down Expand Up @@ -782,12 +811,14 @@ impl Listening {
// (pushed last during listen() construction). We must unbind from its
// port_manager, not inners[0] which may belong to a different iface for
// INADDR_ANY listeners.
self.inners
let owner = self
.inners
.last()
.expect("Listening socket must have at least one inner")
.expect("Listening socket must have at least one inner");
owner
.iface()
.port_manager()
.unbind_port(Types::Tcp, port);
.unbind_tcp_port(port, owner.iface().nic_id(), owner.handle());
}

pub fn release(&mut self) {
Expand Down
Loading
Loading