Sarp Bilgiç commited on
Commit
791579e
·
1 Parent(s): 60917e8

general rate limit for all endpoints

Browse files
src/api/dependencies/clients.py CHANGED
@@ -43,7 +43,6 @@ def get_redis_client() -> RedisClient:
43
 
44
  @lru_cache()
45
  def get_reranker_client() -> RerankerClient | None:
46
- """Get reranker client if enabled."""
47
  if not settings.reranker_enabled:
48
  return None
49
  return RerankerClient(model_name=settings.reranker_model)
 
43
 
44
  @lru_cache()
45
  def get_reranker_client() -> RerankerClient | None:
 
46
  if not settings.reranker_enabled:
47
  return None
48
  return RerankerClient(model_name=settings.reranker_model)
src/api/dependencies/rate_limit.py CHANGED
@@ -29,19 +29,25 @@ async def ip_identifier(request: Request) -> str:
29
  return f"ip:{ip}"
30
 
31
  anonymous_rag_rate_limiter = RateLimiter(
32
- times=7 if settings.env == "production" else 100,
33
  seconds=3600,
34
  identifier=ip_identifier
35
  )
36
 
37
  authenticated_rag_rate_limiter = RateLimiter(
38
- times=15 if settings.env == "production" else 100,
39
  seconds=3600,
40
  identifier=user_id_identifier
41
  )
42
 
43
  login_rate_limiter = RateLimiter(
44
- times=10,
45
  seconds=300,
46
  identifier=ip_identifier
 
 
 
 
 
 
47
  )
 
29
  return f"ip:{ip}"
30
 
31
  anonymous_rag_rate_limiter = RateLimiter(
32
+ times=10 if settings.env == "production" else 100,
33
  seconds=3600,
34
  identifier=ip_identifier
35
  )
36
 
37
  authenticated_rag_rate_limiter = RateLimiter(
38
+ times=25 if settings.env == "production" else 100,
39
  seconds=3600,
40
  identifier=user_id_identifier
41
  )
42
 
43
  login_rate_limiter = RateLimiter(
44
+ times=25,
45
  seconds=300,
46
  identifier=ip_identifier
47
+ )
48
+
49
+ general_rate_limiter = RateLimiter(
50
+ times=100,
51
+ seconds=3600,
52
+ identifier=user_id_identifier
53
  )
src/api/routers/auth.py CHANGED
@@ -12,6 +12,7 @@ from src.api.selectors.user.add_user import add_user
12
  from src.api.models.user import User
13
  from src.api.schemas.auth import RegisterRequest, LoginRequest, TokenResponse
14
  import redis.asyncio as redis
 
15
 
16
 
