"""Run with the pinned Supervision checkout installed; no weights or GPU needed."""
import json
from unittest.mock import patch

import numpy as np
import supervision as sv
from supervision.detection.compact_mask import CompactMask


def segment(tile):
    """Supply deterministic masks in place of model inference."""
    h, w = tile.shape[:2]
    masks = np.zeros((2, h, w), dtype=bool)
    masks[0, 8:24, 8:24] = True
    masks[1, h - 32:h - 8, w - 32:w - 8] = True
    return sv.Detections(
        xyxy=np.array([[8, 8, 23, 23], [w - 32, h - 32, w - 9, h - 9]], dtype=float),
        mask=masks,
        confidence=np.array([0.9, 0.8]),
        class_id=np.array([0, 1]),
    )


def payload(mask):
    """Count NumPy backing buffers, excluding Python overhead and peak RSS."""
    return sum(r.nbytes for r in mask._rles) + mask._offsets.nbytes + mask._crop_shapes.nbytes


def order(detections):
    return np.lexsort((detections.class_id, detections.xyxy[:, 1], detections.xyxy[:, 0]))


image = np.zeros((1024, 1024, 3), dtype=np.uint8)
args = dict(callback=segment, slice_wh=256, overlap_wh=0, thread_workers=1,
            overlap_filter=sv.OverlapFilter.NON_MAX_SUPPRESSION, iou_threshold=0.3)
dense = sv.InferenceSlicer(**args, compact_masks=False)(image)
with patch.object(CompactMask, 'to_dense', side_effect=AssertionError('dense conversion')), \
     patch.object(CompactMask, '__array__', side_effect=AssertionError('array conversion')):
    compact = sv.InferenceSlicer(**args, compact_masks=True)(image)
assert isinstance(compact.mask, CompactMask)
repacked = compact.mask.repack()
equal = np.array_equal(dense.mask[order(dense)], compact.mask.to_dense()[order(compact)])
assert equal
np.testing.assert_array_equal(dense.xyxy[order(dense)], compact.xyxy[order(compact)])
np.testing.assert_array_equal(repacked.to_dense(), compact.mask.to_dense())
print(json.dumps({'fixture': '1024x1024; 16 tiles; two rectangles per tile',
                  'detections': len(compact), 'dense_mask_bytes': dense.mask.nbytes,
                  'compact_buffer_bytes': payload(compact.mask),
                  'repacked_buffer_bytes': payload(repacked),
                  'pixel_equal': equal, 'compact_pipeline_dense_guard': 'passed'}, indent=2))


def outside_box(tile):
    h, w = tile.shape[:2]
    masks = np.zeros((1, h, w), dtype=bool)
    masks[0, 0, 0] = masks[0, h - 1, w - 1] = True
    return sv.Detections(xyxy=np.array([[0, 0, 10, 10]], dtype=float), mask=masks,
                         confidence=np.array([0.9]), class_id=np.array([0]))


outside = sv.InferenceSlicer(callback=outside_box, slice_wh=100, overlap_wh=0,
    overlap_filter=sv.OverlapFilter.NONE, compact_masks=True)(np.zeros((100, 100, 3), np.uint8))
corner = bool(outside.mask.to_dense()[0, 99, 99])
corner_repacked = bool(outside.mask.repack().to_dense()[0, 99, 99])
assert corner and corner_repacked
print(json.dumps({'outside_detector_box_survives': corner,
                  'outside_detector_box_survives_repack': corner_repacked}))

checker = (np.indices((128, 128)).sum(axis=0) % 2 == 0)[None]
encoded = CompactMask.from_dense(checker, np.array([[0, 0, 127, 127]]), (128, 128))
np.testing.assert_array_equal(encoded.to_dense(), checker)
print(json.dumps({'fixture': '128x128 checkerboard', 'dense_mask_bytes': checker.nbytes,
                  'compact_buffer_bytes': payload(encoded)}, indent=2))

# A compact input does not guarantee a compact implementation of every operation.
boxes = np.array([[0, 0, 7, 7], [0, 0, 7, 7]], dtype=float)
duplicates = sv.Detections(xyxy=boxes,
    mask=CompactMask.from_dense(np.ones((2, 8, 8), bool), boxes, (8, 8)),
    confidence=np.array([0.9, 0.8]), class_id=np.array([0, 0]))
original_dense = CompactMask.to_dense
calls = []


def count_dense(self):
    calls.append(self.shape)
    return original_dense(self)


with patch.object(CompactMask, 'to_dense', count_dense):
    suppressed = duplicates.with_nms(threshold=0.3)
nms_calls = len(calls)
calls.clear()
with patch.object(CompactMask, 'to_dense', count_dense):
    merged = duplicates.with_nmm(threshold=0.3)
print(json.dumps({'duplicate_masks': 2, 'nms_output': len(suppressed),
                  'nms_to_dense_calls': nms_calls, 'nmm_output': len(merged),
                  'nmm_to_dense_calls': len(calls)}, indent=2))
assert nms_calls == 0 and len(calls) > 0
