From 33055d6695c3b56c0866c8b4b3991e27ead43847 Mon Sep 17 00:00:00 2001 From: tsygankoviva Date: Wed, 20 May 2026 18:35:26 +0300 Subject: [PATCH] =?UTF-8?q?ws:=20=D0=B2=D0=B5=D0=B1=D1=81=D0=BE=D0=BA?= =?UTF-8?q?=D0=B5=D1=82=D1=8B=20=D0=BD=D0=B0=20=D0=BD=D0=BE=D0=B2=D0=BE?= =?UTF-8?q?=D0=B9=20=D0=BC=D0=BE=D0=B4=D0=B5=D0=BB=D0=B8=20=D0=B4=D0=B0?= =?UTF-8?q?=D0=BD=D0=BD=D1=8B=D1=85?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- api/src/api/v1/deps.py | 13 +- api/src/api/v1/router.py | 2 + api/src/api/v1/websocket.py | 607 +++++++++++++++++++++++ api/src/repository/project_repository.py | 14 +- api/src/repository/sheet_repository.py | 190 +++---- api/src/services/user_service.py | 4 +- 6 files changed, 724 insertions(+), 106 deletions(-) create mode 100644 api/src/api/v1/websocket.py diff --git a/api/src/api/v1/deps.py b/api/src/api/v1/deps.py index 79e62d8..f837ec5 100644 --- a/api/src/api/v1/deps.py +++ b/api/src/api/v1/deps.py @@ -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 diff --git a/api/src/api/v1/router.py b/api/src/api/v1/router.py index 3b9b4b7..915dead 100644 --- a/api/src/api/v1/router.py +++ b/api/src/api/v1/router.py @@ -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) diff --git a/api/src/api/v1/websocket.py b/api/src/api/v1/websocket.py new file mode 100644 index 0000000..4d3609e --- /dev/null +++ b/api/src/api/v1/websocket.py @@ -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, + ) diff --git a/api/src/repository/project_repository.py b/api/src/repository/project_repository.py index 565045b..1faae5e 100644 --- a/api/src/repository/project_repository.py +++ b/api/src/repository/project_repository.py @@ -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]) diff --git a/api/src/repository/sheet_repository.py b/api/src/repository/sheet_repository.py index 12b5e2b..cad3323 100644 --- a/api/src/repository/sheet_repository.py +++ b/api/src/repository/sheet_repository.py @@ -93,26 +93,28 @@ class SheetRepository: else: func_query = "v3.upd_form_cell(:form_id, :sheet, :line_id, :column, :value, :direction, :sections)" - return ( - await self.db.execute( - text( - f""" - SELECT row_type, depth, sort_order, data - FROM {func_query} - """ - ), - { - "form_id": form_id, - "sheet": sheet, - "direction": direction, - "sections": sections, - "line_id": line_id, - "column": column, - "value": json.dumps(value), - "user_id": user_id, - } - ) - ).all() + return [ + tuple(r) for r in ( + await self.db.execute( + text( + f""" + SELECT row_type, depth, sort_order, data + FROM {func_query} + """ + ), + { + "form_id": form_id, + "sheet": sheet, + "direction": direction, + "sections": sections, + "line_id": line_id, + "column": column, + "value": json.dumps(value), + "user_id": user_id, + } + ) + ).all() + ] async def update_cells( self, @@ -138,24 +140,26 @@ class SheetRepository: else: func_query = "v3.upd_form_cells(:form_id, :sheet, :changes, :direction, :sections)" - return ( - await self.db.execute( - text( - f""" - SELECT row_type, depth, sort_order, data - FROM {func_query} - """ - ), - { - "form_id": form_id, - "sheet": sheet, - "direction": direction, - "sections": sections, - "changes": changes, - "user_id": user_id, - } - ) - ).all() + return [ + tuple(r) for r in ( + await self.db.execute( + text( + f""" + SELECT row_type, depth, sort_order, data + FROM {func_query} + """ + ), + { + "form_id": form_id, + "sheet": sheet, + "direction": direction, + "sections": sections, + "changes": changes, + "user_id": user_id, + } + ) + ).all() + ] async def add_line( self, @@ -174,45 +178,47 @@ class SheetRepository: contract_end_date: str | None = None, user_id: int | None = None, ) -> list[tuple]: - return ( - await self.db.execute( - text( - """ - SELECT row_type, depth, sort_order, data - FROM v3.add_budget_line( - :p_form_id, - :p_expense_item_id, - :p_sheet, - :p_item_id, - :p_section_code, - :p_direction, - :p_name, - :p_internal_order, - :p_vsp_id, - :p_project_id, - :p_justification, - :p_contract_number, - :p_contract_end_date - ) - """ - ), - { - "p_form_id": form_id, - "p_expense_item_id": expense_item_id, - "p_sheet": sheet, - "p_item_id": item_id, - "p_section_code": section_code, - "p_direction": direction, - "p_name": name, - "p_internal_order": internal_order, - "p_vsp_id": vsp_id, - "p_project_id": project_id, - "p_justification": justification, - "p_contract_number": contract_number, - "p_contract_end_date": contract_end_date, - } - ) - ).all() + return [ + tuple(r) for r in ( + await self.db.execute( + text( + """ + SELECT row_type, depth, sort_order, data + FROM v3.add_budget_line( + :p_form_id, + :p_expense_item_id, + :p_sheet, + :p_item_id, + :p_section_code, + :p_direction, + :p_name, + :p_internal_order, + :p_vsp_id, + :p_project_id, + :p_justification, + :p_contract_number, + :p_contract_end_date + ) + """ + ), + { + "p_form_id": form_id, + "p_expense_item_id": expense_item_id, + "p_sheet": sheet, + "p_item_id": item_id, + "p_section_code": section_code, + "p_direction": direction, + "p_name": name, + "p_internal_order": internal_order, + "p_vsp_id": vsp_id, + "p_project_id": project_id, + "p_justification": justification, + "p_contract_number": contract_number, + "p_contract_end_date": contract_end_date, + } + ) + ).all() + ] async def delete_line( self, @@ -220,16 +226,20 @@ class SheetRepository: row_id: int, direction: str | None = None, ) -> list[tuple]: - return ( - await self.db.execute( - text( - "SELECT row_type, depth, sort_order, data " - "FROM v3.del_budget_line(:p_row_id, :p_direction, :p_sheet)" - ), - { - "p_row_id": row_id, - "p_direction": direction, - "p_sheet": sheet, - } - ) - ).all() + + return [ + tuple(r) for r in ( + await self.db.execute( + text( + "SELECT row_type, depth, sort_order, data " + "FROM v3.del_budget_line(:p_row_id, :p_direction, :p_sheet)" + ), + { + "p_row_id": row_id, + "p_direction": direction, + "p_sheet": sheet, + } + ) + ).all() + ] + diff --git a/api/src/services/user_service.py b/api/src/services/user_service.py index de3a9c6..9432fcf 100644 --- a/api/src/services/user_service.py +++ b/api/src/services/user_service.py @@ -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)