Skip to content
Merged
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
14 changes: 5 additions & 9 deletions bumble/device.py
Original file line number Diff line number Diff line change
Expand Up @@ -2171,7 +2171,7 @@ def with_connection_from_address(function):
@functools.wraps(function)
def wrapper(device: Device, address: hci.Address, *args, **kwargs):
if connection := device.pending_connections.get(address):
return function(device, connection, address, *args, **kwargs)
return function(device, connection, *args, **kwargs)
for connection in device.connections.values():
if connection.peer_address == address:
return function(device, connection, *args, **kwargs)
Expand Down Expand Up @@ -6443,18 +6443,14 @@ def on_cs_procedure(self, event: hci.HCI_LE_CS_Procedure_Enable_Complete_Event):

# [Classic only]
@host_event_handler
@try_with_connection_from_address
@with_connection_from_address
def on_role_change(
self,
connection: Optional[Connection],
peer_address: hci.Address,
connection: Connection,
new_role: hci.Role,
):
if connection:
connection.role = new_role
connection.emit(connection.EVENT_ROLE_CHANGE, new_role)
else:
logger.warning("Role change to unknown connection %s", peer_address)
connection.role = new_role
connection.emit(connection.EVENT_ROLE_CHANGE, new_role)

# [Classic only]
@host_event_handler
Expand Down
28 changes: 28 additions & 0 deletions tests/device_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -761,6 +761,34 @@ async def test_inquiry_result_with_rssi():
m.assert_called_with(hci.Address("00:11:22:33:44:55/P"), 3, mock.ANY, 5)


# -----------------------------------------------------------------------------
@pytest.mark.parametrize(
"roles",
(
(hci.Role.PERIPHERAL, hci.Role.CENTRAL),
(hci.Role.CENTRAL, hci.Role.PERIPHERAL),
),
)
@pytest.mark.asyncio
async def test_accept_classic_connection(roles: tuple[hci.Role, hci.Role]):
devices = TwoDevices()
devices[0].classic_enabled = True
devices[1].classic_enabled = True
await devices[0].power_on()
await devices[1].power_on()

accept_task = asyncio.create_task(devices[1].accept(role=roles[1]))
await devices[0].connect(
devices[1].public_address, transport=PhysicalTransport.BR_EDR
)
await accept_task

assert devices.connections[0]
assert devices.connections[0].role == roles[0]
assert devices.connections[1]
assert devices.connections[1].role == roles[1]


# -----------------------------------------------------------------------------
async def run_test_device():
await test_device_connect_parallel()
Expand Down
Loading