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#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
14pub enum SetMode {
15 #[default]
17 Always,
18 IfNotExists,
20 IfExists,
22}
23
24#[derive(Debug, Clone)]
26pub struct ClientOptions {
27 pub host: String,
29 pub port: u16,
31 pub username: Option<String>,
33 pub password: Option<String>,
35 pub database: u32,
37 pub connect_timeout: Duration,
39 pub discover_cluster: bool,
42 pub reconnect: bool,
46 pub max_reconnect_attempts: u32,
48 pub reconnect_base_delay: Duration,
50 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#[derive(Debug, Clone, PartialEq, Eq)]
74pub struct ConnectionNotice {
75 pub host: String,
76 pub port: u16,
77 pub subscriber: bool,
79}
80
81pub 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 pub fn connect(host: impl Into<String>, port: u16) -> Result<Self, Error> {
110 Self::connect_database(host, port, 0)
111 }
112
113 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 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 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 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 pub fn is_cluster(&self) -> bool {
183 matches!(self.transport, Transport::Cluster(_))
184 }
185
186 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 pub fn ping(&mut self) -> Result<String, Error> {
196 read_ping(self.run(["PING"])?)
197 }
198
199 pub fn ping_message(&mut self, message: &str) -> Result<String, Error> {
201 read_ping(self.run(["PING", message])?)
202 }
203
204 pub fn database(&self) -> u32 {
206 self.database
207 }
208
209 pub fn subscriber(&self) -> Result<Self, Error> {
212 self.open_pubsub(&["SUBSCRIBE"])
213 }
214
215 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 pub fn into_database(mut self, database: u32) -> Result<Self, Error> {
227 self.select(database)?;
228
229 Ok(self)
230 }
231
232 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 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 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 pub fn set(&mut self, key: &str, value: &str) -> Result<bool, Error> {
270 self.set_with(key, value, None, SetMode::Always)
271 }
272
273 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 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 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 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 pub fn incr(&mut self, key: &str) -> Result<i64, Error> {
327 self.integer(["INCR", key])
328 }
329
330 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 pub fn del_key(&mut self, key: &str) -> Result<i64, Error> {
349 self.del(&[key])
350 }
351
352 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 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 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 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 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 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 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 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 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}