from __future__ import annotations
from typing import Iterable, List, Sequence, Union, Optional

import requests

from ..localization.helpers import create_translated_finding
from ..model import Finding, ScanRequest

SQL_ERROR_SIGNATURES = [
    "you have an error in your sql syntax;",
    "warning: mysql",
    "unclosed quotation mark after the character string",
    "quoted string not properly terminated",
    "psql:",
    "sqlite error",
    "pg_query",
]

REFLECTION_MARKER = "<web-sentinel-xss>"


def evaluate_injection_surface(request: ScanRequest) -> Iterable[Finding]:
    if not request.allow_invasive:
        return [
            create_translated_finding(
                check="injection",
                i18n_key="injection.disabled",
                severity="info",
            )
        ]

    payloads = {
        "__sentinel_sql": "' OR '1'='1--",
        "__sentinel_union": "' UNION SELECT NULL--",
        "__sentinel_xss": REFLECTION_MARKER,
    }

    findings: List[Finding] = []

    endpoints = _discover_openapi_endpoints(request)
    if not endpoints:
        endpoints = [
            {
                "method": "GET",
                "url": f"https://{request.domain}",
                "query_params": ["__sentinel"],
                "json_fields": [],
            }
        ]

    for endpoint in endpoints[:8]:  # cap to protect against huge specifications
        response, error = _probe_endpoint(request, endpoint, payloads)
        if error:
            findings.append(
                create_translated_finding(
                    check="injection",
                    i18n_key="injection.probe_failed",
                    severity="medium",
                    i18n_params={
                        "url": endpoint["url"],
                        "error": str(error),
                    },
                )
            )
            continue

        if response is None:
            continue

        lower_body = response.text.lower()
        for signature in SQL_ERROR_SIGNATURES:
            if signature in lower_body:
                findings.append(
                    create_translated_finding(
                        check="injection",
                        i18n_key="injection.sql_error",
                        severity="high",
                        i18n_params={
                            "url": endpoint["url"],
                            "signature": signature,
                        },
                        evidence=signature
                    )
                )
                break

        if REFLECTION_MARKER.lower() in lower_body:
            findings.append(
                create_translated_finding(
                    check="injection",
                    i18n_key="injection.reflected_xss",
                    severity="high",
                    i18n_params={"url": endpoint["url"]},
                    evidence=REFLECTION_MARKER
                )
            )

        if response.status_code >= 500:
            findings.append(
                create_translated_finding(
                    check="injection",
                    i18n_key="injection.server_error",
                    severity="medium",
                    i18n_params={
                        "url": endpoint["url"],
                        "status": response.status_code,
                    },
                    evidence=str(response.status_code)
                )
            )

    if not findings:
        findings.append(
            create_translated_finding(
                check="injection",
                i18n_key="injection.none_detected",
                severity="info",
            )
        )

    return findings


def _discover_openapi_endpoints(request: ScanRequest) -> List[dict]:
    candidates = [
        f"https://{request.domain}/openapi.json",
        f"https://{request.domain}/.well-known/openapi.json",
        f"https://{request.domain}/swagger.json",
    ]

    headers = {"User-Agent": request.user_agent}
    for url in candidates:
        try:
            response = requests.get(url, headers=headers, timeout=request.timeout)
            if response.status_code == 404:
                continue
            data = response.json()
        except (requests.RequestException, ValueError):
            continue

        return _extract_endpoints_from_openapi(data, request.domain)

    return []


def _extract_endpoints_from_openapi(document: dict, domain: str) -> List[dict]:
    endpoints: List[dict] = []
    paths = document.get("paths", {})

    for path, methods in paths.items():
        if "{" in path:  # avoid templated parameters without context
            continue

        for method, operation in methods.items():
            method_upper = method.upper()
            if method_upper not in {"GET", "POST"}:
                continue

            parameters = operation.get("parameters", [])
            query_params = [
                param["name"]
                for param in parameters
                if param.get("in") == "query" and isinstance(param.get("name"), str)
            ]

            body_props = []
            request_body = operation.get("requestBody", {})
            content = request_body.get("content", {})
            json_schema = content.get("application/json", {}).get("schema", {})
            props = json_schema.get("properties", {})
            for name in props.keys():
                if isinstance(name, str):
                    body_props.append(name)

            endpoints.append(
                {
                    "method": method_upper,
                    "url": f"https://{domain}{path}",
                    "query_params": query_params[:5],
                    "json_fields": body_props[:5],
                }
            )

    return endpoints


def _probe_endpoint(
    request: ScanRequest, endpoint: dict, payloads: dict
) -> tuple[Optional[requests.Response], Optional[Exception]]:
    url = endpoint["url"]
    method = endpoint["method"]

    params = _assign_payloads(endpoint.get("query_params", []), payloads)
    json_body = _assign_payloads(endpoint.get("json_fields", []), payloads) or None

    try:
        response = requests.request(
            method=method,
            url=url,
            params=params if params else None,
            json=json_body,
            headers={"User-Agent": request.user_agent},
            timeout=request.timeout,
            allow_redirects=True,
        )
    except requests.RequestException as exc:
        return None, exc

    return response, None


def _assign_payloads(names: Sequence[str], payloads: dict) -> dict:
    values = [
        payloads["__sentinel_sql"],
        payloads["__sentinel_union"],
        payloads["__sentinel_xss"],
    ]
    assigned = {}
    for idx, name in enumerate(names):
        assigned[name] = values[min(idx, len(values) - 1)]
    return assigned
