Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
56 changes: 56 additions & 0 deletions .github/workflows/ci.yml
Original file line number Diff line number Diff line change
@@ -0,0 +1,56 @@
name: CI

on:
push:
branches: [main]
pull_request:

concurrency:
group: ci-${{ github.ref }}
cancel-in-progress: true

jobs:
test:
name: test (py${{ matrix.python-version }})
runs-on: ubuntu-latest
strategy:
fail-fast: false
matrix:
python-version: ["3.10", "3.11", "3.12", "3.13"]
steps:
- uses: actions/checkout@v4

- name: Install uv
uses: astral-sh/setup-uv@v6
with:
python-version: ${{ matrix.python-version }}
enable-cache: true

- name: Install dependencies
run: uv sync --all-extras

- name: Run tests
env:
# Optional: when set, the live API tests also run; otherwise skipped.
OPENELECTRICITY_API_KEY: ${{ secrets.OPENELECTRICITY_API_KEY }}
run: uv run pytest

lint:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v4

- name: Install uv
uses: astral-sh/setup-uv@v6
with:
python-version: "3.12"
enable-cache: true

- name: Install dependencies
run: uv sync --all-extras

- name: Ruff lint
run: uv run ruff check openelectricity tests

- name: Ruff format
run: uv run ruff format --check openelectricity tests
4 changes: 4 additions & 0 deletions .worktreeinclude
Original file line number Diff line number Diff line change
@@ -0,0 +1,4 @@
# gitignored files copied into `claude --worktree` worktrees so an isolated
# checkout has a working dev environment (API key for tests, publish tokens).
.env
.envrc
150 changes: 131 additions & 19 deletions openelectricity/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,10 +6,11 @@

import asyncio
import os
import ssl
from datetime import datetime
from typing import Any, TypeVar, cast

from aiohttp import ClientResponse, ClientSession
from aiohttp import BasicAuth, ClientResponse, ClientSession, TCPConnector

from openelectricity.logging import get_logger
from openelectricity.models.facilities import FacilityResponse
Expand Down Expand Up @@ -54,9 +55,34 @@ class BaseOEClient:
api_key: Optional API key for authentication. If not provided, will look for
OPENELECTRICITY_API_KEY environment variable.
base_url: Optional base URL for the API. Defaults to production API.
proxy: Optional proxy URL (e.g. "http://proxy.corp:8080") for all requests.
proxy_auth: Optional aiohttp.BasicAuth for proxy authentication.
verify_ssl: Whether to verify TLS certificates. Defaults to True. Set to
False to disable verification (not recommended).
ssl_context: Optional pre-built ssl.SSLContext, e.g. for a corporate CA.
ca_cert: Optional path to a CA certificate bundle. A convenience over
ssl_context for the common "trust this extra CA" case.
trust_env: Whether aiohttp should read proxy settings, .netrc and TLS
config from the environment (HTTP_PROXY/HTTPS_PROXY etc).
Defaults to False.

Note:
ssl_context and ca_cert are mutually exclusive, and neither can be
combined with verify_ssl=False.
"""

def __init__(self, api_key: str | None = None, base_url: str | None = None) -> None:
def __init__(
self,
api_key: str | None = None,
base_url: str | None = None,
*,
proxy: str | None = None,
proxy_auth: BasicAuth | None = None,
verify_ssl: bool = True,
ssl_context: ssl.SSLContext | None = None,
ca_cert: str | os.PathLike[str] | None = None,
trust_env: bool = False,
) -> None:
# Ensure base_url has a trailing slash for aiohttp ClientSession
if base_url:
self.base_url = base_url.rstrip("/") + "/"
Expand All @@ -77,8 +103,60 @@ def __init__(self, api_key: str | None = None, base_url: str | None = None) -> N
"Accept": "application/json",
"Content-Type": "application/json",
}

# Proxy and TLS configuration, applied in _build_session().
self.proxy = proxy
self.proxy_auth = proxy_auth
self.trust_env = trust_env
self._ssl = self._resolve_ssl(verify_ssl, ssl_context, ca_cert)

logger.debug("Initialized client with base URL: %s", self.base_url)

@staticmethod
def _resolve_ssl(
verify_ssl: bool,
ssl_context: ssl.SSLContext | None,
ca_cert: str | os.PathLike[str] | None,
) -> ssl.SSLContext | bool | None:
"""Resolve TLS settings into a value for aiohttp's TCPConnector.

