"""
Unified authentication dependency — gateway_auth.

Routes all three caller types through a single FastAPI Depends() entry point:

  Authorization: ApiKey <key>    → Caller Type 3 (cron)
  X-Service-Key: <uuid>          → Caller Type 2 (HMAC external)
  Authorization: Bearer <jwt>    → Caller Type 1 (human user)

Security rules enforced here:
  • 401 (not 403) when identity cannot be determined — never reveal which auth
    method is expected or whether a key/user exists.
  • hmac.compare_digest() for all secret comparisons — never ==.
  • Timestamp window (±5 min) checked before any DB lookup (replay protection).
  • Auth failures always written to auth_audit_log with IP + caller_type.
  • last_used_at updated asynchronously — never blocks the hot path.

RBAC helpers
  require_role(*roles)          — user-only gate (role in JWT)
  require_scope(*scopes)        — api_key / cron gate (scope in api_keys row)
  require_user_or_scope(...)    — mixed gate (data endpoints used by both)
"""

from __future__ import annotations

import asyncio
import uuid
from datetime import datetime, timezone

import structlog
from fastapi import Depends, HTTPException, Request
from jose import JWTError
from sqlalchemy import select, update
from sqlalchemy.ext.asyncio import AsyncSession

from app.core.security import (
    _FERNET_PREFIX,
    decrypt_api_secret,
    decode_token,
    is_timestamp_fresh,
    verify_api_secret_bcrypt,
    verify_api_secret_plain,
    verify_hmac_signature,
)
from app.db.models import ApiKey, AuthAuditLog
from app.db.session import AsyncSessionLocal, get_db
from app.schemas.auth import (
    ApiKeyCallerContext,
    CallerContext,
    CronCallerContext,
    UserCallerContext,
)

log = structlog.get_logger(__name__)


# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------

def get_client_ip(request: Request) -> str:
    """
    Extract the real client IP.
    Trusts X-Forwarded-For only if the app sits behind a known reverse proxy.
    In production, configure your nginx/ALB to set this header correctly.
    """
    xff = request.headers.get("X-Forwarded-For")
    if xff:
        return xff.split(",")[0].strip()
    return request.client.host if request.client else "unknown"


async def _write_audit_event(
    db: AsyncSession,
    event_type: str,
    caller_type: str | None = None,
    user_id: str | None = None,
    key_id: str | None = None,
    ip_address: str | None = None,
    user_agent: str | None = None,
    metadata: dict | None = None,
) -> None:
    """
    Append an immutable row to auth_audit_log.
    Swallows all exceptions — logging must never crash the request.
    The session is committed by the parent get_db() context.
    """
    try:
        entry = AuthAuditLog(
            event_type=event_type,
            caller_type=caller_type,
            user_id=uuid.UUID(user_id) if user_id else None,
            key_id=uuid.UUID(key_id) if key_id else None,
            ip_address=ip_address,
            user_agent=(user_agent or "")[:512],  # guard against oversized UA
            event_data=metadata or {},             # DB column is 'metadata'
        )
        db.add(entry)
    except Exception:
        log.exception("audit_log_write_failed", event_type=event_type)


async def _update_key_last_used(key_id: uuid.UUID) -> None:
    """
    Fire-and-forget: update api_keys.last_used_at in a dedicated session.
    Runs as a background asyncio task — never delays the response.
    Uses a fresh session so it does not share state with the request session.
    """
    async with AsyncSessionLocal() as session:
        try:
            await session.execute(
                update(ApiKey)
                .where(ApiKey.key_id == key_id)
                .values(last_used_at=datetime.now(timezone.utc))
            )
            await session.commit()
        except Exception:
            log.exception("last_used_at_update_failed", key_id=str(key_id))


# ===========================================================================
# CALLER TYPE 1: Human user — JWT Bearer
# ===========================================================================

