import json
from pprint import pprint
import psycopg
from psycopg_pool import ConnectionPool
from pathlib import Path
from hashlib import md5
from binascii import hexlify
import sqlite3
import re
import os

dbpool = ConnectionPool('host=localhost dbname=composer_new_dataset_v3', min_size=2, max_size=2)

def load_by_basename(cur, dataset_name, source, gen):
    for basename, names in gen:
        cur.execute('select file_id from file_names where basename = %s and dataset_name = %s',
                    (basename, dataset_name))
        rows = cur.fetchall()
        if len(rows) == 0:
            continue
        file_id = rows[0][0]
        cur.execute('insert into file_names (file_id, source, name) select %s, %s, unnest(%s::text[]) on conflict (file_id, name) do nothing',
                    (file_id, source, names))


def load_lmd_paths(cur):
    with open('lmd/md5_to_paths.json', 'r') as f:
        load_by_basename(cur, 'lmd_full', 'lmd_paths',
                         ((f'{k}.mid', v) for k, v in json.load(f).items()))

def load_lmd_clean(cur):
    def gen():
        for path in Path('lmd/clean_midi').rglob('*'):
            if not re.match('.*\.midi?$', str(path), flags=re.IGNORECASE) or not os.path.isfile(path):
                continue
            with open(path, 'rb') as f:
                data = f.read()
            lmd_full_name = hexlify(md5(data).digest()).decode('utf-8') + '.mid'
            clean_name = '/'.join(str(path).split('/')[2:])
            yield lmd_full_name, [clean_name]
    load_by_basename(cur, 'lmd_full', 'lmd_clean', gen())

def load_lmd_matched(cur):
    with sqlite3.connect('lmd/track_metadata.db') as scon:
        scur = scon.cursor()
        def gen():
            for path in Path('lmd/lmd_matched').rglob('*'):
                mat = re.match('.*/([^/]+)/([^/]+\\.mid)$', str(path))
                if not mat:
                    continue
                scur.execute('select artist_name, title from songs where track_id = ?', (mat[1],))
                res = scur.fetchall()
                if len(res) == 0:
                    continue
                row = res[0]
                yield mat[2], [f'{row[0]} - {row[1]}']
        load_by_basename(cur, 'lmd_full', 'lmd_matched', gen())

def load_bitmidi_names(cur):
    with open('bitmidi/urls.json', 'r') as f:
        load_by_basename(cur, 'bitmidi', 'bitmidi_names',
                         ((x['downloadUrl'].split('/')[-1], [x['name']]) for x in json.load(f)))

with dbpool.connection() as conn:
    conn.autocommit = True
    for func in [
        load_lmd_paths,
        load_lmd_clean,
        load_bitmidi_names,
        load_lmd_matched,
    ]:
        with conn.transaction(), conn.cursor() as cur:
            func(cur)
