#!/usr/bin/env python
import rospy
import smach
import rospkg
import smach_ros

from bsc_receptionist import states, data, utils

import typer
from typing import Optional
from typing_extensions import Annotated
from geometry_msgs.msg import Pose, Point, Quaternion
from os import listdir, path
from PIL import Image
import numpy as np

import json
import cv2_img

from lasr_vector_database_msgs.msg import Property
from lasr_vector_database_msgs.srv import CreateCollection, InsertVector, QueryVector
from lasr_vision_msgs.srv import DeepFaceDetection

r = rospkg.RosPack()

SPLIT = [1, 2, 5]
DATABASE_NAME = "receptionist"
IMAGE_ROOT = path.join(r.get_path('bsc_receptionist'),
                       'datasets', 'reidentification')


def main():
    rospy.init_node("unit_test_2")

    # configure three collections
    create_collection = rospy.ServiceProxy(
        f'/database/vectors/{DATABASE_NAME}/create_collection', CreateCollection)

    for learn_count in SPLIT:
        create_collection(name=f"Face{learn_count}",
                          skip_if_exists=False, clear_if_exists=True)

    # process all face images and load into memory
    total = {}
    failures = {}
    vectors = {}
    detection_confidence = {}
    vectorise = rospy.ServiceProxy(f'/deepface/detect', DeepFaceDetection)
    for jdx, name in enumerate(listdir(IMAGE_ROOT)):
        # if jdx > 0:
        # continue
        if name == '.DS_Store':
            continue

        vectors[name] = []
        total[name] = 0
        failures[name] = 0
        detection_confidence[name] = []
        files = sorted([f for f in
                        listdir(path.join(IMAGE_ROOT, name)) if f != '.DS_Store'], key=lambda i: int(i.split('.')[0]))
        for fn in files:
            img = Image.open(path.join(IMAGE_ROOT, name, fn))
            msg = cv2_img.pillow_img_to_msg(img)
            print(f'Process {fn} for {name}')
            detections = vectorise(
                image_raw=msg, model="VGG-Face").detected_faces
            total[name] += 1
            if len(detections) == 0:
                failures[name] += 1
            else:
                det = detections[0]
                detection_confidence[name].append(det.confidence)
                vectors[name].append(det.vector)

    # insert vectors
    insert_vector = rospy.ServiceProxy(
        f'/database/vectors/{DATABASE_NAME}/insert', InsertVector)
    for learn_count in SPLIT:
        for name, vecs in vectors.items():
            for idx, vector in enumerate(vecs):
                if idx < learn_count:
                    print("inserting vector for", idx, name)
                    insert_vector(
                        name=f'Face{learn_count}',
                        properties=[
                            Property(key="name", value=name)
                        ],
                        vector=vector
                    )

    # query vectors
    query_vector = rospy.ServiceProxy(
        f'/database/vectors/{DATABASE_NAME}/query', QueryVector)
    success = {}
    confusion_matrix = {}
    for learn_count in SPLIT:
        success[learn_count] = {}
        confusion_matrix[learn_count] = {}
        for name, vecs in vectors.items():
            confusion_matrix[learn_count][name] = {}
            for nn in vectors.keys():
                confusion_matrix[learn_count][name][nn] = 0

            success[learn_count][name] = []
            for idx, vector in enumerate(vecs):
                if idx >= learn_count:
                    print("searching vector for", idx, name)
                    results = query_vector(
                        name=f'Face{learn_count}',
                        vector=vector,
                        limit=1
                    ).results
                    if len(results) == 1:
                        assert results[0].properties[0].key == 'name'
                        if results[0].properties[0].value == name:
                            success[learn_count][name].append(
                                results[0].certainty)
                            confusion_matrix[learn_count][name][name] += 1
                        else:
                            confusion_matrix[learn_count][name][results[0].properties[0].value] += 1
                    else:
                        confusion_matrix[learn_count][name]["no match"] += 1

    results = {
        "raw": {
            "success": success,
            "det_conf": detection_confidence
        },
        "totals": total,
        "failures": failures,
        "det_conf_avg": {name: np.mean(values) for name, values in detection_confidence.items()},
        "reidentify_success_rate": {
            learn_count: {
                name: len(runs) / (total[name] - failures[name] - learn_count) for name, runs in names.items()
            } for learn_count, names in success.items()
        },
        "reidentify_confidence": {
            learn_count: {
                name: {
                    "min": np.min(runs),
                    "mean": np.mean(runs),
                    "max": np.mean(runs)
                } for name, runs in names.items()
            } for learn_count, names in success.items()
        },
        "confusion_matrix": confusion_matrix
    }
    utils.module_log("Unit 02", "Evaluate Reidentification",
                     "[purple]Test results are ready![/purple]", results)

    fr = open('results2.json', 'w')
    fr.write(json.dumps(results))
    fr.close()


if __name__ == "__main__":
    import warnings

    with warnings.catch_warnings():
        warnings.simplefilter("ignore")
        try:
            typer.run(main)
        except Exception as e:
            print(e)
