|
12 | 12 | # See the License for the specific language governing permissions and |
13 | 13 | # limitations under the License. |
14 | 14 | import itertools |
| 15 | +import socket |
15 | 16 | import unittest |
16 | 17 | from io import BytesIO |
17 | 18 | import time |
18 | | -from threading import Lock |
| 19 | +from threading import Event, Lock, Thread |
19 | 20 | from unittest.mock import Mock, ANY, call, patch |
20 | 21 |
|
21 | 22 | from cassandra import OperationTimedOut |
@@ -383,6 +384,142 @@ def test_wait_for_responses_shutdown_includes_last_error(self): |
383 | 384 | assert "already closed" in error_message |
384 | 385 | assert "Bad file descriptor" in error_message |
385 | 386 |
|
| 387 | + def test_factory_returns_maintenance_mode_clean_startup_close(self): |
| 388 | + """ |
| 389 | + Maintenance mode accepts regular CQL sockets and closes them during |
| 390 | + startup. Direct callers keep that clean close observable, but |
| 391 | + pool-owned startup connections and reconnection probes must be rejected |
| 392 | + so unusable connections cannot be stored or reported as host recovery. |
| 393 | + """ |
| 394 | + |
| 395 | + class MaintenanceModeCqlServer(object): |
| 396 | + def __init__(self): |
| 397 | + self._sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) |
| 398 | + self._sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) |
| 399 | + self._sock.bind(('127.0.0.1', 0)) |
| 400 | + self._sock.listen(1) |
| 401 | + self._sock.settimeout(2) |
| 402 | + self.port = self._sock.getsockname()[1] |
| 403 | + self.first_frame = b'' |
| 404 | + self.ready = Event() |
| 405 | + self.received_frame = Event() |
| 406 | + self.error = None |
| 407 | + self.thread = Thread(target=self._run) |
| 408 | + self.thread.daemon = True |
| 409 | + self.thread.start() |
| 410 | + |
| 411 | + def _run(self): |
| 412 | + self.ready.set() |
| 413 | + try: |
| 414 | + client, _ = self._sock.accept() |
| 415 | + with client: |
| 416 | + client.settimeout(2) |
| 417 | + while len(self.first_frame) < 9: |
| 418 | + chunk = client.recv(9 - len(self.first_frame)) |
| 419 | + if not chunk: |
| 420 | + break |
| 421 | + self.first_frame += chunk |
| 422 | + except Exception as exc: |
| 423 | + self.error = exc |
| 424 | + finally: |
| 425 | + self.received_frame.set() |
| 426 | + |
| 427 | + def close(self): |
| 428 | + self._sock.close() |
| 429 | + self.thread.join(2) |
| 430 | + |
| 431 | + class MaintenanceModeConnection(Connection): |
| 432 | + def __init__(self, *args, **kwargs): |
| 433 | + super(MaintenanceModeConnection, self).__init__(*args, **kwargs) |
| 434 | + self._reader = None |
| 435 | + self._connect_socket() |
| 436 | + self._send_options_message() |
| 437 | + self._reader = Thread(target=self._read_until_server_closes) |
| 438 | + self._reader.daemon = True |
| 439 | + self._reader.start() |
| 440 | + |
| 441 | + def push(self, data): |
| 442 | + self._socket.sendall(data) |
| 443 | + |
| 444 | + def close(self): |
| 445 | + with self.lock: |
| 446 | + if self.is_closed: |
| 447 | + return |
| 448 | + self.is_closed = True |
| 449 | + |
| 450 | + if self._socket: |
| 451 | + self._socket.close() |
| 452 | + |
| 453 | + if not self.is_defunct: |
| 454 | + self.error_all_requests( |
| 455 | + ConnectionShutdown("Connection to %s was closed" % self.endpoint)) |
| 456 | + self.connected_event.set() |
| 457 | + |
| 458 | + def _read_until_server_closes(self): |
| 459 | + try: |
| 460 | + while True: |
| 461 | + data = self._socket.recv(self.in_buffer_size) |
| 462 | + if not data: |
| 463 | + self.close() |
| 464 | + return |
| 465 | + self._iobuf.write(data) |
| 466 | + self.process_io_buffer() |
| 467 | + except socket.error as exc: |
| 468 | + if not self.is_closed: |
| 469 | + self.defunct(exc) |
| 470 | + |
| 471 | + server = MaintenanceModeCqlServer() |
| 472 | + try: |
| 473 | + assert server.ready.wait(2) |
| 474 | + conn = MaintenanceModeConnection.factory( |
| 475 | + DefaultEndPoint('127.0.0.1', server.port), timeout=2) |
| 476 | + |
| 477 | + assert conn.is_closed |
| 478 | + assert server.received_frame.wait(2) |
| 479 | + assert server.error is None |
| 480 | + assert server.first_frame[4] == 0x05 # OPTIONS |
| 481 | + finally: |
| 482 | + server.close() |
| 483 | + |
| 484 | + server = MaintenanceModeCqlServer() |
| 485 | + try: |
| 486 | + assert server.ready.wait(2) |
| 487 | + host_conn = Mock() |
| 488 | + host_conn.is_shutdown = False |
| 489 | + host_conn._pending_connections = [] |
| 490 | + |
| 491 | + with pytest.raises(ConnectionShutdown) as exc_info: |
| 492 | + MaintenanceModeConnection.factory( |
| 493 | + DefaultEndPoint('127.0.0.1', server.port), |
| 494 | + timeout=2, |
| 495 | + host_conn=host_conn) |
| 496 | + |
| 497 | + assert "closed during the startup handshake" in str(exc_info.value) |
| 498 | + assert server.received_frame.wait(2) |
| 499 | + assert server.error is None |
| 500 | + assert server.first_frame[4] == 0x05 # OPTIONS |
| 501 | + assert len(host_conn._pending_connections) == 1 |
| 502 | + assert host_conn._pending_connections[0].is_closed |
| 503 | + finally: |
| 504 | + server.close() |
| 505 | + |
| 506 | + server = MaintenanceModeCqlServer() |
| 507 | + try: |
| 508 | + assert server.ready.wait(2) |
| 509 | + |
| 510 | + with pytest.raises(ConnectionShutdown) as exc_info: |
| 511 | + MaintenanceModeConnection.factory( |
| 512 | + DefaultEndPoint('127.0.0.1', server.port), |
| 513 | + timeout=2, |
| 514 | + _raise_on_startup_close=True) |
| 515 | + |
| 516 | + assert "closed during the startup handshake" in str(exc_info.value) |
| 517 | + assert server.received_frame.wait(2) |
| 518 | + assert server.error is None |
| 519 | + assert server.first_frame[4] == 0x05 # OPTIONS |
| 520 | + finally: |
| 521 | + server.close() |
| 522 | + |
386 | 523 |
|
387 | 524 | @patch('cassandra.connection.ConnectionHeartbeat._raise_if_stopped') |
388 | 525 | class ConnectionHeartbeatTest(unittest.TestCase): |
|
0 commit comments