[PATCH v2 5/6] tests: hwsim: add async infrastructure

Johannes Berg johannes at sipsolutions.net
Fri Sep 11 15:59:53 PDT 2026


From: Johannes Berg <johannes.berg at intel.com>

Add infrastructure to be able to write async tests.
Current tests are fine with sync code, however, the
new relay infrastructure is much better as async.

Add the basic infrastructure (including the hwsim
connectivity check, which takes a bit of effort)
that allows for writing async tests.

Signed-off-by: Johannes Berg <johannes.berg at intel.com>
---
 tests/hwsim/hostapd.py       |  31 +++++++++
 tests/hwsim/hwsim_utils.py   | 118 +++++++++++++++++++++++++++--------
 tests/hwsim/run-tests.py     |  11 +++-
 tests/hwsim/wpasupplicant.py |  29 +++++++++
 wpaspy/wpaspy.py             |  39 ++++++++++++
 5 files changed, 199 insertions(+), 29 deletions(-)

diff --git a/tests/hwsim/hostapd.py b/tests/hwsim/hostapd.py
index 6ce6701667b9..e855390c7fd6 100644
--- a/tests/hwsim/hostapd.py
+++ b/tests/hwsim/hostapd.py
@@ -219,6 +219,10 @@ class Hostapd:
         logger.debug(self.dbg + ": CTRL: " + cmd)
         return self.ctrl.request(cmd)
 
+    async def request_async(self, cmd):
+        logger.debug(self.dbg + ": CTRL: " + cmd)
+        return await self.ctrl.request_async(cmd)
+
     def ping(self):
         return "PONG" in self.request("PING")
 
@@ -293,6 +297,10 @@ class Hostapd:
         return wpaspy.wait_event(self.mon, events, timeout,
                                  log_prefix=self.dbg + ": ")
 
+    async def wait_event_async(self, events, timeout):
+        return await wpaspy.wait_event_async(self.mon, events, timeout,
+                                             log_prefix=self.dbg + ": ")
+
     @staticmethod
     def _check_sta_event(ev, addr, what):
         if ev is None:
@@ -312,6 +320,17 @@ class Hostapd:
                 addr, "4-way handshake completion")
         return ev
 
