Ушли от уровня орм к уровню бд, позже добавлю функции в rnd репо

This commit is contained in:
Raykov-MS 2026-05-21 10:42:40 +03:00
parent 8e43920deb
commit 411773b6e9
3 changed files with 165 additions and 70 deletions

View File

@ -1,6 +1,6 @@
from datetime import datetime from datetime import datetime
from sqlalchemy import delete, select, update from sqlalchemy import select, text
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from src.db.models.form_phase import FormPhase from src.db.models.form_phase import FormPhase
@ -45,40 +45,79 @@ class FormPhaseRepository:
opens_at: datetime, opens_at: datetime,
closes_at: datetime, closes_at: datetime,
) -> FormPhase: ) -> FormPhase:
form_phase = FormPhase( result = await self.db.execute(
budget_form_id=budget_form_id, text(
sheet=sheet, """
phase_code=phase_code, SELECT
role=role, (v3.add_form_phase(
column_keys=column_keys, :budget_form_id,
opens_at=opens_at, :sheet,
closes_at=closes_at, :phase_code,
:role,
:column_keys,
:opens_at,
:closes_at
)).budget_form_id
"""
),
{
"budget_form_id": budget_form_id,
"sheet": sheet,
"phase_code": phase_code,
"role": role,
"column_keys": column_keys,
"opens_at": opens_at,
"closes_at": closes_at,
},
) )
self.db.add(form_phase) if result.scalar_one_or_none() is None:
await self.db.flush() return None
return form_phase return await self.get(budget_form_id, sheet, phase_code)
async def update( async def update(
self, budget_form_id: int, sheet: str, phase_code: str, data: dict self, budget_form_id: int, sheet: str, phase_code: str, data: dict
) -> FormPhase | None: ) -> FormPhase | None:
query = ( result = await self.db.execute(
update(FormPhase) text(
.where(FormPhase.budget_form_id == budget_form_id) """
.where(FormPhase.sheet == sheet) SELECT
.where(FormPhase.phase_code == phase_code) (v3.upd_form_phase(
.values(**data) :budget_form_id,
.returning(FormPhase) :sheet,
:phase_code,
:role,
:column_keys,
:opens_at,
:closes_at
)).budget_form_id
"""
),
{
"budget_form_id": budget_form_id,
"sheet": sheet,
"phase_code": phase_code,
"role": data.get("role"),
"column_keys": data.get("column_keys"),
"opens_at": data.get("opens_at"),
"closes_at": data.get("closes_at"),
},
) )
return (await self.db.execute(query)).scalar_one_or_none() if result.scalar_one_or_none() is None:
return None
return await self.get(budget_form_id, sheet, phase_code)
async def delete( async def delete(
self, budget_form_id: int, sheet: str, phase_code: str self, budget_form_id: int, sheet: str, phase_code: str
) -> bool: ) -> bool:
query = ( result = await self.db.execute(
delete(FormPhase) text(
.where(FormPhase.budget_form_id == budget_form_id) "SELECT v3.del_form_phase(:budget_form_id, :sheet, :phase_code)"
.where(FormPhase.sheet == sheet) ),
.where(FormPhase.phase_code == phase_code) {
"budget_form_id": budget_form_id,
"sheet": sheet,
"phase_code": phase_code,
},
) )
result = await self.db.execute(query) deleted = result.scalar_one_or_none()
return result.rowcount > 0 return bool(deleted)

View File

