From ed18c6fd5a7e917d59bfa635268bc560b04b4f93 Mon Sep 17 00:00:00 2001 From: Jonathan Giffard Date: Tue, 29 Sep 2026 19:46:38 +0100 Subject: [PATCH 1/8] Add core client: config, errors, transport, auth and request layers Phase 2 of PLAN.md. Adds AuraClient (options, from_env, context manager), the AuraError exception hierarchy, the HttpTransport protocol with an httpx implementation, network-only retries with a per-call deadline, OAuth token caching, and the authenticated request service. Co-Authored-By: Claude Opus 5.5 --- PLAN.md | 18 ++ pyproject.toml | 4 +- src/aura_python_sdk/__init__.py | 47 +++- src/aura_python_sdk/_client.py | 166 ++++++++++++ src/aura_python_sdk/_config.py | 126 +++++++++ src/aura_python_sdk/_errors.py | 244 ++++++++++++++++++ src/aura_python_sdk/_internal/__init__.py | 0 src/aura_python_sdk/_internal/_auth.py | 119 +++++++++ src/aura_python_sdk/_internal/_request.py | 124 +++++++++ .../_internal/http/__init__.py | 0 src/aura_python_sdk/_internal/http/_httpx.py | 79 ++++++ .../_internal/http/_service.py | 102 ++++++++ src/aura_python_sdk/_transport.py | 57 ++++ tests/fakes.py | 75 ++++++ tests/transport/__init__.py | 0 tests/transport/test_httpx_transport.py | 134 ++++++++++ tests/unit/test_auth.py | 174 +++++++++++++ tests/unit/test_client.py | 123 +++++++++ tests/unit/test_config.py | 128 +++++++++ tests/unit/test_errors.py | 139 ++++++++++ tests/unit/test_http_service.py | 147 +++++++++++ tests/unit/test_request_service.py | 174 +++++++++++++ 22 files changed, 2177 insertions(+), 3 deletions(-) create mode 100644 src/aura_python_sdk/_client.py create mode 100644 src/aura_python_sdk/_config.py create mode 100644 src/aura_python_sdk/_errors.py create mode 100644 src/aura_python_sdk/_internal/__init__.py create mode 100644 src/aura_python_sdk/_internal/_auth.py create mode 100644 src/aura_python_sdk/_internal/_request.py create mode 100644 src/aura_python_sdk/_internal/http/__init__.py create mode 100644 src/aura_python_sdk/_internal/http/_httpx.py create mode 100644 src/aura_python_sdk/_internal/http/_service.py create mode 100644 src/aura_python_sdk/_transport.py create mode 100644 tests/fakes.py create mode 100644 tests/transport/__init__.py create mode 100644 tests/transport/test_httpx_transport.py create mode 100644 tests/unit/test_auth.py create mode 100644 tests/unit/test_client.py create mode 100644 tests/unit/test_config.py create mode 100644 tests/unit/test_errors.py create mode 100644 tests/unit/test_http_service.py create mode 100644 tests/unit/test_request_service.py diff --git a/PLAN.md b/PLAN.md index 6de406f..f0ff5bc 100644 --- a/PLAN.md +++ b/PLAN.md @@ -185,6 +185,24 @@ parser accepts the spec's `{"errors": [...]}`, the middleware `{"error": "..."}` - Logging goes to the stdlib `logging` module, at debug level for requests and info level for mutations. Credentials, tokens and passwords are never logged. +### 2.7 Deliberate differences from Go (decided in phase 2) + +- **Retry safety.** Go retries every network error for every method. Here, if the request may have + reached the server (read timeout, connection reset), only idempotent methods (GET, PUT, DELETE, + HEAD, OPTIONS) are retried. This stops a `POST /instances` from being sent twice and creating a + duplicate billable instance. Errors that happen before anything is sent (connect errors, pool + timeouts) are retried for every method. Transports report which case applies through + `AuraConnectionError.request_sent`. +- **One deadline per call.** `timeout` covers the token fetch, every attempt and every backoff, + matching Go's `context.WithTimeout` per method. A retry is skipped if its backoff would pass the + deadline. +- **`max_retries=0` is allowed** and means a single attempt. Go requires at least 1. +- **A 401 clears the cached token**, so the next call fetches a new one. The failed call is not + retried. +- **Token endpoint errors.** Any 4xx from `/oauth/token` raises `AuthenticationError`; 429 and 5xx + keep their usual types. A lower-case `bearer` token type is accepted. +- **Transport ownership.** `close()` closes only a transport the client created itself. + ## 3. Package layout ``` diff --git a/pyproject.toml b/pyproject.toml index a487ca8..3c65226 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -29,7 +29,7 @@ dependencies = ["httpx>=0.27,<1"] prometheus = ["prometheus-client>=0.20"] [project.urls] -Homepage = "https://github.com/neo4j-contrib/aura-python-sdk" +Homepage = "https://github.com/LackOfMorals/aura-python-sdk" "Aura API documentation" = "https://neo4j.com/docs/aura/api/overview/" [dependency-groups] @@ -57,7 +57,7 @@ select = ["E", "W", "F", "I", "B", "UP", "SIM", "RUF", "N", "S", "PT", "RET", "T ignore = [] [tool.ruff.lint.per-file-ignores] -"tests/**" = ["S101", "S105", "S106"] +"tests/**" = ["S101", "S105", "S106", "S107"] [tool.mypy] python_version = "3.11" diff --git a/src/aura_python_sdk/__init__.py b/src/aura_python_sdk/__init__.py index ab84b59..7d6077e 100644 --- a/src/aura_python_sdk/__init__.py +++ b/src/aura_python_sdk/__init__.py @@ -9,6 +9,51 @@ print(instance.id, instance.name) """ +import logging + +from aura_python_sdk._client import AuraClient +from aura_python_sdk._errors import ( + AuraAPIError, + AuraConfigurationError, + AuraConnectionError, + AuraError, + AuraResponseError, + AuraTimeoutError, + AuraValidationError, + AuthenticationError, + BadRequestError, + ConflictError, + ErrorDetail, + NotFoundError, + PermissionDeniedError, + RateLimitError, + ServerError, +) +from aura_python_sdk._transport import HttpRequest, HttpResponse, HttpTransport from aura_python_sdk._version import __version__ -__all__ = ["__version__"] +# Library convention: emit nothing unless the application configures logging. +logging.getLogger(__name__).addHandler(logging.NullHandler()) + +__all__ = [ + "AuraAPIError", + "AuraClient", + "AuraConfigurationError", + "AuraConnectionError", + "AuraError", + "AuraResponseError", + "AuraTimeoutError", + "AuraValidationError", + "AuthenticationError", + "BadRequestError", + "ConflictError", + "ErrorDetail", + "HttpRequest", + "HttpResponse", + "HttpTransport", + "NotFoundError", + "PermissionDeniedError", + "RateLimitError", + "ServerError", + "__version__", +] diff --git a/src/aura_python_sdk/_client.py b/src/aura_python_sdk/_client.py new file mode 100644 index 0000000..237681f --- /dev/null +++ b/src/aura_python_sdk/_client.py @@ -0,0 +1,166 @@ +"""The AuraClient entry point (Go: client.go).""" + +from __future__ import annotations + +import logging +import os +from collections.abc import Mapping +from types import TracebackType +from typing import Self + +from aura_python_sdk._config import ( + API_VERSION, + DEFAULT_BASE_URL, + DEFAULT_MAX_RESPONSE_SIZE, + DEFAULT_MAX_RETRIES, + DEFAULT_TIMEOUT, + DEFAULT_USER_AGENT, + ClientConfig, + build_config, +) +from aura_python_sdk._errors import AuraConfigurationError +from aura_python_sdk._internal._auth import TokenManager +from aura_python_sdk._internal._request import RequestService +from aura_python_sdk._internal.http._httpx import HttpxTransport +from aura_python_sdk._internal.http._service import HttpService +from aura_python_sdk._transport import HttpTransport + +ENV_CLIENT_ID = "AURA_CLIENT_ID" +ENV_CLIENT_SECRET = "AURA_CLIENT_SECRET" # noqa: S105 - environment variable name, not a secret + +_LOGGER_NAME = "aura_python_sdk" + + +class AuraClient: + """Client for the Neo4j Aura API v1. + + Example:: + + with AuraClient(client_id="...", client_secret="...") as client: + ... + + Every option is keyword-only. Invalid options raise :class:`AuraConfigurationError`. + + Args: + client_id: Aura API client ID. + client_secret: Aura API client secret. + base_url: API base URL. It must use HTTPS unless ``allow_insecure_base_url`` is set. + allow_insecure_base_url: Allow an ``http://`` base URL. Only for local test servers, + because credentials would be sent in cleartext. + timeout: Seconds allowed for each API call, covering the token fetch, retries and backoff. + max_retries: How many times to retry after a network failure. Responses with an HTTP + status are never retried. + max_response_size: Largest response body accepted, in bytes. + user_agent: Overrides the ``User-Agent`` header. + default_headers: Extra headers sent with every API request. ``Authorization``, + ``Content-Type`` and ``User-Agent`` are ignored. + logger: Logger for SDK diagnostics. Defaults to the ``aura_python_sdk`` logger. + transport: Custom :class:`HttpTransport`. The client does not close a transport it + did not create. + """ + + def __init__( + self, + *, + client_id: str, + client_secret: str, + base_url: str = DEFAULT_BASE_URL, + allow_insecure_base_url: bool = False, + timeout: float = DEFAULT_TIMEOUT, + max_retries: int = DEFAULT_MAX_RETRIES, + max_response_size: int = DEFAULT_MAX_RESPONSE_SIZE, + user_agent: str = DEFAULT_USER_AGENT, + default_headers: Mapping[str, str] | None = None, + logger: logging.Logger | None = None, + transport: HttpTransport | None = None, + ) -> None: + self._config: ClientConfig = build_config( + client_id=client_id, + client_secret=client_secret, + base_url=base_url, + allow_insecure_base_url=allow_insecure_base_url, + timeout=timeout, + max_retries=max_retries, + max_response_size=max_response_size, + user_agent=user_agent, + default_headers=default_headers, + ) + if transport is not None and not isinstance(transport, HttpTransport): + raise AuraConfigurationError("transport must implement send() and close()") + if logger is not None and not isinstance(logger, logging.Logger): + raise AuraConfigurationError("logger must be a logging.Logger") + + self._logger = logger or logging.getLogger(_LOGGER_NAME) + self._owns_transport = transport is None + self._transport: HttpTransport = transport or HttpxTransport() + self._closed = False + + http = HttpService( + self._transport, + max_retries=self._config.max_retries, + max_response_size=self._config.max_response_size, + logger=self._logger.getChild("http"), + ) + auth = TokenManager( + client_id=self._config.client_id, + client_secret=self._config.client_secret, + token_url=f"{self._config.base_url}/oauth/token", + user_agent=self._config.user_agent, + http=http, + logger=self._logger.getChild("auth"), + ) + self._api = RequestService( + http=http, + auth=auth, + base_url=self._config.base_url, + api_version=API_VERSION, + user_agent=self._config.user_agent, + default_headers=self._config.default_headers, + timeout=self._config.timeout, + logger=self._logger.getChild("api"), + ) + + self._logger.debug( + "Aura API client initialized", + extra={"base_url": self._config.base_url, "api_version": API_VERSION}, + ) + + @classmethod + def from_env(cls, **options: object) -> Self: + """Build a client with credentials from ``AURA_CLIENT_ID`` and ``AURA_CLIENT_SECRET``. + + Any other keyword option is passed through to :class:`AuraClient`. + """ + client_id = os.environ.get(ENV_CLIENT_ID, "") + client_secret = os.environ.get(ENV_CLIENT_SECRET, "") + if not client_id or not client_secret: + raise AuraConfigurationError( + f"{ENV_CLIENT_ID} and {ENV_CLIENT_SECRET} must both be set" + ) + return cls(client_id=client_id, client_secret=client_secret, **options) # type: ignore[arg-type] + + @property + def base_url(self) -> str: + return self._config.base_url + + def close(self) -> None: + """Release pooled connections. Safe to call more than once.""" + if self._closed: + return + self._closed = True + if self._owns_transport: + self._transport.close() + + def __enter__(self) -> Self: + return self + + def __exit__( + self, + exc_type: type[BaseException] | None, + exc: BaseException | None, + tb: TracebackType | None, + ) -> None: + self.close() + + def __repr__(self) -> str: + return f"AuraClient(base_url={self._config.base_url!r})" diff --git a/src/aura_python_sdk/_config.py b/src/aura_python_sdk/_config.py new file mode 100644 index 0000000..4b7aaa4 --- /dev/null +++ b/src/aura_python_sdk/_config.py @@ -0,0 +1,126 @@ +"""Client option defaults and validation (Go: the With* functional options in client.go).""" + +from __future__ import annotations + +import math +from collections.abc import Mapping +from dataclasses import dataclass, field +from urllib.parse import urlsplit + +from aura_python_sdk._errors import AuraConfigurationError +from aura_python_sdk._version import __version__ + +# The Aura API version this client targets. It is deliberately not configurable. +API_VERSION = "v1" + +DEFAULT_BASE_URL = "https://api.neo4j.io" +DEFAULT_TIMEOUT = 120.0 +DEFAULT_MAX_RETRIES = 3 +DEFAULT_MAX_RESPONSE_SIZE = 10 * 1024 * 1024 +DEFAULT_USER_AGENT = f"aura-python-sdk/{__version__}" + +# Headers that default_headers may not override (compared case-insensitively). +PROTECTED_HEADERS = frozenset({"authorization", "content-type", "user-agent"}) + + +@dataclass(frozen=True, slots=True) +class ClientConfig: + client_id: str + client_secret: str = field(repr=False) + base_url: str + timeout: float + max_retries: int + max_response_size: int + user_agent: str + default_headers: Mapping[str, str] + + +def build_config( + *, + client_id: str, + client_secret: str, + base_url: str, + allow_insecure_base_url: bool, + timeout: float, + max_retries: int, + max_response_size: int, + user_agent: str, + default_headers: Mapping[str, str] | None, +) -> ClientConfig: + """Validate every option and raise AuraConfigurationError on the first bad one.""" + if not isinstance(client_id, str) or not client_id: + raise AuraConfigurationError("client ID must not be empty") + if not isinstance(client_secret, str) or not client_secret: + raise AuraConfigurationError("client secret must not be empty") + + return ClientConfig( + client_id=client_id, + client_secret=client_secret, + base_url=_validate_base_url(base_url, allow_insecure=allow_insecure_base_url), + timeout=_validate_timeout(timeout), + max_retries=_validate_non_negative_int("max retries", max_retries), + max_response_size=_validate_positive_int("max response size", max_response_size), + user_agent=_validate_header_value("user agent", user_agent, allow_empty=False), + default_headers=_filter_default_headers(default_headers), + ) + + +def _validate_base_url(base_url: str, *, allow_insecure: bool) -> str: + if not isinstance(base_url, str) or not base_url: + raise AuraConfigurationError("base URL must not be empty") + parts = urlsplit(base_url) + if parts.scheme not in ("https", "http") or not parts.netloc: + raise AuraConfigurationError(f"base URL is not a valid http(s) URL: {base_url!r}") + if parts.scheme != "https" and not allow_insecure: + raise AuraConfigurationError( + "base URL must use HTTPS to protect credentials in transit " + "(pass allow_insecure_base_url=True only for local testing)" + ) + if parts.query or parts.fragment: + raise AuraConfigurationError("base URL must not contain a query string or fragment") + return base_url.rstrip("/") + + +def _validate_timeout(timeout: float) -> float: + if ( + isinstance(timeout, bool) + or not isinstance(timeout, int | float) + or not math.isfinite(timeout) + or timeout <= 0 + ): + raise AuraConfigurationError("timeout must be a finite number of seconds greater than zero") + return float(timeout) + + +def _validate_non_negative_int(name: str, value: int) -> int: + if isinstance(value, bool) or not isinstance(value, int) or value < 0: + raise AuraConfigurationError(f"{name} must be an integer of zero or more") + return value + + +def _validate_positive_int(name: str, value: int) -> int: + if isinstance(value, bool) or not isinstance(value, int) or value <= 0: + raise AuraConfigurationError(f"{name} must be an integer greater than zero") + return value + + +def _validate_header_value(name: str, value: str, *, allow_empty: bool) -> str: + if not isinstance(value, str) or (not value and not allow_empty): + raise AuraConfigurationError(f"{name} must be a non-empty string") + if "\r" in value or "\n" in value: + raise AuraConfigurationError(f"{name} must not contain line breaks") + return value + + +def _filter_default_headers(headers: Mapping[str, str] | None) -> Mapping[str, str]: + """Drop protected headers silently, as the Go SDK does, and reject malformed ones.""" + if not headers: + return {} + filtered: dict[str, str] = {} + for key, value in headers.items(): + if not isinstance(key, str) or not key or any(c in key for c in "\r\n:"): + raise AuraConfigurationError(f"invalid default header name: {key!r}") + _validate_header_value(f"default header {key!r}", value, allow_empty=True) + if key.lower() not in PROTECTED_HEADERS: + filtered[key] = value + return filtered diff --git a/src/aura_python_sdk/_errors.py b/src/aura_python_sdk/_errors.py new file mode 100644 index 0000000..928a01b --- /dev/null +++ b/src/aura_python_sdk/_errors.py @@ -0,0 +1,244 @@ +"""Exceptions raised by the SDK. + +Every exception derives from :class:`AuraError`. Errors returned by the Aura API are +:class:`AuraAPIError` subclasses chosen by HTTP status, so callers can write +``except NotFoundError:`` where the Go SDK uses ``aura.IsNotFound(err)``. +""" + +from __future__ import annotations + +import email.utils +import json +import time +from collections.abc import Mapping, Sequence +from dataclasses import dataclass +from http import HTTPStatus + + +class AuraError(Exception): + """Base class for all errors raised by this SDK.""" + + +class AuraConfigurationError(AuraError, ValueError): + """The client was constructed with invalid options.""" + + +class AuraValidationError(AuraError, ValueError): + """An argument failed client-side validation; no request was sent.""" + + +class AuraConnectionError(AuraError): + """The request could not be completed because of a network failure. + + ``request_sent`` is False when the failure happened before the request reached the server + (for example DNS or connect errors), so retrying cannot duplicate the operation. + """ + + def __init__(self, message: str, *, request_sent: bool) -> None: + super().__init__(message) + self.request_sent = request_sent + + +class AuraTimeoutError(AuraConnectionError): + """The request did not complete within the configured timeout.""" + + +class AuraResponseError(AuraError): + """The API response could not be used: too large, not valid JSON, or an unexpected shape.""" + + +@dataclass(frozen=True, slots=True) +class ErrorDetail: + """One entry from the ``errors`` array of an Aura API error response.""" + + message: str + reason: str | None = None + field: str | None = None + + +class AuraAPIError(AuraError): + """The Aura API returned a non-2xx response.""" + + def __init__( + self, + status_code: int, + message: str, + details: Sequence[ErrorDetail] = (), + *, + request_id: str | None = None, + ) -> None: + self.status_code = status_code + self.message = message + self.details: tuple[ErrorDetail, ...] = tuple(details) + self.request_id = request_id + super().__init__(self._format()) + + def _format(self) -> str: + text = f"API error (status {self.status_code}): {self.message}" + if self.details: + text += f" - {self.details[0].message}" + if len(self.details) > 1: + text += f" (and {len(self.details) - 1} more error(s))" + return text + + def all_errors(self) -> list[str]: + """The top-level message followed by every detail message.""" + return [self.message, *(detail.message for detail in self.details)] + + @property + def has_multiple_errors(self) -> bool: + return len(self.details) > 1 + + @property + def is_not_found(self) -> bool: + return self.status_code == HTTPStatus.NOT_FOUND + + @property + def is_unauthorized(self) -> bool: + return self.status_code == HTTPStatus.UNAUTHORIZED + + @property + def is_bad_request(self) -> bool: + return self.status_code == HTTPStatus.BAD_REQUEST + + +class BadRequestError(AuraAPIError): + """HTTP 400.""" + + +class AuthenticationError(AuraAPIError): + """HTTP 401, or the OAuth token request was rejected.""" + + +class PermissionDeniedError(AuraAPIError): + """HTTP 403.""" + + +class NotFoundError(AuraAPIError): + """HTTP 404.""" + + +class ConflictError(AuraAPIError): + """HTTP 409.""" + + +class RateLimitError(AuraAPIError): + """HTTP 429. ``retry_after`` is the server's suggested wait in seconds, if it sent one.""" + + def __init__( + self, + status_code: int, + message: str, + details: Sequence[ErrorDetail] = (), + *, + request_id: str | None = None, + retry_after: float | None = None, + ) -> None: + self.retry_after = retry_after + super().__init__(status_code, message, details, request_id=request_id) + + +class ServerError(AuraAPIError): + """HTTP 5xx.""" + + +_STATUS_TO_ERROR: dict[int, type[AuraAPIError]] = { + HTTPStatus.BAD_REQUEST: BadRequestError, + HTTPStatus.UNAUTHORIZED: AuthenticationError, + HTTPStatus.FORBIDDEN: PermissionDeniedError, + HTTPStatus.NOT_FOUND: NotFoundError, + HTTPStatus.CONFLICT: ConflictError, +} + + +def api_error_from_response( + status_code: int, + body: bytes, + headers: Mapping[str, str], + *, + error_class: type[AuraAPIError] | None = None, +) -> AuraAPIError: + """Build the exception for a non-2xx response. + + Understands the spec's ``{"errors": [...]}`` shape, the middleware ``{"error": "..."}`` shape, + and ``message`` / ``details`` keys. ``headers`` must have lower-case keys. + """ + message, details = _parse_error_body(body) + if message is None: + message = _status_phrase(status_code) + request_id = headers.get("x-request-id") + + if status_code == HTTPStatus.TOO_MANY_REQUESTS: + return RateLimitError( + status_code, + message, + details, + request_id=request_id, + retry_after=_parse_retry_after(headers.get("retry-after")), + ) + if error_class is None: + error_class = _STATUS_TO_ERROR.get(status_code) + if error_class is None: + error_class = ServerError if status_code >= 500 else AuraAPIError + return error_class(status_code, message, details, request_id=request_id) + + +def _status_phrase(status_code: int) -> str: + try: + return HTTPStatus(status_code).phrase + except ValueError: + return f"HTTP {status_code}" + + +def _parse_error_body(body: bytes) -> tuple[str | None, list[ErrorDetail]]: + if not body: + return None, [] + try: + payload = json.loads(body) + except ValueError: + return None, [] + if not isinstance(payload, dict): + return None, [] + + message = payload.get("message") + if not isinstance(message, str) or not message: + middleware_error = payload.get("error") + message = ( + middleware_error if isinstance(middleware_error, str) and middleware_error else None + ) + + raw_details = payload.get("errors") or payload.get("details") or [] + details = ( + [ + ErrorDetail( + message=str(item.get("message", "")), + reason=_optional_str(item.get("reason")), + field=_optional_str(item.get("field")), + ) + for item in raw_details + if isinstance(item, dict) + ] + if isinstance(raw_details, list) + else [] + ) + return message, details + + +def _optional_str(value: object) -> str | None: + return value if isinstance(value, str) else None + + +def _parse_retry_after(value: str | None) -> float | None: + """Retry-After is either delta-seconds or an HTTP date.""" + if not value: + return None + value = value.strip() + try: + return max(0.0, float(value)) + except ValueError: + pass + try: + parsed = email.utils.parsedate_to_datetime(value) + except (TypeError, ValueError): + return None + return max(0.0, parsed.timestamp() - time.time()) diff --git a/src/aura_python_sdk/_internal/__init__.py b/src/aura_python_sdk/_internal/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/aura_python_sdk/_internal/_auth.py b/src/aura_python_sdk/_internal/_auth.py new file mode 100644 index 0000000..97a226b --- /dev/null +++ b/src/aura_python_sdk/_internal/_auth.py @@ -0,0 +1,119 @@ +"""OAuth client-credentials token management (Go: internal/api authManager).""" + +from __future__ import annotations + +import base64 +import json +import logging +import threading +from dataclasses import dataclass +from urllib.parse import urlencode + +from aura_python_sdk._errors import AuraResponseError, AuthenticationError, api_error_from_response +from aura_python_sdk._internal.http._service import HttpService + +# Refresh this many seconds before the token actually expires. +REFRESH_MARGIN = 60.0 +MAX_EXPIRES_IN = 86400 * 365 + + +@dataclass(frozen=True, slots=True) +class _Token: + token_type: str + access_token: str + expires_at: float # on the HttpService clock (monotonic) + + +class TokenManager: + """Obtains and caches a bearer token from ``{base_url}/oauth/token``. + + Thread-safe. Concurrent callers that find the token missing or near expiry trigger a single + refresh between them. + """ + + def __init__( + self, + *, + client_id: str, + client_secret: str, + token_url: str, + user_agent: str, + http: HttpService, + logger: logging.Logger, + ) -> None: + credentials = f"{client_id}:{client_secret}".encode() + self._basic_auth = "Basic " + base64.b64encode(credentials).decode("ascii") + self._token_url = token_url + self._user_agent = user_agent + self._http = http + self._logger = logger + self._lock = threading.Lock() + self._token: _Token | None = None + + def authorization_header(self, *, deadline: float) -> str: + """Return a valid ``Authorization`` header value, fetching a new token if needed.""" + token = self._token + if token is None or not self._is_fresh(token): + with self._lock: + token = self._token + if token is None or not self._is_fresh(token): + token = self._fetch(deadline=deadline) + self._token = token + return f"{token.token_type} {token.access_token}" + + def invalidate(self) -> None: + """Drop the cached token so the next request fetches a new one (e.g. after a 401).""" + with self._lock: + self._token = None + + def _is_fresh(self, token: _Token) -> bool: + return self._http.clock() < token.expires_at - REFRESH_MARGIN + + def _fetch(self, *, deadline: float) -> _Token: + self._logger.debug("obtaining new authentication token") + response = self._http.send( + "POST", + self._token_url, + { + "Authorization": self._basic_auth, + "Content-Type": "application/x-www-form-urlencoded", + "User-Agent": self._user_agent, + }, + urlencode({"grant_type": "client_credentials"}).encode("ascii"), + deadline=deadline, + ) + if not 200 <= response.status_code < 300: + status = response.status_code + # Any client error from the token endpoint means the credentials were rejected. + # Rate limits and server errors keep their usual types. + error_class = None if status == 429 or status >= 500 else AuthenticationError + self._logger.debug("token request failed", extra={"status": status}) + raise api_error_from_response( + status, response.body, response.headers, error_class=error_class + ) + + try: + payload = json.loads(response.body) + token_type = payload["token_type"] + access_token = payload["access_token"] + expires_in = payload["expires_in"] + except (ValueError, KeyError, TypeError) as exc: + raise AuraResponseError("failed to parse token response") from exc + + if not isinstance(token_type, str) or token_type.lower() != "bearer": + raise AuraResponseError(f"token type is not valid: {token_type!r}") + if not isinstance(access_token, str) or not access_token: + raise AuraResponseError("token response did not contain an access token") + if ( + isinstance(expires_in, bool) + or not isinstance(expires_in, int | float) + or not 0 < expires_in <= MAX_EXPIRES_IN + ): + raise AuraResponseError(f"invalid expires_in value: {expires_in!r}") + + self._logger.debug("token obtained", extra={"expires_in": expires_in}) + return _Token( + token_type="Bearer", # noqa: S106 - the OAuth scheme name, not a secret + access_token=access_token, + expires_at=self._http.clock() + float(expires_in), + ) diff --git a/src/aura_python_sdk/_internal/_request.py b/src/aura_python_sdk/_internal/_request.py new file mode 100644 index 0000000..6bee12c --- /dev/null +++ b/src/aura_python_sdk/_internal/_request.py @@ -0,0 +1,124 @@ +"""Authenticated Aura API requests (Go: internal/api RequestService).""" + +from __future__ import annotations + +import json +import logging +from collections.abc import Mapping +from dataclasses import dataclass, field +from urllib.parse import quote, urlencode + +from aura_python_sdk._errors import AuraResponseError, api_error_from_response +from aura_python_sdk._internal._auth import TokenManager +from aura_python_sdk._internal.http._service import HttpService + +QueryParams = Mapping[str, str | None] + + +def build_path(*segments: str) -> str: + """Join path segments, percent-encoding each so an ID can never alter the path.""" + return "/".join(quote(segment, safe="") for segment in segments) + + +@dataclass(frozen=True, slots=True) +class ApiResponse: + status_code: int + headers: Mapping[str, str] = field(default_factory=dict) + body: bytes = b"" + + def json(self) -> object: + try: + return json.loads(self.body) + except ValueError as exc: + raise AuraResponseError("response body is not valid JSON") from exc + + +class RequestService: + """Adds authentication, headers and URL handling, and maps error responses to exceptions. + + A relative path such as ``instances/abc`` resolves to + ``{base_url}/{api_version}/instances/abc``. + An absolute ``http(s)://`` URL, such as a Prometheus metrics endpoint, is used unchanged but + still gets the Aura bearer token. + """ + + def __init__( + self, + *, + http: HttpService, + auth: TokenManager, + base_url: str, + api_version: str, + user_agent: str, + default_headers: Mapping[str, str], + timeout: float, + logger: logging.Logger, + ) -> None: + self._http = http + self._auth = auth + self._endpoint_base = f"{base_url}/{api_version}" + self._user_agent = user_agent + self._default_headers = dict(default_headers) + self._timeout = timeout + self._logger = logger + + def get(self, path: str, *, params: QueryParams | None = None) -> ApiResponse: + return self.request("GET", path, params=params) + + def post(self, path: str, *, json_body: object = None) -> ApiResponse: + return self.request("POST", path, json_body=json_body) + + def patch(self, path: str, *, json_body: object = None) -> ApiResponse: + return self.request("PATCH", path, json_body=json_body) + + def put(self, path: str, *, json_body: object = None) -> ApiResponse: + return self.request("PUT", path, json_body=json_body) + + def delete(self, path: str) -> ApiResponse: + return self.request("DELETE", path) + + def request( + self, + method: str, + path: str, + *, + params: QueryParams | None = None, + json_body: object = None, + ) -> ApiResponse: + # One deadline covers the token fetch, every attempt and every backoff, like the + # context.WithTimeout that wraps each Go service method. + deadline = self._http.clock() + self._timeout + url = self._resolve_url(path, params) + + headers = dict(self._default_headers) + headers["Content-Type"] = "application/json" + headers["User-Agent"] = self._user_agent + headers["Authorization"] = self._auth.authorization_header(deadline=deadline) + + body = None if json_body is None else json.dumps(json_body, separators=(",", ":")).encode() + + self._logger.debug("making authenticated API request", extra={"method": method, "url": url}) + response = self._http.send(method, url, headers, body, deadline=deadline) + + if not 200 <= response.status_code < 300: + if response.status_code == 401: + # The token may have been revoked; make the next call fetch a fresh one. + self._auth.invalidate() + error = api_error_from_response(response.status_code, response.body, response.headers) + self._logger.debug( + "API returned error", + extra={"method": method, "url": url, "status": response.status_code}, + ) + raise error + + return ApiResponse(response.status_code, response.headers, response.body) + + def _resolve_url(self, path: str, params: QueryParams | None) -> str: + if path.startswith(("https://", "http://")): + url = path + else: + url = f"{self._endpoint_base}/{path.lstrip('/')}" + query = {key: value for key, value in (params or {}).items() if value is not None} + if query: + url += ("&" if "?" in url else "?") + urlencode(query) + return url diff --git a/src/aura_python_sdk/_internal/http/__init__.py b/src/aura_python_sdk/_internal/http/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/aura_python_sdk/_internal/http/_httpx.py b/src/aura_python_sdk/_internal/http/_httpx.py new file mode 100644 index 0000000..8166382 --- /dev/null +++ b/src/aura_python_sdk/_internal/http/_httpx.py @@ -0,0 +1,79 @@ +"""The default transport, backed by httpx. + +This is the only module in the SDK that imports httpx (enforced by +tests/unit/test_import_boundaries.py). Every httpx type and exception is translated to the SDK's +own types at this boundary. +""" + +from __future__ import annotations + +import ssl + +import httpx + +from aura_python_sdk._errors import AuraConnectionError, AuraResponseError, AuraTimeoutError +from aura_python_sdk._transport import HttpRequest, HttpResponse + +# Mirrors the Go SDK's http.Transport settings. +_LIMITS = httpx.Limits(max_connections=100, max_keepalive_connections=20, keepalive_expiry=90.0) + +# httpx errors raised before any request bytes reach the server, so retrying cannot duplicate +# a mutation. +_NOT_SENT_ERRORS = (httpx.ConnectError, httpx.ConnectTimeout, httpx.PoolTimeout) + + +def _tls_context() -> ssl.SSLContext: + context = ssl.create_default_context() + context.minimum_version = ssl.TLSVersion.TLSv1_2 + return context + + +class HttpxTransport: + """An :class:`~aura_python_sdk.HttpTransport` backed by a pooled ``httpx.Client``.""" + + def __init__(self, *, _httpx_transport: httpx.BaseTransport | None = None) -> None: + # _httpx_transport is only for tests; it replaces the network layer below httpx. + self._client = httpx.Client( + verify=_tls_context(), + limits=_LIMITS, + follow_redirects=True, + transport=_httpx_transport, + ) + + def send(self, request: HttpRequest) -> HttpResponse: + try: + with self._client.stream( + request.method, + request.url, + headers=dict(request.headers), + content=request.body, + timeout=httpx.Timeout(request.timeout), + ) as response: + body = self._read_limited(response, request.max_response_size) + return HttpResponse( + status_code=response.status_code, + headers=dict(response.headers.items()), + body=body, + ) + except httpx.TimeoutException as exc: + raise AuraTimeoutError( + f"request timed out: {exc}", request_sent=not isinstance(exc, _NOT_SENT_ERRORS) + ) from exc + except httpx.TransportError as exc: + raise AuraConnectionError( + f"request failed: {exc}", request_sent=not isinstance(exc, _NOT_SENT_ERRORS) + ) from exc + + @staticmethod + def _read_limited(response: httpx.Response, limit: int) -> bytes: + chunks: list[bytes] = [] + size = 0 + for chunk in response.iter_bytes(): + size += len(chunk) + if size > limit: + raise AuraResponseError(f"response body exceeded limit of {limit} bytes") + chunks.append(chunk) + return b"".join(chunks) + + def close(self) -> None: + self._client.close() diff --git a/src/aura_python_sdk/_internal/http/_service.py b/src/aura_python_sdk/_internal/http/_service.py new file mode 100644 index 0000000..5c299c4 --- /dev/null +++ b/src/aura_python_sdk/_internal/http/_service.py @@ -0,0 +1,102 @@ +"""Retries and response limits on top of an HttpTransport (Go: internal/httpclient).""" + +from __future__ import annotations + +import logging +import time +from collections.abc import Callable, Mapping + +from aura_python_sdk._errors import AuraConnectionError, AuraResponseError, AuraTimeoutError +from aura_python_sdk._transport import HttpRequest, HttpResponse, HttpTransport + +# Methods that are safe to repeat when the server may already have received the request. +_IDEMPOTENT_METHODS = frozenset({"GET", "HEAD", "OPTIONS", "PUT", "DELETE"}) + +RETRY_WAIT_MIN = 1.0 +RETRY_WAIT_MAX = 5.0 + + +class HttpService: + """Sends requests through a transport, retrying network failures only. + + As in the Go SDK, a response with any HTTP status is final and is never retried. A network + failure is retried up to ``max_retries`` times with exponential backoff (1 s doubling to 5 s). + If the request may have reached the server, only idempotent methods are retried, so a + ``POST /instances`` is never sent twice. No attempt or backoff runs past ``deadline``. + """ + + def __init__( + self, + transport: HttpTransport, + *, + max_retries: int, + max_response_size: int, + logger: logging.Logger, + clock: Callable[[], float] = time.monotonic, + sleep: Callable[[float], None] = time.sleep, + ) -> None: + self._transport = transport + self._max_retries = max_retries + self._max_response_size = max_response_size + self._logger = logger + self._clock = clock + self._sleep = sleep + + @property + def clock(self) -> Callable[[], float]: + return self._clock + + def send( + self, + method: str, + url: str, + headers: Mapping[str, str], + body: bytes | None, + *, + deadline: float, + ) -> HttpResponse: + attempt = 0 + while True: + remaining = deadline - self._clock() + if remaining <= 0: + raise AuraTimeoutError("request deadline exceeded", request_sent=False) + request = HttpRequest( + method=method, + url=url, + headers=headers, + body=body, + timeout=remaining, + max_response_size=self._max_response_size, + ) + self._logger.debug("sending HTTP request", extra={"method": method, "url": url}) + try: + response = self._transport.send(request) + except AuraConnectionError as exc: + wait = min(RETRY_WAIT_MAX, RETRY_WAIT_MIN * 2**attempt) + if ( + attempt >= self._max_retries + or not self._is_retryable(method, exc) + or self._clock() + wait >= deadline + ): + raise + self._logger.debug( + "retrying HTTP request after network error", + extra={"method": method, "url": url, "attempt": attempt + 1, "error": str(exc)}, + ) + self._sleep(wait) + attempt += 1 + continue + + if len(response.body) > self._max_response_size: + raise AuraResponseError( + f"response body exceeded limit of {self._max_response_size} bytes" + ) + self._logger.debug( + "HTTP response received", + extra={"method": method, "url": url, "status": response.status_code}, + ) + return response + + @staticmethod + def _is_retryable(method: str, exc: AuraConnectionError) -> bool: + return not exc.request_sent or method.upper() in _IDEMPOTENT_METHODS diff --git a/src/aura_python_sdk/_transport.py b/src/aura_python_sdk/_transport.py new file mode 100644 index 0000000..fe55315 --- /dev/null +++ b/src/aura_python_sdk/_transport.py @@ -0,0 +1,57 @@ +"""The HTTP transport interface. + +A transport sends exactly one HTTP request and returns the response. Retries, authentication and +error mapping happen above it, so a custom transport only has to move bytes. Pass one to +``AuraClient(transport=...)`` to control proxies, TLS or connection handling, or to fake the network +in tests. This is the Python equivalent of the Go SDK's ``WithHTTPClient`` option. +""" + +from __future__ import annotations + +from collections.abc import Mapping +from dataclasses import dataclass, field +from typing import Protocol, runtime_checkable + + +@dataclass(frozen=True, slots=True) +class HttpRequest: + """A single HTTP request to send. + + ``timeout`` is in seconds. A transport must not read more than ``max_response_size`` bytes + of the response body. If the body is larger, it raises + :class:`~aura_python_sdk.AuraResponseError`. + """ + + method: str + url: str + headers: Mapping[str, str] + body: bytes | None + timeout: float + max_response_size: int + + +@dataclass(frozen=True, slots=True) +class HttpResponse: + """The status, headers and fully read body of a response. Header names are lower-cased.""" + + status_code: int + headers: Mapping[str, str] = field(default_factory=dict) + body: bytes = b"" + + def __post_init__(self) -> None: + object.__setattr__(self, "headers", {k.lower(): v for k, v in self.headers.items()}) + + +@runtime_checkable +class HttpTransport(Protocol): + """Sends HTTP requests. + + On a network failure, ``send`` raises :class:`~aura_python_sdk.AuraConnectionError`, or + :class:`~aura_python_sdk.AuraTimeoutError` for timeouts. It sets ``request_sent=False`` only + when it is certain the server never received the request. It returns non-2xx responses + normally instead of raising. + """ + + def send(self, request: HttpRequest) -> HttpResponse: ... + + def close(self) -> None: ... diff --git a/tests/fakes.py b/tests/fakes.py new file mode 100644 index 0000000..474e188 --- /dev/null +++ b/tests/fakes.py @@ -0,0 +1,75 @@ +"""Test doubles shared by the unit tests. No network access.""" + +from __future__ import annotations + +import json +from collections import deque +from collections.abc import Callable, Iterable +from dataclasses import dataclass, field + +from aura_python_sdk import HttpRequest, HttpResponse + +Reply = HttpResponse | Exception | Callable[[HttpRequest], HttpResponse] + + +def json_response( + status_code: int, payload: object, headers: dict[str, str] | None = None +) -> HttpResponse: + return HttpResponse( + status_code=status_code, + headers={"Content-Type": "application/json", **(headers or {})}, + body=json.dumps(payload).encode(), + ) + + +def token_response( + access_token: str = "token-1", expires_in: int = 3600, token_type: str = "Bearer" +) -> HttpResponse: + return json_response( + 200, {"access_token": access_token, "expires_in": expires_in, "token_type": token_type} + ) + + +@dataclass +class FakeClock: + """A monotonic clock that only moves when told to; ``sleep`` advances it.""" + + now: float = 1000.0 + sleeps: list[float] = field(default_factory=list) + + def __call__(self) -> float: + return self.now + + def sleep(self, seconds: float) -> None: + self.sleeps.append(seconds) + self.now += seconds + + +class FakeTransport: + """Replays queued replies in order and records every request.""" + + def __init__(self, replies: Iterable[Reply] = ()) -> None: + self.replies: deque[Reply] = deque(replies) + self.requests: list[HttpRequest] = [] + self.closed = False + + def queue(self, *replies: Reply) -> None: + self.replies.extend(replies) + + def send(self, request: HttpRequest) -> HttpResponse: + self.requests.append(request) + if not self.replies: + raise AssertionError(f"unexpected request: {request.method} {request.url}") + reply = self.replies.popleft() + if isinstance(reply, Exception): + raise reply + if callable(reply): + return reply(request) + return reply + + def close(self) -> None: + self.closed = True + + @property + def api_requests(self) -> list[HttpRequest]: + return [r for r in self.requests if not r.url.endswith("/oauth/token")] diff --git a/tests/transport/__init__.py b/tests/transport/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/transport/test_httpx_transport.py b/tests/transport/test_httpx_transport.py new file mode 100644 index 0000000..3340095 --- /dev/null +++ b/tests/transport/test_httpx_transport.py @@ -0,0 +1,134 @@ +"""HttpxTransport against httpx.MockTransport (tests may import httpx; src may not).""" + +import ssl +from collections.abc import Iterator + +import httpx +import pytest + +from aura_python_sdk import ( + AuraConnectionError, + AuraResponseError, + AuraTimeoutError, + HttpRequest, + HttpTransport, +) +from aura_python_sdk._internal.http._httpx import HttpxTransport + + +def _request(**overrides: object) -> HttpRequest: + values: dict[str, object] = { + "method": "POST", + "url": "https://api.neo4j.io/v1/instances", + "headers": {"Authorization": "Bearer t", "Content-Type": "application/json"}, + "body": b'{"name":"x"}', + "timeout": 12.5, + "max_response_size": 1024, + } + values.update(overrides) + return HttpRequest(**values) # type: ignore[arg-type] + + +def _transport(handler: object) -> HttpxTransport: + return HttpxTransport(_httpx_transport=httpx.MockTransport(handler)) # type: ignore[arg-type] + + +def test_satisfies_protocol() -> None: + assert isinstance(HttpxTransport(), HttpTransport) + + +def test_request_and_response_are_translated() -> None: + seen: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + seen.append(request) + return httpx.Response(202, headers={"X-Request-Id": "r1"}, content=b'{"data":{}}') + + response = _transport(handler).send(_request()) + + [request] = seen + assert request.method == "POST" + assert str(request.url) == "https://api.neo4j.io/v1/instances" + assert request.headers["authorization"] == "Bearer t" + assert request.content == b'{"name":"x"}' + assert request.extensions["timeout"] == { + "connect": 12.5, + "read": 12.5, + "write": 12.5, + "pool": 12.5, + } + assert response.status_code == 202 + assert response.headers["x-request-id"] == "r1" + assert response.body == b'{"data":{}}' + + +def test_error_statuses_are_returned_not_raised() -> None: + response = _transport(lambda r: httpx.Response(500, content=b"boom")).send(_request()) + assert response.status_code == 500 + assert response.body == b"boom" + + +def test_body_over_limit_is_rejected_while_streaming() -> None: + chunks_read = 0 + + def stream() -> Iterator[bytes]: + nonlocal chunks_read + for _ in range(100): + chunks_read += 1 + yield b"x" * 512 + + handler = lambda r: httpx.Response(200, content=stream()) # noqa: E731 + with pytest.raises(AuraResponseError, match="exceeded limit"): + _transport(handler).send(_request(max_response_size=1024)) + assert chunks_read < 100 + + +def test_body_at_limit_is_accepted() -> None: + response = _transport(lambda r: httpx.Response(200, content=b"x" * 1024)).send(_request()) + assert len(response.body) == 1024 + + +def test_redirects_are_followed() -> None: + def handler(request: httpx.Request) -> httpx.Response: + if request.url.path == "/v1/old": + return httpx.Response(308, headers={"Location": "https://api.neo4j.io/v1/new"}) + return httpx.Response(200, content=b"moved") + + response = _transport(handler).send( + _request(method="GET", body=None, url="https://api.neo4j.io/v1/old") + ) + assert response.body == b"moved" + + +@pytest.mark.parametrize( + ("exc", "expected_type", "request_sent"), + [ + (httpx.ConnectError("refused"), AuraConnectionError, False), + (httpx.ConnectTimeout("slow connect"), AuraTimeoutError, False), + (httpx.PoolTimeout("pool"), AuraTimeoutError, False), + (httpx.ReadTimeout("slow read"), AuraTimeoutError, True), + (httpx.WriteTimeout("slow write"), AuraTimeoutError, True), + (httpx.ReadError("reset"), AuraConnectionError, True), + (httpx.RemoteProtocolError("bad"), AuraConnectionError, True), + ], +) +def test_network_errors_are_translated( + exc: Exception, expected_type: type[AuraConnectionError], request_sent: bool +) -> None: + def handler(request: httpx.Request) -> httpx.Response: + raise exc + + with pytest.raises(expected_type) as info: + _transport(handler).send(_request()) + assert type(info.value) is expected_type + assert info.value.request_sent is request_sent + assert isinstance(info.value.__cause__, httpx.HTTPError) + + +def test_tls_minimum_is_1_2() -> None: + transport = HttpxTransport() + pool = transport._client._transport._pool # type: ignore[attr-defined] + context: ssl.SSLContext = pool._ssl_context + assert context.minimum_version == ssl.TLSVersion.TLSv1_2 + assert context.verify_mode == ssl.CERT_REQUIRED + transport.close() diff --git a/tests/unit/test_auth.py b/tests/unit/test_auth.py new file mode 100644 index 0000000..f3f8f9d --- /dev/null +++ b/tests/unit/test_auth.py @@ -0,0 +1,174 @@ +import base64 +import logging +import threading +import time +from urllib.parse import parse_qs + +import pytest + +from aura_python_sdk import ( + AuraResponseError, + AuthenticationError, + HttpRequest, + HttpResponse, + RateLimitError, + ServerError, +) +from aura_python_sdk._internal._auth import TokenManager +from aura_python_sdk._internal.http._service import HttpService +from tests.fakes import FakeClock, FakeTransport, json_response, token_response + +TOKEN_URL = "https://api.neo4j.io/oauth/token" + + +def _manager(transport: FakeTransport, clock: FakeClock | None = None) -> TokenManager: + clock = clock or FakeClock() + http = HttpService( + transport, + max_retries=0, + max_response_size=1024, + logger=logging.getLogger("test"), + clock=clock, + sleep=clock.sleep, + ) + return TokenManager( + client_id="my-id", + client_secret="my-secret", + token_url=TOKEN_URL, + user_agent="ua/1", + http=http, + logger=logging.getLogger("test"), + ) + + +def _deadline() -> float: + return float("inf") + + +def test_token_request_shape() -> None: + transport = FakeTransport([token_response("abc")]) + header = _manager(transport).authorization_header(deadline=_deadline()) + + assert header == "Bearer abc" + [request] = transport.requests + assert request.method == "POST" + assert request.url == TOKEN_URL + expected_basic = base64.b64encode(b"my-id:my-secret").decode() + assert request.headers["Authorization"] == f"Basic {expected_basic}" + assert request.headers["Content-Type"] == "application/x-www-form-urlencoded" + assert request.headers["User-Agent"] == "ua/1" + assert request.body is not None + assert parse_qs(request.body.decode()) == {"grant_type": ["client_credentials"]} + + +def test_token_is_cached() -> None: + transport = FakeTransport([token_response("abc")]) + manager = _manager(transport) + for _ in range(3): + assert manager.authorization_header(deadline=_deadline()) == "Bearer abc" + assert len(transport.requests) == 1 + + +def test_token_refreshed_sixty_seconds_before_expiry() -> None: + clock = FakeClock() + transport = FakeTransport([token_response("first", expires_in=3600), token_response("second")]) + manager = _manager(transport, clock) + + assert manager.authorization_header(deadline=_deadline()) == "Bearer first" + clock.now += 3600 - 61 + assert manager.authorization_header(deadline=_deadline()) == "Bearer first" + clock.now += 1 + assert manager.authorization_header(deadline=_deadline()) == "Bearer second" + + +def test_invalidate_forces_refetch() -> None: + transport = FakeTransport([token_response("first"), token_response("second")]) + manager = _manager(transport) + manager.authorization_header(deadline=_deadline()) + manager.invalidate() + assert manager.authorization_header(deadline=_deadline()) == "Bearer second" + + +def test_lowercase_bearer_is_normalised() -> None: + transport = FakeTransport([token_response("abc", token_type="bearer")]) + assert _manager(transport).authorization_header(deadline=_deadline()) == "Bearer abc" + + +@pytest.mark.parametrize("status", [400, 401, 403]) +def test_client_errors_raise_authentication_error(status: int) -> None: + body = {"errors": [{"message": "invalid client credentials", "reason": "invalid_client"}]} + transport = FakeTransport([json_response(status, body)]) + with pytest.raises(AuthenticationError) as info: + _manager(transport).authorization_header(deadline=_deadline()) + assert info.value.status_code == status + assert info.value.details[0].reason == "invalid_client" + + +def test_rate_limit_and_server_errors_keep_their_type() -> None: + transport = FakeTransport([HttpResponse(429), HttpResponse(503)]) + manager = _manager(transport) + with pytest.raises(RateLimitError): + manager.authorization_header(deadline=_deadline()) + with pytest.raises(ServerError): + manager.authorization_header(deadline=_deadline()) + + +@pytest.mark.parametrize( + "payload", + [ + {"access_token": "a", "expires_in": 3600, "token_type": "MAC"}, + {"access_token": "", "expires_in": 3600, "token_type": "Bearer"}, + {"access_token": "a", "expires_in": 0, "token_type": "Bearer"}, + {"access_token": "a", "expires_in": -5, "token_type": "Bearer"}, + {"access_token": "a", "expires_in": 86400 * 365 + 1, "token_type": "Bearer"}, + {"access_token": "a", "expires_in": "3600", "token_type": "Bearer"}, + {"access_token": "a", "expires_in": True, "token_type": "Bearer"}, + {"access_token": "a", "token_type": "Bearer"}, + ["not", "an", "object"], + ], +) +def test_invalid_token_responses(payload: object) -> None: + transport = FakeTransport([json_response(200, payload)]) + with pytest.raises(AuraResponseError): + _manager(transport).authorization_header(deadline=_deadline()) + + +def test_non_json_token_response() -> None: + transport = FakeTransport([HttpResponse(200, body=b"")]) + with pytest.raises(AuraResponseError): + _manager(transport).authorization_header(deadline=_deadline()) + + +def test_failed_fetch_does_not_cache() -> None: + transport = FakeTransport([HttpResponse(503), token_response("ok")]) + manager = _manager(transport) + with pytest.raises(ServerError): + manager.authorization_header(deadline=_deadline()) + assert manager.authorization_header(deadline=_deadline()) == "Bearer ok" + + +def test_concurrent_callers_share_one_fetch() -> None: + calls = 0 + + def slow_token(request: HttpRequest) -> HttpResponse: + nonlocal calls + calls += 1 + time.sleep(0.05) + return token_response("shared") + + transport = FakeTransport([slow_token]) + manager = _manager(transport) + results: list[str] = [] + threads = [ + threading.Thread( + target=lambda: results.append(manager.authorization_header(deadline=_deadline())) + ) + for _ in range(10) + ] + for thread in threads: + thread.start() + for thread in threads: + thread.join() + + assert calls == 1 + assert results == ["Bearer shared"] * 10 diff --git a/tests/unit/test_client.py b/tests/unit/test_client.py new file mode 100644 index 0000000..1fcf31e --- /dev/null +++ b/tests/unit/test_client.py @@ -0,0 +1,123 @@ +import logging + +import pytest + +import aura_python_sdk as aura +from aura_python_sdk import AuraClient, AuraConfigurationError +from aura_python_sdk._internal.http._httpx import HttpxTransport +from tests.fakes import FakeTransport, json_response, token_response + + +def test_construct_with_fake_transport_and_make_a_call() -> None: + transport = FakeTransport([token_response("tok"), json_response(200, {"data": []})]) + client = AuraClient(client_id="id", client_secret="secret", transport=transport) + + client._api.get("tenants") + + [request] = transport.api_requests + assert request.url == "https://api.neo4j.io/v1/tenants" + assert request.headers["User-Agent"] == f"aura-python-sdk/{aura.__version__}" + assert transport.requests[0].url == "https://api.neo4j.io/oauth/token" + + +def test_options_are_wired_through() -> None: + transport = FakeTransport([token_response(), json_response(200, {})]) + client = AuraClient( + client_id="id", + client_secret="secret", + base_url="http://localhost:9000", + allow_insecure_base_url=True, + timeout=7, + max_response_size=2048, + user_agent="my-app/2", + default_headers={"X-Team": "db"}, + transport=transport, + ) + client._api.get("instances") + + token_request, api_request = transport.requests + assert token_request.url == "http://localhost:9000/oauth/token" + assert api_request.url == "http://localhost:9000/v1/instances" + assert api_request.timeout == pytest.approx(7, abs=0.5) + assert api_request.max_response_size == 2048 + assert api_request.headers["User-Agent"] == "my-app/2" + assert api_request.headers["X-Team"] == "db" + assert client.base_url == "http://localhost:9000" + + +def test_invalid_option_raises_configuration_error() -> None: + with pytest.raises(AuraConfigurationError): + AuraClient(client_id="", client_secret="secret") + assert issubclass(AuraConfigurationError, ValueError) + + +def test_rejects_non_transport() -> None: + with pytest.raises(AuraConfigurationError, match="transport"): + AuraClient(client_id="id", client_secret="s", transport=object()) # type: ignore[arg-type] + + +def test_rejects_non_logger() -> None: + with pytest.raises(AuraConfigurationError, match="logger"): + AuraClient(client_id="id", client_secret="s", logger="debug") # type: ignore[arg-type] + + +def test_default_transport_is_httpx_and_owned() -> None: + client = AuraClient(client_id="id", client_secret="secret") + assert isinstance(client._transport, HttpxTransport) + client.close() + client.close() # idempotent + + +def test_context_manager_closes_owned_transport(monkeypatch: pytest.MonkeyPatch) -> None: + closed: list[bool] = [] + monkeypatch.setattr(HttpxTransport, "close", lambda self: closed.append(True)) + with AuraClient(client_id="id", client_secret="secret"): + pass + assert closed == [True] + + +def test_does_not_close_caller_transport() -> None: + transport = FakeTransport() + with AuraClient(client_id="id", client_secret="secret", transport=transport): + pass + assert transport.closed is False + + +def test_repr_hides_credentials() -> None: + client = AuraClient(client_id="id-123", client_secret="s3cr3t", transport=FakeTransport()) + assert "s3cr3t" not in repr(client) + assert repr(client) == "AuraClient(base_url='https://api.neo4j.io')" + + +def test_from_env(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("AURA_CLIENT_ID", "env-id") + monkeypatch.setenv("AURA_CLIENT_SECRET", "env-secret") + transport = FakeTransport([token_response(), json_response(200, {})]) + client = AuraClient.from_env(transport=transport, timeout=10) + client._api.get("tenants") + assert transport.requests[0].headers["Authorization"].startswith("Basic ") + + +@pytest.mark.parametrize("missing", ["AURA_CLIENT_ID", "AURA_CLIENT_SECRET"]) +def test_from_env_requires_both(monkeypatch: pytest.MonkeyPatch, missing: str) -> None: + monkeypatch.setenv("AURA_CLIENT_ID", "env-id") + monkeypatch.setenv("AURA_CLIENT_SECRET", "env-secret") + monkeypatch.delenv(missing) + with pytest.raises(AuraConfigurationError, match="must both be set"): + AuraClient.from_env() + + +def test_library_logger_has_null_handler() -> None: + handlers = logging.getLogger("aura_python_sdk").handlers + assert any(isinstance(h, logging.NullHandler) for h in handlers) + + +def test_secrets_never_logged(caplog: pytest.LogCaptureFixture) -> None: + transport = FakeTransport([token_response("tok-value"), json_response(200, {})]) + client = AuraClient(client_id="id", client_secret="s3cr3t", transport=transport) + with caplog.at_level(logging.DEBUG, logger="aura_python_sdk"): + client._api.get("instances") + text = "\n".join(f"{record.getMessage()} {record.__dict__}" for record in caplog.records) + assert caplog.records + assert "s3cr3t" not in text + assert "tok-value" not in text diff --git a/tests/unit/test_config.py b/tests/unit/test_config.py new file mode 100644 index 0000000..b23958c --- /dev/null +++ b/tests/unit/test_config.py @@ -0,0 +1,128 @@ +from typing import Any + +import pytest + +from aura_python_sdk import AuraConfigurationError +from aura_python_sdk._config import ( + DEFAULT_BASE_URL, + DEFAULT_MAX_RESPONSE_SIZE, + DEFAULT_MAX_RETRIES, + DEFAULT_TIMEOUT, + DEFAULT_USER_AGENT, + ClientConfig, + build_config, +) + + +def _config(**overrides: Any) -> ClientConfig: + options: dict[str, Any] = { + "client_id": "id", + "client_secret": "secret", + "base_url": DEFAULT_BASE_URL, + "allow_insecure_base_url": False, + "timeout": DEFAULT_TIMEOUT, + "max_retries": DEFAULT_MAX_RETRIES, + "max_response_size": DEFAULT_MAX_RESPONSE_SIZE, + "user_agent": DEFAULT_USER_AGENT, + "default_headers": None, + } + options.update(overrides) + return build_config(**options) + + +def test_defaults_match_go_sdk() -> None: + config = _config() + assert config.base_url == "https://api.neo4j.io" + assert config.timeout == 120.0 + assert config.max_retries == 3 + assert config.max_response_size == 10 * 1024 * 1024 + assert config.user_agent.startswith("aura-python-sdk/") + assert config.default_headers == {} + + +def test_secret_is_not_in_repr() -> None: + assert "super-secret" not in repr(_config(client_secret="super-secret")) + + +@pytest.mark.parametrize("field", ["client_id", "client_secret"]) +@pytest.mark.parametrize("value", ["", None]) +def test_credentials_required(field: str, value: object) -> None: + with pytest.raises(AuraConfigurationError, match="must not be empty"): + _config(**{field: value}) + + +def test_base_url_requires_https() -> None: + with pytest.raises(AuraConfigurationError, match="HTTPS"): + _config(base_url="http://localhost:8080") + + +def test_insecure_base_url_allowed_when_opted_in() -> None: + config = _config(base_url="http://localhost:8080/", allow_insecure_base_url=True) + assert config.base_url == "http://localhost:8080" + + +@pytest.mark.parametrize( + "base_url", + ["", "api.neo4j.io", "ftp://api.neo4j.io", "https://", "https://x?y=1", "https://x#f"], +) +def test_invalid_base_urls(base_url: str) -> None: + with pytest.raises(AuraConfigurationError): + _config(base_url=base_url, allow_insecure_base_url=True) + + +def test_trailing_slash_is_stripped() -> None: + assert ( + _config(base_url="https://staging.example.com/").base_url == "https://staging.example.com" + ) + + +@pytest.mark.parametrize("timeout", [0, -1, float("inf"), float("nan"), True, "10"]) +def test_invalid_timeout(timeout: object) -> None: + with pytest.raises(AuraConfigurationError, match="timeout"): + _config(timeout=timeout) + + +def test_integer_timeout_is_accepted() -> None: + assert _config(timeout=5).timeout == 5.0 + + +def test_zero_retries_is_allowed() -> None: + assert _config(max_retries=0).max_retries == 0 + + +@pytest.mark.parametrize("value", [-1, 1.5, True]) +def test_invalid_max_retries(value: object) -> None: + with pytest.raises(AuraConfigurationError, match="max retries"): + _config(max_retries=value) + + +@pytest.mark.parametrize("value", [0, -1, 1.5]) +def test_invalid_max_response_size(value: object) -> None: + with pytest.raises(AuraConfigurationError, match="max response size"): + _config(max_response_size=value) + + +@pytest.mark.parametrize("value", ["", "agent\r\nX-Evil: 1"]) +def test_invalid_user_agent(value: str) -> None: + with pytest.raises(AuraConfigurationError, match="user agent"): + _config(user_agent=value) + + +def test_protected_default_headers_are_dropped() -> None: + config = _config( + default_headers={ + "authorization": "Bearer stolen", + "Content-Type": "text/plain", + "USER-AGENT": "other", + "X-Trace": "abc", + } + ) + assert config.default_headers == {"X-Trace": "abc"} + + +@pytest.mark.parametrize( + "headers", [{"": "v"}, {"X-A:B": "v"}, {"X-A": "line\nbreak"}, {"X-A": 1}, {1: "v"}] +) +def test_malformed_default_headers_rejected(headers: dict[Any, Any]) -> None: + with pytest.raises(AuraConfigurationError): + _config(default_headers=headers) diff --git a/tests/unit/test_errors.py b/tests/unit/test_errors.py new file mode 100644 index 0000000..b257ea8 --- /dev/null +++ b/tests/unit/test_errors.py @@ -0,0 +1,139 @@ +import json +from datetime import UTC, datetime, timedelta +from email.utils import format_datetime + +import pytest + +from aura_python_sdk import ( + AuraAPIError, + AuraError, + AuthenticationError, + BadRequestError, + ConflictError, + ErrorDetail, + NotFoundError, + PermissionDeniedError, + RateLimitError, + ServerError, +) +from aura_python_sdk._errors import api_error_from_response + + +def _body(payload: object) -> bytes: + return json.dumps(payload).encode() + + +@pytest.mark.parametrize( + ("status", "expected"), + [ + (400, BadRequestError), + (401, AuthenticationError), + (403, PermissionDeniedError), + (404, NotFoundError), + (409, ConflictError), + (429, RateLimitError), + (500, ServerError), + (503, ServerError), + (405, AuraAPIError), + (415, AuraAPIError), + (420, AuraAPIError), + ], +) +def test_status_maps_to_error_class(status: int, expected: type[AuraAPIError]) -> None: + error = api_error_from_response(status, b"", {}) + assert type(error) is expected + assert isinstance(error, AuraError) + assert error.status_code == status + + +def test_spec_errors_shape() -> None: + body = _body( + { + "errors": [ + {"message": "Instance not found", "reason": "instance-not-found"}, + {"message": "second", "reason": "other", "field": "name"}, + ] + } + ) + error = api_error_from_response(404, body, {"x-request-id": "req-123"}) + + assert error.message == "Not Found" + assert error.details == ( + ErrorDetail("Instance not found", "instance-not-found"), + ErrorDetail("second", "other", "name"), + ) + assert error.request_id == "req-123" + assert error.is_not_found + assert error.has_multiple_errors + assert error.all_errors() == ["Not Found", "Instance not found", "second"] + # Same format as the Go SDK's Error() string. + assert ( + str(error) == "API error (status 404): Not Found - Instance not found (and 1 more error(s))" + ) + + +def test_single_detail_message_format() -> None: + error = api_error_from_response(400, _body({"errors": [{"message": "bad name"}]}), {}) + assert str(error) == "API error (status 400): Bad Request - bad name" + assert error.is_bad_request + assert not error.has_multiple_errors + + +def test_message_and_details_keys() -> None: + body = _body({"message": "Validation failed", "details": [{"message": "memory invalid"}]}) + error = api_error_from_response(400, body, {}) + assert error.message == "Validation failed" + assert [d.message for d in error.details] == ["memory invalid"] + + +def test_middleware_error_shape() -> None: + error = api_error_from_response(429, _body({"error": "Rate limit exceeded"}), {}) + assert error.message == "Rate limit exceeded" + assert str(error) == "API error (status 429): Rate limit exceeded" + + +@pytest.mark.parametrize("body", [b"", b"not json", b"[1, 2]", b'"text"', _body({"errors": "x"})]) +def test_unparseable_bodies_fall_back_to_status_phrase(body: bytes) -> None: + error = api_error_from_response(502, body, {}) + assert error.message == "Bad Gateway" + assert error.details == () + + +def test_unknown_status_code_phrase() -> None: + assert api_error_from_response(420, b"", {}).message == "HTTP 420" + + +def test_retry_after_seconds() -> None: + error = api_error_from_response(429, b"", {"retry-after": "30"}) + assert isinstance(error, RateLimitError) + assert error.retry_after == 30.0 + + +def test_retry_after_http_date() -> None: + when = datetime.now(UTC) + timedelta(seconds=120) + error = api_error_from_response(429, b"", {"retry-after": format_datetime(when, usegmt=True)}) + assert isinstance(error, RateLimitError) + assert error.retry_after is not None + assert 100 < error.retry_after <= 120 + + +@pytest.mark.parametrize("value", [None, "", "soon"]) +def test_retry_after_missing_or_invalid(value: str | None) -> None: + headers = {} if value is None else {"retry-after": value} + error = api_error_from_response(429, b"", headers) + assert isinstance(error, RateLimitError) + assert error.retry_after is None + + +def test_error_class_override_keeps_details() -> None: + error = api_error_from_response( + 400, _body({"errors": [{"message": "invalid_client"}]}), {}, error_class=AuthenticationError + ) + assert type(error) is AuthenticationError + assert error.status_code == 400 + assert error.details[0].message == "invalid_client" + + +def test_is_unauthorized() -> None: + assert api_error_from_response(401, b"", {}).is_unauthorized + assert not api_error_from_response(403, b"", {}).is_unauthorized diff --git a/tests/unit/test_http_service.py b/tests/unit/test_http_service.py new file mode 100644 index 0000000..f5c157b --- /dev/null +++ b/tests/unit/test_http_service.py @@ -0,0 +1,147 @@ +import logging + +import pytest + +from aura_python_sdk import AuraConnectionError, AuraResponseError, AuraTimeoutError, HttpResponse +from aura_python_sdk._internal.http._service import HttpService +from tests.fakes import FakeClock, FakeTransport + +URL = "https://api.neo4j.io/v1/instances" + + +def _service( + transport: FakeTransport, clock: FakeClock, *, max_retries: int = 3, max_size: int = 1024 +) -> HttpService: + return HttpService( + transport, + max_retries=max_retries, + max_response_size=max_size, + logger=logging.getLogger("test"), + clock=clock, + sleep=clock.sleep, + ) + + +def _not_sent() -> AuraConnectionError: + return AuraConnectionError("connect failed", request_sent=False) + + +def _sent() -> AuraConnectionError: + return AuraConnectionError("connection reset", request_sent=True) + + +def test_passes_request_through() -> None: + clock = FakeClock() + transport = FakeTransport([HttpResponse(200, body=b"ok")]) + response = _service(transport, clock).send( + "POST", URL, {"X": "1"}, b"{}", deadline=clock.now + 30 + ) + + assert response.body == b"ok" + [request] = transport.requests + assert (request.method, request.url, request.body) == ("POST", URL, b"{}") + assert request.headers == {"X": "1"} + assert request.timeout == 30 + assert request.max_response_size == 1024 + + +@pytest.mark.parametrize("status", [429, 500, 502, 503, 504, 404]) +def test_http_status_responses_are_never_retried(status: int) -> None: + clock = FakeClock() + transport = FakeTransport([HttpResponse(status)]) + response = _service(transport, clock).send("GET", URL, {}, None, deadline=clock.now + 30) + assert response.status_code == status + assert len(transport.requests) == 1 + + +def test_retries_network_errors_with_backoff() -> None: + clock = FakeClock() + transport = FakeTransport([_sent(), _sent(), _sent(), HttpResponse(200)]) + response = _service(transport, clock).send("GET", URL, {}, None, deadline=clock.now + 60) + + assert response.status_code == 200 + assert len(transport.requests) == 4 + assert clock.sleeps == [1.0, 2.0, 4.0] + + +def test_backoff_is_capped_at_five_seconds() -> None: + clock = FakeClock() + transport = FakeTransport([_not_sent()] * 5 + [HttpResponse(200)]) + _service(transport, clock, max_retries=5).send("GET", URL, {}, None, deadline=clock.now + 60) + assert clock.sleeps == [1.0, 2.0, 4.0, 5.0, 5.0] + + +def test_gives_up_after_max_retries() -> None: + clock = FakeClock() + transport = FakeTransport([_not_sent()] * 3) + with pytest.raises(AuraConnectionError): + _service(transport, clock, max_retries=2).send( + "GET", URL, {}, None, deadline=clock.now + 60 + ) + assert len(transport.requests) == 3 + + +def test_zero_retries_means_single_attempt() -> None: + clock = FakeClock() + transport = FakeTransport([_not_sent()]) + with pytest.raises(AuraConnectionError): + _service(transport, clock, max_retries=0).send( + "GET", URL, {}, None, deadline=clock.now + 60 + ) + assert len(transport.requests) == 1 + + +@pytest.mark.parametrize("method", ["POST", "PATCH"]) +def test_non_idempotent_request_not_retried_once_it_may_have_been_sent(method: str) -> None: + clock = FakeClock() + transport = FakeTransport([_sent()]) + with pytest.raises(AuraConnectionError): + _service(transport, clock).send(method, URL, {}, b"{}", deadline=clock.now + 60) + assert len(transport.requests) == 1 + + +def test_non_idempotent_request_retried_when_never_sent() -> None: + clock = FakeClock() + transport = FakeTransport([_not_sent(), HttpResponse(202)]) + response = _service(transport, clock).send("POST", URL, {}, b"{}", deadline=clock.now + 60) + assert response.status_code == 202 + assert len(transport.requests) == 2 + + +def test_timeouts_are_retried_for_idempotent_methods() -> None: + clock = FakeClock() + transport = FakeTransport( + [AuraTimeoutError("read timeout", request_sent=True), HttpResponse(200)] + ) + _service(transport, clock).send("DELETE", URL, {}, None, deadline=clock.now + 60) + assert len(transport.requests) == 2 + + +def test_no_retry_when_backoff_would_pass_deadline() -> None: + clock = FakeClock() + transport = FakeTransport([_not_sent(), HttpResponse(200)]) + with pytest.raises(AuraConnectionError): + _service(transport, clock).send("GET", URL, {}, None, deadline=clock.now + 0.5) + assert len(transport.requests) == 1 + + +def test_each_attempt_gets_the_remaining_time() -> None: + clock = FakeClock() + transport = FakeTransport([_not_sent(), HttpResponse(200)]) + _service(transport, clock).send("GET", URL, {}, None, deadline=clock.now + 10) + assert [r.timeout for r in transport.requests] == [10.0, 9.0] + + +def test_expired_deadline_raises_timeout_without_sending() -> None: + clock = FakeClock() + transport = FakeTransport() + with pytest.raises(AuraTimeoutError): + _service(transport, clock).send("GET", URL, {}, None, deadline=clock.now) + assert transport.requests == [] + + +def test_oversized_body_rejected_even_if_transport_ignores_limit() -> None: + clock = FakeClock() + transport = FakeTransport([HttpResponse(200, body=b"x" * 11)]) + with pytest.raises(AuraResponseError, match="exceeded limit"): + _service(transport, clock, max_size=10).send("GET", URL, {}, None, deadline=clock.now + 5) diff --git a/tests/unit/test_request_service.py b/tests/unit/test_request_service.py new file mode 100644 index 0000000..8755d32 --- /dev/null +++ b/tests/unit/test_request_service.py @@ -0,0 +1,174 @@ +import json +import logging + +import pytest + +from aura_python_sdk import AuraResponseError, AuthenticationError, HttpResponse, NotFoundError +from aura_python_sdk._internal._auth import TokenManager +from aura_python_sdk._internal._request import RequestService, build_path +from aura_python_sdk._internal.http._service import HttpService +from tests.fakes import FakeClock, FakeTransport, json_response, token_response + +BASE = "https://api.neo4j.io" + + +def _service( + transport: FakeTransport, + *, + timeout: float = 30.0, + default_headers: dict[str, str] | None = None, +) -> RequestService: + clock = FakeClock() + logger = logging.getLogger("test") + http = HttpService( + transport, + max_retries=0, + max_response_size=1 << 20, + logger=logger, + clock=clock, + sleep=clock.sleep, + ) + auth = TokenManager( + client_id="id", + client_secret="secret", + token_url=f"{BASE}/oauth/token", + user_agent="ua/1", + http=http, + logger=logger, + ) + return RequestService( + http=http, + auth=auth, + base_url=BASE, + api_version="v1", + user_agent="ua/1", + default_headers=default_headers or {}, + timeout=timeout, + logger=logger, + ) + + +def test_relative_path_gets_versioned_base_url() -> None: + transport = FakeTransport([token_response(), json_response(200, {"data": []})]) + _service(transport).get("instances") + assert transport.api_requests[0].url == f"{BASE}/v1/instances" + + +def test_leading_slash_is_tolerated() -> None: + transport = FakeTransport([token_response(), json_response(200, {})]) + _service(transport).get("/tenants") + assert transport.api_requests[0].url == f"{BASE}/v1/tenants" + + +def test_absolute_url_passes_through_with_auth() -> None: + prometheus = "https://abc.metrics.neo4j.io/prometheus/metrics" + transport = FakeTransport([token_response("tok"), HttpResponse(200, body=b"metric 1")]) + response = _service(transport).get(prometheus) + + [request] = transport.api_requests + assert request.url == prometheus + assert request.headers["Authorization"] == "Bearer tok" + assert response.body == b"metric 1" + + +def test_query_params_encoded_and_none_dropped() -> None: + transport = FakeTransport([token_response(), json_response(200, {}), json_response(200, {})]) + service = _service(transport) + service.get("customer-managed-keys", params={"tenantId": "a b&c", "other": None}) + service.get("x?y=1", params={"z": "2"}) + urls = [r.url for r in transport.api_requests] + assert urls == [f"{BASE}/v1/customer-managed-keys?tenantId=a+b%26c", f"{BASE}/v1/x?y=1&z=2"] + + +def test_headers() -> None: + transport = FakeTransport([token_response("tok"), json_response(200, {})]) + _service(transport, default_headers={"X-Trace": "t1"}).get("instances") + headers = transport.api_requests[0].headers + assert headers == { + "X-Trace": "t1", + "Content-Type": "application/json", + "User-Agent": "ua/1", + "Authorization": "Bearer tok", + } + + +def test_json_body_is_serialised() -> None: + transport = FakeTransport([token_response(), json_response(202, {"data": {}})]) + response = _service(transport).post("instances", json_body={"name": "Instance01", "n": 1}) + request = transport.api_requests[0] + assert request.method == "POST" + assert request.body is not None + assert json.loads(request.body) == {"name": "Instance01", "n": 1} + assert response.status_code == 202 + + +def test_no_body_when_json_body_is_none() -> None: + transport = FakeTransport([token_response(), json_response(202, {})]) + _service(transport).post("instances/abcd1234/pause") + assert transport.api_requests[0].body is None + + +@pytest.mark.parametrize( + ("method", "call"), + [ + ("GET", lambda s: s.get("p")), + ("POST", lambda s: s.post("p")), + ("PATCH", lambda s: s.patch("p", json_body={})), + ("PUT", lambda s: s.put("p", json_body={})), + ("DELETE", lambda s: s.delete("p")), + ], +) +def test_verbs(method: str, call: object) -> None: + transport = FakeTransport([token_response(), json_response(200, {})]) + call(_service(transport)) # type: ignore[operator] + assert transport.api_requests[0].method == method + + +def test_error_response_raises_mapped_exception() -> None: + body = {"errors": [{"message": "Instance not found", "reason": "instance-not-found"}]} + transport = FakeTransport([token_response(), json_response(404, body, {"X-Request-Id": "r1"})]) + with pytest.raises(NotFoundError) as info: + _service(transport).get("instances/abcd1234") + assert info.value.request_id == "r1" + assert info.value.details[0].message == "Instance not found" + + +def test_401_invalidates_cached_token() -> None: + transport = FakeTransport( + [ + token_response("old"), + json_response(401, {"errors": [{"message": "expired"}]}), + token_response("new"), + json_response(200, {}), + ] + ) + service = _service(transport) + with pytest.raises(AuthenticationError): + service.get("instances") + service.get("instances") + assert [r.headers["Authorization"] for r in transport.api_requests] == [ + "Bearer old", + "Bearer new", + ] + + +def test_token_fetch_shares_the_call_deadline() -> None: + transport = FakeTransport([token_response(), json_response(200, {})]) + _service(transport, timeout=12.0).get("instances") + assert [r.timeout for r in transport.requests] == [12.0, 12.0] + + +def test_response_json() -> None: + transport = FakeTransport([token_response(), json_response(200, {"data": [1]})]) + assert _service(transport).get("x").json() == {"data": [1]} + + +def test_response_json_invalid() -> None: + transport = FakeTransport([token_response(), HttpResponse(200, body=b"")]) + with pytest.raises(AuraResponseError, match="not valid JSON"): + _service(transport).get("x").json() + + +def test_build_path_encodes_segments() -> None: + assert build_path("instances", "abcd1234", "snapshots") == "instances/abcd1234/snapshots" + assert build_path("sessions", "../x?y") == "sessions/..%2Fx%3Fy" From 51c560753155904283b300e988d86e7023c08213 Mon Sep 17 00:00:00 2001 From: Jonathan Giffard Date: Tue, 29 Sep 2026 20:08:24 +0100 Subject: [PATCH 2/8] Add v1 data models and JSON conversion Phase 3 of PLAN.md. Adds frozen dataclass models for every v1 request and response, StrEnums with tolerant parsing, a stdlib-only serde module, and a test that parses every 2xx example in the OpenAPI spec with its model. Co-Authored-By: Claude Opus 5.5 --- PLAN.md | 6 + pyproject.toml | 2 + src/aura_python_sdk/__init__.py | 50 ++++ src/aura_python_sdk/_internal/_serde.py | 212 +++++++++++++++ src/aura_python_sdk/models/__init__.py | 59 +++++ src/aura_python_sdk/models/_common.py | 22 ++ src/aura_python_sdk/models/cmek.py | 36 +++ src/aura_python_sdk/models/graph_analytics.py | 74 ++++++ src/aura_python_sdk/models/instances.py | 128 +++++++++ src/aura_python_sdk/models/snapshots.py | 42 +++ src/aura_python_sdk/models/tenants.py | 44 ++++ tests/unit/test_models.py | 110 ++++++++ tests/unit/test_serde.py | 244 ++++++++++++++++++ tests/unit/test_spec_examples.py | 116 +++++++++ uv.lock | 68 +++++ 15 files changed, 1213 insertions(+) create mode 100644 src/aura_python_sdk/_internal/_serde.py create mode 100644 src/aura_python_sdk/models/__init__.py create mode 100644 src/aura_python_sdk/models/_common.py create mode 100644 src/aura_python_sdk/models/cmek.py create mode 100644 src/aura_python_sdk/models/graph_analytics.py create mode 100644 src/aura_python_sdk/models/instances.py create mode 100644 src/aura_python_sdk/models/snapshots.py create mode 100644 src/aura_python_sdk/models/tenants.py create mode 100644 tests/unit/test_models.py create mode 100644 tests/unit/test_serde.py create mode 100644 tests/unit/test_spec_examples.py diff --git a/PLAN.md b/PLAN.md index f0ff5bc..f0d7193 100644 --- a/PLAN.md +++ b/PLAN.md @@ -288,3 +288,9 @@ packages. not its schema. Go sends them, so we keep them. - **Instance status `stopped` / `available`**: present in Go but not in the spec enum. Keep them for parity; tolerant parsing makes this harmless. +- **Snapshot ID format**: resolved. Snapshot IDs are UUIDs, and the spec's list example + (`snapshot_id: '2023-01-20T13:44:42Z'`) is wrong. We keep Go's UUID validation. +- **Required fields on responses**: models follow the spec's `required` lists, with two + exceptions. Instance `storage` is optional because it isn't returned for Free instances. GDS + session `status` is optional because the spec's 202 example returns `null`. A missing required + field raises `AuraResponseError` and names the field. diff --git a/pyproject.toml b/pyproject.toml index 3c65226..ec0b9ac 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -37,7 +37,9 @@ dev = [ "mypy>=1.11", "pytest>=8", "pytest-cov>=5", + "pyyaml>=6.0.3", "ruff>=0.6", + "types-pyyaml>=6.0.12.20260906", ] [tool.hatch.version] diff --git a/src/aura_python_sdk/__init__.py b/src/aura_python_sdk/__init__.py index 7d6077e..ee2c8f2 100644 --- a/src/aura_python_sdk/__init__.py +++ b/src/aura_python_sdk/__init__.py @@ -31,6 +31,32 @@ ) from aura_python_sdk._transport import HttpRequest, HttpResponse, HttpTransport from aura_python_sdk._version import __version__ +from aura_python_sdk.models import ( + CDCEnrichmentMode, + CloudProvider, + CreatedInstance, + CreatedSnapshot, + CustomerManagedKey, + CustomerManagedKeySummary, + DeletedGDSSession, + GDSSession, + GDSSessionConfig, + GDSSessionSizeEstimate, + GDSSessionStatus, + Instance, + InstanceConfig, + InstanceConfiguration, + InstanceSizeEstimate, + InstanceStatus, + InstanceSummary, + InstanceType, + MetricsIntegration, + Snapshot, + SnapshotProfile, + SnapshotStatus, + Tenant, + TenantSummary, +) # Library convention: emit nothing unless the application configures logging. logging.getLogger(__name__).addHandler(logging.NullHandler()) @@ -46,14 +72,38 @@ "AuraValidationError", "AuthenticationError", "BadRequestError", + "CDCEnrichmentMode", + "CloudProvider", "ConflictError", + "CreatedInstance", + "CreatedSnapshot", + "CustomerManagedKey", + "CustomerManagedKeySummary", + "DeletedGDSSession", "ErrorDetail", + "GDSSession", + "GDSSessionConfig", + "GDSSessionSizeEstimate", + "GDSSessionStatus", "HttpRequest", "HttpResponse", "HttpTransport", + "Instance", + "InstanceConfig", + "InstanceConfiguration", + "InstanceSizeEstimate", + "InstanceStatus", + "InstanceSummary", + "InstanceType", + "MetricsIntegration", "NotFoundError", "PermissionDeniedError", "RateLimitError", "ServerError", + "Snapshot", + "SnapshotProfile", + "SnapshotStatus", + "Tenant", + "TenantSummary", "__version__", ] diff --git a/src/aura_python_sdk/_internal/_serde.py b/src/aura_python_sdk/_internal/_serde.py new file mode 100644 index 0000000..24a08dd --- /dev/null +++ b/src/aura_python_sdk/_internal/_serde.py @@ -0,0 +1,212 @@ +"""Conversion between JSON values and the SDK's dataclass models. + +Field names match the JSON keys. Conversion is driven by each field's type hint: + +- ``X | None`` fields accept a missing key or ``null``. Other fields without a default are + required, and a missing key raises :class:`AuraResponseError`. +- ``SomeEnum | str`` fields hold the enum member when the value is known and the raw string + otherwise, so a new status from the API never breaks parsing. +- Unknown JSON keys are ignored. +- Small spec inconsistencies are tolerated: a numeric string for an ``int`` field, or a number + for a ``str`` field. +""" + +from __future__ import annotations + +import dataclasses +import re +import types +import typing +from collections.abc import Mapping +from datetime import date, datetime +from enum import Enum +from typing import Any, TypeVar, cast + +from aura_python_sdk._errors import AuraResponseError + +T = TypeVar("T") + +_NONE_TYPE = type(None) +# Python 3.11's fromisoformat accepts at most 6 fractional digits; the API (Go) may send 9. +_EXCESS_FRACTION = re.compile(r"(\.\d{6})\d+") + + +class _MismatchError(Exception): + def __init__(self, path: str, message: str) -> None: + super().__init__(f"{path or ''}: {message}") + + +def from_json(cls: type[T], value: object) -> T: + """Build ``cls`` (a dataclass) from a decoded JSON value.""" + try: + return cast(T, _convert(cls, value, "")) + except _MismatchError as exc: + raise AuraResponseError(f"unexpected response shape at {exc}") from None + + +def parse_data(cls: type[T], payload: object) -> T: + """Unwrap a ``{"data": {...}}`` response into ``cls``.""" + return from_json(cls, _data(payload)) + + +def parse_data_list(cls: type[T], payload: object) -> list[T]: + """Unwrap a ``{"data": [...]}`` response into a list of ``cls``.""" + data = _data(payload) + if not isinstance(data, list): + raise AuraResponseError("unexpected response shape at data: expected a list") + return [from_json(cls, item) for item in data] + + +def _data(payload: object) -> object: + if not isinstance(payload, Mapping) or "data" not in payload: + raise AuraResponseError("unexpected response shape: missing 'data'") + return payload["data"] + + +_FieldSpec = tuple[str, Any, bool] +_FIELD_CACHE: dict[type[Any], tuple[_FieldSpec, ...]] = {} + + +def _fields(cls: type[Any]) -> tuple[_FieldSpec, ...]: + """(name, resolved type, required) for each init field of a dataclass.""" + cached = _FIELD_CACHE.get(cls) + if cached is None: + hints = typing.get_type_hints(cls) + cached = tuple( + ( + f.name, + hints[f.name], + f.default is dataclasses.MISSING and f.default_factory is dataclasses.MISSING, + ) + for f in dataclasses.fields(cls) + if f.init + ) + _FIELD_CACHE[cls] = cached + return cached + + +def _convert(tp: Any, value: object, path: str) -> object: + origin = typing.get_origin(tp) + + if origin is typing.Union or origin is types.UnionType: + return _convert_union(typing.get_args(tp), value, path) + if value is None: + raise _MismatchError(path, "value must not be null") + if origin is tuple: + item_type = typing.get_args(tp)[0] + return tuple( + _convert(item_type, v, f"{path}[{i}]") for i, v in enumerate(_list(value, path)) + ) + if origin is list: + item_type = typing.get_args(tp)[0] + return [_convert(item_type, v, f"{path}[{i}]") for i, v in enumerate(_list(value, path))] + if dataclasses.is_dataclass(tp) and isinstance(tp, type): + return _convert_dataclass(tp, value, path) + if isinstance(tp, type) and issubclass(tp, Enum): + try: + return tp(value) + except ValueError: + raise _MismatchError(path, f"{value!r} is not a valid {tp.__name__}") from None + if tp is bool: + if isinstance(value, bool): + return value + raise _MismatchError(path, f"expected a boolean, got {value!r}") + if tp is int: + return _to_int(value, path) + if tp is float: + if isinstance(value, int | float) and not isinstance(value, bool): + return float(value) + raise _MismatchError(path, f"expected a number, got {value!r}") + if tp is str: + if isinstance(value, str): + return value + if isinstance(value, int | float) and not isinstance(value, bool): + return str(value) + raise _MismatchError(path, f"expected a string, got {value!r}") + if tp is datetime: + return _to_datetime(value, path) + if tp is date: + if isinstance(value, str): + try: + return date.fromisoformat(value) + except ValueError: + pass + raise _MismatchError(path, f"expected an ISO date, got {value!r}") + raise TypeError(f"unsupported model field type {tp!r} at {path}") + + +def _convert_union(args: tuple[Any, ...], value: object, path: str) -> object: + if value is None: + if _NONE_TYPE in args: + return None + raise _MismatchError(path, "value must not be null") + candidates = [a for a in args if a is not _NONE_TYPE] + # Optional timestamps: treat an empty string like null. + if value == "" and _NONE_TYPE in args and all(a in (datetime, date) for a in candidates): + return None + last_error: _MismatchError | None = None + for candidate in candidates: + try: + return _convert(candidate, value, path) + except _MismatchError as exc: + last_error = exc + raise last_error or _MismatchError(path, "no matching type") + + +def _convert_dataclass(cls: type[Any], value: object, path: str) -> object: + if not isinstance(value, Mapping): + raise _MismatchError(path, f"expected an object, got {type(value).__name__}") + kwargs: dict[str, object] = {} + for name, field_type, required in _fields(cls): + field_path = f"{path}.{name}" if path else name + if name in value: + kwargs[name] = _convert(field_type, value[name], field_path) + elif required: + raise _MismatchError(field_path, "required field is missing") + return cls(**kwargs) + + +def _list(value: object, path: str) -> list[object]: + if not isinstance(value, list): + raise _MismatchError(path, f"expected a list, got {type(value).__name__}") + return value + + +def _to_int(value: object, path: str) -> int: + if isinstance(value, bool): + raise _MismatchError(path, f"expected an integer, got {value!r}") + if isinstance(value, int): + return value + if isinstance(value, float) and value.is_integer(): + return int(value) + if isinstance(value, str) and value.strip().lstrip("-").isdigit(): + return int(value) + raise _MismatchError(path, f"expected an integer, got {value!r}") + + +def _to_datetime(value: object, path: str) -> datetime: + if isinstance(value, str): + try: + return datetime.fromisoformat(_EXCESS_FRACTION.sub(r"\1", value)) + except ValueError: + pass + raise _MismatchError(path, f"expected an ISO 8601 timestamp, got {value!r}") + + +def to_json(obj: object) -> object: + """Convert a request model (or plain values) to JSON-ready data, omitting None fields.""" + if dataclasses.is_dataclass(obj) and not isinstance(obj, type): + return { + f.name: to_json(getattr(obj, f.name)) + for f in dataclasses.fields(obj) + if getattr(obj, f.name) is not None + } + if isinstance(obj, Mapping): + return {str(k): to_json(v) for k, v in obj.items() if v is not None} + if isinstance(obj, Enum): + return obj.value + if isinstance(obj, datetime | date): + return obj.isoformat() + if isinstance(obj, list | tuple): + return [to_json(v) for v in obj] + return obj diff --git a/src/aura_python_sdk/models/__init__.py b/src/aura_python_sdk/models/__init__.py new file mode 100644 index 0000000..089755d --- /dev/null +++ b/src/aura_python_sdk/models/__init__.py @@ -0,0 +1,59 @@ +"""Data models returned by and passed to the Aura API v1.""" + +from aura_python_sdk.models._common import CloudProvider, InstanceType +from aura_python_sdk.models.cmek import CustomerManagedKey, CustomerManagedKeySummary +from aura_python_sdk.models.graph_analytics import ( + DeletedGDSSession, + GDSSession, + GDSSessionConfig, + GDSSessionSizeEstimate, + GDSSessionStatus, +) +from aura_python_sdk.models.instances import ( + CDCEnrichmentMode, + CreatedInstance, + Instance, + InstanceConfig, + InstanceSizeEstimate, + InstanceStatus, + InstanceSummary, +) +from aura_python_sdk.models.snapshots import ( + CreatedSnapshot, + Snapshot, + SnapshotProfile, + SnapshotStatus, +) +from aura_python_sdk.models.tenants import ( + InstanceConfiguration, + MetricsIntegration, + Tenant, + TenantSummary, +) + +__all__ = [ + "CDCEnrichmentMode", + "CloudProvider", + "CreatedInstance", + "CreatedSnapshot", + "CustomerManagedKey", + "CustomerManagedKeySummary", + "DeletedGDSSession", + "GDSSession", + "GDSSessionConfig", + "GDSSessionSizeEstimate", + "GDSSessionStatus", + "Instance", + "InstanceConfig", + "InstanceConfiguration", + "InstanceSizeEstimate", + "InstanceStatus", + "InstanceSummary", + "InstanceType", + "MetricsIntegration", + "Snapshot", + "SnapshotProfile", + "SnapshotStatus", + "Tenant", + "TenantSummary", +] diff --git a/src/aura_python_sdk/models/_common.py b/src/aura_python_sdk/models/_common.py new file mode 100644 index 0000000..99d185e --- /dev/null +++ b/src/aura_python_sdk/models/_common.py @@ -0,0 +1,22 @@ +"""Enums shared across API areas.""" + +from __future__ import annotations + +from enum import StrEnum + + +class CloudProvider(StrEnum): + GCP = "gcp" + AWS = "aws" + AZURE = "azure" + + +class InstanceType(StrEnum): + """Instance types. ``ENTERPRISE_DB`` is AuraDB Virtual Dedicated Cloud.""" + + ENTERPRISE_DB = "enterprise-db" + ENTERPRISE_DS = "enterprise-ds" + BUSINESS_CRITICAL = "business-critical" + PROFESSIONAL_DB = "professional-db" + PROFESSIONAL_DS = "professional-ds" + FREE_DB = "free-db" diff --git a/src/aura_python_sdk/models/cmek.py b/src/aura_python_sdk/models/cmek.py new file mode 100644 index 0000000..80a7cfd --- /dev/null +++ b/src/aura_python_sdk/models/cmek.py @@ -0,0 +1,36 @@ +"""Customer-managed encryption key models (Go: cmek.go).""" + +from __future__ import annotations + +from dataclasses import dataclass +from datetime import datetime + +from aura_python_sdk.models._common import CloudProvider, InstanceType + + +@dataclass(frozen=True, slots=True, kw_only=True) +class CustomerManagedKeySummary: + """A key as returned by ``GET /customer-managed-keys``.""" + + id: str + name: str + tenant_id: str + + +@dataclass(frozen=True, slots=True, kw_only=True) +class CustomerManagedKey: + """Full details of a customer-managed key. + + ``key_id`` is the key's ID in your cloud provider (the key ARN on AWS). The key can only encrypt + instances of ``instance_type`` in ``region``. + """ + + id: str + name: str + tenant_id: str + cloud_provider: CloudProvider | str + region: str + instance_type: InstanceType | str + key_id: str + status: str + created: datetime | None = None diff --git a/src/aura_python_sdk/models/graph_analytics.py b/src/aura_python_sdk/models/graph_analytics.py new file mode 100644 index 0000000..124c7fc --- /dev/null +++ b/src/aura_python_sdk/models/graph_analytics.py @@ -0,0 +1,74 @@ +"""Graph Analytics (GDS) session models (Go: graphanalytics.go).""" + +from __future__ import annotations + +from dataclasses import dataclass +from datetime import datetime +from enum import StrEnum + +from aura_python_sdk.models._common import CloudProvider + + +class GDSSessionStatus(StrEnum): + CREATING = "Creating" + READY = "Ready" + EXPIRED = "Expired" + FAILED = "Failed" + + +@dataclass(frozen=True, slots=True, kw_only=True) +class GDSSession: + """A Graph Analytics session. + + ``instance_id`` and ``database_uuid`` are empty for a standalone session. ``ttl`` is a + duration string such as ``"20m0s"``. + """ + + id: str + name: str + memory: str + host: str + tenant_id: str + user_id: str + status: GDSSessionStatus | str | None = None + instance_id: str | None = None + database_uuid: str | None = None + cloud_provider: CloudProvider | str | None = None + region: str | None = None + created_at: datetime | None = None + expiry_date: datetime | None = None + ttl: str | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class GDSSessionConfig: + """Settings for a new session (Go: ``CreateGDSSessionConfigData``). + + Set ``instance_id`` and ``database_uuid`` to attach the session to an AuraDB instance, or + ``cloud_provider`` and ``region`` for a standalone session. ``ttl`` is a duration string + such as ``"1h"``. + """ + + name: str + memory: str + tenant_id: str | None = None + ttl: str | None = None + instance_id: str | None = None + database_uuid: str | None = None + cloud_provider: CloudProvider | str | None = None + region: str | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class GDSSessionSizeEstimate: + """Result of ``POST /graph-analytics/sessions/sizing``.""" + + estimated_memory: str + recommended_size: str + + +@dataclass(frozen=True, slots=True, kw_only=True) +class DeletedGDSSession: + """Returned when a session is deleted.""" + + id: str diff --git a/src/aura_python_sdk/models/instances.py b/src/aura_python_sdk/models/instances.py new file mode 100644 index 0000000..f940115 --- /dev/null +++ b/src/aura_python_sdk/models/instances.py @@ -0,0 +1,128 @@ +"""Instance models (Go: instances.go).""" + +from __future__ import annotations + +from dataclasses import dataclass, field +from datetime import datetime +from enum import StrEnum + +from aura_python_sdk.models._common import CloudProvider, InstanceType + + +class InstanceStatus(StrEnum): + """Lifecycle states of an instance.""" + + CREATING = "creating" + DESTROYING = "destroying" + RUNNING = "running" + PAUSING = "pausing" + PAUSED = "paused" + SUSPENDING = "suspending" + SUSPENDED = "suspended" + RESUMING = "resuming" + LOADING = "loading" + LOADING_FAILED = "loading failed" + RESTORING = "restoring" + UPDATING = "updating" + OVERWRITING = "overwriting" + # Not in the v1 spec's enum, but defined by the Go SDK. + STOPPED = "stopped" + AVAILABLE = "available" + + +class CDCEnrichmentMode(StrEnum): + OFF = "OFF" + DIFF = "DIFF" + FULL = "FULL" + + +@dataclass(frozen=True, slots=True, kw_only=True) +class InstanceSummary: + """An instance as returned by ``GET /instances``.""" + + id: str + name: str + tenant_id: str + cloud_provider: CloudProvider | str + created_at: datetime | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class Instance: + """Full details of an instance. + + ``storage`` is not returned for AuraDB Free. ``graph_nodes`` and ``graph_relationships`` are + returned only for Free instances. ``secondaries_count`` is returned only for Virtual + Dedicated Cloud, and ``cdc_enrichment_mode`` only for Virtual Dedicated Cloud and Business + Critical. + """ + + id: str + name: str + status: InstanceStatus | str + tenant_id: str + cloud_provider: CloudProvider | str + connection_url: str + region: str + type: InstanceType | str + memory: str + storage: str | None = None + created_at: datetime | None = None + metrics_integration_url: str | None = None + customer_managed_key_id: str | None = None + graph_nodes: int | None = None + graph_relationships: int | None = None + secondaries_count: int | None = None + cdc_enrichment_mode: CDCEnrichmentMode | str | None = None + vector_optimized: bool | None = None + graph_analytics_plugin: bool | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class CreatedInstance: + """Returned when an instance is created, including its initial credentials. + + ``password`` is shown only once and is left out of ``repr()``. Store it securely. + """ + + id: str + name: str + tenant_id: str + cloud_provider: CloudProvider | str + region: str + type: InstanceType | str + connection_url: str + username: str + password: str = field(repr=False) + created_at: datetime | None = None + vector_optimized: bool | None = None + graph_analytics_plugin: bool | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class InstanceConfig: + """Settings for a new instance (Go: ``CreateInstanceConfigData``). + + Valid combinations of cloud provider, region, type, version and memory for a tenant come from + ``client.tenants.get(tenant_id).instance_configurations``. + """ + + name: str + tenant_id: str + cloud_provider: CloudProvider | str + region: str + type: InstanceType | str + version: str + memory: str + vector_optimized: bool | None = None + graph_analytics_plugin: bool | None = None + customer_managed_key_id: str | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class InstanceSizeEstimate: + """Result of ``POST /instances/sizing``.""" + + recommended_size: str + min_required_memory: str + did_exceed_maximum: bool diff --git a/src/aura_python_sdk/models/snapshots.py b/src/aura_python_sdk/models/snapshots.py new file mode 100644 index 0000000..5bef20d --- /dev/null +++ b/src/aura_python_sdk/models/snapshots.py @@ -0,0 +1,42 @@ +"""Snapshot models (Go: snapshots.go).""" + +from __future__ import annotations + +from dataclasses import dataclass +from datetime import datetime +from enum import StrEnum + + +class SnapshotStatus(StrEnum): + COMPLETED = "Completed" + IN_PROGRESS = "InProgress" + FAILED = "Failed" + PENDING = "Pending" + CANCELLED = "Cancelled" + + +class SnapshotProfile(StrEnum): + AD_HOC = "AdHoc" + SCHEDULED = "Scheduled" + + +@dataclass(frozen=True, slots=True, kw_only=True) +class Snapshot: + """A snapshot of an instance. + + Only snapshots with ``exportable`` set can be used to create a new instance. + """ + + snapshot_id: str + instance_id: str + status: SnapshotStatus | str + profile: SnapshotProfile | str | None = None + timestamp: datetime | None = None + exportable: bool | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class CreatedSnapshot: + """Returned when an on-demand snapshot is started.""" + + snapshot_id: str diff --git a/src/aura_python_sdk/models/tenants.py b/src/aura_python_sdk/models/tenants.py new file mode 100644 index 0000000..b97cc39 --- /dev/null +++ b/src/aura_python_sdk/models/tenants.py @@ -0,0 +1,44 @@ +"""Tenant (project) models (Go: tenants.go).""" + +from __future__ import annotations + +from dataclasses import dataclass + +from aura_python_sdk.models._common import CloudProvider, InstanceType + + +@dataclass(frozen=True, slots=True, kw_only=True) +class TenantSummary: + """A tenant as returned by ``GET /tenants``.""" + + id: str + name: str + + +@dataclass(frozen=True, slots=True, kw_only=True) +class InstanceConfiguration: + """An instance configuration the tenant is allowed to create.""" + + cloud_provider: CloudProvider | str + region: str + region_name: str + type: InstanceType | str + memory: str + version: str + storage: str | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class Tenant: + """A tenant and the instance configurations available to it (``GET /tenants/{id}``).""" + + id: str + name: str + instance_configurations: tuple[InstanceConfiguration, ...] = () + + +@dataclass(frozen=True, slots=True, kw_only=True) +class MetricsIntegration: + """The project-level Prometheus metrics endpoint.""" + + endpoint: str diff --git a/tests/unit/test_models.py b/tests/unit/test_models.py new file mode 100644 index 0000000..1532f3b --- /dev/null +++ b/tests/unit/test_models.py @@ -0,0 +1,110 @@ +import dataclasses + +import pytest + +import aura_python_sdk as aura +from aura_python_sdk import models +from aura_python_sdk._internal._serde import from_json, to_json + +CREATED = { + "id": "db1d1234", + "name": "Instance01", + "tenant_id": "t", + "cloud_provider": "gcp", + "region": "europe-west1", + "type": "enterprise-db", + "connection_url": "neo4j+s://db1d1234.databases.neo4j.io", + "username": "neo4j", + "password": "letMeIn123!", +} + + +def test_created_instance_password_is_redacted_from_repr() -> None: + created = from_json(models.CreatedInstance, CREATED) + assert created.password == "letMeIn123!" + assert "letMeIn123!" not in repr(created) + assert "letMeIn123!" not in str(created) + + +def test_models_are_frozen() -> None: + summary = models.TenantSummary(id="t", name="n") + with pytest.raises(dataclasses.FrozenInstanceError): + summary.name = "other" # type: ignore[misc] + + +def test_models_are_keyword_only() -> None: + with pytest.raises(TypeError): + models.TenantSummary("t", "n") # type: ignore[call-arg] + + +def test_instance_status_covers_spec_and_go_values() -> None: + spec_values = { + "creating", "destroying", "running", "pausing", "paused", "suspending", "suspended", + "resuming", "loading", "loading failed", "restoring", "updating", "overwriting", + } # fmt: skip + assert {s.value for s in models.InstanceStatus} == spec_values | {"stopped", "available"} + + +def test_free_instance_without_storage_and_with_graph_counts() -> None: + instance = from_json( + models.Instance, + { + "id": "abcd1234", + "name": "Free", + "status": "running", + "tenant_id": "t", + "cloud_provider": "gcp", + "connection_url": "neo4j+s://abcd1234.databases.neo4j.io", + "region": "europe-west1", + "type": "free-db", + "memory": "1GB", + "graph_nodes": "1234", + "graph_relationships": "5678", + }, + ) + assert instance.storage is None + assert instance.type is models.InstanceType.FREE_DB + assert (instance.graph_nodes, instance.graph_relationships) == (1234, 5678) + + +def test_gds_ttl_integer_is_accepted_as_string() -> None: + session = from_json( + models.GDSSession, + {"id": "s", "name": "n", "memory": "8GB", "host": "h", "tenant_id": "t", "user_id": "u", + "ttl": 3600}, + ) # fmt: skip + assert session.ttl == "3600" + + +def test_instance_config_to_json_omits_unset_options() -> None: + config = models.InstanceConfig( + name="Instance01", + tenant_id="t", + cloud_provider=models.CloudProvider.GCP, + region="europe-west1", + type=models.InstanceType.ENTERPRISE_DB, + version="5", + memory="8GB", + vector_optimized=False, + ) + assert to_json(config) == { + "name": "Instance01", + "tenant_id": "t", + "cloud_provider": "gcp", + "region": "europe-west1", + "type": "enterprise-db", + "version": "5", + "memory": "8GB", + "vector_optimized": False, + } + + +def test_gds_session_config_to_json() -> None: + config = models.GDSSessionConfig(name="s", memory="8GB", ttl="1h", cloud_provider="aws") + assert to_json(config) == {"name": "s", "memory": "8GB", "ttl": "1h", "cloud_provider": "aws"} + + +def test_all_models_exported_at_top_level() -> None: + for name in models.__all__: + assert getattr(aura, name) is getattr(models, name) + assert name in aura.__all__ diff --git a/tests/unit/test_serde.py b/tests/unit/test_serde.py new file mode 100644 index 0000000..cdc3ca1 --- /dev/null +++ b/tests/unit/test_serde.py @@ -0,0 +1,244 @@ +from __future__ import annotations + +from dataclasses import dataclass, field +from datetime import UTC, date, datetime +from enum import StrEnum + +import pytest + +from aura_python_sdk import AuraResponseError +from aura_python_sdk._internal._serde import from_json, parse_data, parse_data_list, to_json + + +class Colour(StrEnum): + RED = "red" + BLUE = "blue" + + +@dataclass(frozen=True, slots=True, kw_only=True) +class Child: + name: str + + +@dataclass(frozen=True, slots=True, kw_only=True) +class Sample: + id: str + count: int + ratio: float + flag: bool + colour: Colour | str + strict_colour: Colour | None = None + when: datetime | None = None + day: date | None = None + note: str | None = None + children: tuple[Child, ...] = () + tags: list[str] = field(default_factory=list) + + +BASE = {"id": "a", "count": 1, "ratio": 0.5, "flag": True, "colour": "red"} + + +def _sample(**overrides: object) -> Sample: + return from_json(Sample, {**BASE, **overrides}) + + +def test_basic_conversion() -> None: + sample = _sample( + when="2024-01-31T14:06:57Z", + day="2024-01-31", + children=[{"name": "x"}, {"name": "y"}], + tags=["t1"], + ) + assert sample == Sample( + id="a", + count=1, + ratio=0.5, + flag=True, + colour=Colour.RED, + when=datetime(2024, 1, 31, 14, 6, 57, tzinfo=UTC), + day=date(2024, 1, 31), + children=(Child(name="x"), Child(name="y")), + tags=["t1"], + ) + + +def test_unknown_enum_value_is_kept_as_string() -> None: + sample = _sample(colour="green") + assert sample.colour == "green" + assert not isinstance(sample.colour, Colour) + + +def test_known_enum_value_is_member_and_compares_as_string() -> None: + assert _sample(colour="blue").colour is Colour.BLUE + assert _sample(colour="blue").colour == "blue" + + +def test_enum_without_str_fallback_rejects_unknown() -> None: + with pytest.raises(AuraResponseError, match=r"strict_colour: 'green' is not a valid Colour"): + _sample(strict_colour="green") + + +def test_unknown_keys_are_ignored() -> None: + assert _sample(brand_new_field={"x": 1}).id == "a" + + +def test_missing_optional_uses_default() -> None: + sample = _sample() + assert sample.note is None + assert sample.children == () + assert sample.tags == [] + + +def test_null_optional_is_none() -> None: + assert _sample(note=None, when=None).note is None + + +def test_empty_string_timestamp_is_none() -> None: + assert _sample(when="").when is None + + +def test_missing_required_field() -> None: + payload = dict(BASE) + del payload["count"] + with pytest.raises(AuraResponseError, match="count: required field is missing"): + from_json(Sample, payload) + + +def test_null_required_field() -> None: + with pytest.raises(AuraResponseError, match="id: value must not be null"): + _sample(id=None) + + +def test_nested_error_path() -> None: + with pytest.raises(AuraResponseError, match=r"children\[1\]\.name: required field is missing"): + _sample(children=[{"name": "ok"}, {}]) + + +@pytest.mark.parametrize(("value", "expected"), [(3, 3), (3.0, 3), ("42", 42), ("-2", -2)]) +def test_int_tolerance(value: object, expected: int) -> None: + assert _sample(count=value).count == expected + + +@pytest.mark.parametrize("value", [True, 3.5, "4GB", "", [1]]) +def test_int_rejects(value: object) -> None: + with pytest.raises(AuraResponseError, match="count: expected an integer"): + _sample(count=value) + + +def test_float_accepts_int_but_not_bool() -> None: + assert _sample(ratio=2).ratio == 2.0 + with pytest.raises(AuraResponseError, match="ratio"): + _sample(ratio=False) + + +def test_str_accepts_number() -> None: + assert _sample(id=123).id == "123" + + +@pytest.mark.parametrize("value", [True, {"a": 1}, ["x"]]) +def test_str_rejects(value: object) -> None: + with pytest.raises(AuraResponseError, match="id: expected a string"): + _sample(id=value) + + +@pytest.mark.parametrize("value", ["true", 1, 0]) +def test_bool_is_strict(value: object) -> None: + with pytest.raises(AuraResponseError, match="flag: expected a boolean"): + _sample(flag=value) + + +@pytest.mark.parametrize( + ("value", "expected"), + [ + ("2023-01-20T13:44:42Z", datetime(2023, 1, 20, 13, 44, 42, tzinfo=UTC)), + ("2023-01-20T13:44:42.123Z", datetime(2023, 1, 20, 13, 44, 42, 123000, tzinfo=UTC)), + # Go RFC3339Nano: nine fractional digits are truncated to microseconds. + ( + "2023-01-20T13:44:42.123456789Z", + datetime(2023, 1, 20, 13, 44, 42, 123456, tzinfo=UTC), + ), + ("2023-01-20T13:44:42+00:00", datetime(2023, 1, 20, 13, 44, 42, tzinfo=UTC)), + ], +) +def test_timestamps(value: str, expected: datetime) -> None: + assert _sample(when=value).when == expected + + +@pytest.mark.parametrize("value", ["yesterday", 1700000000, "2023-13-01T00:00:00Z"]) +def test_invalid_timestamp(value: object) -> None: + with pytest.raises(AuraResponseError, match="when: expected an ISO 8601 timestamp"): + _sample(when=value) + + +@pytest.mark.parametrize("value", ["31/01/2024", 20240131]) +def test_invalid_date(value: object) -> None: + with pytest.raises(AuraResponseError, match="day: expected an ISO date"): + _sample(day=value) + + +def test_null_for_non_optional_union() -> None: + with pytest.raises(AuraResponseError, match="colour: value must not be null"): + _sample(colour=None) + + +def test_non_object_and_non_list() -> None: + with pytest.raises(AuraResponseError, match=": expected an object, got list"): + from_json(Sample, []) + with pytest.raises(AuraResponseError, match="children: expected a list"): + _sample(children={"name": "x"}) + + +def test_parse_data_unwraps() -> None: + assert parse_data(Child, {"data": {"name": "x"}}) == Child(name="x") + assert parse_data_list(Child, {"data": [{"name": "x"}]}) == [Child(name="x")] + + +@pytest.mark.parametrize("payload", [{}, {"items": []}, [], None, "data"]) +def test_parse_data_requires_data_key(payload: object) -> None: + with pytest.raises(AuraResponseError, match="missing 'data'"): + parse_data(Child, payload) + + +def test_parse_data_list_requires_list() -> None: + with pytest.raises(AuraResponseError, match="expected a list"): + parse_data_list(Child, {"data": {"name": "x"}}) + + +def test_unsupported_field_type_is_a_programming_error() -> None: + @dataclass + class Bad: + value: bytes + + with pytest.raises(TypeError, match="unsupported model field type"): + from_json(Bad, {"value": "x"}) + + +def test_to_json_omits_none_and_converts_values() -> None: + sample = Sample( + id="a", + count=1, + ratio=0.5, + flag=False, + colour=Colour.BLUE, + when=datetime(2024, 1, 31, 14, 6, 57, tzinfo=UTC), + day=date(2024, 1, 31), + children=(Child(name="x"),), + ) + assert to_json(sample) == { + "id": "a", + "count": 1, + "ratio": 0.5, + "flag": False, + "colour": "blue", + "when": "2024-01-31T14:06:57+00:00", + "day": "2024-01-31", + "children": [{"name": "x"}], + "tags": [], + } + + +def test_to_json_mapping_drops_none() -> None: + assert to_json({"name": "x", "memory": None, "colour": Colour.RED}) == { + "name": "x", + "colour": "red", + } diff --git a/tests/unit/test_spec_examples.py b/tests/unit/test_spec_examples.py new file mode 100644 index 0000000..d3f05dc --- /dev/null +++ b/tests/unit/test_spec_examples.py @@ -0,0 +1,116 @@ +"""Parse every 2xx response example in the v1 OpenAPI spec with its model. + +If the spec adds an operation or a success response with a JSON body, the coverage test fails +until it is mapped to a model below. +""" + +from __future__ import annotations + +from pathlib import Path +from typing import Any + +import pytest +import yaml + +from aura_python_sdk import models +from aura_python_sdk._internal._serde import parse_data, parse_data_list + +SPEC_PATH = Path(__file__).resolve().parents[2] / "aura_api_spec_v1 .yaml" +HTTP_METHODS = {"get", "post", "put", "patch", "delete"} + + +class _SpecLoader(yaml.SafeLoader): + """SafeLoader that leaves timestamps as strings, as they would arrive in JSON.""" + + +_SpecLoader.yaml_implicit_resolvers = { + key: [(tag, regex) for tag, regex in resolvers if tag != "tag:yaml.org,2002:timestamp"] + for key, resolvers in yaml.SafeLoader.yaml_implicit_resolvers.items() +} + +# (operationId, status) -> (model, is_list) +RESPONSE_MODELS: dict[tuple[str, str], tuple[type[Any], bool]] = { + ("get-instances", "200"): (models.InstanceSummary, True), + ("post-instances", "202"): (models.CreatedInstance, False), + ("post-instances-sizing", "200"): (models.InstanceSizeEstimate, False), + ("get-instance-id", "200"): (models.Instance, False), + ("delete-instance-id", "202"): (models.Instance, False), + ("patch-instance-id", "200"): (models.Instance, False), + ("patch-instance-id", "202"): (models.Instance, False), + ("post-overwrite-instance", "202"): (models.Instance, False), + ("post-pause-instance", "202"): (models.Instance, False), + ("post-resume-instance", "202"): (models.Instance, False), + ("get-snapshot-snapshotid", "200"): (models.Snapshot, False), + ("post-restore-snapshot", "202"): (models.Instance, False), + ("get-snapshots", "200"): (models.Snapshot, True), + ("post-snapshots", "202"): (models.CreatedSnapshot, False), + ("post-upgrade-instance", "200"): (models.Instance, False), + ("get-projects", "200"): (models.TenantSummary, True), + ("get-project-id", "200"): (models.Tenant, False), + ("get-customer-managed-keys", "200"): (models.CustomerManagedKeySummary, True), + ("post-customer-managed-keys", "202"): (models.CustomerManagedKey, False), + ("get-customer-managed-key-id", "200"): (models.CustomerManagedKey, False), + ("get-project-metrics-integration-details", "200"): (models.MetricsIntegration, False), + ("get-sessions", "200"): (models.GDSSession, True), + ("post-session", "200"): (models.GDSSession, False), + ("post-session", "202"): (models.GDSSession, False), + ("post-sessions-sizing", "200"): (models.GDSSessionSizeEstimate, False), + ("get-session", "200"): (models.GDSSession, False), + ("delete-session", "202"): (models.DeletedGDSSession, False), +} + + +def _load_spec() -> dict[str, Any]: + with SPEC_PATH.open(encoding="utf-8") as handle: + spec: dict[str, Any] = yaml.load(handle, Loader=_SpecLoader) # noqa: S506 - SafeLoader subclass + return spec + + +def _success_responses() -> list[tuple[str, str, dict[str, Any]]]: + """(operationId, status, application/json content) for every 2xx response with a body.""" + found = [] + for path_item in _load_spec()["paths"].values(): + for method, operation in path_item.items(): + if method not in HTTP_METHODS: + continue + for status, response in operation.get("responses", {}).items(): + content = (response or {}).get("content", {}).get("application/json") + if str(status).startswith("2") and content: + found.append((operation["operationId"], str(status), content)) + return found + + +def _examples() -> list[Any]: + cases = [] + for operation_id, status, content in _success_responses(): + values = [] + if "example" in content: + values.append(("example", content["example"])) + for name, example in (content.get("examples") or {}).items(): + values.append((name, example["value"])) + for name, value in values: + cases.append( + pytest.param(operation_id, status, value, id=f"{operation_id}-{status}-{name}") + ) + return cases + + +def test_every_success_response_has_a_model() -> None: + documented = {(operation_id, status) for operation_id, status, _ in _success_responses()} + assert documented - RESPONSE_MODELS.keys() == set() + assert RESPONSE_MODELS.keys() - documented == set() + + +def test_spec_has_examples() -> None: + assert len(_examples()) >= 25 + + +@pytest.mark.parametrize(("operation_id", "status", "example"), _examples()) +def test_spec_example_parses(operation_id: str, status: str, example: Any) -> None: + model, is_list = RESPONSE_MODELS[(operation_id, status)] + if is_list: + parsed = parse_data_list(model, example) + assert len(parsed) == len(example["data"]) + assert all(isinstance(item, model) for item in parsed) + else: + assert isinstance(parse_data(model, example), model) diff --git a/uv.lock b/uv.lock index 662319b..3f69992 100644 --- a/uv.lock +++ b/uv.lock @@ -100,7 +100,9 @@ dev = [ { name = "mypy" }, { name = "pytest" }, { name = "pytest-cov" }, + { name = "pyyaml" }, { name = "ruff" }, + { name = "types-pyyaml" }, ] [package.metadata] @@ -115,7 +117,9 @@ dev = [ { name = "mypy", specifier = ">=1.11" }, { name = "pytest", specifier = ">=8" }, { name = "pytest-cov", specifier = ">=5" }, + { name = "pyyaml", specifier = ">=6.0.3" }, { name = "ruff", specifier = ">=0.6" }, + { name = "types-pyyaml", specifier = ">=6.0.12.20260906" }, ] [[package]] @@ -570,6 +574,61 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/9d/7a/d968e294073affff457b041c2be9868a40c1c71f4a35fcc1e45e5493067b/pytest_cov-7.1.0-py3-none-any.whl", hash = "sha256:a0461110b7865f9a271aa1b51e516c9a95de9d696734a2f71e3e78f46e1d4678", size = 22876, upload-time = "2026-03-21T20:11:14.438Z" }, ] +[[package]] +name = "pyyaml" +version = "6.0.3" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/05/8e/961c0007c59b8dd7729d542c61a4d537767a59645b82a0b521206e1e25c2/pyyaml-6.0.3.tar.gz", hash = "sha256:d76623373421df22fb4cf8817020cbb7ef15c725b9d5e45f17e189bfc384190f", size = 130960, upload-time = "2025-09-25T21:33:16.546Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/6d/16/a95b6757765b7b031c9374925bb718d55e0a9ba8a1b6a12d25962ea44347/pyyaml-6.0.3-cp311-cp311-macosx_10_13_x86_64.whl", hash = "sha256:44edc647873928551a01e7a563d7452ccdebee747728c1080d881d68af7b997e", size = 185826, upload-time = "2025-09-25T21:31:58.655Z" }, + { url = "https://files.pythonhosted.org/packages/16/19/13de8e4377ed53079ee996e1ab0a9c33ec2faf808a4647b7b4c0d46dd239/pyyaml-6.0.3-cp311-cp311-macosx_11_0_arm64.whl", hash = "sha256:652cb6edd41e718550aad172851962662ff2681490a8a711af6a4d288dd96824", size = 175577, upload-time = "2025-09-25T21:32:00.088Z" }, + { url = "https://files.pythonhosted.org/packages/0c/62/d2eb46264d4b157dae1275b573017abec435397aa59cbcdab6fc978a8af4/pyyaml-6.0.3-cp311-cp311-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:10892704fc220243f5305762e276552a0395f7beb4dbf9b14ec8fd43b57f126c", size = 775556, upload-time = "2025-09-25T21:32:01.31Z" }, + { url = "https://files.pythonhosted.org/packages/10/cb/16c3f2cf3266edd25aaa00d6c4350381c8b012ed6f5276675b9eba8d9ff4/pyyaml-6.0.3-cp311-cp311-manylinux2014_s390x.manylinux_2_17_s390x.manylinux_2_28_s390x.whl", hash = "sha256:850774a7879607d3a6f50d36d04f00ee69e7fc816450e5f7e58d7f17f1ae5c00", size = 882114, upload-time = "2025-09-25T21:32:03.376Z" }, + { url = "https://files.pythonhosted.org/packages/71/60/917329f640924b18ff085ab889a11c763e0b573da888e8404ff486657602/pyyaml-6.0.3-cp311-cp311-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:b8bb0864c5a28024fac8a632c443c87c5aa6f215c0b126c449ae1a150412f31d", size = 806638, upload-time = "2025-09-25T21:32:04.553Z" }, + { url = "https://files.pythonhosted.org/packages/dd/6f/529b0f316a9fd167281a6c3826b5583e6192dba792dd55e3203d3f8e655a/pyyaml-6.0.3-cp311-cp311-musllinux_1_2_aarch64.whl", hash = "sha256:1d37d57ad971609cf3c53ba6a7e365e40660e3be0e5175fa9f2365a379d6095a", size = 767463, upload-time = "2025-09-25T21:32:06.152Z" }, + { url = "https://files.pythonhosted.org/packages/f2/6a/b627b4e0c1dd03718543519ffb2f1deea4a1e6d42fbab8021936a4d22589/pyyaml-6.0.3-cp311-cp311-musllinux_1_2_x86_64.whl", hash = "sha256:37503bfbfc9d2c40b344d06b2199cf0e96e97957ab1c1b546fd4f87e53e5d3e4", size = 794986, upload-time = "2025-09-25T21:32:07.367Z" }, + { url = "https://files.pythonhosted.org/packages/45/91/47a6e1c42d9ee337c4839208f30d9f09caa9f720ec7582917b264defc875/pyyaml-6.0.3-cp311-cp311-win32.whl", hash = "sha256:8098f252adfa6c80ab48096053f512f2321f0b998f98150cea9bd23d83e1467b", size = 142543, upload-time = "2025-09-25T21:32:08.95Z" }, + { url = "https://files.pythonhosted.org/packages/da/e3/ea007450a105ae919a72393cb06f122f288ef60bba2dc64b26e2646fa315/pyyaml-6.0.3-cp311-cp311-win_amd64.whl", hash = "sha256:9f3bfb4965eb874431221a3ff3fdcddc7e74e3b07799e0e84ca4a0f867d449bf", size = 158763, upload-time = "2025-09-25T21:32:09.96Z" }, + { url = "https://files.pythonhosted.org/packages/d1/33/422b98d2195232ca1826284a76852ad5a86fe23e31b009c9886b2d0fb8b2/pyyaml-6.0.3-cp312-cp312-macosx_10_13_x86_64.whl", hash = "sha256:7f047e29dcae44602496db43be01ad42fc6f1cc0d8cd6c83d342306c32270196", size = 182063, upload-time = "2025-09-25T21:32:11.445Z" }, + { url = "https://files.pythonhosted.org/packages/89/a0/6cf41a19a1f2f3feab0e9c0b74134aa2ce6849093d5517a0c550fe37a648/pyyaml-6.0.3-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:fc09d0aa354569bc501d4e787133afc08552722d3ab34836a80547331bb5d4a0", size = 173973, upload-time = "2025-09-25T21:32:12.492Z" }, + { url = "https://files.pythonhosted.org/packages/ed/23/7a778b6bd0b9a8039df8b1b1d80e2e2ad78aa04171592c8a5c43a56a6af4/pyyaml-6.0.3-cp312-cp312-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:9149cad251584d5fb4981be1ecde53a1ca46c891a79788c0df828d2f166bda28", size = 775116, upload-time = "2025-09-25T21:32:13.652Z" }, + { url = "https://files.pythonhosted.org/packages/65/30/d7353c338e12baef4ecc1b09e877c1970bd3382789c159b4f89d6a70dc09/pyyaml-6.0.3-cp312-cp312-manylinux2014_s390x.manylinux_2_17_s390x.manylinux_2_28_s390x.whl", hash = "sha256:5fdec68f91a0c6739b380c83b951e2c72ac0197ace422360e6d5a959d8d97b2c", size = 844011, upload-time = "2025-09-25T21:32:15.21Z" }, + { url = "https://files.pythonhosted.org/packages/8b/9d/b3589d3877982d4f2329302ef98a8026e7f4443c765c46cfecc8858c6b4b/pyyaml-6.0.3-cp312-cp312-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:ba1cc08a7ccde2d2ec775841541641e4548226580ab850948cbfda66a1befcdc", size = 807870, upload-time = "2025-09-25T21:32:16.431Z" }, + { url = "https://files.pythonhosted.org/packages/05/c0/b3be26a015601b822b97d9149ff8cb5ead58c66f981e04fedf4e762f4bd4/pyyaml-6.0.3-cp312-cp312-musllinux_1_2_aarch64.whl", hash = "sha256:8dc52c23056b9ddd46818a57b78404882310fb473d63f17b07d5c40421e47f8e", size = 761089, upload-time = "2025-09-25T21:32:17.56Z" }, + { url = "https://files.pythonhosted.org/packages/be/8e/98435a21d1d4b46590d5459a22d88128103f8da4c2d4cb8f14f2a96504e1/pyyaml-6.0.3-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:41715c910c881bc081f1e8872880d3c650acf13dfa8214bad49ed4cede7c34ea", size = 790181, upload-time = "2025-09-25T21:32:18.834Z" }, + { url = "https://files.pythonhosted.org/packages/74/93/7baea19427dcfbe1e5a372d81473250b379f04b1bd3c4c5ff825e2327202/pyyaml-6.0.3-cp312-cp312-win32.whl", hash = "sha256:96b533f0e99f6579b3d4d4995707cf36df9100d67e0c8303a0c55b27b5f99bc5", size = 137658, upload-time = "2025-09-25T21:32:20.209Z" }, + { url = "https://files.pythonhosted.org/packages/86/bf/899e81e4cce32febab4fb42bb97dcdf66bc135272882d1987881a4b519e9/pyyaml-6.0.3-cp312-cp312-win_amd64.whl", hash = "sha256:5fcd34e47f6e0b794d17de1b4ff496c00986e1c83f7ab2fb8fcfe9616ff7477b", size = 154003, upload-time = "2025-09-25T21:32:21.167Z" }, + { url = "https://files.pythonhosted.org/packages/1a/08/67bd04656199bbb51dbed1439b7f27601dfb576fb864099c7ef0c3e55531/pyyaml-6.0.3-cp312-cp312-win_arm64.whl", hash = "sha256:64386e5e707d03a7e172c0701abfb7e10f0fb753ee1d773128192742712a98fd", size = 140344, upload-time = "2025-09-25T21:32:22.617Z" }, + { url = "https://files.pythonhosted.org/packages/d1/11/0fd08f8192109f7169db964b5707a2f1e8b745d4e239b784a5a1dd80d1db/pyyaml-6.0.3-cp313-cp313-macosx_10_13_x86_64.whl", hash = "sha256:8da9669d359f02c0b91ccc01cac4a67f16afec0dac22c2ad09f46bee0697eba8", size = 181669, upload-time = "2025-09-25T21:32:23.673Z" }, + { url = "https://files.pythonhosted.org/packages/b1/16/95309993f1d3748cd644e02e38b75d50cbc0d9561d21f390a76242ce073f/pyyaml-6.0.3-cp313-cp313-macosx_11_0_arm64.whl", hash = "sha256:2283a07e2c21a2aa78d9c4442724ec1eb15f5e42a723b99cb3d822d48f5f7ad1", size = 173252, upload-time = "2025-09-25T21:32:25.149Z" }, + { url = "https://files.pythonhosted.org/packages/50/31/b20f376d3f810b9b2371e72ef5adb33879b25edb7a6d072cb7ca0c486398/pyyaml-6.0.3-cp313-cp313-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:ee2922902c45ae8ccada2c5b501ab86c36525b883eff4255313a253a3160861c", size = 767081, upload-time = "2025-09-25T21:32:26.575Z" }, + { url = "https://files.pythonhosted.org/packages/49/1e/a55ca81e949270d5d4432fbbd19dfea5321eda7c41a849d443dc92fd1ff7/pyyaml-6.0.3-cp313-cp313-manylinux2014_s390x.manylinux_2_17_s390x.manylinux_2_28_s390x.whl", hash = "sha256:a33284e20b78bd4a18c8c2282d549d10bc8408a2a7ff57653c0cf0b9be0afce5", size = 841159, upload-time = "2025-09-25T21:32:27.727Z" }, + { url = "https://files.pythonhosted.org/packages/74/27/e5b8f34d02d9995b80abcef563ea1f8b56d20134d8f4e5e81733b1feceb2/pyyaml-6.0.3-cp313-cp313-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:0f29edc409a6392443abf94b9cf89ce99889a1dd5376d94316ae5145dfedd5d6", size = 801626, upload-time = "2025-09-25T21:32:28.878Z" }, + { url = "https://files.pythonhosted.org/packages/f9/11/ba845c23988798f40e52ba45f34849aa8a1f2d4af4b798588010792ebad6/pyyaml-6.0.3-cp313-cp313-musllinux_1_2_aarch64.whl", hash = "sha256:f7057c9a337546edc7973c0d3ba84ddcdf0daa14533c2065749c9075001090e6", size = 753613, upload-time = "2025-09-25T21:32:30.178Z" }, + { url = "https://files.pythonhosted.org/packages/3d/e0/7966e1a7bfc0a45bf0a7fb6b98ea03fc9b8d84fa7f2229e9659680b69ee3/pyyaml-6.0.3-cp313-cp313-musllinux_1_2_x86_64.whl", hash = "sha256:eda16858a3cab07b80edaf74336ece1f986ba330fdb8ee0d6c0d68fe82bc96be", size = 794115, upload-time = "2025-09-25T21:32:31.353Z" }, + { url = "https://files.pythonhosted.org/packages/de/94/980b50a6531b3019e45ddeada0626d45fa85cbe22300844a7983285bed3b/pyyaml-6.0.3-cp313-cp313-win32.whl", hash = "sha256:d0eae10f8159e8fdad514efdc92d74fd8d682c933a6dd088030f3834bc8e6b26", size = 137427, upload-time = "2025-09-25T21:32:32.58Z" }, + { url = "https://files.pythonhosted.org/packages/97/c9/39d5b874e8b28845e4ec2202b5da735d0199dbe5b8fb85f91398814a9a46/pyyaml-6.0.3-cp313-cp313-win_amd64.whl", hash = "sha256:79005a0d97d5ddabfeeea4cf676af11e647e41d81c9a7722a193022accdb6b7c", size = 154090, upload-time = "2025-09-25T21:32:33.659Z" }, + { url = "https://files.pythonhosted.org/packages/73/e8/2bdf3ca2090f68bb3d75b44da7bbc71843b19c9f2b9cb9b0f4ab7a5a4329/pyyaml-6.0.3-cp313-cp313-win_arm64.whl", hash = "sha256:5498cd1645aa724a7c71c8f378eb29ebe23da2fc0d7a08071d89469bf1d2defb", size = 140246, upload-time = "2025-09-25T21:32:34.663Z" }, + { url = "https://files.pythonhosted.org/packages/9d/8c/f4bd7f6465179953d3ac9bc44ac1a8a3e6122cf8ada906b4f96c60172d43/pyyaml-6.0.3-cp314-cp314-macosx_10_13_x86_64.whl", hash = "sha256:8d1fab6bb153a416f9aeb4b8763bc0f22a5586065f86f7664fc23339fc1c1fac", size = 181814, upload-time = "2025-09-25T21:32:35.712Z" }, + { url = "https://files.pythonhosted.org/packages/bd/9c/4d95bb87eb2063d20db7b60faa3840c1b18025517ae857371c4dd55a6b3a/pyyaml-6.0.3-cp314-cp314-macosx_11_0_arm64.whl", hash = "sha256:34d5fcd24b8445fadc33f9cf348c1047101756fd760b4dacb5c3e99755703310", size = 173809, upload-time = "2025-09-25T21:32:36.789Z" }, + { url = "https://files.pythonhosted.org/packages/92/b5/47e807c2623074914e29dabd16cbbdd4bf5e9b2db9f8090fa64411fc5382/pyyaml-6.0.3-cp314-cp314-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:501a031947e3a9025ed4405a168e6ef5ae3126c59f90ce0cd6f2bfc477be31b7", size = 766454, upload-time = "2025-09-25T21:32:37.966Z" }, + { url = "https://files.pythonhosted.org/packages/02/9e/e5e9b168be58564121efb3de6859c452fccde0ab093d8438905899a3a483/pyyaml-6.0.3-cp314-cp314-manylinux2014_s390x.manylinux_2_17_s390x.manylinux_2_28_s390x.whl", hash = "sha256:b3bc83488de33889877a0f2543ade9f70c67d66d9ebb4ac959502e12de895788", size = 836355, upload-time = "2025-09-25T21:32:39.178Z" }, + { url = "https://files.pythonhosted.org/packages/88/f9/16491d7ed2a919954993e48aa941b200f38040928474c9e85ea9e64222c3/pyyaml-6.0.3-cp314-cp314-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:c458b6d084f9b935061bc36216e8a69a7e293a2f1e68bf956dcd9e6cbcd143f5", size = 794175, upload-time = "2025-09-25T21:32:40.865Z" }, + { url = "https://files.pythonhosted.org/packages/dd/3f/5989debef34dc6397317802b527dbbafb2b4760878a53d4166579111411e/pyyaml-6.0.3-cp314-cp314-musllinux_1_2_aarch64.whl", hash = "sha256:7c6610def4f163542a622a73fb39f534f8c101d690126992300bf3207eab9764", size = 755228, upload-time = "2025-09-25T21:32:42.084Z" }, + { url = "https://files.pythonhosted.org/packages/d7/ce/af88a49043cd2e265be63d083fc75b27b6ed062f5f9fd6cdc223ad62f03e/pyyaml-6.0.3-cp314-cp314-musllinux_1_2_x86_64.whl", hash = "sha256:5190d403f121660ce8d1d2c1bb2ef1bd05b5f68533fc5c2ea899bd15f4399b35", size = 789194, upload-time = "2025-09-25T21:32:43.362Z" }, + { url = "https://files.pythonhosted.org/packages/23/20/bb6982b26a40bb43951265ba29d4c246ef0ff59c9fdcdf0ed04e0687de4d/pyyaml-6.0.3-cp314-cp314-win_amd64.whl", hash = "sha256:4a2e8cebe2ff6ab7d1050ecd59c25d4c8bd7e6f400f5f82b96557ac0abafd0ac", size = 156429, upload-time = "2025-09-25T21:32:57.844Z" }, + { url = "https://files.pythonhosted.org/packages/f4/f4/a4541072bb9422c8a883ab55255f918fa378ecf083f5b85e87fc2b4eda1b/pyyaml-6.0.3-cp314-cp314-win_arm64.whl", hash = "sha256:93dda82c9c22deb0a405ea4dc5f2d0cda384168e466364dec6255b293923b2f3", size = 143912, upload-time = "2025-09-25T21:32:59.247Z" }, + { url = "https://files.pythonhosted.org/packages/7c/f9/07dd09ae774e4616edf6cda684ee78f97777bdd15847253637a6f052a62f/pyyaml-6.0.3-cp314-cp314t-macosx_10_13_x86_64.whl", hash = "sha256:02893d100e99e03eda1c8fd5c441d8c60103fd175728e23e431db1b589cf5ab3", size = 189108, upload-time = "2025-09-25T21:32:44.377Z" }, + { url = "https://files.pythonhosted.org/packages/4e/78/8d08c9fb7ce09ad8c38ad533c1191cf27f7ae1effe5bb9400a46d9437fcf/pyyaml-6.0.3-cp314-cp314t-macosx_11_0_arm64.whl", hash = "sha256:c1ff362665ae507275af2853520967820d9124984e0f7466736aea23d8611fba", size = 183641, upload-time = "2025-09-25T21:32:45.407Z" }, + { url = "https://files.pythonhosted.org/packages/7b/5b/3babb19104a46945cf816d047db2788bcaf8c94527a805610b0289a01c6b/pyyaml-6.0.3-cp314-cp314t-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:6adc77889b628398debc7b65c073bcb99c4a0237b248cacaf3fe8a557563ef6c", size = 831901, upload-time = "2025-09-25T21:32:48.83Z" }, + { url = "https://files.pythonhosted.org/packages/8b/cc/dff0684d8dc44da4d22a13f35f073d558c268780ce3c6ba1b87055bb0b87/pyyaml-6.0.3-cp314-cp314t-manylinux2014_s390x.manylinux_2_17_s390x.manylinux_2_28_s390x.whl", hash = "sha256:a80cb027f6b349846a3bf6d73b5e95e782175e52f22108cfa17876aaeff93702", size = 861132, upload-time = "2025-09-25T21:32:50.149Z" }, + { url = "https://files.pythonhosted.org/packages/b1/5e/f77dc6b9036943e285ba76b49e118d9ea929885becb0a29ba8a7c75e29fe/pyyaml-6.0.3-cp314-cp314t-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:00c4bdeba853cc34e7dd471f16b4114f4162dc03e6b7afcc2128711f0eca823c", size = 839261, upload-time = "2025-09-25T21:32:51.808Z" }, + { url = "https://files.pythonhosted.org/packages/ce/88/a9db1376aa2a228197c58b37302f284b5617f56a5d959fd1763fb1675ce6/pyyaml-6.0.3-cp314-cp314t-musllinux_1_2_aarch64.whl", hash = "sha256:66e1674c3ef6f541c35191caae2d429b967b99e02040f5ba928632d9a7f0f065", size = 805272, upload-time = "2025-09-25T21:32:52.941Z" }, + { url = "https://files.pythonhosted.org/packages/da/92/1446574745d74df0c92e6aa4a7b0b3130706a4142b2d1a5869f2eaa423c6/pyyaml-6.0.3-cp314-cp314t-musllinux_1_2_x86_64.whl", hash = "sha256:16249ee61e95f858e83976573de0f5b2893b3677ba71c9dd36b9cf8be9ac6d65", size = 829923, upload-time = "2025-09-25T21:32:54.537Z" }, + { url = "https://files.pythonhosted.org/packages/f0/7a/1c7270340330e575b92f397352af856a8c06f230aa3e76f86b39d01b416a/pyyaml-6.0.3-cp314-cp314t-win_amd64.whl", hash = "sha256:4ad1906908f2f5ae4e5a8ddfce73c320c2a1429ec52eafd27138b7f1cbe341c9", size = 174062, upload-time = "2025-09-25T21:32:55.767Z" }, + { url = "https://files.pythonhosted.org/packages/f1/12/de94a39c2ef588c7e6455cfbe7343d3b2dc9d6b6b2f40c4c6565744c873d/pyyaml-6.0.3-cp314-cp314t-win_arm64.whl", hash = "sha256:ebc55a14a21cb14062aa4162f906cd962b28e2e9ea38f9b4391244cd8de4ae0b", size = 149341, upload-time = "2025-09-25T21:32:56.828Z" }, +] + [[package]] name = "ruff" version = "0.16.9" @@ -649,6 +708,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/7b/61/cceae43728b7de99d9b847560c262873a1f6c98202171fd5ed62640b494b/tomli-2.4.1-py3-none-any.whl", hash = "sha256:0d85819802132122da43cb86656f8d1f8c6587d54ae7dcaf30e90533028b49fe", size = 14583, upload-time = "2026-03-25T20:22:03.012Z" }, ] +[[package]] +name = "types-pyyaml" +version = "6.0.12.20260906" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/90/6e/abec85b9013db5b934b0280a6dd104904d84f7bcbaab2e2f3def87ac7463/types_pyyaml-6.0.12.20260906.tar.gz", hash = "sha256:f59c1cc05010b833d2d72287bbaa72610106b28d42d89a907313117faba85212", size = 18649, upload-time = "2026-09-06T06:35:35.362Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/15/c0/fc0644b7ddcfb969e95845837143cb5173ddd6e06ee4ba5fc493cd9329b7/types_pyyaml-6.0.12.20260906-py3-none-any.whl", hash = "sha256:bca893ff0d51df5c9053137d5d0e6ccd36e939a196356f1d5c16372422f5137b", size = 21282, upload-time = "2026-09-06T06:35:34.372Z" }, +] + [[package]] name = "typing-extensions" version = "4.16.0" From 1e6a5a301c66d9268b7b168ed5cadcb44b959a6e Mon Sep 17 00:00:00 2001 From: Jonathan Giffard Date: Tue, 29 Sep 2026 20:15:29 +0100 Subject: [PATCH 3/8] Add tenant, instance, snapshot, CMEK and GDS session services Phase 4 of PLAN.md: every method of the Go SDK's v1 service interfaces, wired onto AuraClient, with Go's client-side ID and config validation run before any request. Overwrite responses and list filters follow the spec. Co-Authored-By: Claude Opus 5.5 --- PLAN.md | 6 +- src/aura_python_sdk/_client.py | 21 +- src/aura_python_sdk/_validation.py | 63 ++++ src/aura_python_sdk/services/__init__.py | 15 + src/aura_python_sdk/services/_base.py | 13 + src/aura_python_sdk/services/cmek.py | 27 ++ .../services/graph_analytics.py | 107 +++++++ src/aura_python_sdk/services/instances.py | 181 ++++++++++++ src/aura_python_sdk/services/snapshots.py | 71 +++++ src/aura_python_sdk/services/tenants.py | 35 +++ tests/unit/conftest.py | 74 +++++ tests/unit/test_cmek_service.py | 24 ++ tests/unit/test_graph_analytics_service.py | 149 ++++++++++ tests/unit/test_instances_service.py | 273 ++++++++++++++++++ tests/unit/test_snapshots_service.py | 82 ++++++ tests/unit/test_tenants_service.py | 47 +++ tests/unit/test_validation.py | 61 ++++ 17 files changed, 1246 insertions(+), 3 deletions(-) create mode 100644 src/aura_python_sdk/_validation.py create mode 100644 src/aura_python_sdk/services/__init__.py create mode 100644 src/aura_python_sdk/services/_base.py create mode 100644 src/aura_python_sdk/services/cmek.py create mode 100644 src/aura_python_sdk/services/graph_analytics.py create mode 100644 src/aura_python_sdk/services/instances.py create mode 100644 src/aura_python_sdk/services/snapshots.py create mode 100644 src/aura_python_sdk/services/tenants.py create mode 100644 tests/unit/conftest.py create mode 100644 tests/unit/test_cmek_service.py create mode 100644 tests/unit/test_graph_analytics_service.py create mode 100644 tests/unit/test_instances_service.py create mode 100644 tests/unit/test_snapshots_service.py create mode 100644 tests/unit/test_tenants_service.py create mode 100644 tests/unit/test_validation.py diff --git a/PLAN.md b/PLAN.md index f0d7193..fc3f2c7 100644 --- a/PLAN.md +++ b/PLAN.md @@ -277,9 +277,11 @@ packages. ## 8. Spec and Go discrepancies to resolve during implementation - **Query parameter name**: the spec names the list-filter parameter `tenantId`, but Go sends - `tenant_id` (CMEK list). Check against the live API. + `tenant_id` (CMEK list). *Resolved: follow the spec. `tenantId` is used for + every list filter (defined once in `services/cmek.py`).* - **Overwrite response**: Go models it as `{"data": ""}`, but the spec says - `Instance`. Parse tolerantly and confirm. + `Instance`. The Go tests only use mocks. *Resolved: follow the spec. `overwrite_from_instance` + and `overwrite_from_snapshot` return `Instance`.* - **GDS `ttl` type**: the spec says `integer` in the session details but `string` in the create request. Go uses string throughout. - **GDS create response**: the spec has an odd `data: {type: object, items: ...}` shape. Treat it as diff --git a/src/aura_python_sdk/_client.py b/src/aura_python_sdk/_client.py index 237681f..f703c61 100644 --- a/src/aura_python_sdk/_client.py +++ b/src/aura_python_sdk/_client.py @@ -24,6 +24,13 @@ from aura_python_sdk._internal.http._httpx import HttpxTransport from aura_python_sdk._internal.http._service import HttpService from aura_python_sdk._transport import HttpTransport +from aura_python_sdk.services import ( + CMEKService, + GDSSessionService, + InstanceService, + SnapshotService, + TenantService, +) ENV_CLIENT_ID = "AURA_CLIENT_ID" ENV_CLIENT_SECRET = "AURA_CLIENT_SECRET" # noqa: S105 - environment variable name, not a secret @@ -37,7 +44,11 @@ class AuraClient: Example:: with AuraClient(client_id="...", client_secret="...") as client: - ... + for instance in client.instances.list(): + print(instance.id, instance.name) + + Services, mirroring the Go SDK: ``tenants``, ``instances``, ``snapshots``, ``cmek`` and + ``graph_analytics``. Every option is keyword-only. Invalid options raise :class:`AuraConfigurationError`. @@ -120,6 +131,14 @@ def __init__( logger=self._logger.getChild("api"), ) + self.tenants = TenantService(self._api, self._logger.getChild("tenants")) + self.instances = InstanceService(self._api, self._logger.getChild("instances")) + self.snapshots = SnapshotService(self._api, self._logger.getChild("snapshots")) + self.cmek = CMEKService(self._api, self._logger.getChild("cmek")) + self.graph_analytics = GDSSessionService( + self._api, self._logger.getChild("graph_analytics") + ) + self._logger.debug( "Aura API client initialized", extra={"base_url": self._config.base_url, "api_version": API_VERSION}, diff --git a/src/aura_python_sdk/_validation.py b/src/aura_python_sdk/_validation.py new file mode 100644 index 0000000..cb136b7 --- /dev/null +++ b/src/aura_python_sdk/_validation.py @@ -0,0 +1,63 @@ +"""Client-side argument validation, run before any request is sent (Go: internal/utils).""" + +from __future__ import annotations + +import re + +from aura_python_sdk._errors import AuraValidationError + +_UUID = re.compile(r"[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}") +_INSTANCE_ID = re.compile(r"[0-9a-fA-F]{8}") + +MAX_INSTANCE_NAME_LENGTH = 30 + + +def require_non_empty(name: str, value: object) -> str: + if not isinstance(value, str) or not value.strip(): + raise AuraValidationError(f"{name} must not be empty") + return value + + +def _uuid(name: str, value: object) -> str: + value = require_non_empty(name, value) + if not _UUID.fullmatch(value): + raise AuraValidationError( + f"{name} must be a valid UUID format (xxxxxxxx-xxxx-xxxx-xxxx-xxxxxxxxxxxx)" + ) + return value + + +def instance_id(value: object, name: str = "instance ID") -> str: + value = require_non_empty(name, value) + if not _INSTANCE_ID.fullmatch(value): + raise AuraValidationError( + f"{name} must be in the format of a 8-character hex string (xxxxxxxx)" + ) + return value + + +def tenant_id(value: object, name: str = "tenant ID") -> str: + return _uuid(name, value) + + +def snapshot_id(value: object, name: str = "snapshot ID") -> str: + return _uuid(name, value) + + +def session_id(value: object) -> str: + return require_non_empty("GDS session ID", value) + + +def instance_name(value: object) -> str: + value = require_non_empty("instance name", value) + if len(value) > MAX_INSTANCE_NAME_LENGTH: + raise AuraValidationError( + f"instance name must be at most {MAX_INSTANCE_NAME_LENGTH} characters long" + ) + return value + + +def non_negative_int(name: str, value: object) -> int: + if isinstance(value, bool) or not isinstance(value, int) or value < 0: + raise AuraValidationError(f"{name} must be an integer of zero or more") + return value diff --git a/src/aura_python_sdk/services/__init__.py b/src/aura_python_sdk/services/__init__.py new file mode 100644 index 0000000..f084f56 --- /dev/null +++ b/src/aura_python_sdk/services/__init__.py @@ -0,0 +1,15 @@ +"""The grouped services exposed on :class:`~aura_python_sdk.AuraClient`.""" + +from aura_python_sdk.services.cmek import CMEKService +from aura_python_sdk.services.graph_analytics import GDSSessionService +from aura_python_sdk.services.instances import InstanceService +from aura_python_sdk.services.snapshots import SnapshotService +from aura_python_sdk.services.tenants import TenantService + +__all__ = [ + "CMEKService", + "GDSSessionService", + "InstanceService", + "SnapshotService", + "TenantService", +] diff --git a/src/aura_python_sdk/services/_base.py b/src/aura_python_sdk/services/_base.py new file mode 100644 index 0000000..218fc71 --- /dev/null +++ b/src/aura_python_sdk/services/_base.py @@ -0,0 +1,13 @@ +from __future__ import annotations + +import logging + +from aura_python_sdk._internal._request import RequestService + + +class Service: + """Shared plumbing for the grouped services on :class:`AuraClient`.""" + + def __init__(self, api: RequestService, logger: logging.Logger) -> None: + self._api = api + self._logger = logger diff --git a/src/aura_python_sdk/services/cmek.py b/src/aura_python_sdk/services/cmek.py new file mode 100644 index 0000000..e8a531b --- /dev/null +++ b/src/aura_python_sdk/services/cmek.py @@ -0,0 +1,27 @@ +"""``client.cmek`` (Go: CMEKService).""" + +from __future__ import annotations + +import builtins + +from aura_python_sdk import _validation as validate +from aura_python_sdk._internal._serde import parse_data_list +from aura_python_sdk.models.cmek import CustomerManagedKeySummary +from aura_python_sdk.services._base import Service + +# The spec names this query parameter tenantId. The Go SDK sends tenant_id. +TENANT_FILTER_PARAM = "tenantId" + + +class CMEKService(Service): + """Customer-managed encryption keys.""" + + def list(self, tenant_id: str | None = None) -> builtins.list[CustomerManagedKeySummary]: + """Every key the credentials can access, optionally only those in one tenant.""" + if tenant_id is not None: + tenant_id = validate.tenant_id(tenant_id) + self._logger.debug("listing customer managed keys", extra={"tenant_id": tenant_id}) + response = self._api.get("customer-managed-keys", params={TENANT_FILTER_PARAM: tenant_id}) + keys = parse_data_list(CustomerManagedKeySummary, response.json()) + self._logger.debug("customer managed keys listed", extra={"count": len(keys)}) + return keys diff --git a/src/aura_python_sdk/services/graph_analytics.py b/src/aura_python_sdk/services/graph_analytics.py new file mode 100644 index 0000000..19a45fd --- /dev/null +++ b/src/aura_python_sdk/services/graph_analytics.py @@ -0,0 +1,107 @@ +"""``client.graph_analytics`` (Go: GDSSessionService).""" + +from __future__ import annotations + +import builtins +from collections.abc import Sequence + +from aura_python_sdk import _validation as validate +from aura_python_sdk._errors import AuraValidationError +from aura_python_sdk._internal._request import build_path +from aura_python_sdk._internal._serde import parse_data, parse_data_list, to_json +from aura_python_sdk.models.graph_analytics import ( + DeletedGDSSession, + GDSSession, + GDSSessionConfig, + GDSSessionSizeEstimate, +) +from aura_python_sdk.services._base import Service + +_SESSIONS = "graph-analytics/sessions" + + +class GDSSessionService(Service): + """Graph Analytics (GDS) sessions.""" + + def list(self) -> builtins.list[GDSSession]: + """Every session the credentials can access.""" + self._logger.debug("listing GDS sessions") + sessions = parse_data_list(GDSSession, self._api.get(_SESSIONS).json()) + self._logger.debug("GDS sessions listed", extra={"count": len(sessions)}) + return sessions + + def estimate_size( + self, + *, + node_count: int, + relationship_count: int, + node_property_count: int | None = None, + node_label_count: int | None = None, + relationship_property_count: int | None = None, + algorithm_categories: Sequence[str] | None = None, + ) -> GDSSessionSizeEstimate: + """Estimate the session size needed for a graph (Go: ``Estimate``).""" + body: dict[str, object] = { + "node_count": validate.non_negative_int("node count", node_count), + "relationship_count": validate.non_negative_int( + "relationship count", relationship_count + ), + } + optional_counts = { + "node_property_count": node_property_count, + "node_label_count": node_label_count, + "relationship_property_count": relationship_property_count, + } + for key, value in optional_counts.items(): + if value is not None: + body[key] = validate.non_negative_int(key.replace("_", " "), value) + if algorithm_categories is not None: + if isinstance(algorithm_categories, str): + raise AuraValidationError("algorithm categories must be a sequence of strings") + body["algorithm_categories"] = [ + validate.require_non_empty("algorithm category", category) + for category in algorithm_categories + ] + + self._logger.debug("estimating GDS session size") + response = self._api.post(f"{_SESSIONS}/sizing", json_body=body) + return parse_data(GDSSessionSizeEstimate, response.json()) + + def create(self, config: GDSSessionConfig) -> GDSSession: + """Create a session, or return the matching existing one. + + Attach it to an instance with ``instance_id`` and ``database_uuid``, or make a standalone + session with ``cloud_provider`` and ``region``. + """ + if not isinstance(config, GDSSessionConfig): + raise AuraValidationError("config must be a GDSSessionConfig") + validate.require_non_empty("session name", config.name) + validate.require_non_empty("memory", config.memory) + if config.tenant_id is not None: + validate.tenant_id(config.tenant_id) + if config.instance_id is not None: + validate.instance_id(config.instance_id) + + self._logger.debug("creating GDS session", extra={"session_name": config.name}) + session = parse_data( + GDSSession, self._api.post(_SESSIONS, json_body=to_json(config)).json() + ) + self._logger.info("GDS session created", extra={"session_id": session.id}) + return session + + def get(self, session_id: str) -> GDSSession: + """Details of one session.""" + session_id = validate.session_id(session_id) + self._logger.debug("getting GDS session", extra={"session_id": session_id}) + return parse_data( + GDSSession, self._api.get(build_path("graph-analytics", "sessions", session_id)).json() + ) + + def delete(self, session_id: str) -> DeletedGDSSession: + """Delete a session.""" + session_id = validate.session_id(session_id) + self._logger.debug("deleting GDS session", extra={"session_id": session_id}) + response = self._api.delete(build_path("graph-analytics", "sessions", session_id)) + deleted = parse_data(DeletedGDSSession, response.json()) + self._logger.info("GDS session deleted", extra={"session_id": session_id}) + return deleted diff --git a/src/aura_python_sdk/services/instances.py b/src/aura_python_sdk/services/instances.py new file mode 100644 index 0000000..3cf95c3 --- /dev/null +++ b/src/aura_python_sdk/services/instances.py @@ -0,0 +1,181 @@ +"""``client.instances`` (Go: InstanceService).""" + +from __future__ import annotations + +import builtins + +from aura_python_sdk import _validation as validate +from aura_python_sdk._errors import AuraValidationError +from aura_python_sdk._internal._request import build_path +from aura_python_sdk._internal._serde import parse_data, parse_data_list, to_json +from aura_python_sdk.models.instances import ( + CDCEnrichmentMode, + CreatedInstance, + Instance, + InstanceConfig, + InstanceSummary, +) +from aura_python_sdk.services._base import Service + + +class InstanceService(Service): + """AuraDB and AuraDS instances.""" + + def list(self) -> builtins.list[InstanceSummary]: + """Every instance the credentials can access.""" + self._logger.debug("listing instances") + instances = parse_data_list(InstanceSummary, self._api.get("instances").json()) + self._logger.debug("instances listed", extra={"count": len(instances)}) + return instances + + def get(self, instance_id: str) -> Instance: + """Full details of one instance.""" + instance_id = validate.instance_id(instance_id) + self._logger.debug("getting instance", extra={"instance_id": instance_id}) + return parse_data(Instance, self._api.get(build_path("instances", instance_id)).json()) + + def create(self, config: InstanceConfig) -> CreatedInstance: + """Start creating an instance. + + Creation is asynchronous. Poll :meth:`get` until ``status`` is ``running``. The returned + password is shown only once. + """ + body = _create_body(config) + return self._create(body) + + def create_from_instance( + self, source_instance_id: str, config: InstanceConfig + ) -> CreatedInstance: + """Create an instance cloned from the current data of another instance.""" + source_instance_id = validate.instance_id(source_instance_id, "source instance ID") + body = _create_body(config) + body["source_instance_id"] = source_instance_id + return self._create(body) + + def create_from_snapshot( + self, source_instance_id: str, source_snapshot_id: str, config: InstanceConfig + ) -> CreatedInstance: + """Create an instance from a snapshot. + + The snapshot must belong to ``source_instance_id`` and be exportable. + """ + source_instance_id = validate.instance_id(source_instance_id, "source instance ID") + source_snapshot_id = validate.snapshot_id(source_snapshot_id, "source snapshot ID") + body = _create_body(config) + body["source_instance_id"] = source_instance_id + body["source_snapshot_id"] = source_snapshot_id + return self._create(body) + + def _create(self, body: dict[str, object]) -> CreatedInstance: + self._logger.debug( + "creating instance", + extra={"instance_name": body["name"], "tenant_id": body["tenant_id"]}, + ) + created = parse_data(CreatedInstance, self._api.post("instances", json_body=body).json()) + self._logger.info( + "instance creation started", + extra={"instance_id": created.id, "instance_name": created.name}, + ) + return created + + def update( + self, + instance_id: str, + *, + name: str | None = None, + memory: str | None = None, + cdc_enrichment_mode: CDCEnrichmentMode | str | None = None, + secondaries_count: int | None = None, + ) -> Instance: + """Rename, resize or reconfigure an instance. Only the arguments given are changed. + + The update is asynchronous, and the instance stays available throughout. + """ + instance_id = validate.instance_id(instance_id) + changes: dict[str, object] = {} + if name is not None: + changes["name"] = validate.instance_name(name) + if memory is not None: + changes["memory"] = validate.require_non_empty("memory", memory) + if cdc_enrichment_mode is not None: + changes["cdc_enrichment_mode"] = validate.require_non_empty( + "CDC enrichment mode", cdc_enrichment_mode + ) + if secondaries_count is not None: + changes["secondaries_count"] = validate.non_negative_int( + "secondaries count", secondaries_count + ) + if not changes: + raise AuraValidationError("update requires at least one field to change") + + self._logger.debug( + "updating instance", extra={"instance_id": instance_id, "fields": sorted(changes)} + ) + response = self._api.patch(build_path("instances", instance_id), json_body=to_json(changes)) + instance = parse_data(Instance, response.json()) + self._logger.info("instance update started", extra={"instance_id": instance_id}) + return instance + + def delete(self, instance_id: str) -> Instance: + """Start deleting an instance. This cannot be undone.""" + instance_id = validate.instance_id(instance_id) + self._logger.debug("deleting instance", extra={"instance_id": instance_id}) + instance = parse_data( + Instance, self._api.delete(build_path("instances", instance_id)).json() + ) + self._logger.info("instance deletion started", extra={"instance_id": instance_id}) + return instance + + def pause(self, instance_id: str) -> Instance: + """Pause a running instance.""" + return self._lifecycle(instance_id, "pause") + + def resume(self, instance_id: str) -> Instance: + """Resume a paused instance.""" + return self._lifecycle(instance_id, "resume") + + def _lifecycle(self, instance_id: str, action: str) -> Instance: + instance_id = validate.instance_id(instance_id) + self._logger.debug("%s instance", action, extra={"instance_id": instance_id}) + response = self._api.post(build_path("instances", instance_id, action)) + instance = parse_data(Instance, response.json()) + self._logger.info("instance %s started", action, extra={"instance_id": instance_id}) + return instance + + def overwrite_from_instance(self, instance_id: str, source_instance_id: str) -> Instance: + """Replace an instance's data with the current data of another instance.""" + instance_id = validate.instance_id(instance_id) + source_instance_id = validate.instance_id(source_instance_id, "source instance ID") + return self._overwrite(instance_id, {"source_instance_id": source_instance_id}) + + def overwrite_from_snapshot(self, instance_id: str, source_snapshot_id: str) -> Instance: + """Replace an instance's data with a snapshot.""" + instance_id = validate.instance_id(instance_id) + source_snapshot_id = validate.snapshot_id(source_snapshot_id, "source snapshot ID") + return self._overwrite(instance_id, {"source_snapshot_id": source_snapshot_id}) + + def _overwrite(self, instance_id: str, body: dict[str, object]) -> Instance: + self._logger.debug("overwriting instance", extra={"instance_id": instance_id, **body}) + response = self._api.post(build_path("instances", instance_id, "overwrite"), json_body=body) + instance = parse_data(Instance, response.json()) + self._logger.info("instance overwrite started", extra={"instance_id": instance_id}) + return instance + + +def _create_body(config: InstanceConfig) -> dict[str, object]: + """Validate a create request as the Go SDK's validateCreateInstanceConfig does.""" + if not isinstance(config, InstanceConfig): + raise AuraValidationError("config must be an InstanceConfig") + validate.instance_name(config.name) + validate.tenant_id(config.tenant_id) + validate.require_non_empty("cloud provider", config.cloud_provider) + validate.require_non_empty("region", config.region) + validate.require_non_empty("instance type", config.type) + validate.require_non_empty("version", config.version) + validate.require_non_empty("memory", config.memory) + if config.customer_managed_key_id is not None: + validate.require_non_empty("customer managed key ID", config.customer_managed_key_id) + body = to_json(config) + if not isinstance(body, dict): # pragma: no cover - to_json of a dataclass is a dict + raise TypeError("expected a JSON object") + return body diff --git a/src/aura_python_sdk/services/snapshots.py b/src/aura_python_sdk/services/snapshots.py new file mode 100644 index 0000000..88b1315 --- /dev/null +++ b/src/aura_python_sdk/services/snapshots.py @@ -0,0 +1,71 @@ +"""``client.snapshots`` (Go: SnapshotService).""" + +from __future__ import annotations + +import builtins +import datetime as dt + +from aura_python_sdk import _validation as validate +from aura_python_sdk._errors import AuraValidationError +from aura_python_sdk._internal._request import build_path +from aura_python_sdk._internal._serde import parse_data, parse_data_list +from aura_python_sdk.models.instances import Instance +from aura_python_sdk.models.snapshots import CreatedSnapshot, Snapshot +from aura_python_sdk.services._base import Service + + +class SnapshotService(Service): + """Instance snapshots.""" + + def list(self, instance_id: str, date: dt.date | None = None) -> builtins.list[Snapshot]: + """Snapshots of an instance taken on ``date``. The API defaults to today.""" + instance_id = validate.instance_id(instance_id) + if date is not None and (not isinstance(date, dt.date) or isinstance(date, dt.datetime)): + raise AuraValidationError("date must be a datetime.date") + self._logger.debug("listing snapshots", extra={"instance_id": instance_id}) + response = self._api.get( + build_path("instances", instance_id, "snapshots"), + params={"date": date.isoformat() if date else None}, + ) + snapshots = parse_data_list(Snapshot, response.json()) + self._logger.debug("snapshots listed", extra={"count": len(snapshots)}) + return snapshots + + def get(self, instance_id: str, snapshot_id: str) -> Snapshot: + """Details of one snapshot.""" + instance_id = validate.instance_id(instance_id) + snapshot_id = validate.snapshot_id(snapshot_id) + self._logger.debug( + "getting snapshot", extra={"instance_id": instance_id, "snapshot_id": snapshot_id} + ) + response = self._api.get(build_path("instances", instance_id, "snapshots", snapshot_id)) + return parse_data(Snapshot, response.json()) + + def create(self, instance_id: str) -> CreatedSnapshot: + """Start an on-demand snapshot.""" + instance_id = validate.instance_id(instance_id) + self._logger.debug("creating snapshot", extra={"instance_id": instance_id}) + response = self._api.post(build_path("instances", instance_id, "snapshots")) + created = parse_data(CreatedSnapshot, response.json()) + self._logger.info( + "snapshot started", + extra={"instance_id": instance_id, "snapshot_id": created.snapshot_id}, + ) + return created + + def restore(self, instance_id: str, snapshot_id: str) -> Instance: + """Restore an instance from one of its own snapshots, replacing its current data.""" + instance_id = validate.instance_id(instance_id) + snapshot_id = validate.snapshot_id(snapshot_id) + self._logger.debug( + "restoring snapshot", extra={"instance_id": instance_id, "snapshot_id": snapshot_id} + ) + response = self._api.post( + build_path("instances", instance_id, "snapshots", snapshot_id, "restore") + ) + instance = parse_data(Instance, response.json()) + self._logger.info( + "snapshot restore started", + extra={"instance_id": instance_id, "snapshot_id": snapshot_id}, + ) + return instance diff --git a/src/aura_python_sdk/services/tenants.py b/src/aura_python_sdk/services/tenants.py new file mode 100644 index 0000000..14e10c9 --- /dev/null +++ b/src/aura_python_sdk/services/tenants.py @@ -0,0 +1,35 @@ +"""``client.tenants`` (Go: TenantService).""" + +from __future__ import annotations + +import builtins + +from aura_python_sdk import _validation as validate +from aura_python_sdk._internal._request import build_path +from aura_python_sdk._internal._serde import parse_data, parse_data_list +from aura_python_sdk.models.tenants import MetricsIntegration, Tenant, TenantSummary +from aura_python_sdk.services._base import Service + + +class TenantService(Service): + """Tenants (shown as projects in the Aura Console).""" + + def list(self) -> builtins.list[TenantSummary]: + """Every tenant the credentials can access.""" + self._logger.debug("listing tenants") + tenants = parse_data_list(TenantSummary, self._api.get("tenants").json()) + self._logger.debug("tenants listed", extra={"count": len(tenants)}) + return tenants + + def get(self, tenant_id: str) -> Tenant: + """A tenant and the instance configurations it can create.""" + tenant_id = validate.tenant_id(tenant_id) + self._logger.debug("getting tenant", extra={"tenant_id": tenant_id}) + return parse_data(Tenant, self._api.get(build_path("tenants", tenant_id)).json()) + + def get_metrics_integration(self, tenant_id: str) -> MetricsIntegration: + """The project-level Prometheus metrics endpoint (Go: ``GetMetrics``).""" + tenant_id = validate.tenant_id(tenant_id) + self._logger.debug("getting tenant metrics integration", extra={"tenant_id": tenant_id}) + response = self._api.get(build_path("tenants", tenant_id, "metrics-integration")) + return parse_data(MetricsIntegration, response.json()) diff --git a/tests/unit/conftest.py b/tests/unit/conftest.py new file mode 100644 index 0000000..59ef5a4 --- /dev/null +++ b/tests/unit/conftest.py @@ -0,0 +1,74 @@ +from __future__ import annotations + +import json +from dataclasses import dataclass +from typing import Any + +import pytest + +from aura_python_sdk import AuraClient, HttpRequest +from tests.fakes import FakeTransport, json_response, token_response + +TENANT_ID = "6981ace7-efe8-4f5c-b7c5-267b5162ce91" +INSTANCE_ID = "2f49c2b3" +OTHER_INSTANCE_ID = "b51dc964" +SNAPSHOT_ID = "e9ac0fa5-e1f9-4bb2-b0a2-5d4e5b3d8b43" +BASE = "https://api.neo4j.io/v1" + +INSTANCE = { + "id": INSTANCE_ID, + "name": "Production", + "status": "running", + "tenant_id": TENANT_ID, + "cloud_provider": "gcp", + "connection_url": "neo4j+s://2f49c2b3.databases.neo4j.io", + "region": "europe-west1", + "type": "enterprise-db", + "memory": "8GB", + "storage": "16GB", + "created_at": "2023-01-20T13:44:42Z", +} + +SESSION = { + "id": "s-04de43fe-67ab-4", + "name": "people-and-fruit", + "memory": "8GB", + "host": "s-04de43fe-67ab-4-gds.example.neo4j.io", + "tenant_id": TENANT_ID, + "user_id": "user-1", + "status": "Ready", + "ttl": "20m0s", +} + + +@dataclass +class Api: + """An AuraClient wired to a FakeTransport that already holds a token response.""" + + client: AuraClient + transport: FakeTransport + + def reply(self, status: int, payload: object) -> None: + self.transport.queue(json_response(status, payload)) + + @property + def request(self) -> HttpRequest: + """The single API request sent (excluding the token request).""" + requests = self.transport.api_requests + assert len(requests) == 1, f"expected one API request, got {len(requests)}" + return requests[0] + + @property + def body(self) -> Any: + body = self.request.body + return None if body is None else json.loads(body) + + def assert_no_request(self) -> None: + assert self.transport.requests == [] + + +@pytest.fixture +def api() -> Api: + transport = FakeTransport([token_response()]) + client = AuraClient(client_id="id", client_secret="secret", transport=transport) + return Api(client, transport) diff --git a/tests/unit/test_cmek_service.py b/tests/unit/test_cmek_service.py new file mode 100644 index 0000000..76831bb --- /dev/null +++ b/tests/unit/test_cmek_service.py @@ -0,0 +1,24 @@ +import pytest + +from aura_python_sdk import AuraValidationError, CustomerManagedKeySummary +from tests.unit.conftest import BASE, TENANT_ID, Api + +KEY = {"id": "8c764aed-8eb3-4a1c-92f6-e4ef0c7a6ed9", "name": "Key01", "tenant_id": TENANT_ID} + + +def test_list(api: Api) -> None: + api.reply(200, {"data": [KEY]}) + assert api.client.cmek.list() == [CustomerManagedKeySummary(**KEY)] + assert api.request.url == f"{BASE}/customer-managed-keys" + + +def test_list_filtered_by_tenant(api: Api) -> None: + api.reply(200, {"data": [KEY]}) + api.client.cmek.list(TENANT_ID) + assert api.request.url == f"{BASE}/customer-managed-keys?tenantId={TENANT_ID}" + + +def test_list_invalid_tenant_sends_nothing(api: Api) -> None: + with pytest.raises(AuraValidationError, match="tenant ID"): + api.client.cmek.list("bad") + api.assert_no_request() diff --git a/tests/unit/test_graph_analytics_service.py b/tests/unit/test_graph_analytics_service.py new file mode 100644 index 0000000..3db9459 --- /dev/null +++ b/tests/unit/test_graph_analytics_service.py @@ -0,0 +1,149 @@ +from typing import Any + +import pytest + +from aura_python_sdk import ( + AuraValidationError, + CloudProvider, + DeletedGDSSession, + GDSSessionConfig, + GDSSessionSizeEstimate, + GDSSessionStatus, +) +from tests.unit.conftest import BASE, INSTANCE_ID, SESSION, TENANT_ID, Api + +SESSIONS = f"{BASE}/graph-analytics/sessions" + + +def test_list(api: Api) -> None: + api.reply(200, {"data": [SESSION]}) + [session] = api.client.graph_analytics.list() + assert session.status is GDSSessionStatus.READY + assert api.request.url == SESSIONS + + +def test_get(api: Api) -> None: + api.reply(200, {"data": SESSION}) + assert api.client.graph_analytics.get(SESSION["id"]).ttl == "20m0s" + assert api.request.url == f"{SESSIONS}/{SESSION['id']}" + + +def test_session_id_is_path_encoded(api: Api) -> None: + api.reply(200, {"data": SESSION}) + api.client.graph_analytics.get("../instances") + assert api.request.url == f"{SESSIONS}/..%2Finstances" + + +def test_delete(api: Api) -> None: + api.reply(202, {"data": {"id": SESSION["id"]}}) + assert api.client.graph_analytics.delete(SESSION["id"]) == DeletedGDSSession(id=SESSION["id"]) + assert (api.request.method, api.request.url) == ("DELETE", f"{SESSIONS}/{SESSION['id']}") + + +@pytest.mark.parametrize("call", ["get", "delete"]) +def test_empty_session_id(api: Api, call: str) -> None: + with pytest.raises(AuraValidationError, match="GDS session ID must not be empty"): + getattr(api.client.graph_analytics, call)("") + api.assert_no_request() + + +@pytest.mark.parametrize("status", [200, 202]) +def test_create(api: Api, status: int) -> None: + api.reply(status, {"data": SESSION}) + config = GDSSessionConfig( + name="people-and-fruit", + memory="8GB", + ttl="1h", + tenant_id=TENANT_ID, + cloud_provider=CloudProvider.AZURE, + region="francecentral", + ) + session = api.client.graph_analytics.create(config) + assert session.id == SESSION["id"] + assert (api.request.method, api.request.url) == ("POST", SESSIONS) + assert api.body == { + "name": "people-and-fruit", + "memory": "8GB", + "ttl": "1h", + "tenant_id": TENANT_ID, + "cloud_provider": "azure", + "region": "francecentral", + } + + +@pytest.mark.parametrize( + ("config", "message"), + [ + ({"name": "people", "memory": "8GB", "a": 1}, "config must be a GDSSessionConfig"), + (GDSSessionConfig(name="", memory="8GB"), "session name must not be empty"), + (GDSSessionConfig(name="s", memory=""), "memory must not be empty"), + (GDSSessionConfig(name="s", memory="8GB", tenant_id="bad"), "tenant ID"), + (GDSSessionConfig(name="s", memory="8GB", instance_id="bad"), "instance ID"), + ], +) +def test_create_validation(api: Api, config: Any, message: str) -> None: + with pytest.raises(AuraValidationError, match=message): + api.client.graph_analytics.create(config) + api.assert_no_request() + + +def test_create_attached_to_instance(api: Api) -> None: + api.reply(202, {"data": SESSION}) + config = GDSSessionConfig( + name="s", + memory="4GB", + instance_id=INSTANCE_ID, + database_uuid="ea408a62-c991-490c-96db-2b947003eece", + ) + api.client.graph_analytics.create(config) + assert api.body["instance_id"] == INSTANCE_ID + + +def test_estimate_size(api: Api) -> None: + api.reply(200, {"data": {"estimated_memory": "6GB", "recommended_size": "8GB"}}) + estimate = api.client.graph_analytics.estimate_size( + node_count=1_000_000, + relationship_count=5_000_000, + node_property_count=512, + node_label_count=3, + relationship_property_count=5, + algorithm_categories=["similarity", "community-detection"], + ) + assert estimate == GDSSessionSizeEstimate(estimated_memory="6GB", recommended_size="8GB") + assert (api.request.method, api.request.url) == ("POST", f"{SESSIONS}/sizing") + assert api.body == { + "node_count": 1_000_000, + "relationship_count": 5_000_000, + "node_property_count": 512, + "node_label_count": 3, + "relationship_property_count": 5, + "algorithm_categories": ["similarity", "community-detection"], + } + + +def test_estimate_size_minimal(api: Api) -> None: + api.reply(200, {"data": {"estimated_memory": "1GB", "recommended_size": "2GB"}}) + api.client.graph_analytics.estimate_size(node_count=10, relationship_count=0) + assert api.body == {"node_count": 10, "relationship_count": 0} + + +@pytest.mark.parametrize( + ("kwargs", "message"), + [ + ({"node_count": -1, "relationship_count": 1}, "node count"), + ({"node_count": 1, "relationship_count": 1.5}, "relationship count"), + ({"node_count": 1, "relationship_count": 1, "node_label_count": -3}, "node label count"), + ( + {"node_count": 1, "relationship_count": 1, "algorithm_categories": "similarity"}, + "sequence", + ), + ( + {"node_count": 1, "relationship_count": 1, "algorithm_categories": [""]}, + "algorithm category", + ), + ], +) +def test_estimate_size_validation(api: Api, kwargs: dict[str, Any], message: str) -> None: + with pytest.raises(AuraValidationError, match=message): + api.client.graph_analytics.estimate_size(**kwargs) + api.assert_no_request() diff --git a/tests/unit/test_instances_service.py b/tests/unit/test_instances_service.py new file mode 100644 index 0000000..f8c910b --- /dev/null +++ b/tests/unit/test_instances_service.py @@ -0,0 +1,273 @@ +from dataclasses import replace +from typing import Any + +import pytest + +from aura_python_sdk import ( + AuraValidationError, + CDCEnrichmentMode, + CloudProvider, + ConflictError, + CreatedInstance, + Instance, + InstanceConfig, + InstanceStatus, + InstanceSummary, + InstanceType, + NotFoundError, +) +from tests.unit.conftest import ( + BASE, + INSTANCE, + INSTANCE_ID, + OTHER_INSTANCE_ID, + SNAPSHOT_ID, + TENANT_ID, + Api, +) + +CONFIG = InstanceConfig( + name="Instance01", + tenant_id=TENANT_ID, + cloud_provider=CloudProvider.GCP, + region="europe-west1", + type=InstanceType.ENTERPRISE_DB, + version="5", + memory="8GB", +) + +CONFIG_JSON = { + "name": "Instance01", + "tenant_id": TENANT_ID, + "cloud_provider": "gcp", + "region": "europe-west1", + "type": "enterprise-db", + "version": "5", + "memory": "8GB", +} + +CREATED = { + "id": "db1d1234", + "name": "Instance01", + "tenant_id": TENANT_ID, + "cloud_provider": "gcp", + "region": "europe-west1", + "type": "enterprise-db", + "connection_url": "neo4j+s://db1d1234.databases.neo4j.io", + "username": "neo4j", + "password": "letMeIn123!", + "created_at": "2023-01-20T13:44:42Z", +} + + +def test_list(api: Api) -> None: + api.reply( + 200, + { + "data": [ + {"id": INSTANCE_ID, "name": "P", "tenant_id": TENANT_ID, "cloud_provider": "aws"} + ] + }, + ) + [summary] = api.client.instances.list() + assert isinstance(summary, InstanceSummary) + assert summary.cloud_provider is CloudProvider.AWS + assert (api.request.method, api.request.url) == ("GET", f"{BASE}/instances") + + +def test_get(api: Api) -> None: + api.reply(200, {"data": INSTANCE}) + instance = api.client.instances.get(INSTANCE_ID) + assert isinstance(instance, Instance) + assert instance.status is InstanceStatus.RUNNING + assert api.request.url == f"{BASE}/instances/{INSTANCE_ID}" + + +def test_get_not_found(api: Api) -> None: + api.reply(404, {"errors": [{"message": "Instance not found", "reason": "instance-not-found"}]}) + with pytest.raises(NotFoundError, match="Instance not found"): + api.client.instances.get(INSTANCE_ID) + + +def test_create(api: Api) -> None: + api.reply(202, {"data": CREATED}) + created = api.client.instances.create(CONFIG) + + assert isinstance(created, CreatedInstance) + assert created.password == "letMeIn123!" + assert (api.request.method, api.request.url) == ("POST", f"{BASE}/instances") + assert api.body == CONFIG_JSON + + +def test_create_sends_optional_fields_when_set(api: Api) -> None: + api.reply(202, {"data": CREATED}) + key_id = "8c764aed-8eb3-4a1c-92f6-e4ef0c7a6ed9" + config = replace( + CONFIG, vector_optimized=True, graph_analytics_plugin=False, customer_managed_key_id=key_id + ) + api.client.instances.create(config) + assert api.body == { + **CONFIG_JSON, + "vector_optimized": True, + "graph_analytics_plugin": False, + "customer_managed_key_id": key_id, + } + + +def test_create_from_instance(api: Api) -> None: + api.reply(202, {"data": CREATED}) + api.client.instances.create_from_instance(OTHER_INSTANCE_ID, CONFIG) + assert api.body == {**CONFIG_JSON, "source_instance_id": OTHER_INSTANCE_ID} + + +def test_create_from_snapshot(api: Api) -> None: + api.reply(202, {"data": CREATED}) + api.client.instances.create_from_snapshot(OTHER_INSTANCE_ID, SNAPSHOT_ID, CONFIG) + assert api.body == { + **CONFIG_JSON, + "source_instance_id": OTHER_INSTANCE_ID, + "source_snapshot_id": SNAPSHOT_ID, + } + + +@pytest.mark.parametrize( + ("overrides", "message"), + [ + ({"name": ""}, "instance name must not be empty"), + ({"name": "x" * 31}, "at most 30 characters"), + ({"tenant_id": ""}, "tenant ID must not be empty"), + ({"tenant_id": "abc"}, "tenant ID must be a valid UUID"), + ({"cloud_provider": ""}, "cloud provider must not be empty"), + ({"region": ""}, "region must not be empty"), + ({"type": ""}, "instance type must not be empty"), + ({"version": ""}, "version must not be empty"), + ({"memory": ""}, "memory must not be empty"), + ({"customer_managed_key_id": ""}, "customer managed key ID must not be empty"), + ], +) +def test_create_validation_matches_go(api: Api, overrides: dict[str, Any], message: str) -> None: + with pytest.raises(AuraValidationError, match=message): + api.client.instances.create(replace(CONFIG, **overrides)) + api.assert_no_request() + + +def test_create_requires_instance_config(api: Api) -> None: + with pytest.raises(AuraValidationError, match="config must be an InstanceConfig"): + api.client.instances.create(CONFIG_JSON) # type: ignore[arg-type] + api.assert_no_request() + + +@pytest.mark.parametrize( + "call", + [ + lambda s: s.create_from_instance("", CONFIG), + lambda s: s.create_from_instance("bad", CONFIG), + lambda s: s.create_from_snapshot("bad", SNAPSHOT_ID, CONFIG), + lambda s: s.create_from_snapshot(OTHER_INSTANCE_ID, "bad", CONFIG), + lambda s: s.create_from_snapshot(OTHER_INSTANCE_ID, "", CONFIG), + ], +) +def test_create_from_source_validation(api: Api, call: Any) -> None: + with pytest.raises(AuraValidationError, match=r"source (instance|snapshot) ID"): + call(api.client.instances) + api.assert_no_request() + + +def test_update(api: Api) -> None: + api.reply(202, {"data": {**INSTANCE, "status": "updating"}}) + instance = api.client.instances.update( + INSTANCE_ID, + name="Renamed", + memory="16GB", + cdc_enrichment_mode=CDCEnrichmentMode.FULL, + secondaries_count=2, + ) + assert instance.status is InstanceStatus.UPDATING + assert (api.request.method, api.request.url) == ("PATCH", f"{BASE}/instances/{INSTANCE_ID}") + assert api.body == { + "name": "Renamed", + "memory": "16GB", + "cdc_enrichment_mode": "FULL", + "secondaries_count": 2, + } + + +def test_update_sends_only_given_fields(api: Api) -> None: + api.reply(200, {"data": INSTANCE}) + api.client.instances.update(INSTANCE_ID, secondaries_count=0) + assert api.body == {"secondaries_count": 0} + + +@pytest.mark.parametrize( + ("kwargs", "message"), + [ + ({}, "at least one field"), + ({"name": "x" * 31}, "at most 30 characters"), + ({"memory": ""}, "memory must not be empty"), + ({"cdc_enrichment_mode": ""}, "CDC enrichment mode"), + ({"secondaries_count": -1}, "secondaries count"), + ], +) +def test_update_validation(api: Api, kwargs: dict[str, Any], message: str) -> None: + with pytest.raises(AuraValidationError, match=message): + api.client.instances.update(INSTANCE_ID, **kwargs) + api.assert_no_request() + + +def test_delete(api: Api) -> None: + api.reply(202, {"data": {**INSTANCE, "status": "destroying"}}) + instance = api.client.instances.delete(INSTANCE_ID) + assert instance.status is InstanceStatus.DESTROYING + assert (api.request.method, api.request.url) == ("DELETE", f"{BASE}/instances/{INSTANCE_ID}") + assert api.request.body is None + + +@pytest.mark.parametrize(("action", "status"), [("pause", "pausing"), ("resume", "resuming")]) +def test_pause_and_resume(api: Api, action: str, status: str) -> None: + api.reply(202, {"data": {**INSTANCE, "status": status}}) + instance = getattr(api.client.instances, action)(INSTANCE_ID) + assert instance.status == status + assert (api.request.method, api.request.url) == ( + "POST", + f"{BASE}/instances/{INSTANCE_ID}/{action}", + ) + assert api.request.body is None + + +def test_pause_conflict(api: Api) -> None: + api.reply(409, {"errors": [{"message": "Instance is not running", "reason": "conflict"}]}) + with pytest.raises(ConflictError): + api.client.instances.pause(INSTANCE_ID) + + +def test_overwrite_from_instance(api: Api) -> None: + api.reply(202, {"data": {**INSTANCE, "status": "overwriting"}}) + instance = api.client.instances.overwrite_from_instance(INSTANCE_ID, OTHER_INSTANCE_ID) + assert instance.status is InstanceStatus.OVERWRITING + assert api.request.url == f"{BASE}/instances/{INSTANCE_ID}/overwrite" + assert api.body == {"source_instance_id": OTHER_INSTANCE_ID} + + +def test_overwrite_from_snapshot(api: Api) -> None: + api.reply(202, {"data": {**INSTANCE, "status": "overwriting"}}) + api.client.instances.overwrite_from_snapshot(INSTANCE_ID, SNAPSHOT_ID) + assert api.body == {"source_snapshot_id": SNAPSHOT_ID} + + +@pytest.mark.parametrize( + "call", + [ + lambda s: s.get("nope"), + lambda s: s.delete(""), + lambda s: s.pause("../../x"), + lambda s: s.resume("12345"), + lambda s: s.update("bad", name="x"), + lambda s: s.overwrite_from_instance("bad", OTHER_INSTANCE_ID), + lambda s: s.overwrite_from_instance(INSTANCE_ID, "bad"), + lambda s: s.overwrite_from_snapshot(INSTANCE_ID, "bad"), + ], +) +def test_invalid_ids_send_nothing(api: Api, call: Any) -> None: + with pytest.raises(AuraValidationError): + call(api.client.instances) + api.assert_no_request() diff --git a/tests/unit/test_snapshots_service.py b/tests/unit/test_snapshots_service.py new file mode 100644 index 0000000..f30a781 --- /dev/null +++ b/tests/unit/test_snapshots_service.py @@ -0,0 +1,82 @@ +import datetime as dt +from typing import Any + +import pytest + +from aura_python_sdk import AuraValidationError, InstanceStatus, SnapshotProfile, SnapshotStatus +from tests.unit.conftest import BASE, INSTANCE, INSTANCE_ID, SNAPSHOT_ID, Api + +SNAPSHOT = { + "instance_id": INSTANCE_ID, + "snapshot_id": SNAPSHOT_ID, + "profile": "AdHoc", + "status": "Completed", + "timestamp": "2023-01-20T13:44:42Z", + "exportable": True, +} + + +def test_list(api: Api) -> None: + api.reply(200, {"data": [SNAPSHOT]}) + [snapshot] = api.client.snapshots.list(INSTANCE_ID) + assert snapshot.status is SnapshotStatus.COMPLETED + assert snapshot.profile is SnapshotProfile.AD_HOC + assert api.request.url == f"{BASE}/instances/{INSTANCE_ID}/snapshots" + + +def test_list_with_date(api: Api) -> None: + api.reply(200, {"data": []}) + assert api.client.snapshots.list(INSTANCE_ID, dt.date(2024, 3, 7)) == [] + assert api.request.url == f"{BASE}/instances/{INSTANCE_ID}/snapshots?date=2024-03-07" + + +@pytest.mark.parametrize("value", ["2024-03-07", dt.datetime(2024, 3, 7, tzinfo=dt.UTC)]) +def test_list_rejects_non_date(api: Api, value: object) -> None: + with pytest.raises(AuraValidationError, match=r"datetime\.date"): + api.client.snapshots.list(INSTANCE_ID, value) # type: ignore[arg-type] + api.assert_no_request() + + +def test_get(api: Api) -> None: + api.reply(200, {"data": SNAPSHOT}) + snapshot = api.client.snapshots.get(INSTANCE_ID, SNAPSHOT_ID) + assert snapshot.snapshot_id == SNAPSHOT_ID + assert snapshot.timestamp == dt.datetime(2023, 1, 20, 13, 44, 42, tzinfo=dt.UTC) + assert api.request.url == f"{BASE}/instances/{INSTANCE_ID}/snapshots/{SNAPSHOT_ID}" + + +def test_create(api: Api) -> None: + api.reply(202, {"data": {"snapshot_id": SNAPSHOT_ID}}) + created = api.client.snapshots.create(INSTANCE_ID) + assert created.snapshot_id == SNAPSHOT_ID + assert (api.request.method, api.request.url) == ( + "POST", + f"{BASE}/instances/{INSTANCE_ID}/snapshots", + ) + assert api.request.body is None + + +def test_restore(api: Api) -> None: + api.reply(202, {"data": {**INSTANCE, "status": "restoring"}}) + instance = api.client.snapshots.restore(INSTANCE_ID, SNAPSHOT_ID) + assert instance.status is InstanceStatus.RESTORING + assert (api.request.method, api.request.url) == ( + "POST", + f"{BASE}/instances/{INSTANCE_ID}/snapshots/{SNAPSHOT_ID}/restore", + ) + + +@pytest.mark.parametrize( + "call", + [ + lambda s: s.list("bad"), + lambda s: s.get("bad", SNAPSHOT_ID), + lambda s: s.get(INSTANCE_ID, "bad"), + lambda s: s.create(""), + lambda s: s.restore(INSTANCE_ID, "2023-01-20T13:44:42Z"), + ], +) +def test_invalid_ids_send_nothing(api: Api, call: Any) -> None: + with pytest.raises(AuraValidationError): + call(api.client.snapshots) + api.assert_no_request() diff --git a/tests/unit/test_tenants_service.py b/tests/unit/test_tenants_service.py new file mode 100644 index 0000000..a33936a --- /dev/null +++ b/tests/unit/test_tenants_service.py @@ -0,0 +1,47 @@ +import pytest + +from aura_python_sdk import AuraValidationError, InstanceType, Tenant, TenantSummary +from tests.unit.conftest import BASE, TENANT_ID, Api + + +def test_list(api: Api) -> None: + api.reply(200, {"data": [{"id": TENANT_ID, "name": "Production"}]}) + assert api.client.tenants.list() == [TenantSummary(id=TENANT_ID, name="Production")] + assert (api.request.method, api.request.url) == ("GET", f"{BASE}/tenants") + + +def test_get(api: Api) -> None: + config = { + "cloud_provider": "gcp", + "region": "europe-west1", + "region_name": "Belgium (europe-west1)", + "type": "enterprise-db", + "memory": "8GB", + "storage": "16GB", + "version": "5", + } + api.reply( + 200, {"data": {"id": TENANT_ID, "name": "Production", "instance_configurations": [config]}} + ) + + tenant = api.client.tenants.get(TENANT_ID) + + assert isinstance(tenant, Tenant) + assert tenant.instance_configurations[0].type is InstanceType.ENTERPRISE_DB + assert api.request.url == f"{BASE}/tenants/{TENANT_ID}" + + +def test_get_metrics_integration(api: Api) -> None: + api.reply( + 200, {"data": {"endpoint": "https://customer-metrics-api.neo4j.io/api/v1/abc/metrics"}} + ) + result = api.client.tenants.get_metrics_integration(TENANT_ID) + assert result.endpoint.endswith("/metrics") + assert api.request.url == f"{BASE}/tenants/{TENANT_ID}/metrics-integration" + + +@pytest.mark.parametrize("call", ["get", "get_metrics_integration"]) +def test_invalid_tenant_id_sends_nothing(api: Api, call: str) -> None: + with pytest.raises(AuraValidationError, match="tenant ID"): + getattr(api.client.tenants, call)("not-a-uuid") + api.assert_no_request() diff --git a/tests/unit/test_validation.py b/tests/unit/test_validation.py new file mode 100644 index 0000000..031e774 --- /dev/null +++ b/tests/unit/test_validation.py @@ -0,0 +1,61 @@ +import pytest + +from aura_python_sdk import AuraValidationError +from aura_python_sdk import _validation as validate + + +@pytest.mark.parametrize("value", ["2f49c2b3", "ABCDEF12"]) +def test_valid_instance_ids(value: str) -> None: + assert validate.instance_id(value) == value + + +@pytest.mark.parametrize( + "value", ["", " ", None, 12345678, "2f49c2b", "2f49c2b33", "zzzzzzzz", "../abcde"] +) +def test_invalid_instance_ids(value: object) -> None: + with pytest.raises(AuraValidationError, match="instance ID"): + validate.instance_id(value) + + +def test_instance_id_custom_name_in_message() -> None: + with pytest.raises(AuraValidationError, match="source instance ID must not be empty"): + validate.instance_id("", "source instance ID") + + +@pytest.mark.parametrize("check", [validate.tenant_id, validate.snapshot_id]) +def test_uuid_ids(check: object) -> None: + assert check("6981ace7-efe8-4f5c-b7c5-267b5162ce91") == "6981ace7-efe8-4f5c-b7c5-267b5162ce91" # type: ignore[operator] + for bad in ["", "6981ace7", "6981ace7-efe8-4f5c-b7c5-267b5162ce9Z", "2023-01-20T13:44:42Z"]: + with pytest.raises(AuraValidationError, match="ID"): + check(bad) # type: ignore[operator] + + +def test_uuid_error_message_matches_go() -> None: + with pytest.raises(AuraValidationError) as info: + validate.tenant_id("nope") + assert str(info.value) == ( + "tenant ID must be a valid UUID format (xxxxxxxx-xxxx-xxxx-xxxx-xxxxxxxxxxxx)" + ) + + +def test_session_id_only_needs_to_be_non_empty() -> None: + assert validate.session_id("s-04de43fe-67ab-4") == "s-04de43fe-67ab-4" + with pytest.raises(AuraValidationError, match="GDS session ID must not be empty"): + validate.session_id("") + + +def test_instance_name_length() -> None: + assert validate.instance_name("x" * 30) == "x" * 30 + with pytest.raises(AuraValidationError, match="at most 30 characters"): + validate.instance_name("x" * 31) + + +@pytest.mark.parametrize("value", [-1, 1.5, True, "3"]) +def test_non_negative_int(value: object) -> None: + assert validate.non_negative_int("count", 0) == 0 + with pytest.raises(AuraValidationError, match="count must be an integer"): + validate.non_negative_int("count", value) + + +def test_validation_error_is_value_error() -> None: + assert issubclass(AuraValidationError, ValueError) From cdc98b4d9d4f55855c30c2719a51a73a4e8e1439 Mon Sep 17 00:00:00 2001 From: Jonathan Giffard Date: Tue, 29 Sep 2026 20:23:30 +0100 Subject: [PATCH 4/8] Cover the full v1 spec: sizing, upgrade, CMEK CRUD and list filters Phase 5 of PLAN.md. Adds instance sizing and upgrade, CMEK get/create/delete, tenant/instance/organization list filters, and the storage, vector and graph analytics update fields. A test now maps every spec operation to a method. Co-Authored-By: Claude Opus 5.5 --- PLAN.md | 12 ++ src/aura_python_sdk/_validation.py | 33 ++++- src/aura_python_sdk/services/_base.py | 5 + src/aura_python_sdk/services/cmek.py | 56 ++++++- .../services/graph_analytics.py | 35 +++-- src/aura_python_sdk/services/instances.py | 79 +++++++++- tests/unit/test_cmek_service.py | 114 +++++++++++++- tests/unit/test_graph_analytics_service.py | 26 +++- tests/unit/test_instances_service.py | 140 ++++++++++++++++++ tests/unit/test_spec_examples.py | 49 ++++++ 10 files changed, 519 insertions(+), 30 deletions(-) diff --git a/PLAN.md b/PLAN.md index fc3f2c7..b3df9e4 100644 --- a/PLAN.md +++ b/PLAN.md @@ -203,6 +203,18 @@ parser accepts the spec's `{"errors": [...]}`, the middleware `{"error": "..."}` keep their usual types. A lower-case `bearer` token type is accepted. - **Transport ownership.** `close()` closes only a transport the client created itself. +### 2.8 Decisions made in phase 5 + +- **Names**: instance and CMEK names must be 1–30 characters with no leading or trailing + whitespace, as the spec states. Go checks only length, and only on create. +- **CMEK IDs**: `get` and `delete` only require a non-empty key ID, which is then path-encoded. + The spec doesn't say whether these IDs are UUIDs. +- **`upgrade()`**: `memory` and `storage` must be given together or not at all, as the spec + requires. With neither, it sends `{}`. +- **`cmek.delete()`** returns `None`, since the API responds 204 with no body. +- **Coverage guard**: `test_every_spec_operation_has_a_client_method` fails if the spec gains an + operation that no SDK method covers. + ## 3. Package layout ``` diff --git a/src/aura_python_sdk/_validation.py b/src/aura_python_sdk/_validation.py index cb136b7..74fb84f 100644 --- a/src/aura_python_sdk/_validation.py +++ b/src/aura_python_sdk/_validation.py @@ -3,13 +3,14 @@ from __future__ import annotations import re +from collections.abc import Sequence from aura_python_sdk._errors import AuraValidationError _UUID = re.compile(r"[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}") _INSTANCE_ID = re.compile(r"[0-9a-fA-F]{8}") -MAX_INSTANCE_NAME_LENGTH = 30 +MAX_NAME_LENGTH = 30 def require_non_empty(name: str, value: object) -> str: @@ -48,16 +49,34 @@ def session_id(value: object) -> str: return require_non_empty("GDS session ID", value) -def instance_name(value: object) -> str: - value = require_non_empty("instance name", value) - if len(value) > MAX_INSTANCE_NAME_LENGTH: - raise AuraValidationError( - f"instance name must be at most {MAX_INSTANCE_NAME_LENGTH} characters long" - ) +def display_name(label: str, value: object) -> str: + """An instance or key name: 1-30 characters with no leading or trailing whitespace.""" + value = require_non_empty(label, value) + if len(value) > MAX_NAME_LENGTH: + raise AuraValidationError(f"{label} must be at most {MAX_NAME_LENGTH} characters long") + if value != value.strip(): + raise AuraValidationError(f"{label} must not have leading or trailing whitespace") return value +def instance_name(value: object) -> str: + return display_name("instance name", value) + + def non_negative_int(name: str, value: object) -> int: if isinstance(value, bool) or not isinstance(value, int) or value < 0: raise AuraValidationError(f"{name} must be an integer of zero or more") return value + + +def boolean(name: str, value: object) -> bool: + if not isinstance(value, bool): + raise AuraValidationError(f"{name} must be True or False") + return value + + +def string_list(name: str, value: Sequence[str]) -> list[str]: + """A sequence of non-empty strings. A bare string is rejected, not split into characters.""" + if isinstance(value, str) or not isinstance(value, Sequence): + raise AuraValidationError(f"{name} must be a sequence of strings") + return [require_non_empty(f"{name} entry", item) for item in value] diff --git a/src/aura_python_sdk/services/_base.py b/src/aura_python_sdk/services/_base.py index 218fc71..b416b22 100644 --- a/src/aura_python_sdk/services/_base.py +++ b/src/aura_python_sdk/services/_base.py @@ -4,6 +4,11 @@ from aura_python_sdk._internal._request import RequestService +# List-filter query parameter names, as the v1 spec defines them. (The Go SDK sends tenant_id.) +TENANT_ID_PARAM = "tenantId" +INSTANCE_ID_PARAM = "instanceId" +ORGANIZATION_ID_PARAM = "organizationId" + class Service: """Shared plumbing for the grouped services on :class:`AuraClient`.""" diff --git a/src/aura_python_sdk/services/cmek.py b/src/aura_python_sdk/services/cmek.py index e8a531b..365ee8d 100644 --- a/src/aura_python_sdk/services/cmek.py +++ b/src/aura_python_sdk/services/cmek.py @@ -5,12 +5,13 @@ import builtins from aura_python_sdk import _validation as validate -from aura_python_sdk._internal._serde import parse_data_list -from aura_python_sdk.models.cmek import CustomerManagedKeySummary -from aura_python_sdk.services._base import Service +from aura_python_sdk._internal._request import build_path +from aura_python_sdk._internal._serde import parse_data, parse_data_list, to_json +from aura_python_sdk.models._common import CloudProvider, InstanceType +from aura_python_sdk.models.cmek import CustomerManagedKey, CustomerManagedKeySummary +from aura_python_sdk.services._base import TENANT_ID_PARAM, Service -# The spec names this query parameter tenantId. The Go SDK sends tenant_id. -TENANT_FILTER_PARAM = "tenantId" +_KEYS = "customer-managed-keys" class CMEKService(Service): @@ -21,7 +22,50 @@ def list(self, tenant_id: str | None = None) -> builtins.list[CustomerManagedKey if tenant_id is not None: tenant_id = validate.tenant_id(tenant_id) self._logger.debug("listing customer managed keys", extra={"tenant_id": tenant_id}) - response = self._api.get("customer-managed-keys", params={TENANT_FILTER_PARAM: tenant_id}) + response = self._api.get(_KEYS, params={TENANT_ID_PARAM: tenant_id}) keys = parse_data_list(CustomerManagedKeySummary, response.json()) self._logger.debug("customer managed keys listed", extra={"count": len(keys)}) return keys + + def get(self, key_id: str) -> CustomerManagedKey: + """Full details of one key. ``key_id`` is the Aura key ID, not the cloud provider's.""" + key_id = validate.require_non_empty("customer managed key ID", key_id) + self._logger.debug("getting customer managed key", extra={"key_id": key_id}) + return parse_data(CustomerManagedKey, self._api.get(build_path(_KEYS, key_id)).json()) + + def create( + self, + *, + name: str, + key_id: str, + tenant_id: str, + cloud_provider: CloudProvider | str, + region: str, + instance_type: InstanceType | str, + ) -> CustomerManagedKey: + """Register a key from your cloud provider with Aura. + + ``key_id`` is the key's ID in the cloud provider (the key ARN on AWS). The key can then + encrypt new ``instance_type`` instances in ``region``. It starts in ``pending`` status. + """ + body = { + "name": validate.display_name("key name", name), + "key_id": validate.require_non_empty("cloud provider key ID", key_id), + "tenant_id": validate.tenant_id(tenant_id), + "cloud_provider": validate.require_non_empty("cloud provider", cloud_provider), + "region": validate.require_non_empty("region", region), + "instance_type": validate.require_non_empty("instance type", instance_type), + } + self._logger.debug( + "creating customer managed key", extra={"key_name": name, "tenant_id": tenant_id} + ) + key = parse_data(CustomerManagedKey, self._api.post(_KEYS, json_body=to_json(body)).json()) + self._logger.info("customer managed key created", extra={"key_id": key.id}) + return key + + def delete(self, key_id: str) -> None: + """Delete a key. The API refuses if any instance still uses it.""" + key_id = validate.require_non_empty("customer managed key ID", key_id) + self._logger.debug("deleting customer managed key", extra={"key_id": key_id}) + self._api.delete(build_path(_KEYS, key_id)) + self._logger.info("customer managed key deleted", extra={"key_id": key_id}) diff --git a/src/aura_python_sdk/services/graph_analytics.py b/src/aura_python_sdk/services/graph_analytics.py index 19a45fd..911c8c2 100644 --- a/src/aura_python_sdk/services/graph_analytics.py +++ b/src/aura_python_sdk/services/graph_analytics.py @@ -15,7 +15,12 @@ GDSSessionConfig, GDSSessionSizeEstimate, ) -from aura_python_sdk.services._base import Service +from aura_python_sdk.services._base import ( + INSTANCE_ID_PARAM, + ORGANIZATION_ID_PARAM, + TENANT_ID_PARAM, + Service, +) _SESSIONS = "graph-analytics/sessions" @@ -23,10 +28,23 @@ class GDSSessionService(Service): """Graph Analytics (GDS) sessions.""" - def list(self) -> builtins.list[GDSSession]: - """Every session the credentials can access.""" + def list( + self, + *, + tenant_id: str | None = None, + instance_id: str | None = None, + organization_id: str | None = None, + ) -> builtins.list[GDSSession]: + """Every session the credentials can access, optionally filtered.""" + params = { + TENANT_ID_PARAM: None if tenant_id is None else validate.tenant_id(tenant_id), + INSTANCE_ID_PARAM: None if instance_id is None else validate.instance_id(instance_id), + ORGANIZATION_ID_PARAM: None + if organization_id is None + else validate.require_non_empty("organization ID", organization_id), + } self._logger.debug("listing GDS sessions") - sessions = parse_data_list(GDSSession, self._api.get(_SESSIONS).json()) + sessions = parse_data_list(GDSSession, self._api.get(_SESSIONS, params=params).json()) self._logger.debug("GDS sessions listed", extra={"count": len(sessions)}) return sessions @@ -56,12 +74,9 @@ def estimate_size( if value is not None: body[key] = validate.non_negative_int(key.replace("_", " "), value) if algorithm_categories is not None: - if isinstance(algorithm_categories, str): - raise AuraValidationError("algorithm categories must be a sequence of strings") - body["algorithm_categories"] = [ - validate.require_non_empty("algorithm category", category) - for category in algorithm_categories - ] + body["algorithm_categories"] = validate.string_list( + "algorithm categories", algorithm_categories + ) self._logger.debug("estimating GDS session size") response = self._api.post(f"{_SESSIONS}/sizing", json_body=body) diff --git a/src/aura_python_sdk/services/instances.py b/src/aura_python_sdk/services/instances.py index 3cf95c3..aa0406e 100644 --- a/src/aura_python_sdk/services/instances.py +++ b/src/aura_python_sdk/services/instances.py @@ -3,28 +3,34 @@ from __future__ import annotations import builtins +from collections.abc import Sequence from aura_python_sdk import _validation as validate from aura_python_sdk._errors import AuraValidationError from aura_python_sdk._internal._request import build_path from aura_python_sdk._internal._serde import parse_data, parse_data_list, to_json +from aura_python_sdk.models._common import InstanceType from aura_python_sdk.models.instances import ( CDCEnrichmentMode, CreatedInstance, Instance, InstanceConfig, + InstanceSizeEstimate, InstanceSummary, ) -from aura_python_sdk.services._base import Service +from aura_python_sdk.services._base import TENANT_ID_PARAM, Service class InstanceService(Service): """AuraDB and AuraDS instances.""" - def list(self) -> builtins.list[InstanceSummary]: - """Every instance the credentials can access.""" - self._logger.debug("listing instances") - instances = parse_data_list(InstanceSummary, self._api.get("instances").json()) + def list(self, tenant_id: str | None = None) -> builtins.list[InstanceSummary]: + """Every instance the credentials can access, optionally only those in one tenant.""" + if tenant_id is not None: + tenant_id = validate.tenant_id(tenant_id) + self._logger.debug("listing instances", extra={"tenant_id": tenant_id}) + response = self._api.get("instances", params={TENANT_ID_PARAM: tenant_id}) + instances = parse_data_list(InstanceSummary, response.json()) self._logger.debug("instances listed", extra={"count": len(instances)}) return instances @@ -84,12 +90,17 @@ def update( *, name: str | None = None, memory: str | None = None, + storage: str | None = None, + vector_optimized: bool | None = None, + graph_analytics_plugin: bool | None = None, cdc_enrichment_mode: CDCEnrichmentMode | str | None = None, secondaries_count: int | None = None, ) -> Instance: """Rename, resize or reconfigure an instance. Only the arguments given are changed. The update is asynchronous, and the instance stays available throughout. + ``secondaries_count`` applies only to Virtual Dedicated Cloud, and + ``cdc_enrichment_mode`` only to Virtual Dedicated Cloud and Business Critical. """ instance_id = validate.instance_id(instance_id) changes: dict[str, object] = {} @@ -97,6 +108,14 @@ def update( changes["name"] = validate.instance_name(name) if memory is not None: changes["memory"] = validate.require_non_empty("memory", memory) + if storage is not None: + changes["storage"] = validate.require_non_empty("storage", storage) + if vector_optimized is not None: + changes["vector_optimized"] = validate.boolean("vector optimized", vector_optimized) + if graph_analytics_plugin is not None: + changes["graph_analytics_plugin"] = validate.boolean( + "graph analytics plugin", graph_analytics_plugin + ) if cdc_enrichment_mode is not None: changes["cdc_enrichment_mode"] = validate.require_non_empty( "CDC enrichment mode", cdc_enrichment_mode @@ -116,6 +135,56 @@ def update( self._logger.info("instance update started", extra={"instance_id": instance_id}) return instance + def estimate_size( + self, + *, + node_count: int, + relationship_count: int, + instance_type: InstanceType | str | None = None, + algorithm_categories: Sequence[str] | None = None, + ) -> InstanceSizeEstimate: + """Estimate the instance size needed for a graph. + + Supported for ``enterprise-ds`` and ``professional-ds``. Pass the recommended size as + ``memory`` when creating the instance. + """ + body: dict[str, object] = { + "node_count": validate.non_negative_int("node count", node_count), + "relationship_count": validate.non_negative_int( + "relationship count", relationship_count + ), + } + if instance_type is not None: + body["instance_type"] = validate.require_non_empty("instance type", instance_type) + if algorithm_categories is not None: + body["algorithm_categories"] = validate.string_list( + "algorithm categories", algorithm_categories + ) + self._logger.debug("estimating instance size") + response = self._api.post("instances/sizing", json_body=to_json(body)) + return parse_data(InstanceSizeEstimate, response.json()) + + def upgrade( + self, instance_id: str, *, memory: str | None = None, storage: str | None = None + ) -> Instance: + """Upgrade an AuraDB Professional instance to Business Critical. + + Pass both ``memory`` and ``storage`` to resize as part of the upgrade, or neither to keep + the current size. Not available for Marketplace projects or trial instances. + """ + instance_id = validate.instance_id(instance_id) + if (memory is None) != (storage is None): + raise AuraValidationError("upgrade requires both memory and storage, or neither") + body: dict[str, object] = {} + if memory is not None and storage is not None: + body["memory"] = validate.require_non_empty("memory", memory) + body["storage"] = validate.require_non_empty("storage", storage) + self._logger.debug("upgrading instance", extra={"instance_id": instance_id}) + response = self._api.post(build_path("instances", instance_id, "upgrade"), json_body=body) + instance = parse_data(Instance, response.json()) + self._logger.info("instance upgrade started", extra={"instance_id": instance_id}) + return instance + def delete(self, instance_id: str) -> Instance: """Start deleting an instance. This cannot be undone.""" instance_id = validate.instance_id(instance_id) diff --git a/tests/unit/test_cmek_service.py b/tests/unit/test_cmek_service.py index 76831bb..133af13 100644 --- a/tests/unit/test_cmek_service.py +++ b/tests/unit/test_cmek_service.py @@ -1,6 +1,16 @@ +from typing import Any + import pytest -from aura_python_sdk import AuraValidationError, CustomerManagedKeySummary +from aura_python_sdk import ( + AuraValidationError, + BadRequestError, + CloudProvider, + CustomerManagedKey, + CustomerManagedKeySummary, + HttpResponse, + InstanceType, +) from tests.unit.conftest import BASE, TENANT_ID, Api KEY = {"id": "8c764aed-8eb3-4a1c-92f6-e4ef0c7a6ed9", "name": "Key01", "tenant_id": TENANT_ID} @@ -22,3 +32,105 @@ def test_list_invalid_tenant_sends_nothing(api: Api) -> None: with pytest.raises(AuraValidationError, match="tenant ID"): api.client.cmek.list("bad") api.assert_no_request() + + +FULL_KEY = { + **KEY, + "cloud_provider": "aws", + "region": "us-west-2", + "instance_type": "enterprise-db", + "key_id": "arn:aws:kms:us-west-2:111122223333:key/1234abcd-12ab-34cd-56ef-1234567890ab", + "status": "pending", + "created": "2024-01-31T14:06:57Z", +} + +CREATE_KWARGS = { + "name": "Production Key", + "key_id": FULL_KEY["key_id"], + "tenant_id": TENANT_ID, + "cloud_provider": CloudProvider.AWS, + "region": "us-west-2", + "instance_type": InstanceType.ENTERPRISE_DB, +} + + +def test_get(api: Api) -> None: + api.reply(200, {"data": FULL_KEY}) + key = api.client.cmek.get(KEY["id"]) + assert isinstance(key, CustomerManagedKey) + assert key.cloud_provider is CloudProvider.AWS + assert api.request.url == f"{BASE}/customer-managed-keys/{KEY['id']}" + + +def test_create(api: Api) -> None: + api.reply(202, {"data": FULL_KEY}) + key = api.client.cmek.create(**CREATE_KWARGS) + assert key.status == "pending" + assert (api.request.method, api.request.url) == ("POST", f"{BASE}/customer-managed-keys") + # Matches the spec's "Creates a Customer Managed Key" request example. + assert api.body == { + "name": "Production Key", + "key_id": FULL_KEY["key_id"], + "tenant_id": TENANT_ID, + "cloud_provider": "aws", + "region": "us-west-2", + "instance_type": "enterprise-db", + } + + +@pytest.mark.parametrize( + ("override", "message"), + [ + ({"name": ""}, "key name must not be empty"), + ({"name": "k" * 31}, "at most 30 characters"), + ({"name": "Key "}, "leading or trailing whitespace"), + ({"key_id": ""}, "cloud provider key ID"), + ({"tenant_id": "bad"}, "tenant ID"), + ({"cloud_provider": ""}, "cloud provider must not be empty"), + ({"region": ""}, "region"), + ({"instance_type": ""}, "instance type"), + ], +) +def test_create_validation(api: Api, override: dict[str, Any], message: str) -> None: + with pytest.raises(AuraValidationError, match=message): + api.client.cmek.create(**{**CREATE_KWARGS, **override}) + api.assert_no_request() + + +def test_delete(api: Api) -> None: + api.transport.queue(HttpResponse(204)) + api.client.cmek.delete(KEY["id"]) + assert (api.request.method, api.request.url) == ( + "DELETE", + f"{BASE}/customer-managed-keys/{KEY['id']}", + ) + + +def test_delete_active_key_is_bad_request(api: Api) -> None: + api.reply( + 400, + { + "errors": [ + { + "message": "The key is linked to an active instance.", + "reason": "encryption-key-is-active", + } + ] + }, + ) + with pytest.raises(BadRequestError) as info: + api.client.cmek.delete(KEY["id"]) + assert info.value.details[0].reason == "encryption-key-is-active" + + +@pytest.mark.parametrize("call", ["get", "delete"]) +def test_empty_key_id(api: Api, call: str) -> None: + with pytest.raises(AuraValidationError, match="customer managed key ID must not be empty"): + getattr(api.client.cmek, call)("") + api.assert_no_request() + + +def test_key_id_is_path_encoded(api: Api) -> None: + api.reply(200, {"data": FULL_KEY}) + api.client.cmek.get("a/b") + assert api.request.url == f"{BASE}/customer-managed-keys/a%2Fb" diff --git a/tests/unit/test_graph_analytics_service.py b/tests/unit/test_graph_analytics_service.py index 3db9459..60797af 100644 --- a/tests/unit/test_graph_analytics_service.py +++ b/tests/unit/test_graph_analytics_service.py @@ -139,7 +139,7 @@ def test_estimate_size_minimal(api: Api) -> None: ), ( {"node_count": 1, "relationship_count": 1, "algorithm_categories": [""]}, - "algorithm category", + "algorithm categories entry", ), ], ) @@ -147,3 +147,27 @@ def test_estimate_size_validation(api: Api, kwargs: dict[str, Any], message: str with pytest.raises(AuraValidationError, match=message): api.client.graph_analytics.estimate_size(**kwargs) api.assert_no_request() + + +def test_list_with_filters(api: Api) -> None: + api.reply(200, {"data": []}) + api.client.graph_analytics.list( + tenant_id=TENANT_ID, instance_id=INSTANCE_ID, organization_id="org-1" + ) + assert api.request.url == ( + f"{SESSIONS}?tenantId={TENANT_ID}&instanceId={INSTANCE_ID}&organizationId=org-1" + ) + + +@pytest.mark.parametrize( + ("kwargs", "message"), + [ + ({"tenant_id": "bad"}, "tenant ID"), + ({"instance_id": "bad"}, "instance ID"), + ({"organization_id": ""}, "organization ID"), + ], +) +def test_list_filter_validation(api: Api, kwargs: dict[str, Any], message: str) -> None: + with pytest.raises(AuraValidationError, match=message): + api.client.graph_analytics.list(**kwargs) + api.assert_no_request() diff --git a/tests/unit/test_instances_service.py b/tests/unit/test_instances_service.py index f8c910b..0aa0391 100644 --- a/tests/unit/test_instances_service.py +++ b/tests/unit/test_instances_service.py @@ -271,3 +271,143 @@ def test_invalid_ids_send_nothing(api: Api, call: Any) -> None: with pytest.raises(AuraValidationError): call(api.client.instances) api.assert_no_request() + + +# --- Phase 5: spec coverage beyond the Go SDK --- + + +def test_list_filtered_by_tenant(api: Api) -> None: + api.reply(200, {"data": []}) + api.client.instances.list(TENANT_ID) + assert api.request.url == f"{BASE}/instances?tenantId={TENANT_ID}" + + +def test_list_invalid_tenant_sends_nothing(api: Api) -> None: + with pytest.raises(AuraValidationError, match="tenant ID"): + api.client.instances.list("bad") + api.assert_no_request() + + +def test_update_new_spec_fields(api: Api) -> None: + api.reply(202, {"data": INSTANCE}) + api.client.instances.update( + INSTANCE_ID, storage="32GB", vector_optimized=True, graph_analytics_plugin=False + ) + assert api.body == { + "storage": "32GB", + "vector_optimized": True, + "graph_analytics_plugin": False, + } + + +@pytest.mark.parametrize( + ("kwargs", "message"), + [ + ({"storage": ""}, "storage must not be empty"), + ({"vector_optimized": "yes"}, "vector optimized must be True or False"), + ({"graph_analytics_plugin": 1}, "graph analytics plugin must be True or False"), + ({"name": " padded"}, "leading or trailing whitespace"), + ], +) +def test_update_new_field_validation(api: Api, kwargs: dict[str, Any], message: str) -> None: + with pytest.raises(AuraValidationError, match=message): + api.client.instances.update(INSTANCE_ID, **kwargs) + api.assert_no_request() + + +def test_estimate_size(api: Api) -> None: + api.reply( + 200, + { + "data": { + "did_exceed_maximum": False, + "min_required_memory": "14GB", + "recommended_size": "16GB", + } + }, + ) + estimate = api.client.instances.estimate_size( + node_count=1_000_000, + relationship_count=5_000_000, + instance_type=InstanceType.PROFESSIONAL_DS, + algorithm_categories=["pathfinding", "community-detection"], + ) + assert estimate.recommended_size == "16GB" + assert estimate.did_exceed_maximum is False + assert (api.request.method, api.request.url) == ("POST", f"{BASE}/instances/sizing") + assert api.body == { + "node_count": 1_000_000, + "relationship_count": 5_000_000, + "instance_type": "professional-ds", + "algorithm_categories": ["pathfinding", "community-detection"], + } + + +def test_estimate_size_minimal(api: Api) -> None: + api.reply( + 200, + { + "data": { + "did_exceed_maximum": True, + "min_required_memory": "1TB", + "recommended_size": "1TB", + } + }, + ) + api.client.instances.estimate_size(node_count=1, relationship_count=2) + assert api.body == {"node_count": 1, "relationship_count": 2} + + +@pytest.mark.parametrize( + ("kwargs", "message"), + [ + ({"node_count": -1, "relationship_count": 0}, "node count"), + ({"node_count": 1, "relationship_count": None}, "relationship count"), + ({"node_count": 1, "relationship_count": 1, "instance_type": ""}, "instance type"), + ( + {"node_count": 1, "relationship_count": 1, "algorithm_categories": "pathfinding"}, + "sequence", + ), + ], +) +def test_estimate_size_validation(api: Api, kwargs: dict[str, Any], message: str) -> None: + with pytest.raises(AuraValidationError, match=message): + api.client.instances.estimate_size(**kwargs) + api.assert_no_request() + + +def test_upgrade_keeping_size(api: Api) -> None: + api.reply(200, {"data": {**INSTANCE, "type": "business-critical"}}) + instance = api.client.instances.upgrade(INSTANCE_ID) + assert instance.type is InstanceType.BUSINESS_CRITICAL + assert (api.request.method, api.request.url) == ( + "POST", + f"{BASE}/instances/{INSTANCE_ID}/upgrade", + ) + assert api.body == {} + + +def test_upgrade_with_resize(api: Api) -> None: + api.reply(200, {"data": INSTANCE}) + api.client.instances.upgrade(INSTANCE_ID, memory="16GB", storage="32GB") + assert api.body == {"memory": "16GB", "storage": "32GB"} + + +@pytest.mark.parametrize( + ("kwargs", "message"), + [ + ({"memory": "16GB"}, "both memory and storage"), + ({"storage": "32GB"}, "both memory and storage"), + ({"memory": "", "storage": "32GB"}, "memory must not be empty"), + ], +) +def test_upgrade_validation(api: Api, kwargs: dict[str, Any], message: str) -> None: + with pytest.raises(AuraValidationError, match=message): + api.client.instances.upgrade(INSTANCE_ID, **kwargs) + api.assert_no_request() + + +def test_upgrade_invalid_id(api: Api) -> None: + with pytest.raises(AuraValidationError, match="instance ID"): + api.client.instances.upgrade("bad") + api.assert_no_request() diff --git a/tests/unit/test_spec_examples.py b/tests/unit/test_spec_examples.py index d3f05dc..a72ba14 100644 --- a/tests/unit/test_spec_examples.py +++ b/tests/unit/test_spec_examples.py @@ -114,3 +114,52 @@ def test_spec_example_parses(operation_id: str, status: str, example: Any) -> No assert all(isinstance(item, model) for item in parsed) else: assert isinstance(parse_data(model, example), model) + + +# operationId -> "service.method" on AuraClient +OPERATION_METHODS = { + "get-instances": "instances.list", + "post-instances": "instances.create", + "post-instances-sizing": "instances.estimate_size", + "get-instance-id": "instances.get", + "delete-instance-id": "instances.delete", + "patch-instance-id": "instances.update", + "post-overwrite-instance": "instances.overwrite_from_instance", + "post-pause-instance": "instances.pause", + "post-resume-instance": "instances.resume", + "post-upgrade-instance": "instances.upgrade", + "get-snapshots": "snapshots.list", + "post-snapshots": "snapshots.create", + "get-snapshot-snapshotid": "snapshots.get", + "post-restore-snapshot": "snapshots.restore", + "get-projects": "tenants.list", + "get-project-id": "tenants.get", + "get-project-metrics-integration-details": "tenants.get_metrics_integration", + "get-customer-managed-keys": "cmek.list", + "post-customer-managed-keys": "cmek.create", + "get-customer-managed-key-id": "cmek.get", + "delete-customer-managed-key-id": "cmek.delete", + "get-sessions": "graph_analytics.list", + "post-session": "graph_analytics.create", + "post-sessions-sizing": "graph_analytics.estimate_size", + "get-session": "graph_analytics.get", + "delete-session": "graph_analytics.delete", +} + + +def test_every_spec_operation_has_a_client_method() -> None: + from aura_python_sdk import AuraClient + from tests.fakes import FakeTransport + + operation_ids = { + operation["operationId"] + for path_item in _load_spec()["paths"].values() + for method, operation in path_item.items() + if method in HTTP_METHODS + } + assert operation_ids == OPERATION_METHODS.keys() + + client = AuraClient(client_id="id", client_secret="secret", transport=FakeTransport()) + for dotted in OPERATION_METHODS.values(): + service_name, method_name = dotted.split(".") + assert callable(getattr(getattr(client, service_name), method_name)), dotted From 0a870afe26e7718b77cfac11e08af61c432e5869 Mon Sep 17 00:00:00 2001 From: Jonathan Giffard Date: Tue, 29 Sep 2026 20:35:23 +0100 Subject: [PATCH 5/8] Add Prometheus metrics service with a stdlib exposition parser Phase 6 of PLAN.md. Adds client.prometheus (fetch_raw_metrics, get_metric_value, get_instance_health) with Go's thresholds, a text-format parser whose output matches Go's expfmt, and a guard that only sends the Aura token to https://*.neo4j.io metrics URLs. Drops the prometheus extra. Co-Authored-By: Claude Opus 5.5 --- PLAN.md | 19 +- pyproject.toml | 3 - src/aura_python_sdk/__init__.py | 18 ++ src/aura_python_sdk/_client.py | 13 +- src/aura_python_sdk/_config.py | 2 + src/aura_python_sdk/_errors.py | 4 + .../_internal/metrics/__init__.py | 0 .../_internal/metrics/_parser.py | 145 +++++++++++ src/aura_python_sdk/models/__init__.py | 18 ++ src/aura_python_sdk/models/prometheus.py | 79 ++++++ src/aura_python_sdk/services/__init__.py | 2 + src/aura_python_sdk/services/prometheus.py | 241 ++++++++++++++++++ tests/unit/test_import_boundaries.py | 1 - tests/unit/test_prometheus_parser.py | 126 +++++++++ tests/unit/test_prometheus_service.py | 239 +++++++++++++++++ uv.lock | 20 +- 16 files changed, 900 insertions(+), 30 deletions(-) create mode 100644 src/aura_python_sdk/_internal/metrics/__init__.py create mode 100644 src/aura_python_sdk/_internal/metrics/_parser.py create mode 100644 src/aura_python_sdk/models/prometheus.py create mode 100644 src/aura_python_sdk/services/prometheus.py create mode 100644 tests/unit/test_prometheus_parser.py create mode 100644 tests/unit/test_prometheus_service.py diff --git a/PLAN.md b/PLAN.md index b3df9e4..5129524 100644 --- a/PLAN.md +++ b/PLAN.md @@ -215,6 +215,18 @@ parser accepts the spec's `{"errors": [...]}`, the middleware `{"error": "..."}` - **Coverage guard**: `test_every_spec_operation_has_a_client_method` fails if the spec gains an operation that no SDK method covers. +### 2.9 Decisions made in phase 6 + +- **No `prometheus_client`.** It normalises counter names (`foo` becomes `foo_total`) and + converts timestamps to seconds, so its keys wouldn't match the Go SDK's. A roughly 150-line + stdlib parser gives exactly the same output as Go's `expfmt`, checked with a Go program on the + same input. The `[prometheus]` extra is gone, so httpx is the only runtime dependency. +- **Metrics URL guard.** The Aura bearer token is sent to the metrics URL, so it must be + `https://*.neo4j.io` unless `allow_insecure_base_url=True`. Go sends the token to any URL. +- **Missing metrics are `None`, not `0`.** `InstanceHealth` fields are `None` when the endpoint + didn't report a metric, and threshold checks skip them. The status logic and messages match Go. +- **`get_metric_value`** raises `MetricNotFoundError`, which is also a `LookupError`. + ## 3. Package layout ``` @@ -237,7 +249,7 @@ src/aura_python_sdk/ _types.py # HttpRequest, HttpResponse, HttpTransport Protocol _httpx.py # HttpxTransport, the only httpx import metrics/ - _parser.py # the only prometheus_client import (optional extra) + _parser.py # stdlib Prometheus text-format parser (matches Go expfmt) tests/ unit/ # FakeTransport, no network transport/ # HttpxTransport against httpx.MockTransport @@ -250,7 +262,6 @@ examples/ # ports of go examples/v1/* | Dependency | Purpose | Wrapped in | |---|---|---| | `httpx` | HTTP | `_internal/http/_httpx.py` | -| `prometheus_client` (optional extra `[prometheus]`) | Parse the Prometheus text format | `_internal/metrics/_parser.py` | Dev tooling: `uv`, `ruff` (lint and format), `mypy --strict`, `pytest`, `pytest-cov`. No `respx`: `httpx.MockTransport` plus our own fake transport are enough. @@ -265,14 +276,14 @@ Dev tooling: `uv`, `ruff` (lint and format), `mypy --strict`, `pytest`, `pytest- example payloads. 4. **Services at Go parity**: tenants, instances, snapshots, `cmek.list`, graph_analytics. 5. **Spec gap-fill**: instance sizing and upgrade, CMEK get/create/delete, list filters, extra PATCH fields. -6. **Prometheus**: the optional extra plus the health assessment. +6. **Prometheus**: a stdlib metrics parser plus the health assessment. 7. **Docs and release**: README, the ported examples, CHANGELOG, opt-in integration tests, PyPI publish workflow. 8. *(If chosen)* **Async**: `AsyncAuraClient` over an `AsyncHttpTransport`, reusing request building and parsing. The layering keeps this additive. ## 6. Enforcing "wrap every import" -A unit test walks `src/` with `ast` and fails if `httpx` or `prometheus_client` is imported anywhere +A unit test walks `src/` with `ast` and fails if `httpx` (or any unregistered dependency) is imported anywhere except its designated module. It also checks that no public symbol's annotations reference those packages. diff --git a/pyproject.toml b/pyproject.toml index ec0b9ac..51c2d74 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -25,9 +25,6 @@ classifiers = [ ] dependencies = ["httpx>=0.27,<1"] -[project.optional-dependencies] -prometheus = ["prometheus-client>=0.20"] - [project.urls] Homepage = "https://github.com/LackOfMorals/aura-python-sdk" "Aura API documentation" = "https://neo4j.com/docs/aura/api/overview/" diff --git a/src/aura_python_sdk/__init__.py b/src/aura_python_sdk/__init__.py index ee2c8f2..11b63ff 100644 --- a/src/aura_python_sdk/__init__.py +++ b/src/aura_python_sdk/__init__.py @@ -24,6 +24,7 @@ BadRequestError, ConflictError, ErrorDetail, + MetricNotFoundError, NotFoundError, PermissionDeniedError, RateLimitError, @@ -34,6 +35,7 @@ from aura_python_sdk.models import ( CDCEnrichmentMode, CloudProvider, + ConnectionMetrics, CreatedInstance, CreatedSnapshot, CustomerManagedKey, @@ -43,17 +45,24 @@ GDSSessionConfig, GDSSessionSizeEstimate, GDSSessionStatus, + HealthStatus, Instance, InstanceConfig, InstanceConfiguration, + InstanceHealth, InstanceSizeEstimate, InstanceStatus, InstanceSummary, InstanceType, MetricsIntegration, + PrometheusMetric, + PrometheusMetrics, + QueryMetrics, + ResourceMetrics, Snapshot, SnapshotProfile, SnapshotStatus, + StorageMetrics, Tenant, TenantSummary, ) @@ -75,6 +84,7 @@ "CDCEnrichmentMode", "CloudProvider", "ConflictError", + "ConnectionMetrics", "CreatedInstance", "CreatedSnapshot", "CustomerManagedKey", @@ -85,24 +95,32 @@ "GDSSessionConfig", "GDSSessionSizeEstimate", "GDSSessionStatus", + "HealthStatus", "HttpRequest", "HttpResponse", "HttpTransport", "Instance", "InstanceConfig", "InstanceConfiguration", + "InstanceHealth", "InstanceSizeEstimate", "InstanceStatus", "InstanceSummary", "InstanceType", + "MetricNotFoundError", "MetricsIntegration", "NotFoundError", "PermissionDeniedError", + "PrometheusMetric", + "PrometheusMetrics", + "QueryMetrics", "RateLimitError", + "ResourceMetrics", "ServerError", "Snapshot", "SnapshotProfile", "SnapshotStatus", + "StorageMetrics", "Tenant", "TenantSummary", "__version__", diff --git a/src/aura_python_sdk/_client.py b/src/aura_python_sdk/_client.py index f703c61..744e9b4 100644 --- a/src/aura_python_sdk/_client.py +++ b/src/aura_python_sdk/_client.py @@ -28,6 +28,7 @@ CMEKService, GDSSessionService, InstanceService, + PrometheusService, SnapshotService, TenantService, ) @@ -48,7 +49,7 @@ class AuraClient: print(instance.id, instance.name) Services, mirroring the Go SDK: ``tenants``, ``instances``, ``snapshots``, ``cmek`` and - ``graph_analytics``. + ``graph_analytics``, plus ``prometheus`` for metrics endpoints. Every option is keyword-only. Invalid options raise :class:`AuraConfigurationError`. @@ -56,8 +57,9 @@ class AuraClient: client_id: Aura API client ID. client_secret: Aura API client secret. base_url: API base URL. It must use HTTPS unless ``allow_insecure_base_url`` is set. - allow_insecure_base_url: Allow an ``http://`` base URL. Only for local test servers, - because credentials would be sent in cleartext. + allow_insecure_base_url: Allow an ``http://`` base URL, and Prometheus URLs outside + ``https://*.neo4j.io``. Only for local test servers, because credentials would be sent + in cleartext. timeout: Seconds allowed for each API call, covering the token fetch, retries and backoff. max_retries: How many times to retry after a network failure. Responses with an HTTP status are never retried. @@ -138,6 +140,11 @@ def __init__( self.graph_analytics = GDSSessionService( self._api, self._logger.getChild("graph_analytics") ) + self.prometheus = PrometheusService( + self._api, + self._logger.getChild("prometheus"), + allow_untrusted_urls=self._config.allow_insecure_base_url, + ) self._logger.debug( "Aura API client initialized", diff --git a/src/aura_python_sdk/_config.py b/src/aura_python_sdk/_config.py index 4b7aaa4..c9a3b93 100644 --- a/src/aura_python_sdk/_config.py +++ b/src/aura_python_sdk/_config.py @@ -28,6 +28,7 @@ class ClientConfig: client_id: str client_secret: str = field(repr=False) base_url: str + allow_insecure_base_url: bool timeout: float max_retries: int max_response_size: int @@ -57,6 +58,7 @@ def build_config( client_id=client_id, client_secret=client_secret, base_url=_validate_base_url(base_url, allow_insecure=allow_insecure_base_url), + allow_insecure_base_url=bool(allow_insecure_base_url), timeout=_validate_timeout(timeout), max_retries=_validate_non_negative_int("max retries", max_retries), max_response_size=_validate_positive_int("max response size", max_response_size), diff --git a/src/aura_python_sdk/_errors.py b/src/aura_python_sdk/_errors.py index 928a01b..7df56bd 100644 --- a/src/aura_python_sdk/_errors.py +++ b/src/aura_python_sdk/_errors.py @@ -47,6 +47,10 @@ class AuraResponseError(AuraError): """The API response could not be used: too large, not valid JSON, or an unexpected shape.""" +class MetricNotFoundError(AuraError, LookupError): + """No Prometheus metric matched the requested name and label filters.""" + + @dataclass(frozen=True, slots=True) class ErrorDetail: """One entry from the ``errors`` array of an Aura API error response.""" diff --git a/src/aura_python_sdk/_internal/metrics/__init__.py b/src/aura_python_sdk/_internal/metrics/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/aura_python_sdk/_internal/metrics/_parser.py b/src/aura_python_sdk/_internal/metrics/_parser.py new file mode 100644 index 0000000..45727d3 --- /dev/null +++ b/src/aura_python_sdk/_internal/metrics/_parser.py @@ -0,0 +1,145 @@ +"""Parser for the Prometheus text exposition format (version 0.0.4). + +The keys and values match the Go SDK, which uses ``expfmt.TextParser``: + +- Metrics are keyed by their ``# TYPE`` name. A counter declared as ``foo_total`` stays + ``foo_total``, and a counter declared as ``foo`` stays ``foo``. +- A summary or histogram becomes one entry per label set, keyed by its base name, with the + ``_sum`` sample as its value. The quantile, bucket and ``_count`` lines are skipped. +- A sample without a ``# TYPE`` line is untyped and keyed by its own name. +- Timestamps stay in milliseconds. +""" + +from __future__ import annotations + +import re + +from aura_python_sdk._errors import AuraResponseError +from aura_python_sdk.models.prometheus import PrometheusMetric + +_NAME = re.compile(r"[a-zA-Z_:][a-zA-Z0-9_:]*") +_LABEL_NAME = re.compile(r"[a-zA-Z_][a-zA-Z0-9_]*") +_ESCAPES = {"\\": "\\", '"': '"', "n": "\n"} +_AGGREGATE_TYPES = frozenset({"summary", "histogram"}) +_AGGREGATE_SUFFIXES = ("_sum", "_count", "_bucket") + + +class _ParseError(Exception): + pass + + +def parse_exposition(text: str) -> dict[str, tuple[PrometheusMetric, ...]]: + """Parse exposition text into metrics keyed by name. Raises AuraResponseError if malformed.""" + types: dict[str, str] = {} + grouped: dict[str, list[PrometheusMetric]] = {} + for lineno, raw_line in enumerate(text.splitlines(), start=1): + line = raw_line.strip() + if not line: + continue + if line.startswith("#"): + parts = line[1:].split(None, 2) + if len(parts) == 3 and parts[0] == "TYPE": + types[parts[1]] = parts[2].strip().lower() + continue + try: + name, labels, value, timestamp_ms = _parse_sample(line) + except _ParseError as exc: + raise AuraResponseError(f"invalid Prometheus metrics at line {lineno}: {exc}") from None + + key = _metric_key(name, types) + if key is None: + continue + grouped.setdefault(key, []).append( + PrometheusMetric(name=key, labels=labels, value=value, timestamp_ms=timestamp_ms) + ) + return {name: tuple(samples) for name, samples in grouped.items()} + + +def _metric_key(sample_name: str, types: dict[str, str]) -> str | None: + """The metric key a sample belongs to, or None if the Go SDK would not report it.""" + if types.get(sample_name) in _AGGREGATE_TYPES: + return None # a quantile line of a summary + for suffix in _AGGREGATE_SUFFIXES: + base = sample_name.removesuffix(suffix) + if base != sample_name and types.get(base) in _AGGREGATE_TYPES: + return base if suffix == "_sum" else None + return sample_name + + +def _parse_sample(line: str) -> tuple[str, dict[str, str], float, int | None]: + match = _NAME.match(line) + if not match: + raise _ParseError("expected a metric name") + name = match.group() + pos = match.end() + + labels: dict[str, str] = {} + if pos < len(line) and line[pos] == "{": + labels, pos = _parse_labels(line, pos + 1) + + fields = line[pos:].split() + if len(fields) not in (1, 2): + raise _ParseError("expected a value and an optional timestamp") + try: + value = float(fields[0]) + except ValueError: + raise _ParseError(f"invalid value {fields[0]!r}") from None + timestamp_ms = None + if len(fields) == 2: + try: + timestamp_ms = int(fields[1]) + except ValueError: + raise _ParseError(f"invalid timestamp {fields[1]!r}") from None + return name, labels, value, timestamp_ms + + +def _parse_labels(line: str, pos: int) -> tuple[dict[str, str], int]: + """Parse ``name="value",...}`` starting after the opening brace. Returns (labels, end).""" + labels: dict[str, str] = {} + while True: + pos = _skip_spaces(line, pos) + if pos < len(line) and line[pos] == "}": + return labels, pos + 1 + match = _LABEL_NAME.match(line, pos) + if not match: + raise _ParseError("expected a label name") + label = match.group() + pos = _expect(line, _skip_spaces(line, match.end()), "=", label) + pos = _expect(line, _skip_spaces(line, pos), '"', label) + value, pos = _parse_label_value(line, pos) + labels[label] = value + pos = _skip_spaces(line, pos) + if pos < len(line) and line[pos] == ",": + pos += 1 + elif pos >= len(line) or line[pos] != "}": + raise _ParseError("expected ',' or '}' after a label") + + +def _parse_label_value(line: str, pos: int) -> tuple[str, int]: + chars: list[str] = [] + while pos < len(line): + char = line[pos] + if char == "\\": + escaped = line[pos + 1 : pos + 2] + if escaped not in _ESCAPES: + raise _ParseError(f"invalid escape sequence \\{escaped}") + chars.append(_ESCAPES[escaped]) + pos += 2 + elif char == '"': + return "".join(chars), pos + 1 + else: + chars.append(char) + pos += 1 + raise _ParseError("unterminated label value") + + +def _expect(line: str, pos: int, char: str, label: str) -> int: + if line[pos : pos + 1] != char: + raise _ParseError(f"expected {char!r} in label {label!r}") + return pos + 1 + + +def _skip_spaces(line: str, pos: int) -> int: + while pos < len(line) and line[pos] in " \t": + pos += 1 + return pos diff --git a/src/aura_python_sdk/models/__init__.py b/src/aura_python_sdk/models/__init__.py index 089755d..641f721 100644 --- a/src/aura_python_sdk/models/__init__.py +++ b/src/aura_python_sdk/models/__init__.py @@ -18,6 +18,16 @@ InstanceStatus, InstanceSummary, ) +from aura_python_sdk.models.prometheus import ( + ConnectionMetrics, + HealthStatus, + InstanceHealth, + PrometheusMetric, + PrometheusMetrics, + QueryMetrics, + ResourceMetrics, + StorageMetrics, +) from aura_python_sdk.models.snapshots import ( CreatedSnapshot, Snapshot, @@ -34,6 +44,7 @@ __all__ = [ "CDCEnrichmentMode", "CloudProvider", + "ConnectionMetrics", "CreatedInstance", "CreatedSnapshot", "CustomerManagedKey", @@ -43,17 +54,24 @@ "GDSSessionConfig", "GDSSessionSizeEstimate", "GDSSessionStatus", + "HealthStatus", "Instance", "InstanceConfig", "InstanceConfiguration", + "InstanceHealth", "InstanceSizeEstimate", "InstanceStatus", "InstanceSummary", "InstanceType", "MetricsIntegration", + "PrometheusMetric", + "PrometheusMetrics", + "QueryMetrics", + "ResourceMetrics", "Snapshot", "SnapshotProfile", "SnapshotStatus", + "StorageMetrics", "Tenant", "TenantSummary", ] diff --git a/src/aura_python_sdk/models/prometheus.py b/src/aura_python_sdk/models/prometheus.py new file mode 100644 index 0000000..542bbd1 --- /dev/null +++ b/src/aura_python_sdk/models/prometheus.py @@ -0,0 +1,79 @@ +"""Prometheus metrics models (Go: prometheus.go).""" + +from __future__ import annotations + +from collections.abc import Mapping +from dataclasses import dataclass, field +from datetime import datetime +from enum import StrEnum + + +@dataclass(frozen=True, slots=True, kw_only=True) +class PrometheusMetric: + """One sample. For summaries and histograms, ``value`` is the ``_sum`` sample, as in Go.""" + + name: str + labels: Mapping[str, str] = field(default_factory=dict) + value: float + timestamp_ms: int | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class PrometheusMetrics: + """Every sample from a metrics endpoint, keyed by metric name. + + A counter is keyed by the name on its ``# TYPE`` line (for example + ``neo4j_db_query_execution_success_total``). A summary or histogram is keyed by its base name. + """ + + metrics: Mapping[str, tuple[PrometheusMetric, ...]] = field(default_factory=dict) + + +class HealthStatus(StrEnum): + HEALTHY = "healthy" + WARNING = "warning" + CRITICAL = "critical" + + +@dataclass(frozen=True, slots=True, kw_only=True) +class ResourceMetrics: + cpu_usage_percent: float | None = None + memory_usage_percent: float | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class QueryMetrics: + query_execution_total: float | None = None + # The median (q50) internal query latency, which the Go SDK labels as the average. + avg_latency_ms: float | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class ConnectionMetrics: + active_connections: int | None = None + max_connections: int | None = None + usage_percent: float | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class StorageMetrics: + page_cache_hit_rate: float | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class InstanceHealth: + """A health summary built from an instance's metrics (Go: ``PrometheusHealthMetrics``). + + A value is ``None`` when the endpoint didn't report the metric. Go reports ``0`` in that + case. + """ + + instance_id: str + timestamp: datetime + resources: ResourceMetrics + query: QueryMetrics + connections: ConnectionMetrics + storage: StorageMetrics + overall_status: HealthStatus + issues: tuple[str, ...] = () + recommendations: tuple[str, ...] = () diff --git a/src/aura_python_sdk/services/__init__.py b/src/aura_python_sdk/services/__init__.py index f084f56..881a931 100644 --- a/src/aura_python_sdk/services/__init__.py +++ b/src/aura_python_sdk/services/__init__.py @@ -3,6 +3,7 @@ from aura_python_sdk.services.cmek import CMEKService from aura_python_sdk.services.graph_analytics import GDSSessionService from aura_python_sdk.services.instances import InstanceService +from aura_python_sdk.services.prometheus import PrometheusService from aura_python_sdk.services.snapshots import SnapshotService from aura_python_sdk.services.tenants import TenantService @@ -10,6 +11,7 @@ "CMEKService", "GDSSessionService", "InstanceService", + "PrometheusService", "SnapshotService", "TenantService", ] diff --git a/src/aura_python_sdk/services/prometheus.py b/src/aura_python_sdk/services/prometheus.py new file mode 100644 index 0000000..ce45469 --- /dev/null +++ b/src/aura_python_sdk/services/prometheus.py @@ -0,0 +1,241 @@ +"""``client.prometheus`` (Go: PrometheusService).""" + +from __future__ import annotations + +import logging +from collections.abc import Mapping +from datetime import UTC, datetime +from urllib.parse import urlsplit + +from aura_python_sdk import _validation as validate +from aura_python_sdk._errors import AuraResponseError, AuraValidationError, MetricNotFoundError +from aura_python_sdk._internal._request import RequestService +from aura_python_sdk._internal.metrics._parser import parse_exposition +from aura_python_sdk.models.prometheus import ( + ConnectionMetrics, + HealthStatus, + InstanceHealth, + PrometheusMetrics, + QueryMetrics, + ResourceMetrics, + StorageMetrics, +) +from aura_python_sdk.services._base import Service + +# The Aura bearer token is sent with every metrics request, so by default only Aura's own +# metrics hosts are allowed. +_TRUSTED_METRICS_DOMAIN = "neo4j.io" + + +class PrometheusService(Service): + """Aura's Prometheus metrics endpoints. + + Get an endpoint from ``client.tenants.get_metrics_integration(tenant_id).endpoint`` or + ``client.instances.get(instance_id).metrics_integration_url``. + """ + + def __init__( + self, api: RequestService, logger: logging.Logger, *, allow_untrusted_urls: bool = False + ) -> None: + super().__init__(api, logger) + self._allow_untrusted_urls = allow_untrusted_urls + + def fetch_raw_metrics(self, prometheus_url: str) -> PrometheusMetrics: + """Fetch and parse every metric from a metrics endpoint.""" + url = self._check_url(prometheus_url) + self._logger.debug("fetching Prometheus metrics", extra={"url": url}) + body = self._api.get(url).body + try: + text = body.decode("utf-8") + except UnicodeDecodeError as exc: + raise AuraResponseError("metrics response is not valid UTF-8") from exc + metrics = PrometheusMetrics(metrics=parse_exposition(text)) + self._logger.debug("Prometheus metrics fetched", extra={"count": len(metrics.metrics)}) + return metrics + + def get_metric_value( + self, + metrics: PrometheusMetrics, + name: str, + label_filters: Mapping[str, str] | None = None, + ) -> float: + """The mean value of ``name`` across every sample whose labels match ``label_filters``. + + Raises :class:`MetricNotFoundError` if nothing matches. + """ + if not isinstance(metrics, PrometheusMetrics): + raise AuraValidationError("metrics must be a PrometheusMetrics") + samples = metrics.metrics.get(name) + if not samples: + raise MetricNotFoundError(f"metric {name} not found") + filters = dict(label_filters or {}) + matching = [s for s in samples if all(s.labels.get(k) == v for k, v in filters.items())] + if not matching: + raise MetricNotFoundError( + f"no matching metrics found for {name} with filters {filters}" + ) + return sum(s.value for s in matching) / len(matching) + + def get_instance_health(self, instance_id: str, prometheus_url: str) -> InstanceHealth: + """Summarise an instance's CPU, memory, query, connection and page cache metrics. + + Uses the same metrics, thresholds and status logic as the Go SDK. + """ + instance_id = validate.instance_id(instance_id) + metrics = self.fetch_raw_metrics(prometheus_url) + + def value(name: str) -> float | None: + try: + return self.get_metric_value(metrics, name) + except MetricNotFoundError: + self._logger.warning("metric not available", extra={"metric": name}) + return None + + cpu_usage = value("neo4j_aura_cpu_usage") + cpu_limit = value("neo4j_aura_cpu_limit") if cpu_usage is not None else None + heap_ratio = value("neo4j_dbms_vm_heap_used_ratio") + resources = ResourceMetrics( + cpu_usage_percent=( + cpu_usage / cpu_limit * 100 + if cpu_usage is not None and cpu_limit and cpu_limit > 0 + else None + ), + memory_usage_percent=heap_ratio * 100 if heap_ratio is not None else None, + ) + + query = QueryMetrics( + query_execution_total=value("neo4j_db_query_execution_success_total"), + avg_latency_ms=value("neo4j_db_query_execution_internal_latency_q50"), + ) + + idle = value("neo4j_dbms_bolt_connections_idle") + running = value("neo4j_dbms_bolt_connections_running") + max_connections = value("neo4j_dbms_bolt_connections_max_count") + active = int(idle + running) if idle is not None and running is not None else None + connections = ConnectionMetrics( + active_connections=active, + max_connections=int(max_connections) + if max_connections and max_connections > 0 + else None, + usage_percent=( + active / max_connections * 100 + if active is not None and max_connections and max_connections > 0 + else None + ), + ) + + hit_ratio = value("neo4j_dbms_page_cache_hit_ratio_per_minute") + storage = StorageMetrics( + page_cache_hit_rate=hit_ratio * 100 if hit_ratio is not None else None + ) + + status, issues, recommendations = assess_health(resources, connections, storage) + self._logger.info( + "instance health assessed", extra={"instance_id": instance_id, "status": status} + ) + return InstanceHealth( + instance_id=instance_id, + timestamp=datetime.now(UTC), + resources=resources, + query=query, + connections=connections, + storage=storage, + overall_status=status, + issues=tuple(issues), + recommendations=tuple(recommendations), + ) + + def _check_url(self, prometheus_url: str) -> str: + url = validate.require_non_empty("prometheus URL", prometheus_url) + parts = urlsplit(url) + if parts.scheme not in ("https", "http") or not parts.hostname: + raise AuraValidationError(f"prometheus URL is not a valid http(s) URL: {url!r}") + if self._allow_untrusted_urls: + return url + host = parts.hostname.lower() + trusted = host == _TRUSTED_METRICS_DOMAIN or host.endswith(f".{_TRUSTED_METRICS_DOMAIN}") + if parts.scheme != "https" or not trusted: + raise AuraValidationError( + f"prometheus URL must be an https://*.{_TRUSTED_METRICS_DOMAIN} address, because " + "the Aura API token is sent with the request" + ) + return url + + +def assess_health( + resources: ResourceMetrics, connections: ConnectionMetrics, storage: StorageMetrics +) -> tuple[HealthStatus, list[str], list[str]]: + """Apply the Go SDK's thresholds. Returns (status, issues, recommendations).""" + status = HealthStatus.HEALTHY + issues: list[str] = [] + recommendations: list[str] = [] + + def flag(level: HealthStatus, issue: str, recommendation: str) -> None: + nonlocal status + issues.append(issue) + recommendations.append(recommendation) + if level is HealthStatus.CRITICAL or status is HealthStatus.HEALTHY: + status = level + + cpu = resources.cpu_usage_percent + if cpu is not None: + if cpu > 95: + flag( + HealthStatus.CRITICAL, + f"Critical CPU usage: {cpu:.1f}%", + "Scale to a larger instance size immediately", + ) + elif cpu > 80: + flag( + HealthStatus.WARNING, + f"High CPU usage: {cpu:.1f}%", + "Consider scaling to a larger instance size", + ) + + memory = resources.memory_usage_percent + if memory is not None: + if memory > 95: + flag( + HealthStatus.CRITICAL, + f"Critical memory usage: {memory:.1f}%", + "Scale to a larger memory instance immediately", + ) + elif memory > 85: + flag( + HealthStatus.WARNING, + f"High memory usage: {memory:.1f}%", + "Consider scaling to a larger memory instance", + ) + + usage = connections.usage_percent + if usage is not None and connections.max_connections: + if usage > 95: + flag( + HealthStatus.CRITICAL, + f"Critical connection usage: {usage:.1f}%", + "Reduce active connections immediately; review connection pooling", + ) + elif usage > 80: + flag( + HealthStatus.WARNING, + f"High connection usage: {usage:.1f}%", + "Review connection pooling configuration in your application", + ) + + hit_rate = storage.page_cache_hit_rate + # As in Go, a hit rate of exactly 0 is treated as "no data". + if hit_rate: + if hit_rate < 20: + flag( + HealthStatus.CRITICAL, + f"Critical page cache hit rate: {hit_rate:.1f}%", + "Increase page cache size immediately; query performance is severely degraded", + ) + elif hit_rate < 50: + flag( + HealthStatus.WARNING, + f"Low page cache hit rate: {hit_rate:.1f}%", + "Consider increasing page cache size for better performance", + ) + + return status, issues, recommendations diff --git a/tests/unit/test_import_boundaries.py b/tests/unit/test_import_boundaries.py index b4bedaa..41e2ceb 100644 --- a/tests/unit/test_import_boundaries.py +++ b/tests/unit/test_import_boundaries.py @@ -24,7 +24,6 @@ # third-party top-level module -> the single module (relative to src/) allowed to import it WRAPPED_DEPENDENCIES: dict[str, str] = { "httpx": f"{PACKAGE}/_internal/http/_httpx.py", - "prometheus_client": f"{PACKAGE}/_internal/metrics/_parser.py", } diff --git a/tests/unit/test_prometheus_parser.py b/tests/unit/test_prometheus_parser.py new file mode 100644 index 0000000..b0724db --- /dev/null +++ b/tests/unit/test_prometheus_parser.py @@ -0,0 +1,126 @@ +import math + +import pytest + +from aura_python_sdk import AuraResponseError, PrometheusMetric +from aura_python_sdk._internal.metrics._parser import parse_exposition + + +def test_gauge_with_labels_and_help() -> None: + text = """ +# HELP neo4j_aura_cpu_usage CPU usage (cores) +# TYPE neo4j_aura_cpu_usage gauge +neo4j_aura_cpu_usage{availability_zone="europe-west2-c",instance_id="c9f0d13a"} 0.023206 +neo4j_aura_cpu_usage{availability_zone="europe-west2-b",instance_id="c9f0d13a"} 0.5 +""" + metrics = parse_exposition(text) + assert list(metrics) == ["neo4j_aura_cpu_usage"] + first, second = metrics["neo4j_aura_cpu_usage"] + assert first == PrometheusMetric( + name="neo4j_aura_cpu_usage", + labels={"availability_zone": "europe-west2-c", "instance_id": "c9f0d13a"}, + value=0.023206, + ) + assert second.value == 0.5 + + +def test_counters_keep_their_type_line_name() -> None: + # prometheus_client would rename plain_counter to plain_counter_total; expfmt (Go) does not. + text = """ +# TYPE neo4j_db_query_execution_success_total counter +neo4j_db_query_execution_success_total{db="neo4j"} 42 +# TYPE plain_counter counter +plain_counter 7 +""" + metrics = parse_exposition(text) + assert set(metrics) == {"neo4j_db_query_execution_success_total", "plain_counter"} + assert metrics["plain_counter"][0].value == 7 + + +def test_summary_and_histogram_use_sum_per_label_set() -> None: + text = """ +# TYPE latency summary +latency{db="a",quantile="0.5"} 1 +latency{db="a",quantile="0.99"} 9 +latency_sum{db="a"} 10 +latency_count{db="a"} 4 +latency_sum{db="b"} 20 +latency_count{db="b"} 5 +# TYPE sizes histogram +sizes_bucket{le="1"} 1 +sizes_bucket{le="+Inf"} 2 +sizes_sum 3.5 +sizes_count 2 +""" + metrics = parse_exposition(text) + assert set(metrics) == {"latency", "sizes"} + assert [(m.labels, m.value) for m in metrics["latency"]] == [ + ({"db": "a"}, 10), + ({"db": "b"}, 20), + ] + assert metrics["sizes"][0].value == 3.5 + + +def test_untyped_samples_are_keyed_by_name() -> None: + metrics = parse_exposition("some_metric 1\nsome_metric_sum 2\n") + assert set(metrics) == {"some_metric", "some_metric_sum"} + + +def test_special_values_and_timestamps() -> None: + text = "a NaN\nb +Inf 1700000000000\nc -Inf -5\nd 1.5e3\n" + metrics = parse_exposition(text) + assert math.isnan(metrics["a"][0].value) + assert metrics["b"][0].value == math.inf + assert metrics["b"][0].timestamp_ms == 1700000000000 + assert metrics["c"][0].value == -math.inf + assert metrics["c"][0].timestamp_ms == -5 + assert metrics["d"][0].value == 1500.0 + assert metrics["d"][0].timestamp_ms is None + + +def test_label_escapes_spacing_and_trailing_comma() -> None: + text = r'm{ a = "q\"uote" , b="back\\slash",c="new\nline", } 1' + [metric] = parse_exposition(text)["m"] + assert metric.labels == {"a": 'q"uote', "b": "back\\slash", "c": "new\nline"} + + +def test_empty_labels_and_braces_inside_values() -> None: + [metric] = parse_exposition('m{} 1\nn{path="/a{b}c"} 2')["m"] + assert metric.labels == {} + assert parse_exposition('n{path="/a{b}c"} 2')["n"][0].labels == {"path": "/a{b}c"} + + +def test_comments_and_blank_lines_are_ignored() -> None: + text = "# just a comment\n\n# HELP x help text\n# TYPE\nx 1\n" + assert parse_exposition(text)["x"][0].value == 1 + + +def test_empty_input() -> None: + assert parse_exposition("") == {} + + +@pytest.mark.parametrize( + ("text", "message"), + [ + ("9bad 1", "expected a metric name"), + ("m", "expected a value"), + ("m 1 2 3", "expected a value"), + ("m abc", "invalid value"), + ("m 1 soon", "invalid timestamp"), + ('m{a="1" 1', "expected ',' or '}'"), + ("m{a=1} 1", "expected '\"'"), + ('m{a "1"} 1', "expected '='"), + ('m{="1"} 1', "expected a label name"), + ('m{a="unterminated} 1', "unterminated label value"), + (r'm{a="bad\tescape"} 1', "invalid escape"), + ], +) +def test_malformed_lines(text: str, message: str) -> None: + with pytest.raises(AuraResponseError, match="line 1") as info: + parse_exposition(text) + assert message in str(info.value) + + +def test_error_reports_line_number() -> None: + with pytest.raises(AuraResponseError, match="line 3"): + parse_exposition("a 1\nb 2\nc oops\n") diff --git a/tests/unit/test_prometheus_service.py b/tests/unit/test_prometheus_service.py new file mode 100644 index 0000000..97ebfa1 --- /dev/null +++ b/tests/unit/test_prometheus_service.py @@ -0,0 +1,239 @@ +from typing import Any + +import pytest + +from aura_python_sdk import ( + AuraClient, + AuraResponseError, + AuraValidationError, + ConnectionMetrics, + HealthStatus, + HttpResponse, + MetricNotFoundError, + PrometheusMetric, + PrometheusMetrics, + ResourceMetrics, + StorageMetrics, +) +from aura_python_sdk.services.prometheus import assess_health +from tests.fakes import FakeTransport, token_response +from tests.unit.conftest import INSTANCE_ID, Api + +METRICS_URL = "https://customer-metrics-api.neo4j.io/api/v1/proj/2f49c2b3/metrics" + +HEALTHY = """ +# TYPE neo4j_aura_cpu_usage gauge +neo4j_aura_cpu_usage{instance_mode="PRIMARY"} 0.5 +neo4j_aura_cpu_usage{instance_mode="SECONDARY"} 1.5 +# TYPE neo4j_aura_cpu_limit gauge +neo4j_aura_cpu_limit 4 +# TYPE neo4j_dbms_vm_heap_used_ratio gauge +neo4j_dbms_vm_heap_used_ratio 0.42 +# TYPE neo4j_db_query_execution_success_total counter +neo4j_db_query_execution_success_total 1200 +# TYPE neo4j_db_query_execution_internal_latency_q50 gauge +neo4j_db_query_execution_internal_latency_q50 3.5 +# TYPE neo4j_dbms_bolt_connections_idle gauge +neo4j_dbms_bolt_connections_idle 10 +# TYPE neo4j_dbms_bolt_connections_running gauge +neo4j_dbms_bolt_connections_running 5 +# TYPE neo4j_dbms_bolt_connections_max_count gauge +neo4j_dbms_bolt_connections_max_count 100 +# TYPE neo4j_dbms_page_cache_hit_ratio_per_minute gauge +neo4j_dbms_page_cache_hit_ratio_per_minute 0.98 +""" + + +def _reply_text(api: Api, text: str) -> None: + api.transport.queue(HttpResponse(200, {"Content-Type": "text/plain"}, text.encode())) + + +def _metrics(**samples: list[tuple[dict[str, str], float]]) -> PrometheusMetrics: + return PrometheusMetrics( + metrics={ + name: tuple(PrometheusMetric(name=name, labels=labels, value=v) for labels, v in values) + for name, values in samples.items() + } + ) + + +def test_fetch_raw_metrics_sends_token_to_metrics_url(api: Api) -> None: + _reply_text(api, HEALTHY) + metrics = api.client.prometheus.fetch_raw_metrics(METRICS_URL) + assert len(metrics.metrics) == 9 + assert api.request.url == METRICS_URL + assert api.request.headers["Authorization"].startswith("Bearer ") + + +@pytest.mark.parametrize( + "url", + [ + "http://customer-metrics-api.neo4j.io/metrics", + "https://evil.example.com/metrics", + "https://neo4j.io.evil.com/metrics", + "https://evilneo4j.io/metrics", + "ftp://customer-metrics-api.neo4j.io/metrics", + "not a url", + "", + ], +) +def test_untrusted_urls_are_refused_before_sending_the_token(api: Api, url: str) -> None: + with pytest.raises(AuraValidationError, match="prometheus URL"): + api.client.prometheus.fetch_raw_metrics(url) + api.assert_no_request() + + +def test_apex_domain_is_trusted(api: Api) -> None: + _reply_text(api, "") + api.client.prometheus.fetch_raw_metrics("https://neo4j.io/metrics") + + +def test_insecure_client_allows_local_metrics_urls() -> None: + transport = FakeTransport([token_response(), HttpResponse(200, body=b"up 1")]) + client = AuraClient( + client_id="id", + client_secret="secret", + base_url="http://localhost:9000", + allow_insecure_base_url=True, + transport=transport, + ) + metrics = client.prometheus.fetch_raw_metrics("http://localhost:9100/metrics") + assert metrics.metrics["up"][0].value == 1 + + +def test_non_utf8_body(api: Api) -> None: + api.transport.queue(HttpResponse(200, body=b"\xff\xfe")) + with pytest.raises(AuraResponseError, match="UTF-8"): + api.client.prometheus.fetch_raw_metrics(METRICS_URL) + + +def test_get_metric_value_averages_all_samples(api: Api) -> None: + metrics = _metrics(cpu=[({"zone": "a"}, 1.0), ({"zone": "b"}, 3.0)]) + assert api.client.prometheus.get_metric_value(metrics, "cpu") == 2.0 + + +def test_get_metric_value_with_label_filters(api: Api) -> None: + metrics = _metrics( + cpu=[ + ({"zone": "a", "mode": "PRIMARY"}, 1.0), + ({"zone": "a", "mode": "SECONDARY"}, 5.0), + ({"zone": "b", "mode": "PRIMARY"}, 3.0), + ] + ) + prometheus = api.client.prometheus + assert prometheus.get_metric_value(metrics, "cpu", {"mode": "PRIMARY"}) == 2.0 + assert prometheus.get_metric_value(metrics, "cpu", {"zone": "a", "mode": "SECONDARY"}) == 5.0 + + +def test_get_metric_value_not_found(api: Api) -> None: + metrics = _metrics(cpu=[({"zone": "a"}, 1.0)]) + with pytest.raises(MetricNotFoundError, match="metric memory not found"): + api.client.prometheus.get_metric_value(metrics, "memory") + with pytest.raises(MetricNotFoundError, match="no matching metrics"): + api.client.prometheus.get_metric_value(metrics, "cpu", {"zone": "z"}) + assert issubclass(MetricNotFoundError, LookupError) + + +def test_get_metric_value_requires_metrics(api: Api) -> None: + with pytest.raises(AuraValidationError, match="PrometheusMetrics"): + api.client.prometheus.get_metric_value({"cpu": []}, "cpu") # type: ignore[arg-type] + + +def test_get_instance_health_healthy(api: Api) -> None: + _reply_text(api, HEALTHY) + health = api.client.prometheus.get_instance_health(INSTANCE_ID, METRICS_URL) + + assert health.instance_id == INSTANCE_ID + assert health.overall_status is HealthStatus.HEALTHY + assert health.resources.cpu_usage_percent == pytest.approx(25.0) # mean 1.0 of 4 cores + assert health.resources.memory_usage_percent == pytest.approx(42.0) + assert health.query.query_execution_total == 1200 + assert health.query.avg_latency_ms == 3.5 + assert health.connections == ConnectionMetrics( + active_connections=15, max_connections=100, usage_percent=15.0 + ) + assert health.storage.page_cache_hit_rate == pytest.approx(98.0) + assert health.issues == () + assert health.recommendations == () + assert health.timestamp.tzinfo is not None + + +def test_get_instance_health_with_missing_metrics( + api: Api, caplog: pytest.LogCaptureFixture +) -> None: + _reply_text(api, "# TYPE neo4j_aura_cpu_usage gauge\nneo4j_aura_cpu_usage 3.9\n") + health = api.client.prometheus.get_instance_health(INSTANCE_ID, METRICS_URL) + assert health.resources.cpu_usage_percent is None # no cpu_limit, so no percentage + assert health.resources.memory_usage_percent is None + assert health.connections == ConnectionMetrics() + assert health.overall_status is HealthStatus.HEALTHY + assert any("metric not available" in r.getMessage() for r in caplog.records) + + +def test_get_instance_health_validates_before_fetching(api: Api) -> None: + with pytest.raises(AuraValidationError, match="instance ID"): + api.client.prometheus.get_instance_health("bad", METRICS_URL) + with pytest.raises(AuraValidationError, match="prometheus URL"): + api.client.prometheus.get_instance_health(INSTANCE_ID, "https://example.com/metrics") + api.assert_no_request() + + +def _assess(**kwargs: Any) -> tuple[HealthStatus, list[str], list[str]]: + return assess_health( + ResourceMetrics( + cpu_usage_percent=kwargs.get("cpu"), memory_usage_percent=kwargs.get("memory") + ), + ConnectionMetrics( + max_connections=kwargs.get("max_connections", 100), usage_percent=kwargs.get("conns") + ), + StorageMetrics(page_cache_hit_rate=kwargs.get("hit_rate")), + ) + + +@pytest.mark.parametrize( + ("kwargs", "status", "issue"), + [ + ({"cpu": 80.0}, HealthStatus.HEALTHY, None), + ({"cpu": 80.1}, HealthStatus.WARNING, "High CPU usage: 80.1%"), + ({"cpu": 95.5}, HealthStatus.CRITICAL, "Critical CPU usage: 95.5%"), + ({"memory": 85.0}, HealthStatus.HEALTHY, None), + ({"memory": 90.0}, HealthStatus.WARNING, "High memory usage: 90.0%"), + ({"memory": 99.0}, HealthStatus.CRITICAL, "Critical memory usage: 99.0%"), + ({"conns": 81.0}, HealthStatus.WARNING, "High connection usage: 81.0%"), + ({"conns": 96.0}, HealthStatus.CRITICAL, "Critical connection usage: 96.0%"), + ({"conns": 99.0, "max_connections": None}, HealthStatus.HEALTHY, None), + ({"hit_rate": 50.0}, HealthStatus.HEALTHY, None), + ({"hit_rate": 49.0}, HealthStatus.WARNING, "Low page cache hit rate: 49.0%"), + ({"hit_rate": 10.0}, HealthStatus.CRITICAL, "Critical page cache hit rate: 10.0%"), + ({"hit_rate": 0.0}, HealthStatus.HEALTHY, None), # Go treats 0 as "no data" + ], +) +def test_thresholds_match_go( + kwargs: dict[str, Any], status: HealthStatus, issue: str | None +) -> None: + result_status, issues, recommendations = _assess(**kwargs) + assert result_status is status + assert issues == ([] if issue is None else [issue]) + assert len(recommendations) == len(issues) + + +def test_critical_is_not_downgraded_by_a_later_warning() -> None: + status, issues, _ = _assess(cpu=99.0, memory=90.0) + assert status is HealthStatus.CRITICAL + assert issues == ["Critical CPU usage: 99.0%", "High memory usage: 90.0%"] + + +def test_warning_is_upgraded_by_a_later_critical() -> None: + status, _, _ = _assess(cpu=85.0, hit_rate=5.0) + assert status is HealthStatus.CRITICAL + + +def test_critical_health_end_to_end(api: Api) -> None: + _reply_text( + api, + HEALTHY.replace("neo4j_dbms_vm_heap_used_ratio 0.42", "neo4j_dbms_vm_heap_used_ratio 0.97"), + ) + health = api.client.prometheus.get_instance_health(INSTANCE_ID, METRICS_URL) + assert health.overall_status is HealthStatus.CRITICAL + assert health.issues == ("Critical memory usage: 97.0%",) + assert health.recommendations == ("Scale to a larger memory instance immediately",) diff --git a/uv.lock b/uv.lock index 3f69992..718bd1f 100644 --- a/uv.lock +++ b/uv.lock @@ -90,11 +90,6 @@ dependencies = [ { name = "httpx" }, ] -[package.optional-dependencies] -prometheus = [ - { name = "prometheus-client" }, -] - [package.dev-dependencies] dev = [ { name = "mypy" }, @@ -106,11 +101,7 @@ dev = [ ] [package.metadata] -requires-dist = [ - { name = "httpx", specifier = ">=0.27,<1" }, - { name = "prometheus-client", marker = "extra == 'prometheus'", specifier = ">=0.20" }, -] -provides-extras = ["prometheus"] +requires-dist = [{ name = "httpx", specifier = ">=0.27,<1" }] [package.metadata.requires-dev] dev = [ @@ -526,15 +517,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/54/20/4d324d65cc6d9205fabedc306948156824eb9f0ee1633355a8f7ec5c66bf/pluggy-1.6.0-py3-none-any.whl", hash = "sha256:e920276dd6813095e9377c0bc5566d94c932c33b27a3e3945d8389c374dd4746", size = 20538, upload-time = "2025-05-15T12:30:06.134Z" }, ] -[[package]] -name = "prometheus-client" -version = "0.26.0" -source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/52/73/f1334c29c2af4cd9dba6c7817e61b611bd0215e2eb5565c6064a4de18802/prometheus_client-0.26.0.tar.gz", hash = "sha256:04a91bcf94e2cf74a44a1a874d651a2e853ed354b6e822f3b7487751465d5c2b", size = 92910, upload-time = "2026-07-24T19:36:41.893Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/eb/a3/b69efbf4143b5b9859b977770bbbabcc2796b702fa69dc40271e45cd5a56/prometheus_client-0.26.0-py3-none-any.whl", hash = "sha256:fa93d06737aa02bacd05794768508bb97d2fbee28cb3bca04eaae92f0ca953d6", size = 64494, upload-time = "2026-07-24T19:36:40.854Z" }, -] - [[package]] name = "pygments" version = "2.21.0" From 43ddb78c35dba35ce6c832d00357ac7f756b02b8 Mon Sep 17 00:00:00 2001 From: Jonathan Giffard Date: Tue, 29 Sep 2026 20:44:02 +0100 Subject: [PATCH 6/8] Add README, examples, black-box and live tests, and release workflow Phase 7 of PLAN.md. Full README, Python ports of the Go v1 examples, a black-box suite over real sockets with the httpx transport, opt-in live integration tests (read-only unless writes are enabled), a changelog, and a tag-triggered workflow that tests, builds and publishes to PyPI. Co-Authored-By: Claude Opus 5.5 --- .github/workflows/release.yml | 96 ++++++++ CHANGELOG.md | 26 +++ PLAN.md | 2 + README.md | 353 ++++++++++++++++++++++++++++- examples/create_delete_instance.py | 87 +++++++ examples/get_instance_details.py | 38 ++++ examples/list_instances.py | 29 +++ examples/list_snapshots.py | 41 ++++ examples/list_tenants.py | 30 +++ examples/prometheus.py | 62 +++++ examples/restore_from_snapshot.py | 37 +++ examples/take_snapshot.py | 32 +++ pyproject.toml | 4 +- tests/blackbox/__init__.py | 0 tests/blackbox/conftest.py | 140 ++++++++++++ tests/blackbox/test_blackbox.py | 194 ++++++++++++++++ tests/integration/__init__.py | 0 tests/integration/test_live.py | 123 ++++++++++ 18 files changed, 1287 insertions(+), 7 deletions(-) create mode 100644 .github/workflows/release.yml create mode 100644 CHANGELOG.md create mode 100644 examples/create_delete_instance.py create mode 100644 examples/get_instance_details.py create mode 100644 examples/list_instances.py create mode 100644 examples/list_snapshots.py create mode 100644 examples/list_tenants.py create mode 100644 examples/prometheus.py create mode 100644 examples/restore_from_snapshot.py create mode 100644 examples/take_snapshot.py create mode 100644 tests/blackbox/__init__.py create mode 100644 tests/blackbox/conftest.py create mode 100644 tests/blackbox/test_blackbox.py create mode 100644 tests/integration/__init__.py create mode 100644 tests/integration/test_live.py diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml new file mode 100644 index 0000000..9e189b9 --- /dev/null +++ b/.github/workflows/release.yml @@ -0,0 +1,96 @@ +name: Release + +on: + push: + tags: + - "v[0-9]+.[0-9]+.[0-9]+*" + +permissions: + contents: read + +jobs: + build: + name: Test and build + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - uses: astral-sh/setup-uv@v6 + with: + python-version: "3.11" + + # The tag must match __version__, so the published version is the one in the source. + - name: Check tag matches package version + run: | + VERSION=$(sed -n 's/^__version__ = "\(.*\)"/\1/p' src/aura_python_sdk/_version.py) + if [ "v${VERSION}" != "${GITHUB_REF_NAME}" ]; then + echo "Tag ${GITHUB_REF_NAME} does not match __version__ ${VERSION}" >&2 + exit 1 + fi + + # Gate: the release is only built if lint, types and tests pass. + - run: uv sync --all-extras + - run: uv run ruff format --check + - run: uv run ruff check + - run: uv run mypy + - run: uv run pytest -m "not integration" + + - run: uv build + - name: Smoke-test the wheel in a clean environment + run: | + uv venv /tmp/smoke + uv pip install --python /tmp/smoke/bin/python dist/*.whl + /tmp/smoke/bin/python -c "import aura_python_sdk as aura; print(aura.__version__)" + + - uses: actions/upload-artifact@v4 + with: + name: dist + path: dist/ + + publish: + name: Publish to PyPI + needs: build + runs-on: ubuntu-latest + environment: + name: pypi + url: https://pypi.org/project/aura-python-sdk/ + permissions: + id-token: write # PyPI trusted publishing; no API token is stored in the repo + steps: + - uses: actions/download-artifact@v4 + with: + name: dist + path: dist/ + - uses: pypa/gh-action-pypi-publish@release/v1 + + github-release: + name: Create GitHub release + needs: publish + runs-on: ubuntu-latest + permissions: + contents: write + steps: + - uses: actions/checkout@v4 + - uses: actions/download-artifact@v4 + with: + name: dist + path: dist/ + + # Collect the lines between "## vX.Y.Z" and the next "## " heading in CHANGELOG.md. + - name: Extract release notes + run: | + awk -v ver="${GITHUB_REF_NAME}" ' + /^## / && ($2 == ver) { found=1; next } + found && /^## / { exit } + found { print } + ' CHANGELOG.md | sed '/./,$!d' > release_notes.md + if [ ! -s release_notes.md ]; then + echo "See CHANGELOG.md for details." > release_notes.md + fi + cat release_notes.md + + - uses: softprops/action-gh-release@v2 + with: + name: ${{ github.ref_name }} + body_path: release_notes.md + files: dist/* + prerelease: ${{ contains(github.ref_name, 'a') || contains(github.ref_name, 'b') || contains(github.ref_name, 'rc') }} diff --git a/CHANGELOG.md b/CHANGELOG.md new file mode 100644 index 0000000..648a600 --- /dev/null +++ b/CHANGELOG.md @@ -0,0 +1,26 @@ +# Changelog + +All notable changes to this project are documented here. + +The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.1.0/), and this project +follows [Semantic Versioning](https://semver.org/spec/v2.0.0.html). The release workflow publishes +the `## vX.Y.Z` section that matches the pushed tag as the GitHub release notes. + +## Unreleased + +### Added + +- `AuraClient` for the Aura API v1, with the Go SDK's options as keyword arguments, `from_env()`, + and context-manager support. +- Services matching the Go SDK: `tenants`, `instances`, `snapshots`, `cmek`, `graph_analytics` and + `prometheus`. +- Full v1 spec coverage beyond the Go SDK: `instances.estimate_size`, `instances.upgrade`, + `cmek.get` / `create` / `delete`, list filters, and the `storage`, `vector_optimized` and + `graph_analytics_plugin` update fields. +- Frozen dataclass models and `StrEnum`s that tolerate values the SDK doesn't know yet. +- An exception class per error: `NotFoundError`, `RateLimitError` (with `retry_after`) and others. +- A pluggable `HttpTransport`, with an httpx implementation as the default. +- Only network failures are retried, and a non-idempotent request is never re-sent once it may + have reached the server. +- A stdlib Prometheus text-format parser whose output matches the Go SDK, and + `get_instance_health` with the Go SDK's thresholds. diff --git a/PLAN.md b/PLAN.md index 5129524..080fd8a 100644 --- a/PLAN.md +++ b/PLAN.md @@ -268,6 +268,8 @@ Dev tooling: `uv`, `ruff` (lint and format), `mypy --strict`, `pytest`, `pytest- ## 5. Phases +**Status:** phases 1–7 are done. Phase 8 (async) is not started. + 1. **Scaffold**: pyproject, uv, ruff, mypy, pytest config, CI workflow, and the import-boundary test. 2. **Core**: config/options, errors, `HttpTransport` + `HttpxTransport` (retries, size cap), `TokenManager`, `RequestService`, and `AuraClient` with no services yet. Unit-tested to Go's diff --git a/README.md b/README.md index 6037acf..4993338 100644 --- a/README.md +++ b/README.md @@ -1,11 +1,338 @@ # aura-python-sdk -Python client library for the [Neo4j Aura API](https://neo4j.com/docs/aura/api/overview/) (v1), -modelled on [aura-go-sdk](https://github.com/neo4j-contrib/aura-go-sdk). +A Python client for the [Neo4j Aura API](https://neo4j.com/docs/aura/api/overview/) (v1). For +example, `client.instances.list()` returns your Aura instances. It is modelled on +[aura-go-sdk](https://github.com/neo4j-contrib/aura-go-sdk) and covers the whole v1 API. -> Status: under development. See [PLAN.md](PLAN.md). +- Typed throughout (`py.typed`, checked with `mypy --strict`), using frozen dataclass models. +- One runtime dependency, [httpx](https://www.python-httpx.org/), kept behind the SDK's own + transport interface. +- Client-side validation, automatic OAuth token handling, safe retries, and one exception + class per error. -Requires Python 3.11+. +You need an Aura API client ID and secret. See +[Aura API authentication](https://neo4j.com/docs/aura/api/authentication/). + +## Contents + +- [Installation](#installation) +- [Quick start](#quick-start) +- [Configuration](#configuration) +- [Timeouts and retries](#timeouts-and-retries) +- [Tenants](#tenants) +- [Instances](#instances) +- [Snapshots](#snapshots) +- [Customer-managed keys](#customer-managed-keys) +- [Graph Analytics sessions](#graph-analytics-sessions) +- [Prometheus metrics](#prometheus-metrics) +- [Error handling](#error-handling) +- [Logging](#logging) +- [Custom transports and testing](#custom-transports-and-testing) +- [Coming from the Go SDK](#coming-from-the-go-sdk) +- [Development](#development) + +## Installation + +Requires Python 3.11 or later. + +```sh +pip install aura-python-sdk +``` + +## Quick start + +```python +import aura_python_sdk as aura + +with aura.AuraClient(client_id="your-client-id", client_secret="your-client-secret") as client: + for instance in client.instances.list(): + print(f"{instance.name} ({instance.id})") +``` + +Or read the credentials from the `AURA_CLIENT_ID` and `AURA_CLIENT_SECRET` environment +variables: + +```python +client = aura.AuraClient.from_env() +``` + +Using the client as a context manager (or calling `client.close()`) releases its pooled +connections. + +## Configuration + +Every option is keyword-only. An invalid option raises `AuraConfigurationError` straight away. + +```python +import logging + +client = aura.AuraClient( + client_id="...", + client_secret="...", + timeout=60, # seconds per call (default 120) + max_retries=5, # network-failure retries (default 3) + max_response_size=20 * 1024 * 1024, # bytes (default 10 MB) + base_url="https://api.staging.neo4j.io", + user_agent="my-app/1.0", # default "aura-python-sdk/" + default_headers={"X-Team": "platform"}, # added to every request + logger=logging.getLogger("my-app.aura"), +) +``` + +| Option | Default | Notes | +| --- | --- | --- | +| `client_id`, `client_secret` | required | Must not be empty. | +| `base_url` | `https://api.neo4j.io` | Must be HTTPS. | +| `allow_insecure_base_url` | `False` | Allows an `http://` base URL, and metrics URLs outside `*.neo4j.io`. For local test servers only. | +| `timeout` | `120` | Seconds allowed for each call (see below). | +| `max_retries` | `3` | `0` disables retries. | +| `max_response_size` | 10 MB | Larger responses raise `AuraResponseError`. | +| `user_agent` | `aura-python-sdk/` | | +| `default_headers` | none | `Authorization`, `Content-Type` and `User-Agent` are ignored. | +| `logger` | `logging.getLogger("aura_python_sdk")` | | +| `transport` | built-in httpx transport | See [Custom transports](#custom-transports-and-testing). | + +## Timeouts and retries + +`timeout` is one deadline for the whole call, covering the OAuth token fetch, every retry and +every backoff. This matches the per-call `context.WithTimeout` in the Go SDK. + +Only network failures are retried, with backoff from 1 s doubling to 5 s. A response with an HTTP +status, including 429 and 5xx, is never retried. If a request might already have reached the +server (a read timeout or a dropped connection), only idempotent methods (`GET`, `PUT`, `DELETE`) +are retried. That means a `create` or `pause` is never sent twice. + +## Tenants + +```python +for tenant in client.tenants.list(): + print(tenant.id, tenant.name) + +tenant = client.tenants.get("6981ace7-efe8-4f5c-b7c5-267b5162ce91") +for config in tenant.instance_configurations: + print(config.type, config.cloud_provider, config.region, config.memory, config.version) + +endpoint = client.tenants.get_metrics_integration(tenant.id).endpoint +``` + +## Instances + +```python +from aura_python_sdk import CloudProvider, InstanceConfig, InstanceStatus, InstanceType + +instances = client.instances.list() # or list(tenant_id=...) +instance = client.instances.get("2f49c2b3") +if instance.status == InstanceStatus.RUNNING: + print(instance.connection_url) + +created = client.instances.create( + InstanceConfig( + name="my-instance", + tenant_id="6981ace7-efe8-4f5c-b7c5-267b5162ce91", + cloud_provider=CloudProvider.GCP, + region="europe-west1", + type=InstanceType.PROFESSIONAL_DB, + version="5", + memory="2GB", + ) +) +print(created.id, created.username, created.password) # the password is shown only once +``` + +Creation is asynchronous: poll `get()` until the status is `running`. See +[examples/create_delete_instance.py](examples/create_delete_instance.py). + +| Method | What it does | +| --- | --- | +| `list(tenant_id=None)` | Summaries of every instance, optionally in one tenant. | +| `get(instance_id)` | Full details. | +| `create(config)` | Starts creating an instance. Returns the initial credentials. | +| `create_from_instance(source_instance_id, config)` | Clones another instance's current data. | +| `create_from_snapshot(source_instance_id, source_snapshot_id, config)` | Creates from an exportable snapshot. | +| `update(instance_id, *, name, memory, storage, vector_optimized, graph_analytics_plugin, cdc_enrichment_mode, secondaries_count)` | Changes only the fields you pass. | +| `pause(instance_id)` / `resume(instance_id)` | | +| `delete(instance_id)` | Cannot be undone. | +| `overwrite_from_instance(instance_id, source_instance_id)` | Replaces the data with another instance's. | +| `overwrite_from_snapshot(instance_id, source_snapshot_id)` | Replaces the data with a snapshot. | +| `estimate_size(*, node_count, relationship_count, instance_type, algorithm_categories)` | Sizing for AuraDS instances. | +| `upgrade(instance_id, *, memory, storage)` | Professional to Business Critical. Pass both sizes, or neither. | + +`CreatedInstance.password` is left out of `repr()`, so logging the object doesn't expose it. + +## Snapshots + +```python +import datetime + +snapshots = client.snapshots.list("2f49c2b3") # today +snapshots = client.snapshots.list("2f49c2b3", datetime.date(2026, 9, 1)) + +started = client.snapshots.create("2f49c2b3") +snapshot = client.snapshots.get("2f49c2b3", started.snapshot_id) +client.snapshots.restore("2f49c2b3", snapshot.snapshot_id) +``` + +## Customer-managed keys + +```python +keys = client.cmek.list() # or list(tenant_id=...) +key = client.cmek.create( + name="Production Key", + key_id="arn:aws:kms:us-west-2:111122223333:key/1234abcd-...", + tenant_id="6981ace7-efe8-4f5c-b7c5-267b5162ce91", + cloud_provider=CloudProvider.AWS, + region="us-west-2", + instance_type=InstanceType.ENTERPRISE_DB, +) +print(client.cmek.get(key.id).status) +client.cmek.delete(key.id) +``` + +## Graph Analytics sessions + +```python +from aura_python_sdk import GDSSessionConfig + +estimate = client.graph_analytics.estimate_size(node_count=1_000_000, relationship_count=5_000_000) + +session = client.graph_analytics.create( + GDSSessionConfig( + name="analysis", + memory=estimate.recommended_size, + ttl="1h", + tenant_id="6981ace7-efe8-4f5c-b7c5-267b5162ce91", + cloud_provider=CloudProvider.GCP, + region="europe-west1", + ) +) +sessions = client.graph_analytics.list(tenant_id=session.tenant_id) +client.graph_analytics.delete(session.id) +``` + +## Prometheus metrics + +Get a metrics endpoint from `tenants.get_metrics_integration()` or from an instance's +`metrics_integration_url`. The client sends its Aura token to that endpoint, so only +`https://*.neo4j.io` URLs are accepted. + +```python +instance = client.instances.get("2f49c2b3") +url = instance.metrics_integration_url + +metrics = client.prometheus.fetch_raw_metrics(url) +cpu = client.prometheus.get_metric_value( + metrics, "neo4j_aura_cpu_usage", {"instance_mode": "PRIMARY"} +) + +health = client.prometheus.get_instance_health(instance.id, url) +print(health.overall_status, health.issues, health.recommendations) +``` + +`get_metric_value` averages every matching sample, and raises `MetricNotFoundError` if nothing +matches. `get_instance_health` uses the Go SDK's metrics and thresholds. A metric the endpoint +doesn't report comes back as `None`, not `0`. + +## Error handling + +Every exception derives from `AuraError`: + +```text +AuraError +├── AuraConfigurationError (ValueError) bad client options +├── AuraValidationError (ValueError) bad arguments; nothing was sent +├── AuraConnectionError network failure after retries +│ └── AuraTimeoutError +├── AuraResponseError oversized or malformed response +├── MetricNotFoundError (LookupError) +└── AuraAPIError non-2xx response + ├── BadRequestError 400 + ├── AuthenticationError 401, or rejected credentials + ├── PermissionDeniedError 403 + ├── NotFoundError 404 + ├── ConflictError 409 + ├── RateLimitError 429 (.retry_after in seconds) + └── ServerError 5xx +``` + +```python +try: + client.instances.get("2f49c2b3") +except aura.NotFoundError: + print("no such instance") +except aura.AuraAPIError as err: + print(err.status_code, err.message, err.request_id) + for detail in err.details: + print(detail.reason, detail.field, detail.message) +``` + +`AuraAPIError` also provides the Go SDK's helpers: `is_not_found`, `is_unauthorized`, +`is_bad_request`, `has_multiple_errors` and `all_errors()`. + +## Logging + +The SDK logs through the standard `logging` module under the `aura_python_sdk` logger, and +emits nothing unless your application configures logging. Requests are logged at `DEBUG`, and +started mutations (create, delete, pause and so on) at `INFO`. Credentials, tokens and passwords +are never logged. + +```python +logging.basicConfig() +logging.getLogger("aura_python_sdk").setLevel(logging.DEBUG) +``` + +## Custom transports and testing + +Pass any object with `send(request) -> HttpResponse` and `close()` as `transport=`. This is the +equivalent of the Go SDK's `WithHTTPClient`. The SDK's retries, auth and error mapping still +apply on top. A client never closes a transport it didn't create. + +```python +from aura_python_sdk import AuraClient, HttpRequest, HttpResponse + + +class RecordingTransport: + def __init__(self, responses: list[HttpResponse]) -> None: + self.responses = responses + self.requests: list[HttpRequest] = [] + + def send(self, request: HttpRequest) -> HttpResponse: + self.requests.append(request) + return self.responses.pop(0) + + def close(self) -> None: + pass +``` + +For a network failure, a transport should raise `AuraConnectionError` or `AuraTimeoutError`. Set +`request_sent=False` only when the server certainly never received the request, because that +decides whether a `POST` is retried. + +## Coming from the Go SDK + +| Go | Python | +| --- | --- | +| `aura.NewClient(aura.WithCredentials(id, secret), aura.WithTimeout(t))` | `aura.AuraClient(client_id=id, client_secret=secret, timeout=t)` | +| `defer client.Close()` | `with aura.AuraClient(...) as client:` | +| `client.Instances.List(ctx)` returning `resp.Data` | `client.instances.list()` returns the list | +| `aura.IsNotFound(err)` | `except aura.NotFoundError:` | +| `aura.WithHTTPClient(c)` | `transport=` | +| `aura.WithInsecureBaseURL(u)` | `base_url=u, allow_insecure_base_url=True` | +| `client.Tenants.GetMetrics` | `client.tenants.get_metrics_integration` | +| `client.GraphAnalytics.Estimate` | `client.graph_analytics.estimate_size` | +| `SnapshotDate` / `aura.Today()` | `datetime.date` / omit it for today | + +Python additions: sizing and upgrade for instances; get, create and delete for customer-managed +keys; list filters; and the full set of `update` fields. The design notes are in +[PLAN.md](PLAN.md). + +## Examples + +[examples/](examples/) contains ports of the Go SDK's v1 examples. Each one reads +`AURA_CLIENT_ID` and `AURA_CLIENT_SECRET` from the environment: + +```sh +uv run python examples/list_instances.py +``` ## Development @@ -13,5 +340,21 @@ Requires Python 3.11+. uv sync --all-extras uv run ruff format && uv run ruff check uv run mypy -uv run pytest -m "not integration" +uv run pytest # unit and local black-box tests; no network ``` + +The live tests in `tests/integration/` call the real Aura API, and are skipped unless credentials +are set. They are read-only unless you opt in to creating and deleting an instance: + +```sh +AURA_CLIENT_ID=... AURA_CLIENT_SECRET=... uv run pytest -m integration +AURA_INTEGRATION_WRITE=1 AURA_TENANT_ID=... uv run pytest -m integration # also creates/deletes +``` + +To release, set `__version__` in `src/aura_python_sdk/_version.py`, add a matching +`## vX.Y.Z` section to [CHANGELOG.md](CHANGELOG.md), and push the tag `vX.Y.Z`. The release +workflow runs the tests, builds, publishes to PyPI and creates the GitHub release. + +## License + +MIT. See [LICENSE](LICENSE). diff --git a/examples/create_delete_instance.py b/examples/create_delete_instance.py new file mode 100644 index 0000000..6ceb2fe --- /dev/null +++ b/examples/create_delete_instance.py @@ -0,0 +1,87 @@ +"""Create a free instance, wait until it is running, then delete it. + +Usage: python examples/create_delete_instance.py +Needs AURA_CLIENT_ID, AURA_CLIENT_SECRET and AURA_TENANT_ID. + +Only one free instance can exist per tenant, so the script stops if one already exists. +""" + +import logging +import os +import sys +import time + +import aura_python_sdk as aura + +POLL_INTERVAL = 5.0 +CREATE_TIMEOUT = 10 * 60.0 + + +def wait_for_status( + client: aura.AuraClient, + instance_id: str, + status: aura.InstanceStatus, + timeout: float = CREATE_TIMEOUT, +) -> aura.Instance: + """Poll until the instance reaches ``status``, or raise TimeoutError.""" + deadline = time.monotonic() + timeout + poll = 1 + while True: + try: + instance = client.instances.get(instance_id) + except aura.NotFoundError: + # A brand-new instance can take a moment to appear in the API. + instance = None + if instance is not None and instance.status == status: + return instance + if time.monotonic() > deadline: + raise TimeoutError(f"instance {instance_id} did not reach {status} in {timeout:.0f}s") + current = instance.status if instance else "not visible yet" + print(f" status is {current} (poll {poll})") + poll += 1 + time.sleep(POLL_INTERVAL) + + +def main() -> int: + logging.basicConfig(level=logging.WARNING) + tenant_id = os.environ.get("AURA_TENANT_ID", "") + if not tenant_id: + print("AURA_TENANT_ID must be set", file=sys.stderr) + return 2 + + config = aura.InstanceConfig( + name="auraPythonSdkExample", + tenant_id=tenant_id, + cloud_provider=aura.CloudProvider.GCP, + region="europe-west1", + type=aura.InstanceType.FREE_DB, + version="5", + memory="1GB", + ) + + try: + with aura.AuraClient.from_env() as client: + for summary in client.instances.list(tenant_id): + if client.instances.get(summary.id).type == aura.InstanceType.FREE_DB: + print( + f"{summary.name} ({summary.id}) already uses the free tier", file=sys.stderr + ) + return 1 + + created = client.instances.create(config) + print(f"Created {created.name} ({created.id}) at {created.connection_url}") + print(f" username={created.username}; store the password now, it is shown only once") + + wait_for_status(client, created.id, aura.InstanceStatus.RUNNING) + print(f"Instance {created.id} is running; deleting it") + + deleted = client.instances.delete(created.id) + print(f"Instance {deleted.id} is {deleted.status}") + except (aura.AuraError, TimeoutError) as err: + print(f"error: {err}", file=sys.stderr) + return 1 + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/examples/get_instance_details.py b/examples/get_instance_details.py new file mode 100644 index 0000000..58ac7aa --- /dev/null +++ b/examples/get_instance_details.py @@ -0,0 +1,38 @@ +"""Show the details of one instance. + +Usage: python examples/get_instance_details.py INSTANCE_ID +Needs AURA_CLIENT_ID and AURA_CLIENT_SECRET. +""" + +import sys + +import aura_python_sdk as aura + + +def main() -> int: + if len(sys.argv) != 2: + print(__doc__, file=sys.stderr) + return 2 + try: + with aura.AuraClient.from_env() as client: + instance = client.instances.get(sys.argv[1]) + except aura.NotFoundError: + print(f"instance {sys.argv[1]} not found", file=sys.stderr) + return 1 + except aura.AuraError as err: + print(f"error: {err}", file=sys.stderr) + return 1 + + print(f"Name: {instance.name}") + print(f"Id: {instance.id}") + print(f"Status: {instance.status}") + print(f"Cloud provider: {instance.cloud_provider} ({instance.region})") + print(f"Tier: {instance.type}") + print(f"Memory: {instance.memory}") + print(f"Storage: {instance.storage or 'n/a'}") + print(f"Connection URL: {instance.connection_url}") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/examples/list_instances.py b/examples/list_instances.py new file mode 100644 index 0000000..aae5b0e --- /dev/null +++ b/examples/list_instances.py @@ -0,0 +1,29 @@ +"""List every instance the credentials can access. + +Usage: python examples/list_instances.py [TENANT_ID] +Needs AURA_CLIENT_ID and AURA_CLIENT_SECRET. +""" + +import sys + +import aura_python_sdk as aura + + +def main() -> int: + tenant_id = sys.argv[1] if len(sys.argv) > 1 else None + try: + with aura.AuraClient.from_env() as client: + instances = client.instances.list(tenant_id) + except aura.AuraError as err: + print(f"error: {err}", file=sys.stderr) + return 1 + + print(f"{len(instances)} instance(s)") + for instance in instances: + created = instance.created_at.isoformat() if instance.created_at else "unknown" + print(f"- {instance.name}: {instance.id} {instance.cloud_provider} created {created}") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/examples/list_snapshots.py b/examples/list_snapshots.py new file mode 100644 index 0000000..fd0a5a4 --- /dev/null +++ b/examples/list_snapshots.py @@ -0,0 +1,41 @@ +"""List an instance's snapshots for a day (default: today). + +Usage: python examples/list_snapshots.py INSTANCE_ID [YYYY-MM-DD] +Needs AURA_CLIENT_ID and AURA_CLIENT_SECRET. +""" + +import datetime +import sys + +import aura_python_sdk as aura + + +def main() -> int: + if len(sys.argv) not in (2, 3): + print(__doc__, file=sys.stderr) + return 2 + instance_id = sys.argv[1] + try: + day = datetime.date.fromisoformat(sys.argv[2]) if len(sys.argv) == 3 else None + except ValueError: + print("the date must be in the format YYYY-MM-DD", file=sys.stderr) + return 2 + + try: + with aura.AuraClient.from_env() as client: + snapshots = client.snapshots.list(instance_id, day) + except aura.AuraError as err: + print(f"error: {err}", file=sys.stderr) + return 1 + + for snapshot in snapshots: + taken = snapshot.timestamp.isoformat() if snapshot.timestamp else "unknown" + print( + f"- {snapshot.snapshot_id} {snapshot.status} {snapshot.profile} {taken} " + f"exportable={snapshot.exportable}" + ) + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/examples/list_tenants.py b/examples/list_tenants.py new file mode 100644 index 0000000..366acb0 --- /dev/null +++ b/examples/list_tenants.py @@ -0,0 +1,30 @@ +"""List every tenant and the instance configurations it supports. + +Usage: python examples/list_tenants.py +Needs AURA_CLIENT_ID and AURA_CLIENT_SECRET. +""" + +import sys + +import aura_python_sdk as aura + + +def main() -> int: + try: + with aura.AuraClient.from_env() as client: + for summary in client.tenants.list(): + tenant = client.tenants.get(summary.id) + print(f"{tenant.name} ({tenant.id})") + for config in tenant.instance_configurations: + print( + f" - {config.type} {config.cloud_provider} {config.region} " + f"memory={config.memory} storage={config.storage} version={config.version}" + ) + except aura.AuraError as err: + print(f"error: {err}", file=sys.stderr) + return 1 + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/examples/prometheus.py b/examples/prometheus.py new file mode 100644 index 0000000..25148d0 --- /dev/null +++ b/examples/prometheus.py @@ -0,0 +1,62 @@ +"""Read an instance's Prometheus metrics and print a health summary. + +Usage: python examples/prometheus.py INSTANCE_ID +Needs AURA_CLIENT_ID and AURA_CLIENT_SECRET, and metrics enabled for the instance in the Aura +Console. +""" + +import sys + +import aura_python_sdk as aura + + +def percent(value: float | None) -> str: + return "n/a" if value is None else f"{value:.1f}%" + + +def main() -> int: + if len(sys.argv) != 2: + print(__doc__, file=sys.stderr) + return 2 + instance_id = sys.argv[1] + + try: + with aura.AuraClient.from_env() as client: + url = client.instances.get(instance_id).metrics_integration_url + if not url: + print("metrics are not enabled for this instance", file=sys.stderr) + return 1 + print(f"Metrics URL: {url}") + + metrics = client.prometheus.fetch_raw_metrics(url) + names = sorted(metrics.metrics) + print(f"\nFetched {len(names)} metrics, e.g.:") + for name in names[:10]: + print(f" - {name}") + + try: + nodes = client.prometheus.get_metric_value(metrics, "neo4j_database_count_node") + print(f"\nNodes: {nodes:.0f}") + except aura.MetricNotFoundError: + print("\nNode count is not reported") + + health = client.prometheus.get_instance_health(instance_id, url) + except aura.AuraError as err: + print(f"error: {err}", file=sys.stderr) + return 1 + + print(f"\nOverall status: {health.overall_status} at {health.timestamp:%Y-%m-%d %H:%M:%S}") + print(f" CPU: {percent(health.resources.cpu_usage_percent)}") + print(f" Memory (heap): {percent(health.resources.memory_usage_percent)}") + print( + f" Connections: {health.connections.active_connections}/" + f"{health.connections.max_connections} ({percent(health.connections.usage_percent)})" + ) + print(f" Page cache hit: {percent(health.storage.page_cache_hit_rate)}") + for issue, recommendation in zip(health.issues, health.recommendations, strict=True): + print(f" ! {issue}: {recommendation}") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/examples/restore_from_snapshot.py b/examples/restore_from_snapshot.py new file mode 100644 index 0000000..0b425ef --- /dev/null +++ b/examples/restore_from_snapshot.py @@ -0,0 +1,37 @@ +"""Restore an instance from one of its snapshots, replacing its current data. + +Usage: python examples/restore_from_snapshot.py INSTANCE_ID SNAPSHOT_ID +Needs AURA_CLIENT_ID and AURA_CLIENT_SECRET. + +Run examples/list_snapshots.py first to find a snapshot ID. +""" + +import sys + +import aura_python_sdk as aura + + +def main() -> int: + if len(sys.argv) != 3: + print(__doc__, file=sys.stderr) + return 2 + instance_id, snapshot_id = sys.argv[1], sys.argv[2] + + answer = input(f"This replaces all data in {instance_id}. Type the instance ID to confirm: ") + if answer.strip() != instance_id: + print("not confirmed; nothing was changed") + return 1 + + try: + with aura.AuraClient.from_env() as client: + instance = client.snapshots.restore(instance_id, snapshot_id) + except aura.AuraError as err: + print(f"error: {err}", file=sys.stderr) + return 1 + + print(f"Restore started: {instance.id} is {instance.status}") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/examples/take_snapshot.py b/examples/take_snapshot.py new file mode 100644 index 0000000..5a96c06 --- /dev/null +++ b/examples/take_snapshot.py @@ -0,0 +1,32 @@ +"""Take an on-demand snapshot of an instance and show its details. + +Usage: python examples/take_snapshot.py INSTANCE_ID +Needs AURA_CLIENT_ID and AURA_CLIENT_SECRET. +""" + +import sys + +import aura_python_sdk as aura + + +def main() -> int: + if len(sys.argv) != 2: + print(__doc__, file=sys.stderr) + return 2 + instance_id = sys.argv[1] + try: + with aura.AuraClient.from_env() as client: + started = client.snapshots.create(instance_id) + print(f"Snapshot started: {started.snapshot_id}") + snapshot = client.snapshots.get(instance_id, started.snapshot_id) + except aura.AuraError as err: + print(f"error: {err}", file=sys.stderr) + return 1 + + print(f" instance: {snapshot.instance_id}") + print(f" status: {snapshot.status}") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/pyproject.toml b/pyproject.toml index 51c2d74..e89ae48 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -61,11 +61,11 @@ ignore = [] [tool.mypy] python_version = "3.11" strict = true -files = ["src", "tests"] +files = ["src", "tests", "examples"] [tool.pytest.ini_options] testpaths = ["tests"] -addopts = ["--strict-markers", "--import-mode=importlib"] +addopts = ["--strict-markers", "--import-mode=importlib", "-m", "not integration"] markers = ["integration: talks to a real Aura account; needs AURA_CLIENT_ID / AURA_CLIENT_SECRET"] [tool.coverage.run] diff --git a/tests/blackbox/__init__.py b/tests/blackbox/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/blackbox/conftest.py b/tests/blackbox/conftest.py new file mode 100644 index 0000000..0cf56cc --- /dev/null +++ b/tests/blackbox/conftest.py @@ -0,0 +1,140 @@ +"""A local Aura API stand-in (Go: client_blackbox_test.go's httptest.Server). + +These tests use only the public API and the real httpx transport, over real sockets on +127.0.0.1. Nothing reaches the internet. +""" + +from __future__ import annotations + +import json +import threading +import time +from collections.abc import Callable, Iterator +from dataclasses import dataclass, field +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from typing import Any +from urllib.parse import urlsplit + +import pytest + +import aura_python_sdk as aura + + +@dataclass +class Reply: + status: int = 200 + body: bytes = b"" + headers: dict[str, str] = field(default_factory=dict) + delay: float = 0.0 + + @classmethod + def json(cls, status: int, payload: object, **headers: str) -> Reply: + return cls( + status, json.dumps(payload).encode(), {"Content-Type": "application/json", **headers} + ) + + +@dataclass +class Received: + method: str + path: str + query: str + headers: dict[str, str] + body: bytes + + def json(self) -> Any: + return json.loads(self.body) + + +Route = Callable[[Received], Reply] | Reply + + +class FakeAura: + def __init__(self) -> None: + self.routes: dict[tuple[str, str], Route] = { + ("POST", "/oauth/token"): Reply.json( + 200, {"token_type": "Bearer", "access_token": "local-token", "expires_in": 3600} + ) + } + self.received: list[Received] = [] + self._server = ThreadingHTTPServer(("127.0.0.1", 0), self._handler()) + # Don't wait for slow handlers (e.g. the timeout test) when shutting down. + self._server.daemon_threads = True + self._server.block_on_close = False + self._thread = threading.Thread( + target=self._server.serve_forever, kwargs={"poll_interval": 0.01}, daemon=True + ) + + @property + def url(self) -> str: + host, port = self._server.server_address[:2] + return f"http://{host!s}:{port}" + + def route(self, method: str, path: str, reply: Route) -> None: + self.routes[(method, path)] = reply + + def api_requests(self) -> list[Received]: + return [r for r in self.received if r.path != "/oauth/token"] + + def client(self, **options: Any) -> aura.AuraClient: + options.setdefault("timeout", 5) + return aura.AuraClient( + client_id="local-id", + client_secret="local-secret", + base_url=self.url, + allow_insecure_base_url=True, + **options, + ) + + def _handler(self) -> type[BaseHTTPRequestHandler]: + fake = self + + class Handler(BaseHTTPRequestHandler): + def _serve(self) -> None: + length = int(self.headers.get("Content-Length") or 0) + parts = urlsplit(self.path) + received = Received( + method=self.command, + path=parts.path, + query=parts.query, + headers={k.lower(): v for k, v in self.headers.items()}, + body=self.rfile.read(length) if length else b"", + ) + fake.received.append(received) + route = fake.routes.get((self.command, parts.path)) + reply = ( + Reply.json(404, {"errors": [{"message": "no route"}]}) + if route is None + else (route(received) if callable(route) else route) + ) + if reply.delay: + time.sleep(reply.delay) + self.send_response(reply.status) + for name, value in reply.headers.items(): + self.send_header(name, value) + self.send_header("Content-Length", str(len(reply.body))) + self.end_headers() + if reply.body: + self.wfile.write(reply.body) + + do_GET = do_POST = do_PATCH = do_PUT = do_DELETE = _serve # noqa: N815 - stdlib names + + def log_message(self, format: str, *args: Any) -> None: + pass + + return Handler + + def start(self) -> None: + self._thread.start() + + def stop(self) -> None: + self._server.shutdown() + self._server.server_close() + + +@pytest.fixture +def fake_aura() -> Iterator[FakeAura]: + server = FakeAura() + server.start() + yield server + server.stop() diff --git a/tests/blackbox/test_blackbox.py b/tests/blackbox/test_blackbox.py new file mode 100644 index 0000000..9cde6ad --- /dev/null +++ b/tests/blackbox/test_blackbox.py @@ -0,0 +1,194 @@ +"""End-to-end over real sockets and the real httpx transport (Go: client_blackbox_test.go).""" + +import base64 +import socket +import time + +import pytest + +import aura_python_sdk as aura +from tests.blackbox.conftest import FakeAura, Received, Reply + +TENANT_ID = "6981ace7-efe8-4f5c-b7c5-267b5162ce91" +INSTANCE = { + "id": "2f49c2b3", + "name": "Production", + "status": "running", + "tenant_id": TENANT_ID, + "cloud_provider": "gcp", + "connection_url": "neo4j+s://2f49c2b3.databases.neo4j.io", + "region": "europe-west1", + "type": "enterprise-db", + "memory": "8GB", +} + + +def test_list_instances_end_to_end(fake_aura: FakeAura) -> None: + fake_aura.route( + "GET", + "/v1/instances", + Reply.json( + 200, + { + "data": [ + {"id": "2f49c2b3", "name": "P", "tenant_id": TENANT_ID, "cloud_provider": "gcp"} + ] + }, + ), + ) + with fake_aura.client(default_headers={"X-Team": "db"}) as client: + [instance] = client.instances.list(TENANT_ID) + + assert instance.id == "2f49c2b3" + token_request, api_request = fake_aura.received + assert ( + token_request.headers["authorization"] + == "Basic " + base64.b64encode(b"local-id:local-secret").decode() + ) + assert token_request.body == b"grant_type=client_credentials" + assert api_request.query == f"tenantId={TENANT_ID}" + assert api_request.headers["authorization"] == "Bearer local-token" + assert api_request.headers["user-agent"] == f"aura-python-sdk/{aura.__version__}" + assert api_request.headers["content-type"] == "application/json" + assert api_request.headers["x-team"] == "db" + + +def test_token_is_reused_across_calls(fake_aura: FakeAura) -> None: + fake_aura.route("GET", "/v1/instances/2f49c2b3", Reply.json(200, {"data": INSTANCE})) + with fake_aura.client() as client: + for _ in range(3): + client.instances.get("2f49c2b3") + assert [r.path for r in fake_aura.received].count("/oauth/token") == 1 + + +def test_create_sends_json_body(fake_aura: FakeAura) -> None: + created = { + **INSTANCE, + "id": "db1d1234", + "username": "neo4j", + "password": "secret-pw", + } + fake_aura.route("POST", "/v1/instances", Reply.json(202, {"data": created})) + config = aura.InstanceConfig( + name="Instance01", + tenant_id=TENANT_ID, + cloud_provider=aura.CloudProvider.GCP, + region="europe-west1", + type=aura.InstanceType.ENTERPRISE_DB, + version="5", + memory="8GB", + ) + with fake_aura.client() as client: + result = client.instances.create(config) + assert result.password == "secret-pw" + assert fake_aura.api_requests()[0].json()["type"] == "enterprise-db" + + +def test_api_error_is_mapped(fake_aura: FakeAura) -> None: + fake_aura.route( + "GET", + "/v1/instances/2f49c2b3", + Reply.json( + 404, + {"errors": [{"message": "Instance not found", "reason": "instance-not-found"}]}, + **{"X-Request-Id": "req-42"}, + ), + ) + with fake_aura.client() as client, pytest.raises(aura.NotFoundError) as info: + client.instances.get("2f49c2b3") + assert info.value.request_id == "req-42" + assert info.value.details[0].reason == "instance-not-found" + + +def test_rejected_credentials(fake_aura: FakeAura) -> None: + fake_aura.route("POST", "/oauth/token", Reply.json(401, {"error": "access_denied"})) + with ( + fake_aura.client() as client, + pytest.raises(aura.AuthenticationError, match="access_denied"), + ): + client.tenants.list() + assert fake_aura.api_requests() == [] + + +def test_rate_limit_is_not_retried(fake_aura: FakeAura) -> None: + fake_aura.route( + "GET", + "/v1/tenants", + Reply.json(429, {"error": "Rate limit exceeded"}, **{"Retry-After": "7"}), + ) + with fake_aura.client() as client, pytest.raises(aura.RateLimitError) as info: + client.tenants.list() + assert info.value.retry_after == 7.0 + assert len(fake_aura.api_requests()) == 1 + + +def test_permanent_redirect_is_followed(fake_aura: FakeAura) -> None: + fake_aura.route( + "GET", "/v1/tenants", Reply(308, headers={"Location": f"{fake_aura.url}/v1/tenants-moved"}) + ) + fake_aura.route("GET", "/v1/tenants-moved", Reply.json(200, {"data": []})) + with fake_aura.client() as client: + assert client.tenants.list() == [] + + +def test_response_size_limit(fake_aura: FakeAura) -> None: + fake_aura.route("GET", "/v1/tenants", Reply.json(200, {"data": [], "padding": "x" * 5000})) + with ( + fake_aura.client(max_response_size=1024) as client, + pytest.raises(aura.AuraResponseError, match="exceeded limit"), + ): + client.tenants.list() + + +def test_delete_with_no_content(fake_aura: FakeAura) -> None: + fake_aura.route("DELETE", "/v1/customer-managed-keys/key-1", Reply(204)) + with fake_aura.client() as client: + client.cmek.delete("key-1") + assert fake_aura.api_requests()[0].method == "DELETE" + + +def test_slow_post_times_out_and_is_not_retried(fake_aura: FakeAura) -> None: + fake_aura.route("POST", "/v1/instances/2f49c2b3/pause", Reply(202, b"{}", delay=2.0)) + started = time.monotonic() + with ( + fake_aura.client(timeout=0.5, max_retries=3) as client, + pytest.raises(aura.AuraTimeoutError), + ): + client.instances.pause("2f49c2b3") + assert time.monotonic() - started < 1.9 + assert len(fake_aura.api_requests()) == 1 + + +def test_connection_refused(fake_aura: FakeAura) -> None: + with socket.socket() as sock: + sock.bind(("127.0.0.1", 0)) + closed_port = sock.getsockname()[1] + client = aura.AuraClient( + client_id="id", + client_secret="secret", + base_url=f"http://127.0.0.1:{closed_port}", + allow_insecure_base_url=True, + max_retries=0, + timeout=5, + ) + with client, pytest.raises(aura.AuraConnectionError) as info: + client.tenants.list() + assert info.value.request_sent is False + + +def test_prometheus_over_http(fake_aura: FakeAura) -> None: + text = b"# TYPE neo4j_aura_cpu_usage gauge\nneo4j_aura_cpu_usage 0.5\n" + fake_aura.route("GET", "/metrics", Reply(200, text, {"Content-Type": "text/plain"})) + with fake_aura.client() as client: + metrics = client.prometheus.fetch_raw_metrics(f"{fake_aura.url}/metrics") + assert client.prometheus.get_metric_value(metrics, "neo4j_aura_cpu_usage") == 0.5 + assert fake_aura.api_requests()[0].headers["authorization"] == "Bearer local-token" + + +def test_dynamic_route(fake_aura: FakeAura) -> None: + def echo_patch(request: Received) -> Reply: + return Reply.json(200, {"data": {**INSTANCE, **request.json()}}) + + fake_aura.route("PATCH", "/v1/instances/2f49c2b3", echo_patch) + with fake_aura.client() as client: + assert client.instances.update("2f49c2b3", name="Renamed").name == "Renamed" diff --git a/tests/integration/__init__.py b/tests/integration/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/integration/test_live.py b/tests/integration/test_live.py new file mode 100644 index 0000000..ab0d128 --- /dev/null +++ b/tests/integration/test_live.py @@ -0,0 +1,123 @@ +"""Live tests against the real Aura API. + +Skipped unless AURA_CLIENT_ID and AURA_CLIENT_SECRET are set. Run them with: + + uv run pytest -m integration + +The tests only read unless AURA_INTEGRATION_WRITE=1 and AURA_TENANT_ID are also set. Then +test_create_pause_resume_delete creates a free instance, pauses and resumes it, and deletes it. +""" + +from __future__ import annotations + +import os +import time +from collections.abc import Iterator + +import pytest + +import aura_python_sdk as aura + +pytestmark = [ + pytest.mark.integration, + pytest.mark.skipif( + not (os.environ.get("AURA_CLIENT_ID") and os.environ.get("AURA_CLIENT_SECRET")), + reason="AURA_CLIENT_ID and AURA_CLIENT_SECRET are not set", + ), +] + +WRITES_ENABLED = os.environ.get("AURA_INTEGRATION_WRITE") == "1" +TENANT_ID = os.environ.get("AURA_TENANT_ID", "") + + +@pytest.fixture(scope="module") +def client() -> Iterator[aura.AuraClient]: + with aura.AuraClient.from_env(timeout=60) as live_client: + yield live_client + + +def test_tenants(client: aura.AuraClient) -> None: + tenants = client.tenants.list() + assert tenants, "the credentials should see at least one tenant" + tenant = client.tenants.get(tenants[0].id) + assert tenant.id == tenants[0].id + + +def test_instances_list_and_get(client: aura.AuraClient) -> None: + instances = client.instances.list() + for summary in instances[:3]: + instance = client.instances.get(summary.id) + assert instance.id == summary.id + assert instance.tenant_id == summary.tenant_id + + +def test_instances_list_filtered_by_tenant(client: aura.AuraClient) -> None: + tenant_id = client.tenants.list()[0].id + assert all(i.tenant_id == tenant_id for i in client.instances.list(tenant_id)) + + +def test_snapshots_for_first_instance(client: aura.AuraClient) -> None: + instances = client.instances.list() + if not instances: + pytest.skip("no instances to list snapshots for") + for snapshot in client.snapshots.list(instances[0].id): + assert snapshot.instance_id == instances[0].id + + +def test_cmek_and_sessions_list(client: aura.AuraClient) -> None: + assert isinstance(client.cmek.list(), list) + assert isinstance(client.graph_analytics.list(), list) + + +def test_unknown_instance_is_not_found(client: aura.AuraClient) -> None: + with pytest.raises(aura.NotFoundError): + client.instances.get("00000000") + + +def test_bad_credentials_are_rejected() -> None: + with ( + aura.AuraClient(client_id="not-real", client_secret="not-real") as bad, + pytest.raises(aura.AuthenticationError), + ): + bad.tenants.list() + + +def _wait_for( + client: aura.AuraClient, instance_id: str, status: aura.InstanceStatus, timeout: float = 900 +) -> aura.Instance: + deadline = time.monotonic() + timeout + while True: + try: + instance = client.instances.get(instance_id) + if instance.status == status: + return instance + except aura.NotFoundError: + pass # a new instance can take a moment to appear + if time.monotonic() > deadline: + pytest.fail(f"instance {instance_id} did not reach {status} within {timeout:.0f}s") + time.sleep(10) + + +@pytest.mark.skipif( + not (WRITES_ENABLED and TENANT_ID), reason="needs AURA_INTEGRATION_WRITE=1 and AURA_TENANT_ID" +) +def test_create_pause_resume_delete(client: aura.AuraClient) -> None: + created = client.instances.create( + aura.InstanceConfig( + name="aura-python-sdk-it", + tenant_id=TENANT_ID, + cloud_provider=aura.CloudProvider.GCP, + region="europe-west1", + type=aura.InstanceType.FREE_DB, + version="5", + memory="1GB", + ) + ) + try: + _wait_for(client, created.id, aura.InstanceStatus.RUNNING) + client.instances.pause(created.id) + _wait_for(client, created.id, aura.InstanceStatus.PAUSED) + client.instances.resume(created.id) + _wait_for(client, created.id, aura.InstanceStatus.RUNNING) + finally: + client.instances.delete(created.id) From 17134112451adab6e35f9c2c78fcbd015d641ed0 Mon Sep 17 00:00:00 2001 From: Jonathan Giffard Date: Tue, 29 Sep 2026 20:56:20 +0100 Subject: [PATCH 7/8] Add AsyncAuraClient with shared operations and parity tests Phase 8 of PLAN.md. Service operations are now pure Call descriptions run by sync and async executors; retry policy, token handling and request building are shared. Adds AsyncAuraClient, AsyncHttpTransport and an httpx async transport, plus tests that every async method matches its sync twin. Co-Authored-By: Claude Opus 5.5 --- CHANGELOG.md | 2 + PLAN.md | 21 +- README.md | 33 +- examples/async_instance_details.py | 29 ++ src/aura_python_sdk/__init__.py | 8 +- src/aura_python_sdk/_client.py | 177 ++++++- src/aura_python_sdk/_internal/_auth.py | 188 +++++-- src/aura_python_sdk/_internal/_call.py | 45 ++ src/aura_python_sdk/_internal/_request.py | 166 +++++-- src/aura_python_sdk/_internal/http/_httpx.py | 88 ++-- .../_internal/http/_service.py | 176 +++++-- src/aura_python_sdk/_transport.py | 12 + src/aura_python_sdk/services/__init__.py | 21 +- src/aura_python_sdk/services/_base.py | 36 +- src/aura_python_sdk/services/cmek.py | 151 ++++-- .../services/graph_analytics.py | 223 ++++++--- src/aura_python_sdk/services/instances.py | 469 ++++++++++++------ src/aura_python_sdk/services/prometheus.py | 253 ++++++---- src/aura_python_sdk/services/snapshots.py | 125 +++-- src/aura_python_sdk/services/tenants.py | 65 ++- tests/blackbox/conftest.py | 10 + tests/blackbox/test_blackbox_async.py | 58 +++ tests/conftest.py | 7 + tests/fakes.py | 20 + tests/integration/test_live.py | 7 + tests/transport/test_httpx_transport.py | 70 ++- tests/unit/test_async_core.py | 197 ++++++++ tests/unit/test_async_parity.py | 293 +++++++++++ 28 files changed, 2367 insertions(+), 583 deletions(-) create mode 100644 examples/async_instance_details.py create mode 100644 src/aura_python_sdk/_internal/_call.py create mode 100644 tests/blackbox/test_blackbox_async.py create mode 100644 tests/conftest.py create mode 100644 tests/unit/test_async_core.py create mode 100644 tests/unit/test_async_parity.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 648a600..8c0be9f 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -10,6 +10,8 @@ the `## vX.Y.Z` section that matches the pushed tag as the GitHub release notes. ### Added +- `AsyncAuraClient` for asyncio. It has the same options and services as `AuraClient`, shares + their validation and parsing, and uses an `AsyncHttpTransport` (httpx by default). - `AuraClient` for the Aura API v1, with the Go SDK's options as keyword arguments, `from_env()`, and context-manager support. - Services matching the Go SDK: `tenants`, `instances`, `snapshots`, `cmek`, `graph_analytics` and diff --git a/PLAN.md b/PLAN.md index 080fd8a..a2908c4 100644 --- a/PLAN.md +++ b/PLAN.md @@ -227,6 +227,25 @@ parser accepts the spec's `{"errors": [...]}`, the middleware `{"error": "..."}` didn't report a metric, and threshold checks skip them. The status logic and messages match Go. - **`get_metric_value`** raises `MetricNotFoundError`, which is also a `LookupError`. +### 2.10 Decisions made in phase 8 (async) + +- **Written once, run two ways.** Each service operation is a pure function that validates its + arguments and returns a `Call` (method, path, params, body, parser, log text). `Service._run` + sends it synchronously and `AsyncService._run` awaits it. The retry policy, token parsing, + header building and error mapping are shared the same way, and only the I/O loops are + duplicated. +- **Thin async classes.** `AsyncInstanceService` and the other async services repeat only the + signatures, and their docstrings point to the sync methods. +- **Parity is enforced.** `tests/unit/test_async_parity.py` runs every method on both clients + against the same responses and asserts identical requests and results. It also checks the + signatures match, and fails if a method has no case. Two deliberately broken methods were + caught. +- **Transports can't be mixed up.** `AuraClient` rejects a transport whose `send` is a coroutine, + and `AsyncAuraClient` requires one. mypy catches the same mistake statically. +- **`asyncio.Lock`** guards the token refresh, so concurrent tasks share one token fetch. +- **Test tooling:** async tests use anyio's pytest plugin, which is already installed with httpx. + No new dependency. + ## 3. Package layout ``` @@ -268,7 +287,7 @@ Dev tooling: `uv`, `ruff` (lint and format), `mypy --strict`, `pytest`, `pytest- ## 5. Phases -**Status:** phases 1–7 are done. Phase 8 (async) is not started. +**Status:** all eight phases are done. 1. **Scaffold**: pyproject, uv, ruff, mypy, pytest config, CI workflow, and the import-boundary test. 2. **Core**: config/options, errors, `HttpTransport` + `HttpxTransport` (retries, size cap), diff --git a/README.md b/README.md index 4993338..55e495b 100644 --- a/README.md +++ b/README.md @@ -4,6 +4,7 @@ A Python client for the [Neo4j Aura API](https://neo4j.com/docs/aura/api/overvie example, `client.instances.list()` returns your Aura instances. It is modelled on [aura-go-sdk](https://github.com/neo4j-contrib/aura-go-sdk) and covers the whole v1 API. +- Sync (`AuraClient`) and asyncio (`AsyncAuraClient`) clients with the same services. - Typed throughout (`py.typed`, checked with `mypy --strict`), using frozen dataclass models. - One runtime dependency, [httpx](https://www.python-httpx.org/), kept behind the SDK's own transport interface. @@ -19,6 +20,7 @@ You need an Aura API client ID and secret. See - [Quick start](#quick-start) - [Configuration](#configuration) - [Timeouts and retries](#timeouts-and-retries) +- [Async](#async) - [Tenants](#tenants) - [Instances](#instances) - [Snapshots](#snapshots) @@ -102,6 +104,32 @@ status, including 429 and 5xx, is never retried. If a request might already have server (a read timeout or a dropped connection), only idempotent methods (`GET`, `PUT`, `DELETE`) are retried. That means a `create` or `pause` is never sent twice. +## Async + +`AsyncAuraClient` takes the same options, and its services have the same methods, which you +await. Concurrent calls share one OAuth token. + +```python +import asyncio + +import aura_python_sdk as aura + + +async def main() -> None: + async with aura.AsyncAuraClient.from_env() as client: + summaries = await client.instances.list() + instances = await asyncio.gather(*(client.instances.get(s.id) for s in summaries)) + for instance in instances: + print(instance.name, instance.status) + + +asyncio.run(main()) +``` + +Use `async with` or `await client.aclose()` to release connections. `prometheus.get_metric_value` +does no I/O, so it is a plain method on both clients. A custom transport for the async client +implements `AsyncHttpTransport` (`async send()` and `async aclose()`). + ## Tenants ```python @@ -282,7 +310,9 @@ logging.getLogger("aura_python_sdk").setLevel(logging.DEBUG) ## Custom transports and testing -Pass any object with `send(request) -> HttpResponse` and `close()` as `transport=`. This is the +Pass any object with `send(request) -> HttpResponse` and `close()` as `transport=`. For +`AsyncAuraClient`, pass one with `async send()` and `async aclose()`. Each client rejects the +other kind. This is the equivalent of the Go SDK's `WithHTTPClient`. The SDK's retries, auth and error mapping still apply on top. A client never closes a transport it didn't create. @@ -313,6 +343,7 @@ decides whether a `POST` is retried. | --- | --- | | `aura.NewClient(aura.WithCredentials(id, secret), aura.WithTimeout(t))` | `aura.AuraClient(client_id=id, client_secret=secret, timeout=t)` | | `defer client.Close()` | `with aura.AuraClient(...) as client:` | +| goroutines with a shared client | `AsyncAuraClient` with `asyncio.gather` | | `client.Instances.List(ctx)` returning `resp.Data` | `client.instances.list()` returns the list | | `aura.IsNotFound(err)` | `except aura.NotFoundError:` | | `aura.WithHTTPClient(c)` | `transport=` | diff --git a/examples/async_instance_details.py b/examples/async_instance_details.py new file mode 100644 index 0000000..ff3377d --- /dev/null +++ b/examples/async_instance_details.py @@ -0,0 +1,29 @@ +"""Fetch every instance's details concurrently with AsyncAuraClient. + +Usage: python examples/async_instance_details.py +Needs AURA_CLIENT_ID and AURA_CLIENT_SECRET. +""" + +import asyncio +import sys + +import aura_python_sdk as aura + + +async def main() -> int: + try: + async with aura.AsyncAuraClient.from_env() as client: + summaries = await client.instances.list() + # All the GET requests run concurrently and share one OAuth token. + instances = await asyncio.gather(*(client.instances.get(s.id) for s in summaries)) + except aura.AuraError as err: + print(f"error: {err}", file=sys.stderr) + return 1 + + for instance in instances: + print(f"- {instance.name} ({instance.id}): {instance.status}, {instance.type}") + return 0 + + +if __name__ == "__main__": + sys.exit(asyncio.run(main())) diff --git a/src/aura_python_sdk/__init__.py b/src/aura_python_sdk/__init__.py index 11b63ff..50c26f9 100644 --- a/src/aura_python_sdk/__init__.py +++ b/src/aura_python_sdk/__init__.py @@ -7,11 +7,13 @@ with aura.AuraClient(client_id="...", client_secret="...") as client: for instance in client.instances.list(): print(instance.id, instance.name) + +For asyncio, use :class:`AsyncAuraClient`, which has the same services with awaitable methods. """ import logging -from aura_python_sdk._client import AuraClient +from aura_python_sdk._client import AsyncAuraClient, AuraClient from aura_python_sdk._errors import ( AuraAPIError, AuraConfigurationError, @@ -30,7 +32,7 @@ RateLimitError, ServerError, ) -from aura_python_sdk._transport import HttpRequest, HttpResponse, HttpTransport +from aura_python_sdk._transport import AsyncHttpTransport, HttpRequest, HttpResponse, HttpTransport from aura_python_sdk._version import __version__ from aura_python_sdk.models import ( CDCEnrichmentMode, @@ -71,6 +73,8 @@ logging.getLogger(__name__).addHandler(logging.NullHandler()) __all__ = [ + "AsyncAuraClient", + "AsyncHttpTransport", "AuraAPIError", "AuraClient", "AuraConfigurationError", diff --git a/src/aura_python_sdk/_client.py b/src/aura_python_sdk/_client.py index 744e9b4..15a67af 100644 --- a/src/aura_python_sdk/_client.py +++ b/src/aura_python_sdk/_client.py @@ -1,7 +1,8 @@ -"""The AuraClient entry point (Go: client.go).""" +"""The AuraClient and AsyncAuraClient entry points (Go: client.go).""" from __future__ import annotations +import inspect import logging import os from collections.abc import Mapping @@ -19,12 +20,18 @@ build_config, ) from aura_python_sdk._errors import AuraConfigurationError -from aura_python_sdk._internal._auth import TokenManager -from aura_python_sdk._internal._request import RequestService -from aura_python_sdk._internal.http._httpx import HttpxTransport -from aura_python_sdk._internal.http._service import HttpService -from aura_python_sdk._transport import HttpTransport +from aura_python_sdk._internal._auth import AsyncTokenManager, TokenManager +from aura_python_sdk._internal._request import AsyncRequestService, RequestService +from aura_python_sdk._internal.http._httpx import AsyncHttpxTransport, HttpxTransport +from aura_python_sdk._internal.http._service import AsyncHttpService, HttpService +from aura_python_sdk._transport import AsyncHttpTransport, HttpTransport from aura_python_sdk.services import ( + AsyncCMEKService, + AsyncGDSSessionService, + AsyncInstanceService, + AsyncPrometheusService, + AsyncSnapshotService, + AsyncTenantService, CMEKService, GDSSessionService, InstanceService, @@ -39,6 +46,20 @@ _LOGGER_NAME = "aura_python_sdk" +def _resolve_logger(logger: logging.Logger | None) -> logging.Logger: + if logger is not None and not isinstance(logger, logging.Logger): + raise AuraConfigurationError("logger must be a logging.Logger") + return logger or logging.getLogger(_LOGGER_NAME) + + +def _env_credentials() -> tuple[str, str]: + client_id = os.environ.get(ENV_CLIENT_ID, "") + client_secret = os.environ.get(ENV_CLIENT_SECRET, "") + if not client_id or not client_secret: + raise AuraConfigurationError(f"{ENV_CLIENT_ID} and {ENV_CLIENT_SECRET} must both be set") + return client_id, client_secret + + class AuraClient: """Client for the Neo4j Aura API v1. @@ -98,12 +119,14 @@ def __init__( user_agent=user_agent, default_headers=default_headers, ) - if transport is not None and not isinstance(transport, HttpTransport): - raise AuraConfigurationError("transport must implement send() and close()") - if logger is not None and not isinstance(logger, logging.Logger): - raise AuraConfigurationError("logger must be a logging.Logger") - - self._logger = logger or logging.getLogger(_LOGGER_NAME) + if transport is not None and ( + not isinstance(transport, HttpTransport) or inspect.iscoroutinefunction(transport.send) + ): + raise AuraConfigurationError( + "transport must implement send() and close(); use AsyncAuraClient for an " + "async transport" + ) + self._logger = _resolve_logger(logger) self._owns_transport = transport is None self._transport: HttpTransport = transport or HttpxTransport() self._closed = False @@ -157,12 +180,7 @@ def from_env(cls, **options: object) -> Self: Any other keyword option is passed through to :class:`AuraClient`. """ - client_id = os.environ.get(ENV_CLIENT_ID, "") - client_secret = os.environ.get(ENV_CLIENT_SECRET, "") - if not client_id or not client_secret: - raise AuraConfigurationError( - f"{ENV_CLIENT_ID} and {ENV_CLIENT_SECRET} must both be set" - ) + client_id, client_secret = _env_credentials() return cls(client_id=client_id, client_secret=client_secret, **options) # type: ignore[arg-type] @property @@ -190,3 +208,126 @@ def __exit__( def __repr__(self) -> str: return f"AuraClient(base_url={self._config.base_url!r})" + + +class AsyncAuraClient: + """Async client for the Neo4j Aura API v1, for use with ``asyncio``. + + Takes the same options as :class:`AuraClient`, and its services have the same methods, + which are awaited:: + + async with AsyncAuraClient(client_id="...", client_secret="...") as client: + instances = await client.instances.list() + + ``transport`` must be an :class:`AsyncHttpTransport`. Call :meth:`aclose`, or use + ``async with``, to release connections. + """ + + def __init__( + self, + *, + client_id: str, + client_secret: str, + base_url: str = DEFAULT_BASE_URL, + allow_insecure_base_url: bool = False, + timeout: float = DEFAULT_TIMEOUT, + max_retries: int = DEFAULT_MAX_RETRIES, + max_response_size: int = DEFAULT_MAX_RESPONSE_SIZE, + user_agent: str = DEFAULT_USER_AGENT, + default_headers: Mapping[str, str] | None = None, + logger: logging.Logger | None = None, + transport: AsyncHttpTransport | None = None, + ) -> None: + self._config: ClientConfig = build_config( + client_id=client_id, + client_secret=client_secret, + base_url=base_url, + allow_insecure_base_url=allow_insecure_base_url, + timeout=timeout, + max_retries=max_retries, + max_response_size=max_response_size, + user_agent=user_agent, + default_headers=default_headers, + ) + if transport is not None and ( + not isinstance(transport, AsyncHttpTransport) + or not inspect.iscoroutinefunction(transport.send) + ): + raise AuraConfigurationError( + "transport must implement async send() and aclose(); use AuraClient for a " + "sync transport" + ) + self._logger = _resolve_logger(logger) + self._owns_transport = transport is None + self._transport: AsyncHttpTransport = transport or AsyncHttpxTransport() + self._closed = False + + http = AsyncHttpService( + self._transport, + max_retries=self._config.max_retries, + max_response_size=self._config.max_response_size, + logger=self._logger.getChild("http"), + ) + auth = AsyncTokenManager( + client_id=self._config.client_id, + client_secret=self._config.client_secret, + token_url=f"{self._config.base_url}/oauth/token", + user_agent=self._config.user_agent, + http=http, + logger=self._logger.getChild("auth"), + ) + self._api = AsyncRequestService( + http=http, + auth=auth, + base_url=self._config.base_url, + api_version=API_VERSION, + user_agent=self._config.user_agent, + default_headers=self._config.default_headers, + timeout=self._config.timeout, + logger=self._logger.getChild("api"), + ) + + self.tenants = AsyncTenantService(self._api, self._logger.getChild("tenants")) + self.instances = AsyncInstanceService(self._api, self._logger.getChild("instances")) + self.snapshots = AsyncSnapshotService(self._api, self._logger.getChild("snapshots")) + self.cmek = AsyncCMEKService(self._api, self._logger.getChild("cmek")) + self.graph_analytics = AsyncGDSSessionService( + self._api, self._logger.getChild("graph_analytics") + ) + self.prometheus = AsyncPrometheusService( + self._api, + self._logger.getChild("prometheus"), + allow_untrusted_urls=self._config.allow_insecure_base_url, + ) + + @classmethod + def from_env(cls, **options: object) -> Self: + """Build a client with credentials from ``AURA_CLIENT_ID`` and ``AURA_CLIENT_SECRET``.""" + client_id, client_secret = _env_credentials() + return cls(client_id=client_id, client_secret=client_secret, **options) # type: ignore[arg-type] + + @property + def base_url(self) -> str: + return self._config.base_url + + async def aclose(self) -> None: + """Release pooled connections. Safe to call more than once.""" + if self._closed: + return + self._closed = True + if self._owns_transport: + await self._transport.aclose() + + async def __aenter__(self) -> Self: + return self + + async def __aexit__( + self, + exc_type: type[BaseException] | None, + exc: BaseException | None, + tb: TracebackType | None, + ) -> None: + await self.aclose() + + def __repr__(self) -> str: + return f"AsyncAuraClient(base_url={self._config.base_url!r})" diff --git a/src/aura_python_sdk/_internal/_auth.py b/src/aura_python_sdk/_internal/_auth.py index 97a226b..aac2ab6 100644 --- a/src/aura_python_sdk/_internal/_auth.py +++ b/src/aura_python_sdk/_internal/_auth.py @@ -1,16 +1,24 @@ -"""OAuth client-credentials token management (Go: internal/api authManager).""" +"""OAuth client-credentials token management (Go: internal/api authManager). + +``_TokenSource`` holds everything except I/O and locking: the token request, response +validation, and freshness checks. ``TokenManager`` (threads) and ``AsyncTokenManager`` (asyncio) +add only a lock and the send. +""" from __future__ import annotations +import asyncio import base64 import json import logging import threading +from collections.abc import Callable from dataclasses import dataclass from urllib.parse import urlencode from aura_python_sdk._errors import AuraResponseError, AuthenticationError, api_error_from_response -from aura_python_sdk._internal.http._service import HttpService +from aura_python_sdk._internal.http._service import AsyncHttpService, HttpService +from aura_python_sdk._transport import HttpResponse # Refresh this many seconds before the token actually expires. REFRESH_MARGIN = 60.0 @@ -23,14 +31,12 @@ class _Token: access_token: str expires_at: float # on the HttpService clock (monotonic) + @property + def header(self) -> str: + return f"{self.token_type} {self.access_token}" -class TokenManager: - """Obtains and caches a bearer token from ``{base_url}/oauth/token``. - - Thread-safe. Concurrent callers that find the token missing or near expiry trigger a single - refresh between them. - """ +class _TokenSource: def __init__( self, *, @@ -38,56 +44,35 @@ def __init__( client_secret: str, token_url: str, user_agent: str, - http: HttpService, + clock: Callable[[], float], logger: logging.Logger, ) -> None: credentials = f"{client_id}:{client_secret}".encode() - self._basic_auth = "Basic " + base64.b64encode(credentials).decode("ascii") - self._token_url = token_url - self._user_agent = user_agent - self._http = http - self._logger = logger - self._lock = threading.Lock() - self._token: _Token | None = None - - def authorization_header(self, *, deadline: float) -> str: - """Return a valid ``Authorization`` header value, fetching a new token if needed.""" - token = self._token - if token is None or not self._is_fresh(token): - with self._lock: - token = self._token - if token is None or not self._is_fresh(token): - token = self._fetch(deadline=deadline) - self._token = token - return f"{token.token_type} {token.access_token}" - - def invalidate(self) -> None: - """Drop the cached token so the next request fetches a new one (e.g. after a 401).""" - with self._lock: - self._token = None - - def _is_fresh(self, token: _Token) -> bool: - return self._http.clock() < token.expires_at - REFRESH_MARGIN - - def _fetch(self, *, deadline: float) -> _Token: - self._logger.debug("obtaining new authentication token") - response = self._http.send( - "POST", - self._token_url, - { - "Authorization": self._basic_auth, - "Content-Type": "application/x-www-form-urlencoded", - "User-Agent": self._user_agent, - }, - urlencode({"grant_type": "client_credentials"}).encode("ascii"), - deadline=deadline, - ) + self.url = token_url + self.headers = { + "Authorization": "Basic " + base64.b64encode(credentials).decode("ascii"), + "Content-Type": "application/x-www-form-urlencoded", + "User-Agent": user_agent, + } + self.body = urlencode({"grant_type": "client_credentials"}).encode("ascii") + self.clock = clock + self.logger = logger + self.token: _Token | None = None + + def cached(self) -> _Token | None: + token = self.token + if token is not None and self.clock() < token.expires_at - REFRESH_MARGIN: + return token + return None + + def accept(self, response: HttpResponse) -> _Token: + """Validate a token response, cache the token, and return it.""" if not 200 <= response.status_code < 300: status = response.status_code # Any client error from the token endpoint means the credentials were rejected. # Rate limits and server errors keep their usual types. error_class = None if status == 429 or status >= 500 else AuthenticationError - self._logger.debug("token request failed", extra={"status": status}) + self.logger.debug("token request failed", extra={"status": status}) raise api_error_from_response( status, response.body, response.headers, error_class=error_class ) @@ -111,9 +96,106 @@ def _fetch(self, *, deadline: float) -> _Token: ): raise AuraResponseError(f"invalid expires_in value: {expires_in!r}") - self._logger.debug("token obtained", extra={"expires_in": expires_in}) - return _Token( + self.logger.debug("token obtained", extra={"expires_in": expires_in}) + self.token = _Token( token_type="Bearer", # noqa: S106 - the OAuth scheme name, not a secret access_token=access_token, - expires_at=self._http.clock() + float(expires_in), + expires_at=self.clock() + float(expires_in), + ) + return self.token + + +class TokenManager: + """Obtains and caches a bearer token from ``{base_url}/oauth/token``. + + Thread-safe. Concurrent callers that find the token missing or near expiry trigger a single + refresh between them. + """ + + def __init__( + self, + *, + client_id: str, + client_secret: str, + token_url: str, + user_agent: str, + http: HttpService, + logger: logging.Logger, + ) -> None: + self._source = _TokenSource( + client_id=client_id, + client_secret=client_secret, + token_url=token_url, + user_agent=user_agent, + clock=http.clock, + logger=logger, + ) + self._http = http + self._lock = threading.Lock() + + def authorization_header(self, *, deadline: float) -> str: + """Return a valid ``Authorization`` header value, fetching a new token if needed.""" + token = self._source.cached() + if token is None: + with self._lock: + # Another thread may have refreshed the token while this one waited. + token = self._source.cached() or self._fetch(deadline) + return token.header + + def _fetch(self, deadline: float) -> _Token: + source = self._source + source.logger.debug("obtaining new authentication token") + response = self._http.send( + "POST", source.url, source.headers, source.body, deadline=deadline + ) + return source.accept(response) + + def invalidate(self) -> None: + """Drop the cached token so the next request fetches a new one (e.g. after a 401).""" + with self._lock: + self._source.token = None + + +class AsyncTokenManager: + """The asyncio version of :class:`TokenManager`. Concurrent tasks share one refresh.""" + + def __init__( + self, + *, + client_id: str, + client_secret: str, + token_url: str, + user_agent: str, + http: AsyncHttpService, + logger: logging.Logger, + ) -> None: + self._source = _TokenSource( + client_id=client_id, + client_secret=client_secret, + token_url=token_url, + user_agent=user_agent, + clock=http.clock, + logger=logger, ) + self._http = http + self._lock = asyncio.Lock() + + async def authorization_header(self, *, deadline: float) -> str: + token = self._source.cached() + if token is None: + async with self._lock: + # Another task may have refreshed the token while this one waited. + token = self._source.cached() or await self._fetch(deadline) + return token.header + + async def _fetch(self, deadline: float) -> _Token: + source = self._source + source.logger.debug("obtaining new authentication token") + response = await self._http.send( + "POST", source.url, source.headers, source.body, deadline=deadline + ) + return source.accept(response) + + def invalidate(self) -> None: + # Safe without the lock: asyncio runs this between awaits, never mid-refresh. + self._source.token = None diff --git a/src/aura_python_sdk/_internal/_call.py b/src/aura_python_sdk/_internal/_call.py new file mode 100644 index 0000000..a65c9d0 --- /dev/null +++ b/src/aura_python_sdk/_internal/_call.py @@ -0,0 +1,45 @@ +"""A description of one API call, shared by the sync and async services. + +Each service operation validates its arguments and returns a ``Call``, without doing any I/O. +The sync and async services then run the same ``Call`` through their own request service, so +validation, paths, bodies and parsing are written once. +""" + +from __future__ import annotations + +from collections.abc import Callable, Mapping +from dataclasses import dataclass, field +from typing import Generic, TypeVar + +from aura_python_sdk._internal._request import ApiResponse, QueryParams +from aura_python_sdk._internal._serde import parse_data, parse_data_list + +T = TypeVar("T") + + +@dataclass(frozen=True, slots=True, kw_only=True) +class Call(Generic[T]): + method: str + path: str + parse: Callable[[ApiResponse], T] + params: QueryParams | None = None + json_body: object = None + # Logged at DEBUG before sending. + describe: str + # Logged at INFO after success. Used for operations that change something. + done: str | None = None + context: Mapping[str, object] = field(default_factory=dict) + + +def one(cls: type[T]) -> Callable[[ApiResponse], T]: + """Parse ``{"data": {...}}`` into ``cls``.""" + return lambda response: parse_data(cls, response.json()) + + +def many(cls: type[T]) -> Callable[[ApiResponse], list[T]]: + """Parse ``{"data": [...]}`` into a list of ``cls``.""" + return lambda response: parse_data_list(cls, response.json()) + + +def nothing(response: ApiResponse) -> None: + """For endpoints that return no body (204).""" diff --git a/src/aura_python_sdk/_internal/_request.py b/src/aura_python_sdk/_internal/_request.py index 6bee12c..1c3274a 100644 --- a/src/aura_python_sdk/_internal/_request.py +++ b/src/aura_python_sdk/_internal/_request.py @@ -4,13 +4,14 @@ import json import logging -from collections.abc import Mapping +from collections.abc import Callable, Mapping from dataclasses import dataclass, field from urllib.parse import quote, urlencode from aura_python_sdk._errors import AuraResponseError, api_error_from_response -from aura_python_sdk._internal._auth import TokenManager -from aura_python_sdk._internal.http._service import HttpService +from aura_python_sdk._internal._auth import AsyncTokenManager, TokenManager +from aura_python_sdk._internal.http._service import AsyncHttpService, HttpService +from aura_python_sdk._transport import HttpResponse QueryParams = Mapping[str, str | None] @@ -33,8 +34,8 @@ def json(self) -> object: raise AuraResponseError("response body is not valid JSON") from exc -class RequestService: - """Adds authentication, headers and URL handling, and maps error responses to exceptions. +class _Requests: + """URL, header and body handling plus error mapping, shared by the sync and async services. A relative path such as ``instances/abc`` resolves to ``{base_url}/{api_version}/instances/abc``. @@ -42,6 +43,62 @@ class RequestService: still gets the Aura bearer token. """ + def __init__( + self, + *, + base_url: str, + api_version: str, + user_agent: str, + default_headers: Mapping[str, str], + timeout: float, + logger: logging.Logger, + ) -> None: + self.endpoint_base = f"{base_url}/{api_version}" + self.user_agent = user_agent + self.default_headers = dict(default_headers) + self.timeout = timeout + self.logger = logger + + def resolve_url(self, path: str, params: QueryParams | None) -> str: + if path.startswith(("https://", "http://")): + url = path + else: + url = f"{self.endpoint_base}/{path.lstrip('/')}" + query = {key: value for key, value in (params or {}).items() if value is not None} + if query: + url += ("&" if "?" in url else "?") + urlencode(query) + return url + + def headers(self, authorization: str) -> dict[str, str]: + headers = dict(self.default_headers) + headers["Content-Type"] = "application/json" + headers["User-Agent"] = self.user_agent + headers["Authorization"] = authorization + return headers + + @staticmethod + def body(json_body: object) -> bytes | None: + return None if json_body is None else json.dumps(json_body, separators=(",", ":")).encode() + + def finish( + self, method: str, url: str, response: HttpResponse, invalidate: Callable[[], None] + ) -> ApiResponse: + if not 200 <= response.status_code < 300: + if response.status_code == 401: + # The token may have been revoked; make the next call fetch a fresh one. + invalidate() + error = api_error_from_response(response.status_code, response.body, response.headers) + self.logger.debug( + "API returned error", + extra={"method": method, "url": url, "status": response.status_code}, + ) + raise error + return ApiResponse(response.status_code, response.headers, response.body) + + +class RequestService: + """Adds authentication, headers and URL handling, and maps error responses to exceptions.""" + def __init__( self, *, @@ -56,11 +113,14 @@ def __init__( ) -> None: self._http = http self._auth = auth - self._endpoint_base = f"{base_url}/{api_version}" - self._user_agent = user_agent - self._default_headers = dict(default_headers) - self._timeout = timeout - self._logger = logger + self._requests = _Requests( + base_url=base_url, + api_version=api_version, + user_agent=user_agent, + default_headers=default_headers, + timeout=timeout, + logger=logger, + ) def get(self, path: str, *, params: QueryParams | None = None) -> ApiResponse: return self.request("GET", path, params=params) @@ -87,38 +147,60 @@ def request( ) -> ApiResponse: # One deadline covers the token fetch, every attempt and every backoff, like the # context.WithTimeout that wraps each Go service method. - deadline = self._http.clock() + self._timeout - url = self._resolve_url(path, params) - - headers = dict(self._default_headers) - headers["Content-Type"] = "application/json" - headers["User-Agent"] = self._user_agent - headers["Authorization"] = self._auth.authorization_header(deadline=deadline) - - body = None if json_body is None else json.dumps(json_body, separators=(",", ":")).encode() + deadline = self._http.clock() + self._requests.timeout + url = self._requests.resolve_url(path, params) + headers = self._requests.headers(self._auth.authorization_header(deadline=deadline)) + self._requests.logger.debug( + "making authenticated API request", extra={"method": method, "url": url} + ) + response = self._http.send( + method, url, headers, self._requests.body(json_body), deadline=deadline + ) + return self._requests.finish(method, url, response, self._auth.invalidate) - self._logger.debug("making authenticated API request", extra={"method": method, "url": url}) - response = self._http.send(method, url, headers, body, deadline=deadline) - if not 200 <= response.status_code < 300: - if response.status_code == 401: - # The token may have been revoked; make the next call fetch a fresh one. - self._auth.invalidate() - error = api_error_from_response(response.status_code, response.body, response.headers) - self._logger.debug( - "API returned error", - extra={"method": method, "url": url, "status": response.status_code}, - ) - raise error +class AsyncRequestService: + """The asyncio version of :class:`RequestService`.""" - return ApiResponse(response.status_code, response.headers, response.body) - - def _resolve_url(self, path: str, params: QueryParams | None) -> str: - if path.startswith(("https://", "http://")): - url = path - else: - url = f"{self._endpoint_base}/{path.lstrip('/')}" - query = {key: value for key, value in (params or {}).items() if value is not None} - if query: - url += ("&" if "?" in url else "?") + urlencode(query) - return url + def __init__( + self, + *, + http: AsyncHttpService, + auth: AsyncTokenManager, + base_url: str, + api_version: str, + user_agent: str, + default_headers: Mapping[str, str], + timeout: float, + logger: logging.Logger, + ) -> None: + self._http = http + self._auth = auth + self._requests = _Requests( + base_url=base_url, + api_version=api_version, + user_agent=user_agent, + default_headers=default_headers, + timeout=timeout, + logger=logger, + ) + + async def request( + self, + method: str, + path: str, + *, + params: QueryParams | None = None, + json_body: object = None, + ) -> ApiResponse: + deadline = self._http.clock() + self._requests.timeout + url = self._requests.resolve_url(path, params) + authorization = await self._auth.authorization_header(deadline=deadline) + headers = self._requests.headers(authorization) + self._requests.logger.debug( + "making authenticated API request", extra={"method": method, "url": url} + ) + response = await self._http.send( + method, url, headers, self._requests.body(json_body), deadline=deadline + ) + return self._requests.finish(method, url, response, self._auth.invalidate) diff --git a/src/aura_python_sdk/_internal/http/_httpx.py b/src/aura_python_sdk/_internal/http/_httpx.py index 8166382..a7283ed 100644 --- a/src/aura_python_sdk/_internal/http/_httpx.py +++ b/src/aura_python_sdk/_internal/http/_httpx.py @@ -28,16 +28,31 @@ def _tls_context() -> ssl.SSLContext: return context +def _translate(exc: httpx.TransportError) -> AuraConnectionError: + """Map an httpx network error to the SDK's own exception.""" + request_sent = not isinstance(exc, _NOT_SENT_ERRORS) + if isinstance(exc, httpx.TimeoutException): + return AuraTimeoutError(f"request timed out: {exc}", request_sent=request_sent) + return AuraConnectionError(f"request failed: {exc}", request_sent=request_sent) + + +def _response(response: httpx.Response, body: bytes) -> HttpResponse: + return HttpResponse( + status_code=response.status_code, headers=dict(response.headers.items()), body=body + ) + + +def _too_large(limit: int) -> AuraResponseError: + return AuraResponseError(f"response body exceeded limit of {limit} bytes") + + class HttpxTransport: """An :class:`~aura_python_sdk.HttpTransport` backed by a pooled ``httpx.Client``.""" def __init__(self, *, _httpx_transport: httpx.BaseTransport | None = None) -> None: # _httpx_transport is only for tests; it replaces the network layer below httpx. self._client = httpx.Client( - verify=_tls_context(), - limits=_LIMITS, - follow_redirects=True, - transport=_httpx_transport, + verify=_tls_context(), limits=_LIMITS, follow_redirects=True, transport=_httpx_transport ) def send(self, request: HttpRequest) -> HttpResponse: @@ -49,31 +64,48 @@ def send(self, request: HttpRequest) -> HttpResponse: content=request.body, timeout=httpx.Timeout(request.timeout), ) as response: - body = self._read_limited(response, request.max_response_size) - return HttpResponse( - status_code=response.status_code, - headers=dict(response.headers.items()), - body=body, - ) - except httpx.TimeoutException as exc: - raise AuraTimeoutError( - f"request timed out: {exc}", request_sent=not isinstance(exc, _NOT_SENT_ERRORS) - ) from exc + chunks: list[bytes] = [] + size = 0 + for chunk in response.iter_bytes(): + size += len(chunk) + if size > request.max_response_size: + raise _too_large(request.max_response_size) + chunks.append(chunk) + return _response(response, b"".join(chunks)) except httpx.TransportError as exc: - raise AuraConnectionError( - f"request failed: {exc}", request_sent=not isinstance(exc, _NOT_SENT_ERRORS) - ) from exc - - @staticmethod - def _read_limited(response: httpx.Response, limit: int) -> bytes: - chunks: list[bytes] = [] - size = 0 - for chunk in response.iter_bytes(): - size += len(chunk) - if size > limit: - raise AuraResponseError(f"response body exceeded limit of {limit} bytes") - chunks.append(chunk) - return b"".join(chunks) + raise _translate(exc) from exc def close(self) -> None: self._client.close() + + +class AsyncHttpxTransport: + """An :class:`~aura_python_sdk.AsyncHttpTransport` backed by a pooled ``httpx.AsyncClient``.""" + + def __init__(self, *, _httpx_transport: httpx.AsyncBaseTransport | None = None) -> None: + self._client = httpx.AsyncClient( + verify=_tls_context(), limits=_LIMITS, follow_redirects=True, transport=_httpx_transport + ) + + async def send(self, request: HttpRequest) -> HttpResponse: + try: + async with self._client.stream( + request.method, + request.url, + headers=dict(request.headers), + content=request.body, + timeout=httpx.Timeout(request.timeout), + ) as response: + chunks: list[bytes] = [] + size = 0 + async for chunk in response.aiter_bytes(): + size += len(chunk) + if size > request.max_response_size: + raise _too_large(request.max_response_size) + chunks.append(chunk) + return _response(response, b"".join(chunks)) + except httpx.TransportError as exc: + raise _translate(exc) from exc + + async def aclose(self) -> None: + await self._client.aclose() diff --git a/src/aura_python_sdk/_internal/http/_service.py b/src/aura_python_sdk/_internal/http/_service.py index 5c299c4..96be5ec 100644 --- a/src/aura_python_sdk/_internal/http/_service.py +++ b/src/aura_python_sdk/_internal/http/_service.py @@ -1,13 +1,18 @@ -"""Retries and response limits on top of an HttpTransport (Go: internal/httpclient).""" +"""Retries and response limits on top of a transport (Go: internal/httpclient). + +The retry policy is written once, as pure functions. ``HttpService`` and ``AsyncHttpService`` are +thin loops around it. +""" from __future__ import annotations +import asyncio import logging import time -from collections.abc import Callable, Mapping +from collections.abc import Awaitable, Callable, Mapping from aura_python_sdk._errors import AuraConnectionError, AuraResponseError, AuraTimeoutError -from aura_python_sdk._transport import HttpRequest, HttpResponse, HttpTransport +from aura_python_sdk._transport import AsyncHttpTransport, HttpRequest, HttpResponse, HttpTransport # Methods that are safe to repeat when the server may already have received the request. _IDEMPOTENT_METHODS = frozenset({"GET", "HEAD", "OPTIONS", "PUT", "DELETE"}) @@ -16,15 +21,76 @@ RETRY_WAIT_MAX = 5.0 -class HttpService: - """Sends requests through a transport, retrying network failures only. +class _RetryPolicy: + """Network failures only: exponential backoff (1 s doubling to 5 s), never past the deadline. - As in the Go SDK, a response with any HTTP status is final and is never retried. A network - failure is retried up to ``max_retries`` times with exponential backoff (1 s doubling to 5 s). - If the request may have reached the server, only idempotent methods are retried, so a - ``POST /instances`` is never sent twice. No attempt or backoff runs past ``deadline``. + As in the Go SDK, a response with any HTTP status is final. If the request may have reached + the server, only idempotent methods are retried, so a ``POST /instances`` is never sent twice. """ + def __init__( + self, + *, + max_retries: int, + max_response_size: int, + logger: logging.Logger, + clock: Callable[[], float], + ) -> None: + self.max_retries = max_retries + self.max_response_size = max_response_size + self.logger = logger + self.clock = clock + + def build_request( + self, + method: str, + url: str, + headers: Mapping[str, str], + body: bytes | None, + deadline: float, + ) -> HttpRequest: + remaining = deadline - self.clock() + if remaining <= 0: + raise AuraTimeoutError("request deadline exceeded", request_sent=False) + self.logger.debug("sending HTTP request", extra={"method": method, "url": url}) + return HttpRequest( + method=method, + url=url, + headers=headers, + body=body, + timeout=remaining, + max_response_size=self.max_response_size, + ) + + def retry_wait( + self, method: str, url: str, exc: AuraConnectionError, attempt: int, deadline: float + ) -> float | None: + """Seconds to wait before retrying, or None to give up and re-raise.""" + wait = min(RETRY_WAIT_MAX, RETRY_WAIT_MIN * 2.0**attempt) + retryable = not exc.request_sent or method.upper() in _IDEMPOTENT_METHODS + if attempt >= self.max_retries or not retryable or self.clock() + wait >= deadline: + return None + self.logger.debug( + "retrying HTTP request after network error", + extra={"method": method, "url": url, "attempt": attempt + 1, "error": str(exc)}, + ) + return wait + + def check_response(self, method: str, url: str, response: HttpResponse) -> HttpResponse: + if len(response.body) > self.max_response_size: + raise AuraResponseError( + f"response body exceeded limit of {self.max_response_size} bytes" + ) + self.logger.debug( + "HTTP response received", + extra={"method": method, "url": url, "status": response.status_code}, + ) + return response + + +class HttpService: + """Sends requests through a sync transport with the shared retry policy.""" + def __init__( self, transport: HttpTransport, @@ -36,15 +102,14 @@ def __init__( sleep: Callable[[float], None] = time.sleep, ) -> None: self._transport = transport - self._max_retries = max_retries - self._max_response_size = max_response_size - self._logger = logger - self._clock = clock + self._policy = _RetryPolicy( + max_retries=max_retries, max_response_size=max_response_size, logger=logger, clock=clock + ) self._sleep = sleep @property def clock(self) -> Callable[[], float]: - return self._clock + return self._policy.clock def send( self, @@ -57,46 +122,61 @@ def send( ) -> HttpResponse: attempt = 0 while True: - remaining = deadline - self._clock() - if remaining <= 0: - raise AuraTimeoutError("request deadline exceeded", request_sent=False) - request = HttpRequest( - method=method, - url=url, - headers=headers, - body=body, - timeout=remaining, - max_response_size=self._max_response_size, - ) - self._logger.debug("sending HTTP request", extra={"method": method, "url": url}) + request = self._policy.build_request(method, url, headers, body, deadline) try: response = self._transport.send(request) except AuraConnectionError as exc: - wait = min(RETRY_WAIT_MAX, RETRY_WAIT_MIN * 2**attempt) - if ( - attempt >= self._max_retries - or not self._is_retryable(method, exc) - or self._clock() + wait >= deadline - ): + wait = self._policy.retry_wait(method, url, exc, attempt, deadline) + if wait is None: raise - self._logger.debug( - "retrying HTTP request after network error", - extra={"method": method, "url": url, "attempt": attempt + 1, "error": str(exc)}, - ) self._sleep(wait) attempt += 1 continue + return self._policy.check_response(method, url, response) - if len(response.body) > self._max_response_size: - raise AuraResponseError( - f"response body exceeded limit of {self._max_response_size} bytes" - ) - self._logger.debug( - "HTTP response received", - extra={"method": method, "url": url, "status": response.status_code}, - ) - return response - @staticmethod - def _is_retryable(method: str, exc: AuraConnectionError) -> bool: - return not exc.request_sent or method.upper() in _IDEMPOTENT_METHODS +class AsyncHttpService: + """Sends requests through an async transport with the shared retry policy.""" + + def __init__( + self, + transport: AsyncHttpTransport, + *, + max_retries: int, + max_response_size: int, + logger: logging.Logger, + clock: Callable[[], float] = time.monotonic, + sleep: Callable[[float], Awaitable[None]] = asyncio.sleep, + ) -> None: + self._transport = transport + self._policy = _RetryPolicy( + max_retries=max_retries, max_response_size=max_response_size, logger=logger, clock=clock + ) + self._sleep = sleep + + @property + def clock(self) -> Callable[[], float]: + return self._policy.clock + + async def send( + self, + method: str, + url: str, + headers: Mapping[str, str], + body: bytes | None, + *, + deadline: float, + ) -> HttpResponse: + attempt = 0 + while True: + request = self._policy.build_request(method, url, headers, body, deadline) + try: + response = await self._transport.send(request) + except AuraConnectionError as exc: + wait = self._policy.retry_wait(method, url, exc, attempt, deadline) + if wait is None: + raise + await self._sleep(wait) + attempt += 1 + continue + return self._policy.check_response(method, url, response) diff --git a/src/aura_python_sdk/_transport.py b/src/aura_python_sdk/_transport.py index fe55315..f9c5fe7 100644 --- a/src/aura_python_sdk/_transport.py +++ b/src/aura_python_sdk/_transport.py @@ -55,3 +55,15 @@ class HttpTransport(Protocol): def send(self, request: HttpRequest) -> HttpResponse: ... def close(self) -> None: ... + + +@runtime_checkable +class AsyncHttpTransport(Protocol): + """The async counterpart of :class:`HttpTransport`, for :class:`AsyncAuraClient`. + + ``send`` follows the same error rules as :meth:`HttpTransport.send`. + """ + + async def send(self, request: HttpRequest) -> HttpResponse: ... + + async def aclose(self) -> None: ... diff --git a/src/aura_python_sdk/services/__init__.py b/src/aura_python_sdk/services/__init__.py index 881a931..6aac086 100644 --- a/src/aura_python_sdk/services/__init__.py +++ b/src/aura_python_sdk/services/__init__.py @@ -1,13 +1,20 @@ -"""The grouped services exposed on :class:`~aura_python_sdk.AuraClient`.""" +"""The grouped services exposed on :class:`~aura_python_sdk.AuraClient` and +:class:`~aura_python_sdk.AsyncAuraClient`.""" -from aura_python_sdk.services.cmek import CMEKService -from aura_python_sdk.services.graph_analytics import GDSSessionService -from aura_python_sdk.services.instances import InstanceService -from aura_python_sdk.services.prometheus import PrometheusService -from aura_python_sdk.services.snapshots import SnapshotService -from aura_python_sdk.services.tenants import TenantService +from aura_python_sdk.services.cmek import AsyncCMEKService, CMEKService +from aura_python_sdk.services.graph_analytics import AsyncGDSSessionService, GDSSessionService +from aura_python_sdk.services.instances import AsyncInstanceService, InstanceService +from aura_python_sdk.services.prometheus import AsyncPrometheusService, PrometheusService +from aura_python_sdk.services.snapshots import AsyncSnapshotService, SnapshotService +from aura_python_sdk.services.tenants import AsyncTenantService, TenantService __all__ = [ + "AsyncCMEKService", + "AsyncGDSSessionService", + "AsyncInstanceService", + "AsyncPrometheusService", + "AsyncSnapshotService", + "AsyncTenantService", "CMEKService", "GDSSessionService", "InstanceService", diff --git a/src/aura_python_sdk/services/_base.py b/src/aura_python_sdk/services/_base.py index b416b22..7b6663b 100644 --- a/src/aura_python_sdk/services/_base.py +++ b/src/aura_python_sdk/services/_base.py @@ -1,8 +1,12 @@ from __future__ import annotations import logging +from typing import TypeVar -from aura_python_sdk._internal._request import RequestService +from aura_python_sdk._internal._call import Call +from aura_python_sdk._internal._request import AsyncRequestService, RequestService + +T = TypeVar("T") # List-filter query parameter names, as the v1 spec defines them. (The Go SDK sends tenant_id.) TENANT_ID_PARAM = "tenantId" @@ -11,8 +15,36 @@ class Service: - """Shared plumbing for the grouped services on :class:`AuraClient`.""" + """Base for the sync services on :class:`AuraClient`: runs each operation's ``Call``.""" def __init__(self, api: RequestService, logger: logging.Logger) -> None: self._api = api self._logger = logger + + def _run(self, call: Call[T]) -> T: + self._logger.debug(call.describe, extra=dict(call.context)) + response = self._api.request( + call.method, call.path, params=call.params, json_body=call.json_body + ) + result = call.parse(response) + if call.done: + self._logger.info(call.done, extra=dict(call.context)) + return result + + +class AsyncService: + """Base for the async services on :class:`AsyncAuraClient`: awaits each operation's ``Call``.""" + + def __init__(self, api: AsyncRequestService, logger: logging.Logger) -> None: + self._api = api + self._logger = logger + + async def _run(self, call: Call[T]) -> T: + self._logger.debug(call.describe, extra=dict(call.context)) + response = await self._api.request( + call.method, call.path, params=call.params, json_body=call.json_body + ) + result = call.parse(response) + if call.done: + self._logger.info(call.done, extra=dict(call.context)) + return result diff --git a/src/aura_python_sdk/services/cmek.py b/src/aura_python_sdk/services/cmek.py index 365ee8d..8f75781 100644 --- a/src/aura_python_sdk/services/cmek.py +++ b/src/aura_python_sdk/services/cmek.py @@ -5,33 +5,95 @@ import builtins from aura_python_sdk import _validation as validate +from aura_python_sdk._internal._call import Call, many, nothing, one from aura_python_sdk._internal._request import build_path -from aura_python_sdk._internal._serde import parse_data, parse_data_list, to_json +from aura_python_sdk._internal._serde import to_json from aura_python_sdk.models._common import CloudProvider, InstanceType from aura_python_sdk.models.cmek import CustomerManagedKey, CustomerManagedKeySummary -from aura_python_sdk.services._base import TENANT_ID_PARAM, Service +from aura_python_sdk.services._base import TENANT_ID_PARAM, AsyncService, Service _KEYS = "customer-managed-keys" +# --- Operations (validation, request and parsing; no I/O) --- + + +def _list(tenant_id: str | None) -> Call[list[CustomerManagedKeySummary]]: + if tenant_id is not None: + tenant_id = validate.tenant_id(tenant_id) + return Call( + method="GET", + path=_KEYS, + params={TENANT_ID_PARAM: tenant_id}, + parse=many(CustomerManagedKeySummary), + describe="listing customer managed keys", + context={"tenant_id": tenant_id}, + ) + + +def _get(key_id: str) -> Call[CustomerManagedKey]: + key_id = validate.require_non_empty("customer managed key ID", key_id) + return Call( + method="GET", + path=build_path(_KEYS, key_id), + parse=one(CustomerManagedKey), + describe="getting customer managed key", + context={"key_id": key_id}, + ) + + +def _create( + *, + name: str, + key_id: str, + tenant_id: str, + cloud_provider: CloudProvider | str, + region: str, + instance_type: InstanceType | str, +) -> Call[CustomerManagedKey]: + body = { + "name": validate.display_name("key name", name), + "key_id": validate.require_non_empty("cloud provider key ID", key_id), + "tenant_id": validate.tenant_id(tenant_id), + "cloud_provider": validate.require_non_empty("cloud provider", cloud_provider), + "region": validate.require_non_empty("region", region), + "instance_type": validate.require_non_empty("instance type", instance_type), + } + return Call( + method="POST", + path=_KEYS, + json_body=to_json(body), + parse=one(CustomerManagedKey), + describe="creating customer managed key", + done="customer managed key created", + context={"key_name": name, "tenant_id": tenant_id}, + ) + + +def _delete(key_id: str) -> Call[None]: + key_id = validate.require_non_empty("customer managed key ID", key_id) + return Call( + method="DELETE", + path=build_path(_KEYS, key_id), + parse=nothing, + describe="deleting customer managed key", + done="customer managed key deleted", + context={"key_id": key_id}, + ) + + +# --- Services --- + class CMEKService(Service): """Customer-managed encryption keys.""" def list(self, tenant_id: str | None = None) -> builtins.list[CustomerManagedKeySummary]: """Every key the credentials can access, optionally only those in one tenant.""" - if tenant_id is not None: - tenant_id = validate.tenant_id(tenant_id) - self._logger.debug("listing customer managed keys", extra={"tenant_id": tenant_id}) - response = self._api.get(_KEYS, params={TENANT_ID_PARAM: tenant_id}) - keys = parse_data_list(CustomerManagedKeySummary, response.json()) - self._logger.debug("customer managed keys listed", extra={"count": len(keys)}) - return keys + return self._run(_list(tenant_id)) def get(self, key_id: str) -> CustomerManagedKey: """Full details of one key. ``key_id`` is the Aura key ID, not the cloud provider's.""" - key_id = validate.require_non_empty("customer managed key ID", key_id) - self._logger.debug("getting customer managed key", extra={"key_id": key_id}) - return parse_data(CustomerManagedKey, self._api.get(build_path(_KEYS, key_id)).json()) + return self._run(_get(key_id)) def create( self, @@ -48,24 +110,55 @@ def create( ``key_id`` is the key's ID in the cloud provider (the key ARN on AWS). The key can then encrypt new ``instance_type`` instances in ``region``. It starts in ``pending`` status. """ - body = { - "name": validate.display_name("key name", name), - "key_id": validate.require_non_empty("cloud provider key ID", key_id), - "tenant_id": validate.tenant_id(tenant_id), - "cloud_provider": validate.require_non_empty("cloud provider", cloud_provider), - "region": validate.require_non_empty("region", region), - "instance_type": validate.require_non_empty("instance type", instance_type), - } - self._logger.debug( - "creating customer managed key", extra={"key_name": name, "tenant_id": tenant_id} + return self._run( + _create( + name=name, + key_id=key_id, + tenant_id=tenant_id, + cloud_provider=cloud_provider, + region=region, + instance_type=instance_type, + ) ) - key = parse_data(CustomerManagedKey, self._api.post(_KEYS, json_body=to_json(body)).json()) - self._logger.info("customer managed key created", extra={"key_id": key.id}) - return key def delete(self, key_id: str) -> None: """Delete a key. The API refuses if any instance still uses it.""" - key_id = validate.require_non_empty("customer managed key ID", key_id) - self._logger.debug("deleting customer managed key", extra={"key_id": key_id}) - self._api.delete(build_path(_KEYS, key_id)) - self._logger.info("customer managed key deleted", extra={"key_id": key_id}) + self._run(_delete(key_id)) + + +class AsyncCMEKService(AsyncService): + """Async version of :class:`CMEKService`, with the same arguments and behaviour.""" + + async def list(self, tenant_id: str | None = None) -> builtins.list[CustomerManagedKeySummary]: + """See :meth:`CMEKService.list`.""" + return await self._run(_list(tenant_id)) + + async def get(self, key_id: str) -> CustomerManagedKey: + """See :meth:`CMEKService.get`.""" + return await self._run(_get(key_id)) + + async def create( + self, + *, + name: str, + key_id: str, + tenant_id: str, + cloud_provider: CloudProvider | str, + region: str, + instance_type: InstanceType | str, + ) -> CustomerManagedKey: + """See :meth:`CMEKService.create`.""" + return await self._run( + _create( + name=name, + key_id=key_id, + tenant_id=tenant_id, + cloud_provider=cloud_provider, + region=region, + instance_type=instance_type, + ) + ) + + async def delete(self, key_id: str) -> None: + """See :meth:`CMEKService.delete`.""" + await self._run(_delete(key_id)) diff --git a/src/aura_python_sdk/services/graph_analytics.py b/src/aura_python_sdk/services/graph_analytics.py index 911c8c2..c09171e 100644 --- a/src/aura_python_sdk/services/graph_analytics.py +++ b/src/aura_python_sdk/services/graph_analytics.py @@ -7,8 +7,9 @@ from aura_python_sdk import _validation as validate from aura_python_sdk._errors import AuraValidationError +from aura_python_sdk._internal._call import Call, many, one from aura_python_sdk._internal._request import build_path -from aura_python_sdk._internal._serde import parse_data, parse_data_list, to_json +from aura_python_sdk._internal._serde import to_json from aura_python_sdk.models.graph_analytics import ( DeletedGDSSession, GDSSession, @@ -19,11 +20,112 @@ INSTANCE_ID_PARAM, ORGANIZATION_ID_PARAM, TENANT_ID_PARAM, + AsyncService, Service, ) _SESSIONS = "graph-analytics/sessions" +# --- Operations (validation, request and parsing; no I/O) --- + + +def _list( + tenant_id: str | None, instance_id: str | None, organization_id: str | None +) -> Call[list[GDSSession]]: + params = { + TENANT_ID_PARAM: None if tenant_id is None else validate.tenant_id(tenant_id), + INSTANCE_ID_PARAM: None if instance_id is None else validate.instance_id(instance_id), + ORGANIZATION_ID_PARAM: None + if organization_id is None + else validate.require_non_empty("organization ID", organization_id), + } + return Call( + method="GET", + path=_SESSIONS, + params=params, + parse=many(GDSSession), + describe="listing GDS sessions", + ) + + +def _estimate_size( + node_count: int, + relationship_count: int, + node_property_count: int | None, + node_label_count: int | None, + relationship_property_count: int | None, + algorithm_categories: Sequence[str] | None, +) -> Call[GDSSessionSizeEstimate]: + body: dict[str, object] = { + "node_count": validate.non_negative_int("node count", node_count), + "relationship_count": validate.non_negative_int("relationship count", relationship_count), + } + optional_counts = { + "node_property_count": node_property_count, + "node_label_count": node_label_count, + "relationship_property_count": relationship_property_count, + } + for key, value in optional_counts.items(): + if value is not None: + body[key] = validate.non_negative_int(key.replace("_", " "), value) + if algorithm_categories is not None: + body["algorithm_categories"] = validate.string_list( + "algorithm categories", algorithm_categories + ) + return Call( + method="POST", + path=f"{_SESSIONS}/sizing", + json_body=body, + parse=one(GDSSessionSizeEstimate), + describe="estimating GDS session size", + ) + + +def _create(config: GDSSessionConfig) -> Call[GDSSession]: + if not isinstance(config, GDSSessionConfig): + raise AuraValidationError("config must be a GDSSessionConfig") + validate.require_non_empty("session name", config.name) + validate.require_non_empty("memory", config.memory) + if config.tenant_id is not None: + validate.tenant_id(config.tenant_id) + if config.instance_id is not None: + validate.instance_id(config.instance_id) + return Call( + method="POST", + path=_SESSIONS, + json_body=to_json(config), + parse=one(GDSSession), + describe="creating GDS session", + done="GDS session created", + context={"session_name": config.name}, + ) + + +def _get(session_id: str) -> Call[GDSSession]: + session_id = validate.session_id(session_id) + return Call( + method="GET", + path=build_path("graph-analytics", "sessions", session_id), + parse=one(GDSSession), + describe="getting GDS session", + context={"session_id": session_id}, + ) + + +def _delete(session_id: str) -> Call[DeletedGDSSession]: + session_id = validate.session_id(session_id) + return Call( + method="DELETE", + path=build_path("graph-analytics", "sessions", session_id), + parse=one(DeletedGDSSession), + describe="deleting GDS session", + done="GDS session deleted", + context={"session_id": session_id}, + ) + + +# --- Services --- + class GDSSessionService(Service): """Graph Analytics (GDS) sessions.""" @@ -36,17 +138,7 @@ def list( organization_id: str | None = None, ) -> builtins.list[GDSSession]: """Every session the credentials can access, optionally filtered.""" - params = { - TENANT_ID_PARAM: None if tenant_id is None else validate.tenant_id(tenant_id), - INSTANCE_ID_PARAM: None if instance_id is None else validate.instance_id(instance_id), - ORGANIZATION_ID_PARAM: None - if organization_id is None - else validate.require_non_empty("organization ID", organization_id), - } - self._logger.debug("listing GDS sessions") - sessions = parse_data_list(GDSSession, self._api.get(_SESSIONS, params=params).json()) - self._logger.debug("GDS sessions listed", extra={"count": len(sessions)}) - return sessions + return self._run(_list(tenant_id, instance_id, organization_id)) def estimate_size( self, @@ -59,28 +151,16 @@ def estimate_size( algorithm_categories: Sequence[str] | None = None, ) -> GDSSessionSizeEstimate: """Estimate the session size needed for a graph (Go: ``Estimate``).""" - body: dict[str, object] = { - "node_count": validate.non_negative_int("node count", node_count), - "relationship_count": validate.non_negative_int( - "relationship count", relationship_count - ), - } - optional_counts = { - "node_property_count": node_property_count, - "node_label_count": node_label_count, - "relationship_property_count": relationship_property_count, - } - for key, value in optional_counts.items(): - if value is not None: - body[key] = validate.non_negative_int(key.replace("_", " "), value) - if algorithm_categories is not None: - body["algorithm_categories"] = validate.string_list( - "algorithm categories", algorithm_categories + return self._run( + _estimate_size( + node_count, + relationship_count, + node_property_count, + node_label_count, + relationship_property_count, + algorithm_categories, ) - - self._logger.debug("estimating GDS session size") - response = self._api.post(f"{_SESSIONS}/sizing", json_body=body) - return parse_data(GDSSessionSizeEstimate, response.json()) + ) def create(self, config: GDSSessionConfig) -> GDSSession: """Create a session, or return the matching existing one. @@ -88,35 +168,60 @@ def create(self, config: GDSSessionConfig) -> GDSSession: Attach it to an instance with ``instance_id`` and ``database_uuid``, or make a standalone session with ``cloud_provider`` and ``region``. """ - if not isinstance(config, GDSSessionConfig): - raise AuraValidationError("config must be a GDSSessionConfig") - validate.require_non_empty("session name", config.name) - validate.require_non_empty("memory", config.memory) - if config.tenant_id is not None: - validate.tenant_id(config.tenant_id) - if config.instance_id is not None: - validate.instance_id(config.instance_id) - - self._logger.debug("creating GDS session", extra={"session_name": config.name}) - session = parse_data( - GDSSession, self._api.post(_SESSIONS, json_body=to_json(config)).json() - ) - self._logger.info("GDS session created", extra={"session_id": session.id}) - return session + return self._run(_create(config)) def get(self, session_id: str) -> GDSSession: """Details of one session.""" - session_id = validate.session_id(session_id) - self._logger.debug("getting GDS session", extra={"session_id": session_id}) - return parse_data( - GDSSession, self._api.get(build_path("graph-analytics", "sessions", session_id)).json() - ) + return self._run(_get(session_id)) def delete(self, session_id: str) -> DeletedGDSSession: """Delete a session.""" - session_id = validate.session_id(session_id) - self._logger.debug("deleting GDS session", extra={"session_id": session_id}) - response = self._api.delete(build_path("graph-analytics", "sessions", session_id)) - deleted = parse_data(DeletedGDSSession, response.json()) - self._logger.info("GDS session deleted", extra={"session_id": session_id}) - return deleted + return self._run(_delete(session_id)) + + +class AsyncGDSSessionService(AsyncService): + """Async version of :class:`GDSSessionService`, with the same arguments and behaviour.""" + + async def list( + self, + *, + tenant_id: str | None = None, + instance_id: str | None = None, + organization_id: str | None = None, + ) -> builtins.list[GDSSession]: + """See :meth:`GDSSessionService.list`.""" + return await self._run(_list(tenant_id, instance_id, organization_id)) + + async def estimate_size( + self, + *, + node_count: int, + relationship_count: int, + node_property_count: int | None = None, + node_label_count: int | None = None, + relationship_property_count: int | None = None, + algorithm_categories: Sequence[str] | None = None, + ) -> GDSSessionSizeEstimate: + """See :meth:`GDSSessionService.estimate_size`.""" + return await self._run( + _estimate_size( + node_count, + relationship_count, + node_property_count, + node_label_count, + relationship_property_count, + algorithm_categories, + ) + ) + + async def create(self, config: GDSSessionConfig) -> GDSSession: + """See :meth:`GDSSessionService.create`.""" + return await self._run(_create(config)) + + async def get(self, session_id: str) -> GDSSession: + """See :meth:`GDSSessionService.get`.""" + return await self._run(_get(session_id)) + + async def delete(self, session_id: str) -> DeletedGDSSession: + """See :meth:`GDSSessionService.delete`.""" + return await self._run(_delete(session_id)) diff --git a/src/aura_python_sdk/services/instances.py b/src/aura_python_sdk/services/instances.py index aa0406e..2f52a66 100644 --- a/src/aura_python_sdk/services/instances.py +++ b/src/aura_python_sdk/services/instances.py @@ -7,8 +7,9 @@ from aura_python_sdk import _validation as validate from aura_python_sdk._errors import AuraValidationError +from aura_python_sdk._internal._call import Call, many, one from aura_python_sdk._internal._request import build_path -from aura_python_sdk._internal._serde import parse_data, parse_data_list, to_json +from aura_python_sdk._internal._serde import to_json from aura_python_sdk.models._common import InstanceType from aura_python_sdk.models.instances import ( CDCEnrichmentMode, @@ -18,7 +19,220 @@ InstanceSizeEstimate, InstanceSummary, ) -from aura_python_sdk.services._base import TENANT_ID_PARAM, Service +from aura_python_sdk.services._base import TENANT_ID_PARAM, AsyncService, Service + +# --- Operations (validation, request and parsing; no I/O) --- + + +def _list(tenant_id: str | None) -> Call[list[InstanceSummary]]: + if tenant_id is not None: + tenant_id = validate.tenant_id(tenant_id) + return Call( + method="GET", + path="instances", + params={TENANT_ID_PARAM: tenant_id}, + parse=many(InstanceSummary), + describe="listing instances", + context={"tenant_id": tenant_id}, + ) + + +def _get(instance_id: str) -> Call[Instance]: + instance_id = validate.instance_id(instance_id) + return Call( + method="GET", + path=build_path("instances", instance_id), + parse=one(Instance), + describe="getting instance", + context={"instance_id": instance_id}, + ) + + +def _create( + config: InstanceConfig, + source_instance_id: str | None = None, + source_snapshot_id: str | None = None, +) -> Call[CreatedInstance]: + if source_instance_id is not None: + source_instance_id = validate.instance_id(source_instance_id, "source instance ID") + if source_snapshot_id is not None: + source_snapshot_id = validate.snapshot_id(source_snapshot_id, "source snapshot ID") + body = _create_body(config) + if source_instance_id is not None: + body["source_instance_id"] = source_instance_id + if source_snapshot_id is not None: + body["source_snapshot_id"] = source_snapshot_id + return Call( + method="POST", + path="instances", + json_body=body, + parse=one(CreatedInstance), + describe="creating instance", + done="instance creation started", + context={"instance_name": config.name, "tenant_id": config.tenant_id}, + ) + + +def _update( + instance_id: str, + *, + name: str | None, + memory: str | None, + storage: str | None, + vector_optimized: bool | None, + graph_analytics_plugin: bool | None, + cdc_enrichment_mode: CDCEnrichmentMode | str | None, + secondaries_count: int | None, +) -> Call[Instance]: + instance_id = validate.instance_id(instance_id) + changes: dict[str, object] = {} + if name is not None: + changes["name"] = validate.instance_name(name) + if memory is not None: + changes["memory"] = validate.require_non_empty("memory", memory) + if storage is not None: + changes["storage"] = validate.require_non_empty("storage", storage) + if vector_optimized is not None: + changes["vector_optimized"] = validate.boolean("vector optimized", vector_optimized) + if graph_analytics_plugin is not None: + changes["graph_analytics_plugin"] = validate.boolean( + "graph analytics plugin", graph_analytics_plugin + ) + if cdc_enrichment_mode is not None: + changes["cdc_enrichment_mode"] = validate.require_non_empty( + "CDC enrichment mode", cdc_enrichment_mode + ) + if secondaries_count is not None: + changes["secondaries_count"] = validate.non_negative_int( + "secondaries count", secondaries_count + ) + if not changes: + raise AuraValidationError("update requires at least one field to change") + return Call( + method="PATCH", + path=build_path("instances", instance_id), + json_body=to_json(changes), + parse=one(Instance), + describe="updating instance", + done="instance update started", + context={"instance_id": instance_id, "fields": sorted(changes)}, + ) + + +def _estimate_size( + node_count: int, + relationship_count: int, + instance_type: InstanceType | str | None, + algorithm_categories: Sequence[str] | None, +) -> Call[InstanceSizeEstimate]: + body: dict[str, object] = { + "node_count": validate.non_negative_int("node count", node_count), + "relationship_count": validate.non_negative_int("relationship count", relationship_count), + } + if instance_type is not None: + body["instance_type"] = validate.require_non_empty("instance type", instance_type) + if algorithm_categories is not None: + body["algorithm_categories"] = validate.string_list( + "algorithm categories", algorithm_categories + ) + return Call( + method="POST", + path="instances/sizing", + json_body=to_json(body), + parse=one(InstanceSizeEstimate), + describe="estimating instance size", + ) + + +def _upgrade(instance_id: str, memory: str | None, storage: str | None) -> Call[Instance]: + instance_id = validate.instance_id(instance_id) + if (memory is None) != (storage is None): + raise AuraValidationError("upgrade requires both memory and storage, or neither") + body: dict[str, object] = {} + if memory is not None and storage is not None: + body["memory"] = validate.require_non_empty("memory", memory) + body["storage"] = validate.require_non_empty("storage", storage) + return Call( + method="POST", + path=build_path("instances", instance_id, "upgrade"), + json_body=body, + parse=one(Instance), + describe="upgrading instance", + done="instance upgrade started", + context={"instance_id": instance_id}, + ) + + +def _delete(instance_id: str) -> Call[Instance]: + instance_id = validate.instance_id(instance_id) + return Call( + method="DELETE", + path=build_path("instances", instance_id), + parse=one(Instance), + describe="deleting instance", + done="instance deletion started", + context={"instance_id": instance_id}, + ) + + +def _lifecycle(instance_id: str, action: str) -> Call[Instance]: + instance_id = validate.instance_id(instance_id) + return Call( + method="POST", + path=build_path("instances", instance_id, action), + parse=one(Instance), + describe=f"{action} instance", + done=f"instance {action} started", + context={"instance_id": instance_id}, + ) + + +def _overwrite( + instance_id: str, + *, + source_instance_id: str | None = None, + source_snapshot_id: str | None = None, +) -> Call[Instance]: + instance_id = validate.instance_id(instance_id) + if source_instance_id is not None: + body = { + "source_instance_id": validate.instance_id(source_instance_id, "source instance ID") + } + else: + body = { + "source_snapshot_id": validate.snapshot_id(source_snapshot_id, "source snapshot ID") + } + return Call( + method="POST", + path=build_path("instances", instance_id, "overwrite"), + json_body=body, + parse=one(Instance), + describe="overwriting instance", + done="instance overwrite started", + context={"instance_id": instance_id, **body}, + ) + + +def _create_body(config: InstanceConfig) -> dict[str, object]: + """Validate a create request as the Go SDK's validateCreateInstanceConfig does.""" + if not isinstance(config, InstanceConfig): + raise AuraValidationError("config must be an InstanceConfig") + validate.instance_name(config.name) + validate.tenant_id(config.tenant_id) + validate.require_non_empty("cloud provider", config.cloud_provider) + validate.require_non_empty("region", config.region) + validate.require_non_empty("instance type", config.type) + validate.require_non_empty("version", config.version) + validate.require_non_empty("memory", config.memory) + if config.customer_managed_key_id is not None: + validate.require_non_empty("customer managed key ID", config.customer_managed_key_id) + body = to_json(config) + if not isinstance(body, dict): # pragma: no cover - to_json of a dataclass is a dict + raise TypeError("expected a JSON object") + return body + + +# --- Services --- class InstanceService(Service): @@ -26,19 +240,11 @@ class InstanceService(Service): def list(self, tenant_id: str | None = None) -> builtins.list[InstanceSummary]: """Every instance the credentials can access, optionally only those in one tenant.""" - if tenant_id is not None: - tenant_id = validate.tenant_id(tenant_id) - self._logger.debug("listing instances", extra={"tenant_id": tenant_id}) - response = self._api.get("instances", params={TENANT_ID_PARAM: tenant_id}) - instances = parse_data_list(InstanceSummary, response.json()) - self._logger.debug("instances listed", extra={"count": len(instances)}) - return instances + return self._run(_list(tenant_id)) def get(self, instance_id: str) -> Instance: """Full details of one instance.""" - instance_id = validate.instance_id(instance_id) - self._logger.debug("getting instance", extra={"instance_id": instance_id}) - return parse_data(Instance, self._api.get(build_path("instances", instance_id)).json()) + return self._run(_get(instance_id)) def create(self, config: InstanceConfig) -> CreatedInstance: """Start creating an instance. @@ -46,17 +252,13 @@ def create(self, config: InstanceConfig) -> CreatedInstance: Creation is asynchronous. Poll :meth:`get` until ``status`` is ``running``. The returned password is shown only once. """ - body = _create_body(config) - return self._create(body) + return self._run(_create(config)) def create_from_instance( self, source_instance_id: str, config: InstanceConfig ) -> CreatedInstance: """Create an instance cloned from the current data of another instance.""" - source_instance_id = validate.instance_id(source_instance_id, "source instance ID") - body = _create_body(config) - body["source_instance_id"] = source_instance_id - return self._create(body) + return self._run(_create(config, source_instance_id=source_instance_id)) def create_from_snapshot( self, source_instance_id: str, source_snapshot_id: str, config: InstanceConfig @@ -65,24 +267,7 @@ def create_from_snapshot( The snapshot must belong to ``source_instance_id`` and be exportable. """ - source_instance_id = validate.instance_id(source_instance_id, "source instance ID") - source_snapshot_id = validate.snapshot_id(source_snapshot_id, "source snapshot ID") - body = _create_body(config) - body["source_instance_id"] = source_instance_id - body["source_snapshot_id"] = source_snapshot_id - return self._create(body) - - def _create(self, body: dict[str, object]) -> CreatedInstance: - self._logger.debug( - "creating instance", - extra={"instance_name": body["name"], "tenant_id": body["tenant_id"]}, - ) - created = parse_data(CreatedInstance, self._api.post("instances", json_body=body).json()) - self._logger.info( - "instance creation started", - extra={"instance_id": created.id, "instance_name": created.name}, - ) - return created + return self._run(_create(config, source_instance_id, source_snapshot_id)) def update( self, @@ -102,38 +287,18 @@ def update( ``secondaries_count`` applies only to Virtual Dedicated Cloud, and ``cdc_enrichment_mode`` only to Virtual Dedicated Cloud and Business Critical. """ - instance_id = validate.instance_id(instance_id) - changes: dict[str, object] = {} - if name is not None: - changes["name"] = validate.instance_name(name) - if memory is not None: - changes["memory"] = validate.require_non_empty("memory", memory) - if storage is not None: - changes["storage"] = validate.require_non_empty("storage", storage) - if vector_optimized is not None: - changes["vector_optimized"] = validate.boolean("vector optimized", vector_optimized) - if graph_analytics_plugin is not None: - changes["graph_analytics_plugin"] = validate.boolean( - "graph analytics plugin", graph_analytics_plugin - ) - if cdc_enrichment_mode is not None: - changes["cdc_enrichment_mode"] = validate.require_non_empty( - "CDC enrichment mode", cdc_enrichment_mode + return self._run( + _update( + instance_id, + name=name, + memory=memory, + storage=storage, + vector_optimized=vector_optimized, + graph_analytics_plugin=graph_analytics_plugin, + cdc_enrichment_mode=cdc_enrichment_mode, + secondaries_count=secondaries_count, ) - if secondaries_count is not None: - changes["secondaries_count"] = validate.non_negative_int( - "secondaries count", secondaries_count - ) - if not changes: - raise AuraValidationError("update requires at least one field to change") - - self._logger.debug( - "updating instance", extra={"instance_id": instance_id, "fields": sorted(changes)} ) - response = self._api.patch(build_path("instances", instance_id), json_body=to_json(changes)) - instance = parse_data(Instance, response.json()) - self._logger.info("instance update started", extra={"instance_id": instance_id}) - return instance def estimate_size( self, @@ -148,21 +313,9 @@ def estimate_size( Supported for ``enterprise-ds`` and ``professional-ds``. Pass the recommended size as ``memory`` when creating the instance. """ - body: dict[str, object] = { - "node_count": validate.non_negative_int("node count", node_count), - "relationship_count": validate.non_negative_int( - "relationship count", relationship_count - ), - } - if instance_type is not None: - body["instance_type"] = validate.require_non_empty("instance type", instance_type) - if algorithm_categories is not None: - body["algorithm_categories"] = validate.string_list( - "algorithm categories", algorithm_categories - ) - self._logger.debug("estimating instance size") - response = self._api.post("instances/sizing", json_body=to_json(body)) - return parse_data(InstanceSizeEstimate, response.json()) + return self._run( + _estimate_size(node_count, relationship_count, instance_type, algorithm_categories) + ) def upgrade( self, instance_id: str, *, memory: str | None = None, storage: str | None = None @@ -172,79 +325,117 @@ def upgrade( Pass both ``memory`` and ``storage`` to resize as part of the upgrade, or neither to keep the current size. Not available for Marketplace projects or trial instances. """ - instance_id = validate.instance_id(instance_id) - if (memory is None) != (storage is None): - raise AuraValidationError("upgrade requires both memory and storage, or neither") - body: dict[str, object] = {} - if memory is not None and storage is not None: - body["memory"] = validate.require_non_empty("memory", memory) - body["storage"] = validate.require_non_empty("storage", storage) - self._logger.debug("upgrading instance", extra={"instance_id": instance_id}) - response = self._api.post(build_path("instances", instance_id, "upgrade"), json_body=body) - instance = parse_data(Instance, response.json()) - self._logger.info("instance upgrade started", extra={"instance_id": instance_id}) - return instance + return self._run(_upgrade(instance_id, memory, storage)) def delete(self, instance_id: str) -> Instance: """Start deleting an instance. This cannot be undone.""" - instance_id = validate.instance_id(instance_id) - self._logger.debug("deleting instance", extra={"instance_id": instance_id}) - instance = parse_data( - Instance, self._api.delete(build_path("instances", instance_id)).json() - ) - self._logger.info("instance deletion started", extra={"instance_id": instance_id}) - return instance + return self._run(_delete(instance_id)) def pause(self, instance_id: str) -> Instance: """Pause a running instance.""" - return self._lifecycle(instance_id, "pause") + return self._run(_lifecycle(instance_id, "pause")) def resume(self, instance_id: str) -> Instance: """Resume a paused instance.""" - return self._lifecycle(instance_id, "resume") - - def _lifecycle(self, instance_id: str, action: str) -> Instance: - instance_id = validate.instance_id(instance_id) - self._logger.debug("%s instance", action, extra={"instance_id": instance_id}) - response = self._api.post(build_path("instances", instance_id, action)) - instance = parse_data(Instance, response.json()) - self._logger.info("instance %s started", action, extra={"instance_id": instance_id}) - return instance + return self._run(_lifecycle(instance_id, "resume")) def overwrite_from_instance(self, instance_id: str, source_instance_id: str) -> Instance: """Replace an instance's data with the current data of another instance.""" - instance_id = validate.instance_id(instance_id) - source_instance_id = validate.instance_id(source_instance_id, "source instance ID") - return self._overwrite(instance_id, {"source_instance_id": source_instance_id}) + return self._run(_overwrite(instance_id, source_instance_id=source_instance_id)) def overwrite_from_snapshot(self, instance_id: str, source_snapshot_id: str) -> Instance: """Replace an instance's data with a snapshot.""" - instance_id = validate.instance_id(instance_id) - source_snapshot_id = validate.snapshot_id(source_snapshot_id, "source snapshot ID") - return self._overwrite(instance_id, {"source_snapshot_id": source_snapshot_id}) + return self._run(_overwrite(instance_id, source_snapshot_id=source_snapshot_id)) - def _overwrite(self, instance_id: str, body: dict[str, object]) -> Instance: - self._logger.debug("overwriting instance", extra={"instance_id": instance_id, **body}) - response = self._api.post(build_path("instances", instance_id, "overwrite"), json_body=body) - instance = parse_data(Instance, response.json()) - self._logger.info("instance overwrite started", extra={"instance_id": instance_id}) - return instance +class AsyncInstanceService(AsyncService): + """Async version of :class:`InstanceService`, with the same arguments and behaviour.""" -def _create_body(config: InstanceConfig) -> dict[str, object]: - """Validate a create request as the Go SDK's validateCreateInstanceConfig does.""" - if not isinstance(config, InstanceConfig): - raise AuraValidationError("config must be an InstanceConfig") - validate.instance_name(config.name) - validate.tenant_id(config.tenant_id) - validate.require_non_empty("cloud provider", config.cloud_provider) - validate.require_non_empty("region", config.region) - validate.require_non_empty("instance type", config.type) - validate.require_non_empty("version", config.version) - validate.require_non_empty("memory", config.memory) - if config.customer_managed_key_id is not None: - validate.require_non_empty("customer managed key ID", config.customer_managed_key_id) - body = to_json(config) - if not isinstance(body, dict): # pragma: no cover - to_json of a dataclass is a dict - raise TypeError("expected a JSON object") - return body + async def list(self, tenant_id: str | None = None) -> builtins.list[InstanceSummary]: + """See :meth:`InstanceService.list`.""" + return await self._run(_list(tenant_id)) + + async def get(self, instance_id: str) -> Instance: + """See :meth:`InstanceService.get`.""" + return await self._run(_get(instance_id)) + + async def create(self, config: InstanceConfig) -> CreatedInstance: + """See :meth:`InstanceService.create`.""" + return await self._run(_create(config)) + + async def create_from_instance( + self, source_instance_id: str, config: InstanceConfig + ) -> CreatedInstance: + """See :meth:`InstanceService.create_from_instance`.""" + return await self._run(_create(config, source_instance_id=source_instance_id)) + + async def create_from_snapshot( + self, source_instance_id: str, source_snapshot_id: str, config: InstanceConfig + ) -> CreatedInstance: + """See :meth:`InstanceService.create_from_snapshot`.""" + return await self._run(_create(config, source_instance_id, source_snapshot_id)) + + async def update( + self, + instance_id: str, + *, + name: str | None = None, + memory: str | None = None, + storage: str | None = None, + vector_optimized: bool | None = None, + graph_analytics_plugin: bool | None = None, + cdc_enrichment_mode: CDCEnrichmentMode | str | None = None, + secondaries_count: int | None = None, + ) -> Instance: + """See :meth:`InstanceService.update`.""" + return await self._run( + _update( + instance_id, + name=name, + memory=memory, + storage=storage, + vector_optimized=vector_optimized, + graph_analytics_plugin=graph_analytics_plugin, + cdc_enrichment_mode=cdc_enrichment_mode, + secondaries_count=secondaries_count, + ) + ) + + async def estimate_size( + self, + *, + node_count: int, + relationship_count: int, + instance_type: InstanceType | str | None = None, + algorithm_categories: Sequence[str] | None = None, + ) -> InstanceSizeEstimate: + """See :meth:`InstanceService.estimate_size`.""" + return await self._run( + _estimate_size(node_count, relationship_count, instance_type, algorithm_categories) + ) + + async def upgrade( + self, instance_id: str, *, memory: str | None = None, storage: str | None = None + ) -> Instance: + """See :meth:`InstanceService.upgrade`.""" + return await self._run(_upgrade(instance_id, memory, storage)) + + async def delete(self, instance_id: str) -> Instance: + """See :meth:`InstanceService.delete`.""" + return await self._run(_delete(instance_id)) + + async def pause(self, instance_id: str) -> Instance: + """See :meth:`InstanceService.pause`.""" + return await self._run(_lifecycle(instance_id, "pause")) + + async def resume(self, instance_id: str) -> Instance: + """See :meth:`InstanceService.resume`.""" + return await self._run(_lifecycle(instance_id, "resume")) + + async def overwrite_from_instance(self, instance_id: str, source_instance_id: str) -> Instance: + """See :meth:`InstanceService.overwrite_from_instance`.""" + return await self._run(_overwrite(instance_id, source_instance_id=source_instance_id)) + + async def overwrite_from_snapshot(self, instance_id: str, source_snapshot_id: str) -> Instance: + """See :meth:`InstanceService.overwrite_from_snapshot`.""" + return await self._run(_overwrite(instance_id, source_snapshot_id=source_snapshot_id)) diff --git a/src/aura_python_sdk/services/prometheus.py b/src/aura_python_sdk/services/prometheus.py index ce45469..eb0e0d3 100644 --- a/src/aura_python_sdk/services/prometheus.py +++ b/src/aura_python_sdk/services/prometheus.py @@ -9,7 +9,8 @@ from aura_python_sdk import _validation as validate from aura_python_sdk._errors import AuraResponseError, AuraValidationError, MetricNotFoundError -from aura_python_sdk._internal._request import RequestService +from aura_python_sdk._internal._call import Call +from aura_python_sdk._internal._request import ApiResponse, AsyncRequestService, RequestService from aura_python_sdk._internal.metrics._parser import parse_exposition from aura_python_sdk.models.prometheus import ( ConnectionMetrics, @@ -20,12 +21,130 @@ ResourceMetrics, StorageMetrics, ) -from aura_python_sdk.services._base import Service +from aura_python_sdk.services._base import AsyncService, Service # The Aura bearer token is sent with every metrics request, so by default only Aura's own # metrics hosts are allowed. _TRUSTED_METRICS_DOMAIN = "neo4j.io" +# --- Operations and pure helpers (no I/O) --- + + +def _check_url(prometheus_url: str, *, allow_untrusted: bool) -> str: + url = validate.require_non_empty("prometheus URL", prometheus_url) + parts = urlsplit(url) + if parts.scheme not in ("https", "http") or not parts.hostname: + raise AuraValidationError(f"prometheus URL is not a valid http(s) URL: {url!r}") + if allow_untrusted: + return url + host = parts.hostname.lower() + trusted = host == _TRUSTED_METRICS_DOMAIN or host.endswith(f".{_TRUSTED_METRICS_DOMAIN}") + if parts.scheme != "https" or not trusted: + raise AuraValidationError( + f"prometheus URL must be an https://*.{_TRUSTED_METRICS_DOMAIN} address, because " + "the Aura API token is sent with the request" + ) + return url + + +def _parse_metrics(response: ApiResponse) -> PrometheusMetrics: + try: + text = response.body.decode("utf-8") + except UnicodeDecodeError as exc: + raise AuraResponseError("metrics response is not valid UTF-8") from exc + return PrometheusMetrics(metrics=parse_exposition(text)) + + +def _fetch(prometheus_url: str, *, allow_untrusted: bool) -> Call[PrometheusMetrics]: + url = _check_url(prometheus_url, allow_untrusted=allow_untrusted) + return Call( + method="GET", + path=url, + parse=_parse_metrics, + describe="fetching Prometheus metrics", + context={"url": url}, + ) + + +def metric_value( + metrics: PrometheusMetrics, name: str, label_filters: Mapping[str, str] | None = None +) -> float: + """The mean value of ``name`` across samples matching ``label_filters`` (Go semantics).""" + if not isinstance(metrics, PrometheusMetrics): + raise AuraValidationError("metrics must be a PrometheusMetrics") + samples = metrics.metrics.get(name) + if not samples: + raise MetricNotFoundError(f"metric {name} not found") + filters = dict(label_filters or {}) + matching = [s for s in samples if all(s.labels.get(k) == v for k, v in filters.items())] + if not matching: + raise MetricNotFoundError(f"no matching metrics found for {name} with filters {filters}") + return sum(s.value for s in matching) / len(matching) + + +def build_health( + instance_id: str, metrics: PrometheusMetrics, logger: logging.Logger +) -> InstanceHealth: + """Build the health summary from fetched metrics, using the Go SDK's metric names.""" + + def value(name: str) -> float | None: + try: + return metric_value(metrics, name) + except MetricNotFoundError: + logger.warning("metric not available", extra={"metric": name}) + return None + + cpu_usage = value("neo4j_aura_cpu_usage") + cpu_limit = value("neo4j_aura_cpu_limit") if cpu_usage is not None else None + heap_ratio = value("neo4j_dbms_vm_heap_used_ratio") + resources = ResourceMetrics( + cpu_usage_percent=( + cpu_usage / cpu_limit * 100 + if cpu_usage is not None and cpu_limit and cpu_limit > 0 + else None + ), + memory_usage_percent=heap_ratio * 100 if heap_ratio is not None else None, + ) + + query = QueryMetrics( + query_execution_total=value("neo4j_db_query_execution_success_total"), + avg_latency_ms=value("neo4j_db_query_execution_internal_latency_q50"), + ) + + idle = value("neo4j_dbms_bolt_connections_idle") + running = value("neo4j_dbms_bolt_connections_running") + max_connections = value("neo4j_dbms_bolt_connections_max_count") + active = int(idle + running) if idle is not None and running is not None else None + connections = ConnectionMetrics( + active_connections=active, + max_connections=int(max_connections) if max_connections and max_connections > 0 else None, + usage_percent=( + active / max_connections * 100 + if active is not None and max_connections and max_connections > 0 + else None + ), + ) + + hit_ratio = value("neo4j_dbms_page_cache_hit_ratio_per_minute") + storage = StorageMetrics(page_cache_hit_rate=hit_ratio * 100 if hit_ratio is not None else None) + + status, issues, recommendations = assess_health(resources, connections, storage) + logger.info("instance health assessed", extra={"instance_id": instance_id, "status": status}) + return InstanceHealth( + instance_id=instance_id, + timestamp=datetime.now(UTC), + resources=resources, + query=query, + connections=connections, + storage=storage, + overall_status=status, + issues=tuple(issues), + recommendations=tuple(recommendations), + ) + + +# --- Services --- + class PrometheusService(Service): """Aura's Prometheus metrics endpoints. @@ -42,16 +161,7 @@ def __init__( def fetch_raw_metrics(self, prometheus_url: str) -> PrometheusMetrics: """Fetch and parse every metric from a metrics endpoint.""" - url = self._check_url(prometheus_url) - self._logger.debug("fetching Prometheus metrics", extra={"url": url}) - body = self._api.get(url).body - try: - text = body.decode("utf-8") - except UnicodeDecodeError as exc: - raise AuraResponseError("metrics response is not valid UTF-8") from exc - metrics = PrometheusMetrics(metrics=parse_exposition(text)) - self._logger.debug("Prometheus metrics fetched", extra={"count": len(metrics.metrics)}) - return metrics + return self._run(_fetch(prometheus_url, allow_untrusted=self._allow_untrusted_urls)) def get_metric_value( self, @@ -63,18 +173,7 @@ def get_metric_value( Raises :class:`MetricNotFoundError` if nothing matches. """ - if not isinstance(metrics, PrometheusMetrics): - raise AuraValidationError("metrics must be a PrometheusMetrics") - samples = metrics.metrics.get(name) - if not samples: - raise MetricNotFoundError(f"metric {name} not found") - filters = dict(label_filters or {}) - matching = [s for s in samples if all(s.labels.get(k) == v for k, v in filters.items())] - if not matching: - raise MetricNotFoundError( - f"no matching metrics found for {name} with filters {filters}" - ) - return sum(s.value for s in matching) / len(matching) + return metric_value(metrics, name, label_filters) def get_instance_health(self, instance_id: str, prometheus_url: str) -> InstanceHealth: """Summarise an instance's CPU, memory, query, connection and page cache metrics. @@ -82,84 +181,44 @@ def get_instance_health(self, instance_id: str, prometheus_url: str) -> Instance Uses the same metrics, thresholds and status logic as the Go SDK. """ instance_id = validate.instance_id(instance_id) - metrics = self.fetch_raw_metrics(prometheus_url) - - def value(name: str) -> float | None: - try: - return self.get_metric_value(metrics, name) - except MetricNotFoundError: - self._logger.warning("metric not available", extra={"metric": name}) - return None - - cpu_usage = value("neo4j_aura_cpu_usage") - cpu_limit = value("neo4j_aura_cpu_limit") if cpu_usage is not None else None - heap_ratio = value("neo4j_dbms_vm_heap_used_ratio") - resources = ResourceMetrics( - cpu_usage_percent=( - cpu_usage / cpu_limit * 100 - if cpu_usage is not None and cpu_limit and cpu_limit > 0 - else None - ), - memory_usage_percent=heap_ratio * 100 if heap_ratio is not None else None, - ) + call = _fetch(prometheus_url, allow_untrusted=self._allow_untrusted_urls) + return build_health(instance_id, self._run(call), self._logger) - query = QueryMetrics( - query_execution_total=value("neo4j_db_query_execution_success_total"), - avg_latency_ms=value("neo4j_db_query_execution_internal_latency_q50"), - ) - idle = value("neo4j_dbms_bolt_connections_idle") - running = value("neo4j_dbms_bolt_connections_running") - max_connections = value("neo4j_dbms_bolt_connections_max_count") - active = int(idle + running) if idle is not None and running is not None else None - connections = ConnectionMetrics( - active_connections=active, - max_connections=int(max_connections) - if max_connections and max_connections > 0 - else None, - usage_percent=( - active / max_connections * 100 - if active is not None and max_connections and max_connections > 0 - else None - ), - ) +class AsyncPrometheusService(AsyncService): + """Async version of :class:`PrometheusService`, with the same arguments and behaviour. - hit_ratio = value("neo4j_dbms_page_cache_hit_ratio_per_minute") - storage = StorageMetrics( - page_cache_hit_rate=hit_ratio * 100 if hit_ratio is not None else None - ) + ``get_metric_value`` does no I/O, so it is a plain (non-async) method here too. + """ - status, issues, recommendations = assess_health(resources, connections, storage) - self._logger.info( - "instance health assessed", extra={"instance_id": instance_id, "status": status} - ) - return InstanceHealth( - instance_id=instance_id, - timestamp=datetime.now(UTC), - resources=resources, - query=query, - connections=connections, - storage=storage, - overall_status=status, - issues=tuple(issues), - recommendations=tuple(recommendations), - ) + def __init__( + self, + api: AsyncRequestService, + logger: logging.Logger, + *, + allow_untrusted_urls: bool = False, + ) -> None: + super().__init__(api, logger) + self._allow_untrusted_urls = allow_untrusted_urls - def _check_url(self, prometheus_url: str) -> str: - url = validate.require_non_empty("prometheus URL", prometheus_url) - parts = urlsplit(url) - if parts.scheme not in ("https", "http") or not parts.hostname: - raise AuraValidationError(f"prometheus URL is not a valid http(s) URL: {url!r}") - if self._allow_untrusted_urls: - return url - host = parts.hostname.lower() - trusted = host == _TRUSTED_METRICS_DOMAIN or host.endswith(f".{_TRUSTED_METRICS_DOMAIN}") - if parts.scheme != "https" or not trusted: - raise AuraValidationError( - f"prometheus URL must be an https://*.{_TRUSTED_METRICS_DOMAIN} address, because " - "the Aura API token is sent with the request" - ) - return url + async def fetch_raw_metrics(self, prometheus_url: str) -> PrometheusMetrics: + """See :meth:`PrometheusService.fetch_raw_metrics`.""" + return await self._run(_fetch(prometheus_url, allow_untrusted=self._allow_untrusted_urls)) + + def get_metric_value( + self, + metrics: PrometheusMetrics, + name: str, + label_filters: Mapping[str, str] | None = None, + ) -> float: + """See :meth:`PrometheusService.get_metric_value`.""" + return metric_value(metrics, name, label_filters) + + async def get_instance_health(self, instance_id: str, prometheus_url: str) -> InstanceHealth: + """See :meth:`PrometheusService.get_instance_health`.""" + instance_id = validate.instance_id(instance_id) + call = _fetch(prometheus_url, allow_untrusted=self._allow_untrusted_urls) + return build_health(instance_id, await self._run(call), self._logger) def assess_health( diff --git a/src/aura_python_sdk/services/snapshots.py b/src/aura_python_sdk/services/snapshots.py index 88b1315..8d1ef1e 100644 --- a/src/aura_python_sdk/services/snapshots.py +++ b/src/aura_python_sdk/services/snapshots.py @@ -7,11 +7,67 @@ from aura_python_sdk import _validation as validate from aura_python_sdk._errors import AuraValidationError +from aura_python_sdk._internal._call import Call, many, one from aura_python_sdk._internal._request import build_path -from aura_python_sdk._internal._serde import parse_data, parse_data_list from aura_python_sdk.models.instances import Instance from aura_python_sdk.models.snapshots import CreatedSnapshot, Snapshot -from aura_python_sdk.services._base import Service +from aura_python_sdk.services._base import AsyncService, Service + +# --- Operations (validation, request and parsing; no I/O) --- + + +def _list(instance_id: str, date: dt.date | None) -> Call[list[Snapshot]]: + instance_id = validate.instance_id(instance_id) + if date is not None and (not isinstance(date, dt.date) or isinstance(date, dt.datetime)): + raise AuraValidationError("date must be a datetime.date") + return Call( + method="GET", + path=build_path("instances", instance_id, "snapshots"), + params={"date": date.isoformat() if date else None}, + parse=many(Snapshot), + describe="listing snapshots", + context={"instance_id": instance_id}, + ) + + +def _get(instance_id: str, snapshot_id: str) -> Call[Snapshot]: + instance_id = validate.instance_id(instance_id) + snapshot_id = validate.snapshot_id(snapshot_id) + return Call( + method="GET", + path=build_path("instances", instance_id, "snapshots", snapshot_id), + parse=one(Snapshot), + describe="getting snapshot", + context={"instance_id": instance_id, "snapshot_id": snapshot_id}, + ) + + +def _create(instance_id: str) -> Call[CreatedSnapshot]: + instance_id = validate.instance_id(instance_id) + return Call( + method="POST", + path=build_path("instances", instance_id, "snapshots"), + parse=one(CreatedSnapshot), + describe="creating snapshot", + done="snapshot started", + context={"instance_id": instance_id}, + ) + + +def _restore(instance_id: str, snapshot_id: str) -> Call[Instance]: + instance_id = validate.instance_id(instance_id) + snapshot_id = validate.snapshot_id(snapshot_id) + return Call( + method="POST", + path=build_path("instances", instance_id, "snapshots", snapshot_id, "restore"), + parse=one(Instance), + describe="restoring snapshot", + done="snapshot restore started", + context={"instance_id": instance_id, "snapshot_id": snapshot_id}, + ) + + +# --- Services --- class SnapshotService(Service): @@ -19,53 +75,36 @@ class SnapshotService(Service): def list(self, instance_id: str, date: dt.date | None = None) -> builtins.list[Snapshot]: """Snapshots of an instance taken on ``date``. The API defaults to today.""" - instance_id = validate.instance_id(instance_id) - if date is not None and (not isinstance(date, dt.date) or isinstance(date, dt.datetime)): - raise AuraValidationError("date must be a datetime.date") - self._logger.debug("listing snapshots", extra={"instance_id": instance_id}) - response = self._api.get( - build_path("instances", instance_id, "snapshots"), - params={"date": date.isoformat() if date else None}, - ) - snapshots = parse_data_list(Snapshot, response.json()) - self._logger.debug("snapshots listed", extra={"count": len(snapshots)}) - return snapshots + return self._run(_list(instance_id, date)) def get(self, instance_id: str, snapshot_id: str) -> Snapshot: """Details of one snapshot.""" - instance_id = validate.instance_id(instance_id) - snapshot_id = validate.snapshot_id(snapshot_id) - self._logger.debug( - "getting snapshot", extra={"instance_id": instance_id, "snapshot_id": snapshot_id} - ) - response = self._api.get(build_path("instances", instance_id, "snapshots", snapshot_id)) - return parse_data(Snapshot, response.json()) + return self._run(_get(instance_id, snapshot_id)) def create(self, instance_id: str) -> CreatedSnapshot: """Start an on-demand snapshot.""" - instance_id = validate.instance_id(instance_id) - self._logger.debug("creating snapshot", extra={"instance_id": instance_id}) - response = self._api.post(build_path("instances", instance_id, "snapshots")) - created = parse_data(CreatedSnapshot, response.json()) - self._logger.info( - "snapshot started", - extra={"instance_id": instance_id, "snapshot_id": created.snapshot_id}, - ) - return created + return self._run(_create(instance_id)) def restore(self, instance_id: str, snapshot_id: str) -> Instance: """Restore an instance from one of its own snapshots, replacing its current data.""" - instance_id = validate.instance_id(instance_id) - snapshot_id = validate.snapshot_id(snapshot_id) - self._logger.debug( - "restoring snapshot", extra={"instance_id": instance_id, "snapshot_id": snapshot_id} - ) - response = self._api.post( - build_path("instances", instance_id, "snapshots", snapshot_id, "restore") - ) - instance = parse_data(Instance, response.json()) - self._logger.info( - "snapshot restore started", - extra={"instance_id": instance_id, "snapshot_id": snapshot_id}, - ) - return instance + return self._run(_restore(instance_id, snapshot_id)) + + +class AsyncSnapshotService(AsyncService): + """Async version of :class:`SnapshotService`, with the same arguments and behaviour.""" + + async def list(self, instance_id: str, date: dt.date | None = None) -> builtins.list[Snapshot]: + """See :meth:`SnapshotService.list`.""" + return await self._run(_list(instance_id, date)) + + async def get(self, instance_id: str, snapshot_id: str) -> Snapshot: + """See :meth:`SnapshotService.get`.""" + return await self._run(_get(instance_id, snapshot_id)) + + async def create(self, instance_id: str) -> CreatedSnapshot: + """See :meth:`SnapshotService.create`.""" + return await self._run(_create(instance_id)) + + async def restore(self, instance_id: str, snapshot_id: str) -> Instance: + """See :meth:`SnapshotService.restore`.""" + return await self._run(_restore(instance_id, snapshot_id)) diff --git a/src/aura_python_sdk/services/tenants.py b/src/aura_python_sdk/services/tenants.py index 14e10c9..4e0bb2f 100644 --- a/src/aura_python_sdk/services/tenants.py +++ b/src/aura_python_sdk/services/tenants.py @@ -5,10 +5,41 @@ import builtins from aura_python_sdk import _validation as validate +from aura_python_sdk._internal._call import Call, many, one from aura_python_sdk._internal._request import build_path -from aura_python_sdk._internal._serde import parse_data, parse_data_list from aura_python_sdk.models.tenants import MetricsIntegration, Tenant, TenantSummary -from aura_python_sdk.services._base import Service +from aura_python_sdk.services._base import AsyncService, Service + +# --- Operations (validation, request and parsing; no I/O) --- + + +def _list() -> Call[list[TenantSummary]]: + return Call(method="GET", path="tenants", parse=many(TenantSummary), describe="listing tenants") + + +def _get(tenant_id: str) -> Call[Tenant]: + tenant_id = validate.tenant_id(tenant_id) + return Call( + method="GET", + path=build_path("tenants", tenant_id), + parse=one(Tenant), + describe="getting tenant", + context={"tenant_id": tenant_id}, + ) + + +def _get_metrics_integration(tenant_id: str) -> Call[MetricsIntegration]: + tenant_id = validate.tenant_id(tenant_id) + return Call( + method="GET", + path=build_path("tenants", tenant_id, "metrics-integration"), + parse=one(MetricsIntegration), + describe="getting tenant metrics integration", + context={"tenant_id": tenant_id}, + ) + + +# --- Services --- class TenantService(Service): @@ -16,20 +47,28 @@ class TenantService(Service): def list(self) -> builtins.list[TenantSummary]: """Every tenant the credentials can access.""" - self._logger.debug("listing tenants") - tenants = parse_data_list(TenantSummary, self._api.get("tenants").json()) - self._logger.debug("tenants listed", extra={"count": len(tenants)}) - return tenants + return self._run(_list()) def get(self, tenant_id: str) -> Tenant: """A tenant and the instance configurations it can create.""" - tenant_id = validate.tenant_id(tenant_id) - self._logger.debug("getting tenant", extra={"tenant_id": tenant_id}) - return parse_data(Tenant, self._api.get(build_path("tenants", tenant_id)).json()) + return self._run(_get(tenant_id)) def get_metrics_integration(self, tenant_id: str) -> MetricsIntegration: """The project-level Prometheus metrics endpoint (Go: ``GetMetrics``).""" - tenant_id = validate.tenant_id(tenant_id) - self._logger.debug("getting tenant metrics integration", extra={"tenant_id": tenant_id}) - response = self._api.get(build_path("tenants", tenant_id, "metrics-integration")) - return parse_data(MetricsIntegration, response.json()) + return self._run(_get_metrics_integration(tenant_id)) + + +class AsyncTenantService(AsyncService): + """Async version of :class:`TenantService`.""" + + async def list(self) -> builtins.list[TenantSummary]: + """Every tenant the credentials can access.""" + return await self._run(_list()) + + async def get(self, tenant_id: str) -> Tenant: + """A tenant and the instance configurations it can create.""" + return await self._run(_get(tenant_id)) + + async def get_metrics_integration(self, tenant_id: str) -> MetricsIntegration: + """The project-level Prometheus metrics endpoint (Go: ``GetMetrics``).""" + return await self._run(_get_metrics_integration(tenant_id)) diff --git a/tests/blackbox/conftest.py b/tests/blackbox/conftest.py index 0cf56cc..2b5591a 100644 --- a/tests/blackbox/conftest.py +++ b/tests/blackbox/conftest.py @@ -86,6 +86,16 @@ def client(self, **options: Any) -> aura.AuraClient: **options, ) + def async_client(self, **options: Any) -> aura.AsyncAuraClient: + options.setdefault("timeout", 5) + return aura.AsyncAuraClient( + client_id="local-id", + client_secret="local-secret", + base_url=self.url, + allow_insecure_base_url=True, + **options, + ) + def _handler(self) -> type[BaseHTTPRequestHandler]: fake = self diff --git a/tests/blackbox/test_blackbox_async.py b/tests/blackbox/test_blackbox_async.py new file mode 100644 index 0000000..4265496 --- /dev/null +++ b/tests/blackbox/test_blackbox_async.py @@ -0,0 +1,58 @@ +"""AsyncAuraClient end-to-end over real sockets and the real async httpx transport.""" + +import asyncio +import time + +import pytest + +import aura_python_sdk as aura +from tests.blackbox.conftest import FakeAura, Reply +from tests.blackbox.test_blackbox import INSTANCE, TENANT_ID + +pytestmark = pytest.mark.anyio + + +async def test_concurrent_gets_share_one_token(fake_aura: FakeAura) -> None: + fake_aura.route("GET", "/v1/instances/2f49c2b3", Reply.json(200, {"data": INSTANCE})) + async with fake_aura.async_client() as client: + results = await asyncio.gather(*(client.instances.get("2f49c2b3") for _ in range(5))) + assert {r.id for r in results} == {"2f49c2b3"} + assert [r.path for r in fake_aura.received].count("/oauth/token") == 1 + assert len(fake_aura.api_requests()) == 5 + + +async def test_list_with_filter_and_headers(fake_aura: FakeAura) -> None: + fake_aura.route("GET", "/v1/instances", Reply.json(200, {"data": []})) + async with fake_aura.async_client(user_agent="my-app/1") as client: + assert await client.instances.list(TENANT_ID) == [] + [request] = fake_aura.api_requests() + assert request.query == f"tenantId={TENANT_ID}" + assert request.headers["user-agent"] == "my-app/1" + assert request.headers["authorization"] == "Bearer local-token" + + +async def test_api_error_is_mapped(fake_aura: FakeAura) -> None: + fake_aura.route( + "POST", + "/v1/instances/2f49c2b3/pause", + Reply.json(409, {"errors": [{"message": "Instance is not running"}]}), + ) + async with fake_aura.async_client() as client: + with pytest.raises(aura.ConflictError, match="Instance is not running"): + await client.instances.pause("2f49c2b3") + + +async def test_slow_post_times_out_and_is_not_retried(fake_aura: FakeAura) -> None: + fake_aura.route("POST", "/v1/instances/2f49c2b3/pause", Reply(202, b"{}", delay=2.0)) + started = time.monotonic() + async with fake_aura.async_client(timeout=0.5, max_retries=3) as client: + with pytest.raises(aura.AuraTimeoutError): + await client.instances.pause("2f49c2b3") + assert time.monotonic() - started < 1.9 + assert len(fake_aura.api_requests()) == 1 + + +async def test_delete_with_no_content(fake_aura: FakeAura) -> None: + fake_aura.route("DELETE", "/v1/customer-managed-keys/key-1", Reply(204)) + async with fake_aura.async_client() as client: + await client.cmek.delete("key-1") diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 0000000..c7bce25 --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,7 @@ +import pytest + + +@pytest.fixture +def anyio_backend() -> str: + """Run @pytest.mark.anyio tests on asyncio only.""" + return "asyncio" diff --git a/tests/fakes.py b/tests/fakes.py index 474e188..8ab04f4 100644 --- a/tests/fakes.py +++ b/tests/fakes.py @@ -73,3 +73,23 @@ def close(self) -> None: @property def api_requests(self) -> list[HttpRequest]: return [r for r in self.requests if not r.url.endswith("/oauth/token")] + + +class FakeAsyncTransport(FakeTransport): + """The async counterpart of FakeTransport: same queue and recording, awaitable methods.""" + + async def send(self, request: HttpRequest) -> HttpResponse: # type: ignore[override] + return FakeTransport.send(self, request) + + async def aclose(self) -> None: + self.closed = True + + +@dataclass +class FakeAsyncSleep: + """Advances a FakeClock instead of sleeping.""" + + clock: FakeClock + + async def __call__(self, seconds: float) -> None: + self.clock.sleep(seconds) diff --git a/tests/integration/test_live.py b/tests/integration/test_live.py index ab0d128..f5e59ca 100644 --- a/tests/integration/test_live.py +++ b/tests/integration/test_live.py @@ -121,3 +121,10 @@ def test_create_pause_resume_delete(client: aura.AuraClient) -> None: _wait_for(client, created.id, aura.InstanceStatus.RUNNING) finally: client.instances.delete(created.id) + + +@pytest.mark.anyio +async def test_async_client_reads_the_same_data(client: aura.AuraClient) -> None: + async with aura.AsyncAuraClient.from_env(timeout=60) as async_client: + async_tenants = await async_client.tenants.list() + assert {t.id for t in async_tenants} == {t.id for t in client.tenants.list()} diff --git a/tests/transport/test_httpx_transport.py b/tests/transport/test_httpx_transport.py index 3340095..640c369 100644 --- a/tests/transport/test_httpx_transport.py +++ b/tests/transport/test_httpx_transport.py @@ -1,7 +1,7 @@ """HttpxTransport against httpx.MockTransport (tests may import httpx; src may not).""" import ssl -from collections.abc import Iterator +from collections.abc import AsyncIterator, Iterator import httpx import pytest @@ -132,3 +132,71 @@ def test_tls_minimum_is_1_2() -> None: assert context.minimum_version == ssl.TLSVersion.TLSv1_2 assert context.verify_mode == ssl.CERT_REQUIRED transport.close() + + +# --- AsyncHttpxTransport --- + +from aura_python_sdk import AsyncHttpTransport # noqa: E402 +from aura_python_sdk._internal.http._httpx import AsyncHttpxTransport # noqa: E402 + + +def _async_transport(handler: object) -> AsyncHttpxTransport: + return AsyncHttpxTransport(_httpx_transport=httpx.MockTransport(handler)) # type: ignore[arg-type] + + +def test_async_satisfies_protocol() -> None: + assert isinstance(AsyncHttpxTransport(), AsyncHttpTransport) + + +@pytest.mark.anyio +async def test_async_request_and_response_are_translated() -> None: + seen: list[httpx.Request] = [] + + async def handler(request: httpx.Request) -> httpx.Response: + seen.append(request) + return httpx.Response(202, headers={"X-Request-Id": "r1"}, content=b'{"data":{}}') + + transport = _async_transport(handler) + response = await transport.send(_request()) + await transport.aclose() + + [request] = seen + assert request.method == "POST" + assert request.headers["authorization"] == "Bearer t" + assert request.content == b'{"name":"x"}' + assert response.status_code == 202 + assert response.headers["x-request-id"] == "r1" + assert response.body == b'{"data":{}}' + + +@pytest.mark.anyio +async def test_async_body_over_limit_is_rejected() -> None: + async def stream() -> AsyncIterator[bytes]: + for _ in range(100): + yield b"x" * 512 + + handler = lambda r: httpx.Response(200, content=stream()) # noqa: E731 + with pytest.raises(AuraResponseError, match="exceeded limit"): + await _async_transport(handler).send(_request(max_response_size=1024)) + + +@pytest.mark.anyio +@pytest.mark.parametrize( + ("exc", "expected_type", "request_sent"), + [ + (httpx.ConnectError("refused"), AuraConnectionError, False), + (httpx.ReadTimeout("slow read"), AuraTimeoutError, True), + (httpx.PoolTimeout("pool"), AuraTimeoutError, False), + (httpx.RemoteProtocolError("bad"), AuraConnectionError, True), + ], +) +async def test_async_network_errors_are_translated( + exc: Exception, expected_type: type[AuraConnectionError], request_sent: bool +) -> None: + def handler(request: httpx.Request) -> httpx.Response: + raise exc + + with pytest.raises(expected_type) as info: + await _async_transport(handler).send(_request()) + assert type(info.value) is expected_type + assert info.value.request_sent is request_sent diff --git a/tests/unit/test_async_core.py b/tests/unit/test_async_core.py new file mode 100644 index 0000000..535a054 --- /dev/null +++ b/tests/unit/test_async_core.py @@ -0,0 +1,197 @@ +"""The async plumbing: retries, token sharing, 401 handling and client lifecycle.""" + +from __future__ import annotations + +import asyncio +import logging + +import pytest + +import aura_python_sdk as aura +from aura_python_sdk import AuraConnectionError, HttpRequest, HttpResponse +from aura_python_sdk._internal._auth import AsyncTokenManager +from aura_python_sdk._internal.http._httpx import AsyncHttpxTransport +from aura_python_sdk._internal.http._service import AsyncHttpService +from tests.fakes import ( + FakeAsyncSleep, + FakeAsyncTransport, + FakeClock, + FakeTransport, + json_response, + token_response, +) + +URL = "https://api.neo4j.io/v1/instances" +pytestmark = pytest.mark.anyio + + +def _http( + transport: FakeAsyncTransport, clock: FakeClock, max_retries: int = 3 +) -> AsyncHttpService: + return AsyncHttpService( + transport, + max_retries=max_retries, + max_response_size=1024, + logger=logging.getLogger("test"), + clock=clock, + sleep=FakeAsyncSleep(clock), + ) + + +async def test_retries_network_errors_with_backoff() -> None: + clock = FakeClock() + error = AuraConnectionError("reset", request_sent=True) + transport = FakeAsyncTransport([error, error, HttpResponse(200)]) + response = await _http(transport, clock).send("GET", URL, {}, None, deadline=clock.now + 60) + assert response.status_code == 200 + assert clock.sleeps == [1.0, 2.0] + + +async def test_post_not_retried_once_sent() -> None: + clock = FakeClock() + transport = FakeAsyncTransport([AuraConnectionError("reset", request_sent=True)]) + with pytest.raises(AuraConnectionError): + await _http(transport, clock).send("POST", URL, {}, b"{}", deadline=clock.now + 60) + assert len(transport.requests) == 1 + + +async def test_post_retried_when_never_sent() -> None: + clock = FakeClock() + transport = FakeAsyncTransport( + [AuraConnectionError("refused", request_sent=False), HttpResponse(202)] + ) + response = await _http(transport, clock).send("POST", URL, {}, b"{}", deadline=clock.now + 60) + assert response.status_code == 202 + + +async def test_deadline_stops_retries() -> None: + clock = FakeClock() + transport = FakeAsyncTransport([AuraConnectionError("refused", request_sent=False)]) + with pytest.raises(AuraConnectionError): + await _http(transport, clock).send("GET", URL, {}, None, deadline=clock.now + 0.5) + assert clock.sleeps == [] + + +async def test_expired_deadline() -> None: + clock = FakeClock() + with pytest.raises(aura.AuraTimeoutError): + await _http(FakeAsyncTransport(), clock).send("GET", URL, {}, None, deadline=clock.now) + + +async def test_oversized_body_rejected() -> None: + clock = FakeClock() + transport = FakeAsyncTransport([HttpResponse(200, body=b"x" * 2000)]) + with pytest.raises(aura.AuraResponseError, match="exceeded limit"): + await _http(transport, clock).send("GET", URL, {}, None, deadline=clock.now + 5) + + +async def test_concurrent_tasks_share_one_token_fetch() -> None: + fetches = 0 + + class SlowTokenTransport(FakeAsyncTransport): + async def send(self, request: HttpRequest) -> HttpResponse: # type: ignore[override] + nonlocal fetches + fetches += 1 + await asyncio.sleep(0.02) + return token_response("shared") + + clock = FakeClock() + manager = AsyncTokenManager( + client_id="id", + client_secret="secret", + token_url="https://api.neo4j.io/oauth/token", + user_agent="ua", + http=_http(SlowTokenTransport(), clock), + logger=logging.getLogger("test"), + ) + headers = await asyncio.gather( + *(manager.authorization_header(deadline=clock.now + 30) for _ in range(10)) + ) + assert fetches == 1 + assert headers == ["Bearer shared"] * 10 + + +async def test_token_refreshed_after_401() -> None: + transport = FakeAsyncTransport( + [ + token_response("old"), + json_response(401, {"errors": [{"message": "expired"}]}), + token_response("new"), + json_response(200, {"data": []}), + ] + ) + client = aura.AsyncAuraClient(client_id="id", client_secret="secret", transport=transport) + with pytest.raises(aura.AuthenticationError): + await client.tenants.list() + assert await client.tenants.list() == [] + assert [r.headers["Authorization"] for r in transport.api_requests] == [ + "Bearer old", + "Bearer new", + ] + + +async def test_rejected_credentials() -> None: + transport = FakeAsyncTransport([json_response(401, {"error": "access_denied"})]) + client = aura.AsyncAuraClient(client_id="id", client_secret="secret", transport=transport) + with pytest.raises(aura.AuthenticationError, match="access_denied"): + await client.instances.list() + + +async def test_async_context_manager_closes_owned_transport( + monkeypatch: pytest.MonkeyPatch, +) -> None: + closed: list[bool] = [] + + async def fake_aclose(self: AsyncHttpxTransport) -> None: + closed.append(True) + + monkeypatch.setattr(AsyncHttpxTransport, "aclose", fake_aclose) + async with aura.AsyncAuraClient(client_id="id", client_secret="secret") as client: + assert isinstance(client._transport, AsyncHttpxTransport) + await client.aclose() # idempotent + assert closed == [True] + + +async def test_does_not_close_caller_transport() -> None: + transport = FakeAsyncTransport() + async with aura.AsyncAuraClient(client_id="id", client_secret="secret", transport=transport): + pass + assert transport.closed is False + + +def test_transport_kinds_are_not_interchangeable() -> None: + with pytest.raises(aura.AuraConfigurationError, match="use AuraClient"): + aura.AsyncAuraClient(client_id="id", client_secret="s", transport=FakeTransport()) # type: ignore[arg-type] + with pytest.raises(aura.AuraConfigurationError, match="use AsyncAuraClient"): + aura.AuraClient(client_id="id", client_secret="s", transport=FakeAsyncTransport()) # type: ignore[arg-type] + with pytest.raises(aura.AuraConfigurationError, match="use AsyncAuraClient"): + aura.AuraClient(client_id="id", client_secret="s", transport=AsyncHttpxTransport()) # type: ignore[arg-type] + + +def test_async_client_validates_options_like_sync() -> None: + with pytest.raises(aura.AuraConfigurationError, match="HTTPS"): + aura.AsyncAuraClient(client_id="id", client_secret="s", base_url="http://x") + with pytest.raises(aura.AuraConfigurationError, match="logger"): + aura.AsyncAuraClient(client_id="id", client_secret="s", logger="x") # type: ignore[arg-type] + + +def test_async_from_env_and_repr(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("AURA_CLIENT_ID", "env-id") + monkeypatch.setenv("AURA_CLIENT_SECRET", "env-secret") + client = aura.AsyncAuraClient.from_env(transport=FakeAsyncTransport(), timeout=5) + assert repr(client) == "AsyncAuraClient(base_url='https://api.neo4j.io')" + assert client.base_url == "https://api.neo4j.io" + assert "env-secret" not in repr(client) + monkeypatch.delenv("AURA_CLIENT_SECRET") + with pytest.raises(aura.AuraConfigurationError, match="must both be set"): + aura.AsyncAuraClient.from_env() + + +async def test_token_is_reused_across_calls() -> None: + transport = FakeAsyncTransport( + [token_response("tok"), json_response(200, {"data": []}), json_response(200, {"data": []})] + ) + client = aura.AsyncAuraClient(client_id="id", client_secret="secret", transport=transport) + await client.tenants.list() + await client.tenants.list() + assert [r.url.endswith("/oauth/token") for r in transport.requests] == [True, False, False] diff --git a/tests/unit/test_async_parity.py b/tests/unit/test_async_parity.py new file mode 100644 index 0000000..f3f4a74 --- /dev/null +++ b/tests/unit/test_async_parity.py @@ -0,0 +1,293 @@ +"""The async client must behave exactly like the sync client. + +Every public method of every service runs through both clients against the same canned +responses. The test asserts that the requests on the wire and the parsed results are identical, +and that the method signatures match. Adding a method without a case here fails the coverage +test. +""" + +from __future__ import annotations + +import dataclasses +import datetime as dt +import inspect +from collections.abc import Callable +from typing import Any + +import pytest + +import aura_python_sdk as aura +from aura_python_sdk import services +from aura_python_sdk._transport import HttpResponse +from tests.fakes import FakeAsyncTransport, FakeTransport, json_response, token_response +from tests.unit.conftest import ( + INSTANCE, + INSTANCE_ID, + OTHER_INSTANCE_ID, + SESSION, + SNAPSHOT_ID, + TENANT_ID, +) + +SERVICE_PAIRS = [ + (services.TenantService, services.AsyncTenantService), + (services.InstanceService, services.AsyncInstanceService), + (services.SnapshotService, services.AsyncSnapshotService), + (services.CMEKService, services.AsyncCMEKService), + (services.GDSSessionService, services.AsyncGDSSessionService), + (services.PrometheusService, services.AsyncPrometheusService), +] + +CONFIG = aura.InstanceConfig( + name="Instance01", + tenant_id=TENANT_ID, + cloud_provider=aura.CloudProvider.GCP, + region="europe-west1", + type=aura.InstanceType.ENTERPRISE_DB, + version="5", + memory="8GB", +) +CREATED = {**INSTANCE, "username": "neo4j", "password": "pw"} +SNAPSHOT = {"instance_id": INSTANCE_ID, "snapshot_id": SNAPSHOT_ID, "status": "Completed"} +KEY = { + "id": "8c764aed-8eb3-4a1c-92f6-e4ef0c7a6ed9", + "name": "Key01", + "tenant_id": TENANT_ID, + "cloud_provider": "aws", + "region": "us-west-2", + "instance_type": "enterprise-db", + "key_id": "arn:aws:kms:us-west-2:1:key/abc", + "status": "pending", +} +METRICS_URL = "https://customer-metrics-api.neo4j.io/api/v1/p/2f49c2b3/metrics" +METRICS = HttpResponse( + 200, + body=b"# TYPE neo4j_aura_cpu_usage gauge\nneo4j_aura_cpu_usage 3\n" + b"# TYPE neo4j_aura_cpu_limit gauge\nneo4j_aura_cpu_limit 4\n", +) + + +def data(payload: object, status: int = 200) -> HttpResponse: + return json_response(status, {"data": payload}) + + +# (service, method, args, kwargs, replies) +CASES: list[tuple[str, str, tuple[Any, ...], dict[str, Any], list[HttpResponse]]] = [ + ("tenants", "list", (), {}, [data([{"id": TENANT_ID, "name": "t"}])]), + ("tenants", "get", (TENANT_ID,), {}, [data({"id": TENANT_ID, "name": "t"})]), + ("tenants", "get_metrics_integration", (TENANT_ID,), {}, [data({"endpoint": METRICS_URL})]), + ("instances", "list", (TENANT_ID,), {}, [data([])]), + ("instances", "get", (INSTANCE_ID,), {}, [data(INSTANCE)]), + ("instances", "create", (CONFIG,), {}, [data(CREATED, 202)]), + ("instances", "create_from_instance", (OTHER_INSTANCE_ID, CONFIG), {}, [data(CREATED, 202)]), + ( + "instances", + "create_from_snapshot", + (OTHER_INSTANCE_ID, SNAPSHOT_ID, CONFIG), + {}, + [data(CREATED, 202)], + ), + ( + "instances", + "update", + (INSTANCE_ID,), + {"name": "Renamed", "storage": "32GB", "vector_optimized": True, "secondaries_count": 1}, + [data(INSTANCE, 202)], + ), + ( + "instances", + "estimate_size", + (), + {"node_count": 10, "relationship_count": 20, "algorithm_categories": ["pathfinding"]}, + [ + data( + { + "did_exceed_maximum": False, + "min_required_memory": "1GB", + "recommended_size": "2GB", + } + ) + ], + ), + ( + "instances", + "upgrade", + (INSTANCE_ID,), + {"memory": "16GB", "storage": "32GB"}, + [data(INSTANCE)], + ), + ("instances", "delete", (INSTANCE_ID,), {}, [data(INSTANCE, 202)]), + ("instances", "pause", (INSTANCE_ID,), {}, [data(INSTANCE, 202)]), + ("instances", "resume", (INSTANCE_ID,), {}, [data(INSTANCE, 202)]), + ( + "instances", + "overwrite_from_instance", + (INSTANCE_ID, OTHER_INSTANCE_ID), + {}, + [data(INSTANCE, 202)], + ), + ("instances", "overwrite_from_snapshot", (INSTANCE_ID, SNAPSHOT_ID), {}, [data(INSTANCE, 202)]), + ("snapshots", "list", (INSTANCE_ID, dt.date(2026, 1, 2)), {}, [data([SNAPSHOT])]), + ("snapshots", "get", (INSTANCE_ID, SNAPSHOT_ID), {}, [data(SNAPSHOT)]), + ("snapshots", "create", (INSTANCE_ID,), {}, [data({"snapshot_id": SNAPSHOT_ID}, 202)]), + ("snapshots", "restore", (INSTANCE_ID, SNAPSHOT_ID), {}, [data(INSTANCE, 202)]), + ("cmek", "list", (TENANT_ID,), {}, [data([])]), + ("cmek", "get", (KEY["id"],), {}, [data(KEY)]), + ( + "cmek", + "create", + (), + { + "name": "Key01", + "key_id": KEY["key_id"], + "tenant_id": TENANT_ID, + "cloud_provider": "aws", + "region": "us-west-2", + "instance_type": "enterprise-db", + }, + [data(KEY, 202)], + ), + ("cmek", "delete", (KEY["id"],), {}, [HttpResponse(204)]), + ( + "graph_analytics", + "list", + (), + {"tenant_id": TENANT_ID, "instance_id": INSTANCE_ID, "organization_id": "org"}, + [data([SESSION])], + ), + ( + "graph_analytics", + "estimate_size", + (), + {"node_count": 1, "relationship_count": 2, "node_label_count": 3}, + [data({"estimated_memory": "1GB", "recommended_size": "2GB"})], + ), + ( + "graph_analytics", + "create", + (aura.GDSSessionConfig(name="s", memory="8GB", ttl="1h"),), + {}, + [data(SESSION, 202)], + ), + ("graph_analytics", "get", (SESSION["id"],), {}, [data(SESSION)]), + ("graph_analytics", "delete", (SESSION["id"],), {}, [data({"id": SESSION["id"]}, 202)]), + ("prometheus", "fetch_raw_metrics", (METRICS_URL,), {}, [METRICS]), + ("prometheus", "get_instance_health", (INSTANCE_ID, METRICS_URL), {}, [METRICS]), +] + +# Methods without I/O. They are plain methods on both clients, and checked separately. +NO_IO_METHODS = {("prometheus", "get_metric_value")} + + +def _wire(transport: FakeTransport) -> list[tuple[str, str, dict[str, str], bytes | None]]: + return [(r.method, r.url, dict(r.headers), r.body) for r in transport.requests] + + +def _comparable(result: object) -> object: + if isinstance(result, aura.InstanceHealth): + return dataclasses.replace(result, timestamp=dt.datetime(2000, 1, 1, tzinfo=dt.UTC)) + return result + + +def _public_methods(cls: type) -> dict[str, Callable[..., Any]]: + return { + name: member + for name, member in inspect.getmembers(cls, inspect.isfunction) + if not name.startswith("_") + } + + +@pytest.mark.parametrize(("sync_cls", "async_cls"), SERVICE_PAIRS, ids=lambda c: c.__name__) +def test_signatures_match(sync_cls: type, async_cls: type) -> None: + sync_methods = _public_methods(sync_cls) + async_methods = _public_methods(async_cls) + assert sync_methods.keys() == async_methods.keys() + for name, sync_method in sync_methods.items(): + async_method = async_methods[name] + assert inspect.signature(sync_method) == inspect.signature(async_method), name + assert inspect.iscoroutinefunction(async_method) != ( + (sync_cls.__name__, name) in {("PrometheusService", "get_metric_value")} + ), f"{async_cls.__name__}.{name} should be async" + assert not inspect.iscoroutinefunction(sync_method), name + + +def test_every_method_has_a_parity_case() -> None: + client = aura.AuraClient(client_id="id", client_secret="secret", transport=FakeTransport()) + attribute_for = { + type(getattr(client, attr)): attr + for attr in ("tenants", "instances", "snapshots", "cmek", "graph_analytics", "prometheus") + } + expected = { + (attribute_for[sync_cls], name) + for sync_cls, _ in SERVICE_PAIRS + for name in _public_methods(sync_cls) + } + covered = {(service, method) for service, method, *_ in CASES} | NO_IO_METHODS + assert expected == covered + + +@pytest.mark.anyio +@pytest.mark.parametrize( + ("service", "method", "args", "kwargs", "replies"), + CASES, + ids=[f"{c[0]}.{c[1]}" for c in CASES], +) +async def test_async_matches_sync( + service: str, + method: str, + args: tuple[Any, ...], + kwargs: dict[str, Any], + replies: list[HttpResponse], +) -> None: + sync_transport = FakeTransport([token_response(), *replies]) + sync_client = aura.AuraClient(client_id="id", client_secret="secret", transport=sync_transport) + sync_result = getattr(getattr(sync_client, service), method)(*args, **kwargs) + + async_transport = FakeAsyncTransport([token_response(), *replies]) + async_client = aura.AsyncAuraClient( + client_id="id", client_secret="secret", transport=async_transport + ) + async_result = await getattr(getattr(async_client, service), method)(*args, **kwargs) + + assert _wire(async_transport) == _wire(sync_transport) + assert _comparable(async_result) == _comparable(sync_result) + + +@pytest.mark.anyio +@pytest.mark.parametrize( + ("service", "method", "args", "kwargs"), + [ + ("instances", "get", ("bad",), {}), + ("instances", "update", (INSTANCE_ID,), {}), + ("instances", "upgrade", (INSTANCE_ID,), {"memory": "16GB"}), + ("snapshots", "list", (INSTANCE_ID, "2026-01-02"), {}), + ("cmek", "get", ("",), {}), + ("graph_analytics", "list", (), {"tenant_id": "bad"}), + ("prometheus", "fetch_raw_metrics", ("https://evil.example.com/metrics",), {}), + ("prometheus", "get_instance_health", ("bad", METRICS_URL), {}), + ], +) +async def test_async_validates_like_sync( + service: str, method: str, args: tuple[Any, ...], kwargs: dict[str, Any] +) -> None: + sync_client = aura.AuraClient(client_id="id", client_secret="secret", transport=FakeTransport()) + with pytest.raises(aura.AuraValidationError) as sync_error: + getattr(getattr(sync_client, service), method)(*args, **kwargs) + + async_transport = FakeAsyncTransport() + async_client = aura.AsyncAuraClient( + client_id="id", client_secret="secret", transport=async_transport + ) + with pytest.raises(aura.AuraValidationError) as async_error: + await getattr(getattr(async_client, service), method)(*args, **kwargs) + + assert str(async_error.value) == str(sync_error.value) + assert async_transport.requests == [] + + +def test_get_metric_value_is_shared() -> None: + client = aura.AsyncAuraClient( + client_id="id", client_secret="secret", transport=FakeAsyncTransport() + ) + metrics = aura.PrometheusMetrics(metrics={"m": (aura.PrometheusMetric(name="m", value=2.0),)}) + assert client.prometheus.get_metric_value(metrics, "m") == 2.0 From 8d7381fe536a126a65dbd372737e7a5fc6e40d17 Mon Sep 17 00:00:00 2001 From: Jonathan Giffard Date: Tue, 29 Sep 2026 21:33:48 +0100 Subject: [PATCH 8/8] Tolerate null connection_url and shorten SDK error tracebacks The first live run showed GET /instances/{id} returning connection_url: null for some instances, which the spec marks required; Instance.connection_url is now optional. SDK errors now print under their public name and their tracebacks stop at the public method instead of listing internal frames. Live CMEK and session tests skip on 403 instead of failing. Co-Authored-By: Claude Opus 5.5 --- CHANGELOG.md | 13 +++++ PLAN.md | 4 ++ examples/get_instance_details.py | 2 +- src/aura_python_sdk/_errors.py | 24 +++++++++ src/aura_python_sdk/models/instances.py | 4 +- src/aura_python_sdk/services/_base.py | 37 ++++++++++--- tests/integration/test_live.py | 19 +++++-- tests/unit/test_error_presentation.py | 70 +++++++++++++++++++++++++ tests/unit/test_models.py | 16 ++++++ 9 files changed, 175 insertions(+), 14 deletions(-) create mode 100644 tests/unit/test_error_presentation.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 8c0be9f..09bd952 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -26,3 +26,16 @@ the `## vX.Y.Z` section that matches the pushed tag as the GitHub release notes. have reached the server. - A stdlib Prometheus text-format parser whose output matches the Go SDK, and `get_instance_health` with the Go SDK's thresholds. + +### Fixed + +- `Instance.connection_url` is now optional. The live API returns `null` for some instances, + although the spec marks the field as required, and that made `instances.get()` fail. + +### Changed + +- SDK errors now report their public name (for example `aura_python_sdk.NotFoundError`), and + their tracebacks stop at the public method you called instead of listing the SDK's internal + frames. Unexpected exceptions still show a full traceback. +- The live integration tests skip, instead of failing, when the credentials lack permission for + an endpoint (HTTP 403). diff --git a/PLAN.md b/PLAN.md index a2908c4..069df5f 100644 --- a/PLAN.md +++ b/PLAN.md @@ -336,6 +336,10 @@ packages. parity; tolerant parsing makes this harmless. - **Snapshot ID format**: resolved. Snapshot IDs are UUIDs, and the spec's list example (`snapshot_id: '2023-01-20T13:44:42Z'`) is wrong. We keep Go's UUID validation. +- **`connection_url` can be null**: the first live run showed that `GET /instances/{id}` + returns `connection_url: null` for some instances, although the spec marks it as required. + `Instance.connection_url` is now optional. Run the live tests again after spec updates, to + catch fields that are required in the spec but missing in practice. - **Required fields on responses**: models follow the spec's `required` lists, with two exceptions. Instance `storage` is optional because it isn't returned for Free instances. GDS session `status` is optional because the spec's 202 example returns `null`. A missing required diff --git a/examples/get_instance_details.py b/examples/get_instance_details.py index 58ac7aa..2b20982 100644 --- a/examples/get_instance_details.py +++ b/examples/get_instance_details.py @@ -30,7 +30,7 @@ def main() -> int: print(f"Tier: {instance.type}") print(f"Memory: {instance.memory}") print(f"Storage: {instance.storage or 'n/a'}") - print(f"Connection URL: {instance.connection_url}") + print(f"Connection URL: {instance.connection_url or 'n/a'}") return 0 diff --git a/src/aura_python_sdk/_errors.py b/src/aura_python_sdk/_errors.py index 7df56bd..9a083d5 100644 --- a/src/aura_python_sdk/_errors.py +++ b/src/aura_python_sdk/_errors.py @@ -246,3 +246,27 @@ def _parse_retry_after(value: str | None) -> float | None: except (TypeError, ValueError): return None return max(0.0, parsed.timestamp() - time.time()) + + +# Report the public import path in tracebacks and reprs: "aura_python_sdk.NotFoundError", not +# "aura_python_sdk._errors.NotFoundError". Every class listed here is exported from the package. +for _public in ( + AuraError, + AuraConfigurationError, + AuraValidationError, + AuraConnectionError, + AuraTimeoutError, + AuraResponseError, + MetricNotFoundError, + ErrorDetail, + AuraAPIError, + BadRequestError, + AuthenticationError, + PermissionDeniedError, + NotFoundError, + ConflictError, + RateLimitError, + ServerError, +): + _public.__module__ = "aura_python_sdk" +del _public diff --git a/src/aura_python_sdk/models/instances.py b/src/aura_python_sdk/models/instances.py index f940115..7c8b513 100644 --- a/src/aura_python_sdk/models/instances.py +++ b/src/aura_python_sdk/models/instances.py @@ -51,6 +51,7 @@ class InstanceSummary: class Instance: """Full details of an instance. + ``connection_url`` can be ``None`` (the live API sends null for some instances). ``storage`` is not returned for AuraDB Free. ``graph_nodes`` and ``graph_relationships`` are returned only for Free instances. ``secondaries_count`` is returned only for Virtual Dedicated Cloud, and ``cdc_enrichment_mode`` only for Virtual Dedicated Cloud and Business @@ -62,7 +63,8 @@ class Instance: status: InstanceStatus | str tenant_id: str cloud_provider: CloudProvider | str - connection_url: str + # Required by the spec, but the live API returns null for some instances. + connection_url: str | None = None region: str type: InstanceType | str memory: str diff --git a/src/aura_python_sdk/services/_base.py b/src/aura_python_sdk/services/_base.py index 7b6663b..360d714 100644 --- a/src/aura_python_sdk/services/_base.py +++ b/src/aura_python_sdk/services/_base.py @@ -3,6 +3,7 @@ import logging from typing import TypeVar +from aura_python_sdk._errors import AuraError from aura_python_sdk._internal._call import Call from aura_python_sdk._internal._request import AsyncRequestService, RequestService @@ -14,6 +15,16 @@ ORGANIZATION_ID_PARAM = "organizationId" +def _without_internal_frames(exc: AuraError) -> AuraError: + """Drop the SDK's internal frames from an SDK error's traceback. + + The error message already says what went wrong, so the traceback starts at the service + method the caller used. Unexpected exceptions (bugs) are not caught, and keep their full + traceback. + """ + return exc.with_traceback(None) + + class Service: """Base for the sync services on :class:`AuraClient`: runs each operation's ``Call``.""" @@ -22,11 +33,16 @@ def __init__(self, api: RequestService, logger: logging.Logger) -> None: self._logger = logger def _run(self, call: Call[T]) -> T: + __tracebackhide__ = True # pytest: leave this frame out of failure reports self._logger.debug(call.describe, extra=dict(call.context)) - response = self._api.request( - call.method, call.path, params=call.params, json_body=call.json_body - ) - result = call.parse(response) + try: + response = self._api.request( + call.method, call.path, params=call.params, json_body=call.json_body + ) + result = call.parse(response) + except AuraError as exc: + # Re-raising the same object keeps its __cause__ (e.g. the network error). + raise _without_internal_frames(exc) # noqa: B904 if call.done: self._logger.info(call.done, extra=dict(call.context)) return result @@ -40,11 +56,16 @@ def __init__(self, api: AsyncRequestService, logger: logging.Logger) -> None: self._logger = logger async def _run(self, call: Call[T]) -> T: + __tracebackhide__ = True # pytest: leave this frame out of failure reports self._logger.debug(call.describe, extra=dict(call.context)) - response = await self._api.request( - call.method, call.path, params=call.params, json_body=call.json_body - ) - result = call.parse(response) + try: + response = await self._api.request( + call.method, call.path, params=call.params, json_body=call.json_body + ) + result = call.parse(response) + except AuraError as exc: + # Re-raising the same object keeps its __cause__ (e.g. the network error). + raise _without_internal_frames(exc) # noqa: B904 if call.done: self._logger.info(call.done, extra=dict(call.context)) return result diff --git a/tests/integration/test_live.py b/tests/integration/test_live.py index f5e59ca..600a7f2 100644 --- a/tests/integration/test_live.py +++ b/tests/integration/test_live.py @@ -12,7 +12,7 @@ import os import time -from collections.abc import Iterator +from collections.abc import Callable, Iterator import pytest @@ -64,9 +64,20 @@ def test_snapshots_for_first_instance(client: aura.AuraClient) -> None: assert snapshot.instance_id == instances[0].id -def test_cmek_and_sessions_list(client: aura.AuraClient) -> None: - assert isinstance(client.cmek.list(), list) - assert isinstance(client.graph_analytics.list(), list) +def _skip_if_forbidden(call: Callable[[], object]) -> object: + """Run ``call``, but skip the test if these credentials lack permission for it.""" + try: + return call() + except aura.PermissionDeniedError as err: + pytest.skip(f"credentials lack permission: {err.message}") + + +def test_cmek_list(client: aura.AuraClient) -> None: + assert isinstance(_skip_if_forbidden(client.cmek.list), list) + + +def test_sessions_list(client: aura.AuraClient) -> None: + assert isinstance(_skip_if_forbidden(client.graph_analytics.list), list) def test_unknown_instance_is_not_found(client: aura.AuraClient) -> None: diff --git a/tests/unit/test_error_presentation.py b/tests/unit/test_error_presentation.py new file mode 100644 index 0000000..776c3fb --- /dev/null +++ b/tests/unit/test_error_presentation.py @@ -0,0 +1,70 @@ +"""SDK errors should read cleanly: a public class name and a short traceback.""" + +import traceback + +import pytest + +import aura_python_sdk as aura +from tests.fakes import FakeAsyncTransport, FakeTransport, json_response, token_response + +FORBIDDEN = json_response( + 403, {"errors": [{"message": "Insufficient permissions", "reason": "unauthorized"}]} +) + + +def _frames(exc: BaseException) -> list[str]: + return [frame.filename for frame in traceback.extract_tb(exc.__traceback__)] + + +def test_public_exceptions_report_the_package_path() -> None: + for name in aura.__all__: + obj = getattr(aura, name) + if isinstance(obj, type) and issubclass(obj, BaseException): + assert obj.__module__ == "aura_python_sdk", name + line = traceback.format_exception_only(aura.NotFoundError(404, "Not Found"))[-1] + assert line.startswith("aura_python_sdk.NotFoundError: API error (status 404)") + + +def test_api_error_traceback_skips_internal_frames() -> None: + transport = FakeTransport([token_response(), FORBIDDEN]) + client = aura.AuraClient(client_id="id", client_secret="secret", transport=transport) + with pytest.raises(aura.PermissionDeniedError) as info: + client.cmek.list() + + frames = _frames(info.value) + assert not any("/_internal/" in f for f in frames), frames + # What remains: this test, the public service method, and the re-raise in _run. + assert any(f.endswith("services/cmek.py") for f in frames) + + +@pytest.mark.anyio +async def test_async_api_error_traceback_skips_internal_frames() -> None: + transport = FakeAsyncTransport([token_response(), FORBIDDEN]) + client = aura.AsyncAuraClient(client_id="id", client_secret="secret", transport=transport) + with pytest.raises(aura.PermissionDeniedError) as info: + await client.cmek.list() + assert not any("/_internal/" in f for f in _frames(info.value)) + + +def test_network_error_keeps_its_cause() -> None: + cause = OSError("connection reset") + error = aura.AuraConnectionError("request failed", request_sent=False) + error.__cause__ = cause + transport = FakeTransport([token_response(), error]) + client = aura.AuraClient( + client_id="id", client_secret="secret", transport=transport, max_retries=0 + ) + with pytest.raises(aura.AuraConnectionError) as info: + client.tenants.list() + assert info.value.__cause__ is cause + + +def test_unexpected_errors_keep_their_full_traceback() -> None: + def explode(request: aura.HttpRequest) -> aura.HttpResponse: + raise RuntimeError("bug in a custom transport") + + transport = FakeTransport([token_response(), explode]) + client = aura.AuraClient(client_id="id", client_secret="secret", transport=transport) + with pytest.raises(RuntimeError) as info: + client.tenants.list() + assert any("/_internal/" in f for f in _frames(info.value)) diff --git a/tests/unit/test_models.py b/tests/unit/test_models.py index 1532f3b..6a8a7e3 100644 --- a/tests/unit/test_models.py +++ b/tests/unit/test_models.py @@ -108,3 +108,19 @@ def test_all_models_exported_at_top_level() -> None: for name in models.__all__: assert getattr(aura, name) is getattr(models, name) assert name in aura.__all__ + + +@pytest.mark.parametrize("payload", [{"connection_url": None}, {}]) +def test_instance_connection_url_may_be_null_or_missing(payload: dict[str, object]) -> None: + # The spec marks it required, but the live API returns null for some instances. + base = { + "id": "abcd1234", + "name": "x", + "status": "creating", + "tenant_id": "t", + "cloud_provider": "gcp", + "region": "europe-west1", + "type": "enterprise-db", + "memory": "8GB", + } + assert from_json(models.Instance, {**base, **payload}).connection_url is None