From aa8581dab1e0a50ee06b3b08cf2a0bdfa1d48a5a Mon Sep 17 00:00:00 2001 From: Claude Date: Sat, 27 Jun 2026 10:08:33 +0000 Subject: [PATCH 1/4] Replace httpx with aiohttp HTTP backend Switch the transport layer from httpx to aiohttp so the package fits the Home Assistant ecosystem, which mandates aiohttp and a shared session per config entry. - HttpClient is now built on aiohttp.ClientSession and accepts an optional injected session; injected sessions are never closed by the client. - RohlikAPI gains a `session=` parameter and a `session` property (replacing the httpx-specific `client` property) to forward an external session. - A small buffered Response wrapper keeps the service layer synchronous (`.json()` / `.raise_for_status()`), decoupling services from aiohttp's streaming semantics. - Query params are coerced to aiohttp-acceptable strings (bool/None/nested). - Connection/timeout failures are caught via a shared HTTP_ERRORS tuple (aiohttp.ClientError plus TimeoutError). - Update dependencies, README, and tests; add coverage for session injection. Co-Authored-By: Claude Opus 4.8 Claude-Session: https://claude.ai/code/session_01PmotTwydT558t4JHd5Cnwm --- README.md | 30 +++++- pyproject.toml | 2 +- requirements.txt | 2 +- rohlik_api/__init__.py | 2 +- rohlik_api/auth.py | 8 +- rohlik_api/client.py | 23 +++-- rohlik_api/http_client.py | 162 ++++++++++++++++++++++++++------ rohlik_api/services/account.py | 5 +- rohlik_api/services/base.py | 6 +- rohlik_api/services/cart.py | 9 +- rohlik_api/services/orders.py | 5 +- rohlik_api/services/products.py | 11 +-- rohlik_api/services/recipes.py | 9 +- tests/test_auth.py | 2 +- tests/test_client.py | 19 ++-- tests/test_http_client.py | 76 ++++++++++----- tests/test_recipes.py | 12 +-- tests/test_services.py | 16 ++-- 18 files changed, 279 insertions(+), 120 deletions(-) diff --git a/README.md b/README.md index fc8a69c..ed27435 100644 --- a/README.md +++ b/README.md @@ -33,7 +33,7 @@ online grocery service β€” search products, manage your cart, browse recipes ## Features -- πŸš€ HTTP/2 support for fast, connection-reused requests +- πŸš€ Built on aiohttp; bring your own session (e.g. Home Assistant's shared session) - πŸ” Automatic login/logout and session management - 🎯 Clean, service-based API (`client.cart`, `client.products`, …) - 🧩 Fully typed dataclass models for parsed responses (`py.typed`) @@ -44,7 +44,7 @@ online grocery service β€” search products, manage your cart, browse recipes ## Requirements - Python 3.13+ -- [httpx](https://www.python-httpx.org/) with HTTP/2 (installed automatically) +- [aiohttp](https://docs.aiohttp.org/) (installed automatically) ## Installation @@ -280,6 +280,32 @@ async def main(): await client.close() ``` +### Reusing an existing aiohttp session + +The client is built on [aiohttp](https://docs.aiohttp.org/). By default it +creates and owns its own `ClientSession`, but you can inject an externally +managed session instead β€” useful inside a Home Assistant integration, where +the recommended pattern is to share a single session per instance. An injected +session is **never** closed by the client; its lifecycle stays with the owner. + +```python +import aiohttp +from rohlik_api import RohlikAPI + +async def main(session: aiohttp.ClientSession): + client = RohlikAPI( + username="email@example.com", + password="password", + session=session, # reuse the caller's session + ) + async with client: + cart = await client.cart.get_content() + # `session` is left open for the caller to close. +``` + +Inside a Home Assistant integration you would pass the shared session, for +example `RohlikAPI(..., session=async_get_clientsession(hass))`. + ## Development ```bash diff --git a/pyproject.toml b/pyproject.toml index a293774..0046d28 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -26,7 +26,7 @@ classifiers = [ "Typing :: Typed", ] dependencies = [ - "httpx[http2]>=0.28.1", + "aiohttp>=3.10", ] [project.optional-dependencies] diff --git a/requirements.txt b/requirements.txt index d6ca241..c916146 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,5 +1,5 @@ # Development requirements -httpx[http2]==0.28.1 +aiohttp>=3.10 pytest==9.0.2 pytest-cov==7.0.0 pytest-asyncio==1.3.0 diff --git a/rohlik_api/__init__.py b/rohlik_api/__init__.py index ffd6d8c..cc877c9 100644 --- a/rohlik_api/__init__.py +++ b/rohlik_api/__init__.py @@ -1,6 +1,6 @@ """Rohlik.cz API Python Client. -An async Python client for the Rohlik.cz API, built on httpx with HTTP/2 support. +An async Python client for the Rohlik.cz API, built on aiohttp. """ from .auth import AuthManager diff --git a/rohlik_api/auth.py b/rohlik_api/auth.py index 5b9bf09..54b0212 100644 --- a/rohlik_api/auth.py +++ b/rohlik_api/auth.py @@ -5,12 +5,10 @@ import logging from typing import Any -import httpx - from .endpoints import Endpoints from .errors import APIRequestFailedError, InvalidCredentialsError, RohlikAPIError from .helpers import mask_data -from .http_client import HttpClient +from .http_client import HTTP_ERRORS, HttpClient _LOGGER = logging.getLogger(__name__) @@ -112,7 +110,7 @@ async def login(self) -> dict[str, Any]: return login_response - except httpx.HTTPError as err: + except HTTP_ERRORS as err: raise APIRequestFailedError( f"Cannot connect to website! Check your internet connection " f"and try again: {err}" @@ -138,7 +136,7 @@ async def logout(self) -> None: self._reset_session() - except httpx.HTTPError as err: + except HTTP_ERRORS as err: self._reset_session() # Reset state even on error raise APIRequestFailedError( f"Cannot connect to website! Check your internet connection " diff --git a/rohlik_api/client.py b/rohlik_api/client.py index 37e51cd..d394926 100644 --- a/rohlik_api/client.py +++ b/rohlik_api/client.py @@ -6,12 +6,12 @@ from types import TracebackType from typing import Any -import httpx +import aiohttp from .auth import AuthManager from .endpoints import BASE_URL from .errors import APIRequestFailedError -from .http_client import HttpClient +from .http_client import HTTP_ERRORS, HttpClient from .services import ( AccountService, CartService, @@ -27,9 +27,9 @@ class RohlikAPI: """Async client for interacting with the Rohlik.cz API. - The client uses httpx with HTTP/2 support and exposes a service-based API - for all operations. When used as an async context manager with - ``auto_login=True`` (the default), it logs in on entry and logs out on exit. + The client is built on aiohttp and exposes a service-based API for all + operations. When used as an async context manager with ``auto_login=True`` + (the default), it logs in on entry and logs out on exit. Args: username: Email address used for Rohlik.cz login (required). @@ -39,6 +39,9 @@ class RohlikAPI: headers: Optional custom headers to include in all requests. auto_login: If True (default), log in automatically when used as a context manager. + session: Optional externally managed :class:`aiohttp.ClientSession` to + reuse (for example Home Assistant's shared session). When provided, + the session is not closed by this client. Attributes: cart (CartService): Cart operations (get_content, add_items, delete_item). @@ -62,6 +65,7 @@ def __init__( timeout: float = 30.0, headers: dict[str, str] | None = None, auto_login: bool = True, + session: aiohttp.ClientSession | None = None, ) -> None: # Credential validation is owned by AuthManager (constructed below), # which raises ValueError on empty username/password. @@ -74,6 +78,7 @@ def __init__( base_url=base_url, timeout=timeout, headers=headers, + session=session, ) # Initialize auth manager @@ -126,9 +131,9 @@ def recipes(self) -> RecipeService: return self._recipes @property - def client(self) -> httpx.AsyncClient: - """Get or create the underlying async HTTP client.""" - return self._http.client + def session(self) -> aiohttp.ClientSession: + """Get or create the underlying aiohttp session.""" + return self._http.session @property def is_logged_in(self) -> bool: @@ -232,7 +237,7 @@ async def get_data(self) -> dict[str, Any]: return result - except httpx.HTTPError as err: + except HTTP_ERRORS as err: raise APIRequestFailedError( f"Cannot connect to website! Check your internet connection " f"and try again: {err}" diff --git a/rohlik_api/http_client.py b/rohlik_api/http_client.py index 88daad4..ecb6585 100644 --- a/rohlik_api/http_client.py +++ b/rohlik_api/http_client.py @@ -1,12 +1,13 @@ -"""HTTP client for Rohlik.cz API.""" +"""HTTP client for the Rohlik.cz API.""" from __future__ import annotations +import json import logging from importlib.metadata import PackageNotFoundError, version from typing import Any -import httpx +import aiohttp from .endpoints import BASE_URL @@ -17,9 +18,52 @@ except PackageNotFoundError: # pragma: no cover - package not installed _VERSION = "0.0.0" +# Transport-level errors that callers treat as a failed request. aiohttp raises +# ``asyncio.TimeoutError`` (an alias of the builtin ``TimeoutError``) on +# timeouts, which is not a subclass of ``ClientError``, so it must be listed +# explicitly alongside it. +HTTP_ERRORS: tuple[type[Exception], ...] = (aiohttp.ClientError, TimeoutError) + + +class Response: + """Lightweight wrapper around an aiohttp response. + + The body is buffered when the response is created, so :meth:`json` and + :meth:`raise_for_status` are synchronous and can be called after the + underlying connection has been released back to the pool. This keeps the + service layer decoupled from aiohttp's streaming semantics. + """ + + __slots__ = ("status", "_body", "_response") + + def __init__(self, status: int, body: bytes, response: aiohttp.ClientResponse) -> None: + self.status = status + self._body = body + self._response = response + + def json(self) -> Any: + """Decode the response body as JSON, ignoring the content type.""" + return json.loads(self._body) + + def raise_for_status(self) -> None: + """Raise :class:`aiohttp.ClientResponseError` for a 4xx/5xx status.""" + self._response.raise_for_status() + class HttpClient: - """Async HTTP client with HTTP/2 support for Rohlik.cz API.""" + """Async HTTP client for the Rohlik.cz API, built on aiohttp. + + By default the client creates and owns its own :class:`aiohttp.ClientSession`. + A session may instead be injected (for example Home Assistant's shared + session obtained via ``homeassistant.helpers.aiohttp_client``); an injected + session is never closed by this client, leaving its lifecycle to the owner. + + Args: + base_url: Base URL for the Rohlik.cz API. + timeout: Request timeout in seconds. + headers: Optional custom headers added to every request. + session: Optional externally managed aiohttp session to reuse. + """ DEFAULT_USER_AGENT = f"rohlik-api-python/{_VERSION}" @@ -28,9 +72,11 @@ def __init__( base_url: str = BASE_URL, timeout: float = 30.0, headers: dict[str, str] | None = None, + session: aiohttp.ClientSession | None = None, ) -> None: self.base_url = base_url.rstrip("/") self.timeout = timeout + self._timeout = aiohttp.ClientTimeout(total=timeout) self._headers = { "User-Agent": self.DEFAULT_USER_AGENT, @@ -39,40 +85,102 @@ def __init__( if headers: self._headers.update(headers) - self._client: httpx.AsyncClient | None = None + self._session = session + self._owns_session = session is None @property - def client(self) -> httpx.AsyncClient: - """Get or create the async HTTP client.""" - if self._client is None or self._client.is_closed: - self._client = httpx.AsyncClient( - base_url=self.base_url, - timeout=self.timeout, - headers=self._headers, - http2=True, - follow_redirects=True, - ) - return self._client + def session(self) -> aiohttp.ClientSession: + """Get or lazily create the underlying aiohttp session.""" + if self._session is None or self._session.closed: + self._session = aiohttp.ClientSession() + self._owns_session = True + return self._session @property def is_closed(self) -> bool: - """Check if the client is closed.""" - return self._client is None or self._client.is_closed + """Check whether the underlying session is closed or absent.""" + return self._session is None or self._session.closed async def close(self) -> None: - """Close the HTTP client and release resources.""" - if self._client is not None and not self._client.is_closed: - await self._client.aclose() - self._client = None + """Close the session and release resources. + + Only sessions created (owned) by this client are closed; an injected + session is left untouched for its owner to manage. + """ + if self._owns_session: + if self._session is not None and not self._session.closed: + await self._session.close() + self._session = None + + def _build_url(self, endpoint: str) -> str: + """Resolve an endpoint path against the base URL.""" + if endpoint.startswith(("http://", "https://")): + return endpoint + return f"{self.base_url}{endpoint}" + + def _merge_headers(self, headers: dict[str, str] | None) -> dict[str, str]: + """Combine the default headers with any per-request overrides.""" + if not headers: + return self._headers + merged = dict(self._headers) + merged.update(headers) + return merged + + @staticmethod + def _prepare_params(params: dict[str, Any] | None) -> dict[str, str] | None: + """Coerce query parameters into the string values aiohttp accepts. + + aiohttp rejects ``bool``, ``None`` and nested structures in query + params, so booleans become ``"true"``/``"false"``, ``None`` values are + dropped and dicts/lists are JSON-encoded. + """ + if not params: + return None + prepared: dict[str, str] = {} + for key, value in params.items(): + if value is None: + continue + if isinstance(value, bool): + prepared[key] = "true" if value else "false" + elif isinstance(value, (dict, list)): + prepared[key] = json.dumps(value, separators=(",", ":")) + else: + prepared[key] = str(value) + return prepared + + async def _request( + self, + method: str, + endpoint: str, + *, + params: dict[str, Any] | None = None, + data: dict[str, Any] | None = None, + json_data: dict[str, Any] | None = None, + headers: dict[str, str] | None = None, + ) -> Response: + """Perform a request and return a fully buffered :class:`Response`.""" + response = await self.session.request( + method, + self._build_url(endpoint), + params=self._prepare_params(params), + data=data, + json=json_data, + headers=self._merge_headers(headers), + timeout=self._timeout, + ) + # Buffer the body so the connection is released regardless of how the + # caller consumes the response, and so ``Response.json`` works later. + body = await response.read() + return Response(response.status, body, response) async def get( self, endpoint: str, params: dict[str, Any] | None = None, headers: dict[str, str] | None = None, - ) -> httpx.Response: + ) -> Response: """Make a GET request.""" - return await self.client.get(endpoint, params=params, headers=headers) + return await self._request("GET", endpoint, params=params, headers=headers) async def post( self, @@ -80,18 +188,18 @@ async def post( data: dict[str, Any] | None = None, json: dict[str, Any] | None = None, headers: dict[str, str] | None = None, - ) -> httpx.Response: + ) -> Response: """Make a POST request.""" - return await self.client.post(endpoint, data=data, json=json, headers=headers) + return await self._request("POST", endpoint, data=data, json_data=json, headers=headers) async def delete( self, endpoint: str, params: dict[str, Any] | None = None, headers: dict[str, str] | None = None, - ) -> httpx.Response: + ) -> Response: """Make a DELETE request.""" - return await self.client.delete(endpoint, params=params, headers=headers) + return await self._request("DELETE", endpoint, params=params, headers=headers) async def __aenter__(self) -> HttpClient: """Async context manager entry.""" diff --git a/rohlik_api/services/account.py b/rohlik_api/services/account.py index 7ee8346..f955906 100644 --- a/rohlik_api/services/account.py +++ b/rohlik_api/services/account.py @@ -5,10 +5,9 @@ import logging from typing import Any -import httpx - from ..endpoints import Endpoints from ..errors import APIRequestFailedError +from ..http_client import HTTP_ERRORS from ..models import ShoppingList from .base import BaseService @@ -66,6 +65,6 @@ async def get_shopping_list(self, shopping_list_id: str) -> ShoppingList: response = await self._http.get(url) response.raise_for_status() return ShoppingList.from_api(response.json()) - except httpx.HTTPError as err: + except HTTP_ERRORS as err: _LOGGER.error("Request failed: %s", err) raise APIRequestFailedError(f"Request failed: {err}") from err diff --git a/rohlik_api/services/base.py b/rohlik_api/services/base.py index d935533..4860ba2 100644 --- a/rohlik_api/services/base.py +++ b/rohlik_api/services/base.py @@ -5,10 +5,8 @@ import logging from typing import Any -import httpx - from ..auth import AuthManager -from ..http_client import HttpClient +from ..http_client import HTTP_ERRORS, HttpClient _LOGGER = logging.getLogger(__name__) @@ -60,6 +58,6 @@ async def _fetch_endpoint( response.raise_for_status() data: dict[str, Any] = response.json() return data - except httpx.HTTPError as err: + except HTTP_ERRORS as err: _LOGGER.warning("Error fetching %s: %s", error_context, err) return None diff --git a/rohlik_api/services/cart.py b/rohlik_api/services/cart.py index b875683..330349c 100644 --- a/rohlik_api/services/cart.py +++ b/rohlik_api/services/cart.py @@ -5,10 +5,9 @@ import logging from typing import Any -import httpx - from ..endpoints import Endpoints from ..errors import APIRequestFailedError +from ..http_client import HTTP_ERRORS from ..models import Cart from .base import BaseService @@ -33,7 +32,7 @@ async def get_content(self) -> Cart: response = await self._http.get(Endpoints.CART) response.raise_for_status() return Cart.from_api(response.json()) - except httpx.HTTPError as err: + except HTTP_ERRORS as err: _LOGGER.error("Request failed: %s", err) raise APIRequestFailedError(f"Failed to fetch cart: {err}") from err @@ -64,7 +63,7 @@ async def add_items(self, product_list: list[dict[str, Any]]) -> list[int]: response = await self._http.post(Endpoints.CART, json=cart_payload) response.raise_for_status() added_products.append(product_id) - except httpx.HTTPError as err: + except HTTP_ERRORS as err: _LOGGER.warning("Error adding %s due to %s", product_id, err) return added_products @@ -86,6 +85,6 @@ async def delete_item(self, order_field_id: str) -> None: Endpoints.CART, params={"orderFieldId": order_field_id} ) response.raise_for_status() - except httpx.HTTPError as err: + except HTTP_ERRORS as err: _LOGGER.error("Error deleting item with orderFieldId %s: %s", order_field_id, err) raise APIRequestFailedError(f"Failed to delete item from cart: {err}") from err diff --git a/rohlik_api/services/orders.py b/rohlik_api/services/orders.py index 055fdc4..9082fa8 100644 --- a/rohlik_api/services/orders.py +++ b/rohlik_api/services/orders.py @@ -5,9 +5,8 @@ import logging from typing import Any -import httpx - from ..endpoints import Endpoints +from ..http_client import HTTP_ERRORS from .base import BaseService _LOGGER = logging.getLogger(__name__) @@ -50,6 +49,6 @@ async def get_delivered(self, limit: int = 50, offset: int = 0) -> list[dict[str response.raise_for_status() orders: list[dict[str, Any]] = response.json() return orders - except httpx.HTTPError as err: + except HTTP_ERRORS as err: _LOGGER.warning("Error fetching delivered orders: %s", err) return None diff --git a/rohlik_api/services/products.py b/rohlik_api/services/products.py index a3c68e2..4693eaf 100644 --- a/rohlik_api/services/products.py +++ b/rohlik_api/services/products.py @@ -4,9 +4,8 @@ import logging -import httpx - from ..endpoints import Endpoints +from ..http_client import HTTP_ERRORS from ..models import AISummary, ProductComposition, ProductPrice, ProductSearchResult, SearchResults from .base import BaseService @@ -48,7 +47,7 @@ async def search( response = await self._http.get(Endpoints.SEARCH, params=search_payload) response.raise_for_status() found_products = response.json().get("data", {}).get("productList", []) - except httpx.HTTPError as err: + except HTTP_ERRORS as err: _LOGGER.warning("Request failed: %s", err) return None @@ -85,7 +84,7 @@ async def get_ai_summary(self, product_id: int) -> AISummary | None: response = await self._http.get(Endpoints.product_ai_summary(product_id)) response.raise_for_status() return AISummary.from_api(response.json()) - except httpx.HTTPError as err: + except HTTP_ERRORS as err: _LOGGER.warning("Error fetching AI summary for product %s: %s", product_id, err) return None @@ -104,7 +103,7 @@ async def get_composition(self, product_id: int) -> ProductComposition | None: response = await self._http.get(Endpoints.product_composition(product_id)) response.raise_for_status() return ProductComposition.from_api(response.json()) - except httpx.HTTPError as err: + except HTTP_ERRORS as err: _LOGGER.warning("Error fetching composition for product %s: %s", product_id, err) return None @@ -123,6 +122,6 @@ async def get_price(self, product_id: int) -> ProductPrice | None: response = await self._http.get(Endpoints.product_price(product_id)) response.raise_for_status() return ProductPrice.from_api(response.json()) - except httpx.HTTPError as err: + except HTTP_ERRORS as err: _LOGGER.warning("Error fetching price for product %s: %s", product_id, err) return None diff --git a/rohlik_api/services/recipes.py b/rohlik_api/services/recipes.py index b673c22..855b22a 100644 --- a/rohlik_api/services/recipes.py +++ b/rohlik_api/services/recipes.py @@ -4,9 +4,8 @@ import logging -import httpx - from ..endpoints import Endpoints +from ..http_client import HTTP_ERRORS from ..models import IngredientProducts, RecipeDetail, RecipeSearchResults from .base import BaseService @@ -37,7 +36,7 @@ async def search( response = await self._http.get(url) response.raise_for_status() return RecipeSearchResults.from_api(response.json()) - except httpx.HTTPError as err: + except HTTP_ERRORS as err: _LOGGER.warning("Error searching recipes: %s", err) return None @@ -57,7 +56,7 @@ async def get_detail(self, recipe_id: int) -> RecipeDetail | None: response = await self._http.get(Endpoints.recipe_detail(recipe_id)) response.raise_for_status() return RecipeDetail.from_api(response.json()) - except httpx.HTTPError as err: + except HTTP_ERRORS as err: _LOGGER.warning("Error fetching recipe detail: %s", err) return None @@ -83,6 +82,6 @@ async def get_ingredient_products( response = await self._http.post(Endpoints.INGREDIENT_PRODUCTS, json=payload) response.raise_for_status() return IngredientProducts.from_api(response.json()) - except httpx.HTTPError as err: + except HTTP_ERRORS as err: _LOGGER.warning("Error fetching ingredient products: %s", err) return None diff --git a/tests/test_auth.py b/tests/test_auth.py index 7bb7095..642c446 100644 --- a/tests/test_auth.py +++ b/tests/test_auth.py @@ -9,7 +9,7 @@ def _response(payload): - """Build a mock httpx response returning the given JSON payload.""" + """Build a mock response returning the given JSON payload.""" resp = MagicMock() resp.json.return_value = payload return resp diff --git a/tests/test_client.py b/tests/test_client.py index e84f1f1..10551b2 100644 --- a/tests/test_client.py +++ b/tests/test_client.py @@ -19,7 +19,7 @@ def test_client_initialization(self): client = RohlikAPI(username=TEST_USERNAME, password=TEST_PASSWORD) assert client.base_url == "https://www.rohlik.cz" assert client.timeout == 30.0 - assert client._http._client is None # Lazy initialization + assert client._http._session is None # Lazy initialization def test_client_with_credentials(self): """Test client initializes with username and password.""" @@ -65,13 +65,14 @@ def test_client_default_headers(self): assert "Accept" in client._http._headers assert client._http._headers["Accept"] == "application/json" - def test_client_lazy_initialization(self): - """Test that client is lazily initialized.""" + async def test_client_lazy_initialization(self): + """Test that the session is lazily initialized.""" client = RohlikAPI(username=TEST_USERNAME, password=TEST_PASSWORD) - assert client._http._client is None - # Accessing client property creates the client - _ = client.client - assert client._http._client is not None + assert client._http._session is None + # Accessing the session property creates the session + _ = client.session + assert client._http._session is not None + await client.close() def test_client_base_url_trailing_slash(self): """Test that trailing slash is removed from base URL.""" @@ -94,9 +95,9 @@ async def test_async_context_manager(self): async def test_client_close(self): """Test that client closes without error.""" client = RohlikAPI(username=TEST_USERNAME, password=TEST_PASSWORD, auto_login=False) - _ = client.client # Create the client + _ = client.session # Create the session await client.close() - assert client._http._client is None + assert client._http._session is None class TestMaskData: diff --git a/tests/test_http_client.py b/tests/test_http_client.py index 497ec58..76a6951 100644 --- a/tests/test_http_client.py +++ b/tests/test_http_client.py @@ -1,5 +1,7 @@ """Tests for the HttpClient class.""" +import aiohttp + from rohlik_api import BASE_URL from rohlik_api.http_client import HttpClient @@ -12,7 +14,7 @@ def test_default_initialization(self): client = HttpClient() assert client.base_url == BASE_URL assert client.timeout == 30.0 - assert client._client is None + assert client._session is None def test_custom_base_url(self): """Test HttpClient with custom base URL.""" @@ -49,52 +51,78 @@ def test_custom_headers_merged(self): class TestHttpClientLazyInitialization: """Tests for HttpClient lazy initialization.""" - def test_client_is_none_initially(self): - """Test that internal client is None before first use.""" + def test_session_is_none_initially(self): + """Test that internal session is None before first use.""" http = HttpClient() - assert http._client is None + assert http._session is None - def test_client_created_on_access(self): - """Test that accessing client property creates the client.""" + async def test_session_created_on_access(self): + """Test that accessing the session property creates the session.""" http = HttpClient() - _ = http.client - assert http._client is not None + assert http.session is not None + assert http._session is not None + await http.close() def test_is_closed_initially_true(self): - """Test that is_closed returns True when client not created.""" + """Test that is_closed returns True when session not created.""" http = HttpClient() assert http.is_closed is True - def test_is_closed_false_after_access(self): - """Test that is_closed returns False after client created.""" + async def test_is_closed_false_after_access(self): + """Test that is_closed returns False after session created.""" http = HttpClient() - _ = http.client + _ = http.session assert http.is_closed is False + await http.close() class TestHttpClientClose: """Tests for HttpClient close functionality.""" - async def test_close_without_client(self): - """Test closing when client was never created.""" + async def test_close_without_session(self): + """Test closing when session was never created.""" http = HttpClient() await http.close() # Should not raise - assert http._client is None + assert http._session is None - async def test_close_with_client(self): - """Test closing after client was created.""" + async def test_close_with_session(self): + """Test closing after the session was created.""" http = HttpClient() - _ = http.client # Create client + _ = http.session # Create session await http.close() - assert http._client is None + assert http._session is None async def test_close_multiple_times(self): """Test that closing multiple times is safe.""" http = HttpClient() - _ = http.client + _ = http.session await http.close() await http.close() # Should not raise - assert http._client is None + assert http._session is None + + +class TestHttpClientInjectedSession: + """Tests for reusing an externally managed aiohttp session.""" + + async def test_injected_session_is_used(self): + """An injected session is returned instead of creating a new one.""" + session = aiohttp.ClientSession() + try: + http = HttpClient(session=session) + assert http.session is session + assert http.is_closed is False + finally: + await session.close() + + async def test_close_does_not_close_injected_session(self): + """Closing the client must not close an externally owned session.""" + session = aiohttp.ClientSession() + try: + http = HttpClient(session=session) + await http.close() + assert session.closed is False + finally: + await session.close() class TestHttpClientContextManager: @@ -107,8 +135,8 @@ async def test_context_manager_entry(self): assert isinstance(http, HttpClient) async def test_context_manager_closes_on_exit(self): - """Test that context manager closes client on exit.""" + """Test that context manager closes the session on exit.""" http = HttpClient() async with http: - _ = http.client # Ensure client is created - assert http._client is None + _ = http.session # Ensure session is created + assert http._session is None diff --git a/tests/test_recipes.py b/tests/test_recipes.py index 57b579a..150b263 100644 --- a/tests/test_recipes.py +++ b/tests/test_recipes.py @@ -75,9 +75,9 @@ async def test_search_returns_recipes(self, mock_http, mock_auth): async def test_search_returns_none_on_error(self, mock_http, mock_auth): """Test search returns None when request fails.""" - import httpx + import aiohttp - mock_http.get.side_effect = httpx.HTTPError("Connection failed") + mock_http.get.side_effect = aiohttp.ClientError("Connection failed") service = RecipeService(mock_http, mock_auth) result = await service.search("test") @@ -164,9 +164,9 @@ async def test_get_detail_returns_recipe(self, mock_http, mock_auth): async def test_get_detail_returns_none_on_error(self, mock_http, mock_auth): """Test get_detail returns None when request fails.""" - import httpx + import aiohttp - mock_http.get.side_effect = httpx.HTTPError("Connection failed") + mock_http.get.side_effect = aiohttp.ClientError("Connection failed") service = RecipeService(mock_http, mock_auth) result = await service.get_detail(59) @@ -246,9 +246,9 @@ async def test_get_ingredient_products_sends_correct_payload(self, mock_http, mo async def test_get_ingredient_products_returns_none_on_error(self, mock_http, mock_auth): """Test get_ingredient_products returns None when request fails.""" - import httpx + import aiohttp - mock_http.post.side_effect = httpx.HTTPError("Connection failed") + mock_http.post.side_effect = aiohttp.ClientError("Connection failed") service = RecipeService(mock_http, mock_auth) result = await service.get_ingredient_products([102]) diff --git a/tests/test_services.py b/tests/test_services.py index cb73475..6c59ef3 100644 --- a/tests/test_services.py +++ b/tests/test_services.py @@ -141,9 +141,9 @@ async def test_search_returns_empty_when_no_products(self, mock_http, mock_auth) async def test_search_returns_none_on_error(self, mock_http, mock_auth): """Test search returns None when the request fails.""" - import httpx + import aiohttp - mock_http.get.side_effect = httpx.HTTPError("Connection failed") + mock_http.get.side_effect = aiohttp.ClientError("Connection failed") service = ProductService(mock_http, mock_auth) result = await service.search("test") @@ -202,9 +202,9 @@ async def test_get_ai_summary_returns_data(self, mock_http, mock_auth): async def test_get_ai_summary_returns_none_on_error(self, mock_http, mock_auth): """Test get_ai_summary returns None on error.""" - import httpx + import aiohttp - mock_http.get.side_effect = httpx.HTTPError("Connection failed") + mock_http.get.side_effect = aiohttp.ClientError("Connection failed") service = ProductService(mock_http, mock_auth) result = await service.get_ai_summary(1384964) @@ -255,9 +255,9 @@ async def test_get_composition_returns_data(self, mock_http, mock_auth): async def test_get_composition_returns_none_on_error(self, mock_http, mock_auth): """Test get_composition returns None on error.""" - import httpx + import aiohttp - mock_http.get.side_effect = httpx.HTTPError("Connection failed") + mock_http.get.side_effect = aiohttp.ClientError("Connection failed") service = ProductService(mock_http, mock_auth) result = await service.get_composition(1425155) @@ -288,9 +288,9 @@ async def test_get_price_returns_data(self, mock_http, mock_auth): async def test_get_price_returns_none_on_error(self, mock_http, mock_auth): """Test get_price returns None on error.""" - import httpx + import aiohttp - mock_http.get.side_effect = httpx.HTTPError("Connection failed") + mock_http.get.side_effect = aiohttp.ClientError("Connection failed") service = ProductService(mock_http, mock_auth) result = await service.get_price(1425155) From e105b833432625495ffb4fa742a4681885d4295a Mon Sep 17 00:00:00 2001 From: Claude Date: Sat, 27 Jun 2026 15:01:41 +0000 Subject: [PATCH 2/4] Add session re-auth and order/product endpoints from HA integration Upstream improvements proven out in the HA-RohlikCZ integration: - Transparent re-authentication: HttpClient now retries a request once after invoking an on_unauthorized callback when it receives HTTP 401, so a long-lived client recovers from an expired session instead of failing every call. RohlikAPI wires this to AuthManager.relogin (login/logout requests are exempt to avoid recursion; concurrent 401s are serialized). - New endpoints: orders.get_detail(id), orders.get_all_delivered() (paginates the delivered-orders endpoint), products.get_detail(id), and products.get_categories(id). - 404 -> None convention for optional-resource fetches (order/product detail, product categories), so a discontinued product reads as "not found" rather than an error. Add tests for the re-auth retry path, relogin, the new endpoints, and the 404 handling. Update README. Co-Authored-By: Claude Opus 4.8 Claude-Session: https://claude.ai/code/session_01PmotTwydT558t4JHd5Cnwm --- README.md | 12 +++- rohlik_api/auth.py | 12 ++++ rohlik_api/client.py | 4 ++ rohlik_api/endpoints.py | 15 +++++ rohlik_api/http_client.py | 72 +++++++++++++++++++----- rohlik_api/services/orders.py | 48 ++++++++++++++++ rohlik_api/services/products.py | 49 +++++++++++++++++ tests/test_auth.py | 22 ++++++++ tests/test_client.py | 5 ++ tests/test_http_client.py | 85 ++++++++++++++++++++++++++++- tests/test_services.py | 97 +++++++++++++++++++++++++++++++++ 11 files changed, 403 insertions(+), 18 deletions(-) diff --git a/README.md b/README.md index ed27435..11b4676 100644 --- a/README.md +++ b/README.md @@ -34,7 +34,7 @@ online grocery service β€” search products, manage your cart, browse recipes ## Features - πŸš€ Built on aiohttp; bring your own session (e.g. Home Assistant's shared session) -- πŸ” Automatic login/logout and session management +- πŸ” Automatic login/logout, plus transparent re-authentication when a session expires (HTTP 401) - 🎯 Clean, service-based API (`client.cart`, `client.products`, …) - 🧩 Fully typed dataclass models for parsed responses (`py.typed`) - πŸ”„ Works as an async context manager @@ -171,6 +171,12 @@ composition = await client.products.get_composition(product_id=1425155) # Current price -> ProductPrice | None price = await client.products.get_price(product_id=1425155) + +# Raw product detail (brand, attributes, …) -> dict | None (None on 404) +detail = await client.products.get_detail(product_id=1425155) + +# Category hierarchy -> list[dict] | None (None if discontinued / 404) +categories = await client.products.get_categories(product_id=1425155) ``` ### Orders service (`client.orders`) @@ -178,7 +184,9 @@ price = await client.products.get_price(product_id=1425155) ```python next_order = await client.orders.get_next() # upcoming order last_order = await client.orders.get_last() # last delivered order -orders = await client.orders.get_delivered(limit=50, offset=0) # history +orders = await client.orders.get_delivered(limit=50, offset=0) # one history page +all_orders = await client.orders.get_all_delivered() # every page, paginated +detail = await client.orders.get_detail(order_id=12345678) # full order incl. items ``` ### Delivery service (`client.delivery`) diff --git a/rohlik_api/auth.py b/rohlik_api/auth.py index 54b0212..df2a300 100644 --- a/rohlik_api/auth.py +++ b/rohlik_api/auth.py @@ -148,6 +148,18 @@ async def ensure_logged_in(self) -> None: if not self._is_logged_in: await self.login() + async def relogin(self) -> dict[str, Any]: + """Force a fresh login after a session expiry (HTTP 401). + + Clears the cached session state so :meth:`login` performs a new request + instead of returning the stale cached response, then logs in again. + + Returns: + dict: The JSON response from the new login. + """ + self._reset_session() + return await self.login() + def _reset_session(self) -> None: """Clear all session state so the next login re-fetches it. diff --git a/rohlik_api/client.py b/rohlik_api/client.py index d394926..e28984b 100644 --- a/rohlik_api/client.py +++ b/rohlik_api/client.py @@ -88,6 +88,10 @@ def __init__( password=password, ) + # Wire up transparent re-authentication: when any request hits HTTP 401 + # (expired session), the HTTP client re-logs in and retries once. + self._http.set_unauthorized_handler(self._auth.relogin) + # Initialize services self._cart = CartService(self._http, self._auth) self._products = ProductService(self._http, self._auth) diff --git a/rohlik_api/endpoints.py b/rohlik_api/endpoints.py index cb27015..f268af6 100644 --- a/rohlik_api/endpoints.py +++ b/rohlik_api/endpoints.py @@ -55,6 +55,21 @@ def recipe_detail(cls, recipe_id: int) -> str: """Build recipe detail endpoint URL.""" return f"/services/frontend-service/recipe/{recipe_id}" + @classmethod + def order_detail(cls, order_id: int) -> str: + """Build order detail endpoint URL (full order including items).""" + return f"/api/v3/orders/{order_id}" + + @classmethod + def product_detail(cls, product_id: int) -> str: + """Build product detail endpoint URL.""" + return f"/api/v1/products/{product_id}" + + @classmethod + def product_categories(cls, product_id: int) -> str: + """Build product category-hierarchy endpoint URL.""" + return f"/api/v1/products/{product_id}/categories" + @classmethod def product_ai_summary(cls, product_id: int) -> str: """Build product AI summary endpoint URL.""" diff --git a/rohlik_api/http_client.py b/rohlik_api/http_client.py index ecb6585..bffec87 100644 --- a/rohlik_api/http_client.py +++ b/rohlik_api/http_client.py @@ -2,14 +2,16 @@ from __future__ import annotations +import asyncio import json import logging +from collections.abc import Awaitable, Callable from importlib.metadata import PackageNotFoundError, version from typing import Any import aiohttp -from .endpoints import BASE_URL +from .endpoints import BASE_URL, Endpoints _LOGGER = logging.getLogger(__name__) @@ -63,16 +65,24 @@ class HttpClient: timeout: Request timeout in seconds. headers: Optional custom headers added to every request. session: Optional externally managed aiohttp session to reuse. + on_unauthorized: Optional async callback invoked when a request returns + HTTP 401 (an expired session). After it runs, the request is retried + once. Login/logout requests are exempt to avoid recursion. """ DEFAULT_USER_AGENT = f"rohlik-api-python/{_VERSION}" + # Endpoints that must never trigger the re-auth retry (the callback logs in + # via the login endpoint, so retrying it would recurse). + _NO_REAUTH_ENDPOINTS = (Endpoints.LOGIN, Endpoints.LOGOUT) + def __init__( self, base_url: str = BASE_URL, timeout: float = 30.0, headers: dict[str, str] | None = None, session: aiohttp.ClientSession | None = None, + on_unauthorized: Callable[[], Awaitable[Any]] | None = None, ) -> None: self.base_url = base_url.rstrip("/") self.timeout = timeout @@ -88,6 +98,13 @@ def __init__( self._session = session self._owns_session = session is None + self._on_unauthorized = on_unauthorized + self._reauth_lock = asyncio.Lock() + + def set_unauthorized_handler(self, handler: Callable[[], Awaitable[Any]] | None) -> None: + """Register the callback used to re-authenticate on an HTTP 401.""" + self._on_unauthorized = handler + @property def session(self) -> aiohttp.ClientSession: """Get or lazily create the underlying aiohttp session.""" @@ -158,20 +175,45 @@ async def _request( json_data: dict[str, Any] | None = None, headers: dict[str, str] | None = None, ) -> Response: - """Perform a request and return a fully buffered :class:`Response`.""" - response = await self.session.request( - method, - self._build_url(endpoint), - params=self._prepare_params(params), - data=data, - json=json_data, - headers=self._merge_headers(headers), - timeout=self._timeout, - ) - # Buffer the body so the connection is released regardless of how the - # caller consumes the response, and so ``Response.json`` works later. - body = await response.read() - return Response(response.status, body, response) + """Perform a request and return a fully buffered :class:`Response`. + + On an HTTP 401 the registered re-auth callback (if any) is invoked once + and the request is retried, transparently recovering from an expired + session on a long-lived client. + """ + url = self._build_url(endpoint) + prepared_params = self._prepare_params(params) + merged_headers = self._merge_headers(headers) + + async def _send() -> tuple[int, bytes, aiohttp.ClientResponse]: + response = await self.session.request( + method, + url, + params=prepared_params, + data=data, + json=json_data, + headers=merged_headers, + timeout=self._timeout, + ) + # Buffer the body so the connection is released regardless of how + # the caller consumes the response, and so ``Response.json`` works + # later. + body = await response.read() + return response.status, body, response + + status, body, response = await _send() + + if ( + status == 401 + and self._on_unauthorized is not None + and endpoint not in self._NO_REAUTH_ENDPOINTS + ): + # Serialize re-auth so concurrent 401s trigger a single login. + async with self._reauth_lock: + await self._on_unauthorized() + status, body, response = await _send() + + return Response(status, body, response) async def get( self, diff --git a/rohlik_api/services/orders.py b/rohlik_api/services/orders.py index 9082fa8..a8f6881 100644 --- a/rohlik_api/services/orders.py +++ b/rohlik_api/services/orders.py @@ -52,3 +52,51 @@ async def get_delivered(self, limit: int = 50, offset: int = 0) -> list[dict[str except HTTP_ERRORS as err: _LOGGER.warning("Error fetching delivered orders: %s", err) return None + + async def get_all_delivered(self, page_size: int = 50) -> list[dict[str, Any]]: + """Get every delivered order by paginating until the list is exhausted. + + Args: + page_size: Number of orders fetched per request. + + Returns: + list: All delivered orders (empty if there are none or the first + page fails). + """ + await self._ensure_logged_in() + + all_orders: list[dict[str, Any]] = [] + offset = 0 + while True: + page = await self.get_delivered(limit=page_size, offset=offset) + if not page: + break + all_orders.extend(page) + if len(page) < page_size: + break + offset += page_size + + return all_orders + + async def get_detail(self, order_id: int) -> dict[str, Any] | None: + """Get full detail for a single order, including its line items. + + Args: + order_id: The ID of the order. + + Returns: + dict: The order detail, or None if the order does not exist (404) + or the request fails. + """ + await self._ensure_logged_in() + + try: + response = await self._http.get(Endpoints.order_detail(order_id)) + if response.status == 404: + return None + response.raise_for_status() + detail: dict[str, Any] = response.json() + return detail + except HTTP_ERRORS as err: + _LOGGER.warning("Error fetching order detail for %s: %s", order_id, err) + return None diff --git a/rohlik_api/services/products.py b/rohlik_api/services/products.py index 4693eaf..758922e 100644 --- a/rohlik_api/services/products.py +++ b/rohlik_api/services/products.py @@ -3,6 +3,7 @@ from __future__ import annotations import logging +from typing import Any from ..endpoints import Endpoints from ..http_client import HTTP_ERRORS @@ -125,3 +126,51 @@ async def get_price(self, product_id: int) -> ProductPrice | None: except HTTP_ERRORS as err: _LOGGER.warning("Error fetching price for product %s: %s", product_id, err) return None + + async def get_detail(self, product_id: int) -> dict[str, Any] | None: + """Get the full product detail (brand, attributes, etc.). + + Args: + product_id: The ID of the product. + + Returns: + dict: The raw product detail, or None if the product does not exist + (404) or the request fails. + """ + await self._ensure_logged_in() + + try: + response = await self._http.get(Endpoints.product_detail(product_id)) + if response.status == 404: + return None + response.raise_for_status() + detail: dict[str, Any] = response.json() + return detail + except HTTP_ERRORS as err: + _LOGGER.warning("Error fetching detail for product %s: %s", product_id, err) + return None + + async def get_categories(self, product_id: int) -> list[dict[str, Any]] | None: + """Get the category hierarchy for a product. + + Args: + product_id: The ID of the product. + + Returns: + list: The category hierarchy (possibly empty), or None if the + product no longer exists (404) or the request fails. A 404 typically + means the product has been discontinued. + """ + await self._ensure_logged_in() + + try: + response = await self._http.get(Endpoints.product_categories(product_id)) + if response.status == 404: + _LOGGER.debug("Product %s not found (discontinued)", product_id) + return None + response.raise_for_status() + categories: list[dict[str, Any]] = response.json().get("categories", []) + return categories + except HTTP_ERRORS as err: + _LOGGER.warning("Error fetching categories for product %s: %s", product_id, err) + return None diff --git a/tests/test_auth.py b/tests/test_auth.py index 642c446..d0a5d88 100644 --- a/tests/test_auth.py +++ b/tests/test_auth.py @@ -125,3 +125,25 @@ async def test_logout_resets_session_state(self): assert auth.is_logged_in is False assert auth.user_id is None assert auth.address_id is None + + async def test_relogin_forces_fresh_login(self): + """relogin clears cached state and performs a new login request.""" + http = MagicMock(spec=HttpClient) + http.post = AsyncMock( + side_effect=[ + _response({"status": 200, "data": {"user": {"id": 1}, "address": {"id": 2}}}), + _response({"status": 200, "data": {"user": {"id": 9}, "address": {"id": 8}}}), + ] + ) + auth = AuthManager(http, "user@example.com", "password123") + + await auth.login() + assert auth.user_id == 1 + + # relogin must bypass the "already logged in" short-circuit. + await auth.relogin() + + assert auth.is_logged_in is True + assert auth.user_id == 9 + assert auth.address_id == 8 + assert http.post.call_count == 2 diff --git a/tests/test_client.py b/tests/test_client.py index 10551b2..9bfaf8f 100644 --- a/tests/test_client.py +++ b/tests/test_client.py @@ -183,6 +183,11 @@ async def test_context_manager_auto_login(self): async with client: client._auth.login.assert_awaited_once() + def test_reauth_handler_is_wired(self): + """The HTTP client's 401 handler delegates to the auth manager's relogin.""" + client = RohlikAPI(username=TEST_USERNAME, password=TEST_PASSWORD) + assert client._http._on_unauthorized == client._auth.relogin + def test_user_and_address_id_properties(self): """Test that user_id and address_id properties expose auth state.""" client = RohlikAPI(username=TEST_USERNAME, password=TEST_PASSWORD) diff --git a/tests/test_http_client.py b/tests/test_http_client.py index 76a6951..0eaaaae 100644 --- a/tests/test_http_client.py +++ b/tests/test_http_client.py @@ -1,11 +1,41 @@ """Tests for the HttpClient class.""" +from unittest.mock import AsyncMock + import aiohttp -from rohlik_api import BASE_URL +from rohlik_api import BASE_URL, Endpoints from rohlik_api.http_client import HttpClient +class _FakeResponse: + """Minimal stand-in for an aiohttp ClientResponse.""" + + def __init__(self, status: int, body: bytes = b"{}") -> None: + self.status = status + self._body = body + + async def read(self) -> bytes: + return self._body + + def raise_for_status(self) -> None: + if self.status >= 400: + raise aiohttp.ClientResponseError(None, (), status=self.status) + + +class _FakeSession: + """Fake aiohttp session that yields a scripted sequence of responses.""" + + def __init__(self, responses: list[_FakeResponse]) -> None: + self._responses = list(responses) + self.closed = False + self.calls: list[tuple[str, str]] = [] + + async def request(self, method: str, url: str, **kwargs: object) -> _FakeResponse: + self.calls.append((method, url)) + return self._responses.pop(0) + + class TestHttpClientInitialization: """Tests for HttpClient initialization.""" @@ -125,6 +155,59 @@ async def test_close_does_not_close_injected_session(self): await session.close() +class TestHttpClientReauth: + """Tests for transparent re-authentication on HTTP 401.""" + + async def test_retries_once_after_reauth_on_401(self): + """A 401 triggers the handler and the request is retried once.""" + session = _FakeSession([_FakeResponse(401), _FakeResponse(200, b'{"ok": true}')]) + handler = AsyncMock() + http = HttpClient(session=session) + http.set_unauthorized_handler(handler) + + response = await http.get("/api/v3/orders/upcoming") + + handler.assert_awaited_once() + assert response.status == 200 + assert response.json() == {"ok": True} + assert len(session.calls) == 2 + + async def test_no_retry_without_handler(self): + """Without a handler, a 401 is returned untouched.""" + session = _FakeSession([_FakeResponse(401)]) + http = HttpClient(session=session) + + response = await http.get("/api/v3/orders/upcoming") + + assert response.status == 401 + assert len(session.calls) == 1 + + async def test_login_endpoint_is_exempt_from_reauth(self): + """A 401 on the login endpoint must not invoke the handler (no recursion).""" + session = _FakeSession([_FakeResponse(401)]) + handler = AsyncMock() + http = HttpClient(session=session) + http.set_unauthorized_handler(handler) + + response = await http.post(Endpoints.LOGIN, json={"email": "a", "password": "b"}) + + handler.assert_not_awaited() + assert response.status == 401 + assert len(session.calls) == 1 + + async def test_non_401_does_not_trigger_reauth(self): + """A successful request never invokes the re-auth handler.""" + session = _FakeSession([_FakeResponse(200)]) + handler = AsyncMock() + http = HttpClient(session=session) + http.set_unauthorized_handler(handler) + + await http.get("/api/v3/orders/upcoming") + + handler.assert_not_awaited() + assert len(session.calls) == 1 + + class TestHttpClientContextManager: """Tests for HttpClient async context manager.""" diff --git a/tests/test_services.py b/tests/test_services.py index 6c59ef3..2ad1646 100644 --- a/tests/test_services.py +++ b/tests/test_services.py @@ -297,6 +297,58 @@ async def test_get_price_returns_none_on_error(self, mock_http, mock_auth): assert result is None + async def test_get_detail_returns_data(self, mock_http, mock_auth): + """Test get_detail returns the raw product detail.""" + mock_response = MagicMock() + mock_response.json.return_value = {"id": 123, "brand": "TestBrand"} + mock_response.raise_for_status = MagicMock() + mock_response.status = 200 + mock_http.get.return_value = mock_response + + service = ProductService(mock_http, mock_auth) + result = await service.get_detail(123) + + assert result["brand"] == "TestBrand" + + async def test_get_detail_returns_none_on_404(self, mock_http, mock_auth): + """Test get_detail returns None for a discontinued product.""" + mock_response = MagicMock() + mock_response.status = 404 + mock_response.raise_for_status = MagicMock() + mock_http.get.return_value = mock_response + + service = ProductService(mock_http, mock_auth) + result = await service.get_detail(123) + + assert result is None + + async def test_get_categories_returns_hierarchy(self, mock_http, mock_auth): + """Test get_categories returns the inner category list.""" + mock_response = MagicMock() + mock_response.json.return_value = { + "categories": [{"level": 0, "name": "Food"}, {"level": 1, "name": "Dairy"}] + } + mock_response.raise_for_status = MagicMock() + mock_response.status = 200 + mock_http.get.return_value = mock_response + + service = ProductService(mock_http, mock_auth) + result = await service.get_categories(123) + + assert [c["name"] for c in result] == ["Food", "Dairy"] + + async def test_get_categories_returns_none_on_404(self, mock_http, mock_auth): + """Test get_categories returns None for a discontinued product.""" + mock_response = MagicMock() + mock_response.status = 404 + mock_response.raise_for_status = MagicMock() + mock_http.get.return_value = mock_response + + service = ProductService(mock_http, mock_auth) + result = await service.get_categories(123) + + assert result is None + class TestOrderService: """Tests for OrderService.""" @@ -330,6 +382,51 @@ async def test_get_delivered_with_pagination(self, mock_http, mock_auth): assert len(result) == 2 + async def test_get_all_delivered_paginates(self, mock_http, mock_auth): + """Test get_all_delivered walks pages until a short page ends it.""" + page1 = MagicMock() + page1.json.return_value = [{"id": 1}, {"id": 2}] + page1.raise_for_status = MagicMock() + page1.status = 200 + page2 = MagicMock() + page2.json.return_value = [{"id": 3}] + page2.raise_for_status = MagicMock() + page2.status = 200 + mock_http.get.side_effect = [page1, page2] + + service = OrderService(mock_http, mock_auth) + result = await service.get_all_delivered(page_size=2) + + assert [o["id"] for o in result] == [1, 2, 3] + assert mock_http.get.call_count == 2 + + async def test_get_detail_returns_data(self, mock_http, mock_auth): + """Test get_detail returns the order detail.""" + mock_response = MagicMock() + mock_response.json.return_value = {"id": 42, "items": [{"name": "Milk"}]} + mock_response.raise_for_status = MagicMock() + mock_response.status = 200 + mock_http.get.return_value = mock_response + + service = OrderService(mock_http, mock_auth) + result = await service.get_detail(42) + + assert result["id"] == 42 + assert result["items"][0]["name"] == "Milk" + + async def test_get_detail_returns_none_on_404(self, mock_http, mock_auth): + """Test get_detail returns None when the order does not exist.""" + mock_response = MagicMock() + mock_response.status = 404 + mock_response.raise_for_status = MagicMock() + mock_http.get.return_value = mock_response + + service = OrderService(mock_http, mock_auth) + result = await service.get_detail(999) + + assert result is None + mock_response.raise_for_status.assert_not_called() + class TestDeliveryService: """Tests for DeliveryService.""" From 7fc93b1dcb52f29a22650c40fd581a7b04d5bc0b Mon Sep 17 00:00:00 2001 From: Claude Date: Sat, 27 Jun 2026 15:18:07 +0000 Subject: [PATCH 3/4] Address PR review: re-auth dedup, connection release, session safety - HttpClient._request now buffers the body inside `async with session.request(...)` so the connection is released even if read() raises mid-response. - Concurrent 401s now re-authenticate once: a generation counter lets coroutines queued behind the in-flight re-auth skip a redundant login. - session property only (re)creates owned sessions; a closed injected session is returned as-is so the next request raises "Session is closed" instead of silently spawning a session that bypasses the owner's connector/SSL. - orders.get_all_delivered distinguishes a failed page (None -> warn, stop, possibly-incomplete result) from an empty page (done), and documents it. Add tests for single re-auth under concurrent 401s and for a closed injected session not being replaced. Co-Authored-By: Claude Opus 4.8 Claude-Session: https://claude.ai/code/session_01PmotTwydT558t4JHd5Cnwm --- rohlik_api/http_client.py | 40 ++++++++++++++------- rohlik_api/services/orders.py | 15 ++++++-- tests/test_http_client.py | 67 +++++++++++++++++++++++++++++++++-- 3 files changed, 104 insertions(+), 18 deletions(-) diff --git a/rohlik_api/http_client.py b/rohlik_api/http_client.py index bffec87..d8eea89 100644 --- a/rohlik_api/http_client.py +++ b/rohlik_api/http_client.py @@ -100,6 +100,9 @@ def __init__( self._on_unauthorized = on_unauthorized self._reauth_lock = asyncio.Lock() + # Bumped each time a re-auth completes, so coroutines that queued on the + # lock behind an in-flight re-auth can skip a redundant second login. + self._reauth_generation = 0 def set_unauthorized_handler(self, handler: Callable[[], Awaitable[Any]] | None) -> None: """Register the callback used to re-authenticate on an HTTP 401.""" @@ -107,10 +110,17 @@ def set_unauthorized_handler(self, handler: Callable[[], Awaitable[Any]] | None) @property def session(self) -> aiohttp.ClientSession: - """Get or lazily create the underlying aiohttp session.""" - if self._session is None or self._session.closed: + """Get or lazily create the underlying aiohttp session. + + Only sessions this client owns are (re)created. An injected session that + has been closed by its owner is returned as-is, so the next request + surfaces aiohttp's "Session is closed" error instead of silently + spawning a new session that bypasses the owner's connector/SSL config. + """ + if self._owns_session and (self._session is None or self._session.closed): self._session = aiohttp.ClientSession() - self._owns_session = True + # Owned sessions are created above; injected ones are set in __init__. + assert self._session is not None return self._session @property @@ -186,7 +196,10 @@ async def _request( merged_headers = self._merge_headers(headers) async def _send() -> tuple[int, bytes, aiohttp.ClientResponse]: - response = await self.session.request( + # ``async with`` guarantees the connection is released on every + # path, including if ``read()`` raises mid-response. Buffering the + # body here also lets ``Response.json`` work after release. + async with self.session.request( method, url, params=prepared_params, @@ -194,12 +207,9 @@ async def _send() -> tuple[int, bytes, aiohttp.ClientResponse]: json=json_data, headers=merged_headers, timeout=self._timeout, - ) - # Buffer the body so the connection is released regardless of how - # the caller consumes the response, and so ``Response.json`` works - # later. - body = await response.read() - return response.status, body, response + ) as response: + body = await response.read() + return response.status, body, response status, body, response = await _send() @@ -208,9 +218,15 @@ async def _send() -> tuple[int, bytes, aiohttp.ClientResponse]: and self._on_unauthorized is not None and endpoint not in self._NO_REAUTH_ENDPOINTS ): - # Serialize re-auth so concurrent 401s trigger a single login. + # Serialize re-auth, and use a generation counter so that several + # requests that all hit a 401 at once trigger only one login: the + # first through the lock re-authenticates, the rest see the bumped + # generation and just retry. + seen_generation = self._reauth_generation async with self._reauth_lock: - await self._on_unauthorized() + if self._reauth_generation == seen_generation: + await self._on_unauthorized() + self._reauth_generation += 1 status, body, response = await _send() return Response(status, body, response) diff --git a/rohlik_api/services/orders.py b/rohlik_api/services/orders.py index a8f6881..42fdcb5 100644 --- a/rohlik_api/services/orders.py +++ b/rohlik_api/services/orders.py @@ -60,8 +60,9 @@ async def get_all_delivered(self, page_size: int = 50) -> list[dict[str, Any]]: page_size: Number of orders fetched per request. Returns: - list: All delivered orders (empty if there are none or the first - page fails). + list: All delivered orders (empty if there are none). If a request + fails partway through pagination, the orders gathered so far are + returned and a warning is logged, so the result may be incomplete. """ await self._ensure_logged_in() @@ -69,8 +70,16 @@ async def get_all_delivered(self, page_size: int = 50) -> list[dict[str, Any]]: offset = 0 while True: page = await self.get_delivered(limit=page_size, offset=offset) - if not page: + if page is None: + # Request error (already logged by get_delivered): stop and + # return what we have rather than silently looping forever. + _LOGGER.warning( + "Stopped paginating delivered orders at offset %s; result may be incomplete", + offset, + ) break + if not page: + break # Empty page: genuinely no more orders. all_orders.extend(page) if len(page) < page_size: break diff --git a/tests/test_http_client.py b/tests/test_http_client.py index 0eaaaae..5e3dc7c 100644 --- a/tests/test_http_client.py +++ b/tests/test_http_client.py @@ -1,5 +1,6 @@ """Tests for the HttpClient class.""" +import asyncio from unittest.mock import AsyncMock import aiohttp @@ -9,12 +10,20 @@ class _FakeResponse: - """Minimal stand-in for an aiohttp ClientResponse.""" + """Minimal stand-in for an aiohttp ClientResponse (also its own CM).""" def __init__(self, status: int, body: bytes = b"{}") -> None: self.status = status self._body = body + async def __aenter__(self) -> "_FakeResponse": + # Yield control so concurrent requests can interleave deterministically. + await asyncio.sleep(0) + return self + + async def __aexit__(self, *exc: object) -> bool: + return False + async def read(self) -> bytes: return self._body @@ -24,18 +33,36 @@ def raise_for_status(self) -> None: class _FakeSession: - """Fake aiohttp session that yields a scripted sequence of responses.""" + """Fake aiohttp session that yields a scripted sequence of responses. + + ``request`` returns the response object directly (not a coroutine), so it + works with ``async with session.request(...)`` like the real client. + """ def __init__(self, responses: list[_FakeResponse]) -> None: self._responses = list(responses) self.closed = False self.calls: list[tuple[str, str]] = [] - async def request(self, method: str, url: str, **kwargs: object) -> _FakeResponse: + def request(self, method: str, url: str, **kwargs: object) -> _FakeResponse: self.calls.append((method, url)) return self._responses.pop(0) +class _CountingSession: + """Fake session whose first ``fail_first`` requests return 401, rest 200.""" + + def __init__(self, fail_first: int) -> None: + self.fail_first = fail_first + self.count = 0 + self.closed = False + + def request(self, method: str, url: str, **kwargs: object) -> _FakeResponse: + self.count += 1 + status = 401 if self.count <= self.fail_first else 200 + return _FakeResponse(status) + + class TestHttpClientInitialization: """Tests for HttpClient initialization.""" @@ -154,6 +181,17 @@ async def test_close_does_not_close_injected_session(self): finally: await session.close() + async def test_closed_injected_session_is_not_replaced(self): + """A closed injected session is returned as-is, never silently replaced.""" + session = aiohttp.ClientSession() + await session.close() + http = HttpClient(session=session) + + # The client must not spawn a new owned session in place of the + # caller's (now closed) one. + assert http.session is session + assert http.is_closed is True + class TestHttpClientReauth: """Tests for transparent re-authentication on HTTP 401.""" @@ -207,6 +245,29 @@ async def test_non_401_does_not_trigger_reauth(self): handler.assert_not_awaited() assert len(session.calls) == 1 + async def test_concurrent_401s_trigger_single_reauth(self): + """Several requests hitting 401 at once re-authenticate only once.""" + # Both initial requests 401; both retries succeed -> 4 requests total. + session = _CountingSession(fail_first=2) + calls = 0 + + async def handler() -> None: + nonlocal calls + calls += 1 + await asyncio.sleep(0) # hold the lock long enough to overlap + + http = HttpClient(session=session) + http.set_unauthorized_handler(handler) + + responses = await asyncio.gather( + http.get("/api/v3/orders/upcoming"), + http.get("/api/v3/orders/delivered"), + ) + + assert calls == 1 + assert all(r.status == 200 for r in responses) + assert session.count == 4 + class TestHttpClientContextManager: """Tests for HttpClient async context manager.""" From f75fb6bb19b979f44726e943c01e212c25399da1 Mon Sep 17 00:00:00 2001 From: Claude Date: Sat, 27 Jun 2026 18:25:20 +0000 Subject: [PATCH 4/4] Document model fields and tidy two review nits MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Make the API self-describing for callers (and LLMs) reading only the code: - models.py: add an Attributes section to every dataclass documenting each field's meaning and units (energy in kJ/kcal, prices in CZK, formatted vs numeric price fields, IDs and how they feed back into other calls). Fix RecipeDetail.duration type from int to str β€” the API returns a display string like "Do hodinky", not a number. - client.py: fix the RohlikAPI example to use the typed Cart model (cart.total_price / cart.products) instead of stale dict access. Address the two optional notes from the re-review: - http_client.session: replace the assert with an explicit raise so the "session present" invariant survives `python -O`. - _request: document that if the re-auth callback raises, the error propagates and the request is not retried. Co-Authored-By: Claude Opus 4.8 Claude-Session: https://claude.ai/code/session_01PmotTwydT558t4JHd5Cnwm --- rohlik_api/client.py | 4 +- rohlik_api/http_client.py | 11 +- rohlik_api/models.py | 226 ++++++++++++++++++++++++++++++++++---- 3 files changed, 215 insertions(+), 26 deletions(-) diff --git a/rohlik_api/client.py b/rohlik_api/client.py index e28984b..b700a60 100644 --- a/rohlik_api/client.py +++ b/rohlik_api/client.py @@ -54,7 +54,9 @@ class RohlikAPI: Example: >>> async with RohlikAPI("user@example.com", "password") as client: ... cart = await client.cart.get_content() - ... print(cart["total_price"]) + ... print(cart.total_price, cart.total_items) + ... for item in cart.products: + ... print(item.name, item.quantity, item.price) """ def __init__( diff --git a/rohlik_api/http_client.py b/rohlik_api/http_client.py index d8eea89..7f67e02 100644 --- a/rohlik_api/http_client.py +++ b/rohlik_api/http_client.py @@ -119,8 +119,11 @@ def session(self) -> aiohttp.ClientSession: """ if self._owns_session and (self._session is None or self._session.closed): self._session = aiohttp.ClientSession() - # Owned sessions are created above; injected ones are set in __init__. - assert self._session is not None + if self._session is None: # pragma: no cover - unreachable by construction + # Owned sessions are created above; injected ones are set in + # __init__. A plain ``raise`` (rather than ``assert``) keeps the + # invariant enforced even under ``python -O``. + raise RuntimeError("HTTP session is unexpectedly missing") return self._session @property @@ -189,7 +192,9 @@ async def _request( On an HTTP 401 the registered re-auth callback (if any) is invoked once and the request is retried, transparently recovering from an expired - session on a long-lived client. + session on a long-lived client. If the callback raises (re-auth itself + failed), the error propagates to the caller and the request is not + retried; a subsequent request will attempt re-auth again. """ url = self._build_url(endpoint) prepared_params = self._prepare_params(params) diff --git a/rohlik_api/models.py b/rohlik_api/models.py index 28e296c..d01abb9 100644 --- a/rohlik_api/models.py +++ b/rohlik_api/models.py @@ -5,6 +5,11 @@ API response. All models are plain dataclasses, so ``dataclasses.asdict`` can be used to convert them back to JSON-serialisable dictionaries (useful for the Home Assistant integration and the MCP server). + +Monetary amounts come in two shapes: ``price`` fields typed as ``str`` are +pre-formatted for display (for example ``"29.90 Kč"``), while numeric ``price`` +fields are raw amounts paired with a separate ``currency``. Czech crowns (CZK, +"Kč") are the usual currency. """ from __future__ import annotations @@ -21,7 +26,19 @@ @dataclass(slots=True) class CartItem: - """A single item in the shopping cart.""" + """A single item (line) in the shopping cart. + + Attributes: + id: The product ID, as a string. + cart_item_id: The cart-line identifier (``orderFieldId``). Pass this to + :meth:`~rohlik_api.RohlikAPI.cart`'s ``delete_item`` to remove the + line from the cart. + name: Product name. + quantity: Number of units of this product in the cart. + price: Line price for this product, in the account currency (CZK). + category_name: Primary category name of the product. + brand: Brand name, or an empty string if unknown. + """ id: str cart_item_id: str @@ -47,7 +64,16 @@ def from_api(cls, item_id: str, data: dict[str, Any]) -> CartItem: @dataclass(slots=True) class Cart: - """The current shopping cart.""" + """The current shopping cart. + + Attributes: + total_price: Total price of the cart, in the account currency (CZK). + total_items: Number of distinct products in the cart (line count, not + the summed quantity). + can_make_order: Whether the cart currently satisfies the conditions to + place an order (e.g. the minimum order value is met). + products: The cart's line items. + """ total_price: float total_items: int @@ -74,7 +100,17 @@ def from_api(cls, payload: dict[str, Any]) -> Cart: @dataclass(slots=True) class ProductSearchResult: - """A product entry from a search response.""" + """A single product entry from a search response. + + Attributes: + id: Product ID. Use it with ``cart.add_items`` or the + ``products.get_*`` lookups. + name: Product name. + price: Pre-formatted price string, e.g. ``"29.90 Kč"`` (empty if the + API omitted price information). + brand: Brand name, if known. + amount: Textual packaging/amount, e.g. ``"500 g"``. + """ id: int | None name: str | None @@ -96,14 +132,26 @@ def from_api(cls, data: dict[str, Any]) -> ProductSearchResult: @dataclass(slots=True) class SearchResults: - """Container for product search results.""" + """Container for product search results. + + Attributes: + results: The matched products, in ranked order. Empty if nothing + matched the search term. + """ results: list[ProductSearchResult] = field(default_factory=list) @dataclass(slots=True) class AISummary: - """AI-generated summary for a product.""" + """AI-generated summary for a product. + + Attributes: + product_id: The product the summary is for. + rating: Rohlik's rating bucket for the summary (e.g. ``"EMPTY"``). + title: Summary title (localised, e.g. ``"AI Souhrn"``). + content: The summary text. + """ product_id: int | None rating: str | None @@ -123,7 +171,23 @@ def from_api(cls, data: dict[str, Any]) -> AISummary: @dataclass(slots=True) class NutritionalValue: - """Nutritional values for a single portion.""" + """Nutritional values for a single portion. + + Every amount is expressed for the stated :attr:`portion`. Any value the API + omits is ``None``. + + Attributes: + portion: The reference portion these values describe, e.g. ``"100 g"``. + energy_kj: Energy in kilojoules (kJ). + energy_kcal: Energy in kilocalories (kcal). + fats: Total fat, in grams. + saturated_fats: Saturated fat, in grams. + carbohydrates: Carbohydrates, in grams. + sugars: Sugars, in grams. + protein: Protein, in grams. + salt: Salt, in grams. + fiber: Fibre, in grams. + """ portion: str | None energy_kj: float | None @@ -161,7 +225,13 @@ def amount(key: str) -> float | None: @dataclass(slots=True) class Allergens: - """Allergen information for a product.""" + """Allergen information for a product. + + Attributes: + contained: Allergens the product definitely contains. + possibly_contained: Allergens that may be present (e.g. traces from + shared production lines). + """ contained: list[str] = field(default_factory=list) possibly_contained: list[str] = field(default_factory=list) @@ -169,7 +239,15 @@ class Allergens: @dataclass(slots=True) class ProductComposition: - """Composition and nutritional information for a product.""" + """Composition and nutritional information for a product. + + Attributes: + product_id: The product this composition is for. + nutritional_values: Nutrition broken down by portion (one entry per + portion size the API provides). + ingredients: Plain-text ingredient list, if available. + allergens: Allergen information. + """ product_id: int | None nutritional_values: list[NutritionalValue] = field(default_factory=list) @@ -195,7 +273,16 @@ def from_api(cls, data: dict[str, Any]) -> ProductComposition: @dataclass(slots=True) class ProductPrice: - """Current price information for a product.""" + """Current price information for a product. + + Attributes: + product_id: The product this price is for. + price: Current price as a number, expressed in :attr:`currency`. + currency: ISO currency code, e.g. ``"CZK"``. + price_per_unit: Price per base unit (e.g. per kg or per litre), in + :attr:`currency`. + sales: Raw list of active sales/discounts, in the API's own shape. + """ product_id: int | None price: float | None @@ -223,7 +310,17 @@ def from_api(cls, data: dict[str, Any]) -> ProductPrice: @dataclass(slots=True) class RecipeSummary: - """A recipe entry from a recipe search response.""" + """A recipe entry from a recipe search response. + + Attributes: + id: Recipe ID. Use it with ``recipes.get_detail``. + name: Recipe name. + link: Relative web path to the recipe on rohlik.cz. + image: Relative path to the recipe image. + is_favorite: Whether the recipe is marked as a favourite by the user. + is_new: Whether the recipe is flagged as new. + is_best_seller: Whether the recipe is flagged as a best seller. + """ id: int | None name: str | None @@ -249,7 +346,13 @@ def from_api(cls, data: dict[str, Any]) -> RecipeSummary: @dataclass(slots=True) class RecipeSearchResults: - """Container for recipe search results.""" + """Container for recipe search results. + + Attributes: + recipes: The matched recipes for this page of results. + total_hits: Total number of matching recipes (may exceed + ``len(recipes)`` because results are paginated). + """ recipes: list[RecipeSummary] = field(default_factory=list) total_hits: int = 0 @@ -266,7 +369,17 @@ def from_api(cls, payload: dict[str, Any]) -> RecipeSearchResults: @dataclass(slots=True) class IngredientItem: - """A single ingredient within a recipe ingredient group.""" + """A single ingredient within a recipe ingredient group. + + Attributes: + name: Ingredient display name. + ingredient_id: Ingredient ID. Pass it to ``recipes.get_ingredient_products`` + to find purchasable products for this ingredient. + ingredient_name: Textual amount and name, e.g. ``"2 vΔ›tΕ‘Γ­ mrkve"``. + products_count: Number of purchasable products available for this + ingredient. + image: Relative path to the ingredient image. + """ name: str | None ingredient_id: int | None @@ -288,7 +401,13 @@ def from_api(cls, data: dict[str, Any]) -> IngredientItem: @dataclass(slots=True) class IngredientGroup: - """A named group of recipe ingredients.""" + """A named group of recipe ingredients. + + Attributes: + name: Group name, e.g. ``"HOVĚZÍ VÝVAR"``. + position: Ordering index of the group within the recipe. + items: The ingredients belonging to this group. + """ name: str | None position: int | None @@ -306,7 +425,12 @@ def from_api(cls, data: dict[str, Any]) -> IngredientGroup: @dataclass(slots=True) class DirectionStep: - """A single step in a recipe direction section.""" + """A single step in a recipe direction section. + + Attributes: + step_number: 1-based step number within its section. + content: The step's instruction text. + """ step_number: int | None content: str | None @@ -319,7 +443,13 @@ def from_api(cls, data: dict[str, Any]) -> DirectionStep: @dataclass(slots=True) class DirectionSection: - """A named section of recipe directions.""" + """A named section of recipe directions. + + Attributes: + name: Section name, e.g. ``"POSTUP"``. + position: Ordering index of the section within the recipe. + steps: The ordered steps in this section. + """ name: str | None position: int | None @@ -337,7 +467,12 @@ def from_api(cls, data: dict[str, Any]) -> DirectionSection: @dataclass(slots=True) class RecipeAuthor: - """Author of a recipe.""" + """Author of a recipe. + + Attributes: + name: Author name. + annotation: Short note or bio about the author. + """ name: str | None annotation: str | None @@ -345,11 +480,27 @@ class RecipeAuthor: @dataclass(slots=True) class RecipeDetail: - """Detailed information about a recipe.""" + """Detailed information about a recipe. + + Attributes: + id: Recipe ID. + name: Recipe name. + duration: Human-readable preparation time, e.g. ``"Do hodinky"`` + ("within an hour"). This is a display string, not a number. + servings: Raw list of serving options, in the API's own shape (each + entry typically has ``name`` and ``default``). + image: Relative path to the recipe image. + author: The recipe's author. + tips: Free-text tips for preparing the recipe. + ingredients: Ingredients, grouped into named sections. + directions: Cooking directions, grouped into named sections of steps. + is_favorite: Whether the recipe is marked as a favourite by the user. + link: Relative web path to the recipe on rohlik.cz. + """ id: int | None name: str | None - duration: int | None + duration: str | None servings: list[Any] = field(default_factory=list) image: str | None = None author: RecipeAuthor = field(default_factory=lambda: RecipeAuthor(None, None)) @@ -383,7 +534,21 @@ def from_api(cls, payload: dict[str, Any]) -> RecipeDetail: @dataclass(slots=True) class IngredientProduct: - """A purchasable product matched to a recipe ingredient.""" + """A purchasable product matched to a recipe ingredient. + + Attributes: + product_id: Product ID. Use it with ``cart.add_items``. + name: Product name. + image: Relative path to the product image. + price: Pre-formatted price string, e.g. ``"41.88 Kč"``. + price_value: The raw numeric price (the ``full`` amount) behind + :attr:`price`. + unit: Base unit the product is sold in, e.g. ``"kg"`` or ``"ks"`` + (pieces). + amount: Textual amount, e.g. ``"cca 1,2 kg"``. + in_stock: Whether the product is currently in stock. + is_favorite: Whether the product is marked as a favourite by the user. + """ product_id: int | None name: str | None @@ -414,7 +579,14 @@ def from_api(cls, data: dict[str, Any]) -> IngredientProduct: @dataclass(slots=True) class IngredientProductGroup: - """Products available for a single ingredient.""" + """Products available for a single recipe ingredient. + + Attributes: + ingredient_id: The ingredient these products are matched to. + products: The purchasable products for this ingredient (this page). + total_hits: Total number of products available for the ingredient (may + exceed ``len(products)`` because results are paginated). + """ ingredient_id: int | None products: list[IngredientProduct] = field(default_factory=list) @@ -432,7 +604,11 @@ def from_api(cls, data: dict[str, Any]) -> IngredientProductGroup: @dataclass(slots=True) class IngredientProducts: - """Container for ingredient product groups.""" + """Container mapping recipe ingredients to purchasable products. + + Attributes: + ingredients: One group per requested ingredient ID. + """ ingredients: list[IngredientProductGroup] = field(default_factory=list) @@ -454,7 +630,13 @@ def from_api(cls, payload: dict[str, Any]) -> IngredientProducts: @dataclass(slots=True) class ShoppingList: - """A saved shopping list.""" + """A saved shopping list. + + Attributes: + name: The shopping list's name. + products_in_list: Raw list of product entries on the list, in the API's + own shape (each entry typically has ``productId`` and ``quantity``). + """ name: str | None products_in_list: list[Any] = field(default_factory=list)