Files
spotify-voice-assistant/test_prompt.py
T
bhethermanandClaude Sonnet 5 2ed1ca1e99 Fix cold-start play_media failure and swap reversed room IDs
The Spotify integration only exposes the media_player.play_media
service once it has an active playback context, which doesn't exist
yet on a cold start; the fixed 2s delay after select_source wasn't
long enough in that case even though it always worked once something
had already played. Poll supported_features for the PLAY_MEDIA bit
(up to 8s) before calling play_media instead of guessing a fixed delay.

Also swap the Living Room/Bedroom source_list ID mapping in both the
prompt (for display) and the functions (for select_source) -- it was
backwards.

Adds a test-harness regression case exercising the opaque Spotify
Connect device IDs the source_list actually reports.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
2026-09-05 13:40:40 -04:00

499 lines
18 KiB
Python

#!/usr/bin/env python3
"""Test search_spotify/control_playback tool-call disambiguation directly against
the Ollama endpoint, bypassing Home Assistant entirely.
Renders current-prompt.txt (with mock HA template values) and current-functions.txt
(as OpenAI-style tool definitions) exactly as Extended OpenAI Conversation would send
them, then checks whether the model picks the right tool + arguments for each test
phrase. Only evaluates the FIRST tool call the model returns per test (i.e. whether
it correctly picks search_spotify vs control_playback and fills in query/type/action) —
it does not simulate the full multi-turn tool-result loop.
"""
import json
import os
import sys
from datetime import datetime
import requests
import yaml
from jinja2 import Template
HERE = os.path.dirname(os.path.abspath(__file__))
OLLAMA_BASE_URL = os.environ.get("OLLAMA_BASE_URL", "http://192.168.50.214:11434/v1")
OLLAMA_API_KEY = os.environ.get("OLLAMA_API_KEY", "38ecbb5e71adf1306291e810826d965afe51fa27ff42e60f")
MODEL = os.environ.get("OLLAMA_MODEL") # set this to the exact `ollama list` tag
MOCK_SOURCE_LIST = ["Living Room", "Web Player (Firefox)", "Bedroom"]
# Mirrors the real (currently broken) source_list this entity actually reports —
# Spotify Connect device IDs instead of friendly names — to exercise the
# opaque-ID fallback mapping added to current-prompt.txt's Room section.
MOCK_SOURCE_LIST_IDS = [
"a334c8ba989367b52303eaa9dbc424cdd83012b5",
"d0166d6f83796b0f9059b0adf0f6cb0d81fa6376",
]
ROOM_NAME_TO_ID = {
"Living Room": "a334c8ba989367b52303eaa9dbc424cdd83012b5",
"Bedroom": "d0166d6f83796b0f9059b0adf0f6cb0d81fa6376",
}
def render_prompt(source_list=None):
if source_list is None:
source_list = MOCK_SOURCE_LIST
with open(os.path.join(HERE, "current-prompt.txt")) as f:
raw = f.read()
def state_attr(entity_id, attr):
if entity_id == "media_player.spotify_bhetherman" and attr == "source_list":
return source_list
return None
template = Template(raw)
return template.render(
now=lambda: datetime.now(),
exposed_entities=[], # not relevant to these tests; kept empty on purpose
state_attr=state_attr,
)
def load_tools():
with open(os.path.join(HERE, "current-functions.txt")) as f:
functions = yaml.safe_load(f)
tools = []
for entry in functions:
spec = entry["spec"]
tools.append({
"type": "function",
"function": {
"name": spec["name"],
"description": spec["description"],
"parameters": spec["parameters"],
},
})
return tools
def call_model(system_prompt, user_message, tools):
if not MODEL:
sys.exit("Set OLLAMA_MODEL env var to the exact `ollama list` tag before running.")
resp = requests.post(
f"{OLLAMA_BASE_URL}/chat/completions",
headers={"Authorization": f"Bearer {OLLAMA_API_KEY}"},
json={
"model": MODEL,
"messages": [
{"role": "system", "content": system_prompt},
{"role": "user", "content": user_message},
],
"tools": tools,
"tool_choice": "auto",
},
timeout=120,
)
resp.raise_for_status()
return resp.json()
def first_tool_call(response):
message = response["choices"][0]["message"]
tool_calls = message.get("tool_calls") or []
if not tool_calls:
return None, None, message.get("content")
call = tool_calls[0]
name = call["function"]["name"]
try:
args = json.loads(call["function"]["arguments"])
except (json.JSONDecodeError, TypeError):
args = call["function"]["arguments"]
return name, args, None
REPEAT = int(os.environ.get("TEST_REPEAT", "15"))
def _check_once(user_message, expect_name, expect_fn, system_prompt, tools):
response = call_model(system_prompt, user_message, tools)
name, args, content = first_tool_call(response)
if name is None:
return False, f"NO TOOL CALL — model replied with text instead: {content!r}"
if name != expect_name:
return False, f"WRONG FUNCTION — expected {expect_name!r}, got {name!r} ({args})"
ok, reason = expect_fn(args)
if ok:
return True, f"{name}({args})"
return False, f"{name}({args}) — {reason}"
def _summarize(label, user_message, passed, attempts):
n = len(attempts)
print(f"\n=== {label} ===")
print(f" user: {user_message!r}")
for i, (ok, detail) in enumerate(attempts, 1):
mark = "✅" if ok else "❌"
print(f" [{i}/{n}] {mark} {detail}")
if passed == n:
print(f" RESULT: {passed}/{n} — reliable pass")
elif passed == 0:
print(f" RESULT: {passed}/{n} — reliable FAIL")
else:
print(f" RESULT: {passed}/{n} — FLAKY (inconsistent across runs)")
return passed, n
def check(label, user_message, expect_name, expect_fn, system_prompt, tools):
attempts = [_check_once(user_message, expect_name, expect_fn, system_prompt, tools) for _ in range(REPEAT)]
passed = sum(1 for ok, _ in attempts if ok)
return _summarize(label, user_message, passed, attempts)
def eq_ci(a, b):
return isinstance(a, str) and a.strip().lower() == b.strip().lower()
def contains_ci(a, b):
return isinstance(a, str) and b.strip().lower() in a.strip().lower()
TESTS = [
(
"artist-only, generic 'music by X' phrasing",
"play music by Jungle in the living room",
"search_spotify",
lambda args: (True, None) if eq_ci(args.get("query"), "Jungle") and eq_ci(args.get("type"), "artist")
else (False, f"expected query='Jungle' type='artist', got {args}"),
),
(
"song title containing a number (previously misread as volume)",
"play back on 74 in the living room",
"search_spotify",
lambda args: (True, None) if eq_ci(args.get("query"), "Back on 74") and eq_ci(args.get("type"), "track")
else (False, f"expected query='Back on 74' type='track', got {args}"),
),
(
"specific title + artist (regression check — previously worked)",
"play Yellow by Coldplay",
"search_spotify",
lambda args: (True, None) if eq_ci(args.get("query"), "Yellow") and eq_ci(args.get("type"), "track")
else (False, f"expected query='Yellow' type='track', got {args}"),
),
(
"explicit volume command (must NOT be treated as a search)",
"set the volume to 50",
"control_playback",
lambda args: (True, None) if args.get("action") == "volume_set" and float(args.get("volume_level", -1)) == 50
else (False, f"expected action='volume_set' volume_level=50, got {args}"),
),
(
"album title, different room",
"queue Dark Side of the Moon in the bedroom",
"search_spotify",
lambda args: (True, None) if eq_ci(args.get("query"), "Dark Side of the Moon") and args.get("type") in ("album", "track")
else (False, f"expected query='Dark Side of the Moon' type in (album,track), got {args}"),
),
(
"genre/mood, no artist",
"play some jazz",
"search_spotify",
lambda args: (True, None) if contains_ci(args.get("query"), "jazz") and eq_ci(args.get("type"), "playlist")
else (False, f"expected query~'jazz' type='playlist', got {args}"),
),
(
"named playlist phrase",
"play my daily mix",
"search_spotify",
lambda args: (True, None) if contains_ci(args.get("query"), "daily mix") and eq_ci(args.get("type"), "playlist")
else (False, f"expected query~'daily mix' type='playlist', got {args}"),
),
(
"number-titled track + explicit artist (by-stripping regression variant)",
"play 1979 by the Smashing Pumpkins",
"search_spotify",
lambda args: (True, None) if eq_ci(args.get("query"), "1979") and eq_ci(args.get("type"), "track")
else (False, f"expected query='1979' type='track', got {args}"),
),
(
"number-titled track, no artist given",
"play 99 Luftballons",
"search_spotify",
lambda args: (True, None) if eq_ci(args.get("query"), "99 Luftballons") and eq_ci(args.get("type"), "track")
else (False, f"expected query='99 Luftballons' type='track', got {args}"),
),
(
"artist name that itself contains a number",
"queue some Blink 182",
"search_spotify",
lambda args: (True, None) if eq_ci(args.get("query"), "Blink 182") and eq_ci(args.get("type"), "artist")
else (False, f"expected query='Blink 182' type='artist', got {args}"),
),
(
"colloquial phrasing ('put on')",
"put on some Fred again",
"search_spotify",
lambda args: (True, None) if eq_ci(args.get("query"), "Fred again") and eq_ci(args.get("type"), "artist")
else (False, f"expected query='Fred again' type='artist', got {args}"),
),
(
"explicit 'album' keyword + by-stripping",
"throw on the album Parachutes by Coldplay",
"search_spotify",
lambda args: (True, None) if eq_ci(args.get("query"), "Parachutes") and eq_ci(args.get("type"), "album")
else (False, f"expected query='Parachutes' type='album', got {args}"),
),
(
"explicit 'track' keyword",
"can you play the track Yellow",
"search_spotify",
lambda args: (True, None) if eq_ci(args.get("query"), "Yellow") and eq_ci(args.get("type"), "track")
else (False, f"expected query='Yellow' type='track', got {args}"),
),
(
"pause",
"pause the music",
"control_playback",
lambda args: (True, None) if args.get("action") == "pause"
else (False, f"expected action='pause', got {args}"),
),
(
"next track",
"skip to the next song",
"control_playback",
lambda args: (True, None) if args.get("action") == "next_track"
else (False, f"expected action='next_track', got {args}"),
),
(
"stop",
"stop the music in the living room",
"control_playback",
lambda args: (True, None) if args.get("action") == "stop"
else (False, f"expected action='stop', got {args}"),
),
(
"shuffle on",
"turn on shuffle",
"control_playback",
lambda args: (True, None) if args.get("action") == "shuffle_on"
else (False, f"expected action='shuffle_on', got {args}"),
),
(
"volume with 'percent' phrasing",
"set volume to 80 percent",
"control_playback",
lambda args: (True, None) if args.get("action") == "volume_set" and float(args.get("volume_level", -1)) == 80
else (False, f"expected action='volume_set' volume_level=80, got {args}"),
),
]
def call_model_step2(system_prompt, user_message, tools, step1_name, step1_args):
"""Feed a fake search_spotify tool result back and get the model's next tool call
(should be play_music or queue_music)."""
fake_result = {
"uri": f"spotify:{step1_args.get('type', 'track')}:FAKEURI0000000000000000",
"name": step1_args.get("query", ""),
"type": step1_args.get("type", "track"),
}
tool_call_id = "call_1"
messages = [
{"role": "system", "content": system_prompt},
{"role": "user", "content": user_message},
{
"role": "assistant",
"content": None,
"tool_calls": [{
"id": tool_call_id,
"type": "function",
"function": {"name": step1_name, "arguments": json.dumps(step1_args)},
}],
},
{
"role": "tool",
"tool_call_id": tool_call_id,
"name": step1_name,
"content": json.dumps(fake_result),
},
]
resp = requests.post(
f"{OLLAMA_BASE_URL}/chat/completions",
headers={"Authorization": f"Bearer {OLLAMA_API_KEY}"},
json={"model": MODEL, "messages": messages, "tools": tools, "tool_choice": "auto"},
timeout=120,
)
resp.raise_for_status()
return resp.json()
def _check_flow_once(user_message, expect_step2_name, expect_room, expect_type, system_prompt, tools):
response1 = call_model(system_prompt, user_message, tools)
name1, args1, content1 = first_tool_call(response1)
if name1 != "search_spotify":
return False, f"STEP 1 FAILED — expected search_spotify, got {name1!r} ({args1 or content1})"
response2 = call_model_step2(system_prompt, user_message, tools, name1, args1)
name2, args2, content2 = first_tool_call(response2)
if name2 is None:
return False, f"step1={name1}({args1}) -> STEP 2 no tool call, model replied with text: {content2!r}"
detail = f"step1={name1}({args1}) -> step2={name2}({args2})"
if name2 != expect_step2_name:
return False, f"{detail} — WRONG FUNCTION, expected {expect_step2_name!r}"
if not eq_ci(args2.get("room"), expect_room):
return False, f"{detail} — expected room={expect_room!r}"
if not eq_ci(args2.get("type"), expect_type):
return False, f"{detail} — expected type={expect_type!r} (must be passed through from step1, not hardcoded)"
return True, detail
def check_flow(label, user_message, expect_step2_name, expect_room, expect_type, system_prompt, tools):
"""Runs search_spotify -> (fake result) -> play_music/queue_music, repeated REPEAT
times, and checks the second call's function name, room, and type parameters."""
attempts = [
_check_flow_once(user_message, expect_step2_name, expect_room, expect_type, system_prompt, tools)
for _ in range(REPEAT)
]
passed = sum(1 for ok, _ in attempts if ok)
return _summarize(label, user_message, passed, attempts)
TESTS_FLOW = [
# (label, user_message, expect_step2_name, expect_room, expect_type)
(
"play + explicit room",
"play Fred again in the bedroom",
"play_music",
"Bedroom",
"artist",
),
(
"queue + artist -> still calls queue_music (fallback handled inside the script now)",
"queue Fred again in the bedroom",
"queue_music",
"Bedroom",
"artist",
),
(
"play + no room specified -> default room",
"play Yellow by Coldplay",
"play_music",
"Living Room",
"track",
),
(
"playlist/genre + explicit room",
"play jazz in the bedroom",
"play_music",
"Bedroom",
"playlist",
),
(
"queue + album -> still calls queue_music (fallback handled inside the script now)",
"queue Dark Side of the Moon",
"queue_music",
"Living Room",
"album",
),
(
"queue + playlist -> still calls queue_music (fallback handled inside the script now)",
"queue my workout playlist in the bedroom",
"queue_music",
"Bedroom",
"playlist",
),
(
"'add to queue' phrasing + artist -> still calls queue_music",
"add some Blink 182 to the queue in the living room",
"queue_music",
"Living Room",
"artist",
),
(
"queue + track, no room specified -> default room, stays queue_music",
"queue Yellow by Coldplay",
"queue_music",
"Living Room",
"track",
),
(
"artist ('music by X' pattern) + explicit room, carried through to step 2",
"play some music by Jungle in the bedroom",
"play_music",
"Bedroom",
"artist",
),
(
"number-titled track + explicit room, carried through to step 2",
"play back on 74 in the bedroom",
"play_music",
"Bedroom",
"track",
),
(
"queue + number-titled track, explicit room",
"queue 99 Luftballons in the living room",
"queue_music",
"Living Room",
"track",
),
]
def main():
system_prompt = render_prompt()
tools = load_tools()
print(f"Model: {MODEL}")
print(f"Endpoint: {OLLAMA_BASE_URL}")
print(f"Repeat per test: {REPEAT}")
print(f"Loaded {len(tools)} tools: {[t['function']['name'] for t in tools]}")
results = []
for label, user_message, expect_name, expect_fn in TESTS:
results.append((label, check(label, user_message, expect_name, expect_fn, system_prompt, tools)))
print(f"\n{'#' * 40}\n# Two-step flow tests (search -> play/queue)\n{'#' * 40}")
for label, user_message, expect_step2_name, expect_room, expect_type in TESTS_FLOW:
results.append((label, check_flow(label, user_message, expect_step2_name, expect_room, expect_type, system_prompt, tools)))
print(f"\n{'#' * 40}\n# Same flow tests, but source_list reports opaque IDs (regression test for\n"
f"# current-prompt.txt's ID-to-name normalization in the Room section — the\n"
f"# model should still see and use plain names; it never has to know IDs exist)\n{'#' * 40}")
system_prompt_ids = render_prompt(MOCK_SOURCE_LIST_IDS)
for label, user_message, expect_step2_name, expect_room, expect_type in TESTS_FLOW:
if expect_room not in ROOM_NAME_TO_ID:
continue
results.append((
f"{label} [opaque-ID rooms]",
check_flow(label, user_message, expect_step2_name, expect_room, expect_type, system_prompt_ids, tools),
))
total_passed = sum(passed for _, (passed, n) in results)
total_attempts = sum(n for _, (passed, n) in results)
reliable = [label for label, (passed, n) in results if passed == n]
flaky = [(label, passed, n) for label, (passed, n) in results if 0 < passed < n]
failing = [label for label, (passed, n) in results if passed == 0]
print(f"\n{'=' * 40}")
print(f"{total_passed}/{total_attempts} individual attempts passed across {len(results)} tests (x{REPEAT} reps)")
print(f" {len(reliable)}/{len(results)} tests fully reliable (all {REPEAT} reps passed)")
if flaky:
print(f" {len(flaky)} FLAKY tests (inconsistent across reps):")
for label, passed, n in flaky:
print(f" - {label} ({passed}/{n})")
if failing:
print(f" {len(failing)} tests failing every rep:")
for label in failing:
print(f" - {label}")
sys.exit(0 if not flaky and not failing else 1)
if __name__ == "__main__":
main()