From 4228bf24462bda9fd005bc33bd83e54230a2b9a1 Mon Sep 17 00:00:00 2001 From: Kim Burgaard Date: Mon, 5 Oct 2026 07:06:38 -0700 Subject: [PATCH] Added optional timeout parameter --- .github/copilot-instructions.md | 4 +- CHANGELOG.md | 9 ++ README.md | 14 +- seclai/seclai.py | 228 +++++++++++----------------- tests/test_auth_and_headers.py | 18 --- tests/test_convenience_methods.py | 68 +++++---- tests/test_error_and_guard_paths.py | 15 +- tests/test_gated_lists.py | 45 +++--- tests/test_generated_bridge.py | 181 +++++++++++++++++++++- 9 files changed, 337 insertions(+), 245 deletions(-) diff --git a/.github/copilot-instructions.md b/.github/copilot-instructions.md index 7bab976..faf3151 100644 --- a/.github/copilot-instructions.md +++ b/.github/copilot-instructions.md @@ -54,5 +54,5 @@ Individual commands: - Auth modes: `api_key`, `bearer_static`, `bearer_provider`, `sso`. - `_build_default_headers()` only sets static auth; dynamic modes (`bearer_provider`, `sso`) are resolved per-request in `_merge_request_headers` / `_merge_request_headers_async`. -- Typed wrapper methods use `_sync_generated_client()` / `_async_generated_client()` which return a `GeneratedClient` with its own internal `httpx.Client` (separate from the SDK's `self._client`). -- In tests, to mock typed methods you must wire the mock transport into the generated client via `gc.set_httpx_client(httpx.Client(transport=transport, base_url=..., headers=dict(gc._headers)))`. +- Typed wrapper methods pass `self._generated`, a `GeneratedClient` that sends through the SDK's `self._client` (`_SyncSender` / `_AsyncSender`), so one httpx client serves every method. +- In tests, mock any method by passing `http_client=httpx.Client(transport=transport, base_url=...)`. diff --git a/CHANGELOG.md b/CHANGELOG.md index 5625a8e..0c441d4 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,5 +1,13 @@ # Changelog +## [1.7.2] - 2026-10-05 + +### Fixed + +- Apply a `timeout` passed to `Seclai(...)` or `AsyncSeclai(...)` to `run_agent()`, `list_agent_runs()`, `get_agent_run()`, `delete_agent_run()`, `list_sources()`, `get_content_detail()`, `delete_content()`, `list_content_embeddings()`, `upload_file_to_source()` and `upload_file_to_content()`. They sent through a second HTTP client that had no timeout, so the value was ignored and a stalled connection never returned. Without a `timeout` argument they still wait without limit: the 30-second default does not apply to them, so an upload that worked before still does +- Send those ten methods through a supplied `http_client`. They ignored it, so its transport, proxy, certificates, default headers and `base_url` did not apply to them, and they went to `SECLAI_API_URL` whatever its `base_url` said. Its timeout still does not apply to them. A supplied client with no `base_url` still reaches the Seclai API from these methods, as it did before +- Apply the unknown-version guard on each of those ten calls to a `Seclai-Version` in a supplied `http_client`'s default headers, which they now send. A value set on the client after `Seclai(...)` was built, and that this release was not built against, raises `SeclaiConfigurationError` unless `allow_unknown_api_version=True` + ## [1.7.1] - 2026-10-05 ### Changed @@ -209,6 +217,7 @@ _Stable release. Packaging, CI, and documentation deployment only; no API change _Initial release._ +[1.7.2]: https://github.com/seclai/seclai-python/releases/tag/1.7.2 [1.7.1]: https://github.com/seclai/seclai-python/releases/tag/1.7.1 [1.7.0]: https://github.com/seclai/seclai-python/releases/tag/1.7.0 [1.6.0]: https://github.com/seclai/seclai-python/releases/tag/1.6.0 diff --git a/README.md b/README.md index 1f2243e..1523412 100644 --- a/README.md +++ b/README.md @@ -79,6 +79,13 @@ asyncio.run(main()) | `default_headers` | — | `None` | | `http_client` | — | `None` (auto-created `httpx.Client`) | +Ten methods wait without limit unless you pass `timeout` yourself: `run_agent()`, +`list_agent_runs()`, `get_agent_run()`, `delete_agent_run()`, `list_sources()`, +`get_content_detail()`, `delete_content()`, `list_content_embeddings()`, +`upload_file_to_source()` and `upload_file_to_content()`. The 30-second default +does not apply to them, and neither does the timeout of an `http_client` you +supply. Requests from every other method use a supplied client's own timeout. + Set `SECLAI_API_URL` to point at a different API host (e.g., staging): ```bash @@ -178,12 +185,7 @@ The guard covers the header however it reaches the wire: `api_version`, `default_headers`, a per-request `headers` argument, or the default headers of an `http_client` you supply, which are checked at construction and again on each request, exactly as a value in `default_headers` is. It covers nothing -else. The typed methods that go through the generated client — `run_agent()`, -`list_agent_runs()`, `get_agent_run()`, `delete_agent_run()`, `list_sources()`, -`get_content_detail()`, `delete_content()`, `list_content_embeddings()`, -`upload_file_to_source()` and `upload_file_to_content()` — do not use a supplied -`http_client` at all, so nothing it carries reaches them. An account pinned -server-side can still be +else. An account pinned server-side can still be newer than this release — `get_api_version()` reports the `effective_version` the request resolved to, and comparing it against `LATEST_API_VERSION` is how you detect the gap. diff --git a/seclai/seclai.py b/seclai/seclai.py index d32a751..3c2b425 100644 --- a/seclai/seclai.py +++ b/seclai/seclai.py @@ -62,6 +62,7 @@ from seclai._generated.models.file_upload_response import FileUploadResponse from seclai._generated.models.http_validation_error import HTTPValidationError from seclai._generated.models.source_list_response import SourceListResponse +from seclai._generated.types import UNSET, Unset from seclai._generated.types import Response as OpenAPIResponse from seclai.auth import ( AuthState, @@ -415,6 +416,8 @@ class ClientOptions: default_headers: Mapping[str, str] api_version: str | None = None allow_unknown_api_version: bool = False + #: The timeout the caller passed, or ``None`` if they left the default. + explicit_timeout: float | None = None def _merge_headers( @@ -580,18 +583,49 @@ def _raise_for_status(response: httpx.Response) -> None: ) -def _raise_on_error_status(response: httpx.Response) -> None: - """httpx response hook: raise for an error status before its body is decoded.""" - if response.status_code >= 400: - response.read() +def _send_url(client: httpx.Client | httpx.AsyncClient, path: str) -> str: + """Return ``path`` for a client with a base URL, else the absolute API URL.""" + return path if str(client.base_url) else f"{SECLAI_API_URL.rstrip('/')}{path}" + + +class _SyncSender: + """What the generated client sends through: the client every method uses.""" + + def __init__(self, options: ClientOptions, client: httpx.Client) -> None: + self._options = options + self._client = client + + def request(self, method: str, url: str, **kwargs: Any) -> httpx.Response: + kwargs["timeout"] = self._options.explicit_timeout + kwargs["headers"] = _merge_request_headers( + options=self._options, + client_headers=self._client.headers, + request_headers=kwargs.get("headers"), + ) + response = self._client.request(method, _send_url(self._client, url), **kwargs) _raise_for_status(response) + return response + +class _AsyncSender: + """Async twin of :class:`_SyncSender`.""" -async def _raise_on_error_status_async(response: httpx.Response) -> None: - """Async twin of :func:`_raise_on_error_status`.""" - if response.status_code >= 400: - await response.aread() + def __init__(self, options: ClientOptions, client: httpx.AsyncClient) -> None: + self._options = options + self._client = client + + async def request(self, method: str, url: str, **kwargs: Any) -> httpx.Response: + kwargs["timeout"] = self._options.explicit_timeout + kwargs["headers"] = await _merge_request_headers_async( + options=self._options, + client_headers=self._client.headers, + request_headers=kwargs.get("headers"), + ) + response = await self._client.request( + method, _send_url(self._client, url), **kwargs + ) _raise_for_status(response) + return response def _validate_request_version( @@ -638,18 +672,6 @@ def _validate_client_version( ) from e -def _with_auth_headers( - client: GeneratedClient, auth_headers: Mapping[str, str] -) -> GeneratedClient: - """Add headers as ``with_headers`` does, replacing one whatever its spelling.""" - evolved = client.with_headers(dict(auth_headers)) - replaced = {name.lower() for name in auth_headers} - for name in [k for k in evolved._headers if k not in auth_headers]: - if name.lower() in replaced: - del evolved._headers[name] - return evolved - - class _SeclaiBase: """Shared implementation for Seclai sync/async clients. @@ -662,7 +684,7 @@ def __init__( *, api_key: str | None, access_token: str | Callable[[], str | Awaitable[str]] | None, - timeout: float, + timeout: float | Unset, api_key_header: str, default_headers: Mapping[str, str] | None, profile: str | None, @@ -756,16 +778,14 @@ def __init__( self._options = ClientOptions( auth_state=auth_state, - timeout=timeout, + timeout=30.0 if isinstance(timeout, Unset) else timeout, + explicit_timeout=None if isinstance(timeout, Unset) else timeout, api_key_header=api_key_header, default_headers=frozen_headers, api_version=api_version, allow_unknown_api_version=allow_unknown_api_version, ) - self._generated_client_instance: GeneratedClient | None = None - self._owns_generated_client = False - @property def api_key(self) -> str | None: """Return the resolved API key used for authentication, or ``None`` for bearer auth.""" @@ -783,82 +803,6 @@ def _default_headers(self) -> dict[str, str]: api_version=self._options.api_version, ) - def _generated_client(self) -> GeneratedClient: - """Return a cached generated OpenAPI client configured with this client's auth. - - Used by the wrapper methods in this file. Once one of them has run, the - client's httpx client raises `SeclaiAPIStatusError` for a 4xx/5xx rather - than returning it to a generated request helper. - """ - if self._generated_client_instance is None: - self._generated_client_instance = GeneratedClient( - base_url=SECLAI_API_URL, - headers=self._default_headers(), - ) - self._owns_generated_client = True - return self._generated_client_instance - - def _sync_generated_client(self) -> GeneratedClient: - """Return the generated client, ready to send (sync). - - Resolves fresh auth headers for the ``bearer_provider`` and ``sso`` modes, - and makes its httpx client raise through :func:`_raise_for_status`. - """ - gc = self._generated_client() - if self._options.auth_state.mode in ("bearer_provider", "sso"): - try: - auth_headers = resolve_auth_headers_sync(self._options.auth_state) - except (RuntimeError, TypeError) as exc: - raise SeclaiConfigurationError(str(exc)) from exc - except Exception as exc: - raise SeclaiConfigurationError( - f"Auth resolution failed: {exc}" - ) from exc - # with_headers() mutates existing httpx clients in-place and - # returns an evolved copy with updated _headers. We keep the - # original instance (preserving any user-set httpx client) but - # store the evolved copy so get_httpx_client() will use the - # refreshed headers when it lazily creates an httpx.Client. - evolved = _with_auth_headers(gc, auth_headers) - self._generated_client_instance = evolved - # Re-attach the existing httpx client to the evolved instance - # so pre-configured transports (e.g. in tests) aren't lost. - if gc._client is not None: - evolved.set_httpx_client(gc._client) - client = self._generated_client() - hooks = client.get_httpx_client().event_hooks["response"] - if _raise_on_error_status not in hooks: - hooks.append(_raise_on_error_status) - return client - - async def _async_generated_client(self) -> GeneratedClient: - """Return the generated client, ready to send (async). - - Resolves fresh auth headers for the ``bearer_provider`` and ``sso`` modes, - and makes its httpx client raise through :func:`_raise_for_status`. - """ - gc = self._generated_client() - if self._options.auth_state.mode in ("bearer_provider", "sso"): - try: - auth_headers = await resolve_auth_headers_async( - self._options.auth_state - ) - except (RuntimeError, TypeError) as exc: - raise SeclaiConfigurationError(str(exc)) from exc - except Exception as exc: - raise SeclaiConfigurationError( - f"Auth resolution failed: {exc}" - ) from exc - evolved = _with_auth_headers(gc, auth_headers) - self._generated_client_instance = evolved - if gc._async_client is not None: - evolved.set_async_httpx_client(gc._async_client) - client = self._generated_client() - hooks = client.get_async_httpx_client().event_hooks["response"] - if _raise_on_error_status_async not in hooks: - hooks.append(_raise_on_error_status_async) - return client - def _build_url(self, path: str) -> str: """Build an absolute URL string from a request path. @@ -937,7 +881,7 @@ def __init__( *, api_key: str | None = None, access_token: str | Callable[[], str] | None = None, - timeout: float = 30.0, + timeout: float | Unset = UNSET, api_key_header: str = "x-api-key", default_headers: Mapping[str, str] | None = None, http_client: httpx.Client | None = None, @@ -961,14 +905,14 @@ def __init__( api_key: API key used for authentication. If omitted, ``SECLAI_API_KEY`` is used. access_token: Static bearer token string or a callable returning one. Mutually exclusive with ``api_key``. - timeout: Request timeout (seconds). + timeout: Request timeout (seconds), 30 by default; a supplied + http_client keeps its own. The typed methods the README lists + wait without limit unless you pass one. api_key_header: Header name to use for the API key. default_headers: Extra headers to include on every request. http_client: Optional pre-configured ``httpx.Client`` to use. A ``Seclai-Version`` among its default headers is checked as one in ``default_headers`` is, at construction and on each request. - The typed methods that go through the generated client do not - use this client. profile: SSO profile name from ``~/.seclai/config``. config_dir: Override the config directory path. auto_refresh: Auto-refresh expired SSO tokens. Defaults to ``True``. @@ -999,6 +943,9 @@ def __init__( timeout=self._options.timeout, headers=self._default_headers(), ) + self._generated = GeneratedClient(base_url=SECLAI_API_URL).set_httpx_client( + cast(httpx.Client, _SyncSender(self._options, self._client)) + ) self._owns_client = http_client is None if http_client is not None: _validate_client_version( @@ -1013,8 +960,6 @@ def close(self) -> None: """ if self._owns_client: self._client.close() - if self._owns_generated_client and self._generated_client_instance is not None: - self._generated_client_instance.get_httpx_client().close() def __enter__(self) -> Self: """Enter a context manager and return self.""" @@ -1097,9 +1042,7 @@ def run_agent(self, agent_id: str, body: AgentRunRequest) -> AgentRunResponse: ) path = f"/agents/{agent_id}/runs" - response = sync_detailed( - agent_id=agent_id, client=self._sync_generated_client(), body=body - ) + response = sync_detailed(agent_id=agent_id, client=self._generated, body=body) self._raise_for_openapi_response( method="POST", path=path, @@ -1269,7 +1212,7 @@ def list_agent_runs( path = f"/agents/{agent_id}/runs" response = sync_detailed( agent_id=agent_id, - client=self._sync_generated_client(), + client=self._generated, page=page, limit=limit, ) @@ -1344,7 +1287,7 @@ def get_agent_run( path = f"/agents/runs/{run_id}" response = sync_detailed( run_id=run_id, - client=self._sync_generated_client(), + client=self._generated, include_step_outputs=include_step_outputs, ) self._raise_for_openapi_response( @@ -1408,7 +1351,7 @@ def delete_agent_run(self, *args: str) -> AgentRunResponse: ) path = f"/agents/runs/{run_id}" - response = sync_detailed(run_id=run_id, client=self._sync_generated_client()) + response = sync_detailed(run_id=run_id, client=self._generated) self._raise_for_openapi_response( method="DELETE", path=path, @@ -1465,7 +1408,7 @@ def get_content_detail( path = f"/contents/{source_connection_content_version}" response = sync_detailed( source_connection_content_version=source_connection_content_version, - client=self._sync_generated_client(), + client=self._generated, start=start, end=end, ) @@ -1514,7 +1457,7 @@ def delete_content(self, source_connection_content_version: str) -> None: path = f"/contents/{source_connection_content_version}" response = sync_detailed( source_connection_content_version=source_connection_content_version, - client=self._sync_generated_client(), + client=self._generated, ) self._raise_for_openapi_response( method="DELETE", @@ -1554,7 +1497,7 @@ def list_content_embeddings( path = f"/contents/{source_connection_content_version}/embeddings" response = sync_detailed( source_connection_content_version=source_connection_content_version, - client=self._sync_generated_client(), + client=self._generated, page=page, limit=limit, ) @@ -1617,7 +1560,7 @@ def list_sources( path = "/sources" response = sync_detailed( - client=self._sync_generated_client(), + client=self._generated, page=page, limit=limit, sort=sort, @@ -1774,7 +1717,7 @@ def upload_file_to_source( # Note: openapi-python-client currently struggles with Seclai's spec for this endpoint # due to duplicate schema names, so we send the multipart request directly and parse # into our SDK model types. - http = self._sync_generated_client().get_httpx_client() + http = self._generated.get_httpx_client() raw = http.request( "POST", endpoint_path, @@ -1913,7 +1856,7 @@ def upload_file_to_content( endpoint_path = f"/contents/{source_connection_content_version}/upload" response = sync_detailed( source_connection_content_version=source_connection_content_version, - client=self._sync_generated_client(), + client=self._generated, body=body, ) self._raise_for_openapi_response( @@ -2212,10 +2155,8 @@ def cancel_agent_run(self, run_id: str) -> dict[str, Any]: ``POST .../cancel`` route, and no operation that deletes a run. Rejected when the run has already reached a terminal state. - Hits the same endpoint as :meth:`delete_agent_run`, but through - :meth:`request` rather than the generated client, so a caller-supplied - ``http_client`` applies. Both now surface a 422 as - :class:`SeclaiAPIValidationError`. + Hits the same endpoint as :meth:`delete_agent_run`. Both surface a 422 + as :class:`SeclaiAPIValidationError`. Args: run_id: Run identifier. @@ -5243,7 +5184,7 @@ def __init__( *, api_key: str | None = None, access_token: str | Callable[[], str | Awaitable[str]] | None = None, - timeout: float = 30.0, + timeout: float | Unset = UNSET, api_key_header: str = "x-api-key", default_headers: Mapping[str, str] | None = None, http_client: httpx.AsyncClient | None = None, @@ -5267,14 +5208,14 @@ def __init__( api_key: API key used for authentication. If omitted, ``SECLAI_API_KEY`` is used. access_token: Static bearer token string or a callable returning one (sync or async). Mutually exclusive with ``api_key``. - timeout: Request timeout (seconds). + timeout: Request timeout (seconds), 30 by default; a supplied + http_client keeps its own. The typed methods the README lists + wait without limit unless you pass one. api_key_header: Header name to use for the API key. default_headers: Extra headers to include on every request. http_client: Optional pre-configured ``httpx.AsyncClient`` to use. A ``Seclai-Version`` among its default headers is checked as one in ``default_headers`` is, at construction and on each request. - The typed methods that go through the generated client do not - use this client. profile: SSO profile name from ``~/.seclai/config``. config_dir: Override the config directory path. auto_refresh: Auto-refresh expired SSO tokens. Defaults to ``True``. @@ -5305,6 +5246,11 @@ def __init__( timeout=self._options.timeout, headers=self._default_headers(), ) + self._generated = GeneratedClient( + base_url=SECLAI_API_URL + ).set_async_httpx_client( + cast(httpx.AsyncClient, _AsyncSender(self._options, self._client)) + ) self._owns_client = http_client is None if http_client is not None: _validate_client_version( @@ -5319,8 +5265,6 @@ async def aclose(self) -> None: """ if self._owns_client: await self._client.aclose() - if self._owns_generated_client and self._generated_client_instance is not None: - await self._generated_client_instance.get_async_httpx_client().aclose() async def __aenter__(self) -> Self: """Enter an async context manager and return self.""" @@ -5404,7 +5348,7 @@ async def run_agent(self, agent_id: str, body: AgentRunRequest) -> AgentRunRespo path = f"/agents/{agent_id}/runs" response = await asyncio_detailed( - agent_id=agent_id, client=(await self._async_generated_client()), body=body + agent_id=agent_id, client=self._generated, body=body ) self._raise_for_openapi_response( method="POST", @@ -5560,7 +5504,7 @@ async def list_agent_runs( path = f"/agents/{agent_id}/runs" response = await asyncio_detailed( agent_id=agent_id, - client=(await self._async_generated_client()), + client=self._generated, page=page, limit=limit, ) @@ -5635,7 +5579,7 @@ async def get_agent_run( path = f"/agents/runs/{run_id}" response = await asyncio_detailed( run_id=run_id, - client=(await self._async_generated_client()), + client=self._generated, include_step_outputs=include_step_outputs, ) self._raise_for_openapi_response( @@ -5703,7 +5647,7 @@ async def delete_agent_run(self, *args: str) -> AgentRunResponse: path = f"/agents/runs/{run_id}" response = await asyncio_detailed( run_id=run_id, - client=(await self._async_generated_client()), + client=self._generated, ) self._raise_for_openapi_response( method="DELETE", @@ -5761,7 +5705,7 @@ async def get_content_detail( path = f"/contents/{source_connection_content_version}" response = await asyncio_detailed( source_connection_content_version=source_connection_content_version, - client=(await self._async_generated_client()), + client=self._generated, start=start, end=end, ) @@ -5810,7 +5754,7 @@ async def delete_content(self, source_connection_content_version: str) -> None: path = f"/contents/{source_connection_content_version}" response = await asyncio_detailed( source_connection_content_version=source_connection_content_version, - client=(await self._async_generated_client()), + client=self._generated, ) self._raise_for_openapi_response( method="DELETE", @@ -5850,7 +5794,7 @@ async def list_content_embeddings( path = f"/contents/{source_connection_content_version}/embeddings" response = await asyncio_detailed( source_connection_content_version=source_connection_content_version, - client=(await self._async_generated_client()), + client=self._generated, page=page, limit=limit, ) @@ -5913,7 +5857,7 @@ async def list_sources( path = "/sources" response = await asyncio_detailed( - client=(await self._async_generated_client()), + client=self._generated, page=page, limit=limit, sort=sort, @@ -6067,7 +6011,7 @@ async def upload_file_to_source( endpoint_path = f"/sources/{source_connection_id}/upload" - http = (await self._async_generated_client()).get_async_httpx_client() + http = self._generated.get_async_httpx_client() raw = await http.request( "POST", endpoint_path, @@ -6206,7 +6150,7 @@ async def upload_file_to_content( endpoint_path = f"/contents/{source_connection_content_version}/upload" response = await asyncio_detailed( source_connection_content_version=source_connection_content_version, - client=(await self._async_generated_client()), + client=self._generated, body=body, ) self._raise_for_openapi_response( @@ -6513,10 +6457,8 @@ async def cancel_agent_run(self, run_id: str) -> dict[str, Any]: ``POST .../cancel`` route, and no operation that deletes a run. Rejected when the run has already reached a terminal state. - Hits the same endpoint as :meth:`delete_agent_run`, but through - :meth:`request` rather than the generated client, so a caller-supplied - ``http_client`` applies. Both now surface a 422 as - :class:`SeclaiAPIValidationError`. + Hits the same endpoint as :meth:`delete_agent_run`. Both surface a 422 + as :class:`SeclaiAPIValidationError`. Args: run_id: Run identifier. diff --git a/tests/test_auth_and_headers.py b/tests/test_auth_and_headers.py index e5efb21..9d45042 100644 --- a/tests/test_auth_and_headers.py +++ b/tests/test_auth_and_headers.py @@ -160,15 +160,6 @@ def handler(request: httpx.Request) -> httpx.Response: transport = httpx.MockTransport(handler) http_client = httpx.Client(transport=transport, base_url="https://example.invalid") client = Seclai(access_token="typed-jwt", http_client=http_client) - # Also wire mock transport into the generated client used by typed methods - gc = client._generated_client() - gc.set_httpx_client( - httpx.Client( - transport=transport, - base_url="https://example.invalid", - headers=dict(gc._headers), - ) - ) result = client.list_sources() assert result.data == [] @@ -204,15 +195,6 @@ def handler(request: httpx.Request) -> httpx.Response: transport = httpx.MockTransport(handler) http_client = httpx.Client(transport=transport, base_url="https://example.invalid") client = Seclai(access_token=provider, http_client=http_client) - # Also wire mock transport into the generated client used by typed methods - gc = client._generated_client() - gc.set_httpx_client( - httpx.Client( - transport=transport, - base_url="https://example.invalid", - headers=dict(gc._headers), - ) - ) client.list_sources() client.list_sources() assert call_count == 2 diff --git a/tests/test_convenience_methods.py b/tests/test_convenience_methods.py index 773dba8..212edba 100644 --- a/tests/test_convenience_methods.py +++ b/tests/test_convenience_methods.py @@ -25,7 +25,7 @@ def _extract_multipart_field( return None -def test_convenience_delete_content_uses_generated_client_transport() -> None: +def test_convenience_delete_content_uses_supplied_http_client() -> None: seen: dict[str, str] = {} def handler(request: httpx.Request) -> httpx.Response: @@ -35,17 +35,18 @@ def handler(request: httpx.Request) -> httpx.Response: transport = httpx.MockTransport(handler) - client = Seclai(api_key="test") - gen = client._generated_client() - gen.set_httpx_client( - httpx.Client(base_url="https://api.seclai.com", transport=transport) + client = Seclai( + api_key="test", + http_client=httpx.Client( + base_url="https://api.seclai.com", transport=transport + ), ) client.delete_content("sc_cv_123") assert seen == {"method": "DELETE", "path": "/contents/sc_cv_123"} -def test_convenience_upload_file_to_source_sends_metadata_and_uses_generated_client_transport() -> ( +def test_convenience_upload_file_to_source_sends_metadata_and_uses_supplied_http_client() -> ( None ): seen: dict[str, object] = {} @@ -73,10 +74,11 @@ def handler(request: httpx.Request) -> httpx.Response: transport = httpx.MockTransport(handler) - client = Seclai(api_key="test") - gen = client._generated_client() - gen.set_httpx_client( - httpx.Client(base_url="https://api.seclai.com", transport=transport) + client = Seclai( + api_key="test", + http_client=httpx.Client( + base_url="https://api.seclai.com", transport=transport + ), ) resp = client.upload_file_to_source( @@ -94,7 +96,7 @@ def handler(request: httpx.Request) -> httpx.Response: } -def test_convenience_upload_file_to_content_sends_metadata_and_uses_generated_client_transport() -> ( +def test_convenience_upload_file_to_content_sends_metadata_and_uses_supplied_http_client() -> ( None ): seen: dict[str, object] = {} @@ -122,10 +124,11 @@ def handler(request: httpx.Request) -> httpx.Response: transport = httpx.MockTransport(handler) - client = Seclai(api_key="test") - gen = client._generated_client() - gen.set_httpx_client( - httpx.Client(base_url="https://api.seclai.com", transport=transport) + client = Seclai( + api_key="test", + http_client=httpx.Client( + base_url="https://api.seclai.com", transport=transport + ), ) resp = client.upload_file_to_content( @@ -144,9 +147,7 @@ def handler(request: httpx.Request) -> httpx.Response: @pytest.mark.asyncio -async def test_async_convenience_delete_content_uses_generated_client_transport() -> ( - None -): +async def test_async_convenience_delete_content_uses_supplied_http_client() -> None: seen: dict[str, str] = {} async def handler(request: httpx.Request) -> httpx.Response: @@ -156,10 +157,11 @@ async def handler(request: httpx.Request) -> httpx.Response: transport = httpx.MockTransport(handler) - client = AsyncSeclai(api_key="test") - gen = client._generated_client() - gen.set_async_httpx_client( - httpx.AsyncClient(base_url="https://api.seclai.com", transport=transport) + client = AsyncSeclai( + api_key="test", + http_client=httpx.AsyncClient( + base_url="https://api.seclai.com", transport=transport + ), ) await client.delete_content("sc_cv_123") @@ -167,7 +169,7 @@ async def handler(request: httpx.Request) -> httpx.Response: @pytest.mark.asyncio -async def test_async_convenience_upload_file_to_source_sends_metadata_and_uses_generated_client_transport() -> ( +async def test_async_convenience_upload_file_to_source_sends_metadata_and_uses_supplied_http_client() -> ( None ): seen: dict[str, object] = {} @@ -195,10 +197,11 @@ async def handler(request: httpx.Request) -> httpx.Response: transport = httpx.MockTransport(handler) - client = AsyncSeclai(api_key="test") - gen = client._generated_client() - gen.set_async_httpx_client( - httpx.AsyncClient(base_url="https://api.seclai.com", transport=transport) + client = AsyncSeclai( + api_key="test", + http_client=httpx.AsyncClient( + base_url="https://api.seclai.com", transport=transport + ), ) resp = await client.upload_file_to_source( @@ -217,7 +220,7 @@ async def handler(request: httpx.Request) -> httpx.Response: @pytest.mark.asyncio -async def test_async_convenience_upload_file_to_content_sends_metadata_and_uses_generated_client_transport() -> ( +async def test_async_convenience_upload_file_to_content_sends_metadata_and_uses_supplied_http_client() -> ( None ): seen: dict[str, object] = {} @@ -245,10 +248,11 @@ async def handler(request: httpx.Request) -> httpx.Response: transport = httpx.MockTransport(handler) - client = AsyncSeclai(api_key="test") - gen = client._generated_client() - gen.set_async_httpx_client( - httpx.AsyncClient(base_url="https://api.seclai.com", transport=transport) + client = AsyncSeclai( + api_key="test", + http_client=httpx.AsyncClient( + base_url="https://api.seclai.com", transport=transport + ), ) resp = await client.upload_file_to_content( diff --git a/tests/test_error_and_guard_paths.py b/tests/test_error_and_guard_paths.py index 1aa5cba..c2b9986 100644 --- a/tests/test_error_and_guard_paths.py +++ b/tests/test_error_and_guard_paths.py @@ -17,12 +17,7 @@ def _typed_client(handler: Any) -> Seclai: - """A client whose typed (generated-client) methods answer from ``handler``.""" - client = Seclai(api_key="k") - client._generated_client().set_httpx_client( - httpx.Client(transport=httpx.MockTransport(handler), base_url=_BASE_URL) - ) - return client + return _request_client(handler) def _request_client(handler: Any, **options: Any) -> Seclai: @@ -93,14 +88,14 @@ def test_a_success_still_decodes(self) -> None: @pytest.mark.asyncio async def test_async_503_html_is_a_status_error(self) -> None: - client = AsyncSeclai(api_key="k") - client._generated_client().set_async_httpx_client( - httpx.AsyncClient( + client = AsyncSeclai( + api_key="k", + http_client=httpx.AsyncClient( transport=httpx.MockTransport( lambda req: httpx.Response(503, text="down") ), base_url=_BASE_URL, - ) + ), ) with pytest.raises(seclai.SeclaiAPIStatusError) as exc: await client.list_sources() diff --git a/tests/test_gated_lists.py b/tests/test_gated_lists.py index 1f951bf..9e47c22 100644 --- a/tests/test_gated_lists.py +++ b/tests/test_gated_lists.py @@ -658,47 +658,40 @@ async def resolve(state: Any) -> str: @pytest.mark.parametrize("mode", ["bearer_provider", "sso"]) -class TestGeneratedClientCredentials: +class TestTypedMethodCredentials: """One credential per name on the typed-method path, from the first call.""" EXPECTED = [("authorization", "Bearer TOKEN"), ("x-account-id", "ACCT")] + def _transport(self, seen: list[list[tuple[str, str]]]) -> httpx.MockTransport: + def handler(request: httpx.Request) -> httpx.Response: + seen.append(_credentials(request)) + return httpx.Response(200, json={"data": [], "pagination": PAGED}) + + return httpx.MockTransport(handler) + def test_sync(self, mode: str, dynamic_modes: dict[str, dict[str, Any]]) -> None: + seen: list[list[tuple[str, str]]] = [] client = Seclai(default_headers=SPELLED, **dynamic_modes[mode]) assert client._options.auth_state.mode == mode - for _call in ("first", "second"): - http = client._sync_generated_client().get_httpx_client() - assert _credentials(http.build_request("GET", "/x")) == self.EXPECTED + # Swap only the transport: the headers are the ones the SDK built. + client._client._transport = self._transport(seen) + client.list_sources() + client.list_sources() + assert seen == [self.EXPECTED, self.EXPECTED] client.close() async def test_async( self, mode: str, dynamic_modes: dict[str, dict[str, Any]] ) -> None: + seen: list[list[tuple[str, str]]] = [] client = AsyncSeclai(default_headers=SPELLED, **dynamic_modes[mode]) assert client._options.auth_state.mode == mode - for _call in ("first", "second"): - generated = await client._async_generated_client() - http = generated.get_async_httpx_client() - assert _credentials(http.build_request("GET", "/x")) == self.EXPECTED - await client.aclose() - - def test_typed_method_sends_one_of_each( - self, mode: str, dynamic_modes: dict[str, dict[str, Any]] - ) -> None: - seen: list[list[tuple[str, str]]] = [] - - def handler(request: httpx.Request) -> httpx.Response: - seen.append(_credentials(request)) - return httpx.Response(200, json={"data": [], "pagination": PAGED}) - - client = Seclai(default_headers=SPELLED, **dynamic_modes[mode]) - generated = client._sync_generated_client() - # Swap only the transport: the headers are the ones the SDK built. - generated.get_httpx_client()._transport = httpx.MockTransport(handler) - client.list_sources() - client.list_sources() + client._client._transport = self._transport(seen) + await client.list_sources() + await client.list_sources() assert seen == [self.EXPECTED, self.EXPECTED] - client.close() + await client.aclose() UNKNOWN = "2099-01-01" diff --git a/tests/test_generated_bridge.py b/tests/test_generated_bridge.py index b3a223d..925a8d8 100644 --- a/tests/test_generated_bridge.py +++ b/tests/test_generated_bridge.py @@ -1,17 +1,182 @@ +"""The typed methods send through the same httpx client as every other method.""" + +from typing import Any + +import httpx import pytest -from seclai import Seclai +from seclai import AsyncSeclai, Seclai from seclai import seclai as seclai_module +_PAGE = { + "data": [], + "pagination": { + "page": 1, + "limit": 20, + "total": 0, + "pages": 0, + "has_next": False, + "has_prev": False, + }, +} + + +def _recorder(seen: list[httpx.Request]) -> httpx.MockTransport: + def handler(request: httpx.Request) -> httpx.Response: + seen.append(request) + return httpx.Response(200, json=_PAGE) + + return httpx.MockTransport(handler) + + +def test_typed_method_uses_the_configured_timeout() -> None: + seen: list[httpx.Request] = [] + client = Seclai(api_key="k", timeout=7.0) + client._client._transport = _recorder(seen) + + client.list_sources() + + assert set(seen[0].extensions["timeout"].values()) == {7.0} + + +def test_typed_method_sends_through_a_supplied_client() -> None: + seen: list[httpx.Request] = [] + http_client = httpx.Client( + transport=_recorder(seen), + base_url="https://proxy.invalid", + timeout=3.0, + headers={"x-via": "supplied"}, + ) + client = Seclai(api_key="k", http_client=http_client) + + client.list_sources() + + request = seen[0] + assert request.url.host == "proxy.invalid" + assert request.headers["x-api-key"] == "k" + assert request.headers["x-via"] == "supplied" + assert set(request.extensions["timeout"].values()) == {None} -def test_generated_client_get_httpx_client_has_api_key_header( + +def test_typed_method_reaches_the_api_through_a_client_with_no_base_url( monkeypatch: pytest.MonkeyPatch, ) -> None: - monkeypatch.setenv("SECLAI_API_KEY", "k") - monkeypatch.setattr(seclai_module, "SECLAI_API_URL", "https://example.invalid") - client = Seclai() + monkeypatch.setattr(seclai_module, "SECLAI_API_URL", "https://api.invalid/") + seen: list[httpx.Request] = [] + client = Seclai(api_key="k", http_client=httpx.Client(transport=_recorder(seen))) + + client.list_sources() + + assert str(seen[0].url).startswith("https://api.invalid/sources") + + +def test_upload_sends_through_a_supplied_client() -> None: + seen: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + seen.append(request) + return httpx.Response(503, text="down") + + http_client = httpx.Client( + transport=httpx.MockTransport(handler), base_url="https://proxy.invalid" + ) + client = Seclai(api_key="k", http_client=http_client) + + with pytest.raises(seclai_module.SeclaiAPIStatusError): + client.upload_file_to_source("sc_1", file=b"x", file_name="a.txt") + + assert seen[0].url.host == "proxy.invalid" + assert seen[0].headers["x-api-key"] == "k" + + +def test_closing_leaves_a_supplied_client_open() -> None: + http_client = httpx.Client(transport=_recorder([])) + Seclai(api_key="k", http_client=http_client).close() + assert not http_client.is_closed + + +def test_unknown_version_on_a_supplied_client_is_refused_by_a_typed_method() -> None: + http_client = httpx.Client( + transport=_recorder([]), base_url="https://proxy.invalid" + ) + client = Seclai(api_key="k", http_client=http_client) + http_client.headers["Seclai-Version"] = "2099-01-01" + + with pytest.raises(seclai_module.SeclaiConfigurationError): + client.list_sources() + + +@pytest.mark.asyncio +async def test_async_typed_method_uses_the_configured_timeout() -> None: + seen: list[httpx.Request] = [] + client = AsyncSeclai(api_key="k", timeout=7.0) + client._client._transport = _recorder(seen) + + await client.list_sources() + + assert set(seen[0].extensions["timeout"].values()) == {7.0} + + +@pytest.mark.asyncio +async def test_async_typed_method_sends_through_a_supplied_client() -> None: + seen: list[httpx.Request] = [] + http_client = httpx.AsyncClient( + transport=_recorder(seen), base_url="https://proxy.invalid", timeout=3.0 + ) + client = AsyncSeclai(api_key="k", http_client=http_client) + + await client.list_sources() + + assert seen[0].url.host == "proxy.invalid" + assert seen[0].headers["x-api-key"] == "k" + assert set(seen[0].extensions["timeout"].values()) == {None} + + +def test_typed_method_has_no_timeout_unless_one_is_passed() -> None: + seen: list[httpx.Request] = [] + client = Seclai(api_key="k") + client._client._transport = _recorder(seen) + + client.list_sources() + client.request("GET", "/sources/") + + assert set(seen[0].extensions["timeout"].values()) == {None} + assert set(seen[1].extensions["timeout"].values()) == {30.0} + + +def test_a_passed_timeout_applies_through_a_supplied_client() -> None: + seen: list[httpx.Request] = [] + http_client = httpx.Client( + transport=_recorder(seen), base_url="https://proxy.invalid", timeout=3.0 + ) + client = Seclai(api_key="k", timeout=7.0, http_client=http_client) + + client.list_sources() + + assert set(seen[0].extensions["timeout"].values()) == {7.0} + + +@pytest.mark.asyncio +async def test_async_typed_method_has_no_timeout_unless_one_is_passed() -> None: + seen: list[httpx.Request] = [] + client = AsyncSeclai(api_key="k") + client._client._transport = _recorder(seen) + + await client.list_sources() + await client.request("GET", "/sources/") + + assert set(seen[0].extensions["timeout"].values()) == {None} + assert set(seen[1].extensions["timeout"].values()) == {30.0} + + +def test_an_explicit_none_timeout_still_means_no_timeout_everywhere() -> None: + seen: list[httpx.Request] = [] + no_timeout: Any = None + client = Seclai(api_key="k", timeout=no_timeout) + client._client._transport = _recorder(seen) - gen = client._generated_client() - http = gen.get_httpx_client() + client.list_sources() + client.request("GET", "/sources/") - assert http.headers.get("x-api-key") == "k" + assert set(seen[0].extensions["timeout"].values()) == {None} + assert set(seen[1].extensions["timeout"].values()) == {None}