"""
SQLAlchemy 2.0 ORM models — FastAPI Data Gateway
=================================================

Uses the modern Mapped[T] + mapped_column() declarative style required by
SQLAlchemy 2.0's DeclarativeBase. Bare Python type annotations (e.g.
  id: uuid.UUID = Column(...)
) are rejected by DeclarativeBase — every mapped attribute must be wrapped
in Mapped[T].

Tables
------
  users            Human user accounts (Caller Type 1)
  refresh_tokens   JWT refresh-token family tree for token rotation
  api_keys         External subscriber keys + internal cron keys (Types 2 & 3)
  auth_audit_log   Immutable security event journal

Security decisions (inline)
---------------------------
- Passwords/secrets NEVER stored in plaintext — bcrypt/Fernet only.
- UUID primary keys via gen_random_uuid() — non-sequential, prevents
  row-count enumeration attacks.
- TIMESTAMPTZ for all timestamps — avoids DST bugs, enforces UTC.
- inet type for IP addresses — allows CIDR subnet queries.
- TEXT[] + GIN index on api_keys.scopes — O(1) scope containment check.
- CASCADE DELETE on refresh_tokens → users — atomic revocation on delete.
- auth_audit_log has no FK cascade — rows survive account deletion.
"""

from __future__ import annotations

import uuid
from datetime import datetime
from typing import Optional

from sqlalchemy import (
    BigInteger,
    Boolean,
    DateTime,
    ForeignKey,
    Index,
    Integer,
    String,
    Text,
    text,
)
from sqlalchemy.dialects.postgresql import ARRAY, INET, JSONB, UUID
from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column, relationship
from sqlalchemy.sql import func


class Base(DeclarativeBase):
    """Shared declarative base — imported by alembic/env.py for autogenerate."""
    pass


# ---------------------------------------------------------------------------
# USERS  (Caller Type 1 — human browser / mobile)
# ---------------------------------------------------------------------------

class User(Base):
    __tablename__ = "users"

    # UUIDv4 primary key — non-sequential, prevents row-count enumeration
    id: Mapped[uuid.UUID] = mapped_column(
        UUID(as_uuid=True),
        primary_key=True,
        server_default=func.gen_random_uuid(),
        comment="UUIDv4 — non-sequential, prevents row-count enumeration",
    )

    email: Mapped[str] = mapped_column(
        Text,
        nullable=False,
        unique=True,
        index=True,
        comment="Stored lowercased + trimmed; unique constraint enforced here",
    )

    # bcrypt cost ≥ 12 enforced at the application layer (passlib CryptContext)
    password_hash: Mapped[str] = mapped_column(
        Text,
        nullable=False,
        comment="bcrypt hash, cost>=12. Never store plaintext.",
    )

    role: Mapped[str] = mapped_column(
        String(16),
        nullable=False,
        default="viewer",
        server_default=text("'viewer'"),
        comment="RBAC role: viewer | analyst | admin",
    )

    is_active: Mapped[bool] = mapped_column(
        Boolean,
        nullable=False,
        default=True,
        server_default=text("true"),
        comment="Soft-disable without deleting the account",
    )

    # Brute-force lockout: track consecutive login failures
    failed_login_attempts: Mapped[int] = mapped_column(
        Integer,
        nullable=False,
        default=0,
        server_default=text("0"),
        comment="Reset to 0 on successful login",
    )

    # NULL = not locked; set to a future timestamp after N failures
    locked_until: Mapped[Optional[datetime]] = mapped_column(
        DateTime(timezone=True),
        nullable=True,
        comment="Account locked until this UTC timestamp. NULL = not locked.",
    )

    last_login_at: Mapped[Optional[datetime]] = mapped_column(
        DateTime(timezone=True),
        nullable=True,
    )

    created_at: Mapped[datetime] = mapped_column(
        DateTime(timezone=True),
        nullable=False,
        server_default=func.now(),
    )

    updated_at: Mapped[datetime] = mapped_column(
        DateTime(timezone=True),
        nullable=False,
        server_default=func.now(),
        # DB trigger trg_users_updated_at handles this at the DB level;
        # onupdate here is the Python-side fallback.
        onupdate=func.now(),
    )

    # Relationship — cascade DELETE at ORM level (DB-level ON DELETE CASCADE
    # is also defined in the migration for belt-and-suspenders safety).
    # lazy="selectin" is safe for async sessions — avoids implicit IO errors.
    refresh_tokens: Mapped[list["RefreshToken"]] = relationship(
        "RefreshToken",
        back_populates="user",
        cascade="all, delete-orphan",
        lazy="selectin",
    )

    __table_args__ = (
        # Partial index: only active users — login lookup skips inactive rows
        Index(
            "ix_users_email_active",
            "email",
            postgresql_where=text("is_active = true"),
        ),
        Index("ix_users_role", "role"),
    )

    def __repr__(self) -> str:
        return f"<User id={self.id} email={self.email} role={self.role}>"


