Skip to content
Open
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
1 change: 1 addition & 0 deletions Makefile
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,7 @@ tests: .install-poetry
poetry run black --check --diff ocpp tests
poetry run isort --check-only ocpp tests
poetry run flake8 ocpp tests
poetry run mypy --strict ocpp/
poetry run py.test -vvv --cov=ocpp --cov-report=term-missing tests/

build: .install-poetry
Expand Down
12 changes: 12 additions & 0 deletions ocpp/_types.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,12 @@
from typing import Any, Callable, Literal, TypedDict

OCPPVersion = Literal["1.6", "2.0", "2.0.1", "2.1"]


Handler = Callable[..., Any]


class Route(TypedDict, total=False):
_on_action: Handler
_after_action: Handler
_skip_schema_validation: bool
142 changes: 114 additions & 28 deletions ocpp/charge_point.py
Original file line number Diff line number Diff line change
@@ -1,20 +1,54 @@
from __future__ import annotations

import asyncio
import inspect
import logging
import re
import time
import uuid
from dataclasses import Field, asdict, is_dataclass
from typing import Any, Dict, List, Optional, Union, get_args, get_origin
from types import ModuleType
from typing import (
TYPE_CHECKING,
Any,
Callable,
Dict,
List,
NoReturn,
Optional,
Protocol,
TypeGuard,
Union,
get_args,
get_origin,
overload,
)
from urllib.parse import urlparse

from ocpp._types import OCPPVersion, Route
from ocpp.exceptions import NotImplementedError, NotSupportedError, OCPPError
from ocpp.messages import Call, MessageType, unpack, validate_payload
from ocpp.messages import (
Call,
CallError,
CallResult,
MessageType,
unpack,
validate_payload,
)
from ocpp.routing import create_route_map

if TYPE_CHECKING:
from _typeshed import DataclassInstance


LOGGER = logging.getLogger("ocpp")


class WebSocket(Protocol):
async def recv(self) -> Union[str, bytes]: ...
async def send(self, payload: Any) -> None: ...


def extract_charge_point_id(path: Optional[str]) -> Optional[str]:
"""Extract the charge point ID from a WebSocket URL path.

Expand Down Expand Up @@ -60,7 +94,19 @@
return charge_point_id


def camel_to_snake_case(data):
@overload
def camel_to_snake_case(data: Dict[str, Any]) -> Dict[str, Any]: ...


@overload
def camel_to_snake_case(data: List[Any]) -> List[Any]: ...


@overload
def camel_to_snake_case(data: DataclassInstance) -> Dict[str, Any]: ...


def camel_to_snake_case(data: Any) -> Any:
"""
Convert all keys of all dictionaries inside the given argument from
camelCase to snake_case.
Expand Down Expand Up @@ -90,7 +136,17 @@
return data


def snake_to_camel_case(data):
@overload
def snake_to_camel_case(data: Dict[str, Any]) -> Dict[str, Any]: ...


@overload
def snake_to_camel_case(data: List[Any]) -> List[Any]: ...


def snake_to_camel_case(
data: Union[Dict[str, Any], List[Any]],
) -> Union[Dict[str, Any], List[Any]]:
"""
Convert all keys of all dictionaries inside given argument from
snake_case to camelCase.
Expand Down Expand Up @@ -127,12 +183,12 @@
return data


def _is_dataclass_instance(input: Any) -> bool:
def _is_dataclass_instance(input: Any) -> TypeGuard[DataclassInstance]:
"""Verify if given `input` is a dataclass."""
return is_dataclass(input) and not isinstance(input, type)


def _is_optional_field(field: Field) -> bool:
def _is_optional_field(field: Field[Any]) -> bool:
"""Verify if given `field` allows `None` as value.

The fields `schema` and `host` on the following class would return `False`.
Expand All @@ -149,7 +205,7 @@
return get_origin(field.type) is Union and type(None) in get_args(field.type)


def serialize_as_dict(dataclass):
def serialize_as_dict(dataclass: DataclassInstance) -> Dict[str, Any]:
"""Serialize the given `dataclass` as a `dict` recursively.

@dataclass
Expand Down Expand Up @@ -194,7 +250,17 @@
return serialized


def remove_nones(data: Union[List, Dict]) -> Union[List, Dict]:
@overload
def remove_nones(data: List[Any]) -> List[Any]: ...


@overload
def remove_nones(data: Dict[str, Any]) -> Dict[str, Any]: ...


def remove_nones(
data: Union[List[Any], Dict[str, Any]],
) -> Union[List[Any], Dict[str, Any]]:
if isinstance(data, dict):
return {k: remove_nones(v) for k, v in data.items() if v is not None}

Expand All @@ -204,7 +270,7 @@
return data


def _raise_key_error(action, version):
def _raise_key_error(action: str, version: OCPPVersion) -> NoReturn | None:
"""
Checks whether a keyerror returned by _handle_call
is supported by the OCPP version or is simply
Expand Down Expand Up @@ -236,7 +302,7 @@
details={"cause": f"{action} not supported by OCPP{version}."}
)

return
return None


