from datetime import datetime

from sqlalchemy import case, func, or_, select
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import selectinload

from app.domain.entities.users_roles import CompanyUser, Role, RolePermission, UsersRolesSummary
from app.domain.enums import CompanyUserStatus
from app.domain.permission_menus import canonical_permission_key
from app.domain.repositories.users_roles_repository import UsersRolesRepository
from app.infrastructure.db.models import (
    CompanyUserModel,
    DepartmentModel,
    RoleModel,
    RolePermissionModel,
    new_id,
)

_USER_SORT_COLUMNS = {
    "full_name": CompanyUserModel.full_name,
    "email": CompanyUserModel.email,
    "status": CompanyUserModel.status,
    "last_login": CompanyUserModel.last_login,
    "created_at": CompanyUserModel.created_at,
    "employee_code": CompanyUserModel.employee_code,
    "username": CompanyUserModel.username,
}


class MySQLUsersRolesRepository(UsersRolesRepository):
    def __init__(self, session: AsyncSession) -> None:
        self._session = session

    def _to_permission(self, row: RolePermissionModel) -> RolePermission:
        return RolePermission(
            id=row.id,
            module=canonical_permission_key(row.module),
            can_view=bool(row.can_view),
            can_add=bool(row.can_add),
            can_edit=bool(row.can_edit),
            can_delete=bool(row.can_delete),
            can_export=bool(row.can_export),
        )

    def _to_role(self, row: RoleModel, *, users_count: int = 0) -> Role:
        by_module: dict[str, RolePermission] = {}
        for perm_row in row.permissions or []:
            perm = self._to_permission(perm_row)
            by_module[perm.module] = perm
        return Role(
            id=row.id,
            company_id=row.company_id,
            name=row.name,
            slug=row.slug,
            description=row.description,
            permissions=list(by_module.values()),
            users_count=users_count,
            created_by=row.created_by,
            created_at=row.created_at,
            updated_at=row.updated_at,
        )

    def _normalize_status(self, status: str) -> CompanyUserStatus:
        value = (status or "").strip().lower()
        if value == "inactive":
            value = CompanyUserStatus.DISABLED.value
        return CompanyUserStatus(value)

    def _to_user(
        self,
        row: CompanyUserModel,
        *,
        role_name: str | None = None,
        role_slug: str | None = None,
        department_code: str | None = None,
        department_name: str | None = None,
        include_secret: bool = False,
    ) -> CompanyUser:
        return CompanyUser(
            id=row.id,
            company_id=row.company_id,
            full_name=row.full_name,
            username=row.username,
            email=row.email,
            phone=row.phone,
            employee_code=row.employee_code,
            password_hash=row.password_hash if include_secret else None,
            role_id=row.role_id,
            role_name=role_name,
            role_slug=role_slug,
            department_id=row.department_id,
            department_code=department_code,
            department_name=department_name,
            status=self._normalize_status(row.status),
            send_welcome_email=bool(row.send_welcome_email),
            address=row.address,
            notes=row.notes,
            last_login=row.last_login,
            created_by=row.created_by,
            created_at=row.created_at,
            updated_at=row.updated_at,
        )

    def _user_from_join(self, row, *, include_secret: bool = False) -> CompanyUser:
        return self._to_user(
            row.CompanyUserModel,
            role_name=row.role_name,
            role_slug=row.role_slug,
            department_code=row.department_code,
            department_name=row.department_name,
            include_secret=include_secret,
        )

    def _user_lookup_query(self):
        return (
            select(
                CompanyUserModel,
                RoleModel.name.label("role_name"),
                RoleModel.slug.label("role_slug"),
                DepartmentModel.code.label("department_code"),
                DepartmentModel.name.label("department_name"),
            )
            .join(RoleModel, RoleModel.id == CompanyUserModel.role_id)
            .outerjoin(
                DepartmentModel, DepartmentModel.id == CompanyUserModel.department_id
            )
        )

    async def _users_count_map(self, company_id: str, role_ids: list[str]) -> dict[str, int]:
        if not role_ids:
            return {}
        result = await self._session.execute(
            select(CompanyUserModel.role_id, func.count())
            .where(
                CompanyUserModel.company_id == company_id,
                CompanyUserModel.role_id.in_(role_ids),
            )
            .group_by(CompanyUserModel.role_id)
        )
        return {role_id: int(count or 0) for role_id, count in result.all()}

    async def create_role(self, role: Role) -> Role:
        role_id = role.id or new_id()
        row = RoleModel(
            id=role_id,
            company_id=role.company_id,
            name=role.name,
            slug=role.slug,
            description=role.description,
            created_by=role.created_by,
            created_at=role.created_at,
            updated_at=role.updated_at,
            permissions=[
                RolePermissionModel(
                    id=new_id(),
                    role_id=role_id,
                    module=p.module if isinstance(p.module, str) else p.module.value,
                    can_view=p.can_view,
                    can_add=p.can_add,
                    can_edit=p.can_edit,
                    can_delete=p.can_delete,
                    can_export=p.can_export,
                )
                for p in role.permissions
            ],
        )
        self._session.add(row)
        await self._session.commit()
        return await self.get_role(role_id, role.company_id)  # type: ignore[return-value]

    async def get_role(self, role_id: str, company_id: str) -> Role | None:
        result = await self._session.execute(
            select(RoleModel)
            .options(selectinload(RoleModel.permissions))
            .where(RoleModel.id == role_id, RoleModel.company_id == company_id)
        )
        row = result.scalar_one_or_none()
        if not row:
            return None
        counts = await self._users_count_map(company_id, [row.id])
        return self._to_role(row, users_count=counts.get(row.id, 0))

    async def get_role_by_slug(self, company_id: str, slug: str) -> Role | None:
        result = await self._session.execute(
            select(RoleModel)
            .options(selectinload(RoleModel.permissions))
            .where(RoleModel.company_id == company_id, RoleModel.slug == slug)
        )
        row = result.scalar_one_or_none()
        if not row:
            return None
        counts = await self._users_count_map(company_id, [row.id])
        return self._to_role(row, users_count=counts.get(row.id, 0))

    async def get_role_by_name(self, company_id: str, name: str) -> Role | None:
        result = await self._session.execute(
            select(RoleModel)
            .options(selectinload(RoleModel.permissions))
            .where(RoleModel.company_id == company_id, RoleModel.name == name)
        )
        row = result.scalar_one_or_none()
        if not row:
            return None
        counts = await self._users_count_map(company_id, [row.id])
        return self._to_role(row, users_count=counts.get(row.id, 0))

    async def list_roles(
        self, company_id: str, *, search: str | None = None, skip: int = 0, limit: int = 100
    ) -> list[Role]:
        filters = [RoleModel.company_id == company_id]
        if search:
            term = f"%{search.strip().lower()}%"
            filters.append(
                or_(
                    func.lower(RoleModel.name).like(term),
                    func.lower(RoleModel.slug).like(term),
                )
            )
        result = await self._session.execute(
            select(RoleModel)
            .options(selectinload(RoleModel.permissions))
            .where(*filters)
            .order_by(RoleModel.name.asc())
            .offset(skip)
            .limit(limit)
        )
        rows = list(result.scalars().unique().all())
        counts = await self._users_count_map(company_id, [r.id for r in rows])
        return [self._to_role(row, users_count=counts.get(row.id, 0)) for row in rows]

    async def count_roles(self, company_id: str, *, search: str | None = None) -> int:
        filters = [RoleModel.company_id == company_id]
        if search:
            term = f"%{search.strip().lower()}%"
            filters.append(
                or_(
                    func.lower(RoleModel.name).like(term),
                    func.lower(RoleModel.slug).like(term),
                )
            )
        result = await self._session.execute(
            select(func.count()).select_from(RoleModel).where(*filters)
        )
        return int(result.scalar_one() or 0)

    async def update_role(self, role_id: str, role: Role) -> Role | None:
        result = await self._session.execute(
            select(RoleModel)
            .options(selectinload(RoleModel.permissions))
            .where(RoleModel.id == role_id, RoleModel.company_id == role.company_id)
        )
        row = result.scalar_one_or_none()
        if not row:
            return None
        row.name = role.name
        row.slug = role.slug
        row.description = role.description
        row.updated_at = role.updated_at
        row.permissions.clear()
        for p in role.permissions:
            row.permissions.append(
                RolePermissionModel(
                    id=new_id(),
                    role_id=row.id,
                    module=p.module if isinstance(p.module, str) else p.module.value,
                    can_view=p.can_view,
                    can_add=p.can_add,
                    can_edit=p.can_edit,
                    can_delete=p.can_delete,
                    can_export=p.can_export,
                )
            )
        await self._session.commit()
        return await self.get_role(role_id, role.company_id)

    async def delete_role(self, role_id: str, company_id: str) -> bool:
        row = await self._session.get(RoleModel, role_id)
        if not row or row.company_id != company_id:
            return False
        await self._session.delete(row)
        await self._session.commit()
        return True

    async def count_users_for_role(self, role_id: str, company_id: str) -> int:
        result = await self._session.execute(
            select(func.count())
            .select_from(CompanyUserModel)
            .where(
                CompanyUserModel.role_id == role_id,
                CompanyUserModel.company_id == company_id,
            )
        )
        return int(result.scalar_one() or 0)

    async def get_summary(self, company_id: str) -> UsersRolesSummary:
        total_users = await self.count_users(company_id)
        active_users = await self.count_users(
            company_id, status=CompanyUserStatus.ACTIVE.value
        )
        disabled_users = await self.count_users(
            company_id, status=CompanyUserStatus.DISABLED.value
        )
        total_roles = await self.count_roles(company_id)
        return UsersRolesSummary(
            total_users=total_users,
            active_users=active_users,
            disabled_users=disabled_users,
            total_roles=total_roles,
        )

    def _user_filters(
        self,
        company_id: str,
        *,
        role_id: str | None = None,
        department_id: str | None = None,
        status: str | None = None,
        search: str | None = None,
    ) -> list:
        filters = [CompanyUserModel.company_id == company_id]
        if role_id:
            filters.append(CompanyUserModel.role_id == role_id)
        if department_id:
            filters.append(CompanyUserModel.department_id == department_id)
        if status:
            normalized = status.strip().lower()
            if normalized == "inactive":
                normalized = CompanyUserStatus.DISABLED.value
            filters.append(CompanyUserModel.status == normalized)
        if search:
            term = f"%{search.strip().lower()}%"
            filters.append(
                or_(
                    func.lower(CompanyUserModel.full_name).like(term),
                    func.lower(CompanyUserModel.username).like(term),
                    func.lower(CompanyUserModel.email).like(term),
                    func.lower(func.coalesce(CompanyUserModel.employee_code, "")).like(
                        term
                    ),
                    func.lower(CompanyUserModel.id).like(term),
                )
            )
        return filters

    async def create_user(self, user: CompanyUser) -> CompanyUser:
        row = CompanyUserModel(
            id=user.id or new_id(),
            company_id=user.company_id,
            full_name=user.full_name,
            username=user.username,
            email=user.email,
            phone=user.phone,
            employee_code=user.employee_code,
            password_hash=user.password_hash or "",
            role_id=user.role_id,
            department_id=user.department_id,
            status=user.status.value,
            send_welcome_email=user.send_welcome_email,
            address=user.address,
            notes=user.notes,
            last_login=user.last_login,
            created_by=user.created_by,
            created_at=user.created_at,
            updated_at=user.updated_at,
        )
        self._session.add(row)
        await self._session.commit()
        return await self.get_user(row.id, user.company_id)  # type: ignore[return-value]

    async def get_user(self, user_id: str, company_id: str) -> CompanyUser | None:
        result = await self._session.execute(
            self._user_lookup_query().where(
                CompanyUserModel.id == user_id,
                CompanyUserModel.company_id == company_id,
            )
        )
        row = result.first()
        return self._user_from_join(row) if row else None

    async def get_user_global(self, user_id: str) -> CompanyUser | None:
        result = await self._session.execute(
            self._user_lookup_query().where(CompanyUserModel.id == user_id)
        )
        row = result.first()
        return self._user_from_join(row, include_secret=True) if row else None

    async def find_by_login(self, identifier: str) -> CompanyUser | None:
        ident = (identifier or "").strip().lower()
        if not ident:
            return None
        result = await self._session.execute(
            self._user_lookup_query()
            .where(
                or_(
                    func.lower(CompanyUserModel.username) == ident,
                    func.lower(CompanyUserModel.email) == ident,
                )
            )
            .order_by(
                case((CompanyUserModel.status == CompanyUserStatus.ACTIVE.value, 0), else_=1),
                CompanyUserModel.updated_at.desc(),
            )
            .limit(1)
        )
        row = result.first()
        return self._user_from_join(row, include_secret=True) if row else None

    async def update_last_login(self, user_id: str) -> None:
        row = await self._session.get(CompanyUserModel, user_id)
        if not row:
            return
        row.last_login = datetime.utcnow()
        await self._session.commit()

    async def get_user_by_username(
        self, company_id: str, username: str
    ) -> CompanyUser | None:
        result = await self._session.execute(
            select(CompanyUserModel).where(
                CompanyUserModel.company_id == company_id,
                CompanyUserModel.username == username,
            )
        )
        row = result.scalar_one_or_none()
        return self._to_user(row) if row else None

    async def get_user_by_email(self, company_id: str, email: str) -> CompanyUser | None:
        result = await self._session.execute(
            select(CompanyUserModel).where(
                CompanyUserModel.company_id == company_id,
                CompanyUserModel.email == email,
            )
        )
        row = result.scalar_one_or_none()
        return self._to_user(row) if row else None

    async def get_user_by_employee_code(
        self, company_id: str, employee_code: str
    ) -> CompanyUser | None:
        result = await self._session.execute(
            select(CompanyUserModel).where(
                CompanyUserModel.company_id == company_id,
                CompanyUserModel.employee_code == employee_code,
            )
        )
        row = result.scalar_one_or_none()
        return self._to_user(row) if row else None

    async def list_users(
        self,
        company_id: str,
        *,
        role_id: str | None = None,
        department_id: str | None = None,
        status: str | None = None,
        search: str | None = None,
        sort_by: str = "full_name",
        sort_dir: str = "asc",
        skip: int = 0,
        limit: int = 100,
    ) -> list[CompanyUser]:
        filters = self._user_filters(
            company_id,
            role_id=role_id,
            department_id=department_id,
            status=status,
            search=search,
        )
        column = _USER_SORT_COLUMNS.get(sort_by, CompanyUserModel.full_name)
        order = column.desc() if sort_dir.lower() == "desc" else column.asc()
        result = await self._session.execute(
            select(
                CompanyUserModel,
                RoleModel.name.label("role_name"),
                RoleModel.slug.label("role_slug"),
                DepartmentModel.code.label("department_code"),
                DepartmentModel.name.label("department_name"),
            )
            .join(RoleModel, RoleModel.id == CompanyUserModel.role_id)
            .outerjoin(
                DepartmentModel, DepartmentModel.id == CompanyUserModel.department_id
            )
            .where(*filters)
            .order_by(order, CompanyUserModel.id.asc())
            .offset(skip)
            .limit(limit)
        )
        return [
            self._to_user(
                row.CompanyUserModel,
                role_name=row.role_name,
                role_slug=row.role_slug,
                department_code=row.department_code,
                department_name=row.department_name,
            )
            for row in result.all()
        ]

    async def count_users(
        self,
        company_id: str,
        *,
        role_id: str | None = None,
        department_id: str | None = None,
        status: str | None = None,
        search: str | None = None,
    ) -> int:
        filters = self._user_filters(
            company_id,
            role_id=role_id,
            department_id=department_id,
            status=status,
            search=search,
        )
        result = await self._session.execute(
            select(func.count()).select_from(CompanyUserModel).where(*filters)
        )
        return int(result.scalar_one() or 0)

    async def update_user(self, user_id: str, user: CompanyUser) -> CompanyUser | None:
        row = await self._session.get(CompanyUserModel, user_id)
        if not row or row.company_id != user.company_id:
            return None
        row.full_name = user.full_name
        row.username = user.username
        row.email = str(user.email)
        row.phone = user.phone
        row.employee_code = user.employee_code
        if user.password_hash:
            row.password_hash = user.password_hash
        row.role_id = user.role_id
        row.department_id = user.department_id
        row.status = user.status.value
        row.send_welcome_email = user.send_welcome_email
        row.address = user.address
        row.notes = user.notes
        row.last_login = user.last_login
        row.updated_at = user.updated_at
        await self._session.commit()
        return await self.get_user(user_id, user.company_id)

    async def delete_user(self, user_id: str, company_id: str) -> bool:
        row = await self._session.get(CompanyUserModel, user_id)
        if not row or row.company_id != company_id:
            return False
        await self._session.delete(row)
        await self._session.commit()
        return True
