diff --git a/components/aic8800/src/fdrv/thread/rx.rs b/components/aic8800/src/fdrv/thread/rx.rs index e03d696915..7e3b789811 100644 --- a/components/aic8800/src/fdrv/thread/rx.rs +++ b/components/aic8800/src/fdrv/thread/rx.rs @@ -24,7 +24,7 @@ pub static RX_WAKE_COUNT: AtomicU64 = AtomicU64::new(0); /// AIC8800 是 SDIO WiFi,RX 走自己的线程并独占 SDIO CARD_INT (IRQ#38), /// 不经过 ax_net 的以太网 IRQ 框架。因此数据帧入队后,需主动通知网络栈 /// 来驱动一轮 poll(否则进来的 ARP/ICMP/数据包无人处理)。上层把此回调 -/// 设为 `ax_net::poll_interfaces`,反转依赖,避免本 crate 直接依赖网络栈。 +/// 设为 `ax_net::wake_net_task_irq`,反转依赖,避免本 crate 直接依赖网络栈。 static RX_DATA_CALLBACK: AtomicUsize = AtomicUsize::new(0); /// 本批 RX 是否有数据帧入队(由 `build_and_enqueue_eth_frame` 置位, diff --git a/docs/docs/architecture/net/api.md b/docs/docs/architecture/net/api.md index ba55bdbe47..88ad681be9 100644 --- a/docs/docs/architecture/net/api.md +++ b/docs/docs/architecture/net/api.md @@ -1,5 +1,5 @@ --- -sidebar_position: 7 +sidebar_position: 8 sidebar_label: "对外接口" --- @@ -131,7 +131,6 @@ pub fn init_network(net_devs: EthernetDeviceList, config: NetworkConfig); ```rust pub fn request_poll(); -pub fn poll_interfaces(); ``` `request_poll()` 是 socket、设备和控制路径使用的轻量进度请求入口: @@ -143,8 +142,6 @@ pub fn request_poll() { } ``` -`poll_interfaces()` 保留为 public trigger/debug API,内部只转调 `request_poll()`,不会在调用者线程中同步执行完整 `Service::poll()`。 - ### Vsock 初始化 ```rust @@ -249,14 +246,6 @@ pub struct ArpEntry { `device` 字段是真实接口名。loopback 不产生 ARP entry。 -### 兼容 helper - -```rust -pub fn eth0_ipv4_config() -> Option; -``` - -该函数是旧调用方的 convenience helper。新代码应优先使用 `ipv4_config(name)` 或接口 registry API,避免重新引入固定 `eth0` 假设。 - ## Socket Facade socket API 统一 AF_INET、AF_UNIX 和 AF_VSOCK 的公共操作形状。协议细节由具体 backend 负责。 @@ -355,7 +344,7 @@ pub enum Socket { | `udp::UdpSocket` | `UdpSocket::new()` | AF_INET / SOCK_DGRAM | | `raw::RawSocket` | `RawSocket::new(ip_version, ip_protocol)` | AF_INET / SOCK_RAW | | `unix::UnixSocket` | `UnixSocket::new(Transport)` | AF_UNIX / stream,dgram | -| `vsock::VsockSocket` | `VsockSocket::new(VsockTransport)` | AF_VSOCK / stream | +| `vsock::VsockSocket` | `VsockSocket::new()` | AF_VSOCK / stream | ### 设备绑定 @@ -621,7 +610,7 @@ pub struct NetConfig { } pub fn register_device_with_config(dev: Box, config: NetConfig); -pub fn notify_oob_rx(); +pub fn wake_net_task_irq(); ``` `register_device_with_config()` 用于运行期加入静态 IPv4 Ethernet 设备,例如 Wi-Fi AP 模式设备。它会: @@ -632,7 +621,7 @@ pub fn notify_oob_rx(); - 启动该设备的 RX/TX worker。 - 可选启用内置单客户端 DHCP server。 -`dedicated_poll = true` 时,设备 RX readiness 不走 Ethernet IRQ registrar,而由外部驱动线程调用 `notify_oob_rx()` 唤醒 OOB poll task。 +`dedicated_poll = true` 时,设备 RX readiness 不走 Ethernet IRQ registrar,而由外部驱动线程调用 `wake_net_task_irq()` 唤醒 OOB poll task。 ## Unix Namespace API @@ -677,7 +666,7 @@ TCP/UDP 在 bind port 为 `0` 时分配临时端口。临时端口范围从 `491 ### API 使用建议 -- 新代码使用 `interfaces()`、`interface_by_name()`、`interface_by_id()` 和 `ipv4_config(name)`,不要依赖 `eth0_ipv4_config()`。 +- 使用 `interfaces()`、`interface_by_name()`、`interface_by_id()` 和 `ipv4_config(name)` 查询接口状态。 - socket 发送路径不要直接使用 `default_routes()` 自行选路,应交给 TCP/UDP/raw backend。 - 设备驱动只实现 `EthernetDriver`,不要直接接触 `Router` 或 `SocketSet`。 - 需要唤醒协议栈进度时调用 `request_poll()`,不要从调用者上下文同步 poll smoltcp。 diff --git a/docs/docs/architecture/net/architecture.md b/docs/docs/architecture/net/architecture.md index b448a4a0ec..8dca4ae971 100644 --- a/docs/docs/architecture/net/architecture.md +++ b/docs/docs/architecture/net/architecture.md @@ -40,7 +40,7 @@ flowchart TB Query["interfaces() / default_routes() / arp_entries()"] DnsApi["dns_servers() / dns_query()"] Sockets["TcpSocket / UdpSocket / RawSocket / UnixSocket / VsockSocket"] - PollApi["request_poll() / poll_interfaces()"] + PollApi["request_poll()"] end subgraph Control["Control plane"] @@ -108,6 +108,7 @@ flowchart TB | Single protocol core | 一个 smoltcp `Interface`、全局 `SocketSet`、socket backend、DHCP、orphan 回收、poll 调度 | [service.rs](net/ax-net/src/service.rs), [wrapper.rs](net/ax-net/src/wrapper.rs), [tcp.rs](net/ax-net/src/tcp.rs), [udp.rs](net/ax-net/src/udp.rs), [listen_table.rs](net/ax-net/src/listen_table.rs), [orphan.rs](net/ax-net/src/orphan.rs) | 本文、[Socket 系统](sockets.md) | | Multi-device Router | smoltcp `Device` 适配、TX 路由、RX 汇聚、loopback 快速路径 | [router.rs](net/ax-net/src/router.rs) | [多设备实现](devices.md) | | Device layer | Ethernet 封装/解封装、ARP、IRQ/OOB RX、rd-net 适配 | [device/](net/ax-net/src/device/) | [多设备实现](devices.md) | +| Locking and concurrency | 全局锁顺序、设备/协议核心解耦、原子状态和 waker 协调 | [lib.rs](net/ax-net/src/lib.rs), [service.rs](net/ax-net/src/service.rs), [router.rs](net/ax-net/src/router.rs) | [锁与并发](locks.md) | | Configuration | 静态网络配置、DHCP、MTU、缓冲区、feature | [config.rs](net/ax-net/src/config.rs), [consts.rs](net/ax-net/src/consts.rs), `Cargo.toml` | [配置参考](configuration.md) | | Integration and tests | OS 集成、启动流程、测试范围 | `ax-runtime`, `starry-kernel`, `ax-api` | [集成](integration.md), [测试](testing.md) | @@ -132,7 +133,7 @@ Public API 是上层 OS 模块进入 `ax-net` 的边界,主要定义在 [lib.r | 接口查询 | `interfaces()`、`interface_by_name()`、`ipv4_config()`、`default_routes()`、`arp_entries()` | 从控制面或设备层返回只读快照 | | DNS | `dns_servers()`、`dns_query()`、`dns_query_timeout()` | 读取 DNS registry,并通过临时 smoltcp DNS socket 查询 | | Socket facade | `TcpSocket`、`UdpSocket`、`RawSocket`、`UnixSocket`、`VsockSocket` | 为 syscall/POSIX 层提供统一 socket backend | -| Poll 触发 | `request_poll()`、`poll_interfaces()` | 唤醒专用 net-poll worker,避免应用线程同步驱动协议栈 | +| Poll 触发 | `request_poll()` | 唤醒专用 net-poll worker,避免应用线程同步驱动协议栈 | | Socket options | `GetSocketOption`、`SetSocketOption`、`Configurable` | 覆盖通用 `SO_*`、`TCP_*`、`IP_*` 选项 | Public API 的职责是做边界收敛:上层不需要知道某个 socket 是否由 smoltcp、Unix transport 或 vsock transport 实现,也不需要直接操作 `Service`、`Router` 或 `SocketSet`。具体 API 列表见 [API 参考](api.md)。 diff --git a/docs/docs/architecture/net/configuration.md b/docs/docs/architecture/net/configuration.md index 23914800c3..7e37b0a5fb 100644 --- a/docs/docs/architecture/net/configuration.md +++ b/docs/docs/architecture/net/configuration.md @@ -1,5 +1,5 @@ --- -sidebar_position: 8 +sidebar_position: 9 sidebar_label: "配置参考" --- @@ -235,7 +235,7 @@ pub struct NetConfig { ```rust pub fn register_device_with_config(dev: Box, config: NetConfig); -pub fn notify_oob_rx(); +pub fn wake_net_task_irq(); ``` 注册过程: @@ -247,7 +247,7 @@ pub fn notify_oob_rx(); - `dhcp_server_client_ip` 存在时启用内置单客户端 DHCP server。 - 调用 `request_poll()` 让 net-poll worker 看到新状态。 -`dedicated_poll = true` 时,驱动侧收到 out-of-band RX 事件后调用 `notify_oob_rx()`。 +`dedicated_poll = true` 时,驱动侧收到 out-of-band RX 事件后调用 `wake_net_task_irq()`。 ## 资源预算 @@ -414,5 +414,5 @@ let config = NetworkConfig { - 多网口默认路由通过 metric 控制,主出口使用较小 metric。 - `gateway = 0.0.0.0` 用于只有直连路由的静态接口。 - 需要稳定接口名时优先使用 `ByMac` 或 `ByDriverName`,避免依赖探测顺序。 -- 新代码通过 `ipv4_config(name)` 查询地址,不依赖 `eth0_ipv4_config()`。 +- 通过 `ipv4_config(name)` 查询指定接口地址,避免固定 `eth0` 假设。 - 提高队列常量时应按“每 socket”或“每设备”的乘数估算内存,而不是只看单个 buffer。 diff --git a/docs/docs/architecture/net/devices.md b/docs/docs/architecture/net/devices.md index c8fb07cbcc..2358c1f964 100644 --- a/docs/docs/architecture/net/devices.md +++ b/docs/docs/architecture/net/devices.md @@ -553,7 +553,7 @@ Ethernet 设备在 IP packet 与真实 Ethernet frame 之间转换,并维护 A Ethernet 支持两种 RX readiness 模式: - IRQ 模式:`EthernetIrqRegistrar` 注册硬件 IRQ,IRQ 到来后 `handle_ethernet_irq()` 唤醒 `poll_ready`。 -- OOB RX 模式:用于 SDIO Wi-Fi 等设备,RX 就绪由设备外部线程调用 `notify_oob_rx()`,再唤醒 `{ifname}-oob-poll` 和设备 worker。 +- OOB RX 模式:用于 SDIO Wi-Fi 等设备,RX 就绪由设备外部线程调用 `wake_net_task_irq()`,再唤醒 `{ifname}-oob-poll` 和设备 worker。 `register_waker()` 只在存在 IRQ registration 或 OOB RX wake source 时注册: @@ -622,7 +622,7 @@ fn poll_until_idle() { } while POLL_AGAIN.swap(false, Ordering::AcqRel) { - while poll_once() {} + while get_service().poll(&mut SOCKET_SET.inner.lock()) {} } POLLING_INTERFACES.store(false, Ordering::Release); if !POLL_AGAIN.load(Ordering::Acquire) { diff --git a/docs/docs/architecture/net/flows.md b/docs/docs/architecture/net/flows.md index c4f496fd18..318388fa73 100644 --- a/docs/docs/architecture/net/flows.md +++ b/docs/docs/architecture/net/flows.md @@ -1,5 +1,5 @@ --- -sidebar_position: 9 +sidebar_position: 10 sidebar_label: "运行时流程" --- @@ -146,7 +146,7 @@ fn poll_until_idle() { } while POLL_AGAIN.swap(false, Ordering::AcqRel) { - while poll_once() {} + while get_service().poll(&mut SOCKET_SET.inner.lock()) {} } POLLING_INTERFACES.store(false, Ordering::Release); if !POLL_AGAIN.load(Ordering::Acquire) { @@ -156,7 +156,7 @@ fn poll_until_idle() { } ``` -`poll_once()` 的锁顺序是: +`poll_until_idle()` 内联 poll 调用的锁顺序是: ```text SERVICE -> SOCKET_SET.inner -> Service::poll() @@ -342,7 +342,7 @@ incoming SYN -> snoop_tcp_packet() -> LISTEN_TABLE.incoming_tcp_packet() -> create child smoltcp TCP socket - -> enqueue PendingTcp + -> enqueue AcceptedTcp -> smoltcp consumes SYN and advances child state TcpSocket::accept() @@ -507,7 +507,7 @@ vsock 只在 `vsock` feature 下启用: ```text VsockSocket - -> VsockTransport::Stream + -> VsockStreamTransport -> vsock::connection_manager -> rdif_vsock::Interface event path ``` diff --git a/docs/docs/architecture/net/integration.md b/docs/docs/architecture/net/integration.md index e7262f05ef..6605b44267 100644 --- a/docs/docs/architecture/net/integration.md +++ b/docs/docs/architecture/net/integration.md @@ -1,5 +1,5 @@ --- -sidebar_position: 10 +sidebar_position: 11 sidebar_label: "系统集成" --- @@ -114,10 +114,10 @@ ax_net::unix::register_unix_namespace(crate::unix_ns::AxFsUnixNamespace); ### 动态 Wi-Fi 与 SoftAP -带 `wifi_control()` 的设备会走动态注册路径。runtime 从驱动读取 link policy,并把 OOB RX 唤醒函数设置为 `ax_net::notify_oob_rx`: +带 `wifi_control()` 的设备会走动态注册路径。runtime 从驱动读取 link policy,并把 OOB RX 唤醒函数设置为 `ax_net::wake_net_task_irq`: ```rust -ctrl.set_rx_wake(ax_net::notify_oob_rx); +ctrl.set_rx_wake(ax_net::wake_net_task_irq); let policy = ctrl.link_policy(); ``` diff --git a/docs/docs/architecture/net/locks.md b/docs/docs/architecture/net/locks.md new file mode 100644 index 0000000000..648814a398 --- /dev/null +++ b/docs/docs/architecture/net/locks.md @@ -0,0 +1,878 @@ +--- +sidebar_position: 7 +sidebar_label: "锁与并发" +--- + +# 锁与并发 + +`ax-net` 的并发模型是:**协议核心串行推进,设备收发、控制面查询和用户 socket 调用通过短锁、队列、原子状态与 waker 解耦**。smoltcp 的 `Interface` 和 `SocketSet` 不是多线程并发对象,因此 `ax-net` 保持单协议核心,由专用 `net-poll` worker 独占执行完整 poll;应用线程和设备 worker 不直接成为协议栈驱动者。 + +本文按锁所在层级说明每个同步对象负责的状态、实际源码位置、常见获取路径和不应跨越的边界。代码片段只保留与锁边界相关的部分,完整实现以链接源码为准。 + +## 总体并发模型 + +```mermaid +flowchart TB + App["应用/系统调用线程
send/recv/connect/accept/ioctl"] + PollWorker["net-poll worker
Service::poll()"] + RxWorker["每设备 RX worker
Device::recv()"] + TxWorker["每设备 TX worker
Device::send()"] + Irq["IRQ / OOB RX
wake only"] + + ServiceLock["SERVICE
Mutex"] + SocketLock["SOCKET_SET.inner
Mutex"] + ControlLock["NetControl.state
RwLock"] + RouteLock["SharedRouteTable
RwLock"] + RxQueue["shared RX queue
Mutex + AtomicUsize"] + TxQueue["per-device TX queue
Mutex + AtomicUsize"] + DevLock["DeviceHandle.inner
Mutex"] + DriverLock["driver state
SpinNoIrq"] + Wake["WaitQueue / PollSet / AtomicBool"] + + App -->|"socket state"| SocketLock + App -->|"request_poll()"| Wake + App -->|"interface/DNS query"| ControlLock + App -->|"route query"| RouteLock + + PollWorker --> ServiceLock + ServiceLock --> SocketLock + ServiceLock --> ControlLock + ServiceLock --> RouteLock + ServiceLock --> RxQueue + ServiceLock --> TxQueue + + RxWorker --> DevLock + DevLock --> DriverLock + RxWorker --> RxQueue + RxWorker --> Wake + + TxWorker --> TxQueue + TxWorker --> DevLock + DevLock --> DriverLock + + Irq --> DriverLock + Irq --> Wake +``` + +图中的箭头表示常见访问方向,不表示所有对象都在同一个调用栈中嵌套。关键边界是: + +- 应用线程可以修改 socket 状态并 `request_poll()`,但不直接执行完整 `Service::poll()`。 +- 设备 worker 只在设备和 Router queue 之间搬运 packet,不反向进入 `SERVICE` 或 `SOCKET_SET`。 +- IRQ 路径只做 driver 短操作和 wake,不进入 smoltcp、Router 或 socket set。 +- 控制面查询返回快照,不持锁暴露内部对象引用。 + +## 协议核心锁 + +协议核心由 [lib.rs](net/ax-net/src/lib.rs#L113) 中的全局单例组织: + +```rust +// lib.rs:113-123 +static LISTEN_TABLE: LazyLock = LazyLock::new(ListenTable::new); +static SOCKET_SET: LazyLock = LazyLock::new(SocketSetWrapper::new); + +static SERVICE: Once> = Once::new(); +static NET_CONTROL: Once> = Once::new(); +static POLLING_INTERFACES: AtomicBool = AtomicBool::new(false); +static POLL_AGAIN: AtomicBool = AtomicBool::new(false); +static NET_POLL_REQUESTED: AtomicBool = AtomicBool::new(false); +static NET_POLL_WAKE: WaitQueue = WaitQueue::new(); +``` + +### `SERVICE` + +`SERVICE: Mutex` 是协议核心最外层锁,保护 [Service](net/ax-net/src/service.rs#L310) 内的 smoltcp `Interface`、`Router`、DHCP client/server 和 orphan reaper 状态。完整 poll 只通过 [poll_until_idle()](net/ax-net/src/lib.rs#L429) 进入: + +```rust +// lib.rs:439-440 +while POLL_AGAIN.swap(false, Ordering::AcqRel) { + while get_service().poll(&mut SOCKET_SET.inner.lock()) {} +} +``` + +这行代码定义了主锁顺序:`SERVICE -> SOCKET_SET.inner -> Service::poll()`。因此任何已经持有 `SOCKET_SET.inner` 的路径都不能反向获取 `SERVICE`。 + +`Service::poll()` 的主体在 [service.rs](net/ax-net/src/service.rs#L737)。它在同一轮 poll 内处理 Router RX、DHCP event、DHCP server reply、smoltcp poll、DHCP 定时器、orphan reaper 和 Router TX dispatch: + +```rust +// service.rs:737-785, 摘要 +pub fn poll(&mut self, sockets: &mut SocketSet) -> bool { + router_rx_pending = self.router.poll(timestamp, sockets, |interface_id, packet| { + if let Some(event) = state.process_packet(interface_id, packet, timestamp) { + dhcp_events.push(event); + } + }); + for event in dhcp_events { + self.handle_dhcp_event(event); + } + let socket_state_changed = + self.iface.poll(timestamp, &mut self.router, sockets) == PollResult::SocketStateChanged; + let dhcp_poll_next = self.poll_dhcp(timestamp); + crate::orphan::reap_orphans(timestamp, sockets); + + self.router.dispatch(timestamp, sockets) + || dhcp_poll_next + || socket_state_changed + || router_rx_pending +} +``` + +`SERVICE` 是必要的全局串行点,因为 smoltcp `Interface`、Router 的 smoltcp-facing buffers 和 DHCP 状态必须作为一个协议核心一起推进。它不应包围用户态阻塞 I/O、设备驱动等待或长时间 sleep。 + +### `SOCKET_SET.inner` + +`SOCKET_SET.inner` 定义在 [wrapper.rs](net/ax-net/src/wrapper.rs#L44),保护全局 smoltcp `SocketSet`: + +```rust +// wrapper.rs:44-50 +pub(crate) struct SocketSetWrapper<'a> { + pub inner: Mutex>, + udp_binds: Mutex>, +} +``` + +socket API 经常只需要 `SOCKET_SET.inner`,例如 `with_socket_mut()` 在 [wrapper.rs](net/ax-net/src/wrapper.rs#L68) 只短暂进入某个 smoltcp socket: + +```rust +// wrapper.rs:68-75 +pub fn with_socket_mut, R, F>(&self, handle: SocketHandle, f: F) -> R +where + F: FnOnce(&mut T) -> R, +{ + let mut set = self.inner.lock(); + let socket = set.get_mut(handle); + f(socket) +} +``` + +`SOCKET_SET.inner` 保护的是 smoltcp socket 内部状态,不保护 TCP/UDP public bind side table,也不保护控制面接口 registry。这样做可以避免所有 POSIX 语义都挤进一个全局 socket set 锁。 + +### poll 原子量与 worker 唤醒 + +[poll_until_idle()](net/ax-net/src/lib.rs#L415) 使用 `POLLING_INTERFACES` 防重入,用 `POLL_AGAIN` 合并 poll 过程中新到的请求: + +```rust +// lib.rs:415-429 +fn poll_until_idle() { + POLL_AGAIN.store(true, Ordering::Release); + loop { + if POLLING_INTERFACES + .compare_exchange(false, true, Ordering::Acquire, Ordering::Acquire) + .is_err() + { + return; + } + + while POLL_AGAIN.swap(false, Ordering::AcqRel) { + while get_service().poll(&mut SOCKET_SET.inner.lock()) {} + } + POLLING_INTERFACES.store(false, Ordering::Release); +``` + +这些原子量不是数据结构锁。它们只表达“是否已有线程在 poll”“poll 过程中是否又有新事件”,从而让 socket path 和 device worker 只需轻量 `request_poll()`。 + +## 控制面与路由锁 + +控制面状态定义在 [service.rs](net/ax-net/src/service.rs#L99)。`NetControl.state` 保护接口 registry 和 DNS registry,`routes` 指向共享路由表: + +```rust +// service.rs:99-107 +struct ControlState { + interfaces: Vec, + dns: Vec, +} + +pub struct NetControl { + state: RwLock, + pub(crate) routes: SharedRouteTable, +} +``` + +路由表共享类型定义在 [router.rs](net/ax-net/src/router.rs#L417),Router 和控制面持有同一份 `SharedRouteTable`: + +```rust +// router.rs:417-425 +pub(crate) type SharedRouteTable = Arc>; + +pub struct Router { + rx_buffer: PacketBuffer, + tx_buffer: PacketBuffer, + queues: Arc, + devices: Vec>, + table: SharedRouteTable, +} +``` + +### `NetControl.state` + +`NetControl.state` 是读多写少锁。接口查询、DNS 查询和本地地址绑定推导只持读锁并返回快照,例如 `interfaces()` 在 [service.rs](net/ax-net/src/service.rs#L126): + +```rust +// service.rs:126-130 +pub fn interfaces(&self) -> Vec { + let state = self.state.read(); + state.interfaces.iter().map(NetInterface::to_info).collect() +} +``` + +运行期 DHCP 或静态设备注册会写入接口/DNS 状态。DHCP commit 的关键更新在 [service.rs](net/ax-net/src/service.rs#L254): + +```rust +// service.rs:254-279, 摘要 +let mut state = self.state.write(); +if let Some(interface) = state + .interfaces + .iter_mut() + .find(|interface| interface.id == update.interface_id) +{ + interface.ipv4 = update.ipv4; + interface.gateway = update.gateway; +} +state.dns.retain(|entry| { + entry.interface_id != update.interface_id || entry.source != update.dns_source +}); +self.routes + .write() + .replace_ipv4_rules_for_interface(update.interface_id, routes); +``` + +这里写锁范围只覆盖接口和 DNS registry 的更新;路由表用独立 `SharedRouteTable` 锁。控制面查询路径不进入设备锁,也不需要获取 `SERVICE`。 + +### `SharedRouteTable` + +`SharedRouteTable` 是 route lookup 和 TX dispatch 的共享边界。socket connect/send 通过控制面查询 route;Router dispatch 在 [router.rs](net/ax-net/src/router.rs#L672) 直接读路由表并根据 smoltcp 已选择的源地址决定出接口: + +```rust +// router.rs:672-695, 摘要 +let routes = self.table.read(); +let Some(route) = routes.select_route_for_source(&dst_addr, &src_addr) else { + warn!("No route found for source {} destination {}", src_addr, dst_addr); + continue; +}; + +let dev = &self.devices[route.dev]; +if dev.interface_id == InterfaceId::LOOPBACK { + poll_next |= inject_loopback_rx_direct( + &mut self.rx_buffer, + dst_addr, + packet.into_inner(), + sockets, + ); +} else { + poll_next |= dev.enqueue_tx(route.next_hop, packet.into_inner()); +} +``` + +因此 `SharedRouteTable` 是 TX 热路径锁,但只做规则查找,不访问 driver,不访问 socket payload。接口配置或 DHCP 更新通过写锁替换某接口的 IPv4 路由规则。 + +## Socket 层锁 + +Socket 层锁分为三类:全局 smoltcp socket set、协议 public side table、单 socket 局部状态。它们不是 `SERVICE` 的重复,而是为了让不同语义有不同粒度。 + +### TCP public state、端口表与 listen bucket + +TCP socket 的 public 状态不完全等同于 smoltcp TCP 状态。`TcpSocket` 在 [tcp.rs](net/ax-net/src/tcp.rs#L82) 中用 `StateLock`、endpoint mutex 和原子 option 保存 POSIX 可见状态: + +```rust +// tcp.rs:82-92 +pub struct TcpSocket { + state: StateLock, + handle: SocketHandle, + bound_endpoint: Mutex, + peer_endpoint: Mutex>, + bound_registered: AtomicBool, +``` + +`StateLock` 在 [state.rs](net/ax-net/src/state.rs#L47) 用 `AtomicU8` 做 public state CAS gate: + +```rust +// state.rs:47-65 +pub struct StateLock(AtomicU8); + +impl StateLock { + pub fn get(&self) -> State { + self.0 + .load(Ordering::Acquire) + .try_into() + .expect("invalid state") + } +} +``` + +TCP 端口占用表在 [tcp.rs](net/ax-net/src/tcp.rs#L953): + +```rust +// tcp.rs:953-970 +static TCP_BOUND_PORTS: LazyLock>>>> = + LazyLock::new(|| Mutex::new(HashMap::new())); + +fn register_tcp_bound(endpoint: IpListenEndpoint) -> AxResult { + let mut bound_ports = TCP_BOUND_PORTS.lock(); + let bound_addrs = bound_ports.entry(endpoint.port).or_default(); + if bound_addrs + .iter() + .any(|&addr| listen_addrs_conflict(addr, endpoint.addr)) + { + return Err(AxError::AddrInUse); + } + bound_addrs.insert(endpoint.addr); + Ok(()) +} +``` + +它只记录 public bind ownership,避免每次 ephemeral port 或 bind 冲突检查都扫描整个 `SocketSet`。listen table 按端口懒创建 bucket,每个 bucket 仍是独立短锁,定义在 [listen_table.rs](net/ax-net/src/listen_table.rs#L111): + +```rust +// listen_table.rs:108-112 +type ListenTableEntry = Arc>>; + +pub struct ListenTable { + tcp: Mutex>, +} +``` + +SYN snoop 在 Router RX 阶段进入对应 bucket,并在已经持有 `SOCKET_SET.inner` 的 poll 上下文里创建 child socket,见 [listen_table.rs](net/ax-net/src/listen_table.rs#L267): + +```rust +// listen_table.rs:274-315, 摘要 +let Some(entries) = self.listen_entry(dst.port) else { + return; +}; +let mut table = entries.lock(); +if let Some(entry) = table + .iter_mut() + .find(|entry| entry.can_accept_endpoint(dst)) +{ + if entry.syn_queue.len() >= entry.backlog { + return; + } + let handle = sockets.add(socket); + entry.syn_queue.push_back(AcceptedTcp { + handle, + local_endpoint: dst, + remote_endpoint: src, + }); + entry.accept_poll.wake(); +} +``` + +对应的 accept 路径在 [tcp.rs](net/ax-net/src/tcp.rs#L517),顺序是先进入 `SOCKET_SET.inner`,再进入 `LISTEN_TABLE` bucket: + +```rust +// tcp.rs:522-528 +let bound_endpoint = self.bound_endpoint()?; +self.general.recv_poller(self, || { + request_poll(); + let accepted = { + let mut sockets = SOCKET_SET.inner.lock(); + LISTEN_TABLE.accept(bound_endpoint, &mut sockets)? + }; +``` + +### UDP bind side table 与局部状态 + +UDP 的 local/peer 状态定义在 [udp.rs](net/ax-net/src/udp.rs#L76): + +```rust +// udp.rs:76-88 +pub struct UdpSocket { + handle: SocketHandle, + local_addr: RwLock>, + peer_addr: RwLock>, + general: GeneralOptions, + cork: Mutex>, +} +``` + +UDP public bind side table 放在 `SocketSetWrapper.udp_binds`,bind 路径在 [udp.rs](net/ax-net/src/udp.rs#L181) 体现了实际顺序:先写本地地址状态,进入 smoltcp bind,再登记 public bind ownership。 + +```rust +// udp.rs:183-220, 摘要 +fn bind(&self, local_addr: SocketAddrEx) -> AxResult { + let mut guard = self.local_addr.write(); + let binding = get_control().local_binding_for(&endpoint)?; + + self.with_smol_socket(|socket| { + socket.bind(endpoint).map_err(|e| /* ... */) + })?; + if !self.general.reuse_address() + && let Err(err) = + SOCKET_SET.udp_bind(self.handle, local_endpoint.addr, local_endpoint.port) + { + self.with_smol_socket(|socket| socket.close()); + return Err(err); + } + *guard = Some(local_endpoint); + Ok(()) +} +``` + +这里 `local_addr` 和 `udp_binds` 是不同层级:前者是单 socket public state,后者是全局 UDP 端口占用 side table。它们不能简单合并到 `SOCKET_SET.inner`,否则 bind 冲突检查和 socket payload 访问会共享同一个重锁。 + +### raw socket 暂存锁 + +raw socket 使用读写锁保存 filter/TTL,并用 `SpinNoIrq` 保护本地暂存包。文件顶部在 [raw.rs](net/ax-net/src/raw.rs#L35) 把 `SpinNoIrq` 别名为 `Mutex`,字段定义在 [raw.rs](net/ax-net/src/raw.rs#L70): + +```rust +// raw.rs:35,70-84 +use ax_kspin::SpinNoIrq as Mutex; + +pub struct RawSocket { + handle: SocketHandle, + ip_version: IpVersion, + local_addr: RwLock>, + peer_addr: RwLock>, + loopback_rx: Mutex)>>, + deferred_rx: Mutex)>>, + ttl: RwLock>, +``` + +`deferred_rx` 的写入在 [raw.rs](net/ax-net/src/raw.rs#L488),只保存一个被 peer filter 跳过的 wire packet,不跨越阻塞等待: + +```rust +// raw.rs:488-490 +if !self.source_matches_peer(source) { + *self.deferred_rx.lock() = Some((source, wire_packet.to_vec())); + return Err(AxError::WouldBlock); +} +``` + +### 通用 socket option + +`GeneralOptions` 在 [general.rs](net/ax-net/src/general.rs#L34) 用原子字段保存 nonblocking、reuseaddr、timeout、`SO_BINDTODEVICE` 和 socket identity: + +```rust +// general.rs:34-49 +pub(crate) struct GeneralOptions { + nonblock: AtomicBool, + reuse_address: AtomicBool, + send_timeout_nanos: AtomicU64, + recv_timeout_nanos: AtomicU64, + bound_if: AtomicU32, + socket_type: AtomicI32, +``` + +这些字段是单值状态,用原子量可以避免 option get/set 每次进入全局 socket set。它们不保护 smoltcp socket state,也不保护复合 bind 语义。 + +## Router 与设备队列锁 + +Router 层是协议核心和设备 worker 之间的内存边界。它用有界队列解耦设备收发,不让设备 worker 直接持有 `SERVICE` 或 `SOCKET_SET`。 + +### `BoundedPacketQueue` + +队列定义在 [router.rs](net/ax-net/src/router.rs#L128)。`inner` 保护 `VecDeque`,`len` 是 wait predicate 和快速空队列检查用的长度快照: + +```rust +// router.rs:128-163 +struct BoundedPacketQueue { + inner: Mutex>, + capacity: usize, + len: AtomicUsize, +} + +fn push(&self, packet: T) -> Result<(), T> { + let mut inner = self.inner.lock(); + if inner.len() >= self.capacity { + return Err(packet); + } + inner.push_back(packet); + self.len.store(inner.len(), Ordering::Release); + Ok(()) +} + +fn pop(&self) -> Option { + let mut inner = self.inner.lock(); + let packet = inner.pop_front(); + self.len.store(inner.len(), Ordering::Release); + packet +} +``` + +`len` 不保护队列内容,因此任何真正 push/pop 都必须进入 `inner`。它的作用是减少 worker 等待判断时的无谓加锁。 + +### `DeviceHandle` + +每个设备一个 `DeviceHandle`,定义在 [router.rs](net/ax-net/src/router.rs#L212): + +```rust +// router.rs:212-228 +struct DeviceHandle { + interface_id: InterfaceId, + name: String, + inner: Arc>>, + rx_queue: Arc>, + tx_queue: Arc>, + rx_wake: Arc, + tx_wake: Arc, + rx_waker: Waker, +} +``` + +`DeviceHandle.inner` 只保护具体设备对象,例如 loopback、Ethernet 或 OOB 设备。它不保护 smoltcp `Interface`、`SocketSet`、路由表或控制面。 + +### RX worker + +RX worker 在 [router.rs](net/ax-net/src/router.rs#L805) 先短暂持有 `DeviceHandle.inner` 调用 `Device::recv()`,再把 packet 推入共享 RX queue 并 `request_poll()`: + +```rust +// router.rs:810-847, 摘要 +let mut rx_buffer = PacketBuffer::new( + vec![PacketMetadata::EMPTY; DEVICE_RX_WORKER_BATCH], + vec![0u8; STANDARD_MTU * DEVICE_RX_WORKER_BATCH], +); + +{ + let mut device_inner = device.inner.lock(); + while !rx_buffer.is_full() + && device_inner.recv(device.interface_id, &mut rx_buffer, now(), &mut snoop) + { + received = true; + } +} + +while let Ok((interface_id, packet)) = rx_buffer.dequeue() { + if device.rx_queue.push(rx).is_err() { + warn!("{}: RX queue is full, dropping packet", device.name); + crate::request_poll(); + break; + } + crate::request_poll(); +} + +if !received { + device.inner.lock().register_waker(&device.rx_waker); + device.rx_wake.wait(); +} +``` + +这里的关键点是:设备锁释放后才向 Router queue 搬运 packet;整个路径不进入 `SERVICE` 或 `SOCKET_SET`。 + +### TX worker + +TX worker 在 [router.rs](net/ax-net/src/router.rs#L787) 从 per-device TX queue 弹出 packet,然后持设备锁调用 `Device::send()`: + +```rust +// router.rs:787-800 +fn device_tx_worker(device: Arc) { + loop { + if let Some(packet) = device.tx_queue.pop() { + let poll_next = + device + .inner + .lock() + .send(packet.next_hop, packet.bytes.as_slice(), now()); + if poll_next { + crate::request_poll(); + } + } else { + device.tx_wake.wait_until(|| !device.tx_queue.is_empty()); + } + } +} +``` + +这保证慢设备发送不会持有协议核心锁。Router dispatch 只负责路由选择和入队,真实发送由 TX worker 完成。 + +### Router poll 与 dispatch + +`Router::poll()` 在 [router.rs](net/ax-net/src/router.rs#L577) 把 worker RX queue 搬到 smoltcp-facing `rx_buffer`: + +```rust +// router.rs:586-600 +while !self.rx_buffer.is_full() { + let Some(packet) = self.queues.rx.pop() else { + break; + }; + let bytes = packet.bytes.as_slice(); + snoop_tcp_packet(bytes, sockets); + snoop(packet.interface_id, bytes); + let Ok(dst) = self.rx_buffer.enqueue(bytes.len(), packet.interface_id) else { + break; + }; + dst.copy_from_slice(bytes); + moved_rx = true; +} +``` + +`Router::dispatch()` 在 smoltcp poll 之后处理 `tx_buffer`。普通设备走 per-device TX queue;loopback 直接注入 `rx_buffer`,避免设备队列 hop。相关逻辑见 [router.rs](net/ax-net/src/router.rs#L653)。 + +## 设备驱动短锁 + +设备驱动层锁比 Router 队列更底层,通常需要覆盖 IRQ 与任务上下文同时访问 driver state 的场景,因此使用短临界区。 + +### Ethernet IRQ state + +Ethernet IRQ 共享状态定义在 [device/ethernet.rs](net/ax-net/src/device/ethernet.rs#L123): + +```rust +// device/ethernet.rs:123-138 +struct EthernetIrqState { + irq: Option, + irq_registration: spin::Once>, + oob_rx: bool, + driver: SpinNoIrq>, + poll_ready: PollSet, +} + +impl EthernetIrqState { + fn handle_irq(&self) -> NetIrqEvents { + self.driver.lock().handle_irq() + } +} +``` + +`SpinNoIrq` 下只允许 driver 短操作,例如 IRQ event 读取、TX/RX queue 操作、poll-ready 注册。不得在 guard 内 sleep、wait、调用 socket API 或进入 `Service::poll()`。 + +### rd-net adapter state + +`rd-net` 适配器在 [device/driver.rs](net/ax-net/src/device/driver.rs#L179) 用 `SpinNoIrq` 保护底层 TX/RX queue 和 `pending_rx`: + +```rust +// device/driver.rs:179-204 +pub struct RdNetDriver { + name: String, + mac: [u8; 6], + irq: Option, + irq_handler: Option, + state: SpinNoIrq, +} + +state: SpinNoIrq::new(RdNetState { + tx_queue, + rx_queue, + pending_rx: VecDeque::with_capacity(RX_PREFETCH_TARGET), +}), +``` + +这个锁只保护 `rd-net` ownership 和 queue state,不保护 Router queue,也不保护 smoltcp 状态。 + +## Unix 与 vsock 本地传输锁 + +Unix socket 和 vsock 不经过 smoltcp `SocketSet`,但仍复用 socket facade、`GeneralOptions` 和 readiness/poll 机制。它们有自己的局部锁。 + +Unix abstract namespace 的 bind slot 在 [unix/mod.rs](net/ax-net/src/unix/mod.rs#L112): + +```rust +// unix/mod.rs:112-121 +pub struct BindSlot { + stream: Mutex>, + dgram: Mutex>, +} + +static ABSTRACT_BINDS: LazyLock, BindSlot>>> = + LazyLock::new(|| Mutex::new(HashMap::new())); +``` + +vsock 设备和 pending events 在 [device/vsock.rs](net/ax-net/src/device/vsock.rs#L36),连接管理器在 [vsock/connection_manager.rs](net/ax-net/src/vsock/connection_manager.rs#L568): + +```rust +// device/vsock.rs:36-53 +static VSOCK_DEVICE: Mutex> = Mutex::new(None); +static PENDING_EVENTS: Mutex> = Mutex::new(VecDeque::new()); +static POLL_REF_COUNT: Mutex = Mutex::new(0); +static POLL_TASK_RUNNING: AtomicBool = AtomicBool::new(false); + +// vsock/connection_manager.rs:568-569 +pub static VSOCK_CONN_MANAGER: Mutex = + Mutex::new(VsockConnectionManager::new()); +``` + +这部分锁不进入 `SERVICE`,也不要求 smoltcp poll。它们的锁顺序只需要在 Unix/vsock 局部模块内保持一致。 + +## 锁顺序与关键路径 + +### 全局锁顺序 + +[service.rs](net/ax-net/src/service.rs#L3) 文件头给出了协议核心的全局顺序: + +```text +SERVICE + -> SOCKET_SET.inner + -> TCP_BOUND_PORTS + -> LISTEN_TABLE.tcp[port] +``` + +这是允许嵌套时必须遵守的顺序,不表示每条路径都会同时持有全部锁。实际代码会尽量拆短临界区,例如 TCP bind 会分开检查 listen table、控制面绑定和 `TCP_BOUND_PORTS` 登记。 + +禁止模式: + +```text +SOCKET_SET.inner -> SERVICE // 反向获取协议核心锁 +DeviceHandle.inner -> SERVICE // 设备 worker 反向进入协议核心 +SpinNoIrq guard -> block_on/wait // 禁 IRQ/抢占状态下阻塞 +SERVICE/SOCKET_SET -> long sleep // 阻塞协议栈推进 +``` + +### net-poll worker + +```text +NET_POLL_WAKE wait + -> POLLING_INTERFACES CAS + -> SERVICE.lock() + -> SOCKET_SET.inner.lock() + -> Service::poll() + -> Router::poll(): shared RX queue lock + -> DHCP events may commit NetControl/RouteTable + -> smoltcp Interface::poll() + -> orphan reaper: ORPHAN_SOCKETS.lock() + -> Router::dispatch(): RouteTable.read + per-device TX queue lock +``` + +`poll_until_idle()` 是唯一执行完整 smoltcp poll 的路径。socket 调用者只 `request_poll()`,不会在热路径同步抢 `SERVICE` 推进整个协议栈。 + +### TCP bind / listen / accept + +```text +bind(): + StateLock CAS + -> bound_endpoint Mutex + -> LISTEN_TABLE.can_listen() + -> NetControl.local_binding_for() + -> TCP_BOUND_PORTS.lock() + +listen(): + StateLock CAS + -> bound_endpoint Mutex + -> NetControl.local_binding_for() + -> TCP_BOUND_PORTS.lock() // only if bind() did not register it earlier + -> LISTEN_TABLE.tcp[port].lock() + +accept(): + SOCKET_SET.inner.lock() + -> LISTEN_TABLE.tcp[port].lock() + +SYN snoop during Router RX: + SOCKET_SET.inner is already held by poll_until_idle()'s inline poll call + -> LISTEN_TABLE.tcp[port].lock() + -> SocketSet::add(child socket) +``` + +TCP 的 public bind 语义由 `TCP_BOUND_PORTS` 维护,passive open 的 pending child 由 `LISTEN_TABLE` 维护。两者分离是为了避免每次 bind 都扫描整个 `SocketSet`。 + +### UDP / raw socket + +```text +UDP bind: + local_addr.write() + -> smoltcp socket bind through SOCKET_SET.inner + -> SocketSetWrapper.udp_binds.lock() + +UDP send: + peer_addr.read() + -> SOCKET_SET.inner.lock() + -> request_poll() + +raw recv with peer filter: + SOCKET_SET.inner.lock() + -> deferred_rx SpinNoIrq only for local stash +``` + +UDP/raw 的本地地址、peer 地址、TTL 使用读写锁或短 `SpinNoIrq`,因为这些状态和 smoltcp socket payload 的生命周期不同。 + +### 控制面查询与提交 + +```text +interfaces()/dns_servers()/ipv4_config(): + NetControl.state.read() + -> clone snapshot + -> unlock + +select_route_with_binding(): + NetControl.state.read() + -> SharedRouteTable.read() + -> RouteDecision + +DHCP/static commit: + SERVICE.lock() + -> update smoltcp Interface address list + -> NetControl.state.write() and/or SharedRouteTable.write() +``` + +查询路径通常不获取 `SERVICE`,因此接口查询、DNS server 查询和 route snapshot 不会被 smoltcp poll 长时间阻塞。提交路径由 `Service` 协调,是为了让 smoltcp address list 与控制面快照保持一致。 + +### 设备 RX/TX worker + +```text +RX worker: + DeviceHandle.inner.lock() + -> Device::recv() + -> EthernetIrqState.driver SpinNoIrq + -> RdNetDriver.state SpinNoIrq + -> shared RX queue lock + -> request_poll() + +TX worker: + per-device TX queue lock + -> DeviceHandle.inner.lock() + -> Device::send() + -> EthernetIrqState.driver SpinNoIrq + -> RdNetDriver.state SpinNoIrq +``` + +设备 worker 不持有 `SERVICE` 或 `SOCKET_SET`。它们只在设备和 Router queue 之间搬运 packet。队列满时丢包并唤醒 net-poll worker,不创建无界 backlog。 + +### IRQ / OOB RX + +```text +Ethernet IRQ: + EthernetIrqState.driver SpinNoIrq + -> poll_ready.wake() + -> return Wake + +OOB RX: + wake_net_task_irq() + -> OOB_RX_SIGNAL.wake() + -> {ifname}-oob-poll task + -> request_poll() +``` + +IRQ handler 不进入 `SERVICE`、`SOCKET_SET` 或 Router。它只读取 driver IRQ event 并唤醒任务上下文,由 worker 或 net-poll 在可阻塞上下文中继续处理。 + +## 设计约束 + +### 为什么不能只保留全局锁 + +`SERVICE` 和 `SOCKET_SET.inner` 只能保护 smoltcp 协议核心与 socket set。内部仍需要局部锁,原因是: + +- 设备 worker 必须能在不进入协议核心的情况下收发 packet,否则慢设备会阻塞 smoltcp poll。 +- 控制面查询需要在不持 `SERVICE` 的情况下返回接口、DNS 和路由快照,否则 `ifconfig`、netlink、DNS server 查询会和协议 poll 强耦合。 +- TCP/UDP bind side table 是 POSIX public 语义,不等同于 smoltcp socket payload state,单独维护可以避免扫描整个 `SocketSet`。 +- Unix/vsock 不经过 smoltcp,不能依赖 `SERVICE` 表达本地传输状态。 +- `SpinNoIrq` 只服务 IRQ/driver 短临界区,不能和任务级 `Mutex` 混用成一个大锁。 + +因此当前锁不是冗余叠加,而是按所有权边界拆分:协议核心、控制面、socket public state、Router queue、设备驱动、本地传输各自保护自己的状态。 + +### 为什么 socket 热路径不直接 poll + +如果 socket send/recv/connect 在持有 socket 或 `SocketSet` 状态时同步调用 `Service::poll()`,容易形成: + +- 应用线程与协议核心互相阻塞。 +- 多线程抢全局 poll 锁。 +- 设备 worker / net-poll worker 被应用线程时序影响。 +- 单核上出现 busy-loop 或不稳定调度依赖。 + +当前模型是: + +```text +socket path: + mutate socket state + -> request_poll() + -> optional wait on poller/waker + +net-poll worker: + wake + -> poll_until_idle() +``` + +这更接近 lwIP `tcpip_thread` 和 Linux softirq/NAPI 的职责分离:应用线程不成为临时协议栈驱动者。 + +## 检查清单 + +修改 `ax-net` 锁相关代码时,应确认: + +- 新路径没有 `SOCKET_SET.inner -> SERVICE` 的反向获取。 +- 设备 worker 没有在持有 `DeviceHandle.inner` 或 driver `SpinNoIrq` 时进入 `SERVICE` / `SOCKET_SET`。 +- `SpinNoIrq` guard 内没有 sleep、wait、block_on、DNS 查询或 socket API。 +- 新增全局表优先使用短临界区,并说明它与 `SERVICE` / `SOCKET_SET` 的顺序。 +- 新增 socket 局部状态优先用原子或 `RwLock`,避免把整个 POSIX 操作包在全局 `SocketSet` 锁内。 +- 新增 worker wake 路径只设置原子/waker,不直接执行 smoltcp poll。 +- 文档中的锁顺序与 [service.rs](net/ax-net/src/service.rs#L3) 文件头保持一致。 diff --git a/docs/docs/architecture/net/overview.md b/docs/docs/architecture/net/overview.md index 1e2d4812f0..1ee96341c6 100644 --- a/docs/docs/architecture/net/overview.md +++ b/docs/docs/architecture/net/overview.md @@ -25,7 +25,7 @@ TGOSKits 的网络能力收敛在 `net/ax-net`。它是 ArceOS、StarryOS 和 Ax | [orphan.rs](net/ax-net/src/orphan.rs) | TCP orphan socket 回收(RFC 793 TIME_WAIT) | `add_orphan`, `reap_orphans` | | [dhcp_server.rs](net/ax-net/src/dhcp_server.rs) | 最简 DHCP 服务器(SoftAP 模式) | `DhcpServer` | | [unix/](net/ax-net/src/unix/) | Unix domain socket | `UnixSocket`, `Transport` | -| [vsock/](net/ax-net/src/vsock/) | 可选 vsock 支持(`vsock` feature) | `VsockSocket`, `VsockTransport` | +| [vsock/](net/ax-net/src/vsock/) | 可选 vsock 支持(`vsock` feature) | `VsockSocket`, `VsockStreamTransport` | | [device/](net/ax-net/src/device/) | loopback、Ethernet、rd-net、vsock 设备适配 | `Device`, `EthernetDevice`, `RdNetDriver` | | [consts.rs](net/ax-net/src/consts.rs) | 缓冲区大小等常量 | `STANDARD_MTU`, `SOCKET_BUFFER_SIZE` | @@ -47,7 +47,7 @@ TGOSKits 的网络能力收敛在 `net/ax-net`。它是 ArceOS、StarryOS 和 Ax | Loopback | 零状态 `LoopbackDevice` + `Router::dispatch()` 快速路径 inline 注入 `rx_buffer`,不经设备 worker 和队列分配 | 完整 | | TCP orphan 回收 | `orphan.rs`:Drop 后保留 smoltcp socket 直到 FIN/TIME_WAIT 完成,RFC 793 合规 | 完整 | | DHCP 服务器(SoftAP) | `dhcp_server.rs`:最简单客户端 DHCP 服务器,Discover→Offer、Request→Ack | 完整 | -| OOB RX(SDIO Wi-Fi) | `EthernetDevice::new_oob_rx()` + `notify_oob_rx()` + 独立 poll task | 完整 | +| OOB RX(SDIO Wi-Fi) | `EthernetDevice::new_oob_rx()` + `wake_net_task_irq()` + 独立 poll task | 完整 | | 动态设备注册 | `register_device_with_config()` 运行时添加静态 IP 设备(Wi-Fi AP) | 完整 | ## 设计原则 @@ -72,6 +72,7 @@ TGOSKits 的网络能力收敛在 `net/ax-net`。它是 ArceOS、StarryOS 和 Ax | `{ifname}-oob-poll` | OOB RX 设备(如 SDIO Wi-Fi)的专用 poll task | `OOB_RX_SIGNAL.wait()` | `NET_POLL_DEVICE_WAKER` 是全局设备 readiness waker。Router 会把它注册给所有允许触发全局协议栈推进的设备;设备 RX/IRQ/OOB 路径只唤醒 worker 和设置 poll 请求,不直接进入 smoltcp `Interface::poll()`。 +完整锁类型、锁顺序和禁止模式见[锁与并发](locks.md)。 ### 全局锁顺序 @@ -90,6 +91,7 @@ SERVICE (Mutex) - `NET_CONTROL.state` 是独立 RwLock,接口查询(只读)可以在不持有 `SERVICE` 的情况下进行。 - `ListenTable` 条目锁在 `SOCKET_SET` 锁内获取,保证 accept/snoop 的一致性。 - 设备锁(`DeviceHandle.inner`)主要由 `{ifname}-rx` / `{ifname}-tx` worker 独立获取。worker 不应在持有设备锁时反向进入 `SERVICE` 或 `SOCKET_SET`,避免设备路径与协议核心互相阻塞。 +更细的控制面、Router、socket、Unix 和 vsock 锁划分见[锁与并发](locks.md)。 ## 核心方案 diff --git a/docs/docs/architecture/net/sockets.md b/docs/docs/architecture/net/sockets.md index 6dab54f30f..a5e8cd5557 100644 --- a/docs/docs/architecture/net/sockets.md +++ b/docs/docs/architecture/net/sockets.md @@ -303,12 +303,25 @@ fn udp_bind_available(binds: &HashMap, key: UdpBindKey TCP 除了 listen table,还需要记录“已经 bind 但还没有 listen/connect 完成”的端口所有权: ```rust -static TCP_BOUND_PORTS: LazyLock>>>> = +static TCP_BOUND_PORTS: LazyLock>>>> = LazyLock::new(|| Mutex::new(HashMap::new())); fn listen_addrs_conflict(a: Option, b: Option) -> bool { a.is_none() || b.is_none() || a == b } + +fn register_tcp_bound(endpoint: IpListenEndpoint) -> AxResult { + let mut bound_ports = TCP_BOUND_PORTS.lock(); + let bound_addrs = bound_ports.entry(endpoint.port).or_default(); + if bound_addrs + .iter() + .any(|&addr| listen_addrs_conflict(addr, endpoint.addr)) + { + return Err(AxError::AddrInUse); + } + bound_addrs.insert(endpoint.addr); + Ok(()) +} ``` 语义是 wildcard 与所有地址冲突,两个具体地址仅在相等时冲突。ephemeral TCP 端口分配同时检查 listen table 和 bound table: @@ -330,21 +343,22 @@ fn tcp_port_available(port: u16) -> bool { struct ListenTableEntryInner { listen_endpoint: IpListenEndpoint, backlog: usize, - syn_queue: VecDeque, + syn_queue: VecDeque, accept_poll: Arc, } pub struct ListenTable { - tcp: Box<[ListenTableEntry]>, + tcp: Mutex>, } ``` -`tcp` 是 65536 个端口 bucket,每个 bucket 存放该端口下的多个具体地址 listener。`listen()` 检查 wildcard/specific 冲突后插入 entry: +`tcp` 按端口懒创建 listen bucket,每个 bucket 存放该端口下的多个具体地址 listener。`listen()` 检查 wildcard/specific 冲突后插入 entry: ```rust pub fn listen(&self, listen_endpoint: IpListenEndpoint, backlog: usize) -> AxResult { let port = listen_endpoint.port; - let mut entries = self.tcp[port as usize].lock(); + let entries = self.listen_entry_or_create(port); + let mut entries = entries.lock(); if entries .iter() .any(|entry| listen_addrs_conflict(entry.listen_endpoint.addr, listen_endpoint.addr)) @@ -368,7 +382,9 @@ pub fn accept( listen_endpoint: IpListenEndpoint, sockets: &mut SocketSet<'_>, ) -> AxResult { - let entries = self.listen_entry(listen_endpoint.port); + let Some(entries) = self.listen_entry(listen_endpoint.port) else { + return Err(AxError::InvalidInput); + }; let mut table = entries.lock(); let Some(entry) = table .iter_mut() @@ -377,17 +393,17 @@ pub fn accept( return Err(AxError::InvalidInput); }; - let syn_queue: &mut VecDeque = &mut entry.syn_queue; + let syn_queue: &mut VecDeque = &mut entry.syn_queue; let mut idx = 0; while idx < syn_queue.len() { - let handle = syn_queue[idx].accepted.handle; + let handle = syn_queue[idx].handle; if is_closed_without_data(sockets, handle) { syn_queue.swap_remove_front(idx); sockets.remove(handle); continue; } if is_acceptable(sockets, handle) { - return Ok(syn_queue.swap_remove_front(idx).unwrap().accepted); + return Ok(syn_queue.swap_remove_front(idx).unwrap()); } idx += 1; } @@ -598,19 +614,15 @@ Unix transport 支持 Linux 风格的 credentials 查询: ### Vsock Socket -vsock 是可选 feature,不属于 IP 协议,也不通过 smoltcp poll。facade 只把 `SocketOps` 映射到 `VsockTransport`: +vsock 是可选 feature,不属于 IP 协议,也不通过 smoltcp poll。facade 只把 `SocketOps` 映射到 stream transport: ```rust -pub enum VsockTransport { - Stream(VsockStreamTransport), -} - pub struct VsockSocket { - transport: VsockTransport, + transport: VsockStreamTransport, } ``` -核心连接状态位于 `vsock::connection_manager`,设备事件由 vsock 设备层推进。transport enum 提供 stream variant,并为后续 datagram 扩展保留结构。 +核心连接状态位于 `vsock::connection_manager`,设备事件由 vsock 设备层推进。 #### Vsock Connection Manager @@ -646,7 +658,6 @@ vsock 设备层还有一个临时 RX buffer 和 pending event queue: - `VSOCK_RX_TMPBUF_SIZE = 4 KiB`:poll task 从 `rdif_vsock::Interface` 拉取事件时使用的临时接收缓冲。 - `PENDING_EVENTS`:当事件暂时无法完整交付给 manager(例如目标连接 RX ring 空间不足)时保存事件,后续 poll 周期继续处理,避免直接丢弃设备事件。 -- `VsockStats` / `get_vsock_stats()`:导出 connection manager 的连接数、监听数和队列状态,用于诊断 vsock 连接泄漏或 accept backlog 问题。 #### Vsock Poll Worker diff --git a/docs/docs/architecture/net/testing.md b/docs/docs/architecture/net/testing.md index 4db42f05d2..1054656692 100644 --- a/docs/docs/architecture/net/testing.md +++ b/docs/docs/architecture/net/testing.md @@ -1,5 +1,5 @@ --- -sidebar_position: 11 +sidebar_position: 12 sidebar_label: "测试与限制" --- diff --git a/docs/docs/components/crates/ax-net.md b/docs/docs/components/crates/ax-net.md index 7475040d4c..897a26dd70 100644 --- a/docs/docs/components/crates/ax-net.md +++ b/docs/docs/components/crates/ax-net.md @@ -28,7 +28,7 @@ - `init_network(net_devs)`:由 `ax-runtime` 调用,接入 Ethernet 设备、loopback、router、smoltcp service,并完成静态网络或 DHCP 初始化。 - `init_vsock(vsock_devs)`:在 `vsock` feature 下启用,由 `ax-runtime` 调用。 -- `poll_interfaces()`:推动收包、发包、协议状态机和 readiness 更新。 +- `request_poll()`:唤醒 net-poll worker 推动收包、发包、协议状态机和 readiness 更新。 socket 入口: diff --git a/drivers/net/rd-net/src/lib.rs b/drivers/net/rd-net/src/lib.rs index 740b460bc6..26a083e659 100644 --- a/drivers/net/rd-net/src/lib.rs +++ b/drivers/net/rd-net/src/lib.rs @@ -200,6 +200,14 @@ pub struct WifiControlHandle { inner: Arc, } +impl Clone for WifiControlHandle { + fn clone(&self) -> Self { + Self { + inner: self.inner.clone(), + } + } +} + unsafe impl Send for WifiControlHandle {} unsafe impl Sync for WifiControlHandle {} diff --git a/net/ax-net/src/addr.rs b/net/ax-net/src/addr.rs new file mode 100644 index 0000000000..b067be2a1a --- /dev/null +++ b/net/ax-net/src/addr.rs @@ -0,0 +1,44 @@ +//! Shared address and ephemeral-port helpers. + +use ax_errno::{AxResult, ax_bail}; +use ax_sync::Mutex; +use smoltcp::wire::{IpAddress, Ipv4Address}; + +const EPHEMERAL_PORT_START: u16 = 0xc000; +const EPHEMERAL_PORT_END: u16 = 0xffff; + +/// Returns whether two wildcard/specific local addresses conflict on one port. +pub(crate) fn listen_addrs_conflict(a: Option, b: Option) -> bool { + a.is_none() || b.is_none() || a == b +} + +/// Allocates an ephemeral port accepted by `check_available`. +pub(crate) fn allocate_ephemeral_port(check_available: impl Fn(u16) -> bool) -> AxResult { + static CURR: Mutex = Mutex::new(EPHEMERAL_PORT_START); + + let mut curr = CURR.lock(); + let mut tries = 0; + while tries <= EPHEMERAL_PORT_END - EPHEMERAL_PORT_START { + let port = *curr; + if *curr == EPHEMERAL_PORT_END { + *curr = EPHEMERAL_PORT_START; + } else { + *curr += 1; + } + if check_available(port) { + return Ok(port); + } + tries += 1; + } + ax_bail!(AddrInUse, "no available ports"); +} + +/// Builds an IPv4 netmask from a CIDR prefix length. +pub(crate) fn mask_from_prefix(prefix_len: u8) -> Ipv4Address { + let bits: u32 = if prefix_len == 0 { + 0 + } else { + u32::MAX << (32 - prefix_len.min(32) as u32) + }; + Ipv4Address::from_bits(bits) +} diff --git a/net/ax-net/src/device/ethernet.rs b/net/ax-net/src/device/ethernet.rs index 51b0623b8c..5467967f27 100644 --- a/net/ax-net/src/device/ethernet.rs +++ b/net/ax-net/src/device/ethernet.rs @@ -126,7 +126,7 @@ struct EthernetIrqState { /// RX readiness is delivered out-of-band (outside the ethernet IRQ /// framework) via the device readiness poll set, e.g. an SDIO Wi-Fi chip /// that owns its own card interrupt and pokes the stack through - /// `notify_oob_rx`. + /// `wake_net_task_irq`. oob_rx: bool, driver: SpinNoIrq>, poll_ready: Arc, diff --git a/net/ax-net/src/dhcp_server.rs b/net/ax-net/src/dhcp_server.rs index 9d87ae245d..c790228e3e 100644 --- a/net/ax-net/src/dhcp_server.rs +++ b/net/ax-net/src/dhcp_server.rs @@ -29,6 +29,55 @@ use crate::config::InterfaceId; /// Lease duration advertised in Offer/Ack replies, in seconds. const LEASE_SECS: u32 = 86400; +/// Parsed DHCP-over-IPv4/UDP packet. +pub(crate) struct ParsedDhcp { + pub(crate) src_addr: Ipv4Address, + pub(crate) udp: UdpRepr, + pub(crate) message_type: DhcpMessageType, + pub(crate) transaction_id: u32, + pub(crate) client_hardware_address: EthernetAddress, + pub(crate) your_ip: Ipv4Address, + pub(crate) server_identifier: Option, + pub(crate) subnet_mask: Option, + pub(crate) router: Option, + pub(crate) dns_servers: Vec, +} + +/// Parses an IPv4 UDP DHCP packet, leaving direction checks to the caller. +pub(crate) fn parse_dhcp_packet(packet: &[u8]) -> Option { + let ipv4_packet = Ipv4Packet::new_checked(packet).ok()?; + let ipv4_repr = Ipv4Repr::parse(&ipv4_packet, &ChecksumCapabilities::default()).ok()?; + if ipv4_repr.next_header != IpProtocol::Udp { + return None; + } + + let udp_packet = UdpPacket::new_checked(ipv4_packet.payload()).ok()?; + let udp = UdpRepr::parse( + &udp_packet, + &IpAddress::Ipv4(ipv4_repr.src_addr), + &IpAddress::Ipv4(ipv4_repr.dst_addr), + &ChecksumCapabilities::default(), + ) + .ok()?; + let dhcp_packet = DhcpPacket::new_checked(udp_packet.payload()).ok()?; + let dhcp = DhcpRepr::parse(&dhcp_packet).ok()?; + Some(ParsedDhcp { + src_addr: ipv4_repr.src_addr, + udp, + message_type: dhcp.message_type, + transaction_id: dhcp.transaction_id, + client_hardware_address: dhcp.client_hardware_address, + your_ip: dhcp.your_ip, + server_identifier: dhcp.server_identifier, + subnet_mask: dhcp.subnet_mask, + router: dhcp.router, + dns_servers: dhcp + .dns_servers + .map(|servers| servers.iter().copied().collect()) + .unwrap_or_default(), + }) +} + /// Minimal DHCP server configuration and one-client lease state. pub struct DhcpServer { /// Server address, also advertised as router and server identifier. @@ -73,32 +122,16 @@ impl DhcpServer { return None; } - let ipv4_packet = Ipv4Packet::new_checked(packet).ok()?; - let ipv4_repr = Ipv4Repr::parse(&ipv4_packet, &ChecksumCapabilities::default()).ok()?; - if ipv4_repr.next_header != IpProtocol::Udp { - return None; - } - - let udp_packet = UdpPacket::new_checked(ipv4_packet.payload()).ok()?; - let udp_repr = UdpRepr::parse( - &udp_packet, - &IpAddress::Ipv4(ipv4_repr.src_addr), - &IpAddress::Ipv4(ipv4_repr.dst_addr), - &ChecksumCapabilities::default(), - ) - .ok()?; + let parsed = parse_dhcp_packet(packet)?; // Client -> server uses UDP src=68, dst=67. - if udp_repr.src_port != DHCP_CLIENT_PORT || udp_repr.dst_port != DHCP_SERVER_PORT { + if parsed.udp.src_port != DHCP_CLIENT_PORT || parsed.udp.dst_port != DHCP_SERVER_PORT { return None; } - let dhcp_packet = DhcpPacket::new_checked(udp_packet.payload()).ok()?; - let dhcp_repr = DhcpRepr::parse(&dhcp_packet).ok()?; - - let client_mac = dhcp_repr.client_hardware_address; - let xid = dhcp_repr.transaction_id; + let client_mac = parsed.client_hardware_address; + let xid = parsed.transaction_id; - let reply_type = match dhcp_repr.message_type { + let reply_type = match parsed.message_type { DhcpMessageType::Discover => { info!( "[dhcp-srv] Discover from {client_mac} -> Offer {}", diff --git a/net/ax-net/src/lib.rs b/net/ax-net/src/lib.rs index 4e8ce23571..42e0d4e8ae 100644 --- a/net/ax-net/src/lib.rs +++ b/net/ax-net/src/lib.rs @@ -37,6 +37,7 @@ extern crate alloc; #[cfg(test)] extern crate std; +mod addr; mod config; mod consts; mod device; @@ -85,6 +86,14 @@ use spin::{LazyLock, Once}; #[cfg(feature = "vsock")] pub use self::device::{VsockDevice, VsockDeviceList}; +use self::{ + addr::mask_from_prefix, + device::{EthernetDevice, LoopbackDevice}, + listen_table::ListenTable, + router::{RouteTable, Router, Rule, SharedRouteTable}, + service::{NetControl, NetInterface, Service}, + wrapper::SocketSetWrapper, +}; pub use self::{ config::{ DeviceBinding, InterfaceConfig, InterfaceFlags, InterfaceId, InterfaceInfo, InterfaceKind, @@ -101,13 +110,6 @@ pub use self::{ SocketOps, }, }; -use self::{ - device::{EthernetDevice, LoopbackDevice}, - listen_table::ListenTable, - router::{RouteTable, Router, Rule, SharedRouteTable}, - service::{NetControl, NetInterface, Service}, - wrapper::SocketSetWrapper, -}; static LISTEN_TABLE: LazyLock = LazyLock::new(ListenTable::new); static SOCKET_SET: LazyLock = LazyLock::new(SocketSetWrapper::new); @@ -120,11 +122,29 @@ static NET_POLL_REQUESTED: AtomicBool = AtomicBool::new(false); static NET_POLL_WAKE: WaitQueue = WaitQueue::new(); static NET_POLL_DEVICE_WAKER: LazyLock = LazyLock::new(|| Waker::from(Arc::new(NetPollWake))); -type DeferredPollWake = (Arc, IoEvents); +type DeferredPollEntry = (Arc, IoEvents); static DEFERRED_POLL_WAKE_PENDING: AtomicBool = AtomicBool::new(false); -static DEFERRED_POLL_WAKES: LazyLock>> = +static DEFERRED_POLL_WAKES: LazyLock>> = LazyLock::new(|| Mutex::new(Vec::new())); +pub(crate) struct DeferPollWake { + pub(crate) poll: Arc, + pub(crate) ready: IoEvents, +} + +impl Wake for DeferPollWake { + fn wake(self: Arc) { + self.wake_by_ref(); + } + + fn wake_by_ref(self: &Arc) { + // smoltcp invokes socket wakers from the net poll task context after + // updating readiness. The socket set may still be locked there, so + // defer the actual PollSet wake to the net worker outer loop. + defer_poll_wake(self.poll.clone(), self.ready); + } +} + /// Registry of wireless control-plane handles, keyed by interface name. /// /// Populated when a wireless device is registered (the runtime captures a @@ -165,67 +185,14 @@ pub fn init_network(mut net_devs: EthernetDeviceList, config: NetworkConfig) { info!("Initialize network subsystem..."); - for cfg in &config.interfaces { - if cfg.name == "lo" { - panic!("interface name 'lo' is reserved"); - } - if cfg.dhcp && cfg.static_ip.is_some() { - panic!( - "interface {} has both DHCP and static IP configured", - cfg.name - ); - } - if let Some(static_cfg) = &cfg.static_ip { - if static_cfg.ip.is_unspecified() { - panic!("Invalid static IP for {}: unspecified address", cfg.name); - } - if static_cfg.prefix_len > 32 { - panic!("Invalid static IP for {}: prefix length > 32", cfg.name); - } - } - for (i, dns) in cfg.dns_servers.iter().enumerate() { - if dns.is_unspecified() { - panic!( - "Invalid DNS server for {} at index {}: unspecified address", - cfg.name, i - ); - } - } - } - for (i, dns) in config.default_dns_servers.iter().enumerate() { - if dns.is_unspecified() { - panic!("Invalid DNS server at index {}: unspecified address", i); - } - } + validate_config(&config); let routes: SharedRouteTable = Arc::new(spin::RwLock::new(RouteTable::new())); let mut router = Router::new(routes.clone()); let mut interfaces = Vec::new(); let mut dns = Vec::new(); - let lo_id = InterfaceId::LOOPBACK; - let lo_dev = router.add_device(lo_id, Box::new(LoopbackDevice::new())); - - let lo_ip = Ipv4Cidr::new(Ipv4Address::new(127, 0, 0, 1), 8); - router.add_rule(Rule::new( - lo_ip.into(), - None, - lo_dev, - lo_id, - lo_ip.address().into(), - 0, - )); - interfaces.push(NetInterface { - id: lo_id, - name: "lo".to_owned(), - kind: InterfaceKind::Loopback, - mac: None, - ipv4: Some(lo_ip), - gateway: None, - mtu: consts::STANDARD_MTU, - metric: 0, - flags: InterfaceFlags::UP | InterfaceFlags::RUNNING | InterfaceFlags::LOOPBACK, - }); + let lo_ip = register_loopback(&mut router, &mut interfaces); if net_devs.is_empty() { warn!(" No network device found!"); @@ -314,27 +281,9 @@ pub fn init_network(mut net_devs: EthernetDeviceList, config: NetworkConfig) { }); } - for (i, used) in used_configs.iter().enumerate() { - if !used { - panic!( - "interface config {} did not match any device", - config.interfaces[i].name - ); - } - } + ensure_all_interface_configs_used(&config, &used_configs); - dns.extend( - config - .default_dns_servers - .iter() - .copied() - .map(|server| config::DnsServerEntry { - server: Ipv4Address::from(server.octets()), - interface_id: lo_id, - metric: u32::MAX, - source: config::DnsSource::Fallback, - }), - ); + add_default_dns_servers(&config, &mut dns); for name in router.device_names() { info!("Device: {}", name); @@ -363,6 +312,94 @@ pub fn init_network(mut net_devs: EthernetDeviceList, config: NetworkConfig) { } } +fn validate_config(config: &NetworkConfig) { + for cfg in &config.interfaces { + if cfg.name == "lo" { + panic!("interface name 'lo' is reserved"); + } + if cfg.dhcp && cfg.static_ip.is_some() { + panic!( + "interface {} has both DHCP and static IP configured", + cfg.name + ); + } + if let Some(static_cfg) = &cfg.static_ip { + if static_cfg.ip.is_unspecified() { + panic!("Invalid static IP for {}: unspecified address", cfg.name); + } + if static_cfg.prefix_len > 32 { + panic!("Invalid static IP for {}: prefix length > 32", cfg.name); + } + } + for (i, dns) in cfg.dns_servers.iter().enumerate() { + if dns.is_unspecified() { + panic!( + "Invalid DNS server for {} at index {}: unspecified address", + cfg.name, i + ); + } + } + } + for (i, dns) in config.default_dns_servers.iter().enumerate() { + if dns.is_unspecified() { + panic!("Invalid DNS server at index {}: unspecified address", i); + } + } +} + +fn register_loopback(router: &mut Router, interfaces: &mut Vec) -> Ipv4Cidr { + let lo_id = InterfaceId::LOOPBACK; + let lo_dev = router.add_device(lo_id, Box::new(LoopbackDevice::new())); + + let lo_ip = Ipv4Cidr::new(Ipv4Address::new(127, 0, 0, 1), 8); + router.add_rule(Rule::new( + lo_ip.into(), + None, + lo_dev, + lo_id, + lo_ip.address().into(), + 0, + )); + interfaces.push(NetInterface { + id: lo_id, + name: "lo".to_owned(), + kind: InterfaceKind::Loopback, + mac: None, + ipv4: Some(lo_ip), + gateway: None, + mtu: consts::STANDARD_MTU, + metric: 0, + flags: InterfaceFlags::UP | InterfaceFlags::RUNNING | InterfaceFlags::LOOPBACK, + }); + lo_ip +} + +fn ensure_all_interface_configs_used(config: &NetworkConfig, used_configs: &[bool]) { + for (i, used) in used_configs.iter().enumerate() { + if !used { + panic!( + "interface config {} did not match any device", + config.interfaces[i].name + ); + } + } +} + +fn add_default_dns_servers(config: &NetworkConfig, dns: &mut Vec) { + dns.extend( + config + .default_dns_servers + .iter() + .copied() + .map(|server| config::DnsServerEntry { + server: Ipv4Address::from(server.octets()), + interface_id: InterfaceId::LOOPBACK, + metric: u32::MAX, + source: config::DnsSource::Fallback, + }), + ); +} + fn find_interface_config( configs: &[InterfaceConfig], used: &mut [bool], @@ -408,10 +445,6 @@ pub fn init_vsock(mut vsock_devs: device::VsockDeviceList) { } } -fn poll_once() -> bool { - get_service().poll(&mut SOCKET_SET.inner.lock()) -} - fn poll_until_idle() { POLL_AGAIN.store(true, Ordering::Release); loop { @@ -423,7 +456,7 @@ fn poll_until_idle() { } while POLL_AGAIN.swap(false, Ordering::AcqRel) { - while poll_once() {} + while get_service().poll(&mut SOCKET_SET.inner.lock()) {} } POLLING_INTERFACES.store(false, Ordering::Release); if !POLL_AGAIN.load(Ordering::Acquire) { @@ -432,26 +465,20 @@ fn poll_until_idle() { } } -/// Request network polling from the dedicated net-poll worker. -/// -/// This function is retained as a public trigger/debug entry. It no longer -/// synchronously drives the whole protocol stack from the caller's context. -pub fn poll_interfaces() { - request_poll(); -} - /// Request network polling. /// /// This is the lightweight entry used by socket and device paths. pub fn request_poll() { - NET_POLL_REQUESTED.store(true, Ordering::Release); - NET_POLL_WAKE.notify_one(true); + if !NET_POLL_REQUESTED.swap(true, Ordering::AcqRel) { + NET_POLL_WAKE.notify_one(true); + } } pub(crate) fn defer_poll_wake(poll: Arc, ready: IoEvents) { DEFERRED_POLL_WAKES.lock().push((poll, ready)); - DEFERRED_POLL_WAKE_PENDING.store(true, Ordering::Release); - NET_POLL_WAKE.notify_one(true); + if !DEFERRED_POLL_WAKE_PENDING.swap(true, Ordering::AcqRel) { + NET_POLL_WAKE.notify_one(true); + } } fn drain_deferred_poll_wakes() { @@ -521,13 +548,13 @@ pub struct NetConfig { /// Registers an extra Ethernet device with a static IPv4 address. /// -/// If `dedicated_poll` is set, RX readiness is driven by [`notify_oob_rx`] +/// If `dedicated_poll` is set, RX readiness is driven by [`wake_net_task_irq`] /// instead of the shared Ethernet IRQ framework. pub fn register_device_with_config(dev: Box, config: NetConfig) { let mac = EthernetAddress(dev.mac_address()); - let server_ip = Ipv4Address::new(config.ip[0], config.ip[1], config.ip[2], config.ip[3]); + let server_ip = Ipv4Address::from(config.ip); let cidr = Ipv4Cidr::new(server_ip, config.prefix_len); - // A dedicated-poll device gets RX out-of-band (via `notify_oob_rx` and the + // A dedicated-poll device gets RX out-of-band (via `wake_net_task_irq` and the // shared net poll task), so its socket wakers must be armed even though it // has no ethernet IRQ registration. let eth_dev = if config.dedicated_poll { @@ -537,13 +564,9 @@ pub fn register_device_with_config(dev: Box, config: NetConf }; let dev_idx = get_service().register_static_device(config.name.clone(), eth_dev, mac, cidr); if let Some(client_ip) = config.dhcp_server_client_ip { - let client_ip = Ipv4Address::new(client_ip[0], client_ip[1], client_ip[2], client_ip[3]); - get_service().enable_dhcp_server( - dev_idx, - server_ip, - client_ip, - prefix_to_mask(config.prefix_len), - ); + let client_ip = Ipv4Address::from(client_ip); + let subnet_mask = mask_from_prefix(config.prefix_len); + get_service().enable_dhcp_server(dev_idx, server_ip, client_ip, subnet_mask); } info!("{}: up, mac {mac}, ip {cidr}", config.name); @@ -603,11 +626,14 @@ pub fn reconfigure_wifi(name: &str, mode: WifiMode<'_>) -> AxResult<()> { // (possibly new) MAC. The registry lock is released before touching the // stack service to avoid holding two locks across the blocking path. let mac = { - let controls = WIFI_CONTROLS.lock(); - let (_, handle) = controls - .iter() - .find(|(n, _)| n == name) - .ok_or(AxError::NoSuchDevice)?; + let handle = { + let controls = WIFI_CONTROLS.lock(); + controls + .iter() + .find(|(n, _)| n == name) + .map(|(_, handle)| handle.clone()) + .ok_or(AxError::NoSuchDevice)? + }; let ctrl = handle.wifi_control().ok_or(AxError::NoSuchDevice)?; match &mode { WifiMode::Station { ssid, password } => ctrl @@ -632,15 +658,15 @@ pub fn reconfigure_wifi(name: &str, mode: WifiMode<'_>) -> AxResult<()> { dhcp_client_ip, .. } => { - let server_ip = Ipv4Address::new(ip[0], ip[1], ip[2], ip[3]); - let client_ip = dhcp_client_ip.map(|c| Ipv4Address::new(c[0], c[1], c[2], c[3])); + let server_ip = Ipv4Address::from(ip); + let client_ip = dhcp_client_ip.map(Ipv4Address::from); service.reconfigure_as_ap(dev, server_ip, prefix_len, client_ip); } } } // Kick a poll so the new addressing takes effect immediately. - poll_interfaces(); + request_poll(); info!("{name}: wifi mode switch complete"); Ok(()) } @@ -648,37 +674,13 @@ pub fn reconfigure_wifi(name: &str, mode: WifiMode<'_>) -> AxResult<()> { /// Wakes the net poll task from a hard IRQ callback. /// /// The IRQ path must only publish small pending state and call this wrapper. -/// The deferred net task performs `poll_interfaces()` and wakes socket waiters -/// from ordinary task context. +/// The deferred net task requests polling and wakes socket waiters from ordinary +/// task context. pub fn wake_net_task_irq() { NET_IRQ_NOTIFY.notify_irq(); NET_POLL_WAKE.notify_one_from_irq(); } -/// Wakes the out-of-band RX poll task; intended as a device RX-data callback. -/// -/// A device whose RX path sits outside the ethernet IRQ framework (e.g. an SDIO -/// chip owning its own card interrupt) registers this as its RX callback. It -/// only signals here; the dedicated poll task does the actual stack polling, so -/// the device's RX thread is never blocked on the stack. -pub fn notify_oob_rx() { - wake_net_task_irq(); -} - -/// Convenience helper for retrieving `eth0` IPv4 configuration. -pub fn eth0_ipv4_config() -> Option { - get_service().eth0_ipv4_config() -} - -fn prefix_to_mask(prefix_len: u8) -> Ipv4Address { - let bits = if prefix_len == 0 { - 0 - } else { - u32::MAX << (32 - prefix_len.min(32) as u32) - }; - Ipv4Address::from_bits(bits) -} - fn next_poll_delay() -> Duration { const IDLE_POLL_INTERVAL: Duration = Duration::from_millis(100); let next = { @@ -849,15 +851,6 @@ fn wait_for_dhcp_bootstrap() { warn!("DHCP bootstrap timed out"); } -pub(crate) fn endpoint_from_ip_endpoint( - endpoint: smoltcp::wire::IpEndpoint, -) -> smoltcp::wire::IpListenEndpoint { - smoltcp::wire::IpListenEndpoint { - addr: Some(endpoint.addr), - port: endpoint.port, - } -} - #[cfg(test)] pub(crate) mod test_support { use alloc::{boxed::Box, sync::Arc, vec, vec::Vec}; diff --git a/net/ax-net/src/listen_table.rs b/net/ax-net/src/listen_table.rs index 395a5b4bf6..862f598f56 100644 --- a/net/ax-net/src/listen_table.rs +++ b/net/ax-net/src/listen_table.rs @@ -25,12 +25,13 @@ //! module. The required order is `SOCKET_SET -> listen-table bucket`; this file //! must never acquire the outer service lock. -use alloc::{boxed::Box, collections::VecDeque, sync::Arc, task::Wake, vec, vec::Vec}; +use alloc::{collections::VecDeque, sync::Arc, vec, vec::Vec}; use core::task::Waker; use ax_errno::{AxError, AxResult}; use ax_sync::Mutex; use axpoll::{IoEvents, PollSet}; +use hashbrown::HashMap; use smoltcp::{ iface::{SocketHandle, SocketSet}, socket::tcp::{self, SocketBuffer, State}, @@ -38,19 +39,18 @@ use smoltcp::{ }; use crate::{ - SOCKET_SET, + DeferPollWake, SOCKET_SET, + addr::listen_addrs_conflict, consts::{LISTEN_QUEUE_SIZE, TCP_RX_BUF_LEN, TCP_TX_BUF_LEN}, }; -const PORT_NUM: usize = 65536; - struct ListenTableEntryInner { /// Local endpoint accepted by this listener. listen_endpoint: IpListenEndpoint, /// Maximum pending child sockets. backlog: usize, /// Pending smoltcp child sockets waiting for accept(). - syn_queue: VecDeque, + syn_queue: VecDeque, /// Wakes accept/poll waiters when child readiness changes. accept_poll: Arc, } @@ -66,28 +66,6 @@ pub(crate) struct AcceptedTcp { pub(crate) remote_endpoint: IpEndpoint, } -#[derive(Clone, Copy)] -struct PendingTcp { - accepted: AcceptedTcp, -} - -struct AcceptWake { - poll: Arc, -} - -impl Wake for AcceptWake { - fn wake(self: Arc) { - self.wake_by_ref(); - } - - fn wake_by_ref(self: &Arc) { - // smoltcp invokes this from the net poll task context after publishing - // child socket readiness. Defer the actual PollSet wake until the net - // worker has released service/socket locks. - crate::defer_poll_wake(self.poll.clone(), IoEvents::IN); - } -} - impl ListenTableEntryInner { /// Creates a listener entry and clamps backlog to the global limit. pub fn new(listen_endpoint: IpListenEndpoint, backlog: usize) -> Self { @@ -115,15 +93,15 @@ impl ListenTableEntryInner { fn into_handles(self) -> Vec { self.syn_queue .into_iter() - .map(|pending| pending.accepted.handle) + .map(|pending| pending.handle) .collect() } /// Returns whether a child socket for this endpoint pair already exists. fn has_pending(&self, src: IpEndpoint, dst: IpEndpoint) -> bool { - self.syn_queue.iter().any(|pending| { - pending.accepted.local_endpoint == dst && pending.accepted.remote_endpoint == src - }) + self.syn_queue + .iter() + .any(|pending| pending.local_endpoint == dst && pending.remote_endpoint == src) } } @@ -131,25 +109,23 @@ type ListenTableEntry = Arc>>; /// Per-port table of active TCP listeners. pub struct ListenTable { - tcp: Box<[ListenTableEntry]>, + tcp: Mutex>, } impl ListenTable { /// Creates an empty listen table indexed by TCP port. pub fn new() -> Self { - let tcp = unsafe { - let mut buf = Box::new_uninit_slice(PORT_NUM); - for i in 0..PORT_NUM { - buf[i].write(Arc::default()); - } - buf.assume_init() - }; - Self { tcp } + Self { + tcp: Mutex::new(HashMap::new()), + } } /// Checks whether a listen endpoint can be registered. pub fn can_listen(&self, listen_endpoint: IpListenEndpoint) -> bool { - self.tcp[listen_endpoint.port as usize] + let Some(entries) = self.listen_entry(listen_endpoint.port) else { + return true; + }; + entries .lock() .iter() .all(|entry| !listen_addrs_conflict(entry.listen_endpoint.addr, listen_endpoint.addr)) @@ -159,7 +135,8 @@ impl ListenTable { pub fn listen(&self, listen_endpoint: IpListenEndpoint, backlog: usize) -> AxResult { let port = listen_endpoint.port; assert_ne!(port, 0); - let mut entries = self.tcp[port as usize].lock(); + let entries = self.listen_entry_or_create(port); + let mut entries = entries.lock(); if entries .iter() .any(|entry| listen_addrs_conflict(entry.listen_endpoint.addr, listen_endpoint.addr)) @@ -174,23 +151,34 @@ impl ListenTable { /// Removes a listener and destroys any unaccepted child sockets. pub fn unlisten(&self, listen_endpoint: IpListenEndpoint) { debug!("TCP socket unlisten on {}", listen_endpoint); - let handles = { - let mut entries = self.tcp[listen_endpoint.port as usize].lock(); + let (handles, remove_port) = { + let Some(entries) = self.listen_entry(listen_endpoint.port) else { + return; + }; + let mut entries = entries.lock(); let Some(idx) = entries .iter() .position(|entry| entry.listen_endpoint == listen_endpoint) else { return; }; - entries.swap_remove(idx).into_handles() + let handles = entries.swap_remove(idx).into_handles(); + (handles, entries.is_empty()) }; + if remove_port { + self.tcp.lock().remove(&listen_endpoint.port); + } for handle in handles { SOCKET_SET.remove(handle); } } - fn listen_entry(&self, port: u16) -> Arc>> { - self.tcp[port as usize].clone() + fn listen_entry(&self, port: u16) -> Option { + self.tcp.lock().get(&port).cloned() + } + + fn listen_entry_or_create(&self, port: u16) -> ListenTableEntry { + self.tcp.lock().entry(port).or_default().clone() } // Callers pass the locked SocketSet to keep the global order: @@ -201,18 +189,20 @@ impl ListenTable { listen_endpoint: IpListenEndpoint, sockets: &SocketSet<'_>, ) -> AxResult { - let entries = self.listen_entry(listen_endpoint.port); + let Some(entries) = self.listen_entry(listen_endpoint.port) else { + warn!("accept before listen"); + return Err(AxError::InvalidInput); + }; let table = entries.lock(); if let Some(entry) = table .iter() .find(|entry| entry.listen_endpoint == listen_endpoint) { - return Ok(entry + Ok(entry .syn_queue .iter() - .any(|pending| is_acceptable(sockets, pending.accepted.handle))); - } - { + .any(|pending| is_acceptable(sockets, pending.handle))) + } else { warn!("accept before listen"); Err(AxError::InvalidInput) } @@ -224,7 +214,10 @@ impl ListenTable { listen_endpoint: IpListenEndpoint, sockets: &mut SocketSet<'_>, ) -> AxResult { - let entries = self.listen_entry(listen_endpoint.port); + let Some(entries) = self.listen_entry(listen_endpoint.port) else { + warn!("accept before listen"); + return Err(AxError::InvalidInput); + }; let mut table = entries.lock(); let Some(entry) = table .iter_mut() @@ -234,10 +227,10 @@ impl ListenTable { return Err(AxError::InvalidInput); }; - let syn_queue: &mut VecDeque = &mut entry.syn_queue; + let syn_queue: &mut VecDeque = &mut entry.syn_queue; let mut idx = 0; while idx < syn_queue.len() { - let handle = syn_queue[idx].accepted.handle; + let handle = syn_queue[idx].handle; if is_closed_without_data(sockets, handle) { syn_queue.swap_remove_front(idx); sockets.remove(handle); @@ -251,7 +244,7 @@ impl ListenTable { syn_queue.len() ); } - return Ok(syn_queue.swap_remove_front(idx).unwrap().accepted); + return Ok(syn_queue.swap_remove_front(idx).unwrap()); } idx += 1; } @@ -260,7 +253,7 @@ impl ListenTable { /// Returns the listener readiness poll set for lock-free registration. pub fn accept_poll(&self, listen_endpoint: IpListenEndpoint) -> Option> { - let entries = self.listen_entry(listen_endpoint.port); + let entries = self.listen_entry(listen_endpoint.port)?; let table = entries.lock(); table .iter() @@ -270,7 +263,10 @@ impl ListenTable { /// Builds the smoltcp child-progress waker for a listener poll set. pub fn accept_waker(&self, accept_poll: Arc) -> Waker { - Waker::from(Arc::new(AcceptWake { poll: accept_poll })) + Waker::from(Arc::new(DeferPollWake { + poll: accept_poll, + ready: IoEvents::IN, + })) } /// Registers a waker for queued child progress. @@ -281,7 +277,9 @@ impl ListenTable { accept_poll: &Arc, waker: &Waker, ) { - let entries = self.listen_entry(listen_endpoint.port); + let Some(entries) = self.listen_entry(listen_endpoint.port) else { + return; + }; let table = entries.lock(); let Some(entry) = table.iter().find(|entry| { entry.listen_endpoint == listen_endpoint && Arc::ptr_eq(&entry.accept_poll, accept_poll) @@ -289,7 +287,7 @@ impl ListenTable { return; }; for pending in &entry.syn_queue { - let socket: &mut tcp::Socket = sockets.get_mut(pending.accepted.handle); + let socket: &mut tcp::Socket = sockets.get_mut(pending.handle); socket.register_recv_waker(waker); socket.register_send_waker(waker); } @@ -302,7 +300,9 @@ impl ListenTable { dst: IpEndpoint, sockets: &mut SocketSet<'_>, ) { - let entries = self.listen_entry(dst.port); + let Some(entries) = self.listen_entry(dst.port) else { + return; + }; let wake_poll = { let mut table = entries.lock(); let Some(entry) = table @@ -339,12 +339,10 @@ impl ListenTable { "TCP socket {}: prepare for connection {} -> {}", handle, src, entry.listen_endpoint ); - entry.syn_queue.push_back(PendingTcp { - accepted: AcceptedTcp { - handle, - local_endpoint: dst, - remote_endpoint: src, - }, + entry.syn_queue.push_back(AcceptedTcp { + handle, + local_endpoint: dst, + remote_endpoint: src, }); entry.accept_poll.clone() }; @@ -355,13 +353,6 @@ impl ListenTable { } } -fn listen_addrs_conflict( - a: Option, - b: Option, -) -> bool { - a.is_none() || b.is_none() || a == b -} - fn is_acceptable(sockets: &SocketSet<'_>, handle: SocketHandle) -> bool { let socket: &tcp::Socket = sockets.get(handle); match socket.state() { diff --git a/net/ax-net/src/orphan.rs b/net/ax-net/src/orphan.rs index d58a17955f..70cd8a29b2 100644 --- a/net/ax-net/src/orphan.rs +++ b/net/ax-net/src/orphan.rs @@ -164,9 +164,3 @@ pub(crate) fn reap_orphans(timestamp: Instant, sockets: &mut SocketSet<'_>) { ); } } - -/// Get current orphan socket count (for diagnostics). -#[allow(dead_code)] -pub(crate) fn orphan_count() -> usize { - ORPHAN_SOCKETS.lock().len() -} diff --git a/net/ax-net/src/raw.rs b/net/ax-net/src/raw.rs index 6aa6c9e5be..994603d94a 100644 --- a/net/ax-net/src/raw.rs +++ b/net/ax-net/src/raw.rs @@ -22,6 +22,14 @@ //! Loopback ICMP-style traffic may be delivered through a local fast path. For //! connected raw sockets, packets from other peers can be skipped or deferred //! without corrupting the smoltcp receive queue format. +//! +//! # Locking +//! +//! Raw sockets keep their small deferred-packet slots behind IRQ-off spin locks +//! because packet delivery may be inspected while the net poll worker is +//! servicing device-originated receive work. These locks are only held around +//! `Option>` swaps and never across route lookup, smoltcp polling, or +//! userspace buffer I/O. use alloc::vec; use core::{ @@ -53,17 +61,28 @@ use crate::{ request_poll, }; -/// Allocates a smoltcp raw socket for one IP version and protocol. -pub(crate) fn new_raw_socket( - ip_version: IpVersion, - ip_protocol: IpProtocol, -) -> smol::Socket<'static> { - smol::Socket::new( - Some(ip_version), - Some(ip_protocol), - smol::PacketBuffer::new(vec![PacketMetadata::EMPTY; 256], vec![0; RAW_RX_BUF_LEN]), - smol::PacketBuffer::new(vec![PacketMetadata::EMPTY; 256], vec![0; RAW_TX_BUF_LEN]), - ) +enum RawIpHeader { + Ipv4(Ipv4Repr), + Ipv6(Ipv6Repr), +} + +impl RawIpHeader { + fn buffer_len(&self) -> usize { + match self { + Self::Ipv4(header) => header.buffer_len(), + Self::Ipv6(header) => header.buffer_len(), + } + } + + fn emit(&self, buf: &mut [u8]) { + match self { + Self::Ipv4(header) => header.emit( + &mut Ipv4Packet::new_unchecked(buf), + &smoltcp::phy::ChecksumCapabilities::ignored(), + ), + Self::Ipv6(header) => header.emit(&mut Ipv6Packet::new_unchecked(buf)), + } + } } /// A raw IP socket used for ICMP and ICMPv6 traffic. @@ -93,11 +112,15 @@ pub struct RawSocket { impl RawSocket { /// Creates a raw socket for the given IP version and protocol. pub fn new(ip_version: IpVersion, ip_protocol: IpProtocol) -> Self { - let handle = SOCKET_SET.add(new_raw_socket(ip_version, ip_protocol)); let general = GeneralOptions::new(3, 2, u8::from(ip_protocol) as i32); // SOCK_RAW general.set_device_binding(DeviceBinding::default()); Self { - handle, + handle: SOCKET_SET.add(smol::Socket::new( + Some(ip_version), + Some(ip_protocol), + smol::PacketBuffer::new(vec![PacketMetadata::EMPTY; 256], vec![0; RAW_RX_BUF_LEN]), + smol::PacketBuffer::new(vec![PacketMetadata::EMPTY; 256], vec![0; RAW_TX_BUF_LEN]), + )), ip_version, local_addr: RwLock::new(None), peer_addr: RwLock::new(None), @@ -126,6 +149,37 @@ impl RawSocket { SOCKET_SET.with_socket_mut::(self.handle, f) } + fn outgoing_ip_header( + &self, + local: IpAddress, + remote: IpAddress, + next_header: IpProtocol, + payload_len: usize, + hop_limit: u8, + ) -> RawIpHeader { + match (self.ip_version, local, remote) { + (IpVersion::Ipv4, IpAddress::Ipv4(src_addr), IpAddress::Ipv4(dst_addr)) => { + RawIpHeader::Ipv4(Ipv4Repr { + src_addr, + dst_addr, + next_header, + payload_len, + hop_limit, + }) + } + (IpVersion::Ipv6, IpAddress::Ipv6(src_addr), IpAddress::Ipv6(dst_addr)) => { + RawIpHeader::Ipv6(Ipv6Repr { + src_addr, + dst_addr, + next_header, + payload_len, + hop_limit, + }) + } + _ => unreachable!(), + } + } + /// Validates that an address belongs to this socket's IP version. fn check_ip_version(&self, addr: IpAddress) -> AxResult { match (self.ip_version, addr) { @@ -160,8 +214,12 @@ impl RawSocket { .source) } - /// Parses a complete IP packet and returns its source plus deliverable bytes. - fn parse_ip_packet<'a>(&self, packet: &'a [u8]) -> AxResult<(IpAddress, &'a [u8])> { + /// Splits a complete IP packet into source and bytes returned to userspace. + /// + /// Linux raw IPv4 receive returns the IP header plus payload, while raw IPv6 + /// receive returns only the transport payload. The returned slice preserves + /// that ABI difference. + fn split_packet_for_delivery<'a>(&self, packet: &'a [u8]) -> AxResult<(IpAddress, &'a [u8])> { match self.ip_version { IpVersion::Ipv4 => { let packet = Ipv4Packet::new_checked(packet) @@ -342,77 +400,14 @@ impl SocketOps for RawSocket { let next_header = socket.ip_protocol().expect("raw socket protocol"); let hop_limit = (*self.ttl.read()).unwrap_or(64); - let header_len = match self.ip_version { - IpVersion::Ipv4 => Ipv4Repr { - src_addr: match local { - IpAddress::Ipv4(addr) => addr, - _ => unreachable!(), - }, - dst_addr: match remote { - IpAddress::Ipv4(addr) => addr, - _ => unreachable!(), - }, - next_header, - payload_len, - hop_limit, - } - .buffer_len(), - IpVersion::Ipv6 => Ipv6Repr { - src_addr: match local { - IpAddress::Ipv6(addr) => addr, - _ => unreachable!(), - }, - dst_addr: match remote { - IpAddress::Ipv6(addr) => addr, - _ => unreachable!(), - }, - next_header, - payload_len, - hop_limit, - } - .buffer_len(), - }; + let header = + self.outgoing_ip_header(local, remote, next_header, payload_len, hop_limit); + let header_len = header.buffer_len(); let buf = socket .send(header_len + payload_len) .map_err(|_| AxError::WouldBlock)?; - match self.ip_version { - IpVersion::Ipv4 => { - let header = Ipv4Repr { - src_addr: match local { - IpAddress::Ipv4(addr) => addr, - _ => unreachable!(), - }, - dst_addr: match remote { - IpAddress::Ipv4(addr) => addr, - _ => unreachable!(), - }, - next_header, - payload_len, - hop_limit, - }; - header.emit( - &mut Ipv4Packet::new_unchecked(&mut *buf), - &smoltcp::phy::ChecksumCapabilities::ignored(), - ); - } - IpVersion::Ipv6 => { - let header = Ipv6Repr { - src_addr: match local { - IpAddress::Ipv6(addr) => addr, - _ => unreachable!(), - }, - dst_addr: match remote { - IpAddress::Ipv6(addr) => addr, - _ => unreachable!(), - }, - next_header, - payload_len, - hop_limit, - }; - header.emit(&mut Ipv6Packet::new_unchecked(&mut *buf)); - } - } + header.emit(&mut *buf); let written = src.read(&mut buf[header_len..])?; if next_header == IpProtocol::Icmpv6 { @@ -455,7 +450,7 @@ impl SocketOps for RawSocket { *self.deferred_rx.lock() = Some((source, packet)); return Err(AxError::WouldBlock); } - let (_, payload) = self.parse_ip_packet(&packet)?; + let (_, payload) = self.split_packet_for_delivery(&packet)?; return self.deliver_packet(source, payload, &mut dst, &mut options); } @@ -473,7 +468,7 @@ impl SocketOps for RawSocket { let wire_packet = if options.flags.contains(RecvFlags::PEEK) { let packet = socket.peek().map_err(|_| AxError::WouldBlock)?; - let (source, _) = self.parse_ip_packet(packet)?; + let (source, _) = self.split_packet_for_delivery(packet)?; if let Some(peer) = *self.peer_addr.read() && source != peer { @@ -483,7 +478,7 @@ impl SocketOps for RawSocket { } else { socket.recv().map_err(|_| AxError::WouldBlock)? }; - let (source, packet) = self.parse_ip_packet(wire_packet)?; + let (source, packet) = self.split_packet_for_delivery(wire_packet)?; if !self.source_matches_peer(source) { *self.deferred_rx.lock() = Some((source, wire_packet.to_vec())); diff --git a/net/ax-net/src/router.rs b/net/ax-net/src/router.rs index c808267627..f5ad6ae91f 100644 --- a/net/ax-net/src/router.rs +++ b/net/ax-net/src/router.rs @@ -47,7 +47,7 @@ use core::{ task::Waker, }; -use ax_hal::time::{NANOS_PER_MICROS, wall_time_nanos}; +use ax_hal::time::{NANOS_PER_MICROS, monotonic_time_nanos}; use ax_sync::Mutex; use ax_task::WaitQueue; use axpoll::IoEvents; @@ -70,6 +70,8 @@ use crate::{ device::{ArpEntry, Device}, }; +const DEVICE_RX_WORKER_BATCH: usize = 16; + #[derive(Debug)] pub struct Rule { /// Destination prefix matched by this route. @@ -308,7 +310,7 @@ impl Wake for DeviceRxWake { } fn now() -> Instant { - Instant::from_micros_const((wall_time_nanos() / NANOS_PER_MICROS) as i64) + Instant::from_micros_const((monotonic_time_nanos() / NANOS_PER_MICROS) as i64) } #[derive(Debug, Clone, Copy)] @@ -672,7 +674,14 @@ impl Router { /// Routes smoltcp-emitted TX packets to loopback or device workers. pub fn dispatch(&mut self, _timestamp: Instant, sockets: &mut SocketSet<'_>) -> bool { let mut poll_next = false; - while let Ok((_, packet)) = self.tx_buffer.dequeue() { + let Router { + rx_buffer, + tx_buffer, + devices, + table, + .. + } = self; + while let Ok((_, packet)) = tx_buffer.dequeue() { match IpVersion::of_packet(packet).expect("got invalid IP packet") { IpVersion::Ipv4 => { let packet = smoltcp::wire::Ipv4Packet::new_checked(packet) @@ -680,39 +689,18 @@ impl Router { let src_addr = IpAddress::Ipv4(packet.src_addr()); let dst_addr = IpAddress::Ipv4(packet.dst_addr()); if packet.dst_addr().is_broadcast() { - let buf = packet.into_inner(); - // Broadcast only to Ethernet devices (not loopback) - for dev in &self.devices { - if dev.interface_id != InterfaceId::LOOPBACK { - poll_next |= dev.enqueue_tx(dst_addr, buf); - } - } + poll_next |= + dispatch_link_local_fanout(devices, dst_addr, packet.into_inner()); } else { - let routes = self.table.read(); - let Some(route) = routes.select_route_for_source(&dst_addr, &src_addr) - else { - warn!( - "No route found for source {} destination {}", - src_addr, dst_addr - ); - continue; - }; - - let dev = &self.devices[route.dev]; - if dev.interface_id == InterfaceId::LOOPBACK { - // Loopback packets are copied directly from the TX - // buffer into the RX buffer. This avoids the - // per-device worker and shared RX queue used by - // real devices. - poll_next |= inject_loopback_rx_direct( - &mut self.rx_buffer, - dst_addr, - packet.into_inner(), - sockets, - ); - } else { - poll_next |= dev.enqueue_tx(route.next_hop, packet.into_inner()); - } + poll_next |= dispatch_unicast_packet( + rx_buffer, + devices, + table, + src_addr, + dst_addr, + packet.into_inner(), + sockets, + ); } } IpVersion::Ipv6 => { @@ -721,35 +709,18 @@ impl Router { let src_addr = IpAddress::Ipv6(packet.src_addr()); let dst_addr = IpAddress::Ipv6(packet.dst_addr()); if packet.dst_addr().is_multicast() { - let buf = packet.into_inner(); - // Multicast only to Ethernet devices (not loopback) - for dev in &self.devices { - if dev.interface_id != InterfaceId::LOOPBACK { - poll_next |= dev.enqueue_tx(dst_addr, buf); - } - } + poll_next |= + dispatch_link_local_fanout(devices, dst_addr, packet.into_inner()); } else { - let routes = self.table.read(); - let Some(route) = routes.select_route_for_source(&dst_addr, &src_addr) - else { - warn!( - "No route found for source {} destination {}", - src_addr, dst_addr - ); - continue; - }; - - let dev = &self.devices[route.dev]; - if dev.interface_id == InterfaceId::LOOPBACK { - poll_next |= inject_loopback_rx_direct( - &mut self.rx_buffer, - dst_addr, - packet.into_inner(), - sockets, - ); - } else { - poll_next |= dev.enqueue_tx(route.next_hop, packet.into_inner()); - } + poll_next |= dispatch_unicast_packet( + rx_buffer, + devices, + table, + src_addr, + dst_addr, + packet.into_inner(), + sockets, + ); } } } @@ -758,6 +729,48 @@ impl Router { } } +fn dispatch_link_local_fanout( + devices: &[Arc], + dst_addr: IpAddress, + packet: &[u8], +) -> bool { + let mut poll_next = false; + for dev in devices { + if dev.interface_id != InterfaceId::LOOPBACK { + poll_next |= dev.enqueue_tx(dst_addr, packet); + } + } + poll_next +} + +fn dispatch_unicast_packet( + rx_buffer: &mut PacketBuffer, + devices: &[Arc], + table: &SharedRouteTable, + src_addr: IpAddress, + dst_addr: IpAddress, + packet: &[u8], + sockets: &mut SocketSet<'_>, +) -> bool { + let routes = table.read(); + let Some(route) = routes.select_route_for_source(&dst_addr, &src_addr) else { + warn!( + "No route found for source {} destination {}", + src_addr, dst_addr + ); + return false; + }; + + let dev = &devices[route.dev]; + if dev.interface_id == InterfaceId::LOOPBACK { + // Loopback packets are copied directly from the TX buffer into the RX + // buffer, bypassing per-device workers and the shared RX queue. + inject_loopback_rx_direct(rx_buffer, dst_addr, packet, sockets) + } else { + dev.enqueue_tx(route.next_hop, packet) + } +} + /// Injects a loopback packet directly into the smoltcp-facing RX buffer. fn inject_loopback_rx_direct( rx_buffer: &mut PacketBuffer, @@ -822,14 +835,17 @@ fn device_tx_worker(device: Arc) { /// Dedicated worker that polls one device and forwards packets to router RX. fn device_rx_worker(device: Arc) { - let mut rx_buffer = PacketBuffer::new(vec![PacketMetadata::EMPTY; 1], vec![0u8; STANDARD_MTU]); + let mut rx_buffer = PacketBuffer::new( + vec![PacketMetadata::EMPTY; DEVICE_RX_WORKER_BATCH], + vec![0u8; STANDARD_MTU * DEVICE_RX_WORKER_BATCH], + ); loop { let mut received = false; { let mut device_inner = device.inner.lock(); let mut snoop = |_packet: &[u8]| {}; - while rx_buffer.is_empty() + while !rx_buffer.is_full() && device_inner.recv(device.interface_id, &mut rx_buffer, now(), &mut snoop) { received = true; @@ -887,11 +903,13 @@ impl smoltcp::phy::TxToken for TxToken<'_> { /// Detects passive TCP opens before smoltcp consumes the incoming packet. fn snoop_tcp_packet(buf: &[u8], sockets: &mut SocketSet<'_>) { - let (protocol, src_addr, dst_addr, payload) = match IpVersion::of_packet(buf).unwrap() { + let (src_addr, dst_addr, payload) = match IpVersion::of_packet(buf).unwrap() { IpVersion::Ipv4 => { let packet = Ipv4Packet::new_unchecked(buf); + if packet.next_header() != IpProtocol::Tcp { + return; + } ( - packet.next_header(), IpAddress::Ipv4(packet.src_addr()), IpAddress::Ipv4(packet.dst_addr()), packet.payload(), @@ -899,22 +917,22 @@ fn snoop_tcp_packet(buf: &[u8], sockets: &mut SocketSet<'_>) { } IpVersion::Ipv6 => { let packet = Ipv6Packet::new_unchecked(buf); + if packet.next_header() != IpProtocol::Tcp { + return; + } ( - packet.next_header(), IpAddress::Ipv6(packet.src_addr()), IpAddress::Ipv6(packet.dst_addr()), packet.payload(), ) } }; - if protocol == IpProtocol::Tcp { - let tcp_packet = TcpPacket::new_unchecked(payload); - let src_addr = (src_addr, tcp_packet.src_port()).into(); - let dst_addr = (dst_addr, tcp_packet.dst_port()).into(); - let is_first = tcp_packet.syn() && !tcp_packet.ack(); - if is_first { - LISTEN_TABLE.incoming_tcp_packet(src_addr, dst_addr, sockets); - } + let tcp_packet = TcpPacket::new_unchecked(payload); + let src_addr = (src_addr, tcp_packet.src_port()).into(); + let dst_addr = (dst_addr, tcp_packet.dst_port()).into(); + let is_first = tcp_packet.syn() && !tcp_packet.ack(); + if is_first { + LISTEN_TABLE.incoming_tcp_packet(src_addr, dst_addr, sockets); } } diff --git a/net/ax-net/src/service.rs b/net/ax-net/src/service.rs index ad08ab5754..db237fdcb5 100644 --- a/net/ax-net/src/service.rs +++ b/net/ax-net/src/service.rs @@ -56,7 +56,7 @@ //! waker.wake(); // WRONG: potential self-deadlock //! ``` -use alloc::{boxed::Box, format, string::String, vec, vec::Vec}; +use alloc::{boxed::Box, format, string::String, sync::Arc, vec, vec::Vec}; use core::{ pin::Pin, task::{Context, Waker}, @@ -75,16 +75,18 @@ use smoltcp::{ Ipv4Packet, Ipv4Repr, UdpPacket, UdpRepr, }, }; +use spin::RwLock; use crate::{ SOCKET_SET, + addr::mask_from_prefix, config::{ DeviceBinding, DnsServerEntry, DnsSource, InterfaceFlags, InterfaceId, InterfaceInfo, InterfaceKind, Ipv4InterfaceConfig, RouteInfo, }, consts::STANDARD_MTU, device::{ArpEntry, EthernetDevice}, - dhcp_server::DhcpServer, + dhcp_server::{DhcpServer, parse_dhcp_packet}, router::{RouteDecision, Router, SharedRouteTable}, }; @@ -92,10 +94,6 @@ fn now() -> Instant { Instant::from_micros_const((monotonic_time_nanos() / NANOS_PER_MICROS) as i64) } -use alloc::sync::Arc; - -use spin::RwLock; - struct ControlState { interfaces: Vec, dns: Vec, @@ -174,7 +172,6 @@ impl NetControl { } pub fn default_routes(&self) -> Vec { - let _state = self.state.read(); self.routes.read().default_routes() } @@ -311,9 +308,16 @@ pub struct Service { pub iface: Interface, router: Router, control: Arc, - timeout: Option + Send>>>, + timeouts: Vec, dhcp: Vec, dhcp_server: Option, + dhcp_events: Vec, + dhcp_server_replies: Vec<(usize, Vec)>, +} + +struct TimeoutRegistration { + deadline: Instant, + _future: Pin + Send>>, } #[derive(Clone)] @@ -421,72 +425,53 @@ impl DhcpState { return None; } - let ipv4_packet = Ipv4Packet::new_checked(packet).ok()?; - let ipv4_repr = Ipv4Repr::parse(&ipv4_packet, &ChecksumCapabilities::default()).ok()?; - if ipv4_repr.next_header != IpProtocol::Udp { - return None; - } - - let udp_packet = UdpPacket::new_checked(ipv4_packet.payload()).ok()?; - let udp_repr = UdpRepr::parse( - &udp_packet, - &IpAddress::Ipv4(ipv4_repr.src_addr), - &IpAddress::Ipv4(ipv4_repr.dst_addr), - &ChecksumCapabilities::default(), - ) - .ok()?; - if udp_repr.src_port != DHCP_SERVER_PORT || udp_repr.dst_port != DHCP_CLIENT_PORT { + let parsed = parse_dhcp_packet(packet)?; + if parsed.udp.src_port != DHCP_SERVER_PORT || parsed.udp.dst_port != DHCP_CLIENT_PORT { return None; } - let dhcp_packet = DhcpPacket::new_checked(udp_packet.payload()).ok()?; - let dhcp_repr = DhcpRepr::parse(&dhcp_packet).ok()?; - if dhcp_repr.client_hardware_address != self.mac - || dhcp_repr.transaction_id != self.transaction_id + if parsed.client_hardware_address != self.mac + || parsed.transaction_id != self.transaction_id { return None; } - match (self.phase, dhcp_repr.message_type) { + match (self.phase, parsed.message_type) { (DhcpPhase::Discovering, DhcpMessageType::Offer) => { - if !is_unicast_ipv4(dhcp_repr.your_ip) { + if !is_unicast_ipv4(parsed.your_ip) { return None; } - self.offered_address = Some(dhcp_repr.your_ip); - self.server_identifier = dhcp_repr.server_identifier.or(Some(ipv4_repr.src_addr)); + self.offered_address = Some(parsed.your_ip); + self.server_identifier = parsed.server_identifier.or(Some(parsed.src_addr)); self.phase = DhcpPhase::Requesting; self.retry = 0; self.retry_at = timestamp; info!( "{}: DHCP offered address {} from {}", self.ifname, - dhcp_repr.your_ip, - self.server_identifier.unwrap_or(ipv4_repr.src_addr) + parsed.your_ip, + self.server_identifier.unwrap_or(parsed.src_addr) ); None } (DhcpPhase::Requesting, DhcpMessageType::Ack) | (DhcpPhase::Bound, DhcpMessageType::Ack) => { - let subnet_mask = dhcp_repr.subnet_mask?; + let subnet_mask = parsed.subnet_mask?; let prefix_len = IpAddress::Ipv4(subnet_mask).prefix_len()?; - if !is_unicast_ipv4(dhcp_repr.your_ip) { + if !is_unicast_ipv4(parsed.your_ip) { return None; } self.phase = DhcpPhase::Bound; self.retry = 0; - let address = Ipv4Cidr::new(dhcp_repr.your_ip, prefix_len); + let address = Ipv4Cidr::new(parsed.your_ip, prefix_len); Some(DhcpEvent::Configured { interface_id: self.interface_id, dev: self.dev, ifname: self.ifname.clone(), metric: self.metric, address, - router: dhcp_repr.router, - dns_servers: dhcp_repr - .dns_servers - .as_ref() - .map(|servers| servers.iter().copied().collect()) - .unwrap_or_default(), + router: parsed.router, + dns_servers: parsed.dns_servers, }) } (_, DhcpMessageType::Nak) => { @@ -556,9 +541,11 @@ impl Service { iface, router, control, - timeout: None, + timeouts: Vec::new(), dhcp: Vec::new(), dhcp_server: None, + dhcp_events: Vec::new(), + dhcp_server_replies: Vec::new(), } } @@ -736,8 +723,10 @@ impl Service { pub fn poll(&mut self, sockets: &mut SocketSet) -> bool { let timestamp = now(); - let mut dhcp_events = Vec::new(); - let mut dhcp_server_replies = Vec::new(); + let mut dhcp_events = core::mem::take(&mut self.dhcp_events); + let mut dhcp_server_replies = core::mem::take(&mut self.dhcp_server_replies); + dhcp_events.clear(); + dhcp_server_replies.clear(); let router_rx_pending; { @@ -758,23 +747,26 @@ impl Service { } }); } - for event in dhcp_events { + for event in dhcp_events.drain(..) { self.handle_dhcp_event(event); } let mut dhcp_server_sent = false; - for (dev, reply) in dhcp_server_replies { + for (dev, reply) in &dhcp_server_replies { dhcp_server_sent |= self.router.send_on_device( - dev, + *dev, IpAddress::Ipv4(Ipv4Address::BROADCAST), - &reply, + reply, timestamp, ); } + dhcp_server_replies.clear(); + self.dhcp_events = dhcp_events; + self.dhcp_server_replies = dhcp_server_replies; let socket_state_changed = self.iface.poll(timestamp, &mut self.router, sockets) == PollResult::SocketStateChanged; let dhcp_poll_next = self.poll_dhcp(timestamp); - // Reap orphaned TCP sockets using the SocketSet already held by poll_once(). + // Reap orphaned TCP sockets using the SocketSet already held by poll_until_idle(). crate::orphan::reap_orphans(timestamp, sockets); self.router.dispatch(timestamp, sockets) @@ -925,10 +917,6 @@ impl Service { self.router.arp_entries(now()) } - pub fn eth0_ipv4_config(&self) -> Option { - self.control.ipv4_config("eth0") - } - pub fn wake_all_devices(&self) { self.router.wake_all_devices(); } @@ -939,9 +927,6 @@ impl Service { if let Some(t) = next { let next = TimeValue::from_micros(t.total_micros() as _); - // drop old timeout future - self.timeout = None; - let mut fut = Box::pin(sleep_until(next)); let mut cx = Context::from_waker(waker); @@ -949,7 +934,12 @@ impl Service { waker.wake_by_ref(); return; } else { - self.timeout = Some(fut); + let now = now(); + self.timeouts.retain(|timeout| timeout.deadline > now); + self.timeouts.push(TimeoutRegistration { + deadline: t, + _future: fut, + }); } } @@ -991,17 +981,6 @@ fn is_unicast_ipv4(addr: Ipv4Address) -> bool { addr != Ipv4Address::UNSPECIFIED && addr != Ipv4Address::BROADCAST && !addr.is_multicast() } -/// 由前缀长度构造 IPv4 子网掩码(与 lib.rs 的 `prefix_to_mask` 等价, -/// 这里独立提供以免暴露跨模块的私有函数)。 -fn mask_from_prefix(prefix_len: u8) -> Ipv4Address { - let bits: u32 = if prefix_len == 0 { - 0 - } else { - u32::MAX << (32 - prefix_len.min(32) as u32) - }; - Ipv4Address::from_bits(bits) -} - fn build_dhcp_packet( mac: EthernetAddress, transaction_id: u32, diff --git a/net/ax-net/src/socket.rs b/net/ax-net/src/socket.rs index 3edb7f0b28..c9c80753d2 100644 --- a/net/ax-net/src/socket.rs +++ b/net/ax-net/src/socket.rs @@ -30,6 +30,7 @@ use ax_errno::{AxError, AxResult, LinuxError}; use ax_io::prelude::*; use axpoll::{IoEvents, Pollable}; use bitflags::bitflags; +use enum_dispatch::enum_dispatch; #[cfg(feature = "vsock")] use crate::vsock::{VsockAddr, VsockSocket}; @@ -194,6 +195,7 @@ impl Shutdown { } /// Operations that can be performed on a socket. +#[enum_dispatch] pub trait SocketOps: Configurable { /// Binds an unbound socket to the given address and port. fn bind(&self, local_addr: SocketAddrEx) -> AxResult; @@ -270,6 +272,7 @@ impl SocketOps for Box { } /// Network socket abstraction. +#[enum_dispatch(Configurable, SocketOps)] pub enum Socket { /// UDP socket. Udp(Box), @@ -309,142 +312,6 @@ impl From for Socket { } } -impl Configurable for Socket { - fn get_option_inner(&self, opt: &mut GetSocketOption) -> AxResult { - match self { - Socket::Tcp(tcp) => tcp.get_option_inner(opt), - Socket::Udp(udp) => udp.get_option_inner(opt), - Socket::Raw(raw) => raw.get_option_inner(opt), - Socket::Unix(unix) => unix.get_option_inner(opt), - #[cfg(feature = "vsock")] - Socket::Vsock(vsock) => vsock.get_option_inner(opt), - } - } - - fn set_option_inner(&self, opt: SetSocketOption) -> AxResult { - match self { - Socket::Tcp(tcp) => tcp.set_option_inner(opt), - Socket::Udp(udp) => udp.set_option_inner(opt), - Socket::Raw(raw) => raw.set_option_inner(opt), - Socket::Unix(unix) => unix.set_option_inner(opt), - #[cfg(feature = "vsock")] - Socket::Vsock(vsock) => vsock.set_option_inner(opt), - } - } -} - -impl SocketOps for Socket { - fn bind(&self, local_addr: SocketAddrEx) -> AxResult { - match self { - Socket::Tcp(tcp) => tcp.bind(local_addr), - Socket::Udp(udp) => udp.bind(local_addr), - Socket::Raw(raw) => raw.bind(local_addr), - Socket::Unix(unix) => unix.bind(local_addr), - #[cfg(feature = "vsock")] - Socket::Vsock(vsock) => vsock.bind(local_addr), - } - } - - fn connect(&self, remote_addr: SocketAddrEx) -> AxResult { - match self { - Socket::Tcp(tcp) => tcp.connect(remote_addr), - Socket::Udp(udp) => udp.connect(remote_addr), - Socket::Raw(raw) => raw.connect(remote_addr), - Socket::Unix(unix) => unix.connect(remote_addr), - #[cfg(feature = "vsock")] - Socket::Vsock(vsock) => vsock.connect(remote_addr), - } - } - - fn listen(&self, backlog: usize) -> AxResult { - match self { - Socket::Tcp(tcp) => tcp.listen(backlog), - Socket::Udp(udp) => udp.listen(backlog), - Socket::Raw(raw) => raw.listen(backlog), - Socket::Unix(unix) => unix.listen(backlog), - #[cfg(feature = "vsock")] - Socket::Vsock(vsock) => vsock.listen(backlog), - } - } - - fn accept(&self) -> AxResult { - match self { - Socket::Tcp(tcp) => tcp.accept(), - Socket::Udp(udp) => udp.accept(), - Socket::Raw(raw) => raw.accept(), - Socket::Unix(unix) => unix.accept(), - #[cfg(feature = "vsock")] - Socket::Vsock(vsock) => vsock.accept(), - } - } - - fn send(&self, src: impl Read + IoBuf, options: SendOptions) -> AxResult { - match self { - Socket::Tcp(tcp) => tcp.send(src, options), - Socket::Udp(udp) => udp.send(src, options), - Socket::Raw(raw) => raw.send(src, options), - Socket::Unix(unix) => unix.send(src, options), - #[cfg(feature = "vsock")] - Socket::Vsock(vsock) => vsock.send(src, options), - } - } - - fn recv(&self, dst: impl Write + IoBufMut, options: RecvOptions<'_>) -> AxResult { - match self { - Socket::Tcp(tcp) => tcp.recv(dst, options), - Socket::Udp(udp) => udp.recv(dst, options), - Socket::Raw(raw) => raw.recv(dst, options), - Socket::Unix(unix) => unix.recv(dst, options), - #[cfg(feature = "vsock")] - Socket::Vsock(vsock) => vsock.recv(dst, options), - } - } - - fn recv_available(&self) -> AxResult { - match self { - Socket::Tcp(tcp) => tcp.recv_available(), - Socket::Udp(udp) => udp.recv_available(), - Socket::Raw(raw) => raw.recv_available(), - Socket::Unix(unix) => unix.recv_available(), - #[cfg(feature = "vsock")] - Socket::Vsock(vsock) => vsock.recv_available(), - } - } - - fn local_addr(&self) -> AxResult { - match self { - Socket::Tcp(tcp) => tcp.local_addr(), - Socket::Udp(udp) => udp.local_addr(), - Socket::Raw(raw) => raw.local_addr(), - Socket::Unix(unix) => unix.local_addr(), - #[cfg(feature = "vsock")] - Socket::Vsock(vsock) => vsock.local_addr(), - } - } - - fn peer_addr(&self) -> AxResult { - match self { - Socket::Tcp(tcp) => tcp.peer_addr(), - Socket::Udp(udp) => udp.peer_addr(), - Socket::Raw(raw) => raw.peer_addr(), - Socket::Unix(unix) => unix.peer_addr(), - #[cfg(feature = "vsock")] - Socket::Vsock(vsock) => vsock.peer_addr(), - } - } - - fn shutdown(&self, how: Shutdown) -> AxResult { - match self { - Socket::Tcp(tcp) => tcp.shutdown(how), - Socket::Udp(udp) => udp.shutdown(how), - Socket::Raw(raw) => raw.shutdown(how), - Socket::Unix(unix) => unix.shutdown(how), - #[cfg(feature = "vsock")] - Socket::Vsock(vsock) => vsock.shutdown(how), - } - } -} - impl Pollable for Socket { fn poll(&self) -> IoEvents { match self { diff --git a/net/ax-net/src/tcp.rs b/net/ax-net/src/tcp.rs index 4be93433ca..1c4eca76da 100644 --- a/net/ax-net/src/tcp.rs +++ b/net/ax-net/src/tcp.rs @@ -25,7 +25,7 @@ //! - `LISTEN_TABLE` owns passive-open child sockets and accept wakeups. //! - `orphan` keeps dropped sockets alive long enough for FIN/TIME-WAIT cleanup. -use alloc::{sync::Arc, task::Wake, vec, vec::Vec}; +use alloc::{sync::Arc, vec}; use core::{ net::{Ipv4Addr, SocketAddr}, sync::atomic::{AtomicBool, AtomicI32, AtomicU32, Ordering}, @@ -36,7 +36,7 @@ use ax_errno::{AxError, AxResult, LinuxError, ax_bail, ax_err_type}; use ax_io::prelude::*; use ax_sync::Mutex; use axpoll::{IoEvents, PollSet, Pollable}; -use hashbrown::HashMap; +use hashbrown::{HashMap, HashSet}; use smoltcp::{ iface::SocketHandle, socket::tcp as smol, @@ -46,11 +46,11 @@ use smoltcp::{ use spin::LazyLock; use crate::{ - LISTEN_TABLE, RecvFlags, RecvOptions, SOCKET_SET, SendOptions, Shutdown, Socket, SocketAddrEx, - SocketOps, + DeferPollWake, LISTEN_TABLE, RecvFlags, RecvOptions, SOCKET_SET, SendOptions, Shutdown, Socket, + SocketAddrEx, SocketOps, + addr::{allocate_ephemeral_port, listen_addrs_conflict}, config::{DeviceBinding, InterfaceId}, consts::{TCP_RX_BUF_LEN, TCP_TX_BUF_LEN}, - endpoint_from_ip_endpoint, general::GeneralOptions, get_control, get_service, interface_by_id, options::{Configurable, GetSocketOption, SetSocketOption, TcpInfo, TcpInfoOptions, TcpState}, @@ -58,14 +58,6 @@ use crate::{ state::*, }; -/// Allocates a smoltcp TCP socket with ax-net's default buffers. -pub(crate) fn new_tcp_socket() -> smol::Socket<'static> { - smol::Socket::new( - smol::SocketBuffer::new(vec![0; TCP_RX_BUF_LEN]), - smol::SocketBuffer::new(vec![0; TCP_TX_BUF_LEN]), - ) -} - const TCP_KEEPIDLE_DEFAULT_SECS: u32 = 7200; const TCP_KEEPINTVL_DEFAULT_SECS: u32 = 75; const TCP_KEEPCNT_DEFAULT: u32 = 9; @@ -115,30 +107,15 @@ pub struct TcpSocket { unsafe impl Sync for TcpSocket {} -struct TcpPollWake { - poll: Arc, - ready: IoEvents, -} - -impl Wake for TcpPollWake { - fn wake(self: Arc) { - self.wake_by_ref(); - } - - fn wake_by_ref(self: &Arc) { - // smoltcp invokes socket wakers from the net poll task context after - // updating socket readiness. The socket set may still be locked there, - // so defer the actual PollSet wake to the net worker outer loop. - crate::defer_poll_wake(self.poll.clone(), self.ready); - } -} - impl TcpSocket { /// Creates a new TCP socket. pub fn new() -> Self { Self { state: StateLock::new(State::Idle), - handle: SOCKET_SET.add(new_tcp_socket()), + handle: SOCKET_SET.add(smol::Socket::new( + smol::SocketBuffer::new(vec![0; TCP_RX_BUF_LEN]), + smol::SocketBuffer::new(vec![0; TCP_TX_BUF_LEN]), + )), bound_endpoint: Mutex::new(empty_endpoint()), peer_endpoint: Mutex::new(None), bound_registered: AtomicBool::new(false), @@ -191,7 +168,10 @@ impl TcpSocket { poll_tx: Arc::new(PollSet::new()), poll_rx_closed: PollSet::new(), }; - let endpoint = endpoint_from_ip_endpoint(local_endpoint); + let endpoint = IpListenEndpoint { + addr: Some(local_endpoint.addr), + port: local_endpoint.port, + }; *result.bound_endpoint.lock() = endpoint; result.general.set_device_binding( get_control() @@ -210,26 +190,13 @@ impl Default for TcpSocket { /// Private methods impl TcpSocket { - fn state(&self) -> State { - self.state.get() - } - - #[inline] - fn is_listening(&self) -> bool { - self.state() == State::Listening - } - fn with_smol_socket(&self, f: impl FnOnce(&mut smol::Socket) -> R) -> R { SOCKET_SET.with_socket_mut::(self.handle, f) } - fn keep_alive_interval(&self) -> Duration { - Duration::from_secs(self.keep_idle_secs.load(Ordering::Relaxed) as u64) - } - fn tcp_info_snapshot(&self) -> TcpInfo { self.with_smol_socket(|socket| { - let send_queue = saturating_u32(socket.send_queue()); + let send_queue = socket.send_queue().min(u32::MAX as usize) as u32; let snd_mss = TCP_INFO_DEFAULT_MSS; let mut options = TcpInfoOptions::empty(); @@ -265,23 +232,6 @@ impl TcpSocket { Ok(endpoint) } - fn store_pending_error(&self, err: LinuxError) { - self.pending_error.store(err.code(), Ordering::Release); - } - - fn clear_pending_error(&self) { - self.pending_error.store(0, Ordering::Release); - } - - fn take_pending_error(&self) -> i32 { - self.pending_error.swap(0, Ordering::AcqRel) - } - - fn connect_error(&self) -> AxError { - LinuxError::try_from(self.pending_error.load(Ordering::Acquire)) - .map_or(AxError::ConnectionRefused, AxError::from) - } - fn poll_connect(&self) -> IoEvents { let mut events = IoEvents::empty(); self.with_smol_socket(|socket| match socket.state() { @@ -289,7 +239,7 @@ impl TcpSocket { // wait for connection } smol::State::Established => { - self.clear_pending_error(); + self.pending_error.store(0, Ordering::Release); self.state.set(State::Connected); // connected *self.peer_endpoint.lock() = socket.remote_endpoint(); debug!( @@ -301,7 +251,8 @@ impl TcpSocket { } state => { *self.peer_endpoint.lock() = None; - self.store_pending_error(LinuxError::ECONNREFUSED); + self.pending_error + .store(LinuxError::ECONNREFUSED.code(), Ordering::Release); self.state.set(State::Closed); // connection failed debug!( "TCP socket {}: connect failed in state {:?}", @@ -345,7 +296,7 @@ impl Configurable for TcpSocket { use GetSocketOption as O; if let O::Error(error) = option { - **error = self.take_pending_error(); + **error = self.pending_error.swap(0, Ordering::AcqRel); return Ok(true); } @@ -404,7 +355,8 @@ impl Configurable for TcpSocket { }); } O::KeepAlive(keep_alive) => { - let interval = self.keep_alive_interval(); + let interval = + Duration::from_secs(self.keep_idle_secs.load(Ordering::Relaxed) as u64); self.with_smol_socket(|socket| { socket.set_keep_alive(keep_alive.then_some(interval)); }); @@ -490,10 +442,13 @@ impl SocketOps for TcpSocket { let events = self.poll_connect(); if !events.contains(IoEvents::OUT) { Err(AxError::WouldBlock) - } else if self.state() == State::Connected { + } else if self.state.get() == State::Connected { Ok(()) } else { - Err(self.connect_error()) + Err( + LinuxError::try_from(self.pending_error.load(Ordering::Acquire)) + .map_or(AxError::ConnectionRefused, AxError::from), + ) } }) } @@ -506,20 +461,10 @@ impl SocketOps for TcpSocket { bound_endpoint.port = get_ephemeral_port()?; } let binding = get_control().local_binding_for(&bound_endpoint)?; - let register_bound = !self.bound_registered.load(Ordering::Acquire); - if register_bound { - register_tcp_bound(bound_endpoint)?; - } - if let Err(err) = LISTEN_TABLE.listen(bound_endpoint, backlog) { - if register_bound { - unregister_tcp_bound(bound_endpoint); - } - return Err(err); - } + self.with_bound_endpoint_registered(bound_endpoint, || { + LISTEN_TABLE.listen(bound_endpoint, backlog) + })?; *self.bound_endpoint.lock() = bound_endpoint; - if register_bound { - self.bound_registered.store(true, Ordering::Release); - } if binding.bound_if.is_some() { self.general.set_device_binding(binding); } @@ -533,7 +478,7 @@ impl SocketOps for TcpSocket { } fn accept(&self) -> AxResult { - if !self.is_listening() { + if self.state.get() != State::Listening { ax_bail!(InvalidInput, "not listening"); } @@ -592,7 +537,7 @@ impl SocketOps for TcpSocket { if self.rx_closed.load(Ordering::Acquire) { return Err(AxError::NotConnected); } - if self.state() == State::Closed { + if self.state.get() == State::Closed { return Err(AxError::NotConnected); } let extra_nb = options.flags.contains(RecvFlags::DONTWAIT); @@ -637,7 +582,7 @@ impl SocketOps for TcpSocket { } fn recv_available(&self) -> AxResult { - if self.is_listening() { + if self.state.get() == State::Listening { return Err(AxError::InvalidInput); } let available = self.with_smol_socket(|socket| socket.recv_queue()); @@ -652,7 +597,10 @@ impl SocketOps for TcpSocket { let endpoint = self.with_smol_socket(|socket| { socket .local_endpoint() - .map(endpoint_from_ip_endpoint) + .map(|endpoint| IpListenEndpoint { + addr: Some(endpoint.addr), + port: endpoint.port, + }) .unwrap_or_else(|| *self.bound_endpoint.lock()) }); Ok(SocketAddrEx::Ip(SocketAddr::new( @@ -724,7 +672,7 @@ impl SocketOps for TcpSocket { impl Pollable for TcpSocket { fn poll(&self) -> IoEvents { request_poll(); - let mut events = match self.state() { + let mut events = match self.state.get() { State::Connecting => self.poll_connect(), State::Connected | State::Idle | State::Closed => self.poll_stream(), State::Listening => self.poll_listener(), @@ -736,7 +684,8 @@ impl Pollable for TcpSocket { fn register(&self, context: &mut Context<'_>, events: IoEvents) { let mut accept_registration = None; - if self.is_listening() && events.intersects(IoEvents::IN | IoEvents::RDHUP) { + if self.state.get() == State::Listening && events.intersects(IoEvents::IN | IoEvents::RDHUP) + { let port = self.bound_endpoint.lock().port; if port != 0 { let endpoint = *self.bound_endpoint.lock(); @@ -756,7 +705,7 @@ impl Pollable for TcpSocket { self.poll_rx .register(context.waker(), IoEvents::IN | IoEvents::RDHUP) }; - Some(Waker::from(Arc::new(TcpPollWake { + Some(Waker::from(Arc::new(DeferPollWake { poll: self.poll_rx.clone(), ready: IoEvents::IN | IoEvents::RDHUP, }))) @@ -767,7 +716,7 @@ impl Pollable for TcpSocket { // Socket registration runs from task poll context before taking the // socket-set lock. unsafe { self.poll_tx.register(context.waker(), IoEvents::OUT) }; - Some(Waker::from(Arc::new(TcpPollWake { + Some(Waker::from(Arc::new(DeferPollWake { poll: self.poll_tx.clone(), ready: IoEvents::OUT, }))) @@ -806,9 +755,15 @@ impl Pollable for TcpSocket { impl Drop for TcpSocket { fn drop(&mut self) { + let endpoint = *self.bound_endpoint.lock(); + if self.state.get() == State::Listening && endpoint.port != 0 { + LISTEN_TABLE.unlisten(endpoint); + } + let should_orphan = self.with_smol_socket(|socket| { - matches!( - socket.state(), + let state = socket.state(); + let should_orphan = matches!( + state, smol::State::Established | smol::State::CloseWait | smol::State::FinWait1 @@ -816,14 +771,24 @@ impl Drop for TcpSocket { | smol::State::Closing | smol::State::LastAck | smol::State::TimeWait - ) || socket.send_queue() > 0 + ) || socket.send_queue() > 0; + if matches!( + state, + smol::State::Established + | smol::State::SynSent + | smol::State::SynReceived + | smol::State::CloseWait + | smol::State::FinWait1 + | smol::State::FinWait2 + | smol::State::Closing + | smol::State::LastAck + ) { + debug!("TCP socket {}: closing on drop", self.handle); + socket.close(); + } + should_orphan }); - // Initiate graceful shutdown (send FIN if connected) - if let Err(err) = self.shutdown(Shutdown::Both) { - warn!("TCP socket {}: shutdown failed: {}", self.handle, err); - } - // Unbind from API layer (port registry, etc.) self.unregister_bound_endpoint(); @@ -842,10 +807,6 @@ impl Drop for TcpSocket { } } -fn saturating_u32(value: usize) -> u32 { - value.min(u32::MAX as usize) as u32 -} - fn duration_micros_u32(value: Duration) -> u32 { value.total_micros().min(u32::MAX as u64) as u32 } @@ -887,7 +848,7 @@ impl TcpSocket { } })? .transit(State::Connecting, || { - self.clear_pending_error(); + self.pending_error.store(0, Ordering::Release); // TODO: check remote addr unreachable // let (bound_endpoint, remote_endpoint) = self.get_endpoint_pair(remote_addr)?; let remote_endpoint = IpEndpoint::from(remote_addr); @@ -916,12 +877,7 @@ impl TcpSocket { "TCP connection from {} to {}", bound_endpoint, remote_endpoint ); - let register_bound = !self.bound_registered.load(Ordering::Acquire); - if register_bound { - register_tcp_bound(bound_endpoint)?; - } - - let result = { + self.with_bound_endpoint_registered(bound_endpoint, || { let mut service = get_service(); let context = service.iface.context(); self.with_smol_socket(|socket| { @@ -937,17 +893,8 @@ impl TcpSocket { })?; Ok::<(), AxError>(()) }) - }; - if let Err(err) = result { - if register_bound { - unregister_tcp_bound(bound_endpoint); - } - return Err(err); - } + })?; *self.bound_endpoint.lock() = bound_endpoint; - if register_bound { - self.bound_registered.store(true, Ordering::Release); - } // Only set device binding if was originally unbound or bound to 0.0.0.0 // Binding to a specific IP should lock the interface @@ -970,6 +917,31 @@ impl TcpSocket { Ok(()) } + fn with_bound_endpoint_registered( + &self, + endpoint: IpListenEndpoint, + f: impl FnOnce() -> AxResult, + ) -> AxResult { + let register_bound = !self.bound_registered.load(Ordering::Acquire); + if register_bound { + register_tcp_bound(endpoint)?; + } + match f() { + Ok(value) => { + if register_bound { + self.bound_registered.store(true, Ordering::Release); + } + Ok(value) + } + Err(err) => { + if register_bound { + unregister_tcp_bound(endpoint); + } + Err(err) + } + } + } + /// Removes the public TCP bind side-table entry, if present. fn unregister_bound_endpoint(&self) { if self.bound_registered.swap(false, Ordering::AcqRel) { @@ -978,7 +950,7 @@ impl TcpSocket { } } -static TCP_BOUND_PORTS: LazyLock>>>> = +static TCP_BOUND_PORTS: LazyLock>>>> = LazyLock::new(|| Mutex::new(HashMap::new())); /// Registers TCP bind ownership with wildcard/specific address conflicts. @@ -995,7 +967,7 @@ fn register_tcp_bound(endpoint: IpListenEndpoint) -> AxResult { { return Err(AxError::AddrInUse); } - bound_addrs.push(endpoint.addr); + bound_addrs.insert(endpoint.addr); Ok(()) } @@ -1004,9 +976,7 @@ fn unregister_tcp_bound(endpoint: IpListenEndpoint) { if endpoint.port != 0 { let mut bound_ports = TCP_BOUND_PORTS.lock(); if let Some(bound_addrs) = bound_ports.get_mut(&endpoint.port) { - if let Some(idx) = bound_addrs.iter().position(|&addr| addr == endpoint.addr) { - bound_addrs.swap_remove(idx); - } + bound_addrs.remove(&endpoint.addr); if bound_addrs.is_empty() { bound_ports.remove(&endpoint.port); } @@ -1022,35 +992,8 @@ fn tcp_port_available(port: u16) -> bool { && !TCP_BOUND_PORTS.lock().contains_key(&port) } -/// Returns whether two listen/bind addresses conflict on the same port. -fn listen_addrs_conflict( - a: Option, - b: Option, -) -> bool { - a.is_none() || b.is_none() || a == b -} - fn get_ephemeral_port() -> AxResult { - const PORT_START: u16 = 0xc000; - const PORT_END: u16 = 0xffff; - static CURR: Mutex = Mutex::new(PORT_START); - - let mut curr = CURR.lock(); - let mut tries = 0; - // TODO: more robust - while tries <= PORT_END - PORT_START { - let port = *curr; - if *curr == PORT_END { - *curr = PORT_START; - } else { - *curr += 1; - } - if tcp_port_available(port) { - return Ok(port); - } - tries += 1; - } - ax_bail!(AddrInUse, "no available ports"); + allocate_ephemeral_port(tcp_port_available) } #[cfg(test)] diff --git a/net/ax-net/src/udp.rs b/net/ax-net/src/udp.rs index 6d88d392f1..21ff816c54 100644 --- a/net/ax-net/src/udp.rs +++ b/net/ax-net/src/udp.rs @@ -45,6 +45,7 @@ use spin::RwLock; use crate::{ RecvFlags, RecvOptions, SOCKET_SET, SendFlags, SendOptions, Shutdown, SocketAddrEx, SocketOps, + addr::allocate_ephemeral_port, config::{DeviceBinding, InterfaceId}, consts::{UDP_RX_BUF_LEN, UDP_TX_BUF_LEN}, general::GeneralOptions, @@ -63,15 +64,6 @@ struct CorkState { source: IpAddress, } -/// Allocates a smoltcp UDP socket with ax-net's default packet buffers. -pub(crate) fn new_udp_socket() -> smol::Socket<'static> { - // TODO(mivik): buffer size - smol::Socket::new( - smol::PacketBuffer::new(vec![PacketMetadata::EMPTY; 256], vec![0; UDP_RX_BUF_LEN]), - smol::PacketBuffer::new(vec![PacketMetadata::EMPTY; 256], vec![0; UDP_TX_BUF_LEN]), - ) -} - /// A UDP socket that provides POSIX-like APIs. pub struct UdpSocket { /// Handle into the global smoltcp socket set. @@ -90,13 +82,12 @@ pub struct UdpSocket { impl UdpSocket { /// Creates a new UDP socket. - #[allow(clippy::new_without_default)] pub fn new() -> Self { - let socket = new_udp_socket(); - let handle = SOCKET_SET.add(socket); - Self { - handle, + handle: SOCKET_SET.add(smol::Socket::new( + smol::PacketBuffer::new(vec![PacketMetadata::EMPTY; 256], vec![0; UDP_RX_BUF_LEN]), + smol::PacketBuffer::new(vec![PacketMetadata::EMPTY; 256], vec![0; UDP_TX_BUF_LEN]), + )), local_addr: RwLock::new(None), peer_addr: RwLock::new(None), @@ -135,6 +126,29 @@ impl UdpSocket { .select_route_with_binding(remote, self.general.device_binding())? .source) } + + fn send_source_for_remote(&self, remote: &IpAddress) -> AxResult { + if let Some(local_ep) = *self.local_addr.read() + && !local_ep.addr.is_unspecified() + { + Ok(local_ep.addr) + } else { + self.source_for_remote(remote) + } + } + + fn source_and_binding_update_for_remote( + &self, + remote: &IpAddress, + ) -> AxResult<(IpAddress, bool)> { + if let Some(local_ep) = *self.local_addr.read() + && !local_ep.addr.is_unspecified() + { + Ok((local_ep.addr, false)) + } else { + Ok((self.source_for_remote(remote)?, true)) + } + } } impl Configurable for UdpSocket { @@ -233,20 +247,9 @@ impl SocketOps for UdpSocket { } let remote_addr = IpEndpoint::from(remote_addr); - let local = self.local_addr.read(); - - // Determine source address and device binding based on bind state - let (src, should_update_binding) = if let Some(local_ep) = *local { - if local_ep.addr.is_unspecified() { - // Bound to 0.0.0.0, use route decision - (self.source_for_remote(&remote_addr.addr)?, true) - } else { - // Bound to specific IP, use that address and keep interface - (local_ep.addr, false) - } - } else { - (self.source_for_remote(&remote_addr.addr)?, true) - }; + let local_port = self.local_addr.read().map_or(0, |endpoint| endpoint.port); + let (src, should_update_binding) = + self.source_and_binding_update_for_remote(&remote_addr.addr)?; *guard = Some((remote_addr, src)); @@ -254,7 +257,7 @@ impl SocketOps for UdpSocket { self.general .set_device_binding(get_control().local_binding_for(&IpListenEndpoint { addr: Some(src), - port: (*local).map_or(0, |endpoint| endpoint.port), + port: local_port, })?); } @@ -285,16 +288,7 @@ impl SocketOps for UdpSocket { let (remote_addr, source_addr) = match options.to { Some(addr) => { let addr = IpEndpoint::from(addr.into_ip()?); - // Use bound address if bound to specific IP - let src = if let Some(local_ep) = *self.local_addr.read() { - if local_ep.addr.is_unspecified() { - self.source_for_remote(&addr.addr)? - } else { - local_ep.addr - } - } else { - self.source_for_remote(&addr.addr)? - }; + let src = self.send_source_for_remote(&addr.addr)?; (addr, src) } None => match self.remote_endpoint() { @@ -335,16 +329,7 @@ impl SocketOps for UdpSocket { let resolved = match options.to { Some(addr) => { let addr = IpEndpoint::from(addr.into_ip()?); - // Use bound address if bound to specific IP, otherwise route decision - let src = if let Some(local_ep) = *self.local_addr.read() { - if local_ep.addr.is_unspecified() { - self.source_for_remote(&addr.addr)? - } else { - local_ep.addr - } - } else { - self.source_for_remote(&addr.addr)? - }; + let src = self.send_source_for_remote(&addr.addr)?; Some((addr, src)) } None => self.remote_endpoint().ok(), @@ -571,6 +556,12 @@ impl Pollable for UdpSocket { } } +impl Default for UdpSocket { + fn default() -> Self { + Self::new() + } +} + impl Drop for UdpSocket { fn drop(&mut self) { self.shutdown(Shutdown::Both).ok(); @@ -579,25 +570,9 @@ impl Drop for UdpSocket { } fn get_ephemeral_port() -> AxResult { - const PORT_START: u16 = 0xc000; - const PORT_END: u16 = 0xffff; - static CURR: Mutex = Mutex::new(PORT_START); - let mut curr = CURR.lock(); - - let mut tries = 0; - while tries <= PORT_END - PORT_START { - let port = *curr; - if *curr == PORT_END { - *curr = PORT_START; - } else { - *curr += 1; - } - if SOCKET_SET.udp_port_available(IpAddress::Ipv4(Ipv4Addr::UNSPECIFIED), port) { - return Ok(port); - } - tries += 1; - } - ax_bail!(AddrInUse, "no available ports") + allocate_ephemeral_port(|port| { + SOCKET_SET.udp_port_available(IpAddress::Ipv4(Ipv4Addr::UNSPECIFIED), port) + }) } #[cfg(test)] diff --git a/net/ax-net/src/vsock/connection_manager.rs b/net/ax-net/src/vsock/connection_manager.rs index 945f953a9f..0ccb3cef87 100644 --- a/net/ax-net/src/vsock/connection_manager.rs +++ b/net/ax-net/src/vsock/connection_manager.rs @@ -548,40 +548,7 @@ impl VsockConnectionManager { } Ok(()) } - - /// statistics - #[allow(dead_code)] - pub fn get_stats(&self) -> VsockStats { - VsockStats { - total_connections: self.connections.len(), - listening_ports: self.listen_queues.len(), - total_rx_bytes: self.connections.values().map(|c| c.lock().rx_bytes).sum(), - total_tx_bytes: self.connections.values().map(|c| c.lock().tx_bytes).sum(), - total_dropped_bytes: self - .connections - .values() - .map(|c| c.lock().dropped_bytes) - .sum(), - } - } -} - -/// Vsock statistics -#[allow(dead_code)] -#[derive(Debug, Clone)] -pub struct VsockStats { - pub total_connections: usize, - pub listening_ports: usize, - pub total_rx_bytes: usize, - pub total_tx_bytes: usize, - pub total_dropped_bytes: usize, } pub static VSOCK_CONN_MANAGER: Mutex = Mutex::new(VsockConnectionManager::new()); - -/// for debug -#[allow(dead_code)] -pub fn get_vsock_stats() -> VsockStats { - VSOCK_CONN_MANAGER.lock().get_stats() -} diff --git a/net/ax-net/src/vsock/mod.rs b/net/ax-net/src/vsock/mod.rs index 6241ce9469..f593771bcd 100644 --- a/net/ax-net/src/vsock/mod.rs +++ b/net/ax-net/src/vsock/mod.rs @@ -1,8 +1,6 @@ //! Vsock socket facade. //! -//! This module exposes vsock transports through the common socket API. Stream -//! transport is implemented today; the transport enum leaves room for future -//! datagram support without changing the public socket wrapper. +//! This module exposes stream-oriented vsock through the common socket API. //! //! # Stack Boundary //! @@ -11,8 +9,6 @@ //! sockets, but actual connection state lives in `connection_manager` and the //! device event loop in `device::vsock`. -// pub(crate) mod dgram; todo - pub(crate) mod connection_manager; pub(crate) mod stream; @@ -21,7 +17,6 @@ use core::task::Context; use ax_errno::{AxError, AxResult}; use ax_io::{IoBuf, IoBufMut, Read, Write}; use axpoll::{IoEvents, Pollable}; -use enum_dispatch::enum_dispatch; pub use rdif_vsock::{VsockAddr, VsockConnId}; pub use self::stream::VsockStreamTransport; @@ -30,65 +25,28 @@ use crate::{ options::{Configurable, GetSocketOption, SetSocketOption}, }; -/// Abstract transport trait for vsock. -#[enum_dispatch] -pub trait VsockTransportOps: Configurable + Pollable + Send + Sync { - /// Bind the transport to a local address. - fn bind(&self, local_addr: VsockAddr) -> AxResult; - /// Start listening for incoming connections. - fn listen(&self) -> AxResult; - /// Connect to a remote peer address. - fn connect(&self, peer_addr: VsockAddr) -> AxResult; - /// Accept an incoming connection. - fn accept(&self) -> AxResult<(VsockTransport, VsockAddr)>; - /// Send data through the transport. - fn send(&self, src: impl Read + IoBuf, options: SendOptions) -> AxResult; - /// Receive data from the transport. - fn recv(&self, dst: impl Write, options: RecvOptions<'_>) -> AxResult; - /// Shutdown the transport. - fn shutdown(&self, _how: Shutdown) -> AxResult; - /// Get the local address, if bound. - fn local_addr(&self) -> AxResult>; - /// Get the peer address, if connected. - fn peer_addr(&self) -> AxResult>; -} - -/// Vsock transport type. -#[enum_dispatch(Configurable, VsockTransportOps)] -pub enum VsockTransport { +/// A network socket using the vsock protocol. +pub struct VsockSocket { /// Stream-oriented vsock transport. - Stream(VsockStreamTransport), - // Dgram(VsockDgramVsockTransport), + transport: VsockStreamTransport, } -impl Pollable for VsockTransport { - fn poll(&self) -> IoEvents { - match self { - VsockTransport::Stream(stream) => stream.poll(), - // VsockTransport::Dgram(dgram) => dgram.poll(), +impl VsockSocket { + /// Create a new stream-oriented vsock socket. + pub fn new() -> Self { + Self { + transport: VsockStreamTransport::new(), } } - fn register(&self, context: &mut core::task::Context<'_>, events: IoEvents) { - match self { - VsockTransport::Stream(stream) => stream.register(context, events), - // VsockTransport::Dgram(dgram) => dgram.register(context, events), - } + fn from_transport(transport: VsockStreamTransport) -> Self { + Self { transport } } } -/// A network socket using the vsock protocol. -pub struct VsockSocket { - /// Concrete vsock transport. - transport: VsockTransport, -} - -impl VsockSocket { - /// Create a new vsock socket with the given transport. - pub fn new(transport: impl Into) -> Self { - Self { - transport: transport.into(), - } +impl Default for VsockSocket { + fn default() -> Self { + Self::new() } } @@ -119,7 +77,7 @@ impl SocketOps for VsockSocket { fn accept(&self) -> AxResult { self.transport.accept().map(|(transport, _addr)| { - let socket = VsockSocket::new(transport); + let socket = VsockSocket::from_transport(transport); socket.into() }) } diff --git a/net/ax-net/src/vsock/stream.rs b/net/ax-net/src/vsock/stream.rs index 2ab6b15881..43634f91f2 100644 --- a/net/ax-net/src/vsock/stream.rs +++ b/net/ax-net/src/vsock/stream.rs @@ -32,7 +32,7 @@ use crate::{ general::GeneralOptions, options::{Configurable, GetSocketOption, SetSocketOption}, state::*, - vsock::{VsockAddr, VsockConnId, VsockTransport, VsockTransportOps}, + vsock::{VsockAddr, VsockConnId}, }; /// Stream transport for vsock sockets. @@ -80,8 +80,8 @@ impl Configurable for VsockStreamTransport { } } -impl VsockTransportOps for VsockStreamTransport { - fn bind(&self, mut local_addr: VsockAddr) -> AxResult<()> { +impl VsockStreamTransport { + pub(super) fn bind(&self, mut local_addr: VsockAddr) -> AxResult<()> { self.state .lock(State::Idle) .map_err(|_| ax_err_type!(InvalidInput, "already bound"))? @@ -102,7 +102,7 @@ impl VsockTransportOps for VsockStreamTransport { Ok(()) } - fn listen(&self) -> AxResult<()> { + pub(super) fn listen(&self) -> AxResult<()> { let guard = self .state .lock(State::Idle) @@ -122,7 +122,7 @@ impl VsockTransportOps for VsockStreamTransport { }) } - fn accept(&self) -> AxResult<(VsockTransport, VsockAddr)> { + pub(super) fn accept(&self) -> AxResult<(VsockStreamTransport, VsockAddr)> { if self.state.get() != State::Listening { ax_bail!(InvalidInput, "not listening"); } @@ -149,11 +149,11 @@ impl VsockTransportOps for VsockStreamTransport { general: GeneralOptions::new(1, 40, 0), // SOCK_STREAM }; - Ok((VsockTransport::Stream(new_transport), peer_addr)) + Ok((new_transport, peer_addr)) }) } - fn connect(&self, peer_addr: VsockAddr) -> AxResult<()> { + pub(super) fn connect(&self, peer_addr: VsockAddr) -> AxResult<()> { let guard = self.state.lock(State::Idle).map_err(|state| match state { State::Idle => unreachable!(), State::Listening => ax_err_type!(InvalidInput, "already listening"), @@ -224,7 +224,11 @@ impl VsockTransportOps for VsockStreamTransport { }) } - fn send(&self, mut src: impl Read + IoBuf, _options: SendOptions) -> AxResult { + pub(super) fn send( + &self, + mut src: impl Read + IoBuf, + _options: SendOptions, + ) -> AxResult { let conn = self.get_connection()?; let conn_guard = conn.lock(); @@ -245,7 +249,7 @@ impl VsockTransportOps for VsockStreamTransport { result } - fn recv(&self, mut dst: impl Write, options: RecvOptions) -> AxResult { + pub(super) fn recv(&self, mut dst: impl Write, options: RecvOptions) -> AxResult { let conn = self.get_connection()?; let extra_nb = options.flags.contains(RecvFlags::DONTWAIT); @@ -292,7 +296,7 @@ impl VsockTransportOps for VsockStreamTransport { }) } - fn shutdown(&self, how: Shutdown) -> AxResult<()> { + pub(super) fn shutdown(&self, how: Shutdown) -> AxResult<()> { let conn = self.get_connection()?; let mut conn = conn.lock(); @@ -315,14 +319,14 @@ impl VsockTransportOps for VsockStreamTransport { Ok(()) } - fn local_addr(&self) -> AxResult> { + pub(super) fn local_addr(&self) -> AxResult> { Ok(self .get_connection() .ok() .map(|conn| conn.lock().local_addr())) } - fn peer_addr(&self) -> AxResult> { + pub(super) fn peer_addr(&self) -> AxResult> { Ok(self .get_connection() .ok() diff --git a/os/StarryOS/kernel/src/syscall/net/socket.rs b/os/StarryOS/kernel/src/syscall/net/socket.rs index d198482974..87e6ad52e5 100644 --- a/os/StarryOS/kernel/src/syscall/net/socket.rs +++ b/os/StarryOS/kernel/src/syscall/net/socket.rs @@ -3,7 +3,7 @@ use alloc::boxed::Box; use ax_errno::{AxError, AxResult, LinuxError}; use ax_fs_ng::vfs::FS_CONTEXT; #[cfg(feature = "vsock")] -use ax_net::vsock::{VsockSocket, VsockStreamTransport}; +use ax_net::vsock::VsockSocket; use ax_net::{ Shutdown, Socket as SocketInner, SocketAddrEx, SocketOps, raw::{IpProtocol, IpVersion, RawSocket}, @@ -90,7 +90,7 @@ pub fn sys_socket(domain: u32, raw_ty: u32, proto: u32) -> AxResult { return add_file_like(socket as _, cloexec).map(|fd| fd as isize); } #[cfg(feature = "vsock")] - (AF_VSOCK, SOCK_STREAM) => VsockSocket::new(VsockStreamTransport::new()).into(), + (AF_VSOCK, SOCK_STREAM) => VsockSocket::new().into(), (AF_INET, SOCK_RAW) => { if proto != IPPROTO_ICMP as u32 { return Err(AxError::from(LinuxError::EPROTONOSUPPORT)); diff --git a/os/arceos/api/arceos_api/src/imp/net.rs b/os/arceos/api/arceos_api/src/imp/net.rs index c18807103a..898778045c 100644 --- a/os/arceos/api/arceos_api/src/imp/net.rs +++ b/os/arceos/api/arceos_api/src/imp/net.rs @@ -160,7 +160,7 @@ pub fn ax_dns_query(domain_name: &str) -> AxResult> { } pub fn ax_poll_interfaces() -> AxResult { - ax_net::poll_interfaces(); + ax_net::request_poll(); Ok(()) } diff --git a/os/arceos/api/arceos_posix_api/src/imp/io_mpx/epoll.rs b/os/arceos/api/arceos_posix_api/src/imp/io_mpx/epoll.rs index 9537f130ca..3c563295f3 100644 --- a/os/arceos/api/arceos_posix_api/src/imp/io_mpx/epoll.rs +++ b/os/arceos/api/arceos_posix_api/src/imp/io_mpx/epoll.rs @@ -321,7 +321,7 @@ pub unsafe fn sys_epoll_wait( let epoll_instance = EpollInstance::from_fd(epfd)?; loop { #[cfg(feature = "net")] - ax_net::poll_interfaces(); + ax_net::request_poll(); let events_num = epoll_instance.poll_all(events)?; if events_num > 0 { return Ok(events_num as c_int); diff --git a/os/arceos/api/arceos_posix_api/src/imp/io_mpx/poll.rs b/os/arceos/api/arceos_posix_api/src/imp/io_mpx/poll.rs index f639abee95..6193848b5c 100644 --- a/os/arceos/api/arceos_posix_api/src/imp/io_mpx/poll.rs +++ b/os/arceos/api/arceos_posix_api/src/imp/io_mpx/poll.rs @@ -46,7 +46,7 @@ pub fn sys_poll(fds: *mut ctypes::pollfd, nfds: ctypes::nfds_t, timeout: c_int) loop { #[cfg(feature = "net")] - ax_net::poll_interfaces(); + ax_net::request_poll(); let mut ready_count: usize = 0; diff --git a/os/arceos/api/arceos_posix_api/src/imp/io_mpx/select.rs b/os/arceos/api/arceos_posix_api/src/imp/io_mpx/select.rs index 04b16120c6..5c58d644fa 100644 --- a/os/arceos/api/arceos_posix_api/src/imp/io_mpx/select.rs +++ b/os/arceos/api/arceos_posix_api/src/imp/io_mpx/select.rs @@ -135,7 +135,7 @@ pub unsafe fn sys_select( loop { #[cfg(feature = "net")] - ax_net::poll_interfaces(); + ax_net::request_poll(); let res = fd_sets.poll_all(readfds, writefds, exceptfds)?; if res > 0 { return Ok(res); diff --git a/os/arceos/modules/axruntime/src/devices.rs b/os/arceos/modules/axruntime/src/devices.rs index 7b615dac24..7c38330312 100644 --- a/os/arceos/modules/axruntime/src/devices.rs +++ b/os/arceos/modules/axruntime/src/devices.rs @@ -128,7 +128,7 @@ fn adapt_net_device( let policy = if let Some(ctrl) = net.wifi_control() { // SDIO Wi-Fi RX is out-of-band (not the ethernet IRQ framework); the // chip's RX-data callback wakes the stack's dedicated poll task. - ctrl.set_rx_wake(ax_net::notify_oob_rx); + ctrl.set_rx_wake(ax_net::wake_net_task_irq); ctrl.link_policy() } else { None