diff --git a/tests/scripts/thread-cert/Cert_6_1_02_REEDAttach_MED.py b/tests/scripts/thread-cert/Cert_6_1_02_REEDAttach_MED.py index bfa768ab5..8c07faf3b 100755 --- a/tests/scripts/thread-cert/Cert_6_1_02_REEDAttach_MED.py +++ b/tests/scripts/thread-cert/Cert_6_1_02_REEDAttach_MED.py @@ -83,8 +83,8 @@ class Cert_6_1_2_REEDAttach_MED(unittest.TestCase): self.assertEqual(self.nodes[REED].get_state(), 'child') self.nodes[MED].start() - - self.simulator.go(5) + + self.simulator.go(5) self.assertEqual(self.nodes[MED].get_state(), 'child') self.assertEqual(self.nodes[REED].get_state(), 'router') med_messages = self.simulator.get_messages_sent_by(MED) @@ -106,7 +106,7 @@ class Cert_6_1_2_REEDAttach_MED(unittest.TestCase): # Wait additional DEFAULT_CHILD_TIMEOUT to ensure the keep-alive message (child update request from MED) happens. self.simulator.go(config.DEFAULT_CHILD_TIMEOUT) med_messages = self.simulator.get_messages_sent_by(MED) - + # Step 8 - DUT sends Child Update messages msg = med_messages.next_mle_message(mle.CommandType.CHILD_UPDATE_REQUEST) check_child_update_request_from_child(msg, source_address=CheckType.CONTAIN, leader_data=CheckType.CONTAIN) diff --git a/tests/scripts/thread-cert/Cert_7_1_01_BorderRouterAsLeader.py b/tests/scripts/thread-cert/Cert_7_1_01_BorderRouterAsLeader.py index 8b1a64e0a..64a6715e0 100755 --- a/tests/scripts/thread-cert/Cert_7_1_01_BorderRouterAsLeader.py +++ b/tests/scripts/thread-cert/Cert_7_1_01_BorderRouterAsLeader.py @@ -31,12 +31,11 @@ import functools import time import unittest -from command import check_child_id_response -from command import check_child_update_response -from command import check_child_update_request_from_child -from command import check_data_response +from command import check_child_id_response, check_child_update_response, check_child_update_request_from_child, check_data_response from command import CheckType +from command import CommissioningDataCheck, NetworkDataCheck, PrefixesCheck, SinglePrefixCheck from command import NetworkDataCheckType + import config import mle import network_data @@ -135,22 +134,38 @@ class Cert_7_1_1_BorderRouterAsLeader(unittest.TestCase): # Step 2 - DUT creates network data msg = leader_messages.next_mle_message(mle.CommandType.DATA_RESPONSE) - check_data_response(msg, network_data_check=(NetworkDataCheckType.PREFIX_CONTENT, - [{network_data.TlvType.PREFIX:b'2001000200000001'}, {network_data.TlvType.PREFIX:b'2001000200000002'}])) + check_data_response(msg, + network_data_check=NetworkDataCheck( + prefixes_check=PrefixesCheck(prefix_check_list=[ SinglePrefixCheck(prefix=b'2001000200000001'), SinglePrefixCheck(prefix=b'2001000200000002') ]) + ) + ) # Step 4 - DUT sends a MLE Child ID Response to Router1 msg = leader_messages.next_mle_message(mle.CommandType.CHILD_ID_RESPONSE) - check_child_id_response(msg, network_data_check=(NetworkDataCheckType.PREFIX_CNT, 2)) + check_child_id_response(msg, + network_data_check=NetworkDataCheck( + prefixes_check=PrefixesCheck(prefix_cnt=2) + ) + ) # Step 6 - DUT sends a MLE Child ID Response to SED1 msg = leader_messages.next_mle_message(mle.CommandType.CHILD_ID_RESPONSE) - check_child_id_response(msg, network_data_check=(NetworkDataCheckType.PREFIX_CONTENT, [{network_data.TlvType.BORDER_ROUTER:0xFFFE}])) + check_child_id_response(msg, + network_data_check=NetworkDataCheck( + prefixes_check=PrefixesCheck(prefix_check_list=[ SinglePrefixCheck(border_router_16=0xFFFE) ]) + ) + ) + # For Step 10 msg_chd_upd_res_to_sed = leader_messages.next_mle_message(mle.CommandType.CHILD_UPDATE_RESPONSE) # Step 8 - DUT sends a MLE Child ID Response to MED1 msg = leader_messages.next_mle_message(mle.CommandType.CHILD_ID_RESPONSE) - check_child_id_response(msg, network_data_check=(NetworkDataCheckType.PREFIX_CNT, 2)) + check_child_id_response(msg, + network_data_check=NetworkDataCheck( + prefixes_check=PrefixesCheck(prefix_cnt=2) + ) + ) # Step 10 - DUT sends Child Update Response msg_chd_upd_res_to_med = leader_messages.next_mle_message(mle.CommandType.CHILD_UPDATE_RESPONSE) diff --git a/tests/scripts/thread-cert/Cert_7_1_03_BorderRouterAsLeader.py b/tests/scripts/thread-cert/Cert_7_1_03_BorderRouterAsLeader.py index 88beb7453..469687faa 100755 --- a/tests/scripts/thread-cert/Cert_7_1_03_BorderRouterAsLeader.py +++ b/tests/scripts/thread-cert/Cert_7_1_03_BorderRouterAsLeader.py @@ -30,10 +30,12 @@ import time import unittest +from command import check_child_update_request_from_child, check_child_update_request_from_parent, check_child_update_response, check_data_response from command import CheckType +from command import NetworkDataCheck, PrefixesCheck, SinglePrefixCheck from command import NetworkDataCheckType + import config -import command import mle import network_data import node @@ -136,37 +138,40 @@ class Cert_7_1_3_BorderRouterAsLeader(unittest.TestCase): # 3 - Leader msg = leader_messages.next_mle_message(mle.CommandType.DATA_RESPONSE) - command.check_data_response(msg, network_data_check=(NetworkDataCheckType.PREFIX_CONTENT, - [{network_data.TlvType.PREFIX:b'2001000200000001'}, {network_data.TlvType.PREFIX:b'2001000200000002'}])) + check_data_response(msg, + network_data_check=NetworkDataCheck( + prefixes_check=PrefixesCheck(prefix_check_list=[ SinglePrefixCheck(b'2001000200000001'), SinglePrefixCheck(b'2001000200000002')]) + ) + ) # 4 - N/A # Get addresses registered by MED1 msg = med1_messages.next_mle_message(mle.CommandType.CHILD_UPDATE_REQUEST) - command.check_child_update_request_from_child(msg, address_registration=CheckType.CONTAIN, CIDs=[0, 1, 2]) + check_child_update_request_from_child(msg, address_registration=CheckType.CONTAIN, CIDs=[0, 1, 2]) # 5 - Leader # Make a copy of leader's messages to ensure that we don't miss messages to SED1 leader_messages_copy = leader_messages.clone() msg = leader_messages_copy.next_mle_message(mle.CommandType.CHILD_UPDATE_RESPONSE, sent_to_node=self.nodes[MED1]) - command.check_child_update_response(msg, address_registration=CheckType.CONTAIN, CIDs=[1, 2]) + check_child_update_response(msg, address_registration=CheckType.CONTAIN, CIDs=[1, 2]) # 6A & 6B - Leader if config.LEADER_NOTIFY_SED_BY_CHILD_UPDATE_REQUEST: msg = leader_messages.next_mle_message(mle.CommandType.CHILD_UPDATE_REQUEST, sent_to_node=self.nodes[SED1]) - command.check_child_update_request_from_parent(msg, + check_child_update_request_from_parent(msg, leader_data=CheckType.CONTAIN, network_data=CheckType.CONTAIN, active_timestamp=CheckType.CONTAIN) else: msg = leader_messages.next_mle_message(mle.CommandType.DATA_RESPONSE, sent_to_node=self.nodes[SED1]) - command.check_data_response(msg, network_data=CheckType.CONTAIN, active_timestamp=CheckType.CONTAIN) + check_data_response(msg, network_data_check=command.NetworkDataCheck()) # 7 - N/A # Get addresses registered by SED1 msg = sed1_messages.next_mle_message(mle.CommandType.CHILD_UPDATE_REQUEST) - command.check_child_update_request_from_child(msg, address_registration=CheckType.CONTAIN, CIDs=[0, 1]) + check_child_update_request_from_child(msg, address_registration=CheckType.CONTAIN, CIDs=[0, 1]) # 8 - Leader msg = leader_messages.next_mle_message(mle.CommandType.CHILD_UPDATE_RESPONSE, sent_to_node=self.nodes[SED1]) - command.check_child_update_response(msg, address_registration=CheckType.CONTAIN, CIDs=[1]) + check_child_update_response(msg, address_registration=CheckType.CONTAIN, CIDs=[1]) if __name__ == '__main__': diff --git a/tests/scripts/thread-cert/Cert_9_2_02_MGMTCommissionerSet.py b/tests/scripts/thread-cert/Cert_9_2_02_MGMTCommissionerSet.py new file mode 100755 index 000000000..5ab114699 --- /dev/null +++ b/tests/scripts/thread-cert/Cert_9_2_02_MGMTCommissionerSet.py @@ -0,0 +1,172 @@ +#!/usr/bin/env python +# +# Copyright (c) 2019, The OpenThread Authors. +# All rights reserved. +# +# Redistribution and use in source and binary forms, with or without +# modification, are permitted provided that the following conditions are met: +# 1. Redistributions of source code must retain the above copyright +# notice, this list of conditions and the following disclaimer. +# 2. Redistributions in binary form must reproduce the above copyright +# notice, this list of conditions and the following disclaimer in the +# documentation and/or other materials provided with the distribution. +# 3. Neither the name of the copyright holder nor the +# names of its contributors may be used to endorse or promote products +# derived from this software without specific prior written permission. +# +# THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" +# AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE +# IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE +# ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE +# LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR +# CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF +# SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS +# INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN +# CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) +# ARISING IN ANY WAY OUT OF THE USE OF THIS SOFTWARE, EVEN IF ADVISED OF THE +# POSSIBILITY OF SUCH DAMAGE. +# + +from ipaddress import ip_address +import unittest + +from mle import NetworkData +from network_data import CommissioningData +import command +import config +import mesh_cop +import mle +import node + +COMMISSIONER = 1 +LEADER = 2 + + +class Cert_9_2_02_MGMTCommissionerSet(unittest.TestCase): + + def setUp(self): + self.simulator = config.create_default_simulator() + + self.nodes = {} + for i in range(1,3): + self.nodes[i] = node.Node(i, simulator=self.simulator) + + self.nodes[COMMISSIONER].set_panid(0xface) + self.nodes[COMMISSIONER].set_mode('rsdn') + self.nodes[COMMISSIONER].add_whitelist(self.nodes[LEADER].get_addr64()) + self.nodes[COMMISSIONER].enable_whitelist() + self.nodes[COMMISSIONER].set_router_selection_jitter(1) + + self.nodes[LEADER].set_panid(0xface) + self.nodes[LEADER].set_mode('rsdn') + self.nodes[LEADER].add_whitelist(self.nodes[COMMISSIONER].get_addr64()) + self.nodes[LEADER].enable_whitelist() + self.nodes[LEADER].set_router_selection_jitter(1) + + def tearDown(self): + for node in list(self.nodes.values()): + node.stop() + node.destroy() + self.simulator.stop() + + def test(self): + self.nodes[LEADER].start() + self.simulator.go(5) + self.assertEqual(self.nodes[LEADER].get_state(), 'leader') + + self.nodes[COMMISSIONER].start() + self.simulator.go(5) + self.assertEqual(self.nodes[COMMISSIONER].get_state(), 'router') + + # Skip all other Coaps sent by Leader + self.simulator.get_messages_sent_by(COMMISSIONER) + self.simulator.get_messages_sent_by(LEADER) + + # Commissioner start + self.nodes[COMMISSIONER].commissioner_start() + self.simulator.go(3) + self.simulator.get_messages_sent_by(COMMISSIONER) # Skip LEAD_PET.req + + # Get CommissionerSesssionId from LEAD_PET.rsp + leader_messages = self.simulator.get_messages_sent_by(LEADER) + msg = leader_messages.next_coap_message('2.04', assert_enabled=True) + commissioner_session_id_tlv = command.get_sub_tlv(msg.coap.payload, mesh_cop.CommissionerSessionId) + + # Step 2 - Harness instructs commissioner to send MGMT_COMMISSIONER_SET.req to Leader + steering_data_tlv = mesh_cop.SteeringData(bytes([0xFF])) + self.nodes[COMMISSIONER].commissioner_mgmtset_with_tlvs([steering_data_tlv]) + self.simulator.go(5) + + # Step 3 - Leader responds to MGMT_COMMISSIONER_SET.req with MGMT_COMMISSIONER_SET.rsp + leader_messages = self.simulator.get_messages_sent_by(LEADER) + msg = leader_messages.next_coap_message('2.04') + command.check_coap_message(msg, [mesh_cop.State(mesh_cop.MeshCopState.REJECT)]) # (mesh_cop.State(mesh_cop.MeshCopState.REJECT),) <- this a tuple, don't delete the comma + self.simulator.get_messages_sent_by(COMMISSIONER) # Skip LEAD_PET.req + + # Step 4 - Harness instructs commissioner to send MGMT_COMMISSIONER_SET.req to Leader + self.nodes[COMMISSIONER].commissioner_mgmtset_with_tlvs([steering_data_tlv, commissioner_session_id_tlv]) + self.simulator.go(5) + commissioner_messages = self.simulator.get_messages_sent_by(COMMISSIONER) + msg = commissioner_messages.next_coap_message('0.02', uri_path='/c/cs') + rloc = ip_address(self.nodes[LEADER].get_addr_rloc()) + leader_aloc = ip_address(self.nodes[LEADER].get_addr_leader_aloc()) + command.check_coap_message(msg, [steering_data_tlv, commissioner_session_id_tlv], dest_addrs=[rloc, leader_aloc]) + + # Step 5 - Leader sends MGMT_COMMISSIONER_SET.rsp to commissioner + leader_messages = self.simulator.get_messages_sent_by(LEADER) + msg = leader_messages.next_coap_message('2.04') + command.check_coap_message(msg, [mesh_cop.State(mesh_cop.MeshCopState.ACCEPT)]) + + # Step 6 - Leader sends a multicast MLE Data Response + msg = leader_messages.next_mle_message(mle.CommandType.DATA_RESPONSE) + command.check_data_response(msg, command.NetworkDataCheck( + commissioning_data_check=command.CommissioningDataCheck(stable=0, sub_tlv_type_list=[mesh_cop.CommissionerSessionId, mesh_cop.SteeringData, mesh_cop.BorderAgentLocator]))) + + # Step 7 - Harness instructs commissioner to send MGMT_COMMISSIONER_SET.req to Leader + border_agent_locator_tlv = mesh_cop.BorderAgentLocator(0x0400) + self.nodes[COMMISSIONER].commissioner_mgmtset_with_tlvs( + [commissioner_session_id_tlv, border_agent_locator_tlv]) + self.simulator.go(5) + + # Step 8 - Leader responds to MGMT_COMMISSIONER_SET.req with MGMT_COMMISSIONER_SET.rsp + leader_messages = self.simulator.get_messages_sent_by(LEADER) + msg = leader_messages.next_coap_message('2.04') + command.check_coap_message(msg, [mesh_cop.State(mesh_cop.MeshCopState.REJECT)]) + + # Step 9 - Harness instructs commissioner to send MGMT_COMMISSIONER_SET.req to Leader + self.nodes[COMMISSIONER].commissioner_mgmtset_with_tlvs( + [steering_data_tlv, commissioner_session_id_tlv, border_agent_locator_tlv]) + self.simulator.go(5) + + # Step 10 - Leader responds to MGMT_COMMISSIONER_SET.req with MGMT_COMMISSIONER_SET.rsp + leader_messages = self.simulator.get_messages_sent_by(LEADER) + msg = leader_messages.next_coap_message('2.04') + command.check_coap_message(msg, [mesh_cop.State(mesh_cop.MeshCopState.REJECT)]) + + # Step 11 - Harness instructs commissioner to send MGMT_COMMISSIONER_SET.req to Leader + self.nodes[COMMISSIONER].commissioner_mgmtset_with_tlvs( + [mesh_cop.CommissionerSessionId(0xFFFF), steering_data_tlv]) + self.simulator.go(5) + + # Step 12 - Leader responds to MGMT_COMMISSIONER_SET.req with MGMT_COMMISSIONER_SET.rsp + leader_messages = self.simulator.get_messages_sent_by(LEADER) + msg = leader_messages.next_coap_message('2.04') + command.check_coap_message(msg, [mesh_cop.State(mesh_cop.MeshCopState.REJECT)]) + + # Step 13 - Harness instructs commissioner to send MGMT_COMMISSIONER_SET.req to Leader + self.nodes[COMMISSIONER].commissioner_mgmtset_with_tlvs( + [commissioner_session_id_tlv, steering_data_tlv, mesh_cop.Channel(0x0, 0x0)]) + self.simulator.go(5) + + # Step 14 - Leader responds to MGMT_COMMISSIONER_SET.req with MGMT_COMMISSIONER_SET.rsp + leader_messages = self.simulator.get_messages_sent_by(LEADER) + msg = leader_messages.next_coap_message('2.04') + command.check_coap_message(msg, [mesh_cop.State(mesh_cop.MeshCopState.ACCEPT)]) + + # Step 15 - Send ICMPv6 Echo Request to Leader + leader_rloc = self.nodes[LEADER].get_addr_rloc() + self.assertTrue(self.nodes[COMMISSIONER].ping(leader_rloc)) + + +if __name__ == '__main__': + unittest.main() diff --git a/tests/scripts/thread-cert/Makefile.am b/tests/scripts/thread-cert/Makefile.am index f1c39271a..c5bf54817 100644 --- a/tests/scripts/thread-cert/Makefile.am +++ b/tests/scripts/thread-cert/Makefile.am @@ -108,6 +108,7 @@ EXTRA_DIST = \ Cert_8_1_02_Commissioning.py \ Cert_8_2_01_JoinerRouter.py \ Cert_8_2_02_JoinerRouter.py \ + Cert_9_2_02_MGMTCommissionerSet.py \ Cert_9_2_04_ActiveDataset.py \ Cert_9_2_07_DelayTimer.py \ Cert_9_2_08_PersistentDatasets.py \ @@ -153,6 +154,7 @@ EXTRA_DIST = \ test_service.py \ test_network_data.py \ test_network_layer.py \ + tlvs_parsing.py \ $(NULL) check_PROGRAMS = \ @@ -247,6 +249,7 @@ check_SCRIPTS = \ Cert_8_1_02_Commissioning.py \ Cert_8_2_01_JoinerRouter.py \ Cert_8_2_02_JoinerRouter.py \ + Cert_9_2_02_MGMTCommissionerSet.py \ Cert_9_2_04_ActiveDataset.py \ Cert_9_2_07_DelayTimer.py \ Cert_9_2_08_PersistentDatasets.py \ @@ -288,6 +291,7 @@ XFAIL_NCP_TESTS = \ Cert_8_1_02_Commissioning.py \ Cert_8_2_01_JoinerRouter.py \ Cert_8_2_02_JoinerRouter.py \ + Cert_9_2_02_MGMTCommissionerSet.py \ Cert_9_2_04_ActiveDataset.py \ Cert_9_2_07_DelayTimer.py \ Cert_9_2_08_PersistentDatasets.py \ diff --git a/tests/scripts/thread-cert/coap.py b/tests/scripts/thread-cert/coap.py index 7fa3cc4f3..43e46b67f 100644 --- a/tests/scripts/thread-cert/coap.py +++ b/tests/scripts/thread-cert/coap.py @@ -257,10 +257,10 @@ class CoapMessage(object): class CoapMessageProxy(object): - """ Proxy class of CoAP message. + """ Proxy class of CoAP message. - The main idea behind this class is to delay parsing payload. Due to architecture of the existing solution - it is possible to process confirmation message before a request message. In such case it is not possible + The main idea behind this class is to delay parsing payload. Due to architecture of the existing solution + it is possible to process confirmation message before a request message. In such case it is not possible to get URI path to get proper payload parser. """ diff --git a/tests/scripts/thread-cert/command.py b/tests/scripts/thread-cert/command.py index 5be88973e..d986acf94 100644 --- a/tests/scripts/thread-cert/command.py +++ b/tests/scripts/thread-cert/command.py @@ -273,7 +273,6 @@ def check_parent_request(command_msg, is_first_request): elif not scan_mask.end_device: raise ValueError("Second parent request without E bit set") - def check_parent_response(command_msg, mle_frame_counter = CheckType.OPTIONAL): """Verify a properly formatted Parent Response command message. """ @@ -313,41 +312,6 @@ def check_child_id_request(command_msg, tlv_request = CheckType.OPTIONAL, \ check_tlv_request_tlv(command_msg, CheckType.CONTAIN, mle.TlvType.ADDRESS16) check_tlv_request_tlv(command_msg, CheckType.CONTAIN, mle.TlvType.NETWORK_DATA) -def find_prefix_tlv(tlvs, cond_map): - """Find a prefix tlv in tlvs which matchs some conditions specified by cond_map - """ - for tlv in tlvs: - if network_data.TlvType.PREFIX in cond_map: - if binascii.hexlify(tlv.prefix) != cond_map[network_data.TlvType.PREFIX]: - continue - if network_data.TlvType.BORDER_ROUTER in cond_map: - border_router_tlv = get_sub_tlv(tlv.sub_tlvs, network_data.BorderRouter) - if border_router_tlv.border_router_16 != cond_map[network_data.TlvType.BORDER_ROUTER]: - continue - return tlv - return None - -def check_network_data(data, check_detail): - check_type = check_detail[0] - prefixes = [tlv for tlv in data.tlvs if isinstance(tlv, network_data.Prefix)] - if check_type == NetworkDataCheckType.PREFIX_CNT: - # check_detail[1] should be a integer number representing the minimum count of prefixes should be - min_cnt = check_detail[1] - assert len(prefixes) >= min_cnt, 'Network data should contain at least {} prefixes'.format(mn_cnt) - for prefix in prefixes: - check_prefix(prefix) - elif check_type == NetworkDataCheckType.PREFIX_CONTENT: - # check_detail[1] should be a list of dictionary(like - # [{network_data.TlvType.PREFIX='...', network_data.TlvType.BORDER_ROUTER='...'}]) - # each entry of the dictionary represents one thing to check of the prefix tlv - assert len(prefixes) >= len(check_detail[1]), 'Network data seems to have less prefixes than expected' - # basic check of prefixes - for prefix in prefixes: - check_prefix(prefix) - for cond_map in check_detail[1]: - tlv = find_prefix_tlv(prefixes, cond_map) - assert tlv is not None, 'Some prefix sub-tlv is not found:{}'.format(cond_map) - def check_child_id_response(command_msg, route64 = CheckType.OPTIONAL, network_data = CheckType.OPTIONAL, \ address_registration = CheckType.OPTIONAL, active_timestamp = CheckType.OPTIONAL, \ pending_timestamp = CheckType.OPTIONAL, active_operational_dataset = CheckType.OPTIONAL, \ @@ -367,9 +331,9 @@ def check_child_id_response(command_msg, route64 = CheckType.OPTIONAL, network_d check_mle_optional_tlv(command_msg, active_operational_dataset, mle.ActiveOperationalDataset) check_mle_optional_tlv(command_msg, pending_operational_dataset, mle.PendingOperationalDataset) - if network_data_check != None: + if network_data_check is not None: network_data_tlv = command_msg.assertMleMessageContainsTlv(mle.NetworkData) - check_network_data(network_data_tlv, network_data_check) + network_data_check.check(network_data_tlv) def check_prefix(prefix): """Verify if a prefix contains 6loWPAN sub-TLV and border router sub-TLV @@ -418,27 +382,27 @@ def contains_tlv(sub_tlvs, tlv_type): """ return any(isinstance(sub_tlv, tlv_type) for sub_tlv in sub_tlvs) +def contains_tlvs(sub_tlvs, tlv_types): + """Verify if all types of tlv in a list are included in a sub-tlv list. + """ + return all((any(isinstance(sub_tlv, tlv_type) for sub_tlv in sub_tlvs)) for tlv_type in tlv_types) + def check_secure_mle_key_id_mode(command_msg, key_id_mode): """Verify if the mle command message sets the right key id mode. """ assert isinstance(command_msg.mle, mle.MleMessageSecured) assert command_msg.mle.aux_sec_hdr.key_id_mode == key_id_mode -def check_data_response(command_msg, network_data_opt=CheckType.OPTIONAL, - active_timestamp=CheckType.OPTIONAL, - network_data_check=None): +def check_data_response(command_msg, network_data_check=None, active_timestamp=CheckType.OPTIONAL): """Verify a properly formatted Data Response command message. """ check_secure_mle_key_id_mode(command_msg, 0x02) - command_msg.assertMleMessageContainsTlv(mle.SourceAddress) command_msg.assertMleMessageContainsTlv(mle.LeaderData) - check_mle_optional_tlv(command_msg, network_data_opt, mle.NetworkData) check_mle_optional_tlv(command_msg, active_timestamp, mle.ActiveTimestamp) - - if network_data_check != None: + if network_data_check is not None: network_data_tlv = command_msg.assertMleMessageContainsTlv(mle.NetworkData) - check_network_data(network_data_tlv, network_data_check) + network_data_check.check(network_data_tlv) def check_child_update_request_from_parent(command_msg, leader_data=CheckType.OPTIONAL, network_data=CheckType.OPTIONAL, challenge=CheckType.OPTIONAL, @@ -532,11 +496,11 @@ def check_discovery_response(command_msg, request_src_addr, steering_data=CheckT assert response.version == config.PROTOCOL_VERSION assert_contains_tlv(tlvs, CheckType.CONTAIN, mesh_cop.ExtendedPanid) assert_contains_tlv(tlvs, CheckType.CONTAIN, mesh_cop.NetworkName) - assert_contains_tlv(tlvs, steering_data, network_data.SteeringData) + assert_contains_tlv(tlvs, steering_data, mesh_cop.SteeringData) assert_contains_tlv(tlvs, steering_data, mesh_cop.JoinerUdpPort) check_type = CheckType.CONTAIN if response.native_flag else CheckType.OPTIONAL - assert_contains_tlv(tlvs, check_type, network_data.CommissionerUdpPort) + assert_contains_tlv(tlvs, check_type, mesh_cop.CommissionerUdpPort) def get_joiner_udp_port_in_discovery_response(command_msg): """Get the udp port specified in a DISCOVERY RESPONSE message @@ -565,3 +529,87 @@ def check_joiner_router_commissioning_messages(commissioning_messages): """Verify COAP messages sent by joiner router while commissioning process. """ assert any(msg.type == mesh_cop.MeshCopMessageType.JOIN_ENT_NTF for msg in commissioning_messages) + return None + +def check_payload_same(tp1, tp2): + """Verfiy two payloads are totally the same. + A payload is a tuple of tlvs. + """ + assert len(tp1) == len(tp2) + for tlv in tp2: + peer_tlv = get_sub_tlv(tp1, type(tlv)) + assert peer_tlv is not None and peer_tlv == tlv, 'peer_tlv:{}, tlv:{} type:{}'.format(peer_tlv, tlv, type(tlv)) + +def check_coap_message(msg, payloads, dest_addrs=None): + if dest_addrs is not None: + found = False + for dest in dest_addrs: + if msg.ipv6_packet.ipv6_header.destination_address == dest: + found = True + break + assert found, 'Destination address incorrect' + check_payload_same(msg.coap.payload, payloads) + +class SinglePrefixCheck: + + def __init__(self, prefix=None, border_router_16=None): + self._prefix = prefix + self._border_router_16 = border_router_16 + + def check(self, prefix_tlv): + border_router_tlv = assert_contains_tlv(prefix_tlv.sub_tlvs, CheckType.CONTAIN, network_data.BorderRouter) + lowpan_id_tlv = assert_contains_tlv(prefix_tlv.sub_tlvs, CheckType.CONTAIN, network_data.LowpanId) + result = True + if self._prefix is not None: + result &= (self._prefix == binascii.hexlify(prefix_tlv.prefix)) + if self._border_router_16 is not None: + result &= (self._border_router_16 == border_router_tlv.border_router_16) + return result + + +class PrefixesCheck: + + def __init__(self, prefix_cnt=0, prefix_check_list=[]): + self._prefix_cnt = prefix_cnt + self._prefix_check_list = prefix_check_list + + def check(self, prefix_tlvs): + # if prefix_cnt is given, then check count only + if self._prefix_cnt > 0: + assert len(prefix_tlvs) >= self._prefix_cnt, 'prefix count is less than expected' + else: + for prefix_check in self._prefix_check_list: + found = False + for prefix_tlv in prefix_tlvs: + if prefix_check.check(prefix_tlv): + found = True + break + assert found, 'Some prefix is absent: {}'.format(prefix_check) + + +class CommissioningDataCheck: + + def __init__(self, stable=None, sub_tlv_type_list=[]): + self._stable = stable + self._sub_tlv_type_list = sub_tlv_type_list + + def check(self, commissioning_data_tlv): + if self._stable is not None: + assert self._stable == commissioning_data_tlv.stable, 'Commissioning Data stable flag is not correct' + assert contains_tlvs(commissioning_data_tlv.sub_tlvs, self._sub_tlv_type_list), 'Some sub tlvs are missing in Commissioning Data' + + +class NetworkDataCheck: + + def __init__(self, prefixes_check=None, commissioning_data_check=None): + self._prefixes_check = prefixes_check + self._commissioning_data_check = commissioning_data_check + + def check(self, network_data_tlv): + if self._prefixes_check is not None: + prefix_tlvs = [tlv for tlv in network_data_tlv.tlvs if isinstance(tlv, network_data.Prefix)] + self._prefixes_check.check(prefix_tlvs) + if self._commissioning_data_check is not None: + commissioning_data_tlv = assert_contains_tlv(network_data_tlv.tlvs, CheckType.CONTAIN, network_data.CommissioningData) + self._commissioning_data_check.check(commissioning_data_tlv) + diff --git a/tests/scripts/thread-cert/config.py b/tests/scripts/thread-cert/config.py index a4dbd26e7..dd9323516 100644 --- a/tests/scripts/thread-cert/config.py +++ b/tests/scripts/thread-cert/config.py @@ -26,9 +26,10 @@ # ARISING IN ANY WAY OUT OF THE USE OF THIS SOFTWARE, EVEN IF ADVISED OF THE # POSSIBILITY OF SUCH DAMAGE. # -import os from enum import Enum +import os +from tlvs_parsing import SubTlvsFactory import coap import dtls import ipv6 @@ -108,10 +109,11 @@ def create_default_network_data_service_sub_tlvs_factory(): def create_default_network_data_commissioning_data_sub_tlvs_factories(): return { - network_data.MeshcopTlvType.STEERING_DATA: network_data.SteeringDataFactory(), - network_data.MeshcopTlvType.BORDER_AGENT_LOCATOR: network_data.BorderAgentLocatorFactory(), - network_data.MeshcopTlvType.COMMISSIONER_SESSION_ID: network_data.CommissionerSessionIdFactory(), - network_data.MeshcopTlvType.COMMISSIONER_UDP_PORT: network_data.CommissionerUdpPortFactory(), + mesh_cop.TlvType.CHANNEL: mesh_cop.ChannelFactory(), + mesh_cop.TlvType.STEERING_DATA: mesh_cop.SteeringDataFactory(), + mesh_cop.TlvType.BORDER_AGENT_LOCATOR: mesh_cop.BorderAgentLocatorFactory(), + mesh_cop.TlvType.COMMISSIONER_SESSION_ID: mesh_cop.CommissionerSessionIdFactory(), + mesh_cop.TlvType.COMMISSIONER_UDP_PORT: mesh_cop.CommissionerUdpPortFactory(), } @@ -168,9 +170,9 @@ def create_default_thread_discovery_sub_tlvs_factories(): mesh_cop.TlvType.DISCOVERY_RESPONSE: mesh_cop.DiscoveryResponseFactory(), mesh_cop.TlvType.EXTENDED_PANID: mesh_cop.ExtendedPanidFactory(), mesh_cop.TlvType.NETWORK_NAME: mesh_cop.NetworkNameFactory(), - mesh_cop.TlvType.STEERING_DATA: network_data.SteeringDataFactory(), + mesh_cop.TlvType.STEERING_DATA: mesh_cop.SteeringDataFactory(), mesh_cop.TlvType.JOINER_UDP_PORT: mesh_cop.JoinerUdpPortFactory(), - mesh_cop.TlvType.COMMISSIONER_UDP_PORT: network_data.CommissionerUdpPortFactory() + mesh_cop.TlvType.COMMISSIONER_UDP_PORT: mesh_cop.CommissionerUdpPortFactory() } def create_default_mle_tlvs_factories(): @@ -238,20 +240,64 @@ def create_deafult_network_tlvs_factories(): def create_default_network_tlvs_factory(): - return network_layer.NetworkLayerTlvsFactory( - tlvs_factories=create_deafult_network_tlvs_factories()) + return SubTlvsFactory(sub_tlvs_factories=create_deafult_network_tlvs_factories()) +def create_default_mesh_cop_tlvs_factories(): + return { + mesh_cop.TlvType.CHANNEL: mesh_cop.ChannelFactory(), + mesh_cop.TlvType.PAN_ID: mesh_cop.PanidFactory(), + mesh_cop.TlvType.EXTENDED_PANID: mesh_cop.ExtendedPanidFactory(), + mesh_cop.TlvType.NETWORK_NAME: mesh_cop.NetworkNameFactory(), + mesh_cop.TlvType.PSKC: mesh_cop.PSKcFactory(), + mesh_cop.TlvType.NETWORK_MASTER_KEY: mesh_cop.NetworkMasterKeyFactory(), + mesh_cop.TlvType.NETWORK_KEY_SEQUENCE_COUNTER: mesh_cop.NetworkKeySequenceCounterFactory(), + mesh_cop.TlvType.NETWORK_MESH_LOCAL_PREFIX: mesh_cop.NetworkMeshLocalPrefixFactory(), + mesh_cop.TlvType.STEERING_DATA: mesh_cop.SteeringDataFactory(), + mesh_cop.TlvType.BORDER_AGENT_LOCATOR: mesh_cop.BorderAgentLocatorFactory(), + mesh_cop.TlvType.COMMISSIONER_ID: mesh_cop.CommissionerIdFactory(), + mesh_cop.TlvType.COMMISSIONER_SESSION_ID: mesh_cop.CommissionerSessionIdFactory(), + mesh_cop.TlvType.SECURITY_POLICY: mesh_cop.SecurityPolicyFactory(), + mesh_cop.TlvType.GET: mesh_cop.GetFactory(), + mesh_cop.TlvType.ACTIVE_TIMESTAMP: mesh_cop.ActiveTimestampFactory(), + mesh_cop.TlvType.COMMISSIONER_UDP_PORT: mesh_cop.CommissionerUdpPortFactory(), + mesh_cop.TlvType.STATE: mesh_cop.StateFactory(), + mesh_cop.TlvType.JOINER_DTLS_ENCAPSULATION: mesh_cop.JoinerDtlsEncapsulationFactory(), + mesh_cop.TlvType.JOINER_UDP_PORT: mesh_cop.JoinerUdpPortFactory(), + mesh_cop.TlvType.JOINER_IID: mesh_cop.JoinerIIDFactory(), + mesh_cop.TlvType.JOINER_ROUTER_LOCATOR: mesh_cop.JoinerRouterLocatorFactory(), + mesh_cop.TlvType.JOINER_ROUTER_KEK: mesh_cop.JoinerRouterKEKFactory(), + mesh_cop.TlvType.PROVISIONING_URL: mesh_cop.ProvisioningUrlFactory(), + mesh_cop.TlvType.VENDOR_NAME: mesh_cop.VendorNameFactory(), + mesh_cop.TlvType.VENDOR_MODEL: mesh_cop.VendorModelFactory(), + mesh_cop.TlvType.VENDOR_SW_VERSION: mesh_cop.VendorSWVersionFactory(), + mesh_cop.TlvType.VENDOR_DATA: mesh_cop.VendorDataFactory(), + mesh_cop.TlvType.VENDOR_STACK_VERSION: mesh_cop.VendorStackVersionFactory(), + mesh_cop.TlvType.UDP_ENCAPSULATION: mesh_cop.UdpEncapsulationFactory(), + mesh_cop.TlvType.IPV6_ADDRESS: mesh_cop.Ipv6AddressFactory(), + mesh_cop.TlvType.PENDING_TIMESTAMP: mesh_cop.PendingTimestampFactory(), + mesh_cop.TlvType.DELAY_TIMER: mesh_cop.DelayTimerFactory(), + mesh_cop.TlvType.CHANNEL_MASK: mesh_cop.ChannelMaskFactory(), + mesh_cop.TlvType.COUNT: mesh_cop.CountFactory(), + mesh_cop.TlvType.PERIOD: mesh_cop.PeriodFactory(), + mesh_cop.TlvType.SCAN_DURATION: mesh_cop.ScanDurationFactory(), + mesh_cop.TlvType.ENERGY_LIST: mesh_cop.EnergyListFactory() + } + +def create_default_mesh_cop_tlvs_factory(): + return SubTlvsFactory(sub_tlvs_factories=create_default_mesh_cop_tlvs_factories()) def create_default_uri_path_based_payload_factories(): network_layer_tlvs_factory = create_default_network_tlvs_factory() - + mesh_cop_tlvs_factory = create_default_mesh_cop_tlvs_factory() return { "/a/as": network_layer_tlvs_factory, "/a/aq": network_layer_tlvs_factory, "/a/ar": network_layer_tlvs_factory, "/a/ae": network_layer_tlvs_factory, "/a/an": network_layer_tlvs_factory, - "/a/sd": network_layer_tlvs_factory + "/a/sd": network_layer_tlvs_factory, + "/c/lp": mesh_cop_tlvs_factory, + "/c/cs": mesh_cop_tlvs_factory } diff --git a/tests/scripts/thread-cert/mesh_cop.py b/tests/scripts/thread-cert/mesh_cop.py index 2485b10b1..250bb4405 100644 --- a/tests/scripts/thread-cert/mesh_cop.py +++ b/tests/scripts/thread-cert/mesh_cop.py @@ -27,33 +27,65 @@ # POSSIBILITY OF SUCH DAMAGE. # +from binascii import hexlify from enum import IntEnum import io import struct from network_data import SubTlvsFactory +import common class TlvType(IntEnum): + CHANNEL = 0 + PAN_ID = 1 EXTENDED_PANID = 2 NETWORK_NAME = 3 + PSKC = 4 + NETWORK_MASTER_KEY = 5 + NETWORK_KEY_SEQUENCE_COUNTER = 6 + NETWORK_MESH_LOCAL_PREFIX = 7 STEERING_DATA = 8 + BORDER_AGENT_LOCATOR = 9 + COMMISSIONER_ID = 10 + COMMISSIONER_SESSION_ID = 11 + SECURITY_POLICY = 12 + GET = 13 + ACTIVE_TIMESTAMP = 14 COMMISSIONER_UDP_PORT = 15 STATE = 16 + JOINER_DTLS_ENCAPSULATION = 17 JOINER_UDP_PORT = 18 + JOINER_IID = 19 + JOINER_ROUTER_LOCATOR = 20 + JOINER_ROUTER_KEK = 21 PROVISIONING_URL = 32 VENDOR_NAME = 33 VENDOR_MODEL = 34 VENDOR_SW_VERSION = 35 VENDOR_DATA = 36 VENDOR_STACK_VERSION = 37 + UDP_ENCAPSULATION = 48 + IPV6_ADDRESS = 49 + PENDING_TIMESTAMP = 51 + DELAY_TIMER = 52 + CHANNEL_MASK = 53 + COUNT = 54 + PERIOD = 55 + SCAN_DURATION = 56 + ENERGY_LIST = 57 DISCOVERY_REQUEST = 128 DISCOVERY_RESPONSE = 129 +class MeshCopState(IntEnum): + ACCEPT = 0x1 + REJECT = 0xFF + + class MeshCopMessageType(IntEnum): - JOIN_FIN_REQ = 1 - JOIN_FIN_RSP = 2 - JOIN_ENT_NTF = 3 + JOIN_FIN_REQ = 1, + JOIN_FIN_RSP = 2, + JOIN_ENT_NTF = 3, JOIN_ENT_RSP = 4 @@ -64,6 +96,335 @@ def create_mesh_cop_message_type_set(): MeshCopMessageType.JOIN_ENT_RSP ] +# Channel TLV (0) +class Channel(object): + + def __init__(self, channel_page, channel): + self._channel_page = channel_page + self._channel = channel + + @property + def channel_page(self): + return self._channel_page + + @property + def channel(self): + return self._channel + + def __eq__(self, other): + common.expect_the_same_class(self, other) + + return self._channel_page == other._channel_page and self._channel == other.__channel + + def __repr__(self): + return 'Channel(channel_page={},channel={})'.format(self._channel_page, self._channel) + + def to_hex(self): + return struct.pack('>BBBH', TlvType.CHANNEL, 3, self.channel_page, self.channel) + + +class ChannelFactory(object): + + def parse(self, data, message_info): + data_tp = struct.unpack('>BH', data.read(3)) + channel_page = data_tp[0] + channel = data_tp[1] + return Channel(channel_page, channel) + + +# PanId TLV (1) +class Panid(object): + # TODO: Not implemented yet + pass + + +class PanidFactory(object): + # TODO: Not implemented yet + + def parse(self, data, message_info): + raise NotImplementedError("TODO: Not implemented yet") + + +# ExtendedPanid TLV (2) +class ExtendedPanid(object): + + def __init__(self, extended_panid): + self._extended_panid = extended_panid + + @property + def extended_panid(self): + return self._extended_panid + + def __eq__(self, other): + return (type(self) is type(other) + and self.extended_panid == other.extended_panid) + + def __repr__(self): + return "ExtendedPanid(extended_panid={})".format(self.extended_panid) + + +class ExtendedPanidFactory(object): + + def parse(self, data, message_info): + extended_panid = struct.unpack(">Q", data.read(8))[0] + return ExtendedPanid(extended_panid) + + +# NetworkName TLV (3) +class NetworkName(object): + + def __init__(self, network_name): + self._network_name = network_name + + @property + def network_name(self): + return self._network_name + + def __eq__(self, other): + return (type(self) is type(other) + and self.network_name == other.network_name) + + def __repr__(self): + return "NetworkName(network_name={})".format(self.network_name) + + +class NetworkNameFactory(object): + + def parse(self, data, message_info): + len = message_info.length + network_name = struct.unpack("{}s".format(10), data.read(len))[0] + return NetworkName(network_name) + + +# PSKc TLV (4) +class PSKc(object): + # TODO: Not implemented yet + pass + + +class PSKcFactory(object): + # TODO: Not implemented yet + + def parse(self, data, message_info): + raise NotImplementedError("TODO: Not implemented yet") + + +# NetworkMasterKey TLV (5) +class NetworkMasterKey(object): + # TODO: Not implemented yet + pass + + +class NetworkMasterKeyFactory(object): + # TODO: Not implemented yet + + def parse(self, data, message_info): + raise NotImplementedError("TODO: Not implemented yet") + + +# NetworkKeySequenceCounter TLV (6) +class NetworkKeySequenceCounter(object): + # TODO: Not implemented yet + pass + + +class NetworkKeySequenceCounterFactory(object): + # TODO: Not implemented yet + + def parse(self, data, message_info): + raise NotImplementedError("TODO: Not implemented yet") + + +# NetworkMeshLocalPrefix TLV (7) +class NetworkMeshLocalPrefix(object): + # TODO: Not implemented yet + pass + + +class NetworkMeshLocalPrefixFactory(object): + # TODO: Not implemented yet + + def parse(self, data, message_info): + raise NotImplementedError("TODO: Not implemented yet") + + +# Steering Data TLV (8) +class SteeringData(object): + + def __init__(self, bloom_filter): + self._bloom_filter = bloom_filter + + @property + def bloom_filter(self): + return self._bloom_filter + + def __eq__(self, other): + common.expect_the_same_class(self, other) + + return self._bloom_filter == other._bloom_filter + + def __repr__(self): + return "SteeringData(bloom_filter={})".format(hexlify(self._bloom_filter)) + + def to_hex(self): + bloom_filter_len = len(self.bloom_filter) + return struct.pack('>BB', TlvType.STEERING_DATA, bloom_filter_len) + self.bloom_filter + + +class SteeringDataFactory: + + def parse(self, data, message_info): + bloom_filter = data.read(message_info.length) + return SteeringData(bloom_filter) + + +# Border Agent Locator TLV (9) +class BorderAgentLocator(object): + + def __init__(self, address): + self._border_agent_locator = address + + @property + def border_agent_locator(self): + return self._border_agent_locator + + def __eq__(self, other): + common.expect_the_same_class(self, other) + + return self._border_agent_locator == other._border_agent_locator + + def __repr__(self): + return "BorderAgentLocator(rloc16={})".format(hex(self._border_agent_locator)) + + def to_hex(self): + return struct.pack('>BBH', TlvType.BORDER_AGENT_LOCATOR, 2, self.border_agent_locator) + + +class BorderAgentLocatorFactory: + + def parse(self, data, message_info): + border_agent_locator = struct.unpack(">H", data.read(2))[0] + return BorderAgentLocator(border_agent_locator) + + +# CommissionerId TLV (10) +class CommissionerId(object): + + def __init__(self, commissioner_id): + self._commissioner_id = commissioner_id + + @property + def commissioner_id(self): + return self._commissioner_id + + def __eq__(self, other): + return self.commissioner_id == other.commissioner_id + + def __repr__(self): + return "CommissionerId(commissioner_id={})".format(self.commissioner_id) + + +class CommissionerIdFactory(object): + def parse(self, data, message_info): + commissioner_id = data.getvalue().decode('utf-8') + return CommissionerId(commissioner_id) + + +# Commissioner Session ID TLV (11) +class CommissionerSessionId(object): + + def __init__(self, commissioner_session_id): + self._commissioner_session_id = commissioner_session_id + + @property + def commissioner_session_id(self): + return self._commissioner_session_id + + def __eq__(self, other): + common.expect_the_same_class(self, other) + + return self._commissioner_session_id == other._commissioner_session_id + + def __repr__(self): + return "CommissionerSessionId(commissioner_session_id={})".format(self._commissioner_session_id) + + def to_hex(self): + return struct.pack('>BBH', TlvType.COMMISSIONER_SESSION_ID, 2, self.commissioner_session_id) + + +class CommissionerSessionIdFactory: + + def parse(self, data, message_info): + session_id = struct.unpack(">H", data.read(2))[0] + return CommissionerSessionId(session_id) + + +# SecurityPolicy TLV (12) +class SecurityPolicy(object): + # TODO: Not implemented yet + pass + + +class SecurityPolicyFactory(object): + # TODO: Not implemented yet + + def parse(self, data, message_info): + raise NotImplementedError("TODO: Not implemented yet") + + +# Get TLV (13) +class Get(object): + # TODO: Not implemented yet + pass + + +class GetFactory(object): + # TODO: Not implemented yet + + def parse(self, data, message_info): + raise NotImplementedError("TODO: Not implemented yet") + + +# ActiveTimestamp TLV (14) +class ActiveTimestamp(object): + # TODO: Not implemented yet + pass + + +class ActiveTimestampFactory(object): + # TODO: Not implemented yet + + def parse(self, data, message_info): + raise NotImplementedError("TODO: Not implemented yet") + + +# Commissioner UDP Port TLV (15) +class CommissionerUdpPort(object): + + def __init__(self, udp_port): + self._udp_port = udp_port + + @property + def udp_port(self): + return self._udp_port + + def __eq__(self, other): + common.expect_the_same_class(self, other) + + return self._udp_port == other._udp_port + + def __repr__(self): + return "CommissionerUdpPort(udp_port={})".format(self._udp_port) + + +class CommissionerUdpPortFactory: + + def parse(self, data, message_info): + udp_port = struct.unpack(">H", data.read(2))[0] + return CommissionerUdpPort(udp_port) + + +# State TLV (16) class State(object): def __init__(self, state): @@ -82,11 +443,109 @@ class State(object): class StateFactory: - def parse(self, data): + def parse(self, data, message_info): state = ord(data.read(1)) return State(state) +# JoinerDtlsEncapsulation TLV (17) +class JoinerDtlsEncapsulation(object): + # TODO: Not implemented yet + pass + + +class JoinerDtlsEncapsulationFactory(object): + # TODO: Not implemented yet + + def parse(self, data, message_info): + raise NotImplementedError("TODO: Not implemented yet") + + +# JoinerUdpPort TLV (18) +class JoinerUdpPort(object): + + def __init__(self, udp_port): + self._udp_port = udp_port + + @property + def udp_port(self): + return self._udp_port + + def __eq__(self, other): + return type(self) is type(other) and self.udp_port == other.udp_port + + def __repr__(self): + return "JoinerUdpPort(udp_port={})".format(self.udp_port) + + +class JoinerUdpPortFactory(object): + + def parse(self, data, message_info): + udp_port = struct.unpack(">H", data.read(2))[0] + return JoinerUdpPort(udp_port) + + +# JoinerIID TLV (19) +class JoinerIID(object): + # TODO: Not implemented yet + pass + + +class JoinerIIDFactory(object): + # TODO: Not implemented yet + + def parse(self, data, message_info): + raise NotImplementedError("TODO: Not implemented yet") + + +# JoinerRouterLocator TLV (20) +class JoinerRouterLocator(object): + # TODO: Not implemented yet + pass + + +class JoinerRouterLocatorFactory(object): + # TODO: Not implemented yet + + def parse(self, data, message_info): + raise NotImplementedError("TODO: Not implemented yet") + + +# JoinerRouterKEK TLV (21) +class JoinerRouterKEK(object): + # TODO: Not implemented yet + pass + + +class JoinerRouterKEKFactory(object): + # TODO: Not implemented yet + + def parse(self, data, message_info): + raise NotImplementedError("TODO: Not implemented yet") + + +# ProvisioningURL TLV (32) +class ProvisioningUrl(object): + + def __init__(self, url): + self._url = url + + @property + def url(self): + return self._url + + def __repr__(self): + return "ProvisioningUrl(url={})".format(self.url) + + +class ProvisioningUrlFactory: + + def parse(self, data, message_info): + url = data.decode('utf-8') + return ProvisioningUrl(url) + + +# VendorName TLV (33) class VendorName(object): def __init__(self, vendor_name): @@ -102,14 +561,14 @@ class VendorName(object): def __repr__(self): return "VendorName(vendor_name={})".format(self.vendor_name) - class VendorNameFactory: - def parse(self, data): + def parse(self, data, message_info): vendor_name = data.getvalue().decode('utf-8') return VendorName(vendor_name) +# VendorModel TLV (34) class VendorModel(object): def __init__(self, vendor_model): @@ -128,10 +587,12 @@ class VendorModel(object): class VendorModelFactory: - def parse(self, data): + def parse(self, data, message_info): vendor_model = data.getvalue().decode('utf-8') return VendorModel(vendor_model) + +# VendorSWVersion TLV (35) class VendorSWVersion(object): def __init__(self, vendor_sw_version): @@ -150,11 +611,31 @@ class VendorSWVersion(object): class VendorSWVersionFactory: - def parse(self, data): + def parse(self, data, message_info): vendor_sw_version = data.getvalue() return VendorSWVersion(vendor_sw_version) +# VendorData TLV (36) +class VendorData(object): + + def __init__(self, data): + self._vendor_data = data + + @property + def vendor_data(self): + return self._vendor_data + + def __repr__(self): + return "Vendor(url={})".format(self.vendor_data) + + +class VendorDataFactory(object): + + def parse(self, data, message_info): + return VendorData(data) + + # VendorStackVersion TLV (37) class VendorStackVersion(object): @@ -189,10 +670,9 @@ class VendorStackVersion(object): def __repr__(self): return "VendorStackVersion(vendor_stack_version={}, build={}, rev={}, minor={}, major={})".format(self.stack_vendor_oui, self.build, self.rev, self.minor, self.major) - class VendorStackVersionFactory: - def parse(self, data): + def parse(self, data, message_info): stack_vendor_oui = struct.unpack(">H", data.read(2))[0] rest = struct.unpack(">BBBB", data.read(4)) build = rest[1] << 4 | (0xf0 & rest[2]) @@ -202,42 +682,191 @@ class VendorStackVersionFactory: return VendorStackVersion(stack_vendor_oui, build, rev, minor, major) -class ProvisioningUrl(object): +# UdpEncapsulation TLV (48) +class UdpEncapsulation(object): + # TODO: Not implemented yet + pass - def __init__(self, url): - self._url = url + +class UdpEncapsulationFactory(object): + # TODO: Not implemented yet + + def parse(self, data, message_info): + raise NotImplementedError("TODO: Not implemented yet") + + +# Ipv6Address TLV (49) +class Ipv6Address(object): + # TODO: Not implemented yet + pass + + +class Ipv6AddressFactory(object): + # TODO: Not implemented yet + + def parse(self, data, message_info): + raise NotImplementedError("TODO: Not implemented yet") + + +# PendingTimestamp TLV (51) +class PendingTimestamp(object): + # TODO: Not implemented yet + pass + + +class PendingTimestampFactory(object): + # TODO: Not implemented yet + + def parse(self, data, message_info): + raise NotImplementedError("TODO: Not implemented yet") + + +# DelayTimer TLV (52) +class DelayTimer(object): + # TODO: Not implemented yet + pass + + +class DelayTimerFactory(object): + # TODO: Not implemented yet + + def parse(self, data, message_info): + raise NotImplementedError("TODO: Not implemented yet") + + +# ChannelMask TLV (53) +class ChannelMask(object): + # TODO: Not implemented yet + pass + + +class ChannelMaskFactory(object): + # TODO: Not implemented yet + + def parse(self, data, message_info): + raise NotImplementedError("TODO: Not implemented yet") + + +# Count TLV (54) +class Count(object): + # TODO: Not implemented yet + pass + + +class CountFactory(object): + # TODO: Not implemented yet + + def parse(self, data, message_info): + raise NotImplementedError("TODO: Not implemented yet") + + +# Period TLV (55) +class Period(object): + # TODO: Not implemented yet + pass + + +class PeriodFactory(object): + # TODO: Not implemented yet + + def parse(self, data, message_info): + raise NotImplementedError("TODO: Not implemented yet") + + +# ScanDuration TLV (56) +class ScanDuration(object): + # TODO: Not implemented yet + pass + + +class ScanDurationFactory(object): + # TODO: Not implemented yet + + def parse(self, data, message_info): + raise NotImplementedError("TODO: Not implemented yet") + + +# EnergyList TLV (57) +class EnergyList(object): + # TODO: Not implemented yet + pass + + +class EnergyListFactory(object): + # TODO: Not implemented yet + + def parse(self, data, message_info): + raise NotImplementedError("TODO: Not implemented yet") + + +# Discovery Request TLV (128) +class DiscoveryRequest(object): + + def __init__(self, version, joiner_flag): + self._version = version + self._joiner_flag = joiner_flag @property - def url(self): - return self._url + def version(self): + return self._version - def __repr__(self): - return "ProvisioningUrl(url={})".format(self.url) - - -class ProvisioningUrlFactory: - - def parse(self, data): - url = data.decode('utf-8') - return ProvisioningUrl(url) - - -class VendorData(object): - - def __init__(self, data): - self._vendor_data = data @property - def vendor_data(self): - return self._vendor_data + def joiner_flag(self): + return self._joiner_flag + + def __eq__(self, other): + return (type(self) is type(other) + and self.version == other.version + and self.joiner_flag == other.joiner_flag) def __repr__(self): - return "Vendor(url={})".format(self.vendor_data) + return "DiscoveryRequest(version={}, joiner_flag={})".format( + self.version, self.joiner_flag) -class VendorDataFactory(object): +class DiscoveryRequestFactory(object): - def parse(self, data): - return VendorData(data) + def parse(self, data, message_info): + data_byte = struct.unpack(">B", data.read(1))[0] + version = (data_byte & 0xf0) >> 4 + joiner_flag = (data_byte & 0x08) >> 3 + + return DiscoveryRequest(version, joiner_flag) + + +# Discovery Response TLV (128) +class DiscoveryResponse(object): + + def __init__(self, version, native_flag): + self._version = version + self._native_flag = native_flag + + @property + def version(self): + return self._version + + @property + def native_flag(self): + return self._native_flag + + def __eq__(self, other): + return (type(self) is type(other) + and self.version == other.version + and self.native_flag == other.native_flag) + + def __repr__(self): + return "DiscoveryResponse(version={}, native_flag={})".format( + self.version, self.native_flag) + + +class DiscoveryResponseFactory(object): + + def parse(self, data, message_info): + data_byte = struct.unpack(">B", data.read(1))[0] + version = (data_byte & 0xf0) >> 4 + native_flag = (data_byte & 0x08) >> 3 + + return DiscoveryResponse(version, native_flag) class MeshCopCommand(object): @@ -288,15 +917,13 @@ class MeshCopCommandFactory: length = self._get_length(data) value = data.read(length) factory = self._get_tlv_factory(_type) - if factory == None: - return None - return factory.parse(io.BytesIO(value)) + return factory.parse(io.BytesIO(value), None) # message_info not needed here def _get_mesh_cop_msg_type(self, msg_type_str): - tp = self._mesh_cop_msg_type_map[msg_type_str] - if tp == None: - raise RuntimeError('Mesh cop message type not found: {}'.format(msg_type_str)) - return tp + try: + return self._mesh_cop_msg_type_map[msg_type_str] + except KeyError: + raise KeyError('Mesh cop message type not found: {}'.format(msg_type_str)) def parse(self, cmd_type_str, data): cmd_type = self._get_mesh_cop_msg_type(cmd_type_str) @@ -325,143 +952,3 @@ class ThreadDiscoveryTlvsFactory(SubTlvsFactory): def __init__(self, sub_tlvs_factories): super(ThreadDiscoveryTlvsFactory, self).__init__(sub_tlvs_factories) - - -class DiscoveryRequest(object): - - def __init__(self, version, joiner_flag): - self._version = version - self._joiner_flag = joiner_flag - - @property - def version(self): - return self._version - - @property - def joiner_flag(self): - return self._joiner_flag - - def __eq__(self, other): - return (type(self) is type(other) - and self.version == other.version - and self.joiner_flag == other.joiner_flag) - - def __repr__(self): - return "DiscoveryRequest(version={}, joiner_flag={})".format( - self.version, self.joiner_flag) - - -class DiscoveryRequestFactory(object): - - def parse(self, data, message_info): - data_byte = struct.unpack(">B", data.read(1))[0] - version = (data_byte & 0xf0) >> 4 - joiner_flag = (data_byte & 0x08) >> 3 - - return DiscoveryRequest(version, joiner_flag) - - -class DiscoveryResponse(object): - - def __init__(self, version, native_flag): - self._version = version - self._native_flag = native_flag - - @property - def version(self): - return self._version - - @property - def native_flag(self): - return self._native_flag - - def __eq__(self, other): - return (type(self) is type(other) - and self.version == other.version - and self.native_flag == other.native_flag) - - def __repr__(self): - return "DiscoveryResponse(version={}, native_flag={})".format( - self.version, self.native_flag) - - -class DiscoveryResponseFactory(object): - - def parse(self, data, message_info): - data_byte = struct.unpack(">B", data.read(1))[0] - version = (data_byte & 0xf0) >> 4 - native_flag = (data_byte & 0x08) >> 3 - - return DiscoveryResponse(version, native_flag) - - -class ExtendedPanid(object): - - def __init__(self, extended_panid): - self._extended_panid = extended_panid - - @property - def extended_panid(self): - return self._extended_panid - - def __eq__(self, other): - return (type(self) is type(other) - and self.extended_panid == other.extended_panid) - - def __repr__(self): - return "ExtendedPanid(extended_panid={})".format(self.extended_panid) - - -class ExtendedPanidFactory(object): - - def parse(self, data, message_info): - extended_panid = struct.unpack(">Q", data.read(8))[0] - return ExtendedPanid(extended_panid) - - -class NetworkName(object): - - def __init__(self, network_name): - self._network_name = network_name - - @property - def network_name(self): - return self._network_name - - def __eq__(self, other): - return (type(self) is type(other) - and self.network_name == other.network_name) - - def __repr__(self): - return "NetworkName(network_name={})".format(self.network_name) - - -class NetworkNameFactory(object): - - def parse(self, data, message_info): - len = message_info.length - network_name = struct.unpack("{}s".format(10), data.read(len))[0] - return NetworkName(network_name) - - -class JoinerUdpPort(object): - - def __init__(self, udp_port): - self._udp_port = udp_port - - @property - def udp_port(self): - return self._udp_port - - def __eq__(self, other): - return type(self) is type(other) and self.udp_port == other.udp_port - - def __repr__(self): - return "JoinerUdpPort(udp_port={})".format(self.udp_port) - - -class JoinerUdpPortFactory(object): - - def parse(self, data, message_info): - udp_port = struct.unpack(">H", data.read(2))[0] - return JoinerUdpPort(udp_port) diff --git a/tests/scripts/thread-cert/network_data.py b/tests/scripts/thread-cert/network_data.py index 0f077c1b9..c0e6e797d 100644 --- a/tests/scripts/thread-cert/network_data.py +++ b/tests/scripts/thread-cert/network_data.py @@ -33,6 +33,7 @@ import struct from binascii import hexlify from enum import IntEnum +from tlvs_parsing import SubTlvsFactory import common @@ -47,13 +48,6 @@ class TlvType(IntEnum): SERVER = 6 -class MeshcopTlvType(IntEnum): - STEERING_DATA = 8 - BORDER_AGENT_LOCATOR = 9 - COMMISSIONER_SESSION_ID = 11 - COMMISSIONER_UDP_PORT = 15 - - class NetworkData(object): def __init__(self, stable): @@ -64,36 +58,6 @@ class NetworkData(object): return self._stable -class SubTlvsFactory(object): - - def __init__(self, sub_tlvs_factories): - self._sub_tlvs_factories = sub_tlvs_factories - - def _get_factory(self, _type): - try: - return self._sub_tlvs_factories[_type] - except KeyError: - raise RuntimeError("Could not find factory. Factory type = {}.".format(_type)) - - def parse(self, data, message_info): - sub_tlvs = [] - - while data.tell() < len(data.getvalue()): - _type = ord(data.read(1)) - - length = ord(data.read(1)) - value = data.read(length) - - factory = self._get_factory(_type) - - message_info.length = length - tlv = factory.parse(io.BytesIO(value), message_info) - - sub_tlvs.append(tlv) - - return sub_tlvs - - class NetworkDataSubTlvsFactory(SubTlvsFactory): def parse(self, data, message_info): @@ -434,107 +398,6 @@ class CommissioningDataFactory(object): return CommissioningData(sub_tlvs, message_info.stable) - -class SteeringData(object): - - def __init__(self, bloom_filter): - self._bloom_filter = bloom_filter - - @property - def bloom_filter(self): - return self._bloom_filter - - def __eq__(self, other): - common.expect_the_same_class(self, other) - - return self._bloom_filter == other._bloom_filter - - def __repr__(self): - return "SteeringData(bloom_filter={})".format(hexlify(self._bloom_filter)) - - -class SteeringDataFactory: - - def parse(self, data, message_info): - bloom_filter = data.read(message_info.length) - return SteeringData(bloom_filter) - - -class BorderAgentLocator(object): - - def __init__(self, address): - self._udp_port = address - - @property - def udp_port(self): - return self._udp_port - - def __eq__(self, other): - common.expect_the_same_class(self, other) - - return self._udp_port == other._udp_port - - def __repr__(self): - return "BorderAgentLocator(rloc16={})".format(hex(self._udp_port)) - - -class BorderAgentLocatorFactory: - - def parse(self, data, message_info): - border_agent_locator = struct.unpack(">H", data.read(2))[0] - return BorderAgentLocator(border_agent_locator) - - -class CommissionerSessionId(object): - - def __init__(self, commissioner_session_id): - self._udp_port = commissioner_session_id - - @property - def udp_port(self): - return self._udp_port - - def __eq__(self, other): - common.expect_the_same_class(self, other) - - return self._udp_port == other._udp_port - - def __repr__(self): - return "CommissionerSessionId(id={})".format(hex(self._udp_port)) - - -class CommissionerSessionIdFactory: - - def parse(self, data, message_info): - session_id = struct.unpack(">H", data.read(2))[0] - return CommissionerSessionId(session_id) - - -class CommissionerUdpPort(object): - - def __init__(self, udp_port): - self._udp_port = udp_port - - @property - def udp_port(self): - return self._udp_port - - def __eq__(self, other): - common.expect_the_same_class(self, other) - - return self._udp_port == other._udp_port - - def __repr__(self): - return "CommissionerUdpPort(udp_port={})".format(self._udp_port) - - -class CommissionerUdpPortFactory: - - def parse(self, data, message_info): - udp_port = struct.unpack(">H", data.read(2))[0] - return CommissionerUdpPort(udp_port) - - class Service(NetworkData): def __init__(self, t, _id, enterprise_number, service_data_length, service_data, sub_tlvs, stable): diff --git a/tests/scripts/thread-cert/network_layer.py b/tests/scripts/thread-cert/network_layer.py index d4cde02ab..796b3af18 100644 --- a/tests/scripts/thread-cert/network_layer.py +++ b/tests/scripts/thread-cert/network_layer.py @@ -314,22 +314,3 @@ class ThreadNetworkDataFactory(object): tlvs = self._network_data_tlvs_factory.parse(data, message_info) return ThreadNetworkData(tlvs) - -class NetworkLayerTlvsFactory(object): - - def __init__(self, tlvs_factories): - self._tlvs_factories = tlvs_factories - - def parse(self, data, message_info): - tlvs = [] - - while data.tell() < len(data.getvalue()): - _type = ord(data.read(1)) - length = ord(data.read(1)) - - factory = self._tlvs_factories[_type] - tlv = factory.parse(io.BytesIO(data.read(length)), message_info) - - tlvs.append(tlv) - - return tlvs diff --git a/tests/scripts/thread-cert/node.py b/tests/scripts/thread-cert/node.py index e3c7d7866..0d283a16e 100755 --- a/tests/scripts/thread-cert/node.py +++ b/tests/scripts/thread-cert/node.py @@ -205,6 +205,12 @@ class Node: def get_addr(self, prefix): return self.interface.get_addr(prefix) + def get_addr_rloc(self): + return self.interface.get_addr_rloc() + + def get_addr_leader_aloc(self): + return self.interface.get_addr_leader_aloc() + def get_eidcaches(self): return self.interface.get_eidcaches() @@ -296,5 +302,11 @@ class Node: def coaps_get(self): self.interface.coaps_get() + def commissioner_mgmtset(self, tlvs_binary): + self.interface.commissioner_mgmtset(tlvs_binary) + + def commissioner_mgmtset_with_tlvs(self, tlvs): + self.interface.commissioner_mgmtset_with_tlvs(tlvs) + if __name__ == '__main__': unittest.main() diff --git a/tests/scripts/thread-cert/node_cli.py b/tests/scripts/thread-cert/node_cli.py index f868a6dcd..9aabea497 100644 --- a/tests/scripts/thread-cert/node_cli.py +++ b/tests/scripts/thread-cert/node_cli.py @@ -511,6 +511,22 @@ class otCli: return None + def get_addr_rloc(self): + addrs = self.get_addrs() + for addr in addrs: + segs = addr.split(':') + if segs[4] == '0' and segs[5] == 'ff' and segs[6] == 'fe00' and segs[7] != 'fc00': + return addr + return None + + def get_addr_leader_aloc(self): + addrs = self.get_addrs() + for addr in addrs: + segs = addr.split(':') + if segs[4] == '0' and segs[5] == 'ff' and segs[6] == 'fe00' and segs[7] == 'fc00': + return addr + return None + def get_eidcaches(self): eidcaches = [] self.send_command('eidcache') @@ -901,3 +917,18 @@ class otCli: timeout = 5 self._expect('Received coap secure response', timeout=timeout) + + def commissioner_mgmtset(self, tlvs_binary): + cmd = 'commissioner mgmtset binary ' + tlvs_binary + self.send_command(cmd) + self._expect('Done') + + def bytes_to_hex_str(self, src): + return ''.join(format(x, '02x') for x in src) + + def commissioner_mgmtset_with_tlvs(self, tlvs): + payload = bytearray() + for tlv in tlvs: + payload += tlv.to_hex() + self.commissioner_mgmtset(self.bytes_to_hex_str(payload)) + diff --git a/tests/scripts/thread-cert/test_network_layer.py b/tests/scripts/thread-cert/test_network_layer.py index d0bad26f2..7b5a0498e 100755 --- a/tests/scripts/thread-cert/test_network_layer.py +++ b/tests/scripts/thread-cert/test_network_layer.py @@ -481,24 +481,5 @@ class TestThreadNetworkDataFactory(unittest.TestCase): self.assertEqual(tlvs, thread_network_data.tlvs) -class TestNetworkLayerTlvsFactory(unittest.TestCase): - - def test_should_create_tlv_when_parse_method_is_called(self): - # GIVEN - class DummyNetworkDataTlvsFactory: - - def parse(self, data, message_info): - return bytearray(data.read()) - - tlv = any_tlv_data() - - factory = network_layer.NetworkLayerTlvsFactory({tlv[0]: DummyNetworkDataTlvsFactory()}) - - # WHEN - actual_tlvs = factory.parse(io.BytesIO(tlv), None) - - # THEN - self.assertEqual(tlv[2:], actual_tlvs[0]) - if __name__ == "__main__": unittest.main() diff --git a/tests/scripts/thread-cert/tlvs_parsing.py b/tests/scripts/thread-cert/tlvs_parsing.py new file mode 100644 index 000000000..cfe46cb10 --- /dev/null +++ b/tests/scripts/thread-cert/tlvs_parsing.py @@ -0,0 +1,59 @@ +#!/usr/bin/env python +# +# Copyright (c) 2019, The OpenThread Authors. +# All rights reserved. +# +# Redistribution and use in source and binary forms, with or without +# modification, are permitted provided that the following conditions are met: +# 1. Redistributions of source code must retain the above copyright +# notice, this list of conditions and the following disclaimer. +# 2. Redistributions in binary form must reproduce the above copyright +# notice, this list of conditions and the following disclaimer in the +# documentation and/or other materials provided with the distribution. +# 3. Neither the name of the copyright holder nor the +# names of its contributors may be used to endorse or promote products +# derived from this software without specific prior written permission. +# +# THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" +# AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE +# IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE +# ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE +# LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR +# CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF +# SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS +# INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN +# CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) +# ARISING IN ANY WAY OUT OF THE USE OF THIS SOFTWARE, EVEN IF ADVISED OF THE +# POSSIBILITY OF SUCH DAMAGE. + +import io + +class SubTlvsFactory(object): + + def __init__(self, sub_tlvs_factories): + self._sub_tlvs_factories = sub_tlvs_factories + + def _get_factory(self, _type): + try: + return self._sub_tlvs_factories[_type] + except KeyError: + raise RuntimeError("Could not find factory. Factory type = {}.".format(_type)) + + def parse(self, data, message_info): + sub_tlvs = [] + + while data.tell() < len(data.getvalue()): + _type = ord(data.read(1)) + + length = ord(data.read(1)) + value = data.read(length) + + factory = self._get_factory(_type) + + message_info.length = length + tlv = factory.parse(io.BytesIO(value), message_info) + + sub_tlvs.append(tlv) + + return sub_tlvs +