from __future__ import annotations

import base64
from pathlib import Path
from typing import Any

from autoshorts.nvidia_image_client import NvidiaImageClient, NvidiaImageRequest, load_nvidia_api_key

ONE_PIXEL_PNG = base64.b64decode(
    "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mP8/x8AAwMCAO+/p9sAAAAASUVORK5CYII="
)


class FakeTransport:
    def __init__(self, statuses: list[int] | None = None):
        self.statuses = statuses or [200]
        self.calls: list[dict[str, Any]] = []

    def post_json(self, endpoint: str, api_key: str, payload: dict[str, Any], *, timeout_seconds: int) -> tuple[int, bytes, str]:
        self.calls.append({"endpoint": endpoint, "api_key": api_key, "payload": payload, "timeout_seconds": timeout_seconds})
        status = self.statuses[min(len(self.calls) - 1, len(self.statuses) - 1)]
        if status != 200:
            return status, b"404 page not found", "text/plain"
        encoded = base64.b64encode(ONE_PIXEL_PNG).decode()
        return 200, (f'{{"artifacts":[{{"base64":"{encoded}"}}]}}').encode(), "application/json"


def test_load_nvidia_api_key_supports_env_file_format(tmp_path: Path) -> None:
    secret = tmp_path / "nvidia_api_key"
    secret.write_text("NVIDIA_API_KEY=test-key\n", encoding="utf-8")

    assert load_nvidia_api_key(secret) == "test-key"


def test_client_writes_image_from_artifacts_response(tmp_path: Path) -> None:
    transport = FakeTransport()
    client = NvidiaImageClient(api_key="secret", transport=transport)
    output = tmp_path / "scene_01.png"

    result = client.generate(NvidiaImageRequest(prompt="red car", output_path=output, width=416, height=736), model_order=["flux.1-dev"])

    assert result.status == "success"
    assert result.model == "flux.1-dev"
    assert output.read_bytes() == ONE_PIXEL_PNG
    assert transport.calls[0]["payload"]["width"] == 416
    assert transport.calls[0]["payload"]["height"] == 736
    assert transport.calls[0]["payload"]["mode"] == "base"


def test_client_falls_back_when_first_model_is_unavailable(tmp_path: Path) -> None:
    transport = FakeTransport(statuses=[404, 200])
    client = NvidiaImageClient(api_key="secret", transport=transport)
    output = tmp_path / "scene_01.png"

    result = client.generate(
        NvidiaImageRequest(prompt="red car", output_path=output),
        model_order=["qwen-image", "flux.1-dev"],
    )

    assert result.status == "success"
    assert result.model == "flux.1-dev"
    assert len(transport.calls) == 2
    assert transport.calls[0]["endpoint"].endswith("/qwen/qwen-image")
    assert transport.calls[1]["endpoint"].endswith("/black-forest-labs/flux.1-dev")
