import asyncio
import types
import unittest
from dataclasses import replace

from mrjunejune.inference.copilot_sidecar import Sidecar, SidecarConfig


def event(event_type, **data):
    return types.SimpleNamespace(
        type=types.SimpleNamespace(value=event_type),
        data=types.SimpleNamespace(**data),
    )


class FakeSession:
    def __init__(self, session_id, behavior="complete"):
        self.session_id = session_id
        self.behavior = behavior
        self.handlers = []
        self.prompts = []
        self.abort_calls = 0
        self.disconnect_calls = 0
        self.disconnect_error = False

    def on(self, handler):
        self.handlers.append(handler)

        def unsubscribe():
            if handler in self.handlers:
                self.handlers.remove(handler)

        return unsubscribe

    def emit(self, value):
        for handler in list(self.handlers):
            handler(value)

    async def send(self, prompt):
        self.prompts.append(prompt)
        message_id = f"message-{len(self.prompts)}"
        if self.behavior == "complete":
            self.emit(
                event(
                    "assistant.message_delta",
                    delta_content=f"{prompt}-delta",
                    message_id=message_id,
                )
            )
            self.emit(
                event(
                    "assistant.message",
                    content=f"{prompt}-answer",
                    message_id=message_id,
                    model="test-model",
                )
            )
            self.emit(
                event(
                    "assistant.usage",
                    model="test-model",
                    input_tokens=3,
                    output_tokens=5,
                )
            )
            self.emit(event("session.idle", aborted=False))
        elif self.behavior == "error":
            self.emit(
                event(
                    "session.error",
                    error_type="provider",
                    error_code="upstream_error",
                    message="upstream failed",
                    status_code=502,
                )
            )
        return message_id

    async def abort(self):
        self.abort_calls += 1
        if self.behavior == "abort_error":
            raise RuntimeError("abort failed")
        self.emit(event("session.idle", aborted=True))

    async def disconnect(self):
        self.disconnect_calls += 1
        if self.disconnect_error:
            raise RuntimeError("disconnect failed")


class FakeClient:
    def __init__(self, behavior="complete"):
        self.behavior = behavior
        self.sessions = {}
        self.create_calls = []
        self.resume_calls = []
        self.delete_calls = []
        self.started = False
        self.stopped = False
        self.resume_started = None
        self.resume_release = None

    async def start(self):
        self.started = True

    async def stop(self):
        self.stopped = True

    async def resume_session(self, session_id, **kwargs):
        self.resume_calls.append((session_id, kwargs))
        if self.resume_started is not None:
            self.resume_started.set()
        if self.resume_release is not None:
            await self.resume_release.wait()
        if session_id not in self.sessions:
            raise LookupError(session_id)
        return self.sessions[session_id]

    async def create_session(self, session_id, **kwargs):
        self.create_calls.append((session_id, kwargs))
        session = FakeSession(session_id, self.behavior)
        self.sessions[session_id] = session
        return session

    async def delete_session(self, session_id):
        self.delete_calls.append(session_id)
        self.sessions.pop(session_id, None)


