From 6018c7feb3af29621878826598bee6014495e2cf Mon Sep 17 00:00:00 2001 From: yhirose Date: Fri, 7 Aug 2026 18:02:29 -0400 Subject: [PATCH] Return ws::Result from WebSocketClient::connect() instead of bool Issue #2531 asked for connect() to expose error detail the way ClientImpl/SSLClient do via Result, instead of collapsing every failure into a bare bool. The groundwork (detail::ClientTlsSessionError) was already laid during the WebSocketClient/SSLClient dedup but left unwired. - Add httplib::ws::Result: explicit operator bool(), error(), and flattened upgrade-response accessors (status(), headers(), get_header_value(), has_header()); ssl_error()/ssl_backend_error() on SSL builds. - Add Error::WebSocketHandshake for upgrade-validation failures (non-101 status, bad Sec-WebSocket-Accept, bad Upgrade/Connection headers). - Extract detail::parse_status_line from ClientImpl::read_response_line and reuse it in read_websocket_upgrade_response, replacing the previous "HTTP/1.1 101" substring match with a proper parse. Non-101 responses now surface their status and headers instead of being read and discarded. - Wire WebSocketClient::create_stream() to capture ClientTlsSessionError so TLS failures (SSLServerVerification, SSLServerHostnameVerification, ...) reach the caller with backend error codes. - Update tests and README-websocket.md accordingly. This is a source-breaking change for callers that assign the result to bool (e.g. bool ok = cli.connect();); if (cli.connect()) and gtest's ASSERT_TRUE/EXPECT_FALSE(...) macros are unaffected since operator bool still participates in contextual conversion. --- README-websocket.md | 31 ++++++- httplib.h | 220 +++++++++++++++++++++++++++++++++----------- test/test.cc | 40 +++++++- 3 files changed, 232 insertions(+), 59 deletions(-) diff --git a/README-websocket.md b/README-websocket.md index 3c6ac31..d16239d 100644 --- a/README-websocket.md +++ b/README-websocket.md @@ -151,8 +151,15 @@ explicit WebSocketClient(const std::string &scheme_host_port_path, // Check if the URL was parsed successfully bool is_valid() const; -// Connect (performs HTTP upgrade handshake) -bool connect(); +// Connect (performs HTTP upgrade handshake). The returned Result is truthy +// only when the handshake fully succeeded; on failure it describes what went +// wrong: +// res.error() httplib::Error identifying the failing layer +// res.status() HTTP status of the upgrade response (-1 if none) +// res.headers() headers of the upgrade response +// res.ssl_error() TLS error detail (wss://, SSL builds only) +// res.ssl_backend_error() backend-specific TLS error code (SSL builds only) +Result connect(); // Get the subprotocol selected by the server (empty if none) const std::string &subprotocol() const; @@ -221,6 +228,26 @@ if (ws.connect()) { } ``` +### Inspecting Connection Failures + +`connect()` returns a `Result` that tells you why a connection attempt failed. +`error()` distinguishes network problems (`Connection`, `ConnectionTimeout`), +TLS problems (`SSLConnection`, `SSLServerVerification`, +`SSLServerHostnameVerification`), and upgrade rejections +(`WebSocketHandshake`). When the server answered with something other than +`101 Switching Protocols`, `status()` and `headers()` carry that response: + +```cpp +auto res = ws.connect(); +if (!res) { + std::cerr << "connect failed: " << httplib::to_string(res.error()) << std::endl; + if (res.status() != -1) { + // The server responded but refused the upgrade (e.g. 401, 404) + std::cerr << "HTTP status: " << res.status() << std::endl; + } +} +``` + ### Text and Binary Messages Check the `ReadResult` return value to distinguish between text and binary: diff --git a/httplib.h b/httplib.h index 231c37c..71c4364 100644 --- a/httplib.h +++ b/httplib.h @@ -1805,6 +1805,7 @@ enum class Error { HTTPParsing, InvalidRangeHeader, UnsupportedContentEncoding, + WebSocketHandshake, // For internal use only SSLPeerCouldBeClosed_, @@ -4202,6 +4203,50 @@ enum class CloseStatus : uint16_t { enum ReadResult : int { Fail = 0, Text = 1, Binary = 2 }; +// Result of WebSocketClient::connect(). Truthy only when the WebSocket +// upgrade handshake fully succeeded. On failure error() identifies the +// failing layer; status()/headers() expose the server's upgrade response +// when one was received (status() is -1 otherwise). +class Result { +public: + Result() = default; + Result(Error err, int status, Headers &&headers) + : err_(err), status_(status), headers_(std::move(headers)) {} + + explicit operator bool() const { return err_ == Error::Success; } + Error error() const { return err_; } + + // Upgrade response info + int status() const { return status_; } + const Headers &headers() const { return headers_; } + std::string get_header_value(const std::string &key, + const char *def = "") const { + return detail::get_header_value(headers_, key, def, 0); + } + bool has_header(const std::string &key) const { + return headers_.find(key) != headers_.end(); + } + +#ifdef CPPHTTPLIB_SSL_ENABLED + Result(Error err, int status, Headers &&headers, int ssl_error, + uint64_t ssl_backend_error) + : err_(err), status_(status), headers_(std::move(headers)), + ssl_error_(ssl_error), ssl_backend_error_(ssl_backend_error) {} + + int ssl_error() const { return ssl_error_; } + uint64_t ssl_backend_error() const { return ssl_backend_error_; } +#endif + +private: + Error err_ = Error::Unknown; // a default-constructed Result is falsy + int status_ = -1; + Headers headers_; +#ifdef CPPHTTPLIB_SSL_ENABLED + int ssl_error_ = 0; + uint64_t ssl_backend_error_ = 0; +#endif +}; + class WebSocket { public: WebSocket(const WebSocket &) = delete; @@ -4268,7 +4313,7 @@ public: bool is_valid() const; - bool connect(); + Result connect(); ReadResult read(std::string &msg); bool send(const std::string &data); bool send(const char *data, size_t len); @@ -4320,7 +4365,8 @@ public: private: void shutdown_and_close(); - bool create_stream(std::unique_ptr &strm); + bool create_stream(std::unique_ptr &strm, Error &error, + int &ssl_error, uint64_t &ssl_backend_error); void prepare_default_headers(Request &req); std::string host_; @@ -7716,42 +7762,94 @@ inline bool read_headers(Stream &strm, Headers &headers) { return true; } +inline bool parse_status_line(const char *line, std::string &version, + int &status, std::string &reason) { +#ifdef CPPHTTPLIB_ALLOW_LF_AS_LINE_TERMINATOR + thread_local const std::regex re("(HTTP/1\\.[01]) (\\d{3})(?: (.*?))?\r?\n"); +#else + thread_local const std::regex re("(HTTP/1\\.[01]) (\\d{3})(?: (.*?))?\r\n"); +#endif + + std::cmatch m; + if (!std::regex_match(line, m, re)) { return false; } + version = std::string(m[1]); + status = std::stoi(std::string(m[2])); + reason = std::string(m[3]); + return true; +} + +// Everything WebSocketClient::connect() reports about the upgrade exchange. +// status stays -1 until a status line is parsed, mirroring stream::Result. +struct WebSocketUpgradeResponse { + Error error = Error::Success; + int status = -1; + Headers headers; + std::string selected_subprotocol; +}; + inline bool read_websocket_upgrade_response(Stream &strm, const std::string &expected_accept, - std::string &selected_subprotocol) { + WebSocketUpgradeResponse &upgrade) { // Read status line const auto bufsiz = 2048; char buf[bufsiz]; stream_line_reader line_reader(strm, buf, bufsiz); - if (!line_reader.getline()) { return false; } + if (!line_reader.getline()) { + upgrade.error = Error::Read; + return false; + } - // Check for "HTTP/1.1 101" - auto line = std::string(line_reader.ptr(), line_reader.size()); - if (line.find("HTTP/1.1 101") == std::string::npos) { return false; } + std::string version; + std::string reason; + if (!parse_status_line(line_reader.ptr(), version, upgrade.status, reason)) { + upgrade.error = Error::WebSocketHandshake; + return false; + } - // Parse headers using existing read_headers - Headers headers; - if (!read_headers(strm, headers)) { return false; } + // Read the headers even for a rejection so the caller can see why the + // server refused the upgrade. A non-101 response may carry a body; it is + // deliberately left unread since the caller closes the socket right away. + if (!read_headers(strm, upgrade.headers)) { + upgrade.error = Error::Read; + return false; + } + + const auto &headers = upgrade.headers; + + if (upgrade.status != StatusCode::SwitchingProtocol_101) { + upgrade.error = Error::WebSocketHandshake; + return false; + } // Verify Upgrade: websocket (case-insensitive) auto upgrade_it = headers.find("Upgrade"); - if (upgrade_it == headers.end()) { return false; } - auto upgrade_val = case_ignore::to_lower(upgrade_it->second); - if (upgrade_val != "websocket") { return false; } + if (upgrade_it == headers.end() || + case_ignore::to_lower(upgrade_it->second) != "websocket") { + upgrade.error = Error::WebSocketHandshake; + return false; + } // Verify Connection header contains "Upgrade" (case-insensitive) auto connection_it = headers.find("Connection"); - if (connection_it == headers.end()) { return false; } - auto connection_val = case_ignore::to_lower(connection_it->second); - if (connection_val.find("upgrade") == std::string::npos) { return false; } + if (connection_it == headers.end() || + case_ignore::to_lower(connection_it->second).find("upgrade") == + std::string::npos) { + upgrade.error = Error::WebSocketHandshake; + return false; + } // Verify Sec-WebSocket-Accept header value auto it = headers.find("Sec-WebSocket-Accept"); - if (it == headers.end() || it->second != expected_accept) { return false; } + if (it == headers.end() || it->second != expected_accept) { + upgrade.error = Error::WebSocketHandshake; + return false; + } // Extract negotiated subprotocol auto proto_it = headers.find("Sec-WebSocket-Protocol"); - if (proto_it != headers.end()) { selected_subprotocol = proto_it->second; } + if (proto_it != headers.end()) { + upgrade.selected_subprotocol = proto_it->second; + } return true; } @@ -9501,7 +9599,7 @@ inline bool is_field_valid(const std::string &name, const std::string &value) { } // namespace fields inline bool perform_websocket_handshake(Stream &strm, Request &req, - std::string &selected_subprotocol) { + WebSocketUpgradeResponse &upgrade) { // Generate random Sec-WebSocket-Key thread_local std::mt19937 rng(std::random_device{}()); std::string key_bytes(16, '\0'); @@ -9526,20 +9624,26 @@ inline bool perform_websocket_handshake(Stream &strm, Request &req, // and would emit one small write per header. BufferStream bstrm; - if (write_request_line(bstrm, req.method, req.path) < 0) { return false; } + if (write_request_line(bstrm, req.method, req.path) < 0) { + upgrade.error = Error::Write; + return false; + } auto error = Error::Success; if (!check_and_write_headers(bstrm, req.headers, write_headers, error)) { + upgrade.error = error; return false; } const auto &data = bstrm.get_buffer(); - if (!write_data(strm, data.data(), data.size())) { return false; } + if (!write_data(strm, data.data(), data.size())) { + upgrade.error = Error::Write; + return false; + } // Verify 101 response and Sec-WebSocket-Accept header auto expected_accept = websocket_accept_key(client_key); - return read_websocket_upgrade_response(strm, expected_accept, - selected_subprotocol); + return read_websocket_upgrade_response(strm, expected_accept, upgrade); } inline bool is_ip_address(const std::string &host) { @@ -10322,6 +10426,7 @@ inline std::string to_string(const Error error) { case Error::HTTPParsing: return "HTTP parsing failed"; case Error::InvalidRangeHeader: return "Invalid Range header"; case Error::UnsupportedContentEncoding: return "Unsupported Content-Encoding"; + case Error::WebSocketHandshake: return "WebSocket handshake failed"; default: break; } @@ -13713,29 +13818,20 @@ inline bool ClientImpl::read_response_line(Stream &strm, const Request &req, if (!line_reader.getline()) { return false; } -#ifdef CPPHTTPLIB_ALLOW_LF_AS_LINE_TERMINATOR - thread_local const std::regex re("(HTTP/1\\.[01]) (\\d{3})(?: (.*?))?\r?\n"); -#else - thread_local const std::regex re("(HTTP/1\\.[01]) (\\d{3})(?: (.*?))?\r\n"); -#endif - - std::cmatch m; - if (!std::regex_match(line_reader.ptr(), m, re)) { + if (!detail::parse_status_line(line_reader.ptr(), res.version, res.status, + res.reason)) { return req.method == "CONNECT"; } - res.version = std::string(m[1]); - res.status = std::stoi(std::string(m[2])); - res.reason = std::string(m[3]); // Ignore '100 Continue' (only when not using Expect: 100-continue explicitly) while (skip_100_continue && res.status == StatusCode::Continue_100) { if (!line_reader.getline()) { return false; } // CRLF if (!line_reader.getline()) { return false; } // next response line - if (!std::regex_match(line_reader.ptr(), m, re)) { return false; } - res.version = std::string(m[1]); - res.status = std::stoi(std::string(m[2])); - res.reason = std::string(m[3]); + if (!detail::parse_status_line(line_reader.ptr(), res.version, res.status, + res.reason)) { + return false; + } } return true; @@ -21351,7 +21447,9 @@ inline void WebSocketClient::shutdown_and_close() { } } -inline bool WebSocketClient::create_stream(std::unique_ptr &strm) { +inline bool WebSocketClient::create_stream(std::unique_ptr &strm, + Error &error, int &ssl_error, + uint64_t &ssl_backend_error) { #ifdef CPPHTTPLIB_SSL_ENABLED if (is_ssl_) { // A plain flag rather than SSLClient::load_certs()'s call_once: connect() @@ -21365,10 +21463,14 @@ inline bool WebSocketClient::create_stream(std::unique_ptr &strm) { certs_loaded_ = true; } + detail::ClientTlsSessionError tls_error; if (!detail::setup_client_tls_session(host_, tls_ctx_, tls_session_, sock_, server_certificate_verification_, - read_timeout_sec_, - read_timeout_usec_)) { + read_timeout_sec_, read_timeout_usec_, + &tls_error)) { + error = tls_error.error; + ssl_error = tls_error.ssl_error; + ssl_backend_error = tls_error.backend_error; return false; } @@ -21377,6 +21479,10 @@ inline bool WebSocketClient::create_stream(std::unique_ptr &strm) { write_timeout_sec_, write_timeout_usec_)); return true; } +#else + (void)error; + (void)ssl_error; + (void)ssl_backend_error; #endif strm = std::unique_ptr( new detail::SocketStream(sock_, read_timeout_sec_, read_timeout_usec_, @@ -21399,8 +21505,8 @@ inline void WebSocketClient::prepare_default_headers(Request &req) { detail::add_default_user_agent_header(req); } -inline bool WebSocketClient::connect() { - if (!is_valid_) { return false; } +inline Result WebSocketClient::connect() { + if (!is_valid_) { return Result{Error::Connection, -1, Headers{}}; } shutdown_and_close(); // Check is custom IP or hostname specified for host_ @@ -21408,19 +21514,29 @@ inline bool WebSocketClient::connect() { std::string ip; detail::apply_addr_map(addr_map_, host_, connect_host, ip); - Error error; + auto error = Error::Success; sock_ = detail::create_client_socket( connect_host, ip, port_, address_family_, tcp_nodelay_, ipv6_v6only_, socket_options_, connection_timeout_sec_, connection_timeout_usec_, read_timeout_sec_, read_timeout_usec_, write_timeout_sec_, write_timeout_usec_, interface_, error); - if (sock_ == INVALID_SOCKET) { return false; } + if (sock_ == INVALID_SOCKET) { + if (error == Error::Success) { error = Error::Connection; } + return Result{error, -1, Headers{}}; + } std::unique_ptr strm; - if (!create_stream(strm)) { + auto stream_error = Error::SSLConnection; + int ssl_error = 0; + uint64_t ssl_backend_error = 0; + if (!create_stream(strm, stream_error, ssl_error, ssl_backend_error)) { shutdown_and_close(); - return false; +#ifdef CPPHTTPLIB_SSL_ENABLED + return Result{stream_error, -1, Headers{}, ssl_error, ssl_backend_error}; +#else + return Result{stream_error, -1, Headers{}}; +#endif } Request req; @@ -21429,17 +21545,17 @@ inline bool WebSocketClient::connect() { req.headers = headers_; prepare_default_headers(req); - std::string selected_subprotocol; - if (!detail::perform_websocket_handshake(*strm, req, selected_subprotocol)) { + detail::WebSocketUpgradeResponse upgrade; + if (!detail::perform_websocket_handshake(*strm, req, upgrade)) { shutdown_and_close(); - return false; + return Result{upgrade.error, upgrade.status, std::move(upgrade.headers)}; } - subprotocol_ = std::move(selected_subprotocol); + subprotocol_ = std::move(upgrade.selected_subprotocol); ws_ = std::unique_ptr(new WebSocket(std::move(strm), req, false, websocket_ping_interval_sec_, websocket_max_missed_pongs_)); - return true; + return Result{Error::Success, upgrade.status, std::move(upgrade.headers)}; } inline ReadResult WebSocketClient::read(std::string &msg) { diff --git a/test/test.cc b/test/test.cc index 689973c..08ca72d 100644 --- a/test/test.cc +++ b/test/test.cc @@ -19896,7 +19896,11 @@ TEST(WebSocketTest, ConnectAndDisconnect) { svr.wait_until_ready(); ws::WebSocketClient client("ws://localhost:" + std::to_string(port) + "/ws"); - ASSERT_TRUE(client.connect()); + auto res = client.connect(); + ASSERT_TRUE(res); + EXPECT_EQ(Error::Success, res.error()); + EXPECT_EQ(StatusCode::SwitchingProtocol_101, res.status()); + EXPECT_TRUE(res.has_header("Sec-WebSocket-Accept")); EXPECT_TRUE(client.is_open()); client.close(); EXPECT_FALSE(client.is_open()); @@ -19972,7 +19976,23 @@ TEST(WebSocketTest, UnsupportedScheme) { TEST(WebSocketTest, ConnectWhenInvalid) { ws::WebSocketClient ws("not a valid url"); EXPECT_FALSE(ws.is_valid()); - EXPECT_FALSE(ws.connect()); + auto res = ws.connect(); + EXPECT_FALSE(res); + EXPECT_EQ(Error::Connection, res.error()); + EXPECT_EQ(-1, res.status()); +} + +TEST(WebSocketTest, ConnectRefusedReportsError) { + // Grab a port that is free, then close it again so nothing listens there + Server svr; + auto port = svr.bind_to_any_port(HOST); + svr.stop(); + + ws::WebSocketClient client("ws://localhost:" + std::to_string(port) + "/ws"); + auto res = client.connect(); + ASSERT_FALSE(res); + EXPECT_EQ(Error::Connection, res.error()); + EXPECT_EQ(-1, res.status()); } TEST(WebSocketTest, DefaultPort) { @@ -20368,7 +20388,10 @@ TEST_F(WebSocketIntegrationTest, MaxPayloadAtLimit) { TEST_F(WebSocketIntegrationTest, ConnectToInvalidPath) { ws::WebSocketClient client("ws://localhost:" + std::to_string(port_) + "/nonexistent"); - EXPECT_FALSE(client.connect()); + auto res = client.connect(); + EXPECT_FALSE(res); + EXPECT_EQ(Error::WebSocketHandshake, res.error()); + EXPECT_EQ(StatusCode::NotFound_404, res.status()); EXPECT_FALSE(client.is_open()); } @@ -20955,7 +20978,11 @@ TEST_F(WebSocketSSLCATest, WrongCustomCaFailsVerification) { read_file(CLIENT_CA_CERT_FILE, cert); client.load_ca_cert_store(cert.c_str(), cert.size()); - ASSERT_FALSE(client.connect()); + auto res = client.connect(); + ASSERT_FALSE(res); + EXPECT_EQ(Error::SSLServerVerification, res.error()); + EXPECT_EQ(-1, res.status()); + EXPECT_NE(0u, res.ssl_backend_error()); } // The same CA as a file path rather than PEM in memory @@ -21090,7 +21117,10 @@ TEST_F(WebSocketSSLDnsHostTest, TrustedChainWrongNameFails) { ws::WebSocketClient client(url()); client.set_ca_cert_path(SERVER_CERT_FILE); - ASSERT_FALSE(client.connect()); + auto res = client.connect(); + ASSERT_FALSE(res); + EXPECT_EQ(Error::SSLServerHostnameVerification, res.error()); + EXPECT_EQ(-1, res.status()); } // A CA that did not sign the server certificate fails the chain, even though