Skip to main content

couchbase_core/memdx/
client.rs

1/*
2 *
3 *  * Copyright (c) 2025 Couchbase, Inc.
4 *  *
5 *  * Licensed under the Apache License, Version 2.0 (the "License");
6 *  * you may not use this file except in compliance with the License.
7 *  * You may obtain a copy of the License at
8 *  *
9 *  *    http://www.apache.org/licenses/LICENSE-2.0
10 *  *
11 *  * Unless required by applicable law or agreed to in writing, software
12 *  * distributed under the License is distributed on an "AS IS" BASIS,
13 *  * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
14 *  * See the License for the specific language governing permissions and
15 *  * limitations under the License.
16 *
17 */
18
19use 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 the client is being shut down, the receiver may have already been dropped
163        // (e.g. the owning kvclient is tearing down too), which is expected and not an error.
164        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
296/// A guard that removes an opaque entry from the map when dropped, unless disarmed.
297/// This prevents opaque map leaks if the dispatch future is cancelled/dropped at an
298/// `.await` point before a `ClientPendingOp` is created to take over cleanup.
299struct 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    /// Disarm the guard so that dropping it will not remove the opaque entry.
313    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        // Create a guard that will remove the opaque entry from the map if the future is
403        // dropped before we successfully construct a ClientPendingOp (which takes over
404        // cleanup responsibility via its own Drop impl).
405        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                // Disarm the guard — the ClientPendingOp now owns cleanup responsibility.
418                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                // opaque_guard will remove the entry from the opaque map when dropped.
433
434                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}