"""Lyrics copyright infringement detection through minhash"""

import logging
import os
import time
import pickle
import re
import string
import json

import modal
from suno_utils.utils.s3 import _download_s3_file
from suno_utils.worker.loader import S3Loader
from suno_utils.worker.modal_base import MODAL_MOUNTS, get_modal_base_image
from suno_utils.worker.settings import s3_client
from suno_utils.worker.utils import retry_decorator
from suno_utils.gpt import chirp_v2
from suno_utils.gpt import chirp_v2
from suno_utils.harvest.youtube.constants.base import NON_WHITESPACE_LANGS
from suno_utils.utils.text import normalize_whitespace
from suno_utils.harvest.youtube.language_classify import get_text_lang_p
from suno_utils.worker.schema import QueueItem
from collections import Counter, defaultdict

logger = logging.getLogger(__name__)
logging.basicConfig()
logger.setLevel(logging.INFO)

############## CHANGE THESE ##############
DEPLOYMENT_TYPE = "dev"

##########################################
DEPLOYMENT_TYPES = {"dev", "prod"}
assert DEPLOYMENT_TYPE in DEPLOYMENT_TYPES
APP_NAME = f"lyrics-check-reverse-index-{DEPLOYMENT_TYPE}"
assert APP_NAME.endswith(DEPLOYMENT_TYPE)
MOUNT_PATH = "/suno/models"
UPLOADS_S3_BUCKET = "suno-data-uploads"
INVERTED_INDEX_NAME = "inverted_index_v1"
INVERTED_INDEX_PATH = os.path.join(MOUNT_PATH, f"{INVERTED_INDEX_NAME}.pickle")

retry_s3_download = retry_decorator(3, wait_seconds=20)(s3_client.download_fileobj)

N_CPU = 2
aws_secret = modal.Secret.from_name("studio-aws")
DD_SAMPLE_RATE = "1" if DEPLOYMENT_TYPE == "dev" else "0.0001"
SECRETS = [
    aws_secret,
    modal.Secret.from_dict(
        {
            "SUNO_ASSETS_PATH": "/suno/models/assets",
            "XDG_CACHE_HOME": "/suno/models/",
        },
    ),
    modal.Secret.from_name("openai-secret"),
    modal.Secret.from_dict(
        {
            "DD_SITE": "datadoghq.com",
            "DD_ENV": DEPLOYMENT_TYPE,
            "DD_SERVICE": "chatgpt-worker",
            "DD_LOGS_ENABLED": "true",
            "DD_TRACE_ENABLED": "true",
            "DD_TRACE_SAMPLE_RATE": DD_SAMPLE_RATE,
        },
    ),
    modal.Secret.from_name("datadog-metrics"),
    modal.Secret.from_name("api-callback-token"),
]
# for soundclash
WHITELISTED_DOCS = ["doc_1369399", "doc_1892304", "doc_1862978", "doc_8681331"] + [
    "doc_9028903",
    "doc_4340105",
    "doc_5384282",
    "doc_6656167",
    "doc_1797383",
    "doc_7182424",
    "doc_9100110",
    "doc_6217677",
    "doc_6680769",
    "doc_5017032",
    "doc_4173209",
]


