#!/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()