"""Tests for Lab 01 — A2A delegation and a protocol-agnostic core.

    pytest test_lab.py -v
    LAB_MODULE=solution pytest test_lab.py -v    # the reference — must be green
"""

import importlib
import os
from dataclasses import replace

import pytest

lab = importlib.import_module(os.environ.get("LAB_MODULE", "lab"))


# ======================================================================================
# 1. Parts and messages
# ======================================================================================


def test_part_constructors():
    assert lab.Part.text_part("hi").kind is lab.PartKind.TEXT
    assert lab.Part.file_part("s3://x", "application/pdf").kind is lab.PartKind.FILE
    assert lab.Part.data_part({"a": 1}).kind is lab.PartKind.DATA


@pytest.mark.parametrize("kwargs", [
    {"kind": lab.PartKind.TEXT},
    {"kind": lab.PartKind.FILE},
    {"kind": lab.PartKind.DATA},
])
def test_parts_validate_their_payload(kwargs):
    with pytest.raises(ValueError):
        lab.Part(**kwargs)


def test_message_text_joins_only_text_parts():
    message = lab.Message("m1", lab.Role.USER,
                          (lab.Part.text_part("screen"), lab.Part.data_part({"x": 1}),
                           lab.Part.text_part("this")))
    assert message.text() == "screen this"


def test_data_part_is_copied_not_aliased():
    payload = {"a": 1}
    part = lab.Part.data_part(payload)
    payload["a"] = 2
    assert part.data["a"] == 1


# ======================================================================================
# 2. The task lifecycle
# ======================================================================================


def test_happy_path():
    state = lab.advance(lab.TaskState.SUBMITTED, lab.TaskState.WORKING)
    state = lab.advance(state, lab.TaskState.COMPLETED)
    assert state is lab.TaskState.COMPLETED


def test_terminal_states_are_absorbing():
    for terminal in lab.TERMINAL_STATES:
        for target in lab.TaskState:
            with pytest.raises(lab.IllegalTaskTransition):
                lab.advance(terminal, target)


def test_submitted_cannot_complete_directly():
    """Work has to be observed to have happened."""
    with pytest.raises(lab.IllegalTaskTransition):
        lab.advance(lab.TaskState.SUBMITTED, lab.TaskState.COMPLETED)


def test_input_required_round_trip():
    state = lab.advance(lab.TaskState.WORKING, lab.TaskState.INPUT_REQUIRED)
    assert lab.advance(state, lab.TaskState.WORKING) is lab.TaskState.WORKING


def test_auth_required_may_be_rejected():
    state = lab.advance(lab.TaskState.SUBMITTED, lab.TaskState.AUTH_REQUIRED)
    assert lab.advance(state, lab.TaskState.REJECTED) is lab.TaskState.REJECTED


def test_every_non_terminal_state_can_terminate():
    for state, targets in lab.TASK_TRANSITIONS.items():
        assert targets & lab.TERMINAL_STATES, f"{state} cannot terminate"


def test_input_required_cannot_be_rejected():
    """Rejection is an admission decision, not a mid-flight one."""
    with pytest.raises(lab.IllegalTaskTransition):
        lab.advance(lab.TaskState.INPUT_REQUIRED, lab.TaskState.REJECTED)


# ======================================================================================
# 3. Agent cards and discovery
# ======================================================================================


def card(name, **kwargs):
    kwargs.setdefault("description", f"{name} does things")
    kwargs.setdefault("url", f"https://agents.bank.ae/{name}")
    kwargs.setdefault("version", "1.0.0")
    return lab.AgentCard(name=name, **kwargs)


def directory():
    d = lab.AgentDirectory()
    d.register(card("sanctions",
                    skills=(lab.Skill("screen", "Screen", "screen an entity",
                                      tags=("sanctions", "compliance")),),
                    tenants=("wholesale",), max_data_classification="restricted"))
    d.register(card("kyc",
                    skills=(lab.Skill("refresh", "Refresh", "refresh kyc",
                                      tags=("kyc", "compliance")),),
                    tenants=("retail",), max_data_classification="confidential"))
    d.register(card("marketing",
                    skills=(lab.Skill("copy", "Copy", "write copy", tags=("marketing",)),),
                    max_data_classification="internal",
                    security_schemes=("apikey",)))
    return d


def caller(**kwargs):
    kwargs.setdefault("agent_id", "investigator")
    kwargs.setdefault("tenant", "wholesale")
    kwargs.setdefault("user_id", "u-1")
    kwargs.setdefault("data_classification", "internal")
    return lab.CallerContext(**kwargs)


def test_duplicate_registration_is_refused():
    d = lab.AgentDirectory()
    d.register(card("a"))
    with pytest.raises(ValueError):
        d.register(card("a"))


