import json

from django.core.management.base import BaseCommand, CommandError
from django.db.models import F, Max, Window
from django.db.models.functions import RowNumber

from api.models import IntegrationOutbox


class Command(BaseCommand):
    help = "Export a stable, versioned canonical reconciliation manifest as JSON Lines."

    def add_arguments(self, parser):
        parser.add_argument("--aggregate", choices=["staff", "location", "activity"], required=True)
        parser.add_argument("--after-key", default="")
        parser.add_argument("--snapshot-id", type=int)
        parser.add_argument("--limit", type=int, default=1000)

    def handle(self, *args, **options):
        if options["limit"] < 1 or options["limit"] > 10000:
            raise CommandError("--limit must be between 1 and 10000")
        snapshot_id = options["snapshot_id"]
        if snapshot_id is None:
            snapshot_id = IntegrationOutbox.objects.aggregate(value=Max("id"))["value"] or 0
        queryset = (
            IntegrationOutbox.objects.filter(
                id__lte=snapshot_id,
                aggregate_type=options["aggregate"],
                aggregate_id__gt=options["after_key"],
            )
            .annotate(
                manifest_rank=Window(
                    expression=RowNumber(),
                    partition_by=[F("aggregate_type"), F("aggregate_id")],
                    order_by=F("event_version").desc(),
                )
            )
            .filter(manifest_rank=1)
            .order_by("aggregate_id")[: options["limit"]]
        )
        self.stderr.write(f"snapshot_id={snapshot_id}")
        last_key = options["after_key"]
        for event in queryset:
            last_key = event.aggregate_id
            self.stdout.write(
                json.dumps(
                    {
                        "snapshot_id": snapshot_id,
                        "aggregate_type": event.aggregate_type,
                        "aggregate_id": event.aggregate_id,
                        "version": event.event_version,
                        "event_id": str(event.event_id),
                        "state": event.payload.get("committed_state"),
                    },
                    sort_keys=True,
                    default=str,
                )
            )
        self.stderr.write(f"next_after_key={last_key}")
