from dataclasses import dataclass
from typing import Any, Iterable, cast

from pants.backend.python.subsystems.python_tool_base import PythonToolBase
from pants.backend.python.target_types import (
    EntryPoint,
    InterpreterConstraintsField,
    PythonResolveField,
    PythonTestsBatchCompatibilityTagField,
    PythonTestsExtraEnvVarsField,
    PythonTestSourceField,
    PythonTestsTimeoutField,
    PythonTestsXdistConcurrencyField,
    RuntimePackageDependenciesField,
    SkipPythonTestsField,
)
from pants.backend.python.util_rules.pex import Pex, PexRequest, VenvPex, VenvPexProcess
from pants.backend.python.util_rules.pex_from_targets import RequirementsPexRequest
from pants.backend.python.util_rules.python_sources import (
    PythonSourceFiles,
    PythonSourceFilesRequest,
)
from pants.core.goals.test import TestFieldSet, TestRequest, TestResult, TestSubsystem
from pants.core.util_rules.environments import EnvironmentField
from pants.engine.process import Process, ProcessResultWithRetries, ProcessWithRetries
from pants.engine.rules import (
    Get,
    MultiGet,
    Rule,
    collect_rules,
    rule,  # type: ignore
)
from pants.engine.target import (
    COMMON_TARGET_FIELDS,
    Dependencies,
    Target,
    TransitiveTargets,
    TransitiveTargetsRequest,
)
from pants.option.global_options import GlobalOptions
from pants.option.option_types import SkipOption


@dataclass(frozen=True)
class NoseTestFieldSet(TestFieldSet):
    required_fields = (PythonTestSourceField,)

    source: PythonTestSourceField
    interpreter_constraints: InterpreterConstraintsField
    timeout: PythonTestsTimeoutField
    runtime_package_dependencies: RuntimePackageDependenciesField
    extra_env_vars: PythonTestsExtraEnvVarsField
    resolve: PythonResolveField
    environment: EnvironmentField
    xdist_concurrency: PythonTestsXdistConcurrencyField
    batch_compatibility_tag: PythonTestsBatchCompatibilityTagField

    @classmethod
    def opt_out(cls, tgt: Target) -> bool:
        return tgt.get(SkipPythonTestsField).value


class NoseTest(Target):
    alias = "nose_test"
    help = "Run Nose (Python) tests."
    core_fields = (
        *COMMON_TARGET_FIELDS,
        Dependencies,
        PythonTestSourceField,
        PythonResolveField,
    )


class NoseTool(PythonToolBase):
    name = "Nose"
    options_scope = "nose"
    help = "The Nose test runner (and framework) for Python (https://nose.readthedocs.io/en/latest/)."

    register_interpreter_constraints = True

    default_requirements = ["pynose==1.5.1"]
    default_main = EntryPoint.parse("nose.core:console_main")
    default_lockfile_resource = ("nose", "nose.lock")

    skip = SkipOption("test")


@dataclass(frozen=True)
class NoseRequest(TestRequest):
    tool_subsystem = NoseTool
    field_set_type = NoseTestFieldSet


@rule(desc="Run Python test(s) with Nose")  # type: ignore
async def run_nose_test(
    batch: NoseRequest.Batch[NoseTestFieldSet, Any],
    nose: NoseTool,
    test_subsystem: TestSubsystem,
    global_options: GlobalOptions,
) -> TestResult:
    transitive_targets = await Get(
        TransitiveTargets, TransitiveTargetsRequest((batch.single_element.address,))
    )
    requirements_pex_request = Get(
        Pex, RequirementsPexRequest(tgt.address for tgt in transitive_targets.closure)
    )

    nose_pex_request = Get(Pex, PexRequest, nose.to_pex_request())
    sources_request = Get(
        PythonSourceFiles,
        PythonSourceFilesRequest(transitive_targets.closure, include_files=True),
    )
    nose_pex, requirements_pex, sources = await MultiGet(
        nose_pex_request, requirements_pex_request, sources_request
    )

    nose_runner_pex = await Get(
        VenvPex,
        PexRequest(
            output_filename="nose_runner.pex",
            main=nose.main,
            internal_only=True,
            pex_path=[nose_pex, requirements_pex],
        ),
    )

    process = await Get(
        Process,
        VenvPexProcess(
            nose_runner_pex,
            description=f"Run Nose for {batch.single_element.address}",
            argv=("--verbose", "--nocapture", *sources.source_files.files),
        ),
    )
    results = await Get(
        ProcessResultWithRetries,
        ProcessWithRetries(process, test_subsystem.attempts_default),
    )

    return TestResult.from_batched_fallible_process_result(
        results.results,
        batch=batch,
        output_setting=test_subsystem.output,
        output_simplifier=global_options.output_simplifier(),
    )


def target_types() -> Iterable[type[Target]]:
    return [NoseTest]


def rules() -> Iterable[Rule]:
    return [
        *collect_rules(),
        *cast(Iterable[Rule], NoseRequest.rules()),  # type: ignore
        *NoseTool.rules(),
    ]
