#!/usr/bin/env python
import rospy
import smach
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

# TODO: move into common package
import numpy as np


def euler_to_quaternion(yaw, pitch, roll):
    qx = np.sin(roll/2) * np.cos(pitch/2) * np.cos(yaw/2) - \
        np.cos(roll/2) * np.sin(pitch/2) * np.sin(yaw/2)
    qy = np.cos(roll/2) * np.sin(pitch/2) * np.cos(yaw/2) + \
        np.sin(roll/2) * np.cos(pitch/2) * np.sin(yaw/2)
    qz = np.cos(roll/2) * np.cos(pitch/2) * np.sin(yaw/2) - \
        np.sin(roll/2) * np.sin(pitch/2) * np.cos(yaw/2)
    qw = np.cos(roll/2) * np.cos(pitch/2) * np.cos(yaw/2) + \
        np.sin(roll/2) * np.sin(pitch/2) * np.sin(yaw/2)
    return [qx, qy, qz, qw]


def main(using_simulator: Annotated[Optional[bool], typer.Argument()] = None):
    rospy.init_node("unit_test_01")
    sm = smach.StateMachine(outcomes=['failed', 'end'], output_keys=[])

    class SwapGuests(smach.State):
        def __init__(self):
            smach.State.__init__(self, outcomes=['succeeded'])

        def execute(self, userdata):
            utils.module_log('Unit 01', 'waiting for guests to swap...',
                             '[blue]Please position yourself.[/blue]', 'Continuing in 5 seconds...')
            rospy.sleep(5)
            return 'succeeded'

    class AlertNextStage(smach.State):
        def __init__(self):
            smach.State.__init__(self, outcomes=['succeeded'])

        def execute(self, userdata):
            utils.module_log('Unit 01', 'will now start continously detecting who is in front of the robot...',
                             '[green]Two guests have been successfully registered![/green]', 'Continuing in 3 seconds...')
            rospy.sleep(3)
            return 'succeeded'

    class PrepareGuestData(smach.State):
        def __init__(self):
            smach.State.__init__(self, outcomes=['succeeded'], input_keys=[
                                 'index'], output_keys=['name', 'model'])

        def execute(self, userdata):
            userdata['name'] = f"learn_guest_{userdata['index']}"
            userdata['model'] = 'female2' if userdata['index'] == 3 else 'male' if userdata['index'] == 1 else 'female'
            return 'succeeded'

    class MapFirstFace(smach.State):
        def __init__(self):
            smach.State.__init__(self, outcomes=['succeeded'], input_keys=[
                                 'detected_faces'], output_keys=['vector'])

        def execute(self, userdata):
            assert len(userdata['detected_faces']) == 1, 'expected 1 face'
            userdata['vector'] = userdata['detected_faces'][0].vector
            return 'succeeded'

    class MapIndexFace(smach.State):
        def __init__(self):
            smach.State.__init__(self, outcomes=['succeeded'], input_keys=[
                                 'index', 'detected_faces'], output_keys=['vector'])

        def execute(self, userdata):
            userdata['vector'] = userdata['detected_faces'][userdata['index']].vector
            return 'succeeded'

    class ValidateResult(smach.State):
        def __init__(self):
            smach.State.__init__(
                self, outcomes=['succeeded'], input_keys=['results', 'index'])

        def execute(self, userdata):
            assert len(userdata['results']
                       ) != 0 or userdata['index'] == 3, 'no results!'

            result = userdata['results'][0]
            indices = [idx for idx, prop in enumerate(
                result.properties) if prop.key == 'guest']

            if userdata['index'] != 3 or result.certainty > 0.95:
                assert int(
                    result.properties[indices[0]].value) == userdata['index'], 'query result does not match expected guest!'

            return 'succeeded'

    with sm:
        if using_simulator:
            smach.StateMachine.add('RESET_SIMULATION', utils.SimulationReset(), transitions={
                                   'succeeded': 'SETUP'})

        smach.StateMachine.add('SETUP', utils.Setup(), transitions={
                               'succeeded': 'LEARN_TWO_GUESTS', 'failed': 'failed'})

        run_for_two_guests = smach.Iterator(
            outcomes=['succeeded', 'failed'],
            it=lambda: range(1, 3),
            input_keys=[],
            output_keys=[],
            it_label='index',
            exhausted_outcome='succeeded'
        )

        with run_for_two_guests:
            container_sm = smach.StateMachine(
                outcomes=['succeeded', 'failed', 'continue'], input_keys=['index'], output_keys=[])

            with container_sm:
                if using_simulator:
                    smach.StateMachine.add('PREPARE_GUEST_DATA',
                                           PrepareGuestData(),
                                           transitions={
                                               'succeeded': 'SPAWN_GUEST'
                                           })

                    smach.StateMachine.add('SPAWN_GUEST',
                                           utils.SimulationSpawnGuest(
                                               Pose(
                                                   Point(2, 0, 0),
                                                   Quaternion(
                                                       *euler_to_quaternion(-1.53, 0, 0)
                                                   )
                                               )
                                           ),
                                           transitions={
                                               'succeeded': 'COLLECT_FACE_DATA'
                                           })
                else:
                    smach.StateMachine.add('WAIT_FOR_SWAP',
                                           SwapGuests(),
                                           transitions={
                                               'succeeded': 'COLLECT_FACE_DATA'
                                           })

                smach.StateMachine.add('COLLECT_FACE_DATA',
                                       states.GuestSetup(),
                                       transitions={
                                           'succeeded': 'DELETE_GUEST' if using_simulator else 'continue'
                                       })

                if using_simulator:
                    smach.StateMachine.add('DELETE_GUEST',
                                           utils.SimulationDeleteModel(),
                                           transitions={
                                               'succeeded': 'continue'
                                           })

            smach.Iterator.set_contained_state(
                'CONTAINER_STATE', container_sm, loop_outcomes=['continue'])

        smach.StateMachine.add('LEARN_TWO_GUESTS', run_for_two_guests, transitions={
            'succeeded': 'ASSESS_DETECTIONS' if using_simulator else 'ALERT_NEXT', 'failed': 'failed'
        })

        if using_simulator:
            run_for_three_guests = smach.Iterator(
                outcomes=['succeeded', 'failed'],
                it=lambda: range(1, 4),
                input_keys=[],
                output_keys=[],
                it_label='index',
                exhausted_outcome='succeeded'
            )

            with run_for_three_guests:
                container_sm = smach.StateMachine(
                    outcomes=['succeeded', 'failed', 'continue'], input_keys=['index'], output_keys=[])

                with container_sm:
                    smach.StateMachine.add('PREPARE_GUEST_DATA',
                                           PrepareGuestData(),
                                           transitions={
                                               'succeeded': 'SPAWN_GUEST'
                                           })

                    smach.StateMachine.add('SPAWN_GUEST',
                                           utils.SimulationSpawnGuest(
                                               Pose(
                                                   Point(2, 0, 0),
                                                   Quaternion(
                                                       *euler_to_quaternion(-1.53, 0, 0)
                                                   )
                                               )
                                           ),
                                           transitions={
                                               'succeeded': 'DETECT_FACES'
                                           })

                    smach.StateMachine.add('DETECT_FACES',
                                           utils.DetectFaces(),
                                           transitions={
                                               'succeeded': 'MAP_FACE',
                                               'preempted': 'failed',
                                               'failed': 'failed'
                                           })

                    smach.StateMachine.add('MAP_FACE', MapFirstFace(), transitions={
                                           'succeeded': 'QUERY_FACE'})

                    smach.StateMachine.add('QUERY_FACE', utils.VectorQuery(), transitions={
                                           'succeeded': 'CHECK_GUEST'})

                    smach.StateMachine.add('CHECK_GUEST', ValidateResult(), transitions={
                                           'succeeded': 'DELETE_GUEST'})

                    smach.StateMachine.add('DELETE_GUEST',
                                           utils.SimulationDeleteModel(),
                                           transitions={'succeeded': 'continue'})

                smach.Iterator.set_contained_state(
                    'CONTAINER_STATE', container_sm, loop_outcomes=['continue'])

            smach.StateMachine.add('ASSESS_DETECTIONS', run_for_three_guests, transitions={
                'succeeded': 'end', 'failed': 'failed'
            })
        else:
            smach.StateMachine.add('ALERT_NEXT',
                                   AlertNextStage(),
                                   transitions={
                                       'succeeded': 'DETECT_FACES'
                                   })

            smach.StateMachine.add('DETECT_FACES',
                                   utils.DetectFaces(),
                                   transitions={
                                       'succeeded': 'SEARCH_EACH_FACE',
                                       'preempted': 'failed',
                                       'failed': 'failed'
                                   })

            run_for_each_detection = smach.Iterator(
                outcomes=['succeeded', 'failed'],
                it=lambda: range(0, len(sm.userdata['detected_faces'])),
                input_keys=['detected_faces'],
                output_keys=[],
                it_label='index',
                exhausted_outcome='succeeded'
            )

            with run_for_each_detection:
                container_sm = smach.StateMachine(
                    outcomes=['succeeded', 'failed', 'continue'], input_keys=['detected_faces', 'index'], output_keys=[])

                with container_sm:
                    smach.StateMachine.add('MAP_FACE', MapIndexFace(), transitions={
                                           'succeeded': 'QUERY_FACE'})

                    smach.StateMachine.add('QUERY_FACE', utils.VectorQuery(), transitions={
                                           'succeeded': 'continue'})

                smach.Iterator.set_contained_state(
                    'CONTAINER_STATE', container_sm, loop_outcomes=['continue'])

            smach.StateMachine.add('SEARCH_EACH_FACE', run_for_each_detection, transitions={
                'succeeded': 'SLEEP', 'failed': 'failed'
            })

            smach.StateMachine.add('SLEEP', utils.Sleep(
                2), transitions={'succeeded': 'DETECT_FACES'})

    sis = smach_ros.IntrospectionServer('unit_test_01', sm, '/SM_ROOT')
    sis.start()
    sm.execute()
    rospy.signal_shutdown("down")
    sis.stop()

    utils.module_log("Unit 01", "Learn and Recognise People",
                     "[green]Test passed[/green]")


if __name__ == "__main__":
    import warnings

    with warnings.catch_warnings():
        warnings.simplefilter("ignore")
        typer.run(main)
