import asyncio import re from datetime import datetime, timedelta, timezone import pytest from jam_session.models import Track from jam_session.session import ( NicknameTakenError, SessionManager, generate_nickname, ) def make_track(title="Test Track", source_id="abc123", **kwargs): defaults = { "title": title, "duration": 180, "stream_url": "https://example.com/audio.opus", "thumbnail": None, "source_id": source_id, "webpage_url": f"https://youtube.com/watch?v={source_id}", } defaults.update(kwargs) return Track(**defaults) class TestGenerateNickname: def test_format_matches_pattern(self): name = generate_nickname() match = re.match(r"^[A-Z][a-z]+-[A-Z][a-z]+-\d{1,2}$", name) assert match is not None, f"'{name}' does not match expected pattern" def test_returns_different_names(self): names = {generate_nickname() for _ in range(50)} assert len(names) > 1 class TestComputeSkipThreshold: def test_one_nickname_returns_one(self): assert SessionManager._compute_skip_threshold(1, 50) == 1 def test_two_nicknames_returns_two(self): assert SessionManager._compute_skip_threshold(2, 50) == 2 def test_three_nicknames_returns_two(self): assert SessionManager._compute_skip_threshold(3, 50) == 2 def test_ten_nicknames_fifty_pct(self): assert SessionManager._compute_skip_threshold(10, 50) == 6 def test_four_nicknames_fifty_pct(self): assert SessionManager._compute_skip_threshold(4, 50) == 3 def test_ten_nicknames_one_hundred_pct(self): assert SessionManager._compute_skip_threshold(10, 100) == 11 def test_ten_nicknames_ten_pct(self): assert SessionManager._compute_skip_threshold(10, 10) == 2 class TestAddTrack: @pytest.mark.asyncio async def test_appends_track(self): sm = SessionManager() track = make_track() result = await sm.add_track(track) assert result is track assert sm._queue == [track] @pytest.mark.asyncio async def test_multiple_appends_in_order(self): sm = SessionManager() t1 = make_track(source_id="1") t2 = make_track(source_id="2") await sm.add_track(t1) await sm.add_track(t2) assert len(sm._queue) == 2 assert sm._queue[0].source_id == "1" assert sm._queue[1].source_id == "2" class TestRemoveTrack: @pytest.mark.asyncio async def test_removes_existing(self): sm = SessionManager() track = make_track(source_id="remove-me") await sm.add_track(track) removed = await sm.remove_track(track.id) assert removed is track assert sm._queue == [] @pytest.mark.asyncio async def test_returns_none_for_missing(self): sm = SessionManager() removed = await sm.remove_track("nonexistent") assert removed is None @pytest.mark.asyncio async def test_preserves_order_after_removal(self): sm = SessionManager() t1 = make_track(source_id="1") t2 = make_track(source_id="2") t3 = make_track(source_id="3") await sm.add_track(t1) await sm.add_track(t2) await sm.add_track(t3) await sm.remove_track(t2.id) assert [t.source_id for t in sm._queue] == ["1", "3"] @pytest.mark.asyncio async def test_silently_removes_votes_for_track(self): sm = SessionManager() track = make_track(source_id="voted") await sm.add_track(track) sm._votes["conn1"] = {track.id} await sm.remove_track(track.id) assert track.id not in sm._votes["conn1"] class TestReorder: @pytest.mark.asyncio async def test_reorders_queue(self): sm = SessionManager() t1 = make_track(source_id="1") t2 = make_track(source_id="2") t3 = make_track(source_id="3") await sm.add_track(t1) await sm.add_track(t2) await sm.add_track(t3) result = await sm.reorder([t3.id, t1.id, t2.id]) assert [t.source_id for t in result] == ["3", "1", "2"] @pytest.mark.asyncio async def test_mismatch_raises_value_error(self): sm = SessionManager() t1 = make_track(source_id="1") await sm.add_track(t1) with pytest.raises(ValueError, match="does not match queue content"): await sm.reorder(["nonexistent"]) @pytest.mark.asyncio async def test_partial_order_raises_value_error(self): sm = SessionManager() t1 = make_track(source_id="1") t2 = make_track(source_id="2") await sm.add_track(t1) await sm.add_track(t2) with pytest.raises(ValueError, match="does not match queue content"): await sm.reorder([t1.id]) class TestNickname: @pytest.mark.asyncio async def test_set_assigns_nickname(self): sm = SessionManager() result = await sm.set_nickname("conn1", "Cool-Cat-42") assert result == "Cool-Cat-42" assert sm.get_nickname("conn1") == "Cool-Cat-42" @pytest.mark.asyncio async def test_set_raises_when_taken(self): sm = SessionManager() await sm.set_nickname("conn1", "Cool-Cat-42") with pytest.raises(NicknameTakenError, match="already taken"): await sm.set_nickname("conn2", "Cool-Cat-42") @pytest.mark.asyncio async def test_case_insensitive_collision(self): sm = SessionManager() await sm.set_nickname("conn1", "Cool-Cat-42") with pytest.raises(NicknameTakenError): await sm.set_nickname("conn2", "cool-cat-42") @pytest.mark.asyncio async def test_allow_same_conn_reassign(self): sm = SessionManager() await sm.set_nickname("conn1", "Old-Name-1") result = await sm.set_nickname("conn1", "New-Name-2") assert result == "New-Name-2" assert sm.get_nickname("conn1") == "New-Name-2" @pytest.mark.asyncio async def test_reassign_frees_old_name(self): sm = SessionManager() await sm.set_nickname("conn1", "Old-Name-1") await sm.set_nickname("conn1", "New-Name-2") await sm.set_nickname("conn2", "Old-Name-1") assert sm.get_nickname("conn2") == "Old-Name-1" @pytest.mark.asyncio async def test_get_none_for_unknown(self): sm = SessionManager() assert sm.get_nickname("nobody") is None class TestHoldNickname: @pytest.mark.asyncio async def test_hold_removes_from_active(self): sm = SessionManager() await sm.set_nickname("conn1", "Ghost-1") await sm.hold_nickname("conn1") assert sm.get_nickname("conn1") is None assert "Ghost-1" in sm._held_nicknames @pytest.mark.asyncio async def test_hold_frees_for_other_after_release(self): sm = SessionManager() await sm.set_nickname("conn1", "Shared-Name-7") await sm.hold_nickname("conn1") await sm.release_expired_holds() await sm.set_nickname("conn2", "Shared-Name-7") assert sm.get_nickname("conn2") == "Shared-Name-7" class TestToggleVote: @pytest.mark.asyncio async def test_add_vote_increases_count(self): sm = SessionManager() track = make_track() await sm.add_track(track) count, upvoted = await sm.toggle_vote("conn1", track.id) assert count == 1 assert upvoted is True @pytest.mark.asyncio async def test_remove_vote_decreases_count(self): sm = SessionManager() track = make_track() await sm.add_track(track) await sm.toggle_vote("conn1", track.id) count, upvoted = await sm.toggle_vote("conn1", track.id) assert count == 0 assert upvoted is False @pytest.mark.asyncio async def test_multiple_users_vote_same_track(self): sm = SessionManager() track = make_track() await sm.add_track(track) await sm.toggle_vote("conn1", track.id) await sm.toggle_vote("conn2", track.id) count, _ = await sm.toggle_vote("conn3", track.id) assert count == 3 @pytest.mark.asyncio async def test_raises_for_missing_track(self): sm = SessionManager() with pytest.raises(ValueError, match="not found"): await sm.toggle_vote("conn1", "nonexistent") class TestSortByVotes: @pytest.mark.asyncio async def test_higher_votes_sort_first(self): sm = SessionManager() t1 = make_track(source_id="low") t2 = make_track(source_id="high") await sm.add_track(t1) await sm.add_track(t2) sm._votes["conn1"] = {t2.id, t1.id} sm._votes["conn2"] = {t2.id} sm._sort_by_votes() assert sm._queue[0].source_id == "high" assert sm._queue[1].source_id == "low" @pytest.mark.asyncio async def test_tiebreak_by_added_at(self): sm = SessionManager() t1 = make_track(source_id="older") t2 = make_track(source_id="newer") t1.added_at = datetime(2020, 1, 1, tzinfo=timezone.utc) t2.added_at = datetime(2020, 1, 2, tzinfo=timezone.utc) await sm.add_track(t1) await sm.add_track(t2) sm._sort_by_votes() assert sm._queue[0].source_id == "older" assert sm._queue[1].source_id == "newer" class TestGetVotes: def test_get_votes_for_returns_list(self): sm = SessionManager() sm._votes["conn1"] = {"id1", "id2"} assert sorted(sm.get_votes_for("conn1")) == ["id1", "id2"] def test_get_votes_for_empty_connection(self): sm = SessionManager() assert sm.get_votes_for("nobody") == [] @pytest.mark.asyncio async def test_get_vote_counts_maps_all(self): sm = SessionManager() t1 = make_track(source_id="a") t2 = make_track(source_id="b") await sm.add_track(t1) await sm.add_track(t2) sm._votes["conn1"] = {t1.id} counts = sm.get_vote_counts() assert counts[t1.id] == 1 assert counts[t2.id] == 0 class TestSkipVote: @pytest.mark.asyncio async def test_cast_vote_counts_up(self): sm = SessionManager() sm._skip_denominator = 3 state = await sm.toggle_skip_vote("conn1", 50) assert state["vote_count"] == 1 assert state["triggered"] is False @pytest.mark.asyncio async def test_trigger_when_threshold_reached(self): sm = SessionManager() sm._skip_denominator = 2 await sm.toggle_skip_vote("conn1", 50) state = await sm.toggle_skip_vote("conn2", 50) assert state["triggered"] is True @pytest.mark.asyncio async def test_cooldown_blocks_new_votes(self): sm = SessionManager() import time loop = asyncio.get_running_loop() sm._skip_cooldown_until = loop.time() + 100 sm._skip_denominator = 5 state = await sm.toggle_skip_vote("conn1", 50) assert state["vote_count"] == 0 assert state["triggered"] is False assert state["cooldown_remaining"] > 0 @pytest.mark.asyncio async def test_remove_vote_during_cooldown_also_blocked(self): sm = SessionManager() loop = asyncio.get_running_loop() sm._skip_cooldown_until = loop.time() + 100 sm._skip_votes.add("conn1") sm._skip_denominator = 3 state = await sm.toggle_skip_vote("conn1", 50) assert state["vote_count"] == 1 assert state["triggered"] is False assert state["cooldown_remaining"] > 0 @pytest.mark.asyncio async def test_get_skip_state_my_vote(self): sm = SessionManager() sm._skip_denominator = 5 sm._skip_votes.add("conn1") state = sm.get_skip_state("conn1", 50) assert state["my_skip_vote"] is True assert state["skip_vote_count"] == 1 assert state["skip_threshold"] == 3 @pytest.mark.asyncio async def test_get_skip_state_no_vote(self): sm = SessionManager() sm._skip_denominator = 5 state = sm.get_skip_state("conn1", 50) assert state["my_skip_vote"] is False def test_reset_skip_votes_clears_all(self): sm = SessionManager() sm._skip_votes.add("conn1") sm._skip_votes.add("conn2") sm._skip_triggered = True sm.reset_skip_votes() assert sm._skip_votes == set() assert sm._skip_triggered is False assert sm._skip_denominator == 0 class TestSkipCooldown: @pytest.mark.asyncio async def test_start_cooldown_sets_until(self): sm = SessionManager() loop = asyncio.get_running_loop() sm.start_skip_cooldown(5) assert sm._skip_cooldown_until > loop.time() @pytest.mark.asyncio async def test_cooldown_callback_fires(self): sm = SessionManager() sm.start_skip_cooldown(0.05) await asyncio.sleep(0.15) assert sm._cooldown_handle is None class TestHostToken: def test_set_and_get_token(self): sm = SessionManager() sm.set_host_token("secret") assert sm.get_host_token() == "secret" def test_validate_correct_token(self): sm = SessionManager() sm.set_host_token("secret") assert sm.validate_host_token("secret") is True def test_validate_wrong_token(self): sm = SessionManager() sm.set_host_token("secret") assert sm.validate_host_token("wrong") is False def test_validate_no_token_set(self): sm = SessionManager() assert sm.validate_host_token("anything") is False def test_set_and_get_host_conn_id(self): sm = SessionManager() sm.set_host_conn_id("conn1") assert sm.get_host_conn_id() == "conn1" def test_is_host_true(self): sm = SessionManager() sm.set_host_conn_id("conn1") assert sm.is_host("conn1") is True def test_is_host_false(self): sm = SessionManager() sm.set_host_conn_id("conn1") assert sm.is_host("conn2") is False def test_is_host_no_host_set(self): sm = SessionManager() assert sm.is_host("conn1") is False class TestAdvance: @pytest.mark.asyncio async def test_pops_next_track(self): sm = SessionManager() t1 = make_track(source_id="next") await sm.add_track(t1) result = await sm.advance() assert result is t1 assert sm._current_track is t1 @pytest.mark.asyncio async def test_empty_queue_returns_none(self): sm = SessionManager() result = await sm.advance() assert result is None assert sm._current_track is None @pytest.mark.asyncio async def test_resets_playback_time(self): sm = SessionManager() sm.playback_time = 99.0 t1 = make_track() await sm.add_track(t1) await sm.advance() assert sm.playback_time == 0.0 class TestManualOrder: def test_set_manual_order(self): sm = SessionManager() sm.set_manual_order(True) assert sm.is_manual_order() is True def test_default_not_manual(self): sm = SessionManager() assert sm.is_manual_order() is False class TestGetState: @pytest.mark.asyncio async def test_returns_queue_state_snapshot(self): sm = SessionManager() t1 = make_track(source_id="1") await sm.add_track(t1) sm.paused = True sm.playback_time = 42.0 state = sm.get_state() assert len(state.queue) == 1 assert state.paused is True assert state.playback_time == 42.0 @pytest.mark.asyncio async def test_snapshot_includes_nicknames(self): sm = SessionManager() await sm.set_nickname("conn1", "DJ-Fox-5") state = sm.get_state() assert state.nicknames == {"conn1": "DJ-Fox-5"}