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