From d2e1de63e02f500c9cb9f5361cc4d213dd51a2e8 Mon Sep 17 00:00:00 2001 From: Florian Valeye Date: Thu, 30 Jul 2026 23:47:16 +0200 Subject: [PATCH] fix(client): preserve prepared null parameter types --- src/altertable_flightsql/client.py | 18 +++++++++++++++--- tests/test_client.py | 21 ++++++++++++++++++++- 2 files changed, 35 insertions(+), 4 deletions(-) diff --git a/src/altertable_flightsql/client.py b/src/altertable_flightsql/client.py index a5b61ac..ef097ac 100644 --- a/src/altertable_flightsql/client.py +++ b/src/altertable_flightsql/client.py @@ -697,7 +697,18 @@ def _get_parameter_as_pyarrow( elif isinstance(parameters, pa.RecordBatch): return parameters elif isinstance(parameters, Mapping): - return pa.record_batch({key: [value] for (key, value) in parameters.items()}) + return pa.RecordBatch.from_pydict( + { + key: ( + pa.array([value], type=self._parameter_schema.field(key).type) + if value is None + and self._parameter_schema is not None + and key in self._parameter_schema.names + else [value] + ) + for key, value in parameters.items() + } + ) elif isinstance(parameters, Sequence): if self._parameter_schema is None: raise ValueError( @@ -711,10 +722,11 @@ def _get_parameter_as_pyarrow( f"Expected {len(self._parameter_schema)} parameters, but got {len(parameters)}" ) param_dict = { - field.name: [value] for field, value in zip(self._parameter_schema, parameters) + field.name: pa.array([value], type=field.type) if value is None else [value] + for field, value in zip(self._parameter_schema, parameters) } - return pa.record_batch(param_dict) + return pa.RecordBatch.from_pydict(param_dict) else: raise TypeError( f"Unsupported parameter type: {type(parameters)}. " diff --git a/tests/test_client.py b/tests/test_client.py index b0ebd02..edafdd9 100644 --- a/tests/test_client.py +++ b/tests/test_client.py @@ -1,10 +1,11 @@ from types import SimpleNamespace +import pyarrow as pa import pyarrow.flight as flight import pytest from google.protobuf import any_pb2 -from altertable_flightsql.client import Client +from altertable_flightsql.client import Client, PreparedStatement from altertable_flightsql.generated import arrow_flight_pb2 as flight_pb2 @@ -69,6 +70,24 @@ def _action_body_bytes(action) -> bytes: return bytes(body) +@pytest.mark.parametrize( + "values", + [ + [None], + {"amount": None}, + ], + ids=["positional", "mapping"], +) +def test_python_null_parameters_use_prepared_type(values): + parameter_schema = pa.schema([("amount", pa.float64())]) + statement = PreparedStatement(None, b"handle", parameter_schema=parameter_schema) + + parameters = statement._get_parameter_as_pyarrow(values) + + assert parameters.schema.equals(parameter_schema) + assert parameters.to_pydict() == {"amount": [None]} + + def test_set_options_serializes_flight_session_request_without_any(): flight_client = FakeFlightClient() client = _client_backed_by(flight_client)