Merge pull request #30 from pselamy/fix/3

feat: implement WebSocket trade stream handler (#3)
This commit is contained in:
Patrick Selamy
2026-01-04 14:47:14 -05:00
committed by GitHub
6 changed files with 938 additions and 2 deletions
+1
View File
@@ -13,6 +13,7 @@ dependencies = [
"alembic>=1.13.0",
"pydantic>=2.0.0",
"python-dotenv>=1.0.0",
"websockets>=12.0",
]
[project.optional-dependencies]
@@ -10,14 +10,29 @@ from polymarket_insider_tracker.ingestor.models import (
Orderbook,
OrderbookLevel,
Token,
TradeEvent,
)
from polymarket_insider_tracker.ingestor.websocket import (
ConnectionState,
StreamStats,
TradeStreamError,
TradeStreamHandler,
)
__all__ = [
# CLOB Client
"ClobClient",
"ClobClientError",
"RetryError",
# Models
"Market",
"Orderbook",
"OrderbookLevel",
"Token",
"TradeEvent",
# WebSocket
"ConnectionState",
"StreamStats",
"TradeStreamError",
"TradeStreamHandler",
]
@@ -1,9 +1,9 @@
"""Data models for the ingestor module."""
from dataclasses import dataclass, field
from datetime import datetime
from datetime import datetime, timezone
from decimal import Decimal
from typing import Any
from typing import Any, Literal
@dataclass(frozen=True)
@@ -138,3 +138,91 @@ class Orderbook:
if self.best_bid is not None and self.best_ask is not None:
return (self.best_bid + self.best_ask) / 2
return None
@dataclass(frozen=True)
class TradeEvent:
"""Represents a trade event from the Polymarket WebSocket feed.
This captures all the information about a single trade execution,
including the market, wallet, trade details, and metadata.
"""
# Core trade identifiers
market_id: str # conditionId - the market/CTF condition ID
trade_id: str # transactionHash - unique trade identifier
wallet_address: str # proxyWallet - trader's wallet address
# Trade details
side: Literal["BUY", "SELL"]
outcome: str # Human-readable outcome (e.g., "Yes", "No")
outcome_index: int # Index of the outcome (0 or 1)
price: Decimal
size: Decimal # Number of shares traded
timestamp: datetime
# Asset information
asset_id: str # ERC1155 token ID
# Market metadata
market_slug: str = ""
event_slug: str = ""
event_title: str = ""
# Trader metadata (optional - may not be available for all trades)
trader_name: str = ""
trader_pseudonym: str = ""
@classmethod
def from_websocket_message(cls, data: dict[str, Any]) -> "TradeEvent":
"""Create a TradeEvent from a WebSocket activity/trade message.
Args:
data: The payload from a WebSocket trade message.
Returns:
TradeEvent instance.
"""
# Parse timestamp - it's a Unix timestamp in seconds
raw_timestamp = data.get("timestamp", 0)
if isinstance(raw_timestamp, int):
timestamp = datetime.fromtimestamp(raw_timestamp, tz=timezone.utc)
else:
timestamp = datetime.now(timezone.utc)
# Parse side - normalize to uppercase
side_raw = str(data.get("side", "BUY")).upper()
side: Literal["BUY", "SELL"] = "BUY" if side_raw == "BUY" else "SELL"
return cls(
market_id=str(data.get("conditionId", "")),
trade_id=str(data.get("transactionHash", "")),
wallet_address=str(data.get("proxyWallet", "")),
side=side,
outcome=str(data.get("outcome", "")),
outcome_index=int(data.get("outcomeIndex", 0)),
price=Decimal(str(data.get("price", 0))),
size=Decimal(str(data.get("size", 0))),
timestamp=timestamp,
asset_id=str(data.get("asset", "")),
market_slug=str(data.get("slug", "")),
event_slug=str(data.get("eventSlug", "")),
event_title=str(data.get("title", "")),
trader_name=str(data.get("name", "")),
trader_pseudonym=str(data.get("pseudonym", "")),
)
@property
def is_buy(self) -> bool:
"""Return True if this is a buy trade."""
return self.side == "BUY"
@property
def is_sell(self) -> bool:
"""Return True if this is a sell trade."""
return self.side == "SELL"
@property
def notional_value(self) -> Decimal:
"""Return the notional value of the trade (price * size)."""
return self.price * self.size
@@ -0,0 +1,338 @@
"""WebSocket client for streaming Polymarket trade events."""
import asyncio
import json
import logging
import time
from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from enum import Enum
from typing import Any
import websockets
from websockets.asyncio.client import ClientConnection
from polymarket_insider_tracker.ingestor.models import TradeEvent
logger = logging.getLogger(__name__)
# Constants
DEFAULT_WS_HOST = "wss://ws-live-data.polymarket.com"
DEFAULT_PING_INTERVAL = 30 # seconds
DEFAULT_MAX_RECONNECT_DELAY = 30 # seconds
DEFAULT_INITIAL_RECONNECT_DELAY = 1 # seconds
class ConnectionState(Enum):
"""WebSocket connection states."""
DISCONNECTED = "disconnected"
CONNECTING = "connecting"
CONNECTED = "connected"
RECONNECTING = "reconnecting"
@dataclass
class StreamStats:
"""Statistics about the trade stream."""
trades_received: int = 0
reconnect_count: int = 0
last_trade_time: float | None = None
connected_since: float | None = None
last_error: str | None = None
class TradeStreamError(Exception):
"""Base exception for trade stream errors."""
class ConnectionError(TradeStreamError):
"""Raised when connection to WebSocket fails."""
TradeCallback = Callable[[TradeEvent], Awaitable[None]]
StateCallback = Callable[[ConnectionState], Awaitable[None]]
class TradeStreamHandler:
"""WebSocket client for streaming Polymarket trade events.
This handler maintains a persistent connection to Polymarket's real-time
trade feed, automatically reconnecting on disconnection with exponential
backoff.
Example:
>>> async def on_trade(trade: TradeEvent):
... print(f"Trade: {trade.side} {trade.size} @ {trade.price}")
...
>>> handler = TradeStreamHandler(on_trade=on_trade)
>>> await handler.start() # Blocks until stop() is called
Attributes:
state: Current connection state.
stats: Statistics about the stream (trades received, reconnects, etc.).
"""
def __init__(
self,
on_trade: TradeCallback,
*,
host: str = DEFAULT_WS_HOST,
on_state_change: StateCallback | None = None,
ping_interval: int = DEFAULT_PING_INTERVAL,
max_reconnect_delay: int = DEFAULT_MAX_RECONNECT_DELAY,
initial_reconnect_delay: int = DEFAULT_INITIAL_RECONNECT_DELAY,
event_filter: str | None = None,
market_filter: str | None = None,
) -> None:
"""Initialize the trade stream handler.
Args:
on_trade: Async callback invoked for each trade event.
host: WebSocket endpoint URL.
on_state_change: Optional callback for connection state changes.
ping_interval: Seconds between heartbeat pings.
max_reconnect_delay: Maximum delay between reconnection attempts.
initial_reconnect_delay: Initial delay for reconnection backoff.
event_filter: Optional event slug to filter trades by event.
market_filter: Optional market slug to filter trades by market.
"""
self._on_trade = on_trade
self._on_state_change = on_state_change
self._host = host
self._ping_interval = ping_interval
self._max_reconnect_delay = max_reconnect_delay
self._initial_reconnect_delay = initial_reconnect_delay
self._event_filter = event_filter
self._market_filter = market_filter
self._state = ConnectionState.DISCONNECTED
self._stats = StreamStats()
self._ws: ClientConnection | None = None
self._running = False
self._stop_event: asyncio.Event | None = None
@property
def state(self) -> ConnectionState:
"""Current connection state."""
return self._state
@property
def stats(self) -> StreamStats:
"""Stream statistics."""
return self._stats
async def _set_state(self, new_state: ConnectionState) -> None:
"""Update state and notify callback."""
if self._state != new_state:
old_state = self._state
self._state = new_state
logger.info("Connection state: %s -> %s", old_state.value, new_state.value)
if self._on_state_change:
try:
await self._on_state_change(new_state)
except Exception as e:
logger.error("Error in state change callback: %s", e)
def _build_subscription_message(self) -> dict[str, Any]:
"""Build the WebSocket subscription message."""
subscription: dict[str, Any] = {
"topic": "activity",
"type": "trades",
}
# Add filters if specified
if self._event_filter:
subscription["filters"] = json.dumps({"event_slug": self._event_filter})
elif self._market_filter:
subscription["filters"] = json.dumps({"market_slug": self._market_filter})
return {"subscriptions": [subscription]}
async def _connect(self) -> ClientConnection:
"""Establish WebSocket connection."""
await self._set_state(ConnectionState.CONNECTING)
try:
ws = await websockets.connect(
self._host,
ping_interval=self._ping_interval,
ping_timeout=self._ping_interval * 2,
)
# Send subscription message
subscribe_msg = self._build_subscription_message()
await ws.send(json.dumps(subscribe_msg))
logger.info("Connected to %s and subscribed to trades", self._host)
await self._set_state(ConnectionState.CONNECTED)
self._stats.connected_since = time.time()
return ws
except Exception as e:
logger.error("Failed to connect: %s", e)
self._stats.last_error = str(e)
raise ConnectionError(f"Failed to connect to {self._host}: {e}") from e
async def _handle_message(self, message: str) -> None:
"""Parse and process an incoming WebSocket message."""
try:
data = json.loads(message)
# Check if this is a trade message
topic = data.get("topic")
msg_type = data.get("type")
if topic == "activity" and msg_type == "trades":
payload = data.get("payload", {})
trade = TradeEvent.from_websocket_message(payload)
self._stats.trades_received += 1
self._stats.last_trade_time = time.time()
logger.debug(
"Trade: %s %s @ %s on %s",
trade.side,
trade.size,
trade.price,
trade.market_slug,
)
try:
await self._on_trade(trade)
except Exception as e:
logger.error("Error in trade callback: %s", e)
else:
# Log other message types for debugging
logger.debug("Received message: topic=%s type=%s", topic, msg_type)
except json.JSONDecodeError as e:
logger.warning("Invalid JSON message: %s", e)
except Exception as e:
logger.error("Error processing message: %s", e)
async def _listen(self, ws: ClientConnection) -> None:
"""Listen for messages on the WebSocket."""
try:
async for message in ws:
if not self._running:
break
if isinstance(message, str):
await self._handle_message(message)
else:
logger.debug("Received binary message (%d bytes)", len(message))
except websockets.ConnectionClosed as e:
logger.warning("Connection closed: %s", e)
raise
except Exception as e:
logger.error("Error in message loop: %s", e)
raise
async def _reconnect_loop(self) -> None:
"""Handle reconnection with exponential backoff."""
delay = self._initial_reconnect_delay
while self._running:
try:
await self._set_state(ConnectionState.RECONNECTING)
logger.info("Reconnecting in %.1f seconds...", delay)
await asyncio.sleep(delay)
if not self._running:
break
self._ws = await self._connect()
self._stats.reconnect_count += 1
delay = self._initial_reconnect_delay # Reset delay on success
return
except Exception as e:
logger.error("Reconnection failed: %s", e)
self._stats.last_error = str(e)
# Exponential backoff with jitter
delay = min(delay * 2, self._max_reconnect_delay)
async def start(self) -> None:
"""Connect and begin streaming trades.
This method blocks until stop() is called. It automatically
handles reconnection on disconnection.
Raises:
ConnectionError: If initial connection fails.
"""
if self._running:
logger.warning("Handler already running")
return
self._running = True
self._stop_event = asyncio.Event()
try:
# Initial connection
self._ws = await self._connect()
# Main loop
while self._running:
try:
await self._listen(self._ws)
except (websockets.ConnectionClosed, Exception) as e:
if not self._running:
break
logger.warning("Connection lost: %s", e)
await self._set_state(ConnectionState.DISCONNECTED)
# Attempt reconnection
await self._reconnect_loop()
if not self._running or self._ws is None:
break
finally:
await self._cleanup()
async def stop(self) -> None:
"""Gracefully disconnect from the WebSocket.
This signals the handler to stop and cleanly close the connection.
"""
if not self._running:
return
logger.info("Stopping trade stream handler...")
self._running = False
if self._stop_event:
self._stop_event.set()
await self._cleanup()
async def _cleanup(self) -> None:
"""Clean up resources."""
if self._ws:
try:
await self._ws.close()
except Exception as e:
logger.debug("Error closing WebSocket: %s", e)
finally:
self._ws = None
await self._set_state(ConnectionState.DISCONNECTED)
logger.info("Trade stream handler stopped")
async def __aenter__(self) -> "TradeStreamHandler":
"""Async context manager entry."""
return self
async def __aexit__(self, *args: Any) -> None:
"""Async context manager exit."""
await self.stop()
+151
View File
@@ -11,6 +11,7 @@ from polymarket_insider_tracker.ingestor.models import (
Orderbook,
OrderbookLevel,
Token,
TradeEvent,
)
@@ -267,3 +268,153 @@ class TestOrderbook:
)
assert orderbook.midpoint is None
class TestTradeEvent:
"""Tests for TradeEvent model."""
def test_from_websocket_message_full(self) -> None:
"""Test creating TradeEvent from a full WebSocket message."""
data = {
"conditionId": "0xmarket123",
"transactionHash": "0xtx456",
"proxyWallet": "0xwallet789",
"side": "BUY",
"outcome": "Yes",
"outcomeIndex": 0,
"price": 0.65,
"size": 100,
"timestamp": 1704067200, # 2024-01-01 00:00:00 UTC
"asset": "token123",
"slug": "will-it-rain",
"eventSlug": "weather-markets",
"title": "Weather Predictions",
"name": "Alice",
"pseudonym": "AliceTrader",
}
trade = TradeEvent.from_websocket_message(data)
assert trade.market_id == "0xmarket123"
assert trade.trade_id == "0xtx456"
assert trade.wallet_address == "0xwallet789"
assert trade.side == "BUY"
assert trade.outcome == "Yes"
assert trade.outcome_index == 0
assert trade.price == Decimal("0.65")
assert trade.size == Decimal("100")
assert trade.timestamp == datetime(2024, 1, 1, 0, 0, 0, tzinfo=timezone.utc)
assert trade.asset_id == "token123"
assert trade.market_slug == "will-it-rain"
assert trade.event_slug == "weather-markets"
assert trade.event_title == "Weather Predictions"
assert trade.trader_name == "Alice"
assert trade.trader_pseudonym == "AliceTrader"
def test_from_websocket_message_minimal(self) -> None:
"""Test creating TradeEvent with minimal data."""
data = {
"conditionId": "0x123",
"transactionHash": "0xtx",
"proxyWallet": "0xwallet",
"side": "SELL",
"outcome": "No",
"price": 0.35,
"size": 50,
}
trade = TradeEvent.from_websocket_message(data)
assert trade.market_id == "0x123"
assert trade.side == "SELL"
assert trade.outcome == "No"
assert trade.price == Decimal("0.35")
assert trade.size == Decimal("50")
assert trade.market_slug == ""
assert trade.trader_name == ""
def test_from_websocket_message_lowercase_side(self) -> None:
"""Test that lowercase side is normalized."""
data = {
"side": "buy",
"price": 0.5,
"size": 10,
}
trade = TradeEvent.from_websocket_message(data)
assert trade.side == "BUY"
def test_from_websocket_message_sell_side(self) -> None:
"""Test SELL side handling."""
data = {
"side": "sell",
"price": 0.5,
"size": 10,
}
trade = TradeEvent.from_websocket_message(data)
assert trade.side == "SELL"
def test_is_buy(self) -> None:
"""Test is_buy property."""
buy_trade = TradeEvent(
market_id="",
trade_id="",
wallet_address="",
side="BUY",
outcome="",
outcome_index=0,
price=Decimal("0.5"),
size=Decimal("10"),
timestamp=datetime.now(timezone.utc),
asset_id="",
)
sell_trade = TradeEvent(
market_id="",
trade_id="",
wallet_address="",
side="SELL",
outcome="",
outcome_index=0,
price=Decimal("0.5"),
size=Decimal("10"),
timestamp=datetime.now(timezone.utc),
asset_id="",
)
assert buy_trade.is_buy is True
assert buy_trade.is_sell is False
assert sell_trade.is_buy is False
assert sell_trade.is_sell is True
def test_notional_value(self) -> None:
"""Test notional value calculation."""
trade = TradeEvent(
market_id="",
trade_id="",
wallet_address="",
side="BUY",
outcome="",
outcome_index=0,
price=Decimal("0.65"),
size=Decimal("100"),
timestamp=datetime.now(timezone.utc),
asset_id="",
)
assert trade.notional_value == Decimal("65")
def test_frozen(self) -> None:
"""Test that TradeEvent is immutable."""
trade = TradeEvent(
market_id="0x123",
trade_id="0xtx",
wallet_address="0xwallet",
side="BUY",
outcome="Yes",
outcome_index=0,
price=Decimal("0.5"),
size=Decimal("10"),
timestamp=datetime.now(timezone.utc),
asset_id="token",
)
with pytest.raises(AttributeError):
trade.market_id = "0x456" # type: ignore[misc]
+343
View File
@@ -0,0 +1,343 @@
"""Tests for WebSocket trade stream handler."""
import asyncio
import json
from decimal import Decimal
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from polymarket_insider_tracker.ingestor.models import TradeEvent
from polymarket_insider_tracker.ingestor.websocket import (
ConnectionState,
StreamStats,
TradeStreamHandler,
)
class TestStreamStats:
"""Tests for StreamStats."""
def test_defaults(self) -> None:
"""Test default values."""
stats = StreamStats()
assert stats.trades_received == 0
assert stats.reconnect_count == 0
assert stats.last_trade_time is None
assert stats.connected_since is None
assert stats.last_error is None
class TestTradeStreamHandler:
"""Tests for TradeStreamHandler."""
@pytest.fixture
def on_trade_mock(self) -> AsyncMock:
"""Create mock trade callback."""
return AsyncMock()
@pytest.fixture
def on_state_change_mock(self) -> AsyncMock:
"""Create mock state change callback."""
return AsyncMock()
@pytest.fixture
def handler(
self, on_trade_mock: AsyncMock, on_state_change_mock: AsyncMock
) -> TradeStreamHandler:
"""Create handler with mocks."""
return TradeStreamHandler(
on_trade=on_trade_mock,
on_state_change=on_state_change_mock,
initial_reconnect_delay=0.01, # Fast reconnect for tests
max_reconnect_delay=0.1,
)
def test_init_defaults(self, on_trade_mock: AsyncMock) -> None:
"""Test handler initialization with defaults."""
handler = TradeStreamHandler(on_trade=on_trade_mock)
assert handler.state == ConnectionState.DISCONNECTED
assert handler.stats.trades_received == 0
assert handler._host == "wss://ws-live-data.polymarket.com"
def test_init_custom_host(self, on_trade_mock: AsyncMock) -> None:
"""Test handler with custom host."""
handler = TradeStreamHandler(
on_trade=on_trade_mock,
host="wss://custom.example.com",
)
assert handler._host == "wss://custom.example.com"
def test_init_with_event_filter(self, on_trade_mock: AsyncMock) -> None:
"""Test handler with event filter."""
handler = TradeStreamHandler(
on_trade=on_trade_mock,
event_filter="presidential-election-2024",
)
assert handler._event_filter == "presidential-election-2024"
def test_build_subscription_message_no_filter(
self, handler: TradeStreamHandler
) -> None:
"""Test building subscription message without filters."""
msg = handler._build_subscription_message()
assert msg == {
"subscriptions": [{"topic": "activity", "type": "trades"}]
}
def test_build_subscription_message_with_event_filter(
self, on_trade_mock: AsyncMock
) -> None:
"""Test building subscription message with event filter."""
handler = TradeStreamHandler(
on_trade=on_trade_mock,
event_filter="test-event",
)
msg = handler._build_subscription_message()
assert msg["subscriptions"][0]["filters"] == json.dumps(
{"event_slug": "test-event"}
)
def test_build_subscription_message_with_market_filter(
self, on_trade_mock: AsyncMock
) -> None:
"""Test building subscription message with market filter."""
handler = TradeStreamHandler(
on_trade=on_trade_mock,
market_filter="test-market",
)
msg = handler._build_subscription_message()
assert msg["subscriptions"][0]["filters"] == json.dumps(
{"market_slug": "test-market"}
)
@pytest.mark.asyncio
async def test_handle_message_trade(
self, handler: TradeStreamHandler, on_trade_mock: AsyncMock
) -> None:
"""Test handling a valid trade message."""
message = json.dumps(
{
"topic": "activity",
"type": "trades",
"payload": {
"conditionId": "0xmarket",
"transactionHash": "0xtx",
"proxyWallet": "0xwallet",
"side": "BUY",
"outcome": "Yes",
"price": 0.65,
"size": 100,
"timestamp": 1704067200,
"asset": "token123",
},
}
)
await handler._handle_message(message)
on_trade_mock.assert_called_once()
trade: TradeEvent = on_trade_mock.call_args[0][0]
assert trade.market_id == "0xmarket"
assert trade.side == "BUY"
assert trade.price == Decimal("0.65")
assert handler.stats.trades_received == 1
@pytest.mark.asyncio
async def test_handle_message_non_trade(
self, handler: TradeStreamHandler, on_trade_mock: AsyncMock
) -> None:
"""Test handling a non-trade message."""
message = json.dumps(
{
"topic": "comments",
"type": "comment_created",
"payload": {"body": "Hello"},
}
)
await handler._handle_message(message)
on_trade_mock.assert_not_called()
assert handler.stats.trades_received == 0
@pytest.mark.asyncio
async def test_handle_message_invalid_json(
self, handler: TradeStreamHandler, on_trade_mock: AsyncMock
) -> None:
"""Test handling invalid JSON message."""
await handler._handle_message("not valid json")
on_trade_mock.assert_not_called()
assert handler.stats.trades_received == 0
@pytest.mark.asyncio
async def test_handle_message_callback_error(
self, handler: TradeStreamHandler, on_trade_mock: AsyncMock
) -> None:
"""Test that callback errors don't crash the handler."""
on_trade_mock.side_effect = ValueError("Callback error")
message = json.dumps(
{
"topic": "activity",
"type": "trades",
"payload": {
"conditionId": "0x",
"transactionHash": "0x",
"proxyWallet": "0x",
"side": "BUY",
"price": 0.5,
"size": 10,
},
}
)
# Should not raise
await handler._handle_message(message)
# Trade was still counted
assert handler.stats.trades_received == 1
@pytest.mark.asyncio
async def test_set_state_calls_callback(
self, handler: TradeStreamHandler, on_state_change_mock: AsyncMock
) -> None:
"""Test that state changes trigger callback."""
await handler._set_state(ConnectionState.CONNECTING)
on_state_change_mock.assert_called_once_with(ConnectionState.CONNECTING)
assert handler.state == ConnectionState.CONNECTING
@pytest.mark.asyncio
async def test_set_state_same_state_no_callback(
self, handler: TradeStreamHandler, on_state_change_mock: AsyncMock
) -> None:
"""Test that same state doesn't trigger callback."""
handler._state = ConnectionState.CONNECTED
await handler._set_state(ConnectionState.CONNECTED)
on_state_change_mock.assert_not_called()
@pytest.mark.asyncio
async def test_stop_when_not_running(self, handler: TradeStreamHandler) -> None:
"""Test stop when handler is not running."""
# Should not raise
await handler.stop()
@pytest.mark.asyncio
async def test_connect_sends_subscription(
self, handler: TradeStreamHandler, on_state_change_mock: AsyncMock
) -> None:
"""Test that connection sends subscription message."""
mock_ws = AsyncMock()
mock_ws.send = AsyncMock()
with patch("websockets.connect", AsyncMock(return_value=mock_ws)):
ws = await handler._connect()
assert ws is mock_ws
mock_ws.send.assert_called_once()
# Verify subscription message
sent_msg = json.loads(mock_ws.send.call_args[0][0])
assert "subscriptions" in sent_msg
assert sent_msg["subscriptions"][0]["topic"] == "activity"
assert sent_msg["subscriptions"][0]["type"] == "trades"
@pytest.mark.asyncio
async def test_cleanup_closes_websocket(
self, handler: TradeStreamHandler
) -> None:
"""Test that cleanup closes the WebSocket."""
mock_ws = AsyncMock()
mock_ws.close = AsyncMock()
handler._ws = mock_ws
await handler._cleanup()
mock_ws.close.assert_called_once()
assert handler._ws is None
assert handler.state == ConnectionState.DISCONNECTED
@pytest.mark.asyncio
async def test_context_manager(self, on_trade_mock: AsyncMock) -> None:
"""Test async context manager."""
handler = TradeStreamHandler(on_trade=on_trade_mock)
async with handler:
pass
# Should be stopped after exiting context
assert handler._running is False
class TestTradeStreamHandlerIntegration:
"""Integration tests for TradeStreamHandler.
These tests verify the full message flow with mocked WebSocket.
"""
@pytest.mark.asyncio
async def test_start_and_receive_trades(self) -> None:
"""Test starting handler and receiving trades."""
received_trades: list[TradeEvent] = []
async def on_trade(trade: TradeEvent) -> None:
received_trades.append(trade)
handler = TradeStreamHandler(
on_trade=on_trade,
initial_reconnect_delay=0.01,
)
# Create mock WebSocket that sends one trade then closes
mock_ws = MagicMock()
mock_ws.send = AsyncMock()
mock_ws.close = AsyncMock()
trade_message = json.dumps(
{
"topic": "activity",
"type": "trades",
"payload": {
"conditionId": "0xtest",
"transactionHash": "0xtx",
"proxyWallet": "0xwallet",
"side": "BUY",
"outcome": "Yes",
"price": 0.75,
"size": 50,
"timestamp": 1704067200,
"asset": "token",
},
}
)
# Mock async iteration
async def mock_iter() -> None:
yield trade_message
await handler.stop() # Stop after first message
mock_ws.__aiter__ = mock_iter
with patch("websockets.connect", AsyncMock(return_value=mock_ws)):
# Run with timeout to prevent hanging
try:
await asyncio.wait_for(handler.start(), timeout=1.0)
except asyncio.TimeoutError:
await handler.stop()
# Verify trade was received
assert len(received_trades) == 1
assert received_trades[0].market_id == "0xtest"
assert received_trades[0].side == "BUY"
assert received_trades[0].price == Decimal("0.75")