Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
100 changes: 85 additions & 15 deletions bumble/gatt_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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:
Expand All @@ -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(
Expand All @@ -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(
(
Expand Down
159 changes: 159 additions & 0 deletions tests/gatt_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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())
Expand Down
12 changes: 10 additions & 2 deletions tests/profiles/heart_rate_service_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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


# -----------------------------------------------------------------------------
Expand Down
Loading