diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 131582e6..1ca21a1e 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -104,6 +104,19 @@ project sees an installed permit, and fails while `permit/_sync_types.pyi` is ou (see [Regenerating the sync stubs](#regenerating-the-sync-stubs)). The `mypy` pre-commit hook type-checks the SDK itself, strictly and with the pydantic plugin (see [Setup](#setup)). +### Connection reuse + +pytest-httpserver closes each connection after its response, so the tests of how the +clients keep and close their connections (`tests/test_async_session_lifecycle.py` and +`tests/test_sync_lifecycle.py`) use `tests/keepalive_server.py`, a local HTTP/1.1 server +that keeps every connection open and counts the connections it accepted and those that were +closed. The benchmark runs on it too: it times sequential `check()` calls of the async and +the blocking client, and prints how many connections each opened. + +```sh +uv run --locked python -m tests.benchmark_connection_reuse --calls 500 +``` + ### The migration skill's tests `skills/tests` checks `MIGRATION.md` and the permit-python-3-migration skill against each diff --git a/README.md b/README.md index a0d6ae33..f7def99d 100644 --- a/README.md +++ b/README.md @@ -21,6 +21,74 @@ every breaking change, who it affects and what to change. To have an AI agent su do the upgrade, use the [permit-python-3-migration skill](https://github.com/permitio/permit-python/tree/main/skills/permit-python-3-migration). +## Connections + +Both clients keep the HTTP connections they open and reuse them for their next requests, so a +request does not pay for a new connection, and a TLS handshake, each time. A client keeps one +pool of connections for the Permit API and one for the PDP, opened by its first request. + +### The async client + +`permit.Permit` keeps its pools per event loop it is used on, each opened by the first +request from that loop. + +```py +async with Permit(token="") as permit: + allowed = await permit.check("alice", "read", "document") +``` + +- `await permit.close()` closes the connections, as leaving the `async with` block does. + Calling it again does nothing more, and the client stays usable: a request sent after it + opens new connections. +- A client you never close leaves nothing open when its loop shuts down through + `asyncio.run()`, `asyncio.Runner` or anything else that shuts down the loop's async + generators before closing it: the client's connections on that loop are closed then. A + client that is garbage collected while its loop runs closes its connections on that + loop. As the interpreter exits, the client closes what is still open, so aiohttp reports + no unclosed session. +- If you drive an event loop yourself, run `await permit.close()` on it before you close it. + A loop closed with `loop.close()` alone cannot close its connections any more: they stay + open until the client's next request, from any loop, lets the garbage collector free + them, and Python reports each one with a `ResourceWarning`. Python's default warning + filters hide it, but a test suite that turns warnings into errors, such as pytest with + `filterwarnings = error`, fails on it. +- Close the client once no request is in flight: a request in flight when `close()` runs + fails. + +### The blocking client + +`permit.sync.Permit` runs its calls on an event loop in a background daemon thread of its +own, which it starts on its first call. Calls from every thread that uses the client are +handed to that thread and waited for, so they share the client's connections. + +```py +from permit.sync import Permit + +with Permit(token="") as permit: + allowed = permit.check("alice", "read", "document") +``` + +- `permit.close()` waits for the calls other threads have in flight, closes the connections + and stops the thread, as leaving the `with` block does. A call or a `close()` another + thread makes meanwhile waits for it to finish. Calling it again does nothing more, and the + client stays usable: its next call starts a new thread and opens new connections. +- A client you never close is cleaned up when it is garbage collected, or as the + interpreter exits. The thread never holds up the exit. +- Do not call the blocking client from code that runs on its own background thread, such as + a callback scheduled on its loop: such a call, and `close()`, raise `RuntimeError` rather + than wait for themselves. + +### Both clients + +- With `proxy_facts_via_pdp` on, `wait_for_sync()` yields a client that uses the connections + of the client it is called on, and on the blocking client its thread too. That client's + `close()` closes them; the yielded one's `close()` does nothing. With it off, the default, + `wait_for_sync()` logs a warning and yields the client itself, whose `close()` closes them. +- A child process made by `fork()` leaves the connections it inherits to its parent, and + opens its own; the blocking client starts a thread of its own in the child. +- The number of connections open at once is not capped, as before. An idle connection is + closed after aiohttp's keep-alive timeout of 15 seconds. + ## Groups `permit.api.groups` manages groups. A group is a resource instance, of the `group` resource diff --git a/permit/api/base.py b/permit/api/base.py index ea404d50..cb34c96e 100644 --- a/permit/api/base.py +++ b/permit/api/base.py @@ -1,9 +1,11 @@ from typing import TYPE_CHECKING, Any, TypeVar, cast, overload -import aiohttp from aiohttp import ClientTimeout +from multidict import CIMultiDict +from yarl import URL from permit.api.encoders import jsonable_encoder +from permit.utils.http_sessions import LoopSessions from permit.utils.pydantic_version import PYDANTIC_VERSION from permit.utils.sdk_logger import sdk_logger @@ -56,16 +58,100 @@ class Config: ) +# What a SimpleHttpClient's client_config may set. The other options of an aiohttp session +# cannot be set per client, since the client sends its requests through shared sessions. +_CLIENT_CONFIG_KEYS = frozenset({"base_url", "headers", "timeout"}) + + +def _session_base_url(base_url: str | URL) -> URL: + """``base_url`` as ``aiohttp.ClientSession(base_url=...)`` reads it, raising what it raises. + + Raises: + ValueError: If ``base_url`` has no scheme or host, or its path does not end with "/". + """ + if isinstance(base_url, URL): + url = base_url + else: + url = URL(base_url) + url.origin() # raises ValueError for a URL without a scheme and a host + if not url.path.endswith("/"): + msg = "base_url must have a trailing '/'" + raise ValueError(msg) + return url + + class SimpleHttpClient: - """wraps aiohttp client to reduce boilerplace.""" + """Sends requests to one endpoint and parses their JSON responses. + + The requests go through ``sessions``, which keep their connections open for the next + request. Everything else a request carries comes from this client and the call itself, + so the sessions can serve every client of an SDK client. + + Args: + client_config: Optional request settings: ``base_url``, the server a relative request + URL is resolved against, as ``aiohttp.ClientSession(base_url=...)`` resolves it; + ``headers``, sent with every request; and ``timeout``, an + ``aiohttp.ClientTimeout``. + base_url: The endpoint's path, put before the URL of every request. + timeout: The total timeout of each request in seconds, in place of + ``client_config["timeout"]``. + sessions: The sessions to send the requests through. Without them, the client has + sessions of its own. + + Raises: + TypeError: If ``client_config`` has a key other than those above. + """ def __init__( - self, client_config: dict[str, Any], base_url: str = "", timeout: int | None = None + self, + client_config: dict[str, Any], + base_url: str = "", + timeout: int | None = None, + *, + sessions: LoopSessions | None = None, ) -> None: - self._client_config = client_config + unsupported = sorted(set(client_config) - _CLIENT_CONFIG_KEYS) + if unsupported: + msg = ( + f"SimpleHttpClient does not take the client_config keys {unsupported}: " + f"it sets only {sorted(_CLIENT_CONFIG_KEYS)} on its requests." + ) + raise TypeError(msg) + self._server_url: str | URL | None = client_config.get("base_url") + self._headers: dict[str, str] | None = client_config.get("headers") + self._timeout: ClientTimeout | None = ( + ClientTimeout(total=timeout) if timeout is not None else client_config.get("timeout") + ) self._base_url = base_url - if timeout is not None: - self._client_config["timeout"] = ClientTimeout(total=timeout) + self._sessions = sessions if sessions is not None else LoopSessions() + + def _use_sessions(self, sessions: LoopSessions) -> None: + """Send the requests through ``sessions`` from now on.""" + self._sessions = sessions + + def _request_url(self, url: str) -> URL: + """``url`` resolved against the client's ``base_url``, as an aiohttp session does it. + + Raises: + ValueError: If the client's ``base_url`` is not one an aiohttp session takes. + """ + target = URL(url) + if self._server_url is None: + return target + server_url = _session_base_url(self._server_url) + return target if target.absolute else server_url.join(target) + + def _request_options(self, options: dict[str, Any]) -> dict[str, Any]: + """The client's headers and timeout, with a request's own aiohttp ``options`` over them. + + The request's options win, as they did over the options of a session of the + client's own: a header in ``options["headers"]`` replaces the client's header of + that name. + """ + headers = CIMultiDict(self._headers or {}) + headers.update(options.get("headers") or {}) + defaults = {} if self._timeout is None else {"timeout": self._timeout} + return {**defaults, **options, "headers": headers} def _log_request(self, url: str, method: str) -> None: sdk_logger.debug(f"Sending HTTP request: {method} {url}") @@ -100,13 +186,14 @@ def _prepare_json( async def get(self, url: str, model: type[TModel], **kwargs: Any) -> TModel: """Send a GET request and parse the JSON response into `model`.""" url = f"{self._base_url}{url}" - async with aiohttp.ClientSession(**self._client_config) as client: - self._log_request(url, "GET") - async with client.get(url, **kwargs) as response: - await handle_api_error(response) - self._log_response(url, "GET", response.status) - data = await response.json() - return parse_obj_as(model, data) + target = self._request_url(url) + client = await self._sessions.current() + self._log_request(url, "GET") + async with client.get(target, **self._request_options(kwargs)) as response: + await handle_api_error(response) + self._log_response(url, "GET", response.status) + data = await response.json() + return parse_obj_as(model, data) @handle_client_error async def post( @@ -118,13 +205,16 @@ async def post( ) -> TModel: """Send a POST request with a JSON body and parse the JSON response into `model`.""" url = f"{self._base_url}{url}" - async with aiohttp.ClientSession(**self._client_config) as client: - self._log_request(url, "POST") - async with client.post(url, json=self._prepare_json(json), **kwargs) as response: - await handle_api_error(response) - self._log_response(url, "POST", response.status) - data = await response.json() - return parse_obj_as(model, data) + target = self._request_url(url) + client = await self._sessions.current() + self._log_request(url, "POST") + async with client.post( + target, json=self._prepare_json(json), **self._request_options(kwargs) + ) as response: + await handle_api_error(response) + self._log_response(url, "POST", response.status) + data = await response.json() + return parse_obj_as(model, data) @handle_client_error async def put( @@ -136,13 +226,16 @@ async def put( ) -> TModel: """Send a PUT request with a JSON body and parse the JSON response into `model`.""" url = f"{self._base_url}{url}" - async with aiohttp.ClientSession(**self._client_config) as client: - self._log_request(url, "PUT") - async with client.put(url, json=self._prepare_json(json), **kwargs) as response: - await handle_api_error(response) - self._log_response(url, "PUT", response.status) - data = await response.json() - return parse_obj_as(model, data) + target = self._request_url(url) + client = await self._sessions.current() + self._log_request(url, "PUT") + async with client.put( + target, json=self._prepare_json(json), **self._request_options(kwargs) + ) as response: + await handle_api_error(response) + self._log_response(url, "PUT", response.status) + data = await response.json() + return parse_obj_as(model, data) @handle_client_error async def patch( @@ -154,13 +247,16 @@ async def patch( ) -> TModel: """Send a PATCH request with a JSON body and parse the JSON response into `model`.""" url = f"{self._base_url}{url}" - async with aiohttp.ClientSession(**self._client_config) as client: - self._log_request(url, "PATCH") - async with client.patch(url, json=self._prepare_json(json), **kwargs) as response: - await handle_api_error(response) - self._log_response(url, "PATCH", response.status) - data = await response.json() - return parse_obj_as(model, data) + target = self._request_url(url) + client = await self._sessions.current() + self._log_request(url, "PATCH") + async with client.patch( + target, json=self._prepare_json(json), **self._request_options(kwargs) + ) as response: + await handle_api_error(response) + self._log_response(url, "PATCH", response.status) + data = await response.json() + return parse_obj_as(model, data) @overload async def delete( @@ -190,15 +286,18 @@ async def delete( ) -> TModel | None: """Send a DELETE request; parse the JSON response into `model` if one is given.""" url = f"{self._base_url}{url}" - async with aiohttp.ClientSession(**self._client_config) as client: - self._log_request(url, "DELETE") - async with client.delete(url, json=self._prepare_json(json), **kwargs) as response: - await handle_api_error(response) - self._log_response(url, "DELETE", response.status) - if model is None: - return None - data = await response.json() - return parse_obj_as(model, data) + target = self._request_url(url) + client = await self._sessions.current() + self._log_request(url, "DELETE") + async with client.delete( + target, json=self._prepare_json(json), **self._request_options(kwargs) + ) as response: + await handle_api_error(response) + self._log_response(url, "DELETE", response.status) + if model is None: + return None + data = await response.json() + return parse_obj_as(model, data) class BasePermitApi: @@ -211,10 +310,21 @@ def __init__(self, config: PermitConfig) -> None: config: The Permit SDK configuration. """ self.config = config + self._sessions = LoopSessions() self.__api_keys = self._build_http_client("/v2/api-key") + def _use_sessions(self, sessions: LoopSessions) -> None: + """Send the requests of this API and of the APIs and clients it holds through ``sessions``. + + A Permit client calls it so that all of its APIs share one session per event loop. + """ + self._sessions = sessions + for value in vars(self).values(): + if isinstance(value, (BasePermitApi, SimpleHttpClient)): + value._use_sessions(sessions) # noqa: SLF001 - SDK-internal + def _build_http_client( - self, endpoint_url: str = "", *, use_pdp: bool = False, **kwargs: Any + self, endpoint_url: str = "", *, use_pdp: bool = False ) -> SimpleHttpClient: optional_headers = {} if self.config.proxy_facts_via_pdp: @@ -231,12 +341,11 @@ def _build_http_client( **optional_headers, }, ) - client_config_dict = client_config.dict() - client_config_dict.update(kwargs) return SimpleHttpClient( - client_config_dict, + client_config.dict(), base_url=endpoint_url, timeout=self.config.api_timeout, + sessions=self._sessions, ) async def _set_context_from_api_key(self) -> None: diff --git a/permit/enforcement/enforcer.py b/permit/enforcement/enforcer.py index e8e7393b..32c5bf71 100644 --- a/permit/enforcement/enforcer.py +++ b/permit/enforcement/enforcer.py @@ -17,6 +17,7 @@ from permit.exceptions import PermitConnectionError from permit.utils.context import Context, ContextStore from permit.utils.dicts import deep_merge +from permit.utils.http_sessions import LoopSessions from permit.utils.pydantic_version import PYDANTIC_VERSION from permit.utils.sdk_logger import sdk_logger from permit.utils.sync import SyncClass @@ -105,6 +106,11 @@ def __init__(self, config: PermitConfig) -> None: "Authorization": f"Bearer {self._config.token}", } self._base_url = self._config.pdp + self._sessions = LoopSessions() + + def _use_sessions(self, sessions: LoopSessions) -> None: + """Send the queries through ``sessions`` from now on.""" + self._sessions = sessions @property def context_store(self) -> ContextStore: @@ -169,71 +175,73 @@ async def authorized_users( "context": query_context, } - async with aiohttp.ClientSession(headers=self._headers, **self._timeout_config) as session: - check_url = f"{self._base_url}/authorized_users" - try: - async with session.post( - check_url, - data=json.dumps(request_body), - ) as response: - if response.status != HTTPStatus.OK: - if response.status == HTTPStatus.NOT_IMPLEMENTED: - msg = ( - f"Permit SDK got an error: {response.status}, " - f"and cannot connect to the PDP container." - f"\nPlease ensure you are not using ABAC/ReBAC policies," - f"as the cloud PDP is not compatible with these kinds " - f"of policies.\n" - f"Also, please check your configuration and " - f"make sure it's running at {self._base_url} " - f"and accepting requests.\n" - f"Read more about setting up the PDP at {SETUP_PDP_DOCS_LINK}" - ) - raise PermitConnectionError(msg) - - error_body = await read_error_body(response) - sdk_logger.error( - "error in permit.authorized_users({}, {}):\n{}\n{}".format( - action, - self._resource_repr(normalized_resource), - f"status code: {response.status}", - error_body, - ) - ) + session = await self._sessions.current() + check_url = f"{self._base_url}/authorized_users" + try: + async with session.post( + check_url, + data=json.dumps(request_body), + headers=self._headers, + **self._timeout_config, + ) as response: + if response.status != HTTPStatus.OK: + if response.status == HTTPStatus.NOT_IMPLEMENTED: msg = ( - f"Permit SDK got unexpected status code: {response.status} " - f"from the PDP at {self._base_url}.\nResponse body: {error_body}\n" - f"The PDP is reachable, so this is a rejected request rather than a " - f"connectivity problem -- a 401/403 usually means the PDP was started " - f"with a different API key than the SDK is using.\n" + f"Permit SDK got an error: {response.status}, " + f"and cannot connect to the PDP container." + f"\nPlease ensure you are not using ABAC/ReBAC policies," + f"as the cloud PDP is not compatible with these kinds " + f"of policies.\n" + f"Also, please check your configuration and " + f"make sure it's running at {self._base_url} " + f"and accepting requests.\n" f"Read more about setting up the PDP at {SETUP_PDP_DOCS_LINK}" ) raise PermitConnectionError(msg) - content: dict[str, Any] = await response.json() - sdk_logger.debug( - f"permit.authorized_users() response:" - f"\ninput: {pformat(request_body, indent=2)}" - f"\nresponse status: {response.status}" - f"\nresponse data: {pformat(content, indent=2)}" + error_body = await read_error_body(response) + sdk_logger.error( + "error in permit.authorized_users({}, {}):\n{}\n{}".format( + action, + self._resource_repr(normalized_resource), + f"status code: {response.status}", + error_body, + ) ) - result: AuthorizedUsersResult = parse_obj_as(AuthorizedUsersResult, content) - return result - except aiohttp.ClientError as err: - sdk_logger.error( - f"error in permit.authorized_users({action}, " - f"{self._resource_repr(normalized_resource)}):\n{err}" - ) - msg = ( - f"Permit SDK got error: {err}, and cannot connect to the PDP container.\n" - f"Please check your configuration and make sure it's running at " - f"{self._base_url} and accepting requests.\n " - f"Read more about setting up the PDP at {SETUP_PDP_DOCS_LINK}" + msg = ( + f"Permit SDK got unexpected status code: {response.status} " + f"from the PDP at {self._base_url}.\nResponse body: {error_body}\n" + f"The PDP is reachable, so this is a rejected request rather than a " + f"connectivity problem -- a 401/403 usually means the PDP was started " + f"with a different API key than the SDK is using.\n" + f"Read more about setting up the PDP at {SETUP_PDP_DOCS_LINK}" + ) + raise PermitConnectionError(msg) + + content: dict[str, Any] = await response.json() + sdk_logger.debug( + f"permit.authorized_users() response:" + f"\ninput: {pformat(request_body, indent=2)}" + f"\nresponse status: {response.status}" + f"\nresponse data: {pformat(content, indent=2)}" ) - raise PermitConnectionError( - msg, - error=err, - ) from err + result: AuthorizedUsersResult = parse_obj_as(AuthorizedUsersResult, content) + return result + except aiohttp.ClientError as err: + sdk_logger.error( + f"error in permit.authorized_users({action}, " + f"{self._resource_repr(normalized_resource)}):\n{err}" + ) + msg = ( + f"Permit SDK got error: {err}, and cannot connect to the PDP container.\n" + f"Please check your configuration and make sure it's running at " + f"{self._base_url} and accepting requests.\n " + f"Read more about setting up the PDP at {SETUP_PDP_DOCS_LINK}" + ) + raise PermitConnectionError( + msg, + error=err, + ) from err async def bulk_check( self, @@ -304,57 +312,59 @@ async def bulk_check( } ) - async with aiohttp.ClientSession(headers=self._headers, **self._timeout_config) as session: - check_url = f"{self._base_url}/allowed/bulk" - try: - async with session.post( - check_url, - data=json.dumps(request_body), - ) as response: - if response.status != HTTPStatus.OK: - error_body = await read_error_body(response) - msg = "error in permit.check({}):\n{}\n{}".format( - ( + session = await self._sessions.current() + check_url = f"{self._base_url}/allowed/bulk" + try: + async with session.post( + check_url, + data=json.dumps(request_body), + headers=self._headers, + **self._timeout_config, + ) as response: + if response.status != HTTPStatus.OK: + error_body = await read_error_body(response) + msg = "error in permit.check({}):\n{}\n{}".format( + ( + [ [ - [ - check.get("user"), - check.get("action"), - check.get("resource"), - ] - for check in request_body + check.get("user"), + check.get("action"), + check.get("resource"), ] - ), - f"status code: {response.status}", - error_body, - ) - sdk_logger.error(msg) - raise PermitConnectionError(msg) - content: dict[str, Any] = await response.json() - sdk_logger.debug( - f"permit.check() response:\n" - f"input: {pformat(request_body, indent=2)}\n" - f"response status: {response.status}\n" - f"response data: {pformat(content, indent=2)}" + for check in request_body + ] + ), + f"status code: {response.status}", + error_body, ) - data = content.get("allow", content.get("result", {}).get("allow", [])) - decisions: list[bool] = [bool(item.get("allow", False)) for item in data] - except aiohttp.ClientError as err: - msg = "error in permit.check({}):\n{}".format( - ( + sdk_logger.error(msg) + raise PermitConnectionError(msg) + content: dict[str, Any] = await response.json() + sdk_logger.debug( + f"permit.check() response:\n" + f"input: {pformat(request_body, indent=2)}\n" + f"response status: {response.status}\n" + f"response data: {pformat(content, indent=2)}" + ) + data = content.get("allow", content.get("result", {}).get("allow", [])) + decisions: list[bool] = [bool(item.get("allow", False)) for item in data] + except aiohttp.ClientError as err: + msg = "error in permit.check({}):\n{}".format( + ( + [ [ - [ - check.get("user"), - check.get("action"), - check.get("resource"), - ] - for check in request_body + check.get("user"), + check.get("action"), + check.get("resource"), ] - ), - err, - ) - sdk_logger.error(msg) - raise PermitConnectionError(msg, error=err) from err - return decisions + for check in request_body + ] + ), + err, + ) + sdk_logger.error(msg) + raise PermitConnectionError(msg, error=err) from err + return decisions async def check( self, @@ -407,73 +417,75 @@ async def check( "resource": normalized_resource.dict(exclude_unset=True), "context": query_context, } - async with aiohttp.ClientSession(headers=self._headers, **self._timeout_config) as session: - check_url = f"{self._base_url}/allowed" - try: - async with session.post( - check_url, - data=json.dumps(body), - ) as response: - if response.status != HTTPStatus.OK: - if response.status == HTTPStatus.NOT_IMPLEMENTED: - msg = ( - f"Permit SDK got an error: {response.status}, " - f"and cannot connect to the PDP container." - f"\nPlease ensure you are not using ABAC/ReBAC policies,\n" - f"as the cloud PDP is not compatible with these kinds " - f"of policies.\n" - f"Also, please check your configuration and make sure it's running " - f"at {self._base_url} and accepting requests.\n" - f"Read more about setting up the PDP at {SETUP_PDP_DOCS_LINK}" - ) - raise PermitConnectionError(msg) - - error_body = await read_error_body(response) - sdk_logger.error( - "error in permit.check({}, {}, {}):\n{}\n{}".format( - normalized_user, - action, - self._resource_repr(normalized_resource), - f"status code: {response.status}", - error_body, - ) - ) + session = await self._sessions.current() + check_url = f"{self._base_url}/allowed" + try: + async with session.post( + check_url, + data=json.dumps(body), + headers=self._headers, + **self._timeout_config, + ) as response: + if response.status != HTTPStatus.OK: + if response.status == HTTPStatus.NOT_IMPLEMENTED: msg = ( - f"Permit SDK got unexpected status code: {response.status} " - f"from the PDP at {self._base_url}.\nResponse body: {error_body}\n" - f"The PDP is reachable, so this is a rejected request rather than a " - f"connectivity problem -- a 401/403 usually means the PDP was started " - f"with a different API key than the SDK is using.\n" + f"Permit SDK got an error: {response.status}, " + f"and cannot connect to the PDP container." + f"\nPlease ensure you are not using ABAC/ReBAC policies,\n" + f"as the cloud PDP is not compatible with these kinds " + f"of policies.\n" + f"Also, please check your configuration and make sure it's running " + f"at {self._base_url} and accepting requests.\n" f"Read more about setting up the PDP at {SETUP_PDP_DOCS_LINK}" ) raise PermitConnectionError(msg) - content: dict[str, Any] = await response.json() - sdk_logger.debug( - f"permit.check() response:\n" - f"body: {pformat(body, indent=2)}\n" - f"response status: {response.status}\n" - f"response data: {pformat(content, indent=2)}" + error_body = await read_error_body(response) + sdk_logger.error( + "error in permit.check({}, {}, {}):\n{}\n{}".format( + normalized_user, + action, + self._resource_repr(normalized_resource), + f"status code: {response.status}", + error_body, + ) ) - decision: bool = bool(content.get("allow", False)) - return decision - except aiohttp.ClientError as err: - sdk_logger.error( - f"error in permit.check({normalized_user}, {action}, " - f"{self._resource_repr(normalized_resource)}):" - f"\n{err}" - ) - msg = ( - f"Permit SDK got error: {err}, \n" - f"and cannot connect to the PDP container, please check your configuration " - f"and make sure it's " - f"running at {self._base_url} and accepting requests. \n" - f"Read more about setting up the PDP at {SETUP_PDP_DOCS_LINK}" + msg = ( + f"Permit SDK got unexpected status code: {response.status} " + f"from the PDP at {self._base_url}.\nResponse body: {error_body}\n" + f"The PDP is reachable, so this is a rejected request rather than a " + f"connectivity problem -- a 401/403 usually means the PDP was started " + f"with a different API key than the SDK is using.\n" + f"Read more about setting up the PDP at {SETUP_PDP_DOCS_LINK}" + ) + raise PermitConnectionError(msg) + + content: dict[str, Any] = await response.json() + sdk_logger.debug( + f"permit.check() response:\n" + f"body: {pformat(body, indent=2)}\n" + f"response status: {response.status}\n" + f"response data: {pformat(content, indent=2)}" ) - raise PermitConnectionError( - msg, - error=err, - ) from err + decision: bool = bool(content.get("allow", False)) + return decision + except aiohttp.ClientError as err: + sdk_logger.error( + f"error in permit.check({normalized_user}, {action}, " + f"{self._resource_repr(normalized_resource)}):" + f"\n{err}" + ) + msg = ( + f"Permit SDK got error: {err}, \n" + f"and cannot connect to the PDP container, please check your configuration " + f"and make sure it's " + f"running at {self._base_url} and accepting requests. \n" + f"Read more about setting up the PDP at {SETUP_PDP_DOCS_LINK}" + ) + raise PermitConnectionError( + msg, + error=err, + ) from err async def get_user_permissions( self, @@ -510,50 +522,52 @@ async def get_user_permissions( if context is not None: input_data["context"] = self._context_store.get_derived_context(context) - async with aiohttp.ClientSession(headers=self._headers, **self._timeout_config) as session: - url = f"{self._base_url}/user-permissions" - try: - async with session.post( - url, - data=json.dumps(input_data), - ) as response: - if response.status != HTTPStatus.OK: - msg = ( - f"Permit.getUserPermissions() got an unexpected status code: " - f"{response.status}, " - f"please check your SDK init and make sure the PDP sidecar " - f"is configured correctly.\n" - f"Read more about setting up the PDP at {SETUP_PDP_DOCS_LINK}" - ) - raise PermitConnectionError(msg) - - content = await response.json() - permissions: dict[str, Any] = ( - content.get("result", {}).get("permissions", {}) - if "result" in content - else content + session = await self._sessions.current() + url = f"{self._base_url}/user-permissions" + try: + async with session.post( + url, + data=json.dumps(input_data), + headers=self._headers, + **self._timeout_config, + ) as response: + if response.status != HTTPStatus.OK: + msg = ( + f"Permit.getUserPermissions() got an unexpected status code: " + f"{response.status}, " + f"please check your SDK init and make sure the PDP sidecar " + f"is configured correctly.\n" + f"Read more about setting up the PDP at {SETUP_PDP_DOCS_LINK}" ) + raise PermitConnectionError(msg) - sdk_logger.debug( - f"permit.get_user_permissions() response:\n" - f"input: {pformat(input_data, indent=2)}\n" - f"response data: {pformat(permissions, indent=2)}" - ) - return permissions - - except aiohttp.ClientError as err: - sdk_logger.error(f"Error in permit.get_user_permissions(): {err}") - msg = ( - f"Permit SDK got error: {err}, \n" - f"and cannot connect to the PDP container, please check your configuration " - f"and make sure it's " - f"running at {self._base_url} and accepting requests. \n" - f"Read more about setting up the PDP at {SETUP_PDP_DOCS_LINK}" + content = await response.json() + permissions: dict[str, Any] = ( + content.get("result", {}).get("permissions", {}) + if "result" in content + else content ) - raise PermitConnectionError( - msg, - error=err, - ) from err + + sdk_logger.debug( + f"permit.get_user_permissions() response:\n" + f"input: {pformat(input_data, indent=2)}\n" + f"response data: {pformat(permissions, indent=2)}" + ) + return permissions + + except aiohttp.ClientError as err: + sdk_logger.error(f"Error in permit.get_user_permissions(): {err}") + msg = ( + f"Permit SDK got error: {err}, \n" + f"and cannot connect to the PDP container, please check your configuration " + f"and make sure it's " + f"running at {self._base_url} and accepting requests. \n" + f"Read more about setting up the PDP at {SETUP_PDP_DOCS_LINK}" + ) + raise PermitConnectionError( + msg, + error=err, + ) from err async def get_user_tenants( self, user: User, context: Context | None = None @@ -591,38 +605,40 @@ async def get_user_tenants( "context": self._context_store.get_derived_context(context or {}), } - async with aiohttp.ClientSession(headers=self._headers, **self._timeout_config) as session: - url = f"{self._base_url}/user-tenants" - try: - async with session.post(url, data=json.dumps(body)) as response: - if response.status == HTTPStatus.NOT_FOUND: - msg = ( - f"permit.get_user_tenants() got status code 404 from the PDP at " - f"{self._base_url}: only the container PDP serves /user-tenants, " - f"and the cloud PDP does not.\n" - f"Point the SDK's `pdp` setting at a container PDP to use it.\n" - f"Read more about setting up the PDP at {SETUP_PDP_DOCS_LINK}" - ) - raise PermitConnectionError(msg) - if response.status != HTTPStatus.OK: - error_body = await read_error_body(response) - msg = ( - f"permit.get_user_tenants() got an unexpected status code: " - f"{response.status} from the PDP at {self._base_url}.\n" - f"Response body: {error_body}\n" - f"Read more about setting up the PDP at {SETUP_PDP_DOCS_LINK}" - ) - raise PermitConnectionError(msg) - content = await response.json() - except aiohttp.ClientError as err: - sdk_logger.error(f"Error in permit.get_user_tenants(): {err}") - msg = ( - f"Permit SDK got error: {err}, \n" - f"and cannot connect to the PDP container, please check your configuration " - f"and make sure it's running at {self._base_url} and accepting requests. \n" - f"Read more about setting up the PDP at {SETUP_PDP_DOCS_LINK}" - ) - raise PermitConnectionError(msg, error=err) from err + session = await self._sessions.current() + url = f"{self._base_url}/user-tenants" + try: + async with session.post( + url, data=json.dumps(body), headers=self._headers, **self._timeout_config + ) as response: + if response.status == HTTPStatus.NOT_FOUND: + msg = ( + f"permit.get_user_tenants() got status code 404 from the PDP at " + f"{self._base_url}: only the container PDP serves /user-tenants, " + f"and the cloud PDP does not.\n" + f"Point the SDK's `pdp` setting at a container PDP to use it.\n" + f"Read more about setting up the PDP at {SETUP_PDP_DOCS_LINK}" + ) + raise PermitConnectionError(msg) + if response.status != HTTPStatus.OK: + error_body = await read_error_body(response) + msg = ( + f"permit.get_user_tenants() got an unexpected status code: " + f"{response.status} from the PDP at {self._base_url}.\n" + f"Response body: {error_body}\n" + f"Read more about setting up the PDP at {SETUP_PDP_DOCS_LINK}" + ) + raise PermitConnectionError(msg) + content = await response.json() + except aiohttp.ClientError as err: + sdk_logger.error(f"Error in permit.get_user_tenants(): {err}") + msg = ( + f"Permit SDK got error: {err}, \n" + f"and cannot connect to the PDP container, please check your configuration " + f"and make sure it's running at {self._base_url} and accepting requests. \n" + f"Read more about setting up the PDP at {SETUP_PDP_DOCS_LINK}" + ) + raise PermitConnectionError(msg, error=err) from err sdk_logger.debug( f"permit.get_user_tenants() response:\n" diff --git a/permit/pdp_api/base.py b/permit/pdp_api/base.py index 1002226f..7e2edd07 100644 --- a/permit/pdp_api/base.py +++ b/permit/pdp_api/base.py @@ -1,7 +1,6 @@ -from typing import Any - from permit import PermitConfig from permit.api.base import ClientConfig, SimpleHttpClient, pagination_params +from permit.utils.http_sessions import LoopSessions __all__ = ["BasePdpPermitApi", "ClientConfig", "pagination_params"] @@ -16,8 +15,13 @@ def __init__(self, config: PermitConfig) -> None: config: The Permit SDK configuration. """ self.config = config + self._sessions = LoopSessions() + + def _use_sessions(self, sessions: LoopSessions) -> None: + """Send this API's requests through ``sessions`` from now on.""" + self._sessions = sessions - def _build_http_client(self, endpoint_url: str = "", **kwargs: Any) -> SimpleHttpClient: + def _build_http_client(self, endpoint_url: str = "") -> SimpleHttpClient: client_config = ClientConfig( base_url=f"{self.config.pdp}", headers={ @@ -25,13 +29,12 @@ def _build_http_client(self, endpoint_url: str = "", **kwargs: Any) -> SimpleHtt "Authorization": f"Bearer {self.config.token}", }, ) - client_config_dict = client_config.dict() - client_config_dict.update(kwargs) return SimpleHttpClient( - client_config_dict, + client_config.dict(), base_url=endpoint_url, # pdp_timeout was documented on PermitConfig and honoured by the # enforcer, but silently ignored here, so every permit.pdp_api.* # call used aiohttp's default timeout instead of the configured one. timeout=self.config.pdp_timeout, + sessions=self._sessions, ) diff --git a/permit/pdp_api/pdp_api_client.py b/permit/pdp_api/pdp_api_client.py index 6417858f..014cc4f3 100644 --- a/permit/pdp_api/pdp_api_client.py +++ b/permit/pdp_api/pdp_api_client.py @@ -2,6 +2,7 @@ from permit.config import PermitConfig from permit.pdp_api.role_assignments import RoleAssignmentsApi +from permit.utils.http_sessions import LoopSessions from permit.utils.sync import SyncClass # Type checkers read this class from a generated stub: the SyncClass metaclass @@ -36,6 +37,10 @@ def __init__(self, config: PermitConfig) -> None: self._role_assignments = RoleAssignmentsApi(config) + def _use_sessions(self, sessions: LoopSessions) -> None: + """Send the requests of every API of this client through ``sessions`` from now on.""" + self._role_assignments._use_sessions(sessions) # noqa: SLF001 - SDK-internal + @property def role_assignments(self) -> RoleAssignmentsApi: """Role assignments as the PDP currently sees them.""" diff --git a/permit/permit.py b/permit/permit.py index c6836346..1754b6db 100644 --- a/permit/permit.py +++ b/permit/permit.py @@ -1,6 +1,7 @@ import copy from collections.abc import Generator from contextlib import contextmanager +from types import TracebackType from typing import Any, Literal from typing_extensions import Self @@ -19,12 +20,25 @@ from permit.logger import configure_logger from permit.pdp_api.pdp_api_client import PermitPdpApiClient from permit.utils.context import Context +from permit.utils.http_sessions import LoopSessions from permit.utils.sdk_logger import sdk_logger class Permit: """The Permit SDK client (asyncio): authorization checks and the Permit REST API. + The client keeps its HTTP connections open and reuses them: one aiohttp session, with + its own pool of connections, for the Permit API and one for the PDP, per event loop it + is used on. They are created by the first request from each loop. Close them with + ``await permit.close()``, or use the client as an async context manager:: + + async with Permit(token="") as permit: + await permit.check("user", "read", "document") + + A client that is never closed leaves nothing open behind it under ``asyncio.run()``, + which closes the loop's sessions as it shuts the loop down, nor once it is garbage + collected while its loop runs. + Args: config: The SDK configuration. **options: `PermitConfig` fields, used to build the configuration when `config` @@ -35,7 +49,13 @@ def __init__(self, config: PermitConfig | None = None, **options: Any) -> None: self._config: PermitConfig = config if config is not None else PermitConfig(**options) configure_logger(self._config) + self._api_sessions = LoopSessions() + self._pdp_sessions = LoopSessions() + # A copy made by wait_for_sync() shares the sessions of the client it copies, and + # leaves closing them to that client. + self._owns_sessions = True self._connect() + self._share_sessions() sdk_logger.debug( f"Permit SDK initialized: api_url={self._config.api_url}, pdp={self._config.pdp}" ) @@ -47,6 +67,50 @@ def _connect(self) -> None: self._elements = ElementsApi(self._config) self._pdp_api = PermitPdpApiClient(self._config) + def _share_sessions(self) -> None: + """Make the clients `_connect()` created send their requests through the sessions.""" + self._enforcer._use_sessions(self._pdp_sessions) # noqa: SLF001 - SDK-internal + self._pdp_api._use_sessions(self._pdp_sessions) # noqa: SLF001 - SDK-internal + self._api._use_sessions(self._api_sessions) # noqa: SLF001 - SDK-internal + self._elements._use_sessions(self._api_sessions) # noqa: SLF001 - SDK-internal + + async def close(self) -> None: + """Close the HTTP connections this client keeps open. + + It closes the sessions of the event loop it runs on, of loops already closed, and of + loops running in other threads, on those loops, waiting for each while its loop + runs. The session of a loop that is neither running nor closed, or that stops before + it has closed its session, stays open until that loop shuts down its async + generators, as ``asyncio.run()`` does, or ``close()`` runs on it. When one session + fails to close, the others are still closed before the error is raised. + + A request still in flight when ``close()`` runs fails. Calling ``close()`` again + closes nothing more. The client stays usable: a request sent after ``close()`` + opens new connections, which a later ``close()`` closes. + + With ``proxy_facts_via_pdp`` on, a client yielded by ``wait_for_sync()`` sends its + requests over the connections of the client it was made from: its ``close()`` does + nothing, and the other client's ``close()`` closes them. With it off, the default, + ``wait_for_sync()`` yields the client itself, whose ``close()`` closes them. + """ + if not self._owns_sessions: + return + try: + await self._api_sessions.close() + finally: + await self._pdp_sessions.close() + + async def __aenter__(self) -> Self: + return self + + async def __aexit__( + self, + exc_type: type[BaseException] | None, + exc: BaseException | None, + traceback: TracebackType | None, + ) -> None: + await self.close() + @property def config(self) -> PermitConfig: """Access the SDK configuration using this property. @@ -78,7 +142,11 @@ def wait_for_sync( PDP. Yields: - Permit: A Permit instance that is configured to wait for facts to be synced. + Permit: A Permit instance that is configured to wait for facts to be synced. It + sends its requests over this client's connections, so it needs no ``close()``: + closing this client closes them, and its own ``close()`` does nothing. With + ``proxy_facts_via_pdp`` off, it logs a warning and yields this client itself, + whose ``close()`` closes them. See Also: https://docs.permit.io/how-to/manage-data/local-facts-uploader @@ -97,7 +165,9 @@ def wait_for_sync( # client instead would apply its log settings to the whole process again. waiting: Self = copy.copy(self) waiting._config = contextualized_config + waiting._owns_sessions = False waiting._connect() + waiting._share_sessions() yield waiting @property diff --git a/permit/sync.py b/permit/sync.py index a0c556d4..ed3861b2 100644 --- a/permit/sync.py +++ b/permit/sync.py @@ -1,5 +1,9 @@ +import weakref +from types import TracebackType from typing import Any +from typing_extensions import Self + from permit.api.elements import SyncElementsApi from permit.api.sync_api_client import SyncPermitApiClient from permit.config import PermitConfig @@ -14,6 +18,7 @@ from permit.pdp_api.pdp_api_client import SyncPDPApi from permit.permit import Permit as AsyncPermit from permit.utils.context import Context +from permit.utils.sync import _BackgroundLoop # The blocking client keeps the blocking twins of the async client's helpers in the @@ -23,20 +28,103 @@ class Permit(AsyncPermit): """The Permit SDK client with a blocking interface. + The client runs every blocking call on an event loop in a background daemon thread of + its own, which it starts on the first call. Calls from any number of threads are handed + to that thread and waited for, so they share the client's HTTP connections instead of + each opening its own. Calling it from a thread that runs an event loop works too, and + blocks that loop until the call returns, as any blocking call does. + + Close the client when done with it, with `close()` or a `with` block, to close its + connections and stop the thread. A client that is never closed is cleaned up when it is + garbage collected, or at interpreter exit; the thread never holds up the exit. + Args: config: The SDK configuration. **options: `PermitConfig` fields, used to build the configuration when `config` is not given. + + Examples: + with Permit(token="") as permit: + permit.check("user", "read", "document") """ def __init__(self, config: PermitConfig | None = None, **options: Any) -> None: + # Before super().__init__, which calls _connect. + self._background_loop = _BackgroundLoop() super().__init__(config, **options) + # close() and the exit hook close the sessions on the loop while the client is + # alive; once it is collected, the sessions close themselves there. Copies made by + # wait_for_sync() use the sessions and the loop of the client that made them, and + # leave closing both to it. + self._background_loop.set_closer(weakref.WeakMethod(self._close_sessions)) def _connect(self) -> None: self._enforcer = SyncEnforcer(self._config) # type: ignore[assignment] self._api = SyncPermitApiClient(self._config) # type: ignore[assignment] self._elements = SyncElementsApi(self._config) # type: ignore[assignment] self._pdp_api = SyncPDPApi(self._config) + self._background_loop.bind(self._enforcer, self._api, self._elements, self._pdp_api) + + async def _close_sessions(self) -> None: + """Close the HTTP sessions this client opened. Runs on its background loop.""" + await AsyncPermit.close(self) + + def close(self) -> None: # type: ignore[override] + """Close the client's HTTP connections and stop its background thread. + + It waits for the calls that other threads have in flight to return first. A call or a + `close()` that another thread makes meanwhile waits until this one has finished. + Calling it again does nothing. The client stays usable: the next call starts a new + thread and opens new connections. + + With `proxy_facts_via_pdp` on, a client yielded by `wait_for_sync()` runs its calls on + the thread and over the connections of the client it was made from: its `close()` + does nothing, and the other client's `close()` closes them. With it off, the default, + `wait_for_sync()` yields the client itself, whose `close()` closes them. + + Raises: + RuntimeError: If called on the client's own background thread, which it has to + stop and join. + + Examples: + permit = Permit(token="") + try: + permit.check("user", "read", "document") + finally: + permit.close() + """ + if not self._owns_sessions: + return + self._background_loop.close() + + def __enter__(self) -> Self: + """Return the client itself, which the end of the `with` block closes.""" + return self + + def __exit__( + self, + exc_type: type[BaseException] | None, + exc_value: BaseException | None, + traceback: TracebackType | None, + ) -> None: + """Close the client, as `close()` does.""" + self.close() + + def __aenter__(self) -> None: # type: ignore[override] + """Refuse `async with`, which the blocking client does not support. + + It is annotated to return None rather than an awaitable, so that type checkers reject + `async with` on this client too, as they reject `with` on the async client. + + Raises: + TypeError: Always. A `with` block closes this client; `async with` is for the + async client, `permit.Permit`. + """ + msg = ( + "permit.sync.Permit is a blocking client: use `with Permit(...) as permit:`, not " + "`async with`. In async code, use the async client, permit.Permit." + ) + raise TypeError(msg) @property def api(self) -> SyncPermitApiClient: # type: ignore[override] diff --git a/permit/utils/http_sessions.py b/permit/utils/http_sessions.py new file mode 100644 index 00000000..6f40fd4f --- /dev/null +++ b/permit/utils/http_sessions.py @@ -0,0 +1,365 @@ +import asyncio +import atexit +import concurrent.futures +import contextlib +import functools +import os +import sys +import threading +import weakref +from collections.abc import AsyncGenerator +from typing import NamedTuple + +import aiohttp + + +class _LoopSession(NamedTuple): + """A loop's session, and the async generator that closes it when the loop shuts down.""" + + session: aiohttp.ClientSession + closer: AsyncGenerator[None, None] + + +class LoopSessions: + """The aiohttp sessions an SDK client sends its requests through, one per event loop. + + An aiohttp session, and the connections it keeps open for reuse, belong to the event + loop that created them. So the client has one session per loop it is used on: a single + one in an application that runs one loop, a new one for each ``asyncio.run()`` call. + Each is created by the first request sent from its loop. + + A session is closed: + + - by ``close()``; + - when its loop shuts down its async generators, as ``asyncio.run()`` and + ``asyncio.Runner`` do before they close the loop, so a program that never calls + ``close()`` does not leave it open; + - as the interpreter exits, if its loop is still open then: on that loop if it is not + running, or by that loop's thread if it runs in another thread; + - once this object is garbage collected, on its loop if that loop is running. Until + then, and until the session is closed, a finalizer holds it apart from this object: + the garbage collector never finds an open session unreachable, so aiohttp never + reports one unclosed, even when the client that holds this object ends up in a + reference cycle. + + A child process made by ``fork()`` sets the sessions it inherits aside, untouched: their + loops cannot run in the child, and their connections are the parent's. The child's + requests open sessions of their own. + + The sessions carry no headers, base URL or timeout: each request brings its own, so one + session serves every request sent from its loop. They keep no cookies either, so a + request carries exactly the headers it would carry through a session of its own. + + Every API object, HTTP client and enforcer builds one of these for itself, so that it + works when used alone. A ``Permit`` client then gives all of them its own two, one for + the Permit API and one for the PDP, through their ``_use_sessions()``: the ones they + built open no session, and are collected right away. + """ + + def __init__(self) -> None: + self._lock = threading.Lock() + # Changed in place only: the finalizer holds this dict. + self._sessions: dict[asyncio.AbstractEventLoop, _LoopSession] = {} + _open_at_exit.add(self) + finalizer = weakref.finalize(self, _orphan, self._sessions) + # Writable, as the weakref documentation says; typeshed declares __slots__ = () on + # it. At exit, the exit hook below closes the sessions of the objects still alive. + finalizer.atexit = False # type: ignore[misc] + + async def current(self) -> aiohttp.ClientSession: + """The session of the running event loop, created by the first call from that loop. + + Returns: + The session to send the request through. + """ + loop = asyncio.get_running_loop() + with self._lock: + existing = self._sessions.get(loop) + # A session close() kept, because its loop stopped first, may be closing now. + if existing is not None and not existing.session.closed: + return existing.session + abandoned = self._take_sessions_of_closed_loops() + _take_orphans_of_closed_loops() + session = aiohttp.ClientSession( + # No limit on concurrent connections, as when every request had a session of + # its own; idle connections are kept open for the next request. + connector=aiohttp.TCPConnector(limit=0), + cookie_jar=aiohttp.DummyCookieJar(), + ) + closer = _close_with_loop(weakref.ref(self), loop, session) + self._sessions[loop] = _LoopSession(session, closer) + # Runs the generator up to its `yield`, which registers it with the loop: the loop + # closes it, and so the session, when it shuts down its async generators. + await anext(closer) + for stale in abandoned: + await stale.session.close() + return session + + async def close(self) -> None: + """Close the sessions of every loop that can close them now. + + The session of the running loop, and those of loops already closed, are closed here. + The session of a loop running in another thread is closed on that loop, and this + waits for it while that loop runs. A loop that stops before it has closed its + session, or that is neither running nor closed, cannot run anything now: its session + stays open until that loop shuts down its async generators, or ``close()`` runs on + it. A request in flight on a session being closed fails. + + The sessions are closed side by side: when one fails to close, the others are still + closed, and the first error is raised then. A session that is still open afterwards, + including when the task running this is cancelled, is kept for the next ``close()``. + """ + running = asyncio.get_running_loop() + with self._lock: + closable = { + loop: entry + for loop, entry in self._sessions.items() + if loop is running or loop.is_closed() or loop.is_running() + } + for loop in closable: + del self._sessions[loop] + try: + outcomes = await asyncio.gather( + *(_close_from(running, loop, entry) for loop, entry in closable.items()), + return_exceptions=True, + ) + finally: + for loop, entry in closable.items(): + self._keep_if_open(loop, entry) + # Errors only: gather has raised the cancellation of this task already. + errors = [outcome for outcome in outcomes if isinstance(outcome, Exception)] + if errors: + raise errors[0] + + def _keep_if_open(self, loop: asyncio.AbstractEventLoop, entry: _LoopSession) -> None: + """Keep ``entry``, which ``close()`` took, if its session is still open.""" + with self._lock: + if entry.session.closed: + return + if loop in self._sessions: + # The loop opened a new session meanwhile: keep this one apart, for its loop + # to close when it shuts down its async generators. + _orphaned[id(entry.session)] = (loop, entry) + else: + self._sessions[loop] = entry + + def _forget(self, loop: asyncio.AbstractEventLoop, session: aiohttp.ClientSession) -> None: + """Drop ``session`` from the sessions, if it is still the one of ``loop``.""" + with self._lock: + entry = self._sessions.get(loop) + if entry is not None and entry.session is session: + del self._sessions[loop] + + def _take_sessions_of_closed_loops(self) -> list[_LoopSession]: + """Remove and return the sessions of loops closed without shutting them down. + + The caller holds the lock. + """ + closed = [loop for loop in self._sessions if loop.is_closed()] + return [self._sessions.pop(loop) for loop in closed] + + def _set_aside_after_fork(self) -> None: + """In a child made by ``fork()``: keep the inherited sessions, but never use them. + + Only the thread that forked runs in the child, so the lock may be held by a thread + that is gone, and no loop of the parent runs. Closing an inherited session would + close connections the parent still uses, so they stay open, and referenced, for the + life of the child. + """ + self._lock = threading.Lock() + _sessions_lost_to_fork.extend(self._sessions.values()) + self._sessions.clear() + + def _close_at_exit(self) -> None: + """Close every session as the interpreter exits, from a thread that runs no loop.""" + with self._lock: + entries = list(self._sessions.items()) + self._sessions.clear() + of_closed_loops = [ + entry.session for loop, entry in entries if not _close_at_exit_on(loop, entry) + ] + if of_closed_loops: + # The connections of a closed loop cannot be closed, but its sessions can be + # marked closed from any loop, which keeps aiohttp from reporting them unclosed. + asyncio.run(_close_all(of_closed_loops)) + + +async def _close_with_loop( + sessions: weakref.ref[LoopSessions], + loop: asyncio.AbstractEventLoop, + session: aiohttp.ClientSession, +) -> AsyncGenerator[None, None]: + """An async generator that closes ``session`` when it is closed. + + It holds ``sessions`` weakly, so that a client dropped without ``close()`` is garbage + collected. + """ + try: + yield + finally: + owner = sessions() + if owner is not None: + owner._forget(loop, session) # noqa: SLF001 - this module's own class + try: + await session.close() + finally: + _orphaned.pop(id(session), None) + + +async def _close_all(sessions: list[aiohttp.ClientSession]) -> None: + for session in sessions: + await session.close() + + +async def _close_from( + running: asyncio.AbstractEventLoop, loop: asyncio.AbstractEventLoop, entry: _LoopSession +) -> None: + """Close ``entry``, the session of ``loop``, from the ``running`` loop.""" + if loop is running: + await _close_on_its_loop(entry) + return + closing = None if loop.is_closed() else _hand_close_to(loop, entry) + if closing is None: + # Nothing touches the closed loop: its connections cannot be closed any more, and + # this only marks the session closed. + await entry.session.close() + return + await _wait_while_running(loop, closing) + # Not done, or cancelled: the loop stopped first, and closes the session as it shuts + # down its async generators. + if closing.done() and not closing.cancelled(): + error = closing.exception() + if error is not None: + raise error + + +def _hand_close_to( + loop: asyncio.AbstractEventLoop, entry: _LoopSession +) -> concurrent.futures.Future[None] | None: + """Start closing ``entry`` on ``loop``, from another thread; None if ``loop`` is closed. + + The coroutine is created on the loop: one the loop never runs, because it closes first, + would be reported as never awaited. + + Returns: + The future of the close, which is cancelled if the loop cancels the close. + """ + closing: concurrent.futures.Future[None] = concurrent.futures.Future() + try: + loop.call_soon_threadsafe(_start_closing_task, loop, entry, closing) + except RuntimeError: # the loop is closed + return None + return closing + + +def _start_closing_task( + loop: asyncio.AbstractEventLoop, + entry: _LoopSession, + closing: concurrent.futures.Future[None], +) -> None: + task = loop.create_task(_close_on_its_loop(entry)) + # The loop holds its tasks weakly: this keeps the task until it is done. + _closing_tasks.add(task) + task.add_done_callback(_closing_tasks.discard) + task.add_done_callback(functools.partial(_report_close, closing)) + + +async def _close_on_its_loop(entry: _LoopSession) -> None: + """Close the session, then its closer. + + In this order, a loop that shuts down its async generators while the session closes + finds the session closed already, rather than its closer running; and a session kept + after its closer failed to close it is closed all the same. + """ + await entry.session.close() + await entry.closer.aclose() + + +def _report_close(closing: concurrent.futures.Future[None], task: asyncio.Task[None]) -> None: + """Give ``closing`` the outcome of ``task``, the close it stands for.""" + if task.cancelled(): + closing.cancel() + elif (error := task.exception()) is not None: + closing.set_exception(error) + else: + closing.set_result(None) + + +async def _wait_while_running( + loop: asyncio.AbstractEventLoop, closing: concurrent.futures.Future[None] +) -> None: + """Wait until ``closing`` is done, or until ``loop``, which runs it, stops running.""" + here = asyncio.get_running_loop() + done = asyncio.Event() + + def wake(_: concurrent.futures.Future[None]) -> None: + with contextlib.suppress(RuntimeError): # this loop closed meanwhile + here.call_soon_threadsafe(done.set) + + closing.add_done_callback(wake) + while not done.is_set() and loop.is_running(): + with contextlib.suppress(asyncio.TimeoutError): + await asyncio.wait_for(done.wait(), _STOPPED_LOOP_POLL_SECONDS) + + +def _close_at_exit_on(loop: asyncio.AbstractEventLoop, entry: _LoopSession) -> bool: + """Close ``entry`` on ``loop`` as the interpreter exits; False if ``loop`` is closed. + + A loop running in another thread gets the close to run, and is not waited for: its + thread, if it is a daemon, may be stopped first, which leaves nothing to report. + """ + if loop.is_closed(): + return False + if loop.is_running(): + return _hand_close_to(loop, entry) is not None + loop.run_until_complete(_close_on_its_loop(entry)) + return True + + +def _orphan(sessions: dict[asyncio.AbstractEventLoop, _LoopSession]) -> None: + """Keep the sessions of a collected `LoopSessions` until they are closed. + + The session of a running loop is closed on that loop now; the session of a loop that is + not running is closed when that loop shuts down its async generators; one of a closed + loop is marked closed by the next request from any loop. As a finalizer, this may run + in any thread, so it only hands the closes to the loops. + """ + entries = list(sessions.items()) + sessions.clear() + for loop, entry in entries: + _orphaned[id(entry.session)] = (loop, entry) + if loop.is_running(): + _hand_close_to(loop, entry) + + +def _take_orphans_of_closed_loops() -> list[_LoopSession]: + """Remove and return the orphaned sessions of loops closed without shutting them down.""" + taken = [] + for key, (loop, entry) in list(_orphaned.items()): + if loop.is_closed() and _orphaned.pop(key, None) is not None: + taken.append(entry) + return taken + + +_open_at_exit: weakref.WeakSet[LoopSessions] = weakref.WeakSet() +# The open sessions no LoopSessions holds any more, by the id of the session: those of +# collected LoopSessions, and those close() kept while their loop had a new one. +_orphaned: dict[int, tuple[asyncio.AbstractEventLoop, _LoopSession]] = {} +_closing_tasks: set[asyncio.Task[None]] = set() +# How often close() looks whether a loop it waits for in another thread still runs. +_STOPPED_LOOP_POLL_SECONDS = 0.05 +_sessions_lost_to_fork: list[_LoopSession] = [] + + +@atexit.register +def _close_open_sessions_at_exit() -> None: + for sessions in list(_open_at_exit): + sessions._close_at_exit() # noqa: SLF001 - this module's own class + + +def _set_aside_sessions_after_fork() -> None: + for sessions in list(_open_at_exit): + sessions._set_aside_after_fork() # noqa: SLF001 - this module's own class + + +if sys.platform != "win32": + os.register_at_fork(after_in_child=_set_aside_sessions_after_fork) diff --git a/permit/utils/sync.py b/permit/utils/sync.py index 9e7572e7..3d80aa58 100644 --- a/permit/utils/sync.py +++ b/permit/utils/sync.py @@ -1,8 +1,14 @@ import asyncio +import atexit +import concurrent.futures +import contextlib import functools import inspect +import os import sys +import threading import warnings +import weakref from collections.abc import Awaitable, Callable, Coroutine from concurrent.futures import ThreadPoolExecutor from contextvars import ContextVar @@ -18,9 +24,14 @@ from typing_extensions import ParamSpec +from permit.utils.sdk_logger import sdk_logger + P = ParamSpec("P") T = TypeVar("T") +CloseSessions = Callable[[], Coroutine[Any, Any, None]] +"""A coroutine function that closes the HTTP sessions of a sync client.""" + SYNC_WRAPPER_MARKER = "__permit_sync_wrapper__" """Attribute set on every wrapper produced by :func:`async_to_sync`. @@ -122,6 +133,390 @@ def run_coroutine_sync(coroutine: Coroutine[Any, Any, T]) -> T: return _run_blocking(coroutine, _CallSite.from_frame(caller)) +_CALL_ON_LOOP_THREAD = ( + "A blocking call of permit.sync.Permit was made on the client's own event loop thread " + "({thread}), where it would wait for itself forever. Make the call from another thread, " + "or await the async client, permit.Permit." +) +_CLOSE_ON_LOOP_THREAD = ( + "permit.sync.Permit.close() was called on the client's own event loop thread ({thread}), " + "which close() stops and joins. Call it from another thread." +) + +_BACKGROUND_LOOP_ATTRIBUTE = "_permit_background_loop" +"""The attribute through which an object of a `SyncClass` class reaches its client's loop. + +`_BackgroundLoop.bind` sets it. A blocking method of an object without it runs its coroutine +in an event loop of its own, as every blocking call did before the sync client had a +background loop. +""" + + +class _Raised(NamedTuple): + """The exception a blocking call's coroutine raised, carried to the caller as a result. + + `asyncio.run_coroutine_threadsafe` copies an exception into the caller's future through + asyncio's own conversion, which on Python 3.11 and 3.12 replaces a `TimeoutError` with a + new one that has neither its traceback nor its cause. As a result it is not converted, so + the caller raises the exception the coroutine raised. + """ + + error: Exception + + +class _LoopThread: + """An event loop that runs in a daemon thread until it is shut down. + + It tracks the tasks it runs for blocking calls, so that a shutdown can wait for them, or + cancel them. + """ + + def __init__(self) -> None: + self.loop = asyncio.new_event_loop() + # Read and written on the loop's thread only. + self._tasks: set[asyncio.Task[Any]] = set() + self._stopping: asyncio.Task[None] | None = None + self.thread = threading.Thread(target=self._serve, name="permit-sync-loop", daemon=True) + self.thread.start() + + def _serve(self) -> None: + try: + self.loop.run_forever() + self.loop.run_until_complete(self._settle()) + self.loop.run_until_complete(self.loop.shutdown_asyncgens()) + finally: + self.loop.close() + + async def _settle(self) -> None: + """Finish the tasks still on the stopped loop: tracked ones run, any other is cancelled. + + Another task may be one that a call left running, or the close of an HTTP session + that a finalizer handed to the loop as it stopped; shutting down the loop's async + generators next closes any session such a close left open. + """ + current = asyncio.current_task() + while others := asyncio.all_tasks() - {current}: + for task in others - self._tasks: + task.cancel() + await asyncio.wait(others) + + async def _track(self, coroutine: Coroutine[Any, Any, T], call_site: _CallSite) -> T | _Raised: + """Await `coroutine` as a tracked task that runs for the blocking call made at `call_site`. + + The task runs in a copy of the context of the thread that submitted it, as + `call_soon_threadsafe` documents, and `_blocking_call_site` is set in that copy. + + Returns: + What the coroutine returns, or the exception it raises, as a `_Raised`. A + cancellation, of the task or from the coroutine, is raised. + """ + # Never None: this coroutine only ever runs as a task. + task = cast("asyncio.Task[Any]", asyncio.current_task()) + self._tasks.add(task) + try: + _blocking_call_site.set(call_site) + return await coroutine + except Exception as error: # noqa: BLE001 - the blocking caller raises it + return _Raised(error) + finally: + self._tasks.discard(task) + # An exception the coroutine raised keeps this frame in its traceback, and the + # task keeps the exception: without this, the three form a cycle that holds the + # coroutine's objects until the cyclic garbage collector runs. + del task + + def submit( + self, coroutine: Coroutine[Any, Any, T], call_site: _CallSite + ) -> concurrent.futures.Future[T | _Raised]: + """Start `coroutine` on the loop, for the blocking call made at `call_site`. + + Args: + coroutine: The coroutine to run. + call_site: The line that made the blocking call. + + Returns: + The future of the coroutine's result, or of the exception it raised, as a + `_Raised`. + + Raises: + RuntimeError: If the loop is closed. `coroutine` is closed, never started. + """ + tracked = self._track(coroutine, call_site) + try: + return asyncio.run_coroutine_threadsafe(tracked, self.loop) + except BaseException: + tracked.close() + coroutine.close() + raise + + async def drain(self, *, cancel: bool) -> None: + """Wait until no tracked task is left, cancelling each one first when `cancel` is True.""" + current = asyncio.current_task() + while pending := self._tasks - {current}: + if cancel: + for task in pending: + task.cancel() + await asyncio.wait(pending) + + async def _drain_and_close( + self, close_sessions: CloseSessions | None, *, cancel_calls: bool, call_site: _CallSite + ) -> None: + # The sessions' close() may await the client's own converted methods, which must hand + # back their coroutines rather than block, as they do in any blocking call's coroutine. + _blocking_call_site.set(call_site) + await self.drain(cancel=cancel_calls) + if close_sessions is not None: + await close_sessions() + + def close( + self, close_sessions: CloseSessions | None, *, cancel_calls: bool, call_site: _CallSite + ) -> None: + """Wait for (or cancel) the tracked tasks, run `close_sessions`, then stop and join. + + Args: + close_sessions: The coroutine function that closes the sessions opened on this + loop, if any. + cancel_calls: Cancel the blocking calls in flight instead of waiting for them. + call_site: The line that called close(). + """ + drained = self._drain_and_close( + close_sessions, cancel_calls=cancel_calls, call_site=call_site + ) + try: + asyncio.run_coroutine_threadsafe(drained, self.loop).result() + finally: + self.loop.call_soon_threadsafe(self.loop.stop) + self.thread.join() + + def stop_soon(self) -> None: + """Stop the loop once its tracked tasks are done, without waiting; safe in a finalizer.""" + # A closed loop raises; there is nothing left to stop then. + with contextlib.suppress(RuntimeError): + self.loop.call_soon_threadsafe(self._start_stopping) + + def _start_stopping(self) -> None: + self._stopping = self.loop.create_task(self._drain_and_stop()) + + async def _drain_and_stop(self) -> None: + await self.drain(cancel=False) + self.loop.stop() + + +class _BackgroundLoop: + """The event loop on which a sync client runs its blocking calls, in a daemon thread. + + The thread starts on the first call. Calls from any number of threads are submitted to + it and waited for, so they share the client's HTTP sessions and connections. `close()` + waits for the calls in flight, closes the sessions and stops the thread; the next call + starts a new one. While a `close()` runs, a call or another `close()` from another + thread waits for it to finish. A client that is never closed has its thread stopped + once nothing references this object any more, or at interpreter exit. + """ + + def __init__(self) -> None: + self._lock = threading.Lock() + # Notified, with the lock held, when a close() finishes. + self._closed = threading.Condition(self._lock) + self._thread: _LoopThread | None = None + # The thread a close() is stopping, until it has stopped. + self._closing: _LoopThread | None = None + self._stop_when_collected: weakref.finalize[[], _BackgroundLoop] | None = None + self._closer: weakref.WeakMethod[CloseSessions] | None = None + _background_loops.add(self) + + def bind(self, *roots: object) -> None: + """Run the blocking calls of `roots`, and of every `SyncClass` object they hold, here. + + Args: + roots: The objects a sync client exposes, such as its enforcer and API clients. + """ + pending = list(roots) + seen: set[int] = set() + while pending: + obj = pending.pop() + if id(obj) in seen: + continue + seen.add(id(obj)) + if isinstance(type(obj), SyncClass): + setattr(obj, _BACKGROUND_LOOP_ATTRIBUTE, self) + pending.extend( + value for value in vars(obj).values() if isinstance(type(value), SyncClass) + ) + + def set_closer(self, close_sessions: "weakref.WeakMethod[CloseSessions]") -> None: + """Run `close_sessions` on the loop when it is closed, while its client is alive. + + Args: + close_sessions: A weak reference to the client's method that closes its sessions. + """ + with self._lock: + self._closer = close_sessions + + def run(self, coroutine: Coroutine[Any, Any, T], call_site: _CallSite) -> T: + """Run `coroutine` on the loop for the blocking call made at `call_site`, and wait. + + Args: + coroutine: The coroutine of the blocking call. + call_site: The line that made the blocking call. + + Returns: + Whatever the coroutine returns. + + Raises: + RuntimeError: If called from the loop's own thread, where waiting would deadlock. + """ + future: concurrent.futures.Future[T | _Raised] | None = None + try: + with self._lock: + future = self._thread_for_call().submit(coroutine, call_site) + outcome = future.result() + except BaseException: + if future is None: + coroutine.close() + raise + finally: + if future is not None: + # A no-op once the call is done. When waiting was interrupted, such as by + # KeyboardInterrupt, it cancels the call, as asyncio.run() would. + future.cancel() + if not isinstance(outcome, _Raised): + return outcome + error = outcome.error + # The error's traceback will hold this frame: drop what would lead back to the error. + del outcome, future + try: + raise error + finally: + del error + + def _thread_for_call(self) -> _LoopThread: + """The loop thread to run a call on, started first if there is none. + + Called with the lock held. While a close() stops the thread, this waits for it, then + starts a new one. + + Raises: + RuntimeError: If the caller is the loop thread, or the one a close() is stopping, + where waiting would deadlock. + """ + while self._thread is None and self._closing is not None: + self._refuse_on(self._closing, _CALL_ON_LOOP_THREAD) + self._closed.wait() + if self._thread is None: + self._thread = _LoopThread() + stop_when_collected = weakref.finalize(self, self._thread.stop_soon) + # Only when collected: at exit, _close_running_loops closes the loop, with the + # client's sessions. Writable, as the weakref documentation says; typeshed + # declares __slots__ = () on it. + stop_when_collected.atexit = False # type: ignore[misc] + self._stop_when_collected = stop_when_collected + _running_loops.add(self) + else: + self._refuse_on(self._thread, _CALL_ON_LOOP_THREAD) + return self._thread + + @staticmethod + def _refuse_on(loop_thread: _LoopThread, message: str) -> None: + """Raise RuntimeError with `message` if the caller runs on `loop_thread`.""" + if loop_thread.thread is threading.current_thread(): + raise RuntimeError(message.format(thread=loop_thread.thread.name)) + + def close(self, *, cancel_calls: bool = False) -> None: + """Close the sessions opened on the loop and stop its thread, if it is running. + + A close() that another thread runs is waited for first. So when this returns, the + thread has stopped, unless a call started a new one since. + + Args: + cancel_calls: Cancel the blocking calls in flight instead of waiting for them. + + Raises: + RuntimeError: If called from the loop's own thread, which it would have to join. + """ + caller = sys._getframe(0).f_back # noqa: SLF001 - see run_coroutine_sync + call_site = _CallSite.from_frame(caller) + with self._lock: + while self._closing is not None: + self._refuse_on(self._closing, _CLOSE_ON_LOOP_THREAD) + self._closed.wait() + loop_thread = self._thread + if loop_thread is None: + return + self._refuse_on(loop_thread, _CLOSE_ON_LOOP_THREAD) + self._thread = None + self._closing = loop_thread + if self._stop_when_collected is not None: + self._stop_when_collected.detach() + self._stop_when_collected = None + _running_loops.discard(self) + close_sessions = None if self._closer is None else self._closer() + try: + loop_thread.close(close_sessions, cancel_calls=cancel_calls, call_site=call_site) + finally: + with self._lock: + self._closing = None + self._closed.notify_all() + + def forget_thread(self) -> None: + """In a child process made by fork(): drop the thread, which the fork did not copy. + + The next call starts a new thread. The old loop is kept referenced, not closed: it + still looks like it is running, so closing it would raise, and collecting it would + report it, and the sessions bound to it, as unclosed. + """ + self._lock = threading.Lock() + self._closed = threading.Condition(self._lock) + _loops_lost_to_fork.extend( + lost for lost in (self._thread, self._closing) if lost is not None + ) + self._thread = None + self._closing = None + if self._stop_when_collected is not None: + self._stop_when_collected.detach() + self._stop_when_collected = None + + +_running_loops: "weakref.WeakSet[_BackgroundLoop]" = weakref.WeakSet() +# Every background loop, including those a close() is stopping, which _running_loops leaves +# out so that the exit hook does not wait for them. +_background_loops: "weakref.WeakSet[_BackgroundLoop]" = weakref.WeakSet() +_loops_lost_to_fork: list[_LoopThread] = [] + + +def _close_running_loops() -> None: + """At interpreter exit, close every running background loop, with its client's sessions. + + The calls still in flight can only come from daemon threads by then, and are cancelled + rather than waited for, so they cannot hold up the exit. + """ + for background_loop in list(_running_loops): + _close_at_exit(background_loop) + + +def _close_at_exit(background_loop: _BackgroundLoop) -> None: + try: + background_loop.close(cancel_calls=True) + except Exception as error: # noqa: BLE001 - logged; the other clients still get closed + sdk_logger.error(f"Could not close a Permit sync client at exit: {error!r}") + + +def _forget_threads_after_fork() -> None: + for background_loop in list(_background_loops): + background_loop.forget_thread() + _running_loops.clear() + + +atexit.register(_close_running_loops) +if sys.platform != "win32": + os.register_at_fork(after_in_child=_forget_threads_after_fork) + + +def _background_loop_of(obj: object) -> _BackgroundLoop | None: + """The background loop `obj` was bound to, if it was.""" + candidate = getattr(obj, _BACKGROUND_LOOP_ATTRIBUTE, None) + return candidate if isinstance(candidate, _BackgroundLoop) else None + + def async_to_sync(func: Callable[P, Coroutine[Any, Any, T]]) -> Callable[P, T]: """Turn an async callable into a blocking one. @@ -129,10 +524,12 @@ def async_to_sync(func: Callable[P, Coroutine[Any, Any, T]]) -> Callable[P, T]: func: The coroutine function to convert. Returns: - A callable that runs `func` to completion and returns its result. When it - is called from inside a coroutine that a blocking call is already driving, - the coroutine is handed back untouched instead, so that internal - `await self.public_method(...)` calls keep working on a converted class. + A callable that runs `func` to completion and returns its result: on the background + loop of the sync client that the first argument (`self`, for a method) belongs to, + otherwise in an event loop of its own. When it is called from inside a coroutine + that a blocking call is already driving, the coroutine is handed back untouched + instead, so that internal `await self.public_method(...)` calls keep working on a + converted class. """ @wraps(func) @@ -142,7 +539,11 @@ def wrapper(*args: P.args, **kwargs: P.kwargs) -> T: # Read in the caller's thread, while its frame is the one that called us. caller = sys._getframe(0).f_back # noqa: SLF001 - see run_coroutine_sync call_site = _CallSite.from_frame(caller) - return _run_blocking(func(*args, **kwargs), call_site) + background_loop = _background_loop_of(args[0]) if args else None + coroutine = func(*args, **kwargs) + if background_loop is None: + return _run_blocking(coroutine, call_site) + return background_loop.run(coroutine, call_site) setattr(wrapper, SYNC_WRAPPER_MARKER, True) return wrapper diff --git a/pyproject.toml b/pyproject.toml index 3ef409a0..af954771 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -33,6 +33,10 @@ dependencies = [ # 0.7.3 is the first loguru release that imports without a DeprecationWarning on # Python 3.14: earlier ones call asyncio.iscoroutinefunction, which 3.16 removes. "loguru>=0.7.3,<1", + # permit/api/base.py imports multidict and yarl itself; aiohttp also depends on + # them. The floors are the first releases with wheels for every supported Python, + # 3.14 and 3.14t included, so a floor install never builds them from source. + "multidict>=6.7.0,<7", # pydantic has one line per Python range, updated by hand: Dependabot ignores # pydantic (see .github/dependabot.yml). Why each version is excluded: # - CVE-2024-3772 (ReDoS in email validation) affects pydantic 1.x before @@ -66,6 +70,7 @@ dependencies = [ # before 4.6 break `import permit` on 3.12+, before 4.12 on 3.13+, and 4.12-4.13 # lose TypedDict keys on 3.14. "typing-extensions>=4.14.0,<5", + "yarl>=1.21.0,<2", ] [project.urls] diff --git a/tests/benchmark_connection_reuse.py b/tests/benchmark_connection_reuse.py new file mode 100644 index 00000000..2d371859 --- /dev/null +++ b/tests/benchmark_connection_reuse.py @@ -0,0 +1,110 @@ +"""Benchmark: sequential check() calls against a local server that counts TCP connections. + +Run it from the repository root, with the permit package to measure on the path: + + uv run --locked python -m tests.benchmark_connection_reuse --calls 200 + +For the async and the sync client, it makes that many check() calls one after the other, +and prints how many TCP connections they opened and how long each call took. A client that +keeps its HTTP session opens one connection; one that opens a session per call opens one +per call. The server answers on 127.0.0.1 without delay, so the times are the client's own +cost; against a remote PDP each new connection also pays a network round trip, and a TLS +handshake over https. +""" + +import argparse +import asyncio +import statistics +import time +from pathlib import Path +from typing import TYPE_CHECKING + +from loguru import logger + +import permit +from permit import Permit +from permit.sync import Permit as SyncPermit +from tests.keepalive_server import KeepAliveServer +from tests.utils import offline_config + +if TYPE_CHECKING: + from collections.abc import Awaitable, Callable + +COLUMNS = ("client", "calls", "connections", "total ms", "mean ms", "p50 ms", "p95 ms") + + +def time_async_client(url: str, calls: int) -> list[float]: + """The duration of each of `calls` sequential `await permit.check()` calls, in seconds.""" + + async def main() -> list[float]: + client = Permit(offline_config(url)) + durations = [] + for _ in range(calls): + start = time.perf_counter() + allowed = await client.check("user", "read", "document") + durations.append(time.perf_counter() - start) + assert allowed, "the server answers every check with allow: true" + # The benchmark compares versions of permit, and the earlier ones have no close(). + close: Callable[[], Awaitable[None]] | None = getattr(client, "close", None) + if close is not None: + await close() + return durations + + return asyncio.run(main()) + + +def time_sync_client(url: str, calls: int) -> list[float]: + """The duration of each of `calls` sequential blocking `permit.check()` calls, in seconds.""" + client = SyncPermit(offline_config(url)) + durations = [] + for _ in range(calls): + start = time.perf_counter() + allowed = client.check("user", "read", "document") + durations.append(time.perf_counter() - start) + assert allowed, "the server answers every check with allow: true" + # The benchmark compares versions of permit, and the earlier ones have no close(). + close: Callable[[], None] | None = getattr(client, "close", None) + if close is not None: + close() + return durations + + +def row(client: str, connections: int, durations: list[float]) -> tuple[str, ...]: + """One line of the report, with the times in milliseconds.""" + milliseconds = [duration * 1000 for duration in durations] + p95 = statistics.quantiles(milliseconds, n=20)[18] + return ( + client, + str(len(durations)), + str(connections), + f"{sum(milliseconds):.1f}", + f"{statistics.fmean(milliseconds):.3f}", + f"{statistics.median(milliseconds):.3f}", + f"{p95:.3f}", + ) + + +def main() -> None: + """Measure both clients and print the report.""" + parser = argparse.ArgumentParser(description=__doc__.splitlines()[0]) + parser.add_argument("--calls", type=int, default=200, help="check() calls per client") + args = parser.parse_args() + if args.calls < 2: + parser.error("--calls must be at least 2") + + # Only the report goes to the terminal, not the SDK's debug records. + logger.disable("permit") + rows: list[tuple[str, ...]] = [COLUMNS] + for name, measure in (("async", time_async_client), ("sync", time_sync_client)): + with KeepAliveServer() as server: + durations = measure(server.url, args.calls) + rows.append(row(name, server.opened, durations)) + + print(f"permit from {Path(permit.__file__).parent}") + widths = [max(len(line[column]) for line in rows) for column in range(len(COLUMNS))] + for line in rows: + print(" ".join(cell.rjust(width) for cell, width in zip(line, widths, strict=True))) + + +if __name__ == "__main__": + main() diff --git a/tests/keepalive_server.py b/tests/keepalive_server.py new file mode 100644 index 00000000..31de33cf --- /dev/null +++ b/tests/keepalive_server.py @@ -0,0 +1,213 @@ +"""A local HTTP/1.1 server that keeps its connections open and counts them. + +pytest-httpserver answers in HTTP/1.0 and closes each connection after its response, so it +cannot show whether a client reuses connections. This server keeps every connection open +until the client closes it, as the Permit API and the PDP do. It counts the connections it +accepted and those that were closed, and records the requests it read. +""" + +import asyncio +import contextlib +import json +import threading +from types import TracebackType +from typing import NamedTuple + +from typing_extensions import Self + +# How long the server's own startup and shutdown, and a test's wait, may take. +_SERVER_TIMEOUT_SECONDS = 5.0 +_HEADER_END = b"\r\n\r\n" + + +class ServedRequest(NamedTuple): + """A request the server read: its method, path and headers (names as sent).""" + + method: str + path: str + headers: dict[str, str] + + +class _Response(NamedTuple): + body: bytes + headers: dict[str, str] + delay: float + + +_ALLOW = _Response(b'{"allow": true}', {}, 0.0) + + +class KeepAliveServer: + """Answers JSON on 127.0.0.1, from an event loop in a thread of its own. + + Every path answers ``{"allow": true}`` unless ``respond()`` set another answer for it. + Use it as a context manager, or call ``start()`` and ``stop()``. + """ + + def __init__(self) -> None: + self._loop = asyncio.new_event_loop() + self._thread = threading.Thread( + target=self._loop.run_forever, name="keepalive-server", daemon=True + ) + self._changed = threading.Condition() + self._opened = 0 + self._closed = 0 + self._requests: list[ServedRequest] = [] + self._responses: dict[str, _Response] = {} + self._server: asyncio.Server | None = None + # Read and written on the server's loop only. + self._handlers: set[asyncio.Task[None]] = set() + self._writers: set[asyncio.StreamWriter] = set() + + def start(self) -> None: + """Start serving on a free port.""" + self._thread.start() + self._server = asyncio.run_coroutine_threadsafe( + asyncio.start_server(self._serve, "127.0.0.1", 0), self._loop + ).result(_SERVER_TIMEOUT_SECONDS) + + def stop(self) -> None: + """Close every connection, stop serving and stop the server's thread.""" + asyncio.run_coroutine_threadsafe(self._shut_down(), self._loop).result( + _SERVER_TIMEOUT_SECONDS + ) + self._loop.call_soon_threadsafe(self._loop.stop) + self._thread.join(_SERVER_TIMEOUT_SECONDS) + self._loop.close() + + def __enter__(self) -> Self: + self.start() + return self + + def __exit__( + self, + exc_type: type[BaseException] | None, + exc_value: BaseException | None, + traceback: TracebackType | None, + ) -> None: + self.stop() + + @property + def url(self) -> str: + """The server's base URL.""" + assert self._server is not None + host, port = self._server.sockets[0].getsockname()[:2] + return f"http://{host}:{port}" + + @property + def opened(self) -> int: + """How many connections the server accepted.""" + with self._changed: + return self._opened + + @property + def closed(self) -> int: + """How many of those connections were closed, by the client or by a failed write.""" + with self._changed: + return self._closed + + @property + def requests(self) -> list[ServedRequest]: + """The requests the server read, in the order it read them.""" + with self._changed: + return list(self._requests) + + def respond( + self, + path: str, + body: object, + *, + headers: dict[str, str] | None = None, + delay: float = 0.0, + ) -> None: + """Answer requests to ``path`` with ``body`` as JSON, after ``delay`` seconds. + + A client that closes the connection during the delay gets no answer, and the + connection counts as closed then. + """ + self._responses[path] = _Response(json.dumps(body).encode(), headers or {}, delay) + + def wait_until_closed(self, count: int, timeout: float = _SERVER_TIMEOUT_SECONDS) -> int: + """Wait until ``count`` connections were closed, and return how many were. + + It returns once they are, or once ``timeout`` seconds have passed. + """ + with self._changed: + self._changed.wait_for(lambda: self._closed >= count, timeout) + return self._closed + + def wait_for_requests(self, count: int, timeout: float = _SERVER_TIMEOUT_SECONDS) -> bool: + """Wait until the server has read ``count`` requests; False if ``timeout`` passes first.""" + with self._changed: + return self._changed.wait_for(lambda: len(self._requests) >= count, timeout) + + async def _serve(self, reader: asyncio.StreamReader, writer: asyncio.StreamWriter) -> None: + handler = asyncio.current_task() + if handler is not None: + self._handlers.add(handler) + self._writers.add(writer) + with self._changed: + self._opened += 1 + self._changed.notify_all() + try: + while await self._answer_one(reader, writer): + pass + except (asyncio.IncompleteReadError, ConnectionError): + pass # The client closed the connection mid-request, which ends it. + finally: + self._writers.discard(writer) + writer.close() + with contextlib.suppress(ConnectionError): + await writer.wait_closed() + with self._changed: + self._closed += 1 + self._changed.notify_all() + if handler is not None: + self._handlers.discard(handler) + + async def _answer_one(self, reader: asyncio.StreamReader, writer: asyncio.StreamWriter) -> bool: + """Answer the connection's next request; False once the client closed it.""" + try: + head = await reader.readuntil(_HEADER_END) + except asyncio.IncompleteReadError: + return False + request_line, *header_lines = head.decode("latin-1").rstrip("\r\n").split("\r\n") + method, target, _ = request_line.split(" ") + headers = dict(line.split(": ", 1) for line in header_lines) + length = next((v for k, v in headers.items() if k.lower() == "content-length"), "0") + await reader.readexactly(int(length)) + path = target.split("?", 1)[0] + with self._changed: + self._requests.append(ServedRequest(method, path, headers)) + self._changed.notify_all() + + response = self._responses.get(path, _ALLOW) + if response.delay and await _closed_within(reader, response.delay): + return False + extra_headers = "".join(f"{name}: {value}\r\n" for name, value in response.headers.items()) + writer.write( + f"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\n" + f"Content-Length: {len(response.body)}\r\n{extra_headers}\r\n".encode("latin-1") + + response.body + ) + await writer.drain() + return True + + async def _shut_down(self) -> None: + assert self._server is not None + self._server.close() + for writer in list(self._writers): + writer.close() + await asyncio.gather(*self._handlers, return_exceptions=True) + await self._server.wait_closed() + + +async def _closed_within(reader: asyncio.StreamReader, seconds: float) -> bool: + """Whether the client closes the connection within ``seconds``, sending nothing meanwhile. + + A client waiting for its response sends nothing, so this reads nothing it needs later. + """ + try: + return await asyncio.wait_for(reader.read(1), seconds) == b"" + except asyncio.TimeoutError: + return False diff --git a/tests/test_async_session_lifecycle.py b/tests/test_async_session_lifecycle.py new file mode 100644 index 00000000..76fbcd63 --- /dev/null +++ b/tests/test_async_session_lifecycle.py @@ -0,0 +1,957 @@ +"""Offline tests of how the async client reuses and closes its HTTP connections (PER-16344). + +The client sends its requests through one aiohttp session, with its pool of open +connections, for the Permit API and one for the PDP, per event loop. The connection tests +count connections at a local server that keeps them open between requests, as the Permit +API and the PDP do. +""" + +import asyncio +import gc +import os +import subprocess +import sys +import textwrap +import threading +import time +import warnings +import weakref +from collections.abc import Callable, Coroutine, Iterator +from pathlib import Path +from typing import Any + +import aiohttp +import pytest +from pytest_httpserver import HTTPServer +from yarl import URL + +from permit import Permit, PermitConfig +from permit.api.base import BasePermitApi, SimpleHttpClient +from permit.pdp_api.base import BasePdpPermitApi +from permit.sync import Permit as SyncPermit +from permit.utils.http_sessions import LoopSessions +from tests.keepalive_server import KeepAliveServer, ServedRequest +from tests.utils import FACTS, offline_config + +REPO_ROOT = Path(__file__).resolve().parents[1] +USERS_PAGE = {"data": [], "total_count": 0, "page_count": 0} +# How long a test waits for a loop in another thread. +THREAD_TIMEOUT_SECONDS = 5.0 + + +@pytest.fixture +def server() -> Iterator[KeepAliveServer]: + with KeepAliveServer() as server: + yield server + + +@pytest.fixture +def client(server: KeepAliveServer) -> Permit: + server.respond(f"{FACTS}/users", USERS_PAGE) + return Permit(offline_config(server.url)) + + +def header(request: ServedRequest, name: str) -> str | None: + """The value of the request's header ``name``, matched case-insensitively.""" + return next( + (value for key, value in request.headers.items() if key.lower() == name.lower()), None + ) + + +async def check(client: Permit) -> bool: + return await client.check("user-1", "read", "document") + + +# --- connection reuse --------------------------------------------------------- + + +async def test_sequential_checks_share_one_connection( + server: KeepAliveServer, client: Permit +) -> None: + for _ in range(5): + assert await check(client) + + assert (len(server.requests), server.opened) == (5, 1) + + +async def test_the_api_calls_share_one_connection_and_the_pdp_calls_another( + server: KeepAliveServer, client: Permit +) -> None: + server.respond("/allowed/bulk", {"allow": [{"allow": True}]}) + server.respond("/user-permissions", {}) + server.respond("/user-tenants", []) + server.respond("/local/role_assignments", []) + server.respond("/v2/auth/elements_login_as", {"redirect_url": "http://elements.test/login"}) + server.respond(f"{FACTS}/tenants", []) + + assert await check(client) + assert await client.bulk_check([{"user": "user-1", "action": "read", "resource": "doc"}]) + await client.get_user_permissions("user-1") + await client.get_user_tenants("user-1") + await client.pdp_api.role_assignments.list() + pdp_connections = server.opened + await client.api.users.list() + await client.api.tenants.list() + await client.elements.login_as("user-1", "tenant-1") + + assert (len(server.requests), pdp_connections, server.opened) == (8, 1, 2) + + +@pytest.mark.parametrize( + "client_class", [pytest.param(Permit, id="async"), pytest.param(SyncPermit, id="sync")] +) +def test_every_api_of_a_client_sends_through_the_client_sessions( + client_class: type[Permit], +) -> None: + """Each of the client's APIs, and each API they hold, uses the client's sessions.""" + config = offline_config("http://localhost:1") + config.proxy_facts_via_pdp = True + client = client_class(config) + + with client.wait_for_sync() as waiting: + for each in (client, waiting): + api = sessions_reachable_from(each._api) | sessions_reachable_from(each._elements) + pdp = sessions_reachable_from(each._enforcer) | sessions_reachable_from(each._pdp_api) + assert api == {id(client._api_sessions)} + assert pdp == {id(client._pdp_sessions)} + + +def sessions_reachable_from(root: object) -> set[int]: + """The ids of the ``LoopSessions`` that ``root`` and every API and client it holds use.""" + found: set[int] = set() + pending = [root] + while pending: + holder = pending.pop() + for value in vars(holder).values(): + if isinstance(value, LoopSessions): + found.add(id(value)) + elif isinstance(value, (BasePermitApi, BasePdpPermitApi, SimpleHttpClient)): + pending.append(value) + return found + + +async def test_a_request_replaces_a_session_closed_under_the_client( + server: KeepAliveServer, client: Permit +) -> None: + """A request never goes through a closed session. + + close() keeps the session of a loop that stopped before closing it, and the close handed + to that loop closes it once the loop runs again. + """ + assert await check(client) + await (await client._pdp_sessions.current()).close() + + assert await check(client) + assert server.opened == 2 + + +async def test_the_connections_are_not_capped_in_number(client: Permit) -> None: + """As when every request had a session of its own, any number may be open at once.""" + session = await client._pdp_sessions.current() + + assert session.connector is not None + assert session.connector.limit == 0 + + +async def test_a_cookie_the_server_sets_is_not_sent_back(server: KeepAliveServer) -> None: + """The shared session keeps no cookies, so each request carries the headers it did alone.""" + server.respond("/allowed", {"allow": True}, headers={"Set-Cookie": "balancer=a1; Path=/"}) + # By a host name: aiohttp's cookie jar ignores cookies from an IP address. + client = Permit(offline_config(server.url.replace("127.0.0.1", "localhost"))) + + assert await check(client) + assert await check(client) + + assert [header(request, "Cookie") for request in server.requests] == [None, None] + + +# --- event loops --------------------------------------------------------------- + + +def test_one_client_serves_two_successive_asyncio_runs(server: KeepAliveServer) -> None: + client = Permit(offline_config(server.url)) + + async def three_checks() -> list[bool]: + return [await check(client) for _ in range(3)] + + assert asyncio.run(three_checks()) == [True] * 3 + # asyncio.run() closed the connection as it shut its loop down, without close(). + assert server.wait_until_closed(1) == 1 + assert asyncio.run(three_checks()) == [True] * 3 + assert server.wait_until_closed(2) == 2 + assert server.opened == 2 + + +@pytest.mark.parametrize("close", [False, True], ids=["left open", "closed"]) +def test_a_client_does_not_keep_a_finished_loop_alive( + server: KeepAliveServer, *, close: bool +) -> None: + client = Permit(offline_config(server.url)) + loops: list[weakref.ref[asyncio.AbstractEventLoop]] = [] + + async def remember_the_loop_and_check() -> bool: + loops.append(weakref.ref(asyncio.get_running_loop())) + allowed = await check(client) + if close: + await client.close() + return allowed + + assert asyncio.run(remember_the_loop_and_check()) + gc.collect() + + assert loops[0]() is None + assert asyncio.run(check(client)) + + +async def test_a_client_dropped_without_close_closes_its_connection( + server: KeepAliveServer, +) -> None: + """Dropping the last reference to a client closes its connections on their loop.""" + client = Permit(offline_config(server.url)) + assert await check(client) + + # Reference counting alone must free the client: nothing may hold it in a cycle. + gc.disable() + try: + del client + # Waits in another thread, so that this loop runs the close it was given. + assert await asyncio.to_thread(server.wait_until_closed, 1) == 1 + finally: + gc.enable() + + +async def test_a_client_freed_by_the_cycle_collector_closes_its_connection_without_a_warning( + server: KeepAliveServer, +) -> None: + """A client held in a reference cycle, as an exception it raised can hold it, is freed by gc. + + The client's sessions must not be collected with it while they are open, or aiohttp + reports them unclosed. + """ + client = Permit(offline_config(server.url)) + assert await check(client) + cycle: list[object] = [client] + cycle.append(cycle) + del client, cycle + + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + gc.collect() + # Waits in another thread, so that this loop runs the close it was given. + assert await asyncio.to_thread(server.wait_until_closed, 1) == 1 + + assert [f"{w.category.__name__}: {w.message}" for w in caught] == [] + + +def assert_nothing_reported_unclosed(drop: Callable[[], None]) -> None: + """Run ``drop`` and a garbage collection, and check aiohttp reported nothing unclosed.""" + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + drop() + gc.collect() + assert [str(warning.message) for warning in caught] == [] + + +def run_on_a_loop_closed_without_shutting_down(coroutine: Coroutine[Any, Any, bool]) -> bool: + """Run ``coroutine`` the old way: on a new loop, closed without shutting down its generators.""" + loop = asyncio.new_event_loop() + try: + return loop.run_until_complete(coroutine) + finally: + loop.close() + + +# pytest-httpserver closes each connection after its response, so the closed loops in these +# tests keep no connection open, only their session. + + +def test_the_next_request_closes_the_session_of_a_loop_closed_without_shutting_down( + httpserver: HTTPServer, config: PermitConfig +) -> None: + httpserver.expect_request("/allowed", method="POST").respond_with_json({"allow": True}) + client = Permit(config) + assert run_on_a_loop_closed_without_shutting_down(check(client)) + + assert asyncio.run(check(client)) + + def drop() -> None: + nonlocal client + del client + + assert_nothing_reported_unclosed(drop) + + +# --- per-request settings ------------------------------------------------------ + + +async def call_check(client: Permit) -> object: + return await check(client) + + +async def call_pdp_api(client: Permit) -> object: + return await client.pdp_api.role_assignments.list() + + +async def call_api(client: Permit) -> object: + return await client.api.users.list() + + +@pytest.mark.parametrize( + ("timeout_setting", "path", "call"), + [ + ("pdp_timeout", "/allowed", call_check), + ("pdp_timeout", "/local/role_assignments", call_pdp_api), + ("api_timeout", f"{FACTS}/users", call_api), + ], + ids=["check", "pdp_api", "api"], +) +async def test_a_request_slower_than_its_timeout_fails( + server: KeepAliveServer, + timeout_setting: str, + path: str, + call: Callable[[Permit], Coroutine[Any, Any, object]], +) -> None: + """Each request carries the client's timeout, now that the sessions are shared.""" + server.respond(path, [], delay=1.5) + config = offline_config(server.url) + setattr(config, timeout_setting, 1) + client = Permit(config) + + with pytest.raises(asyncio.TimeoutError): + await call(client) + + +async def test_headers_given_to_a_request_go_over_the_client_headers( + httpserver: HTTPServer, +) -> None: + httpserver.expect_request(f"{FACTS}/users", method="GET").respond_with_json(USERS_PAGE) + client = SimpleHttpClient( + { + "base_url": httpserver.url_for("/"), + "headers": {"Content-Type": "application/json", "Authorization": "Bearer a"}, + }, + base_url=FACTS, + ) + + await client.get("/users", model=dict, headers={"authorization": "Bearer b", "X-Extra": "1"}) + + [(request, _)] = httpserver.log + sent = { + name: request.headers.get(name) for name in ("Content-Type", "Authorization", "X-Extra") + } + assert sent == {"Content-Type": "application/json", "Authorization": "Bearer b", "X-Extra": "1"} + + +def test_a_client_config_option_requests_cannot_carry_is_refused() -> None: + with pytest.raises(TypeError, match=r"\['cookies'\]"): + SimpleHttpClient({"headers": {}, "cookies": {"session": "a"}}) + + +BASE_URLS = [ + "http://pdp.test", + "http://pdp.test/", + "http://pdp.test:7766/prefix/", + URL("http://pdp.test/prefix/"), + "http://pdp.test/prefix", + URL("http://pdp.test/prefix"), + "pdp.test:7766", + "//pdp.test/", + "", +] +PATHS = ["/v2/facts/users", "v2/facts/users", "", "http://elsewhere.test/v2/x"] + + +@pytest.mark.parametrize("path", PATHS) +@pytest.mark.parametrize("base_url", BASE_URLS, ids=repr) +async def test_a_request_url_resolves_as_a_session_base_url_resolved_it( + base_url: str | URL, path: str +) -> None: + """The base URL moved from the session to each request without changing where it goes.""" + + async def through_a_session() -> URL: + async with aiohttp.ClientSession(base_url=base_url) as session: + return session._build_url(path) + + async def through_the_client() -> URL: + return SimpleHttpClient({"base_url": base_url})._request_url(path) + + assert await outcome(through_the_client) == await outcome(through_a_session) + + +async def outcome(resolve: Callable[[], Coroutine[Any, Any, URL]]) -> URL | tuple[type, str]: + """What ``resolve`` returns, or the type and message of what it raises.""" + try: + return await resolve() + except ValueError as error: + return type(error), str(error) + + +# --- close() and the context manager ------------------------------------------- + + +async def test_close_closes_the_connections(server: KeepAliveServer, client: Permit) -> None: + assert await check(client) + await client.api.users.list() + + await client.close() + + assert server.wait_until_closed(2) == 2 + + +async def test_async_with_yields_the_client_and_closes_it_on_exit( + server: KeepAliveServer, +) -> None: + client = Permit(offline_config(server.url)) + + async with client as entered: + assert entered is client + assert await check(client) + + assert server.wait_until_closed(1) == 1 + + +async def test_async_with_closes_the_client_when_the_block_raises( + server: KeepAliveServer, +) -> None: + async def check_then_fail() -> None: + async with Permit(offline_config(server.url)) as client: + assert await check(client) + raise LookupError + + with pytest.raises(LookupError): + await check_then_fail() + + assert server.wait_until_closed(1) == 1 + + +async def test_closing_twice_closes_nothing_more(server: KeepAliveServer, client: Permit) -> None: + assert await check(client) + + await client.close() + await client.close() + + assert server.wait_until_closed(1) == 1 + assert server.opened == 1 + + +async def test_close_on_a_client_that_sent_nothing_does_nothing(server: KeepAliveServer) -> None: + await Permit(offline_config(server.url)).close() + + assert (server.opened, server.closed) == (0, 0) + + +async def test_a_request_after_close_opens_a_new_connection( + server: KeepAliveServer, client: Permit +) -> None: + assert await check(client) + await client.close() + + assert await check(client) + assert await check(client) + + assert server.opened == 2 + assert server.wait_until_closed(1) == 1 + await client.close() + assert server.wait_until_closed(2) == 2 + + +async def test_a_wait_for_sync_copy_shares_the_connection_and_leaves_closing_it_to_its_client( + server: KeepAliveServer, +) -> None: + config = offline_config(server.url) + config.proxy_facts_via_pdp = True + client = Permit(config) + + with client.wait_for_sync(timeout=3.0, policy="fail") as waiting: + await waiting.api.tenants.delete("tenant-1") + await waiting.close() + await client.api.tenants.delete("tenant-2") + + waited, not_waited = server.requests + assert (header(waited, "X-Wait-Timeout"), header(waited, "X-Timeout-Policy")) == ("3.0", "fail") + assert (header(not_waited, "X-Wait-Timeout"), header(not_waited, "X-Timeout-Policy")) == ( + None, + None, + ) + # The copy's close() left the connection open: the client's request went over it. + assert server.opened == 1 + + await client.close() + assert server.wait_until_closed(1) == 1 + # The copy still works after its client closed: it opens a new connection. + await waiting.api.tenants.delete("tenant-3") + assert server.opened == 2 + await client.close() + assert server.wait_until_closed(2) == 2 + + +async def test_without_proxy_facts_via_pdp_wait_for_sync_yields_the_client_itself( + server: KeepAliveServer, client: Permit +) -> None: + """So the yielded client's close() closes the connections, unlike a copy's.""" + assert not client.config.proxy_facts_via_pdp + assert await check(client) + + with client.wait_for_sync() as waiting: + assert waiting is client + await waiting.close() + + assert server.wait_until_closed(1) == 1 + + +async def test_close_also_closes_the_connection_of_a_loop_running_in_another_thread( + server: KeepAliveServer, client: Permit +) -> None: + loop = asyncio.new_event_loop() + thread = threading.Thread(target=loop.run_forever, daemon=True) + thread.start() + try: + assert asyncio.run_coroutine_threadsafe(check(client), loop).result(THREAD_TIMEOUT_SECONDS) + assert await check(client) + assert server.opened == 2 + other_session = asyncio.run_coroutine_threadsafe( + client._pdp_sessions.current(), loop + ).result(THREAD_TIMEOUT_SECONDS) + # Keep the other loop busy, so its session is closed only if close() waits for it. + loop.call_soon_threadsafe(time.sleep, 0.2) + + await client.close() + + assert other_session.closed + assert server.wait_until_closed(2) == 2 + finally: + loop.call_soon_threadsafe(loop.stop) + thread.join(THREAD_TIMEOUT_SECONDS) + loop.close() + + +def test_close_leaves_the_connection_of_an_idle_loop_to_that_loop( + server: KeepAliveServer, client: Permit +) -> None: + """A loop that is not running cannot close its session from another loop's close().""" + idle = asyncio.new_event_loop() + try: + assert idle.run_until_complete(check(client)) + + asyncio.run(client.close()) + # The idle loop's connection is still open: its next request goes over it. + assert idle.run_until_complete(check(client)) + assert server.opened == 1 + + idle.run_until_complete(client.close()) + assert server.wait_until_closed(1) == 1 + finally: + idle.close() + + +def keep_the_loop_busy(busy: threading.Event, seconds: float = 0.5) -> None: + """Block the running loop for ``seconds``, as a slow callback does: what it is handed waits.""" + busy.set() + time.sleep(seconds) + + +async def test_close_returns_when_a_loop_in_another_thread_ends_before_closing_its_session( + server: KeepAliveServer, client: Permit +) -> None: + """That loop's asyncio.run() cancels the close handed to it, and closes the session itself.""" + assert await check(client) + busy = threading.Event() + + async def check_then_end_busy() -> None: + assert await check(client) + keep_the_loop_busy(busy) + + other = threading.Thread(target=asyncio.run, args=(check_then_end_busy(),)) + other.start() + try: + assert await asyncio.to_thread(busy.wait, THREAD_TIMEOUT_SECONDS) + await asyncio.wait_for(client.close(), THREAD_TIMEOUT_SECONDS) + finally: + await asyncio.to_thread(other.join, THREAD_TIMEOUT_SECONDS) + + assert await asyncio.to_thread(server.wait_until_closed, 2) == 2 + + +async def test_close_returns_when_a_loop_in_another_thread_stops_before_closing_its_session( + server: KeepAliveServer, client: Permit, caplog: pytest.LogCaptureFixture +) -> None: + """The stopped loop closes the session as it shuts down its async generators.""" + assert await check(client) + loop = asyncio.new_event_loop() + + def run_then_shut_down() -> None: + loop.run_forever() + loop.run_until_complete(loop.shutdown_asyncgens()) + loop.close() + + other = threading.Thread(target=run_then_shut_down, daemon=True) + other.start() + busy = threading.Event() + + def stay_busy_then_stop() -> None: + keep_the_loop_busy(busy) + loop.stop() + + try: + assert asyncio.run_coroutine_threadsafe(check(client), loop).result(THREAD_TIMEOUT_SECONDS) + loop.call_soon_threadsafe(stay_busy_then_stop) + assert await asyncio.to_thread(busy.wait, THREAD_TIMEOUT_SECONDS) + await asyncio.wait_for(client.close(), THREAD_TIMEOUT_SECONDS) + finally: + await asyncio.to_thread(other.join, THREAD_TIMEOUT_SECONDS) + + assert loop.is_closed() + assert await asyncio.to_thread(server.wait_until_closed, 2) == 2 + assert [record.getMessage() for record in caplog.records if record.name == "asyncio"] == [] + + +def test_close_closes_every_other_session_when_one_loop_ends_before_closing_its_own( + httpserver: HTTPServer, config: PermitConfig +) -> None: + httpserver.expect_request("/allowed", method="POST").respond_with_json({"allow": True}) + client = Permit(config) + busy = threading.Event() + + async def check_then_end_busy() -> None: + assert await check(client) + keep_the_loop_busy(busy) + + # The session of the loop that ends first is the first the client opened. + other = threading.Thread(target=asyncio.run, args=(check_then_end_busy(),)) + other.start() + try: + assert busy.wait(THREAD_TIMEOUT_SECONDS) + assert run_on_a_loop_closed_without_shutting_down(check(client)) + asyncio.run(asyncio.wait_for(client.close(), THREAD_TIMEOUT_SECONDS)) + finally: + other.join(THREAD_TIMEOUT_SECONDS) + + def drop() -> None: + nonlocal client + del client + + assert_nothing_reported_unclosed(drop) + + +def test_close_keeps_the_session_of_a_loop_that_stops_and_closes_before_closing_it( + httpserver: HTTPServer, config: PermitConfig +) -> None: + """The next request marks it closed, as it does the session of any closed loop.""" + httpserver.expect_request("/allowed", method="POST").respond_with_json({"allow": True}) + client = Permit(config) + loop = asyncio.new_event_loop() + + def run_then_close() -> None: + loop.run_forever() + loop.close() + + other = threading.Thread(target=run_then_close, daemon=True) + busy = threading.Event() + + def stay_busy_then_stop() -> None: + keep_the_loop_busy(busy) + loop.stop() + + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + other.start() + assert asyncio.run_coroutine_threadsafe(check(client), loop).result(THREAD_TIMEOUT_SECONDS) + loop.call_soon_threadsafe(stay_busy_then_stop) + assert busy.wait(THREAD_TIMEOUT_SECONDS) + + asyncio.run(asyncio.wait_for(client.close(), THREAD_TIMEOUT_SECONDS)) + other.join(THREAD_TIMEOUT_SECONDS) + gc.collect() + assert asyncio.run(check(client)) + del client + gc.collect() + + assert loop.is_closed() + assert [f"{w.category.__name__}: {w.message}" for w in caught] == [] + + +async def test_the_next_close_closes_a_session_close_failed_to_close( + server: KeepAliveServer, client: Permit, monkeypatch: pytest.MonkeyPatch +) -> None: + assert await check(client) + session = await client._pdp_sessions.current() + + async def fail(_: aiohttp.ClientSession) -> None: + msg = "the session did not close" + raise OSError(msg) + + monkeypatch.setattr(aiohttp.ClientSession, "close", fail) + with pytest.raises(OSError, match="the session did not close"): + await client.close() + monkeypatch.undo() + assert server.closed == 0 + + await client.close() + + assert session.closed + assert await asyncio.to_thread(server.wait_until_closed, 1) == 1 + + +def test_close_still_closes_the_other_sessions_when_one_fails_to_close( + httpserver: HTTPServer, config: PermitConfig, monkeypatch: pytest.MonkeyPatch +) -> None: + """The session that failed to close is kept, and the next close() closes it.""" + httpserver.expect_request("/allowed", method="POST").respond_with_json({"allow": True}) + client = Permit(config) + # The first session the client opens is the one that fails to close. Its loop closes + # only once the other session is open: opening a session closes those of closed loops. + closed_later = asyncio.new_event_loop() + assert closed_later.run_until_complete(check(client)) + [failing] = [entry.session for entry in client._pdp_sessions._sessions.values()] + close_session = aiohttp.ClientSession.close + + async def close_or_fail(session: aiohttp.ClientSession) -> None: + if session is failing: + msg = "the session did not close" + raise OSError(msg) + await close_session(session) + + loop = asyncio.new_event_loop() + thread = threading.Thread(target=loop.run_forever, daemon=True) + thread.start() + try: + assert asyncio.run_coroutine_threadsafe(check(client), loop).result(THREAD_TIMEOUT_SECONDS) + other = asyncio.run_coroutine_threadsafe(client._pdp_sessions.current(), loop).result( + THREAD_TIMEOUT_SECONDS + ) + closed_later.close() + monkeypatch.setattr(aiohttp.ClientSession, "close", close_or_fail) + + with pytest.raises(OSError, match="the session did not close"): + asyncio.run(client.close()) + + assert other.closed + assert not failing.closed + monkeypatch.undo() + asyncio.run(client.close()) + assert failing.closed + finally: + loop.call_soon_threadsafe(loop.stop) + thread.join(THREAD_TIMEOUT_SECONDS) + loop.close() + closed_later.close() + + +def test_a_loop_closed_without_shutting_down_leaves_its_connection_to_the_garbage_collector( + server: KeepAliveServer, client: Permit +) -> None: + """The documented limit of a loop closed with ``loop.close()`` alone. + + Nothing can close a connection on a closed loop, so it stays open until the next request + marks the session closed and the garbage collector frees the connection, which Python + reports with a ResourceWarning. Running close() on the loop before closing it closes the + connection instead (test_close_leaves_the_connection_of_an_idle_loop_to_that_loop). + """ + assert run_on_a_loop_closed_without_shutting_down(check(client)) + assert server.wait_until_closed(1, timeout=0.2) == 0 + + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + assert asyncio.run(check(client)) + gc.collect() + + assert caught + assert {warning.category for warning in caught} == {ResourceWarning} + assert all("unclosed" in str(warning.message) for warning in caught) + # The connection of the closed loop, and the one asyncio.run() closed as it ended. + assert server.wait_until_closed(2) == 2 + + +def test_close_closes_the_session_of_a_loop_closed_without_shutting_down( + httpserver: HTTPServer, config: PermitConfig +) -> None: + httpserver.expect_request("/allowed", method="POST").respond_with_json({"allow": True}) + client = Permit(config) + assert run_on_a_loop_closed_without_shutting_down(check(client)) + + asyncio.run(client.close()) + + def drop() -> None: + nonlocal client + del client + + assert_nothing_reported_unclosed(drop) + + +def test_a_close_handed_to_a_loop_that_closes_first_leaves_nothing_unawaited( + httpserver: HTTPServer, config: PermitConfig +) -> None: + """A client collected on its running loop hands it the close, and the loop may stop first.""" + httpserver.expect_request("/allowed", method="POST").respond_with_json({"allow": True}) + client = Permit(config) + loop = asyncio.new_event_loop() + assert loop.run_until_complete(check(client)) + + def drop_and_stop() -> None: + nonlocal client + del client + loop.stop() + + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + loop.call_soon(drop_and_stop) + loop.run_forever() + loop.close() + gc.collect() + # The next request marks the session of the closed loop closed. + assert asyncio.run(check(Permit(config))) + gc.collect() + + assert [f"{w.category.__name__}: {w.message}" for w in caught] == [] + + +def test_a_client_dropped_after_its_loop_closed_without_shutting_down_reports_nothing( + httpserver: HTTPServer, config: PermitConfig +) -> None: + """Its session is kept until the next request, from any client, marks it closed.""" + httpserver.expect_request("/allowed", method="POST").respond_with_json({"allow": True}) + client = Permit(config) + assert run_on_a_loop_closed_without_shutting_down(check(client)) + [entry] = client._pdp_sessions._sessions.values() + session = weakref.ref(entry.session) + del entry + + def drop() -> None: + nonlocal client + del client + + assert_nothing_reported_unclosed(drop) + assert session() is not None + + assert asyncio.run(check(Permit(config))) + gc.collect() + + assert session() is None + + +EXIT_SCRIPT = """ +import asyncio +import sys +import threading + +from permit import Permit + +client = Permit(token="test-token", pdp=sys.argv[1], api_url=sys.argv[1]) +loop = asyncio.new_event_loop() +check = client.check("user-1", "read", "document") +{use} +print("checked") +""" + +EXIT_SCENARIOS = { + "a loop still running in a daemon thread": """ +threading.Thread(target=loop.run_forever, daemon=True).start() +assert asyncio.run_coroutine_threadsafe(check, loop).result(5) +""", + "a loop that is not running": """ +assert loop.run_until_complete(check) +""", + "a loop closed without shutting down": """ +assert loop.run_until_complete(check) +loop.close() +""", + "a client dropped after its loop closed without shutting down": """ +assert loop.run_until_complete(check) +loop.close() +del check, client +""", +} + + +@pytest.mark.parametrize("use", EXIT_SCENARIOS.values(), ids=EXIT_SCENARIOS.keys()) +def test_a_client_never_closed_reports_nothing_unclosed_at_exit( + server: KeepAliveServer, use: str +) -> None: + """The interpreter's exit closes the sessions whose loop it finds still open. + + The script runs under Python's default warning filters, as an application does: they + hide ResourceWarnings, but not what aiohttp logs about a session it finds unclosed. + """ + script = EXIT_SCRIPT.format(use=textwrap.dedent(use)) + env = { + name: value + for name, value in os.environ.items() + if name not in ("PYTHONWARNINGS", "PYTHONDEVMODE") + } + + result = subprocess.run( + [sys.executable, "-c", script, server.url], + cwd=REPO_ROOT, + env=env, + capture_output=True, + text=True, + timeout=120, + check=False, + ) + + assert (result.returncode, result.stdout) == (0, "checked\n"), result.stderr + reported = [ + line + for line in result.stderr.splitlines() + if "nclosed" in line or "Exception ignored" in line + ] + assert reported == [] + + +FORK_SCRIPT = """ +import asyncio +import os +import sys +import threading +import warnings + +from permit import Permit + +client = Permit(token="test-token", pdp=sys.argv[1], api_url=sys.argv[1]) +loop = asyncio.new_event_loop() +threading.Thread(target=loop.run_forever, daemon=True).start() +check = client.check("user-1", "read", "document") +assert asyncio.run_coroutine_threadsafe(check, loop).result(5) +# Python 3.12+ warns that forking a process that runs threads can deadlock the child. +warnings.simplefilter("ignore", DeprecationWarning) +pid = os.fork() +if pid == 0: + print("child:", asyncio.run(client.check("user-1", "read", "document")), flush=True) + asyncio.run(client.close()) + print("child closed", flush=True) + sys.exit(0) +_, status = os.waitpid(pid, 0) +print("child exit status:", status) +""" + + +@pytest.mark.skipif(sys.platform == "win32", reason="os.fork") +def test_a_forked_child_leaves_the_parent_sessions_alone(server: KeepAliveServer) -> None: + """The parent's loop does not run in the child, so close() must not wait for it there.""" + env = { + name: value + for name, value in os.environ.items() + if name not in ("PYTHONWARNINGS", "PYTHONDEVMODE") + } + + result = subprocess.run( + [ + sys.executable, + "-W", + "ignore:Support for pydantic 1 is deprecated:DeprecationWarning", + "-c", + FORK_SCRIPT, + server.url, + ], + cwd=REPO_ROOT, + env=env, + capture_output=True, + text=True, + timeout=60, + check=False, + ) + + assert (result.returncode, result.stderr) == (0, "") + assert result.stdout == "child: True\nchild closed\nchild exit status: 0\n" + # The parent's connection and the child's, which the child's asyncio.run() closed. + assert server.opened == 2 diff --git a/tests/test_benchmark_connection_reuse.py b/tests/test_benchmark_connection_reuse.py new file mode 100644 index 00000000..51c8bd7b --- /dev/null +++ b/tests/test_benchmark_connection_reuse.py @@ -0,0 +1,44 @@ +"""The connection reuse benchmark (PER-16344) runs, and counts one connection per client. + +The benchmark runs in an interpreter of its own, since it turns the SDK's logging off. +""" + +import os +import subprocess +import sys +from pathlib import Path + +REPO_ROOT = Path(__file__).resolve().parents[1] + + +def test_the_benchmark_reports_one_connection_per_client() -> None: + env = { + name: value + for name, value in os.environ.items() + if name not in ("PYTHONWARNINGS", "PYTHONDEVMODE") + } + + result = subprocess.run( + [ + sys.executable, + "-W", + "error", + "-W", + "ignore:Support for pydantic 1 is deprecated:DeprecationWarning", + "-m", + "tests.benchmark_connection_reuse", + "--calls", + "5", + ], + cwd=REPO_ROOT, + env=env, + capture_output=True, + text=True, + timeout=60, + check=False, + ) + + assert (result.returncode, result.stderr) == (0, "") + header, *rows = result.stdout.splitlines()[1:] + assert header.split()[:3] == ["client", "calls", "connections"] + assert [row.split()[:3] for row in rows] == [["async", "5", "1"], ["sync", "5", "1"]] diff --git a/tests/test_sync_lifecycle.py b/tests/test_sync_lifecycle.py new file mode 100644 index 00000000..25454bdb --- /dev/null +++ b/tests/test_sync_lifecycle.py @@ -0,0 +1,1081 @@ +"""Offline tests of the sync client's lifecycle: its background thread, close() and `with`. + +permit.sync.Permit runs every blocking call on an event loop in a daemon thread of its own, +so that its calls share HTTP connections (PER-16344). These tests read what can be observed +from outside: the connections a local keep-alive server counts, the state of the client's +thread, the warnings issued, and how a separate interpreter exits. +""" + +import asyncio +import contextlib +import contextvars +import gc +import os +import subprocess +import sys +import threading +import time +import traceback +import warnings +import weakref +from collections.abc import Callable, Iterator +from concurrent.futures import Future, ThreadPoolExecutor +from contextvars import ContextVar +from datetime import datetime, timezone +from pathlib import Path +from typing import Any +from uuid import uuid4 + +import pytest +from pytest_httpserver import HTTPServer + +from permit.config import PermitConfig +from permit.sync import Permit as SyncPermit +from permit.utils.deprecation import deprecated +from permit.utils.http_sessions import LoopSessions +from permit.utils.sync import SyncClass, _background_loop_of, _BackgroundLoop, _LoopThread +from tests.keepalive_server import KeepAliveServer +from tests.utils import FACTS, offline_config + +REPO_ROOT = Path(__file__).resolve().parents[1] +LOOP_THREAD_NAME = "permit-sync-loop" + + +@pytest.fixture +def server() -> Iterator[KeepAliveServer]: + with KeepAliveServer() as server: + yield server + + +@pytest.fixture +def permit(server: KeepAliveServer) -> Iterator[SyncPermit]: + client = SyncPermit(offline_config(server.url)) + yield client + close_within(client) + + +def close_within(client: SyncPermit, timeout: float = 5.0) -> None: + """Close `client`, failing instead of waiting forever if its loop thread is stuck. + + A regression that deadlocks the thread, such as a blocking call the client lets wait for + its own thread, would otherwise hang the test session here. The stuck thread is a + daemon, and close() has already left the exit hook nothing to wait for. + """ + closing = threading.Thread(target=client.close, daemon=True) + closing.start() + closing.join(timeout) + assert not closing.is_alive(), "close() did not return: the client's loop thread is stuck" + + +def loop_thread(client: SyncPermit) -> threading.Thread | None: + """The client's background thread, or None while it has none.""" + running = client._background_loop._thread + return None if running is None else running.thread + + +def started_loop_threads(before: list[threading.Thread]) -> list[threading.Thread]: + """The background loop threads running now that were not in `before`.""" + return [ + thread + for thread in threading.enumerate() + if thread.name == LOOP_THREAD_NAME and thread not in before + ] + + +def wait_until_stopped(thread: threading.Thread, timeout: float = 5.0) -> None: + """Wait for `thread` to end, a slice at a time. + + On a free-threaded build, an object that another thread releases is freed by the thread + that created it, once that thread runs again. So a collected client's loop is only + stopped when this thread wakes up, which one long join() would not let it do. + """ + deadline = time.monotonic() + timeout + while thread.is_alive() and time.monotonic() < deadline: + thread.join(timeout=0.05) + + +def check(client: SyncPermit) -> bool: + return client.check("user", "read", "document") + + +def user_payload(key: str) -> dict[str, Any]: + now = datetime.now(timezone.utc).isoformat() + ids = {name: str(uuid4()) for name in ("id", "organization_id", "project_id", "environment_id")} + return {"key": key, **ids, "created_at": now, "updated_at": now} + + +# --- the background thread ------------------------------------------------------------ + + +def test_the_client_starts_no_thread_before_its_first_call(permit: SyncPermit) -> None: + assert loop_thread(permit) is None + + +def test_the_calls_of_a_client_run_on_one_daemon_thread(permit: SyncPermit) -> None: + before = threading.enumerate() + + assert check(permit) is True + thread = loop_thread(permit) + assert check(permit) is True + + assert thread is not None + assert thread.is_alive() + assert thread.daemon + assert loop_thread(permit) is thread + assert started_loop_threads(before) == [thread] + + +def test_many_threads_share_one_client_and_its_thread( + permit: SyncPermit, server: KeepAliveServer +) -> None: + threads, calls = 16, 10 + all_started = threading.Barrier(threads) + before = threading.enumerate() + + def caller(index: int) -> list[bool]: + all_started.wait(timeout=10) + return [permit.check(f"user-{index}-{call}", "read", "document") for call in range(calls)] + + with ThreadPoolExecutor(max_workers=threads) as executor: + results = list(executor.map(caller, range(threads))) + + assert results == [[True] * calls] * threads + assert len(server.requests) == threads * calls + assert started_loop_threads(before) == [loop_thread(permit)] + + +def test_a_call_from_a_thread_that_runs_an_event_loop(permit: SyncPermit) -> None: + """The caller's loop is blocked for the call, which runs on the client's thread.""" + + async def main() -> tuple[bool, threading.Thread]: + return check(permit), threading.current_thread() + + allowed, caller = asyncio.run(main()) + + assert allowed is True + assert loop_thread(permit) not in (None, caller) + + +def sync_api_objects(root: object) -> list[object]: + """Every object of a `SyncClass` class that `root` holds, through instance attributes.""" + found: list[object] = [] + pending, seen = [root], set() + while pending: + obj = pending.pop() + if id(obj) in seen or not hasattr(obj, "__dict__"): + continue + seen.add(id(obj)) + if isinstance(type(obj), SyncClass): + found.append(obj) + pending.extend( + value for value in vars(obj).values() if value.__class__.__module__.startswith("permit") + ) + return found + + +@pytest.mark.parametrize("copy", [False, True], ids=["client", "wait_for_sync copy"]) +def test_every_blocking_api_object_of_the_client_runs_on_its_loop( + config: PermitConfig, *, copy: bool +) -> None: + config.proxy_facts_via_pdp = True + client = SyncPermit(config) + with client.wait_for_sync() as waiting: + objects = sync_api_objects(waiting if copy else client) + + # The enforcer, permit.api and its 19 sub-APIs, permit.elements and the PDP's role + # assignments. + assert len(objects) >= 23 + assert [obj for obj in objects if _background_loop_of(obj) is not client._background_loop] == [] + + +@pytest.mark.parametrize( + ("path", "response", "call"), + [ + ("/allowed", {"allow": True}, lambda client: client.check("u", "read", "document")), + (f"{FACTS}/users/u", None, lambda client: client.api.users.get("u")), + ("/local/role_assignments", [], lambda client: client.pdp_api.role_assignments.list()), + ], + ids=["check", "api.users.get", "pdp_api.role_assignments.list"], +) +def test_a_call_through_any_api_starts_the_client_thread( + httpserver: HTTPServer, + config: PermitConfig, + path: str, + response: object, + call: Callable[[SyncPermit], object], +) -> None: + httpserver.expect_oneshot_request(path).respond_with_json( + user_payload("u") if response is None else response + ) + with SyncPermit(config) as client: + call(client) + thread = loop_thread(client) + + assert thread is not None + assert not thread.is_alive() + httpserver.check_assertions() + + +# --- close() and `with` --------------------------------------------------------------- + + +def test_close_stops_and_joins_the_thread(permit: SyncPermit) -> None: + check(permit) + thread = loop_thread(permit) + + permit.close() + + assert thread is not None + assert not thread.is_alive() + assert loop_thread(permit) is None + + +def test_close_can_be_called_twice_and_before_any_call(server: KeepAliveServer) -> None: + unused = SyncPermit(offline_config(server.url)) + unused.close() + unused.close() + used = SyncPermit(offline_config(server.url)) + check(used) + used.close() + used.close() + + assert loop_thread(unused) is None + assert loop_thread(used) is None + + +def test_a_call_after_close_starts_a_new_thread(permit: SyncPermit) -> None: + check(permit) + first = loop_thread(permit) + permit.close() + + assert check(permit) is True + second = loop_thread(permit) + + assert second is not None + assert second is not first + assert second.is_alive() + + +def test_a_with_block_gives_the_client_and_closes_it(server: KeepAliveServer) -> None: + client = SyncPermit(offline_config(server.url)) + + with client as entered: + check(entered) + thread = loop_thread(entered) + + assert entered is client + assert thread is not None + assert not thread.is_alive() + + +def test_a_with_block_that_raises_still_closes_the_client( + server: KeepAliveServer, +) -> None: + client = SyncPermit(offline_config(server.url)) + check(client) + thread = loop_thread(client) + + def fail_in_a_with_block() -> None: + with client: + raise LookupError + + with pytest.raises(LookupError): + fail_in_a_with_block() + + assert thread is not None + assert not thread.is_alive() + + +def test_async_with_is_refused(server: KeepAliveServer) -> None: + client = SyncPermit(offline_config(server.url)) + + async def enter() -> None: + # The mistake a type checker reports too: this checks what it does at runtime. + async with client: # type: ignore[misc] + check(client) + + with pytest.raises(TypeError, match=r"use `with Permit\(\.\.\.\) as permit:`"): + asyncio.run(enter()) + assert loop_thread(client) is None + assert server.opened == 0 + + +def test_close_waits_for_a_call_in_flight(permit: SyncPermit, server: KeepAliveServer) -> None: + server.respond("/allowed", {"allow": True}, delay=0.5) + with ThreadPoolExecutor(max_workers=1) as executor: + in_flight = executor.submit(check, permit) + assert server.wait_for_requests(1) + permit.close() + + assert in_flight.result(timeout=5) is True + + +def wait_until_closing(client: SyncPermit, timeout: float = 5.0) -> None: + """Wait until a close() of `client` has started stopping its thread.""" + deadline = time.monotonic() + timeout + while client._background_loop._closing is None: + assert time.monotonic() < deadline, "close() did not start" + time.sleep(0.01) + + +def in_a_daemon_thread(function: Callable[[], object]) -> Future[object]: + """Call `function` in a daemon thread, which cannot hold up the exit if it gets stuck.""" + outcome: Future[object] = Future() + + def call() -> None: + try: + outcome.set_result(function()) + except Exception as error: + outcome.set_exception(error) + + threading.Thread(target=call, daemon=True).start() + return outcome + + +def test_a_call_made_during_close_waits_for_it_then_starts_a_new_thread( + permit: SyncPermit, server: KeepAliveServer +) -> None: + server.respond("/allowed", {"allow": True}, delay=0.5) + in_flight = in_a_daemon_thread(lambda: check(permit)) + assert server.wait_for_requests(1) + first = loop_thread(permit) + closing = in_a_daemon_thread(permit.close) + wait_until_closing(permit) + + during_close = in_a_daemon_thread(lambda: check(permit)) + + assert in_flight.result(timeout=5) is True + assert closing.result(timeout=5) is None + assert during_close.result(timeout=5) is True + second = loop_thread(permit) + assert first is not None + assert not first.is_alive() + assert second not in (None, first) + assert (server.opened, server.wait_until_closed(1)) == (2, 1) + + +def test_a_second_close_returns_once_the_first_has_stopped_the_thread( + permit: SyncPermit, server: KeepAliveServer +) -> None: + server.respond("/allowed", {"allow": True}, delay=0.5) + in_flight = in_a_daemon_thread(lambda: check(permit)) + assert server.wait_for_requests(1) + thread = loop_thread(permit) + first_close = in_a_daemon_thread(permit.close) + wait_until_closing(permit) + + second_close = in_a_daemon_thread(permit.close) + + assert second_close.result(timeout=5) is None + assert thread is not None + assert not thread.is_alive() + assert first_close.result(timeout=5) is None + assert in_flight.result(timeout=5) is True + assert loop_thread(permit) is None + + +def test_a_blocking_call_on_the_thread_a_close_stops_raises_instead_of_deadlocking( + permit: SyncPermit, server: KeepAliveServer +) -> None: + server.respond("/allowed", {"allow": True}, delay=0.5) + in_flight = in_a_daemon_thread(lambda: check(permit)) + assert server.wait_for_requests(1) + stopping = permit._background_loop._thread + assert stopping is not None + closing = in_a_daemon_thread(permit.close) + wait_until_closing(permit) + outcome: Future[object] = Future() + + def call() -> None: + try: + outcome.set_result(check(permit)) + except Exception as error: + outcome.set_exception(error) + + stopping.loop.call_soon_threadsafe(call, context=contextvars.Context()) + error = outcome.exception(timeout=5) + + assert isinstance(error, RuntimeError) + assert "own event loop thread" in str(error) + assert closing.result(timeout=5) is None + assert in_flight.result(timeout=5) is True + + +def test_close_closes_the_connections(permit: SyncPermit, server: KeepAliveServer) -> None: + check(permit) + + permit.close() + + assert server.wait_until_closed(1) == 1 + + +def test_closing_a_wait_for_sync_copy_leaves_its_client_thread_and_connection_open( + server: KeepAliveServer, +) -> None: + """A copy runs on its client's thread and connections, and leaves closing them to it.""" + config = offline_config(server.url) + config.proxy_facts_via_pdp = True + client = SyncPermit(config) + with client.wait_for_sync() as waiting: + assert check(waiting) is True + thread = loop_thread(client) + assert loop_thread(waiting) is thread + waiting.close() + waiting.close() + + assert thread is not None + assert thread.is_alive() + assert check(client) is True + assert loop_thread(client) is thread + assert (server.opened, server.closed) == (1, 0) + + client.close() + + assert not thread.is_alive() + assert server.wait_until_closed(1) == 1 + + +# --- errors and re-entrancy ------------------------------------------------------------ + + +def test_an_error_keeps_its_type_and_traceback(permit: SyncPermit) -> None: + with pytest.raises(ValueError, match="invalid resource string") as caught: + permit.check("user", "read", "too:many:parts") + + frames = [ + (Path(frame.filename).name, frame.name) + for frame in traceback.extract_tb(caught.value.__traceback__) + ] + assert ("enforcer.py", "_resource_from_string") in frames + assert (Path(__file__).name, test_an_error_keeps_its_type_and_traceback.__name__) in frames + assert loop_thread(permit) is not None + + +def test_a_timeout_keeps_its_traceback_and_cause(server: KeepAliveServer) -> None: + """On Python 3.11 and 3.12, asyncio would hand the caller a bare copy of the TimeoutError.""" + server.respond("/allowed", {"allow": True}, delay=1.5) + config = offline_config(server.url) + config.pdp_timeout = 1 + + with SyncPermit(config) as client, pytest.raises(asyncio.TimeoutError) as caught: + check(client) + + frames = [ + (Path(frame.filename).name, frame.name) + for frame in traceback.extract_tb(caught.value.__traceback__) + ] + assert ("enforcer.py", "check") in frames + assert (Path(__file__).name, test_a_timeout_keeps_its_traceback_and_cause.__name__) in frames + assert caught.value.__cause__ is not None + + +def test_a_client_whose_call_raised_is_freed_by_reference_counting(server: KeepAliveServer) -> None: + """The exception a call raised leaves no reference cycle that would hold the client.""" + client = SyncPermit(offline_config(server.url)) + with pytest.raises(ValueError, match="invalid resource string"): + client.check("user", "read", "too:many:parts") + freed = weakref.ref(client) + running = client._background_loop._thread + assert running is not None + + gc.disable() + try: + del client + # On a free-threaded build, an object that another thread releases is freed by the + # thread that created it, once that thread runs again: wake the client's thread and + # let both threads run. Freeing the client at `del` may already have stopped or closed + # its loop, so the wake-up is best effort and only the client being freed is checked. + with contextlib.suppress(RuntimeError): + running.loop.call_soon_threadsafe(lambda: None) + deadline = time.monotonic() + 5 + while freed() is not None and time.monotonic() < deadline: + time.sleep(0.01) + assert freed() is None + finally: + gc.enable() + + +def run_on_client_thread(client: SyncPermit, function: Callable[[], object]) -> Future[object]: + """Call `function` on the client's thread, in a fresh context, as a loop callback would.""" + running = client._background_loop._thread + assert running is not None + outcome: Future[object] = Future() + + def call() -> None: + try: + outcome.set_result(function()) + except Exception as error: + outcome.set_exception(error) + + running.loop.call_soon_threadsafe(call, context=contextvars.Context()) + return outcome + + +def test_a_blocking_call_on_the_client_thread_raises_instead_of_deadlocking( + permit: SyncPermit, +) -> None: + check(permit) + + outcome = run_on_client_thread(permit, lambda: check(permit)) + error = outcome.exception(timeout=5) + + assert isinstance(error, RuntimeError) + assert "own event loop thread" in str(error) + assert check(permit) is True + # The refused call's coroutine was closed, so collecting it does not report it unawaited. + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + del outcome, error + gc.collect() + assert [f"{w.category.__name__}: {w.message}" for w in caught] == [] + + +def test_close_on_the_client_thread_raises_instead_of_deadlocking(permit: SyncPermit) -> None: + check(permit) + + outcome = run_on_client_thread(permit, permit.close) + + with pytest.raises(RuntimeError, match="own event loop thread"): + outcome.result(timeout=5) + assert check(permit) is True + + +class RecordingPermit(SyncPermit): + """A sync client that records the thread on which its HTTP sessions are closed. + + A wait_for_sync() copy shares the list, and records nothing: it closes no sessions. + """ + + closed_on: list[str] + + async def _close_sessions(self) -> None: + await super()._close_sessions() + if self._owns_sessions: + self.closed_on.append(threading.current_thread().name) + + +def recording_client(url: str) -> RecordingPermit: + client = RecordingPermit(offline_config(url)) + client.closed_on = [] + return client + + +def test_close_closes_the_sessions_once_on_the_client_thread_while_a_copy_is_alive( + config: PermitConfig, httpserver: HTTPServer +) -> None: + """A wait_for_sync() copy shares its client's sessions, so they are closed once.""" + httpserver.expect_request("/allowed").respond_with_json({"allow": True}) + config.proxy_facts_via_pdp = True + client = RecordingPermit(config) + client.closed_on = [] + with client.wait_for_sync() as waiting: + check(waiting) + client.close() + + assert client.closed_on == [LOOP_THREAD_NAME] + + +class FailingPermit(SyncPermit): + """A sync client whose HTTP sessions fail to close.""" + + async def _close_sessions(self) -> None: + await super()._close_sessions() + msg = "the sessions did not close" + raise OSError(msg) + + +def test_close_raises_what_closing_the_sessions_raised_and_still_stops( + server: KeepAliveServer, +) -> None: + client = FailingPermit(offline_config(server.url)) + check(client) + thread = loop_thread(client) + + with pytest.raises(OSError, match="the sessions did not close"): + client.close() + + assert thread is not None + assert not thread.is_alive() + assert loop_thread(client) is None + + +def test_close_without_a_call_closes_no_sessions(server: KeepAliveServer) -> None: + client = recording_client(server.url) + + client.close() + + assert client.closed_on == [] + + +# --- a client that is never closed ----------------------------------------------------- + + +def test_a_client_that_is_garbage_collected_stops_its_thread( + server: KeepAliveServer, +) -> None: + client = SyncPermit(offline_config(server.url)) + check(client) + thread = loop_thread(client) + + del client + gc.collect() + + assert thread is not None + wait_until_stopped(thread) + assert not thread.is_alive() + assert server.wait_until_closed(1) == 1 + + +def test_an_api_object_outliving_its_client_keeps_working() -> None: + with KeepAliveServer() as server: + server.respond("/local/role_assignments", []) + role_assignments = SyncPermit(offline_config(server.url)).pdp_api.role_assignments + gc.collect() + + assert role_assignments.list(user_key="u") == [] + background_loop = _background_loop_of(role_assignments) + assert background_loop is not None + running = background_loop._thread + assert running is not None + + del role_assignments, background_loop + gc.collect() + wait_until_stopped(running.thread) + assert not running.thread.is_alive() + + +def test_a_client_freed_by_the_cycle_collector_closes_its_connection_without_a_warning( + server: KeepAliveServer, +) -> None: + """A client held in a reference cycle, as an exception it raised can hold it, is freed by gc. + + The client's sessions must not be collected with it while they are open, or aiohttp + reports them unclosed. + """ + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + client = SyncPermit(offline_config(server.url)) + check(client) + thread = loop_thread(client) + cycle: list[object] = [client] + cycle.append(cycle) + del client, cycle + gc.collect() + assert thread is not None + wait_until_stopped(thread) + gc.collect() + + assert [f"{w.category.__name__}: {w.message}" for w in caught] == [] + assert not thread.is_alive() + assert server.wait_until_closed(1) == 1 + + +def test_a_client_never_closed_issues_no_warning(server: KeepAliveServer) -> None: + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + client = SyncPermit(offline_config(server.url)) + check(client) + thread = loop_thread(client) + del client + gc.collect() + assert thread is not None + wait_until_stopped(thread) + gc.collect() + + assert [f"{w.category.__name__}: {w.message}" for w in caught] == [] + + +def run_script(script: str, timeout: float = 60) -> subprocess.CompletedProcess[str]: + """Run `script` in a new interpreter that turns every warning into an error.""" + env = { + name: value + for name, value in os.environ.items() + if name not in ("PYTHONWARNINGS", "PYTHONDEVMODE") + } + env["PYTHONPATH"] = str(REPO_ROOT) + return subprocess.run( + [ + sys.executable, + "-W", + "error", + "-W", + "ignore:Support for pydantic 1 is deprecated:DeprecationWarning", + "-c", + script, + ], + env=env, + cwd=REPO_ROOT, + capture_output=True, + text=True, + timeout=timeout, + check=False, + ) + + +SCRIPT_HEADER = """\ +import atexit +import threading + +from loguru import logger + +from tests.keepalive_server import KeepAliveServer + +logger.disable("permit") +server = KeepAliveServer() +server.start() + + +def report() -> None: + opened = server.opened + print("connections closed:", server.wait_until_closed(opened) == opened) + loop_threads = [t for t in threading.enumerate() if t.name == "permit-sync-loop"] + print("loop threads left:", len(loop_threads)) + + +# Registered before permit is imported, so it runs after permit's own exit hook. +atexit.register(report) + +from permit.sync import Permit +from tests.utils import offline_config + + +class Client(Permit): + async def _close_sessions(self) -> None: + await super()._close_sessions() + print("sessions closed on", threading.current_thread().name) + + +client = Client(offline_config(server.url)) +""" +AT_EXIT = "sessions closed on permit-sync-loop\nconnections closed: True\nloop threads left: 0\n" + + +def test_a_client_never_closed_is_closed_at_exit_without_noise() -> None: + result = run_script(SCRIPT_HEADER + "print(client.check('user', 'read', 'document'))\n") + + assert (result.returncode, result.stderr) == (0, "") + assert result.stdout == "True\n" + AT_EXIT + + +def test_a_call_in_flight_does_not_hold_up_the_exit() -> None: + script = SCRIPT_HEADER + ( + "server.respond('/allowed', {'allow': True}, delay=60)\n" + "def call():\n" + " try:\n" + " client.check('user', 'read', 'document')\n" + " except BaseException:\n" + " pass\n" + "threading.Thread(target=call, daemon=True).start()\n" + "print('in flight:', server.wait_for_requests(1))\n" + ) + started = time.monotonic() + + result = run_script(script, timeout=30) + + assert (result.returncode, result.stderr) == (0, "") + assert result.stdout == "in flight: True\n" + AT_EXIT + assert time.monotonic() - started < 30 + + +@pytest.mark.skipif(sys.platform == "win32", reason="POSIX signals") +def test_ctrl_c_while_waiting_cancels_the_call() -> None: + script = """\ +import asyncio +import os +import signal +import threading + +from permit.utils.sync import SyncClass, _BackgroundLoop + +started = threading.Event() +cancelled = threading.Event() + + +class Api(metaclass=SyncClass): + async def wait(self) -> None: + started.set() + try: + await asyncio.sleep(60) + except asyncio.CancelledError: + cancelled.set() + raise + + +def interrupt() -> None: + started.wait(10) + os.kill(os.getpid(), signal.SIGINT) + + +# A process started with SIGINT ignored, as a shell's background job is, inherits that. +signal.signal(signal.SIGINT, signal.default_int_handler) +background_loop = _BackgroundLoop() +api = Api() +background_loop.bind(api) +threading.Thread(target=interrupt).start() +try: + api.wait() +except KeyboardInterrupt: + print("interrupted") +print("cancelled:", cancelled.wait(10)) +background_loop.close() +""" + result = run_script(script, timeout=30) + + assert (result.returncode, result.stderr) == (0, "") + assert result.stdout == "interrupted\ncancelled: True\n" + + +@pytest.mark.skipif(sys.platform == "win32", reason="os.fork") +def test_a_forked_child_starts_a_thread_of_its_own_and_closes_it() -> None: + """The child leaves the parent's loop and connections alone, so close() cannot hang on them.""" + script = SCRIPT_HEADER + ( + "import os\n" + "import sys\n" + "import warnings\n" + "print('parent:', client.check('user', 'read', 'document'), flush=True)\n" + "# Python 3.12+ warns that forking a process that runs threads can deadlock the child.\n" + "warnings.simplefilter('ignore', DeprecationWarning)\n" + "pid = os.fork()\n" + "if pid == 0:\n" + " # The server's thread runs in the parent only.\n" + " atexit.unregister(report)\n" + " print('child:', client.check('user', 'read', 'document'), flush=True)\n" + " client.close()\n" + " print('child:', client.check('user', 'read', 'document'), flush=True)\n" + " sys.exit(0)\n" + "_, status = os.waitpid(pid, 0)\n" + "print('child exit status:', status)\n" + ) + + result = run_script(script, timeout=30) + + assert (result.returncode, result.stderr) == (0, "") + closed_in_the_child = "sessions closed on permit-sync-loop\n" + assert result.stdout == ( + "parent: True\n" + f"child: True\n{closed_in_the_child}" + f"child: True\n{closed_in_the_child}" + "child exit status: 0\n" + AT_EXIT + ) + + +@pytest.mark.skipif(sys.platform == "win32", reason="os.fork") +def test_a_child_forked_while_close_runs_starts_a_thread_of_its_own() -> None: + """The close() the parent runs does not run in the child, so a call must not wait for it.""" + script = SCRIPT_HEADER + ( + "import os\n" + "import sys\n" + "import time\n" + "import warnings\n" + "server.respond('/allowed', {'allow': True}, delay=1)\n" + "threading.Thread(target=client.check, args=('user', 'read', 'document')).start()\n" + "server.wait_for_requests(1)\n" + "closing = threading.Thread(target=client.close)\n" + "closing.start()\n" + "while client._background_loop._closing is None:\n" + " time.sleep(0.01)\n" + "# Python 3.12+ warns that forking a process that runs threads can deadlock the child.\n" + "warnings.simplefilter('ignore', DeprecationWarning)\n" + "pid = os.fork()\n" + "if pid == 0:\n" + " # The server's thread runs in the parent only.\n" + " atexit.unregister(report)\n" + " print('child:', client.check('user', 'read', 'document'), flush=True)\n" + " client.close()\n" + " sys.exit(0)\n" + "_, status = os.waitpid(pid, 0)\n" + "closing.join()\n" + "print('child exit status:', status)\n" + ) + + result = run_script(script, timeout=30) + + assert (result.returncode, result.stderr) == (0, "") + assert "child: True\n" in result.stdout + assert "child exit status: 0\n" in result.stdout + + +# --- the background loop on its own ---------------------------------------------------- + +request_id: ContextVar[str] = ContextVar("request_id", default="") + + +class Probe(metaclass=SyncClass): + """A blocking API object for the tests of the background loop itself.""" + + async def request_id(self) -> str: + return request_id.get() + + async def thread(self) -> threading.Thread: + return threading.current_thread() + + @deprecated("old_fetch() is deprecated") + async def old_fetch(self) -> None: + await asyncio.sleep(0) + + +def blocking(method: Callable[[], object]) -> object: + """Call a method of `Probe`, whose methods mypy sees as returning coroutines.""" + return method() + + +@pytest.fixture +def probe() -> Iterator[Probe]: + background_loop = _BackgroundLoop() + bound = Probe() + background_loop.bind(bound) + yield bound + background_loop.close() + + +def test_close_cancels_the_tasks_a_call_left_running() -> None: + left_running: list[asyncio.Task[None]] = [] + cancelled = threading.Event() + + async def linger() -> None: + try: + await asyncio.sleep(60) + except asyncio.CancelledError: + cancelled.set() + raise + + class Spawner(metaclass=SyncClass): + async def spawn(self) -> None: + left_running.append(asyncio.get_running_loop().create_task(linger())) + await asyncio.sleep(0) + + background_loop = _BackgroundLoop() + spawner = Spawner() + background_loop.bind(spawner) + blocking(spawner.spawn) + + background_loop.close() + + assert cancelled.is_set() + assert left_running[0].cancelled() + + +def test_a_session_close_handed_to_the_loop_as_it_stops_still_closes_the_session( + server: KeepAliveServer, +) -> None: + """The loop cancels that close as it settles; shutting down its async generators closes it.""" + loop_thread = _LoopThread() + sessions = LoopSessions() + + async def open_a_connection(through: LoopSessions) -> None: + session = await through.current() + async with session.post(f"{server.url}/allowed") as response: + await response.read() + + opening = open_a_connection(sessions) + asyncio.run_coroutine_threadsafe(opening, loop_thread.loop).result(timeout=5) + del opening + stopped, release = threading.Event(), threading.Event() + + def stop_and_hold() -> None: + # The loop leaves run_forever() once this callback returns. + loop_thread.loop.stop() + stopped.set() + release.wait(timeout=5) + + loop_thread.loop.call_soon_threadsafe(stop_and_hold) + assert stopped.wait(timeout=5) + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + # Collecting the sessions hands the close of the open one to the stopping loop. + del sessions + release.set() + loop_thread.thread.join(timeout=5) + gc.collect() + + assert not loop_thread.thread.is_alive() + assert loop_thread.loop.is_closed() + assert server.wait_until_closed(1) == 1 + assert [f"{w.category.__name__}: {w.message}" for w in caught] == [] + + +def test_an_object_bound_to_no_client_runs_each_call_in_a_loop_of_its_own() -> None: + unbound = Probe() + + first, second = blocking(unbound.thread), blocking(unbound.thread) + + assert first is threading.current_thread() + assert second is threading.current_thread() + + +def test_the_caller_context_reaches_the_call(probe: Probe) -> None: + token = request_id.set("r-1") + try: + inside = blocking(probe.request_id) + finally: + request_id.reset(token) + + assert inside == "r-1" + assert blocking(probe.request_id) == "" + + +def test_concurrent_calls_each_warn_at_their_own_line(probe: Probe) -> None: + both_calling = threading.Barrier(2) + + def first_caller() -> None: + both_calling.wait(timeout=10) + _ = probe.old_fetch() # Blocking; mypy sees the async def it converts. + + def second_caller() -> None: + both_calling.wait(timeout=10) + _ = probe.old_fetch() # Blocking; mypy sees the async def it converts. + + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + with ThreadPoolExecutor(max_workers=2) as executor: + for future in [executor.submit(first_caller), executor.submit(second_caller)]: + future.result() + + def call_line(function: Callable[..., object]) -> tuple[str, int]: + return function.__code__.co_filename, function.__code__.co_firstlineno + 2 + + assert sorted((w.filename, w.lineno) for w in caught) == sorted( + [call_line(first_caller), call_line(second_caller)] + ) + + +def test_a_deprecated_method_of_the_client_warns_at_the_caller( + httpserver: HTTPServer, config: PermitConfig +) -> None: + httpserver.expect_oneshot_request(f"{FACTS}/users/u").respond_with_json(user_payload("u")) + with SyncPermit(config) as client, warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + line = sys._getframe().f_lineno + 1 + client.api.get_user("u") + + assert [(w.category, w.filename, w.lineno) for w in caught] == [ + (DeprecationWarning, __file__, line) + ] + + +# --- connection reuse ------------------------------------------------------------------ + + +def test_sequential_calls_reuse_one_connection(permit: SyncPermit, server: KeepAliveServer) -> None: + for _ in range(20): + check(permit) + + assert (server.opened, server.closed) == (1, 0) + + +def test_concurrent_threads_open_at_most_one_connection_each( + permit: SyncPermit, server: KeepAliveServer +) -> None: + threads = 8 + all_started = threading.Barrier(threads) + + def caller(_: int) -> None: + all_started.wait(timeout=10) + for _call in range(10): + check(permit) + + with ThreadPoolExecutor(max_workers=threads) as executor: + list(executor.map(caller, range(threads))) + + assert len(server.requests) == threads * 10 + assert server.opened <= threads diff --git a/tests/type_check/consumer.py b/tests/type_check/consumer.py index dbfcf800..d6a5aabd 100644 --- a/tests/type_check/consumer.py +++ b/tests/type_check/consumer.py @@ -244,6 +244,20 @@ def sync_client() -> None: assert_type(listed.key, str) +async def async_client_lifecycle() -> None: + async with Permit(CONFIG) as permit: + assert_type(permit, Permit) + assert_type(await permit.check("user", "read", "document"), bool) + await permit.close() + + +def sync_client_lifecycle() -> None: + with SyncPermit(CONFIG) as permit: + assert_type(permit, SyncPermit) + assert_type(permit.check("user", "read", "document"), bool) + permit.close() + + async def mistakes_stay_errors() -> None: permit = Permit(CONFIG) sync_permit = SyncPermit(CONFIG) @@ -264,5 +278,12 @@ async def mistakes_stay_errors() -> None: # The blocking client returns values, not awaitables. await sync_permit.api.users.get("u") # type: ignore[misc] await sync_permit.get_user_tenants("u") # type: ignore[misc] + await sync_permit.close() # type: ignore[func-returns-value, misc] # The async client returns awaitables, not values. _ = permit.api.users.get("u").email # type: ignore[attr-defined] + permit.close() # type: ignore[unused-coroutine] + # Each client has the context manager of its kind only. + with permit: # type: ignore[attr-defined] + pass + async with sync_permit: # type: ignore[misc] + pass diff --git a/uv.lock b/uv.lock index 5c853a53..4f058c51 100644 --- a/uv.lock +++ b/uv.lock @@ -1012,9 +1012,11 @@ source = { editable = "." } dependencies = [ { name = "aiohttp" }, { name = "loguru" }, + { name = "multidict" }, { name = "pydantic", version = "1.10.26", source = { registry = "https://pypi.org/simple" }, extra = ["email"], marker = "extra == 'group-6-permit-pydantic-v1'" }, { name = "pydantic", version = "2.13.5", source = { registry = "https://pypi.org/simple" }, extra = ["email"], marker = "extra == 'group-6-permit-pydantic-v2' or extra != 'group-6-permit-pydantic-v1'" }, { name = "typing-extensions" }, + { name = "yarl" }, ] [package.dev-dependencies] @@ -1042,10 +1044,12 @@ pydantic-v2 = [ requires-dist = [ { name = "aiohttp", specifier = ">=3.14.3,<4" }, { name = "loguru", specifier = ">=0.7.3,<1" }, + { name = "multidict", specifier = ">=6.7.0,<7" }, { name = "pydantic", extras = ["email"], marker = "python_full_version < '3.13'", specifier = ">=1.10.18,!=2.0.*,!=2.1.*,!=2.2.*,!=2.3.*,!=2.4.0,!=2.4.1" }, { name = "pydantic", extras = ["email"], marker = "python_full_version == '3.13.*'", specifier = ">=1.10.18,!=2.0.*,!=2.1.*,!=2.2.*,!=2.3.*,!=2.4.*,!=2.5.*,!=2.6.*,!=2.7.*" }, { name = "pydantic", extras = ["email"], marker = "python_full_version >= '3.14'", specifier = ">=1.10.25,!=2.0.*,!=2.1.*,!=2.2.*,!=2.3.*,!=2.4.*,!=2.5.*,!=2.6.*,!=2.7.*,!=2.8.*,!=2.9.*,!=2.10.*,!=2.11.*,!=2.12.*" }, { name = "typing-extensions", specifier = ">=4.14.0,<5" }, + { name = "yarl", specifier = ">=1.21.0,<2" }, ] [package.metadata.requires-dev]