from datetime import datetime
from sqlmodel import Session, select
from app.models.models import UserSession


def create_session(db: Session, user_id: int, jti: str, expires_at: datetime) -> UserSession:
    session = UserSession(user_id=user_id, jti=jti, issued_at=datetime.utcnow(), expires_at=expires_at)
    db.add(session)
    db.commit()
    db.refresh(session)
    return session


def revoke_session(db: Session, jti: str) -> None:
    statement = select(UserSession).where(UserSession.jti == jti)
    sess = db.exec(statement).one_or_none()
    if not sess:
        return
    sess.revoked = True
    db.add(sess)
    db.commit()


def revoke_all_sessions_for_user(db: Session, user_id: int) -> int:
    statement = select(UserSession).where(UserSession.user_id == user_id, UserSession.revoked == False)
    sessions = db.exec(statement).all()
    for s in sessions:
        s.revoked = True
        db.add(s)
    db.commit()
    return len(sessions)


def get_active_session_by_jti(db: Session, jti: str):
    statement = select(UserSession).where(UserSession.jti == jti, UserSession.revoked == False)
    return db.exec(statement).one_or_none()
