"""Static analysis integration (PMD) for Apex/trigger review."""

from __future__ import annotations

import asyncio
import json
import xml.etree.ElementTree as ET
from pathlib import Path
from typing import List, Optional, Protocol

from pydantic import BaseModel, Field

from packages.config import Settings, get_settings


class StaticViolation(BaseModel):
    rule: str
    file: str
    line: Optional[int] = None
    message: str
    severity: str = "medium"


class StaticAnalysisResult(BaseModel):
    success: bool
    violations: List[StaticViolation] = Field(default_factory=list)
    message: str = ""


class StaticAnalyzer(Protocol):
    async def analyze(self, *, source_dir: Path) -> StaticAnalysisResult: ...


class MockStaticAnalyzer:
    """No-op analyzer for local/dev and CI."""

    def __init__(self) -> None:
        self.calls: List[str] = []

    async def analyze(self, *, source_dir: Path) -> StaticAnalysisResult:
        self.calls.append(str(source_dir))
        return StaticAnalysisResult(
            success=True,
            violations=[],
            message="Mock static analysis passed",
        )


class PmdStaticAnalyzer:
    """Run PMD CLI against Apex classes and triggers under force-app."""

    def __init__(self, settings: Optional[Settings] = None) -> None:
        self.settings = settings or get_settings()

    async def analyze(self, *, source_dir: Path) -> StaticAnalysisResult:
        if not source_dir.exists():
            return StaticAnalysisResult(
                success=True,
                violations=[],
                message="No force-app source dir — static analysis skipped",
            )

        apex_files = list(source_dir.rglob("*.cls")) + list(source_dir.rglob("*.trigger"))
        if not apex_files:
            return StaticAnalysisResult(
                success=True,
                violations=[],
                message="No Apex files — static analysis skipped",
            )

        pmd_bin = (self.settings.pmd_bin or "pmd").strip()
        violations: List[StaticViolation] = []
        for apex_file in apex_files:
            rel = apex_file.relative_to(source_dir.parent.parent.parent)
            cmd = [
                pmd_bin,
                "check",
                "-d",
                str(apex_file),
                "-R",
                "rulesets/apex/quickstart.xml",
                "-f",
                "xml",
            ]
            proc = await asyncio.create_subprocess_exec(
                *cmd,
                stdout=asyncio.subprocess.PIPE,
                stderr=asyncio.subprocess.PIPE,
            )
            stdout_b, stderr_b = await proc.communicate()
            stdout = (stdout_b or b"").decode("utf-8", errors="replace")
            stderr = (stderr_b or b"").decode("utf-8", errors="replace")
            if proc.returncode not in (0, 4):
                return StaticAnalysisResult(
                    success=False,
                    violations=[],
                    message=f"PMD failed ({proc.returncode}): {stderr[:500]}",
                )
            if stdout.strip():
                violations.extend(_parse_pmd_xml(stdout, str(rel)))

        return StaticAnalysisResult(
            success=True,
            violations=violations,
            message=f"PMD found {len(violations)} violation(s)",
        )


def _parse_pmd_xml(xml_text: str, default_file: str) -> List[StaticViolation]:
    violations: List[StaticViolation] = []
    try:
        root = ET.fromstring(xml_text)
    except ET.ParseError:
        return violations
    ns = {"pmd": "http://pmd.sourceforge.net/report/2.0.0"}
    for file_el in root.findall(".//pmd:file", ns) or root.findall(".//file"):
        file_path = file_el.get("name") or default_file
        for v in file_el.findall("pmd:violation", ns) or file_el.findall("violation"):
            line_raw = v.get("beginline") or v.get("line")
            violations.append(
                StaticViolation(
                    rule=v.get("rule") or "PMD",
                    file=file_path,
                    line=int(line_raw) if line_raw and line_raw.isdigit() else None,
                    message=(v.text or "").strip(),
                    severity=(v.get("priority") or "medium"),
                )
            )
    return violations


_mock_singleton: Optional[MockStaticAnalyzer] = None


def get_static_analyzer(settings: Optional[Settings] = None) -> StaticAnalyzer:
    global _mock_singleton
    settings = settings or get_settings()
    if settings.static_analysis_provider == "pmd":
        return PmdStaticAnalyzer(settings)
    if _mock_singleton is None:
        _mock_singleton = MockStaticAnalyzer()
    return _mock_singleton


def reset_mock_static_analyzer() -> MockStaticAnalyzer:
    global _mock_singleton
    _mock_singleton = MockStaticAnalyzer()
    return _mock_singleton
