diff --git a/crates/ember-cluster/src/topology.rs b/crates/ember-cluster/src/topology.rs index a2d4d1de..bb079850 100644 --- a/crates/ember-cluster/src/topology.rs +++ b/crates/ember-cluster/src/topology.rs @@ -149,7 +149,7 @@ impl ClusterNode { /// Creates a new primary node with a custom bus port offset. pub fn new_primary_with_offset(id: NodeId, addr: SocketAddr, bus_port_offset: u16) -> Self { let cluster_bus_addr = - SocketAddr::new(addr.ip(), addr.port().wrapping_add(bus_port_offset)); + SocketAddr::new(addr.ip(), addr.port().saturating_add(bus_port_offset)); Self { id, addr, @@ -168,7 +168,7 @@ impl ClusterNode { /// Creates a new replica node. pub fn new_replica(id: NodeId, addr: SocketAddr, primary_id: NodeId) -> Self { - let cluster_bus_addr = SocketAddr::new(addr.ip(), addr.port() + 10000); + let cluster_bus_addr = SocketAddr::new(addr.ip(), addr.port().saturating_add(10000)); Self { id, addr, diff --git a/crates/ember-core/src/shard.rs b/crates/ember-core/src/shard.rs index 3aa36552..46effaf7 100644 --- a/crates/ember-core/src/shard.rs +++ b/crates/ember-core/src/shard.rs @@ -1033,6 +1033,20 @@ where resp } +/// Clamps a duration's millisecond value to fit in i64. +/// +/// `Duration::as_millis()` returns u128 which silently wraps when cast +/// to i64 for TTLs longer than ~292 million years. This caps at i64::MAX +/// instead, preserving "very long TTL" semantics without sign corruption. +fn duration_to_expire_ms(d: Duration) -> i64 { + let ms = d.as_millis(); + if ms > i64::MAX as u128 { + i64::MAX + } else { + ms as i64 + } +} + /// Converts a successful mutation request+response pair into an AOF record. /// Returns None for non-mutation requests or failed mutations. fn to_aof_record(req: &ShardRequest, resp: &ShardResponse) -> Option { @@ -1043,7 +1057,7 @@ fn to_aof_record(req: &ShardRequest, resp: &ShardResponse) -> Option }, ShardResponse::Ok, ) => { - let expire_ms = expire.map(|d| d.as_millis() as i64).unwrap_or(-1); + let expire_ms = expire.map(duration_to_expire_ms).unwrap_or(-1); Some(AofRecord::Set { key: key.clone(), value: value.clone(), @@ -1180,7 +1194,7 @@ fn to_aof_record(req: &ShardRequest, resp: &ShardResponse) -> Option }, ShardResponse::Ok, ) => { - let expire_ms = expire.map(|d| d.as_millis() as i64).unwrap_or(-1); + let expire_ms = expire.map(duration_to_expire_ms).unwrap_or(-1); Some(AofRecord::ProtoSet { key: key.clone(), type_name: type_name.clone(), @@ -1205,7 +1219,7 @@ fn to_aof_record(req: &ShardRequest, resp: &ShardResponse) -> Option expire, }, ) => { - let expire_ms = expire.map(|d| d.as_millis() as i64).unwrap_or(-1); + let expire_ms = expire.map(duration_to_expire_ms).unwrap_or(-1); Some(AofRecord::ProtoSet { key: key.clone(), type_name: type_name.clone(), diff --git a/crates/ember-persistence/src/snapshot.rs b/crates/ember-persistence/src/snapshot.rs index b5a594ae..9eb1ba22 100644 --- a/crates/ember-persistence/src/snapshot.rs +++ b/crates/ember-persistence/src/snapshot.rs @@ -876,6 +876,51 @@ mod tests { assert_eq!(p, PathBuf::from("/data/shard-5.snap")); } + #[test] + fn truncated_snapshot_detected() { + let dir = temp_dir(); + let path = dir.path().join("truncated.snap"); + + // write a valid 2-entry snapshot + { + let mut writer = SnapshotWriter::create(&path, 0).unwrap(); + writer + .write_entry(&SnapEntry { + key: "a".into(), + value: SnapValue::String(Bytes::from("1")), + expire_ms: -1, + }) + .unwrap(); + writer + .write_entry(&SnapEntry { + key: "b".into(), + value: SnapValue::String(Bytes::from("2")), + expire_ms: 5000, + }) + .unwrap(); + writer.finish().unwrap(); + } + + // truncate the file mid-way through the second entry + let data = fs::read(&path).unwrap(); + let truncated = &data[..data.len() - 20]; + fs::write(&path, truncated).unwrap(); + + let mut reader = SnapshotReader::open(&path).unwrap(); + assert_eq!(reader.entry_count, 2); + + // first entry should still be readable + let first = reader.read_entry().unwrap(); + assert!(first.is_some()); + + // second entry should fail with an EOF-related error + let err = reader.read_entry().unwrap_err(); + assert!( + matches!(err, FormatError::UnexpectedEof | FormatError::Io(_)), + "expected EOF error, got {err:?}" + ); + } + #[cfg(feature = "encryption")] mod encrypted { use super::*; diff --git a/crates/ember-server/src/cluster.rs b/crates/ember-server/src/cluster.rs index b2d66687..842d49c2 100644 --- a/crates/ember-server/src/cluster.rs +++ b/crates/ember-server/src/cluster.rs @@ -53,7 +53,8 @@ impl ClusterCoordinator { let (event_tx, event_rx) = mpsc::channel(256); let port_offset = gossip_config.gossip_port_offset; - let gossip_addr = SocketAddr::new(bind_addr.ip(), bind_addr.port() + port_offset); + let gossip_addr = + SocketAddr::new(bind_addr.ip(), bind_addr.port().saturating_add(port_offset)); let gossip = GossipEngine::new(local_id, gossip_addr, gossip_config, event_tx); @@ -144,9 +145,16 @@ impl ClusterCoordinator { let mut gossip = self.gossip.lock().await; let new_id = NodeId::new(); - let gossip_port_offset = 10000u16; // default - let gossip_addr = - SocketAddr::new(addr.ip(), addr.port().saturating_add(gossip_port_offset)); + let gossip_port = match port.checked_add(self.gossip_port_offset) { + Some(p) => p, + None => { + return Frame::Error(format!( + "ERR port {port} + offset {} overflows", + self.gossip_port_offset + )) + } + }; + let gossip_addr = SocketAddr::new(addr.ip(), gossip_port); gossip.add_seed(new_id, gossip_addr); @@ -412,8 +420,10 @@ impl ClusterCoordinator { bind_addr: SocketAddr, mut event_rx: mpsc::Receiver, ) { - let gossip_addr = - SocketAddr::new(bind_addr.ip(), bind_addr.port() + self.gossip_port_offset); + let gossip_addr = SocketAddr::new( + bind_addr.ip(), + bind_addr.port().saturating_add(self.gossip_port_offset), + ); let socket = match UdpSocket::bind(gossip_addr).await { Ok(s) => Arc::new(s), diff --git a/crates/ember-server/src/main.rs b/crates/ember-server/src/main.rs index 52c93900..18afb031 100644 --- a/crates/ember-server/src/main.rs +++ b/crates/ember-server/src/main.rs @@ -187,9 +187,13 @@ async fn main() { } } - let addr: SocketAddr = format!("{}:{}", args.host, args.port) - .parse() - .expect("invalid bind address"); + let addr: SocketAddr = match format!("{}:{}", args.host, args.port).parse() { + Ok(a) => a, + Err(e) => { + eprintln!("invalid bind address '{}:{}': {e}", args.host, args.port); + std::process::exit(1); + } + }; let max_memory = args.max_memory.as_deref().map(|s| { parse_byte_size(s).unwrap_or_else(|e| { @@ -301,9 +305,17 @@ async fn main() { // install prometheus metrics exporter if --metrics-port is set if let Some(metrics_port) = args.metrics_port { - let metrics_addr: std::net::SocketAddr = format!("{}:{}", args.host, metrics_port) - .parse() - .expect("invalid metrics bind address"); + let metrics_addr: std::net::SocketAddr = + match format!("{}:{}", args.host, metrics_port).parse() { + Ok(a) => a, + Err(e) => { + eprintln!( + "invalid metrics bind address '{}:{metrics_port}': {e}", + args.host + ); + std::process::exit(1); + } + }; if let Err(e) = metrics::install_exporter(metrics_addr) { eprintln!("failed to start metrics exporter: {e}"); std::process::exit(1); @@ -346,9 +358,16 @@ async fn main() { } }; - let tls_addr: SocketAddr = format!("{}:{}", args.host, tls_port) - .parse() - .expect("invalid TLS bind address"); + let tls_addr: SocketAddr = match format!("{}:{}", args.host, tls_port).parse() { + Ok(a) => a, + Err(e) => { + eprintln!( + "invalid TLS bind address '{}:{tls_port}': {e}", + args.host + ); + std::process::exit(1); + } + }; info!( tls_port = tls_port, @@ -382,6 +401,14 @@ async fn main() { // build cluster coordinator if cluster mode is enabled let cluster: Option> = if args.cluster_enabled { + if args.port.checked_add(args.cluster_port_offset).is_none() { + eprintln!( + "error: port {} + cluster-port-offset {} exceeds u16 range", + args.port, args.cluster_port_offset + ); + std::process::exit(1); + } + let local_id = NodeId::new(); let gossip_config = GossipConfig { gossip_port_offset: args.cluster_port_offset, diff --git a/crates/ember-server/src/server.rs b/crates/ember-server/src/server.rs index 5d559c56..fece1590 100644 --- a/crates/ember-server/src/server.rs +++ b/crates/ember-server/src/server.rs @@ -267,10 +267,13 @@ pub async fn run( } } - // wait for all connection handlers to finish by acquiring all permits + // wait for all connection handlers to finish, with a timeout info!("waiting for active connections to close..."); - let _ = semaphore.acquire_many(max_conn as u32).await; - info!("all connections drained, shutting down"); + let drain = semaphore.acquire_many(max_conn as u32); + match tokio::time::timeout(std::time::Duration::from_secs(30), drain).await { + Ok(_) => info!("all connections drained, shutting down"), + Err(_) => warn!("shutdown timeout after 30s, forcing exit"), + } Ok(()) } @@ -493,8 +496,11 @@ pub async fn run_concurrent( } info!("waiting for active connections to close..."); - let _ = semaphore.acquire_many(max_conn as u32).await; - info!("all connections drained, shutting down"); + let drain = semaphore.acquire_many(max_conn as u32); + match tokio::time::timeout(std::time::Duration::from_secs(30), drain).await { + Ok(_) => info!("all connections drained, shutting down"), + Err(_) => warn!("shutdown timeout after 30s, forcing exit"), + } Ok(()) }