diff --git a/crates/ember-server/src/grpc.rs b/crates/ember-server/src/grpc.rs index 65d6f4ca..8a1fc8ba 100644 --- a/crates/ember-server/src/grpc.rs +++ b/crates/ember-server/src/grpc.rs @@ -14,6 +14,7 @@ use ember_core::{Engine, ShardRequest, ShardResponse, TtlResult, Value}; use tokio_stream::wrappers::ReceiverStream; use tonic::{Request, Response, Status, Streaming}; +use crate::pubsub::PubSubManager; use crate::server::ServerContext; use crate::slowlog::SlowLog; @@ -29,14 +30,21 @@ pub struct EmberService { engine: Engine, ctx: Arc, slow_log: Arc, + pubsub: Arc, } impl EmberService { - pub fn new(engine: Engine, ctx: Arc, slow_log: Arc) -> Self { + pub fn new( + engine: Engine, + ctx: Arc, + slow_log: Arc, + pubsub: Arc, + ) -> Self { Self { engine, ctx, slow_log, + pubsub, } } @@ -1681,6 +1689,241 @@ impl EmberCache for EmberService { Ok(Response::new(InfoResponse { info })) } + // ----------------------------------------------------------------------- + // additional commands + // ----------------------------------------------------------------------- + + async fn echo(&self, request: Request) -> Result, Status> { + Ok(Response::new(EchoResponse { + message: request.into_inner().message, + })) + } + + async fn decr(&self, request: Request) -> Result, Status> { + let start = Instant::now(); + let key = request.into_inner().key; + validate_key(&key)?; + let resp = self + .route(&key, ShardRequest::Decr { key: key.clone() }) + .await?; + self.record_command(start, "DECR"); + + match resp { + ShardResponse::Integer(v) => Ok(Response::new(IntResponse { value: v })), + other => Err(unexpected_response(&other)), + } + } + + async fn unlink( + &self, + request: Request, + ) -> Result, Status> { + let start = Instant::now(); + let keys = request.into_inner().keys; + for k in &keys { + validate_key(k)?; + } + + let responses = self + .engine + .route_multi(&keys, |k| ShardRequest::Unlink { key: k }) + .await + .map_err(|_| Status::unavailable("shard unavailable"))?; + + let mut deleted = 0i64; + for resp in responses { + if let ShardResponse::Bool(true) = resp { + deleted += 1; + } + } + self.record_command(start, "UNLINK"); + Ok(Response::new(DelResponse { deleted })) + } + + async fn bg_save( + &self, + _request: Request, + ) -> Result, Status> { + let start = Instant::now(); + self.broadcast(|| ShardRequest::Snapshot).await?; + self.record_command(start, "BGSAVE"); + Ok(Response::new(StatusResponse { + status: "Background saving started".to_string(), + })) + } + + async fn bg_rewrite_aof( + &self, + _request: Request, + ) -> Result, Status> { + let start = Instant::now(); + self.broadcast(|| ShardRequest::RewriteAof).await?; + self.record_command(start, "BGREWRITEAOF"); + Ok(Response::new(StatusResponse { + status: "Background append only file rewriting started".to_string(), + })) + } + + // ----------------------------------------------------------------------- + // slowlog + // ----------------------------------------------------------------------- + + async fn slow_log_get( + &self, + request: Request, + ) -> Result, Status> { + let count = request.into_inner().count.map(|c| c as usize); + let entries = self.slow_log.get(count); + Ok(Response::new(SlowLogGetResponse { + entries: entries + .into_iter() + .map(|e| { + let ts = e + .timestamp + .duration_since(std::time::UNIX_EPOCH) + .unwrap_or_default() + .as_secs(); + SlowLogEntry { + id: e.id, + timestamp_unix: ts, + duration_micros: e.duration.as_micros() as u64, + command: e.command, + } + }) + .collect(), + })) + } + + async fn slow_log_len( + &self, + _request: Request, + ) -> Result, Status> { + Ok(Response::new(IntResponse { + value: self.slow_log.len() as i64, + })) + } + + async fn slow_log_reset( + &self, + _request: Request, + ) -> Result, Status> { + self.slow_log.reset(); + Ok(Response::new(StatusResponse { + status: "OK".to_string(), + })) + } + + // ----------------------------------------------------------------------- + // pub/sub + // ----------------------------------------------------------------------- + + async fn publish( + &self, + request: Request, + ) -> Result, Status> { + let req = request.into_inner(); + let count = self.pubsub.publish(&req.channel, Bytes::from(req.message)); + Ok(Response::new(IntResponse { + value: count as i64, + })) + } + + type SubscribeStream = ReceiverStream>; + + async fn subscribe( + &self, + request: Request, + ) -> Result, Status> { + let req = request.into_inner(); + if req.channels.is_empty() && req.patterns.is_empty() { + return Err(Status::invalid_argument( + "at least one channel or pattern required", + )); + } + + let (tx, rx) = tokio::sync::mpsc::channel(256); + let pubsub = Arc::clone(&self.pubsub); + + // collect all broadcast receivers + let mut channel_rxs: Vec<(String, tokio::sync::broadcast::Receiver)> = Vec::new(); + let mut pattern_rxs: Vec<(String, tokio::sync::broadcast::Receiver)> = Vec::new(); + + for ch in &req.channels { + channel_rxs.push((ch.clone(), pubsub.subscribe(ch))); + } + for pat in &req.patterns { + pattern_rxs.push((pat.clone(), pubsub.psubscribe(pat))); + } + + tokio::spawn(async move { + loop { + // build a future that races all receivers + let event = tokio::select! { + biased; + result = recv_any_channel(&mut channel_rxs) => result, + result = recv_any_pattern(&mut pattern_rxs) => result, + }; + + match event { + Some(evt) => { + if tx.send(Ok(evt)).await.is_err() { + break; // client disconnected + } + } + None => { + // all receivers closed + break; + } + } + } + + // cleanup subscriptions + for (ch, _) in &channel_rxs { + pubsub.unsubscribe(ch); + } + for (pat, _) in &pattern_rxs { + pubsub.punsubscribe(pat); + } + }); + + Ok(Response::new(ReceiverStream::new(rx))) + } + + async fn pub_sub_channels( + &self, + request: Request, + ) -> Result, Status> { + let pattern = request.into_inner().pattern; + let names = self.pubsub.channel_names(pattern.as_deref()); + Ok(Response::new(KeysResponse { keys: names })) + } + + async fn pub_sub_num_sub( + &self, + request: Request, + ) -> Result, Status> { + let channels = request.into_inner().channels; + let pairs = self.pubsub.numsub(&channels); + Ok(Response::new(PubSubNumSubResponse { + counts: pairs + .into_iter() + .map(|(channel, count)| ChannelCount { + channel, + count: count as i64, + }) + .collect(), + })) + } + + async fn pub_sub_num_pat( + &self, + _request: Request, + ) -> Result, Status> { + Ok(Response::new(IntResponse { + value: self.pubsub.active_patterns() as i64, + })) + } + // ----------------------------------------------------------------------- // pipeline (bidirectional streaming) // ----------------------------------------------------------------------- @@ -1695,11 +1938,12 @@ impl EmberCache for EmberService { let engine = self.engine.clone(); let ctx = Arc::clone(&self.ctx); let slow_log = Arc::clone(&self.slow_log); + let pubsub = Arc::clone(&self.pubsub); let (tx, rx) = tokio::sync::mpsc::channel(256); tokio::spawn(async move { - let svc = EmberService::new(engine, ctx, slow_log); + let svc = EmberService::new(engine, ctx, slow_log, pubsub); while let Ok(Some(req)) = stream.message().await { let id = req.id; let result = handle_pipeline_command(&svc, req).await; @@ -1822,12 +2066,92 @@ async fn handle_pipeline_command( // server commands Ping => ping => Ping, + Echo => echo => Echo, + Decr => decr => IntVal, + Unlink => unlink => Del, Flushdb => flush_db => Status, Dbsize => db_size => IntVal, + Bgsave => bg_save => Status, + Bgrewriteaof => bg_rewrite_aof => Status, Mget => m_get => Mget, Mset => m_set => Mset, Keys => keys => Keys, Rename => rename => Status, Scan => scan => Scan, + + // slowlog + SlowlogGet => slow_log_get => SlowlogGet, + SlowlogLen => slow_log_len => IntVal, + SlowlogReset => slow_log_reset => Status, + + // pub/sub (unary only — Subscribe is streaming) + Publish => publish => IntVal, + PubsubChannels => pub_sub_channels => Keys, + PubsubNumsub => pub_sub_num_sub => PubsubNumsub, + PubsubNumpat => pub_sub_num_pat => IntVal, }) } + +/// Waits for the next message on any channel subscription receiver. +/// Returns None when all receivers are closed. +async fn recv_any_channel( + rxs: &mut [(String, tokio::sync::broadcast::Receiver)], +) -> Option { + if rxs.is_empty() { + // no channel subscriptions — park forever so pattern branch can drive + std::future::pending::<()>().await; + return None; + } + + loop { + // poll each receiver in round-robin (tokio::select! on a slice + // requires a loop because we can't use select! with dynamic count). + for (_, rx) in rxs.iter_mut() { + match rx.try_recv() { + Ok(msg) => { + return Some(SubscribeEvent { + kind: "message".to_string(), + channel: msg.channel, + data: msg.data.to_vec(), + pattern: None, + }); + } + Err(tokio::sync::broadcast::error::TryRecvError::Empty) => continue, + Err(tokio::sync::broadcast::error::TryRecvError::Lagged(_)) => continue, + Err(tokio::sync::broadcast::error::TryRecvError::Closed) => return None, + } + } + // yield to avoid busy-spinning + tokio::time::sleep(Duration::from_millis(1)).await; + } +} + +/// Waits for the next message on any pattern subscription receiver. +/// Returns None when all receivers are closed. +async fn recv_any_pattern( + rxs: &mut [(String, tokio::sync::broadcast::Receiver)], +) -> Option { + if rxs.is_empty() { + std::future::pending::<()>().await; + return None; + } + + loop { + for (pat, rx) in rxs.iter_mut() { + match rx.try_recv() { + Ok(msg) => { + return Some(SubscribeEvent { + kind: "pmessage".to_string(), + channel: msg.channel, + data: msg.data.to_vec(), + pattern: Some(pat.clone()), + }); + } + Err(tokio::sync::broadcast::error::TryRecvError::Empty) => continue, + Err(tokio::sync::broadcast::error::TryRecvError::Lagged(_)) => continue, + Err(tokio::sync::broadcast::error::TryRecvError::Closed) => return None, + } + } + tokio::time::sleep(Duration::from_millis(1)).await; + } +} diff --git a/crates/ember-server/src/server.rs b/crates/ember-server/src/server.rs index 76424624..b3a08e38 100644 --- a/crates/ember-server/src/server.rs +++ b/crates/ember-server/src/server.rs @@ -132,7 +132,7 @@ pub async fn run( #[cfg(feature = "grpc")] let _grpc_handle = if let Some(grpc_addr) = grpc_addr { let svc = - crate::grpc::EmberService::new(engine.clone(), Arc::clone(&ctx), Arc::clone(&slow_log)); + crate::grpc::EmberService::new(engine.clone(), Arc::clone(&ctx), Arc::clone(&slow_log), Arc::clone(&pubsub)); info!("gRPC listening on {grpc_addr}"); let server = tonic::transport::Server::builder() .add_service(svc.into_service()) @@ -397,7 +397,7 @@ pub async fn run_concurrent( #[cfg(feature = "grpc")] let _grpc_handle = if let Some(grpc_addr) = grpc_addr { let svc = - crate::grpc::EmberService::new(engine.clone(), Arc::clone(&ctx), Arc::clone(&slow_log)); + crate::grpc::EmberService::new(engine.clone(), Arc::clone(&ctx), Arc::clone(&slow_log), Arc::clone(&pubsub)); info!("gRPC listening on {grpc_addr}"); let server = tonic::transport::Server::builder() .add_service(svc.into_service()) diff --git a/proto/ember/v1/ember.proto b/proto/ember/v1/ember.proto index 640816dc..e63e7c84 100644 --- a/proto/ember/v1/ember.proto +++ b/proto/ember/v1/ember.proto @@ -87,9 +87,28 @@ service EmberCache { // --- server --- rpc Ping(PingRequest) returns (PingResponse); + rpc Echo(EchoRequest) returns (EchoResponse); + rpc Decr(DecrRequest) returns (IntResponse); + rpc Unlink(UnlinkRequest) returns (DelResponse); rpc FlushDb(FlushDbRequest) returns (StatusResponse); rpc DbSize(DbSizeRequest) returns (IntResponse); rpc Info(InfoRequest) returns (InfoResponse); + rpc BgSave(BgSaveRequest) returns (StatusResponse); + rpc BgRewriteAof(BgRewriteAofRequest) returns (StatusResponse); + + // --- slowlog --- + + rpc SlowLogGet(SlowLogGetRequest) returns (SlowLogGetResponse); + rpc SlowLogLen(SlowLogLenRequest) returns (IntResponse); + rpc SlowLogReset(SlowLogResetRequest) returns (StatusResponse); + + // --- pub/sub --- + + rpc Publish(PublishRequest) returns (IntResponse); + rpc Subscribe(SubscribeRequest) returns (stream SubscribeEvent); + rpc PubSubChannels(PubSubChannelsRequest) returns (KeysResponse); + rpc PubSubNumSub(PubSubNumSubRequest) returns (PubSubNumSubResponse); + rpc PubSubNumPat(PubSubNumPatRequest) returns (IntResponse); // --- streaming --- // bidirectional streaming for batch operations, matching RESP3 pipelining. @@ -560,6 +579,91 @@ message InfoResponse { string info = 1; } +message EchoRequest { + string message = 1; +} + +message EchoResponse { + string message = 1; +} + +message DecrRequest { + string key = 1; +} + +message UnlinkRequest { + repeated string keys = 1; +} + +message BgSaveRequest {} + +message BgRewriteAofRequest {} + +// --------------------------------------------------------------------------- +// slowlog +// --------------------------------------------------------------------------- + +message SlowLogGetRequest { + optional uint32 count = 1; +} + +message SlowLogGetResponse { + repeated SlowLogEntry entries = 1; +} + +message SlowLogEntry { + uint64 id = 1; + uint64 timestamp_unix = 2; + uint64 duration_micros = 3; + string command = 4; +} + +message SlowLogLenRequest {} + +message SlowLogResetRequest {} + +// --------------------------------------------------------------------------- +// pub/sub +// --------------------------------------------------------------------------- + +message PublishRequest { + string channel = 1; + bytes message = 2; +} + +message SubscribeRequest { + repeated string channels = 1; + repeated string patterns = 2; +} + +message SubscribeEvent { + // "message" for exact channel match, "pmessage" for pattern match. + string kind = 1; + string channel = 2; + bytes data = 3; + // the pattern that matched (only set for pmessage). + optional string pattern = 4; +} + +message PubSubChannelsRequest { + optional string pattern = 1; +} + +message PubSubNumSubRequest { + repeated string channels = 1; +} + +message PubSubNumSubResponse { + repeated ChannelCount counts = 1; +} + +message PubSubNumPatRequest {} + +message ChannelCount { + string channel = 1; + int64 count = 2; +} + // --------------------------------------------------------------------------- // pipeline (bidirectional streaming) // --------------------------------------------------------------------------- @@ -625,6 +729,18 @@ message PipelineRequest { KeysRequest keys = 57; RenameRequest rename = 58; ScanRequest scan = 59; + EchoRequest echo = 60; + DecrRequest decr = 61; + UnlinkRequest unlink = 62; + BgSaveRequest bgsave = 63; + BgRewriteAofRequest bgrewriteaof = 64; + SlowLogGetRequest slowlog_get = 65; + SlowLogLenRequest slowlog_len = 66; + SlowLogResetRequest slowlog_reset = 67; + PublishRequest publish = 68; + PubSubChannelsRequest pubsub_channels = 69; + PubSubNumSubRequest pubsub_numsub = 70; + PubSubNumPatRequest pubsub_numpat = 71; } } @@ -656,6 +772,9 @@ message PipelineResponse { PingResponse ping = 24; ErrorResponse error = 25; InfoResponse info = 26; + EchoResponse echo = 27; + SlowLogGetResponse slowlog_get = 28; + PubSubNumSubResponse pubsub_numsub = 29; } }