"""Offline probes for the Companion article; run with upstream src on PYTHONPATH."""

import asyncio
import json

from m_agent.adapters import (
    DeterministicModelAdapter,
    InMemoryRunStore,
    PlaintextPayloadCodec,
)
from m_agent.companion import (
    InMemorySessionStore,
    SessionCommitStatus,
    SessionRunner,
    SessionScope,
)
from m_agent.runtime import AgentDefinition, DefinitionRegistry, Runner, RunStatus


def make_runner(store):
    adapter = DeterministicModelAdapter(responses=("first answer", "second answer"))
    registry = DefinitionRegistry()
    registry.register(
        AgentDefinition.for_adapter(
            definition_id="support",
            version="1",
            instructions="Answer the support question.",
            model_adapter=adapter,
        )
    )
    return Runner(registry=registry, store=store), adapter


class CommitThenRaiseRunStore(InMemoryRunStore):
    """Model a successful Store write whose acknowledgement is lost."""

    def __init__(self):
        super().__init__(payload_codec=PlaintextPayloadCodec())
        self.created_ids = []
        self.fail_once = True

    async def create_run(self, run):
        result = await super().create_run(run)
        self.created_ids.append(run.run_id)
        if self.fail_once:
            self.fail_once = False
            raise RuntimeError("injected acknowledgement loss after create")
        return result


async def main():
    scope = SessionScope(token="offline-test-scope")
    store = InMemoryRunStore(payload_codec=PlaintextPayloadCodec())
    runner, adapter = make_runner(store)
    sessions = InMemorySessionStore()
    await sessions.create_session(scope, "conversation")
    companion = SessionRunner(runner=runner, session_store=sessions)
    first = await companion.submit(scope, "conversation", "support", "1", "Where is it?")
    second = await companion.submit(scope, "conversation", "support", "1", "What next?")
    snapshot = await sessions.read_snapshot(scope, "conversation")
    assert first.commit_status is SessionCommitStatus.COMMITTED
    assert second.commit_status is SessionCommitStatus.COMMITTED
    assert snapshot.version == 2 and len(snapshot.turns) == 2
    assert len(second.run.history) == 2
    assert second.run.history[0].content == "Where is it?"
    assert second.run.history[1].content == first.run.output
    assert await sessions.get_claim(scope, "conversation") is None
    assert adapter.call_count == 2

    faulty_store = CommitThenRaiseRunStore()
    faulty_runner, faulty_adapter = make_runner(faulty_store)
    faulty_sessions = InMemorySessionStore()
    await faulty_sessions.create_session(scope, "ack-loss")
    faulty_companion = SessionRunner(
        runner=faulty_runner, session_store=faulty_sessions
    )
    try:
        await faulty_companion.submit(scope, "ack-loss", "support", "1", "First input")
    except RuntimeError as error:
        assert "acknowledgement loss" in str(error)
    else:
        raise AssertionError("expected injected error")
    orphan = await faulty_runner.get_run(faulty_store.created_ids[0])
    assert orphan.status is RunStatus.CREATED
    assert await faulty_sessions.get_claim(scope, "ack-loss") is None
    assert faulty_adapter.call_count == 0
    followup = await faulty_companion.submit(
        scope, "ack-loss", "support", "1", "Next input"
    )
    assert followup.run.status is RunStatus.SUCCEEDED
    assert followup.run.run_id != orphan.run_id
    assert (await faulty_runner.get_run(orphan.run_id)).status is RunStatus.CREATED
    print(json.dumps({
        "two_turn_smoke": "PASS",
        "snapshot_version": snapshot.version,
        "second_run_history_messages": len(second.run.history),
        "create_ack_loss": {
            "reproduced": True,
            "first_run_status": orphan.status.value,
            "claim_released_after_error": True,
            "followup_admitted": True,
            "model_calls_before_followup": 0,
        },
        "evidence_scope": "offline public API; injected Store acknowledgement loss",
    }, indent=2))


if __name__ == "__main__":
    asyncio.run(main())
