Ушли от уровня орм к уровню бд, позже добавлю функции в 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 sqlalchemy import delete, select, update
from sqlalchemy import select, text
from sqlalchemy.ext.asyncio import AsyncSession
from src.db.models.form_phase import FormPhase
@ -45,40 +45,79 @@ class FormPhaseRepository:
opens_at: datetime,
closes_at: datetime,
) -> FormPhase:
form_phase = FormPhase(
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,
result = await self.db.execute(
text(
"""
SELECT
(v3.add_form_phase(
:budget_form_id,
: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": role,
"column_keys": column_keys,
"opens_at": opens_at,
"closes_at": closes_at,
},
)
self.db.add(form_phase)
await self.db.flush()
return form_phase
if result.scalar_one_or_none() is None:
return None
return await self.get(budget_form_id, sheet, phase_code)
async def update(
self, budget_form_id: int, sheet: str, phase_code: str, data: dict
) -> FormPhase | None:
query = (
update(FormPhase)
.where(FormPhase.budget_form_id == budget_form_id)
.where(FormPhase.sheet == sheet)
.where(FormPhase.phase_code == phase_code)
.values(**data)
.returning(FormPhase)
result = await self.db.execute(
text(
"""
SELECT
(v3.upd_form_phase(
:budget_form_id,
: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(
self, budget_form_id: int, sheet: str, phase_code: str
) -> bool:
query = (
delete(FormPhase)
.where(FormPhase.budget_form_id == budget_form_id)
.where(FormPhase.sheet == sheet)
.where(FormPhase.phase_code == phase_code)
result = await self.db.execute(
text(
"SELECT v3.del_form_phase(:budget_form_id, :sheet, :phase_code)"
),
{
"budget_form_id": budget_form_id,
"sheet": sheet,
"phase_code": phase_code,
},
)
result = await self.db.execute(query)
return result.rowcount > 0
deleted = result.scalar_one_or_none()
return bool(deleted)

View File

@ -1,6 +1,6 @@
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.orm import selectinload
@ -75,36 +75,6 @@ class UserRepository:
)
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:
query = update(AppUser).where(AppUser.id == user_id).values(is_active=False)
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()
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
)
return (await self.db.execute(query)).scalars().all()
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()
async def set_many_ssp(self, user_id: int, ssp_ids: list[int]) -> bool:
try:
for ssp_id in ssp_ids:
link = UserOrg(user_id=user_id, ssp_id=ssp_id)
self.db.add(link)
await self.db.execute(
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()
return True
except Exception:
@ -149,10 +123,12 @@ class UserRepository:
async def unset_many_ssp(self, user_id: int, ssp_ids: list[int]) -> bool:
try:
for ssp_id in ssp_ids:
query = delete(UserOrg).where(
(UserOrg.user_id == user_id) & (UserOrg.ssp_id == ssp_id)
await self.db.execute(
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()
return True
except Exception:
@ -161,10 +137,80 @@ class UserRepository:
async def clear_many_ssp(self, user_id: int) -> bool:
try:
query = delete(UserOrg).where(UserOrg.user_id == user_id)
await self.db.execute(query)
org_unit_ids = await self.get_many_ssp_ids(user_id)
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()
return True
except Exception:
await self.db.rollback()
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
def _result_with_scalar_one(value):
result = MagicMock()
result.scalar_one.return_value = value
return result
@pytest.fixture
def mock_db():
return AsyncMock(spec=AsyncSession)
@ -49,9 +55,12 @@ async def test_get_returns_none_for_missing_id(repository, mock_db):
@pytest.mark.asyncio
async def test_get_all_applies_filters(repository, mock_db):
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,
offset=5,
user_id=2,
@ -64,5 +73,6 @@ async def test_get_all_applies_filters(repository, mock_db):
)
assert result == [row]
mock_db.execute.assert_awaited_once()
assert count == 1
assert mock_db.execute.await_count == 2