# ---------------------------------------------------------------------------
# REFRESH TOKENS  (JWT rotation + token theft detection)
# ---------------------------------------------------------------------------

class RefreshToken(Base):
    __tablename__ = "refresh_tokens"

    # jti is the JWT ID claim — UUID embedded inside the signed refresh JWT
    jti: Mapped[uuid.UUID] = mapped_column(
        UUID(as_uuid=True),
        primary_key=True,
        server_default=func.gen_random_uuid(),
        comment="JWT ID claim. Embedded in the signed refresh token.",
    )

    user_id: Mapped[uuid.UUID] = mapped_column(
        UUID(as_uuid=True),
        ForeignKey("users.id", ondelete="CASCADE"),
        nullable=False,
        index=True,
    )

    # family_id links the entire rotation chain from one login session.
    # On detecting a reused (revoked) jti, revoke ALL rows sharing this family_id.
    family_id: Mapped[uuid.UUID] = mapped_column(
        UUID(as_uuid=True),
        nullable=False,
        index=True,
        comment="Shared across rotation chain. Entire family revoked on theft.",
    )

    revoked: Mapped[bool] = mapped_column(
        Boolean,
        nullable=False,
        default=False,
        server_default=text("false"),
        comment="Set true on: logout, rotation (old token), theft detection",
    )

    # Device fingerprinting — mismatch between issuance and presentation
    # IP/UA is logged but NOT used to auto-reject (avoids blocking mobile users
    # whose IP changes; use anomaly scoring instead).
    user_agent: Mapped[Optional[str]] = mapped_column(
        Text,
        nullable=True,
        comment="User-agent at issuance; used for anomaly detection logging",
    )

    ip_address: Mapped[Optional[str]] = mapped_column(
        INET,
        nullable=True,
        comment="Client IP at issuance. inet type supports CIDR subnet queries.",
    )

    expires_at: Mapped[datetime] = mapped_column(
        DateTime(timezone=True),
        nullable=False,
        comment="Hard expiry = issued_at + REFRESH_TOKEN_TTL (default 7 days)",
    )

    created_at: Mapped[datetime] = mapped_column(
        DateTime(timezone=True),
        nullable=False,
        server_default=func.now(),
    )

    user: Mapped["User"] = relationship(
        "User",
        back_populates="refresh_tokens",
        lazy="selectin",
    )

    __table_args__ = (
        # Nightly cleanup: DELETE WHERE expires_at < NOW() - INTERVAL '1 day'
        Index("ix_refresh_tokens_expires_at", "expires_at"),
        # Theft detection: WHERE family_id=? AND revoked=false
        Index("ix_refresh_tokens_family_active", "family_id", "revoked"),
        # Validation hot path: WHERE jti=? AND revoked=false
        Index("ix_refresh_tokens_jti_active", "jti", "revoked"),
    )

    def __repr__(self) -> str:
        return (
            f"<RefreshToken jti={self.jti} user_id={self.user_id} "
            f"revoked={self.revoked}>"
        )


# ---------------------------------------------------------------------------
# API KEYS  (Caller Type 2 — HMAC-signed external  +  Type 3 — internal cron)
# ---------------------------------------------------------------------------

