Skip to content
Open
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
95 changes: 94 additions & 1 deletion src/cms/cms_command_handler.rs
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
use valkey_module::{
Context, NotifyEvent, ValkeyError, ValkeyResult, ValkeyString, ValkeyValue, VALKEY_OK,
key::ValkeyKey, Context, NotifyEvent, ValkeyError, ValkeyResult, ValkeyString, ValkeyValue,
VALKEY_OK,
};

use crate::cms::data_type::CMS_TYPE;
Expand All @@ -13,6 +14,7 @@ enum Replications {
enum Operation {
Initialization { replications: Replications },
Increment,
Merge,
}

fn replicate_and_notify_events(ctx: &Context, key_name: &ValkeyString, operation: Operation) {
Expand Down Expand Up @@ -52,6 +54,11 @@ fn replicate_and_notify_events(ctx: &Context, key_name: &ValkeyString, operation
ctx.replicate_verbatim();
ctx.notify_keyspace_event(NotifyEvent::GENERIC, utils::INCR_EVENT, key_name);
}
Operation::Merge => {
//TODO::How should this replication be done, and why vs verbatim
ctx.replicate_verbatim();
ctx.notify_keyspace_event(NotifyEvent::GENERIC, utils::MERGE_EVENT, key_name);
}
}
}

Expand Down Expand Up @@ -239,6 +246,92 @@ pub fn cms_query(ctx: &Context, args: Vec<ValkeyString>) -> ValkeyResult {
}
}

/// Function that implements logic to handle the CMS.MERGE command.
pub fn cms_merge(ctx: &Context, args: Vec<ValkeyString>) -> ValkeyResult {
let args_count = args.len();
if args_count < 5 {
return Err(ValkeyError::WrongArity);
}

//This must already be initialized.
let destination_key = &args[1];

let number_of_keys_value = args[2]
.to_string_lossy()
.parse::<usize>()
.map_err(|_| ValkeyError::Str("ERR invalid number of keys value"))?;

//Indexes 3 -> 3 + N-1 are keys to merge
let sketch_end_index = 3 + number_of_keys_value - 1;
//Up to non inclusive grab 3 up to sketch_end_index
let source_keys: Vec<&ValkeyString> = args[3..=sketch_end_index].iter().collect();

//Then Parse the optional WEIGHTS section of the command. WEIGHTS is at sketch_end + 1
let passed_in_weights = if sketch_end_index + 1 == args_count {
Vec::new()
} else {
//Make sure we have at least WEIGHT weight left
let weight_args_left = args_count - sketch_end_index - 1;
if weight_args_left < 2 {
return Err(ValkeyError::WrongArity);
}

//There should be at least 2 args left WEIGHT weight [weight ...]
let weights_keyword_index = sketch_end_index + 1;
let weights_keyword = args[weights_keyword_index].to_string_lossy();
if weights_keyword.to_uppercase() != "WEIGHTS" {
return Err(ValkeyError::Str("ERR invalid argument"));
}
let weights_start = weights_keyword_index + 1;
let weights_args = args[weights_start..].iter();

let weights: Vec<f64> = weights_args
.map(|weight| {
weight
.parse_float()
.map_err(|_| ValkeyError::Str("ERR invalid weight value"))
})
.collect::<Result<Vec<_>, _>>()?;
weights
};

let source_key_handles: Vec<ValkeyKey> =
source_keys.iter().map(|key| ctx.open_key(key)).collect();
let sketches = source_key_handles
.iter()
.map(|key_handle| {
key_handle
.get_value::<CMSObject>(&CMS_TYPE)
.and_then(|opt| opt.ok_or_else(|| ValkeyError::Str("ERR key does not exist")))
})
.collect::<Result<Vec<_>, _>>()?;

let destination_sketch = ctx
.open_key_writable(destination_key)
.get_value::<CMSObject>(&CMS_TYPE)
.and_then(|opt| opt.ok_or_else(|| ValkeyError::Str("ERR key does not exist")))?;

let sketches_with_weights: Vec<(&CMSObject, f64)> = sketches
.into_iter()
.enumerate()
.map(|(i, sketch)| {
let weight = passed_in_weights.get(i).copied().unwrap_or(1.0);
(sketch, weight)
})
.collect();

//Mutates the destination_sketch's internal CMS to be the merge of the sketches_with_weights
//Impl note: We do not handle the weights yet in the called function, as the lib-source needs to change.
destination_sketch
.merge(&sketches_with_weights)
.map_err(|_| {
ValkeyError::Str("ERR destination key is not of the same width and/or depth")
})?;

replicate_and_notify_events(ctx, destination_key, Operation::Merge);
VALKEY_OK
}

