diff --git a/Cargo.lock b/Cargo.lock index 8da64a8a..0fc9f04a 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -653,6 +653,7 @@ version = "0.4.2" dependencies = [ "bytes", "ember-protocol", + "prost-reflect", "tempfile", "tokio", ] diff --git a/crates/ember-server/src/concurrent_handler.rs b/crates/ember-server/src/concurrent_handler.rs index 8ccc0321..060db6c6 100644 --- a/crates/ember-server/src/concurrent_handler.rs +++ b/crates/ember-server/src/concurrent_handler.rs @@ -16,6 +16,8 @@ use std::sync::atomic::Ordering; use std::sync::Arc; use std::time::{Duration, Instant}; +#[cfg(feature = "protobuf")] +use bytes::Bytes; use bytes::BytesMut; use ember_core::{ConcurrentKeyspace, Engine, TtlResult}; use ember_protocol::{parse_frame, Command, Frame, SetExpire}; @@ -290,6 +292,172 @@ async fn execute_concurrent( } }, + // -- protobuf commands -- + #[cfg(feature = "protobuf")] + Command::ProtoRegister { name, descriptor } => { + let registry = match _engine.schema_registry() { + Some(r) => r, + None => return Frame::Error("ERR protobuf support is not enabled".into()), + }; + let result = { + let mut reg = match registry.write() { + Ok(r) => r, + Err(_) => return Frame::Error("ERR schema registry lock poisoned".into()), + }; + reg.register(name.clone(), descriptor.clone()) + }; + match result { + Ok(types) => { + let _ = _engine + .broadcast(|| ember_core::ShardRequest::ProtoRegisterAof { + name: name.clone(), + descriptor: descriptor.clone(), + }) + .await; + Frame::Array( + types + .into_iter() + .map(|t| Frame::Bulk(Bytes::from(t))) + .collect(), + ) + } + Err(e) => Frame::Error(format!("ERR {e}")), + } + } + + #[cfg(feature = "protobuf")] + Command::ProtoSet { + key, + type_name, + data, + expire, + nx, + xx, + } => { + let registry = match _engine.schema_registry() { + Some(r) => r, + None => return Frame::Error("ERR protobuf support is not enabled".into()), + }; + { + let reg = match registry.read() { + Ok(r) => r, + Err(_) => return Frame::Error("ERR schema registry lock poisoned".into()), + }; + if let Err(e) = reg.validate(&type_name, &data) { + return Frame::Error(format!("ERR {e}")); + } + } + let duration = expire.map(|e| match e { + SetExpire::Ex(secs) => Duration::from_secs(secs), + SetExpire::Px(millis) => Duration::from_millis(millis), + }); + let req = ember_core::ShardRequest::ProtoSet { + key: key.clone(), + type_name, + data, + expire: duration, + nx, + xx, + }; + match _engine.route(&key, req).await { + Ok(ember_core::ShardResponse::Ok) => Frame::Simple("OK".into()), + Ok(ember_core::ShardResponse::Value(None)) => Frame::Null, + Ok(ember_core::ShardResponse::OutOfMemory) => { + Frame::Error("OOM command not allowed when used memory > 'maxmemory'".into()) + } + Ok(other) => Frame::Error(format!("ERR unexpected shard response: {other:?}")), + Err(e) => Frame::Error(format!("ERR {e}")), + } + } + + #[cfg(feature = "protobuf")] + Command::ProtoGet { key } => { + if _engine.schema_registry().is_none() { + return Frame::Error("ERR protobuf support is not enabled".into()); + } + let req = ember_core::ShardRequest::ProtoGet { key: key.clone() }; + match _engine.route(&key, req).await { + Ok(ember_core::ShardResponse::ProtoValue(Some((type_name, data)))) => { + Frame::Array(vec![Frame::Bulk(Bytes::from(type_name)), Frame::Bulk(data)]) + } + Ok(ember_core::ShardResponse::ProtoValue(None)) => Frame::Null, + Ok(ember_core::ShardResponse::WrongType) => Frame::Error( + "WRONGTYPE Operation against a key holding the wrong kind of value".into(), + ), + Ok(other) => Frame::Error(format!("ERR unexpected shard response: {other:?}")), + Err(e) => Frame::Error(format!("ERR {e}")), + } + } + + #[cfg(feature = "protobuf")] + Command::ProtoType { key } => { + if _engine.schema_registry().is_none() { + return Frame::Error("ERR protobuf support is not enabled".into()); + } + let req = ember_core::ShardRequest::ProtoType { key: key.clone() }; + match _engine.route(&key, req).await { + Ok(ember_core::ShardResponse::ProtoTypeName(Some(name))) => { + Frame::Bulk(Bytes::from(name)) + } + Ok(ember_core::ShardResponse::ProtoTypeName(None)) => Frame::Null, + Ok(ember_core::ShardResponse::WrongType) => Frame::Error( + "WRONGTYPE Operation against a key holding the wrong kind of value".into(), + ), + Ok(other) => Frame::Error(format!("ERR unexpected shard response: {other:?}")), + Err(e) => Frame::Error(format!("ERR {e}")), + } + } + + #[cfg(feature = "protobuf")] + Command::ProtoSchemas => { + let registry = match _engine.schema_registry() { + Some(r) => r, + None => return Frame::Error("ERR protobuf support is not enabled".into()), + }; + let reg = match registry.read() { + Ok(r) => r, + Err(_) => return Frame::Error("ERR schema registry lock poisoned".into()), + }; + let names = reg.schema_names(); + Frame::Array( + names + .into_iter() + .map(|n| Frame::Bulk(Bytes::from(n))) + .collect(), + ) + } + + #[cfg(feature = "protobuf")] + Command::ProtoDescribe { name } => { + let registry = match _engine.schema_registry() { + Some(r) => r, + None => return Frame::Error("ERR protobuf support is not enabled".into()), + }; + let reg = match registry.read() { + Ok(r) => r, + Err(_) => return Frame::Error("ERR schema registry lock poisoned".into()), + }; + match reg.describe(&name) { + Some(types) => Frame::Array( + types + .into_iter() + .map(|t| Frame::Bulk(Bytes::from(t))) + .collect(), + ), + None => Frame::Error(format!("ERR unknown schema '{name}'")), + } + } + + #[cfg(not(feature = "protobuf"))] + Command::ProtoRegister { .. } + | Command::ProtoSet { .. } + | Command::ProtoGet { .. } + | Command::ProtoType { .. } + | Command::ProtoSchemas + | Command::ProtoDescribe { .. } => { + Frame::Error("ERR unknown command (protobuf support not compiled)".into()) + } + Command::Quit => Frame::Simple("OK".into()), // For unsupported commands, return an error diff --git a/tests/integration/Cargo.toml b/tests/integration/Cargo.toml index 5e8b7e36..406c4001 100644 --- a/tests/integration/Cargo.toml +++ b/tests/integration/Cargo.toml @@ -14,4 +14,5 @@ harness = true tokio = { workspace = true, features = ["rt-multi-thread", "macros", "net", "io-util", "time"] } bytes = { workspace = true } ember-protocol = { workspace = true } +prost-reflect = { workspace = true } tempfile = "3" diff --git a/tests/integration/src/helpers.rs b/tests/integration/src/helpers.rs index 702200e6..ef8343c5 100644 --- a/tests/integration/src/helpers.rs +++ b/tests/integration/src/helpers.rs @@ -31,6 +31,10 @@ pub struct ServerOptions { pub cluster_enabled: bool, /// Bootstrap as a single-node cluster owning all 16384 slots. pub cluster_bootstrap: bool, + /// Enable protobuf value storage. + pub protobuf: bool, + /// Use concurrent (DashMap) mode instead of sharded channels. + pub concurrent: bool, } impl TestServer { @@ -58,6 +62,14 @@ impl TestServer { cmd.arg("--requirepass").arg(pass); } + if opts.protobuf { + cmd.arg("--protobuf"); + } + + if opts.concurrent { + cmd.arg("--concurrent"); + } + if opts.cluster_enabled { cmd.arg("--cluster-enabled"); // use a small offset so the gossip port stays in valid u16 range. @@ -144,6 +156,36 @@ impl TestClient { } } + /// Sends a command with raw byte arguments and returns the parsed response. + /// Useful for binary data like protobuf descriptors. + pub async fn cmd_raw(&mut self, args: &[&[u8]]) -> Frame { + let parts: Vec = args + .iter() + .map(|a| Frame::Bulk(Bytes::copy_from_slice(a))) + .collect(); + let frame = Frame::Array(parts); + + let mut out = BytesMut::new(); + frame.serialize(&mut out); + self.stream.write_all(&out).await.unwrap(); + + loop { + match parse_frame(&self.buf) { + Ok(Some((frame, consumed))) => { + let _ = self.buf.split_to(consumed); + return frame; + } + Ok(None) => { + let n = self.stream.read_buf(&mut self.buf).await.unwrap(); + if n == 0 { + panic!("server closed connection while waiting for response"); + } + } + Err(e) => panic!("protocol error: {e}"), + } + } + } + /// Sends a command and returns the parsed response frame. pub async fn cmd(&mut self, args: &[&str]) -> Frame { // build RESP3 array diff --git a/tests/integration/src/main.rs b/tests/integration/src/main.rs index 7ee5389e..ddc7c156 100644 --- a/tests/integration/src/main.rs +++ b/tests/integration/src/main.rs @@ -6,4 +6,5 @@ mod cli; mod cluster; mod data_types; mod persistence; +mod proto; mod pubsub; diff --git a/tests/integration/src/proto.rs b/tests/integration/src/proto.rs new file mode 100644 index 00000000..60b02bbd --- /dev/null +++ b/tests/integration/src/proto.rs @@ -0,0 +1,439 @@ +//! Integration tests for protobuf value storage (PROTO.* commands). +//! +//! Core tests run against both sharded and concurrent server modes to +//! ensure proto commands work through both code paths. + +use bytes::Bytes; +use ember_protocol::Frame; +use prost_reflect::prost::Message; +use prost_reflect::prost_types::{ + DescriptorProto, FieldDescriptorProto, FileDescriptorProto, FileDescriptorSet, +}; +use prost_reflect::{DescriptorPool, DynamicMessage}; + +use crate::helpers::{ServerOptions, TestServer}; + +/// Builds a minimal FileDescriptorSet with a single message type. +fn make_descriptor(package: &str, message_name: &str, field_name: &str) -> Vec { + let fds = FileDescriptorSet { + file: vec![FileDescriptorProto { + name: Some(format!("{package}.proto")), + package: Some(package.to_owned()), + message_type: vec![DescriptorProto { + name: Some(message_name.to_owned()), + field: vec![FieldDescriptorProto { + name: Some(field_name.to_owned()), + number: Some(1), + r#type: Some(9), // TYPE_STRING + label: Some(1), // LABEL_OPTIONAL + ..Default::default() + }], + ..Default::default() + }], + ..Default::default() + }], + }; + let mut buf = Vec::new(); + fds.encode(&mut buf).expect("encode descriptor"); + buf +} + +/// Encodes a DynamicMessage with a single string field. +fn encode_message(descriptor_bytes: &[u8], type_name: &str, field: &str, value: &str) -> Vec { + let pool = DescriptorPool::decode(descriptor_bytes).expect("decode pool"); + let msg_desc = pool.get_message_by_name(type_name).expect("find message"); + let mut msg = DynamicMessage::new(msg_desc); + msg.set_field_by_name(field, prost_reflect::Value::String(value.into())); + let mut buf = Vec::new(); + msg.encode(&mut buf).expect("encode message"); + buf +} + +fn start_proto_server(concurrent: bool) -> TestServer { + TestServer::start_with(ServerOptions { + protobuf: true, + concurrent, + ..Default::default() + }) +} + +// ---- sharded mode tests ---- + +#[tokio::test] +async fn register_schema_and_describe() { + let server = start_proto_server(false); + let mut c = server.connect().await; + + let desc = make_descriptor("test", "User", "name"); + + // register should return the message type names + let resp = c.cmd_raw(&[b"PROTO.REGISTER", b"users", &desc]).await; + match resp { + Frame::Array(items) => { + assert_eq!(items.len(), 1); + assert_eq!(items[0], Frame::Bulk(Bytes::from("test.User"))); + } + other => panic!("expected array, got {other:?}"), + } + + // PROTO.SCHEMAS should list the registered schema + let resp = c.cmd(&["PROTO.SCHEMAS"]).await; + match resp { + Frame::Array(items) => { + assert_eq!(items.len(), 1); + assert_eq!(items[0], Frame::Bulk(Bytes::from("users"))); + } + other => panic!("expected array, got {other:?}"), + } + + // PROTO.DESCRIBE should list message types + let resp = c.cmd(&["PROTO.DESCRIBE", "users"]).await; + match resp { + Frame::Array(items) => { + assert_eq!(items.len(), 1); + assert_eq!(items[0], Frame::Bulk(Bytes::from("test.User"))); + } + other => panic!("expected array, got {other:?}"), + } +} + +#[tokio::test] +async fn describe_unknown_schema() { + let server = start_proto_server(false); + let mut c = server.connect().await; + + let resp = c.cmd(&["PROTO.DESCRIBE", "nonexistent"]).await; + assert!(matches!(resp, Frame::Error(_))); +} + +#[tokio::test] +async fn set_get_and_type() { + 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; + + let data = encode_message(&desc, "test.User", "name", "alice"); + + // PROTO.SET + let resp = c + .cmd_raw(&[b"PROTO.SET", b"user:1", b"test.User", &data]) + .await; + assert!(matches!(resp, Frame::Simple(ref s) if s == "OK")); + + // PROTO.GET should return [type_name, data] + let resp = c.cmd(&["PROTO.GET", "user:1"]).await; + match resp { + Frame::Array(items) => { + assert_eq!(items.len(), 2); + assert_eq!(items[0], Frame::Bulk(Bytes::from("test.User"))); + assert_eq!(items[1], Frame::Bulk(Bytes::from(data))); + } + other => panic!("expected array, got {other:?}"), + } + + // PROTO.TYPE should return the type name + let resp = c.cmd(&["PROTO.TYPE", "user:1"]).await; + assert_eq!(resp, Frame::Bulk(Bytes::from("test.User"))); +} + +#[tokio::test] +async fn get_missing_key_returns_null() { + 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; + + let resp = c.cmd(&["PROTO.GET", "nonexistent"]).await; + assert!(matches!(resp, Frame::Null)); + + let resp = c.cmd(&["PROTO.TYPE", "nonexistent"]).await; + assert!(matches!(resp, Frame::Null)); +} + +#[tokio::test] +async fn wrong_type_error() { + let server = start_proto_server(false); + let mut c = server.connect().await; + + // set a regular string key + c.ok(&["SET", "str:key", "hello"]).await; + + // PROTO.GET on a string key should return WRONGTYPE + let resp = c.cmd(&["PROTO.GET", "str:key"]).await; + assert!(matches!(resp, Frame::Error(ref s) if s.starts_with("WRONGTYPE"))); + + // PROTO.TYPE on a string key should also return WRONGTYPE + let resp = c.cmd(&["PROTO.TYPE", "str:key"]).await; + assert!(matches!(resp, Frame::Error(ref s) if s.starts_with("WRONGTYPE"))); +} + +#[tokio::test] +async fn invalid_proto_bytes_rejected() { + 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; + + // send garbage bytes as proto data + let resp = c + .cmd_raw(&[b"PROTO.SET", b"bad:1", b"test.User", b"not valid proto"]) + .await; + assert!(matches!(resp, Frame::Error(_))); +} + +#[tokio::test] +async fn invalid_descriptor_rejected() { + let server = start_proto_server(false); + let mut c = server.connect().await; + + let resp = c + .cmd_raw(&[b"PROTO.REGISTER", b"bad", b"not a descriptor"]) + .await; + assert!(matches!(resp, Frame::Error(_))); +} + +#[tokio::test] +async fn set_with_nx_and_xx() { + 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; + + let alice = encode_message(&desc, "test.User", "name", "alice"); + let bob = encode_message(&desc, "test.User", "name", "bob"); + + // NX on a new key should succeed + let resp = c + .cmd_raw(&[b"PROTO.SET", b"user:nx", b"test.User", &alice, b"NX"]) + .await; + assert!(matches!(resp, Frame::Simple(ref s) if s == "OK")); + + // NX on an existing key should return null + let resp = c + .cmd_raw(&[b"PROTO.SET", b"user:nx", b"test.User", &bob, b"NX"]) + .await; + assert!(matches!(resp, Frame::Null)); + + // XX on a missing key should return null + let resp = c + .cmd_raw(&[b"PROTO.SET", b"user:xx", b"test.User", &alice, b"XX"]) + .await; + assert!(matches!(resp, Frame::Null)); + + // XX on an existing key should succeed + let resp = c + .cmd_raw(&[b"PROTO.SET", b"user:nx", b"test.User", &bob, b"XX"]) + .await; + assert!(matches!(resp, Frame::Simple(ref s) if s == "OK")); +} + +#[tokio::test] +async fn set_with_ttl() { + 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; + + let data = encode_message(&desc, "test.User", "name", "alice"); + + // set with 60s expiry + let resp = c + .cmd_raw(&[b"PROTO.SET", b"user:ttl", b"test.User", &data, b"EX", b"60"]) + .await; + assert!(matches!(resp, Frame::Simple(ref s) if s == "OK")); + + // TTL should be set + let ttl = c.get_int(&["TTL", "user:ttl"]).await; + assert!(ttl > 0 && ttl <= 60); +} + +#[tokio::test] +async fn persistence_recovery() { + let data_dir = tempfile::tempdir().unwrap(); + let path = data_dir.path().to_path_buf(); + + let desc = make_descriptor("test", "User", "name"); + let data = encode_message(&desc, "test.User", "name", "alice"); + + // start server, register schema, set proto value + { + let server = TestServer::start_with(ServerOptions { + protobuf: true, + appendonly: true, + data_dir_path: Some(path.clone()), + ..Default::default() + }); + let mut c = server.connect().await; + + c.cmd_raw(&[b"PROTO.REGISTER", b"users", &desc]).await; + + let resp = c + .cmd_raw(&[b"PROTO.SET", b"user:1", b"test.User", &data]) + .await; + assert!(matches!(resp, Frame::Simple(ref s) if s == "OK")); + + // wait for fsync + tokio::time::sleep(std::time::Duration::from_millis(300)).await; + } + + // restart and verify recovery + let server = TestServer::start_with(ServerOptions { + protobuf: true, + appendonly: true, + data_dir_path: Some(path), + ..Default::default() + }); + let mut c = server.connect().await; + + // proto value should be recovered + let resp = c.cmd(&["PROTO.GET", "user:1"]).await; + match resp { + Frame::Array(items) => { + assert_eq!(items.len(), 2); + assert_eq!(items[0], Frame::Bulk(Bytes::from("test.User"))); + assert_eq!(items[1], Frame::Bulk(Bytes::from(data))); + } + other => panic!("expected recovered proto value, got {other:?}"), + } + + drop(data_dir); +} + +// ---- concurrent mode tests ---- +// These mirror the core sharded tests to verify the concurrent handler's +// proto command routing through the engine fallback path. + +#[tokio::test] +async fn concurrent_register_schema_and_describe() { + let server = start_proto_server(true); + let mut c = server.connect().await; + + let desc = make_descriptor("test", "User", "name"); + + let resp = c.cmd_raw(&[b"PROTO.REGISTER", b"users", &desc]).await; + match resp { + Frame::Array(items) => { + assert_eq!(items.len(), 1); + assert_eq!(items[0], Frame::Bulk(Bytes::from("test.User"))); + } + other => panic!("expected array, got {other:?}"), + } + + let resp = c.cmd(&["PROTO.SCHEMAS"]).await; + match resp { + Frame::Array(items) => { + assert_eq!(items.len(), 1); + assert_eq!(items[0], Frame::Bulk(Bytes::from("users"))); + } + other => panic!("expected array, got {other:?}"), + } + + let resp = c.cmd(&["PROTO.DESCRIBE", "users"]).await; + match resp { + Frame::Array(items) => { + assert_eq!(items.len(), 1); + assert_eq!(items[0], Frame::Bulk(Bytes::from("test.User"))); + } + other => panic!("expected array, got {other:?}"), + } +} + +#[tokio::test] +async fn concurrent_set_get_and_type() { + let server = start_proto_server(true); + let mut c = server.connect().await; + + let desc = make_descriptor("test", "User", "name"); + c.cmd_raw(&[b"PROTO.REGISTER", b"users", &desc]).await; + + let data = encode_message(&desc, "test.User", "name", "alice"); + + let resp = c + .cmd_raw(&[b"PROTO.SET", b"user:1", b"test.User", &data]) + .await; + assert!(matches!(resp, Frame::Simple(ref s) if s == "OK")); + + let resp = c.cmd(&["PROTO.GET", "user:1"]).await; + match resp { + Frame::Array(items) => { + assert_eq!(items.len(), 2); + assert_eq!(items[0], Frame::Bulk(Bytes::from("test.User"))); + assert_eq!(items[1], Frame::Bulk(Bytes::from(data))); + } + other => panic!("expected array, got {other:?}"), + } + + let resp = c.cmd(&["PROTO.TYPE", "user:1"]).await; + assert_eq!(resp, Frame::Bulk(Bytes::from("test.User"))); +} + +#[tokio::test] +async fn concurrent_invalid_proto_bytes_rejected() { + let server = start_proto_server(true); + let mut c = server.connect().await; + + let desc = make_descriptor("test", "User", "name"); + c.cmd_raw(&[b"PROTO.REGISTER", b"users", &desc]).await; + + let resp = c + .cmd_raw(&[b"PROTO.SET", b"bad:1", b"test.User", b"not valid proto"]) + .await; + assert!(matches!(resp, Frame::Error(_))); +} + +#[tokio::test] +async fn concurrent_set_with_nx_and_xx() { + let server = start_proto_server(true); + let mut c = server.connect().await; + + let desc = make_descriptor("test", "User", "name"); + c.cmd_raw(&[b"PROTO.REGISTER", b"users", &desc]).await; + + let alice = encode_message(&desc, "test.User", "name", "alice"); + let bob = encode_message(&desc, "test.User", "name", "bob"); + + // NX on a new key should succeed + let resp = c + .cmd_raw(&[b"PROTO.SET", b"user:nx", b"test.User", &alice, b"NX"]) + .await; + assert!(matches!(resp, Frame::Simple(ref s) if s == "OK")); + + // NX on an existing key should return null + let resp = c + .cmd_raw(&[b"PROTO.SET", b"user:nx", b"test.User", &bob, b"NX"]) + .await; + assert!(matches!(resp, Frame::Null)); + + // XX on a missing key should return null + let resp = c + .cmd_raw(&[b"PROTO.SET", b"user:xx", b"test.User", &alice, b"XX"]) + .await; + assert!(matches!(resp, Frame::Null)); + + // XX on an existing key should succeed + let resp = c + .cmd_raw(&[b"PROTO.SET", b"user:nx", b"test.User", &bob, b"XX"]) + .await; + assert!(matches!(resp, Frame::Simple(ref s) if s == "OK")); +} + +#[tokio::test] +async fn concurrent_get_missing_key_returns_null() { + let server = start_proto_server(true); + let mut c = server.connect().await; + + let desc = make_descriptor("test", "User", "name"); + c.cmd_raw(&[b"PROTO.REGISTER", b"users", &desc]).await; + + let resp = c.cmd(&["PROTO.GET", "nonexistent"]).await; + assert!(matches!(resp, Frame::Null)); + + let resp = c.cmd(&["PROTO.TYPE", "nonexistent"]).await; + assert!(matches!(resp, Frame::Null)); +}