class ChargePoint:
Expand All @@ -245,14 +311,24 @@
initiated and received by the Central System
"""

def __init__(self, id, connection, response_timeout=30, logger=LOGGER):
_call: ModuleType
_call_result: ModuleType
_ocpp_version: OCPPVersion

def __init__(
self,
id: str,
connection: WebSocket,
response_timeout: float = 30.0,
logger: logging.Logger = LOGGER,
):
"""

Args:

charger_id (str): ID of the charger.
connection: Connection to CP.
response_timeout (int): When no response on a request is received
response_timeout (float): When no response on a request is received
within this interval, a asyncio.TimeoutError is raised.
logger: Optional Logger instance used for logging.
By default, the 'ocpp' logger is used.
Expand All @@ -273,28 +349,30 @@
# if exists.
self.route_map = create_route_map(self)

self._call_lock = asyncio.Lock()
self._call_lock: asyncio.Lock = asyncio.Lock()

# A queue used to pass CallResults and CallErrors from
# the self.serve() task to the self.call() task.
self._response_queue = asyncio.Queue()
self._response_queue: asyncio.Queue[Union[Call, CallResult, CallError]] = (
asyncio.Queue()
)

# Function used to generate unique ids for CALLs. By default
# uuid.uuid4() is used, but it can be changed. This is meant primarily
# for testing purposes to have predictable unique ids.
self._unique_id_generator = uuid.uuid4
self._unique_id_generator: Callable[[], uuid.UUID] = uuid.uuid4

# The logger used to log messages
self.logger = logger

async def start(self):
async def start(self) -> NoReturn:
while True:
message = await self._connection.recv()
self.logger.debug("%s: receive message %s", self.id, message)

await self.route_message(message)

async def route_message(self, raw_msg):
async def route_message(self, raw_msg: Union[str, bytes]) -> None:
"""
Route a message received from a CP.

Expand All @@ -306,14 +384,14 @@
msg = unpack(raw_msg)
except OCPPError as e:
self.logger.exception(
"Unable to parse message: '%s', it doesn't seem "
"to be valid OCPP: %s",
"Unable to parse message: '%s', it doesn't seem to be valid OCPP: %s",
raw_msg,
e,
)
return

if msg.message_type_id == MessageType.Call:
assert isinstance(msg, Call)
try:
await self._handle_call(msg)
except OCPPError as error:
Expand All @@ -322,9 +400,10 @@
await self._send(response)

elif msg.message_type_id in [MessageType.CallResult, MessageType.CallError]:
assert isinstance(msg, (CallResult, CallError))
self._response_queue.put_nowait(msg)

async def _handle_call(self, msg):
async def _handle_call(self, msg: Call) -> Any:
"""
Execute all hooks installed for based on the Action of the message.

Expand All @@ -337,10 +416,10 @@

"""
try:
handlers = self.route_map[msg.action]
handlers: Route = self.route_map[msg.action]
except KeyError:
_raise_key_error(msg.action, self._ocpp_version)
return
return None

if not handlers.get("_skip_schema_validation", False):
await validate_payload(msg, self._ocpp_version)
Expand All @@ -354,7 +433,7 @@
snake_case_payload = camel_to_snake_case(msg.payload)

try:
handler = handlers["_on_action"]
handler: Callable[..., Any] = handlers["_on_action"]
except KeyError:
_raise_key_error(msg.action, self._ocpp_version)
handler_signature = inspect.signature(handler)
Expand All @@ -373,7 +452,7 @@
response = msg.create_call_error(e).to_json()
await self._send(response)

return
return None

temp_response_payload = serialize_as_dict(response)

Expand Down Expand Up @@ -419,9 +498,14 @@
return response

async def call(
self, payload, suppress=True, unique_id=None, skip_schema_validation=False
):
self,
payload: DataclassInstance,
suppress: bool = True,
unique_id: Optional[str] = None,
skip_schema_validation: bool = False,
) -> Any:
"""

Send Call message to client and return payload of response.

The given payload is transformed into a Call object by looking at the
Expand Down Expand Up @@ -494,7 +578,9 @@
cls = getattr(self._call_result, payload.__class__.__name__) # noqa
return cls(**snake_case_payload)

async def _get_specific_response(self, unique_id, timeout):
async def _get_specific_response(
self, unique_id: Optional[str], timeout: Union[int, float]

Check warning on line 582 in ocpp/charge_point.py

View check run for this annotation

SonarQubeCloud / SonarCloud Code Analysis

Remove this "timeout" parameter and use a timeout context manager instead.

See more on https://sonarcloud.io/project/issues?id=mobilityhouse_ocpp&issues=AaDJFLvdMsgLhz9raEC3&open=AaDJFLvdMsgLhz9raEC3&pullRequest=771
) -> Any:
"""
Return response with given unique ID or raise an asyncio.TimeoutError.
"""
Expand All @@ -516,6 +602,6 @@

return await self._get_specific_response(unique_id, timeout_left)

async def _send(self, message):
async def _send(self, message: Any) -> None:
self.logger.debug("%s: send message %s", self.id, message)
await self._connection.send(message)
Loading