diff --git a/README.md b/README.md index 902a8e6..09b4fc1 100644 --- a/README.md +++ b/README.md @@ -37,13 +37,14 @@ asyncio.run(main()) ## Resources -The client exposes three resource groups — all endpoints, parameters, and DTOs are documented in the [OpenAPI specification](https://github.com/energy-tracker/public-docs/blob/main/public-api/openapi.yml). +The client exposes four resource groups — all endpoints, parameters, and DTOs are documented in the [OpenAPI specification](https://github.com/energy-tracker/public-docs/blob/main/public-api/openapi.yml). | Resource | Methods | |---|---| | `client.devices` | `list_standard()`, `list_virtual()` | | `client.meter_readings` | `list()`, `create()`, `delete()`, `export()` | | `client.environments` | `list()`, `get()`, `create()`, `delete()`, `create_entry()`, `delete_entry()` | +| `client.calculations` | `daily_values()`, `extrapolations()` | ## Configuration @@ -52,6 +53,7 @@ client = EnergyTrackerClient( access_token="your-token", base_url="https://custom-api.example.com", # Optional timeout=30, # Optional, default: 10s + calculation_timeout=60, # Optional, calculations only ) ``` diff --git a/energy_tracker_api/__init__.py b/energy_tracker_api/__init__.py index 39b0cb9..ba06baa 100644 --- a/energy_tracker_api/__init__.py +++ b/energy_tracker_api/__init__.py @@ -11,10 +11,13 @@ NetworkError, RateLimitError, ResourceNotFoundError, + ServiceUnavailableError, TimeoutError, ValidationError, ) from .models import ( + CalculationInterval, + CalculationPointDto, CreateEnvironmentEntryDto, CreateEnvironmentRecordDto, CreateMeterReadingDto, @@ -25,6 +28,7 @@ EnvironmentRecordDto, ExportColumn, ExportMeterReadingsDto, + ExtrapolationMethod, MeterReadingDto, SortDirection, TimestampDto, @@ -34,6 +38,9 @@ __all__ = [ "EnergyTrackerClient", # Models + "CalculationInterval", + "CalculationPointDto", + "ExtrapolationMethod", "CreateMeterReadingDto", "MeterReadingDto", "ExportMeterReadingsDto", @@ -55,6 +62,7 @@ "ResourceNotFoundError", "ConflictError", "RateLimitError", + "ServiceUnavailableError", "NetworkError", "TimeoutError", ] diff --git a/energy_tracker_api/client.py b/energy_tracker_api/client.py index 9e0032b..6be83fd 100644 --- a/energy_tracker_api/client.py +++ b/energy_tracker_api/client.py @@ -1,6 +1,7 @@ """Energy Tracker API client implementation.""" import asyncio +import math from http import HTTPStatus from typing import Any, Literal from urllib.parse import urljoin @@ -15,6 +16,7 @@ NetworkError, RateLimitError, ResourceNotFoundError, + ServiceUnavailableError, TimeoutError, ValidationError, ) @@ -28,6 +30,7 @@ class EnergyTrackerClient: _base_url: str _access_token: str _timeout: aiohttp.ClientTimeout + _calculation_timeout: aiohttp.ClientTimeout _session: aiohttp.ClientSession | None def __init__( @@ -35,6 +38,8 @@ def __init__( access_token: str, base_url: str | None = None, timeout: int = 10, + *, + calculation_timeout: float = 60, ): """Initialize the Energy Tracker API client. @@ -42,19 +47,35 @@ def __init__( access_token: Bearer token for authentication. base_url: Base URL of the API (defaults to production API). timeout: Request timeout in seconds (default: 10). + calculation_timeout: Positive, finite timeout in seconds for calculations + only (default: 60; the backend allows up to 45 seconds). """ + if ( + isinstance(calculation_timeout, bool) + or not isinstance(calculation_timeout, (int, float)) + or not math.isfinite(calculation_timeout) + or calculation_timeout <= 0 + ): + raise ValueError("calculation_timeout must be a positive, finite number") url = base_url or self._DEFAULT_BASE_URL self._base_url = url.strip().rstrip("/") self._access_token = access_token self._timeout = aiohttp.ClientTimeout(total=timeout) + self._calculation_timeout = aiohttp.ClientTimeout(total=calculation_timeout) self._session = None - from .resources import DeviceResource, EnvironmentResource, MeterReadingResource + from .resources import ( + CalculationResource, + DeviceResource, + EnvironmentResource, + MeterReadingResource, + ) self.devices = DeviceResource(self) self.meter_readings = MeterReadingResource(self) self.environments = EnvironmentResource(self) + self.calculations = CalculationResource(self) async def _get_session(self) -> aiohttp.ClientSession: if self._session is None or self._session.closed: @@ -111,10 +132,13 @@ async def _make_request( try: data = await response.json() except (ValueError, aiohttp.ContentTypeError) as e: - raise EnergyTrackerAPIError("Expected a valid JSON response") from e + raise EnergyTrackerAPIError( + "Expected a valid JSON response", status_code=response.status + ) from e if not isinstance(data, (dict, list)): raise EnergyTrackerAPIError( - f"Expected a JSON object or array, got {type(data).__name__}" + f"Expected a JSON object or array, got {type(data).__name__}", + status_code=response.status, ) return data @@ -135,19 +159,29 @@ async def _make_request( message = "Bad Request" if api_message: message += f" ({'; '.join(api_message)})" - raise ValidationError(message, api_message=api_message) + raise ValidationError( + message, api_message=api_message, status_code=response.status + ) elif response.status == 401: raise AuthenticationError( - "Unauthorized: Check your access token", api_message=api_message + "Unauthorized: Check your access token", + api_message=api_message, + status_code=response.status, ) elif response.status == 403: raise ForbiddenError( - "Forbidden: Insufficient permissions", api_message=api_message + "Forbidden: Insufficient permissions", + api_message=api_message, + status_code=response.status, ) elif response.status == 404: - raise ResourceNotFoundError("Not Found", api_message=api_message) + raise ResourceNotFoundError( + "Not Found", api_message=api_message, status_code=response.status + ) elif response.status == 409: - raise ConflictError("Conflict", api_message=api_message) + raise ConflictError( + "Conflict", api_message=api_message, status_code=response.status + ) elif response.status == 429: retry_after = response.headers.get("Retry-After") retry_seconds = ( @@ -157,20 +191,34 @@ async def _make_request( if retry_seconds: message += f" - Retry after {retry_seconds} seconds" raise RateLimitError( - message, api_message=api_message, retry_after=retry_seconds + message, + api_message=api_message, + retry_after=retry_seconds, + status_code=response.status, + ) + elif response.status == HTTPStatus.SERVICE_UNAVAILABLE: + raise ServiceUnavailableError( + f"Server error: {response.status}", + api_message=api_message, + status_code=response.status, ) elif response.status >= 500: raise EnergyTrackerAPIError( - f"Server error: {response.status}", api_message=api_message + f"Server error: {response.status}", + api_message=api_message, + status_code=response.status, ) elif response.status >= 400: raise EnergyTrackerAPIError( - f"HTTP error: {response.status}", api_message=api_message + f"HTTP error: {response.status}", + api_message=api_message, + status_code=response.status, ) else: raise EnergyTrackerAPIError( f"Unexpected HTTP status: {response.status} (expected {expected_status})", api_message=api_message, + status_code=response.status, ) except asyncio.TimeoutError as e: diff --git a/energy_tracker_api/exceptions.py b/energy_tracker_api/exceptions.py index 5891e07..923adce 100644 --- a/energy_tracker_api/exceptions.py +++ b/energy_tracker_api/exceptions.py @@ -6,11 +6,19 @@ class EnergyTrackerAPIError(Exception): Attributes: api_message: List of messages from the API response body. + status_code: HTTP response status, or None for local and transport errors. """ - def __init__(self, message: str, api_message: list[str] | None = None): + def __init__( + self, + message: str, + api_message: list[str] | None = None, + *, + status_code: int | None = None, + ): super().__init__(message) self.api_message = api_message if api_message is not None else [] + self.status_code = status_code class ValidationError(EnergyTrackerAPIError): @@ -72,12 +80,21 @@ class RateLimitError(EnergyTrackerAPIError): """ def __init__( - self, message: str, api_message: list[str] | None = None, retry_after: int | None = None + self, + message: str, + api_message: list[str] | None = None, + retry_after: int | None = None, + *, + status_code: int | None = None, ): - super().__init__(message, api_message) + super().__init__(message, api_message, status_code=status_code) self.retry_after = retry_after +class ServiceUnavailableError(EnergyTrackerAPIError): + """Raised when the service is unavailable (HTTP 503), e.g. a calculation deadline expires.""" + + class NetworkError(EnergyTrackerAPIError): """Raised when a network error occurs (connection issues, DNS, etc.). diff --git a/energy_tracker_api/models/__init__.py b/energy_tracker_api/models/__init__.py index c4e2b2c..8c13141 100644 --- a/energy_tracker_api/models/__init__.py +++ b/energy_tracker_api/models/__init__.py @@ -1,5 +1,6 @@ """Data models for Energy Tracker API.""" +from .calculations import CalculationInterval, CalculationPointDto, ExtrapolationMethod from .common import TimestampDto from .devices import DeviceSummaryDto from .environments import ( @@ -19,6 +20,9 @@ ) __all__ = [ + "CalculationInterval", + "CalculationPointDto", + "ExtrapolationMethod", "TimestampDto", "DeviceSummaryDto", "CreateMeterReadingDto", diff --git a/energy_tracker_api/models/calculations.py b/energy_tracker_api/models/calculations.py new file mode 100644 index 0000000..8635864 --- /dev/null +++ b/energy_tracker_api/models/calculations.py @@ -0,0 +1,65 @@ +"""Models for daily values and calendar interval extrapolations.""" + +import math +from dataclasses import dataclass +from datetime import UTC, datetime +from enum import StrEnum + + +class CalculationInterval(StrEnum): + """Calendar interval for extrapolation; weeks begin on Monday.""" + + DAY = "day" + WEEK = "week" + MONTH = "month" + QUARTER = "quarter" + YEAR = "year" + + +class ExtrapolationMethod(StrEnum): + """Available server-side extrapolation methods.""" + + STANDARD = "standard" + + +def _number(value: object, field: str, *, duration: bool = False) -> float: + if isinstance(value, bool) or not isinstance(value, (int, float)): + raise TypeError(f"{field} must be a JSON number") + number = float(value) + if not math.isfinite(number): + raise ValueError(f"{field} must be finite") + if duration and number < 0: + raise ValueError(f"{field} must not be negative") + return number + + +@dataclass(frozen=True, slots=True) +class CalculationPointDto: + """Consumption or production for a calendar interval, in the device unit. + + Attributes: + date: Interval start in UTC. + actual_value: Value derived from recorded readings by interpolation. + actual_duration: Reading-backed duration in seconds, affected by DST. + expected_value: Estimated total for the interval; do not add actual_value. + expected_duration: Nominal duration in seconds, using 24 hours per calendar day. + """ + + date: datetime + actual_value: float + actual_duration: float + expected_value: float + expected_duration: float + + @classmethod + def _from_dict(cls, data: dict) -> CalculationPointDto: + date = datetime.fromisoformat(data["date"]) + if date.utcoffset() is None: + raise ValueError("date must include a UTC offset") + return cls( + date=date.astimezone(UTC), + actual_value=_number(data["actualValue"], "actualValue"), + actual_duration=_number(data["actualDuration"], "actualDuration", duration=True), + expected_value=_number(data["expectedValue"], "expectedValue"), + expected_duration=_number(data["expectedDuration"], "expectedDuration", duration=True), + ) diff --git a/energy_tracker_api/resources/__init__.py b/energy_tracker_api/resources/__init__.py index c6cab19..e161c99 100644 --- a/energy_tracker_api/resources/__init__.py +++ b/energy_tracker_api/resources/__init__.py @@ -1,10 +1,12 @@ """Resource handlers for Energy Tracker API.""" +from .calculations import CalculationResource from .devices import DeviceResource from .environments import EnvironmentResource from .meter_readings import MeterReadingResource __all__ = [ + "CalculationResource", "DeviceResource", "EnvironmentResource", "MeterReadingResource", diff --git a/energy_tracker_api/resources/calculations.py b/energy_tracker_api/resources/calculations.py new file mode 100644 index 0000000..038775c --- /dev/null +++ b/energy_tracker_api/resources/calculations.py @@ -0,0 +1,116 @@ +"""Server-side daily values and extrapolations for standard devices.""" + +from datetime import UTC, datetime +from http import HTTPStatus + +from ..exceptions import ValidationError +from ..models import CalculationInterval, CalculationPointDto, ExtrapolationMethod +from .base import BaseResource + + +def _query_params( + from_timestamp: datetime | None, + to_timestamp: datetime | None, + time_zone: str | None, +) -> dict[str, str]: + params: dict[str, str] = {} + for key, timestamp in (("from", from_timestamp), ("to", to_timestamp)): + if timestamp is None: + continue + if not isinstance(timestamp, datetime) or timestamp.utcoffset() is None: + raise ValidationError(f"{key} must be a datetime with a UTC offset") + try: + params[key] = timestamp.astimezone(UTC).isoformat() + except (ValueError, OverflowError) as error: + raise ValidationError(f"{key} cannot be represented as a UTC timestamp") from error + if time_zone is not None: + params["timeZone"] = time_zone + return params + + +class CalculationResource(BaseResource): + """Handler for calculations; values and calendar boundaries are computed by the API.""" + + async def daily_values( + self, + device_id: str, + *, + from_timestamp: datetime | None = None, + to_timestamp: datetime | None = None, + time_zone: str | None = None, + ) -> list[CalculationPointDto]: + """Return reading-backed daily points. Requires scope ``read:daily-values``. + + Args: + device_id: Standard device identifier. + from_timestamp: Inclusive start, with a UTC offset. Omit for no lower filter. + to_timestamp: Exclusive end, with a UTC offset. Omit for no upper filter. + time_zone: IANA time zone; defaults to the device location zone, then UTC. + + The server rounds the start down and end up to day boundaries in time_zone. + An explicit end cannot extend beyond tomorrow. Insufficient readings return []. + Points remain in server order, including terminal points with zero duration. + + Raises: + ValidationError: Invalid input or exceeded calculation limits. + AuthenticationError: Invalid access token. + ForbiddenError: Missing scope or blocked access. + ResourceNotFoundError: Device does not exist or is not owned by the user. + ConflictError: Ambiguous or unrepresentable meter history. + RateLimitError: Request limit exceeded. + ServiceUnavailableError: Calculation unavailable or server deadline exceeded. + TimeoutError: Client calculation timeout exceeded. + EnergyTrackerAPIError: Invalid response or another API/transport failure. + """ + params = _query_params(from_timestamp, to_timestamp, time_zone) + return await self._request_model_list( + response_type=CalculationPointDto, + method="GET", + endpoint=f"/v1/devices/standard/{device_id}/daily-values", + expected_status=HTTPStatus.OK, + params=params or None, + timeout=self._client._calculation_timeout, + ) + + async def extrapolations( + self, + device_id: str, + *, + interval: CalculationInterval, + method: ExtrapolationMethod = ExtrapolationMethod.STANDARD, + from_timestamp: datetime | None = None, + to_timestamp: datetime | None = None, + time_zone: str | None = None, + ) -> list[CalculationPointDto]: + """Return estimates for calendar intervals. Requires scope ``read:extrapolation``. + + Args: + device_id: Standard device identifier. + interval: Day, week, month, quarter or year. Weeks begin on Monday. + method: Calculation method; currently only standard is supported. + from_timestamp: Inclusive start with a UTC offset; defaults to current interval start. + to_timestamp: Exclusive end with a UTC offset; defaults to next interval start. + time_zone: IANA time zone; defaults to the device location zone, then UTC. + + The server rounds the start down and end up to interval boundaries. Output + is limited to 366 calendar days and input to 2,000 readings including boundary + readings. The horizon includes the interval containing twelve months after the + latest reading. Insufficient readings return []. expected_value is the entire + estimate for an interval: do not add actual_value. Errors match daily_values(). + """ + try: + interval = CalculationInterval(interval) + method = ExtrapolationMethod(method) + except (ValueError, TypeError) as error: + raise ValidationError( + "Unsupported calculation interval or extrapolation method" + ) from error + params = _query_params(from_timestamp, to_timestamp, time_zone) + return await self._request_model_list( + response_type=CalculationPointDto, + method="GET", + endpoint=f"/v1/devices/standard/{device_id}/extrapolations/{method.value}/{interval.value}", + expected_status=HTTPStatus.OK, + params=params or None, + timeout=self._client._calculation_timeout, + ) diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 0000000..f7d9271 --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,22 @@ +"""Shared local HTTP server fixture.""" + +import pytest +from aiohttp import web +from aiohttp.test_utils import TestServer + + +@pytest.fixture +async def serve(): + servers = [] + + async def start(handler): + app = web.Application() + app.router.add_route("*", "/{path:.*}", handler) + server = TestServer(app) + servers.append(server) + await server.start_server() + return str(server.make_url("/")) + + yield start + for server in servers: + await server.close() diff --git a/tests/test_exceptions.py b/tests/test_exceptions.py index 05da55b..5d1ebc9 100644 --- a/tests/test_exceptions.py +++ b/tests/test_exceptions.py @@ -8,11 +8,25 @@ NetworkError, RateLimitError, ResourceNotFoundError, + ServiceUnavailableError, TimeoutError, ValidationError, ) +def test_http_status_is_optional_and_old_positional_arguments_still_work(): + error = EnergyTrackerAPIError("Error", ["Detail"]) + assert error.status_code is None + rate_limit = RateLimitError("Retry later", ["Limit reached"], 7, status_code=429) + assert rate_limit.retry_after == 7 + assert rate_limit.api_message == ["Limit reached"] + assert rate_limit.status_code == 429 + unavailable = ServiceUnavailableError("Unavailable", ["Deadline exceeded"], status_code=503) + assert unavailable.status_code == 503 + assert unavailable.api_message == ["Deadline exceeded"] + assert isinstance(unavailable, EnergyTrackerAPIError) + + class TestEnergyTrackerAPIError: """Tests for base EnergyTrackerAPIError.""" diff --git a/tests/test_models_calculations.py b/tests/test_models_calculations.py new file mode 100644 index 0000000..258c1bc --- /dev/null +++ b/tests/test_models_calculations.py @@ -0,0 +1,75 @@ +"""Calculation response contract, including DST durations and malformed numbers.""" + +from datetime import UTC, datetime + +import pytest + +from energy_tracker_api import CalculationPointDto + + +@pytest.fixture +def point(): + return { + "date": "2026-03-28T23:00:00.000Z", + "actualValue": 4.2, + "actualDuration": 82800, + "expectedValue": 8.4, + "expectedDuration": 86400, + } + + +def test_point_preserves_values_and_dst_durations(point): + result = CalculationPointDto._from_dict(point) + assert result == CalculationPointDto( + date=datetime(2026, 3, 28, 23, tzinfo=UTC), + actual_value=4.2, + actual_duration=82800.0, + expected_value=8.4, + expected_duration=86400.0, + ) + + +def test_partial_and_terminal_points_are_not_normalized(point): + point.update(actualValue=-1.5, actualDuration=0, expectedValue=-3, expectedDuration=0) + result = CalculationPointDto._from_dict(point) + assert result.actual_value == -1.5 + assert result.expected_value == -3 + assert result.actual_duration == result.expected_duration == 0 + + +def test_fractional_duration_is_preserved(point): + point["actualDuration"] = 0.5 + assert CalculationPointDto._from_dict(point).actual_duration == 0.5 + + +@pytest.mark.parametrize( + "field", ["date", "actualValue", "actualDuration", "expectedValue", "expectedDuration"] +) +def test_every_field_is_required(point, field): + del point[field] + with pytest.raises(KeyError): + CalculationPointDto._from_dict(point) + + +@pytest.mark.parametrize( + "field", ["actualValue", "actualDuration", "expectedValue", "expectedDuration"] +) +@pytest.mark.parametrize("value", [True, "1.5", None, float("nan"), float("inf")]) +def test_numbers_are_not_coerced_from_invalid_json_values(point, field, value): + point[field] = value + with pytest.raises((TypeError, ValueError)): + CalculationPointDto._from_dict(point) + + +@pytest.mark.parametrize("field", ["actualDuration", "expectedDuration"]) +def test_negative_durations_are_invalid(point, field): + point[field] = -1 + with pytest.raises(ValueError): + CalculationPointDto._from_dict(point) + + +@pytest.mark.parametrize("date", [None, "invalid", "2026-03-29", "2026-03-29T00:00:00"]) +def test_response_dates_require_an_offset(point, date): + point["date"] = date + with pytest.raises((TypeError, ValueError)): + CalculationPointDto._from_dict(point) diff --git a/tests/test_resources_calculations.py b/tests/test_resources_calculations.py new file mode 100644 index 0000000..a613e3a --- /dev/null +++ b/tests/test_resources_calculations.py @@ -0,0 +1,327 @@ +"""Exercise calculation requests and failure contracts through a real HTTP transport.""" + +import asyncio +from datetime import UTC, datetime, timedelta, timezone, tzinfo +from unittest.mock import AsyncMock + +import pytest +from aiohttp import web + +from energy_tracker_api import ( + AuthenticationError, + CalculationInterval, + ConflictError, + EnergyTrackerAPIError, + EnergyTrackerClient, + ForbiddenError, + RateLimitError, + ResourceNotFoundError, + ServiceUnavailableError, + TimeoutError, + ValidationError, +) + + +@pytest.fixture(params=["daily_values", "extrapolations"]) +def operation(request): + return request.param + + +async def calculate(client, operation, **kwargs): + if operation == "extrapolations": + kwargs.setdefault("interval", CalculationInterval.MONTH) + return await getattr(client.calculations, operation)("device-id", **kwargs) + + +async def test_default_requests_and_empty_results(serve, operation): + async def handler(request): + suffix = "daily-values" if operation == "daily_values" else "extrapolations/standard/month" + assert request.method == "GET" + assert request.path == f"/v1/devices/standard/device-id/{suffix}" + assert request.query == {} + assert request.headers["Authorization"] == "Bearer test-token" + return web.json_response([]) + + async with EnergyTrackerClient("test-token", base_url=await serve(handler)) as client: + assert await calculate(client, operation) == [] + + +@pytest.mark.parametrize( + "interval", + [ + CalculationInterval.DAY, + CalculationInterval.WEEK, + CalculationInterval.MONTH, + CalculationInterval.QUARTER, + CalculationInterval.YEAR, + ], +) +async def test_all_intervals_are_encoded_in_path(serve, interval): + async def handler(request): + assert ( + request.path + == f"/v1/devices/standard/device-id/extrapolations/standard/{interval.value}" + ) + return web.json_response([]) + + async with EnergyTrackerClient("test-token", base_url=await serve(handler)) as client: + assert await client.calculations.extrapolations("device-id", interval=interval) == [] + + +async def test_time_filters_preserve_instants_without_calendar_rounding(serve, operation): + async def handler(request): + assert dict(request.query) == { + "from": "2026-03-28T23:15:12.123456+00:00", + "to": "2026-03-29T22:45:23.456789+00:00", + "timeZone": "Europe/Berlin", + } + return web.json_response([]) + + async with EnergyTrackerClient("test-token", base_url=await serve(handler)) as client: + await calculate( + client, + operation, + from_timestamp=datetime.fromisoformat("2026-03-29T00:15:12.123456+01:00"), + to_timestamp=datetime.fromisoformat("2026-03-30T00:45:23.456789+02:00"), + time_zone="Europe/Berlin", + ) + + +@pytest.mark.parametrize("microsecond", [1, 999, 1000, 123456, 999999]) +async def test_time_filters_preserve_subsecond_range_at_midnight(serve, operation, microsecond): + start = datetime(2026, 9, 1, tzinfo=UTC) + end = start + timedelta(microseconds=microsecond) + + async def handler(request): + transmitted_start = datetime.fromisoformat(request.query["from"]) + transmitted_end = datetime.fromisoformat(request.query["to"]) + assert transmitted_start == start + assert transmitted_end == end + assert transmitted_start < transmitted_end + return web.json_response([]) + + async with EnergyTrackerClient("test-token", base_url=await serve(handler)) as client: + await calculate(client, operation, from_timestamp=start, to_timestamp=end) + + +async def test_time_filters_serialize_exact_seconds_and_milliseconds(serve, operation): + async def handler(request): + assert dict(request.query) == { + "from": "2026-09-01T00:00:00+00:00", + "to": "2026-09-02T00:00:00.123000+00:00", + } + return web.json_response([]) + + async with EnergyTrackerClient("test-token", base_url=await serve(handler)) as client: + await calculate( + client, + operation, + from_timestamp=datetime(2026, 9, 1, tzinfo=UTC), + to_timestamp=datetime(2026, 9, 2, microsecond=123000, tzinfo=UTC), + ) + + +@pytest.mark.parametrize("parameter", ["from_timestamp", "to_timestamp", "time_zone"]) +async def test_filters_can_be_provided_individually(serve, operation, parameter): + names = {"from_timestamp": "from", "to_timestamp": "to", "time_zone": "timeZone"} + value = "Europe/Berlin" if parameter == "time_zone" else datetime(2026, 1, 1, tzinfo=UTC) + + async def handler(request): + assert set(request.query) == {names[parameter]} + return web.json_response([]) + + async with EnergyTrackerClient("test-token", base_url=await serve(handler)) as client: + await calculate(client, operation, **{parameter: value}) + + +class MissingOffset(tzinfo): + def utcoffset(self, dt): + return None + + +@pytest.mark.parametrize("parameter", ["from_timestamp", "to_timestamp"]) +@pytest.mark.parametrize( + "value", [datetime(2026, 1, 1), datetime(2026, 1, 1, tzinfo=MissingOffset())] +) +async def test_naive_dates_are_rejected_before_request(operation, parameter, value): + client = EnergyTrackerClient("test-token") + client._make_request = AsyncMock() + with pytest.raises(ValidationError) as exc: + await calculate(client, operation, **{parameter: value}) + assert exc.value.status_code is None + client._make_request.assert_not_called() + + +async def test_date_outside_utc_range_is_local_validation_error(operation): + client = EnergyTrackerClient("test-token") + client._make_request = AsyncMock() + with pytest.raises(ValidationError): + await calculate( + client, operation, from_timestamp=datetime(1, 1, 1, tzinfo=timezone(timedelta(hours=1))) + ) + client._make_request.assert_not_called() + + +@pytest.mark.parametrize( + "options", [{"interval": "minute"}, {"interval": "day", "method": "unknown"}] +) +async def test_invalid_path_enums_are_rejected_before_request(options): + client = EnergyTrackerClient("test-token") + client._make_request = AsyncMock() + with pytest.raises(ValidationError): + await client.calculations.extrapolations("device-id", **options) + client._make_request.assert_not_called() + + +async def test_response_order_and_terminal_points_are_preserved(serve, operation): + async def handler(request): + return web.json_response( + [ + { + "date": "2026-03-28T23:00:00Z", + "actualValue": 4.2, + "actualDuration": 82800, + "expectedValue": 8.4, + "expectedDuration": 86400, + }, + { + "date": "2026-03-29T22:00:00Z", + "actualValue": 0, + "actualDuration": 0, + "expectedValue": 0, + "expectedDuration": 0, + }, + ] + ) + + async with EnergyTrackerClient("test-token", base_url=await serve(handler)) as client: + points = await calculate(client, operation) + assert [point.date for point in points] == [ + datetime(2026, 3, 28, 23, tzinfo=UTC), + datetime(2026, 3, 29, 22, tzinfo=UTC), + ] + assert points[0].expected_value == 8.4 + assert points[0].actual_duration == 82800 + assert points[1].expected_duration == 0 + + +@pytest.mark.parametrize("status", [201, 202, 204, 206, 302]) +async def test_only_200_is_accepted(serve, operation, status): + async def handler(request): + return web.json_response([], status=status) + + async with EnergyTrackerClient("test-token", base_url=await serve(handler)) as client: + with pytest.raises(EnergyTrackerAPIError, match="Unexpected HTTP status") as exc: + await calculate(client, operation) + assert exc.value.status_code == status + + +@pytest.mark.parametrize( + "status,error", + [ + (400, ValidationError), + (401, AuthenticationError), + (403, ForbiddenError), + (404, ResourceNotFoundError), + (409, ConflictError), + (429, RateLimitError), + (503, ServiceUnavailableError), + (500, EnergyTrackerAPIError), + (418, EnergyTrackerAPIError), + ], +) +async def test_http_errors_preserve_status_and_do_not_retry(serve, operation, status, error): + requests = 0 + + async def handler(request): + nonlocal requests + requests += 1 + return web.json_response( + {"message": "Calculation unavailable"}, status=status, headers={"Retry-After": "7"} + ) + + async with EnergyTrackerClient("test-token", base_url=await serve(handler)) as client: + with pytest.raises(error) as exc: + await calculate(client, operation) + assert type(exc.value) is error + assert exc.value.status_code == status + assert exc.value.api_message == ["Calculation unavailable"] + assert isinstance(exc.value, EnergyTrackerAPIError) + if status == 429: + assert exc.value.retry_after == 7 + assert requests == 1 + + +@pytest.mark.parametrize( + "payload", + [ + {}, + [None], + [{}], + [ + { + "date": "2026-01-01", + "actualValue": 0, + "actualDuration": 0, + "expectedValue": 0, + "expectedDuration": 0, + } + ], + ], +) +async def test_invalid_response_shape_is_api_error(serve, operation, payload): + async def handler(request): + return web.json_response(payload) + + async with EnergyTrackerClient("test-token", base_url=await serve(handler)) as client: + with pytest.raises(EnergyTrackerAPIError): + await calculate(client, operation) + + +async def test_calculation_timeout_does_not_replace_session_timeout(serve, operation): + async def handler(request): + await asyncio.sleep(0.05) + return web.json_response([]) + + async with EnergyTrackerClient( + "test-token", base_url=await serve(handler), timeout=0.01, calculation_timeout=1 + ) as client: + session = await client._get_session() + assert await calculate(client, operation) == [] + assert session is await client._get_session() + assert session.timeout.total == 0.01 + with pytest.raises(TimeoutError): + await client.devices.list_standard() + + +async def test_calculation_deadline_is_enforced_without_retries(serve, operation): + requests = 0 + + async def handler(request): + nonlocal requests + requests += 1 + await asyncio.sleep(0.05) + return web.json_response([]) + + async with EnergyTrackerClient( + "test-token", base_url=await serve(handler), calculation_timeout=0.01 + ) as client: + with pytest.raises(TimeoutError) as exc: + await calculate(client, operation) + assert exc.value.status_code is None + assert requests == 1 + + +def test_default_and_custom_timeouts(): + default = EnergyTrackerClient("test-token") + assert default._timeout.total == 10 + assert default._calculation_timeout.total == 60 + configured = EnergyTrackerClient("test-token", None, 30, calculation_timeout=90) + assert configured._timeout.total == 30 + assert configured._calculation_timeout.total == 90 + + +@pytest.mark.parametrize("value", [0, -1, float("nan"), float("inf"), True, None, "60"]) +def test_invalid_calculation_timeout(value): + with pytest.raises(ValueError, match="calculation_timeout"): + EnergyTrackerClient("test-token", calculation_timeout=value) diff --git a/tests/test_transport.py b/tests/test_transport.py index 9e99840..03a94b6 100644 --- a/tests/test_transport.py +++ b/tests/test_transport.py @@ -7,7 +7,6 @@ import aiohttp import pytest from aiohttp import web -from aiohttp.test_utils import TestServer from energy_tracker_api import ( CreateEnvironmentEntryDto, @@ -23,23 +22,6 @@ ) -@pytest.fixture -async def serve(): - servers = [] - - async def start(handler): - app = web.Application() - app.router.add_route("*", "/{path:.*}", handler) - server = TestServer(app) - servers.append(server) - await server.start_server() - return str(server.make_url("/")) - - yield start - for server in servers: - await server.close() - - @pytest.mark.parametrize("body", [b'"123.45"', b'\xef\xbb\xbf"123.45"', b"123", b"null", b""]) async def test_export_preserves_exact_bytes(serve, body): async def handler(request):