Fix rate limiting

This commit is contained in:
2025-11-10 20:13:23 +03:00
Unverified
parent 0f86e9541c
commit 6a3b313f2c
3 changed files with 75 additions and 63 deletions
+23 -23
View File
@@ -68,10 +68,10 @@ def check_auth(current_user: User = Depends(get_current_user)):
@router.post("/login") @router.post("/login")
@rate_limit_per_ip("5/minute") @rate_limit_per_ip("5/minute")
def login(request: LoginRequest, http: Request, db: Session = Depends(get_db)): def login(request: Request, login_request: LoginRequest, db: Session = Depends(get_db)):
username = request.username.strip() username = login_request.username.strip()
client_ip = get_client_ip(http) client_ip = get_client_ip(request)
raw_ua = http.headers.get("user-agent") raw_ua = request.headers.get("user-agent")
if is_user_agent_blocked(raw_ua): if is_user_agent_blocked(raw_ua):
log_security( log_security(
@@ -89,7 +89,7 @@ def login(request: LoginRequest, http: Request, db: Session = Depends(get_db)):
user = db.query(User).filter(User.username == username).first() user = db.query(User).filter(User.username == username).first()
if not user or not verify_password(request.password.strip(), user.password_hash): if not user or not verify_password(login_request.password.strip(), user.password_hash):
log_security( log_security(
"login_failed", "login_failed",
severity="warning", severity="warning",
@@ -125,8 +125,8 @@ def login(request: LoginRequest, http: Request, db: Session = Depends(get_db)):
) )
# Create device session and embed into JWT # Create device session and embed into JWT
raw_ua = http.headers.get("user-agent") raw_ua = request.headers.get("user-agent")
device_name = http.headers.get("x-device-name") device_name = request.headers.get("x-device-name")
ua = parse_ua(raw_ua or "") ua = parse_ua(raw_ua or "")
session_id = uuid.uuid4().hex session_id = uuid.uuid4().hex
@@ -181,13 +181,13 @@ def login(request: LoginRequest, http: Request, db: Session = Depends(get_db)):
@router.post("/register") @router.post("/register")
@rate_limit_per_ip("3/hour") @rate_limit_per_ip("3/hour")
def register(request: RegisterRequest, http: Request, db: Session = Depends(get_db)): def register(request: Request, register_request: RegisterRequest, db: Session = Depends(get_db)):
username = request.username.strip() username = register_request.username.strip()
display_name = request.display_name.strip() display_name = register_request.display_name.strip()
password = request.password.strip() password = register_request.password.strip()
confirm_password = request.confirm_password.strip() confirm_password = register_request.confirm_password.strip()
client_ip = get_client_ip(http) client_ip = get_client_ip(request)
raw_ua = http.headers.get("user-agent") raw_ua = request.headers.get("user-agent")
if is_user_agent_blocked(raw_ua): if is_user_agent_blocked(raw_ua):
log_security( log_security(
@@ -281,8 +281,8 @@ def register(request: RegisterRequest, http: Request, db: Session = Depends(get_
db.refresh(new_user) db.refresh(new_user)
# Create initial device session # Create initial device session
raw_ua = http.headers.get("user-agent") raw_ua = request.headers.get("user-agent")
device_name = http.headers.get("x-device-name") device_name = request.headers.get("x-device-name")
ua = parse_ua(raw_ua or "") ua = parse_ua(raw_ua or "")
session_id = uuid.uuid4().hex session_id = uuid.uuid4().hex
device = DeviceSession( device = DeviceSession(
@@ -447,22 +447,22 @@ def logout(
@router.post("/change-password") @router.post("/change-password")
@rate_limit_per_user("5/hour") @rate_limit_per_user("5/hour")
def change_password( def change_password(
request: ChangePasswordRequest, request: Request,
http: Request, password_request: ChangePasswordRequest,
credentials: HTTPAuthorizationCredentials = Depends(HTTPBearer()), credentials: HTTPAuthorizationCredentials = Depends(HTTPBearer()),
current_user: User = Depends(get_current_user), current_user: User = Depends(get_current_user),
db: Session = Depends(get_db) db: Session = Depends(get_db)
): ):
# Verify current derived password against stored hash # Verify current derived password against stored hash
if not verify_password(request.currentPasswordDerived.strip(), current_user.password_hash): if not verify_password(password_request.currentPasswordDerived.strip(), current_user.password_hash):
raise HTTPException(status_code=401, detail="Текущий пароль неверный") raise HTTPException(status_code=401, detail="Текущий пароль неверный")
# Update password hash to hash of new derived password # Update password hash to hash of new derived password
current_user.password_hash = get_password_hash(request.newPasswordDerived.strip()) current_user.password_hash = get_password_hash(password_request.newPasswordDerived.strip())
db.commit() db.commit()
# Optionally revoke all other sessions, keeping the current one # Optionally revoke all other sessions, keeping the current one
if request.logoutAllExceptCurrent: if password_request.logoutAllExceptCurrent:
from utils import verify_token as _verify_token from utils import verify_token as _verify_token
payload = _verify_token(credentials.credentials) payload = _verify_token(credentials.credentials)
if not payload: if not payload:
@@ -474,13 +474,13 @@ def change_password(
).update({DeviceSession.revoked: True}) ).update({DeviceSession.revoked: True})
db.commit() db.commit()
client_ip = get_client_ip(http) client_ip = get_client_ip(request)
log_security( log_security(
"password_changed", "password_changed",
username=current_user.username, username=current_user.username,
user_id=current_user.id, user_id=current_user.id,
ip=client_ip, ip=client_ip,
logout_others=bool(request.logoutAllExceptCurrent), logout_others=bool(password_request.logoutAllExceptCurrent),
) )
return {"status": "success"} return {"status": "success"}
+38 -30
View File
@@ -12,7 +12,7 @@ from collections import defaultdict, deque
from difflib import SequenceMatcher from difflib import SequenceMatcher
from types import SimpleNamespace from types import SimpleNamespace
from typing import Any from typing import Any
from fastapi import APIRouter, Depends, HTTPException, WebSocket, WebSocketDisconnect, UploadFile, File, Form from fastapi import APIRouter, Depends, HTTPException, WebSocket, WebSocketDisconnect, UploadFile, File, Form, Request
from fastapi.responses import FileResponse from fastapi.responses import FileResponse
from fastapi.security import HTTPAuthorizationCredentials from fastapi.security import HTTPAuthorizationCredentials
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
@@ -254,7 +254,8 @@ def convert_dm_envelope(envelope: DMEnvelope) -> dict:
@router.post("/send_message") @router.post("/send_message")
@rate_limit_per_user("30/minute") @rate_limit_per_user("30/minute")
async def send_message( async def send_message(
request: SendMessageRequest | None = None, request: Request,
message_request: SendMessageRequest | None = None,
current_user: User = Depends(get_current_user), current_user: User = Depends(get_current_user),
db: Session = Depends(get_db), db: Session = Depends(get_db),
# Optional multipart form support # Optional multipart form support
@@ -262,23 +263,26 @@ async def send_message(
files: list[UploadFile] = File(default=[]), files: list[UploadFile] = File(default=[]),
): ):
# If payload is provided, prefer it for multipart requests # If payload is provided, prefer it for multipart requests
if payload and request is None: if payload and message_request is None:
# Expect JSON: {"type":"text","data":{"content": str}, "reply_to_id": number|null} # Expect JSON: {"type":"text","data":{"content": str}, "reply_to_id": number|null}
try: try:
obj = json.loads(payload) obj = json.loads(payload)
content = obj.get("content", "") content = obj.get("content", "")
reply_to_id = obj.get("reply_to_id", None) reply_to_id = obj.get("reply_to_id", None)
request = SendMessageRequest(content=content, reply_to_id=reply_to_id) message_request = SendMessageRequest(content=content, reply_to_id=reply_to_id)
except Exception: except Exception:
raise HTTPException(status_code=400, detail="Invalid payload JSON") raise HTTPException(status_code=400, detail="Invalid payload JSON")
if request.reply_to_id: if not message_request:
raise HTTPException(status_code=400, detail="Missing request data")
if message_request.reply_to_id:
# Check if the message being replied to exists # Check if the message being replied to exists
original_message = db.query(Message).filter(Message.id == request.reply_to_id).first() original_message = db.query(Message).filter(Message.id == message_request.reply_to_id).first()
if not original_message: if not original_message:
raise HTTPException(status_code=404, detail="Original message not found") raise HTTPException(status_code=404, detail="Original message not found")
raw_content = request.content.strip() raw_content = message_request.content.strip()
if not raw_content: if not raw_content:
raise HTTPException( raise HTTPException(
@@ -299,7 +303,7 @@ async def send_message(
new_message = Message( new_message = Message(
content=escaped_content, content=escaped_content,
user_id=current_user.id, user_id=current_user.id,
reply_to_id=request.reply_to_id, reply_to_id=message_request.reply_to_id,
timestamp=datetime.now() timestamp=datetime.now()
) )
@@ -412,6 +416,7 @@ async def get_messages(db: Session = Depends(get_db)):
@router.post("/dm/send") @router.post("/dm/send")
@rate_limit_per_user("20/minute") @rate_limit_per_user("20/minute")
async def dm_send( async def dm_send(
request: Request,
payload: dict | None = None, payload: dict | None = None,
current_user: User = Depends(get_current_user), current_user: User = Depends(get_current_user),
db: Session = Depends(get_db), db: Session = Depends(get_db),
@@ -624,8 +629,9 @@ async def get_dm_conversations(current_user: User = Depends(get_current_user), d
@router.put("/edit_message/{message_id}") @router.put("/edit_message/{message_id}")
@rate_limit_per_user("20/minute") @rate_limit_per_user("20/minute")
async def edit_message( async def edit_message(
request: Request,
message_id: int, message_id: int,
request: EditMessageRequest, edit_request: EditMessageRequest,
current_user: User = Depends(get_current_user), current_user: User = Depends(get_current_user),
db: Session = Depends(get_db) db: Session = Depends(get_db)
): ):
@@ -635,7 +641,7 @@ async def edit_message(
raise HTTPException(status_code=404, detail="Message not found") raise HTTPException(status_code=404, detail="Message not found")
if message.user_id != current_user.id: if message.user_id != current_user.id:
raise HTTPException(status_code=403, detail="You can only edit your own messages") raise HTTPException(status_code=403, detail="You can only edit your own messages")
raw_content = request.content.strip() raw_content = edit_request.content.strip()
if not raw_content: if not raw_content:
raise HTTPException(status_code=400, detail="Message content cannot be empty") raise HTTPException(status_code=400, detail="Message content cannot be empty")
@@ -700,20 +706,21 @@ async def delete_message(
@router.post("/add_reaction") @router.post("/add_reaction")
@rate_limit_per_user("50/minute") @rate_limit_per_user("50/minute")
async def add_reaction( async def add_reaction(
request: ReactionRequest, request: Request,
reaction_request: ReactionRequest,
current_user: User = Depends(get_current_user), current_user: User = Depends(get_current_user),
db: Session = Depends(get_db) db: Session = Depends(get_db)
): ):
# Check if message exists # Check if message exists
message = db.query(Message).filter(Message.id == request.message_id).first() message = db.query(Message).filter(Message.id == reaction_request.message_id).first()
if not message: if not message:
raise HTTPException(status_code=404, detail="Message not found") raise HTTPException(status_code=404, detail="Message not found")
# Check if reaction already exists # Check if reaction already exists
existing_reaction = db.query(Reaction).filter( existing_reaction = db.query(Reaction).filter(
Reaction.message_id == request.message_id, Reaction.message_id == reaction_request.message_id,
Reaction.user_id == current_user.id, Reaction.user_id == current_user.id,
Reaction.emoji == request.emoji Reaction.emoji == reaction_request.emoji
).first() ).first()
if existing_reaction: if existing_reaction:
@@ -723,9 +730,9 @@ async def add_reaction(
else: else:
# Add new reaction # Add new reaction
new_reaction = Reaction( new_reaction = Reaction(
message_id=request.message_id, message_id=reaction_request.message_id,
user_id=current_user.id, user_id=current_user.id,
emoji=request.emoji emoji=reaction_request.emoji
) )
db.add(new_reaction) db.add(new_reaction)
action = "added" action = "added"
@@ -742,8 +749,8 @@ async def add_reaction(
await messagingManager.broadcast({ await messagingManager.broadcast({
"type": "reactionUpdate", "type": "reactionUpdate",
"data": { "data": {
"message_id": request.message_id, "message_id": reaction_request.message_id,
"emoji": request.emoji, "emoji": reaction_request.emoji,
"action": action, "action": action,
"user_id": current_user.id, "user_id": current_user.id,
"username": current_user.username, "username": current_user.username,
@@ -755,11 +762,11 @@ async def add_reaction(
log_public_chat( log_public_chat(
"reaction_update", "reaction_update",
message_id=request.message_id, message_id=reaction_request.message_id,
user_id=current_user.id, user_id=current_user.id,
username=current_user.username, username=current_user.username,
action=action, action=action,
emoji=request.emoji, emoji=reaction_request.emoji,
) )
return {"status": "success", "action": action, "reactions": message_data["reactions"]} return {"status": "success", "action": action, "reactions": message_data["reactions"]}
@@ -768,12 +775,13 @@ async def add_reaction(
@router.post("/dm/add_reaction") @router.post("/dm/add_reaction")
@rate_limit_per_user("50/minute") @rate_limit_per_user("50/minute")
async def add_dm_reaction( async def add_dm_reaction(
request: DMReactionRequest, request: Request,
reaction_request: DMReactionRequest,
current_user: User = Depends(get_current_user), current_user: User = Depends(get_current_user),
db: Session = Depends(get_db) db: Session = Depends(get_db)
): ):
# Check if DM envelope exists # Check if DM envelope exists
envelope = db.query(DMEnvelope).filter(DMEnvelope.id == request.dm_envelope_id).first() envelope = db.query(DMEnvelope).filter(DMEnvelope.id == reaction_request.dm_envelope_id).first()
if not envelope: if not envelope:
raise HTTPException(status_code=404, detail="DM envelope not found") raise HTTPException(status_code=404, detail="DM envelope not found")
@@ -783,9 +791,9 @@ async def add_dm_reaction(
# Check if reaction already exists # Check if reaction already exists
existing_reaction = db.query(DMReaction).filter( existing_reaction = db.query(DMReaction).filter(
DMReaction.dm_envelope_id == request.dm_envelope_id, DMReaction.dm_envelope_id == reaction_request.dm_envelope_id,
DMReaction.user_id == current_user.id, DMReaction.user_id == current_user.id,
DMReaction.emoji == request.emoji DMReaction.emoji == reaction_request.emoji
).first() ).first()
if existing_reaction: if existing_reaction:
@@ -795,9 +803,9 @@ async def add_dm_reaction(
else: else:
# Add new reaction # Add new reaction
new_reaction = DMReaction( new_reaction = DMReaction(
dm_envelope_id=request.dm_envelope_id, dm_envelope_id=reaction_request.dm_envelope_id,
user_id=current_user.id, user_id=current_user.id,
emoji=request.emoji emoji=reaction_request.emoji
) )
db.add(new_reaction) db.add(new_reaction)
action = "added" action = "added"
@@ -814,8 +822,8 @@ async def add_dm_reaction(
await messagingManager.broadcast({ await messagingManager.broadcast({
"type": "dmReactionUpdate", "type": "dmReactionUpdate",
"data": { "data": {
"dm_envelope_id": request.dm_envelope_id, "dm_envelope_id": reaction_request.dm_envelope_id,
"emoji": request.emoji, "emoji": reaction_request.emoji,
"action": action, "action": action,
"user_id": current_user.id, "user_id": current_user.id,
"username": current_user.username, "username": current_user.username,
@@ -827,11 +835,11 @@ async def add_dm_reaction(
log_dm( log_dm(
"reaction_update", "reaction_update",
dm_envelope_id=request.dm_envelope_id, dm_envelope_id=reaction_request.dm_envelope_id,
user_id=current_user.id, user_id=current_user.id,
username=current_user.username, username=current_user.username,
action=action, action=action,
emoji=request.emoji, emoji=reaction_request.emoji,
) )
return {"status": "success", "action": action, "reactions": envelope_data["reactions"]} return {"status": "success", "action": action, "reactions": envelope_data["reactions"]}
+14 -10
View File
@@ -7,6 +7,7 @@ from PIL import Image
import os import os
import uuid import uuid
import io import io
from fastapi import Request
from dependencies import get_db, get_current_user from dependencies import get_db, get_current_user
from models import User, UpdateBioRequest, UserProfileResponse from models import User, UpdateBioRequest, UserProfileResponse
@@ -42,6 +43,7 @@ os.makedirs(PROFILE_PICTURES_DIR, exist_ok=True)
@router.post("/upload-profile-picture") @router.post("/upload-profile-picture")
@rate_limit_per_user("10/minute") @rate_limit_per_user("10/minute")
async def upload_profile_picture( async def upload_profile_picture(
request: Request,
profile_picture: UploadFile = File(...), profile_picture: UploadFile = File(...),
current_user: User = Depends(get_current_user), current_user: User = Depends(get_current_user),
db: Session = Depends(get_db) db: Session = Depends(get_db)
@@ -167,7 +169,8 @@ async def list_users(
@router.put("/user/profile") @router.put("/user/profile")
@rate_limit_per_user("10/minute") @rate_limit_per_user("10/minute")
async def update_user_profile( async def update_user_profile(
request: UpdateProfileRequest, request: Request,
update_request: UpdateProfileRequest,
current_user: User = Depends(get_current_user), current_user: User = Depends(get_current_user),
db: Session = Depends(get_db) db: Session = Depends(get_db)
): ):
@@ -177,8 +180,8 @@ async def update_user_profile(
updated = False updated = False
# Update username if provided # Update username if provided
if request.username is not None: if update_request.username is not None:
username = request.username.strip() username = update_request.username.strip()
if not is_valid_username(username): if not is_valid_username(username):
raise HTTPException( raise HTTPException(
status_code=400, status_code=400,
@@ -199,8 +202,8 @@ async def update_user_profile(
updated = True updated = True
# Update display name if provided # Update display name if provided
if request.display_name is not None: if update_request.display_name is not None:
display_name = request.display_name.strip() display_name = update_request.display_name.strip()
if not is_valid_display_name(display_name): if not is_valid_display_name(display_name):
raise HTTPException( raise HTTPException(
status_code=400, status_code=400,
@@ -216,8 +219,8 @@ async def update_user_profile(
updated = True updated = True
# Update bio if provided # Update bio if provided
if request.description is not None: if update_request.description is not None:
bio = request.description.strip() bio = update_request.description.strip()
if len(bio) > 500: if len(bio) > 500:
raise HTTPException(status_code=400, detail="Bio must be 500 characters or less") raise HTTPException(status_code=400, detail="Bio must be 500 characters or less")
@@ -244,17 +247,18 @@ async def update_user_profile(
@router.put("/user/bio") @router.put("/user/bio")
@rate_limit_per_user("10/minute") @rate_limit_per_user("10/minute")
async def update_user_bio( async def update_user_bio(
request: UpdateBioRequest, request: Request,
bio_request: UpdateBioRequest,
current_user: User = Depends(get_current_user), current_user: User = Depends(get_current_user),
db: Session = Depends(get_db) db: Session = Depends(get_db)
): ):
""" """
Update current user's bio Update current user's bio
""" """
if len(request.bio) > 500: # Limit bio to 500 characters if len(bio_request.bio) > 500: # Limit bio to 500 characters
raise HTTPException(status_code=400, detail="Bio must be 500 characters or less") raise HTTPException(status_code=400, detail="Bio must be 500 characters or less")
current_user.bio = request.bio.strip() current_user.bio = bio_request.bio.strip()
db.commit() db.commit()
return { return {