Skip to main content

adb_transport/
socket.rs

1use crate::error::ErrorKind;
2use crate::message::{AdbCommand, AdbMessage};
3use crate::transport::{TransportBackend, MAX_PAYLOAD};
4use crate::util::MaybeDone;
5use crate::{Error, Result};
6use diatomic_waker::{WakeSink, WakeSource};
7use parking_lot::Mutex;
8use std::collections::VecDeque;
9use std::future::Future;
10use std::io;
11use std::io::Cursor;
12use std::pin::Pin;
13use std::sync::Arc;
14use std::task::{ready, Context, Poll};
15use std::time::Duration;
16use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
17use tokio::spawn;
18use tokio::time::{interval, Interval};
19use tracing::span::EnteredSpan;
20use tracing::{debug, info, info_span, warn, Instrument, Span};
21
22fn io_error(cause: impl Into<Error>) -> io::Error {
23    io::Error::new(io::ErrorKind::BrokenPipe, cause.into())
24}
25
26pub(crate) struct State {
27    pub remote_id: u32,
28    pub closed: bool,
29
30    read_buffer: VecDeque<Cursor<Vec<u8>>>,
31    read_fut: MaybeDone<io::Result<()>>,
32
33    initial_acknowledgment: isize,
34    acknowledged: isize,
35    write_fut: MaybeDone<io::Result<()>>,
36    write_throttle: Interval,
37}
38
39pub(crate) struct SocketBackend {
40    transport: Arc<TransportBackend>,
41    local_id: u32,
42    span: Span,
43
44    pub read_waker: WakeSource,
45    pub write_waker: WakeSource,
46
47    pub state: Mutex<State>,
48}
49
50impl SocketBackend {
51    fn check_transport_error(&self) -> io::Result<()> {
52        if let Some(err) = self.transport.error.get() {
53            debug!(?err, "transport error");
54            Err(io_error(ErrorKind::TransportError(err.clone())))
55        } else {
56            Ok(())
57        }
58    }
59
60    fn enter_span(&self) -> EnteredSpan {
61        self.span.clone().entered()
62    }
63
64    fn shutdown(&self) -> Option<impl Future<Output = ()>> {
65        if self.check_transport_error().is_ok() {
66            let mut state = self.state.lock();
67            if !state.closed {
68                state.closed = true;
69
70                let transport = self.transport.clone();
71                let msg =
72                    AdbMessage::new(AdbCommand::CLSE, self.local_id, state.remote_id, Vec::new());
73                return Some(
74                    async move {
75                        if let Err(err) = transport.write_message(msg).await {
76                            warn!(?err, "failed to close socket");
77                        }
78                    }
79                    .instrument(self.span.clone()),
80                );
81            }
82        }
83
84        None
85    }
86
87    pub fn handle_message(&self, msg: AdbMessage) -> Result<bool> {
88        let _span = self.enter_span();
89
90        let mut state = self.state.lock();
91
92        match msg.header.command {
93            AdbCommand::OKAY => {
94                let avail = u32::from_le_bytes(
95                    msg.payload.try_into().map_err(|_| (ErrorKind::Other, "no avail in OKAY"))?,
96                );
97
98                if state.remote_id == 0 {
99                    state.initial_acknowledgment = avail as isize;
100                    state.remote_id = msg.header.arg0;
101
102                    self.span.record("remote_id", state.remote_id);
103                }
104
105                debug!(acknowledged = state.acknowledged, avail);
106                state.acknowledged += avail as isize;
107
108                self.write_waker.notify();
109            }
110            AdbCommand::WRTE => {
111                state.read_buffer.push_back(Cursor::new(msg.payload));
112                self.read_waker.notify();
113            }
114            AdbCommand::CLSE => {
115                state.closed = true;
116                self.write_waker.notify();
117                self.read_waker.notify();
118            }
119            _ => unreachable!("other messages should not end up here"),
120        }
121
122        Ok(state.closed)
123    }
124}
125
126impl Drop for SocketBackend {
127    fn drop(&mut self) {
128        let _span = self.enter_span();
129
130        if let Some(fut) = self.shutdown() {
131            warn!("dropping non closed socket");
132            spawn(fut);
133        } else {
134            debug!("dropping socket backend");
135        }
136    }
137}
138
139pub struct Socket {
140    pub(crate) inner: Arc<SocketBackend>,
141    pub(crate) read_waker: WakeSink,
142    pub(crate) write_waker: WakeSink,
143}
144
145impl Socket {
146    pub(crate) fn new(
147        transport: Arc<TransportBackend>, local_id: u32, remote_id: u32, acknowledged: usize,
148    ) -> Self {
149        let read_waker = WakeSink::new();
150        let write_waker = WakeSink::new();
151
152        Self {
153            inner: SocketBackend {
154                transport,
155                local_id,
156
157                span: info_span!(
158                    "socket",
159                    local_id,
160                    remote_id = if remote_id != 0 { Some(remote_id) } else { None }
161                ),
162
163                read_waker: read_waker.source(),
164                write_waker: write_waker.source(),
165
166                state: State {
167                    remote_id,
168                    closed: false,
169
170                    read_fut: MaybeDone::empty(),
171                    read_buffer: VecDeque::new(),
172
173                    initial_acknowledgment: acknowledged as isize,
174                    acknowledged: acknowledged as isize,
175                    write_fut: MaybeDone::empty(),
176                    write_throttle: interval(Duration::from_millis(10)),
177                }
178                .into(),
179            }
180            .into(),
181
182            read_waker,
183            write_waker,
184        }
185    }
186}
187
188impl Drop for Socket {
189    fn drop(&mut self) {
190        debug!(parent: &self.inner.span, "dropping socket");
191        self.inner.transport.sockets.remove(&self.inner.local_id);
192    }
193}
194
195impl AsyncRead for Socket {
196    fn poll_read(
197        mut self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &mut ReadBuf<'_>,
198    ) -> Poll<io::Result<()>> {
199        let _span = self.inner.enter_span();
200        self.read_waker.register(cx.waker());
201
202        self.inner.check_transport_error()?;
203
204        let mut state = self.inner.state.lock();
205
206        ready!(state.read_fut.poll(cx))?;
207
208        let filled = buf.filled().len();
209        while buf.remaining() > 0 {
210            let Some(front) = state.read_buffer.front_mut() else {
211                break;
212            };
213
214            let _ = Pin::new(&mut *front).poll_read(cx, buf);
215
216            if front.position() == front.get_ref().len() as u64 {
217                state.read_buffer.pop_front();
218            }
219        }
220
221        if state.closed {
222            return Poll::Ready(Ok(()));
223        }
224
225        let filled = buf.filled().len() - filled;
226        if filled == 0 {
227            return Poll::Pending;
228        }
229
230        debug!(acknowledge = filled);
231
232        let msg = AdbMessage::new(
233            AdbCommand::OKAY,
234            self.inner.local_id,
235            state.remote_id,
236            (filled as u32).to_le_bytes().into(),
237        );
238
239        let transport = self.inner.transport.clone();
240        let _ = state
241            .read_fut
242            .set(async move {
243                transport
244                    .write_message(msg)
245                    .await
246                    .map_err(|e| io_error((ErrorKind::msg("OKAY"), e)))
247            })
248            .poll(cx)?;
249
250        Poll::Ready(Ok(()))
251    }
252}
253
254impl AsyncWrite for Socket {
255    fn poll_write(
256        mut self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &[u8],
257    ) -> Poll<io::Result<usize>> {
258        let _span = self.inner.enter_span();
259        self.write_waker.register(cx.waker());
260
261        self.inner.check_transport_error()?;
262
263        let mut state = self.inner.state.lock();
264
265        if state.closed {
266            Err(io_error(ErrorKind::Closed))?;
267        }
268
269        ready!(state.write_fut.poll(cx))?;
270
271        if state.acknowledged <= -state.initial_acknowledgment {
272            return Poll::Pending;
273        }
274
275        if state.acknowledged <= 0 {
276            ready!(state.write_throttle.poll_tick(cx));
277        }
278
279        let len = buf
280            .len()
281            .min((state.acknowledged + state.initial_acknowledgment) as usize)
282            .min(MAX_PAYLOAD as usize);
283
284        state.acknowledged -= len as isize;
285        debug!(write = len, acknowledged = state.acknowledged);
286
287        let msg = AdbMessage::new(
288            AdbCommand::WRTE,
289            self.inner.local_id,
290            state.remote_id,
291            buf[..len].to_vec(),
292        );
293
294        let transport = self.inner.transport.clone();
295        let _ = state
296            .write_fut
297            .set(async move {
298                transport
299                    .write_message(msg)
300                    .await
301                    .map_err(|e| io_error((ErrorKind::msg("WRTE"), e)))
302            })
303            .poll(cx)?;
304
305        Poll::Ready(Ok(len))
306    }
307
308    fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
309        let _span = self.inner.enter_span();
310        self.write_waker.register(cx.waker());
311
312        self.inner.check_transport_error()?;
313
314        let mut state = self.inner.state.lock();
315        ready!(state.write_fut.poll(cx))?;
316
317        if state.acknowledged == state.initial_acknowledgment {
318            Poll::Ready(Ok(()))
319        } else {
320            if state.closed {
321                Err(io_error((ErrorKind::msg("flush"), Error::from(ErrorKind::Closed))))?;
322            }
323
324            Poll::Pending
325        }
326    }
327
328    fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
329        let _span = self.inner.enter_span();
330        info!("shutdown");
331
332        ready!(self.as_mut().poll_flush(cx))?;
333
334        if let Some(fut) = self.inner.shutdown() {
335            self.inner
336                .state
337                .lock()
338                .write_fut
339                .set(async move {
340                    fut.await;
341                    Ok(())
342                })
343                .poll(cx)
344        } else {
345            debug!("shutdown completed");
346            Poll::Ready(Ok(()))
347        }
348    }
349}