class InvertedIndexForLyrics:
    def __init__(self, k=15, num_buckets=2**18):
        self.lyrics = []
        self.inverted_index = defaultdict(list)
        self.k = k
        self.num_buckets = num_buckets
        self.text_lang_model = chirp_v2.text_lang_model

    def preprocess(self, text):
        return self.clean_text(text)

    def create_shingles(self, text):
        if len(text) < self.k:
            return [text]
        return [text[i : i + self.k] for i in range(len(text) - self.k + 1)]

    def hash_shingle(self, shingle):
        import mmh3

        return mmh3.hash(shingle) % self.num_buckets

    def add_lyrics(self, lyric):
        doc_id = len(self.lyrics)
        processed_text = self.preprocess(lyric)
        shingles = self.create_shingles(processed_text)
        hashed_shingles = [self.hash_shingle(shingle) for shingle in shingles]
        self.lyrics.append(lyric)

        for position, hashed_shingle in enumerate(hashed_shingles):
            self.inverted_index[hashed_shingle].append((doc_id, position))

    def simple_query(self, query):
        processed_query = self.preprocess(query)
        shingles = self.create_shingles(processed_query)
        query_hashed_shingles = [self.hash_shingle(shingle) for shingle in shingles]
        query_tf = Counter(query_hashed_shingles)
        candidate = Counter()

        for hashed_term in query_tf:
            # print(hashed_term, self.inverted_index.get(hashed_term, []))
            candidate_pair = self.inverted_index.get(hashed_term, [])
            candidate.update([c[0] for c in candidate_pair])
        candidate = {k: v / len(query_hashed_shingles) for k, v in candidate.items()}

        candidate = {
            f"pdoc_{k}": v for k, v in candidate.items() if k not in WHITELISTED_DOCS and v > 0.9
        }

        return sorted(candidate.items(), key=lambda x: x[1], reverse=True)[:20]

    def dump(self, filename):
        """
        Dump the inverted index and lyrics to a file using JSON.

        :param filename: The name of the file to save the data to.
        """
        data = {
            "lyrics": self.lyrics,
            "inverted_index": {str(k): v for k, v in self.inverted_index.items()},
            "k": self.k,
            "num_buckets": self.num_buckets,
        }
        with open(filename, "w", encoding="utf-8") as f:
            json.dump(data, f, ensure_ascii=False, indent=2)

    def dump_to_pickle(self, filename):
        data = {
            "lyrics": self.lyrics,
            "inverted_index": {str(k): v for k, v in self.inverted_index.items()},
            "k": self.k,
            "num_buckets": self.num_buckets,
        }
        with open(filename, "wb") as f:
            pickle.dump(data, f)

    @classmethod
    def load(cls, filename):
        """
        Load the inverted index and lyrics from a file using JSON.

        :param filename: The name of the file to load the data from.
        :return: An instance of InvertedIndexForLyrics with the loaded data.
        """
        with open(filename, "r", encoding="utf-8") as f:
            data = json.load(f)

        instance = cls(k=data["k"], num_buckets=data["num_buckets"])
        instance.lyrics = data["lyrics"]
        instance.inverted_index = defaultdict(
            list, {int(k): v for k, v in data["inverted_index"].items()}
        )

        return instance

    @classmethod
    def load_from_pickle(cls, filename):
        with open(filename, "rb") as f:
            data = pickle.load(f)

        print(f"Loading from Pickle:Inverted index size: {len(data['inverted_index'])}")

        instance = cls(k=data["k"], num_buckets=data["num_buckets"])
        instance.lyrics = data["lyrics"]
        instance.inverted_index = defaultdict(
            list, {int(k): v for k, v in data["inverted_index"].items()}
        )
        return instance

    def clean_text(self, text: str, preserve_whitespace=True) -> str:
        text = text.lower()
        # strip special tags
        text = re.sub(
            r"""
            \(          # Match an opening parenthesis
            \s*         # Match zero or more whitespace characters
            x           # Match the letter 'x'
            \s*         # Match zero or more whitespace characters
            [0-9]+      # Match one or more digits
            \s*         # Match zero or more whitespace characters
            \)          # Match a closing parenthesis
            """,
            " ",  # Replace with a single space
            text,
            flags=re.VERBOSE,
        )
        text = re.sub(
            r"""
            \[          # Match an opening square bracket
            .{1,40}     # Match any character (except newline) 1 to 40 times
            \]          # Match a closing square bracket
            """,
            " ",  # Replace with a single space
            text,
            flags=re.VERBOSE,
        )
        # replace the punctuation with spaces
        trans_table = str.maketrans(string.punctuation, " " * len(string.punctuation))
        text = text.translate(trans_table)
        # do basic clean
        text = normalize_whitespace(text)

        p_lang = get_text_lang_p(self.text_lang_model, text) if self.text_lang_model else {}
        lang_guess, p_lang_guess = self.get_most_likely_language(p_lang)
        # if english do more clean
        if p_lang.get("en", 0) > 0.9:
            text = re.sub(r"[^a-z]", " ", text)
        # if non-whitespace language then artificially make words
        text = normalize_whitespace(text)
        if len(text) == 0:
            return ""
        if (p_lang_guess >= 0.6 and lang_guess in NON_WHITESPACE_LANGS) or (
            text.count(" ") / len(text) == 0.05 and len(text) > 100
        ):
            text = " ".join(list(text))
        # finalize
        text = normalize_whitespace(text)
        if not preserve_whitespace:
            text = re.sub(r"\s", "", text)
        return text

    def get_most_likely_language(self, p_lang: dict[str, float]) -> tuple[str, float]:
        if p_lang:
            # Sort the language-probability pairs by probability
            sorted_langs = sorted(p_lang.items(), key=lambda x: x[1], reverse=True)
            return sorted_langs[0]
        else:
            return ("", 0)


