diff --git a/crates/ember-cluster/src/gossip.rs b/crates/ember-cluster/src/gossip.rs index a04a8668..859ad71b 100644 --- a/crates/ember-cluster/src/gossip.rs +++ b/crates/ember-cluster/src/gossip.rs @@ -422,7 +422,7 @@ impl GossipEngine { if *node == self.local_id { // Refute suspicion by incrementing our incarnation if *incarnation >= self.incarnation { - self.incarnation = incarnation + 1; + self.incarnation = incarnation.saturating_add(1); self.queue_update(NodeUpdate::Alive { node: self.local_id, addr: self.local_addr, @@ -447,7 +447,7 @@ impl GossipEngine { NodeUpdate::Dead { node, incarnation } => { if *node == self.local_id { // Refute death claim - self.incarnation = incarnation + 1; + self.incarnation = incarnation.saturating_add(1); self.queue_update(NodeUpdate::Alive { node: self.local_id, addr: self.local_addr, diff --git a/crates/ember-cluster/src/message.rs b/crates/ember-cluster/src/message.rs index ff3120b7..fcd0cc76 100644 --- a/crates/ember-cluster/src/message.rs +++ b/crates/ember-cluster/src/message.rs @@ -10,6 +10,33 @@ use bytes::{Buf, BufMut, Bytes, BytesMut}; use crate::{NodeId, SlotRange}; +/// Maximum number of members in a Welcome message or updates in a Ping/Ack. +/// Prevents allocation bombs from crafted messages. +const MAX_COLLECTION_COUNT: usize = 1024; + +// Safe read helpers that return io::Error instead of panicking on truncated input. + +fn safe_get_u8(buf: &mut &[u8]) -> io::Result { + if buf.is_empty() { + return Err(io::Error::new(io::ErrorKind::UnexpectedEof, "need 1 byte")); + } + Ok(buf.get_u8()) +} + +fn safe_get_u16_le(buf: &mut &[u8]) -> io::Result { + if buf.len() < 2 { + return Err(io::Error::new(io::ErrorKind::UnexpectedEof, "need 2 bytes")); + } + Ok(buf.get_u16_le()) +} + +fn safe_get_u64_le(buf: &mut &[u8]) -> io::Result { + if buf.len() < 8 { + return Err(io::Error::new(io::ErrorKind::UnexpectedEof, "need 8 bytes")); + } + Ok(buf.get_u64_le()) +} + /// Message types for the SWIM gossip protocol. #[derive(Debug, Clone, PartialEq)] pub enum GossipMessage { @@ -160,10 +187,10 @@ impl GossipMessage { )); } - let msg_type = buf.get_u8(); + let msg_type = safe_get_u8(&mut buf)?; match msg_type { MSG_PING => { - let seq = buf.get_u64_le(); + let seq = safe_get_u64_le(&mut buf)?; let sender = decode_node_id(&mut buf)?; let updates = decode_updates(&mut buf)?; Ok(GossipMessage::Ping { @@ -173,7 +200,7 @@ impl GossipMessage { }) } MSG_PING_REQ => { - let seq = buf.get_u64_le(); + let seq = safe_get_u64_le(&mut buf)?; let sender = decode_node_id(&mut buf)?; let target = decode_node_id(&mut buf)?; let target_addr = decode_socket_addr(&mut buf)?; @@ -185,7 +212,7 @@ impl GossipMessage { }) } MSG_ACK => { - let seq = buf.get_u64_le(); + let seq = safe_get_u64_le(&mut buf)?; let sender = decode_node_id(&mut buf)?; let updates = decode_updates(&mut buf)?; Ok(GossipMessage::Ack { @@ -204,7 +231,13 @@ impl GossipMessage { } MSG_WELCOME => { let sender = decode_node_id(&mut buf)?; - let count = buf.get_u16_le() as usize; + let count = safe_get_u16_le(&mut buf)? as usize; + if count > MAX_COLLECTION_COUNT { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + format!("member count {count} exceeds limit"), + )); + } let mut members = Vec::with_capacity(count); for _ in 0..count { members.push(decode_member_info(&mut buf)?); @@ -327,13 +360,13 @@ fn encode_update(buf: &mut BytesMut, update: &NodeUpdate) { } fn decode_updates(buf: &mut &[u8]) -> io::Result> { - if buf.len() < 2 { + let count = safe_get_u16_le(buf)? as usize; + if count > MAX_COLLECTION_COUNT { return Err(io::Error::new( - io::ErrorKind::UnexpectedEof, - "not enough bytes for update count", + io::ErrorKind::InvalidData, + format!("update count {count} exceeds limit"), )); } - let count = buf.get_u16_le() as usize; let mut updates = Vec::with_capacity(count); for _ in 0..count { updates.push(decode_update(buf)?); @@ -342,18 +375,12 @@ fn decode_updates(buf: &mut &[u8]) -> io::Result> { } fn decode_update(buf: &mut &[u8]) -> io::Result { - if buf.is_empty() { - return Err(io::Error::new( - io::ErrorKind::UnexpectedEof, - "not enough bytes for update type", - )); - } - let update_type = buf.get_u8(); + let update_type = safe_get_u8(buf)?; match update_type { UPDATE_ALIVE => { let node = decode_node_id(buf)?; let addr = decode_socket_addr(buf)?; - let incarnation = buf.get_u64_le(); + let incarnation = safe_get_u64_le(buf)?; Ok(NodeUpdate::Alive { node, addr, @@ -362,12 +389,12 @@ fn decode_update(buf: &mut &[u8]) -> io::Result { } UPDATE_SUSPECT => { let node = decode_node_id(buf)?; - let incarnation = buf.get_u64_le(); + let incarnation = safe_get_u64_le(buf)?; Ok(NodeUpdate::Suspect { node, incarnation }) } UPDATE_DEAD => { let node = decode_node_id(buf)?; - let incarnation = buf.get_u64_le(); + let incarnation = safe_get_u64_le(buf)?; Ok(NodeUpdate::Dead { node, incarnation }) } UPDATE_LEFT => { @@ -396,13 +423,19 @@ fn encode_member_info(buf: &mut BytesMut, member: &MemberInfo) { fn decode_member_info(buf: &mut &[u8]) -> io::Result { let id = decode_node_id(buf)?; let addr = decode_socket_addr(buf)?; - let incarnation = buf.get_u64_le(); - let is_primary = buf.get_u8() != 0; - let slot_count = buf.get_u16_le() as usize; + let incarnation = safe_get_u64_le(buf)?; + let is_primary = safe_get_u8(buf)? != 0; + let slot_count = safe_get_u16_le(buf)? as usize; + if slot_count > MAX_COLLECTION_COUNT { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + format!("slot range count {slot_count} exceeds limit"), + )); + } let mut slots = Vec::with_capacity(slot_count); for _ in 0..slot_count { - let start = buf.get_u16_le(); - let end = buf.get_u16_le(); + let start = safe_get_u16_le(buf)?; + let end = safe_get_u16_le(buf)?; slots.push(SlotRange::try_new(start, end)?); } Ok(MemberInfo { diff --git a/crates/ember-core/src/concurrent.rs b/crates/ember-core/src/concurrent.rs index 76252294..40bf2a8f 100644 --- a/crates/ember-core/src/concurrent.rs +++ b/crates/ember-core/src/concurrent.rs @@ -169,7 +169,7 @@ impl ConcurrentKeyspace { if entry.is_expired() { return false; } - entry.expires_at_ms = time::now_ms() + seconds * 1000; + entry.expires_at_ms = time::now_ms().saturating_add(seconds.saturating_mul(1000)); true } else { false diff --git a/crates/ember-core/src/engine.rs b/crates/ember-core/src/engine.rs index 400c496a..f5976660 100644 --- a/crates/ember-core/src/engine.rs +++ b/crates/ember-core/src/engine.rs @@ -185,6 +185,11 @@ impl Engine { Ok(results) } + /// Returns true if both keys are owned by the same shard. + pub fn same_shard(&self, key1: &str, key2: &str) -> bool { + self.shard_for_key(key1) == self.shard_for_key(key2) + } + /// Determines which shard owns a given key. fn shard_for_key(&self, key: &str) -> usize { shard_index(key, self.shards.len()) diff --git a/crates/ember-core/src/keyspace.rs b/crates/ember-core/src/keyspace.rs index bc5797db..90c52630 100644 --- a/crates/ember-core/src/keyspace.rs +++ b/crates/ember-core/src/keyspace.rs @@ -513,7 +513,7 @@ impl Keyspace { if entry.expires_at_ms == 0 { self.expiry_count += 1; } - entry.expires_at_ms = time::now_ms() + seconds * 1000; + entry.expires_at_ms = time::now_ms().saturating_add(seconds.saturating_mul(1000)); true } None => false, @@ -589,7 +589,7 @@ impl Keyspace { if entry.expires_at_ms == 0 { self.expiry_count += 1; } - entry.expires_at_ms = time::now_ms() + millis; + entry.expires_at_ms = time::now_ms().saturating_add(millis); true } None => false, @@ -1471,8 +1471,12 @@ impl Keyspace { }; if is_empty { + if let Some(removed_entry) = self.entries.remove(key) { + if removed_entry.expires_at_ms != 0 { + self.expiry_count = self.expiry_count.saturating_sub(1); + } + } self.memory.remove_with_size(old_entry_size); - self.entries.remove(key); } else { let new_entry_size = memory::entry_size(key, &self.entries[key].value); self.memory.adjust(old_entry_size, new_entry_size); @@ -1732,8 +1736,12 @@ impl Keyspace { }; if is_empty { + if let Some(removed_entry) = self.entries.remove(key) { + if removed_entry.expires_at_ms != 0 { + self.expiry_count = self.expiry_count.saturating_sub(1); + } + } self.memory.remove_with_size(old_entry_size); - self.entries.remove(key); } else { let new_entry_size = memory::entry_size(key, &self.entries[key].value); self.memory.adjust(old_entry_size, new_entry_size); diff --git a/crates/ember-core/src/shard.rs b/crates/ember-core/src/shard.rs index 46effaf7..68c1ed76 100644 --- a/crates/ember-core/src/shard.rs +++ b/crates/ember-core/src/shard.rs @@ -705,11 +705,14 @@ fn dispatch( Err(IncrError::OutOfMemory) => ShardResponse::OutOfMemory, Err(e) => ShardResponse::Err(e.to_string()), }, - ShardRequest::DecrBy { 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::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()), + }, + None => ShardResponse::Err("ERR increment or decrement would overflow".into()), }, ShardRequest::IncrByFloat { key, delta } => match ks.incr_by_float(key, *delta) { Ok(val) => ShardResponse::BulkString(val), diff --git a/crates/ember-core/src/time.rs b/crates/ember-core/src/time.rs index 0edf1322..317df4dd 100644 --- a/crates/ember-core/src/time.rs +++ b/crates/ember-core/src/time.rs @@ -26,8 +26,11 @@ pub fn is_expired(expires_at_ms: u64) -> bool { /// Converts a Duration to an absolute expiry timestamp. #[inline] pub fn expiry_from_duration(ttl: Option) -> u64 { - ttl.map(|d| now_ms() + d.as_millis() as u64) - .unwrap_or(NO_EXPIRY) + ttl.map(|d| { + let ms = d.as_millis().min(u64::MAX as u128) as u64; + now_ms().saturating_add(ms) + }) + .unwrap_or(NO_EXPIRY) } /// Returns remaining TTL in seconds, or None if no expiry. diff --git a/crates/ember-persistence/src/aof.rs b/crates/ember-persistence/src/aof.rs index 0b4a4c6d..b667f678 100644 --- a/crates/ember-persistence/src/aof.rs +++ b/crates/ember-persistence/src/aof.rs @@ -304,6 +304,13 @@ 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); @@ -331,7 +338,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(count as usize); + let mut values = Vec::with_capacity(Self::capped_capacity(count)); for _ in 0..count { values.push(Bytes::from(format::read_bytes(&mut cursor)?)); } @@ -352,7 +359,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(count as usize); + let mut members = Vec::with_capacity(Self::capped_capacity(count)); for _ in 0..count { let score = format::read_f64(&mut cursor)?; let member = read_string(&mut cursor, "member")?; @@ -363,7 +370,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(count as usize); + let mut members = Vec::with_capacity(Self::capped_capacity(count)); for _ in 0..count { members.push(read_string(&mut cursor, "member")?); } @@ -389,7 +396,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(count as usize); + let mut fields = Vec::with_capacity(Self::capped_capacity(count)); for _ in 0..count { let field = read_string(&mut cursor, "field")?; let value = Bytes::from(format::read_bytes(&mut cursor)?); @@ -400,7 +407,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(count as usize); + let mut fields = Vec::with_capacity(Self::capped_capacity(count)); for _ in 0..count { fields.push(read_string(&mut cursor, "field")?); } @@ -415,7 +422,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(count as usize); + let mut members = Vec::with_capacity(Self::capped_capacity(count)); for _ in 0..count { members.push(read_string(&mut cursor, "member")?); } @@ -424,7 +431,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(count as usize); + let mut members = Vec::with_capacity(Self::capped_capacity(count)); for _ in 0..count { members.push(read_string(&mut cursor, "member")?); } diff --git a/crates/ember-persistence/src/recovery.rs b/crates/ember-persistence/src/recovery.rs index d8597bd8..90ddb831 100644 --- a/crates/ember-persistence/src/recovery.rs +++ b/crates/ember-persistence/src/recovery.rs @@ -361,7 +361,7 @@ fn replay_aof( apply_incr(map, key, delta); } AofRecord::DecrBy { key, delta } => { - apply_incr(map, key, -delta); + apply_incr(map, key, delta.saturating_neg()); } AofRecord::Append { key, value } => { let entry = map diff --git a/crates/ember-persistence/src/snapshot.rs b/crates/ember-persistence/src/snapshot.rs index 9eb1ba22..1d9f8f36 100644 --- a/crates/ember-persistence/src/snapshot.rs +++ b/crates/ember-persistence/src/snapshot.rs @@ -39,6 +39,13 @@ 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 { @@ -248,6 +255,14 @@ impl SnapshotWriter { // atomic rename fs::rename(&self.tmp_path, &self.final_path)?; + + // fsync the parent directory so the rename is durable across crashes + if let Some(parent) = self.final_path.parent() { + if let Ok(dir) = File::open(parent) { + let _ = dir.sync_all(); + } + } + self.finished = true; Ok(()) } @@ -367,7 +382,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(count as usize); + let mut deque = VecDeque::with_capacity(capped_capacity(count)); for _ in 0..count { let item = format::read_bytes(&mut self.reader)?; format::write_bytes(&mut buf, &item)?; @@ -378,7 +393,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(count as usize); + let mut members = Vec::with_capacity(capped_capacity(count)); for _ in 0..count { let score = format::read_f64(&mut self.reader)?; format::write_f64(&mut buf, score)?; @@ -397,7 +412,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(count as usize); + let mut map = HashMap::with_capacity(capped_capacity(count)); for _ in 0..count { let field_bytes = format::read_bytes(&mut self.reader)?; format::write_bytes(&mut buf, &field_bytes)?; @@ -416,7 +431,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(count as usize); + let mut set = HashSet::with_capacity(capped_capacity(count)); for _ in 0..count { let member_bytes = format::read_bytes(&mut self.reader)?; format::write_bytes(&mut buf, &member_bytes)?; @@ -526,7 +541,7 @@ impl SnapshotReader { } TYPE_LIST => { let count = format::read_u32(&mut cursor)?; - let mut deque = VecDeque::with_capacity(count as usize); + let mut deque = VecDeque::with_capacity(capped_capacity(count)); for _ in 0..count { deque.push_back(Bytes::from(format::read_bytes(&mut cursor)?)); } @@ -534,7 +549,7 @@ impl SnapshotReader { } TYPE_SORTED_SET => { let count = format::read_u32(&mut cursor)?; - let mut members = Vec::with_capacity(count as usize); + let mut members = Vec::with_capacity(capped_capacity(count)); for _ in 0..count { let score = format::read_f64(&mut cursor)?; let member_bytes = format::read_bytes(&mut cursor)?; @@ -550,7 +565,7 @@ impl SnapshotReader { } TYPE_HASH => { let count = format::read_u32(&mut cursor)?; - let mut map = HashMap::with_capacity(count as usize); + let mut map = HashMap::with_capacity(capped_capacity(count)); for _ in 0..count { let field_bytes = format::read_bytes(&mut cursor)?; let field = String::from_utf8(field_bytes).map_err(|_| { @@ -566,7 +581,7 @@ impl SnapshotReader { } TYPE_SET => { let count = format::read_u32(&mut cursor)?; - let mut set = HashSet::with_capacity(count as usize); + let mut set = HashSet::with_capacity(capped_capacity(count)); for _ in 0..count { let member_bytes = format::read_bytes(&mut cursor)?; let member = String::from_utf8(member_bytes).map_err(|_| { diff --git a/crates/ember-protocol/src/error.rs b/crates/ember-protocol/src/error.rs index 50e602e9..e5b0c734 100644 --- a/crates/ember-protocol/src/error.rs +++ b/crates/ember-protocol/src/error.rs @@ -37,4 +37,12 @@ pub enum ProtocolError { /// The frame exceeds the maximum nesting depth. #[error("frame nesting depth exceeds limit of {0}")] NestingTooDeep(usize), + + /// An array or map declares more elements than the allowed maximum. + #[error("array/map element count {0} exceeds limit")] + TooManyElements(usize), + + /// A bulk string exceeds the maximum allowed length. + #[error("bulk string length {0} exceeds limit")] + BulkStringTooLarge(usize), } diff --git a/crates/ember-protocol/src/parse.rs b/crates/ember-protocol/src/parse.rs index e12eeef3..d47d98fc 100644 --- a/crates/ember-protocol/src/parse.rs +++ b/crates/ember-protocol/src/parse.rs @@ -18,6 +18,14 @@ use crate::types::Frame; /// from malicious or malformed deeply-nested frames. const MAX_NESTING_DEPTH: usize = 64; +/// Maximum number of elements in an array or map. Prevents memory +/// amplification attacks where tiny elements (3 bytes each) create +/// disproportionately large Vec allocations. +const MAX_ARRAY_ELEMENTS: usize = 1_048_576; + +/// Maximum length of a bulk string in bytes (512 MB, matching Redis). +const MAX_BULK_LEN: i64 = 512 * 1024 * 1024; + /// Checks whether `buf` contains a complete RESP3 frame and parses it. /// /// Returns `Ok(Some(frame))` if a complete frame was parsed, @@ -73,6 +81,9 @@ fn check_bulk(cursor: &mut Cursor<&[u8]>) -> Result<(), ProtocolError> { if len < 0 { return Err(ProtocolError::InvalidFrameLength(len)); } + if len > MAX_BULK_LEN { + return Err(ProtocolError::BulkStringTooLarge(len as usize)); + } let len = len as usize; // need `len` bytes of data + \r\n @@ -102,6 +113,9 @@ fn check_array(cursor: &mut Cursor<&[u8]>, depth: usize) -> Result<(), ProtocolE if count < 0 { return Err(ProtocolError::InvalidFrameLength(count)); } + if count as usize > MAX_ARRAY_ELEMENTS { + return Err(ProtocolError::TooManyElements(count as usize)); + } for _ in 0..count { check(cursor, next_depth)?; @@ -119,6 +133,9 @@ fn check_map(cursor: &mut Cursor<&[u8]>, depth: usize) -> Result<(), ProtocolErr if count < 0 { return Err(ProtocolError::InvalidFrameLength(count)); } + if count as usize > MAX_ARRAY_ELEMENTS { + return Err(ProtocolError::TooManyElements(count as usize)); + } for _ in 0..count { check(cursor, next_depth)?; // key diff --git a/crates/ember-server/src/connection.rs b/crates/ember-server/src/connection.rs index 31928590..f06819d0 100644 --- a/crates/ember-server/src/connection.rs +++ b/crates/ember-server/src/connection.rs @@ -219,8 +219,21 @@ where } } - // check for new commands from the client - result = stream.read_buf(buf) => { + // check for new commands from the client (with idle timeout) + result = tokio::time::timeout(IDLE_TIMEOUT, stream.read_buf(buf)) => { + let result = match result { + Ok(inner) => inner, + Err(_) => { + // idle timeout — clean up and close + cleanup_subscriptions(pubsub, &channel_rxs, &pattern_rxs); + return Ok(()); + } + }; + // guard against unbounded buffer growth + if buf.len() > MAX_BUF_SIZE { + cleanup_subscriptions(pubsub, &channel_rxs, &pattern_rxs); + return Ok(()); + } match result { Ok(0) => { // client disconnected — clean up subscriptions @@ -951,15 +964,19 @@ async fn execute( } Command::Rename { key, newkey } => { - let req = ShardRequest::Rename { - key: key.clone(), - newkey, - }; - match engine.route(&key, req).await { - Ok(ShardResponse::Ok) => Frame::Simple("OK".into()), - Ok(ShardResponse::Err(msg)) => Frame::Error(msg), - Ok(other) => Frame::Error(format!("ERR unexpected shard response: {other:?}")), - Err(e) => Frame::Error(format!("ERR {e}")), + if !engine.same_shard(&key, &newkey) { + Frame::Error("ERR source and destination keys must hash to the same shard".into()) + } else { + let req = ShardRequest::Rename { + key: key.clone(), + newkey, + }; + match engine.route(&key, req).await { + Ok(ShardResponse::Ok) => Frame::Simple("OK".into()), + Ok(ShardResponse::Err(msg)) => Frame::Error(msg), + Ok(other) => Frame::Error(format!("ERR unexpected shard response: {other:?}")), + Err(e) => Frame::Error(format!("ERR {e}")), + } } } diff --git a/crates/ember-server/src/server.rs b/crates/ember-server/src/server.rs index fece1590..e677172c 100644 --- a/crates/ember-server/src/server.rs +++ b/crates/ember-server/src/server.rs @@ -6,7 +6,7 @@ use std::net::SocketAddr; use std::sync::atomic::{AtomicU64, Ordering}; use std::sync::Arc; -use std::time::Instant; +use std::time::{Duration, Instant}; use ember_core::{ConcurrentKeyspace, Engine, EngineConfig, EvictionPolicy}; use tokio::io::AsyncWriteExt; @@ -246,16 +246,23 @@ pub async fn run( let pubsub = Arc::clone(&pubsub); tokio::spawn(async move { - // perform TLS handshake - match acceptor.accept(stream).await { - Ok(tls_stream) => { + // perform TLS handshake with timeout to prevent slowloris + let handshake = tokio::time::timeout( + Duration::from_secs(10), + acceptor.accept(stream), + ); + match handshake.await { + Ok(Ok(tls_stream)) => { if let Err(e) = connection::handle(tls_stream, engine, &ctx, &slow_log, &pubsub).await { error!("TLS connection error from {peer}: {e}"); } } - Err(e) => { + Ok(Err(e)) => { warn!("TLS handshake failed from {peer}: {e}"); } + Err(_) => { + warn!("TLS handshake timed out from {peer}"); + } } ctx.connections_active.fetch_sub(1, Ordering::Relaxed); if ctx.metrics_enabled { @@ -473,17 +480,24 @@ pub async fn run_concurrent( let pubsub = Arc::clone(&pubsub); tokio::spawn(async move { - match acceptor.accept(stream).await { - Ok(tls_stream) => { + let handshake = tokio::time::timeout( + Duration::from_secs(10), + acceptor.accept(stream), + ); + match handshake.await { + Ok(Ok(tls_stream)) => { if let Err(e) = crate::concurrent_handler::handle( tls_stream, keyspace, engine, &ctx, &slow_log, &pubsub ).await { error!("TLS connection error from {peer}: {e}"); } } - Err(e) => { + Ok(Err(e)) => { warn!("TLS handshake failed from {peer}: {e}"); } + Err(_) => { + warn!("TLS handshake timed out from {peer}"); + } } ctx.connections_active.fetch_sub(1, Ordering::Relaxed); if ctx.metrics_enabled {