from __future__ import annotations import asyncio import json from unittest.mock import AsyncMock, MagicMock, patch import pytest from yuxi.channels.adapters.nostr.relay_manager import RelayManager class TestRelayManagerCB: def test_per_relay_cb_created(self): urls = ["wss://relay1.example.com", "wss://relay2.example.com"] mgr = RelayManager(urls) assert len(mgr._circuit_breakers) == 2 for url in urls: assert mgr.get_circuit_breaker(url) is not None def test_default_relay_scores(self): urls = ["wss://a.relay", "wss://b.relay"] mgr = RelayManager(urls) assert mgr.get_relay_score("wss://a.relay") == 0.5 def test_set_relay_scores(self): urls = ["wss://a.relay", "wss://b.relay"] mgr = RelayManager(urls) mgr.set_relay_scores({"wss://a.relay": 0.9, "wss://b.relay": 0.2}) assert mgr.get_relay_score("wss://a.relay") == 0.9 assert mgr.get_relay_score("wss://b.relay") == 0.2 def test_pubkey_filter_injection(self): urls = ["wss://a.relay"] mgr = RelayManager(urls) mgr.set_pubkey_filter(["abc123"]) filters = mgr._build_subscription_filters([{"kinds": [1, 4], "since": 1000}]) assert "#p" in filters[0] assert filters[0]["#p"] == ["abc123"] def test_pubkey_filter_no_change_when_empty(self): urls = ["wss://a.relay"] mgr = RelayManager(urls) filters = mgr._build_subscription_filters([{"kinds": [1, 4], "since": 1000}]) assert "#p" not in filters[0] class TestRelayManagerConnect: @pytest.mark.asyncio async def test_connect_all_no_relays(self): mgr = RelayManager(relay_urls=[]) await mgr.connect_all() @pytest.mark.asyncio async def test_connect_all_triggers_handler(self): mgr = RelayManager(relay_urls=["wss://a.relay"], timeout=0.1) connect_calls = [] mgr.on_connect(lambda url: connect_calls.append(url) or asyncio.sleep(0)) with patch("yuxi.channels.adapters.nostr.relay_manager.websockets.connect") as mc: mock_ws = AsyncMock() mock_ws.open = True mc.return_value = mock_ws await mgr.connect_all() assert len(mgr._connections) == 1 assert "wss://a.relay" in connect_calls class TestRelayManagerBroadcast: @pytest.mark.asyncio async def test_broadcast_sends_to_all(self): mgr = RelayManager(relay_urls=["wss://a.relay", "wss://b.relay"]) ws_a = AsyncMock() ws_a.open = True ws_b = AsyncMock() ws_b.open = True mgr._connections = {"wss://a.relay": ws_a, "wss://b.relay": ws_b} event = {"id": "test_event", "kind": 1, "content": "hello"} count = await mgr.broadcast(event) assert count == 2 @pytest.mark.asyncio async def test_broadcast_skips_open_cb(self): from yuxi.channels.infra.circuit_breaker import CircuitBreaker mgr = RelayManager(relay_urls=["wss://a.relay", "wss://b.relay"]) ws_a = AsyncMock() ws_a.open = True ws_b = AsyncMock() ws_b.open = True mgr._connections = {"wss://a.relay": ws_a, "wss://b.relay": ws_b} mgr._circuit_breakers["wss://b.relay"] = CircuitBreaker( failure_threshold=5, recovery_timeout=30.0 ) mgr._circuit_breakers["wss://b.relay"].state = "open" event = {"id": "test_event", "kind": 1, "content": "hello"} count = await mgr.broadcast(event) assert count == 1 @pytest.mark.asyncio async def test_broadcast_no_connections(self): mgr = RelayManager(relay_urls=["wss://a.relay"]) mgr._connections = {} event = {"id": "test_event", "kind": 1, "content": "hello"} count = await mgr.broadcast(event) assert count == 0 @pytest.mark.asyncio async def test_broadcast_sorted_by_score(self): mgr = RelayManager(relay_urls=["wss://low.relay", "wss://high.relay"]) mgr.set_relay_scores({"wss://low.relay": 0.2, "wss://high.relay": 0.9}) ws_low = AsyncMock() ws_low.open = True ws_high = AsyncMock() ws_high.open = True mgr._connections = {"wss://low.relay": ws_low, "wss://high.relay": ws_high} send_order = [] async def track_send(msg): send_order.append(msg) ws_low.send = track_send ws_high.send = track_send event = {"id": "test_event", "kind": 1, "content": "hello"} await mgr.broadcast(event) class TestRelayManagerQuery: @pytest.mark.asyncio async def test_query_aggregates_multiple_relays(self): mgr = RelayManager(relay_urls=["wss://a.relay", "wss://b.relay"]) ws_a = AsyncMock() ws_a.open = True ws_a.send = AsyncMock() ws_b = AsyncMock() ws_b.open = True ws_b.send = AsyncMock() mgr._connections = {"wss://a.relay": ws_a, "wss://b.relay": ws_b} eose = json.dumps(["EOSE", "sub"]) ws_a.recv = AsyncMock(side_effect=[ json.dumps(["EVENT", "sub", {"id": "evt_a", "kind": 1, "content": "a"}]), eose, TimeoutError(), ]) ws_b.recv = AsyncMock(side_effect=[ json.dumps(["EVENT", "sub", {"id": "evt_b", "kind": 1, "content": "b"}]), eose, TimeoutError(), ]) results = await mgr.query([{"kinds": [1]}], timeout=1.0) assert len(results) == 2 @pytest.mark.asyncio async def test_query_dedup(self): mgr = RelayManager(relay_urls=["wss://a.relay", "wss://b.relay"]) ws_a = AsyncMock() ws_a.open = True ws_a.send = AsyncMock() ws_b = AsyncMock() ws_b.open = True ws_b.send = AsyncMock() mgr._connections = {"wss://a.relay": ws_a, "wss://b.relay": ws_b} dup = json.dumps(["EVENT", "sub", {"id": "duplicate", "kind": 1, "content": "dup"}]) eose = json.dumps(["EOSE", "sub"]) ws_a.recv = AsyncMock(side_effect=[dup, eose, TimeoutError()]) ws_b.recv = AsyncMock(side_effect=[dup, eose, TimeoutError()]) results = await mgr.query([{"kinds": [1]}], timeout=1.0) assert len(results) == 1 @pytest.mark.asyncio async def test_query_no_active(self): mgr = RelayManager(relay_urls=["wss://a.relay"]) mgr._connections = {} results = await mgr.query([{"kinds": [1]}]) assert results == [] class TestRelayManagerSendAuth: @pytest.mark.asyncio async def test_send_auth_ok_response(self): mgr = RelayManager(relay_urls=["wss://auth.relay"]) ws = AsyncMock() ws.open = True ws.send = AsyncMock() ws.recv = AsyncMock(return_value=json.dumps(["OK", "event_id", True, ""])) mgr._connections = {"wss://auth.relay": ws} result = await mgr.send_auth({"id": "auth_event", "kind": 22242}) assert result["wss://auth.relay"] is True @pytest.mark.asyncio async def test_send_auth_no_connection(self): mgr = RelayManager(relay_urls=["wss://auth.relay"]) mgr._connections = {} result = await mgr.send_auth({"id": "auth_event", "kind": 22242}, urls=["wss://auth.relay"]) assert result["wss://auth.relay"] is False class TestRelayManagerErrorClassification: def test_classify_timeout(self): mgr = RelayManager(relay_urls=["wss://a.relay"]) assert mgr._classify_error("wss://a.relay", TimeoutError()) == "timeout" def test_classify_connection(self): mgr = RelayManager(relay_urls=["wss://a.relay"]) assert mgr._classify_error("wss://a.relay", ConnectionError()) == "connection" def test_classify_unknown(self): mgr = RelayManager(relay_urls=["wss://a.relay"]) assert mgr._classify_error("wss://a.relay", RuntimeError("unknown")) == "unknown"