import logging
import re
import time
import uuid
from data_gen import Dataset, RawExample
from train_target import EMBED, WITH_TEXT_ONLY
import subprocess
import json
import sys
from data import wavtoolm2midi, wavtoolmm2midi
import random
import concurrent.futures
import itertools
import psycopg


class DatasetFromExternalTool(Dataset):
    def __init__(
        self, tool_cmdline, split_num_examples, tool_parallelism, high_watermark=1000
    ):
        super().__init__()
        self.tool_cmdline = tool_cmdline
        self.split_num_examples = {int(k): v for k, v in split_num_examples.items()}
        self.tool_parallelism = tool_parallelism
        self.high_watermark = high_watermark

    @classmethod
    def from_config(cls, common_config, my_config):
        return cls(
            my_config["tool_cmdline"],
            my_config["split_num_examples"],
            my_config["tool_parallelism"],
        )

    def num_examples(self):
        return dict(self.split_num_examples)

    def shuffle(self):
        pass

    def stream_examples_impl(self, split, ranksize, order):
        example_q = []
        example_futs = []
        rank, size = ranksize
        with concurrent.futures.ThreadPoolExecutor(
            max_workers=self.tool_parallelism
        ) as executor:
            for example_idx in (
                itertools.count()
                if split == 0
                else range(self.split_num_examples[split])
            ):
                while True:
                    # harvest futures
                    new_example_futs = []
                    # if empty, we must wait for at least one future to complete
                    if len(example_q) == 0 and len(example_futs) > 0:
                        concurrent.futures.wait(
                            example_futs, return_when=concurrent.futures.FIRST_COMPLETED
                        )
                    for fut in example_futs:
                        if fut.done():
                            example_q.extend(fut.result())
                        else:
                            new_example_futs.append(fut)
                    example_futs = new_example_futs

                    # sow futures
                    while (
                        len(example_q) < self.high_watermark
                        and len(example_futs) < self.tool_parallelism
                    ):
                        example_futs.append(executor.submit(self._get_more_examples))

                    if len(example_q) > 0:
                        break

                example = example_q.pop(0)
                if example_idx % size != rank:
                    continue
                yield RawExample(
                    id=example_idx,
                    descs=[example["prompt"]],
                    example=self.join_with_anchor(
                        wavtoolmm2midi(example["mainContext"]),
                        wavtoolmm2midi(example["mainContinuation"]),
                    ),
                    accomps=[
                        wavtoolmm2midi(accomp)
                        for accomp in random.sample(
                            example["accompaniments"],
                            random.randint(0, min(1, len(example["accompaniments"]))),
                        )
                    ],
                )

    def join_with_anchor(self, a, b):
        xs = a.copy()
        ys = b.copy()
        if len(ys) > 0:
            ys[0] = dict(ys[0])
            ys[0]["descriptionAnchor"] = True
        xs.extend(ys)
        return xs

    def _get_more_examples(self):
        p = subprocess.Popen(
            self.tool_cmdline,
            stdout=subprocess.PIPE,
            shell=True,
        )
        try:
            while p.poll() is not None:
                line = p.stdout.readline()
                if not line:
                    continue
                return json.loads(line)
            line = p.stdout.readline()
            if not line:
                return []
            return json.loads(line)
        except json.JSONDecodeError:
            print(
                f"Warning: failed to decode tool output JSON: {line}",
                file=sys.stderr,
            )
            return []


def _get_randomizer(cur, index_name):
    cur.execute(
        "SELECT indexdef FROM pg_indexes WHERE indexname = %s",
        (index_name,),
    )
    indexdef = cur.fetchone()
    if indexdef is None:
        raise RuntimeError("index not found")
    mat = re.match(r".*hashint4extended\(.*,(.*)\)\)$", indexdef[0])
    if mat is None:
        raise RuntimeError("index has wrong definition")
    return mat.group(1)


def _shuffle_randomizer(cur, table_name, index_name):
    cur.execute(f"DROP INDEX IF EXISTS {index_name}")
    seed = random.randint(-(2**63), 2**63 - 1)
    cur.execute(
        f"CREATE INDEX {index_name} ON {table_name} (hashint4extended(id, {seed}))"
    )
    cur.execute(f"ANALYZE {table_name}")


