Implement M3.5.7: Rate limiting + idempotency (20 tests)
This commit is contained in:
@@ -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("&", "&")
|
||||
body.clone().replace("&", "&")
|
||||
.replace("<", "<")
|
||||
.replace(">", ">")
|
||||
.lines()
|
||||
|
||||
@@ -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,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;
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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());
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user