[agents] Swap agent flow to OpenRouter with model selection
This commit is contained in:
@ -8,7 +8,9 @@ from agents.providers import (
|
||||
_parse_opencode_events,
|
||||
agent_prompt,
|
||||
gemini_agent_prompt,
|
||||
list_openrouter_free_models,
|
||||
opencode_agent_prompt,
|
||||
openrouter_agent_prompt,
|
||||
)
|
||||
from django.conf import settings
|
||||
from django.contrib.auth import get_user_model
|
||||
@ -35,7 +37,7 @@ def user(db):
|
||||
|
||||
def _mk_scrobble(user, *, in_progress=True, turns=None, title="First prompt"):
|
||||
agent_session, _ = AgentSession.find_or_create(
|
||||
provider="gemini", model=settings.LLM_MODEL
|
||||
provider=settings.AGENT_PROVIDER, model=settings.LLM_MODEL
|
||||
)
|
||||
scrobble = Scrobble.create_or_update(
|
||||
agent_session,
|
||||
@ -49,7 +51,7 @@ def _mk_scrobble(user, *, in_progress=True, turns=None, title="First prompt"):
|
||||
skip_in_progress_check=True,
|
||||
)
|
||||
log = scrobble.log if isinstance(scrobble.log, dict) else {}
|
||||
log["provider"] = "gemini"
|
||||
log["provider"] = settings.AGENT_PROVIDER
|
||||
log["model"] = settings.LLM_MODEL
|
||||
log["title"] = title
|
||||
log["turns"] = turns or []
|
||||
@ -214,7 +216,120 @@ def test_agent_prompt_defaults_to_configured_provider(settings):
|
||||
settings.GOOGLE_AI_API_KEY = "test-key"
|
||||
with patch("agents.providers.gemini_agent_prompt", return_value={"text": "x"}) as m:
|
||||
agent_prompt("hi")
|
||||
m.assert_called_once_with("hi", history=None)
|
||||
m.assert_called_once_with("hi", history=None, model=None)
|
||||
|
||||
|
||||
def test_agent_prompt_openrouter_dispatches(settings):
|
||||
settings.AGENT_PROVIDER = "openrouter"
|
||||
settings.OPENROUTER_API_KEY = "sk-test"
|
||||
with patch(
|
||||
"agents.providers.openrouter_agent_prompt", return_value={"text": "x"}
|
||||
) as m:
|
||||
agent_prompt("hi", model="free/model-a")
|
||||
m.assert_called_once_with("hi", history=None, model="free/model-a")
|
||||
|
||||
|
||||
def test_openrouter_agent_prompt(settings):
|
||||
settings.OPENROUTER_API_KEY = "sk-test"
|
||||
settings.LLM_MODEL = "some/model"
|
||||
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 = openrouter_agent_prompt("what is 2+2?")
|
||||
|
||||
mock_post.assert_called_once()
|
||||
kwargs = mock_post.call_args.kwargs
|
||||
assert kwargs["headers"] == {"Authorization": "Bearer sk-test"}
|
||||
assert kwargs["json"] == {
|
||||
"model": "some/model",
|
||||
"messages": [{"role": "user", "content": "what is 2+2?"}],
|
||||
}
|
||||
assert result == {
|
||||
"text": "the answer",
|
||||
"provider": "openrouter",
|
||||
"model": "some/model",
|
||||
}
|
||||
|
||||
|
||||
def test_openrouter_agent_prompt_missing_key(settings):
|
||||
settings.OPENROUTER_API_KEY = ""
|
||||
with pytest.raises(ValueError, match="VROBBLER_OPENROUTER_API_KEY is not set"):
|
||||
openrouter_agent_prompt("hi")
|
||||
|
||||
|
||||
def test_openrouter_agent_prompt_bad_response(settings):
|
||||
settings.OPENROUTER_API_KEY = "sk-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 OpenRouter response"):
|
||||
openrouter_agent_prompt("hi")
|
||||
|
||||
|
||||
def test_openrouter_agent_prompt_with_history(settings):
|
||||
settings.OPENROUTER_API_KEY = "sk-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:
|
||||
openrouter_agent_prompt("times 3?", history=history, model="some/model")
|
||||
|
||||
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()
|
||||
fake.json.return_value = {
|
||||
"data": [
|
||||
{
|
||||
"id": "paid/model",
|
||||
"name": "Paid",
|
||||
"pricing": {"prompt": "0.001", "completion": "0.002"},
|
||||
},
|
||||
{
|
||||
"id": "free/model-a",
|
||||
"name": "Bee",
|
||||
"pricing": {"prompt": "0", "completion": "0"},
|
||||
},
|
||||
{
|
||||
"id": "free/model-b",
|
||||
"name": "Alpha",
|
||||
"pricing": {"prompt": "0", "completion": "0"},
|
||||
},
|
||||
]
|
||||
}
|
||||
with patch("agents.providers.httpx.get", return_value=fake) as mock_get:
|
||||
models = list_openrouter_free_models()
|
||||
|
||||
mock_get.assert_called_once()
|
||||
assert models == [
|
||||
{"id": "free/model-b", "name": "Alpha"},
|
||||
{"id": "free/model-a", "name": "Bee"},
|
||||
]
|
||||
|
||||
|
||||
def test_list_openrouter_free_models_skips_missing_pricing():
|
||||
fake = MagicMock()
|
||||
fake.raise_for_status = MagicMock()
|
||||
fake.json.return_value = {
|
||||
"data": [
|
||||
{"id": "no-pricing/model", "name": "No Pricing"},
|
||||
]
|
||||
}
|
||||
with patch("agents.providers.httpx.get", return_value=fake):
|
||||
assert list_openrouter_free_models() == []
|
||||
|
||||
|
||||
# --- scrobbler ---
|
||||
@ -227,9 +342,12 @@ def test_manual_scrobble_agent_session_creates_new_session(mock_delay, user):
|
||||
assert Scrobble.objects.filter(id=scrobble.id).exists()
|
||||
assert scrobble.media_type == Scrobble.MediaType.AGENT_SESSION
|
||||
assert scrobble.in_progress is True
|
||||
assert scrobble.agent_session.provider == "gemini"
|
||||
assert scrobble.agent_session.provider == settings.AGENT_PROVIDER
|
||||
assert scrobble.agent_session.model == settings.LLM_MODEL
|
||||
assert scrobble.agent_session.title == f"gemini {settings.LLM_MODEL}"
|
||||
assert (
|
||||
scrobble.agent_session.title
|
||||
== f"{settings.AGENT_PROVIDER} {settings.LLM_MODEL}"
|
||||
)
|
||||
assert scrobble.log["title"] == "What is the capital of France?"
|
||||
turns = scrobble.log["turns"]
|
||||
assert len(turns) == 1
|
||||
@ -241,14 +359,14 @@ def test_manual_scrobble_agent_session_creates_new_session(mock_delay, user):
|
||||
@patch("scrobbles.tasks.scrobble_agent_session_prompt.delay")
|
||||
def test_manual_scrobble_agent_session_reuses_agent_session(mock_delay, user):
|
||||
agent_session, created = AgentSession.find_or_create(
|
||||
provider="gemini", model=settings.LLM_MODEL
|
||||
provider=settings.AGENT_PROVIDER, model=settings.LLM_MODEL
|
||||
)
|
||||
assert created is True
|
||||
|
||||
manual_scrobble_agent_session("First question?", user.id)
|
||||
|
||||
agent_session_again, created = AgentSession.find_or_create(
|
||||
provider="gemini", model=settings.LLM_MODEL
|
||||
provider=settings.AGENT_PROVIDER, model=settings.LLM_MODEL
|
||||
)
|
||||
assert created is False
|
||||
assert agent_session_again.id == agent_session.id
|
||||
@ -317,16 +435,18 @@ def test_scrobble_agent_session_prompt_fills_turn(
|
||||
)
|
||||
mock_agent_prompt.return_value = {
|
||||
"text": "hi back",
|
||||
"provider": "gemini",
|
||||
"model": "gemini-2.5-flash",
|
||||
"provider": settings.AGENT_PROVIDER,
|
||||
"model": settings.LLM_MODEL,
|
||||
}
|
||||
|
||||
scrobble_agent_session_prompt(scrobble.id, "abc-123")
|
||||
|
||||
scrobble.refresh_from_db()
|
||||
mock_agent_prompt.assert_called_once_with("hello", provider="gemini", history=[])
|
||||
mock_agent_prompt.assert_called_once_with(
|
||||
"hello", provider=settings.AGENT_PROVIDER, history=[], model=settings.LLM_MODEL
|
||||
)
|
||||
assert scrobble.log["turns"][0]["response"] == "hi back"
|
||||
assert scrobble.log["provider"] == "gemini"
|
||||
assert scrobble.log["provider"] == settings.AGENT_PROVIDER
|
||||
assert scrobble.in_progress is True
|
||||
assert scrobble.played_to_completion is False
|
||||
mock_complete.assert_called_once_with(
|
||||
@ -363,15 +483,16 @@ def test_scrobble_agent_session_prompt_passes_history(
|
||||
|
||||
mock_agent_prompt.return_value = {
|
||||
"text": "12",
|
||||
"provider": "gemini",
|
||||
"provider": settings.AGENT_PROVIDER,
|
||||
"model": settings.LLM_MODEL,
|
||||
}
|
||||
scrobble_agent_session_prompt(scrobble.id, "cur-1")
|
||||
|
||||
mock_agent_prompt.assert_called_once_with(
|
||||
"and times 3?",
|
||||
provider="gemini",
|
||||
provider=settings.AGENT_PROVIDER,
|
||||
history=[{"prompt": "what is 2+2?", "response": "4"}],
|
||||
model=settings.LLM_MODEL,
|
||||
)
|
||||
|
||||
|
||||
@ -452,11 +573,90 @@ def test_manual_scrobble_view_routes_plain_text_to_agent(mock_delay, client, use
|
||||
{"item_id": "what is the meaning of life?"},
|
||||
)
|
||||
|
||||
assert response.status_code == 302
|
||||
assert response.url == reverse("agents:agent_model_select")
|
||||
assert client.session["agent_prompt"] == "what is the meaning of life?"
|
||||
assert not Scrobble.objects.filter(user_id=user.id).exists()
|
||||
|
||||
|
||||
@patch("agents.views.list_openrouter_free_models")
|
||||
def test_agent_model_select_view_lists_models(mock_models, client, user):
|
||||
mock_models.return_value = [
|
||||
{"id": "free/model-a", "name": "Model A"},
|
||||
{"id": "free/model-b", "name": "Model B"},
|
||||
]
|
||||
client.force_login(user)
|
||||
session = client.session
|
||||
session["agent_prompt"] = "what is 2+2?"
|
||||
session.save()
|
||||
response = client.get(reverse("agents:agent_model_select"))
|
||||
|
||||
assert response.status_code == 200
|
||||
assert b"what is 2+2?" in response.content
|
||||
assert b"free/model-a" in response.content
|
||||
assert b"Model A" in response.content
|
||||
assert client.session["agent_models"] == mock_models.return_value
|
||||
|
||||
|
||||
@patch("agents.views.list_openrouter_free_models")
|
||||
def test_agent_model_select_view_handles_errors(mock_models, client, user):
|
||||
mock_models.side_effect = ValueError("boom")
|
||||
client.force_login(user)
|
||||
session = client.session
|
||||
session["agent_prompt"] = "hi"
|
||||
session.save()
|
||||
response = client.get(reverse("agents:agent_model_select"))
|
||||
|
||||
assert response.status_code == 200
|
||||
assert b"Could not load models from OpenRouter" in response.content
|
||||
|
||||
|
||||
@patch("agents.views.list_openrouter_free_models")
|
||||
@patch("scrobbles.tasks.scrobble_agent_session_prompt.delay")
|
||||
def test_agent_scrobble_from_model_creates_session(
|
||||
mock_delay, mock_models, client, user
|
||||
):
|
||||
client.force_login(user)
|
||||
session = client.session
|
||||
session["agent_prompt"] = "what is 2+2?"
|
||||
session.save()
|
||||
response = client.post(
|
||||
reverse("agents:agent_scrobble_from_model"),
|
||||
{"model": "free/model-a", "prompt": "what is 2+2?"},
|
||||
)
|
||||
|
||||
assert response.status_code == 302
|
||||
scrobble = Scrobble.objects.filter(user_id=user.id).first()
|
||||
assert scrobble is not None
|
||||
assert scrobble.media_type == Scrobble.MediaType.AGENT_SESSION
|
||||
assert response.url.startswith(reverse("scrobbles:detail", args=[scrobble.id]))
|
||||
assert scrobble.agent_session.provider == "openrouter"
|
||||
assert scrobble.agent_session.model == "free/model-a"
|
||||
assert scrobble.log["title"] == "what is 2+2?"
|
||||
assert response.url == scrobble.get_absolute_url()
|
||||
assert "agent_prompt" not in client.session
|
||||
|
||||
|
||||
@patch("agents.views.list_openrouter_free_models")
|
||||
@patch("scrobbles.tasks.scrobble_agent_session_prompt.delay")
|
||||
def test_agent_scrobble_from_model_missing_model(mock_delay, mock_models, client, user):
|
||||
client.force_login(user)
|
||||
response = client.post(reverse("agents:agent_scrobble_from_model"), {"model": ""})
|
||||
|
||||
assert response.status_code == 302
|
||||
assert response.url == reverse("agents:agent_model_select")
|
||||
mock_delay.assert_not_called()
|
||||
assert not Scrobble.objects.filter(user_id=user.id).exists()
|
||||
|
||||
|
||||
@patch("scrobbles.tasks.scrobble_agent_session_prompt.delay")
|
||||
def test_manual_scrobble_agent_session_explicit_model(mock_delay, user):
|
||||
scrobble = manual_scrobble_agent_session(
|
||||
"hi", user.id, provider="openrouter", model="free/model-a"
|
||||
)
|
||||
|
||||
assert scrobble.agent_session.provider == "openrouter"
|
||||
assert scrobble.agent_session.model == "free/model-a"
|
||||
assert scrobble.agent_session.title == "openrouter free/model-a"
|
||||
assert scrobble.log["turns"][0]["prompt"] == "hi"
|
||||
|
||||
|
||||
@patch("scrobbles.tasks.scrobble_agent_session_prompt.delay")
|
||||
|
||||
Reference in New Issue
Block a user