class DatasetFromPostgresExtractedClips(Dataset):
    def __init__(
        self,
        conn_str,
        eval_split_idx,
        train_target,
        min_len,
        max_len,
        aug_limit,
        permit_tags=None,  # do not restrict by tag TODO set in config!
        permit_datasets=None,  # all datasets are permitted
        drum_tags=None,  # tags that indicate drum tracks (None means no drums emitted)
    ):
        super().__init__()
        self.conn_str = conn_str
        self.eval_split_idx = eval_split_idx
        self.train_target = train_target
        self.min_len = min_len
        self.max_len = max_len
        self.aug_limit = aug_limit
        self.permit_tags = permit_tags
        self.permit_datasets = permit_datasets
        self.drum_tags = drum_tags

    @classmethod
    def from_config(cls, common_config, my_config):
        return cls(
            my_config["conn_str"],
            int(common_config["eval_split_idx"]),
            common_config["train_target"],
            int(common_config["seq_len_min"]),
            int(common_config["seq_len_max"]),
            int(my_config["aug_limit"]),
            my_config.get("permit_tags", None),
            my_config.get("permit_datasets", None),
            my_config.get("drum_tags", None),
        )

    def num_examples(self):
        with psycopg.connect(self.conn_str) as conn:
            with conn.cursor() as cur:
                cur.execute(
                    """
                    SELECT f.split, count(1) FROM extracted_clips c
                    JOIN files f ON c.file_id = f.id
                    WHERE c.symbolic_length IS NOT NULL AND c.symbolic_length >= %s AND c.symbolic_length <= %s
                    AND (
                        f.split != %s OR NOT EXISTS (
                            SELECT FROM automatic_lane_tags lt
                            JOIN tag_values tv ON lt.tag_value_id = tv.id
                            WHERE lt.lane_id = c.lane_id
                            AND tv.value = 'Contaminated Eval'
                        )
                    )
                    """
                    + (
                        """
                        AND EXISTS (
                            SELECT FROM file_descriptions fd
                            WHERE fd.file_id = f.id
                            AND fd.description IS NOT NULL
                            AND fd.description <> ''
                        )
                        """
                        if self.train_target == WITH_TEXT_ONLY
                        else ""
                    )
                    + (
                        " AND f.dataset_name = ANY(%s)"
                        if self.permit_datasets is not None
                        else ""
                    )
                    + (
                        """
                        AND EXISTS (
                            SELECT FROM automatic_extracted_clip_tags ct
                            JOIN tag_values tv ON ct.tag_value_id = tv.id
                            WHERE tv.value = ANY(%s)
                            AND ct.extracted_clip_id = c.id
                        )
                        """
                        if self.permit_tags is not None
                        else ""
                    )
                    + " GROUP BY f.split",
                    (
                        self.min_len,
                        self.max_len,
                        self.eval_split_idx,
                        *(
                            (self.permit_datasets,)
                            if self.permit_datasets is not None
                            else tuple()
                        ),
                        *(
                            (self.permit_tags,)
                            if self.permit_tags is not None
                            else tuple()
                        ),
                    ),
                )
                return {row[0]: row[1] for row in cur.fetchall()}

    def shuffle(self):
        with psycopg.connect(self.conn_str) as conn:
            conn.autocommit = True
            with conn.transaction(), conn.cursor() as cur:
                _shuffle_randomizer(cur, "extracted_clips", "extracted_clips_rand_idx")

    def _emit_drums(self, notes, tags):
        if any(tag in self.drum_tags for tag in tags):
            return [dict(n, note=n["note"] + 1000) for n in notes]
        return notes

    def _stream_rows(self, split, ranksize, order):
        with psycopg.connect(self.conn_str) as conn:
            with conn.cursor() as cur:
                cur.execute("SET cursor_tuple_fraction = 0.0001")  # dank
                randomizer = _get_randomizer(cur, "extracted_clips_rand_idx")
            with conn.cursor(name=f"fetch_{uuid.uuid4()}") as cur:
                start_time = time.time()
                cur.itersize = 100
                cur.execute(
                    """
                    SELECT
                        c.id,
                        coalesce(fd.descs, '[]'::jsonb) descs,
                        cn.notes example,
                        coalesce(et.tags, '[]'::jsonb) example_all_tags,
                        coalesce(ca.notess, '[]'::jsonb) accomp,
                        coalesce(ca.tagss, '[]'::jsonb) accomp_all_tagss,
                        c.file_id,
                        c.start
                    FROM extracted_clips c
                    JOIN public.files f ON c.file_id = f.id
                    JOIN extracted_clip_notes cn ON cn.extracted_clip_id = c.id
                    LEFT JOIN LATERAL (
                        SELECT jsonb_agg(
                            jsonb_build_object(
                                'description', description,
                                'type', type
                            ) ORDER BY description
                        ) descs
                        FROM file_descriptions fds
                        WHERE fds.file_id = f.id
                        AND description IS NOT NULL
                        AND description <> ''
                    ) fd ON true
                    LEFT JOIN LATERAL (
                        WITH all_tag_ids AS (
                            SELECT lt.tag_value_id FROM lane_tags lt
                            WHERE lt.lane_id = c.lane_id
                            UNION
                            SELECT ct.tag_value_id FROM automatic_extracted_clip_tags ct
                            WHERE ct.extracted_clip_id = c.id
                        )
                        SELECT jsonb_agg(tv.value) tags
                        FROM all_tag_ids atid JOIN tag_values tv ON atid.tag_value_id = tv.id
                    ) et ON true
                    LEFT JOIN LATERAL (
                        WITH cas AS (
                            SELECT can.notes, coalesce(cat.tags, '[]'::jsonb) tags
                            FROM extracted_clips ca
                            LEFT JOIN LATERAL (
                                WITH all_tag_ids AS (
                                    SELECT lt.tag_value_id FROM lane_tags lt
                                    WHERE lt.lane_id = ca.lane_id
                                    UNION
                                    SELECT ct.tag_value_id FROM automatic_extracted_clip_tags ct
                                    WHERE ct.extracted_clip_id = ca.id
                                )
                                SELECT jsonb_agg(tv.value) tags
                                FROM all_tag_ids atid JOIN tag_values tv ON atid.tag_value_id = tv.id
                            ) cat ON true
                            JOIN extracted_clip_notes can ON can.extracted_clip_id = ca.id
                            WHERE ca.file_id = c.file_id
                            AND ca.start = c.start
                            AND ca.lane_id <> c.lane_id
                            AND ca.symbolic_length >= 10
                            AND c.symbolic_length + ca.symbolic_length <= %s
                            """
                    + (
                        """
                            AND NOT EXISTS (
                                SELECT FROM automatic_lane_tags lt
                                JOIN tag_values tv ON lt.tag_value_id = tv.id
                                WHERE lt.lane_id = ca.lane_id AND tv.value = 'Contaminated Eval'
                            )
                        """
                        if split == self.eval_split_idx
                        else ""
                    )
                    + " ORDER BY "
                    + ("random()" if order == "random" else "ca.stripe")
                    + """
                            LIMIT %s
                        ) SELECT jsonb_agg(notes) notess, jsonb_agg(tags) tagss FROM cas
                    ) ca ON true
                    WHERE f.split = %s
                    AND c.stripe %% %s = %s
                    """
                    + (
                        "AND fd.descs IS NOT NULL"
                        if self.train_target == WITH_TEXT_ONLY
                        else ""
                    )
                    + (
                        " AND f.dataset_name = ANY(%s)"
                        if self.permit_datasets is not None
                        else ""
                    )
                    + (
                        """
                        AND EXISTS (
                            SELECT FROM automatic_extracted_clip_tags ct
                            JOIN tag_values tv ON ct.tag_value_id = tv.id
                            WHERE tv.value = ANY(%s)
                            AND ct.extracted_clip_id = c.id
                        )
                        """
                        if self.permit_tags is not None
                        else ""
                    )
                    + (
                        """
                        AND NOT EXISTS (
                            SELECT FROM automatic_lane_tags lt
                            JOIN tag_values tv ON lt.tag_value_id = tv.id
                            WHERE lt.lane_id = c.lane_id AND tv.value = 'Contaminated Eval'
                        )
                        """
                        if split == self.eval_split_idx
                        else ""
                    )
                    + """
                    AND c.symbolic_length >= %s AND c.symbolic_length <= %s
                    ORDER BY """
                    + (
                        f"hashint4extended(c.id, {randomizer})"
                        if order == "random"
                        else "c.stripe"
                    )
                    + " ASC",
                    (
                        self.max_len,
                        self.aug_limit,
                        split,
                        ranksize[1],
                        ranksize[0],
                        *(
                            (self.permit_datasets,)
                            if self.permit_datasets is not None
                            else tuple()
                        ),
                        *(
                            (self.permit_tags,)
                            if self.permit_tags is not None
                            else tuple()
                        ),
                        self.min_len,
                        self.max_len,
                    ),
                )
                for row in cur:
                    if start_time is not None:
                        logging.info(
                            f"split {split} time to first row {time.time() - start_time}s"
                        )
                        start_time = None
                    yield row

    def stream_examples_impl(self, split, ranksize, order):
        for row in self._stream_rows(split, ranksize, order):
            yield RawExample(
                id=row[0],
                descs=row[1],
                example=self._emit_drums(wavtoolm2midi(row[2]), row[3]),
                example_tags=row[3],
                accomps=[
                    self._emit_drums(wavtoolm2midi(accomp), accomp_tags)
                    for accomp, accomp_tags in zip(row[4], row[5])
                ],
                accomp_tagss=row[5],
            )


