Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
24 changes: 21 additions & 3 deletions src/altertable_flightsql/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -549,9 +550,26 @@ 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:
result = next(iter(self._client.do_action(action, options)), None)
if result is not None:
close_result = flight_pb2.CloseSessionResult.FromString(bytes(result.body))
if close_result.status not in (
flight_pb2.CloseSessionResult.CLOSED,
flight_pb2.CloseSessionResult.CLOSING,
):
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."""
Expand Down
91 changes: 87 additions & 4 deletions tests/test_client.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -7,10 +11,36 @@
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 _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:
Expand All @@ -22,8 +52,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"),
Expand All @@ -38,3 +67,57 @@ 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_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():
flight_client = FakeFlightClient()
client = _client_backed_by(flight_client)

client.close()

assert isinstance(flight_client.options[0], flight.FlightCallOptions)
2 changes: 1 addition & 1 deletion uv.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

Loading