use alloc::vec::Vec;
use core::{
mem::{MaybeUninit, offset_of},
task::Poll,
};
use ax_runtime::hal::time::TimeValue;
use axpoll::{IoEvents, Pollable};
use linux_raw_sys::general::{POLLNVAL, RLIMIT_NOFILE, pollfd, timespec};
use starry_signal::SignalSet;
use super::FdPollSet;
use crate::{
StarryError, StarryResult,
file::get_file_like,
mm::{UserConstPtr, UserPtr, vm_read_slice, vm_write_slice},
syscall::signal::check_sigset_size,
task::{
future::{UserWaitOutcome, block_on_user_timeout, poll_shared},
with_blocked_signals,
},
time::TimeValueLike,
};
fn check_nfds_limit(current: &crate::task::UserTaskRef, nfds: usize) -> crate::StarryResult<()> {
let nofile = current.as_thread().proc_data.rlimit_current(RLIMIT_NOFILE);
if !nfds_within_limit(nfds, nofile) {
Err(StarryError::InvalidInput)
} else {
Ok(())
}
}
fn nfds_within_limit(nfds: usize, nofile: u64) -> bool {
nfds as u64 <= nofile
}
fn read_poll_fds(
current: &crate::task::UserTaskRef,
fds: UserPtr<pollfd>,
nfds: usize,
) -> crate::StarryResult<Vec<pollfd>> {
check_nfds_limit(current, nfds)?;
if nfds == 0 {
return Ok(Vec::new());
}
let mut buf = Vec::with_capacity(nfds);
buf.resize_with(nfds, MaybeUninit::uninit);
vm_read_slice(current, fds.as_ptr(), &mut buf)?;
Ok(buf
.into_iter()
.map(|fd| unsafe { fd.assume_init() })
.collect())
}
fn write_poll_revents(
current: &crate::task::UserTaskRef,
fds: UserPtr<pollfd>,
poll_fds: &[pollfd],
) -> crate::StarryResult<()> {
let revents_offset = offset_of!(pollfd, revents);
for (index, poll_fd) in poll_fds.iter().enumerate() {
let revents_ptr = (fds.as_ptr().wrapping_add(index) as *mut u8)
.wrapping_add(revents_offset)
.cast::<_>();
vm_write_slice(
current,
revents_ptr,
core::slice::from_ref(&poll_fd.revents),
)?;
}
Ok(())
}
fn mask_poll_revents(ready: IoEvents, requested: IoEvents) -> IoEvents {
(ready & requested) | (ready & IoEvents::ALWAYS_POLL)
}
fn collect_ready_poll_events(
fds: &FdPollSet,
revent_indices: &[usize],
poll_fds: &mut [pollfd],
) -> usize {
let mut res = 0usize;
for ((fd, events), revent_index) in fds.0.iter().zip(revent_indices.iter()) {
let result = mask_poll_revents(fd.poll(), *events);
let revents = &mut poll_fds[*revent_index].revents;
*revents = result.bits() as _;
if *revents != 0 {
res += 1;
}
}
res
}
fn do_poll(
current: &crate::task::UserTaskRef,
poll_fds: &mut [pollfd],
timeout: Option<TimeValue>,
sigmask: Option<SignalSet>,
) -> StarryResult<isize> {
debug!("do_poll fds={poll_fds:?} timeout={timeout:?}");
let mut invalid_count = 0isize;
let mut fds = Vec::with_capacity(poll_fds.len());
let mut revent_indices = Vec::with_capacity(poll_fds.len());
for (index, fd) in poll_fds.iter_mut().enumerate() {
fd.revents = 0;
if fd.fd < 0 {
continue;
}
match get_file_like(fd.fd) {
Ok(f) => {
fds.push((
f,
IoEvents::from_bits_truncate(u32::from(fd.events as u16))
| IoEvents::ALWAYS_POLL,
));
revent_indices.push(index);
}
Err(_) => {
fd.revents = POLLNVAL as _;
invalid_count += 1;
}
}
}
let fds = FdPollSet(fds);
if invalid_count > 0 {
let ready_count = collect_ready_poll_events(&fds, &revent_indices, poll_fds);
return Ok(invalid_count + ready_count as isize);
}
with_blocked_signals(sigmask, || {
let wait = poll_shared(
|| {
let res = collect_ready_poll_events(&fds, &revent_indices, poll_fds);
if res > 0 {
return Poll::Ready(Ok(res as _));
}
Poll::Pending
},
|registrar| unsafe { fds.register_shared(registrar, IoEvents::empty()) },
);
let task = current;
match block_on_user_timeout(task, timeout, wait) {
UserWaitOutcome::Ready(result) => result,
UserWaitOutcome::TimedOut => Ok(0),
UserWaitOutcome::Interrupted => Err(crate::StarryError::Interrupted),
}
})
}
#[cfg(target_arch = "x86_64")]
pub fn sys_poll(
current: &crate::task::UserTaskRef,
fds: UserPtr<pollfd>,
nfds: u32,
timeout: i32,
) -> crate::StarryResult<isize> {
let nfds = nfds as usize;
let mut poll_fds = read_poll_fds(current, fds, nfds)?;
let timeout = if timeout < 0 {
None
} else {
Some(TimeValue::from_millis(timeout as u64))
};
let res = do_poll(current, &mut poll_fds, timeout, None);
if nfds > 0 {
write_poll_revents(current, fds, &poll_fds)?;
}
res
}
pub fn sys_ppoll(
current: &crate::task::UserTaskRef,
fds: UserPtr<pollfd>,
nfds: i32,
timeout: UserConstPtr<timespec>,
sigmask: UserConstPtr<SignalSet>,
sigsetsize: usize,
) -> StarryResult<isize> {
if !sigmask.is_null() {
check_sigset_size(sigsetsize)?;
}
let nfds = nfds
.try_into()
.map_err(|_| crate::StarryError::InvalidInput)?;
let mut poll_fds = read_poll_fds(current, fds, nfds)?;
let timeout = (if timeout.is_null() {
None
} else {
Some(unsafe { timeout.read_abi(current)? })
})
.map(|ts| ts.try_into_time_value())
.transpose()?;
let sigmask = if sigmask.is_null() {
None
} else {
Some(unsafe { sigmask.read_abi(current)? })
};
let res = do_poll(current, &mut poll_fds, timeout, sigmask);
if nfds > 0 {
write_poll_revents(current, fds, &poll_fds)?;
}
res
}
#[cfg(all(test, not(axtest)))]
fn poll_nfds_validation_rules_hold_for_test() -> bool {
assert!(nfds_within_limit(0, 0));
assert!(nfds_within_limit(1024, 1024));
assert!(!nfds_within_limit(1025, 1024));
const { assert!(POLLNVAL != 0) }
let socket_ready = IoEvents::IN | IoEvents::OUT | IoEvents::RDHUP;
let without_rdhup = mask_poll_revents(socket_ready, IoEvents::IN | IoEvents::OUT);
let with_rdhup = mask_poll_revents(socket_ready, IoEvents::IN | IoEvents::RDHUP);
let always_reported = mask_poll_revents(IoEvents::ERR | IoEvents::HUP, IoEvents::empty());
without_rdhup.bits() == (IoEvents::IN | IoEvents::OUT).bits()
&& with_rdhup.bits() == (IoEvents::IN | IoEvents::RDHUP).bits()
&& always_reported.bits() == (IoEvents::ERR | IoEvents::HUP).bits()
}
#[cfg(all(test, not(axtest)))]
mod tests {
use axpoll::IoEvents;
use super::mask_poll_revents;
#[test]
fn pollrdhup_is_only_reported_when_requested() {
let ready = IoEvents::IN | IoEvents::OUT | IoEvents::RDHUP;
assert_eq!(
mask_poll_revents(ready, IoEvents::IN | IoEvents::OUT).bits(),
(IoEvents::IN | IoEvents::OUT).bits()
);
assert_eq!(
mask_poll_revents(ready, IoEvents::IN | IoEvents::RDHUP).bits(),
(IoEvents::IN | IoEvents::RDHUP).bits()
);
}
#[test]
fn pollerr_and_pollhup_are_reported_without_interest() {
assert_eq!(
mask_poll_revents(IoEvents::ERR | IoEvents::HUP, IoEvents::empty()).bits(),
(IoEvents::ERR | IoEvents::HUP).bits()
);
}
#[test]
fn poll_nfds_validation_rules_hold() {
assert!(super::poll_nfds_validation_rules_hold_for_test());
}
}