Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions crates/ember-cluster/src/topology.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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,
Expand Down
20 changes: 17 additions & 3 deletions crates/ember-core/src/shard.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<AofRecord> {
Expand All @@ -1043,7 +1057,7 @@ fn to_aof_record(req: &ShardRequest, resp: &ShardResponse) -> Option<AofRecord>
},
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(),
Expand Down Expand Up @@ -1180,7 +1194,7 @@ fn to_aof_record(req: &ShardRequest, resp: &ShardResponse) -> Option<AofRecord>
},
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(),
Expand All @@ -1205,7 +1219,7 @@ fn to_aof_record(req: &ShardRequest, resp: &ShardResponse) -> Option<AofRecord>
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(),
Expand Down
45 changes: 45 additions & 0 deletions crates/ember-persistence/src/snapshot.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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::*;
Expand Down
22 changes: 16 additions & 6 deletions crates/ember-server/src/cluster.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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);

Expand Down Expand Up @@ -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);

Expand Down Expand Up @@ -412,8 +420,10 @@ impl ClusterCoordinator {
bind_addr: SocketAddr,
mut event_rx: mpsc::Receiver<GossipEvent>,
) {
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),
Expand Down
45 changes: 36 additions & 9 deletions crates/ember-server/src/main.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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| {
Expand Down Expand Up @@ -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);
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -382,6 +401,14 @@ async fn main() {

// build cluster coordinator if cluster mode is enabled
let cluster: Option<Arc<ClusterCoordinator>> = 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,
Expand Down
16 changes: 11 additions & 5 deletions crates/ember-server/src/server.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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(())
}
Expand Down Expand Up @@ -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(())
}
Expand Down
Loading