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..09bd952 --- /dev/null +++ b/CHANGELOG.md @@ -0,0 +1,41 @@ +# 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 + +- `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 + `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. + +### 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 6de406f..069df5f 100644 --- a/PLAN.md +++ b/PLAN.md @@ -185,6 +185,67 @@ 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. + +### 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. + +### 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`. + +### 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 ``` @@ -207,7 +268,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 @@ -220,13 +281,14 @@ 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. ## 5. Phases +**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), `TokenManager`, `RequestService`, and `AuraClient` with no services yet. Unit-tested to Go's @@ -235,14 +297,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. @@ -259,9 +321,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 @@ -270,3 +334,13 @@ 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. +- **`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 + field raises `AuraResponseError` and names the field. diff --git a/README.md b/README.md index 6037acf..55e495b 100644 --- a/README.md +++ b/README.md @@ -1,11 +1,369 @@ # 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). +- 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. +- 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) +- [Async](#async) +- [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. + +## 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 +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=`. 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. + +```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:` | +| 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=` | +| `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 +371,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/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/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..2b20982 --- /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 or 'n/a'}") + 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 a487ca8..e89ae48 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -25,11 +25,8 @@ classifiers = [ ] dependencies = ["httpx>=0.27,<1"] -[project.optional-dependencies] -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] @@ -37,7 +34,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] @@ -57,16 +56,16 @@ 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" 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/src/aura_python_sdk/__init__.py b/src/aura_python_sdk/__init__.py index ab84b59..50c26f9 100644 --- a/src/aura_python_sdk/__init__.py +++ b/src/aura_python_sdk/__init__.py @@ -7,8 +7,125 @@ 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 AsyncAuraClient, AuraClient +from aura_python_sdk._errors import ( + AuraAPIError, + AuraConfigurationError, + AuraConnectionError, + AuraError, + AuraResponseError, + AuraTimeoutError, + AuraValidationError, + AuthenticationError, + BadRequestError, + ConflictError, + ErrorDetail, + MetricNotFoundError, + NotFoundError, + PermissionDeniedError, + RateLimitError, + ServerError, +) +from aura_python_sdk._transport import AsyncHttpTransport, HttpRequest, HttpResponse, HttpTransport from aura_python_sdk._version import __version__ +from aura_python_sdk.models import ( + CDCEnrichmentMode, + CloudProvider, + ConnectionMetrics, + CreatedInstance, + CreatedSnapshot, + CustomerManagedKey, + CustomerManagedKeySummary, + DeletedGDSSession, + GDSSession, + GDSSessionConfig, + GDSSessionSizeEstimate, + GDSSessionStatus, + HealthStatus, + Instance, + InstanceConfig, + InstanceConfiguration, + InstanceHealth, + InstanceSizeEstimate, + InstanceStatus, + InstanceSummary, + InstanceType, + MetricsIntegration, + PrometheusMetric, + PrometheusMetrics, + QueryMetrics, + ResourceMetrics, + Snapshot, + SnapshotProfile, + SnapshotStatus, + StorageMetrics, + Tenant, + TenantSummary, +) + +# Library convention: emit nothing unless the application configures logging. +logging.getLogger(__name__).addHandler(logging.NullHandler()) -__all__ = ["__version__"] +__all__ = [ + "AsyncAuraClient", + "AsyncHttpTransport", + "AuraAPIError", + "AuraClient", + "AuraConfigurationError", + "AuraConnectionError", + "AuraError", + "AuraResponseError", + "AuraTimeoutError", + "AuraValidationError", + "AuthenticationError", + "BadRequestError", + "CDCEnrichmentMode", + "CloudProvider", + "ConflictError", + "ConnectionMetrics", + "CreatedInstance", + "CreatedSnapshot", + "CustomerManagedKey", + "CustomerManagedKeySummary", + "DeletedGDSSession", + "ErrorDetail", + "GDSSession", + "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 new file mode 100644 index 0000000..15a67af --- /dev/null +++ b/src/aura_python_sdk/_client.py @@ -0,0 +1,333 @@ +"""The AuraClient and AsyncAuraClient entry points (Go: client.go).""" + +from __future__ import annotations + +import inspect +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 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, + PrometheusService, + SnapshotService, + TenantService, +) + +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" + + +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. + + 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``, plus ``prometheus`` for metrics endpoints. + + 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, 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. + 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) 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 + + 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.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.prometheus = PrometheusService( + self._api, + self._logger.getChild("prometheus"), + allow_untrusted_urls=self._config.allow_insecure_base_url, + ) + + 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, 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 + + 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})" + + +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/_config.py b/src/aura_python_sdk/_config.py new file mode 100644 index 0000000..c9a3b93 --- /dev/null +++ b/src/aura_python_sdk/_config.py @@ -0,0 +1,128 @@ +"""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 + allow_insecure_base_url: bool + 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), + 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), + 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..9a083d5 --- /dev/null +++ b/src/aura_python_sdk/_errors.py @@ -0,0 +1,272 @@ +"""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.""" + + +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.""" + + 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()) + + +# 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/_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..aac2ab6 --- /dev/null +++ b/src/aura_python_sdk/_internal/_auth.py @@ -0,0 +1,201 @@ +"""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 AsyncHttpService, HttpService +from aura_python_sdk._transport import HttpResponse + +# 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) + + @property + def header(self) -> str: + return f"{self.token_type} {self.access_token}" + + +class _TokenSource: + def __init__( + self, + *, + client_id: str, + client_secret: str, + token_url: str, + user_agent: str, + clock: Callable[[], float], + logger: logging.Logger, + ) -> None: + credentials = f"{client_id}:{client_secret}".encode() + 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}) + 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}) + self.token = _Token( + token_type="Bearer", # noqa: S106 - the OAuth scheme name, not a secret + access_token=access_token, + 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 new file mode 100644 index 0000000..1c3274a --- /dev/null +++ b/src/aura_python_sdk/_internal/_request.py @@ -0,0 +1,206 @@ +"""Authenticated Aura API requests (Go: internal/api RequestService).""" + +from __future__ import annotations + +import json +import logging +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 AsyncTokenManager, TokenManager +from aura_python_sdk._internal.http._service import AsyncHttpService, HttpService +from aura_python_sdk._transport import HttpResponse + +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 _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``. + An absolute ``http(s)://`` URL, such as a Prometheus metrics endpoint, is used unchanged but + 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, + *, + 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._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) + + 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._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) + + +class AsyncRequestService: + """The asyncio version of :class:`RequestService`.""" + + 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/_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/_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..a7283ed --- /dev/null +++ b/src/aura_python_sdk/_internal/http/_httpx.py @@ -0,0 +1,111 @@ +"""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 + + +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 + ) + + 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: + 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 _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 new file mode 100644 index 0000000..96be5ec --- /dev/null +++ b/src/aura_python_sdk/_internal/http/_service.py @@ -0,0 +1,182 @@ +"""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 Awaitable, Callable, Mapping + +from aura_python_sdk._errors import AuraConnectionError, AuraResponseError, AuraTimeoutError +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"}) + +RETRY_WAIT_MIN = 1.0 +RETRY_WAIT_MAX = 5.0 + + +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. 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, + *, + 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._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 + + 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 = self._transport.send(request) + except AuraConnectionError as exc: + wait = self._policy.retry_wait(method, url, exc, attempt, deadline) + if wait is None: + raise + self._sleep(wait) + attempt += 1 + continue + return self._policy.check_response(method, url, response) + + +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/_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/_transport.py b/src/aura_python_sdk/_transport.py new file mode 100644 index 0000000..f9c5fe7 --- /dev/null +++ b/src/aura_python_sdk/_transport.py @@ -0,0 +1,69 @@ +"""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: ... + + +@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/_validation.py b/src/aura_python_sdk/_validation.py new file mode 100644 index 0000000..74fb84f --- /dev/null +++ b/src/aura_python_sdk/_validation.py @@ -0,0 +1,82 @@ +"""Client-side argument validation, run before any request is sent (Go: internal/utils).""" + +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_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 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/models/__init__.py b/src/aura_python_sdk/models/__init__.py new file mode 100644 index 0000000..641f721 --- /dev/null +++ b/src/aura_python_sdk/models/__init__.py @@ -0,0 +1,77 @@ +"""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.prometheus import ( + ConnectionMetrics, + HealthStatus, + InstanceHealth, + PrometheusMetric, + PrometheusMetrics, + QueryMetrics, + ResourceMetrics, + StorageMetrics, +) +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", + "ConnectionMetrics", + "CreatedInstance", + "CreatedSnapshot", + "CustomerManagedKey", + "CustomerManagedKeySummary", + "DeletedGDSSession", + "GDSSession", + "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/_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..7c8b513 --- /dev/null +++ b/src/aura_python_sdk/models/instances.py @@ -0,0 +1,130 @@ +"""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. + + ``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 + Critical. + """ + + id: str + name: str + status: InstanceStatus | str + tenant_id: str + cloud_provider: CloudProvider | 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 + 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/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/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/src/aura_python_sdk/services/__init__.py b/src/aura_python_sdk/services/__init__.py new file mode 100644 index 0000000..6aac086 --- /dev/null +++ b/src/aura_python_sdk/services/__init__.py @@ -0,0 +1,24 @@ +"""The grouped services exposed on :class:`~aura_python_sdk.AuraClient` and +:class:`~aura_python_sdk.AsyncAuraClient`.""" + +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", + "PrometheusService", + "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..360d714 --- /dev/null +++ b/src/aura_python_sdk/services/_base.py @@ -0,0 +1,71 @@ +from __future__ import annotations + +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 + +T = TypeVar("T") + +# 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" + + +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``.""" + + def __init__(self, api: RequestService, logger: logging.Logger) -> None: + self._api = api + 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)) + 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 + + +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: + __tracebackhide__ = True # pytest: leave this frame out of failure reports + self._logger.debug(call.describe, extra=dict(call.context)) + 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/src/aura_python_sdk/services/cmek.py b/src/aura_python_sdk/services/cmek.py new file mode 100644 index 0000000..8f75781 --- /dev/null +++ b/src/aura_python_sdk/services/cmek.py @@ -0,0 +1,164 @@ +"""``client.cmek`` (Go: CMEKService).""" + +from __future__ import annotations + +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 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, 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.""" + 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.""" + return self._run(_get(key_id)) + + 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. + """ + return self._run( + _create( + name=name, + key_id=key_id, + tenant_id=tenant_id, + cloud_provider=cloud_provider, + region=region, + instance_type=instance_type, + ) + ) + + def delete(self, key_id: str) -> None: + """Delete a key. The API refuses if any instance still uses it.""" + 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 new file mode 100644 index 0000000..c09171e --- /dev/null +++ b/src/aura_python_sdk/services/graph_analytics.py @@ -0,0 +1,227 @@ +"""``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._call import Call, many, one +from aura_python_sdk._internal._request import build_path +from aura_python_sdk._internal._serde import to_json +from aura_python_sdk.models.graph_analytics import ( + DeletedGDSSession, + GDSSession, + GDSSessionConfig, + GDSSessionSizeEstimate, +) +from aura_python_sdk.services._base import ( + 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.""" + + 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.""" + return self._run(_list(tenant_id, instance_id, organization_id)) + + 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``).""" + return self._run( + _estimate_size( + node_count, + relationship_count, + node_property_count, + node_label_count, + relationship_property_count, + algorithm_categories, + ) + ) + + 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``. + """ + return self._run(_create(config)) + + def get(self, session_id: str) -> GDSSession: + """Details of one session.""" + return self._run(_get(session_id)) + + def delete(self, session_id: str) -> DeletedGDSSession: + """Delete a session.""" + 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 new file mode 100644 index 0000000..2f52a66 --- /dev/null +++ b/src/aura_python_sdk/services/instances.py @@ -0,0 +1,441 @@ +"""``client.instances`` (Go: InstanceService).""" + +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._call import Call, many, one +from aura_python_sdk._internal._request import build_path +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, + CreatedInstance, + Instance, + InstanceConfig, + InstanceSizeEstimate, + InstanceSummary, +) +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): + """AuraDB and AuraDS instances.""" + + def list(self, tenant_id: str | None = None) -> builtins.list[InstanceSummary]: + """Every instance the credentials can access, optionally only those in one tenant.""" + return self._run(_list(tenant_id)) + + def get(self, instance_id: str) -> Instance: + """Full details of one instance.""" + return self._run(_get(instance_id)) + + 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. + """ + 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.""" + 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 + ) -> CreatedInstance: + """Create an instance from a snapshot. + + The snapshot must belong to ``source_instance_id`` and be exportable. + """ + return self._run(_create(config, source_instance_id, source_snapshot_id)) + + 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: + """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. + """ + 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, + ) + ) + + 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. + """ + 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 + ) -> 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. + """ + return self._run(_upgrade(instance_id, memory, storage)) + + def delete(self, instance_id: str) -> Instance: + """Start deleting an instance. This cannot be undone.""" + return self._run(_delete(instance_id)) + + def pause(self, instance_id: str) -> Instance: + """Pause a running instance.""" + return self._run(_lifecycle(instance_id, "pause")) + + def resume(self, instance_id: str) -> Instance: + """Resume a paused 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.""" + 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.""" + return self._run(_overwrite(instance_id, source_snapshot_id=source_snapshot_id)) + + +class AsyncInstanceService(AsyncService): + """Async version of :class:`InstanceService`, with the same arguments and behaviour.""" + + 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 new file mode 100644 index 0000000..eb0e0d3 --- /dev/null +++ b/src/aura_python_sdk/services/prometheus.py @@ -0,0 +1,300 @@ +"""``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._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, + HealthStatus, + InstanceHealth, + PrometheusMetrics, + QueryMetrics, + ResourceMetrics, + StorageMetrics, +) +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. + + 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.""" + return 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: + """The mean value of ``name`` across every sample whose labels match ``label_filters``. + + Raises :class:`MetricNotFoundError` if nothing matches. + """ + 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. + + Uses the same metrics, thresholds and status logic as the Go SDK. + """ + instance_id = validate.instance_id(instance_id) + call = _fetch(prometheus_url, allow_untrusted=self._allow_untrusted_urls) + return build_health(instance_id, self._run(call), self._logger) + + +class AsyncPrometheusService(AsyncService): + """Async version of :class:`PrometheusService`, with the same arguments and behaviour. + + ``get_metric_value`` does no I/O, so it is a plain (non-async) method here too. + """ + + 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 + + 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( + 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/src/aura_python_sdk/services/snapshots.py b/src/aura_python_sdk/services/snapshots.py new file mode 100644 index 0000000..8d1ef1e --- /dev/null +++ b/src/aura_python_sdk/services/snapshots.py @@ -0,0 +1,110 @@ +"""``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._call import Call, many, one +from aura_python_sdk._internal._request import build_path +from aura_python_sdk.models.instances import Instance +from aura_python_sdk.models.snapshots import CreatedSnapshot, Snapshot +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): + """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.""" + return self._run(_list(instance_id, date)) + + def get(self, instance_id: str, snapshot_id: str) -> Snapshot: + """Details of one snapshot.""" + return self._run(_get(instance_id, snapshot_id)) + + def create(self, instance_id: str) -> CreatedSnapshot: + """Start an on-demand snapshot.""" + 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.""" + 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 new file mode 100644 index 0000000..4e0bb2f --- /dev/null +++ b/src/aura_python_sdk/services/tenants.py @@ -0,0 +1,74 @@ +"""``client.tenants`` (Go: TenantService).""" + +from __future__ import annotations + +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.models.tenants import MetricsIntegration, Tenant, TenantSummary +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): + """Tenants (shown as projects in the Aura Console).""" + + def list(self) -> builtins.list[TenantSummary]: + """Every tenant the credentials can access.""" + return self._run(_list()) + + def get(self, tenant_id: str) -> Tenant: + """A tenant and the instance configurations it can create.""" + return self._run(_get(tenant_id)) + + def get_metrics_integration(self, tenant_id: str) -> MetricsIntegration: + """The project-level Prometheus metrics endpoint (Go: ``GetMetrics``).""" + 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/__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..2b5591a --- /dev/null +++ b/tests/blackbox/conftest.py @@ -0,0 +1,150 @@ +"""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 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 + + 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/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 new file mode 100644 index 0000000..8ab04f4 --- /dev/null +++ b/tests/fakes.py @@ -0,0 +1,95 @@ +"""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")] + + +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/__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..600a7f2 --- /dev/null +++ b/tests/integration/test_live.py @@ -0,0 +1,141 @@ +"""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 Callable, 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 _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: + 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) + + +@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/__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..640c369 --- /dev/null +++ b/tests/transport/test_httpx_transport.py @@ -0,0 +1,202 @@ +"""HttpxTransport against httpx.MockTransport (tests may import httpx; src may not).""" + +import ssl +from collections.abc import AsyncIterator, 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() + + +# --- 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/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_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 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_cmek_service.py b/tests/unit/test_cmek_service.py new file mode 100644 index 0000000..133af13 --- /dev/null +++ b/tests/unit/test_cmek_service.py @@ -0,0 +1,136 @@ +from typing import Any + +import pytest + +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} + + +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() + + +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_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_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_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_graph_analytics_service.py b/tests/unit/test_graph_analytics_service.py new file mode 100644 index 0000000..60797af --- /dev/null +++ b/tests/unit/test_graph_analytics_service.py @@ -0,0 +1,173 @@ +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 categories entry", + ), + ], +) +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() + + +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_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_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_instances_service.py b/tests/unit/test_instances_service.py new file mode 100644 index 0000000..0aa0391 --- /dev/null +++ b/tests/unit/test_instances_service.py @@ -0,0 +1,413 @@ +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() + + +# --- 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_models.py b/tests/unit/test_models.py new file mode 100644 index 0000000..6a8a7e3 --- /dev/null +++ b/tests/unit/test_models.py @@ -0,0 +1,126 @@ +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__ + + +@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 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/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" 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_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_spec_examples.py b/tests/unit/test_spec_examples.py new file mode 100644 index 0000000..a72ba14 --- /dev/null +++ b/tests/unit/test_spec_examples.py @@ -0,0 +1,165 @@ +"""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) + + +# 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 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) diff --git a/uv.lock b/uv.lock index 662319b..718bd1f 100644 --- a/uv.lock +++ b/uv.lock @@ -90,32 +90,27 @@ dependencies = [ { name = "httpx" }, ] -[package.optional-dependencies] -prometheus = [ - { name = "prometheus-client" }, -] - [package.dev-dependencies] dev = [ { name = "mypy" }, { name = "pytest" }, { name = "pytest-cov" }, + { name = "pyyaml" }, { name = "ruff" }, + { name = "types-pyyaml" }, ] [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 = [ { 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]] @@ -522,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" @@ -570,6 +556,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 +690,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"