from __future__ import annotations import pytest from app.llm.base import ChatResult, ProviderError, RateLimitError from app.llm.router import LLMRouter class FakeProvider: def __init__(self, result: ChatResult | None = None, raise_exc: Exception | None = None): self.result = result self.raise_exc = raise_exc self.calls = 0 self.last_model: str | None = None async def chat(self, messages, tools, model, temperature=0.3, tool_choice=None, max_tokens=None): self.calls += 1 self.last_model = model if self.raise_exc is not None: raise self.raise_exc assert self.result is not None return self.result def _models(): return { "groq": {"small": "gs", "large": "gl"}, "cloudflare": {"small": "cs", "large": "cl"}, } async def test_router_routes_tier_to_correct_model(): groq = FakeProvider(result=ChatResult(content="ok", tool_calls=[], finish_reason="stop")) router = LLMRouter(providers={"groq": groq}, order=["groq"], models=_models()) await router.chat(messages=[{"role": "user", "content": "x"}], tools=[], tier="small") assert groq.last_model == "gs" await router.chat(messages=[{"role": "user", "content": "x"}], tools=[], tier="large") assert groq.last_model == "gl" async def test_router_failover_on_ratelimit(): groq = FakeProvider(raise_exc=RateLimitError()) cf = FakeProvider(result=ChatResult(content="ok", tool_calls=[], finish_reason="stop")) router = LLMRouter( providers={"groq": groq, "cloudflare": cf}, order=["groq", "cloudflare"], models=_models(), ) r = await router.chat(messages=[{"role": "user", "content": "x"}], tools=[], tier="small") assert r.content == "ok" assert groq.calls == 1 assert cf.calls == 1 assert cf.last_model == "cs" def test_build_router_multiple_groq_keys_and_cloudflare(): from types import SimpleNamespace from app.llm.router import build_router_from_settings s = SimpleNamespace( groq_api_key="k1", groq_api_keys=["k2", "k3", "k1"], # k1 deduped groq_base_url="https://api.groq.com/openai/v1", model_small="s", model_large="l", cloudflare_account_id="acct", cloudflare_api_token="tok", cloudflare_enabled=True, # CF is opt-in now (its live key was dead/401) cf_model_small="cs", cf_model_large="cl", ) router = build_router_from_settings(s) assert router.order == ["groq0", "groq1", "groq2", "cloudflare"] # 3 keys + CF assert len(router.providers) == 4 def test_build_router_with_three_providers(): from types import SimpleNamespace from app.llm.router import build_router_from_settings s = SimpleNamespace( groq_api_key="k1", groq_api_keys=[], groq_base_url="u", model_small="s", model_large="l", cloudflare_account_id="a", cloudflare_api_token="t", cloudflare_enabled=True, cf_model_small="cs", cf_model_large="cl", cerebras_api_key="cb", cerebras_base_url="https://api.cerebras.ai/v1", cerebras_model="gpt-oss-120b", ) router = build_router_from_settings(s) assert router.order == ["groq0", "cloudflare", "cerebras"] # Qwen off by default assert len(router.providers) == 3 def test_build_router_adds_qwen_when_enabled(): from types import SimpleNamespace from app.llm.router import build_router_from_settings s = SimpleNamespace( groq_api_key="k1", groq_api_keys=[], groq_base_url="u", model_small="s", model_large="l", cloudflare_account_id="", cloudflare_api_token="", cerebras_api_key="cb", cerebras_base_url="https://api.cerebras.ai/v1", cerebras_model="gpt-oss-120b", cerebras_qwen_enabled=True, ) router = build_router_from_settings(s) assert router.order == ["groq0", "cerebras", "cerebras_qwen"] assert router.models["cerebras_qwen"]["large"] == "qwen-3-235b-a22b" def test_build_router_full_cascade_with_extra_free_providers(): from types import SimpleNamespace from app.llm.router import build_router_from_settings s = SimpleNamespace( groq_api_key="k1", groq_api_keys=[], groq_base_url="u", model_small="s", model_large="l", cloudflare_account_id="a", cloudflare_api_token="t", cloudflare_enabled=True, cf_model_small="cs", cf_model_large="cl", cerebras_api_key="cb", cerebras_base_url="https://api.cerebras.ai/v1", cerebras_model="m", mistral_api_key="mk", mistral_base_url="https://api.mistral.ai/v1", mistral_model="mistral-small-latest", sambanova_api_key="sk", sambanova_base_url="https://api.sambanova.ai/v1", sambanova_model="Meta-Llama-3.3-70B-Instruct", openrouter_api_key="ok", openrouter_base_url="https://openrouter.ai/api/v1", openrouter_model="x:free", ) router = build_router_from_settings(s) # full failover cascade (Qwen3 fallback off by default; CF enabled here) assert router.order == ["groq0", "cloudflare", "cerebras", "sambanova", "mistral", "openrouter"] assert router.models["mistral"]["large"] == "mistral-small-latest" assert router.models["sambanova"]["large"] == "Meta-Llama-3.3-70B-Instruct" def test_multilingual_order_puts_strong_providers_first(): router = LLMRouter( providers={}, order=["groq0", "cloudflare", "cerebras", "sambanova", "mistral", "openrouter"], models={}, ) assert router.multilingual_order() == [ "groq0", "cerebras", # gpt-oss-120b first (best at non-Latin scripts) "cloudflare", "sambanova", "mistral", "openrouter", ] async def test_router_chat_honours_order_override(): a = FakeProvider(result=ChatResult(content="A", tool_calls=[], finish_reason="stop")) b = FakeProvider(result=ChatResult(content="B", tool_calls=[], finish_reason="stop")) router = LLMRouter( providers={"a": a, "b": b}, order=["a", "b"], models={"a": {"large": "m"}, "b": {"large": "m"}}, ) r = await router.chat(messages=[{"role": "user", "content": "x"}], tools=[], tier="large", order=["b", "a"]) assert r.content == "B" # override tried b first async def test_router_tags_provider_name(): groq = FakeProvider(result=ChatResult(content="ok", tool_calls=[], finish_reason="stop")) router = LLMRouter(providers={"groq0": groq}, order=["groq0"], models={"groq0": {"large": "l"}}) r = await router.chat(messages=[{"role": "user", "content": "x"}], tools=[], tier="large") assert r.provider == "groq0" async def test_router_raises_when_all_fail(): groq = FakeProvider(raise_exc=ProviderError("g")) cf = FakeProvider(raise_exc=ProviderError("c")) router = LLMRouter( providers={"groq": groq, "cloudflare": cf}, order=["groq", "cloudflare"], models=_models(), ) with pytest.raises(ProviderError): await router.chat(messages=[{"role": "user", "content": "x"}], tools=[], tier="large")