feat(auth): implement group-role mapping and permission system

Architecture changes:
- Permission format: domain.category:action (e.g., control-room.general:admin)
- Decoupled groups from roles via group_roles mapping table
- Groups are organizational (synced from Authentik)
- Roles are permissions (admin-managed via API)

New features:
- require_permission() and require_any_permission() dependency factories
- Action hierarchy: admin > editor > user > viewer
- Global admin override (admin.general:admin grants all)
- Group-role management endpoints (assign/remove roles)
- GET /auth/roles endpoint to list all roles

Database changes:
- Added category column to roles table (default: general)
- Removed authentik_group column (decoupled)
- Added group_roles association table
- Added user_groups association table
- Migration updates role names to domain.general:action format

Tests:
- 67 new tests for auth service and controller
- Covers token validation, user sync, role sync
- Covers group-role assignment/removal
- Covers schema conversions and permission system

🤖 Generated with [Claude Code](https://claude.com/claude-code)

Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com>
This commit is contained in:
Jeroen Schweitzer
2026-01-03 19:52:10 +01:00
co-authored by Claude Opus 4.5
parent 075b0ec297
commit 7752cd9d23
9 changed files with 1899 additions and 25 deletions
@@ -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")
+14
View File
@@ -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",
]
+92 -2
View File
@@ -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",
+34 -8
View File
@@ -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"<Role {self.name}>"
@@ -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"<Group {self.name}>"
+335 -1
View File
@@ -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
+20 -2
View File
@@ -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")
+119 -12
View File
@@ -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
+517
View File
@@ -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)
+627
View File
@@ -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"