//Function that implements logic to handle the CMS.INFO command.
pub fn cms_info(ctx: &Context, args: Vec<ValkeyString>) -> ValkeyResult {
let args_count = args.len();
Expand Down
36 changes: 35 additions & 1 deletion src/cms/utils.rs
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
use flowstats::CountMinSketch;
use flowstats::{traits::Sketch, CountMinSketch};

/// Client Errors
pub const ERROR: &str = "ERROR";
Expand All @@ -12,19 +12,22 @@ pub const BAD_PROBABILITY: &str = "ERR bad probability";
pub const PROBABILITY_RANGE: &str = "ERR probability rate should be between 0 and 1";
pub const KEY_EXISTS: &str = "ERR Target key name already exists.";
pub const BAD_INCREMENT: &str = "ERR bad increment";
pub const MERGE_FAILURE: &str = "ERR width / depth is not equal across sketches";
pub const INVALID_INFO_VALUE: &str = "ERR invalid information value";

///Keyspace Notification Events
pub const INITBYPROB_EVENT: &str = "countminsketch.initbyprob";
pub const INITBYDIM_EVENT: &str = "countminsketch.initbydim";
pub const INCR_EVENT: &str = "countminsketch.incrby";
pub const MERGE_EVENT: &str = "countminsketch.merge";

#[derive(Debug, PartialEq)]
pub enum CMSError {
InvalidWidth,
InvalidDepth,
InvalidErrorRate,
InvalidProbability,
MergeFailed,
}

impl CMSError {
Expand All @@ -34,6 +37,7 @@ impl CMSError {
CMSError::InvalidDepth => BAD_DEPTH,
CMSError::InvalidErrorRate => ERROR_RATE_RANGE,
CMSError::InvalidProbability => PROBABILITY_RANGE,
CMSError::MergeFailed => MERGE_FAILURE,
}
}
}
Expand Down Expand Up @@ -96,6 +100,36 @@ impl CMSObject {
pub fn estimate(&self, item: &[u8]) -> u64 {
self.cms.estimate_item(item)
}

/// Merges the current CMS structure with a list of others as well as their weights. Weights are applied to the paired structure via multiplication first.
/// The values in each index of each of the K arrays internally then are added together into a single CMS structure.
pub fn merge(&mut self, sketches_and_weights: &[(&CMSObject, f64)]) -> Result<(), CMSError> {
//Pre-check this so dest is the same size as the others.
if sketches_and_weights
.iter()
.any(|sketch| sketch.0.width != self.width || sketch.0.depth != self.depth)
{
return Err(CMSError::MergeFailed);
}

let (head, tail) = sketches_and_weights
.split_first()
.ok_or(CMSError::MergeFailed)?;

let mut merged_sketch = head.0.cms.sketch.clone();

for (cms, _weight) in tail.iter() {
merged_sketch
.merge(&cms.cms.sketch)
.map_err(|_| CMSError::MergeFailed)?;
}

self.cms = CMS {
sketch: merged_sketch,
};

Ok(())
}
}