class DatasetFromPostgresExtractedClipsOverTime(DatasetFromPostgresExtractedClips):
    def __init__(self, *args, **kwargs):
        super().__init__(*args, **kwargs)
        self._num_examples = None

    def num_examples(self):
        if self._num_examples is not None:
            return self._num_examples
        with psycopg.connect(self.conn_str) as conn:
            with conn.cursor() as cur:
                cur.execute(
                    """
                    SELECT f.split, count(distinct (c.file_id, c.start)) FROM extracted_clips c
                    JOIN files f ON c.file_id = f.id
                    WHERE c.symbolic_length IS NOT NULL AND c.symbolic_length >= %s AND c.symbolic_length <= %s
                    AND (
                        f.split != %s OR NOT EXISTS (
                            SELECT FROM automatic_lane_tags lt
                            JOIN tag_values tv ON lt.tag_value_id = tv.id
                            WHERE lt.lane_id = c.lane_id
                            AND tv.value = 'Contaminated Eval'
                        )
                    )
                    """
                    + (
                        """
                        AND EXISTS (
                            SELECT FROM file_descriptions fd
                            WHERE fd.file_id = f.id
                            AND fd.description IS NOT NULL
                            AND fd.description <> ''
                        )
                        """
                        if self.train_target == WITH_TEXT_ONLY
                        else ""
                    )
                    + (
                        " AND f.dataset_name = ANY(%s)"
                        if self.permit_datasets is not None
                        else ""
                    )
                    + (
                        """
                        AND EXISTS (
                            SELECT FROM automatic_extracted_clip_tags ct
                            JOIN tag_values tv ON ct.tag_value_id = tv.id
                            WHERE tv.value = ANY(%s)
                            AND ct.extracted_clip_id = c.id
                        )
                        """
                        if self.permit_tags is not None
                        else ""
                    )
                    + " GROUP BY f.split",
                    (
                        self.min_len,
                        self.max_len,
                        self.eval_split_idx,
                        *(
                            (self.permit_datasets,)
                            if self.permit_datasets is not None
                            else tuple()
                        ),
                        *(
                            (self.permit_tags,)
                            if self.permit_tags is not None
                            else tuple()
                        ),
                    ),
                )
                self._num_examples = {row[0]: row[1] for row in cur.fetchall()}
                return self._num_examples

    def stream_examples_impl(self, split, ranksize, order):
        seen_file_starts = set()
        for row in self._stream_rows(split, ranksize, order):
            # only emit the first example for each (file, start) pair
            file_start = (row[6], row[7])
            if file_start in seen_file_starts:
                continue
            seen_file_starts.add(file_start)
            yield RawExample(
                id=row[0],
                descs=row[1],
                example=self._emit_drums(wavtoolm2midi(row[2]), row[3]),
                example_tags=row[3],
                accomps=[
                    self._emit_drums(wavtoolm2midi(accomp), accomp_tags)
                    for accomp, accomp_tags in zip(row[4], row[5])
                ],
                accomp_tagss=row[5],
            )


