import uuid
from contextlib import contextmanager
from contextvars import ContextVar
from dataclasses import dataclass, field
from typing import Iterable

from django.db import connection

from api.models import TtActivity, TtLocation, TtStaff
from api.services.integration.outbox import (
    append_activity_event,
    append_resource_event,
    capture_activity_state,
    capture_activity_states,
    capture_resource_state,
    finalize_change_set,
)


@dataclass
class MutationTracker:
    change_set_id: uuid.UUID = field(default_factory=uuid.uuid4)
    actor_id: int | None = None
    request_id: str | None = None
    origin: str = "timetabler"
    resource_before: dict[tuple[str, int], dict | None] = field(default_factory=dict)
    dirty_resources: set[tuple[str, int]] = field(default_factory=set)
    deleted_resources: set[tuple[str, int]] = field(default_factory=set)
    activity_before: dict[int, dict | None] = field(default_factory=dict)
    dirty_activities: set[int] = field(default_factory=set)
    deleted_activities: set[int] = field(default_factory=set)


_tracker: ContextVar[MutationTracker | None] = ContextVar(
    "timetabler_integration_mutation_tracker",
    default=None,
)


def current_tracker() -> MutationTracker | None:
    return _tracker.get()


def _resource_key(resource) -> tuple[str, int]:
    if isinstance(resource, TtStaff):
        return "staff", resource.pk
    if isinstance(resource, TtLocation):
        return "location", resource.pk
    raise TypeError(f"Unsupported integration resource: {type(resource)!r}")


def track_resource(resource, *, deleted: bool = False) -> None:
    tracker = current_tracker()
    if tracker is None or resource.pk is None:
        return
    key = _resource_key(resource)
    if key not in tracker.resource_before:
        tracker.resource_before[key] = capture_resource_state(resource)
    tracker.dirty_resources.add(key)
    if deleted:
        tracker.deleted_resources.add(key)


def track_resource_created(resource) -> None:
    tracker = current_tracker()
    if tracker is None or resource.pk is None:
        return
    key = _resource_key(resource)
    tracker.resource_before.setdefault(key, None)
    tracker.dirty_resources.add(key)


def track_resource_queryset(resources: Iterable[TtStaff] | Iterable[TtLocation]) -> None:
    resources = list(resources)
    if not resources:
        return
    model = type(resources[0])
    ids = sorted({resource.pk for resource in resources if resource.pk is not None})
    # Bulk-import callers mutate their in-memory objects before bulk_update().
    # Reload here so the replacement always contains the committed old state.
    for resource in model.objects.filter(pk__in=ids).order_by("id"):
        track_resource(resource)


def track_activity(activity: TtActivity, *, deleted: bool = False) -> None:
    tracker = current_tracker()
    if tracker is None or activity.pk is None:
        return
    if activity.pk not in tracker.activity_before:
        tracker.activity_before[activity.pk] = capture_activity_state(activity)
    tracker.dirty_activities.add(activity.pk)
    if deleted:
        tracker.deleted_activities.add(activity.pk)


def track_activity_created(activity: TtActivity) -> None:
    tracker = current_tracker()
    if tracker is None or activity.pk is None:
        return
    tracker.activity_before.setdefault(activity.pk, None)
    tracker.dirty_activities.add(activity.pk)


def track_activity_ids(activity_ids: Iterable[int]) -> None:
    tracker = current_tracker()
    if tracker is None:
        return
    ids = sorted({int(activity_id) for activity_id in activity_ids})
    missing = [activity_id for activity_id in ids if activity_id not in tracker.activity_before]
    tracker.activity_before.update(capture_activity_states(missing))
    tracker.dirty_activities.update(ids)


def flush_tracked_mutations(tracker: MutationTracker | None = None) -> int:
    tracker = tracker or current_tracker()
    if tracker is None:
        return 0
    if not connection.in_atomic_block:
        raise RuntimeError("Tracked integration mutations must flush before the database transaction commits")

    count = 0
    for aggregate_type, resource_id in sorted(tracker.dirty_resources):
        model = TtStaff if aggregate_type == "staff" else TtLocation
        resource = model.objects.filter(pk=resource_id).first()
        previous = tracker.resource_before.get((aggregate_type, resource_id))
        deleted = (aggregate_type, resource_id) in tracker.deleted_resources or resource is None
        if deleted:
            # A lightweight instance carries the stable primary key and model type;
            # the full old snapshot was captured before the collector removed it.
            resource = model(pk=resource_id)
            action = "deleted"
        elif previous is None:
            action = "created"
        elif previous == capture_resource_state(resource):
            continue
        elif resource.status != model.STATUS_TO_CODE["active"]:
            action = "archived"
        else:
            action = "updated"
        append_resource_event(
            resource,
            action=action,
            previous=previous,
            tombstone=deleted,
            change_set_id=tracker.change_set_id,
            actor_id=tracker.actor_id,
            origin=tracker.origin,
            defer_change_set_finalization=True,
        )
        count += 1

    final_activities = capture_activity_states(
        sorted(tracker.dirty_activities - tracker.deleted_activities)
    )
    for activity_id in sorted(tracker.dirty_activities):
        previous = tracker.activity_before.get(activity_id)
        activity = TtActivity.objects.filter(pk=activity_id).first()
        deleted = activity_id in tracker.deleted_activities or activity is None
        current = final_activities.get(activity_id)
        if not deleted and previous == current:
            continue
        if deleted:
            action = "deleted"
        elif previous is None:
            action = "created"
        elif previous and previous.get("scheduled") and not current.get("scheduled"):
            action = "unscheduled"
        elif current.get("scheduled"):
            action = "allocation_replaced"
        else:
            action = "updated"
        append_activity_event(
            activity=activity,
            activity_id=activity_id,
            action=action,
            previous=previous,
            tombstone=deleted,
            change_set_id=tracker.change_set_id,
            request_id=tracker.request_id,
            actor_id=tracker.actor_id,
            origin=tracker.origin,
            defer_change_set_finalization=True,
        )
        count += 1
    if count:
        finalize_change_set(tracker.change_set_id)
    return count


@contextmanager
def mutation_tracking(*, actor_id=None, request_id=None, origin="timetabler", flush=True):
    tracker = MutationTracker(actor_id=actor_id, request_id=request_id, origin=origin)
    token = _tracker.set(tracker)
    try:
        yield tracker
        if flush:
            flush_tracked_mutations(tracker)
    finally:
        _tracker.reset(token)
