import jwt as pyjwt
import base64
import logging
from functools import wraps
from flask import request, jsonify, g

log = logging.getLogger("JWT_MIDDLEWARE")

# JJWT 0.9.x tronque le secret au multiple de 4 le plus proche
# "camelsoft3200" (13 chars) → "camelsoft320" (12 chars) + "=" padding
JWT_SECRET    = base64.b64decode("camelsoft320=")
JWT_ALGORITHM = "HS512"

log.info(f"[JWT] Secret résolu — {len(JWT_SECRET)} bytes")


def require_jwt(f):
    @wraps(f)
    def decorated(*args, **kwargs):
        auth_header = request.headers.get("Authorization", "")
        if not auth_header.startswith("Bearer "):
            return jsonify({"error": "Token manquant"}), 401

        token = auth_header.split(" ", 1)[1]

        try:
            payload = pyjwt.decode(token, JWT_SECRET, algorithms=[JWT_ALGORITHM])
        except pyjwt.ExpiredSignatureError:
            log.warning("[JWT] Token expiré")
            return jsonify({"error": "Token expiré — reconnectez-vous"}), 401
        except pyjwt.InvalidTokenError as e:
            log.error(f"[JWT] Token invalide : {e}")
            return jsonify({"error": f"Token invalide : {e}"}), 401

        tenant_schema = payload.get("tenantSchema")
        if not tenant_schema or tenant_schema in ("", "public"):
            return jsonify({"error": "Aucune entreprise assignée à ce compte"}), 403

        company_id = payload.get("companyId")
        if not company_id:
            return jsonify({"error": "companyId manquant dans le token"}), 403

        g.username      = payload.get("sub")
        g.tenant_schema = tenant_schema.strip()
        g.company_id    = company_id
        g.is_superadmin = payload.get("isSuperAdmin", False)

        log.info(f"[JWT] ✅ user={g.username} tenant={g.tenant_schema}")
        return f(*args, **kwargs)
    return decorated