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>
499 lines
18 KiB
Python
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()
|