1use async_trait::async_trait;
20use bytes::Bytes;
21use futures::{SinkExt, TryFutureExt};
22use snap::raw::Decoder;
23use std::backtrace::Backtrace;
24use std::cell::RefCell;
25use std::collections::HashMap;
26use std::io::empty;
27use std::net::SocketAddr;
28use std::pin::pin;
29use std::sync::atomic::{AtomicBool, AtomicU32, Ordering};
30use std::sync::Arc;
31use std::thread::spawn;
32use std::{env, mem};
33use tokio::io::{AsyncRead, AsyncWrite, Join, ReadHalf, WriteHalf};
34use tokio::select;
35use tokio::sync::mpsc::unbounded_channel;
36use tokio::sync::mpsc::{Receiver, Sender, UnboundedReceiver, UnboundedSender};
37use tokio::sync::{mpsc, oneshot, Mutex, MutexGuard, RwLock};
38use tokio::task::JoinHandle;
39use tokio_stream::StreamExt;
40use tokio_util::codec::{FramedRead, FramedWrite};
41use tokio_util::sync::{CancellationToken, DropGuard};
42use tracing::{debug, error, info, trace, warn};
43use uuid::Uuid;
44
45use crate::memdx::client_response::ClientResponse;
46use crate::memdx::codec::KeyValueCodec;
47use crate::memdx::connection::{ConnectionType, Stream};
48use crate::memdx::datatype::DataTypeFlag;
49use crate::memdx::dispatcher::{
50 Dispatcher, DispatcherOptions, OnReadLoopCloseHandler, OrphanResponseHandler,
51 UnsolicitedPacketHandler,
52};
53use crate::memdx::error;
54use crate::memdx::error::{CancellationErrorKind, Error};
55use crate::memdx::hello_feature::HelloFeature::DataType;
56use crate::memdx::magic::Magic;
57use crate::memdx::opcode::OpCode;
58use crate::memdx::packet::{RequestPacket, ResponsePacket};
59use crate::memdx::pendingop::ClientPendingOp;
60use crate::memdx::subdoc::SubdocRequestInfo;
61use crate::orphan_reporter::OrphanContext;
62
63pub(crate) type ResponseSender = Sender<error::Result<ClientResponse>>;
64pub(crate) type OpaqueMap = HashMap<u32, SenderContext>;
65
66#[derive(Debug, Clone)]
67pub struct ResponseContext {
68 pub cas: Option<u64>,
69 pub subdoc_info: Option<SubdocRequestInfo>,
70 pub scope_name: Option<String>,
71 pub collection_name: Option<String>,
72}
73
74#[derive(Debug, Clone)]
75pub(crate) struct SenderContext {
76 pub sender: ResponseSender,
77 pub is_persistent: bool,
78 pub context: Option<ResponseContext>,
79}
80
81struct ReadLoopOptions {
82 pub client_id: String,
83 pub unsolicited_packet_handler: UnsolicitedPacketHandler,
84 pub orphan_handler: Option<OrphanResponseHandler>,
85 pub on_read_close_handler: OnReadLoopCloseHandler,
86 pub on_close_cancel: CancellationToken,
87 pub disable_decompression: bool,
88 pub local_addr: SocketAddr,
89 pub peer_addr: SocketAddr,
90 pub closed: Arc<AtomicBool>,
91}
92
93#[derive(Debug)]
94struct ClientReadHandle {
95 read_handle: JoinHandle<()>,
96}
97
98impl ClientReadHandle {
99 pub async fn await_completion(&mut self) {
100 (&mut self.read_handle).await.unwrap_or_default()
101 }
102}
103
104#[derive(Debug)]
105pub struct Client {
106 current_opaque: AtomicU32,
107 opaque_map: Arc<std::sync::Mutex<OpaqueMap>>,
108
109 client_id: String,
110
111 writer: Mutex<FramedWrite<WriteHalf<Box<dyn Stream>>, KeyValueCodec>>,
112 on_close_cancel: DropGuard,
113
114 local_addr: SocketAddr,
115 peer_addr: SocketAddr,
116
117 closed: Arc<AtomicBool>,
118}
119
120impl Client {
121 fn register_handler(&self, response_context: SenderContext) -> u32 {
122 let mut map = self.opaque_map.lock().unwrap();
123
124 let opaque = self.current_opaque.fetch_add(1, Ordering::SeqCst);
125
126 map.insert(opaque, response_context);
127
128 opaque
129 }
130
131 async fn drain_opaque_map(opaque_map: Arc<std::sync::Mutex<OpaqueMap>>) {
132 let mut senders = vec![];
133 {
134 let mut guard = opaque_map.lock().unwrap();
135 guard.drain().for_each(|(_, v)| {
136 senders.push(v);
137 });
138 }
139
140 for sender in senders {
141 sender
142 .sender
143 .send(Err(Error::new_cancelled_error(
144 CancellationErrorKind::ClosedInFlight,
145 )))
146 .await
147 .unwrap_or_default();
148 }
149 }
150
151 async fn on_read_loop_close(
152 client_id: &str,
153 stream: FramedRead<ReadHalf<Box<dyn Stream>>, KeyValueCodec>,
154 opaque_map: Arc<std::sync::Mutex<OpaqueMap>>,
155 on_read_loop_close: OnReadLoopCloseHandler,
156 graceful: bool,
157 ) {
158 drop(stream);
159
160 Self::drain_opaque_map(opaque_map).await;
161
162 if on_read_loop_close.send(()).is_err() && !graceful {
165 warn!("{} failed to notify read loop closure", &client_id);
166 }
167
168 debug!("{client_id} read loop shut down");
169 }
170
171 async fn read_loop(
172 mut stream: FramedRead<ReadHalf<Box<dyn Stream>>, KeyValueCodec>,
173 opaque_map: Arc<std::sync::Mutex<OpaqueMap>>,
174 mut opts: ReadLoopOptions,
175 ) {
176 loop {
177 select! {
178 (_) = opts.on_close_cancel.cancelled() => {
179 Self::on_read_loop_close(&opts.client_id, stream, opaque_map, opts.on_read_close_handler, true).await;
180 return;
181 },
182 (next) = stream.next() => {
183 match next {
184 Some(input) => {
185 match input {
186 Ok(mut packet) => {
187 if packet.magic == Magic::ServerReq {
188
189 trace!(
190 "Handling server request on {}. Opcode={}",
191 opts.client_id,
192 packet.op_code,
193 );
194
195 (opts.unsolicited_packet_handler)(packet).await;
196 continue;
197 }
198
199 trace!(
200 "Resolving response on {}. Opcode={}. Opaque={}. Status={}",
201 opts.client_id,
202 packet.op_code,
203 packet.opaque,
204 packet.status,
205 );
206
207 let opaque = packet.opaque;
208
209 let requests: Arc<std::sync::Mutex<OpaqueMap>> = Arc::clone(&opaque_map);
210 let context = {
211 let mut map = requests.lock().unwrap();
212 map.remove(&opaque)
213 };
214
215 if let Some(mut context) = context {
216 let sender = &context.sender;
217
218 if let Some(value) = &packet.value {
219 if !opts.disable_decompression && (packet.datatype & u8::from(DataTypeFlag::Compressed) != 0) {
220 let mut decoder = Decoder::new();
221 let new_value = match decoder
222 .decompress_vec(value)
223 {
224 Ok(v) => v,
225 Err(e) => {
226 match sender.send(Err(Error::new_decompression_error().with(e))).await{
227 Ok(_) => {}
228 Err(e) => {
229 debug!("Sending response to caller failed: {e}");
230 }
231 };
232 continue;
233 }
234 };
235
236 packet.datatype &= !u8::from(DataTypeFlag::Compressed);
237 packet.value = Some(Bytes::from(new_value));
238 }
239 }
240
241 if context.is_persistent {
242 {
243 let mut map = requests.lock().unwrap();
244 map.insert(opaque, context.clone());
245 }
246 }
247
248 let resp = ClientResponse::new(packet, context.context);
249 match sender.send(Ok(resp)).await {
250 Ok(_) => {}
251 Err(e) => {
252 debug!("Sending response to caller failed: {e}");
253 let graceful = opts.closed.load(Ordering::SeqCst) || opts.on_close_cancel.is_cancelled();
254 Self::on_read_loop_close(&opts.client_id, stream, opaque_map, opts.on_read_close_handler, graceful).await;
255 return;
256 }
257 };
258 } else if let Some(ref orphan_handler) = opts.orphan_handler {
259 orphan_handler(
260 packet,
261 OrphanContext {
262 client_id: opts.client_id.clone(),
263 local_addr: opts.local_addr,
264 peer_addr: opts.peer_addr,
265 },
266 );
267 }
268 drop(requests);
269 }
270 Err(e) => {
271 warn!("{} failed to read frame {}", opts.client_id, e);
272 let graceful = opts.closed.load(Ordering::SeqCst) || opts.on_close_cancel.is_cancelled();
273 Self::on_read_loop_close(&opts.client_id, stream, opaque_map, opts.on_read_close_handler, graceful).await;
274 return;
275 }
276 }
277 }
278 None => {
279 let graceful = opts.closed.load(Ordering::SeqCst) || opts.on_close_cancel.is_cancelled();
280 Self::on_read_loop_close(&opts.client_id, stream, opaque_map, opts.on_read_close_handler, graceful).await;
281 return;
282 }
283 }
284 }
285 }
286 }
287 }
288
289 fn split_stream<StreamType: AsyncRead + AsyncWrite + Send + Unpin>(
290 stream: StreamType,
291 ) -> (ReadHalf<StreamType>, WriteHalf<StreamType>) {
292 tokio::io::split(stream)
293 }
294}
295
296struct DispatchOpaqueGuard {
300 opaque: u32,
301 opaque_map: Option<Arc<std::sync::Mutex<OpaqueMap>>>,
302}
303
304impl DispatchOpaqueGuard {
305 fn new(opaque: u32, opaque_map: Arc<std::sync::Mutex<OpaqueMap>>) -> Self {
306 Self {
307 opaque,
308 opaque_map: Some(opaque_map),
309 }
310 }
311
312 fn disarm(&mut self) {
314 self.opaque_map = None;
315 }
316}
317
318impl Drop for DispatchOpaqueGuard {
319 fn drop(&mut self) {
320 if let Some(opaque_map) = self.opaque_map.take() {
321 let mut map = opaque_map.lock().unwrap();
322 map.remove(&self.opaque);
323 }
324 }
325}
326
327#[async_trait]
328impl Dispatcher for Client {
329 fn new(conn: ConnectionType, opts: DispatcherOptions) -> Self {
330 let local_addr = *conn.local_addr();
331 let peer_addr = *conn.peer_addr();
332
333 let (r, w) = tokio::io::split(conn.into_inner());
334
335 let codec = KeyValueCodec::default();
336 let reader = FramedRead::new(r, codec);
337 let writer = FramedWrite::new(w, codec);
338
339 let cancel_token = CancellationToken::new();
340 let cancel_child = cancel_token.child_token();
341 let cancel_guard = cancel_token.drop_guard();
342
343 let opaque_map = Arc::new(std::sync::Mutex::new(OpaqueMap::default()));
344
345 let read_opaque_map = Arc::clone(&opaque_map);
346 let read_uuid = opts.id.clone();
347
348 let closed = Arc::new(AtomicBool::new(false));
349 let read_closed = Arc::clone(&closed);
350
351 tokio::spawn(async move {
352 Client::read_loop(
353 reader,
354 read_opaque_map,
355 ReadLoopOptions {
356 client_id: read_uuid,
357 unsolicited_packet_handler: opts.unsolicited_packet_handler,
358 orphan_handler: opts.orphan_handler,
359 on_read_close_handler: opts.on_read_close_tx,
360 on_close_cancel: cancel_child,
361 disable_decompression: opts.disable_decompression,
362 local_addr,
363 peer_addr,
364 closed: read_closed,
365 },
366 )
367 .await;
368 });
369
370 Self {
371 current_opaque: AtomicU32::new(1),
372 opaque_map,
373 client_id: opts.id,
374
375 on_close_cancel: cancel_guard,
376
377 writer: Mutex::new(writer),
378
379 local_addr,
380 peer_addr,
381
382 closed,
383 }
384 }
385
386 async fn dispatch<'a>(
387 &self,
388 mut packet: RequestPacket<'a>,
389 is_persistent: bool,
390 response_context: Option<ResponseContext>,
391 ) -> error::Result<ClientPendingOp> {
392 let (response_tx, response_rx) = mpsc::channel(1);
393
394 let opaque = self.register_handler(SenderContext {
395 sender: response_tx,
396 is_persistent,
397 context: response_context,
398 });
399 packet.opaque = Some(opaque);
400 let op_code = packet.op_code;
401
402 let mut opaque_guard = DispatchOpaqueGuard::new(opaque, self.opaque_map.clone());
406
407 trace!(
408 "Writing request on {}. Opcode={}. Opaque={}",
409 &self.client_id,
410 packet.op_code,
411 opaque,
412 );
413
414 let mut writer = self.writer.lock().await;
415 match writer.send(packet).await {
416 Ok(_) => {
417 opaque_guard.disarm();
419 Ok(ClientPendingOp::new(
420 opaque,
421 self.opaque_map.clone(),
422 response_rx,
423 is_persistent,
424 ))
425 }
426 Err(e) => {
427 debug!(
428 "{} failed to write packet {} {} {}",
429 self.client_id, opaque, op_code, e
430 );
431
432 Err(Error::new_dispatch_error(opaque, op_code, Box::new(e)))
435 }
436 }
437 }
438
439 async fn close(&self) -> error::Result<()> {
440 if self.closed.swap(true, Ordering::SeqCst) {
441 return Ok(());
442 }
443
444 info!("Closing client {}", self.client_id);
445
446 let mut close_err = None;
447 let mut writer = self.writer.lock().await;
448 match writer.close().await {
449 Ok(_) => {}
450 Err(e) => {
451 close_err = Some(e);
452 }
453 };
454
455 Self::drain_opaque_map(self.opaque_map.clone()).await;
456
457 if let Some(e) = close_err {
458 return Err(Error::new_close_error(e.to_string(), Box::new(e)));
459 }
460
461 Ok(())
462 }
463}
464
465impl Drop for Client {
466 fn drop(&mut self) {
467 info!("Dropping client {}", self.client_id);
468 }
469}