from fastapi import Depends, HTTPException, status
from fastapi.security import OAuth2PasswordBearer
from jose import JWTError, jwt
from sqlmodel import Session, select
from app.core import security, config
from app.core.db import get_session
from app.models.models import User
from app.models.models import UserSession

oauth2_scheme = OAuth2PasswordBearer(tokenUrl="token")

def get_current_user(token: str = Depends(oauth2_scheme), session: Session = Depends(get_session)):
    credentials_exception = HTTPException(
        status_code=status.HTTP_401_UNAUTHORIZED,
        detail="Could not validate credentials",
        headers={"WWW-Authenticate": "Bearer"},
    )
    try:
        payload = jwt.decode(token, config.settings.SECRET_KEY, algorithms=[config.settings.ALGORITHM])
        username: str = payload.get("sub")
        if username is None:
            raise credentials_exception
    except JWTError:
        raise credentials_exception

    # Verify session/jti hasn't been revoked
    jti = payload.get("jti")
    if not jti:
        raise credentials_exception

    # Check the session record exists and is not revoked and not expired
    session_rec = session.exec(select(UserSession).where(UserSession.jti == jti, UserSession.revoked == False)).one_or_none()
    if not session_rec:
        raise credentials_exception

    # Check expiry
    from datetime import datetime as _dt
    if session_rec.expires_at and session_rec.expires_at < _dt.utcnow():
        raise credentials_exception

    user = session.exec(select(User).where(User.username == username)).one_or_none()

    if user is None:
        raise credentials_exception
    return user

def get_current_active_user(current_user: User = Depends(get_current_user)):
    if not current_user.is_active:
        raise HTTPException(status_code=400, detail="Inactive user")
    return current_user


def get_current_user_optional(token: str = Depends(oauth2_scheme), session: Session = Depends(get_session)):
    """Attempt to resolve the user from token; return None when token is missing/invalid."""
    try:
        payload = jwt.decode(token, config.settings.SECRET_KEY, algorithms=[config.settings.ALGORITHM])
        username: str = payload.get("sub")
        if username is None:
            return None
    except JWTError:
        return None

    # Apply the same jti/session revocation check as get_current_user
    jti = payload.get("jti")
    if jti:
        from datetime import datetime as _dt
        session_rec = session.exec(
            select(UserSession).where(UserSession.jti == jti, UserSession.revoked == False)
        ).one_or_none()
        if not session_rec or (session_rec.expires_at and session_rec.expires_at < _dt.utcnow()):
            return None

    return session.exec(select(User).where(User.username == username)).one_or_none()