+    async def wait_sta_async(self, addr=None, timeout=2, wait_4way_hs=False):
+        ev = self._check_sta_event(
+            await self.wait_event_async(["AP-STA-CONNECT"], timeout=timeout),
+            addr, "STA connection")
+        if wait_4way_hs:
+            self._check_sta_event(
+                await self.wait_event_async(["EAPOL-4WAY-HS-COMPLETED"],
+                                            timeout=timeout),
+                addr, "4-way handshake completion")
+        return ev
+
     def wait_sta_disconnect(self, addr=None, timeout=2):
         return self._check_sta_event(
             self.wait_event(["AP-STA-DISCONNECT"], timeout=timeout), addr,
@@ -689,6 +708,18 @@ def add_ap(apdev, params, wait_enabled=True, no_enable=False, timeout=30,
                 raise Exception("AP startup failed")
         return hapd
 
+async def add_ap_async(apdev, params, wait_enabled=True, timeout=30, **kwargs):
+    hapd = add_ap(apdev, params, wait_enabled=False, **kwargs)
+    if kwargs.get('no_enable') or not wait_enabled:
+        return hapd
+    ev = await hapd.wait_event_async(["AP-ENABLED", "AP-DISABLED"],
+                                     timeout=timeout)
+    if ev is None:
+        raise Exception("AP startup timed out")
+    if "AP-ENABLED" not in ev:
+        raise Exception("AP startup failed")
+    return hapd
+
 def add_bss(apdev, ifname, confname, ignore_error=False):
     phy = utils.get_phy(apdev)
     try:
diff --git a/tests/hwsim/hwsim_utils.py b/tests/hwsim/hwsim_utils.py
index c8d918cb7376..ad6ed37656ce 100644
--- a/tests/hwsim/hwsim_utils.py
+++ b/tests/hwsim/hwsim_utils.py
@@ -6,9 +6,11 @@
 
 import os
 import time
+import asyncio
 import logging
 logger = logging.getLogger()
 
+import wpaspy
 from wpasupplicant import WpaSupplicant
 
 def sync_carrier(dev, ifname=None):
@@ -54,13 +56,29 @@ def config_data_test(dev1, dev2, dev1group, dev2group, ifname1, ifname2):
 
     sync_carrier(dev2, ifname2)
 
-def run_multicast_connectivity_test(dev1, dev2, tos=None,
-                                    dev1group=False, dev2group=False,
-                                    ifname1=None, ifname2=None,
-                                    config=True, timeout=5,
-                                    send_len=None, multicast_to_unicast=False,
-                                    broadcast_retry_c=1,
-                                    check_dup=False):
+def _drive_sync(gen):
+    try:
+        req = next(gen)
+        while True:
+            req = gen.send(wpaspy.wait_event(*req))
+    except StopIteration as e:
+        return e.value
+
+async def _drive_async(gen):
+    try:
+        req = next(gen)
+        while True:
+            req = gen.send(await wpaspy.wait_event_async(*req))
+    except StopIteration as e:
+        return e.value
+
+def _run_multicast_connectivity_test(dev1, dev2, tos=None,
+                                     dev1group=False, dev2group=False,
+                                     ifname1=None, ifname2=None,
+                                     config=True, timeout=5,
+                                     send_len=None, multicast_to_unicast=False,
+                                     broadcast_retry_c=1,
+                                     check_dup=False):
     addr1 = dev1.get_addr(dev1group)
     addr2 = dev2.get_addr(dev2group)
 
@@ -85,7 +103,7 @@ def run_multicast_connectivity_test(dev1, dev2, tos=None,
                 ev = dev2.wait_group_event(["DATA-TEST-RX"],
                                            timeout=timeout)
             else:
-                ev = dev2.wait_event(["DATA-TEST-RX"], timeout=timeout)
+                ev = yield (dev2.mon, ["DATA-TEST-RX"], timeout, dev2.dbg + ": ")
             if ev is None:
                 raise Exception("dev1->dev2 broadcast data delivery failed")
             if multicast_to_unicast:
@@ -122,7 +140,7 @@ def run_multicast_connectivity_test(dev1, dev2, tos=None,
             ev = dev2.wait_group_event(["DATA-TEST-RX"],
                                        timeout=timeout)
         else:
-            ev = dev2.wait_event(["DATA-TEST-RX"], timeout=timeout)
+            ev = yield (dev2.mon, ["DATA-TEST-RX"], timeout, dev2.dbg + ": ")
         if not ev:
             return
         if not " id=" in ev:
@@ -134,10 +152,14 @@ def run_multicast_connectivity_test(dev1, dev2, tos=None,
         if _id in rxed:
             raise Exception("duplicate packet with ID %d received" % _id)
 
-def run_connectivity_test(dev1, dev2, tos, dev1group=False, dev2group=False,
-                          ifname1=None, ifname2=None, config=True, timeout=5,
-                          multicast_to_unicast=False, broadcast=True,
-                          send_len=None, check_bcast_dup=False):
+# see _run_multicast_connectivity_test() for arguments
+def run_multicast_connectivity_test(*args, **kwargs):
+    return _drive_sync(_run_multicast_connectivity_test(*args, **kwargs))
+
+def _run_connectivity_test(dev1, dev2, tos, dev1group=False, dev2group=False,
+                           ifname1=None, ifname2=None, config=True, timeout=5,
+                           multicast_to_unicast=False, broadcast=True,
+                           send_len=None, check_bcast_dup=False):
     addr1 = dev1.get_addr(dev1group)
     addr2 = dev2.get_addr(dev2group)
 
@@ -165,7 +187,7 @@ def run_connectivity_test(dev1, dev2, tos, dev1group=False, dev2group=False,
         if dev2group:
             ev = dev2.wait_group_event(["DATA-TEST-RX"], timeout=timeout)
         else:
-            ev = dev2.wait_event(["DATA-TEST-RX"], timeout=timeout)
+            ev = yield (dev2.mon, ["DATA-TEST-RX"], timeout, dev2.dbg + ": ")
         if ev is None:
             raise Exception("dev1->dev2 unicast data delivery failed")
         if "DATA-TEST-RX {} {}".format(addr2, addr1) not in ev:
@@ -178,11 +200,13 @@ def run_connectivity_test(dev1, dev2, tos, dev1group=False, dev2group=False,
                 raise Exception("Unexpected dev1->dev2 unicast data length")
 
         if broadcast:
-            run_multicast_connectivity_test(dev1, dev2, tos,
-                                            dev1group, dev2group,
-                                            ifname1, ifname2, False, timeout,
-                                            send_len, False, broadcast_retry_c,
-                                            check_dup=check_bcast_dup)
+            yield from _run_multicast_connectivity_test(dev1, dev2, tos,
+                                                        dev1group, dev2group,
+                                                        ifname1, ifname2,
+                                                        False, timeout,
+                                                        send_len, False,
+                                                        broadcast_retry_c,
+                                                        check_dup=check_bcast_dup)
 
         cmd = "DATA_TEST_TX {} {} {}".format(addr1, addr2, tos)
         if send_len is not None:
@@ -194,7 +218,7 @@ def run_connectivity_test(dev1, dev2, tos, dev1group=False, dev2group=False,
         if dev1group:
             ev = dev1.wait_group_event(["DATA-TEST-RX"], timeout=timeout)
         else:
-            ev = dev1.wait_event(["DATA-TEST-RX"], timeout=timeout)
+            ev = yield (dev1.mon, ["DATA-TEST-RX"], timeout, dev1.dbg + ": ")
         if ev is None:
             raise Exception("dev2->dev1 unicast data delivery failed")
         if "DATA-TEST-RX {} {}".format(addr1, addr2) not in ev:
@@ -207,12 +231,14 @@ def run_connectivity_test(dev1, dev2, tos, dev1group=False, dev2group=False,
                 raise Exception("Unexpected dev2->dev1 unicast data length")
 
         if broadcast:
-            run_multicast_connectivity_test(dev2, dev1, tos,
-                                            dev2group, dev1group,
-                                            ifname2, ifname1, False, timeout,
-                                            send_len, multicast_to_unicast,
-                                            broadcast_retry_c,
-                                            check_dup=check_bcast_dup)
+            yield from _run_multicast_connectivity_test(dev2, dev1, tos,
+                                                        dev2group, dev1group,
+                                                        ifname2, ifname1,
+                                                        False, timeout,
+                                                        send_len,
+                                                        multicast_to_unicast,
+                                                        broadcast_retry_c,
+                                                        check_dup=check_bcast_dup)
 
     finally:
         if config:
@@ -225,6 +251,10 @@ def run_connectivity_test(dev1, dev2, tos, dev1group=False, dev2group=False,
             else:
                 dev2.request("DATA_TEST_CONFIG 0")
 
+# see _run_connectivity_test() for arguments
+def run_connectivity_test(*args, **kwargs):
+    return _drive_sync(_run_connectivity_test(*args, **kwargs))
+
 def test_connectivity(dev1, dev2, dscp=None, tos=None, max_tries=1,
                       dev1group=False, dev2group=False,
                       ifname1=None, ifname2=None, config=True, timeout=5,
@@ -261,6 +291,42 @@ def test_connectivity_iface(dev1, dev2, ifname, dscp=None, tos=None,
     test_connectivity(dev1, dev2, dscp, tos, ifname2=ifname,
                       max_tries=max_tries, timeout=timeout)
 
+async def test_connectivity_async(dev1, dev2, dscp=None, tos=None, max_tries=1,
+                                  dev1group=False, dev2group=False,
+                                  ifname1=None, ifname2=None, config=True,
+                                  timeout=5, multicast_to_unicast=False,
+                                  success_expected=True, broadcast=True,
+                                  send_len=None, check_bcast_dup=False):
+    if dscp:
+        tos = dscp << 2
+    if not tos:
+        tos = 0
+
+    success = False
+    last_err = None
+    for i in range(0, max_tries):
+        try:
+            await _drive_async(_run_connectivity_test(dev1, dev2, tos,
+                                                      dev1group, dev2group,
+                                                      ifname1, ifname2,
+                                                      config=config,
+                                                      timeout=timeout,
+                                                      multicast_to_unicast=multicast_to_unicast,
+                                                      broadcast=broadcast,
+                                                      send_len=send_len,
+                                                      check_bcast_dup=check_bcast_dup))
+            success = True
+            break
+        except Exception as e:
+            last_err = e
+            if i + 1 < max_tries:
+                await asyncio.sleep(1)
+    if success_expected and not success:
+        raise Exception(last_err)
+    if not success_expected and success:
+        raise Exception("Unexpected connectivity detected")
+
+
 def test_connectivity_p2p(dev1, dev2, dscp=None, tos=None):
     test_connectivity(dev1, dev2, dscp, tos, dev1group=True, dev2group=True)
 
diff --git a/tests/hwsim/run-tests.py b/tests/hwsim/run-tests.py
index f113b8149347..f35400368813 100755
--- a/tests/hwsim/run-tests.py
+++ b/tests/hwsim/run-tests.py
@@ -37,6 +37,11 @@ from check_kernel import check_kernel
 from wlantest import Wlantest
 from utils import HwsimSkip
 
+def run_test_func(t, *args):
+    if inspect.iscoroutinefunction(t):
+        return asyncio.run(t(*args))
+    return t(*args)
+
 def set_term_echo(fd, enabled):
     [iflag, oflag, cflag, lflag, ispeed, ospeed, cc] = termios.tcgetattr(fd)
     if enabled:
@@ -632,11 +637,11 @@ def main():
                     params['logdir'] = args.logdir
                     params['name'] = name
                     params['prefix'] = os.path.join(args.logdir, name)
-                    t(dev, apdev, params)
+                    run_test_func(t, dev, apdev, params)
                 elif t.__code__.co_argcount > 1:
-                    t(dev, apdev)
+                    run_test_func(t, dev, apdev)
                 else:
-                    t(dev)
+                    run_test_func(t, dev)
                 result = "PASS"
                 if check_country_00:
                     for d in dev:
diff --git a/tests/hwsim/wpasupplicant.py b/tests/hwsim/wpasupplicant.py
index 15178e2b817e..db8eb297111a 100644
--- a/tests/hwsim/wpasupplicant.py
+++ b/tests/hwsim/wpasupplicant.py
@@ -270,6 +270,10 @@ class WpaSupplicant:
         logger.debug(self.dbg + ": CTRL: " + cmd)
         return self.ctrl.request(cmd, timeout=timeout)
 
+    async def request_async(self, cmd, timeout=10):
+        logger.debug(self.dbg + ": CTRL: " + cmd)
+        return await self.ctrl.request_async(cmd, timeout=timeout)
+
     def global_request(self, cmd):
         if self.global_iface is None:
             return self.request(cmd)
@@ -884,6 +888,10 @@ class WpaSupplicant:
         return wpaspy.wait_event(self.mon, events, timeout,
                                  log_prefix=self.dbg + ": ")
 
+    async def wait_event_async(self, events, timeout=10):
+        return await wpaspy.wait_event_async(self.mon, events, timeout,
+                                             log_prefix=self.dbg + ": ")
+
     def wait_global_event(self, events, timeout):
         if self.global_iface is None:
             return self.wait_event(events, timeout)
@@ -1174,6 +1182,27 @@ class WpaSupplicant:
             self.select_network(id)
         return id
 
+    async def connect_async(self, ssid=None, ssid2=None, timeout=None,
+                            wait_connect=True, only_add_network=False,
+                            **kwargs):
+        """
+        Async equivalent of connect(): everything connect() does is quick
+        request/response setup except the final wait for
+        CTRL-EVENT-CONNECTED, so that part is done here as a real
+        (non-thread) asyncio wait instead.
+        """
+        if timeout is None:
+            timeout = 20 if "eap" in kwargs else 15
+        id = self.connect(ssid, ssid2, timeout=timeout, wait_connect=False,
+                          only_add_network=only_add_network, **kwargs)
+        if only_add_network or not wait_connect:
+            return id
+        ev = await self.wait_event_async(["CTRL-EVENT-CONNECTED"],
+                                         timeout=timeout)
+        if ev is None:
+            raise Exception("Connection timed out")
+        return id
+
     def scan(self, type=None, freq=None, no_wait=False, only_new=False,
              passive=False, timeout=15):
         if not no_wait:
diff --git a/wpaspy/wpaspy.py b/wpaspy/wpaspy.py
index eb7f06ded838..9b50b9b3050e 100644
--- a/wpaspy/wpaspy.py
+++ b/wpaspy/wpaspy.py
@@ -10,6 +10,7 @@ import os
 import stat
 import socket
 import select
+import asyncio
 import logging
 
 logger = logging.getLogger()
@@ -64,6 +65,8 @@ class Ctrl:
                 if self.s != None:
                     self.s.close()
                 raise
+        # needed for async and non-async always uses select()
+        self.s.setblocking(False)
         self.started = True
 
     def __del__(self):
@@ -151,6 +154,25 @@ class Ctrl:
             r = res
         return r
 
+    async def recv_async(self):
+        data = await asyncio.get_running_loop().sock_recv(self.s, 4096)
+        return data.decode()
+
+    async def request_async(self, cmd, timeout=10):
+        if type(cmd) == str:
+            try:
+                cmd = cmd.encode()
+            except UnicodeDecodeError as e:
+                pass
+        if self.udp:
+            self.s.sendto(self.cookie + cmd, self.sockaddr)
+        else:
+            self.s.send(cmd)
+        try:
+            return await asyncio.wait_for(self.recv_async(), timeout)
+        except asyncio.TimeoutError:
+            raise Exception("Timeout on waiting response")
+
 def wait_event(sock, events, timeout=10, log_prefix=""):
     assert isinstance(events, list), "'events' must be a list"
     start = os.times()[4]
@@ -168,3 +190,20 @@ def wait_event(sock, events, timeout=10, log_prefix=""):
         if not sock.pending(timeout=remaining):
             break
     return None
+
+async def wait_event_async(sock, events, timeout=10, log_prefix=""):
+    assert isinstance(events, list), "'events' must be a list"
+    loop = asyncio.get_running_loop()
+    deadline = loop.time() + timeout
+    while True:
+        remaining = deadline - loop.time()
+        if remaining <= 0:
+            return None
+        try:
+            ev = await asyncio.wait_for(sock.recv_async(), remaining)
+        except asyncio.TimeoutError:
+            return None
+        logger.debug(log_prefix + ev)
+        for event in events:
+            if event in ev:
+                return ev
-- 
2.55.0




More information about the Hostap mailing list