# pylint: disable=E1129

from pathlib import Path
import re
import music21
import psycopg
from psycopg.errors import DeadlockDetected


def music21_to_wavtool(elem, add_offset=0):
    if isinstance(elem, music21.note.Note):
        return [
            {
                "pitch": elem.pitch.midi,
                "start": float(elem.offset + add_offset),
                "end": float(elem.offset + add_offset + elem.duration.quarterLength),
            }
        ]
    if isinstance(elem, music21.chord.Chord):
        return [
            x
            for xs in (
                music21_to_wavtool(n, add_offset=elem.offset) for n in elem.notes
            )
            for x in xs
        ]
    if isinstance(
        elem,
        (
            music21.spanner.Slur,
            music21.expressions.TextExpression,
            music21.instrument.Instrument,
            music21.layout.SystemLayout,
            music21.layout.StaffLayout,
            music21.meter.TimeSignature,
            music21.note.Rest,
            music21.clef.Clef,
            music21.key.KeySignature,
            music21.dynamics.Crescendo,
            music21.dynamics.Diminuendo,
            music21.dynamics.Dynamic,
            music21.bar.Barline,
            music21.repeat.Coda,
        ),
    ):
        return []
    print(f"Unhandled element: {elem.__class__.__module__}.{elem.__class__.__name__}")
    return []


def collect(score_path):
    pages = []
    for f in sorted(Path(score_path).iterdir()):
        m = re.search(r"page(\d+)\.xml", f.name)
        if m is None:
            continue
        pages.append((int(m.group(1)), f))
    pages.sort()
    parts = {}
    page_start_time = 0
    for _, f in pages:
        s = music21.converter.parse(f)
        for elem in s:
            if isinstance(elem, music21.stream.Part):
                if elem.partName not in parts:
                    p = music21.stream.Part()
                    parts[elem.partName] = p
                    if page_start_time > 0:
                        p.append(music21.note.Rest(quarterLength=page_start_time))
                for x in elem:
                    p.append(x)
        # add rests to make sure all parts have the same length
        highest_time = 0
        for p in parts.values():
            highest_time = max(highest_time, p.highestTime)
        for p in parts.values():
            if p.highestTime < highest_time:
                p.append(music21.note.Rest(quarterLength=highest_time - p.highestTime))
        page_start_time = highest_time

    timesigs = {}

    lanes = {}
    for name, part in parts.items():
        lane = []
        lanes[name] = lane
        for elem in part.flatten():
            if isinstance(p, music21.meter.TimeSignature):
                k = (p.numerator, p.denominator)
                if k not in timesigs:
                    timesigs[k] = 1
                else:
                    timesigs[k] += 1
            lane += music21_to_wavtool(elem)

    if len(timesigs) == 0:
        timesig = (4, 4)
    else:
        timesig = max(timesigs, key=timesigs.get)

    return timesig, lanes


def process(cur, path, path_hash, timesig, lanes, descs):
    if len(lanes) == 0:
        return

    cur.execute(
        """
        insert into files (tempo, time_signature, path, hash, dataset_name)
                values (120, %s, %s, %s, 'imslp')
                on conflict (hash) do nothing returning id
        """,
        (f"{timesig[0]}/{timesig[1]}", str(path), path_hash),
    )

    rows = cur.fetchall()
    if len(rows) == 0:
        raise RuntimeError("File already in")

    file_id = rows[0][0]

    cur.execute(
        "insert into file_names (file_id, name, source) values (%s, %s, 'collect_imslp')",
        (file_id, str(path)),
    )

    for text in descs:
        cur.execute(
            "insert into file_descriptions (file_id, description, type) values (%s, %s, 'imslp_text') on conflict do nothing",
            (file_id, text),
        )

    cur.execute(
        "insert into instruments (file_id) values (%s) returning id", (file_id,)
    )
    inst_id = cur.fetchone()[0]

    for i, lane in enumerate(lanes):
        if len(lane) == 0:
            continue

        cur.execute(
            "insert into tracks (file_id, track_index) values (%s, %s) returning id",
            (file_id, i),
        )
        track_id = cur.fetchone()[0]

        cur.execute(
            "insert into lanes (track_id, instrument_id) values (%s, %s) returning id",
            (track_id, inst_id),
        )
        lane_id = cur.fetchone()[0]

        cur.executemany(
            'insert into notes (lane_id, pitch, "start", "end", velocity, off_velocity) values (%s, %s, %s, %s, 127, 127)',
            (
                (
                    lane_id,
                    note["pitch"],
                    note["start"],
                    note["end"],
                )
                for note in lane
            ),
        )


def get_descs_from_imslp_db(path_hash):
    with psycopg.connect("host=localhost dbname=imslp") as conn:
        with conn.cursor() as cur:
            cur.execute(
                "select composer, worktitle from pages where pdfhash = %s",
                (path_hash,),
            )
            xs = cur.fetchall()
            if xs is None:
                raise RuntimeError("File not found in imslp db")
            return [x for x in xs[0] if x is not None]


def process_insert_imslp(path, path_hash, timesig, lanes, descs):
    for i in range(100):
        try:
            with psycopg.connect(
                "host=localhost dbname=composer_new_dataset_v3"
            ) as conn:
                conn.autocommit = True
                with conn.transaction(), conn.cursor() as cur:
                    process(cur, path, path_hash, timesig, lanes, descs)
        except DeadlockDetected as e:
            if i < 99:
                print("retry on deadlock")
                continue
        break
    else:
        raise RuntimeError("insert retries exceeded")


def collect_and_insert(score_path):
    score_path = Path(score_path)
    with psycopg.connect(
        "host=localhost dbname=composer_new_dataset_v3"
    ) as conn:
        with conn.cursor() as cur:
            cur.execute('select 1 from files where hash = %s', (score_path.name,))
            if cur.fetchone() is not None:
                print(f"Skipping {score_path.name}")
                return
    descs = get_descs_from_imslp_db(score_path.name)
    timesig, lanes = collect(score_path)
    process_insert_imslp(score_path, score_path.name, timesig, lanes.values(), descs)
    print(f"Inserted {score_path.name}")


if __name__ == "__main__":
    import sys

    collect_and_insert(sys.argv[1])
