Implement M3.5.7: Rate limiting + idempotency (20 tests)

This commit is contained in:
Story Crater Bot
2026-08-26 13:35:50 -07:00
parent 1b0bc29027
commit f05565edd0
8 changed files with 1083 additions and 125 deletions
+98 -11
View File
@@ -10,6 +10,8 @@ use std::time::Instant;
use crate::endpoints::IngestRequest;
use crate::ingest_worker::IngestWorker;
use crate::query_worker::QueryWorker;
use crate::rate_limiter::{RateLimiter, LimitConfig};
use crate::idempotency::IdempotencyStore;
/// Server state with database and workers
pub struct AppState {
@@ -20,6 +22,8 @@ pub struct AppState {
pub embeddings: Arc<EmbeddingsClient>,
pub ingest_worker: Arc<IngestWorker>,
pub query_worker: Arc<QueryWorker>,
pub rate_limiter: Arc<RateLimiter>,
pub idempotency_store: Arc<IdempotencyStore>,
}
/// Auth extractor — validates apikey header
@@ -36,6 +40,34 @@ fn check_auth(req: &HttpRequest, state: &AppState) -> Result<(), HttpResponse> {
Ok(())
}
/// Extract apikey from request
fn extract_apikey(req: &HttpRequest) -> Option<String> {
req.headers()
.get("apikey")
.and_then(|h| h.to_str().ok())
.map(|s| s.to_string())
}
/// Rate limit guard — call this in handlers to check rate limit
fn check_rate_limit(req: &HttpRequest, state: &AppState, endpoint: &str) -> Result<(), HttpResponse> {
let apikey = extract_apikey(req).unwrap_or_else(|| "unknown".to_string());
match state.rate_limiter.check(&apikey, endpoint) {
Ok(_) => Ok(()),
Err(rate_limit_err) => {
let retry_after = rate_limit_err.retry_after_seconds.to_string();
Err(HttpResponse::TooManyRequests()
.insert_header(("Retry-After", retry_after))
.json(json!({
"error": "rate_limit_exceeded",
"reason": rate_limit_err.reason.clone(),
"retry_after_seconds": rate_limit_err.retry_after_seconds,
"limit_window": format!("{}s", rate_limit_err.limit_window_secs),
})))
}
}
}
/// Start HTTP server with database initialization
pub async fn start_server(port: u16, api_key: String, database_url: &str) -> Result<()> {
// Create connection pool
@@ -53,6 +85,33 @@ pub async fn start_server(port: u16, api_key: String, database_url: &str) -> Res
let reranker = RerankClient::from_env()?;
let query_worker = Arc::new(QueryWorker::new(VectorStore::new(pool.clone()), (*embeddings).clone(), reranker));
// Initialize rate limiter and idempotency store
let limit_config = LimitConfig {
ingest_per_hour: std::env::var("MEM_RATE_LIMIT_INGEST")
.ok()
.and_then(|s| s.parse().ok())
.unwrap_or(100.0),
query_per_hour: std::env::var("MEM_RATE_LIMIT_QUERY")
.ok()
.and_then(|s| s.parse().ok())
.unwrap_or(1000.0),
projects_per_hour: std::env::var("MEM_RATE_LIMIT_PROJECTS")
.ok()
.and_then(|s| s.parse().ok())
.unwrap_or(100.0),
burst_per_second: std::env::var("MEM_RATE_LIMIT_BURST")
.ok()
.and_then(|s| s.parse().ok())
.unwrap_or(10.0),
};
let rate_limiter = Arc::new(RateLimiter::new(limit_config));
let idempotency_ttl = std::env::var("MEM_IDEMPOTENCY_TTL_SECS")
.ok()
.and_then(|s| s.parse().ok())
.unwrap_or(86400); // 24 hours default
let idempotency_store = Arc::new(IdempotencyStore::new(idempotency_ttl));
let state = web::Data::new(AppState {
api_key,
start_time: Instant::now(),
@@ -61,6 +120,8 @@ pub async fn start_server(port: u16, api_key: String, database_url: &str) -> Res
embeddings,
ingest_worker,
query_worker,
rate_limiter,
idempotency_store,
});
tracing::info!("Starting HTTP server on port {}", port);
@@ -103,6 +164,10 @@ pub async fn ingest_handler(
return e;
}
if let Err(e) = check_rate_limit(&req, &state, "/memory/ingest") {
return e;
}
let project = body.project.clone();
let ingest_id = body.ingest_id.clone();
let records: Vec<(String, String)> = body
@@ -111,6 +176,12 @@ pub async fn ingest_handler(
.map(|r| (r.text.clone(), body.source.clone()))
.collect();
// Check idempotency cache first
if let Some(cached_response) = state.idempotency_store.get(&ingest_id) {
tracing::info!("Returning cached response for ingest_id: {}", ingest_id);
return HttpResponse::Accepted().json(cached_response);
}
// Create ingest job in DB
let job_result = sqlx::query(
"INSERT INTO ingest_jobs (id, project, ingest_id, status, created_at)
@@ -136,18 +207,26 @@ pub async fn ingest_handler(
}
});
HttpResponse::Accepted().json(json!({
let response = json!({
"ingest_id": ingest_id,
"status": "pending",
"status_url": format!("/memory/ingest/{}", ingest_id)
}))
});
// Cache the response for idempotency
state.idempotency_store.set(ingest_id.clone(), response.clone());
HttpResponse::Accepted().json(response)
}
Ok(None) => {
// Already exists
HttpResponse::Conflict().json(json!({
"error": "already_ingesting",
"ingest_id": ingest_id
}))
// Already exists in DB (was inserted concurrently)
let response = json!({
"ingest_id": ingest_id,
"status": "pending",
"status_url": format!("/memory/ingest/{}", ingest_id)
});
state.idempotency_store.set(ingest_id.clone(), response.clone());
HttpResponse::Accepted().json(response)
}
Err(e) => {
tracing::error!("DB error: {}", e);
@@ -203,6 +282,10 @@ pub async fn query_handler(
return e;
}
if let Err(e) = check_rate_limit(&req, &state, "/memory/query") {
return e;
}
let project = match query.get("project") {
Some(p) => p.clone(),
None => {
@@ -246,6 +329,10 @@ pub async fn projects_handler(
return e;
}
if let Err(e) = check_rate_limit(&req, &state, "/memory/projects") {
return e;
}
let result = sqlx::query_as::<_, (String,)>(
"SELECT DISTINCT project FROM memories_l2 ORDER BY project",
)
@@ -544,12 +631,12 @@ pub async fn vault_file_handler(
let (frontmatter, body) = if content.starts_with("---") {
let parts: Vec<&str> = content.split("---").collect();
if parts.len() >= 3 {
(parts[1], parts[2..].join("---"))
(parts[1].to_string(), parts[2..].join("---"))
} else {
("", &content[..])
("".to_string(), content.clone())
}
} else {
("", &content[..])
("".to_string(), content.clone())
};
let html = format!(
@@ -590,7 +677,7 @@ pub async fn vault_file_handler(
} else {
String::new()
},
body.replace("&", "&amp;")
body.clone().replace("&", "&amp;")
.replace("<", "&lt;")
.replace(">", "&gt;")
.lines()
+129
View File
@@ -0,0 +1,129 @@
use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant};
#[cfg(test)]
use serde_json::json;
/// Cached ingest response with expiry
#[derive(Clone, Debug)]
struct CachedResponse {
response: serde_json::Value,
inserted_at: Instant,
ttl: Duration,
}
impl CachedResponse {
fn is_expired(&self) -> bool {
self.inserted_at.elapsed() > self.ttl
}
}
/// Idempotency store for ingest operations
pub struct IdempotencyStore {
cache: Arc<Mutex<HashMap<String, CachedResponse>>>,
ttl: Duration,
}
impl IdempotencyStore {
pub fn new(ttl_seconds: u64) -> Self {
Self {
cache: Arc::new(Mutex::new(HashMap::new())),
ttl: Duration::from_secs(ttl_seconds),
}
}
/// Get cached response for ingest_id. Returns None if not found or expired.
pub fn get(&self, ingest_id: &str) -> Option<serde_json::Value> {
let mut cache = self.cache.lock().unwrap();
if let Some(cached) = cache.get(ingest_id) {
if !cached.is_expired() {
return Some(cached.response.clone());
}
}
// Clean up expired entry
cache.remove(ingest_id);
None
}
/// Store response for ingest_id
pub fn set(&self, ingest_id: String, response: serde_json::Value) {
let mut cache = self.cache.lock().unwrap();
cache.insert(
ingest_id,
CachedResponse {
response,
inserted_at: Instant::now(),
ttl: self.ttl,
},
);
}
/// Evict expired entries (background maintenance)
pub fn evict_expired(&self) {
let mut cache = self.cache.lock().unwrap();
cache.retain(|_, v| !v.is_expired());
}
/// Clear all entries (for testing)
#[cfg(test)]
pub fn clear(&self) {
let mut cache = self.cache.lock().unwrap();
cache.clear();
}
/// Get cache size (for testing)
#[cfg(test)]
pub fn len(&self) -> usize {
let cache = self.cache.lock().unwrap();
cache.len()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_idempotency_store_basic() {
let store = IdempotencyStore::new(60);
let response = json!({"ingest_id": "test-123", "status": "pending"});
store.set("test-123".to_string(), response.clone());
assert_eq!(store.get("test-123"), Some(response));
}
#[test]
fn test_idempotency_store_expiry() {
let store = IdempotencyStore::new(0);
let response = json!({"ingest_id": "test-123", "status": "pending"});
store.set("test-123".to_string(), response);
std::thread::sleep(Duration::from_millis(10));
assert_eq!(store.get("test-123"), None);
}
#[test]
fn test_idempotency_missing_key() {
let store = IdempotencyStore::new(60);
assert_eq!(store.get("nonexistent"), None);
}
#[test]
fn test_idempotency_evict_expired() {
let store = IdempotencyStore::new(1);
store.set("key1".to_string(), json!({"data": "value1"}));
store.set("key2".to_string(), json!({"data": "value2"}));
assert_eq!(store.len(), 2);
std::thread::sleep(Duration::from_secs(1));
std::thread::sleep(Duration::from_millis(100));
store.evict_expired();
assert_eq!(store.len(), 0);
}
}
+2
View File
@@ -2,6 +2,8 @@ pub mod endpoints;
pub mod http_server;
pub mod ingest_worker;
pub mod query_worker;
pub mod rate_limiter;
pub mod idempotency;
pub use endpoints::{IngestQueue, IngestRequest, JobStatus};
pub use ingest_worker::IngestWorker;
+2
View File
@@ -3,6 +3,8 @@ mod http_server;
mod endpoints;
mod ingest_worker;
mod query_worker;
mod rate_limiter;
mod idempotency;
use clap::{Parser, Subcommand};
use mem_chunk::token_counter::CharsOverFourCounter;
+243
View File
@@ -0,0 +1,243 @@
use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use std::time::Instant;
/// Rate limit error with retry guidance
#[derive(Debug, Clone)]
pub struct RateLimitError {
pub retry_after_seconds: u64,
pub limit_window_secs: u64,
pub reason: String,
}
impl RateLimitError {
pub fn reason(&self) -> String {
format!(
"{} (retry after {} seconds, window: {} seconds)",
self.reason, self.retry_after_seconds, self.limit_window_secs
)
}
}
/// Token bucket for a single endpoint
#[derive(Debug, Clone)]
struct TokenBucket {
tokens: f64,
last_refill: Instant,
capacity: f64, // max tokens (per hour)
refill_rate: f64, // tokens per second
}
impl TokenBucket {
fn new(capacity: f64, refill_rate: f64) -> Self {
Self {
tokens: capacity,
last_refill: Instant::now(),
capacity,
refill_rate,
}
}
/// Refill tokens based on elapsed time
fn refill(&mut self) {
let now = Instant::now();
let elapsed = now.duration_since(self.last_refill).as_secs_f64();
let refilled = elapsed * self.refill_rate;
self.tokens = (self.tokens + refilled).min(self.capacity);
self.last_refill = now;
}
/// Try to consume 1 token. Returns Ok if successful, Err(retry_after_secs) if rate limited.
fn try_consume(&mut self) -> Result<(), u64> {
self.refill();
if self.tokens >= 1.0 {
self.tokens -= 1.0;
return Ok(());
}
// Rate limited: estimate time until next token available
let tokens_needed = 1.0 - self.tokens;
let retry_after = (tokens_needed / self.refill_rate).ceil() as u64;
Err(retry_after.max(1))
}
}
/// Rate limiter with per-apikey, per-endpoint buckets
pub struct RateLimiter {
buckets: Arc<Mutex<HashMap<String, Arc<Mutex<TokenBucket>>>>>,
limit_config: LimitConfig,
}
#[derive(Clone, Debug)]
pub struct LimitConfig {
pub ingest_per_hour: f64,
pub query_per_hour: f64,
pub projects_per_hour: f64,
pub burst_per_second: f64, // Currently unused but kept for API compatibility
}
impl Default for LimitConfig {
fn default() -> Self {
Self {
ingest_per_hour: 100.0,
query_per_hour: 1000.0,
projects_per_hour: 100.0,
burst_per_second: 10.0,
}
}
}
impl RateLimiter {
pub fn new(config: LimitConfig) -> Self {
Self {
buckets: Arc::new(Mutex::new(HashMap::new())),
limit_config: config,
}
}
/// Get or create bucket for apikey + endpoint
fn get_or_create_bucket(&self, apikey_endpoint: &str) -> Arc<Mutex<TokenBucket>> {
let mut buckets = self.buckets.lock().unwrap();
let config = &self.limit_config;
if !buckets.contains_key(apikey_endpoint) {
// Determine limit based on endpoint
let capacity = if apikey_endpoint.contains("/memory/ingest") {
config.ingest_per_hour
} else if apikey_endpoint.contains("/memory/query") {
config.query_per_hour
} else if apikey_endpoint.contains("/memory/projects") {
config.projects_per_hour
} else {
// Unlimited for unknown endpoints
f64::INFINITY
};
let refill_rate = if capacity.is_infinite() {
f64::INFINITY
} else {
capacity / 3600.0 // per second
};
let bucket = TokenBucket::new(capacity, refill_rate);
buckets.insert(apikey_endpoint.to_string(), Arc::new(Mutex::new(bucket)));
}
buckets[apikey_endpoint].clone()
}
/// Check rate limit for apikey + endpoint. Returns Ok or Err with retry guidance.
pub fn check(&self, apikey: &str, endpoint: &str) -> Result<(), RateLimitError> {
let key = format!("{}::{}", apikey, endpoint);
let bucket = self.get_or_create_bucket(&key);
let mut b = bucket.lock().unwrap();
match b.try_consume() {
Ok(_) => Ok(()),
Err(retry_after) => {
let window_secs = if endpoint.contains("/memory/ingest") {
3600
} else if endpoint.contains("/memory/query") {
3600
} else if endpoint.contains("/memory/projects") {
3600
} else {
3600
};
Err(RateLimitError {
retry_after_seconds: retry_after,
limit_window_secs: window_secs,
reason: format!(
"rate_limit_exceeded for {}",
endpoint
),
})
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_token_bucket_refill() {
let mut bucket = TokenBucket::new(100.0, 100.0 / 3600.0);
assert!(bucket.try_consume().is_ok());
// After one consumption, should have 99 tokens
assert_eq!((bucket.tokens * 1.0) as i64, 99);
}
#[test]
fn test_rate_limit_within_capacity() {
let config = LimitConfig {
ingest_per_hour: 5.0,
query_per_hour: 10.0,
projects_per_hour: 10.0,
burst_per_second: 10.0,
};
let limiter = RateLimiter::new(config);
// First 5 should succeed
for _ in 0..5 {
assert!(limiter.check("apikey1", "/memory/ingest").is_ok());
}
// 6th should fail
let err = limiter.check("apikey1", "/memory/ingest");
assert!(err.is_err());
if let Err(e) = err {
assert!(e.retry_after_seconds > 0);
}
}
#[test]
fn test_per_apikey_isolation() {
let config = LimitConfig {
ingest_per_hour: 5.0,
query_per_hour: 10.0,
projects_per_hour: 10.0,
burst_per_second: 10.0,
};
let limiter = RateLimiter::new(config);
// apikey1 uses up 5 ingest requests
for _ in 0..5 {
assert!(limiter.check("apikey1", "/memory/ingest").is_ok());
}
assert!(limiter.check("apikey1", "/memory/ingest").is_err());
// apikey2 should have its own 5
for _ in 0..5 {
assert!(limiter.check("apikey2", "/memory/ingest").is_ok());
}
assert!(limiter.check("apikey2", "/memory/ingest").is_err());
}
#[test]
fn test_per_endpoint_isolation() {
let config = LimitConfig {
ingest_per_hour: 5.0,
query_per_hour: 10.0,
projects_per_hour: 10.0,
burst_per_second: 10.0,
};
let limiter = RateLimiter::new(config);
// Use up 5 ingest
for _ in 0..5 {
assert!(limiter.check("apikey1", "/memory/ingest").is_ok());
}
assert!(limiter.check("apikey1", "/memory/ingest").is_err());
// Query should have separate 10 limit
for _ in 0..10 {
assert!(limiter.check("apikey1", "/memory/query").is_ok());
}
assert!(limiter.check("apikey1", "/memory/query").is_err());
}
}