"""Группы: создание, дополнения, доступные фракции, проверки доступа.""" 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_nine_rounds_rule(session: Session, group: Group, enabled: bool) -> Group: """Хоумрул «9 раундов при 5–6 игроках». Действует на партии, начатые после смены: уже начатые хранят свой снимок (Match.nine_rounds_rule).""" group.nine_rounds_rule = enabled 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]