def test_discovery_filters_by_tenant():
    found = [c.name for c in directory().discover(caller())]
    assert "kyc" not in found
    assert "sanctions" in found


def test_discovery_filters_by_classification():
    strict = caller(data_classification="restricted")
    found = [c.name for c in directory().discover(strict)]
    assert found == ["sanctions"]


def test_discovery_ranks_by_skill_tag_overlap():
    d = directory()
    ranked = [c.name for c in d.discover(caller(data_classification="internal"),
                                         tags=["compliance"])]
    assert ranked == ["sanctions"]            # kyc is retail-only


def test_discovery_with_tags_excludes_zero_overlap():
    found = [c.name for c in directory().discover(caller(), tags=["marketing"])]
    assert found == ["marketing"]


def test_discovery_never_returns_the_caller():
    d = lab.AgentDirectory()
    d.register(card("investigator"))
    assert d.discover(caller()) == []


def test_discovery_is_deterministic():
    d = directory()
    first = [c.name for c in d.discover(caller())]
    assert first == [c.name for c in d.discover(caller())]


def test_agent_card_wire_shape_omits_platform_metadata():
    payload = card("x", tenants=("wholesale",), owner="team",
                   max_data_classification="restricted").to_json()
    assert set(payload) == {"protocolVersion", "name", "description", "url", "version",
                            "capabilities", "defaultInputModes", "defaultOutputModes",
                            "securitySchemes", "skills"}
    assert "tenants" not in payload
    assert "owner" not in payload


def test_unknown_classification_raises():
    with pytest.raises(ValueError):
        lab.classification_rank("cosmic")


# ======================================================================================
# 4. Delegation admission
# ======================================================================================


def test_permitted_delegation_has_no_denials():
    target = card("sanctions", tenants=("wholesale",), max_data_classification="restricted")
    assert lab.check_delegation(caller(), target) == []


def test_depth_is_bounded():
    target = card("sanctions", max_data_classification="restricted")
    deep = caller(delegation_chain=("a", "b", "c", "d"))
    assert [d.code for d in lab.check_delegation(deep, target)] == ["DEPTH_EXCEEDED"]


def test_depth_boundary_is_exclusive():
    target = card("sanctions", max_data_classification="restricted")
    at_limit = caller(delegation_chain=("a", "b", "c"))
    assert lab.check_delegation(at_limit, target) == []


def test_cycles_are_detected():
    target = card("sanctions", max_data_classification="restricted")
    cyclic = caller(delegation_chain=("sanctions",))
    assert "CYCLE_DETECTED" in [d.code for d in lab.check_delegation(cyclic, target)]


def test_cross_tenant_delegation_is_denied():
    target = card("kyc", tenants=("retail",), max_data_classification="restricted")
    assert "TENANT_NOT_PERMITTED" in [d.code for d in lab.check_delegation(caller(), target)]


def test_data_may_not_flow_to_a_less_cleared_agent():
    target = card("marketing", max_data_classification="internal")
    strict = caller(data_classification="restricted")
    assert "CLASSIFICATION_EXCEEDED" in [d.code for d in lab.check_delegation(strict, target)]


def test_weak_authentication_is_refused():
    target = card("x", security_schemes=("apikey",), max_data_classification="restricted")
    assert "NO_ACCEPTABLE_AUTH" in [d.code for d in lab.check_delegation(caller(), target)]


def test_all_denials_are_reported():
    target = card("x", tenants=("retail",), security_schemes=("apikey",),
                  max_data_classification="internal")
    codes = {d.code for d in lab.check_delegation(caller(data_classification="restricted"), target)}
    assert {"TENANT_NOT_PERMITTED", "CLASSIFICATION_EXCEEDED", "NO_ACCEPTABLE_AUTH"} <= codes


# ======================================================================================
# 5. Server behaviour
# ======================================================================================


def simple_handler(server, task, message):
    task = server.set_status(task, lab.TaskState.WORKING)
    yield lab.TaskStatusUpdate(task.task_id, task.context_id, task.status)
    artifact = lab.Artifact("art-1", "result", (lab.Part.text_part("done"),
                                                lab.Part.data_part({"score": 1})))
    task = server.add_artifact(task, artifact)
    yield lab.TaskArtifactUpdate(task.task_id, task.context_id, artifact, last_chunk=True)
    done = server.set_status(task, lab.TaskState.COMPLETED,
                             lab.Message("m-out", lab.Role.AGENT,
                                         (lab.Part.text_part("complete"),),
                                         task_id=task.task_id, context_id=task.context_id))
    yield lab.TaskStatusUpdate(done.task_id, done.context_id, done.status, final=True)


