Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
test_ssl_edge_cases.py639 linesDownload Raw Back to tests
1# -*- coding: utf-8 -*-2import unittest3import socket4import ssl5from unittest.mock import Mock, patch, MagicMock6 7from websocket._ssl_compat import (8    SSLError,9    SSLEOFError,10    SSLWantReadError,11    SSLWantWriteError,12    HAVE_SSL,13)14from websocket._http import _ssl_socket, _wrap_sni_socket15from websocket._exceptions import WebSocketException16from websocket._socket import recv, send17 18"""19test_ssl_edge_cases.py20websocket - WebSocket client library for Python21 22Copyright 2025 engn33r23 24Licensed under the Apache License, Version 2.0 (the "License");25you may not use this file except in compliance with the License.26You may obtain a copy of the License at27 28    http://www.apache.org/licenses/LICENSE-2.029 30Unless required by applicable law or agreed to in writing, software31distributed under the License is distributed on an "AS IS" BASIS,32WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.33See the License for the specific language governing permissions and34limitations under the License.35"""36 37class SSLEdgeCasesTest(unittest.TestCase):38 39    def setUp(self):40        if not HAVE_SSL:41            self.skipTest("SSL not available")42 43    def test_ssl_handshake_failure(self):44        """Test SSL handshake failure scenarios"""45        mock_sock = Mock()46 47        # Test SSL handshake timeout48        with patch("ssl.SSLContext") as mock_ssl_context:49            mock_context = Mock()50            mock_ssl_context.return_value = mock_context51            mock_context.wrap_socket.side_effect = socket.timeout(52                "SSL handshake timeout"53            )54 55            sslopt = {"cert_reqs": ssl.CERT_REQUIRED}56 57            with self.assertRaises(socket.timeout):58                _ssl_socket(mock_sock, sslopt, "example.com")59 60    def test_ssl_certificate_verification_failures(self):61        """Test various SSL certificate verification failure scenarios"""62        mock_sock = Mock()63 64        # Test certificate verification failure65        with patch("ssl.SSLContext") as mock_ssl_context:66            mock_context = Mock()67            mock_ssl_context.return_value = mock_context68            mock_context.wrap_socket.side_effect = ssl.SSLCertVerificationError(69                "Certificate verification failed"70            )71 72            sslopt = {"cert_reqs": ssl.CERT_REQUIRED, "check_hostname": True}73 74            with self.assertRaises(ssl.SSLCertVerificationError):75                _ssl_socket(mock_sock, sslopt, "badssl.example")76 77    def test_ssl_context_configuration_edge_cases(self):78        """Test SSL context configuration with various edge cases"""79        mock_sock = Mock()80 81        # Test with pre-created SSL context82        with patch("ssl.SSLContext") as mock_ssl_context:83            existing_context = Mock()84            existing_context.wrap_socket.return_value = Mock()85            mock_ssl_context.return_value = existing_context86 87            sslopt = {"context": existing_context}88 89            # Call _ssl_socket which should use the existing context90            _ssl_socket(mock_sock, sslopt, "example.com")91 92            # Should use the provided context, not create a new one93            existing_context.wrap_socket.assert_called_once()94 95    def test_ssl_ca_bundle_environment_edge_cases(self):96        """Test CA bundle environment variable edge cases"""97        mock_sock = Mock()98 99        # Test with non-existent CA bundle file100        with patch.dict(101            "os.environ", {"WEBSOCKET_CLIENT_CA_BUNDLE": "/nonexistent/ca-bundle.crt"}102        ):103            with patch("os.path.isfile", return_value=False):104                with patch("os.path.isdir", return_value=False):105                    with patch("ssl.SSLContext") as mock_ssl_context:106                        mock_context = Mock()107                        mock_ssl_context.return_value = mock_context108                        mock_context.wrap_socket.return_value = Mock()109 110                        sslopt = {}111                        _ssl_socket(mock_sock, sslopt, "example.com")112 113                        # Should not try to load non-existent CA bundle114                        mock_context.load_verify_locations.assert_not_called()115 116        # Test with CA bundle directory117        with patch.dict("os.environ", {"WEBSOCKET_CLIENT_CA_BUNDLE": "/etc/ssl/certs"}):118            with patch("os.path.isfile", return_value=False):119                with patch("os.path.isdir", return_value=True):120                    with patch("ssl.SSLContext") as mock_ssl_context:121                        mock_context = Mock()122                        mock_ssl_context.return_value = mock_context123                        mock_context.wrap_socket.return_value = Mock()124 125                        sslopt = {}126                        _ssl_socket(mock_sock, sslopt, "example.com")127 128                        # Should load CA directory129                        mock_context.load_verify_locations.assert_called_with(130                            cafile=None, capath="/etc/ssl/certs"131                        )132 133    def test_ssl_cipher_configuration_edge_cases(self):134        """Test SSL cipher configuration edge cases"""135        mock_sock = Mock()136 137        # Test with invalid cipher suite138        with patch("ssl.SSLContext") as mock_ssl_context:139            mock_context = Mock()140            mock_ssl_context.return_value = mock_context141            mock_context.set_ciphers.side_effect = ssl.SSLError(142                "No cipher can be selected"143            )144            mock_context.wrap_socket.return_value = Mock()145 146            sslopt = {"ciphers": "INVALID_CIPHER"}147 148            with self.assertRaises(WebSocketException):149                _ssl_socket(mock_sock, sslopt, "example.com")150 151    def test_ssl_ecdh_curve_edge_cases(self):152        """Test ECDH curve configuration edge cases"""153        mock_sock = Mock()154 155        # Test with invalid ECDH curve156        with patch("ssl.SSLContext") as mock_ssl_context:157            mock_context = Mock()158            mock_ssl_context.return_value = mock_context159            mock_context.set_ecdh_curve.side_effect = ValueError("unknown curve name")160            mock_context.wrap_socket.return_value = Mock()161 162            sslopt = {"ecdh_curve": "invalid_curve"}163 164            with self.assertRaises(WebSocketException):165                _ssl_socket(mock_sock, sslopt, "example.com")166 167    def test_ssl_client_certificate_edge_cases(self):168        """Test client certificate configuration edge cases"""169        mock_sock = Mock()170 171        # Test with non-existent client certificate172        with patch("ssl.SSLContext") as mock_ssl_context:173            mock_context = Mock()174            mock_ssl_context.return_value = mock_context175            mock_context.load_cert_chain.side_effect = FileNotFoundError("No such file")176            mock_context.wrap_socket.return_value = Mock()177 178            sslopt = {"certfile": "/nonexistent/client.crt"}179 180            with self.assertRaises(WebSocketException):181                _ssl_socket(mock_sock, sslopt, "example.com")182 183    def test_ssl_want_read_write_retry_edge_cases(self):184        """Test SSL want read/write retry edge cases"""185        mock_sock = Mock()186 187        # Test SSLWantReadError with multiple retries before success188        read_attempts = [0]  # Use list for mutable reference189 190        def mock_recv(bufsize):191            read_attempts[0] += 1192            if read_attempts[0] == 1:193                raise SSLWantReadError("The operation did not complete")194            elif read_attempts[0] == 2:195                return b"data after retries"196            else:197                return b""198 199        mock_sock.recv.side_effect = mock_recv200        mock_sock.gettimeout.return_value = 30.0201 202        with patch("selectors.DefaultSelector") as mock_selector_class:203            mock_selector = Mock()204            mock_selector_class.return_value = mock_selector205            mock_selector.select.return_value = [True]  # Always ready206 207            result = recv(mock_sock, 100)208 209            self.assertEqual(result, b"data after retries")210            self.assertEqual(read_attempts[0], 2)211            # Should have used selector for retry212            mock_selector.register.assert_called()213            mock_selector.select.assert_called()214 215    def test_ssl_want_write_retry_edge_cases(self):216        """Test SSL want write retry edge cases"""217        mock_sock = Mock()218 219        # Test SSLWantWriteError with multiple retries before success220        write_attempts = [0]  # Use list for mutable reference221 222        def mock_send(data):223            write_attempts[0] += 1224            if write_attempts[0] == 1:225                raise SSLWantWriteError("The operation did not complete")226            elif write_attempts[0] == 2:227                return len(data)228            else:229                return 0230 231        mock_sock.send.side_effect = mock_send232        mock_sock.gettimeout.return_value = 30.0233 234        with patch("selectors.DefaultSelector") as mock_selector_class:235            mock_selector = Mock()236            mock_selector_class.return_value = mock_selector237            mock_selector.select.return_value = [True]  # Always ready238 239            result = send(mock_sock, b"test data")240 241            self.assertEqual(result, 9)  # len("test data")242            self.assertEqual(write_attempts[0], 2)243 244    def test_ssl_eof_error_edge_cases(self):245        """Test SSL EOF error edge cases"""246        mock_sock = Mock()247 248        # Test SSLEOFError during send249        mock_sock.send.side_effect = SSLEOFError("SSL connection has been closed")250        mock_sock.gettimeout.return_value = 30.0251 252        from websocket._exceptions import WebSocketConnectionClosedException253 254        with self.assertRaises(WebSocketConnectionClosedException):255            send(mock_sock, b"test data")256 257    def test_ssl_pending_data_edge_cases(self):258        """Test SSL pending data scenarios"""259        from websocket._dispatcher import SSLDispatcher260        from websocket._app import WebSocketApp261 262        # Mock SSL socket with pending data263        mock_ssl_sock = Mock()264        mock_ssl_sock.pending.return_value = 1024  # Simulates pending SSL data265 266        # Mock WebSocketApp267        mock_app = Mock(spec=WebSocketApp)268        mock_app.sock = Mock()269        mock_app.sock.sock = mock_ssl_sock270 271        dispatcher = SSLDispatcher(mock_app, 5.0)272 273        # When there's pending data, should return immediately without selector274        result = dispatcher.select(mock_ssl_sock, Mock())275 276        # Should return the socket list when there's pending data277        self.assertEqual(result, [mock_ssl_sock])278        mock_ssl_sock.pending.assert_called_once()279 280    def test_ssl_renegotiation_edge_cases(self):281        """Test SSL renegotiation scenarios"""282        mock_sock = Mock()283 284        # Simulate SSL renegotiation during read285        call_count = 0286 287        def mock_recv(bufsize):288            nonlocal call_count289            call_count += 1290            if call_count == 1:291                raise SSLWantReadError("SSL renegotiation required")292            return b"data after renegotiation"293 294        mock_sock.recv.side_effect = mock_recv295        mock_sock.gettimeout.return_value = 30.0296 297        with patch("selectors.DefaultSelector") as mock_selector_class:298            mock_selector = Mock()299            mock_selector_class.return_value = mock_selector300            mock_selector.select.return_value = [True]301 302            result = recv(mock_sock, 100)303 304            self.assertEqual(result, b"data after renegotiation")305            self.assertEqual(call_count, 2)306 307    def test_ssl_server_hostname_override(self):308        """Test SSL server hostname override scenarios"""309        mock_sock = Mock()310 311        with patch("ssl.SSLContext") as mock_ssl_context:312            mock_context = Mock()313            mock_ssl_context.return_value = mock_context314            mock_context.wrap_socket.return_value = Mock()315 316            # Test server_hostname override317            sslopt = {"server_hostname": "override.example.com"}318            _ssl_socket(mock_sock, sslopt, "original.example.com")319 320            # Should use override hostname in wrap_socket call321            mock_context.wrap_socket.assert_called_with(322                mock_sock,323                do_handshake_on_connect=True,324                suppress_ragged_eofs=True,325                server_hostname="override.example.com",326            )327 328    def test_ssl_protocol_version_edge_cases(self):329        """Test SSL protocol version edge cases"""330        mock_sock = Mock()331 332        # Test with deprecated SSL version333        with patch("ssl.SSLContext") as mock_ssl_context:334            mock_context = Mock()335            mock_ssl_context.return_value = mock_context336            mock_context.wrap_socket.return_value = Mock()337 338            # Test that deprecated ssl_version is still handled339            if hasattr(ssl, "PROTOCOL_TLS"):340                sslopt = {"ssl_version": ssl.PROTOCOL_TLS}341                _ssl_socket(mock_sock, sslopt, "example.com")342 343                mock_ssl_context.assert_called_with(ssl.PROTOCOL_TLS)344 345    def test_ssl_keylog_file_edge_cases(self):346        """Test SSL keylog file configuration edge cases"""347        mock_sock = Mock()348 349        # Test with SSLKEYLOGFILE environment variable350        with patch.dict("os.environ", {"SSLKEYLOGFILE": "/tmp/ssl_keys.log"}):351            with patch("ssl.SSLContext") as mock_ssl_context:352                mock_context = Mock()353                mock_ssl_context.return_value = mock_context354                mock_context.wrap_socket.return_value = Mock()355 356                sslopt = {}357                _ssl_socket(mock_sock, sslopt, "example.com")358 359                # Should set keylog_filename360                self.assertEqual(mock_context.keylog_filename, "/tmp/ssl_keys.log")361 362    def test_ssl_context_verification_modes(self):363        """Test different SSL verification mode combinations"""364        mock_sock = Mock()365 366        test_cases = [367            # (cert_reqs, check_hostname, expected_verify_mode, expected_check_hostname)368            (ssl.CERT_NONE, False, ssl.CERT_NONE, False),369            (ssl.CERT_REQUIRED, False, ssl.CERT_REQUIRED, False),370            (ssl.CERT_REQUIRED, True, ssl.CERT_REQUIRED, True),371        ]372 373        for cert_reqs, check_hostname, expected_verify, expected_check in test_cases:374            with self.subTest(cert_reqs=cert_reqs, check_hostname=check_hostname):375                with patch("ssl.SSLContext") as mock_ssl_context:376                    mock_context = Mock()377                    mock_ssl_context.return_value = mock_context378                    mock_context.wrap_socket.return_value = Mock()379 380                    sslopt = {"cert_reqs": cert_reqs, "check_hostname": check_hostname}381                    _ssl_socket(mock_sock, sslopt, "example.com")382 383                    self.assertEqual(mock_context.verify_mode, expected_verify)384                    self.assertEqual(mock_context.check_hostname, expected_check)385 386    def test_ssl_socket_shutdown_edge_cases(self):387        """Test SSL socket shutdown edge cases"""388        from websocket._core import WebSocket389 390        mock_ssl_sock = Mock()391        mock_ssl_sock.shutdown.side_effect = SSLError("SSL shutdown failed")392 393        ws = WebSocket()394        ws.sock = mock_ssl_sock395        ws.connected = True396 397        # Should handle SSL shutdown errors gracefully398        try:399            ws.close()400        except SSLError:401            self.fail("SSL shutdown error should be handled gracefully")402 403    def test_ssl_socket_close_during_operation(self):404        """Test SSL socket being closed during ongoing operations"""405        mock_sock = Mock()406 407        # Simulate SSL socket being closed during recv408        mock_sock.recv.side_effect = SSLError(409            "SSL connection has been closed unexpectedly"410        )411        mock_sock.gettimeout.return_value = 30.0412 413        from websocket._exceptions import WebSocketConnectionClosedException414 415        # Should handle unexpected SSL closure416        with self.assertRaises((SSLError, WebSocketConnectionClosedException)):417            recv(mock_sock, 100)418 419    def test_ssl_compression_edge_cases(self):420        """Test SSL compression configuration edge cases"""421        mock_sock = Mock()422 423        with patch("ssl.SSLContext") as mock_ssl_context:424            mock_context = Mock()425            mock_ssl_context.return_value = mock_context426            mock_context.wrap_socket.return_value = Mock()427 428            # Test SSL compression options (if available)429            sslopt = {"compression": False}  # Some SSL contexts support this430 431            try:432                _ssl_socket(mock_sock, sslopt, "example.com")433                # Should not fail even if compression option is not supported434            except AttributeError:435                # Expected if SSL context doesn't support compression option436                pass437 438    def test_ssl_session_reuse_edge_cases(self):439        """Test SSL session reuse scenarios"""440        mock_sock = Mock()441 442        with patch("ssl.SSLContext") as mock_ssl_context:443            mock_context = Mock()444            mock_ssl_context.return_value = mock_context445            mock_ssl_sock = Mock()446            mock_context.wrap_socket.return_value = mock_ssl_sock447 448            # Test session reuse449            mock_ssl_sock.session = "mock_session"450            mock_ssl_sock.session_reused = True451 452            result = _ssl_socket(mock_sock, {}, "example.com")453 454            # Should handle session reuse without issues455            self.assertIsNotNone(result)456 457    def test_ssl_alpn_protocol_edge_cases(self):458        """Test SSL ALPN (Application Layer Protocol Negotiation) edge cases"""459        mock_sock = Mock()460 461        with patch("ssl.SSLContext") as mock_ssl_context:462            mock_context = Mock()463            mock_ssl_context.return_value = mock_context464            mock_context.wrap_socket.return_value = Mock()465 466            # Test ALPN configuration467            sslopt = {"alpn_protocols": ["http/1.1", "h2"]}468 469            # ALPN protocols are not currently supported in the SSL wrapper470            # but the test should not fail471            result = _ssl_socket(mock_sock, sslopt, "example.com")472            self.assertIsNotNone(result)473            # ALPN would need to be implemented in _wrap_sni_socket function474 475    def test_ssl_sni_edge_cases(self):476        """Test SSL SNI (Server Name Indication) edge cases"""477        mock_sock = Mock()478 479        # Test with IPv6 address (should not use SNI)480        with patch("ssl.SSLContext") as mock_ssl_context:481            mock_context = Mock()482            mock_ssl_context.return_value = mock_context483            mock_context.wrap_socket.return_value = Mock()484 485            # IPv6 addresses should not be used for SNI486            ipv6_hostname = "2001:db8::1"487            _ssl_socket(mock_sock, {}, ipv6_hostname)488 489            # Should use IPv6 address as server_hostname490            mock_context.wrap_socket.assert_called_with(491                mock_sock,492                do_handshake_on_connect=True,493                suppress_ragged_eofs=True,494                server_hostname=ipv6_hostname,495            )496 497    def test_ssl_buffer_size_edge_cases(self):498        """Test SSL buffer size related edge cases"""499        mock_sock = Mock()500 501        def mock_recv(bufsize):502            # SSL should never try to read more than 16KB at once503            if bufsize > 16384:504                raise SSLError("[SSL: BAD_LENGTH] buffer too large")505            return b"A" * min(bufsize, 1024)  # Return smaller chunks506 507        mock_sock.recv.side_effect = mock_recv508        mock_sock.gettimeout.return_value = 30.0509 510        from websocket._abnf import frame_buffer511 512        # Frame buffer should handle large requests by chunking513        fb = frame_buffer(lambda size: recv(mock_sock, size), skip_utf8_validation=True)514 515        # This should work even with large size due to chunking516        result = fb.recv_strict(16384)  # Exactly 16KB517 518        self.assertGreater(len(result), 0)519 520    def test_ssl_protocol_downgrade_protection(self):521        """Test SSL protocol downgrade protection"""522        mock_sock = Mock()523 524        with patch("ssl.SSLContext") as mock_ssl_context:525            mock_context = Mock()526            mock_ssl_context.return_value = mock_context527            mock_context.wrap_socket.side_effect = ssl.SSLError(528                "SSLV3_ALERT_HANDSHAKE_FAILURE"529            )530 531            sslopt = {"ssl_version": ssl.PROTOCOL_TLS_CLIENT}532 533            # Should propagate SSL protocol errors534            with self.assertRaises(ssl.SSLError):535                _ssl_socket(mock_sock, sslopt, "example.com")536 537    def test_ssl_certificate_chain_validation(self):538        """Test SSL certificate chain validation edge cases"""539        mock_sock = Mock()540 541        with patch("ssl.SSLContext") as mock_ssl_context:542            mock_context = Mock()543            mock_ssl_context.return_value = mock_context544 545            # Test certificate chain validation failure546            mock_context.wrap_socket.side_effect = ssl.SSLCertVerificationError(547                "certificate verify failed: certificate has expired"548            )549 550            sslopt = {"cert_reqs": ssl.CERT_REQUIRED, "check_hostname": True}551 552            with self.assertRaises(ssl.SSLCertVerificationError):553                _ssl_socket(mock_sock, sslopt, "expired.badssl.com")554 555    def test_ssl_weak_cipher_rejection(self):556        """Test SSL weak cipher rejection scenarios"""557        mock_sock = Mock()558 559        with patch("ssl.SSLContext") as mock_ssl_context:560            mock_context = Mock()561            mock_ssl_context.return_value = mock_context562            mock_context.wrap_socket.side_effect = ssl.SSLError("no shared cipher")563 564            sslopt = {"ciphers": "RC4-MD5"}  # Intentionally weak cipher565 566            # Should fail with weak ciphers (SSL error is not wrapped by our code)567            with self.assertRaises(ssl.SSLError):568                _ssl_socket(mock_sock, sslopt, "example.com")569 570    def test_ssl_hostname_verification_edge_cases(self):571        """Test SSL hostname verification edge cases"""572        mock_sock = Mock()573 574        # Test with wildcard certificate scenarios575        test_cases = [576            ("*.example.com", "subdomain.example.com"),  # Valid wildcard577            ("*.example.com", "sub.subdomain.example.com"),  # Invalid wildcard depth578            ("example.com", "www.example.com"),  # Hostname mismatch579        ]580 581        for cert_hostname, connect_hostname in test_cases:582            with self.subTest(cert=cert_hostname, hostname=connect_hostname):583                with patch("ssl.SSLContext") as mock_ssl_context:584                    mock_context = Mock()585                    mock_ssl_context.return_value = mock_context586 587                    if (588                        cert_hostname != connect_hostname589                        and "sub.subdomain" in connect_hostname590                    ):591                        # Simulate hostname verification failure for invalid wildcard592                        mock_context.wrap_socket.side_effect = ssl.SSLCertVerificationError(593                            f"hostname '{connect_hostname}' doesn't match '{cert_hostname}'"594                        )595 596                        sslopt = {597                            "cert_reqs": ssl.CERT_REQUIRED,598                            "check_hostname": True,599                        }600 601                        with self.assertRaises(ssl.SSLCertVerificationError):602                            _ssl_socket(mock_sock, sslopt, connect_hostname)603                    else:604                        mock_context.wrap_socket.return_value = Mock()605                        sslopt = {606                            "cert_reqs": ssl.CERT_REQUIRED,607                            "check_hostname": True,608                        }609 610                        # Should succeed for valid cases611                        result = _ssl_socket(mock_sock, sslopt, connect_hostname)612                        self.assertIsNotNone(result)613 614    def test_ssl_memory_bio_edge_cases(self):615        """Test SSL memory BIO edge cases"""616        mock_sock = Mock()617 618        # Test SSL memory BIO scenarios (if available)619        try:620            import ssl621 622            if hasattr(ssl, "MemoryBIO"):623                with patch("ssl.SSLContext") as mock_ssl_context:624                    mock_context = Mock()625                    mock_ssl_context.return_value = mock_context626                    mock_context.wrap_socket.return_value = Mock()627 628                    # Memory BIO should work if available629                    _ssl_socket(mock_sock, {}, "example.com")630 631                    # Standard socket wrapping should still work632                    mock_context.wrap_socket.assert_called_once()633        except (ImportError, AttributeError):634            self.skipTest("SSL MemoryBIO not available")635 636 637if __name__ == "__main__":638    unittest.main()639 
codekingpro/portable-devtools · Team Ai