use alloc::vec::Vec;
use core::{fmt, mem::offset_of, time::Duration};
use axpoll::IoEvents;
use bitmaps::Bitmap;
use linux_raw_sys::{
general::*,
select_macros::{FD_ISSET, FD_SET, FD_ZERO},
};
use starry_signal::SignalSet;
use super::FdPollSet;
use crate::{
StarryError, StarryResult,
file::current_fd_table,
mm::{UserConstPtr, UserPtr},
syscall::signal::check_sigset_size,
task::{
future::{UserWaitOutcome, block_on_user_timeout, poll_io},
with_blocked_signals,
},
time::TimeValueLike,
};
struct FdSet(Bitmap<{ __FD_SETSIZE as usize }>);
impl FdSet {
fn new(nfds: usize, fds: Option<&__kernel_fd_set>) -> Self {
let mut bitmap = Bitmap::new();
if let Some(fds) = fds {
for i in 0..nfds {
if unsafe { FD_ISSET(i as _, fds) } {
bitmap.set(i, true);
}
}
}
Self(bitmap)
}
}
fn write_fd_set(user: Option<&mut __kernel_fd_set>, selected: &FdSet, nfds: usize) {
if let Some(user) = user {
unsafe { FD_ZERO(user) };
for index in selected.0.into_iter().take(nfds) {
unsafe { FD_SET(index as _, user) };
}
}
}
impl fmt::Debug for FdSet {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_list().entries(&self.0).finish()
}
}
fn do_select(
current: &crate::task::UserTaskRef,
nfds: u32,
readfds: UserPtr<__kernel_fd_set>,
writefds: UserPtr<__kernel_fd_set>,
exceptfds: UserPtr<__kernel_fd_set>,
timeout: Option<Duration>,
sigmask: UserConstPtr<SignalSetWithSize>,
) -> StarryResult<isize> {
if nfds > __FD_SETSIZE {
return Err(StarryError::InvalidInput);
}
let sigmask = if sigmask.is_null() {
None
} else {
let sigmask = unsafe { sigmask.read_abi(current)? };
check_sigset_size(sigmask.sigsetsize)?;
let set = UserConstPtr::<SignalSet>::from(sigmask.set);
if set.is_null() {
None
} else {
Some(unsafe { set.read_abi(current)? })
}
};
let mut readfds_value = if readfds.is_null() {
None
} else {
Some(unsafe { readfds.read_abi(current)? })
};
let mut writefds_value = if writefds.is_null() {
None
} else {
Some(unsafe { writefds.read_abi(current)? })
};
let mut exceptfds_value = if exceptfds.is_null() {
None
} else {
Some(unsafe { exceptfds.read_abi(current)? })
};
let read_set = FdSet::new(nfds as _, readfds_value.as_ref());
let write_set = FdSet::new(nfds as _, writefds_value.as_ref());
let except_set = FdSet::new(nfds as _, exceptfds_value.as_ref());
debug!(
"sys_select <= nfds: {nfds} sets: [read: {read_set:?}, write: {write_set:?}, except: \
{except_set:?}] timeout: {timeout:?}"
);
let fd_table_owner = current_fd_table();
let fd_table = fd_table_owner.read();
let fd_bitmap = read_set.0 | write_set.0 | except_set.0;
let fd_count = fd_bitmap.len();
let mut fds = Vec::with_capacity(fd_count);
let mut fd_indices = Vec::with_capacity(fd_count);
for fd in fd_bitmap.into_iter() {
let f = fd_table
.get(fd)
.ok_or(StarryError::BadFileDescriptor)?
.inner
.clone();
let mut events = IoEvents::empty();
events.set(IoEvents::IN, read_set.0.get(fd));
events.set(IoEvents::OUT, write_set.0.get(fd));
events.set(IoEvents::ERR, except_set.0.get(fd));
if !events.is_empty() {
fds.push((f, events));
fd_indices.push(fd);
}
}
drop(fd_table);
let fds = FdPollSet(fds);
let task = current;
let result = with_blocked_signals(sigmask, || {
let result = block_on_user_timeout(
task,
timeout,
poll_io(&fds, IoEvents::empty(), false, || {
let mut res = 0usize;
let mut selected_readfds = FdSet(Bitmap::new());
let mut selected_writefds = FdSet(Bitmap::new());
let mut selected_exceptfds = FdSet(Bitmap::new());
for ((fd, interested), index) in fds.0.iter().zip(fd_indices.iter().copied()) {
let events = fd.poll();
let always_report = events & IoEvents::ALWAYS_POLL;
let write_report = events & IoEvents::ERR;
let selected = events & *interested;
let selected_read = selected.contains(IoEvents::IN)
|| (read_set.0.get(index) && !always_report.is_empty());
let selected_write = selected.contains(IoEvents::OUT)
|| (write_set.0.get(index) && !write_report.is_empty());
let selected_except =
selected.contains(IoEvents::ERR) && except_set.0.get(index);
if selected_read {
res += 1;
selected_readfds.0.set(index, true);
}
if selected_write {
res += 1;
selected_writefds.0.set(index, true);
}
if selected_except {
res += 1;
selected_exceptfds.0.set(index, true);
}
}
if res > 0 {
write_fd_set(readfds_value.as_mut(), &selected_readfds, nfds as _);
write_fd_set(writefds_value.as_mut(), &selected_writefds, nfds as _);
write_fd_set(exceptfds_value.as_mut(), &selected_exceptfds, nfds as _);
return Ok(res as _);
}
Err(StarryError::WouldBlock)
}),
);
match result {
UserWaitOutcome::Ready(result) => result,
UserWaitOutcome::TimedOut => {
let empty = FdSet(Bitmap::new());
write_fd_set(readfds_value.as_mut(), &empty, nfds as _);
write_fd_set(writefds_value.as_mut(), &empty, nfds as _);
write_fd_set(exceptfds_value.as_mut(), &empty, nfds as _);
Ok(0)
}
UserWaitOutcome::Interrupted => Err(crate::StarryError::Interrupted),
}
});
if let Some(value) = readfds_value {
readfds.write_field(
current,
offset_of!(__kernel_fd_set, fds_bits),
value.fds_bits,
)?;
}
if let Some(value) = writefds_value {
writefds.write_field(
current,
offset_of!(__kernel_fd_set, fds_bits),
value.fds_bits,
)?;
}
if let Some(value) = exceptfds_value {
exceptfds.write_field(
current,
offset_of!(__kernel_fd_set, fds_bits),
value.fds_bits,
)?;
}
result
}
#[cfg(target_arch = "x86_64")]
pub fn sys_select(
current: &crate::task::UserTaskRef,
nfds: u32,
readfds: UserPtr<__kernel_fd_set>,
writefds: UserPtr<__kernel_fd_set>,
exceptfds: UserPtr<__kernel_fd_set>,
timeout: UserConstPtr<timeval>,
) -> StarryResult<isize> {
do_select(
current,
nfds,
readfds,
writefds,
exceptfds,
(if timeout.is_null() {
None
} else {
Some(unsafe { timeout.read_abi(current)? })
})
.map(|it| it.try_into_time_value())
.transpose()?,
0.into(),
)
}
#[repr(C)]
#[derive(Clone, Copy, bytemuck::AnyBitPattern)]
pub struct SignalSetWithSize {
set: usize,
sigsetsize: usize,
}
pub fn sys_pselect6(
current: &crate::task::UserTaskRef,
nfds: u32,
readfds: UserPtr<__kernel_fd_set>,
writefds: UserPtr<__kernel_fd_set>,
exceptfds: UserPtr<__kernel_fd_set>,
timeout: UserConstPtr<timespec>,
sigmask: UserConstPtr<SignalSetWithSize>,
) -> StarryResult<isize> {
do_select(
current,
nfds,
readfds,
writefds,
exceptfds,
(if timeout.is_null() {
None
} else {
Some(unsafe { timeout.read_abi(current)? })
})
.map(|ts| ts.try_into_time_value())
.transpose()?,
sigmask,
)
}