17
  router = APIRouter(
@@ -19,7 +20,7 @@ router = APIRouter(
19
  tags=["auth"]
20
  )
21
 
22
- @router.post("/register", response_model=TokenResponse)
23
  async def register(
24
  request: RegisterRequest,
25
  auth_service: Annotated[AuthService, Depends(get_auth_service)],
@@ -43,7 +44,7 @@ async def register(
43
 
44
  return TokenResponse(access_token=access_token)
45
 
46
- @router.post("/login", response_model=TokenResponse)
47
  async def login(
48
  request: LoginRequest,
49
  auth_service: Annotated[AuthService, Depends(get_auth_service)],
@@ -68,7 +69,7 @@ async def login(
68
  access_token = auth_service.create_access_token(data={"sub": user.email})
69
  return TokenResponse(access_token=access_token)
70
 
71
- @router.post("/token", response_model=TokenResponse, include_in_schema=False)
72
  async def login_for_swagger(
73
  form_data: Annotated[OAuth2PasswordRequestForm, Depends()],
74
  db: Annotated[AsyncSession, Depends(get_db)],
 
12
  from src.api.models.user import User
13
  from src.api.schemas.auth import RegisterRequest, LoginRequest, TokenResponse
14
  import redis.asyncio as redis
15
+ from src.api.dependencies.rate_limit import login_rate_limiter
16
 
17
 
18
  router = APIRouter(
 
20
  tags=["auth"]
21
  )
22
 
23
+ @router.post("/register", response_model=TokenResponse, dependencies=[Depends(login_rate_limiter)])
24
  async def register(
25
  request: RegisterRequest,
26
  auth_service: Annotated[AuthService, Depends(get_auth_service)],
 
44
 
45
  return TokenResponse(access_token=access_token)
46
 
47
+ @router.post("/login", response_model=TokenResponse, dependencies=[Depends(login_rate_limiter)])
48
  async def login(
49
  request: LoginRequest,
50
  auth_service: Annotated[AuthService, Depends(get_auth_service)],
 
69
  access_token = auth_service.create_access_token(data={"sub": user.email})
70
  return TokenResponse(access_token=access_token)
71
 
72
+ @router.post("/token", response_model=TokenResponse, include_in_schema=False, dependencies=[Depends(login_rate_limiter)])
73
  async def login_for_swagger(
74
  form_data: Annotated[OAuth2PasswordRequestForm, Depends()],
75
  db: Annotated[AsyncSession, Depends(get_db)],
src/api/routers/sessions.py CHANGED
@@ -10,6 +10,7 @@ from src.api.schemas.session import ChatSessionList, ChatMessageRead
10
  from src.api.services.chat_history_service import ChatHistoryService
11
  from llama_index.core.llms import MessageRole
12
  import uuid
 
13
 
14
  router = APIRouter(
15
  prefix="/api/v1",
@@ -24,7 +25,7 @@ def _llama_role_to_chat_role(role: MessageRole) -> ChatMessageRole:
24
  }
25
  return mapping.get(role, ChatMessageRole.USER)
26
 
27
- @router.get("/sessions", response_model=List[ChatSessionList])
28
  async def list_sessions(
29
  user: Annotated[User, Depends(get_current_user_required)],
30
  db: Annotated[AsyncSession, Depends(get_db)],
@@ -38,7 +39,7 @@ async def list_sessions(
38
  offset=offset
39
  )
40
 
41
- @router.get("/sessions/{session_id}/messages", response_model=List[ChatMessageRead])
42
  async def get_messages(
43
  session_id: uuid.UUID,
44
  chat_history_service: Annotated[ChatHistoryService, Depends(get_chat_history_service)],
@@ -58,7 +59,7 @@ async def get_messages(
58
  for msg in messages
59
  ]
60
 
61
- @router.delete("/sessions/{session_id}")
62
  async def delete_session(
63
  session_id: uuid.UUID,
64
  chat_history_service: Annotated[ChatHistoryService, Depends(get_chat_history_service)],
 
10
  from src.api.services.chat_history_service import ChatHistoryService
11
  from llama_index.core.llms import MessageRole
12
  import uuid
13
+ from src.api.dependencies.rate_limit import general_rate_limiter
14
 
15
  router = APIRouter(
16
  prefix="/api/v1",
 
25
  }
26
  return mapping.get(role, ChatMessageRole.USER)
27
 
28
+ @router.get("/sessions", response_model=List[ChatSessionList], dependencies=[Depends(general_rate_limiter)])
29
  async def list_sessions(
30
  user: Annotated[User, Depends(get_current_user_required)],
31
  db: Annotated[AsyncSession, Depends(get_db)],
 
39
  offset=offset
40
  )
41
 
42
+ @router.get("/sessions/{session_id}/messages", response_model=List[ChatMessageRead], dependencies=[Depends(general_rate_limiter)])
43
  async def get_messages(
44
  session_id: uuid.UUID,
45
  chat_history_service: Annotated[ChatHistoryService, Depends(get_chat_history_service)],
 
59
  for msg in messages
60
  ]
61
 
62
+ @router.delete("/sessions/{session_id}", dependencies=[Depends(general_rate_limiter)])
63
  async def delete_session(
64
  session_id: uuid.UUID,
65
  chat_history_service: Annotated[ChatHistoryService, Depends(get_chat_history_service)],
src/api/routers/user.py CHANGED
@@ -3,12 +3,13 @@ from src.api.dependencies.auth import get_current_user_required
3
  from src.api.models.user import User
4
  from src.api.schemas.user import UserRead
5
  from typing import Annotated
 
6
 
7
  router = APIRouter(
8
  prefix="/api/v1/user",
9
  tags=["user"],
10
  )
11
- @router.get("/me", response_model=UserRead)
12
  async def get_me(
13
  user: Annotated[User, Depends(get_current_user_required)],
14
  ):
 
3
  from src.api.models.user import User
4
  from src.api.schemas.user import UserRead
5
  from typing import Annotated
6
+ from src.api.dependencies.rate_limit import general_rate_limiter
7
 
8
  router = APIRouter(
9
  prefix="/api/v1/user",
10
  tags=["user"],
11
  )
12
+ @router.get("/me", response_model=UserRead, dependencies=[Depends(general_rate_limiter)])
13
  async def get_me(
14
  user: Annotated[User, Depends(get_current_user_required)],
15
  ):