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:
co-authored by
Claude Opus 4.5
parent
075b0ec297
commit
7752cd9d23
@@ -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")
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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)
|
||||
@@ -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"
|
||||
Reference in New Issue
Block a user