from datetime import datetime

from sqlalchemy import func, select
from sqlalchemy.ext.asyncio import AsyncSession

from app.domain.entities.user_registration import UserRegistration
from app.domain.repositories.user_registration_repository import UserRegistrationRepository
from app.infrastructure.db.models import UserRegistrationModel, new_id


class MySQLUserRegistrationRepository(UserRegistrationRepository):
    def __init__(self, session: AsyncSession) -> None:
        self._session = session

    def _to_entity(self, row: UserRegistrationModel) -> UserRegistration:
        return UserRegistration(
            id=row.id,
            username=row.username,
            email=row.email,
            full_name=row.full_name,
            password_hash=row.password_hash,
            is_active=row.is_active,
            created_at=row.created_at,
            updated_at=row.updated_at,
        )

    async def create(self, user: UserRegistration) -> UserRegistration:
        row = UserRegistrationModel(
            id=user.id or new_id(),
            username=user.username.lower(),
            email=str(user.email).lower(),
            full_name=user.full_name,
            password_hash=user.password_hash,
            is_active=user.is_active,
            created_at=user.created_at,
            updated_at=user.updated_at,
        )
        self._session.add(row)
        await self._session.commit()
        await self._session.refresh(row)
        return self._to_entity(row)

    async def get_by_id(self, user_id: str) -> UserRegistration | None:
        row = await self._session.get(UserRegistrationModel, user_id)
        return self._to_entity(row) if row else None

    async def get_by_email(self, email: str) -> UserRegistration | None:
        result = await self._session.execute(
            select(UserRegistrationModel).where(UserRegistrationModel.email == email.lower())
        )
        row = result.scalar_one_or_none()
        return self._to_entity(row) if row else None

    async def get_by_username(self, username: str) -> UserRegistration | None:
        result = await self._session.execute(
            select(UserRegistrationModel).where(UserRegistrationModel.username == username.lower())
        )
        row = result.scalar_one_or_none()
        return self._to_entity(row) if row else None

    async def list_all(
        self,
        is_active: bool | None = None,
        skip: int = 0,
        limit: int = 100,
    ) -> list[UserRegistration]:
        query = select(UserRegistrationModel)
        if is_active is not None:
            query = query.where(UserRegistrationModel.is_active == is_active)
        result = await self._session.execute(
            query.order_by(UserRegistrationModel.created_at.desc()).offset(skip).limit(limit)
        )
        return [self._to_entity(row) for row in result.scalars().all()]

    async def count(self, is_active: bool | None = None) -> int:
        query = select(func.count()).select_from(UserRegistrationModel)
        if is_active is not None:
            query = query.where(UserRegistrationModel.is_active == is_active)
        result = await self._session.execute(query)
        return int(result.scalar_one())
