Fix suspension status

This commit is contained in:
2026-04-21 16:25:27 +03:00
Unverified
parent 3079ff944e
commit b7d884df58
8 changed files with 211 additions and 107 deletions
+16 -8
View File
@@ -9,7 +9,7 @@ from user_agents import parse as parse_ua
from fastapi.security import HTTPBearer, HTTPAuthorizationCredentials
from ..constants import OWNER_USERNAME
from ..dependencies import get_current_user, get_db
from ..dependencies import get_current_user, get_current_user_allow_suspended, get_db
from ..models import LoginRequest, RegisterRequest, ChangePasswordRequest, User, CryptoPublicKey, CryptoBackup, DeviceSession
from ..utils import create_token, get_password_hash, verify_password, get_client_ip
from ..validation import is_valid_password, is_valid_username, is_valid_display_name
@@ -18,6 +18,7 @@ import os
from ..security.audit import log_security
from ..security.profanity import contains_profanity
from ..security.rate_limit import rate_limit_per_ip
from ..key_lifecycle import destroy_message_keys_for_user
router = APIRouter()
_FAILED_ATTEMPT_WINDOW_SECONDS = 300
@@ -314,13 +315,13 @@ def register(
}
@router.get("/crypto/public-key")
def get_public_key(current_user: User = Depends(get_current_user), db: Session = Depends(get_db)):
def get_public_key(current_user: User = Depends(get_current_user_allow_suspended), db: Session = Depends(get_db)):
row = db.query(CryptoPublicKey).filter(CryptoPublicKey.user_id == current_user.id).first()
return {"publicKey": row.public_key_b64 if row else None}
@router.post("/crypto/public-key")
def set_public_key(payload: dict, current_user: User = Depends(get_current_user), db: Session = Depends(get_db)):
def set_public_key(payload: dict, current_user: User = Depends(get_current_user_allow_suspended), db: Session = Depends(get_db)):
pk = payload.get("publicKey")
if not pk:
raise HTTPException(status_code=400, detail="publicKey required")
@@ -337,13 +338,13 @@ def set_public_key(payload: dict, current_user: User = Depends(get_current_user)
@router.get("/crypto/backup")
def get_backup(current_user: User = Depends(get_current_user), db: Session = Depends(get_db)):
def get_backup(current_user: User = Depends(get_current_user_allow_suspended), db: Session = Depends(get_db)):
row = db.query(CryptoBackup).filter(CryptoBackup.user_id == current_user.id).first()
return {"blob": row.blob_json if row else None}
@router.post("/crypto/backup")
def set_backup(payload: dict, current_user: User = Depends(get_current_user), db: Session = Depends(get_db)):
def set_backup(payload: dict, current_user: User = Depends(get_current_user_allow_suspended), db: Session = Depends(get_db)):
blob = payload.get("blob")
if not blob:
raise HTTPException(status_code=400, detail="blob required")
@@ -475,7 +476,7 @@ def change_password(
@router.get("/users")
@rate_limit_per_ip("30/minute") # Per-IP limit to prevent abuse
def list_users(request: Request, current_user: User = Depends(get_current_user), db: Session = Depends(get_db)):
def list_users(request: Request, current_user: User = Depends(get_current_user_allow_suspended), db: Session = Depends(get_db)):
users = db.query(User).order_by(User.username.asc()).all()
return {
"users": [
@@ -486,14 +487,19 @@ def list_users(request: Request, current_user: User = Depends(get_current_user),
@router.get("/crypto/public-key/of/{user_id}")
@rate_limit_per_ip("100/minute") # Per-IP limit to prevent abuse
def get_public_key_of(request: Request, user_id: int, current_user: User = Depends(get_current_user), db: Session = Depends(get_db)):
def get_public_key_of(
request: Request,
user_id: int,
current_user: User = Depends(get_current_user_allow_suspended),
db: Session = Depends(get_db),
):
row = db.query(CryptoPublicKey).filter(CryptoPublicKey.user_id == user_id).first()
return {"publicKey": row.public_key_b64 if row else None}
@router.get("/users/search")
@rate_limit_per_ip("60/minute") # Per-IP limit to prevent abuse
def search_users(request: Request, q: str, current_user: User = Depends(get_current_user), db: Session = Depends(get_db)):
def search_users(request: Request, q: str, current_user: User = Depends(get_current_user_allow_suspended), db: Session = Depends(get_db)):
if len(q.strip()) < 2:
return {"users": []}
@@ -554,6 +560,8 @@ async def _delete_user_data(user: User, db: Session):
if has_user_id:
# Delete all records for this user
db.execute(text(f"DELETE FROM {table_name} WHERE user_id = :uid"), {"uid": user_id})
destroy_message_keys_for_user(db, user_id, commit=False)
db.commit()
except Exception as e:
@@ -23,7 +23,7 @@ from pydantic import BaseModel, Field
from ..db import get_db
from ..models import User, DMEnvelope, DMFile, DMEditHistory, EditMessageRequest
from ..dependencies import get_current_user
from ..dependencies import get_current_user, get_current_user_allow_suspended
from ..security.audit import log_security
from ..service_calls import (
get_messaging_transport_public_key,
@@ -538,7 +538,7 @@ async def get_encrypted_conversation(
other_user_id: int,
limit: int = 50,
offset: int = 0,
current_user: User = Depends(get_current_user),
current_user: User = Depends(get_current_user_allow_suspended),
db: Session = Depends(get_db),
):
"""
+13 -9
View File
@@ -14,7 +14,7 @@ from typing import Any
import httpx
from fastapi import APIRouter, Depends, HTTPException, WebSocket, WebSocketDisconnect, UploadFile, File, Form, Request, status
from sqlalchemy.orm import Session
from ..dependencies import get_current_user, get_db
from ..dependencies import get_current_user, get_current_user_allow_suspended, get_db
from .account import convert_user
from ..constants import OWNER_USERNAME
from ..models import Message, SendMessageRequest, EditMessageRequest, User, DMEnvelope, MessageFile, DMFile, Reaction, ReactionRequest, ReactionResponse, DMReaction, DMReactionRequest, DMReactionResponse, UpdateLog, MessageEditHistory, MessageEditHistoryResponse
@@ -607,7 +607,7 @@ async def push_test(request: Request, current_user: User = Depends(get_current_u
@router.get("/get_messages")
@rate_limit_per_ip("60/minute") # Per-IP limit to prevent abuse
async def get_messages(request: Request, current_user: User = Depends(get_current_user), db: Session = Depends(get_db)):
async def get_messages(request: Request, current_user: User = Depends(get_current_user_allow_suspended), db: Session = Depends(get_db)):
messages = db.query(Message).order_by(Message.timestamp.asc()).all()
messages_data = []
@@ -626,7 +626,7 @@ class MarkReadRequest(BaseModel):
@router.get("/messages/new")
@rate_limit_per_ip("60/minute")
async def get_new_messages(request: Request, current_user: User = Depends(get_current_user), db: Session = Depends(get_db)):
async def get_new_messages(request: Request, current_user: User = Depends(get_current_user_allow_suspended), db: Session = Depends(get_db)):
"""
Return unread public messages (Message.is_read == False).
"""
@@ -661,7 +661,7 @@ async def mark_messages_read(request: Request, read_request: MarkReadRequest, cu
@router.get("/dm/fetch")
@rate_limit_per_ip("60/minute") # Per-IP limit to prevent abuse
async def dm_fetch(request: Request, since: int | None = None, current_user: User = Depends(get_current_user), db: Session = Depends(get_db)):
async def dm_fetch(request: Request, since: int | None = None, current_user: User = Depends(get_current_user_allow_suspended), db: Session = Depends(get_db)):
envelopes = db.query(DMEnvelope).filter(DMEnvelope.recipient_id == current_user.id)
if since:
envelopes = envelopes.filter(DMEnvelope.id > since)
@@ -675,7 +675,7 @@ async def dm_fetch(request: Request, since: int | None = None, current_user: Use
@router.get("/dm/history/{other_user_id}")
@rate_limit_per_ip("60/minute") # Per-IP limit to prevent abuse
async def dm_history(request: Request, other_user_id: int, current_user: User = Depends(get_current_user), db: Session = Depends(get_db)):
async def dm_history(request: Request, other_user_id: int, current_user: User = Depends(get_current_user_allow_suspended), db: Session = Depends(get_db)):
if other_user_id <= 0:
raise HTTPException(status_code=400, detail="Invalid user ID")
@@ -684,7 +684,7 @@ async def dm_history(request: Request, other_user_id: int, current_user: User =
# Verify other user exists
other_user = db.query(User).filter(User.id == other_user_id).first()
if not other_user or other_user.deleted or other_user.suspended:
if not other_user:
raise HTTPException(status_code=404, detail="User not found")
envelopes = db.query(DMEnvelope).filter(
@@ -700,7 +700,7 @@ async def dm_history(request: Request, other_user_id: int, current_user: User =
@router.get("/dm/conversations")
@rate_limit_per_ip("60/minute") # Per-IP limit to prevent abuse
async def get_dm_conversations(request: Request, current_user: User = Depends(get_current_user), db: Session = Depends(get_db)):
async def get_dm_conversations(request: Request, current_user: User = Depends(get_current_user_allow_suspended), db: Session = Depends(get_db)):
# Get all DM conversations where current user is involved
conversations_query = db.query(DMEnvelope).filter(
(DMEnvelope.sender_id == current_user.id) | (DMEnvelope.recipient_id == current_user.id)
@@ -1361,6 +1361,10 @@ class MessaggingSocketManager:
"reason": reason
})
async def send_unsuspension_to_user(self, user_id: int):
"""Send unsuspension message to user's WebSocket connections (as batched update)"""
await self.send_update_to_user(user_id, "unsuspended", {})
async def send_deletion_to_user(self, user_id: int):
"""Send account deletion message to user's WebSocket connections (as batched update)"""
await self.send_update_to_user(user_id, "account_deleted", {})
@@ -1478,7 +1482,7 @@ import httpx
async def proxy_normal_file(
request: Request,
filename: str,
current_user: User = Depends(get_current_user)
current_user: User = Depends(get_current_user_allow_suspended)
):
"""Proxy file requests to file_storage service."""
mod = service_calls._get_file_storage_module()
@@ -1531,7 +1535,7 @@ async def test_proxy():
async def proxy_encrypted_file(
request: Request,
filename: str,
current_user: User = Depends(get_current_user)
current_user: User = Depends(get_current_user_allow_suspended)
):
"""Proxy file requests to file_storage service."""
mod = service_calls._get_file_storage_module()
+52 -67
View File
@@ -9,7 +9,7 @@ import uuid
import io
from fastapi import Request
from ..dependencies import get_db, get_current_user
from ..dependencies import get_current_user, get_current_user_allow_suspended, get_db
from ..models import User, UpdateBioRequest, UserProfileResponse
from pydantic import BaseModel
from ..validation import is_valid_username, is_valid_display_name
@@ -22,6 +22,40 @@ from ..security.rate_limit import rate_limit_per_ip
router = APIRouter()
def _build_user_profile_response(user: User, is_owner_request: bool = False) -> UserProfileResponse:
should_hide_profile = (not is_owner_request) and (user.deleted or user.suspended)
if not should_hide_profile:
return UserProfileResponse(
id=user.id,
username=user.username,
display_name=user.display_name,
profile_picture=user.profile_picture,
bio=user.bio,
online=user.online,
last_seen=user.last_seen,
created_at=user.created_at,
verified=user.verified,
suspended=user.suspended or False,
suspension_reason=user.suspension_reason,
deleted=user.deleted or False,
)
return UserProfileResponse(
id=user.id,
username="deleted",
display_name="Deleted User",
profile_picture=None,
bio=None,
online=False,
last_seen=None,
created_at=None,
verified=False,
suspended=False,
suspension_reason=None,
deleted=True,
)
def _ensure_owner_unsuspended(user: User | None, db: Session):
if user and user.id == 1 and user.suspended:
user.suspended = False
@@ -111,7 +145,7 @@ async def get_profile_picture(filename: str):
@router.get("/user/profile")
async def get_user_profile(
current_user: User = Depends(get_current_user),
current_user: User = Depends(get_current_user_allow_suspended),
db: Session = Depends(get_db)
):
"""
@@ -132,7 +166,7 @@ async def get_user_profile(
verified=current_user.verified,
suspended=current_user.suspended or False,
suspension_reason=current_user.suspension_reason,
deleted=(current_user.deleted or current_user.suspended) or False, # Treat suspended as deleted
deleted=current_user.deleted or False,
)
except Exception as e:
# Log and return a consistent HTTP 500 error with minimal details
@@ -278,7 +312,7 @@ async def update_user_bio(
@router.get("/user/stats/registered-count")
def get_registered_user_count(
current_user: User = Depends(get_current_user),
current_user: User = Depends(get_current_user_allow_suspended),
db: Session = Depends(get_db),
):
"""Number of registered accounts (non-deleted users)."""
@@ -289,6 +323,7 @@ def get_registered_user_count(
@router.get("/user/{username}")
async def get_user_by_username(
username: str,
current_user: User = Depends(get_current_user_allow_suspended),
db: Session = Depends(get_db)
):
"""
@@ -304,41 +339,13 @@ async def get_user_by_username(
_ensure_owner_unsuspended(user, db)
# Handle deleted or suspended users
if user.deleted or user.suspended:
return UserProfileResponse(
id=user.id,
username="deleted",
display_name="Deleted User",
profile_picture=None,
bio=None,
online=False,
last_seen=None, # Clear last seen timestamp
created_at=None, # Clear member since timestamp
verified=False,
suspended=False,
suspension_reason=None,
deleted=True
)
return UserProfileResponse(
id=user.id,
username=user.username,
display_name=user.display_name,
profile_picture=user.profile_picture,
bio=user.bio,
online=user.online,
last_seen=user.last_seen,
created_at=user.created_at,
verified=user.verified,
suspended=user.suspended or False,
suspension_reason=user.suspension_reason,
deleted=(user.deleted or user.suspended) or False, # Treat suspended as deleted
)
is_owner_request = current_user.id == user.id or current_user.id == 1
return _build_user_profile_response(user, is_owner_request=is_owner_request)
@router.get("/user/id/{user_id}")
async def get_user_by_id(
user_id: int,
current_user: User = Depends(get_current_user_allow_suspended),
db: Session = Depends(get_db)
):
"""
@@ -354,37 +361,8 @@ async def get_user_by_id(
_ensure_owner_unsuspended(user, db)
# Handle deleted or suspended users
if user.deleted or user.suspended:
return UserProfileResponse(
id=user.id,
username="deleted",
display_name="Deleted User",
profile_picture=None,
bio=None,
online=False,
last_seen=None, # Clear last seen timestamp
created_at=None, # Clear member since timestamp
verified=False,
suspended=False,
suspension_reason=None,
deleted=True
)
return UserProfileResponse(
id=user.id,
username=user.username,
display_name=user.display_name,
profile_picture=user.profile_picture,
bio=user.bio,
online=user.online,
last_seen=user.last_seen,
created_at=user.created_at,
verified=user.verified,
suspended=user.suspended or False,
suspension_reason=user.suspension_reason,
deleted=user.deleted or False
)
is_owner_request = current_user.id == user.id or current_user.id == 1
return _build_user_profile_response(user, is_owner_request=is_owner_request)
@router.post("/user/{user_id}/verify")
@@ -426,7 +404,7 @@ async def verify_user(
@router.get("/user/check-similarity/{user_id}")
async def check_user_similarity(
user_id: int,
current_user: User = Depends(get_current_user),
current_user: User = Depends(get_current_user_allow_suspended),
db: Session = Depends(get_db)
):
"""
@@ -531,6 +509,13 @@ async def unsuspend_user(
target_user.suspended = False
target_user.suspension_reason = None
db.commit()
# Send WebSocket unsuspension message
try:
await messagingManager.send_unsuspension_to_user(user_id)
except Exception:
# Log error but don't fail the request
pass
log_security(
"admin_unsuspend_user",