async def verify_user_jwt(token: str) -> UserCallerContext:
    """
    Verify JWT signature and expiry. No DB hit — pure cryptographic check.

    Trade-off: a revoked user retains access for up to ACCESS_TOKEN_TTL (15 min).
    Acceptable: revocation is handled at the refresh token level.
    For stricter revocation, add a Redis blocklist checked here.
    """
    try:
        payload = decode_token(token)
    except JWTError as exc:
        # Same error for expired vs tampered — do not distinguish to callers
        raise HTTPException(
            status_code=401,
            detail="Invalid or expired token",
            headers={"WWW-Authenticate": "Bearer"},
        ) from exc

    if payload.get("type") != "access":
        raise HTTPException(
            status_code=401,
            detail="Invalid token type",
            headers={"WWW-Authenticate": "Bearer"},
        )

    sub: str | None = payload.get("sub")
    role: str | None = payload.get("role")

    if not sub or not role:
        raise HTTPException(status_code=401, detail="Malformed token payload")

    return UserCallerContext(sub=sub, role=role)


# ===========================================================================
# CALLER TYPE 2: External subscriber — HMAC-signed
# ===========================================================================

async def verify_api_key(
    request: Request,
    db: AsyncSession,
) -> ApiKeyCallerContext:
    """
    Full HMAC-SHA256 verification for external API keys.

    Verification order (fail fast on cheapest operation first):
      1. Check all required headers present
      2. Timestamp window (no DB, no crypto)
      3. Parse + look up key in DB
      4. Decrypt secret + recompute HMAC
      5. Constant-time signature compare
      6. Scope check deferred to each route handler
    """
    ip = get_client_ip(request)
    ua = request.headers.get("User-Agent", "")

    key_id_str = request.headers.get("X-Service-Key", "")
    timestamp = request.headers.get("X-Timestamp", "")
    signature = request.headers.get("X-Signature", "")

    if not key_id_str or not timestamp or not signature:
        raise HTTPException(status_code=401, detail="Missing authentication headers")

    # Step 1: Timestamp window — cheapest check, no DB touch
    if not is_timestamp_fresh(timestamp):
        await _write_audit_event(
            db, "replay_attack", "api_key",
            ip_address=ip,
            metadata={"key_id": key_id_str, "timestamp": timestamp},
        )
        log.warning("replay_attack", ip=ip, key_id=key_id_str)
        raise HTTPException(status_code=401, detail="Request timestamp expired")

    # Step 2: Look up key
    try:
        key_id = uuid.UUID(key_id_str)
    except ValueError:
        raise HTTPException(status_code=401, detail="Invalid credentials")

    result = await db.execute(
        select(ApiKey).where(
            ApiKey.key_id == key_id,
            ApiKey.is_active == True,  # noqa: E712
        )
    )
    key = result.scalar_one_or_none()

    if not key:
        # 401 not 404 — do not reveal whether key exists
        await _write_audit_event(
            db, "invalid_signature", "api_key",
            ip_address=ip,
            metadata={"key_id": key_id_str, "reason": "key_not_found_or_inactive"},
        )
        raise HTTPException(status_code=401, detail="Invalid credentials")

    # Step 3: Recover plain secret (Fernet-encrypted for HMAC keys)
    plain_secret = decrypt_api_secret(key.secret_hash)
    if plain_secret is None:
        # This key uses bcrypt (cron key type) — HMAC not supported
        await _write_audit_event(
            db, "invalid_signature", "api_key",
            key_id=str(key.key_id),
            ip_address=ip,
            metadata={"reason": "wrong_key_type_for_hmac"},
        )
        raise HTTPException(status_code=401, detail="Invalid credentials")

    # Step 4: Recompute + compare (constant-time via hmac.compare_digest)
    body = await request.body()
    if not verify_hmac_signature(plain_secret, timestamp, body, signature):
        await _write_audit_event(
            db, "invalid_signature", "api_key",
            key_id=str(key.key_id),
            ip_address=ip,
            metadata={"reason": "signature_mismatch"},
        )
        log.warning("invalid_hmac_signature", ip=ip, key_id=str(key.key_id))
        raise HTTPException(status_code=401, detail="Invalid credentials")

    # Step 5: Scope validation deferred — each route calls require_scope()

    # Step 6: Update last_used_at — fire-and-forget background task
    asyncio.create_task(_update_key_last_used(key.key_id))

    return ApiKeyCallerContext(
        key_id=str(key.key_id),
        service_name=key.service_name,
        scopes=list(key.scopes or []),
        rate_limit=key.rate_limit,
    )


