diff --git a/pyproject.toml b/pyproject.toml index 55aef88..244d82b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -37,7 +37,7 @@ dependencies = [ [project.optional-dependencies] dev = [ - "grpcio-tools>=1.73.0", + "grpcio-tools==1.80.0", "pytest>=7.0.0", "pytest-cov>=3.0.0", "black>=22.0.0", diff --git a/src/altertable_flightsql/client.py b/src/altertable_flightsql/client.py index 2adad99..a5b61ac 100644 --- a/src/altertable_flightsql/client.py +++ b/src/altertable_flightsql/client.py @@ -142,6 +142,7 @@ def __init__( self._password = password self._auto_commit = auto_commit self._transaction = None + self._closed = False auth_middleware = BearerAuthMiddlewareFactory() self._client = flight.FlightClient( @@ -549,9 +550,24 @@ def _end_transaction(self, transaction: "Transaction", commit: bool) -> None: action = flight.Action("EndTransaction", _pack_command(request)) list(self._client.do_action(action)) - def close(self) -> None: - """Close the client connection.""" - self._client.close() + def close(self, timeout_seconds: float = 10.0) -> None: + """Close the server session and client transport. Idempotent.""" + if self._closed: + return + self._closed = True + request = flight_pb2.CloseSessionRequest() + action = flight.Action("CloseSession", request.SerializeToString()) + options = flight.FlightCallOptions(timeout=timeout_seconds) + try: + results = list(self._client.do_action(action, options)) + if not results: + raise RuntimeError("Server returned no CloseSessionResult") + close_result = flight_pb2.CloseSessionResult.FromString(bytes(results[0].body)) + if close_result.status != flight_pb2.CloseSessionResult.CLOSED: + status = flight_pb2.CloseSessionResult.Status.Name(close_result.status) + raise RuntimeError(f"Server did not close Flight session: {status}") + finally: + self._client.close() def __enter__(self) -> "Client": """Context manager entry.""" diff --git a/tests/test_client.py b/tests/test_client.py index fae3f3e..b0ebd02 100644 --- a/tests/test_client.py +++ b/tests/test_client.py @@ -1,3 +1,7 @@ +from types import SimpleNamespace + +import pyarrow.flight as flight +import pytest from google.protobuf import any_pb2 from altertable_flightsql.client import Client @@ -7,10 +11,55 @@ class FakeFlightClient: def __init__(self): self.actions = [] + self.options = [] + self.events = [] + self.closed = False + close_result = flight_pb2.CloseSessionResult(status=flight_pb2.CloseSessionResult.CLOSED) + self.action_results = [SimpleNamespace(body=close_result.SerializeToString())] - def do_action(self, action): + def do_action(self, action, options=None): + if self.closed: + raise RuntimeError("FlightClient is closed") self.actions.append(action) - return [] + self.options.append(options) + self.events.append(("action", action.type)) + return self.action_results + + def close(self): + self.closed = True + self.events.append(("close", None)) + + +class FailingCloseSessionFlightClient(FakeFlightClient): + def do_action(self, action, options=None): + super().do_action(action, options) + raise RuntimeError("close session failed") + + +def _record_call_option_timeouts(monkeypatch) -> list: + """`FlightCallOptions.timeout` is unreadable on older PyArrow, so record construction.""" + timeouts = [] + build_options = flight.FlightCallOptions + + def recording_options(**kwargs): + timeouts.append(kwargs.get("timeout")) + return build_options(**kwargs) + + monkeypatch.setattr(flight, "FlightCallOptions", recording_options) + return timeouts + + +def _closed_result_then_error(): + close_result = flight_pb2.CloseSessionResult(status=flight_pb2.CloseSessionResult.CLOSED) + yield SimpleNamespace(body=close_result.SerializeToString()) + raise RuntimeError("server failed after the first result") + + +def _client_backed_by(flight_client) -> Client: + client = Client.__new__(Client) + client._client = flight_client + client._closed = False + return client def _action_body_bytes(action) -> bytes: @@ -22,8 +71,7 @@ def _action_body_bytes(action) -> bytes: def test_set_options_serializes_flight_session_request_without_any(): flight_client = FakeFlightClient() - client = Client.__new__(Client) - client._client = flight_client + client = _client_backed_by(flight_client) session_options = { "catalog": flight_pb2.SessionOptionValue(string_value="test_catalog"), @@ -38,3 +86,100 @@ def test_set_options_serializes_flight_session_request_without_any(): assert action.type == "SetSessionOptions" assert _action_body_bytes(action) == request.SerializeToString() assert _action_body_bytes(action) != wrapped_request.SerializeToString() + + +def test_close_closes_server_session_before_transport(): + flight_client = FakeFlightClient() + client = _client_backed_by(flight_client) + + client.close() + + action = flight_client.actions[0] + request = flight_pb2.CloseSessionRequest() + + assert flight_client.events == [("action", "CloseSession"), ("close", None)] + assert _action_body_bytes(action) == request.SerializeToString() + + +def test_close_closes_transport_when_server_session_close_fails(): + flight_client = FailingCloseSessionFlightClient() + client = _client_backed_by(flight_client) + + with pytest.raises(RuntimeError, match="close session failed"): + client.close() + + assert flight_client.events == [("action", "CloseSession"), ("close", None)] + + +def test_close_rejects_unclosed_server_session(): + flight_client = FakeFlightClient() + close_result = flight_pb2.CloseSessionResult(status=flight_pb2.CloseSessionResult.NOT_CLOSEABLE) + flight_client.action_results = [SimpleNamespace(body=close_result.SerializeToString())] + client = _client_backed_by(flight_client) + + with pytest.raises(RuntimeError, match="NOT_CLOSEABLE"): + client.close() + + assert flight_client.events == [("action", "CloseSession"), ("close", None)] + + +def test_close_rejects_a_session_the_server_is_still_closing(): + flight_client = FakeFlightClient() + close_result = flight_pb2.CloseSessionResult(status=flight_pb2.CloseSessionResult.CLOSING) + flight_client.action_results = [SimpleNamespace(body=close_result.SerializeToString())] + client = _client_backed_by(flight_client) + + with pytest.raises(RuntimeError, match="CLOSING"): + client.close() + + assert flight_client.events == [("action", "CloseSession"), ("close", None)] + + +def test_close_is_idempotent(): + flight_client = FakeFlightClient() + client = _client_backed_by(flight_client) + + client.close() + client.close() + + assert flight_client.events == [("action", "CloseSession"), ("close", None)] + + +def test_close_passes_bounded_timeout_to_do_action(monkeypatch): + timeouts = _record_call_option_timeouts(monkeypatch) + client = _client_backed_by(FakeFlightClient()) + + client.close() + + assert timeouts == [10.0] + + +def test_close_passes_caller_timeout_to_do_action(monkeypatch): + timeouts = _record_call_option_timeouts(monkeypatch) + client = _client_backed_by(FakeFlightClient()) + + client.close(timeout_seconds=2.5) + + assert timeouts == [2.5] + + +def test_close_surfaces_server_error_raised_after_the_first_result(): + flight_client = FakeFlightClient() + flight_client.action_results = _closed_result_then_error() + client = _client_backed_by(flight_client) + + with pytest.raises(RuntimeError, match="server failed after the first result"): + client.close() + + assert flight_client.events == [("action", "CloseSession"), ("close", None)] + + +def test_close_rejects_empty_server_response(): + flight_client = FakeFlightClient() + flight_client.action_results = [] + client = _client_backed_by(flight_client) + + with pytest.raises(RuntimeError, match="no CloseSessionResult"): + client.close() + + assert flight_client.events == [("action", "CloseSession"), ("close", None)] diff --git a/uv.lock b/uv.lock index aaedd47..d32c456 100644 --- a/uv.lock +++ b/uv.lock @@ -43,7 +43,7 @@ dev = [ [package.metadata] requires-dist = [ { name = "black", marker = "extra == 'dev'", specifier = ">=22.0.0" }, - { name = "grpcio-tools", marker = "extra == 'dev'", specifier = ">=1.73.0" }, + { name = "grpcio-tools", marker = "extra == 'dev'", specifier = "==1.80.0" }, { name = "isort", marker = "extra == 'dev'", specifier = ">=5.10.0" }, { name = "pandas", marker = "extra == 'dev'", specifier = ">=1.3.0" }, { name = "protobuf", specifier = ">=5.26.0" },