Hamdy005 commited on
Commit
6ee4b4a
·
1 Parent(s): 7b90b65

feat: implement daily request rate limiting with per-user usage tracking and admin bypass

Browse files
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, waiting for rename")
 
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
- job_store[job.job_id]["status"] = "done"
138
- job_store[job.job_id]["result"] = job_result
 
 
 
 
 
 
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
- job_store[job.job_id]["status"] = "error"
145
- job_store[job.job_id]["error"] = str(e)
146
- job.done.set()
 
 
 
 
 
 
 
147
  finally:
148
  set_request_in_flight(False)
149
 
150
 
151
  # ═══════════════════════ Warmup Loop ════════════════════════
152
 
153
- _WARMUP_INTERVAL_S = 45
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
- - 2 embedding workers (batched SentenceTransformer inference)
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
- for i in range(2):
199
- asyncio.create_task(embedding_worker(), name=f"embedding_worker_{i}")
200
  asyncio.create_task(_warmup_loop(), name="warmup_loop")
201
 
202
- logger.info("Embedding batch workers started (2 workers + warmup loop)")
 
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 is None:
273
  return None
274
- # _make_response returns list; unwrap the first item
275
- if isinstance(result.data, list):
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")