mirror of
https://github.com/fromchat-messenger/web.git
synced 2026-09-22 19:15:08 +03:00
Improve websocket code
This commit is contained in:
+73
-23
@@ -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)
|
||||||
Reference in New Issue
Block a user