Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 7 additions & 6 deletions app/api/v1/endpoints/habits.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
5 changes: 3 additions & 2 deletions app/api/v1/endpoints/log.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,17 +3,18 @@
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')
async def create(request: Request, habit_id: int, data: LogCreate, service: LogService = Depends(get_log_service)):
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
5 changes: 3 additions & 2 deletions app/api/v1/endpoints/stats.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
3 changes: 0 additions & 3 deletions app/core/logger.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 = (
'<green>{time:YYYY-MM-DD HH:mm:ss}</green> | '
'<level>{level: <8}</level> | '
Expand Down
6 changes: 3 additions & 3 deletions app/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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):
Expand Down
37 changes: 26 additions & 11 deletions app/middlewares/rate_limit/limiter.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.'}
)

20 changes: 3 additions & 17 deletions tests/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down