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/schema.rs b/crates/ember-core/src/schema.rs index b4f34c6b..2567592f 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,13 +42,21 @@ 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. 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, @@ -55,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. @@ -65,6 +84,7 @@ impl SchemaRegistry { pub fn new() -> Self { Self { schemas: HashMap::new(), + message_cache: HashMap::new(), } } @@ -86,6 +106,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()))?; @@ -100,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 { @@ -165,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 { @@ -249,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())) } } @@ -281,7 +315,24 @@ fn resolve_field_path( ))); } } + if segments.len() > MAX_FIELD_PATH_DEPTH { + return Err(SchemaError::PathTooDeep( + segments.len(), + MAX_FIELD_PATH_DEPTH, + )); + } + + // 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() { @@ -392,6 +443,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 { @@ -534,6 +591,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() } } @@ -1077,4 +1135,233 @@ 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:?}"), + } + } + + // --- 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)); + } } 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, }; 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(""))); +}