"""The CLI: JSON lines on stderr, one request id per request, and honest exit codes."""

import json
import threading

import pytest

from askcli import DEFAULTS, ask, backend, configure_logging, load_config, main

KEYS = {"ts", "level", "logger", "msg", "request_id"}


def run(argv):
    """main's exit code — whether it returns the int or lets argparse's SystemExit out."""
    try:
        return main(argv)
    except SystemExit as exc:
        return exc.code


def json_rows(stream_text, where="stderr"):
    rows = []
    for i, line in enumerate(ln for ln in stream_text.splitlines() if ln.strip()):
        try:
            obj = json.loads(line)
        except json.JSONDecodeError:
            pytest.fail(f"{where} line {i} is not one JSON object: {line!r}")
        assert isinstance(obj, dict), f"{where} line {i} is JSON but not an object: {line!r}"
        rows.append(obj)
    return rows


def test_json_flag_prints_one_object_per_line(capsys):
    for attempt in (1, 2):  # the second run proves configure_logging did not add a second handler
        code = run(["ask", "what is a logger tree?", "--json"])
        out, err = capsys.readouterr()
        assert code == 0, f"run {attempt}: exit code {code!r}, want 0"
        rows = json_rows(err)
        assert len(rows) == 3, (
            f"run {attempt}: want exactly 3 stderr lines (request start, model call, "
            f"request done), got {len(rows)}: {err!r}"
        )
        for row in rows:
            missing = KEYS - row.keys()
            assert not missing, f"run {attempt}: line {row!r} is missing {sorted(missing)}"
        loggers = [row["logger"] for row in rows]
        assert loggers == ["askcli.cli", "askcli.backend", "askcli.cli"], (
            f"run {attempt}: loggers in order {loggers}, want the cli, the backend, the cli"
        )
        answers = json_rows(out, where="stdout")
        assert len(answers) == 1, (
            f"run {attempt}: stdout should be ONE JSON object, got {out!r}"
        )
        assert "what is a logger tree?" in answers[0].get("answer", ""), (
            f"run {attempt}: stdout object {answers[0]!r} has no 'answer' holding the prompt"
        )


def test_request_id_same_within_a_request_and_distinct_across_threads(capsys, monkeypatch):
    configure_logging(json_lines=True, level="INFO")
    cfg = load_config(DEFAULTS, None, env={})
    both_in_flight = threading.Barrier(2, timeout=5)
    both_logged = threading.Barrier(2, timeout=5)
    real_complete = backend.complete

    def complete_while_the_other_request_runs(prompt, model):
        both_in_flight.wait()   # each request has set its id; neither has finished
        out = real_complete(prompt, model)   # logs "model call" from askcli.backend
        both_logged.wait()      # nobody leaves its request until both model-call lines exist
        return out

    monkeypatch.setattr(backend, "complete", complete_while_the_other_request_runs)
    errors = []

    def serve(prompt):
        try:
            ask(prompt, cfg)
        except Exception as exc:  # a BrokenBarrierError means the requests never overlapped
            errors.append(repr(exc))

    threads = [threading.Thread(target=serve, args=(p,)) for p in ("alpha", "beta")]
    for t in threads:
        t.start()
    for t in threads:
        t.join(10)
    assert not errors, f"a request failed while two were in flight: {errors}"

    rows = json_rows(capsys.readouterr().err)
    ids = [row.get("request_id") for row in rows]
    assert "-" not in ids and None not in ids, (
        f"a line inside a request carries no request id: {ids}"
    )
    groups = {}
    for row in rows:
        groups.setdefault(row["request_id"], []).append(row["msg"])
    assert len(groups) == 2, (
        f"two requests should give two distinct ids, got {len(groups)}: {groups}"
    )
    for rid, msgs in groups.items():
        starts = [m for m in msgs if m.startswith("request start")]
        calls = [m for m in msgs if m.startswith("model call")]
        dones = [m for m in msgs if m.startswith("request done")]
        assert (len(starts), len(calls), len(dones)) == (1, 1, 1), (
            f"request_id {rid} stamps {msgs} — one request's lines must all carry ITS id: "
            "one start, one model call, one done. Another request's line under this id means "
            "the id is shared by the whole process, not held per thread"
        )


def test_exit_code_2_on_bad_args_and_0_on_success(capsys):
    for argv in ([], ["ask"], ["ask", "hi", "--bogus"]):
        code = run(argv)
        assert code == 2, f"askcli {' '.join(argv)!r}: exit code {code!r}, want 2 (bad arguments)"
    capsys.readouterr()
    code = run(["ask", "hi there", "--model", "gpt-mini"])
    out, _ = capsys.readouterr()
    assert code == 0, f"askcli ask 'hi there' --model gpt-mini: exit code {code!r}, want 0"
    assert "gpt-mini" in out and "hi there" in out, (
        f"stdout should hold the answer from model gpt-mini to 'hi there'; got {out!r}"
    )
