diff --git a/bumble/gatt_server.py b/bumble/gatt_server.py index ee7ac7ba..0a9e4297 100644 --- a/bumble/gatt_server.py +++ b/bumble/gatt_server.py @@ -62,6 +62,20 @@ # ----------------------------------------------------------------------------- GATT_SERVER_DEFAULT_MAX_MTU = 517 +# A security requirement implies the access that it protects +_READ_PERMISSIONS = ( + att.Attribute.READABLE + | att.Attribute.READ_REQUIRES_ENCRYPTION + | att.Attribute.READ_REQUIRES_AUTHENTICATION + | att.Attribute.READ_REQUIRES_AUTHORIZATION +) +_WRITE_PERMISSIONS = ( + att.Attribute.WRITEABLE + | att.Attribute.WRITE_REQUIRES_ENCRYPTION + | att.Attribute.WRITE_REQUIRES_AUTHENTICATION + | att.Attribute.WRITE_REQUIRES_AUTHORIZATION +) + # ----------------------------------------------------------------------------- # Helpers @@ -99,6 +113,8 @@ def __init__(self, device: Device) -> None: self.max_mtu = ( GATT_SERVER_DEFAULT_MAX_MTU # The max MTU we're willing to negotiate ) + # When False, READABLE and WRITEABLE are not required for peer reads and writes + self.strict_permissions = True self.subscribers = ( {} ) # Map of subscriber states by connection handle and attribute handle @@ -370,6 +386,16 @@ def send_response(self, bearer: att.Bearer, response: att.ATT_PDU) -> None: logger.debug(f'GATT Response from server: {_bearer_id(bearer)} {response}') self.send_gatt_pdu(bearer, bytes(response)) + async def _read_attribute_value( + self, bearer: att.Bearer, attribute: att.Attribute + ) -> bytes: + # Read a value on behalf of a peer + if self.strict_permissions and not attribute.permissions & _READ_PERMISSIONS: + raise att.ATT_Error( + error_code=att.ATT_READ_NOT_PERMITTED_ERROR, att_handle=attribute.handle + ) + return await attribute.read_value(bearer) + async def notify_subscriber( self, bearer: att.Bearer, @@ -721,16 +747,21 @@ async def on_att_find_by_type_value_request( pdu_space_available = bearer.att_mtu - 2 attributes = [] response: att.ATT_PDU - async for attribute in ( + for attribute in ( attribute for attribute in self.attributes if attribute.handle >= request.starting_handle and attribute.handle <= request.ending_handle and attribute.type == request.attribute_type - and (await attribute.read_value(bearer)) == request.attribute_value and pdu_space_available >= 4 ): - # TODO: check permissions + # Only attributes that can be read are returned + try: + attribute_value = await self._read_attribute_value(bearer, attribute) + except att.ATT_Error: + continue + if attribute_value != request.attribute_value: + continue # Add the attribute to the list attributes.append(attribute) @@ -803,7 +834,7 @@ async def on_att_read_by_type_request( and pdu_space_available ): try: - attribute_value = await attribute.read_value(bearer) + attribute_value = await self._read_attribute_value(bearer, attribute) except att.ATT_Error as error: # If the first attribute is unreadable, return an error # Otherwise return attributes up to this point @@ -856,7 +887,7 @@ async def on_att_read_request( response: att.ATT_PDU if attribute := self.get_attribute(request.attribute_handle): try: - value = await attribute.read_value(bearer) + value = await self._read_attribute_value(bearer, attribute) except att.ATT_Error as error: response = att.ATT_Error_Response( request_opcode_in_error=request.op_code, @@ -885,7 +916,7 @@ async def on_att_read_blob_request( response: att.ATT_PDU if attribute := self.get_attribute(request.attribute_handle): try: - value = await attribute.read_value(bearer) + value = await self._read_attribute_value(bearer, attribute) except att.ATT_Error as error: response = att.ATT_Error_Response( request_opcode_in_error=request.op_code, @@ -1014,9 +1045,16 @@ async def on_att_read_multiple_request( ) self.send_response(bearer, response) return - # No need to catch permission errors here, since these attributes - # must all be world-readable - attribute_value = await attribute.read_value(bearer) + try: + attribute_value = await self._read_attribute_value(bearer, attribute) + except att.ATT_Error as error: + response = att.ATT_Error_Response( + request_opcode_in_error=request.op_code, + attribute_handle_in_error=handle, + error_code=error.error_code, + ) + self.send_response(bearer, response) + return # Check the attribute value size max_attribute_size = min(bearer.att_mtu - 1, 251) if len(attribute_value) > max_attribute_size: @@ -1056,9 +1094,16 @@ async def on_att_read_multiple_variable_request( ) self.send_response(bearer, response) return - # No need to catch permission errors here, since these attributes - # must all be world-readable - attribute_value = await attribute.read_value(bearer) + try: + attribute_value = await self._read_attribute_value(bearer, attribute) + except att.ATT_Error as error: + response = att.ATT_Error_Response( + request_opcode_in_error=request.op_code, + attribute_handle_in_error=handle, + error_code=error.error_code, + ) + self.send_response(bearer, response) + return length = len(attribute_value) # Check the attribute value size max_attribute_size = min(bearer.att_mtu - 3, 251) @@ -1102,7 +1147,17 @@ async def on_att_write_request( ) return - # TODO: check permissions + # Check that the attribute can be written + if self.strict_permissions and not attribute.permissions & _WRITE_PERMISSIONS: + self.send_response( + bearer, + att.ATT_Error_Response( + request_opcode_in_error=request.op_code, + attribute_handle_in_error=request.attribute_handle, + error_code=att.ATT_WRITE_NOT_PERMITTED_ERROR, + ), + ) + return # Check the request parameters if len(request.attribute_value) > GATT_MAX_ATTRIBUTE_VALUE_SIZE: @@ -1144,7 +1199,9 @@ async def on_att_write_command( if attribute is None: return - # TODO: check permissions + # Check that the attribute can be written + if self.strict_permissions and not attribute.permissions & _WRITE_PERMISSIONS: + return # Check the request parameters if len(request.attribute_value) > GATT_MAX_ATTRIBUTE_VALUE_SIZE: @@ -1164,7 +1221,8 @@ def on_att_prepare_write_request( ''' # Check that the attribute exists - if self.get_attribute(request.attribute_handle) is None: + attribute = self.get_attribute(request.attribute_handle) + if attribute is None: self.send_response( bearer, att.ATT_Error_Response( @@ -1175,6 +1233,18 @@ def on_att_prepare_write_request( ) return + # Check that the attribute can be written + if self.strict_permissions and not attribute.permissions & _WRITE_PERMISSIONS: + self.send_response( + bearer, + att.ATT_Error_Response( + request_opcode_in_error=request.op_code, + attribute_handle_in_error=request.attribute_handle, + error_code=att.ATT_WRITE_NOT_PERMITTED_ERROR, + ), + ) + return + # Queue the partial value, to be committed on Execute Write Request self.prepared_writes.setdefault(bearer, []).append( ( diff --git a/tests/gatt_test.py b/tests/gatt_test.py index 64cfa7c4..dfc36060 100644 --- a/tests/gatt_test.py +++ b/tests/gatt_test.py @@ -1865,6 +1865,165 @@ async def test_write_long_value_gap_rejected(): assert characteristic.value is None +# ----------------------------------------------------------------------------- +@pytest.mark.asyncio +async def test_write_not_permitted(): + devices = await TwoDevices.create_with_connection() + + characteristic = Characteristic( + '1234', Characteristic.Properties.READ, Characteristic.READABLE, b'1234' + ) + devices[1].add_service(Service('ABCD', [characteristic])) + client = devices.connections[0].gatt_client + + response = await client.send_request( + att.ATT_Write_Request( + attribute_handle=characteristic.handle, attribute_value=b'5678' + ) + ) + assert isinstance(response, att.ATT_Error_Response) + assert response.error_code == att.ATT_WRITE_NOT_PERMITTED_ERROR + + await client.send_command( + att.ATT_Write_Command( + attribute_handle=characteristic.handle, attribute_value=b'5678' + ) + ) + + response = await client.send_request( + att.ATT_Prepare_Write_Request( + attribute_handle=characteristic.handle, + value_offset=0, + part_attribute_value=b'5678', + ) + ) + assert isinstance(response, att.ATT_Error_Response) + assert response.error_code == att.ATT_WRITE_NOT_PERMITTED_ERROR + + # Nothing was queued, so there is nothing to commit + response = await client.send_request(att.ATT_Execute_Write_Request(flags=0x01)) + assert isinstance(response, att.ATT_Execute_Write_Response) + await async_barrier() + + assert characteristic.value == b'1234' + + +# ----------------------------------------------------------------------------- +@pytest.mark.asyncio +async def test_read_not_permitted(): + devices = await TwoDevices.create_with_connection() + + characteristic = Characteristic( + '1234', Characteristic.Properties.WRITE, Characteristic.WRITEABLE, b'1234' + ) + devices[1].add_service(Service('ABCD', [characteristic])) + client = devices.connections[0].gatt_client + + for request in ( + att.ATT_Read_Request(attribute_handle=characteristic.handle), + att.ATT_Read_Blob_Request( + attribute_handle=characteristic.handle, value_offset=0 + ), + att.ATT_Read_By_Type_Request( + starting_handle=0x0001, + ending_handle=0xFFFF, + attribute_type=characteristic.uuid, + ), + att.ATT_Read_Multiple_Request(set_of_handles=[characteristic.handle]), + att.ATT_Read_Multiple_Variable_Request(set_of_handles=[characteristic.handle]), + ): + response = await client.send_request(request) + assert isinstance(response, att.ATT_Error_Response) + assert response.error_code == att.ATT_READ_NOT_PERMITTED_ERROR + assert response.attribute_handle_in_error == characteristic.handle + + # An attribute that cannot be read is not matched by value + response = await client.send_request( + att.ATT_Find_By_Type_Value_Request( + starting_handle=0x0001, + ending_handle=0xFFFF, + attribute_type=characteristic.uuid, + attribute_value=b'1234', + ) + ) + assert isinstance(response, att.ATT_Error_Response) + assert response.error_code == att.ATT_ATTRIBUTE_NOT_FOUND_ERROR + + +# ----------------------------------------------------------------------------- +@pytest.mark.asyncio +async def test_requires_permissions_imply_access(): + devices = await TwoDevices.create_with_connection() + + # No READABLE or WRITEABLE bit, only the security requirements + characteristic = Characteristic( + '1234', + Characteristic.Properties.READ | Characteristic.Properties.WRITE, + Characteristic.READ_REQUIRES_ENCRYPTION + | Characteristic.WRITE_REQUIRES_ENCRYPTION, + b'1234', + ) + devices[1].add_service(Service('ABCD', [characteristic])) + client = devices.connections[0].gatt_client + + # Not encrypted: the security requirement is reported, not NOT_PERMITTED + response = await client.send_request( + att.ATT_Read_Request(attribute_handle=characteristic.handle) + ) + assert isinstance(response, att.ATT_Error_Response) + assert response.error_code == att.ATT_INSUFFICIENT_ENCRYPTION_ERROR + + response = await client.send_request( + att.ATT_Write_Request( + attribute_handle=characteristic.handle, attribute_value=b'5678' + ) + ) + assert isinstance(response, att.ATT_Error_Response) + assert response.error_code == att.ATT_INSUFFICIENT_ENCRYPTION_ERROR + + # Encrypted: both are allowed + devices.connections[1].encryption = 1 + response = await client.send_request( + att.ATT_Read_Request(attribute_handle=characteristic.handle) + ) + assert isinstance(response, att.ATT_Read_Response) + assert response.attribute_value == b'1234' + + response = await client.send_request( + att.ATT_Write_Request( + attribute_handle=characteristic.handle, attribute_value=b'5678' + ) + ) + assert isinstance(response, att.ATT_Write_Response) + assert characteristic.value == b'5678' + + +# ----------------------------------------------------------------------------- +@pytest.mark.asyncio +async def test_permissions_not_strict(): + devices = await TwoDevices.create_with_connection() + devices[1].gatt_server.strict_permissions = False + + characteristic = Characteristic( + '1234', Characteristic.Properties.NOTIFY, Characteristic.Permissions(0), b'1234' + ) + devices[1].add_service(Service('ABCD', [characteristic])) + client = devices.connections[0].gatt_client + + response = await client.send_request( + att.ATT_Write_Request( + attribute_handle=characteristic.handle, attribute_value=b'5678' + ) + ) + assert isinstance(response, att.ATT_Write_Response) + + response = await client.send_request( + att.ATT_Read_Request(attribute_handle=characteristic.handle) + ) + assert isinstance(response, att.ATT_Read_Response) + assert response.attribute_value == b'5678' + + # ----------------------------------------------------------------------------- if __name__ == '__main__': logging.basicConfig(level=os.environ.get('BUMBLE_LOGLEVEL', 'INFO').upper()) diff --git a/tests/profiles/heart_rate_service_test.py b/tests/profiles/heart_rate_service_test.py index 70119c1f..3bfdff91 100644 --- a/tests/profiles/heart_rate_service_test.py +++ b/tests/profiles/heart_rate_service_test.py @@ -31,7 +31,7 @@ (1, 1000), (True, False, None), (2, None), ((3.0, 4.0, 5.0), None) ), ) -async def test_read_measurement( +async def test_notify_measurement( heart_rate: int, sensor_contact_detected: bool | None, energy_expanded: int | None, @@ -47,7 +47,15 @@ async def test_read_measurement( async with device_module.Peer(devices.connections[1]) as peer: client = peer.create_service_proxy(heart_rate_service.HeartRateServiceProxy) assert client - assert await client.heart_rate_measurement.read_value() == measurement + # The measurement can only be notified, not read + notifications = asyncio.Queue[ + heart_rate_service.HeartRateService.HeartRateMeasurement + ]() + await client.heart_rate_measurement.subscribe(notifications.put_nowait) + await devices[0].notify_subscribers( + service.heart_rate_measurement_characteristic + ) + assert await notifications.get() == measurement # -----------------------------------------------------------------------------