class ApiKey(Base):
    __tablename__ = "api_keys"

    key_id: Mapped[uuid.UUID] = mapped_column(
        UUID(as_uuid=True),
        primary_key=True,
        server_default=func.gen_random_uuid(),
        comment="Sent as X-Service-Key header. Safe to expose publicly.",
    )

    # Storage strategy depends on caller type (see app/core/security.py):
    #   "fernet:..."  → Fernet-encrypted secret for HMAC keys (Type 2)
    #   "$2b$..."     → bcrypt hash for cron keys (Type 3)
    secret_hash: Mapped[str] = mapped_column(
        Text,
        nullable=False,
        comment="Fernet-encrypted (HMAC keys) or bcrypt hash (cron keys).",
    )

    service_name: Mapped[str] = mapped_column(
        Text,
        nullable=False,
        comment="Human label: analytics-svc, cron-collector, etc.",
    )

    # TEXT[] with GIN index — enables O(1) @> containment queries:
    #   WHERE scopes @> ARRAY['data:pull']
    scopes: Mapped[list[str]] = mapped_column(
        ARRAY(Text),
        nullable=False,
        server_default=text("'{}'"),
        comment="Permission list. GIN-indexed for @> containment queries.",
    )

    is_active: Mapped[bool] = mapped_column(
        Boolean,
        nullable=False,
        default=True,
        server_default=text("true"),
        comment="False = immediate revocation. Use 24h grace on rotation.",
    )

    # Per-key rate ceiling stored here; Redis enforces it at runtime.
    rate_limit: Mapped[int] = mapped_column(
        Integer,
        nullable=False,
        default=600,
        server_default=text("600"),
        comment="Max requests/minute for this key. Enforced via Redis.",
    )

    # NULL = never expires; set for time-bounded access grants
    expires_at: Mapped[Optional[datetime]] = mapped_column(
        DateTime(timezone=True),
        nullable=True,
        comment="NULL = never expires. Set for time-bounded access grants.",
    )

    # SET NULL on admin account deletion — key record survives
    created_by: Mapped[Optional[uuid.UUID]] = mapped_column(
        UUID(as_uuid=True),
        ForeignKey("users.id", ondelete="SET NULL"),
        nullable=True,
        comment="Admin who issued the key. NULL after issuer account deleted.",
    )

    description: Mapped[Optional[str]] = mapped_column(
        Text,
        nullable=True,
        comment="Ops-team notes: purpose, owner, rotation schedule",
    )

    created_at: Mapped[datetime] = mapped_column(
        DateTime(timezone=True),
        nullable=False,
        server_default=func.now(),
    )

    # Written asynchronously (fire-and-forget) — never blocks the hot path
    last_used_at: Mapped[Optional[datetime]] = mapped_column(
        DateTime(timezone=True),
        nullable=True,
        comment="Updated async on each request. DO NOT block hot path for this.",
    )

    __table_args__ = (
        # GIN index on scopes — enables fast @> containment operator
        Index(
            "ix_api_keys_scopes_gin",
            "scopes",
            postgresql_using="gin",
        ),
        # Partial index — validation query filters is_active=true only
        Index(
            "ix_api_keys_active",
            "is_active",
            postgresql_where=text("is_active = true"),
        ),
    )

    def __repr__(self) -> str:
        return (
            f"<ApiKey key_id={self.key_id} service={self.service_name} "
            f"active={self.is_active}>"
        )


# ---------------------------------------------------------------------------
# AUTH AUDIT LOG  (immutable security event journal)
# ---------------------------------------------------------------------------