class LyricsReverseIndexWorker(S3Loader):
    """Detect lyrics copyright infringement through minhash"""

    def __init__(
        self,
        n_matches: int = 1,
        version_name: str = "v2-large",
    ):
        self.n_matches = n_matches
        self.popular_reverse_index = None
        self.version_name = version_name

    def preload(self):
        ckpt_path = chirp_v2._get_model_if_needed(chirp_v2.FASTTEXT_CKPT_PATH, cache_dir=MOUNT_PATH)
        chirp_v2.load_fasttext_model(ckpt_path)
        start_time = time.time()
        self.popular_inverted_index = InvertedIndexForLyrics.load_from_pickle(INVERTED_INDEX_PATH)
        finish_time = time.time()
        print(f"Preloading inverted index took {finish_time - start_time}s")
        print(f"Inverted index size: {len(self.popular_inverted_index.inverted_index)}")

    def query_lyrics(self, lyrics: str) -> tuple[bool, str, str]:
        result = []
        if self.popular_inverted_index is not None:
            result = self.popular_inverted_index.simple_query(lyrics)
        if len(result) >= self.n_matches:
            return True, str(result), self.version_name
        return False, str(result), self.version_name

    def query_lyrics_with_details(self, lyrics: str, threshold: float = 0.9) -> dict:
        result = []
        if self.popular_inverted_index is not None:
            partial_match_result = self.popular_inverted_index.simple_query(lyrics)
            result = result + partial_match_result
        return {"result": result, "version": self.version_name, "model": "minhash"}

    def is_lyrics_copyright_infringement_with_popular_index(self, lyrics: str) -> dict:
        if self.popular_inverted_index is None:
            print("Popular inverted index not initialized")
            return {
                "result": [],
                "version": self.version_name,
                "error": "Popular inverted index not initialized",
            }
        result = self.popular_inverted_index.simple_query(lyrics)
        return {"result": result, "version": self.version_name, "model": "inverted_index"}


def download_minhash_wrapper() -> None:
    start_time = time.time()
    target_inverted_index_path = os.path.join(MOUNT_PATH, f"{INVERTED_INDEX_NAME}.pickle")
    print("Downloading inverted index files")
    _download_s3_file(
        f"s3://suno-data/ashe/trained_models/{INVERTED_INDEX_NAME}.pickle",
        target_inverted_index_path,
    )
    finish_time = time.time()
    print(f"Downloaded inverted index files successfully in {finish_time - start_time}s")


base_image = get_modal_base_image().pip_install("datasketch").pip_install("mmh3")
image = base_image.run_function(download_minhash_wrapper, secrets=SECRETS)
app = modal.App(APP_NAME, image=image)


@app.cls(
    cpu=N_CPU,
    secrets=SECRETS,
    timeout=4000,
    scaledown_window=1200,
    mounts=MODAL_MOUNTS,
    retries=modal.Retries(
        max_retries=2,
        backoff_coefficient=2.0,
        initial_delay=5.0,
    ),
    memory=250000,
    min_containers=25,
    allow_concurrent_inputs=60,
)
class LyricsReverseIndexStub:
    """App for detecting lyrics copyright infringement"""

    def __init__(self):
        """Set up LyricsCopyrightDetectApp."""
        self.worker = LyricsReverseIndexWorker()
        self.worker.preload()

    @modal.method()
    def query_lyrics(self, queue_item: str):
        item = QueueItem(**json.loads(queue_item))
        lyrics = item.prompt_text or ""
        if len(lyrics) == 0:
            return False
        result = self.worker.query_lyrics(lyrics)
        if item.callback_url:
            self.worker.notify_finish(
                QueueItem(id=item.id, metadata={}, callback_url=item.callback_url),
                {
                    "id": item.id,
                    "clip_ids": item.ids,
                    "copyright_infringement": result[0],
                    "copyright_infringement_details": result[1],
                    "model_name": result[2],
                },
            )
        return result[0]

    @modal.method()
    def query_lyrics_with_details(self, lyrics: str, threshold: float = 0.9):
        if len(lyrics) == 0:
            return None
        result = self.worker.query_lyrics_with_details(lyrics, threshold)
        return result


