SaarthiAI / tests /test_providers.py
parthmax24's picture
working proto 5
8b96826
Raw
History Blame Contribute Delete
5.72 kB
"""Provider fallback chain: Gemini first, Groq on failure, error if both die."""
import pytest
from app.providers import gemini, groq, llm
from app.providers.gemini import ProviderError
from app.providers.llm import LLMUnavailable
MESSAGES = [{"role": "user", "content": "hello"}]
def test_gemini_success_is_used(monkeypatch, fake_keys):
monkeypatch.setattr(
gemini, "chat", lambda *a, **k: {"text": "from gemini", "tool_calls": [], "provider": "gemini"}
)
result = llm.chat(MESSAGES)
assert result["provider"] == "gemini"
assert result["text"] == "from gemini"
def test_falls_back_to_groq_when_gemini_fails(monkeypatch, fake_keys):
def gemini_fails(*args, **kwargs):
raise ProviderError("rate limited")
monkeypatch.setattr(gemini, "chat", gemini_fails)
monkeypatch.setattr(
groq, "chat", lambda *a, **k: {"text": "from groq", "tool_calls": [], "provider": "groq"}
)
result = llm.chat(MESSAGES)
assert result["provider"] == "groq"
def test_raises_when_both_fail(monkeypatch, fake_keys):
def fails(*args, **kwargs):
raise ProviderError("down")
monkeypatch.setattr(gemini, "chat", fails)
monkeypatch.setattr(groq, "chat", fails)
with pytest.raises(LLMUnavailable):
llm.chat(MESSAGES)
def test_raises_when_no_keys(no_llm_keys):
assert llm.available() is False
with pytest.raises(LLMUnavailable):
llm.chat(MESSAGES)
def test_skips_gemini_without_key(monkeypatch):
from app import config
monkeypatch.setattr(config, "gemini_key", lambda: None)
monkeypatch.setattr(config, "groq_key", lambda: "fake-groq")
monkeypatch.setattr(
groq, "chat", lambda *a, **k: {"text": "groq only", "tool_calls": [], "provider": "groq"}
)
assert llm.chat(MESSAGES)["provider"] == "groq"
def test_gemini_second_model_used_when_first_fails(monkeypatch, fake_keys):
"""gemini-2.5-flash rate-limited -> gemini-2.5-flash-lite answers."""
attempts = []
def gemini_chat(messages, model=None, **kwargs):
attempts.append(model)
if model == "gemini-2.5-flash":
raise ProviderError("429 rate limited")
return {"text": "from lite", "tool_calls": [], "provider": "gemini"}
monkeypatch.setattr(gemini, "chat", gemini_chat)
result = llm.chat(MESSAGES)
assert result["text"] == "from lite"
assert attempts == ["gemini-2.5-flash", "gemini-2.5-flash-lite"]
def test_full_chain_order_gemini_then_groq(monkeypatch, fake_keys):
"""All four (provider, model) pairs are tried in the documented order."""
attempts = []
def failing(provider):
def chat(messages, model=None, **kwargs):
attempts.append(f"{provider}/{model}")
raise ProviderError("down")
return chat
monkeypatch.setattr(gemini, "chat", failing("gemini"))
monkeypatch.setattr(groq, "chat", failing("groq"))
with pytest.raises(LLMUnavailable):
llm.chat(MESSAGES)
assert attempts == [
"gemini/gemini-2.5-flash",
"gemini/gemini-2.5-flash-lite",
"groq/llama-3.3-70b-versatile",
"groq/llama-3.1-8b-instant",
]
# ---- adapter wire-format checks ------------------------------------------------
def test_gemini_parses_function_call(monkeypatch, fake_response, fake_keys):
payload = {
"candidates": [
{
"content": {
"parts": [
{"text": "Let me check."},
{"functionCall": {"name": "get_weather", "args": {"lat": 26.8}}},
]
}
}
]
}
monkeypatch.setattr(gemini.requests, "post", lambda *a, **k: fake_response(payload))
result = gemini.chat(MESSAGES, tools=[{"name": "get_weather", "description": "x", "parameters": {}}])
assert result["tool_calls"] == [{"name": "get_weather", "args": {"lat": 26.8}}]
assert result["text"] == "Let me check."
def test_gemini_http_error_raises_provider_error(monkeypatch, fake_response, fake_keys):
monkeypatch.setattr(
gemini.requests, "post", lambda *a, **k: fake_response({}, status_code=429, text="quota")
)
with pytest.raises(ProviderError, match="429"):
gemini.chat(MESSAGES)
def test_groq_parses_tool_calls(monkeypatch, fake_response, fake_keys):
payload = {
"choices": [
{
"message": {
"content": None,
"tool_calls": [
{
"id": "call_1",
"function": {"name": "geocode_place", "arguments": '{"place": "Hazratganj"}'},
}
],
}
}
]
}
monkeypatch.setattr(groq.requests, "post", lambda *a, **k: fake_response(payload))
result = groq.chat(MESSAGES, tools=[{"name": "geocode_place", "description": "x", "parameters": {}}])
assert result["tool_calls"] == [{"name": "geocode_place", "args": {"place": "Hazratganj"}}]
def test_groq_message_conversion_handles_tool_roundtrip():
messages = [
{"role": "system", "content": "sys"},
{"role": "user", "content": "q"},
{"role": "assistant", "content": "", "tool_calls": [{"name": "get_weather", "args": {"lat": 1}}]},
{"role": "tool", "name": "get_weather", "content": '{"rain": false}'},
]
converted = groq._to_openai_messages(messages)
assert converted[0]["role"] == "system"
assert converted[2]["tool_calls"][0]["function"]["name"] == "get_weather"
assert converted[3]["role"] == "tool"
assert converted[3]["tool_call_id"] == converted[2]["tool_calls"][0]["id"]