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")
@rate_limit_per_ip("5/minute")
def login(request: LoginRequest, http: Request, db: Session = Depends(get_db)):
username = request.username.strip()
client_ip = get_client_ip(http)
raw_ua = http.headers.get("user-agent")
def login(request: Request, login_request: LoginRequest, db: Session = Depends(get_db)):
username = login_request.username.strip()
client_ip = get_client_ip(request)
raw_ua = request.headers.get("user-agent")
if is_user_agent_blocked(raw_ua):
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()
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(
"login_failed",
severity="warning",
@@ -125,8 +125,8 @@ def login(request: LoginRequest, http: Request, db: Session = Depends(get_db)):
)
# Create device session and embed into JWT
raw_ua = http.headers.get("user-agent")
device_name = http.headers.get("x-device-name")
raw_ua = request.headers.get("user-agent")
device_name = request.headers.get("x-device-name")
ua = parse_ua(raw_ua or "")
session_id = uuid.uuid4().hex
@@ -181,13 +181,13 @@ def login(request: LoginRequest, http: Request, db: Session = Depends(get_db)):
@router.post("/register")
@rate_limit_per_ip("3/hour")
def register(request: RegisterRequest, http: Request, db: Session = Depends(get_db)):
username = request.username.strip()
display_name = request.display_name.strip()
password = request.password.strip()
confirm_password = request.confirm_password.strip()
client_ip = get_client_ip(http)
raw_ua = http.headers.get("user-agent")
def register(request: Request, register_request: RegisterRequest, db: Session = Depends(get_db)):
username = register_request.username.strip()
display_name = register_request.display_name.strip()
password = register_request.password.strip()
confirm_password = register_request.confirm_password.strip()
client_ip = get_client_ip(request)
raw_ua = request.headers.get("user-agent")
if is_user_agent_blocked(raw_ua):
log_security(
@@ -281,8 +281,8 @@ def register(request: RegisterRequest, http: Request, db: Session = Depends(get_
db.refresh(new_user)
# Create initial device session
raw_ua = http.headers.get("user-agent")
device_name = http.headers.get("x-device-name")
raw_ua = request.headers.get("user-agent")
device_name = request.headers.get("x-device-name")
ua = parse_ua(raw_ua or "")
session_id = uuid.uuid4().hex
device = DeviceSession(
@@ -447,22 +447,22 @@ def logout(
@router.post("/change-password")
@rate_limit_per_user("5/hour")
def change_password(
request: ChangePasswordRequest,
http: Request,
request: Request,
password_request: ChangePasswordRequest,
credentials: HTTPAuthorizationCredentials = Depends(HTTPBearer()),
current_user: User = Depends(get_current_user),
db: Session = Depends(get_db)
):
# 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="Текущий пароль неверный")
# 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()
# Optionally revoke all other sessions, keeping the current one
if request.logoutAllExceptCurrent:
if password_request.logoutAllExceptCurrent:
from utils import verify_token as _verify_token
payload = _verify_token(credentials.credentials)
if not payload:
@@ -474,13 +474,13 @@ def change_password(
).update({DeviceSession.revoked: True})
db.commit()
client_ip = get_client_ip(http)
client_ip = get_client_ip(request)
log_security(
"password_changed",
username=current_user.username,
user_id=current_user.id,
ip=client_ip,
logout_others=bool(request.logoutAllExceptCurrent),
logout_others=bool(password_request.logoutAllExceptCurrent),
)
return {"status": "success"}
+38 -30
View File
@@ -12,7 +12,7 @@ from collections import defaultdict, deque
from difflib import SequenceMatcher
from types import SimpleNamespace
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.security import HTTPAuthorizationCredentials
from sqlalchemy.orm import Session
@@ -254,7 +254,8 @@ def convert_dm_envelope(envelope: DMEnvelope) -> dict:
@router.post("/send_message")
@rate_limit_per_user("30/minute")
async def send_message(
request: SendMessageRequest | None = None,
request: Request,
message_request: SendMessageRequest | None = None,
current_user: User = Depends(get_current_user),
db: Session = Depends(get_db),
# Optional multipart form support
@@ -262,23 +263,26 @@ async def send_message(
files: list[UploadFile] = File(default=[]),
):
# 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}
try:
obj = json.loads(payload)
content = obj.get("content", "")
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:
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
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:
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:
raise HTTPException(
@@ -299,7 +303,7 @@ async def send_message(
new_message = Message(
content=escaped_content,
user_id=current_user.id,
reply_to_id=request.reply_to_id,
reply_to_id=message_request.reply_to_id,
timestamp=datetime.now()
)
@@ -412,6 +416,7 @@ async def get_messages(db: Session = Depends(get_db)):
@router.post("/dm/send")
@rate_limit_per_user("20/minute")
async def dm_send(
request: Request,
payload: dict | None = None,
current_user: User = Depends(get_current_user),
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}")
@rate_limit_per_user("20/minute")
async def edit_message(
request: Request,
message_id: int,
request: EditMessageRequest,
edit_request: EditMessageRequest,
current_user: User = Depends(get_current_user),
db: Session = Depends(get_db)
):
@@ -635,7 +641,7 @@ async def edit_message(
raise HTTPException(status_code=404, detail="Message not found")
if message.user_id != current_user.id:
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:
raise HTTPException(status_code=400, detail="Message content cannot be empty")
@@ -700,20 +706,21 @@ async def delete_message(
@router.post("/add_reaction")
@rate_limit_per_user("50/minute")
async def add_reaction(
request: ReactionRequest,
request: Request,
reaction_request: ReactionRequest,
current_user: User = Depends(get_current_user),
db: Session = Depends(get_db)
):
# 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:
raise HTTPException(status_code=404, detail="Message not found")
# Check if reaction already exists
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.emoji == request.emoji
Reaction.emoji == reaction_request.emoji
).first()
if existing_reaction:
@@ -723,9 +730,9 @@ async def add_reaction(
else:
# Add new reaction
new_reaction = Reaction(
message_id=request.message_id,
message_id=reaction_request.message_id,
user_id=current_user.id,
emoji=request.emoji
emoji=reaction_request.emoji
)
db.add(new_reaction)
action = "added"
@@ -742,8 +749,8 @@ async def add_reaction(
await messagingManager.broadcast({
"type": "reactionUpdate",
"data": {
"message_id": request.message_id,
"emoji": request.emoji,
"message_id": reaction_request.message_id,
"emoji": reaction_request.emoji,
"action": action,
"user_id": current_user.id,
"username": current_user.username,
@@ -755,11 +762,11 @@ async def add_reaction(
log_public_chat(
"reaction_update",
message_id=request.message_id,
message_id=reaction_request.message_id,
user_id=current_user.id,
username=current_user.username,
action=action,
emoji=request.emoji,
emoji=reaction_request.emoji,
)
return {"status": "success", "action": action, "reactions": message_data["reactions"]}
@@ -768,12 +775,13 @@ async def add_reaction(
@router.post("/dm/add_reaction")
@rate_limit_per_user("50/minute")
async def add_dm_reaction(
request: DMReactionRequest,
request: Request,
reaction_request: DMReactionRequest,
current_user: User = Depends(get_current_user),
db: Session = Depends(get_db)
):
# 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:
raise HTTPException(status_code=404, detail="DM envelope not found")
@@ -783,9 +791,9 @@ async def add_dm_reaction(
# Check if reaction already exists
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.emoji == request.emoji
DMReaction.emoji == reaction_request.emoji
).first()
if existing_reaction:
@@ -795,9 +803,9 @@ async def add_dm_reaction(
else:
# Add new reaction
new_reaction = DMReaction(
dm_envelope_id=request.dm_envelope_id,
dm_envelope_id=reaction_request.dm_envelope_id,
user_id=current_user.id,
emoji=request.emoji
emoji=reaction_request.emoji
)
db.add(new_reaction)
action = "added"
@@ -814,8 +822,8 @@ async def add_dm_reaction(
await messagingManager.broadcast({
"type": "dmReactionUpdate",
"data": {
"dm_envelope_id": request.dm_envelope_id,
"emoji": request.emoji,
"dm_envelope_id": reaction_request.dm_envelope_id,
"emoji": reaction_request.emoji,
"action": action,
"user_id": current_user.id,
"username": current_user.username,
@@ -827,11 +835,11 @@ async def add_dm_reaction(
log_dm(
"reaction_update",
dm_envelope_id=request.dm_envelope_id,
dm_envelope_id=reaction_request.dm_envelope_id,
user_id=current_user.id,
username=current_user.username,
action=action,
emoji=request.emoji,
emoji=reaction_request.emoji,
)
return {"status": "success", "action": action, "reactions": envelope_data["reactions"]}
+14 -10
View File
@@ -7,6 +7,7 @@ from PIL import Image
import os
import uuid
import io
from fastapi import Request
from dependencies import get_db, get_current_user
from models import User, UpdateBioRequest, UserProfileResponse
@@ -42,6 +43,7 @@ os.makedirs(PROFILE_PICTURES_DIR, exist_ok=True)
@router.post("/upload-profile-picture")
@rate_limit_per_user("10/minute")
async def upload_profile_picture(
request: Request,
profile_picture: UploadFile = File(...),
current_user: User = Depends(get_current_user),
db: Session = Depends(get_db)
@@ -167,7 +169,8 @@ async def list_users(
@router.put("/user/profile")
@rate_limit_per_user("10/minute")
async def update_user_profile(
request: UpdateProfileRequest,
request: Request,
update_request: UpdateProfileRequest,
current_user: User = Depends(get_current_user),
db: Session = Depends(get_db)
):
@@ -177,8 +180,8 @@ async def update_user_profile(
updated = False
# Update username if provided
if request.username is not None:
username = request.username.strip()
if update_request.username is not None:
username = update_request.username.strip()
if not is_valid_username(username):
raise HTTPException(
status_code=400,
@@ -199,8 +202,8 @@ async def update_user_profile(
updated = True
# Update display name if provided
if request.display_name is not None:
display_name = request.display_name.strip()
if update_request.display_name is not None:
display_name = update_request.display_name.strip()
if not is_valid_display_name(display_name):
raise HTTPException(
status_code=400,
@@ -216,8 +219,8 @@ async def update_user_profile(
updated = True
# Update bio if provided
if request.description is not None:
bio = request.description.strip()
if update_request.description is not None:
bio = update_request.description.strip()
if len(bio) > 500:
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")
@rate_limit_per_user("10/minute")
async def update_user_bio(
request: UpdateBioRequest,
request: Request,
bio_request: UpdateBioRequest,
current_user: User = Depends(get_current_user),
db: Session = Depends(get_db)
):
"""
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")
current_user.bio = request.bio.strip()
current_user.bio = bio_request.bio.strip()
db.commit()
return {