def clarifying_handler(server, task, message):
    if task.status.state is lab.TaskState.SUBMITTED:
        task = server.set_status(task, lab.TaskState.WORKING)
        asked = server.set_status(task, lab.TaskState.INPUT_REQUIRED,
                                  lab.Message("m-ask", lab.Role.AGENT,
                                              (lab.Part.text_part("which jurisdiction?"),),
                                              task_id=task.task_id,
                                              context_id=task.context_id))
        yield lab.TaskStatusUpdate(asked.task_id, asked.context_id, asked.status)
        return
    task = server.set_status(task, lab.TaskState.WORKING)
    done = server.set_status(task, lab.TaskState.COMPLETED,
                             lab.Message("m-done", lab.Role.AGENT,
                                         (lab.Part.text_part("finished"),),
                                         task_id=task.task_id, context_id=task.context_id))
    yield lab.TaskStatusUpdate(done.task_id, done.context_id, done.status, final=True)


def build(handler=simple_handler, **card_kwargs):
    card_kwargs.setdefault("max_data_classification", "restricted")
    card_kwargs.setdefault("push_notifications", True)
    server = lab.A2AServer(card=card("sanctions", **card_kwargs), handler=handler,
                           allowed_callback_hosts=("agents.bank.ae",))
    return server, lab.A2AClient(server)


def test_delegation_produces_a_completed_task_with_an_artifact():
    server, client = build()
    task = client.delegate(caller(), "screen this")
    assert task.status.state is lab.TaskState.COMPLETED
    assert len(task.artifacts) == 1
    assert task.artifacts[0].name == "result"


def test_the_delegation_chain_is_recorded_on_the_task():
    server, client = build()
    task = client.delegate(caller(delegation_chain=("orchestrator",)), "go")
    assert task.delegation_chain == ("orchestrator", "investigator")


def test_the_callee_cannot_forge_the_chain():
    """The chain comes from the caller's context, not from the message."""
    server, client = build()
    task = client.delegate(caller(delegation_chain=("a", "b")), "go")
    assert task.delegation_chain == ("a", "b", "investigator")


def test_a_denied_delegation_never_creates_a_task():
    server, client = build(tenants=("retail",))
    with pytest.raises(lab.A2AError) as exc:
        client.delegate(caller(), "go")
    assert exc.value.code == "TENANT_NOT_PERMITTED"
    assert server.tasks == {}


def test_streaming_yields_updates_then_the_task():
    server, client = build()
    events = client.delegate_streaming(caller(), "go")
    assert isinstance(events[-1], lab.Task)
    kinds = [type(e).__name__ for e in events]
    assert kinds[0] == "TaskStatusUpdate"
    assert "TaskArtifactUpdate" in kinds


def test_send_and_stream_agree():
    server_a, client_a = build()
    server_b, client_b = build()
    sent = client_a.delegate(caller(), "go")
    streamed = client_b.delegate_streaming(caller(), "go")[-1]
    assert sent.status.state is streamed.status.state
    assert [a.name for a in sent.artifacts] == [a.name for a in streamed.artifacts]


def test_input_required_pauses_and_resumes():
    server, client = build(handler=clarifying_handler)
    task = client.delegate(caller(), "screen")
    assert task.status.state is lab.TaskState.INPUT_REQUIRED
    assert task.status.message.text() == "which jurisdiction?"
    resumed = client.reply(caller(), task, "UAE")
    assert resumed.status.state is lab.TaskState.COMPLETED
    assert [m.role for m in resumed.history] == [lab.Role.USER, lab.Role.AGENT,
                                                 lab.Role.USER, lab.Role.AGENT]


def test_replying_to_a_terminal_task_is_refused():
    server, client = build()
    task = client.delegate(caller(), "go")
    with pytest.raises(lab.A2AError) as exc:
        client.reply(caller(), task, "more")
    assert exc.value.code == "TASK_TERMINAL"


def test_messaging_an_unknown_task_is_refused():
    server, client = build()
    message = lab.Message("m", lab.Role.USER, (lab.Part.text_part("x"),), task_id="nope")
    with pytest.raises(lab.A2AError) as exc:
        server.message_send(message, caller=caller())
    assert exc.value.code == "TASK_NOT_FOUND"


def test_tasks_get_and_cancel():
    server, client = build(handler=clarifying_handler)
    task = client.delegate(caller(), "screen")
    assert server.tasks_get(task.task_id).task_id == task.task_id
    cancelled = server.tasks_cancel(task.task_id)
    assert cancelled.status.state is lab.TaskState.CANCELED
    with pytest.raises(lab.A2AError):
        server.tasks_cancel(task.task_id)
    with pytest.raises(lab.A2AError):
        server.tasks_get("nope")


def test_context_id_groups_related_tasks():
    server, client = build()
    first = client.delegate(caller(), "one")
    second = client.delegate(caller(), "two", context_id=first.context_id)
    assert first.context_id == second.context_id
    assert first.task_id != second.task_id


