Spaces:
Running
Running
feat: implement daily request rate limiting with per-user usage tracking and admin bypass
Browse files- main.py +7 -0
- materials/routes.py +2 -1
- quiz_generator/routes.py +8 -2
- rag/batch_workers.py +23 -10
- rag/rag.py +1 -0
- rag/routes.py +8 -0
- store.py +97 -7
- summary_generator/routes.py +6 -1
main.py
CHANGED
|
@@ -11,6 +11,8 @@ from src.summary_generator.routes import router as summary_router
|
|
| 11 |
from src.rag.routes import router as tutor_router
|
| 12 |
from src.quiz_generator.routes import router as quiz_router
|
| 13 |
from src.auth.routes import router as auth_router
|
|
|
|
|
|
|
| 14 |
from src.config import settings
|
| 15 |
|
| 16 |
# ── Logging Setup ──────────────────────────────────────
|
|
@@ -85,3 +87,8 @@ async def root():
|
|
| 85 |
@app.get("/api/health")
|
| 86 |
async def health_check():
|
| 87 |
return {"status": "ok", "service": "AI Tutor API"}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 11 |
from src.rag.routes import router as tutor_router
|
| 12 |
from src.quiz_generator.routes import router as quiz_router
|
| 13 |
from src.auth.routes import router as auth_router
|
| 14 |
+
from src.store import get_usage
|
| 15 |
+
from src.dependencies import get_current_user_id
|
| 16 |
from src.config import settings
|
| 17 |
|
| 18 |
# ── Logging Setup ──────────────────────────────────────
|
|
|
|
| 87 |
@app.get("/api/health")
|
| 88 |
async def health_check():
|
| 89 |
return {"status": "ok", "service": "AI Tutor API"}
|
| 90 |
+
|
| 91 |
+
|
| 92 |
+
@app.get("/api/usage")
|
| 93 |
+
async def get_user_usage(user_id: str = Depends(get_current_user_id)):
|
| 94 |
+
return get_usage(user_id)
|
materials/routes.py
CHANGED
|
@@ -64,7 +64,8 @@ async def _process_pdf_background(material_id: str, file_content: bytes):
|
|
| 64 |
# Skip processing if this user already has a material with this title
|
| 65 |
mat = get_material(material_id)
|
| 66 |
if mat and is_title_taken(mat.get("title", ""), exclude_id=material_id, user_id=mat.get("user_id")):
|
| 67 |
-
logger.info(f"Skipping processing for {material_id}: duplicate title
|
|
|
|
| 68 |
return
|
| 69 |
|
| 70 |
loop = asyncio.get_event_loop()
|
|
|
|
| 64 |
# Skip processing if this user already has a material with this title
|
| 65 |
mat = get_material(material_id)
|
| 66 |
if mat and is_title_taken(mat.get("title", ""), exclude_id=material_id, user_id=mat.get("user_id")):
|
| 67 |
+
logger.info(f"Skipping processing for {material_id}: duplicate title")
|
| 68 |
+
update_material_status(material_id, "failed", "Duplicate title. Please rename to retry.")
|
| 69 |
return
|
| 70 |
|
| 71 |
loop = asyncio.get_event_loop()
|
quiz_generator/routes.py
CHANGED
|
@@ -4,8 +4,8 @@ from pydantic import BaseModel
|
|
| 4 |
from typing import Optional
|
| 5 |
|
| 6 |
from src.quiz_generator.quiz import smart_quiz_generator
|
| 7 |
-
from src.store import get_material, get_chunks, get_summary, save_quiz, get_quizzes, save_quiz_result, get_quiz_results
|
| 8 |
-
from src.dependencies import get_current_user_id
|
| 9 |
from src.config import settings
|
| 10 |
|
| 11 |
router = APIRouter(prefix="/api/quiz", tags=["Quiz"])
|
|
@@ -37,7 +37,13 @@ async def get_quiz_list(
|
|
| 37 |
async def generate_quiz(
|
| 38 |
body: QuizRequest,
|
| 39 |
user_id: str = Depends(get_current_user_id),
|
|
|
|
| 40 |
):
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 41 |
body.difficulty = body.difficulty.capitalize()
|
| 42 |
if body.mcq_count < 1 or body.mcq_count > 20:
|
| 43 |
raise HTTPException(400, "MCQ count must be between 1 and 20")
|
|
|
|
| 4 |
from typing import Optional
|
| 5 |
|
| 6 |
from src.quiz_generator.quiz import smart_quiz_generator
|
| 7 |
+
from src.store import get_material, get_chunks, get_summary, save_quiz, get_quizzes, save_quiz_result, get_quiz_results, check_and_increment_daily_limit
|
| 8 |
+
from src.dependencies import get_current_user_id, get_current_user
|
| 9 |
from src.config import settings
|
| 10 |
|
| 11 |
router = APIRouter(prefix="/api/quiz", tags=["Quiz"])
|
|
|
|
| 37 |
async def generate_quiz(
|
| 38 |
body: QuizRequest,
|
| 39 |
user_id: str = Depends(get_current_user_id),
|
| 40 |
+
current_user=Depends(get_current_user),
|
| 41 |
):
|
| 42 |
+
# Rate limit check
|
| 43 |
+
user_email = current_user.get("email") if isinstance(current_user, dict) else getattr(current_user, "email", None)
|
| 44 |
+
if not check_and_increment_daily_limit(user_id, email=user_email, limit=10):
|
| 45 |
+
raise HTTPException(429, "Daily limit of 10 requests reached. Come back tomorrow!")
|
| 46 |
+
|
| 47 |
body.difficulty = body.difficulty.capitalize()
|
| 48 |
if body.mcq_count < 1 or body.mcq_count > 20:
|
| 49 |
raise HTTPException(400, "MCQ count must be between 1 and 20")
|
rag/batch_workers.py
CHANGED
|
@@ -134,23 +134,36 @@ async def embedding_worker():
|
|
| 134 |
job_result = all_embeddings[idx: idx + n]
|
| 135 |
idx += n
|
| 136 |
|
| 137 |
-
|
| 138 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 139 |
job.done.set()
|
| 140 |
|
| 141 |
except Exception as e:
|
| 142 |
logger.error(f"Embedding batch failed: {e}", exc_info=True)
|
| 143 |
for job in batch:
|
| 144 |
-
|
| 145 |
-
|
| 146 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 147 |
finally:
|
| 148 |
set_request_in_flight(False)
|
| 149 |
|
| 150 |
|
| 151 |
# ═══════════════════════ Warmup Loop ════════════════════════
|
| 152 |
|
| 153 |
-
_WARMUP_INTERVAL_S =
|
| 154 |
|
| 155 |
|
| 156 |
async def _warmup_loop():
|
|
@@ -187,7 +200,7 @@ def start_workers():
|
|
| 187 |
"""
|
| 188 |
Launch all async worker coroutines. Call once during app startup.
|
| 189 |
|
| 190 |
-
-
|
| 191 |
- 1 warmup loop (keeps OpenMP threads alive)
|
| 192 |
"""
|
| 193 |
global _workers_started
|
|
@@ -195,8 +208,8 @@ def start_workers():
|
|
| 195 |
return
|
| 196 |
_workers_started = True
|
| 197 |
|
| 198 |
-
|
| 199 |
-
|
| 200 |
asyncio.create_task(_warmup_loop(), name="warmup_loop")
|
| 201 |
|
| 202 |
-
logger.info("Embedding batch workers started (
|
|
|
|
| 134 |
job_result = all_embeddings[idx: idx + n]
|
| 135 |
idx += n
|
| 136 |
|
| 137 |
+
# Use .get() or setdefault to avoid KeyError if initialization was missed
|
| 138 |
+
if job.job_id not in job_store:
|
| 139 |
+
job_store[job.job_id] = {"status": "pending", "result": None, "error": None}
|
| 140 |
+
|
| 141 |
+
job_store[job.job_id].update({
|
| 142 |
+
"status": "done",
|
| 143 |
+
"result": job_result
|
| 144 |
+
})
|
| 145 |
job.done.set()
|
| 146 |
|
| 147 |
except Exception as e:
|
| 148 |
logger.error(f"Embedding batch failed: {e}", exc_info=True)
|
| 149 |
for job in batch:
|
| 150 |
+
if job.job_id not in job_store:
|
| 151 |
+
job_store[job.job_id] = {"status": "error", "result": None, "error": str(e)}
|
| 152 |
+
else:
|
| 153 |
+
job_store[job.job_id].update({
|
| 154 |
+
"status": "error",
|
| 155 |
+
"error": str(e)
|
| 156 |
+
})
|
| 157 |
+
# Critical: always set the event so the request doesn't hang
|
| 158 |
+
if not job.done.is_set():
|
| 159 |
+
job.done.set()
|
| 160 |
finally:
|
| 161 |
set_request_in_flight(False)
|
| 162 |
|
| 163 |
|
| 164 |
# ═══════════════════════ Warmup Loop ════════════════════════
|
| 165 |
|
| 166 |
+
_WARMUP_INTERVAL_S = 300 # 5 minutes
|
| 167 |
|
| 168 |
|
| 169 |
async def _warmup_loop():
|
|
|
|
| 200 |
"""
|
| 201 |
Launch all async worker coroutines. Call once during app startup.
|
| 202 |
|
| 203 |
+
- 1 embedding worker (batched SentenceTransformer inference)
|
| 204 |
- 1 warmup loop (keeps OpenMP threads alive)
|
| 205 |
"""
|
| 206 |
global _workers_started
|
|
|
|
| 208 |
return
|
| 209 |
_workers_started = True
|
| 210 |
|
| 211 |
+
# Use only 1 worker to save RAM on this environment
|
| 212 |
+
asyncio.create_task(embedding_worker(), name="embedding_worker_0")
|
| 213 |
asyncio.create_task(_warmup_loop(), name="warmup_loop")
|
| 214 |
|
| 215 |
+
logger.info("Embedding batch workers started (1 worker + warmup loop)")
|
rag/rag.py
CHANGED
|
@@ -71,6 +71,7 @@ async def store_embeddings_async(material_id: str, chunk_ids: list[str], chunks:
|
|
| 71 |
from src.rag.batch_workers import EmbeddingJob, embedding_queue, job_store
|
| 72 |
|
| 73 |
job = EmbeddingJob(job_id=str(uuid.uuid4()), texts=chunks)
|
|
|
|
| 74 |
await embedding_queue.put(job)
|
| 75 |
await job.done.wait()
|
| 76 |
|
|
|
|
| 71 |
from src.rag.batch_workers import EmbeddingJob, embedding_queue, job_store
|
| 72 |
|
| 73 |
job = EmbeddingJob(job_id=str(uuid.uuid4()), texts=chunks)
|
| 74 |
+
job_store[job.job_id] = {"status": "pending", "result": None, "error": None}
|
| 75 |
await embedding_queue.put(job)
|
| 76 |
await job.done.wait()
|
| 77 |
|
rag/routes.py
CHANGED
|
@@ -12,6 +12,7 @@ from src.store import (
|
|
| 12 |
create_chat_session, list_chat_sessions, get_chat_session,
|
| 13 |
rename_chat_session, delete_chat_session,
|
| 14 |
append_session_message, get_session_messages,
|
|
|
|
| 15 |
# Legacy
|
| 16 |
save_chat_messages, get_chat_messages,
|
| 17 |
)
|
|
@@ -43,6 +44,13 @@ async def ask_tutor(
|
|
| 43 |
if not body.query.strip():
|
| 44 |
raise HTTPException(400, "Query cannot be empty")
|
| 45 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 46 |
# Determine memory key (prefer session_id for persistence)
|
| 47 |
mem_key = body.session_id or body.memory_id
|
| 48 |
|
|
|
|
| 12 |
create_chat_session, list_chat_sessions, get_chat_session,
|
| 13 |
rename_chat_session, delete_chat_session,
|
| 14 |
append_session_message, get_session_messages,
|
| 15 |
+
check_and_increment_daily_limit,
|
| 16 |
# Legacy
|
| 17 |
save_chat_messages, get_chat_messages,
|
| 18 |
)
|
|
|
|
| 44 |
if not body.query.strip():
|
| 45 |
raise HTTPException(400, "Query cannot be empty")
|
| 46 |
|
| 47 |
+
# Rate limit check
|
| 48 |
+
user_id = current_user.get("id") if isinstance(current_user, dict) else getattr(current_user, "id", None)
|
| 49 |
+
user_email = current_user.get("email") if isinstance(current_user, dict) else getattr(current_user, "email", None)
|
| 50 |
+
|
| 51 |
+
if user_id and not check_and_increment_daily_limit(user_id, email=user_email, limit=10):
|
| 52 |
+
raise HTTPException(429, "Daily limit of 10 requests reached. Come back tomorrow!")
|
| 53 |
+
|
| 54 |
# Determine memory key (prefer session_id for persistence)
|
| 55 |
mem_key = body.session_id or body.memory_id
|
| 56 |
|
store.py
CHANGED
|
@@ -1,5 +1,6 @@
|
|
|
|
|
| 1 |
from typing import Optional
|
| 2 |
-
from datetime import datetime, timezone
|
| 3 |
from langchain.memory import ConversationBufferMemory, ConversationBufferWindowMemory
|
| 4 |
from src.database import get_supabase
|
| 5 |
|
|
@@ -12,6 +13,12 @@ _in_memory: dict = {
|
|
| 12 |
"next_id": 0,
|
| 13 |
}
|
| 14 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 15 |
|
| 16 |
def _get_next_id() -> str:
|
| 17 |
_in_memory["next_id"] += 1
|
|
@@ -266,15 +273,12 @@ def get_summary(material_id: str) -> Optional[dict]:
|
|
| 266 |
_table_supabase("summaries")
|
| 267 |
.select("*")
|
| 268 |
.eq("material_id", material_id)
|
| 269 |
-
.maybe_single()
|
| 270 |
.execute()
|
| 271 |
)
|
| 272 |
-
if not result or result.data
|
| 273 |
return None
|
| 274 |
-
#
|
| 275 |
-
|
| 276 |
-
return result.data[0] if result.data else None
|
| 277 |
-
return result.data
|
| 278 |
except Exception:
|
| 279 |
return None
|
| 280 |
|
|
@@ -520,3 +524,89 @@ def get_or_create_memory(memory_id: Optional[str] = None, seed_messages: list[di
|
|
| 520 |
)
|
| 521 |
_memories[mid] = mem
|
| 522 |
return mem, mid
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os
|
| 2 |
from typing import Optional
|
| 3 |
+
from datetime import datetime, timezone, date
|
| 4 |
from langchain.memory import ConversationBufferMemory, ConversationBufferWindowMemory
|
| 5 |
from src.database import get_supabase
|
| 6 |
|
|
|
|
| 13 |
"next_id": 0,
|
| 14 |
}
|
| 15 |
|
| 16 |
+
ADMIN_EMAILS = set(
|
| 17 |
+
email.strip()
|
| 18 |
+
for email in os.environ.get("ADMIN_EMAILS", "").split(",")
|
| 19 |
+
if email.strip()
|
| 20 |
+
)
|
| 21 |
+
|
| 22 |
|
| 23 |
def _get_next_id() -> str:
|
| 24 |
_in_memory["next_id"] += 1
|
|
|
|
| 273 |
_table_supabase("summaries")
|
| 274 |
.select("*")
|
| 275 |
.eq("material_id", material_id)
|
|
|
|
| 276 |
.execute()
|
| 277 |
)
|
| 278 |
+
if not result or not result.data:
|
| 279 |
return None
|
| 280 |
+
# Return the most recent one if duplicates exist
|
| 281 |
+
return result.data[0]
|
|
|
|
|
|
|
| 282 |
except Exception:
|
| 283 |
return None
|
| 284 |
|
|
|
|
| 524 |
)
|
| 525 |
_memories[mid] = mem
|
| 526 |
return mem, mid
|
| 527 |
+
|
| 528 |
+
|
| 529 |
+
def check_and_increment_daily_limit(user_id: str, email: Optional[str] = None, limit: int = 10) -> bool:
|
| 530 |
+
"""
|
| 531 |
+
Returns True if request is allowed, False if limit exceeded.
|
| 532 |
+
Excludes Admin Emails from Limits
|
| 533 |
+
"""
|
| 534 |
+
# Exclude specific email from rate limiting
|
| 535 |
+
if email in ADMIN_EMAILS:
|
| 536 |
+
return True
|
| 537 |
+
|
| 538 |
+
today = date.today().isoformat()
|
| 539 |
+
|
| 540 |
+
try:
|
| 541 |
+
# Get current profile
|
| 542 |
+
result = _robust_execute(
|
| 543 |
+
_table_supabase("profiles")
|
| 544 |
+
.select("daily_requests, last_request_date")
|
| 545 |
+
.eq("id", user_id)
|
| 546 |
+
.maybe_single()
|
| 547 |
+
)
|
| 548 |
+
|
| 549 |
+
# If no profile or data, allow (or we could create one, but usually it exists)
|
| 550 |
+
if not result.data:
|
| 551 |
+
return True
|
| 552 |
+
|
| 553 |
+
# Handle both single object or list response from maybe_single/execute
|
| 554 |
+
profile = result.data[0] if isinstance(result.data, list) and result.data else result.data
|
| 555 |
+
if not profile:
|
| 556 |
+
return True
|
| 557 |
+
|
| 558 |
+
last_date = profile.get("last_request_date")
|
| 559 |
+
count = profile.get("daily_requests", 0) or 0
|
| 560 |
+
|
| 561 |
+
# Reset count if it's a new day
|
| 562 |
+
if last_date != today:
|
| 563 |
+
count = 0
|
| 564 |
+
|
| 565 |
+
# Check limit
|
| 566 |
+
if count >= limit:
|
| 567 |
+
return False
|
| 568 |
+
|
| 569 |
+
# Increment
|
| 570 |
+
_robust_execute(
|
| 571 |
+
_table_supabase("profiles")
|
| 572 |
+
.update({
|
| 573 |
+
"daily_requests": count + 1,
|
| 574 |
+
"last_request_date": today,
|
| 575 |
+
})
|
| 576 |
+
.eq("id", user_id)
|
| 577 |
+
)
|
| 578 |
+
return True
|
| 579 |
+
except Exception as e:
|
| 580 |
+
# If DB fails, we default to allowing the request to not break the app
|
| 581 |
+
import logging
|
| 582 |
+
logging.getLogger(__name__).error(f"Rate limit check failed: {e}")
|
| 583 |
+
return True
|
| 584 |
+
|
| 585 |
+
|
| 586 |
+
def get_usage(user_id: str) -> dict:
|
| 587 |
+
"""
|
| 588 |
+
Returns current usage for a user.
|
| 589 |
+
"""
|
| 590 |
+
today = date.today().isoformat()
|
| 591 |
+
try:
|
| 592 |
+
result = _robust_execute(
|
| 593 |
+
_table_supabase("profiles")
|
| 594 |
+
.select("daily_requests, last_request_date")
|
| 595 |
+
.eq("id", user_id)
|
| 596 |
+
.maybe_single()
|
| 597 |
+
)
|
| 598 |
+
if not result.data:
|
| 599 |
+
return {"used": 0, "limit": 10, "remaining": 10}
|
| 600 |
+
|
| 601 |
+
profile = result.data[0] if isinstance(result.data, list) and result.data else result.data
|
| 602 |
+
if not profile:
|
| 603 |
+
return {"used": 0, "limit": 10, "remaining": 10}
|
| 604 |
+
|
| 605 |
+
used = profile.get("daily_requests", 0) if profile.get("last_request_date") == today else 0
|
| 606 |
+
return {
|
| 607 |
+
"used": used,
|
| 608 |
+
"limit": 10,
|
| 609 |
+
"remaining": max(0, 10 - used)
|
| 610 |
+
}
|
| 611 |
+
except Exception:
|
| 612 |
+
return {"used": 0, "limit": 10, "remaining": 10}
|
summary_generator/routes.py
CHANGED
|
@@ -6,7 +6,7 @@ from pydantic import BaseModel
|
|
| 6 |
from typing import Optional
|
| 7 |
|
| 8 |
from src.summary_generator.summary import summarizer
|
| 9 |
-
from src.store import get_material, get_chunks, save_summary, get_summary as get_stored_summary
|
| 10 |
from src.dependencies import get_current_user_id, get_current_user
|
| 11 |
from src.config import settings
|
| 12 |
|
|
@@ -28,6 +28,11 @@ async def generate_summary(
|
|
| 28 |
user_id: str = Depends(get_current_user_id),
|
| 29 |
current_user=Depends(get_current_user),
|
| 30 |
):
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 31 |
mat = get_material(body.material_id)
|
| 32 |
if not mat:
|
| 33 |
raise HTTPException(404, "Material not found")
|
|
|
|
| 6 |
from typing import Optional
|
| 7 |
|
| 8 |
from src.summary_generator.summary import summarizer
|
| 9 |
+
from src.store import get_material, get_chunks, save_summary, get_summary as get_stored_summary, check_and_increment_daily_limit
|
| 10 |
from src.dependencies import get_current_user_id, get_current_user
|
| 11 |
from src.config import settings
|
| 12 |
|
|
|
|
| 28 |
user_id: str = Depends(get_current_user_id),
|
| 29 |
current_user=Depends(get_current_user),
|
| 30 |
):
|
| 31 |
+
# Rate limit check
|
| 32 |
+
user_email = current_user.get("email") if isinstance(current_user, dict) else getattr(current_user, "email", None)
|
| 33 |
+
if not check_and_increment_daily_limit(user_id, email=user_email, limit=10):
|
| 34 |
+
raise HTTPException(429, "Daily limit of 10 requests reached. Come back tomorrow!")
|
| 35 |
+
|
| 36 |
mat = get_material(body.material_id)
|
| 37 |
if not mat:
|
| 38 |
raise HTTPException(404, "Material not found")
|