import json
from api import ComposerAPI, GenerationRequest
from data import midi2wavtool, wavtool2midi
from bottle import Bottle, run, request, response
import threading
from pathlib import Path
from loader import load_by_path
from util import configure_logging, stopwatch
from warmup import warmup

# import gil_load
import sys


def serve(path):
    path = Path(path)

    with stopwatch("startup"):
        api = ComposerAPI(*load_by_path(path))

        warmup(api)

    app = Bottle()
    netlock = threading.Lock()

    @app.hook("after_request")
    def cors():
        response.headers["Access-Control-Allow-Origin"] = "*"
        response.headers[
            "Access-Control-Allow-Methods"
        ] = "PUT, GET, POST, DELETE, OPTIONS"
        response.headers[
            "Access-Control-Allow-Headers"
        ] = "Origin, Accept, Content-Type, X-Requested-With, X-CSRF-Token"

    @app.route("/continue", method=["OPTIONS", "POST"])
    def generate():
        if request.method == "OPTIONS":
            return {}

        obj = request.json
        accompany = obj.get("accompany", None)
        if accompany is not None:
            accompany = wavtool2midi({"notes": accompany})

        text_prompt = obj.get("textPrompt", None)

        with netlock:
            midi_out = api.generate(
                wavtool2midi(obj),
                [GenerationRequest.from_json(x) for x in obj["requests"]],
                accompany=accompany,
                text_prompt=text_prompt,
            )

        response.content_type = "application/json"
        return json.dumps([midi2wavtool(x) for x in midi_out])

    run(app, host="0.0.0.0", port=8081, debug=True)


if __name__ == "__main__":
    import sys

    configure_logging()
    serve(sys.argv[1])
