"""Privacy-bounded administrator audit search and export."""

from __future__ import annotations

import base64
from datetime import datetime, timezone
import hashlib
from typing import Mapping, Optional

from sqlalchemy import and_, func, or_, select
from sqlalchemy.orm import Session

from licensing_shared.canonical_json import canonicalize_json, parse_bounded_json
from licensing_shared.constants import validate_identifier

from .errors import ServerErrorCode, ServerLicensingError
from .models import AuditEvent, User
from .security import new_identifier


AUDIT_EXPORT_SCHEMA = "apolon.licensing.audit-export"
AUDIT_EXPORT_SCHEMA_VERSION = 1
MAXIMUM_AUDIT_SEARCH_EVENTS = 500
MAXIMUM_AUDIT_EXPORT_EVENTS = 5000
MAXIMUM_AUDIT_CURSOR_BYTES = 1024
AUDIT_METADATA_ALLOWLIST = frozenset(
    (
        "activationId",
        "activeSerialCount",
        "alreadyQueued",
        "assignedUserId",
        "assignmentGeneration",
        "billingHold",
        "billingPortalAvailable",
        "cancelAtPeriodEnd",
        "catalogRevision",
        "catalogSha256",
        "checked",
        "checkoutAvailable",
        "deviceId",
        "deliveredCount",
        "deliveryCount",
        "errorCode",
        "expiredIdempotencyDeleted",
        "extensionDays",
        "extensionGrantId",
        "extensionNumber",
        "entitlementCount",
        "failed",
        "failedOutboxDeleted",
        "failedRetentionDays",
        "invitationId",
        "licenseId",
        "membershipId",
        "notificationKind",
        "offlineCertificate",
        "organizationId",
        "preservedEntitlements",
        "previousCatalogRevision",
        "privacyRequestId",
        "privacyRequestStatus",
        "processedOutboxDeleted",
        "processedRetentionDays",
        "provider",
        "providerEventType",
        "quantity",
        "queued",
        "reasonCode",
        "redactedDeviceCount",
        "recipientCount",
        "redeemedLicenseCount",
        "redeemedSerialCount",
        "releasedActivationCount",
        "releasedSeatCount",
        "remainingAutomaticReleases",
        "removedEntitlements",
        "removedMembershipCount",
        "requestType",
        "retainedAuditEventCount",
        "retentionDays",
        "revocationGeneration",
        "revokedSerialCount",
        "role",
        "seatCount",
        "seatId",
        "serialLookupId",
        "skuId",
        "skuCount",
        "snapshotId",
        "stateDigest",
        "skippedCount",
        "subscriptionState",
        "unchanged",
        "ownedLicenseCount",
        "purgedPrivacyRequestCount",
        "webhookRecordsDeleted",
        "wouldRevokeSerialCount",
    )
)


def _utc(value: datetime) -> datetime:
    if not isinstance(value, datetime):
        raise TypeError("value must be a datetime")
    if value.tzinfo is None or value.utcoffset() is None:
        return value.replace(tzinfo=timezone.utc)
    return value.astimezone(timezone.utc)


def _optional_identifier(value: Optional[str], field_name: str) -> Optional[str]:
    if value is None:
        return None
    return validate_identifier(value, field_name)


def _require_limit(value: int, maximum: int) -> int:
    if isinstance(value, bool) or not isinstance(value, int):
        raise TypeError("limit must be an integer")
    if value < 1 or value > maximum:
        raise ValueError(f"limit must be between one and {maximum}")
    return value


def _cursor_value(value: Optional[str]) -> Optional[tuple[datetime, str]]:
    if value is None:
        return None
    if not isinstance(value, str) or not value.strip() or len(value) > 2048:
        raise ValueError("cursor must be a bounded non-empty string or None")
    encoded = value.strip()
    encoded += "=" * (-len(encoded) % 4)
    try:
        raw = base64.urlsafe_b64decode(encoded.encode("ascii"))
    except (UnicodeError, ValueError) as exc:
        raise ValueError("cursor is invalid") from exc
    parsed = parse_bounded_json(raw, maximum_bytes=MAXIMUM_AUDIT_CURSOR_BYTES)
    if not isinstance(parsed, Mapping):
        raise ValueError("cursor is invalid")
    occurred_at = parsed.get("occurredAt")
    event_id = parsed.get("eventId")
    if not isinstance(occurred_at, str):
        raise ValueError("cursor occurrence time is invalid")
    try:
        timestamp = datetime.fromisoformat(occurred_at.replace("Z", "+00:00"))
    except ValueError as exc:
        raise ValueError("cursor occurrence time is invalid") from exc
    return _utc(timestamp), validate_identifier(event_id, "cursor event ID")


