From a17727d7285711f24fc9974ef6b883a91e39321b Mon Sep 17 00:00:00 2001 From: Georgie Kennedy Date: Fri, 11 Sep 2026 13:06:19 +1000 Subject: [PATCH] overall cleanup --- README.md | 2 +- src/pbs_client/cli/main.py | 115 ++++++++++++++++++--- src/pbs_client/config.py | 11 +- src/pbs_client/db/schema.py | 6 ++ src/pbs_client/db/state.py | 4 +- src/pbs_client/errors.py | 25 +++++ src/pbs_client/http/__init__.py | 10 +- src/pbs_client/http/client.py | 108 +++++++++++++++++--- src/pbs_client/sync/orchestrator.py | 36 +++++-- src/pbs_client/toolkit/core/service.py | 31 ++++-- tests/test_config.py | 10 +- tests/test_http.py | 103 ++++++++++++++++++- tests/test_query.py | 37 +++++++ tests/test_sync.py | 133 ++++++++++++++++++++++--- 14 files changed, 562 insertions(+), 69 deletions(-) diff --git a/README.md b/README.md index 0076f38..05d902c 100644 --- a/README.md +++ b/README.md @@ -35,7 +35,7 @@ subscription_key = "your-subscription-key" Set `OA_CONFIG_PATH` before invoking the commands when using a config file outside `~/.config/omop/config.toml`. -The public API is deliberately rate limited to one request per twenty seconds. The client enforces that interval process-wide, including retries and page continuations. Tests use local fixtures and never call the API. +The public API is deliberately rate limited to one request per twenty seconds. The client enforces that interval process-wide, including retries and page continuations, and rejects configured intervals below three seconds because the quota is shared across users. The default page size is 5,000 records: this keeps large responses manageable without creating unnecessary calls against the shared quota. Use `--limit 1000` when an endpoint still returns an empty or non-JSON response; if a page-size change is made during a resume, the affected resource safely restarts from page one. Refreshes are upserts and intentionally retain rows no longer returned by a later response, preserving local PBS history. Tests use local fixtures and never call the API. All configuration — the subscription key, base URL, rate limit, and the shared mirror database — is read from `oa-configurator`. There is no diff --git a/src/pbs_client/cli/main.py b/src/pbs_client/cli/main.py index e6c1fac..b69b2e2 100644 --- a/src/pbs_client/cli/main.py +++ b/src/pbs_client/cli/main.py @@ -2,18 +2,86 @@ from __future__ import annotations +from collections import Counter +from typing import Any + import typer from sqlalchemy.engine import Engine from pbs_client.config import PBSSettings, get_pbs_context from pbs_client.db import init_db, make_session_factory -from pbs_client.errors import PBSSyncError, PBSTransportError -from pbs_client.http import PBSClient +from pbs_client.errors import PBSHTTPError, PBSInvalidResponseError, PBSSyncError, PBSTransportError +from pbs_client.http import DEFAULT_PAGE_SIZE, PBSClient from pbs_client.sync import SyncOrchestrator, mirror_status app = typer.Typer(help="Maintain a local offline mirror of the PBS Public Data API v3.") +def _status_lines(rows: list[dict[str, Any]]) -> list[str]: + """Render a compact operational summary followed by grouped resource details.""" + + counts = Counter(row["status"] for row in rows) + complete = counts.get("complete", 0) + total_rows = sum(int(row["rows"]) for row in rows) + state_summary = " · ".join( + f"{counts.get(status, 0)} {label.lower()}" + for status, label in ( + ("complete", "complete"), + ("in_progress", "in progress"), + ("failed", "failed"), + ("pending", "pending"), + ) + if counts.get(status, 0) + ) + lines = [ + f"PBS mirror: {complete}/{len(rows)} resources complete · {total_rows:,} rows", + f"State: {state_summary}", + ] + attention = [ + row["resource"] + for row in rows + if row["status"] in {"failed", "in_progress"} + ] + if attention: + lines.append(f"Attention: {', '.join(attention)}") + + headings = { + "failed": "Failed", + "in_progress": "In progress", + "pending": "Pending", + "complete": "Complete", + } + for status in ("failed", "in_progress", "pending", "complete"): + group = [row for row in rows if row["status"] == status] + if not group: + continue + lines.extend(["", f"{headings[status]} ({len(group)})"]) + for row in group: + if status == "complete": + detail = ( + f"{int(row['rows']):,} rows · page {row['page']} · " + f"completed {_format_status_time(row['completed_at'])}" + ) + elif status == "in_progress": + detail = ( + f"{int(row['rows']):,} rows · page {row['page']} · " + f"started {_format_status_time(row['started_at'])}" + ) + elif status == "failed": + error = str(row["last_error"] or "unknown error").splitlines()[0] + detail = f"{int(row['rows']):,} rows · page {row['page']} · {error}" + else: + detail = "not started" + lines.append(f" {row['resource']:<30} {detail}") + return lines + + +def _format_status_time(value: Any) -> str: + if value is None: + return "-" + return value.strftime("%Y-%m-%d %H:%M") + + def _runtime() -> tuple[PBSSettings, Engine, str]: """Resolve the shared oa-configurator config, engine, and database name.""" @@ -21,12 +89,14 @@ def _runtime() -> tuple[PBSSettings, Engine, str]: return PBSSettings.from_config(config), database.create_engine(future=True), config.pbs_db -def _find_transport_error(error: BaseException) -> PBSTransportError | None: - """Find a transport failure retained in a wrapped sync exception.""" +def _find_cause[Error: BaseException]( + error: BaseException, expected: type[Error] +) -> Error | None: + """Find a cause of the requested type retained in a wrapped exception.""" current: BaseException | None = error while current is not None: - if isinstance(current, PBSTransportError): + if isinstance(current, expected): return current current = current.__cause__ return None @@ -36,7 +106,9 @@ def _report_sync_failure(error: PBSSyncError) -> None: """Print an actionable, traceback-free message for a failed CLI sync.""" resource = error.resource or "the current resource" - transport = _find_transport_error(error) + transport = _find_cause(error, PBSTransportError) + http_error = _find_cause(error, PBSHTTPError) + invalid_response = _find_cause(error, PBSInvalidResponseError) typer.echo(f"PBS sync paused on {resource}.", err=True) if transport is not None and transport.timed_out: typer.echo( @@ -50,6 +122,27 @@ def _report_sync_failure(error: PBSSyncError) -> None: err=True, ) limit = " --limit 1000" + elif http_error is not None and http_error.status_code == 429: + typer.echo(str(error), err=True) + typer.echo( + "The PBS API rate limit was reached. Wait before retrying; committed " + "pages and completed resources are already saved.", + err=True, + ) + limit = "" + elif invalid_response is not None: + typer.echo(str(error), err=True) + typer.echo( + "The PBS API returned an empty or non-JSON response; this can " + "happen when a page is too large or the service returns an error page.", + err=True, + ) + typer.echo( + "Completed resources and committed pages are already saved. " + "Retry with a smaller page size.", + err=True, + ) + limit = " --limit 1000" else: typer.echo(str(error), err=True) typer.echo("Completed resources are already saved and can be resumed.", err=True) @@ -76,7 +169,7 @@ def sync_command( refresh: bool = typer.Option( True, "--refresh/--resume-only", help="Refresh completed resources." ), - limit: int = typer.Option(100_000, min=1, help="API page size."), + limit: int = typer.Option(DEFAULT_PAGE_SIZE, min=1, help="API page size."), ) -> None: """Synchronize all PBS resources, or one resource, into the local mirror.""" @@ -106,12 +199,8 @@ def status() -> None: _, engine, _ = _runtime() init_db(engine) with make_session_factory(engine)() as session: - for row in mirror_status(session): - completed = row["completed_at"].isoformat() if row["completed_at"] else "-" - typer.echo( - f"{row['resource']}: {row['status']}; rows={row['rows']}; " - f"page={row['page']}; completed={completed}" - ) + for line in _status_lines(mirror_status(session)): + typer.echo(line) def main() -> None: diff --git a/src/pbs_client/config.py b/src/pbs_client/config.py index 28ebde8..f8fe475 100644 --- a/src/pbs_client/config.py +++ b/src/pbs_client/config.py @@ -3,6 +3,7 @@ from __future__ import annotations from dataclasses import dataclass +from math import isfinite from typing import Annotated, ClassVar from oa_configurator import ( @@ -19,6 +20,7 @@ DEFAULT_BASE_URL = "https://data-api.health.gov.au/pbs/api/v3" DEFAULT_PUBLIC_KEY = "2384af7c667342ceb5a736fe29f1dc6b" DEFAULT_RATE_LIMIT_SECONDS = 20.0 +MIN_RATE_LIMIT_SECONDS = 3.0 class PBSClientConfig(PackageConfigBase): @@ -45,7 +47,7 @@ class PBSClientConfig(PackageConfigBase): ) rate_limit_seconds: float = Field( default=DEFAULT_RATE_LIMIT_SECONDS, - gt=0, + ge=MIN_RATE_LIMIT_SECONDS, description="Minimum delay between PBS API requests in seconds.", ) @@ -80,6 +82,13 @@ class PBSSettings: base_url: str = DEFAULT_BASE_URL rate_limit_seconds: float = DEFAULT_RATE_LIMIT_SECONDS + def __post_init__(self) -> None: + if not isfinite(self.rate_limit_seconds) or self.rate_limit_seconds < MIN_RATE_LIMIT_SECONDS: + raise ValueError( + "rate_limit_seconds must be at least " + f"{MIN_RATE_LIMIT_SECONDS:g} seconds to respect the PBS API quota" + ) + @classmethod def from_config(cls, config: PBSClientConfig) -> PBSSettings: """Build HTTP settings from the resolved package configuration.""" diff --git a/src/pbs_client/db/schema.py b/src/pbs_client/db/schema.py index e6d65d5..e3d5dc9 100644 --- a/src/pbs_client/db/schema.py +++ b/src/pbs_client/db/schema.py @@ -189,3 +189,9 @@ def primary_key(self) -> tuple[str, ...]: "SummaryOfChanges", "ApiChangelog", ) + +if len(SYNC_ORDER) != len(set(SYNC_ORDER)): + raise RuntimeError("SYNC_ORDER contains duplicate resources") +unknown = sorted(set(SYNC_ORDER) - set(RESOURCE_BY_NAME)) +if unknown: + raise RuntimeError(f"SYNC_ORDER contains unregistered resources: {unknown}") diff --git a/src/pbs_client/db/state.py b/src/pbs_client/db/state.py index b5b1254..1e71d25 100644 --- a/src/pbs_client/db/state.py +++ b/src/pbs_client/db/state.py @@ -26,10 +26,12 @@ class SyncState(Base): last_error: Mapped[str | None] = mapped_column(Text) metadata_json: Mapped[dict[str, Any]] = mapped_column(JSON, nullable=False, default=dict) - def begin(self) -> None: + def begin(self, *, page_limit: int) -> None: self.status = "in_progress" self.started_at = datetime.now(UTC) + self.completed_at = None self.last_error = None + self.metadata_json = {**self.metadata_json, "page_limit": page_limit} def checkpoint(self, page: int, count: int, metadata: dict[str, Any]) -> None: self.page = page diff --git a/src/pbs_client/errors.py b/src/pbs_client/errors.py index 088b91a..c4a644d 100644 --- a/src/pbs_client/errors.py +++ b/src/pbs_client/errors.py @@ -15,6 +15,31 @@ class PBSAPIError(PBSClientError): """The PBS API returned an error or an unusable response.""" +class PBSHTTPError(PBSAPIError): + """The PBS API returned an HTTP error response.""" + + def __init__( + self, + url: str, + status_code: int, + attempts: int, + *, + retryable: bool, + retry_after_seconds: float | None = None, + ) -> None: + self.url = url + self.status_code = status_code + self.attempts = attempts + self.retryable = retryable + self.retry_after_seconds = retry_after_seconds + suffix = f" after {attempts} attempts" if retryable else "" + super().__init__(f"PBS API returned HTTP {status_code}{suffix}: {url}") + + +class PBSInvalidResponseError(PBSAPIError): + """The PBS API response could not be decoded as the expected document.""" + + class PBSTransportError(PBSAPIError): """The PBS API could not be reached after all transport retries.""" diff --git a/src/pbs_client/http/__init__.py b/src/pbs_client/http/__init__.py index ac44ed5..3cd7e2f 100644 --- a/src/pbs_client/http/__init__.py +++ b/src/pbs_client/http/__init__.py @@ -1,5 +1,11 @@ """PBS API v3 HTTP client.""" -from pbs_client.http.client import GlobalRateLimiter, Page, PBSClient, TransportResponse +from pbs_client.http.client import ( + DEFAULT_PAGE_SIZE, + GlobalRateLimiter, + Page, + PBSClient, + TransportResponse, +) -__all__ = ["GlobalRateLimiter", "PBSClient", "Page", "TransportResponse"] +__all__ = ["DEFAULT_PAGE_SIZE", "GlobalRateLimiter", "PBSClient", "Page", "TransportResponse"] diff --git a/src/pbs_client/http/client.py b/src/pbs_client/http/client.py index 92b4b3e..3aabcd9 100644 --- a/src/pbs_client/http/client.py +++ b/src/pbs_client/http/client.py @@ -8,16 +8,18 @@ import time from collections.abc import Callable, Iterator, Mapping from dataclasses import dataclass, field -from typing import Any, Protocol +from typing import Any, Protocol, cast from urllib.error import HTTPError, URLError from urllib.parse import urlencode, urljoin from urllib.request import Request, urlopen -from pbs_client.config import PBSSettings -from pbs_client.errors import PBSAPIError, PBSTransportError +from pbs_client.config import DEFAULT_RATE_LIMIT_SECONDS, PBSSettings +from pbs_client.errors import PBSAPIError, PBSHTTPError, PBSInvalidResponseError, PBSTransportError logger = logging.getLogger(__name__) +DEFAULT_PAGE_SIZE = 5_000 + class Sleeper(Protocol): def __call__(self, seconds: float, /) -> None: ... @@ -39,13 +41,48 @@ def _urlopen_transport(method: str, url: str, headers: Mapping[str, str]) -> Tra return TransportResponse(response.status, dict(response.headers), response.read()) +def _header_value(headers: Mapping[str, str], name: str) -> str | None: + """Read a response header without depending on its casing.""" + + return next((value for key, value in headers.items() if key.lower() == name.lower()), None) + + +def _response_diagnostic(response: TransportResponse) -> str: + """Return bounded response details suitable for an error message.""" + + content_type = _header_value(response.headers, "Content-Type") or "" + body = response.body.strip() + if not body: + preview = "" + else: + preview = " ".join(body[:200].decode("utf-8", errors="replace").split()) + if len(body) > 200: + preview += "..." + return ( + f"HTTP {response.status_code}; content-type={content_type!r}; " + f"body_bytes={len(response.body)}; body_preview={preview!r}" + ) + + +def _retry_after_seconds(value: str | None) -> float | None: + """Parse the numeric form of ``Retry-After`` without guessing dates.""" + + if value is None: + return None + try: + seconds = float(value) + except ValueError: + return None + return seconds if seconds >= 0 else None + + class GlobalRateLimiter: """One process-wide monotonic gate shared by every PBS client instance.""" _lock = threading.Lock() _next_allowed = 0.0 - def __init__(self, interval: float = 20.0, sleeper: Sleeper = time.sleep) -> None: + def __init__(self, interval: float = DEFAULT_RATE_LIMIT_SECONDS, sleeper: Sleeper = time.sleep) -> None: self.interval = max(0.0, interval) self.sleeper = sleeper @@ -122,11 +159,13 @@ def fetch_page( endpoint: str, *, page: int = 1, - limit: int = 100_000, + limit: int = DEFAULT_PAGE_SIZE, params: Mapping[str, Any] | None = None, ) -> Page: """Fetch and decode one page without adding a schedule filter.""" + if limit < 1: + raise ValueError("limit must be at least 1") query: dict[str, Any] = dict(params or {}) query.update(page=page, limit=limit) query_string = urlencode(query, doseq=True) @@ -137,22 +176,30 @@ def fetch_page( try: document = json.loads(response.body.decode("utf-8")) except (UnicodeDecodeError, json.JSONDecodeError) as exc: - raise PBSAPIError(f"PBS endpoint {endpoint} returned invalid JSON") from exc + raise PBSInvalidResponseError( + f"PBS endpoint {endpoint} returned invalid JSON ({_response_diagnostic(response)})" + ) from exc if not isinstance(document, dict) or not isinstance(document.get("data"), list): - raise PBSAPIError(f"PBS endpoint {endpoint} returned no JSON data array") + raise PBSInvalidResponseError( + f"PBS endpoint {endpoint} returned no JSON data array " + f"({_response_diagnostic(response)})" + ) metadata = document.get("_meta", {}) links = document.get("_links", []) if isinstance(links, dict): links = [links] if not isinstance(metadata, dict) or not isinstance(links, list): - raise PBSAPIError(f"PBS endpoint {endpoint} returned malformed pagination metadata") + raise PBSInvalidResponseError( + f"PBS endpoint {endpoint} returned malformed pagination metadata " + f"({_response_diagnostic(response)})" + ) return Page(endpoint, page, limit, document["data"], metadata, links) def iter_pages( self, endpoint: str, *, - limit: int = 100_000, + limit: int = DEFAULT_PAGE_SIZE, params: Mapping[str, Any] | None = None, start_page: int = 1, ) -> Iterator[Page]: @@ -191,9 +238,17 @@ def _request(self, url: str) -> TransportResponse: try: response = self.transport("GET", url, headers) if response.status_code == 429 or response.status_code >= 500: - raise _RetryableStatus(response.status_code) + raise _RetryableStatus( + response.status_code, + _retry_after_seconds(_header_value(response.headers, "Retry-After")), + ) if response.status_code >= 400: - raise PBSAPIError(f"PBS API returned HTTP {response.status_code} for {url}") + raise PBSHTTPError( + url, + response.status_code, + attempt + 1, + retryable=False, + ) return response except _RetryableStatus as exc: last_error = exc @@ -205,8 +260,21 @@ def _request(self, url: str) -> TransportResponse: ) except HTTPError as exc: if exc.code < 500 and exc.code != 429: - raise PBSAPIError(f"PBS API returned HTTP {exc.code} for {url}") from exc - last_error = exc + raise PBSHTTPError( + url, + exc.code, + attempt + 1, + retryable=False, + ) from exc + last_error = _RetryableStatus( + exc.code, + _retry_after_seconds( + _header_value( + cast(Mapping[str, str], exc.headers or {}), + "Retry-After", + ) + ), + ) except (URLError, TimeoutError, OSError) as exc: last_error = exc if attempt < self.max_retries: @@ -214,7 +282,18 @@ def _request(self, url: str) -> TransportResponse: else: logger.warning("PBS API request exhausted transport retries: %s", exc) if attempt < self.max_retries: - self.limiter.sleeper(self.backoff_base * (2**attempt)) + delay = self.backoff_base * (2**attempt) + if isinstance(last_error, _RetryableStatus): + delay = max(delay, last_error.retry_after_seconds or 0.0) + self.limiter.sleeper(delay) + if isinstance(last_error, _RetryableStatus): + raise PBSHTTPError( + url, + last_error.status_code, + self.max_retries + 1, + retryable=True, + retry_after_seconds=last_error.retry_after_seconds, + ) from last_error if last_error is not None: raise PBSTransportError(url, self.max_retries + 1, last_error) from last_error raise PBSAPIError(f"PBS API request failed after retries: {url}") @@ -223,3 +302,4 @@ def _request(self, url: str) -> TransportResponse: @dataclass(frozen=True, slots=True) class _RetryableStatus(Exception): status_code: int + retry_after_seconds: float | None = None diff --git a/src/pbs_client/sync/orchestrator.py b/src/pbs_client/sync/orchestrator.py index 5c30676..524b07e 100644 --- a/src/pbs_client/sync/orchestrator.py +++ b/src/pbs_client/sync/orchestrator.py @@ -14,7 +14,7 @@ from pbs_client.db.engine import fk_checks_disabled_for_refresh from pbs_client.db.model import Base from pbs_client.errors import PBSSyncError -from pbs_client.http import PBSClient +from pbs_client.http import DEFAULT_PAGE_SIZE, PBSClient logger = logging.getLogger(__name__) @@ -30,10 +30,9 @@ class SyncResult: def upsert_records(session: Session, model: type[Base], records: Iterable[dict[str, Any]]) -> int: """Insert or update records using the model's declared natural key.""" - fields = tuple( - column.name for column in model.__table__.columns if column.name != "raw_payload" - ) - primary_key = tuple(column.name for column in model.__table__.primary_key) + spec = next(spec for spec in RESOURCE_SPECS if spec.model is model) + fields = spec.fields + primary_key = spec.primary_key written = 0 for record in records: if not isinstance(record, dict): @@ -67,7 +66,7 @@ def run( self, *, resource: str | None = None, - limit: int = 100_000, + limit: int = DEFAULT_PAGE_SIZE, refresh_completed: bool = True, ) -> list[SyncResult]: """Sync all resources, or one named resource, and return page summaries. @@ -75,9 +74,12 @@ def run( An incomplete resource resumes at the page after its last committed checkpoint. A completed resource is refreshed from page one by default so newly published schedules are discovered; upsert semantics - keep that refresh idempotent. + keep that refresh idempotent and intentionally retain rows missing from + later responses as local history. """ + if limit < 1: + raise PBSSyncError("sync page limit must be at least 1") names = self._resource_names(resource) probe = self.session_factory() try: @@ -102,11 +104,24 @@ def run( def _sync_resource(self, name: str, state: SyncState, *, limit: int) -> SyncResult: spec = RESOURCE_BY_NAME[name] model = MODEL_BY_NAME[name] - start_page = state.page + 1 if state.status in {"in_progress", "failed"} else 1 + stored_limit = state.metadata_json.get("page_limit") + can_resume = ( + state.status in {"in_progress", "failed"} + and state.page > 0 + and stored_limit == limit + ) + start_page = state.page + 1 if can_resume else 1 + if state.status in {"in_progress", "failed"} and state.page > 0 and not can_resume: + logger.info( + "Restarting %s from page one because the saved page size is %r, not %s", + name, + stored_limit, + limit, + ) if start_page == 1: state.page = 0 state.records_written = 0 - state.begin() + state.begin(page_limit=limit) self._save_state(state) pages = 0 try: @@ -122,6 +137,7 @@ def _sync_resource(self, name: str, state: SyncState, *, limit: int) -> SyncResu "messages": page.messages, "links": page.links, "synced_at": page.metadata.get("synced_at"), + "page_limit": limit, } current.checkpoint(page.page, count, metadata) session.commit() @@ -204,7 +220,9 @@ def mirror_status(session: Session) -> list[dict[str, Any]]: "rows": count, "status": state.status if state else "pending", "page": state.page if state else 0, + "started_at": state.started_at if state else None, "completed_at": state.completed_at if state else None, + "last_error": state.last_error if state else None, } ) return rows diff --git a/src/pbs_client/toolkit/core/service.py b/src/pbs_client/toolkit/core/service.py index 649eb33..f5a172e 100644 --- a/src/pbs_client/toolkit/core/service.py +++ b/src/pbs_client/toolkit/core/service.py @@ -42,7 +42,7 @@ class IndicationText: schedule_code: int res_code: str prescribing_txt_id: int | None - benefit_type_code: BenefitTypeCode + benefit_type_code: BenefitTypeCode | str episodicity: str | None = None severity: str | None = None @@ -81,6 +81,7 @@ def _clean_html(value: str | None) -> str | None: return None parser = _HTMLTextExtractor() parser.feed(value) + parser.close() text = " ".join("".join(parser.parts).split()) return text or None @@ -103,15 +104,33 @@ def _as_date(value: str | date | datetime) -> date: raise ValueError(f"cannot parse PBS date: {value!r}") +def _benefit_type(value: str | None) -> BenefitTypeCode | str: + """Preserve an API benefit code even when a newer value is introduced.""" + + if value is None: + return "UNKNOWN" + try: + return BenefitTypeCode(value) + except ValueError: + return value + + def resolve_schedule(session: Session, as_of: date | datetime | str) -> Schedule | None: """Resolve the latest schedule effective on ``as_of`` by date.""" target = _as_date(as_of) schedules = session.scalars(select(Schedule)).all() - eligible = [schedule for schedule in schedules if _as_date(schedule.effective_date) <= target] + eligible = [ + schedule + for schedule in schedules + if schedule.effective_date is not None and _as_date(schedule.effective_date) <= target + ] if not eligible: return None - return max(eligible, key=lambda schedule: _as_date(schedule.effective_date)) + return max( + eligible, + key=lambda schedule: (_as_date(schedule.effective_date), schedule.schedule_code), + ) def find_items( @@ -209,7 +228,7 @@ def get_item_indication_text(session: Session, item: Item) -> list[IndicationTex restriction = session.get(RestrictionText, (item.schedule_code, link.res_code)) if restriction is None: continue - benefit_type = BenefitTypeCode(link.benefit_type_code) + benefit_type = _benefit_type(link.benefit_type_code) structured = _structured_indications(session, item, link, benefit_type) if structured: results.extend(structured) @@ -223,7 +242,7 @@ def _structured_indications( session: Session, item: Item, link: ItemRestrictionRltd, - benefit_type: BenefitTypeCode, + benefit_type: BenefitTypeCode | str, ) -> list[IndicationText]: text_links = session.scalars( select(RstrctnPrscrbngTxtRltd) @@ -266,7 +285,7 @@ def _fallback_indication( item: Item, link: ItemRestrictionRltd, restriction: RestrictionText, - benefit_type: BenefitTypeCode, + benefit_type: BenefitTypeCode | str, ) -> IndicationText | None: text = _clean_html(restriction.schedule_html_text) or _clean_html(restriction.li_html_text) if not text: diff --git a/tests/test_config.py b/tests/test_config.py index ed926c0..ded5398 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -7,7 +7,7 @@ StackConfig, ) -from pbs_client.config import PBSClientConfig, PBSSettings +from pbs_client.config import MIN_RATE_LIMIT_SECONDS, PBSClientConfig, PBSSettings def test_pbs_config_resolves_shared_generic_database(tmp_path: Path): @@ -25,7 +25,7 @@ def test_pbs_config_resolves_shared_generic_database(tmp_path: Path): "pbs_client": { "pbs_db": "pbs_db", "subscription_key": "configured-key", - "rate_limit_seconds": 1, + "rate_limit_seconds": MIN_RATE_LIMIT_SECONDS, } }, ) @@ -34,7 +34,7 @@ def test_pbs_config_resolves_shared_generic_database(tmp_path: Path): database = Resolver(stack).resolve_database(config.pbs_db) assert config.subscription_key == "configured-key" - assert config.rate_limit_seconds == 1 + assert config.rate_limit_seconds == MIN_RATE_LIMIT_SECONDS assert database.name == "pbs_db" assert database.connection.url == f"sqlite:///{tmp_path / 'pbs.db'}" @@ -43,11 +43,11 @@ def test_settings_can_be_built_from_package_config(): config = PBSClientConfig( subscription_key="configured-key", base_url="https://example.test/", - rate_limit_seconds=2, + rate_limit_seconds=MIN_RATE_LIMIT_SECONDS, ) settings = PBSSettings.from_config(config) assert settings.subscription_key == "configured-key" assert settings.base_url == "https://example.test" - assert settings.rate_limit_seconds == 2 + assert settings.rate_limit_seconds == MIN_RATE_LIMIT_SECONDS diff --git a/tests/test_http.py b/tests/test_http.py index c18ed93..03e52cb 100644 --- a/tests/test_http.py +++ b/tests/test_http.py @@ -1,13 +1,15 @@ from __future__ import annotations import json +from io import BytesIO +from urllib.error import HTTPError from urllib.parse import parse_qs, urlparse import pytest -from pbs_client.config import PBSSettings -from pbs_client.errors import PBSTransportError -from pbs_client.http import GlobalRateLimiter, PBSClient, TransportResponse +from pbs_client.config import MIN_RATE_LIMIT_SECONDS, PBSSettings +from pbs_client.errors import PBSAPIError, PBSHTTPError, PBSTransportError +from pbs_client.http import DEFAULT_PAGE_SIZE, GlobalRateLimiter, PBSClient, TransportResponse def test_pagination_and_headers_are_fixture_driven(fixture_dir): @@ -30,8 +32,13 @@ def transport(method, url, headers): GlobalRateLimiter.reset_for_tests() client = PBSClient( - PBSSettings(subscription_key="test", base_url="https://example.test", rate_limit_seconds=0), + PBSSettings( + subscription_key="test", + base_url="https://example.test", + rate_limit_seconds=MIN_RATE_LIMIT_SECONDS, + ), transport=transport, + limiter=GlobalRateLimiter(0), ) pages = list(client.iter_pages("/schedules", limit=1)) @@ -53,13 +60,99 @@ def transport(*args): assert called is False +def test_default_page_size_is_conservative(): + calls = [] + + def transport(method, url, headers): + calls.append(url) + return TransportResponse(200, {}, b'{"data": [], "_meta": {}, "_links": []}') + + client = PBSClient( + PBSSettings( + subscription_key="test", + base_url="https://example.test", + rate_limit_seconds=MIN_RATE_LIMIT_SECONDS, + ), + transport=transport, + limiter=GlobalRateLimiter(0), + ) + + client.fetch_page("/restrictions") + + assert f"limit={DEFAULT_PAGE_SIZE}" in calls[0] + + +def test_invalid_json_reports_bounded_response_diagnostic(): + def transport(*args): + return TransportResponse( + 200, + {"Content-Type": "text/html"}, + b"temporary upstream error", + ) + + client = PBSClient( + PBSSettings(subscription_key="test", rate_limit_seconds=MIN_RATE_LIMIT_SECONDS), + transport=transport, + limiter=GlobalRateLimiter(0), + ) + + with pytest.raises(PBSAPIError) as caught: + client.fetch_page("/restrictions") + + message = str(caught.value) + assert "HTTP 200" in message + assert "content-type='text/html'" in message + assert "body_bytes=37" in message + assert "temporary upstream error" in message + + +def test_http_error_preserves_status_and_honours_retry_after(): + calls = 0 + sleeps = [] + + def transport(*args): + nonlocal calls + calls += 1 + raise HTTPError( + "https://example.test/restrictions?page=1", + 429, + "too many requests", + {"Retry-After": "7"}, + BytesIO(b"rate limited"), + ) + + client = PBSClient( + PBSSettings(subscription_key="test", rate_limit_seconds=MIN_RATE_LIMIT_SECONDS), + transport=transport, + limiter=GlobalRateLimiter(0, sleeper=sleeps.append), + max_retries=1, + backoff_base=2, + ) + + with pytest.raises(PBSHTTPError) as caught: + client.fetch_page("/restrictions") + + assert calls == 2 + assert caught.value.status_code == 429 + assert caught.value.attempts == 2 + assert caught.value.retryable is True + assert caught.value.retry_after_seconds == 7 + assert sleeps == [7] + + +def test_rate_limit_cannot_be_configured_below_quota_floor(): + with pytest.raises(ValueError, match="at least 3 seconds"): + PBSSettings(subscription_key="test", rate_limit_seconds=2) + + def test_timeout_after_retries_is_actionable(): def transport(*args): raise TimeoutError("read timed out") client = PBSClient( - PBSSettings(subscription_key="test", rate_limit_seconds=0), + PBSSettings(subscription_key="test", rate_limit_seconds=MIN_RATE_LIMIT_SECONDS), transport=transport, + limiter=GlobalRateLimiter(0), sleeper=lambda _: None, max_retries=1, ) diff --git a/tests/test_query.py b/tests/test_query.py index 6c407e8..d53737c 100644 --- a/tests/test_query.py +++ b/tests/test_query.py @@ -203,3 +203,40 @@ def test_item_lookup_with_unknown_date_returns_no_items(session_factory): session.commit() assert find_items(session, "X4", as_of="2025-01-01") == [] + + +def test_schedule_resolution_tie_breaks_by_schedule_code(session_factory): + with session_factory() as session: + session.add_all( + [ + Schedule(schedule_code=10, effective_date="2026-04-01", effective_year=2026), + Schedule(schedule_code=11, effective_date="2026-04-01", effective_year=2026), + ] + ) + session.commit() + + assert resolve_schedule(session, "2026-05-01").schedule_code == 11 + + +def test_unknown_benefit_type_is_preserved(session_factory): + with session_factory() as session: + session.add(Schedule(schedule_code=12, effective_date="2026-05-01", effective_year=2026)) + session.commit() + session.add_all( + [ + Item(schedule_code=12, li_item_id="li-5", pbs_code="X5", drug_name="Drug"), + RestrictionText(schedule_code=12, res_code="R5", schedule_html_text="Use for condition"), + ItemRestrictionRltd( + schedule_code=12, + pbs_code="X5", + res_code="R5", + benefit_type_code="Z", + restriction_indicator="Y", + ), + ] + ) + session.commit() + + indications = get_item_indication_text(session, session.get(Item, (12, "li-5"))) + + assert indications[0].benefit_type_code == "Z" diff --git a/tests/test_sync.py b/tests/test_sync.py index ed6cb0a..17ed481 100644 --- a/tests/test_sync.py +++ b/tests/test_sync.py @@ -1,6 +1,8 @@ from __future__ import annotations -from pbs_client.cli.main import _report_sync_failure +from datetime import UTC, datetime + +from pbs_client.cli.main import _report_sync_failure, _status_lines from pbs_client.db import MODEL_BY_NAME, SyncState from pbs_client.errors import PBSSyncError, PBSTransportError from pbs_client.http import Page @@ -34,6 +36,14 @@ def iter_pages(self, endpoint, *, limit, start_page): raise RuntimeError("interrupted") +def report_wrapped_failure(error, resource, *, capsys): + try: + raise PBSSyncError("sync failed", resource=resource) from error + except PBSSyncError as sync_error: + _report_sync_failure(sync_error) + return capsys.readouterr().err + + def test_sync_is_idempotent_and_upserts(session_factory): page = Page("/schedules", 1, 100, [schedule(1, "2026-01-01")], {"total_records": 1}, []) client = FakeClient([page]) @@ -75,18 +85,117 @@ def test_sync_resumes_after_a_committed_page(session_factory): assert result[0].records_written == 2 -def test_cli_timeout_message_explains_how_to_resume(capsys): +def test_sync_restarts_when_page_size_changes(session_factory): + pages = [ + Page("/schedules", 1, 1, [schedule(1, "2026-01-01")], {"total_records": 2}, [{"rel": "next"}]), + Page("/schedules", 2, 1, [schedule(2, "2026-02-01")], {"total_records": 2}, []), + ] + interrupted = FakeClient(pages, fail=True) try: - raise PBSTransportError( - "https://example.test/restrictions?page=1", 4, TimeoutError("read timed out") - ) - except PBSTransportError as transport_error: - try: - raise PBSSyncError("sync failed", resource="RestrictionText") from transport_error - except PBSSyncError as sync_error: - _report_sync_failure(sync_error) - - output = capsys.readouterr().err + SyncOrchestrator(interrupted, session_factory).run(resource="Schedule", limit=1) + except PBSSyncError: + pass + + resumed = FakeClient(pages) + result = SyncOrchestrator(resumed, session_factory).run(resource="Schedule", limit=2) + + with session_factory() as session: + assert session.query(MODEL_BY_NAME["Schedule"]).count() == 2 + assert resumed.calls == [("/schedules", 1)] + assert result[0].records_written == 2 + + +def test_refresh_retains_rows_missing_from_the_latest_page(session_factory): + first = Page( + "/schedules", + 1, + 100, + [schedule(1, "2026-01-01"), schedule(2, "2026-02-01")], + {"total_records": 2}, + [], + ) + SyncOrchestrator(FakeClient([first]), session_factory).run(resource="Schedule", limit=100) + + second = Page("/schedules", 1, 100, [schedule(1, "2026-01-01")], {"total_records": 1}, []) + SyncOrchestrator(FakeClient([second]), session_factory).run(resource="Schedule", limit=100) + + with session_factory() as session: + assert session.query(MODEL_BY_NAME["Schedule"]).count() == 2 + + +def test_failed_refresh_clears_previous_completion_time(session_factory): + page = Page("/schedules", 1, 100, [schedule(1, "2026-01-01")], {"total_records": 1}, []) + SyncOrchestrator(FakeClient([page]), session_factory).run(resource="Schedule", limit=100) + + failing = FakeClient([page], fail=True) + try: + SyncOrchestrator(failing, session_factory).run(resource="Schedule", limit=100) + except PBSSyncError: + pass + + with session_factory() as session: + state = session.get(SyncState, "Schedule") + assert state.status == "failed" + assert state.completed_at is None + + +def test_cli_timeout_message_explains_how_to_resume(capsys): + transport_error = PBSTransportError( + "https://example.test/restrictions?page=1", 4, TimeoutError("read timed out") + ) + output = report_wrapped_failure(transport_error, "RestrictionText", capsys=capsys) assert "PBS sync paused on RestrictionText" in output assert "did not respond before the request timeout" in output assert "--resume-only --limit 1000" in output + + +def test_cli_invalid_json_message_explains_how_to_reduce_page_size(capsys): + from pbs_client.errors import PBSInvalidResponseError + + api_error = PBSInvalidResponseError( + "PBS endpoint /item-dispensing-rule-relationships returned invalid JSON" + ) + output = report_wrapped_failure(api_error, "ItemDispensingRuleRltd", capsys=capsys) + assert "empty or non-JSON response" in output + assert "--resume-only --limit 1000" in output + + +def test_status_output_summarises_and_groups_resources(): + lines = _status_lines( + [ + { + "resource": "Schedule", + "status": "complete", + "rows": 13, + "page": 1, + "started_at": None, + "completed_at": datetime(2026, 9, 10, 0, 26, tzinfo=UTC), + "last_error": None, + }, + { + "resource": "MarkupBand", + "status": "in_progress", + "rows": 0, + "page": 0, + "started_at": datetime(2026, 9, 11, 3, 5, tzinfo=UTC), + "completed_at": None, + "last_error": None, + }, + { + "resource": "Prescriber", + "status": "pending", + "rows": 0, + "page": 0, + "started_at": None, + "completed_at": None, + "last_error": None, + }, + ] + ) + + output = "\n".join(lines) + assert "PBS mirror: 1/3 resources complete · 13 rows" in output + assert "Attention: MarkupBand" in output + assert "In progress (1)" in output + assert "Complete (1)" in output + assert "Pending (1)" in output