use std::pin::Pin;
use std::sync::Arc;
use futures_core::Stream;
use futures_util::StreamExt;
use prost::Message as _;
use tokio::io::{AsyncRead, AsyncWrite};
use tokio::net::TcpListener;
use tokio::sync::mpsc;
use tokio_rustls::TlsAcceptor;
use dynomite::cluster::admin_rpc::{
ClusterAdmin, ClusterChange, ClusterChangeKind, ClusterError, NoopClusterAdmin, PeerSnapshot,
};
use dynomite::embed::hooks::{DatastoreByteStream, DatastoreError};
use dynomite::embed::Datastore;
use dynomite::msg::{Msg, MsgType};
use crate::aae::status::{AaeStatusProvider, AaeStatusSnapshot, NoopAaeStatusProvider};
use crate::error::RiakError;
use crate::mapreduce::{MrError, PhaseBatch};
use crate::proto::http::object::{HttpIndex, HttpLink, HttpObject};
use crate::proto::pb::framer::{read_frame, write_frame, Frame};
use crate::proto::pb::mapreduce::{RpbMapRedReq, RpbMapRedResp};
use crate::proto::pb::messages::{
DynRpbAaePeerStatus, DynRpbAaeStatusReq, DynRpbAaeStatusResp, DynRpbClusterCommitReq,
DynRpbClusterCommitResp, DynRpbClusterJoinReq, DynRpbClusterJoinResp, DynRpbClusterLeaveReq,
DynRpbClusterLeaveResp, DynRpbClusterPlanReq, DynRpbClusterPlanResp, DynRpbListPeersReq,
DynRpbListPeersResp, DynRpbPeerInfo, DynRpbStagedChange, MessageCode, RpbBucketProps,
RpbContent, RpbDelReq, RpbErrorResp, RpbGetBucketReq, RpbGetBucketResp, RpbGetReq, RpbGetResp,
RpbGetServerInfoResp, RpbIndexReq, RpbIndexResp, RpbLink, RpbListBucketsReq,
RpbListBucketsResp, RpbListKeysReq, RpbListKeysResp, RpbPair, RpbPingReq, RpbPingResp,
RpbPutReq, RpbPutResp, RpbServerInfoReq, RpbSetBucketReq, RpbSetBucketResp,
DYN_STAGED_CHANGE_ADD, DYN_STAGED_CHANGE_REMOVE, INDEX_QUERY_TYPE_EQ, INDEX_QUERY_TYPE_RANGE,
};
use crate::router::{PeerOp, RoutingHooks};
pub(crate) const LIST_CHUNK_SIZE: usize = 256;
pub(crate) type FrameStream = Pin<Box<dyn Stream<Item = Result<Frame, RiakError>> + Send>>;
pub async fn serve_pbc(
listener: TcpListener,
datastore: Arc<dyn Datastore>,
) -> Result<(), RiakError> {
let admin: Arc<dyn ClusterAdmin> = Arc::new(NoopClusterAdmin);
serve_pbc_inner(listener, datastore, admin, None).await
}
pub async fn serve_pbc_with_admin(
listener: TcpListener,
datastore: Arc<dyn Datastore>,
admin: Arc<dyn ClusterAdmin>,
) -> Result<(), RiakError> {
serve_pbc_inner(listener, datastore, admin, None).await
}
pub async fn serve_pbc_tls(
listener: TcpListener,
datastore: Arc<dyn Datastore>,
acceptor: TlsAcceptor,
) -> Result<(), RiakError> {
let admin: Arc<dyn ClusterAdmin> = Arc::new(NoopClusterAdmin);
serve_pbc_inner(listener, datastore, admin, Some(acceptor)).await
}
pub async fn serve_pbc_tls_with_admin(
listener: TcpListener,
datastore: Arc<dyn Datastore>,
admin: Arc<dyn ClusterAdmin>,
acceptor: TlsAcceptor,
) -> Result<(), RiakError> {
serve_pbc_inner(listener, datastore, admin, Some(acceptor)).await
}
#[cfg(feature = "quic")]
pub async fn serve_pbc_quic(
listener: dynomite::net::quic::QuicListener,
datastore: Arc<dyn Datastore>,
) -> Result<(), RiakError> {
let admin: Arc<dyn ClusterAdmin> = Arc::new(NoopClusterAdmin);
serve_pbc_quic_inner(listener, datastore, admin).await
}
#[cfg(feature = "quic")]
pub async fn serve_pbc_quic_with_admin(
listener: dynomite::net::quic::QuicListener,
datastore: Arc<dyn Datastore>,
admin: Arc<dyn ClusterAdmin>,
) -> Result<(), RiakError> {
serve_pbc_quic_inner(listener, datastore, admin).await
}
#[cfg(feature = "quic")]
async fn serve_pbc_quic_inner(
listener: dynomite::net::quic::QuicListener,
datastore: Arc<dyn Datastore>,
admin: Arc<dyn ClusterAdmin>,
) -> Result<(), RiakError> {
let aae_status: Arc<dyn AaeStatusProvider> = Arc::new(NoopAaeStatusProvider);
loop {
let transport = listener.accept().await?;
let peer = transport.peer_addr_socket();
let datastore = Arc::clone(&datastore);
let admin = Arc::clone(&admin);
let aae = Arc::clone(&aae_status);
tokio::spawn(async move {
if let Err(e) = handle_conn_full(transport, datastore, admin, None, aae).await {
tracing::warn!(%peer, error = %e, "riak pbc quic connection ended with error");
}
});
}
}
pub async fn serve_pbc_with_routing(
listener: TcpListener,
datastore: Arc<dyn Datastore>,
admin: Arc<dyn ClusterAdmin>,
hooks: RoutingHooks,
) -> Result<(), RiakError> {
serve_pbc_inner_with_hooks(listener, datastore, admin, None, Some(hooks)).await
}
pub async fn serve_pbc_with_aae_status(
listener: TcpListener,
datastore: Arc<dyn Datastore>,
admin: Arc<dyn ClusterAdmin>,
aae_status: Arc<dyn AaeStatusProvider>,
) -> Result<(), RiakError> {
serve_pbc_full(listener, datastore, admin, None, None, Some(aae_status)).await
}
async fn serve_pbc_inner(
listener: TcpListener,
datastore: Arc<dyn Datastore>,
admin: Arc<dyn ClusterAdmin>,
acceptor: Option<TlsAcceptor>,
) -> Result<(), RiakError> {
serve_pbc_inner_with_hooks(listener, datastore, admin, acceptor, None).await
}
async fn serve_pbc_inner_with_hooks(
listener: TcpListener,
datastore: Arc<dyn Datastore>,
admin: Arc<dyn ClusterAdmin>,
acceptor: Option<TlsAcceptor>,
hooks: Option<RoutingHooks>,
) -> Result<(), RiakError> {
serve_pbc_full(listener, datastore, admin, acceptor, hooks, None).await
}
async fn serve_pbc_full(
listener: TcpListener,
datastore: Arc<dyn Datastore>,
admin: Arc<dyn ClusterAdmin>,
acceptor: Option<TlsAcceptor>,
hooks: Option<RoutingHooks>,
aae_status: Option<Arc<dyn AaeStatusProvider>>,
) -> Result<(), RiakError> {
let aae_status: Arc<dyn AaeStatusProvider> =
aae_status.unwrap_or_else(|| Arc::new(NoopAaeStatusProvider));
loop {
let (sock, peer) = listener.accept().await?;
let datastore = Arc::clone(&datastore);
let admin = Arc::clone(&admin);
let aae = Arc::clone(&aae_status);
let hooks = hooks.clone();
match acceptor.as_ref() {
Some(acc) => {
let acc = acc.clone();
tokio::spawn(async move {
match acc.accept(sock).await {
Ok(tls) => {
if let Err(e) =
handle_conn_full(tls, datastore, admin, hooks, aae).await
{
tracing::warn!(
%peer,
error = %e,
"riak pbc tls connection ended with error"
);
}
}
Err(e) => tracing::warn!(
%peer,
error = %e,
"riak pbc tls handshake failed"
),
}
});
}
None => {
tokio::spawn(async move {
if let Err(e) = handle_conn_full(sock, datastore, admin, hooks, aae).await {
tracing::warn!(%peer, error = %e, "riak pbc connection ended with error");
}
});
}
}
}
}
pub async fn handle_conn<S>(stream: S, datastore: Arc<dyn Datastore>) -> Result<(), RiakError>
where
S: AsyncRead + AsyncWrite + Unpin,
{
let admin: Arc<dyn ClusterAdmin> = Arc::new(NoopClusterAdmin);
handle_conn_with_admin(stream, datastore, admin).await
}
pub async fn handle_conn_with_admin<S>(
stream: S,
datastore: Arc<dyn Datastore>,
admin: Arc<dyn ClusterAdmin>,
) -> Result<(), RiakError>
where
S: AsyncRead + AsyncWrite + Unpin,
{
handle_conn_with_hooks(stream, datastore, admin, None).await
}
pub async fn handle_conn_with_hooks<S>(
stream: S,
datastore: Arc<dyn Datastore>,
admin: Arc<dyn ClusterAdmin>,
hooks: Option<RoutingHooks>,
) -> Result<(), RiakError>
where
S: AsyncRead + AsyncWrite + Unpin,
{
let aae_status: Arc<dyn AaeStatusProvider> = Arc::new(NoopAaeStatusProvider);
handle_conn_full(stream, datastore, admin, hooks, aae_status).await
}
pub async fn handle_conn_with_aae_status<S>(
stream: S,
datastore: Arc<dyn Datastore>,
admin: Arc<dyn ClusterAdmin>,
hooks: Option<RoutingHooks>,
aae_status: Arc<dyn AaeStatusProvider>,
) -> Result<(), RiakError>
where
S: AsyncRead + AsyncWrite + Unpin,
{
handle_conn_full(stream, datastore, admin, hooks, aae_status).await
}
async fn handle_conn_full<S>(
stream: S,
datastore: Arc<dyn Datastore>,
admin: Arc<dyn ClusterAdmin>,
hooks: Option<RoutingHooks>,
aae_status: Arc<dyn AaeStatusProvider>,
) -> Result<(), RiakError>
where
S: AsyncRead + AsyncWrite + Unpin,
{
let (mut reader, mut writer) = tokio::io::split(stream);
loop {
let frame = match read_frame(&mut reader).await {
Ok(f) => f,
Err(RiakError::UnexpectedEof { .. }) => {
return Ok(());
}
Err(other) => return Err(other),
};
let mut response = process_frame(
&frame,
datastore.as_ref(),
admin.as_ref(),
hooks.as_ref(),
aae_status.as_ref(),
)
.await?;
while let Some(item) = response.next().await {
let f = item?;
write_frame(&mut writer, &f).await?;
}
}
}
async fn process_frame(
frame: &Frame,
datastore: &dyn Datastore,
admin: &dyn ClusterAdmin,
hooks: Option<&RoutingHooks>,
aae_status: &dyn AaeStatusProvider,
) -> Result<FrameStream, RiakError> {
let code = MessageCode::from_u8(frame.code).map_err(RiakError::UnknownMessageCode)?;
let stream: FrameStream = match code {
MessageCode::PingReq => single_frame(handle_ping(&frame.body)?),
MessageCode::ServerInfoReq => single_frame(handle_server_info(&frame.body)?),
MessageCode::GetReq => single_frame(handle_get(&frame.body, datastore, hooks).await?),
MessageCode::PutReq => single_frame(handle_put(&frame.body, datastore, hooks).await?),
MessageCode::DelReq => single_frame(handle_del(&frame.body, datastore, hooks).await?),
MessageCode::GetBucketReq => single_frame(handle_get_bucket(&frame.body, hooks)?),
MessageCode::SetBucketReq => single_frame(handle_set_bucket(&frame.body, hooks)?),
MessageCode::ListBucketsReq => handle_list_buckets(&frame.body, datastore)?,
MessageCode::ListKeysReq => handle_list_keys(&frame.body, datastore)?,
MessageCode::IndexReq => handle_index(&frame.body, datastore).await?,
MessageCode::MapRedReq => handle_mapreduce(&frame.body),
MessageCode::DynListPeersReq => single_frame(handle_list_peers(&frame.body, admin)?),
MessageCode::DynClusterJoinReq => single_frame(handle_cluster_join(&frame.body, admin)?),
MessageCode::DynClusterLeaveReq => single_frame(handle_cluster_leave(&frame.body, admin)?),
MessageCode::DynClusterPlanReq => single_frame(handle_cluster_plan(&frame.body, admin)?),
MessageCode::DynClusterCommitReq => {
single_frame(handle_cluster_commit(&frame.body, admin)?)
}
MessageCode::DynAaeStatusReq => single_frame(handle_aae_status(&frame.body, aae_status)?),
MessageCode::DtUpdateReq => {
single_frame(handle_dt_update(&frame.body, datastore, hooks).await?)
}
MessageCode::DtFetchReq => {
single_frame(handle_dt_fetch(&frame.body, datastore, hooks).await?)
}
MessageCode::ErrorResp
| MessageCode::PingResp
| MessageCode::GetServerInfoResp
| MessageCode::GetResp
| MessageCode::PutResp
| MessageCode::DelResp
| MessageCode::ListBucketsResp
| MessageCode::ListKeysResp
| MessageCode::GetBucketResp
| MessageCode::SetBucketResp
| MessageCode::DtUpdateResp
| MessageCode::DtFetchResp
| MessageCode::IndexResp
| MessageCode::MapRedResp
| MessageCode::DynListPeersResp
| MessageCode::DynClusterJoinResp
| MessageCode::DynClusterLeaveResp
| MessageCode::DynClusterPlanResp
| MessageCode::DynClusterCommitResp
| MessageCode::DynAaeStatusResp => {
let body = RpbErrorResp {
errmsg: format!("unsupported inbound message code: {}", frame.code).into_bytes(),
errcode: 0,
}
.encode_to_vec();
single_frame(Frame::new(MessageCode::ErrorResp.as_u8(), body))
}
};
Ok(stream)
}
fn single_frame(f: Frame) -> FrameStream {
Box::pin(futures_util::stream::once(async move { Ok(f) }))
}
fn handle_ping(body: &[u8]) -> Result<Frame, RiakError> {
let _ = RpbPingReq::decode(body)?;
let resp = RpbPingResp::default();
Ok(Frame::new(
MessageCode::PingResp.as_u8(),
resp.encode_to_vec(),
))
}
fn handle_server_info(body: &[u8]) -> Result<Frame, RiakError> {
let _ = RpbServerInfoReq::decode(body)?;
let resp = RpbGetServerInfoResp {
node: Some(b"dyniak".to_vec()),
server_version: Some(format!("dyniak {}", env!("CARGO_PKG_VERSION")).into_bytes()),
};
Ok(Frame::new(
MessageCode::GetServerInfoResp.as_u8(),
resp.encode_to_vec(),
))
}
fn pbc_content_from_storage(stored: &[u8]) -> RpbContent {
match HttpObject::from_storage_bytes(stored) {
Ok(obj) => RpbContent {
value: obj.value,
content_type: obj.content_type.map(String::into_bytes),
links: obj.links.iter().map(http_link_to_rpb).collect(),
indexes: obj
.indexes
.iter()
.map(|i| RpbPair {
key: i.name.clone().into_bytes(),
value: Some(i.value.clone().into_bytes()),
})
.collect(),
..RpbContent::default()
},
Err(_) => RpbContent {
value: stored.to_vec(),
..RpbContent::default()
},
}
}
fn http_link_to_rpb(link: &HttpLink) -> RpbLink {
let opt = |s: &str| {
if s.is_empty() {
None
} else {
Some(s.as_bytes().to_vec())
}
};
RpbLink {
bucket: opt(&link.bucket),
key: opt(&link.key),
tag: opt(&link.tag),
}
}
fn rpb_link_to_http(link: &RpbLink) -> HttpLink {
let text = |b: &Option<Vec<u8>>| {
b.as_deref()
.map(|v| String::from_utf8_lossy(v).into_owned())
.unwrap_or_default()
};
HttpLink {
bucket: text(&link.bucket),
key: text(&link.key),
tag: text(&link.tag),
}
}
async fn handle_get(
body: &[u8],
datastore: &dyn Datastore,
hooks: Option<&RoutingHooks>,
) -> Result<Frame, RiakError> {
let req = RpbGetReq::decode(body)?;
if let Some(hooks) = hooks {
let bucket_type = req.r#type.as_deref().unwrap_or(b"");
let decision = match hooks.router.try_route(bucket_type, &req.bucket, &req.key) {
Ok(d) => d,
Err(e) => return Ok(error_frame(format!("riak get: {e}"))),
};
for replica in decision.replica_list() {
hooks
.outbound
.dispatch(
replica.peer_idx,
PeerOp::Get {
bucket_type: decision.bucket_type.clone(),
bucket: req.bucket.clone(),
key: req.key.clone(),
},
)
.await;
}
}
let routing = Msg::new(0, MsgType::Unknown, true);
datastore.dispatch(routing).await?;
let resp = match datastore.riak_get(&req.bucket, &req.key).await {
Ok(Some(v)) => RpbGetResp {
content: vec![pbc_content_from_storage(&v)],
..RpbGetResp::default()
},
Ok(None) | Err(DatastoreError::Unsupported(_)) => RpbGetResp::default(),
Err(e) => {
return Ok(error_frame(format!("riak get: {e}")));
}
};
Ok(Frame::new(
MessageCode::GetResp.as_u8(),
resp.encode_to_vec(),
))
}
async fn handle_put(
body: &[u8],
datastore: &dyn Datastore,
hooks: Option<&RoutingHooks>,
) -> Result<Frame, RiakError> {
let req = RpbPutReq::decode(body)?;
let routing = Msg::new(0, MsgType::Unknown, true);
datastore.dispatch(routing).await?;
let key = match req.key.as_ref() {
Some(k) if !k.is_empty() => k.clone(),
_ => {
return Ok(error_frame(
"riak put: server-assigned keys not implemented; supply 'key'".into(),
));
}
};
let content = req.content.clone().unwrap_or_default();
if let Some(hooks) = hooks {
let bucket_type = req.r#type.as_deref().unwrap_or(b"");
let decision = match hooks.router.try_route(bucket_type, &req.bucket, &key) {
Ok(d) => d,
Err(e) => return Ok(error_frame(format!("riak put: {e}"))),
};
for replica in decision.replica_list() {
hooks
.outbound
.dispatch(
replica.peer_idx,
PeerOp::Put {
bucket_type: decision.bucket_type.clone(),
bucket: req.bucket.clone(),
key: key.clone(),
value: content.value.clone(),
},
)
.await;
}
}
let indexes: Vec<(Vec<u8>, Vec<u8>)> = content
.indexes
.iter()
.filter_map(|p| p.value.as_ref().map(|v| (p.key.clone(), v.clone())))
.collect();
let envelope = HttpObject {
value: content.value.clone(),
content_type: content
.content_type
.as_deref()
.map(|c| String::from_utf8_lossy(c).into_owned()),
indexes: indexes
.iter()
.map(|(n, v)| HttpIndex {
name: String::from_utf8_lossy(n).into_owned(),
value: String::from_utf8_lossy(v).into_owned(),
})
.collect(),
links: content.links.iter().map(rpb_link_to_http).collect(),
};
let storage = envelope.to_storage_bytes();
match datastore
.riak_put(&req.bucket, &key, &storage, &indexes)
.await
{
Ok(()) | Err(DatastoreError::Unsupported(_)) => Ok(Frame::new(
MessageCode::PutResp.as_u8(),
RpbPutResp::default().encode_to_vec(),
)),
Err(e) => Ok(error_frame(format!("riak put: {e}"))),
}
}
async fn handle_del(
body: &[u8],
datastore: &dyn Datastore,
hooks: Option<&RoutingHooks>,
) -> Result<Frame, RiakError> {
let req = RpbDelReq::decode(body)?;
if let Some(hooks) = hooks {
let bucket_type = req.r#type.as_deref().unwrap_or(b"");
let decision = match hooks.router.try_route(bucket_type, &req.bucket, &req.key) {
Ok(d) => d,
Err(e) => return Ok(error_frame(format!("riak del: {e}"))),
};
for replica in decision.replica_list() {
hooks
.outbound
.dispatch(
replica.peer_idx,
PeerOp::Del {
bucket_type: decision.bucket_type.clone(),
bucket: req.bucket.clone(),
key: req.key.clone(),
},
)
.await;
}
}
let routing = Msg::new(0, MsgType::Unknown, true);
datastore.dispatch(routing).await?;
match datastore.riak_delete(&req.bucket, &req.key).await {
Ok(_) | Err(DatastoreError::Unsupported(_)) => {
Ok(Frame::new(MessageCode::DelResp.as_u8(), Vec::new()))
}
Err(e) => Ok(error_frame(format!("riak del: {e}"))),
}
}
async fn handle_dt_update(
body: &[u8],
datastore: &dyn Datastore,
hooks: Option<&RoutingHooks>,
) -> Result<Frame, RiakError> {
use crate::crdt_store::{CrdtStore, CrdtValue};
use crate::proto::pb::{DtUpdateReq, DtUpdateResp};
let req = DtUpdateReq::decode(body)?;
let key = match req.key.as_ref() {
Some(k) if !k.is_empty() => k.clone(),
_ => {
return Ok(error_frame(
"riak dt_update: server-assigned keys not implemented; supply 'key'".into(),
))
}
};
let actor = hooks.map_or_else(
|| crate::datatypes::ActorId::new("local", "local"),
|h| h.local_actor.clone(),
);
let Some(op) = req.op.as_ref().and_then(|o| dt_op_to_crdt(o, &actor)) else {
return Ok(error_frame(
"riak dt_update: unsupported or empty op (counter/set only)".into(),
));
};
let (value, state_bytes) =
match CrdtStore::apply_borrowed_with_state(datastore, &req.bucket, &key, &op).await {
Ok(r) => r,
Err(e) => return Ok(error_frame(format!("riak dt_update: {e}"))),
};
if let Some(hooks) = hooks {
let bucket_type = req.r#type.as_slice();
if let Ok(decision) = hooks.router.try_route(bucket_type, &req.bucket, &key) {
let wire = crate::crdt_store::to_state_wire(&state_bytes);
for replica in decision.replica_list() {
if replica.peer_idx == hooks.local_peer_idx {
continue;
}
hooks
.outbound
.dispatch(
replica.peer_idx,
crate::router::PeerOp::DtUpdate {
bucket_type: decision.bucket_type.clone(),
bucket: req.bucket.clone(),
key: key.clone(),
op: wire.clone(),
},
)
.await;
}
}
}
let mut resp = DtUpdateResp::default();
match value {
CrdtValue::Counter(n) => resp.counter_value = Some(n),
CrdtValue::Set(elems) => resp.set_value = elems,
CrdtValue::Missing => {}
}
Ok(Frame::new(
MessageCode::DtUpdateResp.as_u8(),
resp.encode_to_vec(),
))
}
#[doc(hidden)]
pub async fn handle_dt_fetch_for_test(
body: &[u8],
datastore: &dyn Datastore,
hooks: Option<&RoutingHooks>,
) -> Result<crate::proto::pb::framer::Frame, RiakError> {
handle_dt_fetch(body, datastore, hooks).await
}
async fn handle_dt_fetch(
body: &[u8],
datastore: &dyn Datastore,
hooks: Option<&RoutingHooks>,
) -> Result<Frame, RiakError> {
use crate::crdt_store::{CrdtStore, CrdtValue};
use crate::datatypes::{TAG_COUNTER, TAG_SET};
use crate::proto::pb::{DtFetchReq, DtFetchResp, DtValue, DATA_TYPE_COUNTER, DATA_TYPE_SET};
let req = DtFetchReq::decode(body)?;
let (tag, dtype) = if req.r#type == b"sets" {
(TAG_SET, DATA_TYPE_SET)
} else {
(TAG_COUNTER, DATA_TYPE_COUNTER)
};
let mut merged_state: Vec<u8> = match datastore.riak_get(&req.bucket, &req.key).await {
Ok(Some(s)) => s,
_ => Vec::new(),
};
if let Some(hooks) = hooks {
if let Ok(decision) = hooks.router.try_route(&req.r#type, &req.bucket, &req.key) {
for replica in decision.replica_list() {
if replica.peer_idx == hooks.local_peer_idx {
continue;
}
let reply = hooks
.outbound
.request(
replica.peer_idx,
crate::router::PeerOp::DtFetch {
bucket_type: decision.bucket_type.clone(),
bucket: req.bucket.clone(),
key: req.key.clone(),
tag,
},
)
.await;
if let Some(state) = reply {
if !state.is_empty() {
merged_state =
match crate::crdt_store::merge_two_states(&merged_state, &state) {
Ok(m) => m,
Err(_) => merged_state,
};
}
}
}
}
}
let value = if merged_state.is_empty() {
CrdtValue::Missing
} else {
crate::crdt_store::project_state(&merged_state, tag).unwrap_or(CrdtValue::Missing)
};
if !matches!(value, CrdtValue::Missing) {
let _ =
CrdtStore::merge_state_borrowed(datastore, &req.bucket, &req.key, &merged_state).await;
}
let mut resp = DtFetchResp {
r#type: dtype,
..DtFetchResp::default()
};
match value {
CrdtValue::Counter(n) => {
resp.value = Some(DtValue {
counter_value: Some(n),
..DtValue::default()
});
}
CrdtValue::Set(elems) => {
resp.value = Some(DtValue {
set_value: elems,
..DtValue::default()
});
}
CrdtValue::Missing => {}
}
Ok(Frame::new(
MessageCode::DtFetchResp.as_u8(),
resp.encode_to_vec(),
))
}
fn dt_op_to_crdt(
op: &crate::proto::pb::DtOp,
actor: &crate::datatypes::ActorId,
) -> Option<crate::crdt_store::CrdtOp> {
use crate::crdt_store::CrdtOp;
if let Some(c) = op.counter_op.as_ref() {
return Some(CrdtOp::Counter {
actor: actor.clone(),
delta: c.increment.unwrap_or(0),
});
}
if let Some(s) = op.set_op.as_ref() {
return Some(CrdtOp::Set {
actor: actor.clone(),
adds: s.adds.clone(),
removes: s.removes.clone(),
});
}
None
}
async fn handle_index(body: &[u8], datastore: &dyn Datastore) -> Result<FrameStream, RiakError> {
let req = RpbIndexReq::decode(body)?;
let result = match req.qtype {
INDEX_QUERY_TYPE_EQ => {
let value = req.key.as_deref().unwrap_or(b"");
datastore
.riak_index_eq(&req.bucket, &req.index, value)
.await
}
INDEX_QUERY_TYPE_RANGE => {
let min = req.range_min.as_deref().unwrap_or(b"");
let max = req.range_max.as_deref().unwrap_or(b"");
datastore
.riak_index_range(&req.bucket, &req.index, min, max)
.await
}
other => {
return Ok(single_frame(error_frame(format!(
"riak index: unsupported qtype {other}; expected 0 (eq) or 1 (range)"
))));
}
};
match result {
Ok(mut keys) => {
if let Some(cap) = req.max_results {
let cap = cap as usize;
if keys.len() > cap {
keys.truncate(cap);
}
}
Ok(Box::pin(index_keys_to_frames(keys)))
}
Err(DatastoreError::Unsupported(_)) => Ok(single_frame(error_frame(
"secondary-index queries not implemented for this datastore".into(),
))),
Err(e) => Ok(single_frame(error_frame(format!("riak index: {e}")))),
}
}
enum IndexChunkState {
Streaming { keys: Vec<Vec<u8>> },
Terminate,
Done,
}
fn index_keys_to_frames(keys: Vec<Vec<u8>>) -> impl Stream<Item = Result<Frame, RiakError>> + Send {
futures_util::stream::unfold(IndexChunkState::Streaming { keys }, |state| async move {
match state {
IndexChunkState::Done => None,
IndexChunkState::Terminate => {
let resp = RpbIndexResp {
keys: Vec::new(),
results: Vec::new(),
continuation: None,
done: Some(true),
};
let frame = Frame::new(MessageCode::IndexResp.as_u8(), resp.encode_to_vec());
Some((Ok(frame), IndexChunkState::Done))
}
IndexChunkState::Streaming { mut keys } => {
if keys.is_empty() {
let resp = RpbIndexResp {
keys: Vec::new(),
results: Vec::new(),
continuation: None,
done: Some(true),
};
let frame = Frame::new(MessageCode::IndexResp.as_u8(), resp.encode_to_vec());
return Some((Ok(frame), IndexChunkState::Done));
}
let take = LIST_CHUNK_SIZE.min(keys.len());
let tail = keys.split_off(take);
let chunk = keys; let next = if tail.is_empty() {
IndexChunkState::Terminate
} else {
IndexChunkState::Streaming { keys: tail }
};
let resp = RpbIndexResp {
keys: chunk,
results: Vec::new(),
continuation: None,
done: Some(false),
};
let frame = Frame::new(MessageCode::IndexResp.as_u8(), resp.encode_to_vec());
Some((Ok(frame), next))
}
}
})
}
fn error_frame(message: String) -> Frame {
let resp = RpbErrorResp {
errmsg: message.into_bytes(),
errcode: 1,
};
Frame::new(MessageCode::ErrorResp.as_u8(), resp.encode_to_vec())
}
fn handle_get_bucket(body: &[u8], hooks: Option<&RoutingHooks>) -> Result<Frame, RiakError> {
let req = RpbGetBucketReq::decode(body)?;
let props = if let Some(hooks) = hooks {
let bucket_type = req.r#type.as_deref().unwrap_or(b"");
let resolved = hooks.router.registry().resolve(bucket_type, &req.bucket);
RpbBucketProps {
n_val: Some(u32::from(resolved.effective_n_val())),
allow_mult: Some(false),
last_write_wins: Some(false),
chash_keyfun: Some(resolved.effective_keyfun().to_wire()),
chash_keyfun_module: resolved
.effective_keyfun()
.custom_module()
.map(|s| s.as_bytes().to_vec()),
replication_strategy: Some(resolved.effective_strategy().to_wire()),
..RpbBucketProps::default()
}
} else {
RpbBucketProps {
n_val: Some(3),
allow_mult: Some(false),
last_write_wins: Some(false),
..RpbBucketProps::default()
}
};
let resp = RpbGetBucketResp { props: Some(props) };
Ok(Frame::new(
MessageCode::GetBucketResp.as_u8(),
resp.encode_to_vec(),
))
}
fn handle_set_bucket(body: &[u8], hooks: Option<&RoutingHooks>) -> Result<Frame, RiakError> {
let req = RpbSetBucketReq::decode(body)?;
if let Some(hooks) = hooks {
let bucket_type = req.r#type.as_deref().unwrap_or(b"");
if let Some(props) = req.props.as_ref() {
let mut bp = crate::bucket_props::BucketProps::default();
if let Some(w) = props.chash_keyfun {
if let Ok(kf) = crate::datatypes::keyfun::KeyFun::from_wire(w) {
if let crate::datatypes::keyfun::KeyFun::Custom(_) = kf {
let module_id = props
.chash_keyfun_module
.as_deref()
.map(|b| String::from_utf8_lossy(b).into_owned())
.unwrap_or_default();
if let Err(msg) = validate_custom_keyfun(hooks, &module_id) {
return Ok(error_frame(msg));
}
bp.keyfun =
Some(crate::datatypes::keyfun::KeyFun::Custom(module_id.clone()));
bp.custom_keyfun_module = Some(module_id);
} else {
bp.keyfun = Some(kf);
}
}
}
if let Some(w) = props.replication_strategy {
if let Ok(s) = crate::replication::ReplicationStrategy::from_wire(w) {
bp.strategy = Some(s);
}
}
if let Some(n) = props.n_val {
bp.n_val = Some(u8::try_from(n).unwrap_or(u8::MAX));
}
hooks.router.registry().set(bucket_type, &req.bucket, bp);
}
}
let resp = RpbSetBucketResp::default();
Ok(Frame::new(
MessageCode::SetBucketResp.as_u8(),
resp.encode_to_vec(),
))
}
#[cfg(feature = "wasm")]
fn validate_custom_keyfun(hooks: &RoutingHooks, module_id: &str) -> Result<(), String> {
if module_id.is_empty() {
return Err(
"set bucket: chash_keyfun CUSTOM requires a non-empty chash_keyfun_module".into(),
);
}
match hooks.router.keyfun_store() {
Some(store) if store.contains(module_id) => Ok(()),
Some(_) => Err(format!(
"set bucket: chash_keyfun CUSTOM module {module_id:?} is not registered"
)),
None => Err(
"set bucket: chash_keyfun CUSTOM selected but no keyfun WASM store is configured"
.into(),
),
}
}
#[cfg(not(feature = "wasm"))]
fn validate_custom_keyfun(_hooks: &RoutingHooks, module_id: &str) -> Result<(), String> {
if module_id.is_empty() {
return Err(
"set bucket: chash_keyfun CUSTOM requires a non-empty chash_keyfun_module".into(),
);
}
Err("set bucket: chash_keyfun CUSTOM requires the 'wasm' feature".into())
}
fn handle_mapreduce(body: &[u8]) -> FrameStream {
use crate::mapreduce::{builtins::default_registry, run_job_streaming, MapReduceJob};
let req = match RpbMapRedReq::decode(body) {
Ok(r) => r,
Err(e) => {
let resp = RpbErrorResp {
errmsg: format!("MapReduce request decode: {e}").into_bytes(),
errcode: 1,
};
return single_frame(Frame::new(
MessageCode::ErrorResp.as_u8(),
resp.encode_to_vec(),
));
}
};
if req.content_type != b"application/json" {
let resp = RpbErrorResp {
errmsg: format!(
"unsupported MapReduce content-type: {}",
String::from_utf8_lossy(&req.content_type)
)
.into_bytes(),
errcode: 1,
};
return single_frame(Frame::new(
MessageCode::ErrorResp.as_u8(),
resp.encode_to_vec(),
));
}
let job: MapReduceJob = match serde_json::from_slice(&req.request) {
Ok(j) => j,
Err(e) => {
let resp = RpbErrorResp {
errmsg: format!("MapReduce job decode: {e}").into_bytes(),
errcode: 1,
};
return single_frame(Frame::new(
MessageCode::ErrorResp.as_u8(),
resp.encode_to_vec(),
));
}
};
let registry = Arc::new(default_registry());
let rx = run_job_streaming(job, registry);
Box::pin(mapreduce_response_stream(rx))
}
enum MrStreamState {
Streaming(mpsc::Receiver<Result<PhaseBatch, MrError>>),
Done,
}
fn mapreduce_response_stream(
rx: mpsc::Receiver<Result<PhaseBatch, MrError>>,
) -> impl Stream<Item = Result<Frame, RiakError>> + Send {
futures_util::stream::unfold(MrStreamState::Streaming(rx), |state| async move {
match state {
MrStreamState::Done => None,
MrStreamState::Streaming(mut rx) => match rx.recv().await {
None => {
let resp = RpbMapRedResp {
phase: None,
response: None,
done: Some(true),
};
let frame = Frame::new(MessageCode::MapRedResp.as_u8(), resp.encode_to_vec());
Some((Ok(frame), MrStreamState::Done))
}
Some(Ok(batch)) => {
let payload = serde_json::json!([{
"phase": batch.phase,
"data": batch.data,
}]);
let body = serde_json::to_vec(&payload).unwrap_or_else(|_| b"[]".to_vec());
let resp = RpbMapRedResp {
phase: Some(batch.phase),
response: Some(body),
done: Some(false),
};
let frame = Frame::new(MessageCode::MapRedResp.as_u8(), resp.encode_to_vec());
Some((Ok(frame), MrStreamState::Streaming(rx)))
}
Some(Err(e)) => {
let resp = RpbErrorResp {
errmsg: format!("MapReduce execution: {e}").into_bytes(),
errcode: 1,
};
let frame = Frame::new(MessageCode::ErrorResp.as_u8(), resp.encode_to_vec());
Some((Ok(frame), MrStreamState::Done))
}
},
}
})
}
fn handle_list_buckets(body: &[u8], datastore: &dyn Datastore) -> Result<FrameStream, RiakError> {
let _req = RpbListBucketsReq::decode(body)?;
let stream = datastore.list_buckets_stream();
Ok(Box::pin(buckets_to_frames(stream)))
}
fn handle_list_keys(body: &[u8], datastore: &dyn Datastore) -> Result<FrameStream, RiakError> {
let req = RpbListKeysReq::decode(body)?;
let stream = datastore.list_keys_stream(&req.bucket);
Ok(Box::pin(keys_to_frames(stream)))
}
enum ListChunkState {
Streaming(DatastoreByteStream, Vec<Vec<u8>>),
Terminate,
Done,
}
fn buckets_to_frames(
s: DatastoreByteStream,
) -> impl Stream<Item = Result<Frame, RiakError>> + Send {
futures_util::stream::unfold(
ListChunkState::Streaming(s, Vec::with_capacity(LIST_CHUNK_SIZE)),
|state| async move {
match state {
ListChunkState::Done => None,
ListChunkState::Terminate => {
let resp = RpbListBucketsResp {
buckets: Vec::new(),
done: Some(true),
};
let frame =
Frame::new(MessageCode::ListBucketsResp.as_u8(), resp.encode_to_vec());
Some((Ok(frame), ListChunkState::Done))
}
ListChunkState::Streaming(mut stream, mut buffer) => loop {
if buffer.len() >= LIST_CHUNK_SIZE {
let resp = RpbListBucketsResp {
buckets: buffer,
done: Some(false),
};
let frame =
Frame::new(MessageCode::ListBucketsResp.as_u8(), resp.encode_to_vec());
return Some((
Ok(frame),
ListChunkState::Streaming(stream, Vec::with_capacity(LIST_CHUNK_SIZE)),
));
}
match stream.next().await {
Some(Ok(b)) => buffer.push(b.to_vec()),
Some(Err(e)) => {
let resp = RpbErrorResp {
errmsg: format!("list-buckets failed: {e}").into_bytes(),
errcode: 1,
};
let frame =
Frame::new(MessageCode::ErrorResp.as_u8(), resp.encode_to_vec());
return Some((Ok(frame), ListChunkState::Done));
}
None => {
if buffer.is_empty() {
let resp = RpbListBucketsResp {
buckets: Vec::new(),
done: Some(true),
};
let frame = Frame::new(
MessageCode::ListBucketsResp.as_u8(),
resp.encode_to_vec(),
);
return Some((Ok(frame), ListChunkState::Done));
}
let resp = RpbListBucketsResp {
buckets: buffer,
done: Some(false),
};
let frame = Frame::new(
MessageCode::ListBucketsResp.as_u8(),
resp.encode_to_vec(),
);
return Some((Ok(frame), ListChunkState::Terminate));
}
}
},
}
},
)
}
fn keys_to_frames(s: DatastoreByteStream) -> impl Stream<Item = Result<Frame, RiakError>> + Send {
futures_util::stream::unfold(
ListChunkState::Streaming(s, Vec::with_capacity(LIST_CHUNK_SIZE)),
|state| async move {
match state {
ListChunkState::Done => None,
ListChunkState::Terminate => {
let resp = RpbListKeysResp {
keys: Vec::new(),
done: Some(true),
};
let frame = Frame::new(MessageCode::ListKeysResp.as_u8(), resp.encode_to_vec());
Some((Ok(frame), ListChunkState::Done))
}
ListChunkState::Streaming(mut stream, mut buffer) => loop {
if buffer.len() >= LIST_CHUNK_SIZE {
let resp = RpbListKeysResp {
keys: buffer,
done: Some(false),
};
let frame =
Frame::new(MessageCode::ListKeysResp.as_u8(), resp.encode_to_vec());
return Some((
Ok(frame),
ListChunkState::Streaming(stream, Vec::with_capacity(LIST_CHUNK_SIZE)),
));
}
match stream.next().await {
Some(Ok(b)) => buffer.push(b.to_vec()),
Some(Err(e)) => {
let resp = RpbErrorResp {
errmsg: format!("list-keys failed: {e}").into_bytes(),
errcode: 1,
};
let frame =
Frame::new(MessageCode::ErrorResp.as_u8(), resp.encode_to_vec());
return Some((Ok(frame), ListChunkState::Done));
}
None => {
if buffer.is_empty() {
let resp = RpbListKeysResp {
keys: Vec::new(),
done: Some(true),
};
let frame = Frame::new(
MessageCode::ListKeysResp.as_u8(),
resp.encode_to_vec(),
);
return Some((Ok(frame), ListChunkState::Done));
}
let resp = RpbListKeysResp {
keys: buffer,
done: Some(false),
};
let frame =
Frame::new(MessageCode::ListKeysResp.as_u8(), resp.encode_to_vec());
return Some((Ok(frame), ListChunkState::Terminate));
}
}
},
}
},
)
}
fn handle_list_peers(body: &[u8], admin: &dyn ClusterAdmin) -> Result<Frame, RiakError> {
let _ = DynRpbListPeersReq::decode(body)?;
let snaps = admin.list_peers();
let resp = DynRpbListPeersResp {
peers: snaps.iter().map(snapshot_to_pb).collect(),
};
Ok(Frame::new(
MessageCode::DynListPeersResp.as_u8(),
resp.encode_to_vec(),
))
}
fn handle_cluster_join(body: &[u8], admin: &dyn ClusterAdmin) -> Result<Frame, RiakError> {
let req = DynRpbClusterJoinReq::decode(body)?;
let Ok(target_str) = std::str::from_utf8(&req.target) else {
return Ok(error_frame("cluster-join: target is not UTF-8".into()));
};
let target = match target_str.parse::<std::net::SocketAddr>() {
Ok(t) => t,
Err(e) => {
return Ok(error_frame(format!(
"cluster-join: invalid target '{target_str}': {e}"
)));
}
};
match admin.cluster_join(target) {
Ok(plan) => {
let resp = DynRpbClusterJoinResp {
change: Some(change_to_pb(&plan.change)),
};
Ok(Frame::new(
MessageCode::DynClusterJoinResp.as_u8(),
resp.encode_to_vec(),
))
}
Err(e) => Ok(error_frame(format_cluster_error("cluster-join", &e))),
}
}
fn handle_cluster_leave(body: &[u8], admin: &dyn ClusterAdmin) -> Result<Frame, RiakError> {
let req = DynRpbClusterLeaveReq::decode(body)?;
match admin.cluster_leave(req.peer_idx) {
Ok(plan) => {
let resp = DynRpbClusterLeaveResp {
change: Some(change_to_pb(&plan.change)),
};
Ok(Frame::new(
MessageCode::DynClusterLeaveResp.as_u8(),
resp.encode_to_vec(),
))
}
Err(e) => Ok(error_frame(format_cluster_error("cluster-leave", &e))),
}
}
fn handle_cluster_plan(body: &[u8], admin: &dyn ClusterAdmin) -> Result<Frame, RiakError> {
let _ = DynRpbClusterPlanReq::decode(body)?;
let pending = admin.cluster_plan_pending();
let resp = DynRpbClusterPlanResp {
changes: pending.iter().map(change_to_pb).collect(),
};
Ok(Frame::new(
MessageCode::DynClusterPlanResp.as_u8(),
resp.encode_to_vec(),
))
}
fn handle_cluster_commit(body: &[u8], admin: &dyn ClusterAdmin) -> Result<Frame, RiakError> {
let _ = DynRpbClusterCommitReq::decode(body)?;
let staged = admin.cluster_plan_pending();
let applied = u32::try_from(staged.len()).unwrap_or(u32::MAX);
match admin.cluster_commit() {
Ok(()) => {
let resp = DynRpbClusterCommitResp { applied };
Ok(Frame::new(
MessageCode::DynClusterCommitResp.as_u8(),
resp.encode_to_vec(),
))
}
Err(e) => Ok(error_frame(format_cluster_error("cluster-commit", &e))),
}
}
fn handle_aae_status(body: &[u8], aae: &dyn AaeStatusProvider) -> Result<Frame, RiakError> {
let _ = DynRpbAaeStatusReq::decode(body)?;
let snap: AaeStatusSnapshot = aae.current_status();
let resp = DynRpbAaeStatusResp {
peers: snap
.peers
.iter()
.map(|p| DynRpbAaePeerStatus {
peer_idx: p.peer_idx,
dc: p.dc.as_bytes().to_vec(),
rack: p.rack.as_bytes().to_vec(),
last_exchange_unix: p.last_exchange_unix,
divergent_keys_since_last_full_sweep: p.divergent_keys_since_last_full_sweep,
repair_dispatched_total: p.repair_dispatched_total,
})
.collect(),
snapshot_path: snap.snapshot_path.into_bytes(),
snapshot_last_save_unix: snap.snapshot_last_save_unix,
snapshot_last_load_unix: snap.snapshot_last_load_unix,
snapshot_save_total: snap.snapshot_save_total,
snapshot_load_total: snap.snapshot_load_total,
snapshot_corruption_total: snap.snapshot_corruption_total,
tree_n_time_buckets: snap.tree_n_time_buckets,
tree_n_segments: snap.tree_n_segments,
tree_time_window_seconds: snap.tree_time_window_seconds,
tree_memory_estimate_bytes: snap.tree_memory_estimate_bytes,
};
Ok(Frame::new(
MessageCode::DynAaeStatusResp.as_u8(),
resp.encode_to_vec(),
))
}
fn snapshot_to_pb(snap: &PeerSnapshot) -> DynRpbPeerInfo {
DynRpbPeerInfo {
idx: snap.idx,
dc: snap.dc.as_bytes().to_vec(),
rack: snap.rack.as_bytes().to_vec(),
host: snap.host.as_bytes().to_vec(),
port: u32::from(snap.port),
tokens: snap
.tokens
.iter()
.map(|t| t.to_string().into_bytes())
.collect(),
state: snap.state.name().as_bytes().to_vec(),
is_local: snap.is_local,
is_secure: None,
}
}
fn change_to_pb(change: &ClusterChange) -> DynRpbStagedChange {
let kind = match change.kind {
ClusterChangeKind::Add => DYN_STAGED_CHANGE_ADD,
ClusterChangeKind::Remove => DYN_STAGED_CHANGE_REMOVE,
};
let peer = change.peer.as_ref().map(|spec| DynRpbPeerInfo {
idx: 0,
dc: spec.dc.as_bytes().to_vec(),
rack: spec.rack.as_bytes().to_vec(),
host: spec.host.as_bytes().to_vec(),
port: u32::from(spec.port),
tokens: spec
.tokens
.iter()
.map(|t| t.to_string().into_bytes())
.collect(),
state: Vec::new(),
is_local: false,
is_secure: Some(spec.is_secure),
});
DynRpbStagedChange {
kind,
peer_idx: change.peer_idx,
peer,
}
}
fn format_cluster_error(op: &str, err: &ClusterError) -> String {
format!("{op}: {err}")
}
#[cfg(test)]
mod tests {
use super::*;
use dynomite::embed::MemoryDatastore;
use tokio::io::duplex;
#[tokio::test]
async fn ping_round_trips_over_duplex() {
let (client, server) = duplex(4096);
let ds: Arc<dyn Datastore> = Arc::new(MemoryDatastore::new());
let server_task = tokio::spawn(handle_conn(server, ds));
let (mut client_r, mut client_w) = tokio::io::split(client);
write_frame(
&mut client_w,
&Frame::new(MessageCode::PingReq.as_u8(), Vec::new()),
)
.await
.unwrap();
let resp = read_frame(&mut client_r).await.unwrap();
assert_eq!(resp.code, MessageCode::PingResp.as_u8());
assert!(resp.body.is_empty());
drop(client_r);
drop(client_w);
let _ = server_task.await.unwrap();
}
#[tokio::test]
async fn unknown_code_surfaces_error_to_caller() {
let frame = Frame::new(99, Vec::new());
let ds = MemoryDatastore::new();
let Err(err) =
process_frame(&frame, &ds, &NoopClusterAdmin, None, &NoopAaeStatusProvider).await
else {
panic!("expected error for unknown code");
};
assert!(matches!(err, RiakError::UnknownMessageCode(99)));
}
async fn collect_frames(mut s: FrameStream) -> Vec<Frame> {
let mut out = Vec::new();
while let Some(item) = s.next().await {
out.push(item.expect("stream item"));
}
out
}
#[tokio::test]
async fn response_codes_inbound_yield_error_resp() {
let frame = Frame::new(MessageCode::GetResp.as_u8(), Vec::new());
let ds = MemoryDatastore::new();
let stream = process_frame(&frame, &ds, &NoopClusterAdmin, None, &NoopAaeStatusProvider)
.await
.expect("ok");
let frames = collect_frames(stream).await;
assert_eq!(frames.len(), 1);
assert_eq!(frames[0].code, MessageCode::ErrorResp.as_u8());
let parsed = RpbErrorResp::decode(frames[0].body.as_slice()).expect("decode");
assert!(!parsed.errmsg.is_empty());
}
#[tokio::test]
async fn malformed_body_reports_decode_error() {
let frame = Frame::new(MessageCode::GetReq.as_u8(), vec![0x0a, 0xff]);
let ds = MemoryDatastore::new();
let Err(err) =
process_frame(&frame, &ds, &NoopClusterAdmin, None, &NoopAaeStatusProvider).await
else {
panic!("expected decode error");
};
assert!(matches!(err, RiakError::Decode(_)));
}
#[tokio::test]
async fn datastore_dispatch_is_invoked_for_kv_ops() {
let ds = Arc::new(MemoryDatastore::new());
let frame = Frame::new(
MessageCode::PutReq.as_u8(),
RpbPutReq {
bucket: b"b".to_vec(),
key: Some(b"k".to_vec()),
content: Some(RpbContent {
value: b"v".to_vec(),
..RpbContent::default()
}),
..RpbPutReq::default()
}
.encode_to_vec(),
);
let _ = process_frame(
&frame,
ds.as_ref(),
&NoopClusterAdmin,
None,
&NoopAaeStatusProvider,
)
.await
.expect("ok");
assert_eq!(ds.dispatch_count(), 1);
}
#[tokio::test]
async fn list_buckets_empty_yields_one_terminator_frame() {
let ds = MemoryDatastore::new();
let frame = Frame::new(
MessageCode::ListBucketsReq.as_u8(),
RpbListBucketsReq::default().encode_to_vec(),
);
let stream = process_frame(&frame, &ds, &NoopClusterAdmin, None, &NoopAaeStatusProvider)
.await
.expect("ok");
let frames = collect_frames(stream).await;
assert_eq!(frames.len(), 1);
assert_eq!(frames[0].code, MessageCode::ListBucketsResp.as_u8());
let resp = RpbListBucketsResp::decode(frames[0].body.as_slice()).expect("decode");
assert_eq!(resp.done, Some(true));
assert!(resp.buckets.is_empty());
}
#[tokio::test]
async fn list_keys_chunks_at_chunk_size() {
let ds = MemoryDatastore::new();
for i in 0..1000u16 {
ds.insert(b"u", format!("k{i:04}").as_bytes());
}
let frame = Frame::new(
MessageCode::ListKeysReq.as_u8(),
RpbListKeysReq {
bucket: b"u".to_vec(),
..RpbListKeysReq::default()
}
.encode_to_vec(),
);
let stream = process_frame(&frame, &ds, &NoopClusterAdmin, None, &NoopAaeStatusProvider)
.await
.expect("ok");
let frames = collect_frames(stream).await;
assert_eq!(frames.len(), 5, "expected 4 chunks plus a terminator");
let mut total_keys = 0usize;
for (i, f) in frames.iter().enumerate() {
assert_eq!(f.code, MessageCode::ListKeysResp.as_u8());
let resp = RpbListKeysResp::decode(f.body.as_slice()).expect("decode");
if i == frames.len() - 1 {
assert_eq!(resp.done, Some(true), "final frame must carry done=true");
assert!(resp.keys.is_empty(), "terminator carries no keys");
} else {
assert!(
resp.done == Some(false) || resp.done.is_none(),
"non-terminator frame must not carry done=true"
);
total_keys += resp.keys.len();
if i < 3 {
assert_eq!(resp.keys.len(), LIST_CHUNK_SIZE);
} else {
assert_eq!(resp.keys.len(), 1000 - 3 * LIST_CHUNK_SIZE);
}
}
}
assert_eq!(total_keys, 1000);
}
#[tokio::test]
async fn list_buckets_streams_multiple_buckets() {
let ds = MemoryDatastore::new();
for i in 0..512u16 {
ds.insert(format!("b{i:04}").as_bytes(), b"k");
}
let frame = Frame::new(
MessageCode::ListBucketsReq.as_u8(),
RpbListBucketsReq::default().encode_to_vec(),
);
let stream = process_frame(&frame, &ds, &NoopClusterAdmin, None, &NoopAaeStatusProvider)
.await
.expect("ok");
let frames = collect_frames(stream).await;
assert_eq!(frames.len(), 3);
let last = RpbListBucketsResp::decode(frames[2].body.as_slice()).expect("decode");
assert_eq!(last.done, Some(true));
assert!(last.buckets.is_empty());
}
#[tokio::test]
async fn list_keys_against_unsupported_datastore_yields_error_frame() {
struct Noop;
impl Datastore for Noop {
fn protocol(&self) -> dynomite::embed::hooks::Protocol {
dynomite::embed::hooks::Protocol::Custom
}
fn dispatch(
&self,
req: Msg,
) -> dynomite::embed::hooks::BoxFuture<
'_,
Result<Msg, dynomite::embed::hooks::DatastoreError>,
> {
Box::pin(async move {
let mut rsp = Msg::new(req.id(), MsgType::Unknown, false);
rsp.set_parent_id(req.id());
Ok(rsp)
})
}
}
let ds = Noop;
let frame = Frame::new(
MessageCode::ListKeysReq.as_u8(),
RpbListKeysReq {
bucket: b"u".to_vec(),
..RpbListKeysReq::default()
}
.encode_to_vec(),
);
let stream = process_frame(&frame, &ds, &NoopClusterAdmin, None, &NoopAaeStatusProvider)
.await
.expect("ok");
let frames = collect_frames(stream).await;
assert_eq!(frames.len(), 1);
assert_eq!(frames[0].code, MessageCode::ErrorResp.as_u8());
let resp = RpbErrorResp::decode(frames[0].body.as_slice()).expect("decode");
assert!(
resp.errmsg
.windows(b"unsupported".len())
.any(|w| w == b"unsupported"),
"errmsg should mention the unsupported variant: {:?}",
String::from_utf8_lossy(&resp.errmsg)
);
}
struct ScriptedIndexStore {
keys: Vec<Vec<u8>>,
}
impl Datastore for ScriptedIndexStore {
fn protocol(&self) -> dynomite::embed::hooks::Protocol {
dynomite::embed::hooks::Protocol::Custom
}
fn dispatch(
&self,
req: Msg,
) -> dynomite::embed::hooks::BoxFuture<
'_,
Result<Msg, dynomite::embed::hooks::DatastoreError>,
> {
Box::pin(async move {
let mut rsp = Msg::new(req.id(), MsgType::Unknown, false);
rsp.set_parent_id(req.id());
Ok(rsp)
})
}
fn riak_index_eq<'a>(
&'a self,
_bucket: &'a [u8],
_index_name: &'a [u8],
_value: &'a [u8],
) -> dynomite::embed::hooks::BoxFuture<
'a,
Result<Vec<Vec<u8>>, dynomite::embed::hooks::DatastoreError>,
> {
let keys = self.keys.clone();
Box::pin(async move { Ok(keys) })
}
fn riak_index_range<'a>(
&'a self,
_bucket: &'a [u8],
_index_name: &'a [u8],
_min: &'a [u8],
_max: &'a [u8],
) -> dynomite::embed::hooks::BoxFuture<
'a,
Result<Vec<Vec<u8>>, dynomite::embed::hooks::DatastoreError>,
> {
let keys = self.keys.clone();
Box::pin(async move { Ok(keys) })
}
}
#[tokio::test]
async fn index_eq_streams_chunks_of_chunk_size() {
let mut keys = Vec::new();
for i in 0..1000u16 {
keys.push(format!("k{i:04}").as_bytes().to_vec());
}
let ds = ScriptedIndexStore { keys };
let frame = Frame::new(
MessageCode::IndexReq.as_u8(),
RpbIndexReq {
bucket: b"u".to_vec(),
index: b"age_int".to_vec(),
qtype: INDEX_QUERY_TYPE_EQ,
key: Some(b"42".to_vec()),
..RpbIndexReq::default()
}
.encode_to_vec(),
);
let stream = process_frame(&frame, &ds, &NoopClusterAdmin, None, &NoopAaeStatusProvider)
.await
.expect("ok");
let frames = collect_frames(stream).await;
assert_eq!(frames.len(), 5, "4 chunks plus a terminator");
let mut total = 0usize;
for (i, f) in frames.iter().enumerate() {
assert_eq!(f.code, MessageCode::IndexResp.as_u8());
let resp = RpbIndexResp::decode(f.body.as_slice()).expect("decode");
if i == frames.len() - 1 {
assert_eq!(resp.done, Some(true));
assert!(resp.keys.is_empty());
} else {
assert_eq!(
resp.done,
Some(false),
"non-terminator frames carry done=false"
);
total += resp.keys.len();
if i < 3 {
assert_eq!(resp.keys.len(), LIST_CHUNK_SIZE);
} else {
assert_eq!(resp.keys.len(), 1000 - 3 * LIST_CHUNK_SIZE);
}
}
}
assert_eq!(total, 1000);
}
#[tokio::test]
async fn index_eq_first_frame_carries_partial_keys_for_old_clients() {
let mut keys = Vec::new();
for i in 0..600u16 {
keys.push(format!("k{i:04}").as_bytes().to_vec());
}
let ds = ScriptedIndexStore { keys };
let frame = Frame::new(
MessageCode::IndexReq.as_u8(),
RpbIndexReq {
bucket: b"u".to_vec(),
index: b"x".to_vec(),
qtype: INDEX_QUERY_TYPE_EQ,
key: Some(b"v".to_vec()),
..RpbIndexReq::default()
}
.encode_to_vec(),
);
let mut stream =
process_frame(&frame, &ds, &NoopClusterAdmin, None, &NoopAaeStatusProvider)
.await
.expect("ok");
let first = stream.next().await.expect("first").expect("frame");
assert_eq!(first.code, MessageCode::IndexResp.as_u8());
let parsed = RpbIndexResp::decode(first.body.as_slice()).expect("decode");
assert_eq!(parsed.keys.len(), LIST_CHUNK_SIZE);
assert_eq!(parsed.done, Some(false));
}
#[tokio::test]
async fn index_empty_yields_single_terminator() {
let ds = ScriptedIndexStore { keys: Vec::new() };
let frame = Frame::new(
MessageCode::IndexReq.as_u8(),
RpbIndexReq {
bucket: b"u".to_vec(),
index: b"x".to_vec(),
qtype: INDEX_QUERY_TYPE_EQ,
key: Some(b"v".to_vec()),
..RpbIndexReq::default()
}
.encode_to_vec(),
);
let stream = process_frame(&frame, &ds, &NoopClusterAdmin, None, &NoopAaeStatusProvider)
.await
.expect("ok");
let frames = collect_frames(stream).await;
assert_eq!(frames.len(), 1);
let resp = RpbIndexResp::decode(frames[0].body.as_slice()).expect("decode");
assert_eq!(resp.done, Some(true));
assert!(resp.keys.is_empty());
}
#[tokio::test]
async fn index_eq_max_results_caps_total_streamed_keys() {
let mut keys = Vec::new();
for i in 0..1000u16 {
keys.push(format!("k{i:04}").as_bytes().to_vec());
}
let ds = ScriptedIndexStore { keys };
let frame = Frame::new(
MessageCode::IndexReq.as_u8(),
RpbIndexReq {
bucket: b"u".to_vec(),
index: b"x".to_vec(),
qtype: INDEX_QUERY_TYPE_EQ,
key: Some(b"v".to_vec()),
max_results: Some(300),
..RpbIndexReq::default()
}
.encode_to_vec(),
);
let stream = process_frame(&frame, &ds, &NoopClusterAdmin, None, &NoopAaeStatusProvider)
.await
.expect("ok");
let frames = collect_frames(stream).await;
let mut total = 0usize;
for f in &frames {
let resp = RpbIndexResp::decode(f.body.as_slice()).expect("decode");
total += resp.keys.len();
}
assert_eq!(total, 300);
let last = frames.last().expect("last");
let last_resp = RpbIndexResp::decode(last.body.as_slice()).expect("decode");
assert_eq!(last_resp.done, Some(true));
}
#[tokio::test]
async fn index_unsupported_datastore_yields_error_frame() {
struct Noop;
impl Datastore for Noop {
fn protocol(&self) -> dynomite::embed::hooks::Protocol {
dynomite::embed::hooks::Protocol::Custom
}
fn dispatch(
&self,
req: Msg,
) -> dynomite::embed::hooks::BoxFuture<
'_,
Result<Msg, dynomite::embed::hooks::DatastoreError>,
> {
Box::pin(async move {
let mut rsp = Msg::new(req.id(), MsgType::Unknown, false);
rsp.set_parent_id(req.id());
Ok(rsp)
})
}
}
let ds = Noop;
let frame = Frame::new(
MessageCode::IndexReq.as_u8(),
RpbIndexReq {
bucket: b"u".to_vec(),
index: b"x".to_vec(),
qtype: INDEX_QUERY_TYPE_EQ,
key: Some(b"v".to_vec()),
..RpbIndexReq::default()
}
.encode_to_vec(),
);
let stream = process_frame(&frame, &ds, &NoopClusterAdmin, None, &NoopAaeStatusProvider)
.await
.expect("ok");
let frames = collect_frames(stream).await;
assert_eq!(frames.len(), 1);
assert_eq!(frames[0].code, MessageCode::ErrorResp.as_u8());
}
fn mapred_req_two_phase(values: &[i64]) -> Vec<u8> {
let inputs: Vec<serde_json::Value> = values
.iter()
.enumerate()
.map(|(i, v)| {
serde_json::json!({
"bucket": "b",
"key": format!("k{i}"),
"value": *v,
})
})
.collect();
let job = serde_json::json!({
"inputs": inputs,
"query": [
{ "map": { "language": "erlang",
"name": "map_object_value",
"keep": true } },
{ "reduce": { "language": "erlang",
"name": "reduce_sum",
"keep": true } },
]
});
let req = RpbMapRedReq {
request: serde_json::to_vec(&job).expect("job json"),
content_type: b"application/json".to_vec(),
};
req.encode_to_vec()
}
#[tokio::test]
async fn process_frame_streams_mapreduce_response_with_per_phase_frames() {
let body = mapred_req_two_phase(&[1, 2, 3]);
let frame = Frame::new(MessageCode::MapRedReq.as_u8(), body);
let ds = MemoryDatastore::new();
let stream = process_frame(&frame, &ds, &NoopClusterAdmin, None, &NoopAaeStatusProvider)
.await
.expect("ok");
let frames = collect_frames(stream).await;
assert_eq!(
frames.len(),
3,
"expected two per-phase frames plus one terminator",
);
for f in &frames {
assert_eq!(f.code, MessageCode::MapRedResp.as_u8());
}
let p0 = RpbMapRedResp::decode(frames[0].body.as_slice()).expect("decode 0");
assert_eq!(p0.phase, Some(0));
assert_eq!(p0.done, Some(false));
let p0_body = p0.response.as_ref().expect("phase 0 body");
let p0_json: serde_json::Value = serde_json::from_slice(p0_body).expect("phase 0 json");
assert_eq!(p0_json[0]["phase"], 0);
assert_eq!(p0_json[0]["data"].as_array().unwrap().len(), 3);
let p1 = RpbMapRedResp::decode(frames[1].body.as_slice()).expect("decode 1");
assert_eq!(p1.phase, Some(1));
assert_eq!(p1.done, Some(false));
let p1_body = p1.response.as_ref().expect("phase 1 body");
let p1_json: serde_json::Value = serde_json::from_slice(p1_body).expect("phase 1 json");
assert_eq!(p1_json[0]["phase"], 1);
assert_eq!(p1_json[0]["data"], serde_json::json!([6]));
}
#[tokio::test]
async fn process_frame_emits_terminator_frame_with_done_true() {
let body = mapred_req_two_phase(&[10, 20]);
let frame = Frame::new(MessageCode::MapRedReq.as_u8(), body);
let ds = MemoryDatastore::new();
let stream = process_frame(&frame, &ds, &NoopClusterAdmin, None, &NoopAaeStatusProvider)
.await
.expect("ok");
let frames = collect_frames(stream).await;
let term = frames.last().expect("at least one frame");
assert_eq!(term.code, MessageCode::MapRedResp.as_u8());
let parsed = RpbMapRedResp::decode(term.body.as_slice()).expect("decode terminator");
assert_eq!(parsed.done, Some(true));
assert_eq!(parsed.phase, None);
assert!(parsed.response.is_none());
}
#[tokio::test]
async fn mapreduce_first_frame_is_a_partial_phase_zero_answer() {
let body = mapred_req_two_phase(&[5, 6, 7, 8]);
let frame = Frame::new(MessageCode::MapRedReq.as_u8(), body);
let ds = MemoryDatastore::new();
let mut stream =
process_frame(&frame, &ds, &NoopClusterAdmin, None, &NoopAaeStatusProvider)
.await
.expect("ok");
let first = stream.next().await.expect("first").expect("frame");
assert_eq!(first.code, MessageCode::MapRedResp.as_u8());
let parsed = RpbMapRedResp::decode(first.body.as_slice()).expect("decode");
assert_eq!(parsed.phase, Some(0));
assert_eq!(parsed.done, Some(false));
let body = parsed.response.expect("first frame carries body");
let json: serde_json::Value = serde_json::from_slice(&body).expect("json");
assert_eq!(json[0]["phase"], 0);
assert_eq!(json[0]["data"].as_array().unwrap().len(), 4);
}
#[tokio::test]
async fn mapreduce_unknown_function_emits_single_error_frame() {
let job = serde_json::json!({
"inputs": [{"bucket": "b", "key": "k"}],
"query": [
{ "map": { "language": "erlang",
"name": "no_such_function",
"keep": true } }
]
});
let req = RpbMapRedReq {
request: serde_json::to_vec(&job).expect("job json"),
content_type: b"application/json".to_vec(),
};
let frame = Frame::new(MessageCode::MapRedReq.as_u8(), req.encode_to_vec());
let ds = MemoryDatastore::new();
let stream = process_frame(&frame, &ds, &NoopClusterAdmin, None, &NoopAaeStatusProvider)
.await
.expect("ok");
let frames = collect_frames(stream).await;
assert_eq!(frames.len(), 1);
assert_eq!(frames[0].code, MessageCode::ErrorResp.as_u8());
let parsed = RpbErrorResp::decode(frames[0].body.as_slice()).expect("decode");
let msg = String::from_utf8_lossy(&parsed.errmsg);
assert!(
msg.contains("no_such_function") || msg.contains("unknown"),
"errmsg: {msg}",
);
}
#[tokio::test]
async fn mapreduce_unsupported_content_type_yields_error_frame() {
let req = RpbMapRedReq {
request: b"<xml/>".to_vec(),
content_type: b"application/xml".to_vec(),
};
let frame = Frame::new(MessageCode::MapRedReq.as_u8(), req.encode_to_vec());
let ds = MemoryDatastore::new();
let stream = process_frame(&frame, &ds, &NoopClusterAdmin, None, &NoopAaeStatusProvider)
.await
.expect("ok");
let frames = collect_frames(stream).await;
assert_eq!(frames.len(), 1);
assert_eq!(frames[0].code, MessageCode::ErrorResp.as_u8());
}
#[tokio::test]
async fn aae_status_default_provider_returns_empty_snapshot() {
let ds = MemoryDatastore::new();
let frame = Frame::new(
MessageCode::DynAaeStatusReq.as_u8(),
DynRpbAaeStatusReq::default().encode_to_vec(),
);
let stream = process_frame(&frame, &ds, &NoopClusterAdmin, None, &NoopAaeStatusProvider)
.await
.expect("ok");
let frames = collect_frames(stream).await;
assert_eq!(frames.len(), 1);
assert_eq!(frames[0].code, MessageCode::DynAaeStatusResp.as_u8());
let resp = DynRpbAaeStatusResp::decode(frames[0].body.as_slice()).expect("decode");
assert!(resp.peers.is_empty());
assert_eq!(resp.snapshot_save_total, 0);
}
#[tokio::test]
async fn aae_status_custom_provider_returns_live_snapshot() {
struct Provider;
impl crate::aae::status::AaeStatusProvider for Provider {
fn current_status(&self) -> crate::aae::status::AaeStatusSnapshot {
crate::aae::status::AaeStatusSnapshot {
peers: vec![crate::aae::status::AaePeerStatus {
peer_idx: 7,
dc: "dc1".into(),
rack: "rA".into(),
last_exchange_unix: 1_700_000_000,
divergent_keys_since_last_full_sweep: 4,
repair_dispatched_total: 3,
}],
snapshot_path: "/var/lib/dynomite/aae/tree.snapshot".into(),
snapshot_last_save_unix: 1_700_000_300,
snapshot_last_load_unix: 1_700_000_100,
snapshot_save_total: 5,
snapshot_load_total: 1,
snapshot_corruption_total: 0,
tree_n_time_buckets: 24,
tree_n_segments: 1024,
tree_time_window_seconds: 3600,
tree_memory_estimate_bytes: 8192,
}
}
}
let ds = MemoryDatastore::new();
let frame = Frame::new(
MessageCode::DynAaeStatusReq.as_u8(),
DynRpbAaeStatusReq::default().encode_to_vec(),
);
let stream = process_frame(&frame, &ds, &NoopClusterAdmin, None, &Provider)
.await
.expect("ok");
let frames = collect_frames(stream).await;
assert_eq!(frames.len(), 1);
assert_eq!(frames[0].code, MessageCode::DynAaeStatusResp.as_u8());
let resp = DynRpbAaeStatusResp::decode(frames[0].body.as_slice()).expect("decode");
assert_eq!(resp.peers.len(), 1);
assert_eq!(resp.peers[0].peer_idx, 7);
assert_eq!(resp.peers[0].dc, b"dc1".to_vec());
assert_eq!(resp.snapshot_save_total, 5);
assert_eq!(resp.tree_n_time_buckets, 24);
assert_eq!(
resp.snapshot_path,
b"/var/lib/dynomite/aae/tree.snapshot".to_vec()
);
}
}