#!/usr/bin/env nix-shell
#!nix-shell -I nixpkgs=./. -i python3 -p "music-assistant.pythonPackages.python.withPackages (ps: music-assistant.dependencies ++ (with ps; [ jinja2 packaging ]))" -p nixfmt pyright ruff isort
import asyncio
import json
import os.path
import re
import sys
import tarfile
import tempfile
from dataclasses import dataclass, field
from functools import cache
from io import BytesIO
from pathlib import Path
from subprocess import check_output, run
from typing import Final, cast
from urllib.request import urlopen

from jinja2 import Environment
from mashumaro.exceptions import MissingField
from music_assistant_models.provider import ProviderManifest  # type: ignore
from packaging.requirements import Requirement

TEMPLATE = """# Do not edit manually, run ./update-providers.py

{
  version = "{{ version }}";
  builtins = [
{%- for builtin in builtins | sort %}
    "{{ builtin }}"
{%- endfor %}
  ];
  providers = {
{%- for provider in providers | sort(attribute='domain') %}
    {{ provider.domain }} = {% if provider.available %}ps: with ps;{% else %}ps:{% endif %} [
{%- for requirement in provider.available | sort %}
    {{ requirement }}
{%- endfor %}
    ]
{%- for requirement in provider.extra_list_deps | sort %}
    ++ {{ requirement }}
{%- endfor %}
;{% if provider.missing %} # missing {{ ", ".join(provider.missing) }}{% endif %}
{%- endfor %}
  };
}

"""


ROOT: Final = (
    check_output(
        [
            "git",
            "rev-parse",
            "--show-toplevel",
        ]
    )
    .decode()
    .strip()
)

PACKAGE_SET = "music-assistant.pythonPackages"
PACKAGE_MAP = {
    "git+https://github.com/MarvinSchenkel/pytube.git": "pytube",
}


EXTRA_DEPS = {
    # Those providers cannot guard pydantic behind TYPE_CHECKING
    "msx_bridge": ["pydantic"],
    "nicovideo": ["pydantic"],
    "ytmusic": [
        # https://github.com/music-assistant/server/blob/2.5.8/music_assistant/providers/ytmusic/__init__.py#L120
        "bgutil-ytdlp-pot-provider",
        "yt-dlp",
    ],
}


EXTRA_LIST_DEPS = {
    "sendspin": ["aiosendspin.optional-dependencies.server"],
}


def run_sync(cmd: list[str]) -> None:
    print(f"$ {' '.join(cmd)}")
    run(cmd, check=True)


async def check_async(cmd: list[str]) -> str:
    print(f"$ {' '.join(cmd)}")
    process = await asyncio.create_subprocess_exec(
        *cmd, stdout=asyncio.subprocess.PIPE, stderr=asyncio.subprocess.PIPE
    )
    stdout, stderr = await process.communicate()

    if process.returncode != 0:
        error = stderr.decode()
        raise RuntimeError(f"{cmd[0]} failed: {error}")

    return stdout.decode().strip()


class Nix:
    base_cmd: Final = [
        "nix",
        "--show-trace",
        "--extra-experimental-features",
        "nix-command",
    ]

    @classmethod
    async def _run(cls, args: list[str]) -> str | None:
        return await check_async(cls.base_cmd + args)

    @classmethod
    async def eval(cls, expr: str) -> list | dict | int | float | str | bool:
        response = await cls._run(["eval", "-f", f"{ROOT}/default.nix", "--json", expr])
        if response is None:
            raise RuntimeError("Nix eval expression returned no response")
        try:
            return json.loads(response)
        except (TypeError, ValueError):
            raise RuntimeError("Nix eval response could not be parsed from JSON")


