"""AWS Rekognition client + face-recognition helpers for HR attendance."""

from functools import lru_cache

import boto3
from botocore.exceptions import ClientError
from django.conf import settings


@lru_cache(maxsize=1)
def get_client():
    return boto3.client(
        'rekognition',
        aws_access_key_id=settings.AWS_ACCESS_KEY_ID,
        aws_secret_access_key=settings.AWS_SECRET_ACCESS_KEY,
        region_name=settings.AWS_REGION,
    )


def ensure_collection(collection_id: str | None = None) -> str:
    collection_id = collection_id or settings.REKOGNITION_COLLECTION_ID
    client = get_client()
    try:
        client.create_collection(CollectionId=collection_id)
    except ClientError as exc:
        if exc.response['Error']['Code'] != 'ResourceAlreadyExistsException':
            raise
    return collection_id


def list_faces(collection_id: str | None = None) -> list[dict]:
    """All faces in the collection."""
    client = get_client()
    collection_id = collection_id or settings.REKOGNITION_COLLECTION_ID
    faces: list[dict] = []
    token = None
    while True:
        kwargs = {'CollectionId': collection_id, 'MaxResults': 4096}
        if token:
            kwargs['NextToken'] = token
        response = client.list_faces(**kwargs)
        faces.extend(response.get('Faces') or [])
        token = response.get('NextToken')
        if not token:
            break
    return faces


def delete_faces_for(external_image_id: str, collection_id: str | None = None) -> int:
    """Remove every face indexed under external_image_id. Returns count deleted."""
    collection_id = collection_id or settings.REKOGNITION_COLLECTION_ID
    face_ids = [
        f['FaceId']
        for f in list_faces(collection_id)
        if f.get('ExternalImageId') == external_image_id
    ]
    if not face_ids:
        return 0
    get_client().delete_faces(CollectionId=collection_id, FaceIds=face_ids)
    return len(face_ids)


def index_face(image_bytes: bytes, external_image_id: str, collection_id: str | None = None) -> dict:
    """Enroll a face. external_image_id should be the employee's stable id (e.g. user uuid).

    Purges any face already indexed under this id first, so re-enrolling
    replaces rather than accumulates (a stale extra face can out-match the
    current one, since search_face only looks at the top result).
    """
    collection_id = collection_id or settings.REKOGNITION_COLLECTION_ID
    result = get_client().index_faces(
        CollectionId=collection_id,
        Image={'Bytes': image_bytes},
        ExternalImageId=external_image_id,
        DetectionAttributes=['DEFAULT'],
        MaxFaces=1,
        QualityFilter='AUTO',
    )
    new_ids = {r['Face']['FaceId'] for r in (result.get('FaceRecords') or [])}
    if new_ids:
        stale = [
            f['FaceId']
            for f in list_faces(collection_id)
            if f.get('ExternalImageId') == external_image_id and f['FaceId'] not in new_ids
        ]
        if stale:
            get_client().delete_faces(CollectionId=collection_id, FaceIds=stale)
    return result


def search_face(image_bytes: bytes, collection_id: str | None = None) -> dict | None:
    """Match a face against the collection. Returns top match or None."""
    response = get_client().search_faces_by_image(
        CollectionId=collection_id or settings.REKOGNITION_COLLECTION_ID,
        Image={'Bytes': image_bytes},
        FaceMatchThreshold=settings.REKOGNITION_FACE_MATCH_THRESHOLD,
        MaxFaces=1,
    )
    matches = response.get('FaceMatches') or []
    return matches[0] if matches else None
