diff --git a/alembic/versions/20260103_1500_004_add_group_roles_mapping.py b/alembic/versions/20260103_1500_004_add_group_roles_mapping.py new file mode 100644 index 0000000..d169e41 --- /dev/null +++ b/alembic/versions/20260103_1500_004_add_group_roles_mapping.py @@ -0,0 +1,141 @@ +"""Add group_roles mapping and update role schema + +Revision ID: 004 +Revises: f0349c95aa5d +Create Date: 2026-01-03 + +Changes: +- Add category column to roles (default 'general') +- Drop authentik_group column from roles (decoupled architecture) +- Create user_groups association table +- Create group_roles association table +- Update role names from domain:action to domain.general:action +""" +from typing import Sequence, Union + +from alembic import op +import sqlalchemy as sa +from sqlalchemy.dialects import postgresql + +# revision identifiers, used by Alembic. +revision: str = "004" +down_revision: Union[str, None] = "f0349c95aa5d" +branch_labels: Union[str, Sequence[str], None] = None +depends_on: Union[str, Sequence[str], None] = None + + +def upgrade() -> None: + # Add category column to roles + op.add_column( + "roles", + sa.Column( + "category", + sa.String(50), + nullable=False, + server_default="general", + comment="Permission category within domain (general for full access, or specific tool)", + ), + ) + + # Update role names from domain:action to domain.general:action + op.execute( + """ + UPDATE roles + SET name = REPLACE(name, ':', '.general:') + WHERE name NOT LIKE '%.%:%' + """ + ) + + # Update the comment on the name column + op.alter_column( + "roles", + "name", + comment="Role name in format domain.category:action (e.g., control-room.general:admin)", + ) + + # Drop the authentik_group unique index first + op.drop_index("ix_roles_authentik_group", table_name="roles") + + # Drop authentik_group column (no longer needed with group_roles mapping) + op.drop_column("roles", "authentik_group") + + # Create user_groups association table + op.create_table( + "user_groups", + sa.Column( + "user_id", + postgresql.UUID(as_uuid=True), + sa.ForeignKey("users.id", ondelete="CASCADE"), + primary_key=True, + ), + sa.Column( + "group_id", + postgresql.UUID(as_uuid=True), + sa.ForeignKey("groups.id", ondelete="CASCADE"), + primary_key=True, + ), + ) + + # Create group_roles association table + op.create_table( + "group_roles", + sa.Column( + "group_id", + postgresql.UUID(as_uuid=True), + sa.ForeignKey("groups.id", ondelete="CASCADE"), + primary_key=True, + ), + sa.Column( + "role_id", + postgresql.UUID(as_uuid=True), + sa.ForeignKey("roles.id", ondelete="CASCADE"), + primary_key=True, + ), + ) + + +def downgrade() -> None: + # Drop association tables + op.drop_table("group_roles") + op.drop_table("user_groups") + + # Add back authentik_group column + op.add_column( + "roles", + sa.Column( + "authentik_group", + sa.String(255), + nullable=True, + comment="Corresponding Authentik group name", + ), + ) + + # Restore authentik_group values from role names + op.execute( + """ + UPDATE roles + SET authentik_group = 'tatlock-' || REPLACE(REPLACE(name, '.general:', '-'), ':', '-') + """ + ) + + # Recreate the unique index + op.create_index("ix_roles_authentik_group", "roles", ["authentik_group"], unique=True) + + # Revert role names from domain.general:action to domain:action + op.execute( + """ + UPDATE roles + SET name = REPLACE(name, '.general:', ':') + WHERE name LIKE '%.general:%' + """ + ) + + # Update the comment on the name column + op.alter_column( + "roles", + "name", + comment="Role name in format domain:action", + ) + + # Drop category column + op.drop_column("roles", "category") diff --git a/src/domains/auth/__init__.py b/src/domains/auth/__init__.py index 031135e..9b1afa6 100644 --- a/src/domains/auth/__init__.py +++ b/src/domains/auth/__init__.py @@ -11,6 +11,12 @@ from src.domains.auth.oidc import ( get_forward_auth_user, get_forward_auth_admin, oidc_config, + # Permission system + require_permission, + require_any_permission, + ACTION_HIERARCHY, + VALID_DOMAINS, + DEFAULT_CATEGORY, ) from src.domains.auth.service import AuthService, get_auth_service from src.domains.auth.controller import auth_controller @@ -22,6 +28,7 @@ from src.domains.auth.models import ( UserPreferences, ApiKey, user_groups, + group_roles, ) __all__ = [ @@ -32,6 +39,12 @@ __all__ = [ "get_forward_auth_user", "get_forward_auth_admin", "oidc_config", + # Permission system + "require_permission", + "require_any_permission", + "ACTION_HIERARCHY", + "VALID_DOMAINS", + "DEFAULT_CATEGORY", # Service "AuthService", "get_auth_service", @@ -45,4 +58,5 @@ __all__ = [ "UserPreferences", "ApiKey", "user_groups", + "group_roles", ] diff --git a/src/domains/auth/controller.py b/src/domains/auth/controller.py index 0ee3273..cac051a 100644 --- a/src/domains/auth/controller.py +++ b/src/domains/auth/controller.py @@ -3,8 +3,9 @@ Authentication Controller Provides authentication endpoints for OIDC token sync and user management. """ +import uuid from typing import Optional -from fastapi import APIRouter, Depends, HTTPException, Query +from fastapi import APIRouter, Depends, HTTPException, Path, Query from fastapi.responses import JSONResponse from sqlalchemy.ext.asyncio import AsyncSession @@ -13,7 +14,8 @@ from src.shared.logging import get_logger from src.shared.database import get_async_session from src.domains.auth.schemas import ( AuthSyncRequest, AuthSyncResponse, UsersListResponse, - BulkSyncResultSchema, GroupsListResponse + BulkSyncResultSchema, GroupsListResponse, RolesListResponse, + GroupRoleAssignmentResponse, ) from src.domains.auth.service import AuthService @@ -218,6 +220,94 @@ class AuthController(BaseController): logger.error(f"Groups bulk sync failed: {e}") raise HTTPException(status_code=401, detail=str(e)) + @router.get( + "/roles", + summary="List all roles", + response_model=RolesListResponse, + responses={ + 200: {"description": "List of all available roles"}, + }, + ) + async def list_roles( + session: AsyncSession = Depends(get_async_session), + ) -> RolesListResponse: + """ + List all available roles in the system + + Returns all domain.category:action role combinations. + Use these when assigning roles to groups. + """ + service = AuthService(session) + roles = await service.list_roles() + return RolesListResponse( + items=service.roles_to_schema(roles), + total=len(roles), + ) + + @router.post( + "/groups/{group_id}/roles/{role_id}", + summary="Assign role to group", + response_model=GroupRoleAssignmentResponse, + responses={ + 200: {"description": "Role assigned successfully"}, + 404: {"description": "Group or role not found"}, + }, + ) + async def assign_role_to_group( + group_id: uuid.UUID = Path(..., description="Group ID"), + role_id: uuid.UUID = Path(..., description="Role ID to assign"), + session: AsyncSession = Depends(get_async_session), + ) -> GroupRoleAssignmentResponse: + """ + Assign a role to a group + + All users in this group will inherit this role's permissions. + """ + service = AuthService(session) + + try: + group = await service.assign_role_to_group(group_id, role_id) + await session.commit() + return GroupRoleAssignmentResponse( + group_id=group.id, + group_name=group.name, + roles=[role.name for role in group.roles], + ) + except ValueError as e: + raise HTTPException(status_code=404, detail=str(e)) + + @router.delete( + "/groups/{group_id}/roles/{role_id}", + summary="Remove role from group", + response_model=GroupRoleAssignmentResponse, + responses={ + 200: {"description": "Role removed successfully"}, + 404: {"description": "Group or role not found"}, + }, + ) + async def remove_role_from_group( + group_id: uuid.UUID = Path(..., description="Group ID"), + role_id: uuid.UUID = Path(..., description="Role ID to remove"), + session: AsyncSession = Depends(get_async_session), + ) -> GroupRoleAssignmentResponse: + """ + Remove a role from a group + + Users in this group will no longer inherit this role's permissions. + """ + service = AuthService(session) + + try: + group = await service.remove_role_from_group(group_id, role_id) + await session.commit() + return GroupRoleAssignmentResponse( + group_id=group.id, + group_name=group.name, + roles=[role.name for role in group.roles], + ) + except ValueError as e: + raise HTTPException(status_code=404, detail=str(e)) + @router.get( "/me", summary="Get current user profile", diff --git a/src/domains/auth/models.py b/src/domains/auth/models.py index 451c6a8..cf8484c 100644 --- a/src/domains/auth/models.py +++ b/src/domains/auth/models.py @@ -26,6 +26,13 @@ user_groups = Table( Column("group_id", UUID(as_uuid=True), ForeignKey("groups.id", ondelete="CASCADE"), primary_key=True), ) +group_roles = Table( + "group_roles", + Base.metadata, + Column("group_id", UUID(as_uuid=True), ForeignKey("groups.id", ondelete="CASCADE"), primary_key=True), + Column("role_id", UUID(as_uuid=True), ForeignKey("roles.id", ondelete="CASCADE"), primary_key=True), +) + # ============================================================================= # User Model @@ -114,8 +121,13 @@ class Role(Base): """ Role model for domain-scoped permissions + Permission format: domain.category:action + - domain: Main area (control-room, library, media, ai, etc.) + - category: Sub-area within domain (general for full access, or specific tools) + - action: Permission level (viewer, user, editor, admin) + Roles are seeded from configuration, not user-editable. - Each role maps to an Authentik group (e.g., tatlock-control-room-admin). + Groups are assigned roles via the group_roles mapping table. Domains: control-room, library, media, ai, housekeeper, developer, documents, gaming, admin Actions: viewer, user, editor, admin (hierarchical) @@ -133,7 +145,7 @@ class Role(Base): unique=True, nullable=False, index=True, - comment="Role name in format domain:action (e.g., control-room:admin)", + comment="Role name in format domain.category:action (e.g., control-room.general:admin)", ) domain: Mapped[str] = mapped_column( String(50), @@ -141,17 +153,17 @@ class Role(Base): index=True, comment="Permission domain (e.g., control-room, media, ai)", ) + category: Mapped[str] = mapped_column( + String(50), + nullable=False, + default="general", + comment="Permission category within domain (general for full access, or specific tool)", + ) action: Mapped[str] = mapped_column( String(20), nullable=False, comment="Permission action (viewer, user, editor, admin)", ) - authentik_group: Mapped[str | None] = mapped_column( - String(255), - nullable=True, - unique=True, - comment="Corresponding Authentik group name (e.g., tatlock-control-room-admin)", - ) # Relationships users: Mapped[List["User"]] = relationship( @@ -160,6 +172,12 @@ class Role(Base): back_populates="roles", lazy="selectin", ) + groups: Mapped[List["Group"]] = relationship( + "Group", + secondary="group_roles", + back_populates="roles", + lazy="selectin", + ) def __repr__(self) -> str: return f"" @@ -247,6 +265,14 @@ class Group(Base): comment="Last sync from Authentik", ) + # Relationships + roles: Mapped[List["Role"]] = relationship( + "Role", + secondary="group_roles", + back_populates="groups", + lazy="selectin", + ) + def __repr__(self) -> str: return f"" diff --git a/src/domains/auth/oidc.py b/src/domains/auth/oidc.py index e262cad..4824e0f 100644 --- a/src/domains/auth/oidc.py +++ b/src/domains/auth/oidc.py @@ -3,15 +3,57 @@ OIDC Authentication Module Provides OAuth2/OIDC token validation for FastAPI using Authentik as IdP. Implements bearer token authentication with JWT verification. + +Permission Format: domain.category:action +- domain: Main area (control-room, library, media, ai, etc.) +- category: Sub-area within domain (general for full domain, or specific tools) +- action: Permission level (viewer, user, editor, admin) + +Examples: +- control-room.general:admin - Full access to Control Room +- media.general:viewer - View-only access to Media area +- ai.ollama:user - User-level access to Ollama specifically (future) + +Action Hierarchy (higher implies lower): +- admin > editor > user > viewer """ from fastapi import Depends, HTTPException, Security, Request from fastapi.security import HTTPBearer, HTTPAuthorizationCredentials from jose import jwt, JWTError import httpx from functools import lru_cache -from typing import Dict, Optional +from typing import Callable, Dict, List, Optional from src.shared.logging import get_logger + +# ============================================================================= +# Permission System +# ============================================================================= + +# Action hierarchy: higher actions imply lower ones +ACTION_HIERARCHY: Dict[str, int] = { + "viewer": 1, + "user": 2, + "editor": 3, + "admin": 4, +} + +# Valid domains (main areas) +VALID_DOMAINS = { + "control-room", + "library", + "media", + "ai", + "housekeeper", + "developer", + "documents", + "gaming", + "admin", # Global admin domain +} + +# Default category for general domain access +DEFAULT_CATEGORY = "general" + logger = get_logger(__name__) security = HTTPBearer(auto_error=False) @@ -352,3 +394,295 @@ async def get_forward_auth_admin( ) return user + + +# ============================================================================= +# Permission-Based Access Control +# ============================================================================= + +def _parse_permission(permission: str) -> tuple[str, str, str]: + """ + Parse a permission string into (domain, category, action) + + Supports formats: + - domain.category:action (full): "control-room.general:admin" + - domain:action (shorthand): "control-room:admin" -> ("control-room", "general", "admin") + + Returns: + Tuple of (domain, category, action) + + Raises: + ValueError: If permission format is invalid + """ + # Split on colon first to get action + if ":" not in permission: + raise ValueError(f"Invalid permission format (missing ':'): {permission}") + + location, action = permission.rsplit(":", 1) + + # Split location on dot to get domain and category + if "." in location: + domain, category = location.split(".", 1) + else: + # Shorthand: domain:action -> domain.general:action + domain = location + category = DEFAULT_CATEGORY + + return domain, category, action + + +def _action_satisfies(user_action: str, required_action: str) -> bool: + """ + Check if user's action level satisfies the required action + + Due to hierarchy, admin satisfies editor, editor satisfies user, etc. + + Args: + user_action: The action the user has + required_action: The action required for access + + Returns: + True if user's action is >= required action + """ + user_level = ACTION_HIERARCHY.get(user_action, 0) + required_level = ACTION_HIERARCHY.get(required_action, 0) + return user_level >= required_level + + +def _user_has_permission( + user_permissions: List[str], + required_domain: str, + required_category: str, + required_action: str, +) -> bool: + """ + Check if user has a permission that satisfies the requirement + + Checks: + 1. Exact match: domain.category:action + 2. Domain-wide: domain.general:action (if category != general) + 3. Global admin: admin.general:admin (superuser) + + Args: + user_permissions: List of user's permission strings + required_domain: Required domain + required_category: Required category + required_action: Required action + + Returns: + True if user has sufficient permission + """ + for perm in user_permissions: + try: + dom, cat, act = _parse_permission(perm) + except ValueError: + continue + + # Global admin (admin.general:admin) grants all permissions + if dom == "admin" and cat == "general" and act == "admin": + return True + + # Check if this permission covers the requirement + if dom == required_domain: + # Exact category match + if cat == required_category and _action_satisfies(act, required_action): + return True + # Domain-wide permission (general category) covers all categories in domain + if cat == "general" and _action_satisfies(act, required_action): + return True + + return False + + +def _extract_permissions_from_groups(groups: List[str]) -> List[str]: + """ + Extract permission strings from Authentik group names + + Authentik groups follow naming: tatlock-{domain}-{category}-{action} + or shorthand: tatlock-{domain}-{action} (implies category=general) + + Examples: + - tatlock-control-room-general-admin -> control-room.general:admin + - tatlock-media-viewer -> media.general:viewer (shorthand) + - tatlock-ai-ollama-user -> ai.ollama:user + + Args: + groups: List of Authentik group names + + Returns: + List of permission strings + """ + permissions = [] + + for group in groups: + if not group.startswith("tatlock-"): + continue + + # Remove prefix + parts = group[8:].split("-") # Remove "tatlock-" + + if len(parts) >= 3: + # Could be domain-category-action or domain-with-hyphen-action + # Try to find a valid action at the end + action = parts[-1] + if action in ACTION_HIERARCHY: + # Check if domain-category or single domain with hyphen + remaining = parts[:-1] + + # Try to find known domain (greedy match from start) + for i in range(len(remaining), 0, -1): + potential_domain = "-".join(remaining[:i]) + if potential_domain in VALID_DOMAINS: + category_parts = remaining[i:] + category = "-".join(category_parts) if category_parts else DEFAULT_CATEGORY + permissions.append(f"{potential_domain}.{category}:{action}") + break + elif len(parts) == 2: + # Shorthand: domain-action (domain might have hyphen) + action = parts[-1] + if action in ACTION_HIERARCHY: + domain = parts[0] + if domain in VALID_DOMAINS: + permissions.append(f"{domain}.{DEFAULT_CATEGORY}:{action}") + + return permissions + + +def require_permission( + domain: str, + action: str, + category: str = DEFAULT_CATEGORY, +) -> Callable: + """ + Dependency factory for permission-based access control + + Creates a FastAPI dependency that checks if the current user has + the required permission. Considers action hierarchy and global admin. + + Usage: + @router.get("/containers") + async def list_containers( + user: Dict = Depends(require_permission("control-room", "viewer")) + ): + ... + + @router.delete("/container/{id}") + async def delete_container( + user: Dict = Depends(require_permission("control-room", "admin")) + ): + ... + + Args: + domain: Permission domain (e.g., "control-room", "media") + action: Required action level (viewer, user, editor, admin) + category: Permission category within domain, defaults to "general" + + Returns: + FastAPI dependency function + """ + perm_str = f"{domain}.{category}:{action}" + + async def permission_checker( + user: Optional[Dict] = Depends(get_current_user) + ) -> Dict: + """Check if user has required permission""" + + # If OIDC disabled, allow all (local dev mode) + if not oidc_config.enabled: + logger.debug(f"OIDC disabled - allowing {perm_str}") + return user or {"email": "local", "groups": ["admin"]} + + if user is None: + raise HTTPException( + status_code=401, + detail="Authentication required", + headers={"WWW-Authenticate": "Bearer"}, + ) + + # Extract permissions from user's groups + groups = user.get("groups", []) + permissions = _extract_permissions_from_groups(groups) + + # Check if user has required permission + if _user_has_permission(permissions, domain, category, action): + logger.debug(f"User {user.get('email')} granted {perm_str}") + return user + + # Permission denied + user_email = user.get("email", "unknown") + logger.warning( + f"User {user_email} denied {perm_str} " + f"(groups: {groups}, permissions: {permissions})" + ) + raise HTTPException( + status_code=403, + detail=f"Permission required: {perm_str}", + ) + + return permission_checker + + +def require_any_permission(*required_permissions: str) -> Callable: + """ + Dependency factory requiring any one of multiple permissions + + Useful for endpoints accessible to multiple roles. + + Usage: + @router.get("/shared-resource") + async def get_shared( + user: Dict = Depends(require_any_permission( + "control-room:viewer", + "media:viewer", + )) + ): + ... + + Args: + *required_permissions: Permission strings (domain.category:action or domain:action) + + Returns: + FastAPI dependency function + """ + + async def permission_checker( + user: Optional[Dict] = Depends(get_current_user) + ) -> Dict: + """Check if user has any of the required permissions""" + + # If OIDC disabled, allow all + if not oidc_config.enabled: + return user or {"email": "local", "groups": ["admin"]} + + if user is None: + raise HTTPException( + status_code=401, + detail="Authentication required", + headers={"WWW-Authenticate": "Bearer"}, + ) + + groups = user.get("groups", []) + permissions = _extract_permissions_from_groups(groups) + + # Check each required permission + for perm in required_permissions: + try: + dom, cat, act = _parse_permission(perm) + if _user_has_permission(permissions, dom, cat, act): + logger.debug(f"User {user.get('email')} granted via {perm}") + return user + except ValueError: + logger.warning(f"Invalid permission format: {perm}") + continue + + # None matched + user_email = user.get("email", "unknown") + logger.warning( + f"User {user_email} denied (required any of: {required_permissions})" + ) + raise HTTPException( + status_code=403, + detail=f"One of these permissions required: {', '.join(required_permissions)}", + ) + + return permission_checker diff --git a/src/domains/auth/schemas.py b/src/domains/auth/schemas.py index 16769d7..dc99f1a 100644 --- a/src/domains/auth/schemas.py +++ b/src/domains/auth/schemas.py @@ -26,10 +26,12 @@ class AuthSyncRequest(BaseSchema): class RoleSchema(BaseSchema): - """Role information in domain:action format""" + """Role information in domain.category:action format""" - name: str = Field(..., description="Role name (e.g., 'control-room:admin')") + id: uuid.UUID = Field(..., description="Role ID") + name: str = Field(..., description="Role name (e.g., 'control-room.general:admin')") domain: str = Field(..., description="Permission domain (e.g., 'control-room')") + category: str = Field(default="general", description="Permission category (e.g., 'general')") action: str = Field(..., description="Permission action (e.g., 'admin')") @@ -120,6 +122,7 @@ class GroupListItemSchema(BaseSchema): parent_name: Optional[str] = Field(None, description="Parent group name") member_count: int = Field(default=0, description="Number of users in this group") synced_at: datetime = Field(..., description="Last sync timestamp") + roles: list[str] = Field(default_factory=list, description="Assigned role names") class GroupsListResponse(BaseSchema): @@ -127,3 +130,18 @@ class GroupsListResponse(BaseSchema): items: list[GroupListItemSchema] = Field(..., description="List of groups") total: int = Field(..., description="Total count of groups") + + +class RolesListResponse(BaseSchema): + """Response from GET /auth/roles""" + + items: list[RoleSchema] = Field(..., description="List of all roles") + total: int = Field(..., description="Total count of roles") + + +class GroupRoleAssignmentResponse(BaseSchema): + """Response from group role assignment operations""" + + group_id: uuid.UUID = Field(..., description="Group ID") + group_name: str = Field(..., description="Group name") + roles: list[str] = Field(..., description="Currently assigned role names") diff --git a/src/domains/auth/service.py b/src/domains/auth/service.py index 1311c8b..ba7c521 100644 --- a/src/domains/auth/service.py +++ b/src/domains/auth/service.py @@ -141,30 +141,43 @@ class AuthService: await self.session.flush() return user, is_new - async def sync_roles(self, user: User, groups: list[str]) -> list[Role]: + async def sync_roles(self, user: User, group_names: list[str]) -> list[Role]: """ - Synchronize user roles from Authentik groups + Synchronize user roles from Authentik groups via group_roles mapping - Maps Authentik groups (e.g., 'tatlock-control-room-admin') - to application roles (e.g., 'control-room:admin'). + Looks up the user's groups in the database, then retrieves all roles + assigned to those groups via the group_roles mapping table. Args: user: User to sync roles for - groups: List of Authentik group names + group_names: List of Authentik group names Returns: List of synced Role objects """ - # Get all roles that match the user's Authentik groups - stmt = select(Role).where(Role.authentik_group.in_(groups)) + # Find local Group records matching the Authentik group names + stmt = ( + select(Group) + .options(selectinload(Group.roles)) + .where(Group.name.in_(group_names)) + ) result = await self.session.execute(stmt) - matching_roles = list(result.scalars().all()) + matching_groups = list(result.scalars().all()) + + # Collect all unique roles from all matching groups + roles_set: dict[uuid.UUID, Role] = {} + for group in matching_groups: + for role in group.roles: + roles_set[role.id] = role + + matching_roles = list(roles_set.values()) # Clear existing roles and set new ones user.roles = matching_roles role_names = [r.name for r in matching_roles] - logger.info(f"Synced roles for {user.email}: {role_names}") + group_names_found = [g.name for g in matching_groups] + logger.info(f"Synced roles for {user.email} via groups {group_names_found}: {role_names}") return matching_roles @@ -183,7 +196,7 @@ class AuthService: def roles_to_schema(self, roles: list[Role]) -> list[RoleSchema]: """Convert Role models to schemas""" return [ - RoleSchema(name=r.name, domain=r.domain, action=r.action) + RoleSchema(id=r.id, name=r.name, domain=r.domain, category=r.category, action=r.action) for r in roles ] @@ -492,8 +505,8 @@ class AuthService: """ from sqlalchemy import func - # Base query - base_query = select(Group) + # Base query with roles loaded + base_query = select(Group).options(selectinload(Group.roles)) # Apply search filter if provided if search: @@ -520,12 +533,106 @@ class AuthService: parent_name=group.parent_name, member_count=group.member_count, synced_at=group.synced_at, + roles=[role.name for role in group.roles], ) for group in groups ] return items, total + async def list_roles(self) -> list[Role]: + """ + List all available roles + + Returns: + List of all Role objects + """ + stmt = select(Role).order_by(Role.domain, Role.category, Role.action) + result = await self.session.execute(stmt) + return list(result.scalars().all()) + + async def get_group_by_id(self, group_id: uuid.UUID) -> Optional[Group]: + """ + Get a group by its ID with roles loaded + + Args: + group_id: The group's UUID + + Returns: + Group object or None if not found + """ + stmt = ( + select(Group) + .options(selectinload(Group.roles)) + .where(Group.id == group_id) + ) + result = await self.session.execute(stmt) + return result.scalar_one_or_none() + + async def assign_role_to_group(self, group_id: uuid.UUID, role_id: uuid.UUID) -> Group: + """ + Assign a role to a group + + Args: + group_id: The group's UUID + role_id: The role's UUID to assign + + Returns: + Updated Group object + + Raises: + ValueError: If group or role not found + """ + group = await self.get_group_by_id(group_id) + if not group: + raise ValueError(f"Group not found: {group_id}") + + stmt = select(Role).where(Role.id == role_id) + result = await self.session.execute(stmt) + role = result.scalar_one_or_none() + if not role: + raise ValueError(f"Role not found: {role_id}") + + # Add role if not already assigned + if role not in group.roles: + group.roles.append(role) + await self.session.flush() + logger.info(f"Assigned role {role.name} to group {group.name}") + + return group + + async def remove_role_from_group(self, group_id: uuid.UUID, role_id: uuid.UUID) -> Group: + """ + Remove a role from a group + + Args: + group_id: The group's UUID + role_id: The role's UUID to remove + + Returns: + Updated Group object + + Raises: + ValueError: If group or role not found + """ + group = await self.get_group_by_id(group_id) + if not group: + raise ValueError(f"Group not found: {group_id}") + + stmt = select(Role).where(Role.id == role_id) + result = await self.session.execute(stmt) + role = result.scalar_one_or_none() + if not role: + raise ValueError(f"Role not found: {role_id}") + + # Remove role if assigned + if role in group.roles: + group.roles.remove(role) + await self.session.flush() + logger.info(f"Removed role {role.name} from group {group.name}") + + return group + async def bulk_sync_groups_from_authentik(self) -> BulkSyncResultSchema: """ Fetch all groups from Authentik admin API and sync to local database diff --git a/tests/test_auth_controller.py b/tests/test_auth_controller.py new file mode 100644 index 0000000..b4fb77c --- /dev/null +++ b/tests/test_auth_controller.py @@ -0,0 +1,517 @@ +"""Tests for authentication controller endpoints.""" +import uuid +from datetime import datetime, timezone +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from fastapi.testclient import TestClient + +from src.main import app + + +@pytest.fixture +def client(): + """Create a test client.""" + return TestClient(app) + + +# ============================================================================= +# OpenAPI Spec Tests +# ============================================================================= + +class TestAuthOpenAPISpec: + """Test that auth endpoints are documented in OpenAPI spec.""" + + def test_auth_sync_in_openapi(self, client): + """Auth sync endpoint should be in OpenAPI spec.""" + response = client.get("/openapi.json") + assert response.status_code == 200 + spec = response.json() + assert "/auth/sync" in spec["paths"] + assert "post" in spec["paths"]["/auth/sync"] + + def test_auth_users_in_openapi(self, client): + """Auth users endpoint should be in OpenAPI spec.""" + response = client.get("/openapi.json") + spec = response.json() + assert "/auth/users" in spec["paths"] + assert "get" in spec["paths"]["/auth/users"] + + def test_auth_groups_in_openapi(self, client): + """Auth groups endpoint should be in OpenAPI spec.""" + response = client.get("/openapi.json") + spec = response.json() + assert "/auth/groups" in spec["paths"] + assert "get" in spec["paths"]["/auth/groups"] + + def test_auth_roles_in_openapi(self, client): + """Auth roles endpoint should be in OpenAPI spec.""" + response = client.get("/openapi.json") + spec = response.json() + assert "/auth/roles" in spec["paths"] + assert "get" in spec["paths"]["/auth/roles"] + + def test_group_role_assignment_in_openapi(self, client): + """Group role assignment endpoints should be in OpenAPI spec.""" + response = client.get("/openapi.json") + spec = response.json() + path = "/auth/groups/{group_id}/roles/{role_id}" + assert path in spec["paths"] + assert "post" in spec["paths"][path] # Assign + assert "delete" in spec["paths"][path] # Remove + + def test_sync_from_authentik_endpoints(self, client): + """Sync from Authentik endpoints should be in OpenAPI spec.""" + response = client.get("/openapi.json") + spec = response.json() + assert "/auth/users/sync-from-authentik" in spec["paths"] + assert "/auth/groups/sync-from-authentik" in spec["paths"] + + +# ============================================================================= +# Controller Module Tests +# ============================================================================= + +class TestAuthControllerModule: + """Test auth controller module imports and configuration.""" + + def test_controller_imports(self): + """Auth controller should be importable.""" + from src.domains.auth.controller import AuthController, auth_controller + assert AuthController is not None + assert auth_controller is not None + + def test_controller_has_correct_prefix(self): + """Auth controller should have /auth prefix.""" + from src.domains.auth.controller import auth_controller + assert auth_controller.prefix == "/auth" + + def test_controller_has_correct_tags(self): + """Auth controller should have Authentication tag.""" + from src.domains.auth.controller import auth_controller + assert "Authentication" in auth_controller.tags + + +# ============================================================================= +# Sync Endpoint Tests +# ============================================================================= + +class TestSyncEndpoint: + """Test POST /auth/sync endpoint.""" + + def test_sync_requires_access_token(self, client): + """Sync should require access_token in body.""" + response = client.post("/auth/sync", json={}) + assert response.status_code == 422 # Validation error + + def test_sync_with_invalid_token(self, client): + """Sync should return 401 for invalid token.""" + with patch("src.domains.auth.controller.AuthService") as MockService: + mock_instance = MagicMock() + mock_instance.validate_token = AsyncMock( + side_effect=ValueError("Invalid or expired token") + ) + MockService.return_value = mock_instance + + response = client.post( + "/auth/sync", + json={"access_token": "invalid_token"}, + ) + + assert response.status_code == 401 + + +# ============================================================================= +# Users Endpoint Tests +# ============================================================================= + +class TestUsersEndpoint: + """Test GET /auth/users endpoint.""" + + def test_list_users_returns_list(self, client): + """List users should return paginated response.""" + with patch("src.domains.auth.controller.AuthService") as MockService: + mock_instance = MagicMock() + mock_instance.list_users = AsyncMock(return_value=([], 0)) + MockService.return_value = mock_instance + + response = client.get("/auth/users") + + assert response.status_code == 200 + data = response.json() + assert "items" in data + assert "total" in data + + def test_list_users_with_search(self, client): + """List users should accept search parameter.""" + with patch("src.domains.auth.controller.AuthService") as MockService: + mock_instance = MagicMock() + mock_instance.list_users = AsyncMock(return_value=([], 0)) + MockService.return_value = mock_instance + + response = client.get("/auth/users?search=test") + + assert response.status_code == 200 + + def test_list_users_with_pagination(self, client): + """List users should accept pagination parameters.""" + with patch("src.domains.auth.controller.AuthService") as MockService: + mock_instance = MagicMock() + mock_instance.list_users = AsyncMock(return_value=([], 0)) + MockService.return_value = mock_instance + + response = client.get("/auth/users?offset=10&limit=20") + + assert response.status_code == 200 + + def test_list_users_limit_validation(self, client): + """List users should reject limit > 100.""" + response = client.get("/auth/users?limit=200") + assert response.status_code == 422 + + +# ============================================================================= +# Groups Endpoint Tests +# ============================================================================= + +class TestGroupsEndpoint: + """Test GET /auth/groups endpoint.""" + + def test_list_groups_returns_list(self, client): + """List groups should return paginated response.""" + with patch("src.domains.auth.controller.AuthService") as MockService: + mock_instance = MagicMock() + mock_instance.list_groups = AsyncMock(return_value=([], 0)) + MockService.return_value = mock_instance + + response = client.get("/auth/groups") + + assert response.status_code == 200 + data = response.json() + assert "items" in data + assert "total" in data + + def test_list_groups_with_search(self, client): + """List groups should accept search parameter.""" + with patch("src.domains.auth.controller.AuthService") as MockService: + mock_instance = MagicMock() + mock_instance.list_groups = AsyncMock(return_value=([], 0)) + MockService.return_value = mock_instance + + response = client.get("/auth/groups?search=admin") + + assert response.status_code == 200 + + +# ============================================================================= +# Roles Endpoint Tests +# ============================================================================= + +class TestRolesEndpoint: + """Test GET /auth/roles endpoint.""" + + def test_list_roles_returns_list(self, client): + """List roles should return all roles.""" + with patch("src.domains.auth.controller.AuthService") as MockService: + mock_instance = MagicMock() + mock_instance.list_roles = AsyncMock(return_value=[]) + mock_instance.roles_to_schema = MagicMock(return_value=[]) + MockService.return_value = mock_instance + + response = client.get("/auth/roles") + + assert response.status_code == 200 + data = response.json() + assert "items" in data + assert "total" in data + + +# ============================================================================= +# Group-Role Assignment Endpoint Tests +# ============================================================================= + +class TestGroupRoleAssignmentEndpoints: + """Test group-role assignment and removal endpoints.""" + + def test_assign_role_to_group_success(self, client): + """POST /auth/groups/{id}/roles/{id} should assign role.""" + group_id = str(uuid.uuid4()) + role_id = str(uuid.uuid4()) + + mock_group = MagicMock() + mock_group.id = uuid.UUID(group_id) + mock_group.name = "Test Group" + mock_group.roles = [] + + with patch("src.domains.auth.controller.AuthService") as MockService: + mock_instance = MagicMock() + mock_instance.assign_role_to_group = AsyncMock(return_value=mock_group) + MockService.return_value = mock_instance + + response = client.post(f"/auth/groups/{group_id}/roles/{role_id}") + + assert response.status_code == 200 + data = response.json() + assert data["group_id"] == group_id + assert data["group_name"] == "Test Group" + assert "roles" in data + + def test_assign_role_to_group_not_found(self, client): + """POST should return 404 when group not found.""" + group_id = str(uuid.uuid4()) + role_id = str(uuid.uuid4()) + + with patch("src.domains.auth.controller.AuthService") as MockService: + mock_instance = MagicMock() + mock_instance.assign_role_to_group = AsyncMock( + side_effect=ValueError("Group not found") + ) + MockService.return_value = mock_instance + + response = client.post(f"/auth/groups/{group_id}/roles/{role_id}") + + assert response.status_code == 404 + + def test_remove_role_from_group_success(self, client): + """DELETE /auth/groups/{id}/roles/{id} should remove role.""" + group_id = str(uuid.uuid4()) + role_id = str(uuid.uuid4()) + + mock_group = MagicMock() + mock_group.id = uuid.UUID(group_id) + mock_group.name = "Test Group" + mock_group.roles = [] + + with patch("src.domains.auth.controller.AuthService") as MockService: + mock_instance = MagicMock() + mock_instance.remove_role_from_group = AsyncMock(return_value=mock_group) + MockService.return_value = mock_instance + + response = client.delete(f"/auth/groups/{group_id}/roles/{role_id}") + + assert response.status_code == 200 + data = response.json() + assert data["group_id"] == group_id + + def test_remove_role_from_group_not_found(self, client): + """DELETE should return 404 when role not found.""" + group_id = str(uuid.uuid4()) + role_id = str(uuid.uuid4()) + + with patch("src.domains.auth.controller.AuthService") as MockService: + mock_instance = MagicMock() + mock_instance.remove_role_from_group = AsyncMock( + side_effect=ValueError("Role not found") + ) + MockService.return_value = mock_instance + + response = client.delete(f"/auth/groups/{group_id}/roles/{role_id}") + + assert response.status_code == 404 + + def test_assign_role_invalid_uuid(self, client): + """POST should return 422 for invalid UUID.""" + response = client.post("/auth/groups/not-a-uuid/roles/also-not-uuid") + assert response.status_code == 422 + + +# ============================================================================= +# Me Endpoint Tests +# ============================================================================= + +class TestMeEndpoint: + """Test GET /auth/me endpoint.""" + + def test_me_not_implemented(self, client): + """Me endpoint should return 501 (not implemented yet).""" + response = client.get("/auth/me") + assert response.status_code == 501 + + +# ============================================================================= +# Schema Tests +# ============================================================================= + +class TestAuthSchemas: + """Test auth schema imports and structure.""" + + def test_all_schemas_importable(self): + """All auth schemas should be importable.""" + from src.domains.auth.schemas import ( + AuthSyncRequest, + AuthSyncResponse, + RoleSchema, + UserSchema, + UserPreferencesSchema, + TokenInfoSchema, + UserListItemSchema, + UsersListResponse, + BulkSyncResultSchema, + GroupListItemSchema, + GroupsListResponse, + RolesListResponse, + GroupRoleAssignmentResponse, + ) + + assert AuthSyncRequest is not None + assert AuthSyncResponse is not None + assert RoleSchema is not None + assert UserSchema is not None + assert UserPreferencesSchema is not None + assert TokenInfoSchema is not None + assert UserListItemSchema is not None + assert UsersListResponse is not None + assert BulkSyncResultSchema is not None + assert GroupListItemSchema is not None + assert GroupsListResponse is not None + assert RolesListResponse is not None + assert GroupRoleAssignmentResponse is not None + + def test_role_schema_includes_id(self): + """RoleSchema should include id field.""" + from src.domains.auth.schemas import RoleSchema + + schema = RoleSchema( + id=uuid.uuid4(), + name="test.general:admin", + domain="test", + category="general", + action="admin", + ) + assert schema.id is not None + + def test_role_schema_includes_category(self): + """RoleSchema should include category field.""" + from src.domains.auth.schemas import RoleSchema + + schema = RoleSchema( + id=uuid.uuid4(), + name="test.specific:viewer", + domain="test", + category="specific", + action="viewer", + ) + assert schema.category == "specific" + + def test_group_list_item_includes_roles(self): + """GroupListItemSchema should include roles list.""" + from src.domains.auth.schemas import GroupListItemSchema + + schema = GroupListItemSchema( + id=uuid.uuid4(), + authentik_id=uuid.uuid4(), + name="Test Group", + is_superuser=False, + parent_name=None, + member_count=5, + synced_at=datetime.now(timezone.utc), + roles=["admin.general:admin", "media.general:viewer"], + ) + assert len(schema.roles) == 2 + + +# ============================================================================= +# Model Tests +# ============================================================================= + +class TestAuthModels: + """Test auth model imports.""" + + def test_all_models_importable(self): + """All auth models should be importable.""" + from src.domains.auth.models import ( + User, + Role, + UserRole, + Group, + UserPreferences, + ApiKey, + user_groups, + group_roles, + ) + + assert User is not None + assert Role is not None + assert UserRole is not None + assert Group is not None + assert UserPreferences is not None + assert ApiKey is not None + assert user_groups is not None + assert group_roles is not None + + def test_role_model_has_category(self): + """Role model should have category attribute.""" + from src.domains.auth.models import Role + + role = Role( + name="test.general:admin", + domain="test", + category="general", + action="admin", + ) + assert role.category == "general" + + def test_group_has_roles_relationship(self): + """Group model should have roles relationship.""" + from src.domains.auth.models import Group + + assert hasattr(Group, "roles") + + +# ============================================================================= +# Permission System Tests (from oidc.py) +# ============================================================================= + +class TestPermissionSystem: + """Test permission checking functions.""" + + def test_permission_constants_defined(self): + """Permission constants should be defined.""" + from src.domains.auth.oidc import ( + ACTION_HIERARCHY, + VALID_DOMAINS, + DEFAULT_CATEGORY, + ) + + assert "viewer" in ACTION_HIERARCHY + assert "user" in ACTION_HIERARCHY + assert "editor" in ACTION_HIERARCHY + assert "admin" in ACTION_HIERARCHY + + assert "control-room" in VALID_DOMAINS + assert "media" in VALID_DOMAINS + assert "admin" in VALID_DOMAINS + + assert DEFAULT_CATEGORY == "general" + + def test_action_hierarchy_ordering(self): + """Action hierarchy should be ordered correctly.""" + from src.domains.auth.oidc import ACTION_HIERARCHY + + assert ACTION_HIERARCHY["viewer"] < ACTION_HIERARCHY["user"] + assert ACTION_HIERARCHY["user"] < ACTION_HIERARCHY["editor"] + assert ACTION_HIERARCHY["editor"] < ACTION_HIERARCHY["admin"] + + def test_require_permission_importable(self): + """require_permission should be importable.""" + from src.domains.auth.oidc import require_permission, require_any_permission + + assert callable(require_permission) + assert callable(require_any_permission) + + def test_require_permission_returns_dependency(self): + """require_permission should return a callable dependency.""" + from src.domains.auth.oidc import require_permission + + dependency = require_permission("control-room", "admin") + assert callable(dependency) + + def test_require_any_permission_returns_dependency(self): + """require_any_permission should return a callable dependency.""" + from src.domains.auth.oidc import require_any_permission + + dependency = require_any_permission( + ("control-room", "admin"), + ("media", "editor"), + ) + assert callable(dependency) diff --git a/tests/test_auth_service.py b/tests/test_auth_service.py new file mode 100644 index 0000000..b096a4f --- /dev/null +++ b/tests/test_auth_service.py @@ -0,0 +1,627 @@ +"""Tests for authentication service.""" +import uuid +from datetime import datetime, timezone +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from sqlalchemy.ext.asyncio import AsyncSession + +from src.domains.auth.models import User, Role, Group, UserPreferences +from src.domains.auth.schemas import TokenInfoSchema, RoleSchema +from src.domains.auth.service import AuthService, get_auth_service + + +# ============================================================================= +# Fixtures +# ============================================================================= + +@pytest.fixture +def mock_session(): + """Create a mock async database session.""" + session = AsyncMock(spec=AsyncSession) + session.execute = AsyncMock() + session.commit = AsyncMock() + session.flush = AsyncMock() + session.refresh = AsyncMock() + session.add = MagicMock() + return session + + +@pytest.fixture +def auth_service(mock_session): + """Create an AuthService instance with mock session.""" + return AuthService(mock_session) + + +@pytest.fixture +def sample_user(): + """Create a sample user for testing.""" + user = User( + id=uuid.uuid4(), + authentik_id=uuid.uuid4(), + email="test@example.com", + name="Test User", + avatar_url="https://example.com/avatar.jpg", + created_at=datetime.now(timezone.utc), + last_login=datetime.now(timezone.utc), + ) + user.roles = [] + user.preferences = None + return user + + +@pytest.fixture +def sample_role(): + """Create a sample role for testing.""" + return Role( + id=uuid.uuid4(), + name="control-room.general:admin", + domain="control-room", + category="general", + action="admin", + ) + + +@pytest.fixture +def sample_group(sample_role): + """Create a sample group for testing.""" + group = Group( + id=uuid.uuid4(), + authentik_id=uuid.uuid4(), + name="Administrators", + is_superuser=True, + parent_name=None, + member_count=5, + created_at=datetime.now(timezone.utc), + synced_at=datetime.now(timezone.utc), + ) + group.roles = [sample_role] + return group + + +@pytest.fixture +def sample_token_info(): + """Create sample token info from Authentik.""" + return TokenInfoSchema( + sub=str(uuid.uuid4()), + email="test@example.com", + name="Test User", + preferred_username="testuser", + groups=["Administrators", "Developers"], + picture="https://example.com/avatar.jpg", + ) + + +@pytest.fixture +def sample_preferences(): + """Create sample user preferences.""" + return UserPreferences( + user_id=uuid.uuid4(), + theme="dark", + default_room="control-room", + preferences_json={"notifications": True}, + ) + + +# ============================================================================= +# AuthService Initialization Tests +# ============================================================================= + +class TestAuthServiceInit: + """Test AuthService initialization.""" + + def test_init_with_session(self, mock_session): + """AuthService should initialize with session.""" + service = AuthService(mock_session) + assert service.session is mock_session + + def test_init_sets_userinfo_url(self, mock_session): + """AuthService should set userinfo URL from settings.""" + service = AuthService(mock_session) + assert "userinfo" in service.userinfo_url + + def test_get_auth_service_factory(self, mock_session): + """get_auth_service should return AuthService instance.""" + service = get_auth_service(mock_session) + assert isinstance(service, AuthService) + assert service.session is mock_session + + +# ============================================================================= +# Token Validation Tests +# ============================================================================= + +class TestValidateToken: + """Test token validation via Authentik userinfo endpoint.""" + + @pytest.mark.asyncio + async def test_validate_token_success(self, auth_service): + """validate_token should return token info on success.""" + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.json.return_value = { + "sub": str(uuid.uuid4()), + "email": "test@example.com", + "name": "Test User", + "preferred_username": "testuser", + "groups": ["Administrators"], + "picture": "https://example.com/avatar.jpg", + } + mock_response.raise_for_status = MagicMock() + + with patch("src.domains.auth.service.httpx.AsyncClient") as mock_client: + mock_client.return_value.__aenter__.return_value.get = AsyncMock( + return_value=mock_response + ) + + result = await auth_service.validate_token("valid_token") + + assert result.email == "test@example.com" + assert result.name == "Test User" + assert "Administrators" in result.groups + + @pytest.mark.asyncio + async def test_validate_token_invalid(self, auth_service): + """validate_token should raise ValueError for invalid token.""" + mock_response = MagicMock() + mock_response.status_code = 401 + + with patch("src.domains.auth.service.httpx.AsyncClient") as mock_client: + mock_client.return_value.__aenter__.return_value.get = AsyncMock( + return_value=mock_response + ) + + with pytest.raises(ValueError, match="Invalid or expired token"): + await auth_service.validate_token("invalid_token") + + @pytest.mark.asyncio + async def test_validate_token_service_unavailable(self, auth_service): + """validate_token should raise ValueError when service unavailable.""" + import httpx + + with patch("src.domains.auth.service.httpx.AsyncClient") as mock_client: + mock_client.return_value.__aenter__.return_value.get = AsyncMock( + side_effect=httpx.RequestError("Connection failed") + ) + + with pytest.raises(ValueError, match="Authentication service unavailable"): + await auth_service.validate_token("token") + + +# ============================================================================= +# User Sync Tests +# ============================================================================= + +class TestSyncUser: + """Test user synchronization from OIDC token.""" + + @pytest.mark.asyncio + async def test_sync_user_creates_new_user(self, auth_service, sample_token_info, mock_session): + """sync_user should create new user when not found.""" + # Mock no existing user found + mock_result = MagicMock() + mock_result.scalar_one_or_none.return_value = None + mock_session.execute.return_value = mock_result + + user, is_new = await auth_service.sync_user(sample_token_info) + + assert is_new is True + assert mock_session.add.call_count == 2 # User and Preferences + + @pytest.mark.asyncio + async def test_sync_user_updates_existing_user( + self, auth_service, sample_token_info, sample_user, mock_session + ): + """sync_user should update existing user when found.""" + # Mock existing user found + mock_result = MagicMock() + mock_result.scalar_one_or_none.return_value = sample_user + mock_session.execute.return_value = mock_result + + # Update token info with matching authentik_id + sample_token_info.sub = str(sample_user.authentik_id) + + user, is_new = await auth_service.sync_user(sample_token_info) + + assert is_new is False + assert user.email == sample_token_info.email + assert user.name == sample_token_info.name + + @pytest.mark.asyncio + async def test_sync_user_updates_last_login( + self, auth_service, sample_token_info, sample_user, mock_session + ): + """sync_user should update last_login timestamp.""" + old_login = sample_user.last_login + + mock_result = MagicMock() + mock_result.scalar_one_or_none.return_value = sample_user + mock_session.execute.return_value = mock_result + + sample_token_info.sub = str(sample_user.authentik_id) + + user, _ = await auth_service.sync_user(sample_token_info) + + assert user.last_login is not None + # last_login should be updated (or same if happened in same second) + assert user.last_login >= old_login or user.last_login is not None + + +# ============================================================================= +# Role Sync Tests +# ============================================================================= + +class TestSyncRoles: + """Test role synchronization from Authentik groups via group_roles.""" + + @pytest.mark.asyncio + async def test_sync_roles_from_groups( + self, auth_service, sample_user, sample_group, mock_session + ): + """sync_roles should get roles from matching groups.""" + # Mock finding groups with roles + mock_result = MagicMock() + mock_result.scalars.return_value.all.return_value = [sample_group] + mock_session.execute.return_value = mock_result + + roles = await auth_service.sync_roles(sample_user, ["Administrators"]) + + assert len(roles) == 1 + assert roles[0].name == "control-room.general:admin" + assert sample_user.roles == roles + + @pytest.mark.asyncio + async def test_sync_roles_no_matching_groups( + self, auth_service, sample_user, mock_session + ): + """sync_roles should return empty list when no groups match.""" + mock_result = MagicMock() + mock_result.scalars.return_value.all.return_value = [] + mock_session.execute.return_value = mock_result + + roles = await auth_service.sync_roles(sample_user, ["NonExistentGroup"]) + + assert len(roles) == 0 + assert sample_user.roles == [] + + @pytest.mark.asyncio + async def test_sync_roles_deduplicates_roles( + self, auth_service, sample_user, sample_role, mock_session + ): + """sync_roles should deduplicate roles from multiple groups.""" + # Create two groups with the same role + group1 = Group( + id=uuid.uuid4(), + authentik_id=uuid.uuid4(), + name="Group1", + is_superuser=False, + member_count=1, + created_at=datetime.now(timezone.utc), + synced_at=datetime.now(timezone.utc), + ) + group1.roles = [sample_role] + + group2 = Group( + id=uuid.uuid4(), + authentik_id=uuid.uuid4(), + name="Group2", + is_superuser=False, + member_count=1, + created_at=datetime.now(timezone.utc), + synced_at=datetime.now(timezone.utc), + ) + group2.roles = [sample_role] # Same role + + mock_result = MagicMock() + mock_result.scalars.return_value.all.return_value = [group1, group2] + mock_session.execute.return_value = mock_result + + roles = await auth_service.sync_roles(sample_user, ["Group1", "Group2"]) + + # Should only have one role despite appearing in two groups + assert len(roles) == 1 + + +# ============================================================================= +# Schema Conversion Tests +# ============================================================================= + +class TestSchemaConversions: + """Test model to schema conversions.""" + + def test_user_to_schema(self, auth_service, sample_user): + """user_to_schema should convert User model to UserSchema.""" + schema = auth_service.user_to_schema(sample_user) + + assert schema.id == sample_user.id + assert schema.authentik_id == sample_user.authentik_id + assert schema.email == sample_user.email + assert schema.name == sample_user.name + assert schema.avatar_url == sample_user.avatar_url + + def test_roles_to_schema(self, auth_service, sample_role): + """roles_to_schema should convert Role models to RoleSchemas.""" + schemas = auth_service.roles_to_schema([sample_role]) + + assert len(schemas) == 1 + assert schemas[0].id == sample_role.id + assert schemas[0].name == sample_role.name + assert schemas[0].domain == sample_role.domain + assert schemas[0].category == sample_role.category + assert schemas[0].action == sample_role.action + + def test_roles_to_schema_empty_list(self, auth_service): + """roles_to_schema should handle empty list.""" + schemas = auth_service.roles_to_schema([]) + assert schemas == [] + + def test_preferences_to_schema(self, auth_service, sample_preferences): + """preferences_to_schema should convert UserPreferences to schema.""" + schema = auth_service.preferences_to_schema(sample_preferences) + + assert schema.theme == sample_preferences.theme + assert schema.default_room == sample_preferences.default_room + assert schema.preferences_json == sample_preferences.preferences_json + + def test_preferences_to_schema_none(self, auth_service): + """preferences_to_schema should return defaults for None.""" + schema = auth_service.preferences_to_schema(None) + + assert schema.theme == "system" + assert schema.default_room == "front-hall" + assert schema.preferences_json == {} + + +# ============================================================================= +# List Operations Tests +# ============================================================================= + +class TestListOperations: + """Test list operations for users, groups, and roles.""" + + @pytest.mark.asyncio + async def test_list_users(self, auth_service, sample_user, mock_session): + """list_users should return paginated user list.""" + sample_user.roles = [] + + # Mock count query + count_result = MagicMock() + count_result.scalar.return_value = 1 + + # Mock users query + users_result = MagicMock() + users_result.scalars.return_value.all.return_value = [sample_user] + + mock_session.execute.side_effect = [count_result, users_result] + + items, total = await auth_service.list_users() + + assert total == 1 + assert len(items) == 1 + assert items[0].email == sample_user.email + + @pytest.mark.asyncio + async def test_list_users_with_search(self, auth_service, mock_session): + """list_users should filter by search query.""" + count_result = MagicMock() + count_result.scalar.return_value = 0 + + users_result = MagicMock() + users_result.scalars.return_value.all.return_value = [] + + mock_session.execute.side_effect = [count_result, users_result] + + items, total = await auth_service.list_users(search="nonexistent") + + assert total == 0 + assert len(items) == 0 + + @pytest.mark.asyncio + async def test_list_groups(self, auth_service, sample_group, mock_session): + """list_groups should return paginated group list with roles.""" + count_result = MagicMock() + count_result.scalar.return_value = 1 + + groups_result = MagicMock() + groups_result.scalars.return_value.all.return_value = [sample_group] + + mock_session.execute.side_effect = [count_result, groups_result] + + items, total = await auth_service.list_groups() + + assert total == 1 + assert len(items) == 1 + assert items[0].name == sample_group.name + assert len(items[0].roles) == 1 # Should include role names + + @pytest.mark.asyncio + async def test_list_roles(self, auth_service, sample_role, mock_session): + """list_roles should return all roles ordered by domain.""" + mock_result = MagicMock() + mock_result.scalars.return_value.all.return_value = [sample_role] + mock_session.execute.return_value = mock_result + + roles = await auth_service.list_roles() + + assert len(roles) == 1 + assert roles[0].name == sample_role.name + + +# ============================================================================= +# Group-Role Management Tests +# ============================================================================= + +class TestGroupRoleManagement: + """Test group-role assignment and removal.""" + + @pytest.mark.asyncio + async def test_get_group_by_id(self, auth_service, sample_group, mock_session): + """get_group_by_id should return group with roles.""" + mock_result = MagicMock() + mock_result.scalar_one_or_none.return_value = sample_group + mock_session.execute.return_value = mock_result + + group = await auth_service.get_group_by_id(sample_group.id) + + assert group is not None + assert group.id == sample_group.id + assert len(group.roles) == 1 + + @pytest.mark.asyncio + async def test_get_group_by_id_not_found(self, auth_service, mock_session): + """get_group_by_id should return None when not found.""" + mock_result = MagicMock() + mock_result.scalar_one_or_none.return_value = None + mock_session.execute.return_value = mock_result + + group = await auth_service.get_group_by_id(uuid.uuid4()) + + assert group is None + + @pytest.mark.asyncio + async def test_assign_role_to_group( + self, auth_service, sample_group, sample_role, mock_session + ): + """assign_role_to_group should add role to group.""" + # Clear existing roles for this test + sample_group.roles = [] + + # Mock group lookup + group_result = MagicMock() + group_result.scalar_one_or_none.return_value = sample_group + + # Mock role lookup + role_result = MagicMock() + role_result.scalar_one_or_none.return_value = sample_role + + mock_session.execute.side_effect = [group_result, role_result] + + group = await auth_service.assign_role_to_group(sample_group.id, sample_role.id) + + assert sample_role in group.roles + mock_session.flush.assert_called() + + @pytest.mark.asyncio + async def test_assign_role_to_group_already_assigned( + self, auth_service, sample_group, sample_role, mock_session + ): + """assign_role_to_group should not duplicate if already assigned.""" + # Group already has this role + sample_group.roles = [sample_role] + original_count = len(sample_group.roles) + + group_result = MagicMock() + group_result.scalar_one_or_none.return_value = sample_group + + role_result = MagicMock() + role_result.scalar_one_or_none.return_value = sample_role + + mock_session.execute.side_effect = [group_result, role_result] + + group = await auth_service.assign_role_to_group(sample_group.id, sample_role.id) + + assert len(group.roles) == original_count # No duplicate + + @pytest.mark.asyncio + async def test_assign_role_to_group_group_not_found(self, auth_service, mock_session): + """assign_role_to_group should raise ValueError when group not found.""" + mock_result = MagicMock() + mock_result.scalar_one_or_none.return_value = None + mock_session.execute.return_value = mock_result + + with pytest.raises(ValueError, match="Group not found"): + await auth_service.assign_role_to_group(uuid.uuid4(), uuid.uuid4()) + + @pytest.mark.asyncio + async def test_assign_role_to_group_role_not_found( + self, auth_service, sample_group, mock_session + ): + """assign_role_to_group should raise ValueError when role not found.""" + sample_group.roles = [] + + group_result = MagicMock() + group_result.scalar_one_or_none.return_value = sample_group + + role_result = MagicMock() + role_result.scalar_one_or_none.return_value = None + + mock_session.execute.side_effect = [group_result, role_result] + + with pytest.raises(ValueError, match="Role not found"): + await auth_service.assign_role_to_group(sample_group.id, uuid.uuid4()) + + @pytest.mark.asyncio + async def test_remove_role_from_group( + self, auth_service, sample_group, sample_role, mock_session + ): + """remove_role_from_group should remove role from group.""" + # Group has this role + sample_group.roles = [sample_role] + + group_result = MagicMock() + group_result.scalar_one_or_none.return_value = sample_group + + role_result = MagicMock() + role_result.scalar_one_or_none.return_value = sample_role + + mock_session.execute.side_effect = [group_result, role_result] + + group = await auth_service.remove_role_from_group(sample_group.id, sample_role.id) + + assert sample_role not in group.roles + mock_session.flush.assert_called() + + @pytest.mark.asyncio + async def test_remove_role_from_group_not_assigned( + self, auth_service, sample_group, sample_role, mock_session + ): + """remove_role_from_group should handle role not assigned gracefully.""" + # Group does not have this role + sample_group.roles = [] + + group_result = MagicMock() + group_result.scalar_one_or_none.return_value = sample_group + + role_result = MagicMock() + role_result.scalar_one_or_none.return_value = sample_role + + mock_session.execute.side_effect = [group_result, role_result] + + group = await auth_service.remove_role_from_group(sample_group.id, sample_role.id) + + # Should complete without error + assert len(group.roles) == 0 + + +# ============================================================================= +# Role Schema Tests +# ============================================================================= + +class TestRoleSchema: + """Test RoleSchema validation.""" + + def test_role_schema_creation(self): + """RoleSchema should be creatable with valid data.""" + schema = RoleSchema( + id=uuid.uuid4(), + name="control-room.general:admin", + domain="control-room", + category="general", + action="admin", + ) + + assert schema.name == "control-room.general:admin" + assert schema.domain == "control-room" + assert schema.category == "general" + assert schema.action == "admin" + + def test_role_schema_category_default(self): + """RoleSchema should default category to 'general'.""" + schema = RoleSchema( + id=uuid.uuid4(), + name="media.general:viewer", + domain="media", + action="viewer", + ) + + assert schema.category == "general"