/
/
/
1"""Tests for SonicSimilarityPlugin CLAP handler and rebuild methods."""
2
3from __future__ import annotations
4
5from typing import TYPE_CHECKING
6from unittest.mock import AsyncMock, MagicMock
7
8import numpy as np
9import pytest
10
11from music_assistant.providers.sonic_similarity.similarity import ScoredCandidate
12from tests.providers.sonic_similarity.conftest import make_analysis_row
13
14if TYPE_CHECKING:
15 from collections.abc import Callable
16 from typing import Any
17
18
19class TestHandleSimilarClap:
20 """Tests for SonicSimilarityPlugin._handle_similar_clap."""
21
22 @pytest.mark.asyncio
23 async def test_clap_index_disabled_when_no_index(self, make_plugin: Callable[..., Any]) -> None:
24 """Returns clap_index_disabled when no CLAP index attached."""
25 plugin = make_plugin(signatures={("spotify", "seed"): [0.0] * 18})
26 result = await plugin._handle_similar_clap("any_id")
27 assert result["analyzed"] is False
28 assert result["reason"] == "clap_index_disabled"
29 assert result["seed_track_id"] == "any_id"
30 assert result["items"] == []
31
32 @pytest.mark.asyncio
33 async def test_seed_not_in_index(self, make_plugin: Callable[..., Any]) -> None:
34 """Returns seed_not_in_index when get_embedding_by_item_id returns None."""
35 plugin = make_plugin(clap_enabled=True, signatures={("spotify", "seed"): [0.0] * 18})
36 result = await plugin._handle_similar_clap("seed")
37 assert result["analyzed"] is False
38 assert result["reason"] == "seed_not_in_index"
39 assert result["items"] == []
40
41 @pytest.mark.asyncio
42 async def test_happy_path_excludes_seed(self, make_plugin: Callable[..., Any]) -> None:
43 """Returns ranked items, drops the seed and respects limit ordering."""
44 plugin = make_plugin(clap_enabled=True, signatures={("spotify", "seed"): [0.0] * 18})
45 plugin._clap_index.get_embedding_by_item_id.return_value = (
46 "spotify",
47 np.zeros(1024, dtype=np.float32),
48 )
49 plugin._clap_index.search = AsyncMock(
50 return_value=[
51 ScoredCandidate("seed", "spotify", 0.0),
52 ScoredCandidate("other1", "tidal", 0.2),
53 ScoredCandidate("other2", "apple", 0.3),
54 ]
55 )
56 result = await plugin._handle_similar_clap("seed", limit=2)
57 assert result["analyzed"] is True
58 assert result["seed_track_id"] == "seed"
59 assert len(result["items"]) == 2
60 ids = [item["item_id"] for item in result["items"]]
61 assert "seed" not in ids
62 assert ids == ["other1", "other2"]
63 assert result["items"][0]["provider"] == "tidal"
64 assert result["items"][0]["distance"] == 0.2
65
66 @pytest.mark.asyncio
67 async def test_respects_limit(self, make_plugin: Callable[..., Any]) -> None:
68 """Caps the items list at limit even when search returns many candidates."""
69 plugin = make_plugin(clap_enabled=True, signatures={("spotify", "seed"): [0.0] * 18})
70 plugin._clap_index.get_embedding_by_item_id.return_value = (
71 "spotify",
72 np.zeros(1024, dtype=np.float32),
73 )
74 plugin._clap_index.search = AsyncMock(
75 return_value=[
76 ScoredCandidate(f"track_{i}", "spotify", float(i) * 0.01) for i in range(10)
77 ]
78 )
79 result = await plugin._handle_similar_clap("seed", limit=3)
80 assert result["analyzed"] is True
81 assert len(result["items"]) == 3
82
83
84class TestRebuildClapIndexFromDatabase:
85 """Tests for SonicSimilarityPlugin._rebuild_clap_index_from_database."""
86
87 @pytest.mark.asyncio
88 async def test_noop_when_clap_index_disabled(
89 self, make_plugin: Callable[..., Any], mock_mass: MagicMock
90 ) -> None:
91 """Early-returns without touching audio_analysis when index disabled."""
92 plugin = make_plugin()
93 await plugin._rebuild_clap_index_from_database()
94 assert mock_mass.streams.audio_analysis.iter_audio_analysis_rows.call_count == 0
95
96 @pytest.mark.asyncio
97 async def test_adds_new_embeddings_and_saves(
98 self, make_plugin: Callable[..., Any], mock_mass: MagicMock
99 ) -> None:
100 """Adds each unique new embedding to the index and persists once."""
101 mock_mass._iter_audio_analysis_rows_data = [
102 make_analysis_row(item_id="a", provider="spotify", clap_embedding=[0.1] * 1024),
103 make_analysis_row(item_id="b", provider="spotify", clap_embedding=[0.2] * 1024),
104 ]
105 plugin = make_plugin(clap_enabled=True)
106 await plugin._rebuild_clap_index_from_database()
107 assert plugin._clap_index.add.await_count == 2
108 call_args = [call.args for call in plugin._clap_index.add.await_args_list]
109 # Each call: (provider, item_id, ndarray)
110 assert call_args[0][0] == "spotify"
111 assert call_args[0][1] == "a"
112 assert isinstance(call_args[0][2], np.ndarray)
113 assert call_args[0][2].shape == (1024,)
114 assert call_args[1][0] == "spotify"
115 assert call_args[1][1] == "b"
116 assert isinstance(call_args[1][2], np.ndarray)
117 plugin._clap_index.save.assert_awaited_once()
118
119 @pytest.mark.asyncio
120 async def test_skips_already_indexed_rows(
121 self, make_plugin: Callable[..., Any], mock_mass: MagicMock
122 ) -> None:
123 """Rows whose (provider, item_id) is contained() are skipped entirely."""
124 mock_mass._iter_audio_analysis_rows_data = [
125 make_analysis_row(item_id="a", clap_embedding=[0.1] * 1024),
126 make_analysis_row(item_id="b", clap_embedding=[0.2] * 1024),
127 ]
128 plugin = make_plugin(clap_enabled=True)
129 plugin._clap_index.contains = MagicMock(return_value=True)
130 await plugin._rebuild_clap_index_from_database()
131 plugin._clap_index.add.assert_not_awaited()
132 plugin._clap_index.save.assert_not_awaited()
133
134 @pytest.mark.asyncio
135 async def test_skips_malformed_json(
136 self, make_plugin: Callable[..., Any], mock_mass: MagicMock
137 ) -> None:
138 """A row with non-JSON analysis_data is skipped without crashing."""
139 malformed_row = {
140 "item_id": "x",
141 "provider": "spotify",
142 "aa_provider_domain": "sonic_analysis",
143 "analysis_data": "not-valid-json",
144 }
145 mock_mass._iter_audio_analysis_rows_data = [malformed_row]
146 plugin = make_plugin(clap_enabled=True)
147 await plugin._rebuild_clap_index_from_database()
148 plugin._clap_index.add.assert_not_awaited()
149 plugin._clap_index.save.assert_not_awaited()
150
151 @pytest.mark.asyncio
152 async def test_skips_rows_missing_clap_embedding(
153 self, make_plugin: Callable[..., Any], mock_mass: MagicMock
154 ) -> None:
155 """Rows whose extra_data lacks clap_embedding are skipped."""
156 mock_mass._iter_audio_analysis_rows_data = [
157 make_analysis_row(item_id="x", clap_embedding=None),
158 ]
159 plugin = make_plugin(clap_enabled=True)
160 await plugin._rebuild_clap_index_from_database()
161 plugin._clap_index.add.assert_not_awaited()
162 plugin._clap_index.save.assert_not_awaited()
163
164 @pytest.mark.asyncio
165 async def test_skips_wrong_shape_embedding(
166 self, make_plugin: Callable[..., Any], mock_mass: MagicMock
167 ) -> None:
168 """Rows whose embedding parses but isn't 1024-dim are rejected."""
169 mock_mass._iter_audio_analysis_rows_data = [
170 make_analysis_row(item_id="x", clap_embedding=[0.1] * 100),
171 ]
172 plugin = make_plugin(clap_enabled=True)
173 await plugin._rebuild_clap_index_from_database()
174 plugin._clap_index.add.assert_not_awaited()
175 plugin._clap_index.save.assert_not_awaited()
176
177 @pytest.mark.asyncio
178 async def test_dedupes_within_single_rebuild(
179 self, make_plugin: Callable[..., Any], mock_mass: MagicMock
180 ) -> None:
181 """Duplicate (provider, item_id) rows in a single rebuild add at most once."""
182 mock_mass._iter_audio_analysis_rows_data = [
183 make_analysis_row(item_id="a", provider="spotify", clap_embedding=[0.1] * 1024),
184 make_analysis_row(item_id="a", provider="spotify", clap_embedding=[0.1] * 1024),
185 ]
186 plugin = make_plugin(clap_enabled=True)
187 await plugin._rebuild_clap_index_from_database()
188 assert plugin._clap_index.add.await_count == 1
189