from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession

from app.domain.entities.refresh_token import RefreshToken
from app.domain.repositories.refresh_token_repository import RefreshTokenRepository
from app.infrastructure.db.models import RefreshTokenModel, new_id


class MySQLRefreshTokenRepository(RefreshTokenRepository):
    def __init__(self, session: AsyncSession) -> None:
        self._session = session

    def _to_entity(self, row: RefreshTokenModel) -> RefreshToken:
        return RefreshToken(
            id=row.id,
            user_id=row.user_id,
            jti=row.jti,
            expires_at=row.expires_at,
            is_revoked=row.is_revoked,
            created_at=row.created_at,
        )

    async def create(self, token: RefreshToken) -> RefreshToken:
        row = RefreshTokenModel(
            id=token.id or new_id(),
            user_id=token.user_id,
            jti=token.jti,
            expires_at=token.expires_at,
            is_revoked=token.is_revoked,
            created_at=token.created_at,
        )
        self._session.add(row)
        await self._session.commit()
        await self._session.refresh(row)
        return self._to_entity(row)

    async def get_by_jti(self, jti: str) -> RefreshToken | None:
        result = await self._session.execute(
            select(RefreshTokenModel).where(
                RefreshTokenModel.jti == jti,
                RefreshTokenModel.is_revoked.is_(False),
            )
        )
        row = result.scalar_one_or_none()
        return self._to_entity(row) if row else None

    async def revoke(self, jti: str) -> None:
        result = await self._session.execute(
            select(RefreshTokenModel).where(RefreshTokenModel.jti == jti)
        )
        row = result.scalar_one_or_none()
        if row:
            row.is_revoked = True
            await self._session.commit()
