[agents] Add Groq and Mistral agent providers
This commit is contained in:
@ -10,9 +10,13 @@ from agents.providers import (
|
||||
agent_prompt,
|
||||
friendly_error_message,
|
||||
gemini_agent_prompt,
|
||||
groq_agent_prompt,
|
||||
is_transient_error,
|
||||
list_gemini_agent_models,
|
||||
list_groq_agent_models,
|
||||
list_mistral_agent_models,
|
||||
list_openrouter_free_models,
|
||||
mistral_agent_prompt,
|
||||
opencode_agent_prompt,
|
||||
openrouter_agent_prompt,
|
||||
retry_after_from,
|
||||
@ -238,6 +242,24 @@ def test_agent_prompt_openrouter_dispatches(settings):
|
||||
m.assert_called_once_with("hi", history=None, model="free/model-a")
|
||||
|
||||
|
||||
def test_agent_prompt_groq_dispatches(settings):
|
||||
settings.AGENT_PROVIDER = "groq"
|
||||
settings.GROQ_API_KEY = "gsk-test"
|
||||
with patch("agents.providers.groq_agent_prompt", return_value={"text": "x"}) as m:
|
||||
agent_prompt("hi", model="llama-3.3-70b-versatile")
|
||||
m.assert_called_once_with("hi", history=None, model="llama-3.3-70b-versatile")
|
||||
|
||||
|
||||
def test_agent_prompt_mistral_dispatches(settings):
|
||||
settings.AGENT_PROVIDER = "mistral"
|
||||
settings.MISTRAL_API_KEY = "msk-test"
|
||||
with patch(
|
||||
"agents.providers.mistral_agent_prompt", return_value={"text": "x"}
|
||||
) as m:
|
||||
agent_prompt("hi", model="mistral-small-latest")
|
||||
m.assert_called_once_with("hi", history=None, model="mistral-small-latest")
|
||||
|
||||
|
||||
def test_openrouter_agent_prompt(settings):
|
||||
settings.OPENROUTER_API_KEY = "sk-test"
|
||||
settings.LLM_MODEL = "some/model"
|
||||
@ -297,6 +319,128 @@ def test_openrouter_agent_prompt_with_history(settings):
|
||||
]
|
||||
|
||||
|
||||
def test_groq_agent_prompt(settings):
|
||||
settings.GROQ_API_KEY = "gsk-test"
|
||||
settings.LLM_MODEL = "llama-3.3-70b-versatile"
|
||||
fake = MagicMock()
|
||||
fake.raise_for_status = MagicMock()
|
||||
fake.json.return_value = {"choices": [{"message": {"content": " the answer "}}]}
|
||||
with patch("agents.providers.httpx.post", return_value=fake) as mock_post:
|
||||
result = groq_agent_prompt("what is 2+2?")
|
||||
|
||||
mock_post.assert_called_once()
|
||||
assert (
|
||||
mock_post.call_args.args[0] == "https://api.groq.com/openai/v1/chat/completions"
|
||||
)
|
||||
kwargs = mock_post.call_args.kwargs
|
||||
assert kwargs["headers"] == {"Authorization": "Bearer gsk-test"}
|
||||
assert kwargs["json"] == {
|
||||
"model": "llama-3.3-70b-versatile",
|
||||
"messages": [{"role": "user", "content": "what is 2+2?"}],
|
||||
}
|
||||
assert result == {
|
||||
"text": "the answer",
|
||||
"provider": "groq",
|
||||
"model": "llama-3.3-70b-versatile",
|
||||
}
|
||||
|
||||
|
||||
def test_groq_agent_prompt_missing_key(settings):
|
||||
settings.GROQ_API_KEY = ""
|
||||
with pytest.raises(ValueError, match="VROBBLER_GROQ_API_KEY is not set"):
|
||||
groq_agent_prompt("hi")
|
||||
|
||||
|
||||
def test_groq_agent_prompt_bad_response(settings):
|
||||
settings.GROQ_API_KEY = "gsk-test"
|
||||
fake = MagicMock()
|
||||
fake.raise_for_status = MagicMock()
|
||||
fake.json.return_value = {"unexpected": True}
|
||||
with patch("agents.providers.httpx.post", return_value=fake):
|
||||
with pytest.raises(ValueError, match="Unexpected Groq response"):
|
||||
groq_agent_prompt("hi")
|
||||
|
||||
|
||||
def test_groq_agent_prompt_with_history(settings):
|
||||
settings.GROQ_API_KEY = "gsk-test"
|
||||
fake = MagicMock()
|
||||
fake.raise_for_status = MagicMock()
|
||||
fake.json.return_value = {"choices": [{"message": {"content": "12"}}]}
|
||||
history = [
|
||||
{"prompt": "what is 2+2?", "response": "4"},
|
||||
{"prompt": "no response yet", "response": None},
|
||||
]
|
||||
with patch("agents.providers.httpx.post", return_value=fake) as mock_post:
|
||||
groq_agent_prompt("times 3?", history=history, model="llama-3.3-70b-versatile")
|
||||
|
||||
messages = mock_post.call_args.kwargs["json"]["messages"]
|
||||
assert messages == [
|
||||
{"role": "user", "content": "what is 2+2?"},
|
||||
{"role": "assistant", "content": "4"},
|
||||
{"role": "user", "content": "times 3?"},
|
||||
]
|
||||
|
||||
|
||||
def test_mistral_agent_prompt(settings):
|
||||
settings.MISTRAL_API_KEY = "msk-test"
|
||||
settings.LLM_MODEL = "mistral-small-latest"
|
||||
fake = MagicMock()
|
||||
fake.raise_for_status = MagicMock()
|
||||
fake.json.return_value = {"choices": [{"message": {"content": " the answer "}}]}
|
||||
with patch("agents.providers.httpx.post", return_value=fake) as mock_post:
|
||||
result = mistral_agent_prompt("what is 2+2?")
|
||||
|
||||
mock_post.assert_called_once()
|
||||
assert mock_post.call_args.args[0] == "https://api.mistral.ai/v1/chat/completions"
|
||||
kwargs = mock_post.call_args.kwargs
|
||||
assert kwargs["headers"] == {"Authorization": "Bearer msk-test"}
|
||||
assert kwargs["json"] == {
|
||||
"model": "mistral-small-latest",
|
||||
"messages": [{"role": "user", "content": "what is 2+2?"}],
|
||||
}
|
||||
assert result == {
|
||||
"text": "the answer",
|
||||
"provider": "mistral",
|
||||
"model": "mistral-small-latest",
|
||||
}
|
||||
|
||||
|
||||
def test_mistral_agent_prompt_missing_key(settings):
|
||||
settings.MISTRAL_API_KEY = ""
|
||||
with pytest.raises(ValueError, match="VROBBLER_MISTRAL_API_KEY is not set"):
|
||||
mistral_agent_prompt("hi")
|
||||
|
||||
|
||||
def test_mistral_agent_prompt_bad_response(settings):
|
||||
settings.MISTRAL_API_KEY = "msk-test"
|
||||
fake = MagicMock()
|
||||
fake.raise_for_status = MagicMock()
|
||||
fake.json.return_value = {"unexpected": True}
|
||||
with patch("agents.providers.httpx.post", return_value=fake):
|
||||
with pytest.raises(ValueError, match="Unexpected Mistral response"):
|
||||
mistral_agent_prompt("hi")
|
||||
|
||||
|
||||
def test_mistral_agent_prompt_with_history(settings):
|
||||
settings.MISTRAL_API_KEY = "msk-test"
|
||||
fake = MagicMock()
|
||||
fake.raise_for_status = MagicMock()
|
||||
fake.json.return_value = {"choices": [{"message": {"content": "12"}}]}
|
||||
history = [
|
||||
{"prompt": "what is 2+2?", "response": "4"},
|
||||
{"prompt": "no response yet", "response": None},
|
||||
]
|
||||
with patch("agents.providers.httpx.post", return_value=fake) as mock_post:
|
||||
mistral_agent_prompt("times 3?", history=history, model="mistral-small-latest")
|
||||
|
||||
messages = mock_post.call_args.kwargs["json"]["messages"]
|
||||
assert messages == [
|
||||
{"role": "user", "content": "what is 2+2?"},
|
||||
{"role": "assistant", "content": "4"},
|
||||
{"role": "user", "content": "times 3?"},
|
||||
]
|
||||
|
||||
|
||||
def test_list_openrouter_free_models_filters_free():
|
||||
fake = MagicMock()
|
||||
fake.raise_for_status = MagicMock()
|
||||
@ -363,6 +507,50 @@ def test_list_gemini_agent_models_no_key(settings):
|
||||
assert list_gemini_agent_models() == []
|
||||
|
||||
|
||||
def test_list_groq_agent_models(settings):
|
||||
settings.GROQ_API_KEY = "gsk-test"
|
||||
settings.GROQ_AGENT_MODELS = "llama-3.3-70b-versatile, openai/gpt-oss-20b"
|
||||
assert list_groq_agent_models() == [
|
||||
{
|
||||
"provider": "groq",
|
||||
"id": "llama-3.3-70b-versatile",
|
||||
"name": "llama-3.3-70b-versatile",
|
||||
},
|
||||
{
|
||||
"provider": "groq",
|
||||
"id": "openai/gpt-oss-20b",
|
||||
"name": "openai/gpt-oss-20b",
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
def test_list_groq_agent_models_no_key(settings):
|
||||
settings.GROQ_API_KEY = ""
|
||||
assert list_groq_agent_models() == []
|
||||
|
||||
|
||||
def test_list_mistral_agent_models(settings):
|
||||
settings.MISTRAL_API_KEY = "msk-test"
|
||||
settings.MISTRAL_AGENT_MODELS = "mistral-small-latest, codestral-latest"
|
||||
assert list_mistral_agent_models() == [
|
||||
{
|
||||
"provider": "mistral",
|
||||
"id": "mistral-small-latest",
|
||||
"name": "mistral-small-latest",
|
||||
},
|
||||
{
|
||||
"provider": "mistral",
|
||||
"id": "codestral-latest",
|
||||
"name": "codestral-latest",
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
def test_list_mistral_agent_models_no_key(settings):
|
||||
settings.MISTRAL_API_KEY = ""
|
||||
assert list_mistral_agent_models() == []
|
||||
|
||||
|
||||
# --- retry helpers ---
|
||||
|
||||
|
||||
@ -928,10 +1116,28 @@ def test_manual_scrobble_view_routes_plain_text_to_agent(mock_delay, client, use
|
||||
|
||||
@patch("agents.views.list_openrouter_free_models")
|
||||
@patch("agents.views.list_gemini_agent_models")
|
||||
def test_agent_model_select_view_lists_models(mock_gemini, mock_or, client, user):
|
||||
@patch("agents.views.list_groq_agent_models")
|
||||
@patch("agents.views.list_mistral_agent_models")
|
||||
def test_agent_model_select_view_lists_models(
|
||||
mock_mistral, mock_groq, mock_gemini, mock_or, client, user
|
||||
):
|
||||
mock_gemini.return_value = [
|
||||
{"provider": "gemini", "id": "gemini-3.6-flash", "name": "gemini-3.6-flash"}
|
||||
]
|
||||
mock_groq.return_value = [
|
||||
{
|
||||
"provider": "groq",
|
||||
"id": "llama-3.3-70b-versatile",
|
||||
"name": "llama-3.3-70b-versatile",
|
||||
}
|
||||
]
|
||||
mock_mistral.return_value = [
|
||||
{
|
||||
"provider": "mistral",
|
||||
"id": "mistral-small-latest",
|
||||
"name": "mistral-small-latest",
|
||||
}
|
||||
]
|
||||
mock_or.return_value = [
|
||||
{"provider": "openrouter", "id": "free/model-a", "name": "Model A"},
|
||||
{"provider": "openrouter", "id": "free/model-b", "name": "Model B"},
|
||||
@ -948,14 +1154,29 @@ def test_agent_model_select_view_lists_models(mock_gemini, mock_or, client, user
|
||||
assert b"Model A" in response.content
|
||||
assert b"gemini-3.6-flash" in response.content
|
||||
assert b"Google Gemini" in response.content
|
||||
expected = mock_gemini.return_value + mock_or.return_value
|
||||
assert b"llama-3.3-70b-versatile" in response.content
|
||||
assert b"Groq" in response.content
|
||||
assert b"mistral-small-latest" in response.content
|
||||
assert b"Mistral" in response.content
|
||||
expected = (
|
||||
mock_gemini.return_value
|
||||
+ mock_or.return_value
|
||||
+ mock_groq.return_value
|
||||
+ mock_mistral.return_value
|
||||
)
|
||||
assert client.session["agent_models"] == expected
|
||||
|
||||
|
||||
@patch("agents.views.list_openrouter_free_models")
|
||||
@patch("agents.views.list_gemini_agent_models")
|
||||
def test_agent_model_select_view_handles_errors(mock_gemini, mock_or, client, user):
|
||||
@patch("agents.views.list_groq_agent_models")
|
||||
@patch("agents.views.list_mistral_agent_models")
|
||||
def test_agent_model_select_view_handles_errors(
|
||||
mock_mistral, mock_groq, mock_gemini, mock_or, client, user
|
||||
):
|
||||
mock_gemini.return_value = []
|
||||
mock_groq.return_value = []
|
||||
mock_mistral.return_value = []
|
||||
mock_or.side_effect = ValueError("boom")
|
||||
client.force_login(user)
|
||||
session = client.session
|
||||
|
||||
Reference in New Issue
Block a user