#[allow(unused_imports)]
pub use log::{debug, error, info, log, trace, warn};
use core::future::{Future, poll_fn};
use core::pin::pin;
use core::sync::atomic::AtomicBool;
use core::sync::atomic::Ordering::{AcqRel, Acquire, Relaxed};
use core::task::{Context, Poll, Poll::Pending, Poll::Ready};
use portable_atomic::AtomicUsize;
use embassy_futures::join;
use embassy_futures::select::select;
#[allow(unused_imports)]
use embassy_sync::blocking_mutex::raw::{CriticalSectionRawMutex, NoopRawMutex};
use embassy_sync::mutex::{Mutex, MutexGuard};
use embassy_sync::signal::Signal;
use embedded_io_async::{Read, Write};
use crate::async_channel::ChanIO;
use sunset::ChanData::{Normal, Stderr};
use sunset::config::MAX_CHANNELS;
use sunset::error::TrapBug;
use sunset::event::Event;
use sunset::{ChanData, ChanHandle, ChanNum, CliServ, Error, Result, Runner, error};
#[cfg(feature = "multi-thread")]
pub type SunsetRawMutex = CriticalSectionRawMutex;
#[cfg(not(feature = "multi-thread"))]
pub type SunsetRawMutex = NoopRawMutex;
pub type SunsetMutex<T> = Mutex<SunsetRawMutex, T>;
struct Inner<'a, CS: CliServ> {
runner: Runner<'a, CS>,
chan_handles: [Option<ChanHandle>; MAX_CHANNELS],
}
impl<'a, CS: CliServ> Inner<'a, CS> {
fn fetch(&mut self, num: ChanNum) -> Result<(&mut Runner<'a, CS>, &ChanHandle)> {
let ch = self
.chan_handles
.get(num.0 as usize)
.ok_or(Error::BadChannel { num })?
.as_ref()
.trap()?;
Ok((&mut self.runner, ch))
}
}
pub struct ProgressHolder<'g, 'a, CS: CliServ> {
guard: Option<MutexGuard<'g, SunsetRawMutex, Inner<'a, CS>>>,
}
impl<'g, 'a, CS: CliServ> ProgressHolder<'g, 'a, CS> {
pub fn new() -> Self {
Self { guard: None }
}
}
impl<CS: CliServ> Default for ProgressHolder<'_, '_, CS> {
fn default() -> Self {
Self::new()
}
}
pub(crate) struct AsyncSunset<'a, CS: CliServ> {
inner: SunsetMutex<Inner<'a, CS>>,
progress_notify: Signal<SunsetRawMutex, ()>,
last_progress_idled: AtomicBool,
moribund: AtomicBool,
chan_refcounts: [AtomicUsize; MAX_CHANNELS],
chan_norm_readcounts: [AtomicUsize; MAX_CHANNELS],
chan_stderr_readcounts: [AtomicUsize; MAX_CHANNELS],
}
impl<CS: CliServ> core::fmt::Debug for AsyncSunset<'_, CS> {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
let mut d = f.debug_struct("AsyncSunset");
if let Ok(i) = self.inner.try_lock() {
d.field("runner", &i.runner);
} else {
d.field("inner", &"(locked)");
}
d.finish_non_exhaustive()
}
}
impl<'a, CS: CliServ> AsyncSunset<'a, CS> {
pub fn new(runner: Runner<'a, CS>) -> Self {
let inner = Inner { runner, chan_handles: Default::default() };
let inner = Mutex::new(inner);
let progress_notify = Signal::new();
Self {
inner,
moribund: AtomicBool::new(false),
progress_notify,
chan_refcounts: Default::default(),
chan_norm_readcounts: Default::default(),
chan_stderr_readcounts: Default::default(),
last_progress_idled: AtomicBool::new(false),
}
}
pub async fn run(
&self,
rsock: &mut impl Read,
wsock: &mut impl Write,
) -> Result<()> {
let tx_stop = Signal::<SunsetRawMutex, ()>::new();
let rx_stop = Signal::<SunsetRawMutex, ()>::new();
let tx = async {
self.output_loop(wsock).await.inspect(|r| warn!("tx complete {r:?}"))
};
let tx = select(tx, tx_stop.wait());
let mut rxbuf = [0; 1024];
let rx = async {
loop {
let l = match rsock.read(&mut rxbuf).await {
Ok(0) => {
debug!("net EOF");
self.with_runner(|r| r.close_input()).await;
self.moribund.store(true, Relaxed);
self.wake_progress();
break Ok(());
}
Ok(l) => l,
Err(_) => {
info!("socket read error");
self.with_runner(|r| r.close_input()).await;
break Err(Error::ChannelEOF);
}
};
let mut rxbuf = &rxbuf[..l];
while !rxbuf.is_empty() {
let n = self.input(rxbuf).await?;
self.wake_progress();
rxbuf = &rxbuf[n..];
}
}
.inspect(|r| warn!("rx complete {r:?}"))
};
let rx = async {
let r = select(rx, rx_stop.wait()).await;
tx_stop.signal(());
r
};
let f = join::join(rx, tx).await;
let (_frx, _ftx) = f;
Ok(())
}
fn wake_progress(&self) {
trace!("wake_progress");
self.progress_notify.signal(())
}
fn discard_channels(&self, inner: &mut Inner<CS>) -> Result<()> {
if let Some((num, dt, _len)) = inner.runner.read_channel_ready() {
if self.chan_readcount(num, dt).load(Acquire) == 0 {
let ch = inner.chan_handles[num.0 as usize].as_ref().trap()?;
inner.runner.discard_read_channel(ch)?;
}
}
Ok(())
}
fn clear_refcounts(&self, inner: &mut Inner<CS>) -> Result<()> {
for (ch, count) in
inner.chan_handles.iter_mut().zip(self.chan_refcounts.iter())
{
let count = count.load(Acquire);
if count > 0 {
debug_assert!(ch.is_some());
continue;
}
if let Some(ch) = ch.take() {
inner.runner.channel_done(ch)?;
}
}
Ok(())
}
pub(crate) async fn progress<'g, 'f>(
&'g self,
ph: &'f mut ProgressHolder<'g, 'a, CS>,
) -> Result<Event<'f, 'a>> {
*ph = ProgressHolder::default();
let need_wait = self.last_progress_idled.load(Relaxed);
if need_wait {
self.last_progress_idled.store(false, Relaxed);
self.progress_notify.wait().await;
}
let inner = ph.guard.insert(self.inner.lock().await);
self.clear_refcounts(inner)?;
self.discard_channels(inner)?;
if self.moribund.load(Relaxed) {
debug!("All data flushed")
}
let ev = inner.runner.progress();
if matches!(ev, Ok(Event::None)) {
self.last_progress_idled.store(true, Relaxed);
}
ev
}
pub(crate) async fn with_runner<F, R>(&self, f: F) -> R
where
F: FnOnce(&mut Runner<CS>) -> R,
{
let mut inner = self.inner.lock().await;
f(&mut inner.runner)
}
fn chan_readcount(&self, num: ChanNum, dt: ChanData) -> &AtomicUsize {
let counts = match dt {
Normal => &self.chan_norm_readcounts,
Stderr => &self.chan_stderr_readcounts,
};
&counts[num.0 as usize]
}
async fn poll_inner<F, T>(&self, mut f: F) -> T
where
F: FnMut(&mut Inner<CS>, &mut Context) -> Poll<T>,
{
poll_fn(|cx| {
let i = self.inner.lock();
let i = pin!(i);
match i.poll(cx) {
Poll::Ready(mut inner) => f(&mut inner, cx),
Poll::Pending => {
Poll::Pending
}
}
})
.await
}
pub async fn output_loop(&self, wsock: &mut impl Write) -> Result<()> {
poll_fn(|cx| {
let i = self.inner.lock();
let i = pin!(i);
let Ready(mut inner) = i.poll(cx) else {
return Pending;
};
loop {
let buf = inner.runner.output_buf();
if buf.is_empty() {
inner.runner.set_output_waker(cx.waker());
return Pending;
}
let res = {
let w = wsock.write(buf);
let w = pin!(w);
w.poll(cx)
};
let r = match res {
Pending => {
Pending
}
Ready(Ok(0)) => {
info!("socket EOF");
inner.runner.close_output();
Ready(error::ChannelEOF.fail())
}
Ready(Ok(write_len)) => {
let buf_len = buf.len();
inner.runner.consume_output(write_len);
if write_len < buf_len {
continue;
}
inner.runner.set_output_waker(cx.waker());
if !inner.runner.is_output_pending() {
self.wake_progress();
}
Pending
}
Ready(Err(_e)) => {
info!("socket write error");
inner.runner.close_output();
Ready(error::ChannelEOF.fail())
}
};
return r;
}
})
.await
}
pub async fn input(&self, buf: &[u8]) -> Result<usize> {
let res = self
.poll_inner(|inner, cx| {
if inner.runner.is_input_ready() {
match inner.runner.input(buf) {
Ok(0) => {
inner.runner.set_input_waker(cx.waker());
Poll::Pending
}
Ok(n) => Poll::Ready(Ok(n)),
Err(e) => Poll::Ready(Err(e)),
}
} else {
inner.runner.set_input_waker(cx.waker());
Poll::Pending
}
})
.await;
self.wake_progress();
res
}
pub(crate) async fn add_channel(
&self,
handle: ChanHandle,
) -> Result<ChanIO<'_>> {
let mut inner = self.inner.lock().await;
let num = handle.num();
let idx = num.0 as usize;
if inner.chan_handles[idx].is_some() {
return error::Bug.fail();
}
inner.chan_handles[idx] = Some(handle);
debug_assert_eq!(self.chan_refcounts[idx].load(Relaxed), 0);
self.chan_refcounts[idx].store(1, Relaxed);
Ok(ChanIO::new_normal(num, self))
}
}
#[cfg(feature = "multi-thread")]
pub(crate) trait MaybeSend: Sync {}
#[cfg(not(feature = "multi-thread"))]
pub(crate) trait MaybeSend {}
impl<'a, CS: CliServ> MaybeSend for AsyncSunset<'a, CS> {}
pub(crate) trait ChanCore: MaybeSend {
fn inc_chan(&self, num: ChanNum);
fn dec_chan(&self, num: ChanNum);
fn inc_read_chan(&self, num: ChanNum, dt: ChanData);
fn dec_read_chan(&self, num: ChanNum, dt: ChanData);
fn poll_until_channel_closed(
&self,
cx: &mut Context,
num: ChanNum,
) -> Poll<Result<()>>;
fn poll_read_channel(
&self,
cx: &mut Context,
num: ChanNum,
dt: ChanData,
buf: &mut [u8],
) -> Poll<Result<usize>>;
fn poll_write_channel(
&self,
cx: &mut Context,
num: ChanNum,
dt: ChanData,
buf: &[u8],
) -> Poll<Result<usize>>;
fn poll_term_window_change(
&self,
cx: &mut Context,
num: ChanNum,
winch: &sunset::packets::WinChange,
) -> Poll<Result<()>>;
}
impl<'a, CS: CliServ> ChanCore for AsyncSunset<'a, CS> {
fn inc_chan(&self, num: ChanNum) {
let c = self.chan_refcounts[num.0 as usize].fetch_add(1, Relaxed);
debug_assert_ne!(c, 0);
debug_assert_ne!(c, usize::MAX);
}
fn dec_chan(&self, num: ChanNum) {
let c = self.chan_refcounts[num.0 as usize].fetch_sub(1, AcqRel);
debug_assert_ne!(c, 0);
if c == 1 {
self.wake_progress();
}
}
fn inc_read_chan(&self, num: ChanNum, dt: ChanData) {
let c = self.chan_readcount(num, dt).fetch_add(1, AcqRel);
debug_assert_ne!(c, usize::MAX);
}
fn dec_read_chan(&self, num: ChanNum, dt: ChanData) {
let c = self.chan_readcount(num, dt).fetch_sub(1, AcqRel);
debug_assert_ne!(c, 0);
if c == 1 {
self.wake_progress();
}
}
fn poll_until_channel_closed(
&self,
cx: &mut Context,
num: ChanNum,
) -> Poll<Result<()>> {
let i = self.inner.lock();
let i = pin!(i);
let Ready(mut inner) = i.poll(cx) else {
return Pending;
};
let (runner, h) = inner.fetch(num)?;
if runner.is_channel_closed(h) {
Poll::Ready(Ok(()))
} else {
runner.set_channel_read_waker(h, Normal, cx.waker());
Poll::Pending
}
}
fn poll_read_channel(
&self,
cx: &mut Context,
num: ChanNum,
dt: ChanData,
buf: &mut [u8],
) -> Poll<Result<usize>> {
let i = self.inner.lock();
let i = pin!(i);
let Ready(mut inner) = i.poll(cx) else {
return Pending;
};
let (runner, h) = inner.fetch(num)?;
let i = match runner.read_channel(h, dt, buf) {
Ok(0) => {
trace!("read ch {num:?} dt {dt:?} pending");
runner.set_channel_read_waker(h, dt, cx.waker());
Poll::Pending
}
Err(Error::ChannelEOF) => Poll::Ready(Ok(0)),
r => {
trace!("read ready ch {num:?} dt {dt:?} {r:?}");
Poll::Ready(r)
}
};
if matches!(i, Poll::Ready(_)) {
self.wake_progress()
}
i
}
fn poll_write_channel(
&self,
cx: &mut Context,
num: ChanNum,
dt: ChanData,
buf: &[u8],
) -> Poll<Result<usize>> {
if buf.is_empty() {
return Poll::Ready(Ok(0));
}
let i = self.inner.lock();
let i = pin!(i);
let Ready(mut inner) = i.poll(cx) else {
return Pending;
};
let (runner, h) = inner.fetch(num)?;
let l = runner.write_channel(h, dt, buf);
if let Ok(0) = l {
trace!("write ch {num:?} dt {dt:?} pending");
runner.set_channel_write_waker(h, dt, cx.waker());
Poll::Pending
} else {
trace!("write ready ch {num:?} dt {dt:?} {l:?}");
self.wake_progress();
Poll::Ready(l)
}
}
fn poll_term_window_change(
&self,
cx: &mut Context,
num: ChanNum,
winch: &sunset::packets::WinChange,
) -> Poll<Result<()>> {
let i = self.inner.lock();
let i = pin!(i);
let Ready(mut inner) = i.poll(cx) else {
return Pending;
};
let (runner, h) = inner.fetch(num)?;
Poll::Ready(runner.term_window_change(h, winch))
}
}