diff --git a/app/api/v1/endpoints/habits.py b/app/api/v1/endpoints/habits.py index ab95a59..20f8f5b 100644 --- a/app/api/v1/endpoints/habits.py +++ b/app/api/v1/endpoints/habits.py @@ -3,38 +3,39 @@ from app.dependencies import get_habit_service from app.services.habit_service import HabitService from app.schemas.habit import HabitCreate, HabitUpdate, HabitResponse -from app.middlewares.rate_limit.limiter import limiter +from app.middlewares.rate_limit.limiter import RateLimit from app.infrastructure.redis.cache import CacheService router = APIRouter() +limiter = RateLimit.get_limiter() @router.get('/habits', response_model = List[HabitResponse]) @CacheService.cached(expire = 50) -@limiter.limit('5/minute') +@limiter.limit('15/minute') async def get_all(request: Request, service: HabitService = Depends(get_habit_service)): return await service.get_all() @router.post('/habit/create', response_model = HabitResponse, status_code = status.HTTP_201_CREATED) -@limiter.limit('5/minute') +@limiter.limit('15/minute') async def create(request: Request, data: HabitCreate, service: HabitService = Depends(get_habit_service)): await CacheService.clear() return await service.create(data) @router.get('/habit/{habit_id}', response_model = HabitResponse) @CacheService.cached(expire = 50) -@limiter.limit('5/minute') +@limiter.limit('15/minute') async def get_by_id(request: Request, habit_id: int, service: HabitService = Depends(get_habit_service)): return await service.get_by_id(habit_id) @router.delete('/habit/{habit_id}', status_code = status.HTTP_204_NO_CONTENT) -@limiter.limit('5/minute') +@limiter.limit('15/minute') async def delete_by_id(request: Request, habit_id: int, service: HabitService = Depends(get_habit_service)): await service.delete(habit_id) await CacheService.clear() return None @router.put('/habit/{habit_id}', response_model = HabitResponse) -@limiter.limit('5/minute') +@limiter.limit('15/minute') async def update(request: Request, habit_id: int, data: HabitUpdate, service: HabitService = Depends(get_habit_service)): await CacheService.clear() return await service.update(habit_id, data) diff --git a/app/api/v1/endpoints/log.py b/app/api/v1/endpoints/log.py index 0ef8fa6..36e3c03 100644 --- a/app/api/v1/endpoints/log.py +++ b/app/api/v1/endpoints/log.py @@ -3,9 +3,10 @@ from app.services.log_service import LogService from app.schemas.log import LogCreate from app.schemas.log import LogResponse -from app.middlewares.rate_limit.limiter import limiter +from app.middlewares.rate_limit.limiter import RateLimit router = APIRouter() +limiter = RateLimit.get_limiter() @router.post('/log/{habit_id}', response_model = LogResponse, status_code = status.HTTP_201_CREATED) @limiter.limit('1/day') @@ -13,7 +14,7 @@ async def create(request: Request, habit_id: int, data: LogCreate, service: LogS return await service.create(habit_id, data) @router.delete('/log/{log_id}', status_code = status.HTTP_204_NO_CONTENT) -@limiter.limit('5/minute') +@limiter.limit('15/minute') async def delete(request: Request, log_id: int, service: LogService = Depends(get_log_service)): await service.delete(log_id) return None diff --git a/app/api/v1/endpoints/stats.py b/app/api/v1/endpoints/stats.py index d573159..410f771 100644 --- a/app/api/v1/endpoints/stats.py +++ b/app/api/v1/endpoints/stats.py @@ -2,13 +2,14 @@ from app.dependencies import get_stats_service from app.services.stats_service import StatsService from app.schemas.stats import StatsResponse -from app.middlewares.rate_limit.limiter import limiter +from app.middlewares.rate_limit.limiter import RateLimit from app.infrastructure.redis.cache import CacheService router = APIRouter() +limiter = RateLimit.get_limiter() @router.get('/stats/{habit_id}', response_model = StatsResponse) @CacheService.cached(expire = 50) -@limiter.limit('5/minute') +@limiter.limit('15/minute') async def get(request: Request, habit_id: int, service: StatsService = Depends(get_stats_service)): return await service.get_stats(habit_id) \ No newline at end of file diff --git a/app/core/logger.py b/app/core/logger.py index b63591b..8488efd 100644 --- a/app/core/logger.py +++ b/app/core/logger.py @@ -4,9 +4,6 @@ logger.remove() # removing the old handler to set everything up yourself -log_dir = Path('logs') -log_dir.mkdir(exist_ok = True) - log_format = ( '{time:YYYY-MM-DD HH:mm:ss} | ' '{level: <8} | ' diff --git a/app/main.py b/app/main.py index c4928a1..af536b0 100644 --- a/app/main.py +++ b/app/main.py @@ -5,7 +5,7 @@ from slowapi.errors import RateLimitExceeded from contextlib import asynccontextmanager from app.core.database import engine, Base -from app.middlewares.rate_limit.limiter import limiter, rate_limit_exceed_handler +from app.middlewares.rate_limit.limiter import RateLimit from app.middlewares.logging.logging_middleware import LoggingMiddleware from app.api import router from app.core.logger import logger @@ -21,8 +21,8 @@ async def lifespan(app: FastAPI): app = FastAPI(lifespan = lifespan) -app.state.limiter = limiter -app.add_exception_handler(RateLimitExceeded, rate_limit_exceed_handler) +app.state.limiter = RateLimit.get_limiter() +app.add_exception_handler(RateLimitExceeded, RateLimit.rate_limit_exceed_handler) @app.exception_handler(Exception) async def global_exeption_handler(request: Request, exc: Exception): diff --git a/app/middlewares/rate_limit/limiter.py b/app/middlewares/rate_limit/limiter.py index a7efacb..0339aa9 100644 --- a/app/middlewares/rate_limit/limiter.py +++ b/app/middlewares/rate_limit/limiter.py @@ -2,19 +2,34 @@ from slowapi.util import get_remote_address from slowapi.errors import RateLimitExceeded from fastapi import Request, status -import redis.asyncio as redis from app.core.config import settings from fastapi.responses import JSONResponse -limiter = Limiter( - key_func = get_remote_address, - storage_uri = settings.REDIS_URL, - default_limits = ['20/minute'] -) +class RateLimit: + _enabled_redis = True -def rate_limit_exceed_handler(request: Request, exc: RateLimitExceeded): - return JSONResponse( - status_code = status.HTTP_429_TOO_MANY_REQUESTS, - content = {'detail' : 'Too many requests.'} - ) + @classmethod + def disable_redis(cls): + cls._enabled_redis = False + + @classmethod + def get_limiter(cls): + if not cls._enabled_redis: + return Limiter( + key_func = get_remote_address, + default_limits = ['20/minute'], + enabled = False # Disabling limits for tests + ) + return Limiter( + key_func = get_remote_address, + storage_uri = settings.REDIS_URL, + default_limits = ['20/minute'] + ) + + @staticmethod + def rate_limit_exceed_handler(request: Request, exc: RateLimitExceeded): + return JSONResponse( + status_code = status.HTTP_429_TOO_MANY_REQUESTS, + content = {'detail' : 'Too many requests.'} + ) diff --git a/tests/conftest.py b/tests/conftest.py index f659ecb..367588e 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -4,27 +4,13 @@ from app.core.database import Base, get_db from app.infrastructure.redis.cache import CacheService -CacheService.disable() - -import app.middlewares.rate_limit.limiter as rate_limit_module -from slowapi import Limiter -from slowapi.util import get_remote_address - -def mock_limit(self, limit_str: str): - def decorator(func): - return func - return decorator +from app.middlewares.rate_limit.limiter import RateLimit -rate_limit_module.limiter = Limiter( - key_func=get_remote_address, - default_limits=["100/minute"] -) -rate_limit_module.limiter.limit = mock_limit.__get__(rate_limit_module.limiter, Limiter) +CacheService.disable() +RateLimit.disable_redis() from app.main import app -app.state.limiter = rate_limit_module.limiter - TEST_DATABASE_URL = "sqlite+aiosqlite:///:memory:" engine = create_async_engine(TEST_DATABASE_URL, echo=False) TestingSessionLocal = async_sessionmaker(engine, class_=AsyncSession, expire_on_commit=False)