Files
ForbiddenStarsApp/backend/app/services/group_service.py
T

139 lines
4.6 KiB
Python

"""Группы: создание, дополнения, доступные фракции, проверки доступа."""
from __future__ import annotations
from sqlalchemy import or_
from sqlmodel import Session, select
from app.core.errors import NotFoundError, NotGroupMemberError, ForbiddenError, ValidationError
from app.models import (
Expansion,
Faction,
Group,
GroupExpansion,
GroupMember,
User,
)
def get_group(session: Session, group_id: int) -> Group:
group = session.get(Group, group_id)
if group is None:
raise NotFoundError("Группа не найдена.")
return group
def get_membership(session: Session, group_id: int, user_id: int) -> GroupMember | None:
return session.exec(
select(GroupMember).where(
GroupMember.group_id == group_id, GroupMember.user_id == user_id
)
).first()
def assert_member(session: Session, group_id: int, user_id: int) -> GroupMember:
member = get_membership(session, group_id, user_id)
if member is None:
raise NotGroupMemberError()
return member
def assert_owner(session: Session, group_id: int, user_id: int) -> GroupMember:
member = assert_member(session, group_id, user_id)
if member.role != "owner":
raise ForbiddenError("Действие доступно только владельцу группы.")
return member
def list_user_groups(session: Session, user_id: int) -> list[tuple[Group, str]]:
rows = session.exec(
select(Group, GroupMember.role)
.join(GroupMember, GroupMember.group_id == Group.id)
.where(GroupMember.user_id == user_id)
.order_by(Group.name)
).all()
return [(g, role) for g, role in rows]
def _valid_non_base_expansion_ids(session: Session, expansion_ids: list[int]) -> list[int]:
if not expansion_ids:
return []
found = session.exec(
select(Expansion.id).where(
Expansion.id.in_(expansion_ids), Expansion.is_base == False # noqa: E712
)
).all()
return list(found)
def create_group(session: Session, owner: User, name: str, expansion_ids: list[int]) -> Group:
name = (name or "").strip()
if not (2 <= len(name) <= 64):
raise ValidationError("Название группы: 2–64 символа.")
group = Group(name=name, owner_id=owner.id) # type: ignore[arg-type]
session.add(group)
session.flush()
session.add(GroupMember(group_id=group.id, user_id=owner.id, role="owner")) # type: ignore[arg-type]
for exp_id in _valid_non_base_expansion_ids(session, expansion_ids):
session.add(GroupExpansion(group_id=group.id, expansion_id=exp_id)) # type: ignore[arg-type]
owner.active_group_id = group.id
session.add(owner)
session.commit()
session.refresh(group)
return group
def rename_group(session: Session, group: Group, name: str) -> Group:
name = (name or "").strip()
if not (2 <= len(name) <= 64):
raise ValidationError("Название группы: 2–64 символа.")
group.name = name
session.add(group)
session.commit()
session.refresh(group)
return group
def set_expansions(session: Session, group: Group, expansion_ids: list[int]) -> Group:
valid = set(_valid_non_base_expansion_ids(session, expansion_ids))
current = session.exec(
select(GroupExpansion).where(GroupExpansion.group_id == group.id)
).all()
current_ids = {ge.expansion_id for ge in current}
for ge in current:
if ge.expansion_id not in valid:
session.delete(ge)
for exp_id in valid - current_ids:
session.add(GroupExpansion(group_id=group.id, expansion_id=exp_id)) # type: ignore[arg-type]
session.commit()
session.refresh(group)
return group
def group_expansion_ids(session: Session, group_id: int) -> list[int]:
return list(
session.exec(
select(GroupExpansion.expansion_id).where(GroupExpansion.group_id == group_id)
).all()
)
def available_factions(session: Session, group_id: int) -> list[Faction]:
owned = select(GroupExpansion.expansion_id).where(GroupExpansion.group_id == group_id)
stmt = (
select(Faction)
.join(Expansion, Expansion.id == Faction.expansion_id)
.where(or_(Expansion.is_base == True, Expansion.id.in_(owned))) # noqa: E712
.order_by(Expansion.sort_order, Faction.sort_order)
)
return list(session.exec(stmt).all())
def available_faction_ids(session: Session, group_id: int) -> set[int]:
return {f.id for f in available_factions(session, group_id)} # type: ignore[misc]