def test_status_timestamps_are_monotone():
    server, client = build()
    task = client.delegate(caller(), "go")
    assert task.status.timestamp > 0


# ======================================================================================
# 6. Push notifications
# ======================================================================================


def test_push_config_requires_support():
    server, client = build(handler=clarifying_handler, push_notifications=False)
    task = client.delegate(caller(), "go")
    with pytest.raises(lab.A2AError) as exc:
        server.set_push_config(task.task_id, lab.PushNotificationConfig(
            "https://agents.bank.ae/cb", token="t"))
    assert exc.value.code == "UNSUPPORTED"


def test_push_config_rejects_hosts_outside_the_allow_list():
    server, client = build(handler=clarifying_handler)
    task = client.delegate(caller(), "go")
    with pytest.raises(lab.A2AError) as exc:
        server.set_push_config(task.task_id,
                               lab.PushNotificationConfig("https://evil.example/x", token="t"))
    assert exc.value.code == "CALLBACK_NOT_ALLOWED"


def test_push_config_requires_a_token():
    server, client = build(handler=clarifying_handler)
    task = client.delegate(caller(), "go")
    with pytest.raises(lab.A2AError) as exc:
        server.set_push_config(task.task_id,
                               lab.PushNotificationConfig("https://agents.bank.ae/cb"))
    assert exc.value.code == "CALLBACK_UNAUTHENTICATED"


def test_push_notifications_fire_on_every_status_change():
    server, client = build(handler=clarifying_handler)
    task = client.delegate(caller(), "go")
    server.set_push_config(task.task_id,
                           lab.PushNotificationConfig("https://agents.bank.ae/cb", token="t"))
    client.reply(caller(), task, "UAE")
    states = [status.state for _, status in server.pushed]
    assert lab.TaskState.COMPLETED in states


def test_push_config_on_an_unknown_task_is_refused():
    server, _ = build(handler=clarifying_handler)
    with pytest.raises(lab.A2AError):
        server.set_push_config("nope", lab.PushNotificationConfig(
            "https://agents.bank.ae/cb", token="t"))


# ======================================================================================
# 7. The protocol-agnostic core
# ======================================================================================


def completed_task():
    server, client = build()
    return client.delegate(caller(), "screen the beneficiary")


def test_a2a_maps_onto_internal_vocabulary():
    internal = lab.a2a_to_internal(completed_task())
    assert internal.state == "succeeded"
    assert internal.input_text == "screen the beneficiary"
    assert internal.output_text == "done"
    assert internal.structured == {"score": 1}
    assert internal.artifacts == ("result",)


def test_every_a2a_state_maps_and_round_trips():
    for state in lab.TaskState:
        internal = lab.A2A_STATE_TO_INTERNAL[state]
        assert lab.INTERNAL_TO_A2A_STATE[internal] is state


def test_every_acp_status_maps():
    for status, internal in lab.ACP_STATUS_TO_INTERNAL.items():
        assert lab.INTERNAL_TO_ACP_STATUS[internal] in lab.ACP_STATUS_TO_INTERNAL


def test_internal_round_trips_through_acp():
    internal = lab.a2a_to_internal(completed_task())
    envelope = lab.internal_to_acp(internal)
    assert lab.acp_to_internal(envelope) == internal


def test_acp_envelope_shape():
    envelope = lab.internal_to_acp(lab.a2a_to_internal(completed_task()))
    assert envelope["status"] == "completed"
    assert envelope["output"][0]["role"] == "agent"
    content_types = {p["content_type"] for p in envelope["output"][0]["parts"]}
    assert content_types == {"text/plain", "application/json"}
    assert envelope["metadata"]["delegation_chain"] == ["investigator"]


def test_acp_rejects_an_unknown_status():
    with pytest.raises(ValueError):
        lab.acp_to_internal({"status": "quantum", "run_id": "r", "session_id": "s"})


def test_states_with_no_acp_equivalent_degrade_predictably():
    """ACP has no 'rejected'; it must map to something, and the choice must be stated."""
    assert lab.INTERNAL_TO_ACP_STATUS["rejected"] == "failed"
    assert lab.INTERNAL_TO_ACP_STATUS["awaiting_auth"] == "awaiting"


def test_an_empty_task_still_converts():
    internal = lab.InternalTask(task_id="t", context_id="c", state="queued",
                                input_text="hi")
    envelope = lab.internal_to_acp(internal)
    assert envelope["status"] == "created"
    assert lab.acp_to_internal(envelope) == internal


# ======================================================================================
# 8. Determinism
# ======================================================================================


def test_two_identical_servers_produce_identical_tasks():
    server_a, client_a = build()
    server_b, client_b = build()
    a = client_a.delegate(caller(), "go")
    b = client_b.delegate(caller(), "go")
    assert a == b
