Skip to main content

ruvio_client/
client.rs

1use std::net::{IpAddr, SocketAddr};
2
3use std::sync::{Arc, Mutex};
4
5use std::time::{Duration, SystemTime, UNIX_EPOCH};
6
7use crate::cluster::{ClusterRouter, Connection, Discovery, throw_error};
8
9use crate::error::Error;
10use crate::value::RespValue;
11
12/// How `SET` treats a key that already exists or is missing.
13#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
14pub enum SetMode {
15    /// Write the key either way.
16    #[default]
17    Always,
18    /// Write only when the key is absent (`NX`).
19    IfNotExists,
20    /// Write only when the key exists (`XX`).
21    IfExists,
22}
23
24/// Where to connect, and an optional password sent as `AUTH`.
25#[derive(Debug, Clone)]
26pub struct ClientOptions {
27    /// Server host. Defaults to loopback.
28    pub host: String,
29    /// RESP port. Defaults to 6379.
30    pub port: u16,
31    /// ACL user. When set together with [`Self::password`], the client sends `AUTH user password`.
32    pub username: Option<String>,
33    /// When set, the client sends `AUTH` before [`Client::connect_with`] returns.
34    pub password: Option<String>,
35    /// Logical database. When not 0, the client sends `SELECT` after `AUTH`.
36    pub database: u32,
37    /// How long TCP connect may take.
38    pub connect_timeout: Duration,
39    /// When true (the default), connect sends `CLUSTER SLOTS`. A slot map turns
40    /// the client into a router; a disabled-cluster error keeps one socket.
41    pub discover_cluster: bool,
42    /// When true (the default), a broken socket is opened again before the next command.
43    /// The command that discovers the drop still fails. Set this to false and call
44    /// [`Client::reconnect`] on your own schedule.
45    pub reconnect: bool,
46    /// How many connect attempts a reconnect makes before giving up. One is a single try.
47    pub max_reconnect_attempts: u32,
48    /// First automatic wait after a failed reconnect attempt.
49    pub reconnect_base_delay: Duration,
50    /// Ceiling for the automatic exponential wait.
51    pub reconnect_max_delay: Duration,
52}
53
54impl Default for ClientOptions {
55    fn default() -> Self {
56        Self {
57            host: "127.0.0.1".to_owned(),
58            port: 6379,
59            username: None,
60            password: None,
61            database: 0,
62            connect_timeout: Duration::from_secs(5),
63            discover_cluster: true,
64            reconnect: true,
65            max_reconnect_attempts: 8,
66            reconnect_base_delay: Duration::from_millis(100),
67            reconnect_max_delay: Duration::from_secs(2),
68        }
69    }
70}
71
72/// A command or subscriber socket dropped, or came back.
73#[derive(Debug, Clone, PartialEq, Eq)]
74pub struct ConnectionNotice {
75    pub host: String,
76    pub port: u16,
77    /// True when the socket is the dedicated Pub/Sub connection.
78    pub subscriber: bool,
79}
80
81/// A Ruvio client. Against a single server it is one TCP connection. Against a
82/// sharded server it opens one connection per shard and sends each command to
83/// the shard that owns its key.
84///
85/// Subscribe commands open a second TCP socket, the same way `Ruvio.Client`
86/// does. `GET`, `SET`, `PING`, and `PUBLISH` stay on the command sockets.
87///
88/// The client speaks RESP2 and does not send `CLIENT SETINFO` or `HELLO`.
89pub struct Client {
90    transport: Transport,
91    database: u32,
92    options: ClientOptions,
93    dedicated_subscriber: bool,
94    pubsub: Option<Box<Client>>,
95    on_lost: Arc<Mutex<Option<Arc<dyn Fn(ConnectionNotice) + Send + Sync>>>>,
96    on_restored: Arc<Mutex<Option<Arc<dyn Fn(ConnectionNotice) + Send + Sync>>>>,
97    reconnect_delay: Option<Arc<dyn Fn(u32) -> Option<Duration> + Send + Sync>>,
98    subscriptions: Vec<Vec<String>>,
99    read_timeout: Option<Option<Duration>>,
100}
101
102enum Transport {
103    Single(Connection),
104    Cluster(Box<ClusterRouter>),
105}
106
107impl Client {
108    /// Opens a connection. Sends `CLUSTER SLOTS` so a sharded server is routed automatically.
109    pub fn connect(host: impl Into<String>, port: u16) -> Result<Self, Error> {
110        Self::connect_database(host, port, 0)
111    }
112
113    /// Opens a connection bound to one logical database. Sends `SELECT` when `database` is not 0.
114    ///
115    /// The returned client is the database handle: every command on it runs in `database`.
116    pub fn connect_database(
117        host: impl Into<String>,
118        port: u16,
119        database: u32,
120    ) -> Result<Self, Error> {
121        Self::connect_with(ClientOptions {
122            host: host.into(),
123            port,
124            database,
125            ..ClientOptions::default()
126        })
127    }
128
129    /// Opens a connection to an IP address and port. Sends `SELECT` when `database` is not 0.
130    pub fn connect_ip(ip: IpAddr, port: u16, database: u32) -> Result<Self, Error> {
131        Self::connect_addr(SocketAddr::new(ip, port), database)
132    }
133
134    /// Opens a connection to a socket address. Sends `SELECT` when `database` is not 0.
135    pub fn connect_addr(address: SocketAddr, database: u32) -> Result<Self, Error> {
136        Self::connect_database(address.ip().to_string(), address.port(), database)
137    }
138
139    /// Opens a connection. When [`ClientOptions::password`] is set, sends `AUTH`.
140    /// When [`ClientOptions::discover_cluster`] is on, sends `CLUSTER SLOTS`.
141    /// When [`ClientOptions::database`] is not 0, a standalone server then gets `SELECT`.
142    pub fn connect_with(options: ClientOptions) -> Result<Self, Error> {
143        if options.host.is_empty() {
144            return Err(Error::Protocol("host is required".to_owned()));
145        }
146
147        let seed = Connection::open(&options.host, options.port, &options)?;
148
149        if !options.discover_cluster {
150            return finish_standalone(seed, options);
151        }
152
153        match ClusterRouter::discover(seed, options.clone())? {
154            Discovery::Cluster(_) if options.database != 0 => Err(Error::Protocol(
155                "SELECT is not supported in cluster mode".to_owned(),
156            )),
157            Discovery::Cluster(mut router) => {
158                let on_lost = Arc::new(Mutex::new(None));
159                let on_restored = Arc::new(Mutex::new(None));
160
161                router.share_hooks(Arc::clone(&on_lost), Arc::clone(&on_restored));
162
163                Ok(Self {
164                    transport: Transport::Cluster(router),
165                    database: 0,
166                    options,
167                    dedicated_subscriber: false,
168                    pubsub: None,
169                    read_timeout: None,
170                    on_lost,
171                    on_restored,
172                    reconnect_delay: None,
173                    subscriptions: Vec::new(),
174                })
175            }
176
177            Discovery::Standalone(seed) => finish_standalone(seed, options),
178        }
179    }
180
181    /// True when the server reported a slot map and commands are routed by slot.
182    pub fn is_cluster(&self) -> bool {
183        matches!(self.transport, Transport::Cluster(_))
184    }
185
186    /// How many shards the client routes to. 1 for a single server.
187    pub fn shard_count(&self) -> usize {
188        match &self.transport {
189            Transport::Single(_) => 1,
190            Transport::Cluster(router) => router.node_count(),
191        }
192    }
193
194    /// Sends `PING` and returns the simple-string reply.
195    pub fn ping(&mut self) -> Result<String, Error> {
196        read_ping(self.run(["PING"])?)
197    }
198
199    /// Sends `PING message` and returns the echoed message.
200    pub fn ping_message(&mut self, message: &str) -> Result<String, Error> {
201        read_ping(self.run(["PING", message])?)
202    }
203
204    /// The logical database this connection last selected.
205    pub fn database(&self) -> u32 {
206        self.database
207    }
208
209    /// Opens a second socket for Pub/Sub. `subscribe` on this client already
210    /// does that automatically; use this when you want the socket in hand.
211    pub fn subscriber(&self) -> Result<Self, Error> {
212        self.open_pubsub(&["SUBSCRIBE"])
213    }
214
215    /// Sends `SELECT`. Indexes run from 0 to the server's `databases` setting minus one (16 by default).
216    pub fn select(&mut self, database: u32) -> Result<(), Error> {
217        self.run(["SELECT".to_owned(), database.to_string()])?;
218        self.database = database;
219
220        Ok(())
221    }
222
223    /// Sends `SELECT` and returns the same connection.
224    ///
225    /// This is not a second TCP socket: every later command on the client uses the new database.
226    pub fn into_database(mut self, database: u32) -> Result<Self, Error> {
227        self.select(database)?;
228
229        Ok(self)
230    }
231
232    /// Sets how long a read may block. `None` waits forever.
233    ///
234    /// A timed-out read returns [`Error::Io`], and the reply may still arrive later,
235    /// so drop the client unless it was only waiting in [`Self::read_message`].
236    pub fn set_read_timeout(&mut self, timeout: Option<Duration>) -> Result<(), Error> {
237        self.read_timeout = Some(timeout);
238
239        match &mut self.transport {
240            Transport::Single(connection) => connection.set_read_timeout(timeout)?,
241            Transport::Cluster(router) => router.set_read_timeout(timeout)?,
242        }
243
244        if let Some(pubsub) = &mut self.pubsub {
245            pubsub.set_read_timeout(timeout)?;
246        }
247
248        Ok(())
249    }
250
251    /// Sends `GET` and returns the raw bulk, or `None` when the key is missing.
252    pub fn get(&mut self, key: &str) -> Result<Option<Vec<u8>>, Error> {
253        match self.execute(&["GET", key])? {
254            RespValue::Null => Ok(None),
255            RespValue::Bulk(bytes) => Ok(Some(bytes)),
256            other => Err(Error::Protocol(format!("GET returned {other:?}"))),
257        }
258    }
259
260    /// Sends `GET` and decodes the bulk as UTF-8. `None` when the key is missing.
261    pub fn get_string(&mut self, key: &str) -> Result<Option<String>, Error> {
262        match self.get(key)? {
263            None => Ok(None),
264            Some(bytes) => RespValue::Bulk(bytes).as_string(),
265        }
266    }
267
268    /// Sends `SET`. Returns false when `NX` or `XX` skips the write.
269    pub fn set(&mut self, key: &str, value: &str) -> Result<bool, Error> {
270        self.set_with(key, value, None, SetMode::Always)
271    }
272
273    /// Sends `SET` with an optional millisecond expiry and `NX` or `XX`.
274    pub fn set_with(
275        &mut self,
276        key: &str,
277        value: &str,
278        expiry: Option<Duration>,
279        mode: SetMode,
280    ) -> Result<bool, Error> {
281        let expiry = expiry.map(|ttl| ["PX".to_owned(), duration_millis(ttl).to_string()]);
282        let arguments = set_arguments(key, value, expiry.as_ref().map(|pair| &pair[..]), mode);
283
284        Ok(!matches!(self.run(arguments)?, RespValue::Null))
285    }
286
287    /// Sends `SET ... PXAT`: the key expires at `expires_at`.
288    pub fn set_expires_at(
289        &mut self,
290        key: &str,
291        value: &str,
292        expires_at: SystemTime,
293        mode: SetMode,
294    ) -> Result<bool, Error> {
295        let expiry = ["PXAT".to_owned(), unix_millis(expires_at)?.to_string()];
296        let arguments = set_arguments(key, value, Some(&expiry), mode);
297
298        Ok(!matches!(self.run(arguments)?, RespValue::Null))
299    }
300
301    /// Sends `SET ... KEEPTTL`: the key keeps its current expiry.
302    pub fn set_keep_ttl(&mut self, key: &str, value: &str, mode: SetMode) -> Result<bool, Error> {
303        let expiry = ["KEEPTTL".to_owned()];
304        let arguments = set_arguments(key, value, Some(&expiry), mode);
305
306        Ok(!matches!(self.run(arguments)?, RespValue::Null))
307    }
308
309    /// Sends `SET ... GET` and returns the previous string, or `None` when the key was missing.
310    /// With `NX` or `XX` the previous value is returned even when the write is skipped.
311    pub fn set_and_get(
312        &mut self,
313        key: &str,
314        value: &str,
315        expiry: Option<Duration>,
316        mode: SetMode,
317    ) -> Result<Option<String>, Error> {
318        let expiry = expiry.map(|ttl| ["PX".to_owned(), duration_millis(ttl).to_string()]);
319        let mut arguments = set_arguments(key, value, expiry.as_ref().map(|pair| &pair[..]), mode);
320
321        arguments.push("GET".to_owned());
322        self.bulk_or_null(arguments)
323    }
324
325    /// Sends `INCR`.
326    pub fn incr(&mut self, key: &str) -> Result<i64, Error> {
327        self.integer(["INCR", key])
328    }
329
330    /// Sends `EXPIRE`, or `PEXPIRE` when `ttl` is not whole seconds. False when the key does not exist.
331    pub fn expire(&mut self, key: &str, ttl: Duration) -> Result<bool, Error> {
332        if ttl.subsec_nanos() == 0 {
333            return self.flag([
334                "EXPIRE".to_owned(),
335                key.to_owned(),
336                duration_secs(ttl).to_string(),
337            ]);
338        }
339
340        self.flag([
341            "PEXPIRE".to_owned(),
342            key.to_owned(),
343            duration_millis(ttl).to_string(),
344        ])
345    }
346
347    /// Sends `DEL` for one key.
348    pub fn del_key(&mut self, key: &str) -> Result<i64, Error> {
349        self.del(&[key])
350    }
351
352    /// Sends `DEL` and returns how many keys were removed.
353    pub fn del(&mut self, keys: &[&str]) -> Result<i64, Error> {
354        if keys.is_empty() {
355            return Err(Error::Protocol("DEL needs at least one key".to_owned()));
356        }
357
358        self.integer(join("DEL", keys))
359    }
360
361    /// Reads the next reply without sending a command. Use this after a subscribe
362    /// for `message`, `pmessage`, and `smessage` pushes.
363    pub fn read_message(&mut self) -> Result<RespValue, Error> {
364        if let Some(pubsub) = &mut self.pubsub {
365            return pubsub.read_message();
366        }
367
368        self.ensure_open()?;
369
370        let reply = match &mut self.transport {
371            Transport::Single(connection) => connection.read()?,
372            Transport::Cluster(router) => router.read_message()?,
373        };
374
375        throw_error(reply)
376    }
377
378    /// Called when a command or subscriber socket drops.
379    pub fn on_connection_lost<F>(&mut self, hook: F)
380    where
381        F: Fn(ConnectionNotice) + Send + Sync + 'static,
382    {
383        *self.on_lost.lock().expect("connection hook") = Some(Arc::new(hook));
384    }
385
386    /// Called after a dropped socket is open again.
387    pub fn on_connection_restored<F>(&mut self, hook: F)
388    where
389        F: Fn(ConnectionNotice) + Send + Sync + 'static,
390    {
391        *self.on_restored.lock().expect("connection hook") = Some(Arc::new(hook));
392    }
393
394    /// Replaces the built-in exponential wait. The attempt starts at 1 after the first failure.
395    /// Return `None` to stop. When set, `max_reconnect_attempts` is not applied.
396    pub fn set_reconnect_delay<F>(&mut self, delay: F)
397    where
398        F: Fn(u32) -> Option<Duration> + Send + Sync + 'static,
399    {
400        self.reconnect_delay = Some(Arc::new(delay));
401
402        if let Transport::Cluster(router) = &mut self.transport {
403            router.set_reconnect_delay(Arc::clone(self.reconnect_delay.as_ref().unwrap()));
404        }
405    }
406
407    /// Opens a dropped command socket again. With [`ClientOptions::reconnect`] on, the next
408    /// command does this itself.
409    pub fn reconnect(&mut self) -> Result<(), Error> {
410        if let Transport::Cluster(router) = &mut self.transport {
411            return router.reconnect_broken();
412        }
413
414        if !self.transport_is_broken() {
415            return Ok(());
416        }
417
418        self.reopen()
419    }
420
421    /// Sends one command. A RESP error becomes [`Error::Server`] and the connection stays open.
422    pub fn execute(&mut self, arguments: &[&str]) -> Result<RespValue, Error> {
423        self.ensure_open()?;
424
425        if is_subscription_command(arguments) && !self.dedicated_subscriber {
426            return self.ensure_pubsub(arguments)?.execute(arguments);
427        }
428
429        let reply = match &mut self.transport {
430            Transport::Single(connection) => connection.send(arguments)?,
431            Transport::Cluster(router) => router.execute(arguments)?,
432        };
433
434        let reply = throw_error(reply)?;
435
436        self.remember_subscription(arguments);
437
438        Ok(reply)
439    }
440
441    /// Writes every command, then reads one reply each. Error replies stay as [`RespValue::Error`].
442    /// On a cluster the commands are grouped by shard; replies keep command order.
443    pub fn execute_many(&mut self, commands: &[&[&str]]) -> Result<Vec<RespValue>, Error> {
444        if commands.is_empty() {
445            return Err(Error::Protocol(
446                "a pipeline needs at least one command".to_owned(),
447            ));
448        }
449
450        match &mut self.transport {
451            Transport::Single(connection) => connection.send_many(commands),
452            Transport::Cluster(router) => router.execute_many(commands),
453        }
454    }
455
456    /// Writes one command and reads `replies` replies. A subscribe confirms each channel separately.
457    pub(crate) fn run_replies<I, S>(
458        &mut self,
459        arguments: I,
460        replies: usize,
461    ) -> Result<Vec<RespValue>, Error>
462    where
463        I: IntoIterator<Item = S>,
464        S: AsRef<str>,
465    {
466        let owned: Vec<S> = arguments.into_iter().collect();
467        let borrowed: Vec<&str> = owned.iter().map(AsRef::as_ref).collect();
468
469        if is_subscription_command(&borrowed) && !self.dedicated_subscriber {
470            return self.ensure_pubsub(&borrowed)?.run_replies(owned, replies);
471        }
472
473        self.ensure_open()?;
474
475        let values = match &mut self.transport {
476            Transport::Single(connection) => connection.send_replies(&borrowed, replies)?,
477            Transport::Cluster(router) => router.run_replies(&borrowed, replies)?,
478        };
479
480        if values.len() < replies {
481            return Err(Error::Protocol("missing subscribe confirmation".to_owned()));
482        }
483
484        for value in &values {
485            throw_error(value.clone())?;
486        }
487
488        self.remember_subscription(&borrowed);
489
490        Ok(values)
491    }
492
493    pub(crate) fn run<I, S>(&mut self, arguments: I) -> Result<RespValue, Error>
494    where
495        I: IntoIterator<Item = S>,
496        S: AsRef<str>,
497    {
498        let owned: Vec<S> = arguments.into_iter().collect();
499        let borrowed: Vec<&str> = owned.iter().map(AsRef::as_ref).collect();
500
501        self.execute(&borrowed)
502    }
503
504    fn ensure_pubsub(&mut self, arguments: &[&str]) -> Result<&mut Client, Error> {
505        if self.pubsub.is_none() {
506            self.pubsub = Some(Box::new(self.open_pubsub(arguments)?));
507        }
508
509        Ok(self.pubsub.as_mut().unwrap())
510    }
511
512    fn open_pubsub(&self, arguments: &[&str]) -> Result<Client, Error> {
513        let mut options = self.options.clone();
514
515        options.discover_cluster = false;
516
517        if let Transport::Cluster(router) = &self.transport {
518            let (host, port) = router.subscription_endpoint(arguments);
519
520            options.host = host;
521            options.port = port;
522        }
523
524        let mut client = Client::connect_with(options)?;
525
526        client.dedicated_subscriber = true;
527        client.on_lost = Arc::clone(&self.on_lost);
528        client.on_restored = Arc::clone(&self.on_restored);
529        client.reconnect_delay = self.reconnect_delay.clone();
530
531        if let Some(timeout) = self.read_timeout {
532            client.set_read_timeout(timeout)?;
533        }
534
535        Ok(client)
536    }
537
538    fn transport_is_broken(&self) -> bool {
539        match &self.transport {
540            Transport::Single(connection) => connection.is_broken(),
541            Transport::Cluster(_) => false,
542        }
543    }
544
545    fn ensure_open(&mut self) -> Result<(), Error> {
546        let broken = self.transport_is_broken();
547
548        if !broken {
549            return Ok(());
550        }
551
552        if !self.options.reconnect {
553            return Err(Error::Io(std::io::Error::new(
554                std::io::ErrorKind::BrokenPipe,
555                "connection is closed",
556            )));
557        }
558
559        self.reopen()
560    }
561
562    fn reopen(&mut self) -> Result<(), Error> {
563        let host = self.options.host.clone();
564        let port = self.options.port;
565
566        self.notify_lost(host.clone(), port);
567
568        let mut attempt = 0_u32;
569
570        let mut opened = loop {
571            attempt += 1;
572
573            match crate::cluster::Connection::open(&host, port, &self.options) {
574                Ok(connection) => break connection,
575                Err(error) => match self.next_delay(attempt) {
576                    Some(delay) => std::thread::sleep(delay),
577                    None => return Err(error),
578                },
579            }
580        };
581
582        if self.database != 0 {
583            let index = self.database.to_string();
584
585            throw_error(opened.send(&["SELECT", &index])?)?;
586        }
587
588        for command in &self.subscriptions {
589            let borrowed: Vec<&str> = command.iter().map(String::as_str).collect();
590
591            throw_error(opened.send(&borrowed)?)?;
592        }
593
594        self.transport = Transport::Single(opened);
595        self.notify_restored(host, port);
596
597        Ok(())
598    }
599
600    fn next_delay(&self, attempt: u32) -> Option<Duration> {
601        if !self.options.reconnect {
602            return None;
603        }
604
605        if let Some(delay) = &self.reconnect_delay {
606            return delay(attempt);
607        }
608
609        if attempt >= self.options.max_reconnect_attempts {
610            return None;
611        }
612
613        let multiplier = 2_u32.saturating_pow(attempt.saturating_sub(1));
614        let delay = self.options.reconnect_base_delay.saturating_mul(multiplier);
615
616        Some(delay.min(self.options.reconnect_max_delay))
617    }
618
619    fn notify_lost(&self, host: String, port: u16) {
620        let notice = ConnectionNotice {
621            host,
622            port,
623            subscriber: self.dedicated_subscriber,
624        };
625
626        if let Some(hook) = self.on_lost.lock().expect("connection hook").as_ref() {
627            hook(notice);
628        }
629    }
630
631    fn notify_restored(&self, host: String, port: u16) {
632        let notice = ConnectionNotice {
633            host,
634            port,
635            subscriber: self.dedicated_subscriber,
636        };
637
638        if let Some(hook) = self.on_restored.lock().expect("connection hook").as_ref() {
639            hook(notice);
640        }
641    }
642
643    fn remember_subscription(&mut self, arguments: &[&str]) {
644        if !self.dedicated_subscriber || arguments.is_empty() {
645            return;
646        }
647
648        let command = arguments[0].to_ascii_uppercase();
649        let (subscribe, dropping) = match command.as_str() {
650            "SUBSCRIBE" => ("SUBSCRIBE", false),
651            "UNSUBSCRIBE" => ("SUBSCRIBE", true),
652            "PSUBSCRIBE" => ("PSUBSCRIBE", false),
653            "PUNSUBSCRIBE" => ("PSUBSCRIBE", true),
654            "SSUBSCRIBE" => ("SSUBSCRIBE", false),
655            "SUNSUBSCRIBE" => ("SSUBSCRIBE", true),
656            _ => return,
657        };
658
659        if dropping && arguments.len() <= 1 {
660            self.subscriptions
661                .retain(|item| item.first().map(String::as_str) != Some(subscribe));
662
663            return;
664        }
665
666        if dropping {
667            for channel in &arguments[1..] {
668                self.subscriptions.retain(|item| {
669                    item.first().map(String::as_str) != Some(subscribe)
670                        || item.get(1).map(String::as_str) != Some(*channel)
671                });
672            }
673
674            return;
675        }
676
677        for channel in &arguments[1..] {
678            let entry = vec![subscribe.to_owned(), (*channel).to_owned()];
679
680            if !self.subscriptions.iter().any(|item| item == &entry) {
681                self.subscriptions.push(entry);
682            }
683        }
684    }
685
686    pub(crate) fn integer<I, S>(&mut self, arguments: I) -> Result<i64, Error>
687    where
688        I: IntoIterator<Item = S>,
689        S: AsRef<str>,
690    {
691        let (command, reply) = self.run_named(arguments)?;
692
693        match reply {
694            RespValue::Integer(value) => Ok(value),
695            other => Err(unexpected(&command, &other)),
696        }
697    }
698
699    pub(crate) fn flag<I, S>(&mut self, arguments: I) -> Result<bool, Error>
700    where
701        I: IntoIterator<Item = S>,
702        S: AsRef<str>,
703    {
704        Ok(self.integer(arguments)? > 0)
705    }
706
707    pub(crate) fn ok<I, S>(&mut self, arguments: I) -> Result<(), Error>
708    where
709        I: IntoIterator<Item = S>,
710        S: AsRef<str>,
711    {
712        self.run(arguments)?;
713
714        Ok(())
715    }
716
717    pub(crate) fn text<I, S>(&mut self, arguments: I) -> Result<String, Error>
718    where
719        I: IntoIterator<Item = S>,
720        S: AsRef<str>,
721    {
722        let (command, reply) = self.run_named(arguments)?;
723
724        reply
725            .as_string()?
726            .ok_or_else(|| Error::Protocol(format!("{command} returned a null reply")))
727    }
728
729    pub(crate) fn bulk_or_null<I, S>(&mut self, arguments: I) -> Result<Option<String>, Error>
730    where
731        I: IntoIterator<Item = S>,
732        S: AsRef<str>,
733    {
734        self.run(arguments)?.as_string()
735    }
736
737    pub(crate) fn strings<I, S>(&mut self, arguments: I) -> Result<Vec<String>, Error>
738    where
739        I: IntoIterator<Item = S>,
740        S: AsRef<str>,
741    {
742        let (command, reply) = self.run_named(arguments)?;
743
744        match reply {
745            RespValue::Null => Ok(Vec::new()),
746            RespValue::Array(items) => items
747                .iter()
748                .map(|item| {
749                    item.as_string()?
750                        .ok_or_else(|| Error::Protocol(format!("{command} returned a null bulk")))
751                })
752                .collect(),
753            other => Err(unexpected(&command, &other)),
754        }
755    }
756
757    pub(crate) fn optional_strings<I, S>(
758        &mut self,
759        arguments: I,
760    ) -> Result<Vec<Option<String>>, Error>
761    where
762        I: IntoIterator<Item = S>,
763        S: AsRef<str>,
764    {
765        let (command, reply) = self.run_named(arguments)?;
766
767        match reply {
768            RespValue::Null => Ok(Vec::new()),
769            RespValue::Array(items) => items.iter().map(RespValue::as_string).collect(),
770            other => Err(unexpected(&command, &other)),
771        }
772    }
773
774    pub(crate) fn run_named<I, S>(&mut self, arguments: I) -> Result<(String, RespValue), Error>
775    where
776        I: IntoIterator<Item = S>,
777        S: AsRef<str>,
778    {
779        let owned: Vec<S> = arguments.into_iter().collect();
780        let command = owned
781            .first()
782            .map(|first| first.as_ref().to_owned())
783            .unwrap_or_default();
784        let borrowed: Vec<&str> = owned.iter().map(AsRef::as_ref).collect();
785
786        Ok((command, self.execute(&borrowed)?))
787    }
788}
789
790fn read_ping(reply: RespValue) -> Result<String, Error> {
791    if let Some(items) = reply.as_array() {
792        if let Some(first) = items.first() {
793            if first
794                .as_string()?
795                .is_some_and(|kind| kind.eq_ignore_ascii_case("pong"))
796            {
797                if items.len() > 1 {
798                    return Ok(items[1].as_string()?.unwrap_or_default());
799                }
800
801                return Ok("PONG".to_owned());
802            }
803        }
804    }
805
806    reply
807        .as_string()?
808        .ok_or_else(|| Error::Protocol("PING returned a null reply".to_owned()))
809}
810
811fn finish_standalone(mut seed: Connection, options: ClientOptions) -> Result<Client, Error> {
812    if options.database != 0 {
813        let index = options.database.to_string();
814
815        throw_error(seed.send(&["SELECT", &index])?)?;
816    }
817
818    Ok(Client {
819        transport: Transport::Single(seed),
820        database: options.database,
821        options,
822        dedicated_subscriber: false,
823        pubsub: None,
824        read_timeout: None,
825        on_lost: Arc::new(Mutex::new(None)),
826        on_restored: Arc::new(Mutex::new(None)),
827        reconnect_delay: None,
828        subscriptions: Vec::new(),
829    })
830}
831
832fn is_subscription_command(arguments: &[&str]) -> bool {
833    arguments.first().is_some_and(|name| {
834        matches!(
835            name.to_ascii_uppercase().as_str(),
836            "SUBSCRIBE"
837                | "UNSUBSCRIBE"
838                | "PSUBSCRIBE"
839                | "PUNSUBSCRIBE"
840                | "SSUBSCRIBE"
841                | "SUNSUBSCRIBE"
842        )
843    })
844}
845
846pub(crate) fn unexpected(command: &str, reply: &RespValue) -> Error {
847    Error::Protocol(format!("{command} returned {reply:?}"))
848}
849
850pub(crate) fn join(command: &str, arguments: &[&str]) -> Vec<String> {
851    let mut values = Vec::with_capacity(arguments.len() + 1);
852
853    values.push(command.to_owned());
854    values.extend(arguments.iter().map(|argument| (*argument).to_owned()));
855
856    values
857}
858
859fn set_arguments(key: &str, value: &str, expiry: Option<&[String]>, mode: SetMode) -> Vec<String> {
860    let mut arguments = vec!["SET".to_owned(), key.to_owned(), value.to_owned()];
861
862    if let Some(expiry) = expiry {
863        arguments.extend_from_slice(expiry);
864    }
865
866    match mode {
867        SetMode::Always => {}
868        SetMode::IfNotExists => arguments.push("NX".to_owned()),
869        SetMode::IfExists => arguments.push("XX".to_owned()),
870    }
871
872    arguments
873}
874
875pub(crate) fn unix_millis(time: SystemTime) -> Result<u128, Error> {
876    time.duration_since(UNIX_EPOCH)
877        .map(|elapsed| elapsed.as_millis())
878        .map_err(|_| Error::Protocol("time is before the Unix epoch".to_owned()))
879}
880
881pub(crate) fn duration_millis(ttl: Duration) -> u128 {
882    let millis = ttl.as_millis();
883
884    if ttl.subsec_nanos().is_multiple_of(1_000_000) {
885        return millis;
886    }
887
888    millis.saturating_add(1)
889}
890
891pub(crate) fn duration_secs(ttl: Duration) -> u64 {
892    let seconds = ttl.as_secs();
893
894    if ttl.subsec_nanos() == 0 {
895        return seconds;
896    }
897
898    seconds.saturating_add(1)
899}
900
901#[cfg(test)]
902mod tests {
903    use std::io::{Read, Write};
904
905    use std::net::{TcpListener, TcpStream};
906
907    use std::thread;
908    use std::time::Duration;
909
910    use super::{Client, ClientOptions, RespValue, read_ping};
911
912    fn bound(stream: TcpStream) -> TcpStream {
913        stream
914            .set_read_timeout(Some(Duration::from_secs(5)))
915            .unwrap();
916        stream
917            .set_write_timeout(Some(Duration::from_secs(5)))
918            .unwrap();
919
920        stream
921    }
922
923    fn reply_standalone_cluster(server: &mut TcpStream) {
924        let expected = b"*2\r\n$7\r\nCLUSTER\r\n$5\r\nSLOTS\r\n";
925        let mut got = vec![0_u8; expected.len()];
926
927        server.read_exact(&mut got).unwrap();
928        assert_eq!(got, expected);
929        server
930            .write_all(b"-ERR cluster mode is disabled\r\n")
931            .unwrap();
932    }
933
934    #[test]
935    fn connect_asks_for_cluster_slots_then_is_quiet() {
936        let listener = TcpListener::bind("127.0.0.1:0").unwrap();
937        let port = listener.local_addr().unwrap().port();
938        let accepted = thread::spawn(move || bound(listener.accept().unwrap().0));
939        let connecting = thread::spawn(move || Client::connect("127.0.0.1", port).unwrap());
940        let mut server = accepted.join().unwrap();
941
942        reply_standalone_cluster(&mut server);
943
944        let mut client = connecting.join().unwrap();
945
946        assert!(!client.is_cluster());
947        thread::sleep(Duration::from_millis(50));
948        server.set_nonblocking(true).unwrap();
949
950        let mut peeked = [0_u8; 1];
951
952        assert!(server.peek(&mut peeked).is_err());
953        server.set_nonblocking(false).unwrap();
954
955        let ping = thread::spawn(move || client.ping().unwrap());
956        let mut got = [0_u8; 14];
957
958        server.read_exact(&mut got).unwrap();
959
960        assert_eq!(&got, b"*1\r\n$4\r\nPING\r\n");
961        server.write_all(b"+PONG\r\n").unwrap();
962
963        assert_eq!(ping.join().unwrap(), "PONG");
964    }
965
966    #[test]
967    fn discover_cluster_off_sends_nothing_until_the_first_command() {
968        let listener = TcpListener::bind("127.0.0.1:0").unwrap();
969        let port = listener.local_addr().unwrap().port();
970        let accepted = thread::spawn(move || bound(listener.accept().unwrap().0));
971        let mut client = Client::connect_with(ClientOptions {
972            host: "127.0.0.1".to_owned(),
973            port,
974            discover_cluster: false,
975            ..ClientOptions::default()
976        })
977        .unwrap();
978        let mut server = accepted.join().unwrap();
979
980        thread::sleep(Duration::from_millis(50));
981        server.set_nonblocking(true).unwrap();
982
983        let mut peeked = [0_u8; 1];
984
985        assert!(server.peek(&mut peeked).is_err());
986        server.set_nonblocking(false).unwrap();
987
988        let ping = thread::spawn(move || client.ping().unwrap());
989        let mut got = [0_u8; 14];
990
991        server.read_exact(&mut got).unwrap();
992
993        assert_eq!(&got, b"*1\r\n$4\r\nPING\r\n");
994        server.write_all(b"+PONG\r\n").unwrap();
995
996        assert_eq!(ping.join().unwrap(), "PONG");
997    }
998
999    #[test]
1000    fn a_server_error_leaves_the_next_command_usable() {
1001        let listener = TcpListener::bind("127.0.0.1:0").unwrap();
1002        let port = listener.local_addr().unwrap().port();
1003        let accepted = thread::spawn(move || bound(listener.accept().unwrap().0));
1004        let connecting = thread::spawn(move || Client::connect("127.0.0.1", port).unwrap());
1005        let mut server = accepted.join().unwrap();
1006
1007        reply_standalone_cluster(&mut server);
1008
1009        let mut client = connecting.join().unwrap();
1010        let error = thread::spawn(move || {
1011            let error = client.get_string("missing").unwrap_err();
1012            let pong = client.ping().unwrap();
1013
1014            (error.to_string(), pong)
1015        });
1016
1017        read_some(&mut server);
1018        server.write_all(b"-ERR no such key\r\n").unwrap();
1019        read_some(&mut server);
1020        server.write_all(b"+PONG\r\n").unwrap();
1021
1022        let (message, pong) = error.join().unwrap();
1023
1024        assert_eq!(message, "ERR no such key");
1025        assert_eq!(pong, "PONG");
1026    }
1027
1028    #[test]
1029    fn password_is_sent_as_auth() {
1030        let listener = TcpListener::bind("127.0.0.1:0").unwrap();
1031        let port = listener.local_addr().unwrap().port();
1032        let accepted = thread::spawn(move || bound(listener.accept().unwrap().0));
1033        let connecting = thread::spawn(move || {
1034            Client::connect_with(ClientOptions {
1035                host: "127.0.0.1".to_owned(),
1036                port,
1037                password: Some("secret".to_owned()),
1038                ..ClientOptions::default()
1039            })
1040            .unwrap()
1041        });
1042
1043        let mut server = accepted.join().unwrap();
1044        let expected = b"*2\r\n$4\r\nAUTH\r\n$6\r\nsecret\r\n";
1045        let mut got = vec![0_u8; expected.len()];
1046
1047        server.read_exact(&mut got).unwrap();
1048
1049        assert_eq!(got, expected);
1050        server.write_all(b"+OK\r\n").unwrap();
1051        reply_standalone_cluster(&mut server);
1052        connecting.join().unwrap();
1053    }
1054
1055    #[test]
1056    fn database_is_selected_after_auth() {
1057        let listener = TcpListener::bind("127.0.0.1:0").unwrap();
1058        let port = listener.local_addr().unwrap().port();
1059        let accepted = thread::spawn(move || bound(listener.accept().unwrap().0));
1060        let connecting = thread::spawn(move || {
1061            Client::connect_with(ClientOptions {
1062                host: "127.0.0.1".to_owned(),
1063                port,
1064                password: Some("secret".to_owned()),
1065                database: 3,
1066                ..ClientOptions::default()
1067            })
1068            .unwrap()
1069        });
1070
1071        let mut server = accepted.join().unwrap();
1072        let auth = b"*2\r\n$4\r\nAUTH\r\n$6\r\nsecret\r\n";
1073        let select = b"*2\r\n$6\r\nSELECT\r\n$1\r\n3\r\n";
1074        let mut got = vec![0_u8; auth.len()];
1075
1076        server.read_exact(&mut got).unwrap();
1077        assert_eq!(got, auth);
1078        server.write_all(b"+OK\r\n").unwrap();
1079        reply_standalone_cluster(&mut server);
1080        got = vec![0_u8; select.len()];
1081        server.read_exact(&mut got).unwrap();
1082        assert_eq!(got, select);
1083        server.write_all(b"+OK\r\n").unwrap();
1084        assert_eq!(connecting.join().unwrap().database(), 3);
1085    }
1086
1087    #[test]
1088    fn subscribe_opens_a_second_socket_and_leaves_ping_on_the_first() {
1089        let listener = TcpListener::bind("127.0.0.1:0").unwrap();
1090        let port = listener.local_addr().unwrap().port();
1091        let connecting = thread::spawn(move || Client::connect("127.0.0.1", port).unwrap());
1092        let mut command = bound(listener.accept().unwrap().0);
1093
1094        reply_standalone_cluster(&mut command);
1095
1096        let mut client = connecting.join().unwrap();
1097        let working = thread::spawn(move || {
1098            let confirms = client.subscribe(&["news"]).unwrap();
1099            let pong = client.ping().unwrap();
1100
1101            (confirms.len(), pong)
1102        });
1103
1104        let mut subscriber = bound(listener.accept().unwrap().0);
1105        let expected = b"*2\r\n$9\r\nSUBSCRIBE\r\n$4\r\nnews\r\n";
1106        let mut got = vec![0_u8; expected.len()];
1107
1108        subscriber.read_exact(&mut got).unwrap();
1109        assert_eq!(got, expected);
1110        subscriber
1111            .write_all(b"*3\r\n$9\r\nsubscribe\r\n$4\r\nnews\r\n:1\r\n")
1112            .unwrap();
1113
1114        let mut ping = [0_u8; 14];
1115
1116        command.read_exact(&mut ping).unwrap();
1117        assert_eq!(&ping, b"*1\r\n$4\r\nPING\r\n");
1118        command.write_all(b"+PONG\r\n").unwrap();
1119
1120        let (confirms, pong) = working.join().unwrap();
1121
1122        assert_eq!(confirms, 1);
1123        assert_eq!(pong, "PONG");
1124    }
1125
1126    #[test]
1127    fn explicit_subscriber_uses_one_extra_socket() {
1128        let listener = TcpListener::bind("127.0.0.1:0").unwrap();
1129        let port = listener.local_addr().unwrap().port();
1130        let connecting = thread::spawn(move || Client::connect("127.0.0.1", port).unwrap());
1131        let mut command = bound(listener.accept().unwrap().0);
1132
1133        reply_standalone_cluster(&mut command);
1134
1135        let client = connecting.join().unwrap();
1136        let working = thread::spawn(move || {
1137            let mut subscriber = client.subscriber().unwrap();
1138
1139            subscriber.subscribe(&["news"]).unwrap();
1140        });
1141
1142        let mut subscriber = bound(listener.accept().unwrap().0);
1143        let expected = b"*2\r\n$9\r\nSUBSCRIBE\r\n$4\r\nnews\r\n";
1144        let mut got = vec![0_u8; expected.len()];
1145
1146        subscriber.read_exact(&mut got).unwrap();
1147        assert_eq!(got, expected);
1148        subscriber
1149            .write_all(b"*3\r\n$9\r\nsubscribe\r\n$4\r\nnews\r\n:1\r\n")
1150            .unwrap();
1151        working.join().unwrap();
1152        listener.set_nonblocking(true).unwrap();
1153
1154        assert!(listener.accept().is_err());
1155    }
1156
1157    #[test]
1158    fn the_command_after_a_dropped_socket_reconnects() {
1159        let listener = TcpListener::bind("127.0.0.1:0").unwrap();
1160        let port = listener.local_addr().unwrap().port();
1161        let connecting = thread::spawn(move || {
1162            Client::connect_with(ClientOptions {
1163                host: "127.0.0.1".to_owned(),
1164                port,
1165                discover_cluster: false,
1166                reconnect_base_delay: std::time::Duration::ZERO,
1167                reconnect_max_delay: std::time::Duration::ZERO,
1168                ..ClientOptions::default()
1169            })
1170            .unwrap()
1171        });
1172
1173        let accepted = bound(listener.accept().unwrap().0);
1174
1175        drop(accepted);
1176
1177        let mut client = connecting.join().unwrap();
1178        let failed = thread::spawn(move || {
1179            assert!(client.ping().is_err());
1180
1181            client.ping().unwrap()
1182        });
1183
1184        let mut again = bound(listener.accept().unwrap().0);
1185        let mut got = [0_u8; 14];
1186
1187        again.read_exact(&mut got).unwrap();
1188        assert_eq!(&got, b"*1\r\n$4\r\nPING\r\n");
1189        again.write_all(b"+PONG\r\n").unwrap();
1190
1191        assert_eq!(failed.join().unwrap(), "PONG");
1192    }
1193
1194    #[test]
1195    fn ping_reads_a_pubsub_array() {
1196        let reply = RespValue::Array(vec![
1197            RespValue::Bulk(b"pong".to_vec()),
1198            RespValue::Bulk(b"hello".to_vec()),
1199        ]);
1200
1201        assert_eq!(read_ping(reply).unwrap(), "hello");
1202    }
1203
1204    fn read_some(stream: &mut TcpStream) {
1205        let mut buffer = [0_u8; 64];
1206
1207        let read = stream.read(&mut buffer).unwrap();
1208
1209        assert!(read > 0);
1210    }
1211}