async def verify_api_key_bearer(
    request: Request,
    db: AsyncSession,
    bearer_secret: str,
) -> ApiKeyCallerContext:
    """
    Simpler S2S auth: X-Service-Key + Authorization: Bearer <secret>.

    Same as legacy market-data.php (static Bearer) but keyed per subscriber.
    No per-request HMAC signing — use full HMAC when you need replay protection.

    If X-Service-Key is omitted, scans active Fernet-backed keys (small fleets).
    """
    ip = get_client_ip(request)
    key_id_str = request.headers.get("X-Service-Key", "").strip()

    if key_id_str:
        try:
            key_id = uuid.UUID(key_id_str)
        except ValueError:
            raise HTTPException(status_code=401, detail="Invalid credentials")

        result = await db.execute(
            select(ApiKey).where(
                ApiKey.key_id == key_id,
                ApiKey.is_active == True,  # noqa: E712
            )
        )
        key = result.scalar_one_or_none()
        if not key:
            raise HTTPException(status_code=401, detail="Invalid credentials")

        plain = decrypt_api_secret(key.secret_hash)
        if plain is None or not verify_api_secret_plain(bearer_secret, plain):
            await _write_audit_event(
                db,
                "invalid_signature",
                "api_key",
                key_id=str(key.key_id),
                ip_address=ip,
                metadata={"reason": "bearer_secret_mismatch"},
            )
            raise HTTPException(status_code=401, detail="Invalid credentials")

        asyncio.create_task(_update_key_last_used(key.key_id))
        return ApiKeyCallerContext(
            key_id=str(key.key_id),
            service_name=key.service_name,
            scopes=list(key.scopes or []),
            rate_limit=key.rate_limit,
        )

    # Bearer secret only — find matching HMAC subscriber key
    result = await db.execute(
        select(ApiKey).where(
            ApiKey.is_active == True,  # noqa: E712
            ApiKey.secret_hash.like(f"{_FERNET_PREFIX}%"),
        )
    )
    for candidate in result.scalars().all():
        plain = decrypt_api_secret(candidate.secret_hash)
        if plain and verify_api_secret_plain(bearer_secret, plain):
            asyncio.create_task(_update_key_last_used(candidate.key_id))
            return ApiKeyCallerContext(
                key_id=str(candidate.key_id),
                service_name=candidate.service_name,
                scopes=list(candidate.scopes or []),
                rate_limit=candidate.rate_limit,
            )

    await _write_audit_event(
        db,
        "invalid_signature",
        "api_key",
        ip_address=ip,
        metadata={"reason": "bearer_secret_no_match"},
    )
    raise HTTPException(status_code=401, detail="Invalid credentials")


def _is_jwt_format(token: str) -> bool:
    """True when the token looks like a JWT (three base64url segments)."""
    parts = token.split(".")
    return len(parts) == 3 and all(parts)


# ===========================================================================
# CALLER TYPE 3: Internal cron job — static ApiKey header
# ===========================================================================

async def verify_cron_key(
    authorization: str,
    db: AsyncSession,
) -> CronCallerContext:
    """
    Verify a static internal API key (bcrypt, no HMAC signing).

    Simpler auth is justified because cron jobs run inside the private VPC.
    The network boundary is the primary control. If jobs ever run outside
    the VPC perimeter, upgrade to full HMAC (Caller Type 2 protocol).
    """
    raw_key = authorization[len("ApiKey "):]

    # Load all active internal:job keys (typically only 1–2 rows — fast)
    result = await db.execute(
        select(ApiKey).where(
            ApiKey.scopes.contains(["internal:job"]),  # type: ignore[arg-type]
            ApiKey.is_active == True,  # noqa: E712
        )
    )
    candidates = result.scalars().all()

    matched: ApiKey | None = None
    for candidate in candidates:
        # Only attempt bcrypt verify — HMAC keys store Fernet, not bcrypt
        if candidate.secret_hash.startswith("$2b$"):
            if verify_api_secret_bcrypt(raw_key, candidate.secret_hash):
                matched = candidate
                break

    if not matched:
        # 401 — identity could not be confirmed
        raise HTTPException(status_code=401, detail="Invalid job key")

    return CronCallerContext(
        key_id=str(matched.key_id),
        service_name=matched.service_name,
        scopes=list(matched.scopes or []),
    )


