from datetime import datetime
from sqlmodel import Session, select
from app.models.models import RefreshToken


def create_refresh_token(db: Session, user_id: int, token_hash: str, expires_at: datetime, session_jti: str = None) -> RefreshToken:
    rt = RefreshToken(
        user_id=user_id,
        token_hash=token_hash,
        session_jti=session_jti,
        issued_at=datetime.utcnow(),
        expires_at=expires_at,
        revoked=False,
    )
    db.add(rt)
    db.commit()
    db.refresh(rt)
    return rt


def get_active_by_hash(db: Session, token_hash: str):
    stmt = select(RefreshToken).where(RefreshToken.token_hash == token_hash, RefreshToken.revoked == False)
    return db.exec(stmt).one_or_none()


def revoke_refresh_token(db: Session, token_id: int):
    stmt = select(RefreshToken).where(RefreshToken.id == token_id)
    rt = db.exec(stmt).one_or_none()
    if not rt:
        return
    rt.revoked = True
    db.add(rt)
    db.commit()


def revoke_all_for_session(db: Session, session_jti: str) -> int:
    stmt = select(RefreshToken).where(RefreshToken.session_jti == session_jti, RefreshToken.revoked == False)
    rts = db.exec(stmt).all()
    for rt in rts:
        rt.revoked = True
        db.add(rt)
    db.commit()
    return len(rts)


def revoke_all_for_user(db: Session, user_id: int) -> int:
    stmt = select(RefreshToken).where(RefreshToken.user_id == user_id, RefreshToken.revoked == False)
    rts = db.exec(stmt).all()
    for rt in rts:
        rt.revoked = True
        db.add(rt)
    db.commit()
    return len(rts)


def rotate_refresh_token(db: Session, old_token_id: int, new_token_hash: str, new_expires_at: datetime):
    # mark old revoked and create replaced_by link after creating new token
    old = db.get(RefreshToken, old_token_id)
    if not old:
        return None
    old.revoked = True
    db.add(old)
    db.commit()
    return
