diff --git a/api/src/api/v1/websocket.py b/api/src/api/v1/websocket.py index 4d3609e..402e6d8 100644 --- a/api/src/api/v1/websocket.py +++ b/api/src/api/v1/websocket.py @@ -11,6 +11,7 @@ from sqlalchemy.exc import IntegrityError from fastapi import APIRouter, WebSocket, WebSocketDisconnect, status from sqlalchemy.ext.asyncio import AsyncSession +from src.services.budget_line_service import BudgetLineService 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 @@ -186,6 +187,7 @@ class FormEventProcess: def __init__(self, db: AsyncSession): self.sheet_service: SheetService = SheetService(db) self.bf_service: BudgetFormService = BudgetFormService(db) + self.bl_service: BudgetLineService = BudgetLineService(db) self.user_service: UserService = UserService(db) @@ -283,7 +285,8 @@ class FormEventProcess: if not form: return None - return await self.sheet_service.add_line( + ids = await self.bl_service.get_ids(budget_form_id=form_id) + result = await self.sheet_service.add_line( form_id=form_id, sheet=sheet, expense_item_id=event_data.get("expense_item_id"), @@ -299,6 +302,16 @@ class FormEventProcess: contract_end_date=event_data.get("contract_end_date"), user=user, ) + final_result = { + "data": result, + "new_line_id": None, + } + for el in result: + if el[3]["line_id"] and el[3]["line_id"] not in ids: + final_result["new_line_id"] = el[3]["line_id"] + return final_result + return final_result + async def __del_row( self, diff --git a/api/src/repository/budget_line_repository.py b/api/src/repository/budget_line_repository.py index 8d4d99f..0ca3c71 100644 --- a/api/src/repository/budget_line_repository.py +++ b/api/src/repository/budget_line_repository.py @@ -14,4 +14,13 @@ class BudgetLineRepository: async def get_list(self, budget_line_ids: list[int]) -> BudgetLine | None: query = select(BudgetLine).where(BudgetLine.id.in_(budget_line_ids)) - return (await self.db.execute(query)).scalars().all() \ No newline at end of file + return (await self.db.execute(query)).scalars().all() + + async def get_ids( + self, + budget_form_id: int, + ) -> list[int]: + + query = select(BudgetLine.id).where(BudgetLine.budget_form_id == budget_form_id) + return (await self.db.execute(query)).scalars().all() + diff --git a/api/src/services/budget_line_service.py b/api/src/services/budget_line_service.py index fca7904..b04cf5b 100644 --- a/api/src/services/budget_line_service.py +++ b/api/src/services/budget_line_service.py @@ -26,3 +26,10 @@ class BudgetLineService: budget_line_ids: list[int], ) -> list[BudgetLine]: return await self.bl_repo.get_list(budget_line_ids=budget_line_ids) + + async def get_ids( + self, + budget_form_id: int, + ) -> list[int]: + return await self.bl_repo.get_ids(budget_form_id=budget_form_id) +