"""No network: exercise the installed classifier, LCEL and async transport."""
import asyncio
import copy
import json
import unittest
import httpx2
from langchain_typesafe import TypeSafeClassifier
from langchain_workflow import DEMO_RESPONSE, build_chain, route_ticket, route_batch


class WorkflowTests(unittest.IsolatedAsyncioTestCase):
    async def asyncSetUp(self):
        self.response = copy.deepcopy(DEMO_RESPONSE)
        self.status = 200
        self.delay = 0
        self.calls = self.active = self.peak = 0
        self.sync = httpx2.Client(transport=httpx2.MockTransport(lambda _: httpx2.Response(500)))
        self.http = httpx2.AsyncClient(transport=httpx2.MockTransport(self.handle), timeout=2)
        self.chain = build_chain(TypeSafeClassifier(model="jev-1.13", api_key="mock-only", client=self.sync, async_client=self.http))

    async def asyncTearDown(self):
        self.sync.close()
        await self.http.aclose()

    async def handle(self, request):
        self.calls += 1
        body = json.loads(request.content)
        self.assertEqual(request.url.path, "/v1/systemone")
        self.assertEqual(set(body["state"]), {"message"})
        self.assertEqual(set(body["questions"]), {"route", "severity", "urgent"})
        self.active += 1
        self.peak = max(self.peak, self.active)
        try:
            await asyncio.sleep(self.delay)
            return httpx2.Response(self.status, json=self.response)
        finally:
            self.active -= 1

    async def run_ticket(self, message="CSV works", **kwargs):
        return await route_ticket(self.chain, {"id": "local", "message": message, "private": "not sent"}, **kwargs)

    async def test_valid_chain(self):
        result = await self.run_ticket()
        self.assertEqual((result["action"], result["queue"]), ("suggest", "technical"))

    async def test_billing_branch(self):
        self.response["answers"]["route"].update(choice="billing", probabilities={"technical": 0.05, "billing": 0.9, "human": 0.05})
        self.assertEqual((await self.run_ticket())["queue"], "billing")

    async def test_policy_boundaries(self):
        for key, field, value, reason in [("route", "confidence", 0.74, "low_confidence"),
                ("severity", "score", 1.5, "high_impact"), ("urgent", "noul", 0.8, "high_impact")]:
            self.response = copy.deepcopy(DEMO_RESPONSE)
            self.response["answers"][key][field] = value
            result = await self.run_ticket()
            self.assertEqual(result["reason"], reason)
            self.assertEqual(result["queue"], "human")

    async def test_invalid_input_no_request(self):
        for message in ["", " ", None, "x" * 4001]:
            self.assertEqual((await self.run_ticket(message))["reason"], "invalid_input")
        self.assertEqual(self.calls, 0)

    async def test_bad_response_never_suggests(self):
        for changed in [{}, {"answers": {}}, {**DEMO_RESPONSE, "answers": {**DEMO_RESPONSE["answers"],
                "route": {**DEMO_RESPONSE["answers"]["route"], "choice": "invented"}}}]:
            self.response = changed
            self.assertEqual((await self.run_ticket())["action"], "review")

    async def test_http_failure_no_implicit_retry(self):
        for status in [401, 429, 529]:
            self.status = status
            before = self.calls
            self.assertEqual((await self.run_ticket())["reason"], "provider_or_validation_failure")
            self.assertEqual(self.calls - before, 1)

    async def test_deadline_cancels_transport(self):
        self.delay = 1
        self.assertEqual((await self.run_ticket(budget=0.03))["action"], "review")
        self.assertEqual(self.active, 0)

    async def test_batch_bounds_concurrency_and_keeps_results(self):
        self.delay = 0.01
        tickets = [{"id": str(i), "message": "CSV works" if i != 2 else ""} for i in range(7)]
        results = await route_batch(self.chain, tickets, concurrency=2)
        self.assertEqual(len(results), 7)
        self.assertEqual(results[2]["reason"], "invalid_input")
        self.assertEqual(self.calls, 6)
        self.assertLessEqual(self.peak, 2)
        with self.assertRaises(ValueError):
            await route_batch(self.chain, [], concurrency=0)


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