1use crate::{
4 CancelToken, DetachedFuture, DetachedReader, DetachedWriter, Error, Gate, Path, Record, Value,
5};
6use std::{
7 collections::VecDeque,
8 sync::{Arc, Mutex},
9};
10
11struct State {
12 bytes: VecDeque<u8>,
13 closed: bool,
14 released: bool,
15}
16struct Pipe {
17 capacity: usize,
18 state: Mutex<State>,
19 gate: Gate,
20}
21impl Pipe {
22 fn new(capacity: usize) -> Arc<Self> {
23 Arc::new(Self {
24 capacity,
25 state: Mutex::new(State {
26 bytes: VecDeque::new(),
27 closed: false,
28 released: false,
29 }),
30 gate: Gate::new(),
31 })
32 }
33 fn close(&self, release: bool) {
34 let mut state = self.state.lock().unwrap_or_else(|e| e.into_inner());
35 state.closed = true;
36 if release {
37 state.released = true;
38 state.bytes.clear();
39 }
40 drop(state);
41 self.gate.notify();
42 }
43 async fn read(&self, max: usize, cancel: &CancelToken) -> Result<Vec<u8>, Error> {
44 if max == 0 {
45 return Err(Error::conflict("read size must be positive"));
46 }
47 let result = self
48 .gate
49 .wait_until_cancellable(cancel, || {
50 let mut state = self.state.lock().unwrap_or_else(|e| e.into_inner());
51 if state.released {
52 return Some(Err(Error::cancelled("stream released")));
53 }
54 if !state.bytes.is_empty() || state.closed {
55 let count = max.min(state.bytes.len());
56 return Some(Ok(state.bytes.drain(..count).collect()));
57 }
58 None
59 })
60 .await
61 .map_err(|e| e.into_error("stream read cancelled"))?;
62 self.gate.notify();
63 result
64 }
65 async fn write(&self, bytes: &[u8], cancel: &CancelToken) -> Result<(), Error> {
66 if bytes.len() > self.capacity {
67 return Err(Error::resource_limit(
68 "write exceeds stream capacity; chunk it",
69 ));
70 }
71 let result = self
72 .gate
73 .wait_until_cancellable(cancel, || {
74 let mut state = self.state.lock().unwrap_or_else(|e| e.into_inner());
75 if state.closed {
76 return Some(Err(Error::cancelled("stream write side closed")));
77 }
78 if bytes.len() <= self.capacity - state.bytes.len() {
79 state.bytes.extend(bytes);
80 return Some(Ok(()));
81 }
82 None
83 })
84 .await
85 .map_err(|e| e.into_error("stream write cancelled"))?;
86 self.gate.notify();
87 result
88 }
89}
90
91#[derive(Clone, Copy, Debug, PartialEq, Eq)]
92pub struct StreamReadiness {
93 pub readable_bytes: usize,
94 pub writable_bytes: usize,
95 pub eof: bool,
96 pub write_closed: bool,
97 pub released: bool,
98}
99pub struct DuplexStream {
103 rx: Arc<Pipe>,
104 tx: Arc<Pipe>,
105}
106impl DuplexStream {
107 pub fn pair(capacity_per_direction: usize) -> Result<(Arc<Self>, Arc<Self>), Error> {
108 if capacity_per_direction == 0 {
109 return Err(Error::resource_limit("stream capacity must be positive"));
110 }
111 let a = Pipe::new(capacity_per_direction);
112 let b = Pipe::new(capacity_per_direction);
113 Ok((
114 Arc::new(Self {
115 rx: a.clone(),
116 tx: b.clone(),
117 }),
118 Arc::new(Self { rx: b, tx: a }),
119 ))
120 }
121 pub async fn read(&self, max: usize, cancel: &CancelToken) -> Result<Vec<u8>, Error> {
122 self.rx.read(max, cancel).await
123 }
124 pub async fn write(&self, bytes: &[u8], cancel: &CancelToken) -> Result<(), Error> {
125 self.tx.write(bytes, cancel).await
126 }
127 pub fn shutdown_write(&self) {
129 self.tx.close(false);
130 }
131 pub fn release(&self) {
133 self.rx.close(true);
134 self.tx.close(true);
135 }
136 pub fn readiness(&self) -> StreamReadiness {
138 let (readable_bytes, eof, released) = {
140 let state = self.rx.state.lock().unwrap_or_else(|e| e.into_inner());
141 (
142 state.bytes.len(),
143 state.closed && state.bytes.is_empty(),
144 state.released,
145 )
146 };
147 let state = self.tx.state.lock().unwrap_or_else(|e| e.into_inner());
148 StreamReadiness {
149 readable_bytes,
150 eof,
151 released: released || state.released,
152 write_closed: state.closed,
153 writable_bytes: if state.closed {
154 0
155 } else {
156 self.tx.capacity - state.bytes.len()
157 },
158 }
159 }
160 pub async fn ready(
163 &self,
164 read: bool,
165 write: bool,
166 cancel: &CancelToken,
167 ) -> Result<StreamReadiness, Error> {
168 if !read && !write {
169 return Err(Error::conflict("readiness needs an interest"));
170 }
171 let check = || {
172 let r = self.readiness();
173 (r.released
174 || (read && (r.readable_bytes > 0 || r.eof))
175 || (write && (r.writable_bytes > 0 || r.write_closed)))
176 .then_some(r)
177 };
178 tokio::select! {
179 r = self.rx.gate.wait_until_cancellable(cancel, check) => r,
180 r = self.tx.gate.wait_until_cancellable(cancel, check) => r,
181 }
182 .map_err(|e| e.into_error("stream readiness cancelled"))
183 }
184 pub fn store(self: &Arc<Self>) -> StreamStore {
186 StreamStore(self.clone())
187 }
188}
189impl Drop for DuplexStream {
190 fn drop(&mut self) {
191 self.release();
192 }
193}
194
195#[derive(Clone)]
200pub struct StreamStore(Arc<DuplexStream>);
201impl DetachedReader for StreamStore {
202 fn read_detached(&mut self, path: &Path) -> DetachedFuture<Option<Record>> {
203 let stream = self.0.clone();
204 let path = path.clone();
205 Box::pin(async move {
206 let parts: Vec<_> = path.iter().collect();
207 let cancel = CancelToken::new();
208 let value = match parts.as_slice() {
209 ["rx", max] => Value::Bytes(
210 stream
211 .read(
212 max.parse()
213 .map_err(|_| Error::conflict("invalid read size"))?,
214 &cancel,
215 )
216 .await?,
217 ),
218 ["ready", interest @ ("read" | "write" | "both")] => {
219 let r = stream
220 .ready(*interest != "write", *interest != "read", &cancel)
221 .await?;
222 Value::Map(std::collections::BTreeMap::from([
223 (
224 "readable_bytes".into(),
225 Value::Integer(r.readable_bytes as i64),
226 ),
227 (
228 "writable_bytes".into(),
229 Value::Integer(r.writable_bytes as i64),
230 ),
231 ("eof".into(), Value::Bool(r.eof)),
232 ("write_closed".into(), Value::Bool(r.write_closed)),
233 ("released".into(), Value::Bool(r.released)),
234 ]))
235 }
236 _ => return Err(Error::not_found(path.clone())),
237 };
238 Ok(Some(Record::parsed(value)))
239 })
240 }
241}
242impl DetachedWriter for StreamStore {
243 fn write_detached(&mut self, path: &Path, data: Record) -> DetachedFuture<Path> {
244 let stream = self.0.clone();
245 let path = path.clone();
246 Box::pin(async move {
247 let value = data.into_value(&structfs_core_store::NoCodec)?;
248 if path.is_empty() && value.is_null() {
249 stream.release();
250 } else if path.to_string() == "shutdown" && value.is_null() {
251 stream.shutdown_write();
252 } else if path.to_string() == "tx" {
253 let Value::Bytes(bytes) = value else {
254 return Err(Error::conflict("tx expects Bytes"));
255 };
256 stream.write(&bytes, &CancelToken::new()).await?;
257 } else {
258 return Err(Error::not_found(path.clone()));
259 }
260 Ok(path)
261 })
262 }
263}
264
265#[cfg(test)]
266mod tests {
267 use super::*;
268 #[tokio::test]
269 async fn backpressure_half_close_and_reverse_traffic() {
270 let (a, b) = DuplexStream::pair(4).unwrap();
271 let cancel = CancelToken::new();
272 a.write(b"abcd", &cancel).await.unwrap();
273 assert_eq!(a.readiness().writable_bytes, 0);
274 assert!(a.write(b"large", &cancel).await.is_err());
275 let write = a.write(b"ef", &cancel);
276 tokio::pin!(write);
277 assert!(
278 tokio::time::timeout(std::time::Duration::from_millis(5), &mut write)
279 .await
280 .is_err()
281 );
282 assert_eq!(b.read(2, &cancel).await.unwrap(), b"ab");
283 write.await.unwrap();
284 a.shutdown_write();
285 assert_eq!(b.read(4, &cancel).await.unwrap(), b"cdef");
286 assert!(b.read(1, &cancel).await.unwrap().is_empty());
287 b.write(b"back", &cancel).await.unwrap();
288 assert_eq!(a.read(4, &cancel).await.unwrap(), b"back");
289 }
290 #[tokio::test]
291 async fn cancel_and_release_wake_without_transferring_bytes() {
292 let (a, b) = DuplexStream::pair(1).unwrap();
293 let cancel = CancelToken::new();
294 a.write(b"a", &cancel).await.unwrap();
295 cancel.cancel();
296 assert!(a.write(b"b", &cancel).await.is_err());
297 assert_eq!(b.read(1, &CancelToken::new()).await.unwrap(), b"a");
298 let token = CancelToken::new();
299 let waiting = b.ready(true, false, &token);
300 drop(a);
302 assert!(waiting.await.unwrap().released);
303 assert!(b.read(1, &CancelToken::new()).await.is_err());
304 }
305}