From 891ea76d3521998e7c8befa2f950bb0d15ebcdda Mon Sep 17 00:00:00 2001 From: Christian Date: Sun, 23 Nov 2025 18:01:38 +0100 Subject: [PATCH 1/6] Support windows, use TCP for daemon/client connection --- src/someipy/__init__.py | 4 ++- src/someipy/_internal/_daemon/__init__.py | 0 .../{ => _daemon}/someipy_daemon_client.py | 25 +++++++++---- .../_internal/{ => _daemon}/uds_messages.py | 2 +- src/someipy/_internal/daemon_client_abcs.py | 5 ++- src/someipy/client_service_instance.py | 4 +-- src/someipy/server_service_instance.py | 4 +-- src/someipy/someipyd.py | 35 +++++++++++++------ 8 files changed, 55 insertions(+), 24 deletions(-) create mode 100644 src/someipy/_internal/_daemon/__init__.py rename src/someipy/_internal/{ => _daemon}/someipy_daemon_client.py (94%) rename src/someipy/_internal/{ => _daemon}/uds_messages.py (99%) diff --git a/src/someipy/__init__.py b/src/someipy/__init__.py index 846ba37..eafd38a 100644 --- a/src/someipy/__init__.py +++ b/src/someipy/__init__.py @@ -6,4 +6,6 @@ from ._internal.method_result import MethodResult # noqa: F401 from ._internal.return_codes import ReturnCode # noqa: F401 from ._internal.message_types import MessageType # noqa: F401 -from ._internal.someipy_daemon_client import connect_to_someipy_daemon # noqa: F401 +from ._internal._daemon.someipy_daemon_client import ( + connect_to_someipy_daemon, +) # noqa: F401 diff --git a/src/someipy/_internal/_daemon/__init__.py b/src/someipy/_internal/_daemon/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/someipy/_internal/someipy_daemon_client.py b/src/someipy/_internal/_daemon/someipy_daemon_client.py similarity index 94% rename from src/someipy/_internal/someipy_daemon_client.py rename to src/someipy/_internal/_daemon/someipy_daemon_client.py index c25194c..8d336cf 100644 --- a/src/someipy/_internal/someipy_daemon_client.py +++ b/src/someipy/_internal/_daemon/someipy_daemon_client.py @@ -17,6 +17,7 @@ import base64 import ipaddress import json +import platform import struct from typing import Dict, List, TypedDict, cast @@ -27,7 +28,7 @@ from someipy._internal.logging import get_logger from someipy._internal.someip_sd_header import SdService from someipy._internal.transport_layer_protocol import TransportLayerProtocol -from someipy._internal.uds_messages import ( +from someipy._internal._daemon.uds_messages import ( BaseMessage, InboundCallMethodRequest, InboundCallMethodResponse, @@ -74,10 +75,14 @@ class SomeIpDaemonClient: def __init__(self, config: dict = None): self._config = config - if self._config == None or "socket_path" not in self._config: + if self._config is None: + self._use_tcp = platform.system() != "Linux" + self._tcp_port = 30500 self._socket_path = "/tmp/someipyd.sock" else: - self._socket_path = self._config["socket_path"] + self._socket_path = self._config.get("socket_path", "/tmp/someipyd.sock") + self._use_tcp = self._config.get("use_tcp", platform.system() != "Linux") + self._tcp_port = self._config.get("tcp_port", 30500) self._rx_message_queue: asyncio.Queue[DaemonMessage] = asyncio.Queue() self._rx_task: asyncio.Task = None @@ -148,15 +153,21 @@ async def _connect_to_daemon(self): self._clear_rx_queue() self._rx_message_queue = asyncio.Queue() - self.reader, self.writer = await asyncio.open_unix_connection( - self._socket_path - ) + if self._use_tcp: + self.reader, self.writer = await asyncio.open_connection( + "127.0.0.1", self._tcp_port + ) + else: + self.reader, self.writer = await asyncio.open_unix_connection( + self._socket_path + ) success = True break except Exception as e: + connection = "TCP" if self._use_tcp else "Unix Domain Socket" get_logger(_logger_name).error( - f"Failed to connect to daemon: {e}. Retries left: {num_retries}" + f"Failed to connect to daemon via {connection}: {e}. Retries left: {num_retries}" ) await asyncio.sleep(1.0) diff --git a/src/someipy/_internal/uds_messages.py b/src/someipy/_internal/_daemon/uds_messages.py similarity index 99% rename from src/someipy/_internal/uds_messages.py rename to src/someipy/_internal/_daemon/uds_messages.py index 1f3e1e0..cc0199b 100644 --- a/src/someipy/_internal/uds_messages.py +++ b/src/someipy/_internal/_daemon/uds_messages.py @@ -1,4 +1,4 @@ -# Copyright (C) 2024 Christian H. +# Copyright (C) 2025 Christian H. # # This program is free software: you can redistribute it and/or modify # it under the terms of the GNU General Public License as published by diff --git a/src/someipy/_internal/daemon_client_abcs.py b/src/someipy/_internal/daemon_client_abcs.py index 72baa9c..1001e42 100644 --- a/src/someipy/_internal/daemon_client_abcs.py +++ b/src/someipy/_internal/daemon_client_abcs.py @@ -16,7 +16,10 @@ from abc import ABC, abstractmethod from typing import Tuple -from someipy._internal.uds_messages import InboundCallMethodResponse, ReceivedEvent +from someipy._internal._daemon.uds_messages import ( + InboundCallMethodResponse, + ReceivedEvent, +) from someipy.service import Service diff --git a/src/someipy/client_service_instance.py b/src/someipy/client_service_instance.py index 92169eb..efe0a5a 100644 --- a/src/someipy/client_service_instance.py +++ b/src/someipy/client_service_instance.py @@ -22,8 +22,8 @@ from someipy._internal.daemon_client_abcs import ClientInstanceInterface from someipy._internal.method_result import MethodResult from someipy._internal.someip_sd_header import SdService -from someipy._internal.someipy_daemon_client import SomeIpDaemonClient -from someipy._internal.uds_messages import ( +from someipy._internal._daemon.someipy_daemon_client import SomeIpDaemonClient +from someipy._internal._daemon.uds_messages import ( OutboundCallMethodRequest, OutboundCallMethodResponse, ReceivedEvent, diff --git a/src/someipy/server_service_instance.py b/src/someipy/server_service_instance.py index ed64eb1..9ab94b4 100644 --- a/src/someipy/server_service_instance.py +++ b/src/someipy/server_service_instance.py @@ -16,8 +16,8 @@ import base64 from typing import List from someipy._internal.daemon_client_abcs import ServerInstanceInterface -from someipy._internal.someipy_daemon_client import SomeIpDaemonClient -from someipy._internal.uds_messages import create_uds_message, SendEventRequest +from someipy._internal._daemon.someipy_daemon_client import SomeIpDaemonClient +from someipy._internal._daemon.uds_messages import create_uds_message, SendEventRequest from someipy.service import EventGroup, Method, Service from someipy._internal.logging import get_logger diff --git a/src/someipy/someipyd.py b/src/someipy/someipyd.py index 14c65cb..8eba696 100644 --- a/src/someipy/someipyd.py +++ b/src/someipy/someipyd.py @@ -22,6 +22,7 @@ import json import logging import os +import platform import struct import sys import ipaddress @@ -66,7 +67,7 @@ SomeIpSdHeader, ) from someipy._internal.subscribers import EventGroupSubscriber, Subscribers -from someipy._internal.uds_messages import ( +from someipy._internal._daemon.uds_messages import ( InboundCallMethodRequest, InboundCallMethodResponse, FindServiceRequest, @@ -96,6 +97,7 @@ DEFAULT_SD_ADDRESS = "224.224.224.245" DEFAULT_INTERFACE_IP = "127.0.0.2" DEFAULT_SD_PORT = 30490 +DEFAULT_TCP_PORT = 30500 class Subscription: @@ -285,6 +287,8 @@ def __init__(self, config_file=None, log_path=None): self.interface = self.config.get("interface", DEFAULT_INTERFACE_IP) log_level = self.config.get("log_level", "DEBUG") self.log_path = log_path if log_path else self.config.get("log_path") + self.use_tcp = self.config.get("use_tcp", platform.system() != "Linux") + self.tcp_port = self.config.get("tcp_port", DEFAULT_TCP_PORT) self.logger = self._configure_logging( log_level=log_level, log_path=self.log_path @@ -298,6 +302,8 @@ def __init__(self, config_file=None, log_path=None): f"Interface: {self.interface}\n" f"Loglevel: {log_level}\n" f"Log path: {self.log_path if self.log_path else 'Console'}\n" + f"Use TCP: {self.use_tcp}\n" + f"TCP Port: {self.tcp_port}\n" ) log_level_mapping = { @@ -1427,27 +1433,36 @@ async def send_to_all_clients(self, message): await writer.wait_closed() async def start_server(self): - if os.path.exists(self.socket_path): - os.unlink(self.socket_path) - server = await asyncio.start_unix_server( - self.handle_client, path=self.socket_path - ) - self.logger.info(f"Unix domain socket server started at {self.socket_path}") + if not self.use_tcp: + if os.path.exists(self.socket_path): + os.unlink(self.socket_path) + + server = await asyncio.start_unix_server( + self.handle_client, path=self.socket_path + ) + self.logger.info(f"Unix domain socket server started at {self.socket_path}") + else: + server = await asyncio.start_server( + self.handle_client, + host="127.0.0.1", + port=self.tcp_port, + reuse_port=True, + ) try: await self.start_sd_listening() async with server: await server.serve_forever() except asyncio.CancelledError: - self.logger.info("UDS server cancelled.") - pass + self.logger.info(f"{"TCP" if self.use_tcp else "UDS"} server cancelled.") finally: if self._mcast_transport: self._mcast_transport.close() if self._ucast_transport: self._ucast_transport.close() - self.logger.info("UDS server stopped.") + + self.logger.info(f"{'TCP' if self.use_tcp else 'UDS'} server stopped.") def _timeout_of_offered_service(self, offered_service: SdService): self.logger.info( From 8258c9672ee93904b9211d653ebcc63caf0ab685 Mon Sep 17 00:00:00 2001 From: Christian Date: Thu, 25 Dec 2025 17:05:34 +0100 Subject: [PATCH 2/6] * Support Windows (TCP for daemon communication) * Start refactor to improve testability * Implement first unit tests of daemon --- .coveragerc | 3 + .gitignore | 4 +- setup.cfg | 2 +- src/someipy/_internal/_common/endpoint.py | 33 + src/someipy/_internal/_common/event.py | 48 ++ .../_internal/_daemon/daemon_server.py | 60 ++ .../_internal/_daemon/daemon_server_client.py | 91 +++ src/someipy/_internal/_sd/__init__.py | 0 .../_internal/_sd/deserialization/__init__.py | 0 .../_sd/deserialization/sd_deserialization.py | 575 ++++++++++++++++++ .../_sd/deserialization/sd_serialization.py | 273 +++++++++ src/someipy/_internal/_sd/entries/__init__.py | 0 .../_sd/entries/find_service_entry.py | 40 ++ .../_sd/entries/offer_service_entry.py | 53 ++ src/someipy/_internal/_sd/entries/sd_entry.py | 37 ++ .../_sd/entries/stop_offer_service_entry.py | 52 ++ .../stop_subscribe_eventgroup_entry.py | 57 ++ .../_sd/entries/subscribe_ack_entry.py | 49 ++ .../_sd/entries/subscribe_eventgroup_entry.py | 60 ++ .../_sd/entries/subscribe_eventgroup_nack.py | 46 ++ src/someipy/_internal/_sd/options/__init__.py | 0 .../_sd/options/configuration_option.py | 22 + src/someipy/_internal/_sd/options/endpoint.py | 34 ++ .../_internal/_sd/options/load_balancing.py | 23 + .../_internal/_sd/options/multicast.py | 34 ++ .../_internal/_sd/options/sd_endpoint.py | 30 + src/someipy/_internal/_sd/sd_message.py | 14 + src/someipy/_internal/_sd/service_instance.py | 33 + src/someipy/_internal/utils.py | 4 +- src/someipy/someipyd.py | 451 +++++++------- tests/sd/__init__.py | 0 tests/sd/test_offer_service_entry.py | 46 ++ tests/sd/test_sd_deserialization.py | 185 ++++++ tests/sd/test_sd_serialization.py | 187 ++++++ tests/test_sd_service_instance.py | 41 ++ tests/test_someipyd.py | 85 +++ 36 files changed, 2440 insertions(+), 232 deletions(-) create mode 100644 .coveragerc create mode 100644 src/someipy/_internal/_common/endpoint.py create mode 100644 src/someipy/_internal/_common/event.py create mode 100644 src/someipy/_internal/_daemon/daemon_server.py create mode 100644 src/someipy/_internal/_daemon/daemon_server_client.py create mode 100644 src/someipy/_internal/_sd/__init__.py create mode 100644 src/someipy/_internal/_sd/deserialization/__init__.py create mode 100644 src/someipy/_internal/_sd/deserialization/sd_deserialization.py create mode 100644 src/someipy/_internal/_sd/deserialization/sd_serialization.py create mode 100644 src/someipy/_internal/_sd/entries/__init__.py create mode 100644 src/someipy/_internal/_sd/entries/find_service_entry.py create mode 100644 src/someipy/_internal/_sd/entries/offer_service_entry.py create mode 100644 src/someipy/_internal/_sd/entries/sd_entry.py create mode 100644 src/someipy/_internal/_sd/entries/stop_offer_service_entry.py create mode 100644 src/someipy/_internal/_sd/entries/stop_subscribe_eventgroup_entry.py create mode 100644 src/someipy/_internal/_sd/entries/subscribe_ack_entry.py create mode 100644 src/someipy/_internal/_sd/entries/subscribe_eventgroup_entry.py create mode 100644 src/someipy/_internal/_sd/entries/subscribe_eventgroup_nack.py create mode 100644 src/someipy/_internal/_sd/options/__init__.py create mode 100644 src/someipy/_internal/_sd/options/configuration_option.py create mode 100644 src/someipy/_internal/_sd/options/endpoint.py create mode 100644 src/someipy/_internal/_sd/options/load_balancing.py create mode 100644 src/someipy/_internal/_sd/options/multicast.py create mode 100644 src/someipy/_internal/_sd/options/sd_endpoint.py create mode 100644 src/someipy/_internal/_sd/sd_message.py create mode 100644 src/someipy/_internal/_sd/service_instance.py create mode 100644 tests/sd/__init__.py create mode 100644 tests/sd/test_offer_service_entry.py create mode 100644 tests/sd/test_sd_deserialization.py create mode 100644 tests/sd/test_sd_serialization.py create mode 100644 tests/test_sd_service_instance.py create mode 100644 tests/test_someipyd.py diff --git a/.coveragerc b/.coveragerc new file mode 100644 index 0000000..c387873 --- /dev/null +++ b/.coveragerc @@ -0,0 +1,3 @@ +[run] +omit = + tests/* \ No newline at end of file diff --git a/.gitignore b/.gitignore index 863f233..2d64693 100644 --- a/.gitignore +++ b/.gitignore @@ -9,4 +9,6 @@ venv/ integration_tests/build integration_tests/install -build/ \ No newline at end of file +build/ + +.coverage \ No newline at end of file diff --git a/setup.cfg b/setup.cfg index 23a72f7..5096cf2 100644 --- a/setup.cfg +++ b/setup.cfg @@ -1,6 +1,6 @@ [metadata] name = someipy -version = 1.0.0 +version = 2.0.0 author = Christian H. author_email = someipy.package@gmail.com description = A Python package implementing the SOME/IP protocol diff --git a/src/someipy/_internal/_common/endpoint.py b/src/someipy/_internal/_common/endpoint.py new file mode 100644 index 0000000..554cadd --- /dev/null +++ b/src/someipy/_internal/_common/endpoint.py @@ -0,0 +1,33 @@ +# Copyright (C) 2024 Christian H. +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# This program is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with this program. If not, see . + +import ipaddress +from typing import NamedTuple, Union + + +class Endpoint(NamedTuple): + """Represents a network endpoint with IP address and port.""" + + ip: Union[ipaddress.IPv4Address, ipaddress.IPv6Address] + port: int + + @property + def is_ipv4(self): + return isinstance(self.ip, ipaddress.IPv4Address) + + def __str__(self): + if isinstance(self.ip, ipaddress.IPv6Address): + return f"[{self.ip}]:{self.port}" + return f"{self.ip}:{self.port}" diff --git a/src/someipy/_internal/_common/event.py b/src/someipy/_internal/_common/event.py new file mode 100644 index 0000000..9a474f3 --- /dev/null +++ b/src/someipy/_internal/_common/event.py @@ -0,0 +1,48 @@ +# Copyright (C) 2024 Christian H. +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# This program is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with this program. If not, see . + +import asyncio +from typing import TypeVar, Generic, Optional, Callable, Any + + +T = TypeVar("T") + + +class Event(Generic[T]): + def __init__(self): + self._handlers: list[Callable[[object, Optional[T]], Any]] = [] + + def add_handler(self, handler: Callable[[object, Optional[T]], Any]) -> None: + if handler not in self._handlers: + self._handlers.append(handler) + + def remove_handler(self, handler: Callable[[object, Optional[T]], Any]) -> None: + if handler in self._handlers: + self._handlers.remove(handler) + + async def invoke(self, sender: object, e: Optional[T] = None) -> None: + """Invoke all handlers, handling both sync and async automatically""" + for handler in self._handlers[:]: + result = handler(sender, e) + if asyncio.iscoroutine(result): + await result + + def __iadd__(self, handler: Callable[[object, Optional[T]], Any]) -> "Event[T]": + self.add_handler(handler) + return self + + def __isub__(self, handler: Callable[[object, Optional[T]], Any]) -> "Event[T]": + self.remove_handler(handler) + return self diff --git a/src/someipy/_internal/_daemon/daemon_server.py b/src/someipy/_internal/_daemon/daemon_server.py new file mode 100644 index 0000000..d3bde11 --- /dev/null +++ b/src/someipy/_internal/_daemon/daemon_server.py @@ -0,0 +1,60 @@ +import asyncio +import logging +import os +from someipy._internal._common.event import Event +from someipy._internal._daemon.daemon_server_client import DaemonServerClient + + +class ClientConnectedEventArgs: + def __init__(self, client: DaemonServerClient): + self.client = client + + +class DaemonServer: + + def __init__(self, logger: logging.Logger): + self._logger = logger + self.client_connected: Event[ClientConnectedEventArgs] = Event() + self.client_disconnected: Event[ClientConnectedEventArgs] = Event() + + async def _handle_client(self, reader, writer): + writer_id = id(writer) + self._logger.info(f"New client connected: {writer_id}") + + client = DaemonServerClient(reader, writer, writer_id, self._logger) + await self.client_connected.invoke(self, ClientConnectedEventArgs(client)) + + while True: + message = await client.read_next_message() + if message is None: + break # Client disconnected + + await self.client_disconnected.invoke(self, ClientConnectedEventArgs(client)) + + async def start( + self, + use_uds: bool = True, + socket_path: str | None = None, + tcp_port: int | None = None, + host: str = "127.0.0.1", + ): + if use_uds: + if os.path.exists(socket_path): + os.unlink(socket_path) + + self._server = await asyncio.start_unix_server( + self._handle_client, path=socket_path + ) + self._logger.info(f"Unix domain socket server started at {socket_path}") + else: + self._server = await asyncio.start_server( + self._handle_client, + host=host, + port=tcp_port, + reuse_port=True, + ) + self._logger.info(f"TCP server started at {host}:{tcp_port}") + + async def serve_forever(self): + async with self._server: + await self._server.serve_forever() diff --git a/src/someipy/_internal/_daemon/daemon_server_client.py b/src/someipy/_internal/_daemon/daemon_server_client.py new file mode 100644 index 0000000..bfe661b --- /dev/null +++ b/src/someipy/_internal/_daemon/daemon_server_client.py @@ -0,0 +1,91 @@ +import asyncio +import json +import logging +import struct +from typing import Optional + +from someipy._internal._common.event import Event + + +class DaemonServerClient: + + def __init__( + self, + reader: asyncio.StreamReader, + writer: asyncio.StreamWriter, + id: int, + logger: logging.Logger = None, + ): + self._reader = reader + self._writer = writer + self._id = id + self._logger = logger + self.message_received: Event[ClientMessageEventArgs] = Event() + + async def read_next_message(self) -> Optional[dict]: + wait_for_header = True + header_buffer = b"" + message_buffer = b"" + message_length = 0 + + while True: + if wait_for_header: + data = await self._reader.read(256 - len(header_buffer)) + if not data: + self._logger.debug( + f"Reading data returned none. Client disconnected." + ) + break # Client disconnected + + header_buffer += data + + if len(header_buffer) == 256: + try: + message_length = struct.unpack(" 256: + self._logger.error(f"Client sent too much header data.") + break + + else: + data = await self._reader.read(message_length - len(message_buffer)) + if not data: + self._logger.debug( + f"Reading data returned none. Client disconnected." + ) + break # Client disconnected + + message_buffer += data + + if len(message_buffer) == message_length: + self._logger.debug(f"Client sent message: {message_buffer}") + json_message = json.loads(message_buffer.decode("utf-8")) + await self.message_received.invoke( + self, ClientMessageEventArgs(self, json_message) + ) + return json_message + + elif len(message_buffer) > message_length: + self._logger.error(f"Client sent too much message data.") + raise Exception("Client sent too much message data.") + return None + + def send(self, message: bytes): + pass + + @property + def id(self) -> str: + return self._id + + +class ClientMessageEventArgs: + def __init__(self, client: DaemonServerClient, message: dict): + self.client = client + self.message = message diff --git a/src/someipy/_internal/_sd/__init__.py b/src/someipy/_internal/_sd/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/someipy/_internal/_sd/deserialization/__init__.py b/src/someipy/_internal/_sd/deserialization/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/someipy/_internal/_sd/deserialization/sd_deserialization.py b/src/someipy/_internal/_sd/deserialization/sd_deserialization.py new file mode 100644 index 0000000..cd0711b --- /dev/null +++ b/src/someipy/_internal/_sd/deserialization/sd_deserialization.py @@ -0,0 +1,575 @@ +# Copyright (C) 2024 Christian H. +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# This program is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with this program. If not, see . + + +from dataclasses import dataclass +from enum import Enum +import ipaddress +import socket +import struct + +from requests import options +from someipy._internal._sd.entries.find_service_entry import FindServiceEntry +from someipy._internal._sd.entries.offer_service_entry import OfferServiceEntry +from someipy._internal._sd.entries.stop_offer_service_entry import StopOfferServiceEntry +from someipy._internal._sd.entries.stop_subscribe_eventgroup_entry import ( + StopSubscribeEventGroupEntry, +) +from someipy._internal._sd.entries.subscribe_ack_entry import ( + SubscribeAckEventGroupEntry, +) +from someipy._internal._sd.entries.subscribe_eventgroup_entry import ( + SubscribeEventGroupEntry, +) +from someipy._internal._sd.entries.subscribe_eventgroup_nack import ( + SubscribeEventGroupNackEntry, +) +from someipy._internal._sd.options.configuration_option import ConfigurationOption +from someipy._internal._sd.options.endpoint import ( + IpV4EndpointOption, + IpV6EndpointOption, +) +from someipy._internal._sd.options.load_balancing import LoadBalancingOption +from someipy._internal._sd.options.multicast import ( + IpV4MulticastOption, + IpV6MulticastOption, +) +from someipy._internal._sd.options.sd_endpoint import ( + IpV4SdEndpointOption, + IpV6SdEndpointOption, +) +from someipy._internal.transport_layer_protocol import TransportLayerProtocol +from someipy._internal.utils import is_bit_set +from someipy._internal._sd.sd_message import SdMessage + +SERVICE_ID_SD = 0xFFFF +METHOD_ID_SD = 0x8100 +CLIENT_ID_SD = 0x0000 +PROTOCOL_VERSION_SD = 0x01 +INTERFACE_VERSION_SD = 0x01 +MESSAGE_TYPE_SD = 0x02 +RETURN_CODE_SD = 0x00 + +MINIMAL_HEADER_SIZE = 16 + + +class SdOptionOnWireType(Enum): + CONFIGURATION = 0x01 + LOAD_BALANCING = 0x02 + IPV4_ENDPOINT = 0x04 + IPV6_ENDPOINT = 0x06 + IPV4_MULTICAST = 0x14 + IPV6_MULTICAST = 0x16 + IPV4_SD_ENDPOINT = 0x24 + IPV6_SD_ENDPOINT = 0x26 + + +class SdEntryOnWireType(Enum): + FIND_SERVICE = 0x00 + OFFER_SERVICE = 0x01 + STOP_OFFER_SERVICE = 0x01 # with TTL to 0x000000 + SUBSCRIBE_EVENT_GROUP = 0x06 + STOP_SUBSCRIBE_EVENT_GROUP = 0x06 # with TTL to 0x000000 + SUBSCRIBE_EVENT_GROUP_ACK = 0x07 + SUBSCRIBE_EVENT_GROUP_NACK = 0x07 # with TTL to 0x000000 + + +def is_sd_message(data: bytes) -> bool: + if len(data) < MINIMAL_HEADER_SIZE: + return False + + service_id, method_id, length = struct.unpack(">HHI", data[0:8]) + + if length < 12: + return False + + if length > len(data) - 8: + return False + + ( + client_id, + session_id, + protocol_version, + interface_version, + message_type, + return_code, + ) = struct.unpack(">HHBBBB", data[8:16]) + + return ( + service_id == SERVICE_ID_SD + and method_id == METHOD_ID_SD + and client_id == CLIENT_ID_SD + and protocol_version == PROTOCOL_VERSION_SD + and interface_version == INTERFACE_VERSION_SD + and message_type == MESSAGE_TYPE_SD + and return_code == RETURN_CODE_SD + and session_id != 0 + ) + + +@dataclass +class CommonEntryData: + type_field_value: int + index_first_option: int + index_second_option: int + num_options_1: int + num_options_2: int + service_id: int + instance_id: int + major_version: int + ttl: int + + +@dataclass +class CommonOptionData: + option_length: int + option_type: SdOptionOnWireType + discardable_flag: bool + + +def deserialize_common_entry_data(data: bytes) -> CommonEntryData: + type_field_value, index_first_option, index_second_option = struct.unpack( + ">BBB", data[0:3] + ) + + num_options_1 = struct.unpack(">B", data[3:4])[0] # higher 4 bits + num_options_1 = (num_options_1 >> 4) & 0x0F + + num_options_2 = struct.unpack(">B", data[3:4])[0] # lower 4 bits + num_options_2 = num_options_2 & 0x0F + + service_id, instance_id, major_version = struct.unpack(">HHB", data[4:9]) + (ttl,) = struct.unpack(">I", data[8:12]) + ttl = ttl & 0xFFFFFF + + return CommonEntryData( + type_field_value, + index_first_option, + index_second_option, + num_options_1, + num_options_2, + service_id, + instance_id, + major_version, + ttl, + ) + + +def deserialize_common_option_data(data: bytes) -> CommonOptionData: + option_length, option_type, discardable_flag_value = struct.unpack( + ">HBB", data[0:4] + ) + option_type = SdOptionOnWireType(option_type) + discardable_flag = is_bit_set(discardable_flag_value, 7) + return CommonOptionData(option_length, option_type, discardable_flag) + + +def deserialize_ipv4_endpoint_option(data: bytes) -> IpV4EndpointOption: + ip1, ip2, ip3, ip4, _, protocol_value, port = struct.unpack(">BBBBBBH", data[0:8]) + address = ipaddress.IPv4Address(f"{ip1}.{ip2}.{ip3}.{ip4}") + protocol = TransportLayerProtocol(protocol_value) + return IpV4EndpointOption(address=address, protocol=protocol, port=port) + + +def deserialize_ipv6_endpoint_option(data: bytes) -> IpV6EndpointOption: + packed = data[0:16] + ipv6_str = socket.inet_ntop(socket.AF_INET6, packed) + address = ipaddress.IPv6Address(ipv6_str) + _, protocol_value, port = struct.unpack(">BBH", data[16:20]) + protocol = TransportLayerProtocol(protocol_value) + return IpV6EndpointOption(address=address, protocol=protocol, port=port) + + +def deserialize_ipv4_multicast_option(data: bytes) -> IpV4MulticastOption: + ip1, ip2, ip3, ip4, _, protocol_value, port = struct.unpack(">BBBBBBH", data[0:8]) + address = ipaddress.IPv4Address(f"{ip1}.{ip2}.{ip3}.{ip4}") + protocol = TransportLayerProtocol(protocol_value) + return IpV4MulticastOption(address=address, protocol=protocol, port=port) + + +def deserialize_ipv6_multicast_option(data: bytes) -> IpV6MulticastOption: + packed = data[0:16] + ipv6_str = socket.inet_ntop(socket.AF_INET6, packed) + address = ipaddress.IPv6Address(ipv6_str) + _, protocol_value, port = struct.unpack(">BBH", data[16:20]) + protocol = TransportLayerProtocol(protocol_value) + return IpV6MulticastOption(address=address, protocol=protocol, port=port) + + +def deserialize_ipv4_sd_endpoint_option(data: bytes) -> IpV4SdEndpointOption: + ip1, ip2, ip3, ip4, _, _, port = struct.unpack(">BBBBBBH", data[0:8]) + address = ipaddress.IPv4Address(f"{ip1}.{ip2}.{ip3}.{ip4}") + return IpV4SdEndpointOption(address=address, port=port) + + +def deserialize_ipv6_sd_endpoint_option(data: bytes) -> IpV6SdEndpointOption: + packed = data[0:16] + ipv6_str = socket.inet_ntop(socket.AF_INET6, packed) + address = ipaddress.IPv6Address(ipv6_str) + _, _, port = struct.unpack(">BBH", data[16:20]) + return IpV6SdEndpointOption(address=address, port=port) + + +def deserialize_load_balancing_option(data: bytes) -> LoadBalancingOption: + priority, weight = struct.unpack(">HH", data[0:4]) + return LoadBalancingOption(priority=priority, weight=weight) + + +def deserialize_sd_message( + data: bytes, source: str, port: int, multicast: bool +) -> SdMessage: + + if not is_sd_message(data): + raise ValueError("The provided data is not a valid SOME/IP-SD message.") + + service_id, method_id, length = struct.unpack(">HHI", data[0:8]) + if length <= 0: + raise ValueError(f"Length in SOME/IP header is <=0 ({length})") + + if length < 8: + raise ValueError(f"Length in SOME/IP header is <8 ({length})") + + ( + client_id, + session_id, + protocol_version, + interface_version, + message_type, + return_code, + ) = struct.unpack(">HHBBBB", data[8:16]) + + (flags,) = struct.unpack(">B", data[16:17]) + reboot_flag = is_bit_set(flags, 7) + unicast_flag = is_bit_set(flags, 6) + + # Constants for byte positions inside the SD header + SD_POSITION_ENTRY_LENGTH = 20 + SD_START_POSITION_ENTRIES = 24 + + # Constants for length of sections in the SD header + SD_SINGLE_ENTRY_LENGTH_BYTES = 16 + + (length_entries,) = struct.unpack( + ">I", data[SD_POSITION_ENTRY_LENGTH : SD_POSITION_ENTRY_LENGTH + 4] + ) + + number_of_entries = int(length_entries / SD_SINGLE_ENTRY_LENGTH_BYTES) + + pos_length_options = SD_POSITION_ENTRY_LENGTH + 4 + length_entries + (length_options,) = struct.unpack( + ">I", data[pos_length_options : pos_length_options + 4] + ) + pos_start_options = pos_length_options + 4 + + current_pos_option = pos_start_options + bytes_options_left = length_options + + options = [] + while bytes_options_left > 0: + + common_option_data = deserialize_common_option_data( + data[current_pos_option : current_pos_option + 4] + ) + + option_type = common_option_data.option_type + + if option_type == SdOptionOnWireType.IPV4_ENDPOINT: + sd_option = deserialize_ipv4_endpoint_option( + data[current_pos_option + 4 : current_pos_option + 12] + ) + options.append(sd_option) + + elif option_type == SdOptionOnWireType.IPV6_ENDPOINT: + sd_option = deserialize_ipv6_endpoint_option( + data[current_pos_option + 4 : current_pos_option + 24] + ) + options.append(sd_option) + elif option_type == SdOptionOnWireType.IPV4_MULTICAST: + sd_option = deserialize_ipv4_multicast_option( + data[current_pos_option + 4 : current_pos_option + 12] + ) + options.append(sd_option) + elif option_type == SdOptionOnWireType.IPV6_MULTICAST: + sd_option = deserialize_ipv6_multicast_option( + data[current_pos_option + 4 : current_pos_option + 24] + ) + options.append(sd_option) + elif option_type == SdOptionOnWireType.IPV4_SD_ENDPOINT: + sd_option = deserialize_ipv4_sd_endpoint_option( + data[current_pos_option + 4 : current_pos_option + 12] + ) + options.append(sd_option) + elif option_type == SdOptionOnWireType.IPV6_SD_ENDPOINT: + sd_option = deserialize_ipv6_sd_endpoint_option( + data[current_pos_option + 4 : current_pos_option + 24] + ) + options.append(sd_option) + elif option_type == SdOptionOnWireType.CONFIGURATION: + dummy_config_option = ConfigurationOption() + options.append(dummy_config_option) + elif option_type == SdOptionOnWireType.LOAD_BALANCING: + sd_option = deserialize_load_balancing_option( + data[current_pos_option + 4 : current_pos_option + 8] + ) + options.append(sd_option) + + # Subtract 3 bytes first for length and type + bytes_options_left -= common_option_data.option_length + 3 + current_pos_option += common_option_data.option_length + 3 + + # Read in all Service and Event Group entries + entries = [] + for i in range(number_of_entries): + start_entry = SD_START_POSITION_ENTRIES + (i * SD_SINGLE_ENTRY_LENGTH_BYTES) + end_entry = start_entry + SD_SINGLE_ENTRY_LENGTH_BYTES + + common_entry_data = deserialize_common_entry_data( + data[start_entry : start_entry + 12] + ) + + if ( + common_entry_data.type == SdEntryOnWireType.OFFER_SERVICE.value + and common_entry_data.ttl != 0 + ): + (minor_version,) = struct.unpack( + ">I", data[start_entry + 12 : start_entry + 16] + ) + + applicable_options = [] + for j in range(common_entry_data.num_options_1): + applicable_options.append( + options[common_entry_data.index_first_option + j] + ) + for j in range(common_entry_data.num_options_2): + applicable_options.append( + options[common_entry_data.index_second_option + j] + ) + + offer_service_entry = OfferServiceEntry( + service_id=common_entry_data.service_id, + instance_id=common_entry_data.instance_id, + major_version=common_entry_data.major_version, + minor_version=minor_version, + ttl=common_entry_data.ttl, + ip_v4_endpoints=[ + o for o in applicable_options if isinstance(o, IpV4EndpointOption) + ], + ip_v6_endpoints=[ + o for o in applicable_options if isinstance(o, IpV6EndpointOption) + ], + ) + entries.append(offer_service_entry) + + elif ( + common_entry_data.type == SdEntryOnWireType.STOP_OFFER_SERVICE.value + and common_entry_data.ttl == 0 + ): + (minor_version,) = struct.unpack( + ">I", data[start_entry + 12 : start_entry + 16] + ) + + applicable_options = [] + for j in range(common_entry_data.num_options_1): + applicable_options.append( + options[common_entry_data.index_first_option + j] + ) + for j in range(common_entry_data.num_options_2): + applicable_options.append( + options[common_entry_data.index_second_option + j] + ) + + stop_offer_service_entry = StopOfferServiceEntry( + service_id=common_entry_data.service_id, + instance_id=common_entry_data.instance_id, + major_version=common_entry_data.major_version, + minor_version=minor_version, + ttl=common_entry_data.ttl, + ip_v4_endpoints=[ + o for o in applicable_options if isinstance(o, IpV4EndpointOption) + ], + ip_v6_endpoints=[ + o for o in applicable_options if isinstance(o, IpV6EndpointOption) + ], + ) + entries.append(stop_offer_service_entry) + + elif common_entry_data.type == SdEntryOnWireType.FIND_SERVICE.value: + (minor_version,) = struct.unpack( + ">I", data[start_entry + 12 : start_entry + 16] + ) + + find_service_entry = FindServiceEntry( + service_id=common_entry_data.service_id, + instance_id=common_entry_data.instance_id, + major_version=common_entry_data.major_version, + minor_version=minor_version, + ttl=common_entry_data.ttl, + ) + entries.append(find_service_entry) + + elif ( + common_entry_data.type == SdEntryOnWireType.SUBSCRIBE_EVENT_GROUP.value + and common_entry_data.ttl != 0 + ): + initial_data_requested_flag_counter_value, eventgroup_id = struct.unpack( + ">BH", data[start_entry + 13 : start_entry + 16] + ) + initial_data_requested_flag = is_bit_set( + initial_data_requested_flag_counter_value, 7 + ) + counter = initial_data_requested_flag_counter_value & 0xF + + applicable_options = [] + for j in range(common_entry_data.num_options_1): + applicable_options.append( + options[common_entry_data.index_first_option + j] + ) + for j in range(common_entry_data.num_options_2): + applicable_options.append( + options[common_entry_data.index_second_option + j] + ) + + subscribe_eventgroup_entry = SubscribeEventGroupEntry( + service_id=common_entry_data.service_id, + instance_id=common_entry_data.instance_id, + major_version=common_entry_data.major_version, + minor_version=minor_version, + ttl=common_entry_data.ttl, + eventgroup_id=eventgroup_id, + counter=counter, + ip_v4_endpoints=[ + o for o in applicable_options if isinstance(o, IpV4EndpointOption) + ], + ip_v6_endpoints=[ + o for o in applicable_options if isinstance(o, IpV6EndpointOption) + ], + ) + entries.append(subscribe_eventgroup_entry) + + elif ( + common_entry_data.type == SdEntryOnWireType.STOP_SUBSCRIBE_EVENT_GROUP.value + and common_entry_data.ttl == 0 + ): + initial_data_requested_flag_counter_value, eventgroup_id = struct.unpack( + ">BH", data[start_entry + 13 : start_entry + 16] + ) + initial_data_requested_flag = is_bit_set( + initial_data_requested_flag_counter_value, 7 + ) + counter = initial_data_requested_flag_counter_value & 0xF + + applicable_options = [] + for j in range(common_entry_data.num_options_1): + applicable_options.append( + options[common_entry_data.index_first_option + j] + ) + for j in range(common_entry_data.num_options_2): + applicable_options.append( + options[common_entry_data.index_second_option + j] + ) + + stop_subscribe_eventgroup_entry = StopSubscribeEventGroupEntry( + service_id=common_entry_data.service_id, + instance_id=common_entry_data.instance_id, + major_version=common_entry_data.major_version, + minor_version=minor_version, + ttl=common_entry_data.ttl, + ip_v4_endpoints=[ + o for o in applicable_options if isinstance(o, IpV4EndpointOption) + ], + ip_v6_endpoints=[ + o for o in applicable_options if isinstance(o, IpV6EndpointOption) + ], + ) + entries.append(stop_subscribe_eventgroup_entry) + + elif ( + common_entry_data.type == SdEntryOnWireType.SUBSCRIBE_EVENT_GROUP_ACK.value + and common_entry_data.ttl != 0 + ): + initial_data_requested_flag_counter_value, eventgroup_id = struct.unpack( + ">BH", data[start_entry + 13 : start_entry + 16] + ) + initial_data_requested_flag = is_bit_set( + initial_data_requested_flag_counter_value, 7 + ) + counter = initial_data_requested_flag_counter_value & 0xF + + applicable_options = [] + for j in range(common_entry_data.num_options_1): + applicable_options.append( + options[common_entry_data.index_first_option + j] + ) + for j in range(common_entry_data.num_options_2): + applicable_options.append( + options[common_entry_data.index_second_option + j] + ) + + subscribe_ack_eventgroup_entry = SubscribeAckEventGroupEntry( + service_id=common_entry_data.service_id, + instance_id=common_entry_data.instance_id, + major_version=common_entry_data.major_version, + minor_version=minor_version, + ttl=common_entry_data.ttl, + eventgroup_id=eventgroup_id, + counter=counter, + ip_v4_endpoints=[ + o for o in applicable_options if isinstance(o, IpV4EndpointOption) + ], + ip_v6_endpoints=[ + o for o in applicable_options if isinstance(o, IpV6EndpointOption) + ], + ) + entries.append(subscribe_ack_eventgroup_entry) + + elif ( + common_entry_data.type == SdEntryOnWireType.SUBSCRIBE_EVENT_GROUP_NACK.value + and common_entry_data.ttl == 0 + ): + initial_data_requested_flag_counter_value, eventgroup_id = struct.unpack( + ">BH", data[start_entry + 13 : start_entry + 16] + ) + initial_data_requested_flag = is_bit_set( + initial_data_requested_flag_counter_value, 7 + ) + counter = initial_data_requested_flag_counter_value & 0xF + + applicable_options = [] + for j in range(common_entry_data.num_options_1): + applicable_options.append( + options[common_entry_data.index_first_option + j] + ) + for j in range(common_entry_data.num_options_2): + applicable_options.append( + options[common_entry_data.index_second_option + j] + ) + + subscribe_ack_eventgroup_entry = SubscribeEventGroupNackEntry( + service_id=common_entry_data.service_id, + instance_id=common_entry_data.instance_id, + major_version=common_entry_data.major_version, + minor_version=minor_version, + eventgroup_id=eventgroup_id, + counter=counter, + ) + entries.append(subscribe_ack_eventgroup_entry) + + sd_message = SdMessage() + sd_message.source = source + sd_message.source_port = port + sd_message.multicast = multicast + sd_message.session_id = session_id + sd_message.entries = entries + return sd_message diff --git a/src/someipy/_internal/_sd/deserialization/sd_serialization.py b/src/someipy/_internal/_sd/deserialization/sd_serialization.py new file mode 100644 index 0000000..d886924 --- /dev/null +++ b/src/someipy/_internal/_sd/deserialization/sd_serialization.py @@ -0,0 +1,273 @@ +# Copyright (C) 2024 Christian H. +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# This program is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with this program. If not, see . + + +import socket +import struct + +from someipy._internal._sd.options.endpoint import ( + IpV4EndpointOption, + IpV6EndpointOption, +) +from someipy._internal._sd.entries.sd_entry import SdEntryType +from someipy._internal.utils import set_bit_at_position +from someipy._internal._sd.sd_message import SdMessage + +SERVICE_ID_SD = 0xFFFF +METHOD_ID_SD = 0x8100 +CLIENT_ID_SD = 0x0000 +PROTOCOL_VERSION_SD = 0x01 +INTERFACE_VERSION_SD = 0x01 +MESSAGE_TYPE_SD = 0x02 +RETURN_CODE_SD = 0x00 + +MINIMAL_HEADER_SIZE = 16 + + +def entry_type_to_on_wire(entry_type: SdEntryType) -> int: + lookup = { + SdEntryType.FIND_SERVICE: 0x00, + SdEntryType.OFFER_SERVICE: 0x01, + SdEntryType.STOP_OFFER_SERVICE: 0x01, # with TTL to 0x000000 + SdEntryType.SUBSCRIBE_EVENT_GROUP: 0x06, + SdEntryType.STOP_SUBSCRIBE_EVENT_GROUP: 0x06, # with TTL to 0x000000 + SdEntryType.SUBSCRIBE_EVENT_GROUP_ACK: 0x07, + SdEntryType.SUBSCRIBE_EVENT_GROUP_NACK: 0x07, + } + if entry_type not in lookup: + raise ValueError(f"Unsupported SdEntryType: {entry_type}") + + return lookup[entry_type] + + +def serialize_ipv4_endpoint_option(option: IpV4EndpointOption) -> bytes: + output = bytes() + LENGTH = 0x0009 + TYPE = 0x04 + + discardable_flag_value = set_bit_at_position(0, 7, False) + output += struct.pack(">HBB", LENGTH, TYPE, discardable_flag_value) + output += struct.pack( + ">IBBH", int(option.address), 0, option.protocol.value, option.port + ) + return output + + +def serialize_ipv6_endpoint_option(option: IpV6EndpointOption) -> bytes: + output = bytes() + LENGTH = 0x0015 + TYPE = 0x06 + + ipv6_str = str(option.address) + packed_ip = socket.inet_pton(socket.AF_INET6, ipv6_str) + + discardable_flag_value = set_bit_at_position(0, 7, False) + output += struct.pack(">HBB", LENGTH, TYPE, discardable_flag_value) + output += packed_ip + output += struct.pack(">BBH", 0, option.protocol.value, option.port) + return output + + +def serialize_sd_message(sd_message: SdMessage) -> bytes: + + output = bytes() + + SERVICE_ID_SD = 0xFFFF + METHOD_ID_SD = 0x8100 + CLIENT_ID_SD = 0x0000 + PROTOCOL_VERSION_SD = 0x01 + INTERFACE_VERSION_SD = 0x01 + MESSAGE_TYPE_SD = 0x02 + RETURN_CODE_SD = 0x00 + + output += struct.pack( + ">HHIHHBBBB", + SERVICE_ID_SD, + METHOD_ID_SD, + 0, + CLIENT_ID_SD, + sd_message.session_id, + PROTOCOL_VERSION_SD, + INTERFACE_VERSION_SD, + MESSAGE_TYPE_SD, + RETURN_CODE_SD, + ) + + # TODO: Set the reboot flag properly + reboot_flag = False + unicast_flag = True + + flags = 0 + flags = set_bit_at_position(flags, 31, reboot_flag) + flags = set_bit_at_position(flags, 30, unicast_flag) + + output += struct.pack(">I", flags) # 8 bit flags + 24 reserved bits + + options = [] + option_set = set() + + for entry in sd_message.entries: + if entry.type in [ + SdEntryType.OFFER_SERVICE, + SdEntryType.STOP_OFFER_SERVICE, + SdEntryType.SUBSCRIBE_EVENT_GROUP, + SdEntryType.STOP_SUBSCRIBE_EVENT_GROUP, + ]: + for endpoint in entry.ip_v4_endpoints: + if endpoint not in option_set: + option_set.add(endpoint) + options.append(endpoint) + for endpoint in entry.ip_v6_endpoints: + if endpoint not in option_set: + option_set.add(endpoint) + options.append(endpoint) + + # Length of the entries array + SD_SINGLE_ENTRY_LENGTH_BYTES = 16 + length_entries_array = len(sd_message.entries) * SD_SINGLE_ENTRY_LENGTH_BYTES + output += struct.pack(">I", length_entries_array) + + for entry in sd_message.entries: + if entry.type in [ + SdEntryType.OFFER_SERVICE, + SdEntryType.STOP_OFFER_SERVICE, + SdEntryType.FIND_SERVICE, + ]: + if entry.type == SdEntryType.OFFER_SERVICE: + ttl = entry.ttl + else: + ttl = 0 + + ttl_high = (ttl & 0xFF0000) >> 16 + ttl_low = ttl & 0xFFFF + + index_first_option = 0 + + if entry.type == SdEntryType.FIND_SERVICE: + num_options_1 = 0 + index_first_option = 0 + else: + num_options_1 = len(entry.ip_v4_endpoints) + len(entry.ip_v6_endpoints) + if num_options_1 > 2: + raise ValueError("Too many options for entry configured.") + + entry_options = set(entry.ip_v4_endpoints + entry.ip_v6_endpoints) + + for i in range(len(options)): + if options[i] in entry_options: + index_first_option = i + break + + num_options_2 = 0 # No second option in this implementation + num_options = (num_options_1 << 4) | num_options_2 + + index_second_option = 0 # No second option in this implementation + + output += struct.pack( + ">BBBBHHBBH", + entry_type_to_on_wire(entry.type), + index_first_option, + index_second_option, + num_options, + entry.service_id, + entry.instance_id, + entry.major_version, + ttl_high, + ttl_low, + ) + + output += struct.pack(">I", entry.minor_version) + + else: + if entry.type in [ + SdEntryType.STOP_SUBSCRIBE_EVENT_GROUP, + SdEntryType.SUBSCRIBE_EVENT_GROUP_NACK, + ]: + ttl = 0 + else: + ttl = entry.ttl + + ttl_high = (ttl & 0xFF0000) >> 16 + ttl_low = ttl & 0xFFFF + + if entry.type in [ + SdEntryType.SUBSCRIBE_EVENT_GROUP, + SdEntryType.STOP_SUBSCRIBE_EVENT_GROUP, + ]: + num_options_1 = len(entry.ip_v4_endpoints) + len(entry.ip_v6_endpoints) + if num_options_1 > 2: + raise ValueError("Too many options for entry configured.") + + entry_options = set(entry.ip_v4_endpoints + entry.ip_v6_endpoints) + + for i in range(len(options)): + if options[i] in entry_options: + index_first_option = i + break + else: + num_options_1 = 0 + index_first_option = 0 + + num_options_2 = 0 # No second option in this implementation + num_options = (num_options_1 << 4) | num_options_2 + + index_second_option = 0 # No second option in this implementation + + output += struct.pack( + ">BBBBHHBBH", + entry_type_to_on_wire(entry.type), + index_first_option, + index_second_option, + num_options, + entry.service_id, + entry.instance_id, + entry.major_version, + ttl_high, + ttl_low, + ) + + initial_data_requested_flag_counter_value = set_bit_at_position(0, 7, True) + initial_data_requested_flag_counter_value = ( + initial_data_requested_flag_counter_value | (entry.counter & 0xF) + ) + output += struct.pack( + ">BBH", + 0, + initial_data_requested_flag_counter_value, + entry.eventgroup_id, + ) + + length_of_options = 0 + + for option in options: + if isinstance(option, IpV4EndpointOption): + length_of_options += 12 + elif isinstance(option, IpV6EndpointOption): + length_of_options += 24 + + output += struct.pack(">I", length_of_options) + + for option in options: + if isinstance(option, IpV4EndpointOption): + output += serialize_ipv4_endpoint_option(option) + elif isinstance(option, IpV6EndpointOption): + output += serialize_ipv6_endpoint_option(option) + + total_length = len(output) - 8 + total_length_bytes = struct.pack(">I", total_length) + + output = output[:4] + total_length_bytes + output[8:] + + return output diff --git a/src/someipy/_internal/_sd/entries/__init__.py b/src/someipy/_internal/_sd/entries/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/someipy/_internal/_sd/entries/find_service_entry.py b/src/someipy/_internal/_sd/entries/find_service_entry.py new file mode 100644 index 0000000..3884a81 --- /dev/null +++ b/src/someipy/_internal/_sd/entries/find_service_entry.py @@ -0,0 +1,40 @@ +# Copyright (C) 2024 Christian H. +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# This program is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with this program. If not, see . + + +from dataclasses import dataclass + +from someipy._internal._sd.entries.sd_entry import SdEntry, SdEntryType + + +@dataclass +class FindServiceEntry(SdEntry): + service_id: int + instance_id: int + major_version: int + minor_version: int + + def __init__( + self, + service_id: int, + instance_id: int, + major_version: int, + minor_version: int, + ): + super().__init__(SdEntryType.FIND_SERVICE) + self.service_id = service_id + self.instance_id = instance_id + self.major_version = major_version + self.minor_version = minor_version diff --git a/src/someipy/_internal/_sd/entries/offer_service_entry.py b/src/someipy/_internal/_sd/entries/offer_service_entry.py new file mode 100644 index 0000000..f1cb24f --- /dev/null +++ b/src/someipy/_internal/_sd/entries/offer_service_entry.py @@ -0,0 +1,53 @@ +# Copyright (C) 2024 Christian H. +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# This program is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with this program. If not, see . + +from dataclasses import dataclass +from typing import List + +from someipy._internal._sd.entries.sd_entry import SdEntry, SdEntryType +from someipy._internal._sd.options.endpoint import ( + IpV4EndpointOption, + IpV6EndpointOption, +) + + +@dataclass +class OfferServiceEntry(SdEntry): + service_id: int + instance_id: int + major_version: int + minor_version: int + ttl: int + ip_v4_endpoints: List[IpV4EndpointOption] + ip_v6_endpoints: List[IpV6EndpointOption] + + def __init__( + self, + service_id: int, + instance_id: int, + major_version: int, + minor_version: int, + ttl: int, + ip_v4_endpoints: List[IpV4EndpointOption], + ip_v6_endpoints: List[IpV6EndpointOption], + ): + super().__init__(SdEntryType.OFFER_SERVICE) + self.service_id = service_id + self.instance_id = instance_id + self.major_version = major_version + self.minor_version = minor_version + self.ttl = ttl + self.ip_v4_endpoints = ip_v4_endpoints + self.ip_v6_endpoints = ip_v6_endpoints diff --git a/src/someipy/_internal/_sd/entries/sd_entry.py b/src/someipy/_internal/_sd/entries/sd_entry.py new file mode 100644 index 0000000..a8bbbba --- /dev/null +++ b/src/someipy/_internal/_sd/entries/sd_entry.py @@ -0,0 +1,37 @@ +# Copyright (C) 2024 Christian H. +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# This program is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with this program. If not, see . + + +from enum import Enum, unique + + +@unique +class SdEntryType(Enum): + FIND_SERVICE = 0 + OFFER_SERVICE = 1 + STOP_OFFER_SERVICE = 2 + SUBSCRIBE_EVENT_GROUP = 3 + STOP_SUBSCRIBE_EVENT_GROUP = 4 + SUBSCRIBE_EVENT_GROUP_ACK = 5 + SUBSCRIBE_EVENT_GROUP_NACK = 6 + + +class SdEntry: + def __init__(self, entry_type: SdEntryType): + self._entry_type = entry_type + + @property + def type(self) -> SdEntryType: + return self._entry_type diff --git a/src/someipy/_internal/_sd/entries/stop_offer_service_entry.py b/src/someipy/_internal/_sd/entries/stop_offer_service_entry.py new file mode 100644 index 0000000..39421c2 --- /dev/null +++ b/src/someipy/_internal/_sd/entries/stop_offer_service_entry.py @@ -0,0 +1,52 @@ +# Copyright (C) 2024 Christian H. +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# This program is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with this program. If not, see . + + +from dataclasses import dataclass +from typing import List + +from someipy._internal._sd.entries.sd_entry import SdEntry, SdEntryType +from someipy._internal._sd.options.endpoint import ( + IpV4EndpointOption, + IpV6EndpointOption, +) + + +@dataclass +class StopOfferServiceEntry(SdEntry): + service_id: int + instance_id: int + major_version: int + minor_version: int + ttl: int + ip_v4_endpoints: List[IpV4EndpointOption] + ip_v6_endpoints: List[IpV6EndpointOption] + + def __init__( + self, + service_id: int, + instance_id: int, + major_version: int, + minor_version: int, + ip_v4_endpoints: List[IpV4EndpointOption], + ip_v6_endpoints: List[IpV6EndpointOption], + ): + super().__init__(SdEntryType.STOP_OFFER_SERVICE) + self.service_id = service_id + self.instance_id = instance_id + self.major_version = major_version + self.minor_version = minor_version + self.ip_v4_endpoints = ip_v4_endpoints + self.ip_v6_endpoints = ip_v6_endpoints diff --git a/src/someipy/_internal/_sd/entries/stop_subscribe_eventgroup_entry.py b/src/someipy/_internal/_sd/entries/stop_subscribe_eventgroup_entry.py new file mode 100644 index 0000000..946f6bf --- /dev/null +++ b/src/someipy/_internal/_sd/entries/stop_subscribe_eventgroup_entry.py @@ -0,0 +1,57 @@ +# Copyright (C) 2024 Christian H. +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# This program is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with this program. If not, see . + + +from dataclasses import dataclass +from typing import List + +from someipy._internal._sd.entries.sd_entry import SdEntry, SdEntryType +from someipy._internal._sd.options.endpoint import ( + IpV4EndpointOption, + IpV6EndpointOption, +) + + +@dataclass +class StopSubscribeEventGroupEntry(SdEntry): + service_id: int + instance_id: int + major_version: int + minor_version: int + eventgroup_id: int + counter: int + ip_v4_endpoints: List[IpV4EndpointOption] + ip_v6_endpoints: List[IpV6EndpointOption] + + def __init__( + self, + service_id: int, + instance_id: int, + major_version: int, + minor_version: int, + eventgroup_id: int, + counter: int, + ip_v4_endpoints: List[IpV4EndpointOption], + ip_v6_endpoints: List[IpV6EndpointOption], + ): + super().__init__(SdEntryType.STOP_SUBSCRIBE_EVENT_GROUP) + self.service_id = service_id + self.instance_id = instance_id + self.major_version = major_version + self.minor_version = minor_version + self.eventgroup_id = eventgroup_id + self.counter = counter + self.ip_v4_endpoints = ip_v4_endpoints + self.ip_v6_endpoints = ip_v6_endpoints diff --git a/src/someipy/_internal/_sd/entries/subscribe_ack_entry.py b/src/someipy/_internal/_sd/entries/subscribe_ack_entry.py new file mode 100644 index 0000000..6550b90 --- /dev/null +++ b/src/someipy/_internal/_sd/entries/subscribe_ack_entry.py @@ -0,0 +1,49 @@ +# Copyright (C) 2024 Christian H. +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# This program is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with this program. If not, see . + + +from dataclasses import dataclass + +from someipy._internal._sd.entries.sd_entry import SdEntry, SdEntryType + + +@dataclass +class SubscribeAckEventGroupEntry(SdEntry): + service_id: int + instance_id: int + major_version: int + minor_version: int + ttl: int + eventgroup_id: int + counter: int + + def __init__( + self, + service_id: int, + instance_id: int, + major_version: int, + minor_version: int, + ttl: int, + eventgroup_id: int, + counter: int, + ): + super().__init__(SdEntryType.SUBSCRIBE_EVENT_GROUP_ACK) + self.service_id = service_id + self.instance_id = instance_id + self.major_version = major_version + self.minor_version = minor_version + self.ttl = ttl + self.eventgroup_id = eventgroup_id + self.counter = counter diff --git a/src/someipy/_internal/_sd/entries/subscribe_eventgroup_entry.py b/src/someipy/_internal/_sd/entries/subscribe_eventgroup_entry.py new file mode 100644 index 0000000..627dac6 --- /dev/null +++ b/src/someipy/_internal/_sd/entries/subscribe_eventgroup_entry.py @@ -0,0 +1,60 @@ +# Copyright (C) 2024 Christian H. +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# This program is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with this program. If not, see . + + +from dataclasses import dataclass +from typing import List + +from someipy._internal._sd.entries.sd_entry import SdEntry, SdEntryType +from someipy._internal._sd.options.endpoint import ( + IpV4EndpointOption, + IpV6EndpointOption, +) + + +@dataclass +class SubscribeEventGroupEntry(SdEntry): + service_id: int + instance_id: int + major_version: int + minor_version: int + ttl: int + eventgroup_id: int + counter: int + ip_v4_endpoints: List[IpV4EndpointOption] + ip_v6_endpoints: List[IpV6EndpointOption] + + def __init__( + self, + service_id: int, + instance_id: int, + major_version: int, + minor_version: int, + ttl: int, + eventgroup_id: int, + counter: int, + ip_v4_endpoints: List[IpV4EndpointOption], + ip_v6_endpoints: List[IpV6EndpointOption], + ): + super().__init__(SdEntryType.SUBSCRIBE_EVENT_GROUP) + self.service_id = service_id + self.instance_id = instance_id + self.major_version = major_version + self.minor_version = minor_version + self.ttl = ttl + self.eventgroup_id = eventgroup_id + self.counter = counter + self.ip_v4_endpoints = ip_v4_endpoints + self.ip_v6_endpoints = ip_v6_endpoints diff --git a/src/someipy/_internal/_sd/entries/subscribe_eventgroup_nack.py b/src/someipy/_internal/_sd/entries/subscribe_eventgroup_nack.py new file mode 100644 index 0000000..42c16ce --- /dev/null +++ b/src/someipy/_internal/_sd/entries/subscribe_eventgroup_nack.py @@ -0,0 +1,46 @@ +# Copyright (C) 2024 Christian H. +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# This program is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with this program. If not, see . + + +from dataclasses import dataclass + +from someipy._internal._sd.entries.sd_entry import SdEntry, SdEntryType + + +@dataclass +class SubscribeEventGroupNackEntry(SdEntry): + service_id: int + instance_id: int + major_version: int + minor_version: int + eventgroup_id: int + counter: int + + def __init__( + self, + service_id: int, + instance_id: int, + major_version: int, + minor_version: int, + eventgroup_id: int, + counter: int, + ): + super().__init__(SdEntryType.SUBSCRIBE_EVENT_GROUP_NACK) + self.service_id = service_id + self.instance_id = instance_id + self.major_version = major_version + self.minor_version = minor_version + self.eventgroup_id = eventgroup_id + self.counter = counter diff --git a/src/someipy/_internal/_sd/options/__init__.py b/src/someipy/_internal/_sd/options/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/someipy/_internal/_sd/options/configuration_option.py b/src/someipy/_internal/_sd/options/configuration_option.py new file mode 100644 index 0000000..fd8d46a --- /dev/null +++ b/src/someipy/_internal/_sd/options/configuration_option.py @@ -0,0 +1,22 @@ +# Copyright (C) 2024 Christian H. +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# This program is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with this program. If not, see . + +from dataclasses import dataclass + + +@dataclass +class ConfigurationOption: + key: str + value: str diff --git a/src/someipy/_internal/_sd/options/endpoint.py b/src/someipy/_internal/_sd/options/endpoint.py new file mode 100644 index 0000000..a00334a --- /dev/null +++ b/src/someipy/_internal/_sd/options/endpoint.py @@ -0,0 +1,34 @@ +# Copyright (C) 2024 Christian H. +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# This program is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with this program. If not, see . + + +from dataclasses import dataclass +import ipaddress + +from someipy._internal.transport_layer_protocol import TransportLayerProtocol + + +@dataclass(eq=True, frozen=True) +class IpV4EndpointOption: + address: ipaddress.IPv4Address + protocol: TransportLayerProtocol + port: int + + +@dataclass(eq=True, frozen=True) +class IpV6EndpointOption: + address: ipaddress.IPv6Address + protocol: TransportLayerProtocol + port: int diff --git a/src/someipy/_internal/_sd/options/load_balancing.py b/src/someipy/_internal/_sd/options/load_balancing.py new file mode 100644 index 0000000..b8ba81a --- /dev/null +++ b/src/someipy/_internal/_sd/options/load_balancing.py @@ -0,0 +1,23 @@ +# Copyright (C) 2024 Christian H. +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# This program is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with this program. If not, see . + + +from dataclasses import dataclass + + +@dataclass +class LoadBalancingOption: + priority: int + weight: int diff --git a/src/someipy/_internal/_sd/options/multicast.py b/src/someipy/_internal/_sd/options/multicast.py new file mode 100644 index 0000000..c6c7c55 --- /dev/null +++ b/src/someipy/_internal/_sd/options/multicast.py @@ -0,0 +1,34 @@ +# Copyright (C) 2024 Christian H. +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# This program is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with this program. If not, see . + + +from dataclasses import dataclass +import ipaddress + +from someipy._internal.transport_layer_protocol import TransportLayerProtocol + + +@dataclass +class IpV4MulticastOption: + address: ipaddress.IPv4Address + protocol: TransportLayerProtocol + port: int + + +@dataclass +class IpV6MulticastOption: + address: ipaddress.IPv6Address + protocol: TransportLayerProtocol + port: int diff --git a/src/someipy/_internal/_sd/options/sd_endpoint.py b/src/someipy/_internal/_sd/options/sd_endpoint.py new file mode 100644 index 0000000..cc89258 --- /dev/null +++ b/src/someipy/_internal/_sd/options/sd_endpoint.py @@ -0,0 +1,30 @@ +# Copyright (C) 2024 Christian H. +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# This program is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with this program. If not, see . + + +from dataclasses import dataclass +import ipaddress + + +@dataclass +class IpV4SdEndpointOption: + address: ipaddress.IPv4Address + port: int + + +@dataclass +class IpV6SdEndpointOption: + address: ipaddress.IPv6Address + port: int diff --git a/src/someipy/_internal/_sd/sd_message.py b/src/someipy/_internal/_sd/sd_message.py new file mode 100644 index 0000000..45803bb --- /dev/null +++ b/src/someipy/_internal/_sd/sd_message.py @@ -0,0 +1,14 @@ +from typing import List +from someipy._internal._sd.entries.sd_entry import SdEntry + + +class SdMessage: + + def __init__(self): + self.source: str = "" + self.source_port: int = 0 + self.multicast: bool = True + self.timestamp: float = 0.0 + + self.session_id: int = 0 + self.entries: List[SdEntry] = [] diff --git a/src/someipy/_internal/_sd/service_instance.py b/src/someipy/_internal/_sd/service_instance.py new file mode 100644 index 0000000..30e4a7c --- /dev/null +++ b/src/someipy/_internal/_sd/service_instance.py @@ -0,0 +1,33 @@ +from dataclasses import dataclass + +from someipy._internal._common.endpoint import Endpoint +from someipy._internal.transport_layer_protocol import TransportLayerProtocol + + +@dataclass +class ServiceInstance: + """This class aggregates data from entries and options and provides a compact interface instead of loose SD entries and options""" + + service_id: int + instance_id: int + major_version: int + minor_version: int + ttl: int + endpoint: Endpoint + protocols: frozenset[TransportLayerProtocol] + timestamp: float + + def __eq__(self, other: object) -> bool: + if not isinstance(other, ServiceInstance): + return False + + # Ignore the timestamp in equality comparison + return ( + self.service_id == other.service_id + and self.instance_id == other.instance_id + and self.major_version == other.major_version + and self.minor_version == other.minor_version + and self.ttl == other.ttl + and self.endpoint == other.endpoint + and self.protocols == other.protocols + ) diff --git a/src/someipy/_internal/utils.py b/src/someipy/_internal/utils.py index 9313e0b..1eec3e1 100644 --- a/src/someipy/_internal/utils.py +++ b/src/someipy/_internal/utils.py @@ -86,6 +86,7 @@ def create_rcv_multicast_socket( sock.setsockopt(socket.IPPROTO_IP, socket.IP_ADD_MEMBERSHIP, mreq) return sock + def create_rcv_broadcast_socket( ip_address: str, port: int, interface_address ) -> socket.socket: @@ -119,6 +120,7 @@ def create_rcv_broadcast_socket( return sock + EndpointType = Tuple[ipaddress.IPv4Address, int] @@ -154,7 +156,7 @@ def connection_lost(self, exc: Exception) -> None: def set_bit_at_position(number: int, position: int, value: bool) -> int: - """Set the bit at the specified position to the given boolean value.""" + """Set the bit at the specified position (0 is the least significant bit) to the given boolean value.""" if value: # Set the bit to 1 return number | (1 << position) diff --git a/src/someipy/someipyd.py b/src/someipy/someipyd.py index 8eba696..3e4070b 100644 --- a/src/someipy/someipyd.py +++ b/src/someipy/someipyd.py @@ -22,13 +22,27 @@ import json import logging import os -import platform import struct import sys import ipaddress import time from typing import Any, Dict, List, Set, Tuple, Union +from someipy._internal._daemon.daemon_server_client import ( + ClientMessageEventArgs, + DaemonServerClient, +) +from someipy._internal._sd.deserialization.sd_deserialization import ( + deserialize_sd_message, + is_sd_message, +) +from someipy._internal._sd.deserialization.sd_serialization import serialize_sd_message +from someipy._internal._sd.entries.subscribe_eventgroup_entry import ( + SubscribeEventGroupEntry, +) +from someipy._internal._sd.options.endpoint import IpV4EndpointOption +from someipy._internal._sd.sd_message import SdMessage +from someipy._internal._sd.service_instance import ServiceInstance from someipy._internal.message_types import MessageType from someipy._internal.someip_endpoint import ( SomeipEndpoint, @@ -51,13 +65,11 @@ build_subscribe_eventgroup_sd_header, ) from someipy._internal.someip_sd_extractors import ( - extract_offered_services, extract_subscribe_ack_eventgroup_entries, extract_subscribe_entries, extract_subscribe_nack_eventgroup_entries, ) from someipy._internal.someip_sd_header import ( - SdEntry, SdEntryType, SdEventGroupEntry, SdService, @@ -90,6 +102,10 @@ ) from someipy._internal.offer_service_storage import OfferServiceStorage, ServiceToOffer from someipy.service import Event, Method, EventGroup +from someipy._internal._daemon.daemon_server import ( + ClientConnectedEventArgs, + DaemonServer, +) DEFAULT_SOCKET_PATH = "/tmp/someipyd.sock" @@ -279,42 +295,18 @@ def __hash__(self): class SomeipDaemon: - def __init__(self, config_file=None, log_path=None): - self.config = self._load_config(config_file) - self.socket_path = self.config.get("socket_path", DEFAULT_SOCKET_PATH) + def __init__( + self, server: DaemonServer, config: dict = None, logger: logging.Logger = None + ): + + self.config = config + self.logger = logger + self.sd_address = self.config.get("sd_address", DEFAULT_SD_ADDRESS) self.sd_port = self.config.get("sd_port", DEFAULT_SD_PORT) self.interface = self.config.get("interface", DEFAULT_INTERFACE_IP) - log_level = self.config.get("log_level", "DEBUG") - self.log_path = log_path if log_path else self.config.get("log_path") - self.use_tcp = self.config.get("use_tcp", platform.system() != "Linux") - self.tcp_port = self.config.get("tcp_port", DEFAULT_TCP_PORT) - - self.logger = self._configure_logging( - log_level=log_level, log_path=self.log_path - ) - - self.logger.info( - f"Starting SOME/IP Daemon with config:\n" - f"Socket path: {self.socket_path}\n" - f"SD address: {self.sd_address}\n" - f"SD port: {self.sd_port}\n" - f"Interface: {self.interface}\n" - f"Loglevel: {log_level}\n" - f"Log path: {self.log_path if self.log_path else 'Console'}\n" - f"Use TCP: {self.use_tcp}\n" - f"TCP Port: {self.tcp_port}\n" - ) - log_level_mapping = { - "DEBUG": logging.DEBUG, - "ERROR": logging.ERROR, - "INFO": logging.INFO, - "FATAL": logging.FATAL, - } - - if log_level in log_level_mapping: - self.logger.setLevel(log_level_mapping[log_level]) + self._server = server self._sd_socket_mcast = None self._sd_socket_ucast = None @@ -322,7 +314,7 @@ def __init__(self, config_file=None, log_path=None): self._ucast_transport = None # Services offered by other ECUs - self._found_services: List[SdServiceWithTimestamp] = [] + self._found_services: List[ServiceInstance] = [] # Services offered by this daemon self._services_to_offer = OfferServiceStorage() @@ -343,57 +335,58 @@ def __init__(self, config_file=None, log_path=None): # Qeueues and tasks stored by id of asyncio.StreamWriter self._tx_queues: Dict[int, asyncio.Queue] = {} self._tx_tasks: Dict[int, asyncio.Task] = {} - self._rx_queues: Dict[int, asyncio.Queue] = {} self._someip_server_endpoints = SomeipEndpointStorage() self._someip_client_endpoints = SomeipEndpointStorage() - self._ttl_task = asyncio.create_task(self._check_services_ttl_task()) + self._ttl_task: asyncio.Task = None self._issued_method_calls: Dict[MethodCall, int] = {} - def _configure_logging(self, log_level=logging.DEBUG, log_path=None): - logger = logging.getLogger(f"someipyd") - logger.setLevel(log_level) - - # Remove any existing handlers to prevent duplicate logs - if logger.hasHandlers(): - logger.handlers.clear() + async def new_client_connected( + self, sender: object, event_args: ClientConnectedEventArgs + ): + self.logger.info(f"New client connected: {event_args.client.id}") - formatter = logging.Formatter( - "%(asctime)s.%(msecs)03d %(name)s [%(levelname)s]: %(message)s", - datefmt="%Y-%m-%d,%H:%M:%S", + # Create the tx_queue and tx_task for the client using the writers id + self._tx_queues[event_args.client.id] = asyncio.Queue() + self._tx_tasks[event_args.client.id] = asyncio.create_task( + self.tx_task(event_args.client) ) - if log_path: - file_handler = logging.FileHandler(log_path) - file_handler.setLevel(log_level) - file_handler.setFormatter(formatter) - logger.addHandler(file_handler) - else: - console_handler = logging.StreamHandler(sys.stdout) - console_handler.setLevel(log_level) - console_handler.setFormatter(formatter) - logger.addHandler(console_handler) - return logger - - def _load_config(self, config_file): - if config_file and os.path.exists(config_file): - try: - with open(config_file, "r") as f: - return json.load(f) - except (FileNotFoundError, json.JSONDecodeError) as e: - self.logger.error(f"Error loading config file: {e}. Using defaults.") - return {} - elif os.path.exists(DEFAULT_CONFIG_FILE): + event_args.client.message_received += self.handle_client_message + + async def client_disconnected( + self, sender: object, event_args: ClientConnectedEventArgs + ): + self.logger.info(f"Client disconnected: {event_args.client.id}") + + writer_id = event_args.client.id + # Remove all subscriptions for the client + self._requested_subscriptions.remove_client(writer_id) + + # Clean up the transmission task for the client. This will also clean up the transmission queue + tx_task = self._tx_tasks.get(writer_id) + if tx_task and not tx_task.cancelled(): + tx_task.cancel() try: - with open(DEFAULT_CONFIG_FILE, "r") as f: - return json.load(f) - except (FileNotFoundError, json.JSONDecodeError) as e: - self.logger.error(f"Error loading config file: {e}. Using defaults.") - return {} - else: - return {} + await tx_task + except asyncio.CancelledError: + pass + + self._services_to_offer.remove_client(writer_id) + self._cleanup_unused_timers() + + client_endpoints = self._someip_server_endpoints.get_endpoints(writer_id) + if client_endpoints is not None: + for endpoint in client_endpoints: + self.logger.debug( + f"Closing endpoint {endpoint.dst_ip()}:{endpoint.dst_port()} for client {writer_id}" + ) + endpoint.shutdown() + self._someip_server_endpoints.remove_endpoint(writer_id, endpoint) + + self.logger.debug(f"Client disconnected") async def _create_server_endpoint( self, ip: str, port: int, protocol: TransportLayerProtocol @@ -672,9 +665,8 @@ def _close_unused_endpoints(self): if not endpoint_used: endpoints_to_close.append(endpoint) - async def tx_task(self, writer: asyncio.StreamWriter): - tx_queue = self._tx_queues[id(writer)] - + async def tx_task(self, client: DaemonServerClient): + tx_queue = self._tx_queues[client.id] try: while True: try: @@ -683,8 +675,8 @@ async def tx_task(self, writer: asyncio.StreamWriter): try: # Send the data - writer.write(data) - await writer.drain() + client.writer.write(data) + await client.writer.drain() tx_queue.task_done() except ConnectionError as e: self.logger.error(f"Error sending data in tx task: {e}") @@ -695,115 +687,24 @@ async def tx_task(self, writer: asyncio.StreamWriter): continue except asyncio.CancelledError: - self.logger.debug(f"TX task for writer {id(writer)} cancelled") + self.logger.debug(f"TX task for writer {client.id} cancelled") # Perform cleanup here try: - writer.close() - await writer.wait_closed() + client.writer.close() + await client.writer.wait_closed() except Exception as e: self.logger.error(f"Error closing writer: {e}") finally: # Always clean up the queue - self._tx_queues.pop(id(writer), None) - self.logger.debug(f"TX task for writer {id(writer)} finished") - - async def handle_client(self, reader, writer): - writer_id = id(writer) - self.logger.info(f"New client connected: {writer_id}") - - # Create the tx_queue and tx_task for the client using the writers id - self._tx_queues[writer_id] = asyncio.Queue() - self._tx_tasks[writer_id] = asyncio.create_task(self.tx_task(writer)) - self._rx_queues[writer_id] = asyncio.Queue() - - try: - wait_for_header = True - header_buffer = b"" - message_buffer = b"" - message_length = 0 - - while True: - if wait_for_header: - data = await reader.read(256 - len(header_buffer)) - if not data: - self.logger.debug(f"Data is none. Client disconnected.") - break # Client disconnected - - header_buffer += data - - if len(header_buffer) == 256: - try: - message_length = struct.unpack(" 256: - self.logger.error(f"Client sent too much header data.") - break - - else: - data = await reader.read(message_length - len(message_buffer)) - if not data: - self.logger.debug(f"Data is none. Client disconnected.") - break # Client disconnected - - message_buffer += data - - if len(message_buffer) == message_length: - - self.logger.debug(f"Client sent message: {message_buffer}") - json_message = json.loads(message_buffer.decode("utf-8")) - await self.handle_client_message(json_message, writer) - - wait_for_header = True - header_buffer = b"" # reset header buffer - message_buffer = b"" # reset message buffer - message_length = 0 # reset message length - elif len(message_buffer) > message_length: - self.logger.error(f"Client sent too much message data.") - break - except ConnectionResetError: - self.logger.error(f"Client disconnected abruptly.") - except Exception as e: - self.logger.error(f"Error handling client: {e}") - finally: + self._tx_queues.pop(client.id, None) + self.logger.debug(f"TX task for writer {client.id} finished") - # Remove all subscriptions for the client - self._requested_subscriptions.remove_client(writer_id) - - # Clean up the transmission task for the client. This will also clean up the transmission queue - tx_task = self._tx_tasks.get(writer_id) - if tx_task and not tx_task.cancelled(): - tx_task.cancel() - try: - await tx_task - except asyncio.CancelledError: - pass - - self._rx_queues.pop(writer_id, None) - - self._services_to_offer.remove_client(writer_id) - self._cleanup_unused_timers() - - client_endpoints = self._someip_server_endpoints.get_endpoints(writer_id) - if client_endpoints is not None: - for endpoint in client_endpoints: - self.logger.debug( - f"Closing endpoint {endpoint.dst_ip()}:{endpoint.dst_port()} for client {writer_id}" - ) - endpoint.shutdown() - self._someip_server_endpoints.remove_endpoint(writer_id, endpoint) - - self.logger.debug(f"Client disconnected") - - async def handle_client_message(self, message: dict, writer: asyncio.StreamWriter): - writer_id = id(writer) - message_type = message.get("type") + async def handle_client_message( + self, sender: object, event_args: ClientMessageEventArgs + ): + writer_id = id(event_args.client.id) + message = event_args.message + message_type = event_args.message.get("type") self.logger.debug(f"Received message type: {message_type}") message_handlers = { @@ -1433,36 +1334,20 @@ async def send_to_all_clients(self, message): await writer.wait_closed() async def start_server(self): - - if not self.use_tcp: - if os.path.exists(self.socket_path): - os.unlink(self.socket_path) - - server = await asyncio.start_unix_server( - self.handle_client, path=self.socket_path - ) - self.logger.info(f"Unix domain socket server started at {self.socket_path}") - else: - server = await asyncio.start_server( - self.handle_client, - host="127.0.0.1", - port=self.tcp_port, - reuse_port=True, - ) - try: + if self._ttl_task is None or self._ttl_task.done(): + self._ttl_task = asyncio.create_task(self._check_services_ttl_task()) await self.start_sd_listening() - async with server: - await server.serve_forever() + await self._server.serve_forever() except asyncio.CancelledError: - self.logger.info(f"{"TCP" if self.use_tcp else "UDS"} server cancelled.") + self.logger.info(f"Server cancelled.") finally: if self._mcast_transport: self._mcast_transport.close() if self._ucast_transport: self._ucast_transport.close() - self.logger.info(f"{'TCP' if self.use_tcp else 'UDS'} server stopped.") + self.logger.info(f"Server stopped.") def _timeout_of_offered_service(self, offered_service: SdService): self.logger.info( @@ -1500,17 +1385,15 @@ async def wait_for_message_in_rx_queue( return found_message - def _handle_offered_service(self, offered_service: SdService2): + def _handle_offered_service(self, offered_service: ServiceInstance): self.logger.info(f"Received offered service: {offered_service}") - new_service = SdServiceWithTimestamp(offered_service, time.time()) - - if new_service not in self._found_services: - self._found_services.append(new_service) + if offered_service not in self._found_services: + self._found_services.append(offered_service) else: # Update the timestamp if the service is already in the list - index = self._found_services.index(new_service) - self._found_services[index].timestamp = time.time() + index = self._found_services.index(offered_service) + self._found_services[index].timestamp = offered_service.timestamp # Check if there is a requested subscription for this service for requested_subscription in self._requested_subscriptions.has_subscriptions( @@ -1556,21 +1439,34 @@ def _handle_offered_service(self, offered_service: SdService2): reboot_flag, ) = self._unicast_session_handler.update_session() - # Improvement: Pack all entries into a single SD message - subscribe_sd_header = build_subscribe_eventgroup_sd_header( + # Build subscribe message + sd_message = SdMessage() + sd_message.session_id = session_id + + options = [] + for protocol in requested_protocols: + options.append( + IpV4EndpointOption( + address=ipaddress.IPv4Address( + requested_subscription[0].client_endpoint_ip + ), + protocol=protocol, + port=requested_subscription[0].client_endpoint_port, + ) + ) + + entry = SubscribeEventGroupEntry( service_id=offered_service.service_id, instance_id=offered_service.instance_id, major_version=offered_service.major_version, - ttl=int(requested_subscription[0].ttl), - event_group_id=requested_subscription[0].eventgroup.id, - session_id=session_id, - reboot_flag=reboot_flag, - endpoint=( - ipaddress.IPv4Address(requested_subscription[0].client_endpoint_ip), - requested_subscription[0].client_endpoint_port, - ), - protocols=requested_protocols, + minor_version=offered_service.minor_version, + ttl=requested_subscription[0].ttl, + eventgroup_id=requested_subscription[0].eventgroup.id, + counter=0, + ip_v4_endpoints=options, + ip_v6_endpoints=[], ) + sd_message.entries.append(entry) pending_subscription = Subscription( service_id=offered_service.service_id, @@ -1588,7 +1484,7 @@ def _handle_offered_service(self, offered_service: SdService2): if self._ucast_transport: self._ucast_transport.sendto( - subscribe_sd_header.to_buffer(), + serialize_sd_message(sd_message), (str(offered_service.endpoint[0]), self.sd_port), ) @@ -1734,14 +1630,22 @@ def datagram_received_mcast( if addr[0] == self.interface and addr[1] == self.sd_port: return - someip_header = SomeIpHeader.from_buffer(data) - if not someip_header.is_sd_header(): + if is_sd_message(data) is False: return - someip_sd_header = SomeIpSdHeader.from_buffer(data) + sd_message = deserialize_sd_message(data) + sd_message.timestamp = time.time() - for offered_service in extract_offered_services(someip_sd_header): - self._handle_offered_service(offered_service) + # someip_header = SomeIpHeader.from_buffer(data) + # if not someip_header.is_sd_header(): + # return + + for offer_service_entry in [ + o for o in sd_message.entries if o.entry_type == SdEntryType.OFFER_SERVICE + ]: + self._handle_offered_service(offer_service_entry, sd_message.timestamp) + + someip_sd_header = SomeIpSdHeader.from_buffer(data) for subscription in extract_subscribe_entries(someip_sd_header): self._handle_subscription(subscription) @@ -1816,13 +1720,102 @@ async def start_sd_listening(self): ) +def _configure_logging(log_level=logging.DEBUG, log_path=None) -> logging.Logger: + logger = logging.getLogger(f"someipyd") + logger.setLevel(log_level) + + # Remove any existing handlers to prevent duplicate logs + if logger.hasHandlers(): + logger.handlers.clear() + + formatter = logging.Formatter( + "%(asctime)s.%(msecs)03d %(name)s [%(levelname)s]: %(message)s", + datefmt="%Y-%m-%d,%H:%M:%S", + ) + + if log_path: + file_handler = logging.FileHandler(log_path) + file_handler.setLevel(log_level) + file_handler.setFormatter(formatter) + logger.addHandler(file_handler) + else: + console_handler = logging.StreamHandler(sys.stdout) + console_handler.setLevel(log_level) + console_handler.setFormatter(formatter) + logger.addHandler(console_handler) + return logger + + +def _load_config(config_file: str) -> dict: + if config_file and os.path.exists(config_file): + try: + with open(config_file, "r") as f: + return json.load(f) + except (FileNotFoundError, json.JSONDecodeError) as e: + print(f"Error loading config file: {e}. Using defaults.") + return {} + elif os.path.exists(DEFAULT_CONFIG_FILE): + try: + with open(DEFAULT_CONFIG_FILE, "r") as f: + return json.load(f) + except (FileNotFoundError, json.JSONDecodeError) as e: + print(f"Error loading config file: {e}. Using defaults.") + return {} + else: + return {} + + async def async_main(): parser = argparse.ArgumentParser(description="SOME/IP Daemon") parser.add_argument("--config", help="Path to configuration file") parser.add_argument("--log-path", help="Path to log file") args = parser.parse_args() - daemon = SomeipDaemon(args.config, args.log_path) + # Load configuration + config = _load_config(args.config) + + # Logging + log_path = args.log_path if args.log_path else config.get("log_path", None) + log_level = config.get("log_level", "INFO") + + log_level_mapping = { + "DEBUG": logging.DEBUG, + "ERROR": logging.ERROR, + "INFO": logging.INFO, + "FATAL": logging.FATAL, + } + + if log_level in log_level_mapping: + log_level = log_level_mapping[log_level] + + logger = _configure_logging(log_level=log_level, log_path=log_path) + + logger.info( + f"Starting SOME/IP Daemon with config:\n" + f"Socket path: {config.get('socket_path', DEFAULT_SOCKET_PATH)}\n" + f"SD address: {config.get('sd_address', DEFAULT_SD_ADDRESS)}\n" + f"SD port: {config.get('sd_port', DEFAULT_SD_PORT)}\n" + f"Interface: {config.get('interface', DEFAULT_INTERFACE_IP)}\n" + f"Loglevel: {log_level}\n" + f"Log path: {log_path if log_path else 'Console'}\n" + f"Use TCP: {config.get('use_tcp', False)}\n" + f"TCP Port: {config.get('tcp_port', None)}\n" + ) + + daemon_server = DaemonServer(logger) + + daemon = SomeipDaemon(daemon_server, config, logger) + + daemon_server.client_connected += daemon.new_client_connected + daemon_server.client_disconnected += daemon.client_disconnected + + await daemon_server.start( + use_uds=config.get("use_uds", True), + socket_path=config.get("socket_path", DEFAULT_SOCKET_PATH), + tcp_port=config.get("tcp_port", None), + host=config.get("host", "127.0.0.1"), + ) + await daemon.start_server() diff --git a/tests/sd/__init__.py b/tests/sd/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/sd/test_offer_service_entry.py b/tests/sd/test_offer_service_entry.py new file mode 100644 index 0000000..5b57d80 --- /dev/null +++ b/tests/sd/test_offer_service_entry.py @@ -0,0 +1,46 @@ +import pytest +from someipy._internal._sd.entries.offer_service_entry import OfferServiceEntry +from someipy._internal._sd.options.endpoint import IpV4EndpointOption + + +def test_base_types_len(): + + ip_endpoint_option_1 = IpV4EndpointOption( + address="192.168.1.1", protocol=1, port=8080 + ) + ip_endpoint_option_2 = IpV4EndpointOption( + address="192.168.1.2", protocol=1, port=8080 + ) + + offer_service_entry_1 = OfferServiceEntry( + service_id=1, + instance_id=1, + major_version=1, + minor_version=0, + ttl=120, + ip_v4_endpoints=[ip_endpoint_option_1], + ip_v6_endpoints=[], + ) + + offer_service_entry_2 = OfferServiceEntry( + service_id=1, + instance_id=1, + major_version=1, + minor_version=0, + ttl=120, + ip_v4_endpoints=[ip_endpoint_option_1], + ip_v6_endpoints=[], + ) + + offer_service_entry_3 = OfferServiceEntry( + service_id=1, + instance_id=1, + major_version=1, + minor_version=0, + ttl=120, + ip_v4_endpoints=[ip_endpoint_option_2], + ip_v6_endpoints=[], + ) + + assert offer_service_entry_1 == offer_service_entry_2 + assert offer_service_entry_1 != offer_service_entry_3 diff --git a/tests/sd/test_sd_deserialization.py b/tests/sd/test_sd_deserialization.py new file mode 100644 index 0000000..c8bd07d --- /dev/null +++ b/tests/sd/test_sd_deserialization.py @@ -0,0 +1,185 @@ +import pytest + +from someipy._internal._sd.deserialization.sd_deserialization import ( + CommonEntryData, + CommonOptionData, + SdOptionOnWireType, + deserialize_common_entry_data, + deserialize_common_option_data, + deserialize_ipv4_endpoint_option, + deserialize_ipv4_multicast_option, + deserialize_ipv4_sd_endpoint_option, + deserialize_ipv6_endpoint_option, + deserialize_ipv6_multicast_option, + deserialize_ipv6_sd_endpoint_option, + deserialize_load_balancing_option, +) +from someipy._internal._sd.options.endpoint import ( + IpV4EndpointOption, + IpV6EndpointOption, +) +from someipy._internal._sd.options.load_balancing import LoadBalancingOption +from someipy._internal._sd.options.multicast import ( + IpV4MulticastOption, + IpV6MulticastOption, +) +from someipy._internal._sd.options.sd_endpoint import ( + IpV4SdEndpointOption, + IpV6SdEndpointOption, +) +from someipy._internal.transport_layer_protocol import TransportLayerProtocol + + +def test_deserialize_common_entry_data(): + data = bytes( + [0x01, 0x02, 0x03, 0x24, 0x00, 0x10, 0x00, 0x20, 0x01, 0x0B, 0x00, 0x0A] + ) + result = deserialize_common_entry_data(data) + assert isinstance(result, CommonEntryData) + assert result.type_field_value == 0x01 + assert result.index_first_option == 0x02 + assert result.index_second_option == 0x03 + assert result.num_options_1 == 2 + assert result.num_options_2 == 4 + assert result.service_id == 0x0010 + assert result.instance_id == 0x0020 + assert result.major_version == 0x01 + assert result.ttl == 0x0B000A + + +def test_deserialize_common_option_data(): + data = bytes([0x01, 0x05, 0x04, 0x80]) + result = deserialize_common_option_data(data) + assert isinstance(result, CommonOptionData) + assert result.option_length == 5 + 256 + assert result.option_type == SdOptionOnWireType.IPV4_ENDPOINT + assert result.discardable_flag is True + + +def test_deserialize_ipv4_endpoint_option(): + data = bytes([192, 168, 1, 10, 0x01, 0x11, 0x10, 0x01]) + result = deserialize_ipv4_endpoint_option(data) + assert isinstance(result, IpV4EndpointOption) + assert str(result.address) == "192.168.1.10" + assert result.protocol == TransportLayerProtocol.UDP + assert result.port == 0x1001 + + +def test_deserialize_ipv6_endpoint_option(): + data = bytes( + [ + 0x20, + 0x01, + 0x0D, + 0xB8, + 0x85, + 0xA3, + 0x00, + 0x00, + 0x00, + 0x00, + 0x8A, + 0x2E, + 0x03, + 0x70, + 0x73, + 0x34, + 0x01, + 0x11, + 0x10, + 0x01, + ] + ) + result = deserialize_ipv6_endpoint_option(data) + assert isinstance(result, IpV6EndpointOption) + assert str(result.address) == "2001:db8:85a3::8a2e:370:7334" + assert result.protocol == TransportLayerProtocol.UDP + assert result.port == 0x1001 + + +def test_deserialize_ipv4_multicast_option(): + data = bytes([192, 168, 1, 10, 0x01, 0x11, 0x10, 0x01]) + result = deserialize_ipv4_multicast_option(data) + assert isinstance(result, IpV4MulticastOption) + assert str(result.address) == "192.168.1.10" + assert result.protocol == TransportLayerProtocol.UDP + assert result.port == 0x1001 + + +def test_deserialize_ipv6_multicast_option(): + data = bytes( + [ + 0x20, + 0x01, + 0x0D, + 0xB8, + 0x85, + 0xA3, + 0x00, + 0x00, + 0x00, + 0x00, + 0x8A, + 0x2E, + 0x03, + 0x70, + 0x73, + 0x34, + 0x01, + 0x11, + 0x10, + 0x01, + ] + ) + result = deserialize_ipv6_multicast_option(data) + assert isinstance(result, IpV6MulticastOption) + assert str(result.address) == "2001:db8:85a3::8a2e:370:7334" + assert result.protocol == TransportLayerProtocol.UDP + assert result.port == 0x1001 + + +def test_deserialize_ipv4_sd_endpoint_option(): + data = bytes([192, 168, 1, 10, 0x01, 0x11, 0x10, 0x01]) + result = deserialize_ipv4_sd_endpoint_option(data) + assert isinstance(result, IpV4SdEndpointOption) + assert str(result.address) == "192.168.1.10" + assert result.port == 0x1001 + + +def test_deserialize_ipv6_sd_endpoint_option(): + data = bytes( + [ + 0x20, + 0x01, + 0x0D, + 0xB8, + 0x85, + 0xA3, + 0x00, + 0x00, + 0x00, + 0x00, + 0x8A, + 0x2E, + 0x03, + 0x70, + 0x73, + 0x34, + 0x01, + 0x11, + 0x10, + 0x01, + ] + ) + result = deserialize_ipv6_sd_endpoint_option(data) + assert isinstance(result, IpV6SdEndpointOption) + assert str(result.address) == "2001:db8:85a3::8a2e:370:7334" + assert result.port == pow(16, 3) + 1 + + +def test_deserialize_load_balancing_option(): + data = bytes([0x01, 0x02, 0x03, 0x04]) + result = deserialize_load_balancing_option(data) + assert isinstance(result, LoadBalancingOption) + assert result.priority == 0x0102 + assert result.weight == 0x0304 diff --git a/tests/sd/test_sd_serialization.py b/tests/sd/test_sd_serialization.py new file mode 100644 index 0000000..f507948 --- /dev/null +++ b/tests/sd/test_sd_serialization.py @@ -0,0 +1,187 @@ +import ipaddress + +from someipy._internal._sd.deserialization.sd_serialization import ( + serialize_ipv4_endpoint_option, + serialize_ipv6_endpoint_option, + serialize_sd_message, +) +from someipy._internal._sd.entries.offer_service_entry import OfferServiceEntry +from someipy._internal._sd.options.endpoint import ( + IpV4EndpointOption, + IpV6EndpointOption, +) +from someipy._internal._sd.sd_message import SdMessage +from someipy._internal.transport_layer_protocol import TransportLayerProtocol + + +def test_serialize_ipv4_endpoint_option(): + option = IpV4EndpointOption( + address=ipaddress.IPv4Address("1.2.3.4"), + protocol=TransportLayerProtocol.UDP, + port=0x1F90, + ) + data = serialize_ipv4_endpoint_option(option) + + expected_data = bytes( + [0x00, 0x09, 0x04, 0x00, 0x01, 0x02, 0x03, 0x04, 0x00, 0x11, 0x1F, 0x90] + ) + + assert data == expected_data + + +def test_serialize_ipv6_endpoint_option(): + option = IpV6EndpointOption( + address=ipaddress.IPv6Address("2001:0db8:85a3:0000:0000:8a2e:0370:7334"), + protocol=TransportLayerProtocol.TCP, + port=0x2328, + ) + data = serialize_ipv6_endpoint_option(option) + + expected_data = bytes( + [ + 0x00, + 0x15, + 0x06, + 0x00, + 0x20, + 0x01, + 0x0D, + 0xB8, + 0x85, + 0xA3, + 0x00, + 0x00, + 0x00, + 0x00, + 0x8A, + 0x2E, + 0x03, + 0x70, + 0x73, + 0x34, + 0x00, + 0x06, + 0x23, + 0x28, + ] + ) + + assert data == expected_data + + +def test_serialize_empty_sd_message(): + sd_message = SdMessage() + sd_message.multicast = False + sd_message.session_id = 0x2 + + data = serialize_sd_message(sd_message) + + # fmt: off + expected_data = bytes( + [ + 0xFF, 0xFF, 0x81, 0x00, + 0x00, 0x00, 0x00, 0x14, # length: 20 + 0x00, 0x00, 0x00, 0x02, + 0x01, 0x01, 0x02, 0x00, + 0x40, 0x00, 0x00, 0x00, + 0x00, 0x00, 0x00, 0x00, # length entries + 0x00, 0x00, 0x00, 0x00]) # length of options + # fmt: on + + assert len(data) == len(expected_data) + assert data == expected_data + + +def test_serialize_sd_message_with_one_offer_service_entry_without_options(): + sd_message = SdMessage() + sd_message.multicast = True + sd_message.session_id = 0x1234 + + sd_message.entries.append( + OfferServiceEntry( + service_id=0x01, + instance_id=0x02, + major_version=0x03, + minor_version=0x04, + ttl=0x10, + ip_v4_endpoints=[], + ip_v6_endpoints=[], + ) + ) + + data = serialize_sd_message(sd_message) + + # fmt: off + expected_data = bytes( + [ + 0xFF, 0xFF, 0x81, 0x00, + 0x00, 0x00, 0x00, 0x24, # length: 36 + 0x00, 0x00, 0x12, 0x34, + 0x01, 0x01, 0x02, 0x00, + 0x40, 0x00, 0x00, 0x00, + 0x00, 0x00, 0x00, 0x10, # length entries: 16 + 0x01, 0x00, 0x00, 0x00, + 0x00, 0x01, 0x00, 0x02, + 0x03, 0x00, 0x00, 0x10, + 0x00, 0x00, 0x00, 0x04, + 0x00, 0x00, 0x00, 0x00]) # length of options: 0 + # fmt: on + + assert len(data) == len(expected_data) + assert data == expected_data + + +def test_serialize_sd_message_with_one_offer_service_entry_with_two_options(): + sd_message = SdMessage() + sd_message.multicast = True + sd_message.session_id = 0x1234 + + sd_message.entries.append( + OfferServiceEntry( + service_id=0x01, + instance_id=0x02, + major_version=0x03, + minor_version=0x04, + ttl=0x10, + ip_v4_endpoints=[ + IpV4EndpointOption( + address=ipaddress.IPv4Address("192.168.0.1"), + protocol=TransportLayerProtocol.TCP, + port=0x2328, + ), + IpV4EndpointOption( + address=ipaddress.IPv4Address("192.168.0.2"), + protocol=TransportLayerProtocol.UDP, + port=0x2429, + ), + ], + ip_v6_endpoints=[], + ) + ) + + data = serialize_sd_message(sd_message) + + # fmt: off + expected_data = bytes( + [ + 0xFF, 0xFF, 0x81, 0x00, + 0x00, 0x00, 0x00, 60, # length: 60 + 0x00, 0x00, 0x12, 0x34, + 0x01, 0x01, 0x02, 0x00, + 0x40, 0x00, 0x00, 0x00, + 0x00, 0x00, 0x00, 0x10, # length entries: 16 + 0x01, 0x00, 0x00, 0x20, + 0x00, 0x01, 0x00, 0x02, + 0x03, 0x00, 0x00, 0x10, + 0x00, 0x00, 0x00, 0x04, + 0x00, 0x00, 0x00, 0x18, # length of options: 2 * 12 + 0x00, 0x09, 0x04, 0x00, + 192, 168, 0, 1, + 0x00, 0x06, 0x23, 0x28, + 0x00, 0x09, 0x04, 0x00, + 192, 168, 0, 2, + 0x00, 0x11, 0x24, 0x29]) + # fmt: on + + assert len(data) == len(expected_data) + assert data == expected_data diff --git a/tests/test_sd_service_instance.py b/tests/test_sd_service_instance.py new file mode 100644 index 0000000..cfbbdda --- /dev/null +++ b/tests/test_sd_service_instance.py @@ -0,0 +1,41 @@ +import ipaddress +import pytest +from someipy._internal._sd.service_instance import ServiceInstance +from someipy._internal.transport_layer_protocol import TransportLayerProtocol + + +def test_equality(): + instance_1 = ServiceInstance( + service_id=1234, + instance_id=5678, + major_version=1, + minor_version=0, + ttl=4, + endpoint=(ipaddress.IPv4Address("127.0.0.1"), 12345), + protocols=frozenset([TransportLayerProtocol.TCP]), + timestamp=1.0, + ) + + instance_2 = ServiceInstance( + service_id=1234, + instance_id=5678, + major_version=1, + minor_version=0, + ttl=4, + endpoint=(ipaddress.IPv4Address("127.0.0.1"), 12345), + protocols=frozenset([TransportLayerProtocol.TCP]), + timestamp=1.0, + ) + + instance_3 = ServiceInstance( + service_id=1234, + instance_id=5678, + major_version=1, + minor_version=0, + ttl=4, + endpoint=(ipaddress.IPv4Address("127.0.0.2"), 12345), + protocols=frozenset([TransportLayerProtocol.TCP]), + timestamp=1.0, + ) + assert instance_1 == instance_2 + assert instance_1 != instance_3 diff --git a/tests/test_someipyd.py b/tests/test_someipyd.py new file mode 100644 index 0000000..454e17f --- /dev/null +++ b/tests/test_someipyd.py @@ -0,0 +1,85 @@ +import logging +import pytest +from unittest.mock import Mock +from someipy._internal._common.endpoint import Endpoint +from someipy._internal._sd.service_instance import ServiceInstance +from someipy.someipyd import DaemonServer, SomeipDaemon + + +@pytest.fixture +def mock_logger() -> logging.Logger: + mock_logger = Mock(spec=logging.Logger) + return mock_logger + + +@pytest.fixture +def mock_daemon_server(mock_logger) -> DaemonServer: + return DaemonServer(mock_logger) + + +@pytest.fixture +def daemon(mock_daemon_server, mock_logger) -> SomeipDaemon: + config = {} + + return SomeipDaemon(mock_daemon_server, config, mock_logger) + + +def test_handle_offered_service_adds_service(daemon: SomeipDaemon): + service_instance = ServiceInstance( + service_id=1, + instance_id=1, + major_version=1, + minor_version=0, + ttl=10, + endpoint=Endpoint("123", 1), + protocols=frozenset(), + timestamp=1000, + ) + + # Initially, the found services list should be empty + assert len(daemon._found_services) == 0 + + # Handle the offered service + daemon._handle_offered_service(service_instance) + + # Now, the found services list should contain the new service + assert len(daemon._found_services) == 1 + assert daemon._found_services[0] == service_instance + + +def test_handle_offered_service_updates_timestamp(daemon: SomeipDaemon): + service_instance = ServiceInstance( + service_id=1, + instance_id=1, + major_version=1, + minor_version=0, + ttl=10, + endpoint=Endpoint("123", 1), + protocols=frozenset(), + timestamp=1000, + ) + + # Initially, the found services list should be empty + assert len(daemon._found_services) == 0 + + # Handle the offered service + daemon._handle_offered_service(service_instance) + + service_instance_2 = ServiceInstance( + service_id=1, + instance_id=1, + major_version=1, + minor_version=0, + ttl=10, + endpoint=Endpoint("123", 1), + protocols=frozenset(), + timestamp=2000, + ) + + # Handle the offered service again with updated timestamp + daemon._handle_offered_service(service_instance_2) + + # Now, the found services list should contain the updated timestamp + assert len(daemon._found_services) == 1 + assert daemon._found_services[0] == service_instance + assert daemon._found_services[0].timestamp == 2000 From 47e0cb28df7da82114a229e1c8e8cb4f76fa96d3 Mon Sep 17 00:00:00 2001 From: Christian Date: Mon, 29 Dec 2025 10:53:39 +0100 Subject: [PATCH 3/6] * Extend unit tests * Extract classes into separate files --- .gitignore | 5 +- src/someipy/_internal/_common/event.py | 2 +- src/someipy/_internal/_daemon/subscription.py | 70 ++++ .../_internal/_daemon/subscription_storage.py | 96 +++++ .../_sd/deserialization/sd_deserialization.py | 2 +- .../_internal/offer_service_storage.py | 13 +- .../_internal/someip_endpoint_factory.py | 95 +++++ src/someipy/_internal/someip_sd_builder.py | 74 ---- src/someipy/someipyd.py | 394 +++++++----------- tests/common/__init__.py | 0 tests/common/test_endpoint.py | 28 ++ tests/daemon/__init__.py | 0 tests/daemon/test_subscription_storage.py | 162 +++++++ tests/test_someipyd.py | 210 ++++++++-- 14 files changed, 789 insertions(+), 362 deletions(-) create mode 100644 src/someipy/_internal/_daemon/subscription.py create mode 100644 src/someipy/_internal/_daemon/subscription_storage.py create mode 100644 src/someipy/_internal/someip_endpoint_factory.py create mode 100644 tests/common/__init__.py create mode 100644 tests/common/test_endpoint.py create mode 100644 tests/daemon/__init__.py create mode 100644 tests/daemon/test_subscription_storage.py diff --git a/.gitignore b/.gitignore index 2d64693..e776ac9 100644 --- a/.gitignore +++ b/.gitignore @@ -11,4 +11,7 @@ integration_tests/install build/ -.coverage \ No newline at end of file +.coverage +htmlcov/ +*.pdf +*.bash diff --git a/src/someipy/_internal/_common/event.py b/src/someipy/_internal/_common/event.py index 9a474f3..6d722f1 100644 --- a/src/someipy/_internal/_common/event.py +++ b/src/someipy/_internal/_common/event.py @@ -1,4 +1,4 @@ -# Copyright (C) 2024 Christian H. +# Copyright (C) 2025 Christian H. # # This program is free software: you can redistribute it and/or modify # it under the terms of the GNU General Public License as published by diff --git a/src/someipy/_internal/_daemon/subscription.py b/src/someipy/_internal/_daemon/subscription.py new file mode 100644 index 0000000..f5ac842 --- /dev/null +++ b/src/someipy/_internal/_daemon/subscription.py @@ -0,0 +1,70 @@ +# Copyright (C) 2025 Christian H. +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# This program is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with this program. If not, see . + +from someipy._internal._common.endpoint import Endpoint +from someipy._internal.transport_layer_protocol import TransportLayerProtocol +from someipy.service import EventGroup + + +class Subscription: + def __init__( + self, + service_id: int, + instance_id: int, + major_version: int, + eventgroup: EventGroup, + ttl_seconds: int, + client_endpoint: Endpoint, + server_endpoint: Endpoint, + protocols: frozenset[TransportLayerProtocol], + timestamp_last_update: float = 0.0, + ): + self.service_id = service_id + self.instance_id = instance_id + self.major_version = major_version + self.eventgroup = eventgroup + self.ttl_seconds = ttl_seconds + + self.client_endpoint = client_endpoint + self.server_endpoint = server_endpoint + self.protocols = protocols + self.timestamp_last_update = timestamp_last_update + + def __eq__(self, value: "Subscription") -> bool: + return ( + self.service_id == value.service_id + and self.instance_id == value.instance_id + and self.major_version == value.major_version + and self.eventgroup == value.eventgroup + and self.ttl_seconds == value.ttl_seconds + and self.client_endpoint == value.client_endpoint + and self.server_endpoint == value.server_endpoint + and self.protocols == value.protocols + ) + + def __hash__(self) -> int: + # Do not include the timestamp in the hash calculation + return hash( + ( + self.service_id, + self.instance_id, + self.major_version, + self.eventgroup, + self.ttl_seconds, + self.client_endpoint, + self.server_endpoint, + self.protocols, + ) + ) diff --git a/src/someipy/_internal/_daemon/subscription_storage.py b/src/someipy/_internal/_daemon/subscription_storage.py new file mode 100644 index 0000000..22cc1a9 --- /dev/null +++ b/src/someipy/_internal/_daemon/subscription_storage.py @@ -0,0 +1,96 @@ +# Copyright (C) 2025 Christian H. +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# This program is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with this program. If not, see . + + +from typing import Dict, List, Tuple + +from someipy._internal._daemon.subscription import Subscription + + +class SubscriptionStorage: + + def __init__(self): + self._subscriptions_by_client: Dict[int, List[Subscription]] = {} + + def add_subscription(self, client_id: int, subscription: Subscription): + if client_id not in self._subscriptions_by_client: + self._subscriptions_by_client[client_id] = [] + self._subscriptions_by_client[client_id].append(subscription) + else: + if subscription not in self._subscriptions_by_client[client_id]: + self._subscriptions_by_client[client_id].append(subscription) + + def remove_subscription(self, client_id: int, subscription: Subscription): + if client_id in self._subscriptions_by_client: + if subscription in self._subscriptions_by_client[client_id]: + self._subscriptions_by_client[client_id].remove(subscription) + + if len(self._subscriptions_by_client[client_id]) == 0: + del self._subscriptions_by_client[client_id] + + @property + def subscriptions(self) -> List[Subscription]: + """ + Get all subscriptions from all clients. + """ + subscriptions = [] + for client_subscriptions in self._subscriptions_by_client.values(): + subscriptions.extend(client_subscriptions) + return subscriptions + + def __len__(self) -> int: + """ + Get the total number of subscriptions across all clients. + """ + total = 0 + for client_subscriptions in self._subscriptions_by_client.values(): + total += len(client_subscriptions) + return total + + def get_client_ids(self, subscription: Subscription) -> List[int]: + """ + Get all client ids (writer ids) that have the given subscription. + """ + client_ids = [] + for client_id, subscriptions in self._subscriptions_by_client.items(): + if subscription in subscriptions: + client_ids.append(client_id) + return client_ids + + def has_subscriptions( + self, + service_id: int, + instance_id: int, + major_version: int, + ) -> List[Tuple[Subscription, int]]: + """ + Check if there are any subscriptions for the given service id, instance id, major version and protocol. + Returns a list of tuples containing the subscription and the writer id (UDS client). + """ + subscriptions_to_return = [] + for writer_id, subscriptions in self._subscriptions_by_client.items(): + for subscription in subscriptions: + if ( + subscription.service_id == service_id + and subscription.instance_id == instance_id + and subscription.major_version == major_version + ): + subscriptions_to_return.append((subscription, writer_id)) + + return subscriptions_to_return + + def remove_client(self, client_id: int): + if client_id in self._subscriptions_by_client: + del self._subscriptions_by_client[client_id] diff --git a/src/someipy/_internal/_sd/deserialization/sd_deserialization.py b/src/someipy/_internal/_sd/deserialization/sd_deserialization.py index cd0711b..428adc0 100644 --- a/src/someipy/_internal/_sd/deserialization/sd_deserialization.py +++ b/src/someipy/_internal/_sd/deserialization/sd_deserialization.py @@ -20,7 +20,7 @@ import socket import struct -from requests import options + from someipy._internal._sd.entries.find_service_entry import FindServiceEntry from someipy._internal._sd.entries.offer_service_entry import OfferServiceEntry from someipy._internal._sd.entries.stop_offer_service_entry import StopOfferServiceEntry diff --git a/src/someipy/_internal/offer_service_storage.py b/src/someipy/_internal/offer_service_storage.py index 0b5ecf2..a587879 100644 --- a/src/someipy/_internal/offer_service_storage.py +++ b/src/someipy/_internal/offer_service_storage.py @@ -14,6 +14,7 @@ # along with this program. If not, see . from typing import List +from someipy._internal._common.endpoint import Endpoint from someipy.service import Event, EventGroup, Method from someipy._internal.transport_layer_protocol import TransportLayerProtocol @@ -28,8 +29,7 @@ def __init__( minor_version: int, offer_ttl_seconds: int, cyclic_offer_delay_ms: int, - endpoint_ip: str, - endpoint_port: int, + endpoint: Endpoint, methods: List[Method], eventgroups: List[EventGroup], ): @@ -40,8 +40,7 @@ def __init__( self.minor_version = minor_version self.offer_ttl_seconds = offer_ttl_seconds self.cyclic_offer_delay_ms = cyclic_offer_delay_ms - self.endpoint_ip = endpoint_ip - self.endpoint_port = endpoint_port + self.endpoint = endpoint self.methods = methods self.eventgroups = eventgroups self.last_offer_time = None # Placeholder for last offer time @@ -57,8 +56,7 @@ def __eq__(self, other: object) -> bool: and self.minor_version == other.minor_version and self.offer_ttl_seconds == other.offer_ttl_seconds and self.cyclic_offer_delay_ms == other.cyclic_offer_delay_ms - and self.endpoint_ip == other.endpoint_ip - and self.endpoint_port == other.endpoint_port + and self.endpoint == other.endpoint and self.methods == other.methods and self.eventgroups == other.eventgroups ) @@ -73,8 +71,7 @@ def __hash__(self) -> int: self.minor_version, self.offer_ttl_seconds, self.cyclic_offer_delay_ms, - self.endpoint_ip, - self.endpoint_port, + self.endpoint, tuple(self.methods), # Convert list to tuple for hashing tuple(self.eventgroups), # Convert list to tuple for hashing ) diff --git a/src/someipy/_internal/someip_endpoint_factory.py b/src/someipy/_internal/someip_endpoint_factory.py new file mode 100644 index 0000000..3e0b211 --- /dev/null +++ b/src/someipy/_internal/someip_endpoint_factory.py @@ -0,0 +1,95 @@ +# Copyright (C) 2025 Christian H. +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# This program is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with this program. If not, see . + +import asyncio +from collections.abc import Callable +import logging +from typing import Tuple +from someipy._internal.someip_endpoint import ( + SomeipEndpoint, + TCPClientSomeipEndpoint, + TCPSomeipEndpoint, + UDPSomeipEndpoint, +) +from someipy._internal.someip_message import SomeIpMessage +from someipy._internal.tcp_client_manager import TcpClientManager, TcpClientProtocol +from someipy._internal.transport_layer_protocol import TransportLayerProtocol +from someipy._internal.utils import create_udp_socket + + +class SomeipEndpointFactory: + + @staticmethod + async def create_server_endpoint( + ip_address: str, + port: int, + protocol: TransportLayerProtocol, + someip_callback: Callable[ + [SomeIpMessage, Tuple[str, int], Tuple[str, int], TransportLayerProtocol], + None, + ], + ) -> SomeipEndpoint: + + if protocol == TransportLayerProtocol.UDP: + loop = asyncio.get_running_loop() + rcv_socket = create_udp_socket(ip_address, port) + + _, udp_endpoint = await loop.create_datagram_endpoint( + lambda: UDPSomeipEndpoint(ip_address, port), sock=rcv_socket + ) + + udp_endpoint.set_someip_callback(someip_callback) + + return udp_endpoint + else: + tcp_client_manager = TcpClientManager(ip_address, port) + loop = asyncio.get_running_loop() + server = await loop.create_server( + lambda: TcpClientProtocol(client_manager=tcp_client_manager), + ip_address, + port, + ) + tcp_someip_endpoint = TCPSomeipEndpoint( + server, tcp_client_manager, ip_address, port + ) + + tcp_someip_endpoint.set_someip_callback(someip_callback) + + return tcp_someip_endpoint + + @staticmethod + async def create_client_endpoint( + dst_ip: str, + dst_port: int, + src_ip: str, + src_port: int, + protocol: TransportLayerProtocol, + someip_message_callback: Callable[[SomeIpMessage], None], + logger: logging.Logger = None, + ) -> SomeipEndpoint: + if protocol == TransportLayerProtocol.UDP: + udp_endpoint = SomeipEndpointFactory.create_server_endpoint( + src_ip, + src_port, + TransportLayerProtocol.UDP, + someip_message_callback, + ) + return udp_endpoint + else: + tcp_endpoint = TCPClientSomeipEndpoint( + dst_ip, dst_port, src_ip, src_port, logger + ) + tcp_endpoint.set_someip_callback(someip_message_callback) + return tcp_endpoint diff --git a/src/someipy/_internal/someip_sd_builder.py b/src/someipy/_internal/someip_sd_builder.py index 34f93e0..37633ff 100644 --- a/src/someipy/_internal/someip_sd_builder.py +++ b/src/someipy/_internal/someip_sd_builder.py @@ -205,80 +205,6 @@ def build_subscribe_eventgroup_ack_sd_header( ) -def build_subscribe_eventgroup_sd_header( - service_id: int, - instance_id: int, - major_version: int, - ttl: int, - event_group_id: int, - session_id: int, - reboot_flag: bool, - endpoint: Tuple[ipaddress.IPv4Address, int], - protocols: Set[TransportLayerProtocol], -) -> SomeIpSdHeader: - sd_entry: SdEntry = SdEntry( - SdEntryType.SUBSCRIBE_EVENT_GROUP, - 0, # index_first_option - 0, # index_second_option - len(protocols), # num_options_1 - 0, # num_options_2 - service_id, - instance_id, - major_version, - ttl, - ) - entry = SdEventGroupEntry( - sd_entry=sd_entry, - initial_data_requested_flag=False, - counter=0, - eventgroup_id=event_group_id, - ) - - option_entry_common = SdOptionCommon( - length=SD_IPV4ENDPOINT_OPTION_LENGTH_VALUE, - type=SdOptionType.IPV4_ENDPOINT, - discardable_flag=False, - ) - - options = [] - for protocol in protocols: - if protocol not in (TransportLayerProtocol.TCP, TransportLayerProtocol.UDP): - raise ValueError( - f"Unsupported protocol {protocol} for SD IPV4 Endpoint option." - ) - - # Create an option for each protocol - sd_option_entry = SdIPV4EndpointOption( - sd_option_common=option_entry_common, - ipv4_address=endpoint[0], - protocol=protocol, - port=endpoint[1], - ) - options.append(sd_option_entry) - - # 20 bytes for header and length values of entries and options - # + length of entries array (1 entry) - # + length of options array - total_length = ( - 20 - + (1 * SD_SINGLE_ENTRY_LENGTH_BYTES) - + (len(options) * SD_BYTE_LENGTH_IP4ENDPOINT_OPTION) - ) - someip_header = SomeIpHeader.generate_sd_header( - length=total_length, session_id=session_id - ) - - return SomeIpSdHeader( - someip_header=someip_header, - reboot_flag=reboot_flag, - unicast_flag=True, - length_entries=(1 * SD_SINGLE_ENTRY_LENGTH_BYTES), - length_options=(len(options) * SD_BYTE_LENGTH_IP4ENDPOINT_OPTION), - service_entries=[entry], - options=options, - ) - - def build_find_service_sd_header( service_id: int, instance_id: int = 0xFFFF, diff --git a/src/someipy/someipyd.py b/src/someipy/someipyd.py index 3e4070b..718fb3f 100644 --- a/src/someipy/someipyd.py +++ b/src/someipy/someipyd.py @@ -28,15 +28,19 @@ import time from typing import Any, Dict, List, Set, Tuple, Union +from someipy._internal._common.endpoint import Endpoint from someipy._internal._daemon.daemon_server_client import ( ClientMessageEventArgs, DaemonServerClient, ) +from someipy._internal._daemon.subscription import Subscription +from someipy._internal._daemon.subscription_storage import SubscriptionStorage from someipy._internal._sd.deserialization.sd_deserialization import ( deserialize_sd_message, is_sd_message, ) from someipy._internal._sd.deserialization.sd_serialization import serialize_sd_message +from someipy._internal._sd.entries.offer_service_entry import OfferServiceEntry from someipy._internal._sd.entries.subscribe_eventgroup_entry import ( SubscribeEventGroupEntry, ) @@ -50,6 +54,7 @@ TCPSomeipEndpoint, UDPSomeipEndpoint, ) +from someipy._internal.someip_endpoint_factory import SomeipEndpointFactory from someipy._internal.someip_endpoint_storage import SomeipEndpointStorage from someipy._internal.someip_message import SomeIpMessage from someipy._internal.tcp_client_manager import TcpClientManager, TcpClientProtocol @@ -58,11 +63,9 @@ from someipy._internal.simple_timer import SimplePeriodicTimer from someipy._internal.someip_header import SomeIpHeader from someipy._internal.someip_sd_builder import ( - build_offer_service_sd_header, build_stop_offer_service_sd_header, build_subscribe_eventgroup_ack_entry, build_subscribe_eventgroup_ack_sd_header, - build_subscribe_eventgroup_sd_header, ) from someipy._internal.someip_sd_extractors import ( extract_subscribe_ack_eventgroup_entries, @@ -116,161 +119,6 @@ DEFAULT_TCP_PORT = 30500 -class Subscription: - def __init__( - self, - service_id: int, - instance_id: int, - major_version: int, - eventgroup: EventGroup, - ttl_seconds: int, - client_endpoint_ip: str, - client_endpoint_port: int, - server_endpoint_ip: str, - server_endpoint_port: int, - timestamp: float = time.time(), - ): - self.service_id = service_id - self.instance_id = instance_id - self.major_version = major_version - self.eventgroup = eventgroup - self.ttl_seconds = ttl_seconds - - self.client_endpoint_ip = client_endpoint_ip - self.client_endpoint_port = client_endpoint_port - self.server_endpoint_ip = server_endpoint_ip - self.server_endpoint_port = server_endpoint_port - self.timestamp = timestamp - - def __eq__(self, value: "Subscription") -> bool: - return ( - self.service_id == value.service_id - and self.instance_id == value.instance_id - and self.major_version == value.major_version - and self.eventgroup == value.eventgroup - and self.client_endpoint_ip == value.client_endpoint_ip - and self.client_endpoint_port == value.client_endpoint_port - and self.server_endpoint_ip == value.server_endpoint_ip - and self.server_endpoint_port == value.server_endpoint_port - ) - - def __hash__(self) -> int: - return hash( - ( - self.service_id, - self.instance_id, - self.major_version, - self.eventgroup, - self.client_endpoint_ip, - self.client_endpoint_port, - self.server_endpoint_ip, - self.server_endpoint_port, - ) - ) - - -class RequestedSubscription: - def __init__( - self, - service_id: int, - instance_id: int, - major_version: int, - client_endpoint_ip: str, - client_endpoint_port: int, - protocols: frozenset[TransportLayerProtocol], - eventgroup: EventGroup, - ttl_subscription: int, - ): - self.service_id = service_id - self.instance_id = instance_id - self.major_version = major_version - self.client_endpoint_ip = client_endpoint_ip - self.client_endpoint_port = client_endpoint_port - self.protocols = protocols - self.eventgroup = eventgroup - self.ttl = ttl_subscription - - def __eq__(self, other: "RequestedSubscription") -> bool: - return ( - self.service_id == other.service_id - and self.instance_id == other.instance_id - and self.major_version == other.major_version - and self.protocols == other.protocols - and self.eventgroup == other.eventgroup - and self.client_endpoint_ip == other.client_endpoint_ip - and self.client_endpoint_port == other.client_endpoint_port - and self.ttl == other.ttl - ) - - -class RequestedSubscriptionStore: - def __init__(self): - self._subscriptions_by_client: Dict[int, List[RequestedSubscription]] = {} - - def add_subscription(self, writer_id: int, subscription: RequestedSubscription): - if writer_id not in self._subscriptions_by_client: - self._subscriptions_by_client[writer_id] = [] - self._subscriptions_by_client[writer_id].append(subscription) - - else: - if subscription not in self._subscriptions_by_client[writer_id]: - self._subscriptions_by_client[writer_id].append(subscription) - - def remove_subscription(self, writer_id: int, subscription: RequestedSubscription): - if writer_id in self._subscriptions_by_client: - if subscription in self._subscriptions_by_client[writer_id]: - self._subscriptions_by_client[writer_id].remove(subscription) - - if len(self._subscriptions_by_client[writer_id]) == 0: - del self._subscriptions_by_client[writer_id] - - @property - def subscriptions(self) -> List[RequestedSubscription]: - """ - Get all subscriptions from all clients. - """ - subscriptions = [] - for client_subscriptions in self._subscriptions_by_client.values(): - subscriptions.extend(client_subscriptions) - return subscriptions - - def get_client_ids(self, subscription: RequestedSubscription) -> List[int]: - """ - Get all client ids (writer ids) that have the given subscription. - """ - client_ids = [] - for writer_id, subscriptions in self._subscriptions_by_client.items(): - if subscription in subscriptions: - client_ids.append(writer_id) - return client_ids - - def has_subscriptions( - self, - service_id: int, - instance_id: int, - major_version: int, - ) -> List[Tuple[RequestedSubscription, int]]: - """ - Check if there are any subscriptions for the given service id, instance id, major version and protocol. - Returns a list of tuples containing the subscription and the writer id (UDS client). - """ - subscriptions_to_return = [] - for writer_id, subscriptions in self._subscriptions_by_client.items(): - for subscription in subscriptions: - if ( - subscription.service_id == service_id - and subscription.instance_id == instance_id - and subscription.major_version == major_version - ): - subscriptions_to_return.append((subscription, writer_id)) - - return subscriptions_to_return - - def remove_client(self, writer_id: int): - if writer_id in self._subscriptions_by_client: - del self._subscriptions_by_client[writer_id] - - @dataclass class MethodCall: service_id: int @@ -296,11 +144,16 @@ def __hash__(self): class SomeipDaemon: def __init__( - self, server: DaemonServer, config: dict = None, logger: logging.Logger = None + self, + server: DaemonServer, + endpoint_factory: SomeipEndpointFactory, + config: dict = None, + logger: logging.Logger = None, ): self.config = config self.logger = logger + self._endpoint_factory = endpoint_factory self.sd_address = self.config.get("sd_address", DEFAULT_SD_ADDRESS) self.sd_port = self.config.get("sd_port", DEFAULT_SD_PORT) @@ -324,7 +177,7 @@ def __init__( self._service_subscribers: Dict[ServiceToOffer, Subscribers] = {} # Subscriptions requested by local clients - self._requested_subscriptions = RequestedSubscriptionStore() + self._requested_subscriptions = SubscriptionStorage() self._pending_subscriptions: Set[Subscription] = set() self._active_subscriptions: Set[Subscription] = set() @@ -388,37 +241,6 @@ async def client_disconnected( self.logger.debug(f"Client disconnected") - async def _create_server_endpoint( - self, ip: str, port: int, protocol: TransportLayerProtocol - ) -> SomeipEndpoint: - - if protocol == TransportLayerProtocol.UDP: - loop = asyncio.get_running_loop() - rcv_socket = create_udp_socket(ip, port) - - _, udp_endpoint = await loop.create_datagram_endpoint( - lambda: UDPSomeipEndpoint(ip, port), sock=rcv_socket - ) - - udp_endpoint.set_someip_callback(self._someip_message_callback) - - return udp_endpoint - else: - tcp_client_manager = TcpClientManager(ip, port) - loop = asyncio.get_running_loop() - server = await loop.create_server( - lambda: TcpClientProtocol(client_manager=tcp_client_manager), - ip, - port, - ) - tcp_someip_endpoint = TCPSomeipEndpoint( - server, tcp_client_manager, ip, port - ) - - tcp_someip_endpoint.set_someip_callback(self._someip_message_callback) - - return tcp_someip_endpoint - async def _check_services_ttl_task(self): try: while True: @@ -570,8 +392,8 @@ def _someip_message_callback( self.logger.debug("check event ids: %s", event_ids) if ( event_id in event_ids - and active_subscription.server_endpoint_ip == src_addr[0] - and active_subscription.server_endpoint_port == src_addr[1] + and str(active_subscription.server_endpoint.ip) == src_addr[0] + and active_subscription.server_endpoint.port == src_addr[1] and active_subscription.service_id == header.service_id ): self.logger.debug( @@ -592,10 +414,10 @@ def _someip_message_callback( == active_subscription.major_version and requested_subscription.eventgroup == active_subscription.eventgroup - and requested_subscription.client_endpoint_ip - == active_subscription.client_endpoint_ip - and requested_subscription.client_endpoint_port - == active_subscription.client_endpoint_port + and requested_subscription.client_endpoint.ip + == active_subscription.client_endpoint.ip + and requested_subscription.client_endpoint.port + == active_subscription.client_endpoint.port ): writer_ids = ( @@ -615,7 +437,10 @@ def _cleanup_active_subscriptions(self): current_time = time.time() subscriptions_to_remove = [] for subscription in self._active_subscriptions: - if current_time - subscription.timestamp > subscription.ttl_seconds: + if ( + current_time - subscription.timestamp_last_update + > subscription.ttl_seconds + ): subscriptions_to_remove.append(subscription) for subscription in subscriptions_to_remove: @@ -629,7 +454,7 @@ def _cleanup_obsolete_pending_subscriptions(self): current_time = time.time() subscriptions_to_remove = [] for subscription in self._pending_subscriptions: - if current_time - subscription.timestamp > 10.0: + if current_time - subscription.timestamp_last_update > 10.0: subscriptions_to_remove.append(subscription) for subscription in subscriptions_to_remove: @@ -733,7 +558,7 @@ async def handle_client_message( ) async def _handle_subscribe_eventgroup_request( - self, message: SubscribeEventGroupRequest, writer_id: int + self, message: SubscribeEventGroupRequest, client_id: int ): protocols = [] @@ -749,33 +574,34 @@ async def _handle_subscribe_eventgroup_request( f"Creating new UDP endpoint for {message['client_endpoint_ip']}:{message['client_endpoint_port']}" ) - udp_endpoint = await self._create_server_endpoint( + udp_endpoint = await self._endpoint_factory.create_server_endpoint( message["client_endpoint_ip"], message["client_endpoint_port"], TransportLayerProtocol.UDP, + self._someip_message_callback, ) - udp_endpoint.set_someip_callback(self._someip_message_callback) - - self._someip_client_endpoints.add_endpoint(writer_id, udp_endpoint) + self._someip_client_endpoints.add_endpoint(client_id, udp_endpoint) if message["tcp"]: protocols.append(TransportLayerProtocol.TCP) event_group = EventGroup.from_json(message["eventgroup"]) - new_subscription = RequestedSubscription( + new_subscription = Subscription( service_id=message["service_id"], instance_id=message["instance_id"], major_version=message["major_version"], - client_endpoint_ip=message["client_endpoint_ip"], - client_endpoint_port=message["client_endpoint_port"], + client_endpoint=Endpoint( + ip=message["client_endpoint_ip"], port=message["client_endpoint_port"] + ), + server_endpoint=None, protocols=frozenset(protocols), eventgroup=event_group, - ttl_subscription=message["ttl_subscription"], + ttl_seconds=message["ttl_subscription"], ) - self._requested_subscriptions.add_subscription(writer_id, new_subscription) + self._requested_subscriptions.add_subscription(client_id, new_subscription) def _handle_stop_subscribe_eventgroup_request( self, message: StopSubscribeEventGroupRequest, writer_id: int @@ -810,8 +636,9 @@ async def _handle_offer_service_request( minor_version=message["minor_version"], offer_ttl_seconds=message["ttl"], cyclic_offer_delay_ms=message["cyclic_offer_delay_ms"], - endpoint_ip=message["endpoint_ip"], - endpoint_port=message["endpoint_port"], + endpoint=Endpoint( + ipaddress.IPv4Address(message["endpoint_ip"]), message["endpoint_port"] + ), methods=methods, eventgroups=eventgroups, ) @@ -821,19 +648,21 @@ async def _handle_offer_service_request( # Check if there is already an endpoint for the ip and port, if not, open a new endpoint if service_to_add.has_udp: if not self._someip_server_endpoints.has_endpoint( - service_to_add.endpoint_ip, - service_to_add.endpoint_port, + str(service_to_add.endpoint.ip), + service_to_add.endpoint.port, TransportLayerProtocol.UDP, ): self.logger.debug( - f"Creating new UDP endpoint for {service_to_add.endpoint_ip}:{service_to_add.endpoint_port}" + f"Creating new UDP endpoint for {service_to_add.endpoint}" ) - udp_endpoint = await self._create_server_endpoint( - service_to_add.endpoint_ip, - service_to_add.endpoint_port, + udp_endpoint = await self._endpoint_factory.create_server_endpoint( + str(service_to_add.endpoint.ip), + service_to_add.endpoint.port, TransportLayerProtocol.UDP, + self._someip_message_callback, ) + self._someip_server_endpoints.add_endpoint(writer_id, udp_endpoint) if service_to_add.has_tcp: @@ -846,11 +675,13 @@ async def _handle_offer_service_request( f"Creating new TCP endpoint for {service_to_add.endpoint_ip}:{service_to_add.endpoint_port}" ) - tcp_endpoint = await self._create_server_endpoint( + tcp_endpoint = await self._endpoint_factory.create_server_endpoint( service_to_add.endpoint_ip, service_to_add.endpoint_port, TransportLayerProtocol.TCP, + self._someip_message_callback, ) + self._someip_server_endpoints.add_endpoint(writer_id, tcp_endpoint) cyclic_offer_delay_ms = message["cyclic_offer_delay_ms"] @@ -989,15 +820,16 @@ async def _handle_outbound_call_method_request( f"Creating new UDP endpoint for {message['src_endpoint_ip']}:{message['src_endpoint_port']}" ) - udp_endpoint = await self._create_server_endpoint( + udp_endpoint = await self._endpoint_factory.create_client_endpoint( + message["dst_endpoint_ip"], + message["dst_endpoint_port"], message["src_endpoint_ip"], message["src_endpoint_port"], TransportLayerProtocol.UDP, + self._someip_message_callback, + self.logger, ) - udp_endpoint.set_someip_callback(self._someip_message_callback) - - self._someip_client_endpoints.add_endpoint(writer_id, udp_endpoint) endpoint = udp_endpoint else: endpoint = self._someip_client_endpoints.get_endpoint_by_ip_port( @@ -1018,16 +850,16 @@ async def _handle_outbound_call_method_request( f"Creating new TCP endpoint for {message['src_endpoint_ip']}:{message['src_endpoint_port']}" ) - tcp_endpoint = TCPClientSomeipEndpoint( + tcp_endpoint = self._endpoint_factory.create_client_endpoint( message["dst_endpoint_ip"], message["dst_endpoint_port"], message["src_endpoint_ip"], message["src_endpoint_port"], + TransportLayerProtocol.TCP, + self._someip_message_callback, self.logger, ) - tcp_endpoint.set_someip_callback(self._someip_message_callback) - self._someip_client_endpoints.add_endpoint(writer_id, tcp_endpoint) endpoint: TCPClientSomeipEndpoint = tcp_endpoint else: @@ -1310,13 +1142,69 @@ def offer_timer_callback(self, cyclic_offer_delay_ms: int): reboot_flag, ) = self._mcast_session_handler.update_session() - sd_message = build_offer_service_sd_header( - services_to_offer, session_id, reboot_flag - ) - buffer = sd_message.to_buffer() + options = set() + for service in services_to_offer: + if service.has_udp: + options.add( + IpV4EndpointOption( + address=service.endpoint.ip, + protocol=TransportLayerProtocol.UDP, + port=service.endpoint.port, + ) + ) + if service.has_tcp: + options.add( + IpV4EndpointOption( + address=service.endpoint.ip, + protocol=TransportLayerProtocol.TCP, + port=service.endpoint.port, + ) + ) + + options = list(options) + sd_message = SdMessage() + sd_message.session_id = session_id + sd_message.reboot_flag = reboot_flag + + for service in services_to_offer: + + endpoints = [] + if service.has_udp: + endpoints.extend( + [ + option + for option in options + if option.protocol == TransportLayerProtocol.UDP + and option.address == service.endpoint.ip + and option.port == service.endpoint.port + ] + ) + if service.has_tcp: + endpoints.extend( + [ + option + for option in options + if option.protocol == TransportLayerProtocol.TCP + and option.address == service.endpoint.ip + and option.port == service.endpoint.port + ] + ) + + new_entry = OfferServiceEntry( + service_id=service.service_id, + instance_id=service.instance_id, + major_version=service.major_version, + minor_version=service.minor_version, + ttl=service.offer_ttl_seconds, + ip_v4_endpoints=endpoints, + ip_v6_endpoints=[], + ) + sd_message.entries.append(new_entry) if self._ucast_transport: - self._ucast_transport.sendto(buffer, (self.sd_address, self.sd_port)) + self._ucast_transport.sendto( + serialize_sd_message(sd_message), (self.sd_address, self.sd_port) + ) def prepare_message(self, message: dict): payload = json.dumps(message).encode("utf-8") @@ -1408,27 +1296,29 @@ def _handle_offered_service(self, offered_service: ServiceInstance): if TransportLayerProtocol.TCP in requested_protocols: if not self._someip_client_endpoints.has_tcp_endpoint( - requested_subscription[0].client_endpoint_ip, - requested_subscription[0].client_endpoint_port, + str(requested_subscription[0].client_endpoint.ip), + requested_subscription[0].client_endpoint.port, str(offered_service.endpoint[0]), offered_service.endpoint[1], ): self.logger.debug( - f"Creating new TCP endpoint for {requested_subscription[0].client_endpoint_ip}:{requested_subscription[0].client_endpoint_port}" + f"Creating new TCP endpoint for {requested_subscription[0].client_endpoint}" ) - tcp_endpoint = TCPClientSomeipEndpoint( - str(offered_service.endpoint[0]), - offered_service.endpoint[1], - requested_subscription[0].client_endpoint_ip, - requested_subscription[0].client_endpoint_port, + tcp_endpoint = self._endpoint_factory.create_client_endpoint( + str(offered_service.endpoint.ip), + offered_service.endpoint.port, + str(requested_subscription[0].client_endpoint.ip), + requested_subscription[0].client_endpoint.port, + TransportLayerProtocol.TCP, + self._someip_message_callback, self.logger, ) - tcp_endpoint.set_someip_callback(self._someip_message_callback) self._someip_client_endpoints.add_endpoint( requested_subscription[1], tcp_endpoint ) + # TODO: This shall not block the handle_client function. A new task shall be created # For TCP wait for the connection to be established # while not tcp_endpoint.is_connected(): @@ -1447,11 +1337,9 @@ def _handle_offered_service(self, offered_service: ServiceInstance): for protocol in requested_protocols: options.append( IpV4EndpointOption( - address=ipaddress.IPv4Address( - requested_subscription[0].client_endpoint_ip - ), + address=requested_subscription[0].client_endpoint.ip, protocol=protocol, - port=requested_subscription[0].client_endpoint_port, + port=requested_subscription[0].client_endpoint.port, ) ) @@ -1460,7 +1348,7 @@ def _handle_offered_service(self, offered_service: ServiceInstance): instance_id=offered_service.instance_id, major_version=offered_service.major_version, minor_version=offered_service.minor_version, - ttl=requested_subscription[0].ttl, + ttl=requested_subscription[0].ttl_seconds, eventgroup_id=requested_subscription[0].eventgroup.id, counter=0, ip_v4_endpoints=options, @@ -1468,24 +1356,26 @@ def _handle_offered_service(self, offered_service: ServiceInstance): ) sd_message.entries.append(entry) + client_endpoint = requested_subscription[0].client_endpoint + server_endpoint = offered_service.endpoint + pending_subscription = Subscription( service_id=offered_service.service_id, instance_id=offered_service.instance_id, major_version=offered_service.major_version, eventgroup=requested_subscription[0].eventgroup, - ttl_seconds=requested_subscription[0].ttl, - client_endpoint_ip=requested_subscription[0].client_endpoint_ip, - client_endpoint_port=requested_subscription[0].client_endpoint_port, - server_endpoint_ip=str(offered_service.endpoint[0]), - server_endpoint_port=offered_service.endpoint[1], - timestamp=time.time(), + ttl_seconds=requested_subscription[0].ttl_seconds, + client_endpoint=client_endpoint, + server_endpoint=server_endpoint, + protocols=frozenset(requested_protocols), + timestamp_last_update=time.time(), ) self._pending_subscriptions.add(pending_subscription) if self._ucast_transport: self._ucast_transport.sendto( serialize_sd_message(sd_message), - (str(offered_service.endpoint[0]), self.sd_port), + (str(offered_service.endpoint.ip), self.sd_port), ) def _handle_subscription( @@ -1608,7 +1498,7 @@ def _handle_sd_subscribe_ack_eventgroup_entry( break if pending_subscription is not None: - pending_subscription.timestamp = time.time() + pending_subscription.timestamp_last_update = time.time() self._active_subscriptions.discard(pending_subscription) self._active_subscriptions.add(pending_subscription) @@ -1804,7 +1694,7 @@ async def async_main(): daemon_server = DaemonServer(logger) - daemon = SomeipDaemon(daemon_server, config, logger) + daemon = SomeipDaemon(daemon_server, SomeipEndpointFactory(), config, logger) daemon_server.client_connected += daemon.new_client_connected daemon_server.client_disconnected += daemon.client_disconnected diff --git a/tests/common/__init__.py b/tests/common/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/common/test_endpoint.py b/tests/common/test_endpoint.py new file mode 100644 index 0000000..0c7f7fd --- /dev/null +++ b/tests/common/test_endpoint.py @@ -0,0 +1,28 @@ +import ipaddress +from someipy._internal._common.endpoint import Endpoint + + +def test_equality(): + endpoint_1 = Endpoint(ipaddress.ip_address("192.168.0.1"), 3000) + endpoint_2 = Endpoint(ipaddress.ip_address("192.168.0.1"), 3000) + + assert endpoint_1 == endpoint_2 + + +def test_inequality(): + endpoint_1 = Endpoint(ipaddress.ip_address("192.168.0.1"), 3000) + endpoint_2 = Endpoint(ipaddress.ip_address("192.168.0.2"), 3000) + + assert endpoint_1 != endpoint_2 + + +def test_endpoint_is_ipv4(): + endpoint_ipv4 = Endpoint(ipaddress.ip_address("192.168.0.1"), 3000) + assert endpoint_ipv4.is_ipv4 == True + + +def test_endpoint_is_ipv6(): + endpoint_ipv6 = Endpoint( + ipaddress.ip_address("2001:0db8:85a3:0000:0000:8a2e:0370:7334"), 3000 + ) + assert endpoint_ipv6.is_ipv4 == False diff --git a/tests/daemon/__init__.py b/tests/daemon/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/daemon/test_subscription_storage.py b/tests/daemon/test_subscription_storage.py new file mode 100644 index 0000000..55b4b6b --- /dev/null +++ b/tests/daemon/test_subscription_storage.py @@ -0,0 +1,162 @@ +import pytest +from someipy._internal._daemon.subscription import Subscription +from someipy._internal._daemon.subscription_storage import SubscriptionStorage +from someipy._internal.transport_layer_protocol import TransportLayerProtocol +from someipy.service import EventGroup + + +@pytest.fixture +def subscription_service_id_1() -> Subscription: + return Subscription( + service_id=1, + instance_id=1, + major_version=1, + eventgroup=EventGroup(0, []), + ttl_seconds=60, + client_endpoint=None, + server_endpoint=None, + protocols=frozenset([TransportLayerProtocol.TCP]), + timestamp_last_update=0, + ) + + +@pytest.fixture +def subscription_service_id_2() -> Subscription: + return Subscription( + service_id=2, + instance_id=1, + major_version=1, + eventgroup=EventGroup(0, []), + ttl_seconds=60, + client_endpoint=None, + server_endpoint=None, + protocols=frozenset([TransportLayerProtocol.TCP]), + timestamp_last_update=0, + ) + + +def test_add_subscription_adds_subscription(subscription_service_id_1): + + storage = SubscriptionStorage() + storage.add_subscription(100, subscription_service_id_1) + assert subscription_service_id_1 in storage.subscriptions + + +def test_add_subscription_adds_multiple_subscriptions_per_client( + subscription_service_id_1, subscription_service_id_2 +): + + storage = SubscriptionStorage() + storage.add_subscription(100, subscription_service_id_1) + storage.add_subscription(100, subscription_service_id_2) + assert subscription_service_id_1 in storage.subscriptions + assert subscription_service_id_2 in storage.subscriptions + + assert storage.get_client_ids(subscription_service_id_1) == [100] + assert storage.get_client_ids(subscription_service_id_2) == [100] + + +def test_get_client_ids_returns_multiple_clients(subscription_service_id_1): + + storage = SubscriptionStorage() + storage.add_subscription(100, subscription_service_id_1) + storage.add_subscription(200, subscription_service_id_1) + + client_ids = storage.get_client_ids(subscription_service_id_1) + assert 100 in client_ids + assert 200 in client_ids + + +def test_remove_subscription_removes_subscription(subscription_service_id_1): + + storage = SubscriptionStorage() + storage.add_subscription(100, subscription_service_id_1) + storage.remove_subscription(100, subscription_service_id_1) + assert subscription_service_id_1 not in storage.subscriptions + + +def test_remove_client_removes_all_subscriptions( + subscription_service_id_1, subscription_service_id_2 +): + + storage = SubscriptionStorage() + storage.add_subscription(100, subscription_service_id_1) + storage.add_subscription(100, subscription_service_id_2) + + storage.remove_client(100) + + assert subscription_service_id_1 not in storage.subscriptions + assert subscription_service_id_2 not in storage.subscriptions + + +def test_has_subscriptions_returns_empty_list_when_no_subscriptions( + subscription_service_id_1, subscription_service_id_2 +): + storage = SubscriptionStorage() + storage.add_subscription(100, subscription_service_id_1) + result = storage.has_subscriptions( + service_id=subscription_service_id_2.service_id, + instance_id=subscription_service_id_1.instance_id, + major_version=subscription_service_id_1.major_version, + ) + assert result == [] + + +def test_has_subscriptions_returns_matching_subscriptions( + subscription_service_id_1, subscription_service_id_2 +): + storage = SubscriptionStorage() + storage.add_subscription(100, subscription_service_id_1) + storage.add_subscription(200, subscription_service_id_2) + + result = storage.has_subscriptions( + service_id=subscription_service_id_1.service_id, + instance_id=subscription_service_id_1.instance_id, + major_version=subscription_service_id_1.major_version, + ) + assert len(result) == 1 + subscription, client_id = result[0] + assert subscription == subscription_service_id_1 + assert client_id == 100 + + +def test_has_subscriptions_returns_multiple_matching_subscriptions( + subscription_service_id_1, subscription_service_id_2 +): + storage = SubscriptionStorage() + storage.add_subscription(100, subscription_service_id_1) + storage.add_subscription(200, subscription_service_id_1) + storage.add_subscription(300, subscription_service_id_2) + + result = storage.has_subscriptions( + service_id=subscription_service_id_1.service_id, + instance_id=subscription_service_id_1.instance_id, + major_version=subscription_service_id_1.major_version, + ) + assert len(result) == 2 + client_ids = [client_id for _, client_id in result] + assert 100 in client_ids + assert 200 in client_ids + + +def test_len(subscription_service_id_1, subscription_service_id_2): + storage = SubscriptionStorage() + assert len(storage) == 0 + + storage.add_subscription(100, subscription_service_id_1) + assert len(storage) == 1 + + storage.add_subscription(200, subscription_service_id_1) + assert len(storage) == 2 + + storage.add_subscription(100, subscription_service_id_2) + assert len(storage) == 3 + + storage.remove_subscription(100, subscription_service_id_1) + assert len(storage) == 2 + + storage.remove_subscription(200, subscription_service_id_1) + assert len(storage) == 1 + + storage.remove_subscription(100, subscription_service_id_2) + assert len(storage) == 0 diff --git a/tests/test_someipyd.py b/tests/test_someipyd.py index 454e17f..c71519d 100644 --- a/tests/test_someipyd.py +++ b/tests/test_someipyd.py @@ -1,8 +1,18 @@ +from asyncio import DatagramTransport import logging import pytest +import pytest_asyncio from unittest.mock import Mock from someipy._internal._common.endpoint import Endpoint +from someipy._internal._daemon.uds_messages import ( + OfferServiceRequest, + SubscribeEventGroupRequest, + create_uds_message, +) from someipy._internal._sd.service_instance import ServiceInstance +from someipy._internal.someip_endpoint_factory import SomeipEndpointFactory +from someipy._internal.transport_layer_protocol import TransportLayerProtocol +from someipy.service import Event, EventGroup, Method from someipy.someipyd import DaemonServer, SomeipDaemon @@ -18,24 +28,120 @@ def mock_daemon_server(mock_logger) -> DaemonServer: @pytest.fixture -def daemon(mock_daemon_server, mock_logger) -> SomeipDaemon: - config = {} +def mock_endpoint_factory() -> Mock: + return Mock(spec=SomeipEndpointFactory) + - return SomeipDaemon(mock_daemon_server, config, mock_logger) +@pytest.fixture +def daemon(mock_daemon_server, mock_endpoint_factory, mock_logger) -> SomeipDaemon: + config = {} + daemon = SomeipDaemon( + mock_daemon_server, mock_endpoint_factory, config, mock_logger + ) + daemon._ucast_transport = Mock(spec=DatagramTransport) + return daemon -def test_handle_offered_service_adds_service(daemon: SomeipDaemon): +@pytest.fixture +def service_instance() -> ServiceInstance: service_instance = ServiceInstance( service_id=1, - instance_id=1, - major_version=1, + instance_id=2, + major_version=3, minor_version=0, ttl=10, endpoint=Endpoint("123", 1), - protocols=frozenset(), + protocols=frozenset([TransportLayerProtocol.UDP]), timestamp=1000, ) + return service_instance + + +@pytest.fixture +def eventgroup() -> EventGroup: + return EventGroup( + id=1, + events=[ + Event(id=1, protocol=TransportLayerProtocol.UDP), + Event(id=2, protocol=TransportLayerProtocol.UDP), + ], + ) + + +@pytest.fixture +def method() -> Method: + return Method( + id=1, + protocol=TransportLayerProtocol.UDP, + ) + + +@pytest.fixture +def subscribe_event_group_request_udp(eventgroup) -> SubscribeEventGroupRequest: + return SubscribeEventGroupRequest( + service_id=1, + instance_id=2, + major_version=3, + ttl_subscription=10, + eventgroup=eventgroup.to_json(), + client_endpoint_ip="123", + client_endpoint_port=1, + udp=True, + tcp=False, + ) + + +@pytest.fixture +def subscribe_event_group_request_tcp(eventgroup) -> SubscribeEventGroupRequest: + return SubscribeEventGroupRequest( + service_id=1, + instance_id=2, + major_version=3, + ttl_subscription=10, + eventgroup=eventgroup.to_json(), + client_endpoint_ip="123", + client_endpoint_port=1, + udp=False, + tcp=True, + ) + + +@pytest.fixture +def subscribe_event_group_request_udp_and_tcp(eventgroup) -> SubscribeEventGroupRequest: + return create_uds_message( + SubscribeEventGroupRequest, + service_id=1, + instance_id=2, + major_version=3, + ttl_subscription=10, + eventgroup=eventgroup.to_json(), + client_endpoint_ip="123", + client_endpoint_port=1, + udp=True, + tcp=True, + ) + + +@pytest.fixture +def offer_service_request(eventgroup, method) -> OfferServiceRequest: + return create_uds_message( + OfferServiceRequest, + service_id=1, + instance_id=2, + major_version=3, + minor_version=0, + endpoint_ip="127.0.0.1", + endpoint_port=1, + ttl=5, + eventgroup_list=[eventgroup.to_json()], + method_list=[method.to_json()], + cyclic_offer_delay_ms=1000, + ) + +def test_handle_offered_service_adds_service( + daemon: SomeipDaemon, service_instance: ServiceInstance +): # Initially, the found services list should be empty assert len(daemon._found_services) == 0 @@ -47,17 +153,9 @@ def test_handle_offered_service_adds_service(daemon: SomeipDaemon): assert daemon._found_services[0] == service_instance -def test_handle_offered_service_updates_timestamp(daemon: SomeipDaemon): - service_instance = ServiceInstance( - service_id=1, - instance_id=1, - major_version=1, - minor_version=0, - ttl=10, - endpoint=Endpoint("123", 1), - protocols=frozenset(), - timestamp=1000, - ) +def test_handle_offered_service_updates_timestamp( + daemon: SomeipDaemon, service_instance: ServiceInstance +): # Initially, the found services list should be empty assert len(daemon._found_services) == 0 @@ -66,13 +164,13 @@ def test_handle_offered_service_updates_timestamp(daemon: SomeipDaemon): daemon._handle_offered_service(service_instance) service_instance_2 = ServiceInstance( - service_id=1, - instance_id=1, - major_version=1, - minor_version=0, - ttl=10, - endpoint=Endpoint("123", 1), - protocols=frozenset(), + service_id=service_instance.service_id, + instance_id=service_instance.instance_id, + major_version=service_instance.major_version, + minor_version=service_instance.minor_version, + ttl=service_instance.ttl, + endpoint=service_instance.endpoint, + protocols=service_instance.protocols, timestamp=2000, ) @@ -83,3 +181,65 @@ def test_handle_offered_service_updates_timestamp(daemon: SomeipDaemon): assert len(daemon._found_services) == 1 assert daemon._found_services[0] == service_instance assert daemon._found_services[0].timestamp == 2000 + + +@pytest.mark.asyncio +async def test_handle_offered_service_opens_tcp_client_endpoint( + daemon: SomeipDaemon, + mock_endpoint_factory: Mock, + service_instance: ServiceInstance, + subscribe_event_group_request_tcp: SubscribeEventGroupRequest, +): + service_instance.protocols = frozenset([TransportLayerProtocol.TCP]) + + # Add a subscription first + await daemon._handle_subscribe_eventgroup_request( + subscribe_event_group_request_tcp, 1 + ) + + # Handle the offered service + daemon._handle_offered_service(service_instance) + + # Verify that create_client_endpoint was called for TCP + mock_endpoint_factory.create_client_endpoint.assert_called_once_with( + str(service_instance.endpoint.ip), + service_instance.endpoint.port, + str(service_instance.endpoint.ip), + service_instance.endpoint.port, + TransportLayerProtocol.TCP, + daemon._someip_message_callback, + daemon.logger, + ) + + assert len(daemon._pending_subscriptions) == 1 + assert len(daemon._someip_client_endpoints) == 1 + daemon._ucast_transport.sendto.assert_called_once() + + +@pytest.mark.asyncio +async def test_handle_subscribe_eventgroup_request_adds_requested_subscription( + daemon: SomeipDaemon, + subscribe_event_group_request_udp: SubscribeEventGroupRequest, +): + # Initially, the requested subscriptions list should be empty + assert len(daemon._requested_subscriptions) == 0 + + # Handle the subscribe event group request + await daemon._handle_subscribe_eventgroup_request( + subscribe_event_group_request_udp, 1 + ) + + # Now, the requested subscriptions list should contain the new subscription + assert len(daemon._requested_subscriptions) == 1 + + +@pytest.mark.asyncio +async def test_handle_offer_service_request_opens_server_endpoint( + daemon: SomeipDaemon, + offer_service_request: OfferServiceRequest, +): + assert len(daemon._someip_server_endpoints) == 0 + + await daemon._handle_offer_service_request(offer_service_request, 1) + + assert len(daemon._someip_server_endpoints) == 1 From 88fb1f69f0b6eea4dbaa915e3deb016288e87dbc Mon Sep 17 00:00:00 2001 From: Christian Date: Mon, 29 Dec 2025 10:58:41 +0100 Subject: [PATCH 4/6] Add coverage report to github actions --- .github/workflows/tests.yml | 26 +++++++++++++++++++++++++- README.md | 4 ++++ 2 files changed, 29 insertions(+), 1 deletion(-) diff --git a/.github/workflows/tests.yml b/.github/workflows/tests.yml index 945bdec..38c606c 100644 --- a/.github/workflows/tests.yml +++ b/.github/workflows/tests.yml @@ -27,6 +27,30 @@ jobs: pip install -e . pip install pytest pip install pytest-asyncio + pip install coverage - name: Run pytest - run: pytest tests + run: coverage run -m pytest tests/ + + - name: Generate coverage HTML report + run: coverage html + + - name: Upload coverage report + uses: actions/upload-artifact@v4 + with: + name: coverage-report + path: './htmlcov' + + - name: Generate coverage json + run: coverage json -o coverage.json + + - name: Upload coverage json to gist + env: + GIST_ID: ${{ secrets.GIST_ID }} + GIST_TOKEN: ${{ secrets.GIST_TOKEN }} + run: | + # Create a GitHub CLI authentication token + echo "$GIST_TOKEN" | gh auth login --with-token + + # Update existing gist + gh gist edit $GIST_ID coverage.json \ No newline at end of file diff --git a/README.md b/README.md index 136bbdd..d49b251 100644 --- a/README.md +++ b/README.md @@ -1,5 +1,9 @@ # someipy - A Python Library implementing the SOME/IP Protocol +![Dynamic JSON Badge](https://img.shields.io/badge/dynamic/json?url=https%3A%2F%2Fgist.githubusercontent.com%2Fchrizog%2F6a2e6f355eedf38ae3af74dcdb7b30a1%2Fraw%2F61630b4d2cc9c04aae78b3cca0cbc9c56e76f0e5%2Fcoverage_someipy_v1.json&query=totals.percent_covered_display&suffix=%20%25&label=Coverage) + + + ## Get in Contact :postbox: If you want to connect, have a feature request, bug report or need support, send me an email or connect on LinkedIn: From f28d8d1ee210af1462790a2a7912cf1090beada9 Mon Sep 17 00:00:00 2001 From: Christian Date: Fri, 2 Jan 2026 20:02:02 +0100 Subject: [PATCH 5/6] * Extend unit tests for daemon * Major fixes in someipyd --- example_apps/receive_events_tcp.py | 2 +- example_apps/receive_events_udp.py | 2 +- .../_internal/_daemon/daemon_server_client.py | 19 +- .../{ => _daemon}/offer_service_storage.py | 3 + .../_sd/deserialization/sd_deserialization.py | 22 +- src/someipy/_internal/_sd/sd_message.py | 17 + .../_internal/_sd/sd_message_creator.py | 156 ++++++++++ src/someipy/_internal/_sd/service_instance.py | 16 + .../_internal/someip_endpoint_factory.py | 17 +- src/someipy/_internal/someip_sd_builder.py | 135 -------- src/someipy/client_service_instance.py | 3 + src/someipy/someipyd.py | 291 ++++++++---------- tests/sd/test_sd_deserialization.py | 70 +++++ tests/test_someipyd.py | 179 ++++++++++- 14 files changed, 619 insertions(+), 313 deletions(-) rename src/someipy/_internal/{ => _daemon}/offer_service_storage.py (98%) create mode 100644 src/someipy/_internal/_sd/sd_message_creator.py diff --git a/example_apps/receive_events_tcp.py b/example_apps/receive_events_tcp.py index d4c670f..5d18987 100644 --- a/example_apps/receive_events_tcp.py +++ b/example_apps/receive_events_tcp.py @@ -87,7 +87,7 @@ async def main(): service_instance_temperature.register_callback(temperature_callback) # The second argument is the time to live (TTL) of the subscription in seconds - service_instance_temperature.subscribe_eventgroup(temperature_eventgroup, 5.0) + service_instance_temperature.subscribe_eventgroup(temperature_eventgroup, 5) try: # Keep the task alive diff --git a/example_apps/receive_events_udp.py b/example_apps/receive_events_udp.py index 7c133ad..b6b7cdb 100644 --- a/example_apps/receive_events_udp.py +++ b/example_apps/receive_events_udp.py @@ -86,7 +86,7 @@ async def main(): service_instance_temperature.register_callback(temperature_callback) # The second argument is the time to live (TTL) of the subscription in seconds - service_instance_temperature.subscribe_eventgroup(temperature_eventgroup, 5.0) + service_instance_temperature.subscribe_eventgroup(temperature_eventgroup, 5) try: # Keep the task alive diff --git a/src/someipy/_internal/_daemon/daemon_server_client.py b/src/someipy/_internal/_daemon/daemon_server_client.py index bfe661b..1a41cc7 100644 --- a/src/someipy/_internal/_daemon/daemon_server_client.py +++ b/src/someipy/_internal/_daemon/daemon_server_client.py @@ -13,12 +13,12 @@ def __init__( self, reader: asyncio.StreamReader, writer: asyncio.StreamWriter, - id: int, + client_id: int, logger: logging.Logger = None, ): self._reader = reader self._writer = writer - self._id = id + self._id = client_id self._logger = logger self.message_received: Event[ClientMessageEventArgs] = Event() @@ -65,7 +65,9 @@ async def read_next_message(self) -> Optional[dict]: message_buffer += data if len(message_buffer) == message_length: - self._logger.debug(f"Client sent message: {message_buffer}") + self._logger.debug( + f"Client {self.id} sent message: {message_buffer}" + ) json_message = json.loads(message_buffer.decode("utf-8")) await self.message_received.invoke( self, ClientMessageEventArgs(self, json_message) @@ -77,11 +79,16 @@ async def read_next_message(self) -> Optional[dict]: raise Exception("Client sent too much message data.") return None - def send(self, message: bytes): - pass + async def send(self, data: bytes): + self._writer.write(data) + await self._writer.drain() + + async def close(self): + self._writer.close() + await self._writer.wait_closed() @property - def id(self) -> str: + def id(self) -> int: return self._id diff --git a/src/someipy/_internal/offer_service_storage.py b/src/someipy/_internal/_daemon/offer_service_storage.py similarity index 98% rename from src/someipy/_internal/offer_service_storage.py rename to src/someipy/_internal/_daemon/offer_service_storage.py index a587879..53e92ac 100644 --- a/src/someipy/_internal/offer_service_storage.py +++ b/src/someipy/_internal/_daemon/offer_service_storage.py @@ -166,6 +166,9 @@ def clear(self) -> None: """Remove all services from the storage""" self._services.clear() + def __len__(self) -> int: + return len(self._services) + @property def cyclic_offer_delays(self) -> List[int]: """Returns a list of all unique cyclic_offer_delay_ms values from stored services""" diff --git a/src/someipy/_internal/_sd/deserialization/sd_deserialization.py b/src/someipy/_internal/_sd/deserialization/sd_deserialization.py index 428adc0..4dd7507 100644 --- a/src/someipy/_internal/_sd/deserialization/sd_deserialization.py +++ b/src/someipy/_internal/_sd/deserialization/sd_deserialization.py @@ -340,7 +340,7 @@ def deserialize_sd_message( ) if ( - common_entry_data.type == SdEntryOnWireType.OFFER_SERVICE.value + common_entry_data.type_field_value == SdEntryOnWireType.OFFER_SERVICE.value and common_entry_data.ttl != 0 ): (minor_version,) = struct.unpack( @@ -373,7 +373,8 @@ def deserialize_sd_message( entries.append(offer_service_entry) elif ( - common_entry_data.type == SdEntryOnWireType.STOP_OFFER_SERVICE.value + common_entry_data.type_field_value + == SdEntryOnWireType.STOP_OFFER_SERVICE.value and common_entry_data.ttl == 0 ): (minor_version,) = struct.unpack( @@ -395,7 +396,6 @@ def deserialize_sd_message( instance_id=common_entry_data.instance_id, major_version=common_entry_data.major_version, minor_version=minor_version, - ttl=common_entry_data.ttl, ip_v4_endpoints=[ o for o in applicable_options if isinstance(o, IpV4EndpointOption) ], @@ -405,7 +405,7 @@ def deserialize_sd_message( ) entries.append(stop_offer_service_entry) - elif common_entry_data.type == SdEntryOnWireType.FIND_SERVICE.value: + elif common_entry_data.type_field_value == SdEntryOnWireType.FIND_SERVICE.value: (minor_version,) = struct.unpack( ">I", data[start_entry + 12 : start_entry + 16] ) @@ -415,12 +415,12 @@ def deserialize_sd_message( instance_id=common_entry_data.instance_id, major_version=common_entry_data.major_version, minor_version=minor_version, - ttl=common_entry_data.ttl, ) entries.append(find_service_entry) elif ( - common_entry_data.type == SdEntryOnWireType.SUBSCRIBE_EVENT_GROUP.value + common_entry_data.type_field_value + == SdEntryOnWireType.SUBSCRIBE_EVENT_GROUP.value and common_entry_data.ttl != 0 ): initial_data_requested_flag_counter_value, eventgroup_id = struct.unpack( @@ -459,7 +459,8 @@ def deserialize_sd_message( entries.append(subscribe_eventgroup_entry) elif ( - common_entry_data.type == SdEntryOnWireType.STOP_SUBSCRIBE_EVENT_GROUP.value + common_entry_data.type_field_value + == SdEntryOnWireType.STOP_SUBSCRIBE_EVENT_GROUP.value and common_entry_data.ttl == 0 ): initial_data_requested_flag_counter_value, eventgroup_id = struct.unpack( @@ -496,7 +497,8 @@ def deserialize_sd_message( entries.append(stop_subscribe_eventgroup_entry) elif ( - common_entry_data.type == SdEntryOnWireType.SUBSCRIBE_EVENT_GROUP_ACK.value + common_entry_data.type_field_value + == SdEntryOnWireType.SUBSCRIBE_EVENT_GROUP_ACK.value and common_entry_data.ttl != 0 ): initial_data_requested_flag_counter_value, eventgroup_id = struct.unpack( @@ -535,7 +537,8 @@ def deserialize_sd_message( entries.append(subscribe_ack_eventgroup_entry) elif ( - common_entry_data.type == SdEntryOnWireType.SUBSCRIBE_EVENT_GROUP_NACK.value + common_entry_data.type_field_value + == SdEntryOnWireType.SUBSCRIBE_EVENT_GROUP_NACK.value and common_entry_data.ttl == 0 ): initial_data_requested_flag_counter_value, eventgroup_id = struct.unpack( @@ -571,5 +574,6 @@ def deserialize_sd_message( sd_message.source_port = port sd_message.multicast = multicast sd_message.session_id = session_id + sd_message.reboot_flag = reboot_flag sd_message.entries = entries return sd_message diff --git a/src/someipy/_internal/_sd/sd_message.py b/src/someipy/_internal/_sd/sd_message.py index 45803bb..30dc129 100644 --- a/src/someipy/_internal/_sd/sd_message.py +++ b/src/someipy/_internal/_sd/sd_message.py @@ -1,3 +1,19 @@ +# Copyright (C) 2025 Christian H. +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# This program is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with this program. If not, see . + + from typing import List from someipy._internal._sd.entries.sd_entry import SdEntry @@ -11,4 +27,5 @@ def __init__(self): self.timestamp: float = 0.0 self.session_id: int = 0 + self.reboot_flag: bool = False self.entries: List[SdEntry] = [] diff --git a/src/someipy/_internal/_sd/sd_message_creator.py b/src/someipy/_internal/_sd/sd_message_creator.py new file mode 100644 index 0000000..d3c7c0d --- /dev/null +++ b/src/someipy/_internal/_sd/sd_message_creator.py @@ -0,0 +1,156 @@ +# Copyright (C) 2025 Christian H. +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# This program is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with this program. If not, see . + + +from typing import Iterable +from someipy._internal._daemon.offer_service_storage import ServiceToOffer +from someipy._internal._sd.entries.offer_service_entry import OfferServiceEntry +from someipy._internal._sd.entries.stop_offer_service_entry import StopOfferServiceEntry +from someipy._internal._sd.options.endpoint import IpV4EndpointOption +from someipy._internal._sd.sd_message import SdMessage +from someipy._internal.transport_layer_protocol import TransportLayerProtocol + + +def create_offer_service_message( + services_to_offer: Iterable[ServiceToOffer], session_id: int, reboot_flag: bool +) -> SdMessage: + + options = set() + + for service in services_to_offer: + if service.has_udp: + options.add( + IpV4EndpointOption( + address=service.endpoint.ip, + protocol=TransportLayerProtocol.UDP, + port=service.endpoint.port, + ) + ) + + if service.has_tcp: + options.add( + IpV4EndpointOption( + address=service.endpoint.ip, + protocol=TransportLayerProtocol.TCP, + port=service.endpoint.port, + ) + ) + + options = list(options) + + sd_message = SdMessage() + sd_message.session_id = session_id + sd_message.reboot_flag = reboot_flag + + for service in services_to_offer: + endpoints = [] + if service.has_udp: + endpoints.extend( + [ + option + for option in options + if option.protocol == TransportLayerProtocol.UDP + and option.address == service.endpoint.ip + and option.port == service.endpoint.port + ] + ) + if service.has_tcp: + endpoints.extend( + [ + option + for option in options + if option.protocol == TransportLayerProtocol.TCP + and option.address == service.endpoint.ip + and option.port == service.endpoint.port + ] + ) + + new_entry = OfferServiceEntry( + service_id=service.service_id, + instance_id=service.instance_id, + major_version=service.major_version, + minor_version=service.minor_version, + ttl=service.offer_ttl_seconds, + ip_v4_endpoints=endpoints, + ip_v6_endpoints=[], + ) + sd_message.entries.append(new_entry) + return sd_message + + +def create_stop_offer_service_message( + services_to_stop: Iterable[ServiceToOffer], session_id: int, reboot_flag: bool +) -> SdMessage: + + options = set() + + for service in services_to_stop: + if service.has_udp: + options.add( + IpV4EndpointOption( + address=service.endpoint.ip, + protocol=TransportLayerProtocol.UDP, + port=service.endpoint.port, + ) + ) + + if service.has_tcp: + options.add( + IpV4EndpointOption( + address=service.endpoint.ip, + protocol=TransportLayerProtocol.TCP, + port=service.endpoint.port, + ) + ) + + options = list(options) + + sd_message = SdMessage() + sd_message.session_id = session_id + sd_message.reboot_flag = reboot_flag + + for service in services_to_stop: + endpoints = [] + if service.has_udp: + endpoints.extend( + [ + option + for option in options + if option.protocol == TransportLayerProtocol.UDP + and option.address == service.endpoint.ip + and option.port == service.endpoint.port + ] + ) + if service.has_tcp: + endpoints.extend( + [ + option + for option in options + if option.protocol == TransportLayerProtocol.TCP + and option.address == service.endpoint.ip + and option.port == service.endpoint.port + ] + ) + + new_entry = StopOfferServiceEntry( + service_id=service.service_id, + instance_id=service.instance_id, + major_version=service.major_version, + minor_version=service.minor_version, + ip_v4_endpoints=endpoints, + ip_v6_endpoints=[], + ) + sd_message.entries.append(new_entry) + return sd_message diff --git a/src/someipy/_internal/_sd/service_instance.py b/src/someipy/_internal/_sd/service_instance.py index 30e4a7c..0af942d 100644 --- a/src/someipy/_internal/_sd/service_instance.py +++ b/src/someipy/_internal/_sd/service_instance.py @@ -1,3 +1,19 @@ +# Copyright (C) 2025 Christian H. +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# This program is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with this program. If not, see . + + from dataclasses import dataclass from someipy._internal._common.endpoint import Endpoint diff --git a/src/someipy/_internal/someip_endpoint_factory.py b/src/someipy/_internal/someip_endpoint_factory.py index 3e0b211..d16a75a 100644 --- a/src/someipy/_internal/someip_endpoint_factory.py +++ b/src/someipy/_internal/someip_endpoint_factory.py @@ -80,7 +80,7 @@ async def create_client_endpoint( logger: logging.Logger = None, ) -> SomeipEndpoint: if protocol == TransportLayerProtocol.UDP: - udp_endpoint = SomeipEndpointFactory.create_server_endpoint( + udp_endpoint = await SomeipEndpointFactory.create_server_endpoint( src_ip, src_port, TransportLayerProtocol.UDP, @@ -93,3 +93,18 @@ async def create_client_endpoint( ) tcp_endpoint.set_someip_callback(someip_message_callback) return tcp_endpoint + + @staticmethod + def create_tcp_client_endpoint( + dst_ip: str, + dst_port: int, + src_ip: str, + src_port: int, + someip_message_callback: Callable[[SomeIpMessage], None], + logger: logging.Logger = None, + ) -> TCPClientSomeipEndpoint: + tcp_endpoint = TCPClientSomeipEndpoint( + dst_ip, dst_port, src_ip, src_port, logger + ) + tcp_endpoint.set_someip_callback(someip_message_callback) + return tcp_endpoint diff --git a/src/someipy/_internal/someip_sd_builder.py b/src/someipy/_internal/someip_sd_builder.py index 37633ff..4909fc5 100644 --- a/src/someipy/_internal/someip_sd_builder.py +++ b/src/someipy/_internal/someip_sd_builder.py @@ -13,153 +13,18 @@ # You should have received a copy of the GNU General Public License # along with this program. If not, see . -import ipaddress -from typing import Iterable, List, Set, Tuple -from someipy._internal.offer_service_storage import ServiceToOffer from someipy._internal.someip_header import SomeIpHeader -from .transport_layer_protocol import TransportLayerProtocol from .someip_sd_header import ( SD_BYTE_LENGTH_IP4ENDPOINT_OPTION, SD_SINGLE_ENTRY_LENGTH_BYTES, - SdService, SomeIpSdHeader, SdEntry, SdEntryType, SdServiceEntry, - SdOptionCommon, - SD_IPV4ENDPOINT_OPTION_LENGTH_VALUE, - SdOptionType, - SdIPV4EndpointOption, SdEventGroupEntry, ) -def build_offer_service_sd_header( - services_to_offer: Iterable[ServiceToOffer], session_id: int, reboot_flag: bool -) -> SomeIpSdHeader: - # Collect all endpoints and create an option for each unique endpoint - options: List[SdIPV4EndpointOption] = [] - for service in services_to_offer: - - if service.has_tcp: - # Check if endpoint is already contained in options - if any( - option.ipv4_address == ipaddress.IPv4Address(service.endpoint_ip) - and option.port == service.endpoint_port - and option.protocol == TransportLayerProtocol.TCP - for option in options - ): - continue - - option_entry_common = SdOptionCommon( - length=SD_IPV4ENDPOINT_OPTION_LENGTH_VALUE, - type=SdOptionType.IPV4_ENDPOINT, - discardable_flag=False, - ) - sd_option_entry = SdIPV4EndpointOption( - sd_option_common=option_entry_common, - ipv4_address=ipaddress.IPv4Address(service.endpoint_ip), - protocol=TransportLayerProtocol.TCP, - port=service.endpoint_port, - ) - options.append(sd_option_entry) - - if service.has_udp: - # Check if endpoint is already contained in options - if any( - option.ipv4_address == ipaddress.IPv4Address(service.endpoint_ip) - and option.port == service.endpoint_port - and option.protocol == TransportLayerProtocol.UDP - for option in options - ): - continue - - option_entry_common = SdOptionCommon( - length=SD_IPV4ENDPOINT_OPTION_LENGTH_VALUE, - type=SdOptionType.IPV4_ENDPOINT, - discardable_flag=False, - ) - sd_option_entry = SdIPV4EndpointOption( - sd_option_common=option_entry_common, - ipv4_address=ipaddress.IPv4Address(service.endpoint_ip), - protocol=TransportLayerProtocol.UDP, - port=service.endpoint_port, - ) - options.append(sd_option_entry) - - # Build the entries that reference the options - entries = [] - for service in services_to_offer: - num_options_1 = 2 if service.has_tcp and service.has_udp else 1 - for i in range(len(options)): - - # Loop through all options and check if the endpoint matches, then the index of the first option is found - # Depending on whether only one protocol or two are used, the num_options_1 is set to 1 or 2 - if ( - options[i].ipv4_address == ipaddress.IPv4Address(service.endpoint_ip) - and options[i].port == service.endpoint_port - and ( - ( - options[i].protocol == TransportLayerProtocol.UDP - and service.has_udp - ) - or ( - options[i].protocol == TransportLayerProtocol.TCP - and service.has_tcp - ) - ) - ): - sd_entry = SdEntry( - SdEntryType.OFFER_SERVICE, - i, # index_first_option - 0, # index_second_option - num_options_1, # num_options_1 - 0, # num_options_2 - service.service_id, - service.instance_id, - service.major_version, - service.offer_ttl_seconds, - ) - service_entry = SdServiceEntry( - sd_entry=sd_entry, minor_version=service.minor_version - ) - entries.append(service_entry) - - # Pack together all entries and options into a single SD message - - # 20 bytes for header and length values of entries and options - # + length of entries array - # + length of options array - total_length = ( - 20 - + (len(entries) * SD_SINGLE_ENTRY_LENGTH_BYTES) - + (len(options) * SD_BYTE_LENGTH_IP4ENDPOINT_OPTION) - ) - someip_header = SomeIpHeader.generate_sd_header( - length=total_length, session_id=session_id - ) - - return SomeIpSdHeader( - someip_header=someip_header, - reboot_flag=reboot_flag, - unicast_flag=True, - length_entries=(len(entries) * SD_SINGLE_ENTRY_LENGTH_BYTES), - length_options=(len(options) * SD_BYTE_LENGTH_IP4ENDPOINT_OPTION), - service_entries=entries, - options=options, - ) - - -def build_stop_offer_service_sd_header( - services: Iterable[SdService], session_id: int, reboot_flag: bool -) -> SomeIpSdHeader: - header = build_offer_service_sd_header(services, session_id, reboot_flag) - for entry in header.service_entries: - entry.sd_entry.entry_type = SdEntryType.STOP_OFFER_SERVICE - entry.sd_entry.ttl = 0 - return header - - def build_subscribe_eventgroup_ack_entry( service_id: int, instance_id: int, major_version: int, ttl: int, event_group_id: int ) -> SdEventGroupEntry: diff --git a/src/someipy/client_service_instance.py b/src/someipy/client_service_instance.py index efe0a5a..6c7c0f8 100644 --- a/src/someipy/client_service_instance.py +++ b/src/someipy/client_service_instance.py @@ -294,6 +294,9 @@ def someip_message_received( def subscribe_eventgroup( self, eventgroup: EventGroup, ttl_subscription_seconds: int ): + if not isinstance(ttl_subscription_seconds, int): + raise ValueError("ttl_subscription_seconds must be an integer value.") + method_request = create_uds_message( SubscribeEventGroupRequest, service_id=self._service.id, diff --git a/src/someipy/someipyd.py b/src/someipy/someipyd.py index 718fb3f..06dd491 100644 --- a/src/someipy/someipyd.py +++ b/src/someipy/someipyd.py @@ -28,6 +28,7 @@ import time from typing import Any, Dict, List, Set, Tuple, Union +import someipy from someipy._internal._common.endpoint import Endpoint from someipy._internal._daemon.daemon_server_client import ( ClientMessageEventArgs, @@ -46,24 +47,24 @@ ) from someipy._internal._sd.options.endpoint import IpV4EndpointOption from someipy._internal._sd.sd_message import SdMessage +from someipy._internal._sd.sd_message_creator import ( + create_offer_service_message, + create_stop_offer_service_message, +) from someipy._internal._sd.service_instance import ServiceInstance from someipy._internal.message_types import MessageType from someipy._internal.someip_endpoint import ( - SomeipEndpoint, TCPClientSomeipEndpoint, TCPSomeipEndpoint, - UDPSomeipEndpoint, ) from someipy._internal.someip_endpoint_factory import SomeipEndpointFactory from someipy._internal.someip_endpoint_storage import SomeipEndpointStorage from someipy._internal.someip_message import SomeIpMessage -from someipy._internal.tcp_client_manager import TcpClientManager, TcpClientProtocol from someipy._internal.transport_layer_protocol import TransportLayerProtocol from someipy._internal.session_handler import SessionHandler from someipy._internal.simple_timer import SimplePeriodicTimer from someipy._internal.someip_header import SomeIpHeader from someipy._internal.someip_sd_builder import ( - build_stop_offer_service_sd_header, build_subscribe_eventgroup_ack_entry, build_subscribe_eventgroup_ack_sd_header, ) @@ -103,7 +104,10 @@ create_rcv_broadcast_socket, create_udp_socket, ) -from someipy._internal.offer_service_storage import OfferServiceStorage, ServiceToOffer +from someipy._internal._daemon.offer_service_storage import ( + OfferServiceStorage, + ServiceToOffer, +) from someipy.service import Event, Method, EventGroup from someipy._internal._daemon.daemon_server import ( ClientConnectedEventArgs, @@ -114,7 +118,7 @@ DEFAULT_SOCKET_PATH = "/tmp/someipyd.sock" DEFAULT_CONFIG_FILE = "someipyd.json" DEFAULT_SD_ADDRESS = "224.224.224.245" -DEFAULT_INTERFACE_IP = "127.0.0.2" +DEFAULT_INTERFACE_IP = "127.0.0.1" DEFAULT_SD_PORT = 30490 DEFAULT_TCP_PORT = 30500 @@ -214,12 +218,15 @@ async def client_disconnected( ): self.logger.info(f"Client disconnected: {event_args.client.id}") - writer_id = event_args.client.id + client_id = event_args.client.id + + event_args.client.message_received -= self.handle_client_message + # Remove all subscriptions for the client - self._requested_subscriptions.remove_client(writer_id) + self._requested_subscriptions.remove_client(client_id) # Clean up the transmission task for the client. This will also clean up the transmission queue - tx_task = self._tx_tasks.get(writer_id) + tx_task = self._tx_tasks.pop(client_id, None) if tx_task and not tx_task.cancelled(): tx_task.cancel() try: @@ -227,17 +234,17 @@ async def client_disconnected( except asyncio.CancelledError: pass - self._services_to_offer.remove_client(writer_id) + self._services_to_offer.remove_client(client_id) self._cleanup_unused_timers() - client_endpoints = self._someip_server_endpoints.get_endpoints(writer_id) + client_endpoints = self._someip_server_endpoints.get_endpoints(client_id) if client_endpoints is not None: for endpoint in client_endpoints: self.logger.debug( - f"Closing endpoint {endpoint.dst_ip()}:{endpoint.dst_port()} for client {writer_id}" + f"Closing endpoint {endpoint.dst_ip()}:{endpoint.dst_port()} for client {client_id}" ) endpoint.shutdown() - self._someip_server_endpoints.remove_endpoint(writer_id, endpoint) + self._someip_server_endpoints.remove_endpoint(client_id, endpoint) self.logger.debug(f"Client disconnected") @@ -245,33 +252,7 @@ async def _check_services_ttl_task(self): try: while True: await asyncio.sleep(0.1) - - self._cleanup_obsolete_pending_subscriptions() - self._cleanup_active_subscriptions() - - count_before = len(self._found_services) - - current_time = time.time() - - # self.logger.debug(f"Current time: {current_time}, checking services...") - - # for service in self._found_services: - # self.logger.debug( - # f"Checking service {service.service_id}, timestamp: {service.timestamp}, ttl: {service.ttl}" - # ) - - # Process timeouts and filter services in one operation - self._found_services = [ - service - for service in self._found_services - if not (current_time - service.timestamp > service.service.ttl) - ] - - count_after = len(self._found_services) - if count_before != count_after: - self.logger.info( - f"Removed {count_before - count_after} timed out services. Remaining: {count_after}" - ) + self._check_service_ttl_task_impl() except asyncio.CancelledError: # Task was cancelled - exit cleanly @@ -280,6 +261,28 @@ async def _check_services_ttl_task(self): self.logger.error(f"Error in service TTL checker task: {e}") pass + def _check_service_ttl_task_impl(self): + + self._cleanup_obsolete_pending_subscriptions() + self._cleanup_active_subscriptions() + + count_before = len(self._found_services) + + current_time = time.time() + + # Process timeouts and filter services in one operation + self._found_services = [ + service + for service in self._found_services + if not (current_time - service.timestamp > service.ttl) + ] + + count_after = len(self._found_services) + if count_before != count_after: + self.logger.info( + f"Removed {count_before - count_after} timed out services. Remaining: {count_after}" + ) + def _someip_message_callback( self, message: SomeIpMessage, @@ -300,8 +303,8 @@ def _someip_message_callback( if ( service.service_id == service_id - and service.endpoint_ip == dst_addr[0] - and service.endpoint_port == dst_addr[1] + and str(service.endpoint.ip) == dst_addr[0] + and service.endpoint.port == dst_addr[1] ): self.logger.debug(f"Found matching service {service.service_id}") for method in service.methods: @@ -420,14 +423,14 @@ def _someip_message_callback( == active_subscription.client_endpoint.port ): - writer_ids = ( + client_ids = ( self._requested_subscriptions.get_client_ids( requested_subscription ) ) - for writer_id in writer_ids: - tx_queue = self._tx_queues.get(writer_id) + for client_id in client_ids: + tx_queue = self._tx_queues.get(client_id) if tx_queue: tx_queue.put_nowait( self.prepare_message(event_msg) @@ -481,7 +484,7 @@ def _close_unused_endpoints(self): endpoint_used = False for service in self._services_to_offer.services: if ( - service.endpoint_ip == endpoint.ip() + service.endpoint.ip == endpoint.ip() and service.endpoint_port == endpoint.port() ): endpoint_used = True @@ -500,8 +503,7 @@ async def tx_task(self, client: DaemonServerClient): try: # Send the data - client.writer.write(data) - await client.writer.drain() + await client.send(data) tx_queue.task_done() except ConnectionError as e: self.logger.error(f"Error sending data in tx task: {e}") @@ -515,8 +517,7 @@ async def tx_task(self, client: DaemonServerClient): self.logger.debug(f"TX task for writer {client.id} cancelled") # Perform cleanup here try: - client.writer.close() - await client.writer.wait_closed() + await client.close() except Exception as e: self.logger.error(f"Error closing writer: {e}") finally: @@ -527,7 +528,8 @@ async def tx_task(self, client: DaemonServerClient): async def handle_client_message( self, sender: object, event_args: ClientMessageEventArgs ): - writer_id = id(event_args.client.id) + client_id = event_args.client.id + message = event_args.message message_type = event_args.message.get("type") self.logger.debug(f"Received message type: {message_type}") @@ -547,10 +549,10 @@ async def handle_client_message( handler = message_handlers[message_type] if asyncio.iscoroutinefunction(handler): - await handler(message, writer_id) + await handler(message, client_id) return else: - handler(message, writer_id) + handler(message, client_id) return else: self.logger.warning( @@ -588,12 +590,14 @@ async def _handle_subscribe_eventgroup_request( event_group = EventGroup.from_json(message["eventgroup"]) + client_address = ipaddress.IPv4Address(message["client_endpoint_ip"]) + new_subscription = Subscription( service_id=message["service_id"], instance_id=message["instance_id"], major_version=message["major_version"], client_endpoint=Endpoint( - ip=message["client_endpoint_ip"], port=message["client_endpoint_port"] + ip=client_address, port=message["client_endpoint_port"] ), server_endpoint=None, protocols=frozenset(protocols), @@ -601,6 +605,7 @@ async def _handle_subscribe_eventgroup_request( ttl_seconds=message["ttl_subscription"], ) + self.logger.debug(f"Add subscription to storage with id {client_id}") self._requested_subscriptions.add_subscription(client_id, new_subscription) def _handle_stop_subscribe_eventgroup_request( @@ -667,17 +672,17 @@ async def _handle_offer_service_request( if service_to_add.has_tcp: if not self._someip_server_endpoints.has_endpoint( - service_to_add.endpoint_ip, - service_to_add.endpoint_port, + str(service_to_add.endpoint.ip), + service_to_add.endpoint.port, TransportLayerProtocol.TCP, ): self.logger.debug( - f"Creating new TCP endpoint for {service_to_add.endpoint_ip}:{service_to_add.endpoint_port}" + f"Creating new TCP endpoint for {service_to_add.endpoint}" ) tcp_endpoint = await self._endpoint_factory.create_server_endpoint( - service_to_add.endpoint_ip, - service_to_add.endpoint_port, + str(service_to_add.endpoint.ip), + service_to_add.endpoint.port, TransportLayerProtocol.TCP, self._someip_message_callback, ) @@ -712,8 +717,9 @@ def _handle_stop_offer_service_request( minor_version=message["minor_version"], offer_ttl_seconds=message["ttl"], cyclic_offer_delay_ms=message["cyclic_offer_delay_ms"], - endpoint_ip=message["endpoint_ip"], - endpoint_port=message["endpoint_port"], + endpoint=Endpoint( + ipaddress.IPv4Address(message["endpoint_ip"]), message["endpoint_port"] + ), methods=methods, eventgroups=eventgroups, ) @@ -727,15 +733,19 @@ def _handle_stop_offer_service_request( reboot_flag, ) = self._mcast_session_handler.update_session() - sd_header = build_stop_offer_service_sd_header( - [service_to_stop], session_id, reboot_flag + sd_message = create_stop_offer_service_message( + services_to_stop=[service_to_stop], + session_id=session_id, + reboot_flag=reboot_flag, ) - buffer = sd_header.to_buffer() + if self._ucast_transport: self.logger.debug( f"Send stop offer message for service 0x{service_to_stop.service_id:04x}, instance 0x{service_to_stop.instance_id:04x} to {self.sd_address}:{self.sd_port}" ) - self._ucast_transport.sendto(buffer, (self.sd_address, self.sd_port)) + self._ucast_transport.sendto( + serialize_sd_message(sd_message), (self.sd_address, self.sd_port) + ) if service_to_stop.has_udp: try: @@ -744,7 +754,7 @@ def _handle_stop_offer_service_request( ) if udp_endpoint: self.logger.debug( - f"Closing UDP endpoint for {service_to_stop.endpoint_ip}:{service_to_stop.endpoint_port}" + f"Closing UDP endpoint for {service_to_stop.endpoint}" ) udp_endpoint.shutdown() self._someip_server_endpoints.remove_endpoint( @@ -752,7 +762,7 @@ def _handle_stop_offer_service_request( ) except Exception as e: self.logger.error( - f"Error closing UDP endpoint for {service_to_stop.endpoint_ip}:{service_to_stop.endpoint_port}: {e}" + f"Error closing UDP endpoint for {service_to_stop.endpoint}: {e}" ) if service_to_stop.has_tcp: @@ -762,7 +772,7 @@ def _handle_stop_offer_service_request( ) if tcp_endpoint: self.logger.debug( - f"Closing TCP endpoint for {service_to_stop.endpoint_ip}:{service_to_stop.endpoint_port}" + f"Closing TCP endpoint for {service_to_stop.endpoint}" ) tcp_endpoint.shutdown() self._someip_server_endpoints.remove_endpoint( @@ -770,7 +780,7 @@ def _handle_stop_offer_service_request( ) except Exception as e: self.logger.error( - f"Error closing TCP endpoint for {service_to_stop.endpoint_ip}:{service_to_stop.endpoint_port}: {e}" + f"Error closing TCP endpoint for {service_to_stop.endpoint}: {e}" ) def _handle_inbound_call_method_response( @@ -850,12 +860,11 @@ async def _handle_outbound_call_method_request( f"Creating new TCP endpoint for {message['src_endpoint_ip']}:{message['src_endpoint_port']}" ) - tcp_endpoint = self._endpoint_factory.create_client_endpoint( + tcp_endpoint = self._endpoint_factory.create_tcp_client_endpoint( message["dst_endpoint_ip"], message["dst_endpoint_port"], message["src_endpoint_ip"], message["src_endpoint_port"], - TransportLayerProtocol.TCP, self._someip_message_callback, self.logger, ) @@ -934,20 +943,17 @@ def _handle_find_service_request(self, message: FindServiceRequest, writer_id: i if service.has_tcp: protocols_to_add.add(TransportLayerProtocol.TCP) - service_to_add = SdService2( + service_to_add = ServiceInstance( service_id=service.service_id, instance_id=service.instance_id, major_version=service.major_version, minor_version=service.minor_version, - ttl=0, - endpoint=( - ipaddress.IPv4Address(service.endpoint_ip), - service.endpoint_port, - ), + ttl=service.offer_ttl_seconds, + endpoint=service.endpoint, protocols=frozenset(protocols_to_add), + timestamp=0.0, ) - service_to_add = SdServiceWithTimestamp(service_to_add, 0.0) all_services.append(service_to_add) for found_service in all_services: @@ -964,17 +970,17 @@ def _handle_find_service_request(self, message: FindServiceRequest, writer_id: i """ if ( - (message["service_id"] == found_service.service.service_id) + (message["service_id"] == found_service.service_id) and ( - message["instance_id"] == found_service.service.instance_id + message["instance_id"] == found_service.instance_id or message["instance_id"] == 0xFFFF ) and ( - message["major_version"] == found_service.service.major_version + message["major_version"] == found_service.major_version or message["major_version"] == 0xFF ) and ( - message["minor_version"] == found_service.service.minor_version + message["minor_version"] == found_service.minor_version or message["minor_version"] == 0xFFFFFFFF ) ): @@ -982,12 +988,12 @@ def _handle_find_service_request(self, message: FindServiceRequest, writer_id: i response = create_uds_message( FindServiceResponse, success=True, - service_id=found_service.service.service_id, - instance_id=found_service.service.instance_id, - major_version=found_service.service.major_version, - minor_version=found_service.service.minor_version, - endpoint_ip=str(found_service.service.endpoint[0]), - endpoint_port=found_service.service.endpoint[1], + service_id=found_service.service_id, + instance_id=found_service.instance_id, + major_version=found_service.major_version, + minor_version=found_service.minor_version, + endpoint_ip=str(found_service.endpoint.ip), + endpoint_port=found_service.endpoint.port, ) tx_queue = self._tx_queues.get(writer_id) @@ -1142,64 +1148,11 @@ def offer_timer_callback(self, cyclic_offer_delay_ms: int): reboot_flag, ) = self._mcast_session_handler.update_session() - options = set() - for service in services_to_offer: - if service.has_udp: - options.add( - IpV4EndpointOption( - address=service.endpoint.ip, - protocol=TransportLayerProtocol.UDP, - port=service.endpoint.port, - ) - ) - if service.has_tcp: - options.add( - IpV4EndpointOption( - address=service.endpoint.ip, - protocol=TransportLayerProtocol.TCP, - port=service.endpoint.port, - ) - ) - - options = list(options) - sd_message = SdMessage() - sd_message.session_id = session_id - sd_message.reboot_flag = reboot_flag - - for service in services_to_offer: - - endpoints = [] - if service.has_udp: - endpoints.extend( - [ - option - for option in options - if option.protocol == TransportLayerProtocol.UDP - and option.address == service.endpoint.ip - and option.port == service.endpoint.port - ] - ) - if service.has_tcp: - endpoints.extend( - [ - option - for option in options - if option.protocol == TransportLayerProtocol.TCP - and option.address == service.endpoint.ip - and option.port == service.endpoint.port - ] - ) - - new_entry = OfferServiceEntry( - service_id=service.service_id, - instance_id=service.instance_id, - major_version=service.major_version, - minor_version=service.minor_version, - ttl=service.offer_ttl_seconds, - ip_v4_endpoints=endpoints, - ip_v6_endpoints=[], - ) - sd_message.entries.append(new_entry) + sd_message = create_offer_service_message( + services_to_offer=services_to_offer, + session_id=session_id, + reboot_flag=reboot_flag, + ) if self._ucast_transport: self._ucast_transport.sendto( @@ -1298,19 +1251,18 @@ def _handle_offered_service(self, offered_service: ServiceInstance): if not self._someip_client_endpoints.has_tcp_endpoint( str(requested_subscription[0].client_endpoint.ip), requested_subscription[0].client_endpoint.port, - str(offered_service.endpoint[0]), - offered_service.endpoint[1], + str(offered_service.endpoint.ip), + offered_service.endpoint.port, ): self.logger.debug( f"Creating new TCP endpoint for {requested_subscription[0].client_endpoint}" ) - tcp_endpoint = self._endpoint_factory.create_client_endpoint( + tcp_endpoint = self._endpoint_factory.create_tcp_client_endpoint( str(offered_service.endpoint.ip), offered_service.endpoint.port, str(requested_subscription[0].client_endpoint.ip), requested_subscription[0].client_endpoint.port, - TransportLayerProtocol.TCP, self._someip_message_callback, self.logger, ) @@ -1332,6 +1284,7 @@ def _handle_offered_service(self, offered_service: ServiceInstance): # Build subscribe message sd_message = SdMessage() sd_message.session_id = session_id + sd_message.reboot_flag = reboot_flag options = [] for protocol in requested_protocols: @@ -1520,20 +1473,42 @@ def datagram_received_mcast( if addr[0] == self.interface and addr[1] == self.sd_port: return - if is_sd_message(data) is False: + if not is_sd_message(data): return - sd_message = deserialize_sd_message(data) + sd_message = deserialize_sd_message(data, addr[0], addr[1], multicast=True) sd_message.timestamp = time.time() - # someip_header = SomeIpHeader.from_buffer(data) - # if not someip_header.is_sd_header(): - # return - for offer_service_entry in [ - o for o in sd_message.entries if o.entry_type == SdEntryType.OFFER_SERVICE + o + for o in sd_message.entries + if o.type + == someipy._internal._sd.entries.sd_entry.SdEntryType.OFFER_SERVICE ]: - self._handle_offered_service(offer_service_entry, sd_message.timestamp) + entry: OfferServiceEntry = offer_service_entry + + protocols = set() + for ep in entry.ip_v4_endpoints: + protocols.add(ep.protocol) + for ep in entry.ip_v6_endpoints: + protocols.add(ep.protocol) + + endpoint = Endpoint( + ip=entry.ip_v4_endpoints[0].address, + port=entry.ip_v4_endpoints[0].port, + ) + + service_instance = ServiceInstance( + service_id=entry.service_id, + instance_id=entry.instance_id, + major_version=entry.major_version, + minor_version=entry.minor_version, + ttl=entry.ttl, + endpoint=endpoint, + protocols=frozenset(protocols), + timestamp=sd_message.timestamp, + ) + self._handle_offered_service(service_instance) someip_sd_header = SomeIpSdHeader.from_buffer(data) diff --git a/tests/sd/test_sd_deserialization.py b/tests/sd/test_sd_deserialization.py index c8bd07d..be8e3fc 100644 --- a/tests/sd/test_sd_deserialization.py +++ b/tests/sd/test_sd_deserialization.py @@ -13,7 +13,9 @@ deserialize_ipv6_multicast_option, deserialize_ipv6_sd_endpoint_option, deserialize_load_balancing_option, + deserialize_sd_message, ) +from someipy._internal._sd.entries.sd_entry import SdEntryType from someipy._internal._sd.options.endpoint import ( IpV4EndpointOption, IpV6EndpointOption, @@ -183,3 +185,71 @@ def test_deserialize_load_balancing_option(): assert isinstance(result, LoadBalancingOption) assert result.priority == 0x0102 assert result.weight == 0x0304 + + +@pytest.fixture +def sd_message_no_entries_no_options() -> bytes: + # fmt: off + data = bytes([ + 0xFF, 0xFF, 0x81, 0x00, + 0x00, 0x00, 0x00, 20, # length + 0x00, 0x00, 0x00, 0x01, # client id, session id + 0x01, 0x01, 0x02, 0x00, + 0x80, 0x00, 0x00, 0x00, # flags 8 bit, reserved 24 bit + 0x00, 0x00, 0x00, 0x00, # entries length + 0x00, 0x00, 0x00, 0x00, # options length + ]) + # fmt: on + return data + + +@pytest.fixture +def sd_message_no_with_stop_offer_service_entry() -> bytes: + # fmt: off + data = bytes([ + 0xFF, 0xFF, 0x81, 0x00, + 0x00, 0x00, 0x00, 20, # length + 0x00, 0x00, 0x00, 0x01, # client id, session id + 0x01, 0x01, 0x02, 0x00, + 0x80, 0x00, 0x00, 0x00, # flags 8 bit, reserved 24 bit + 0x00, 0x00, 0x00, 16, # entries length + 0x01, 0x00, 0x00, 0x01, + 0x00, 0x01, 0x00, 0x02, # service id, instance id + 0x03, 0x00, 0x00, 0x00, # major version, ttl + 0x00, 0x00, 0x00, 0x05, # minor version + 0x00, 0x00, 0x00, 12, # options length + 0x00, 0x09, 0x04, 0x00, # common option data + 192, 168, 1, 10, # ipv4 address + 0x00, 0x11, 0x10, 0x01 # reserved, proto, port + ]) + # fmt: on + return data + + +def test_deserialize_sd_message_reboot_flag(sd_message_no_entries_no_options): + sd_message = deserialize_sd_message(sd_message_no_entries_no_options, "", 0, True) + assert sd_message.reboot_flag == True + + +def test_deserialize_sd_message_with_stop_offer_service_entry( + sd_message_no_with_stop_offer_service_entry, +): + sd_message = deserialize_sd_message( + sd_message_no_with_stop_offer_service_entry, "", 0, True + ) + + assert len(sd_message.entries) == 1 + entry = sd_message.entries[0] + + assert entry.type == SdEntryType.STOP_OFFER_SERVICE + + assert entry.service_id == 0x0001 + assert entry.instance_id == 0x0002 + assert entry.major_version == 0x03 + assert entry.minor_version == 0x05 + assert len(entry.ip_v4_endpoints) == 1 + endpoint = entry.ip_v4_endpoints[0] + assert len(entry.ip_v6_endpoints) == 0 + assert str(endpoint.address) == "192.168.1.10" + assert endpoint.protocol == TransportLayerProtocol.UDP + assert endpoint.port == 0x1001 diff --git a/tests/test_someipyd.py b/tests/test_someipyd.py index c71519d..9aa01bd 100644 --- a/tests/test_someipyd.py +++ b/tests/test_someipyd.py @@ -1,14 +1,24 @@ from asyncio import DatagramTransport +import asyncio +import ipaddress import logging +import time import pytest import pytest_asyncio -from unittest.mock import Mock +from unittest.mock import MagicMock, Mock, patch +import someipy from someipy._internal._common.endpoint import Endpoint +from someipy._internal._daemon.daemon_server import ClientConnectedEventArgs +from someipy._internal._daemon.daemon_server_client import DaemonServerClient from someipy._internal._daemon.uds_messages import ( OfferServiceRequest, + StopOfferServiceRequest, SubscribeEventGroupRequest, create_uds_message, ) +from someipy._internal._sd.deserialization.sd_serialization import serialize_sd_message +from someipy._internal._sd.entries.offer_service_entry import OfferServiceEntry +from someipy._internal._sd.options.endpoint import IpV4EndpointOption from someipy._internal._sd.service_instance import ServiceInstance from someipy._internal.someip_endpoint_factory import SomeipEndpointFactory from someipy._internal.transport_layer_protocol import TransportLayerProtocol @@ -63,7 +73,7 @@ def eventgroup() -> EventGroup: id=1, events=[ Event(id=1, protocol=TransportLayerProtocol.UDP), - Event(id=2, protocol=TransportLayerProtocol.UDP), + Event(id=2, protocol=TransportLayerProtocol.TCP), ], ) @@ -139,6 +149,23 @@ def offer_service_request(eventgroup, method) -> OfferServiceRequest: ) +@pytest.fixture +def stop_offer_service_request(eventgroup, method) -> StopOfferServiceRequest: + return create_uds_message( + StopOfferServiceRequest, + service_id=1, + instance_id=2, + major_version=3, + minor_version=0, + endpoint_ip="127.0.0.1", + endpoint_port=1, + ttl=5, + eventgroup_list=[eventgroup.to_json()], + method_list=[method.to_json()], + cyclic_offer_delay_ms=1000, + ) + + def test_handle_offered_service_adds_service( daemon: SomeipDaemon, service_instance: ServiceInstance ): @@ -243,3 +270,151 @@ async def test_handle_offer_service_request_opens_server_endpoint( await daemon._handle_offer_service_request(offer_service_request, 1) assert len(daemon._someip_server_endpoints) == 1 + + +@pytest.mark.asyncio +async def test_handle_offer_service_does_not_add_service_twice( + daemon: SomeipDaemon, + offer_service_request: OfferServiceRequest, +): + assert len(daemon._services_to_offer) == 0 + + await daemon._handle_offer_service_request(offer_service_request, 1) + await daemon._handle_offer_service_request(offer_service_request, 1) + + assert len(daemon._services_to_offer) == 1 + + +@pytest.mark.asyncio +async def test_handle_stop_offer_service_removes_service_to_offer( + daemon: SomeipDaemon, + offer_service_request: OfferServiceRequest, + stop_offer_service_request: StopOfferServiceRequest, +): + assert len(daemon._services_to_offer) == 0 + + await daemon._handle_offer_service_request(offer_service_request, 1) + assert len(daemon._services_to_offer) == 1 + + daemon._handle_stop_offer_service_request(stop_offer_service_request, 1) + + assert len(daemon._services_to_offer) == 0 + + +@pytest.mark.asyncio +async def test_client_connected_disconnected( + daemon: SomeipDaemon, +): + initial_queue_count = len(daemon._tx_queues) + + # Simulate a new client connection + client_id = 42 + + new_client = Mock(spec=DaemonServerClient) + new_client.id = client_id + new_client.message_received = someipy._internal._common.event.Event() + new_client.message_received.add_handler = Mock() + new_client.message_received.remove_handler = Mock() + + new_client_args = ClientConnectedEventArgs(new_client) + + await daemon.new_client_connected(new_client, new_client_args) + + # Verify that a new queue has been created for the client + assert client_id in daemon._tx_queues.keys() + assert daemon._tx_queues[client_id] is not None + + new_client.message_received.add_handler.assert_called_once() + + await asyncio.sleep(0.0) # Allow tx task to be scheduled + + await daemon.client_disconnected(daemon, new_client_args) + + # Verify that the client's queue and tx task was removed + assert client_id not in daemon._tx_queues.keys() + assert client_id not in daemon._tx_tasks.keys() + + # Message received event handler was removed + new_client.message_received.remove_handler.assert_called_once() + + """ Cleanup of: + - Offer timers + - Services to be offered + - Subscriptions (pending) + - Subscriptions (active) + - Method calls pending + - Pending find calls + - Close server or client endpoints that are only used by the disconnected client + """ + + +def test_datagram_received_mcast_calls_handle_offered_service( + daemon: SomeipDaemon, +): + + # Create an SdMessage with an OfferService entry + sd_message = someipy._internal._sd.sd_message.SdMessage() + sd_message.session_id = 1 + + ip_endpoint_option_1 = IpV4EndpointOption( + address=ipaddress.IPv4Address("192.168.1.1"), + protocol=TransportLayerProtocol.TCP, + port=8080, + ) + ip_endpoint_option_2 = IpV4EndpointOption( + address=ipaddress.IPv4Address("192.168.1.2"), + protocol=TransportLayerProtocol.UDP, + port=8080, + ) + + offer_service_entry = OfferServiceEntry( + service_id=1, + instance_id=1, + major_version=1, + minor_version=0, + ttl=120, + ip_v4_endpoints=[ip_endpoint_option_1], + ip_v6_endpoints=[], + ) + + sd_message.entries.append(offer_service_entry) + data = serialize_sd_message(sd_message) + + # Patch the _handle_offered_service method to monitor its calls + with patch.object( + daemon, "_handle_offered_service", wraps=daemon._handle_offered_service + ) as mock_handle_offered_service: + # Simulate receiving a multicast datagram + daemon.datagram_received_mcast(data, ("127.0.0.1", 5000)) + + mock_handle_offered_service.assert_called_once() + + +def test_check_services_ttl_task_removes_expired_offered_services( + daemon: SomeipDaemon, service_instance: ServiceInstance +): + service_instance.timestamp = time.time() + service_instance.ttl = 10 # seconds + + # Add a service instance with a short TTL + daemon._found_services.append(service_instance) + + # Run the TTL check task once + daemon._check_service_ttl_task_impl() + + # Service is still valid + assert len(daemon._found_services) == 1 + + service_instance.timestamp -= service_instance.ttl + 1 # Simulate time passage + daemon._check_service_ttl_task_impl() + + # Service should be removed due to TTL expiry + assert len(daemon._found_services) == 0 + + +def test_find_service_request_sends_negative_response(daemon: SomeipDaemon): + pass + + +def test_find_service_request_sends_positive_response(daemon: SomeipDaemon): + pass From 45bd4d721476a2a366d95f3ac40e71c0b2f5cbee Mon Sep 17 00:00:00 2001 From: Christian Date: Wed, 7 Jan 2026 21:01:05 +0100 Subject: [PATCH 6/6] Support stop subscribe functionality in someipy daemon --- README.md | 5 +- docs/changelog.rst | 6 + docs/getting_started.rst | 8 + docs/index.rst | 2 +- docs/someipy_daemon.rst | 9 +- example_apps/call_method_tcp.py | 5 + example_apps/call_method_udp.py | 5 + example_apps/offer_method_tcp.py | 5 + example_apps/offer_method_udp.py | 5 + example_apps/offer_multiple_services.py | 5 + example_apps/receive_events_tcp.py | 8 +- example_apps/receive_events_udp.py | 8 +- example_apps/send_events_tcp.py | 5 + example_apps/send_events_udp.py | 5 + integration_tests/automated_tests/run_all.py | 1 - setup.cfg | 2 +- src/someipy/_internal/_common/endpoint.py | 2 +- .../_internal/_daemon/daemon_server.py | 19 +- .../_internal/_daemon/daemon_server_client.py | 15 ++ .../_daemon/offer_service_storage.py | 2 +- .../_daemon/someipy_daemon_client.py | 6 +- src/someipy/_internal/_daemon/subscription.py | 4 +- src/someipy/_internal/_daemon/uds_messages.py | 4 +- .../_sd/deserialization/sd_deserialization.py | 2 +- .../_sd/deserialization/sd_serialization.py | 2 +- .../_sd/entries/find_service_entry.py | 2 +- .../_sd/entries/offer_service_entry.py | 2 +- src/someipy/_internal/_sd/entries/sd_entry.py | 2 +- .../_sd/entries/stop_offer_service_entry.py | 2 +- .../stop_subscribe_eventgroup_entry.py | 5 +- .../_sd/entries/subscribe_ack_entry.py | 4 +- .../_sd/entries/subscribe_eventgroup_entry.py | 5 +- .../_sd/entries/subscribe_eventgroup_nack.py | 4 +- .../_sd/options/configuration_option.py | 2 +- src/someipy/_internal/_sd/options/endpoint.py | 2 +- .../_internal/_sd/options/load_balancing.py | 2 +- .../_internal/_sd/options/multicast.py | 2 +- .../_internal/_sd/options/sd_endpoint.py | 2 +- src/someipy/_internal/daemon_client_abcs.py | 2 +- .../_internal/someip_endpoint_factory.py | 58 ++--- src/someipy/client_service_instance.py | 6 +- src/someipy/someipyd.json | 5 +- src/someipy/someipyd.py | 238 +++++++++++++----- tests/{ => daemon}/test_someipyd.py | 122 ++++++++- tests/{ => sd}/test_sd_service_instance.py | 0 45 files changed, 457 insertions(+), 150 deletions(-) rename tests/{ => daemon}/test_someipyd.py (76%) rename tests/{ => sd}/test_sd_service_instance.py (100%) diff --git a/README.md b/README.md index d49b251..5c23183 100644 --- a/README.md +++ b/README.md @@ -1,7 +1,6 @@ # someipy - A Python Library implementing the SOME/IP Protocol -![Dynamic JSON Badge](https://img.shields.io/badge/dynamic/json?url=https%3A%2F%2Fgist.githubusercontent.com%2Fchrizog%2F6a2e6f355eedf38ae3af74dcdb7b30a1%2Fraw%2F61630b4d2cc9c04aae78b3cca0cbc9c56e76f0e5%2Fcoverage_someipy_v1.json&query=totals.percent_covered_display&suffix=%20%25&label=Coverage) - +![Dynamic JSON Badge](https://img.shields.io/badge/dynamic/json?url=https%3A%2F%2Fgist.githubusercontent.com%2Fchrizog%2F6a2e6f355eedf38ae3af74dcdb7b30a1%2Fraw%2F61630b4d2cc9c04aae78b3cca0cbc9c56e76f0e5%2Fcoverage_someipy_v1.json&query=totals.percent_covered_display&suffix=%20%25&label=Line%20%Coverage) ## Get in Contact :postbox: @@ -27,7 +26,7 @@ someipy is based on the specification version of R22-11: - [SOME/IP Protocol Specification](https://www.autosar.org/fileadmin/standards/R22-11/FO/AUTOSAR_PRS_SOMEIPProtocol.pdf) - [SOME/IP Service Discovery Protocol Specification](https://www.autosar.org/fileadmin/standards/R22-11/FO/AUTOSAR_PRS_SOMEIPServiceDiscoveryProtocol.pdf) -The library is currently developed and tested under Ubuntu 22.04 and Python 3.8. +The library is currently developed and tested under Ubuntu 22.04 and Python 3.8. Windows is supported as well. ## Typical Use Cases diff --git a/docs/changelog.rst b/docs/changelog.rst index 1b62d0e..ad677d0 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -8,6 +8,12 @@ HEAD Code changes to ``master`` that are *not* in the latest release: +Release v2.1.0 +----------------- +- Support Windows: Use tcp socket for daemon communication on Windows systems. +- Add support for broadcasting SOME/IP SD messages over UDP. +- Implement stop subscribe functionality in the `someipyd` daemon. + Release v2.0.0 ----------------- - Introduced a new architecture with a `someipyd` daemon for centralized network handling. diff --git a/docs/getting_started.rst b/docs/getting_started.rst index 7c7f064..7490d4c 100644 --- a/docs/getting_started.rst +++ b/docs/getting_started.rst @@ -113,6 +113,14 @@ The next step is to connect to the someipy daemon. The daemon is a separate proc In case, a non-default Unix Domain Socket path is used, a config dictionary can be passed to the *connect_to_someipy_daemon* function. +If Windows is used, TCP sockets have to be used for communication between the application and the someipy daemon. In this case, the *connect_to_someipy_daemon* function has to be called with a config dictionary containing the keys *use_tcp*, *tcp_host* and optionally *tcp_port*. + +.. code-block:: python + + someipy_daemon = await connect_to_someipy_daemon( + {"use_tcp": True, "tcp_host": interface_ip, "tcp_port": 30500} + ) + Defining the SOME/IP Service ---------------------------- diff --git a/docs/index.rst b/docs/index.rst index 03057ff..ff476c3 100644 --- a/docs/index.rst +++ b/docs/index.rst @@ -15,7 +15,7 @@ someipy is based on the specification version of R22-11: - `SOME/IP Protocol Specification `_ - `SOME/IP Service Discovery Protocol Specification `_ -The library is developed and tested on Ubuntu 22.04 and using Python 3.8. +The library is developed and tested on Ubuntu 22.04 and using Python 3.8. Windows is supported as well. For Inquiries ------------- diff --git a/docs/someipy_daemon.rst b/docs/someipy_daemon.rst index 1063a2f..3041939 100644 --- a/docs/someipy_daemon.rst +++ b/docs/someipy_daemon.rst @@ -29,11 +29,16 @@ The configuration file is a JSON file that allows you to customize the daemon's "sd_port": 30490, "log_level": "DEBUG", "interface": "127.0.0.2", - "log_path": "/var/log/someipy.log" - } + "log_path": "/var/log/someipy.log", + "use_tcp": false, + "tcp_host": "127.0.0.1", + "tcp_port": 30500 - ``sd_address``: The multicast address for Service Discovery. - ``sd_port``: The port for Service Discovery. - ``log_level``: The logging level (e.g., DEBUG, INFO, WARNING, ERROR). - ``interface``: The network interface to bind to. - ``log_path``: The path to the log file. +- ``use_tcp``: Whether to use TCP sockets instead of UDS sockets for communication between the daemon and clients. If Windows is used `use_tcp` must be set to true. +- ``tcp_host``: The host address for TCP communication (only relevant if `use_tcp` is true). +- ``tcp_port``: The port for TCP communication (only relevant if `use_tcp` is true). \ No newline at end of file diff --git a/example_apps/call_method_tcp.py b/example_apps/call_method_tcp.py index 5e32db5..a7ac6ef 100644 --- a/example_apps/call_method_tcp.py +++ b/example_apps/call_method_tcp.py @@ -38,6 +38,11 @@ async def main(): someipy_daemon = await connect_to_someipy_daemon() + # For Windows use: + # someipy_daemon = await connect_to_someipy_daemon( + # {"use_tcp": True, "tcp_host": interface_ip, "tcp_port": 30500} + # ) + addition_method = Method( id=SAMPLE_METHOD_ID, protocol=TransportLayerProtocol.TCP, diff --git a/example_apps/call_method_udp.py b/example_apps/call_method_udp.py index 195f9ce..d546429 100644 --- a/example_apps/call_method_udp.py +++ b/example_apps/call_method_udp.py @@ -38,6 +38,11 @@ async def main(): someipy_daemon = await connect_to_someipy_daemon() + # For Windows use: + # someipy_daemon = await connect_to_someipy_daemon( + # {"use_tcp": True, "tcp_host": interface_ip, "tcp_port": 30500} + # ) + addition_method = Method( id=SAMPLE_METHOD_ID, protocol=TransportLayerProtocol.UDP, diff --git a/example_apps/offer_method_tcp.py b/example_apps/offer_method_tcp.py index 83821cc..8d34957 100644 --- a/example_apps/offer_method_tcp.py +++ b/example_apps/offer_method_tcp.py @@ -70,6 +70,11 @@ async def main(): someipy_daemon = await connect_to_someipy_daemon() + # For Windows use: + # someipy_daemon = await connect_to_someipy_daemon( + # {"use_tcp": True, "tcp_host": interface_ip, "tcp_port": 30500} + # ) + addition_method = Method( id=SAMPLE_METHOD_ID, protocol=TransportLayerProtocol.TCP, diff --git a/example_apps/offer_method_udp.py b/example_apps/offer_method_udp.py index d01fe0e..ca7d500 100644 --- a/example_apps/offer_method_udp.py +++ b/example_apps/offer_method_udp.py @@ -70,6 +70,11 @@ async def main(): someipy_daemon = await connect_to_someipy_daemon() + # For Windows use: + # someipy_daemon = await connect_to_someipy_daemon( + # {"use_tcp": True, "tcp_host": interface_ip, "tcp_port": 30500} + # ) + addition_method = Method( id=SAMPLE_METHOD_ID, protocol=TransportLayerProtocol.UDP, diff --git a/example_apps/offer_multiple_services.py b/example_apps/offer_multiple_services.py index 6387d0c..8136ffa 100644 --- a/example_apps/offer_multiple_services.py +++ b/example_apps/offer_multiple_services.py @@ -45,6 +45,11 @@ async def main(): someipy_daemon = await connect_to_someipy_daemon() + # For Windows use: + # someipy_daemon = await connect_to_someipy_daemon( + # {"use_tcp": True, "tcp_host": interface_ip, "tcp_port": 30500} + # ) + temperature_event = Event(id=SAMPLE_EVENT_ID, protocol=TransportLayerProtocol.UDP) temperature_eventgroup = EventGroup( diff --git a/example_apps/receive_events_tcp.py b/example_apps/receive_events_tcp.py index 5d18987..ef9535d 100644 --- a/example_apps/receive_events_tcp.py +++ b/example_apps/receive_events_tcp.py @@ -59,6 +59,11 @@ async def main(): someipy_daemon = await connect_to_someipy_daemon() + # For Windows use: + # someipy_daemon = await connect_to_someipy_daemon( + # {"use_tcp": True, "tcp_host": interface_ip, "tcp_port": 30500} + # ) + temperature_event = Event(id=SAMPLE_EVENT_ID, protocol=TransportLayerProtocol.TCP) temperature_eventgroup = EventGroup( id=SAMPLE_EVENTGROUP_ID, events=[temperature_event] @@ -95,9 +100,10 @@ async def main(): except asyncio.CancelledError as e: print("Shutdown..") finally: + print("Unsubscribe eventgroup..") + service_instance_temperature.unsubscribe_eventgroup(temperature_eventgroup) print("Shutdown service instance..") - await someipy_daemon.disconnect_from_daemon() print("End main task..") diff --git a/example_apps/receive_events_udp.py b/example_apps/receive_events_udp.py index b6b7cdb..8aeeb04 100644 --- a/example_apps/receive_events_udp.py +++ b/example_apps/receive_events_udp.py @@ -58,6 +58,11 @@ async def main(): someipy_daemon = await connect_to_someipy_daemon() + # For Windows use: + # someipy_daemon = await connect_to_someipy_daemon( + # {"use_tcp": True, "tcp_host": interface_ip, "tcp_port": 30500} + # ) + temperature_event = Event(id=SAMPLE_EVENT_ID, protocol=TransportLayerProtocol.UDP) temperature_eventgroup = EventGroup( id=SAMPLE_EVENTGROUP_ID, events=[temperature_event] @@ -94,9 +99,10 @@ async def main(): except asyncio.CancelledError as e: print("Shutdown..") finally: + print("Unsubscribe eventgroup..") + service_instance_temperature.unsubscribe_eventgroup(temperature_eventgroup) print("Shutdown service instance..") - await someipy_daemon.disconnect_from_daemon() print("End main task..") diff --git a/example_apps/send_events_tcp.py b/example_apps/send_events_tcp.py index 1ac0674..daad165 100644 --- a/example_apps/send_events_tcp.py +++ b/example_apps/send_events_tcp.py @@ -36,6 +36,11 @@ async def main(): someipy_daemon = await connect_to_someipy_daemon() + # For Windows use: + # someipy_daemon = await connect_to_someipy_daemon( + # {"use_tcp": True, "tcp_host": interface_ip, "tcp_port": 30500} + # ) + temperature_event = Event(id=SAMPLE_EVENT_ID, protocol=TransportLayerProtocol.TCP) temperature_eventgroup = EventGroup( diff --git a/example_apps/send_events_udp.py b/example_apps/send_events_udp.py index c2b6212..9e99169 100644 --- a/example_apps/send_events_udp.py +++ b/example_apps/send_events_udp.py @@ -36,6 +36,11 @@ async def main(): someipy_daemon = await connect_to_someipy_daemon() + # For Windows use: + # someipy_daemon = await connect_to_someipy_daemon( + # {"use_tcp": True, "tcp_host": interface_ip, "tcp_port": 30500} + # ) + temperature_event = Event(id=SAMPLE_EVENT_ID, protocol=TransportLayerProtocol.UDP) temperature_eventgroup = EventGroup( diff --git a/integration_tests/automated_tests/run_all.py b/integration_tests/automated_tests/run_all.py index c0cc602..8106dd9 100644 --- a/integration_tests/automated_tests/run_all.py +++ b/integration_tests/automated_tests/run_all.py @@ -14,7 +14,6 @@ vsomeip_library_path = "/home/christian/projects/someip/vsomeip_install/lib/" current_file_path = os.path.abspath(__file__) -repository = os.path.dirname(os.path.dirname(current_file_path)) repository = os.path.dirname(os.path.dirname(os.path.dirname(current_file_path))) test_durations = 60 # duration of each test in seconds diff --git a/setup.cfg b/setup.cfg index 5096cf2..e6172fa 100644 --- a/setup.cfg +++ b/setup.cfg @@ -1,6 +1,6 @@ [metadata] name = someipy -version = 2.0.0 +version = 2.1.0 author = Christian H. author_email = someipy.package@gmail.com description = A Python package implementing the SOME/IP protocol diff --git a/src/someipy/_internal/_common/endpoint.py b/src/someipy/_internal/_common/endpoint.py index 554cadd..59caaf6 100644 --- a/src/someipy/_internal/_common/endpoint.py +++ b/src/someipy/_internal/_common/endpoint.py @@ -1,4 +1,4 @@ -# Copyright (C) 2024 Christian H. +# Copyright (C) 2025 Christian H. # # This program is free software: you can redistribute it and/or modify # it under the terms of the GNU General Public License as published by diff --git a/src/someipy/_internal/_daemon/daemon_server.py b/src/someipy/_internal/_daemon/daemon_server.py index d3bde11..4552b16 100644 --- a/src/someipy/_internal/_daemon/daemon_server.py +++ b/src/someipy/_internal/_daemon/daemon_server.py @@ -1,3 +1,18 @@ +# Copyright (C) 2025 Christian H. +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# This program is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with this program. If not, see . + import asyncio import logging import os @@ -33,12 +48,12 @@ async def _handle_client(self, reader, writer): async def start( self, - use_uds: bool = True, + use_tcp: bool = False, socket_path: str | None = None, tcp_port: int | None = None, host: str = "127.0.0.1", ): - if use_uds: + if not use_tcp: if os.path.exists(socket_path): os.unlink(socket_path) diff --git a/src/someipy/_internal/_daemon/daemon_server_client.py b/src/someipy/_internal/_daemon/daemon_server_client.py index 1a41cc7..e7fbf1d 100644 --- a/src/someipy/_internal/_daemon/daemon_server_client.py +++ b/src/someipy/_internal/_daemon/daemon_server_client.py @@ -1,3 +1,18 @@ +# Copyright (C) 2025 Christian H. +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# This program is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with this program. If not, see . + import asyncio import json import logging diff --git a/src/someipy/_internal/_daemon/offer_service_storage.py b/src/someipy/_internal/_daemon/offer_service_storage.py index 53e92ac..907dbda 100644 --- a/src/someipy/_internal/_daemon/offer_service_storage.py +++ b/src/someipy/_internal/_daemon/offer_service_storage.py @@ -1,4 +1,4 @@ -# Copyright (C) 2024 Christian H. +# Copyright (C) 2025 Christian H. # # This program is free software: you can redistribute it and/or modify # it under the terms of the GNU General Public License as published by diff --git a/src/someipy/_internal/_daemon/someipy_daemon_client.py b/src/someipy/_internal/_daemon/someipy_daemon_client.py index 8d336cf..a35ca06 100644 --- a/src/someipy/_internal/_daemon/someipy_daemon_client.py +++ b/src/someipy/_internal/_daemon/someipy_daemon_client.py @@ -1,4 +1,4 @@ -# Copyright (C) 2024 Christian H. +# Copyright (C) 2025 Christian H. # # This program is free software: you can redistribute it and/or modify # it under the terms of the GNU General Public License as published by @@ -77,12 +77,14 @@ def __init__(self, config: dict = None): if self._config is None: self._use_tcp = platform.system() != "Linux" + self._tcp_host = "127.0.0.1" self._tcp_port = 30500 self._socket_path = "/tmp/someipyd.sock" else: self._socket_path = self._config.get("socket_path", "/tmp/someipyd.sock") self._use_tcp = self._config.get("use_tcp", platform.system() != "Linux") self._tcp_port = self._config.get("tcp_port", 30500) + self._tcp_host = self._config.get("tcp_host", "127.0.0.1") self._rx_message_queue: asyncio.Queue[DaemonMessage] = asyncio.Queue() self._rx_task: asyncio.Task = None @@ -155,7 +157,7 @@ async def _connect_to_daemon(self): if self._use_tcp: self.reader, self.writer = await asyncio.open_connection( - "127.0.0.1", self._tcp_port + self._tcp_host, self._tcp_port ) else: self.reader, self.writer = await asyncio.open_unix_connection( diff --git a/src/someipy/_internal/_daemon/subscription.py b/src/someipy/_internal/_daemon/subscription.py index f5ac842..0919e04 100644 --- a/src/someipy/_internal/_daemon/subscription.py +++ b/src/someipy/_internal/_daemon/subscription.py @@ -48,21 +48,19 @@ def __eq__(self, value: "Subscription") -> bool: and self.instance_id == value.instance_id and self.major_version == value.major_version and self.eventgroup == value.eventgroup - and self.ttl_seconds == value.ttl_seconds and self.client_endpoint == value.client_endpoint and self.server_endpoint == value.server_endpoint and self.protocols == value.protocols ) def __hash__(self) -> int: - # Do not include the timestamp in the hash calculation + # Do not include the timestamp and ttl in the hash calculation return hash( ( self.service_id, self.instance_id, self.major_version, self.eventgroup, - self.ttl_seconds, self.client_endpoint, self.server_endpoint, self.protocols, diff --git a/src/someipy/_internal/_daemon/uds_messages.py b/src/someipy/_internal/_daemon/uds_messages.py index cc0199b..bdaf212 100644 --- a/src/someipy/_internal/_daemon/uds_messages.py +++ b/src/someipy/_internal/_daemon/uds_messages.py @@ -155,9 +155,11 @@ class StopSubscribeEventGroupRequest(BaseMessage): service_id: int instance_id: int major_version: int - eventgroup_id: int + eventgroup: str client_endpoint_ip: str client_endpoint_port: int + udp: bool + tcp: bool class ReceivedEvent(BaseMessage): diff --git a/src/someipy/_internal/_sd/deserialization/sd_deserialization.py b/src/someipy/_internal/_sd/deserialization/sd_deserialization.py index 4dd7507..686003f 100644 --- a/src/someipy/_internal/_sd/deserialization/sd_deserialization.py +++ b/src/someipy/_internal/_sd/deserialization/sd_deserialization.py @@ -1,4 +1,4 @@ -# Copyright (C) 2024 Christian H. +# Copyright (C) 2025 Christian H. # # This program is free software: you can redistribute it and/or modify # it under the terms of the GNU General Public License as published by diff --git a/src/someipy/_internal/_sd/deserialization/sd_serialization.py b/src/someipy/_internal/_sd/deserialization/sd_serialization.py index d886924..fcdc68f 100644 --- a/src/someipy/_internal/_sd/deserialization/sd_serialization.py +++ b/src/someipy/_internal/_sd/deserialization/sd_serialization.py @@ -1,4 +1,4 @@ -# Copyright (C) 2024 Christian H. +# Copyright (C) 2025 Christian H. # # This program is free software: you can redistribute it and/or modify # it under the terms of the GNU General Public License as published by diff --git a/src/someipy/_internal/_sd/entries/find_service_entry.py b/src/someipy/_internal/_sd/entries/find_service_entry.py index 3884a81..83d30ff 100644 --- a/src/someipy/_internal/_sd/entries/find_service_entry.py +++ b/src/someipy/_internal/_sd/entries/find_service_entry.py @@ -1,4 +1,4 @@ -# Copyright (C) 2024 Christian H. +# Copyright (C) 2025 Christian H. # # This program is free software: you can redistribute it and/or modify # it under the terms of the GNU General Public License as published by diff --git a/src/someipy/_internal/_sd/entries/offer_service_entry.py b/src/someipy/_internal/_sd/entries/offer_service_entry.py index f1cb24f..a9cdf78 100644 --- a/src/someipy/_internal/_sd/entries/offer_service_entry.py +++ b/src/someipy/_internal/_sd/entries/offer_service_entry.py @@ -1,4 +1,4 @@ -# Copyright (C) 2024 Christian H. +# Copyright (C) 2025 Christian H. # # This program is free software: you can redistribute it and/or modify # it under the terms of the GNU General Public License as published by diff --git a/src/someipy/_internal/_sd/entries/sd_entry.py b/src/someipy/_internal/_sd/entries/sd_entry.py index a8bbbba..5badb00 100644 --- a/src/someipy/_internal/_sd/entries/sd_entry.py +++ b/src/someipy/_internal/_sd/entries/sd_entry.py @@ -1,4 +1,4 @@ -# Copyright (C) 2024 Christian H. +# Copyright (C) 2025 Christian H. # # This program is free software: you can redistribute it and/or modify # it under the terms of the GNU General Public License as published by diff --git a/src/someipy/_internal/_sd/entries/stop_offer_service_entry.py b/src/someipy/_internal/_sd/entries/stop_offer_service_entry.py index 39421c2..7c804f7 100644 --- a/src/someipy/_internal/_sd/entries/stop_offer_service_entry.py +++ b/src/someipy/_internal/_sd/entries/stop_offer_service_entry.py @@ -1,4 +1,4 @@ -# Copyright (C) 2024 Christian H. +# Copyright (C) 2025 Christian H. # # This program is free software: you can redistribute it and/or modify # it under the terms of the GNU General Public License as published by diff --git a/src/someipy/_internal/_sd/entries/stop_subscribe_eventgroup_entry.py b/src/someipy/_internal/_sd/entries/stop_subscribe_eventgroup_entry.py index 946f6bf..22ded39 100644 --- a/src/someipy/_internal/_sd/entries/stop_subscribe_eventgroup_entry.py +++ b/src/someipy/_internal/_sd/entries/stop_subscribe_eventgroup_entry.py @@ -1,4 +1,4 @@ -# Copyright (C) 2024 Christian H. +# Copyright (C) 2025 Christian H. # # This program is free software: you can redistribute it and/or modify # it under the terms of the GNU General Public License as published by @@ -29,7 +29,6 @@ class StopSubscribeEventGroupEntry(SdEntry): service_id: int instance_id: int major_version: int - minor_version: int eventgroup_id: int counter: int ip_v4_endpoints: List[IpV4EndpointOption] @@ -40,7 +39,6 @@ def __init__( service_id: int, instance_id: int, major_version: int, - minor_version: int, eventgroup_id: int, counter: int, ip_v4_endpoints: List[IpV4EndpointOption], @@ -50,7 +48,6 @@ def __init__( self.service_id = service_id self.instance_id = instance_id self.major_version = major_version - self.minor_version = minor_version self.eventgroup_id = eventgroup_id self.counter = counter self.ip_v4_endpoints = ip_v4_endpoints diff --git a/src/someipy/_internal/_sd/entries/subscribe_ack_entry.py b/src/someipy/_internal/_sd/entries/subscribe_ack_entry.py index 6550b90..5a6b749 100644 --- a/src/someipy/_internal/_sd/entries/subscribe_ack_entry.py +++ b/src/someipy/_internal/_sd/entries/subscribe_ack_entry.py @@ -1,4 +1,4 @@ -# Copyright (C) 2024 Christian H. +# Copyright (C) 2025 Christian H. # # This program is free software: you can redistribute it and/or modify # it under the terms of the GNU General Public License as published by @@ -34,7 +34,6 @@ def __init__( service_id: int, instance_id: int, major_version: int, - minor_version: int, ttl: int, eventgroup_id: int, counter: int, @@ -43,7 +42,6 @@ def __init__( self.service_id = service_id self.instance_id = instance_id self.major_version = major_version - self.minor_version = minor_version self.ttl = ttl self.eventgroup_id = eventgroup_id self.counter = counter diff --git a/src/someipy/_internal/_sd/entries/subscribe_eventgroup_entry.py b/src/someipy/_internal/_sd/entries/subscribe_eventgroup_entry.py index 627dac6..1ef70f0 100644 --- a/src/someipy/_internal/_sd/entries/subscribe_eventgroup_entry.py +++ b/src/someipy/_internal/_sd/entries/subscribe_eventgroup_entry.py @@ -1,4 +1,4 @@ -# Copyright (C) 2024 Christian H. +# Copyright (C) 2025 Christian H. # # This program is free software: you can redistribute it and/or modify # it under the terms of the GNU General Public License as published by @@ -29,7 +29,6 @@ class SubscribeEventGroupEntry(SdEntry): service_id: int instance_id: int major_version: int - minor_version: int ttl: int eventgroup_id: int counter: int @@ -41,7 +40,6 @@ def __init__( service_id: int, instance_id: int, major_version: int, - minor_version: int, ttl: int, eventgroup_id: int, counter: int, @@ -52,7 +50,6 @@ def __init__( self.service_id = service_id self.instance_id = instance_id self.major_version = major_version - self.minor_version = minor_version self.ttl = ttl self.eventgroup_id = eventgroup_id self.counter = counter diff --git a/src/someipy/_internal/_sd/entries/subscribe_eventgroup_nack.py b/src/someipy/_internal/_sd/entries/subscribe_eventgroup_nack.py index 42c16ce..0d77f70 100644 --- a/src/someipy/_internal/_sd/entries/subscribe_eventgroup_nack.py +++ b/src/someipy/_internal/_sd/entries/subscribe_eventgroup_nack.py @@ -1,4 +1,4 @@ -# Copyright (C) 2024 Christian H. +# Copyright (C) 2025 Christian H. # # This program is free software: you can redistribute it and/or modify # it under the terms of the GNU General Public License as published by @@ -33,7 +33,6 @@ def __init__( service_id: int, instance_id: int, major_version: int, - minor_version: int, eventgroup_id: int, counter: int, ): @@ -41,6 +40,5 @@ def __init__( self.service_id = service_id self.instance_id = instance_id self.major_version = major_version - self.minor_version = minor_version self.eventgroup_id = eventgroup_id self.counter = counter diff --git a/src/someipy/_internal/_sd/options/configuration_option.py b/src/someipy/_internal/_sd/options/configuration_option.py index fd8d46a..e2280de 100644 --- a/src/someipy/_internal/_sd/options/configuration_option.py +++ b/src/someipy/_internal/_sd/options/configuration_option.py @@ -1,4 +1,4 @@ -# Copyright (C) 2024 Christian H. +# Copyright (C) 2025 Christian H. # # This program is free software: you can redistribute it and/or modify # it under the terms of the GNU General Public License as published by diff --git a/src/someipy/_internal/_sd/options/endpoint.py b/src/someipy/_internal/_sd/options/endpoint.py index a00334a..f521b69 100644 --- a/src/someipy/_internal/_sd/options/endpoint.py +++ b/src/someipy/_internal/_sd/options/endpoint.py @@ -1,4 +1,4 @@ -# Copyright (C) 2024 Christian H. +# Copyright (C) 2025 Christian H. # # This program is free software: you can redistribute it and/or modify # it under the terms of the GNU General Public License as published by diff --git a/src/someipy/_internal/_sd/options/load_balancing.py b/src/someipy/_internal/_sd/options/load_balancing.py index b8ba81a..4af25e6 100644 --- a/src/someipy/_internal/_sd/options/load_balancing.py +++ b/src/someipy/_internal/_sd/options/load_balancing.py @@ -1,4 +1,4 @@ -# Copyright (C) 2024 Christian H. +# Copyright (C) 2025 Christian H. # # This program is free software: you can redistribute it and/or modify # it under the terms of the GNU General Public License as published by diff --git a/src/someipy/_internal/_sd/options/multicast.py b/src/someipy/_internal/_sd/options/multicast.py index c6c7c55..be8e971 100644 --- a/src/someipy/_internal/_sd/options/multicast.py +++ b/src/someipy/_internal/_sd/options/multicast.py @@ -1,4 +1,4 @@ -# Copyright (C) 2024 Christian H. +# Copyright (C) 2025 Christian H. # # This program is free software: you can redistribute it and/or modify # it under the terms of the GNU General Public License as published by diff --git a/src/someipy/_internal/_sd/options/sd_endpoint.py b/src/someipy/_internal/_sd/options/sd_endpoint.py index cc89258..2d81402 100644 --- a/src/someipy/_internal/_sd/options/sd_endpoint.py +++ b/src/someipy/_internal/_sd/options/sd_endpoint.py @@ -1,4 +1,4 @@ -# Copyright (C) 2024 Christian H. +# Copyright (C) 2025 Christian H. # # This program is free software: you can redistribute it and/or modify # it under the terms of the GNU General Public License as published by diff --git a/src/someipy/_internal/daemon_client_abcs.py b/src/someipy/_internal/daemon_client_abcs.py index 1001e42..5ce9df1 100644 --- a/src/someipy/_internal/daemon_client_abcs.py +++ b/src/someipy/_internal/daemon_client_abcs.py @@ -1,4 +1,4 @@ -# Copyright (C) 2024 Christian H. +# Copyright (C) 2025 Christian H. # # This program is free software: you can redistribute it and/or modify # it under the terms of the GNU General Public License as published by diff --git a/src/someipy/_internal/someip_endpoint_factory.py b/src/someipy/_internal/someip_endpoint_factory.py index d16a75a..f334f26 100644 --- a/src/someipy/_internal/someip_endpoint_factory.py +++ b/src/someipy/_internal/someip_endpoint_factory.py @@ -17,6 +17,7 @@ from collections.abc import Callable import logging from typing import Tuple +from someipy._internal._common.endpoint import Endpoint from someipy._internal.someip_endpoint import ( SomeipEndpoint, TCPClientSomeipEndpoint, @@ -33,8 +34,7 @@ class SomeipEndpointFactory: @staticmethod async def create_server_endpoint( - ip_address: str, - port: int, + endpoint: Endpoint, protocol: TransportLayerProtocol, someip_callback: Callable[ [SomeIpMessage, Tuple[str, int], Tuple[str, int], TransportLayerProtocol], @@ -44,25 +44,26 @@ async def create_server_endpoint( if protocol == TransportLayerProtocol.UDP: loop = asyncio.get_running_loop() - rcv_socket = create_udp_socket(ip_address, port) + rcv_socket = create_udp_socket(str(endpoint.ip), endpoint.port) _, udp_endpoint = await loop.create_datagram_endpoint( - lambda: UDPSomeipEndpoint(ip_address, port), sock=rcv_socket + lambda: UDPSomeipEndpoint(str(endpoint.ip), endpoint.port), + sock=rcv_socket, ) udp_endpoint.set_someip_callback(someip_callback) return udp_endpoint else: - tcp_client_manager = TcpClientManager(ip_address, port) + tcp_client_manager = TcpClientManager(str(endpoint.ip), endpoint.port) loop = asyncio.get_running_loop() server = await loop.create_server( lambda: TcpClientProtocol(client_manager=tcp_client_manager), - ip_address, - port, + str(endpoint.ip), + endpoint.port, ) tcp_someip_endpoint = TCPSomeipEndpoint( - server, tcp_client_manager, ip_address, port + server, tcp_client_manager, str(endpoint.ip), endpoint.port ) tcp_someip_endpoint.set_someip_callback(someip_callback) @@ -70,41 +71,32 @@ async def create_server_endpoint( return tcp_someip_endpoint @staticmethod - async def create_client_endpoint( - dst_ip: str, - dst_port: int, - src_ip: str, - src_port: int, - protocol: TransportLayerProtocol, + async def create_udp_client_endpoint( + dst_endpoint: Endpoint, + src_endpoint: Endpoint, someip_message_callback: Callable[[SomeIpMessage], None], logger: logging.Logger = None, ) -> SomeipEndpoint: - if protocol == TransportLayerProtocol.UDP: - udp_endpoint = await SomeipEndpointFactory.create_server_endpoint( - src_ip, - src_port, - TransportLayerProtocol.UDP, - someip_message_callback, - ) - return udp_endpoint - else: - tcp_endpoint = TCPClientSomeipEndpoint( - dst_ip, dst_port, src_ip, src_port, logger - ) - tcp_endpoint.set_someip_callback(someip_message_callback) - return tcp_endpoint + udp_endpoint = await SomeipEndpointFactory.create_server_endpoint( + src_endpoint, + TransportLayerProtocol.UDP, + someip_message_callback, + ) + return udp_endpoint @staticmethod def create_tcp_client_endpoint( - dst_ip: str, - dst_port: int, - src_ip: str, - src_port: int, + dst_endpoint: Endpoint, + src_endpoint: Endpoint, someip_message_callback: Callable[[SomeIpMessage], None], logger: logging.Logger = None, ) -> TCPClientSomeipEndpoint: tcp_endpoint = TCPClientSomeipEndpoint( - dst_ip, dst_port, src_ip, src_port, logger + str(dst_endpoint.ip), + dst_endpoint.port, + str(src_endpoint.ip), + src_endpoint.port, + logger, ) tcp_endpoint.set_someip_callback(someip_message_callback) return tcp_endpoint diff --git a/src/someipy/client_service_instance.py b/src/someipy/client_service_instance.py index 6c7c0f8..385b5e4 100644 --- a/src/someipy/client_service_instance.py +++ b/src/someipy/client_service_instance.py @@ -312,15 +312,17 @@ def subscribe_eventgroup( self._daemon.transmit_message_to_daemon(method_request) - def unsubscribe_eventgroup(self, eventgroup_id: int): + def unsubscribe_eventgroup(self, eventgroup: EventGroup): method_request = create_uds_message( StopSubscribeEventGroupRequest, service_id=self._service.id, instance_id=self.instance_id, major_version=self._service.major_version, - eventgroup_id=eventgroup_id, + eventgroup=eventgroup.to_json(), client_endpoint_ip=self._endpoint_ip, client_endpoint_port=self._endpoint_port, + udp=eventgroup.has_udp, + tcp=eventgroup.has_tcp, ) self._daemon.transmit_message_to_daemon(method_request) diff --git a/src/someipy/someipyd.json b/src/someipy/someipyd.json index 04cb781..3236467 100644 --- a/src/someipy/someipyd.json +++ b/src/someipy/someipyd.json @@ -2,5 +2,8 @@ "sd_address": "224.224.224.245", "sd_port": 30490, "log_level": "DEBUG", - "interface": "127.0.0.1" + "interface": "127.0.0.1", + "use_tcp": false, + "tcp_host": "127.0.0.1", + "tcp_port": 30500 } \ No newline at end of file diff --git a/src/someipy/someipyd.py b/src/someipy/someipyd.py index 06dd491..354ed9b 100644 --- a/src/someipy/someipyd.py +++ b/src/someipy/someipyd.py @@ -42,6 +42,9 @@ ) from someipy._internal._sd.deserialization.sd_serialization import serialize_sd_message from someipy._internal._sd.entries.offer_service_entry import OfferServiceEntry +from someipy._internal._sd.entries.stop_subscribe_eventgroup_entry import ( + StopSubscribeEventGroupEntry, +) from someipy._internal._sd.entries.subscribe_eventgroup_entry import ( SubscribeEventGroupEntry, ) @@ -577,8 +580,10 @@ async def _handle_subscribe_eventgroup_request( ) udp_endpoint = await self._endpoint_factory.create_server_endpoint( - message["client_endpoint_ip"], - message["client_endpoint_port"], + Endpoint( + ip=ipaddress.IPv4Address(message["client_endpoint_ip"]), + port=message["client_endpoint_port"], + ), TransportLayerProtocol.UDP, self._someip_message_callback, ) @@ -591,32 +596,139 @@ async def _handle_subscribe_eventgroup_request( event_group = EventGroup.from_json(message["eventgroup"]) client_address = ipaddress.IPv4Address(message["client_endpoint_ip"]) + client_endpoint = Endpoint( + ip=client_address, port=message["client_endpoint_port"] + ) new_subscription = Subscription( service_id=message["service_id"], instance_id=message["instance_id"], major_version=message["major_version"], - client_endpoint=Endpoint( - ip=client_address, port=message["client_endpoint_port"] - ), + client_endpoint=client_endpoint, server_endpoint=None, protocols=frozenset(protocols), eventgroup=event_group, ttl_seconds=message["ttl_subscription"], ) + if new_subscription in self._requested_subscriptions.subscriptions: + self.logger.warning( + f"The requested subscription received by client {client_id} is already requested." + ) + self.logger.debug(f"Add subscription to storage with id {client_id}") self._requested_subscriptions.add_subscription(client_id, new_subscription) def _handle_stop_subscribe_eventgroup_request( - self, message: StopSubscribeEventGroupRequest, writer_id: int + self, message: StopSubscribeEventGroupRequest, client_id: int ): - # TODO: Remove from self._requested_subscriptions - # Check if there is an active subscription. If yes, send out a stop subscribe message - pass + client_endpoint = Endpoint( + ip=ipaddress.IPv4Address(message["client_endpoint_ip"]), + port=message["client_endpoint_port"], + ) + + event_group = EventGroup.from_json(message["eventgroup"]) + + protocols = [ + protocol + for flag, protocol in ( + (message["udp"], TransportLayerProtocol.UDP), + (message["tcp"], TransportLayerProtocol.TCP), + ) + if flag + ] + + subscription_to_remove = Subscription( + service_id=message["service_id"], + instance_id=message["instance_id"], + major_version=message["major_version"], + client_endpoint=client_endpoint, + server_endpoint=None, + protocols=frozenset(protocols), + eventgroup=event_group, + ttl_seconds=0, # TTL is not relevant for removal + ) + + self._requested_subscriptions.remove_subscription( + client_id, subscription_to_remove + ) + + if ( + len(self._requested_subscriptions.get_client_ids(subscription_to_remove)) + == 0 + ): + + for pending_subscription in list(self._pending_subscriptions): + if ( + pending_subscription.service_id == subscription_to_remove.service_id + and pending_subscription.instance_id + == subscription_to_remove.instance_id + and pending_subscription.major_version + == subscription_to_remove.major_version + and pending_subscription.eventgroup.id + == subscription_to_remove.eventgroup.id + and pending_subscription.client_endpoint + == subscription_to_remove.client_endpoint + ): + self._pending_subscriptions.remove(pending_subscription) + + for active_subscription in list(self._active_subscriptions): + if ( + active_subscription.service_id == subscription_to_remove.service_id + and active_subscription.instance_id + == subscription_to_remove.instance_id + and active_subscription.major_version + == subscription_to_remove.major_version + and active_subscription.eventgroup.id + == subscription_to_remove.eventgroup.id + and active_subscription.client_endpoint + == subscription_to_remove.client_endpoint + ): + self._active_subscriptions.remove(active_subscription) + + ( + session_id, + reboot_flag, + ) = self._unicast_session_handler.update_session() + + # Build subscribe message + sd_message = SdMessage() + sd_message.session_id = session_id + sd_message.reboot_flag = reboot_flag + + options = [] + for protocol in active_subscription.protocols: + options.append( + IpV4EndpointOption( + address=active_subscription.client_endpoint.ip, + protocol=protocol, + port=active_subscription.client_endpoint.port, + ) + ) + + entry = StopSubscribeEventGroupEntry( + service_id=active_subscription.service_id, + instance_id=active_subscription.instance_id, + major_version=active_subscription.major_version, + eventgroup_id=active_subscription.eventgroup.id, + counter=0, + ip_v4_endpoints=options, + ip_v6_endpoints=[], + ) + sd_message.entries.append(entry) + + # Send to server + if self._ucast_transport: + self._ucast_transport.sendto( + serialize_sd_message(sd_message), + ( + str(active_subscription.server_endpoint.ip), + self.sd_port, + ), + ) async def _handle_offer_service_request( - self, message: OfferServiceRequest, writer_id: int + self, message: OfferServiceRequest, client_id: int ): method_strs = message.get("method_list", []) methods = [Method.from_json(m) for m in method_strs] @@ -634,7 +746,7 @@ async def _handle_offer_service_request( """ service_to_add = ServiceToOffer( - client_writer_id=writer_id, + client_writer_id=client_id, instance_id=message["instance_id"], service_id=message["service_id"], major_version=message["major_version"], @@ -662,13 +774,12 @@ async def _handle_offer_service_request( ) udp_endpoint = await self._endpoint_factory.create_server_endpoint( - str(service_to_add.endpoint.ip), - service_to_add.endpoint.port, + service_to_add.endpoint, TransportLayerProtocol.UDP, self._someip_message_callback, ) - self._someip_server_endpoints.add_endpoint(writer_id, udp_endpoint) + self._someip_server_endpoints.add_endpoint(client_id, udp_endpoint) if service_to_add.has_tcp: if not self._someip_server_endpoints.has_endpoint( @@ -681,13 +792,12 @@ async def _handle_offer_service_request( ) tcp_endpoint = await self._endpoint_factory.create_server_endpoint( - str(service_to_add.endpoint.ip), - service_to_add.endpoint.port, + service_to_add.endpoint, TransportLayerProtocol.TCP, self._someip_message_callback, ) - self._someip_server_endpoints.add_endpoint(writer_id, tcp_endpoint) + self._someip_server_endpoints.add_endpoint(client_id, tcp_endpoint) cyclic_offer_delay_ms = message["cyclic_offer_delay_ms"] @@ -817,7 +927,7 @@ def _handle_inbound_call_method_response( ) async def _handle_outbound_call_method_request( - self, message: OutboundCallMethodRequest, writer_id: int + self, message: OutboundCallMethodRequest, client_id: int ): endpoint = None if TransportLayerProtocol(message["protocol"]) == TransportLayerProtocol.UDP: @@ -830,15 +940,22 @@ async def _handle_outbound_call_method_request( f"Creating new UDP endpoint for {message['src_endpoint_ip']}:{message['src_endpoint_port']}" ) - udp_endpoint = await self._endpoint_factory.create_client_endpoint( - message["dst_endpoint_ip"], - message["dst_endpoint_port"], - message["src_endpoint_ip"], - message["src_endpoint_port"], - TransportLayerProtocol.UDP, + dst_endpoint = Endpoint( + ip=ipaddress.IPv4Address(message["dst_endpoint_ip"]), + port=message["dst_endpoint_port"], + ) + src_endpoint = Endpoint( + ip=ipaddress.IPv4Address(message["src_endpoint_ip"]), + port=message["src_endpoint_port"], + ) + + udp_endpoint = await self._endpoint_factory.create_udp_client_endpoint( + dst_endpoint, + src_endpoint, self._someip_message_callback, self.logger, ) + self._someip_client_endpoints.add_endpoint(client_id, udp_endpoint) endpoint = udp_endpoint else: @@ -861,15 +978,19 @@ async def _handle_outbound_call_method_request( ) tcp_endpoint = self._endpoint_factory.create_tcp_client_endpoint( - message["dst_endpoint_ip"], - message["dst_endpoint_port"], - message["src_endpoint_ip"], - message["src_endpoint_port"], + Endpoint( + ip=ipaddress.IPv4Address(message["dst_endpoint_ip"]), + port=message["dst_endpoint_port"], + ), + Endpoint( + ip=ipaddress.IPv4Address(message["src_endpoint_ip"]), + port=message["src_endpoint_port"], + ), self._someip_message_callback, self.logger, ) - self._someip_client_endpoints.add_endpoint(writer_id, tcp_endpoint) + self._someip_client_endpoints.add_endpoint(client_id, tcp_endpoint) endpoint: TCPClientSomeipEndpoint = tcp_endpoint else: endpoint: TCPClientSomeipEndpoint = ( @@ -924,7 +1045,7 @@ async def _handle_outbound_call_method_request( f"Method call {new_call} already issued. Overwriting writer_id." ) - self._issued_method_calls[new_call] = writer_id + self._issued_method_calls[new_call] = client_id endpoint.sendto( someip_message.serialize(), @@ -1237,39 +1358,38 @@ def _handle_offered_service(self, offered_service: ServiceInstance): self._found_services[index].timestamp = offered_service.timestamp # Check if there is a requested subscription for this service - for requested_subscription in self._requested_subscriptions.has_subscriptions( + for ( + requested_subscription, + client_id, + ) in self._requested_subscriptions.has_subscriptions( offered_service.service_id, offered_service.instance_id, offered_service.major_version, ): - requested_protocols: Set[TransportLayerProtocol] = set() - for protocol in offered_service.protocols: - if protocol in requested_subscription[0].protocols: - requested_protocols.add(protocol) + + requested_protocols: Set[TransportLayerProtocol] = ( + offered_service.protocols & requested_subscription.protocols + ) if TransportLayerProtocol.TCP in requested_protocols: if not self._someip_client_endpoints.has_tcp_endpoint( - str(requested_subscription[0].client_endpoint.ip), - requested_subscription[0].client_endpoint.port, + str(requested_subscription.client_endpoint.ip), + requested_subscription.client_endpoint.port, str(offered_service.endpoint.ip), offered_service.endpoint.port, ): self.logger.debug( - f"Creating new TCP endpoint for {requested_subscription[0].client_endpoint}" + f"Creating new TCP endpoint for {requested_subscription.client_endpoint}" ) tcp_endpoint = self._endpoint_factory.create_tcp_client_endpoint( - str(offered_service.endpoint.ip), - offered_service.endpoint.port, - str(requested_subscription[0].client_endpoint.ip), - requested_subscription[0].client_endpoint.port, + offered_service.endpoint, + requested_subscription.client_endpoint, self._someip_message_callback, self.logger, ) - self._someip_client_endpoints.add_endpoint( - requested_subscription[1], tcp_endpoint - ) + self._someip_client_endpoints.add_endpoint(client_id, tcp_endpoint) # TODO: This shall not block the handle_client function. A new task shall be created # For TCP wait for the connection to be established @@ -1290,9 +1410,9 @@ def _handle_offered_service(self, offered_service: ServiceInstance): for protocol in requested_protocols: options.append( IpV4EndpointOption( - address=requested_subscription[0].client_endpoint.ip, + address=requested_subscription.client_endpoint.ip, protocol=protocol, - port=requested_subscription[0].client_endpoint.port, + port=requested_subscription.client_endpoint.port, ) ) @@ -1300,24 +1420,23 @@ def _handle_offered_service(self, offered_service: ServiceInstance): service_id=offered_service.service_id, instance_id=offered_service.instance_id, major_version=offered_service.major_version, - minor_version=offered_service.minor_version, - ttl=requested_subscription[0].ttl_seconds, - eventgroup_id=requested_subscription[0].eventgroup.id, + ttl=requested_subscription.ttl_seconds, + eventgroup_id=requested_subscription.eventgroup.id, counter=0, ip_v4_endpoints=options, ip_v6_endpoints=[], ) sd_message.entries.append(entry) - client_endpoint = requested_subscription[0].client_endpoint + client_endpoint = requested_subscription.client_endpoint server_endpoint = offered_service.endpoint pending_subscription = Subscription( service_id=offered_service.service_id, instance_id=offered_service.instance_id, major_version=offered_service.major_version, - eventgroup=requested_subscription[0].eventgroup, - ttl_seconds=requested_subscription[0].ttl_seconds, + eventgroup=requested_subscription.eventgroup, + ttl_seconds=requested_subscription.ttl_seconds, client_endpoint=client_endpoint, server_endpoint=server_endpoint, protocols=frozenset(requested_protocols), @@ -1656,15 +1775,16 @@ async def async_main(): logger = _configure_logging(log_level=log_level, log_path=log_path) logger.info( - f"Starting SOME/IP Daemon with config:\n" + f"Starting SOME/IP daemon with config:\n" f"Socket path: {config.get('socket_path', DEFAULT_SOCKET_PATH)}\n" f"SD address: {config.get('sd_address', DEFAULT_SD_ADDRESS)}\n" f"SD port: {config.get('sd_port', DEFAULT_SD_PORT)}\n" f"Interface: {config.get('interface', DEFAULT_INTERFACE_IP)}\n" f"Loglevel: {log_level}\n" f"Log path: {log_path if log_path else 'Console'}\n" - f"Use TCP: {config.get('use_tcp', False)}\n" - f"TCP Port: {config.get('tcp_port', None)}\n" + f"Use tcp: {config.get('use_tcp', False)}\n" + f"Tcp port: {config.get('tcp_port', None)}\n" + f"Tcp host: {config.get('tcp_host', '127.0.0.1')}\n" ) daemon_server = DaemonServer(logger) @@ -1675,10 +1795,10 @@ async def async_main(): daemon_server.client_disconnected += daemon.client_disconnected await daemon_server.start( - use_uds=config.get("use_uds", True), + use_tcp=config.get("use_tcp", False), socket_path=config.get("socket_path", DEFAULT_SOCKET_PATH), - tcp_port=config.get("tcp_port", None), - host=config.get("host", "127.0.0.1"), + tcp_port=config.get("tcp_port", 30500), + host=config.get("tcp_host", "127.0.0.1"), ) await daemon.start_server() diff --git a/tests/test_someipyd.py b/tests/daemon/test_someipyd.py similarity index 76% rename from tests/test_someipyd.py rename to tests/daemon/test_someipyd.py index 9aa01bd..ad5f030 100644 --- a/tests/test_someipyd.py +++ b/tests/daemon/test_someipyd.py @@ -10,9 +10,11 @@ from someipy._internal._common.endpoint import Endpoint from someipy._internal._daemon.daemon_server import ClientConnectedEventArgs from someipy._internal._daemon.daemon_server_client import DaemonServerClient +from someipy._internal._daemon.subscription import Subscription from someipy._internal._daemon.uds_messages import ( OfferServiceRequest, StopOfferServiceRequest, + StopSubscribeEventGroupRequest, SubscribeEventGroupRequest, create_uds_message, ) @@ -60,7 +62,7 @@ def service_instance() -> ServiceInstance: major_version=3, minor_version=0, ttl=10, - endpoint=Endpoint("123", 1), + endpoint=Endpoint(ipaddress.IPv4Address("192.168.1.1"), 1), protocols=frozenset([TransportLayerProtocol.UDP]), timestamp=1000, ) @@ -94,7 +96,7 @@ def subscribe_event_group_request_udp(eventgroup) -> SubscribeEventGroupRequest: major_version=3, ttl_subscription=10, eventgroup=eventgroup.to_json(), - client_endpoint_ip="123", + client_endpoint_ip="192.168.1.1", client_endpoint_port=1, udp=True, tcp=False, @@ -109,7 +111,7 @@ def subscribe_event_group_request_tcp(eventgroup) -> SubscribeEventGroupRequest: major_version=3, ttl_subscription=10, eventgroup=eventgroup.to_json(), - client_endpoint_ip="123", + client_endpoint_ip="192.168.1.1", client_endpoint_port=1, udp=False, tcp=True, @@ -132,6 +134,23 @@ def subscribe_event_group_request_udp_and_tcp(eventgroup) -> SubscribeEventGroup ) +@pytest.fixture +def stop_subscribe_eventgroup_request( + eventgroup: EventGroup, +) -> StopSubscribeEventGroupRequest: + return create_uds_message( + StopSubscribeEventGroupRequest, + service_id=1, + instance_id=2, + major_version=3, + eventgroup=eventgroup.to_json(), + client_endpoint_ip="192.168.1.1", + client_endpoint_port=1, + udp=eventgroup.has_udp, + tcp=eventgroup.has_tcp, + ) + + @pytest.fixture def offer_service_request(eventgroup, method) -> OfferServiceRequest: return create_uds_message( @@ -228,12 +247,9 @@ async def test_handle_offered_service_opens_tcp_client_endpoint( daemon._handle_offered_service(service_instance) # Verify that create_client_endpoint was called for TCP - mock_endpoint_factory.create_client_endpoint.assert_called_once_with( - str(service_instance.endpoint.ip), - service_instance.endpoint.port, - str(service_instance.endpoint.ip), - service_instance.endpoint.port, - TransportLayerProtocol.TCP, + mock_endpoint_factory.create_tcp_client_endpoint.assert_called_once_with( + service_instance.endpoint, + service_instance.endpoint, daemon._someip_message_callback, daemon.logger, ) @@ -418,3 +434,91 @@ def test_find_service_request_sends_negative_response(daemon: SomeipDaemon): def test_find_service_request_sends_positive_response(daemon: SomeipDaemon): pass + + +def test_stop_subscribe_eventgroup_request_removes_requested_subscriptions( + daemon: SomeipDaemon, + eventgroup: EventGroup, + stop_subscribe_eventgroup_request: StopSubscribeEventGroupRequest, +): + protocols = [ + protocol + for flag, protocol in ( + (stop_subscribe_eventgroup_request["udp"], TransportLayerProtocol.UDP), + (stop_subscribe_eventgroup_request["tcp"], TransportLayerProtocol.TCP), + ) + if flag + ] + + new_subscription = Subscription( + service_id=stop_subscribe_eventgroup_request["service_id"], + instance_id=stop_subscribe_eventgroup_request["instance_id"], + major_version=stop_subscribe_eventgroup_request["major_version"], + eventgroup=eventgroup, + ttl_seconds=10, + client_endpoint=Endpoint( + ipaddress.IPv4Address( + stop_subscribe_eventgroup_request["client_endpoint_ip"] + ), + stop_subscribe_eventgroup_request["client_endpoint_port"], + ), + server_endpoint=None, + protocols=frozenset(protocols), + ) + + daemon._requested_subscriptions.add_subscription( + 1, + new_subscription, + ) + + assert len(daemon._requested_subscriptions) == 1 + + daemon._handle_stop_subscribe_eventgroup_request( + stop_subscribe_eventgroup_request, 1 + ) + + assert len(daemon._requested_subscriptions) == 0 + + +def test_stop_subscribe_eventgroup_request_removes_pending_and_active_subscriptions( + daemon: SomeipDaemon, + eventgroup: EventGroup, + stop_subscribe_eventgroup_request: StopSubscribeEventGroupRequest, +): + protocols = [ + protocol + for flag, protocol in ( + (stop_subscribe_eventgroup_request["udp"], TransportLayerProtocol.UDP), + (stop_subscribe_eventgroup_request["tcp"], TransportLayerProtocol.TCP), + ) + if flag + ] + + endpoint = Endpoint( + ipaddress.IPv4Address(stop_subscribe_eventgroup_request["client_endpoint_ip"]), + stop_subscribe_eventgroup_request["client_endpoint_port"], + ) + + new_subscription = Subscription( + service_id=stop_subscribe_eventgroup_request["service_id"], + instance_id=stop_subscribe_eventgroup_request["instance_id"], + major_version=stop_subscribe_eventgroup_request["major_version"], + eventgroup=eventgroup, + ttl_seconds=10, + client_endpoint=endpoint, + server_endpoint=endpoint, + protocols=frozenset(protocols), + ) + + daemon._pending_subscriptions.add(new_subscription) + daemon._active_subscriptions.add(new_subscription) + + assert len(daemon._pending_subscriptions) == 1 + assert len(daemon._active_subscriptions) == 1 + + daemon._handle_stop_subscribe_eventgroup_request( + stop_subscribe_eventgroup_request, 1 + ) + + assert len(daemon._pending_subscriptions) == 0 + assert len(daemon._active_subscriptions) == 0 diff --git a/tests/test_sd_service_instance.py b/tests/sd/test_sd_service_instance.py similarity index 100% rename from tests/test_sd_service_instance.py rename to tests/sd/test_sd_service_instance.py