Returns None when defaults apply (no custom connector needed), False to
disable verification, or an SSLContext for a custom CA.
"""
if ssl_context is not None and ca_cert is not None:
raise OpenElectricityError("Provide either ssl_context or ca_cert, not both")
if (ssl_context is not None or ca_cert is not None) and not verify_ssl:
raise OpenElectricityError("verify_ssl=False conflicts with a custom ssl_context/ca_cert")

if ssl_context is not None:
return ssl_context
if ca_cert is not None:
# Add the extra CA to the default trust store rather than
# replacing it (cafile= on create_default_context would drop the
# system CAs and break normal public TLS).
context = ssl.create_default_context()
context.load_verify_locations(cafile=os.fspath(ca_cert))
return context
if not verify_ssl:
return False
return None

def _build_session(self) -> ClientSession:
"""Create a ClientSession with the configured proxy and TLS options.

Single point of session construction so proxy, custom certificates and
trust_env behaviour stay consistent across the sync and async clients.
"""
connector = TCPConnector(ssl=self._ssl) if self._ssl is not None else None
return ClientSession(
base_url=self.base_url,
headers=self.headers,
trust_env=self.trust_env,
proxy=self.proxy,
proxy_auth=self.proxy_auth,
connector=connector,
)


class OEClient(BaseOEClient):
"""
Expand All @@ -88,8 +166,28 @@ class OEClient(BaseOEClient):
API consistency while using the same underlying HTTP client as the async version.
"""

def __init__(self, api_key: str | None = None, base_url: str | None = None) -> None:
super().__init__(api_key, base_url)
def __init__(
self,
api_key: str | None = None,
base_url: str | None = None,
*,
proxy: str | None = None,
proxy_auth: BasicAuth | None = None,
verify_ssl: bool = True,
ssl_context: ssl.SSLContext | None = None,
ca_cert: str | os.PathLike[str] | None = None,
trust_env: bool = False,
) -> None:
super().__init__(
api_key,
base_url,
proxy=proxy,
proxy_auth=proxy_auth,
verify_ssl=verify_ssl,
ssl_context=ssl_context,
ca_cert=ca_cert,
trust_env=trust_env,
)
self._session: ClientSession | None = None
self._loop: asyncio.AbstractEventLoop | None = None
logger.debug("Initialized synchronous client")
Expand All @@ -98,10 +196,7 @@ def _ensure_session(self) -> None:
"""Ensure session and event loop are initialized."""
if self._session is None or self._session.closed:
logger.debug("Creating new client session")
self._session = ClientSession(
base_url=self.base_url,
headers=self.headers,
)
self._session = self._build_session()

async def _handle_response(self, response: ClientResponse) -> dict[str, Any] | list[dict[str, Any]]:
"""Handle API response and raise appropriate errors."""
Expand Down Expand Up @@ -288,7 +383,7 @@ def get_facilities(
"""Get a list of facilities."""

async def _run():
async with ClientSession(base_url=self.base_url, headers=self.headers) as session:
async with self._build_session() as session:
self._session = session
return await self._async_get_facilities(facility_code, status_id, fueltech_id, network_id, network_region)

Expand Down Expand Up @@ -327,7 +422,7 @@ def get_network_data(
"""

