Skip to content

Commit 2604039

Browse files
committed
test: harden conformance discovery, trait boundaries, and enum resilience
1 parent 4341d34 commit 2604039

3 files changed

Lines changed: 91 additions & 14 deletions

File tree

‎tests/conformance/discovery.py‎

Lines changed: 1 addition & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -21,10 +21,7 @@ def walk_modules(package: types.ModuleType) -> list[types.ModuleType]:
2121
if not hasattr(package, "__path__"):
2222
return modules
2323
for _, modname, _ in pkgutil.walk_packages(package.__path__, package.__name__ + "."):
24-
try:
25-
modules.append(importlib.import_module(modname))
26-
except (ImportError, AttributeError):
27-
continue
24+
modules.append(importlib.import_module(modname))
2825
return modules
2926

3027

‎tests/conformance/test_enum_conformance.py‎

Lines changed: 48 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -8,11 +8,14 @@
88

99
from __future__ import annotations
1010

11+
import enum
12+
import inspect
13+
1114
import pytest
1215

1316
import roborock
1417
from roborock.data.code_mappings import RoborockEnum, RoborockModeEnum
15-
from tests.conformance.discovery import discover_subclasses, to_pytest_params
18+
from tests.conformance.discovery import discover_subclasses, to_pytest_params, walk_modules
1619

1720
# Baseline inventory of legacy RoborockEnum classes that do not yet define an
1821
# explicit `unknown` member. When newer firmware emits an undocumented code,
@@ -87,21 +90,62 @@
8790
_ALL_MODE_ENUMS = discover_subclasses(roborock, RoborockModeEnum, exclude=(RoborockModeEnum,))
8891

8992

93+
# Enums in wire code mapping modules that intentionally remain standard Enum/StrEnum/IntEnum
94+
# (e.g., domain categories, product nicknames, or outgoing command identifiers).
95+
ALLOWED_NON_RESILIENT_CODE_MAPPING_ENUMS = {
96+
"roborock.data.b01_q10.b01_q10_code_mappings.RemoteCommand",
97+
"roborock.data.code_mappings.RoborockCategory",
98+
"roborock.data.code_mappings.RoborockProductNickname",
99+
"roborock.data.v1.v1_code_mappings.RoborockDockState",
100+
}
101+
102+
103+
def _discover_code_mapping_enums() -> list[type[enum.Enum]]:
104+
enums: list[type[enum.Enum]] = []
105+
for mod in walk_modules(roborock):
106+
if "code_mapping" in mod.__name__:
107+
for _, obj in inspect.getmembers(mod, inspect.isclass):
108+
if (
109+
obj.__module__ == mod.__name__
110+
and issubclass(obj, enum.Enum)
111+
and obj not in (enum.Enum, enum.IntEnum, enum.StrEnum, RoborockEnum, RoborockModeEnum)
112+
):
113+
enums.append(obj)
114+
return sorted(enums, key=lambda c: f"{c.__module__}.{c.__name__}")
115+
116+
117+
_ALL_CODE_MAPPING_ENUMS = _discover_code_mapping_enums()
118+
119+
90120
@pytest.mark.parametrize("enum_cls", to_pytest_params(_ALL_ROBOROCK_ENUMS, marks_by_fqn=_XFAIL_MARKS))
91121
def test_roborock_enum_has_unknown_fallback(enum_cls: type[RoborockEnum]) -> None:
92122
"""All RoborockEnum subclasses must define an explicit 'unknown' member."""
93123
assert hasattr(enum_cls, "unknown"), (
94124
f"{enum_cls.__module__}.{enum_cls.__name__} must define an 'unknown' member to prevent "
95125
"crashing or defaulting to arbitrary states on new firmware."
96126
)
97-
# Also verify that resolving an unknown int code returns the unknown member
98-
assert enum_cls(99999) == enum_cls.unknown
127+
# Derive an integer sentinel guaranteed to not exist in the enum
128+
sentinel = max(item.value for item in enum_cls) + 1 if list(enum_cls) else 99999
129+
assert enum_cls(sentinel) == enum_cls.unknown
99130

100131

