From 38082b927abe5039afdebbe63a5a6bc9e607a33e Mon Sep 17 00:00:00 2001 From: Kacy Fortner Date: Thu, 5 Feb 2026 22:17:07 -0500 Subject: [PATCH] feat: add set commands (SADD, SREM, SMEMBERS, SISMEMBER, SCARD) implements redis set data type with the core operations: - SADD: add members to a set, returns count of new members - SREM: remove members from a set, auto-deletes empty sets - SMEMBERS: retrieve all members - SISMEMBER: check membership - SCARD: return set cardinality includes snapshot persistence, memory tracking, and full test coverage. --- README.md | 8 + crates/ember-core/src/keyspace.rs | 245 +++++++++++++++++++++++ crates/ember-core/src/memory.rs | 12 ++ crates/ember-core/src/shard.rs | 41 ++++ crates/ember-core/src/types/mod.rs | 9 +- crates/ember-persistence/src/recovery.rs | 3 + crates/ember-persistence/src/snapshot.rs | 29 ++- crates/ember-protocol/src/command.rs | 192 ++++++++++++++++++ crates/ember-server/src/connection.rs | 66 ++++++ 9 files changed, 602 insertions(+), 3 deletions(-) diff --git a/README.md b/README.md index 90e019a8..7a8cbda7 100644 --- a/README.md +++ b/README.md @@ -13,6 +13,7 @@ a low-latency, memory-efficient, distributed cache written in Rust. designed to - **list operations** — LPUSH, RPUSH, LPOP, RPOP, LRANGE, LLEN - **sorted sets** — ZADD (with NX/XX/GT/LT/CH), ZREM, ZSCORE, ZRANK, ZRANGE, ZCARD - **hashes** — HSET, HGET, HGETALL, HDEL, HEXISTS, HLEN, HINCRBY, HKEYS, HVALS, HMGET +- **sets** — SADD, SREM, SMEMBERS, SISMEMBER, SCARD - **key commands** — DEL, EXISTS, EXPIRE, TTL, PEXPIRE, PTTL, PERSIST, TYPE, SCAN - **server commands** — PING, ECHO, INFO, DBSIZE, FLUSHDB, BGSAVE, BGREWRITEAOF - **sharded engine** — shared-nothing, thread-per-core design with no cross-shard locking @@ -71,6 +72,13 @@ redis-cli HGET user:1 name # => "alice" redis-cli HGETALL user:1 # => 1) "name" 2) "alice" 3) "age" 4) "30" redis-cli HINCRBY user:1 age 1 # => (integer) 31 +# sets +redis-cli SADD tags rust cache fast # => (integer) 3 +redis-cli SMEMBERS tags # => 1) "cache" 2) "fast" 3) "rust" +redis-cli SISMEMBER tags rust # => (integer) 1 +redis-cli SCARD tags # => (integer) 3 +redis-cli SREM tags fast # => (integer) 1 + # iteration redis-cli SCAN 0 MATCH "user:*" COUNT 100 redis-cli DBSIZE # => (integer) 6 diff --git a/crates/ember-core/src/keyspace.rs b/crates/ember-core/src/keyspace.rs index f65b2367..7988558d 100644 --- a/crates/ember-core/src/keyspace.rs +++ b/crates/ember-core/src/keyspace.rs @@ -1366,6 +1366,162 @@ impl Keyspace { } } + // ------------------------------------------------------------------------- + // Set operations + // ------------------------------------------------------------------------- + + /// Adds one or more members to a set. + /// + /// Creates the set if the key doesn't exist. Returns the number of + /// new members added (existing members don't count). + pub fn sadd(&mut self, key: &str, members: &[String]) -> Result { + if members.is_empty() { + return Ok(0); + } + + self.remove_if_expired(key); + + let is_new = !self.entries.contains_key(key); + + if !is_new && !matches!(self.entries[key].value, Value::Set(_)) { + return Err(WriteError::WrongType); + } + + // estimate memory increase before mutating + let member_increase: usize = members + .iter() + .map(|m| m.len() + memory::HASHSET_MEMBER_OVERHEAD) + .sum(); + let estimated_increase = if is_new { + memory::ENTRY_OVERHEAD + key.len() + memory::HASHSET_BASE_OVERHEAD + member_increase + } else { + member_increase + }; + if !self.enforce_memory_limit(estimated_increase) { + return Err(WriteError::OutOfMemory); + } + + if is_new { + let value = Value::Set(std::collections::HashSet::new()); + self.memory.add(key, &value); + self.entries.insert(key.to_owned(), Entry::new(value, None)); + } + + let entry = self + .entries + .get_mut(key) + .expect("just inserted or verified"); + let old_entry_size = memory::entry_size(key, &entry.value); + + let mut added = 0; + if let Value::Set(ref mut set) = entry.value { + for member in members { + if set.insert(member.clone()) { + added += 1; + } + } + } + entry.touch(); + + let new_entry_size = memory::entry_size(key, &entry.value); + self.memory.adjust(old_entry_size, new_entry_size); + + Ok(added) + } + + /// Removes one or more members from a set. + /// + /// Returns the number of members that were actually removed. + pub fn srem(&mut self, key: &str, members: &[String]) -> Result { + if self.remove_if_expired(key) { + return Ok(0); + } + + match self.entries.get(key) { + None => return Ok(0), + Some(e) => { + if !matches!(e.value, Value::Set(_)) { + return Err(WrongType); + } + } + } + + let old_entry_size = memory::entry_size(key, &self.entries[key].value); + let entry = self.entries.get_mut(key).expect("verified above"); + + let mut removed = 0; + let is_empty = if let Value::Set(ref mut set) = entry.value { + for member in members { + if set.remove(member) { + removed += 1; + } + } + set.is_empty() + } else { + false + }; + + if is_empty { + self.memory.remove_with_size(old_entry_size); + self.entries.remove(key); + } else { + let new_entry_size = memory::entry_size(key, &self.entries[key].value); + self.memory.adjust(old_entry_size, new_entry_size); + } + + Ok(removed) + } + + /// Returns all members of a set. + pub fn smembers(&mut self, key: &str) -> Result, WrongType> { + if self.remove_if_expired(key) { + return Ok(vec![]); + } + match self.entries.get_mut(key) { + None => Ok(vec![]), + Some(entry) => match &entry.value { + Value::Set(set) => { + let result = set.iter().cloned().collect(); + entry.touch(); + Ok(result) + } + _ => Err(WrongType), + }, + } + } + + /// Checks if a member exists in a set. + pub fn sismember(&mut self, key: &str, member: &str) -> Result { + if self.remove_if_expired(key) { + return Ok(false); + } + match self.entries.get_mut(key) { + None => Ok(false), + Some(entry) => match &entry.value { + Value::Set(set) => { + let result = set.contains(member); + entry.touch(); + Ok(result) + } + _ => Err(WrongType), + }, + } + } + + /// Returns the cardinality (number of elements) of a set. + pub fn scard(&mut self, key: &str) -> Result { + if self.remove_if_expired(key) { + return Ok(0); + } + match self.entries.get(key) { + None => Ok(0), + Some(entry) => match &entry.value { + Value::Set(set) => Ok(set.len()), + _ => Err(WrongType), + }, + } + } + /// Randomly samples up to `count` keys and removes any that have expired. /// /// Returns the number of keys actually removed. Used by the active @@ -3003,4 +3159,93 @@ mod tests { assert!(ks.hvals("s").is_err()); assert!(ks.hmget("s", &["f".into()]).is_err()); } + + // --- set tests --- + + #[test] + fn sadd_creates_set() { + let mut ks = Keyspace::new(); + let added = ks.sadd("s", &["a".into(), "b".into()]).unwrap(); + assert_eq!(added, 2); + assert_eq!(ks.value_type("s"), "set"); + } + + #[test] + fn sadd_returns_new_member_count() { + let mut ks = Keyspace::new(); + ks.sadd("s", &["a".into(), "b".into()]).unwrap(); + // add one existing, one new + let added = ks.sadd("s", &["b".into(), "c".into()]).unwrap(); + assert_eq!(added, 1); // only "c" is new + } + + #[test] + fn srem_removes_members() { + let mut ks = Keyspace::new(); + ks.sadd("s", &["a".into(), "b".into(), "c".into()]).unwrap(); + let removed = ks.srem("s", &["a".into(), "c".into()]).unwrap(); + assert_eq!(removed, 2); + assert_eq!(ks.scard("s").unwrap(), 1); + } + + #[test] + fn srem_auto_deletes_empty_set() { + let mut ks = Keyspace::new(); + ks.sadd("s", &["only".into()]).unwrap(); + ks.srem("s", &["only".into()]).unwrap(); + assert_eq!(ks.value_type("s"), "none"); + } + + #[test] + fn smembers_returns_all_members() { + let mut ks = Keyspace::new(); + ks.sadd("s", &["a".into(), "b".into(), "c".into()]).unwrap(); + let mut members = ks.smembers("s").unwrap(); + members.sort(); + assert_eq!(members, vec!["a", "b", "c"]); + } + + #[test] + fn smembers_missing_key_returns_empty() { + let mut ks = Keyspace::new(); + assert_eq!(ks.smembers("missing").unwrap(), Vec::::new()); + } + + #[test] + fn sismember_returns_true_for_existing() { + let mut ks = Keyspace::new(); + ks.sadd("s", &["member".into()]).unwrap(); + assert!(ks.sismember("s", "member").unwrap()); + } + + #[test] + fn sismember_returns_false_for_missing() { + let mut ks = Keyspace::new(); + ks.sadd("s", &["a".into()]).unwrap(); + assert!(!ks.sismember("s", "missing").unwrap()); + } + + #[test] + fn scard_returns_count() { + let mut ks = Keyspace::new(); + ks.sadd("s", &["a".into(), "b".into(), "c".into()]).unwrap(); + assert_eq!(ks.scard("s").unwrap(), 3); + } + + #[test] + fn scard_missing_key_returns_zero() { + let mut ks = Keyspace::new(); + assert_eq!(ks.scard("missing").unwrap(), 0); + } + + #[test] + fn set_on_string_key_returns_wrongtype() { + let mut ks = Keyspace::new(); + ks.set("s".into(), Bytes::from("string"), None); + assert!(ks.sadd("s", &["m".into()]).is_err()); + assert!(ks.srem("s", &["m".into()]).is_err()); + assert!(ks.smembers("s").is_err()); + assert!(ks.sismember("s", "m").is_err()); + assert!(ks.scard("s").is_err()); + } } diff --git a/crates/ember-core/src/memory.rs b/crates/ember-core/src/memory.rs index 42bbb21d..3341d162 100644 --- a/crates/ember-core/src/memory.rs +++ b/crates/ember-core/src/memory.rs @@ -131,6 +131,14 @@ pub(crate) const HASHMAP_ENTRY_OVERHEAD: usize = 64; /// Base overhead for an empty HashMap (bucket array pointer + len + capacity). pub(crate) const HASHMAP_BASE_OVERHEAD: usize = 48; +/// Estimated overhead per member in a HashSet. +/// +/// Each member is a String (24 bytes ptr+len+cap) plus bucket overhead. +pub(crate) const HASHSET_MEMBER_OVERHEAD: usize = 40; + +/// Base overhead for an empty HashSet (bucket array pointer + len + capacity). +pub(crate) const HASHSET_BASE_OVERHEAD: usize = 48; + /// Returns the byte size of a value's payload. pub fn value_size(value: &Value) -> usize { match value { @@ -150,6 +158,10 @@ pub fn value_size(value: &Value) -> usize { .sum(); HASHMAP_BASE_OVERHEAD + entry_bytes } + Value::Set(set) => { + let member_bytes: usize = set.iter().map(|m| m.len() + HASHSET_MEMBER_OVERHEAD).sum(); + HASHSET_BASE_OVERHEAD + member_bytes + } } } diff --git a/crates/ember-core/src/shard.rs b/crates/ember-core/src/shard.rs index 2f875cf4..db531ce9 100644 --- a/crates/ember-core/src/shard.rs +++ b/crates/ember-core/src/shard.rs @@ -177,6 +177,24 @@ pub enum ShardRequest { key: String, fields: Vec, }, + SAdd { + key: String, + members: Vec, + }, + SRem { + key: String, + members: Vec, + }, + SMembers { + key: String, + }, + SIsMember { + key: String, + member: String, + }, + SCard { + key: String, + }, /// Returns the key count for this shard. DbSize, /// Returns keyspace stats for this shard. @@ -334,6 +352,7 @@ async fn run_shard( Value::SortedSet(ss) } RecoveredValue::Hash(map) => Value::Hash(map), + RecoveredValue::Set(set) => Value::Set(set), }; keyspace.restore(entry.key, value, entry.expires_at); } @@ -648,6 +667,27 @@ fn dispatch(ks: &mut Keyspace, req: &ShardRequest) -> ShardResponse { Ok(vals) => ShardResponse::OptionalArray(vals), Err(_) => ShardResponse::WrongType, }, + ShardRequest::SAdd { key, members } => match ks.sadd(key, members) { + Ok(count) => ShardResponse::Len(count), + Err(WriteError::WrongType) => ShardResponse::WrongType, + Err(WriteError::OutOfMemory) => ShardResponse::OutOfMemory, + }, + ShardRequest::SRem { key, members } => match ks.srem(key, members) { + Ok(count) => ShardResponse::Len(count), + Err(_) => ShardResponse::WrongType, + }, + ShardRequest::SMembers { key } => match ks.smembers(key) { + Ok(members) => ShardResponse::StringArray(members), + Err(_) => ShardResponse::WrongType, + }, + ShardRequest::SIsMember { key, member } => match ks.sismember(key, member) { + Ok(exists) => ShardResponse::Bool(exists), + Err(_) => ShardResponse::WrongType, + }, + ShardRequest::SCard { key } => match ks.scard(key) { + Ok(count) => ShardResponse::Len(count), + Err(_) => ShardResponse::WrongType, + }, // snapshot/rewrite are handled in the main loop, not here ShardRequest::Snapshot | ShardRequest::RewriteAof => ShardResponse::Ok, } @@ -804,6 +844,7 @@ fn write_snapshot( SnapValue::SortedSet(members) } Value::Hash(map) => SnapValue::Hash(map.clone()), + Value::Set(set) => SnapValue::Set(set.clone()), }; writer.write_entry(&SnapEntry { key: key.to_owned(), diff --git a/crates/ember-core/src/types/mod.rs b/crates/ember-core/src/types/mod.rs index fade45d9..24b08bba 100644 --- a/crates/ember-core/src/types/mod.rs +++ b/crates/ember-core/src/types/mod.rs @@ -1,11 +1,11 @@ //! Data type representations for stored values. //! //! Each variant maps to a Redis-like data type. Strings, lists, sorted -//! sets, and hashes are supported; plain sets will come later. +//! sets, hashes, and sets are supported. pub mod sorted_set; -use std::collections::{HashMap, VecDeque}; +use std::collections::{HashMap, HashSet, VecDeque}; use bytes::Bytes; @@ -33,6 +33,9 @@ pub enum Value { /// Hash map of field names to values. Fields are unique strings, /// values are binary-safe byte sequences. Hash(HashMap), + + /// Unordered set of unique string members. + Set(HashSet), } impl PartialEq for Value { @@ -47,6 +50,7 @@ impl PartialEq for Value { .all(|((m1, s1), (m2, s2))| m1 == m2 && s1 == s2) } (Value::Hash(a), Value::Hash(b)) => a == b, + (Value::Set(a), Value::Set(b)) => a == b, _ => false, } } @@ -59,6 +63,7 @@ pub fn type_name(value: &Value) -> &'static str { Value::List(_) => "list", Value::SortedSet(_) => "zset", Value::Hash(_) => "hash", + Value::Set(_) => "set", } } diff --git a/crates/ember-persistence/src/recovery.rs b/crates/ember-persistence/src/recovery.rs index ae07a385..14f7bdb6 100644 --- a/crates/ember-persistence/src/recovery.rs +++ b/crates/ember-persistence/src/recovery.rs @@ -27,6 +27,8 @@ pub enum RecoveredValue { SortedSet(Vec<(f64, String)>), /// Hash map of field names to values. Hash(HashMap), + /// Unordered set of unique string members. + Set(HashSet), } impl From for RecoveredValue { @@ -36,6 +38,7 @@ impl From for RecoveredValue { SnapValue::List(deque) => RecoveredValue::List(deque), SnapValue::SortedSet(members) => RecoveredValue::SortedSet(members), SnapValue::Hash(map) => RecoveredValue::Hash(map), + SnapValue::Set(set) => RecoveredValue::Set(set), } } } diff --git a/crates/ember-persistence/src/snapshot.rs b/crates/ember-persistence/src/snapshot.rs index 5ebeaff0..197cc16e 100644 --- a/crates/ember-persistence/src/snapshot.rs +++ b/crates/ember-persistence/src/snapshot.rs @@ -21,7 +21,7 @@ //! `expire_ms` is the remaining TTL in milliseconds, or -1 for no expiry. //! v1 entries (no type tag) are still readable for backward compatibility. -use std::collections::{HashMap, VecDeque}; +use std::collections::{HashMap, HashSet, VecDeque}; use std::fs::{self, File}; use std::io::{self, BufReader, BufWriter, Write}; use std::path::{Path, PathBuf}; @@ -35,6 +35,7 @@ const TYPE_STRING: u8 = 0; const TYPE_LIST: u8 = 1; const TYPE_SORTED_SET: u8 = 2; const TYPE_HASH: u8 = 3; +const TYPE_SET: u8 = 4; /// The value stored in a snapshot entry. #[derive(Debug, Clone, PartialEq)] @@ -47,6 +48,8 @@ pub enum SnapValue { SortedSet(Vec<(f64, String)>), /// A hash: map of field names to values. Hash(HashMap), + /// An unordered set of unique string members. + Set(HashSet), } /// A single entry in a snapshot file. @@ -134,6 +137,13 @@ impl SnapshotWriter { format::write_bytes(&mut buf, value)?; } } + SnapValue::Set(set) => { + format::write_u8(&mut buf, TYPE_SET)?; + format::write_u32(&mut buf, set.len() as u32)?; + for member in set { + format::write_bytes(&mut buf, member.as_bytes())?; + } + } } format::write_i64(&mut buf, entry.expire_ms)?; @@ -286,6 +296,23 @@ impl SnapshotReader { } SnapValue::Hash(map) } + TYPE_SET => { + let count = format::read_u32(&mut self.reader)?; + format::write_u32(&mut buf, count).expect("vec write"); + let mut set = HashSet::with_capacity(count as usize); + for _ in 0..count { + let member_bytes = format::read_bytes(&mut self.reader)?; + format::write_bytes(&mut buf, &member_bytes).expect("vec write"); + let member = String::from_utf8(member_bytes).map_err(|_| { + FormatError::Io(io::Error::new( + io::ErrorKind::InvalidData, + "set member is not valid utf-8", + )) + })?; + set.insert(member); + } + SnapValue::Set(set) + } _ => { return Err(FormatError::UnknownTag(type_tag)); } diff --git a/crates/ember-protocol/src/command.rs b/crates/ember-protocol/src/command.rs index 224b2e27..dcd7b3cb 100644 --- a/crates/ember-protocol/src/command.rs +++ b/crates/ember-protocol/src/command.rs @@ -181,6 +181,21 @@ pub enum Command { /// HMGET [field ...]. Gets multiple field values from a hash. HMGet { key: String, fields: Vec }, + /// SADD [member ...]. Adds members to a set. + SAdd { key: String, members: Vec }, + + /// SREM [member ...]. Removes members from a set. + SRem { key: String, members: Vec }, + + /// SMEMBERS . Returns all members of a set. + SMembers { key: String }, + + /// SISMEMBER . Checks if a member exists in a set. + SIsMember { key: String, member: String }, + + /// SCARD . Returns the cardinality (number of members) of a set. + SCard { key: String }, + /// A command we don't recognize (yet). Unknown(String), } @@ -271,6 +286,11 @@ impl Command { "HKEYS" => parse_hkeys(&frames[1..]), "HVALS" => parse_hvals(&frames[1..]), "HMGET" => parse_hmget(&frames[1..]), + "SADD" => parse_sadd(&frames[1..]), + "SREM" => parse_srem(&frames[1..]), + "SMEMBERS" => parse_smembers(&frames[1..]), + "SISMEMBER" => parse_sismember(&frames[1..]), + "SCARD" => parse_scard(&frames[1..]), _ => Ok(Command::Unknown(name)), } } @@ -945,6 +965,57 @@ fn parse_hmget(args: &[Frame]) -> Result { Ok(Command::HMGet { key, fields }) } +// --- set commands --- + +fn parse_sadd(args: &[Frame]) -> Result { + if args.len() < 2 { + return Err(ProtocolError::WrongArity("SADD".into())); + } + let key = extract_string(&args[0])?; + let members = args[1..] + .iter() + .map(extract_string) + .collect::, _>>()?; + Ok(Command::SAdd { key, members }) +} + +fn parse_srem(args: &[Frame]) -> Result { + if args.len() < 2 { + return Err(ProtocolError::WrongArity("SREM".into())); + } + let key = extract_string(&args[0])?; + let members = args[1..] + .iter() + .map(extract_string) + .collect::, _>>()?; + Ok(Command::SRem { key, members }) +} + +fn parse_smembers(args: &[Frame]) -> Result { + if args.len() != 1 { + return Err(ProtocolError::WrongArity("SMEMBERS".into())); + } + let key = extract_string(&args[0])?; + Ok(Command::SMembers { key }) +} + +fn parse_sismember(args: &[Frame]) -> Result { + if args.len() != 2 { + return Err(ProtocolError::WrongArity("SISMEMBER".into())); + } + let key = extract_string(&args[0])?; + let member = extract_string(&args[1])?; + Ok(Command::SIsMember { key, member }) +} + +fn parse_scard(args: &[Frame]) -> Result { + if args.len() != 1 { + return Err(ProtocolError::WrongArity("SCARD".into())); + } + let key = extract_string(&args[0])?; + Ok(Command::SCard { key }) +} + #[cfg(test)] mod tests { use super::*; @@ -2314,4 +2385,125 @@ mod tests { Command::HGetAll { .. } )); } + + // --- set commands --- + + #[test] + fn sadd_single_member() { + assert_eq!( + Command::from_frame(cmd(&["SADD", "s", "member"])).unwrap(), + Command::SAdd { + key: "s".into(), + members: vec!["member".into()], + }, + ); + } + + #[test] + fn sadd_multiple_members() { + let parsed = Command::from_frame(cmd(&["SADD", "s", "a", "b", "c"])).unwrap(); + match parsed { + Command::SAdd { key, members } => { + assert_eq!(key, "s"); + assert_eq!(members.len(), 3); + } + other => panic!("expected SAdd, got {other:?}"), + } + } + + #[test] + fn sadd_wrong_arity() { + let err = Command::from_frame(cmd(&["SADD", "s"])).unwrap_err(); + assert!(matches!(err, ProtocolError::WrongArity(_))); + } + + #[test] + fn srem_single_member() { + assert_eq!( + Command::from_frame(cmd(&["SREM", "s", "member"])).unwrap(), + Command::SRem { + key: "s".into(), + members: vec!["member".into()], + }, + ); + } + + #[test] + fn srem_multiple_members() { + let parsed = Command::from_frame(cmd(&["SREM", "s", "a", "b"])).unwrap(); + match parsed { + Command::SRem { key, members } => { + assert_eq!(key, "s"); + assert_eq!(members.len(), 2); + } + other => panic!("expected SRem, got {other:?}"), + } + } + + #[test] + fn srem_wrong_arity() { + let err = Command::from_frame(cmd(&["SREM", "s"])).unwrap_err(); + assert!(matches!(err, ProtocolError::WrongArity(_))); + } + + #[test] + fn smembers_basic() { + assert_eq!( + Command::from_frame(cmd(&["SMEMBERS", "s"])).unwrap(), + Command::SMembers { key: "s".into() }, + ); + } + + #[test] + fn smembers_wrong_arity() { + let err = Command::from_frame(cmd(&["SMEMBERS"])).unwrap_err(); + assert!(matches!(err, ProtocolError::WrongArity(_))); + } + + #[test] + fn sismember_basic() { + assert_eq!( + Command::from_frame(cmd(&["SISMEMBER", "s", "member"])).unwrap(), + Command::SIsMember { + key: "s".into(), + member: "member".into(), + }, + ); + } + + #[test] + fn sismember_wrong_arity() { + let err = Command::from_frame(cmd(&["SISMEMBER", "s"])).unwrap_err(); + assert!(matches!(err, ProtocolError::WrongArity(_))); + } + + #[test] + fn scard_basic() { + assert_eq!( + Command::from_frame(cmd(&["SCARD", "s"])).unwrap(), + Command::SCard { key: "s".into() }, + ); + } + + #[test] + fn scard_wrong_arity() { + let err = Command::from_frame(cmd(&["SCARD"])).unwrap_err(); + assert!(matches!(err, ProtocolError::WrongArity(_))); + } + + #[test] + fn set_commands_case_insensitive() { + assert!(matches!( + Command::from_frame(cmd(&["sadd", "s", "m"])).unwrap(), + Command::SAdd { .. } + )); + assert!(matches!( + Command::from_frame(cmd(&["srem", "s", "m"])).unwrap(), + Command::SRem { .. } + )); + assert!(matches!( + Command::from_frame(cmd(&["smembers", "s"])).unwrap(), + Command::SMembers { .. } + )); + } } diff --git a/crates/ember-server/src/connection.rs b/crates/ember-server/src/connection.rs index ca9e1efd..e50b358e 100644 --- a/crates/ember-server/src/connection.rs +++ b/crates/ember-server/src/connection.rs @@ -768,6 +768,72 @@ async fn execute(cmd: Command, engine: &Engine) -> Frame { } } + // --- set commands --- + Command::SAdd { key, members } => { + let req = ShardRequest::SAdd { + key: key.clone(), + members, + }; + match engine.route(&key, req).await { + Ok(ShardResponse::Len(n)) => Frame::Integer(n as i64), + Ok(ShardResponse::WrongType) => wrongtype_error(), + Ok(ShardResponse::OutOfMemory) => oom_error(), + Ok(other) => Frame::Error(format!("ERR unexpected shard response: {other:?}")), + Err(e) => Frame::Error(format!("ERR {e}")), + } + } + + Command::SRem { key, members } => { + let req = ShardRequest::SRem { + key: key.clone(), + members, + }; + match engine.route(&key, req).await { + Ok(ShardResponse::Len(n)) => Frame::Integer(n as i64), + Ok(ShardResponse::WrongType) => wrongtype_error(), + Ok(other) => Frame::Error(format!("ERR unexpected shard response: {other:?}")), + Err(e) => Frame::Error(format!("ERR {e}")), + } + } + + Command::SMembers { key } => { + let req = ShardRequest::SMembers { key: key.clone() }; + match engine.route(&key, req).await { + Ok(ShardResponse::StringArray(members)) => Frame::Array( + members + .into_iter() + .map(|m| Frame::Bulk(Bytes::from(m))) + .collect(), + ), + Ok(ShardResponse::WrongType) => wrongtype_error(), + Ok(other) => Frame::Error(format!("ERR unexpected shard response: {other:?}")), + Err(e) => Frame::Error(format!("ERR {e}")), + } + } + + Command::SIsMember { key, member } => { + let req = ShardRequest::SIsMember { + key: key.clone(), + member, + }; + match engine.route(&key, req).await { + Ok(ShardResponse::Bool(b)) => Frame::Integer(if b { 1 } else { 0 }), + Ok(ShardResponse::WrongType) => wrongtype_error(), + Ok(other) => Frame::Error(format!("ERR unexpected shard response: {other:?}")), + Err(e) => Frame::Error(format!("ERR {e}")), + } + } + + Command::SCard { key } => { + let req = ShardRequest::SCard { key: key.clone() }; + match engine.route(&key, req).await { + Ok(ShardResponse::Len(n)) => Frame::Integer(n as i64), + Ok(ShardResponse::WrongType) => wrongtype_error(), + Ok(other) => Frame::Error(format!("ERR unexpected shard response: {other:?}")), + Err(e) => Frame::Error(format!("ERR {e}")), + } + } + Command::Unknown(name) => Frame::Error(format!("ERR unknown command '{name}'")), } }