diff --git a/api/src/repository/user_repository.py b/api/src/repository/user_repository.py index 5ee7413..e841ada 100644 --- a/api/src/repository/user_repository.py +++ b/api/src/repository/user_repository.py @@ -167,7 +167,7 @@ class UserRepository: ), {"user_id": user_id, "org_unit_ids": org_unit_ids}, ) - await self.db.commit() + await self.db.flush() async def create(self, user_data: dict) -> AppUser: hashed_password = ( diff --git a/api/src/services/export_service.py b/api/src/services/export_service.py index e6dd1e9..9c8629f 100644 --- a/api/src/services/export_service.py +++ b/api/src/services/export_service.py @@ -169,13 +169,28 @@ class ExportService: form = await self._get_form_for_export(form_id=form_id, current_user=current_user) sheets = self._resolve_requested_sheets(form=form, requested_sheet=sheet) self._validate_sections_for_form(form=form, sections=sections) - plans = self._build_sheet_plans( - form=form, - sheets=sheets, - direction=direction, - sections=sections, - ignore_non_applicable_query_params=ignore_non_applicable_query_params, - ) + + if direction: + plans = self._build_sheet_plans( + form=form, + sheets=sheets, + direction=direction, + sections=sections, + ignore_non_applicable_query_params=ignore_non_applicable_query_params, + ) + else: + sheets_by_directions = self._resolve_sheets_by_directions(form=form, requested_sheet=sheet) + plans = [] + for dr, shts in sheets_by_directions.items(): + plans += self._build_sheet_plans( + form=form, + sheets=shts, + direction=dr, + sections=sections, + ignore_non_applicable_query_params=ignore_non_applicable_query_params, + ) + + workbook = self._create_workbook_with_metadata(form=form) exported_sheets = 0 @@ -294,6 +309,18 @@ class ExportService: if requested_sheet: return [requested_sheet] return list(form.form_type.sheet_list) + + def _resolve_sheets_by_directions(self, form: Any, requested_sheet: str | None) -> dict[str | None, str]: + result = {} + for sheet in form.form_type.sheet_list_with_directions: + if requested_sheet and sheet["sheet_name"] != requested_sheet: + continue + + if sheet["direction"] in result: + result[sheet["direction"]].append(sheet["sheet_name"]) + else: + result[sheet["direction"]] = [sheet["sheet_name"]] + return result @staticmethod def _validate_sections_for_form(form: Any, sections: list[str] | None) -> None: diff --git a/api/src/services/user_service.py b/api/src/services/user_service.py index 0d0bf27..55db948 100644 --- a/api/src/services/user_service.py +++ b/api/src/services/user_service.py @@ -112,6 +112,7 @@ class UserService: if (had_many_ssp_role and not has_many_ssp_role_now) or is_executor_to_dfip_transition: await self.user_repo.clear_many_ssp(user_id) if user_data.load_orgs: + self.db.expire(updated) updated = await self.user_repo.get( user_id=user_id, load_orgs=user_data.load_orgs,