class SidecarTest(unittest.IsolatedAsyncioTestCase):
    async def asyncSetUp(self):
        self.output = []

        async def capture(payload):
            self.output.append(payload)

        self.capture = capture
        self.config = SidecarConfig(
            base_url="http://litellm.invalid/v1",
            model="test-model",
            wire_api="responses",
            base_directory="/not-used-by-fake",
        )

    async def make_sidecar(self, behavior="complete"):
        client = FakeClient(behavior)
        sidecar = Sidecar(client, self.config, self.capture)
        await sidecar.start()
        return sidecar, client

    async def test_health_and_invalid_protocol(self):
        sidecar, _ = await self.make_sidecar()
        await sidecar.announce_ready()
        await sidecar.dispatch({"command": "health", "request_id": "health-1"})
        await sidecar.dispatch({"command": "unknown", "request_id": "bad-1"})

        self.assertEqual(self.output[0]["type"], "ready")
        self.assertIsNone(self.output[0]["request_id"])
        self.assertEqual(self.output[1]["request_id"], "health-1")
        self.assertEqual(
            [item["type"] for item in self.output[-2:]],
            ["turn.error", "turn.done"],
        )

    async def test_multi_turn_reuses_one_session_and_configures_sdk(self):
        sidecar, client = await self.make_sidecar()
        for request_id, prompt in (("r1", "first"), ("r2", "second")):
            await sidecar.dispatch(
                {
                    "command": "turn.start",
                    "request_id": request_id,
                    "conversation_id": "conversation-a",
                    "prompt": prompt,
                }
            )
            await sidecar.drain_events()

        self.assertEqual(len(client.create_calls), 1)
        self.assertEqual(len(client.resume_calls), 1)
        session = client.sessions["conversation-a"]
        self.assertEqual(session.prompts, ["first", "second"])
        options = client.create_calls[0][1]
        self.assertEqual(options["provider"]["type"], "openai")
        self.assertEqual(
            options["provider"]["base_url"], "http://litellm.invalid/v1"
        )
        self.assertEqual(options["provider"]["wire_api"], "responses")
        self.assertEqual(options["model"], "test-model")
        self.assertEqual(options["available_tools"], [])
        self.assertTrue(options["streaming"])
        self.assertEqual(
            [item["request_id"] for item in self.output if item["type"] == "turn.done"],
            ["r1", "r2"],
        )

    async def test_concurrent_conversations_keep_correlation(self):
        sidecar, _ = await self.make_sidecar()
        await asyncio.gather(
            sidecar.dispatch(
                {
                    "command": "turn.start",
                    "request_id": "left-request",
                    "conversation_id": "left",
                    "prompt": "left",
                }
            ),
            sidecar.dispatch(
                {
                    "command": "turn.start",
                    "request_id": "right-request",
                    "conversation_id": "right",
                    "prompt": "right",
                }
            ),
        )
        await sidecar.drain_events()

        correlated = {
            (item["request_id"], item["conversation_id"])
            for item in self.output
            if item["type"] in ("assistant.delta", "assistant.completed")
        }
        self.assertEqual(
            correlated,
            {("left-request", "left"), ("right-request", "right")},
        )

    async def test_abort_finishes_active_turn(self):
        sidecar, client = await self.make_sidecar("pending")
        await sidecar.dispatch(
            {
                "command": "turn.start",
                "request_id": "turn-request",
                "conversation_id": "abort-me",
                "prompt": "wait",
            }
        )
        await sidecar.dispatch(
            {
                "command": "turn.abort",
                "request_id": "abort-request",
                "conversation_id": "abort-me",
            }
        )
        await sidecar.drain_events()

        self.assertEqual(client.sessions["abort-me"].abort_calls, 1)
        done = [
            item
            for item in self.output
            if item["type"] == "turn.done"
            and item["request_id"] == "turn-request"
        ]
        self.assertEqual(done[0]["aborted"], True)
        abort_accept = [
            item
            for item in self.output
            if item["type"] == "turn.accepted"
            and item["request_id"] == "abort-request"
        ]
        self.assertEqual(abort_accept[0]["target_request_id"], "turn-request")

    async def test_session_error_is_terminal(self):
        sidecar, _ = await self.make_sidecar("error")
        await sidecar.dispatch(
            {
                "command": "turn.start",
                "request_id": "error-request",
                "conversation_id": "error-conversation",
                "prompt": "fail",
            }
        )
        await sidecar.drain_events()

        errors = [item for item in self.output if item["type"] == "turn.error"]
        done = [item for item in self.output if item["type"] == "turn.done"]
        self.assertEqual(errors[0]["error"]["code"], "upstream_error")
        self.assertEqual(errors[0]["error"]["status_code"], 502)
        self.assertEqual(len(done), 1)
        self.assertTrue(done[0]["failed"])

    async def test_abort_error_terminates_abort_request(self):
        sidecar, _ = await self.make_sidecar("abort_error")
        await sidecar.dispatch(
            {
                "command": "turn.start",
                "request_id": "active-request",
                "conversation_id": "abort-error",
                "prompt": "wait",
            }
        )
        await sidecar.dispatch(
            {
                "command": "turn.abort",
                "request_id": "failed-abort",
                "conversation_id": "abort-error",
            }
        )

        abort_events = [
            item for item in self.output if item["request_id"] == "failed-abort"
        ]
        self.assertEqual(
            [item["type"] for item in abort_events],
            ["turn.accepted", "turn.error", "turn.done"],
        )
        self.assertTrue(abort_events[-1]["failed"])

    async def test_delete_and_shutdown_release_resources(self):
        sidecar, client = await self.make_sidecar("pending")
        await sidecar.dispatch(
            {
                "command": "turn.start",
                "request_id": "active",
                "conversation_id": "delete-me",
                "prompt": "wait",
            }
        )
        session = client.sessions["delete-me"]
        await sidecar.dispatch(
            {
                "command": "conversation.delete",
                "request_id": "delete-request",
                "conversation_id": "delete-me",
            }
        )
        await sidecar.dispatch(
            {
                "command": "shutdown",
                "request_id": "shutdown-request",
                "conversation_id": None,
            }
        )

        self.assertEqual(session.disconnect_calls, 1)
        self.assertEqual(client.delete_calls, ["delete-me"])
        self.assertTrue(client.stopped)
        self.assertTrue(sidecar.shutting_down)
        shutdown = [
            item
            for item in self.output
            if item["request_id"] == "shutdown-request"
        ]
        self.assertEqual(shutdown[0]["type"], "turn.done")

    async def test_idle_and_overflow_sessions_are_evicted(self):
        self.config = replace(
            self.config,
            idle_timeout_seconds=3600,
            max_sessions=1,
        )
        sidecar, client = await self.make_sidecar()
        for conversation_id in ("old", "new"):
            await sidecar.dispatch(
                {
                    "command": "turn.start",
                    "request_id": f"request-{conversation_id}",
                    "conversation_id": conversation_id,
                    "prompt": conversation_id,
                }
            )
            await sidecar.drain_events()
            await asyncio.sleep(0.01)

        old_session = client.sessions["old"]
        old_session.disconnect_error = True
        await sidecar.evict_idle_sessions()
        self.assertEqual(old_session.disconnect_calls, 1)
        self.assertNotIn("old", sidecar._conversations)
        self.assertIn("new", sidecar._conversations)

        sidecar._conversations["new"].last_used -= 4000
        await sidecar.evict_idle_sessions()
        self.assertEqual(client.sessions["new"].disconnect_calls, 1)
        self.assertEqual(sidecar._conversations, {})
        await sidecar.dispatch(
            {
                "command": "shutdown",
                "request_id": "shutdown-eviction",
                "conversation_id": None,
            }
        )

    async def test_shutdown_waits_for_inflight_session_creation(self):
        client = FakeClient("pending")
        client.resume_started = asyncio.Event()
        client.resume_release = asyncio.Event()
        sidecar = Sidecar(client, self.config, self.capture)
        await sidecar.start()

        start_task = asyncio.create_task(
            sidecar.dispatch(
                {
                    "command": "turn.start",
                    "request_id": "starting",
                    "conversation_id": "race",
                    "prompt": "wait",
                }
            )
        )
        await client.resume_started.wait()
        shutdown_task = asyncio.create_task(
            sidecar.dispatch(
                {
                    "command": "shutdown",
                    "request_id": "shutdown",
                    "conversation_id": None,
                }
            )
        )
        await asyncio.sleep(0)
        self.assertTrue(sidecar.shutting_down)
        client.resume_release.set()
        await asyncio.gather(start_task, shutdown_task)

        self.assertTrue(client.stopped)
        self.assertEqual(client.sessions["race"].disconnect_calls, 1)
        await sidecar.dispatch(
            {
                "command": "turn.start",
                "request_id": "too-late",
                "conversation_id": "late",
                "prompt": "no",
            }
        )
        late_error = [
            item
            for item in self.output
            if item["request_id"] == "too-late" and item["type"] == "turn.error"
        ]
        self.assertEqual(late_error[0]["error"]["code"], "shutting_down")

    async def test_command_gates_are_released(self):
        sidecar, _ = await self.make_sidecar()
        for index in range(20):
            await sidecar.dispatch(
                {
                    "command": "unknown",
                    "request_id": f"unknown-{index}",
                    "conversation_id": f"conversation-{index}",
                }
            )
        self.assertEqual(sidecar._conversation_gates, {})


if __name__ == "__main__":
    unittest.main()
