diff --git a/src/archastro/platform/runtime/http_client.py b/src/archastro/platform/runtime/http_client.py index 2ba290d..27c52cb 100644 --- a/src/archastro/platform/runtime/http_client.py +++ b/src/archastro/platform/runtime/http_client.py @@ -18,6 +18,27 @@ T = TypeVar("T") +def _encode_query(query: dict[str, Any] | None) -> dict[str, Any] | None: + """Drop unset parameters and encode list values as ``key[]=item``. + + httpx writes a list under a bare repeated key (``source=a&source=b``), + and Plug's query parser keeps only the last value — the server would + silently filter on one element of a multi-value filter. The bracket + suffix parses as a list and matches the TypeScript SDK's encoding. + """ + if not query: + return None + params: dict[str, Any] = {} + for key, value in query.items(): + if value is None: + continue + if isinstance(value, (list, tuple)): + params[f"{key}[]"] = [str(item) for item in value] + else: + params[key] = value + return params or None + + @cache def _type_adapter(tp: Any) -> TypeAdapter[Any]: """Cache adapters per response type; TypeAdapter construction is costly.""" @@ -101,9 +122,7 @@ async def _do_fetch( if headers: req_headers.update(headers) - params = None - if query: - params = {k: v for k, v in query.items() if v is not None} + params = _encode_query(query) return await self._client.request( method, @@ -278,7 +297,7 @@ async def stream_sse( req_headers["Authorization"] = f"Bearer {token}" if headers: req_headers.update(headers) - params = {k: v for k, v in (query or {}).items() if v is not None} or None + params = _encode_query(query) async with self._client.stream( method, @@ -378,9 +397,7 @@ def _do_fetch( if headers: req_headers.update(headers) - params = None - if query: - params = {k: v for k, v in query.items() if v is not None} + params = _encode_query(query) return self._client.request( method, @@ -537,7 +554,7 @@ def stream_sse_sync( req_headers["Authorization"] = f"Bearer {token}" if headers: req_headers.update(headers) - params = {k: v for k, v in (query or {}).items() if v is not None} or None + params = _encode_query(query) with self._client.stream( method, diff --git a/tests/test_http_client.py b/tests/test_http_client.py index 602898b..ee6183d 100644 --- a/tests/test_http_client.py +++ b/tests/test_http_client.py @@ -856,3 +856,90 @@ def test_stream_sse_sync_raises_apierror_on_non_2xx(): with patch.object(client._client, "stream", return_value=_FakeSyncStream(resp)): with pytest.raises(ApiError): list(client.stream_sse_sync("/api/v1/x/stream", method="POST", body={})) + + +# --- query encoding contract ------------------------------------------------- +# +# The platform reads query parameters through Plug, whose parser keeps only the +# last value of a repeated bare key. A multi-value filter sent as +# `source=a&source=b` therefore silently narrows to one value server-side. These +# tests pin the wire bytes rather than the dict handed to httpx, because the +# defect only becomes visible after httpx encodes. + + +async def test_async_list_query_params_encode_as_bracket_suffixed_repeats(): + seen: list[httpx.Request] = [] + + async def handler(request: httpx.Request) -> httpx.Response: + seen.append(request) + return httpx.Response(200, json={}) + + client = HttpClient(base_url="https://api.test") + client._client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + + await client.request( + "/api/v1/knowledge_documents", + query={"source": ["cso_a", "cso_b", "cso_c"], "page_size": 100}, + ) + + assert seen[0].url.query == ( + b"source%5B%5D=cso_a&source%5B%5D=cso_b&source%5B%5D=cso_c&page_size=100" + ) + # Every element survives the round trip as a distinct value under one key. + assert seen[0].url.params.get_list("source[]") == ["cso_a", "cso_b", "cso_c"] + + +async def test_async_scalar_query_params_are_unchanged_and_none_is_dropped(): + seen: list[httpx.Request] = [] + + async def handler(request: httpx.Request) -> httpx.Response: + seen.append(request) + return httpx.Response(200, json={}) + + client = HttpClient(base_url="https://api.test") + client._client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + + await client.request( + "/api/v1/knowledge_documents", + query={"q": "notes", "page_size": 25, "agent": None}, + ) + + assert seen[0].url.query == b"q=notes&page_size=25" + + +def test_sync_list_query_params_encode_as_bracket_suffixed_repeats(): + seen: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + seen.append(request) + return httpx.Response(200, json={}) + + client = SyncHttpClient(base_url="https://api.test") + client._client = httpx.Client(transport=httpx.MockTransport(handler)) + + client.request( + "/api/v1/knowledge_documents", + query={"source": ["cso_a", "cso_b", "cso_c"], "page_size": 100}, + ) + + assert seen[0].url.query == ( + b"source%5B%5D=cso_a&source%5B%5D=cso_b&source%5B%5D=cso_c&page_size=100" + ) + + +def test_sync_scalar_query_params_are_unchanged_and_none_is_dropped(): + seen: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + seen.append(request) + return httpx.Response(200, json={}) + + client = SyncHttpClient(base_url="https://api.test") + client._client = httpx.Client(transport=httpx.MockTransport(handler)) + + client.request( + "/api/v1/knowledge_documents", + query={"q": "notes", "page_size": 25, "agent": None}, + ) + + assert seen[0].url.query == b"q=notes&page_size=25"