From 620fbc11d0c4d671446d91f56c11841b78079e0f Mon Sep 17 00:00:00 2001 From: Kacy Fortner Date: Tue, 10 Feb 2026 20:49:54 -0500 Subject: [PATCH 1/6] fix: add descriptor size and field path depth limits adds two security checks to prevent DoS through the proto API: - MAX_DESCRIPTOR_BYTES (10MB) rejects oversized descriptors before decoding in PROTO.REGISTER - MAX_FIELD_PATH_DEPTH (16) rejects deeply nested field paths in get_field, set_field, and clear_field includes unit tests for both limits plus edge cases (double-dot and trailing-dot paths). --- crates/ember-core/src/schema.rs | 107 ++++++++++++++++++++++++++++++++ 1 file changed, 107 insertions(+) diff --git a/crates/ember-core/src/schema.rs b/crates/ember-core/src/schema.rs index b4f34c6b..99f71779 100644 --- a/crates/ember-core/src/schema.rs +++ b/crates/ember-core/src/schema.rs @@ -17,6 +17,14 @@ use prost_reflect::{ }; use thiserror::Error; +/// Maximum allowed size for a `FileDescriptorSet` in bytes (10 MB). +/// Prevents a single PROTO.REGISTER from consuming unbounded memory. +const MAX_DESCRIPTOR_BYTES: usize = 10 * 1024 * 1024; + +/// Maximum number of segments in a dot-separated field path. +/// Deep nesting beyond this is almost certainly a bug or an abuse vector. +const MAX_FIELD_PATH_DEPTH: usize = 16; + /// Errors that can occur during schema operations. #[derive(Debug, Error)] pub enum SchemaError { @@ -34,6 +42,12 @@ pub enum SchemaError { #[error("field not found: {0}")] FieldNotFound(String), + + #[error("descriptor too large: {0} bytes (max {1})")] + DescriptorTooLarge(usize, usize), + + #[error("field path too deep: {0} segments (max {1})")] + PathTooDeep(usize, usize), } /// A registered schema: the raw descriptor bytes and the parsed pool. @@ -86,6 +100,13 @@ impl SchemaRegistry { return Err(SchemaError::AlreadyExists(name)); } + if descriptor_bytes.len() > MAX_DESCRIPTOR_BYTES { + return Err(SchemaError::DescriptorTooLarge( + descriptor_bytes.len(), + MAX_DESCRIPTOR_BYTES, + )); + } + let pool = DescriptorPool::decode(descriptor_bytes.as_ref()) .map_err(|e| SchemaError::InvalidDescriptor(e.to_string()))?; @@ -281,6 +302,12 @@ fn resolve_field_path( ))); } } + if segments.len() > MAX_FIELD_PATH_DEPTH { + return Err(SchemaError::PathTooDeep( + segments.len(), + MAX_FIELD_PATH_DEPTH, + )); + } let mut current_msg = msg.clone(); @@ -392,6 +419,12 @@ fn resolve_field_path_mut<'a>( ))); } } + if segments.len() > MAX_FIELD_PATH_DEPTH { + return Err(SchemaError::PathTooDeep( + segments.len(), + MAX_FIELD_PATH_DEPTH, + )); + } // for a single segment, just verify the field exists and return if segments.len() == 1 { @@ -1077,4 +1110,78 @@ mod tests { .unwrap(); assert_eq!(frame, Frame::Integer(25)); } + + // --- size / depth limit tests --- + + #[test] + fn descriptor_size_limit_exceeded() { + let mut registry = SchemaRegistry::new(); + // craft a payload larger than MAX_DESCRIPTOR_BYTES + let oversized = Bytes::from(vec![0u8; MAX_DESCRIPTOR_BYTES + 1]); + let err = registry.register("huge".into(), oversized).unwrap_err(); + assert!(matches!(err, SchemaError::DescriptorTooLarge(_, _))); + } + + #[test] + fn field_path_depth_limit_exceeded() { + let desc = make_nested_descriptor(); + let mut registry = SchemaRegistry::new(); + registry.register("nested".into(), desc).unwrap(); + + let pool = ®istry.schemas["nested"].pool; + let outer_desc = pool.get_message_by_name("test.Outer").unwrap(); + let msg = DynamicMessage::new(outer_desc); + let mut buf = Vec::new(); + use prost_reflect::prost::Message; + msg.encode(&mut buf).unwrap(); + + // 17 segments exceeds the limit of 16 + let deep_path = (0..17) + .map(|i| format!("f{i}")) + .collect::>() + .join("."); + + let err = registry + .get_field("test.Outer", &buf, &deep_path) + .unwrap_err(); + assert!(matches!(err, SchemaError::PathTooDeep(17, 16))); + + let err = registry + .set_field("test.Outer", &buf, &deep_path, "val") + .unwrap_err(); + assert!(matches!(err, SchemaError::PathTooDeep(17, 16))); + + let err = registry + .clear_field("test.Outer", &buf, &deep_path) + .unwrap_err(); + assert!(matches!(err, SchemaError::PathTooDeep(17, 16))); + } + + #[test] + fn double_dot_path_returns_error() { + let mut registry = SchemaRegistry::new(); + let desc = make_descriptor("test", "User", "name"); + registry.register("users".into(), desc).unwrap(); + + let data = encode_user(®istry, "alice"); + let err = registry.get_field("test.User", &data, "a..b").unwrap_err(); + match err { + SchemaError::FieldNotFound(msg) => assert!(msg.contains("empty segment"), "{msg}"), + other => panic!("expected FieldNotFound, got {other:?}"), + } + } + + #[test] + fn trailing_dot_path_returns_error() { + let mut registry = SchemaRegistry::new(); + let desc = make_descriptor("test", "User", "name"); + registry.register("users".into(), desc).unwrap(); + + let data = encode_user(®istry, "alice"); + let err = registry.get_field("test.User", &data, "name.").unwrap_err(); + match err { + SchemaError::FieldNotFound(msg) => assert!(msg.contains("empty segment"), "{msg}"), + other => panic!("expected FieldNotFound, got {other:?}"), + } + } } From c8119e3a8fdcd2829bd10a8eabe0814a1926680c Mon Sep 17 00:00:00 2001 From: Kacy Fortner Date: Tue, 10 Feb 2026 20:52:04 -0500 Subject: [PATCH 2/6] perf: add message descriptor cache for O(1) find_message replaces O(n) linear scan across all registered schemas with a HashMap that's populated during register() and restore(). common path (every get_field/set_field/validate call) now does a single hash lookup instead of iterating all pools. --- crates/ember-core/src/schema.rs | 30 ++++++++++++++++++++++-------- 1 file changed, 22 insertions(+), 8 deletions(-) diff --git a/crates/ember-core/src/schema.rs b/crates/ember-core/src/schema.rs index 99f71779..a7e84051 100644 --- a/crates/ember-core/src/schema.rs +++ b/crates/ember-core/src/schema.rs @@ -54,7 +54,9 @@ pub enum SchemaError { struct RegisteredSchema { /// Raw `FileDescriptorSet` bytes, kept for persistence. descriptor_bytes: Bytes, - /// Parsed descriptor pool for message lookup and validation. + /// Parsed descriptor pool. Used during registration to build the + /// message cache, and in tests to construct dynamic messages. + #[cfg_attr(not(test), allow(dead_code))] pool: DescriptorPool, /// All message type full names in this schema. message_types: Vec, @@ -69,6 +71,9 @@ struct RegisteredSchema { /// derive it. pub struct SchemaRegistry { schemas: HashMap, + /// Flattened index of message type full name -> descriptor, built from + /// all registered schemas. Turns `find_message` from O(schemas) to O(1). + message_cache: HashMap, } /// Thread-safe handle to a shared schema registry. @@ -79,6 +84,7 @@ impl SchemaRegistry { pub fn new() -> Self { Self { schemas: HashMap::new(), + message_cache: HashMap::new(), } } @@ -121,6 +127,10 @@ impl SchemaRegistry { )); } + for desc in pool.all_messages() { + self.message_cache.insert(desc.full_name().to_owned(), desc); + } + self.schemas.insert( name, RegisteredSchema { @@ -186,6 +196,10 @@ impl SchemaRegistry { .map(|m| m.full_name().to_owned()) .collect(); + for desc in pool.all_messages() { + self.message_cache.insert(desc.full_name().to_owned(), desc); + } + self.schemas.insert( name, RegisteredSchema { @@ -270,14 +284,13 @@ impl SchemaRegistry { Ok(Bytes::from(buf)) } - /// Looks up a message descriptor by full name across all schemas. + /// Looks up a message descriptor by full name. O(1) via the + /// flattened `message_cache` built during registration. fn find_message(&self, message_type: &str) -> Result { - for schema in self.schemas.values() { - if let Some(desc) = schema.pool.get_message_by_name(message_type) { - return Ok(desc); - } - } - Err(SchemaError::UnknownMessageType(message_type.to_owned())) + self.message_cache + .get(message_type) + .cloned() + .ok_or_else(|| SchemaError::UnknownMessageType(message_type.to_owned())) } } @@ -567,6 +580,7 @@ impl std::fmt::Debug for SchemaRegistry { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("SchemaRegistry") .field("schema_count", &self.schemas.len()) + .field("cached_messages", &self.message_cache.len()) .finish() } } From 6b956199e47136da1749fde944aa2a905e24e352 Mon Sep 17 00:00:00 2001 From: Kacy Fortner Date: Tue, 10 Feb 2026 20:52:43 -0500 Subject: [PATCH 3/6] perf: optimize resolve_field_path for single-segment reads adds an early return for the common case of reading a top-level field (e.g. "name") that borrows directly from the decoded message instead of cloning the entire DynamicMessage. mirrors the pattern already used in resolve_field_path_mut. --- crates/ember-core/src/schema.rs | 11 +++++++++++ 1 file changed, 11 insertions(+) diff --git a/crates/ember-core/src/schema.rs b/crates/ember-core/src/schema.rs index a7e84051..6fb2801a 100644 --- a/crates/ember-core/src/schema.rs +++ b/crates/ember-core/src/schema.rs @@ -322,6 +322,17 @@ fn resolve_field_path( )); } + // fast path: single-segment reads can borrow directly without cloning + if segments.len() == 1 { + let field_desc = msg + .descriptor() + .get_field_by_name(segments[0]) + .ok_or_else(|| SchemaError::FieldNotFound(segments[0].to_string()))?; + let value = msg.get_field(&field_desc).into_owned(); + return Ok((value, field_desc)); + } + + // multi-segment: clone is unavoidable due to prost-reflect API let mut current_msg = msg.clone(); for (i, segment) in segments.iter().enumerate() { From fc5dc296d63fc19b88ae0f42e0811998093f644c Mon Sep 17 00:00:00 2001 From: Kacy Fortner Date: Tue, 10 Feb 2026 20:53:48 -0500 Subject: [PATCH 4/6] test: add nested path and u64 overflow unit tests covers edge cases not previously tested: - set_field on nested paths (inner.value) - clear_field on nested paths - auto-initialization of intermediate messages on nested set - u64::MAX returns bulk string (too large for i64) - u64 that fits in i64 returns integer frame --- crates/ember-core/src/schema.rs | 155 ++++++++++++++++++++++++++++++++ 1 file changed, 155 insertions(+) diff --git a/crates/ember-core/src/schema.rs b/crates/ember-core/src/schema.rs index 6fb2801a..2567592f 100644 --- a/crates/ember-core/src/schema.rs +++ b/crates/ember-core/src/schema.rs @@ -1209,4 +1209,159 @@ mod tests { other => panic!("expected FieldNotFound, got {other:?}"), } } + + // --- nested set/clear and u64 edge case tests --- + + #[test] + fn set_field_nested_path() { + let desc = make_nested_descriptor(); + let mut registry = SchemaRegistry::new(); + registry.register("nested".into(), desc).unwrap(); + + let pool = ®istry.schemas["nested"].pool; + let outer_desc = pool.get_message_by_name("test.Outer").unwrap(); + let inner_desc = pool.get_message_by_name("test.Inner").unwrap(); + + let mut inner = DynamicMessage::new(inner_desc); + inner.set_field_by_name("value", prost_reflect::Value::String("hello".into())); + let mut outer = DynamicMessage::new(outer_desc); + outer.set_field_by_name("inner", prost_reflect::Value::Message(inner)); + + let mut buf = Vec::new(); + use prost_reflect::prost::Message; + outer.encode(&mut buf).unwrap(); + + let new_data = registry + .set_field("test.Outer", &buf, "inner.value", "world") + .unwrap(); + + let frame = registry + .get_field("test.Outer", &new_data, "inner.value") + .unwrap(); + assert_eq!(frame, Frame::Bulk(Bytes::from("world"))); + } + + #[test] + fn clear_field_nested_path() { + let desc = make_nested_descriptor(); + let mut registry = SchemaRegistry::new(); + registry.register("nested".into(), desc).unwrap(); + + let pool = ®istry.schemas["nested"].pool; + let outer_desc = pool.get_message_by_name("test.Outer").unwrap(); + let inner_desc = pool.get_message_by_name("test.Inner").unwrap(); + + let mut inner = DynamicMessage::new(inner_desc); + inner.set_field_by_name("value", prost_reflect::Value::String("hello".into())); + let mut outer = DynamicMessage::new(outer_desc); + outer.set_field_by_name("inner", prost_reflect::Value::Message(inner)); + + let mut buf = Vec::new(); + use prost_reflect::prost::Message; + outer.encode(&mut buf).unwrap(); + + let new_data = registry + .clear_field("test.Outer", &buf, "inner.value") + .unwrap(); + + // cleared string field returns empty default + let frame = registry + .get_field("test.Outer", &new_data, "inner.value") + .unwrap(); + assert_eq!(frame, Frame::Bulk(Bytes::from(""))); + } + + #[test] + fn set_field_nested_creates_intermediate() { + let desc = make_nested_descriptor(); + let mut registry = SchemaRegistry::new(); + registry.register("nested".into(), desc).unwrap(); + + // create an Outer with no inner field set + let pool = ®istry.schemas["nested"].pool; + let outer_desc = pool.get_message_by_name("test.Outer").unwrap(); + let outer = DynamicMessage::new(outer_desc); + + let mut buf = Vec::new(); + use prost_reflect::prost::Message; + outer.encode(&mut buf).unwrap(); + + // set_field should auto-init the intermediate Inner message + let new_data = registry + .set_field("test.Outer", &buf, "inner.value", "auto") + .unwrap(); + + let frame = registry + .get_field("test.Outer", &new_data, "inner.value") + .unwrap(); + assert_eq!(frame, Frame::Bulk(Bytes::from("auto"))); + } + + /// Builds a descriptor with a single uint64 field. + fn make_uint64_descriptor() -> Bytes { + use prost_reflect::prost_types::{ + DescriptorProto, FieldDescriptorProto, FileDescriptorProto, FileDescriptorSet, + }; + + let fds = FileDescriptorSet { + file: vec![FileDescriptorProto { + name: Some("test.proto".into()), + package: Some("test".into()), + message_type: vec![DescriptorProto { + name: Some("BigNum".into()), + field: vec![FieldDescriptorProto { + name: Some("val".into()), + number: Some(1), + r#type: Some(4), // TYPE_UINT64 + label: Some(1), + ..Default::default() + }], + ..Default::default() + }], + ..Default::default() + }], + }; + let mut buf = Vec::new(); + use prost_reflect::prost::Message; + fds.encode(&mut buf).unwrap(); + Bytes::from(buf) + } + + #[test] + fn u64_overflow_returns_bulk_string() { + let desc = make_uint64_descriptor(); + let mut registry = SchemaRegistry::new(); + registry.register("bignums".into(), desc).unwrap(); + + let pool = ®istry.schemas["bignums"].pool; + let msg_desc = pool.get_message_by_name("test.BigNum").unwrap(); + let mut msg = DynamicMessage::new(msg_desc); + msg.set_field_by_name("val", prost_reflect::Value::U64(u64::MAX)); + + let mut buf = Vec::new(); + use prost_reflect::prost::Message; + msg.encode(&mut buf).unwrap(); + + let frame = registry.get_field("test.BigNum", &buf, "val").unwrap(); + assert_eq!(frame, Frame::Bulk(Bytes::from("18446744073709551615"))); + } + + #[test] + fn u64_fits_in_i64_returns_integer() { + let desc = make_uint64_descriptor(); + let mut registry = SchemaRegistry::new(); + registry.register("bignums".into(), desc).unwrap(); + + let pool = ®istry.schemas["bignums"].pool; + let msg_desc = pool.get_message_by_name("test.BigNum").unwrap(); + let mut msg = DynamicMessage::new(msg_desc); + msg.set_field_by_name("val", prost_reflect::Value::U64(42)); + + let mut buf = Vec::new(); + use prost_reflect::prost::Message; + msg.encode(&mut buf).unwrap(); + + let frame = registry.get_field("test.BigNum", &buf, "val").unwrap(); + assert_eq!(frame, Frame::Integer(42)); + } } From 548647942afe216357bc4615a75e2a8b206f42cf Mon Sep 17 00:00:00 2001 From: Kacy Fortner Date: Tue, 10 Feb 2026 20:56:42 -0500 Subject: [PATCH 5/6] fix: preserve TTL across SETFIELD and DELFIELD operations proto_get now returns remaining TTL alongside type_name and data. the SETFIELD/DELFIELD handlers in both sharded and concurrent modes pass the existing TTL through to the write-back ProtoSet, so a key's expiry is no longer silently reset to no-expiry on field mutation. --- crates/ember-core/src/keyspace.rs | 20 ++++++++++++++++--- crates/ember-core/src/shard.rs | 4 ++-- crates/ember-server/src/concurrent_handler.rs | 16 +++++++-------- crates/ember-server/src/connection.rs | 20 +++++++++---------- 4 files changed, 37 insertions(+), 23 deletions(-) diff --git a/crates/ember-core/src/keyspace.rs b/crates/ember-core/src/keyspace.rs index 6528cd34..bc5797db 100644 --- a/crates/ember-core/src/keyspace.rs +++ b/crates/ember-core/src/keyspace.rs @@ -1888,18 +1888,32 @@ impl Keyspace { SetResult::Ok } - /// Retrieves a proto value, returning `(type_name, data)` or `None`. + /// Retrieves a proto value, returning `(type_name, data, remaining_ttl)` + /// or `None`. + /// + /// The remaining TTL is `Some(duration)` if the key has an expiry set, + /// or `None` for keys that never expire. This allows callers to preserve + /// the TTL across read-modify-write cycles (e.g. SETFIELD/DELFIELD). /// /// Returns `Err(WrongType)` if the key holds a different value type. #[cfg(feature = "protobuf")] - pub fn proto_get(&mut self, key: &str) -> Result, WrongType> { + pub fn proto_get( + &mut self, + key: &str, + ) -> Result)>, WrongType> { if self.remove_if_expired(key) { return Ok(None); } match self.entries.get_mut(key) { Some(e) => { if let Value::Proto { type_name, data } = &e.value { - let result = (type_name.clone(), data.clone()); + let remaining = if e.expires_at_ms == 0 { + None + } else { + let now = time::now_ms(); + Some(Duration::from_millis(e.expires_at_ms.saturating_sub(now))) + }; + let result = (type_name.clone(), data.clone(), remaining); e.touch(); Ok(Some(result)) } else { diff --git a/crates/ember-core/src/shard.rs b/crates/ember-core/src/shard.rs index 5380745b..8354a193 100644 --- a/crates/ember-core/src/shard.rs +++ b/crates/ember-core/src/shard.rs @@ -344,9 +344,9 @@ pub enum ShardResponse { StringArray(Vec), /// HMGET result: array of optional values. OptionalArray(Vec>), - /// PROTO.GET result: (type_name, data) or None. + /// PROTO.GET result: (type_name, data, remaining_ttl) or None. #[cfg(feature = "protobuf")] - ProtoValue(Option<(String, Bytes)>), + ProtoValue(Option<(String, Bytes, Option)>), /// PROTO.TYPE result: message type name or None. #[cfg(feature = "protobuf")] ProtoTypeName(Option), diff --git a/crates/ember-server/src/concurrent_handler.rs b/crates/ember-server/src/concurrent_handler.rs index aef45c9d..2267e720 100644 --- a/crates/ember-server/src/concurrent_handler.rs +++ b/crates/ember-server/src/concurrent_handler.rs @@ -377,7 +377,7 @@ async fn execute_concurrent( } let req = ember_core::ShardRequest::ProtoGet { key: key.clone() }; match _engine.route(&key, req).await { - Ok(ember_core::ShardResponse::ProtoValue(Some((type_name, data)))) => { + Ok(ember_core::ShardResponse::ProtoValue(Some((type_name, data, _ttl)))) => { Frame::Array(vec![Frame::Bulk(Bytes::from(type_name)), Frame::Bulk(data)]) } Ok(ember_core::ShardResponse::ProtoValue(None)) => Frame::Null, @@ -456,7 +456,7 @@ async fn execute_concurrent( }; let req = ember_core::ShardRequest::ProtoGet { key: key.clone() }; match _engine.route(&key, req).await { - Ok(ember_core::ShardResponse::ProtoValue(Some((type_name, data)))) => { + Ok(ember_core::ShardResponse::ProtoValue(Some((type_name, data, _ttl)))) => { let reg = match registry.read() { Ok(r) => r, Err(_) => return Frame::Error("ERR schema registry lock poisoned".into()), @@ -486,8 +486,8 @@ async fn execute_concurrent( None => return Frame::Error("ERR protobuf support is not enabled".into()), }; let req = ember_core::ShardRequest::ProtoGet { key: key.clone() }; - let (type_name, data) = match _engine.route(&key, req).await { - Ok(ember_core::ShardResponse::ProtoValue(Some(pair))) => pair, + let (type_name, data, existing_ttl) = match _engine.route(&key, req).await { + Ok(ember_core::ShardResponse::ProtoValue(Some(tuple))) => tuple, Ok(ember_core::ShardResponse::ProtoValue(None)) => return Frame::Null, Ok(ember_core::ShardResponse::WrongType) => { return Frame::Error( @@ -513,7 +513,7 @@ async fn execute_concurrent( key: key.clone(), type_name, data: new_data, - expire: None, + expire: existing_ttl, nx: false, xx: true, }; @@ -535,8 +535,8 @@ async fn execute_concurrent( None => return Frame::Error("ERR protobuf support is not enabled".into()), }; let req = ember_core::ShardRequest::ProtoGet { key: key.clone() }; - let (type_name, data) = match _engine.route(&key, req).await { - Ok(ember_core::ShardResponse::ProtoValue(Some(pair))) => pair, + let (type_name, data, existing_ttl) = match _engine.route(&key, req).await { + Ok(ember_core::ShardResponse::ProtoValue(Some(tuple))) => tuple, Ok(ember_core::ShardResponse::ProtoValue(None)) => return Frame::Null, Ok(ember_core::ShardResponse::WrongType) => { return Frame::Error( @@ -562,7 +562,7 @@ async fn execute_concurrent( key: key.clone(), type_name, data: new_data, - expire: None, + expire: existing_ttl, nx: false, xx: true, }; diff --git a/crates/ember-server/src/connection.rs b/crates/ember-server/src/connection.rs index 50a8b717..20ae9849 100644 --- a/crates/ember-server/src/connection.rs +++ b/crates/ember-server/src/connection.rs @@ -1712,7 +1712,7 @@ async fn execute( } let req = ShardRequest::ProtoGet { key: key.clone() }; match engine.route(&key, req).await { - Ok(ShardResponse::ProtoValue(Some((type_name, data)))) => { + Ok(ShardResponse::ProtoValue(Some((type_name, data, _ttl)))) => { Frame::Array(vec![Frame::Bulk(Bytes::from(type_name)), Frame::Bulk(data)]) } Ok(ShardResponse::ProtoValue(None)) => Frame::Null, @@ -1785,7 +1785,7 @@ async fn execute( }; let req = ShardRequest::ProtoGet { key: key.clone() }; match engine.route(&key, req).await { - Ok(ShardResponse::ProtoValue(Some((type_name, data)))) => { + Ok(ShardResponse::ProtoValue(Some((type_name, data, _ttl)))) => { let reg = match registry.read() { Ok(r) => r, Err(_) => return Frame::Error("ERR schema registry lock poisoned".into()), @@ -1814,8 +1814,8 @@ async fn execute( }; // step 1: fetch current value let req = ShardRequest::ProtoGet { key: key.clone() }; - let (type_name, data) = match engine.route(&key, req).await { - Ok(ShardResponse::ProtoValue(Some(pair))) => pair, + let (type_name, data, existing_ttl) = match engine.route(&key, req).await { + Ok(ShardResponse::ProtoValue(Some(tuple))) => tuple, Ok(ShardResponse::ProtoValue(None)) => return Frame::Null, Ok(ShardResponse::WrongType) => return wrongtype_error(), Ok(other) => { @@ -1834,12 +1834,12 @@ async fn execute( Err(e) => return Frame::Error(format!("ERR {e}")), } }; - // step 3: store back (XX = only if key still exists) + // step 3: store back (XX = only if key still exists), preserving TTL let req = ShardRequest::ProtoSet { key: key.clone(), type_name, data: new_data, - expire: None, + expire: existing_ttl, nx: false, xx: true, }; @@ -1860,8 +1860,8 @@ async fn execute( }; // step 1: fetch current value let req = ShardRequest::ProtoGet { key: key.clone() }; - let (type_name, data) = match engine.route(&key, req).await { - Ok(ShardResponse::ProtoValue(Some(pair))) => pair, + let (type_name, data, existing_ttl) = match engine.route(&key, req).await { + Ok(ShardResponse::ProtoValue(Some(tuple))) => tuple, Ok(ShardResponse::ProtoValue(None)) => return Frame::Null, Ok(ShardResponse::WrongType) => return wrongtype_error(), Ok(other) => { @@ -1880,12 +1880,12 @@ async fn execute( Err(e) => return Frame::Error(format!("ERR {e}")), } }; - // step 3: store back (XX = only if key still exists) + // step 3: store back (XX = only if key still exists), preserving TTL let req = ShardRequest::ProtoSet { key: key.clone(), type_name, data: new_data, - expire: None, + expire: existing_ttl, nx: false, xx: true, }; From 557663d1ccd4838f7580530c3f298e1726be2131 Mon Sep 17 00:00:00 2001 From: Kacy Fortner Date: Tue, 10 Feb 2026 20:59:47 -0500 Subject: [PATCH 6/6] test: add integration tests for TTL preservation and nested paths adds sharded-mode tests verifying that SETFIELD and DELFIELD preserve the key's TTL instead of resetting it. also covers nested field paths (inner.value) for both set and del, and duplicate schema registration rejection. concurrent-mode tests cover nested paths but skip TTL verification since proto values route through engine shards while TTL checks the concurrent keyspace. --- tests/integration/src/proto.rs | 219 +++++++++++++++++++++++++++++++++ 1 file changed, 219 insertions(+) diff --git a/tests/integration/src/proto.rs b/tests/integration/src/proto.rs index f66653af..2701ff0d 100644 --- a/tests/integration/src/proto.rs +++ b/tests/integration/src/proto.rs @@ -794,3 +794,222 @@ async fn concurrent_delfield_clears_field() { let resp = c.cmd(&["PROTO.GETFIELD", "p:1", "name"]).await; assert_eq!(resp, Frame::Bulk(Bytes::from(""))); } + +// ---- TTL preservation tests ---- + +#[tokio::test] +async fn setfield_preserves_ttl() { + let server = start_proto_server(false); + let mut c = server.connect().await; + + let desc = make_multi_field_descriptor(); + c.cmd_raw(&[b"PROTO.REGISTER", b"profiles", &desc]).await; + + let data = encode_profile(&desc, "alice", 25, true); + c.cmd_raw(&[ + b"PROTO.SET", + b"p:ttl", + b"test.Profile", + &data, + b"EX", + b"120", + ]) + .await; + + // verify TTL is set + let ttl_before = c.get_int(&["TTL", "p:ttl"]).await; + assert!(ttl_before > 0 && ttl_before <= 120); + + // mutate a field + c.cmd(&["PROTO.SETFIELD", "p:ttl", "name", "bob"]).await; + + // TTL should still be set + let ttl_after = c.get_int(&["TTL", "p:ttl"]).await; + assert!(ttl_after > 0 && ttl_after <= 120); +} + +#[tokio::test] +async fn delfield_preserves_ttl() { + let server = start_proto_server(false); + let mut c = server.connect().await; + + let desc = make_multi_field_descriptor(); + c.cmd_raw(&[b"PROTO.REGISTER", b"profiles", &desc]).await; + + let data = encode_profile(&desc, "alice", 25, true); + c.cmd_raw(&[ + b"PROTO.SET", + b"p:ttl2", + b"test.Profile", + &data, + b"EX", + b"120", + ]) + .await; + + let ttl_before = c.get_int(&["TTL", "p:ttl2"]).await; + assert!(ttl_before > 0 && ttl_before <= 120); + + c.cmd(&["PROTO.DELFIELD", "p:ttl2", "name"]).await; + + let ttl_after = c.get_int(&["TTL", "p:ttl2"]).await; + assert!(ttl_after > 0 && ttl_after <= 120); +} + +// ---- nested path integration tests ---- + +/// Builds a FileDescriptorSet with Inner { string value = 1; } and +/// Outer { Inner inner = 1; }. +fn make_nested_descriptor() -> Vec { + let fds = FileDescriptorSet { + file: vec![FileDescriptorProto { + name: Some("test.proto".into()), + package: Some("test".into()), + message_type: vec![ + DescriptorProto { + name: Some("Inner".into()), + field: vec![FieldDescriptorProto { + name: Some("value".into()), + number: Some(1), + r#type: Some(9), // TYPE_STRING + label: Some(1), + ..Default::default() + }], + ..Default::default() + }, + DescriptorProto { + name: Some("Outer".into()), + field: vec![FieldDescriptorProto { + name: Some("inner".into()), + number: Some(1), + r#type: Some(11), // TYPE_MESSAGE + label: Some(1), + type_name: Some(".test.Inner".into()), + ..Default::default() + }], + ..Default::default() + }, + ], + ..Default::default() + }], + }; + let mut buf = Vec::new(); + fds.encode(&mut buf).expect("encode nested descriptor"); + buf +} + +/// Encodes an Outer message with the given inner value. +fn encode_outer(descriptor_bytes: &[u8], inner_value: &str) -> Vec { + let pool = DescriptorPool::decode(descriptor_bytes).expect("decode pool"); + let inner_desc = pool.get_message_by_name("test.Inner").expect("find Inner"); + let outer_desc = pool.get_message_by_name("test.Outer").expect("find Outer"); + + let mut inner = DynamicMessage::new(inner_desc); + inner.set_field_by_name("value", prost_reflect::Value::String(inner_value.into())); + + let mut outer = DynamicMessage::new(outer_desc); + outer.set_field_by_name("inner", prost_reflect::Value::Message(inner)); + + let mut buf = Vec::new(); + outer.encode(&mut buf).expect("encode outer"); + buf +} + +#[tokio::test] +async fn setfield_nested() { + let server = start_proto_server(false); + let mut c = server.connect().await; + + let desc = make_nested_descriptor(); + c.cmd_raw(&[b"PROTO.REGISTER", b"nested", &desc]).await; + + let data = encode_outer(&desc, "hello"); + c.cmd_raw(&[b"PROTO.SET", b"o:1", b"test.Outer", &data]) + .await; + + let resp = c + .cmd(&["PROTO.SETFIELD", "o:1", "inner.value", "world"]) + .await; + assert!(matches!(resp, Frame::Simple(ref s) if s == "OK")); + + let resp = c.cmd(&["PROTO.GETFIELD", "o:1", "inner.value"]).await; + assert_eq!(resp, Frame::Bulk(Bytes::from("world"))); +} + +#[tokio::test] +async fn delfield_nested() { + let server = start_proto_server(false); + let mut c = server.connect().await; + + let desc = make_nested_descriptor(); + c.cmd_raw(&[b"PROTO.REGISTER", b"nested", &desc]).await; + + let data = encode_outer(&desc, "hello"); + c.cmd_raw(&[b"PROTO.SET", b"o:1", b"test.Outer", &data]) + .await; + + let resp = c.cmd(&["PROTO.DELFIELD", "o:1", "inner.value"]).await; + assert_eq!(resp, Frame::Integer(1)); + + // cleared field should return default (empty string) + let resp = c.cmd(&["PROTO.GETFIELD", "o:1", "inner.value"]).await; + assert_eq!(resp, Frame::Bulk(Bytes::from(""))); +} + +#[tokio::test] +async fn duplicate_registration_returns_error() { + let server = start_proto_server(false); + let mut c = server.connect().await; + + let desc = make_descriptor("test", "User", "name"); + c.cmd_raw(&[b"PROTO.REGISTER", b"users", &desc]).await; + + // second registration with same name should fail + let resp = c.cmd_raw(&[b"PROTO.REGISTER", b"users", &desc]).await; + assert!(matches!(resp, Frame::Error(ref s) if s.contains("already registered"))); +} + +// ---- concurrent mode nested path and misc tests ---- +// note: TTL preservation tests are sharded-only because concurrent mode +// routes proto commands through engine shards while TTL checks the +// concurrent keyspace — these are separate stores. + +#[tokio::test] +async fn concurrent_setfield_nested() { + let server = start_proto_server(true); + let mut c = server.connect().await; + + let desc = make_nested_descriptor(); + c.cmd_raw(&[b"PROTO.REGISTER", b"nested", &desc]).await; + + let data = encode_outer(&desc, "hello"); + c.cmd_raw(&[b"PROTO.SET", b"o:1", b"test.Outer", &data]) + .await; + + let resp = c + .cmd(&["PROTO.SETFIELD", "o:1", "inner.value", "world"]) + .await; + assert!(matches!(resp, Frame::Simple(ref s) if s == "OK")); + + let resp = c.cmd(&["PROTO.GETFIELD", "o:1", "inner.value"]).await; + assert_eq!(resp, Frame::Bulk(Bytes::from("world"))); +} + +#[tokio::test] +async fn concurrent_delfield_nested() { + let server = start_proto_server(true); + let mut c = server.connect().await; + + let desc = make_nested_descriptor(); + c.cmd_raw(&[b"PROTO.REGISTER", b"nested", &desc]).await; + + let data = encode_outer(&desc, "hello"); + c.cmd_raw(&[b"PROTO.SET", b"o:1", b"test.Outer", &data]) + .await; + + let resp = c.cmd(&["PROTO.DELFIELD", "o:1", "inner.value"]).await; + assert_eq!(resp, Frame::Integer(1)); + + let resp = c.cmd(&["PROTO.GETFIELD", "o:1", "inner.value"]).await; + assert_eq!(resp, Frame::Bulk(Bytes::from(""))); +}