@app.local_entrypoint()
def main():
    testApp = LyricsReverseIndexStub()
    test_data = {
        "id": "123",
        "prompt_text": "my lyrics",
        "metadata": {"tags": "r&b"},
    }
    result = testApp.query_lyrics.remote(json.dumps(test_data))
    print(f"Test1: lyrics copyright infringement result: {result}")

    existing_lyrics = "it might seem crazy what i am about to say sunshine she s here you can take a break i m a hot air balloon that could go to space with the air like i don t care baby by the way because i m happy clap along if you feel like a room without a roof because i m happy clap along if you feel like happiness is the truth because i m happy clap along if you know what happiness is to you because i m happy clap along if you feel like that s what you wanna do here come bad news talking this and that yeah give me all you got don t hold back yeah well i should probably warn you i ll be just fine yeah no offense to you don t waste your time here s why because i m happy clap along if you feel like a room without a roof because i m happy clap along if you feel like happiness is the truth because i m happy clap along if you know what happiness is to you because i m happy clap along if you feel like that s what you wanna do happy bring me down can t nothin happy bring me down my love is to high happy bring me down can t nothin happy bring me down let me tell you now happy happy happy happy bring me down can t nothin happy happy happy happy bring me down my love is too high happy happy happy happy bring me down can t nothin happy happy happy happy bring me down i said because i m happy clap along if you feel like a room without a roof because i m happy clap along if you feel like happiness is the truth because i m happy clap along if you know what happiness is to you because i m happy clap along if you feel like that s what you wanna do because i m happy clap along if you feel like a room without a roof because i m happy clap along if you feel like happiness is the truth because i m happy clap along if you know what happiness is to you because i m happy clap along if you feel like that s what you wanna do happy happy happy happy bring me down can t nothin happy happy happy happy bring me down my love is too high happy happy happy happy bring me down can t nothin happy happy happy happy bring me down i said because i m happy clap along if you feel like a room without a roof because i m happy clap along if you feel like happiness is the truth because i m happy clap along if you know what happiness is to you because i m happy clap along if you feel like that s what you wanna do because i m happy clap along if you feel like a room without a roof because i m happy clap along if you feel like happiness is the truth because i m happy clap along if you know what happiness is to you because i m happy clap along if you feel like that s what you wanna do"
    test_existing_data = {
        "id": "456",
        "prompt_text": existing_lyrics,
        "metadata": {"tags": "jazz"},
    }
    result = testApp.query_lyrics.remote(json.dumps(test_existing_data))
    print(f"Test2: positive lyrics copyright infringement result: {result}")

    result = testApp.query_lyrics_with_details.remote(json.dumps(test_existing_data))
    print(f"Test2: positive lyrics copyright infringement detailed result: {result}")

    existing_lyrics_short = """
        [Verse]
        I've created a monster
        'Cause nobody wants to see Marshall no more, they want Shady, I'm chopped liver
        Well, if you want Shady, this is what I'll give ya
        A little bit of weed mixed with some hard liquor
        Some vodka that'll jump-start my heart quicker
        Than a shock when I get shocked at the hospital
        By the doctor when I'm not cooperating
        When I'm rockin' the table while he's operating (Hey!)
        You waited this long, now stop debating
        'Cause I'm back, I'm on the rag and ovulating
        I know that you got a job, Ms. Cheney
        But your husband's heart problem's complicating
        So the FCC won't let me be
        Or let me be me, so let me see
        They tried to shut me down on MTV
        But it feels so empty without me
        So come on and dip, bum on your lips
        Fuck that, cum on your lips and some on your tits
        And get ready, 'cause this shit's about to get heavy
        I just settled all my lawsuits (Fuck you, Debbie!)
    """

    result = testApp.query_lyrics_with_details.remote(existing_lyrics_short)
    print(f"Test3: positive lyrics copyright infringement detailed result: {result}")

    test_lyrics = """[Verse]
    一盏离愁孤灯伫立在窗口
    我在门后假装你人还没走
    旧地如重游月圆更寂寞
    夜半清醒的烛火不忍苛责我

    [Verse]
    一壶漂泊浪迹天涯难入喉
    你走之后酒暖回忆思念瘦
    水向东流时间怎么偷
    花开就一次成熟我却错过

    [Chorus]
    谁在用琵琶弹奏一曲东风破
    岁月在墙上剥落看见小时候
    犹记得那年我们都还很年幼
    而如今琴声幽幽我的等候你没听过"""

    result = testApp.query_lyrics_with_details.remote(test_lyrics)
    print(f"Test4: positive lyrics copyright infringement detailed result: {result}")