async def _run():
async with ClientSession(base_url=self.base_url, headers=self.headers) as session:
async with self._build_session() as session:
self._session = session
return await self._async_get_network_data(
network_code,
Expand Down Expand Up @@ -357,7 +452,7 @@ def get_facility_data(
"""Get facility data for specified metrics."""

async def _run():
async with ClientSession(base_url=self.base_url, headers=self.headers) as session:
async with self._build_session() as session:
self._session = session
return await self._async_get_facility_data(
network_code, facility_code, metrics, interval, date_start, date_end, unit_code
Expand All @@ -378,7 +473,7 @@ def get_market(
"""Get market data for specified metrics."""

async def _run():
async with ClientSession(base_url=self.base_url, headers=self.headers) as session:
async with self._build_session() as session:
self._session = session
return await self._async_get_market(
network_code, metrics, interval, date_start, date_end, primary_grouping, network_region
Expand All @@ -390,7 +485,7 @@ def get_current_user(self) -> OpennemUserResponse:
"""Get current user information."""

async def _run():
async with ClientSession(base_url=self.base_url, headers=self.headers) as session:
async with self._build_session() as session:
self._session = session
return await self._async_get_current_user()

Expand All @@ -417,19 +512,36 @@ class AsyncOEClient(BaseOEClient):
Asynchronous client for the OpenElectricity API.
"""

def __init__(self, api_key: str | None = None, base_url: str | None = None) -> None:
super().__init__(api_key, base_url)
def __init__(
self,
api_key: str | None = None,
base_url: str | None = None,
*,
proxy: str | None = None,
proxy_auth: BasicAuth | None = None,
verify_ssl: bool = True,
ssl_context: ssl.SSLContext | None = None,
ca_cert: str | os.PathLike[str] | None = None,
trust_env: bool = False,
) -> None:
super().__init__(
api_key,
base_url,
proxy=proxy,
proxy_auth=proxy_auth,
verify_ssl=verify_ssl,
ssl_context=ssl_context,
ca_cert=ca_cert,
trust_env=trust_env,
)
self.client: ClientSession | None = None
logger.debug("Initialized asynchronous client")

async def _ensure_client(self) -> None:
"""Ensure client session is initialized."""
if self.client is None or self.client.closed:
logger.debug("Creating new async client session")
self.client = ClientSession(
base_url=self.base_url,
headers=self.headers,
)
self.client = self._build_session()

async def _handle_response(self, response: ClientResponse) -> dict[str, Any] | list[dict[str, Any]]:
"""Handle API response and raise appropriate errors."""
Expand Down
14 changes: 13 additions & 1 deletion openelectricity/types.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,9 +5,21 @@
Matches the TypeScript definitions from the official client.
"""

from enum import StrEnum
import sys
from typing import Any, Literal

if sys.version_info >= (3, 11):
from enum import StrEnum
else: # pragma: no cover - Python 3.10 compatibility shim
from enum import Enum

class StrEnum(str, Enum):
"""Minimal backport of enum.StrEnum for Python 3.10."""

__str__ = str.__str__
__format__ = str.__format__


# Network and Data Types
NetworkCode = Literal["NEM", "WEM", "AU"]
DataInterval = Literal["5m", "1h", "1d", "7d", "1M", "3M", "season", "1y", "fy"]
Expand Down
5 changes: 4 additions & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@ classifiers = [
"Development Status :: 4 - Beta",
"Intended Audience :: Developers",
"Programming Language :: Python :: 3",
"Programming Language :: Python :: 3.13",
"Programming Language :: Python :: 3.12",
"Programming Language :: Python :: 3.11",
"Programming Language :: Python :: 3.10",
Expand Down Expand Up @@ -45,7 +46,9 @@ build-backend = "hatchling.build"
path = "openelectricity/__init__.py"

[tool.ruff]
target-version = "py312"
# Match requires-python (>=3.10) so lint never suggests 3.12-only syntax
# such as PEP 695 generics, which break the advertised 3.10/3.11 support.
target-version = "py310"
line-length = 130
exclude = [
"docs",
Expand Down
6 changes: 3 additions & 3 deletions tests/models/test_timeseries.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@
using real API response examples.
"""

from datetime import UTC, datetime
from datetime import datetime, timezone

import pytest

Expand Down Expand Up @@ -148,7 +148,7 @@ def test_timeseries_result_parsing(facility_response):

# Check first data point
first_point = result.data[0]
assert first_point.timestamp == datetime(2025, 2, 12, 13, tzinfo=UTC)
assert first_point.timestamp == datetime(2025, 2, 12, 13, tzinfo=timezone.utc)
assert first_point.value == 931.4554


Expand All @@ -157,7 +157,7 @@ def test_timeseries_datapoint_parsing():
data = ["2025-02-12T13:00:00Z", 931.4554]
point = TimeSeriesDataPoint.model_validate(data)

assert point.timestamp == datetime(2025, 2, 12, 13, tzinfo=UTC)
assert point.timestamp == datetime(2025, 2, 12, 13, tzinfo=timezone.utc)
assert point.value == 931.4554


Expand Down
Loading
Loading