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}