class DatasetFromPostgresLanes(Dataset):
    def __init__(
        self,
        conn_str,
        eval_split_idx,
        permit_datasets=None,  # all datasets are permitted
    ):
        super().__init__()
        self.conn_str = conn_str
        self.eval_split_idx = eval_split_idx
        self.permit_datasets = permit_datasets

    @classmethod
    def from_config(cls, common_config, my_config):
        return cls(
            my_config["conn_str"],
            common_config["eval_split_idx"],
            my_config.get("permit_datasets", None),
        )

    def num_examples(self):
        with psycopg.connect(self.conn_str) as conn:
            with conn.cursor() as cur:
                cur.execute(
                    """
                    SELECT f.split, count(1) FROM lanes l
                    JOIN instruments i ON i.id = l.instrument_id
                    JOIN files f ON i.file_id = f.id
                    WHERE (f.split != %s OR NOT EXISTS (
                        SELECT FROM automatic_lane_tags lt
                        JOIN tag_values tv ON lt.tag_value_id = tv.id
                        WHERE lt.lane_id = l.id
                        AND tv.value = 'Contaminated Eval'
                    ))
                    """
                    + (
                        " AND f.dataset_name = ANY(%s)"
                        if self.permit_datasets is not None
                        else ""
                    )
                    + " GROUP BY f.split",
                    (
                        self.eval_split_idx,
                        *(
                            (self.permit_datasets,)
                            if self.permit_datasets is not None
                            else tuple()
                        ),
                    ),
                )
                return {row[0]: row[1] for row in cur.fetchall()}

    def shuffle(self):
        with psycopg.connect(self.conn_str) as conn:
            conn.autocommit = True
            with conn.transaction(), conn.cursor() as cur:
                _shuffle_randomizer(cur, "lanes", "lanes_rand_idx")

    def stream_examples_impl(self, split, ranksize, order):
        with psycopg.connect(self.conn_str) as conn:
            with conn.cursor() as cur:
                cur.execute("SET cursor_tuple_fraction = 0.0001")  # dank
                randomizer = _get_randomizer(cur, "lanes_rand_idx")
            with conn.cursor(name=f"fetch_{uuid.uuid4()}") as cur:
                start_time = time.time()
                cur.itersize = 100
                cur.execute(
                    """
                    SELECT l.id, a.notes
                    FROM lanes l
                    JOIN LATERAL (
                        SELECT jsonb_agg(jsonb_build_object('pitch', n.pitch, 'start', n.start, 'end', n."end", 'velocity', n.velocity)) AS notes
                        FROM notes n
                        WHERE n.lane_id = l.id
                    ) a ON true
                    JOIN instruments i ON i.id = l.instrument_id
                    JOIN files f ON i.file_id = f.id
                    """
                    + (
                        "WHERE f.split = %s AND l.stripe %% %s = %s "
                        if split is not None
                        else "WHERE l.stripe %% %s = %s "
                    )
                    + (
                        """
                        AND NOT EXISTS (
                            SELECT FROM automatic_lane_tags lt
                            JOIN tag_values tv ON lt.tag_value_id = tv.id
                            WHERE lt.lane_id = l.id AND tv.value = 'Contaminated Eval'
                        )
                        """
                        if split == self.eval_split_idx
                        else ""
                    )
                    + (
                        " AND f.dataset_name = ANY(%s)"
                        if self.permit_datasets is not None
                        else ""
                    )
                    + " ORDER BY "
                    + (
                        f"hashint4extended(l.id, {randomizer})"
                        if order == "random"
                        else "l.stripe"
                    )
                    + " ASC",
                    (
                        *((split,) if split is not None else tuple()),
                        ranksize[1],
                        ranksize[0],
                        *(
                            (self.permit_datasets,)
                            if self.permit_datasets is not None
                            else tuple()
                        ),
                    ),
                )
                for row in cur:
                    if start_time is not None:
                        logging.info(
                            f"split {split} time to first row {time.time() - start_time}s"
                        )
                        start_time = None
                    yield RawExample(
                        id=row[0],
                        descs=[],
                        example=wavtoolm2midi(row[1]),
                        accomps=[],
                    )


# d = DatasetFromPostgres(
#     "host=localhost dbname=composer_new_dataset_v3", WITH_TEXT_ONLY, 30, 300, 1
# )
# print(d.num_examples())
# for e, i in zip(d.stream_examples(0), range(10)):
#     pprint(e)
