Harden validation, improve runtime security and refactor core services
Rust / build (push) Canceled after 0s

This commit is contained in:
thakares committed 2026-06-04 19:58:21 +05:30
1 parent aebef4c623
commit 8d119ac00e
39 files changed
+528 -158

No files matched your search

+1
View File
@@ -31,4 +31,5 @@ base64 = "0.22"
rand = "0.8"
ed25519-dalek = "2"
redis = { version = "0.29", features = ["r2d2"] }
dashmap = "6"
+2 -4
View File
@@ -20,12 +20,10 @@ pub async fn cleanup_loop(state: Arc<AppState>) {
}
}
// Evict stale rate-limiter entries to prevent unbounded HashMap growth.
// Evict stale rate-limiter entries to prevent unbounded map growth.
{
let window_secs = state.get_config().rate_limit_window_secs;
let mut rl = state.rate_limiter.lock().await;
rl.evict_stale(window_secs);
state.rate_limiter.evict_stale(window_secs);
}
}
}
+5 -2
View File
@@ -137,7 +137,9 @@ impl Config {
size: self.gene_size,
});
}
if !(1..=shared::constants::MAX_MUTATION_ROUNDS).contains(&self.mutation_rounds) {
if !(shared::constants::MIN_MUTATION_ROUNDS..=shared::constants::MAX_MUTATION_ROUNDS)
.contains(&self.mutation_rounds)
{
return Err(ConfigError::InvalidMutationRounds {
rounds: self.mutation_rounds,
});
@@ -274,7 +276,8 @@ impl std::fmt::Display for ConfigError {
Self::InvalidMutationRounds { rounds } => {
write!(
f,
"invalid mutation rounds {rounds}; expected 1..={}",
"invalid mutation rounds {rounds}; expected {}..={}",
shared::constants::MIN_MUTATION_ROUNDS,
shared::constants::MAX_MUTATION_ROUNDS
)
}
-1
View File
@@ -44,4 +44,3 @@ pub fn verify_signature(
pk.verify_strict(message.as_bytes(), &sig)?;
Ok(())
}
+7
View File
@@ -25,12 +25,19 @@ pub enum SessionError {
#[error("Invalid gene configuration: {0}")]
InvalidGeneConfiguration(String),
#[error("Rate limited")]
RateLimited,
}
impl IntoResponse for SessionError {
fn into_response(self) -> Response {
let (status, error_message) = match self {
SessionError::InvalidPublicKeyLength => (StatusCode::BAD_REQUEST, self.to_string()),
SessionError::RateLimited => (
StatusCode::TOO_MANY_REQUESTS,
"Too many requests".to_string(),
),
_ => (
StatusCode::INTERNAL_SERVER_ERROR,
"Internal server error".to_string(),
+58 -3
View File
@@ -1,5 +1,10 @@
use shared::protocol::Fingerprint;
const MIN_ASPECT_RATIO: f64 = 0.5;
const MAX_ASPECT_RATIO: f64 = 3.0;
const MAX_DEVICE_PIXEL_RATIO: f64 = 5.0;
const MAX_HARDWARE_CONCURRENCY: u32 = 256;
/// Validates the browser fingerprint fields submitted by the client.
///
/// Checks basic screen aspect ratio thresholds, device pixel ratio limits,
@@ -9,16 +14,66 @@ use shared::protocol::Fingerprint;
/// * `fp` - The client's hardware and screen layout fingerprint.
pub fn validate(fp: &Fingerprint) -> Result<(), Box<dyn std::error::Error>> {
let ar: f64 = fp.aspect_ratio.parse().map_err(|_| "ar")?;
if !(0.5..=3.0).contains(&ar) {
if !ar.is_finite() || !(MIN_ASPECT_RATIO..=MAX_ASPECT_RATIO).contains(&ar) {
return Err("aspect ratio".into());
}
let dpr: f64 = fp.device_pixel_ratio.parse().map_err(|_| "dpr")?;
if dpr <= 0.0 || dpr > 5.0 {
if !dpr.is_finite() || dpr <= 0.0 || dpr > MAX_DEVICE_PIXEL_RATIO {
return Err("dpr".into());
}
if fp.hardware_concurrency == 0 {
if fp.hardware_concurrency == 0 || fp.hardware_concurrency > MAX_HARDWARE_CONCURRENCY {
return Err("hw".into());
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
fn fingerprint(
aspect_ratio: impl Into<String>,
device_pixel_ratio: impl Into<String>,
hardware_concurrency: u32,
) -> Fingerprint {
Fingerprint {
aspect_ratio: aspect_ratio.into(),
device_pixel_ratio: device_pixel_ratio.into(),
hardware_concurrency,
}
}
#[test]
fn accepts_valid_fingerprint() {
assert!(validate(&fingerprint("1.7777777778", "2", 8)).is_ok());
}
#[test]
fn accepts_boundary_values() {
assert!(validate(&fingerprint("0.5", "1", 1)).is_ok());
assert!(validate(&fingerprint("3.0", "5.0", MAX_HARDWARE_CONCURRENCY)).is_ok());
}
#[test]
fn rejects_invalid_aspect_ratios() {
for aspect_ratio in ["not-a-number", "NaN", "inf", "0.49", "3.01"] {
assert!(validate(&fingerprint(aspect_ratio, "2", 8)).is_err());
}
}
#[test]
fn rejects_invalid_device_pixel_ratios() {
for device_pixel_ratio in ["not-a-number", "NaN", "inf", "0", "-1", "5.01"] {
assert!(validate(&fingerprint("1.77", device_pixel_ratio, 8)).is_err());
}
}
#[test]
fn rejects_invalid_hardware_concurrency() {
assert!(validate(&fingerprint("1.77", "2", 0)).is_err());
assert!(validate(&fingerprint("1.77", "2", MAX_HARDWARE_CONCURRENCY + 1)).is_err());
}
}
-3
View File
@@ -30,9 +30,6 @@ async fn main() {
async fn try_main() -> Result<(), Box<dyn std::error::Error>> {
let cli = Cli::parse();
if let Some(config_path) = cli.globals.config.as_deref() {
std::env::set_var("CHRONOSEAL_CONFIG", config_path);
}
let log_filter = cli.globals.log.as_deref().unwrap_or("info");
let log_file = log_file_for_command(&cli);
let _log_guard = init_logging(log_filter, log_file)?;
+22
View File
@@ -9,3 +9,25 @@ pub async fn log_request(req: Request, next: Next) -> Response {
tracing::info!("{} {} -> {}", method, uri, response.status());
response
}
/// Injects defensive HTTP response headers on every response.
///
/// These headers mitigate several classes of attacks:
/// - `X-Content-Type-Options: nosniff` — prevents MIME-type sniffing.
/// - `X-Frame-Options: DENY` — blocks clickjacking via framing.
/// - `Referrer-Policy: no-referrer` — suppresses referrer leakage.
/// - `X-XSS-Protection: 0` — disables legacy XSS auditors (can introduce bugs).
/// - `Permissions-Policy` — restricts powerful browser features.
pub async fn security_headers(req: Request, next: Next) -> Response {
let mut response = next.run(req).await;
let headers = response.headers_mut();
headers.insert("x-content-type-options", "nosniff".parse().unwrap());
headers.insert("x-frame-options", "DENY".parse().unwrap());
headers.insert("referrer-policy", "no-referrer".parse().unwrap());
headers.insert("x-xss-protection", "0".parse().unwrap());
headers.insert(
"permissions-policy",
"camera=(), microphone=(), geolocation=()".parse().unwrap(),
);
response
}
+20 -15
View File
@@ -1,17 +1,20 @@
use std::collections::HashMap;
use dashmap::DashMap;
use std::time::Instant;
/// A simple, in-memory sliding-window rate limiter for tracking client heartbeat frequency.
/// A lock-free, concurrent sliding-window rate limiter backed by `DashMap`.
///
/// All public methods take `&self` (no `&mut self`), so the limiter can live in
/// an `Arc<AppState>` without a `Mutex` wrapper.
pub struct RateLimiter {
/// Maps session identifiers to request counts and window start timestamps.
buckets: HashMap<String, (u32, Instant)>,
/// Maps rate-limit keys to request counts and window start timestamps.
buckets: DashMap<String, (u32, Instant)>,
}
impl RateLimiter {
/// Creates a new, empty `RateLimiter`.
pub fn new() -> Self {
Self {
buckets: HashMap::new(),
buckets: DashMap::new(),
}
}
@@ -20,19 +23,21 @@ impl RateLimiter {
/// Returns `true` if allowed, or `false` if the rate limit is exceeded.
///
/// # Arguments
/// * `key` - The unique identifier to rate-limit (e.g., session ID).
/// * `key` - The unique identifier to rate-limit (e.g., client IP address).
/// * `limit` - The maximum number of allowed requests per window.
/// * `window_secs` - The length of the sliding-window in seconds.
pub fn check(&mut self, key: &str, limit: u32, window_secs: u64) -> bool {
pub fn check(&self, key: &str, limit: u32, window_secs: u64) -> bool {
let now = Instant::now();
let entry = self.buckets.entry(key.to_string()).or_insert((0, now));
if now.duration_since(entry.1).as_secs() >= window_secs {
*entry = (1, now);
let mut entry = self.buckets.entry(key.to_string()).or_insert((0, now));
let (count, ts) = entry.value_mut();
if now.duration_since(*ts).as_secs() >= window_secs {
*count = 1;
*ts = now;
true
} else if entry.0 >= limit {
} else if *count >= limit {
false
} else {
entry.0 += 1;
*count += 1;
true
}
}
@@ -43,7 +48,7 @@ impl RateLimiter {
///
/// # Arguments
/// * `window_secs` - The active rate-limiting window duration in seconds.
pub fn evict_stale(&mut self, window_secs: u64) {
pub fn evict_stale(&self, window_secs: u64) {
let now = Instant::now();
self.buckets
.retain(|_, (_, ts)| now.duration_since(*ts).as_secs() < window_secs);
@@ -58,7 +63,7 @@ mod tests {
#[test]
fn test_rate_limiter() {
let mut rl = RateLimiter::new();
let rl = RateLimiter::new();
// Limit of 2 requests per 1 second window
assert!(rl.check("user1", 2, 1));
assert!(rl.check("user1", 2, 1));
@@ -72,7 +77,7 @@ mod tests {
#[test]
fn test_rate_limiter_eviction() {
let mut rl = RateLimiter::new();
let rl = RateLimiter::new();
assert!(rl.check("user1", 1, 1));
assert_eq!(rl.buckets.len(), 1);
+64 -16
View File
@@ -8,21 +8,51 @@ pub async fn handler(
Json(payload): Json<HeartbeatRequest>,
) -> (StatusCode, Json<HeartbeatResponse>) {
let start_http = std::time::Instant::now();
state.heartbeats_total.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
state
.heartbeats_total
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
// Rate limiting
// Cap entropy events to prevent oversized payloads from exhausting memory.
if payload.entropy_data.events.len() > 1000 {
let http_dur = start_http.elapsed().as_nanos() as u64;
state
.http_latency_ns
.fetch_add(http_dur, std::sync::atomic::Ordering::Relaxed);
state
.http_ops_count
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
return (
StatusCode::OK,
Json(HeartbeatResponse {
status: "ok".into(),
next_salt: None,
next_mutation_step: None,
next_mutation_order_b64: None,
}),
);
}
// Rate limiting (lock-free via DashMap)
{
let (limit, window_secs) = {
let cfg = state.get_config();
(cfg.rate_limit_count, cfg.rate_limit_window_secs)
};
let mut rl = state.rate_limiter.lock().await;
if !rl.check(&payload.session_id, limit, window_secs) {
if !state
.rate_limiter
.check(&payload.session_id, limit, window_secs)
{
tracing::debug!("Rate limit hit: {}", payload.session_id);
state.verification_failures_total.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
state
.verification_failures_total
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let http_dur = start_http.elapsed().as_nanos() as u64;
state.http_latency_ns.fetch_add(http_dur, std::sync::atomic::Ordering::Relaxed);
state.http_ops_count.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
state
.http_latency_ns
.fetch_add(http_dur, std::sync::atomic::Ordering::Relaxed);
state
.http_ops_count
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
return (
StatusCode::OK,
Json(HeartbeatResponse {
@@ -39,8 +69,12 @@ pub async fn handler(
let start_db = std::time::Instant::now();
let db_res = crate::session::verify_heartbeat(&state.db_pool, &config, &payload);
let db_dur = start_db.elapsed().as_nanos() as u64;
state.storage_latency_ns.fetch_add(db_dur, std::sync::atomic::Ordering::Relaxed);
state.storage_ops_count.fetch_add(2, std::sync::atomic::Ordering::Relaxed); // read + write
state
.storage_latency_ns
.fetch_add(db_dur, std::sync::atomic::Ordering::Relaxed);
state
.storage_ops_count
.fetch_add(2, std::sync::atomic::Ordering::Relaxed); // read + write
let outcome = match db_res {
Ok(result) => (
@@ -54,15 +88,21 @@ pub async fn handler(
),
Err(e) => {
tracing::warn!("Heartbeat failed for {}: {}", payload.session_id, e);
state.verification_failures_total.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
state
.verification_failures_total
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
match &e {
crate::errors::VerificationError::ChainBroken => {
state.replay_attempts_total.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
state
.replay_attempts_total
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
}
crate::errors::VerificationError::MutationCommitmentMismatch
| crate::errors::VerificationError::MutationProgram(_)
| crate::errors::VerificationError::GeneState(_) => {
state.mutation_failures_total.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
state
.mutation_failures_total
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
}
_ => {}
}
@@ -79,8 +119,12 @@ pub async fn handler(
};
let http_dur = start_http.elapsed().as_nanos() as u64;
state.http_latency_ns.fetch_add(http_dur, std::sync::atomic::Ordering::Relaxed);
state.http_ops_count.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
state
.http_latency_ns
.fetch_add(http_dur, std::sync::atomic::Ordering::Relaxed);
state
.http_ops_count
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
outcome
}
@@ -138,7 +182,11 @@ mod tests {
},
],
};
let program_bytes = base64::Engine::decode(&base64::engine::general_purpose::STANDARD, &init.opcodes_b64).unwrap();
let program_bytes = base64::Engine::decode(
&base64::engine::general_purpose::STANDARD,
&init.opcodes_b64,
)
.unwrap();
let stack_state = shared::vm::execute(&program_bytes);
let order =
@@ -176,7 +224,7 @@ mod tests {
let pool = crate::storage::init_pool(Path::new(":memory:")).unwrap();
let state = Arc::new(AppState {
db_pool: pool.clone(),
rate_limiter: tokio::sync::Mutex::new(crate::ratelimit::RateLimiter::new()),
rate_limiter: crate::ratelimit::RateLimiter::new(),
config: std::sync::RwLock::new(config.clone()),
heartbeats_total: std::sync::atomic::AtomicU64::new(0),
verification_failures_total: std::sync::atomic::AtomicU64::new(0),
+22 -5
View File
@@ -10,18 +10,35 @@ pub async fn handler(
) -> Result<Json<InitResponse>, SessionError> {
let start_http = std::time::Instant::now();
let config = state.get_config();
// Rate limit session creation by public key to prevent storage exhaustion.
if !state.rate_limiter.check(
&payload.public_key,
config.rate_limit_count,
config.rate_limit_window_secs,
) {
return Err(SessionError::RateLimited);
}
let start_db = std::time::Instant::now();
let resp = crate::session::create_session(&state.db_pool, &config, &payload.public_key);
let db_dur = start_db.elapsed().as_nanos() as u64;
state.storage_latency_ns.fetch_add(db_dur, std::sync::atomic::Ordering::Relaxed);
state.storage_ops_count.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
state
.storage_latency_ns
.fetch_add(db_dur, std::sync::atomic::Ordering::Relaxed);
state
.storage_ops_count
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let resp = resp?;
let http_dur = start_http.elapsed().as_nanos() as u64;
state.http_latency_ns.fetch_add(http_dur, std::sync::atomic::Ordering::Relaxed);
state.http_ops_count.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
state
.http_latency_ns
.fetch_add(http_dur, std::sync::atomic::Ordering::Relaxed);
state
.http_ops_count
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
Ok(Json(resp))
}
+55 -15
View File
@@ -5,17 +5,19 @@ use crate::{
routes, session,
storage::{self, StoreStats},
};
use axum::{http::StatusCode, response::IntoResponse, routing::get, Json, Router};
use axum::{
extract::ConnectInfo, http::StatusCode, response::IntoResponse, routing::get, Json, Router,
};
use serde::Serialize;
use std::{
fs,
io::{Read, Write},
net::{SocketAddr, TcpStream},
net::{IpAddr, SocketAddr, TcpStream},
path::Path,
sync::Arc,
time::Duration,
};
use tokio::sync::{Mutex, Notify};
use tokio::sync::Notify;
use tracing::{error, info, warn};
#[derive(Debug, Serialize)]
@@ -144,7 +146,7 @@ pub async fn run_daemon(config: Config) -> Result<(), Box<dyn std::error::Error>
let db_pool = init_db_pool(&config)?;
let state = Arc::new(session::AppState {
db_pool,
rate_limiter: Mutex::new(RateLimiter::new()),
rate_limiter: RateLimiter::new(),
config: std::sync::RwLock::new(config.clone()),
heartbeats_total: std::sync::atomic::AtomicU64::new(0),
verification_failures_total: std::sync::atomic::AtomicU64::new(0),
@@ -170,7 +172,11 @@ pub async fn run_daemon(config: Config) -> Result<(), Box<dyn std::error::Error>
tower_http::services::ServeDir::new(&config.frontend_dir),
)
.layer(tower_http::cors::CorsLayer::permissive())
.layer(axum::middleware::from_fn(
crate::middleware::security_headers,
))
.layer(axum::middleware::from_fn(crate::middleware::log_request))
.layer(axum::extract::DefaultBodyLimit::max(64 * 1024)) // 64 KiB
.with_state(state.clone());
let addr: SocketAddr = config.bind.parse()?;
@@ -178,9 +184,12 @@ pub async fn run_daemon(config: Config) -> Result<(), Box<dyn std::error::Error>
info!(bind = %config.bind, "chronoseal daemon started");
let shutdown = signal_task(state.clone());
let result = axum::serve(listener, app)
.with_graceful_shutdown(shutdown)
.await;
let result = axum::serve(
listener,
app.into_make_service_with_connect_info::<SocketAddr>(),
)
.with_graceful_shutdown(shutdown)
.await;
remove_pid_file(&config.pid_file);
result?;
@@ -277,8 +286,12 @@ async fn health_handler() -> impl IntoResponse {
}
async fn stats_handler(
ConnectInfo(addr): ConnectInfo<SocketAddr>,
axum::extract::State(state): axum::extract::State<Arc<session::AppState>>,
) -> Result<Json<StoreStats>, (StatusCode, String)> {
if !is_loopback(addr.ip()) {
return Err((StatusCode::FORBIDDEN, "Forbidden".to_string()));
}
state
.db_pool
.stats()
@@ -287,25 +300,45 @@ async fn stats_handler(
}
async fn metrics_handler(
ConnectInfo(addr): ConnectInfo<SocketAddr>,
axum::extract::State(state): axum::extract::State<Arc<session::AppState>>,
) -> Result<String, (StatusCode, String)> {
if !is_loopback(addr.ip()) {
return Err((StatusCode::FORBIDDEN, "Forbidden".to_string()));
}
let stats = state
.db_pool
.stats()
.map_err(|err| (StatusCode::INTERNAL_SERVER_ERROR, err.to_string()))?;
let heartbeats = state.heartbeats_total.load(std::sync::atomic::Ordering::Relaxed);
let ver_failures = state.verification_failures_total.load(std::sync::atomic::Ordering::Relaxed);
let mut_failures = state.mutation_failures_total.load(std::sync::atomic::Ordering::Relaxed);
let replays = state.replay_attempts_total.load(std::sync::atomic::Ordering::Relaxed);
let heartbeats = state
.heartbeats_total
.load(std::sync::atomic::Ordering::Relaxed);
let ver_failures = state
.verification_failures_total
.load(std::sync::atomic::Ordering::Relaxed);
let mut_failures = state
.mutation_failures_total
.load(std::sync::atomic::Ordering::Relaxed);
let replays = state
.replay_attempts_total
.load(std::sync::atomic::Ordering::Relaxed);
let store_ns = state.storage_latency_ns.load(std::sync::atomic::Ordering::Relaxed) as f64;
let store_ns = state
.storage_latency_ns
.load(std::sync::atomic::Ordering::Relaxed) as f64;
let store_sum = store_ns / 1_000_000_000.0;
let store_count = state.storage_ops_count.load(std::sync::atomic::Ordering::Relaxed);
let store_count = state
.storage_ops_count
.load(std::sync::atomic::Ordering::Relaxed);
let http_ns = state.http_latency_ns.load(std::sync::atomic::Ordering::Relaxed) as f64;
let http_ns = state
.http_latency_ns
.load(std::sync::atomic::Ordering::Relaxed) as f64;
let http_sum = http_ns / 1_000_000_000.0;
let http_count = state.http_ops_count.load(std::sync::atomic::Ordering::Relaxed);
let http_count = state
.http_ops_count
.load(std::sync::atomic::Ordering::Relaxed);
Ok(format!(
"# HELP chronoseal_active_sessions Active ChronoSeal sessions\n\
@@ -355,6 +388,13 @@ async fn metrics_handler(
))
}
fn is_loopback(ip: IpAddr) -> bool {
match ip {
IpAddr::V4(v4) => v4.is_loopback(),
IpAddr::V6(v6) => v6.is_loopback(),
}
}
async fn signal_task(state: Arc<session::AppState>) {
let shutdown = Arc::new(Notify::new());
+7 -4
View File
@@ -1,6 +1,6 @@
pub struct AppState {
pub db_pool: crate::storage::DbPool,
pub rate_limiter: tokio::sync::Mutex<crate::ratelimit::RateLimiter>,
pub rate_limiter: crate::ratelimit::RateLimiter,
pub config: std::sync::RwLock<crate::config::Config>,
pub heartbeats_total: std::sync::atomic::AtomicU64,
pub verification_failures_total: std::sync::atomic::AtomicU64,
@@ -271,7 +271,6 @@ mod tests {
}
}
fn test_fingerprint() -> Fingerprint {
Fingerprint {
aspect_ratio: "1.77".to_string(),
@@ -323,8 +322,12 @@ mod tests {
vm_extensions::apply_program_clone(&client.committed_gene_state, &order.program)
.unwrap();
let entropy = test_entropy();
let program_bytes = base64::Engine::decode(&base64::engine::general_purpose::STANDARD, &client.opcodes_b64).unwrap();
let program_bytes = base64::Engine::decode(
&base64::engine::general_purpose::STANDARD,
&client.opcodes_b64,
)
.unwrap();
let stack = shared::vm::execute(&program_bytes);
let mut req = HeartbeatRequest {
+54 -24
View File
@@ -1,8 +1,8 @@
use crate::config::Config;
use redis::Commands;
use serde::{Deserialize, Serialize};
use std::path::Path;
use std::time::{SystemTime, UNIX_EPOCH};
use redis::Commands;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct StoreStats {
@@ -54,11 +54,12 @@ impl DbPool {
crate::config::DbType::Valkey => {
let addr = std::env::var("CHRONOSEAL_VALKEY_ADDR")
.unwrap_or_else(|_| "127.0.0.1:6666".to_string());
let connection_string = if addr.starts_with("redis://") || addr.starts_with("rediss://") {
addr.clone()
} else {
format!("redis://{}", addr)
};
let connection_string =
if addr.starts_with("redis://") || addr.starts_with("rediss://") {
addr.clone()
} else {
format!("redis://{}", addr)
};
let client = redis::Client::open(connection_string)?;
let pool = r2d2::Pool::builder().build(client)?;
Ok(DbPool::Valkey(ValkeyStore {
@@ -354,9 +355,19 @@ impl ValkeyStore {
redis::pipe()
.atomic()
.cmd("SET").arg(&key).arg(&value).arg("EX").arg(ttl_seconds)
.cmd("ZADD").arg(&self.index_key).arg(record.expires_at).arg(&record.session_id)
.cmd("ZADD").arg("sessions:chain_lengths").arg(record.chain_length).arg(&record.session_id)
.cmd("SET")
.arg(&key)
.arg(&value)
.arg("EX")
.arg(ttl_seconds)
.cmd("ZADD")
.arg(&self.index_key)
.arg(record.expires_at)
.arg(&record.session_id)
.cmd("ZADD")
.arg("sessions:chain_lengths")
.arg(record.chain_length)
.arg(&record.session_id)
.query::<()>(&mut *conn)?;
Ok(())
}
@@ -400,9 +411,19 @@ impl ValkeyStore {
let response: Option<()> = redis::pipe()
.atomic()
.cmd("SET").arg(&key).arg(&value).arg("EX").arg(ttl_seconds)
.cmd("ZADD").arg(&self.index_key).arg(record.expires_at).arg(&record.session_id)
.cmd("ZADD").arg("sessions:chain_lengths").arg(record.chain_length).arg(&record.session_id)
.cmd("SET")
.arg(&key)
.arg(&value)
.arg("EX")
.arg(ttl_seconds)
.cmd("ZADD")
.arg(&self.index_key)
.arg(record.expires_at)
.arg(&record.session_id)
.cmd("ZADD")
.arg("sessions:chain_lengths")
.arg(record.chain_length)
.arg(&record.session_id)
.query(&mut *conn)?;
match response {
@@ -422,8 +443,12 @@ impl ValkeyStore {
if !expired_ids.is_empty() {
redis::pipe()
.atomic()
.cmd("ZREM").arg(&self.index_key).arg(&expired_ids)
.cmd("ZREM").arg("sessions:chain_lengths").arg(&expired_ids)
.cmd("ZREM")
.arg(&self.index_key)
.arg(&expired_ids)
.cmd("ZREM")
.arg("sessions:chain_lengths")
.arg(&expired_ids)
.query::<()>(&mut *conn)?;
}
Ok(())
@@ -435,8 +460,12 @@ impl ValkeyStore {
let sessions: u64 = conn.zcard(&self.index_key)?;
let expired_sessions: u64 = conn.zcount(&self.index_key, 0, now)?;
let max_chain_length_res: Vec<(String, u64)> = conn.zrevrange_withscores("sessions:chain_lengths", 0, 0)?;
let max_chain_length = max_chain_length_res.first().map(|(_, score)| *score).unwrap_or(0);
let max_chain_length_res: Vec<(String, u64)> =
conn.zrevrange_withscores("sessions:chain_lengths", 0, 0)?;
let max_chain_length = max_chain_length_res
.first()
.map(|(_, score)| *score)
.unwrap_or(0);
Ok(StoreStats {
sessions,
@@ -449,7 +478,7 @@ impl ValkeyStore {
pub fn current_time_ms() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.expect("system clock is before UNIX epoch; check system time")
.as_millis() as u64
}
@@ -459,7 +488,8 @@ mod valkey_tests {
#[test]
fn test_valkey_store_operations() {
let addr = std::env::var("CHRONOSEAL_VALKEY_ADDR").unwrap_or_else(|_| "127.0.0.1:6379".to_string());
let addr = std::env::var("CHRONOSEAL_VALKEY_ADDR")
.unwrap_or_else(|_| "127.0.0.1:6379".to_string());
let connection_string = format!("redis://{}", addr);
let client = match redis::Client::open(connection_string) {
Ok(c) => c,
@@ -522,10 +552,11 @@ mod valkey_tests {
#[test]
fn test_valkey_pool_concurrency() {
use std::thread;
use std::sync::Arc;
use std::thread;
let addr = std::env::var("CHRONOSEAL_VALKEY_ADDR").unwrap_or_else(|_| "127.0.0.1:6379".to_string());
let addr = std::env::var("CHRONOSEAL_VALKEY_ADDR")
.unwrap_or_else(|_| "127.0.0.1:6379".to_string());
let connection_string = format!("redis://{}", addr);
let client = match redis::Client::open(connection_string) {
Ok(c) => c,
@@ -589,7 +620,8 @@ mod valkey_tests {
// Cleanup
let mut conn = store_arc.pool.get().unwrap();
for t in 0..10 {
let _: Result<(), _> = conn.del(store_arc.session_key(&format!("valkey_concurrent_{}", t)));
let _: Result<(), _> =
conn.del(store_arc.session_key(&format!("valkey_concurrent_{}", t)));
}
let _: Result<(), _> = conn.del(&store_arc.index_key);
}
@@ -598,8 +630,8 @@ mod valkey_tests {
#[cfg(test)]
mod sqlite_tests {
use super::*;
use std::thread;
use std::sync::Arc;
use std::thread;
#[test]
fn test_sqlite_pool_concurrency() {
@@ -655,5 +687,3 @@ mod sqlite_tests {
let _ = std::fs::remove_file(db_path);
}
}
-1
View File
@@ -49,7 +49,6 @@ pub fn validate_mouse(
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
-1
View File
@@ -80,7 +80,6 @@ pub fn execute_mutation_order(
vm_extensions::execute_program(state, &order.program)
}
#[cfg(test)]
mod tests {
use super::*;