class AuthAuditLog(Base):
    """
    Append-only security log.  Application role has INSERT-only permission.
    Never UPDATE or DELETE rows here.

    Used for:
      - Login failure rate analysis (brute-force detection)
      - Token theft forensics
      - HMAC replay attack monitoring
      - Compliance audit trails

    Partition by month in production once volume > ~10M rows/month:
      PARTITION BY RANGE (created_at)
    """
    __tablename__ = "auth_audit_log"

    # BIGSERIAL — sequential for fast append and time-ordered range scans
    id: Mapped[int] = mapped_column(
        BigInteger,
        primary_key=True,
        autoincrement=True,
        comment="Sequential BIGINT — fast append and time-range scans",
    )

    event_type: Mapped[str] = mapped_column(
        String(64),
        nullable=False,
        index=True,
        comment=(
            "login_success | login_failure | token_theft | token_expired | "
            "logout | key_revoked | replay_attack | invalid_signature | "
            "scope_denied | account_locked | job_run"
        ),
    )

    caller_type: Mapped[Optional[str]] = mapped_column(
        String(16),
        nullable=True,
        comment="user | api_key | cron",
    )

    # No FK constraints — log rows must survive account/key deletion
    user_id: Mapped[Optional[uuid.UUID]] = mapped_column(
        UUID(as_uuid=True),
        nullable=True,
        index=True,
        comment="No FK constraint — log survives user deletion",
    )

    key_id: Mapped[Optional[uuid.UUID]] = mapped_column(
        UUID(as_uuid=True),
        nullable=True,
        index=True,
        comment="No FK constraint — log survives key deletion",
    )

    # inet type — supports CIDR subnet queries: ip_address << '10.0.0.0/8'
    ip_address: Mapped[Optional[str]] = mapped_column(
        INET,
        nullable=True,
        index=True,
    )

    user_agent: Mapped[Optional[str]] = mapped_column(
        Text,
        nullable=True,
        comment="Truncated to 512 chars in application layer",
    )

    # Flexible JSONB blob — event-specific context.
    # NEVER log passwords, tokens, or secrets in this column.
    # Note: Python attribute is 'event_data'; DB column stays 'metadata'.
    # ('metadata' is reserved by SQLAlchemy's DeclarativeBase.)
    event_data: Mapped[dict] = mapped_column(
        "metadata",   # explicit DB column name
        JSONB,
        nullable=False,
        server_default=text("'{}'::jsonb"),
        comment="Event detail blob. Never log secrets here.",
    )

    created_at: Mapped[datetime] = mapped_column(
        DateTime(timezone=True),
        nullable=False,
        server_default=func.now(),
        index=True,
    )

    __table_args__ = (
        # Brute-force detection: COUNT failures per IP in a time window
        Index("ix_audit_ip_event_time", "ip_address", "event_type", "created_at"),
        # User history: list events for a specific account
        Index("ix_audit_user_time", "user_id", "created_at"),
        # Event stream by type in time order (alerting system)
        Index("ix_audit_event_type_time", "event_type", "created_at"),
    )

    def __repr__(self) -> str:
        return (
            f"<AuthAuditLog id={self.id} event={self.event_type} "
            f"caller={self.caller_type} at={self.created_at}>"
        )


# ---------------------------------------------------------------------------
# SCHEDULER JOB LOG  (one row per scheduler invocation)
# ---------------------------------------------------------------------------

class SchedulerJobLog(Base):
    """
    Execution history for APScheduler jobs.
    Created by migration bfd7a23f9437.

    One row per invocation — records status, duration, error, and
    a JSON metadata blob with per-stage context from the service.
    """
    __tablename__ = "scheduler_job_log"

    id: Mapped[int] = mapped_column(
        BigInteger, primary_key=True, autoincrement=True
    )
    job_id: Mapped[str] = mapped_column(String(128), nullable=False)
    status: Mapped[str] = mapped_column(
        String(16), nullable=False,
        comment="success | error | skipped",
    )
    duration_ms: Mapped[Optional[int]] = mapped_column(
        Integer, nullable=True,
        comment="Wall-clock time in ms. NULL when skipped.",
    )
    error_msg: Mapped[Optional[str]] = mapped_column(
        Text, nullable=True,
        comment="Exception message / traceback on status=error",
    )
    job_metadata: Mapped[dict] = mapped_column(
        "job_metadata",
        JSONB,
        nullable=False,
        server_default=text("'{}'::jsonb"),
    )
    started_at: Mapped[datetime] = mapped_column(
        DateTime(timezone=True),
        nullable=False,
        server_default=text("NOW()"),
    )

    def __repr__(self) -> str:
        return (
            f"<SchedulerJobLog id={self.id} job={self.job_id!r} "
            f"status={self.status} at={self.started_at}>"
        )