def _encode_cursor(event: AuditEvent) -> str:
    document = canonicalize_json(
        {
            "occurredAt": _utc(event.occurred_at).isoformat(),
            "eventId": event.id,
        }
    )
    return base64.urlsafe_b64encode(document).decode("ascii").rstrip("=")


def _safe_metadata(value: object) -> dict[str, object]:
    if not isinstance(value, Mapping):
        return {}
    result = {}
    for key, item in value.items():
        if key not in AUDIT_METADATA_ALLOWLIST:
            continue
        if item is None or isinstance(item, (bool, int, float, str)):
            result[key] = item
        elif isinstance(item, list) and all(isinstance(entry, str) for entry in item):
            result[key] = item[:500]
    return result


def _event_mapping(event: AuditEvent) -> dict[str, object]:
    return {
        "eventId": event.id,
        "occurredAt": _utc(event.occurred_at).isoformat(),
        "actorType": event.actor_type,
        "actorId": event.actor_id,
        "action": event.action,
        "targetType": event.target_type,
        "targetId": event.target_id,
        "reason": event.reason,
        "correlationId": event.correlation_id,
        "metadata": _safe_metadata(event.metadata_json),
    }


class AuditAdministrationService:
    def __init__(self, session: Session) -> None:
        if not isinstance(session, Session):
            raise TypeError("session must be a Session")
        self._session = session

    def search(
        self,
        actor_user_id: str,
        *,
        action: Optional[str] = None,
        target_type: Optional[str] = None,
        target_id: Optional[str] = None,
        event_actor_id: Optional[str] = None,
        occurred_from: Optional[datetime] = None,
        occurred_to: Optional[datetime] = None,
        cursor: Optional[str] = None,
        limit: int = 100,
    ) -> dict[str, object]:
        self._require_server_admin(actor_user_id)
        normalized_action = _optional_identifier(action, "action")
        normalized_target_type = _optional_identifier(target_type, "target_type")
        normalized_target_id = _optional_identifier(target_id, "target_id")
        normalized_event_actor = _optional_identifier(event_actor_id, "event_actor_id")
        start = None if occurred_from is None else _utc(occurred_from)
        end = None if occurred_to is None else _utc(occurred_to)
        if start is not None and end is not None and end <= start:
            raise ValueError("occurred_to must be later than occurred_from")
        cursor_value = _cursor_value(cursor)
        bounded_limit = _require_limit(limit, MAXIMUM_AUDIT_SEARCH_EVENTS)
        events = self._query(
            action=normalized_action,
            target_type=normalized_target_type,
            target_id=normalized_target_id,
            event_actor_id=normalized_event_actor,
            occurred_from=start,
            occurred_to=end,
            cursor=cursor_value,
            limit=bounded_limit + 1,
        )
        has_more = len(events) > bounded_limit
        returned = events[:bounded_limit]
        return {
            "events": [_event_mapping(value) for value in returned],
            "hasMore": has_more,
            "nextCursor": (
                _encode_cursor(returned[-1]) if has_more and returned else None
            ),
        }

    def export(
        self,
        actor_user_id: str,
        correlation_id: str,
        *,
        action: Optional[str] = None,
        target_type: Optional[str] = None,
        target_id: Optional[str] = None,
        event_actor_id: Optional[str] = None,
        occurred_from: Optional[datetime] = None,
        occurred_to: Optional[datetime] = None,
        limit: int = MAXIMUM_AUDIT_EXPORT_EVENTS,
    ) -> dict[str, object]:
        actor = self._require_server_admin(actor_user_id)
        normalized_correlation = validate_identifier(correlation_id, "correlation_id")
        normalized_action = _optional_identifier(action, "action")
        normalized_target_type = _optional_identifier(target_type, "target_type")
        normalized_target_id = _optional_identifier(target_id, "target_id")
        normalized_event_actor = _optional_identifier(event_actor_id, "event_actor_id")
        start = None if occurred_from is None else _utc(occurred_from)
        end = None if occurred_to is None else _utc(occurred_to)
        if start is not None and end is not None and end <= start:
            raise ValueError("occurred_to must be later than occurred_from")
        bounded_limit = _require_limit(limit, MAXIMUM_AUDIT_EXPORT_EVENTS)
        events = self._query(
            action=normalized_action,
            target_type=normalized_target_type,
            target_id=normalized_target_id,
            event_actor_id=normalized_event_actor,
            occurred_from=start,
            occurred_to=end,
            cursor=None,
            limit=bounded_limit,
        )
        generated_at = datetime.now(timezone.utc)
        document: dict[str, object] = {
            "schema": AUDIT_EXPORT_SCHEMA,
            "schemaVersion": AUDIT_EXPORT_SCHEMA_VERSION,
            "generatedAt": generated_at.isoformat(),
            "filters": {
                "action": normalized_action,
                "targetType": normalized_target_type,
                "targetId": normalized_target_id,
                "actorId": normalized_event_actor,
                "occurredFrom": None if start is None else start.isoformat(),
                "occurredTo": None if end is None else end.isoformat(),
                "limit": bounded_limit,
            },
            "events": [_event_mapping(value) for value in events],
        }
        document["sha256"] = hashlib.sha256(canonicalize_json(document)).hexdigest()
        self._session.add(
            AuditEvent(
                id=new_identifier("audit"),
                actor_type="user",
                actor_id=actor.id,
                action="audit.exported",
                target_type="audit_log",
                target_id=None,
                reason=None,
                correlation_id=normalized_correlation,
                source_address_digest=None,
                metadata_json={"quantity": len(events)},
            )
        )
        self._session.commit()
        return document

    def _require_server_admin(self, actor_user_id: str) -> User:
        normalized_actor = validate_identifier(actor_user_id, "actor_user_id")
        actor = self._session.get(User, normalized_actor)
        if actor is None or actor.status != "active" or not actor.is_server_admin:
            raise ServerLicensingError(
                ServerErrorCode.AUTHORIZATION_DENIED,
                "an active server administrator is required",
                status_code=403,
            )
        return actor

    def _query(
        self,
        *,
        action: Optional[str],
        target_type: Optional[str],
        target_id: Optional[str],
        event_actor_id: Optional[str],
        occurred_from: Optional[datetime],
        occurred_to: Optional[datetime],
        cursor: Optional[tuple[datetime, str]],
        limit: int,
    ) -> list[AuditEvent]:
        query = select(AuditEvent)
        dialect_name = self._session.get_bind().dialect.name
        database_start = (
            None
            if occurred_from is None
            else occurred_from.replace(tzinfo=None)
            if dialect_name == "sqlite"
            else occurred_from
        )
        database_end = (
            None
            if occurred_to is None
            else occurred_to.replace(tzinfo=None)
            if dialect_name == "sqlite"
            else occurred_to
        )
        if action is not None:
            query = query.where(AuditEvent.action == action)
        if target_type is not None:
            query = query.where(AuditEvent.target_type == target_type)
        if target_id is not None:
            query = query.where(AuditEvent.target_id == target_id)
        if event_actor_id is not None:
            query = query.where(AuditEvent.actor_id == event_actor_id)
        occurred_expression = (
            func.datetime(AuditEvent.occurred_at)
            if dialect_name == "sqlite"
            else AuditEvent.occurred_at
        )
        if database_start is not None:
            start_expression = (
                func.datetime(database_start)
                if dialect_name == "sqlite"
                else database_start
            )
            query = query.where(occurred_expression >= start_expression)
        if database_end is not None:
            end_expression = (
                func.datetime(database_end)
                if dialect_name == "sqlite"
                else database_end
            )
            query = query.where(occurred_expression < end_expression)
        if cursor is not None:
            cursor_time, cursor_id = cursor
            database_cursor_time = (
                cursor_time.replace(tzinfo=None)
                if dialect_name == "sqlite"
                else cursor_time
            )
            cursor_time_expression = (
                func.datetime(database_cursor_time)
                if dialect_name == "sqlite"
                else database_cursor_time
            )
            query = query.where(
                or_(
                    occurred_expression < cursor_time_expression,
                    and_(
                        occurred_expression == cursor_time_expression,
                        AuditEvent.id < cursor_id,
                    ),
                )
            )
        return self._session.scalars(
            query.order_by(AuditEvent.occurred_at.desc(), AuditEvent.id.desc()).limit(limit)
        ).all()


__all__ = [
    "AUDIT_EXPORT_SCHEMA",
    "AuditAdministrationService",
    "MAXIMUM_AUDIT_EXPORT_EVENTS",
    "MAXIMUM_AUDIT_SEARCH_EVENTS",
]
