#!/usr/bin/env python
import rospkg
import rospy
import typer
import json
import os

from rich import print
from enum import Enum
from shapely.geometry.polygon import Polygon
from geometry_msgs.msg import PoseWithCovarianceStamped, PointStamped

from bsc_receptionist import data, utils

rospy.init_node("configure_receptionist", anonymous=True)

points = typer.Typer()


@points.command()
def set(name: data.Points):
    msg = rospy.wait_for_message('/amcl_pose', PoseWithCovarianceStamped)
    data.Points.set_pose(name, msg.pose.pose)
    print(
        f"\n\n[bright_black]Set pose for {name} to:[/bright_black]\n", msg.pose.pose, "\n", sep='')


@points.command()
def get(name: data.Points):
    pose = data.Points.get_pose(name)
    print(
        f"\n\n[bright_black]The current pose for {name} is:[/bright_black]\n", pose, "\n", sep='')


polygons = typer.Typer()


@polygons.command()
def set(name: data.Polygons, point_count: int, topic: str = '/clicked_point'):
    print(
        f"\n\nWaiting for {point_count} published points on topic {topic}...")

    points = []
    for i in range(1, point_count + 1):
        msg = rospy.wait_for_message('/clicked_point', PointStamped)
        points.append(msg.point)
        print(f"\nReceived point {i} / {point_count}:\n", msg.point, sep='')

    polygon = [
        [point.x, point.y]
        for point
        in points
    ]

    data.Polygons.set_polygon(name, polygon)
    print(
        f"\n\n[bright_black]Set polygon for {name} to:[/bright_black]\n", Polygon(polygon), "\n", sep='')


@polygons.command()
def get(name: data.Polygons):
    polygon = data.Polygons.get_polygon(name)
    print(
        f"\n\n[bright_black]The current polygon for {name} is:[/bright_black]\n", polygon, "\n", sep='')


app = typer.Typer()
app.add_typer(points, name="points")
app.add_typer(polygons, name="polygons")

rp = rospkg.RosPack()
package_path = rp.get_path("bsc_receptionist")


@app.command()
def learn_host(name: str):
    import smach
    import lasr_skills
    sm = smach.StateMachine(outcomes=['failed', 'end'], output_keys=[])

    vectors = []

    @smach.cb_interface(input_keys=['detected_faces'], output_keys=[], outcomes=['succeeded'])
    def learn_vector(userdata):
        assert len(userdata['detected_faces']) == 1, 'expected 1 face'
        vectors.append(userdata['detected_faces'][0].vector)
        return 'succeeded'

    with sm:
        smach.StateMachine.add('DISABLE_HEAD', utils.HeadManagerStop(), transitions={
            'succeeded': 'RESET_HEAD'})

        smach.StateMachine.add('RESET_HEAD', lasr_skills.PlayMotion(motion_name='look_center'), transitions={
            'succeeded': 'LEARN_FACE_INFORMATION', 'preempted': 'failed', 'aborted': 'failed'})

        collect_five_faces = smach.Iterator(
            outcomes=['succeeded', 'failed'],
            it=lambda: range(0, 5),
            input_keys=[],
            output_keys=[],
            it_label='index',
            exhausted_outcome='succeeded'
        )

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

            with container_sm:
                smach.StateMachine.add('DETECT_FACES',
                                       utils.DetectFaces(),
                                       transitions={
                                           'succeeded': 'SAVE_VECTOR',
                                           'preempted': 'continue',
                                           'failed': 'continue'
                                       })

                smach.StateMachine.add('SAVE_VECTOR',
                                       smach.CBState(learn_vector),
                                       transitions={'succeeded': 'continue'})

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

        smach.StateMachine.add('LEARN_FACE_INFORMATION', collect_five_faces, transitions={
            'succeeded': 'end', 'failed': 'failed'
        })

    sm.execute()

    if len(vectors) == 0:
        print(f"\n\n[red]Failed to learn host![/red]\n")
        return

    data.save_host_data(name, vectors)
    print(
        f"\n\n[purple]Successfully learnt host![/purple]\n[bright_black]{len(vectors)} vectors were collected and saved.[/bright_black]\n")


@app.command()
def load(name: str):
    with open(os.path.abspath(os.path.join(package_path, 'configs', f'{name}.json')), 'r') as f:
        rospy.set_param('/receptionist', json.loads(f.read()))

    print(f"\n\n[purple]Successfully loaded configuration![/purple]\n")


@app.command()
def save(name: str):
    with open(os.path.abspath(os.path.join(package_path, 'configs', f'{name}.json')), 'w') as f:
        f.write(json.dumps(rospy.get_param('/receptionist')))

    print(f"\n\n[purple]Successfully saved configuration![/purple]\n")


if __name__ == "__main__":
    import warnings

    with warnings.catch_warnings():
        warnings.simplefilter("ignore")
        app()
