use super::callbacks::{stream_exit_trampoline, stream_io_trampoline};
use crate::channel::AsyncReceiver;
use crate::error::WslcError;
use dashmap::DashMap;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex, OnceLock, Weak};
use tokio::sync::{mpsc, oneshot};
use wslcsdk_sys::types::WslcProcessCallbacks;
static NEXT_STREAM_ID: AtomicU64 = AtomicU64::new(1);
type StreamRegistry = DashMap<u64, Weak<StreamState>>;
pub(super) fn stream_registry() -> &'static StreamRegistry {
static REGISTRY: OnceLock<StreamRegistry> = OnceLock::new();
REGISTRY.get_or_init(DashMap::new)
}
#[derive(Debug)]
pub struct ProcessStreams {
pub stdout: AsyncReceiver<Vec<u8>>,
pub stderr: AsyncReceiver<Vec<u8>>,
exit_rx: Option<oneshot::Receiver<i32>>,
pub(crate) stream_state: Arc<StreamState>,
}
#[derive(Debug)]
pub(crate) struct StreamState {
pub(crate) id: u64,
pub(super) stdout_tx: Mutex<Option<mpsc::Sender<Vec<u8>>>>,
pub(super) stderr_tx: Mutex<Option<mpsc::Sender<Vec<u8>>>>,
pub(crate) exit_tx: Mutex<Option<oneshot::Sender<i32>>>,
pub(crate) stdout_dropped_bytes: AtomicU64,
pub(crate) stderr_dropped_bytes: AtomicU64,
}
impl StreamState {
pub(super) fn close_io_channels(&self) {
if let Ok(mut guard) = self.stdout_tx.lock() {
guard.take();
}
if let Ok(mut guard) = self.stderr_tx.lock() {
guard.take();
}
}
}
impl Drop for StreamState {
fn drop(&mut self) {
stream_registry().remove(&self.id);
let stdout_dropped = self.stdout_dropped_bytes.load(Ordering::Relaxed);
let stderr_dropped = self.stderr_dropped_bytes.load(Ordering::Relaxed);
if stdout_dropped > 0 || stderr_dropped > 0 {
log::warn!(
"流式会话结束但存在被丢弃的数据,流 ID: {},标准输出丢弃 {} 字节,标准错误丢弃 {} 字节,请考虑增大通道容量或加快消费速度",
self.id,
stdout_dropped,
stderr_dropped
);
}
}
}
impl ProcessStreams {
pub fn stdout_dropped_bytes(&self) -> u64 {
self.stream_state
.stdout_dropped_bytes
.load(Ordering::Relaxed)
}
pub fn stderr_dropped_bytes(&self) -> u64 {
self.stream_state
.stderr_dropped_bytes
.load(Ordering::Relaxed)
}
pub async fn wait_exit(&mut self) -> Result<i32, WslcError> {
let Some(exit_rx) = self.exit_rx.take() else {
return Err(WslcError::AlreadyConsumed("进程退出通知".to_string()));
};
exit_rx
.await
.map_err(|_| WslcError::ChannelTerminated("进程退出通知通道".to_string()))
}
}
pub(crate) fn setup_streaming_channels(
capacity: usize,
) -> (
WslcProcessCallbacks,
usize,
Arc<StreamState>,
ProcessStreams,
) {
let cap = capacity.max(1);
let (stdout_tx, stdout_rx) = mpsc::channel::<Vec<u8>>(cap);
let (stderr_tx, stderr_rx) = mpsc::channel::<Vec<u8>>(cap);
let (exit_tx, exit_rx) = oneshot::channel::<i32>();
let id = NEXT_STREAM_ID.fetch_add(1, Ordering::Relaxed);
let state = Arc::new(StreamState {
id,
stdout_tx: Mutex::new(Some(stdout_tx)),
stderr_tx: Mutex::new(Some(stderr_tx)),
exit_tx: Mutex::new(Some(exit_tx)),
stdout_dropped_bytes: AtomicU64::new(0),
stderr_dropped_bytes: AtomicU64::new(0),
});
stream_registry().insert(id, Arc::downgrade(&state));
let callbacks = WslcProcessCallbacks {
on_stdout: Some(stream_io_trampoline),
on_stderr: Some(stream_io_trampoline),
on_exit: Some(stream_exit_trampoline),
};
let streams = ProcessStreams {
stdout: AsyncReceiver::new(stdout_rx),
stderr: AsyncReceiver::new(stderr_rx),
exit_rx: Some(exit_rx),
stream_state: state.clone(),
};
(callbacks, id as usize, state, streams)
}
#[cfg(test)]
mod tests {
use super::*;
use core::ffi::c_void;
use wslcsdk_sys::types::WslcProcessIOHandle;
fn setup(
capacity: usize,
) -> (
WslcProcessCallbacks,
usize,
Arc<StreamState>,
ProcessStreams,
) {
setup_streaming_channels(capacity)
}
unsafe fn emit_stdout(context: usize, payload: &[u8]) {
let ctx = context as *mut c_void;
unsafe {
stream_io_trampoline(
WslcProcessIOHandle::Stdout,
payload.as_ptr(),
payload.len() as u32,
ctx,
)
};
}
unsafe fn emit_stderr(context: usize, payload: &[u8]) {
let ctx = context as *mut c_void;
unsafe {
stream_io_trampoline(
WslcProcessIOHandle::Stderr,
payload.as_ptr(),
payload.len() as u32,
ctx,
)
};
}
fn test_runtime() -> tokio::runtime::Runtime {
tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.expect("创建测试运行时失败")
}
#[test]
fn test_registry_entry_lifecycle() {
let (callbacks, context, state, streams) = setup(8);
assert!(callbacks.on_stdout.is_some(), "必须注册标准输出回调");
assert!(callbacks.on_stderr.is_some(), "必须注册标准错误回调");
assert!(callbacks.on_exit.is_some(), "必须注册退出回调");
assert!(stream_registry().contains_key(&(context as u64)));
let id = state.id;
drop(state);
assert!(
stream_registry().contains_key(&id),
"仅释放一份强引用时注册表条目应仍然存在"
);
drop(streams);
assert!(
!stream_registry().contains_key(&id),
"流状态析构后注册表条目应被移除"
);
}
#[test]
fn test_trampoline_routes_payload_to_matching_channel() {
let rt = test_runtime();
let (_callbacks, context, _state, mut streams) = setup(8);
rt.block_on(async {
unsafe { emit_stdout(context, b"hello-stdout") };
unsafe { emit_stderr(context, b"hello-stderr") };
let out = streams.stdout.recv().await.expect("应收到标准输出");
let err = streams.stderr.recv().await.expect("应收到标准错误");
assert_eq!(out, b"hello-stdout".to_vec());
assert_eq!(err, b"hello-stderr".to_vec());
});
}
#[test]
fn test_backpressure_drops_and_accumulates_byte_count() {
let rt = test_runtime();
let (_callbacks, context, state, mut streams) = setup(1);
rt.block_on(async {
unsafe { emit_stdout(context, b"first") };
unsafe { emit_stdout(context, b"second") };
unsafe { emit_stdout(context, b"third") };
let first = streams.stdout.recv().await.expect("应收到首条数据");
assert_eq!(first, b"first".to_vec());
assert_eq!(
streams.stdout_dropped_bytes(),
11,
"被丢弃的两条分别为 6 与 5 字节,合计 11 字节"
);
});
assert_eq!(state.stdout_dropped_bytes.load(Ordering::Relaxed), 11);
assert_eq!(
state.stderr_dropped_bytes.load(Ordering::Relaxed),
0,
"标准错误方向不应受到标准输出丢弃的影响"
);
}
#[test]
fn test_trampoline_ignores_invalid_inputs() {
let (_callbacks, context, _state, mut streams) = setup(4);
let ctx = context as *mut c_void;
unsafe {
stream_io_trampoline(WslcProcessIOHandle::Stdout, b"data".as_ptr(), 0, ctx);
stream_io_trampoline(
WslcProcessIOHandle::Stdout,
b"data".as_ptr(),
4,
std::ptr::null_mut(),
);
stream_io_trampoline(WslcProcessIOHandle::Stdout, std::ptr::null(), 4, ctx);
}
assert!(
streams.stdout.try_recv().is_err(),
"空载荷、空上下文与空数据指针均不应产生数据"
);
assert_eq!(streams.stdout_dropped_bytes(), 0);
}
#[test]
fn test_callback_after_state_released_is_ignored() {
let (_callbacks, context, state, streams) = setup(4);
let id = state.id;
drop(state);
drop(streams);
unsafe { emit_stdout(context, b"late-data") };
assert!(!stream_registry().contains_key(&id));
}
#[test]
fn test_exit_trampoline_delivers_code() {
let rt = test_runtime();
let (_callbacks, context, _state, mut streams) = setup(4);
rt.block_on(async {
unsafe { stream_exit_trampoline(7, context as *mut c_void) };
assert_eq!(streams.wait_exit().await.expect("应收到退出码"), 7);
});
}
#[test]
fn test_exit_trampoline_with_null_context_is_ignored() {
let rt = test_runtime();
let (_callbacks, _context, _state, mut streams) = setup(4);
rt.block_on(async {
unsafe { stream_exit_trampoline(0, std::ptr::null_mut()) };
let res =
tokio::time::timeout(std::time::Duration::from_millis(50), streams.wait_exit())
.await;
assert!(res.is_err(), "空上下文不应触发退出通知");
});
}
#[test]
fn test_exit_closes_io_channels_so_consumer_loop_terminates() {
let rt = test_runtime();
let (_callbacks, context, _state, mut streams) = setup(4);
rt.block_on(async {
unsafe { emit_stdout(context, b"before-exit") };
unsafe { emit_stderr(context, b"before-exit") };
unsafe { stream_exit_trampoline(0, context as *mut c_void) };
assert_eq!(streams.stdout.recv().await, Some(b"before-exit".to_vec()));
assert!(
streams.stdout.recv().await.is_none(),
"缓冲耗尽后应返回 None"
);
assert_eq!(streams.stderr.recv().await, Some(b"before-exit".to_vec()));
assert!(
streams.stderr.recv().await.is_none(),
"缓冲耗尽后应返回 None"
);
});
}
#[test]
fn test_zero_capacity_is_normalized_to_one() {
let (_callbacks, _context, _state, mut streams) = setup(0);
let _ = &mut streams;
assert_eq!(streams.stdout_dropped_bytes(), 0);
}
}