|
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,137 @@ 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_startup_close(self): |
| 388 | + """ |
| 389 | + Maintenance mode accepts regular CQL sockets and closes them during |
| 390 | + startup. The low-level factory keeps that close observable while |
| 391 | + still tracking pool-owned startup connections for shutdown cleanup. |
| 392 | + """ |
| 393 | + |
| 394 | + class MaintenanceModeCqlServer(object): |
| 395 | + def __init__(self): |
| 396 | + self._sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) |
| 397 | + self._sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) |
| 398 | + self._sock.bind(('127.0.0.1', 0)) |
| 399 | + self._sock.listen(1) |
| 400 | + self._sock.settimeout(2) |
| 401 | + self.port = self._sock.getsockname()[1] |
| 402 | + self.first_frame = b'' |
| 403 | + self.ready = Event() |
| 404 | + self.received_frame = Event() |
| 405 | + self.error = None |
| 406 | + self.thread = Thread(target=self._run) |
| 407 | + self.thread.daemon = True |
| 408 | + self.thread.start() |
| 409 | + |
| 410 | + def _run(self): |
| 411 | + self.ready.set() |
| 412 | + try: |
| 413 | + client, _ = self._sock.accept() |
| 414 | + with client: |
| 415 | + client.settimeout(2) |
| 416 | + while len(self.first_frame) < 9: |
| 417 | + chunk = client.recv(9 - len(self.first_frame)) |
| 418 | + if not chunk: |
| 419 | + break |
| 420 | + self.first_frame += chunk |
| 421 | + except Exception as exc: |
| 422 | + self.error = exc |
| 423 | + finally: |
| 424 | + self.received_frame.set() |
| 425 | + |
| 426 | + def close(self): |
| 427 | + self._sock.close() |
| 428 | + self.thread.join(2) |
| 429 | + |
| 430 | + class MaintenanceModeConnection(Connection): |
| 431 | + def __init__(self, *args, **kwargs): |
| 432 | + super(MaintenanceModeConnection, self).__init__(*args, **kwargs) |
| 433 | + self._reader = None |
| 434 | + self._connect_socket() |
| 435 | + self._send_options_message() |
| 436 | + self._reader = Thread(target=self._read_until_server_closes) |
| 437 | + self._reader.daemon = True |
| 438 | + self._reader.start() |
| 439 | + |
| 440 | + def push(self, data): |
| 441 | + self._socket.sendall(data) |
| 442 | + |
| 443 | + def close(self): |
| 444 | + with self.lock: |
| 445 | + if self.is_closed: |
| 446 | + return |
| 447 | + self.is_closed = True |
| 448 | + |
| 449 | + if self._socket: |
| 450 | + self._socket.close() |
| 451 | + |
| 452 | + if not self.is_defunct: |
| 453 | + shutdown_error = ConnectionShutdown("Connection to %s was closed" % self.endpoint) |
| 454 | + self.error_all_requests(shutdown_error) |
| 455 | + if not self.connected_event.is_set(): |
| 456 | + self.last_error = shutdown_error |
| 457 | + self.connected_event.set() |
| 458 | + |
| 459 | + def _read_until_server_closes(self): |
| 460 | + try: |
| 461 | + while True: |
| 462 | + data = self._socket.recv(self.in_buffer_size) |
| 463 | + if not data: |
| 464 | + self.close() |
| 465 | + return |
| 466 | + self._iobuf.write(data) |
| 467 | + self.process_io_buffer() |
| 468 | + except socket.error as exc: |
| 469 | + if not self.is_closed: |
| 470 | + self.defunct(exc) |
| 471 | + |
| 472 | + class PendingConnections(list): |
| 473 | + def __init__(self): |
| 474 | + super(PendingConnections, self).__init__() |
| 475 | + self.appended = [] |
| 476 | + |
| 477 | + def append(self, conn): |
| 478 | + self.appended.append(conn) |
| 479 | + super(PendingConnections, self).append(conn) |
| 480 | + |
| 481 | + server = MaintenanceModeCqlServer() |
| 482 | + try: |
| 483 | + assert server.ready.wait(2) |
| 484 | + conn = MaintenanceModeConnection.factory( |
| 485 | + DefaultEndPoint('127.0.0.1', server.port), timeout=2) |
| 486 | + |
| 487 | + assert conn.is_closed |
| 488 | + assert server.received_frame.wait(2) |
| 489 | + assert server.error is None |
| 490 | + assert len(server.first_frame) >= 5 |
| 491 | + assert server.first_frame[4] == 0x05 # OPTIONS |
| 492 | + finally: |
| 493 | + server.close() |
| 494 | + |
| 495 | + server = MaintenanceModeCqlServer() |
| 496 | + try: |
| 497 | + assert server.ready.wait(2) |
| 498 | + host_conn = Mock() |
| 499 | + host_conn.is_shutdown = False |
| 500 | + host_conn._pending_connections = PendingConnections() |
| 501 | + |
| 502 | + conn = MaintenanceModeConnection.factory( |
| 503 | + DefaultEndPoint('127.0.0.1', server.port), |
| 504 | + timeout=2, |
| 505 | + host_conn=host_conn) |
| 506 | + |
| 507 | + assert conn.is_closed |
| 508 | + assert server.received_frame.wait(2) |
| 509 | + assert server.error is None |
| 510 | + assert len(server.first_frame) >= 5 |
| 511 | + assert server.first_frame[4] == 0x05 # OPTIONS |
| 512 | + assert host_conn._pending_connections == [] |
| 513 | + assert len(host_conn._pending_connections.appended) == 1 |
| 514 | + assert host_conn._pending_connections.appended[0] is conn |
| 515 | + finally: |
| 516 | + server.close() |
| 517 | + |
386 | 518 |
|
387 | 519 | @patch('cassandra.connection.ConnectionHeartbeat._raise_if_stopped') |
388 | 520 | class ConnectionHeartbeatTest(unittest.TestCase): |
|
0 commit comments