Improve websocket code

This commit is contained in:
2025-08-18 16:37:50 +03:00
Unverified
parent 4e54b8bec2
commit 6ecd885e7a
+73 -23
View File
@@ -1,5 +1,6 @@
from datetime import datetime from datetime import datetime
from fastapi import APIRouter, Depends, HTTPException, WebSocket from email.policy import HTTP
from fastapi import APIRouter, Depends, HTTPException, WebSocket, WebSocketDisconnect
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from dependencies import get_current_user, get_db from dependencies import get_current_user, get_db
from models import Message, SendMessageRequest from models import Message, SendMessageRequest
@@ -7,32 +8,29 @@ from models import Message, SendMessageRequest
router = APIRouter() router = APIRouter()
async def get_messages_inner(current_user: dict, db: Session): def convert_message(msg: Message, current_user: dict) -> dict:
messages = db.query(Message).order_by(Message.timestamp.asc()).all() return {
messages_data = []
for msg in messages:
messages_data.append({
"id": msg.id, "id": msg.id,
"content": msg.content, "content": msg.content,
"timestamp": msg.timestamp.isoformat(), "timestamp": msg.timestamp.isoformat(),
"is_author": msg.user_id == current_user["user_id"], "is_author": msg.user_id == current_user["user_id"],
"is_read": msg.is_read, "is_read": msg.is_read,
"username": msg.author.username "username": msg.author.username
}) }
async def get_messages_inner(current_user: dict, db: Session):
messages = db.query(Message).order_by(Message.timestamp.asc()).all()
messages_data = []
for msg in messages:
messages_data.append(convert_message(msg, current_user))
return { return {
"status": "success", "status": "success",
"messages": messages_data "messages": messages_data
} }
async def send_message_inner(request: SendMessageRequest, current_user: dict, db: Session):
@router.post("/send_message")
async def send_message(
request: SendMessageRequest,
current_user: dict = Depends(get_current_user),
db: Session = Depends(get_db)
):
if not request.content.strip(): if not request.content.strip():
raise HTTPException( raise HTTPException(
status_code=400, status_code=400,
@@ -49,7 +47,15 @@ async def send_message(
db.commit() db.commit()
db.refresh(new_message) db.refresh(new_message)
return {"status": "success"} return {"status": "success", "message": convert_message(new_message, current_user)}
@router.post("/send_message")
async def send_message(
request: SendMessageRequest,
current_user: dict = Depends(get_current_user),
db: Session = Depends(get_db)
):
return await send_message_inner(request, current_user, db)
@router.get("/get_messages") @router.get("/get_messages")
@@ -60,10 +66,17 @@ async def get_messages(
return await get_messages_inner(current_user, db) return await get_messages_inner(current_user, db)
@router.websocket("/chat/ws") class MessaggingSocketManager:
async def messaging(websocket: WebSocket): def __init__(self) -> None:
await websocket.accept() self.connections: list[WebSocket] = []
def get_dependencies(self):
return get_current_user(), next(get_db())
async def send_error(self, websocket: WebSocket, type: str, e: HTTPException):
await websocket.send_json({"type": type, "error": {"code": e.status_code, "detail": e.detail}})
async def handle_connection(self, websocket: WebSocket):
while True: while True:
data = await websocket.receive_json() data = await websocket.receive_json()
@@ -71,11 +84,48 @@ async def messaging(websocket: WebSocket):
await websocket.send_json({"type": "ping", "data": {"status": "success"}}) await websocket.send_json({"type": "ping", "data": {"status": "success"}})
elif data.type == "getMessages": elif data.type == "getMessages":
try: try:
current_user = get_current_user() current_user, db = self.get_dependencies()
db = next(get_db())
await websocket.send_json({"type": "getMessages", "data": await get_messages_inner(current_user, db)}) await websocket.send_json({"type": data.type, "data": await get_messages_inner(current_user, db)})
except HTTPException as e: except HTTPException as e:
await websocket.send_json({"type": "getMessages", "error": {"code": e.status_code, "detail": e.detail}}) await self.send_error(websocket, data.type, e)
elif data.type == "sendMessage":
try:
current_user, db = self.get_dependencies()
request: SendMessageRequest = SendMessageRequest.model_validate(data.data)
response = await send_message_inner(request, current_user, db)
await self.broadcast({
"type": "newMessage",
"data": response.message
})
await websocket.send_json({"type": data.type, "data": response})
except HTTPException as e:
await self.send_error(websocket, data.type, e)
else: else:
await websocket.send_json({"type": data.type, "error": {"code": 400, "detail": "Invalid type"}}) await websocket.send_json({"type": data.type, "error": {"code": 400, "detail": "Invalid type"}})
async def disconnect(self, websocket: WebSocket, code: int = 1000, message: str | None = None):
try:
await websocket.close(code=code, reason=message)
finally:
self.connections.remove(websocket)
async def connect(self, websocket: WebSocket):
await websocket.accept()
self.connections.append(websocket)
try:
await self.handle_connection(websocket)
finally:
self.connections.remove(websocket)
async def broadcast(self, message: dict):
for websocket in self.connections:
await websocket.send_json(message)
messagingManager = MessaggingSocketManager()
@router.websocket("/chat/ws")
async def messaging(websocket: WebSocket):
await messagingManager.connect(websocket)