use regex::bytes::Regex;
use serde::{Deserialize, Serialize};
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use crate::error::{Error, Result};
use crate::types::{Key, Value};
pub const MAX_KEY_PATTERN_BYTES: usize = 4 * 1024;
pub const MAX_KEY_SEARCH_LIMIT: usize = 10_000;
pub const MAX_KEY_SEARCH_SCAN_BUDGET: usize = 1_000_000;
pub const DEFAULT_KEY_SEARCH_RESPONSE_BYTES: usize = 16 * 1024 * 1024;
pub const MAX_KEY_SEARCH_RESPONSE_BYTES: usize = 100 * 1024 * 1024;
#[derive(Debug, Clone, Default)]
pub struct KeySearchCancellation(Arc<AtomicBool>);
impl KeySearchCancellation {
pub fn cancel(&self) {
self.0.store(true, Ordering::Release);
}
pub fn is_cancelled(&self) -> bool {
self.0.load(Ordering::Acquire)
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(tag = "mode", rename_all = "snake_case")]
pub enum KeyPattern {
Glob {
pattern: Vec<u8>,
},
Regex {
pattern: String,
},
}
impl KeyPattern {
pub fn glob(pattern: impl AsRef<[u8]>) -> Self {
Self::Glob {
pattern: pattern.as_ref().to_vec(),
}
}
pub fn regex(pattern: impl Into<String>) -> Self {
Self::Regex {
pattern: pattern.into(),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct KeySearchRequest {
pub pattern: KeyPattern,
pub cursor: Option<Key>,
pub limit: usize,
pub scan_budget: usize,
#[serde(default = "default_response_bytes")]
pub max_bytes: usize,
}
impl KeySearchRequest {
pub fn new(pattern: KeyPattern, limit: usize, scan_budget: usize) -> Self {
Self {
pattern,
cursor: None,
limit,
scan_budget,
max_bytes: DEFAULT_KEY_SEARCH_RESPONSE_BYTES,
}
}
pub fn after(mut self, cursor: Key) -> Self {
self.cursor = Some(cursor);
self
}
pub fn with_max_bytes(mut self, max_bytes: usize) -> Self {
self.max_bytes = max_bytes;
self
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct KeySearchEntry {
pub key: Key,
pub value: Value,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct KeySearchPage {
pub entries: Vec<KeySearchEntry>,
pub next_cursor: Option<Key>,
pub scanned: usize,
}
pub(crate) struct PreparedKeySearch {
matcher: Regex,
prefix: Key,
}
impl PreparedKeySearch {
pub(crate) fn new(request: &KeySearchRequest) -> Result<Self> {
validate_request(request)?;
let (source, prefix) = match &request.pattern {
KeyPattern::Glob { pattern } => (glob_regex(pattern)?, glob_prefix(pattern)?),
KeyPattern::Regex { pattern } => (pattern.clone(), Vec::new()),
};
let matcher = Regex::new(&source).map_err(|error| Error::InvalidParameter {
param: "pattern".into(),
reason: error.to_string(),
})?;
Ok(Self { matcher, prefix })
}
pub(crate) fn prefix(&self) -> &[u8] {
&self.prefix
}
pub(crate) fn collect(
&self,
mut next: impl FnMut() -> Result<Option<(Key, Option<Value>)>>,
request: &KeySearchRequest,
cancellation: &KeySearchCancellation,
) -> Result<KeySearchPage> {
let mut entries = Vec::with_capacity(request.limit);
let mut scanned = 0usize;
let mut scanned_bytes = 0usize;
let mut response_bytes = 0usize;
loop {
if cancellation.is_cancelled() {
return Err(Error::SearchCancelled);
}
if scanned == request.scan_budget {
return Err(Error::SearchBudgetExceeded {
limit: request.scan_budget,
});
}
let Some((key, value)) = next()? else {
break;
};
if !self.prefix.is_empty() && !key.starts_with(&self.prefix) {
break;
}
scanned += 1;
scanned_bytes = scanned_bytes
.saturating_add(key.len())
.saturating_add(value.as_ref().map_or(0, Vec::len));
if scanned_bytes > request.max_bytes {
return Err(Error::SearchResponseTooLarge {
limit: request.max_bytes,
requested: scanned_bytes,
});
}
if let Some(value) = value.filter(|_| self.matcher.is_match(&key)) {
let cursor_bytes = if entries.len() + 1 == request.limit {
key.len()
} else {
0
};
let requested = response_bytes
.saturating_add(key.len())
.saturating_add(value.len())
.saturating_add(cursor_bytes);
if requested > request.max_bytes {
return Err(Error::SearchResponseTooLarge {
limit: request.max_bytes,
requested,
});
}
response_bytes = requested;
entries.push(KeySearchEntry { key, value });
if entries.len() == request.limit {
break;
}
}
}
let next_cursor = (entries.len() == request.limit)
.then(|| entries.last().expect("non-empty full page").key.clone());
Ok(KeySearchPage {
entries,
next_cursor,
scanned,
})
}
}
fn validate_request(request: &KeySearchRequest) -> Result<()> {
let pattern_len = match &request.pattern {
KeyPattern::Glob { pattern } => pattern.len(),
KeyPattern::Regex { pattern } => pattern.len(),
};
if pattern_len > MAX_KEY_PATTERN_BYTES {
return invalid(
"pattern",
format!("must be at most {MAX_KEY_PATTERN_BYTES} bytes"),
);
}
if !(1..=MAX_KEY_SEARCH_LIMIT).contains(&request.limit) {
return invalid(
"limit",
format!("must be between 1 and {MAX_KEY_SEARCH_LIMIT}"),
);
}
if !(1..=MAX_KEY_SEARCH_SCAN_BUDGET).contains(&request.scan_budget) {
return invalid(
"scan_budget",
format!("must be between 1 and {MAX_KEY_SEARCH_SCAN_BUDGET}"),
);
}
if !(1..=MAX_KEY_SEARCH_RESPONSE_BYTES).contains(&request.max_bytes) {
return invalid(
"max_bytes",
format!("must be between 1 and {MAX_KEY_SEARCH_RESPONSE_BYTES}"),
);
}
Ok(())
}
fn default_response_bytes() -> usize {
DEFAULT_KEY_SEARCH_RESPONSE_BYTES
}
fn invalid<T>(param: &str, reason: String) -> Result<T> {
Err(Error::InvalidParameter {
param: param.into(),
reason,
})
}
fn glob_regex(pattern: &[u8]) -> Result<String> {
let mut source = String::from("^(?-u:");
let mut escaped = false;
for &byte in pattern {
if escaped {
push_byte(&mut source, byte);
escaped = false;
} else {
match byte {
b'\\' => escaped = true,
b'*' => source.push_str("(?s:.*)"),
b'?' => source.push_str("(?s:.)"),
literal => push_byte(&mut source, literal),
}
}
}
if escaped {
return invalid("pattern", "glob ends with an escape byte".into());
}
source.push_str(")$");
Ok(source)
}
fn glob_prefix(pattern: &[u8]) -> Result<Key> {
let mut prefix = Vec::new();
let mut escaped = false;
for &byte in pattern {
if escaped {
prefix.push(byte);
escaped = false;
} else {
match byte {
b'\\' => escaped = true,
b'*' | b'?' => break,
literal => prefix.push(literal),
}
}
}
if escaped {
return invalid("pattern", "glob ends with an escape byte".into());
}
Ok(prefix)
}
fn push_byte(output: &mut String, byte: u8) {
use std::fmt::Write;
write!(output, "\\x{byte:02X}").expect("writing to String cannot fail");
}