diff --git a/README.md b/README.md index 516817ed..90e019a8 100644 --- a/README.md +++ b/README.md @@ -12,6 +12,7 @@ a low-latency, memory-efficient, distributed cache written in Rust. designed to - **string commands** — GET, SET (with NX/XX/EX/PX), MGET, MSET, INCR, DECR - **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 - **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 @@ -64,6 +65,12 @@ redis-cli ZADD board 100 alice 200 bob redis-cli ZRANGE board 0 -1 WITHSCORES redis-cli ZCARD board # => (integer) 2 +# hashes +redis-cli HSET user:1 name alice age 30 +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 + # 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 a9c9eb20..f65b2367 100644 --- a/crates/ember-core/src/keyspace.rs +++ b/crates/ember-core/src/keyspace.rs @@ -1065,6 +1065,307 @@ impl Keyspace { } } + // ------------------------------------------------------------------------- + // Hash operations + // ------------------------------------------------------------------------- + + /// Sets one or more field-value pairs in a hash. + /// + /// Creates the hash if the key doesn't exist. Returns the number of + /// new fields added (fields that were updated don't count). + pub fn hset(&mut self, key: &str, fields: &[(String, Bytes)]) -> Result { + if fields.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::Hash(_)) { + return Err(WriteError::WrongType); + } + + // estimate memory increase before mutating + let field_increase: usize = fields + .iter() + .map(|(f, v)| f.len() + v.len() + memory::HASHMAP_ENTRY_OVERHEAD) + .sum(); + let estimated_increase = if is_new { + memory::ENTRY_OVERHEAD + key.len() + memory::HASHMAP_BASE_OVERHEAD + field_increase + } else { + field_increase + }; + if !self.enforce_memory_limit(estimated_increase) { + return Err(WriteError::OutOfMemory); + } + + if is_new { + let value = Value::Hash(HashMap::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::Hash(ref mut map) = entry.value { + for (field, value) in fields { + if map.insert(field.clone(), value.clone()).is_none() { + 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) + } + + /// Gets the value of a field in a hash. + /// + /// Returns `None` if the key or field doesn't exist. + pub fn hget(&mut self, key: &str, field: &str) -> Result, WrongType> { + if self.remove_if_expired(key) { + return Ok(None); + } + match self.entries.get_mut(key) { + None => Ok(None), + Some(entry) => match &entry.value { + Value::Hash(map) => { + let result = map.get(field).cloned(); + entry.touch(); + Ok(result) + } + _ => Err(WrongType), + }, + } + } + + /// Gets all field-value pairs from a hash. + /// + /// Returns an empty vec if the key doesn't exist. + pub fn hgetall(&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::Hash(map) => { + let result: Vec<_> = map.iter().map(|(k, v)| (k.clone(), v.clone())).collect(); + entry.touch(); + Ok(result) + } + _ => Err(WrongType), + }, + } + } + + /// Deletes one or more fields from a hash. + /// + /// Returns the fields that were actually removed. + pub fn hdel(&mut self, key: &str, fields: &[String]) -> Result, WrongType> { + if self.remove_if_expired(key) { + return Ok(vec![]); + } + + match self.entries.get(key) { + None => return Ok(vec![]), + Some(e) => { + if !matches!(e.value, Value::Hash(_)) { + 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 = Vec::new(); + let is_empty = if let Value::Hash(ref mut map) = entry.value { + for field in fields { + if map.remove(field).is_some() { + removed.push(field.clone()); + } + } + map.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) + } + + /// Checks if a field exists in a hash. + pub fn hexists(&mut self, key: &str, field: &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::Hash(map) => { + let result = map.contains_key(field); + entry.touch(); + Ok(result) + } + _ => Err(WrongType), + }, + } + } + + /// Returns the number of fields in a hash. + pub fn hlen(&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::Hash(map) => Ok(map.len()), + _ => Err(WrongType), + }, + } + } + + /// Increments a field's integer value by the given amount. + /// + /// Creates the hash and field if they don't exist, starting from 0. + pub fn hincrby(&mut self, key: &str, field: &str, delta: i64) -> Result { + self.remove_if_expired(key); + + let is_new = !self.entries.contains_key(key); + + if !is_new { + match &self.entries[key].value { + Value::Hash(_) => {} + _ => return Err(IncrError::WrongType), + } + } + + // estimate memory for new field (worst case: new hash + new field) + let val_str_len = 20; // max i64 string length + let estimated_increase = if is_new { + memory::ENTRY_OVERHEAD + + key.len() + + memory::HASHMAP_BASE_OVERHEAD + + field.len() + + val_str_len + + memory::HASHMAP_ENTRY_OVERHEAD + } else { + field.len() + val_str_len + memory::HASHMAP_ENTRY_OVERHEAD + }; + + if !self.enforce_memory_limit(estimated_increase) { + return Err(IncrError::OutOfMemory); + } + + if is_new { + let value = Value::Hash(HashMap::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 new_val = if let Value::Hash(ref mut map) = entry.value { + let current_val = match map.get(field) { + Some(data) => { + let s = std::str::from_utf8(data).map_err(|_| IncrError::NotAnInteger)?; + s.parse::().map_err(|_| IncrError::NotAnInteger)? + } + None => 0, + }; + + let new_val = current_val.checked_add(delta).ok_or(IncrError::Overflow)?; + map.insert(field.to_owned(), Bytes::from(new_val.to_string())); + new_val + } else { + unreachable!() + }; + entry.touch(); + + let new_entry_size = memory::entry_size(key, &entry.value); + self.memory.adjust(old_entry_size, new_entry_size); + + Ok(new_val) + } + + /// Returns all field names in a hash. + pub fn hkeys(&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::Hash(map) => { + let result = map.keys().cloned().collect(); + entry.touch(); + Ok(result) + } + _ => Err(WrongType), + }, + } + } + + /// Returns all values in a hash. + pub fn hvals(&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::Hash(map) => { + let result = map.values().cloned().collect(); + entry.touch(); + Ok(result) + } + _ => Err(WrongType), + }, + } + } + + /// Gets multiple field values from a hash. + /// + /// Returns `None` for fields that don't exist. + pub fn hmget(&mut self, key: &str, fields: &[String]) -> Result>, WrongType> { + if self.remove_if_expired(key) { + return Ok(fields.iter().map(|_| None).collect()); + } + match self.entries.get_mut(key) { + None => Ok(fields.iter().map(|_| None).collect()), + Some(entry) => match &entry.value { + Value::Hash(map) => { + let result = fields.iter().map(|f| map.get(f).cloned()).collect(); + entry.touch(); + Ok(result) + } + _ => 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 @@ -2455,4 +2756,251 @@ mod tests { assert!(!super::glob_match("exact", "exactnot")); assert!(!super::glob_match("exact", "notexact")); } + + // --- hash tests --- + + #[test] + fn hset_creates_hash() { + let mut ks = Keyspace::new(); + let count = ks + .hset("h", &[("field1".into(), Bytes::from("value1"))]) + .unwrap(); + assert_eq!(count, 1); + assert_eq!(ks.value_type("h"), "hash"); + } + + #[test] + fn hset_returns_new_field_count() { + let mut ks = Keyspace::new(); + // add two new fields + let count = ks + .hset( + "h", + &[ + ("f1".into(), Bytes::from("v1")), + ("f2".into(), Bytes::from("v2")), + ], + ) + .unwrap(); + assert_eq!(count, 2); + + // update one, add one new + let count = ks + .hset( + "h", + &[ + ("f1".into(), Bytes::from("updated")), + ("f3".into(), Bytes::from("v3")), + ], + ) + .unwrap(); + assert_eq!(count, 1); // only f3 is new + } + + #[test] + fn hget_returns_value() { + let mut ks = Keyspace::new(); + ks.hset("h", &[("name".into(), Bytes::from("alice"))]) + .unwrap(); + let val = ks.hget("h", "name").unwrap(); + assert_eq!(val, Some(Bytes::from("alice"))); + } + + #[test] + fn hget_missing_field_returns_none() { + let mut ks = Keyspace::new(); + ks.hset("h", &[("a".into(), Bytes::from("1"))]).unwrap(); + assert_eq!(ks.hget("h", "b").unwrap(), None); + } + + #[test] + fn hget_missing_key_returns_none() { + let mut ks = Keyspace::new(); + assert_eq!(ks.hget("missing", "field").unwrap(), None); + } + + #[test] + fn hgetall_returns_all_fields() { + let mut ks = Keyspace::new(); + ks.hset( + "h", + &[ + ("a".into(), Bytes::from("1")), + ("b".into(), Bytes::from("2")), + ], + ) + .unwrap(); + let mut fields = ks.hgetall("h").unwrap(); + fields.sort_by(|a, b| a.0.cmp(&b.0)); + assert_eq!(fields.len(), 2); + assert_eq!(fields[0], ("a".into(), Bytes::from("1"))); + assert_eq!(fields[1], ("b".into(), Bytes::from("2"))); + } + + #[test] + fn hdel_removes_fields() { + let mut ks = Keyspace::new(); + ks.hset( + "h", + &[ + ("a".into(), Bytes::from("1")), + ("b".into(), Bytes::from("2")), + ("c".into(), Bytes::from("3")), + ], + ) + .unwrap(); + let removed = ks.hdel("h", &["a".into(), "c".into()]).unwrap(); + assert_eq!(removed.len(), 2); + assert!(removed.contains(&"a".into())); + assert!(removed.contains(&"c".into())); + assert_eq!(ks.hlen("h").unwrap(), 1); + } + + #[test] + fn hdel_auto_deletes_empty_hash() { + let mut ks = Keyspace::new(); + ks.hset("h", &[("only".into(), Bytes::from("field"))]) + .unwrap(); + ks.hdel("h", &["only".into()]).unwrap(); + assert_eq!(ks.value_type("h"), "none"); + } + + #[test] + fn hexists_returns_true_for_existing_field() { + let mut ks = Keyspace::new(); + ks.hset("h", &[("field".into(), Bytes::from("val"))]) + .unwrap(); + assert!(ks.hexists("h", "field").unwrap()); + } + + #[test] + fn hexists_returns_false_for_missing_field() { + let mut ks = Keyspace::new(); + ks.hset("h", &[("a".into(), Bytes::from("1"))]).unwrap(); + assert!(!ks.hexists("h", "missing").unwrap()); + } + + #[test] + fn hlen_returns_field_count() { + let mut ks = Keyspace::new(); + ks.hset( + "h", + &[ + ("a".into(), Bytes::from("1")), + ("b".into(), Bytes::from("2")), + ], + ) + .unwrap(); + assert_eq!(ks.hlen("h").unwrap(), 2); + } + + #[test] + fn hlen_missing_key_returns_zero() { + let mut ks = Keyspace::new(); + assert_eq!(ks.hlen("missing").unwrap(), 0); + } + + #[test] + fn hincrby_new_field() { + let mut ks = Keyspace::new(); + ks.hset("h", &[("x".into(), Bytes::from("ignored"))]) + .unwrap(); + let val = ks.hincrby("h", "counter", 5).unwrap(); + assert_eq!(val, 5); + } + + #[test] + fn hincrby_existing_field() { + let mut ks = Keyspace::new(); + ks.hset("h", &[("n".into(), Bytes::from("10"))]).unwrap(); + let val = ks.hincrby("h", "n", 3).unwrap(); + assert_eq!(val, 13); + } + + #[test] + fn hincrby_negative_delta() { + let mut ks = Keyspace::new(); + ks.hset("h", &[("n".into(), Bytes::from("10"))]).unwrap(); + let val = ks.hincrby("h", "n", -7).unwrap(); + assert_eq!(val, 3); + } + + #[test] + fn hincrby_non_integer_returns_error() { + let mut ks = Keyspace::new(); + ks.hset("h", &[("s".into(), Bytes::from("notanumber"))]) + .unwrap(); + assert_eq!( + ks.hincrby("h", "s", 1).unwrap_err(), + IncrError::NotAnInteger + ); + } + + #[test] + fn hkeys_returns_field_names() { + let mut ks = Keyspace::new(); + ks.hset( + "h", + &[ + ("alpha".into(), Bytes::from("1")), + ("beta".into(), Bytes::from("2")), + ], + ) + .unwrap(); + let mut keys = ks.hkeys("h").unwrap(); + keys.sort(); + assert_eq!(keys, vec!["alpha", "beta"]); + } + + #[test] + fn hvals_returns_values() { + let mut ks = Keyspace::new(); + ks.hset( + "h", + &[ + ("a".into(), Bytes::from("x")), + ("b".into(), Bytes::from("y")), + ], + ) + .unwrap(); + let mut vals = ks.hvals("h").unwrap(); + vals.sort(); + assert_eq!(vals, vec![Bytes::from("x"), Bytes::from("y")]); + } + + #[test] + fn hmget_returns_values_for_existing_fields() { + let mut ks = Keyspace::new(); + ks.hset( + "h", + &[ + ("a".into(), Bytes::from("1")), + ("b".into(), Bytes::from("2")), + ], + ) + .unwrap(); + let vals = ks + .hmget("h", &["a".into(), "missing".into(), "b".into()]) + .unwrap(); + assert_eq!(vals.len(), 3); + assert_eq!(vals[0], Some(Bytes::from("1"))); + assert_eq!(vals[1], None); + assert_eq!(vals[2], Some(Bytes::from("2"))); + } + + #[test] + fn hash_on_string_key_returns_wrongtype() { + let mut ks = Keyspace::new(); + ks.set("s".into(), Bytes::from("string"), None); + assert!(ks.hset("s", &[("f".into(), Bytes::from("v"))]).is_err()); + assert!(ks.hget("s", "f").is_err()); + assert!(ks.hgetall("s").is_err()); + assert!(ks.hdel("s", &["f".into()]).is_err()); + assert!(ks.hexists("s", "f").is_err()); + assert!(ks.hlen("s").is_err()); + assert!(ks.hincrby("s", "f", 1).is_err()); + assert!(ks.hkeys("s").is_err()); + assert!(ks.hvals("s").is_err()); + assert!(ks.hmget("s", &["f".into()]).is_err()); + } } diff --git a/crates/ember-core/src/memory.rs b/crates/ember-core/src/memory.rs index 320c6d79..42bbb21d 100644 --- a/crates/ember-core/src/memory.rs +++ b/crates/ember-core/src/memory.rs @@ -122,6 +122,15 @@ pub(crate) const VECDEQUE_ELEMENT_OVERHEAD: usize = 32; /// Base overhead for an empty VecDeque (internal buffer pointer + head/len). pub(crate) const VECDEQUE_BASE_OVERHEAD: usize = 24; +/// Estimated overhead per entry in a HashMap (for hash type). +/// +/// Each entry has: key String (24 bytes ptr+len+cap), value Bytes (24 bytes), +/// plus HashMap bucket overhead (~16 bytes for hash + next pointer). +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; + /// Returns the byte size of a value's payload. pub fn value_size(value: &Value) -> usize { match value { @@ -134,6 +143,13 @@ pub fn value_size(value: &Value) -> usize { VECDEQUE_BASE_OVERHEAD + element_bytes } Value::SortedSet(ss) => ss.memory_usage(), + Value::Hash(map) => { + let entry_bytes: usize = map + .iter() + .map(|(k, v)| k.len() + v.len() + HASHMAP_ENTRY_OVERHEAD) + .sum(); + HASHMAP_BASE_OVERHEAD + entry_bytes + } } } diff --git a/crates/ember-core/src/shard.rs b/crates/ember-core/src/shard.rs index 0939d8a9..2f875cf4 100644 --- a/crates/ember-core/src/shard.rs +++ b/crates/ember-core/src/shard.rs @@ -140,6 +140,43 @@ pub enum ShardRequest { stop: i64, with_scores: bool, }, + HSet { + key: String, + fields: Vec<(String, Bytes)>, + }, + HGet { + key: String, + field: String, + }, + HGetAll { + key: String, + }, + HDel { + key: String, + fields: Vec, + }, + HExists { + key: String, + field: String, + }, + HLen { + key: String, + }, + HIncrBy { + key: String, + field: String, + delta: i64, + }, + HKeys { + key: String, + }, + HVals { + key: String, + }, + HMGet { + key: String, + fields: Vec, + }, /// Returns the key count for this shard. DbSize, /// Returns keyspace stats for this shard. @@ -202,6 +239,14 @@ pub enum ShardResponse { Err(String), /// Scan result: next cursor and list of keys. Scan { cursor: u64, keys: Vec }, + /// HGETALL result: all field-value pairs. + HashFields(Vec<(String, Bytes)>), + /// HDEL result: removed count + field names for AOF. + HDelLen { count: usize, removed: Vec }, + /// Array of strings (e.g. HKEYS). + StringArray(Vec), + /// HMGET result: array of optional values. + OptionalArray(Vec>), } /// A request bundled with its reply channel. @@ -288,6 +333,7 @@ async fn run_shard( } Value::SortedSet(ss) } + RecoveredValue::Hash(map) => Value::Hash(map), }; keyspace.restore(entry.key, value, entry.expires_at); } @@ -556,6 +602,52 @@ fn dispatch(ks: &mut Keyspace, req: &ShardRequest) -> ShardResponse { keys, } } + ShardRequest::HSet { key, fields } => match ks.hset(key, fields) { + Ok(count) => ShardResponse::Len(count), + Err(WriteError::WrongType) => ShardResponse::WrongType, + Err(WriteError::OutOfMemory) => ShardResponse::OutOfMemory, + }, + ShardRequest::HGet { key, field } => match ks.hget(key, field) { + Ok(val) => ShardResponse::Value(val.map(Value::String)), + Err(_) => ShardResponse::WrongType, + }, + ShardRequest::HGetAll { key } => match ks.hgetall(key) { + Ok(fields) => ShardResponse::HashFields(fields), + Err(_) => ShardResponse::WrongType, + }, + ShardRequest::HDel { key, fields } => match ks.hdel(key, fields) { + Ok(removed) => ShardResponse::HDelLen { + count: removed.len(), + removed, + }, + Err(_) => ShardResponse::WrongType, + }, + ShardRequest::HExists { key, field } => match ks.hexists(key, field) { + Ok(exists) => ShardResponse::Bool(exists), + Err(_) => ShardResponse::WrongType, + }, + ShardRequest::HLen { key } => match ks.hlen(key) { + Ok(len) => ShardResponse::Len(len), + Err(_) => ShardResponse::WrongType, + }, + ShardRequest::HIncrBy { key, field, delta } => match ks.hincrby(key, field, *delta) { + Ok(val) => ShardResponse::Integer(val), + Err(IncrError::WrongType) => ShardResponse::WrongType, + Err(IncrError::OutOfMemory) => ShardResponse::OutOfMemory, + Err(e) => ShardResponse::Err(e.to_string()), + }, + ShardRequest::HKeys { key } => match ks.hkeys(key) { + Ok(keys) => ShardResponse::StringArray(keys), + Err(_) => ShardResponse::WrongType, + }, + ShardRequest::HVals { key } => match ks.hvals(key) { + Ok(vals) => ShardResponse::Array(vals), + Err(_) => ShardResponse::WrongType, + }, + ShardRequest::HMGet { key, fields } => match ks.hmget(key, fields) { + Ok(vals) => ShardResponse::OptionalArray(vals), + Err(_) => ShardResponse::WrongType, + }, // snapshot/rewrite are handled in the main loop, not here ShardRequest::Snapshot | ShardRequest::RewriteAof => ShardResponse::Ok, } @@ -711,6 +803,7 @@ fn write_snapshot( .collect(); SnapValue::SortedSet(members) } + Value::Hash(map) => SnapValue::Hash(map.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 49112536..fade45d9 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, and -//! sorted sets are supported; plain sets and hashes will come later. +//! Each variant maps to a Redis-like data type. Strings, lists, sorted +//! sets, and hashes are supported; plain sets will come later. pub mod sorted_set; -use std::collections::VecDeque; +use std::collections::{HashMap, VecDeque}; use bytes::Bytes; @@ -29,6 +29,10 @@ pub enum Value { /// Sorted set of unique string members, each with a float score. /// Members are ordered by (score, member_name). SortedSet(SortedSet), + + /// Hash map of field names to values. Fields are unique strings, + /// values are binary-safe byte sequences. + Hash(HashMap), } impl PartialEq for Value { @@ -42,6 +46,7 @@ impl PartialEq for Value { .zip(b.iter()) .all(|((m1, s1), (m2, s2))| m1 == m2 && s1 == s2) } + (Value::Hash(a), Value::Hash(b)) => a == b, _ => false, } } @@ -53,6 +58,7 @@ pub fn type_name(value: &Value) -> &'static str { Value::String(_) => "string", Value::List(_) => "list", Value::SortedSet(_) => "zset", + Value::Hash(_) => "hash", } } diff --git a/crates/ember-persistence/src/recovery.rs b/crates/ember-persistence/src/recovery.rs index a1636c4f..ae07a385 100644 --- a/crates/ember-persistence/src/recovery.rs +++ b/crates/ember-persistence/src/recovery.rs @@ -25,6 +25,8 @@ pub enum RecoveredValue { List(VecDeque), /// Sorted set stored as (score, member) pairs. SortedSet(Vec<(f64, String)>), + /// Hash map of field names to values. + Hash(HashMap), } impl From for RecoveredValue { @@ -33,6 +35,7 @@ impl From for RecoveredValue { SnapValue::String(data) => RecoveredValue::String(data), SnapValue::List(deque) => RecoveredValue::List(deque), SnapValue::SortedSet(members) => RecoveredValue::SortedSet(members), + SnapValue::Hash(map) => RecoveredValue::Hash(map), } } } diff --git a/crates/ember-persistence/src/snapshot.rs b/crates/ember-persistence/src/snapshot.rs index 63028314..5ebeaff0 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::VecDeque; +use std::collections::{HashMap, VecDeque}; use std::fs::{self, File}; use std::io::{self, BufReader, BufWriter, Write}; use std::path::{Path, PathBuf}; @@ -34,6 +34,7 @@ use crate::format::{self, FormatError}; const TYPE_STRING: u8 = 0; const TYPE_LIST: u8 = 1; const TYPE_SORTED_SET: u8 = 2; +const TYPE_HASH: u8 = 3; /// The value stored in a snapshot entry. #[derive(Debug, Clone, PartialEq)] @@ -44,6 +45,8 @@ pub enum SnapValue { List(VecDeque), /// A sorted set: vec of (score, member) pairs. SortedSet(Vec<(f64, String)>), + /// A hash: map of field names to values. + Hash(HashMap), } /// A single entry in a snapshot file. @@ -123,6 +126,14 @@ impl SnapshotWriter { format::write_bytes(&mut buf, member.as_bytes())?; } } + SnapValue::Hash(map) => { + format::write_u8(&mut buf, TYPE_HASH)?; + format::write_u32(&mut buf, map.len() as u32)?; + for (field, value) in map { + format::write_bytes(&mut buf, field.as_bytes())?; + format::write_bytes(&mut buf, value)?; + } + } } format::write_i64(&mut buf, entry.expire_ms)?; @@ -256,6 +267,25 @@ impl SnapshotReader { } SnapValue::SortedSet(members) } + TYPE_HASH => { + let count = format::read_u32(&mut self.reader)?; + format::write_u32(&mut buf, count).expect("vec write"); + let mut map = HashMap::with_capacity(count as usize); + for _ in 0..count { + let field_bytes = format::read_bytes(&mut self.reader)?; + format::write_bytes(&mut buf, &field_bytes).expect("vec write"); + let field = String::from_utf8(field_bytes).map_err(|_| { + FormatError::Io(io::Error::new( + io::ErrorKind::InvalidData, + "hash field is not valid utf-8", + )) + })?; + let value_bytes = format::read_bytes(&mut self.reader)?; + format::write_bytes(&mut buf, &value_bytes).expect("vec write"); + map.insert(field, Bytes::from(value_bytes)); + } + SnapValue::Hash(map) + } _ => { return Err(FormatError::UnknownTag(type_tag)); } diff --git a/crates/ember-protocol/src/command.rs b/crates/ember-protocol/src/command.rs index 1fccff18..224b2e27 100644 --- a/crates/ember-protocol/src/command.rs +++ b/crates/ember-protocol/src/command.rs @@ -144,6 +144,43 @@ pub enum Command { with_scores: bool, }, + /// HSET [field value ...]. Sets field-value pairs in a hash. + HSet { + key: String, + fields: Vec<(String, Bytes)>, + }, + + /// HGET . Gets a field's value from a hash. + HGet { key: String, field: String }, + + /// HGETALL . Gets all field-value pairs from a hash. + HGetAll { key: String }, + + /// HDEL [field ...]. Deletes fields from a hash. + HDel { key: String, fields: Vec }, + + /// HEXISTS . Checks if a field exists in a hash. + HExists { key: String, field: String }, + + /// HLEN . Returns the number of fields in a hash. + HLen { key: String }, + + /// HINCRBY . Increments a hash field's integer value. + HIncrBy { + key: String, + field: String, + delta: i64, + }, + + /// HKEYS . Returns all field names in a hash. + HKeys { key: String }, + + /// HVALS . Returns all values in a hash. + HVals { key: String }, + + /// HMGET [field ...]. Gets multiple field values from a hash. + HMGet { key: String, fields: Vec }, + /// A command we don't recognize (yet). Unknown(String), } @@ -224,6 +261,16 @@ impl Command { "ZRANK" => parse_zrank(&frames[1..]), "ZCARD" => parse_zcard(&frames[1..]), "ZRANGE" => parse_zrange(&frames[1..]), + "HSET" => parse_hset(&frames[1..]), + "HGET" => parse_hget(&frames[1..]), + "HGETALL" => parse_hgetall(&frames[1..]), + "HDEL" => parse_hdel(&frames[1..]), + "HEXISTS" => parse_hexists(&frames[1..]), + "HLEN" => parse_hlen(&frames[1..]), + "HINCRBY" => parse_hincrby(&frames[1..]), + "HKEYS" => parse_hkeys(&frames[1..]), + "HVALS" => parse_hvals(&frames[1..]), + "HMGET" => parse_hmget(&frames[1..]), _ => Ok(Command::Unknown(name)), } } @@ -792,6 +839,112 @@ fn parse_zrange(args: &[Frame]) -> Result { }) } +// --- hash commands --- + +fn parse_hset(args: &[Frame]) -> Result { + // HSET key field value [field value ...] + // args = [key, field, value, ...] + // Need at least 3 args, and after key we need pairs (so remaining count must be even) + if args.len() < 3 || !(args.len() - 1).is_multiple_of(2) { + return Err(ProtocolError::WrongArity("HSET".into())); + } + + let key = extract_string(&args[0])?; + let mut fields = Vec::with_capacity((args.len() - 1) / 2); + + for chunk in args[1..].chunks(2) { + let field = extract_string(&chunk[0])?; + let value = extract_bytes(&chunk[1])?; + fields.push((field, value)); + } + + Ok(Command::HSet { key, fields }) +} + +fn parse_hget(args: &[Frame]) -> Result { + if args.len() != 2 { + return Err(ProtocolError::WrongArity("HGET".into())); + } + let key = extract_string(&args[0])?; + let field = extract_string(&args[1])?; + Ok(Command::HGet { key, field }) +} + +fn parse_hgetall(args: &[Frame]) -> Result { + if args.len() != 1 { + return Err(ProtocolError::WrongArity("HGETALL".into())); + } + let key = extract_string(&args[0])?; + Ok(Command::HGetAll { key }) +} + +fn parse_hdel(args: &[Frame]) -> Result { + if args.len() < 2 { + return Err(ProtocolError::WrongArity("HDEL".into())); + } + let key = extract_string(&args[0])?; + let fields = args[1..] + .iter() + .map(extract_string) + .collect::, _>>()?; + Ok(Command::HDel { key, fields }) +} + +fn parse_hexists(args: &[Frame]) -> Result { + if args.len() != 2 { + return Err(ProtocolError::WrongArity("HEXISTS".into())); + } + let key = extract_string(&args[0])?; + let field = extract_string(&args[1])?; + Ok(Command::HExists { key, field }) +} + +fn parse_hlen(args: &[Frame]) -> Result { + if args.len() != 1 { + return Err(ProtocolError::WrongArity("HLEN".into())); + } + let key = extract_string(&args[0])?; + Ok(Command::HLen { key }) +} + +fn parse_hincrby(args: &[Frame]) -> Result { + if args.len() != 3 { + return Err(ProtocolError::WrongArity("HINCRBY".into())); + } + let key = extract_string(&args[0])?; + let field = extract_string(&args[1])?; + let delta = parse_i64(&args[2], "HINCRBY")?; + Ok(Command::HIncrBy { key, field, delta }) +} + +fn parse_hkeys(args: &[Frame]) -> Result { + if args.len() != 1 { + return Err(ProtocolError::WrongArity("HKEYS".into())); + } + let key = extract_string(&args[0])?; + Ok(Command::HKeys { key }) +} + +fn parse_hvals(args: &[Frame]) -> Result { + if args.len() != 1 { + return Err(ProtocolError::WrongArity("HVALS".into())); + } + let key = extract_string(&args[0])?; + Ok(Command::HVals { key }) +} + +fn parse_hmget(args: &[Frame]) -> Result { + if args.len() < 2 { + return Err(ProtocolError::WrongArity("HMGET".into())); + } + let key = extract_string(&args[0])?; + let fields = args[1..] + .iter() + .map(extract_string) + .collect::, _>>()?; + Ok(Command::HMGet { key, fields }) +} + #[cfg(test)] mod tests { use super::*; @@ -1973,4 +2126,192 @@ mod tests { let err = Command::from_frame(cmd(&["SCAN", "0", "COUNT"])).unwrap_err(); assert!(matches!(err, ProtocolError::WrongArity(_))); } + + // --- hash commands --- + + #[test] + fn hset_single_field() { + assert_eq!( + Command::from_frame(cmd(&["HSET", "h", "field", "value"])).unwrap(), + Command::HSet { + key: "h".into(), + fields: vec![("field".into(), Bytes::from("value"))], + }, + ); + } + + #[test] + fn hset_multiple_fields() { + let parsed = Command::from_frame(cmd(&["HSET", "h", "f1", "v1", "f2", "v2"])).unwrap(); + match parsed { + Command::HSet { key, fields } => { + assert_eq!(key, "h"); + assert_eq!(fields.len(), 2); + } + other => panic!("expected HSet, got {other:?}"), + } + } + + #[test] + fn hset_wrong_arity() { + let err = Command::from_frame(cmd(&["HSET", "h"])).unwrap_err(); + assert!(matches!(err, ProtocolError::WrongArity(_))); + let err = Command::from_frame(cmd(&["HSET", "h", "f"])).unwrap_err(); + assert!(matches!(err, ProtocolError::WrongArity(_))); + } + + #[test] + fn hget_basic() { + assert_eq!( + Command::from_frame(cmd(&["HGET", "h", "field"])).unwrap(), + Command::HGet { + key: "h".into(), + field: "field".into(), + }, + ); + } + + #[test] + fn hget_wrong_arity() { + let err = Command::from_frame(cmd(&["HGET", "h"])).unwrap_err(); + assert!(matches!(err, ProtocolError::WrongArity(_))); + } + + #[test] + fn hgetall_basic() { + assert_eq!( + Command::from_frame(cmd(&["HGETALL", "h"])).unwrap(), + Command::HGetAll { key: "h".into() }, + ); + } + + #[test] + fn hgetall_wrong_arity() { + let err = Command::from_frame(cmd(&["HGETALL"])).unwrap_err(); + assert!(matches!(err, ProtocolError::WrongArity(_))); + } + + #[test] + fn hdel_single() { + assert_eq!( + Command::from_frame(cmd(&["HDEL", "h", "f"])).unwrap(), + Command::HDel { + key: "h".into(), + fields: vec!["f".into()], + }, + ); + } + + #[test] + fn hdel_multiple() { + let parsed = Command::from_frame(cmd(&["HDEL", "h", "f1", "f2", "f3"])).unwrap(); + match parsed { + Command::HDel { fields, .. } => assert_eq!(fields.len(), 3), + other => panic!("expected HDel, got {other:?}"), + } + } + + #[test] + fn hdel_wrong_arity() { + let err = Command::from_frame(cmd(&["HDEL", "h"])).unwrap_err(); + assert!(matches!(err, ProtocolError::WrongArity(_))); + } + + #[test] + fn hexists_basic() { + assert_eq!( + Command::from_frame(cmd(&["HEXISTS", "h", "f"])).unwrap(), + Command::HExists { + key: "h".into(), + field: "f".into(), + }, + ); + } + + #[test] + fn hlen_basic() { + assert_eq!( + Command::from_frame(cmd(&["HLEN", "h"])).unwrap(), + Command::HLen { key: "h".into() }, + ); + } + + #[test] + fn hincrby_basic() { + assert_eq!( + Command::from_frame(cmd(&["HINCRBY", "h", "f", "5"])).unwrap(), + Command::HIncrBy { + key: "h".into(), + field: "f".into(), + delta: 5, + }, + ); + } + + #[test] + fn hincrby_negative() { + assert_eq!( + Command::from_frame(cmd(&["HINCRBY", "h", "f", "-3"])).unwrap(), + Command::HIncrBy { + key: "h".into(), + field: "f".into(), + delta: -3, + }, + ); + } + + #[test] + fn hincrby_wrong_arity() { + let err = Command::from_frame(cmd(&["HINCRBY", "h", "f"])).unwrap_err(); + assert!(matches!(err, ProtocolError::WrongArity(_))); + } + + #[test] + fn hkeys_basic() { + assert_eq!( + Command::from_frame(cmd(&["HKEYS", "h"])).unwrap(), + Command::HKeys { key: "h".into() }, + ); + } + + #[test] + fn hvals_basic() { + assert_eq!( + Command::from_frame(cmd(&["HVALS", "h"])).unwrap(), + Command::HVals { key: "h".into() }, + ); + } + + #[test] + fn hmget_basic() { + assert_eq!( + Command::from_frame(cmd(&["HMGET", "h", "f1", "f2"])).unwrap(), + Command::HMGet { + key: "h".into(), + fields: vec!["f1".into(), "f2".into()], + }, + ); + } + + #[test] + fn hmget_wrong_arity() { + let err = Command::from_frame(cmd(&["HMGET", "h"])).unwrap_err(); + assert!(matches!(err, ProtocolError::WrongArity(_))); + } + + #[test] + fn hash_commands_case_insensitive() { + assert!(matches!( + Command::from_frame(cmd(&["hset", "h", "f", "v"])).unwrap(), + Command::HSet { .. } + )); + assert!(matches!( + Command::from_frame(cmd(&["hget", "h", "f"])).unwrap(), + Command::HGet { .. } + )); + assert!(matches!( + Command::from_frame(cmd(&["hgetall", "h"])).unwrap(), + Command::HGetAll { .. } + )); + } } diff --git a/crates/ember-server/src/connection.rs b/crates/ember-server/src/connection.rs index c50e2d5b..ca9e1efd 100644 --- a/crates/ember-server/src/connection.rs +++ b/crates/ember-server/src/connection.rs @@ -624,6 +624,150 @@ async fn execute(cmd: Command, engine: &Engine) -> Frame { } } + // --- hash commands --- + Command::HSet { key, fields } => { + let req = ShardRequest::HSet { + key: key.clone(), + fields, + }; + 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::HGet { key, field } => { + let req = ShardRequest::HGet { + key: key.clone(), + field, + }; + match engine.route(&key, req).await { + Ok(ShardResponse::Value(Some(Value::String(data)))) => Frame::Bulk(data), + Ok(ShardResponse::Value(None)) => Frame::Null, + Ok(ShardResponse::WrongType) => wrongtype_error(), + Ok(other) => Frame::Error(format!("ERR unexpected shard response: {other:?}")), + Err(e) => Frame::Error(format!("ERR {e}")), + } + } + + Command::HGetAll { key } => { + let req = ShardRequest::HGetAll { key: key.clone() }; + match engine.route(&key, req).await { + Ok(ShardResponse::HashFields(fields)) => { + let mut frames = Vec::with_capacity(fields.len() * 2); + for (field, value) in fields { + frames.push(Frame::Bulk(Bytes::from(field))); + frames.push(Frame::Bulk(value)); + } + Frame::Array(frames) + } + Ok(ShardResponse::WrongType) => wrongtype_error(), + Ok(other) => Frame::Error(format!("ERR unexpected shard response: {other:?}")), + Err(e) => Frame::Error(format!("ERR {e}")), + } + } + + Command::HDel { key, fields } => { + let req = ShardRequest::HDel { + key: key.clone(), + fields, + }; + match engine.route(&key, req).await { + Ok(ShardResponse::HDelLen { count, .. }) => Frame::Integer(count 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::HExists { key, field } => { + let req = ShardRequest::HExists { + key: key.clone(), + field, + }; + 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::HLen { key } => { + let req = ShardRequest::HLen { 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::HIncrBy { key, field, delta } => { + let req = ShardRequest::HIncrBy { + key: key.clone(), + field, + delta, + }; + match engine.route(&key, req).await { + Ok(ShardResponse::Integer(n)) => Frame::Integer(n), + Ok(ShardResponse::WrongType) => wrongtype_error(), + Ok(ShardResponse::OutOfMemory) => oom_error(), + Ok(ShardResponse::Err(msg)) => Frame::Error(format!("ERR {msg}")), + Ok(other) => Frame::Error(format!("ERR unexpected shard response: {other:?}")), + Err(e) => Frame::Error(format!("ERR {e}")), + } + } + + Command::HKeys { key } => { + let req = ShardRequest::HKeys { key: key.clone() }; + match engine.route(&key, req).await { + Ok(ShardResponse::StringArray(keys)) => Frame::Array( + keys.into_iter() + .map(|k| Frame::Bulk(Bytes::from(k))) + .collect(), + ), + Ok(ShardResponse::WrongType) => wrongtype_error(), + Ok(other) => Frame::Error(format!("ERR unexpected shard response: {other:?}")), + Err(e) => Frame::Error(format!("ERR {e}")), + } + } + + Command::HVals { key } => { + let req = ShardRequest::HVals { key: key.clone() }; + match engine.route(&key, req).await { + Ok(ShardResponse::Array(vals)) => { + Frame::Array(vals.into_iter().map(Frame::Bulk).collect()) + } + Ok(ShardResponse::WrongType) => wrongtype_error(), + Ok(other) => Frame::Error(format!("ERR unexpected shard response: {other:?}")), + Err(e) => Frame::Error(format!("ERR {e}")), + } + } + + Command::HMGet { key, fields } => { + let req = ShardRequest::HMGet { + key: key.clone(), + fields, + }; + match engine.route(&key, req).await { + Ok(ShardResponse::OptionalArray(vals)) => Frame::Array( + vals.into_iter() + .map(|v| match v { + Some(data) => Frame::Bulk(data), + None => Frame::Null, + }) + .collect(), + ), + 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}'")), } }