#[cfg(feature = "redis")]
use crate::error::{OxCacheError, OxCacheResult};
#[cfg(feature = "redis")]
pub(super) const MAX_LUA_SCRIPT_LENGTH: usize = 10 * 1024;
#[cfg(feature = "redis")]
pub(super) const MAX_LUA_SCRIPT_KEYS: usize = 100;
#[cfg(feature = "redis")]
pub(super) const MAX_SCAN_PATTERN_LENGTH: usize = 256;
#[cfg(feature = "redis")]
pub(super) const MAX_SCAN_WILDCARDS: usize = 10;
#[cfg(feature = "redis")]
pub(super) const SCAN_COUNT_MIN: usize = 1;
#[cfg(feature = "redis")]
pub(super) const SCAN_COUNT_MAX: usize = 1000;
#[cfg(feature = "redis")]
pub(super) static LUA_LOOP_PATTERNS: &[(&str, &str)] = &[
(r"WHILE\s+TRUE", "WHILE TRUE 循环"),
(r"WHILE\s+1", "WHILE 1 循环"),
(r"REPEAT", "REPEAT 循环"),
(r"GOTO", "GOTO 语句"),
];
#[cfg(feature = "redis")]
pub(super) static LUA_LOOP_REGEXES: ::once_cell::sync::Lazy<Vec<::regex::Regex>> = ::once_cell::sync::Lazy::new(|| {
LUA_LOOP_PATTERNS
.iter()
.map(|(pattern, _)| ::regex::Regex::new(pattern).expect("Invalid loop pattern regex"))
.collect()
});
#[cfg(feature = "redis")]
pub(super) static WHITESPACE_REGEX: ::once_cell::sync::Lazy<::regex::Regex> =
::once_cell::sync::Lazy::new(|| ::regex::Regex::new(r"\s+").expect("Invalid whitespace regex"));
#[cfg(feature = "redis")]
const SQL_INJECTION_PATTERNS: &[(&str, &str)] = &[
("' OR '", "单引号后跟 OR 模式"),
("'--", "SQL 注释模式"),
("'; DROP", "SQL DROP 语句"),
("'; DELETE", "SQL DELETE 语句"),
("'; INSERT", "SQL INSERT 语句"),
("UNION SELECT", "SQL UNION 查询"),
("xp_cmdshell", "SQL Server 命令执行"),
("' OR '1'='1", "经典 SQL 注入永真条件"),
("admin'--", "SQL 认证绕过"),
];
#[cfg(feature = "redis")]
const PATH_TRAVERSAL_PATTERNS: &[&str] = &[
"../",
"..\\",
"%2e%2e",
"%252e%252e",
"..%2f",
"..%5c",
"%2e%2e%2f",
"%2e%2e%5c",
];
#[cfg(feature = "redis")]
const COMMAND_INJECTION_CHARS: &[char] = &[';', '|', '&', '`'];
#[cfg(feature = "redis")]
#[cfg_attr(docsrs, doc(cfg(feature = "security")))]
pub fn validate_redis_key(key: &str) -> OxCacheResult<()> {
use crate::security::{DANGEROUS_CHARS, MAX_KEY_LENGTH};
use crate::security::{validate_max_length, validate_no_dangerous_chars, validate_not_empty};
validate_not_empty(key, "Redis key")?;
validate_max_length(key, MAX_KEY_LENGTH, "Redis key")?;
validate_no_dangerous_chars(key, &DANGEROUS_CHARS, "Redis key")?;
check_control_characters(key, &DANGEROUS_CHARS)?;
check_sql_injection(key)?;
check_path_traversal(key)?;
check_command_injection(key)?;
Ok(())
}
#[cfg(feature = "redis")]
fn check_control_characters(key: &str, dangerous_chars: &[char]) -> OxCacheResult<()> {
for c in key.chars() {
if c.is_control() && !dangerous_chars.contains(&c) && c != '\t' {
return Err(OxCacheError::InvalidInput(format!(
"Redis key contains control character: U+{:04X}",
c as u32
)));
}
}
Ok(())
}
#[cfg(feature = "redis")]
fn check_sql_injection(key: &str) -> OxCacheResult<()> {
let key_upper = key.to_uppercase();
for (pattern, description) in SQL_INJECTION_PATTERNS {
if key_upper.contains(&pattern.to_uppercase()) {
return Err(OxCacheError::InvalidInput(format!(
"Redis key contains suspicious SQL injection pattern: {}",
description
)));
}
}
Ok(())
}
#[cfg(feature = "redis")]
fn check_path_traversal(key: &str) -> OxCacheResult<()> {
let key_lower = key.to_lowercase();
for pattern in PATH_TRAVERSAL_PATTERNS {
if key_lower.contains(&pattern.to_lowercase()) {
return Err(OxCacheError::InvalidInput(format!(
"Redis key contains path traversal pattern: {}",
pattern
)));
}
}
Ok(())
}
#[cfg(feature = "redis")]
fn check_command_injection(key: &str) -> OxCacheResult<()> {
for c in key.chars() {
if COMMAND_INJECTION_CHARS.contains(&c) {
return Err(OxCacheError::InvalidInput(format!(
"Redis key contains potential command injection character: {:?}",
c
)));
}
}
Ok(())
}
#[cfg(feature = "redis")]
#[cfg_attr(docsrs, doc(cfg(feature = "security")))]
pub fn validate_lua_script(script: &str, key_count: usize) -> OxCacheResult<()> {
if script.len() > MAX_LUA_SCRIPT_LENGTH {
return Err(OxCacheError::InvalidInput(format!(
"Lua script exceeds maximum length of {} bytes (got {} bytes)",
MAX_LUA_SCRIPT_LENGTH,
script.len()
)));
}
if key_count > MAX_LUA_SCRIPT_KEYS {
return Err(OxCacheError::InvalidInput(format!(
"Lua script exceeds maximum key count of {} (got {} keys)",
MAX_LUA_SCRIPT_KEYS, key_count
)));
}
let cleaned = preprocess_lua_script(script);
let cleaned_upper = cleaned.to_uppercase();
let forbidden_patterns = [
("REDIS.CALL('FLUSHALL')", "FLUSHALL"),
("REDIS.CALL(\"FLUSHALL\")", "FLUSHALL"),
("REDIS.PCALL('FLUSHALL')", "FLUSHALL via PCALL"),
("REDIS.PCALL(\"FLUSHALL\")", "FLUSHALL via PCALL"),
("REDIS.CALL('FLUSHDB')", "FLUSHDB"),
("REDIS.CALL(\"FLUSHDB\")", "FLUSHDB"),
("REDIS.PCALL('FLUSHDB')", "FLUSHDB via PCALL"),
("REDIS.PCALL(\"FLUSHDB\")", "FLUSHDB via PCALL"),
("REDIS.CALL('KEYS'", "KEYS"),
("REDIS.CALL(\"KEYS\"", "KEYS"),
("REDIS.PCALL('KEYS'", "KEYS via PCALL"),
("REDIS.PCALL(\"KEYS\"", "KEYS via PCALL"),
("REDIS.CALL('SHUTDOWN')", "SHUTDOWN"),
("REDIS.CALL(\"SHUTDOWN\")", "SHUTDOWN"),
("REDIS.CALL('CONFIG'", "CONFIG"),
("REDIS.CALL(\"CONFIG\"", "CONFIG"),
("REDIS.CALL('DEBUG'", "DEBUG"),
("REDIS.CALL(\"DEBUG\"", "DEBUG"),
("REDIS.CALL('SAVE')", "SAVE"),
("REDIS.CALL(\"SAVE\")", "SAVE"),
("REDIS.CALL('BGSAVE')", "BGSAVE"),
("REDIS.CALL(\"BGSAVE\")", "BGSAVE"),
("REDIS.CALL('MONITOR')", "MONITOR"),
("REDIS.CALL(\"MONITOR\")", "MONITOR"),
("OS.EXECUTE", "os.execute"),
("OS.EXEC", "os.exec"),
("IO.POPEN", "io.popen"),
("LOADSTRING", "loadstring"),
("LOAD(", "load()"),
];
for (pattern, description) in &forbidden_patterns {
if cleaned_upper.contains(pattern) {
return Err(OxCacheError::InvalidInput(format!(
"Lua script contains forbidden pattern: {}",
description
)));
}
}
if cleaned_upper.contains("REDIS.EVAL") || cleaned_upper.contains("REDIS.EVALSHA") {
return Err(OxCacheError::InvalidInput(
"Lua script contains nested redis.eval/evalsha".to_string(),
));
}
for re in LUA_LOOP_REGEXES.iter() {
if re.is_match(&cleaned_upper) {
return Err(OxCacheError::InvalidInput(
"Lua script contains potential infinite loop patterns".to_string(),
));
}
}
Ok(())
}
#[cfg(feature = "redis")]
pub(super) fn preprocess_lua_script(script: &str) -> String {
let mut result = String::with_capacity(script.len());
let mut chars = script.chars().peekable();
while let Some(c) = chars.next() {
if c == '-' && chars.peek() == Some(&'-') {
chars.next();
skip_lua_comment(&mut chars);
} else if c == '[' {
if !try_skip_long_string(&mut chars) {
result.push('[');
}
} else if c == '"' || c == '\'' {
result.push(c);
scan_quoted_string(&mut chars, &mut result, c);
} else if c.is_whitespace() {
if !result.is_empty() && !result.ends_with(' ') {
result.push(' ');
}
} else {
result.push(c);
}
}
WHITESPACE_REGEX.replace_all(&result, " ").to_string()
}
#[cfg(feature = "redis")]
fn skip_lua_comment(chars: &mut std::iter::Peekable<std::str::Chars>) {
let level = count_lua_long_string_level(chars, 1);
if level > 0 {
skip_lua_long_string(chars, level);
} else {
while let Some(&next_c) = chars.peek() {
if next_c == '\n' {
break;
}
chars.next();
}
}
}
#[cfg(feature = "redis")]
fn try_skip_long_string(chars: &mut std::iter::Peekable<std::str::Chars>) -> bool {
let level = count_lua_long_string_level(chars, 0);
if level > 0 {
skip_lua_long_string(chars, level);
true
} else {
false
}
}
#[cfg(feature = "redis")]
fn scan_quoted_string(chars: &mut std::iter::Peekable<std::str::Chars>, result: &mut String, quote: char) {
while let Some(&next_c) = chars.peek() {
if next_c == quote {
chars.next();
result.push(quote);
break;
} else if next_c == '\\' {
chars.next();
if let Some(escaped) = chars.next() {
if escaped.is_alphanumeric() || escaped == '_' {
result.push(escaped);
}
}
} else if next_c == '\n' {
break; } else if next_c.is_alphanumeric() || next_c == '_' {
result.push(next_c);
chars.next();
} else {
chars.next();
}
}
}
#[cfg(feature = "redis")]
pub(super) fn count_lua_long_string_level(
chars: &mut std::iter::Peekable<std::str::Chars>,
start_level: usize,
) -> usize {
let mut level = start_level;
while let Some(&c) = chars.peek() {
if c == '=' {
level += 1;
chars.next();
} else if c == '[' {
chars.next();
return level; } else {
break; }
}
0 }
#[cfg(feature = "redis")]
pub(super) fn skip_lua_long_string(chars: &mut std::iter::Peekable<std::str::Chars>, level: usize) {
let closing: String = format!("]{}{}]", "=".repeat(level - 1), "=".repeat(level - 1));
let closing_chars: Vec<char> = closing.chars().collect();
let mut pos = 0;
let closing_len = closing.len();
while let Some(c) = chars.next() {
if c == ']' {
let mut check_pos = 1;
let mut is_match = true;
while check_pos < closing_len {
if let Some(&next_c) = chars.peek() {
if next_c == closing_chars[check_pos] {
chars.next();
check_pos += 1;
} else {
is_match = false;
break;
}
} else {
is_match = false;
break;
}
}
if is_match && check_pos == closing_len {
break; }
}
pos += 1;
if pos > 1_000_000 {
break; }
}
}
#[cfg(feature = "redis")]
#[cfg_attr(docsrs, doc(cfg(feature = "security")))]
pub fn validate_scan_pattern(pattern: &str) -> OxCacheResult<()> {
if pattern.len() > MAX_SCAN_PATTERN_LENGTH {
return Err(OxCacheError::InvalidInput(format!(
"SCAN pattern exceeds maximum length of {} characters (got {} characters)",
MAX_SCAN_PATTERN_LENGTH,
pattern.len()
)));
}
let wildcard_count = pattern.chars().filter(|c| *c == '*').count();
if wildcard_count > MAX_SCAN_WILDCARDS {
return Err(OxCacheError::InvalidInput(format!(
"SCAN pattern contains too many wildcards (max {}, got {})",
MAX_SCAN_WILDCARDS, wildcard_count
)));
}
Ok(())
}
#[cfg(feature = "redis")]
#[cfg_attr(docsrs, doc(cfg(feature = "security")))]
pub fn clamp_scan_count(count: usize) -> usize {
count.clamp(SCAN_COUNT_MIN, SCAN_COUNT_MAX)
}