diff --git a/src/server/communication/message_data.py b/src/server/communication/message_data.py index c11bb52..445840d 100644 --- a/src/server/communication/message_data.py +++ b/src/server/communication/message_data.py @@ -2,6 +2,8 @@ import csv from dataclasses import dataclass +import numpy as np + from server.logger.log import logger @@ -9,20 +11,20 @@ class MessageData: topic: str payload: str - device_id: int - message_id: int - message_content: str - timestamp: str + device_id: str + message_id: str + message_content: dict + timestamp: float - received_timestamp = None - avg_speed = None - latency = None - synthetic_latency = None - payload_size = None - offloading_layer_index = None - layer_output = None - device_layers_inference_time = None - device_cpu_percent = None + received_timestamp: float | None = None + avg_speed: float | None = None + latency: float | None = None + synthetic_latency: float | None = None + payload_size: int | None = None + offloading_layer_index: int | None = None + layer_output: np.ndarray | None = None + device_layers_inference_time: np.ndarray | None = None + device_cpu_percent: float | None = None def to_dict(self): return self.__dict__ diff --git a/tests/unit/test_message_data.py b/tests/unit/test_message_data.py index ab73ef0..2afcd47 100644 --- a/tests/unit/test_message_data.py +++ b/tests/unit/test_message_data.py @@ -1,8 +1,50 @@ +import dataclasses from typing import get_type_hints from server.communication.message_data import MessageData +def test_core_field_annotations_match_runtime_types(): + annotations = get_type_hints(MessageData) + + assert annotations["device_id"] is str + assert annotations["message_id"] is str + assert annotations["message_content"] is dict + assert annotations["timestamp"] is float + + +def test_dynamically_populated_attributes_are_real_dataclass_fields(): + field_names = {f.name for f in dataclasses.fields(MessageData)} + + for attribute in ( + "received_timestamp", + "avg_speed", + "latency", + "synthetic_latency", + "payload_size", + "offloading_layer_index", + "layer_output", + "device_layers_inference_time", + "device_cpu_percent", + ): + assert attribute in field_names, f"{attribute} is not a dataclass field" + + +def test_optional_fields_default_to_none(): + message_data = MessageData( + topic="topic", + payload="", + device_id="device-1", + message_id="msg-1", + message_content={}, + timestamp=123.456, + ) + + assert message_data.received_timestamp is None + assert message_data.avg_speed is None + assert message_data.device_cpu_percent is None + + def test_get_latency_returns_duration_as_float(): latency = MessageData.get_latency("10.5", "12.75")