From 73455d232d36ee727d2c0314fb3c6057d0ddf7ed Mon Sep 17 00:00:00 2001 From: Kacy Fortner Date: Wed, 11 Feb 2026 22:38:42 -0500 Subject: [PATCH 1/6] refactor: move capped_capacity to shared format module identical function defined in both aof.rs and snapshot.rs. now lives in format.rs as a pub fn, both callers import from there. --- crates/ember-persistence/src/aof.rs | 21 +++++++-------------- crates/ember-persistence/src/format.rs | 7 +++++++ crates/ember-persistence/src/snapshot.rs | 23 ++++++++--------------- 3 files changed, 22 insertions(+), 29 deletions(-) diff --git a/crates/ember-persistence/src/aof.rs b/crates/ember-persistence/src/aof.rs index b667f678..fb3fefb2 100644 --- a/crates/ember-persistence/src/aof.rs +++ b/crates/ember-persistence/src/aof.rs @@ -304,13 +304,6 @@ impl AofRecord { Ok(buf) } - /// Cap pre-allocation to avoid huge allocations from corrupt count fields. - /// The loop will still iterate `count` times — this just limits the - /// up-front reservation so a bogus u32 can't exhaust memory. - fn capped_capacity(count: u32) -> usize { - (count as usize).min(65_536) - } - /// Deserializes a record from a byte slice (tag + payload, no CRC). fn from_bytes(data: &[u8]) -> Result { let mut cursor = io::Cursor::new(data); @@ -338,7 +331,7 @@ impl AofRecord { TAG_LPUSH | TAG_RPUSH => { let key = read_string(&mut cursor, "key")?; let count = format::read_u32(&mut cursor)?; - let mut values = Vec::with_capacity(Self::capped_capacity(count)); + let mut values = Vec::with_capacity(format::capped_capacity(count)); for _ in 0..count { values.push(Bytes::from(format::read_bytes(&mut cursor)?)); } @@ -359,7 +352,7 @@ impl AofRecord { TAG_ZADD => { let key = read_string(&mut cursor, "key")?; let count = format::read_u32(&mut cursor)?; - let mut members = Vec::with_capacity(Self::capped_capacity(count)); + let mut members = Vec::with_capacity(format::capped_capacity(count)); for _ in 0..count { let score = format::read_f64(&mut cursor)?; let member = read_string(&mut cursor, "member")?; @@ -370,7 +363,7 @@ impl AofRecord { TAG_ZREM => { let key = read_string(&mut cursor, "key")?; let count = format::read_u32(&mut cursor)?; - let mut members = Vec::with_capacity(Self::capped_capacity(count)); + let mut members = Vec::with_capacity(format::capped_capacity(count)); for _ in 0..count { members.push(read_string(&mut cursor, "member")?); } @@ -396,7 +389,7 @@ impl AofRecord { TAG_HSET => { let key = read_string(&mut cursor, "key")?; let count = format::read_u32(&mut cursor)?; - let mut fields = Vec::with_capacity(Self::capped_capacity(count)); + let mut fields = Vec::with_capacity(format::capped_capacity(count)); for _ in 0..count { let field = read_string(&mut cursor, "field")?; let value = Bytes::from(format::read_bytes(&mut cursor)?); @@ -407,7 +400,7 @@ impl AofRecord { TAG_HDEL => { let key = read_string(&mut cursor, "key")?; let count = format::read_u32(&mut cursor)?; - let mut fields = Vec::with_capacity(Self::capped_capacity(count)); + let mut fields = Vec::with_capacity(format::capped_capacity(count)); for _ in 0..count { fields.push(read_string(&mut cursor, "field")?); } @@ -422,7 +415,7 @@ impl AofRecord { TAG_SADD => { let key = read_string(&mut cursor, "key")?; let count = format::read_u32(&mut cursor)?; - let mut members = Vec::with_capacity(Self::capped_capacity(count)); + let mut members = Vec::with_capacity(format::capped_capacity(count)); for _ in 0..count { members.push(read_string(&mut cursor, "member")?); } @@ -431,7 +424,7 @@ impl AofRecord { TAG_SREM => { let key = read_string(&mut cursor, "key")?; let count = format::read_u32(&mut cursor)?; - let mut members = Vec::with_capacity(Self::capped_capacity(count)); + let mut members = Vec::with_capacity(format::capped_capacity(count)); for _ in 0..count { members.push(read_string(&mut cursor, "member")?); } diff --git a/crates/ember-persistence/src/format.rs b/crates/ember-persistence/src/format.rs index 7e1cb7de..1b6a75cc 100644 --- a/crates/ember-persistence/src/format.rs +++ b/crates/ember-persistence/src/format.rs @@ -209,6 +209,13 @@ pub fn verify_crc32(data: &[u8], expected: u32) -> Result<(), FormatError> { verify_crc32_values(actual, expected) } +/// Caps pre-allocation to avoid huge allocations from corrupt count fields. +/// The loop will still iterate `count` times — this just limits the +/// up-front reservation so a bogus u32 can't exhaust memory. +pub fn capped_capacity(count: u32) -> usize { + (count as usize).min(65_536) +} + /// Verifies that two CRC32 values match. pub fn verify_crc32_values(computed: u32, stored: u32) -> Result<(), FormatError> { if computed != stored { diff --git a/crates/ember-persistence/src/snapshot.rs b/crates/ember-persistence/src/snapshot.rs index 1d9f8f36..6814fee6 100644 --- a/crates/ember-persistence/src/snapshot.rs +++ b/crates/ember-persistence/src/snapshot.rs @@ -39,13 +39,6 @@ const TYPE_SET: u8 = 4; #[cfg(feature = "protobuf")] const TYPE_PROTO: u8 = 5; -/// Cap pre-allocation to avoid huge allocations from corrupt count fields. -/// The loop will still iterate `count` times — this just limits the -/// up-front reservation so a bogus u32 can't exhaust memory. -fn capped_capacity(count: u32) -> usize { - (count as usize).min(65_536) -} - /// The value stored in a snapshot entry. #[derive(Debug, Clone, PartialEq)] pub enum SnapValue { @@ -382,7 +375,7 @@ impl SnapshotReader { TYPE_LIST => { let count = format::read_u32(&mut self.reader)?; format::write_u32(&mut buf, count)?; - let mut deque = VecDeque::with_capacity(capped_capacity(count)); + let mut deque = VecDeque::with_capacity(format::capped_capacity(count)); for _ in 0..count { let item = format::read_bytes(&mut self.reader)?; format::write_bytes(&mut buf, &item)?; @@ -393,7 +386,7 @@ impl SnapshotReader { TYPE_SORTED_SET => { let count = format::read_u32(&mut self.reader)?; format::write_u32(&mut buf, count)?; - let mut members = Vec::with_capacity(capped_capacity(count)); + let mut members = Vec::with_capacity(format::capped_capacity(count)); for _ in 0..count { let score = format::read_f64(&mut self.reader)?; format::write_f64(&mut buf, score)?; @@ -412,7 +405,7 @@ impl SnapshotReader { TYPE_HASH => { let count = format::read_u32(&mut self.reader)?; format::write_u32(&mut buf, count)?; - let mut map = HashMap::with_capacity(capped_capacity(count)); + let mut map = HashMap::with_capacity(format::capped_capacity(count)); for _ in 0..count { let field_bytes = format::read_bytes(&mut self.reader)?; format::write_bytes(&mut buf, &field_bytes)?; @@ -431,7 +424,7 @@ impl SnapshotReader { TYPE_SET => { let count = format::read_u32(&mut self.reader)?; format::write_u32(&mut buf, count)?; - let mut set = HashSet::with_capacity(capped_capacity(count)); + let mut set = HashSet::with_capacity(format::capped_capacity(count)); for _ in 0..count { let member_bytes = format::read_bytes(&mut self.reader)?; format::write_bytes(&mut buf, &member_bytes)?; @@ -541,7 +534,7 @@ impl SnapshotReader { } TYPE_LIST => { let count = format::read_u32(&mut cursor)?; - let mut deque = VecDeque::with_capacity(capped_capacity(count)); + let mut deque = VecDeque::with_capacity(format::capped_capacity(count)); for _ in 0..count { deque.push_back(Bytes::from(format::read_bytes(&mut cursor)?)); } @@ -549,7 +542,7 @@ impl SnapshotReader { } TYPE_SORTED_SET => { let count = format::read_u32(&mut cursor)?; - let mut members = Vec::with_capacity(capped_capacity(count)); + let mut members = Vec::with_capacity(format::capped_capacity(count)); for _ in 0..count { let score = format::read_f64(&mut cursor)?; let member_bytes = format::read_bytes(&mut cursor)?; @@ -565,7 +558,7 @@ impl SnapshotReader { } TYPE_HASH => { let count = format::read_u32(&mut cursor)?; - let mut map = HashMap::with_capacity(capped_capacity(count)); + let mut map = HashMap::with_capacity(format::capped_capacity(count)); for _ in 0..count { let field_bytes = format::read_bytes(&mut cursor)?; let field = String::from_utf8(field_bytes).map_err(|_| { @@ -581,7 +574,7 @@ impl SnapshotReader { } TYPE_SET => { let count = format::read_u32(&mut cursor)?; - let mut set = HashSet::with_capacity(capped_capacity(count)); + let mut set = HashSet::with_capacity(format::capped_capacity(count)); for _ in 0..count { let member_bytes = format::read_bytes(&mut cursor)?; let member = String::from_utf8(member_bytes).map_err(|_| { From 55fd45f75fb1c4a946933ff286dfc7eade0d7184 Mon Sep 17 00:00:00 2001 From: Kacy Fortner Date: Wed, 11 Feb 2026 22:39:07 -0500 Subject: [PATCH 2/6] docs: group aof tag constants by data type 24 flat constants reorganized into logical sections: string, list, sorted set, hash, set, key lifecycle, protobuf. values unchanged. --- crates/ember-persistence/src/aof.rs | 35 ++++++++++++++++++++--------- 1 file changed, 25 insertions(+), 10 deletions(-) diff --git a/crates/ember-persistence/src/aof.rs b/crates/ember-persistence/src/aof.rs index fb3fefb2..f2344d02 100644 --- a/crates/ember-persistence/src/aof.rs +++ b/crates/ember-persistence/src/aof.rs @@ -38,29 +38,44 @@ fn read_string(r: &mut impl io::Read, field: &str) -> Result Date: Wed, 11 Feb 2026 22:39:49 -0500 Subject: [PATCH 3/6] refactor: extract read_string_list helper in aof deserialization ZREM, HDEL, SADD, and SREM all did identical work: read count, alloc vec with capped capacity, loop read_string. shared helper replaces ~7 lines per call site with ~3. --- crates/ember-persistence/src/aof.rs | 35 +++++++++++++---------------- 1 file changed, 15 insertions(+), 20 deletions(-) diff --git a/crates/ember-persistence/src/aof.rs b/crates/ember-persistence/src/aof.rs index f2344d02..c7317a7b 100644 --- a/crates/ember-persistence/src/aof.rs +++ b/crates/ember-persistence/src/aof.rs @@ -38,6 +38,17 @@ fn read_string(r: &mut impl io::Read, field: &str) -> Result Result, FormatError> { + let count = format::read_u32(r)?; + let mut items = Vec::with_capacity(format::capped_capacity(count)); + for _ in 0..count { + items.push(read_string(r, field)?); + } + Ok(items) +} + // -- record tags -- // values are stable and must not change (on-disk format). @@ -377,11 +388,7 @@ impl AofRecord { } TAG_ZREM => { let key = read_string(&mut cursor, "key")?; - let count = format::read_u32(&mut cursor)?; - let mut members = Vec::with_capacity(format::capped_capacity(count)); - for _ in 0..count { - members.push(read_string(&mut cursor, "member")?); - } + let members = read_string_list(&mut cursor, "member")?; Ok(AofRecord::ZRem { key, members }) } TAG_PERSIST => { @@ -414,11 +421,7 @@ impl AofRecord { } TAG_HDEL => { let key = read_string(&mut cursor, "key")?; - let count = format::read_u32(&mut cursor)?; - let mut fields = Vec::with_capacity(format::capped_capacity(count)); - for _ in 0..count { - fields.push(read_string(&mut cursor, "field")?); - } + let fields = read_string_list(&mut cursor, "field")?; Ok(AofRecord::HDel { key, fields }) } TAG_HINCRBY => { @@ -429,20 +432,12 @@ impl AofRecord { } TAG_SADD => { let key = read_string(&mut cursor, "key")?; - let count = format::read_u32(&mut cursor)?; - let mut members = Vec::with_capacity(format::capped_capacity(count)); - for _ in 0..count { - members.push(read_string(&mut cursor, "member")?); - } + let members = read_string_list(&mut cursor, "member")?; Ok(AofRecord::SAdd { key, members }) } TAG_SREM => { let key = read_string(&mut cursor, "key")?; - let count = format::read_u32(&mut cursor)?; - let mut members = Vec::with_capacity(format::capped_capacity(count)); - for _ in 0..count { - members.push(read_string(&mut cursor, "member")?); - } + let members = read_string_list(&mut cursor, "member")?; Ok(AofRecord::SRem { key, members }) } TAG_INCRBY => { From 9fd4b67d31c9f876a8ada1a241f41bc78d2eec57 Mon Sep 17 00:00:00 2001 From: Kacy Fortner Date: Wed, 11 Feb 2026 22:41:07 -0500 Subject: [PATCH 4/6] refactor: extract parse_snap_value for encrypted snapshot reading MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit type-tag → SnapValue parsing was duplicated between read_plaintext_entry and read_encrypted_entry. extracted the clean version as parse_snap_value, used in the encrypted path. plaintext path stays inline because it interleaves CRC buffer mirroring. --- crates/ember-persistence/src/snapshot.rs | 164 +++++++++++------------ 1 file changed, 77 insertions(+), 87 deletions(-) diff --git a/crates/ember-persistence/src/snapshot.rs b/crates/ember-persistence/src/snapshot.rs index 6814fee6..20291dd2 100644 --- a/crates/ember-persistence/src/snapshot.rs +++ b/crates/ember-persistence/src/snapshot.rs @@ -39,6 +39,81 @@ const TYPE_SET: u8 = 4; #[cfg(feature = "protobuf")] const TYPE_PROTO: u8 = 5; +/// Reads a UTF-8 string from a length-prefixed byte field. +#[cfg(feature = "encryption")] +fn read_snap_string(r: &mut impl io::Read, field: &str) -> Result { + let bytes = format::read_bytes(r)?; + String::from_utf8(bytes).map_err(|_| { + FormatError::Io(io::Error::new( + io::ErrorKind::InvalidData, + format!("{field} is not valid utf-8"), + )) + }) +} + +/// Parses a type-tagged SnapValue from a reader (v2+ format). +/// +/// Used by `read_encrypted_entry` to parse the `[type_tag][payload]` +/// portion of a decrypted entry. The plaintext path has parallel logic +/// but interleaves CRC buffer mirroring, so it stays inline. +#[cfg(feature = "encryption")] +fn parse_snap_value(r: &mut impl io::Read) -> Result { + let type_tag = format::read_u8(r)?; + match type_tag { + TYPE_STRING => { + let v = format::read_bytes(r)?; + Ok(SnapValue::String(Bytes::from(v))) + } + TYPE_LIST => { + let count = format::read_u32(r)?; + let mut deque = VecDeque::with_capacity(format::capped_capacity(count)); + for _ in 0..count { + deque.push_back(Bytes::from(format::read_bytes(r)?)); + } + Ok(SnapValue::List(deque)) + } + TYPE_SORTED_SET => { + let count = format::read_u32(r)?; + let mut members = Vec::with_capacity(format::capped_capacity(count)); + for _ in 0..count { + let score = format::read_f64(r)?; + let member = read_snap_string(r, "member")?; + members.push((score, member)); + } + Ok(SnapValue::SortedSet(members)) + } + TYPE_HASH => { + let count = format::read_u32(r)?; + let mut map = HashMap::with_capacity(format::capped_capacity(count)); + for _ in 0..count { + let field = read_snap_string(r, "hash field")?; + let value = format::read_bytes(r)?; + map.insert(field, Bytes::from(value)); + } + Ok(SnapValue::Hash(map)) + } + TYPE_SET => { + let count = format::read_u32(r)?; + let mut set = HashSet::with_capacity(format::capped_capacity(count)); + for _ in 0..count { + let member = read_snap_string(r, "set member")?; + set.insert(member); + } + Ok(SnapValue::Set(set)) + } + #[cfg(feature = "protobuf")] + TYPE_PROTO => { + let type_name = read_snap_string(r, "proto type_name")?; + let data = format::read_bytes(r)?; + Ok(SnapValue::Proto { + type_name, + data: Bytes::from(data), + }) + } + _ => Err(FormatError::UnknownTag(type_tag)), + } +} + /// The value stored in a snapshot entry. #[derive(Debug, Clone, PartialEq)] pub enum SnapValue { @@ -523,96 +598,11 @@ impl SnapshotReader { let plaintext = crate::encryption::decrypt_record(key, &nonce, &ciphertext)?; - // parse the decrypted bytes using the same logic as v2 let mut cursor = io::Cursor::new(&plaintext); - let key_bytes = format::read_bytes(&mut cursor)?; - let type_tag = format::read_u8(&mut cursor)?; - let value = match type_tag { - TYPE_STRING => { - let v = format::read_bytes(&mut cursor)?; - SnapValue::String(Bytes::from(v)) - } - TYPE_LIST => { - let count = format::read_u32(&mut cursor)?; - let mut deque = VecDeque::with_capacity(format::capped_capacity(count)); - for _ in 0..count { - deque.push_back(Bytes::from(format::read_bytes(&mut cursor)?)); - } - SnapValue::List(deque) - } - TYPE_SORTED_SET => { - let count = format::read_u32(&mut cursor)?; - let mut members = Vec::with_capacity(format::capped_capacity(count)); - for _ in 0..count { - let score = format::read_f64(&mut cursor)?; - let member_bytes = format::read_bytes(&mut cursor)?; - let member = String::from_utf8(member_bytes).map_err(|_| { - FormatError::Io(io::Error::new( - io::ErrorKind::InvalidData, - "member is not valid utf-8", - )) - })?; - members.push((score, member)); - } - SnapValue::SortedSet(members) - } - TYPE_HASH => { - let count = format::read_u32(&mut cursor)?; - let mut map = HashMap::with_capacity(format::capped_capacity(count)); - for _ in 0..count { - let field_bytes = format::read_bytes(&mut cursor)?; - 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 cursor)?; - map.insert(field, Bytes::from(value_bytes)); - } - SnapValue::Hash(map) - } - TYPE_SET => { - let count = format::read_u32(&mut cursor)?; - let mut set = HashSet::with_capacity(format::capped_capacity(count)); - for _ in 0..count { - let member_bytes = format::read_bytes(&mut cursor)?; - 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) - } - #[cfg(feature = "protobuf")] - TYPE_PROTO => { - let type_name_bytes = format::read_bytes(&mut cursor)?; - let type_name = String::from_utf8(type_name_bytes).map_err(|_| { - FormatError::Io(io::Error::new( - io::ErrorKind::InvalidData, - "proto type_name is not valid utf-8", - )) - })?; - let data = format::read_bytes(&mut cursor)?; - SnapValue::Proto { - type_name, - data: Bytes::from(data), - } - } - _ => return Err(FormatError::UnknownTag(type_tag)), - }; + let entry_key = read_snap_string(&mut cursor, "key")?; + let value = parse_snap_value(&mut cursor)?; let expire_ms = format::read_i64(&mut cursor)?; - let entry_key = String::from_utf8(key_bytes).map_err(|_| { - FormatError::Io(io::Error::new( - io::ErrorKind::InvalidData, - "key is not valid utf-8", - )) - })?; - self.read_so_far += 1; Ok(Some(SnapEntry { key: entry_key, From 3d6de5f909e5c25a3f3aee7bd1ed9bacc7279cc0 Mon Sep 17 00:00:00 2001 From: Kacy Fortner Date: Wed, 11 Feb 2026 22:42:38 -0500 Subject: [PATCH 5/6] refactor: extract incr/write error conversion helpers in shard dispatch MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit incr_result() handles the IncrError → ShardResponse pattern (5 call sites), write_result_len() handles WriteError → Len (4 call sites). each call site shrinks from 4-5 lines to 1. --- crates/ember-core/src/shard.rs | 84 ++++++++++++---------------------- 1 file changed, 29 insertions(+), 55 deletions(-) diff --git a/crates/ember-core/src/shard.rs b/crates/ember-core/src/shard.rs index 68c1ed76..8bdbc5cf 100644 --- a/crates/ember-core/src/shard.rs +++ b/crates/ember-core/src/shard.rs @@ -656,6 +656,25 @@ fn describe_request(req: &ShardRequest) -> RequestKind { } } +/// Converts an `IncrError` result into a `ShardResponse::Integer`. +fn incr_result(result: Result) -> ShardResponse { + match result { + Ok(val) => ShardResponse::Integer(val), + Err(IncrError::WrongType) => ShardResponse::WrongType, + Err(IncrError::OutOfMemory) => ShardResponse::OutOfMemory, + Err(e) => ShardResponse::Err(e.to_string()), + } +} + +/// Converts a `WriteError` result into a `ShardResponse::Len`. +fn write_result_len(result: Result) -> ShardResponse { + match result { + Ok(len) => ShardResponse::Len(len), + Err(WriteError::WrongType) => ShardResponse::WrongType, + Err(WriteError::OutOfMemory) => ShardResponse::OutOfMemory, + } +} + /// Executes a single request against the keyspace. fn dispatch( ks: &mut Keyspace, @@ -687,31 +706,11 @@ fn dispatch( SetResult::OutOfMemory => ShardResponse::OutOfMemory, } } - ShardRequest::Incr { key } => match ks.incr(key) { - Ok(val) => ShardResponse::Integer(val), - Err(IncrError::WrongType) => ShardResponse::WrongType, - Err(IncrError::OutOfMemory) => ShardResponse::OutOfMemory, - Err(e) => ShardResponse::Err(e.to_string()), - }, - ShardRequest::Decr { key } => match ks.decr(key) { - Ok(val) => ShardResponse::Integer(val), - Err(IncrError::WrongType) => ShardResponse::WrongType, - Err(IncrError::OutOfMemory) => ShardResponse::OutOfMemory, - Err(e) => ShardResponse::Err(e.to_string()), - }, - ShardRequest::IncrBy { key, delta } => match ks.incr_by(key, *delta) { - Ok(val) => ShardResponse::Integer(val), - Err(IncrError::WrongType) => ShardResponse::WrongType, - Err(IncrError::OutOfMemory) => ShardResponse::OutOfMemory, - Err(e) => ShardResponse::Err(e.to_string()), - }, + ShardRequest::Incr { key } => incr_result(ks.incr(key)), + ShardRequest::Decr { key } => incr_result(ks.decr(key)), + ShardRequest::IncrBy { key, delta } => incr_result(ks.incr_by(key, *delta)), ShardRequest::DecrBy { key, delta } => match delta.checked_neg() { - Some(neg) => match ks.incr_by(key, neg) { - Ok(val) => ShardResponse::Integer(val), - Err(IncrError::WrongType) => ShardResponse::WrongType, - Err(IncrError::OutOfMemory) => ShardResponse::OutOfMemory, - Err(e) => ShardResponse::Err(e.to_string()), - }, + Some(neg) => incr_result(ks.incr_by(key, neg)), None => ShardResponse::Err("ERR increment or decrement would overflow".into()), }, ShardRequest::IncrByFloat { key, delta } => match ks.incr_by_float(key, *delta) { @@ -720,11 +719,7 @@ fn dispatch( Err(IncrFloatError::OutOfMemory) => ShardResponse::OutOfMemory, Err(e) => ShardResponse::Err(e.to_string()), }, - ShardRequest::Append { key, value } => match ks.append(key, value) { - Ok(len) => ShardResponse::Len(len), - Err(WriteError::WrongType) => ShardResponse::WrongType, - Err(WriteError::OutOfMemory) => ShardResponse::OutOfMemory, - }, + ShardRequest::Append { key, value } => write_result_len(ks.append(key, value)), ShardRequest::Strlen { key } => match ks.strlen(key) { Ok(len) => ShardResponse::Len(len), Err(_) => ShardResponse::WrongType, @@ -750,16 +745,8 @@ fn dispatch( ShardRequest::Pexpire { key, milliseconds } => { ShardResponse::Bool(ks.pexpire(key, *milliseconds)) } - ShardRequest::LPush { key, values } => match ks.lpush(key, values) { - Ok(len) => ShardResponse::Len(len), - Err(WriteError::WrongType) => ShardResponse::WrongType, - Err(WriteError::OutOfMemory) => ShardResponse::OutOfMemory, - }, - ShardRequest::RPush { key, values } => match ks.rpush(key, values) { - Ok(len) => ShardResponse::Len(len), - Err(WriteError::WrongType) => ShardResponse::WrongType, - Err(WriteError::OutOfMemory) => ShardResponse::OutOfMemory, - }, + ShardRequest::LPush { key, values } => write_result_len(ks.lpush(key, values)), + ShardRequest::RPush { key, values } => write_result_len(ks.rpush(key, values)), ShardRequest::LPop { key } => match ks.lpop(key) { Ok(val) => ShardResponse::Value(val.map(Value::String)), Err(_) => ShardResponse::WrongType, @@ -844,11 +831,7 @@ fn dispatch( 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::HSet { key, fields } => write_result_len(ks.hset(key, fields)), ShardRequest::HGet { key, field } => match ks.hget(key, field) { Ok(val) => ShardResponse::Value(val.map(Value::String)), Err(_) => ShardResponse::WrongType, @@ -872,12 +855,7 @@ fn dispatch( 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::HIncrBy { key, field, delta } => incr_result(ks.hincrby(key, field, *delta)), ShardRequest::HKeys { key } => match ks.hkeys(key) { Ok(keys) => ShardResponse::StringArray(keys), Err(_) => ShardResponse::WrongType, @@ -890,11 +868,7 @@ fn dispatch( 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::SAdd { key, members } => write_result_len(ks.sadd(key, members)), ShardRequest::SRem { key, members } => match ks.srem(key, members) { Ok(count) => ShardResponse::Len(count), Err(_) => ShardResponse::WrongType, From ad0eaec64421d97d95a21e08401582d629a7709d Mon Sep 17 00:00:00 2001 From: Kacy Fortner Date: Wed, 11 Feb 2026 22:44:17 -0500 Subject: [PATCH 6/6] refactor: extract startup helpers from main resolve_password(), parse_bind_addr(), and build_persistence_config() pulled out of main() so the startup flow is scannable at a glance. three bind-address parse blocks replaced with a shared helper. --- crates/ember-server/src/main.rs | 128 +++++++++++++++++--------------- 1 file changed, 70 insertions(+), 58 deletions(-) diff --git a/crates/ember-server/src/main.rs b/crates/ember-server/src/main.rs index 34a31030..01507f0f 100644 --- a/crates/ember-server/src/main.rs +++ b/crates/ember-server/src/main.rs @@ -151,18 +151,9 @@ struct Args { cluster_node_timeout: u64, } -#[tokio::main] -async fn main() { - tracing_subscriber::fmt() - .with_env_filter( - tracing_subscriber::EnvFilter::try_from_default_env() - .unwrap_or_else(|_| "ember=info".into()), - ) - .init(); - - let mut args = Args::parse(); - - // resolve password: --requirepass-file takes the same role as --requirepass +/// Resolves the password from either `--requirepass` or `--requirepass-file`. +/// The two options are mutually exclusive. Exits on error. +fn resolve_password(args: &mut Args) { if args.requirepass.is_some() && args.requirepass_file.is_some() { eprintln!("error: --requirepass and --requirepass-file are mutually exclusive"); std::process::exit(1); @@ -186,14 +177,71 @@ async fn main() { } } } +} - let addr: SocketAddr = match format!("{}:{}", args.host, args.port).parse() { +/// Parses a `host:port` pair into a `SocketAddr`. Exits with a message on failure. +fn parse_bind_addr(host: &str, port: u16, label: &str) -> SocketAddr { + match format!("{host}:{port}").parse() { Ok(a) => a, Err(e) => { - eprintln!("invalid bind address '{}:{}': {e}", args.host, args.port); + if label.is_empty() { + eprintln!("invalid bind address '{host}:{port}': {e}"); + } else { + eprintln!("invalid {label} bind address '{host}:{port}': {e}"); + } std::process::exit(1); } - }; + } +} + +/// Builds the persistence config from CLI args. Returns `None` if persistence +/// is not enabled. Exits on validation errors. +fn build_persistence_config( + args: &mut Args, + #[cfg(feature = "encryption")] encryption_key: Option< + ember_persistence::encryption::EncryptionKey, + >, +) -> Option { + if !args.appendonly && args.data_dir.is_none() { + return None; + } + + let data_dir = args.data_dir.take().unwrap_or_else(|| { + if args.appendonly { + eprintln!("--data-dir is required when --appendonly is set"); + std::process::exit(1); + } + PathBuf::from(".") + }); + + let fsync_policy = parse_fsync_policy(&args.appendfsync).unwrap_or_else(|e| { + eprintln!("invalid --appendfsync value: {e}"); + std::process::exit(1); + }); + + Some(ShardPersistenceConfig { + data_dir, + append_only: args.appendonly, + fsync_policy, + #[cfg(feature = "encryption")] + encryption_key, + }) +} + +#[tokio::main] +async fn main() { + tracing_subscriber::fmt() + .with_env_filter( + tracing_subscriber::EnvFilter::try_from_default_env() + .unwrap_or_else(|_| "ember=info".into()), + ) + .init(); + + let mut args = Args::parse(); + + resolve_password(&mut args); + + let addr = parse_bind_addr(&args.host, args.port, ""); let max_memory = args.max_memory.as_deref().map(|s| { parse_byte_size(s).unwrap_or_else(|e| { @@ -238,31 +286,11 @@ async fn main() { std::process::exit(1); } - // build persistence config if data-dir is set or appendonly is enabled - let persistence = if args.appendonly || args.data_dir.is_some() { - let data_dir = args.data_dir.unwrap_or_else(|| { - if args.appendonly { - eprintln!("--data-dir is required when --appendonly is set"); - std::process::exit(1); - } - PathBuf::from(".") - }); - - let fsync_policy = parse_fsync_policy(&args.appendfsync).unwrap_or_else(|e| { - eprintln!("invalid --appendfsync value: {e}"); - std::process::exit(1); - }); - - Some(ShardPersistenceConfig { - data_dir, - append_only: args.appendonly, - fsync_policy, - #[cfg(feature = "encryption")] - encryption_key, - }) - } else { - None - }; + let persistence = build_persistence_config( + &mut args, + #[cfg(feature = "encryption")] + encryption_key, + ); #[allow(unused_mut)] let mut engine_config = @@ -305,17 +333,7 @@ async fn main() { // install prometheus metrics exporter if --metrics-port is set if let Some(metrics_port) = args.metrics_port { - let metrics_addr: std::net::SocketAddr = - match format!("{}:{}", args.host, metrics_port).parse() { - Ok(a) => a, - Err(e) => { - eprintln!( - "invalid metrics bind address '{}:{metrics_port}': {e}", - args.host - ); - std::process::exit(1); - } - }; + let metrics_addr = parse_bind_addr(&args.host, metrics_port, "metrics"); if let Err(e) = metrics::install_exporter(metrics_addr) { eprintln!("failed to start metrics exporter: {e}"); std::process::exit(1); @@ -358,13 +376,7 @@ async fn main() { } }; - let tls_addr: SocketAddr = match format!("{}:{}", args.host, tls_port).parse() { - Ok(a) => a, - Err(e) => { - eprintln!("invalid TLS bind address '{}:{tls_port}': {e}", args.host); - std::process::exit(1); - } - }; + let tls_addr = parse_bind_addr(&args.host, tls_port, "TLS"); info!( tls_port = tls_port,