# ===========================================================================
# UNIFIED GATEWAY  (single Depends() entry point for all routes)
# ===========================================================================

async def gateway_auth(
    request: Request,
    db: AsyncSession = Depends(get_db),
) -> CallerContext:
    """
    Single FastAPI dependency that dispatches to the correct verifier.

    Routing order is significant:
      ApiKey prefix  → cron (most specific internal header)
      X-Service-Key  → HMAC external
      Bearer         → JWT user

    Always returns a CallerContext or raises HTTPException(401).
    Never returns 403 here — that is left to RBAC helpers below.
    """
    auth = request.headers.get("Authorization", "")

    if auth.startswith("ApiKey "):
        return await verify_cron_key(auth, db)

    if "X-Service-Key" in request.headers:
        if request.headers.get("X-Signature", "").strip():
            return await verify_api_key(request, db)
        if auth.startswith("Bearer "):
            return await verify_api_key_bearer(
                request, db, auth[len("Bearer "):].strip()
            )
        raise HTTPException(status_code=401, detail="Missing authentication headers")

    if auth.startswith("Bearer "):
        token = auth[len("Bearer "):].strip()
        if _is_jwt_format(token):
            return await verify_user_jwt(token)
        return await verify_api_key_bearer(request, db, token)

    # 401 not 403 — do not reveal which auth scheme is required
    raise HTTPException(
        status_code=401,
        detail="Authentication required",
        headers={"WWW-Authenticate": "Bearer"},
    )


# ===========================================================================
# RBAC composable dependencies
# ===========================================================================

def require_role(*roles: str):
    """
    Restrict a route to users with one of the specified roles.
    API key / cron callers are passed through (they use scope-based auth).

    Usage:
        @router.get("/admin/users")
        async def list_users(caller = Depends(require_role("admin"))):
    """
    async def _dep(
        caller: CallerContext = Depends(gateway_auth),
    ) -> CallerContext:
        if caller.caller_type == "user" and caller.role not in roles:
            raise HTTPException(
                status_code=403,
                detail=f"Role '{caller.role}' is not authorised for this resource",
            )
        return caller

    return _dep


def require_scope(*required_scopes: str):
    """
    Restrict API key / cron callers to those holding all required scopes.
    User (JWT) callers are passed through (they use role-based auth).

    Usage:
        @router.get("/data/pull/{dataset}")
        async def pull(caller = Depends(require_scope("data:pull"))):
    """
    async def _dep(
        caller: CallerContext = Depends(gateway_auth),
    ) -> CallerContext:
        if caller.caller_type in ("api_key", "cron"):
            for scope in required_scopes:
                if scope not in caller.scopes:
                    raise HTTPException(
                        status_code=403,
                        detail=f"Missing required scope: '{scope}'",
                    )
        return caller

    return _dep


def require_user_or_scope(roles: tuple[str, ...], scope: str):
    """
    Mixed gate for endpoints accessible by both human users and API keys.

    Examples:
      GET /data/pull/{dataset}  — user with viewer/analyst/admin OR api_key with data:pull
      WS  /data/push/subscribe  — user with analyst/admin OR api_key with feed:subscribe

    Cron callers are explicitly rejected from public data endpoints.
    """
    async def _dep(
        caller: CallerContext = Depends(gateway_auth),
    ) -> CallerContext:
        if caller.caller_type == "user":
            if caller.role not in roles:
                raise HTTPException(
                    status_code=403,
                    detail=f"Role '{caller.role}' cannot access this resource",
                )
        elif caller.caller_type == "api_key":
            if scope not in caller.scopes:
                raise HTTPException(
                    status_code=403,
                    detail=f"Missing required scope: '{scope}'",
                )
        else:
            # Cron callers must not reach public data endpoints
            raise HTTPException(
                status_code=403,
                detail="Internal job keys cannot access public data endpoints",
            )
        return caller

    return _dep
