diff --git a/CMakeLists.txt b/CMakeLists.txt index 876647b..1f2ab82 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -123,6 +123,7 @@ set(POLYMARKET_CLIENT_SOURCES src/websocket_client_transport.cpp src/websocket_resilience.cpp src/websocket_market_data.cpp + src/websocket_user_data.cpp src/market_fetcher.cpp src/market_discovery.cpp src/market_time.cpp @@ -132,6 +133,10 @@ set(POLYMARKET_CLIENT_SOURCES src/orderbook_runtime_stream.cpp src/orderbook_subscription.cpp src/orderbook_stream.cpp + src/user_stream.cpp + src/user_stream_runtime.cpp + src/user_stream_runtime_events.cpp + src/user_stream_subscription.cpp src/order_amounts.cpp src/clob_order_execution.cpp src/clob_order_submission.cpp diff --git a/README.md b/README.md index f5b091f..af682d1 100644 --- a/README.md +++ b/README.md @@ -10,6 +10,7 @@ Reusable C++20 client for Polymarket: REST, WebSocket streaming, and order signi - **REST**: market discovery, orderbook/price queries, auth key management, and trading endpoints. - **Transport controls**: configurable libcurl timeouts, keepalive, connection reuse, proxy/user-agent, request metrics, and cumulative stats. - **WebSocket**: orderbook streaming via IXWebSocket with reconnect, subscription replay, typed callbacks, and backpressure counters. +- **User stream**: authenticated CLOB user channel (`UserStream`) with typed order/trade events, per-market or all-market subscriptions, plus gap callbacks (invalidate state) and recovery callbacks (reconcile via REST once the subscription is restored). - **Signing**: CLOB V2 EIP-712 order signing (secp256k1, keccak). - **Decimal math**: shared scaled-integer conversion for trading amounts. - **Structured errors**: opt-in `Result` APIs with typed SDK error classification. @@ -135,6 +136,7 @@ int main() { - `rest_example`: fetch markets from CLOB REST - `sign_example`: sign a dummy order (requires `PRIVATE_KEY`) - `ws_example`: connect to Polymarket WS and subscribe to orderbook agg +- `user_stream_example`: stream your own order and trade events (requires `PRIVATE_KEY`; optional condition IDs as arguments) - `uma_oracle_watch`: stream UMA adapter lifecycle events over Polygon JSON-RPC - `condition_resolution_watch`: stream Conditional Tokens resolution/redemption events - `evm_event_indexer_example`: persistent HTTP catch-up + live WS indexer with a cursor file diff --git a/cmake/PolymarketExamples.cmake b/cmake/PolymarketExamples.cmake index 3db3cd6..3c3fa2e 100644 --- a/cmake/PolymarketExamples.cmake +++ b/cmake/PolymarketExamples.cmake @@ -36,6 +36,9 @@ if(POLYMARKET_CLIENT_BUILD_EXAMPLES) add_executable(ws_example examples/ws_example.cpp) target_link_libraries(ws_example PRIVATE polymarket::client) + add_executable(user_stream_example examples/user_stream_example.cpp) + target_link_libraries(user_stream_example PRIVATE polymarket::client) + add_executable(uma_oracle_watch examples/uma_oracle_watch.cpp) target_link_libraries(uma_oracle_watch PRIVATE polymarket::client) diff --git a/cmake/PolymarketTests.cmake b/cmake/PolymarketTests.cmake index a5091f1..95c63b9 100644 --- a/cmake/PolymarketTests.cmake +++ b/cmake/PolymarketTests.cmake @@ -113,6 +113,16 @@ if(POLYMARKET_CLIENT_BUILD_TESTS) target_link_libraries(test_orderbook_stream PRIVATE polymarket::client) add_test(NAME test_orderbook_stream COMMAND test_orderbook_stream) + add_executable(test_user_stream + tests/test_user_stream.cpp + tests/test_user_stream_connection.cpp + tests/test_user_stream_failures.cpp) + target_include_directories(test_user_stream PRIVATE + ${CMAKE_CURRENT_SOURCE_DIR}/src + ${CMAKE_CURRENT_SOURCE_DIR}/tests) + target_link_libraries(test_user_stream PRIVATE polymarket::client) + add_test(NAME test_user_stream COMMAND test_user_stream) + add_executable(test_orderbook_owner_reset tests/test_orderbook_owner_reset.cpp) target_include_directories(test_orderbook_owner_reset PRIVATE ${CMAKE_CURRENT_SOURCE_DIR}/tests) diff --git a/examples/user_stream_example.cpp b/examples/user_stream_example.cpp new file mode 100644 index 0000000..4cfb0ee --- /dev/null +++ b/examples/user_stream_example.cpp @@ -0,0 +1,88 @@ +#include "clob_client.hpp" +#include "user_stream.hpp" + +#include +#include +#include +#include + +int main(int argc, char **argv) +{ + using namespace polymarket; + + const char *pk_env = std::getenv("PRIVATE_KEY"); + const char *key_env = std::getenv("POLY_API_KEY"); + const char *secret_env = std::getenv("POLY_API_SECRET"); + const char *passphrase_env = std::getenv("POLY_API_PASSPHRASE"); + const bool has_api_credentials = key_env && secret_env && passphrase_env; + if (!pk_env && !has_api_credentials) + { + std::cout << "PRIVATE_KEY or POLY_API_KEY/POLY_API_SECRET/POLY_API_PASSPHRASE " + "not set; skipping user stream example.\n"; + return 0; + } + + try + { + ApiCredentials credentials; + if (has_api_credentials) + { + credentials = {key_env, secret_env, passphrase_env}; + } + else + { + // Derive L2 API credentials from the signer (no orders are placed). + ClobClient client("https://clob.polymarket.com", 137, pk_env); + credentials = client.create_or_derive_api_key(); + } + + UserStream stream(Config{}, credentials); + stream.on_order([](const UserOrderEvent &order) + { std::cout << "[order] " << order.type << ' ' << order.side << ' ' + << order.size_matched << '/' << order.original_size + << " @ " << order.price << " id=" << order.id << '\n'; }); + stream.on_trade([](const UserTradeEvent &trade) + { std::cout << "[trade] " << trade.status << ' ' << trade.side << ' ' + << trade.size << " @ " << trade.price << " id=" << trade.id << '\n'; }); + stream.on_error([](const std::string &error) + { std::cerr << "[error] " << error << '\n'; }); + stream.on_stream_gap([] + { std::cout << "[gap] events may have been missed; local state is stale\n"; }); + stream.on_stream_recovered([] + { std::cout << "[recovered] subscription restored; reconcile via REST\n"; }); + + // Optional condition IDs narrow the stream; none means all markets. + if (argc > 1) + { + for (int i = 1; i < argc; ++i) + stream.subscribe(argv[i]); + } + else + { + stream.subscribe_all_markets(); + } + + if (!stream.connect()) + { + std::cerr << "failed to connect user stream\n"; + return 1; + } + + const char *seconds_env = std::getenv("USER_STREAM_SECONDS"); + const auto deadline = std::chrono::steady_clock::now() + + std::chrono::seconds(seconds_env ? std::atoi(seconds_env) : 60); + while (std::chrono::steady_clock::now() < deadline && !stream.authentication_failed()) + std::this_thread::sleep_for(std::chrono::milliseconds(100)); + stream.stop(); + if (stream.authentication_failed()) + return 1; + std::cout << "orders=" << stream.order_events() + << " trades=" << stream.trade_events() << '\n'; + } + catch (const std::exception &e) + { + std::cerr << "Error: " << e.what() << "\n"; + return 1; + } + return 0; +} diff --git a/include/types.hpp b/include/types.hpp index ccc30fd..873c9eb 100644 --- a/include/types.hpp +++ b/include/types.hpp @@ -262,6 +262,7 @@ namespace polymarket std::string clob_ws_url = "wss://ws-subscriptions-clob.polymarket.com/ws/market"; std::string gamma_api_url = "https://gamma-api.polymarket.com"; std::string rtds_ws_url = "wss://ws-live-data.polymarket.com"; + std::string clob_user_ws_url = "wss://ws-subscriptions-clob.polymarket.com/ws/user"; // Trading parameters double trigger_combined = 0.98; diff --git a/include/user_stream.hpp b/include/user_stream.hpp new file mode 100644 index 0000000..dfd74d4 --- /dev/null +++ b/include/user_stream.hpp @@ -0,0 +1,137 @@ +#pragma once + +#include "clob_types.hpp" +#include "order_signer.hpp" +#include "types.hpp" + +#include +#include +#include +#include +#include +#include + +namespace polymarket +{ + namespace detail + { + class UserStreamRuntime; + } + + // Order lifecycle event from the authenticated CLOB user channel. + struct UserOrderEvent + { + std::string id; + std::string owner; + std::string market; // condition ID + std::string asset_id; // token ID + std::string side; // BUY or SELL + std::string original_size; + std::string size_matched; + std::string price; + std::string type; // PLACEMENT, UPDATE, or CANCELLATION + std::string status; + std::string order_type; + std::string maker_address; + std::string order_owner; + std::vector associate_trades; + std::string outcome; + std::string created_at; + std::string expiration; + std::string timestamp; + }; + + // Trade lifecycle event (MATCHED, MINED, CONFIRMED, RETRYING, FAILED). + struct UserTradeEvent + { + std::string id; + std::string taker_order_id; + std::string market; // condition ID + std::string asset_id; // token ID + std::string side; + std::string size; + std::string price; + std::string status; + std::string owner; + std::string fee_rate_bps; + std::string match_time; + std::string last_update; + std::string timestamp; + std::string trade_owner; + std::string maker_address; + std::string transaction_hash; + std::optional bucket_index; + std::vector maker_orders; + std::string trader_side; // TAKER or MAKER + std::string outcome; + }; + + using UserOrderCallback = std::function; + using UserTradeCallback = std::function; + // Fired immediately whenever events may have been missed (every connect, + // disconnect, queue overflow, or invalid payload). Treat local order and + // trade state as stale; do not reconcile here, because on reconnect it + // runs before the subscription is restored and the server does not replay + // events missed in between. + using UserStreamGapCallback = std::function; + // Fired once the authenticated subscription has been resent after a gap. + // Reconcile order and trade state via REST (`ClobClient::get_open_orders`, + // `get_trades`) here and merge it with events delivered from this point. + using UserStreamRecoveredCallback = std::function; + // Fired when the server rejects the session (close code 1008, e.g. invalid + // API credentials). The stream stops reconnecting and run() returns; call + // connect() to retry. + using UserStreamErrorCallback = std::function; + + // Authenticated user-channel stream. Subscriptions are replayed with the + // API credentials after every reconnect. + class UserStream + { + public: + UserStream(const Config &config, const ApiCredentials &credentials); + ~UserStream(); + + UserStream(const UserStream &) = delete; + UserStream &operator=(const UserStream &) = delete; + + // Receive events for every market the API key trades. + void subscribe_all_markets(); + // Receive events only for these condition IDs. Ignored for markets + // already covered by subscribe_all_markets(). + void subscribe(const std::vector &condition_ids); + void subscribe(const std::string &condition_id); + void unsubscribe(const std::vector &condition_ids); + void unsubscribe(const std::string &condition_id); + void unsubscribe_all(); + + bool is_subscribed_to_all_markets() const; + std::vector subscribed_markets() const; + + // Callbacks + void on_order(UserOrderCallback callback); + void on_trade(UserTradeCallback callback); + void on_stream_gap(UserStreamGapCallback callback); + void on_stream_recovered(UserStreamRecoveredCallback callback); + void on_error(UserStreamErrorCallback callback); + + // Connection + bool connect(); + void disconnect(); + bool is_connected() const; + bool authentication_failed() const; + + // Run event loop (blocking) + void run(); + + // Stop + void stop(); + + // Statistics + uint64_t order_events() const; + uint64_t trade_events() const; + + private: + std::shared_ptr runtime_; + }; + +} // namespace polymarket diff --git a/include/websocket_client.hpp b/include/websocket_client.hpp index 65b3089..c00e5aa 100644 --- a/include/websocket_client.hpp +++ b/include/websocket_client.hpp @@ -26,6 +26,7 @@ namespace polymarket using OnMessageCallback = std::function; using OnConnectCallback = std::function; using OnDisconnectCallback = std::function; + using OnCloseCallback = std::function; using OnErrorCallback = std::function; using OnSequencedMessageCallback = std::function; using OnStreamGapCallback = std::function; @@ -92,6 +93,7 @@ namespace polymarket void on_typed_message(OnTypedMessageCallback callback); void on_connect(OnConnectCallback callback); void on_disconnect(OnDisconnectCallback callback); + void on_close(OnCloseCallback callback); void on_error(OnErrorCallback callback); void on_stream_gap(OnStreamGapCallback callback); diff --git a/src/user_stream.cpp b/src/user_stream.cpp new file mode 100644 index 0000000..3d8025b --- /dev/null +++ b/src/user_stream.cpp @@ -0,0 +1,140 @@ +#include "user_stream.hpp" +#include "user_stream_runtime.hpp" + +namespace polymarket +{ + UserStream::UserStream(const Config &config, const ApiCredentials &credentials) + : runtime_(detail::UserStreamRuntime::create(config, credentials)) + { + } + + UserStream::~UserStream() + { + auto runtime = std::move(runtime_); + if (runtime) runtime->shutdown(); + } + + void UserStream::subscribe_all_markets() + { + auto runtime = runtime_; + if (runtime) runtime->subscribe_all_markets(); + } + + void UserStream::subscribe(const std::vector &condition_ids) + { + auto runtime = runtime_; + if (runtime) runtime->subscribe(condition_ids); + } + + void UserStream::subscribe(const std::string &condition_id) + { + subscribe(std::vector{condition_id}); + } + + void UserStream::unsubscribe(const std::vector &condition_ids) + { + auto runtime = runtime_; + if (runtime) runtime->unsubscribe(condition_ids); + } + + void UserStream::unsubscribe(const std::string &condition_id) + { + unsubscribe(std::vector{condition_id}); + } + + void UserStream::unsubscribe_all() + { + auto runtime = runtime_; + if (runtime) runtime->unsubscribe_all(); + } + + bool UserStream::is_subscribed_to_all_markets() const + { + auto runtime = runtime_; + return runtime && runtime->is_subscribed_to_all_markets(); + } + + std::vector UserStream::subscribed_markets() const + { + auto runtime = runtime_; + return runtime ? runtime->subscribed_markets() : std::vector{}; + } + + void UserStream::on_order(UserOrderCallback callback) + { + auto runtime = runtime_; + if (runtime) runtime->on_order(std::move(callback)); + } + + void UserStream::on_trade(UserTradeCallback callback) + { + auto runtime = runtime_; + if (runtime) runtime->on_trade(std::move(callback)); + } + + void UserStream::on_stream_gap(UserStreamGapCallback callback) + { + auto runtime = runtime_; + if (runtime) runtime->on_stream_gap(std::move(callback)); + } + + void UserStream::on_stream_recovered(UserStreamRecoveredCallback callback) + { + auto runtime = runtime_; + if (runtime) runtime->on_stream_recovered(std::move(callback)); + } + + void UserStream::on_error(UserStreamErrorCallback callback) + { + auto runtime = runtime_; + if (runtime) runtime->on_error(std::move(callback)); + } + + bool UserStream::connect() + { + auto runtime = runtime_; + return runtime && runtime->connect(); + } + + void UserStream::disconnect() + { + auto runtime = runtime_; + if (runtime) runtime->disconnect(); + } + + bool UserStream::is_connected() const + { + auto runtime = runtime_; + return runtime && runtime->is_connected(); + } + + bool UserStream::authentication_failed() const + { + auto runtime = runtime_; + return runtime && runtime->authentication_failed(); + } + + void UserStream::run() + { + auto runtime = runtime_; + if (runtime) runtime->run(); + } + + void UserStream::stop() + { + auto runtime = runtime_; + if (runtime) runtime->stop(); + } + + uint64_t UserStream::order_events() const + { + auto runtime = runtime_; + return runtime ? runtime->order_events() : 0; + } + + uint64_t UserStream::trade_events() const + { + auto runtime = runtime_; + return runtime ? runtime->trade_events() : 0; + } +} diff --git a/src/user_stream_protocol.hpp b/src/user_stream_protocol.hpp new file mode 100644 index 0000000..7ed249d --- /dev/null +++ b/src/user_stream_protocol.hpp @@ -0,0 +1,38 @@ +#pragma once + +#include "user_stream.hpp" + +#include +#include +#include + +namespace polymarket::detail +{ + struct UserSubscriptionState + { + bool all_markets{false}; + std::vector markets; + + bool active() const { return all_markets || !markets.empty(); } + }; + + // Messages to bring an open socket from one subscription state to another. + // `reconnect` means the server state cannot be narrowed in place, so the + // socket must be reopened and the tracked subscription replayed. + struct UserSubscriptionSync + { + std::vector messages; + bool reconnect{false}; + }; + + using UserStreamEvent = std::variant; + + std::string user_subscription_message(const UserSubscriptionState &state, + const ApiCredentials &credentials); + std::string user_subscription_update_message(const std::vector &markets, + bool subscribe); + UserSubscriptionSync plan_user_subscription_sync(const UserSubscriptionState &before, + const UserSubscriptionState &after, + const ApiCredentials &credentials); + std::vector parse_user_events(const std::string &message); +} diff --git a/src/user_stream_runtime.cpp b/src/user_stream_runtime.cpp new file mode 100644 index 0000000..5d2d690 --- /dev/null +++ b/src/user_stream_runtime.cpp @@ -0,0 +1,277 @@ +#include "user_stream_runtime.hpp" + +#include +#include +#include +#include +#include + +namespace polymarket::detail +{ + namespace + { + Config validated_config(Config config) + { + if (config.clob_user_ws_url.empty()) + { + throw std::invalid_argument("clob_user_ws_url must not be empty"); + } + if (config.ws_ping_interval_ms <= 0) + { + throw std::invalid_argument( + "ws_ping_interval_ms must be positive"); + } + if (config.ws_connect_timeout_ms <= 0) + { + throw std::invalid_argument( + "ws_connect_timeout_ms must be positive"); + } + return config; + } + + ApiCredentials validated_credentials(ApiCredentials credentials) + { + if (credentials.api_key.empty() || credentials.api_secret.empty() || + credentials.api_passphrase.empty()) + { + throw std::invalid_argument( + "user stream requires API key, secret, and passphrase"); + } + return credentials; + } + + void add_markets(std::vector &markets, + const std::vector &condition_ids) + { + for (const auto &condition_id : condition_ids) + { + if (condition_id.empty() || + std::find(markets.begin(), markets.end(), condition_id) != + markets.end()) + continue; + markets.push_back(condition_id); + } + } + } + + std::shared_ptr UserStreamRuntime::create( + const Config &config, const ApiCredentials &credentials) + { + auto runtime = std::shared_ptr( + new UserStreamRuntime(config, credentials)); + runtime->bind_websocket_callbacks(); + return runtime; + } + + UserStreamRuntime::UserStreamRuntime(const Config &config, + const ApiCredentials &credentials) + : config_(validated_config(config)), + credentials_(validated_credentials(credentials)) + { + websocket_.set_url(config_.clob_user_ws_url); + websocket_.set_ping_interval_ms(config_.ws_ping_interval_ms); + websocket_.set_auto_reconnect(true); + } + + UserStreamRuntime::~UserStreamRuntime() + { + shutdown(); + } + + void UserStreamRuntime::bind_websocket_callbacks() + { + std::weak_ptr weak = shared_from_this(); + websocket_.on_sequenced_message( + [weak](const std::string &message, uint64_t generation) + { + if (auto runtime = weak.lock()) + runtime->handle_message(message, generation); + }); + websocket_.on_stream_gap( + [weak](uint64_t generation) + { + if (auto runtime = weak.lock()) + runtime->handle_stream_gap(generation); + }); + websocket_.on_close( + [weak](uint16_t code, const std::string &reason) + { + if (auto runtime = weak.lock()) + runtime->handle_close(code, reason); + }); + // WebSocketClient fires connect only after replaying the tracked + // subscriptions, so recovery here cannot precede subscription restore. + websocket_.on_connect( + [weak] + { + std::cout << "[WS] Connected to CLOB user stream\n"; + if (auto runtime = weak.lock()) + runtime->handle_connect(); + }); + websocket_.on_disconnect([] + { std::cout << "[WS] Disconnected from CLOB user stream\n"; }); + websocket_.on_error([](const std::string &error) + { std::cerr << "[WS] Error: " << error << '\n'; }); + } + + void UserStreamRuntime::shutdown() + { + if (shutdown_started_.exchange(true)) return; + owner_active_.store(false, std::memory_order_release); + deactivate_stream(); + clear_callbacks(); + websocket_.stop(); + } + + void UserStreamRuntime::subscribe_all_markets() + { + update_subscription([](UserSubscriptionState &state) + { state.all_markets = true; }); + } + + void UserStreamRuntime::subscribe(const std::vector &condition_ids) + { + update_subscription([&condition_ids](UserSubscriptionState &state) + { add_markets(state.markets, condition_ids); }); + } + + void UserStreamRuntime::unsubscribe(const std::vector &condition_ids) + { + update_subscription( + [&condition_ids](UserSubscriptionState &state) + { + auto &markets = state.markets; + markets.erase(std::remove_if(markets.begin(), markets.end(), + [&condition_ids](const auto &market) + { + return std::find(condition_ids.begin(), + condition_ids.end(), + market) != condition_ids.end(); + }), + markets.end()); + }); + } + + void UserStreamRuntime::unsubscribe_all() + { + update_subscription([](UserSubscriptionState &state) + { state = {}; }); + } + + bool UserStreamRuntime::is_subscribed_to_all_markets() const + { + std::lock_guard lock(subscriptions_mutex_); + return subscription_.all_markets; + } + + std::vector UserStreamRuntime::subscribed_markets() const + { + std::lock_guard lock(subscriptions_mutex_); + return subscription_.markets; + } + + std::vector UserStreamRuntime::tracked_messages() const + { + return subscription_.active() + ? std::vector{user_subscription_message( + subscription_, credentials_)} + : std::vector{}; + } + + void UserStreamRuntime::on_order(UserOrderCallback callback) + { + update_callbacks([callback = std::move(callback)](auto &callbacks) mutable + { callbacks.order = std::move(callback); }); + } + + void UserStreamRuntime::on_trade(UserTradeCallback callback) + { + update_callbacks([callback = std::move(callback)](auto &callbacks) mutable + { callbacks.trade = std::move(callback); }); + } + + void UserStreamRuntime::on_stream_gap(UserStreamGapCallback callback) + { + update_callbacks([callback = std::move(callback)](auto &callbacks) mutable + { callbacks.gap = std::move(callback); }); + } + + void UserStreamRuntime::on_stream_recovered(UserStreamRecoveredCallback callback) + { + update_callbacks([callback = std::move(callback)](auto &callbacks) mutable + { callbacks.recovered = std::move(callback); }); + } + + void UserStreamRuntime::on_error(UserStreamErrorCallback callback) + { + update_callbacks([callback = std::move(callback)](auto &callbacks) mutable + { callbacks.error = std::move(callback); }); + } + + void UserStreamRuntime::clear_callbacks() + { + std::lock_guard lock(callback_update_mutex_); + std::shared_ptr empty = + std::make_shared(); + std::atomic_store_explicit(&callbacks_, std::move(empty), + std::memory_order_release); + } + + bool UserStreamRuntime::stream_is_current(uint64_t user_generation, + uint64_t websocket_generation) const + { + return owner_active_.load(std::memory_order_acquire) && + stream_active_.load(std::memory_order_acquire) && + user_generation == stream_generation_.load() && + websocket_generation == websocket_.stream_generation(); + } + + bool UserStreamRuntime::connect() + { + if (!owner_active_.load(std::memory_order_acquire)) return false; + authentication_failed_.store(false, std::memory_order_release); + stream_active_.store(true, std::memory_order_release); + const bool connected = websocket_.connect() && + websocket_.wait_until_connected( + std::chrono::milliseconds( + config_.ws_connect_timeout_ms)); + // A timed-out handshake may still complete later; stop the transport + // so it cannot come up with event delivery already deactivated. + if (!connected) disconnect(); + return connected; + } + + void UserStreamRuntime::disconnect() + { + deactivate_stream(); + websocket_.disconnect(); + } + + bool UserStreamRuntime::is_connected() const + { + return owner_active_.load(std::memory_order_acquire) && + websocket_.is_connected(); + } + + bool UserStreamRuntime::authentication_failed() const + { + return authentication_failed_.load(std::memory_order_acquire); + } + + void UserStreamRuntime::run() + { + if (owner_active_.load(std::memory_order_acquire)) websocket_.run(); + } + + void UserStreamRuntime::stop() + { + deactivate_stream(); + websocket_.stop(); + } + + void UserStreamRuntime::deactivate_stream() + { + if (stream_active_.exchange(false, std::memory_order_acq_rel)) + stream_generation_.fetch_add(1); + } +} diff --git a/src/user_stream_runtime.hpp b/src/user_stream_runtime.hpp new file mode 100644 index 0000000..a40cc1a --- /dev/null +++ b/src/user_stream_runtime.hpp @@ -0,0 +1,146 @@ +#pragma once + +#include "user_stream.hpp" +#include "user_stream_protocol.hpp" +#include "websocket_client.hpp" + +#include +#include +#include +#include +#include + +namespace polymarket::detail +{ + struct UserStreamCallbacks + { + UserOrderCallback order; + UserTradeCallback trade; + UserStreamGapCallback gap; + UserStreamRecoveredCallback recovered; + UserStreamErrorCallback error; + }; + + class UserStreamRuntime final + : public std::enable_shared_from_this + { + public: + static std::shared_ptr create(const Config &config, + const ApiCredentials &credentials); + ~UserStreamRuntime(); + + void shutdown(); + void subscribe_all_markets(); + void subscribe(const std::vector &condition_ids); + void unsubscribe(const std::vector &condition_ids); + void unsubscribe_all(); + bool is_subscribed_to_all_markets() const; + std::vector subscribed_markets() const; + + void on_order(UserOrderCallback callback); + void on_trade(UserTradeCallback callback); + void on_stream_gap(UserStreamGapCallback callback); + void on_stream_recovered(UserStreamRecoveredCallback callback); + void on_error(UserStreamErrorCallback callback); + + bool connect(); + void disconnect(); + bool is_connected() const; + bool authentication_failed() const; + void run(); + void stop(); + uint64_t order_events() const; + uint64_t trade_events() const; + + private: + UserStreamRuntime(const Config &config, const ApiCredentials &credentials); + void bind_websocket_callbacks(); + + void handle_message(const std::string &message, + uint64_t websocket_generation); + void handle_stream_gap(uint64_t websocket_generation); + void handle_connect(); + void handle_close(uint16_t code, const std::string &reason); + void deactivate_stream(); + bool stream_is_current(uint64_t user_generation, + uint64_t websocket_generation) const; + + template + void update_subscription(Update &&update) + { + if (!owner_active_.load(std::memory_order_acquire)) return; + bool recover = false; + bool reconnect = false; + { + std::lock_guard lock(subscriptions_mutex_); + const auto before = subscription_; + update(subscription_); + if (before.all_markets == subscription_.all_markets && + before.markets == subscription_.markets) + return; + websocket_.replace_subscriptions(tracked_messages()); + if (!websocket_.is_connected()) return; + const auto sync = plan_user_subscription_sync( + before, subscription_, credentials_); + reconnect = sync.reconnect; + for (const auto &message : sync.messages) + { + if (!websocket_.send(message)) + { + recover = true; + break; + } + } + } + // Recovery fires the gap callback, so it must run unlocked. + if (reconnect) + websocket_.request_resnapshot( + "user subscription removed; reconnect required"); + else if (recover) + websocket_.request_resnapshot( + "user subscription update send failed; reconnect required"); + } + + std::vector tracked_messages() const; + + template + void update_callbacks(Update &&update) + { + std::lock_guard lock(callback_update_mutex_); + if (!owner_active_.load(std::memory_order_acquire)) return; + auto next = std::make_shared(*callbacks_snapshot()); + update(*next); + std::shared_ptr immutable = std::move(next); + std::atomic_store_explicit(&callbacks_, std::move(immutable), + std::memory_order_release); + } + + std::shared_ptr callbacks_snapshot() const + { + return std::atomic_load_explicit(&callbacks_, std::memory_order_acquire); + } + + void clear_callbacks(); + + Config config_; + ApiCredentials credentials_; + WebSocketClient websocket_; + + mutable std::mutex subscriptions_mutex_; + UserSubscriptionState subscription_; + + mutable std::mutex callback_update_mutex_; + std::shared_ptr callbacks_{ + std::make_shared()}; + + std::atomic owner_active_{true}; + std::atomic stream_active_{false}; + std::atomic shutdown_started_{false}; + std::atomic authentication_failed_{false}; + std::atomic order_events_{0}; + std::atomic trade_events_{0}; + std::atomic stream_generation_{0}; + // WebSocket generation of the gap awaiting recovery; 0 when none. + std::atomic recovery_generation_{0}; + }; +} diff --git a/src/user_stream_runtime_events.cpp b/src/user_stream_runtime_events.cpp new file mode 100644 index 0000000..a860a77 --- /dev/null +++ b/src/user_stream_runtime_events.cpp @@ -0,0 +1,120 @@ +#include "user_stream_runtime.hpp" + +#include +#include +#include +#include +#include + +namespace polymarket::detail +{ + namespace + { + // CLOB closes the user channel with policy violation on bad credentials. + constexpr uint16_t policy_violation_close_code = 1008; + } + + void UserStreamRuntime::handle_stream_gap(uint64_t websocket_generation) + { + if (!owner_active_.load(std::memory_order_acquire) || + !stream_active_.load(std::memory_order_acquire)) + return; + stream_generation_.fetch_add(1); + // Every gap closes or reopens the transport, so the matching connect + // arrives on this websocket generation once subscriptions are replayed. + recovery_generation_.store(websocket_generation, std::memory_order_release); + const auto callbacks = callbacks_snapshot(); + if (callbacks->gap) callbacks->gap(); + } + + void UserStreamRuntime::handle_connect() + { + if (!owner_active_.load(std::memory_order_acquire) || + !stream_active_.load(std::memory_order_acquire)) + return; + auto pending = recovery_generation_.load(std::memory_order_acquire); + // A newer gap means this connection is already stale; its own + // reconnect will report recovery instead. + if (pending == 0 || pending != websocket_.stream_generation() || + !recovery_generation_.compare_exchange_strong(pending, 0)) + return; + const auto callbacks = callbacks_snapshot(); + if (callbacks->recovered) callbacks->recovered(); + } + + void UserStreamRuntime::handle_close(uint16_t code, const std::string &reason) + { + if (code != policy_violation_close_code || + !owner_active_.load(std::memory_order_acquire) || + !stream_active_.load(std::memory_order_acquire)) + return; + // Reconnecting with rejected credentials would loop forever. stop() + // rather than disconnect() so a blocked or later run() returns too; + // connect() clears the stop for a retry. + authentication_failed_.store(true, std::memory_order_release); + deactivate_stream(); + websocket_.stop(); + const auto error = "user stream rejected by server (" + std::to_string(code) + + "): " + reason; + std::cerr << "[WS] " << error << '\n'; + const auto callbacks = callbacks_snapshot(); + if (callbacks->error) callbacks->error(error); + } + + void UserStreamRuntime::handle_message(const std::string &message, + uint64_t websocket_generation) + { + const auto generation = stream_generation_.load(); + if (!stream_is_current(generation, websocket_generation) || + message.empty() || message == "PONG") + return; + + std::vector events; + try + { + events = parse_user_events(message); + } + catch (const std::exception &error) + { + if (!owner_active_.load(std::memory_order_acquire)) return; + std::cerr << "[WS] User stream parse error: " << error.what() << '\n'; + websocket_.request_resnapshot( + "invalid user data; reconnect required"); + return; + } + + // Callback exceptions propagate to WebSocketClient, which counts them + // and requests a reconnect. + for (const auto &event : events) + { + if (!stream_is_current(generation, websocket_generation)) return; + const auto callbacks = callbacks_snapshot(); + std::visit( + [this, &callbacks](const auto &payload) + { + using Event = std::decay_t; + if constexpr (std::is_same_v) + { + order_events_++; + if (callbacks->order) callbacks->order(payload); + } + else + { + trade_events_++; + if (callbacks->trade) callbacks->trade(payload); + } + }, + event); + } + } + + uint64_t UserStreamRuntime::order_events() const + { + return order_events_.load(); + } + + uint64_t UserStreamRuntime::trade_events() const + { + return trade_events_.load(); + } +} diff --git a/src/user_stream_subscription.cpp b/src/user_stream_subscription.cpp new file mode 100644 index 0000000..c80d549 --- /dev/null +++ b/src/user_stream_subscription.cpp @@ -0,0 +1,87 @@ +#include "user_stream_protocol.hpp" + +#include + +#include + +using json = nlohmann::json; + +namespace polymarket::detail +{ + namespace + { + std::vector difference(const std::vector &left, + const std::vector &right) + { + std::vector result; + for (const auto &value : left) + { + if (std::find(right.begin(), right.end(), value) == right.end()) + result.push_back(value); + } + return result; + } + } + + std::string user_subscription_message(const UserSubscriptionState &state, + const ApiCredentials &credentials) + { + json message{{"auth", {{"apiKey", credentials.api_key}, + {"secret", credentials.api_secret}, + {"passphrase", credentials.api_passphrase}}}, + {"type", "user"}}; + if (!state.all_markets) + { + message["markets"] = state.markets; + } + return message.dump(); + } + + std::string user_subscription_update_message(const std::vector &markets, + bool subscribe) + { + return json{{"markets", markets}, + {"operation", subscribe ? "subscribe" : "unsubscribe"}} + .dump(); + } + + UserSubscriptionSync plan_user_subscription_sync(const UserSubscriptionState &before, + const UserSubscriptionState &after, + const ApiCredentials &credentials) + { + UserSubscriptionSync sync; + if (!after.active()) + { + // An empty market list means "all markets" to the server, so the + // last unsubscribe must drop the authenticated session instead. + sync.reconnect = before.active(); + return sync; + } + if (!before.active()) + { + sync.messages.push_back(user_subscription_message(after, credentials)); + return sync; + } + if (after.all_markets) + { + if (!before.all_markets) + sync.messages.push_back( + user_subscription_update_message(before.markets, false)); + return sync; + } + if (before.all_markets) + { + sync.messages.push_back( + user_subscription_update_message(after.markets, true)); + return sync; + } + + const auto added = difference(after.markets, before.markets); + const auto removed = difference(before.markets, after.markets); + if (!added.empty()) + sync.messages.push_back(user_subscription_update_message(added, true)); + if (!removed.empty()) + sync.messages.push_back(user_subscription_update_message(removed, false)); + return sync; + } +} diff --git a/src/websocket_client.cpp b/src/websocket_client.cpp index 8090007..a67ed4a 100644 --- a/src/websocket_client.cpp +++ b/src/websocket_client.cpp @@ -74,6 +74,12 @@ namespace polymarket state->on_disconnect(std::move(callback)); } + void WebSocketClient::on_close(OnCloseCallback callback) + { + auto state = state_; + state->on_close(std::move(callback)); + } + void WebSocketClient::on_error(OnErrorCallback callback) { auto state = state_; diff --git a/src/websocket_client_callbacks.cpp b/src/websocket_client_callbacks.cpp index 2d8ac44..ab937ca 100644 --- a/src/websocket_client_callbacks.cpp +++ b/src/websocket_client_callbacks.cpp @@ -35,6 +35,12 @@ namespace polymarket::detail { callbacks.disconnect = std::move(callback); }); } + void WebSocketClientState::on_close(OnCloseCallback callback) + { + update_callbacks([callback = std::move(callback)](auto &callbacks) mutable + { callbacks.close = std::move(callback); }); + } + void WebSocketClientState::on_error(OnErrorCallback callback) { update_callbacks([callback = std::move(callback)](auto &callbacks) mutable diff --git a/src/websocket_client_state.hpp b/src/websocket_client_state.hpp index 49fb2e6..1b77058 100644 --- a/src/websocket_client_state.hpp +++ b/src/websocket_client_state.hpp @@ -23,6 +23,7 @@ namespace polymarket::detail OnTypedMessageCallback typed_message; OnConnectCallback connect; OnDisconnectCallback disconnect; + OnCloseCallback close; OnErrorCallback error; OnStreamGapCallback stream_gap; }; @@ -45,6 +46,7 @@ namespace polymarket::detail void on_typed_message(OnTypedMessageCallback callback); void on_connect(OnConnectCallback callback); void on_disconnect(OnDisconnectCallback callback); + void on_close(OnCloseCallback callback); void on_error(OnErrorCallback callback); void on_stream_gap(OnStreamGapCallback callback); @@ -78,7 +80,7 @@ namespace polymarket::detail void install_transport_callback(); void handle_transport_message(const ix::WebSocketMessagePtr &message); void handle_open(); - void handle_close(); + void handle_close(const ix::WebSocketMessagePtr &message); void handle_error(const ix::WebSocketMessagePtr &message); void start_message_worker(); diff --git a/src/websocket_client_transport.cpp b/src/websocket_client_transport.cpp index 688b3c3..b604711 100644 --- a/src/websocket_client_transport.cpp +++ b/src/websocket_client_transport.cpp @@ -10,7 +10,10 @@ namespace polymarket::detail } WebSocketClientState::WebSocketClientState() - : message_queue_(std::make_unique(options_.message_queue_limit)) + // options_ is declared after message_queue_, so read the default limit + // from a fresh options value rather than the not-yet-constructed member. + : message_queue_(std::make_unique( + WebSocketOptions{}.message_queue_limit)) { apply_options_locked(); } @@ -71,6 +74,7 @@ namespace polymarket::detail { start_transport_worker(); install_transport_callback(); + if (options_.reconnect_enabled) ws_.enableAutomaticReconnection(); state_.store(WsState::CONNECTING); resnapshot_pending_.store(false); should_stop_.store(false); @@ -114,7 +118,7 @@ namespace polymarket::detail handle_open(); break; case ix::WebSocketMessageType::Close: - handle_close(); + handle_close(message); break; case ix::WebSocketMessageType::Error: handle_error(message); @@ -150,7 +154,7 @@ namespace polymarket::detail invoke_user_callback("connect", callbacks->connect); } - void WebSocketClientState::handle_close() + void WebSocketClientState::handle_close(const ix::WebSocketMessagePtr &message) { if (!transport_stop_pending_.load()) mark_stream_gap(); { @@ -167,6 +171,9 @@ namespace polymarket::detail if (transport_stop_pending_.load()) return; auto callbacks = callbacks_snapshot(); + invoke_user_callback("close", callbacks->close, + message->closeInfo.code, message->closeInfo.reason); + if (transport_stop_pending_.load()) return; invoke_user_callback("disconnect", callbacks->disconnect); } @@ -208,7 +215,9 @@ namespace polymarket::detail std::lock_guard lock(lifecycle_mutex_); state_.store(WsState::CLOSING); resnapshot_pending_.store(false); - if (!options_.reconnect_enabled) ws_.disableAutomaticReconnection(); + // The transport stop may be deferred to the worker thread; block + // ix from reconnecting in the meantime. connect() re-enables it. + ws_.disableAutomaticReconnection(); request_message_worker_stop(); deferred = request_transport_stop(); } diff --git a/src/websocket_user_data.cpp b/src/websocket_user_data.cpp new file mode 100644 index 0000000..dbc50a1 --- /dev/null +++ b/src/websocket_user_data.cpp @@ -0,0 +1,233 @@ +#include "user_stream_protocol.hpp" + +#include + +#include +#include +#include +#include +#include + +using json = nlohmann::json; + +namespace polymarket::detail +{ + namespace + { + std::string text_field(const json &item, const char *field, + const char *context, bool required) + { + const auto value = item.find(field); + if (value == item.end() || value->is_null()) + { + if (required) + { + throw std::invalid_argument(std::string(context) + " requires " + + field); + } + return {}; + } + if (value->is_string()) return value->get(); + if (value->is_number()) return value->dump(); + throw std::invalid_argument(std::string(context) + " " + field + + " must be a string or number"); + } + + std::string required_text(const json &item, const char *field, + const char *context) + { + auto value = text_field(item, field, context, true); + if (value.empty()) + { + throw std::invalid_argument(std::string(context) + " " + field + + " must not be empty"); + } + return value; + } + + std::string optional_text(const json &item, const char *field, + const char *context) + { + return text_field(item, field, context, false); + } + + std::string asset_id(const json &item, const char *context) + { + if (item.contains("asset_id")) return required_text(item, "asset_id", context); + return required_text(item, "token_id", context); + } + + std::string uppercase(std::string value) + { + std::transform(value.begin(), value.end(), value.begin(), + [](unsigned char c) + { return static_cast(std::toupper(c)); }); + return value; + } + + std::string side(const json &item, const char *field, const char *context, + bool required) + { + auto value = uppercase(text_field(item, field, context, required)); + if (required && value != "BUY" && value != "SELL") + { + throw std::invalid_argument(std::string(context) + " " + field + + " must be BUY or SELL"); + } + return value; + } + + std::vector string_array(const json &item, const char *field, + const char *context) + { + const auto value = item.find(field); + if (value == item.end() || value->is_null()) return {}; + if (!value->is_array()) + { + throw std::invalid_argument(std::string(context) + " " + field + + " must be an array or null"); + } + std::vector result; + result.reserve(value->size()); + for (const auto &entry : *value) + { + if (!entry.is_string()) + { + throw std::invalid_argument(std::string(context) + " " + field + + " entries must be strings"); + } + result.push_back(entry.get()); + } + return result; + } + + std::optional bucket_index(const json &item) + { + const auto value = item.find("bucket_index"); + if (value == item.end() || value->is_null()) return std::nullopt; + if (value->is_number_unsigned()) + { + const auto index = value->get(); + if (index <= std::numeric_limits::max()) + return static_cast(index); + } + else if (value->is_number_integer()) + { + const auto index = value->get(); + if (index >= 0 && index <= std::numeric_limits::max()) + return static_cast(index); + } + throw std::invalid_argument( + "user trade bucket_index must be a uint32 integer"); + } + + MakerOrder parse_maker_order(const json &item) + { + constexpr const char *context = "user trade maker order"; + if (!item.is_object()) + throw std::invalid_argument("user trade maker order must be an object"); + MakerOrder order; + order.order_id = required_text(item, "order_id", context); + order.owner = required_text(item, "owner", context); + order.maker_address = optional_text(item, "maker_address", context); + order.matched_amount = required_text(item, "matched_amount", context); + order.price = required_text(item, "price", context); + order.fee_rate_bps = optional_text(item, "fee_rate_bps", context); + order.asset_id = asset_id(item, context); + order.outcome = optional_text(item, "outcome", context); + order.side = side(item, "side", context, true); + return order; + } + + std::vector maker_orders(const json &item) + { + const auto orders = item.find("maker_orders"); + if (orders == item.end() || orders->is_null()) return {}; + if (!orders->is_array()) + { + throw std::invalid_argument( + "user trade maker_orders must be an array or null"); + } + std::vector result; + result.reserve(orders->size()); + for (const auto &order : *orders) + result.push_back(parse_maker_order(order)); + return result; + } + + UserOrderEvent parse_order(const json &item) + { + constexpr const char *context = "user order"; + UserOrderEvent order; + order.id = required_text(item, "id", context); + order.owner = required_text(item, "owner", context); + order.market = required_text(item, "market", context); + order.asset_id = asset_id(item, context); + order.side = side(item, "side", context, true); + order.original_size = required_text(item, "original_size", context); + order.size_matched = required_text(item, "size_matched", context); + order.price = required_text(item, "price", context); + order.type = uppercase(required_text(item, "type", context)); + order.status = optional_text(item, "status", context); + order.order_type = optional_text(item, "order_type", context); + order.maker_address = optional_text(item, "maker_address", context); + order.order_owner = optional_text(item, "order_owner", context); + order.associate_trades = string_array(item, "associate_trades", context); + order.outcome = optional_text(item, "outcome", context); + order.created_at = optional_text(item, "created_at", context); + order.expiration = optional_text(item, "expiration", context); + order.timestamp = optional_text(item, "timestamp", context); + return order; + } + + UserTradeEvent parse_trade(const json &item) + { + constexpr const char *context = "user trade"; + UserTradeEvent trade; + trade.id = required_text(item, "id", context); + trade.taker_order_id = required_text(item, "taker_order_id", context); + trade.market = required_text(item, "market", context); + trade.asset_id = asset_id(item, context); + trade.side = side(item, "side", context, true); + trade.size = required_text(item, "size", context); + trade.price = required_text(item, "price", context); + trade.status = required_text(item, "status", context); + trade.owner = required_text(item, "owner", context); + trade.fee_rate_bps = optional_text(item, "fee_rate_bps", context); + trade.match_time = item.contains("match_time") + ? optional_text(item, "match_time", context) + : optional_text(item, "matchtime", context); + trade.last_update = optional_text(item, "last_update", context); + trade.timestamp = optional_text(item, "timestamp", context); + trade.trade_owner = optional_text(item, "trade_owner", context); + trade.maker_address = optional_text(item, "maker_address", context); + trade.transaction_hash = optional_text(item, "transaction_hash", context); + trade.bucket_index = bucket_index(item); + trade.maker_orders = maker_orders(item); + trade.trader_side = side(item, "trader_side", context, false); + trade.outcome = optional_text(item, "outcome", context); + return trade; + } + } + + std::vector parse_user_events(const std::string &message) + { + const auto parsed = json::parse(message); + const auto items = parsed.is_array() ? parsed : json::array({parsed}); + std::vector events; + events.reserve(items.size()); + + for (const auto &item : items) + { + if (!item.is_object()) + throw std::invalid_argument("user event must be an object"); + + const std::string kind = item.value("event_type", ""); + if (kind == "order") + events.emplace_back(parse_order(item)); + else if (kind == "trade") + events.emplace_back(parse_trade(item)); + } + return events; + } +} diff --git a/tests/test_user_stream.cpp b/tests/test_user_stream.cpp new file mode 100644 index 0000000..09807ba --- /dev/null +++ b/tests/test_user_stream.cpp @@ -0,0 +1,171 @@ +#include "user_stream_test_support.hpp" +#include "user_stream_protocol.hpp" + +#include + +#include +#include +#include +#include +#include + +using namespace polymarket; +using namespace std::chrono_literals; +using namespace user_stream_test; +using json = nlohmann::json; + +namespace user_stream_test +{ + bool protocol_tests() + { + const auto credentials = test_credentials(); + + const auto filtered = json::parse(detail::user_subscription_message( + {false, {"cond-1", "cond-2"}}, credentials)); + const auto all = json::parse( + detail::user_subscription_message({true, {"ignored"}}, credentials)); + if (!check(filtered["type"] == "user", "subscription type is user") || + !check(filtered["auth"]["apiKey"] == "key" && + filtered["auth"]["secret"] == "secret" && + filtered["auth"]["passphrase"] == "passphrase", + "subscription carries API credentials") || + !check(filtered["markets"] == json::array({"cond-1", "cond-2"}), + "filtered subscription lists markets") || + !check(!all.contains("markets"), + "all-markets subscription omits the market filter")) + return false; + + const auto update = json::parse( + detail::user_subscription_update_message({"cond-3"}, false)); + if (!check(update["operation"] == "unsubscribe" && + update["markets"] == json::array({"cond-3"}), + "update message shape")) + return false; + + using detail::plan_user_subscription_sync; + const detail::UserSubscriptionState empty; + const detail::UserSubscriptionState one{false, {"a"}}; + const detail::UserSubscriptionState two{false, {"a", "b"}}; + const detail::UserSubscriptionState other{false, {"b"}}; + const detail::UserSubscriptionState everything{true, {}}; + + const auto first = plan_user_subscription_sync(empty, one, credentials); + const auto added = plan_user_subscription_sync(one, two, credentials); + const auto swapped = plan_user_subscription_sync(one, other, credentials); + const auto removed_last = plan_user_subscription_sync(one, empty, credentials); + const auto widened = plan_user_subscription_sync(two, everything, credentials); + const auto narrowed = plan_user_subscription_sync(everything, one, credentials); + const auto idle = plan_user_subscription_sync(empty, empty, credentials); + if (!check(first.messages.size() == 1 && !first.reconnect && + json::parse(first.messages[0])["type"] == "user", + "first subscription sends the authenticated message") || + !check(added.messages == + std::vector{ + detail::user_subscription_update_message({"b"}, true)}, + "added market sends a subscribe update") || + !check(swapped.messages == + std::vector{ + detail::user_subscription_update_message({"b"}, true), + detail::user_subscription_update_message({"a"}, false)}, + "swapped market subscribes before unsubscribing") || + !check(removed_last.messages.empty() && removed_last.reconnect, + "removing the last market reconnects instead of widening") || + !check(widened.messages == + std::vector{ + detail::user_subscription_update_message({"a", "b"}, false)}, + "widening to all markets clears the market filter") || + !check(narrowed.messages == + std::vector{ + detail::user_subscription_update_message({"a"}, true)}, + "narrowing from all markets subscribes the filter") || + !check(idle.messages.empty() && !idle.reconnect, + "inactive to inactive is a no-op")) + return false; + + const auto orders = detail::parse_user_events(order_message); + if (!check(orders.size() == 1 && + std::holds_alternative(orders[0]), + "order event parses")) + return false; + const auto &order = std::get(orders[0]); + if (!check(order.id == "0xorder" && order.market == "cond-1" && + order.asset_id == "yes" && order.side == "BUY" && + order.type == "PLACEMENT" && order.status == "LIVE" && + order.associate_trades.empty() && + order.timestamp == "1782753357257", + "order fields are normalized")) + return false; + + const auto trades = detail::parse_user_events(trade_message); + if (!check(trades.size() == 1 && + std::holds_alternative(trades[0]), + "trade event array parses")) + return false; + const auto &trade = std::get(trades[0]); + if (!check(trade.status == "MATCHED" && trade.match_time == "1782753357" && + trade.timestamp == "1782753357257" && + trade.trader_side == "MAKER" && trade.bucket_index == 2u && + trade.maker_orders.size() == 1 && + trade.maker_orders[0].asset_id == "yes" && + trade.maker_orders[0].side == "BUY" && + trade.maker_orders[0].maker_address.empty(), + "trade fields and maker orders are normalized")) + return false; + + // Only the CLOB wire shape is supported; nothing produces envelopes. + if (!check(detail::parse_user_events( + R"({"topic":"user","type":"order","payload":{"id":"o","owner":"k","market":"m","asset_id":"t","side":"SELL","original_size":"1","size_matched":"1","price":"0.5","type":"CANCELLATION"}})") + .empty(), + "topic/payload envelopes are not parsed")) + return false; + + if (!check(detail::parse_user_events( + R"({"event_type":"heartbeat","id":"x"})") + .empty(), + "unknown event types are ignored")) + return false; + + for (const auto *invalid : { + R"({"event_type":"order","owner":"k","market":"m","asset_id":"t","side":"BUY","original_size":"1","size_matched":"0","price":"0.5","type":"PLACEMENT"})", + R"({"event_type":"order","id":"o","owner":"k","market":"m","asset_id":"t","side":"HOLD","original_size":"1","size_matched":"0","price":"0.5","type":"PLACEMENT"})", + R"({"event_type":"trade","id":"t","taker_order_id":"o","market":"m","asset_id":"t","side":"BUY","size":"1","price":"0.5","status":"MATCHED","owner":"k","bucket_index":-1})", + R"({"event_type":"trade","id":"t","taker_order_id":"o","market":"m","asset_id":"t","side":"BUY","size":"1","price":"0.5","status":"MATCHED","owner":"k","maker_orders":{}})", + R"(["not-an-object"])"}) + { + bool rejected = false; + try + { + (void)detail::parse_user_events(invalid); + } + catch (const std::exception &) + { + rejected = true; + } + if (!check(rejected, "malformed user events are rejected")) + return false; + } + + bool rejected_credentials = false; + try + { + UserStream stream(Config{}, ApiCredentials{"key", "", "passphrase"}); + } + catch (const std::invalid_argument &) + { + rejected_credentials = true; + } + return check(rejected_credentials, "incomplete credentials are rejected"); + } +} + +int main() +{ + if (!protocol_tests()) return 1; + if (!stream_tests()) return 1; + if (!recovery_ordering_tests()) return 1; + if (!delayed_handshake_tests()) return 1; + if (!authentication_failure_tests(true)) return 1; + if (!authentication_failure_tests(false)) return 1; + std::cout << "user stream tests passed\n"; + return 0; +} diff --git a/tests/test_user_stream_connection.cpp b/tests/test_user_stream_connection.cpp new file mode 100644 index 0000000..18b6ba1 --- /dev/null +++ b/tests/test_user_stream_connection.cpp @@ -0,0 +1,201 @@ +#include "user_stream_test_support.hpp" +#include "user_stream_protocol.hpp" +#include "websocket_test_server.hpp" + +#include + +#include +#include +#include +#include +#include + +using namespace polymarket; +using namespace std::chrono_literals; +using namespace user_stream_test; +using json = nlohmann::json; + +namespace user_stream_test +{ + bool stream_tests() + { + websocket_test::LocalWebSocketServer server; + Config config; + config.clob_user_ws_url = server.url(); + config.ws_connect_timeout_ms = 1'000; + const auto credentials = test_credentials(); + + std::atomic gaps{0}; + std::mutex events_mutex; + std::vector orders; + std::vector trades; + + bool passed = false; + { + UserStream stream(config, credentials); + stream.on_stream_gap([&gaps] + { ++gaps; }); + stream.on_order([&](const UserOrderEvent &order) + { + std::lock_guard lock(events_mutex); + orders.push_back(order); + }); + stream.on_trade([&](const UserTradeEvent &trade) + { + std::lock_guard lock(events_mutex); + trades.push_back(trade); + }); + stream.subscribe("cond-1"); + + const auto initial = detail::user_subscription_message( + {false, {"cond-1"}}, credentials); + passed = + check(stream.connect(), "user stream connects") && + check(server.wait_for_message_count(initial, 1, 2s), + "connect sends the tracked authenticated subscription") && + check(wait_until([&gaps] + { return gaps.load() >= 1; }), + "connect reports a reconciliation gap") && + check(server.send_to_clients(order_message) && + server.send_to_clients(trade_message), + "server sends user events") && + check(wait_until([&] + { + std::lock_guard lock(events_mutex); + return orders.size() == 1 && trades.size() == 1; + }), + "order and trade callbacks fire") && + check(stream.order_events() == 1 && stream.trade_events() == 1, + "event counters advance"); + if (passed) + { + stream.subscribe(std::vector{"cond-2", "cond-1"}); + stream.unsubscribe("cond-1"); + passed = + check(server.wait_for_message_count( + detail::user_subscription_update_message({"cond-2"}, true), + 1, 2s), + "subscribe sends an incremental update") && + check(server.wait_for_message_count( + detail::user_subscription_update_message({"cond-1"}, false), + 1, 2s), + "unsubscribe sends an incremental update") && + check(stream.subscribed_markets() == + std::vector{"cond-2"}, + "subscribed markets are tracked"); + } + if (passed) + { + const auto gaps_before = gaps.load(); + server.close_clients(); + passed = + check(server.wait_for_message_count( + detail::user_subscription_message({false, {"cond-2"}}, + credentials), + 1, 5s), + "reconnect replays the current subscription") && + check(wait_until([&] + { return gaps.load() > gaps_before; }), + "reconnect reports a reconciliation gap"); + } + if (passed) + { + const auto gaps_before = gaps.load(); + passed = + check(wait_until([&stream] + { return stream.is_connected(); }, 5s), + "stream is connected before invalid payload") && + check(server.send_to_clients( + R"({"event_type":"order","id":"missing-fields"})"), + "server sends invalid user event") && + check(wait_until([&] + { return gaps.load() > gaps_before; }), + "invalid payload forces a reconcile gap"); + } + if (passed) + { + passed = check(wait_until([&stream] + { return stream.is_connected(); }, 5s), + "stream reconnects after invalid payload"); + stream.subscribe_all_markets(); + passed = passed && + check(server.wait_for_message_count( + detail::user_subscription_update_message({"cond-2"}, false), + 1, 2s), + "all-markets subscription clears the market filter") && + check(stream.is_subscribed_to_all_markets(), + "all-markets mode is tracked"); + } + stream.stop(); + } + return passed; + } + + bool recovery_ordering_tests() + { + websocket_test::LocalWebSocketServer server; + Config config; + config.clob_user_ws_url = server.url(); + config.ws_connect_timeout_ms = 3'000; + const auto credentials = test_credentials(); + const auto subscription = detail::user_subscription_message( + {false, {"cond-1"}}, credentials); + + std::mutex log_mutex; + std::vector log; + std::atomic recoveries{0}; + bool passed = false; + { + UserStream stream(config, credentials); + stream.on_stream_gap([&] + { + std::lock_guard lock(log_mutex); + log.push_back("gap"); + }); + stream.on_stream_recovered( + [&] + { + // The transport thread blocks here, so this only succeeds + // if the subscription for this connection was sent first. + const bool subscribed = server.wait_for_message_count( + subscription, recoveries.load() + 1, 1s); + { + std::lock_guard lock(log_mutex); + log.push_back(subscribed ? "recovered" : "recovered-early"); + } + ++recoveries; + }); + stream.subscribe("cond-1"); + + passed = check(stream.connect(), "user stream connects") && + check(wait_until([&] + { return recoveries.load() == 1; }), + "connect reports recovery"); + if (passed) + { + server.close_clients(); + passed = check(wait_until([&] + { return recoveries.load() == 2; }, 5s), + "reconnect reports recovery"); + } + stream.stop(); + } + + std::lock_guard lock(log_mutex); + if (!passed) return false; + const auto recovered_entries = std::count(log.begin(), log.end(), + std::string("recovered")); + bool gap_before_each_recovery = !log.empty() && log.front() == "gap"; + for (std::size_t i = 1; i < log.size(); ++i) + { + if (log[i] == "recovered" && log[i - 1] != "gap") + gap_before_each_recovery = false; + } + return check(recovered_entries == 2 && + std::find(log.begin(), log.end(), "recovered-early") == + log.end(), + "recovery fires only after the subscription is restored") && + check(gap_before_each_recovery, + "gap notification precedes each recovery"); + } +} diff --git a/tests/test_user_stream_failures.cpp b/tests/test_user_stream_failures.cpp new file mode 100644 index 0000000..e661030 --- /dev/null +++ b/tests/test_user_stream_failures.cpp @@ -0,0 +1,140 @@ +#include "user_stream_test_support.hpp" +#include "user_stream_protocol.hpp" +#include "websocket_test_server.hpp" + +#include + +#include +#include +#include +#include +#include + +using namespace polymarket; +using namespace std::chrono_literals; +using namespace user_stream_test; +using json = nlohmann::json; + +namespace user_stream_test +{ + bool delayed_handshake_tests() + { + // The first handshake completes well after connect() gives up. + websocket_test::LocalWebSocketServer server(1'500ms); + Config config; + config.clob_user_ws_url = server.url(); + config.ws_connect_timeout_ms = 300; + const auto credentials = test_credentials(); + + std::atomic orders{0}; + bool passed = false; + { + UserStream stream(config, credentials); + stream.on_order([&orders](const UserOrderEvent &) + { ++orders; }); + stream.subscribe("cond-1"); + + passed = + check(!stream.connect(), "connect times out on a delayed handshake") && + check(!wait_until([&stream] + { return stream.is_connected(); }, 2'500ms), + "timed-out connect stops the transport"); + if (passed) + { + passed = + check(stream.connect(), "retry connects after a timed-out attempt") && + check(server.wait_for_message_count( + detail::user_subscription_message({false, {"cond-1"}}, + credentials), + 1, 2s), + "retry sends the authenticated subscription") && + check(server.send_to_clients(order_message), + "server sends an order after retry") && + check(wait_until([&orders] + { return orders.load() == 1; }), + "retried stream delivers events"); + } + stream.stop(); + } + return passed; + } + + // Rejection is terminal: run() must return whether it was already + // blocking when the server closed with 1008 or is entered afterwards. + bool authentication_failure_tests(bool reject_during_run) + { + websocket_test::LocalWebSocketServer server; + Config config; + config.clob_user_ws_url = server.url(); + config.ws_connect_timeout_ms = 1'000; + const auto credentials = test_credentials(); + + std::mutex errors_mutex; + std::vector errors; + std::atomic run_returned{false}; + std::thread runner; + bool passed = false; + { + UserStream stream(config, credentials); + stream.on_error([&](const std::string &error) + { + std::lock_guard lock(errors_mutex); + errors.push_back(error); + }); + stream.subscribe_all_markets(); + const auto start_run = [&] + { + runner = std::thread([&] + { + stream.run(); + run_returned = true; + }); + }; + passed = + check(stream.connect(), "user stream connects before auth reply") && + check(server.wait_for_message_count( + detail::user_subscription_message({true, {}}, credentials), + 1, 2s), + "subscription is sent before rejection"); + if (passed && reject_during_run) + { + start_run(); + // Let run() enter its blocking loop before the rejection. + passed = check(!wait_until([&] + { return run_returned.load(); }, 300ms), + "run() blocks while the stream is live"); + } + if (passed) + { + server.close_clients(1008, "authentication failed"); + passed = + check(wait_until([&] + { + std::lock_guard lock(errors_mutex); + return errors.size() == 1; + }), + "rejection is reported through on_error") && + check(errors[0].find("authentication failed") != std::string::npos, + "rejection error carries the server reason") && + check(stream.authentication_failed(), + "authentication failure is observable"); + } + if (passed && !reject_during_run) start_run(); + if (passed) + { + passed = + check(wait_until([&] + { return run_returned.load(); }), + reject_during_run ? "rejection ends a blocked run()" + : "run() returns after an earlier rejection") && + check(!server.wait_for_connections(2, 1s), + "rejected stream does not reconnect") && + check(!stream.is_connected(), "rejected stream is disconnected"); + } + // Unblocks run() if the rejection failed to end it. + stream.stop(); + if (runner.joinable()) runner.join(); + } + return passed; + } +} diff --git a/tests/user_stream_test_support.hpp b/tests/user_stream_test_support.hpp new file mode 100644 index 0000000..720bec9 --- /dev/null +++ b/tests/user_stream_test_support.hpp @@ -0,0 +1,49 @@ +#pragma once + +#include "user_stream.hpp" + +#include +#include +#include + +namespace user_stream_test +{ + inline bool check(bool value, const char *name) + { + if (!value) + { + std::cerr << "failed: " << name << '\n'; + } + return value; + } + + template + bool wait_until(Predicate predicate, + std::chrono::milliseconds timeout = std::chrono::seconds(2)) + { + const auto deadline = std::chrono::steady_clock::now() + timeout; + while (std::chrono::steady_clock::now() < deadline) + { + if (predicate()) return true; + std::this_thread::sleep_for(std::chrono::milliseconds(2)); + } + return predicate(); + } + + inline polymarket::ApiCredentials test_credentials() + { + return {"key", "secret", "passphrase"}; + } + + inline constexpr const char *order_message = + R"({"event_type":"order","id":"0xorder","owner":"owner-key","market":"cond-1","asset_id":"yes","side":"buy","original_size":"10","size_matched":"0","price":"0.42","type":"PLACEMENT","status":"LIVE","order_type":"GTC","associate_trades":null,"outcome":"Yes","created_at":"1782753357","expiration":"0","timestamp":"1782753357257"})"; + + inline constexpr const char *trade_message = + R"([{"event_type":"trade","id":"trade-1","taker_order_id":"0xtaker","market":"cond-1","asset_id":"yes","side":"SELL","size":"5","price":"0.42","status":"MATCHED","owner":"owner-key","fee_rate_bps":"0","matchtime":"1782753357","last_update":"1782753358","timestamp":1782753357257,"trader_side":"maker","bucket_index":2,"maker_orders":[{"order_id":"0xorder","owner":"owner-key","matched_amount":"5","price":"0.42","token_id":"yes","side":"buy","outcome":"Yes"}]}])"; + + bool protocol_tests(); + bool stream_tests(); + bool recovery_ordering_tests(); + bool delayed_handshake_tests(); + bool authentication_failure_tests(bool reject_during_run); +} diff --git a/tests/websocket_test_server.hpp b/tests/websocket_test_server.hpp index d4782fb..93782f6 100644 --- a/tests/websocket_test_server.hpp +++ b/tests/websocket_test_server.hpp @@ -19,9 +19,26 @@ namespace websocket_test class LocalWebSocketServer { public: - LocalWebSocketServer() + // first_handshake_delay stalls the first accepted connection before + // the WebSocket handshake, simulating a slow upgrade response. + explicit LocalWebSocketServer( + std::chrono::milliseconds first_handshake_delay = std::chrono::milliseconds(0)) : port_(ix::getFreePort()), server_(port_, "127.0.0.1") { + if (first_handshake_delay.count() > 0) + { + server_.setConnectionStateFactory( + [this, first_handshake_delay] + { + if (!handshake_delayed_) + { + handshake_delayed_ = true; + std::this_thread::sleep_for(first_handshake_delay); + } + return ix::ConnectionState::createConnectionState(); + }); + } + server_.setOnClientMessageCallback( [this](std::shared_ptr, ix::WebSocket &, @@ -81,6 +98,14 @@ namespace websocket_test } } + void close_clients(uint16_t code, const std::string &reason) + { + for (const auto &client : server_.getClients()) + { + client->close(code, reason); + } + } + bool wait_for_connections(std::size_t count, std::chrono::milliseconds timeout) { @@ -122,6 +147,7 @@ namespace websocket_test std::mutex mutex_; std::condition_variable cv_; std::size_t connections_{0}; + bool handshake_delayed_{false}; // accept thread only std::vector received_; }; }