Files
vrobbler/tests/agents_tests/test_agent_sessions.py
Colin Powell 0c92174440
All checks were successful
ci / test (push) Successful in 3m2s
ci / build-and-deploy (push) Has been skipped
[agents] Add agent-session scrobbling for Gemini and opencode providers
2026-07-31 00:21:24 -04:00

593 lines
19 KiB
Python

import json
import subprocess
from unittest.mock import MagicMock, patch
import pytest
from agents.models import AgentSession, AgentSessionLogData
from agents.providers import (
_parse_opencode_events,
agent_prompt,
gemini_agent_prompt,
opencode_agent_prompt,
)
from django.conf import settings
from django.contrib.auth import get_user_model
from django.urls import reverse
from django.utils import timezone
from scrobbles.models import Scrobble
from scrobbles.scrobblers import (
manual_scrobble_agent_follow_up,
manual_scrobble_agent_session,
)
from scrobbles.tasks import scrobble_agent_session_prompt
User = get_user_model()
@pytest.fixture
def user(db):
return User.objects.create_user(username="agentuser", password="pw")
def _mk_scrobble(user, *, in_progress=True, turns=None, title="First prompt"):
agent_session, _ = AgentSession.find_or_create(
provider="gemini", model=settings.LLM_MODEL
)
scrobble = Scrobble.create_or_update(
agent_session,
user.id,
{
"user_id": user.id,
"timestamp": timezone.now(),
"playback_position_seconds": 0,
"source": "Vrobbler",
},
skip_in_progress_check=True,
)
log = scrobble.log if isinstance(scrobble.log, dict) else {}
log["provider"] = "gemini"
log["model"] = settings.LLM_MODEL
log["title"] = title
log["turns"] = turns or []
scrobble.log = log
scrobble.in_progress = in_progress
scrobble.save(update_fields=["log", "in_progress"])
return scrobble
# --- providers ---
def test_parse_opencode_events_joins_text_parts():
raw = "\n".join(
[
json.dumps({"type": "step_start"}),
json.dumps({"type": "text", "part": {"type": "text", "text": "hello"}}),
json.dumps({"type": "text", "part": {"type": "text", "text": " world"}}),
json.dumps({"type": "step_finish"}),
"not-json",
]
)
assert _parse_opencode_events(raw) == "hello\n world"
def test_opencode_agent_prompt(settings):
settings.OPENCODE_CMD = "opencode"
events = json.dumps(
{"type": "text", "part": {"type": "text", "text": "the answer"}}
)
with patch(
"agents.providers.subprocess.run",
return_value=MagicMock(returncode=0, stdout=events, stderr=""),
) as mock_run:
result = opencode_agent_prompt("what is 2+2?")
mock_run.assert_called_once()
cmd = mock_run.call_args.args[0]
assert cmd[:2] == ["opencode", "run"]
assert cmd[2:5] == ["--format", "json", "--dir"]
assert cmd[5].startswith("/tmp/")
assert cmd[6] == "what is 2+2?"
assert result == {"text": "the answer", "provider": "opencode", "model": "opencode"}
def test_opencode_agent_prompt_nonzero_exit():
with patch(
"agents.providers.subprocess.run",
return_value=MagicMock(returncode=1, stdout="", stderr="boom"),
):
with pytest.raises(ValueError, match="opencode exited with 1"):
opencode_agent_prompt("hi")
def test_opencode_agent_prompt_timeout():
with patch(
"agents.providers.subprocess.run",
side_effect=subprocess.TimeoutExpired("cmd", 900),
):
with pytest.raises(ValueError, match="timed out"):
opencode_agent_prompt("hi")
def test_opencode_agent_prompt_no_text_output():
with patch(
"agents.providers.subprocess.run",
return_value=MagicMock(
returncode=0, stdout=json.dumps({"type": "step_start"}), stderr=""
),
):
with pytest.raises(ValueError, match="no text output"):
opencode_agent_prompt("hi")
def test_gemini_agent_prompt(settings):
settings.GOOGLE_API_KEY = "test-key"
settings.LLM_MODEL = "gemini-2.5-flash"
fake = MagicMock()
fake.raise_for_status = MagicMock()
fake.json.return_value = {
"candidates": [{"content": {"parts": [{"text": " the answer "}]}}]
}
with patch("agents.providers.httpx.post", return_value=fake) as mock_post:
result = gemini_agent_prompt("what is 2+2?")
mock_post.assert_called_once()
kwargs = mock_post.call_args.kwargs
assert kwargs["params"] == {"key": "test-key"}
assert kwargs["json"] == {"contents": [{"parts": [{"text": "what is 2+2?"}]}]}
assert result == {
"text": "the answer",
"provider": "gemini",
"model": "gemini-2.5-flash",
}
def test_gemini_agent_prompt_missing_key(settings):
settings.GOOGLE_API_KEY = ""
with pytest.raises(ValueError, match="VROBBLER_GOOGLE_API_KEY is not set"):
gemini_agent_prompt("hi")
def test_gemini_agent_prompt_bad_response(settings):
settings.GOOGLE_API_KEY = "test-key"
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 Gemini response"):
gemini_agent_prompt("hi")
def test_gemini_agent_prompt_with_history(settings):
settings.GOOGLE_API_KEY = "test-key"
fake = MagicMock()
fake.raise_for_status = MagicMock()
fake.json.return_value = {"candidates": [{"content": {"parts": [{"text": "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:
gemini_agent_prompt("times 3?", history=history)
contents = mock_post.call_args.kwargs["json"]["contents"]
assert contents == [
{"parts": [{"text": "what is 2+2?"}]},
{"parts": [{"text": "4"}]},
{"parts": [{"text": "times 3?"}]},
]
def test_opencode_agent_prompt_with_history(settings):
settings.OPENCODE_CMD = "opencode"
with patch(
"agents.providers.subprocess.run",
return_value=MagicMock(
returncode=0,
stdout=json.dumps({"type": "text", "part": {"type": "text", "text": "x"}}),
stderr="",
),
) as mock_run:
opencode_agent_prompt(
"times 3?", history=[{"prompt": "what is 2+2?", "response": "4"}]
)
prompt = mock_run.call_args.args[0][6]
assert "Previous conversation:" in prompt
assert "User: what is 2+2?" in prompt
assert "Assistant: 4" in prompt
assert "times 3?" in prompt
def test_agent_prompt_unknown_provider(settings):
settings.AGENT_PROVIDER = "nope"
with pytest.raises(ValueError, match="Unknown agent provider"):
agent_prompt("hi")
def test_agent_prompt_defaults_to_configured_provider(settings):
settings.AGENT_PROVIDER = "gemini"
settings.GOOGLE_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)
# --- scrobbler ---
@patch("scrobbles.tasks.scrobble_agent_session_prompt.delay")
def test_manual_scrobble_agent_session_creates_new_session(mock_delay, user):
scrobble = manual_scrobble_agent_session("What is the capital of France?", user.id)
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.model == settings.LLM_MODEL
assert scrobble.agent_session.title == f"gemini {settings.LLM_MODEL}"
assert scrobble.log["title"] == "What is the capital of France?"
turns = scrobble.log["turns"]
assert len(turns) == 1
assert turns[0]["prompt"] == "What is the capital of France?"
assert turns[0]["response"] is None
mock_delay.assert_called_once_with(scrobble.id, turns[0]["prompt_id"])
@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
)
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
)
assert created is False
assert agent_session_again.id == agent_session.id
assert AgentSession.objects.count() == 1
@patch("scrobbles.tasks.scrobble_agent_session_prompt.delay")
def test_manual_scrobble_agent_session_appends_to_in_progress(mock_delay, user):
_mk_scrobble(
user,
in_progress=True,
turns=[
{"prompt_id": "first-1", "prompt": "First prompt", "response": "answer"}
],
)
scrobble = manual_scrobble_agent_session("And the second question?", user.id)
turns = scrobble.log["turns"]
assert len(turns) == 2
assert turns[0]["prompt"] == "First prompt"
assert turns[1]["prompt"] == "And the second question?"
assert turns[1]["response"] is None
assert scrobble.log["title"] == "First prompt"
@patch("scrobbles.tasks.scrobble_agent_session_prompt.delay")
def test_manual_scrobble_agent_session_new_session_after_completed(mock_delay, user):
_mk_scrobble(user, in_progress=False)
scrobble = manual_scrobble_agent_session("A new session?", user.id)
assert scrobble.media_type == Scrobble.MediaType.AGENT_SESSION
assert len(scrobble.log["turns"]) == 1
assert Scrobble.objects.filter(user_id=user.id).count() == 2
@patch("scrobbles.tasks.scrobble_agent_session_prompt.delay")
def test_manual_scrobble_agent_session_opencode_model_empty(mock_delay, user, settings):
settings.AGENT_PROVIDER = "opencode"
scrobble = manual_scrobble_agent_session("hi", user.id)
assert scrobble.agent_session.provider == "opencode"
assert scrobble.agent_session.model == "opencode"
assert scrobble.agent_session.title == "opencode opencode"
# --- celery task ---
@patch("agents.providers.agent_prompt")
def test_scrobble_agent_session_prompt_fills_turn(mock_agent_prompt, user):
scrobble = _mk_scrobble(
user,
turns=[
{
"prompt_id": "abc-123",
"prompt": "hello",
"response": None,
}
],
)
mock_agent_prompt.return_value = {
"text": "hi back",
"provider": "gemini",
"model": "gemini-2.5-flash",
}
scrobble_agent_session_prompt(scrobble.id, "abc-123")
scrobble.refresh_from_db()
mock_agent_prompt.assert_called_once_with("hello", provider="gemini", history=[])
assert scrobble.log["turns"][0]["response"] == "hi back"
assert scrobble.log["provider"] == "gemini"
assert scrobble.in_progress is False
assert scrobble.played_to_completion is True
@patch("agents.providers.agent_prompt")
def test_scrobble_agent_session_prompt_passes_history(mock_agent_prompt, user):
scrobble = _mk_scrobble(
user,
turns=[
{
"prompt_id": "first-1",
"prompt": "what is 2+2?",
"response": "4",
},
{
"prompt_id": "err-1",
"prompt": "boom?",
"response": "Error: boom",
"error": True,
},
{
"prompt_id": "cur-1",
"prompt": "and times 3?",
"response": None,
},
],
)
mock_agent_prompt.return_value = {
"text": "12",
"provider": "gemini",
"model": settings.LLM_MODEL,
}
scrobble_agent_session_prompt(scrobble.id, "cur-1")
mock_agent_prompt.assert_called_once_with(
"and times 3?",
provider="gemini",
history=[{"prompt": "what is 2+2?", "response": "4"}],
)
@patch("agents.providers.agent_prompt")
def test_scrobble_agent_session_prompt_stores_error(mock_agent_prompt, user):
scrobble = _mk_scrobble(
user,
turns=[
{
"prompt_id": "abc-123",
"prompt": "hello",
"response": None,
}
],
)
mock_agent_prompt.side_effect = ValueError("boom")
scrobble_agent_session_prompt(scrobble.id, "abc-123")
scrobble.refresh_from_db()
assert scrobble.log["turns"][0]["response"] == "Error: boom"
assert scrobble.log["turns"][0]["error"] is True
assert scrobble.in_progress is False
# --- views ---
@patch("scrobbles.tasks.scrobble_agent_session_prompt.delay")
def test_manual_scrobble_view_routes_plain_text_to_agent(mock_delay, client, user):
client.force_login(user)
response = client.post(
reverse("scrobbles:lookup-manual-scrobble"),
{"item_id": "what is the meaning of life?"},
)
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]))
@patch("scrobbles.tasks.scrobble_agent_session_prompt.delay")
def test_agent_session_partial_pending_polls(mock_delay, client, user):
scrobble = _mk_scrobble(
user,
turns=[{"prompt_id": "abc-123", "prompt": "hello", "response": None}],
)
client.force_login(user)
response = client.get(
reverse("scrobbles:agent-session-partial", args=[scrobble.id])
)
assert response.status_code == 200
assert b"every 5s" in response.content
assert b"Thinking" in response.content
@patch("scrobbles.tasks.scrobble_agent_session_prompt.delay")
def test_agent_session_partial_completed_stops_polling(mock_delay, client, user):
scrobble = _mk_scrobble(
user,
in_progress=False,
turns=[
{
"prompt_id": "abc-123",
"prompt": "hello",
"response": "**hi back**",
}
],
)
client.force_login(user)
response = client.get(
reverse("scrobbles:agent-session-partial", args=[scrobble.id])
)
assert response.status_code == 200
assert b"every 5s" not in response.content
assert b"hi back" in response.content
@patch("scrobbles.tasks.scrobble_agent_session_prompt.delay")
def test_agent_session_list_view(mock_delay, client, user):
_mk_scrobble(user, in_progress=True)
client.force_login(user)
response = client.get(reverse("agents:agent_session_list"))
assert response.status_code == 200
assert settings.LLM_MODEL.encode() in response.content
@patch("scrobbles.tasks.scrobble_agent_session_prompt.delay")
def test_agent_session_detail_view(mock_delay, client, user):
scrobble = _mk_scrobble(user, in_progress=True)
client.force_login(user)
response = client.get(scrobble.media_obj.get_absolute_url())
assert response.status_code == 200
assert settings.LLM_MODEL.encode() in response.content
@patch("scrobbles.tasks.scrobble_agent_session_prompt.delay")
def test_manual_scrobble_agent_follow_up_appends_turn(mock_delay, user):
scrobble = _mk_scrobble(
user,
in_progress=False,
turns=[
{
"prompt_id": "abc-123",
"prompt": "hello",
"response": "hi back",
}
],
)
result = manual_scrobble_agent_follow_up("tell me more", scrobble.id, user.id)
assert result.id == scrobble.id
scrobble.refresh_from_db()
turns = scrobble.log["turns"]
assert len(turns) == 2
assert turns[1]["prompt"] == "tell me more"
assert turns[1]["response"] is None
assert scrobble.in_progress is True
mock_delay.assert_called_once_with(scrobble.id, turns[1]["prompt_id"])
@patch("scrobbles.tasks.scrobble_agent_session_prompt.delay")
def test_manual_scrobble_agent_follow_up_unknown_scrobble(mock_delay, user):
assert manual_scrobble_agent_follow_up("tell me more", 999999, user.id) is None
mock_delay.assert_not_called()
@patch("scrobbles.tasks.scrobble_agent_session_prompt.delay")
def test_agent_session_followup_post_appends_turn(mock_delay, client, user):
scrobble = _mk_scrobble(
user,
in_progress=False,
turns=[
{
"prompt_id": "abc-123",
"prompt": "hello",
"response": "hi back",
}
],
)
client.force_login(user)
response = client.post(
reverse("scrobbles:agent-session-followup", args=[scrobble.id]),
{"prompt": "tell me more"},
)
assert response.status_code == 200
assert b"tell me more" in response.content
assert b"Thinking" in response.content
assert b"every 5s" in response.content
scrobble.refresh_from_db()
assert len(scrobble.log["turns"]) == 2
assert scrobble.in_progress is True
@patch("scrobbles.tasks.scrobble_agent_session_prompt.delay")
def test_agent_session_followup_post_empty_prompt(mock_delay, client, user):
scrobble = _mk_scrobble(user, in_progress=False)
client.force_login(user)
response = client.post(
reverse("scrobbles:agent-session-followup", args=[scrobble.id]),
{"prompt": " "},
)
assert response.status_code == 302
mock_delay.assert_not_called()
@patch("scrobbles.tasks.scrobble_agent_session_prompt.delay")
def test_agent_session_followup_form_shown_when_complete(mock_delay, client, user):
scrobble = _mk_scrobble(
user,
in_progress=False,
turns=[
{
"prompt_id": "abc-123",
"prompt": "hello",
"response": "hi back",
}
],
)
client.force_login(user)
response = client.get(
reverse("scrobbles:agent-session-partial", args=[scrobble.id])
)
assert response.status_code == 200
assert b"Ask a follow-up" in response.content
assert b"/follow-up/" in response.content
assert b"csrfmiddlewaretoken" in response.content
# --- log data ---
def test_agent_session_logdata_pending():
log = AgentSessionLogData(
provider="gemini",
model="gemini-2.5-flash",
turns=[{"prompt_id": "1", "prompt": "hi", "response": None}],
)
assert log.pending is True
log.turns[0]["response"] = "hi back"
assert log.pending is False
assert AgentSessionLogData(provider="gemini").pending is False
def test_agent_session_logdata_as_html():
log = AgentSessionLogData(
provider="gemini",
model="gemini-2.5-flash",
turns=[
{"prompt_id": "1", "prompt": "**bold?**", "response": "yes"},
{"prompt_id": "2", "prompt": "again", "response": None},
],
)
html = log.as_html()
assert '<div class="agent-turn">' in html
assert "Prompt 1" in html
assert "<strong>bold?</strong>" in html
assert "yes" in html
assert "Prompt 2" in html
assert "Thinking" in html