async def get_provider_manifests(version: str = "master") -> list:
    manifests = []
    with tempfile.TemporaryDirectory() as tmp:
        with (
            urlopen(  # noqa: ASYNC210
                f"https://github.com/music-assistant/music-assistant/archive/refs/tags/{version}.tar.gz"
            ) as response,
            tarfile.open(fileobj=BytesIO(response.read())) as tar,
        ):
            tar.extractall(tmp, filter="data")

        basedir = Path(os.path.join(tmp, f"server-{version}"))
        sys.path.append(str(basedir))

        for fn in basedir.glob("**/providers/*/manifest.json"):
            if "_demo_" in str(fn):
                continue
            try:
                manifests.append(await ProviderManifest.parse(str(fn)))
            except MissingField as ex:
                print(f"Error parsing {fn}", ex)

    return manifests


@cache
def packageset_attributes():
    output = check_output(
        [
            "nix-env",
            "-f",
            ROOT,
            "-qa",
            "-A",
            "music-assistant.pythonPackages",
            "--arg",
            "config",
            "{ allowAliases = false; }",
            "--json",
        ]
    )
    return json.loads(output)


class TooManyMatches(Exception):
    pass


class NoMatch(Exception):
    pass


def resolve_package_attribute(package: str) -> str:
    pattern = re.compile(
        rf"^music-assistant\.pythonPackages\.{package}$", re.IGNORECASE
    )
    packages = packageset_attributes()
    matches = []
    for attr in packages:
        if pattern.match(attr):
            matches.append(attr.split(".")[-1])

    if len(matches) > 1:
        raise TooManyMatches(
            f"Too many matching attributes for {package}: {' '.join(matches)}"
        )
    if not matches:
        raise NoMatch(f"No matching attribute for {package}")

    return matches.pop()


async def get_package_version(package: str) -> str:
    version = cast(str, await Nix.eval(f"{PACKAGE_SET}.{package}.version"))
    return version


@dataclass
class Provider:
    domain: str
    available: list[str] = field(default_factory=list)
    missing: list[str] = field(default_factory=list)
    extra_list_deps: list[str] = field(default_factory=list)

    def __eq__(self, other):
        return self.domain == other.domain

    def __hash__(self):
        return hash(self.domain)


async def resolve_providers(manifests) -> tuple[set, set]:
    errors = []
    providers = set()
    for manifest in manifests:
        provider = Provider(manifest.domain)
        requirements = manifest.requirements
        for requirement in requirements:
            # allow substituting requirement specifications that packaging cannot parse
            if requirement in PACKAGE_MAP:
                requirement = PACKAGE_MAP[requirement]

            requirement = Requirement(requirement)
            try:
                attr = resolve_package_attribute(requirement.name)
                provider.available.append(attr)
            except TooManyMatches as ex:
                print(ex, file=sys.stderr)
                provider.missing.append(requirement.name)
                continue
            except NoMatch:
                provider.missing.append(requirement.name)
                continue

            version = await get_package_version(attr)
            if version not in requirement.specifier:
                errors.append(f"{requirement} not satisfied by version {version}")
        if manifest.domain in EXTRA_DEPS:
            for requirement in EXTRA_DEPS[manifest.domain]:
                provider.available.append(requirement)
        if manifest.domain in EXTRA_LIST_DEPS:
            for requirement in EXTRA_LIST_DEPS[manifest.domain]:
                provider.extra_list_deps.append(requirement)
        providers.add(provider)
    if errors:
        print("\n - ", end="")
        print("\n - ".join(errors))

    builtins = {manifest.domain for manifest in manifests if manifest.builtin}

    return providers, builtins


def render(outpath: str, version: str, providers: set, builtins: set):
    env = Environment()
    template = env.from_string(TEMPLATE)
    template.stream(version=version, providers=providers, builtins=builtins).dump(
        outpath
    )


async def main():
    version: str = cast(str, await Nix.eval("music-assistant.version"))
    manifests = await get_provider_manifests(version)
    providers, builtins = await resolve_providers(manifests)

    outpath = os.path.join(ROOT, "pkgs/by-name/mu/music-assistant/providers.nix")
    render(outpath, version, providers, builtins)

    run_sync(["nixfmt", outpath])


if __name__ == "__main__":
    run_sync(["pyright", __file__])
    run_sync(["ruff", "check", "--ignore=EXE005", __file__])
    run_sync(["isort", __file__])
    run_sync(["ruff", "format", __file__])
    asyncio.run(main())