struct CMS {
Expand Down
6 changes: 6 additions & 0 deletions src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -114,6 +114,11 @@ fn cms_query_command(ctx: &Context, args: Vec<ValkeyString>) -> ValkeyResult {
cms_command_handler::cms_query(ctx, args)
}

/// Command handler for CMS.MERGE <destination> <numkeys> <source> [<source> ...] [WEIGHTS <weight> [<weight> ...]]
fn cms_merge_command(ctx: &Context, args: Vec<ValkeyString>) -> ValkeyResult {
cms_command_handler::cms_merge(ctx, args)
}

/// Command handler for CMS.INFO <key>
fn cms_info_command(ctx: &Context, args: Vec<ValkeyString>) -> ValkeyResult {
cms_command_handler::cms_info(ctx, args)
Expand Down Expand Up @@ -157,6 +162,7 @@ valkey_module! {
["CMS.INITBYPROB", cms_initbyprob_command, "write fast deny-oom", 1, 1, 1, "fast write cms"],
["CMS.INCRBY", cms_incrby_command, "write fast deny-oom", 1, 1, 1, "write cms"],
["CMS.QUERY", cms_query_command, "readonly fast", 1, 1, 1, "read cms"],
["CMS.MERGE", cms_merge_command, "write deny-oom", 1, 1, 1, "write cms"],
["CMS.INFO", cms_info_command, "readonly fast", 1, 1, 1, "fast read cms"],
],
configurations: [
Expand Down
21 changes: 21 additions & 0 deletions tests/test_cms_basic.py
Original file line number Diff line number Diff line change
Expand Up @@ -58,6 +58,27 @@ def test_basic_prob(self):
#CMS guarantees that we have the frequency at LEAST the size of the increment for the item
assert client.execute_command('CMS.QUERY sketch1 item1')[0] >= 1


def test_merge(self):
client = self.server.get_new_client()
module_loaded = False
module_list_data = client.execute_command('MODULE LIST')
for module in module_list_data:
if (module[b'name'] == b'bf'):
module_loaded = True
break
assert(module_loaded)
#Create the destination and sketches to be merged into the destination key
assert client.execute_command('CMS.INITBYDIM dest 10 5') == b'OK'
assert client.execute_command('CMS.INITBYDIM s1 10 5') == b'OK'
assert client.execute_command('CMS.INITBYDIM s2 10 5') == b'OK'
assert client.execute_command('CMS.INCRBY s1 a 1 b 2') == [1, 2]
assert client.execute_command('CMS.INCRBY s2 a 1 b 3') == [1, 3]
assert client.execute_command('CMS.MERGE dest 2 s1 s2') == b'OK'
assert client.execute_command('CMS.QUERY dest a') == [2]
assert client.execute_command('CMS.QUERY dest b') == [5]


def test_module_data_type(self):
# Validate the name of the Module data type.
client = self.server.get_new_client()
Expand Down
13 changes: 10 additions & 3 deletions tests/test_cms_command.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,9 @@ def test_cms_command_error(self):
('CMS.INCRBY sketch item abc', 'bad increment'),
('CMS.INCRBY sketch item -1', 'bad increment'),

('CMS.MERGE dest 2 s1 s2 WEIGHTSS 1', 'invalid argument'),
('CMS.MERGE dest 2 s1 s2 WEIGHTS a', 'invalid weight value'),

# wrong number of arguments
('CMS.INITBYDIM', "wrong number of arguments for 'CMS.INITBYDIM' command"),
('CMS.INITBYDIM key', "wrong number of arguments for 'CMS.INITBYDIM' command"),
Expand All @@ -50,14 +53,18 @@ def test_cms_command_error(self):

('CMS.QUERY', "wrong number of arguments for 'CMS.QUERY' command"),
('CMS.QUERY key', "wrong number of arguments for 'CMS.QUERY' command"),

('CMS.MERGE', "wrong number of arguments for 'CMS.MERGE' command"),
('CMS.MERGE dest', "wrong number of arguments for 'CMS.MERGE' command"),
('CMS.MERGE dest 2', "wrong number of arguments for 'CMS.MERGE' command"),
('CMS.MERGE dest 2 s1', "wrong number of arguments for 'CMS.MERGE' command"),
('CMS.MERGE dest 2 s1 s2 WEIGHTS', "wrong number of arguments for 'CMS.MERGE' command"),

('CMS.INFO', "wrong number of arguments for 'CMS.INFO' command"),
('CMS.INFO sketch WIDTH WIDTH', "wrong number of arguments for 'CMS.INFO' command"),

# Invalid parameter name (WIDTH, DEPTH, COUNT)
('CMS.INFO sketch NOTAPARAM', "invalid information value"),




]

Expand Down
Loading