ws: вебсокеты на новой модели данных

This commit is contained in:
tsygankoviva 2026-05-20 18:35:26 +03:00
parent 06ad5e0e93
commit 33055d6695
6 changed files with 724 additions and 106 deletions

View File

@ -14,11 +14,10 @@ security = HTTPBearer()
if settings.DEBUG:
async def get_current_user(
credentials: HTTPAuthorizationCredentials = Depends(security),
db: AsyncSession = Depends(get_db),
async def get_user_by_token(
token: str,
db: AsyncSession,
) -> AppUser:
token = credentials.credentials
payload = verify_token(token)
if payload is None:
raise HTTPException(
@ -47,6 +46,12 @@ if settings.DEBUG:
)
return user
async def get_current_user(
credentials: HTTPAuthorizationCredentials = Depends(security),
db: AsyncSession = Depends(get_db),
) -> AppUser:
return await get_user_by_token(token=credentials.credentials, db=db)
else:
from raisa_fastapi_protected_api import UserInfo, get_user_dependency

View File

@ -1,6 +1,7 @@
from fastapi import APIRouter
from src.api.v1 import auth, users, admin, audit, forms, form_phases, projects
from src.api.v1 import websocket
api_router = APIRouter()
@ -11,3 +12,4 @@ api_router.include_router(audit.router)
api_router.include_router(forms.router)
api_router.include_router(projects.router)
api_router.include_router(form_phases.router)
api_router.include_router(websocket.router)

607
api/src/api/v1/websocket.py Normal file
View File

@ -0,0 +1,607 @@
from contextlib import asynccontextmanager
from dataclasses import asdict
import dataclasses
import enum
from json import dumps, loads
from typing import Any, Optional
#
from asyncpg import UniqueViolationError
from sqlalchemy.exc import IntegrityError
from fastapi import APIRouter, WebSocket, WebSocketDisconnect, status
from sqlalchemy.ext.asyncio import AsyncSession
from src.api.v1.deps import get_user_by_token
from src.db.models.app_user import AppUser
from src.db.models.form_type import FormTypeEnum
from src.services.budget_form_service import BudgetFormService
from src.services.project_service import ProjectService
from src.services.sheet_service import SheetService
from src.core.errors import BasicAppException, ValidationsError
from src.db.session import SessionLocal
from src.services.user_service import UserService
@asynccontextmanager
async def get_db_session():
"""Контекстный менеджер для получения сессии базы данных."""
db = SessionLocal()
try:
yield db
finally:
await db.close()
router = APIRouter(prefix="/ws", tags=["websocket"])
@dataclasses.dataclass
class ConnectionInfo:
ws: WebSocket
user_id: int | None
@dataclasses.dataclass
class FormConnectionInfo(ConnectionInfo):
form_id: int | None = None
sheet: str | None = None
direction: str | None = None
@dataclasses.dataclass
class ProjectConnectionInfo(ConnectionInfo):
project_id: int | None = None
year: int | None = None
report_type: str | None = None
class ConnectionKeyEnum(str, enum.Enum):
FORM = "FORM"
PROJECT = "PROJECT"
class ConnectionManager:
con_info_mapping = {
ConnectionKeyEnum.FORM: FormConnectionInfo,
ConnectionKeyEnum.PROJECT: ProjectConnectionInfo,
}
def __init__(self):
self.connections: dict[int, dict[ConnectionKeyEnum, list[ConnectionInfo]]] = {
ConnectionKeyEnum.FORM: {},
ConnectionKeyEnum.PROJECT: {},
}
async def connect(
self,
websocket: WebSocket,
con_key: ConnectionKeyEnum,
**kwargs,
# form_id: int,
# sheet: str,
# direction: str | None = None,
):
await websocket.accept()
key = frozenset(kwargs.items())
cls = self.con_info_mapping[con_key]
if key not in self.connections[con_key]:
self.connections[con_key][key] = [
cls(
ws=websocket,
user_id=None,
**kwargs,
)
]
else:
self.connections[con_key][key].append(
cls(
ws=websocket,
user_id=None,
**kwargs,
)
)
def set_user(
self,
websocket: WebSocket,
user_id: int,
con_key: ConnectionKeyEnum,
**kwargs,
):
key = frozenset(kwargs.items())
if key not in self.connections[con_key]:
return
for con_info in self.connections[con_key][key]:
if con_info.ws == websocket:
con_info.user_id = user_id
async def disconnect(
self,
websocket: WebSocket,
con_key: ConnectionKeyEnum,
code: int = status.WS_1008_POLICY_VIOLATION,
reason: str = "Ошибка",
**kwargs,
):
key = frozenset(kwargs.items())
if key not in self.connections[con_key]:
return
for con_info in self.connections[con_key][key]:
if con_info.ws == websocket:
self.connections[con_key][key].remove(con_info)
if not len(self.connections[con_key][key]):
del self.connections[con_key][key]
await websocket.close(
code=code, reason=reason,
)
async def broadcast_to_other(
self,
message: str,
user_id: int,
con_key: ConnectionKeyEnum,
**kwargs,
):
key = frozenset(kwargs.items())
if key not in self.connections:
return
for con_info in self.connections[con_key][key]:
if con_info.user_id != user_id:
try:
await con_info.ws.send_text(message)
except:
pass # Игнорируем недоступные соединения
async def broadcast_to_all(
self,
message: str,
con_key: ConnectionKeyEnum,
**kwargs,
):
key = frozenset(kwargs.items())
if key not in self.connections[con_key]:
return
for con_info in self.connections[con_key][key]:
try:
await con_info.ws.send_text(message)
except:
pass # Игнорируем недоступные соединения
async def send_back(self, message: str, websocket: WebSocket):
await websocket.send_text(message)
async def get_data(self, websocket: WebSocket) -> dict:
return loads(await websocket.receive_text())
class FormEventProcess:
def __init__(self, db: AsyncSession):
self.sheet_service: SheetService = SheetService(db)
self.bf_service: BudgetFormService = BudgetFormService(db)
self.user_service: UserService = UserService(db)
async def process(
self,
event_data: dict,
form_id: int,
user_id: int,
sheet: str,
direction: str | None = None,
) -> int | bool | dict:
curr_user: AppUser = await self.user_service.get(user_id)
match event_data["event"]:
case "cell_updated":
return await self.__update_cell(
event_data=event_data["data"],
form_id=form_id,
user=curr_user,
sheet=sheet,
direction=direction,
)
case "row_added":
return await self.__add_row(
event_data=event_data["data"],
form_id=form_id,
user=curr_user,
sheet=sheet,
direction=direction,
)
case "row_deleted":
return await self.__del_row(
event_data=event_data["data"],
form_id=form_id,
user=curr_user,
sheet=sheet,
direction=direction,
)
case _:
pass
async def __update_cell(
self,
event_data: dict,
form_id: int,
sheet: str,
user: AppUser,
direction: str | None = None,
) -> list[tuple]:
"""
event_data: {
"line_id": int,
"column": str,
"value": any,
}
"""
form = await self.bf_service.get(budget_form_id=form_id, user=user)
if not form:
return None
return await self.sheet_service.update_cell(
form_id=form_id,
sheet=sheet,
direction=direction,
sections=None,
line_id=event_data["line_id"],
column=event_data["column"],
value=event_data["value"],
user=user,
)
async def __add_row(
self,
event_data: dict,
form_id: int,
sheet: str,
user: AppUser,
direction: str | None,
) -> list[tuple]:
"""
data = {
expense_item_id: Optional[int] = None
item_id: Optional[str] = None
section_code: Optional[str] = None
name: Optional[str] = None
internal_order: Optional[str] = None
vsp_id: Optional[int] = None
project_id: Optional[int] = None
justification: Optional[str] = None
contract_number: Optional[str] = None
contract_end_date: Optional[datetime] = None
}
"""
form = await self.bf_service.get(budget_form_id=form_id, user=user)
if not form:
return None
return await self.sheet_service.add_line(
form_id=form_id,
sheet=sheet,
expense_item_id=event_data.get("expense_item_id"),
item_id=event_data.get("item_id"),
section_code=event_data.get("section_code"),
direction=direction,
name=event_data.get("name"),
internal_order=event_data.get("internal_order"),
vsp_id=event_data.get("vsp_id"),
project_id=event_data.get("project_id"),
justification=event_data.get("justification"),
contract_number=event_data.get("contract_number"),
contract_end_date=event_data.get("contract_end_date"),
user=user,
)
async def __del_row(
self,
event_data: dict,
form_id: int,
sheet: str,
user: AppUser,
direction: str | None,
) -> list[tuple]:
"""
data = {
"row_id": int,
}
"""
form = await self.bf_service.get(budget_form_id=form_id, user=user)
if not form:
return None
result = await self.sheet_service.delete_line(
form_id=form_id,
sheet=sheet,
row_id=event_data["row_id"],
direction=direction,
user=user,
)
return result
class ProjectEventProcess:
def __init__(self, db: AsyncSession):
self.project_service: ProjectService = ProjectService(db)
self.bf_service: BudgetFormService = BudgetFormService(db)
self.user_service: UserService = UserService(db)
async def process(
self,
event_data: dict,
project_id: int,
user_id: int,
year: int,
report_type: str,
) -> int | bool | dict:
curr_user: AppUser = await self.user_service.get(user_id)
match event_data["event"]:
case "cell_updated":
return await self.__update_cell(
event_data=event_data["data"],
project_id=project_id,
year=year,
report_type=report_type,
user=curr_user,
)
case "row_added":
return await self.__add_row(
event_data=event_data["data"],
project_id=project_id,
year=year,
report_type=report_type,
user=curr_user,
)
case "row_deleted":
return await self.__del_row(
event_data=event_data["data"],
project_id=project_id,
year=year,
report_type=report_type,
user=curr_user,
)
case _:
pass
async def __update_cell(
self,
event_data: dict,
project_id: int,
year: int,
report_type: str,
user: AppUser,
) -> list[tuple]:
"""
event_data: {
"line_id": int,
"column": str,
"value": any,
}
"""
project = await self.project_service.get(project_id=project_id, user=user)
if not project:
return None
return await self.project_service.upd_form3_cell(
project_id=project_id,
year=year,
report_type=report_type,
line_id=event_data["line_id"],
column=event_data["column"],
value=event_data["value"],
user=user,
)
async def __add_row(
self,
event_data: dict,
project_id: int,
year: int,
report_type: str,
user: AppUser,
) -> list[tuple]:
"""
data = {
expense_item_id: Optional[int] = None
}
"""
project = await self.project_service.get(project_id=project_id, user=user)
if not project:
return None
return await self.project_service.add_form3_line(
project_id=project_id,
year=year,
report_type=report_type,
expense_item_id=event_data["expense_item_id"],
user=user,
)
async def __del_row(
self,
event_data: dict,
project_id: int,
year: int,
report_type: str,
user: AppUser,
) -> list[tuple]:
"""
data = {
"line_id": int,
}
"""
project = await self.project_service.get(project_id=project_id, user=user)
if not project:
return None
return await self.project_service.del_form3_line(
project_id=project_id,
year=year,
report_type=report_type,
line_id=event_data["line_id"],
user=user,
)
manager = ConnectionManager()
def _convert_error(o):
try:
return asdict(o)
except TypeError:
return o
async def login(websocket: WebSocket, **kwargs) -> int:
user_data = await manager.get_data(websocket=websocket)
if user_data.get("event") != "user_login":
return None
async with get_db_session() as db:
user = await get_user_by_token(token=user_data["data"].get("token"), db=db)
if not user:
return None
manager.set_user(
websocket=websocket,
user_id=user.id,
**kwargs,
)
return user.id
async def process_websocket(
websocket: WebSocket,
processor_cls,
**kwargs,
):
await manager.connect(
websocket=websocket,
**kwargs,
)
try:
user_id = await login(
websocket=websocket,
**kwargs,
)
if user_id is None:
await manager.disconnect(
websocket=websocket,
reason="Ошибка авторизации",
**kwargs,
)
return
while True:
data = loads(await websocket.receive_text())
kwargs_process = kwargs.copy()
if "con_key" in kwargs_process:
kwargs_process.pop("con_key")
try:
async with get_db_session() as db:
processor = processor_cls(db)
data["result"] = await processor.process(
event_data=data,
user_id=user_id,
**kwargs_process
)
await db.commit()
if data.get("event") in [
"cell_updated",
"row_added",
"row_deleted",
]:
await manager.broadcast_to_all(
message=dumps(data, default=_convert_error),
**kwargs
)
else:
await manager.broadcast_to_other(
message=dumps(data, default=_convert_error),
user_id=user_id,
**kwargs,
)
except BasicAppException as e:
data["error"] = e.description or "Неизвестная ошибка"
await manager.send_back(
message=dumps(data, default=_convert_error),
websocket=websocket,
)
except IntegrityError as e:
data["error"] = str(e)
await manager.send_back(
message=dumps(data, default=_convert_error),
websocket=websocket,
)
except Exception as e:
tp = type(e)
handlers = websocket.app.exception_handlers
if tp in handlers:
data["error"] = loads((await handlers[tp](request=None, exc=e)).body)
await manager.send_back(
message=dumps(data, default=_convert_error),
websocket=websocket,
)
else:
raise e
except WebSocketDisconnect:
await manager.disconnect(
websocket=websocket,
**kwargs
)
except Exception as e:
# Логируем ошибку, но не бросаем HTTPException — это WebSocket
print(f"Error: {e}")
await manager.disconnect(
websocket=websocket,
**kwargs,
)
@router.websocket("/form/{form_id}/sheet/{sheet}")
async def websocket_form(
websocket: WebSocket,
form_id: int,
sheet: str,
direction: Optional[str] = None,
):
await process_websocket(
websocket=websocket,
form_id=form_id,
sheet=sheet,
direction=direction,
processor_cls=FormEventProcess,
con_key=ConnectionKeyEnum.FORM,
)
@router.websocket("/projects/{project_id}/report/{year}/{report_type}")
async def websocket_project(
websocket: WebSocket,
project_id: int,
year: int,
report_type: str,
):
await process_websocket(
websocket=websocket,
project_id=project_id,
year=year,
report_type=report_type,
processor_cls=ProjectEventProcess,
con_key=ConnectionKeyEnum.PROJECT,
)

View File

@ -235,8 +235,7 @@ class ProjectRepository:
},
)
).all()
await self.db.commit()
return rows
return [tuple(r) for r in rows]
async def upd_form3_cells(self, report_id: int, changes: list[dict]) -> list[tuple]:
query = text(
@ -254,8 +253,7 @@ class ProjectRepository:
},
)
).all()
await self.db.commit()
return rows
return [tuple(r) for r in rows]
async def add_form3_line(self, report_id: int, expense_item_id: int) -> list[tuple]:
query = text(
@ -273,8 +271,7 @@ class ProjectRepository:
},
)
).all()
await self.db.commit()
return rows
return [tuple(r) for r in rows]
async def del_form3_line(self, line_id: int) -> list[tuple]:
query = text(
@ -284,8 +281,7 @@ class ProjectRepository:
"""
)
rows = (await self.db.execute(query, {"line_id": line_id})).all()
await self.db.commit()
return rows
return [tuple(r) for r in rows]
async def upd_project(self, project_id: int, column: str, value):
query = text(
@ -303,7 +299,6 @@ class ProjectRepository:
},
)
).scalar_one()
await self.db.commit()
return result
async def add_project(
@ -356,5 +351,4 @@ class ProjectRepository:
},
)
).first()
await self.db.commit()
return int(row[0]), int(row[1]), int(row[2])

View File

@ -93,7 +93,8 @@ class SheetRepository:
else:
func_query = "v3.upd_form_cell(:form_id, :sheet, :line_id, :column, :value, :direction, :sections)"
return (
return [
tuple(r) for r in (
await self.db.execute(
text(
f"""
@ -113,6 +114,7 @@ class SheetRepository:
}
)
).all()
]
async def update_cells(
self,
@ -138,7 +140,8 @@ class SheetRepository:
else:
func_query = "v3.upd_form_cells(:form_id, :sheet, :changes, :direction, :sections)"
return (
return [
tuple(r) for r in (
await self.db.execute(
text(
f"""
@ -156,6 +159,7 @@ class SheetRepository:
}
)
).all()
]
async def add_line(
self,
@ -174,7 +178,8 @@ class SheetRepository:
contract_end_date: str | None = None,
user_id: int | None = None,
) -> list[tuple]:
return (
return [
tuple(r) for r in (
await self.db.execute(
text(
"""
@ -213,6 +218,7 @@ class SheetRepository:
}
)
).all()
]
async def delete_line(
self,
@ -220,7 +226,9 @@ class SheetRepository:
row_id: int,
direction: str | None = None,
) -> list[tuple]:
return (
return [
tuple(r) for r in (
await self.db.execute(
text(
"SELECT row_type, depth, sort_order, data "
@ -233,3 +241,5 @@ class SheetRepository:
}
)
).all()
]

View File

@ -22,8 +22,8 @@ class UserService:
raise AccessDeniedException()
return await self.user_repo.get_list(skip=skip, limit=limit)
async def get(self, user_id: int, current_user: AppUser):
if current_user.role_id != UserRoleEnum.ADMIN:
async def get(self, user_id: int, current_user: AppUser | None = None):
if current_user is not None and current_user.role_id != UserRoleEnum.ADMIN:
raise AccessDeniedException()
return await self.user_repo.get(user_id)