@ -1,6 +1,6 @@
from typing import Optional from typing import Optional
from sqlalchemy import delete, select, text, update from sqlalchemy import select, text, update
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import selectinload from sqlalchemy.orm import selectinload
@ -75,36 +75,6 @@ class UserRepository:
) )
return (await self.db.execute(query)).scalars().all() return (await self.db.execute(query)).scalars().all()
async def create(self, user_data: dict) -> AppUser:
hashed_password = (
get_password_hash(user_data["password"]) if user_data.get("password") else None
)
db_user = AppUser(
email=user_data["email"],
username=user_data["username"],
hashed_password=hashed_password,
full_name=user_data.get("full_name"),
role_id=user_data.get("role_id"),
)
self.db.add(db_user)
await self.db.commit()
await self.db.refresh(db_user)
return db_user
async def update(self, user_id: int, user_data: dict) -> Optional[AppUser]:
user = await self.get(user_id)
if not user:
return None
for field, value in user_data.items():
if field == "password" and value:
value = get_password_hash(value)
setattr(user, field, value)
await self.db.commit()
await self.db.refresh(user)
return user
async def logical_delete(self, user_id: int) -> bool: async def logical_delete(self, user_id: int) -> bool:
query = update(AppUser).where(AppUser.id == user_id).values(is_active=False) query = update(AppUser).where(AppUser.id == user_id).values(is_active=False)
result = await self.db.execute(query) result = await self.db.execute(query)
@ -126,20 +96,24 @@ class UserRepository:
return (await self.db.execute(select(Role).order_by(Role.id))).scalars().all() return (await self.db.execute(select(Role).order_by(Role.id))).scalars().all()
async def get_many_ssp(self, user_id: int) -> list[OrgUnit]: async def get_many_ssp(self, user_id: int) -> list[OrgUnit]:
query = select(OrgUnit).join(UserOrg, UserOrg.ssp_id == OrgUnit.id).where( query = select(OrgUnit).join(UserOrg, UserOrg.org_unit_id == OrgUnit.id).where(
UserOrg.user_id == user_id UserOrg.user_id == user_id
) )
return (await self.db.execute(query)).scalars().all() return (await self.db.execute(query)).scalars().all()
async def get_many_ssp_ids(self, user_id: int) -> list[int]: async def get_many_ssp_ids(self, user_id: int) -> list[int]:
query = select(UserOrg.ssp_id).where(UserOrg.user_id == user_id) query = select(UserOrg.org_unit_id).where(UserOrg.user_id == user_id)
return (await self.db.execute(query)).scalars().all() return (await self.db.execute(query)).scalars().all()
async def set_many_ssp(self, user_id: int, ssp_ids: list[int]) -> bool: async def set_many_ssp(self, user_id: int, ssp_ids: list[int]) -> bool:
try: try:
for ssp_id in ssp_ids: for ssp_id in ssp_ids:
link = UserOrg(user_id=user_id, ssp_id=ssp_id) await self.db.execute(
self.db.add(link) text(
"SELECT v3.grant_user_org_access(:user_id, :org_unit_id)"
),
{"user_id": user_id, "org_unit_id": ssp_id},
)
await self.db.commit() await self.db.commit()
return True return True
except Exception: except Exception:
@ -149,10 +123,12 @@ class UserRepository:
async def unset_many_ssp(self, user_id: int, ssp_ids: list[int]) -> bool: async def unset_many_ssp(self, user_id: int, ssp_ids: list[int]) -> bool:
try: try:
for ssp_id in ssp_ids: for ssp_id in ssp_ids:
query = delete(UserOrg).where( await self.db.execute(
(UserOrg.user_id == user_id) & (UserOrg.ssp_id == ssp_id) text(
"SELECT v3.revoke_user_org_access(:user_id, :org_unit_id)"
),
{"user_id": user_id, "org_unit_id": ssp_id},
) )
await self.db.execute(query)
await self.db.commit() await self.db.commit()
return True return True
except Exception: except Exception:
@ -161,10 +137,80 @@ class UserRepository:
async def clear_many_ssp(self, user_id: int) -> bool: async def clear_many_ssp(self, user_id: int) -> bool:
try: try:
query = delete(UserOrg).where(UserOrg.user_id == user_id) org_unit_ids = await self.get_many_ssp_ids(user_id)
await self.db.execute(query) await self.db.execute(
text(
"SELECT v3.revoke_many_user_org_access(:user_id, :org_unit_ids)"
),
{"user_id": user_id, "org_unit_ids": org_unit_ids},
)
await self.db.commit() await self.db.commit()
return True return True
except Exception: except Exception:
await self.db.rollback() await self.db.rollback()
return False return False
async def create(self, user_data: dict) -> AppUser:
hashed_password = (
get_password_hash(user_data["password"]) if user_data.get("password") else None
)
user_id = (
await self.db.execute(
text(
"""
SELECT (v3.add_user(
:email,
:username,
:hashed_password,
:full_name,
:role_id
)).id
"""
),
{
"email": user_data["email"],
"username": user_data["username"],
"hashed_password": hashed_password,
"full_name": user_data.get("full_name"),
"role_id": user_data.get("role_id"),
},
)
).scalar_one()
await self.db.commit()
return await self.get(user_id)
async def update(self, user_id: int, user_data: dict) -> Optional[AppUser]:
user = await self.get(user_id)
if not user:
return None
hashed_password = None
if "password" in user_data and user_data["password"]:
hashed_password = get_password_hash(user_data["password"])
await self.db.execute(
text(
"""
SELECT (v3.upd_user(
:user_id,
:email,
:username,
:hashed_password,
:full_name,
:role_id,
:is_active
)).id
"""
),
{
"user_id": user_id,
"email": user_data.get("email"),
"username": user_data.get("username"),
"hashed_password": hashed_password,
"full_name": user_data.get("full_name"),
"role_id": user_data.get("role_id"),
"is_active": user_data.get("is_active"),
},
)
await self.db.commit()
return await self.get(user_id)

View File

@ -16,6 +16,12 @@ def _result_with_scalars(items):
return result return result
def _result_with_scalar_one(value):
result = MagicMock()
result.scalar_one.return_value = value
return result
@pytest.fixture @pytest.fixture
def mock_db(): def mock_db():
return AsyncMock(spec=AsyncSession) return AsyncMock(spec=AsyncSession)
@ -49,9 +55,12 @@ async def test_get_returns_none_for_missing_id(repository, mock_db):
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_get_all_applies_filters(repository, mock_db): async def test_get_all_applies_filters(repository, mock_db):
row = MagicMock(id=11, user_id=2, org_unit_id=7, event_type="UPDATE") row = MagicMock(id=11, user_id=2, org_unit_id=7, event_type="UPDATE")
mock_db.execute.return_value = _result_with_scalars([row]) mock_db.execute.side_effect = [
_result_with_scalars([row]),
_result_with_scalar_one(1),
]
result = await repository.get_all( result, count = await repository.get_all(
limit=10, limit=10,
offset=5, offset=5,
user_id=2, user_id=2,
@ -64,5 +73,6 @@ async def test_get_all_applies_filters(repository, mock_db):
) )
assert result == [row] assert result == [row]
mock_db.execute.assert_awaited_once() assert count == 1
assert mock_db.execute.await_count == 2