101132
@pytest.mark.parametrize("mode_enum_cls", to_pytest_params(_ALL_MODE_ENUMS))
102133
def test_roborock_mode_enum_handles_unknown_code(mode_enum_cls: type[RoborockModeEnum]) -> None:
103134
"""RoborockModeEnum subclasses must return None when an unknown code is provided."""
104-
assert mode_enum_cls.from_code_optional(99999) is None
135+
sentinel = max(member.code for member in mode_enum_cls) + 1 if list(mode_enum_cls) else 99999
136+
assert mode_enum_cls.from_code_optional(sentinel) is None
137+
138+
139+
@pytest.mark.parametrize("enum_cls", to_pytest_params(_ALL_CODE_MAPPING_ENUMS))
140+
def test_code_mapping_enums_inherit_resilient_bases(enum_cls: type[enum.Enum]) -> None:
141+
"""Enums in wire code mapping modules must inherit from RoborockEnum or RoborockModeEnum."""
142+
fqn = f"{enum_cls.__module__}.{enum_cls.__name__}"
143+
if fqn in ALLOWED_NON_RESILIENT_CODE_MAPPING_ENUMS:
144+
return
145+
assert issubclass(enum_cls, (RoborockEnum, RoborockModeEnum)), (
146+
f"{fqn} in a code mapping module does not inherit from RoborockEnum or RoborockModeEnum. "
147+
"Wire status and error code enums must use resilient enum bases to handle unknown firmware codes."
148+
)
105149

106150

107151
def test_known_missing_unknown_baseline_inventory() -> None:

‎tests/conformance/test_trait_boundary_conformance.py‎

Lines changed: 42 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -7,13 +7,18 @@
77

88
from __future__ import annotations
99

10+
import asyncio
1011
import inspect
12+
import socket
13+
import typing
1114

1215
import pytest
1316

1417
import roborock.devices.traits
1518
from roborock.devices.traits import Trait
1619
from roborock.devices.traits.v1.common import V1TraitMixin
20+
from roborock.devices.transport.local_channel import LocalChannel, LocalChannelParams
21+
from roborock.devices.transport.mqtt_channel import MqttParams, MqttSession
1722
from tests.conformance.discovery import discover_classes, to_pytest_params
1823

1924
FORBIDDEN_PARAMS = {
@@ -27,8 +32,29 @@
2732
"aes_key",
2833
"password",
2934
"secret",
35+
"key",
3036
}
3137

38+
FORBIDDEN_TYPES = (
39+
socket.socket,
40+
asyncio.BaseTransport,
41+
LocalChannel,
42+
LocalChannelParams,
43+
MqttParams,
44+
MqttSession,
45+
)
46+
47+
48+
def _extract_types(annotation: typing.Any) -> set[type]:
49+
"""Recursively extract underlying concrete types from type hints and unions."""
50+
if isinstance(annotation, type):
51+
return {annotation}
52+
args = typing.get_args(annotation)
53+
types_found: set[type] = set()
54+
for arg in args:
55+
types_found |= _extract_types(arg)
56+
return types_found
57+
3258

3359
def _is_trait_class(cls: type) -> bool:
3460
if cls in (Trait, V1TraitMixin):
@@ -43,9 +69,19 @@ def _is_trait_class(cls: type) -> bool:
4369
def test_trait_constructor_does_not_leak_transport_or_credentials(trait_cls: type) -> None:
4470
"""Trait constructor parameters must never include transport sockets, IPs, or credentials."""
4571
sig = inspect.signature(trait_cls)
46-
param_names = set(sig.parameters.keys()) - {"self", "args", "kwargs"}
47-
violations = param_names & FORBIDDEN_PARAMS
48-
assert not violations, (
49-
f"{trait_cls.__module__}.{trait_cls.__name__}.__init__ accepts forbidden transport/credential "
50-
f"parameters {violations}. Per AGENTS.md, traits must receive only abstract channels or domain models."
51-
)
72+
for param_name, param in sig.parameters.items():
73+
if param_name in ("self", "args", "kwargs"):
74+
continue
75+
76+
assert param_name not in FORBIDDEN_PARAMS, (
77+
f"{trait_cls.__module__}.{trait_cls.__name__}.__init__ accepts forbidden transport/credential "
78+
f"parameter '{param_name}'. Per AGENTS.md, traits must receive only abstract channels or domain models."
79+
)
80+
81+
types_found = _extract_types(param.annotation)
82+
for t in types_found:
83+
assert not issubclass(t, FORBIDDEN_TYPES), (
84+
f"{trait_cls.__module__}.{trait_cls.__name__}.__init__ accepts parameter '{param_name}' "
85+
f"typed with forbidden low-level transport/credential class {t.__name__}. "
86+
"Traits must depend only on abstract channels (e.g. Channel, RpcChannel)."
87+
)

0 commit comments

Comments
 (0)