xref: /linux/tools/testing/selftests/drivers/net/lib/py/env.py (revision c98610c2eb7002436f6fca55b21c63200644549a)
1# SPDX-License-Identifier: GPL-2.0
2
3import ipaddress
4import os
5import sys
6import time
7import json
8from pathlib import Path
9from lib.py import KsftSkipEx, KsftXfailEx
10from lib.py import ksft_setup, wait_file
11from lib.py import cmd, ethtool, ip, CmdExitFailure
12from lib.py import NetNS, NetdevSimDev, UserNetNS
13from .remote import Remote
14from . import bpftool, RtnlFamily, Netlink
15
16
17class NetDrvEnvBase:
18    """
19    Base class for a NIC / host environments
20
21    Attributes:
22      test_dir: Path to the source directory of the test
23      net_lib_dir: Path to the net/lib directory
24    """
25    def __init__(self, src_path):
26        self.src_path = Path(src_path)
27        self.test_dir = self.src_path.parent.resolve()
28        self.net_lib_dir = (Path(__file__).parent / "../../../../net/lib").resolve()
29
30        self.env = self._load_env_file()
31
32        # Following attrs must be set be inheriting classes
33        self.dev = None
34
35    def _load_env_file(self):
36        env = os.environ.copy()
37
38        src_dir = Path(self.src_path).parent.resolve()
39        if not (src_dir / "net.config").exists():
40            return ksft_setup(env)
41
42        with open((src_dir / "net.config").as_posix(), 'r') as fp:
43            for line in fp.readlines():
44                full_file = line
45                # Strip comments
46                pos = line.find("#")
47                if pos >= 0:
48                    line = line[:pos]
49                line = line.strip()
50                if not line:
51                    continue
52                pair = line.split('=', maxsplit=1)
53                if len(pair) != 2:
54                    raise Exception("Can't parse configuration line:", full_file)
55                env[pair[0]] = pair[1]
56        return ksft_setup(env)
57
58    def __del__(self):
59        pass
60
61    def __enter__(self):
62        ip(f"link set dev {self.dev['ifname']} up")
63        wait_file(f"/sys/class/net/{self.dev['ifname']}/carrier",
64                  lambda x: x.strip() == "1")
65
66        return self
67
68    def __exit__(self, ex_type, ex_value, ex_tb):
69        """
70        __exit__ gets called at the end of a "with" block.
71        """
72        self.__del__()
73
74
75class NetDrvEnv(NetDrvEnvBase):
76    """
77    Class for a single NIC / host env, with no remote end
78    """
79    def __init__(self, src_path, nsim_test=None, **kwargs):
80        super().__init__(src_path)
81
82        self._ns = None
83
84        if 'NETIF' in self.env:
85            if nsim_test is True:
86                raise KsftXfailEx("Test only works on netdevsim")
87
88            self.dev = ip("-d link show dev " + self.env['NETIF'], json=True)[0]
89        else:
90            if nsim_test is False:
91                raise KsftXfailEx("Test does not work on netdevsim")
92
93            self._ns = NetdevSimDev(**kwargs)
94            self.dev = self._ns.nsims[0].dev
95        self.ifname = self.dev['ifname']
96        self.ifindex = self.dev['ifindex']
97
98    def __del__(self):
99        if self._ns:
100            self._ns.remove()
101            self._ns = None
102
103
104class NetDrvEpEnv(NetDrvEnvBase):
105    """
106    Class for an environment with a local device and "remote endpoint"
107    which can be used to send traffic in.
108
109    For local testing it creates two network namespaces and a pair
110    of netdevsim devices.
111    """
112
113    # Network prefixes used for local tests
114    nsim_v4_pfx = "192.0.2."
115    nsim_v6_pfx = "2001:db8::"
116
117    def __init__(self, src_path, nsim_test=None, queue_count=None):
118        super().__init__(src_path)
119
120        self._stats_settle_time = None
121        self._queue_count = queue_count
122
123        # Things we try to destroy
124        self.remote = None
125        # These are for local testing state
126        self._netns = None
127        self._ns = None
128        self._ns_peer = None
129
130        self.addr_v        = { "4": None, "6": None }
131        self.remote_addr_v = { "4": None, "6": None }
132
133        if "NETIF" in self.env:
134            if nsim_test is True:
135                raise KsftXfailEx("Test only works on netdevsim")
136            self._check_env()
137
138            self.dev = ip("-d link show dev " + self.env['NETIF'], json=True)[0]
139
140            self.addr_v["4"] = self.env.get("LOCAL_V4")
141            self.addr_v["6"] = self.env.get("LOCAL_V6")
142            self.remote_addr_v["4"] = self.env.get("REMOTE_V4")
143            self.remote_addr_v["6"] = self.env.get("REMOTE_V6")
144            kind = self.env["REMOTE_TYPE"]
145            args = self.env["REMOTE_ARGS"]
146        else:
147            if nsim_test is False:
148                raise KsftXfailEx("Test does not work on netdevsim")
149
150            self.create_local()
151
152            self.dev = self._ns.nsims[0].dev
153
154            self.addr_v["4"] = self.nsim_v4_pfx + "1"
155            self.addr_v["6"] = self.nsim_v6_pfx + "1"
156            self.remote_addr_v["4"] = self.nsim_v4_pfx + "2"
157            self.remote_addr_v["6"] = self.nsim_v6_pfx + "2"
158            kind = "netns"
159            args = self._netns.name
160
161        self.remote = Remote(kind, args, src_path)
162
163        self.set_ipver("6" if self.addr_v["6"] else "4")
164
165        self.ifname = self.dev['ifname']
166        self.ifindex = self.dev['ifindex']
167
168        # resolve remote interface name
169        self.remote_ifname = self.resolve_remote_ifc()
170        self.remote_dev = ip("-d link show dev " + self.remote_ifname,
171                             host=self.remote, json=True)[0]
172        self.remote_ifindex = self.remote_dev['ifindex']
173
174        self._required_cmd = {}
175
176    def create_local(self):
177        nsim_kwargs = {}
178        if self._queue_count:
179            nsim_kwargs["queue_count"] = self._queue_count
180
181        self._netns = NetNS()
182        self._ns = NetdevSimDev(**nsim_kwargs)
183        self._ns_peer = NetdevSimDev(ns=self._netns, **nsim_kwargs)
184
185        with open("/proc/self/ns/net") as nsfd0, \
186             open("/var/run/netns/" + self._netns.name) as nsfd1:
187            ifi0 = self._ns.nsims[0].ifindex
188            ifi1 = self._ns_peer.nsims[0].ifindex
189            NetdevSimDev.ctrl_write('link_device',
190                                    f'{nsfd0.fileno()}:{ifi0} {nsfd1.fileno()}:{ifi1}')
191
192        ip(f"   addr add dev {self._ns.nsims[0].ifname} {self.nsim_v4_pfx}1/24")
193        ip(f"-6 addr add dev {self._ns.nsims[0].ifname} {self.nsim_v6_pfx}1/64 nodad")
194        ip(f"   link set dev {self._ns.nsims[0].ifname} up")
195
196        ip(f"   addr add dev {self._ns_peer.nsims[0].ifname} {self.nsim_v4_pfx}2/24", ns=self._netns)
197        ip(f"-6 addr add dev {self._ns_peer.nsims[0].ifname} {self.nsim_v6_pfx}2/64 nodad", ns=self._netns)
198        ip(f"   link set dev {self._ns_peer.nsims[0].ifname} up", ns=self._netns)
199
200    def _check_env(self):
201        vars_needed = [
202            ["LOCAL_V4", "LOCAL_V6"],
203            ["REMOTE_V4", "REMOTE_V6"],
204            ["REMOTE_TYPE"],
205            ["REMOTE_ARGS"]
206        ]
207        missing = []
208
209        for choice in vars_needed:
210            for entry in choice:
211                if entry in self.env:
212                    break
213            else:
214                missing.append(choice)
215        # Make sure v4 / v6 configs are symmetric
216        if ("LOCAL_V6" in self.env) != ("REMOTE_V6" in self.env):
217            missing.append(["LOCAL_V6", "REMOTE_V6"])
218        if ("LOCAL_V4" in self.env) != ("REMOTE_V4" in self.env):
219            missing.append(["LOCAL_V4", "REMOTE_V4"])
220        if missing:
221            raise Exception("Invalid environment, missing configuration:", missing,
222                            "Please see tools/testing/selftests/drivers/net/README.rst")
223
224    def resolve_remote_ifc(self):
225        v4 = v6 = None
226        if self.remote_addr_v["4"]:
227            v4 = ip("addr show to " + self.remote_addr_v["4"], json=True, host=self.remote)
228        if self.remote_addr_v["6"]:
229            v6 = ip("addr show to " + self.remote_addr_v["6"], json=True, host=self.remote)
230        if v4 and v6 and v4[0]["ifname"] != v6[0]["ifname"]:
231            raise Exception("Can't resolve remote interface name, v4 and v6 don't match")
232        if (v4 and len(v4) > 1) or (v6 and len(v6) > 1):
233            raise Exception("Can't resolve remote interface name, multiple interfaces match")
234        return v6[0]["ifname"] if v6 else v4[0]["ifname"]
235
236    def __del__(self):
237        if self._ns:
238            self._ns.remove()
239            self._ns = None
240        if self._ns_peer:
241            self._ns_peer.remove()
242            self._ns_peer = None
243        if self._netns:
244            del self._netns
245            self._netns = None
246        if self.remote:
247            del self.remote
248            self.remote = None
249
250    def require_ipver(self, ipver):
251        if not self.addr_v[ipver] or not self.remote_addr_v[ipver]:
252            raise KsftSkipEx(f"Test requires IPv{ipver} connectivity")
253
254    def set_ipver(self, ipver):
255        """
256        Modify the IP version used by the generic address fields.
257        """
258        if ipver == getattr(self, "addr_ipver", None):
259            return
260
261        self.require_ipver(ipver)
262
263        self.addr_ipver = ipver
264        self.addr = self.addr_v[ipver]
265        self.remote_addr = self.remote_addr_v[ipver]
266
267        # Bracketed addresses, some commands need IPv6 to be inside []
268        self.baddr = (f"[{self.addr_v['6']}]" if ipver == "6"
269                      else self.addr_v["4"])
270        self.remote_baddr = (f"[{self.remote_addr_v['6']}]" if ipver == "6"
271                             else self.remote_addr_v["4"])
272
273    def require_nsim(self, nsim_test=True):
274        """Require or exclude netdevsim for this test"""
275        if nsim_test and self._ns is None:
276            raise KsftXfailEx("Test only works on netdevsim")
277        if nsim_test is False and self._ns is not None:
278            raise KsftXfailEx("Test does not work on netdevsim")
279
280    def get_local_nsim_dev(self):
281        """Returns the local netdevsim device or None.
282           Using this method is discouraged, as it makes tests nsim-specific.
283           Standard interfaces available on all HW should ideally be used.
284           This method is intended for the few cases where nsim-specific
285           assertions need to be verified which cannot be verified otherwise.
286        """
287        return self._ns
288
289    def _require_cmd(self, comm, key, host=None):
290        cached = self._required_cmd.get(comm, {})
291        if cached.get(key) is None:
292            cached[key] = cmd("command -v -- " + comm, fail=False,
293                              shell=True, host=host).ret == 0
294        self._required_cmd[comm] = cached
295        return cached[key]
296
297    def require_cmd(self, comm, local=True, remote=False):
298        if local:
299            if not self._require_cmd(comm, "local"):
300                raise KsftSkipEx("Test requires command: " + comm)
301        if remote:
302            if not self._require_cmd(comm, "remote", host=self.remote):
303                raise KsftSkipEx("Test requires (remote) command: " + comm)
304
305    def wait_hw_stats_settle(self):
306        """
307        Wait for HW stats to become consistent, some devices DMA HW stats
308        periodically so events won't be reflected until next sync.
309        Good drivers will tell us via ethtool what their sync period is.
310        """
311        if self._stats_settle_time is None:
312            data = {}
313            try:
314                data = ethtool("-c " + self.ifname, json=True)[0]
315            except CmdExitFailure as e:
316                if "Operation not supported" not in e.cmd.stderr:
317                    raise
318
319            self._stats_settle_time = \
320                1.25 * data.get('stats-block-usecs', 20000) / 1000 / 1000
321
322        time.sleep(self._stats_settle_time)
323
324
325class NetDrvContEnv(NetDrvEpEnv):
326    """
327    Class for an environment with a netkit pair setup for forwarding traffic
328    between the physical interface and a network namespace.
329      NETIF           = "eth0"
330      LOCAL_V6        = "2001:db8:1::1"
331      REMOTE_V6       = "2001:db8:1::2"
332      LOCAL_PREFIX_V6 = "2001:db8:2::0/64"
333
334              +-----------------------------+        +------------------------------+
335      dst     | INIT NS                     |        | TEST NS                      |
336      2001:   | +---------------+           |        |                              |
337      db8:2::2| | NETIF         |           |  bpf   |                              |
338          +---|>| 2001:db8:1::1 |           |redirect| +-------------------------+  |
339          |   | |               |-----------|--------|>| Netkit                  |  |
340          |   | +---------------+           | _peer  | | nk_guest                |  |
341          |   | +-------------+ Netkit pair |        | | fe80::2/64              |  |
342          |   | | Netkit      |.............|........|>| 2001:db8:2::2/64        |  |
343          |   | | nk_host     |             |        | +-------------------------+  |
344          |   | | fe80::1/64  |             |        |                              |
345          |   | +-------------+             |        | route:                       |
346          |   |                             |        |   default                    |
347          |   | route:                      |        |     via fe80::1 dev nk_guest |
348          |   |   2001:db8:2::2/128         |        +------------------------------+
349          |   |     via fe80::2 dev nk_host |
350          |   +-----------------------------+
351          |
352          |   +---------------+
353          |   | REMOTE        |
354          +---| 2001:db8:1::2 |
355              +---------------+
356    """
357
358    def __init__(self, src_path, rxqueues=1, primary_rx_redirect=False,
359                 userns=False, **kwargs):
360        self.netns = None
361        self._userns = userns
362        self.nk_host_ifname = None
363        self.nk_guest_ifname = None
364        self._tc_clsact_added = False
365        self._tc_attached = False
366        self._primary_rx_redirect_attached = False
367        self._primary_rx_redirect_clsact_added = False
368        self._bpf_prog_pref = None
369        self._bpf_prog_id = None
370        self._init_ns_attached = False
371        self._remote_route_added = False
372        self._old_fwd = None
373        self._old_accept_ra = None
374
375        super().__init__(src_path, **kwargs)
376
377        self.require_ipver("6")
378        local_prefix = self.env.get("LOCAL_PREFIX_V6")
379        if not local_prefix:
380            raise KsftSkipEx("LOCAL_PREFIX_V6 required")
381
382        net = ipaddress.IPv6Network(local_prefix, strict=False)
383        self.ipv6_prefix = str(net.network_address)
384        self.nk_host_ipv6 = f"{self.ipv6_prefix}2:1"
385        self.nk_guest_ipv6 = f"{self.ipv6_prefix}2:2"
386
387        local_v6 = ipaddress.IPv6Address(self.addr_v["6"])
388        if local_v6 in net:
389            raise KsftSkipEx("LOCAL_V6 must not fall within LOCAL_PREFIX_V6")
390
391        rtnl = RtnlFamily()
392        rtnl.newlink(
393            {
394                "linkinfo": {
395                    "kind": "netkit",
396                    "data": {
397                        "mode": "l2",
398                        "policy": "forward",
399                        "peer-policy": "forward",
400                    },
401                },
402                "num-rx-queues": rxqueues,
403            },
404            flags=[Netlink.NLM_F_CREATE, Netlink.NLM_F_EXCL],
405        )
406
407        all_links = ip("-d link show", json=True)
408        netkit_links = [link for link in all_links
409                        if link.get('linkinfo', {}).get('info_kind') == 'netkit'
410                        and 'UP' not in link.get('flags', [])]
411
412        if len(netkit_links) != 2:
413            raise KsftSkipEx("Failed to create netkit pair")
414
415        netkit_links.sort(key=lambda x: x['ifindex'])
416        self.nk_host_ifname = netkit_links[1]['ifname']
417        self.nk_guest_ifname = netkit_links[0]['ifname']
418        self.nk_host_ifindex = netkit_links[1]['ifindex']
419        self.nk_guest_ifindex = netkit_links[0]['ifindex']
420
421        self._setup_ns()
422        self.attach_bpf()
423        if primary_rx_redirect:
424            self._attach_primary_rx_redirect_bpf()
425
426    def __del__(self):
427        if self._primary_rx_redirect_attached:
428            cmd(f"tc filter del dev {self.nk_host_ifname} ingress", fail=False)
429            self._primary_rx_redirect_attached = False
430
431        if self._primary_rx_redirect_clsact_added:
432            cmd(f"tc qdisc del dev {self.nk_host_ifname} clsact", fail=False)
433            self._primary_rx_redirect_clsact_added = False
434
435        if self._tc_attached:
436            cmd(f"tc filter del dev {self.ifname} ingress pref {self._bpf_prog_pref}")
437            self._tc_attached = False
438
439        if self._tc_clsact_added:
440            cmd(f"tc qdisc del dev {self.ifname} clsact")
441            self._tc_clsact_added = False
442
443        if self._remote_route_added:
444            cmd(f"ip -6 route del {self.nk_guest_ipv6}/128",
445                host=self.remote, fail=False)
446            self._remote_route_added = False
447
448        if self.nk_host_ifname:
449            cmd(f"ip link del dev {self.nk_host_ifname}")
450            self.nk_host_ifname = None
451            self.nk_guest_ifname = None
452
453        if self._init_ns_attached:
454            cmd("ip netns del init", fail=False)
455            self._init_ns_attached = False
456
457        if self.netns:
458            del self.netns
459            self.netns = None
460
461        if self._old_fwd is not None:
462            with open("/proc/sys/net/ipv6/conf/all/forwarding", "w",
463                      encoding="utf-8") as f:
464                f.write(self._old_fwd)
465            self._old_fwd = None
466        if self._old_accept_ra is not None:
467            with open("/proc/sys/net/ipv6/conf/all/accept_ra", "w",
468                      encoding="utf-8") as f:
469                f.write(self._old_accept_ra)
470            self._old_accept_ra = None
471
472        super().__del__()
473
474    def _setup_ns(self):
475        fwd_path = "/proc/sys/net/ipv6/conf/all/forwarding"
476        ra_path = "/proc/sys/net/ipv6/conf/all/accept_ra"
477        with open(fwd_path, encoding="utf-8") as f:
478            self._old_fwd = f.read().strip()
479        with open(ra_path, encoding="utf-8") as f:
480            self._old_accept_ra = f.read().strip()
481        with open(fwd_path, "w", encoding="utf-8") as f:
482            f.write("1")
483        with open(ra_path, "w", encoding="utf-8") as f:
484            f.write("2")
485
486        self.netns = UserNetNS() if self._userns else NetNS()
487        cmd("ip netns attach init 1")
488        self._init_ns_attached = True
489        ip("netns set init 0", ns=self.netns)
490        ip(f"link set dev {self.nk_guest_ifname} netns {self.netns.name}")
491        nk_guest_dev = ip(f"link show dev {self.nk_guest_ifname}",
492                          json=True, ns=self.netns)[0]
493        self.nk_guest_ifindex = nk_guest_dev['ifindex']
494        ip(f"link set dev {self.nk_host_ifname} up")
495        ip(f"-6 addr add fe80::1/64 dev {self.nk_host_ifname} nodad")
496        ip(f"-6 route add {self.nk_guest_ipv6}/128 via fe80::2 dev {self.nk_host_ifname}")
497
498        ip("link set lo up", ns=self.netns)
499        ip(f"link set dev {self.nk_guest_ifname} up", ns=self.netns)
500        ip(f"-6 addr add fe80::2/64 dev {self.nk_guest_ifname}", ns=self.netns)
501        ip(f"-6 addr add {self.nk_guest_ipv6}/64 dev {self.nk_guest_ifname} nodad", ns=self.netns)
502        ip(f"-6 route add default via fe80::1 dev {self.nk_guest_ifname}", ns=self.netns)
503
504    def _tc_ensure_clsact(self, ifname=None):
505        """Ensure a clsact qdisc exists on @ifname.
506
507        Returns True if this call added the qdisc, otherwise returns False.
508        """
509        if ifname is None:
510            ifname = self.ifname
511        qdisc = json.loads(cmd(f"tc -j qdisc show dev {ifname}").stdout)
512        for q in qdisc:
513            if q['kind'] == 'clsact':
514                return False
515        cmd(f"tc qdisc add dev {ifname} clsact")
516        return True
517
518    def _get_bpf_prog_ids(self):
519        filters = json.loads(cmd(f"tc -j filter show dev {self.ifname} ingress").stdout)
520        for bpf in filters:
521            if 'options' not in bpf:
522                continue
523            if bpf['options']['bpf_name'].startswith('nk_forward.bpf'):
524                return (bpf['pref'], bpf['options']['prog']['id'])
525        raise Exception("Failed to get BPF prog ID")
526
527    def _find_bss_map_id(self, prog_id):
528        """Find the .bss map ID for a loaded BPF program."""
529        prog_info = bpftool(f"prog show id {prog_id}", json=True)
530        for map_id in prog_info.get("map_ids", []):
531            map_info = bpftool(f"map show id {map_id}", json=True)
532            if map_info.get("name", "").endswith("bss"):
533                return map_id
534        raise Exception(f"Failed to find .bss map for prog {prog_id}")
535
536    def _find_bpf_obj(self, name):
537        bpf_obj = self.test_dir / name
538        if bpf_obj.exists():
539            return bpf_obj
540        bpf_obj = self.test_dir / "hw" / name
541        if bpf_obj.exists():
542            return bpf_obj
543        return None
544
545    def detach_bpf(self):
546        if self._tc_attached:
547            cmd(f"tc filter del dev {self.ifname} ingress pref "
548                f"{self._bpf_prog_pref}", fail=False)
549            self._tc_attached = False
550
551    def attach_bpf(self):
552        bpf_obj = self._find_bpf_obj("nk_forward.bpf.o")
553        if not bpf_obj:
554            raise KsftSkipEx("BPF prog nk_forward.bpf.o not found")
555
556        if self._tc_ensure_clsact():
557            self._tc_clsact_added = True
558        cmd(f"tc filter add dev {self.ifname} ingress bpf obj {bpf_obj}"
559            " sec tc/ingress direct-action")
560        self._tc_attached = True
561
562        (self._bpf_prog_pref, self._bpf_prog_id) = self._get_bpf_prog_ids()
563        bss_map_id = self._find_bss_map_id(self._bpf_prog_id)
564
565        ipv6_addr = ipaddress.IPv6Address(self.ipv6_prefix)
566        ipv6_bytes = ipv6_addr.packed
567        ifindex_bytes = self.nk_host_ifindex.to_bytes(4, byteorder='little')
568        value = ipv6_bytes + ifindex_bytes
569        value_hex = ' '.join(f'{b:02x}' for b in value)
570        bpftool(f"map update id {bss_map_id} key hex 00 00 00 00 value hex {value_hex}")
571
572    def _attach_primary_rx_redirect_bpf(self):
573        """Attach BPF redirect program on the primary netkit ingress."""
574        bpf_obj = self._find_bpf_obj("nk_primary_rx_redirect.bpf.o")
575        if not bpf_obj:
576            raise KsftSkipEx("nk_primary_rx_redirect.bpf.o not found")
577
578        if self._tc_ensure_clsact(self.nk_host_ifname):
579            self._primary_rx_redirect_clsact_added = True
580        cmd(f"tc filter add dev {self.nk_host_ifname} ingress"
581            f" bpf obj {bpf_obj} sec tc/ingress direct-action")
582        self._primary_rx_redirect_attached = True
583
584        ip(f"-6 route add {self.nk_guest_ipv6}/128 via {self.addr_v['6']}",
585           host=self.remote)
586        self._remote_route_added = True
587
588        filters = json.loads(
589            cmd(f"tc -j filter show dev {self.nk_host_ifname} ingress").stdout)
590        redirect_prog_id = None
591        for bpf in filters:
592            if 'options' not in bpf:
593                continue
594            if bpf['options']['bpf_name'].startswith('nk_primary_rx_redirect'):
595                redirect_prog_id = bpf['options']['prog']['id']
596                break
597        if redirect_prog_id is None:
598            raise Exception("Failed to get primary RX redirect BPF prog ID")
599
600        bss_map_id = self._find_bss_map_id(redirect_prog_id)
601        phys_ifindex_bytes = self.ifindex.to_bytes(4, byteorder=sys.byteorder)
602        value_hex = ' '.join(f'{b:02x}' for b in phys_ifindex_bytes)
603        bpftool(f"map update id {bss_map_id} key hex 00 00 00 00 value hex {value_hex}")
604