diff --git a/.github/dependabot.yml b/.github/dependabot.yml index 075307dea..b68d7daa6 100644 --- a/.github/dependabot.yml +++ b/.github/dependabot.yml @@ -13,3 +13,8 @@ updates: directory: /src/netbrowse schedule: interval: weekly + + - package-ecosystem: gomod + directory: /src/yangerd + schedule: + interval: weekly diff --git a/board/common/rootfs/etc/finit.d/available/firewalld.conf b/board/common/rootfs/etc/finit.d/available/firewalld.conf index 18581c052..439395090 100644 --- a/board/common/rootfs/etc/finit.d/available/firewalld.conf +++ b/board/common/rootfs/etc/finit.d/available/firewalld.conf @@ -1,3 +1,3 @@ -service [2345] reload:'firewall-cmd -q --reload' \ +service [2345] reload:'firewall reload' \ firewalld --nofork --log-target syslog \ -- Firewall daemon diff --git a/configs/aarch64_defconfig b/configs/aarch64_defconfig index b599df64a..613ded6e8 100644 --- a/configs/aarch64_defconfig +++ b/configs/aarch64_defconfig @@ -166,6 +166,7 @@ BR2_PACKAGE_CURIOS_NFTABLES=y BR2_PACKAGE_GENCERT=y BR2_PACKAGE_STATD=y BR2_PACKAGE_SUPPORT_ENCRYPT=y +BR2_PACKAGE_YANGERD=y BR2_PACKAGE_FACTORY=y BR2_PACKAGE_FINIT_PLUGIN_HOTPLUG=y BR2_PACKAGE_FINIT_PLUGIN_HOOK_SCRIPTS=y diff --git a/configs/aarch64_minimal_defconfig b/configs/aarch64_minimal_defconfig index b7429a85d..0cbfbc1d6 100644 --- a/configs/aarch64_minimal_defconfig +++ b/configs/aarch64_minimal_defconfig @@ -130,6 +130,7 @@ BR2_PACKAGE_CONFD_TEST_MODE=y BR2_PACKAGE_GENCERT=y BR2_PACKAGE_STATD=y BR2_PACKAGE_SUPPORT=y +BR2_PACKAGE_YANGERD=y BR2_PACKAGE_FACTORY=y BR2_PACKAGE_FINIT_PLUGIN_HOTPLUG=y BR2_PACKAGE_FINIT_PLUGIN_HOOK_SCRIPTS=y diff --git a/configs/arm_defconfig b/configs/arm_defconfig index 108521113..850bcedc1 100644 --- a/configs/arm_defconfig +++ b/configs/arm_defconfig @@ -152,6 +152,7 @@ BR2_PACKAGE_CONFD_TEST_MODE=y BR2_PACKAGE_GENCERT=y BR2_PACKAGE_STATD=y BR2_PACKAGE_SUPPORT_ENCRYPT=y +BR2_PACKAGE_YANGERD=y BR2_PACKAGE_FACTORY=y BR2_PACKAGE_FINIT_PLUGIN_HOTPLUG=y BR2_PACKAGE_FINIT_PLUGIN_HOOK_SCRIPTS=y diff --git a/configs/arm_minimal_defconfig b/configs/arm_minimal_defconfig index f38f8b719..415fb279e 100644 --- a/configs/arm_minimal_defconfig +++ b/configs/arm_minimal_defconfig @@ -130,6 +130,7 @@ BR2_PACKAGE_CONFD_TEST_MODE=y BR2_PACKAGE_GENCERT=y BR2_PACKAGE_STATD=y BR2_PACKAGE_SUPPORT=y +BR2_PACKAGE_YANGERD=y BR2_PACKAGE_FACTORY=y BR2_PACKAGE_FINIT_PLUGIN_HOTPLUG=y BR2_PACKAGE_FINIT_PLUGIN_HOOK_SCRIPTS=y diff --git a/configs/riscv64_defconfig b/configs/riscv64_defconfig index d2e3478bb..ab2a73db1 100644 --- a/configs/riscv64_defconfig +++ b/configs/riscv64_defconfig @@ -184,6 +184,7 @@ BR2_PACKAGE_NETD=y BR2_PACKAGE_GENCERT=y BR2_PACKAGE_STATD=y BR2_PACKAGE_SUPPORT_ENCRYPT=y +BR2_PACKAGE_YANGERD=y BR2_PACKAGE_FACTORY=y BR2_PACKAGE_FINIT_PLUGIN_HOTPLUG=y BR2_PACKAGE_FINIT_PLUGIN_HOOK_SCRIPTS=y diff --git a/configs/x86_64_defconfig b/configs/x86_64_defconfig index 6cac556ff..f5bf4e9bb 100644 --- a/configs/x86_64_defconfig +++ b/configs/x86_64_defconfig @@ -159,6 +159,7 @@ BR2_PACKAGE_CURIOS_NFTABLES=y BR2_PACKAGE_GENCERT=y BR2_PACKAGE_STATD=y BR2_PACKAGE_SUPPORT_ENCRYPT=y +BR2_PACKAGE_YANGERD=y BR2_PACKAGE_FACTORY=y BR2_PACKAGE_FINIT_PLUGIN_HOTPLUG=y BR2_PACKAGE_FINIT_PLUGIN_HOOK_SCRIPTS=y diff --git a/configs/x86_64_minimal_defconfig b/configs/x86_64_minimal_defconfig index f3db8f305..c392c58a8 100644 --- a/configs/x86_64_minimal_defconfig +++ b/configs/x86_64_minimal_defconfig @@ -127,6 +127,7 @@ BR2_PACKAGE_CONFD_TEST_MODE=y BR2_PACKAGE_GENCERT=y BR2_PACKAGE_STATD=y BR2_PACKAGE_SUPPORT=y +BR2_PACKAGE_YANGERD=y BR2_PACKAGE_FACTORY=y BR2_PACKAGE_FINIT_PLUGIN_HOTPLUG=y BR2_PACKAGE_FINIT_PLUGIN_HOOK_SCRIPTS=y diff --git a/package/Config.in b/package/Config.in index f91fb5580..382044291 100644 --- a/package/Config.in +++ b/package/Config.in @@ -14,6 +14,7 @@ source "$BR2_EXTERNAL_INFIX_PATH/package/curios-nftables/Config.in" source "$BR2_EXTERNAL_INFIX_PATH/package/gencert/Config.in" source "$BR2_EXTERNAL_INFIX_PATH/package/statd/Config.in" source "$BR2_EXTERNAL_INFIX_PATH/package/support/Config.in" +source "$BR2_EXTERNAL_INFIX_PATH/package/yangerd/Config.in" source "$BR2_EXTERNAL_INFIX_PATH/package/factory/Config.in" source "$BR2_EXTERNAL_INFIX_PATH/package/faux/Config.in" source "$BR2_EXTERNAL_INFIX_PATH/package/finit/Config.in" diff --git a/package/statd/Config.in b/package/statd/Config.in index acec4cd7a..83e42cc3e 100644 --- a/package/statd/Config.in +++ b/package/statd/Config.in @@ -1,5 +1,7 @@ config BR2_PACKAGE_STATD bool "statd" + depends on BR2_PACKAGE_HOST_GO_TARGET_ARCH_SUPPORTS # yangerd + select BR2_PACKAGE_YANGERD select BR2_PACKAGE_JANSSON select BR2_PACKAGE_LIBEV select BR2_PACKAGE_SYSREPO diff --git a/package/yangerd/Config.in b/package/yangerd/Config.in new file mode 100644 index 000000000..1720c401b --- /dev/null +++ b/package/yangerd/Config.in @@ -0,0 +1,7 @@ +config BR2_PACKAGE_YANGERD + bool "yangerd" + depends on BR2_PACKAGE_HOST_GO_TARGET_ARCH_SUPPORTS + help + Operational data daemon for YANG/NETCONF/RESTCONF. + Replaces Python yanger scripts with a persistent Go daemon + serving operational data over a Unix socket IPC protocol. diff --git a/package/yangerd/yangerd.conf b/package/yangerd/yangerd.conf new file mode 100644 index 000000000..760114375 --- /dev/null +++ b/package/yangerd/yangerd.conf @@ -0,0 +1,3 @@ +service <> name:yangerd notify:pid log:prio:daemon.notice,tag:yangerd \ + env:-/etc/default/yangerd \ + [2345] yangerd -- Operational data daemon diff --git a/package/yangerd/yangerd.mk b/package/yangerd/yangerd.mk new file mode 100644 index 000000000..a150a4638 --- /dev/null +++ b/package/yangerd/yangerd.mk @@ -0,0 +1,50 @@ +################################################################################ +# +# yangerd +# +################################################################################ + +YANGERD_VERSION = 1.0.0 +YANGERD_SITE = $(BR2_EXTERNAL_INFIX_PATH)/src/yangerd +YANGERD_SITE_METHOD = local +YANGERD_GOMOD = github.com/kernelkit/infix/src/yangerd +YANGERD_LICENSE = BSD-2-Clause +YANGERD_LICENSE_FILES = LICENSE +YANGERD_REDISTRIBUTE = NO + +YANGERD_BUILD_TARGETS = cmd/yangerd cmd/yangerctl +YANGERD_INSTALL_BINS = yangerd yangerctl + +define YANGERD_INSTALL_EXTRA + $(INSTALL) -D -m 0644 $(YANGERD_PKGDIR)/yangerd.conf \ + $(FINIT_D)/available/yangerd.conf + ln -sf ../available/yangerd.conf $(FINIT_D)/enabled/yangerd.conf + $(INSTALL) -d $(TARGET_DIR)/etc/default + echo '# yangerd build-time feature flags (generated by yangerd.mk)' \ + > $(TARGET_DIR)/etc/default/yangerd + echo 'YANGERD_ENABLE_WIFI=$(if $(BR2_PACKAGE_IW),true,false)' \ + >> $(TARGET_DIR)/etc/default/yangerd + echo 'YANGERD_ENABLE_CONTAINERS=$(if $(BR2_PACKAGE_PODMAN),true,false)' \ + >> $(TARGET_DIR)/etc/default/yangerd + echo 'YANGERD_ENABLE_GPS=$(if $(BR2_PACKAGE_GPSD),true,false)' \ + >> $(TARGET_DIR)/etc/default/yangerd + echo 'YANGERD_ENABLE_LLDP=$(if $(BR2_PACKAGE_LLDPD),true,false)' \ + >> $(TARGET_DIR)/etc/default/yangerd + echo 'YANGERD_ENABLE_FIREWALL=$(if $(BR2_PACKAGE_FIREWALLD),true,false)' \ + >> $(TARGET_DIR)/etc/default/yangerd + echo 'YANGERD_ENABLE_DHCP=$(if $(BR2_PACKAGE_DNSMASQ),true,false)' \ + >> $(TARGET_DIR)/etc/default/yangerd + echo 'YANGERD_ENABLE_FRR=$(if $(BR2_PACKAGE_FRR),true,false)' \ + >> $(TARGET_DIR)/etc/default/yangerd + echo 'YANGERD_LOG_LEVEL=info' >> $(TARGET_DIR)/etc/default/yangerd + echo '# Soft heap limit, the live data is well under 1 MiB and the' \ + >> $(TARGET_DIR)/etc/default/yangerd + echo '# rest is garbage; without it RSS drifts past 100 MiB on a' \ + >> $(TARGET_DIR)/etc/default/yangerd + echo '# 512 MiB box and the OOM killer picks yangerd during upgrades' \ + >> $(TARGET_DIR)/etc/default/yangerd + echo 'GOMEMLIMIT=64MiB' >> $(TARGET_DIR)/etc/default/yangerd +endef +YANGERD_POST_INSTALL_TARGET_HOOKS += YANGERD_INSTALL_EXTRA + +$(eval $(golang-package)) diff --git a/src/confd/bin/firewall b/src/confd/bin/firewall index 28a9834d3..8c76f053e 100755 --- a/src/confd/bin/firewall +++ b/src/confd/bin/firewall @@ -6,6 +6,7 @@ DEST="org.fedoraproject.FirewallD1" OBJECT="/org/fedoraproject/FirewallD1" INTERFACE="org.fedoraproject.FirewallD1" +ADDRSET_DIR="/run/confd/address-sets" VERBOSE=0 print() { @@ -117,6 +118,42 @@ ipset_call() fi } +# Dynamic address-set entries only exist in the runtime config; a reload +# rebuilds the sets from the generated ipset XML, which holds static +# entries only. Re-apply the dynamic entries tracked by confd's +# add/remove action handlers, so they survive the reload without being +# baked into the XML (which would resurrect entries removed while the +# reload was in flight). +# +# An entry firewalld rejects as invalid, e.g., one now overlapping a +# static entry, is dropped from the shadow file, or it could never be +# removed with the remove action again. +addrset_resync() +{ + for file in "$ADDRSET_DIR"/*; do + case "$file" in *.resync) continue ;; esac + [ -f "$file" ] || continue + name=$(basename "$file") + keep="$file.resync" + : > "$keep" + + while IFS= read -r entry; do + [ -n "$entry" ] || continue + if ! ipset_call addEntry "$name" "$entry"; then + case "$output" in + *INVALID_ENTRY*) + logger -t firewall -p daemon.warn "ipset $name: dropping rejected dynamic entry $entry" + continue + ;; + esac + fi + printf '%s\n' "$entry" >> "$keep" + done < "$file" + + mv "$keep" "$file" + done +} + panic_status() { if is_panic_enabled; then @@ -377,6 +414,8 @@ main() exit 1 fi fi + + addrset_resync ;; panic) if ! check_firewalld; then diff --git a/src/confd/src/core.c b/src/confd/src/core.c index b89cf9e02..e08f5d195 100644 --- a/src/confd/src/core.c +++ b/src/confd/src/core.c @@ -805,6 +805,10 @@ static int change_cb(sr_session_ctx_t *session, uint32_t sub_id, const char *mod return SR_ERR_SYS; } + /* Nudge yangerd to re-poll, best effort: it may not be installed */ + if (systemf("initctl -bq reload yangerd")) + DEBUG("yangerd not reloaded, not running?"); + AUDIT("The new configuration has been applied."); } diff --git a/src/confd/src/firewall.c b/src/confd/src/firewall.c index be6a18311..ec4dc9a8e 100644 --- a/src/confd/src/firewall.c +++ b/src/confd/src/firewall.c @@ -106,29 +106,6 @@ static int prefix_parse(const char *str, struct prefix *p) return -1; } -static bool prefix_overlap(const char *a, const char *b) -{ - struct prefix pa, pb; - int len, i; - - if (prefix_parse(a, &pa) || prefix_parse(b, &pb) || pa.af != pb.af) - return false; - - len = pa.len < pb.len ? pa.len : pb.len; - for (i = 0; i < len / 8; i++) { - if (pa.addr[i] != pb.addr[i]) - return false; - } - if (len % 8) { - uint8_t mask = 0xff << (8 - len % 8); - - if ((pa.addr[i] & mask) != (pb.addr[i] & mask)) - return false; - } - - return true; -} - static bool shadow_has(const char *name, const char *entry) { char line[ENTRY_STRLEN]; @@ -371,47 +348,12 @@ static int generate_zone(struct lyd_node *cfg, const char *name, char **ifaces) } /* - * Dynamic entries, added at runtime with the add action, are folded - * into the generated ipset as regular entries so they survive the - * firewalld reload triggered by configuration changes. Entries that - * overlap new static configuration are dropped -- config wins, and - * nftables refuses overlapping elements in interval sets. + * Only static entries go into the generated ipset. Dynamic entries, + * added at runtime with the add action, are re-applied from the shadow + * files by 'firewall reload' after firewalld has reloaded. Baking them + * into the XML would resurrect entries removed while a reload was in + * flight -- the reload is asynchronous to the action handlers. */ -static void merge_dynamic(FILE *fp, struct lyd_node *cfg, const char *name) -{ - char line[ENTRY_STRLEN]; - FILE *sf; - - sf = fopenf("r", ADDRSET_RUNDIR "/%s", name); - if (!sf) - return; - - while (fgets(line, sizeof(line), sf)) { - struct lyd_node *node; - bool skip = false; - - chomp(line); - if (!line[0]) - continue; - - LYX_LIST_FOR_EACH(lyd_child(cfg), node, "entry") { - if (prefix_overlap(line, lyd_get_value(node))) { - skip = true; - break; - } - } - - if (skip) { - NOTE("address-set %s: dropping dynamic entry %s, overlaps static entry", - name, line); - continue; - } - - fprintf(fp, " %s\n", line); - } - fclose(sf); -} - static int generate_ipset(struct lyd_node *cfg, const char *name) { const char *family, *timeout, *desc; @@ -441,9 +383,6 @@ static int generate_ipset(struct lyd_node *cfg, const char *name) LYX_LIST_FOR_EACH(lyd_child(cfg), node, "entry") fprintf(fp, " %s\n", lyd_get_value(node)); - if (!timeout) - merge_dynamic(fp, cfg, name); - fprintf(fp, "\n"); return close_file(fp); @@ -745,10 +684,17 @@ int firewall_change(sr_session_ctx_t *session, struct lyd_node *config, struct l return SR_ERR_OK; } - /* Drop dynamic state of deleted address-sets */ + /* + * Drop dynamic state of deleted address-sets, and of sets + * that got a timeout: their entries expire on their own and + * must not be re-applied on reload. + */ clist = lydx_get_descendant(diff, "firewall", "address-set", NULL); LYX_LIST_FOR_EACH(clist, cnode, "address-set") { - if (lydx_get_op(cnode) == LYDX_OP_DELETE) + struct lyd_node *timeout = lydx_get_child(cnode, "timeout"); + + if (lydx_get_op(cnode) == LYDX_OP_DELETE || + (timeout && lydx_get_op(timeout) != LYDX_OP_DELETE)) erasef(ADDRSET_RUNDIR "/%s", lydx_get_cattr(cnode, "name")); } @@ -861,7 +807,7 @@ int firewall_change(sr_session_ctx_t *session, struct lyd_node *config, struct l LYX_LIST_FOR_EACH(clist, cnode, "service") generate_service(cnode, lydx_get_cattr(cnode, "name")); - /* Regenerate all address-sets, incl. dynamic entries */ + /* Regenerate all address-sets (static entries only) */ clist = lydx_get_descendant(tree, "firewall", "address-set", NULL); LYX_LIST_FOR_EACH(clist, cnode, "address-set") generate_ipset(cnode, lydx_get_cattr(cnode, "name")); diff --git a/src/statd/Makefile.am b/src/statd/Makefile.am index 3c5fddbfb..0a68dd6f7 100644 --- a/src/statd/Makefile.am +++ b/src/statd/Makefile.am @@ -2,7 +2,7 @@ DISTCLEANFILES = *~ *.d ACLOCAL_AMFLAGS = -I m4 sbin_PROGRAMS = statd -statd_SOURCES = statd.c shared.c shared.h journal.c journal_retention.c journal.h avahi.c avahi.h iface.c iface.h +statd_SOURCES = statd.c shared.c shared.h journal.c journal_retention.c journal.h avahi.c avahi.h iface.c iface.h yangerd.c yangerd.h statd_CPPFLAGS = -D_DEFAULT_SOURCE -D_GNU_SOURCE statd_CPPFLAGS += -DSTATD_VERSION=\"$(PACKAGE_VERSION)\" statd_CFLAGS = -W -Wall -Wextra diff --git a/src/statd/statd.c b/src/statd/statd.c index 34757fd12..ea6a685a4 100644 --- a/src/statd/statd.c +++ b/src/statd/statd.c @@ -14,43 +14,28 @@ #include #include -#include -#include #include #include #include #include #include -#include #include #include #include -#include #include "shared.h" #include "journal.h" -#include "iface.h" #include "avahi.h" +#include "yangerd.h" -/* New kernel feature, not in sys/mman.h yet */ -#ifndef MFD_NOEXEC_SEAL -#define MFD_NOEXEC_SEAL 0x0008U -#endif - -#define YANGER_BINPATH YANGER_DIR"/yanger" #define XPATH_MAX PATH_MAX #define XPATH_IFACE_BASE "/ietf-interfaces:interfaces" -#define XPATH_ROUTING_BASE "/ietf-routing:routing/control-plane-protocols/control-plane-protocol" +#define XPATH_ROUTING_PROTOCOLS "/ietf-routing:routing/control-plane-protocols" #define XPATH_ROUTING_TABLE "/ietf-routing:routing/ribs" #define XPATH_HARDWARE_BASE "/ietf-hardware:hardware" #define XPATH_SYSTEM_BASE "/ietf-system" -#ifdef HAVE_FRR -#define XPATH_ROUTING_OSPF XPATH_ROUTING_BASE "/ospf" -#define XPATH_ROUTING_RIP XPATH_ROUTING_BASE "/rip" -#define XPATH_ROUTING_BFD XPATH_ROUTING_BASE "/bfd" -#endif #define XPATH_CONTAIN_BASE "/infix-containers:containers" #define XPATH_DHCP_SERVER_BASE "/infix-dhcp-server:dhcp-server" #define XPATH_TFTP_FILES "/infix-services:tftp/files" @@ -64,6 +49,8 @@ TAILQ_HEAD(sub_head, sub); struct sub { struct ev_io watcher; sr_subscription_ctx_t *sr_sub; + char key[XPATH_MAX]; /* yangerd key, derived from the subscription xpath */ + struct statd *statd; /* owning daemon context */ TAILQ_ENTRY(sub) entries; @@ -76,261 +63,220 @@ struct statd { struct ev_loop *ev_loop; struct journal_ctx journal; /* Periodic operational snapshots */ struct mdns_ctx mdns; /* mDNS neighbor monitor */ - struct iface_ctx iface; /* Interface state change tracking */ }; -static int ly_add_yanger_data(const struct ly_ctx *ctx, struct lyd_node **parent, - char *yanger_args[]) +/* + * The name of the node a subscription provides, the last step of its + * path without prefix or predicates: "/ietf-routing:routing/ribs" -> "ribs" + */ +static const char *sub_node_name(const char *path, char *buf, size_t len) { - FILE *stream; - int err; - int fd; - - fd = memfd_create("yanger_tmpfile", MFD_CLOEXEC | MFD_NOEXEC_SEAL); - if (fd == -1) { - ERROR("Error, unable to create memfd"); - return SR_ERR_SYS; - } + const char *p, *colon; + size_t n; + + p = strrchr(path, '/'); + p = p ? p + 1 : path; + colon = strchr(p, ':'); + if (colon) + p = colon + 1; + + n = strcspn(p, "["); + if (n >= len) + n = len - 1; + memcpy(buf, p, n); + buf[n] = 0; + + return buf; +} - /* Wrap the file descriptor in a FILE stream for fwrite */ - stream = fdopen(fd, "w+"); - if (stream == NULL) { - ERROR("Error, unable to fdopen memfd"); - close(fd); - return SR_ERR_SYS; +/* + * For a nested subscription sysrepo hands us the parent instance and + * expects the requested nodes appended to it. yangerd answers with the + * whole module tree, so find the same parent in it and move over only + * the children this subscription provides. + */ +static int graft(struct lyd_node *parent, struct lyd_node *tree, const char *path) +{ + struct lyd_node *match, *node, *next; + char name[64]; + char *xpath; + LY_ERR err; + + xpath = lyd_path(parent, LYD_PATH_STD, NULL, 0); + if (!xpath) + return SR_ERR_NO_MEMORY; + + err = lyd_find_path(tree, xpath, 0, &match); + if (err == LY_ENOTFOUND || err == LY_EINCOMPLETE) { + free(xpath); + return SR_ERR_OK; } - - err = fsystemv(yanger_args, NULL, stream, NULL); if (err) { - ERROR("Error calling yanger %s%s%s, exit code %d", yanger_args[1], - yanger_args[3] ? " " : "", yanger_args[3] ?: "", err); - fclose(stream); - return SR_ERR_SYS; + ERROR("yangerd: cannot find %s in its answer: %s", xpath, ly_last_logmsg()); + free(xpath); + return SR_ERR_INTERNAL; } - - fflush(stream); - - if (lseek(fd, 0, SEEK_SET) == (off_t)-1) { - ERROR("Error, unable reset stream (seek)"); - fclose(stream); - return SR_ERR_SYS; + free(xpath); + + sub_node_name(path, name, sizeof(name)); + LY_LIST_FOR_SAFE(lyd_child(match), next, node) { + if (strcmp(node->schema->name, name)) + continue; + + lyd_unlink_tree(node); + if (lyd_insert_child(parent, node)) { + ERROR("yangerd: cannot add %s for %s: %s", name, path, ly_last_logmsg()); + lyd_free_tree(node); + return SR_ERR_INTERNAL; + } } - err = lyd_parse_data_fd(ctx, fd, LYD_JSON, LYD_PARSE_ONLY, 0, parent); - if (err) - ERROR("Failed parsing yanger data (%d): %s", err, ly_errmsg(ctx)); - - fclose(stream); - /* Note: fclose() already closes the underlying fd from fdopen() */ - - return err; + return SR_ERR_OK; } -static char *xpath_extract(const char *xpath, const char *key) +/* + * One bad value must not cost every GET of the module. When yangerd's + * answer does not parse, try each entry of every list under the top + * container on its own, drop the ones libyang rejects, loudly, and parse + * the rest. Only runs on the error path. + */ +static int parse_salvage(const struct ly_ctx *ctx, const char *json, const char *key, + struct lyd_node **tree) { - char *res = NULL; - const char *ptr; - const char *end; - - /* (also checks if key exist) */ - ptr = strstr(xpath, key); - if (!ptr) - return NULL; - - ptr += strlen(key); - - end = strchr(ptr, '\''); - if (!end) { - ERROR("Cannot find end quote for %s (sanity check)", key); - return NULL; - } + json_t *root, *top, *list, *good, *one, *entry, *name; + const char *lname; + size_t i, dropped = 0; + char *text; + int rc = -1; + + root = json_loads(json, 0, NULL); + top = json_object_get(root, key); + if (!json_is_object(top)) + goto out; + + json_object_foreach(top, lname, list) { + if (!json_is_array(list)) + continue; + + good = json_array(); + json_array_foreach(list, i, entry) { + struct lyd_node *probe = NULL; + LY_ERR err; + + one = json_pack("{s:{s:[O]}}", key, lname, entry); + text = json_dumps(one, JSON_COMPACT); + json_decref(one); + err = text ? lyd_parse_data_mem(ctx, text, LYD_JSON, LYD_PARSE_ONLY, 0, &probe) : LY_EMEM; + free(text); + lyd_free_all(probe); + if (!err) { + json_array_append(good, entry); + continue; + } - if ((end - ptr) >= XPATH_MAX) { - ERROR("Value for %s is too long (sanity check)", key); - return NULL; + name = json_object_get(entry, "name"); + ERROR("yangerd: dropping %s %s[%s]: %s", key, lname, + json_is_string(name) ? json_string_value(name) : "?", ly_last_logmsg()); + dropped++; + } + json_object_set_new(top, lname, good); } - res = calloc((end - ptr) + 1, sizeof(char)); - if (!res) - return NULL; + if (!dropped) + goto out; - strncpy(res, ptr, end - ptr); - res[end - ptr] = '\0'; - - return res; + text = json_dumps(root, JSON_COMPACT); + if (text && !lyd_parse_data_mem(ctx, text, LYD_JSON, LYD_PARSE_ONLY, 0, tree)) + rc = 0; + free(text); +out: + json_decref(root); + return rc; } -static int sr_iface_cb(sr_session_ctx_t *session, uint32_t, const char *model, - const char *, const char *xpath, uint32_t, - struct lyd_node **parent, void *priv) +static int ly_add_yangerd_data(const struct ly_ctx *ctx, struct lyd_node **parent, + const char *path, const char *key) { - char *yanger_args[5] = { - YANGER_BINPATH, - (char *)model, - NULL, - NULL, - NULL - }; - struct statd *statd = priv; - char *ifname = NULL; - const struct ly_ctx *ctx; - sr_conn_ctx_t *con; - int err; - - DEBUG("Incoming interface query for xpath: %s", xpath); - - con = sr_session_get_connection(session); - if (!con) { - ERROR("Error getting sysrepo connection"); - return SR_ERR_INTERNAL; + struct lyd_node *tree = NULL; + char *json = NULL; + size_t len = 0; + int rc; + + rc = yangerd_query(key, &json, &len); + if (rc > 0) { + WARN("yangerd: no data for %s yet, not ready", key); + return SR_ERR_OK; } - - ctx = sr_acquire_context(con); - if (!ctx) { - ERROR("Failed acquiring sysrepo context"); + if (rc) { + ERROR("yangerd: query failed for %s", key); return SR_ERR_INTERNAL; } - ifname = xpath_extract(xpath, "[name='"); - if (ifname) { - yanger_args[2] = "-p"; - yanger_args[3] = ifname; + DEBUG("yangerd: got %zu bytes JSON for %s", len, key); + if (!json || !len) { + free(json); + return SR_ERR_OK; /* feature not active, no data */ } - err = ly_add_yanger_data(ctx, parent, yanger_args); - if (err) - ERROR("Failed adding yanger data for %s", ifname ?: model); - else - iface_annotate(&statd->iface, *parent); - - free(ifname); - sr_release_context(con); - - return SR_ERR_OK; -} -static int sr_generic_cb(sr_session_ctx_t *session, uint32_t, const char *model, - const char *, const char *xpath, uint32_t, - struct lyd_node **parent, __attribute__((unused)) void *priv) -{ - char *yanger_args[5] = { - YANGER_BINPATH, - (char *)model, - NULL - }; - const struct ly_ctx *ctx; - sr_conn_ctx_t *con; - sr_error_t err; - - DEBUG("Incoming generic query for xpath: %s", xpath); - - con = sr_session_get_connection(session); - if (!con) { - ERROR("Error getting sysrepo connection"); - return SR_ERR_INTERNAL; + if (lyd_parse_data_mem(ctx, json, LYD_JSON, LYD_PARSE_ONLY, 0, &tree)) { + ERROR("Failed parsing yangerd data for %s: %s", key, ly_errmsg(ctx)); + tree = NULL; + if (parse_salvage(ctx, json, key, &tree)) { + free(json); + return SR_ERR_INTERNAL; + } } + free(json); + if (!tree) + return SR_ERR_OK; /* "{}", nothing to add */ - ctx = sr_acquire_context(con); - if (!ctx) { - ERROR("Failed acquiring sysrepo context"); - return SR_ERR_INTERNAL; + if (!*parent) { + *parent = tree; + return SR_ERR_OK; } - err = ly_add_yanger_data(ctx, parent, yanger_args); - if (err) - ERROR("Failed adding yanger data for %s", yanger_args[1]); + rc = graft(*parent, tree, path); + lyd_free_all(tree); - sr_release_context(con); - - return err; + return rc; } -#ifdef HAVE_FRR -static int sr_ospf_cb(sr_session_ctx_t *session, uint32_t, const char *, - const char *, const char *xpath, uint32_t, - struct lyd_node **parent, __attribute__((unused)) void *priv) +static const char *xpath_to_yangerd_path(const char *xpath, char *buf, size_t bufsz) { - char *yanger_args[5] = { - YANGER_BINPATH, - "ietf-ospf", - NULL - }; - const struct ly_ctx *ctx; - sr_conn_ctx_t *con; - sr_error_t err; - - DEBUG("Incoming ospf query for xpath: %s", xpath); + const char *start, *slash; + size_t len; - con = sr_session_get_connection(session); - if (!con) { - ERROR("Error getting sysrepo connection"); - return SR_ERR_INTERNAL; + if (!xpath || !*xpath || !strcmp(xpath, "*") || !strcmp(xpath, "/*")) { + buf[0] = '\0'; + return buf; } - ctx = sr_acquire_context(con); - if (!ctx) { - ERROR("Failed acquiring sysrepo context"); - return SR_ERR_INTERNAL; - } + start = xpath; + if (*start == '/') + start++; - err = ly_add_yanger_data(ctx, parent, yanger_args); - if (err) - ERROR("Failed adding yanger data for %s", yanger_args[1]); + slash = strchr(start, '/'); + len = slash ? (size_t)(slash - start) : strlen(start); - sr_release_context(con); + if (len >= bufsz) + len = bufsz - 1; - return err; -} + memcpy(buf, start, len); + buf[len] = '\0'; -static int sr_rip_cb(sr_session_ctx_t *session, uint32_t, const char *, - const char *, const char *xpath, uint32_t, - struct lyd_node **parent, __attribute__((unused)) void *priv) -{ - char *yanger_args[5] = { - YANGER_BINPATH, - "ietf-rip", - NULL - }; - const struct ly_ctx *ctx; - sr_conn_ctx_t *con; - sr_error_t err; - - DEBUG("Incoming RIP query for xpath: %s", xpath); - - con = sr_session_get_connection(session); - if (!con) { - ERROR("Error getting sysrepo connection"); - return SR_ERR_INTERNAL; - } - - ctx = sr_acquire_context(con); - if (!ctx) { - ERROR("Failed acquiring sysrepo context"); - return SR_ERR_INTERNAL; - } - - err = ly_add_yanger_data(ctx, parent, yanger_args); - if (err) - ERROR("Failed adding yanger data for %s", yanger_args[1]); - - sr_release_context(con); - - return err; + return buf; } -static int sr_bfd_cb(sr_session_ctx_t *session, uint32_t, const char *, - const char *, const char *xpath, uint32_t, - struct lyd_node **parent, __attribute__((unused)) void *priv) +static int sr_generic_cb(sr_session_ctx_t *session, uint32_t, const char *, + const char *path, const char *xpath, uint32_t, + struct lyd_node **parent, void *priv) { - char *yanger_args[5] = { - YANGER_BINPATH, - "ietf-bfd-ip-sh", - NULL - }; + struct sub *sub = priv; const struct ly_ctx *ctx; sr_conn_ctx_t *con; - sr_error_t err; + int err; - DEBUG("Incoming BFD query for xpath: %s", xpath); + DEBUG("Incoming query for xpath: %s -> key %s", xpath, sub->key); con = sr_session_get_connection(session); if (!con) { @@ -344,15 +290,11 @@ static int sr_bfd_cb(sr_session_ctx_t *session, uint32_t, const char *, return SR_ERR_INTERNAL; } - err = ly_add_yanger_data(ctx, parent, yanger_args); - if (err) - ERROR("Failed adding yanger data for %s", yanger_args[1]); - + err = ly_add_yangerd_data(ctx, parent, path, sub->key); sr_release_context(con); return err; } -#endif /* HAVE_FRR */ static void sigint_cb(struct ev_loop *loop, struct ev_signal *, int) @@ -380,9 +322,9 @@ static void sr_event_cb(struct ev_loop *, struct ev_io *w, int) sr_subscription_process_events(sub->sr_sub, NULL, NULL); } -static int subscribe(struct statd *statd, char *model, char *xpath, - int (*cb)(sr_session_ctx_t *session, uint32_t, const char *, const char *, - const char *, uint32_t, struct lyd_node **parent, void *priv)) +static int subscribe_opts(struct statd *statd, char *model, char *xpath, uint32_t opts, + int (*cb)(sr_session_ctx_t *session, uint32_t, const char *, const char *, + const char *, uint32_t, struct lyd_node **parent, void *priv)) { struct sub *sub; int sr_ev_pipe; @@ -390,10 +332,20 @@ static int subscribe(struct statd *statd, char *model, char *xpath, sub = malloc(sizeof(struct sub)); memset(sub, 0, sizeof(struct sub)); - - DEBUG("Subscribe to events for \"%s\"", xpath); - err = sr_oper_get_subscribe(statd->sr_ses, model, xpath, cb, statd, - SR_SUBSCR_DEFAULT | SR_SUBSCR_NO_THREAD | SR_SUBSCR_DONE_ONLY, + sub->statd = statd; + + /* + * Derive the yangerd key from the (static) subscription xpath here, + * once. The generic callback must NOT derive it from the runtime + * request xpath sysrepo hands it -- that is unreliable and yields a + * bare "system-state" for /ietf-system:system-state, which yangerd + * (keyed "ietf-system:system-state") cannot match. + */ + xpath_to_yangerd_path(xpath, sub->key, sizeof(sub->key)); + + DEBUG("Subscribe to events for \"%s\" (key \"%s\")", xpath, sub->key); + err = sr_oper_get_subscribe(statd->sr_ses, model, xpath, cb, sub, + SR_SUBSCR_DEFAULT | SR_SUBSCR_NO_THREAD | SR_SUBSCR_DONE_ONLY | opts, &sub->sr_sub); if (err) { ERROR("Failed subscribing to path \"%s\": %s", xpath, sr_strerror(err)); @@ -418,6 +370,13 @@ static int subscribe(struct statd *statd, char *model, char *xpath, return SR_ERR_OK; } +static int subscribe(struct statd *statd, char *model, char *xpath, + int (*cb)(sr_session_ctx_t *session, uint32_t, const char *, const char *, + const char *, uint32_t, struct lyd_node **parent, void *priv)) +{ + return subscribe_opts(statd, model, xpath, 0, cb); +} + static void sub_delete(struct ev_loop *loop, struct sub_head *subs, struct sub *sub) { TAILQ_REMOVE(subs, sub, entries); @@ -442,14 +401,15 @@ static int subscribe_to_all(struct statd *statd) if (subscribe(statd, "ietf-routing", XPATH_ROUTING_TABLE, sr_generic_cb)) return SR_ERR_INTERNAL; - if (subscribe(statd, "ietf-interfaces", XPATH_IFACE_BASE, sr_iface_cb)) + if (subscribe(statd, "ietf-interfaces", XPATH_IFACE_BASE, sr_generic_cb)) return SR_ERR_INTERNAL; #ifdef HAVE_FRR - if (subscribe(statd, "ietf-routing", XPATH_ROUTING_OSPF, sr_ospf_cb)) - return SR_ERR_INTERNAL; - if (subscribe(statd, "ietf-routing", XPATH_ROUTING_RIP, sr_rip_cb)) - return SR_ERR_INTERNAL; - if (subscribe(statd, "ietf-routing", XPATH_ROUTING_BFD, sr_bfd_cb)) + /* + * Merged, not replaced: the BFD instance exists only in operational, + * and config-only instances like static routes must stay. One call + * per GET, not one per control-plane-protocol instance. + */ + if (subscribe_opts(statd, "ietf-routing", XPATH_ROUTING_PROTOCOLS, SR_SUBSCR_OPER_MERGE, sr_generic_cb)) return SR_ERR_INTERNAL; #endif if (subscribe(statd, "ietf-hardware", XPATH_HARDWARE_BASE, sr_generic_cb)) @@ -610,9 +570,6 @@ int main(int argc, char *argv[]) if (mdns_ctx_init(&statd.mdns, statd.ev_loop, statd.sr_conn)) INFO("mDNS neighbor monitoring not available"); - if (iface_ctx_init(&statd.iface, statd.ev_loop)) - WARN("Interface state change tracking not available"); - /* Signal readiness to Finit */ pidfile(NULL); @@ -622,7 +579,6 @@ int main(int argc, char *argv[]) /* We should never get here during normal operation */ INFO("Status daemon shutting down"); - iface_ctx_exit(&statd.iface); mdns_ctx_exit(&statd.mdns); journal_stop(&statd.journal); diff --git a/src/statd/yangerd.c b/src/statd/yangerd.c new file mode 100644 index 000000000..618e069c2 --- /dev/null +++ b/src/statd/yangerd.c @@ -0,0 +1,269 @@ +/* SPDX-License-Identifier: BSD-3-Clause */ + +#include +#include +#include +#include +#include +#include +#include +#include + +#include +#include + +#include "yangerd.h" + +static const char *yangerd_socket_path(void) +{ + const char *env; + + env = getenv("YANGERD_SOCKET"); + if (env && *env) + return env; + + return YANGERD_SOCKET_DEFAULT; +} + +static int yangerd_connect(void) +{ + struct sockaddr_un addr = { .sun_family = AF_UNIX }; + struct timeval tv = { .tv_sec = YANGERD_TIMEOUT_SEC }; + const char *path; + int fd; + + path = yangerd_socket_path(); + if (strlen(path) >= sizeof(addr.sun_path)) { + ERROR("yangerd socket path too long: %s", path); + return -1; + } + strncpy(addr.sun_path, path, sizeof(addr.sun_path) - 1); + + fd = socket(AF_UNIX, SOCK_STREAM | SOCK_CLOEXEC, 0); + if (fd < 0) { + ERROR("yangerd: socket(): %s", strerror(errno)); + return -1; + } + + setsockopt(fd, SOL_SOCKET, SO_RCVTIMEO, &tv, sizeof(tv)); + setsockopt(fd, SOL_SOCKET, SO_SNDTIMEO, &tv, sizeof(tv)); + + if (connect(fd, (struct sockaddr *)&addr, sizeof(addr)) < 0) { + int err = errno; + + DEBUG("yangerd: connect(%s): %s", path, strerror(err)); + close(fd); + /* Not started yet, or restarting */ + if (err == ENOENT || err == ECONNREFUSED) + return -2; + return -1; + } + + return fd; +} + +static int yangerd_write_all(int fd, const void *buf, size_t len) +{ + const unsigned char *p = buf; + + while (len > 0) { + ssize_t n = write(fd, p, len); + + if (n < 0) { + if (errno == EINTR) + continue; + return -1; + } + p += n; + len -= n; + } + + return 0; +} + +static int yangerd_read_all(int fd, void *buf, size_t len) +{ + unsigned char *p = buf; + + while (len > 0) { + ssize_t n = read(fd, p, len); + + if (n < 0) { + if (errno == EINTR) + continue; + return -1; + } + if (n == 0) { + errno = ECONNRESET; + return -1; + } + p += n; + len -= n; + } + + return 0; +} + +static int yangerd_send_request(int fd, const char *path) +{ + json_t *req; + char *json_str; + size_t json_len; + unsigned char hdr[5]; + int rc = -1; + + req = json_pack("{s:s, s:s}", "method", "get", "path", path); + if (!req) + return -1; + + json_str = json_dumps(req, JSON_COMPACT); + json_decref(req); + if (!json_str) + return -1; + + json_len = strlen(json_str); + if (json_len > YANGERD_MAX_PAYLOAD) { + free(json_str); + return -1; + } + + hdr[0] = YANGERD_PROTO_VERSION; + hdr[1] = (json_len >> 24) & 0xff; + hdr[2] = (json_len >> 16) & 0xff; + hdr[3] = (json_len >> 8) & 0xff; + hdr[4] = (json_len >> 0) & 0xff; + + if (yangerd_write_all(fd, hdr, sizeof(hdr)) < 0) + goto out; + if (yangerd_write_all(fd, json_str, json_len) < 0) + goto out; + + rc = 0; +out: + free(json_str); + return rc; +} + +static int yangerd_recv_frame(int fd, char **buf, size_t *len) +{ + unsigned char hdr[5]; + uint32_t payload_len; + char *body; + + if (yangerd_read_all(fd, hdr, sizeof(hdr)) < 0) + return -1; + + if (hdr[0] != YANGERD_PROTO_VERSION) { + ERROR("yangerd: protocol version mismatch: got %u, want %u", + hdr[0], YANGERD_PROTO_VERSION); + return -1; + } + + payload_len = ((uint32_t)hdr[1] << 24) | + ((uint32_t)hdr[2] << 16) | + ((uint32_t)hdr[3] << 8) | + ((uint32_t)hdr[4]); + + if (payload_len > YANGERD_MAX_PAYLOAD) { + ERROR("yangerd: payload too large: %u", payload_len); + return -1; + } + + body = malloc(payload_len + 1); + if (!body) + return -1; + + if (yangerd_read_all(fd, body, payload_len) < 0) { + free(body); + return -1; + } + body[payload_len] = '\0'; + + *buf = body; + *len = payload_len; + + return 0; +} + +/* + * The response header is a small JSON object, the data, if any, + * follows raw in a frame of its own and goes to the caller untouched. + */ +static int yangerd_recv_response(int fd, char **buf, size_t *len) +{ + json_t *resp, *status, *raw; + json_error_t jerr; + size_t hdr_len; + char *hdr; + int rc; + + *buf = NULL; + *len = 0; + + if (yangerd_recv_frame(fd, &hdr, &hdr_len)) + return -1; + + resp = json_loadb(hdr, hdr_len, 0, &jerr); + free(hdr); + if (!resp) { + ERROR("yangerd: invalid response JSON: %s", jerr.text); + return -1; + } + + status = json_object_get(resp, "status"); + if (json_is_string(status) && !strcmp(json_string_value(status), "starting")) { + json_decref(resp); + return 1; + } + if (!json_is_string(status) || strcmp(json_string_value(status), "ok")) { + json_t *msg = json_object_get(resp, "message"); + + ERROR("yangerd: request failed: %s", + json_is_string(msg) ? json_string_value(msg) : "unknown"); + json_decref(resp); + return -1; + } + + raw = json_object_get(resp, "raw"); + rc = json_is_true(raw); + json_decref(resp); + + if (rc) + return yangerd_recv_frame(fd, buf, len); + + *buf = strdup("{}"); + if (!*buf) + return -1; + *len = 2; + + return 0; +} + +int yangerd_query(const char *path, char **buf, size_t *len) +{ + int fd; + int rc; + + *buf = NULL; + *len = 0; + + fd = yangerd_connect(); + if (fd == -2) + return 1; + if (fd < 0) + return -1; + + if (yangerd_send_request(fd, path) < 0) { + ERROR("yangerd: failed sending request for %s", path); + close(fd); + return -1; + } + + rc = yangerd_recv_response(fd, buf, len); + if (rc < 0) + ERROR("yangerd: failed reading response for %s", path); + + close(fd); + + return rc; +} diff --git a/src/statd/yangerd.h b/src/statd/yangerd.h new file mode 100644 index 000000000..a7275c102 --- /dev/null +++ b/src/statd/yangerd.h @@ -0,0 +1,29 @@ +/* SPDX-License-Identifier: BSD-3-Clause */ + +#ifndef STATD_YANGERD_H_ +#define STATD_YANGERD_H_ + +#include + +#define YANGERD_SOCKET_DEFAULT "/run/yangerd.sock" +#define YANGERD_TIMEOUT_SEC 5 +#define YANGERD_MAX_PAYLOAD (4 << 20) /* 4 MiB, matches Go side */ +#define YANGERD_PROTO_VERSION 0x02 + +/** + * yangerd_query() - Query yangerd daemon for operational YANG data + * @path: YANG model path, e.g. "ietf-interfaces:interfaces" + * @buf: Output pointer to malloc'd JSON string (caller must free) + * @len: Output length of JSON data + * + * Connects to the yangerd Unix socket, sends a "get" request for @path, + * and returns the data frame of the response. The socket path defaults + * to %YANGERD_SOCKET_DEFAULT but can be overridden with the + * YANGERD_SOCKET environment variable. + * + * Return: 0 on success, 1 when yangerd is not running or still starting + * (no data yet), -1 on error. buf is NULL unless 0 is returned. + */ +int yangerd_query(const char *path, char **buf, size_t *len); + +#endif diff --git a/src/yangerd/.gitignore b/src/yangerd/.gitignore new file mode 100644 index 000000000..d5b437a9f --- /dev/null +++ b/src/yangerd/.gitignore @@ -0,0 +1,3 @@ +# Build artifacts (root-level binaries only; do not match cmd/ source dirs) +/yangerd +/yangerctl diff --git a/src/yangerd/LICENSE b/src/yangerd/LICENSE new file mode 100644 index 000000000..bf7aa8c9e --- /dev/null +++ b/src/yangerd/LICENSE @@ -0,0 +1,27 @@ +Copyright (c) 2025 The KernelKit Authors +All rights reserved. + +Redistribution and use in source and binary forms, with or without +modification, are permitted provided that the following conditions are met: + +* Redistributions of source code must retain the above copyright notice, this + list of conditions and the following disclaimer. + +* Redistributions in binary form must reproduce the above copyright notice, + this list of conditions and the following disclaimer in the documentation + and/or other materials provided with the distribution. + +* Neither the name of copyright holders nor the names of + contributors may be used to endorse or promote products derived from + this software without specific prior written permission. + +THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" +AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE +IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE +DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE +FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL +DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR +SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER +CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, +OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE +OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. diff --git a/src/yangerd/README.md b/src/yangerd/README.md new file mode 100644 index 000000000..f59fee83b --- /dev/null +++ b/src/yangerd/README.md @@ -0,0 +1,244 @@ +# yangerd + +yangerd collects Infix operational data and keeps it ready as YANG-shaped +JSON. statd asks it for a module's subtree when sysrepo asks statd, and +hands the answer to libyang. yangerd replaces the Python yanger scripts, +which forked a collector per GET. + +``` + NETCONF / RESTCONF client + | + sysrepo <-- oper-get subscriptions -- statd (C) + | + /run/yangerd.sock (IPC) + | + yangerd (Go) + | + netlink, nl80211, ethtool, ZAPI, D-Bus, wpa_supplicant/hostapd, + ptp4l, chronyd, lldpcli, podman, FRR vty, inotify, sysfs/procfs +``` + +statd still owns everything sysrepo-facing: subscriptions, libyang +parsing and error codes. yangerd has no sysrepo or libyang dependency. +It only knows the JSON its consumers expect. + + +## Design decisions + +**Reactive first.** A value with an event source is updated when the +event fires, never on a timer. Netlink, nl80211, ZAPI, D-Bus signals, +wpa_supplicant and hostapd attach events, ptp4l subscriptions, inotify +and `podman events` all push. Between events a monitor does nothing, and a GET returns +what the last event left without forking anything. + +**Poll only what has no events.** FRR's OSPF, RIP and BFD state, the +service list and gpsd don't announce changes, so they are polled. The intervals are in the table below and can be +changed through the environment. + +**On demand when a value changes all the time, or costs nothing to +read.** Sensor readings, radio channel surveys, NTP source selection, +uptime, memory and load are computed at GET time by a tree provider, and +so is the hardware inventory, which is a read of `/run/system.json` and +sysfs. chronyd is asked over its local command socket, so a GET sees a +source the moment chrony selects it. +Polling a temperature every ten seconds that nobody reads is wasted +work. Providers run on the request path, so they must be cheap. + +**One JSON blob per module.** The tree (`internal/tree`) stores one +document per top-level YANG node, keyed like `ietf-interfaces:interfaces`, +each behind its own lock. Writers replace (`Set`) or shallow-merge +(`Merge`) their part. A key several sources share, such as +`ietf-routing:routing`, has each source owning distinct top-level members +and merging. Providers overlay their output on the cached blob at read +time without changing it. + +**Absent means no data.** A feature that isn't active has no key. +yangerd deletes the key rather than storing an empty container, because +libyang would instantiate a presence container from `{}`. statd treats +an empty answer as "nothing to add", not as an error. + +**Long-lived helpers instead of forks.** `ip`, `ip -s -d` and `bridge` +run as `-json -force -batch -` subprocesses (`internal/ipbatch`), fed one +command per event. chronyd is queried over cmdmon, ptp4l over its +management socket, FRR over the daemons' vty sockets, all in-process. + +**Restart, don't die.** Every monitor runs under `backoff.Retry`: when +its source goes away (FRR restart, dnsmasq exit, `lldpcli watch` ending), +it reconnects with exponential backoff and rebuilds its state from a full +read. A monitor's failure never takes the daemon down. + +**Readiness is explicit.** yangerd writes `/run/yangerd.pid` after the +first netlink dump, which is when `notify:pid` in finit marks it ready. +Until then the IPC answers "starting", which statd turns into no data +plus a warning. + +**No CGo, Buildroot's Go, vendored.** Use only the Go version Buildroot +ships, 1.26 at the time of writing. Dependencies are vendored +(`GOFLAGS=-mod=vendor`). + + +## What is reactive, polled and on demand + +| Tree key | Content | How | Source | +|---|---|---|---| +| `ietf-interfaces:interfaces` | links, addresses, neighbours, bridge FDB/MDB, VLANs, LAG | event | rtnetlink, re-read with `ip`/`bridge` batch | +| | Ethernet speed, duplex, PMD | event | ethtool genetlink monitor, sweep at start | +| | WiFi station, AP, mesh point | event | nl80211 and wpa_supplicant/hostapd control sockets | +| | STP port and bridge state | poll 5 s | mstpd | +| | WireGuard peers | poll 10 s | `wg` netlink | +| | `last-change` | on demand | oper-state transitions seen since start | +| `ietf-routing:routing` | `ribs` | event | zebra ZAPI redistribution, or rtnetlink route events without FRR | +| | `control-plane-protocols` | poll 10 s | ospfd, ripd, bfdd vty sockets | +| | forwarding per interface | event | inotify on `/proc/sys/net/*/conf/*/forwarding` | +| `ietf-system:system` | hostname, timezone, users, SSH keys | event | inotify | +| `ietf-system:system-state` | services | poll 60 s | `initctl -j` | +| | software slots, boot order | event | RAUC D-Bus signal, bootloader env files | +| | platform | once | at start | +| | clock, memory, load, filesystems, installer | on demand | procfs, statfs, RAUC D-Bus | +| | NTP sources | on demand | chronyd cmdmon | +| `ietf-hardware:hardware` | mainboard, VPD, USB ports | on demand | `/run/system.json`, sysfs | +| | radio capabilities | event | nl80211 phy, regulatory and interface events | +| | radio channel survey | on demand | nl80211 survey dump | +| | sensors | on demand | hwmon, thermal zones | +| | GPS receivers | poll 10 s | gpsd | +| `ietf-ntp:ntp` | associations, clock state, server stats | on demand | chronyd cmdmon | +| | presence, listening port | poll 60 s | chronyd cmdmon, `ss` | +| `ieee1588-ptp-tt:ptp` | port state, time status, parent | event | ptp4l subscription, near-static sets refreshed every 30 s | +| `ieee802-dot1ab-lldp:lldp` | neighbours | event | `lldpcli watch` | +| `infix-containers:containers` | containers | event | `podman events`, re-read with `podman ps` | +| `infix-dhcp-server:dhcp-server` | leases | event | dnsmasq D-Bus signals | +| `infix-firewall:firewall` | zones, policies, services | event | firewalld D-Bus signals | +| | address-set entries | on demand | nft, only sets with timeouts | +| `infix-services:tftp` | files served | event | inotify on the root, mount table changes | + +A SIGHUP pokes every polled collector once. confd sends it after a +commit, best effort, so polled data catches up with new config. + + +## IPC + +AF_UNIX stream socket, `/run/yangerd.sock`, root only (0660, root:root). +One request per connection, every connection with a 5 s deadline. + +``` +| ver (1) | length (4, big endian) | JSON | +``` + +Version is 2. A request is `{"method": "get", "path": ""}`. +The response header is `{"status": "ok", "raw": true}`, and the data +follows in a second frame, so statd can pass it to +`lyd_parse_data_mem()` without decoding the envelope. `dump` returns all +keys, and `health` returns each key's size and last update time. The +maximum frame is 4 MiB. + +`yangerctl` speaks the same protocol for debugging: + +``` +yangerctl get /ietf-interfaces:interfaces +yangerctl dump +yangerctl health +``` + + +## The statd side + +statd subscribes once per module and derives the yangerd key from the +subscription xpath, not from the request. When you add a module, keep +in mind: + +- **Operational replaces running.** Subscriptions don't pass + `SR_SUBSCR_OPER_MERGE`, so what yangerd returns is all a client sees + under that path. Config-only leaves are absent. The exception is + `ietf-routing:routing/control-plane-protocols`, which is merged, + because the BFD instance exists only in operational and static route + instances only in running. +- **Nested subscriptions graft.** For a nested path such as + `/infix-services:tftp/files`, statd parses the whole module answer and + moves only the requested node under sysrepo's parent. +- **One bad value costs one entry.** If the answer doesn't parse, statd + parses each list entry alone, drops the ones libyang rejects, and logs + them as `yangerd: dropping []: `. If you see + that line, fix the value in yangerd. + + +## Configuration + +`/etc/default/yangerd`, written by `package/yangerd/yangerd.mk` from the +Buildroot selection. + +| Variable | Default | +|---|---| +| `YANGERD_SOCKET` | `/run/yangerd.sock` | +| `YANGERD_LOG_LEVEL` | `info` | +| `YANGERD_POLL_INTERVAL_SYSTEM` | `60s` | +| `YANGERD_POLL_INTERVAL_ROUTING` | `10s` | +| `YANGERD_POLL_INTERVAL_NTP` | `60s`, presence and port only | +| `YANGERD_POLL_INTERVAL_HARDWARE` | `10s`, GPS only | +| `YANGERD_POLL_INTERVAL_STP` | `5s` | +| `YANGERD_ENABLE_WIFI` | `false` | +| `YANGERD_ENABLE_LLDP` | `true` | +| `YANGERD_ENABLE_FIREWALL` | `true` | +| `YANGERD_ENABLE_DHCP` | `true` | +| `YANGERD_ENABLE_CONTAINERS` | `false` | +| `YANGERD_ENABLE_GPS` | `false` | + +A disabled feature's monitor is not started at all. + + +## Adding a data source + +1. **Find the event.** If the source can tell you when it changes, write + a monitor with a `Run(ctx) error` method and start it with `spawn()` in + `cmd/yangerd/main.go`. Put the reconnect loop in `backoff.Retry`, and + rebuild the full state after each reconnect, not just what later events + report. +2. **No event, value moves slowly:** implement `collector.Collector` and add + it to the polled list. Give the interval an environment variable. +3. **No event, value moves all the time:** register a tree provider. It + runs on every GET, so no forks and no network calls. +4. Emit JSON in the shape statd's libyang expects: module-prefixed top + node, RFC 7951 encoding (64-bit integers as strings, numeric union + members as numbers), no empty containers. +5. Delete the key when the feature is inactive. +6. Add the subscription in `src/statd/statd.c` if the module is new. +7. Unit test with `testutil.MockRunner` and `testutil.MockFileReader`, + or with a fake subprocess, see `internal/ipbatch`. + + +## Gotchas + +- **iproute2 caches interface names.** A long-lived `ip -batch` resolves + a name through a cache that never expires, so after an interface is + deleted and recreated, `dev wifi0` points at the dead index. Address, + neighbour and FDB queries therefore name devices `if`. + `link show` doesn't accept that form, so a failed link query for an + interface the kernel still has restarts the `ip` process and asks again. +- **`ip` can answer `[{}]`.** It opens the JSON object before checking + the netlink message, so a link racing a delete comes back empty. Only + stage a link answer with the queried ifindex and a name. +- **Link and address events use separate sockets.** They can be handled + out of order: a delete for an old interface may arrive after a new one + took its name. Staging is keyed by ifindex, and name-keyed state is only + dropped when no interface carries the name any more. +- **A dead batch means one re-dump.** While `ip` restarts, every event + fails. Failures ask for a re-dump, coalesced into one that waits until + the batches are back. +- **nl80211 delete events.** On `DEL_INTERFACE` the ifindex no longer + resolves, so take the name from the message. +- **FRR daemons differ.** bfdd installs its show commands in the enable + node only, so the vty client sends `enable` first, as vtysh does. +- **Operational lags the commit.** Data follows events, so a test that + reads back a value right after a commit has to poll with `until()`. + + +## Building and testing + +``` +go vet -mod=vendor ./... +go test -race -mod=vendor ./... +make yangerd-rebuild all # from the repo root, rebuilds the image +``` + +On a DUT, `YANGERD_LOG_LEVEL=debug` in `/etc/default/yangerd` followed by +`initctl restart yangerd` logs every batch command that fails and every +event that is dropped. diff --git a/src/yangerd/cmd/yangerctl/main.go b/src/yangerd/cmd/yangerctl/main.go new file mode 100644 index 000000000..0bc358b3c --- /dev/null +++ b/src/yangerd/cmd/yangerctl/main.go @@ -0,0 +1,107 @@ +package main + +import ( + "encoding/json" + "fmt" + "os" + "time" + + "github.com/kernelkit/infix/src/yangerd/internal/ipc" +) + +const defaultSocket = "/run/yangerd.sock" +const defaultTimeout = 5 * time.Second + +func main() { + socket := defaultSocket + timeout := defaultTimeout + + args := os.Args[1:] + for len(args) > 0 && len(args[0]) > 0 && args[0][0] == '-' { + switch args[0] { + case "--socket": + if len(args) < 2 { + die("--socket requires an argument") + } + socket = args[1] + args = args[2:] + case "--timeout": + if len(args) < 2 { + die("--timeout requires an argument") + } + d, err := time.ParseDuration(args[1]) + if err != nil { + die("invalid duration: %v", err) + } + timeout = d + args = args[2:] + default: + die("unknown flag: %s", args[0]) + } + } + + if len(args) == 0 { + usage() + } + + client := ipc.NewClient(socket, timeout) + + switch args[0] { + case "get": + if len(args) < 2 { + die("get requires a path argument") + } + resp, err := client.Get(args[1]) + if err != nil { + die("get: %v", err) + } + printResponse(resp) + case "dump": + resp, err := client.Get("/") + if err != nil { + die("dump: %v", err) + } + printResponse(resp) + case "health": + resp, err := client.Health() + if err != nil { + die("health: %v", err) + } + printResponse(resp) + default: + die("unknown command: %s", args[0]) + } +} + +func printResponse(resp *ipc.Response) { + if resp.Code == 503 { + fmt.Fprintf(os.Stderr, "yangerd is starting up\n") + os.Exit(3) + } + if resp.Status == "error" { + fmt.Fprintf(os.Stderr, "error %d: %s\n", resp.Code, resp.Message) + os.Exit(1) + } + + var out []byte + if resp.Data != nil { + out, _ = json.MarshalIndent(json.RawMessage(resp.Data), "", " ") + } else { + out, _ = json.MarshalIndent(resp, "", " ") + } + fmt.Println(string(out)) +} + +func usage() { + fmt.Fprintf(os.Stderr, "Usage: yangerctl [--socket path] [--timeout dur] [args]\n\n") + fmt.Fprintf(os.Stderr, "Commands:\n") + fmt.Fprintf(os.Stderr, " get Query a YANG subtree\n") + fmt.Fprintf(os.Stderr, " dump Dump entire tree\n") + fmt.Fprintf(os.Stderr, " health Show daemon health\n") + os.Exit(1) +} + +func die(format string, args ...interface{}) { + fmt.Fprintf(os.Stderr, "yangerctl: "+format+"\n", args...) + os.Exit(1) +} diff --git a/src/yangerd/cmd/yangerd/main.go b/src/yangerd/cmd/yangerd/main.go new file mode 100644 index 000000000..f2761c5d4 --- /dev/null +++ b/src/yangerd/cmd/yangerd/main.go @@ -0,0 +1,416 @@ +package main + +import ( + "context" + "encoding/json" + "log" + "log/slog" + "os" + "os/signal" + "strconv" + "strings" + "sync" + "sync/atomic" + "syscall" + "time" + + "github.com/kernelkit/infix/src/yangerd/internal/backoff" + "github.com/kernelkit/infix/src/yangerd/internal/collector" + "github.com/kernelkit/infix/src/yangerd/internal/config" + "github.com/kernelkit/infix/src/yangerd/internal/containermonitor" + "github.com/kernelkit/infix/src/yangerd/internal/dbusmonitor" + "github.com/kernelkit/infix/src/yangerd/internal/ethmonitor" + "github.com/kernelkit/infix/src/yangerd/internal/frrvty" + "github.com/kernelkit/infix/src/yangerd/internal/fswatcher" + "github.com/kernelkit/infix/src/yangerd/internal/ipbatch" + "github.com/kernelkit/infix/src/yangerd/internal/ipc" + "github.com/kernelkit/infix/src/yangerd/internal/iwmonitor" + "github.com/kernelkit/infix/src/yangerd/internal/kernelrib" + "github.com/kernelkit/infix/src/yangerd/internal/lldpmonitor" + "github.com/kernelkit/infix/src/yangerd/internal/ptpmonitor" + "github.com/kernelkit/infix/src/yangerd/internal/monitor" + "github.com/kernelkit/infix/src/yangerd/internal/sysreaders" + "github.com/kernelkit/infix/src/yangerd/internal/tftpmonitor" + "github.com/kernelkit/infix/src/yangerd/internal/tree" + "github.com/kernelkit/infix/src/yangerd/internal/wgquery" + "github.com/kernelkit/infix/src/yangerd/internal/zapiwatcher" +) + +// osFileChecker implements iface.FileChecker using the real filesystem. +type osFileChecker struct{} + +func (osFileChecker) Exists(path string) bool { + _, err := os.Stat(path) + return err == nil +} + +func (osFileChecker) ReadFile(path string) (string, error) { + b, err := os.ReadFile(path) + if err != nil { + return "", err + } + return string(b), nil +} + +func (osFileChecker) ListDir(path string) []string { + entries, err := os.ReadDir(path) + if err != nil { + return nil + } + names := make([]string, 0, len(entries)) + for _, e := range entries { + names = append(names, e.Name()) + } + return names +} + +func main() { + cfg := config.Load() + log.SetFlags(0) + + t := tree.New() + ready := &atomic.Bool{} + + srv := ipc.NewServer(t, ready) + if err := srv.Listen(cfg.Socket); err != nil { + log.Fatalf("listen %s: %v", cfg.Socket, err) + } + defer os.Remove(cfg.Socket) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + slogLog := slog.New(slog.NewTextHandler(os.Stderr, &slog.HandlerOptions{Level: slogLevel(cfg.LogLevel)})) + + var wg sync.WaitGroup + // spawn runs a monitor until ctx ends; an exit before that is a bug + // worth a log line, the monitor itself owns its restarts. + spawn := func(name string, run func(context.Context) error) { + wg.Add(1) + go func() { + defer wg.Done() + if err := run(ctx); err != nil && ctx.Err() == nil { + slogLog.Error(name+" exited", "err", err) + } + }() + } + cmd := collector.ExecRunner{} + fs := collector.OSFileReader{} + hardware := collector.NewHardwareCollector(cmd, fs, cfg.PollHardware, cfg.EnableWifi, cfg.EnableGPS) + ntp := collector.NewNTPCollector(cmd, cfg.PollNTP) + t.RegisterProvider("ietf-ntp:ntp", ntp.Live) + t.RegisterProvider("ietf-hardware:hardware", hardware.Live) + collectors := []collector.Collector{ + collector.NewSystemCollector(cmd, fs, cfg.PollSystem), + ntp, + hardware, + } + if cfg.EnableFRR { + collectors = append(collectors, collector.NewRoutingCollector(cfg.PollRouting)) + } + pokes := collector.RunAll(ctx, &wg, t, collectors) + + inst := collector.DBusInstaller{} + t.RegisterProvider("ietf-system:system-state", func() json.RawMessage { + live := tree.ShallowMerge(collector.LiveSystemState(fs), ntp.LiveSources()) + installerOverlay := collector.MergeInstaller(t.GetCached("ietf-system:system-state"), inst) + if installerOverlay == nil { + return live + } + return tree.ShallowMerge(live, installerOverlay) + }) + + if data := collector.BootPlatform(fs); data != nil { + t.Merge("ietf-system:system-state", data) + } + if data := collector.BootSoftware(ctx, cmd); data != nil { + t.Merge("ietf-system:system-state", data) + } + + linkBatch, err := ipbatch.New(ctx, slogLog, ipbatch.WithStats(), ipbatch.WithDetails()) + if err != nil { + log.Fatalf("start link batch: %v", err) + } + defer linkBatch.Close() + + addrBatch, err := ipbatch.New(ctx, slogLog, ipbatch.WithDetails()) + if err != nil { + log.Fatalf("start addr batch: %v", err) + } + defer addrBatch.Close() + + neighBatch, err := ipbatch.New(ctx, slogLog) + if err != nil { + log.Fatalf("start neigh batch: %v", err) + } + defer neighBatch.Close() + + brBatch, err := ipbatch.NewBridge(ctx, slogLog) + if err != nil { + log.Fatalf("start bridge batch: %v", err) + } + defer brBatch.Close() + + nlmon := monitor.New(linkBatch, addrBatch, neighBatch, brBatch, t, osFileChecker{}, slogLog) + t.RegisterProvider("ietf-interfaces:interfaces", nlmon.LastChange) + + // inotify says nothing when a /proc/sys/net/*/conf/ directory + // comes or goes, so follow the interface set from netlink instead. + linkSetCh := make(chan struct{}, 1) + nlmon.SetLinkSetChange(func() { + select { + case linkSetCh <- struct{}{}: + default: + } + }) + + ethMon := ethmonitor.New(slogLog, cmd) + ethMon.SetOnUpdate(nlmon.SetEthernetData) + nlmon.SetEthRefresh(ethMon.RefreshInterface) + spawn("ethmonitor", ethMon.Run) + + spawn("wireguard", poll(nlmon.WaitReady(), 10*time.Second, func() { + nlmon.SetWireguardAll(wgquery.Query(nlmon.Links())) + })) + spawn("stp", poll(nlmon.WaitReady(), cfg.PollSTP, nlmon.RefreshSTP)) + spawn("nlmonitor", func(ctx context.Context) error { + return backoff.Retry(ctx, slogLog, "nlmonitor", nlmon.Run) + }) + + if cfg.EnableWifi { + iwmon := iwmonitor.New(slogLog) + iwmon.SetOnUpdate(nlmon.SetWifiData) + iwmon.SetOnRadioChange(hardware.RequestRadioRefresh) + spawn("iwmonitor", iwmon.Run) + } + + spawn("radios", hardware.RunRadios) + + if cfg.EnableLLDP { + lldpmon := lldpmonitor.New(t, slogLog) + spawn("lldpmonitor", lldpmon.Run) + } + + if cfg.EnableContainers { + ctrmon := containermonitor.New(t, cmd, fs, slogLog) + spawn("containermonitor", ctrmon.Run) + } + + tftpmon := tftpmonitor.New(t, slogLog) + spawn("tftpmonitor", tftpmon.Run) + + ptpmon := ptpmonitor.New(t, slogLog) + spawn("ptpmonitor", ptpmon.Run) + + // The RIB comes from zebra when built with FRR, otherwise straight + // from the kernel table that netd installs into. + if cfg.EnableFRR { + zapi := zapiwatcher.New(t, frrvty.New(""), slogLog) + zapi.SetOnChange(func() { pokes.Poke("routing") }) + spawn("zapiwatcher", zapi.Run) + } else { + rib := kernelrib.New(t, slogLog) + spawn("kernelrib", rib.Run) + } + + if cfg.EnableDHCP || cfg.EnableFirewall { + dbusMon := dbusmonitor.New(t, slogLog) + spawn("dbusmonitor", dbusMon.Run) + } + + fsw, err := fswatcher.New(t, slogLog) + if err != nil { + log.Fatalf("start fswatcher: %v", err) + } + + fwdAgg := sysreaders.NewForwardingAggregator() + forwardingPaths := []string{ + "/proc/sys/net/ipv4/conf/*/forwarding", + "/proc/sys/net/ipv6/conf/*/forwarding", + } + forwarding := fswatcher.WatchHandler{ + TreeKey: routingTreeKey, + ReadFunc: fwdAgg.HandleForwardingChange, + Debounce: 100 * time.Millisecond, + UseMerge: true, + } + syncForwarding := func() { + for _, pattern := range forwardingPaths { + if _, _, err := fsw.SyncGlob(pattern, forwarding); err != nil { + slogLog.Warn("fswatcher glob failed", "pattern", pattern, "err", err) + } + } + } + syncForwarding() + + spawn("forwarding-sync", func(ctx context.Context) error { + for { + select { + case <-ctx.Done(): + return ctx.Err() + case <-linkSetCh: + syncForwarding() + } + } + }) + + if err := fsw.Watch("/etc/hostname", fswatcher.WatchHandler{ + TreeKey: "ietf-system:system", + ReadFunc: sysreaders.ReadHostname, + Debounce: 200 * time.Millisecond, + UseMerge: true, + }); err != nil { + slogLog.Warn("fswatcher watch failed", "path", "/etc/hostname", "err", err) + } + if err := fsw.WatchSymlink("/etc/localtime", fswatcher.WatchHandler{ + TreeKey: "ietf-system:system", + ReadFunc: sysreaders.ReadTimezone, + Debounce: 200 * time.Millisecond, + UseMerge: true, + }); err != nil { + slogLog.Warn("fswatcher watch failed", "path", "/etc/localtime", "err", err) + } + usersHandler := fswatcher.WatchHandler{ + TreeKey: "ietf-system:system", + ReadFunc: sysreaders.ReadUsers, + Debounce: 200 * time.Millisecond, + UseMerge: true, + } + if err := fsw.Watch("/etc/shadow", usersHandler); err != nil { + slogLog.Warn("fswatcher watch failed", "path", "/etc/shadow", "err", err) + } + if err := fsw.WatchDir(sysreaders.SSHDKeysDir, usersHandler); err != nil { + slogLog.Warn("fswatcher watch failed", "path", sysreaders.SSHDKeysDir, "err", err) + } + bootOrderHandler := fswatcher.WatchHandler{ + TreeKey: "ietf-system:system-state", + ReadFunc: makeBootOrderReader(t, cmd), + Debounce: 200 * time.Millisecond, + UseMerge: true, + } + // Watch the parent directory, not the file: fw_setenv (U-Boot) and + // grub-editenv may rewrite the env via a temp file + rename, which + // gives it a new inode that a direct file watch never sees. Watching + // the directory catches the Create/Rename (and still catches in-place + // writes), so a boot-order change after a RAUC install is reflected + // without waiting for a reboot. + for _, path := range []string{"/mnt/aux/grub/grubenv", "/mnt/aux/uboot.env"} { + if err := fsw.WatchSymlink(path, bootOrderHandler); err != nil { + slogLog.Debug("fswatcher boot-order watch skipped", "path", path, "err", err) + } + } + dnsHandler := fswatcher.WatchHandler{ + TreeKey: "ietf-system:system-state", + ReadFunc: sysreaders.ReadDNSResolver, + Debounce: 200 * time.Millisecond, + UseMerge: true, + } + for _, path := range []string{"/etc/resolv.conf.head", "/var/lib/misc/resolv.conf"} { + if err := fsw.WatchSymlink(path, dnsHandler); err != nil { + slogLog.Warn("fswatcher dns watch failed", "path", path, "err", err) + } + } + // Container operational data is handled by containermonitor (a + // `podman events` stream), not the fswatcher. + fsw.InitialRead() + spawn("fswatcher", fsw.Run) + + go func() { + <-nlmon.WaitReady() + ready.Store(true) + // finit's notify:pid marks the service ready when this appears + if err := os.WriteFile(pidFile, []byte(strconv.Itoa(os.Getpid())+"\n"), 0644); err != nil { + slogLog.Warn("write pidfile", "path", pidFile, "err", err) + } + }() + defer os.Remove(pidFile) + + sigCh := make(chan os.Signal, 1) + signal.Notify(sigCh, syscall.SIGTERM, syscall.SIGINT, syscall.SIGHUP) + + go func() { + for sig := range sigCh { + if sig == syscall.SIGHUP { + log.Printf("SIGHUP: triggering immediate re-poll") + pokes.PokeAll() + continue + } + cancel() + return + } + }() + + if err := srv.Serve(ctx); err != nil { + log.Fatalf("serve: %v", err) + } + + wg.Wait() +} + +// poll runs fn once ready is closed and then every interval until the +// context ends. +func poll(ready <-chan struct{}, every time.Duration, fn func()) func(context.Context) error { + return func(ctx context.Context) error { + select { + case <-ctx.Done(): + return ctx.Err() + case <-ready: + } + ticker := time.NewTicker(every) + defer ticker.Stop() + for { + fn() + select { + case <-ctx.Done(): + return ctx.Err() + case <-ticker.C: + } + } + } +} + +func slogLevel(s string) slog.Level { + switch strings.ToLower(s) { + case "debug": + return slog.LevelDebug + case "warn", "warning": + return slog.LevelWarn + case "error": + return slog.LevelError + default: + return slog.LevelInfo + } +} + +const ( + routingTreeKey = "ietf-routing:routing" + pidFile = "/run/yangerd.pid" +) + +func makeBootOrderReader(t *tree.Tree, cmd collector.CommandRunner) func(string) (json.RawMessage, error) { + return func(_ string) (json.RawMessage, error) { + bootOrder := collector.ReadBootOrder(context.TODO(), cmd) + + raw := t.GetCached("ietf-system:system-state") + var state map[string]interface{} + if raw != nil { + json.Unmarshal(raw, &state) + } + if state == nil { + state = make(map[string]interface{}) + } + + sw, _ := state["infix-system:software"].(map[string]interface{}) + if sw == nil { + sw = make(map[string]interface{}) + } + + if bootOrder != nil { + sw["boot-order"] = bootOrder + } else { + delete(sw, "boot-order") + } + + return json.Marshal(map[string]interface{}{"infix-system:software": sw}) + } +} diff --git a/src/yangerd/go.mod b/src/yangerd/go.mod new file mode 100644 index 000000000..0dbb2e2ed --- /dev/null +++ b/src/yangerd/go.mod @@ -0,0 +1,20 @@ +module github.com/kernelkit/infix/src/yangerd + +go 1.23.0 + +require ( + github.com/facebook/time v0.0.0-20250531133328-3ef67721da27 + github.com/godbus/dbus/v5 v5.2.2 + github.com/mdlayher/genetlink v1.3.2 + github.com/mdlayher/netlink v1.8.0 + github.com/vishvananda/netlink v1.3.1 + golang.org/x/sys v0.35.0 +) + +require ( + github.com/google/go-cmp v0.7.0 // indirect + github.com/mdlayher/socket v0.5.1 // indirect + github.com/vishvananda/netns v0.0.5 // indirect + golang.org/x/net v0.43.0 // indirect + golang.org/x/sync v0.10.0 // indirect +) diff --git a/src/yangerd/go.sum b/src/yangerd/go.sum new file mode 100644 index 000000000..dd45514bd --- /dev/null +++ b/src/yangerd/go.sum @@ -0,0 +1,32 @@ +github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= +github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/facebook/time v0.0.0-20250531133328-3ef67721da27 h1:bnAJYUd13zV2nsvYrXxhx6faPhJqd0wdbyTqjJrF+18= +github.com/facebook/time v0.0.0-20250531133328-3ef67721da27/go.mod h1:bp0KsBqhjbum8LF5Canem6aGBHWk5EGZRRN7z3rMp0E= +github.com/godbus/dbus/v5 v5.2.2 h1:TUR3TgtSVDmjiXOgAAyaZbYmIeP3DPkld3jgKGV8mXQ= +github.com/godbus/dbus/v5 v5.2.2/go.mod h1:3AAv2+hPq5rdnr5txxxRwiGjPXamgoIHgz9FPBfOp3c= +github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= +github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= +github.com/mdlayher/genetlink v1.3.2 h1:KdrNKe+CTu+IbZnm/GVUMXSqBBLqcGpRDa0xkQy56gw= +github.com/mdlayher/genetlink v1.3.2/go.mod h1:tcC3pkCrPUGIKKsCsp0B3AdaaKuHtaxoJRz3cc+528o= +github.com/mdlayher/netlink v1.8.0 h1:e7XNIYJKD7hUct3Px04RuIGJbBxy1/c4nX7D5YyvvlM= +github.com/mdlayher/netlink v1.8.0/go.mod h1:UhgKXUlDQhzb09DrCl2GuRNEglHmhYoWAHid9HK3594= +github.com/mdlayher/socket v0.5.1 h1:VZaqt6RkGkt2OE9l3GcC6nZkqD3xKeQLyfleW/uBcos= +github.com/mdlayher/socket v0.5.1/go.mod h1:TjPLHI1UgwEv5J1B5q0zTZq12A/6H7nKmtTanQE37IQ= +github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= +github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= +github.com/stretchr/testify v1.8.4 h1:CcVxjf3Q8PM0mHUKJCdn+eZZtm5yQwehR5yeSVQQcUk= +github.com/stretchr/testify v1.8.4/go.mod h1:sz/lmYIOXD/1dqDmKjjqLyZ2RngseejIcXlSw2iwfAo= +github.com/vishvananda/netlink v1.3.1 h1:3AEMt62VKqz90r0tmNhog0r/PpWKmrEShJU0wJW6bV0= +github.com/vishvananda/netlink v1.3.1/go.mod h1:ARtKouGSTGchR8aMwmkzC0qiNPrrWO5JS/XMVl45+b4= +github.com/vishvananda/netns v0.0.5 h1:DfiHV+j8bA32MFM7bfEunvT8IAqQ/NzSJHtcmW5zdEY= +github.com/vishvananda/netns v0.0.5/go.mod h1:SpkAiCQRtJ6TvvxPnOSyH3BMl6unz3xZlaprSwhNNJM= +golang.org/x/net v0.43.0 h1:lat02VYK2j4aLzMzecihNvTlJNQUq316m2Mr9rnM6YE= +golang.org/x/net v0.43.0/go.mod h1:vhO1fvI4dGsIjh73sWfUVjj3N7CA9WkKJNQm2svM6Jg= +golang.org/x/sync v0.10.0 h1:3NQrjDixjgGwUOCaF8w2+VYHv0Ve/vGYSbdkTa98gmQ= +golang.org/x/sync v0.10.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk= +golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/sys v0.35.0 h1:vz1N37gP5bs89s7He8XuIYXpyY0+QlsKmzipCbUtyxI= +golang.org/x/sys v0.35.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k= +gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= +gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/src/yangerd/internal/backoff/backoff.go b/src/yangerd/internal/backoff/backoff.go new file mode 100644 index 000000000..c975a0c15 --- /dev/null +++ b/src/yangerd/internal/backoff/backoff.go @@ -0,0 +1,80 @@ +// Package backoff provides exponential backoff retry logic with +// context-aware sleep, shared across reactive monitors. +package backoff + +import ( + "context" + "log/slog" + "math" + "time" +) + +// Backoff implements exponential backoff with a configurable initial +// delay, maximum delay, and growth factor. +type Backoff struct { + Initial time.Duration + Max time.Duration + Factor float64 +} + +// Default returns a Backoff with the standard yangerd parameters: +// 100ms initial, 30s max, factor 2. +func Default() *Backoff { + return &Backoff{ + Initial: 100 * time.Millisecond, + Max: 30 * time.Second, + Factor: 2.0, + } +} + +// Next returns the next delay value after current. If current is +// zero, Initial is returned. +func (b *Backoff) Next(current time.Duration) time.Duration { + if current <= 0 { + return b.Initial + } + next := time.Duration(math.Min(float64(current)*b.Factor, float64(b.Max))) + if next <= 0 { + return b.Initial + } + return next +} + +// Sleep waits for duration d or until ctx is cancelled, whichever +// comes first. Returns ctx.Err() if the context was cancelled. +func Sleep(ctx context.Context, d time.Duration) error { + select { + case <-ctx.Done(): + return ctx.Err() + case <-time.After(d): + return nil + } +} + +// Retry runs fn until ctx is cancelled, restarting it with exponential +// backoff each time it returns. A run that lasted longer than the +// maximum delay counts as healthy and resets the delay, so a source +// that flaps after hours of service is retried quickly, while one that +// dies at once backs off to Max. This is the restart loop every +// reactive monitor needs; name labels the log line. +func Retry(ctx context.Context, log *slog.Logger, name string, fn func(ctx context.Context) error) error { + bo := Default() + delay := bo.Initial + + for { + started := time.Now() + err := fn(ctx) + if ctx.Err() != nil { + return ctx.Err() + } + if time.Since(started) > bo.Max { + delay = bo.Initial + } + + log.Warn(name+": exited, restarting", "err", err, "delay", delay) + if err := Sleep(ctx, delay); err != nil { + return err + } + delay = bo.Next(delay) + } +} diff --git a/src/yangerd/internal/backoff/backoff_test.go b/src/yangerd/internal/backoff/backoff_test.go new file mode 100644 index 000000000..54bc4a40d --- /dev/null +++ b/src/yangerd/internal/backoff/backoff_test.go @@ -0,0 +1,43 @@ +package backoff + +import ( + "context" + "errors" + "log/slog" + "testing" + "time" +) + +func TestNextGrowsToMax(t *testing.T) { + b := &Backoff{Initial: 100 * time.Millisecond, Max: 300 * time.Millisecond, Factor: 2} + got := []time.Duration{b.Next(0)} + for i := 0; i < 3; i++ { + got = append(got, b.Next(got[len(got)-1])) + } + want := []time.Duration{100 * time.Millisecond, 200 * time.Millisecond, 300 * time.Millisecond, 300 * time.Millisecond} + for i := range want { + if got[i] != want[i] { + t.Fatalf("step %d = %v, want %v", i, got[i], want[i]) + } + } +} + +// Retry keeps calling fn until the context ends, and returns the +// context error rather than fn's. +func TestRetryRestartsUntilCancelled(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + calls := 0 + err := Retry(ctx, slog.Default(), "test", func(ctx context.Context) error { + calls++ + if calls == 3 { + cancel() + } + return errors.New("boom") + }) + if !errors.Is(err, context.Canceled) { + t.Fatalf("err = %v, want context.Canceled", err) + } + if calls != 3 { + t.Fatalf("calls = %d, want 3", calls) + } +} diff --git a/src/yangerd/internal/collector/boot.go b/src/yangerd/internal/collector/boot.go new file mode 100644 index 000000000..0d18be6c6 --- /dev/null +++ b/src/yangerd/internal/collector/boot.go @@ -0,0 +1,151 @@ +package collector + +import ( + "context" + "encoding/json" + "log" + "strconv" + "strings" +) + +func BootPlatform(fs FileReader) json.RawMessage { + data, err := fs.ReadFile("/etc/os-release") + if err != nil { + log.Printf("boot: os-release: %v", err) + return nil + } + platform := make(map[string]interface{}) + for _, line := range strings.Split(string(data), "\n") { + idx := strings.IndexByte(line, '=') + if idx < 0 { + continue + } + key := line[:idx] + val := strings.Trim(line[idx+1:], "\"") + if mapped, ok := platformKeyMap[key]; ok { + platform[mapped] = val + } + } + result, _ := json.Marshal(map[string]interface{}{"platform": platform}) + return result +} + +func BootSoftware(ctx context.Context, cmd CommandRunner) json.RawMessage { + software := make(map[string]interface{}) + + var raucData map[string]interface{} + if err := runJSON(ctx, cmd, &raucData, "rauc", "status", "--detailed", "--output-format=json"); err != nil { + log.Printf("boot: %v", err) + } else { + for _, key := range []string{"compatible", "variant", "booted"} { + if v, ok := raucData[key]; ok { + software[key] = v + } + } + if slots := softwareSlots(raucData); slots != nil { + software["slot"] = slots + } + } + + bootOrder := ReadBootOrder(ctx, cmd) + if bootOrder != nil { + software["boot-order"] = bootOrder + } + + result, _ := json.Marshal(map[string]interface{}{"infix-system:software": software}) + return result +} + +func ReadBootOrder(ctx context.Context, cmd CommandRunner) []string { + out, err := cmd.Run(ctx, "fw_printenv", "BOOT_ORDER") + if err == nil { + for _, line := range strings.Split(string(out), "\n") { + if strings.Contains(line, "BOOT_ORDER") { + parts := strings.SplitN(line, "=", 2) + if len(parts) == 2 { + return strings.Fields(parts[1]) + } + } + } + } + + out, err = cmd.Run(ctx, "grub-editenv", "/mnt/aux/grub/grubenv", "list") + if err == nil { + for _, line := range strings.Split(string(out), "\n") { + if strings.Contains(line, "ORDER") { + parts := strings.SplitN(line, "=", 2) + if len(parts) == 2 { + return strings.Fields(strings.TrimSpace(parts[1])) + } + } + } + } + + return nil +} + +// softwareSlots lists the RAUC slots, nil when rauc reported none. +func softwareSlots(raucData map[string]interface{}) []interface{} { + slotsArr, ok := raucData["slots"].([]interface{}) + if !ok { + return nil + } + + slots := []interface{}{} + for _, slotItem := range slotsArr { + slotMap, ok := slotItem.(map[string]interface{}) + if !ok { + continue + } + for name, valRaw := range slotMap { + if val, ok := valRaw.(map[string]interface{}); ok { + slots = append(slots, softwareSlot(name, val)) + } + } + } + return slots +} + +func softwareSlot(name string, val map[string]interface{}) map[string]interface{} { + s := map[string]interface{}{ + "name": name, + "bootname": val["bootname"], + "class": val["class"], + "state": val["state"], + } + + slotStatus, _ := val["slot_status"].(map[string]interface{}) + if slotStatus == nil { + return s + } + + bundle := make(map[string]interface{}) + if b, ok := slotStatus["bundle"].(map[string]interface{}); ok { + setIfPresent(bundle, "compatible", b, "compatible") + setIfPresent(bundle, "version", b, "version") + } + s["bundle"] = bundle + + if ck, ok := slotStatus["checksum"].(map[string]interface{}); ok { + if v := ck["size"]; v != nil { + s["size"] = strconv.FormatInt(int64(toInt(v)), 10) + } + setIfPresent(s, "sha256", ck, "sha256") + } + + s["installed"] = slotEvent(slotStatus["installed"]) + s["activated"] = slotEvent(slotStatus["activated"]) + + return s +} + +// slotEvent maps a RAUC {timestamp, count} record, as for the last +// install or activation of a slot. +func slotEvent(raw interface{}) map[string]interface{} { + event := make(map[string]interface{}) + if rec, ok := raw.(map[string]interface{}); ok { + setIfPresent(event, "datetime", rec, "timestamp") + setIfPresentInt(event, "count", rec, "count") + } + return event +} diff --git a/src/yangerd/internal/collector/boot_test.go b/src/yangerd/internal/collector/boot_test.go new file mode 100644 index 000000000..9b53c0bf2 --- /dev/null +++ b/src/yangerd/internal/collector/boot_test.go @@ -0,0 +1,236 @@ +package collector + +import ( + "context" + "encoding/json" + "fmt" + "testing" + + "github.com/kernelkit/infix/src/yangerd/internal/testutil" +) + +const ( + testOSRelease = `NAME="Infix" +VERSION_ID="25.01.0" +BUILD_ID="v25.01.0" +ARCHITECTURE="x86_64" +HOME_URL="https://kernelkit.github.io" +` + + testRaucStatus = `{ + "compatible": "Infix x86_64", + "variant": "", + "booted": "rootfs.0", + "slots": [ + { + "rootfs.0": { + "bootname": "A", + "class": "rootfs", + "state": "booted", + "slot_status": { + "bundle": { + "compatible": "Infix x86_64", + "version": "25.01.0" + }, + "checksum": { + "sha256": "abc123", + "size": 134217728 + }, + "installed": { + "timestamp": "2025-01-15T10:30:00Z", + "count": 3 + }, + "activated": { + "timestamp": "2025-01-15T10:31:00Z", + "count": 3 + } + } + } + }, + { + "rootfs.1": { + "bootname": "B", + "class": "rootfs", + "state": "inactive", + "slot_status": { + "bundle": { + "compatible": "Infix x86_64", + "version": "24.10.0" + }, + "checksum": { + "sha256": "def456", + "size": 130000000 + }, + "installed": { + "timestamp": "2024-10-01T08:00:00Z", + "count": 1 + }, + "activated": { + "timestamp": "2024-10-01T08:01:00Z", + "count": 1 + } + } + } + } + ] +}` + + testBootOrder = "BOOT_ORDER=A B\n" +) + +func TestBootPlatform(t *testing.T) { + fs := &testutil.MockFileReader{ + Files: map[string][]byte{ + "/etc/os-release": []byte(testOSRelease), + }, + } + + raw := BootPlatform(fs) + if raw == nil { + t.Fatal("BootPlatform returned nil") + } + + var result map[string]interface{} + if err := json.Unmarshal(raw, &result); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + platform, ok := result["platform"].(map[string]interface{}) + if !ok { + t.Fatal("missing platform key") + } + + checks := map[string]string{ + "os-name": "Infix", + "os-version": "25.01.0", + "os-release": "v25.01.0", + "machine": "x86_64", + } + for key, expected := range checks { + got, ok := platform[key].(string) + if !ok || got != expected { + t.Fatalf("platform[%q]: expected %q, got %v", key, expected, platform[key]) + } + } +} + +func TestBootPlatformMissingFile(t *testing.T) { + fs := &testutil.MockFileReader{Files: map[string][]byte{}} + raw := BootPlatform(fs) + if raw != nil { + t.Fatalf("expected nil for missing os-release, got %s", raw) + } +} + +func TestBootSoftware(t *testing.T) { + runner := &testutil.MockRunner{ + Results: map[string][]byte{ + "rauc status --detailed --output-format=json": []byte(testRaucStatus), + "fw_printenv BOOT_ORDER": []byte(testBootOrder), + }, + Errors: map[string]error{}, + } + + raw := BootSoftware(context.Background(), runner) + if raw == nil { + t.Fatal("BootSoftware returned nil") + } + + var result map[string]interface{} + if err := json.Unmarshal(raw, &result); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + sw, ok := result["infix-system:software"].(map[string]interface{}) + if !ok { + t.Fatal("missing infix-system:software key") + } + + if sw["compatible"] != "Infix x86_64" { + t.Fatalf("compatible: expected 'Infix x86_64', got %v", sw["compatible"]) + } + if sw["booted"] != "rootfs.0" { + t.Fatalf("booted: expected 'rootfs.0', got %v", sw["booted"]) + } + + bootOrder, ok := sw["boot-order"].([]interface{}) + if !ok || len(bootOrder) != 2 { + t.Fatalf("expected boot-order [A B], got %v", sw["boot-order"]) + } + if bootOrder[0] != "A" || bootOrder[1] != "B" { + t.Fatalf("boot-order: expected [A B], got %v", bootOrder) + } + + slots, ok := sw["slot"].([]interface{}) + if !ok || len(slots) != 2 { + t.Fatalf("expected 2 slots, got %v", sw["slot"]) + } +} + +func TestBootSoftwareAllCommandsFail(t *testing.T) { + runner := &testutil.MockRunner{ + Results: map[string][]byte{}, + Errors: map[string]error{ + "rauc status --detailed --output-format=json": fmt.Errorf("not found"), + "fw_printenv BOOT_ORDER": fmt.Errorf("not found"), + "grub-editenv /mnt/aux/grub/grubenv list": fmt.Errorf("not found"), + }, + } + + raw := BootSoftware(context.Background(), runner) + if raw == nil { + t.Fatal("BootSoftware should return non-nil even when all commands fail") + } + + var result map[string]interface{} + json.Unmarshal(raw, &result) + sw := result["infix-system:software"].(map[string]interface{}) + if _, ok := sw["boot-order"]; ok { + t.Fatal("boot-order should not be present when commands fail") + } +} + +func TestReadBootOrderFwPrintenv(t *testing.T) { + runner := &testutil.MockRunner{ + Results: map[string][]byte{ + "fw_printenv BOOT_ORDER": []byte("BOOT_ORDER=A B\n"), + }, + Errors: map[string]error{}, + } + + order := ReadBootOrder(context.Background(), runner) + if len(order) != 2 || order[0] != "A" || order[1] != "B" { + t.Fatalf("expected [A B], got %v", order) + } +} + +func TestReadBootOrderGrubFallback(t *testing.T) { + runner := &testutil.MockRunner{ + Results: map[string][]byte{ + "grub-editenv /mnt/aux/grub/grubenv list": []byte("ORDER=B A\n"), + }, + Errors: map[string]error{ + "fw_printenv BOOT_ORDER": fmt.Errorf("command not found"), + }, + } + + order := ReadBootOrder(context.Background(), runner) + if len(order) != 2 || order[0] != "B" || order[1] != "A" { + t.Fatalf("expected [B A], got %v", order) + } +} + +func TestReadBootOrderBothFail(t *testing.T) { + runner := &testutil.MockRunner{ + Results: map[string][]byte{}, + Errors: map[string]error{ + "fw_printenv BOOT_ORDER": fmt.Errorf("not found"), + "grub-editenv /mnt/aux/grub/grubenv list": fmt.Errorf("not found"), + }, + } + + order := ReadBootOrder(context.Background(), runner) + if order != nil { + t.Fatalf("expected nil, got %v", order) + } +} diff --git a/src/yangerd/internal/collector/collector.go b/src/yangerd/internal/collector/collector.go new file mode 100644 index 000000000..6651c68d6 --- /dev/null +++ b/src/yangerd/internal/collector/collector.go @@ -0,0 +1,87 @@ +// Package collector defines the Collector interface and the RunAll +// scheduler that drives periodic data collection into the Tree. +package collector + +import ( + "context" + "log" + "sync" + "time" + + "github.com/kernelkit/infix/src/yangerd/internal/tree" +) + +// Collector gathers operational data and writes it to the Tree. +type Collector interface { + Name() string + Interval() time.Duration + Collect(ctx context.Context, t *tree.Tree) error +} + +// Pokes asks running collectors for an immediate collection. Each +// collector has its own channel, so a poke reaches the one it is meant +// for, and pokes that arrive while it is busy collapse into one. +type Pokes struct { + chans map[string]chan struct{} +} + +// Poke asks the named collector to collect now. +func (p *Pokes) Poke(name string) { + if ch, ok := p.chans[name]; ok { + poke(ch) + } +} + +// PokeAll asks every collector to collect now. +func (p *Pokes) PokeAll() { + for _, ch := range p.chans { + poke(ch) + } +} + +func poke(ch chan struct{}) { + select { + case ch <- struct{}{}: + default: + } +} + +// RunAll starts one goroutine per Collector, each ticking at the +// collector's configured interval. A failed Collect is logged and +// retried on the next tick. All goroutines exit when ctx is cancelled. +func RunAll(ctx context.Context, wg *sync.WaitGroup, t *tree.Tree, collectors []Collector) *Pokes { + p := &Pokes{chans: make(map[string]chan struct{}, len(collectors))} + for _, c := range collectors { + ch := make(chan struct{}, 1) + p.chans[c.Name()] = ch + wg.Add(1) + go runOne(ctx, wg, t, c, ch) + } + return p +} + +func runOne(ctx context.Context, wg *sync.WaitGroup, t *tree.Tree, c Collector, pokeCh <-chan struct{}) { + defer wg.Done() + + if err := c.Collect(ctx, t); err != nil { + log.Printf("collector %s: initial: %v", c.Name(), err) + } + + ticker := time.NewTicker(c.Interval()) + defer ticker.Stop() + + for { + select { + case <-ctx.Done(): + return + case <-ticker.C: + if err := c.Collect(ctx, t); err != nil { + log.Printf("collector %s: %v", c.Name(), err) + } + case <-pokeCh: + if err := c.Collect(ctx, t); err != nil { + log.Printf("collector %s: poke: %v", c.Name(), err) + } + } + } +} diff --git a/src/yangerd/internal/collector/containers.go b/src/yangerd/internal/collector/containers.go new file mode 100644 index 000000000..40db70079 --- /dev/null +++ b/src/yangerd/internal/collector/containers.go @@ -0,0 +1,448 @@ +package collector + +import ( + "context" + "encoding/json" + "fmt" + "log" + "path/filepath" + "regexp" + "strconv" + "strings" + "time" +) + +var sizeRe = regexp.MustCompile(`(?i)^\s*([0-9.]+)\s*([KMGT]?I?B)?\s*$`) + +// collectTimeout bounds one full collection: a podman stats stuck on a +// container in a bad state must not wedge the container monitor. +const collectTimeout = 60 * time.Second + +// containerCollector gathers infix-containers operational data. +type containerCollector struct { + cmd CommandRunner + fs FileReader +} + +// CollectContainers runs a full container collection and returns the +// result as JSON suitable for tree.Set("infix-containers:containers"), +// or nil when no container exists, so an enabled but idle container +// feature does not surface as operational data. +func CollectContainers(cmd CommandRunner, fs FileReader) json.RawMessage { + ctx, cancel := context.WithTimeout(context.Background(), collectTimeout) + defer cancel() + + c := &containerCollector{cmd: cmd, fs: fs} + containers := []interface{}{} + for _, ps := range c.podmanPS(ctx) { + if cont := c.container(ctx, ps); cont != nil { + containers = append(containers, cont) + } + } + if len(containers) == 0 { + return nil + } + + data, err := json.Marshal(map[string]interface{}{"container": containers}) + if err != nil { + return nil + } + return data +} + +func (c *containerCollector) podmanPS(ctx context.Context) []map[string]interface{} { + var list []map[string]interface{} + if err := runJSON(ctx, c.cmd, &list, "podman", "ps", "-a", "--format=json"); err != nil { + log.Printf("collector containers: %v", err) + return nil + } + return list +} + +func (c *containerCollector) podmanInspect(ctx context.Context, name string) map[string]interface{} { + out, err := c.cmd.Run(ctx, "podman", "inspect", name) + if err != nil { + log.Printf("collector containers: inspect %s: %v", name, err) + return map[string]interface{}{} + } + + var list []map[string]interface{} + if err := json.Unmarshal(out, &list); err == nil && len(list) > 0 { + return list[0] + } + + var generic []interface{} + if err := json.Unmarshal(out, &generic); err == nil { + for _, item := range generic { + if m, ok := item.(map[string]interface{}); ok { + return m + } + } + } + + var single map[string]interface{} + if err := json.Unmarshal(out, &single); err == nil { + return single + } + + log.Printf("collector containers: inspect %s parse: invalid json", name) + return map[string]interface{}{} +} + +func (c *containerCollector) resourceStats(ctx context.Context, name string) map[string]interface{} { + out, err := c.cmd.Run(ctx, "podman", "stats", "--no-stream", "--format", "json", "--no-reset", name) + if err != nil { + log.Printf("collector containers: stats %s: %v", name, err) + return nil + } + + var statsList []map[string]interface{} + if err := json.Unmarshal(out, &statsList); err != nil { + var single map[string]interface{} + if err2 := json.Unmarshal(out, &single); err2 != nil { + log.Printf("collector containers: stats %s parse: %v", name, err) + return nil + } + statsList = append(statsList, single) + } + + if len(statsList) == 0 { + return nil + } + + stat := statsList[0] + rusage := make(map[string]interface{}) + + if memUsage, ok := stat["mem_usage"].(string); ok { + parts := strings.SplitN(memUsage, "/", 2) + if len(parts) == 2 { + memKiB := parseSizeKiB(strings.TrimSpace(parts[0])) + rusage["memory"] = strconv.Itoa(memKiB) + } + } + + if cpuPercent, ok := stat["cpu_percent"].(string); ok { + cpuPercent = strings.TrimSpace(strings.TrimSuffix(cpuPercent, "%")) + if cpuVal, err := strconv.ParseFloat(cpuPercent, 64); err == nil { + rusage["cpu"] = fmt.Sprintf("%.2f", cpuVal) + } + } + + if blockIO, ok := stat["block_io"].(string); ok { + parts := strings.SplitN(blockIO, "/", 2) + if len(parts) == 2 { + readKiB := parseSizeKiB(strings.TrimSpace(parts[0])) + writeKiB := parseSizeKiB(strings.TrimSpace(parts[1])) + + bio := make(map[string]interface{}) + if readKiB > 0 { + bio["read"] = strconv.Itoa(readKiB) + } + if writeKiB > 0 { + bio["write"] = strconv.Itoa(writeKiB) + } + rusage["block-io"] = bio + } + } + + if netIO, ok := stat["net_io"].(string); ok { + parts := strings.SplitN(netIO, "/", 2) + if len(parts) == 2 { + rxKiB := parseSizeKiB(strings.TrimSpace(parts[0])) + txKiB := parseSizeKiB(strings.TrimSpace(parts[1])) + + nio := make(map[string]interface{}) + if rxKiB > 0 { + nio["received"] = strconv.Itoa(rxKiB) + } + if txKiB > 0 { + nio["sent"] = strconv.Itoa(txKiB) + } + rusage["net-io"] = nio + } + } + + if pids, ok := stat["pids"]; ok { + pidInt := toInt(pids) + rusage["pids"] = pidInt + } + + if len(rusage) == 0 { + return nil + } + + return rusage +} + +func (c *containerCollector) readCgroupLimits(inspect map[string]interface{}) map[string]interface{} { + stateRaw, ok := inspect["State"] + if !ok { + return nil + } + state, ok := stateRaw.(map[string]interface{}) + if !ok { + return nil + } + + cgroupPath, ok := state["CgroupPath"].(string) + if !ok || cgroupPath == "" { + return nil + } + + cgroupBase := "/sys/fs/cgroup" + cgroupPath + memVal := 0 + cpuVal := 0 + + if data, err := c.fs.ReadFile(filepath.Join(cgroupBase, "memory.max")); err == nil { + memVal = parseCgroupMemory(strings.TrimSpace(string(data))) + } + + if data, err := c.fs.ReadFile(filepath.Join(cgroupBase, "cpu.max")); err == nil { + cpuVal = parseCgroupCPU(strings.TrimSpace(string(data))) + } + + if memVal <= 0 && cpuVal <= 0 { + return nil + } + + result := make(map[string]interface{}) + if memVal > 0 { + result["memory"] = strconv.Itoa(memVal) + } + if cpuVal > 0 { + result["cpu"] = cpuVal + } + + return result +} + +func (c *containerCollector) network(ps map[string]interface{}, inspect map[string]interface{}) map[string]interface{} { + networkSettingsRaw, hasNetworkSettings := inspect["NetworkSettings"] + if hasNetworkSettings { + if networkSettings, ok := networkSettingsRaw.(map[string]interface{}); ok { + if networksRaw, ok := networkSettings["Networks"]; ok { + if networks, ok := networksRaw.(map[string]interface{}); ok { + if _, ok := networks["host"]; ok { + return map[string]interface{}{"host": true} + } + } + } + } + } + + net := map[string]interface{}{ + "interface": []interface{}{}, + "publish": []interface{}{}, + } + + networks := asStringSlice(ps["Networks"]) + ifaces := net["interface"].([]interface{}) + for _, n := range networks { + ifaces = append(ifaces, map[string]interface{}{"name": n}) + } + net["interface"] = ifaces + + running := strings.EqualFold(asString(ps["State"]), "running") + if !running { + return net + } + + portsRaw, ok := ps["Ports"] + if !ok { + return net + } + + ports, ok := portsRaw.([]interface{}) + if !ok || len(ports) == 0 { + return net + } + + publish := net["publish"].([]interface{}) + for _, portRaw := range ports { + port, ok := portRaw.(map[string]interface{}) + if !ok { + continue + } + + hostIP := asString(port["host_ip"]) + hostPort := asString(port["host_port"]) + if hostPort == "" { + hostPort = strconv.Itoa(toInt(port["host_port"])) + } + containerPort := asString(port["container_port"]) + if containerPort == "" { + containerPort = strconv.Itoa(toInt(port["container_port"])) + } + protocol := asString(port["protocol"]) + + if hostPort == "0" || hostPort == "" || containerPort == "0" || containerPort == "" || protocol == "" { + continue + } + + addr := "" + if hostIP != "" { + addr = hostIP + ":" + } + + publish = append(publish, fmt.Sprintf("%s%s:%s/%s", addr, hostPort, containerPort, protocol)) + } + net["publish"] = publish + + return net +} + +func (c *containerCollector) container(ctx context.Context, ps map[string]interface{}) map[string]interface{} { + names := asStringSlice(ps["Names"]) + if len(names) == 0 { + return nil + } + + name := names[0] + running := strings.EqualFold(asString(ps["State"]), "running") + + out := map[string]interface{}{ + "name": name, + "id": asString(ps["Id"]), + "image": asString(ps["Image"]), + "image-id": asString(ps["ImageID"]), + "running": running, + "status": asString(ps["Status"]), + } + + inspect := c.podmanInspect(ctx, name) + + // Report the actual running command line as the config-false + // "cmdline" leaf (an unrestricted string), built from inspect's + // Path + Args like the legacy yanger collector. Do NOT report it + // into the config-true "command" leaf: that leaf has a restrictive + // pattern and a real command line (e.g. one containing "&&" or + // quotes) fails YANG validation, which rejects the entire + // containers subtree on read. + if path := asString(inspect["Path"]); path != "" { + parts := append([]string{path}, asStringSlice(inspect["Args"])...) + out["cmdline"] = strings.Join(parts, " ") + } + + if net := c.network(ps, inspect); len(net) > 0 { + out["network"] = net + } + + if limits := c.readCgroupLimits(inspect); limits != nil { + out["resource-limit"] = limits + } + + if running { + if usage := c.resourceStats(ctx, name); usage != nil { + out["resource-usage"] = usage + } + } + + return out +} + +func asString(v interface{}) string { + s, ok := v.(string) + if ok { + return s + } + return "" +} + +func asStringSlice(v interface{}) []string { + switch vv := v.(type) { + case []string: + return vv + case []interface{}: + out := make([]string, 0, len(vv)) + for _, e := range vv { + if s, ok := e.(string); ok && s != "" { + out = append(out, s) + } + } + return out + case string: + if vv == "" { + return nil + } + return splitLines(vv) + default: + return nil + } +} + +func parseSizeKiB(sizeStr string) int { + if strings.TrimSpace(sizeStr) == "" { + return 0 + } + + m := sizeRe.FindStringSubmatch(strings.ToUpper(strings.TrimSpace(sizeStr))) + if len(m) < 2 { + return 0 + } + + value, err := strconv.ParseFloat(m[1], 64) + if err != nil { + return 0 + } + + unit := "B" + if len(m) >= 3 && m[2] != "" { + unit = strings.ToUpper(m[2]) + } + + multipliers := map[string]float64{ + "B": 1.0 / 1024.0, + "KB": 1000.0 / 1024.0, + "KIB": 1, + "MB": (1000.0 * 1000.0) / 1024.0, + "MIB": 1024, + "GB": (1000.0 * 1000.0 * 1000.0) / 1024.0, + "GIB": 1024 * 1024, + "TB": (1000.0 * 1000.0 * 1000.0 * 1000.0) / 1024.0, + "TIB": 1024 * 1024 * 1024, + } + + mult, ok := multipliers[unit] + if !ok { + mult = 1 + } + + return int(value * mult) +} + +func parseCgroupMemory(memStr string) int { + memStr = strings.TrimSpace(memStr) + if memStr == "" || memStr == "max" { + return 0 + } + + memBytes, err := strconv.ParseUint(memStr, 10, 64) + if err != nil { + return 0 + } + + return int(memBytes / 1024) +} + +func parseCgroupCPU(cpuStr string) int { + cpuStr = strings.TrimSpace(cpuStr) + if cpuStr == "" { + return 0 + } + + parts := strings.Fields(cpuStr) + if len(parts) != 2 || parts[0] == "max" { + return 0 + } + + quota, err := strconv.Atoi(parts[0]) + if err != nil { + return 0 + } + period, err := strconv.Atoi(parts[1]) + if err != nil || period == 0 { + return 0 + } + + return (quota * 1000) / period +} diff --git a/src/yangerd/internal/collector/containers_test.go b/src/yangerd/internal/collector/containers_test.go new file mode 100644 index 000000000..b0dde4bc5 --- /dev/null +++ b/src/yangerd/internal/collector/containers_test.go @@ -0,0 +1,437 @@ +package collector + +import ( + "encoding/json" + "fmt" + "testing" + + "github.com/kernelkit/infix/src/yangerd/internal/testutil" +) + +func collectContainers(t *testing.T, runner *testutil.MockRunner, fs *testutil.MockFileReader) map[string]interface{} { + t.Helper() + + raw := CollectContainers(runner, fs) + if raw == nil { + t.Fatal("no container data collected") + } + + out := make(map[string]interface{}) + if err := json.Unmarshal(raw, &out); err != nil { + t.Fatalf("unmarshal containers: %v", err) + } + + return out +} + +func containerList(t *testing.T, data map[string]interface{}) []interface{} { + t.Helper() + + containers, ok := data["container"].([]interface{}) + if !ok { + t.Fatalf("missing container list: %v", data) + } + + return containers +} + +func TestContainerBasicInfo(t *testing.T) { + runner := &testutil.MockRunner{ + Results: map[string][]byte{ + "podman ps -a --format=json": []byte(`[ + { + "Names": ["web"], + "Id": "abc123", + "Image": "docker.io/library/nginx:latest", + "ImageID": "sha256:image", + "State": "running", + "Status": "Up 2 hours", + "Command": ["nginx", "-g", "daemon off;"], + "Networks": ["podman0"], + "Ports": [] + } + ]`), + "podman inspect web": []byte(`[{"Path":"nginx","Args":["-g","daemon off;"]}]`), + "podman stats --no-stream --format json --no-reset web": []byte(`[]`), + }, + Errors: map[string]error{}, + } + + fs := &testutil.MockFileReader{Files: map[string][]byte{}, Globs: map[string][]string{}} + + out := collectContainers(t, runner, fs) + containers := containerList(t, out) + if len(containers) != 1 { + t.Fatalf("expected 1 container, got %d", len(containers)) + } + + c := containers[0].(map[string]interface{}) + if c["name"] != "web" { + t.Fatalf("name: expected web, got %v", c["name"]) + } + if c["id"] != "abc123" { + t.Fatalf("id: expected abc123, got %v", c["id"]) + } + if c["image"] != "docker.io/library/nginx:latest" { + t.Fatalf("image mismatch: %v", c["image"]) + } + if c["status"] != "Up 2 hours" { + t.Fatalf("status mismatch: %v", c["status"]) + } + if c["cmdline"] != "nginx -g daemon off;" { + t.Fatalf("cmdline mismatch: %v", c["cmdline"]) + } + if c["running"] != true { + t.Fatalf("running expected true, got %v", c["running"]) + } +} + +func TestContainerHostNetwork(t *testing.T) { + runner := &testutil.MockRunner{ + Results: map[string][]byte{ + "podman ps -a --format=json": []byte(`[ + { + "Names": ["hostnet"], + "Id": "id1", + "Image": "img", + "ImageID": "sha256:1", + "State": "running", + "Status": "Up", + "Command": ["sleep", "60"], + "Networks": ["podman0"], + "Ports": [{"host_ip":"", "host_port":8080, "container_port":80, "protocol":"tcp"}] + } + ]`), + "podman inspect hostnet": []byte(`[{"NetworkSettings":{"Networks":{"host":{}}}}]`), + "podman stats --no-stream --format json --no-reset hostnet": []byte(`[]`), + }, + Errors: map[string]error{}, + } + + fs := &testutil.MockFileReader{Files: map[string][]byte{}, Globs: map[string][]string{}} + out := collectContainers(t, runner, fs) + + c := containerList(t, out)[0].(map[string]interface{}) + net, ok := c["network"].(map[string]interface{}) + if !ok { + t.Fatalf("missing network: %v", c) + } + if net["host"] != true { + t.Fatalf("expected host network true, got %v", net["host"]) + } +} + +func TestContainerBridgeNetwork(t *testing.T) { + runner := &testutil.MockRunner{ + Results: map[string][]byte{ + "podman ps -a --format=json": []byte(`[ + { + "Names": ["bridge"], + "Id": "id2", + "Image": "img", + "ImageID": "sha256:2", + "State": "running", + "Status": "Up", + "Command": ["app"], + "Networks": ["podman0", "br0"], + "Ports": [ + {"host_ip":"127.0.0.1", "host_port":8080, "container_port":80, "protocol":"tcp"}, + {"host_ip":"", "host_port":8443, "container_port":443, "protocol":"tcp"} + ] + } + ]`), + "podman inspect bridge": []byte(`[{"NetworkSettings":{"Networks":{"bridge":{}}}}]`), + "podman stats --no-stream --format json --no-reset bridge": []byte(`[]`), + }, + Errors: map[string]error{}, + } + + fs := &testutil.MockFileReader{Files: map[string][]byte{}, Globs: map[string][]string{}} + out := collectContainers(t, runner, fs) + + c := containerList(t, out)[0].(map[string]interface{}) + net := c["network"].(map[string]interface{}) + + ifaces := net["interface"].([]interface{}) + if len(ifaces) != 2 { + t.Fatalf("expected 2 interfaces, got %d", len(ifaces)) + } + if ifaces[0].(map[string]interface{})["name"] != "podman0" { + t.Fatalf("interface[0] mismatch: %v", ifaces[0]) + } + + publish := net["publish"].([]interface{}) + if len(publish) != 2 { + t.Fatalf("expected 2 published ports, got %d", len(publish)) + } + if publish[0] != "127.0.0.1:8080:80/tcp" { + t.Fatalf("publish[0] mismatch: %v", publish[0]) + } + if publish[1] != "8443:443/tcp" { + t.Fatalf("publish[1] mismatch: %v", publish[1]) + } +} + +func TestContainerCgroupLimits(t *testing.T) { + cgroupPath := "/machine.slice/libpod-abc.scope" + runner := &testutil.MockRunner{ + Results: map[string][]byte{ + "podman ps -a --format=json": []byte(`[ + {"Names":["limited"],"Id":"id3","Image":"img","ImageID":"sha256:3","State":"exited","Status":"Exited","Command":["app"],"Networks":[],"Ports":[]} + ]`), + "podman inspect limited": []byte(fmt.Sprintf(`[{"State":{"CgroupPath":%q},"NetworkSettings":{"Networks":{"bridge":{}}}}]`, cgroupPath)), + }, + Errors: map[string]error{}, + } + + fs := &testutil.MockFileReader{ + Files: map[string][]byte{ + "/sys/fs/cgroup/machine.slice/libpod-abc.scope/memory.max": []byte("1073741824\n"), + "/sys/fs/cgroup/machine.slice/libpod-abc.scope/cpu.max": []byte("200000 100000\n"), + }, + Globs: map[string][]string{}, + } + + out := collectContainers(t, runner, fs) + c := containerList(t, out)[0].(map[string]interface{}) + + limit, ok := c["resource-limit"].(map[string]interface{}) + if !ok { + t.Fatalf("missing resource-limit: %v", c) + } + if limit["memory"] != "1048576" { + t.Fatalf("memory limit expected 1048576, got %v", limit["memory"]) + } + if toInt(limit["cpu"]) != 2000 { + t.Fatalf("cpu limit expected 2000, got %v", limit["cpu"]) + } +} + +func TestContainerCgroupUnlimited(t *testing.T) { + cgroupPath := "/machine.slice/libpod-unlimited.scope" + runner := &testutil.MockRunner{ + Results: map[string][]byte{ + "podman ps -a --format=json": []byte(`[ + {"Names":["nolimit"],"Id":"id4","Image":"img","ImageID":"sha256:4","State":"exited","Status":"Exited","Command":["app"],"Networks":[],"Ports":[]} + ]`), + "podman inspect nolimit": []byte(fmt.Sprintf(`[{"State":{"CgroupPath":%q},"NetworkSettings":{"Networks":{"bridge":{}}}}]`, cgroupPath)), + }, + Errors: map[string]error{}, + } + + fs := &testutil.MockFileReader{ + Files: map[string][]byte{ + "/sys/fs/cgroup/machine.slice/libpod-unlimited.scope/memory.max": []byte("max\n"), + "/sys/fs/cgroup/machine.slice/libpod-unlimited.scope/cpu.max": []byte("max 100000\n"), + }, + Globs: map[string][]string{}, + } + + out := collectContainers(t, runner, fs) + c := containerList(t, out)[0].(map[string]interface{}) + if _, ok := c["resource-limit"]; ok { + t.Fatalf("resource-limit should be omitted for unlimited cgroup values: %v", c["resource-limit"]) + } +} + +func TestContainerResourceStats(t *testing.T) { + runner := &testutil.MockRunner{ + Results: map[string][]byte{ + "podman ps -a --format=json": []byte(`[ + { + "Names":["stats"],"Id":"id5","Image":"img","ImageID":"sha256:5", + "State":"running","Status":"Up","Command":["app"],"Networks":["podman0"],"Ports":[] + } + ]`), + "podman inspect stats": []byte(`[{"NetworkSettings":{"Networks":{"bridge":{}}}}]`), + "podman stats --no-stream --format json --no-reset stats": []byte(`[ + { + "mem_usage":"123.4MB / 1.5GB", + "cpu_percent":"12.34%", + "block_io":"1.2MB / 3.4GB", + "net_io":"1.2MB / 3.4GB", + "pids":5 + } + ]`), + }, + Errors: map[string]error{}, + } + + fs := &testutil.MockFileReader{Files: map[string][]byte{}, Globs: map[string][]string{}} + out := collectContainers(t, runner, fs) + + c := containerList(t, out)[0].(map[string]interface{}) + usage, ok := c["resource-usage"].(map[string]interface{}) + if !ok { + t.Fatalf("missing resource-usage: %v", c) + } + + if usage["memory"] != "120507" { + t.Fatalf("memory usage expected 120507, got %v", usage["memory"]) + } + if usage["cpu"] != "12.34" { + t.Fatalf("cpu usage expected 12.34, got %v", usage["cpu"]) + } + bio := usage["block-io"].(map[string]interface{}) + if bio["read"] != "1171" { + t.Fatalf("block-io read expected 1171, got %v", bio["read"]) + } + if bio["write"] != "3320312" { + t.Fatalf("block-io write expected 3320312, got %v", bio["write"]) + } + nio := usage["net-io"].(map[string]interface{}) + if nio["received"] != "1171" { + t.Fatalf("net-io received expected 1171, got %v", nio["received"]) + } + if nio["sent"] != "3320312" { + t.Fatalf("net-io sent expected 3320312, got %v", nio["sent"]) + } + if toInt(usage["pids"]) != 5 { + t.Fatalf("pids expected 5, got %v", usage["pids"]) + } +} + +func TestContainerStopped(t *testing.T) { + runner := &testutil.MockRunner{ + Results: map[string][]byte{ + "podman ps -a --format=json": []byte(`[ + { + "Names":["stopped"],"Id":"id6","Image":"img","ImageID":"sha256:6", + "State":"exited","Status":"Exited (0)","Command":["app"], + "Networks":["podman0"],"Ports":[{"host_ip":"","host_port":8080,"container_port":80,"protocol":"tcp"}] + } + ]`), + "podman inspect stopped": []byte(`[{"NetworkSettings":{"Networks":{"bridge":{}}}}]`), + }, + Errors: map[string]error{}, + } + + fs := &testutil.MockFileReader{Files: map[string][]byte{}, Globs: map[string][]string{}} + out := collectContainers(t, runner, fs) + + c := containerList(t, out)[0].(map[string]interface{}) + if _, ok := c["resource-usage"]; ok { + t.Fatalf("stopped container must not include resource-usage: %v", c["resource-usage"]) + } + + net := c["network"].(map[string]interface{}) + publish := net["publish"].([]interface{}) + if len(publish) != 0 { + t.Fatalf("stopped container must not include published ports, got %v", publish) + } +} + +func TestContainerMultiple(t *testing.T) { + runner := &testutil.MockRunner{ + Results: map[string][]byte{ + "podman ps -a --format=json": []byte(`[ + {"Names":["one"],"Id":"id7","Image":"img1","ImageID":"sha256:7","State":"running","Status":"Up","Command":["a"],"Networks":[],"Ports":[]}, + {"Names":["two"],"Id":"id8","Image":"img2","ImageID":"sha256:8","State":"exited","Status":"Exited","Command":["b"],"Networks":[],"Ports":[]} + ]`), + "podman inspect one": []byte(`[{}]`), + "podman stats --no-stream --format json --no-reset one": []byte(`[]`), + "podman inspect two": []byte(`[{}]`), + }, + Errors: map[string]error{}, + } + + fs := &testutil.MockFileReader{Files: map[string][]byte{}, Globs: map[string][]string{}} + out := collectContainers(t, runner, fs) + + containers := containerList(t, out) + if len(containers) != 2 { + t.Fatalf("expected 2 containers, got %d", len(containers)) + } + if containers[0].(map[string]interface{})["name"] != "one" { + t.Fatalf("first container name mismatch: %v", containers[0]) + } + if containers[1].(map[string]interface{})["name"] != "two" { + t.Fatalf("second container name mismatch: %v", containers[1]) + } +} + +func TestParseSizeKiB(t *testing.T) { + tests := []struct { + name string + input string + expected int + }{ + {name: "mb", input: "1.5MB", expected: 1464}, + {name: "kb", input: "512kB", expected: 500}, + {name: "gib", input: "2GiB", expected: 2097152}, + {name: "mib", input: "64MiB", expected: 65536}, + {name: "bytes", input: "2048B", expected: 2}, + {name: "empty", input: "", expected: 0}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := parseSizeKiB(tt.input) + if got != tt.expected { + t.Errorf("parseSizeKiB(%q): expected %d, got %d", tt.input, tt.expected, got) + } + }) + } +} + +func TestParseCgroupMemory(t *testing.T) { + tests := []struct { + name string + input string + expected int + }{ + {name: "max", input: "max", expected: 0}, + {name: "bytes", input: "1073741824", expected: 1048576}, + {name: "empty", input: "", expected: 0}, + {name: "invalid", input: "abc", expected: 0}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := parseCgroupMemory(tt.input) + if got != tt.expected { + t.Errorf("parseCgroupMemory(%q): expected %d, got %d", tt.input, tt.expected, got) + } + }) + } +} + +func TestParseCgroupCPU(t *testing.T) { + tests := []struct { + name string + input string + expected int + }{ + {name: "max", input: "max 100000", expected: 0}, + {name: "limited", input: "50000 100000", expected: 500}, + {name: "empty", input: "", expected: 0}, + {name: "invalid", input: "abc", expected: 0}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := parseCgroupCPU(tt.input) + if got != tt.expected { + t.Errorf("parseCgroupCPU(%q): expected %d, got %d", tt.input, tt.expected, got) + } + }) + } +} + +func TestContainerGracefulDegradation(t *testing.T) { + runner := &testutil.MockRunner{ + Results: map[string][]byte{}, + Errors: map[string]error{ + "podman ps -a --format=json": fmt.Errorf("podman not found"), + }, + } + + fs := &testutil.MockFileReader{Files: map[string][]byte{}, Globs: map[string][]string{}} + + // With no containers there must be no data, not a bare + // {"container":[]} node -- otherwise an enabled-but-idle container + // feature surfaces as operational data. + if raw := CollectContainers(runner, fs); raw != nil { + t.Fatalf("expected no container data when podman ps fails, got %s", raw) + } +} diff --git a/src/yangerd/internal/collector/hardware.go b/src/yangerd/internal/collector/hardware.go new file mode 100644 index 000000000..79dc7218f --- /dev/null +++ b/src/yangerd/internal/collector/hardware.go @@ -0,0 +1,1408 @@ +package collector + +import ( + "bufio" + "bytes" + "context" + "encoding/json" + "fmt" + "log" + "net" + "path/filepath" + "regexp" + "sort" + "strconv" + "strings" + "sync" + "time" + "unicode" + + "github.com/kernelkit/infix/src/yangerd/internal/nl80211" + "github.com/kernelkit/infix/src/yangerd/internal/tree" +) + +// Per-port devices: mt7915_phy0 -> phy0, marvell_alaska_phy7 -> phy7 +var hwPortDeviceRe = regexp.MustCompile(`.*_((?:phy|sfp)\d*)$`) +var hwTrailingNumUnderscoreRe = regexp.MustCompile(`_(\d+)$`) +var hwPhyNumRe = regexp.MustCompile(`(\d+)$`) + +const cpuComponent = "cpu" + +const hardwareKey = "ietf-hardware:hardware" + +// hwmon device names and thermal zone types that report an SoC die +// temperature, after normalization: a plain cpu/soc/core, Intel and AMD +// (coretemp, k10temp), Microchip SparX-5 and LAN969x (s5-temp), or a +// Marvell CN913x application (ap) or communication (cp) processor +// cluster, optionally as the "-thermal" zone the DT names it. Anything +// else after the dash (cpu-fan, soc-vdd) is a different device. +// +// Recognizing vendor names cannot be avoided, but this is the only place +// it happens. Northbound, the sensors are found through the class of +// their parent component, see doc/hardware.md. +var socTempSourceRe = regexp.MustCompile(`^(cpu\d*|soc\d*|core\d*|coretemp|k10temp|s5-temp|ap|cp\d+)(-thermal(-.*)?)?$`) + +// HardwareCollector gathers ietf-hardware operational data. The +// inventory, radios and GPS receivers are polled; sensors are read when +// someone asks, see Live. +type HardwareCollector struct { + cmd CommandRunner + fs FileReader + interval time.Duration + + enableWifi bool + enableGPS bool + + mu sync.Mutex + gps []interface{} // last GPS poll + radios []interface{} // radio capabilities, without survey + radioIf map[string]string // radio name -> interface for its survey + wifiInfo map[string]map[string]interface{} // PHY info the radios were built from + reported map[string]bool // duplicate names already warned about + + radioRefresh chan struct{} +} + +// NewHardwareCollector creates a HardwareCollector with the given dependencies. +func NewHardwareCollector(cmd CommandRunner, fs FileReader, interval time.Duration, enableWifi, enableGPS bool) *HardwareCollector { + return &HardwareCollector{ + cmd: cmd, + fs: fs, + interval: interval, + enableWifi: enableWifi, + enableGPS: enableGPS, + reported: make(map[string]bool), + + radioRefresh: make(chan struct{}, 1), + } +} + +// Name implements Collector. +func (c *HardwareCollector) Name() string { return "hardware" } + +// Interval implements Collector. +func (c *HardwareCollector) Interval() time.Duration { return c.interval } + +// Collect implements Collector. It produces one tree key: +// "ietf-hardware:hardware". +func (c *HardwareCollector) Collect(ctx context.Context, t *tree.Tree) error { + var gps []interface{} + if c.enableGPS { + gps = c.gpsReceiverComponents(ctx) + } + + c.mu.Lock() + c.gps = gps + c.mu.Unlock() + + // Everything else is read by Live when asked, but the tree only + // calls a provider for a key that exists. + if t.GetCached(hardwareKey) == nil { + t.Set(hardwareKey, json.RawMessage(`{}`)) + } + return nil +} + +// RequestRadioRefresh asks RunRadios to rebuild the radio capabilities. +// Called from nl80211 events: a phy came or went, the regulatory domain +// or the set of interfaces on a phy changed. +func (c *HardwareCollector) RequestRadioRefresh() { + select { + case c.radioRefresh <- struct{}{}: + default: + } +} + +// RunRadios keeps the radio capabilities, which only change on nl80211 +// events, built from the kernel. It builds them once at start. +func (c *HardwareCollector) RunRadios(ctx context.Context) error { + if !c.enableWifi { + <-ctx.Done() + return ctx.Err() + } + for { + radios, ifaces, wifiInfo := c.radioCapabilities(ctx) + c.mu.Lock() + c.radios, c.radioIf, c.wifiInfo = radios, ifaces, wifiInfo + c.mu.Unlock() + + select { + case <-ctx.Done(): + return ctx.Err() + case <-c.radioRefresh: + } + } +} + +// Live is the tree provider for the hardware key. The inventory is a +// cheap read of system.json and sysfs, sensors and radio surveys have +// no events, so all of it is read when asked. +func (c *HardwareCollector) Live() json.RawMessage { + return c.assemble(context.Background()) +} + +func (c *HardwareCollector) assemble(ctx context.Context) json.RawMessage { + systemjson := c.readSystemJSON() + inventory := make([]interface{}, 0) + inventory = append(inventory, c.motherboardComponent(systemjson)...) + inventory = append(inventory, c.vpdComponents(systemjson)...) + inventory = append(inventory, c.usbPortComponents(systemjson)...) + + c.mu.Lock() + radios := cloneComponents(c.radios) + ifaces := c.radioIf + gps := cloneComponents(c.gps) + wifiInfo := c.wifiInfo + c.mu.Unlock() + + c.addSurveys(ctx, radios, ifaces) + sensors := c.sensorComponents(ctx, wifiInfo) + + components := make([]interface{}, 0, len(inventory)+len(sensors)+len(radios)+len(gps)+1) + components = append(components, inventory...) + components = append(components, cpuComponentFor(sensors)...) + components = append(components, sensors...) + components = append(components, radios...) + components = append(components, gps...) + + data, err := json.Marshal(map[string]interface{}{ + "component": c.uniqueNames(components), + }) + if err != nil { + return nil + } + return data +} + +// addSurveys adds each radio's channel survey, read now: the counters +// move all the time and nl80211 sends no events for them. radios are +// clones, the wifi-radio container is copied before it is changed. +func (c *HardwareCollector) addSurveys(ctx context.Context, radios []interface{}, ifaces map[string]string) { + if len(radios) == 0 { + return + } + client, err := nl80211.Dial() + if err != nil { + return + } + defer client.Close() + + for _, raw := range radios { + component, ok := raw.(map[string]interface{}) + if !ok { + continue + } + name, _ := component["name"].(string) + channels := c.surveyData(ctx, client, ifaces[name]) + if len(channels) == 0 { + continue + } + radio := map[string]interface{}{} + if old, ok := component["infix-hardware:wifi-radio"].(map[string]interface{}); ok { + for k, v := range old { + radio[k] = v + } + } + radio["survey"] = map[string]interface{}{"channel": channels} + component["infix-hardware:wifi-radio"] = radio + } +} + +// cloneComponents copies the component maps, uniqueNames may rename +// one and the snapshot must stay as polled. +func cloneComponents(components []interface{}) []interface{} { + out := make([]interface{}, 0, len(components)) + for _, raw := range components { + if component, ok := raw.(map[string]interface{}); ok { + clone := make(map[string]interface{}, len(component)) + for k, v := range component { + clone[k] = v + } + out = append(out, clone) + continue + } + out = append(out, raw) + } + return out +} + +// sensorComponents reads every sensor. Thermal zones first: the +// kernel mirrors each one as an hwmon device, which carries nothing the +// zone does not. +func (c *HardwareCollector) sensorComponents(ctx context.Context, wifiInfo map[string]map[string]interface{}) []interface{} { + sensors := c.thermalSensorComponents(ctx) + mirrored := make(map[string]bool, len(sensors)) + for _, raw := range sensors { + if name, ok := raw.(map[string]interface{})["name"].(string); ok { + mirrored[name] = true + } + } + sensors = append(sensors, c.hwmonSensorComponents(ctx, mirrored)...) + return adoptWifiSensors(sensors, wifiInfo) +} + +func (c *HardwareCollector) readSystemJSON() map[string]interface{} { + data, err := c.fs.ReadFile("/run/system.json") + if err != nil { + return map[string]interface{}{} + } + + out := make(map[string]interface{}) + if err := json.Unmarshal(data, &out); err != nil { + log.Printf("collector hardware: system.json: %v", err) + return map[string]interface{}{} + } + + return out +} + +func (c *HardwareCollector) motherboardComponent(systemjson map[string]interface{}) []interface{} { + if len(systemjson) == 0 { + return nil + } + + component := map[string]interface{}{ + "name": "mainboard", + "class": "iana-hardware:chassis", + "state": map[string]interface{}{ + "admin-state": "unknown", + "oper-state": "enabled", + }, + } + + if v, ok := systemjson["vendor"].(string); ok && v != "" { + component["mfg-name"] = v + } + if v, ok := systemjson["product-name"].(string); ok && v != "" { + component["model-name"] = v + } + if v, ok := systemjson["serial-number"].(string); ok && v != "" { + component["serial-num"] = v + } + if v, ok := systemjson["part-number"].(string); ok && v != "" { + component["hardware-rev"] = v + } + if v, ok := systemjson["mac-address"].(string); ok && v != "" { + component["infix-hardware:phys-address"] = v + } + + return []interface{}{component} +} + +func vpdVendorExtensions(data interface{}) []interface{} { + raw, ok := data.([]interface{}) + if !ok { + return nil + } + + vendorExtensions := make([]interface{}, 0, len(raw)) + for _, item := range raw { + pair, ok := item.([]interface{}) + if !ok || len(pair) < 2 { + continue + } + vendorExtensions = append(vendorExtensions, map[string]interface{}{ + "iana-enterprise-number": pair[0], + "extension-data": pair[1], + }) + } + + return vendorExtensions +} + +func (c *HardwareCollector) vpdComponents(systemjson map[string]interface{}) []interface{} { + vpdRaw, ok := systemjson["vpd"].(map[string]interface{}) + if !ok { + return nil + } + + components := make([]interface{}, 0, len(vpdRaw)) + for _, vpdItemRaw := range vpdRaw { + vpdItem, ok := vpdItemRaw.(map[string]interface{}) + if !ok { + continue + } + + component := map[string]interface{}{ + "class": "infix-hardware:vpd", + "infix-hardware:vpd-data": map[string]interface{}{}, + } + + // Board authors name these in the device tree, as "cpu", "power", + // "product", short words that collide with everything else sharing + // the component namespace. Say what they are. + if board, ok := vpdItem["board"].(string); ok && board != "" { + component["name"] = "vpd-" + board + } + + dataRaw, ok := vpdItem["data"].(map[string]interface{}) + if ok { + if mfgDateStr, ok := dataRaw["manufacture-date"].(string); ok && mfgDateStr != "" { + if mfgDate, err := time.Parse("01/02/2006 15:04:05", mfgDateStr); err == nil { + component["mfg-date"] = mfgDate.UTC().Format("2006-01-02T15:04:05Z") + } + } + + if mfg, ok := dataRaw["manufacturer"].(string); ok && mfg != "" { + component["mfg-name"] = mfg + } + if model, ok := dataRaw["product-name"].(string); ok && model != "" { + component["model-name"] = model + } + if serial, ok := dataRaw["serial-number"].(string); ok && serial != "" { + component["serial-num"] = serial + } + + vpdData, ok := component["infix-hardware:vpd-data"].(map[string]interface{}) + if !ok { + vpdData = make(map[string]interface{}) + component["infix-hardware:vpd-data"] = vpdData + } + for key, val := range dataRaw { + if val == nil { + continue + } + if key == "vendor-extension" { + if ext := vpdVendorExtensions(val); len(ext) > 0 { + vpdData["infix-hardware:vendor-extension"] = ext + } + continue + } + vpdData[key] = val + } + } + + if _, ok := component["name"]; ok { + components = append(components, component) + } + } + + return components +} + +func (c *HardwareCollector) usbPortComponents(systemjson map[string]interface{}) []interface{} { + usbPortsRaw, ok := systemjson["usb-ports"].([]interface{}) + if !ok { + return nil + } + + components := make([]interface{}, 0, len(usbPortsRaw)) + for _, usbPortRaw := range usbPortsRaw { + usbPort, ok := usbPortRaw.(map[string]interface{}) + if !ok { + continue + } + + name, ok := usbPort["name"].(string) + if !ok || name == "" { + continue + } + path, ok := usbPort["path"].(string) + if !ok || path == "" { + continue + } + + authorizedDefault, err := c.fs.ReadFile(path + "/authorized_default") + if err != nil { + continue + } + + state := "locked" + if strings.TrimSpace(string(authorizedDefault)) == "1" { + state = "unlocked" + } + + components = append(components, map[string]interface{}{ + "name": name, + "class": "infix-hardware:usb", + "state": map[string]interface{}{ + "admin-state": state, + "oper-state": "enabled", + }, + }) + } + + return components +} + +// normalizeSensorName makes a list key out of a device name: +// sfp_2 -> sfp2, mt7915_phy0 -> phy0, cpu_thermal -> cpu-thermal. A +// thermal zone and the hwmon device the kernel mirrors it as differ only +// in their separators, so the two spellings of one sensor come out +// identical, which is how the mirror is spotted. +func normalizeSensorName(name string) string { + if m := hwPortDeviceRe.FindStringSubmatch(name); len(m) > 1 { + name = m[1] + } + + name = hwTrailingNumUnderscoreRe.ReplaceAllString(name, "$1") + return strings.ReplaceAll(name, "_", "-") +} + +// cpuComponentFor is the SoC that die temperature sensors belong to. Only +// created when something references it, boards without a die sensor have +// nothing to say about their SoC. +func cpuComponentFor(sensors []interface{}) []interface{} { + for _, raw := range sensors { + if sensor, ok := raw.(map[string]interface{}); ok && sensor["parent"] == cpuComponent { + return []interface{}{map[string]interface{}{ + "name": cpuComponent, + "class": "iana-hardware:cpu", + "parent": "mainboard", + "state": map[string]interface{}{ + "admin-state": "unknown", + "oper-state": "enabled", + }, + }} + } + } + return nil +} + +// uniqueNames renames duplicate component names "-1", "-2" +// and so on. Components are keyed by name, so a duplicate fails every +// client parsing the tree. Producers avoid collisions by construction, +// this is the net under them. A renamed component keeps any children +// pointing at the original name, so it is a last resort, not a mechanism +// to rely on. Each name is warned about once, this runs on every GET. +func (c *HardwareCollector) uniqueNames(components []interface{}) []interface{} { + taken := make(map[string]bool, len(components)) + + for _, raw := range components { + component, ok := raw.(map[string]interface{}) + if !ok { + continue + } + name, _ := component["name"].(string) + if !taken[name] { + taken[name] = true + continue + } + + unique := name + for seq := 1; taken[unique]; seq++ { + unique = fmt.Sprintf("%s-%d", name, seq) + } + c.mu.Lock() + reported := c.reported[name] + c.reported[name] = true + c.mu.Unlock() + if !reported { + log.Printf("collector hardware: duplicate component %q, renaming one of them %q", name, unique) + } + component["name"] = unique + taken[unique] = true + } + + return components +} + +// titleCase capitalizes the first letter of every word, like Python's +// str.title(): "volts-DC" -> "Volts-Dc". +func titleCase(s string) string { + out := []rune(strings.ToLower(s)) + start := true + for i, r := range out { + if start && unicode.IsLetter(r) { + out[i] = unicode.ToUpper(r) + } + start = !unicode.IsLetter(r) + } + return string(out) +} + +func humanizeSensorLabel(label string) string { + if label == "" { + return "" + } + parts := strings.Fields(strings.ReplaceAll(label, "_", " ")) + out := make([]string, 0, len(parts)) + for _, part := range parts { + if part == strings.ToUpper(part) { + out = append(out, part) + continue + } + r := []rune(strings.ToLower(part)) + if len(r) == 0 { + continue + } + r[0] = []rune(strings.ToUpper(string(r[0])))[0] + out = append(out, string(r)) + } + return strings.Join(out, " ") +} + +func sensorComponent(name string, value int, valueType, valueScale, label string) map[string]interface{} { + component := map[string]interface{}{ + "name": name, + "class": "iana-hardware:sensor", + "sensor-data": map[string]interface{}{ + "value": value, + "value-type": valueType, + "value-scale": valueScale, + "value-precision": 0, + "value-timestamp": yangDateTime(time.Now()), + "oper-status": "ok", + }, + } + + if d := humanizeSensorLabel(label); d != "" { + component["description"] = d + } + + return component +} + +// listDir names the entries of dir, sorted as the filesystem lists them. +func (c *HardwareCollector) listDir(dir string) ([]string, error) { + matches, err := c.fs.Glob(dir + "/*") + if err != nil { + return nil, err + } + names := make([]string, 0, len(matches)) + for _, match := range matches { + names = append(names, filepath.Base(match)) + } + return names, nil +} + +func (c *HardwareCollector) readSensorString(path string) (string, bool) { + data, err := c.fs.ReadFile(path) + if err != nil { + return "", false + } + return strings.TrimSpace(string(data)), true +} + +func (c *HardwareCollector) readSensorInt(path string) (int, bool) { + data, err := c.fs.ReadFile(path) + if err != nil { + return 0, false + } + v, err := strconv.Atoi(strings.TrimSpace(string(data))) + if err != nil { + return 0, false + } + return v, true +} + +func sensorName(baseName, sensorNum string) string { + if sensorNum == "1" || sensorNum == "0" { + return baseName + } + return baseName + sensorNum +} + +func (c *HardwareCollector) wifiPhyInfo(ctx context.Context, client *nl80211.Client) map[string]map[string]interface{} { + phyInfo := make(map[string]map[string]interface{}) + if err := ctx.Err(); err != nil { + return phyInfo + } + + phys, err := client.ListPhys() + if err != nil { + return phyInfo + } + + for _, phy := range phys { + if phy == "" { + continue + } + phyInfo[phy] = map[string]interface{}{ + "band": "Unknown", + "iface": "", + "description": "WiFi Radio", + } + } + + phyNumToName := make(map[string]string) + for phyName := range phyInfo { + m := hwPhyNumRe.FindStringSubmatch(phyName) + if len(m) > 1 { + phyNumToName[m[1]] = phyName + } + } + + devMap, err := client.PhyInterfaces() + if err == nil { + for phyNum, ifaces := range devMap { + phyName, ok := phyNumToName[phyNum] + if !ok { + continue + } + if len(ifaces) == 0 { + continue + } + if entry, ok := phyInfo[phyName]; ok { + entry["iface"] = ifaces[0] + } + } + } + + for phy, info := range phyInfo { + band := strDefault(info["band"], "Unknown") + iface := strDefault(info["iface"], "") + switch { + case iface != "" && band != "Unknown": + info["description"] = "WiFi Radio " + phy + case band != "Unknown": + info["description"] = "WiFi Radio (" + band + ")" + case iface != "": + info["description"] = "WiFi Radio " + phy + default: + info["description"] = "WiFi Radio" + } + } + + return phyInfo +} + +// hwmonSensorComponents lists the hwmon sensors. Devices named in +// mirrored are skipped, see the thermal zones. +func (c *HardwareCollector) hwmonSensorComponents(ctx context.Context, mirrored map[string]bool) []interface{} { + components := make([]interface{}, 0) + deviceSensors := make(map[string][]map[string]interface{}) + order := make([]string, 0) // devices in discovery order, like the kernel lists them + + hwmonEntries, err := c.listDir("/sys/class/hwmon") + if err != nil { + return components + } + + for _, entry := range hwmonEntries { + if !strings.HasPrefix(entry, "hwmon") { + continue + } + hwmonPath := "/sys/class/hwmon/" + entry + + deviceName, ok := c.readSensorString(hwmonPath + "/name") + if !ok || deviceName == "" { + continue + } + // With THERMAL_HWMON the kernel mirrors every thermal zone as an + // hwmon device, named after the zone with the separators changed. + if mirrored[normalizeSensorName(deviceName)] { + continue + } + if devName, ok := c.readSensorString(hwmonPath + "/device/name"); ok && devName != "" { + deviceName = devName + } + + baseName := normalizeSensorName(deviceName) + if baseName == "" { + continue + } + if _, seen := deviceSensors[baseName]; !seen { + deviceSensors[baseName] = nil + order = append(order, baseName) + } + + entries, err := c.listDir(hwmonPath) + if err != nil { + continue + } + + fanFiles := make([]string, 0) + for _, e := range entries { + if strings.HasPrefix(e, "fan") && strings.HasSuffix(e, "_input") { + fanFiles = append(fanFiles, e) + } + } + + for _, e := range entries { + if !strings.HasPrefix(e, "temp") || !strings.HasSuffix(e, "_input") { + continue + } + sensorNum := strings.TrimPrefix(strings.SplitN(e, "_", 2)[0], "temp") + value, ok := c.readSensorInt(hwmonPath + "/" + e) + if !ok { + continue + } + label := "" + sensor := "" + if rawLabel, ok := c.readSensorString(fmt.Sprintf("%s/temp%s_label", hwmonPath, sensorNum)); ok { + label = rawLabel + sensor = baseName + "-" + normalizeSensorName(rawLabel) + } else { + sensor = sensorName(baseName, sensorNum) + } + deviceSensors[baseName] = append(deviceSensors[baseName], sensorComponent(sensor, value, "celsius", "milli", label)) + } + + for _, e := range fanFiles { + sensorNum := strings.TrimPrefix(strings.SplitN(e, "_", 2)[0], "fan") + value, ok := c.readSensorInt(hwmonPath + "/" + e) + if !ok { + continue + } + label := "" + sensor := "" + if rawLabel, ok := c.readSensorString(fmt.Sprintf("%s/fan%s_label", hwmonPath, sensorNum)); ok { + label = rawLabel + sensor = baseName + "-" + normalizeSensorName(rawLabel) + } else { + sensor = sensorName(baseName, sensorNum) + } + deviceSensors[baseName] = append(deviceSensors[baseName], sensorComponent(sensor, value, "rpm", "units", label)) + } + + if len(fanFiles) == 0 { + for _, e := range entries { + if !strings.HasPrefix(e, "pwm") { + continue + } + n := strings.TrimPrefix(e, "pwm") + if _, err := strconv.Atoi(n); err != nil { + continue + } + pwmRaw, ok := c.readSensorInt(hwmonPath + "/" + e) + if !ok { + continue + } + sensorNum := n + value := int((float64(pwmRaw) / 255.0) * 100.0 * 1000.0) + label := "PWM Fan" + sensor := "" + if rawLabel, ok := c.readSensorString(fmt.Sprintf("%s/pwm%s_label", hwmonPath, sensorNum)); ok { + label = rawLabel + sensor = baseName + "-" + normalizeSensorName(rawLabel) + } else { + sensor = sensorName(baseName, sensorNum) + } + deviceSensors[baseName] = append(deviceSensors[baseName], sensorComponent(sensor, value, "other", "milli", label)) + } + } + + for _, e := range entries { + if !strings.HasPrefix(e, "in") || !strings.HasSuffix(e, "_input") { + continue + } + sensorNum := strings.TrimPrefix(strings.SplitN(e, "_", 2)[0], "in") + value, ok := c.readSensorInt(hwmonPath + "/" + e) + if !ok { + continue + } + label := "voltage" + sensor := "" + if rawLabel, ok := c.readSensorString(fmt.Sprintf("%s/in%s_label", hwmonPath, sensorNum)); ok { + label = rawLabel + sensor = baseName + "-" + normalizeSensorName(rawLabel) + } else { + if sensorNum == "0" { + sensor = baseName + "-voltage" + } else { + sensor = baseName + "-voltage" + sensorNum + } + } + deviceSensors[baseName] = append(deviceSensors[baseName], sensorComponent(sensor, value, "volts-DC", "milli", label)) + } + + for _, e := range entries { + if !strings.HasPrefix(e, "curr") || !strings.HasSuffix(e, "_input") { + continue + } + sensorNum := strings.TrimPrefix(strings.SplitN(e, "_", 2)[0], "curr") + value, ok := c.readSensorInt(hwmonPath + "/" + e) + if !ok { + continue + } + label := "current" + sensor := "" + if rawLabel, ok := c.readSensorString(fmt.Sprintf("%s/curr%s_label", hwmonPath, sensorNum)); ok { + label = rawLabel + sensor = baseName + "-" + normalizeSensorName(rawLabel) + } else { + if sensorNum == "1" { + sensor = baseName + "-current" + } else { + sensor = baseName + "-current" + sensorNum + } + } + deviceSensors[baseName] = append(deviceSensors[baseName], sensorComponent(sensor, value, "amperes", "milli", label)) + } + + for _, e := range entries { + if !strings.HasPrefix(e, "power") || !strings.HasSuffix(e, "_input") { + continue + } + sensorNum := strings.TrimPrefix(strings.SplitN(e, "_", 2)[0], "power") + value, ok := c.readSensorInt(hwmonPath + "/" + e) + if !ok { + continue + } + label := "power" + sensor := "" + if rawLabel, ok := c.readSensorString(fmt.Sprintf("%s/power%s_label", hwmonPath, sensorNum)); ok { + label = rawLabel + sensor = baseName + "-" + normalizeSensorName(rawLabel) + } else { + if sensorNum == "1" { + sensor = baseName + "-power" + } else { + sensor = baseName + "-power" + sensorNum + } + } + deviceSensors[baseName] = append(deviceSensors[baseName], sensorComponent(sensor, value, "watts", "micro", label)) + } + } + + for _, baseName := range order { + sensors := deviceSensors[baseName] + if len(sensors) == 0 { + continue + } + parent := "" + switch { + case socTempSourceRe.MatchString(baseName): + // SoC die sensors belong to the CPU, whatever the vendor + // called the hwmon device + parent = cpuComponent + case len(sensors) > 1: + // Multi-sensor devices, like SFP modules, head their own + parent = baseName + components = append(components, map[string]interface{}{ + "name": baseName, + "class": "iana-hardware:module", + }) + } + + for _, sensor := range sensors { + if parent != "" { + sensor["parent"] = parent + } + components = append(components, sensor) + } + } + + return components +} + +// adoptWifiSensors gives a radio its name back. A WiFi PHY's hwmon +// device is named after the radio, so whatever hwmonSensorComponents +// built for it took the radio's name before wifiRadioComponents gets +// there, and uniqueNames would rename the radio instead, breaking the +// wifi/radio leafref that interfaces are bound by. Its sensors hang off +// it, the way the die sensors hang off the CPU, and the module head a +// multi-sensor device would get is dropped: the radio component already +// is one. +func adoptWifiSensors(components []interface{}, wifiInfo map[string]map[string]interface{}) []interface{} { + out := make([]interface{}, 0, len(components)) + for _, raw := range components { + component, ok := raw.(map[string]interface{}) + if !ok { + continue + } + name, _ := component["name"].(string) + if _, radio := wifiInfo[name]; !radio { + out = append(out, component) + continue + } + + if component["class"] == "iana-hardware:module" { + continue // the radio heads its own sensors + } + + kind := "sensor" + if sd, ok := component["sensor-data"].(map[string]interface{}); ok { + kind = strDefault(sd["value-type"], kind) + } + if kind == "celsius" { + component["name"] = name + "-temp" + component["description"] = "Temperature" + } else { + component["name"] = name + "-" + kind + component["description"] = titleCase(kind) + } + component["parent"] = name + out = append(out, component) + } + + return out +} + +func (c *HardwareCollector) thermalSensorComponents(ctx context.Context) []interface{} { + components := make([]interface{}, 0) + + entries, err := c.listDir("/sys/class/thermal") + if err != nil { + return components + } + + for _, entry := range entries { + if !strings.HasPrefix(entry, "thermal_zone") { + continue + } + zonePath := "/sys/class/thermal/" + entry + zoneType, ok := c.readSensorString(zonePath + "/type") + if !ok || zoneType == "" { + continue + } + temp, ok := c.readSensorInt(zonePath + "/temp") + if !ok { + continue + } + + component := sensorComponent(normalizeSensorName(zoneType), temp, "celsius", "milli", "") + if socTempSourceRe.MatchString(component["name"].(string)) { + component["parent"] = cpuComponent + } + components = append(components, component) + } + + return components +} + +func (c *HardwareCollector) surveyData(ctx context.Context, client *nl80211.Client, ifname string) []interface{} { + if err := ctx.Err(); err != nil { + return nil + } + if ifname == "" { + return nil + } + iface, err := net.InterfaceByName(ifname) + if err != nil { + return nil + } + + survey, err := client.Survey(iface.Index) + if err != nil { + return nil + } + + channels := make([]interface{}, 0, len(survey)) + for _, entry := range survey { + channel := map[string]interface{}{ + "frequency": entry["frequency"], + "in-use": entry["in_use"], + } + setIfPresent(channel, "noise", entry, "noise") + setIfPresent(channel, "active-time", entry, "active_time") + setIfPresent(channel, "busy-time", entry, "busy_time") + setIfPresent(channel, "receive-time", entry, "receive_time") + setIfPresent(channel, "transmit-time", entry, "transmit_time") + channels = append(channels, channel) + } + + return channels +} + +func (c *HardwareCollector) phyInfo(ctx context.Context, client *nl80211.Client, phyName string) map[string]interface{} { + if err := ctx.Err(); err != nil { + return map[string]interface{}{} + } + phyInfo, err := client.PhyInfo(phyName) + if err != nil { + return map[string]interface{}{} + } + + return phyInfo +} + +func convertPhyInfo(phyInfo map[string]interface{}) map[string]interface{} { + result := map[string]interface{}{ + "bands": []interface{}{}, + "driver": nil, + "manufacturer": "Unknown", + "max-interfaces": map[string]interface{}{}, + } + + bandsRaw, _ := phyInfo["bands"].([]interface{}) + bands := make([]interface{}, 0, len(bandsRaw)) + for _, bandRaw := range bandsRaw { + band, ok := bandRaw.(map[string]interface{}) + if !ok { + continue + } + bandData := map[string]interface{}{ + "band": strconv.Itoa(toInt(band["band"])), + } + if name := strDefault(band["name"], ""); name != "" { + bandData["name"] = name + } + if v, ok := band["ht_capable"].(bool); ok && v { + bandData["ht-capable"] = true + } + if v, ok := band["vht_capable"].(bool); ok && v { + bandData["vht-capable"] = true + } + if v, ok := band["he_capable"].(bool); ok && v { + bandData["he-capable"] = true + } + bands = append(bands, bandData) + } + result["bands"] = bands + + if driver, ok := phyInfo["driver"].(string); ok && driver != "" { + result["driver"] = driver + } + if manufacturer, ok := phyInfo["manufacturer"].(string); ok && manufacturer != "" { + result["manufacturer"] = manufacturer + } + + maxInterfaces := make(map[string]interface{}) + ifCombRaw, _ := phyInfo["interface_combinations"].([]interface{}) + for _, combRaw := range ifCombRaw { + comb, ok := combRaw.(map[string]interface{}) + if !ok { + continue + } + limitsRaw, _ := comb["limits"].([]interface{}) + for _, limitRaw := range limitsRaw { + limit, ok := limitRaw.(map[string]interface{}) + if !ok { + continue + } + typesRaw, _ := limit["types"].([]interface{}) + hasAP := false + for _, t := range typesRaw { + if s, ok := t.(string); ok && s == "AP" { + hasAP = true + break + } + } + if !hasAP { + continue + } + apMax := toInt(limit["max"]) + if cur, ok := maxInterfaces["ap"]; !ok || apMax > toInt(cur) { + maxInterfaces["ap"] = apMax + } + } + } + result["max-interfaces"] = maxInterfaces + + return result +} + +func channelFromFrequency(freq int) (int, bool) { + switch { + case freq >= 2412 && freq <= 2484: + return (freq - 2407) / 5, true + case freq >= 5170 && freq <= 5825: + return (freq - 5000) / 5, true + case freq >= 5955 && freq <= 7115: + return (freq - 5950) / 5, true + default: + return 0, false + } +} + +// radioCapabilities lists the radios without their survey, the +// interface each survey is read from, and the PHY info they were built +// from so the sensors can be matched to them. +func (c *HardwareCollector) radioCapabilities(ctx context.Context) ([]interface{}, map[string]string, map[string]map[string]interface{}) { + components := make([]interface{}, 0) + ifaces := map[string]string{} + wifiInfo := map[string]map[string]interface{}{} + client, err := nl80211.Dial() + if err != nil { + return components, ifaces, wifiInfo + } + defer client.Close() + + wifiInfo = c.wifiPhyInfo(ctx, client) + + for phyName, phyData := range wifiInfo { + component := map[string]interface{}{ + "name": phyName, + "class": "infix-hardware:wifi", + "description": strDefault(phyData["description"], "WiFi Radio"), + } + + wifiRadioData := make(map[string]interface{}) + iwInfo := c.phyInfo(ctx, client, phyName) + phyDetails := convertPhyInfo(iwInfo) + + if manufacturer := strDefault(phyDetails["manufacturer"], "Unknown"); manufacturer != "Unknown" { + component["mfg-name"] = manufacturer + } + + if bands, ok := phyDetails["bands"].([]interface{}); ok && len(bands) > 0 { + wifiRadioData["bands"] = bands + } + if driver := strDefault(phyDetails["driver"], ""); driver != "" { + wifiRadioData["driver"] = driver + } + if maxIf, ok := phyDetails["max-interfaces"].(map[string]interface{}); ok && len(maxIf) > 0 { + wifiRadioData["max-interfaces"] = maxIf + } + + setIfPresent(wifiRadioData, "max-txpower", iwInfo, "max_txpower") + + supportedChannelsMap := make(map[int]bool) + bandsRaw, _ := iwInfo["bands"].([]interface{}) + for _, bandRaw := range bandsRaw { + band, ok := bandRaw.(map[string]interface{}) + if !ok { + continue + } + freqsRaw, _ := band["frequencies"].([]interface{}) + for _, freqRaw := range freqsRaw { + freq := toInt(freqRaw) + if channel, ok := channelFromFrequency(freq); ok { + supportedChannelsMap[channel] = true + } + } + } + if len(supportedChannelsMap) > 0 { + supported := make([]int, 0, len(supportedChannelsMap)) + for ch := range supportedChannelsMap { + supported = append(supported, ch) + } + sort.Ints(supported) + supportedIface := make([]interface{}, 0, len(supported)) + for _, ch := range supported { + supportedIface = append(supportedIface, ch) + } + wifiRadioData["supported-channels"] = supportedIface + } + + wifiRadioData["num-virtual-interfaces"] = toInt(iwInfo["num_virtual_interfaces"]) + + ifaces[phyName] = strDefault(phyData["iface"], "") + + if len(wifiRadioData) > 0 { + component["infix-hardware:wifi-radio"] = wifiRadioData + } + + components = append(components, component) + } + + return components, ifaces, wifiInfo +} + +func gpsdPoll(ctx context.Context) map[string]interface{} { + dialer := &net.Dialer{Timeout: 500 * time.Millisecond} + conn, err := dialer.DialContext(ctx, "tcp", "127.0.0.1:2947") + if err != nil { + return map[string]interface{}{} + } + defer conn.Close() + + _ = conn.SetDeadline(time.Now().Add(500 * time.Millisecond)) + + reader := bufio.NewReader(conn) + _, _ = reader.ReadBytes('\n') + + if _, err := conn.Write([]byte("?WATCH={\"enable\":true,\"json\":true};\n?POLL;\n")); err != nil { + return map[string]interface{}{} + } + + buf := bytes.Buffer{} + for i := 0; i < 5; i++ { + chunk := make([]byte, 4096) + n, err := conn.Read(chunk) + if err != nil || n == 0 { + break + } + buf.Write(chunk[:n]) + for _, line := range splitLines(buf.String()) { + var msg map[string]interface{} + if json.Unmarshal([]byte(line), &msg) != nil { + continue + } + if cls, ok := msg["class"].(string); ok && cls == "POLL" { + return msg + } + } + } + + return map[string]interface{}{} +} + +func countUsedSatellites(sats []interface{}) int { + used := 0 + for _, satRaw := range sats { + sat, ok := satRaw.(map[string]interface{}) + if !ok { + continue + } + if v, ok := sat["used"].(bool); ok && v { + used++ + } + } + return used +} + +func (c *HardwareCollector) gpsReceiverComponents(ctx context.Context) []interface{} { + components := make([]interface{}, 0) + gpsDevices := make(map[string]map[string]string) + + devPaths, _ := c.fs.Glob("/dev/gps[0-3]") + for _, devPath := range devPaths { + actual, err := c.cmd.Run(ctx, "readlink", "-f", devPath) + if err != nil { + continue + } + actualPath := strings.TrimSpace(string(actual)) + if actualPath == "" { + continue + } + gpsDevices[actualPath] = map[string]string{ + "name": filepath.Base(devPath), + "symlink": devPath, + } + } + + if len(gpsDevices) == 0 { + return components + } + + poll := gpsdPoll(ctx) + active := toInt(poll["active"]) + + tpvByDev := make(map[string]map[string]interface{}) + tpvRaw, _ := poll["tpv"].([]interface{}) + for _, itemRaw := range tpvRaw { + item, ok := itemRaw.(map[string]interface{}) + if !ok { + continue + } + dev, _ := item["device"].(string) + if dev != "" { + tpvByDev[dev] = item + } + } + + skyByDev := make(map[string]map[string]interface{}) + skyRaw, _ := poll["sky"].([]interface{}) + for _, itemRaw := range skyRaw { + item, ok := itemRaw.(map[string]interface{}) + if !ok { + continue + } + dev, _ := item["device"].(string) + if dev != "" { + skyByDev[dev] = item + } + } + + for actualPath, dev := range gpsDevices { + name := dev["name"] + symlink := dev["symlink"] + + component := map[string]interface{}{ + "name": name, + "class": "infix-hardware:gps", + "description": "GPS/GNSS Receiver", + } + + gpsData := make(map[string]interface{}) + gpsData["device"] = symlink + + tpv := tpvByDev[actualPath] + if tpv == nil { + tpv = tpvByDev[symlink] + } + if tpv == nil && len(tpvByDev) == 1 { + for _, v := range tpvByDev { + tpv = v + } + } + + sky := skyByDev[actualPath] + if sky == nil { + sky = skyByDev[symlink] + } + if sky == nil && len(skyByDev) == 1 { + for _, v := range skyByDev { + sky = v + } + } + + gpsData["activated"] = active > 0 && len(tpv) > 0 + + if driver, ok := tpv["driver"].(string); ok && driver != "" { + gpsData["driver"] = driver + } + + switch toInt(tpv["mode"]) { + case 2: + gpsData["fix-mode"] = "2d" + case 3: + gpsData["fix-mode"] = "3d" + default: + gpsData["fix-mode"] = "none" + } + + if lat, ok := tpv["lat"]; ok { + gpsData["latitude"] = fmt.Sprintf("%.6f", toFloat64(lat)) + } + if lon, ok := tpv["lon"]; ok { + gpsData["longitude"] = fmt.Sprintf("%.6f", toFloat64(lon)) + } + if alt, ok := tpv["altHAE"]; ok { + gpsData["altitude"] = fmt.Sprintf("%.1f", toFloat64(alt)) + } + + satVis := 0 + satUsed := 0 + if sky != nil { + sats, _ := sky["satellites"].([]interface{}) + if len(sats) > 0 { + satVis = len(sats) + satUsed = countUsedSatellites(sats) + } + if satVis == 0 { + satVis = toInt(zeroIfNil(sky["nSat"])) + if satVis == 0 { + satVis = toInt(zeroIfNil(sky["satellites_visible"])) + } + } + if satUsed == 0 { + satUsed = toInt(zeroIfNil(sky["uSat"])) + if satUsed == 0 { + satUsed = toInt(zeroIfNil(sky["satellites_used"])) + } + } + } + + if satVis == 0 { + satVis = toInt(zeroIfNil(tpv["nSat"])) + if satVis == 0 { + satVis = toInt(zeroIfNil(tpv["satellites_visible"])) + } + } + if satUsed == 0 { + satUsed = toInt(zeroIfNil(tpv["uSat"])) + if satUsed == 0 { + satUsed = toInt(zeroIfNil(tpv["satellites_used"])) + } + } + + if satUsed > satVis { + satVis = satUsed + } + gpsData["satellites-visible"] = satVis + gpsData["satellites-used"] = satUsed + + pps, _ := c.fs.Glob(fmt.Sprintf("/dev/pps%s", strings.TrimPrefix(name, "gps"))) + gpsData["pps-available"] = len(pps) > 0 + + component["infix-hardware:gps-receiver"] = gpsData + components = append(components, component) + } + + return components +} + +func toFloat64(v interface{}) float64 { + switch n := v.(type) { + case float64: + return n + case float32: + return float64(n) + case int: + return float64(n) + case int64: + return float64(n) + case json.Number: + f, _ := n.Float64() + return f + case string: + f, _ := strconv.ParseFloat(n, 64) + return f + default: + return 0 + } +} diff --git a/src/yangerd/internal/collector/hardware_test.go b/src/yangerd/internal/collector/hardware_test.go new file mode 100644 index 000000000..797c5bbc6 --- /dev/null +++ b/src/yangerd/internal/collector/hardware_test.go @@ -0,0 +1,692 @@ +package collector + +import ( + "context" + "encoding/json" + "fmt" + "testing" + "time" + + "github.com/kernelkit/infix/src/yangerd/internal/testutil" + "github.com/kernelkit/infix/src/yangerd/internal/tree" +) + +func collectHardware(t *testing.T, c *HardwareCollector) []interface{} { + t.Helper() + + tr := tree.New() + tr.RegisterProvider(hardwareKey, c.Live) + if err := c.Collect(context.Background(), tr); err != nil { + t.Fatalf("Collect failed: %v", err) + } + + raw := tr.Get("ietf-hardware:hardware") + if raw == nil { + t.Fatal("missing ietf-hardware:hardware in tree") + } + + var out map[string]interface{} + if err := json.Unmarshal(raw, &out); err != nil { + t.Fatalf("unmarshal hardware: %v", err) + } + + components, ok := out["component"].([]interface{}) + if !ok { + t.Fatalf("component list missing or invalid: %v", out["component"]) + } + + return components +} + +func getComponentByName(components []interface{}, name string) map[string]interface{} { + for _, c := range components { + m, ok := c.(map[string]interface{}) + if !ok { + continue + } + if m["name"] == name { + return m + } + } + return nil +} + +func containsComponentWithClass(components []interface{}, class string) bool { + for _, c := range components { + m, ok := c.(map[string]interface{}) + if !ok { + continue + } + if m["class"] == class { + return true + } + } + return false +} + +func newHardwareCollector(r *testutil.MockRunner, fs *testutil.MockFileReader) *HardwareCollector { + return NewHardwareCollector(r, fs, 30*time.Second, false, false) +} + +func TestHardwareMotherboard(t *testing.T) { + runner := &testutil.MockRunner{Results: map[string][]byte{}, Errors: map[string]error{}} + fs := &testutil.MockFileReader{Files: map[string][]byte{ + "/run/system.json": []byte(`{"vendor":"Acme","product-name":"Router-1","serial-number":"SN123","part-number":"PN99","mac-address":"00:11:22:33:44:55"}`), + }, Globs: map[string][]string{}} + + components := collectHardware(t, newHardwareCollector(runner, fs)) + mb := getComponentByName(components, "mainboard") + if mb == nil { + t.Fatal("mainboard component not found") + } + + if mb["class"] != "iana-hardware:chassis" { + t.Fatalf("mainboard class: expected chassis, got %v", mb["class"]) + } + if mb["mfg-name"] != "Acme" || mb["model-name"] != "Router-1" || mb["serial-num"] != "SN123" { + t.Fatalf("mainboard identity fields mismatch: %v", mb) + } + if mb["hardware-rev"] != "PN99" { + t.Fatalf("mainboard hardware-rev mismatch: %v", mb["hardware-rev"]) + } + if mb["infix-hardware:phys-address"] != "00:11:22:33:44:55" { + t.Fatalf("mainboard phys-address mismatch: %v", mb["infix-hardware:phys-address"]) + } + + state, ok := mb["state"].(map[string]interface{}) + if !ok { + t.Fatalf("mainboard state missing: %v", mb["state"]) + } + if state["admin-state"] != "unknown" || state["oper-state"] != "enabled" { + t.Fatalf("mainboard state mismatch: %v", state) + } +} + +func TestHardwareVPD(t *testing.T) { + runner := &testutil.MockRunner{Results: map[string][]byte{}, Errors: map[string]error{}} + fs := &testutil.MockFileReader{Files: map[string][]byte{ + "/run/system.json": []byte(`{ + "vpd": { + "slot0": { + "board": "board0", + "data": { + "manufacture-date": "04/11/2026 13:14:15", + "manufacturer": "VPD Inc", + "product-name": "X1", + "serial-number": "VPD-123", + "foo": "bar", + "vendor-extension": [[32473, "aa55"]] + } + } + } + }`), + }, Globs: map[string][]string{}} + + components := collectHardware(t, newHardwareCollector(runner, fs)) + vpd := getComponentByName(components, "vpd-board0") + if vpd == nil { + t.Fatal("vpd component vpd-board0 not found") + } + if vpd["class"] != "infix-hardware:vpd" { + t.Fatalf("vpd class mismatch: %v", vpd["class"]) + } + if vpd["mfg-date"] != "2026-04-11T13:14:15Z" { + t.Fatalf("mfg-date mismatch: %v", vpd["mfg-date"]) + } + if vpd["serial-num"] != "VPD-123" { + t.Fatalf("serial-num mismatch: %v", vpd["serial-num"]) + } + + vpdData, ok := vpd["infix-hardware:vpd-data"].(map[string]interface{}) + if !ok { + t.Fatalf("vpd-data missing: %v", vpd["infix-hardware:vpd-data"]) + } + if vpdData["foo"] != "bar" { + t.Fatalf("vpd-data foo mismatch: %v", vpdData["foo"]) + } + extList, ok := vpdData["infix-hardware:vendor-extension"].([]interface{}) + if !ok || len(extList) != 1 { + t.Fatalf("vendor-extension missing: %v", vpdData["infix-hardware:vendor-extension"]) + } + ext := extList[0].(map[string]interface{}) + if toInt(ext["iana-enterprise-number"]) != 32473 || ext["extension-data"] != "aa55" { + t.Fatalf("vendor-extension mismatch: %v", ext) + } +} + +func TestHardwareUSBPorts(t *testing.T) { + runner := &testutil.MockRunner{Results: map[string][]byte{}, Errors: map[string]error{}} + fs := &testutil.MockFileReader{Files: map[string][]byte{ + "/run/system.json": []byte(`{"usb-ports":[{"name":"usb-a","path":"/sys/devices/usb-a"},{"name":"usb-b","path":"/sys/devices/usb-b"}]}`), + "/sys/devices/usb-a/authorized_default": []byte("1\n"), + "/sys/devices/usb-b/authorized_default": []byte("0\n"), + }, Globs: map[string][]string{}} + + components := collectHardware(t, newHardwareCollector(runner, fs)) + usbA := getComponentByName(components, "usb-a") + usbB := getComponentByName(components, "usb-b") + if usbA == nil || usbB == nil { + t.Fatalf("usb components missing: usb-a=%v usb-b=%v", usbA, usbB) + } + aState := usbA["state"].(map[string]interface{}) + bState := usbB["state"].(map[string]interface{}) + if aState["admin-state"] != "unlocked" || bState["admin-state"] != "locked" { + t.Fatalf("usb admin-state mismatch: a=%v b=%v", aState, bState) + } +} + +func TestHardwareHwmonTemp(t *testing.T) { + runner := &testutil.MockRunner{Results: map[string][]byte{}, Errors: map[string]error{}} + fs := &testutil.MockFileReader{Files: map[string][]byte{ + "/run/system.json": []byte(`{}`), + "/sys/class/hwmon/hwmon0/name": []byte("cpu_thermal\n"), + "/sys/class/hwmon/hwmon0/temp1_input": []byte("42000\n"), + "/sys/class/hwmon/hwmon0/temp1_label": []byte("cpu_temp\n"), + }, Globs: map[string][]string{ + "/sys/class/hwmon/*": {"/sys/class/hwmon/hwmon0"}, + "/sys/class/hwmon/hwmon0/*": {"/sys/class/hwmon/hwmon0/name", "/sys/class/hwmon/hwmon0/temp1_input", "/sys/class/hwmon/hwmon0/temp1_label"}, + }} + + components := collectHardware(t, newHardwareCollector(runner, fs)) + if !containsComponentWithClass(components, "iana-hardware:sensor") { + t.Fatalf("expected at least one sensor component: %v", components) + } + sensor := getComponentByName(components, "cpu-thermal-cpu-temp") + if sensor == nil { + t.Fatalf("expected temp sensor cpu-thermal-cpu-temp, got: %v", components) + } + sd := sensor["sensor-data"].(map[string]interface{}) + if toInt(sd["value"]) != 42000 || sd["value-type"] != "celsius" || sd["value-scale"] != "milli" { + t.Fatalf("temp sensor-data mismatch: %v", sd) + } + if sensor["parent"] != "cpu" { + t.Fatalf("die sensor must hang off the cpu, got parent %v", sensor["parent"]) + } + if cpu := getComponentByName(components, "cpu"); cpu == nil || cpu["class"] != "iana-hardware:cpu" || cpu["parent"] != "mainboard" { + t.Fatalf("expected cpu component under mainboard, got %v", cpu) + } +} + +func TestHardwareHwmonFan(t *testing.T) { + runner := &testutil.MockRunner{Results: map[string][]byte{}, Errors: map[string]error{}} + fs := &testutil.MockFileReader{Files: map[string][]byte{ + "/run/system.json": []byte(`{}`), + "/sys/class/hwmon/hwmon1/name": []byte("pwmfan\n"), + "/sys/class/hwmon/hwmon1/fan1_input": []byte("3200\n"), + }, Globs: map[string][]string{ + "/sys/class/hwmon/*": {"/sys/class/hwmon/hwmon1"}, + "/sys/class/hwmon/hwmon1/*": {"/sys/class/hwmon/hwmon1/name", "/sys/class/hwmon/hwmon1/fan1_input"}, + }} + + components := collectHardware(t, newHardwareCollector(runner, fs)) + sensor := getComponentByName(components, "pwmfan") + if sensor == nil { + t.Fatalf("expected fan sensor pwmfan, got: %v", components) + } + sd := sensor["sensor-data"].(map[string]interface{}) + if toInt(sd["value"]) != 3200 || sd["value-type"] != "rpm" || sd["value-scale"] != "units" { + t.Fatalf("fan sensor-data mismatch: %v", sd) + } + if _, ok := sensor["parent"]; ok { + t.Fatalf("a lone fan has no parent, got %v", sensor["parent"]) + } + if getComponentByName(components, "cpu") != nil { + t.Fatal("no die sensor, so no cpu component") + } +} + +func TestHardwareHwmonVoltage(t *testing.T) { + runner := &testutil.MockRunner{Results: map[string][]byte{}, Errors: map[string]error{}} + fs := &testutil.MockFileReader{Files: map[string][]byte{ + "/run/system.json": []byte(`{}`), + "/sys/class/hwmon/hwmon2/name": []byte("ina3221\n"), + "/sys/class/hwmon/hwmon2/in1_input": []byte("12000\n"), + "/sys/class/hwmon/hwmon2/in1_label": []byte("VCC\n"), + }, Globs: map[string][]string{ + "/sys/class/hwmon/*": {"/sys/class/hwmon/hwmon2"}, + "/sys/class/hwmon/hwmon2/*": {"/sys/class/hwmon/hwmon2/name", "/sys/class/hwmon/hwmon2/in1_input", "/sys/class/hwmon/hwmon2/in1_label"}, + }} + + components := collectHardware(t, newHardwareCollector(runner, fs)) + sensor := getComponentByName(components, "ina3221-VCC") + if sensor == nil { + t.Fatalf("expected voltage sensor ina3221-VCC, got: %v", components) + } + if sensor["description"] != "VCC" { + t.Fatalf("expected VCC description, got %v", sensor["description"]) + } + sd := sensor["sensor-data"].(map[string]interface{}) + if toInt(sd["value"]) != 12000 || sd["value-type"] != "volts-DC" || sd["value-scale"] != "milli" { + t.Fatalf("voltage sensor-data mismatch: %v", sd) + } +} + +func TestHardwareHwmonMultiSensor(t *testing.T) { + runner := &testutil.MockRunner{Results: map[string][]byte{}, Errors: map[string]error{}} + fs := &testutil.MockFileReader{Files: map[string][]byte{ + "/run/system.json": []byte(`{}`), + "/sys/class/hwmon/hwmon3/name": []byte("sfp_2\n"), + "/sys/class/hwmon/hwmon3/temp1_input": []byte("33000\n"), + "/sys/class/hwmon/hwmon3/temp1_label": []byte("temp1\n"), + "/sys/class/hwmon/hwmon3/fan1_input": []byte("2000\n"), + "/sys/class/hwmon/hwmon3/fan1_label": []byte("fan1\n"), + "/sys/class/hwmon/hwmon3/curr1_input": []byte("1500\n"), + "/sys/class/hwmon/hwmon3/power1_input": []byte("2500000\n"), + }, Globs: map[string][]string{ + "/sys/class/hwmon/*": {"/sys/class/hwmon/hwmon3"}, + "/sys/class/hwmon/hwmon3/*": {"/sys/class/hwmon/hwmon3/name", "/sys/class/hwmon/hwmon3/temp1_input", "/sys/class/hwmon/hwmon3/temp1_label", "/sys/class/hwmon/hwmon3/fan1_input", "/sys/class/hwmon/hwmon3/fan1_label", "/sys/class/hwmon/hwmon3/curr1_input", "/sys/class/hwmon/hwmon3/power1_input"}, + }} + + components := collectHardware(t, newHardwareCollector(runner, fs)) + parent := getComponentByName(components, "sfp2") + if parent == nil || parent["class"] != "iana-hardware:module" { + t.Fatalf("expected sfp2 parent module, got: %v", parent) + } + + children := 0 + hasCurrent := false + hasPower := false + for _, compRaw := range components { + comp, ok := compRaw.(map[string]interface{}) + if !ok { + continue + } + if comp["parent"] != "sfp2" { + continue + } + children++ + sd, _ := comp["sensor-data"].(map[string]interface{}) + if sd != nil && sd["value-type"] == "amperes" { + hasCurrent = true + } + if sd != nil && sd["value-type"] == "watts" { + hasPower = true + } + } + if children < 4 { + t.Fatalf("expected at least 4 child sensors, got %d", children) + } + if !hasCurrent || !hasPower { + t.Fatalf("expected current and power sensors under parent: current=%v power=%v", hasCurrent, hasPower) + } +} + +func TestHardwareThermalZone(t *testing.T) { + runner := &testutil.MockRunner{Results: map[string][]byte{}, Errors: map[string]error{}} + fs := &testutil.MockFileReader{Files: map[string][]byte{ + "/run/system.json": []byte(`{}`), + "/sys/class/thermal/thermal_zone0/type": []byte("cpu-thermal\n"), + "/sys/class/thermal/thermal_zone0/temp": []byte("39000\n"), + }, Globs: map[string][]string{ + "/sys/class/thermal/*": {"/sys/class/thermal/thermal_zone0"}, + }} + + components := collectHardware(t, newHardwareCollector(runner, fs)) + sensor := getComponentByName(components, "cpu-thermal") + if sensor == nil { + t.Fatalf("expected thermal sensor cpu-thermal, got %v", components) + } + sd := sensor["sensor-data"].(map[string]interface{}) + if toInt(sd["value"]) != 39000 || sd["value-type"] != "celsius" { + t.Fatalf("thermal sensor mismatch: %v", sd) + } + if sensor["parent"] != "cpu" { + t.Fatalf("thermal die sensor must hang off the cpu, got %v", sensor["parent"]) + } +} + +func TestHardwareNormalizeSensorName(t *testing.T) { + tests := []struct { + in string + want string + }{ + {in: "sfp_2", want: "sfp2"}, + {in: "mt7915_phy0", want: "phy0"}, + {in: "marvell_alaska_tomte_phy7", want: "phy7"}, + {in: "cpu_thermal", want: "cpu-thermal"}, + {in: "gpu-thermal", want: "gpu-thermal"}, + {in: "s5_temp", want: "s5-temp"}, + {in: "pwmfan", want: "pwmfan"}, + } + + for _, tt := range tests { + tt := tt + t.Run(tt.in, func(t *testing.T) { + got := normalizeSensorName(tt.in) + if got != tt.want { + t.Errorf("normalizeSensorName(%q): expected %q, got %q", tt.in, tt.want, got) + } + }) + } +} + +func TestHardwareGracefulDegradation(t *testing.T) { + runner := &testutil.MockRunner{Results: map[string][]byte{}, Errors: map[string]error{}} + fs := &testutil.MockFileReader{Files: map[string][]byte{}, Globs: map[string][]string{}} + + tr := tree.New() + c := newHardwareCollector(runner, fs) + tr.RegisterProvider(hardwareKey, c.Live) + if err := c.Collect(context.Background(), tr); err != nil { + t.Fatalf("Collect should not fail when all probes fail: %v", err) + } + + raw := tr.Get("ietf-hardware:hardware") + if raw == nil { + t.Fatal("expected ietf-hardware:hardware key even on probe failures") + } + + var out map[string]interface{} + if err := json.Unmarshal(raw, &out); err != nil { + t.Fatalf("unmarshal hardware: %v", err) + } + components, ok := out["component"].([]interface{}) + if !ok { + t.Fatalf("component list missing: %v", out["component"]) + } + if len(components) != 0 { + t.Fatalf("expected empty component list on total failure, got %d (%v)", len(components), components) + } +} + +func TestHardwareGPSDeviceNotFound(t *testing.T) { + // Bug 5: When /dev/gps* doesn't exist, readlink -f still succeeds + // (returns canonical form of non-existent path). Verify the existence + // check prevents phantom GPS components. + runner := &testutil.MockRunner{Results: map[string][]byte{}, Errors: map[string]error{}} + fs := &testutil.MockFileReader{Files: map[string][]byte{ + "/run/system.json": []byte(`{}`), + }, Globs: map[string][]string{}} + + components := collectHardware(t, newHardwareCollector(runner, fs)) + for _, c := range components { + m, ok := c.(map[string]interface{}) + if !ok { + continue + } + if m["class"] == "infix-hardware:gps" { + t.Fatalf("phantom GPS component should not exist when /dev/gps* missing: %v", m) + } + } +} + +// The synthetic case from test/case/statd/sensors: an SoC die +// temperature from hwmon (s5_temp) and one from a thermal zone +// (cpu-thermal) both land under the cpu component, the fan does not, the +// hwmon mirror of the zone is skipped, and the VPD board named "cpu" +// does not collide with the cpu component. +func TestHardwareSyntheticSensors(t *testing.T) { + runner := &testutil.MockRunner{Results: map[string][]byte{}, Errors: map[string]error{}} + fs := &testutil.MockFileReader{Files: map[string][]byte{ + "/run/system.json": []byte(`{ + "vendor": "Microchip", + "product-name": "EV23X71A", + "mac-address": "00:a0:85:00:03:00", + "vpd": { + "cpu": { + "board": "cpu", + "available": true, + "trusted": true, + "data": { + "product-name": "CPU board", + "serial-number": "0123456789" + } + } + } + }`), + "/sys/class/hwmon/hwmon0/name": []byte("s5_temp\n"), + "/sys/class/hwmon/hwmon0/temp1_input": []byte("59500\n"), + "/sys/class/hwmon/hwmon1/name": []byte("pwmfan\n"), + "/sys/class/hwmon/hwmon1/fan1_input": []byte("3200\n"), + "/sys/class/hwmon/hwmon2/name": []byte("cpu_thermal\n"), + "/sys/class/hwmon/hwmon2/temp1_input": []byte("48000\n"), + "/sys/class/thermal/thermal_zone0/type": []byte("cpu-thermal\n"), + "/sys/class/thermal/thermal_zone0/temp": []byte("48000\n"), + }, Globs: map[string][]string{ + "/sys/class/hwmon/*": {"/sys/class/hwmon/hwmon0", "/sys/class/hwmon/hwmon1", "/sys/class/hwmon/hwmon2"}, + "/sys/class/hwmon/hwmon0/*": {"/sys/class/hwmon/hwmon0/name", "/sys/class/hwmon/hwmon0/temp1_input"}, + "/sys/class/hwmon/hwmon1/*": {"/sys/class/hwmon/hwmon1/name", "/sys/class/hwmon/hwmon1/fan1_input"}, + "/sys/class/hwmon/hwmon2/*": {"/sys/class/hwmon/hwmon2/name", "/sys/class/hwmon/hwmon2/temp1_input"}, + "/sys/class/thermal/*": {"/sys/class/thermal/thermal_zone0", "/sys/class/thermal/cooling_device0"}, + }} + + components := collectHardware(t, newHardwareCollector(runner, fs)) + + var names []string + for _, raw := range components { + names = append(names, raw.(map[string]interface{})["name"].(string)) + } + want := []string{"mainboard", "vpd-cpu", "cpu", "cpu-thermal", "s5-temp", "pwmfan"} + if fmt.Sprint(names) != fmt.Sprint(want) { + t.Fatalf("component order: want %v, got %v", want, names) + } + + if vpd := getComponentByName(components, "vpd-cpu"); vpd["class"] != "infix-hardware:vpd" || vpd["model-name"] != "CPU board" { + t.Fatalf("vpd-cpu mismatch: %v", vpd) + } + cpu := getComponentByName(components, "cpu") + if cpu["class"] != "iana-hardware:cpu" || cpu["parent"] != "mainboard" { + t.Fatalf("cpu mismatch: %v", cpu) + } + if st := cpu["state"].(map[string]interface{}); st["admin-state"] != "unknown" || st["oper-state"] != "enabled" { + t.Fatalf("cpu state mismatch: %v", st) + } + for name, value := range map[string]int{"cpu-thermal": 48000, "s5-temp": 59500} { + sensor := getComponentByName(components, name) + sd := sensor["sensor-data"].(map[string]interface{}) + if sensor["parent"] != "cpu" || toInt(sd["value"]) != value || sd["value-type"] != "celsius" { + t.Fatalf("%s mismatch: %v", name, sensor) + } + } + fan := getComponentByName(components, "pwmfan") + if _, ok := fan["parent"]; ok { + t.Fatalf("fan must not be under the cpu: %v", fan) + } + if sd := fan["sensor-data"].(map[string]interface{}); toInt(sd["value"]) != 3200 || sd["value-type"] != "rpm" { + t.Fatalf("fan mismatch: %v", fan) + } +} + +// A hwmon device named after a WiFi radio must not steal the radio's +// component name: its sensors hang off the radio, and the module head a +// multi-sensor device would get is dropped. +func TestHardwareAdoptWifiSensors(t *testing.T) { + components := []interface{}{ + map[string]interface{}{"name": "radio0", "class": "iana-hardware:module"}, + sensorComponent("radio0", 41000, "celsius", "milli", ""), + sensorComponent("radio0-VCC", 3300, "volts-DC", "milli", "VCC"), + sensorComponent("pwmfan", 3200, "rpm", "units", ""), + } + components[1].(map[string]interface{})["parent"] = "radio0" + components[2].(map[string]interface{})["parent"] = "radio0" + wifiInfo := map[string]map[string]interface{}{"radio0": {}} + + got := adoptWifiSensors(components, wifiInfo) + if len(got) != 3 { + t.Fatalf("module head must be dropped, got %v", got) + } + temp := getComponentByName(got, "radio0-temp") + if temp == nil || temp["parent"] != "radio0" || temp["description"] != "Temperature" { + t.Fatalf("radio temp sensor mismatch: %v", temp) + } + if getComponentByName(got, "radio0") != nil { + t.Fatal("no sensor may keep the radio's name") + } + if getComponentByName(got, "pwmfan") == nil { + t.Fatal("unrelated sensors must pass through") + } +} + +func TestHardwareUniqueNames(t *testing.T) { + components := []interface{}{ + map[string]interface{}{"name": "cpu"}, + map[string]interface{}{"name": "cpu"}, + map[string]interface{}{"name": "cpu-1"}, + map[string]interface{}{"name": "cpu"}, + } + + c := newHardwareCollector(&testutil.MockRunner{}, &testutil.MockFileReader{}) + var names []string + for _, raw := range c.uniqueNames(components) { + names = append(names, raw.(map[string]interface{})["name"].(string)) + } + want := []string{"cpu", "cpu-1", "cpu-1-1", "cpu-2"} + if fmt.Sprint(names) != fmt.Sprint(want) { + t.Fatalf("want %v, got %v", want, names) + } + if !c.reported["cpu"] || !c.reported["cpu-1"] || len(c.reported) != 2 { + t.Fatalf("duplicates remembered, got %v", c.reported) + } +} + +// A band whose name is unknown carries no name leaf rather than "Unknown". +func TestHardwareBandNameOptional(t *testing.T) { + out := convertPhyInfo(map[string]interface{}{ + "bands": []interface{}{ + map[string]interface{}{"band": 1, "name": "2.4 GHz"}, + map[string]interface{}{"band": 2}, + }, + }) + bands := out["bands"].([]interface{}) + if bands[0].(map[string]interface{})["name"] != "2.4 GHz" { + t.Fatalf("known band name lost: %v", bands[0]) + } + if _, ok := bands[1].(map[string]interface{})["name"]; ok { + t.Fatalf("unknown band must have no name: %v", bands[1]) + } +} + +// Sensor readings come from the GET, not the last poll: with Live +// registered as provider, a changed sysfs value shows without Collect. +func TestHardwareLiveReadsSensorsOnGet(t *testing.T) { + runner := &testutil.MockRunner{Results: map[string][]byte{}, Errors: map[string]error{}} + fs := &testutil.MockFileReader{Files: map[string][]byte{ + "/run/system.json": []byte(`{"vendor": "ACME", "product-name": "Box"}`), + "/sys/class/hwmon/hwmon0/name": []byte("s5_temp\n"), + "/sys/class/hwmon/hwmon0/temp1_input": []byte("40000\n"), + }, Globs: map[string][]string{ + "/sys/class/hwmon/*": {"/sys/class/hwmon/hwmon0"}, + "/sys/class/hwmon/hwmon0/*": {"/sys/class/hwmon/hwmon0/name", "/sys/class/hwmon/hwmon0/temp1_input"}, + "/sys/class/thermal/*": {}, + }} + + tr := tree.New() + c := newHardwareCollector(runner, fs) + tr.RegisterProvider("ietf-hardware:hardware", c.Live) + if err := c.Collect(context.Background(), tr); err != nil { + t.Fatalf("Collect failed: %v", err) + } + + reading := func() interface{} { + var out map[string]interface{} + if err := json.Unmarshal(tr.Get("ietf-hardware:hardware"), &out); err != nil { + t.Fatalf("unmarshal hardware: %v", err) + } + sensor := getComponentByName(out["component"].([]interface{}), "s5-temp") + if sensor == nil { + t.Fatalf("s5-temp missing: %v", out) + } + return sensor["sensor-data"].(map[string]interface{})["value"] + } + + if v := reading(); v != float64(40000) { + t.Fatalf("initial reading = %v, want 40000", v) + } + + fs.Files["/sys/class/hwmon/hwmon0/temp1_input"] = []byte("61000\n") + if v := reading(); v != float64(61000) { + t.Fatalf("reading after sysfs change = %v, want 61000 without a new poll", v) + } + + mainboard := func() map[string]interface{} { + var out map[string]interface{} + json.Unmarshal(tr.Get("ietf-hardware:hardware"), &out) + return getComponentByName(out["component"].([]interface{}), "mainboard") + } + if mainboard() == nil { + t.Fatal("polled inventory must be kept in the live view") + } +} + +// Only die sensors, and the "-thermal" zones the DT names them, belong +// to the CPU; a fan or a rail that merely starts with cpu/soc does not. +func TestSocTempSource(t *testing.T) { + yes := []string{"cpu", "cpu0", "soc", "core1", "coretemp", "k10temp", "s5-temp", "ap", "cp0", "cpu-thermal", "soc-thermal-1"} + no := []string{"cpu-fan", "soc-vdd", "ap-power", "cpufreq", "pwmfan", "sfp2"} + for _, name := range yes { + if !socTempSourceRe.MatchString(name) { + t.Errorf("%q must be a die sensor", name) + } + } + for _, name := range no { + if socTempSourceRe.MatchString(name) { + t.Errorf("%q must not be a die sensor", name) + } + } +} + +// The poll no longer builds the inventory: a GET reads system.json when +// asked, so a change shows without another poll. +func TestHardwareInventoryReadOnGet(t *testing.T) { + runner := &testutil.MockRunner{Results: map[string][]byte{}, Errors: map[string]error{}} + fs := &testutil.MockFileReader{Files: map[string][]byte{ + "/run/system.json": []byte(`{"vendor":"Acme","product-name":"Router-1"}`), + }, Globs: map[string][]string{}} + + tr := tree.New() + c := newHardwareCollector(runner, fs) + tr.RegisterProvider(hardwareKey, c.Live) + if err := c.Collect(context.Background(), tr); err != nil { + t.Fatalf("Collect failed: %v", err) + } + if cached := string(tr.GetCached(hardwareKey)); cached != "{}" { + t.Fatalf("poll stored %s, want the {} placeholder", cached) + } + + model := func() interface{} { + var out map[string]interface{} + if err := json.Unmarshal(tr.Get(hardwareKey), &out); err != nil { + t.Fatalf("unmarshal hardware: %v", err) + } + mb := getComponentByName(out["component"].([]interface{}), "mainboard") + if mb == nil { + t.Fatal("mainboard missing") + } + return mb["model-name"] + } + + if got := model(); got != "Router-1" { + t.Fatalf("model-name = %v, want Router-1", got) + } + fs.Files["/run/system.json"] = []byte(`{"vendor":"Acme","product-name":"Router-2"}`) + if got := model(); got != "Router-2" { + t.Fatalf("model-name = %v, want Router-2 without a new poll", got) + } +} + +// Without WiFi there are no radios to keep: RunRadios only waits for +// shutdown. +func TestHardwareRunRadiosIdleWithoutWifi(t *testing.T) { + runner := &testutil.MockRunner{Results: map[string][]byte{}, Errors: map[string]error{}} + fs := &testutil.MockFileReader{Files: map[string][]byte{}, Globs: map[string][]string{}} + c := newHardwareCollector(runner, fs) + + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan error, 1) + go func() { done <- c.RunRadios(ctx) }() + + c.RequestRadioRefresh() + c.RequestRadioRefresh() + select { + case err := <-done: + t.Fatalf("RunRadios returned early: %v", err) + case <-time.After(50 * time.Millisecond): + } + + cancel() + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("RunRadios did not stop on cancel") + } +} diff --git a/src/yangerd/internal/collector/live.go b/src/yangerd/internal/collector/live.go new file mode 100644 index 000000000..3e379e031 --- /dev/null +++ b/src/yangerd/internal/collector/live.go @@ -0,0 +1,209 @@ +package collector + +import ( + "encoding/json" + "strconv" + "strings" + "syscall" + "time" +) + +// LiveSystemState computes the on-demand portion of ietf-system:system-state. +// It reads uptime, current time, memory, load average from procfs and +// filesystem usage via statfs — all computed fresh on each call. +// +// Installer status is handled separately via MergeInstaller to avoid +// shallow-merge clobbering the boot-time software data. +func LiveSystemState(fs FileReader) json.RawMessage { + state := make(map[string]interface{}) + + if clock := liveClock(fs); len(clock) > 0 { + state["clock"] = clock + } + + resource := make(map[string]interface{}) + if mem := liveMemory(fs); len(mem) > 0 { + resource["memory"] = mem + } + if la := liveLoadAvg(fs); len(la) > 0 { + resource["load-average"] = la + } + if filesys := liveFilesystems(); len(filesys) > 0 { + resource["filesystem"] = filesys + } + if len(resource) > 0 { + state["infix-system:resource-usage"] = resource + } + + data, err := json.Marshal(state) + if err != nil { + return nil + } + return data +} + +// MergeInstaller reads the cached software data from the tree and +// overlays the live installer status into it, returning the merged +// infix-system:software object as a top-level system-state fragment. +func MergeInstaller(cached json.RawMessage, inst InstallerStatus) json.RawMessage { + if inst == nil { + return nil + } + installer := liveInstaller(inst) + if len(installer) == 0 { + return nil + } + + // On a decode error no overlay is better than one that replaces the + // software object with only the installer status, dropping the slots. + var base map[string]json.RawMessage + if len(cached) > 0 && json.Unmarshal(cached, &base) != nil { + return nil + } + + sw := make(map[string]interface{}) + if raw, ok := base["infix-system:software"]; ok && json.Unmarshal(raw, &sw) != nil { + return nil + } + if sw == nil { + sw = make(map[string]interface{}) + } + sw["installer"] = installer + + swJSON, err := json.Marshal(sw) + if err != nil { + return nil + } + + result := map[string]json.RawMessage{ + "infix-system:software": swJSON, + } + out, err := json.Marshal(result) + if err != nil { + return nil + } + return out +} + +func liveInstaller(inst InstallerStatus) map[string]interface{} { + op, lastErr, pct, msg, err := inst.GetInstallStatus() + if err != nil { + return nil + } + installer := make(map[string]interface{}) + if op != "" { + installer["operation"] = op + } + if lastErr != "" { + installer["last-error"] = lastErr + } + if pct > 0 || msg != "" { + progress := make(map[string]interface{}) + if pct > 0 { + progress["percentage"] = pct + } + if msg != "" { + progress["message"] = msg + } + installer["progress"] = progress + } + return installer +} + +func liveClock(fs FileReader) map[string]interface{} { + data, err := fs.ReadFile("/proc/uptime") + if err != nil { + return nil + } + parts := strings.Fields(string(data)) + if len(parts) < 1 { + return nil + } + upSec, err := strconv.ParseFloat(parts[0], 64) + if err != nil { + return nil + } + + now := time.Now() + boot := now.Add(-time.Duration(upSec * float64(time.Second))) + + return map[string]interface{}{ + "current-datetime": yangDateTime(now), + "boot-datetime": yangDateTime(boot), + } +} + +func liveMemory(fs FileReader) map[string]interface{} { + data, err := fs.ReadFile("/proc/meminfo") + if err != nil { + return nil + } + + memFields := map[string]string{ + "MemTotal": "total", + "MemFree": "free", + "MemAvailable": "available", + } + + memory := make(map[string]interface{}) + for _, line := range strings.Split(string(data), "\n") { + parts := strings.SplitN(line, ":", 2) + if len(parts) != 2 { + continue + } + key := strings.TrimSpace(parts[0]) + jsonKey, ok := memFields[key] + if !ok { + continue + } + valStr := strings.TrimSpace(parts[1]) + fields := strings.Fields(valStr) + if len(fields) < 1 { + continue + } + memory[jsonKey] = fields[0] + } + return memory +} + +func liveLoadAvg(fs FileReader) map[string]interface{} { + data, err := fs.ReadFile("/proc/loadavg") + if err != nil { + return nil + } + fields := strings.Fields(string(data)) + if len(fields) < 3 { + return nil + } + return map[string]interface{}{ + "load-1min": fields[0], + "load-5min": fields[1], + "load-15min": fields[2], + } +} + +func liveFilesystems() []interface{} { + // /run and /tmp are RAM-backed tmpfs, the scarcest writable storage + // on small boards. + mounts := []string{"/", "/var", "/cfg", "/run", "/tmp"} + var filesystems []interface{} + + for _, mount := range mounts { + var stat syscall.Statfs_t + if err := syscall.Statfs(mount, &stat); err != nil { + continue + } + bsize := uint64(stat.Bsize) + sizeKB := (stat.Blocks * bsize) / 1024 + availKB := (stat.Bavail * bsize) / 1024 + usedKB := sizeKB - (stat.Bfree*bsize)/1024 + + filesystems = append(filesystems, map[string]interface{}{ + "mount-point": mount, + "size": strconv.FormatUint(sizeKB, 10), + "used": strconv.FormatUint(usedKB, 10), + "available": strconv.FormatUint(availKB, 10), + }) + } + return filesystems +} diff --git a/src/yangerd/internal/collector/live_test.go b/src/yangerd/internal/collector/live_test.go new file mode 100644 index 000000000..5daa0b06f --- /dev/null +++ b/src/yangerd/internal/collector/live_test.go @@ -0,0 +1,264 @@ +package collector + +import ( + "encoding/json" + "testing" + "time" + + "github.com/kernelkit/infix/src/yangerd/internal/testutil" +) + +func TestLiveClock(t *testing.T) { + fs := &testutil.MockFileReader{ + Files: map[string][]byte{ + "/proc/uptime": []byte("12345.67 23456.78\n"), + }, + } + + before := time.Now().Truncate(time.Second) + clock := liveClock(fs) + after := time.Now().Truncate(time.Second).Add(time.Second) + + if clock == nil { + t.Fatal("expected non-nil clock") + } + + cur, ok := clock["current-datetime"].(string) + if !ok || cur == "" { + t.Fatal("missing current-datetime") + } + parsed, err := time.Parse("2006-01-02T15:04:05-07:00", cur) + if err != nil { + t.Fatalf("invalid datetime format: %v", err) + } + if parsed.Before(before) || parsed.After(after) { + t.Fatalf("current-datetime %v not between %v and %v", parsed, before, after) + } + + boot, ok := clock["boot-datetime"].(string) + if !ok || boot == "" { + t.Fatal("missing boot-datetime") + } + _, err = time.Parse("2006-01-02T15:04:05-07:00", boot) + if err != nil { + t.Fatalf("invalid boot-datetime format: %v", err) + } +} + +func TestLiveClockMissingFile(t *testing.T) { + fs := &testutil.MockFileReader{ + Files: map[string][]byte{}, + } + if clock := liveClock(fs); clock != nil { + t.Fatalf("expected nil on missing /proc/uptime, got %v", clock) + } +} + +func TestLiveMemory(t *testing.T) { + fs := &testutil.MockFileReader{ + Files: map[string][]byte{ + "/proc/meminfo": []byte("MemTotal: 1024000 kB\nMemFree: 512000 kB\nMemAvailable: 768000 kB\nBuffers: 64000 kB\n"), + }, + } + + mem := liveMemory(fs) + if mem == nil { + t.Fatal("expected non-nil memory") + } + + checks := map[string]string{ + "total": "1024000", + "free": "512000", + "available": "768000", + } + for key, expected := range checks { + got, ok := mem[key].(string) + if !ok || got != expected { + t.Fatalf("memory[%q]: expected %q, got %v", key, expected, mem[key]) + } + } + + if _, has := mem["Buffers"]; has { + t.Fatal("unexpected Buffers field in memory output") + } +} + +func TestLiveMemoryMissingFile(t *testing.T) { + fs := &testutil.MockFileReader{ + Files: map[string][]byte{}, + } + if mem := liveMemory(fs); mem != nil { + t.Fatalf("expected nil on missing /proc/meminfo, got %v", mem) + } +} + +func TestLiveLoadAvg(t *testing.T) { + fs := &testutil.MockFileReader{ + Files: map[string][]byte{ + "/proc/loadavg": []byte("0.42 0.31 0.15 2/123 4567\n"), + }, + } + + la := liveLoadAvg(fs) + if la == nil { + t.Fatal("expected non-nil load average") + } + + checks := map[string]string{ + "load-1min": "0.42", + "load-5min": "0.31", + "load-15min": "0.15", + } + for key, expected := range checks { + got, ok := la[key].(string) + if !ok || got != expected { + t.Fatalf("load-average[%q]: expected %q, got %v", key, expected, la[key]) + } + } +} + +func TestLiveLoadAvgMissingFile(t *testing.T) { + fs := &testutil.MockFileReader{ + Files: map[string][]byte{}, + } + if la := liveLoadAvg(fs); la != nil { + t.Fatalf("expected nil on missing /proc/loadavg, got %v", la) + } +} + +func TestLiveSystemState(t *testing.T) { + fs := &testutil.MockFileReader{ + Files: map[string][]byte{ + "/proc/uptime": []byte("100.0 200.0\n"), + "/proc/meminfo": []byte("MemTotal: 2048000 kB\nMemFree: 1024000 kB\nMemAvailable: 1536000 kB\n"), + "/proc/loadavg": []byte("1.00 0.50 0.25 3/200 9999\n"), + }, + } + + raw := LiveSystemState(fs) + if raw == nil { + t.Fatal("expected non-nil LiveSystemState output") + } + + var state map[string]interface{} + if err := json.Unmarshal(raw, &state); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + if _, ok := state["clock"]; !ok { + t.Fatal("missing clock in live state") + } + + resource, ok := state["infix-system:resource-usage"].(map[string]interface{}) + if !ok { + t.Fatal("missing infix-system:resource-usage in live state") + } + if _, ok := resource["memory"]; !ok { + t.Fatal("missing memory in resource-usage") + } + if _, ok := resource["load-average"]; !ok { + t.Fatal("missing load-average in resource-usage") + } +} + +func TestLiveSystemStatePartialFailure(t *testing.T) { + fs := &testutil.MockFileReader{ + Files: map[string][]byte{ + "/proc/loadavg": []byte("0.10 0.20 0.30 1/50 1234\n"), + }, + } + + raw := LiveSystemState(fs) + if raw == nil { + t.Fatal("expected non-nil even with partial data") + } + + var state map[string]interface{} + json.Unmarshal(raw, &state) + + if _, ok := state["clock"]; ok { + t.Fatal("clock should be absent when /proc/uptime is missing") + } + + resource := state["infix-system:resource-usage"].(map[string]interface{}) + if _, ok := resource["memory"]; ok { + t.Fatal("memory should be absent when /proc/meminfo is missing") + } + if _, ok := resource["load-average"]; !ok { + t.Fatal("load-average should be present") + } +} + +type mockInstaller struct { + op, lastErr, msg string + pct int + err error +} + +func (m *mockInstaller) GetInstallStatus() (string, string, int, string, error) { + return m.op, m.lastErr, m.pct, m.msg, m.err +} + +func TestMergeInstaller(t *testing.T) { + cached := json.RawMessage(`{"infix-system:software":{"compatible":"infix-x86_64","booted":{"slot":"rootfs.0"}}}`) + inst := &mockInstaller{op: "installing", pct: 45, msg: "Writing rootfs"} + + raw := MergeInstaller(cached, inst) + if raw == nil { + t.Fatal("expected non-nil") + } + + var result map[string]json.RawMessage + json.Unmarshal(raw, &result) + + var sw map[string]interface{} + json.Unmarshal(result["infix-system:software"], &sw) + if sw == nil { + t.Fatal("missing infix-system:software") + } + if sw["compatible"] != "infix-x86_64" { + t.Fatalf("cached 'compatible' was lost: %v", sw["compatible"]) + } + if sw["booted"] == nil { + t.Fatal("cached 'booted' was lost") + } + installer, ok := sw["installer"].(map[string]interface{}) + if !ok { + t.Fatal("missing installer") + } + if installer["operation"] != "installing" { + t.Fatalf("operation = %v, want 'installing'", installer["operation"]) + } + progress, ok := installer["progress"].(map[string]interface{}) + if !ok { + t.Fatal("missing progress") + } + if toInt(progress["percentage"]) != 45 { + t.Fatalf("percentage = %v, want 45", progress["percentage"]) + } + if progress["message"] != "Writing rootfs" { + t.Fatalf("message = %v", progress["message"]) + } +} + +func TestMergeInstallerNilCached(t *testing.T) { + inst := &mockInstaller{op: "idle"} + raw := MergeInstaller(nil, inst) + if raw == nil { + t.Fatal("expected non-nil even with nil cached") + } + var result map[string]json.RawMessage + json.Unmarshal(raw, &result) + var sw map[string]interface{} + json.Unmarshal(result["infix-system:software"], &sw) + if sw["installer"] == nil { + t.Fatal("missing installer") + } +} + +func TestMergeInstallerNilInst(t *testing.T) { + cached := json.RawMessage(`{"infix-system:software":{"compatible":"infix-x86_64"}}`) + if raw := MergeInstaller(cached, nil); raw != nil { + t.Fatalf("expected nil with nil installer, got %s", raw) + } +} diff --git a/src/yangerd/internal/collector/ntp.go b/src/yangerd/internal/collector/ntp.go new file mode 100644 index 000000000..75ea8dd23 --- /dev/null +++ b/src/yangerd/internal/collector/ntp.go @@ -0,0 +1,570 @@ +package collector + +import ( + "context" + "encoding/json" + "fmt" + "math" + "net" + "os" + "strconv" + "strings" + "sync" + "time" + + "github.com/facebook/time/ntp/chrony" + "github.com/kernelkit/infix/src/yangerd/internal/tree" +) + +// chronySock is chronyd's command (cmdmon) Unix socket. The UDP +// command port would suffice for most queries, but serverstats is +// PERMIT_AUTH in chronyd, which only the Unix socket satisfies. +const chronySock = "/run/chrony/chronyd.sock" + +// cmdmonClient is the subset of chrony.Client used by NTPCollector, +// broken out so tests can fake chronyd replies. +type cmdmonClient interface { + Communicate(packet chrony.RequestPacket) (chrony.ResponsePacket, error) +} + +// dialChrony connects to chronyd's cmdmon socket. SOCK_DGRAM over +// AF_UNIX has no connection state, so the client must bind its own +// socket for the replies, in a directory chronyd can write to -- the +// same dance chronyc does. +func dialChrony() (cmdmonClient, func() error, error) { + local := &net.UnixAddr{ + Name: fmt.Sprintf("/run/chrony/yangerd.%d.sock", os.Getpid()), + Net: "unixgram", + } + remote := &net.UnixAddr{Name: chronySock, Net: "unixgram"} + + os.Remove(local.Name) /* stale socket from a crashed run */ + conn, err := net.DialUnix("unixgram", local, remote) + if err != nil { + return nil, nil, err + } + + closeFn := func() error { + err := conn.Close() + os.Remove(local.Name) + return err + } + + // chronyd runs unprivileged and must be able to send replies here + if err := os.Chmod(local.Name, 0666); err != nil { + closeFn() + return nil, nil, err + } + if err := conn.SetDeadline(time.Now().Add(2 * time.Second)); err != nil { + closeFn() + return nil, nil, err + } + return &chrony.Client{Connection: conn}, closeFn, nil +} + +// ntpSource pairs a source's data and stats replies, fetched by the +// same cmdmon source index. stats may be nil if that request failed. +type ntpSource struct { + data *chrony.ReplySourceData + stats *chrony.ReplySourceStats +} + +// NTPCollector gathers ietf-ntp operational data from chronyd over the +// native cmdmon protocol (the same channel chronyc uses), plus ss to +// detect the NTP listening port. +type NTPCollector struct { + cmd CommandRunner + dial func() (cmdmonClient, func() error, error) + interval time.Duration + + // chronyd is asked from the poll and from GETs at once, and every + // dial binds the same client socket path. + mu sync.Mutex +} + +// NewNTPCollector creates an NTPCollector with the given dependencies. +func NewNTPCollector(cmd CommandRunner, interval time.Duration) *NTPCollector { + return &NTPCollector{cmd: cmd, dial: dialChrony, interval: interval} +} + +// Name implements Collector. +func (c *NTPCollector) Name() string { return "ntp" } + +// Interval implements Collector. +func (c *NTPCollector) Interval() time.Duration { return c.interval } + +// Collect implements Collector. It produces two tree keys: +// - "ietf-ntp:ntp" — associations, clock state, server status, and +// server statistics (RFC 9249). +// - "ietf-system:system-state" — merged infix-system:ntp/sources/source +// list with address, mode, state, stratum and poll for each chrony +// source (Infix augmentation of ietf-system). +func (c *NTPCollector) Collect(ctx context.Context, t *tree.Tree) error { + ntp, srcs := c.query(true) + + // Only probe the listening port when chronyd actually answered; + // otherwise a stale ss line would keep the tree key alive after + // chronyd stopped. + if len(ntp) > 0 { + c.addServerStatus(ctx, ntp) + } + + if len(ntp) > 0 { + if data, err := json.Marshal(ntp); err == nil { + t.Set("ietf-ntp:ntp", data) + } + } else { + // chronyd is not running (NTP unconfigured, or disabled by a + // config change). Drop the key so data from a previous run does + // not linger -- yangerd outlives config resets, so stale state + // would otherwise survive until restart. + t.Delete("ietf-ntp:ntp") + } + + if data := sourcesOverlay(srcs); data != nil { + t.Merge("ietf-system:system-state", data) + } + return nil +} + +// query asks chronyd for its sources and, with full, for the rest of +// ietf-ntp:ntp. An empty ntp means chronyd did not answer. +func (c *NTPCollector) query(full bool) (map[string]interface{}, []ntpSource) { + c.mu.Lock() + defer c.mu.Unlock() + + ntp := make(map[string]interface{}) + client, closeConn, err := c.dial() + if err != nil { + return ntp, nil + } + defer closeConn() + + srcs := getSources(client) + if full { + addAssociations(ntp, srcs) + addClockState(client, ntp) + addServerStats(client, ntp) + } + return ntp, srcs +} + +// Live is the tree provider for ietf-ntp:ntp. Source selection and the +// clock state move without any event from chronyd, and asking it is a +// local socket round trip, so a GET reads them now. The poll keeps the +// key present while chronyd runs, and the listening port, which takes +// a fork. +func (c *NTPCollector) Live() json.RawMessage { + ntp, _ := c.query(true) + if len(ntp) == 0 { + return nil + } + data, err := json.Marshal(ntp) + if err != nil { + return nil + } + return data +} + +// LiveSources is the infix-system:ntp overlay for ietf-system:system-state, +// read at GET time for the same reason as Live. +func (c *NTPCollector) LiveSources() json.RawMessage { + _, srcs := c.query(false) + return sourcesOverlay(srcs) +} + +// sourcesOverlay always carries a source list, empty when chronyd has +// none. Otherwise a source that disappears from chrony (e.g. a DHCP +// lease without option 42, or NTP turned off) lingers as stale +// operational data -- a phantom "selected" server chronyc no longer +// reports. Merge only overwrites the keys it is given, so it must be +// handed an empty source list to clear a previously-populated one. +func sourcesOverlay(srcs []ntpSource) json.RawMessage { + sources := addSources(srcs) + if sources == nil { + sources = map[string]interface{}{ + "sources": map[string]interface{}{ + "source": []interface{}{}, + }, + } + } + data, err := json.Marshal(map[string]interface{}{ + "infix-system:ntp": sources, + }) + if err != nil { + return nil + } + return data +} + +// getSources fetches source data and stats for every chrony source. +// Data and stats share the same cmdmon index space, so no address +// matching is needed. +func getSources(client cmdmonClient) []ntpSource { + resp, err := client.Communicate(chrony.NewSourcesPacket()) + if err != nil { + return nil + } + sources, ok := resp.(*chrony.ReplySources) + if !ok { + return nil + } + + srcs := make([]ntpSource, 0, sources.NSources) + for i := 0; i < sources.NSources; i++ { + resp, err := client.Communicate(chrony.NewSourceDataPacket(int32(i))) + if err != nil { + continue + } + data, ok := resp.(*chrony.ReplySourceData) + if !ok { + continue + } + + src := ntpSource{data: data} + if resp, err := client.Communicate(chrony.NewSourceStatsPacket(int32(i))); err == nil { + if stats, ok := resp.(*chrony.ReplySourceStats); ok { + src.stats = stats + } + } + srcs = append(srcs, src) + } + + return srcs +} + +// addAssociations builds the associations/association list from chrony +// source data and stats. +func addAssociations(ntp map[string]interface{}, srcs []ntpSource) { + modeMap := map[chrony.ModeType]string{ + chrony.SourceModeClient: "ietf-ntp:client", + chrony.SourceModePeer: "ietf-ntp:active", + } + + var associations []interface{} + for _, src := range srcs { + d := src.data + + // Skip reference clocks — they have refids, not addresses + if d.Mode == chrony.SourceModeRef { + continue + } + + // YANG requires stratum 1..16 + stratum := int(d.Stratum) + if stratum < 1 || stratum > 16 { + continue + } + + mode := modeMap[d.Mode] + if mode == "" { + mode = "ietf-ntp:client" + } + + assoc := map[string]interface{}{ + "address": d.IPAddr.String(), + "local-mode": mode, + "isconfigured": true, + "stratum": stratum, + "reach": int(d.Reachability), + "poll": int(d.Poll), + "now": int(d.SinceSample), + } + + // Current sync source + if d.State == chrony.SourceStateSync { + assoc["prefer"] = true + } + + // Offset: prefer sourcestats estimate over the last sample. + // Convert seconds → milliseconds with 3 fraction digits + if src.stats != nil { + assoc["offset"] = fmt.Sprintf("%.3f", src.stats.EstimatedOffset*1000.0) + assoc["dispersion"] = fmt.Sprintf("%.3f", src.stats.StandardDeviation*1000.0) + } else { + assoc["offset"] = fmt.Sprintf("%.3f", d.LatestMeas*1000.0) + } + + // Delay: error estimate of the last sample, seconds → milliseconds + assoc["delay"] = fmt.Sprintf("%.3f", math.Abs(d.LatestMeasErr)*1000.0) + + associations = append(associations, assoc) + } + + if len(associations) > 0 { + ntp["associations"] = map[string]interface{}{ + "association": associations, + } + } +} + +// sourceStateMap maps chrony source states to YANG infix-system +// source-state enum values. +var sourceStateMap = map[chrony.SourceStateType]string{ + chrony.SourceStateSync: "selected", + // The library names these after chrony 3. In chrony 4, 4 is + // "unselected" ('-') and 5 is "selectable" ('+'), the other way + // round from what the names say. + chrony.SourceStateType(5): "candidate", + chrony.SourceStateType(4): "outlier", + chrony.SourceStateUnreach: "unusable", + chrony.SourceStateFalseTicker: "falseticker", + chrony.SourceStateJittery: "unstable", +} + +// sourceModeMap maps chrony source modes to YANG infix-system +// source-mode enum values. +var sourceModeMap = map[chrony.ModeType]string{ + chrony.SourceModeClient: "server", + chrony.SourceModePeer: "peer", + chrony.SourceModeRef: "local-clock", +} + +// addSources builds the infix-system:ntp/sources/source list. +// Reference clocks and sources with invalid stratum are skipped, +// matching the Python yanger ietf_system.py add_ntp() behaviour. +func addSources(srcs []ntpSource) map[string]interface{} { + var sources []interface{} + for _, src := range srcs { + d := src.data + + if d.Mode == chrony.SourceModeRef { + continue + } + if d.Stratum > 16 { + continue + } + + mode := sourceModeMap[d.Mode] + if mode == "" { + mode = "server" + } + state := sourceStateMap[d.State] + if state == "" { + continue + } + + sources = append(sources, map[string]interface{}{ + "address": d.IPAddr.String(), + "mode": mode, + "state": state, + "stratum": int(d.Stratum), + "poll": int(d.Poll), + }) + } + + if len(sources) == 0 { + return nil + } + + return map[string]interface{}{ + "sources": map[string]interface{}{ + "source": sources, + }, + } +} + +// chrony LeapStatus from tracking: 0 normal, 1 insert, 2 delete, +// 3 not synchronised. +const leapUnsynchronised = 3 + +// clockRefid renders the tracking reference ID in a form the RFC 9249 +// refid union accepts: an IPv4 address, a uint32, or exactly four +// characters. The uint32 member must be a JSON number -- libyang +// rejects number-typed union members encoded as strings. +func clockRefid(t *chrony.Tracking) interface{} { + if t.IPAddr != nil && !t.IPAddr.IsUnspecified() { + if ip4 := t.IPAddr.To4(); ip4 != nil { + return ip4.String() + } + // IPv6 sources have no representable address: chrony + // stores a hash of it in the refid + return t.RefID + } + if t.RefID != 0 { + // Reference clock, e.g. "GPS": RFC 5905 refids are four + // bytes, space-padded + if s := refidToASCII(t.RefID); s != "" { + return (s + " ")[:4] + } + // Non-printable refid, e.g. chronyd's local reference + // 0x7F7F0101: render as the pseudo-IP it encodes + refid := t.RefID + return fmt.Sprintf("%d.%d.%d.%d", + refid>>24, refid>>16&0xff, refid>>8&0xff, refid&0xff) + } + return "0.0.0.0" +} + +// refidToASCII decodes a printable refid name like "GPS", or returns "" +// when any byte is non-printable (a hash or pseudo-IP, not a name). +func refidToASCII(refid uint32) string { + var s []byte + + for i := 3; i >= 0; i-- { + c := byte(refid >> (8 * i)) + if c == 0 { + continue + } + if c < ' ' || c > '~' { + return "" + } + s = append(s, c) + } + + return string(s) +} + +// addClockState fills the clock-state container from chrony tracking. +func addClockState(client cmdmonClient, ntp map[string]interface{}) { + resp, err := client.Communicate(chrony.NewTrackingPacket()) + if err != nil { + return + } + tracking, ok := resp.(*chrony.ReplyTracking) + if !ok { + return + } + + ss := make(map[string]interface{}) + + // Stratum: chronyd uses 0 for "not synchronized", YANG requires 1-16 + stratum := int(tracking.Stratum) + if stratum == 0 { + stratum = 16 + } + + if stratum == 16 { + ss["clock-state"] = "ietf-ntp:unsynchronized" + } else { + ss["clock-state"] = "ietf-ntp:synchronized" + } + ss["clock-stratum"] = stratum + + ss["clock-refid"] = clockRefid(&tracking.Tracking) + + // Frequencies (ppm → Hz with nominal 1GHz) + nominal := 1000000000.0 + actual := nominal * (1.0 + tracking.FreqPPM/1000000.0) + ss["nominal-freq"] = fmt.Sprintf("%.4f", nominal) + ss["actual-freq"] = fmt.Sprintf("%.4f", actual) + + // Clock precision (fixed estimate, ~1µs) + ss["clock-precision"] = -20 + + // Clock offset (seconds → milliseconds) + ss["clock-offset"] = fmt.Sprintf("%.3f", tracking.CurrentCorrection*1000.0) + + // Root delay and dispersion (seconds → milliseconds) + ss["root-delay"] = fmt.Sprintf("%.3f", tracking.RootDelay*1000.0) + ss["root-dispersion"] = fmt.Sprintf("%.3f", tracking.RootDispersion*1000.0) + + // Reference time (ISO 8601) + if !tracking.RefTime.IsZero() && tracking.RefTime.Unix() > 0 { + ss["reference-time"] = tracking.RefTime.UTC().Format("2006-01-02T15:04:05.000") + "Z" + } + + // Sync state based on leap status + if tracking.LeapStatus == leapUnsynchronised || stratum == 16 { + ss["sync-state"] = "ietf-ntp:clock-never-set" + } else { + ss["sync-state"] = "ietf-ntp:clock-synchronized" + } + + // Infix augmentations + ss["infix-ntp:last-offset"] = fmt.Sprintf("%.9f", tracking.LastOffset) + ss["infix-ntp:rms-offset"] = fmt.Sprintf("%.9f", tracking.RMSOffset) + ss["infix-ntp:residual-freq"] = fmt.Sprintf("%.3f", tracking.ResidFreqPPM) + ss["infix-ntp:skew"] = fmt.Sprintf("%.3f", tracking.SkewPPM) + ss["infix-ntp:update-interval"] = fmt.Sprintf("%.1f", tracking.LastUpdateInterval) + + ntp["clock-state"] = map[string]interface{}{ + "system-status": ss, + } +} + +// addServerStatus adds the refclock-master stratum and listening port. +// Must be called after addClockState so clock-state is available. +func (c *NTPCollector) addServerStatus(ctx context.Context, ntp map[string]interface{}) { + // Reuse stratum from clock-state if already populated + if cs, ok := ntp["clock-state"].(map[string]interface{}); ok { + if ss, ok := cs["system-status"].(map[string]interface{}); ok { + if stratum, ok := ss["clock-stratum"]; ok { + ntp["refclock-master"] = map[string]interface{}{ + "master-stratum": stratum, + } + } + } + } + + // Detect NTP listening port via ss + ssOut, err := c.cmd.Run(ctx, "ss", "-ulnp") + if err != nil { + return + } + + for _, line := range splitLines(string(ssOut)) { + if !strings.Contains(line, "chronyd") { + continue + } + // Skip loopback (command socket) + if strings.Contains(line, "127.0.0.1") || strings.Contains(line, "[::1]") { + continue + } + + fields := strings.Fields(line) + if len(fields) >= 5 { + localAddr := fields[3] + idx := strings.LastIndex(localAddr, ":") + if idx >= 0 { + portStr := localAddr[idx+1:] + if port, err := strconv.Atoi(portStr); err == nil { + ntp["port"] = port + break + } + } + } + } +} + +// addServerStats fills ntp-statistics from chrony server stats. The +// reply version depends on the chronyd version, but all carry the NTP +// packets received/dropped counters. chronyd does not count sent +// packets, so packet-sent/packet-sent-fail are not reported. +func addServerStats(client cmdmonClient, ntp map[string]interface{}) { + resp, err := client.Communicate(chrony.NewServerStatsPacket()) + if err != nil { + return + } + + var received, dropped uint64 + switch r := resp.(type) { + case *chrony.ReplyServerStats: + received, dropped = uint64(r.NTPHits), uint64(r.NTPDrops) + case *chrony.ReplyServerStats2: + received, dropped = uint64(r.NTPHits), uint64(r.NTPDrops) + case *chrony.ReplyServerStats3: + received, dropped = uint64(r.NTPHits), uint64(r.NTPDrops) + case *chrony.ReplyServerStats4: + received, dropped = r.NTPHits, r.NTPDrops + default: + return + } + + ntp["ntp-statistics"] = map[string]interface{}{ + "packet-received": received, + "packet-dropped": dropped, + } +} + +// splitLines splits text into non-empty lines. +func splitLines(text string) []string { + var lines []string + for _, line := range strings.Split(text, "\n") { + line = strings.TrimSpace(line) + if line != "" { + lines = append(lines, line) + } + } + return lines +} diff --git a/src/yangerd/internal/collector/ntp_test.go b/src/yangerd/internal/collector/ntp_test.go new file mode 100644 index 000000000..97e0c2d58 --- /dev/null +++ b/src/yangerd/internal/collector/ntp_test.go @@ -0,0 +1,704 @@ +package collector + +import ( + "context" + "encoding/json" + "errors" + "net" + "strings" + "testing" + "time" + + "github.com/facebook/time/ntp/chrony" + "github.com/kernelkit/infix/src/yangerd/internal/testutil" + "github.com/kernelkit/infix/src/yangerd/internal/tree" +) + +const testSSOutput = `State Recv-Q Send-Q Local Address:Port Peer Address:Port Process +UNCONN 0 0 0.0.0.0:123 0.0.0.0:* users:(("chronyd",pid=5441,fd=5)) +UNCONN 0 0 127.0.0.1:323 0.0.0.0:* users:(("chronyd",pid=5441,fd=1)) +` + +// fakeChrony answers cmdmon requests from canned replies. A nil reply +// (or err set) makes the corresponding request fail, mimicking a dead +// or partially-responding chronyd. +type fakeChrony struct { + err error + data []*chrony.ReplySourceData + stats []*chrony.ReplySourceStats + tracking *chrony.ReplyTracking + serverStats chrony.ResponsePacket +} + +func (f *fakeChrony) Communicate(packet chrony.RequestPacket) (chrony.ResponsePacket, error) { + if f.err != nil { + return nil, f.err + } + switch p := packet.(type) { + case *chrony.RequestSources: + return &chrony.ReplySources{NSources: len(f.data)}, nil + case *chrony.RequestSourceData: + if int(p.Index) >= len(f.data) || f.data[p.Index] == nil { + return nil, errors.New("no such source") + } + return f.data[p.Index], nil + case *chrony.RequestSourceStats: + if int(p.Index) >= len(f.stats) || f.stats[p.Index] == nil { + return nil, errors.New("no stats") + } + return f.stats[p.Index], nil + case *chrony.RequestTracking: + if f.tracking == nil { + return nil, errors.New("no tracking") + } + return f.tracking, nil + case *chrony.RequestServerStats: + if f.serverStats == nil { + return nil, errors.New("no serverstats") + } + return f.serverStats, nil + } + return nil, errors.New("unexpected request") +} + +func v4(a, b, c, d uint8) net.IP { + return net.IPv4(a, b, c, d) +} + +func sourceData(ip net.IP, mode chrony.ModeType, state chrony.SourceStateType, + stratum uint16, poll int16, reach uint16, since uint32, + latestMeas, latestMeasErr float64) *chrony.ReplySourceData { + return &chrony.ReplySourceData{ + SourceData: chrony.SourceData{ + IPAddr: ip, + Poll: poll, + Stratum: stratum, + State: state, + Mode: mode, + Reachability: reach, + SinceSample: since, + LatestMeas: latestMeas, + LatestMeasErr: latestMeasErr, + }, + } +} + +func sourceStats(offset, stddev float64) *chrony.ReplySourceStats { + return &chrony.ReplySourceStats{ + SourceStats: chrony.SourceStats{ + EstimatedOffset: offset, + StandardDeviation: stddev, + }, + } +} + +// fullFakeChrony mirrors the source mix of the old chronyc fixtures: +// selected/candidate servers, an outlier peer, a GPS refclock, and an +// unreachable stratum-0 server (no stats). +func fullFakeChrony() *fakeChrony { + return &fakeChrony{ + data: []*chrony.ReplySourceData{ + sourceData(v4(10, 0, 0, 1), chrony.SourceModeClient, chrony.SourceStateSync, + 2, 6, 0o377, 32, 0.000123, 0.000456), + sourceData(v4(10, 0, 0, 2), chrony.SourceModeClient, chronySelectable, + 3, 7, 0o377, 64, -0.000789, 0.001234), + sourceData(v4(10, 0, 0, 3), chrony.SourceModePeer, chronyUnselected, + 4, 6, 0o177, 128, 0.001500, 0.002000), + sourceData(nil, chrony.SourceModeRef, chrony.SourceStateSync, + 1, 4, 0o377, 16, 0.000001, 0.000010), + sourceData(v4(10, 0, 0, 4), chrony.SourceModeClient, chrony.SourceStateUnreach, + 0, 6, 0, 0, 0, 0), + }, + stats: []*chrony.ReplySourceStats{ + sourceStats(0.000050, 0.000100), + sourceStats(-0.000300, 0.000200), + sourceStats(0.001000, 0.000500), + nil, + nil, + }, + tracking: &chrony.ReplyTracking{ + Tracking: chrony.Tracking{ + RefID: 0xC0A80001, + IPAddr: []byte{192, 168, 0, 1}, + Stratum: 2, + LeapStatus: 0, + RefTime: time.Unix(1700000000, 123000000), + CurrentCorrection: 0.000045, + LastOffset: -0.000012, + RMSOffset: 0.000025, + FreqPPM: -1.5, + ResidFreqPPM: 0.003, + SkewPPM: 0.050, + RootDelay: 0.004500, + RootDispersion: 0.001200, + LastUpdateInterval: 64.0, + }, + }, + serverStats: &chrony.ReplyServerStats4{ + ServerStats4: chrony.ServerStats4{ + NTPHits: 1000, + NTPDrops: 5, + }, + }, + } +} + +func unsyncFakeChrony() *fakeChrony { + return &fakeChrony{ + tracking: &chrony.ReplyTracking{ + Tracking: chrony.Tracking{ + LeapStatus: 3, // not synchronised + }, + }, + } +} + +func newNTPCollector(runner *testutil.MockRunner, fake *fakeChrony) *NTPCollector { + c := NewNTPCollector(runner, 60*time.Second) + if fake != nil { + c.dial = func() (cmdmonClient, func() error, error) { + return fake, func() error { return nil }, nil + } + } else { + c.dial = func() (cmdmonClient, func() error, error) { + return nil, nil, errors.New("connection refused") + } + } + return c +} + +func ssRunner() *testutil.MockRunner { + return &testutil.MockRunner{ + Results: map[string][]byte{ + "ss -ulnp": []byte(testSSOutput), + }, + Errors: map[string]error{}, + } +} + +func ntpCollect(t *testing.T, fake *fakeChrony) (map[string]interface{}, *tree.Tree) { + t.Helper() + c := newNTPCollector(ssRunner(), fake) + tr := tree.New() + if err := c.Collect(context.Background(), tr); err != nil { + t.Fatalf("Collect failed: %v", err) + } + raw := tr.Get("ietf-ntp:ntp") + if raw == nil { + t.Fatal("missing ietf-ntp:ntp in tree") + } + var out map[string]interface{} + if err := json.Unmarshal(raw, &out); err != nil { + t.Fatalf("unmarshal ntp: %v", err) + } + return out, tr +} + +func TestNTPCollectorNameAndInterval(t *testing.T) { + c := newNTPCollector(ssRunner(), fullFakeChrony()) + if c.Name() != "ntp" { + t.Fatalf("expected name 'ntp', got %q", c.Name()) + } + if c.Interval() != 60*time.Second { + t.Fatalf("expected interval 60s, got %v", c.Interval()) + } +} + +func TestNTPAssociations(t *testing.T) { + out, _ := ntpCollect(t, fullFakeChrony()) + assocContainer := out["associations"].(map[string]interface{}) + assocs := assocContainer["association"].([]interface{}) + + // 5 sources minus GPS refclock minus stratum-0 (10.0.0.4) = 3 + if len(assocs) != 3 { + t.Fatalf("expected 3 associations (refclock+stratum0 filtered), got %d", len(assocs)) + } + + byAddr := make(map[string]map[string]interface{}) + for _, a := range assocs { + am := a.(map[string]interface{}) + byAddr[am["address"].(string)] = am + } + + // 10.0.0.1: selected server, stratum 2 + a1 := byAddr["10.0.0.1"] + if a1 == nil { + t.Fatal("missing association for 10.0.0.1") + } + if a1["local-mode"] != "ietf-ntp:client" { + t.Fatalf("10.0.0.1 mode: expected ietf-ntp:client, got %v", a1["local-mode"]) + } + if a1["prefer"] != true { + t.Fatalf("10.0.0.1 should be preferred (selected source)") + } + if toInt(a1["stratum"]) != 2 { + t.Fatalf("10.0.0.1 stratum: expected 2, got %v", a1["stratum"]) + } + // Reach: 377 octal = 255 decimal + if toInt(a1["reach"]) != 255 { + t.Fatalf("10.0.0.1 reach: expected 255, got %v", a1["reach"]) + } + // Offset should come from sourcestats (0.000050s → 0.050ms) + if a1["offset"] != "0.050" { + t.Fatalf("10.0.0.1 offset: expected '0.050', got %v", a1["offset"]) + } + // Dispersion from sourcestats std_dev (0.000100s → 0.100ms) + if a1["dispersion"] != "0.100" { + t.Fatalf("10.0.0.1 dispersion: expected '0.100', got %v", a1["dispersion"]) + } + // Delay from last sample error estimate (0.000456s → 0.456ms) + if a1["delay"] != "0.456" { + t.Fatalf("10.0.0.1 delay: expected '0.456', got %v", a1["delay"]) + } + + // 10.0.0.3: peer mode + a3 := byAddr["10.0.0.3"] + if a3 == nil { + t.Fatal("missing association for 10.0.0.3") + } + if a3["local-mode"] != "ietf-ntp:active" { + t.Fatalf("10.0.0.3 mode: expected ietf-ntp:active, got %v", a3["local-mode"]) + } + // Should NOT be preferred (outlier) + if _, hasPrefer := a3["prefer"]; hasPrefer { + t.Fatal("10.0.0.3 should not be preferred") + } +} + +func TestNTPSources(t *testing.T) { + _, tr := ntpCollect(t, fullFakeChrony()) + + raw := tr.Get("ietf-system:system-state") + if raw == nil { + t.Fatal("missing ietf-system:system-state in tree") + } + var state map[string]interface{} + if err := json.Unmarshal(raw, &state); err != nil { + t.Fatalf("unmarshal system-state: %v", err) + } + + ntpData, ok := state["infix-system:ntp"].(map[string]interface{}) + if !ok { + t.Fatal("missing infix-system:ntp in system-state") + } + sourcesContainer, ok := ntpData["sources"].(map[string]interface{}) + if !ok { + t.Fatal("missing sources in infix-system:ntp") + } + sources, ok := sourcesContainer["source"].([]interface{}) + if !ok { + t.Fatal("missing source list in sources") + } + + // 5 sources minus GPS refclock = 4 (stratum 0 is kept) + if len(sources) != 4 { + t.Fatalf("expected 4 sources, got %d", len(sources)) + } + + byAddr := make(map[string]map[string]interface{}) + for _, s := range sources { + sm := s.(map[string]interface{}) + byAddr[sm["address"].(string)] = sm + } + + // 10.0.0.1: selected server + s1 := byAddr["10.0.0.1"] + if s1 == nil { + t.Fatal("missing source 10.0.0.1") + } + if s1["state"] != "selected" { + t.Fatalf("10.0.0.1 state: expected selected, got %v", s1["state"]) + } + if s1["mode"] != "server" { + t.Fatalf("10.0.0.1 mode: expected server, got %v", s1["mode"]) + } + if toInt(s1["stratum"]) != 2 { + t.Fatalf("10.0.0.1 stratum: expected 2, got %v", s1["stratum"]) + } + if toInt(s1["poll"]) != 6 { + t.Fatalf("10.0.0.1 poll: expected 6, got %v", s1["poll"]) + } + + // 10.0.0.2: candidate server + s2 := byAddr["10.0.0.2"] + if s2 == nil { + t.Fatal("missing source 10.0.0.2") + } + if s2["state"] != "candidate" { + t.Fatalf("10.0.0.2 state: expected candidate, got %v", s2["state"]) + } + if s2["mode"] != "server" { + t.Fatalf("10.0.0.2 mode: expected server, got %v", s2["mode"]) + } + + // 10.0.0.3: outlier peer + s3 := byAddr["10.0.0.3"] + if s3 == nil { + t.Fatal("missing source 10.0.0.3") + } + if s3["state"] != "outlier" { + t.Fatalf("10.0.0.3 state: expected outlier, got %v", s3["state"]) + } + if s3["mode"] != "peer" { + t.Fatalf("10.0.0.3 mode: expected peer, got %v", s3["mode"]) + } + + // 10.0.0.4: unreachable server (stratum 0) + s4 := byAddr["10.0.0.4"] + if s4 == nil { + t.Fatal("missing source 10.0.0.4") + } + if s4["state"] != "unusable" { + t.Fatalf("10.0.0.4 state: expected unusable, got %v", s4["state"]) + } + if toInt(s4["stratum"]) != 0 { + t.Fatalf("10.0.0.4 stratum: expected 0, got %v", s4["stratum"]) + } +} + +// ntpSourceCount returns the number of infix-system:ntp sources in the +// system-state tree key, failing the test if the subtree is missing. +func ntpSourceCount(t *testing.T, tr *tree.Tree) int { + t.Helper() + + raw := tr.Get("ietf-system:system-state") + if raw == nil { + t.Fatal("system-state not set") + } + + var data map[string]json.RawMessage + if err := json.Unmarshal(raw, &data); err != nil { + t.Fatalf("unmarshal system-state: %v", err) + } + ntpRaw, ok := data["infix-system:ntp"] + if !ok { + t.Fatal("infix-system:ntp not present") + } + + var ntp struct { + Sources struct { + Source []json.RawMessage `json:"source"` + } `json:"sources"` + } + if err := json.Unmarshal(ntpRaw, &ntp); err != nil { + t.Fatalf("unmarshal infix-system:ntp: %v", err) + } + return len(ntp.Sources.Source) +} + +func TestNTPSourcesEmpty(t *testing.T) { + fake := &fakeChrony{tracking: fullFakeChrony().tracking} + + c := newNTPCollector(ssRunner(), fake) + tr := tree.New() + c.Collect(context.Background(), tr) + + // With no chrony sources the collector must still write an empty + // source list, so a previously-reported source cannot linger as + // stale operational data. + if n := ntpSourceCount(t, tr); n != 0 { + t.Fatalf("expected empty NTP source list, got %d", n) + } +} + +// When chronyd stops (NTP disabled via config reset), the whole +// ietf-ntp:ntp key must disappear -- yangerd outlives config resets, so +// a key that is only ever Set when non-empty would keep stale data from +// a previous run forever. +func TestNTPTreeKeyRemovedWhenChronydStops(t *testing.T) { + tr := tree.New() + + newNTPCollector(ssRunner(), fullFakeChrony()).Collect(context.Background(), tr) + if tr.Get("ietf-ntp:ntp") == nil { + t.Fatal("expected ietf-ntp:ntp after first poll") + } + + stopped := &fakeChrony{err: errors.New("read timeout")} + newNTPCollector(ssRunner(), stopped).Collect(context.Background(), tr) + if data := tr.Get("ietf-ntp:ntp"); data != nil { + t.Fatalf("stale ietf-ntp:ntp survived chronyd stop: %s", data) + } +} + +// A source that disappears from chrony (e.g. a DHCP NTP server that is no +// longer offered) must be cleared from operational, not left stale. +func TestNTPSourcesClearedWhenGone(t *testing.T) { + tr := tree.New() + + newNTPCollector(ssRunner(), fullFakeChrony()).Collect(context.Background(), tr) + if ntpSourceCount(t, tr) == 0 { + t.Fatal("expected NTP sources after first poll") + } + + noSources := &fakeChrony{tracking: fullFakeChrony().tracking} + newNTPCollector(ssRunner(), noSources).Collect(context.Background(), tr) + if n := ntpSourceCount(t, tr); n != 0 { + t.Fatalf("stale NTP sources not cleared: got %d, want 0", n) + } +} + +func TestNTPClockStateSynchronized(t *testing.T) { + out, _ := ntpCollect(t, fullFakeChrony()) + cs := out["clock-state"].(map[string]interface{}) + ss := cs["system-status"].(map[string]interface{}) + + if ss["clock-state"] != "ietf-ntp:synchronized" { + t.Fatalf("clock-state: expected synchronized, got %v", ss["clock-state"]) + } + if toInt(ss["clock-stratum"]) != 2 { + t.Fatalf("clock-stratum: expected 2, got %v", ss["clock-stratum"]) + } + // refid is the sync source address + if ss["clock-refid"] != "192.168.0.1" { + t.Fatalf("clock-refid: expected '192.168.0.1', got %v", ss["clock-refid"]) + } + if ss["sync-state"] != "ietf-ntp:clock-synchronized" { + t.Fatalf("sync-state: expected clock-synchronized, got %v", ss["sync-state"]) + } + if toInt(ss["clock-precision"]) != -20 { + t.Fatalf("clock-precision: expected -20, got %v", ss["clock-precision"]) + } + + // Verify nominal/actual freq strings + if ss["nominal-freq"] != "1000000000.0000" { + t.Fatalf("nominal-freq: expected '1000000000.0000', got %v", ss["nominal-freq"]) + } + if ss["actual-freq"] != "999998500.0000" { + t.Fatalf("actual-freq: expected '999998500.0000', got %v", ss["actual-freq"]) + } + + // Clock offset (0.000045s → 0.045ms) + if ss["clock-offset"] != "0.045" { + t.Fatalf("clock-offset: expected '0.045', got %v", ss["clock-offset"]) + } + + // Infix augmentations + if ss["infix-ntp:update-interval"] != "64.0" { + t.Fatalf("update-interval: expected '64.0', got %v", ss["infix-ntp:update-interval"]) + } + + // Reference time should be an ISO timestamp + refTime, ok := ss["reference-time"].(string) + if !ok || !strings.HasPrefix(refTime, "2023-") { + t.Fatalf("reference-time should be 2023-* ISO timestamp, got %v", ss["reference-time"]) + } +} + +func TestNTPClockStateUnsynchronized(t *testing.T) { + c := newNTPCollector(ssRunner(), unsyncFakeChrony()) + tr := tree.New() + c.Collect(context.Background(), tr) + + raw := tr.Get("ietf-ntp:ntp") + if raw == nil { + t.Fatal("expected ietf-ntp:ntp even when unsynchronized") + } + var out map[string]interface{} + json.Unmarshal(raw, &out) + + cs := out["clock-state"].(map[string]interface{}) + ss := cs["system-status"].(map[string]interface{}) + + if ss["clock-state"] != "ietf-ntp:unsynchronized" { + t.Fatalf("clock-state: expected unsynchronized, got %v", ss["clock-state"]) + } + // Stratum 0 → 16 + if toInt(ss["clock-stratum"]) != 16 { + t.Fatalf("clock-stratum: expected 16 (mapped from 0), got %v", ss["clock-stratum"]) + } + if ss["sync-state"] != "ietf-ntp:clock-never-set" { + t.Fatalf("sync-state: expected clock-never-set, got %v", ss["sync-state"]) + } + if ss["clock-refid"] != "0.0.0.0" { + t.Fatalf("clock-refid: expected '0.0.0.0', got %v", ss["clock-refid"]) + } +} + +func TestNTPServerPort(t *testing.T) { + out, _ := ntpCollect(t, fullFakeChrony()) + + // Should find port 123 from the non-loopback ss line + if toInt(out["port"]) != 123 { + t.Fatalf("port: expected 123, got %v", out["port"]) + } +} + +func TestNTPRefclockMaster(t *testing.T) { + out, _ := ntpCollect(t, fullFakeChrony()) + master := out["refclock-master"].(map[string]interface{}) + if toInt(master["master-stratum"]) != 2 { + t.Fatalf("master-stratum: expected 2, got %v", master["master-stratum"]) + } +} + +func TestNTPServerStats(t *testing.T) { + out, _ := ntpCollect(t, fullFakeChrony()) + stats := out["ntp-statistics"].(map[string]interface{}) + + if toInt(stats["packet-received"]) != 1000 { + t.Fatalf("packet-received: expected 1000, got %v", stats["packet-received"]) + } + if toInt(stats["packet-dropped"]) != 5 { + t.Fatalf("packet-dropped: expected 5, got %v", stats["packet-dropped"]) + } + // chronyd does not count sent packets; the old chronyc CSV parsing + // reported auth/interleaved counters under these names by mistake. + if _, ok := stats["packet-sent"]; ok { + t.Fatal("packet-sent should not be reported") + } + if _, ok := stats["packet-sent-fail"]; ok { + t.Fatal("packet-sent-fail should not be reported") + } +} + +func TestNTPChronydUnreachable(t *testing.T) { + c := newNTPCollector(ssRunner(), nil) // dial fails + tr := tree.New() + err := c.Collect(context.Background(), tr) + if err != nil { + t.Fatalf("Collect should not error when chronyd unavailable: %v", err) + } + if tr.Get("ietf-ntp:ntp") != nil { + t.Fatal("should not set ietf-ntp:ntp when nothing to report") + } +} + +func TestNTPRefclockRefid(t *testing.T) { + // A refclock-synced chronyd has no source IP; the ASCII refid + // (e.g. "GPS") is reported instead + fake := &fakeChrony{ + tracking: &chrony.ReplyTracking{ + Tracking: chrony.Tracking{ + RefID: 0x47505300, // "GPS\0" + Stratum: 1, + RefTime: time.Unix(1700000000, 0), + }, + }, + } + + c := newNTPCollector(ssRunner(), fake) + tr := tree.New() + c.Collect(context.Background(), tr) + + var out map[string]interface{} + json.Unmarshal(tr.Get("ietf-ntp:ntp"), &out) + cs := out["clock-state"].(map[string]interface{}) + ss := cs["system-status"].(map[string]interface{}) + + // Space-padded to four chars: the RFC 9249 refid union only + // accepts an IPv4 address, a uint32, or exactly four characters + if ss["clock-refid"] != "GPS " { + t.Fatalf("clock-refid: expected 'GPS ', got %q", ss["clock-refid"]) + } +} + +func TestClockRefid(t *testing.T) { + tests := []struct { + name string + in chrony.Tracking + want interface{} + }{ + { + name: "ipv4 source", + in: chrony.Tracking{IPAddr: []byte{192, 168, 0, 1}, RefID: 0xC0A80001}, + want: "192.168.0.1", + }, + { + // The uint32 union member must be a JSON number; libyang + // rejects "1126301404" as a string + name: "ipv6 source falls back to numeric refid", + in: chrony.Tracking{ + IPAddr: []byte{0x20, 0x01, 0x0d, 0xb8, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1}, + RefID: 0x4321FEDC, + }, + want: uint32(0x4321FEDC), + }, + { + name: "refclock ascii refid", + in: chrony.Tracking{RefID: 0x47505300}, // "GPS\0" + want: "GPS ", + }, + { + name: "local reference renders as pseudo-IP", + in: chrony.Tracking{RefID: 0x7F7F0101}, + want: "127.127.1.1", + }, + { + name: "unsynchronized", + in: chrony.Tracking{}, + want: "0.0.0.0", + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + if got := clockRefid(&tc.in); got != tc.want { + t.Fatalf("clockRefid() = %v, want %v", got, tc.want) + } + }) + } +} + +func TestSplitLines(t *testing.T) { + input := "line1\n\nline2\n \nline3\n" + got := splitLines(input) + if len(got) != 3 { + t.Fatalf("expected 3 lines, got %d: %v", len(got), got) + } + if got[0] != "line1" || got[1] != "line2" || got[2] != "line3" { + t.Fatalf("unexpected lines: %v", got) + } +} + +// chrony 4 state codes: 5 is shown as '+' (selectable), 4 as '-' +// (unselected). The library's names for them come from chrony 3. +const ( + chronyUnselected = chrony.SourceStateType(4) + chronySelectable = chrony.SourceStateType(5) +) + +// A GET reads source selection from chronyd now: a source chrony selects +// after the last poll shows as selected without another poll. +func TestNTPLiveSourcesFollowChrony(t *testing.T) { + fake := fullFakeChrony() + c := newNTPCollector(ssRunner(), fake) + tr := tree.New() + tr.RegisterProvider("ietf-system:system-state", c.LiveSources) + if err := c.Collect(context.Background(), tr); err != nil { + t.Fatalf("Collect: %v", err) + } + + stateOf := func(addr string) interface{} { + var state map[string]interface{} + if err := json.Unmarshal(tr.Get("ietf-system:system-state"), &state); err != nil { + t.Fatalf("unmarshal system-state: %v", err) + } + srcs := state["infix-system:ntp"].(map[string]interface{})["sources"].(map[string]interface{})["source"].([]interface{}) + for _, raw := range srcs { + if s := raw.(map[string]interface{}); s["address"] == addr { + return s["state"] + } + } + return nil + } + + if got := stateOf("10.0.0.2"); got != "candidate" { + t.Fatalf("10.0.0.2 = %v before, want candidate", got) + } + fake.data[0].State = chronySelectable + fake.data[1].State = chrony.SourceStateSync + if got := stateOf("10.0.0.2"); got != "selected" { + t.Fatalf("10.0.0.2 = %v after chrony selected it, want selected without a poll", got) + } +} + +// With chronyd gone, Live adds nothing and LiveSources clears the list. +func TestNTPLiveWithoutChrony(t *testing.T) { + c := newNTPCollector(ssRunner(), nil) + if got := c.Live(); got != nil { + t.Fatalf("Live = %s, want nil", got) + } + if got := string(c.LiveSources()); got != `{"infix-system:ntp":{"sources":{"source":[]}}}` { + t.Fatalf("LiveSources = %s, want an empty source list", got) + } +} diff --git a/src/yangerd/internal/collector/pokes_test.go b/src/yangerd/internal/collector/pokes_test.go new file mode 100644 index 000000000..7573ce50d --- /dev/null +++ b/src/yangerd/internal/collector/pokes_test.go @@ -0,0 +1,59 @@ +package collector + +import ( + "context" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/kernelkit/infix/src/yangerd/internal/tree" +) + +type countingCollector struct { + name string + count atomic.Int32 +} + +func (c *countingCollector) Name() string { return c.name } +func (c *countingCollector) Interval() time.Duration { return time.Hour } +func (c *countingCollector) Collect(context.Context, *tree.Tree) error { + c.count.Add(1) + return nil +} + +func waitCount(t *testing.T, c *countingCollector, want int32) { + t.Helper() + deadline := time.Now().Add(2 * time.Second) + for time.Now().Before(deadline) { + if c.count.Load() == want { + return + } + time.Sleep(10 * time.Millisecond) + } + t.Fatalf("%s collected %d times, want %d", c.name, c.count.Load(), want) +} + +// A poke reaches the collector it names, and PokeAll reaches every one. +func TestPokesReachTheirCollector(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + var wg sync.WaitGroup + defer func() { cancel(); wg.Wait() }() + + a := &countingCollector{name: "a"} + b := &countingCollector{name: "b"} + pokes := RunAll(ctx, &wg, tree.New(), []Collector{a, b}) + waitCount(t, a, 1) + waitCount(t, b, 1) + + pokes.Poke("b") + waitCount(t, b, 2) + time.Sleep(50 * time.Millisecond) + if a.count.Load() != 1 { + t.Fatalf("poke for b ran a") + } + + pokes.PokeAll() + waitCount(t, a, 2) + waitCount(t, b, 3) +} diff --git a/src/yangerd/internal/collector/routing.go b/src/yangerd/internal/collector/routing.go new file mode 100644 index 000000000..e4d349a3c --- /dev/null +++ b/src/yangerd/internal/collector/routing.go @@ -0,0 +1,1225 @@ +package collector + +import ( + "context" + "encoding/json" + "path/filepath" + "regexp" + "sort" + "strconv" + "strings" + "time" + + "github.com/kernelkit/infix/src/yangerd/internal/frrvty" + "github.com/kernelkit/infix/src/yangerd/internal/tree" +) + +const ( + frrRunDir = "/var/run/frr" + + // vtyTimeout bounds one show command, so a wedged daemon cannot + // stall the poll. + vtyTimeout = 5 * time.Second +) + +// VtyQuery runs a show command against one FRR daemon, named as its vty +// socket is ("ospfd", "ospf6d", "ripd", "ripngd", "bfdd", "zebra"), and +// returns the output. +type VtyQuery func(ctx context.Context, daemon, command string) ([]byte, error) + +// FRRVty is the production VtyQuery. A daemon that is not running has +// no socket, so the dial fails and the protocol is skipped. +func FRRVty(ctx context.Context, daemon, command string) ([]byte, error) { + ctx, cancel := context.WithTimeout(ctx, vtyTimeout) + defer cancel() + return frrvty.New(filepath.Join(frrRunDir, daemon+".vty")).Query(ctx, command) +} + +// RoutingCollector gathers ietf-routing operational data by merging +// OSPF, RIP, and BFD control-plane protocols into a single tree key. +// Each protocol contributes entries to the control-plane-protocol list +// under ietf-routing:routing. +type RoutingCollector struct { + vty VtyQuery + interval time.Duration +} + +// NewRoutingCollector creates a RoutingCollector querying FRR over vty. +// The runner is no longer used; it stays until the caller drops it. +func NewRoutingCollector(interval time.Duration) *RoutingCollector { + return &RoutingCollector{vty: FRRVty, interval: interval} +} + +// Name implements Collector. +func (c *RoutingCollector) Name() string { return "routing" } + +// Interval implements Collector. +func (c *RoutingCollector) Interval() time.Duration { return c.interval } + +// vtyJSON runs a show command and decodes its JSON output into dst. +func (c *RoutingCollector) vtyJSON(ctx context.Context, daemon, command string, dst interface{}) error { + out, err := c.vty(ctx, daemon, command) + if err != nil { + return err + } + return json.Unmarshal(out, dst) +} + +// Collect implements Collector. It produces one tree key: +// "ietf-routing:routing" containing merged OSPF, RIP, and BFD data. +func (c *RoutingCollector) Collect(ctx context.Context, t *tree.Tree) error { + protocols := []interface{}{} + + if p := c.collectOSPF(ctx); p != nil { + protocols = append(protocols, p) + } + if p := c.collectOSPF6(ctx); p != nil { + protocols = append(protocols, p) + } + if p := c.collectRIP(ctx); p != nil { + protocols = append(protocols, p) + } + if p := c.collectRIPNG(ctx); p != nil { + protocols = append(protocols, p) + } + if p := c.collectBFD(ctx); p != nil { + protocols = append(protocols, p) + } + + // Always written, also when empty, so a protocol that stops running + // disappears instead of leaving its last state behind. + routing := map[string]interface{}{ + "control-plane-protocols": map[string]interface{}{ + "control-plane-protocol": protocols, + }, + } + + if data, err := json.Marshal(routing); err == nil { + t.Merge("ietf-routing:routing", data) + } + return nil +} + +// --- OSPF --- + +// ospfIfaceStateMap covers both ospfd's and ospf6d's spellings. The +// ietf if-state-type has no point-to-multipoint member, so ospf6d's +// "PtMultipoint" is left out. +var ospfIfaceStateMap = map[string]string{ + "DependUpon": "down", + "Down": "down", + "Waiting": "waiting", + "Loopback": "loopback", + "Point-To-Point": "point-to-point", + "PointToPoint": "point-to-point", + "DROther": "dr-other", + "Backup": "bdr", + "BDR": "bdr", + "DR": "dr", +} + +func frrToIETFNeighborState(state string) string { + parts := strings.SplitN(state, "/", 2) + s := parts[0] + // ospfd spells it "TwoWay", ospf6d "Twoway". + if strings.EqualFold(s, "TwoWay") { + return "2-way" + } + return strings.ToLower(s) +} + +func frrToIETFNeighborRole(role string) string { + if role == "Backup" { + return "BDR" + } + return role +} + +func ospfNetworkType(nt string, p2mpNonBroadcast bool) string { + switch nt { + case "POINTOPOINT": + return "point-to-point" + case "BROADCAST": + return "broadcast" + case "POINTOMULTIPOINT": + if p2mpNonBroadcast { + return "point-to-multipoint" + } + return "hybrid" + case "NBMA": + return "non-broadcast" + default: + return "" + } +} + +// ospfAreaSuffix is what ospfd appends to the area of a stub or NSSA +// area, e.g. "0.0.0.1 [Stub]". +var ospfAreaSuffix = map[string]string{ + " [Stub]": "stub-area", + " [NSSA]": "nssa-area", +} + +// ospfArea splits ospfd's area string into the bare area id and its +// area-type. +func ospfArea(area string) (string, string) { + for suffix, areaType := range ospfAreaSuffix { + if id, ok := strings.CutSuffix(area, suffix); ok { + return id, areaType + } + } + return area, "normal-area" +} + +// ospfStatus merges ospfd's three views, areas, interfaces and +// neighbors, into areas holding their interfaces and each interface its +// neighbors, the way the ietf-ospf model nests them. It is a port of +// the ospf-status helper yanger used. +func ospfStatus(ospf, ifaces, neighbors map[string]interface{}) map[string]interface{} { + areas, _ := ospf["areas"].(map[string]interface{}) + ifaceMap, _ := ifaces["interfaces"].(map[string]interface{}) + nbrMap, _ := neighbors["neighbors"].(map[string]interface{}) + + names := make([]string, 0, len(ifaceMap)) + for name := range ifaceMap { + names = append(names, name) + } + sort.Strings(names) + + for _, name := range names { + iface, ok := ifaceMap[name].(map[string]interface{}) + if !ok { + continue + } + enabled, _ := iface["ospfEnabled"].(bool) + areaStr, _ := iface["area"].(string) + if !enabled || areaStr == "" { + continue + } + + areaID, areaType := ospfArea(areaStr) + area, ok := areas[areaID].(map[string]interface{}) + if !ok { + continue + } + area["area-type"] = areaType + + iface["name"] = name + iface["area"] = areaID + iface["neighbors"] = ospfIfaceNeighbors(nbrMap, name, areaID) + + list, _ := area["interfaces"].([]interface{}) + area["interfaces"] = append(list, iface) + } + + return ospf +} + +// ospfIfaceNeighbors picks the neighbors seen on one interface in one +// area, tagging each with the neighbor's router id. +func ospfIfaceNeighbors(nbrMap map[string]interface{}, ifname, areaID string) []interface{} { + ids := make([]string, 0, len(nbrMap)) + for id := range nbrMap { + ids = append(ids, id) + } + sort.Strings(ids) + + out := []interface{}{} + for _, id := range ids { + list, _ := nbrMap[id].([]interface{}) + for _, raw := range list { + nbr, ok := raw.(map[string]interface{}) + if !ok { + continue + } + nbrArea, _ := nbr["areaId"].(string) + nbrArea, _ = ospfArea(nbrArea) + if nbr["ifaceName"] != ifname || nbrArea != areaID { + continue + } + nbr["areaId"] = nbrArea + nbr["neighborIp"] = id + out = append(out, nbr) + } + } + return out +} + +func (c *RoutingCollector) collectOSPF(ctx context.Context) interface{} { + var ospfData, ifaces, neighbors map[string]interface{} + if c.vtyJSON(ctx, "ospfd", "show ip ospf json", &ospfData) != nil || + c.vtyJSON(ctx, "ospfd", "show ip ospf interface json", &ifaces) != nil || + c.vtyJSON(ctx, "ospfd", "show ip ospf neighbor detail json", &neighbors) != nil { + return nil + } + if len(ospfData) == 0 { + return nil + } + data := ospfStatus(ospfData, ifaces, neighbors) + + ospf := map[string]interface{}{ + "ietf-ospf:address-family": "ipv4", + } + if rid, ok := data["routerId"]; ok { + ospf["ietf-ospf:router-id"] = rid + } + + areas := make([]interface{}, 0) + areasRaw, _ := data["areas"].(map[string]interface{}) + for areaID, valRaw := range areasRaw { + values, ok := valRaw.(map[string]interface{}) + if !ok { + continue + } + + area := map[string]interface{}{ + "ietf-ospf:area-id": areaID, + } + if at, ok := values["area-type"]; ok && at != nil { + area["ietf-ospf:area-type"] = at + } + + interfaces := make([]interface{}, 0) + ifacesRaw, _ := values["interfaces"].([]interface{}) + for _, ifaceRaw := range ifacesRaw { + if iface, ok := ifaceRaw.(map[string]interface{}); ok { + interfaces = append(interfaces, ospfInterface(iface)) + } + } + + area["ietf-ospf:interfaces"] = map[string]interface{}{ + "ietf-ospf:interface": interfaces, + } + areas = append(areas, area) + } + + if routes := c.ospfRoutes(ctx); len(routes) > 0 { + ospf["ietf-ospf:local-rib"] = map[string]interface{}{ + "ietf-ospf:route": routes, + } + } + + ospf["ietf-ospf:areas"] = map[string]interface{}{ + "ietf-ospf:area": areas, + } + + return map[string]interface{}{ + "type": "infix-routing:ospfv2", + "name": "default", + "ietf-ospf:ospf": ospf, + } +} + +// secondsAtLeast1 converts a millisecond FRR timer to whole seconds, +// dropping values that round to zero. +func secondsAtLeast1(dst map[string]interface{}, key string, msec interface{}) { + if msec == nil { + return + } + if sec := toInt(msec) / 1000; sec >= 1 { + dst[key] = sec + } +} + +func ospfInterface(iface map[string]interface{}) map[string]interface{} { + intf := map[string]interface{}{ + "name": iface["name"], + } + + setIfPresent(intf, "dr-router-id", iface, "drId") + setIfPresent(intf, "dr-ip-addr", iface, "drAddress") + setIfPresent(intf, "bdr-router-id", iface, "bdrId") + setIfPresent(intf, "bdr-ip-addr", iface, "bdrAddress") + + passive, ok := iface["timerPassiveIface"] + intf["passive"] = ok && passive != nil + + if v, ok := iface["ospfEnabled"]; ok { + intf["enabled"] = v + } + + if nt, ok := iface["networkType"].(string); ok { + p2mpNB, _ := iface["p2mpNonBroadcast"].(bool) + if it := ospfNetworkType(nt, p2mpNB); it != "" { + intf["interface-type"] = it + } + } + + if s, ok := iface["state"].(string); ok { + if mapped, ok := ospfIfaceStateMap[s]; ok { + intf["state"] = mapped + } else { + intf["state"] = "unknown" + } + } + + setIfPresentInt(intf, "priority", iface, "priority") + setIfPresentInt(intf, "cost", iface, "cost") + setIfPresentInt(intf, "dead-interval", iface, "timerDeadSecs") + setIfPresentInt(intf, "retransmit-interval", iface, "timerRetransmitSecs") + setIfPresentInt(intf, "transmit-delay", iface, "transmitDelaySecs") + + secondsAtLeast1(intf, "hello-interval", iface["timerMsecs"]) + secondsAtLeast1(intf, "hello-timer", iface["timerHelloInMsecs"]) + if v := iface["timerWaitSecs"]; v != nil && toInt(v) >= 1 { + intf["wait-timer"] = toInt(v) + } + + neighbors := make([]interface{}, 0) + neighsRaw, _ := iface["neighbors"].([]interface{}) + for _, neighRaw := range neighsRaw { + if neigh, ok := neighRaw.(map[string]interface{}); ok { + neighbors = append(neighbors, ospfNeighbor(neigh)) + } + } + intf["ietf-ospf:neighbors"] = map[string]interface{}{ + "ietf-ospf:neighbor": neighbors, + } + + return intf +} + +func ospfNeighbor(neigh map[string]interface{}) map[string]interface{} { + neighbor := map[string]interface{}{ + "neighbor-router-id": neigh["neighborIp"], + "address": neigh["ifaceAddress"], + } + + setIfPresentInt(neighbor, "priority", neigh, "nbrPriority") + + if v := neigh["lastPrgrsvChangeMsec"]; v != nil { + neighbor["infix-routing:uptime"] = toInt(v) / 1000 + } + secondsAtLeast1(neighbor, "dead-timer", neigh["routerDeadIntervalTimerDueMsec"]) + + if s, ok := neigh["nbrState"].(string); ok { + neighbor["state"] = frrToIETFNeighborState(s) + } + if role, ok := neigh["role"].(string); ok && role != "" { + neighbor["infix-routing:role"] = frrToIETFNeighborRole(role) + } + + ifName, _ := neigh["ifaceName"].(string) + localAddr, _ := neigh["localIfaceAddress"].(string) + if ifName != "" && localAddr != "" { + neighbor["infix-routing:interface-name"] = ifName + ":" + localAddr + } else if ifName != "" { + neighbor["infix-routing:interface-name"] = ifName + } + + setIfPresent(neighbor, "dr-router-id", neigh, "routerDesignatedId") + setIfPresent(neighbor, "bdr-router-id", neigh, "routerDesignatedBackupId") + + return neighbor +} + +func (c *RoutingCollector) ospfRoutes(ctx context.Context) []interface{} { + var data map[string]interface{} + if c.vtyJSON(ctx, "ospfd", "show ip ospf route json", &data) != nil { + return nil + } + + var routes []interface{} + for prefix, infoRaw := range data { + if !strings.Contains(prefix, "/") { + continue + } + if info, ok := infoRaw.(map[string]interface{}); ok { + routes = append(routes, ospfRoute(prefix, info)) + } + } + return routes +} + +func ospfRoute(prefix string, info map[string]interface{}) map[string]interface{} { + route := map[string]interface{}{ + "prefix": prefix, + } + + if rt, ok := info["routeType"].(string); ok { + parts := strings.Fields(rt) + if len(parts) > 1 { + switch parts[1] { + case "E1": + route["route-type"] = "external-1" + case "E2": + route["route-type"] = "external-2" + case "IA": + route["route-type"] = "inter-area" + } + } else if len(parts) > 0 && parts[0] == "N" { + route["route-type"] = "intra-area" + } + } + + if v := info["area"]; v != nil { + route["infix-routing:area-id"] = v + } + if v := info["cost"]; v != nil { + route["metric"] = v + } else if v := info["metric"]; v != nil { + route["metric"] = v + } + if v := info["tag"]; v != nil { + route["route-tag"] = v + } + + nexthops := make([]interface{}, 0) + hopsRaw, _ := info["nexthops"].([]interface{}) + for _, hopRaw := range hopsRaw { + hop, ok := hopRaw.(map[string]interface{}) + if !ok { + continue + } + nh := make(map[string]interface{}) + ip, _ := hop["ip"].(string) + if ip != "" && ip != " " { + nh["next-hop"] = ip + } else if da, ok := hop["directlyAttachedTo"].(string); ok { + nh["outgoing-interface"] = da + } + nexthops = append(nexthops, nh) + } + route["next-hops"] = map[string]interface{}{ + "next-hop": nexthops, + } + + return route +} + +// --- OSPFv3 --- + +// ospf6AreaType reads the flags ospf6d sets on an area in +// "show ipv6 ospf6 json". +func ospf6AreaType(area map[string]interface{}) string { + if v, _ := area["areaIsNSSA"].(bool); v { + return "nssa-area" + } + if v, _ := area["areaIsStub"].(bool); v { + return "stub-area" + } + return "normal-area" +} + +// collectOSPF6 is collectOSPF for ospf6d, whose JSON uses other key +// names: an interface carries its area id, and the neighbor list has no +// area, so neighbors are grouped by interface name alone (an interface +// is in exactly one area). +func (c *RoutingCollector) collectOSPF6(ctx context.Context) interface{} { + var top, ifaces, neighbors map[string]interface{} + if c.vtyJSON(ctx, "ospf6d", "show ipv6 ospf6 json", &top) != nil || + c.vtyJSON(ctx, "ospf6d", "show ipv6 ospf6 interface json", &ifaces) != nil || + c.vtyJSON(ctx, "ospf6d", "show ipv6 ospf6 neighbor json", &neighbors) != nil { + return nil + } + if len(top) == 0 { + return nil + } + + areaTypes := map[string]string{} + areasRaw, _ := top["areas"].(map[string]interface{}) + for id, raw := range areasRaw { + area, _ := raw.(map[string]interface{}) + areaTypes[id] = ospf6AreaType(area) + } + + nbrsByIface := map[string][]interface{}{} + nbrList, _ := neighbors["neighbors"].([]interface{}) + for _, raw := range nbrList { + if nbr, ok := raw.(map[string]interface{}); ok { + ifname, _ := nbr["interfaceName"].(string) + nbrsByIface[ifname] = append(nbrsByIface[ifname], ospf6Neighbor(nbr)) + } + } + + ifaceMap := ifaces + if inner, ok := ifaces["interfaces"].(map[string]interface{}); ok { + ifaceMap = inner + } + byArea := map[string][]interface{}{} + for _, name := range sortedKeys(ifaceMap) { + iface, ok := ifaceMap[name].(map[string]interface{}) + if !ok { + continue + } + areaID, _ := iface["areaId"].(string) + if attached, ok := iface["attachedToArea"].(bool); areaID == "" || (ok && !attached) { + continue + } + if _, ok := areaTypes[areaID]; !ok { + areaTypes[areaID] = "normal-area" + } + byArea[areaID] = append(byArea[areaID], ospf6Interface(name, iface, nbrsByIface[name])) + } + + areas := make([]interface{}, 0, len(areaTypes)) + for _, areaID := range sortedKeys(areaTypes) { + interfaces := byArea[areaID] + if interfaces == nil { + interfaces = []interface{}{} + } + areas = append(areas, map[string]interface{}{ + "ietf-ospf:area-id": areaID, + "ietf-ospf:area-type": areaTypes[areaID], + "ietf-ospf:interfaces": map[string]interface{}{ + "ietf-ospf:interface": interfaces, + }, + }) + } + + ospf := map[string]interface{}{ + "ietf-ospf:address-family": "ipv6", + } + if rid, ok := top["routerId"]; ok { + ospf["ietf-ospf:router-id"] = rid + } + if routes := c.ospf6Routes(ctx); len(routes) > 0 { + ospf["ietf-ospf:local-rib"] = map[string]interface{}{ + "ietf-ospf:route": routes, + } + } + ospf["ietf-ospf:areas"] = map[string]interface{}{ + "ietf-ospf:area": areas, + } + + return map[string]interface{}{ + "type": "infix-routing:ospfv3", + "name": "default", + "ietf-ospf:ospf": ospf, + } +} + +func ospf6Interface(name string, iface map[string]interface{}, neighbors []interface{}) map[string]interface{} { + intf := map[string]interface{}{ + "name": name, + "enabled": true, + } + + // "type" is the link type, always broadcast on Ethernet; + // "operatingAsType" is set when the OSPF network type differs. + // ospf6d has no NBMA or unicast point-to-multipoint. + nt, _ := iface["operatingAsType"].(string) + if nt == "" { + nt, _ = iface["type"].(string) + } + if it := ospfNetworkType(nt, false); it != "" { + intf["interface-type"] = it + } + + passive, _ := iface["timerPassiveIface"].(bool) + intf["passive"] = passive + + setIfPresentInt(intf, "cost", iface, "cost") + setIfPresentInt(intf, "priority", iface, "priority") + + if s, ok := iface["ospf6InterfaceState"].(string); ok { + if mapped, ok := ospfIfaceStateMap[s]; ok { + intf["state"] = mapped + } + } + + setIfPresentInt(intf, "dead-interval", iface, "timerIntervalsConfigDead") + setIfPresentInt(intf, "retransmit-interval", iface, "timerIntervalsConfigRetransmit") + setIfPresentInt(intf, "transmit-delay", iface, "transmitDelaySec") + setIfPresentInt(intf, "hello-interval", iface, "timerIntervalsConfigHello") + + if neighbors == nil { + neighbors = []interface{}{} + } + intf["ietf-ospf:neighbors"] = map[string]interface{}{ + "ietf-ospf:neighbor": neighbors, + } + + return intf +} + +func ospf6Neighbor(nbr map[string]interface{}) map[string]interface{} { + neighbor := map[string]interface{}{ + "neighbor-router-id": nbr["neighborId"], + } + + setIfPresent(neighbor, "address", nbr, "linkLocalAddress") + setIfPresentInt(neighbor, "priority", nbr, "priority") + + if s, ok := nbr["state"].(string); ok { + neighbor["state"] = frrToIETFNeighborState(s) + } + // ospf6d's ifState is the neighbor's DR election outcome on a + // broadcast link; on other link types it is the interface state. + switch role := nbr["ifState"]; role { + case "DR", "BDR", "DROther": + neighbor["infix-routing:role"] = role + } + if ifname, _ := nbr["interfaceName"].(string); ifname != "" { + neighbor["infix-routing:interface-name"] = ifname + } + + return neighbor +} + +// ospf6PathType maps ospf6d's abbreviated path types. +var ospf6PathType = map[string]string{ + "IA": "intra-area", + "IE": "inter-area", + "E1": "external-1", + "E2": "external-2", +} + +func (c *RoutingCollector) ospf6Routes(ctx context.Context) []interface{} { + var data map[string]interface{} + if c.vtyJSON(ctx, "ospf6d", "show ipv6 ospf6 route json", &data) != nil { + return nil + } + + var routes []interface{} + table, _ := data["routes"].(map[string]interface{}) + for prefix, pathsRaw := range table { + if !strings.Contains(prefix, "/") { + continue + } + // ospf6d lists one entry per path; keep the best one so the + // local-rib stays keyed uniquely by prefix. + var best map[string]interface{} + paths, _ := pathsRaw.([]interface{}) + for _, raw := range paths { + path, ok := raw.(map[string]interface{}) + if !ok { + continue + } + if best == nil { + best = path + } + if b, _ := path["isBestRoute"].(bool); b { + best = path + break + } + } + if best != nil { + routes = append(routes, ospf6Route(prefix, best)) + } + } + return routes +} + +func ospf6Route(prefix string, info map[string]interface{}) map[string]interface{} { + route := map[string]interface{}{ + "prefix": prefix, + } + + if pt, _ := info["pathType"].(string); pt != "" { + if rt, ok := ospf6PathType[pt]; ok { + route["route-type"] = rt + } + } + if v := info["area"]; v != nil { + route["infix-routing:area-id"] = v + } + if v := info["cost"]; v != nil { + route["metric"] = v + } else if v := info["metric"]; v != nil { + route["metric"] = v + } + + nexthops := make([]interface{}, 0) + hopsRaw, _ := info["nextHops"].([]interface{}) + for _, hopRaw := range hopsRaw { + hop, ok := hopRaw.(map[string]interface{}) + if !ok { + continue + } + nh := make(map[string]interface{}) + // "::" marks a directly connected prefix. + if ip, _ := hop["nextHop"].(string); ip != "" && ip != "::" { + nh["next-hop"] = ip + } else if ifname, _ := hop["interfaceName"].(string); ifname != "" { + nh["outgoing-interface"] = ifname + } + if len(nh) > 0 { + nexthops = append(nexthops, nh) + } + } + if len(nexthops) > 0 { + route["next-hops"] = map[string]interface{}{ + "next-hop": nexthops, + } + } + + return route +} + +// --- RIP --- + +// ripStatusScalars are the single-value lines of 'show ip rip status'. +var ripStatusScalars = []struct { + re *regexp.Regexp + key string +}{ + {regexp.MustCompile(`Sending updates every (\d+) seconds`), "update-interval"}, + {regexp.MustCompile(`Timeout after (\d+) seconds`), "invalid-interval"}, + {regexp.MustCompile(`garbage collect after (\d+) seconds`), "flush-interval"}, + {regexp.MustCompile(`Default redistribution metric is (\d+)`), "default-metric"}, + {regexp.MustCompile(`Distance: \(default is (\d+)\)`), "distance"}, +} + +// ripVersion maps FRR's ri_version_msg to the infix-routing enum. +var ripVersion = map[string]string{ + "1": "1", + "2": "2", + "1 2": "1-2", +} + +func (c *RoutingCollector) collectRIP(ctx context.Context) interface{} { + statusOut, err := c.vty(ctx, "ripd", "show ip rip status") + if err != nil || len(statusOut) == 0 { + return nil + } + + status := parseRIPStatus(string(statusOut)) + if len(status) == 0 { + return nil + } + + rip := ripGlobals(status) + if ifaces := ripInterfaces(status, true); len(ifaces) > 0 { + rip["interfaces"] = map[string]interface{}{ + "interface": ifaces, + } + } + ripAddressFamily(rip, "ipv4", c.ripRoutes(ctx, "ipv4"), ripNeighbors(status, "ipv4-address")) + + return map[string]interface{}{ + "type": "infix-routing:ripv2", + "name": "default", + "ietf-rip:rip": rip, + } +} + +// collectRIPNG is collectRIP for ripngd. Its status text has the same +// layout apart from the peer table, which ripngd prints as two lines per +// peer, and RIPng has no protocol version to report. +func (c *RoutingCollector) collectRIPNG(ctx context.Context) interface{} { + statusOut, err := c.vty(ctx, "ripngd", "show ipv6 ripng status") + if err != nil || len(statusOut) == 0 { + return nil + } + + status := parseRIPStatus(string(statusOut)) + if len(status) == 0 { + return nil + } + status["neighbors"] = parseRIPNGNeighbors(strings.Split(string(statusOut), "\n")) + + rip := ripGlobals(status) + if ifaces := ripInterfaces(status, false); len(ifaces) > 0 { + rip["interfaces"] = map[string]interface{}{ + "interface": ifaces, + } + } + ripAddressFamily(rip, "ipv6", c.ripRoutes(ctx, "ipv6"), ripNeighbors(status, "ipv6-address")) + + return map[string]interface{}{ + "type": "infix-routing:ripng", + "name": "default", + "ietf-rip:rip": rip, + } +} + +func ripGlobals(status map[string]interface{}) map[string]interface{} { + rip := make(map[string]interface{}) + setIfPresent(rip, "distance", status, "distance") + setIfPresent(rip, "default-metric", status, "default-metric") + + timers := make(map[string]interface{}) + for _, key := range []string{"update-interval", "invalid-interval", "flush-interval"} { + setIfPresent(timers, key, status, key) + } + if len(timers) > 0 { + rip["timers"] = timers + } + return rip +} + +// ripAddressFamily fills the ipv4 or ipv6 container with the learned +// routes and the peers. +func ripAddressFamily(rip map[string]interface{}, family string, routes, neighbors []interface{}) { + af := make(map[string]interface{}) + if len(routes) > 0 { + af["routes"] = map[string]interface{}{ + "route": routes, + } + rip["num-of-routes"] = len(routes) + } + if len(neighbors) > 0 { + af["neighbors"] = map[string]interface{}{ + "neighbor": neighbors, + } + } + if len(af) > 0 { + rip[family] = af + } +} + +func ripInterfaces(status map[string]interface{}, withVersion bool) []interface{} { + var out []interface{} + ifaces, _ := status["interfaces"].([]interface{}) + for _, ifRaw := range ifaces { + ifData, ok := ifRaw.(map[string]interface{}) + if !ok { + continue + } + entry := map[string]interface{}{ + "interface": ifData["name"], + "oper-status": "up", + } + if withVersion { + setIfPresent(entry, "send-version", ifData, "send-version") + setIfPresent(entry, "receive-version", ifData, "recv-version") + } + out = append(out, entry) + } + return out +} + +func ripNeighbors(status map[string]interface{}, addressKey string) []interface{} { + var out []interface{} + neighs, _ := status["neighbors"].([]interface{}) + for _, nRaw := range neighs { + nd, ok := nRaw.(map[string]interface{}) + if !ok { + continue + } + entry := map[string]interface{}{ + addressKey: nd["address"], + } + setIfPresent(entry, "bad-packets-rcvd", nd, "bad-packets") + setIfPresent(entry, "bad-routes-rcvd", nd, "bad-routes") + out = append(out, entry) + } + return out +} + +func (c *RoutingCollector) ripRoutes(ctx context.Context, family string) []interface{} { + command, prefixKey := "show ip route rip json", "ipv4-prefix" + if family == "ipv6" { + command, prefixKey = "show ipv6 route ripng json", "ipv6-prefix" + } + + var routeData map[string]interface{} + if c.vtyJSON(ctx, "zebra", command, &routeData) != nil { + return nil + } + + var routes []interface{} + for prefix, entriesRaw := range routeData { + if !strings.Contains(prefix, "/") { + continue + } + entries, _ := entriesRaw.([]interface{}) + if len(entries) == 0 { + continue + } + entry, ok := entries[0].(map[string]interface{}) + if !ok { + continue + } + + route := map[string]interface{}{ + prefixKey: prefix, + "route-type": "rip", + } + if m, ok := entry["metric"]; ok { + route["metric"] = toInt(m) + } + + nexthops, _ := entry["nexthops"].([]interface{}) + if len(nexthops) > 0 { + firstHop, _ := nexthops[0].(map[string]interface{}) + if ip, ok := firstHop["ip"].(string); ok && ip != "" { + route["next-hop"] = ip + } + if ifName, ok := firstHop["interfaceName"].(string); ok && ifName != "" { + route["interface"] = ifName + } + } + routes = append(routes, route) + } + return routes +} + +// parseRIPStatus parses the text output of 'show ip rip status'. +func parseRIPStatus(text string) map[string]interface{} { + status := make(map[string]interface{}) + + for _, sc := range ripStatusScalars { + if m := sc.re.FindStringSubmatch(text); m != nil { + v, _ := strconv.Atoi(m[1]) + status[sc.key] = v + } + } + + lines := strings.Split(text, "\n") + if interfaces := parseRIPInterfaces(lines); len(interfaces) > 0 { + status["interfaces"] = interfaces + } + if neighbors := parseRIPNeighbors(lines); len(neighbors) > 0 { + status["neighbors"] = neighbors + } + + return status +} + +// parseRIPInterfaces reads the interface table. ripd prints it with +// "%-17s%-3s %-3s", and a version can be "1 2", so the columns are +// taken by position rather than split on whitespace. +func parseRIPInterfaces(lines []string) []interface{} { + var interfaces []interface{} + inTable := false + + for _, raw := range lines { + line := strings.TrimSpace(raw) + if strings.HasPrefix(line, "Interface") && strings.Contains(line, "Send") && strings.Contains(line, "Recv") { + inTable = true + continue + } + if !inTable { + continue + } + if line == "" || strings.HasPrefix(line, "Routing for Networks:") || strings.HasPrefix(line, "Routing Information Sources:") { + break + } + + row := strings.TrimLeft(raw, " ") + name := strings.Fields(row)[0] + col := len(name) + if col < 17 { + col = 17 + } + send, ok1 := ripVersion[column(row, col, 3)] + recv, ok2 := ripVersion[column(row, col+6, 3)] + if !ok1 || !ok2 { + continue + } + interfaces = append(interfaces, map[string]interface{}{ + "name": name, + "send-version": send, + "recv-version": recv, + }) + } + + return interfaces +} + +// parseRIPNeighbors reads the "Routing Information Sources" table. +func parseRIPNeighbors(lines []string) []interface{} { + var neighbors []interface{} + inTable := false + + for _, raw := range lines { + line := strings.TrimSpace(raw) + if strings.HasPrefix(line, "Routing Information Sources:") { + inTable = true + continue + } + if !inTable || strings.HasPrefix(line, "Gateway") { + continue + } + if strings.HasPrefix(line, "Distance:") || (line == "" && len(neighbors) > 0) { + break + } + + parts := strings.Fields(line) + if len(parts) < 5 { + continue + } + badPkts, err1 := strconv.Atoi(parts[1]) + badRoutes, err2 := strconv.Atoi(parts[2]) + if err1 != nil || err2 != nil { + continue + } + neighbors = append(neighbors, map[string]interface{}{ + "address": parts[0], + "bad-packets": badPkts, + "bad-routes": badRoutes, + }) + } + + return neighbors +} + +// parseRIPNGNeighbors reads ripngd's "Routing Information Sources" +// table, printed as two lines per peer: the address, then its counters. +func parseRIPNGNeighbors(lines []string) []interface{} { + var neighbors []interface{} + inTable := false + pending := "" + + for _, raw := range lines { + line := strings.TrimSpace(raw) + if strings.HasPrefix(line, "Routing Information Sources:") { + inTable = true + continue + } + if !inTable || strings.HasPrefix(line, "Gateway") { + continue + } + if line == "" { + if len(neighbors) > 0 && pending == "" { + break + } + continue + } + + if pending == "" { + if strings.Contains(line, ":") { + pending, _, _ = strings.Cut(strings.Fields(line)[0], "%") + } + continue + } + + parts := strings.Fields(line) + if len(parts) >= 3 { + badPkts, err1 := strconv.Atoi(parts[0]) + badRoutes, err2 := strconv.Atoi(parts[1]) + if err1 == nil && err2 == nil { + neighbors = append(neighbors, map[string]interface{}{ + "address": pending, + "bad-packets": badPkts, + "bad-routes": badRoutes, + }) + } + } + pending = "" + } + + return neighbors +} + +// column returns the trimmed text in [start, start+width) of s. +func column(s string, start, width int) string { + if start >= len(s) { + return "" + } + end := start + width + if end > len(s) { + end = len(s) + } + return strings.TrimSpace(s[start:end]) +} + +// --- BFD --- + +var bfdStateMap = map[string]string{ + "up": "up", + "down": "down", + "init": "init", + "adminDown": "adminDown", +} + +func (c *RoutingCollector) collectBFD(ctx context.Context) interface{} { + var data []interface{} + if c.vtyJSON(ctx, "bfdd", "show bfd peers json", &data) != nil || len(data) == 0 { + return nil + } + + var sessions []interface{} + for _, peerRaw := range data { + peer, ok := peerRaw.(map[string]interface{}) + if !ok { + continue + } + // Only process single-hop sessions (multihop == false) + if mh, _ := peer["multihop"].(bool); mh { + continue + } + + session := map[string]interface{}{ + "interface": strDefault(peer["interface"], "unknown"), + "dest-addr": strDefault(peer["peer"], "0.0.0.0"), + } + + if v := peer["id"]; v != nil { + session["local-discriminator"] = v + } + if v := peer["remote-id"]; v != nil { + session["remote-discriminator"] = v + } + + state := strDefault(peer["status"], "down") + ietfState := bfdStateMap[state] + if ietfState == "" { + ietfState = "down" + } + + sessionRunning := map[string]interface{}{ + "local-state": ietfState, + "remote-state": ietfState, + "local-diagnostic": "none", + "detection-mode": "async-without-echo", + } + + if v := peer["receive-interval"]; v != nil { + sessionRunning["negotiated-rx-interval"] = toInt(v) * 1000 + } + if v := peer["transmit-interval"]; v != nil { + sessionRunning["negotiated-tx-interval"] = toInt(v) * 1000 + } + if dm := peer["detect-multiplier"]; dm != nil { + if ri := peer["receive-interval"]; ri != nil { + detectionTimeMs := toInt(dm) * toInt(ri) + sessionRunning["detection-time"] = detectionTimeMs * 1000 + } + } + + session["session-running"] = sessionRunning + session["path-type"] = "ietf-bfd-types:path-ip-sh" + session["ip-encapsulation"] = true + + sessions = append(sessions, session) + } + + if len(sessions) == 0 { + return nil + } + + return map[string]interface{}{ + "type": "infix-routing:bfdv1", + "name": "bfd", + "ietf-bfd:bfd": map[string]interface{}{ + "ietf-bfd-ip-sh:ip-sh": map[string]interface{}{ + "sessions": map[string]interface{}{ + "session": sessions, + }, + }, + }, + } +} + +// --- Helpers --- + +func setIfPresent(dst map[string]interface{}, dstKey string, src map[string]interface{}, srcKey string) { + if v, ok := src[srcKey]; ok && v != nil { + dst[dstKey] = v + } +} + +func setIfPresentInt(dst map[string]interface{}, dstKey string, src map[string]interface{}, srcKey string) { + if v, ok := src[srcKey]; ok && v != nil { + dst[dstKey] = toInt(v) + } +} + +func sortedKeys[V any](m map[string]V) []string { + keys := make([]string, 0, len(m)) + for k := range m { + keys = append(keys, k) + } + sort.Strings(keys) + return keys +} + +func strDefault(v interface{}, def string) string { + if s, ok := v.(string); ok && s != "" { + return s + } + return def +} diff --git a/src/yangerd/internal/collector/routing_test.go b/src/yangerd/internal/collector/routing_test.go new file mode 100644 index 000000000..d07ff920e --- /dev/null +++ b/src/yangerd/internal/collector/routing_test.go @@ -0,0 +1,1122 @@ +package collector + +import ( + "context" + "encoding/json" + "fmt" + "testing" + "time" + + "github.com/kernelkit/infix/src/yangerd/internal/tree" +) + +// Canned ospfd JSON: the areas, interfaces and neighbors views that +// ospfStatus merges. +const testOSPFGlobal = `{ + "routerId": "10.0.0.1", + "areas": { + "0.0.0.0": {} + } +}` + +const testOSPFInterfaces = `{ + "interfaces": { + "e0": { + "area": "0.0.0.0", + "state": "DR", + "ospfEnabled": true, + "networkType": "BROADCAST", + "cost": 10, + "priority": 1, + "timerDeadSecs": 40, + "timerRetransmitSecs": 5, + "transmitDelaySecs": 1, + "timerMsecs": 10000, + "timerHelloInMsecs": 7000, + "timerWaitSecs": 40, + "drId": "10.0.0.1", + "drAddress": "192.168.1.1", + "bdrId": "10.0.0.2", + "bdrAddress": "192.168.1.2" + }, + "lo": { + "area": "0.0.0.0", + "state": "Loopback", + "ospfEnabled": true, + "networkType": "POINTOPOINT", + "cost": 0, + "priority": 0, + "timerPassiveIface": true, + "timerDeadSecs": 0, + "timerRetransmitSecs": 0, + "transmitDelaySecs": 0, + "timerMsecs": 10000 + }, + "e9": { + "ospfEnabled": false + } + } +}` + +const testOSPFNeighbors = `{ + "neighbors": { + "10.0.0.2": [ + { + "areaId": "0.0.0.0", + "ifaceAddress": "192.168.1.2", + "nbrPriority": 1, + "nbrState": "Full/DR", + "role": "Backup", + "lastPrgrsvChangeMsec": 120000, + "routerDeadIntervalTimerDueMsec": 35000, + "routerDesignatedId": "10.0.0.1", + "routerDesignatedBackupId": "10.0.0.2", + "ifaceName": "e0", + "localIfaceAddress": "192.168.1.1" + } + ] + } +}` + +const testOSPFRoutes = `{ + "10.0.0.0/24": { + "routeType": "N IA", + "area": "0.0.0.0", + "cost": 20, + "nexthops": [ + {"ip": "192.168.1.2", "via": "e0"} + ] + }, + "10.0.1.0/24": { + "routeType": "N E2", + "area": "0.0.0.0", + "cost": 100, + "tag": 42, + "nexthops": [ + {"ip": " ", "directlyAttachedTo": "e0"} + ] + } +}` + +const testRIPStatus = `Routing Protocol is "rip" + Sending updates every 30 seconds with +/-50%, next due in 12 seconds + Timeout after 180 seconds, garbage collect after 120 seconds + Outgoing update filter list for all interface is not set + Incoming update filter list for all interface is not set + Default redistribution metric is 1 + Redistributing: + Default version control: send version 2, receive version 2 + Interface Send Recv Key-chain + e0 2 2 + e1 2 2 + Routing for Networks: + 10.0.0.0/24 + 10.0.1.0/24 + Routing Information Sources: + Gateway BadPackets BadRoutes Distance Last Update + 10.0.0.2 0 0 120 00:00:12 + 10.0.0.3 1 2 120 00:00:25 + Distance: (default is 120) +` + +const testRIPRoutes = `{ + "10.0.0.0/24": [ + { + "prefix": "10.0.0.0/24", + "protocol": "rip", + "metric": 1, + "nexthops": [ + {"ip": "10.0.0.2", "interfaceName": "e0"} + ] + } + ], + "10.0.1.0/24": [ + { + "prefix": "10.0.1.0/24", + "protocol": "rip", + "metric": 2, + "nexthops": [ + {"ip": "10.0.0.3", "interfaceName": "e1"} + ] + } + ] +}` + +const testBFDPeers = `[ + { + "multihop": false, + "peer": "10.0.0.2", + "interface": "e0", + "id": 1, + "remote-id": 2, + "status": "up", + "receive-interval": 300, + "transmit-interval": 300, + "detect-multiplier": 3 + }, + { + "multihop": true, + "peer": "10.0.0.99", + "interface": "e1", + "id": 5, + "remote-id": 6, + "status": "down" + } +]` + +// fakeVty answers show commands from canned output keyed by +// "daemon command"; any other daemon is not running. +type fakeVty map[string]string + +func (f fakeVty) query(_ context.Context, daemon, command string) ([]byte, error) { + if out, ok := f[daemon+" "+command]; ok { + return []byte(out), nil + } + return nil, fmt.Errorf("dial %s.vty: no such file or directory", daemon) +} + +func newRoutingCollector(vty fakeVty) *RoutingCollector { + c := NewRoutingCollector(10 * time.Second) + c.vty = vty.query + return c +} + +func routingCollect(t *testing.T, vty fakeVty) map[string]interface{} { + t.Helper() + c := newRoutingCollector(vty) + tr := tree.New() + if err := c.Collect(context.Background(), tr); err != nil { + t.Fatalf("Collect failed: %v", err) + } + raw := tr.Get("ietf-routing:routing") + if raw == nil { + t.Fatal("missing ietf-routing:routing in tree") + } + var out map[string]interface{} + if err := json.Unmarshal(raw, &out); err != nil { + t.Fatalf("unmarshal routing: %v", err) + } + return out +} + +func ospfOnly() fakeVty { + return fakeVty{ + "ospfd show ip ospf json": testOSPFGlobal, + "ospfd show ip ospf interface json": testOSPFInterfaces, + "ospfd show ip ospf neighbor detail json": testOSPFNeighbors, + "ospfd show ip ospf route json": testOSPFRoutes, + } +} + +func fullRunner() fakeVty { + v := ospfOnly() + v["ripd show ip rip status"] = testRIPStatus + v["zebra show ip route rip json"] = testRIPRoutes + v["bfdd show bfd peers json"] = testBFDPeers + return v +} + +func TestRoutingCollectorNameAndInterval(t *testing.T) { + c := newRoutingCollector(fullRunner()) + if c.Name() != "routing" { + t.Fatalf("expected name 'routing', got %q", c.Name()) + } + if c.Interval() != 10*time.Second { + t.Fatalf("expected interval 10s, got %v", c.Interval()) + } +} + +func TestRoutingCollectorMergesThreeProtocols(t *testing.T) { + out := routingCollect(t, fullRunner()) + cpp := out["control-plane-protocols"].(map[string]interface{}) + protocols := cpp["control-plane-protocol"].([]interface{}) + if len(protocols) != 3 { + t.Fatalf("expected 3 protocols (OSPF+RIP+BFD), got %d", len(protocols)) + } + + types := make(map[string]bool) + for _, p := range protocols { + pm := p.(map[string]interface{}) + types[pm["type"].(string)] = true + } + for _, expected := range []string{"infix-routing:ospfv2", "infix-routing:ripv2", "infix-routing:bfdv1"} { + if !types[expected] { + t.Fatalf("missing protocol type %q; got %v", expected, types) + } + } +} + +// --- OSPF tests --- + +func getOSPFProtocol(t *testing.T, out map[string]interface{}) map[string]interface{} { + t.Helper() + cpp := out["control-plane-protocols"].(map[string]interface{}) + for _, p := range cpp["control-plane-protocol"].([]interface{}) { + pm := p.(map[string]interface{}) + if pm["type"] == "infix-routing:ospfv2" { + return pm + } + } + t.Fatal("OSPF protocol not found") + return nil +} + +func TestOSPFRouterID(t *testing.T) { + out := routingCollect(t, fullRunner()) + ospfProto := getOSPFProtocol(t, out) + ospf := ospfProto["ietf-ospf:ospf"].(map[string]interface{}) + if ospf["ietf-ospf:router-id"] != "10.0.0.1" { + t.Fatalf("router-id: expected 10.0.0.1, got %v", ospf["ietf-ospf:router-id"]) + } + if ospf["ietf-ospf:address-family"] != "ipv4" { + t.Fatalf("address-family: expected ipv4, got %v", ospf["ietf-ospf:address-family"]) + } +} + +func TestOSPFAreaAndInterfaces(t *testing.T) { + out := routingCollect(t, fullRunner()) + ospfProto := getOSPFProtocol(t, out) + ospf := ospfProto["ietf-ospf:ospf"].(map[string]interface{}) + areasContainer := ospf["ietf-ospf:areas"].(map[string]interface{}) + areas := areasContainer["ietf-ospf:area"].([]interface{}) + if len(areas) != 1 { + t.Fatalf("expected 1 area, got %d", len(areas)) + } + + area := areas[0].(map[string]interface{}) + if area["ietf-ospf:area-id"] != "0.0.0.0" { + t.Fatalf("area-id: expected 0.0.0.0, got %v", area["ietf-ospf:area-id"]) + } + + ifacesContainer := area["ietf-ospf:interfaces"].(map[string]interface{}) + ifaces := ifacesContainer["ietf-ospf:interface"].([]interface{}) + if len(ifaces) != 2 { + t.Fatalf("expected 2 interfaces, got %d", len(ifaces)) + } + + // First interface: e0 (DR) + e0 := ifaces[0].(map[string]interface{}) + if e0["name"] != "e0" { + t.Fatalf("interface[0] name: expected e0, got %v", e0["name"]) + } + if e0["state"] != "dr" { + t.Fatalf("interface[0] state: expected dr, got %v", e0["state"]) + } + if e0["interface-type"] != "broadcast" { + t.Fatalf("interface[0] type: expected broadcast, got %v", e0["interface-type"]) + } + if e0["passive"] != false { + t.Fatalf("e0 passive: expected false, got %v", e0["passive"]) + } + if e0["enabled"] != true { + t.Fatalf("e0 enabled: expected true, got %v", e0["enabled"]) + } + if e0["dr-router-id"] != "10.0.0.1" { + t.Fatalf("e0 dr-router-id: expected 10.0.0.1, got %v", e0["dr-router-id"]) + } + + // Check timer conversions (ms → seconds) + if toInt(e0["hello-interval"]) != 10 { + t.Fatalf("e0 hello-interval: expected 10, got %v", e0["hello-interval"]) + } + if toInt(e0["hello-timer"]) != 7 { + t.Fatalf("e0 hello-timer: expected 7, got %v", e0["hello-timer"]) + } + if toInt(e0["cost"]) != 10 { + t.Fatalf("e0 cost: expected 10, got %v", e0["cost"]) + } + + // Second interface: lo (passive loopback) + lo := ifaces[1].(map[string]interface{}) + if lo["name"] != "lo" { + t.Fatalf("interface[1] name: expected lo, got %v", lo["name"]) + } + if lo["state"] != "loopback" { + t.Fatalf("lo state: expected loopback, got %v", lo["state"]) + } + if lo["passive"] != true { + t.Fatalf("lo passive: expected true, got %v", lo["passive"]) + } + if lo["interface-type"] != "point-to-point" { + t.Fatalf("lo type: expected point-to-point, got %v", lo["interface-type"]) + } +} + +func TestOSPFNeighbors(t *testing.T) { + out := routingCollect(t, fullRunner()) + ospfProto := getOSPFProtocol(t, out) + ospf := ospfProto["ietf-ospf:ospf"].(map[string]interface{}) + areasContainer := ospf["ietf-ospf:areas"].(map[string]interface{}) + area := areasContainer["ietf-ospf:area"].([]interface{})[0].(map[string]interface{}) + ifacesContainer := area["ietf-ospf:interfaces"].(map[string]interface{}) + e0 := ifacesContainer["ietf-ospf:interface"].([]interface{})[0].(map[string]interface{}) + + neighborsContainer := e0["ietf-ospf:neighbors"].(map[string]interface{}) + neighbors := neighborsContainer["ietf-ospf:neighbor"].([]interface{}) + if len(neighbors) != 1 { + t.Fatalf("expected 1 neighbor, got %d", len(neighbors)) + } + + n := neighbors[0].(map[string]interface{}) + if n["neighbor-router-id"] != "10.0.0.2" { + t.Fatalf("neighbor router-id: expected 10.0.0.2, got %v", n["neighbor-router-id"]) + } + if n["address"] != "192.168.1.2" { + t.Fatalf("neighbor address: expected 192.168.1.2, got %v", n["address"]) + } + if n["state"] != "full" { + t.Fatalf("neighbor state: expected full, got %v", n["state"]) + } + if n["infix-routing:role"] != "BDR" { + t.Fatalf("neighbor role: expected BDR, got %v", n["infix-routing:role"]) + } + // Uptime: 120000ms → 120s + if toInt(n["infix-routing:uptime"]) != 120 { + t.Fatalf("neighbor uptime: expected 120, got %v", n["infix-routing:uptime"]) + } + // Dead timer: 35000ms → 35s + if toInt(n["dead-timer"]) != 35 { + t.Fatalf("neighbor dead-timer: expected 35, got %v", n["dead-timer"]) + } + // Interface name augmentation + if n["infix-routing:interface-name"] != "e0:192.168.1.1" { + t.Fatalf("neighbor interface-name: expected e0:192.168.1.1, got %v", n["infix-routing:interface-name"]) + } +} + +func TestOSPFRoutes(t *testing.T) { + out := routingCollect(t, fullRunner()) + ospfProto := getOSPFProtocol(t, out) + ospf := ospfProto["ietf-ospf:ospf"].(map[string]interface{}) + rib := ospf["ietf-ospf:local-rib"].(map[string]interface{}) + routes := rib["ietf-ospf:route"].([]interface{}) + if len(routes) != 2 { + t.Fatalf("expected 2 OSPF routes, got %d", len(routes)) + } + + routeByPrefix := make(map[string]map[string]interface{}) + for _, r := range routes { + rm := r.(map[string]interface{}) + routeByPrefix[rm["prefix"].(string)] = rm + } + + // Inter-area route + r1 := routeByPrefix["10.0.0.0/24"] + if r1 == nil { + t.Fatal("missing route 10.0.0.0/24") + } + if r1["route-type"] != "inter-area" { + t.Fatalf("route 10.0.0.0/24 type: expected inter-area, got %v", r1["route-type"]) + } + + // External-2 route with tag + r2 := routeByPrefix["10.0.1.0/24"] + if r2 == nil { + t.Fatal("missing route 10.0.1.0/24") + } + if r2["route-type"] != "external-2" { + t.Fatalf("route 10.0.1.0/24 type: expected external-2, got %v", r2["route-type"]) + } + // tag should be present + if r2["route-tag"] == nil { + t.Fatal("route 10.0.1.0/24 should have route-tag") + } + // Directly attached nexthop + nhs := r2["next-hops"].(map[string]interface{}) + nhList := nhs["next-hop"].([]interface{}) + if len(nhList) != 1 { + t.Fatalf("expected 1 nexthop, got %d", len(nhList)) + } + nh := nhList[0].(map[string]interface{}) + if nh["outgoing-interface"] != "e0" { + t.Fatalf("nexthop outgoing-interface: expected e0, got %v", nh["outgoing-interface"]) + } +} + +// --- RIP tests --- + +func getRIPProtocol(t *testing.T, out map[string]interface{}) map[string]interface{} { + t.Helper() + cpp := out["control-plane-protocols"].(map[string]interface{}) + for _, p := range cpp["control-plane-protocol"].([]interface{}) { + pm := p.(map[string]interface{}) + if pm["type"] == "infix-routing:ripv2" { + return pm + } + } + t.Fatal("RIP protocol not found") + return nil +} + +func TestRIPTimers(t *testing.T) { + out := routingCollect(t, fullRunner()) + ripProto := getRIPProtocol(t, out) + rip := ripProto["ietf-rip:rip"].(map[string]interface{}) + + timers := rip["timers"].(map[string]interface{}) + if toInt(timers["update-interval"]) != 30 { + t.Fatalf("RIP update-interval: expected 30, got %v", timers["update-interval"]) + } + if toInt(timers["invalid-interval"]) != 180 { + t.Fatalf("RIP invalid-interval: expected 180, got %v", timers["invalid-interval"]) + } + if toInt(timers["flush-interval"]) != 120 { + t.Fatalf("RIP flush-interval: expected 120, got %v", timers["flush-interval"]) + } + + if toInt(rip["default-metric"]) != 1 { + t.Fatalf("RIP default-metric: expected 1, got %v", rip["default-metric"]) + } + if toInt(rip["distance"]) != 120 { + t.Fatalf("RIP distance: expected 120, got %v", rip["distance"]) + } +} + +func TestRIPInterfaces(t *testing.T) { + out := routingCollect(t, fullRunner()) + ripProto := getRIPProtocol(t, out) + rip := ripProto["ietf-rip:rip"].(map[string]interface{}) + + ifContainer := rip["interfaces"].(map[string]interface{}) + ifaces := ifContainer["interface"].([]interface{}) + if len(ifaces) != 2 { + t.Fatalf("expected 2 RIP interfaces, got %d", len(ifaces)) + } + + iface0 := ifaces[0].(map[string]interface{}) + if iface0["interface"] != "e0" { + t.Fatalf("RIP iface[0]: expected e0, got %v", iface0["interface"]) + } + if iface0["oper-status"] != "up" { + t.Fatalf("RIP iface[0] status: expected up, got %v", iface0["oper-status"]) + } + if iface0["send-version"] != "2" { + t.Fatalf("RIP iface[0] send-version: expected '2', got %v", iface0["send-version"]) + } +} + +func TestRIPRoutes(t *testing.T) { + out := routingCollect(t, fullRunner()) + ripProto := getRIPProtocol(t, out) + rip := ripProto["ietf-rip:rip"].(map[string]interface{}) + + ipv4 := rip["ipv4"].(map[string]interface{}) + routesContainer := ipv4["routes"].(map[string]interface{}) + routes := routesContainer["route"].([]interface{}) + if len(routes) != 2 { + t.Fatalf("expected 2 RIP routes, got %d", len(routes)) + } + + if toInt(rip["num-of-routes"]) != 2 { + t.Fatalf("RIP num-of-routes: expected 2, got %v", rip["num-of-routes"]) + } +} + +func TestRIPNeighbors(t *testing.T) { + out := routingCollect(t, fullRunner()) + ripProto := getRIPProtocol(t, out) + rip := ripProto["ietf-rip:rip"].(map[string]interface{}) + + ipv4 := rip["ipv4"].(map[string]interface{}) + neighContainer := ipv4["neighbors"].(map[string]interface{}) + neighs := neighContainer["neighbor"].([]interface{}) + if len(neighs) != 2 { + t.Fatalf("expected 2 RIP neighbors, got %d", len(neighs)) + } + + n0 := neighs[0].(map[string]interface{}) + if n0["ipv4-address"] != "10.0.0.2" { + t.Fatalf("RIP neighbor[0] address: expected 10.0.0.2, got %v", n0["ipv4-address"]) + } + // Bad packets/routes should be int (from text parse) + if toInt(n0["bad-packets-rcvd"]) != 0 { + t.Fatalf("RIP neighbor[0] bad-packets: expected 0, got %v", n0["bad-packets-rcvd"]) + } + + n1 := neighs[1].(map[string]interface{}) + if toInt(n1["bad-packets-rcvd"]) != 1 { + t.Fatalf("RIP neighbor[1] bad-packets: expected 1, got %v", n1["bad-packets-rcvd"]) + } + if toInt(n1["bad-routes-rcvd"]) != 2 { + t.Fatalf("RIP neighbor[1] bad-routes: expected 2, got %v", n1["bad-routes-rcvd"]) + } +} + +// --- BFD tests --- + +func getBFDProtocol(t *testing.T, out map[string]interface{}) map[string]interface{} { + t.Helper() + cpp := out["control-plane-protocols"].(map[string]interface{}) + for _, p := range cpp["control-plane-protocol"].([]interface{}) { + pm := p.(map[string]interface{}) + if pm["type"] == "infix-routing:bfdv1" { + return pm + } + } + t.Fatal("BFD protocol not found") + return nil +} + +func TestBFDSessions(t *testing.T) { + out := routingCollect(t, fullRunner()) + bfdProto := getBFDProtocol(t, out) + bfd := bfdProto["ietf-bfd:bfd"].(map[string]interface{}) + ipsh := bfd["ietf-bfd-ip-sh:ip-sh"].(map[string]interface{}) + sessionsContainer := ipsh["sessions"].(map[string]interface{}) + sessions := sessionsContainer["session"].([]interface{}) + + // Only single-hop sessions included (multihop=true is filtered) + if len(sessions) != 1 { + t.Fatalf("expected 1 BFD session (multihop filtered), got %d", len(sessions)) + } + + s := sessions[0].(map[string]interface{}) + if s["interface"] != "e0" { + t.Fatalf("BFD session interface: expected e0, got %v", s["interface"]) + } + if s["dest-addr"] != "10.0.0.2" { + t.Fatalf("BFD session dest-addr: expected 10.0.0.2, got %v", s["dest-addr"]) + } + if s["path-type"] != "ietf-bfd-types:path-ip-sh" { + t.Fatalf("BFD path-type: expected ietf-bfd-types:path-ip-sh, got %v", s["path-type"]) + } + + running := s["session-running"].(map[string]interface{}) + if running["local-state"] != "up" { + t.Fatalf("BFD local-state: expected up, got %v", running["local-state"]) + } + if running["detection-mode"] != "async-without-echo" { + t.Fatalf("BFD detection-mode: expected async-without-echo, got %v", running["detection-mode"]) + } + + // Intervals: ms → µs (×1000) + // receive-interval=300ms → 300000µs + if toInt(running["negotiated-rx-interval"]) != 300000 { + t.Fatalf("BFD rx-interval: expected 300000, got %v", running["negotiated-rx-interval"]) + } + if toInt(running["negotiated-tx-interval"]) != 300000 { + t.Fatalf("BFD tx-interval: expected 300000, got %v", running["negotiated-tx-interval"]) + } + // detection-time = detect-multiplier * receive-interval * 1000 = 3 * 300 * 1000 = 900000 + if toInt(running["detection-time"]) != 900000 { + t.Fatalf("BFD detection-time: expected 900000, got %v", running["detection-time"]) + } +} + +// --- Graceful degradation tests --- + +func TestRoutingCollectorOSPFOnly(t *testing.T) { + runner := ospfOnly() + + out := routingCollect(t, runner) + cpp := out["control-plane-protocols"].(map[string]interface{}) + protocols := cpp["control-plane-protocol"].([]interface{}) + if len(protocols) != 1 { + t.Fatalf("expected 1 protocol when only OSPF available, got %d", len(protocols)) + } + pm := protocols[0].(map[string]interface{}) + if pm["type"] != "infix-routing:ospfv2" { + t.Fatalf("expected OSPF protocol, got %v", pm["type"]) + } +} + +func TestRoutingCollectorAllFail(t *testing.T) { + out := routingCollect(t, fakeVty{}) + protocols := out["control-plane-protocols"].(map[string]interface{})["control-plane-protocol"].([]interface{}) + if len(protocols) != 0 { + t.Fatalf("expected an empty protocol list when nothing runs, got %v", protocols) + } +} + +// A protocol that stops must disappear from the tree, not linger with +// its last reported state. +func TestRoutingCollectorProtocolsDisappear(t *testing.T) { + tr := tree.New() + if err := newRoutingCollector(fullRunner()).Collect(context.Background(), tr); err != nil { + t.Fatalf("Collect failed: %v", err) + } + if err := newRoutingCollector(fakeVty{}).Collect(context.Background(), tr); err != nil { + t.Fatalf("Collect failed: %v", err) + } + + var out map[string]interface{} + if err := json.Unmarshal(tr.Get("ietf-routing:routing"), &out); err != nil { + t.Fatalf("unmarshal routing: %v", err) + } + protocols := out["control-plane-protocols"].(map[string]interface{})["control-plane-protocol"].([]interface{}) + if len(protocols) != 0 { + t.Fatalf("stopped protocols still reported: %v", protocols) + } +} + +func TestRIPStatusParsing(t *testing.T) { + status := parseRIPStatus(testRIPStatus) + + if status["update-interval"] != 30 { + t.Fatalf("update-interval: expected 30, got %v", status["update-interval"]) + } + if status["invalid-interval"] != 180 { + t.Fatalf("invalid-interval: expected 180, got %v", status["invalid-interval"]) + } + if status["flush-interval"] != 120 { + t.Fatalf("flush-interval: expected 120, got %v", status["flush-interval"]) + } + if status["default-metric"] != 1 { + t.Fatalf("default-metric: expected 1, got %v", status["default-metric"]) + } + if status["distance"] != 120 { + t.Fatalf("distance: expected 120, got %v", status["distance"]) + } + + ifaces := status["interfaces"].([]interface{}) + if len(ifaces) != 2 { + t.Fatalf("expected 2 parsed interfaces, got %d", len(ifaces)) + } + + neighs := status["neighbors"].([]interface{}) + if len(neighs) != 2 { + t.Fatalf("expected 2 parsed neighbors, got %d", len(neighs)) + } +} + +func TestFrrToIETFNeighborState(t *testing.T) { + tests := []struct { + input string + expected string + }{ + {"Full/DR", "full"}, + {"TwoWay/DROther", "2-way"}, + {"Init/DROther", "init"}, + {"Down/DROther", "down"}, + {"ExStart", "exstart"}, + } + for _, tt := range tests { + got := frrToIETFNeighborState(tt.input) + if got != tt.expected { + t.Fatalf("frrToIETFNeighborState(%q): expected %q, got %q", tt.input, tt.expected, got) + } + } +} + +func TestOSPFNetworkType(t *testing.T) { + tests := []struct { + nt string + p2mpNB bool + expected string + }{ + {"POINTOPOINT", false, "point-to-point"}, + {"BROADCAST", false, "broadcast"}, + {"POINTOMULTIPOINT", false, "hybrid"}, + {"POINTOMULTIPOINT", true, "point-to-multipoint"}, + {"NBMA", false, "non-broadcast"}, + {"UNKNOWN", false, ""}, + } + for _, tt := range tests { + got := ospfNetworkType(tt.nt, tt.p2mpNB) + if got != tt.expected { + t.Fatalf("ospfNetworkType(%q, %v): expected %q, got %q", tt.nt, tt.p2mpNB, tt.expected, got) + } + } +} + +func TestBFDMultihopFiltered(t *testing.T) { + // Ensure multihop peers don't appear in output + runner := fakeVty{ + "bfdd show bfd peers json": `[ + {"multihop": true, "peer": "10.0.0.99", "interface": "e1", "id": 5, "status": "up"} + ]`, + } + + c := newRoutingCollector(runner) + tr := tree.New() + c.Collect(context.Background(), tr) + + // BFD should not set anything when all peers are multihop + raw := tr.Get("ietf-routing:routing") + if raw != nil { + // If routing is set, BFD should not be present + var out map[string]interface{} + json.Unmarshal(raw, &out) + cpp := out["control-plane-protocols"].(map[string]interface{}) + protocols := cpp["control-plane-protocol"].([]interface{}) + for _, p := range protocols { + pm := p.(map[string]interface{}) + if pm["type"] == "infix-routing:bfdv1" { + t.Fatal("multihop-only BFD should not produce a protocol entry") + } + } + } +} + +// ripd prints RIPv1+2 as "1 2" in a fixed-width column, and FRR's +// default receive version is both. +func TestRIPStatusBothVersions(t *testing.T) { + text := ` Default version control: send version 2, receive version 1 2 + Interface Send Recv Key-chain + e0 2 1 2 + e1 1 2 1 2 mykeys + ethernet-long-01 1 2 + Routing for Networks: +` + ifaces := parseRIPStatus(text)["interfaces"].([]interface{}) + want := [][3]string{ + {"e0", "2", "1-2"}, + {"e1", "1-2", "1-2"}, + {"ethernet-long-01", "1", "2"}, + } + if len(ifaces) != len(want) { + t.Fatalf("parsed %d interfaces, want %d: %v", len(ifaces), len(want), ifaces) + } + for i, w := range want { + got := ifaces[i].(map[string]interface{}) + if got["name"] != w[0] || got["send-version"] != w[1] || got["recv-version"] != w[2] { + t.Errorf("row %d = %v, want %v", i, got, w) + } + } +} + +// ospfd appends " [Stub]" or " [NSSA]" to the area of interfaces and +// neighbors; the merge strips it and records the area type, and +// interfaces without OSPF are left out. +func TestOSPFStatusStubArea(t *testing.T) { + var ospf, ifaces, nbrs map[string]interface{} + json.Unmarshal([]byte(`{"areas":{"0.0.0.1":{}}}`), &ospf) + json.Unmarshal([]byte(`{"interfaces":{ + "e1":{"ospfEnabled":true,"area":"0.0.0.1 [Stub]"}, + "e2":{"ospfEnabled":false,"area":"0.0.0.1 [Stub]"}}}`), &ifaces) + json.Unmarshal([]byte(`{"neighbors":{ + "10.0.0.7":[{"ifaceName":"e1","areaId":"0.0.0.1 [Stub]"}], + "10.0.0.8":[{"ifaceName":"e2","areaId":"0.0.0.1 [Stub]"}]}}`), &nbrs) + + area := ospfStatus(ospf, ifaces, nbrs)["areas"].(map[string]interface{})["0.0.0.1"].(map[string]interface{}) + if area["area-type"] != "stub-area" { + t.Fatalf("area-type = %v, want stub-area", area["area-type"]) + } + list := area["interfaces"].([]interface{}) + if len(list) != 1 { + t.Fatalf("expected only the OSPF-enabled interface, got %v", list) + } + iface := list[0].(map[string]interface{}) + if iface["name"] != "e1" || iface["area"] != "0.0.0.1" { + t.Fatalf("interface = %v", iface) + } + peers := iface["neighbors"].([]interface{}) + if len(peers) != 1 || peers[0].(map[string]interface{})["neighborIp"] != "10.0.0.7" { + t.Fatalf("neighbors = %v", peers) + } +} + +// --- OSPFv3 and RIPng --- + +const testOSPF6Global = `{ + "routerId": "2.2.2.2", + "areas": { + "0.0.0.0": {}, + "0.0.0.1": {"areaIsStub": true} + } +}` + +// ospf6d keys interfaces by name at the top level and sets +// operatingAsType only when the OSPF network type differs from the link. +const testOSPF6Interfaces = `{ + "e6": { + "status": "up", + "type": "BROADCAST", + "attachedToArea": true, + "areaId": "0.0.0.0", + "cost": 10, + "ospf6InterfaceState": "DR", + "transmitDelaySec": 1, + "priority": 1, + "timerIntervalsConfigHello": 10, + "timerIntervalsConfigDead": 40, + "timerIntervalsConfigRetransmit": 5, + "timerPassiveIface": false + }, + "e7": { + "status": "up", + "type": "BROADCAST", + "operatingAsType": "POINTOMULTIPOINT", + "attachedToArea": true, + "areaId": "0.0.0.1", + "cost": 20, + "ospf6InterfaceState": "PtMultipoint", + "priority": 1, + "timerPassiveIface": true + }, + "e8": { + "status": "up", + "type": "BROADCAST", + "attachedToArea": false + } +}` + +const testOSPF6Neighbors = `{ + "neighbors": [ + { + "neighborId": "1.1.1.1", + "priority": 1, + "deadTime": "00:00:38", + "state": "Full", + "ifState": "BDR", + "duration": "00:01:08", + "interfaceName": "e6", + "interfaceState": "DR" + }, + { + "neighborId": "3.3.3.3", + "priority": 1, + "deadTime": "00:00:31", + "state": "Twoway", + "ifState": "DROther", + "duration": "00:00:12", + "interfaceName": "e6", + "interfaceState": "DR" + }, + { + "neighborId": "4.4.4.4", + "priority": 1, + "deadTime": "00:00:35", + "state": "Full", + "ifState": "PointToPoint", + "duration": "00:03:02", + "interfaceName": "e7", + "interfaceState": "PtMultipoint" + } + ] +}` + +const testOSPF6Routes = `{ + "routes": { + "2001:db8:1::1/128": [ + { + "isBestRoute": true, + "destinationType": "N", + "pathType": "IA", + "nextHops": [ + {"nextHop": "fe80::2a0:85ff:fe00:301", "interfaceName": "e6"} + ] + } + ], + "2001:db8:33::/64": [ + { + "isBestRoute": false, + "destinationType": "N", + "pathType": "E2", + "nextHops": [ + {"nextHop": "fe80::2a0:85ff:fe00:401", "interfaceName": "e7"} + ] + }, + { + "isBestRoute": true, + "destinationType": "N", + "pathType": "E1", + "nextHops": [ + {"nextHop": "::", "interfaceName": "e7"} + ] + } + ] + } +}` + +// ripngd prints the peer table as two lines per peer and has no +// Distance line. +const testRIPNGStatus = `Routing Protocol is "RIPng" + Sending updates every 30 seconds with +/-50%, next due in 20 seconds + Timeout after 180 seconds, garbage collect after 120 seconds + Outgoing update filter list for all interface is not set + Incoming update filter list for all interface is not set + Default redistribution metric is 1 + Redistributing: + Default version control: send version 1, receive version 1 + Interface Send Recv + e2 1 1 + e7 1 1 + Routing for Networks: + e2 + e7 + Routing Information Sources: + Gateway BadPackets BadRoutes Distance Last Update + fe80::2a0:85ff:fe00:306 + 0 0 120 00:00:12 +` + +const testRIPNGRoutes = `{ + "2001:db8:60::/64": [ + { + "prefix": "2001:db8:60::/64", + "protocol": "ripng", + "metric": 2, + "nexthops": [ + {"ip": "fe80::2a0:85ff:fe00:306", "interfaceName": "e7"} + ] + } + ] +}` + +func ipv6Runner() fakeVty { + return fakeVty{ + "ospf6d show ipv6 ospf6 json": testOSPF6Global, + "ospf6d show ipv6 ospf6 interface json": testOSPF6Interfaces, + "ospf6d show ipv6 ospf6 neighbor json": testOSPF6Neighbors, + "ospf6d show ipv6 ospf6 route json": testOSPF6Routes, + "ripngd show ipv6 ripng status": testRIPNGStatus, + "zebra show ipv6 route ripng json": testRIPNGRoutes, + } +} + +func getProtocol(t *testing.T, out map[string]interface{}, typ string) map[string]interface{} { + t.Helper() + cpp := out["control-plane-protocols"].(map[string]interface{}) + for _, p := range cpp["control-plane-protocol"].([]interface{}) { + pm := p.(map[string]interface{}) + if pm["type"] == typ { + return pm + } + } + t.Fatalf("protocol %s not found", typ) + return nil +} + +func TestRoutingCollectorIPv6Protocols(t *testing.T) { + out := routingCollect(t, ipv6Runner()) + cpp := out["control-plane-protocols"].(map[string]interface{}) + if n := len(cpp["control-plane-protocol"].([]interface{})); n != 2 { + t.Fatalf("expected ospfv3 and ripng only, got %d protocols", n) + } + getProtocol(t, out, "infix-routing:ospfv3") + getProtocol(t, out, "infix-routing:ripng") +} + +func TestOSPF6AreasInterfacesNeighbors(t *testing.T) { + out := routingCollect(t, ipv6Runner()) + ospf := getProtocol(t, out, "infix-routing:ospfv3")["ietf-ospf:ospf"].(map[string]interface{}) + + if ospf["ietf-ospf:router-id"] != "2.2.2.2" { + t.Fatalf("router-id = %v", ospf["ietf-ospf:router-id"]) + } + areas := ospf["ietf-ospf:areas"].(map[string]interface{})["ietf-ospf:area"].([]interface{}) + if len(areas) != 2 { + t.Fatalf("expected 2 areas, got %d", len(areas)) + } + + backbone := areas[0].(map[string]interface{}) + if backbone["ietf-ospf:area-id"] != "0.0.0.0" || backbone["ietf-ospf:area-type"] != "normal-area" { + t.Fatalf("backbone = %v", backbone) + } + ifaces := backbone["ietf-ospf:interfaces"].(map[string]interface{})["ietf-ospf:interface"].([]interface{}) + if len(ifaces) != 1 { + t.Fatalf("expected 1 backbone interface, got %d", len(ifaces)) + } + e6 := ifaces[0].(map[string]interface{}) + if e6["name"] != "e6" || e6["state"] != "dr" || e6["interface-type"] != "broadcast" || e6["passive"] != false { + t.Fatalf("e6 = %v", e6) + } + if toInt(e6["cost"]) != 10 || toInt(e6["hello-interval"]) != 10 || toInt(e6["dead-interval"]) != 40 { + t.Fatalf("e6 timers = %v", e6) + } + nbrs := e6["ietf-ospf:neighbors"].(map[string]interface{})["ietf-ospf:neighbor"].([]interface{}) + if len(nbrs) != 2 { + t.Fatalf("expected 2 neighbors on e6, got %d", len(nbrs)) + } + full := nbrs[0].(map[string]interface{}) + if full["neighbor-router-id"] != "1.1.1.1" || full["state"] != "full" || full["infix-routing:role"] != "BDR" { + t.Fatalf("neighbor = %v", full) + } + if two := nbrs[1].(map[string]interface{}); two["state"] != "2-way" { + t.Fatalf("ospf6d Twoway not mapped: %v", two) + } + + stub := areas[1].(map[string]interface{}) + if stub["ietf-ospf:area-type"] != "stub-area" { + t.Fatalf("stub = %v", stub) + } + e7 := stub["ietf-ospf:interfaces"].(map[string]interface{})["ietf-ospf:interface"].([]interface{})[0].(map[string]interface{}) + if e7["interface-type"] != "hybrid" || e7["passive"] != true { + t.Fatalf("e7 = %v", e7) + } + if _, ok := e7["state"]; ok { + t.Fatalf("PtMultipoint has no ietf state, got %v", e7["state"]) + } + p2p := e7["ietf-ospf:neighbors"].(map[string]interface{})["ietf-ospf:neighbor"].([]interface{})[0].(map[string]interface{}) + if p2p["neighbor-router-id"] != "4.4.4.4" || p2p["infix-routing:role"] != nil { + t.Fatalf("ifState on a non-broadcast link is not a role: %v", p2p) + } +} + +func TestOSPF6Routes(t *testing.T) { + out := routingCollect(t, ipv6Runner()) + ospf := getProtocol(t, out, "infix-routing:ospfv3")["ietf-ospf:ospf"].(map[string]interface{}) + routes := ospf["ietf-ospf:local-rib"].(map[string]interface{})["ietf-ospf:route"].([]interface{}) + if len(routes) != 2 { + t.Fatalf("expected one route per prefix, got %d", len(routes)) + } + + byPrefix := map[string]map[string]interface{}{} + for _, r := range routes { + rm := r.(map[string]interface{}) + byPrefix[rm["prefix"].(string)] = rm + } + + host := byPrefix["2001:db8:1::1/128"] + if host["route-type"] != "intra-area" { + t.Fatalf("host route = %v", host) + } + hop := host["next-hops"].(map[string]interface{})["next-hop"].([]interface{})[0].(map[string]interface{}) + if hop["next-hop"] != "fe80::2a0:85ff:fe00:301" { + t.Fatalf("host next-hop = %v", hop) + } + + ext := byPrefix["2001:db8:33::/64"] + if ext["route-type"] != "external-1" { + t.Fatalf("best path not chosen: %v", ext) + } + hop = ext["next-hops"].(map[string]interface{})["next-hop"].([]interface{})[0].(map[string]interface{}) + if hop["outgoing-interface"] != "e7" || hop["next-hop"] != nil { + t.Fatalf("connected next-hop = %v", hop) + } +} + +func TestRIPNG(t *testing.T) { + out := routingCollect(t, ipv6Runner()) + rip := getProtocol(t, out, "infix-routing:ripng")["ietf-rip:rip"].(map[string]interface{}) + + timers := rip["timers"].(map[string]interface{}) + if toInt(timers["update-interval"]) != 30 || toInt(timers["flush-interval"]) != 120 { + t.Fatalf("timers = %v", timers) + } + if _, ok := rip["distance"]; ok { + t.Fatalf("ripngd prints no distance, got %v", rip["distance"]) + } + + ifaces := rip["interfaces"].(map[string]interface{})["interface"].([]interface{}) + if len(ifaces) != 2 { + t.Fatalf("expected 2 interfaces, got %d", len(ifaces)) + } + e2 := ifaces[0].(map[string]interface{}) + if e2["interface"] != "e2" || e2["oper-status"] != "up" { + t.Fatalf("e2 = %v", e2) + } + if _, ok := e2["send-version"]; ok { + t.Fatalf("RIPng has no protocol version, got %v", e2) + } + + ipv6 := rip["ipv6"].(map[string]interface{}) + routes := ipv6["routes"].(map[string]interface{})["route"].([]interface{}) + if len(routes) != 1 || toInt(rip["num-of-routes"]) != 1 { + t.Fatalf("routes = %v", routes) + } + route := routes[0].(map[string]interface{}) + if route["ipv6-prefix"] != "2001:db8:60::/64" || route["next-hop"] != "fe80::2a0:85ff:fe00:306" || route["interface"] != "e7" { + t.Fatalf("route = %v", route) + } + + nbrs := ipv6["neighbors"].(map[string]interface{})["neighbor"].([]interface{}) + if len(nbrs) != 1 { + t.Fatalf("expected 1 neighbor, got %d", len(nbrs)) + } + nbr := nbrs[0].(map[string]interface{}) + if nbr["ipv6-address"] != "fe80::2a0:85ff:fe00:306" || toInt(nbr["bad-packets-rcvd"]) != 0 { + t.Fatalf("neighbor = %v", nbr) + } +} diff --git a/src/yangerd/internal/collector/runner.go b/src/yangerd/internal/collector/runner.go new file mode 100644 index 000000000..d2c511484 --- /dev/null +++ b/src/yangerd/internal/collector/runner.go @@ -0,0 +1,99 @@ +package collector + +import ( + "context" + "encoding/json" + "fmt" + "os" + "os/exec" + "path/filepath" + + "github.com/godbus/dbus/v5" +) + +// CommandRunner executes external commands and returns their stdout. +type CommandRunner interface { + Run(ctx context.Context, name string, args ...string) ([]byte, error) +} + +// FileReader reads files and globs paths on the filesystem. +type FileReader interface { + ReadFile(path string) ([]byte, error) + Glob(pattern string) ([]string, error) +} + +// InstallerStatus queries RAUC installation progress. +type InstallerStatus interface { + GetInstallStatus() (operation string, lastError string, percentage int, message string, err error) +} + +// runJSON runs a command and decodes its JSON output into dst. +func runJSON(ctx context.Context, cmd CommandRunner, dst interface{}, name string, args ...string) error { + out, err := cmd.Run(ctx, name, args...) + if err != nil { + return fmt.Errorf("%s: %w", name, err) + } + if err := json.Unmarshal(out, dst); err != nil { + return fmt.Errorf("%s: parse output: %w", name, err) + } + return nil +} + +// ExecRunner is the production CommandRunner using os/exec. +type ExecRunner struct{} + +func (ExecRunner) Run(ctx context.Context, name string, args ...string) ([]byte, error) { + return exec.CommandContext(ctx, name, args...).Output() +} + +// OSFileReader is the production FileReader using the os package. +type OSFileReader struct{} + +func (OSFileReader) ReadFile(path string) ([]byte, error) { + return os.ReadFile(path) +} + +func (OSFileReader) Glob(pattern string) ([]string, error) { + return filepath.Glob(pattern) +} + +// DBusInstaller reads RAUC installation status from D-Bus properties. +// It runs on every system-state GET, so it uses the process-wide shared +// bus connection, which godbus re-establishes if it drops, instead of +// paying a connect and auth handshake per query. +type DBusInstaller struct{} + +func (DBusInstaller) GetInstallStatus() (string, string, int, string, error) { + conn, err := dbus.SystemBus() + if err != nil { + return "", "", 0, "", err + } + + obj := conn.Object("de.pengutronix.rauc", "/") + + operation, _ := obj.GetProperty("de.pengutronix.rauc.Installer.Operation") + lastError, _ := obj.GetProperty("de.pengutronix.rauc.Installer.LastError") + + var pct int + var msg string + progress, err := obj.GetProperty("de.pengutronix.rauc.Installer.Progress") + if err == nil { + if vals, ok := progress.Value().([]interface{}); ok && len(vals) >= 2 { + if p, ok := vals[0].(int32); ok { + pct = int(p) + } + if s, ok := vals[1].(string); ok { + msg = s + } + } + } + + return variantString(operation), variantString(lastError), pct, msg, nil +} + +func variantString(v dbus.Variant) string { + if s, ok := v.Value().(string); ok { + return s + } + return "" +} diff --git a/src/yangerd/internal/collector/system.go b/src/yangerd/internal/collector/system.go new file mode 100644 index 000000000..e8c66cbb4 --- /dev/null +++ b/src/yangerd/internal/collector/system.go @@ -0,0 +1,108 @@ +package collector + +import ( + "context" + "encoding/json" + "strconv" + "time" + + "github.com/kernelkit/infix/src/yangerd/internal/numconv" + "github.com/kernelkit/infix/src/yangerd/internal/tree" +) + +var platformKeyMap = map[string]string{ + "NAME": "os-name", + "VERSION_ID": "os-version", + "BUILD_ID": "os-release", + "ARCHITECTURE": "machine", +} + +// SystemCollector gathers ietf-system operational data. +type SystemCollector struct { + cmd CommandRunner + fs FileReader + interval time.Duration +} + +// NewSystemCollector creates a SystemCollector with the given dependencies. +func NewSystemCollector(cmd CommandRunner, fs FileReader, interval time.Duration) *SystemCollector { + return &SystemCollector{cmd: cmd, fs: fs, interval: interval} +} + +// Name implements Collector. +func (c *SystemCollector) Name() string { return "system" } + +// Interval implements Collector. +func (c *SystemCollector) Interval() time.Duration { return c.interval } + +// Collect implements Collector. It merges service data into +// "ietf-system:system-state". DNS is handled reactively by +// fswatcher on /var/lib/misc/resolv.conf. Other system-state +// subtrees (platform, software, users, hostname, timezone, clock, +// memory, load, filesystems) are populated by boot-once, reactive, +// or on-demand providers. +func (c *SystemCollector) Collect(ctx context.Context, t *tree.Tree) error { + state := make(map[string]interface{}) + + if err := c.addServices(ctx, state); err != nil { + return err + } + + data, err := json.Marshal(state) + if err != nil { + return err + } + t.Merge("ietf-system:system-state", data) + return nil +} + +func (c *SystemCollector) addServices(ctx context.Context, state map[string]interface{}) error { + var initData []map[string]interface{} + if err := runJSON(ctx, c.cmd, &initData, "initctl", "-j"); err != nil { + return err + } + + var services []interface{} + for _, d := range initData { + pid, ok := d["pid"] + if !ok { + continue + } + identity, ok := d["identity"] + if !ok { + continue + } + svc := map[string]interface{}{ + "pid": toInt(pid), + "name": identity, + "status": d["status"], + "description": d["description"], + "statistics": map[string]interface{}{ + "memory-usage": strconv.Itoa(toInt(zeroIfNil(d["memory"]))), + "uptime": strconv.Itoa(toInt(zeroIfNil(d["uptime"]))), + "restart-count": toInt(zeroIfNil(d["restarts"])), + }, + } + services = append(services, svc) + } + + state["infix-system:services"] = map[string]interface{}{ + "service": services, + } + return nil +} + +func yangDateTime(t time.Time) string { + return t.Format("2006-01-02T15:04:05-07:00") +} + +func toInt(v interface{}) int { + return numconv.IntOrZero(v) +} + +func zeroIfNil(v interface{}) interface{} { + if v == nil { + return 0 + } + return v +} diff --git a/src/yangerd/internal/collector/system_test.go b/src/yangerd/internal/collector/system_test.go new file mode 100644 index 000000000..a69c7d699 --- /dev/null +++ b/src/yangerd/internal/collector/system_test.go @@ -0,0 +1,194 @@ +package collector + +import ( + "context" + "encoding/json" + "fmt" + "testing" + "time" + + "github.com/kernelkit/infix/src/yangerd/internal/testutil" + "github.com/kernelkit/infix/src/yangerd/internal/tree" +) + +const ( + testInitctlJSON = `[ + { + "identity": "sshd", + "pid": 123, + "status": "running", + "description": "OpenSSH daemon", + "memory": 4096000, + "uptime": 3600, + "restarts": 2 + }, + { + "identity": "sysklogd", + "pid": 456, + "status": "running", + "description": "System logger", + "memory": 2048000, + "uptime": 7200, + "restarts": 0 + } +]` +) + +func newTestCollector() (*SystemCollector, *testutil.MockRunner, *testutil.MockFileReader) { + runner := &testutil.MockRunner{ + Results: map[string][]byte{ + "initctl -j": []byte(testInitctlJSON), + }, + Errors: map[string]error{}, + } + + fs := &testutil.MockFileReader{ + Files: map[string][]byte{}, + Globs: map[string][]string{}, + } + + c := NewSystemCollector(runner, fs, 60*time.Second) + return c, runner, fs +} + +func collectToState(t *testing.T, c *SystemCollector) map[string]interface{} { + t.Helper() + tr := tree.New() + if err := c.Collect(context.Background(), tr); err != nil { + t.Fatalf("Collect failed: %v", err) + } + + stateRaw := tr.Get("ietf-system:system-state") + if stateRaw == nil { + t.Fatal("missing ietf-system:system-state in tree") + } + + state := make(map[string]interface{}) + if err := json.Unmarshal(stateRaw, &state); err != nil { + t.Fatalf("unmarshal system-state: %v", err) + } + return state +} + +func TestSystemCollectorName(t *testing.T) { + c, _, _ := newTestCollector() + if c.Name() != "system" { + t.Fatalf("expected name 'system', got %q", c.Name()) + } +} + +func TestSystemCollectorInterval(t *testing.T) { + c, _, _ := newTestCollector() + if c.Interval() != 60*time.Second { + t.Fatalf("expected interval 60s, got %v", c.Interval()) + } +} + +func TestSystemCollectorServices(t *testing.T) { + c, _, _ := newTestCollector() + state := collectToState(t, c) + + svcs, ok := state["infix-system:services"].(map[string]interface{}) + if !ok { + t.Fatal("missing infix-system:services in system-state") + } + + serviceList, ok := svcs["service"].([]interface{}) + if !ok || len(serviceList) != 2 { + t.Fatalf("expected 2 services, got %v", svcs["service"]) + } + + svc0 := serviceList[0].(map[string]interface{}) + if svc0["name"] != "sshd" { + t.Fatalf("service[0] name: expected sshd, got %v", svc0["name"]) + } + if int(svc0["pid"].(float64)) != 123 { + t.Fatalf("service[0] pid: expected 123, got %v", svc0["pid"]) + } + + stats := svc0["statistics"].(map[string]interface{}) + if stats["memory-usage"] != "4096000" { + t.Fatalf("service[0] memory-usage: expected '4096000', got %v", stats["memory-usage"]) + } + if stats["uptime"] != "3600" { + t.Fatalf("service[0] uptime: expected '3600', got %v", stats["uptime"]) + } + if int(stats["restart-count"].(float64)) != 2 { + t.Fatalf("service[0] restart-count: expected 2, got %v", stats["restart-count"]) + } +} + +func TestSystemCollectorInitctlFailure(t *testing.T) { + runner := &testutil.MockRunner{ + Results: map[string][]byte{}, + Errors: map[string]error{ + "initctl -j": fmt.Errorf("not available"), + }, + } + + fs := &testutil.MockFileReader{ + Files: map[string][]byte{}, + Globs: map[string][]string{}, + } + + c := NewSystemCollector(runner, fs, 60*time.Second) + tr := tree.New() + tr.Set("ietf-system:system-state", json.RawMessage(`{"infix-system:services":{"service":[{"name":"sshd"}]}}`)) + + if err := c.Collect(context.Background(), tr); err == nil { + t.Fatal("an initctl failure must be reported, it is all this collector does") + } + if got := string(tr.Get("ietf-system:system-state")); got != `{"infix-system:services":{"service":[{"name":"sshd"}]}}` { + t.Fatalf("a failed poll must leave the last services in place, got %s", got) + } +} + +func TestSystemCollectorTreeKeys(t *testing.T) { + c, _, _ := newTestCollector() + tr := tree.New() + c.Collect(context.Background(), tr) + + keys := tr.Keys() + if len(keys) != 1 { + t.Fatalf("expected exactly 1 tree key, got %d: %v", len(keys), keys) + } + if keys[0] != "ietf-system:system-state" { + t.Fatalf("expected tree key 'ietf-system:system-state', got %q", keys[0]) + } +} + +func TestSystemCollectorServicesNilFields(t *testing.T) { + runner := &testutil.MockRunner{ + Results: map[string][]byte{ + "initctl -j": []byte(`[{"identity":"minimal","pid":999,"status":"running","description":"Minimal service"}]`), + }, + Errors: map[string]error{}, + } + + fs := &testutil.MockFileReader{ + Files: map[string][]byte{}, + Globs: map[string][]string{}, + } + + c := NewSystemCollector(runner, fs, 60*time.Second) + state := collectToState(t, c) + + svcs := state["infix-system:services"].(map[string]interface{}) + serviceList := svcs["service"].([]interface{}) + if len(serviceList) != 1 { + t.Fatalf("expected 1 service, got %d", len(serviceList)) + } + + svc := serviceList[0].(map[string]interface{}) + stats := svc["statistics"].(map[string]interface{}) + + if stats["memory-usage"] != "0" { + t.Fatalf("nil memory should become '0', got %v", stats["memory-usage"]) + } + if stats["uptime"] != "0" { + t.Fatalf("nil uptime should become '0', got %v", stats["uptime"]) + } + if int(stats["restart-count"].(float64)) != 0 { + t.Fatalf("nil restarts should become 0, got %v", stats["restart-count"]) + } +} diff --git a/src/yangerd/internal/config/config.go b/src/yangerd/internal/config/config.go new file mode 100644 index 000000000..7178480e8 --- /dev/null +++ b/src/yangerd/internal/config/config.go @@ -0,0 +1,77 @@ +package config + +import ( + "os" + "strconv" + "time" +) + +// Config holds all yangerd runtime configuration, populated from +// environment variables with sensible defaults. +type Config struct { + Socket string + LogLevel string + PollSystem time.Duration + PollRouting time.Duration + PollNTP time.Duration + PollHardware time.Duration + PollSTP time.Duration + EnableWifi bool + EnableLLDP bool + EnableFirewall bool + EnableDHCP bool + EnableContainers bool + EnableGPS bool + EnableFRR bool +} + +// Load reads configuration from the environment. +func Load() *Config { + return &Config{ + Socket: envStr("YANGERD_SOCKET", "/run/yangerd.sock"), + LogLevel: envStr("YANGERD_LOG_LEVEL", "info"), + PollSystem: envDur("YANGERD_POLL_INTERVAL_SYSTEM", 60*time.Second), + PollRouting: envDur("YANGERD_POLL_INTERVAL_ROUTING", 10*time.Second), + PollNTP: envDur("YANGERD_POLL_INTERVAL_NTP", 60*time.Second), + PollHardware: envDur("YANGERD_POLL_INTERVAL_HARDWARE", 10*time.Second), + PollSTP: envDur("YANGERD_POLL_INTERVAL_STP", 5*time.Second), + EnableWifi: envBool("YANGERD_ENABLE_WIFI", false), + EnableLLDP: envBool("YANGERD_ENABLE_LLDP", true), + EnableFirewall: envBool("YANGERD_ENABLE_FIREWALL", true), + EnableDHCP: envBool("YANGERD_ENABLE_DHCP", true), + EnableContainers: envBool("YANGERD_ENABLE_CONTAINERS", false), + EnableGPS: envBool("YANGERD_ENABLE_GPS", false), + EnableFRR: envBool("YANGERD_ENABLE_FRR", true), + } +} + +func envStr(key, def string) string { + if v := os.Getenv(key); v != "" { + return v + } + return def +} + +func envBool(key string, def bool) bool { + v := os.Getenv(key) + if v == "" { + return def + } + b, err := strconv.ParseBool(v) + if err != nil { + return def + } + return b +} + +func envDur(key string, def time.Duration) time.Duration { + v := os.Getenv(key) + if v == "" { + return def + } + d, err := time.ParseDuration(v) + if err != nil { + return def + } + return d +} diff --git a/src/yangerd/internal/containermonitor/containermonitor.go b/src/yangerd/internal/containermonitor/containermonitor.go new file mode 100644 index 000000000..6b5dfabeb --- /dev/null +++ b/src/yangerd/internal/containermonitor/containermonitor.go @@ -0,0 +1,172 @@ +// Package containermonitor keeps the infix-containers subtree in the tree +// in sync with podman. A persistent `podman events` subprocess is used +// purely as a change trigger; on every event the full container table is +// re-read with `podman ps` (via collector.CollectContainers) and the +// subtree replaced, so removed containers disappear and containers present +// before yangerd started are picked up. +// +// This replaces an earlier inotify watch on /run/libpod/events, which was +// reactive-only and silently went stale whenever an event was missed +// (debounce coalescing, inotify overflow, a removal racing the re-read, or +// yangerd starting after the container). `podman events` reads whichever +// events backend podman is configured for (file or journald), so it does +// not depend on a specific on-disk layout. +package containermonitor + +import ( + "bufio" + "context" + "encoding/json" + "fmt" + "io" + "log/slog" + "os/exec" + "time" + + "github.com/kernelkit/infix/src/yangerd/internal/backoff" + "github.com/kernelkit/infix/src/yangerd/internal/collector" + "github.com/kernelkit/infix/src/yangerd/internal/tree" +) + +const ( + treeKey = "infix-containers:containers" + + // debounceDelay coalesces bursts of events into one re-read. It is + // deliberately generous: container lifecycle events fire while confd + // is still running its own `podman` start/stop/rm operations, so + // re-reading too eagerly makes yangerd's `podman ps/inspect/stats` + // contend with confd for the libpod lock on a CPU-starved guest. + // Waiting for the churn to settle keeps yangerd off confd's back + // during config apply/reset; a couple of seconds of staleness in + // operational data is harmless. + debounceDelay = 2 * time.Second +) + +// ContainerMonitor subscribes to container lifecycle events via a +// persistent `podman events` subprocess and re-reads the full container +// table on every event. +type ContainerMonitor struct { + tree *tree.Tree + log *slog.Logger + refresh chan struct{} + + // collect returns the current container subtree, or nil when there are + // no containers; overridable in tests. + collect func() json.RawMessage +} + +// New creates a ContainerMonitor. +func New(t *tree.Tree, cmd collector.CommandRunner, fs collector.FileReader, log *slog.Logger) *ContainerMonitor { + if log == nil { + log = slog.Default() + } + return &ContainerMonitor{ + tree: t, + log: log, + refresh: make(chan struct{}, 1), + collect: func() json.RawMessage { return collector.CollectContainers(cmd, fs) }, + } +} + +// Run starts the container monitor. It blocks until ctx is cancelled, +// restarting the events subprocess with backoff if it exits. +func (m *ContainerMonitor) Run(ctx context.Context) error { + go m.refreshLoop(ctx) + return backoff.Retry(ctx, m.log, "container monitor", m.runOnce) +} + +func (m *ContainerMonitor) runOnce(ctx context.Context) error { + cmd := exec.CommandContext(ctx, "podman", "events", "--filter", "type=container", "--format", "json") + stdout, err := cmd.StdoutPipe() + if err != nil { + return fmt.Errorf("stdout pipe: %w", err) + } + if err := cmd.Start(); err != nil { + return fmt.Errorf("start podman events: %w", err) + } + defer cmd.Wait() + + // Pick up containers that existed before we attached. + m.triggerRefresh() + + return m.readEvents(stdout) +} + +// readEvents consumes the newline-delimited JSON event stream. Each event +// is only a trigger; the payload is never used to build state. +func (m *ContainerMonitor) readEvents(r io.Reader) error { + scanner := bufio.NewScanner(r) + scanner.Buffer(make([]byte, 0, 64*1024), 1*1024*1024) + + for scanner.Scan() { + line := scanner.Bytes() + if len(line) == 0 { + continue + } + if status := eventStatus(line); status != "" { + m.log.Debug("container monitor: event", "status", status) + } + m.triggerRefresh() + } + if err := scanner.Err(); err != nil { + return fmt.Errorf("read podman events: %w", err) + } + return fmt.Errorf("podman events process exited") +} + +// eventStatus extracts the event status for logging; best-effort only. +func eventStatus(line []byte) string { + var ev struct { + Status string `json:"Status"` + } + if json.Unmarshal(line, &ev) != nil { + return "" + } + return ev.Status +} + +// triggerRefresh requests a table re-read; the buffered channel collapses +// pending requests into one. +func (m *ContainerMonitor) triggerRefresh() { + select { + case m.refresh <- struct{}{}: + default: + } +} + +func (m *ContainerMonitor) refreshLoop(ctx context.Context) { + for { + select { + case <-ctx.Done(): + return + case <-m.refresh: + } + + // Let a burst of events settle before reading. + select { + case <-ctx.Done(): + return + case <-time.After(debounceDelay): + } + select { + case <-m.refresh: + default: + } + + m.updateTree() + } +} + +// updateTree re-reads the full container table and replaces the subtree. +// With no containers the key is deleted rather than left as an empty node, +// so an idle-but-enabled container feature reads as absent. +func (m *ContainerMonitor) updateTree() { + data := m.collect() + if len(data) == 0 { + m.tree.Delete(treeKey) + m.log.Debug("container monitor: no containers, key removed") + return + } + m.tree.Set(treeKey, data) + m.log.Debug("container monitor: tree updated") +} diff --git a/src/yangerd/internal/containermonitor/containermonitor_test.go b/src/yangerd/internal/containermonitor/containermonitor_test.go new file mode 100644 index 000000000..d70c0f60e --- /dev/null +++ b/src/yangerd/internal/containermonitor/containermonitor_test.go @@ -0,0 +1,89 @@ +package containermonitor + +import ( + "context" + "encoding/json" + "strings" + "testing" + "time" + + "github.com/kernelkit/infix/src/yangerd/internal/tree" +) + +// newTestMonitor builds a monitor whose collect() is driven by the test. +// cmd/fs are nil since collect is overridden, so the default closure that +// would use them is never called. +func newTestMonitor(t *testing.T, collect func() json.RawMessage) (*ContainerMonitor, *tree.Tree) { + t.Helper() + tr := tree.New() + m := New(tr, nil, nil, nil) + m.collect = collect + return m, tr +} + +func TestUpdateTreeSetsContainers(t *testing.T) { + m, tr := newTestMonitor(t, func() json.RawMessage { + return json.RawMessage(`{"container":[{"name":"web"}]}`) + }) + + m.updateTree() + + got := tr.Get(treeKey) + if got == nil || !strings.Contains(string(got), "web") { + t.Fatalf("expected container data, got %s", got) + } +} + +// With no containers the key must be deleted, not left as an empty node, +// so an idle-but-enabled container feature reads as absent. +func TestUpdateTreeDeletesWhenEmpty(t *testing.T) { + m, tr := newTestMonitor(t, func() json.RawMessage { return nil }) + + tr.Set(treeKey, json.RawMessage(`{"container":[{"name":"old"}]}`)) + m.updateTree() + + if got := tr.Get(treeKey); got != nil { + t.Fatalf("expected key removed when no containers, got %s", got) + } +} + +// An event in the stream must trigger a re-read; here the re-read clears a +// previously-present container, proving the stream drives reconciliation. +func TestEventTriggersRefresh(t *testing.T) { + calls := 0 + m, tr := newTestMonitor(t, func() json.RawMessage { + calls++ + return nil // container is gone + }) + tr.Set(treeKey, json.RawMessage(`{"container":[{"name":"gone"}]}`)) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + go m.refreshLoop(ctx) + + // A container "died" event, newline-framed as podman emits it. + go m.readEvents(strings.NewReader(`{"Type":"container","Status":"died","Name":"gone"}` + "\n")) + + // Must comfortably exceed debounceDelay, or this races the re-read. + deadline := time.After(debounceDelay + 3*time.Second) + for { + if tr.Get(treeKey) == nil && calls > 0 { + break + } + select { + case <-deadline: + t.Fatalf("event did not trigger reconcile; calls=%d tree=%s", calls, tr.Get(treeKey)) + default: + time.Sleep(10 * time.Millisecond) + } + } +} + +func TestEventStatus(t *testing.T) { + if s := eventStatus([]byte(`{"Status":"start"}`)); s != "start" { + t.Errorf("eventStatus = %q, want start", s) + } + if s := eventStatus([]byte(`not json`)); s != "" { + t.Errorf("eventStatus on garbage = %q, want empty", s) + } +} diff --git a/src/yangerd/internal/dbusmonitor/dbusmonitor.go b/src/yangerd/internal/dbusmonitor/dbusmonitor.go new file mode 100644 index 000000000..36011c296 --- /dev/null +++ b/src/yangerd/internal/dbusmonitor/dbusmonitor.go @@ -0,0 +1,1327 @@ +// Package dbusmonitor watches D-Bus signals from dnsmasq and firewalld +// and keeps their operational YANG subtrees updated. +package dbusmonitor + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io/fs" + "log/slog" + "net/netip" + "os" + "os/exec" + "path/filepath" + "strconv" + "strings" + "sync" + "time" + + "github.com/godbus/dbus/v5" + "github.com/kernelkit/infix/src/yangerd/internal/backoff" + "github.com/kernelkit/infix/src/yangerd/internal/collector" + "github.com/kernelkit/infix/src/yangerd/internal/numconv" + "github.com/kernelkit/infix/src/yangerd/internal/tree" +) + +const ( + dnsmasqBusName = "uk.org.thekelleys.dnsmasq" + dnsmasqInterface = "uk.org.thekelleys.dnsmasq" + dnsmasqPath = "/uk/org/thekelleys/dnsmasq" + + firewalldBusName = "org.fedoraproject.FirewallD1" + firewalldInterface = "org.fedoraproject.FirewallD1" + firewalldPath = "/org/fedoraproject/FirewallD1" + + dbusInterface = "org.freedesktop.DBus" + dbusPath = "/org/freedesktop/DBus" + + raucBusName = "de.pengutronix.rauc" + raucInstallerInterface = "de.pengutronix.rauc.Installer" + + dnsmasqLeaseFile = "/var/lib/misc/dnsmasq.leases" + + // Written by confd's address-set add/remove actions: one file per + // set, listing the entries that are dynamic (not in the config). + addrsetShadowDir = "/run/confd/address-sets" + + dhcpTreeKey = "infix-dhcp-server:dhcp-server" + firewallTreeKey = "infix-firewall:firewall" + systemStateKey = "ietf-system:system-state" + + // softwareTimeout bounds the rauc status and bootloader env reads + // run after an install completes. + softwareTimeout = 30 * time.Second + + // nftTimeout bounds one nft set listing on the GET path; nft waits + // for the nftables lock while firewalld reloads. + nftTimeout = 3 * time.Second + + // leaseRetry is how long to wait before reading the lease file again + // after catching dnsmasq in the middle of rewriting it. + leaseRetry = 200 * time.Millisecond +) + +// DBusMonitor subscribes to dnsmasq and firewalld D-Bus signals and +// updates the shared operational tree. +type DBusMonitor struct { + tree *tree.Tree + log *slog.Logger + + mu sync.Mutex + conn *dbus.Conn // current bus connection, nil while disconnected + + // software reads the infix-system:software object; overridable in + // tests. + software func(ctx context.Context) json.RawMessage + + // Last address-set overlay read in full, served while firewalld is + // too busy to answer, e.g. during a reload. + overlayMu sync.Mutex + lastOverlay json.RawMessage +} + +// New creates a DBusMonitor. Address-set contents are served through +// an on-demand tree provider rather than the cached firewall tree: +// dynamic entries come and go without any firewalld signal (add/remove +// actions, per-entry timeouts expiring in the kernel), so they must be +// read fresh on every query. +func New(t *tree.Tree, log *slog.Logger) *DBusMonitor { + m := &DBusMonitor{ + tree: t, + log: log, + software: func(ctx context.Context) json.RawMessage { + return collector.BootSoftware(ctx, collector.ExecRunner{}) + }, + } + t.RegisterProvider(firewallTreeKey, m.addressSetOverlay) + return m +} + +func (m *DBusMonitor) setConn(conn *dbus.Conn) { + m.mu.Lock() + m.conn = conn + m.mu.Unlock() +} + +func (m *DBusMonitor) getConn() *dbus.Conn { + m.mu.Lock() + defer m.mu.Unlock() + return m.conn +} + +// Run starts the monitor loop. It connects to the system bus, subscribes +// to relevant signals, loads initial DHCP/firewall data, and reconnects +// with exponential backoff on failures until ctx is cancelled. +func (m *DBusMonitor) Run(ctx context.Context) error { + return backoff.Retry(ctx, m.log, "dbus monitor", m.session) +} + +// session runs one system bus connection until it drops or ctx ends. +func (m *DBusMonitor) session(ctx context.Context) error { + conn, err := dbus.ConnectSystemBus() + if err != nil { + return fmt.Errorf("connect system bus: %w", err) + } + defer conn.Close() + + if err := m.subscribe(conn); err != nil { + return err + } + + m.setConn(conn) + defer m.setConn(nil) + + if err := m.refreshDHCP(conn); err != nil { + m.log.Warn("dbus monitor: initial dhcp refresh failed", "err", err) + } + if err := m.refreshFirewall(conn); err != nil { + m.log.Warn("dbus monitor: initial firewall refresh failed", "err", err) + } + // The boot-time read of rauc status fails when rauc is not up yet, + // and nothing else fills in booted and the slots. + m.refreshSoftware() + + return m.processSignals(ctx, conn) +} + +func (m *DBusMonitor) subscribe(conn *dbus.Conn) error { + if err := conn.AddMatchSignal( + dbus.WithMatchInterface(dnsmasqInterface), + dbus.WithMatchMember("DHCPLeaseAdded"), + ); err != nil { + return fmt.Errorf("add dnsmasq DHCPLeaseAdded match: %w", err) + } + + if err := conn.AddMatchSignal( + dbus.WithMatchInterface(dnsmasqInterface), + dbus.WithMatchMember("DHCPLeaseDeleted"), + ); err != nil { + return fmt.Errorf("add dnsmasq DHCPLeaseDeleted match: %w", err) + } + + if err := conn.AddMatchSignal( + dbus.WithMatchInterface(dnsmasqInterface), + dbus.WithMatchMember("DHCPLeaseUpdated"), + ); err != nil { + return fmt.Errorf("add dnsmasq DHCPLeaseUpdated match: %w", err) + } + + if err := conn.AddMatchSignal( + dbus.WithMatchInterface(firewalldInterface), + dbus.WithMatchMember("Reloaded"), + ); err != nil { + return fmt.Errorf("add firewalld Reloaded match: %w", err) + } + + if err := conn.AddMatchSignal( + dbus.WithMatchInterface(raucInstallerInterface), + dbus.WithMatchMember("Completed"), + ); err != nil { + return fmt.Errorf("add rauc Completed match: %w", err) + } + + if err := conn.AddMatchSignal( + dbus.WithMatchInterface(dbusInterface), + dbus.WithMatchMember("NameOwnerChanged"), + dbus.WithMatchArg(0, dnsmasqBusName), + ); err != nil { + return fmt.Errorf("add NameOwnerChanged dnsmasq match: %w", err) + } + + if err := conn.AddMatchSignal( + dbus.WithMatchInterface(dbusInterface), + dbus.WithMatchMember("NameOwnerChanged"), + dbus.WithMatchArg(0, firewalldBusName), + ); err != nil { + return fmt.Errorf("add NameOwnerChanged firewalld match: %w", err) + } + + if err := conn.AddMatchSignal( + dbus.WithMatchInterface(dbusInterface), + dbus.WithMatchMember("NameOwnerChanged"), + dbus.WithMatchArg(0, raucBusName), + ); err != nil { + return fmt.Errorf("add NameOwnerChanged rauc match: %w", err) + } + + return nil +} + +func (m *DBusMonitor) processSignals(ctx context.Context, conn *dbus.Conn) error { + sigCh := make(chan *dbus.Signal, 128) + conn.Signal(sigCh) + defer conn.RemoveSignal(sigCh) + + for { + select { + case <-ctx.Done(): + return ctx.Err() + case sig, ok := <-sigCh: + if !ok { + return fmt.Errorf("dbus signal channel closed") + } + if sig == nil { + continue + } + if err := m.handleSignal(conn, sig); err != nil { + m.log.Warn("dbus monitor: failed handling signal", "name", sig.Name, "path", sig.Path, "err", err) + } + } + } +} + +func (m *DBusMonitor) handleSignal(conn *dbus.Conn, sig *dbus.Signal) error { + switch sig.Name { + case dnsmasqInterface + ".DHCPLeaseAdded", + dnsmasqInterface + ".DHCPLeaseDeleted", + dnsmasqInterface + ".DHCPLeaseUpdated": + if sig.Path != "" && string(sig.Path) != dnsmasqPath { + return nil + } + return m.refreshDHCP(conn) + + case raucInstallerInterface + ".Completed": + m.refreshSoftware() + return nil + + case firewalldInterface + ".Reloaded": + if sig.Path != "" && string(sig.Path) != firewalldPath { + return nil + } + return m.refreshFirewall(conn) + + case dbusInterface + ".NameOwnerChanged": + if sig.Path != "" && string(sig.Path) != dbusPath { + return nil + } + if len(sig.Body) < 3 { + return fmt.Errorf("NameOwnerChanged: expected 3 args, got %d", len(sig.Body)) + } + + name, ok1 := sig.Body[0].(string) + oldOwner, ok2 := sig.Body[1].(string) + newOwner, ok3 := sig.Body[2].(string) + if !ok1 || !ok2 || !ok3 { + return fmt.Errorf("NameOwnerChanged: unexpected arg types") + } + + switch name { + case dnsmasqBusName: + if newOwner == "" { + m.clearTreeKey(dhcpTreeKey) + return nil + } + if oldOwner == "" { + return m.refreshDHCP(conn) + } + case firewalldBusName: + if newOwner == "" { + m.clearTreeKey(firewallTreeKey) + return nil + } + if oldOwner == "" { + return m.refreshFirewall(conn) + } + case raucBusName: + if oldOwner == "" && newOwner != "" { + m.refreshSoftware() + } + } + } + + return nil +} + +// refreshSoftware re-reads the slots after an install, successful or +// not: the inactive slot's bundle, checksum and install count change, +// and nothing else reports it until the next boot. +func (m *DBusMonitor) refreshSoftware() { + ctx, cancel := context.WithTimeout(context.Background(), softwareTimeout) + defer cancel() + + if data := m.software(ctx); data != nil { + m.tree.Merge(systemStateKey, data) + } +} + +func (m *DBusMonitor) refreshDHCP(conn *dbus.Conn) error { + return m.loadDHCP(conn, true) +} + +// loadDHCP publishes the leases and server statistics. A failed or +// torn read keeps the previous leases rather than publishing an empty +// list, and is retried once after leaseRetry. +func (m *DBusMonitor) loadDHCP(conn *dbus.Conn, retry bool) error { + data, err := readLeases(dnsmasqLeaseFile) + if err != nil { + if retry { + time.AfterFunc(leaseRetry, func() { + if c := m.getConn(); c != nil { + if err := m.loadDHCP(c, false); err != nil { + m.log.Warn("dbus monitor: dhcp refresh retry failed", "err", err) + } + } + }) + } + return err + } + + leases := parseDnsmasqLeases(data) + stats := defaultDHCPStats() + + obj := conn.Object(dnsmasqBusName, dbus.ObjectPath(dnsmasqPath)) + call := obj.Call(dnsmasqInterface+".GetMetrics", 0) + if call.Err != nil { + m.log.Warn("dbus monitor: dnsmasq GetMetrics failed", "err", call.Err) + } else if len(call.Body) > 0 { + stats = mergeDHCPStats(stats, decodeDHCPMetrics(call.Body[0])) + } + + m.tree.Set(dhcpTreeKey, buildDHCPTree(leases, stats)) + return nil +} + +// readLeases reads the dnsmasq lease file. A missing file means no +// leases. dnsmasq rewrites the file in place, truncate then print, so +// a read can land mid-rewrite; every line ends in a newline, so a +// missing final one gives the torn read away. +func readLeases(path string) (string, error) { + data, err := os.ReadFile(path) + if errors.Is(err, fs.ErrNotExist) { + return "", nil + } + if err != nil { + return "", fmt.Errorf("read %s: %w", path, err) + } + if len(data) > 0 && data[len(data)-1] != '\n' { + return "", fmt.Errorf("read %s: partial, dnsmasq is rewriting it", path) + } + return string(data), nil +} + +func (m *DBusMonitor) refreshFirewall(conn *dbus.Conn) error { + obj := conn.Object(firewalldBusName, dbus.ObjectPath(firewalldPath)) + + defaultZone := "" + if call := obj.Call(firewalldInterface+".getDefaultZone", 0); call.Err != nil { + m.log.Info("dbus monitor: firewalld not reachable, skipping", "err", call.Err) + return nil + } else if err := call.Store(&defaultZone); err != nil { + m.log.Warn("dbus monitor: firewalld getDefaultZone decode failed", "err", err) + return nil + } + + logDenied := "" + if call := obj.Call(firewalldInterface+".getLogDenied", 0); call.Err != nil { + m.log.Warn("dbus monitor: firewalld getLogDenied failed", "err", call.Err) + } else if err := call.Store(&logDenied); err != nil { + m.log.Warn("dbus monitor: firewalld getLogDenied decode failed", "err", err) + } + + lockdown := false + if call := obj.Call(firewalldInterface+".queryPanicMode", 0); call.Err != nil { + m.log.Warn("dbus monitor: firewalld queryPanicMode failed", "err", call.Err) + } else if len(call.Body) > 0 { + lockdown = asBool(call.Body[0]) + } + + zones := m.getFirewallZones(obj) + policies := m.getFirewallPolicies(obj) + services := m.getFirewallServices(obj) + + m.tree.Set(firewallTreeKey, buildFirewallTree(defaultZone, logDenied, lockdown, zones, policies, services)) + return nil +} + +func (m *DBusMonitor) getFirewallZones(obj dbus.BusObject) []map[string]any { + active := make(map[string]map[string]any) + if call := obj.Call(firewalldInterface+".zone.getActiveZones", 0); call.Err != nil { + m.log.Warn("dbus monitor: firewalld zone.getActiveZones failed", "err", call.Err) + return nil + } else if len(call.Body) > 0 { + active = decodeActiveZones(call.Body[0]) + } + + zones := make([]map[string]any, 0, len(active)) + for name, zoneInfo := range active { + settings := map[string]any{} + if call := obj.Call(firewalldInterface+".zone.getZoneSettings2", 0, name); call.Err != nil { + m.log.Warn("dbus monitor: firewalld zone.getZoneSettings2 failed", "zone", name, "err", call.Err) + continue + } else if len(call.Body) > 0 { + settings = variantMap(call.Body[0]) + } + + zone := map[string]any{ + "name": name, + "immutable": hasImmutableTag(getString(settings, "short")), + "action": mapZoneTarget(getString(settings, "target")), + } + if ifaces := firstStringList(zoneInfo, "interfaces", getStringList(settings, "interfaces")); len(ifaces) > 0 { + zone["interface"] = ifaces + } + sources := firstStringList(zoneInfo, "sources", getStringList(settings, "sources")) + networks, ipsets := splitSources(sources) + if len(networks) > 0 { + zone["network"] = networks + } + if len(ipsets) > 0 { + zone["address-set"] = ipsets + } + if services := getStringList(settings, "services"); len(services) > 0 { + zone["service"] = services + } + if desc := getString(settings, "description"); desc != "" { + zone["description"] = desc + } + + if forwards := getForwardPorts(settings); len(forwards) > 0 { + zone["port-forward"] = forwards + } + + zones = append(zones, zone) + } + + return zones +} + +func (m *DBusMonitor) getFirewallPolicies(obj dbus.BusObject) []map[string]any { + var names []string + if call := obj.Call(firewalldInterface+".policy.getPolicies", 0); call.Err != nil { + m.log.Warn("dbus monitor: firewalld policy.getPolicies failed", "err", call.Err) + } else if err := call.Store(&names); err != nil { + m.log.Warn("dbus monitor: firewalld policy.getPolicies decode failed", "err", err) + } + + policies := make([]map[string]any, 0, len(names)+1) + for _, name := range names { + settings := map[string]any{} + if call := obj.Call(firewalldInterface+".policy.getPolicySettings", 0, name); call.Err != nil { + m.log.Warn("dbus monitor: firewalld policy.getPolicySettings failed", "policy", name, "err", call.Err) + continue + } else if len(call.Body) > 0 { + settings = variantMap(call.Body[0]) + } + + policy := map[string]any{ + "name": name, + "action": mapPolicyTarget(getString(settings, "target")), + "priority": getInt(settings, "priority", 32767), + "immutable": hasImmutableTag(getString(settings, "short")), + "masquerade": asBool(settings["masquerade"]), + } + if ingress := getStringList(settings, "ingress_zones"); len(ingress) > 0 { + policy["ingress"] = ingress + } + if egress := getStringList(settings, "egress_zones"); len(egress) > 0 { + policy["egress"] = egress + } + if desc := getString(settings, "description"); desc != "" { + policy["description"] = desc + } + if services := getStringList(settings, "services"); len(services) > 0 { + policy["service"] = services + } + if custom := parsePolicyCustomFilters(getStringList(settings, "rich_rules")); len(custom) > 0 { + policy["custom"] = map[string]any{"filter": custom} + } + + policies = append(policies, policy) + } + + policies = append(policies, map[string]any{ + "name": "default-drop", + "description": "Default deny rule - drops all unmatched traffic", + "action": "drop", + "priority": 32767, + "ingress": []string{"ANY"}, + "egress": []string{"ANY"}, + "immutable": true, + }) + + return policies +} + +// getFirewallServices lists every service firewalld knows, not only the +// ones a zone uses: show firewall service looks any of them up. +func (m *DBusMonitor) getFirewallServices(obj dbus.BusObject) []map[string]any { + var names []string + if call := obj.Call(firewalldInterface+".listServices", 0); call.Err != nil { + m.log.Warn("dbus monitor: firewalld listServices failed", "err", call.Err) + return nil + } else if err := call.Store(&names); err != nil { + m.log.Warn("dbus monitor: firewalld listServices decode failed", "err", err) + return nil + } + + services := make([]map[string]any, 0, len(names)) + for _, name := range names { + settings := map[string]any{} + if call := obj.Call(firewalldInterface+".getServiceSettings2", 0, name); call.Err != nil { + m.log.Warn("dbus monitor: firewalld getServiceSettings2 failed", "service", name, "err", call.Err) + continue + } else if len(call.Body) > 0 { + settings = variantMap(call.Body[0]) + } + + service := map[string]any{"name": name} + if ports := parseServicePorts(settings); len(ports) > 0 { + service["port"] = ports + } + if desc := getString(settings, "description"); desc != "" { + service["description"] = desc + } + + services = append(services, service) + } + + return services +} + +// addressSetOverlay is the on-demand tree provider for the firewall +// subtree. It returns a fresh {"address-set": [...]} overlay, or nil +// when firewalld is unreachable or has no sets. +// overlayBudget bounds what a GET waits for firewalld. It is single +// threaded and does not answer while it reloads, which on slow hardware +// takes longer than statd waits for yangerd. +const overlayBudget = 2 * time.Second + +func (m *DBusMonitor) addressSetOverlay() json.RawMessage { + conn := m.getConn() + if conn == nil { + return nil + } + + ctx, cancel := context.WithTimeout(context.Background(), overlayBudget) + defer cancel() + + obj := conn.Object(firewalldBusName, dbus.ObjectPath(firewalldPath)) + sets, ok := m.getAddressSets(ctx, obj) + if !ok { + m.overlayMu.Lock() + defer m.overlayMu.Unlock() + return m.lastOverlay + } + + var raw json.RawMessage + if len(sets) > 0 { + var err error + if raw, err = json.Marshal(map[string]any{"address-set": sets}); err != nil { + return nil + } + } + + m.overlayMu.Lock() + m.lastOverlay = raw + m.overlayMu.Unlock() + return raw +} + +// getAddressSets reads every address-set, ok is false when firewalld did +// not answer in full. +func (m *DBusMonitor) getAddressSets(ctx context.Context, obj dbus.BusObject) ([]map[string]any, bool) { + var names []string + if call := obj.CallWithContext(ctx, firewalldInterface+".ipset.getIPSets", 0); call.Err != nil { + m.log.Debug("dbus monitor: firewalld ipset.getIPSets failed", "err", call.Err) + return nil, false + } else if err := call.Store(&names); err != nil { + m.log.Warn("dbus monitor: firewalld ipset.getIPSets decode failed", "err", err) + return nil, false + } + + sets := make([]map[string]any, 0, len(names)) + for _, name := range names { + aset, ok := m.getAddressSet(ctx, obj, name) + if !ok { + return nil, false + } + if aset != nil { + sets = append(sets, aset) + } + } + return sets, true +} + +func (m *DBusMonitor) getAddressSet(ctx context.Context, obj dbus.BusObject, name string) (map[string]any, bool) { + call := obj.CallWithContext(ctx, firewalldInterface+".ipset.getIPSetSettings", 0, name) + if call.Err != nil { + m.log.Warn("dbus monitor: firewalld ipset.getIPSetSettings failed", "ipset", name, "err", call.Err) + return nil, false + } + if len(call.Body) == 0 { + return nil, true + } + + // (version, short, description, type, options, entries) + fields, ok := call.Body[0].([]any) + if !ok || len(fields) < 6 { + m.log.Warn("dbus monitor: firewalld ipset settings: unexpected shape", "ipset", name) + return nil, true + } + + options := variantMap(fields[4]) + tracked := toStringSlice(fields[5]) + + aset := map[string]any{"name": name} + + if desc := fmt.Sprint(fields[2]); desc != "" { + aset["description"] = desc + } + + family := "ipv4" + if getString(options, "family") == "inet6" { + family = "ipv6" + } + aset["family"] = family + + timeout := getInt(options, "timeout", 0) + if timeout > 0 { + aset["timeout"] = timeout + } + + shadow := readShadowEntries(name) + + static := []string{} + for _, e := range tracked { + e = normalizeEntry(e) + if !shadow[e] { + static = append(static, e) + } + } + if len(static) > 0 { + aset["entry"] = static + } + + var current []map[string]any + if timeout > 0 { + current = kernelEntries(ctx, name) + } else { + current = trackedEntries(tracked, shadow) + } + if len(current) > 0 { + aset["current"] = current + } + + return aset, true +} + +func readShadowEntries(name string) map[string]bool { + shadow := map[string]bool{} + data, err := os.ReadFile(filepath.Join(addrsetShadowDir, name)) + if err != nil { + return shadow + } + for _, line := range strings.Split(string(data), "\n") { + line = strings.TrimSpace(line) + if line != "" { + shadow[normalizeEntry(line)] = true + } + } + return shadow +} + +// trackedEntries is the current contents of a set without a timeout: +// firewalld tracks every member, the runtime-added ones included, so +// the kernel need not be asked. +func trackedEntries(tracked []string, shadow map[string]bool) []map[string]any { + current := make([]map[string]any, 0, len(tracked)) + for _, e := range tracked { + e = normalizeEntry(e) + current = append(current, map[string]any{ + "entry": e, + "dynamic": shadow[e], + }) + } + return current +} + +// kernelEntries is the current contents of a timeout set. firewalld +// does not track those members, so the kernel is the only source, and +// the only one knowing the expiry. Every member of a timeout set is +// dynamic by definition. +func kernelEntries(ctx context.Context, name string) []map[string]any { + current := []map[string]any{} + for _, elem := range nftSetElems(ctx, name) { + entry, expires := nftElemParse(elem) + cur := map[string]any{ + "entry": entry, + "dynamic": true, + } + if expires >= 0 { + cur["expires"] = expires + } + current = append(current, cur) + } + return current +} + +// nftSetElems returns the live contents of firewalld's nftables set. +// The firewalld table is owner-protected, but reading is fine. +func nftSetElems(ctx context.Context, name string) []any { + ctx, cancel := context.WithTimeout(ctx, nftTimeout) + defer cancel() + + out, err := exec.CommandContext(ctx, "nft", "-j", "list", "set", "inet", "firewalld", name).Output() + if err != nil { + return nil + } + return parseNftSetElems(out) +} + +func parseNftSetElems(out []byte) []any { + var doc struct { + Nftables []map[string]json.RawMessage `json:"nftables"` + } + if json.Unmarshal(out, &doc) != nil { + return nil + } + for _, obj := range doc.Nftables { + raw, ok := obj["set"] + if !ok { + continue + } + var set struct { + Elem []any `json:"elem"` + } + if json.Unmarshal(raw, &set) == nil { + return set.Elem + } + } + return nil +} + +// nftElemParse returns (entry, expires) from an nft JSON set element. +// expires is -1 when the element carries no expiry. +func nftElemParse(elem any) (string, int) { + expires := -1 + if wrap, ok := elem.(map[string]any); ok { + if inner, ok := wrap["elem"].(map[string]any); ok { + if e, ok := inner["expires"]; ok { + expires = getNum(e) + } + elem = inner["val"] + } + } + + var entry string + switch v := elem.(type) { + case map[string]any: + if p, ok := v["prefix"].(map[string]any); ok { + entry = fmt.Sprintf("%v/%d", p["addr"], getNum(p["len"])) + } else if r, ok := v["range"].([]any); ok && len(r) == 2 { + entry = fmt.Sprintf("%v-%v", r[0], r[1]) + } else { + entry = fmt.Sprint(v) + } + default: + entry = fmt.Sprint(v) + } + + return normalizeEntry(entry), expires +} + +func getNum(v any) int { + if n, ok := numconv.Int(v); ok { + return n + } + return -1 +} + +// normalizeEntry matches firewalld's entry normalization: host bits are +// masked off prefixes and full-length prefixes reduce to bare addresses. +func normalizeEntry(entry string) string { + if p, err := netip.ParsePrefix(entry); err == nil { + p = p.Masked() + if p.Bits() == p.Addr().BitLen() { + return p.Addr().String() + } + return p.String() + } + if a, err := netip.ParseAddr(entry); err == nil { + return a.String() + } + return entry +} + +// splitSources separates zone sources into IP networks and +// "ipset:NAME" address-set references. +func splitSources(sources []string) (networks, ipsets []string) { + for _, src := range sources { + if name, ok := strings.CutPrefix(src, "ipset:"); ok { + ipsets = append(ipsets, name) + } else { + networks = append(networks, src) + } + } + return networks, ipsets +} + +func (m *DBusMonitor) clearTreeKey(key string) { + m.tree.Delete(key) +} + +func parseDnsmasqLeases(data string) []map[string]any { + leases := make([]map[string]any, 0) + for _, line := range strings.Split(data, "\n") { + line = strings.TrimSpace(line) + if line == "" { + continue + } + + fields := strings.Fields(line) + if len(fields) != 5 { + continue + } + + expires := "never" + if fields[0] != "0" { + ts, err := strconv.ParseInt(fields[0], 10, 64) + if err != nil { + continue + } + expires = time.Unix(ts, 0).UTC().Format(time.RFC3339) + } + + hostname := "" + if fields[3] != "*" { + hostname = fields[3] + } + + clientID := "" + if fields[4] != "*" { + clientID = fields[4] + } + + leases = append(leases, map[string]any{ + "expires": expires, + "address": fields[2], + "phys-address": fields[1], + "hostname": hostname, + "client-id": clientID, + }) + } + + return leases +} + +func buildDHCPTree(leases []map[string]any, stats map[string]any) json.RawMessage { + root := map[string]any{ + "statistics": stats, + "leases": map[string]any{ + "lease": leases, + }, + } + raw, err := json.Marshal(root) + if err != nil { + return json.RawMessage(`{}`) + } + return raw +} + +func buildFirewallTree(defaultZone, logDenied string, lockdown bool, zones, policies, services []map[string]any) json.RawMessage { + fw := map[string]any{ + "default": defaultZone, + "logging": logDenied, + "lockdown": lockdown, + } + if len(zones) > 0 { + fw["zone"] = zones + } + if len(policies) > 0 { + fw["policy"] = policies + } + if len(services) > 0 { + fw["service"] = services + } + + raw, err := json.Marshal(fw) + if err != nil { + return json.RawMessage(`{}`) + } + return raw +} + +func defaultDHCPStats() map[string]any { + return map[string]any{ + "out-offers": uint64(0), + "out-acks": uint64(0), + "out-naks": uint64(0), + "in-declines": uint64(0), + "in-discovers": uint64(0), + "in-requests": uint64(0), + "in-releases": uint64(0), + "in-informs": uint64(0), + } +} + +func decodeDHCPMetrics(v any) map[string]any { + metrics := map[string]any{} + + switch raw := v.(type) { + case map[string]dbus.Variant: + for k, val := range raw { + metrics[k] = val.Value() + } + case map[string]any: + for k, val := range raw { + metrics[k] = val + } + } + + return map[string]any{ + "out-offers": numconv.Uint64(metrics["dhcp_offer"]), + "out-acks": numconv.Uint64(metrics["dhcp_ack"]), + "out-naks": numconv.Uint64(metrics["dhcp_nak"]), + "in-declines": numconv.Uint64(metrics["dhcp_decline"]), + "in-discovers": numconv.Uint64(metrics["dhcp_discover"]), + "in-requests": numconv.Uint64(metrics["dhcp_request"]), + "in-releases": numconv.Uint64(metrics["dhcp_release"]), + "in-informs": numconv.Uint64(metrics["dhcp_inform"]), + } +} + +func mergeDHCPStats(base, override map[string]any) map[string]any { + out := map[string]any{} + for k, v := range base { + out[k] = v + } + for k, v := range override { + out[k] = v + } + return out +} + +func parseServicePorts(settings map[string]any) []map[string]any { + rawPorts, ok := settings["ports"] + if !ok { + return []map[string]any{} + } + + out := []map[string]any{} + for _, entry := range toAnySlice(rawPorts) { + pair := toAnySlice(entry) + if len(pair) < 2 { + continue + } + + portSpec := fmt.Sprint(pair[0]) + proto := fmt.Sprint(pair[1]) + if portSpec == "" || proto == "" { + continue + } + + port := map[string]any{"proto": proto} + if !setPortRange(port, portSpec) { + continue + } + + out = append(out, port) + } + + return out +} + +// setPortRange sets lower, and upper for a range, from a firewalld port +// spec, "80" or "8000-8080". It reports whether the spec was valid. +func setPortRange(dst map[string]any, spec string) bool { + lo, hi, isRange := strings.Cut(spec, "-") + lower, err := strconv.Atoi(strings.TrimSpace(lo)) + if err != nil { + return false + } + if isRange { + upper, err := strconv.Atoi(strings.TrimSpace(hi)) + if err != nil { + return false + } + dst["upper"] = upper + } + dst["lower"] = lower + return true +} + +func parsePolicyCustomFilters(rules []string) []map[string]any { + filters := []map[string]any{} + for _, rule := range rules { + family := "both" + if strings.Contains(rule, `family="ipv4"`) { + family = "ipv4" + } else if strings.Contains(rule, `family="ipv6"`) { + family = "ipv6" + } + + icmpType := "" + action := "" + prio := -1 + + if idx := strings.Index(rule, "priority="); idx >= 0 { + prio = parsePriority(rule[idx+len("priority="):]) + } + + if strings.Contains(rule, "icmp-type") && strings.Contains(rule, `name="`) { + icmpType = parseQuotedName(rule) + action = "accept" + if strings.Contains(rule, " drop") { + action = "drop" + } else if strings.Contains(rule, " reject") { + action = "reject" + } + } else if strings.Contains(rule, "icmp-block") && strings.Contains(rule, `name="`) { + icmpType = parseQuotedName(rule) + action = "reject" + } + + if icmpType == "" || action == "" { + continue + } + + filters = append(filters, map[string]any{ + "name": "icmp-" + icmpType, + "priority": prio, + "family": family, + "action": action, + "icmp": map[string]any{ + "type": icmpType, + }, + }) + } + + return filters +} + +func getForwardPorts(settings map[string]any) []map[string]any { + raw, ok := settings["forward_ports"] + if !ok { + return nil + } + + out := []map[string]any{} + for _, item := range toAnySlice(raw) { + vals := toAnySlice(item) + if len(vals) < 4 { + continue + } + + portStr := fmt.Sprint(vals[0]) + proto := fmt.Sprint(vals[1]) + toPortStr := strings.TrimSpace(fmt.Sprint(vals[2])) + toAddr := fmt.Sprint(vals[3]) + + if portStr == "" || proto == "" { + continue + } + + entry := map[string]any{"proto": proto} + if !setPortRange(entry, portStr) { + continue + } + + to := map[string]any{"addr": toAddr} + if toPortStr != "" && !strings.ContainsAny(toPortStr, ".:") { + if p, err := strconv.Atoi(toPortStr); err == nil { + to["port"] = p + } + } + if _, ok := to["port"]; !ok { + to["port"] = entry["lower"] + } + + entry["to"] = to + out = append(out, entry) + } + + return out +} + +func decodeActiveZones(v any) map[string]map[string]any { + out := map[string]map[string]any{} + + switch m := v.(type) { + case map[string]map[string]dbus.Variant: + for zone, data := range m { + inner := map[string]any{} + for k, vv := range data { + inner[k] = vv.Value() + } + out[zone] = inner + } + case map[string]map[string]any: + for zone, data := range m { + out[zone] = data + } + case map[string]map[string][]string: + for zone, data := range m { + inner := map[string]any{} + for k, v := range data { + inner[k] = v + } + out[zone] = inner + } + case map[string]any: + for zone, raw := range m { + if mm, ok := raw.(map[string]any); ok { + out[zone] = mm + } + } + } + + return out +} + +func variantMap(v any) map[string]any { + out := map[string]any{} + switch m := v.(type) { + case map[string]dbus.Variant: + for k, vv := range m { + out[k] = vv.Value() + } + case map[string]string: + for k, vv := range m { + out[k] = vv + } + case map[string]any: + for k, vv := range m { + if dv, ok := vv.(dbus.Variant); ok { + out[k] = dv.Value() + } else { + out[k] = vv + } + } + } + return out +} + +func getString(m map[string]any, key string) string { + v, ok := m[key] + if !ok || v == nil { + return "" + } + return fmt.Sprint(v) +} + +func getInt(m map[string]any, key string, def int) int { + if n, ok := numconv.Int(m[key]); ok { + return n + } + return def +} + +func getStringList(m map[string]any, key string) []string { + v, ok := m[key] + if !ok { + return nil + } + return toStringSlice(v) +} + +func toStringSlice(v any) []string { + vals := toAnySlice(v) + if len(vals) == 0 { + if s, ok := v.(string); ok { + if s == "" { + return nil + } + return []string{s} + } + return nil + } + + out := make([]string, 0, len(vals)) + for _, item := range vals { + s := strings.TrimSpace(fmt.Sprint(item)) + if s != "" { + out = append(out, s) + } + } + return out +} + +func toAnySlice(v any) []any { + switch a := v.(type) { + case []any: + return a + case []string: + out := make([]any, 0, len(a)) + for _, item := range a { + out = append(out, item) + } + return out + case [][]any: + out := make([]any, 0, len(a)) + for _, item := range a { + out = append(out, any(item)) + } + return out + case [][]string: + out := make([]any, 0, len(a)) + for _, item := range a { + inner := make([]any, 0, len(item)) + for _, p := range item { + inner = append(inner, p) + } + out = append(out, inner) + } + return out + } + return nil +} + +func firstStringList(a map[string]any, key string, fallback []string) []string { + if list := getStringList(a, key); len(list) > 0 { + return list + } + return fallback +} + +func hasImmutableTag(short string) bool { + return strings.Contains(short, "(immutable)") +} + +func mapZoneTarget(target string) string { + switch strings.ToUpper(strings.TrimSpace(target)) { + case "%%REJECT%%", "REJECT": + return "reject" + case "DROP": + return "drop" + case "ACCEPT", "DEFAULT", "": + return "accept" + default: + return "accept" + } +} + +func mapPolicyTarget(target string) string { + switch strings.ToUpper(strings.TrimSpace(target)) { + case "CONTINUE": + return "continue" + case "ACCEPT": + return "accept" + case "DROP": + return "drop" + case "REJECT", "": + return "reject" + default: + return "reject" + } +} + +func parseQuotedName(rule string) string { + idx := strings.Index(rule, `name="`) + if idx < 0 { + return "" + } + start := idx + len(`name="`) + end := strings.Index(rule[start:], `"`) + if end < 0 { + return "" + } + return rule[start : start+end] +} + +func parsePriority(fragment string) int { + fragment = strings.TrimSpace(fragment) + if fragment == "" { + return -1 + } + fields := strings.Fields(fragment) + if len(fields) == 0 { + return -1 + } + p, err := strconv.Atoi(strings.Trim(fields[0], `"`)) + if err != nil { + return -1 + } + return p +} + +func asBool(v any) bool { + switch x := v.(type) { + case bool: + return x + case uint8: + return x != 0 + case uint16: + return x != 0 + case uint32: + return x != 0 + case uint64: + return x != 0 + case int8: + return x != 0 + case int16: + return x != 0 + case int32: + return x != 0 + case int64: + return x != 0 + case int: + return x != 0 + case string: + x = strings.TrimSpace(strings.ToLower(x)) + return x == "1" || x == "true" || x == "yes" || x == "on" + default: + return false + } +} diff --git a/src/yangerd/internal/dbusmonitor/dbusmonitor_test.go b/src/yangerd/internal/dbusmonitor/dbusmonitor_test.go new file mode 100644 index 000000000..968ef447f --- /dev/null +++ b/src/yangerd/internal/dbusmonitor/dbusmonitor_test.go @@ -0,0 +1,844 @@ +package dbusmonitor + +import ( + "context" + "encoding/json" + "log/slog" + "os" + "path/filepath" + "reflect" + "testing" + "time" + + "github.com/godbus/dbus/v5" + "github.com/kernelkit/infix/src/yangerd/internal/tree" +) + +func TestParseDnsmasqLeases(t *testing.T) { + tests := []struct { + name string + input string + want []map[string]any + }{ + { + name: "normal lease line", + input: "1711900000 aa:bb:cc:dd:ee:ff 192.168.1.100 myhost 01:aa:bb:cc:dd:ee:ff", + want: []map[string]any{{ + "expires": time.Unix(1711900000, 0).UTC().Format(time.RFC3339), + "address": "192.168.1.100", + "phys-address": "aa:bb:cc:dd:ee:ff", + "hostname": "myhost", + "client-id": "01:aa:bb:cc:dd:ee:ff", + }}, + }, + { + name: "wildcard hostname and client id", + input: "1711900000 aa:bb:cc:dd:ee:ff 192.168.1.100 * *", + want: []map[string]any{{ + "expires": time.Unix(1711900000, 0).UTC().Format(time.RFC3339), + "address": "192.168.1.100", + "phys-address": "aa:bb:cc:dd:ee:ff", + "hostname": "", + "client-id": "", + }}, + }, + { + name: "never expiring lease", + input: "0 aa:bb:cc:dd:ee:ff 192.168.1.100 host *", + want: []map[string]any{{ + "expires": "never", + "address": "192.168.1.100", + "phys-address": "aa:bb:cc:dd:ee:ff", + "hostname": "host", + "client-id": "", + }}, + }, + { + name: "multiple leases with malformed lines skipped", + input: "1711900000 aa:bb:cc:dd:ee:ff 192.168.1.100 myhost 01:aa:bb:cc:dd:ee:ff\n" + + "bad line with too few fields\n" + + "1711900100 11:22:33:44:55:66 192.168.1.101 host2 *\n", + want: []map[string]any{ + { + "expires": time.Unix(1711900000, 0).UTC().Format(time.RFC3339), + "address": "192.168.1.100", + "phys-address": "aa:bb:cc:dd:ee:ff", + "hostname": "myhost", + "client-id": "01:aa:bb:cc:dd:ee:ff", + }, + { + "expires": time.Unix(1711900100, 0).UTC().Format(time.RFC3339), + "address": "192.168.1.101", + "phys-address": "11:22:33:44:55:66", + "hostname": "host2", + "client-id": "", + }, + }, + }, + { + name: "empty input", + input: "", + want: []map[string]any{}, + }, + { + name: "invalid timestamp skipped", + input: "abc aa:bb:cc:dd:ee:ff 192.168.1.100 host *", + want: []map[string]any{}, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + got := parseDnsmasqLeases(tc.input) + if !reflect.DeepEqual(got, tc.want) { + t.Fatalf("parseDnsmasqLeases() mismatch\nwant: %#v\n got: %#v", tc.want, got) + } + }) + } +} + +func TestBuildDHCPTree(t *testing.T) { + tests := []struct { + name string + leases []map[string]any + stats map[string]any + check func(t *testing.T, root map[string]any) + }{ + { + name: "with leases and stats", + leases: []map[string]any{{ + "expires": "never", + "address": "192.168.1.100", + "phys-address": "aa:bb:cc:dd:ee:ff", + "hostname": "host", + "client-id": "", + }}, + stats: map[string]any{"out-offers": 3, "in-requests": 4}, + check: func(t *testing.T, root map[string]any) { + t.Helper() + stats, ok := root["statistics"].(map[string]any) + if !ok { + t.Fatalf("missing statistics map") + } + if stats["out-offers"] != float64(3) || stats["in-requests"] != float64(4) { + t.Fatalf("unexpected statistics: %#v", stats) + } + + leasesNode, ok := root["leases"].(map[string]any) + if !ok { + t.Fatalf("missing leases map") + } + leaseList, ok := leasesNode["lease"].([]any) + if !ok || len(leaseList) != 1 { + t.Fatalf("unexpected lease list: %#v", leasesNode["lease"]) + } + lease, ok := leaseList[0].(map[string]any) + if !ok || lease["address"] != "192.168.1.100" { + t.Fatalf("unexpected lease entry: %#v", leaseList[0]) + } + }, + }, + { + name: "with empty leases", + leases: []map[string]any{}, + stats: map[string]any{"out-offers": 0}, + check: func(t *testing.T, root map[string]any) { + t.Helper() + leasesNode := root["leases"].(map[string]any) + leaseList, ok := leasesNode["lease"].([]any) + if !ok || len(leaseList) != 0 { + t.Fatalf("expected empty lease list, got %#v", leasesNode["lease"]) + } + }, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + raw := buildDHCPTree(tc.leases, tc.stats) + var root map[string]any + if err := json.Unmarshal(raw, &root); err != nil { + t.Fatalf("unmarshal buildDHCPTree output: %v", err) + } + tc.check(t, root) + }) + } +} + +func TestBuildFirewallTree(t *testing.T) { + tests := []struct { + name string + defaultZ string + logDenied string + lockdown bool + zones []map[string]any + policies []map[string]any + services []map[string]any + expectKeys map[string]bool + }{ + { + name: "with zones policies and services", + defaultZ: "public", + logDenied: "all", + lockdown: true, + zones: []map[string]any{{"name": "public"}}, + policies: []map[string]any{{"name": "default-drop"}}, + services: []map[string]any{{"name": "ssh"}}, + expectKeys: map[string]bool{ + "zone": true, + "policy": true, + "service": true, + }, + }, + { + name: "omits empty zone policy service keys", + defaultZ: "trusted", + logDenied: "off", + lockdown: false, + expectKeys: map[string]bool{ + "zone": false, + "policy": false, + "service": false, + }, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + raw := buildFirewallTree(tc.defaultZ, tc.logDenied, tc.lockdown, tc.zones, tc.policies, tc.services) + var root map[string]any + if err := json.Unmarshal(raw, &root); err != nil { + t.Fatalf("unmarshal buildFirewallTree output: %v", err) + } + if root["default"] != tc.defaultZ || root["logging"] != tc.logDenied || root["lockdown"] != tc.lockdown { + t.Fatalf("default/logging/lockdown mismatch: %#v", root) + } + for k, shouldExist := range tc.expectKeys { + _, exists := root[k] + if exists != shouldExist { + t.Fatalf("key %q exists=%v, want %v", k, exists, shouldExist) + } + } + }) + } +} + +func TestParseServicePorts(t *testing.T) { + tests := []struct { + name string + settings map[string]any + want []map[string]any + }{ + { + name: "single port", + settings: map[string]any{"ports": []any{[]any{"80", "tcp"}}}, + want: []map[string]any{{"proto": "tcp", "lower": 80}}, + }, + { + name: "port range", + settings: map[string]any{"ports": []any{[]any{"8080-8090", "tcp"}}}, + want: []map[string]any{{"proto": "tcp", "lower": 8080, "upper": 8090}}, + }, + { + name: "multiple ports", + settings: map[string]any{"ports": []any{ + []any{"80", "tcp"}, + []any{"53", "udp"}, + }}, + want: []map[string]any{ + {"proto": "tcp", "lower": 80}, + {"proto": "udp", "lower": 53}, + }, + }, + { + name: "missing ports", + settings: map[string]any{}, + want: []map[string]any{}, + }, + { + name: "empty ports", + settings: map[string]any{"ports": []any{}}, + want: []map[string]any{}, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + got := parseServicePorts(tc.settings) + if !reflect.DeepEqual(got, tc.want) { + t.Fatalf("parseServicePorts mismatch\nwant: %#v\n got: %#v", tc.want, got) + } + }) + } +} + +func TestParsePolicyCustomFilters(t *testing.T) { + tests := []struct { + name string + rules []string + want []map[string]any + }{ + { + name: "rich rule icmp type accept", + rules: []string{`rule priority="0" family="ipv4" icmp-type name="echo-request" accept`}, + want: []map[string]any{{ + "name": "icmp-echo-request", + "priority": 0, + "family": "ipv4", + "action": "accept", + "icmp": map[string]any{"type": "echo-request"}, + }}, + }, + { + name: "rich rule icmp block reject", + rules: []string{`rule family="ipv6" icmp-block name="router-advertisement" reject`}, + want: []map[string]any{{ + "name": "icmp-router-advertisement", + "priority": -1, + "family": "ipv6", + "action": "reject", + "icmp": map[string]any{"type": "router-advertisement"}, + }}, + }, + { + name: "rule without icmp skipped", + rules: []string{`rule family="ipv4" service name="ssh" accept`}, + want: []map[string]any{}, + }, + { + name: "multiple rules include only icmp", + rules: []string{ + `rule priority="10" family="ipv4" icmp-type name="echo-reply" drop`, + `rule family="ipv4" service name="http" accept`, + `rule family="ipv6" icmp-block name="router-advertisement" reject`, + }, + want: []map[string]any{ + { + "name": "icmp-echo-reply", + "priority": 10, + "family": "ipv4", + "action": "drop", + "icmp": map[string]any{"type": "echo-reply"}, + }, + { + "name": "icmp-router-advertisement", + "priority": -1, + "family": "ipv6", + "action": "reject", + "icmp": map[string]any{"type": "router-advertisement"}, + }, + }, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + got := parsePolicyCustomFilters(tc.rules) + if !reflect.DeepEqual(got, tc.want) { + t.Fatalf("parsePolicyCustomFilters mismatch\nwant: %#v\n got: %#v", tc.want, got) + } + }) + } +} + +func TestGetForwardPorts(t *testing.T) { + tests := []struct { + name string + settings map[string]any + want []map[string]any + wantNil bool + }{ + { + name: "single port forward", + settings: map[string]any{"forward_ports": []any{[]any{"80", "tcp", "8080", "192.168.1.1"}}}, + want: []map[string]any{{ + "proto": "tcp", + "lower": 80, + "to": map[string]any{"addr": "192.168.1.1", "port": 8080}, + }}, + }, + { + name: "port range forward", + settings: map[string]any{"forward_ports": []any{[]any{"1000-1005", "udp", "2000", "10.0.0.2"}}}, + want: []map[string]any{{ + "proto": "udp", + "lower": 1000, + "upper": 1005, + "to": map[string]any{"addr": "10.0.0.2", "port": 2000}, + }}, + }, + { + name: "missing to port defaults to lower", + settings: map[string]any{"forward_ports": []any{[]any{"8081", "tcp", "", "192.168.1.1"}}}, + want: []map[string]any{{ + "proto": "tcp", + "lower": 8081, + "to": map[string]any{"addr": "192.168.1.1", "port": 8081}, + }}, + }, + { + name: "missing forward ports", + settings: map[string]any{}, + wantNil: true, + }, + { + name: "empty forward ports", + settings: map[string]any{"forward_ports": []any{}}, + want: []map[string]any{}, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + got := getForwardPorts(tc.settings) + if tc.wantNil { + if got != nil { + t.Fatalf("expected nil, got %#v", got) + } + return + } + if !reflect.DeepEqual(got, tc.want) { + t.Fatalf("getForwardPorts mismatch\nwant: %#v\n got: %#v", tc.want, got) + } + }) + } +} + +func TestMapZoneTarget(t *testing.T) { + tests := []struct { + name string + in string + want string + }{ + {name: "percent reject", in: "%%REJECT%%", want: "reject"}, + {name: "reject", in: "REJECT", want: "reject"}, + {name: "drop", in: "DROP", want: "drop"}, + {name: "accept", in: "ACCEPT", want: "accept"}, + {name: "default", in: "DEFAULT", want: "accept"}, + {name: "empty", in: "", want: "accept"}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + if got := mapZoneTarget(tc.in); got != tc.want { + t.Fatalf("mapZoneTarget(%q) = %q, want %q", tc.in, got, tc.want) + } + }) + } +} + +func TestMapPolicyTarget(t *testing.T) { + tests := []struct { + name string + in string + want string + }{ + {name: "continue", in: "CONTINUE", want: "continue"}, + {name: "accept", in: "ACCEPT", want: "accept"}, + {name: "drop", in: "DROP", want: "drop"}, + {name: "reject", in: "REJECT", want: "reject"}, + {name: "empty", in: "", want: "reject"}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + if got := mapPolicyTarget(tc.in); got != tc.want { + t.Fatalf("mapPolicyTarget(%q) = %q, want %q", tc.in, got, tc.want) + } + }) + } +} + +func TestAsBool(t *testing.T) { + tests := []struct { + name string + in any + want bool + }{ + {name: "bool true", in: true, want: true}, + {name: "bool false", in: false, want: false}, + {name: "int one", in: 1, want: true}, + {name: "int zero", in: 0, want: false}, + {name: "int8 one", in: int8(1), want: true}, + {name: "int16 zero", in: int16(0), want: false}, + {name: "int64 one", in: int64(1), want: true}, + {name: "uint32 zero", in: uint32(0), want: false}, + {name: "uint64 one", in: uint64(1), want: true}, + {name: "string true", in: "true", want: true}, + {name: "string false", in: "false", want: false}, + {name: "string one", in: "1", want: true}, + {name: "string zero", in: "0", want: false}, + {name: "string yes", in: "yes", want: true}, + {name: "string no", in: "no", want: false}, + {name: "string on", in: "on", want: true}, + {name: "trim and case", in: " TRUE ", want: true}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + if got := asBool(tc.in); got != tc.want { + t.Fatalf("asBool(%#v) = %v, want %v", tc.in, got, tc.want) + } + }) + } +} + +func TestParseHelpers(t *testing.T) { + t.Run("parseQuotedName", func(t *testing.T) { + tests := []struct { + name string + rule string + want string + }{ + {name: "extract name", rule: `rule icmp-type name="echo-request" accept`, want: "echo-request"}, + {name: "missing name", rule: `rule icmp-type accept`, want: ""}, + {name: "unterminated quote", rule: `rule icmp-type name="echo-request accept`, want: ""}, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + if got := parseQuotedName(tc.rule); got != tc.want { + t.Fatalf("parseQuotedName(%q) = %q, want %q", tc.rule, got, tc.want) + } + }) + } + }) + + t.Run("parsePriority", func(t *testing.T) { + tests := []struct { + name string + in string + want int + }{ + {name: "quoted value", in: `"0" family="ipv4"`, want: 0}, + {name: "plain value", in: `10 family="ipv6"`, want: 10}, + {name: "empty", in: ``, want: -1}, + {name: "invalid", in: `abc family="ipv4"`, want: -1}, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + if got := parsePriority(tc.in); got != tc.want { + t.Fatalf("parsePriority(%q) = %d, want %d", tc.in, got, tc.want) + } + }) + } + }) + + t.Run("hasImmutableTag", func(t *testing.T) { + tests := []struct { + name string + in string + want bool + }{ + {name: "has tag", in: "Public (immutable)", want: true}, + {name: "no tag", in: "Public", want: false}, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + if got := hasImmutableTag(tc.in); got != tc.want { + t.Fatalf("hasImmutableTag(%q) = %v, want %v", tc.in, got, tc.want) + } + }) + } + }) +} + +func TestDecodeActiveZones(t *testing.T) { + tests := []struct { + name string + in any + want map[string]map[string]any + }{ + { + name: "godbus concrete type (a{sa{sas}})", + in: map[string]map[string][]string{ + "public": { + "interfaces": {"eth0", "eth1"}, + "sources": {"10.0.0.0/8"}, + }, + "mgmt": { + "interfaces": {"eth2"}, + }, + }, + want: map[string]map[string]any{ + "public": { + "interfaces": []string{"eth0", "eth1"}, + "sources": []string{"10.0.0.0/8"}, + }, + "mgmt": { + "interfaces": []string{"eth2"}, + }, + }, + }, + { + name: "pre-decoded map[string]map[string]any", + in: map[string]map[string]any{ + "home": { + "interfaces": []string{"wlan0"}, + }, + }, + want: map[string]map[string]any{ + "home": { + "interfaces": []string{"wlan0"}, + }, + }, + }, + { + name: "nil input", + in: nil, + want: map[string]map[string]any{}, + }, + { + name: "unsupported type", + in: "garbage", + want: map[string]map[string]any{}, + }, + { + name: "empty map", + in: map[string]map[string][]string{}, + want: map[string]map[string]any{}, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + got := decodeActiveZones(tc.in) + if !reflect.DeepEqual(got, tc.want) { + t.Fatalf("decodeActiveZones() =\n %v\nwant:\n %v", got, tc.want) + } + }) + } +} + +func TestSplitSources(t *testing.T) { + tests := []struct { + name string + in []string + networks []string + ipsets []string + }{ + { + name: "mixed sources", + in: []string{"10.0.0.0/8", "ipset:allowed", "192.168.1.0/24"}, + networks: []string{"10.0.0.0/8", "192.168.1.0/24"}, + ipsets: []string{"allowed"}, + }, + { + name: "only ipsets", + in: []string{"ipset:allowed", "ipset:greylist"}, + ipsets: []string{"allowed", "greylist"}, + }, + { + name: "only networks", + in: []string{"10.0.0.0/8"}, + networks: []string{"10.0.0.0/8"}, + }, + {name: "empty"}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + networks, ipsets := splitSources(tc.in) + if !reflect.DeepEqual(networks, tc.networks) || !reflect.DeepEqual(ipsets, tc.ipsets) { + t.Fatalf("splitSources(%v) = %v, %v, want %v, %v", + tc.in, networks, ipsets, tc.networks, tc.ipsets) + } + }) + } +} + +func TestNormalizeEntry(t *testing.T) { + tests := []struct { + in string + want string + }{ + {"192.168.1.40", "192.168.1.40"}, + {"192.168.1.40/32", "192.168.1.40"}, + {"192.168.1.40/24", "192.168.1.0/24"}, + {"10.0.0.0/8", "10.0.0.0/8"}, + {"2001:db8::1/128", "2001:db8::1"}, + {"2001:db8::/64", "2001:db8::/64"}, + {"1.2.3.4-1.2.3.9", "1.2.3.4-1.2.3.9"}, + } + + for _, tc := range tests { + t.Run(tc.in, func(t *testing.T) { + if got := normalizeEntry(tc.in); got != tc.want { + t.Fatalf("normalizeEntry(%q) = %q, want %q", tc.in, got, tc.want) + } + }) + } +} + +func TestNftElemParse(t *testing.T) { + tests := []struct { + name string + in any + entry string + expires int + }{ + { + name: "plain address", + in: "192.168.1.42", + entry: "192.168.1.42", + expires: -1, + }, + { + name: "prefix", + in: map[string]any{"prefix": map[string]any{"addr": "10.0.0.0", "len": float64(8)}}, + entry: "10.0.0.0/8", + expires: -1, + }, + { + name: "timeout element", + in: map[string]any{"elem": map[string]any{ + "val": "10.0.0.1", + "timeout": float64(10), + "expires": float64(7), + }}, + entry: "10.0.0.1", + expires: 7, + }, + { + name: "wrapped prefix with expiry", + in: map[string]any{"elem": map[string]any{ + "val": map[string]any{"prefix": map[string]any{"addr": "10.1.0.0", "len": float64(16)}}, + "expires": float64(42), + }}, + entry: "10.1.0.0/16", + expires: 42, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + entry, expires := nftElemParse(tc.in) + if entry != tc.entry || expires != tc.expires { + t.Fatalf("nftElemParse() = %q, %d, want %q, %d", + entry, expires, tc.entry, tc.expires) + } + }) + } +} + +func TestParseNftSetElems(t *testing.T) { + out := []byte(`{"nftables": [ + {"metainfo": {"version": "1.0.9"}}, + {"set": {"family": "inet", "name": "allowed", "table": "firewalld", + "type": "ipv4_addr", + "elem": ["192.168.1.40", {"elem": {"val": "192.168.1.42", "expires": 5}}]}} + ]}`) + + elems := parseNftSetElems(out) + if len(elems) != 2 { + t.Fatalf("expected 2 elements, got %d: %#v", len(elems), elems) + } + + entry, expires := nftElemParse(elems[0]) + if entry != "192.168.1.40" || expires != -1 { + t.Fatalf("elem 0: got %q, %d", entry, expires) + } + entry, expires = nftElemParse(elems[1]) + if entry != "192.168.1.42" || expires != 5 { + t.Fatalf("elem 1: got %q, %d", entry, expires) + } + + if elems := parseNftSetElems([]byte(`{}`)); elems != nil { + t.Fatalf("expected nil for empty document, got %#v", elems) + } + if elems := parseNftSetElems([]byte(`garbage`)); elems != nil { + t.Fatalf("expected nil for invalid JSON, got %#v", elems) + } +} + +// An install replaces the whole software object, so the new bundle on +// the inactive slot shows up without a reboot, and other system-state +// subtrees are left alone. +func TestRaucCompletedRefreshesSoftware(t *testing.T) { + tr := tree.New() + tr.Set(systemStateKey, json.RawMessage(`{"platform":{"os-name":"Infix"},"infix-system:software":{"slot":[{"name":"rootfs.1","bundle":{"version":"v1"}}]}}`)) + + m := New(tr, slog.Default()) + m.software = func(context.Context) json.RawMessage { + return json.RawMessage(`{"infix-system:software":{"slot":[{"name":"rootfs.1","bundle":{"version":"v2"}}]}}`) + } + + sig := &dbus.Signal{Name: raucInstallerInterface + ".Completed", Body: []any{int32(0)}} + if err := m.handleSignal(nil, sig); err != nil { + t.Fatal(err) + } + + var state map[string]json.RawMessage + if err := json.Unmarshal(tr.GetCached(systemStateKey), &state); err != nil { + t.Fatal(err) + } + if got := string(state["infix-system:software"]); got != `{"slot":[{"name":"rootfs.1","bundle":{"version":"v2"}}]}` { + t.Fatalf("software not replaced: %s", got) + } + if _, ok := state["platform"]; !ok { + t.Fatal("platform dropped by the software refresh") + } +} + +// Sets without a timeout are served from firewalld's own list, with +// the shadow file telling runtime-added members apart. +func TestTrackedEntries(t *testing.T) { + got := trackedEntries([]string{"10.0.0.1", "192.168.1.0/24", "10.0.0.9/32"}, + map[string]bool{"10.0.0.9": true}) + want := []map[string]any{ + {"entry": "10.0.0.1", "dynamic": false}, + {"entry": "192.168.1.0/24", "dynamic": false}, + {"entry": "10.0.0.9", "dynamic": true}, + } + if !reflect.DeepEqual(got, want) { + t.Fatalf("got %v, want %v", got, want) + } +} + +func TestReadLeases(t *testing.T) { + dir := t.TempDir() + full := filepath.Join(dir, "full") + torn := filepath.Join(dir, "torn") + os.WriteFile(full, []byte("1711900000 aa:bb:cc:dd:ee:ff 10.0.0.5 host *\n"), 0644) + os.WriteFile(torn, []byte("1711900000 aa:bb:cc:dd:ee:ff 10.0.0.5 ho"), 0644) + + if data, err := readLeases(full); err != nil || data == "" { + t.Fatalf("complete file: data %q, err %v", data, err) + } + if _, err := readLeases(torn); err == nil { + t.Fatal("a file without its final newline must be reported as torn") + } + if data, err := readLeases(filepath.Join(dir, "missing")); err != nil || data != "" { + t.Fatalf("missing file must mean no leases: data %q, err %v", data, err) + } +} + +// yangerd often starts before rauc, and the boot-time rauc status read +// then fails, so the software object is re-read when rauc comes up. +func TestRaucAppearingRefreshesSoftware(t *testing.T) { + tr := tree.New() + tr.Set(systemStateKey, json.RawMessage(`{"infix-system:software":{"boot-order":["primary","secondary"]}}`)) + + calls := 0 + m := New(tr, slog.Default()) + m.software = func(context.Context) json.RawMessage { + calls++ + return json.RawMessage(`{"infix-system:software":{"booted":"primary","boot-order":["primary","secondary"]}}`) + } + + gone := &dbus.Signal{Name: dbusInterface + ".NameOwnerChanged", Body: []any{raucBusName, ":1.7", ""}} + if err := m.handleSignal(nil, gone); err != nil || calls != 0 { + t.Fatalf("rauc leaving: err=%v calls=%d", err, calls) + } + up := &dbus.Signal{Name: dbusInterface + ".NameOwnerChanged", Body: []any{raucBusName, "", ":1.8"}} + if err := m.handleSignal(nil, up); err != nil || calls != 1 { + t.Fatalf("rauc appearing: err=%v calls=%d", err, calls) + } + + var state map[string]map[string]any + if err := json.Unmarshal(tr.GetCached(systemStateKey), &state); err != nil { + t.Fatal(err) + } + if state["infix-system:software"]["booted"] != "primary" { + t.Fatalf("booted missing after rauc appeared: %v", state) + } +} diff --git a/src/yangerd/internal/ethmonitor/ethmonitor.go b/src/yangerd/internal/ethmonitor/ethmonitor.go new file mode 100644 index 000000000..bdf45a2ad --- /dev/null +++ b/src/yangerd/internal/ethmonitor/ethmonitor.go @@ -0,0 +1,434 @@ +// Package ethmonitor subscribes to ethtool genetlink notifications and +// keeps per-interface ethernet settings updated via a callback. +// +// Data is fetched by shelling out to `ethtool --json ` (matching +// the Python yanger approach) while genetlink provides reactive change +// notifications. The exec runs on a worker fed by a set of pending +// interface names, so a burst of link events on many ports costs one +// ethtool per port and never blocks the caller. +package ethmonitor + +import ( + "context" + "encoding/json" + "fmt" + "log/slog" + "net" + "strings" + "sync" + + "github.com/kernelkit/infix/src/yangerd/internal/backoff" + "github.com/kernelkit/infix/src/yangerd/internal/collector" + "github.com/mdlayher/genetlink" + "github.com/mdlayher/netlink" +) + +// Kernel-to-user ethtool message types, from +// include/uapi/linux/ethtool_netlink_generated.h. +const ( + ETHTOOL_MSG_LINKINFO_NTF = 3 + ETHTOOL_MSG_LINKMODES_NTF = 5 + + ethtoolFamilyName = "ethtool" + ethtoolMonitorGroupName = "monitor" + + nlaHeaderIfindex = 1 + + ethtoolSpeedUnknown = (1 << 32) - 1 +) + +// EthMonitor listens for ethtool genetlink monitor events and updates +// interface ethernet operational state via a callback. +type EthMonitor struct { + cmd collector.CommandRunner + log *slog.Logger + onUpdate func(ifname string, data json.RawMessage) + + // ifaceName resolves a kernel ifindex; overridable in tests. + ifaceName func(index int) (string, error) + + mu sync.Mutex + pending map[string]struct{} + kick chan struct{} +} + +// New creates an EthMonitor. The genetlink socket is opened by Run. +func New(log *slog.Logger, cmd collector.CommandRunner) *EthMonitor { + return &EthMonitor{ + cmd: cmd, + log: log, + ifaceName: ifNameByIndex, + pending: make(map[string]struct{}), + kick: make(chan struct{}, 1), + } +} + +// SetOnUpdate sets the callback invoked when ethernet data changes. +func (m *EthMonitor) SetOnUpdate(fn func(string, json.RawMessage)) { + m.onUpdate = fn +} + +// Run serves refresh requests and the ethtool notification stream until +// ctx is cancelled, reconnecting with backoff if the socket fails. +func (m *EthMonitor) Run(ctx context.Context) error { + go m.worker(ctx) + return backoff.Retry(ctx, m.log, "ethmonitor", m.receive) +} + +func (m *EthMonitor) receive(ctx context.Context) error { + conn, err := genetlink.Dial(nil) + if err != nil { + return fmt.Errorf("dial genetlink: %w", err) + } + defer conn.Close() + + family, err := conn.GetFamily(ethtoolFamilyName) + if err != nil { + return fmt.Errorf("resolve %q genetlink family: %w", ethtoolFamilyName, err) + } + var groupID uint32 + for _, g := range family.Groups { + if g.Name == ethtoolMonitorGroupName { + groupID = g.ID + break + } + } + if groupID == 0 { + return fmt.Errorf("multicast group %q not found in family %q", ethtoolMonitorGroupName, ethtoolFamilyName) + } + if err := conn.JoinGroup(groupID); err != nil { + return fmt.Errorf("join ethtool monitor group %d: %w", groupID, err) + } + + // Receive has no deadline; closing the socket is what ends it. + stop := context.AfterFunc(ctx, func() { conn.Close() }) + defer stop() + + for { + msgs, _, err := conn.Receive() + if err != nil { + if ctx.Err() != nil { + return ctx.Err() + } + return fmt.Errorf("receive ethtool genetlink message: %w", err) + } + for _, msg := range msgs { + m.dispatch(msg) + } + } +} + +// dispatch queues a refresh for the interface a link notification is +// about. +func (m *EthMonitor) dispatch(msg genetlink.Message) { + switch msg.Header.Command { + case ETHTOOL_MSG_LINKINFO_NTF, ETHTOOL_MSG_LINKMODES_NTF: + default: + return + } + index, err := headerIfindex(msg.Data) + if err != nil { + m.log.Warn("ethmonitor: decode notification", "err", err) + return + } + ifname, err := m.ifaceName(index) + if err != nil { + m.log.Debug("ethmonitor: notification for unknown ifindex", "index", index, "err", err) + return + } + m.RefreshInterface(ifname) +} + +// RefreshInterface queues a refresh of the ethernet settings for ifname. +// Called by the notification stream and by nlmonitor on link events; +// repeated requests before the worker gets to them collapse into one. +func (m *EthMonitor) RefreshInterface(ifname string) { + m.mu.Lock() + m.pending[ifname] = struct{}{} + m.mu.Unlock() + + select { + case m.kick <- struct{}{}: + default: + } +} + +func (m *EthMonitor) worker(ctx context.Context) { + for { + select { + case <-ctx.Done(): + return + case <-m.kick: + } + for { + ifname, ok := m.next() + if !ok { + break + } + m.refreshEthernetSettings(ctx, ifname) + } + } +} + +func (m *EthMonitor) next() (string, bool) { + m.mu.Lock() + defer m.mu.Unlock() + for ifname := range m.pending { + delete(m.pending, ifname) + return ifname, true + } + return "", false +} + +// ethtoolJSON represents the relevant fields from `ethtool --json `. +type ethtoolJSON struct { + // Speed must be a 64-bit type: ethtool reports unknown speed as + // 0xFFFFFFFF, which overflows int on 32-bit targets (arm). + Speed int64 `json:"speed"` + Duplex string `json:"duplex"` + Port string `json:"port"` + AutoNegotiation bool `json:"auto-negotiation"` + SupportedLinkModes []string `json:"supported-link-modes"` + AdvertisedLinkModes []string `json:"advertised-link-modes"` +} + +func (m *EthMonitor) refreshEthernetSettings(ctx context.Context, ifname string) { + out, err := m.cmd.Run(ctx, "ethtool", "--json", ifname) + if err != nil { + m.log.Warn("ethmonitor: run ethtool", "ifname", ifname, "err", err) + return + } + + var results []ethtoolJSON + if err := json.Unmarshal(out, &results); err != nil { + m.log.Warn("ethmonitor: parse ethtool json", "ifname", ifname, "err", err) + return + } + if len(results) == 0 { + return + } + + data := results[0] + eth, speedBPS := buildEthernetContainer(data) + + // Marshal the result; include interface-level speed as a special key + // that mergeAugments will lift onto the interface object. + result := map[string]any{"ethernet": eth} + if speedBPS > 0 { + result["speed"] = fmt.Sprintf("%d", speedBPS) + } + + raw, err := json.Marshal(result) + if err != nil { + m.log.Warn("ethmonitor: marshal ethernet settings", "ifname", ifname, "err", err) + return + } + + if m.onUpdate != nil { + m.onUpdate(ifname, json.RawMessage(raw)) + } +} + +// buildEthernetContainer builds the ieee802-ethernet-interface:ethernet +// container and returns (container, interface speed in bits/s or 0). +func buildEthernetContainer(data ethtoolJSON) (map[string]any, int64) { + autoneg := map[string]any{"enable": data.AutoNegotiation} + eth := map[string]any{"auto-negotiation": autoneg} + + duplex := strings.ToLower(data.Duplex) + if duplex == "full" || duplex == "half" { + eth["duplex"] = duplex + } + + // Supported PMD types (config-false leaf-list). + supported := ethtoolModesToPMD(data.SupportedLinkModes) + if len(supported) > 0 { + eth["infix-ethernet-interface:supported-pmd-types"] = supported + } + + // Advertised PMD types — suppress when identical to supported (default). + advertised := ethtoolModesToPMD(data.AdvertisedLinkModes) + if len(advertised) > 0 && !stringSliceEqual(advertised, supported) { + autoneg["infix-ethernet-interface:advertised-pmd-types"] = advertised + } + + // Speed, phy-type, pmd-type. + var speedBPS int64 + speedMbps := data.Speed + if speedMbps > 0 && speedMbps < ethtoolSpeedUnknown { + speedBPS = int64(speedMbps) * 1_000_000 + + // Speed inside the ethernet container (decimal64, Gb/s). + eth["speed"] = fmt.Sprintf("%.3f", float64(speedMbps)/1000.0) + + key := linkModeKey{Port: data.Port, SpeedMbps: speedMbps, Duplex: duplex} + if mapping, ok := linkModes[key]; ok { + eth["phy-type"] = "ieee802-ethernet-phy-type:phy-type-" + mapping.PhyType + if mapping.PMDType != "" { + eth["pmd-type"] = "ieee802-ethernet-phy-type:pmd-type-" + mapping.PMDType + } + } + + // Refine pmd-type when exactly one supported mode (specific SFP). + if len(supported) == 1 { + eth["pmd-type"] = supported[0] + } + } + + return eth, speedBPS +} + +// headerIfindex returns the ifindex from the request header nest that +// every ethtool notification carries. +func headerIfindex(data []byte) (int, error) { + ad, err := netlink.NewAttributeDecoder(data) + if err != nil { + return 0, fmt.Errorf("new decoder: %w", err) + } + + for ad.Next() { + nested, err := netlink.NewAttributeDecoder(ad.Bytes()) + if err != nil { + continue + } + + for nested.Next() { + if nested.Type() == nlaHeaderIfindex { + return int(nested.Uint32()), nil + } + } + if err := nested.Err(); err != nil { + return 0, fmt.Errorf("decode nested attrs: %w", err) + } + } + + if err := ad.Err(); err != nil { + return 0, fmt.Errorf("decode attrs: %w", err) + } + + return 0, fmt.Errorf("header ifindex attribute not found") +} + +func ifNameByIndex(index int) (string, error) { + iface, err := net.InterfaceByIndex(index) + if err != nil { + return "", err + } + return iface.Name, nil +} + +// linkModeKey is the lookup key for phy-type/pmd-type mapping. +type linkModeKey struct { + Port string + SpeedMbps int64 + Duplex string +} + +// linkModeMapping holds the IEEE identity suffixes. +type linkModeMapping struct { + PhyType string + PMDType string // empty means "cannot determine from this tuple alone" +} + +// linkModes maps (port, speed, duplex) → (phy-type, pmd-type) per +// IEEE Std 802.3.2-2025 (ieee802-ethernet-phy-type). +var linkModes = map[linkModeKey]linkModeMapping{ + {"Twisted Pair", 10, "full"}: {"10BASE-T", "10BASE-T"}, + {"Twisted Pair", 10, "half"}: {"10BASE-T", "10BASE-T"}, + {"Twisted Pair", 100, "full"}: {"100BASE-X", "100BASE-TX"}, + {"Twisted Pair", 100, "half"}: {"100BASE-X", "100BASE-TX"}, + {"Twisted Pair", 1000, "full"}: {"1000BASE-T", "1000BASE-T"}, + {"Twisted Pair", 1000, "half"}: {"1000BASE-T", "1000BASE-T"}, + {"Twisted Pair", 2500, "full"}: {"2.5GBASE-T", "2.5GBASE-T"}, + {"Twisted Pair", 5000, "full"}: {"5GBASE-T", "5GBASE-T"}, + {"Twisted Pair", 10000, "full"}: {"10GBASE-T", "10GBASE-T"}, + {"Twisted Pair", 25000, "full"}: {"25GBASE-T", "25GBASE-T"}, + {"Twisted Pair", 40000, "full"}: {"40GBASE-T", "40GBASE-T"}, + {"MII", 10, "full"}: {"10BASE-T", "10BASE-T"}, + {"MII", 10, "half"}: {"10BASE-T", "10BASE-T"}, + {"MII", 100, "full"}: {"100BASE-X", "100BASE-TX"}, + {"MII", 100, "half"}: {"100BASE-X", "100BASE-TX"}, + {"FIBRE", 100, "full"}: {"100BASE-X", ""}, + {"FIBRE", 1000, "full"}: {"1000BASE-X", ""}, + {"FIBRE", 10000, "full"}: {"10GBASE-R", ""}, + {"FIBRE", 25000, "full"}: {"25GBASE-R", ""}, + {"FIBRE", 40000, "full"}: {"40GBASE-R", ""}, + {"FIBRE", 100000, "full"}: {"100GBASE-R", ""}, + {"Direct Attach Copper", 10000, "full"}: {"10GBASE-R", ""}, + {"Direct Attach Copper", 25000, "full"}: {"25GBASE-R", "25GBASE-CR"}, + {"Direct Attach Copper", 40000, "full"}: {"40GBASE-R", "40GBASE-CR4"}, + {"Direct Attach Copper", 100000, "full"}: {"100GBASE-R", "100GBASE-CR4"}, +} + +// ethtoolToPMD maps kernel link-mode base names to IEEE pmd-type +// identity suffixes. The kernel reports modes like "1000baseT/Full"; +// we strip the "/Full" or "/Half" suffix before lookup. +var ethtoolToPMD = map[string]string{ + "10baseT": "10BASE-T", + "10baseT1L": "10BASE-T1L", + "100baseT": "100BASE-TX", + "100baseT1": "100BASE-T1", + "100baseFX": "100BASE-FX", + "1000baseT": "1000BASE-T", + "1000baseT1": "1000BASE-T1", + "1000baseX": "1000BASE-LX", + "1000baseKX": "1000BASE-KX", + "2500baseT": "2.5GBASE-T", + "2500baseX": "2.5GBASE-X", + "5000baseT": "5GBASE-T", + "10000baseT": "10GBASE-T", + "10000baseSR": "10GBASE-SR", + "10000baseLR": "10GBASE-LR", + "10000baseLRM": "10GBASE-LRM", + "10000baseER": "10GBASE-ER", + "10000baseKR": "10GBASE-KR", + "10000baseKX4": "10GBASE-KX4", + "25000baseCR": "25GBASE-CR", + "25000baseSR": "25GBASE-SR", + "25000baseKR": "25GBASE-KR", + "40000baseCR4": "40GBASE-CR4", + "40000baseSR4": "40GBASE-SR4", + "40000baseLR4": "40GBASE-LR4", + "40000baseKR4": "40GBASE-KR4", + "100000baseCR4": "100GBASE-CR4", + "100000baseSR4": "100GBASE-SR4", + "100000baseLR4_ER4": "100GBASE-LR4", + "100000baseKR4": "100GBASE-KR4", +} + +// ethtoolModesToPMD translates a list of ethtool link-mode strings +// (e.g. "1000baseT/Full") into deduped, order-preserving PMD identity +// strings. +func ethtoolModesToPMD(modes []string) []string { + seen := make(map[string]bool) + var out []string + for _, entry := range modes { + base := entry + if idx := strings.IndexByte(entry, '/'); idx >= 0 { + base = entry[:idx] + } + pmd, ok := ethtoolToPMD[base] + if !ok || seen[pmd] { + continue + } + seen[pmd] = true + out = append(out, "ieee802-ethernet-phy-type:pmd-type-"+pmd) + } + return out +} + +func stringSliceEqual(a, b []string) bool { + if len(a) != len(b) { + return false + } + set := make(map[string]bool, len(a)) + for _, s := range a { + set[s] = true + } + for _, s := range b { + if !set[s] { + return false + } + } + return true +} diff --git a/src/yangerd/internal/ethmonitor/ethmonitor_test.go b/src/yangerd/internal/ethmonitor/ethmonitor_test.go new file mode 100644 index 000000000..7c9134613 --- /dev/null +++ b/src/yangerd/internal/ethmonitor/ethmonitor_test.go @@ -0,0 +1,210 @@ +package ethmonitor + +import ( + "context" + "encoding/json" + "github.com/mdlayher/genetlink" + "github.com/mdlayher/netlink" + "log/slog" + "strings" + "sync" + "testing" + "time" +) + +func TestBuildEthernetContainerCopper1G(t *testing.T) { + data := ethtoolJSON{ + Speed: 1000, + Duplex: "Full", + Port: "Twisted Pair", + AutoNegotiation: true, + SupportedLinkModes: []string{ + "10baseT/Half", "10baseT/Full", + "100baseT/Half", "100baseT/Full", + "1000baseT/Full", + }, + AdvertisedLinkModes: []string{ + "10baseT/Half", "10baseT/Full", + "100baseT/Half", "100baseT/Full", + "1000baseT/Full", + }, + } + + eth, speedBPS := buildEthernetContainer(data) + + if speedBPS != 1_000_000_000 { + t.Fatalf("speed = %d, want 1000000000", speedBPS) + } + if eth["phy-type"] != "ieee802-ethernet-phy-type:phy-type-1000BASE-T" { + t.Fatalf("phy-type = %v", eth["phy-type"]) + } + if eth["pmd-type"] != "ieee802-ethernet-phy-type:pmd-type-1000BASE-T" { + t.Fatalf("pmd-type = %v", eth["pmd-type"]) + } + if eth["duplex"] != "full" { + t.Fatalf("duplex = %v", eth["duplex"]) + } + autoneg := eth["auto-negotiation"].(map[string]any) + if autoneg["enable"] != true { + t.Fatal("autoneg should be true") + } + // advertised == supported → no advertised-pmd-types key + if _, ok := autoneg["infix-ethernet-interface:advertised-pmd-types"]; ok { + t.Fatal("advertised-pmd-types should be suppressed when equal to supported") + } +} + +func TestBuildEthernetContainerFibre10G(t *testing.T) { + data := ethtoolJSON{ + Speed: 10000, + Duplex: "Full", + Port: "FIBRE", + AutoNegotiation: false, + SupportedLinkModes: []string{"10000baseSR/Full"}, + AdvertisedLinkModes: []string{"10000baseSR/Full"}, + } + + eth, speedBPS := buildEthernetContainer(data) + + if speedBPS != 10_000_000_000 { + t.Fatalf("speed = %d, want 10000000000", speedBPS) + } + // Fibre 10G → phy-type 10GBASE-R, no pmd-type from lookup table + // But exactly one supported mode → pmd-type refined from supported list + if eth["pmd-type"] != "ieee802-ethernet-phy-type:pmd-type-10GBASE-SR" { + t.Fatalf("pmd-type = %v, want refined from single supported mode", eth["pmd-type"]) + } + if eth["phy-type"] != "ieee802-ethernet-phy-type:phy-type-10GBASE-R" { + t.Fatalf("phy-type = %v", eth["phy-type"]) + } +} + +func TestBuildEthernetContainerSpeedUnknown(t *testing.T) { + data := ethtoolJSON{ + Speed: ethtoolSpeedUnknown, + Duplex: "Unknown! (255)", + Port: "Twisted Pair", + AutoNegotiation: true, + } + + eth, speedBPS := buildEthernetContainer(data) + + if speedBPS != 0 { + t.Fatalf("speed = %d, want 0 for unknown", speedBPS) + } + if _, ok := eth["speed"]; ok { + t.Fatal("speed should not be set when unknown") + } + if _, ok := eth["phy-type"]; ok { + t.Fatal("phy-type should not be set when speed unknown") + } +} + +func TestBuildEthernetContainerAdvertisedDiffers(t *testing.T) { + data := ethtoolJSON{ + Speed: 1000, + Duplex: "Full", + Port: "Twisted Pair", + SupportedLinkModes: []string{ + "10baseT/Full", "100baseT/Full", "1000baseT/Full", + }, + AdvertisedLinkModes: []string{"1000baseT/Full"}, + } + + eth, _ := buildEthernetContainer(data) + + autoneg := eth["auto-negotiation"].(map[string]any) + adv, ok := autoneg["infix-ethernet-interface:advertised-pmd-types"] + if !ok { + t.Fatal("advertised-pmd-types should be present when != supported") + } + advList := adv.([]string) + if len(advList) != 1 || advList[0] != "ieee802-ethernet-phy-type:pmd-type-1000BASE-T" { + t.Fatalf("advertised = %v", advList) + } +} + +func TestEthtoolModesToPMD(t *testing.T) { + modes := []string{ + "10baseT/Half", "10baseT/Full", + "1000baseT/Full", + "Autoneg", "TP", + } + got := ethtoolModesToPMD(modes) + want := []string{ + "ieee802-ethernet-phy-type:pmd-type-10BASE-T", + "ieee802-ethernet-phy-type:pmd-type-1000BASE-T", + } + if len(got) != len(want) { + t.Fatalf("got %v, want %v", got, want) + } + for i := range want { + if got[i] != want[i] { + t.Fatalf("got[%d] = %q, want %q", i, got[i], want[i]) + } + } +} + +type recordingRunner struct { + mu sync.Mutex + calls []string +} + +func (r *recordingRunner) Run(_ context.Context, name string, args ...string) ([]byte, error) { + r.mu.Lock() + r.calls = append(r.calls, strings.Join(append([]string{name}, args...), " ")) + r.mu.Unlock() + return []byte(`[{"speed":1000,"duplex":"Full","port":"Twisted Pair","auto-negotiation":true}]`), nil +} + +// A link-modes notification must reach the ethtool refresh for the +// interface named by the header ifindex, and land in the callback. +func TestDispatchNotificationRefreshesInterface(t *testing.T) { + ae := netlink.NewAttributeEncoder() + ae.Nested(1, func(nae *netlink.AttributeEncoder) error { + nae.Uint32(nlaHeaderIfindex, 7) + return nil + }) + data, err := ae.Encode() + if err != nil { + t.Fatal(err) + } + + runner := &recordingRunner{} + m := New(slog.Default(), runner) + m.ifaceName = func(index int) (string, error) { + if index != 7 { + t.Fatalf("ifindex = %d, want 7", index) + } + return "eth3", nil + } + got := make(chan string, 1) + m.SetOnUpdate(func(ifname string, raw json.RawMessage) { + if !strings.Contains(string(raw), `"duplex":"full"`) { + t.Errorf("unexpected ethernet data: %s", raw) + } + got <- ifname + }) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + go m.worker(ctx) + + m.dispatch(genetlink.Message{Header: genetlink.Header{Command: ETHTOOL_MSG_LINKMODES_NTF}, Data: data}) + m.dispatch(genetlink.Message{Header: genetlink.Header{Command: 28}, Data: data}) // not a link notification + + select { + case ifname := <-got: + if ifname != "eth3" { + t.Fatalf("refreshed %q, want eth3", ifname) + } + case <-time.After(2 * time.Second): + t.Fatal("notification never reached the callback") + } + time.Sleep(50 * time.Millisecond) + runner.mu.Lock() + defer runner.mu.Unlock() + if len(runner.calls) != 1 || runner.calls[0] != "ethtool --json eth3" { + t.Fatalf("ethtool calls = %v", runner.calls) + } +} diff --git a/src/yangerd/internal/frrvty/frrvty.go b/src/yangerd/internal/frrvty/frrvty.go new file mode 100644 index 000000000..650ddf7ae --- /dev/null +++ b/src/yangerd/internal/frrvty/frrvty.go @@ -0,0 +1,94 @@ +// Package frrvty is a minimal in-process client for an FRR daemon's vty +// Unix socket. +// +// It speaks the same protocol vtysh uses, so yangerd can run "show ..." +// commands (e.g. "show ip route json") against zebra without forking +// vtysh. The command is written NUL-terminated; the daemon streams the +// command output followed by a four-byte trailer of three NUL bytes and a +// one-byte CLI return code (\0\0\0). Each query enters the +// enable node first, as vtysh does. +package frrvty + +import ( + "bytes" + "context" + "fmt" + "io" + "net" +) + +// ZebraVtySocket is the default path to zebra's vty socket. It lives in +// the same runstatedir as the zserv API socket. +const ZebraVtySocket = "/var/run/frr/zebra.vty" + +// Client runs commands against a single FRR daemon vty socket. A fresh +// connection is opened per query, matching vtysh's behaviour. +type Client struct { + socket string +} + +// New returns a Client for the given vty socket path. An empty path +// selects the zebra socket. +func New(socket string) *Client { + if socket == "" { + socket = ZebraVtySocket + } + return &Client{socket: socket} +} + +// Query connects to the vty socket, runs one command, and returns its raw +// output with the protocol trailer stripped. A non-zero CLI return code +// is reported as an error (the partial output is still returned). +func (c *Client) Query(ctx context.Context, command string) ([]byte, error) { + var d net.Dialer + conn, err := d.DialContext(ctx, "unix", c.socket) + if err != nil { + return nil, fmt.Errorf("dial %s: %w", c.socket, err) + } + defer conn.Close() + + if deadline, ok := ctx.Deadline(); ok { + _ = conn.SetDeadline(deadline) + } + + // A fresh vty starts in the view node, but some daemons, bfdd for + // one, install their show commands in the enable node only. vtysh + // sends "enable" first for the same reason. + if _, err := send(conn, "enable"); err != nil { + return nil, err + } + return send(conn, command) +} + +// send writes one NUL-terminated command and reads its reply up to the +// \0\0\0 trailer, which it strips. +func send(conn net.Conn, command string) ([]byte, error) { + if _, err := conn.Write(append([]byte(command), 0)); err != nil { + return nil, fmt.Errorf("write %q: %w", command, err) + } + + var buf bytes.Buffer + tmp := make([]byte, 4096) + for { + n, rerr := conn.Read(tmp) + if n > 0 { + buf.Write(tmp[:n]) + // The payload is text/JSON and never contains NUL, so + // testing the last four accumulated bytes is unambiguous. + if b := buf.Bytes(); len(b) >= 4 && + b[len(b)-4] == 0 && b[len(b)-3] == 0 && b[len(b)-2] == 0 { + payload := b[:len(b)-4] + if ret := b[len(b)-1]; ret != 0 { + return payload, fmt.Errorf("vty command %q: status %d", command, ret) + } + return payload, nil + } + } + if rerr != nil { + if rerr == io.EOF { + return nil, fmt.Errorf("vty command %q: closed before trailer", command) + } + return nil, fmt.Errorf("vty command %q: read: %w", command, rerr) + } + } +} diff --git a/src/yangerd/internal/frrvty/frrvty_test.go b/src/yangerd/internal/frrvty/frrvty_test.go new file mode 100644 index 000000000..34e3b78fd --- /dev/null +++ b/src/yangerd/internal/frrvty/frrvty_test.go @@ -0,0 +1,139 @@ +package frrvty + +import ( + "bytes" + "context" + "net" + "path/filepath" + "sync" + "testing" + "time" +) + +// fakeZebra serves a single vty connection like an FRR daemon whose show +// commands live in the view node: "enable" and the command both work. +func fakeZebra(t *testing.T, reply string, ret byte) string { + return fakeDaemon(t, reply, ret, false) +} + +// fakeDaemon serves a single vty connection. It answers each +// NUL-terminated command with a \0\0\0 trailer: "enable" with an +// empty reply, anything else with the configured reply. With +// enableOnly it behaves like bfdd, whose show commands exist only in +// the enable node, and rejects a command sent before "enable". +func fakeDaemon(t *testing.T, reply string, ret byte, enableOnly bool) string { + t.Helper() + + sock := filepath.Join(t.TempDir(), "daemon.vty") + ln, err := net.Listen("unix", sock) + if err != nil { + t.Fatalf("listen: %v", err) + } + + var wg sync.WaitGroup + wg.Add(1) + go func() { + defer wg.Done() + conn, err := ln.Accept() + if err != nil { + return + } + defer conn.Close() + + enabled := false + var pending []byte + buf := make([]byte, 256) + for { + n, err := conn.Read(buf) + pending = append(pending, buf[:n]...) + for { + i := bytes.IndexByte(pending, 0) + if i < 0 { + break + } + cmd := string(pending[:i]) + pending = pending[i+1:] + + var out []byte + switch { + case cmd == "enable": + enabled = true + out = []byte{0, 0, 0, 0} + case enableOnly && !enabled: + out = append([]byte("% Unknown command: "+cmd+"\n"), 0, 0, 0, 2) + default: + out = append([]byte(reply), 0, 0, 0, ret) + } + if _, werr := conn.Write(out); werr != nil { + return + } + } + if err != nil { + return + } + } + }() + + t.Cleanup(func() { + ln.Close() + wg.Wait() + }) + return sock +} + +// bfdd installs "show bfd peers" in the enable node only, so the query +// must enter it first or the daemon answers "Unknown command". +func TestQueryEntersEnableNode(t *testing.T) { + sock := fakeDaemon(t, `[{"peer":"192.168.100.2","status":"up"}]`, 0, true) + + out, err := New(sock).Query(context.Background(), "show bfd peers json") + if err != nil { + t.Fatalf("Query: %v", err) + } + if got := string(out); got != `[{"peer":"192.168.100.2","status":"up"}]` { + t.Fatalf("output = %q", got) + } +} + +func TestQueryStripsTrailer(t *testing.T) { + sock := fakeZebra(t, `{"a":1}`, 0) + c := New(sock) + + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + + out, err := c.Query(ctx, "show ip route json") + if err != nil { + t.Fatalf("Query: %v", err) + } + if string(out) != `{"a":1}` { + t.Errorf("output = %q, want %q", out, `{"a":1}`) + } +} + +func TestQueryNonZeroStatus(t *testing.T) { + sock := fakeZebra(t, "Unknown command", 1) + c := New(sock) + + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + + out, err := c.Query(ctx, "bogus") + if err == nil { + t.Fatal("expected error for non-zero status") + } + if string(out) != "Unknown command" { + t.Errorf("partial output = %q, want %q", out, "Unknown command") + } +} + +func TestQueryDialError(t *testing.T) { + c := New(filepath.Join(t.TempDir(), "does-not-exist.vty")) + + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + + if _, err := c.Query(ctx, "show ip route json"); err == nil { + t.Fatal("expected dial error") + } +} diff --git a/src/yangerd/internal/fswatcher/fswatcher.go b/src/yangerd/internal/fswatcher/fswatcher.go new file mode 100644 index 000000000..9b47d8a62 --- /dev/null +++ b/src/yangerd/internal/fswatcher/fswatcher.go @@ -0,0 +1,311 @@ +// Package fswatcher provides inotify-based reactive monitoring of +// filesystem paths. It replaces polling for procfs files that support +// inotify (e.g. IP forwarding flags). Each watched path has a +// handler that reads the file and updates the tree, with per-path +// debouncing to coalesce burst writes. +package fswatcher + +import ( + "context" + "encoding/json" + "fmt" + "log/slog" + "path/filepath" + "sync" + "time" + + "github.com/kernelkit/infix/src/yangerd/internal/inotify" + "github.com/kernelkit/infix/src/yangerd/internal/tree" +) + +// WatchHandler defines the callback for a watched path. +type WatchHandler struct { + TreeKey string + ReadFunc func(path string) (json.RawMessage, error) + Debounce time.Duration + // UseMerge causes the watcher to call tree.Merge instead of + // tree.Set, performing a shallow first-level JSON merge into + // the existing blob at TreeKey. + UseMerge bool +} + +// FSWatcher monitors filesystem paths via inotify and updates the +// tree when files change. +type FSWatcher struct { + watcher *inotify.Watcher + tree *tree.Tree + handlers map[string]WatchHandler + dirHandlers map[string]WatchHandler // directory path → handler + pending map[string]bool // removed files waiting to reappear + debounce map[string]*time.Timer + mu sync.Mutex + log *slog.Logger +} + +// New creates an FSWatcher backed by an inotify instance. +func New(t *tree.Tree, log *slog.Logger) (*FSWatcher, error) { + w, err := inotify.NewWatcher() + if err != nil { + return nil, fmt.Errorf("inotify: %w", err) + } + return &FSWatcher{ + watcher: w, + tree: t, + handlers: make(map[string]WatchHandler), + dirHandlers: make(map[string]WatchHandler), + pending: make(map[string]bool), + debounce: make(map[string]*time.Timer), + log: log, + }, nil +} + +// Watch registers a handler for a specific filesystem path and adds +// the inotify watch. +func (fw *FSWatcher) Watch(path string, handler WatchHandler) error { + fw.mu.Lock() + fw.handlers[path] = handler + fw.mu.Unlock() + return fw.watcher.Add(path) +} + +// SyncGlob makes the watched paths matching pattern equal to what the +// pattern expands to now: new matches are watched with handler and read +// once, vanished ones dropped. inotify reports nothing when a /proc/sys +// directory appears or goes away with its interface, so the caller runs +// this on link add/del. +func (fw *FSWatcher) SyncGlob(pattern string, handler WatchHandler) (added, removed int, err error) { + matches, err := filepath.Glob(pattern) + if err != nil { + return 0, 0, fmt.Errorf("glob %s: %w", pattern, err) + } + want := make(map[string]bool, len(matches)) + for _, path := range matches { + want[path] = true + } + + fw.mu.Lock() + var gone, fresh []string + for path := range fw.handlers { + if ok, _ := filepath.Match(pattern, path); ok && !want[path] { + gone = append(gone, path) + } + } + for path := range want { + if _, ok := fw.handlers[path]; !ok { + fresh = append(fresh, path) + } + } + for _, path := range gone { + fw.dropLocked(path) + } + fw.mu.Unlock() + + for _, path := range fresh { + if err := fw.Watch(path, handler); err != nil { + fw.log.Warn("fswatcher: watch failed, skipping", "path", path, "err", err) + continue + } + fw.fireHandler(path, handler) + added++ + } + return added, len(gone), nil +} + +// dropLocked forgets a file handler. fw.mu must be held. +func (fw *FSWatcher) dropLocked(path string) { + delete(fw.handlers, path) + delete(fw.pending, path) + if timer, ok := fw.debounce[path]; ok { + timer.Stop() + delete(fw.debounce, path) + } + _ = fw.watcher.Remove(path) +} + +// WatchSymlink registers a handler for a symlink by watching its parent +// directory. fsnotify follows symlinks to the target inode, so replacing +// a symlink (ln -sf) would not trigger events on a direct watch. Watching +// the parent directory catches Create and Rename events for the symlink +// entry itself. +func (fw *FSWatcher) WatchSymlink(path string, handler WatchHandler) error { + dir := filepath.Dir(path) + fw.mu.Lock() + fw.handlers[path] = handler + fw.mu.Unlock() + return fw.watcher.Add(dir) +} + +// WatchDir registers a handler for an entire directory. Any file +// create/write/remove event inside the directory triggers the handler +// with the directory path. The handler's ReadFunc receives the directory +// path (not the individual file), so it can rescan all contents. +func (fw *FSWatcher) WatchDir(dir string, handler WatchHandler) error { + fw.mu.Lock() + fw.dirHandlers[dir] = handler + fw.mu.Unlock() + return fw.watcher.Add(dir) +} + +// InitialRead reads the current value of every watched file and +// populates the tree. Called once after all Watch() calls and glob +// expansion, before Run(). +func (fw *FSWatcher) InitialRead() { + fw.mu.Lock() + defer fw.mu.Unlock() + for path, handler := range fw.handlers { + data, err := handler.ReadFunc(path) + if err != nil { + fw.log.Warn("fswatcher: initial read failed", "path", path, "err", err) + continue + } + fw.apply(handler, data) + fw.log.Debug("fswatcher: initial read", "path", path, "key", handler.TreeKey) + } + for dir, handler := range fw.dirHandlers { + data, err := handler.ReadFunc(dir) + if err != nil { + fw.log.Warn("fswatcher: initial read failed", "path", dir, "err", err) + continue + } + fw.apply(handler, data) + fw.log.Debug("fswatcher: initial read", "path", dir, "key", handler.TreeKey) + } +} + +// Run processes inotify events until ctx is cancelled. +func (fw *FSWatcher) Run(ctx context.Context) error { + defer fw.watcher.Close() + for { + select { + case <-ctx.Done(): + return ctx.Err() + case event, ok := <-fw.watcher.Events: + if !ok { + return fmt.Errorf("watcher closed") + } + if event.Has(inotify.Write) || event.Has(inotify.Create) { + fw.handleEvent(event.Name) + } + if event.Has(inotify.Remove) { + fw.handleRemove(event.Name) + } + case err, ok := <-fw.watcher.Errors: + if !ok { + return fmt.Errorf("watcher error channel closed") + } + fw.log.Warn("inotify error", "err", err) + } + } +} + +// Close shuts down the inotify watcher and cancels pending timers. +func (fw *FSWatcher) Close() { + fw.mu.Lock() + defer fw.mu.Unlock() + for _, timer := range fw.debounce { + timer.Stop() + } + fw.watcher.Close() +} + +func (fw *FSWatcher) handleEvent(path string) { + fw.mu.Lock() + handler, ok := fw.handlers[path] + if ok && fw.pending[path] { + if err := fw.watcher.Add(path); err == nil { + delete(fw.pending, path) + } + } + handlerPath := path + if !ok { + dir := filepath.Dir(path) + handler, ok = fw.dirHandlers[dir] + handlerPath = dir + if !ok { + fw.mu.Unlock() + return + } + } + + if handler.Debounce > 0 { + if timer, exists := fw.debounce[handlerPath]; exists { + timer.Reset(handler.Debounce) + fw.mu.Unlock() + return + } + fw.debounce[handlerPath] = time.AfterFunc(handler.Debounce, func() { + fw.fireHandler(handlerPath, handler) + }) + fw.mu.Unlock() + return + } + fw.mu.Unlock() + fw.fireHandler(handlerPath, handler) +} + +func (fw *FSWatcher) handleRemove(path string) { + fw.mu.Lock() + handler, ok := fw.handlers[path] + if !ok { + dir := filepath.Dir(path) + handler, ok = fw.dirHandlers[dir] + if ok { + fw.mu.Unlock() + fw.fireHandler(dir, handler) + return + } + fw.mu.Unlock() + return + } + fw.mu.Unlock() + + if handler.UseMerge { + fw.fireHandler(path, handler) + } else { + fw.tree.Delete(handler.TreeKey) + fw.log.Debug("fswatcher: removed", "path", path, "key", handler.TreeKey) + } + + if err := fw.watcher.Add(path); err == nil { + return + } + + // Gone for now: watch the directory to catch it being recreated, + // or give up when the directory went away too. + fw.mu.Lock() + defer fw.mu.Unlock() + if err := fw.watcher.Add(filepath.Dir(path)); err != nil { + fw.dropLocked(path) + fw.log.Debug("fswatcher: file and directory gone, handler removed", "path", path) + return + } + fw.pending[path] = true + fw.log.Debug("fswatcher: file gone, waiting for it to reappear", "path", path) +} + +func (fw *FSWatcher) fireHandler(path string, handler WatchHandler) { + data, err := handler.ReadFunc(path) + if err != nil { + fw.log.Warn("fswatcher: read failed", "path", path, "err", err) + return + } + fw.apply(handler, data) + fw.log.Debug("fswatcher: updated", "path", path, "key", handler.TreeKey) +} + +// apply writes a ReadFunc result into the tree. For merge handlers it +// merges; for plain handlers an empty result deletes the key rather than +// writing an empty node -- a collector that has "nothing" to report (e.g. +// the containers feature is enabled but no container is running) must not +// leave a bare subtree behind, or clients would see operational data where +// the feature is effectively absent. +func (fw *FSWatcher) apply(handler WatchHandler, data json.RawMessage) { + switch { + case handler.UseMerge: + fw.tree.Merge(handler.TreeKey, data) + case len(data) == 0: + fw.tree.Delete(handler.TreeKey) + default: + fw.tree.Set(handler.TreeKey, data) + } +} diff --git a/src/yangerd/internal/fswatcher/fswatcher_test.go b/src/yangerd/internal/fswatcher/fswatcher_test.go new file mode 100644 index 000000000..99f1e96fe --- /dev/null +++ b/src/yangerd/internal/fswatcher/fswatcher_test.go @@ -0,0 +1,739 @@ +package fswatcher + +import ( + "context" + "encoding/json" + "fmt" + "log/slog" + "os" + "path/filepath" + "strings" + "testing" + "time" + + "github.com/kernelkit/infix/src/yangerd/internal/tree" +) + +func newTestFSWatcher(t *testing.T) (*FSWatcher, *tree.Tree) { + t.Helper() + tr := tree.New() + fw, err := New(tr, slog.Default()) + if err != nil { + t.Fatalf("New: %v", err) + } + t.Cleanup(func() { fw.Close() }) + return fw, tr +} + +func TestNew(t *testing.T) { + tr := tree.New() + fw, err := New(tr, slog.Default()) + if err != nil { + t.Fatalf("New: %v", err) + } + defer fw.Close() + + if fw.tree != tr { + t.Error("tree not stored") + } + if fw.handlers == nil { + t.Error("handlers map nil") + } + if fw.debounce == nil { + t.Error("debounce map nil") + } +} + +func TestWatch(t *testing.T) { + fw, _ := newTestFSWatcher(t) + + tmp := t.TempDir() + path := filepath.Join(tmp, "test.txt") + if err := os.WriteFile(path, []byte("hello"), 0644); err != nil { + t.Fatal(err) + } + + handler := WatchHandler{ + TreeKey: "test/key", + ReadFunc: func(p string) (json.RawMessage, error) { return json.RawMessage(`"ok"`), nil }, + } + + if err := fw.Watch(path, handler); err != nil { + t.Fatalf("Watch: %v", err) + } + + fw.mu.Lock() + _, ok := fw.handlers[path] + fw.mu.Unlock() + if !ok { + t.Error("handler not registered") + } +} + +func TestInitialRead(t *testing.T) { + fw, tr := newTestFSWatcher(t) + + tmp := t.TempDir() + p1 := filepath.Join(tmp, "a.txt") + p2 := filepath.Join(tmp, "b.txt") + os.WriteFile(p1, []byte("1"), 0644) + os.WriteFile(p2, []byte("2"), 0644) + + fw.Watch(p1, WatchHandler{ + TreeKey: "key/a", + ReadFunc: func(path string) (json.RawMessage, error) { + return json.RawMessage(`"value-a"`), nil + }, + }) + fw.Watch(p2, WatchHandler{ + TreeKey: "key/b", + ReadFunc: func(path string) (json.RawMessage, error) { + return json.RawMessage(`"value-b"`), nil + }, + }) + + fw.InitialRead() + + if got := tr.Get("key/a"); string(got) != `"value-a"` { + t.Errorf("key/a = %s, want %q", got, `"value-a"`) + } + if got := tr.Get("key/b"); string(got) != `"value-b"` { + t.Errorf("key/b = %s, want %q", got, `"value-b"`) + } +} + +func TestInitialReadError(t *testing.T) { + fw, tr := newTestFSWatcher(t) + + tmp := t.TempDir() + p := filepath.Join(tmp, "fail.txt") + os.WriteFile(p, []byte("x"), 0644) + + fw.Watch(p, WatchHandler{ + TreeKey: "key/fail", + ReadFunc: func(path string) (json.RawMessage, error) { + return nil, fmt.Errorf("read error") + }, + }) + + fw.InitialRead() + + if got := tr.Get("key/fail"); got != nil { + t.Errorf("expected nil for failed read, got %s", got) + } +} + +func TestSyncGlob(t *testing.T) { + fw, tr := newTestFSWatcher(t) + + tmp := t.TempDir() + for _, name := range []string{"x1.conf", "x2.conf"} { + os.WriteFile(filepath.Join(tmp, name), []byte("data"), 0644) + } + os.WriteFile(filepath.Join(tmp, "y.txt"), []byte("data"), 0644) + + handler := WatchHandler{ + TreeKey: "glob/test", + ReadFunc: func(p string) (json.RawMessage, error) { return json.RawMessage(`"g"`), nil }, + } + pattern := filepath.Join(tmp, "x*.conf") + + if added, removed, err := fw.SyncGlob(pattern, handler); err != nil || added != 2 || removed != 0 { + t.Fatalf("first sync = +%d -%d %v, want +2 -0", added, removed, err) + } + if got := tr.Get("glob/test"); string(got) != `"g"` { + t.Fatalf("new matches must be read at once, got %s", got) + } + + os.Remove(filepath.Join(tmp, "x1.conf")) + os.WriteFile(filepath.Join(tmp, "x3.conf"), []byte("data"), 0644) + if added, removed, err := fw.SyncGlob(pattern, handler); err != nil || added != 1 || removed != 1 { + t.Fatalf("second sync = +%d -%d %v, want +1 -1", added, removed, err) + } + + fw.mu.Lock() + _, has1 := fw.handlers[filepath.Join(tmp, "x1.conf")] + _, has3 := fw.handlers[filepath.Join(tmp, "x3.conf")] + count := len(fw.handlers) + fw.mu.Unlock() + if has1 || !has3 || count != 2 { + t.Errorf("handlers after sync: x1=%v x3=%v count=%d", has1, has3, count) + } +} + +// A file removed and later recreated is picked up again. +func TestRunFileReappears(t *testing.T) { + fw, tr := newTestFSWatcher(t) + + tmp := t.TempDir() + path := filepath.Join(tmp, "hostname") + os.WriteFile(path, []byte("one"), 0644) + + fw.Watch(path, WatchHandler{ + TreeKey: "host", + ReadFunc: func(p string) (json.RawMessage, error) { + b, err := os.ReadFile(p) + if err != nil { + return nil, err + } + return json.RawMessage(fmt.Sprintf("%q", b)), nil + }, + }) + fw.InitialRead() + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + go fw.Run(ctx) + + os.Remove(path) + deadline := time.Now().Add(2 * time.Second) + for tr.Get("host") != nil && time.Now().Before(deadline) { + time.Sleep(10 * time.Millisecond) + } + + os.WriteFile(path, []byte("two"), 0644) + for time.Now().Before(deadline.Add(2 * time.Second)) { + if string(tr.Get("host")) == `"two"` { + return + } + time.Sleep(10 * time.Millisecond) + } + t.Fatalf("recreated file not picked up, have %s", tr.Get("host")) +} + +func TestRunWriteEvent(t *testing.T) { + fw, tr := newTestFSWatcher(t) + + tmp := t.TempDir() + path := filepath.Join(tmp, "watched.txt") + os.WriteFile(path, []byte("initial"), 0644) + + callCount := 0 + fw.Watch(path, WatchHandler{ + TreeKey: "run/test", + ReadFunc: func(p string) (json.RawMessage, error) { + callCount++ + return json.RawMessage(`"updated"`), nil + }, + }) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + errCh := make(chan error, 1) + go func() { errCh <- fw.Run(ctx) }() + + time.Sleep(50 * time.Millisecond) + + os.WriteFile(path, []byte("changed"), 0644) + + deadline := time.After(2 * time.Second) + for { + if got := tr.Get("run/test"); string(got) == `"updated"` { + break + } + select { + case <-deadline: + t.Fatal("timed out waiting for tree update after write event") + default: + time.Sleep(10 * time.Millisecond) + } + } + + cancel() + err := <-errCh + if err != nil && err != context.Canceled { + t.Errorf("Run returned unexpected error: %v", err) + } +} + +func TestFireHandler(t *testing.T) { + fw, tr := newTestFSWatcher(t) + + handler := WatchHandler{ + TreeKey: "fire/test", + ReadFunc: func(path string) (json.RawMessage, error) { + return json.RawMessage(`{"fired":true}`), nil + }, + } + + fw.fireHandler("/fake/path", handler) + + if got := tr.Get("fire/test"); string(got) != `{"fired":true}` { + t.Errorf("tree value = %s, want %s", got, `{"fired":true}`) + } +} + +func TestFireHandlerReadError(t *testing.T) { + fw, tr := newTestFSWatcher(t) + + handler := WatchHandler{ + TreeKey: "fire/err", + ReadFunc: func(path string) (json.RawMessage, error) { + return nil, fmt.Errorf("broken") + }, + } + + fw.fireHandler("/fake/path", handler) + + if got := tr.Get("fire/err"); got != nil { + t.Errorf("expected nil for errored handler, got %s", got) + } +} + +// A plain (non-merge) handler that returns an empty result must delete its +// key, not leave a stale or empty node behind -- e.g. the containers +// collector returns nil when no container is running. +func TestFireHandlerEmptyDeletes(t *testing.T) { + fw, tr := newTestFSWatcher(t) + + tr.Set("fire/gone", json.RawMessage(`{"container":[{"name":"old"}]}`)) + + empty := json.RawMessage(nil) + handler := WatchHandler{ + TreeKey: "fire/gone", + ReadFunc: func(string) (json.RawMessage, error) { return empty, nil }, + } + + fw.fireHandler("/fake/path", handler) + + if got := tr.Get("fire/gone"); got != nil { + t.Errorf("expected key deleted on empty result, got %s", got) + } +} + +func TestDebounce(t *testing.T) { + fw, tr := newTestFSWatcher(t) + + tmp := t.TempDir() + path := filepath.Join(tmp, "debounce.txt") + os.WriteFile(path, []byte("init"), 0644) + + callCount := 0 + fw.Watch(path, WatchHandler{ + TreeKey: "debounce/test", + Debounce: 100 * time.Millisecond, + ReadFunc: func(p string) (json.RawMessage, error) { + callCount++ + return json.RawMessage(fmt.Sprintf(`"call-%d"`, callCount)), nil + }, + }) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + go fw.Run(ctx) + time.Sleep(50 * time.Millisecond) + + for i := 0; i < 5; i++ { + os.WriteFile(path, []byte(fmt.Sprintf("data-%d", i)), 0644) + time.Sleep(10 * time.Millisecond) + } + + time.Sleep(300 * time.Millisecond) + + got := tr.Get("debounce/test") + if got == nil { + t.Fatal("tree not updated after debounced writes") + } + + if callCount > 3 { + t.Errorf("expected debounce to coalesce writes, but handler called %d times", callCount) + } + + cancel() +} + +func TestRunContextCancellation(t *testing.T) { + fw, _ := newTestFSWatcher(t) + + ctx, cancel := context.WithCancel(context.Background()) + + errCh := make(chan error, 1) + go func() { errCh <- fw.Run(ctx) }() + + time.Sleep(20 * time.Millisecond) + cancel() + + err := <-errCh + if err != context.Canceled { + t.Errorf("Run error = %v, want context.Canceled", err) + } +} + +func TestClose(t *testing.T) { + tr := tree.New() + fw, err := New(tr, slog.Default()) + if err != nil { + t.Fatal(err) + } + + tmp := t.TempDir() + path := filepath.Join(tmp, "close.txt") + os.WriteFile(path, []byte("x"), 0644) + + fw.Watch(path, WatchHandler{ + TreeKey: "close/test", + Debounce: time.Second, + ReadFunc: func(p string) (json.RawMessage, error) { return json.RawMessage(`"x"`), nil }, + }) + + fw.handleEvent(path) + + fw.mu.Lock() + timerCount := len(fw.debounce) + fw.mu.Unlock() + if timerCount != 1 { + t.Errorf("expected 1 debounce timer, got %d", timerCount) + } + + fw.Close() +} + +func TestHandleRemoveMergeHandler(t *testing.T) { + fw, tr := newTestFSWatcher(t) + + tmp := t.TempDir() + path := filepath.Join(tmp, "forwarding") + os.WriteFile(path, []byte("1"), 0644) + + fw.Watch(path, WatchHandler{ + TreeKey: "routing", + ReadFunc: func(_ string) (json.RawMessage, error) { + return json.RawMessage(`{"interfaces":{"interface":["eth0"]}}`), nil + }, + UseMerge: true, + }) + + fw.InitialRead() + + got := tr.Get("routing") + if got == nil { + t.Fatal("tree not populated after InitialRead") + } + + os.Remove(path) + fw.handleRemove(path) + + got = tr.Get("routing") + if got == nil { + t.Fatal("tree entry should still exist after merge-remove") + } + if string(got) != `{"interfaces":{"interface":["eth0"]}}` { + t.Errorf("got %s, want updated merge data", got) + } + + fw.mu.Lock() + pending := fw.pending[path] + fw.mu.Unlock() + if !pending { + t.Error("handler should wait for the file to reappear") + } +} + +func TestHandleRemovePlainHandler(t *testing.T) { + fw, tr := newTestFSWatcher(t) + + tmp := t.TempDir() + path := filepath.Join(tmp, "value.txt") + os.WriteFile(path, []byte("data"), 0644) + + fw.Watch(path, WatchHandler{ + TreeKey: "plain/key", + ReadFunc: func(p string) (json.RawMessage, error) { + return json.RawMessage(`"hello"`), nil + }, + }) + + fw.InitialRead() + + if got := tr.Get("plain/key"); string(got) != `"hello"` { + t.Fatalf("initial = %s, want %q", got, `"hello"`) + } + + os.Remove(path) + fw.handleRemove(path) + + if got := tr.Get("plain/key"); got != nil { + t.Errorf("tree entry should be deleted after remove, got %s", got) + } + + fw.mu.Lock() + pending := fw.pending[path] + fw.mu.Unlock() + if !pending { + t.Error("handler should wait for the file to reappear") + } + + os.RemoveAll(tmp) + fw.handleRemove(path) + fw.mu.Lock() + _, handlerExists := fw.handlers[path] + fw.mu.Unlock() + if handlerExists { + t.Error("handler should be cleaned up once its directory is gone too") + } +} + +func TestHandleRemoveUnknownPath(t *testing.T) { + fw, _ := newTestFSWatcher(t) + fw.handleRemove("/nonexistent/path") +} + +func TestHandleRemoveRewatchSuccess(t *testing.T) { + fw, tr := newTestFSWatcher(t) + + tmp := t.TempDir() + path := filepath.Join(tmp, "ephemeral.txt") + os.WriteFile(path, []byte("1"), 0644) + + calls := 0 + fw.Watch(path, WatchHandler{ + TreeKey: "ephem", + ReadFunc: func(_ string) (json.RawMessage, error) { + calls++ + return json.RawMessage(fmt.Sprintf(`"v%d"`, calls)), nil + }, + UseMerge: true, + }) + + fw.InitialRead() + + fw.handleRemove(path) + + fw.mu.Lock() + _, handlerExists := fw.handlers[path] + fw.mu.Unlock() + if !handlerExists { + t.Error("handler should still exist when file still exists (rewatch succeeds)") + } + + got := tr.Get("ephem") + if string(got) != `"v2"` { + t.Errorf("got %s, want %q (handler should have been called again)", got, `"v2"`) + } +} + +func TestWatchSymlink(t *testing.T) { + fw, _ := newTestFSWatcher(t) + + tmp := t.TempDir() + targetA := filepath.Join(tmp, "target-a") + targetB := filepath.Join(tmp, "target-b") + link := filepath.Join(tmp, "link") + os.WriteFile(targetA, []byte("a"), 0644) + os.WriteFile(targetB, []byte("b"), 0644) + os.Symlink(targetA, link) + + handler := WatchHandler{ + TreeKey: "sym/test", + ReadFunc: func(p string) (json.RawMessage, error) { return json.RawMessage(`"sym"`), nil }, + } + + if err := fw.WatchSymlink(link, handler); err != nil { + t.Fatalf("WatchSymlink: %v", err) + } + + fw.mu.Lock() + _, ok := fw.handlers[link] + fw.mu.Unlock() + if !ok { + t.Error("handler not registered under symlink path") + } +} + +func TestWatchSymlinkReplace(t *testing.T) { + fw, tr := newTestFSWatcher(t) + + tmp := t.TempDir() + targetA := filepath.Join(tmp, "zone-a") + targetB := filepath.Join(tmp, "zone-b") + link := filepath.Join(tmp, "current") + os.WriteFile(targetA, []byte("a"), 0644) + os.WriteFile(targetB, []byte("b"), 0644) + os.Symlink(targetA, link) + + calls := 0 + fw.WatchSymlink(link, WatchHandler{ + TreeKey: "sym/replace", + ReadFunc: func(p string) (json.RawMessage, error) { + calls++ + target, _ := os.Readlink(p) + return json.RawMessage(fmt.Sprintf(`"target-%d-%s"`, calls, filepath.Base(target))), nil + }, + }) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + go fw.Run(ctx) + time.Sleep(50 * time.Millisecond) + + os.Remove(link) + os.Symlink(targetB, link) + + deadline := time.After(2 * time.Second) + for { + got := tr.Get("sym/replace") + if got != nil && strings.Contains(string(got), "zone-b") { + break + } + select { + case <-deadline: + t.Fatalf("timed out waiting for symlink replace event; tree = %s", tr.Get("sym/replace")) + default: + time.Sleep(10 * time.Millisecond) + } + } + + cancel() +} + +// fw_setenv (U-Boot) and grub-editenv rewrite the env via a temp file + +// atomic rename, so the env gets a new inode. A direct file watch misses +// that; the parent-directory watch used for boot-order must catch it, +// otherwise operational boot-order stays stale until the next reboot. +func TestWatchSymlinkAtomicRename(t *testing.T) { + fw, tr := newTestFSWatcher(t) + + tmp := t.TempDir() + env := filepath.Join(tmp, "uboot.env") + os.WriteFile(env, []byte("BOOT_ORDER=net\n"), 0644) + + fw.WatchSymlink(env, WatchHandler{ + TreeKey: "boot/env", + ReadFunc: func(p string) (json.RawMessage, error) { + data, _ := os.ReadFile(p) + return json.RawMessage(fmt.Sprintf("%q", strings.TrimSpace(string(data)))), nil + }, + }) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + go fw.Run(ctx) + time.Sleep(50 * time.Millisecond) + + // Replace the file the way fw_setenv does: write a temp, rename over. + tmpEnv := env + ".tmp" + os.WriteFile(tmpEnv, []byte("BOOT_ORDER=primary net\n"), 0644) + if err := os.Rename(tmpEnv, env); err != nil { + t.Fatalf("rename: %v", err) + } + + deadline := time.After(2 * time.Second) + for { + if got := tr.Get("boot/env"); got != nil && strings.Contains(string(got), "primary net") { + break + } + select { + case <-deadline: + t.Fatalf("timed out waiting for atomic-rename event; tree = %s", tr.Get("boot/env")) + default: + time.Sleep(10 * time.Millisecond) + } + } + + cancel() +} + +func TestWatchDir(t *testing.T) { + fw, tr := newTestFSWatcher(t) + + tmp := t.TempDir() + os.WriteFile(filepath.Join(tmp, "a.keys"), []byte("key-a"), 0644) + + fw.WatchDir(tmp, WatchHandler{ + TreeKey: "dir/test", + ReadFunc: func(dir string) (json.RawMessage, error) { + entries, _ := os.ReadDir(dir) + names := make([]string, 0, len(entries)) + for _, e := range entries { + names = append(names, e.Name()) + } + return json.Marshal(map[string]interface{}{"files": names}) + }, + Debounce: 50 * time.Millisecond, + UseMerge: true, + }) + + fw.InitialRead() + got := tr.Get("dir/test") + if got == nil { + t.Fatal("tree not populated after InitialRead for dir handler") + } + if !strings.Contains(string(got), "a.keys") { + t.Fatalf("initial read missing a.keys: %s", got) + } + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + go fw.Run(ctx) + time.Sleep(50 * time.Millisecond) + + os.WriteFile(filepath.Join(tmp, "b.keys"), []byte("key-b"), 0644) + + deadline := time.After(2 * time.Second) + for { + got = tr.Get("dir/test") + if got != nil && strings.Contains(string(got), "b.keys") { + break + } + select { + case <-deadline: + t.Fatalf("timed out waiting for dir event; tree = %s", tr.Get("dir/test")) + default: + time.Sleep(10 * time.Millisecond) + } + } + + cancel() +} + +func TestWatchDirRemoveFile(t *testing.T) { + fw, tr := newTestFSWatcher(t) + + tmp := t.TempDir() + os.WriteFile(filepath.Join(tmp, "x.keys"), []byte("data"), 0644) + os.WriteFile(filepath.Join(tmp, "y.keys"), []byte("data"), 0644) + + fw.WatchDir(tmp, WatchHandler{ + TreeKey: "dir/rm", + ReadFunc: func(dir string) (json.RawMessage, error) { + entries, _ := os.ReadDir(dir) + names := make([]string, 0, len(entries)) + for _, e := range entries { + names = append(names, e.Name()) + } + return json.Marshal(map[string]interface{}{"files": names}) + }, + UseMerge: true, + }) + + fw.InitialRead() + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + go fw.Run(ctx) + time.Sleep(50 * time.Millisecond) + + os.Remove(filepath.Join(tmp, "x.keys")) + + deadline := time.After(2 * time.Second) + for { + got := tr.Get("dir/rm") + if got != nil && !strings.Contains(string(got), "x.keys") && strings.Contains(string(got), "y.keys") { + break + } + select { + case <-deadline: + t.Fatalf("timed out waiting for dir remove event; tree = %s", tr.Get("dir/rm")) + default: + time.Sleep(10 * time.Millisecond) + } + } + + cancel() +} diff --git a/src/yangerd/internal/iface/iface.go b/src/yangerd/internal/iface/iface.go new file mode 100644 index 000000000..156b62688 --- /dev/null +++ b/src/yangerd/internal/iface/iface.go @@ -0,0 +1,983 @@ +// Package iface transforms raw `ip -json` data into YANG-shaped +// ietf-interfaces JSON. +package iface + +import ( + "bytes" + "encoding/json" + "fmt" + "sort" + "strconv" + "strings" +) + +// FileChecker abstracts filesystem probes needed during interface transformation. +type FileChecker interface { + Exists(path string) bool + ReadFile(path string) (string, error) + // ListDir returns the entry names of a directory, nil if unreadable. + ListDir(path string) []string +} + +// Transform converts raw `ip -json` link/address/statistics arrays into +// `{"interface":[...]}`. The caller (NLMonitor) stores this at tree key +// "ietf-interfaces:interfaces"; the IPC server adds the module-qualified +// wrapper when responding to clients. +// +// neighData is the output of `ip -json neigh show` — an array of objects +// with keys: dst, dev, lladdr, state (array of strings like "REACHABLE", +// "STALE", "PERMANENT", etc.). May be nil if unavailable. +func Transform(linkData, addrData, neighData json.RawMessage, fc FileChecker) json.RawMessage { + links := decodeObjects(linkData) + addrs := decodeObjects(addrData) + neighs := decodeObjects(neighData) + + // By ifindex, not name: a name can be reused by a new interface + // while rows of the old one are still on their way out. + addrByIndex := make(map[int]map[string]any, len(addrs)) + addrByName := make(map[string]map[string]any) + for _, addr := range addrs { + if idx := getIntOrZero(addr, "ifindex"); idx != 0 { + addrByIndex[idx] = addr + } else if ifname := getString(addr, "ifname"); ifname != "" { + addrByName[ifname] = addr + } + } + + neighByName := make(map[string][]map[string]any) + for _, n := range neighs { + dev := getString(n, "dev") + if dev == "" { + continue + } + neighByName[dev] = append(neighByName[dev], n) + } + + linkByName := make(map[string]map[string]any, len(links)) + for _, link := range links { + linkByName[getString(link, "ifname")] = link + } + + interfaces := make([]map[string]any, 0, len(links)) + for _, iplink := range links { + // A row without name or index is not an interface, whatever + // ip printed it for; the list key must not come out empty. + if getString(iplink, "ifname") == "" || getIntOrZero(iplink, "ifindex") == 0 { + continue + } + if skipInterface(iplink) { + continue + } + + ifname := getString(iplink, "ifname") + ipaddr, ok := addrByIndex[getIntOrZero(iplink, "ifindex")] + if !ok { + ipaddr, ok = addrByName[ifname] + } + if !ok { + ipaddr = map[string]any{} + } + + iface := interfaceCommon(iplink, ipaddr, neighByName[ifname], fc) + yangType := getString(iface, "type") + + higher, lower := layers(ifname, linkByName, fc) + if len(higher) > 0 { + iface["higher-layer-if"] = higher + } + if len(lower) > 0 { + iface["lower-layer-if"] = lower + } + + switch yangType { + case "infix-if-type:vlan": + if v := vlanAugment(iplink); len(v) > 0 { + iface["infix-interfaces:vlan"] = v + } + case "infix-if-type:veth": + if v := vethAugment(iplink); len(v) > 0 { + iface["infix-interfaces:veth"] = v + } + case "infix-if-type:gre", "infix-if-type:gretap": + if v := greAugment(iplink); len(v) > 0 { + iface["infix-interfaces:gre"] = v + } + case "infix-if-type:vxlan": + if v := vxlanAugment(iplink); len(v) > 0 { + iface["infix-interfaces:vxlan"] = v + } + case "infix-if-type:lag": + if v := lagAugment(iplink); len(v) > 0 { + iface["infix-interfaces:lag"] = v + } + } + + switch iplink2yangLower(iplink) { + case "infix-interfaces:bridge-port": + if lower := bridgePortLower(iplink); len(lower) > 0 { + iface["infix-interfaces:bridge-port"] = lower + } + case "infix-interfaces:lag-port": + if lower := lagPortLower(iplink); len(lower) > 0 { + iface["infix-interfaces:lag-port"] = lower + } + } + + interfaces = append(interfaces, iface) + } + + out := map[string]any{ + "interface": interfaces, + } + + raw, err := json.Marshal(out) + if err != nil { + return json.RawMessage(`{"interface":[]}`) + } + + return raw +} + +func decodeObjects(raw json.RawMessage) []map[string]any { + if len(raw) == 0 { + return nil + } + + // Numbers stay json.Number: counter64 values above 2^53 must not + // pass through a float64. + dec := json.NewDecoder(bytes.NewReader(raw)) + dec.UseNumber() + var entries []any + if err := dec.Decode(&entries); err != nil { + return nil + } + + out := make([]map[string]any, 0, len(entries)) + for _, entry := range entries { + obj, ok := asMap(entry) + if !ok { + continue + } + out = append(out, obj) + } + + return out +} + +func skipInterface(iplink map[string]any) bool { + if getString(iplink, "group") == "internal" { + return true + } + + switch getString(iplink, "link_type") { + case "can", "vcan": + return true + default: + return false + } +} + +// layers returns the directly adjacent interfaces, from the kernel's +// sysfs upper_/lower_ links. Neighbors not shown in operational are +// not referenced either. +func layers(ifname string, links map[string]map[string]any, fc FileChecker) ([]string, []string) { + var higher, lower []string + + if fc == nil { + return nil, nil + } + + visible := func(name string) bool { + link, ok := links[name] + return ok && !skipInterface(link) + } + + for _, entry := range fc.ListDir("/sys/class/net/" + ifname) { + switch { + case strings.HasPrefix(entry, "upper_"): + if name := strings.TrimPrefix(entry, "upper_"); visible(name) { + higher = append(higher, name) + } + case strings.HasPrefix(entry, "lower_"): + if name := strings.TrimPrefix(entry, "lower_"); visible(name) { + lower = append(lower, name) + } + } + } + + sort.Strings(higher) + sort.Strings(lower) + return higher, lower +} + +func interfaceCommon(iplink, ipaddr map[string]any, neighEntries []map[string]any, fc FileChecker) map[string]any { + flags := getStrings(iplink, "flags") + + iface := map[string]any{ + "type": iplink2yangType(iplink, fc), + "name": getString(iplink, "ifname"), + "if-index": getIntOrZero(iplink, "ifindex"), + "admin-status": boolToStatus(contains(flags, "UP"), "up", "down"), + "oper-status": iplink2yangOperstate(getString(iplink, "operstate")), + } + + if _, ok := iplink["ifalias"]; ok { + iface["description"] = getString(iplink, "ifalias") + } + + if !contains(flags, "POINTOPOINT") { + if address, ok := iplink["address"]; ok { + iface["phys-address"] = fmt.Sprintf("%v", address) + } + } + + if stats := statistics(iplink); len(stats) > 0 { + iface["statistics"] = stats + } + + if ipv4 := ipv4Data(ipaddr, neighEntries); len(ipv4) > 0 { + iface["ietf-ip:ipv4"] = ipv4 + } + + if ipv6 := ipv6Data(ipaddr, neighEntries, fc); len(ipv6) > 0 { + iface["ietf-ip:ipv6"] = ipv6 + } + + return iface +} + +// IsEthernet reports whether an `ip -json link` row is reported as an +// ethernet interface, the only type that carries ethtool data. +func IsEthernet(row json.RawMessage, fc FileChecker) bool { + rows := decodeObjects(row) + return len(rows) > 0 && iplink2yangType(rows[0], fc) == "infix-if-type:ethernet" +} + +func iplink2yangType(iplink map[string]any, fc FileChecker) string { + ifname := getString(iplink, "ifname") + + switch getString(iplink, "link_type") { + case "loopback": + return "infix-if-type:loopback" + case "gre", "gre6": + return "infix-if-type:gre" + case "ether": + if fc != nil { + if fc.Exists(fmt.Sprintf("/sys/class/net/%s/wireless/", ifname)) { + return "infix-if-type:wifi" + } + } + case "none": + default: + return "infix-if-type:other" + } + + linkinfo, _ := asMap(iplink["linkinfo"]) + switch getString(linkinfo, "info_kind") { + case "bond": + return "infix-if-type:lag" + case "bridge": + return "infix-if-type:bridge" + case "dummy": + return "infix-if-type:dummy" + case "gretap", "ip6gretap": + return "infix-if-type:gretap" + case "vxlan": + return "infix-if-type:vxlan" + case "veth": + return "infix-if-type:veth" + case "vlan": + return "infix-if-type:vlan" + case "wireguard": + return "infix-if-type:wireguard" + default: + return "infix-if-type:ethernet" + } +} + +func iplink2yangLower(iplink map[string]any) string { + linkinfo, _ := asMap(iplink["linkinfo"]) + switch getString(linkinfo, "info_slave_kind") { + case "bridge": + return "infix-interfaces:bridge-port" + case "bond": + return "infix-interfaces:lag-port" + default: + return "" + } +} + +func iplink2yangOperstate(oper string) string { + switch oper { + case "DOWN": + return "down" + case "UP": + return "up" + case "DORMANT": + return "dormant" + case "TESTING": + return "testing" + case "LOWERLAYERDOWN": + return "lower-layer-down" + case "NOTPRESENT": + return "not-present" + default: + return "unknown" + } +} + +func statistics(iplink map[string]any) map[string]any { + out := map[string]any{} + + stats64, _ := asMap(iplink["stats64"]) + rx, _ := asMap(stats64["rx"]) + tx, _ := asMap(stats64["tx"]) + + if octets, ok := rx["bytes"]; ok && isTruthy(octets) { + out["in-octets"] = toCounterString(octets) + } + + if octets, ok := tx["bytes"]; ok && isTruthy(octets) { + out["out-octets"] = toCounterString(octets) + } + + return out +} + +func ipv4Data(ipaddr map[string]any, neighEntries []map[string]any) map[string]any { + if len(ipaddr) == 0 && len(neighEntries) == 0 { + return nil + } + + out := map[string]any{} + if len(ipaddr) > 0 { + if mtu, ok := getInt(ipaddr, "mtu"); ok && mtu != 0 && getString(ipaddr, "ifname") != "lo" { + out["mtu"] = mtu + } + + if addr := addresses(ipaddr, "inet"); len(addr) > 0 { + out["address"] = addr + } + } + + if n := neighbors(neighEntries, 4); len(n) > 0 { + out["neighbor"] = n + } + + return out +} + +func ipv6Data(ipaddr map[string]any, neighEntries []map[string]any, fc FileChecker) map[string]any { + if len(ipaddr) == 0 && len(neighEntries) == 0 { + return nil + } + + out := map[string]any{} + if len(ipaddr) > 0 { + ifname := getString(ipaddr, "ifname") + if ifname != "" && fc != nil { + path := fmt.Sprintf("/proc/sys/net/ipv6/conf/%s/mtu", ifname) + if raw, err := fc.ReadFile(path); err == nil { + trimmed := strings.TrimSpace(raw) + if mtu, err := strconv.Atoi(trimmed); err == nil { + out["mtu"] = mtu + } + } + } + + if addr := addresses(ipaddr, "inet6"); len(addr) > 0 { + out["address"] = addr + } + } + + if n := neighbors(neighEntries, 6); len(n) > 0 { + out["neighbor"] = n + } + + return out +} + +func addresses(ipaddr map[string]any, family string) []map[string]any { + addrInfo, ok := ipaddr["addr_info"] + if !ok { + return nil + } + + arr, ok := asArray(addrInfo) + if !ok { + return nil + } + + out := make([]map[string]any, 0, len(arr)) + for _, entry := range arr { + inet, ok := asMap(entry) + if !ok { + continue + } + + if getString(inet, "family") != family { + continue + } + + address := map[string]any{ + "ip": inet["local"], + "prefix-length": getIntOrZero(inet, "prefixlen"), + "origin": inet2yangOrigin(inet), + } + out = append(out, address) + } + + return out +} + +func neighbors(entries []map[string]any, ipVersion int) []map[string]any { + out := make([]map[string]any, 0, len(entries)) + for _, entry := range entries { + dst := getString(entry, "dst") + if dst == "" { + continue + } + + if !neighMatchesFamily(dst, ipVersion) { + continue + } + + lladdr := getString(entry, "lladdr") + if lladdr == "" { + continue + } + + states := getStrings(entry, "state") + origin := "dynamic" + if contains(states, "PERMANENT") { + origin = "static" + } + + neigh := map[string]any{ + "ip": dst, + "link-layer-address": lladdr, + "origin": origin, + } + + if ipVersion == 6 { + if state := neighState(states); state != "" { + neigh["state"] = state + } + if getBool(entry, "router") { + neigh["is-router"] = []any{nil} + } + } + + out = append(out, neigh) + } + + if len(out) == 0 { + return nil + } + return out +} + +func neighMatchesFamily(dst string, ipVersion int) bool { + for i := 0; i < len(dst); i++ { + if dst[i] == '.' { + return ipVersion == 4 + } + if dst[i] == ':' { + return ipVersion == 6 + } + } + return false +} + +func neighState(states []string) string { + xlate := map[string]string{ + "REACHABLE": "reachable", + "STALE": "stale", + "DELAY": "delay", + "PROBE": "probe", + "INCOMPLETE": "incomplete", + } + for _, s := range states { + if v, ok := xlate[s]; ok { + return v + } + } + return "" +} + +func inet2yangOrigin(inet map[string]any) string { + proto := getString(inet, "protocol") + if proto == "kernel_ll" || proto == "kernel_ra" { + if _, ok := inet["stable-privacy"]; ok { + return "random" + } + } + + switch proto { + case "kernel_ll", "kernel_ra": + return "link-layer" + case "static": + return "static" + case "dhcp": + return "dhcp" + case "random": + return "random" + default: + return "other" + } +} + +func vlanAugment(iplink map[string]any) map[string]any { + info := infoData(iplink) + if len(info) == 0 { + return nil + } + + vlan := map[string]any{ + "tag-type": proto2yang(getString(info, "protocol")), + "id": getIntOrZero(info, "id"), + } + + if lower := getString(iplink, "link"); lower != "" { + vlan["lower-layer-if"] = lower + } + + return vlan +} + +func vethAugment(iplink map[string]any) map[string]any { + peer := getString(iplink, "link") + if peer == "" { + return nil + } + + return map[string]any{"peer": peer} +} + +func greAugment(iplink map[string]any) map[string]any { + info := infoData(iplink) + if len(info) == 0 { + return nil + } + + return map[string]any{ + "local": firstAny(info["local"], info["local6"]), + "remote": firstAny(info["remote"], info["remote6"]), + } +} + +func vxlanAugment(iplink map[string]any) map[string]any { + vxlan := greAugment(iplink) + if len(vxlan) == 0 { + return nil + } + + info := infoData(iplink) + if vni, ok := info["id"]; ok { + vxlan["vni"] = vni + } + + return vxlan +} + +func lagAugment(iplink map[string]any) map[string]any { + info := infoData(iplink) + if len(info) == 0 { + return nil + } + + mode := lagMode(getString(info, "mode")) + bond := map[string]any{ + "mode": mode, + "link-monitor": map[string]any{ + "debounce": map[string]any{ + "up": getIntOrZero(info, "updelay"), + "down": getIntOrZero(info, "downdelay"), + }, + }, + } + + if mode == "lacp" { + lacp := map[string]any{ + "mode": boolToStatus(getString(info, "ad_lacp_active") == "on", "active", "passive"), + "rate": getString(info, "ad_lacp_rate"), + "hash": lagHash(getString(info, "xmit_hash_policy")), + } + + adInfo, ok := asMap(info["ad_info"]) + if ok { + if v, ok := adInfo["aggregator"]; ok { + lacp["aggregator-id"] = v + } + if v, ok := adInfo["actor_key"]; ok { + lacp["actor-key"] = v + } + if v, ok := adInfo["partner_key"]; ok { + lacp["partner-key"] = v + } + if v, ok := adInfo["partner_mac"]; ok { + lacp["partner-mac"] = v + } + } + + if v, ok := info["ad_actor_sys_prio"]; ok { + lacp["system-priority"] = v + } + + bond["lacp"] = lacp + } else { + bond["static"] = map[string]any{ + "mode": getString(info, "mode"), + "hash": getString(info, "xmit_hash_policy"), + } + } + + return bond +} + +func bridgePortSTP(info map[string]any) map[string]any { + state := getString(info, "state") + if state == "" { + return map[string]any{} + } + + return map[string]any{ + "cist": map[string]any{ + "state": state, + }, + } +} + +func bridgePortLower(iplink map[string]any) map[string]any { + master := getString(iplink, "master") + if master == "" { + return nil + } + + linkinfo, _ := asMap(iplink["linkinfo"]) + info, _ := asMap(linkinfo["info_slave_data"]) + if len(info) == 0 { + return nil + } + + return map[string]any{ + "bridge": master, + "flood": map[string]any{ + "broadcast": getBool(info, "bcast_flood"), + "unicast": getBool(info, "flood"), + "multicast": getBool(info, "mcast_flood"), + }, + "multicast": map[string]any{ + "fast-leave": getBool(info, "fastleave"), + "router": bridgeRouterMode(getIntOrZero(info, "multicast_router")), + }, + "stp": bridgePortSTP(info), + } +} + +func lagPortLower(iplink map[string]any) map[string]any { + master := getString(iplink, "master") + if master == "" { + return nil + } + + port := map[string]any{"lag": master} + + linkinfo, _ := asMap(iplink["linkinfo"]) + info, _ := asMap(linkinfo["info_slave_data"]) + if len(info) == 0 { + port["state"] = "backup" + port["link-failures"] = 0 + return port + } + + port["state"] = strings.ToLower(getString(info, "state")) + port["link-failures"] = getIntOrZero(info, "link_failure_count") + + if _, ok := info["ad_aggregator_id"]; ok { + port["lacp"] = map[string]any{ + "aggregator-id": info["ad_aggregator_id"], + "actor-state": getString(info, "ad_actor_oper_port_state_str"), + "partner-state": getString(info, "ad_partner_oper_port_state_str"), + } + } + + return port +} + +func infoData(iplink map[string]any) map[string]any { + linkinfo, ok := asMap(iplink["linkinfo"]) + if !ok { + return nil + } + + data, ok := asMap(linkinfo["info_data"]) + if !ok { + return nil + } + + return data +} + +func proto2yang(proto string) string { + switch proto { + case "802.1Q": + return "ieee802-dot1q-types:c-vlan" + case "802.1ad": + return "ieee802-dot1q-types:s-vlan" + default: + return "other" + } +} + +func lagMode(mode string) string { + switch mode { + case "802.3ad": + return "lacp" + case "balance-xor": + return "static" + default: + return "static" + } +} + +func lagHash(hash string) string { + switch hash { + case "layer2": + return "layer2" + case "layer3+4": + return "layer3-4" + case "layer2+3": + return "layer2-3" + case "encap2+3": + return "encap2-3" + case "encap3+4": + return "encap3-4" + case "vlan+srcmac": + return "vlan-srcmac" + default: + return "layer2" + } +} + +func bridgeRouterMode(v int) string { + switch v { + case 0: + return "off" + case 1: + return "auto" + case 2: + return "permanent" + default: + return "UNKNOWN" + } +} + +func getString(obj map[string]any, key string) string { + v, ok := obj[key] + if !ok || v == nil { + return "" + } + + s, ok := v.(string) + if ok { + return s + } + + return fmt.Sprintf("%v", v) +} + +func getInt(obj map[string]any, key string) (int, bool) { + v, ok := obj[key] + if !ok || v == nil { + return 0, false + } + + switch n := v.(type) { + case int: + return n, true + case int8: + return int(n), true + case int16: + return int(n), true + case int32: + return int(n), true + case int64: + return int(n), true + case uint: + return int(n), true + case uint8: + return int(n), true + case uint16: + return int(n), true + case uint32: + return int(n), true + case uint64: + return int(n), true + case float64: + return int(n), true + case json.Number: + i, err := n.Int64() + if err != nil { + return 0, false + } + return int(i), true + case string: + i, err := strconv.Atoi(strings.TrimSpace(n)) + if err != nil { + return 0, false + } + return i, true + default: + return 0, false + } +} + +func getIntOrZero(obj map[string]any, key string) int { + v, ok := getInt(obj, key) + if !ok { + return 0 + } + return v +} + +func getBool(obj map[string]any, key string) bool { + v, ok := obj[key] + if !ok || v == nil { + return false + } + + b, ok := v.(bool) + if ok { + return b + } + + s := strings.ToLower(strings.TrimSpace(fmt.Sprintf("%v", v))) + return s == "1" || s == "true" || s == "on" || s == "yes" +} + +func getStrings(obj map[string]any, key string) []string { + v, ok := obj[key] + if !ok || v == nil { + return nil + } + + if direct, ok := v.([]string); ok { + return direct + } + + arr, ok := asArray(v) + if !ok { + return nil + } + + out := make([]string, 0, len(arr)) + for _, item := range arr { + out = append(out, fmt.Sprintf("%v", item)) + } + return out +} + +func asMap(v any) (map[string]any, bool) { + if v == nil { + return nil, false + } + + m, ok := v.(map[string]any) + if ok { + return m, true + } + + m2, ok := v.(map[string]interface{}) + if ok { + return map[string]any(m2), true + } + + return nil, false +} + +func asArray(v any) ([]any, bool) { + if v == nil { + return nil, false + } + + arr, ok := v.([]any) + if ok { + return arr, true + } + + arr2, ok := v.([]interface{}) + if ok { + return []any(arr2), true + } + + return nil, false +} + +func contains(values []string, needle string) bool { + for _, value := range values { + if value == needle { + return true + } + } + return false +} + +func isTruthy(v any) bool { + if v == nil { + return false + } + + s := strings.TrimSpace(fmt.Sprintf("%v", v)) + return s != "" && s != "0" +} + +func toCounterString(v any) string { + switch n := v.(type) { + case int: + return strconv.FormatInt(int64(n), 10) + case int8: + return strconv.FormatInt(int64(n), 10) + case int16: + return strconv.FormatInt(int64(n), 10) + case int32: + return strconv.FormatInt(int64(n), 10) + case int64: + return strconv.FormatInt(n, 10) + case uint: + return strconv.FormatUint(uint64(n), 10) + case uint8: + return strconv.FormatUint(uint64(n), 10) + case uint16: + return strconv.FormatUint(uint64(n), 10) + case uint32: + return strconv.FormatUint(uint64(n), 10) + case uint64: + return strconv.FormatUint(n, 10) + case float64: + return strconv.FormatInt(int64(n), 10) + case json.Number: + return n.String() + case string: + return strings.TrimSpace(n) + default: + return fmt.Sprintf("%v", v) + } +} + +func firstAny(a, b any) any { + if a != nil { + s := strings.TrimSpace(fmt.Sprintf("%v", a)) + if s != "" { + return a + } + } + return b +} + +func boolToStatus(cond bool, yes, no string) string { + if cond { + return yes + } + return no +} diff --git a/src/yangerd/internal/iface/iface_test.go b/src/yangerd/internal/iface/iface_test.go new file mode 100644 index 000000000..9be5629f5 --- /dev/null +++ b/src/yangerd/internal/iface/iface_test.go @@ -0,0 +1,973 @@ +package iface + +import ( + "encoding/json" + "errors" + "reflect" + "testing" +) + +type mockFileChecker struct { + exists map[string]bool + files map[string]string + readErr map[string]error + dirs map[string][]string +} + +func (m *mockFileChecker) Exists(path string) bool { + if m == nil || m.exists == nil { + return false + } + return m.exists[path] +} + +func (m *mockFileChecker) ReadFile(path string) (string, error) { + if m == nil { + return "", errors.New("nil file checker") + } + if err, ok := m.readErr[path]; ok { + return "", err + } + if v, ok := m.files[path]; ok { + return v, nil + } + return "", errors.New("not found") +} + +func (m *mockFileChecker) ListDir(path string) []string { + if m == nil { + return nil + } + return m.dirs[path] +} + +func mustRaw(t *testing.T, v any) json.RawMessage { + t.Helper() + b, err := json.Marshal(v) + if err != nil { + t.Fatalf("marshal: %v", err) + } + return b +} + +func mustInterfaces(t *testing.T, raw json.RawMessage) []map[string]any { + t.Helper() + + var root map[string]any + if err := json.Unmarshal(raw, &root); err != nil { + t.Fatalf("unmarshal transform output: %v", err) + } + + arr, ok := root["interface"].([]any) + if !ok { + t.Fatalf("missing interface list: %#v", root) + } + + out := make([]map[string]any, 0, len(arr)) + for _, v := range arr { + m, ok := v.(map[string]any) + if !ok { + t.Fatalf("interface entry not object: %T", v) + } + out = append(out, m) + } + + return out +} + +func mustIfaceByName(t *testing.T, ifaces []map[string]any, name string) map[string]any { + t.Helper() + for _, iface := range ifaces { + if iface["name"] == name { + return iface + } + } + t.Fatalf("interface %q not found", name) + return nil +} + +func TestTransformEmptyInputs(t *testing.T) { + tests := []struct { + name string + linkData json.RawMessage + addrData json.RawMessage + stats json.RawMessage + }{ + {name: "nil raw messages"}, + { + name: "empty arrays", + linkData: mustRaw(t, []map[string]any{}), + addrData: mustRaw(t, []map[string]any{}), + stats: mustRaw(t, []map[string]any{}), + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ifaces := mustInterfaces(t, Transform(tt.linkData, tt.addrData, nil, nil)) + if len(ifaces) != 0 { + t.Fatalf("expected empty interface list, got %d", len(ifaces)) + } + }) + } +} + +func TestTransformSingleLoopback(t *testing.T) { + link := []map[string]any{{ + "ifindex": 1, + "ifname": "lo", + "flags": []any{"LOOPBACK", "UP"}, + "link_type": "loopback", + "operstate": "UNKNOWN", + "address": "00:00:00:00:00:00", + "statistics": map[string]any{}, + }} + + ifaces := mustInterfaces(t, Transform(mustRaw(t, link), nil, nil, nil)) + if len(ifaces) != 1 { + t.Fatalf("expected 1 interface, got %d", len(ifaces)) + } + + lo := ifaces[0] + if lo["name"] != "lo" { + t.Fatalf("name = %v", lo["name"]) + } + if lo["type"] != "infix-if-type:loopback" { + t.Fatalf("type = %v", lo["type"]) + } + if lo["admin-status"] != "up" || lo["oper-status"] != "unknown" { + t.Fatalf("admin/oper mismatch: %v/%v", lo["admin-status"], lo["oper-status"]) + } +} + +func TestTransformSingleEthernetWithIPv4IPv6(t *testing.T) { + link := []map[string]any{{ + "ifindex": 2, + "ifname": "eth0", + "flags": []any{"UP"}, + "link_type": "ether", + "operstate": "UP", + "address": "52:54:00:12:34:56", + }} + + addr := []map[string]any{{ + "ifname": "eth0", + "mtu": 1500, + "addr_info": []map[string]any{ + {"family": "inet", "local": "192.0.2.10", "prefixlen": 24, "protocol": "static"}, + {"family": "inet6", "local": "2001:db8::10", "prefixlen": 64, "protocol": "kernel_ra"}, + }, + }} + + fc := &mockFileChecker{files: map[string]string{"/proc/sys/net/ipv6/conf/eth0/mtu": "1400\n"}} + + ifaces := mustInterfaces(t, Transform(mustRaw(t, link), mustRaw(t, addr), nil, fc)) + eth0 := mustIfaceByName(t, ifaces, "eth0") + + if eth0["type"] != "infix-if-type:ethernet" { + t.Fatalf("unexpected type: %v", eth0["type"]) + } + + ipv4, ok := eth0["ietf-ip:ipv4"].(map[string]any) + if !ok { + t.Fatalf("missing ipv4 container: %#v", eth0) + } + if ipv4["mtu"] != float64(1500) { + t.Fatalf("ipv4 mtu = %v", ipv4["mtu"]) + } + v4addrs := ipv4["address"].([]any) + v4 := v4addrs[0].(map[string]any) + if v4["ip"] != "192.0.2.10" || v4["prefix-length"] != float64(24) || v4["origin"] != "static" { + t.Fatalf("unexpected ipv4 address entry: %#v", v4) + } + + ipv6, ok := eth0["ietf-ip:ipv6"].(map[string]any) + if !ok { + t.Fatalf("missing ipv6 container: %#v", eth0) + } + if ipv6["mtu"] != float64(1400) { + t.Fatalf("ipv6 mtu = %v", ipv6["mtu"]) + } + v6addrs := ipv6["address"].([]any) + v6 := v6addrs[0].(map[string]any) + if v6["ip"] != "2001:db8::10" || v6["prefix-length"] != float64(64) || v6["origin"] != "link-layer" { + t.Fatalf("unexpected ipv6 address entry: %#v", v6) + } +} + +func TestTransformStatisticsCountersAsStrings(t *testing.T) { + link := []map[string]any{{ + "ifindex": 3, + "ifname": "eth1", + "flags": []any{"UP"}, + "link_type": "ether", + "operstate": "UP", + "stats64": map[string]any{ + "rx": map[string]any{"bytes": uint64(1234567890)}, + "tx": map[string]any{"bytes": uint64(9876543210)}, + }, + }} + + ifaces := mustInterfaces(t, Transform(mustRaw(t, link), nil, nil, nil)) + eth1 := mustIfaceByName(t, ifaces, "eth1") + st, ok := eth1["statistics"].(map[string]any) + if !ok { + t.Fatalf("missing statistics: %#v", eth1) + } + + if _, ok := st["in-octets"].(string); !ok { + t.Fatalf("in-octets must be string, got %T", st["in-octets"]) + } + if _, ok := st["out-octets"].(string); !ok { + t.Fatalf("out-octets must be string, got %T", st["out-octets"]) + } +} + +func TestTransformVLANAugment(t *testing.T) { + link := []map[string]any{{ + "ifindex": 10, + "ifname": "eth0.100", + "flags": []any{"UP"}, + "link_type": "none", + "operstate": "UP", + "link": "eth0", + "linkinfo": map[string]any{ + "info_kind": "vlan", + "info_data": map[string]any{"protocol": "802.1Q", "id": 100}, + }, + }} + + ifaces := mustInterfaces(t, Transform(mustRaw(t, link), nil, nil, nil)) + vlan := mustIfaceByName(t, ifaces, "eth0.100") + + if vlan["type"] != "infix-if-type:vlan" { + t.Fatalf("type = %v", vlan["type"]) + } + v, ok := vlan["infix-interfaces:vlan"].(map[string]any) + if !ok { + t.Fatalf("missing vlan augment: %#v", vlan) + } + if v["tag-type"] != "ieee802-dot1q-types:c-vlan" || v["id"] != float64(100) || v["lower-layer-if"] != "eth0" { + t.Fatalf("unexpected vlan augment: %#v", v) + } +} + +func TestTransformVethAugment(t *testing.T) { + link := []map[string]any{{ + "ifname": "veth0", + "ifindex": 11, + "flags": []any{"UP"}, + "link_type": "none", + "operstate": "UP", + "link": "veth1", + "linkinfo": map[string]any{"info_kind": "veth"}, + }} + + ifaces := mustInterfaces(t, Transform(mustRaw(t, link), nil, nil, nil)) + veth := mustIfaceByName(t, ifaces, "veth0") + v, ok := veth["infix-interfaces:veth"].(map[string]any) + if !ok || v["peer"] != "veth1" { + t.Fatalf("unexpected veth augment: %#v", veth) + } +} + +func TestTransformGREAndVXLANAugments(t *testing.T) { + link := []map[string]any{ + { + "ifname": "gre1", + "ifindex": 12, + "flags": []any{"UP"}, + "link_type": "gre", + "operstate": "UP", + "linkinfo": map[string]any{ + "info_data": map[string]any{"local": "192.0.2.1", "remote": "198.51.100.1"}, + }, + }, + { + "ifname": "vxlan10", + "ifindex": 13, + "flags": []any{"UP"}, + "link_type": "none", + "operstate": "UP", + "linkinfo": map[string]any{ + "info_kind": "vxlan", + "info_data": map[string]any{"local": "10.0.0.1", "remote": "10.0.0.2", "id": 10}, + }, + }, + } + + ifaces := mustInterfaces(t, Transform(mustRaw(t, link), nil, nil, nil)) + + gre := mustIfaceByName(t, ifaces, "gre1") + if gre["type"] != "infix-if-type:gre" { + t.Fatalf("gre type = %v", gre["type"]) + } + g, ok := gre["infix-interfaces:gre"].(map[string]any) + if !ok || g["local"] != "192.0.2.1" || g["remote"] != "198.51.100.1" { + t.Fatalf("unexpected gre augment: %#v", g) + } + + vx := mustIfaceByName(t, ifaces, "vxlan10") + if vx["type"] != "infix-if-type:vxlan" { + t.Fatalf("vxlan type = %v", vx["type"]) + } + v, ok := vx["infix-interfaces:vxlan"].(map[string]any) + if !ok || v["local"] != "10.0.0.1" || v["remote"] != "10.0.0.2" || v["vni"] != float64(10) { + t.Fatalf("unexpected vxlan augment: %#v", v) + } +} + +func TestTransformLAGAugmentModes(t *testing.T) { + link := []map[string]any{ + { + "ifname": "bond0", + "ifindex": 20, + "flags": []any{"UP"}, + "link_type": "none", + "operstate": "UP", + "linkinfo": map[string]any{ + "info_kind": "bond", + "info_data": map[string]any{ + "mode": "802.3ad", + "updelay": 10, + "downdelay": 20, + "ad_lacp_active": "on", + "ad_lacp_rate": "fast", + "xmit_hash_policy": "layer3+4", + "ad_actor_sys_prio": 100, + "ad_info": map[string]any{ + "aggregator": 7, + "actor_key": 1000, + "partner_key": 2000, + "partner_mac": "02:00:00:00:00:01", + }, + }, + }, + }, + { + "ifname": "bond1", + "ifindex": 21, + "flags": []any{"UP"}, + "link_type": "none", + "operstate": "UP", + "linkinfo": map[string]any{ + "info_kind": "bond", + "info_data": map[string]any{ + "mode": "balance-xor", + "xmit_hash_policy": "layer2", + }, + }, + }, + } + + ifaces := mustInterfaces(t, Transform(mustRaw(t, link), nil, nil, nil)) + + bond0 := mustIfaceByName(t, ifaces, "bond0") + b0 := bond0["infix-interfaces:lag"].(map[string]any) + if b0["mode"] != "lacp" { + t.Fatalf("bond0 mode = %v", b0["mode"]) + } + lacp := b0["lacp"].(map[string]any) + if lacp["mode"] != "active" || lacp["rate"] != "fast" || lacp["hash"] != "layer3-4" { + t.Fatalf("unexpected bond0 lacp: %#v", lacp) + } + + bond1 := mustIfaceByName(t, ifaces, "bond1") + b1 := bond1["infix-interfaces:lag"].(map[string]any) + if b1["mode"] != "static" { + t.Fatalf("bond1 mode = %v", b1["mode"]) + } + static := b1["static"].(map[string]any) + if static["mode"] != "balance-xor" || static["hash"] != "layer2" { + t.Fatalf("unexpected bond1 static: %#v", static) + } +} + +func TestTransformBridgePortLowerLayer(t *testing.T) { + link := []map[string]any{{ + "ifname": "eth2", + "ifindex": 30, + "flags": []any{"UP"}, + "link_type": "ether", + "operstate": "UP", + "master": "br0", + "linkinfo": map[string]any{ + "info_slave_kind": "bridge", + "info_slave_data": map[string]any{ + "bcast_flood": true, + "flood": false, + "mcast_flood": true, + "fastleave": true, + "multicast_router": 2, + }, + }, + }} + + ifaces := mustInterfaces(t, Transform(mustRaw(t, link), nil, nil, nil)) + eth2 := mustIfaceByName(t, ifaces, "eth2") + lower := eth2["infix-interfaces:bridge-port"].(map[string]any) + if lower["bridge"] != "br0" { + t.Fatalf("bridge lower bridge = %v", lower["bridge"]) + } + mcast := lower["multicast"].(map[string]any) + if mcast["router"] != "permanent" { + t.Fatalf("bridge router mode = %v", mcast["router"]) + } +} + +func TestTransformLagPortLowerLayer(t *testing.T) { + link := []map[string]any{{ + "ifname": "eth3", + "ifindex": 31, + "flags": []any{"UP"}, + "link_type": "ether", + "operstate": "UP", + "master": "bond0", + "linkinfo": map[string]any{ + "info_slave_kind": "bond", + "info_slave_data": map[string]any{ + "state": "ACTIVE", + "link_failure_count": 5, + "ad_aggregator_id": 42, + "ad_actor_oper_port_state_str": "collecting_distributing", + "ad_partner_oper_port_state_str": "collecting_distributing", + }, + }, + }} + + ifaces := mustInterfaces(t, Transform(mustRaw(t, link), nil, nil, nil)) + eth3 := mustIfaceByName(t, ifaces, "eth3") + lower := eth3["infix-interfaces:lag-port"].(map[string]any) + if lower["lag"] != "bond0" || lower["state"] != "active" || lower["link-failures"] != float64(5) { + t.Fatalf("unexpected lag-port lower-layer: %#v", lower) + } + lacp := lower["lacp"].(map[string]any) + if lacp["aggregator-id"] != float64(42) { + t.Fatalf("lag-port lacp aggregator-id = %v", lacp["aggregator-id"]) + } +} + +func TestTransformFilteredInterfaces(t *testing.T) { + link := []map[string]any{ + {"ifname": "dummy0", "group": "internal", "link_type": "none"}, + {"ifname": "can0", "link_type": "can"}, + {"ifname": "vcan0", "link_type": "vcan"}, + {"ifname": "eth9", "ifindex": 99, "flags": []any{"UP"}, "link_type": "ether", "operstate": "UP"}, + } + + ifaces := mustInterfaces(t, Transform(mustRaw(t, link), nil, nil, nil)) + if len(ifaces) != 1 { + t.Fatalf("expected only one surviving interface, got %d", len(ifaces)) + } + if ifaces[0]["name"] != "eth9" { + t.Fatalf("surviving interface = %v", ifaces[0]["name"]) + } +} + +func TestTransformWiFiType(t *testing.T) { + link := []map[string]any{{ + "ifname": "wlan0", + "ifindex": 40, + "flags": []any{"UP"}, + "link_type": "ether", + "operstate": "UP", + }} + + fc := &mockFileChecker{exists: map[string]bool{"/sys/class/net/wlan0/wireless/": true}} + ifaces := mustInterfaces(t, Transform(mustRaw(t, link), nil, nil, fc)) + wlan0 := mustIfaceByName(t, ifaces, "wlan0") + if wlan0["type"] != "infix-if-type:wifi" { + t.Fatalf("wlan0 type = %v", wlan0["type"]) + } +} + +func TestIplink2yangTypeMappings(t *testing.T) { + fc := &mockFileChecker{exists: map[string]bool{"/sys/class/net/wlan0/wireless/": true}} + + tests := []struct { + name string + iplink map[string]any + want string + }{ + {name: "loopback", iplink: map[string]any{"ifname": "lo", "link_type": "loopback"}, want: "infix-if-type:loopback"}, + {name: "gre", iplink: map[string]any{"ifname": "gre0", "link_type": "gre"}, want: "infix-if-type:gre"}, + {name: "gre6", iplink: map[string]any{"ifname": "gre6", "link_type": "gre6"}, want: "infix-if-type:gre"}, + {name: "wifi via ether", iplink: map[string]any{"ifname": "wlan0", "link_type": "ether"}, want: "infix-if-type:wifi"}, + {name: "bond", iplink: map[string]any{"ifname": "bond0", "link_type": "none", "linkinfo": map[string]any{"info_kind": "bond"}}, want: "infix-if-type:lag"}, + {name: "bridge", iplink: map[string]any{"ifname": "br0", "link_type": "none", "linkinfo": map[string]any{"info_kind": "bridge"}}, want: "infix-if-type:bridge"}, + {name: "dummy", iplink: map[string]any{"ifname": "dummy0", "link_type": "none", "linkinfo": map[string]any{"info_kind": "dummy"}}, want: "infix-if-type:dummy"}, + {name: "gretap", iplink: map[string]any{"ifname": "gretap0", "link_type": "none", "linkinfo": map[string]any{"info_kind": "gretap"}}, want: "infix-if-type:gretap"}, + {name: "vxlan", iplink: map[string]any{"ifname": "vxlan10", "link_type": "none", "linkinfo": map[string]any{"info_kind": "vxlan"}}, want: "infix-if-type:vxlan"}, + {name: "veth", iplink: map[string]any{"ifname": "veth0", "link_type": "none", "linkinfo": map[string]any{"info_kind": "veth"}}, want: "infix-if-type:veth"}, + {name: "vlan", iplink: map[string]any{"ifname": "eth0.10", "link_type": "none", "linkinfo": map[string]any{"info_kind": "vlan"}}, want: "infix-if-type:vlan"}, + {name: "wireguard", iplink: map[string]any{"ifname": "wg0", "link_type": "none", "linkinfo": map[string]any{"info_kind": "wireguard"}}, want: "infix-if-type:wireguard"}, + {name: "default ethernet", iplink: map[string]any{"ifname": "eth0", "link_type": "none", "linkinfo": map[string]any{"info_kind": "unknown"}}, want: "infix-if-type:ethernet"}, + {name: "unknown link type", iplink: map[string]any{"ifname": "x", "link_type": "strange"}, want: "infix-if-type:other"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := iplink2yangType(tt.iplink, fc) + if got != tt.want { + t.Fatalf("iplink2yangType() = %q, want %q", got, tt.want) + } + }) + } +} + +func TestIplink2yangOperstateMappings(t *testing.T) { + tests := []struct { + in string + want string + }{ + {"DOWN", "down"}, + {"UP", "up"}, + {"DORMANT", "dormant"}, + {"TESTING", "testing"}, + {"LOWERLAYERDOWN", "lower-layer-down"}, + {"NOTPRESENT", "not-present"}, + {"WHATEVER", "unknown"}, + } + + for _, tt := range tests { + if got := iplink2yangOperstate(tt.in); got != tt.want { + t.Fatalf("iplink2yangOperstate(%q) = %q, want %q", tt.in, got, tt.want) + } + } +} + +func TestSkipInterface(t *testing.T) { + tests := []struct { + name string + iplink map[string]any + want bool + }{ + {name: "internal group", iplink: map[string]any{"group": "internal"}, want: true}, + {name: "can", iplink: map[string]any{"link_type": "can"}, want: true}, + {name: "vcan", iplink: map[string]any{"link_type": "vcan"}, want: true}, + {name: "normal", iplink: map[string]any{"link_type": "ether"}, want: false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := skipInterface(tt.iplink); got != tt.want { + t.Fatalf("skipInterface() = %v, want %v", got, tt.want) + } + }) + } +} + +func TestInet2yangOrigin(t *testing.T) { + tests := []struct { + name string + inet map[string]any + want string + }{ + {name: "kernel_ll", inet: map[string]any{"protocol": "kernel_ll"}, want: "link-layer"}, + {name: "kernel_ra", inet: map[string]any{"protocol": "kernel_ra"}, want: "link-layer"}, + {name: "stable privacy kernel_ll", inet: map[string]any{"protocol": "kernel_ll", "stable-privacy": true}, want: "random"}, + {name: "static", inet: map[string]any{"protocol": "static"}, want: "static"}, + {name: "dhcp", inet: map[string]any{"protocol": "dhcp"}, want: "dhcp"}, + {name: "random", inet: map[string]any{"protocol": "random"}, want: "random"}, + {name: "other", inet: map[string]any{"protocol": "kernel_lo"}, want: "other"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := inet2yangOrigin(tt.inet); got != tt.want { + t.Fatalf("inet2yangOrigin() = %q, want %q", got, tt.want) + } + }) + } +} + +func TestProto2yang(t *testing.T) { + tests := []struct { + in string + want string + }{ + {"802.1Q", "ieee802-dot1q-types:c-vlan"}, + {"802.1ad", "ieee802-dot1q-types:s-vlan"}, + {"something", "other"}, + } + + for _, tt := range tests { + if got := proto2yang(tt.in); got != tt.want { + t.Fatalf("proto2yang(%q) = %q, want %q", tt.in, got, tt.want) + } + } +} + +func TestLagMode(t *testing.T) { + tests := []struct { + in string + want string + }{ + {"802.3ad", "lacp"}, + {"balance-xor", "static"}, + {"active-backup", "static"}, + } + + for _, tt := range tests { + if got := lagMode(tt.in); got != tt.want { + t.Fatalf("lagMode(%q) = %q, want %q", tt.in, got, tt.want) + } + } +} + +func TestLagHash(t *testing.T) { + tests := []struct { + in string + want string + }{ + {"layer2", "layer2"}, + {"layer3+4", "layer3-4"}, + {"layer2+3", "layer2-3"}, + {"encap2+3", "encap2-3"}, + {"encap3+4", "encap3-4"}, + {"vlan+srcmac", "vlan-srcmac"}, + {"something-else", "layer2"}, + } + + for _, tt := range tests { + if got := lagHash(tt.in); got != tt.want { + t.Fatalf("lagHash(%q) = %q, want %q", tt.in, got, tt.want) + } + } +} + +func TestBridgeRouterMode(t *testing.T) { + tests := []struct { + in int + want string + }{ + {0, "off"}, + {1, "auto"}, + {2, "permanent"}, + {9, "UNKNOWN"}, + } + + for _, tt := range tests { + if got := bridgeRouterMode(tt.in); got != tt.want { + t.Fatalf("bridgeRouterMode(%d) = %q, want %q", tt.in, got, tt.want) + } + } +} + +func TestStatistics(t *testing.T) { + t.Run("with stats64", func(t *testing.T) { + st := statistics(map[string]any{ + "stats64": map[string]any{ + "rx": map[string]any{"bytes": json.Number("123")}, + "tx": map[string]any{"bytes": uint64(456)}, + }, + }) + if st["in-octets"] != "123" || st["out-octets"] != "456" { + t.Fatalf("unexpected statistics map: %#v", st) + } + }) + + t.Run("without stats64", func(t *testing.T) { + st := statistics(map[string]any{}) + if len(st) != 0 { + t.Fatalf("expected empty stats, got %#v", st) + } + }) +} + +func TestToCounterString(t *testing.T) { + tests := []struct { + name string + in any + want string + }{ + {name: "int", in: int(7), want: "7"}, + {name: "int64", in: int64(8), want: "8"}, + {name: "uint64", in: uint64(9), want: "9"}, + {name: "float64", in: float64(10.9), want: "10"}, + {name: "json number", in: json.Number("11"), want: "11"}, + {name: "string", in: " 12 ", want: "12"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := toCounterString(tt.in); got != tt.want { + t.Fatalf("toCounterString(%v) = %q, want %q", tt.in, got, tt.want) + } + }) + } +} + +func TestAddressesFamilyFilter(t *testing.T) { + ipaddr := map[string]any{ + "addr_info": []any{ + map[string]any{"family": "inet", "local": "192.0.2.1", "prefixlen": 24, "protocol": "dhcp"}, + map[string]any{"family": "inet6", "local": "2001:db8::1", "prefixlen": 64, "protocol": "kernel_ra"}, + }, + } + + v4 := addresses(ipaddr, "inet") + if len(v4) != 1 || v4[0]["ip"] != "192.0.2.1" || v4[0]["prefix-length"] != 24 || v4[0]["origin"] != "dhcp" { + t.Fatalf("unexpected inet addresses: %#v", v4) + } + + v6 := addresses(ipaddr, "inet6") + if len(v6) != 1 || v6[0]["ip"] != "2001:db8::1" || v6[0]["prefix-length"] != 64 || v6[0]["origin"] != "link-layer" { + t.Fatalf("unexpected inet6 addresses: %#v", v6) + } +} + +func TestIPv4Data(t *testing.T) { + t.Run("with mtu and addresses", func(t *testing.T) { + in := map[string]any{ + "ifname": "eth0", + "mtu": 1500, + "addr_info": []any{ + map[string]any{"family": "inet", "local": "10.0.0.1", "prefixlen": 24, "protocol": "static"}, + }, + } + out := ipv4Data(in, nil) + if out["mtu"] != 1500 { + t.Fatalf("unexpected mtu: %#v", out) + } + if _, ok := out["address"]; !ok { + t.Fatalf("missing address list: %#v", out) + } + }) + + t.Run("without mtu", func(t *testing.T) { + in := map[string]any{ + "ifname": "eth0", + "addr_info": []any{ + map[string]any{"family": "inet", "local": "10.0.0.2", "prefixlen": 24, "protocol": "static"}, + }, + } + out := ipv4Data(in, nil) + if _, ok := out["mtu"]; ok { + t.Fatalf("did not expect mtu in %#v", out) + } + }) + + t.Run("loopback omits mtu", func(t *testing.T) { + in := map[string]any{"ifname": "lo", "mtu": 65536} + out := ipv4Data(in, nil) + if _, ok := out["mtu"]; ok { + t.Fatalf("loopback must not include mtu: %#v", out) + } + }) +} + +func TestIPv6Data(t *testing.T) { + t.Run("with mtu and addresses", func(t *testing.T) { + in := map[string]any{ + "ifname": "eth0", + "addr_info": []any{ + map[string]any{"family": "inet6", "local": "2001:db8::1", "prefixlen": 64, "protocol": "static"}, + }, + } + fc := &mockFileChecker{files: map[string]string{"/proc/sys/net/ipv6/conf/eth0/mtu": "1280\n"}} + out := ipv6Data(in, nil, fc) + if out["mtu"] != 1280 { + t.Fatalf("unexpected mtu: %#v", out) + } + if _, ok := out["address"]; !ok { + t.Fatalf("missing address list: %#v", out) + } + }) + + t.Run("without mtu from filechecker", func(t *testing.T) { + in := map[string]any{"ifname": "eth1"} + fc := &mockFileChecker{readErr: map[string]error{"/proc/sys/net/ipv6/conf/eth1/mtu": errors.New("no file")}} + out := ipv6Data(in, nil, fc) + if _, ok := out["mtu"]; ok { + t.Fatalf("did not expect mtu in %#v", out) + } + }) + + t.Run("without addresses", func(t *testing.T) { + in := map[string]any{"ifname": "eth2"} + out := ipv6Data(in, nil, nil) + if len(out) != 0 { + t.Fatalf("expected empty ipv6 map, got %#v", out) + } + }) +} + +func TestNeighbors(t *testing.T) { + t.Run("ipv4 static and dynamic", func(t *testing.T) { + link := []map[string]any{ + {"ifindex": 2, "ifname": "eth0", "flags": []any{"UP"}, "link_type": "ether", "operstate": "UP", "address": "02:00:00:00:00:01"}, + } + neighs := []map[string]any{ + {"dst": "192.168.1.1", "dev": "eth0", "lladdr": "aa:bb:cc:dd:ee:ff", "state": []any{"REACHABLE"}}, + {"dst": "192.168.1.2", "dev": "eth0", "lladdr": "11:22:33:44:55:66", "state": []any{"PERMANENT"}}, + {"dst": "2001:db8::1", "dev": "eth0", "lladdr": "aa:bb:cc:dd:ee:01", "state": []any{"STALE"}}, + } + + ifaces := mustInterfaces(t, Transform(mustRaw(t, link), nil, mustRaw(t, neighs), nil)) + eth0 := mustIfaceByName(t, ifaces, "eth0") + + ipv4, ok := eth0["ietf-ip:ipv4"].(map[string]any) + if !ok { + t.Fatalf("missing ipv4: %#v", eth0) + } + v4neighs, ok := ipv4["neighbor"].([]any) + if !ok || len(v4neighs) != 2 { + t.Fatalf("expected 2 ipv4 neighbors, got %#v", ipv4["neighbor"]) + } + + n0 := v4neighs[0].(map[string]any) + if n0["ip"] != "192.168.1.1" || n0["link-layer-address"] != "aa:bb:cc:dd:ee:ff" || n0["origin"] != "dynamic" { + t.Fatalf("unexpected neighbor[0]: %#v", n0) + } + n1 := v4neighs[1].(map[string]any) + if n1["origin"] != "static" { + t.Fatalf("expected static origin: %#v", n1) + } + }) + + t.Run("ipv6 with state and is-router", func(t *testing.T) { + link := []map[string]any{ + {"ifindex": 2, "ifname": "eth0", "flags": []any{"UP"}, "link_type": "ether", "operstate": "UP", "address": "02:00:00:00:00:01"}, + } + neighs := []map[string]any{ + {"dst": "2001:db8::1", "dev": "eth0", "lladdr": "aa:bb:cc:dd:ee:01", "state": []any{"STALE"}, "router": true}, + } + + ifaces := mustInterfaces(t, Transform(mustRaw(t, link), nil, mustRaw(t, neighs), nil)) + eth0 := mustIfaceByName(t, ifaces, "eth0") + + ipv6, ok := eth0["ietf-ip:ipv6"].(map[string]any) + if !ok { + t.Fatalf("missing ipv6: %#v", eth0) + } + v6neighs, ok := ipv6["neighbor"].([]any) + if !ok || len(v6neighs) != 1 { + t.Fatalf("expected 1 ipv6 neighbor, got %#v", ipv6["neighbor"]) + } + + n := v6neighs[0].(map[string]any) + if n["state"] != "stale" { + t.Fatalf("expected stale state: %#v", n) + } + if _, ok := n["is-router"]; !ok { + t.Fatalf("expected is-router: %#v", n) + } + }) + + t.Run("skips entries without lladdr", func(t *testing.T) { + link := []map[string]any{ + {"ifindex": 2, "ifname": "eth0", "flags": []any{"UP"}, "link_type": "ether", "operstate": "UP", "address": "02:00:00:00:00:01"}, + } + neighs := []map[string]any{ + {"dst": "192.168.1.1", "dev": "eth0", "state": []any{"INCOMPLETE"}}, + } + + ifaces := mustInterfaces(t, Transform(mustRaw(t, link), nil, mustRaw(t, neighs), nil)) + eth0 := mustIfaceByName(t, ifaces, "eth0") + + if _, ok := eth0["ietf-ip:ipv4"]; ok { + t.Fatalf("should not have ipv4 with no valid neighbors: %#v", eth0) + } + }) +} + +// Layers come from sysfs upper_/lower_ links, sorted, and never name an +// interface that operational hides. +func TestLayers(t *testing.T) { + links := mustRaw(t, []map[string]any{ + {"ifname": "br0", "ifindex": 3, "link_type": "ether", "flags": []string{"UP"}, "operstate": "UP", "linkinfo": map[string]any{"info_kind": "bridge"}}, + {"ifname": "eth0", "ifindex": 2, "link_type": "ether", "flags": []string{"UP"}, "operstate": "UP", "master": "br0"}, + {"ifname": "eth0.10", "ifindex": 4, "link_type": "ether", "flags": []string{"UP"}, "operstate": "UP", "link": "eth0", "linkinfo": map[string]any{"info_kind": "vlan"}}, + {"ifname": "dsa0", "ifindex": 5, "link_type": "ether", "group": "internal", "flags": []string{"UP"}, "operstate": "UP"}, + }) + fc := &mockFileChecker{dirs: map[string][]string{ + "/sys/class/net/br0": {"lower_eth0", "brif", "lower_dsa0"}, + "/sys/class/net/eth0": {"upper_eth0.10", "upper_br0", "statistics"}, + "/sys/class/net/eth0.10": {"lower_eth0"}, + }} + + ifaces := mustInterfaces(t, Transform(links, nil, nil, fc)) + + want := map[string][2][]string{ + "br0": {nil, {"eth0"}}, + "eth0": {{"br0", "eth0.10"}, nil}, + "eth0.10": {nil, {"eth0"}}, + } + for name, exp := range want { + entry := mustIfaceByName(t, ifaces, name) + got := [2][]string{toStrings(entry["higher-layer-if"]), toStrings(entry["lower-layer-if"])} + if !reflect.DeepEqual(got, exp) { + t.Errorf("%s layers = %v, want %v", name, got, exp) + } + } +} + +func toStrings(v any) []string { + arr, ok := v.([]any) + if !ok { + return nil + } + out := make([]string, 0, len(arr)) + for _, e := range arr { + out = append(out, e.(string)) + } + return out +} + +// Counters above 2^53 must come out exact. +func TestTransformCounterPrecision(t *testing.T) { + link := json.RawMessage(`[{"ifindex":2,"ifname":"eth0","flags":["UP"],"link_type":"ether","operstate":"UP",` + + `"stats64":{"rx":{"bytes":18446744073709551615},"tx":{"bytes":9007199254740993}}}]`) + eth0 := mustIfaceByName(t, mustInterfaces(t, Transform(link, nil, nil, nil)), "eth0") + st := eth0["statistics"].(map[string]any) + if st["in-octets"] != "18446744073709551615" || st["out-octets"] != "9007199254740993" { + t.Fatalf("counters lost precision: %v / %v", st["in-octets"], st["out-octets"]) + } +} + +func TestIsEthernet(t *testing.T) { + for raw, want := range map[string]bool{ + `[{"ifname":"eth0","link_type":"ether"}]`: true, + `[{"ifname":"e1","link_type":"ether","linkinfo":{"info_kind":"dsa"}}]`: true, + `[{"ifname":"br0","link_type":"ether","linkinfo":{"info_kind":"bridge"}}]`: false, + `[{"ifname":"veth0","link_type":"ether","linkinfo":{"info_kind":"veth"}}]`: false, + `[{"ifname":"lo","link_type":"loopback"}]`: false, + `[]`: false, + } { + if got := IsEthernet(json.RawMessage(raw), nil); got != want { + t.Errorf("IsEthernet(%s) = %v, want %v", raw, got, want) + } + } +} + +// Addresses belong to the interface with the same ifindex, not to any +// interface that happens to share the name. +func TestTransformAddressesByIndex(t *testing.T) { + links := mustRaw(t, []map[string]any{ + {"ifname": "wifi0", "ifindex": 16, "link_type": "ether", "flags": []string{"UP"}, "operstate": "UP"}, + }) + addrs := mustRaw(t, []map[string]any{ + {"ifname": "wifi0", "ifindex": 16, "addr_info": []map[string]any{ + {"family": "inet", "local": "192.168.20.101", "prefixlen": 24, "protocol": "dhcp"}}}, + {"ifname": "wifi0", "ifindex": 14, "addr_info": []map[string]any{}}, + }) + + wifi0 := mustIfaceByName(t, mustInterfaces(t, Transform(links, addrs, nil, &mockFileChecker{})), "wifi0") + if _, ok := wifi0["ietf-ip:ipv4"]; !ok { + t.Fatalf("address of ifindex 16 lost to the stale ifindex 14 row: %v", wifi0) + } +} + +// A row ip printed without name or index must not become an interface +// with an empty key and if-index 0. +func TestTransformSkipsNamelessRow(t *testing.T) { + links := mustRaw(t, []map[string]any{ + {}, + {"ifname": "e1", "ifindex": 2, "link_type": "ether", "flags": []string{"UP"}, "operstate": "UP"}, + }) + ifaces := mustInterfaces(t, Transform(links, nil, nil, &mockFileChecker{})) + if len(ifaces) != 1 || ifaces[0]["name"] != "e1" { + t.Fatalf("interfaces = %v, want only e1", ifaces) + } +} diff --git a/src/yangerd/internal/inotify/inotify.go b/src/yangerd/internal/inotify/inotify.go new file mode 100644 index 000000000..901bd2856 --- /dev/null +++ b/src/yangerd/internal/inotify/inotify.go @@ -0,0 +1,283 @@ +// Package inotify is a thin Linux inotify watcher with the fsnotify +// shape the monitors use: Add, Remove, Events and Errors. Events carry +// the full path, a directory watch reports its entries as dir/name. +package inotify + +import ( + "errors" + "fmt" + "os" + "path/filepath" + "strings" + "sync" + "unsafe" + + "golang.org/x/sys/unix" +) + +// Op is a set of event types. +type Op uint32 + +const ( + Create Op = 1 << iota + Write + Remove + Rename + Chmod +) + +// Has reports whether op contains every bit of want. +func (op Op) Has(want Op) bool { return op&want == want } + +func (op Op) String() string { + var parts []string + for _, f := range []struct { + op Op + name string + }{{Create, "CREATE"}, {Write, "WRITE"}, {Remove, "REMOVE"}, {Rename, "RENAME"}, {Chmod, "CHMOD"}} { + if op.Has(f.op) { + parts = append(parts, f.name) + } + } + if len(parts) == 0 { + return "NONE" + } + return strings.Join(parts, "|") +} + +// Event is one change on a watched path. +type Event struct { + Name string + Op Op +} + +// Has reports whether the event has every bit of op. +func (e Event) Has(op Op) bool { return e.Op.Has(op) } + +// ErrEventOverflow is sent on Errors when the kernel dropped events. +var ErrEventOverflow = errors.New("inotify queue overflow") + +// ErrNonExistentWatch is returned by Remove for a path not watched. +var ErrNonExistentWatch = errors.New("can't remove non-existent watch") + +const mask = unix.IN_CREATE | unix.IN_MODIFY | unix.IN_DELETE | unix.IN_DELETE_SELF | + unix.IN_MOVED_FROM | unix.IN_MOVED_TO | unix.IN_MOVE_SELF | unix.IN_ATTRIB + +// Watcher delivers inotify events on Events until Close. +type Watcher struct { + Events chan Event + Errors chan error + + fd int + closeR int + closeW int + mu sync.Mutex + paths map[int]string // wd -> path + watches map[string]int // path -> wd + closing chan struct{} // closed by Close, before readLoop is told + done chan struct{} // closed when readLoop has exited + once sync.Once +} + +// NewWatcher creates an inotify instance and starts reading it. +func NewWatcher() (*Watcher, error) { + fd, err := unix.InotifyInit1(unix.IN_CLOEXEC | unix.IN_NONBLOCK) + if err != nil { + return nil, fmt.Errorf("inotify_init: %w", err) + } + + var p [2]int + if err := unix.Pipe2(p[:], unix.O_CLOEXEC|unix.O_NONBLOCK); err != nil { + unix.Close(fd) + return nil, fmt.Errorf("pipe: %w", err) + } + + w := &Watcher{ + Events: make(chan Event), + Errors: make(chan error), + fd: fd, + closeR: p[0], + closeW: p[1], + paths: make(map[int]string), + watches: make(map[string]int), + closing: make(chan struct{}), + done: make(chan struct{}), + } + go w.readLoop() + return w, nil +} + +// Add watches path. Adding a path already watched is a no-op. +func (w *Watcher) Add(path string) error { + path = filepath.Clean(path) + w.mu.Lock() + defer w.mu.Unlock() + + wd, err := unix.InotifyAddWatch(w.fd, path, mask) + if err != nil { + return &os.PathError{Op: "inotify_add_watch", Path: path, Err: err} + } + if old, ok := w.paths[wd]; ok && old != path { + delete(w.watches, old) + } + w.paths[wd] = path + w.watches[path] = wd + return nil +} + +// Remove stops watching path. +func (w *Watcher) Remove(path string) error { + path = filepath.Clean(path) + w.mu.Lock() + defer w.mu.Unlock() + + wd, ok := w.watches[path] + if !ok { + return fmt.Errorf("%w: %s", ErrNonExistentWatch, path) + } + delete(w.watches, path) + delete(w.paths, wd) + + // The kernel already dropped the watch of a deleted path, so + // EINVAL here is expected and not an error. + if _, err := unix.InotifyRmWatch(w.fd, uint32(wd)); err != nil && err != unix.EINVAL { + return &os.PathError{Op: "inotify_rm_watch", Path: path, Err: err} + } + return nil +} + +// Close stops the watcher and closes Events and Errors. +func (w *Watcher) Close() error { + w.once.Do(func() { + close(w.closing) + unix.Write(w.closeW, []byte{0}) + <-w.done + unix.Close(w.closeW) + unix.Close(w.closeR) + unix.Close(w.fd) + }) + return nil +} + +func (w *Watcher) readLoop() { + defer close(w.done) + defer close(w.Errors) + defer close(w.Events) + + buf := make([]byte, 64*1024) + fds := []unix.PollFd{ + {Fd: int32(w.fd), Events: unix.POLLIN}, + {Fd: int32(w.closeR), Events: unix.POLLIN}, + } + + for { + if _, err := unix.Poll(fds, -1); err != nil { + if err == unix.EINTR { + continue + } + w.sendError(fmt.Errorf("poll: %w", err)) + return + } + if fds[1].Revents != 0 { + return + } + if fds[0].Revents == 0 { + continue + } + + for { + n, err := unix.Read(w.fd, buf) + if err == unix.EAGAIN || err == unix.EINTR { + break + } + if err != nil { + w.sendError(fmt.Errorf("read: %w", err)) + return + } + if !w.dispatch(buf[:n]) { + return + } + } + } +} + +// dispatch decodes one read worth of events. It returns false when the +// watcher was closed while an event was waiting to be received. +func (w *Watcher) dispatch(buf []byte) bool { + for len(buf) >= unix.SizeofInotifyEvent { + raw := (*unix.InotifyEvent)(unsafe.Pointer(&buf[0])) + size := unix.SizeofInotifyEvent + int(raw.Len) + if size > len(buf) { + return true + } + name := strings.TrimRight(string(buf[unix.SizeofInotifyEvent:size]), "\x00") + buf = buf[size:] + + if raw.Mask&unix.IN_Q_OVERFLOW != 0 { + if !w.sendError(ErrEventOverflow) { + return false + } + continue + } + + w.mu.Lock() + path, ok := w.paths[int(raw.Wd)] + if raw.Mask&(unix.IN_IGNORED|unix.IN_DELETE_SELF|unix.IN_MOVE_SELF) != 0 && ok { + delete(w.paths, int(raw.Wd)) + delete(w.watches, path) + } + w.mu.Unlock() + if !ok || raw.Mask&unix.IN_IGNORED != 0 { + continue + } + + if name != "" { + path = filepath.Join(path, name) + } + op := opFromMask(raw.Mask) + if op == 0 { + continue + } + + select { + case w.Events <- Event{Name: path, Op: op}: + case <-w.closed(): + return false + } + } + return true +} + +func opFromMask(m uint32) Op { + var op Op + if m&(unix.IN_CREATE|unix.IN_MOVED_TO) != 0 { + op |= Create + } + if m&unix.IN_MODIFY != 0 { + op |= Write + } + if m&(unix.IN_DELETE|unix.IN_DELETE_SELF) != 0 { + op |= Remove + } + if m&(unix.IN_MOVED_FROM|unix.IN_MOVE_SELF) != 0 { + op |= Rename + } + if m&unix.IN_ATTRIB != 0 { + op |= Chmod + } + return op +} + +func (w *Watcher) sendError(err error) bool { + select { + case w.Errors <- err: + return true + case <-w.closed(): + return false + } +} + +// closed yields a channel that is readable once Close has been called. +func (w *Watcher) closed() <-chan struct{} { + return w.closing +} diff --git a/src/yangerd/internal/inotify/inotify_test.go b/src/yangerd/internal/inotify/inotify_test.go new file mode 100644 index 000000000..75d3d3820 --- /dev/null +++ b/src/yangerd/internal/inotify/inotify_test.go @@ -0,0 +1,184 @@ +package inotify + +import ( + "os" + "path/filepath" + "runtime" + "testing" + "time" +) + +func next(t *testing.T, w *Watcher) Event { + t.Helper() + select { + case ev := <-w.Events: + return ev + case err := <-w.Errors: + t.Fatalf("unexpected error: %v", err) + case <-time.After(2 * time.Second): + t.Fatal("timeout waiting for event") + } + return Event{} +} + +func collect(t *testing.T, w *Watcher, until func(Event) bool) []Event { + t.Helper() + var evs []Event + for { + ev := next(t, w) + evs = append(evs, ev) + if until(ev) { + return evs + } + } +} + +func TestDirectoryEventsCarryFullPath(t *testing.T) { + dir := t.TempDir() + w, err := NewWatcher() + if err != nil { + t.Fatal(err) + } + defer w.Close() + if err := w.Add(dir); err != nil { + t.Fatal(err) + } + + file := filepath.Join(dir, "f") + if err := os.WriteFile(file, []byte("x"), 0644); err != nil { + t.Fatal(err) + } + evs := collect(t, w, func(e Event) bool { return e.Has(Write) }) + if evs[0].Name != file || !evs[0].Has(Create) { + t.Errorf("first event = %+v, want Create on %s", evs[0], file) + } + + if err := os.Remove(file); err != nil { + t.Fatal(err) + } + if ev := next(t, w); ev.Name != file || !ev.Has(Remove) { + t.Errorf("event = %+v, want Remove on %s", ev, file) + } +} + +func TestFileWatchRemoveSelf(t *testing.T) { + dir := t.TempDir() + file := filepath.Join(dir, "f") + if err := os.WriteFile(file, nil, 0644); err != nil { + t.Fatal(err) + } + w, err := NewWatcher() + if err != nil { + t.Fatal(err) + } + defer w.Close() + if err := w.Add(file); err != nil { + t.Fatal(err) + } + if err := w.Add(file); err != nil { + t.Errorf("second Add of a live watch: %v", err) + } + + if err := os.Remove(file); err != nil { + t.Fatal(err) + } + evs := collect(t, w, func(e Event) bool { return e.Has(Remove) }) + if evs[len(evs)-1].Name != file { + t.Errorf("remove event on %s, want %s", evs[len(evs)-1].Name, file) + } + if err := w.Remove(file); err == nil { + t.Error("Remove after the kernel dropped the watch should report it unknown") + } +} + +func TestRenameIntoWatchedDirIsCreate(t *testing.T) { + dir := t.TempDir() + other := t.TempDir() + src := filepath.Join(other, "f") + if err := os.WriteFile(src, nil, 0644); err != nil { + t.Fatal(err) + } + w, err := NewWatcher() + if err != nil { + t.Fatal(err) + } + defer w.Close() + if err := w.Add(dir); err != nil { + t.Fatal(err) + } + + dst := filepath.Join(dir, "f") + if err := os.Rename(src, dst); err != nil { + t.Fatal(err) + } + if ev := next(t, w); ev.Name != dst || !ev.Has(Create) { + t.Errorf("event = %+v, want Create on %s", ev, dst) + } +} + +func TestAddMissingPathFails(t *testing.T) { + w, err := NewWatcher() + if err != nil { + t.Fatal(err) + } + defer w.Close() + if err := w.Add(filepath.Join(t.TempDir(), "missing")); err == nil { + t.Error("Add of a missing path should fail") + } + if err := w.Remove("/never/watched"); err == nil { + t.Error("Remove of an unwatched path should fail") + } +} + +func TestCloseEndsChannels(t *testing.T) { + w, err := NewWatcher() + if err != nil { + t.Fatal(err) + } + if err := w.Add(t.TempDir()); err != nil { + t.Fatal(err) + } + w.Close() + w.Close() + if _, ok := <-w.Events; ok { + t.Error("Events still open after Close") + } + if _, ok := <-w.Errors; ok { + t.Error("Errors still open after Close") + } +} + +// Delivering events must not leave goroutines behind: each one used to +// spawn a poller that lived until Close, pinning an OS thread. +func TestEventsLeaveNoGoroutines(t *testing.T) { + dir := t.TempDir() + w, err := NewWatcher() + if err != nil { + t.Fatal(err) + } + defer w.Close() + if err := w.Add(dir); err != nil { + t.Fatal(err) + } + + before := runtime.NumGoroutine() + for i := 0; i < 200; i++ { + name := filepath.Join(dir, "f") + if err := os.WriteFile(name, []byte("x"), 0644); err != nil { + t.Fatal(err) + } + os.Remove(name) + for drained := false; !drained; { + select { + case <-w.Events: + case <-time.After(20 * time.Millisecond): + drained = true + } + } + } + runtime.GC() + time.Sleep(50 * time.Millisecond) + if after := runtime.NumGoroutine(); after > before+2 { + t.Fatalf("goroutines grew from %d to %d over 200 events", before, after) + } +} diff --git a/src/yangerd/internal/ipbatch/ipbatch.go b/src/yangerd/internal/ipbatch/ipbatch.go new file mode 100644 index 000000000..fa676605f --- /dev/null +++ b/src/yangerd/internal/ipbatch/ipbatch.go @@ -0,0 +1,357 @@ +// Package ipbatch manages a persistent `ip -json ... -force -batch -` or +// `bridge -json -force -batch -` subprocess. Commands sent via Query are +// serialized by a mutex and each is paired with the JSON line the +// subprocess writes for it. +// +// Pairing cannot rely on one line per command: depending on the iproute2 +// version a failing command writes nothing, "[]" or "[{}]" before its +// "Command failed -:" (N is the line number in the batch), and with +// stdout and stderr on separate pipes the two race. So the subprocess +// gets both on one pipe, which keeps its write order, and every command +// is followed by a sentinel that always fails. A query is answered by +// the lines that arrive before the sentinel's failure: a failure for the +// command's own line number means ErrCommandFailed, otherwise the last +// JSON line is the answer. +// +// When -s is present, `link show` commands produce multiple lines of +// output, breaking the one-command-one-line protocol, so address queries +// must use a separate instance without WithStats. +package ipbatch + +import ( + "bufio" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "log/slog" + "os" + "os/exec" + "regexp" + "strconv" + "sync" + "sync/atomic" + "time" + + "github.com/kernelkit/infix/src/yangerd/internal/backoff" +) + +// ErrBatchDead is returned by Query when the subprocess is not running. +// Callers should treat it as transient and retry on the next event. +var ErrBatchDead = errors.New("batch process is dead") + +// ErrCommandFailed is returned by Query when the subprocess rejected the +// command, typically because the device it names does not exist. +var ErrCommandFailed = errors.New("batch command failed") + +const ( + queryTimeout = 5 * time.Second + + // sentinel is an object neither ip nor bridge knows, so it fails + // with only a "Command failed" line and never any JSON. + sentinel = "sentinel" +) + +var failedRe = regexp.MustCompile(`^Command failed -:(\d+)$`) + +// Option configures an ip batch instance. +type Option func(*[]string) + +// WithStats adds -s (statistics) to the ip command. +func WithStats() Option { return func(a *[]string) { *a = append(*a, "-s") } } + +// WithDetails adds -d (details) to the ip command. +func WithDetails() Option { return func(a *[]string) { *a = append(*a, "-d") } } + +// Batch wraps one persistent batch subprocess. +type Batch struct { + argv []string + canary string + log *slog.Logger + ctx context.Context + cancel context.CancelFunc + + mu sync.Mutex // serializes queries, guards the fields below + cmd *exec.Cmd + stdin io.WriteCloser + lines chan []byte + quit chan struct{} // closed when this subprocess is replaced + seq int // lines written to the current subprocess + + // restartMu serializes Refresh and the restart loop, so a process + // one of them just started is never killed by the other. + restartMu sync.Mutex + + alive atomic.Bool + gen atomic.Int64 // bumped per subprocess, so a stale reader cannot kill a new one + died chan struct{} // kicked when the subprocess goes away +} + +// New starts `ip -json [opts] -force -batch -`. +func New(ctx context.Context, log *slog.Logger, opts ...Option) (*Batch, error) { + args := []string{"-json"} + for _, o := range opts { + o(&args) + } + return Start(ctx, log, append([]string{"ip"}, append(args, "-force", "-batch", "-")...), "link show lo") +} + +// NewBridge starts `bridge -json -force -batch -`. +func NewBridge(ctx context.Context, log *slog.Logger) (*Batch, error) { + return Start(ctx, log, []string{"bridge", "-json", "-force", "-batch", "-"}, "vlan show dev lo") +} + +// Start runs argv as a batch subprocess and keeps it running, restarting +// it with backoff when it dies. canary is a command that must succeed, +// used to validate a restarted subprocess. +func Start(ctx context.Context, log *slog.Logger, argv []string, canary string) (*Batch, error) { + ctx, cancel := context.WithCancel(ctx) + b := &Batch{ + argv: argv, + canary: canary, + log: log, + ctx: ctx, + cancel: cancel, + died: make(chan struct{}, 1), + } + if err := b.start(); err != nil { + cancel() + return nil, err + } + go b.restartLoop() + return b, nil +} + +// start spawns a new subprocess and makes it current. The previous +// one, if any, is left to the caller to terminate. +func (b *Batch) start() error { + cmd := exec.CommandContext(b.ctx, b.argv[0], b.argv[1:]...) + stdin, err := cmd.StdinPipe() + if err != nil { + return fmt.Errorf("stdin pipe: %w", err) + } + r, w, err := os.Pipe() + if err != nil { + return fmt.Errorf("output pipe: %w", err) + } + cmd.Stdout = w + cmd.Stderr = w + err = cmd.Start() + w.Close() + if err != nil { + r.Close() + return fmt.Errorf("start %s batch: %w", b.argv[0], err) + } + + lines := make(chan []byte, 8) + quit := make(chan struct{}) + gen := b.gen.Add(1) + b.mu.Lock() + if b.quit != nil { + close(b.quit) + } + b.cmd = cmd + b.stdin = stdin + b.lines = lines + b.quit = quit + b.seq = 0 + b.alive.Store(true) + b.mu.Unlock() + + go b.readLines(r, lines, quit, gen) + return nil +} + +// readLines feeds the merged stdout and stderr of one subprocess to +// lines until it exits or is replaced. +func (b *Batch) readLines(r io.ReadCloser, lines chan<- []byte, quit <-chan struct{}, gen int64) { + defer r.Close() + scanner := bufio.NewScanner(r) + scanner.Buffer(make([]byte, 0, 64*1024), 4*1024*1024) + for scanner.Scan() { + select { + case lines <- append([]byte(nil), scanner.Bytes()...): + case <-quit: + return + } + } + close(lines) + if b.gen.Load() == gen { + b.markDead() + } +} + +func (b *Batch) markDead() { + b.alive.Store(false) + select { + case b.died <- struct{}{}: + default: + } +} + +// Query sends a command to the batch process and returns its JSON +// response, ErrCommandFailed if the subprocess rejected it, or +// ErrBatchDead if the subprocess is gone. +func (b *Batch) Query(command string) (json.RawMessage, error) { + if !b.alive.Load() { + return nil, ErrBatchDead + } + b.mu.Lock() + defer b.mu.Unlock() + if !b.alive.Load() { + return nil, ErrBatchDead + } + + if _, err := fmt.Fprintf(b.stdin, "%s\n%s\n", command, sentinel); err != nil { + b.markDead() + return nil, fmt.Errorf("write command: %w", err) + } + b.seq += 2 + cmdNo, endNo := b.seq-1, b.seq + + var answer json.RawMessage + failed := false + timeout := time.NewTimer(queryTimeout) + defer timeout.Stop() + for { + select { + case line, ok := <-b.lines: + if !ok { + return nil, ErrBatchDead + } + if m := failedRe.FindSubmatch(line); m != nil { + n, _ := strconv.Atoi(string(m[1])) + switch n { + case endNo: + if failed { + return nil, fmt.Errorf("%w: %s", ErrCommandFailed, command) + } + return answer, nil + case cmdNo: + failed = true + } + continue + } + if len(line) > 0 && (line[0] == '[' || line[0] == '{') { + answer = json.RawMessage(line) + } + // Anything else is an error message, e.g. "Device "x" does + // not exist.", and the failure line that follows says so. + case <-timeout.C: + b.log.Warn(b.argv[0]+" batch query timeout, killing subprocess", "cmd", command) + b.markDead() + if b.cmd.Process != nil { + b.cmd.Process.Kill() + } + return nil, fmt.Errorf("timeout waiting for response to: %s", command) + } + } +} + +// Refresh replaces the subprocess with a fresh one and validates it with +// the canary. iproute2 caches name-to-index lookups for the life of the +// process, so once a name is reused by a new interface only a new +// process resolves it right. +func (b *Batch) Refresh() error { + b.restartMu.Lock() + defer b.restartMu.Unlock() + + b.mu.Lock() + cmd, stdin := b.cmd, b.stdin + b.mu.Unlock() + + if err := b.start(); err != nil { + b.markDead() + return err + } + reap(cmd, stdin) + + _, err := b.Query(b.canary) + return err +} + +// reap terminates a subprocess that is no longer current. +func reap(cmd *exec.Cmd, stdin io.Closer) { + if stdin != nil { + stdin.Close() + } + if cmd != nil && cmd.Process != nil { + cmd.Process.Kill() + go cmd.Wait() + } +} + +// Alive tells whether the subprocess is running and answering. +func (b *Batch) Alive() bool { + return b.alive.Load() +} + +// Close terminates the subprocess and cancels the restart loop. +func (b *Batch) Close() { + b.cancel() + b.mu.Lock() + defer b.mu.Unlock() + if b.stdin != nil { + b.stdin.Close() + } + if b.cmd != nil && b.cmd.Process != nil { + b.cmd.Process.Kill() + } + b.alive.Store(false) +} + +// restartLoop respawns the subprocess when it dies, with exponential +// backoff, validating each new process with the canary command. +func (b *Batch) restartLoop() { + bo := backoff.Default() + delay := bo.Initial + for { + select { + case <-b.ctx.Done(): + return + case <-b.died: + } + + for !b.alive.Load() { + b.log.Info(b.argv[0]+" batch: subprocess died, restarting", "delay", delay) + if backoff.Sleep(b.ctx, delay) != nil { + return + } + delay = bo.Next(delay) + + if b.restart() { + b.log.Info(b.argv[0] + " batch: restarted") + delay = bo.Initial + } + } + } +} + +// restart replaces the dead subprocess, unless a Refresh got there +// first. It reports whether a live subprocess is in place. +func (b *Batch) restart() bool { + b.restartMu.Lock() + defer b.restartMu.Unlock() + + if b.alive.Load() { + return true + } + + b.mu.Lock() + cmd, stdin := b.cmd, b.stdin + b.mu.Unlock() + + if err := b.start(); err != nil { + b.log.Warn(b.argv[0]+" batch: restart failed", "err", err) + return false + } + reap(cmd, stdin) + + if _, err := b.Query(b.canary); err != nil { + b.log.Warn(b.argv[0]+" batch: canary query failed", "err", err) + b.markDead() + return false + } + return true +} diff --git a/src/yangerd/internal/ipbatch/ipbatch_test.go b/src/yangerd/internal/ipbatch/ipbatch_test.go new file mode 100644 index 000000000..6e6a50c16 --- /dev/null +++ b/src/yangerd/internal/ipbatch/ipbatch_test.go @@ -0,0 +1,173 @@ +package ipbatch + +import ( + "context" + "errors" + "log/slog" + "os" + "path/filepath" + "testing" + "time" +) + +// fakeBatch behaves like `ip -force -batch -`: one JSON line per good +// command, "Command failed -:N" on stderr for a bad one (after a stray +// "[]" for "junk"), a stall for "hang", and exit for "die". +const fakeBatch = `#!/bin/sh +n=0 +while read -r cmd; do + n=$((n+1)) + case "$cmd" in + sentinel) echo "Object \"sentinel\" is unknown" >&2; echo "Command failed -:$n" >&2 ;; + fail*) echo "Cannot find device" >&2; echo "Command failed -:$n" >&2 ;; + junk*) echo "[]"; echo "Cannot find device" >&2; echo "Command failed -:$n" >&2 ;; + hang*) sleep 30 ;; + die*) exit 1 ;; + *) echo "[\"$cmd\"]" ;; + esac +done +` + +func startFake(t *testing.T) *Batch { + t.Helper() + script := filepath.Join(t.TempDir(), "fake") + if err := os.WriteFile(script, []byte(fakeBatch), 0755); err != nil { + t.Fatal(err) + } + b, err := Start(context.Background(), slog.Default(), []string{script}, "canary") + if err != nil { + t.Fatal(err) + } + t.Cleanup(b.Close) + return b +} + +func TestQueryAnswers(t *testing.T) { + b := startFake(t) + got, err := b.Query("link show dev eth0") + if err != nil || string(got) != `["link show dev eth0"]` { + t.Fatalf("Query = %s, %v", got, err) + } +} + +// A rejected command is reported at once, and the stream stays in step. +func TestQueryFailedCommandIsImmediate(t *testing.T) { + b := startFake(t) + start := time.Now() + if _, err := b.Query("fail dev gone0"); !errors.Is(err, ErrCommandFailed) { + t.Fatalf("err = %v, want ErrCommandFailed", err) + } + if time.Since(start) > time.Second { + t.Fatal("failed command waited for the timeout") + } + got, err := b.Query("next") + if err != nil || string(got) != `["next"]` { + t.Fatalf("follow-up Query = %s, %v", got, err) + } +} + +// A dead subprocess fails the next query fast and is restarted. +func TestQueryDeadThenRestart(t *testing.T) { + b := startFake(t) + start := time.Now() + if _, err := b.Query("die"); !errors.Is(err, ErrBatchDead) { + t.Fatalf("err = %v, want ErrBatchDead", err) + } + if time.Since(start) > time.Second { + t.Fatal("dead subprocess waited for the timeout") + } + + deadline := time.Now().Add(5 * time.Second) + for time.Now().Before(deadline) { + if got, err := b.Query("again"); err == nil { + if string(got) != `["again"]` { + t.Fatalf("after restart got %s", got) + } + return + } + time.Sleep(50 * time.Millisecond) + } + t.Fatal("subprocess was not restarted") +} + +// A command that gets no answer at all times out and kills the process. +func TestQueryTimeout(t *testing.T) { + if testing.Short() { + t.Skip("waits for the query timeout") + } + b := startFake(t) + if _, err := b.Query("hang"); err == nil || errors.Is(err, ErrCommandFailed) { + t.Fatalf("err = %v, want a timeout", err) + } + if b.alive.Load() { + t.Fatal("process still marked alive after timeout") + } +} + +// Refresh swaps in a new subprocess, which starts its line count over, +// and the old one's exit must not mark the new one dead. +func TestRefreshReplacesProcess(t *testing.T) { + b := startFake(t) + if _, err := b.Query("link show dev eth0"); err != nil { + t.Fatal(err) + } + b.mu.Lock() + oldPid := b.cmd.Process.Pid + b.mu.Unlock() + + if err := b.Refresh(); err != nil { + t.Fatalf("Refresh: %v", err) + } + + b.mu.Lock() + newPid := b.cmd.Process.Pid + b.mu.Unlock() + if newPid == oldPid { + t.Fatal("Refresh kept the old process") + } + + time.Sleep(100 * time.Millisecond) + if _, err := b.Query("fail dev gone0"); !errors.Is(err, ErrCommandFailed) { + t.Fatalf("line numbering out of step after Refresh: %v", err) + } + if got, err := b.Query("link show dev eth1"); err != nil || string(got) != `["link show dev eth1"]` { + t.Fatalf("Query after Refresh = %s, %v", got, err) + } +} + +// A failing command that also prints JSON, as some iproute2 versions do, +// is still a failure, and its stray line is not handed to the next query. +func TestQueryJunkBeforeFailure(t *testing.T) { + b := startFake(t) + if _, err := b.Query("junk dev gone0"); !errors.Is(err, ErrCommandFailed) { + t.Fatalf("err = %v, want ErrCommandFailed", err) + } + got, err := b.Query("next") + if err != nil || string(got) != `["next"]` { + t.Fatalf("follow-up Query = %s, %v", got, err) + } +} + +// Refresh and the restart loop must not kill each other's process. +func TestRefreshDuringRestart(t *testing.T) { + b := startFake(t) + if _, err := b.Query("die"); !errors.Is(err, ErrBatchDead) { + t.Fatalf("err = %v, want ErrBatchDead", err) + } + for i := 0; i < 20; i++ { + b.Refresh() + time.Sleep(10 * time.Millisecond) + } + deadline := time.Now().Add(5 * time.Second) + for time.Now().Before(deadline) { + if got, err := b.Query("again"); err == nil && string(got) == `["again"]` { + time.Sleep(300 * time.Millisecond) + if !b.alive.Load() { + t.Fatal("process killed after it was restarted") + } + return + } + time.Sleep(50 * time.Millisecond) + } + t.Fatal("no live subprocess after refreshes") +} diff --git a/src/yangerd/internal/ipc/client.go b/src/yangerd/internal/ipc/client.go new file mode 100644 index 000000000..0a9de57dc --- /dev/null +++ b/src/yangerd/internal/ipc/client.go @@ -0,0 +1,56 @@ +package ipc + +import ( + "encoding/json" + "fmt" + "net" + "time" +) + +// Client connects to a yangerd Unix socket and issues IPC requests. +type Client struct { + addr string + timeout time.Duration +} + +// NewClient returns a Client that connects to the given socket path +// with per-request timeout. +func NewClient(socketPath string, timeout time.Duration) *Client { + return &Client{ + addr: socketPath, + timeout: timeout, + } +} + +// Get queries a YANG subtree by path. Path "/" returns all models. +func (c *Client) Get(path string) (*Response, error) { + return c.call(&Request{Method: "get", Path: path}) +} + +// Health returns per-model freshness metadata. +func (c *Client) Health() (*Response, error) { + return c.call(&Request{Method: "health"}) +} + +func (c *Client) call(req *Request) (*Response, error) { + conn, err := net.DialTimeout("unix", c.addr, c.timeout) + if err != nil { + return nil, fmt.Errorf("connect %s: %w", c.addr, err) + } + defer conn.Close() + + conn.SetDeadline(time.Now().Add(c.timeout)) + + payload, err := json.Marshal(req) + if err != nil { + return nil, err + } + if err := WriteFrame(conn, payload); err != nil { + return nil, fmt.Errorf("write request: %w", err) + } + resp, err := ReadResponse(conn) + if err != nil { + return nil, fmt.Errorf("read response: %w", err) + } + return resp, nil +} diff --git a/src/yangerd/internal/ipc/conn_test.go b/src/yangerd/internal/ipc/conn_test.go new file mode 100644 index 000000000..e1bbc02df --- /dev/null +++ b/src/yangerd/internal/ipc/conn_test.go @@ -0,0 +1,83 @@ +package ipc + +import ( + "context" + "encoding/json" + "net" + "path/filepath" + "sync/atomic" + "testing" + "time" + + "github.com/kernelkit/infix/src/yangerd/internal/tree" +) + +func startServer(t *testing.T, timeout time.Duration) (string, context.CancelFunc, chan error) { + t.Helper() + sockPath := filepath.Join(t.TempDir(), "conn.sock") + ready := &atomic.Bool{} + ready.Store(true) + srv := NewServer(tree.New(), ready) + srv.timeout = timeout + if err := srv.Listen(sockPath); err != nil { + t.Fatal(err) + } + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan error, 1) + go func() { done <- srv.Serve(ctx) }() + return sockPath, cancel, done +} + +// A client that connects and never sends is dropped at the deadline. +func TestServerIdleConnDeadline(t *testing.T) { + sockPath, cancel, _ := startServer(t, 200*time.Millisecond) + defer cancel() + + conn, err := net.Dial("unix", sockPath) + if err != nil { + t.Fatal(err) + } + defer conn.Close() + + conn.SetReadDeadline(time.Now().Add(2 * time.Second)) + _, err = conn.Read(make([]byte, 1)) + if err == nil { + t.Fatal("expected the server to close an idle connection") + } + if ne, ok := err.(net.Error); ok && ne.Timeout() { + t.Fatal("server kept an idle connection open past its deadline") + } +} + +// Shutdown must not wait for a connection that is still open. +func TestServerShutdownClosesConns(t *testing.T) { + sockPath, cancel, done := startServer(t, time.Hour) + + conn, err := net.Dial("unix", sockPath) + if err != nil { + t.Fatal(err) + } + defer conn.Close() + time.Sleep(20 * time.Millisecond) + + cancel() + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("Serve did not return with a connection still open") + } +} + +// Data travels in its own frame, byte for byte. +func TestServerGetRawDataFrame(t *testing.T) { + tr := tree.New() + tr.Set("ietf-system:system", json.RawMessage(`{"hostname":"r1"}`)) + + resp := serverRoundTrip(t, tr, true, &Request{Method: "get", Path: "/ietf-system:system"}) + if !resp.Raw { + t.Fatal("expected raw data frame") + } + if string(resp.Data) != `{"ietf-system:system":{"hostname":"r1"}}` { + t.Fatalf("data = %s", resp.Data) + } +} diff --git a/src/yangerd/internal/ipc/protocol.go b/src/yangerd/internal/ipc/protocol.go new file mode 100644 index 000000000..2d502f048 --- /dev/null +++ b/src/yangerd/internal/ipc/protocol.go @@ -0,0 +1,135 @@ +// Package ipc implements the yangerd IPC protocol: a versioned, +// length-prefixed JSON framing over AF_UNIX SOCK_STREAM. +// +// Wire format: +// +// +--------+--------+--------+--------+--------+--- ... ---+ +// | ver(1) | length (uint32 big-endian, bytes) | JSON body | +// +--------+--------+--------+--------+--------+--- ... ---+ +// +// A response carrying data sets "raw" and sends the data as a second +// frame, so a client can hand it to a parser without decoding the +// envelope around it. +package ipc + +import ( + "encoding/binary" + "encoding/json" + "fmt" + "io" +) + +const ( + // Version is the current protocol version. + Version byte = 2 + + // MaxPayload is the maximum JSON body size (4 MiB). + MaxPayload = 4 << 20 + + headerSize = 5 // 1 byte version + 4 bytes length +) + +// Request is the IPC request from a client. +type Request struct { + Method string `json:"method"` + Path string `json:"path,omitempty"` +} + +// Response is the IPC response to a client. +type Response struct { + Status string `json:"status"` + Code int `json:"code,omitempty"` + Message string `json:"message,omitempty"` + + // Raw is set on the wire when Data follows in its own frame. + Raw bool `json:"raw,omitempty"` + + // Used by "get" responses, sent as the second frame. + Data json.RawMessage `json:"-"` + + // Used by "health" responses. + Models map[string]json.RawMessage `json:"models,omitempty"` +} + +// WriteFrame writes a versioned, length-prefixed frame to w. +func WriteFrame(w io.Writer, payload []byte) error { + if len(payload) > MaxPayload { + return fmt.Errorf("payload size %d exceeds maximum %d", len(payload), MaxPayload) + } + hdr := [headerSize]byte{Version} + binary.BigEndian.PutUint32(hdr[1:], uint32(len(payload))) + if _, err := w.Write(hdr[:]); err != nil { + return err + } + _, err := w.Write(payload) + return err +} + +// ReadFrame reads a versioned, length-prefixed frame from r. +func ReadFrame(r io.Reader) ([]byte, error) { + var hdr [headerSize]byte + if _, err := io.ReadFull(r, hdr[:]); err != nil { + return nil, err + } + if hdr[0] != Version { + return nil, fmt.Errorf("protocol version mismatch: got %d, want %d", hdr[0], Version) + } + length := binary.BigEndian.Uint32(hdr[1:]) + if length > MaxPayload { + return nil, fmt.Errorf("payload size %d exceeds maximum %d", length, MaxPayload) + } + buf := make([]byte, length) + if _, err := io.ReadFull(r, buf); err != nil { + return nil, err + } + return buf, nil +} + +// WriteResponse writes a Response frame, followed by a data frame when +// the response carries data. +func WriteResponse(w io.Writer, resp *Response) error { + hdr := *resp + hdr.Raw = resp.Data != nil + data, err := json.Marshal(&hdr) + if err != nil { + return err + } + if err := WriteFrame(w, data); err != nil { + return err + } + if !hdr.Raw { + return nil + } + return WriteFrame(w, resp.Data) +} + +// ReadRequest reads and unmarshals a framed Request. +func ReadRequest(r io.Reader) (*Request, error) { + data, err := ReadFrame(r) + if err != nil { + return nil, err + } + var req Request + if err := json.Unmarshal(data, &req); err != nil { + return nil, fmt.Errorf("invalid request JSON: %w", err) + } + return &req, nil +} + +// ReadResponse reads and unmarshals a framed Response. +func ReadResponse(r io.Reader) (*Response, error) { + data, err := ReadFrame(r) + if err != nil { + return nil, err + } + var resp Response + if err := json.Unmarshal(data, &resp); err != nil { + return nil, fmt.Errorf("invalid response JSON: %w", err) + } + if resp.Raw { + if resp.Data, err = ReadFrame(r); err != nil { + return nil, fmt.Errorf("read data frame: %w", err) + } + } + return &resp, nil +} diff --git a/src/yangerd/internal/ipc/protocol_test.go b/src/yangerd/internal/ipc/protocol_test.go new file mode 100644 index 000000000..5ce47b7c1 --- /dev/null +++ b/src/yangerd/internal/ipc/protocol_test.go @@ -0,0 +1,89 @@ +package ipc + +import ( + "bytes" + "encoding/json" + "testing" +) + +func TestFrameRoundTrip(t *testing.T) { + payload := []byte(`{"method":"get","path":"/test"}`) + var buf bytes.Buffer + + if err := WriteFrame(&buf, payload); err != nil { + t.Fatal(err) + } + + got, err := ReadFrame(&buf) + if err != nil { + t.Fatal(err) + } + if !bytes.Equal(got, payload) { + t.Fatalf("mismatch: %s vs %s", got, payload) + } +} + +func TestFrameVersionMismatch(t *testing.T) { + var buf bytes.Buffer + buf.Write([]byte{99, 0, 0, 0, 2, '{', '}'}) + + _, err := ReadFrame(&buf) + if err == nil { + t.Fatal("expected version mismatch error") + } +} + +func TestFrameOversized(t *testing.T) { + var buf bytes.Buffer + huge := make([]byte, MaxPayload+1) + if err := WriteFrame(&buf, huge); err == nil { + t.Fatal("expected oversized payload error") + } +} + +func TestRequestResponseRoundTrip(t *testing.T) { + var buf bytes.Buffer + + req := &Request{Method: "get", Path: "/ietf-system:system-state"} + data, _ := json.Marshal(req) + WriteFrame(&buf, data) + + got, err := ReadRequest(&buf) + if err != nil { + t.Fatal(err) + } + if got.Method != "get" || got.Path != "/ietf-system:system-state" { + t.Fatalf("unexpected request: %+v", got) + } +} + +func TestResponseRoundTrip(t *testing.T) { + var buf bytes.Buffer + + resp := &Response{ + Status: "ok", + Data: json.RawMessage(`{"hostname":"r1"}`), + } + WriteResponse(&buf, resp) + + got, err := ReadResponse(&buf) + if err != nil { + t.Fatal(err) + } + if got.Status != "ok" || string(got.Data) != `{"hostname":"r1"}` { + t.Fatalf("unexpected response: %+v", got) + } +} + +func TestEmptyFrame(t *testing.T) { + var buf bytes.Buffer + WriteFrame(&buf, []byte{}) + + got, err := ReadFrame(&buf) + if err != nil { + t.Fatal(err) + } + if len(got) != 0 { + t.Fatalf("expected empty, got %d bytes", len(got)) + } +} diff --git a/src/yangerd/internal/ipc/server.go b/src/yangerd/internal/ipc/server.go new file mode 100644 index 000000000..c6f595273 --- /dev/null +++ b/src/yangerd/internal/ipc/server.go @@ -0,0 +1,229 @@ +package ipc + +import ( + "context" + "encoding/json" + "log" + "net" + "os" + "sync" + "sync/atomic" + "time" + + "github.com/kernelkit/infix/src/yangerd/internal/tree" +) + +// Server listens on an AF_UNIX SOCK_STREAM socket and serves +// YANG operational data from an in-memory Tree. +type Server struct { + tree *tree.Tree + listener net.Listener + ready *atomic.Bool + wg sync.WaitGroup + + // timeout bounds a whole request/response exchange, the C + // client in statd uses the same. + timeout time.Duration + + mu sync.Mutex + conns map[net.Conn]struct{} +} + +// NewServer creates a Server that serves data from the given Tree. +// While ready is false, all requests receive a 503 "starting" response. +func NewServer(t *tree.Tree, ready *atomic.Bool) *Server { + return &Server{ + tree: t, + ready: ready, + timeout: 5 * time.Second, + conns: make(map[net.Conn]struct{}), + } +} + +// Listen creates and binds a Unix domain socket at path, root only. +// A stale socket file is removed before binding. +func (s *Server) Listen(path string) error { + if err := os.Remove(path); err != nil && !os.IsNotExist(err) { + return err + } + ln, err := net.Listen("unix", path) + if err != nil { + return err + } + if err := os.Chmod(path, 0660); err != nil { + ln.Close() + return err + } + s.listener = ln + return nil +} + +// Serve accepts connections until ctx is cancelled. Each connection +// is handled in its own goroutine. +func (s *Server) Serve(ctx context.Context) error { + go func() { + <-ctx.Done() + s.listener.Close() + s.mu.Lock() + for conn := range s.conns { + conn.Close() + } + s.mu.Unlock() + }() + + for { + conn, err := s.listener.Accept() + if err != nil { + // Listener closed by context cancellation — normal shutdown. + select { + case <-ctx.Done(): + s.wg.Wait() + return nil + default: + return err + } + } + s.mu.Lock() + s.conns[conn] = struct{}{} + s.mu.Unlock() + + s.wg.Add(1) + go func() { + defer s.wg.Done() + s.handleConn(conn) + s.mu.Lock() + delete(s.conns, conn) + s.mu.Unlock() + }() + } +} + +func (s *Server) handleConn(conn net.Conn) { + defer conn.Close() + conn.SetDeadline(time.Now().Add(s.timeout)) + + req, err := ReadRequest(conn) + if err != nil { + log.Printf("ipc: read request: %v", err) + return + } + + if !s.ready.Load() { + WriteResponse(conn, &Response{ + Status: "starting", + Code: 503, + Message: "yangerd is starting up", + }) + return + } + + switch req.Method { + case "get": + s.handleGet(conn, req) + case "health": + s.handleHealth(conn) + default: + WriteResponse(conn, &Response{ + Status: "error", + Code: 400, + Message: "unknown method: " + req.Method, + }) + } +} + +func (s *Server) handleGet(conn net.Conn, req *Request) { + path := req.Path + if path == "" || path == "/" { + s.handleDump(conn) + return + } + + key := path + if key[0] == '/' { + key = key[1:] + } + + data := s.tree.Get(key) + if data == nil { + // An absent subtree is a normal answer for operational data -- + // the feature is simply not active (e.g. NTP unconfigured). + // Answer ok with an empty object rather than an error, so every + // client gets "no data" without special-casing. Deliberately + // NOT {"": {}}: that would make libyang instantiate the + // container, which for presence containers is real data. + WriteResponse(conn, &Response{ + Status: "ok", + Data: json.RawMessage(`{}`), + }) + return + } + + envelope := map[string]json.RawMessage{key: data} + body, err := json.Marshal(envelope) + if err != nil { + WriteResponse(conn, &Response{ + Status: "error", + Code: 500, + Message: "marshal error: " + err.Error(), + }) + return + } + + WriteResponse(conn, &Response{ + Status: "ok", + Data: body, + }) +} + +func (s *Server) handleDump(conn net.Conn) { + keys := s.tree.Keys() + blobs := s.tree.GetMulti(keys) + + all := make(map[string]json.RawMessage, len(keys)) + for i, k := range keys { + if i < len(blobs) { + all[k] = blobs[i] + } + } + + body, err := json.Marshal(all) + if err != nil { + WriteResponse(conn, &Response{ + Status: "error", + Code: 500, + Message: "marshal error: " + err.Error(), + }) + return + } + + WriteResponse(conn, &Response{ + Status: "ok", + Data: body, + }) +} + +func (s *Server) handleHealth(conn net.Conn) { + keys := s.tree.Keys() + models := make(map[string]json.RawMessage, len(keys)) + + for _, k := range keys { + info, ok := s.tree.Info(k) + if !ok { + continue + } + entry := struct { + LastUpdated string `json:"last_updated"` + SizeBytes int `json:"size_bytes"` + }{ + LastUpdated: info.LastUpdated.UTC().Format(time.RFC3339), + SizeBytes: info.SizeBytes, + } + b, _ := json.Marshal(entry) + models[k] = b + } + + WriteResponse(conn, &Response{ + Status: "ok", + Models: models, + }) +} diff --git a/src/yangerd/internal/ipc/server_test.go b/src/yangerd/internal/ipc/server_test.go new file mode 100644 index 000000000..f329308b3 --- /dev/null +++ b/src/yangerd/internal/ipc/server_test.go @@ -0,0 +1,140 @@ +package ipc + +import ( + "context" + "encoding/json" + "net" + "os" + "path/filepath" + "sync/atomic" + "testing" + "time" + + "github.com/kernelkit/infix/src/yangerd/internal/tree" +) + +func TestServerGetSingle(t *testing.T) { + tr := tree.New() + tr.Set("ietf-system:system-state", json.RawMessage(`{"platform":{"os-name":"Infix"}}`)) + + resp := serverRoundTrip(t, tr, true, &Request{Method: "get", Path: "/ietf-system:system-state"}) + + if resp.Status != "ok" { + t.Fatalf("expected ok, got %s: %s", resp.Status, resp.Message) + } + var data map[string]json.RawMessage + json.Unmarshal(resp.Data, &data) + if _, ok := data["ietf-system:system-state"]; !ok { + t.Fatalf("missing key in response data: %s", resp.Data) + } +} + +func TestServerGetNotFound(t *testing.T) { + tr := tree.New() + resp := serverRoundTrip(t, tr, true, &Request{Method: "get", Path: "/nonexistent"}) + + // An absent subtree is "no data", not an error: ok + empty object, + // so clients (statd, yangerctl) need no special-casing. + if resp.Status != "ok" { + t.Fatalf("expected ok, got %+v", resp) + } + if string(resp.Data) != "{}" { + t.Fatalf("expected empty object data, got %s", resp.Data) + } +} + +func TestServerDump(t *testing.T) { + tr := tree.New() + tr.Set("a", json.RawMessage(`1`)) + tr.Set("b", json.RawMessage(`2`)) + + resp := serverRoundTrip(t, tr, true, &Request{Method: "get", Path: "/"}) + + if resp.Status != "ok" { + t.Fatalf("expected ok, got %s: %s", resp.Status, resp.Message) + } + var data map[string]json.RawMessage + json.Unmarshal(resp.Data, &data) + if len(data) != 2 { + t.Fatalf("expected 2 models in dump, got %d", len(data)) + } +} + +func TestServerHealth(t *testing.T) { + tr := tree.New() + tr.Set("model-a", json.RawMessage(`{}`)) + + resp := serverRoundTrip(t, tr, true, &Request{Method: "health"}) + + if resp.Status != "ok" { + t.Fatalf("expected ok, got %s", resp.Status) + } + if _, ok := resp.Models["model-a"]; !ok { + t.Fatalf("expected model-a in health models, got %v", resp.Models) + } +} + +func TestServerNotReady(t *testing.T) { + tr := tree.New() + resp := serverRoundTrip(t, tr, false, &Request{Method: "get", Path: "/"}) + + if resp.Status != "starting" || resp.Code != 503 { + t.Fatalf("expected 503 starting, got %+v", resp) + } +} + +func TestServerUnknownMethod(t *testing.T) { + tr := tree.New() + resp := serverRoundTrip(t, tr, true, &Request{Method: "invalid"}) + + if resp.Status != "error" || resp.Code != 400 { + t.Fatalf("expected 400 error, got %+v", resp) + } +} + +func serverRoundTrip(t *testing.T, tr *tree.Tree, ready bool, req *Request) *Response { + t.Helper() + + sockPath := filepath.Join(t.TempDir(), "test.sock") + readyFlag := &atomic.Bool{} + readyFlag.Store(ready) + + srv := NewServer(tr, readyFlag) + if err := srv.Listen(sockPath); err != nil { + t.Fatal(err) + } + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + done := make(chan error, 1) + go func() { + done <- srv.Serve(ctx) + }() + + time.Sleep(10 * time.Millisecond) + + conn, err := net.Dial("unix", sockPath) + if err != nil { + t.Fatal(err) + } + defer conn.Close() + + payload, _ := json.Marshal(req) + if err := WriteFrame(conn, payload); err != nil { + t.Fatal(err) + } + + resp, err := ReadResponse(conn) + if err != nil { + t.Fatal(err) + } + + cancel() + + if _, err := os.Stat(sockPath); err == nil { + os.Remove(sockPath) + } + + return resp +} diff --git a/src/yangerd/internal/iwmonitor/ap_test.go b/src/yangerd/internal/iwmonitor/ap_test.go new file mode 100644 index 000000000..ccdd3238d --- /dev/null +++ b/src/yangerd/internal/iwmonitor/ap_test.go @@ -0,0 +1,126 @@ +package iwmonitor + +import ( + "encoding/json" + "testing" + + "github.com/kernelkit/infix/src/yangerd/internal/wpactrl" +) + +func TestFormatStations(t *testing.T) { + m := &IWMonitor{} + stas := []map[string]string{ + { + "addr": "02:00:00:00:00:01", + "signal": "-57", + "connected_time": "120", + "rx_packets": "1500", + "tx_packets": "2500", + "rx_bytes": "4825331939", + "tx_bytes": "216392802676", + "rx_rate_info": "1560 vhtmcs 8 vhtnss 2", + "tx_rate_info": "1733 vhtmcs 9 vhtnss 2", + }, + } + + raw, err := json.Marshal(m.formatStations(stas)) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + var parsed struct { + Station []struct { + MAC string `json:"mac-address"` + Signal int16 `json:"signal-strength"` + ConnectedTime uint32 `json:"connected-time"` + RxPackets string `json:"rx-packets"` + TxPackets string `json:"tx-packets"` + RxBytes string `json:"rx-bytes"` + TxBytes string `json:"tx-bytes"` + RxSpeed uint32 `json:"rx-speed"` + TxSpeed uint32 `json:"tx-speed"` + } `json:"station"` + } + if err := json.Unmarshal(raw, &parsed); err != nil { + t.Fatalf("unmarshal: %v", err) + } + if len(parsed.Station) != 1 { + t.Fatalf("got %d stations, want 1", len(parsed.Station)) + } + + s := parsed.Station[0] + if s.MAC != "02:00:00:00:00:01" { + t.Errorf("mac-address = %q", s.MAC) + } + if s.Signal != -57 { + t.Errorf("signal-strength = %d, want -57", s.Signal) + } + if s.ConnectedTime != 120 { + t.Errorf("connected-time = %d, want 120", s.ConnectedTime) + } + if s.RxBytes != "4825331939" || s.TxBytes != "216392802676" { + t.Errorf("bytes = %q/%q", s.RxBytes, s.TxBytes) + } + if s.RxPackets != "1500" || s.TxPackets != "2500" { + t.Errorf("packets = %q/%q", s.RxPackets, s.TxPackets) + } + // hostapd rate info is already in 100kbps units + if s.RxSpeed != 1560 || s.TxSpeed != 1733 { + t.Errorf("speed = %d/%d, want 1560/1733", s.RxSpeed, s.TxSpeed) + } +} + +func TestFilterAuthorized(t *testing.T) { + stas := []map[string]string{ + {"addr": "02:00:00:00:00:01", "flags": "[AUTH][ASSOC][AUTHORIZED]"}, + {"addr": "02:00:00:00:00:02", "flags": "[AUTH][ASSOC]"}, // mid-handshake + {"addr": "02:00:00:00:00:03", "flags": "[AUTH][ASSOC][AUTHORIZED][SHORT_PREAMBLE]"}, + } + + out := filterAuthorized(stas) + if len(out) != 2 { + t.Fatalf("got %d stations, want 2", len(out)) + } + if out[0]["addr"] != "02:00:00:00:00:01" || out[1]["addr"] != "02:00:00:00:00:03" { + t.Errorf("addrs = %q, %q", out[0]["addr"], out[1]["addr"]) + } +} + +func TestResolveSSIDHostapd(t *testing.T) { + // hostapd STATUS reports bss[N]= / ssid[N]= pairs; + // multi-BSS setups must resolve by interface name. + status := map[string]string{ + "state": "ENABLED", + "bss[0]": "wlan0", + "ssid[0]": "Lobby", + "bss[1]": "wlan0_1", + "ssid[1]": "Office", + } + + si := wpactrl.SocketInfo{Iface: "wlan0_1", Daemon: "hostapd"} + if got := resolveSSID("wlan0_1", si, status); got != "Office" { + t.Errorf("resolveSSID(wlan0_1) = %q, want Office", got) + } + si = wpactrl.SocketInfo{Iface: "wlan0", Daemon: "hostapd"} + if got := resolveSSID("wlan0", si, status); got != "Lobby" { + t.Errorf("resolveSSID(wlan0) = %q, want Lobby", got) + } + if got := resolveSSID("wlan9", si, status); got != "" { + t.Errorf("resolveSSID(wlan9) = %q, want empty", got) + } +} + +func TestParseBitrate(t *testing.T) { + cases := map[string]uint32{ + "1560 vhtmcs 8 vhtnss 2": 1560, // hostapd: 100kbps units + "866.7 MBit/s VHT-MCS 9": 8667, // iw: MBit/s -> 100kbps + "54.0 MBit/s": 540, + "": 0, + "garbage rate": 0, + } + for in, want := range cases { + if got := parseBitrate(in); got != want { + t.Errorf("parseBitrate(%q) = %d, want %d", in, got, want) + } + } +} diff --git a/src/yangerd/internal/iwmonitor/data.go b/src/yangerd/internal/iwmonitor/data.go new file mode 100644 index 000000000..47709a913 --- /dev/null +++ b/src/yangerd/internal/iwmonitor/data.go @@ -0,0 +1,271 @@ +package iwmonitor + +import ( + "encoding/json" + "strconv" + "strings" + + "github.com/kernelkit/infix/src/yangerd/internal/wpactrl" +) + +func parseIWInfo(output string) json.RawMessage { + info := make(map[string]string) + for _, line := range strings.Split(output, "\n") { + if k, v, ok := parseKV(strings.TrimSpace(line)); ok { + switch k { + case "ssid": + info["ssid"] = v + case "type": + info["type"] = v + case "channel": + info["channel"] = v + case "txpower": + info["tx-power"] = v + } + } + } + data, _ := json.Marshal(info) + return json.RawMessage(data) +} + +func parseIWDevList(output string) []string { + var ifaces []string + for _, line := range strings.Split(output, "\n") { + line = strings.TrimSpace(line) + if strings.HasPrefix(line, "Interface ") { + if name := strings.TrimPrefix(line, "Interface "); name != "" { + ifaces = append(ifaces, name) + } + } + } + return ifaces +} + +func parseKV(line string) (string, string, bool) { + idx := strings.Index(line, ":") + if idx < 0 { + return "", "", false + } + k := strings.TrimSpace(line[:idx]) + v := strings.TrimSpace(line[idx+1:]) + return k, v, k != "" +} + +func parseIWLink(output string) map[string]string { + m := make(map[string]string) + for _, line := range strings.Split(output, "\n") { + if k, v, ok := parseKV(strings.TrimSpace(line)); ok { + m[k] = v + } + } + return m +} + +func parseStationDump(output string) json.RawMessage { + type station struct { + MAC string `json:"mac"` + Signal string `json:"signal,omitempty"` + RxBytes string `json:"rx-bytes,omitempty"` + TxBytes string `json:"tx-bytes,omitempty"` + Connected string `json:"connected-time,omitempty"` + Inactive string `json:"inactive-time,omitempty"` + RxBitrate string `json:"rx-bitrate,omitempty"` + TxBitrate string `json:"tx-bitrate,omitempty"` + Authorized string `json:"authorized,omitempty"` + } + var stations []station + var current *station + + for _, line := range strings.Split(output, "\n") { + line = strings.TrimSpace(line) + if strings.HasPrefix(line, "Station ") { + parts := strings.Fields(line) + if len(parts) >= 2 { + s := station{MAC: parts[1]} + stations = append(stations, s) + current = &stations[len(stations)-1] + } + continue + } + if current == nil { + continue + } + if k, v, ok := parseKV(line); ok { + switch k { + case "signal": + current.Signal = v + case "rx bytes": + current.RxBytes = v + case "tx bytes": + current.TxBytes = v + case "connected time": + current.Connected = v + case "inactive time": + current.Inactive = v + case "rx bitrate": + current.RxBitrate = v + case "tx bitrate": + current.TxBitrate = v + case "authorized": + current.Authorized = v + } + } + } + + data, _ := json.Marshal(stations) + return json.RawMessage(data) +} + +// parseBitrate extracts the speed in 100kbps units from iw/hostapd rate info. +// iw link: "866.7 MBit/s VHT-MCS 9 ..." +// hostapd: "1560 vhtmcs 8 vhtnss 2" (value in 100kbps) +func parseBitrate(s string) uint32 { + s = strings.TrimSpace(s) + if s == "" { + return 0 + } + fields := strings.Fields(s) + if len(fields) == 0 { + return 0 + } + if strings.Contains(s, "MBit/s") { + val, err := strconv.ParseFloat(fields[0], 64) + if err != nil { + return 0 + } + return uint32(val * 10) + } + val, err := strconv.ParseUint(fields[0], 10, 32) + if err != nil { + return 0 + } + return uint32(val) +} + +func extractEncryption(flags string) []string { + flags = strings.ToUpper(flags) + var result []string + if strings.Contains(flags, "WPA3") || strings.Contains(flags, "SAE") { + result = append(result, "WPA3-Personal") + } + if strings.Contains(flags, "WPA2") { + if strings.Contains(flags, "EAP") { + result = append(result, "WPA2-Enterprise") + } else { + result = append(result, "WPA2-Personal") + } + } + if strings.Contains(flags, "WEP") { + return []string{"WEP"} + } + if len(result) == 0 && strings.Contains(flags, "ESS") { + return []string{"Open"} + } + if len(result) == 0 { + return []string{"Unknown"} + } + return result +} + +func formatScanResults(results []wpactrl.ScanResult) []map[string]any { + seen := make(map[string]int) + var out []map[string]any + + for _, r := range results { + if r.SSID == "" { + continue + } + entry := map[string]any{ + "ssid": r.SSID, + "bssid": r.BSSID, + "signal-strength": r.Signal, + "channel": wpactrl.FrequencyToChannel(r.Frequency), + } + if enc := extractEncryption(r.Flags); len(enc) > 0 { + entry["encryption"] = enc + } + + if idx, dup := seen[r.SSID]; dup { + prev := out[idx]["signal-strength"].(int) + if r.Signal > prev { + out[idx] = entry + } + continue + } + seen[r.SSID] = len(out) + out = append(out, entry) + } + return out +} + +// ParseIWEvent parses a single line from `iw event -t` output. +// Retained for tests; no longer used in the main event loop. +func ParseIWEvent(line string) (IWEvent, bool) { + parts := strings.SplitN(line, ": ", 3) + if len(parts) < 3 { + return IWEvent{}, false + } + + ts, err := strconv.ParseFloat(parts[0], 64) + if err != nil { + return IWEvent{}, false + } + + ifacePhy := parts[1] + parenIdx := strings.Index(ifacePhy, " (") + if parenIdx < 0 { + return IWEvent{}, false + } + iface := ifacePhy[:parenIdx] + phy := strings.Trim(ifacePhy[parenIdx+2:], ")") + + eventStr := parts[2] + ev := IWEvent{Timestamp: ts, Interface: iface, Phy: phy} + + switch { + case strings.HasPrefix(eventStr, "new station "): + ev.Type = "new station" + ev.Addr = strings.TrimPrefix(eventStr, "new station ") + case strings.HasPrefix(eventStr, "del station "): + ev.Type = "del station" + ev.Addr = strings.TrimPrefix(eventStr, "del station ") + case strings.HasPrefix(eventStr, "connected to "): + ev.Type = "connected" + ev.Addr = strings.TrimPrefix(eventStr, "connected to ") + case eventStr == "disconnected": + ev.Type = "disconnected" + case strings.HasPrefix(eventStr, "ch_switch_started_notify"): + ev.Type = "ch_switch_started_notify" + case eventStr == "scan started": + ev.Type = "scan started" + case eventStr == "scan aborted": + ev.Type = "scan aborted" + case strings.HasPrefix(eventStr, "reg_change"): + ev.Type = "reg_change" + case strings.HasPrefix(eventStr, "auth"): + ev.Type = "auth" + default: + ev.Type = eventStr + } + + return ev, true +} + +// resolveSSID extracts the SSID for an interface. +// wpa_supplicant STATUS has "ssid=". +// hostapd STATUS has "bss[N]=" / "ssid[N]=" pairs. +func resolveSSID(iface string, si wpactrl.SocketInfo, status map[string]string) string { + if v := status["ssid"]; v != "" { + return v + } + for i := 0; i < 16; i++ { + idx := strconv.Itoa(i) + if status["bss["+idx+"]"] == iface { + if v := status["ssid["+idx+"]"]; v != "" { + return v + } + break + } + } + return "" +} diff --git a/src/yangerd/internal/iwmonitor/iwmonitor.go b/src/yangerd/internal/iwmonitor/iwmonitor.go new file mode 100644 index 000000000..3b830dd3d --- /dev/null +++ b/src/yangerd/internal/iwmonitor/iwmonitor.go @@ -0,0 +1,530 @@ +package iwmonitor + +import ( + "context" + "encoding/json" + "fmt" + "log/slog" + "math" + "net" + "strconv" + "strings" + "sync" + "time" + + "github.com/kernelkit/infix/src/yangerd/internal/wpactrl" + "github.com/mdlayher/genetlink" + "github.com/mdlayher/netlink" + "golang.org/x/sys/unix" +) + +const ( + reconnectInitial = 500 * time.Millisecond + reconnectMax = 30 * time.Second + reconnectFactor = 2.0 + queryTimeout = 5 * time.Second +) + +// IWEvent is retained for ParseIWEvent compatibility (used in tests). +type IWEvent struct { + Timestamp float64 + Interface string + Phy string + Type string + Addr string +} + +type IWMonitor struct { + log *slog.Logger + onUpdate func(ifname string, data json.RawMessage) + onRadioChange func() + + mu sync.Mutex + attached map[string]context.CancelFunc + + // meshState reads iftype, mesh forwarding and peers from the + // kernel; overridable in tests. + meshState func(ifname string) meshState +} + +func New(log *slog.Logger) *IWMonitor { + return &IWMonitor{ + log: log, + attached: make(map[string]context.CancelFunc), + meshState: kernelMeshState, + } +} + +func (m *IWMonitor) SetOnUpdate(fn func(string, json.RawMessage)) { + m.onUpdate = fn +} + +// SetOnRadioChange sets a callback for changes to what a radio can do: +// a phy came or went, the regulatory domain or its interfaces changed. +func (m *IWMonitor) SetOnRadioChange(fn func()) { + m.onRadioChange = fn +} + +// radioChanged tells the owner of the radio capabilities to rebuild them. +func (m *IWMonitor) radioChanged() { + if m.onRadioChange != nil { + m.onRadioChange() + } +} + +func (m *IWMonitor) Run(ctx context.Context) error { + conn, family, err := m.dialNL80211() + if err != nil { + return fmt.Errorf("nl80211 setup: %w", err) + } + defer conn.Close() + + m.refreshAllInterfaces(ctx) + + go func() { + <-ctx.Done() + conn.Close() + }() + + for { + msgs, _, err := conn.Receive() + if err != nil { + if ctx.Err() != nil { + return ctx.Err() + } + return fmt.Errorf("nl80211 receive: %w", err) + } + + for _, msg := range msgs { + m.handleNL80211(ctx, msg, family) + } + } +} + +func (m *IWMonitor) dialNL80211() (*genetlink.Conn, genetlink.Family, error) { + conn, err := genetlink.Dial(nil) + if err != nil { + return nil, genetlink.Family{}, fmt.Errorf("dial genetlink: %w", err) + } + + family, err := conn.GetFamily(unix.NL80211_GENL_NAME) + if err != nil { + conn.Close() + return nil, genetlink.Family{}, fmt.Errorf("resolve nl80211: %w", err) + } + + groups := map[string]bool{ + unix.NL80211_MULTICAST_GROUP_MLME: true, + unix.NL80211_MULTICAST_GROUP_REG: true, + unix.NL80211_MULTICAST_GROUP_CONFIG: true, + } + for _, g := range family.Groups { + if groups[g.Name] { + if err := conn.JoinGroup(g.ID); err != nil { + conn.Close() + return nil, genetlink.Family{}, fmt.Errorf("join %s: %w", g.Name, err) + } + m.log.Info("nl80211: joined multicast group", "name", g.Name, "id", g.ID) + } + } + + return conn, family, nil +} + +func (m *IWMonitor) handleNL80211(ctx context.Context, msg genetlink.Message, family genetlink.Family) { + ifname := m.extractIfname(msg.Data) + cmd := msg.Header.Command + + switch cmd { + case unix.NL80211_CMD_NEW_STATION, unix.NL80211_CMD_DEL_STATION: + if ifname != "" { + m.log.Debug("nl80211: station event", "cmd", cmd, "iface", ifname) + m.refreshInterface(ctx, ifname) + } + case unix.NL80211_CMD_CONNECT: + if ifname != "" { + m.log.Debug("nl80211: connect", "iface", ifname) + m.refreshInterface(ctx, ifname) + } + case unix.NL80211_CMD_DISCONNECT: + if ifname != "" { + m.log.Debug("nl80211: disconnect", "iface", ifname) + m.publishWifi(ifname, nil) + } + case unix.NL80211_CMD_REG_CHANGE: + m.log.Debug("nl80211: reg_change") + m.refreshAllInterfaces(ctx) + m.radioChanged() + case unix.NL80211_CMD_NEW_INTERFACE: + if ifname != "" { + m.log.Info("nl80211: new interface", "iface", ifname) + // A new interface gets a new daemon, drop any loop still + // waiting on the previous one's socket. + m.stopAttach(ifname) + m.startAttach(ctx, ifname) + m.refreshInterface(ctx, ifname) + } + m.radioChanged() + case unix.NL80211_CMD_DEL_INTERFACE: + if ifname != "" { + m.log.Info("nl80211: del interface", "iface", ifname) + m.stopAttach(ifname) + m.publishWifi(ifname, nil) + } + m.radioChanged() + case unix.NL80211_CMD_NEW_WIPHY, unix.NL80211_CMD_DEL_WIPHY: + m.log.Info("nl80211: phy change", "cmd", cmd) + m.radioChanged() + } +} + +// extractIfname takes the name from the message itself when it carries +// one: interface events do, and on DEL_INTERFACE the index no longer +// resolves, the interface is already gone. +func (m *IWMonitor) extractIfname(data []byte) string { + ad, err := netlink.NewAttributeDecoder(data) + if err != nil { + return "" + } + index := -1 + for ad.Next() { + switch ad.Type() { + case unix.NL80211_ATTR_IFNAME: + if name := strings.TrimRight(ad.String(), "\x00"); name != "" { + return name + } + case unix.NL80211_ATTR_IFINDEX: + index = int(ad.Uint32()) + } + } + if index < 0 { + return "" + } + iface, err := net.InterfaceByIndex(index) + if err != nil { + return "" + } + return iface.Name +} + +func (m *IWMonitor) startAttach(ctx context.Context, ifname string) { + m.mu.Lock() + if _, exists := m.attached[ifname]; exists { + m.mu.Unlock() + return + } + attachCtx, cancel := context.WithCancel(ctx) + m.attached[ifname] = cancel + m.mu.Unlock() + + go m.attachLoop(attachCtx, ifname) +} + +func (m *IWMonitor) stopAttach(ifname string) { + m.mu.Lock() + if cancel, ok := m.attached[ifname]; ok { + cancel() + delete(m.attached, ifname) + } + m.mu.Unlock() +} + +func (m *IWMonitor) attachLoop(ctx context.Context, ifname string) { + delay := reconnectInitial + + for { + if ctx.Err() != nil { + return + } + + socks := wpactrl.ScanSockets() + si, ok := socks[ifname] + if !ok { + select { + case <-ctx.Done(): + return + case <-time.After(delay): + } + delay = nextDelay(delay) + continue + } + + ac, err := wpactrl.Attach(si.Path) + if err != nil { + m.log.Debug("attach failed", "iface", ifname, "err", err) + select { + case <-ctx.Done(): + return + case <-time.After(delay): + } + delay = nextDelay(delay) + continue + } + + delay = reconnectInitial + m.log.Info("attached to control socket", "iface", ifname, "daemon", si.Daemon) + m.refreshInterface(ctx, ifname) + + ac.SetHandler(func(ev wpactrl.Event) { + m.handleAttachEvent(ctx, ifname, ev) + }) + + err = ac.Run(ctx) + ac.Close() + + if ctx.Err() != nil { + return + } + + m.log.Warn("control socket lost", "iface", ifname, "err", err) + m.publishWifi(ifname, nil) + } +} + +func (m *IWMonitor) handleAttachEvent(ctx context.Context, ifname string, ev wpactrl.Event) { + switch { + case ev.Name == "AP-STA-CONNECTED" || ev.Name == "AP-STA-DISCONNECTED": + m.refreshInterface(ctx, ifname) + case ev.Name == "CTRL-EVENT-CONNECTED": + m.refreshInterface(ctx, ifname) + case meshEvent(ev): + m.refreshInterface(ctx, ifname) + case ev.Name == "CTRL-EVENT-DISCONNECTED": + // A mesh point stays one after leaving the mesh, it only loses + // its mesh id and peers. + if m.meshState(ifname).iftype == "mesh_point" { + m.refreshInterface(ctx, ifname) + } else { + m.publishWifi(ifname, nil) + } + case ev.Name == "CTRL-EVENT-SCAN-RESULTS": + m.refreshInterface(ctx, ifname) + case ev.Name == "CTRL-EVENT-SIGNAL-CHANGE": + m.handleSignalChange(ifname, ev.Data) + case ev.Name == "CTRL-EVENT-TERMINATING": + // Daemon is shutting down; the read loop will get an error next + } +} + +func (m *IWMonitor) handleSignalChange(ifname string, data string) { + // TODO(lazzer): update signal in-place without full rebuild + // For now, this is a no-op; signal is read during refreshInterface. + _ = ifname + _ = data +} + +func (m *IWMonitor) refreshInterface(ctx context.Context, iface string) { + wifi := m.buildWifiData(ctx, iface) + m.publishWifi(iface, wifi) +} + +func (m *IWMonitor) buildWifiData(ctx context.Context, iface string) map[string]any { + socks := wpactrl.ScanSockets() + si, ok := socks[iface] + + var status map[string]string + if ok { + conn, err := wpactrl.Dial(si.Path) + if err == nil { + status, _ = conn.Status() + conn.Close() + } + } + + ms := m.kernelState(si)(iface) + result := make(map[string]any) + + switch detectMode(si, status, ms.iftype) { + case "ap": + result["access-point"] = m.buildAPData(ctx, iface, si, status) + case "mesh": + result["mesh-point"] = m.buildMeshData(iface, si, status, ms) + default: + result["station"] = m.buildStationData(ctx, iface, si, status) + } + + return result +} + +// kernelState is how to read the kernel side of an interface. hostapd +// only serves access points, so asking nl80211 there is wasted work. +func (m *IWMonitor) kernelState(si wpactrl.SocketInfo) func(string) meshState { + if si.Daemon == "hostapd" { + return func(string) meshState { return meshState{} } + } + return m.meshState +} + +// detectMode picks the operational container. The kernel iftype is +// authoritative, a mesh interface is one before wpa_supplicant has +// joined anything and STATUS says so only afterwards (mode=mesh). +func detectMode(si wpactrl.SocketInfo, status map[string]string, iftype string) string { + if si.Daemon == "hostapd" { + return "ap" + } + if iftype == "mesh_point" || status["mode"] == "mesh" { + return "mesh" + } + return "station" +} + +func (m *IWMonitor) buildAPData(ctx context.Context, iface string, si wpactrl.SocketInfo, status map[string]string) map[string]any { + ap := make(map[string]any) + + ssid := resolveSSID(iface, si, status) + if ssid != "" { + ap["ssid"] = ssid + } + + if si.Daemon == "hostapd" { + conn, err := wpactrl.Dial(si.Path) + if err != nil { + m.log.Warn("hostapd dial for stations", "iface", iface, "err", err) + } else { + stas, err := conn.AllStations() + conn.Close() + if err != nil { + m.log.Warn("hostapd AllStations", "iface", iface, "err", err) + } + if err == nil { + stas = filterAuthorized(stas) + if len(stas) > 0 { + ap["stations"] = m.formatStations(stas) + } + } + } + } + + return ap +} + +func (m *IWMonitor) buildStationData(ctx context.Context, iface string, si wpactrl.SocketInfo, status map[string]string) map[string]any { + sta := make(map[string]any) + + ssid := resolveSSID(iface, si, status) + if ssid != "" { + sta["ssid"] = ssid + } + if bssid := status["bssid"]; bssid != "" && bssid != "00:00:00:00:00:00" { + sta["bssid"] = bssid + } + + if si.Daemon == "wpa_supplicant" { + conn, err := wpactrl.Dial(si.Path) + if err == nil { + poll, err := conn.SignalPoll() + if err == nil { + if v, ok := poll["RSSI"]; ok { + if sig, err := strconv.Atoi(v); err == nil { + sta["signal-strength"] = sig + } + } + if v, ok := poll["LINKSPEED"]; ok { + if speed, err := strconv.ParseUint(v, 10, 32); err == nil { + sta["tx-speed"] = uint32(speed * 10) + } + } + } + results, err := conn.ScanResults() + conn.Close() + if err == nil && len(results) > 0 { + sta["scan-results"] = formatScanResults(results) + } + } + } + + return sta +} + +func (m *IWMonitor) formatStations(stas []map[string]string) map[string]any { + type stationEntry struct { + MAC string `json:"mac-address"` + Signal int16 `json:"signal-strength,omitempty"` + ConnectedTime uint32 `json:"connected-time,omitempty"` + RxPackets string `json:"rx-packets,omitempty"` + TxPackets string `json:"tx-packets,omitempty"` + RxBytes string `json:"rx-bytes,omitempty"` + TxBytes string `json:"tx-bytes,omitempty"` + RxSpeed uint32 `json:"rx-speed,omitempty"` + TxSpeed uint32 `json:"tx-speed,omitempty"` + } + + var out []stationEntry + for _, st := range stas { + s := stationEntry{MAC: st["addr"]} + if v := st["signal"]; v != "" { + if sig, err := strconv.ParseInt(v, 10, 16); err == nil { + s.Signal = int16(sig) + } + } + if v := st["connected_time"]; v != "" { + if ct, err := strconv.ParseUint(v, 10, 32); err == nil { + s.ConnectedTime = uint32(ct) + } + } + if v := st["rx_packets"]; v != "" { + s.RxPackets = v + } + if v := st["tx_packets"]; v != "" { + s.TxPackets = v + } + if v := st["rx_bytes"]; v != "" { + s.RxBytes = v + } + if v := st["tx_bytes"]; v != "" { + s.TxBytes = v + } + if v := st["rx_rate_info"]; v != "" { + if speed := parseBitrate(v); speed > 0 { + s.RxSpeed = speed + } + } + if v := st["tx_rate_info"]; v != "" { + if speed := parseBitrate(v); speed > 0 { + s.TxSpeed = speed + } + } + out = append(out, s) + } + return map[string]any{"station": out} +} + +func (m *IWMonitor) publishWifi(iface string, data map[string]any) { + if m.onUpdate == nil { + return + } + + if data == nil { + m.onUpdate(iface, json.RawMessage(`{}`)) + return + } + + raw, err := json.Marshal(data) + if err != nil { + m.log.Warn("marshal wifi data", "iface", iface, "err", err) + return + } + m.onUpdate(iface, json.RawMessage(raw)) +} + +func (m *IWMonitor) refreshAllInterfaces(ctx context.Context) { + for ifname := range wpactrl.ScanSockets() { + m.startAttach(ctx, ifname) + m.refreshInterface(ctx, ifname) + } +} + +func nextDelay(current time.Duration) time.Duration { + d := time.Duration(math.Min(float64(current)*reconnectFactor, float64(reconnectMax))) + return d +} + +func filterAuthorized(stas []map[string]string) []map[string]string { + var out []map[string]string + for _, st := range stas { + if strings.Contains(st["flags"], "AUTHORIZED") { + out = append(out, st) + } + } + return out +} diff --git a/src/yangerd/internal/iwmonitor/iwmonitor_test.go b/src/yangerd/internal/iwmonitor/iwmonitor_test.go new file mode 100644 index 000000000..4044fa1f4 --- /dev/null +++ b/src/yangerd/internal/iwmonitor/iwmonitor_test.go @@ -0,0 +1,221 @@ +package iwmonitor + +import ( + "encoding/json" + "reflect" + "testing" +) + +func TestParseIWEvent(t *testing.T) { + tests := []struct { + name string + line string + wantOK bool + wantType string + wantIface string + wantPhy string + wantAddr string + }{ + { + name: "new station", + line: "1234567890.123456: wlan0 (phy#0): new station aa:bb:cc:dd:ee:ff", + wantOK: true, + wantType: "new station", + wantIface: "wlan0", + wantPhy: "phy#0", + wantAddr: "aa:bb:cc:dd:ee:ff", + }, + { + name: "del station", + line: "1234567890.123456: wlan0 (phy#0): del station aa:bb:cc:dd:ee:ff", + wantOK: true, + wantType: "del station", + wantIface: "wlan0", + wantPhy: "phy#0", + wantAddr: "aa:bb:cc:dd:ee:ff", + }, + { + name: "connected", + line: "1234567890.123456: wlan0 (phy#0): connected to aa:bb:cc:dd:ee:ff", + wantOK: true, + wantType: "connected", + wantIface: "wlan0", + wantPhy: "phy#0", + wantAddr: "aa:bb:cc:dd:ee:ff", + }, + { + name: "disconnected", + line: "1234567890.123456: wlan0 (phy#0): disconnected", + wantOK: true, + wantType: "disconnected", + wantIface: "wlan0", + wantPhy: "phy#0", + }, + { + name: "channel switch", + line: "1234567890.123456: wlan0 (phy#0): ch_switch_started_notify", + wantOK: true, + wantType: "ch_switch_started_notify", + wantIface: "wlan0", + wantPhy: "phy#0", + }, + { + name: "scan started", + line: "1234567890.123456: wlan0 (phy#0): scan started", + wantOK: true, + wantType: "scan started", + wantIface: "wlan0", + wantPhy: "phy#0", + }, + { + name: "reg change", + line: "1234567890.123456: wlan0 (phy#0): reg_change", + wantOK: true, + wantType: "reg_change", + wantIface: "wlan0", + wantPhy: "phy#0", + }, + { + name: "malformed missing separators", + line: "1234567890.123456 wlan0 (phy#0) new station aa:bb:cc:dd:ee:ff", + wantOK: false, + }, + { + name: "malformed bad timestamp", + line: "not-a-float: wlan0 (phy#0): disconnected", + wantOK: false, + }, + { + name: "malformed missing phy", + line: "1234567890.123456: wlan0 phy#0: disconnected", + wantOK: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, ok := ParseIWEvent(tt.line) + if ok != tt.wantOK { + t.Fatalf("ok = %v, want %v", ok, tt.wantOK) + } + if !tt.wantOK { + return + } + + if got.Type != tt.wantType { + t.Fatalf("Type = %q, want %q", got.Type, tt.wantType) + } + if got.Interface != tt.wantIface { + t.Fatalf("Interface = %q, want %q", got.Interface, tt.wantIface) + } + if got.Phy != tt.wantPhy { + t.Fatalf("Phy = %q, want %q", got.Phy, tt.wantPhy) + } + if got.Addr != tt.wantAddr { + t.Fatalf("Addr = %q, want %q", got.Addr, tt.wantAddr) + } + }) + } +} + +func TestParseStationDump(t *testing.T) { + input := `Station aa:bb:cc:dd:ee:ff (on wlan0) + inactive time: 10 ms + rx bytes: 1234 + tx bytes: 5678 + connected time: 42 seconds + signal: -40 dBm + rx bitrate: 6.5 MBit/s + tx bitrate: 130.0 MBit/s + authorized: yes + +Station 11:22:33:44:55:66 (on wlan0) + inactive time: 20 ms + rx bytes: 9876 + tx bytes: 5432 + connected time: 84 seconds + signal: -55 dBm + authorized: no` + + got := parseStationDump(input) + + var gotDecoded []map[string]string + if err := json.Unmarshal(got, &gotDecoded); err != nil { + t.Fatalf("unmarshal got: %v", err) + } + + want := []map[string]string{ + { + "mac": "aa:bb:cc:dd:ee:ff", + "inactive-time": "10 ms", + "rx-bytes": "1234", + "tx-bytes": "5678", + "connected-time": "42 seconds", + "signal": "-40 dBm", + "rx-bitrate": "6.5 MBit/s", + "tx-bitrate": "130.0 MBit/s", + "authorized": "yes", + }, + { + "mac": "11:22:33:44:55:66", + "inactive-time": "20 ms", + "rx-bytes": "9876", + "tx-bytes": "5432", + "connected-time": "84 seconds", + "signal": "-55 dBm", + "authorized": "no", + }, + } + + if !reflect.DeepEqual(gotDecoded, want) { + t.Fatalf("parseStationDump mismatch\n got: %#v\nwant: %#v", gotDecoded, want) + } +} + +func TestParseIWInfo(t *testing.T) { + input := `Interface wlan0 + ifindex: 4 + wdev: 0x1 + addr: 12:34:56:78:9a:bc + ssid: MyWiFi + type: managed + channel: 11 (2462 MHz), width: 20 MHz, center1: 2462 MHz + txpower: 20.00 dBm` + + got := parseIWInfo(input) + + var gotDecoded map[string]string + if err := json.Unmarshal(got, &gotDecoded); err != nil { + t.Fatalf("unmarshal got: %v", err) + } + + want := map[string]string{ + "ssid": "MyWiFi", + "type": "managed", + "channel": "11 (2462 MHz), width: 20 MHz, center1: 2462 MHz", + "tx-power": "20.00 dBm", + } + + if !reflect.DeepEqual(gotDecoded, want) { + t.Fatalf("parseIWInfo mismatch\n got: %#v\nwant: %#v", gotDecoded, want) + } +} + +func TestParseIWDevList(t *testing.T) { + input := `phy#0 + Interface wlan0 + ifindex 4 + wdev 0x1 + +phy#1 + Interface wlan1 + ifindex 5 + wdev 0x2` + + got := parseIWDevList(input) + want := []string{"wlan0", "wlan1"} + + if !reflect.DeepEqual(got, want) { + t.Fatalf("parseIWDevList mismatch\n got: %#v\nwant: %#v", got, want) + } +} diff --git a/src/yangerd/internal/iwmonitor/mesh.go b/src/yangerd/internal/iwmonitor/mesh.go new file mode 100644 index 000000000..aafef7979 --- /dev/null +++ b/src/yangerd/internal/iwmonitor/mesh.go @@ -0,0 +1,164 @@ +package iwmonitor + +import ( + "net" + "strconv" + "strings" + "unicode" + "unicode/utf8" + + "github.com/kernelkit/infix/src/yangerd/internal/nl80211" + "github.com/kernelkit/infix/src/yangerd/internal/wpactrl" +) + +// meshState is what the kernel knows about a mesh interface, read on +// every refresh. The default implementation asks nl80211, tests +// substitute their own. +type meshState struct { + iftype string + forwarding *bool + peers []nl80211.Station +} + +func kernelMeshState(ifname string) meshState { + var ms meshState + + iface, err := net.InterfaceByName(ifname) + if err != nil { + return ms + } + client, err := nl80211.Dial() + if err != nil { + return ms + } + defer client.Close() + + if iftype, err := client.InterfaceType(iface.Index); err == nil { + ms.iftype = iftype + } + if ms.iftype != "mesh_point" { + return ms + } + if fwd, err := client.MeshForwarding(iface.Index); err == nil { + ms.forwarding = &fwd + } + if peers, err := client.Stations(iface.Index); err == nil { + ms.peers = peers + } + return ms +} + +// meshEvent tells whether a wpa_supplicant event changes mesh state: +// a peer link came or went, or the mesh itself was joined or left. +func meshEvent(ev wpactrl.Event) bool { + return strings.HasPrefix(ev.Name, "MESH-PEER-") || strings.HasPrefix(ev.Name, "MESH-GROUP-") +} + +// buildMeshData assembles the mesh-point container. The mesh id comes +// from wpa_supplicant, which is what joined the mesh, and is absent +// until it has. Forwarding and peers come from the kernel. +func (m *IWMonitor) buildMeshData(iface string, si wpactrl.SocketInfo, status map[string]string, ms meshState) map[string]any { + mesh := make(map[string]any) + + if id := decodeWPASSID(resolveSSID(iface, si, status)); id != "" { + mesh["mesh-id"] = id + } + if ms.forwarding != nil { + mesh["forwarding"] = *ms.forwarding + } + if len(ms.peers) > 0 { + mesh["peers"] = map[string]any{"peer": formatPeers(ms.peers)} + } + + return mesh +} + +type peerEntry struct { + MAC string `json:"mac-address"` + Signal *int16 `json:"signal-strength,omitempty"` + ConnectedTime uint32 `json:"connected-time"` + RxPackets string `json:"rx-packets"` + TxPackets string `json:"tx-packets"` + RxBytes string `json:"rx-bytes"` + TxBytes string `json:"tx-bytes"` + RxSpeed uint32 `json:"rx-speed,omitempty"` + TxSpeed uint32 `json:"tx-speed,omitempty"` +} + +func formatPeers(stas []nl80211.Station) []peerEntry { + out := make([]peerEntry, 0, len(stas)) + for _, st := range stas { + p := peerEntry{ + MAC: st.MAC, + ConnectedTime: st.ConnectedTime, + RxPackets: strconv.FormatUint(st.RxPackets, 10), + TxPackets: strconv.FormatUint(st.TxPackets, 10), + RxBytes: strconv.FormatUint(st.RxBytes, 10), + TxBytes: strconv.FormatUint(st.TxBytes, 10), + RxSpeed: st.RxBitrate, + TxSpeed: st.TxBitrate, + } + if st.HasSignal { + sig := int16(st.Signal) + p.Signal = &sig + } + out = append(out, p) + } + return out +} + +// decodeWPASSID undoes wpa_supplicant's printf_encode(): bytes outside +// printable ASCII arrive as \xHH, plus \\ \" \e \n \r \t. The result is +// taken as UTF-8 when it is valid, otherwise the escaped form is kept. +// Control characters are dropped either way, a rogue peer must not get +// to write escape sequences to a terminal. +func decodeWPASSID(s string) string { + out := s + if raw := unescapePrintf(s); utf8.Valid(raw) { + out = string(raw) + } + return strings.Map(func(r rune) rune { + if unicode.IsPrint(r) { + return r + } + return -1 + }, out) +} + +func unescapePrintf(s string) []byte { + raw := make([]byte, 0, len(s)) + for i := 0; i < len(s); i++ { + c := s[i] + if c != '\\' || i+1 >= len(s) { + raw = append(raw, c) + continue + } + i++ + switch s[i] { + case 'x': + if i+2 < len(s) { + if v, err := strconv.ParseUint(s[i+1:i+3], 16, 8); err == nil { + raw = append(raw, byte(v)) + i += 2 + continue + } + } + raw = append(raw, '\\', 'x') + case '\\': + raw = append(raw, '\\') + case '"': + raw = append(raw, '"') + case 'e': + raw = append(raw, 0x1b) + case 'n': + raw = append(raw, '\n') + case 'r': + raw = append(raw, '\r') + case 't': + raw = append(raw, '\t') + default: + raw = append(raw, '\\', s[i]) + } + } + return raw +} diff --git a/src/yangerd/internal/iwmonitor/mesh_test.go b/src/yangerd/internal/iwmonitor/mesh_test.go new file mode 100644 index 000000000..70a12a739 --- /dev/null +++ b/src/yangerd/internal/iwmonitor/mesh_test.go @@ -0,0 +1,213 @@ +package iwmonitor + +import ( + "context" + "encoding/json" + "github.com/mdlayher/netlink" + "golang.org/x/sys/unix" + "log/slog" + "testing" + + "github.com/kernelkit/infix/src/yangerd/internal/nl80211" + "github.com/kernelkit/infix/src/yangerd/internal/wpactrl" +) + +func TestDetectMode(t *testing.T) { + sup := wpactrl.SocketInfo{Iface: "wifi0", Daemon: "wpa_supplicant"} + ap := wpactrl.SocketInfo{Iface: "wifi0", Daemon: "hostapd"} + + tests := []struct { + name string + si wpactrl.SocketInfo + status map[string]string + iftype string + want string + }{ + {"hostapd", ap, nil, "AP", "ap"}, + {"station joined", sup, map[string]string{"mode": "station"}, "station", "station"}, + {"station idle", sup, map[string]string{}, "station", "station"}, + {"mesh joined", sup, map[string]string{"mode": "mesh"}, "mesh_point", "mesh"}, + {"mesh not yet joined", sup, map[string]string{"wpa_state": "SCANNING"}, "mesh_point", "mesh"}, + {"mesh, kernel unreadable", sup, map[string]string{"mode": "mesh"}, "", "mesh"}, + } + for _, tt := range tests { + if got := detectMode(tt.si, tt.status, tt.iftype); got != tt.want { + t.Errorf("%s: detectMode = %q, want %q", tt.name, got, tt.want) + } + } +} + +func TestDecodeWPASSID(t *testing.T) { + tests := []struct{ in, want string }{ + {"plain", "plain"}, + {"caf\\xc3\\xa9", "café"}, + {"tab\\there", "tabhere"}, + {"esc\\e[31mred", "esc[31mred"}, + {"quote\\\"d", "quote\"d"}, + {"back\\\\slash", "back\\slash"}, + {"bad\\xff\\xfeutf8", "bad\\xff\\xfeutf8"}, + {"trailing\\", "trailing\\"}, + {"", ""}, + } + for _, tt := range tests { + if got := decodeWPASSID(tt.in); got != tt.want { + t.Errorf("decodeWPASSID(%q) = %q, want %q", tt.in, got, tt.want) + } + } +} + +func TestBuildMeshData(t *testing.T) { + m := New(slog.Default()) + si := wpactrl.SocketInfo{Iface: "wifi0", Daemon: "wpa_supplicant"} + fwd := false + sig := int8(-51) + ms := meshState{ + iftype: "mesh_point", + forwarding: &fwd, + peers: []nl80211.Station{{ + MAC: "02:00:00:00:00:02", Signal: sig, HasSignal: true, + ConnectedTime: 42, RxBytes: 1000, TxBytes: 2000, + RxPackets: 10, TxPackets: 20, RxBitrate: 650, TxBitrate: 1200, + }, { + MAC: "02:00:00:00:00:03", + }}, + } + status := map[string]string{"mode": "mesh", "ssid": "backhaul\\x2d1"} + + raw, err := json.Marshal(m.buildMeshData("wifi0", si, status, ms)) + if err != nil { + t.Fatal(err) + } + + var got struct { + MeshID string `json:"mesh-id"` + Forwarding *bool `json:"forwarding"` + Peers struct { + Peer []map[string]any `json:"peer"` + } `json:"peers"` + } + if err := json.Unmarshal(raw, &got); err != nil { + t.Fatal(err) + } + if got.MeshID != "backhaul-1" { + t.Errorf("mesh-id = %q", got.MeshID) + } + if got.Forwarding == nil || *got.Forwarding { + t.Errorf("forwarding = %v, want false", got.Forwarding) + } + if len(got.Peers.Peer) != 2 { + t.Fatalf("peers = %v", got.Peers.Peer) + } + p := got.Peers.Peer[0] + if p["mac-address"] != "02:00:00:00:00:02" || p["signal-strength"] != float64(-51) || + p["connected-time"] != float64(42) || p["rx-bytes"] != "1000" || p["tx-packets"] != "20" || + p["rx-speed"] != float64(650) || p["tx-speed"] != float64(1200) { + t.Errorf("peer = %v", p) + } + if _, has := got.Peers.Peer[1]["signal-strength"]; has { + t.Errorf("peer without signal must omit signal-strength: %v", got.Peers.Peer[1]) + } +} + +// Before wpa_supplicant has joined, the mesh interface has no mesh-id and +// no peers, but is still reported as a mesh point. +func TestBuildMeshDataNotJoined(t *testing.T) { + m := New(slog.Default()) + si := wpactrl.SocketInfo{Iface: "wifi0", Daemon: "wpa_supplicant"} + fwd := true + + data := m.buildMeshData("wifi0", si, map[string]string{"wpa_state": "SCANNING"}, + meshState{iftype: "mesh_point", forwarding: &fwd}) + + if _, has := data["mesh-id"]; has { + t.Errorf("mesh-id must be absent until joined: %v", data) + } + if _, has := data["peers"]; has { + t.Errorf("peers must be absent with no peers: %v", data) + } + if data["forwarding"] != true { + t.Errorf("forwarding = %v", data["forwarding"]) + } +} + +func TestMeshEventsRefresh(t *testing.T) { + for _, line := range []string{ + "<3>MESH-PEER-CONNECTED 02:00:00:00:00:02", + "<3>MESH-PEER-DISCONNECTED 02:00:00:00:00:02", + "<3>MESH-GROUP-STARTED ssid=\"backhaul\" id=0", + "<3>MESH-GROUP-REMOVED wifi0", + } { + ev, ok := wpactrl.ParseEvent(line) + if !ok { + t.Fatalf("ParseEvent(%q) failed", line) + } + if got := meshEvent(ev); !got { + t.Errorf("%q must trigger a refresh", ev.Name) + } + } + ev, _ := wpactrl.ParseEvent("<3>CTRL-EVENT-SCAN-RESULTS ") + if meshEvent(ev) { + t.Errorf("%q is not a mesh event", ev.Name) + } +} + +// hostapd serves access points only, so the kernel is not asked. +func TestKernelStateSkipsHostapd(t *testing.T) { + calls := 0 + m := New(slog.Default()) + m.meshState = func(string) meshState { + calls++ + return meshState{iftype: "mesh_point"} + } + + if ms := m.kernelState(wpactrl.SocketInfo{Daemon: "hostapd"})("wifi0"); ms.iftype != "" || calls != 0 { + t.Fatalf("hostapd: iftype %q, %d kernel queries", ms.iftype, calls) + } + if ms := m.kernelState(wpactrl.SocketInfo{Daemon: "wpa_supplicant"})("wifi0"); ms.iftype != "mesh_point" || calls != 1 { + t.Fatalf("wpa_supplicant: iftype %q, %d kernel queries", ms.iftype, calls) + } +} + +// Leaving the mesh keeps the mesh-point container, a station that +// disconnects loses its container. +func TestDisconnectedMeshStaysMeshPoint(t *testing.T) { + fwd := true + for _, tc := range []struct { + iftype string + want string + }{ + {"mesh_point", `{"mesh-point":{"forwarding":true}}`}, + {"station", `{}`}, + } { + m := New(slog.Default()) + m.meshState = func(string) meshState { + return meshState{iftype: tc.iftype, forwarding: &fwd} + } + var got string + m.SetOnUpdate(func(_ string, raw json.RawMessage) { got = string(raw) }) + + ev, _ := wpactrl.ParseEvent("<3>CTRL-EVENT-DISCONNECTED bssid=02:00:00:00:00:02 reason=3") + m.handleAttachEvent(context.Background(), "wifi-test-none", ev) + + if got != tc.want { + t.Errorf("%s: published %s, want %s", tc.iftype, got, tc.want) + } + } +} + +// DEL_INTERFACE arrives after the interface is gone, so its index no +// longer resolves; the name must come from the message. +func TestExtractIfnameFromDeletedInterface(t *testing.T) { + ae := netlink.NewAttributeEncoder() + ae.Uint32(unix.NL80211_ATTR_IFINDEX, 999999) + ae.String(unix.NL80211_ATTR_IFNAME, "wifi0") + data, err := ae.Encode() + if err != nil { + t.Fatal(err) + } + + m := New(slog.Default()) + if got := m.extractIfname(data); got != "wifi0" { + t.Fatalf("extractIfname = %q, want wifi0", got) + } +} diff --git a/src/yangerd/internal/kernelrib/kernelrib.go b/src/yangerd/internal/kernelrib/kernelrib.go new file mode 100644 index 000000000..d459be0fa --- /dev/null +++ b/src/yangerd/internal/kernelrib/kernelrib.go @@ -0,0 +1,355 @@ +// Package kernelrib fills the ietf-routing ribs from the kernel FIB over +// rtnetlink. It is the RIB source for builds without FRR, where netd +// installs static and DHCP routes straight into the kernel. +// +// Route notifications are only a trigger: on each burst the main table +// is re-read in full and the ribs subtree replaced, the same shape the +// zapiwatcher uses with zebra. Every route in the kernel table is in +// the FIB, so every next-hop is installed, and of the routes to one +// prefix the lowest metric is the active one. netd installs static +// routes with the configured route preference as metric. +package kernelrib + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "log/slog" + "net" + "strconv" + "strings" + "time" + + "github.com/vishvananda/netlink" + "golang.org/x/sys/unix" + + "github.com/kernelkit/infix/src/yangerd/internal/backoff" + "github.com/kernelkit/infix/src/yangerd/internal/tree" +) + +const ( + routingTreeKey = "ietf-routing:routing" + debounceDelay = 100 * time.Millisecond +) + +// Lister returns the routes of one address family from the main table. +type Lister func(family int) ([]netlink.Route, error) + +// IfName resolves an interface index to its name, "" when unknown. +type IfName func(index int) string + +// Watcher mirrors the kernel main routing table into the ribs subtree. +type Watcher struct { + tree *tree.Tree + log *slog.Logger + list Lister + ifname IfName + now func() time.Time + refresh chan struct{} + seen map[string]time.Time // first sighting per route, for last-updated +} + +// New creates a Watcher reading the kernel table over rtnetlink. +func New(t *tree.Tree, log *slog.Logger) *Watcher { + if log == nil { + log = slog.Default() + } + return &Watcher{ + tree: t, + log: log, + list: listMainTable, + ifname: ifNameByIndex, + now: time.Now, + refresh: make(chan struct{}, 1), + seen: map[string]time.Time{}, + } +} + +func listMainTable(family int) ([]netlink.Route, error) { + filter := &netlink.Route{Table: unix.RT_TABLE_MAIN} + return netlink.RouteListFiltered(family, filter, netlink.RT_FILTER_TABLE) +} + +func ifNameByIndex(index int) string { + iface, err := net.InterfaceByIndex(index) + if err != nil { + return "" + } + return iface.Name +} + +// Run subscribes to route changes and keeps the ribs current until ctx +// ends. A broken subscription is re-established with backoff. +func (w *Watcher) Run(ctx context.Context) error { + // The refresh worker owns all writes to the tree and runs for the + // lifetime of the watcher, independent of the subscription. + go w.refreshLoop(ctx) + + return backoff.Retry(ctx, w.log, "kernel rib", w.session) +} + +// session runs one rtnetlink subscription until it fails or ctx ends. +func (w *Watcher) session(ctx context.Context) error { + sessCtx, cancel := context.WithCancel(ctx) + defer cancel() + + done := make(chan struct{}) + defer close(done) + + var subErr error + errorCallback := func(err error) { + if err == nil { + return + } + subErr = err + cancel() + } + + updates := make(chan netlink.RouteUpdate, 256) + err := netlink.RouteSubscribeWithOptions(updates, done, netlink.RouteSubscribeOptions{ + ErrorCallback: errorCallback, + ReceiveBufferSize: 1 << 20, + ReceiveBufferForceSize: true, + }) + if err != nil { + return fmt.Errorf("subscribe route updates: %w", err) + } + + w.log.Info("kernel rib: subscribed to route changes") + + // Read the current table now that we are subscribed, so we have + // data even if no further events arrive. + w.triggerRefresh() + + for { + select { + case <-sessCtx.Done(): + if ctx.Err() != nil { + return ctx.Err() + } + if subErr != nil { + return fmt.Errorf("route subscription: %w", subErr) + } + return errors.New("route subscription ended") + case _, ok := <-updates: + if !ok { + return errors.New("route subscription closed") + } + w.triggerRefresh() + } + } +} + +func (w *Watcher) triggerRefresh() { + select { + case w.refresh <- struct{}{}: + default: + } +} + +func (w *Watcher) refreshLoop(ctx context.Context) { + for { + select { + case <-ctx.Done(): + return + case <-w.refresh: + } + + // Let a burst of notifications settle before reading. + select { + case <-ctx.Done(): + return + case <-time.After(debounceDelay): + } + // Drain a request that arrived during the debounce window; the + // upcoming read already reflects it. + select { + case <-w.refresh: + default: + } + + w.writeRibs() + } +} + +// writeRibs reads the IPv4 and IPv6 main tables and replaces the ribs +// subtree. On a read error it leaves the previous data untouched +// rather than blanking the table. +func (w *Watcher) writeRibs() { + v4, err := w.list(netlink.FAMILY_V4) + if err != nil { + w.log.Warn("kernel rib: read ipv4 routes", "err", err) + return + } + v6, err := w.list(netlink.FAMILY_V6) + if err != nil { + w.log.Warn("kernel rib: read ipv6 routes", "err", err) + return + } + + now := w.now() + seen := make(map[string]time.Time, len(v4)+len(v6)) + ipv4 := w.routes("ipv4", v4, now, seen) + ipv6 := w.routes("ipv6", v6, now, seen) + w.seen = seen + + ribs := map[string]any{ + "rib": []map[string]any{ + { + "name": "ipv4", + "address-family": "ietf-routing:ipv4", + "routes": map[string]any{"route": ipv4}, + }, + { + "name": "ipv6", + "address-family": "ietf-routing:ipv6", + "routes": map[string]any{"route": ipv6}, + }, + }, + } + + data, err := json.Marshal(map[string]any{"ribs": ribs}) + if err != nil { + w.log.Error("kernel rib: marshal ribs", "err", err) + return + } + + w.tree.Merge(routingTreeKey, data) +} + +// routes converts one family's kernel routes into ietf-routing route +// nodes. seen collects the first-sighting time of every route kept. +func (w *Watcher) routes(family string, routes []netlink.Route, now time.Time, seen map[string]time.Time) []map[string]any { + best := map[string]int{} + for _, r := range routes { + dst := prefix(family, r.Dst) + if m, ok := best[dst]; !ok || r.Priority < m { + best[dst] = r.Priority + } + } + + out := make([]map[string]any, 0, len(routes)) + for _, r := range routes { + node := w.transform(family, r) + dst := node[destinationKey(family)].(string) + if r.Priority == best[dst] { + node["active"] = []any{nil} + } + + key := routeKey(family, r) + first, ok := w.seen[key] + if !ok { + first = now + } + seen[key] = first + node["last-updated"] = first.Format(time.RFC3339) + + out = append(out, node) + } + return out +} + +func destinationKey(family string) string { + return "ietf-" + family + "-unicast-routing:destination-prefix" +} + +// prefix renders the destination; a nil destination is the default route. +func prefix(family string, dst *net.IPNet) string { + if dst == nil { + if family == "ipv6" { + return "::/0" + } + return "0.0.0.0/0" + } + return dst.String() +} + +// protocolName maps rtnetlink route protocols to IETF routing-protocol +// identities. Connected routes are what the kernel installs itself; +// everything not modelled falls back to kernel so it still validates. +func protocolName(p netlink.RouteProtocol) string { + switch int(p) { + case unix.RTPROT_KERNEL: + return "ietf-routing:direct" + case unix.RTPROT_STATIC: + return "ietf-routing:static" + default: + return "infix-routing:kernel" + } +} + +// transform converts one kernel route into an ietf-routing route node, +// without the active and last-updated leaves which need the whole table. +func (w *Watcher) transform(family string, r netlink.Route) map[string]any { + addrKey := "ietf-" + family + "-unicast-routing:address" + + node := map[string]any{ + destinationKey(family): prefix(family, r.Dst), + "source-protocol": protocolName(r.Protocol), + "route-preference": r.Priority, + } + + switch r.Type { + case unix.RTN_BLACKHOLE: + node["next-hop"] = map[string]any{"special-next-hop": "blackhole"} + return node + case unix.RTN_UNREACHABLE: + node["next-hop"] = map[string]any{"special-next-hop": "unreachable"} + return node + case unix.RTN_PROHIBIT: + node["next-hop"] = map[string]any{"special-next-hop": "prohibit"} + return node + } + + hops := make([]map[string]any, 0, 1) + add := func(gw net.IP, index int) { + hop := map[string]any{"infix-routing:installed": []any{nil}} + if len(gw) > 0 { + hop[addrKey] = gw.String() + } else if name := w.ifname(index); name != "" { + hop["outgoing-interface"] = name + } else { + return + } + hops = append(hops, hop) + } + + if len(r.MultiPath) > 0 { + for _, nh := range r.MultiPath { + add(nh.Gw, nh.LinkIndex) + } + } else { + add(r.Gw, r.LinkIndex) + } + + if len(hops) > 0 { + node["next-hop"] = map[string]any{ + "next-hop-list": map[string]any{"next-hop": hops}, + } + } + return node +} + +// routeKey identifies a route across refreshes so its first sighting +// survives: destination, metric, protocol and the set of next-hops. +func routeKey(family string, r netlink.Route) string { + var sb strings.Builder + sb.WriteString(family) + sb.WriteByte('|') + sb.WriteString(prefix(family, r.Dst)) + sb.WriteByte('|') + sb.WriteString(strconv.Itoa(r.Priority)) + sb.WriteByte('|') + sb.WriteString(strconv.Itoa(int(r.Protocol))) + sb.WriteByte('|') + sb.WriteString(strconv.Itoa(r.Type)) + if len(r.MultiPath) > 0 { + for _, nh := range r.MultiPath { + fmt.Fprintf(&sb, "|%s@%d", nh.Gw, nh.LinkIndex) + } + } else { + fmt.Fprintf(&sb, "|%s@%d", r.Gw, r.LinkIndex) + } + return sb.String() +} diff --git a/src/yangerd/internal/kernelrib/kernelrib_test.go b/src/yangerd/internal/kernelrib/kernelrib_test.go new file mode 100644 index 000000000..bab2f5ba3 --- /dev/null +++ b/src/yangerd/internal/kernelrib/kernelrib_test.go @@ -0,0 +1,282 @@ +package kernelrib + +import ( + "encoding/json" + "errors" + "log/slog" + "net" + "testing" + "time" + + "github.com/vishvananda/netlink" + "golang.org/x/sys/unix" + + "github.com/kernelkit/infix/src/yangerd/internal/tree" +) + +func cidr(t *testing.T, s string) *net.IPNet { + t.Helper() + _, n, err := net.ParseCIDR(s) + if err != nil { + t.Fatal(err) + } + return n +} + +func fakeIfName(index int) string { + switch index { + case 2: + return "e0" + case 3: + return "e1" + } + return "" +} + +func newWatcher(t *testing.T, v4, v6 []netlink.Route, err error) (*Watcher, *tree.Tree) { + t.Helper() + tr := tree.New() + w := New(tr, slog.New(slog.NewTextHandler(testWriter{t}, nil))) + w.ifname = fakeIfName + w.list = func(family int) ([]netlink.Route, error) { + if err != nil { + return nil, err + } + if family == netlink.FAMILY_V6 { + return v6, nil + } + return v4, nil + } + return w, tr +} + +type testWriter struct{ t *testing.T } + +func (w testWriter) Write(p []byte) (int, error) { w.t.Log(string(p)); return len(p), nil } + +func ribRoutes(t *testing.T, tr *tree.Tree, name string) []map[string]any { + t.Helper() + data := tr.Get(routingTreeKey) + if data == nil { + t.Fatal("routing tree key not set") + } + var routing map[string]any + if err := json.Unmarshal(data, &routing); err != nil { + t.Fatalf("unmarshal routing: %v", err) + } + for _, rib := range routing["ribs"].(map[string]any)["rib"].([]any) { + rm := rib.(map[string]any) + if rm["name"] != name { + continue + } + out := []map[string]any{} + for _, r := range rm["routes"].(map[string]any)["route"].([]any) { + out = append(out, r.(map[string]any)) + } + return out + } + t.Fatalf("rib %s missing", name) + return nil +} + +func findRoute(routes []map[string]any, family, dst string) map[string]any { + for _, r := range routes { + if r["ietf-"+family+"-unicast-routing:destination-prefix"] == dst { + return r + } + } + return nil +} + +func hops(r map[string]any) []map[string]any { + nh := r["next-hop"].(map[string]any) + list := nh["next-hop-list"].(map[string]any)["next-hop"].([]any) + out := make([]map[string]any, 0, len(list)) + for _, h := range list { + out = append(out, h.(map[string]any)) + } + return out +} + +func TestStaticAndConnectedRoutes(t *testing.T) { + v4 := []netlink.Route{ + {Dst: nil, Gw: net.ParseIP("192.168.1.1"), LinkIndex: 2, Protocol: unix.RTPROT_STATIC, Priority: 5}, + {Dst: cidr(t, "192.168.1.0/24"), LinkIndex: 2, Protocol: unix.RTPROT_KERNEL, Scope: unix.RT_SCOPE_LINK}, + {Dst: cidr(t, "10.0.0.0/8"), LinkIndex: 3, Protocol: unix.RTPROT_BOOT, Priority: 100}, + } + w, tr := newWatcher(t, v4, nil, nil) + w.writeRibs() + + routes := ribRoutes(t, tr, "ipv4") + if len(routes) != 3 { + t.Fatalf("want 3 routes, got %d", len(routes)) + } + + def := findRoute(routes, "ipv4", "0.0.0.0/0") + if def == nil { + t.Fatal("default route missing") + } + if def["source-protocol"] != "ietf-routing:static" { + t.Errorf("default source-protocol = %v", def["source-protocol"]) + } + if def["route-preference"] != float64(5) { + t.Errorf("default route-preference = %v", def["route-preference"]) + } + if _, ok := def["active"]; !ok { + t.Error("default route not active") + } + h := hops(def) + if len(h) != 1 || h[0]["ietf-ipv4-unicast-routing:address"] != "192.168.1.1" { + t.Errorf("default next-hop = %v", h) + } + if _, ok := h[0]["infix-routing:installed"]; !ok { + t.Error("default next-hop not installed") + } + if _, ok := h[0]["outgoing-interface"]; ok { + t.Error("gateway hop must not also carry the interface") + } + + conn := findRoute(routes, "ipv4", "192.168.1.0/24") + if conn["source-protocol"] != "ietf-routing:direct" { + t.Errorf("connected source-protocol = %v", conn["source-protocol"]) + } + if h := hops(conn); len(h) != 1 || h[0]["outgoing-interface"] != "e0" { + t.Errorf("connected next-hop = %v", h) + } + + boot := findRoute(routes, "ipv4", "10.0.0.0/8") + if boot["source-protocol"] != "infix-routing:kernel" { + t.Errorf("boot source-protocol = %v", boot["source-protocol"]) + } + + if v6 := ribRoutes(t, tr, "ipv6"); len(v6) != 0 { + t.Errorf("want empty ipv6 rib, got %v", v6) + } +} + +func TestLowestMetricIsActive(t *testing.T) { + v4 := []netlink.Route{ + {Gw: net.ParseIP("192.168.1.1"), LinkIndex: 2, Protocol: unix.RTPROT_STATIC, Priority: 120}, + {Gw: net.ParseIP("192.168.2.1"), LinkIndex: 3, Protocol: unix.RTPROT_STATIC, Priority: 5}, + } + w, tr := newWatcher(t, v4, nil, nil) + w.writeRibs() + + active := 0 + for _, r := range ribRoutes(t, tr, "ipv4") { + _, isActive := r["active"] + if isActive { + active++ + if r["route-preference"] != float64(5) { + t.Errorf("active route has preference %v, want 5", r["route-preference"]) + } + } + } + if active != 1 { + t.Errorf("want exactly one active default route, got %d", active) + } +} + +func TestSpecialAndMultipathNextHops(t *testing.T) { + v4 := []netlink.Route{ + {Dst: cidr(t, "10.1.0.0/16"), Type: unix.RTN_BLACKHOLE, Protocol: unix.RTPROT_STATIC}, + {Dst: cidr(t, "10.2.0.0/16"), Type: unix.RTN_UNREACHABLE, Protocol: unix.RTPROT_STATIC}, + {Dst: cidr(t, "10.3.0.0/16"), Type: unix.RTN_PROHIBIT, Protocol: unix.RTPROT_STATIC}, + {Dst: cidr(t, "10.4.0.0/16"), Protocol: unix.RTPROT_STATIC, MultiPath: []*netlink.NexthopInfo{ + {Gw: net.ParseIP("192.168.1.1"), LinkIndex: 2}, + {LinkIndex: 3}, + }}, + } + w, tr := newWatcher(t, v4, nil, nil) + w.writeRibs() + routes := ribRoutes(t, tr, "ipv4") + + for dst, want := range map[string]string{ + "10.1.0.0/16": "blackhole", + "10.2.0.0/16": "unreachable", + "10.3.0.0/16": "prohibit", + } { + r := findRoute(routes, "ipv4", dst) + if got := r["next-hop"].(map[string]any)["special-next-hop"]; got != want { + t.Errorf("%s special-next-hop = %v, want %s", dst, got, want) + } + } + + h := hops(findRoute(routes, "ipv4", "10.4.0.0/16")) + if len(h) != 2 { + t.Fatalf("want 2 multipath hops, got %v", h) + } + if h[0]["ietf-ipv4-unicast-routing:address"] != "192.168.1.1" || h[1]["outgoing-interface"] != "e1" { + t.Errorf("multipath hops = %v", h) + } +} + +func TestIPv6Routes(t *testing.T) { + v6 := []netlink.Route{ + {Gw: net.ParseIP("fe80::1"), LinkIndex: 2, Protocol: unix.RTPROT_STATIC, Priority: 1}, + {Dst: cidr(t, "2001:db8::/64"), LinkIndex: 2, Protocol: unix.RTPROT_KERNEL, Priority: 256}, + } + w, tr := newWatcher(t, nil, v6, nil) + w.writeRibs() + routes := ribRoutes(t, tr, "ipv6") + + def := findRoute(routes, "ipv6", "::/0") + if def == nil { + t.Fatal("ipv6 default route missing") + } + if h := hops(def); h[0]["ietf-ipv6-unicast-routing:address"] != "fe80::1" { + t.Errorf("ipv6 default next-hop = %v", h) + } + if conn := findRoute(routes, "ipv6", "2001:db8::/64"); conn["source-protocol"] != "ietf-routing:direct" { + t.Errorf("ipv6 connected source-protocol = %v", conn["source-protocol"]) + } +} + +func TestReadErrorKeepsPreviousData(t *testing.T) { + v4 := []netlink.Route{{Gw: net.ParseIP("192.168.1.1"), LinkIndex: 2, Protocol: unix.RTPROT_STATIC}} + w, tr := newWatcher(t, v4, nil, nil) + w.writeRibs() + + w.list = func(int) ([]netlink.Route, error) { return nil, errors.New("netlink down") } + w.writeRibs() + + if routes := ribRoutes(t, tr, "ipv4"); len(routes) != 1 { + t.Errorf("want previous route kept, got %v", routes) + } +} + +func TestLastUpdatedSurvivesRefresh(t *testing.T) { + v4 := []netlink.Route{{Gw: net.ParseIP("192.168.1.1"), LinkIndex: 2, Protocol: unix.RTPROT_STATIC}} + w, tr := newWatcher(t, v4, nil, nil) + + t0 := time.Date(2026, 10, 5, 8, 0, 0, 0, time.UTC) + w.now = func() time.Time { return t0 } + w.writeRibs() + first := ribRoutes(t, tr, "ipv4")[0]["last-updated"] + + w.now = func() time.Time { return t0.Add(time.Hour) } + w.writeRibs() + if again := ribRoutes(t, tr, "ipv4")[0]["last-updated"]; again != first { + t.Errorf("last-updated changed on refresh: %v -> %v", first, again) + } + + // A changed next-hop is a new route and gets a new timestamp. + w.list = func(int) ([]netlink.Route, error) { + return []netlink.Route{{Gw: net.ParseIP("192.168.1.2"), LinkIndex: 2, Protocol: unix.RTPROT_STATIC}}, nil + } + w.writeRibs() + if changed := ribRoutes(t, tr, "ipv4")[0]["last-updated"]; changed == first { + t.Error("last-updated not refreshed for a replaced route") + } +} + +func TestUnknownInterfaceHopIsDropped(t *testing.T) { + v4 := []netlink.Route{{Dst: cidr(t, "10.0.0.0/8"), LinkIndex: 99, Protocol: unix.RTPROT_KERNEL}} + w, tr := newWatcher(t, v4, nil, nil) + w.writeRibs() + + r := ribRoutes(t, tr, "ipv4")[0] + if _, ok := r["next-hop"]; ok { + t.Errorf("want no next-hop for unresolvable interface, got %v", r["next-hop"]) + } +} diff --git a/src/yangerd/internal/lldpmonitor/events_test.go b/src/yangerd/internal/lldpmonitor/events_test.go new file mode 100644 index 000000000..3a94e416f --- /dev/null +++ b/src/yangerd/internal/lldpmonitor/events_test.go @@ -0,0 +1,32 @@ +package lldpmonitor + +import ( + "strings" + "testing" + + "github.com/kernelkit/infix/src/yangerd/internal/tree" +) + +// A brace inside a peer's free-text description must not stall the +// framing: the event after it still triggers a refresh. +func TestReadEventsBraceInDescription(t *testing.T) { + m := New(tree.New(), nil) + m.refresh = make(chan struct{}, 8) + + stream := `{ + "lldp-added": { + "lldp": [{"interface": [{"name": "e1", + "chassis": [{"descr": [{"value": "switch {rack 4"}]}]}]}] + } +} + +{"lldp-deleted": {"lldp": []}} +` + err := m.readEvents(strings.NewReader(stream)) + if err == nil || !strings.Contains(err.Error(), "exited") { + t.Fatalf("readEvents at EOF = %v", err) + } + if n := len(m.refresh); n != 2 { + t.Fatalf("refreshes triggered = %d, want 2", n) + } +} diff --git a/src/yangerd/internal/lldpmonitor/lldpmonitor.go b/src/yangerd/internal/lldpmonitor/lldpmonitor.go new file mode 100644 index 000000000..db494acb9 --- /dev/null +++ b/src/yangerd/internal/lldpmonitor/lldpmonitor.go @@ -0,0 +1,435 @@ +// Package lldpmonitor keeps the LLDP neighbor table in the tree in sync +// with lldpd. A persistent `lldpcli -f json0 watch` subprocess is used +// purely as a change trigger -- its events carry only the changed +// neighbor, so they cannot be used to rebuild state (a delete event +// would re-add the neighbor, and an event for one port would wipe the +// others). On every event the full table is re-read with +// `lldpcli -f json0 show neighbors` and the subtree replaced, so removed +// neighbors disappear and neighbors present before yangerd started are +// picked up. +package lldpmonitor + +import ( + "context" + "encoding/json" + "fmt" + "io" + "log/slog" + "os/exec" + "regexp" + "strconv" + "time" + + "github.com/kernelkit/infix/src/yangerd/internal/backoff" + "github.com/kernelkit/infix/src/yangerd/internal/tree" +) + +const ( + lldpMulticastMAC = "01:80:C2:00:00:0E" + treeKey = "ieee802-dot1ab-lldp:lldp" + + // debounceDelay coalesces bursts of watch events into one re-read. + debounceDelay = 200 * time.Millisecond + queryTimeout = 5 * time.Second +) + +// LLDPMonitor subscribes to LLDP neighbor events via a persistent +// lldpcli watch subprocess and re-reads the full neighbor table on +// every event. +type LLDPMonitor struct { + tree *tree.Tree + log *slog.Logger + refresh chan struct{} + + // query returns the current full neighbor table; overridable in tests. + query func(ctx context.Context) ([]byte, error) +} + +// New creates an LLDPMonitor. +func New(t *tree.Tree, log *slog.Logger) *LLDPMonitor { + if log == nil { + log = slog.Default() + } + return &LLDPMonitor{ + tree: t, + log: log, + refresh: make(chan struct{}, 1), + query: queryNeighbors, + } +} + +func queryNeighbors(ctx context.Context) ([]byte, error) { + ctx, cancel := context.WithTimeout(ctx, queryTimeout) + defer cancel() + return exec.CommandContext(ctx, "lldpcli", "-f", "json0", "show", "neighbors").Output() +} + +// Run starts the LLDP monitor. It blocks until ctx is cancelled. +func (m *LLDPMonitor) Run(ctx context.Context) error { + go m.refreshLoop(ctx) + return backoff.Retry(ctx, m.log, "lldp monitor", m.runOnce) +} + +func (m *LLDPMonitor) runOnce(ctx context.Context) error { + cmd := exec.CommandContext(ctx, "lldpcli", "-f", "json0", "watch") + stdout, err := cmd.StdoutPipe() + if err != nil { + return fmt.Errorf("stdout pipe: %w", err) + } + if err := cmd.Start(); err != nil { + return fmt.Errorf("start lldpcli watch: %w", err) + } + defer cmd.Wait() + + // Pick up neighbors that existed before we attached. + m.triggerRefresh() + + return m.readEvents(stdout) +} + +// readEvents dispatches each JSON value lldpcli prints. A decoder, not +// line or brace counting, frames them: neighbor descriptions are free +// text from the peer and may hold anything. +func (m *LLDPMonitor) readEvents(r io.Reader) error { + dec := json.NewDecoder(r) + for { + var ev json.RawMessage + if err := dec.Decode(&ev); err != nil { + if err == io.EOF { + return fmt.Errorf("lldpcli watch process exited") + } + return fmt.Errorf("read lldpcli: %w", err) + } + m.processEvent(ev) + } +} + +// processEvent inspects a watch event and triggers a full table re-read. +// The event payload itself is never used to build state. +func (m *LLDPMonitor) processEvent(data []byte) { + var raw map[string]json.RawMessage + if err := json.Unmarshal(data, &raw); err != nil { + m.log.Warn("lldp monitor: parse event", "err", err) + return + } + + for key := range raw { + switch key { + case "lldp-added", "lldp-updated", "lldp-deleted": + m.log.Debug("lldp monitor: neighbor change", "event", key) + m.triggerRefresh() + return + } + } + m.log.Debug("lldp monitor: unknown event keys", "keys", keysOf(raw)) +} + +// triggerRefresh requests a table re-read; the buffered channel collapses +// pending requests into one. +func (m *LLDPMonitor) triggerRefresh() { + select { + case m.refresh <- struct{}{}: + default: + } +} + +func (m *LLDPMonitor) refreshLoop(ctx context.Context) { + for { + select { + case <-ctx.Done(): + return + case <-m.refresh: + } + + // Let a burst of events settle before reading. + select { + case <-ctx.Done(): + return + case <-time.After(debounceDelay): + } + select { + case <-m.refresh: + default: + } + + m.updateTree(ctx) + } +} + +// updateTree reads the full neighbor table and replaces the subtree. +// On a query error the previous data is left untouched. +func (m *LLDPMonitor) updateTree(ctx context.Context) { + out, err := m.query(ctx) + if err != nil { + m.log.Warn("lldp monitor: show neighbors", "err", err) + return + } + + m.tree.Set(treeKey, transformNeighbors(out)) + m.log.Debug("lldp monitor: tree updated") +} + +// j0ID is a chassis/port id element: {"type": "mac", "value": "..."}. +type j0ID struct { + Type string `json:"type"` + Value string `json:"value"` +} + +// j0Iface is one neighbor entry on an interface. In json0 format the +// interface name, rid and age are plain string fields and chassis/port +// are arrays. In the older keyed json format chassis/port are objects; +// the custom unmarshallers accept both. +type j0Iface struct { + Name string `json:"name"` + RID interface{} `json:"rid"` + Age string `json:"age"` + Chassis j0IDHolder `json:"chassis"` + Port j0IDHolder `json:"port"` +} + +// j0IDHolder extracts the first id from a chassis/port node in either +// json0 (array) or json (object) form. +type j0IDHolder struct { + ID j0ID +} + +func (h *j0IDHolder) UnmarshalJSON(data []byte) error { + var asArray []struct { + ID json.RawMessage `json:"id"` + } + if err := json.Unmarshal(data, &asArray); err == nil { + for _, e := range asArray { + if id, ok := parseID(e.ID); ok { + h.ID = id + return nil + } + } + return nil + } + + var asObject struct { + ID json.RawMessage `json:"id"` + } + if err := json.Unmarshal(data, &asObject); err != nil { + return nil // tolerate unknown shapes + } + if id, ok := parseID(asObject.ID); ok { + h.ID = id + } + return nil +} + +// parseID accepts an id as object {"type","value"} or array of such. +func parseID(raw json.RawMessage) (j0ID, bool) { + if len(raw) == 0 { + return j0ID{}, false + } + var one j0ID + if err := json.Unmarshal(raw, &one); err == nil && (one.Type != "" || one.Value != "") { + return one, true + } + var many []j0ID + if err := json.Unmarshal(raw, &many); err == nil && len(many) > 0 { + return many[0], true + } + return j0ID{}, false +} + +// collectIfaces extracts all neighbor interface entries from a show +// neighbors document, accepting both json0 ("lldp" is an array, entries +// carry a "name" field) and json ("lldp" is an object, entries are keyed +// by interface name) output formats. +func collectIfaces(data []byte) []j0Iface { + var ifaces []j0Iface + + addRaw := func(name string, raw json.RawMessage) { + var iface j0Iface + if err := json.Unmarshal(raw, &iface); err != nil { + return + } + if iface.Name == "" { + iface.Name = name + } + if iface.Name != "" { + ifaces = append(ifaces, iface) + } + } + + collectInterfaceNode := func(raw json.RawMessage) { + // json0: array of {"name": "eth0", ...} + var asArray []json.RawMessage + if err := json.Unmarshal(raw, &asArray); err == nil { + for _, e := range asArray { + // Either a direct entry with "name", or a keyed map + // {"eth0": {...}} from the older json format. + var iface j0Iface + if err := json.Unmarshal(e, &iface); err == nil && iface.Name != "" { + ifaces = append(ifaces, iface) + continue + } + var keyed map[string]json.RawMessage + if err := json.Unmarshal(e, &keyed); err == nil { + for name, v := range keyed { + addRaw(name, v) + } + } + } + return + } + // json: single keyed map {"eth0": {...}} + var keyed map[string]json.RawMessage + if err := json.Unmarshal(raw, &keyed); err == nil { + for name, v := range keyed { + addRaw(name, v) + } + } + } + + var doc map[string]json.RawMessage + if err := json.Unmarshal(data, &doc); err != nil { + return nil + } + lldpRaw, ok := doc["lldp"] + if !ok { + return nil + } + + // json0: "lldp" is an array of {"interface": [...]}; json: an object. + var lldpArray []map[string]json.RawMessage + if err := json.Unmarshal(lldpRaw, &lldpArray); err == nil { + for _, entry := range lldpArray { + if ifRaw, ok := entry["interface"]; ok { + collectInterfaceNode(ifRaw) + } + } + return ifaces + } + var lldpObject map[string]json.RawMessage + if err := json.Unmarshal(lldpRaw, &lldpObject); err == nil { + if ifRaw, ok := lldpObject["interface"]; ok { + collectInterfaceNode(ifRaw) + } + } + return ifaces +} + +// transformNeighbors converts `lldpcli show neighbors` output to the +// YANG ieee802-dot1ab-lldp subtree (unwrapped -- the IPC layer adds the +// module envelope when serving the key). +func transformNeighbors(data []byte) json.RawMessage { + type remoteEntry struct { + TimeMark int `json:"time-mark"` + RemoteIndex int `json:"remote-index"` + ChassisIDSubtype string `json:"chassis-id-subtype"` + ChassisID string `json:"chassis-id"` + PortIDSubtype string `json:"port-id-subtype"` + PortID string `json:"port-id"` + } + type portEntry struct { + Name string `json:"name"` + DestMACAddress string `json:"dest-mac-address"` + RemoteSystems []remoteEntry `json:"remote-systems-data"` + } + + portMap := make(map[string]*portEntry) + seen := make(map[string]map[[2]int]bool) // per port: {time-mark, rid} + var order []string + + for _, iface := range collectIfaces(data) { + port, ok := portMap[iface.Name] + if !ok { + port = &portEntry{ + Name: iface.Name, + DestMACAddress: lldpMulticastMAC, + } + portMap[iface.Name] = port + order = append(order, iface.Name) + } + + rid := 0 + switch v := iface.RID.(type) { + case float64: + rid = int(v) + case string: + rid, _ = strconv.Atoi(v) + } + + // lldpd's rid is per remote chassis, so one chassis heard on + // two of its ports collides when the ages match too. Nudge + // time-mark to keep the list keys unique, rid stays true to + // lldpcli output. + timeMark := parseAge(iface.Age) + for seen[iface.Name][[2]int{timeMark, rid}] { + timeMark++ + } + if seen[iface.Name] == nil { + seen[iface.Name] = make(map[[2]int]bool) + } + seen[iface.Name][[2]int{timeMark, rid}] = true + + port.RemoteSystems = append(port.RemoteSystems, remoteEntry{ + TimeMark: timeMark, + RemoteIndex: rid, + ChassisIDSubtype: chassisIDSubtype(iface.Chassis.ID.Type), + ChassisID: iface.Chassis.ID.Value, + PortIDSubtype: portIDSubtype(iface.Port.ID.Type), + PortID: iface.Port.ID.Value, + }) + } + + ports := make([]portEntry, 0, len(order)) + for _, name := range order { + ports = append(ports, *portMap[name]) + } + + if len(ports) == 0 { + return json.RawMessage(`{}`) + } + + out, _ := json.Marshal(map[string]interface{}{"port": ports}) + return json.RawMessage(out) +} + +var idSubtypeMap = map[string]string{ + "ifalias": "interface-alias", + "mac": "mac-address", + "ip": "network-address", + "ifname": "interface-name", + "local": "local", +} + +func chassisIDSubtype(t string) string { + if v, ok := idSubtypeMap[t]; ok { + return v + } + return "unknown" +} + +func portIDSubtype(t string) string { + if v, ok := idSubtypeMap[t]; ok { + return v + } + return "unknown" +} + +var ageRe = regexp.MustCompile(`(\d+)\s*day[s]*,\s*(\d+):(\d+):(\d+)`) + +func parseAge(s string) int { + m := ageRe.FindStringSubmatch(s) + if m == nil { + return 0 + } + days, _ := strconv.Atoi(m[1]) + hours, _ := strconv.Atoi(m[2]) + mins, _ := strconv.Atoi(m[3]) + secs, _ := strconv.Atoi(m[4]) + return days*86400 + hours*3600 + mins*60 + secs +} + +func keysOf(m map[string]json.RawMessage) []string { + keys := make([]string, 0, len(m)) + for k := range m { + keys = append(keys, k) + } + return keys +} diff --git a/src/yangerd/internal/lldpmonitor/lldpmonitor_test.go b/src/yangerd/internal/lldpmonitor/lldpmonitor_test.go new file mode 100644 index 000000000..7608ee72e --- /dev/null +++ b/src/yangerd/internal/lldpmonitor/lldpmonitor_test.go @@ -0,0 +1,337 @@ +package lldpmonitor + +import ( + "context" + "encoding/json" + "errors" + "reflect" + "testing" + + "github.com/kernelkit/infix/src/yangerd/internal/tree" +) + +type remote struct { + TimeMark int `json:"time-mark"` + RemoteIndex int `json:"remote-index"` + ChassisIDSubtype string `json:"chassis-id-subtype"` + ChassisID string `json:"chassis-id"` + PortIDSubtype string `json:"port-id-subtype"` + PortID string `json:"port-id"` +} +type port struct { + Name string `json:"name"` + DestMAC string `json:"dest-mac-address"` + RemoteSystems []remote `json:"remote-systems-data"` +} + +// outShape is the stored (unwrapped) subtree: the IPC layer adds the +// module envelope when serving the key. +type outShape struct { + Port []port `json:"port"` +} + +// json0 format: "lldp" is an array, interface entries carry a "name" +// field, rid/age are string attributes, chassis/port/id are arrays. +const showNeighborsJSON0 = `{ + "lldp": [{ + "interface": [ + { + "name": "eth0", + "via": "LLDP", + "rid": "7", + "age": "0 day, 00:05:30", + "chassis": [{ + "id": [{"type": "mac", "value": "aa:bb:cc:dd:ee:ff"}], + "name": [{"value": "switch1"}] + }], + "port": [{ + "id": [{"type": "ifname", "value": "swp1"}] + }] + }, + { + "name": "eth1", + "via": "LLDP", + "rid": "9", + "age": "1 day, 02:30:15", + "chassis": [{ + "id": [{"type": "local", "value": "Chassis ID 007"}] + }], + "port": [{ + "id": [{"type": "mac", "value": "02:01:02:03:04:05"}] + }] + } + ] + }] +}` + +// Older keyed json format: "lldp" is an object, interfaces are keyed by +// name, chassis/port are objects. +const showNeighborsJSON = `{ + "lldp": { + "interface": [ + { + "eth0": { + "rid": 7, + "age": "0 day, 00:05:30", + "chassis": {"id": {"type": "mac", "value": "aa:bb:cc:dd:ee:ff"}}, + "port": {"id": {"type": "ifname", "value": "swp1"}} + } + } + ] + } +}` + +func decode(t *testing.T, raw json.RawMessage) outShape { + t.Helper() + var out outShape + if err := json.Unmarshal(raw, &out); err != nil { + t.Fatalf("unmarshal output: %v", err) + } + return out +} + +func TestTransformNeighborsJSON0(t *testing.T) { + out := decode(t, transformNeighbors([]byte(showNeighborsJSON0))) + + if len(out.Port) != 2 { + t.Fatalf("port count = %d, want 2", len(out.Port)) + } + + byIf := make(map[string]port) + for _, p := range out.Port { + if p.DestMAC != lldpMulticastMAC { + t.Fatalf("dest-mac-address = %q, want %q", p.DestMAC, lldpMulticastMAC) + } + byIf[p.Name] = p + } + + eth0, ok := byIf["eth0"] + if !ok || len(eth0.RemoteSystems) != 1 { + t.Fatalf("eth0 missing or wrong neighbor count: %#v", byIf) + } + want := remote{ + TimeMark: 330, + RemoteIndex: 7, + ChassisIDSubtype: "mac-address", + ChassisID: "aa:bb:cc:dd:ee:ff", + PortIDSubtype: "interface-name", + PortID: "swp1", + } + if !reflect.DeepEqual(eth0.RemoteSystems[0], want) { + t.Fatalf("eth0 remote\n got: %#v\nwant: %#v", eth0.RemoteSystems[0], want) + } + + eth1 := byIf["eth1"] + if len(eth1.RemoteSystems) != 1 { + t.Fatalf("eth1 neighbor count = %d", len(eth1.RemoteSystems)) + } + r := eth1.RemoteSystems[0] + if r.ChassisIDSubtype != "local" || r.ChassisID != "Chassis ID 007" { + t.Errorf("eth1 chassis = %s/%s", r.ChassisIDSubtype, r.ChassisID) + } + if r.PortIDSubtype != "mac-address" || r.PortID != "02:01:02:03:04:05" { + t.Errorf("eth1 port = %s/%s", r.PortIDSubtype, r.PortID) + } + if r.RemoteIndex != 9 || r.TimeMark != 95415 { + t.Errorf("eth1 rid/age = %d/%d", r.RemoteIndex, r.TimeMark) + } +} + +func TestTransformNeighborsKeyedJSON(t *testing.T) { + out := decode(t, transformNeighbors([]byte(showNeighborsJSON))) + + if len(out.Port) != 1 { + t.Fatalf("port count = %d, want 1", len(out.Port)) + } + p := out.Port[0] + if p.Name != "eth0" || len(p.RemoteSystems) != 1 { + t.Fatalf("unexpected port: %#v", p) + } + r := p.RemoteSystems[0] + if r.ChassisID != "aa:bb:cc:dd:ee:ff" || r.PortID != "swp1" || r.RemoteIndex != 7 { + t.Fatalf("unexpected remote: %#v", r) + } +} + +func TestTransformNeighborsEmpty(t *testing.T) { + for name, in := range map[string]string{ + "empty table json0": `{"lldp": [{}]}`, + "empty object": `{}`, + "malformed": `{not-json`, + } { + raw := transformNeighbors([]byte(in)) + if string(raw) != "{}" { + t.Errorf("%s: got %s, want {}", name, raw) + } + } +} + +// A neighbor that disappears between reads must vanish from the tree: +// every update is a full-table replace. +func TestUpdateTreeClearsRemovedNeighbors(t *testing.T) { + tr := tree.New() + m := New(tr, nil) + + m.query = func(context.Context) ([]byte, error) { + return []byte(showNeighborsJSON0), nil + } + m.updateTree(context.Background()) + if out := decode(t, tr.Get(treeKey)); len(out.Port) != 2 { + t.Fatalf("expected 2 ports after first read, got %d", len(out.Port)) + } + + m.query = func(context.Context) ([]byte, error) { + return []byte(`{"lldp": [{}]}`), nil + } + m.updateTree(context.Background()) + if out := decode(t, tr.Get(treeKey)); len(out.Port) != 0 { + t.Fatalf("stale neighbors not cleared: %d ports remain", len(out.Port)) + } +} + +// A failing query must leave the previous data untouched. +func TestUpdateTreeQueryErrorKeepsData(t *testing.T) { + tr := tree.New() + m := New(tr, nil) + + m.query = func(context.Context) ([]byte, error) { + return []byte(showNeighborsJSON0), nil + } + m.updateTree(context.Background()) + before := string(tr.Get(treeKey)) + + m.query = func(context.Context) ([]byte, error) { + return nil, errors.New("lldpcli gone") + } + m.updateTree(context.Background()) + + if after := string(tr.Get(treeKey)); after != before { + t.Fatal("query error overwrote existing lldp data") + } +} + +// Watch events are triggers only: added/updated/deleted all request a +// refresh, unknown events do not. +func TestProcessEventTriggers(t *testing.T) { + m := New(tree.New(), nil) + + drain := func() { + select { + case <-m.refresh: + default: + } + } + + for _, ev := range []string{"lldp-added", "lldp-updated", "lldp-deleted"} { + drain() + m.processEvent([]byte(`{"` + ev + `": {"lldp": {}}}`)) + select { + case <-m.refresh: + default: + t.Errorf("%s did not trigger refresh", ev) + } + } + + drain() + m.processEvent([]byte(`{"lldp-unknown": {}}`)) + select { + case <-m.refresh: + t.Error("unknown event triggered refresh") + default: + } +} + +func TestParseAge(t *testing.T) { + tests := []struct { + name string + in string + want int + }{ + {name: "zero day", in: "0 day, 00:05:30", want: 330}, + {name: "one day", in: "1 day, 02:30:15", want: 95415}, + {name: "ten days plural", in: "10 days, 00:00:00", want: 864000}, + {name: "empty", in: "", want: 0}, + {name: "invalid", in: "n/a", want: 0}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := parseAge(tt.in); got != tt.want { + t.Fatalf("parseAge(%q) = %d, want %d", tt.in, got, tt.want) + } + }) + } +} + +func TestSubtypeMappings(t *testing.T) { + tests := []struct { + name string + in string + want string + }{ + {name: "ifalias", in: "ifalias", want: "interface-alias"}, + {name: "mac", in: "mac", want: "mac-address"}, + {name: "ip", in: "ip", want: "network-address"}, + {name: "ifname", in: "ifname", want: "interface-name"}, + {name: "local", in: "local", want: "local"}, + {name: "unknown", in: "foo", want: "unknown"}, + } + + for _, tt := range tests { + t.Run("chassis_"+tt.name, func(t *testing.T) { + if got := chassisIDSubtype(tt.in); got != tt.want { + t.Fatalf("chassisIDSubtype(%q) = %q, want %q", tt.in, got, tt.want) + } + }) + t.Run("port_"+tt.name, func(t *testing.T) { + if got := portIDSubtype(tt.in); got != tt.want { + t.Fatalf("portIDSubtype(%q) = %q, want %q", tt.in, got, tt.want) + } + }) + } +} + +// One chassis on two of its ports shares rid, and the same age would +// make the (time-mark, remote-index) key collide. +func TestTransformNeighborsUniqueKeys(t *testing.T) { + const twoPorts = `{ + "lldp": [{ + "interface": [ + { + "name": "eth0", "rid": "7", "age": "0 day, 00:05:30", + "chassis": [{"id": [{"type": "mac", "value": "aa:bb:cc:dd:ee:ff"}]}], + "port": [{"id": [{"type": "ifname", "value": "swp1"}]}] + }, + { + "name": "eth0", "rid": "7", "age": "0 day, 00:05:30", + "chassis": [{"id": [{"type": "mac", "value": "aa:bb:cc:dd:ee:ff"}]}], + "port": [{"id": [{"type": "ifname", "value": "swp2"}]}] + }, + { + "name": "eth1", "rid": "7", "age": "0 day, 00:05:30", + "chassis": [{"id": [{"type": "mac", "value": "aa:bb:cc:dd:ee:ff"}]}], + "port": [{"id": [{"type": "ifname", "value": "swp3"}]}] + } + ] + }] +}` + out := decode(t, transformNeighbors([]byte(twoPorts))) + + byIf := make(map[string]port) + for _, p := range out.Port { + byIf[p.Name] = p + } + eth0 := byIf["eth0"].RemoteSystems + if len(eth0) != 2 { + t.Fatalf("eth0 neighbor count = %d, want 2", len(eth0)) + } + if eth0[0].TimeMark != 330 || eth0[1].TimeMark != 331 { + t.Errorf("eth0 time-marks = %d/%d, want 330/331", eth0[0].TimeMark, eth0[1].TimeMark) + } + if eth0[0].RemoteIndex != 7 || eth0[1].RemoteIndex != 7 { + t.Errorf("rid must stay 7, got %d/%d", eth0[0].RemoteIndex, eth0[1].RemoteIndex) + } + if eth1 := byIf["eth1"].RemoteSystems; len(eth1) != 1 || eth1[0].TimeMark != 330 { + t.Errorf("eth1 must not be nudged: %#v", eth1) + } +} diff --git a/src/yangerd/internal/monitor/monitor.go b/src/yangerd/internal/monitor/monitor.go new file mode 100644 index 000000000..8c636e91b --- /dev/null +++ b/src/yangerd/internal/monitor/monitor.go @@ -0,0 +1,1282 @@ +package monitor + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "log/slog" + "net" + "sort" + "strconv" + "strings" + "sync" + "syscall" + "time" + + "github.com/kernelkit/infix/src/yangerd/internal/iface" + "github.com/kernelkit/infix/src/yangerd/internal/ipbatch" + "github.com/kernelkit/infix/src/yangerd/internal/stpquery" + "github.com/kernelkit/infix/src/yangerd/internal/tree" + "github.com/vishvananda/netlink" + "github.com/vishvananda/netlink/nl" +) + +// treeKey is the single YANG module key where the complete +// ietf-interfaces document is stored. +const treeKey = "ietf-interfaces:interfaces" + +// NLMonitor subscribes to netlink link/address/neighbor events and +// keeps interface operational data in the in-memory tree up to date. +// It is the central coordinator for all interface data — raw ip-json +// staging data is transformed via iface.Transform() and augmented +// with ethernet/wifi/bridge data before being stored as a single +// complete YANG document. +// +// Staged link and address rows are keyed by ifindex, which survives a +// rename; everything keyed by name is dropped when the name changes. +type NLMonitor struct { + linkBatch *ipbatch.Batch + addrBatch *ipbatch.Batch + neighBatch *ipbatch.Batch + brBatch *ipbatch.Batch + tree *tree.Tree + ethRefresh func(string) + linkSet func() + log *slog.Logger + fc iface.FileChecker + + // initDone is closed after the first initialDump completes. + // initDoneOnce ensures it is only closed once across restarts. + initDone chan struct{} + initDoneOnce sync.Once + + // redumpCh asks the event loop to run initialDump, e.g. after a + // batch subprocess died. Buffer of 1 so requests coalesce. + redumpCh chan struct{} + + // staging holds raw ip-json data used as input to iface.Transform(). + // Protected by mu. + mu sync.Mutex + links json.RawMessage // ip -json -s -d link show (includes stats+details) + addrs json.RawMessage // ip -json -d addr show (details only, no stats) + neighs json.RawMessage // ip -json neigh show + fdb map[string]json.RawMessage + mdb map[string]json.RawMessage + ethernet map[string]json.RawMessage // ifname → ethtool JSON + wifi map[string]json.RawMessage // ifname → wifi JSON + wireguard map[string]json.RawMessage // ifname → WireGuard peer-status JSON + lastWG string + + lastOperStatus map[string]string + // lastChange stamps the operstate transitions seen since start. + // Kept as time.Time for the monotonic reading: the wall clock is + // derived at GET time, so a later NTP step corrects old stamps. + lastChange map[string]time.Time + + // stpMu guards lastSTP, a fingerprint of the most recent mstpd STP + // query, letting the periodic poll rebuild only on actual change. + stpMu sync.Mutex + lastSTP string +} + +// coalesceDelay is how long the event loop collects further netlink +// events before rebuilding the document once for all of them. +const coalesceDelay = 20 * time.Millisecond + +const ( + // redumpSettle lets a burst of failures and the batch restart that + // caused them settle before one re-dump, redumpRetry follows a + // failed one. + redumpSettle = 200 * time.Millisecond + redumpRetry = 2 * time.Second +) + +// New creates a netlink monitor backed by ip/bridge batch query workers. +// linkBatch should include -s -d flags; addrBatch should include -d only +// (no -s, which causes multi-line output for link commands). +func New(linkBatch, addrBatch, neighBatch, brBatch *ipbatch.Batch, t *tree.Tree, fc iface.FileChecker, log *slog.Logger) *NLMonitor { + return &NLMonitor{ + linkBatch: linkBatch, + addrBatch: addrBatch, + neighBatch: neighBatch, + brBatch: brBatch, + tree: t, + fc: fc, + log: log, + initDone: make(chan struct{}), + redumpCh: make(chan struct{}, 1), + fdb: make(map[string]json.RawMessage), + mdb: make(map[string]json.RawMessage), + ethernet: make(map[string]json.RawMessage), + wifi: make(map[string]json.RawMessage), + wireguard: make(map[string]json.RawMessage), + lastOperStatus: make(map[string]string), + lastChange: make(map[string]time.Time), + } +} + +// SetEthRefresh sets an optional callback used to refresh ethtool data +// when an ethernet interface sees a link event. +func (m *NLMonitor) SetEthRefresh(fn func(string)) { + m.ethRefresh = fn +} + +// SetLinkSetChange sets an optional callback run when an interface +// appears, goes away or is renamed. +func (m *NLMonitor) SetLinkSetChange(fn func()) { + m.linkSet = fn +} + +func (m *NLMonitor) linkSetChanged() { + if m.linkSet != nil { + m.linkSet() + } +} + +// WaitReady returns a channel that is closed after initialDump completes. +func (m *NLMonitor) WaitReady() <-chan struct{} { + return m.initDone +} + +// SetEthernetData updates the staged ethernet data for an interface +// and triggers a full rebuild of the YANG document. +func (m *NLMonitor) SetEthernetData(ifname string, data json.RawMessage) { + m.mu.Lock() + m.ethernet[ifname] = data + m.mu.Unlock() + m.rebuild() +} + +// SetWifiData updates the staged wifi data for an interface +// and triggers a full rebuild of the YANG document. +func (m *NLMonitor) SetWifiData(ifname string, data json.RawMessage) { + m.mu.Lock() + m.wifi[ifname] = data + m.mu.Unlock() + m.rebuild() +} + +// SetWireguardAll replaces the WireGuard peer-status data of every +// interface, so a tunnel that lost its last peer loses its status too. +// The document is rebuilt only when something changed. +func (m *NLMonitor) SetWireguardAll(data map[string]json.RawMessage) { + var b strings.Builder + writeSortedRaw(&b, "w", data) + fp := b.String() + + m.mu.Lock() + changed := fp != m.lastWG + m.lastWG = fp + m.wireguard = copyStringMap(data) + if m.wireguard == nil { + m.wireguard = make(map[string]json.RawMessage) + } + m.mu.Unlock() + + if changed { + m.rebuild() + } +} + +// Links returns a copy of the current staged links data. +func (m *NLMonitor) Links() json.RawMessage { + m.mu.Lock() + cp := append(json.RawMessage{}, m.links...) + m.mu.Unlock() + return cp +} + +// Run starts the netlink monitor loop and returns on context +// cancellation, channel closure, or subscription errors. The netlink +// library ends a subscription on any receive error, ENOBUFS included, +// so every error means resubscribe and re-dump: return and let the +// caller restart us. +func (m *NLMonitor) Run(ctx context.Context) error { + runCtx, cancel := context.WithCancel(ctx) + defer cancel() + + done := make(chan struct{}) + defer close(done) + + errorCallback := func(err error) { + if err == nil { + return + } + if strings.Contains(err.Error(), syscall.ENOBUFS.Error()) { + m.log.Warn("netlink events dropped, resubscribing", "err", err) + } else { + m.log.Error("netlink subscription error", "err", err) + } + cancel() + } + + linkCh := make(chan netlink.LinkUpdate, 64) + addrCh := make(chan netlink.AddrUpdate, 64) + neighCh := make(chan netlink.NeighUpdate, 64) + mdbCh := make(chan struct{}, 32) + + if err := netlink.LinkSubscribeWithOptions(linkCh, done, netlink.LinkSubscribeOptions{ + ErrorCallback: errorCallback, + ReceiveBufferSize: 32 * 1024 * 1024, + ReceiveBufferForceSize: true, + }); err != nil { + return fmt.Errorf("subscribe link updates: %w", err) + } + if err := netlink.AddrSubscribeWithOptions(addrCh, done, netlink.AddrSubscribeOptions{ + ErrorCallback: errorCallback, + ReceiveBufferSize: 32 * 1024 * 1024, + ReceiveBufferForceSize: true, + }); err != nil { + return fmt.Errorf("subscribe addr updates: %w", err) + } + if err := netlink.NeighSubscribeWithOptions(neighCh, done, netlink.NeighSubscribeOptions{ + ErrorCallback: errorCallback, + ReceiveBufferSize: 32 * 1024 * 1024, + ReceiveBufferForceSize: true, + }); err != nil { + return fmt.Errorf("subscribe neigh updates: %w", err) + } + if err := m.subscribeBridgeMDB(runCtx, mdbCh, errorCallback); err != nil { + return fmt.Errorf("subscribe bridge mdb updates: %w", err) + } + + if err := m.initialDump(); err != nil { + m.log.Error("initial dump failed", "err", err) + } + m.initDoneOnce.Do(func() { close(m.initDone) }) + + var flush, redump <-chan time.Time + for { + changed := false + select { + case <-runCtx.Done(): + if ctx.Err() != nil { + return ctx.Err() + } + return runCtx.Err() + case lu, ok := <-linkCh: + if !ok { + return fmt.Errorf("link update channel closed") + } + changed = m.handleLinkUpdate(lu) + case au, ok := <-addrCh: + if !ok { + return fmt.Errorf("addr update channel closed") + } + changed = m.handleAddrUpdate(au) + case nu, ok := <-neighCh: + if !ok { + return fmt.Errorf("neigh update channel closed") + } + changed = m.handleNeighUpdate(nu) + case _, ok := <-mdbCh: + if !ok { + return fmt.Errorf("bridge mdb update channel closed") + } + changed = m.handleMDBUpdate() + case <-flush: + flush = nil + m.rebuild() + case <-m.redumpCh: + // Every event that hits a dead batch asks for this, so + // requests while one is pending are the same request. + if redump == nil { + redump = time.After(redumpSettle) + } + case <-redump: + redump = nil + if !m.batchesAlive() { + redump = time.After(redumpSettle) + break + } + m.log.Warn("re-dumping all interfaces") + if err := m.initialDump(); err != nil { + m.log.Error("re-dump failed, will retry", "err", err, "in", redumpRetry) + redump = time.After(redumpRetry) + } + } + if changed && flush == nil { + flush = time.After(coalesceDelay) + } + } +} + +func (m *NLMonitor) initialDump() error { + linkRaw, err := m.query(m.linkBatch, "link show") + if err != nil { + return err + } + addrRaw, err := m.query(m.addrBatch, "addr show") + if err != nil { + return err + } + neighRaw, err := m.query(m.neighBatch, "neigh show") + if err != nil { + neighRaw = json.RawMessage(`[]`) + } + mdbRaw, err := m.query(m.brBatch, "mdb show") + if err != nil { + mdbRaw = json.RawMessage(`[]`) + } + + m.log.Debug("initialDump", "linkBytes", len(linkRaw), "addrBytes", len(addrRaw), "neighBytes", len(neighRaw)) + m.validateAddrData("initialDump", addrRaw) + + m.mu.Lock() + m.links = linkRaw + m.addrs = addrRaw + m.neighs = neighRaw + m.mdb = mdbByBridge(mdbRaw) + for _, row := range decodeRows(linkRaw) { + if name := rowString(row, "ifname"); name != "" { + if st := rowString(row, "operstate"); st != "" { + if prev, had := m.lastOperStatus[name]; had && prev != st { + m.lastChange[name] = time.Now() + } + m.lastOperStatus[name] = st + } + } + } + m.mu.Unlock() + + m.rebuild() + + // Ports that are up when yangerd starts send no link event, so ask + // for their ethtool data here or they stay without speed and duplex. + if m.ethRefresh != nil { + for _, name := range ethernetNames(linkRaw, m.fc) { + m.ethRefresh(name) + } + } + return nil +} + +// ethernetNames lists the interfaces in an `ip -json link` dump that +// carry ethtool data. +func ethernetNames(linkRaw json.RawMessage, fc iface.FileChecker) []string { + var rows []json.RawMessage + if json.Unmarshal(linkRaw, &rows) != nil { + return nil + } + + var names []string + for _, row := range rows { + one := append(append(json.RawMessage{'['}, row...), ']') + if !iface.IsEthernet(one, fc) { + continue + } + var link struct { + Name string `json:"ifname"` + } + if json.Unmarshal(row, &link) == nil && link.Name != "" { + names = append(names, link.Name) + } + } + return names +} + +func (m *NLMonitor) handleLinkUpdate(update netlink.LinkUpdate) bool { + index := int(update.Index) + if update.Header.Type == syscall.RTM_DELLINK { + defer m.linkSetChanged() + return m.removeInterface(index) + } + + name, ok := linkNameFromUpdate(update) + if !ok || name == "" { + m.log.Warn("link update without interface name", "index", index) + return false + } + return m.refreshInterface(index, name) +} + +func (m *NLMonitor) handleAddrUpdate(update netlink.AddrUpdate) bool { + ifname, err := ifNameByIndex(update.LinkIndex) + if err != nil { + m.log.Debug("addr update: interface gone", "index", update.LinkIndex, "err", err) + return false + } + + raw, err := m.query(m.addrBatch, "addr show dev "+devRef(update.LinkIndex)) + if err != nil { + return false + } + if !m.validateAddrData("handleAddrUpdate/"+ifname, raw) { + m.log.Error("handleAddrUpdate: REFUSING to store invalid addr data", "ifname", ifname) + return false + } + + m.mu.Lock() + m.addrs = replaceRows(m.addrs, "ifindex", update.LinkIndex, raw) + m.mu.Unlock() + return true +} + +func (m *NLMonitor) handleNeighUpdate(update netlink.NeighUpdate) bool { + if isBridgeFDB(update) { + bridgeName, bridgeIndex, ok := bridgeNameFromNeigh(update) + if !ok { + m.log.Warn("fdb update: bridge name not found", "link-index", update.LinkIndex) + return false + } + + raw, err := m.query(m.brBatch, "fdb show br "+devRef(bridgeIndex)) + if err != nil { + return false + } + + m.mu.Lock() + m.fdb[bridgeName] = raw + m.mu.Unlock() + return true + } + + ifname, err := ifNameByIndex(update.LinkIndex) + if err != nil { + raw, err := m.query(m.neighBatch, "neigh show") + if err != nil { + return false + } + m.mu.Lock() + m.neighs = raw + m.mu.Unlock() + return true + } + + raw, err := m.query(m.neighBatch, "neigh show dev "+devRef(update.LinkIndex)) + if err != nil { + return false + } + rows := withField(raw, "dev", ifname) + + m.mu.Lock() + m.neighs = replaceRows(m.neighs, "dev", ifname, rows) + m.mu.Unlock() + return true +} + +func (m *NLMonitor) handleMDBUpdate() bool { + raw, err := m.query(m.brBatch, "mdb show") + if err != nil { + return false + } + + m.mu.Lock() + m.mdb = mdbByBridge(raw) + m.mu.Unlock() + return true +} + +// removeInterface purges all staged data for the interface at index. +func (m *NLMonitor) removeInterface(index int) bool { + m.mu.Lock() + defer m.mu.Unlock() + + name := nameByIndex(m.links, index) + m.links = replaceRows(m.links, "ifindex", index, nil) + m.addrs = replaceRows(m.addrs, "ifindex", index, nil) + // Link and address events arrive on separate channels, so the delete + // of an old interface can be handled after a new one took its name. + if name != "" && !nameInUse(m.links, name) { + m.forgetName(name) + } + m.log.Debug("removeInterface", "index", index, "ifname", name) + return true +} + +// forgetName drops everything staged under an interface name. Caller +// holds m.mu. +func (m *NLMonitor) forgetName(name string) { + m.neighs = replaceRows(m.neighs, "dev", name, nil) + delete(m.fdb, name) + delete(m.mdb, name) + delete(m.ethernet, name) + delete(m.wifi, name) + delete(m.wireguard, name) + delete(m.lastOperStatus, name) + delete(m.lastChange, name) +} + +// batchesAlive tells whether every batch subprocess is up, so a re-dump +// can succeed. +func (m *NLMonitor) batchesAlive() bool { + for _, b := range []*ipbatch.Batch{m.linkBatch, m.addrBatch, m.neighBatch, m.brBatch} { + if b != nil && !b.Alive() { + return false + } + } + return true +} + +func (m *NLMonitor) requestRedump() { + select { + case m.redumpCh <- struct{}{}: + default: + } +} + +func (m *NLMonitor) refreshInterface(index int, name string) bool { + linkRaw, err := m.query(m.linkBatch, "link show dev "+name) + if errors.Is(err, ipbatch.ErrCommandFailed) { + if _, gone := net.InterfaceByIndex(index); gone != nil { + return m.removeInterface(index) // gone before we asked + } + // Still there: ip resolved the name from its cache, to an + // interface that had it before. See devRef. + m.log.Debug("link query failed for a live interface, refreshing ip", "ifname", name) + if rerr := m.linkBatch.Refresh(); rerr != nil { + m.log.Warn("refreshing ip batch failed", "err", rerr) + return false + } + linkRaw, err = m.query(m.linkBatch, "link show dev "+name) + } + if err != nil { + return false + } + if !linkRowFor(linkRaw, index) { + // ip prints the object before it checks the message, so a + // link racing a delete or rename can come back as [{}]. + m.log.Warn("link query returned no row for the interface, re-dumping", "ifname", name, "ifindex", index) + m.requestRedump() + return false + } + + addrRaw, err := m.query(m.addrBatch, "addr show dev "+devRef(index)) + if err != nil { + addrRaw = nil + } + if addrRaw != nil && !m.validateAddrData("refreshInterface/"+name, addrRaw) { + m.log.Error("refreshInterface: REFUSING to store invalid addr data", "ifname", name) + addrRaw = nil + } + + m.mu.Lock() + old := nameByIndex(m.links, index) + if old != "" && old != name { + m.log.Info("interface renamed", "from", old, "to", name) + if !nameInUse(replaceRows(m.links, "ifindex", index, nil), old) { + m.forgetName(old) + } + } + m.updateOperStatus(name, linkRaw) + m.links = replaceRows(m.links, "ifindex", index, linkRaw) + if addrRaw != nil { + m.addrs = replaceRows(m.addrs, "ifindex", index, addrRaw) + } + m.mu.Unlock() + + if old != name { + m.linkSetChanged() + } + if m.ethRefresh != nil && iface.IsEthernet(linkRaw, m.fc) { + m.ethRefresh(name) + } + return true +} + +// rebuild runs iface.Transform on all staged data, merges augments +// (ethernet, wifi, bridge fdb/mdb), and stores the result. +// Caller must NOT hold m.mu. +func (m *NLMonitor) rebuild() { + m.mu.Lock() + linksCopy := append(json.RawMessage{}, m.links...) + addrsCopy := append(json.RawMessage{}, m.addrs...) + neighsCopy := append(json.RawMessage{}, m.neighs...) + doc := iface.Transform(linksCopy, addrsCopy, neighsCopy, m.fc) + eth := copyStringMap(m.ethernet) + wfi := copyStringMap(m.wifi) + fdb := copyStringMap(m.fdb) + mdb := copyStringMap(m.mdb) + wg := copyStringMap(m.wireguard) + m.mu.Unlock() + + var brSTP, ptSTP map[string]json.RawMessage + resolver := stpquery.NewLinksIfIndexResolver(linksCopy) + brSTP, ptSTP = stpquery.Query(linksCopy, resolver) + + doc = mergeAugments(doc, eth, wfi, fdb, mdb, brSTP, ptSTP, wg) + m.tree.Set(treeKey, doc) +} + +// RefreshSTP re-queries mstpd and rebuilds only when STP data changed. +// mstpd's control socket is request/response with no event channel, and +// the bridge-level root-id settles via BPDU exchange without any netlink +// event, so STP state must be polled to stay current. An empty result +// is a result too: mstpd going away must drop the stale STP data, and +// its return must bring it back. +func (m *NLMonitor) RefreshSTP() { + links := m.Links() + resolver := stpquery.NewLinksIfIndexResolver(links) + brSTP, ptSTP := stpquery.Query(links, resolver) + + fp := stpFingerprint(brSTP, ptSTP) + m.stpMu.Lock() + changed := fp != m.lastSTP + m.lastSTP = fp + m.stpMu.Unlock() + + if changed { + m.rebuild() + } +} + +func stpFingerprint(brSTP, ptSTP map[string]json.RawMessage) string { + var b strings.Builder + writeSortedRaw(&b, "b", brSTP) + writeSortedRaw(&b, "p", ptSTP) + return b.String() +} + +func writeSortedRaw(b *strings.Builder, prefix string, m map[string]json.RawMessage) { + keys := make([]string, 0, len(m)) + for k := range m { + keys = append(keys, k) + } + sort.Strings(keys) + for _, k := range keys { + b.WriteString(prefix) + b.WriteByte(':') + b.WriteString(k) + b.WriteByte('=') + b.Write(m[k]) + b.WriteByte('\n') + } +} + +// mergeAugments adds ethernet, wifi, and bridge data into the +// complete ietf-interfaces document produced by iface.Transform(). +func mergeAugments(doc json.RawMessage, ethernet, wifi, fdb, mdb, bridgeSTP, portSTP, wireguard map[string]json.RawMessage) json.RawMessage { + if len(ethernet) == 0 && len(wifi) == 0 && len(fdb) == 0 && len(mdb) == 0 && len(bridgeSTP) == 0 && len(portSTP) == 0 && len(wireguard) == 0 { + return doc + } + + var root map[string]any + if err := json.Unmarshal(doc, &root); err != nil { + return doc + } + + ifaceList, ok := root["interface"] + if !ok { + return doc + } + ifaceArr, ok := ifaceList.([]any) + if !ok { + return doc + } + + for i, entry := range ifaceArr { + ifaceObj, ok := entry.(map[string]any) + if !ok { + continue + } + name, _ := ifaceObj["name"].(string) + if name == "" { + continue + } + + if ethData, ok := ethernet[name]; ok { + var wrapper map[string]json.RawMessage + if err := json.Unmarshal(ethData, &wrapper); err == nil { + if ethRaw, ok := wrapper["ethernet"]; ok { + var ethObj any + if err := json.Unmarshal(ethRaw, ðObj); err == nil { + ifaceObj["ieee802-ethernet-interface:ethernet"] = ethObj + } + } + if speedRaw, ok := wrapper["speed"]; ok { + var speed string + if err := json.Unmarshal(speedRaw, &speed); err == nil { + ifaceObj["speed"] = speed + } + } + } + } + + if wifiData, ok := wifi[name]; ok { + var wifiObj any + if err := json.Unmarshal(wifiData, &wifiObj); err == nil { + ifaceObj["infix-interfaces:wifi"] = wifiObj + } + } + + if fdbData, ok := fdb[name]; ok { + bridgeObj := ensureBridgeAugment(ifaceObj) + var fdbObj any + if err := json.Unmarshal(fdbData, &fdbObj); err == nil { + bridgeObj["fdb"] = fdbObj + } + } + + if mdbData, ok := mdb[name]; ok { + bridgeObj := ensureBridgeAugment(ifaceObj) + if mf := transformMDB(mdbData); mf != nil { + bridgeObj["multicast-filters"] = mf + } + } + + if stpData, ok := bridgeSTP[name]; ok { + bridgeObj := ensureBridgeAugment(ifaceObj) + var stpObj any + if err := json.Unmarshal(stpData, &stpObj); err == nil { + bridgeObj["stp"] = stpObj + } + } + + if stpData, ok := portSTP[name]; ok { + bpObj := ensureBridgePortAugment(ifaceObj) + var stpObj any + if err := json.Unmarshal(stpData, &stpObj); err == nil { + deepMergeSTP(bpObj, stpObj) + } + } + + if wgData, ok := wireguard[name]; ok { + var wgObj any + if err := json.Unmarshal(wgData, &wgObj); err == nil { + ifaceObj["infix-interfaces:wireguard"] = wgObj + } + } + + ifaceArr[i] = ifaceObj + } + + out, err := json.Marshal(root) + if err != nil { + return doc + } + return json.RawMessage(out) +} + +// ensureBridgeAugment returns the bridge augment object within an +// interface, creating it if necessary. +func ensureBridgeAugment(ifaceObj map[string]any) map[string]any { + key := "infix-interfaces:bridge" + if existing, ok := ifaceObj[key]; ok { + if m, ok := existing.(map[string]any); ok { + return m + } + } + bridgeObj := map[string]any{} + ifaceObj[key] = bridgeObj + return bridgeObj +} + +func ensureBridgePortAugment(ifaceObj map[string]any) map[string]any { + key := "infix-interfaces:bridge-port" + if existing, ok := ifaceObj[key]; ok { + if m, ok := existing.(map[string]any); ok { + return m + } + } + obj := map[string]any{} + ifaceObj[key] = obj + return obj +} + +// deepMergeSTP merges mstpd STP data into the bridge-port augment. +// The kernel already provides stp.cist.state via iface.Transform; +// mstpd adds role, port-id, designated, etc. We deep-merge to +// preserve the kernel state field while adding mstpd fields. +func deepMergeSTP(bpObj map[string]any, stpData any) { + stpMap, ok := stpData.(map[string]any) + if !ok { + return + } + + existing, _ := bpObj["stp"].(map[string]any) + if existing == nil { + bpObj["stp"] = stpMap + return + } + + if newCist, ok := stpMap["cist"].(map[string]any); ok { + existingCist, _ := existing["cist"].(map[string]any) + if existingCist == nil { + existing["cist"] = newCist + } else { + for k, v := range newCist { + existingCist[k] = v + } + } + } + + for k, v := range stpMap { + if k != "cist" { + existing[k] = v + } + } +} + +func copyStringMap(m map[string]json.RawMessage) map[string]json.RawMessage { + if len(m) == 0 { + return nil + } + cp := make(map[string]json.RawMessage, len(m)) + for k, v := range m { + cp[k] = v + } + return cp +} + +func (m *NLMonitor) updateOperStatus(ifname string, raw json.RawMessage) { + status, ok := extractOperStatus(raw) + if !ok { + return + } + + prev, had := m.lastOperStatus[ifname] + m.lastOperStatus[ifname] = status + if !had || prev != status { + m.lastChange[ifname] = time.Now() + } + if had && prev != status { + m.log.Info("oper-status transition", "ifname", ifname, "from", prev, "to", status) + } +} + +// LastChange is a tree provider adding last-change to every interface +// whose operstate changed since start. Interfaces up since before +// yangerd started carry no stamp, their change predates us. +func (m *NLMonitor) LastChange() json.RawMessage { + m.mu.Lock() + stamps := make(map[string]time.Time, len(m.lastChange)) + for k, v := range m.lastChange { + stamps[k] = v + } + m.mu.Unlock() + if len(stamps) == 0 { + return nil + } + + return stampLastChange(m.tree.GetCached(treeKey), stamps, time.Now()) +} + +func stampLastChange(doc json.RawMessage, stamps map[string]time.Time, now time.Time) json.RawMessage { + var root map[string]any + if err := json.Unmarshal(doc, &root); err != nil { + return nil + } + ifaceArr, ok := root["interface"].([]any) + if !ok { + return nil + } + + for _, entry := range ifaceArr { + ifaceObj, ok := entry.(map[string]any) + if !ok { + continue + } + name, _ := ifaceObj["name"].(string) + at, ok := stamps[name] + if !ok { + continue + } + // Elapsed time is monotonic, so the stamp follows the wall + // clock as it is now, not as it was when the link changed. + when := now.Add(-now.Sub(at)) + ifaceObj["last-change"] = when.UTC().Format("2006-01-02T15:04:05+00:00") + } + + out, err := json.Marshal(map[string]any{"interface": ifaceArr}) + if err != nil { + return nil + } + return out +} + +// query runs one command on a batch worker. A dead worker asks for a +// re-dump once it is back; a rejected command is left to the caller. +func (m *NLMonitor) query(b *ipbatch.Batch, command string) (json.RawMessage, error) { + raw, err := b.Query(command) + switch { + case err == nil: + return raw, nil + case errors.Is(err, ipbatch.ErrCommandFailed): + m.log.Debug("batch command failed", "command", command) + case errors.Is(err, ipbatch.ErrBatchDead): + m.log.Warn("batch dead", "command", command) + m.requestRedump() + default: + m.log.Error("batch query failed", "command", command, "err", err) + } + return nil, err +} + +func (m *NLMonitor) subscribeBridgeMDB(ctx context.Context, ch chan<- struct{}, errorCallback func(error)) error { + sock, err := nl.Subscribe(syscall.NETLINK_ROUTE, 26) + if err != nil { + return err + } + + // Receive blocks until a message arrives; closing the socket on + // cancel is what ends it. Close once, the fd number may be reused. + var once sync.Once + closeSock := func() { once.Do(sock.Close) } + stop := context.AfterFunc(ctx, closeSock) + + go func() { + defer close(ch) + defer stop() + defer closeSock() + + for { + msgs, _, err := sock.Receive() + if err != nil { + if ctx.Err() != nil { + return + } + errorCallback(err) + return + } + if len(msgs) == 0 { + continue + } + + select { + case ch <- struct{}{}: + default: + } + } + }() + + return nil +} + +func ifNameByIndex(index int) (string, error) { + iface, err := net.InterfaceByIndex(index) + if err != nil { + return "", err + } + return iface.Name, nil +} + +func linkNameFromUpdate(update netlink.LinkUpdate) (string, bool) { + if update.Link != nil && update.Link.Attrs() != nil && update.Link.Attrs().Name != "" { + return update.Link.Attrs().Name, true + } + if update.Index <= 0 { + return "", false + } + name, err := ifNameByIndex(int(update.Index)) + if err != nil { + return "", false + } + return name, true +} + +func isBridgeFDB(update netlink.NeighUpdate) bool { + if update.Family == syscall.AF_BRIDGE { + return true + } + if update.MasterIndex > 0 { + return true + } + if update.Flags&netlink.NTF_MASTER != 0 { + return true + } + return false +} + +func bridgeNameFromNeigh(update netlink.NeighUpdate) (string, int, bool) { + if update.MasterIndex > 0 { + name, err := ifNameByIndex(update.MasterIndex) + if err == nil { + return name, update.MasterIndex, true + } + } + + if update.LinkIndex <= 0 { + return "", 0, false + } + link, err := netlink.LinkByIndex(update.LinkIndex) + if err == nil && link != nil && link.Attrs() != nil && link.Attrs().MasterIndex > 0 { + name, err := ifNameByIndex(link.Attrs().MasterIndex) + if err == nil { + return name, link.Attrs().MasterIndex, true + } + } + return "", 0, false +} + +// devRef names a device by index for a batch command. iproute2 caches +// name-to-index lookups for the life of the process, so once a name is +// reused by a new interface, "dev wifi0" keeps resolving to the deleted +// one. "if" bypasses that cache. "link show" is the exception: it +// sends the name to the kernel and does not accept this form. +func devRef(index int) string { + return "if" + strconv.Itoa(index) +} + +// validateAddrData checks whether a JSON response from "addr show" contains +// addr_info entries. "ip -json addr show" always includes an "addr_info" +// array for every interface object; its absence means we got link-format +// data instead. Returns true if the data looks valid (has addr_info). +func (m *NLMonitor) validateAddrData(caller string, raw json.RawMessage) bool { + if len(raw) == 0 { + m.log.Error("addr data is EMPTY", "caller", caller) + return false + } + + var rows []map[string]json.RawMessage + if err := json.Unmarshal(raw, &rows); err != nil { + m.log.Error("addr data unmarshal failed", "caller", caller, "err", err, "raw", string(raw)) + return false + } + + if len(rows) == 0 { + // Empty array is valid — interface exists but has no addresses. + return true + } + + for _, row := range rows { + ifnRaw, _ := row["ifname"] + var ifn string + json.Unmarshal(ifnRaw, &ifn) + + if _, ok := row["addr_info"]; !ok { + m.log.Error("addr data MISSING addr_info — got link-format data", + "caller", caller, + "ifname", ifn, + "keys", mapKeys(row), + "raw", string(raw), + ) + return false + } + } + return true +} + +// mapKeys returns the JSON object keys from a map for diagnostic logging. +func mapKeys(m map[string]json.RawMessage) []string { + keys := make([]string, 0, len(m)) + for k := range m { + keys = append(keys, k) + } + return keys +} + +func extractOperStatus(raw json.RawMessage) (string, bool) { + var rows []map[string]json.RawMessage + if err := json.Unmarshal(raw, &rows); err != nil || len(rows) == 0 { + return "", false + } + stateRaw, ok := rows[0]["operstate"] + if !ok { + return "", false + } + var state string + if err := json.Unmarshal(stateRaw, &state); err != nil || state == "" { + return "", false + } + return state, true +} + +// decodeRows splits an ip/bridge -json array into raw rows. +func decodeRows(raw json.RawMessage) []map[string]json.RawMessage { + var rows []map[string]json.RawMessage + if json.Unmarshal(raw, &rows) != nil { + return nil + } + return rows +} + +func rowString(row map[string]json.RawMessage, key string) string { + var s string + if json.Unmarshal(row[key], &s) != nil { + return "" + } + return s +} + +// rowMatches compares a row field against a Go value by its JSON form, +// so an ifindex matches as a number and a dev name as a string. +func rowMatches(row map[string]json.RawMessage, key, want string) bool { + v, ok := row[key] + return ok && string(v) == want +} + +func jsonKey(v any) string { + b, _ := json.Marshal(v) + return string(b) +} + +// replaceRows drops every row of bulk whose key equals match, appends +// the rows of replacement (nil to only drop), and returns the result. +func replaceRows(bulk json.RawMessage, key string, match any, replacement json.RawMessage) json.RawMessage { + want := jsonKey(match) + var kept []json.RawMessage + if json.Unmarshal(bulk, &kept) != nil { + kept = nil + } + + out := kept[:0] + for _, raw := range kept { + var row map[string]json.RawMessage + if json.Unmarshal(raw, &row) == nil && rowMatches(row, key, want) { + continue + } + out = append(out, raw) + } + + var add []json.RawMessage + if json.Unmarshal(replacement, &add) == nil { + out = append(out, add...) + } + if out == nil { + out = []json.RawMessage{} + } + + merged, err := json.Marshal(out) + if err != nil { + return bulk + } + return merged +} + +// nameByIndex returns the staged ifname for index, "" if none. +// linkRowFor tells whether raw is a link answer for index: one row with +// that ifindex and a name. +func linkRowFor(raw json.RawMessage, index int) bool { + rows := decodeRows(raw) + return len(rows) == 1 && rowMatches(rows[0], "ifindex", jsonKey(index)) && + rowString(rows[0], "ifname") != "" +} + +// nameInUse tells whether any staged link row carries name. +func nameInUse(links json.RawMessage, name string) bool { + for _, row := range decodeRows(links) { + if rowString(row, "ifname") == name { + return true + } + } + return false +} + +func nameByIndex(links json.RawMessage, index int) string { + want := jsonKey(index) + for _, row := range decodeRows(links) { + if rowMatches(row, "ifindex", want) { + return rowString(row, "ifname") + } + } + return "" +} + +// withField sets key to value on every row of raw that lacks it. +// `ip neigh show dev X` leaves out the dev it was asked about. +func withField(raw json.RawMessage, key, value string) json.RawMessage { + rows := decodeRows(raw) + v := json.RawMessage(jsonKey(value)) + for _, row := range rows { + if _, ok := row[key]; !ok { + row[key] = v + } + } + if rows == nil { + rows = []map[string]json.RawMessage{} + } + out, err := json.Marshal(rows) + if err != nil { + return raw + } + return out +} + +// parseMDBEntries extracts the flat list of MDB entries from the +// bridge batch output format: [{"mdb":[{entries...}],"router":{}}] +func parseMDBEntries(raw json.RawMessage) []map[string]any { + var wrappers []map[string]json.RawMessage + if err := json.Unmarshal(raw, &wrappers); err != nil { + return nil + } + + var all []map[string]any + for _, w := range wrappers { + mdbRaw, ok := w["mdb"] + if !ok { + continue + } + var entries []map[string]any + if err := json.Unmarshal(mdbRaw, &entries); err != nil { + continue + } + all = append(all, entries...) + } + return all +} + +// mdbByBridge splits `bridge mdb show` into per-bridge entry lists. A +// bridge without entries is simply absent, which clears its filters. +func mdbByBridge(raw json.RawMessage) map[string]json.RawMessage { + groups := make(map[string][]map[string]any) + for _, e := range parseMDBEntries(raw) { + if dev, _ := e["dev"].(string); dev != "" { + groups[dev] = append(groups[dev], e) + } + } + out := make(map[string]json.RawMessage, len(groups)) + for dev, entries := range groups { + if b, err := json.Marshal(entries); err == nil { + out[dev] = b + } + } + return out +} + +func transformMDB(raw json.RawMessage) map[string]any { + var entries []map[string]any + if err := json.Unmarshal(raw, &entries); err != nil || len(entries) == 0 { + return nil + } + + type portEntry struct { + Port string `json:"port"` + State string `json:"state"` + } + + groups := make(map[string][]portEntry) + var order []string + + for _, e := range entries { + grp, _ := e["grp"].(string) + port, _ := e["port"].(string) + state, _ := e["state"].(string) + if grp == "" || port == "" { + continue + } + + if _, seen := groups[grp]; !seen { + order = append(order, grp) + } + groups[grp] = append(groups[grp], portEntry{ + Port: port, + State: mdbStateToYANG(state), + }) + } + + if len(groups) == 0 { + return nil + } + + filters := make([]map[string]any, 0, len(order)) + for _, grp := range order { + filters = append(filters, map[string]any{ + "group": grp, + "ports": groups[grp], + }) + } + + return map[string]any{"multicast-filter": filters} +} + +func mdbStateToYANG(state string) string { + switch state { + case "temp": + return "temporary" + case "permanent": + return "permanent" + default: + return state + } +} diff --git a/src/yangerd/internal/monitor/monitor_test.go b/src/yangerd/internal/monitor/monitor_test.go new file mode 100644 index 000000000..9ef4abfc3 --- /dev/null +++ b/src/yangerd/internal/monitor/monitor_test.go @@ -0,0 +1,417 @@ +package monitor + +import ( + "encoding/json" + "github.com/kernelkit/infix/src/yangerd/internal/tree" + "log/slog" + "os" + "reflect" + "strings" + "syscall" + "testing" + "time" + + "github.com/vishvananda/netlink" +) + +func TestExtractOperStatus(t *testing.T) { + tests := []struct { + name string + raw json.RawMessage + want string + wantOK bool + }{ + { + name: "valid single entry", + raw: json.RawMessage(`[{"operstate":"UP"}]`), + want: "UP", + wantOK: true, + }, + { + name: "multiple entries first wins", + raw: json.RawMessage(`[{"operstate":"DOWN"},{"operstate":"UP"}]`), + want: "DOWN", + wantOK: true, + }, + { + name: "missing operstate", + raw: json.RawMessage(`[{"ifname":"eth0"}]`), + want: "", + wantOK: false, + }, + { + name: "empty array", + raw: json.RawMessage(`[]`), + want: "", + wantOK: false, + }, + { + name: "invalid json", + raw: json.RawMessage(`{`), + want: "", + wantOK: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, ok := extractOperStatus(tt.raw) + if got != tt.want || ok != tt.wantOK { + t.Fatalf("extractOperStatus() = (%q, %v), want (%q, %v)", got, ok, tt.want, tt.wantOK) + } + }) + } +} + +func TestIsBridgeFDB(t *testing.T) { + tests := []struct { + name string + update netlink.NeighUpdate + want bool + }{ + { + name: "bridge family", + update: netlink.NeighUpdate{Neigh: netlink.Neigh{Family: syscall.AF_BRIDGE}}, + want: true, + }, + { + name: "master index set", + update: netlink.NeighUpdate{Neigh: netlink.Neigh{MasterIndex: 10}}, + want: true, + }, + { + name: "master flag set", + update: netlink.NeighUpdate{Neigh: netlink.Neigh{Flags: netlink.NTF_MASTER}}, + want: true, + }, + { + name: "non-bridge", + update: netlink.NeighUpdate{}, + want: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := isBridgeFDB(tt.update); got != tt.want { + t.Fatalf("isBridgeFDB() = %v, want %v", got, tt.want) + } + }) + } +} + +func TestMergeAugments(t *testing.T) { + doc := json.RawMessage(`{"interface":[{"name":"eth0","type":"infix-if-type:ethernet"},{"name":"br0","type":"infix-if-type:bridge"}]}`) + + eth := map[string]json.RawMessage{ + "eth0": json.RawMessage(`{"ethernet":{"speed":"1.000","duplex":"full"},"speed":"1000000000"}`), + } + fdb := map[string]json.RawMessage{ + "br0": json.RawMessage(`[{"mac":"00:11:22:33:44:55"}]`), + } + + got := mergeAugments(doc, eth, nil, fdb, nil, nil, nil, nil) + + var root map[string]any + if err := json.Unmarshal(got, &root); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + ifaces := root["interface"].([]any) + eth0 := ifaces[0].(map[string]any) + if _, ok := eth0["ieee802-ethernet-interface:ethernet"]; !ok { + t.Fatal("ethernet augment not merged into eth0") + } + + br0 := ifaces[1].(map[string]any) + bridge, ok := br0["infix-interfaces:bridge"] + if !ok { + t.Fatal("bridge augment not created for br0") + } + bridgeMap := bridge.(map[string]any) + if _, ok := bridgeMap["fdb"]; !ok { + t.Fatal("fdb not merged into bridge augment") + } +} + +func TestMergeAugmentsNoOp(t *testing.T) { + doc := json.RawMessage(`{"interface":[{"name":"lo"}]}`) + got := mergeAugments(doc, nil, nil, nil, nil, nil, nil, nil) + if string(got) != string(doc) { + t.Fatalf("expected no-op, got %s", string(got)) + } +} + +func TestMergeAugmentsInvalidDoc(t *testing.T) { + doc := json.RawMessage(`{invalid`) + eth := map[string]json.RawMessage{"eth0": json.RawMessage(`{}`)} + got := mergeAugments(doc, eth, nil, nil, nil, nil, nil, nil) + if string(got) != string(doc) { + t.Fatalf("expected passthrough on invalid doc, got %s", string(got)) + } +} + +func TestTreeKey(t *testing.T) { + if treeKey != "ietf-interfaces:interfaces" { + t.Fatalf("treeKey = %q, want %q", treeKey, "ietf-interfaces:interfaces") + } +} + +func TestTransformMDB(t *testing.T) { + // transformMDB receives the output of filterByMDBBridge: a flat array of entries + raw := json.RawMessage(`[{"dev":"br0","port":"e3","grp":"224.1.1.1","state":"temp"},{"dev":"br0","port":"e4","grp":"224.1.1.1","state":"permanent"},{"dev":"br0","port":"e3","grp":"ff02::6a","state":"temp"}]`) + + result := transformMDB(raw) + if result == nil { + t.Fatal("expected non-nil result") + } + + filters, ok := result["multicast-filter"].([]map[string]any) + if !ok { + t.Fatalf("unexpected type: %T", result["multicast-filter"]) + } + if len(filters) != 2 { + t.Fatalf("expected 2 filters, got %d", len(filters)) + } + + if filters[0]["group"] != "224.1.1.1" { + t.Fatalf("unexpected group: %v", filters[0]["group"]) + } + + out, _ := json.Marshal(result) + if !json.Valid(out) { + t.Fatalf("invalid JSON: %s", out) + } +} + +func TestTransformMDBEmpty(t *testing.T) { + if transformMDB(json.RawMessage(`[]`)) != nil { + t.Fatal("expected nil for empty") + } + if transformMDB(json.RawMessage(`[{"dev":"br0","port":"br0","grp":"ff02::6a","state":"temp"}]`)) == nil { + t.Fatal("expected non-nil for router-only entry") + } +} + +func TestSTPFingerprintDeterministic(t *testing.T) { + br1 := map[string]json.RawMessage{ + "br0": json.RawMessage(`{"root-id":"1.000.00:a0:85:00:01:00"}`), + "br1": json.RawMessage(`{"root-id":"2.000.00:a0:85:00:02:00"}`), + } + br2 := map[string]json.RawMessage{ + "br1": json.RawMessage(`{"root-id":"2.000.00:a0:85:00:02:00"}`), + "br0": json.RawMessage(`{"root-id":"1.000.00:a0:85:00:01:00"}`), + } + pt1 := map[string]json.RawMessage{"e1": json.RawMessage(`{"state":"forwarding"}`)} + pt2 := map[string]json.RawMessage{"e1": json.RawMessage(`{"state":"forwarding"}`)} + + if stpFingerprint(br1, pt1) != stpFingerprint(br2, pt2) { + t.Fatal("fingerprint must be independent of map iteration order") + } + + br2["br0"] = json.RawMessage(`{"root-id":"8.000.00:a0:85:00:01:00"}`) + if stpFingerprint(br1, pt1) == stpFingerprint(br2, pt2) { + t.Fatal("fingerprint must change when STP data changes") + } +} + +func names(t *testing.T, raw json.RawMessage) []string { + t.Helper() + var out []string + for _, row := range decodeRows(raw) { + out = append(out, rowString(row, "ifname")) + } + return out +} + +// Rows are keyed by ifindex: a rename replaces the old row instead of +// leaving it behind, and a later delete leaves nothing. +func TestReplaceRowsByIfindexRenameThenDelete(t *testing.T) { + links := json.RawMessage(`[{"ifindex":1,"ifname":"lo"},{"ifindex":5,"ifname":"foo"}]`) + + if got := nameByIndex(links, 5); got != "foo" { + t.Fatalf("nameByIndex = %q, want foo", got) + } + + links = replaceRows(links, "ifindex", 5, json.RawMessage(`[{"ifindex":5,"ifname":"bar"}]`)) + if got := names(t, links); !reflect.DeepEqual(got, []string{"lo", "bar"}) { + t.Fatalf("after rename = %v, want [lo bar]", got) + } + + links = replaceRows(links, "ifindex", 5, nil) + if got := names(t, links); !reflect.DeepEqual(got, []string{"lo"}) { + t.Fatalf("after delete = %v, want [lo]", got) + } + + if got := string(replaceRows(links, "ifindex", 1, nil)); got != "[]" { + t.Fatalf("empty result = %s, want []", got) + } +} + +func TestReplaceRowsByDev(t *testing.T) { + neighs := json.RawMessage(`[{"dst":"10.0.0.1","dev":"eth0"},{"dst":"10.0.0.2","dev":"eth1"}]`) + fresh := withField(json.RawMessage(`[{"dst":"10.0.0.9"}]`), "dev", "eth0") + got := replaceRows(neighs, "dev", "eth0", fresh) + + var rows []map[string]string + if err := json.Unmarshal(got, &rows); err != nil { + t.Fatal(err) + } + want := []map[string]string{{"dst": "10.0.0.2", "dev": "eth1"}, {"dst": "10.0.0.9", "dev": "eth0"}} + if !reflect.DeepEqual(rows, want) { + t.Fatalf("got %v, want %v", rows, want) + } +} + +// A bridge whose last group left has no entry, so its filters go away. +func TestMDBByBridge(t *testing.T) { + raw := json.RawMessage(`[{"mdb":[{"dev":"br0","port":"e1","grp":"239.1.1.1","state":"temp"},` + + `{"dev":"br1","port":"e2","grp":"239.2.2.2","state":"permanent"}],"router":{}}]`) + got := mdbByBridge(raw) + if len(got) != 2 || got["br0"] == nil || got["br1"] == nil { + t.Fatalf("got %v", got) + } + if got := mdbByBridge(json.RawMessage(`[{"mdb":[],"router":{}}]`)); len(got) != 0 { + t.Fatalf("empty dump gave %v", got) + } +} + +func TestSetWireguardAllClearsAndSkipsUnchanged(t *testing.T) { + m := New(nil, nil, nil, nil, tree.New(), nil, slog.Default()) + peer := map[string]json.RawMessage{"wg0": json.RawMessage(`{"peer-status":{}}`)} + + m.SetWireguardAll(peer) + if m.wireguard["wg0"] == nil { + t.Fatal("wg0 not staged") + } + before, _ := m.tree.Info(treeKey) + m.SetWireguardAll(peer) + if after, _ := m.tree.Info(treeKey); !after.LastUpdated.Equal(before.LastUpdated) { + t.Fatal("unchanged WireGuard data rebuilt the document") + } + + m.SetWireguardAll(nil) + if len(m.wireguard) != 0 { + t.Fatalf("wireguard staging not cleared: %v", m.wireguard) + } +} + +type noFiles struct{} + +func (noFiles) Exists(string) bool { return false } +func (noFiles) ReadFile(string) (string, error) { return "", os.ErrNotExist } +func (noFiles) ListDir(string) []string { return nil } + +// The startup sweep asks ethtool for ethernet ports only, not for the +// loopback, bridges or VLANs in the same dump. +func TestEthernetNames(t *testing.T) { + dump := json.RawMessage(`[ + {"ifindex": 1, "ifname": "lo", "link_type": "loopback"}, + {"ifindex": 2, "ifname": "eth0", "link_type": "ether"}, + {"ifindex": 3, "ifname": "br0", "link_type": "ether", "linkinfo": {"info_kind": "bridge"}}, + {"ifindex": 4, "ifname": "eth0.10", "link_type": "ether", "linkinfo": {"info_kind": "vlan"}}, + {"ifindex": 5, "ifname": "eth1", "link_type": "ether"} + ]`) + + got := ethernetNames(dump, noFiles{}) + want := []string{"eth0", "eth1"} + if !reflect.DeepEqual(got, want) { + t.Fatalf("ethernetNames = %v, want %v", got, want) + } +} + +// A delete of the old wifi0 handled after the new wifi0 took the name +// must not drop the new interface's name-keyed state. +func TestLateDeleteKeepsReusedName(t *testing.T) { + m := New(nil, nil, nil, nil, tree.New(), nil, slog.Default()) + m.links = json.RawMessage(`[{"ifindex":14,"ifname":"wifi0"},{"ifindex":16,"ifname":"wifi0"}]`) + m.wifi["wifi0"] = json.RawMessage(`{"station":{"ssid":"campus"}}`) + + m.removeInterface(14) + + if m.wifi["wifi0"] == nil { + t.Fatal("late delete of the old wifi0 dropped the new one's wifi data") + } + if got := nameByIndex(m.links, 16); got != "wifi0" { + t.Fatalf("new wifi0 row lost, links %s", m.links) + } + + m.removeInterface(16) + if m.wifi["wifi0"] != nil { + t.Fatal("deleting the last wifi0 kept its wifi data") + } +} + +func TestLinkRowFor(t *testing.T) { + for _, tc := range []struct { + raw string + want bool + }{ + {`[{"ifindex":16,"ifname":"wifi0"}]`, true}, + {`[{}]`, false}, + {`[{"ifindex":14,"ifname":"wifi0"}]`, false}, + {`[{"ifindex":16}]`, false}, + {`[]`, false}, + } { + if got := linkRowFor(json.RawMessage(tc.raw), 16); got != tc.want { + t.Errorf("linkRowFor(%s) = %v, want %v", tc.raw, got, tc.want) + } + } +} + +func TestLastChangeStampsTransitionsOnly(t *testing.T) { + m := New(nil, nil, nil, nil, tree.New(), nil, slog.Default()) + m.tree.Set(treeKey, json.RawMessage(`{"interface":[{"name":"e1","oper-status":"up"},{"name":"e2","oper-status":"down"}]}`)) + + if got := m.LastChange(); got != nil { + t.Fatalf("no transitions yet, want nil overlay, got %s", got) + } + + m.mu.Lock() + m.lastOperStatus["e1"] = "DOWN" + m.lastOperStatus["e2"] = "DOWN" + m.updateOperStatus("e1", json.RawMessage(`[{"ifname":"e1","operstate":"UP"}]`)) + m.updateOperStatus("e2", json.RawMessage(`[{"ifname":"e2","operstate":"DOWN"}]`)) + m.mu.Unlock() + + var out struct { + Interface []map[string]any `json:"interface"` + } + if err := json.Unmarshal(m.LastChange(), &out); err != nil { + t.Fatal(err) + } + if len(out.Interface) != 2 { + t.Fatalf("overlay must carry the whole list, got %d entries", len(out.Interface)) + } + if _, ok := out.Interface[0]["last-change"].(string); !ok { + t.Errorf("e1 changed state, want last-change, got %v", out.Interface[0]) + } + if _, ok := out.Interface[1]["last-change"]; ok { + t.Errorf("e2 kept its state, want no last-change, got %v", out.Interface[1]) + } + + m.mu.Lock() + m.forgetName("e1") + m.mu.Unlock() + if got := m.LastChange(); got != nil { + t.Errorf("stamp must go with the interface, got %s", got) + } +} + +func TestStampLastChange(t *testing.T) { + doc := json.RawMessage(`{"interface":[{"name":"e1"},{"name":"e2"}]}`) + changed := time.Now() + stamps := map[string]time.Time{"e1": changed} + + out := stampLastChange(doc, stamps, changed.Add(time.Minute)) + want := changed.UTC().Format("2006-01-02T15:04:05+00:00") + if !strings.Contains(string(out), `{"last-change":"`+want+`","name":"e1"}`) { + t.Errorf("want %s on e1 in %s", want, out) + } + if !strings.Contains(string(out), `{"name":"e2"}`) { + t.Errorf("e2 must be untouched in %s", out) + } + if got := stampLastChange(json.RawMessage(`garbage`), stamps, changed); got != nil { + t.Errorf("bad doc must yield no overlay, got %s", got) + } +} diff --git a/src/yangerd/internal/nl80211/bands_test.go b/src/yangerd/internal/nl80211/bands_test.go new file mode 100644 index 000000000..35a64779b --- /dev/null +++ b/src/yangerd/internal/nl80211/bands_test.go @@ -0,0 +1,35 @@ +package nl80211 + +import "testing" + +// Only 2.4, 5 and 6 GHz are reported; hwsim's S1G and 60 GHz bands are +// dropped rather than listed as "Unknown". +func TestFinalizeBandsDropsUnsupported(t *testing.T) { + bands := map[uint16]*bandInfo{ + 0: {frequencies: []interface{}{2412, 2437}, htCapable: true}, + 1: {frequencies: []interface{}{5180, 5200}, vhtCapable: true}, + 2: {frequencies: []interface{}{58320, 60480}}, + 4: {frequencies: []interface{}{902, 904}}, + } + + var names []string + for _, raw := range finalizeBands(bands) { + names = append(names, raw.(map[string]interface{})["name"].(string)) + } + if len(names) != 2 || names[0] != "2.4 GHz" || names[1] != "5 GHz" { + t.Fatalf("bands = %v, want [2.4 GHz 5 GHz]", names) + } +} + +func TestManufacturerFor(t *testing.T) { + for driver, want := range map[string]string{ + "mac80211_hwsim": "Virtual (hwsim)", + "mt7915e": "MediaTek Inc.", + "ath11k_pci": "Qualcomm Atheros", + "": "Unknown", + } { + if got := manufacturerFor(driver); got != want { + t.Errorf("manufacturerFor(%q) = %q, want %q", driver, got, want) + } + } +} diff --git a/src/yangerd/internal/nl80211/nl80211.go b/src/yangerd/internal/nl80211/nl80211.go new file mode 100644 index 000000000..cec04842e --- /dev/null +++ b/src/yangerd/internal/nl80211/nl80211.go @@ -0,0 +1,668 @@ +package nl80211 + +import ( + "encoding/binary" + "fmt" + "os" + "path/filepath" + "sort" + "strconv" + "strings" + + "github.com/mdlayher/genetlink" + "github.com/mdlayher/netlink" + "golang.org/x/sys/unix" +) + +type Client struct { + conn *genetlink.Conn + family genetlink.Family +} + +func Dial() (*Client, error) { + conn, err := genetlink.Dial(nil) + if err != nil { + return nil, fmt.Errorf("dial genetlink: %w", err) + } + + family, err := conn.GetFamily("nl80211") + if err != nil { + _ = conn.Close() + return nil, fmt.Errorf("resolve nl80211 family: %w", err) + } + + return &Client{conn: conn, family: family}, nil +} + +func (c *Client) Close() error { + if c == nil || c.conn == nil { + return nil + } + return c.conn.Close() +} + +func (c *Client) ListPhys() ([]string, error) { + msgs, err := c.execute(unix.NL80211_CMD_GET_WIPHY, nil, netlink.Request|netlink.Dump) + if err != nil { + return nil, err + } + + set := make(map[string]bool) + for _, msg := range msgs { + attrs, err := netlink.NewAttributeDecoder(msg.Data) + if err != nil { + continue + } + for attrs.Next() { + if attrs.Type() != unix.NL80211_ATTR_WIPHY_NAME { + continue + } + name := attrs.String() + if name != "" { + set[name] = true + } + } + if err := attrs.Err(); err != nil { + continue + } + } + + out := make([]string, 0, len(set)) + for name := range set { + out = append(out, name) + } + sort.Strings(out) + + return out, nil +} + +func (c *Client) PhyInterfaces() (map[string][]string, error) { + msgs, err := c.execute(unix.NL80211_CMD_GET_INTERFACE, nil, netlink.Request|netlink.Dump) + if err != nil { + return nil, err + } + + out := make(map[string][]string) + for _, msg := range msgs { + ad, err := netlink.NewAttributeDecoder(msg.Data) + if err != nil { + continue + } + + phyIdx := -1 + ifname := "" + + for ad.Next() { + switch ad.Type() { + case unix.NL80211_ATTR_WIPHY: + phyIdx = int(ad.Uint32()) + case unix.NL80211_ATTR_IFNAME: + ifname = ad.String() + } + } + if err := ad.Err(); err != nil { + continue + } + if phyIdx < 0 || ifname == "" { + continue + } + + k := strconv.Itoa(phyIdx) + out[k] = append(out[k], ifname) + } + + for k := range out { + sort.Strings(out[k]) + } + + return out, nil +} + +func (c *Client) PhyInfo(phyName string) (map[string]interface{}, error) { + ae := netlink.NewAttributeEncoder() + ae.Flag(unix.NL80211_ATTR_SPLIT_WIPHY_DUMP, true) + req, _ := ae.Encode() + msgs, err := c.execute(unix.NL80211_CMD_GET_WIPHY, req, netlink.Request|netlink.Dump) + if err != nil { + return nil, err + } + + // Collect all messages belonging to the target PHY. The kernel + // identifies fragments by repeating NL80211_ATTR_WIPHY (index) + // or NL80211_ATTR_WIPHY_NAME in each fragment. + targetIdx := -1 + var phyMsgs [][]byte + for _, msg := range msgs { + idx, name := parseWiphyIdent(msg.Data) + if name == phyName { + targetIdx = idx + phyMsgs = append(phyMsgs, msg.Data) + } else if targetIdx >= 0 && idx == targetIdx { + phyMsgs = append(phyMsgs, msg.Data) + } + } + if len(phyMsgs) == 0 { + return nil, fmt.Errorf("phy %q not found", phyName) + } + + info := map[string]interface{}{ + "bands": []interface{}{}, + "driver": readDriver(phyName), + "manufacturer": readManufacturer(phyName), + "interface_combinations": []interface{}{}, + "max_txpower": 0, + "num_virtual_interfaces": 0, + } + + phyIdx := -1 + bandMap := make(map[uint16]*bandInfo) + + for _, data := range phyMsgs { + ad, err := netlink.NewAttributeDecoder(data) + if err != nil { + continue + } + for ad.Next() { + switch ad.Type() { + case unix.NL80211_ATTR_WIPHY: + phyIdx = int(ad.Uint32()) + case unix.NL80211_ATTR_WIPHY_BANDS: + mergeBands(bandMap, ad.Bytes()) + case unix.NL80211_ATTR_INTERFACE_COMBINATIONS: + if combs := parseInterfaceCombinations(ad.Bytes()); len(combs) > 0 { + info["interface_combinations"] = combs + } + case unix.NL80211_ATTR_WIPHY_TX_POWER_LEVEL: + info["max_txpower"] = int(ad.Uint32() / 100) + } + } + } + if bands := finalizeBands(bandMap); len(bands) > 0 { + info["bands"] = bands + } + + if phyIdx >= 0 { + ifs, err := c.PhyInterfaces() + if err == nil { + info["num_virtual_interfaces"] = len(ifs[strconv.Itoa(phyIdx)]) + } + } + + return info, nil +} + +func (c *Client) Survey(ifindex int) ([]map[string]interface{}, error) { + ae := netlink.NewAttributeEncoder() + ae.Uint32(unix.NL80211_ATTR_IFINDEX, uint32(ifindex)) + req, err := ae.Encode() + if err != nil { + return nil, fmt.Errorf("encode get_survey request: %w", err) + } + + msgs, err := c.execute(unix.NL80211_CMD_GET_SURVEY, req, netlink.Request|netlink.Dump) + if err != nil { + return nil, err + } + + out := make([]map[string]interface{}, 0) + for _, msg := range msgs { + ad, err := netlink.NewAttributeDecoder(msg.Data) + if err != nil { + continue + } + for ad.Next() { + if ad.Type() != unix.NL80211_ATTR_SURVEY_INFO { + continue + } + entry := parseSurveyEntry(ad.Bytes()) + if entry != nil { + out = append(out, entry) + } + } + if err := ad.Err(); err != nil { + continue + } + } + + return out, nil +} + +func (c *Client) execute(cmd uint8, data []byte, flags netlink.HeaderFlags) ([]genetlink.Message, error) { + msgs, err := c.conn.Execute( + genetlink.Message{Header: genetlink.Header{Command: cmd, Version: c.family.Version}, Data: data}, + c.family.ID, + flags, + ) + if err != nil { + return nil, fmt.Errorf("nl80211 command %d: %w", cmd, err) + } + + return msgs, nil +} + +func parseWiphyIdent(data []byte) (int, string) { + ad, err := netlink.NewAttributeDecoder(data) + if err != nil { + return -1, "" + } + idx := -1 + name := "" + for ad.Next() { + switch ad.Type() { + case unix.NL80211_ATTR_WIPHY: + idx = int(ad.Uint32()) + case unix.NL80211_ATTR_WIPHY_NAME: + name = ad.String() + } + } + return idx, name +} + +type bandInfo struct { + frequencies []interface{} + htCapable bool + vhtCapable bool + heCapable bool +} + +func mergeBands(m map[uint16]*bandInfo, data []byte) { + ad, err := netlink.NewAttributeDecoder(data) + if err != nil { + return + } + for ad.Next() { + bandType := ad.Type() + bi, ok := m[bandType] + if !ok { + bi = &bandInfo{} + m[bandType] = bi + } + mergeBandAttrs(bi, ad.Bytes()) + } +} + +func mergeBandAttrs(bi *bandInfo, data []byte) { + nad, err := netlink.NewAttributeDecoder(data) + if err != nil { + return + } + for nad.Next() { + switch nad.Type() { + case unix.NL80211_BAND_ATTR_FREQS: + if freqs := parseBandFrequencies(nad.Bytes()); len(freqs) > 0 { + bi.frequencies = freqs + } + case unix.NL80211_BAND_ATTR_HT_CAPA: + if nad.Uint16() != 0 { + bi.htCapable = true + } + case unix.NL80211_BAND_ATTR_VHT_CAPA: + if nad.Uint32() != 0 { + bi.vhtCapable = true + } + case unix.NL80211_BAND_ATTR_IFTYPE_DATA: + if len(nad.Bytes()) > 0 { + bi.heCapable = true + } + } + } +} + +func finalizeBands(m map[uint16]*bandInfo) []interface{} { + keys := make([]int, 0, len(m)) + for k := range m { + keys = append(keys, int(k)) + } + sort.Ints(keys) + + out := make([]interface{}, 0, len(keys)) + for _, k := range keys { + bi := m[uint16(k)] + if len(bi.frequencies) == 0 && !bi.htCapable && !bi.vhtCapable && !bi.heCapable { + continue + } + // Infix supports 2.4, 5 and 6 GHz only. Hardware may expose + // others, like S1G and 60 GHz on hwsim, that are neither + // configured nor reported. + name := detectBandName(bi.frequencies) + if name == "Unknown" { + continue + } + out = append(out, map[string]interface{}{ + "band": k, + "name": name, + "ht_capable": bi.htCapable, + "vht_capable": bi.vhtCapable, + "he_capable": bi.heCapable, + "frequencies": bi.frequencies, + }) + } + return out +} + +func parseBandFrequencies(data []byte) []interface{} { + out := make([]interface{}, 0) + + ad, err := netlink.NewAttributeDecoder(data) + if err != nil { + return out + } + + for ad.Next() { + freq, ok := parseFrequencyEntry(ad.Bytes()) + if ok { + out = append(out, freq) + } + } + + return out +} + +func parseFrequencyEntry(data []byte) (int, bool) { + ad, err := netlink.NewAttributeDecoder(data) + if err != nil { + return 0, false + } + + freq := 0 + disabled := false + + for ad.Next() { + switch ad.Type() { + case unix.NL80211_FREQUENCY_ATTR_FREQ: + freq = int(ad.Uint32()) + case unix.NL80211_FREQUENCY_ATTR_DISABLED: + disabled = true + } + } + + if disabled || freq == 0 { + return 0, false + } + + return freq, true +} + +func detectBandName(freqs []interface{}) string { + has24 := false + has5 := false + has6 := false + + for _, f := range freqs { + freq, ok := f.(int) + if !ok { + continue + } + switch { + case freq >= 2400 && freq <= 2500: + has24 = true + case freq >= 5000 && freq <= 5900: + has5 = true + case freq >= 5925 && freq <= 7125: + has6 = true + } + } + + switch { + case has24: + return "2.4 GHz" + case has5: + return "5 GHz" + case has6: + return "6 GHz" + default: + return "Unknown" + } +} + +func parseInterfaceCombinations(data []byte) []interface{} { + out := make([]interface{}, 0) + + ad, err := netlink.NewAttributeDecoder(data) + if err != nil { + return out + } + + for ad.Next() { + comb := parseInterfaceCombination(ad.Bytes()) + if comb != nil { + out = append(out, comb) + } + } + + return out +} + +func parseInterfaceCombination(data []byte) map[string]interface{} { + limits := make([]interface{}, 0) + + ad, err := netlink.NewAttributeDecoder(data) + if err != nil { + return nil + } + + for ad.Next() { + if ad.Type() != unix.NL80211_IFACE_COMB_LIMITS { + continue + } + limits = parseInterfaceLimits(ad.Bytes()) + } + + if len(limits) == 0 { + return nil + } + + return map[string]interface{}{"limits": limits} +} + +func parseInterfaceLimits(data []byte) []interface{} { + out := make([]interface{}, 0) + + ad, err := netlink.NewAttributeDecoder(data) + if err != nil { + return out + } + + for ad.Next() { + entry := parseInterfaceLimitEntry(ad.Bytes()) + if entry != nil { + out = append(out, entry) + } + } + + return out +} + +func parseInterfaceLimitEntry(data []byte) map[string]interface{} { + max := 0 + types := make([]interface{}, 0) + + ad, err := netlink.NewAttributeDecoder(data) + if err != nil { + return nil + } + + for ad.Next() { + switch ad.Type() { + case unix.NL80211_IFACE_LIMIT_MAX: + max = int(ad.Uint32()) + case unix.NL80211_IFACE_LIMIT_TYPES: + types = parseIfaceLimitTypes(ad.Bytes()) + } + } + + if max == 0 || len(types) == 0 { + return nil + } + + return map[string]interface{}{ + "max": max, + "types": types, + } +} + +func parseIfaceLimitTypes(data []byte) []interface{} { + out := make([]interface{}, 0) + seen := make(map[string]bool) + + ad, err := netlink.NewAttributeDecoder(data) + if err != nil { + return out + } + + for ad.Next() { + if len(ad.Bytes()) == 4 { + iftype := int(ad.Uint32()) + if s, ok := iftypeName(iftype); ok { + if !seen[s] { + seen[s] = true + out = append(out, s) + } + } + continue + } + + iftype := int(ad.Type()) + if s, ok := iftypeName(iftype); ok { + if !seen[s] { + seen[s] = true + out = append(out, s) + } + } + } + + sort.Slice(out, func(i, j int) bool { + return out[i].(string) < out[j].(string) + }) + + return out +} + +func iftypeName(v int) (string, bool) { + switch v { + case 0: + return "unspecified", true + case 1: + return "adhoc", true + case 2: + return "station", true + case 3: + return "AP", true + case 4: + return "AP_VLAN", true + case 5: + return "WDS", true + case 6: + return "monitor", true + case 7: + return "mesh_point", true + case 8: + return "P2P_client", true + case 9: + return "P2P_GO", true + case 10: + return "P2P_device", true + default: + return "", false + } +} + +func parseSurveyEntry(data []byte) map[string]interface{} { + entry := map[string]interface{}{ + "frequency": 0, + "noise": 0, + "in_use": false, + "active_time": 0, + "busy_time": 0, + "receive_time": 0, + "transmit_time": 0, + } + + ad, err := netlink.NewAttributeDecoder(data) + if err != nil { + return nil + } + + hasFrequency := false + + for ad.Next() { + switch ad.Type() { + case unix.NL80211_SURVEY_INFO_FREQUENCY: + entry["frequency"] = int(readUint(ad.Bytes())) + hasFrequency = true + case unix.NL80211_SURVEY_INFO_NOISE: + entry["noise"] = int(ad.Int8()) + case unix.NL80211_SURVEY_INFO_IN_USE: + entry["in_use"] = true + case unix.NL80211_SURVEY_INFO_TIME: + entry["active_time"] = int(readUint(ad.Bytes())) + case unix.NL80211_SURVEY_INFO_TIME_BUSY: + entry["busy_time"] = int(readUint(ad.Bytes())) + case unix.NL80211_SURVEY_INFO_TIME_RX: + entry["receive_time"] = int(readUint(ad.Bytes())) + case unix.NL80211_SURVEY_INFO_TIME_TX: + entry["transmit_time"] = int(readUint(ad.Bytes())) + } + } + + if !hasFrequency { + return nil + } + + return entry +} + +func readUint(b []byte) uint64 { + switch len(b) { + case 1: + return uint64(b[0]) + case 2: + return uint64(binary.NativeEndian.Uint16(b)) + case 4: + return uint64(binary.NativeEndian.Uint32(b)) + case 8: + return binary.NativeEndian.Uint64(b) + default: + return 0 + } +} + +func readDriver(phyName string) string { + path := filepath.Join("/sys/class/ieee80211", phyName, "device", "driver") + target, err := os.Readlink(path) + if err != nil { + return "" + } + if base := filepath.Base(target); base != "." && base != "/" { + return base + } + return "" +} + +func readManufacturer(phyName string) string { + return manufacturerFor(readDriver(phyName)) +} + +// manufacturerFor names the vendor behind a wireless driver. +func manufacturerFor(driver string) string { + if driver == "" { + return "Unknown" + } + d := strings.ToLower(driver) + switch { + case strings.Contains(d, "hwsim"): + return "Virtual (hwsim)" + case strings.Contains(d, "mt") || strings.Contains(d, "mediatek"): + return "MediaTek Inc." + case strings.Contains(d, "rtw") || strings.Contains(d, "realtek"): + return "Realtek Semiconductor Corp." + case strings.Contains(d, "ath") || strings.Contains(d, "qca"): + return "Qualcomm Atheros" + case strings.Contains(d, "iwl") || strings.Contains(d, "intel"): + return "Intel Corporation" + case strings.Contains(d, "brcm") || strings.Contains(d, "broadcom"): + return "Broadcom Inc." + default: + return "Unknown" + } +} diff --git a/src/yangerd/internal/nl80211/station.go b/src/yangerd/internal/nl80211/station.go new file mode 100644 index 000000000..6f7dbbe39 --- /dev/null +++ b/src/yangerd/internal/nl80211/station.go @@ -0,0 +1,201 @@ +package nl80211 + +import ( + "fmt" + "net" + + "github.com/mdlayher/netlink" + "golang.org/x/sys/unix" +) + +// Station is one entry of the kernel station table for an interface: +// an associated client in AP mode, a peer in mesh mode. +type Station struct { + MAC string + Signal int8 + HasSignal bool + ConnectedTime uint32 + RxBytes uint64 + TxBytes uint64 + RxPackets uint64 + TxPackets uint64 + RxBitrate uint32 // 100 kbps + TxBitrate uint32 // 100 kbps +} + +// Stations dumps the station table of ifindex. +func (c *Client) Stations(ifindex int) ([]Station, error) { + ae := netlink.NewAttributeEncoder() + ae.Uint32(unix.NL80211_ATTR_IFINDEX, uint32(ifindex)) + req, err := ae.Encode() + if err != nil { + return nil, fmt.Errorf("encode get_station request: %w", err) + } + + msgs, err := c.execute(unix.NL80211_CMD_GET_STATION, req, netlink.Request|netlink.Dump) + if err != nil { + return nil, err + } + + out := make([]Station, 0, len(msgs)) + for _, msg := range msgs { + if st, ok := parseStation(msg.Data); ok { + out = append(out, st) + } + } + return out, nil +} + +// parseStation decodes one NL80211_CMD_NEW_STATION message. +func parseStation(data []byte) (Station, bool) { + var st Station + + ad, err := netlink.NewAttributeDecoder(data) + if err != nil { + return st, false + } + for ad.Next() { + switch ad.Type() { + case unix.NL80211_ATTR_MAC: + st.MAC = net.HardwareAddr(ad.Bytes()).String() + case unix.NL80211_ATTR_STA_INFO: + parseStationInfo(&st, ad.Bytes()) + } + } + return st, st.MAC != "" +} + +func parseStationInfo(st *Station, data []byte) { + ad, err := netlink.NewAttributeDecoder(data) + if err != nil { + return + } + for ad.Next() { + switch ad.Type() { + case unix.NL80211_STA_INFO_SIGNAL: + st.Signal = ad.Int8() + st.HasSignal = true + case unix.NL80211_STA_INFO_CONNECTED_TIME: + st.ConnectedTime = ad.Uint32() + case unix.NL80211_STA_INFO_RX_BYTES: + if st.RxBytes == 0 { + st.RxBytes = uint64(ad.Uint32()) + } + case unix.NL80211_STA_INFO_TX_BYTES: + if st.TxBytes == 0 { + st.TxBytes = uint64(ad.Uint32()) + } + case unix.NL80211_STA_INFO_RX_BYTES64: + st.RxBytes = ad.Uint64() + case unix.NL80211_STA_INFO_TX_BYTES64: + st.TxBytes = ad.Uint64() + case unix.NL80211_STA_INFO_RX_PACKETS: + st.RxPackets = uint64(ad.Uint32()) + case unix.NL80211_STA_INFO_TX_PACKETS: + st.TxPackets = uint64(ad.Uint32()) + case unix.NL80211_STA_INFO_RX_BITRATE: + st.RxBitrate = parseBitrate(ad.Bytes()) + case unix.NL80211_STA_INFO_TX_BITRATE: + st.TxBitrate = parseBitrate(ad.Bytes()) + } + } +} + +// parseBitrate returns the rate in 100 kbps from a rate_info nest. The +// 32-bit attribute is preferred, the 16-bit one saturates at 6.5 Gbps. +func parseBitrate(data []byte) uint32 { + ad, err := netlink.NewAttributeDecoder(data) + if err != nil { + return 0 + } + var rate16 uint16 + for ad.Next() { + switch ad.Type() { + case unix.NL80211_RATE_INFO_BITRATE32: + return ad.Uint32() + case unix.NL80211_RATE_INFO_BITRATE: + rate16 = ad.Uint16() + } + } + return uint32(rate16) +} + +// InterfaceType returns the iftype of ifindex as iw names it, e.g. +// "station", "AP" or "mesh_point". +func (c *Client) InterfaceType(ifindex int) (string, error) { + ae := netlink.NewAttributeEncoder() + ae.Uint32(unix.NL80211_ATTR_IFINDEX, uint32(ifindex)) + req, err := ae.Encode() + if err != nil { + return "", fmt.Errorf("encode get_interface request: %w", err) + } + + msgs, err := c.execute(unix.NL80211_CMD_GET_INTERFACE, req, netlink.Request) + if err != nil { + return "", err + } + + for _, msg := range msgs { + if name, ok := parseInterfaceType(msg.Data); ok { + return name, nil + } + } + return "", fmt.Errorf("no iftype for ifindex %d", ifindex) +} + +func parseInterfaceType(data []byte) (string, bool) { + ad, err := netlink.NewAttributeDecoder(data) + if err != nil { + return "", false + } + for ad.Next() { + if ad.Type() == unix.NL80211_ATTR_IFTYPE { + return iftypeName(int(ad.Uint32())) + } + } + return "", false +} + +// MeshForwarding returns the mesh_fwding parameter of a mesh interface. +func (c *Client) MeshForwarding(ifindex int) (bool, error) { + ae := netlink.NewAttributeEncoder() + ae.Uint32(unix.NL80211_ATTR_IFINDEX, uint32(ifindex)) + req, err := ae.Encode() + if err != nil { + return false, fmt.Errorf("encode get_mesh_config request: %w", err) + } + + msgs, err := c.execute(unix.NL80211_CMD_GET_MESH_CONFIG, req, netlink.Request) + if err != nil { + return false, err + } + + for _, msg := range msgs { + if fwd, ok := parseMeshForwarding(msg.Data); ok { + return fwd, nil + } + } + return false, fmt.Errorf("no mesh config for ifindex %d", ifindex) +} + +func parseMeshForwarding(data []byte) (bool, bool) { + ad, err := netlink.NewAttributeDecoder(data) + if err != nil { + return false, false + } + for ad.Next() { + if ad.Type() != unix.NL80211_ATTR_MESH_CONFIG { + continue + } + nested, err := netlink.NewAttributeDecoder(ad.Bytes()) + if err != nil { + return false, false + } + for nested.Next() { + if nested.Type() == unix.NL80211_MESHCONF_FORWARDING { + return nested.Uint8() != 0, true + } + } + } + return false, false +} diff --git a/src/yangerd/internal/nl80211/station_test.go b/src/yangerd/internal/nl80211/station_test.go new file mode 100644 index 000000000..1d6f5d717 --- /dev/null +++ b/src/yangerd/internal/nl80211/station_test.go @@ -0,0 +1,105 @@ +package nl80211 + +import ( + "testing" + + "github.com/mdlayher/netlink" + "golang.org/x/sys/unix" +) + +func encode(t *testing.T, fn func(ae *netlink.AttributeEncoder)) []byte { + t.Helper() + ae := netlink.NewAttributeEncoder() + fn(ae) + b, err := ae.Encode() + if err != nil { + t.Fatal(err) + } + return b +} + +func TestParseStation(t *testing.T) { + msg := encode(t, func(ae *netlink.AttributeEncoder) { + ae.Bytes(unix.NL80211_ATTR_MAC, []byte{0x02, 0x00, 0x00, 0x00, 0x00, 0x02}) + ae.Nested(unix.NL80211_ATTR_STA_INFO, func(sta *netlink.AttributeEncoder) error { + sta.Int8(unix.NL80211_STA_INFO_SIGNAL, -47) + sta.Uint32(unix.NL80211_STA_INFO_CONNECTED_TIME, 321) + sta.Uint32(unix.NL80211_STA_INFO_RX_BYTES, 1) + sta.Uint64(unix.NL80211_STA_INFO_RX_BYTES64, 5000000000) + sta.Uint32(unix.NL80211_STA_INFO_TX_BYTES, 2) + sta.Uint64(unix.NL80211_STA_INFO_TX_BYTES64, 6000000000) + sta.Uint32(unix.NL80211_STA_INFO_RX_PACKETS, 10) + sta.Uint32(unix.NL80211_STA_INFO_TX_PACKETS, 20) + sta.Nested(unix.NL80211_STA_INFO_TX_BITRATE, func(r *netlink.AttributeEncoder) error { + r.Uint16(unix.NL80211_RATE_INFO_BITRATE, 8667) + r.Uint32(unix.NL80211_RATE_INFO_BITRATE32, 8667) + return nil + }) + sta.Nested(unix.NL80211_STA_INFO_RX_BITRATE, func(r *netlink.AttributeEncoder) error { + r.Uint16(unix.NL80211_RATE_INFO_BITRATE, 650) + return nil + }) + return nil + }) + }) + + st, ok := parseStation(msg) + if !ok { + t.Fatal("station not parsed") + } + if st.MAC != "02:00:00:00:00:02" { + t.Errorf("MAC = %q", st.MAC) + } + if !st.HasSignal || st.Signal != -47 { + t.Errorf("Signal = %d (has %v)", st.Signal, st.HasSignal) + } + if st.ConnectedTime != 321 { + t.Errorf("ConnectedTime = %d", st.ConnectedTime) + } + if st.RxBytes != 5000000000 || st.TxBytes != 6000000000 { + t.Errorf("bytes = %d/%d, 64-bit counters must win", st.RxBytes, st.TxBytes) + } + if st.RxPackets != 10 || st.TxPackets != 20 { + t.Errorf("packets = %d/%d", st.RxPackets, st.TxPackets) + } + if st.TxBitrate != 8667 || st.RxBitrate != 650 { + t.Errorf("bitrate = %d/%d", st.TxBitrate, st.RxBitrate) + } +} + +func TestParseStationWithoutMAC(t *testing.T) { + msg := encode(t, func(ae *netlink.AttributeEncoder) { + ae.Uint32(unix.NL80211_ATTR_IFINDEX, 3) + }) + if _, ok := parseStation(msg); ok { + t.Error("a message without a MAC is not a station") + } +} + +func TestParseMeshForwarding(t *testing.T) { + for _, want := range []bool{true, false} { + msg := encode(t, func(ae *netlink.AttributeEncoder) { + ae.Uint32(unix.NL80211_ATTR_IFINDEX, 3) + ae.Nested(unix.NL80211_ATTR_MESH_CONFIG, func(mc *netlink.AttributeEncoder) error { + mc.Uint16(unix.NL80211_MESHCONF_RETRY_TIMEOUT, 100) + if want { + mc.Uint8(unix.NL80211_MESHCONF_FORWARDING, 1) + } else { + mc.Uint8(unix.NL80211_MESHCONF_FORWARDING, 0) + } + return nil + }) + }) + got, ok := parseMeshForwarding(msg) + if !ok || got != want { + t.Errorf("forwarding = %v (ok %v), want %v", got, ok, want) + } + } + + msg := encode(t, func(ae *netlink.AttributeEncoder) { + ae.Uint32(unix.NL80211_ATTR_IFINDEX, 3) + }) + if _, ok := parseMeshForwarding(msg); ok { + t.Error("no mesh config must not parse") + } +} diff --git a/src/yangerd/internal/numconv/numconv.go b/src/yangerd/internal/numconv/numconv.go new file mode 100644 index 000000000..6a9d91bba --- /dev/null +++ b/src/yangerd/internal/numconv/numconv.go @@ -0,0 +1,88 @@ +// Package numconv converts the loosely typed numbers yangerd gets from +// decoded JSON, D-Bus variants and command output into Go integers. +package numconv + +import ( + "encoding/json" + "strconv" + "strings" +) + +// Int returns v as an int. ok is false when v is neither a number nor +// a string holding one. Floats are truncated. +func Int(v any) (int, bool) { + switch n := v.(type) { + case int: + return n, true + case int8: + return int(n), true + case int16: + return int(n), true + case int32: + return int(n), true + case int64: + return int(n), true + case uint: + return int(n), true + case uint8: + return int(n), true + case uint16: + return int(n), true + case uint32: + return int(n), true + case uint64: + return int(n), true + case float32: + return int(n), true + case float64: + return int(n), true + case json.Number: + if i, err := n.Int64(); err == nil { + return int(i), true + } + if f, err := n.Float64(); err == nil { + return int(f), true + } + case string: + if i, err := strconv.Atoi(strings.TrimSpace(n)); err == nil { + return i, true + } + } + return 0, false +} + +// IntOrZero is Int for callers where a missing value reads as 0. +func IntOrZero(v any) int { + n, _ := Int(v) + return n +} + +// Uint64 returns v as a uint64, for counters. Negative values and +// anything that is not a number are 0. +func Uint64(v any) uint64 { + switch n := v.(type) { + case uint: + return uint64(n) + case uint8: + return uint64(n) + case uint16: + return uint64(n) + case uint32: + return uint64(n) + case uint64: + return n + case json.Number: + if u, err := strconv.ParseUint(n.String(), 10, 64); err == nil { + return u + } + case string: + if u, err := strconv.ParseUint(strings.TrimSpace(n), 10, 64); err == nil { + return u + } + return 0 + } + if i, ok := Int(v); ok && i > 0 { + return uint64(i) + } + return 0 +} diff --git a/src/yangerd/internal/numconv/numconv_test.go b/src/yangerd/internal/numconv/numconv_test.go new file mode 100644 index 000000000..acd633a60 --- /dev/null +++ b/src/yangerd/internal/numconv/numconv_test.go @@ -0,0 +1,55 @@ +package numconv + +import ( + "encoding/json" + "testing" +) + +func TestInt(t *testing.T) { + tests := []struct { + in any + want int + ok bool + }{ + {42, 42, true}, + {int32(-7), -7, true}, + {uint16(16), 16, true}, + {float64(99.9), 99, true}, + {json.Number("12"), 12, true}, + {json.Number("1.5"), 1, true}, + {" 42 ", 42, true}, + {"nope", 0, false}, + {nil, 0, false}, + {true, 0, false}, + } + for _, tc := range tests { + got, ok := Int(tc.in) + if got != tc.want || ok != tc.ok { + t.Errorf("Int(%#v) = %d, %v; want %d, %v", tc.in, got, ok, tc.want, tc.ok) + } + } +} + +func TestUint64(t *testing.T) { + tests := []struct { + in any + want uint64 + }{ + {uint8(8), 8}, + {uint64(1 << 60), 1 << 60}, + {42, 42}, + {-1, 0}, + {int64(-9), 0}, + {float64(99.9), 99}, + {float64(-0.1), 0}, + {"42", 42}, + {"18446744073709551615", 18446744073709551615}, + {json.Number("18446744073709551615"), 18446744073709551615}, + {"nope", 0}, + } + for _, tc := range tests { + if got := Uint64(tc.in); got != tc.want { + t.Errorf("Uint64(%#v) = %d, want %d", tc.in, got, tc.want) + } + } +} diff --git a/src/yangerd/internal/ptpmonitor/mapping.go b/src/yangerd/internal/ptpmonitor/mapping.go new file mode 100644 index 000000000..4b91171f8 --- /dev/null +++ b/src/yangerd/internal/ptpmonitor/mapping.go @@ -0,0 +1,243 @@ +package ptpmonitor + +import ( + "fmt" + "strconv" + + ptp "github.com/facebook/time/ptp/protocol" +) + +// clockClassIdentity maps IEEE 1588 clockClass values to +// ieee1588-ptp-tt identity names (identityref, not uint8). +var clockClassIdentity = map[uint8]string{ + 6: "ieee1588-ptp-tt:cc-primary-sync", + 7: "ieee1588-ptp-tt:cc-primary-sync-lost", + 13: "ieee1588-ptp-tt:cc-application-specific-sync", + 14: "ieee1588-ptp-tt:cc-application-specific-sync-lost", + 52: "ieee1588-ptp-tt:cc-primary-sync-alternative-a", + 58: "ieee1588-ptp-tt:cc-application-specific-alternative-a", + 187: "ieee1588-ptp-tt:cc-primary-sync-alternative-b", + 193: "ieee1588-ptp-tt:cc-application-specific-alternative-b", + 248: "ieee1588-ptp-tt:cc-default", + 255: "ieee1588-ptp-tt:cc-time-receiver-only", +} + +// clockAccuracyIdentity maps IEEE 1588 clockAccuracy values to +// ieee1588-ptp-tt identity names. 0xfe (unknown) has no identity and +// is omitted. +var clockAccuracyIdentity = map[uint8]string{ + 0x17: "ieee1588-ptp-tt:ca-time-accurate-to-1000-fs", + 0x18: "ieee1588-ptp-tt:ca-time-accurate-to-2500-fs", + 0x19: "ieee1588-ptp-tt:ca-time-accurate-to-10-ps", + 0x1a: "ieee1588-ptp-tt:ca-time-accurate-to-25ps", + 0x1b: "ieee1588-ptp-tt:ca-time-accurate-to-100-ps", + 0x1c: "ieee1588-ptp-tt:ca-time-accurate-to-250-ps", + 0x1d: "ieee1588-ptp-tt:ca-time-accurate-to-1000-ps", + 0x1e: "ieee1588-ptp-tt:ca-time-accurate-to-2500-ps", + 0x1f: "ieee1588-ptp-tt:ca-time-accurate-to-10-ns", + 0x20: "ieee1588-ptp-tt:ca-time-accurate-to-25-ns", + 0x21: "ieee1588-ptp-tt:ca-time-accurate-to-100-ns", + 0x22: "ieee1588-ptp-tt:ca-time-accurate-to-250-ns", + 0x23: "ieee1588-ptp-tt:ca-time-accurate-to-1000-ns", + 0x24: "ieee1588-ptp-tt:ca-time-accurate-to-2500-ns", + 0x25: "ieee1588-ptp-tt:ca-time-accurate-to-10-us", + 0x26: "ieee1588-ptp-tt:ca-time-accurate-to-25-us", + 0x27: "ieee1588-ptp-tt:ca-time-accurate-to-100-us", + 0x28: "ieee1588-ptp-tt:ca-time-accurate-to-250-us", + 0x29: "ieee1588-ptp-tt:ca-time-accurate-to-1000-us", + 0x2a: "ieee1588-ptp-tt:ca-time-accurate-to-2500-us", + 0x2b: "ieee1588-ptp-tt:ca-time-accurate-to-10-ms", + 0x2c: "ieee1588-ptp-tt:ca-time-accurate-to-25-ms", + 0x2d: "ieee1588-ptp-tt:ca-time-accurate-to-100-ms", + 0x2e: "ieee1588-ptp-tt:ca-time-accurate-to-250-ms", + 0x2f: "ieee1588-ptp-tt:ca-time-accurate-to-1-s", + 0x30: "ieee1588-ptp-tt:ca-time-accurate-to-10-s", + 0x31: "ieee1588-ptp-tt:ca-time-accurate-to-gt-10-s", +} + +var timeSourceIdentity = map[uint8]string{ + 0x10: "ieee1588-ptp-tt:atomic-clock", + 0x20: "ieee1588-ptp-tt:gnss", + 0x30: "ieee1588-ptp-tt:terrestrial-radio", + 0x39: "ieee1588-ptp-tt:serial-time-code", + 0x40: "ieee1588-ptp-tt:ptp", + 0x50: "ieee1588-ptp-tt:ntp", + 0x60: "ieee1588-ptp-tt:hand-set", + 0x90: "ieee1588-ptp-tt:other", + 0xa0: "ieee1588-ptp-tt:internal-oscillator", +} + +// portStateName maps IEEE 1588 portState values to YANG enum names. +var portStateName = map[uint8]string{ + 1: "initializing", + 2: "faulty", + 3: "disabled", + 4: "listening", + 5: "pre-time-transmitter", + 6: "time-transmitter", + 7: "passive", + 8: "uncalibrated", + 9: "time-receiver", +} + +func timeSourceName(ts uint8) string { + if s, ok := timeSourceIdentity[ts]; ok { + return s + } + return "ieee1588-ptp-tt:internal-oscillator" +} + +func delayMechanismName(dm uint8) string { + switch dm { + case 2: + return "p2p" + case 0: /* linuxptp DM_AUTO */ + return "no-mechanism" + default: + return "e2e" + } +} + +// fmtClockIdentity renders a clock identity in the YANG clock-identity +// pattern [0-9A-F]{2}(-[0-9A-F]{2}){7}, e.g. "00-51-82-FF-FE-11-22-02". +func fmtClockIdentity(cid ptp.ClockIdentity) string { + b := make([]byte, 0, 23) + for i := 7; i >= 0; i-- { + if len(b) > 0 { + b = append(b, '-') + } + b = append(b, fmt.Sprintf("%02X", uint8(uint64(cid)>>(8*i)))...) + } + return string(b) +} + +func fmtPortIdentity(pid ptp.PortIdentity) map[string]interface{} { + return map[string]interface{}{ + "clock-identity": fmtClockIdentity(pid.ClockIdentity), + "port-number": int(pid.PortNumber), + } +} + +// timeIntervalString renders a raw IEEE 1588 TimeInterval (nanoseconds +// scaled by 2^16) for a YANG time-interval leaf. RFC 7951 requires +// int64 to be JSON-encoded as a string. +func timeIntervalString(ti ptp.TimeInterval) string { + return strconv.FormatInt(int64(ti), 10) +} + +func clockQualityMap(cq ptp.ClockQuality) map[string]interface{} { + out := map[string]interface{}{} + if cc, ok := clockClassIdentity[uint8(cq.ClockClass)]; ok { + out["clock-class"] = cc + } + if ca, ok := clockAccuracyIdentity[uint8(cq.ClockAccuracy)]; ok { + out["clock-accuracy"] = ca + } + out["offset-scaled-log-variance"] = int(cq.OffsetScaledLogVariance) + return out +} + +func defaultDSMap(d *ptp.DefaultDataSetTLV, instanceType string) map[string]interface{} { + ds := map[string]interface{}{ + "clock-identity": fmtClockIdentity(d.ClockIdentity), + "number-ports": int(d.NumberPorts), + "clock-quality": clockQualityMap(d.ClockQuality), + "priority1": int(d.Priority1), + "priority2": int(d.Priority2), + "domain-number": int(d.DomainNumber), + "time-receiver-only": d.SoTSC&flagDefaultDSSOnly != 0, + } + if instanceType != "" { + ds["instance-type"] = instanceType + } + return ds +} + +func currentDSMap(d *ptp.CurrentDataSetTLV) map[string]interface{} { + return map[string]interface{}{ + "steps-removed": int(d.StepsRemoved), + "offset-from-time-transmitter": timeIntervalString(d.OffsetFromMaster), + "mean-delay": timeIntervalString(d.MeanPathDelay), + } +} + +func parentDSMap(d *ptp.ParentDataSetTLV) map[string]interface{} { + return map[string]interface{}{ + "parent-port-identity": fmtPortIdentity(d.ParentPortIdentity), + "parent-stats": d.PS&1 != 0, + "observed-parent-offset-scaled-log-variance": int(d.ObservedParentOffsetScaledLogVariance), + "observed-parent-clock-phase-change-rate": int64(int32(d.ObservedParentClockPhaseChangeRate)), + "grandmaster-identity": fmtClockIdentity(d.GrandmasterIdentity), + "grandmaster-clock-quality": clockQualityMap(d.GrandmasterClockQuality), + "grandmaster-priority1": int(d.GrandmasterPriority1), + "grandmaster-priority2": int(d.GrandmasterPriority2), + } +} + +func timePropertiesDSMap(d *timePropertiesDataSetTLV) map[string]interface{} { + ds := map[string]interface{}{ + "leap61": d.Flags&flagLeap61 != 0, + "leap59": d.Flags&flagLeap59 != 0, + "current-utc-offset-valid": d.Flags&flagUtcOffValid != 0, + "ptp-timescale": d.Flags&flagPtpTimescale != 0, + "time-traceable": d.Flags&flagTimeTraceable != 0, + "frequency-traceable": d.Flags&flagFreqTraceable != 0, + "time-source": timeSourceName(uint8(d.TimeSource)), + } + // current-utc-offset has a when-condition on + // current-utc-offset-valid = 'true' + if d.Flags&flagUtcOffValid != 0 { + ds["current-utc-offset"] = int(d.CurrentUtcOffset) + } + return ds +} + +func portDSMap(d *portDataSetTLV) map[string]interface{} { + state, ok := portStateName[d.PortState] + if !ok { + state = "disabled" + } + return map[string]interface{}{ + "port-identity": fmtPortIdentity(d.PortIdentity), + "port-state": state, + "log-min-delay-req-interval": int(d.LogMinDelayReqInterval), + "mean-link-delay": timeIntervalString(d.PeerMeanPathDelay), + "log-announce-interval": int(d.LogAnnounceInterval), + "announce-receipt-timeout": int(d.AnnounceReceiptTimeout), + "log-sync-interval": int(d.LogSyncInterval), + "delay-mechanism": delayMechanismName(d.DelayMechanism), + "log-min-pdelay-req-interval": int(d.LogMinPdelayReqInterval), + "version-number": int(d.VersionNumber), + } +} + +// IEEE 1588 message type indices in the PORT_STATS_NP counter arrays. +const ( + msgSync = 0x0 + msgPdelayReq = 0x2 + msgPdelayResp = 0x3 + msgFollowUp = 0x8 + msgPdelayRespFollowUp = 0xA + msgAnnounce = 0xB +) + +// portStatsMap maps PORT_STATS_NP counters to the ieee802-dot1as-gptp +// port-statistics-ds counters. +func portStatsMap(s *ptp.PortStatsNPTLV) map[string]interface{} { + rx := s.PortStats.RXMsgType + tx := s.PortStats.TXMsgType + return map[string]interface{}{ + "rx-sync-count": rx[msgSync], + "rx-follow-up-count": rx[msgFollowUp], + "rx-pdelay-req-count": rx[msgPdelayReq], + "rx-pdelay-resp-count": rx[msgPdelayResp], + "rx-pdelay-resp-follow-up-count": rx[msgPdelayRespFollowUp], + "rx-announce-count": rx[msgAnnounce], + "tx-sync-count": tx[msgSync], + "tx-follow-up-count": tx[msgFollowUp], + "tx-pdelay-req-count": tx[msgPdelayReq], + "tx-pdelay-resp-count": tx[msgPdelayResp], + "tx-pdelay-resp-follow-up-count": tx[msgPdelayRespFollowUp], + "tx-announce-count": tx[msgAnnounce], + } +} diff --git a/src/yangerd/internal/ptpmonitor/ptpmonitor.go b/src/yangerd/internal/ptpmonitor/ptpmonitor.go new file mode 100644 index 000000000..d14afde6c --- /dev/null +++ b/src/yangerd/internal/ptpmonitor/ptpmonitor.go @@ -0,0 +1,547 @@ +// Package ptpmonitor keeps the ieee1588-ptp-tt subtree in sync with the +// running ptp4l instances. One ptp4l process runs per instance-index, +// with its config at /etc/linuxptp/ptp4l-.conf and its management +// (UDS) socket at /var/run/ptp4l-. +// +// Instead of polling, each instance connection subscribes to ptp4l's +// push notifications (SUBSCRIBE_EVENTS_NP): port state transitions +// deliver PORT_DATA_SET, every clock update delivers TIME_STATUS_NP, +// and BMCA reselection delivers PARENT_DATA_SET. The near-static data +// sets are fetched at connect and refreshed on a slow timer, which also +// renews the subscription. +package ptpmonitor + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "log/slog" + "net" + "os" + "path/filepath" + "regexp" + "sort" + "strconv" + "strings" + "sync" + "time" + + ptp "github.com/facebook/time/ptp/protocol" + "github.com/kernelkit/infix/src/yangerd/internal/backoff" + "github.com/kernelkit/infix/src/yangerd/internal/tree" + "github.com/kernelkit/infix/src/yangerd/internal/unixgram" +) + +const ( + treeKey = "ieee1588-ptp-tt:ptp" + + // scanInterval is how often the set of configured instances is + // re-discovered. confd (re)writes the config files and + // (re)starts ptp4l on configuration changes. + scanInterval = 10 * time.Second + + // refreshInterval bounds the staleness of the data sets that + // have no notification (default-ds, time-properties-ds, port + // statistics), and paces subscription renewal. + refreshInterval = 30 * time.Second + + // subscribeDuration must comfortably exceed refreshInterval so + // the subscription never lapses between renewals. + subscribeDuration = 180 /* seconds */ +) + +// Overridable in tests. +var ( + confDir = "/etc/linuxptp" + sockDir = "/var/run" +) + +var confRe = regexp.MustCompile(`ptp4l-(\d+)\.conf$`) + +// PTPMonitor supervises one connection per running ptp4l instance and +// publishes the combined operational state. +type PTPMonitor struct { + tree *tree.Tree + log *slog.Logger + + mu sync.Mutex + instances map[uint16]*instance +} + +// New creates a PTPMonitor. +func New(t *tree.Tree, log *slog.Logger) *PTPMonitor { + if log == nil { + log = slog.Default() + } + return &PTPMonitor{ + tree: t, + log: log, + instances: make(map[uint16]*instance), + } +} + +// Run discovers ptp4l instances and supervises their connections until +// ctx is cancelled. +func (m *PTPMonitor) Run(ctx context.Context) error { + tick := time.NewTicker(scanInterval) + defer tick.Stop() + + for { + m.scan(ctx) + select { + case <-ctx.Done(): + m.mu.Lock() + insts := make([]*instance, 0, len(m.instances)) + for _, inst := range m.instances { + inst.stop() + insts = append(insts, inst) + } + m.mu.Unlock() + for _, inst := range insts { + <-inst.done + } + return ctx.Err() + case <-tick.C: + } + } +} + +// scan reconciles the running instance connections with the configured +// ptp4l instances. +func (m *PTPMonitor) scan(ctx context.Context) { + confs, _ := filepath.Glob(confDir + "/ptp4l-*.conf") + + found := make(map[uint16]string) + for _, conf := range confs { + match := confRe.FindStringSubmatch(conf) + if match == nil { + continue + } + var idx uint16 + fmt.Sscanf(match[1], "%d", &idx) + found[idx] = conf + } + + m.mu.Lock() + defer m.mu.Unlock() + + for idx, inst := range m.instances { + if _, ok := found[idx]; !ok { + inst.stop() + delete(m.instances, idx) + } + } + + changed := false + for idx, conf := range found { + if _, ok := m.instances[idx]; ok { + continue + } + inst := newInstance(idx, conf, m.log, m.publish) + ictx, cancel := context.WithCancel(ctx) + inst.cancel = cancel + m.instances[idx] = inst + go inst.run(ictx) + changed = true + } + + if changed || len(found) < len(m.instances) { + m.publishLocked() + } +} + +// publish rebuilds the whole subtree from the current instance states. +func (m *PTPMonitor) publish() { + m.mu.Lock() + defer m.mu.Unlock() + m.publishLocked() +} + +func (m *PTPMonitor) publishLocked() { + var idxs []int + for idx := range m.instances { + idxs = append(idxs, int(idx)) + } + sort.Ints(idxs) + + var list []interface{} + for _, idx := range idxs { + if data := m.instances[uint16(idx)].snapshot(); data != nil { + list = append(list, data) + } + } + + if len(list) == 0 { + m.tree.Delete(treeKey) + return + } + + data, err := json.Marshal(map[string]interface{}{ + "instances": map[string]interface{}{ + "instance": list, + }, + }) + if err != nil { + return + } + m.tree.Set(treeKey, data) +} + +// instance is one supervised ptp4l management connection. +type instance struct { + idx uint16 + conf string + sock string + log *slog.Logger + publish func() + cancel context.CancelFunc + done chan struct{} // closed when run returns + + mu sync.Mutex + state *instanceState +} + +// instanceState holds the last known data sets for one instance. +type instanceState struct { + instanceType string + ifaces []string + + defaultDS *ptp.DefaultDataSetTLV + currentDS *ptp.CurrentDataSetTLV + parentDS *ptp.ParentDataSetTLV + timeProps *timePropertiesDataSetTLV + ports map[uint16]*portDataSetTLV + portStats map[uint16]*ptp.PortStatsNPTLV +} + +func newInstance(idx uint16, conf string, log *slog.Logger, publish func()) *instance { + return &instance{ + idx: idx, + conf: conf, + sock: fmt.Sprintf("%s/ptp4l-%d", sockDir, idx), + log: log.With("ptp-instance", idx), + publish: publish, + done: make(chan struct{}), + } +} + +// stop ends the supervisor. cancel is set under the monitor lock +// before run starts, and only read under it. +func (i *instance) stop() { + if i.cancel != nil { + i.cancel() + } +} + +// snapshot renders the current state as a YANG instance list entry, or +// nil when the instance has no data (ptp4l not answering). +func (i *instance) snapshot() map[string]interface{} { + i.mu.Lock() + defer i.mu.Unlock() + + s := i.state + if s == nil || s.defaultDS == nil { + return nil + } + + inst := map[string]interface{}{ + "instance-index": int(i.idx), + "default-ds": defaultDSMap(s.defaultDS, s.instanceType), + } + if s.currentDS != nil { + inst["current-ds"] = currentDSMap(s.currentDS) + } + if s.parentDS != nil { + inst["parent-ds"] = parentDSMap(s.parentDS) + } + if s.timeProps != nil { + inst["time-properties-ds"] = timePropertiesDSMap(s.timeProps) + } + + var nums []int + for num := range s.ports { + nums = append(nums, int(num)) + } + sort.Ints(nums) + + var ports []interface{} + for _, num := range nums { + pd := s.ports[uint16(num)] + entry := map[string]interface{}{ + "port-index": num, + "port-ds": portDSMap(pd), + } + if num >= 1 && num <= len(s.ifaces) { + entry["underlying-interface"] = s.ifaces[num-1] + } + if stats, ok := s.portStats[uint16(num)]; ok { + entry["ieee802-dot1as-gptp:port-statistics-ds"] = portStatsMap(stats) + } + ports = append(ports, entry) + } + if len(ports) > 0 { + inst["ports"] = map[string]interface{}{"port": ports} + } + + return inst +} + +// run supervises the connection: connect, fill, subscribe, consume +// events, and reconnect with backoff when ptp4l goes away. +func (i *instance) run(ctx context.Context) { + defer close(i.done) + + name := fmt.Sprintf("ptp instance %d", i.idx) + backoff.Retry(ctx, i.log, name, func(ctx context.Context) error { + err := i.session(ctx) + + // Drop stale state so a dead ptp4l disappears from + // operational instead of lingering. + i.mu.Lock() + i.state = nil + i.mu.Unlock() + i.publish() + + return err + }) +} + +// session runs one connected episode against ptp4l. +func (i *instance) session(ctx context.Context) error { + // ptp4l runs unprivileged and must be able to send replies here + local := fmt.Sprintf("%s/yangerd-ptp-%d.sock", sockDir, i.idx) + conn, err := unixgram.Dial(local, i.sock, 0666) + if err != nil { + return err + } + defer conn.Close() + + // Terminate the blocking read loop on shutdown + done := make(chan struct{}) + defer close(done) + go func() { + select { + case <-ctx.Done(): + conn.SetReadDeadline(time.Now()) + case <-done: + } + }() + + // gPTP (802.1AS) instances only accept management messages + // carrying their transport-specific nibble + sdoID := transportSpecificFromConf(i.conf) + + i.mu.Lock() + i.state = &instanceState{ + instanceType: instanceTypeFromConf(i.conf), + ifaces: portInterfaces(i.conf), + ports: make(map[uint16]*portDataSetTLV), + portStats: make(map[uint16]*ptp.PortStatsNPTLV), + } + i.mu.Unlock() + + var seq uint16 + send := func(pkt *ptp.Management) error { + seq++ + pkt.SetSequence(seq) + b, err := pkt.MarshalBinary() + if err != nil { + return err + } + _, err = conn.Write(b) + return err + } + + refresh := func() error { + pkts := getRequests(sdoID) + pkts = append(pkts, subscribeRequest(sdoID, subscribeDuration, + notifyPortState, notifyTimeSync, notifyParentDataSet)) + for _, pkt := range pkts { + if err := send(pkt); err != nil { + return err + } + } + return nil + } + + if err := refresh(); err != nil { + return err + } + + buf := make([]byte, 2048) + lastData := time.Now() + lastRefresh := time.Now() + for { + // Anchor the deadline to the last refresh, not the last + // read: a syncing instance pushes TIME_SYNC events many + // times per second, which must not starve the refresh. + conn.SetReadDeadline(lastRefresh.Add(refreshInterval)) + n, err := conn.Read(buf) + if ctx.Err() != nil { + return ctx.Err() + } + if err != nil { + var nerr net.Error + if errors.As(err, &nerr) && nerr.Timeout() { + // Give up if ptp4l has not answered + // anything for a while + if time.Since(lastData) > 3*refreshInterval { + return fmt.Errorf("ptp4l not responding") + } + if err := refresh(); err != nil { + return err + } + lastRefresh = time.Now() + continue + } + return err + } + + tlv, err := decodePacket(buf[:n]) + if err != nil { + var merr *mgmtError + if errors.As(err, &merr) { + // Expected for data sets an instance type + // does not implement (e.g. TCs have no + // current/parent DS) + i.log.Debug("ptp: management error", "err", merr) + lastData = time.Now() + continue + } + i.log.Debug("ptp: decode failed", "err", err) + continue + } + + lastData = time.Now() + if i.apply(tlv) { + i.publish() + } + + // steps-removed lives in current-ds but changes with the + // parent selection: fetch it right away instead of + // letting it wait for the slow refresh + if _, ok := tlv.(*ptp.ParentDataSetTLV); ok { + if err := send(getRequest(sdoID, ptp.IDCurrentDataSet)); err != nil { + return err + } + } + + if time.Since(lastRefresh) >= refreshInterval { + if err := refresh(); err != nil { + return err + } + lastRefresh = time.Now() + } + } +} + +// apply folds a received TLV into the instance state. Returns true +// when the state changed in a way worth publishing. +func (i *instance) apply(tlv ptp.ManagementTLV) bool { + i.mu.Lock() + defer i.mu.Unlock() + + s := i.state + if s == nil { + return false + } + + switch t := tlv.(type) { + case *ptp.DefaultDataSetTLV: + if s.instanceType == "" { + if t.NumberPorts > 1 { + s.instanceType = "bc" + } else { + s.instanceType = "oc" + } + } + s.defaultDS = t + case *ptp.CurrentDataSetTLV: + s.currentDS = t + case *ptp.ParentDataSetTLV: + s.parentDS = t + case *timePropertiesDataSetTLV: + s.timeProps = t + case *portDataSetTLV: + s.ports[t.PortIdentity.PortNumber] = t + case *ptp.PortStatsNPTLV: + s.portStats[t.PortIdentity.PortNumber] = t + case *ptp.TimeStatusNPTLV: + // Pushed on every clock update: keep the current-ds + // offset live between CURRENT_DATA_SET refreshes + if s.currentDS != nil { + s.currentDS.OffsetFromMaster = ptp.TimeInterval(t.MasterOffsetNS << 16) + } + case *subscribeEventsNPTLV: + // Acknowledgement of our subscription + return false + default: + return false + } + + return true +} + +// portInterfaces returns the ordered interface names from a ptp4l conf +// (non-global section headers). +func portInterfaces(conf string) []string { + data, err := os.ReadFile(conf) + if err != nil { + return nil + } + + var ifaces []string + for _, line := range strings.Split(string(data), "\n") { + s := strings.TrimSpace(line) + if strings.HasPrefix(s, "[") && strings.HasSuffix(s, "]") && s != "[global]" { + ifaces = append(ifaces, s[1:len(s)-1]) + } + } + return ifaces +} + +// transportSpecificFromConf reads the transportSpecific keyword from a +// ptp4l conf. confd writes 1 for the ieee802-dot1as profile and 0 +// otherwise. +func transportSpecificFromConf(conf string) uint8 { + data, err := os.ReadFile(conf) + if err != nil { + return 0 + } + + for _, line := range strings.Split(string(data), "\n") { + fields := strings.Fields(line) + if len(fields) == 2 && fields[0] == "transportSpecific" { + v, err := strconv.ParseUint(strings.TrimPrefix(fields[1], "0x"), 16, 4) + if err == nil { + return uint8(v) + } + } + } + return 0 +} + +// instanceTypeFromConf reads the clockType keyword from a ptp4l conf. +// Returns "" when not present (OC/BC derived from numberPorts instead). +func instanceTypeFromConf(conf string) string { + data, err := os.ReadFile(conf) + if err != nil { + return "" + } + + for _, line := range strings.Split(string(data), "\n") { + fields := strings.Fields(line) + if len(fields) == 2 && fields[0] == "clockType" { + switch strings.ToUpper(fields[1]) { + case "P2P_TC": + return "p2p-tc" + case "E2E_TC": + return "e2e-tc" + case "BOUNDARY_CLOCK": + return "bc" + } + } + } + return "" +} diff --git a/src/yangerd/internal/ptpmonitor/ptpmonitor_test.go b/src/yangerd/internal/ptpmonitor/ptpmonitor_test.go new file mode 100644 index 000000000..8f2564d0d --- /dev/null +++ b/src/yangerd/internal/ptpmonitor/ptpmonitor_test.go @@ -0,0 +1,484 @@ +package ptpmonitor + +import ( + "context" + "encoding/json" + "fmt" + "log/slog" + "net" + "os" + "path/filepath" + "testing" + "time" + + ptp "github.com/facebook/time/ptp/protocol" + "github.com/kernelkit/infix/src/yangerd/internal/tree" +) + +// reply wraps a response TLV in a management packet and marshals it to +// wire format, the way ptp4l would. +func reply(t *testing.T, tlv ptp.ManagementTLV) []byte { + t.Helper() + + headerSize := uint16(48) /* binary.Size(ManagementMsgHead{}) */ + pkt := &ptp.Management{ + ManagementMsgHead: ptp.ManagementMsgHead{ + Header: ptp.Header{ + SdoIDAndMsgType: ptp.NewSdoIDAndMsgType(ptp.MessageManagement, 0), + Version: ptp.Version, + MessageLength: headerSize, + LogMessageInterval: ptp.MgmtLogMessageInterval, + }, + TargetPortIdentity: ptp.DefaultTargetPortIdentity, + ActionField: ptp.RESPONSE, + }, + TLV: tlv, + } + b, err := pkt.MarshalBinary() + if err != nil { + t.Fatalf("marshal reply: %v", err) + } + return b +} + +func defaultDSFixture() *ptp.DefaultDataSetTLV { + size := uint16(24 + 4) /* payload + TLV head, unused by decode */ + return &ptp.DefaultDataSetTLV{ + ManagementTLVHead: ptp.ManagementTLVHead{ + TLVHead: ptp.TLVHead{TLVType: ptp.TLVManagement, LengthField: size}, + ManagementID: ptp.IDDefaultDataSet, + }, + SoTSC: 2, /* slaveOnly */ + NumberPorts: 1, + Priority1: 128, + Priority2: 127, + ClockQuality: ptp.ClockQuality{ + ClockClass: ptp.ClockClass(248), + ClockAccuracy: ptp.ClockAccuracy(0x21), + OffsetScaledLogVariance: 0xfffe, + }, + ClockIdentity: 0x005182FFFE112202, + DomainNumber: 0, + } +} + +func TestDecodeRoundTrip(t *testing.T) { + tests := []struct { + name string + tlv ptp.ManagementTLV + }{ + {"default-ds", defaultDSFixture()}, + { + "current-ds", + &ptp.CurrentDataSetTLV{ + ManagementTLVHead: ptp.ManagementTLVHead{ + TLVHead: ptp.TLVHead{TLVType: ptp.TLVManagement}, + ManagementID: ptp.IDCurrentDataSet, + }, + StepsRemoved: 1, + OffsetFromMaster: ptp.TimeInterval(42 << 16), + MeanPathDelay: ptp.TimeInterval(1000 << 16), + }, + }, + { + "port-ds", + &portDataSetTLV{ + ManagementTLVHead: ptp.ManagementTLVHead{ + TLVHead: ptp.TLVHead{TLVType: ptp.TLVManagement}, + ManagementID: ptp.IDPortDataSet, + }, + PortIdentity: ptp.PortIdentity{ + ClockIdentity: 0x005182FFFE112202, + PortNumber: 1, + }, + PortState: 9, /* SLAVE */ + DelayMechanism: 1, /* E2E */ + VersionNumber: 2, + }, + }, + { + "time-properties-ds", + &timePropertiesDataSetTLV{ + ManagementTLVHead: ptp.ManagementTLVHead{ + TLVHead: ptp.TLVHead{TLVType: ptp.TLVManagement}, + ManagementID: ptp.IDTimePropertiesDataSet, + }, + CurrentUtcOffset: 37, + Flags: flagUtcOffValid | flagPtpTimescale, + TimeSource: ptp.TimeSource(0xa0), + }, + }, + { + "time-status-np", + &ptp.TimeStatusNPTLV{ + ManagementTLVHead: ptp.ManagementTLVHead{ + TLVHead: ptp.TLVHead{TLVType: ptp.TLVManagement}, + ManagementID: ptp.IDTimeStatusNP, + }, + MasterOffsetNS: -1234, + GMPresent: 1, + GMIdentity: 0x005182FFFE112202, + }, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + got, err := decodePacket(reply(t, tc.tlv)) + if err != nil { + t.Fatalf("decodePacket: %v", err) + } + if got.MgmtID() != tc.tlv.MgmtID() { + t.Fatalf("management ID: got 0x%04x, want 0x%04x", + uint16(got.MgmtID()), uint16(tc.tlv.MgmtID())) + } + }) + } +} + +func TestDecodePortStatsEndianness(t *testing.T) { + tlv := &ptp.PortStatsNPTLV{ + ManagementTLVHead: ptp.ManagementTLVHead{ + TLVHead: ptp.TLVHead{TLVType: ptp.TLVManagement}, + ManagementID: ptp.IDPortStatsNP, + }, + PortIdentity: ptp.PortIdentity{ClockIdentity: 1, PortNumber: 1}, + } + tlv.PortStats.RXMsgType[msgSync] = 42 + tlv.PortStats.TXMsgType[msgAnnounce] = 7 + + // PortStatsNPTLV.MarshalBinary writes counters in host byte + // order, exactly like ptp4l + got, err := decodePacket(reply(t, tlv)) + if err != nil { + t.Fatalf("decodePacket: %v", err) + } + stats, ok := got.(*ptp.PortStatsNPTLV) + if !ok { + t.Fatalf("got %T", got) + } + if stats.PortStats.RXMsgType[msgSync] != 42 || stats.PortStats.TXMsgType[msgAnnounce] != 7 { + t.Fatalf("counters mangled: rx=%d tx=%d", + stats.PortStats.RXMsgType[msgSync], stats.PortStats.TXMsgType[msgAnnounce]) + } +} + +func TestFmtClockIdentity(t *testing.T) { + got := fmtClockIdentity(ptp.ClockIdentity(0x005182FFFE112202)) + want := "00-51-82-FF-FE-11-22-02" + if got != want { + t.Fatalf("fmtClockIdentity: got %q, want %q", got, want) + } +} + +func TestCurrentDSMapRFC7951(t *testing.T) { + ds := currentDSMap(&ptp.CurrentDataSetTLV{ + StepsRemoved: 1, + OffsetFromMaster: ptp.TimeInterval(-42 << 16), + MeanPathDelay: ptp.TimeInterval(1000 << 16), + }) + + // int64 leafs must be JSON strings (RFC 7951) + if ds["offset-from-time-transmitter"] != fmt.Sprint(-42<<16) { + t.Fatalf("offset: got %v", ds["offset-from-time-transmitter"]) + } + if ds["mean-delay"] != fmt.Sprint(1000<<16) { + t.Fatalf("mean-delay: got %v", ds["mean-delay"]) + } + if ds["steps-removed"] != 1 { + t.Fatalf("steps-removed: got %v", ds["steps-removed"]) + } +} + +func TestPortDSMapStates(t *testing.T) { + tests := []struct { + state uint8 + mech uint8 + want string + mechS string + }{ + {6, 1, "time-transmitter", "e2e"}, + {9, 2, "time-receiver", "p2p"}, + {4, 0, "listening", "no-mechanism"}, + {99, 1, "disabled", "e2e"}, /* unknown state */ + } + + for _, tc := range tests { + ds := portDSMap(&portDataSetTLV{PortState: tc.state, DelayMechanism: tc.mech}) + if ds["port-state"] != tc.want { + t.Fatalf("state %d: got %v, want %v", tc.state, ds["port-state"], tc.want) + } + if ds["delay-mechanism"] != tc.mechS { + t.Fatalf("mech %d: got %v, want %v", tc.mech, ds["delay-mechanism"], tc.mechS) + } + } +} + +func TestTimePropertiesUtcOffsetWhenCondition(t *testing.T) { + // current-utc-offset must be omitted unless + // current-utc-offset-valid is true (YANG when-condition) + ds := timePropertiesDSMap(&timePropertiesDataSetTLV{CurrentUtcOffset: 37}) + if _, ok := ds["current-utc-offset"]; ok { + t.Fatal("current-utc-offset present despite valid=false") + } + + ds = timePropertiesDSMap(&timePropertiesDataSetTLV{ + CurrentUtcOffset: 37, + Flags: flagUtcOffValid, + }) + if ds["current-utc-offset"] != 37 { + t.Fatalf("current-utc-offset: got %v", ds["current-utc-offset"]) + } +} + +func TestInstanceTypeFromConf(t *testing.T) { + dir := t.TempDir() + + tests := []struct { + conf string + want string + }{ + {"[global]\nclockType BOUNDARY_CLOCK\n[e1]\n[e2]\n", "bc"}, + {"[global]\nclockType E2E_TC\n", "e2e-tc"}, + {"[global]\nclockType P2P_TC\n", "p2p-tc"}, + {"[global]\nuds_address /var/run/ptp4l-0\n[e1]\n", ""}, + } + + for i, tc := range tests { + path := filepath.Join(dir, fmt.Sprintf("ptp4l-%d.conf", i)) + os.WriteFile(path, []byte(tc.conf), 0644) + if got := instanceTypeFromConf(path); got != tc.want { + t.Fatalf("conf %d: got %q, want %q", i, got, tc.want) + } + } +} + +func TestPortInterfaces(t *testing.T) { + dir := t.TempDir() + path := filepath.Join(dir, "ptp4l-0.conf") + os.WriteFile(path, []byte("[global]\nfoo 1\n[e1]\n[e2]\nbar 2\n"), 0644) + + got := portInterfaces(path) + if len(got) != 2 || got[0] != "e1" || got[1] != "e2" { + t.Fatalf("portInterfaces: got %v", got) + } +} + +// fakePtp4l answers management requests with canned replies, mimicking +// a running ptp4l instance. Like the real thing, it silently drops +// requests whose transport-specific nibble does not match sdo. +func fakePtp4l(t *testing.T, sock string, sdo uint8) { + t.Helper() + + addr := &net.UnixAddr{Name: sock, Net: "unixgram"} + conn, err := net.ListenUnixgram("unixgram", addr) + if err != nil { + t.Fatalf("fake ptp4l listen: %v", err) + } + t.Cleanup(func() { conn.Close() }) + + respond := func(raddr *net.UnixAddr, resp ptp.ManagementTLV) { + pkt := &ptp.Management{ + ManagementMsgHead: ptp.ManagementMsgHead{ + Header: ptp.Header{ + SdoIDAndMsgType: ptp.NewSdoIDAndMsgType(ptp.MessageManagement, 0), + Version: ptp.Version, + LogMessageInterval: ptp.MgmtLogMessageInterval, + }, + TargetPortIdentity: ptp.DefaultTargetPortIdentity, + ActionField: ptp.RESPONSE, + }, + TLV: resp, + } + if b, err := pkt.MarshalBinary(); err == nil { + conn.WriteToUnix(b, raddr) + } + } + + go func() { + buf := make([]byte, 2048) + var stepsRemoved uint16 = 1 + + for { + n, raddr, err := conn.ReadFromUnix(buf) + if err != nil { + return + } + if n < 1 || buf[0]>>4 != sdo { + continue + } + req, err := decodePacket(buf[:n]) + if err != nil { + continue + } + + var resp ptp.ManagementTLV + switch req.MgmtID() { + case ptp.IDDefaultDataSet: + resp = defaultDSFixture() + case ptp.IDCurrentDataSet: + resp = &ptp.CurrentDataSetTLV{ + ManagementTLVHead: ptp.ManagementTLVHead{ + TLVHead: ptp.TLVHead{TLVType: ptp.TLVManagement}, + ManagementID: ptp.IDCurrentDataSet, + }, + StepsRemoved: stepsRemoved, + OffsetFromMaster: ptp.TimeInterval(42 << 16), + } + case ptp.IDPortDataSet: + resp = &portDataSetTLV{ + ManagementTLVHead: ptp.ManagementTLVHead{ + TLVHead: ptp.TLVHead{TLVType: ptp.TLVManagement}, + ManagementID: ptp.IDPortDataSet, + }, + PortIdentity: ptp.PortIdentity{ + ClockIdentity: 0x005182FFFE112202, + PortNumber: 1, + }, + PortState: 9, + DelayMechanism: 1, + VersionNumber: 2, + } + case idSubscribeEventsNP: + resp = &subscribeEventsNPTLV{ + ManagementTLVHead: ptp.ManagementTLVHead{ + TLVHead: ptp.TLVHead{TLVType: ptp.TLVManagement}, + ManagementID: idSubscribeEventsNP, + }, + } + default: + continue + } + respond(raddr, resp) + + // After the subscription lands, simulate a BMCA + // parent change: steps-removed bumps and ptp4l + // pushes a PARENT_DATA_SET notification. The + // monitor must react by re-fetching current-ds. + if req.MgmtID() == idSubscribeEventsNP && stepsRemoved == 1 { + stepsRemoved = 2 + respond(raddr, &ptp.ParentDataSetTLV{ + ManagementTLVHead: ptp.ManagementTLVHead{ + TLVHead: ptp.TLVHead{TLVType: ptp.TLVManagement}, + ManagementID: ptp.IDParentDataSet, + }, + GrandmasterIdentity: 0x005182FFFE112202, + }) + } + } + }() +} + +func TestTransportSpecificFromConf(t *testing.T) { + dir := t.TempDir() + + tests := []struct { + conf string + want uint8 + }{ + {"[global]\ntransportSpecific 1\n[e1]\n", 1}, + {"[global]\ntransportSpecific 0x1\n[e1]\n", 1}, + {"[global]\ntransportSpecific 0\n[e1]\n", 0}, + {"[global]\n[e1]\n", 0}, + } + + for i, tc := range tests { + path := filepath.Join(dir, fmt.Sprintf("ptp4l-%d.conf", i)) + os.WriteFile(path, []byte(tc.conf), 0644) + if got := transportSpecificFromConf(path); got != tc.want { + t.Fatalf("conf %d: got %d, want %d", i, got, tc.want) + } + } +} + +func monitorEndToEnd(t *testing.T, transportSpecific string, sdo uint8) { + dir := t.TempDir() + oldConf, oldSock := confDir, sockDir + confDir, sockDir = dir, dir + t.Cleanup(func() { confDir, sockDir = oldConf, oldSock }) + + os.WriteFile(filepath.Join(dir, "ptp4l-0.conf"), + []byte("[global]\ntransportSpecific "+transportSpecific+ + "\nuds_address "+dir+"/ptp4l-0\n[e1]\n"), 0644) + fakePtp4l(t, filepath.Join(dir, "ptp4l-0"), sdo) + + tr := tree.New() + m := New(tr, slog.Default()) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + go m.Run(ctx) + + deadline := time.Now().Add(5 * time.Second) + for time.Now().Before(deadline) { + raw := tr.Get(treeKey) + if raw == nil { + time.Sleep(50 * time.Millisecond) + continue + } + + var out struct { + Instances struct { + Instance []struct { + InstanceIndex int `json:"instance-index"` + DefaultDS struct { + ClockIdentity string `json:"clock-identity"` + InstanceType string `json:"instance-type"` + } `json:"default-ds"` + CurrentDS struct { + Offset string `json:"offset-from-time-transmitter"` + StepsRemoved int `json:"steps-removed"` + } `json:"current-ds"` + Ports struct { + Port []struct { + PortIndex int `json:"port-index"` + Underlying string `json:"underlying-interface"` + PortDS struct { + PortState string `json:"port-state"` + } `json:"port-ds"` + } `json:"port"` + } `json:"ports"` + } `json:"instance"` + } `json:"instances"` + } + if err := json.Unmarshal(raw, &out); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + insts := out.Instances.Instance + if len(insts) == 1 && + insts[0].DefaultDS.ClockIdentity == "00-51-82-FF-FE-11-22-02" && + len(insts[0].Ports.Port) == 1 && + // Converges to 2 only via the PARENT_DATA_SET + // notification triggering a current-ds re-fetch + insts[0].CurrentDS.StepsRemoved == 2 { + inst := insts[0] + if inst.DefaultDS.InstanceType != "oc" { + t.Fatalf("instance-type: got %q", inst.DefaultDS.InstanceType) + } + if inst.CurrentDS.Offset != fmt.Sprint(42<<16) { + t.Fatalf("offset: got %q", inst.CurrentDS.Offset) + } + port := inst.Ports.Port[0] + if port.PortDS.PortState != "time-receiver" { + t.Fatalf("port-state: got %q", port.PortDS.PortState) + } + if port.Underlying != "e1" { + t.Fatalf("underlying-interface: got %q", port.Underlying) + } + return + } + time.Sleep(50 * time.Millisecond) + } + t.Fatalf("tree never converged, last: %s", tr.Get(treeKey)) +} + +func TestMonitorEndToEnd(t *testing.T) { + monitorEndToEnd(t, "0", 0) +} + +// gPTP instances only answer management messages carrying the 0x1 +// transport-specific nibble. +func TestMonitorEndToEndGPTP(t *testing.T) { + monitorEndToEnd(t, "1", 1) +} diff --git a/src/yangerd/internal/ptpmonitor/scan_test.go b/src/yangerd/internal/ptpmonitor/scan_test.go new file mode 100644 index 000000000..8ce6fa58a --- /dev/null +++ b/src/yangerd/internal/ptpmonitor/scan_test.go @@ -0,0 +1,53 @@ +package ptpmonitor + +import ( + "context" + "log/slog" + "os" + "path/filepath" + "testing" + + "github.com/kernelkit/infix/src/yangerd/internal/tree" +) + +// An instance removed and re-added across scans is stopped and started +// cleanly: every start has its cancel func before run sees it. +func TestScanStopStart(t *testing.T) { + dir := t.TempDir() + oldConf, oldSock := confDir, sockDir + confDir, sockDir = dir, dir + t.Cleanup(func() { confDir, sockDir = oldConf, oldSock }) + + conf := filepath.Join(dir, "ptp4l-0.conf") + m := New(tree.New(), slog.Default()) + ctx, cancel := context.WithCancel(context.Background()) + var started []*instance + defer func() { + cancel() + for _, inst := range started { + <-inst.done + } + }() + + for i := 0; i < 20; i++ { + os.WriteFile(conf, []byte("[global]\n[e1]\n"), 0644) + m.scan(ctx) + m.mu.Lock() + inst := m.instances[0] + m.mu.Unlock() + if inst == nil || inst.cancel == nil { + t.Fatal("instance started without a cancel func") + } + started = append(started, inst) + + os.Remove(conf) + m.scan(ctx) + } + + m.mu.Lock() + n := len(m.instances) + m.mu.Unlock() + if n != 0 { + t.Fatalf("instances left after removal: %d", n) + } +} diff --git a/src/yangerd/internal/ptpmonitor/tlvs.go b/src/yangerd/internal/ptpmonitor/tlvs.go new file mode 100644 index 000000000..e0845f28f --- /dev/null +++ b/src/yangerd/internal/ptpmonitor/tlvs.go @@ -0,0 +1,258 @@ +package ptpmonitor + +import ( + "bytes" + "encoding/binary" + "fmt" + + ptp "github.com/facebook/time/ptp/protocol" +) + +// Management IDs not defined by the library. +const idSubscribeEventsNP ptp.ManagementID = 0xC003 + +// Event bit numbers in subscribe_events_np.bitmask (linuxptp +// notification.h). +const ( + notifyPortState = iota + notifyTimeSync + notifyParentDataSet +) + +// portDataSetTLV mirrors linuxptp struct portDS (PORT_DATA_SET, +// IEEE 1588 Table 95). Not implemented by the library. +type portDataSetTLV struct { + ptp.ManagementTLVHead + + PortIdentity ptp.PortIdentity + PortState uint8 + LogMinDelayReqInterval int8 + PeerMeanPathDelay ptp.TimeInterval + LogAnnounceInterval int8 + AnnounceReceiptTimeout uint8 + LogSyncInterval int8 + DelayMechanism uint8 + LogMinPdelayReqInterval int8 + VersionNumber uint8 +} + +// timePropertiesDataSetTLV mirrors linuxptp struct timePropertiesDS +// (TIME_PROPERTIES_DATA_SET, IEEE 1588 Table 92). The flags byte packs +// the Announce flag-field bits (msg.h): LEAP_61, LEAP_59, UTC_OFF_VALID, +// PTP_TIMESCALE, TIME_TRACEABLE, FREQ_TRACEABLE. +type timePropertiesDataSetTLV struct { + ptp.ManagementTLVHead + + CurrentUtcOffset int16 + Flags uint8 + TimeSource ptp.TimeSource +} + +const ( + flagLeap61 = 1 << 0 + flagLeap59 = 1 << 1 + flagUtcOffValid = 1 << 2 + flagPtpTimescale = 1 << 3 + flagTimeTraceable = 1 << 4 + flagFreqTraceable = 1 << 5 + flagDefaultDSSOnly = 1 << 1 /* DDS_SLAVE_ONLY in DEFAULT_DATA_SET flags */ +) + +// subscribeEventsNPTLV mirrors linuxptp struct subscribe_events_np +// (SUBSCRIBE_EVENTS_NP, linuxptp tlv.h). +type subscribeEventsNPTLV struct { + ptp.ManagementTLVHead + + Duration uint16 + Bitmask [64]uint8 +} + +// tlvHead builds a management TLV head for the given ID and full TLV +// struct size. +func tlvHead(id ptp.ManagementID, size uint16) ptp.ManagementTLVHead { + return ptp.ManagementTLVHead{ + TLVHead: ptp.TLVHead{ + TLVType: ptp.TLVManagement, + LengthField: size - uint16(binary.Size(ptp.TLVHead{})), + }, + ManagementID: id, + } +} + +// request wraps a TLV in a management message. sdoID carries the +// transport-specific nibble: ptp4l silently drops management messages +// whose nibble does not match its transportSpecific setting, so gPTP +// (802.1AS) instances require sdoID 1. +func request(sdoID uint8, action ptp.Action, tlv ptp.ManagementTLV, size uint16) *ptp.Management { + headerSize := uint16(binary.Size(ptp.ManagementMsgHead{})) + + return &ptp.Management{ + ManagementMsgHead: ptp.ManagementMsgHead{ + Header: ptp.Header{ + SdoIDAndMsgType: ptp.NewSdoIDAndMsgType(ptp.MessageManagement, sdoID), + Version: ptp.Version, + MessageLength: headerSize + size, + LogMessageInterval: ptp.MgmtLogMessageInterval, + }, + TargetPortIdentity: ptp.DefaultTargetPortIdentity, + ActionField: action, + }, + TLV: tlv, + } +} + +// getRequest builds a GET request for one data set. The payload is a +// zero-padded full TLV struct, mirroring the library's own request +// builders. +func getRequest(sdoID uint8, id ptp.ManagementID) *ptp.Management { + var tlv ptp.ManagementTLV + var size uint16 + + switch id { + case ptp.IDDefaultDataSet: + size = uint16(binary.Size(ptp.DefaultDataSetTLV{})) + tlv = &ptp.DefaultDataSetTLV{ManagementTLVHead: tlvHead(id, size)} + case ptp.IDCurrentDataSet: + size = uint16(binary.Size(ptp.CurrentDataSetTLV{})) + tlv = &ptp.CurrentDataSetTLV{ManagementTLVHead: tlvHead(id, size)} + case ptp.IDParentDataSet: + size = uint16(binary.Size(ptp.ParentDataSetTLV{})) + tlv = &ptp.ParentDataSetTLV{ManagementTLVHead: tlvHead(id, size)} + case ptp.IDTimePropertiesDataSet: + size = uint16(binary.Size(timePropertiesDataSetTLV{})) + tlv = &timePropertiesDataSetTLV{ManagementTLVHead: tlvHead(id, size)} + case ptp.IDPortDataSet: + size = uint16(binary.Size(portDataSetTLV{})) + tlv = &portDataSetTLV{ManagementTLVHead: tlvHead(id, size)} + default: + // Send just the TLV head, like pmc does + size = uint16(binary.Size(ptp.ManagementTLVHead{})) + h := tlvHead(id, size) + tlv = &h + } + + return request(sdoID, ptp.GET, tlv, size) +} + +// getRequests builds the GET requests for all data sets the monitor +// tracks. +func getRequests(sdoID uint8) []*ptp.Management { + var reqs []*ptp.Management + + for _, id := range []ptp.ManagementID{ + ptp.IDDefaultDataSet, + ptp.IDCurrentDataSet, + ptp.IDParentDataSet, + ptp.IDTimePropertiesDataSet, + ptp.IDPortDataSet, + ptp.IDPortStatsNP, + } { + reqs = append(reqs, getRequest(sdoID, id)) + } + + return reqs +} + +// subscribeRequest builds a SET SUBSCRIBE_EVENTS_NP request asking +// ptp4l to push notifications for the given event bits for duration +// seconds. +func subscribeRequest(sdoID uint8, duration uint16, events ...int) *ptp.Management { + size := uint16(binary.Size(subscribeEventsNPTLV{})) + + tlv := &subscribeEventsNPTLV{ + ManagementTLVHead: tlvHead(idSubscribeEventsNP, size), + Duration: duration, + } + for _, ev := range events { + tlv.Bitmask[ev/8] |= 1 << (ev % 8) + } + + return request(sdoID, ptp.SET, tlv, size) +} + +// mgmtError is a MANAGEMENT_ERROR_STATUS response. +type mgmtError struct { + id ptp.ManagementID + err ptp.ManagementErrorID +} + +func (e *mgmtError) Error() string { + return fmt.Sprintf("management error for 0x%04x: %v", uint16(e.id), e.err) +} + +// decodePacket parses a management message from ptp4l into one of the +// supported TLVs. The library's own decoder is not extensible (its TLV +// registry is package-private), so all TLVs are decoded here. +func decodePacket(data []byte) (ptp.ManagementTLV, error) { + var head ptp.ManagementMsgHead + var tlvHead ptp.ManagementTLVHead + + r := bytes.NewReader(data) + if err := binary.Read(r, binary.BigEndian, &head); err != nil { + return nil, fmt.Errorf("management header: %w", err) + } + if err := binary.Read(r, binary.BigEndian, &tlvHead.TLVHead); err != nil { + return nil, fmt.Errorf("TLV header: %w", err) + } + + if tlvHead.TLVType == ptp.TLVManagementErrorStatus { + e := &mgmtError{} + if err := binary.Read(r, binary.BigEndian, &e.err); err != nil { + return nil, fmt.Errorf("error status: %w", err) + } + if err := binary.Read(r, binary.BigEndian, &e.id); err != nil { + return nil, fmt.Errorf("error status ID: %w", err) + } + return nil, e + } + if tlvHead.TLVType != ptp.TLVManagement { + return nil, fmt.Errorf("unexpected TLV type 0x%04x", uint16(tlvHead.TLVType)) + } + + if err := binary.Read(r, binary.BigEndian, &tlvHead.ManagementID); err != nil { + return nil, fmt.Errorf("management ID: %w", err) + } + + // Rewind to the start of the TLV so the full struct (embedded + // head included) can be read in one go. + tlvStart := int64(binary.Size(head)) + tlvData := data[tlvStart:] + tr := bytes.NewReader(tlvData) + + switch tlvHead.ManagementID { + case ptp.IDDefaultDataSet: + tlv := &ptp.DefaultDataSetTLV{} + return tlv, binary.Read(tr, binary.BigEndian, tlv) + case ptp.IDCurrentDataSet: + tlv := &ptp.CurrentDataSetTLV{} + return tlv, binary.Read(tr, binary.BigEndian, tlv) + case ptp.IDParentDataSet: + tlv := &ptp.ParentDataSetTLV{} + return tlv, binary.Read(tr, binary.BigEndian, tlv) + case ptp.IDTimePropertiesDataSet: + tlv := &timePropertiesDataSetTLV{} + return tlv, binary.Read(tr, binary.BigEndian, tlv) + case ptp.IDPortDataSet: + tlv := &portDataSetTLV{} + return tlv, binary.Read(tr, binary.BigEndian, tlv) + case ptp.IDTimeStatusNP: + tlv := &ptp.TimeStatusNPTLV{} + return tlv, binary.Read(tr, binary.BigEndian, tlv) + case ptp.IDPortStatsNP: + tlv := &ptp.PortStatsNPTLV{} + if err := binary.Read(tr, binary.BigEndian, &tlv.ManagementTLVHead); err != nil { + return nil, err + } + if err := binary.Read(tr, binary.BigEndian, &tlv.PortIdentity); err != nil { + return nil, err + } + // linuxptp sends these counters in host byte order, + // everything else is big-endian + return tlv, binary.Read(tr, binary.NativeEndian, &tlv.PortStats) + case idSubscribeEventsNP: + tlv := &subscribeEventsNPTLV{} + return tlv, binary.Read(tr, binary.BigEndian, tlv) + } + + return nil, fmt.Errorf("unsupported management TLV 0x%04x", uint16(tlvHead.ManagementID)) +} diff --git a/src/yangerd/internal/stpquery/stpquery.go b/src/yangerd/internal/stpquery/stpquery.go new file mode 100644 index 000000000..03a0e3876 --- /dev/null +++ b/src/yangerd/internal/stpquery/stpquery.go @@ -0,0 +1,644 @@ +// Package stpquery provides a native Go client for querying mstpd's +// operational data over its abstract Unix datagram socket. It decodes +// the binary wire protocol directly — no subprocess, no CGo. +package stpquery + +import ( + "encoding/binary" + "encoding/json" + "fmt" + "os" + "strconv" + "sync" + "sync/atomic" + "syscall" + "time" + "unsafe" +) + +const ( + cmdGetCISTBridgeStatus = 101 + cmdGetCISTPortStatus = 105 + serverSocketName = ".mstp_server" +) + +type ctlMsgHdr struct { + Cmd, Lin, Lout, Llog, Res int32 +} + +const hdrSize = 20 + +type Client struct { + fd int +} + +// sockaddrUN is the full-size struct sockaddr_un used by mstpd. +// mstpd passes sizeof(struct sockaddr_un) to bind/connect, so the +// abstract name is zero-padded to fill the entire sun_path[108]. +// Go's net package uses minimal length, which produces a different +// abstract socket name. We must match mstpd's behavior exactly. +type sockaddrUN struct { + Family uint16 + Path [108]byte +} + +func setSockAddr(sa *sockaddrUN, name string) { + sa.Family = syscall.AF_UNIX + copy(sa.Path[1:], name) +} + +var connSeq atomic.Uint64 + +func New() (*Client, error) { + fd, err := syscall.Socket(syscall.AF_UNIX, syscall.SOCK_DGRAM, 0) + if err != nil { + return nil, fmt.Errorf("socket: %w", err) + } + + seq := connSeq.Add(1) + var local sockaddrUN + setSockAddr(&local, fmt.Sprintf("MSTPCTL_%d_%d", os.Getpid(), seq)) + _, _, errno := syscall.Syscall(syscall.SYS_BIND, uintptr(fd), + uintptr(unsafe.Pointer(&local)), unsafe.Sizeof(local)) + if errno != 0 { + syscall.Close(fd) + return nil, fmt.Errorf("bind: %w", errno) + } + + var remote sockaddrUN + setSockAddr(&remote, serverSocketName) + _, _, errno = syscall.Syscall(syscall.SYS_CONNECT, uintptr(fd), + uintptr(unsafe.Pointer(&remote)), unsafe.Sizeof(remote)) + if errno != 0 { + syscall.Close(fd) + return nil, fmt.Errorf("connect to mstpd: %w", errno) + } + + return &Client{fd: fd}, nil +} + +func (c *Client) Close() error { + if c.fd >= 0 { + err := syscall.Close(c.fd) + c.fd = -1 + return err + } + return nil +} + +func (c *Client) roundTrip(cmd int32, in []byte, outSize int) ([]byte, error) { + hdr := ctlMsgHdr{ + Cmd: cmd, + Lin: int32(len(in)), + Lout: int32(outSize), + } + + buf := make([]byte, hdrSize+len(in)) + binary.NativeEndian.PutUint32(buf[0:4], uint32(hdr.Cmd)) + binary.NativeEndian.PutUint32(buf[4:8], uint32(hdr.Lin)) + binary.NativeEndian.PutUint32(buf[8:12], uint32(hdr.Lout)) + binary.NativeEndian.PutUint32(buf[12:16], uint32(hdr.Llog)) + binary.NativeEndian.PutUint32(buf[16:20], uint32(hdr.Res)) + copy(buf[hdrSize:], in) + + tv := syscall.Timeval{Sec: 5} + syscall.SetsockoptTimeval(c.fd, syscall.SOL_SOCKET, syscall.SO_SNDTIMEO, &tv) + syscall.SetsockoptTimeval(c.fd, syscall.SOL_SOCKET, syscall.SO_RCVTIMEO, &tv) + + if err := syscall.Sendmsg(c.fd, buf, nil, nil, 0); err != nil { + return nil, fmt.Errorf("write to mstpd: %w", err) + } + + resp := make([]byte, hdrSize+outSize+4096) + n, _, _, _, err := syscall.Recvmsg(c.fd, resp, nil, 0) + if err != nil { + return nil, fmt.Errorf("read from mstpd: %w", err) + } + if n < hdrSize { + return nil, fmt.Errorf("mstpd response too short: %d bytes", n) + } + + resCode := int32(binary.NativeEndian.Uint32(resp[16:20])) + if resCode != 0 { + return nil, fmt.Errorf("mstpd error: res=%d", resCode) + } + + lout := int(int32(binary.NativeEndian.Uint32(resp[8:12]))) + if n < hdrSize+lout { + return nil, fmt.Errorf("mstpd response truncated: got %d, need %d", n, hdrSize+lout) + } + + return resp[hdrSize : hdrSize+lout], nil +} + +// CISTBridgeStatus holds decoded bridge-level STP data from mstpd. +type CISTBridgeStatus struct { + BridgeID BridgeID + TimeSinceTopologyChange uint32 + TopologyChangeCount uint32 + TopologyChange bool + TopologyChangePort string // max 16 chars + LastTopologyChangePort string // max 16 chars + DesignatedRoot BridgeID + RootPathCost uint32 + RootPortID PortID + RootMaxAge uint8 + RootForwardDelay uint8 + BridgeMaxAge uint8 + BridgeForwardDelay uint8 + TxHoldCount uint32 + ProtocolVersion uint32 + RegionalRoot BridgeID + InternalPathCost uint32 + Enabled bool + AgeingTime uint32 + MaxHops uint8 + BridgeHelloTime uint8 + RootPortName string // from get_cist_bridge_status_OUT tail +} + +// CISTPortStatus holds decoded port-level STP data from mstpd. +type CISTPortStatus struct { + Uptime uint32 + State uint32 + PortID PortID + AdminExternalPortPathCost uint32 + ExternalPortPathCost uint32 + DesignatedRoot BridgeID + DesignatedExternalCost uint32 + DesignatedBridge BridgeID + DesignatedPort PortID + TcAck bool + PortHelloTime uint8 + AdminEdgePort bool + AutoEdgePort bool + OperEdgePort bool + Enabled bool + AdminP2P uint32 + OperP2P bool + RestrictedRole bool + RestrictedTCN bool + Role uint32 + Disputed bool + DesignatedRegionalRoot BridgeID + DesignatedInternalCost uint32 + AdminInternalPortPathCost uint32 + InternalPortPathCost uint32 + BPDUGuardPort bool + BPDUGuardError bool + BPDUFilterPort bool + NetworkPort bool + BAInconsistent bool + NumRxBPDUFiltered uint32 + NumRxBPDU uint32 + NumRxTCN uint32 + NumTxBPDU uint32 + NumTxTCN uint32 + NumTransFwd uint32 + NumTransBlk uint32 + RcvdBpdu bool + RcvdRSTP bool + RcvdSTP bool + RcvdTcAck bool + RcvdTcn bool + SendRSTP bool +} + +// BridgeID is an 8-byte STP bridge identifier. +type BridgeID [8]byte + +// Priority returns the 4-bit priority value (0-15). +func (b BridgeID) Priority() int { + return int(b[0]) >> 4 +} + +// SystemID returns the 12-bit system extension. +func (b BridgeID) SystemID() int { + return (int(b[0])&0x0f)<<8 | int(b[1]) +} + +// Address returns the 6-byte MAC address as a colon-separated string. +func (b BridgeID) Address() string { + return fmt.Sprintf("%02x:%02x:%02x:%02x:%02x:%02x", b[2], b[3], b[4], b[5], b[6], b[7]) +} + +// PortID is a 2-byte STP port identifier (big-endian on wire). +type PortID [2]byte + +// Priority returns the 4-bit port priority (0-15). +func (p PortID) Priority() int { + return int(p[0]) >> 4 +} + +// Number returns the 12-bit port number. +func (p PortID) Number() int { + return (int(p[0])&0x0f)<<8 | int(p[1]) +} + +// GetBridgeStatus queries mstpd for CIST bridge status. +// brIndex is the kernel interface index of the bridge. +func (c *Client) GetBridgeStatus(brIndex int) (*CISTBridgeStatus, error) { + // Input: 4-byte int32 br_index + in := make([]byte, 4) + binary.NativeEndian.PutUint32(in, uint32(int32(brIndex))) + + // Output: 128 bytes = 112 (CIST_BridgeStatus) + 16 (root_port_name) + out, err := c.roundTrip(cmdGetCISTBridgeStatus, in, 128) + if err != nil { + return nil, err + } + if len(out) < 128 { + return nil, fmt.Errorf("bridge status response too short: %d", len(out)) + } + + s := &CISTBridgeStatus{} + copy(s.BridgeID[:], out[0:8]) + s.TimeSinceTopologyChange = binary.NativeEndian.Uint32(out[8:12]) + s.TopologyChangeCount = binary.NativeEndian.Uint32(out[12:16]) + s.TopologyChange = out[16] != 0 + s.TopologyChangePort = cString(out[17:33]) + s.LastTopologyChangePort = cString(out[33:49]) + copy(s.DesignatedRoot[:], out[56:64]) + s.RootPathCost = binary.NativeEndian.Uint32(out[64:68]) + s.RootPortID = PortID{out[68], out[69]} + s.RootMaxAge = out[70] + s.RootForwardDelay = out[71] + s.BridgeMaxAge = out[72] + s.BridgeForwardDelay = out[73] + s.TxHoldCount = binary.NativeEndian.Uint32(out[76:80]) + s.ProtocolVersion = binary.NativeEndian.Uint32(out[80:84]) + copy(s.RegionalRoot[:], out[88:96]) + s.InternalPathCost = binary.NativeEndian.Uint32(out[96:100]) + s.Enabled = out[100] != 0 + s.AgeingTime = binary.NativeEndian.Uint32(out[104:108]) + s.MaxHops = out[108] + s.BridgeHelloTime = out[109] + // Bytes 112..127 = root_port_name[16] + s.RootPortName = cString(out[112:128]) + + return s, nil +} + +// GetPortStatus queries mstpd for CIST port status. +// brIndex and portIndex are kernel interface indices. +func (c *Client) GetPortStatus(brIndex, portIndex int) (*CISTPortStatus, error) { + // Input: 8 bytes = 2x int32 + in := make([]byte, 8) + binary.NativeEndian.PutUint32(in[0:4], uint32(int32(brIndex))) + binary.NativeEndian.PutUint32(in[4:8], uint32(int32(portIndex))) + + // Output: 136 bytes (CIST_PortStatus) + out, err := c.roundTrip(cmdGetCISTPortStatus, in, 136) + if err != nil { + return nil, err + } + if len(out) < 136 { + return nil, fmt.Errorf("port status response too short: %d", len(out)) + } + + s := &CISTPortStatus{} + s.Uptime = binary.NativeEndian.Uint32(out[0:4]) + s.State = binary.NativeEndian.Uint32(out[4:8]) + s.PortID = PortID{out[8], out[9]} + s.AdminExternalPortPathCost = binary.NativeEndian.Uint32(out[12:16]) + s.ExternalPortPathCost = binary.NativeEndian.Uint32(out[16:20]) + copy(s.DesignatedRoot[:], out[24:32]) + s.DesignatedExternalCost = binary.NativeEndian.Uint32(out[32:36]) + copy(s.DesignatedBridge[:], out[40:48]) + s.DesignatedPort = PortID{out[48], out[49]} + s.TcAck = out[50] != 0 + s.PortHelloTime = out[51] + s.AdminEdgePort = out[52] != 0 + s.AutoEdgePort = out[53] != 0 + s.OperEdgePort = out[54] != 0 + s.Enabled = out[55] != 0 + s.AdminP2P = binary.NativeEndian.Uint32(out[56:60]) + s.OperP2P = out[60] != 0 + s.RestrictedRole = out[61] != 0 + s.RestrictedTCN = out[62] != 0 + s.Role = binary.NativeEndian.Uint32(out[64:68]) + s.Disputed = out[68] != 0 + copy(s.DesignatedRegionalRoot[:], out[72:80]) + s.DesignatedInternalCost = binary.NativeEndian.Uint32(out[80:84]) + s.AdminInternalPortPathCost = binary.NativeEndian.Uint32(out[84:88]) + s.InternalPortPathCost = binary.NativeEndian.Uint32(out[88:92]) + s.BPDUGuardPort = out[92] != 0 + s.BPDUGuardError = out[93] != 0 + s.BPDUFilterPort = out[94] != 0 + s.NetworkPort = out[95] != 0 + s.BAInconsistent = out[96] != 0 + s.NumRxBPDUFiltered = binary.NativeEndian.Uint32(out[100:104]) + s.NumRxBPDU = binary.NativeEndian.Uint32(out[104:108]) + s.NumRxTCN = binary.NativeEndian.Uint32(out[108:112]) + s.NumTxBPDU = binary.NativeEndian.Uint32(out[112:116]) + s.NumTxTCN = binary.NativeEndian.Uint32(out[116:120]) + s.NumTransFwd = binary.NativeEndian.Uint32(out[120:124]) + s.NumTransBlk = binary.NativeEndian.Uint32(out[124:128]) + s.RcvdBpdu = out[128] != 0 + s.RcvdRSTP = out[129] != 0 + s.RcvdSTP = out[130] != 0 + s.RcvdTcAck = out[131] != 0 + s.RcvdTcn = out[132] != 0 + s.SendRSTP = out[133] != 0 + + return s, nil +} + +// cString extracts a NUL-terminated C string from a byte slice. +func cString(b []byte) string { + for i, c := range b { + if c == 0 { + return string(b[:i]) + } + } + return string(b) +} + +// protocolName maps mstpd protocol_version to YANG force-protocol value. +func protocolName(v uint32) string { + switch v { + case 0: + return "stp" + case 2: + return "rstp" + default: + return "rstp" + } +} + +// roleName maps mstpd port role to YANG role value. +func roleName(v uint32) string { + switch v { + case 0: + return "disabled" + case 1: + return "root" + case 2: + return "designated" + case 3: + return "alternate" + case 4: + return "backup" + case 5: + return "master" + default: + return "disabled" + } +} + +// bridgeIDMap returns a YANG bridge-id object. +func bridgeIDMap(b BridgeID) map[string]any { + return map[string]any{ + "priority": b.Priority(), + "system-id": b.SystemID(), + "address": b.Address(), + } +} + +// portIDMap returns a YANG port-id object. +func portIDMap(p PortID) map[string]any { + return map[string]any{ + "priority": p.Priority(), + "port-id": p.Number(), + } +} + +// IfIndexResolver looks up kernel interface indices by name. +type IfIndexResolver interface { + IfIndex(name string) (int, bool) +} + +// Query queries mstpd for STP data on all bridges found in the ip-json +// links data. Returns per-bridge and per-port STP JSON fragments ready +// for merging into the YANG interface tree. +// +// Query connects to mstpd, queries STP data for all bridges in links, +// and returns per-bridge and per-port STP JSON fragments. A fresh +// connection is established per call so that late-starting or restarted +// mstpd instances are handled gracefully. +func Query(links json.RawMessage, resolver IfIndexResolver) (bridgeSTP, portSTP map[string]json.RawMessage) { + brs := findBridges(links) + if len(brs) == 0 { + return nil, nil + } + + client, err := New() + if err != nil { + return nil, nil + } + defer client.Close() + + bridgeSTP = make(map[string]json.RawMessage) + portSTP = make(map[string]json.RawMessage) + + for _, br := range brs { + brIdx, ok := resolver.IfIndex(br.name) + if !ok { + continue + } + + bs, err := client.GetBridgeStatus(brIdx) + if err != nil { + continue + } + + stp := buildBridgeSTP(br.name, bs) + if data, err := json.Marshal(stp); err == nil { + bridgeSTP[br.name] = data + } + + for _, port := range br.ports { + portIdx, ok := resolver.IfIndex(port) + if !ok { + continue + } + ps, err := client.GetPortStatus(brIdx, portIdx) + if err != nil { + continue + } + pstp := buildPortSTP(ps) + if data, err := json.Marshal(pstp); err == nil { + portSTP[port] = data + } + } + } + + return bridgeSTP, portSTP +} + +// tcSeen remembers when each bridge's last topology change happened. +// mstpd reports whole seconds since the change, so now minus that jitters +// by a second between polls; the time is only recomputed when the change +// count moves, which keeps the document, and its fingerprint, stable. +var tcSeen = struct { + sync.Mutex + m map[string]tcStamp +}{m: make(map[string]tcStamp)} + +type tcStamp struct { + count uint32 + at time.Time +} + +func tcTime(bridge string, count, since uint32) time.Time { + tcSeen.Lock() + defer tcSeen.Unlock() + if st, ok := tcSeen.m[bridge]; ok && st.count == count { + return st.at + } + at := time.Now().UTC().Add(-time.Duration(since) * time.Second).Truncate(time.Second) + tcSeen.m[bridge] = tcStamp{count, at} + return at +} + +func buildBridgeSTP(name string, bs *CISTBridgeStatus) map[string]any { + cist := map[string]any{ + "bridge-id": bridgeIDMap(bs.BridgeID), + "root-id": bridgeIDMap(bs.DesignatedRoot), + } + + bid := bridgeIDMap(bs.BridgeID) + if prio, ok := bid["priority"]; ok { + cist["priority"] = prio + } + + if bs.RootPortName != "" { + cist["root-port"] = bs.RootPortName + } + + if bs.TopologyChangeCount > 0 { + tc := map[string]any{ + "count": bs.TopologyChangeCount, + "in-progress": bs.TopologyChange, + } + if bs.TopologyChangePort != "" { + tc["port"] = bs.TopologyChangePort + } + if bs.TimeSinceTopologyChange > 0 { + tc["time"] = tcTime(name, bs.TopologyChangeCount, bs.TimeSinceTopologyChange).Format(time.RFC3339) + } + cist["topology-change"] = tc + } + + stp := map[string]any{ + "force-protocol": protocolName(bs.ProtocolVersion), + "hello-time": int(bs.BridgeHelloTime), + "forward-delay": int(bs.BridgeForwardDelay), + "max-age": int(bs.BridgeMaxAge), + "transmit-hold-count": int(bs.TxHoldCount), + "max-hops": int(bs.MaxHops), + "cist": cist, + } + + return stp +} + +func buildPortSTP(ps *CISTPortStatus) map[string]any { + cist := map[string]any{ + "port-id": portIDMap(ps.PortID), + "role": roleName(ps.Role), + "disputed": ps.Disputed, + "external-path-cost": int(ps.ExternalPortPathCost), + "designated": map[string]any{ + "bridge-id": bridgeIDMap(ps.DesignatedBridge), + "port-id": portIDMap(ps.DesignatedPort), + }, + } + + stp := map[string]any{ + "edge": ps.OperEdgePort, + "cist": cist, + "statistics": map[string]any{ + "in-bpdus": strconv.FormatUint(uint64(ps.NumRxBPDU), 10), + "in-bpdus-filtered": strconv.FormatUint(uint64(ps.NumRxBPDUFiltered), 10), + "in-tcns": strconv.FormatUint(uint64(ps.NumRxTCN), 10), + "out-bpdus": strconv.FormatUint(uint64(ps.NumTxBPDU), 10), + "out-tcns": strconv.FormatUint(uint64(ps.NumTxTCN), 10), + "to-blocking": strconv.FormatUint(uint64(ps.NumTransBlk), 10), + "to-forwarding": strconv.FormatUint(uint64(ps.NumTransFwd), 10), + }, + } + + return stp +} + +type bridgeInfo struct { + name string + ports []string +} + +func findBridges(links json.RawMessage) []bridgeInfo { + var ifaces []map[string]any + if json.Unmarshal(links, &ifaces) != nil { + return nil + } + + bridges := make(map[string]*bridgeInfo) + for _, iface := range ifaces { + linkinfo, _ := iface["linkinfo"].(map[string]any) + if linkinfo == nil { + continue + } + + name, _ := iface["ifname"].(string) + if name == "" { + continue + } + + if kind, _ := linkinfo["info_kind"].(string); kind == "bridge" { + if bridges[name] == nil { + bridges[name] = &bridgeInfo{name: name} + } + } + + if master, _ := iface["master"].(string); master != "" { + br := bridges[master] + if br == nil { + br = &bridgeInfo{name: master} + bridges[master] = br + } + br.ports = append(br.ports, name) + } + } + + var result []bridgeInfo + for _, br := range bridges { + if len(br.ports) > 0 { + result = append(result, *br) + } + } + return result +} + +// LinksIfIndexResolver resolves interface names to indices from ip-json link data. +type LinksIfIndexResolver struct { + idx map[string]int +} + +// NewLinksIfIndexResolver builds a resolver from ip-json link data. +func NewLinksIfIndexResolver(links json.RawMessage) *LinksIfIndexResolver { + r := &LinksIfIndexResolver{idx: make(map[string]int)} + var ifaces []map[string]any + if json.Unmarshal(links, &ifaces) != nil { + return r + } + for _, iface := range ifaces { + name, _ := iface["ifname"].(string) + if name == "" { + continue + } + switch v := iface["ifindex"].(type) { + case float64: + r.idx[name] = int(v) + case int: + r.idx[name] = v + } + } + return r +} + +// IfIndex returns the kernel interface index for the given name. +func (r *LinksIfIndexResolver) IfIndex(name string) (int, bool) { + idx, ok := r.idx[name] + return idx, ok +} diff --git a/src/yangerd/internal/stpquery/stpquery_test.go b/src/yangerd/internal/stpquery/stpquery_test.go new file mode 100644 index 000000000..2afca04c0 --- /dev/null +++ b/src/yangerd/internal/stpquery/stpquery_test.go @@ -0,0 +1,18 @@ +package stpquery + +import ( + "testing" + "time" +) + +// The topology-change time only moves when the change count does. +func TestTCTimeStableUntilCountChanges(t *testing.T) { + first := tcTime("br-test", 3, 10) + time.Sleep(1100 * time.Millisecond) + if again := tcTime("br-test", 3, 11); !again.Equal(first) { + t.Fatalf("time moved without a new change: %v -> %v", first, again) + } + if next := tcTime("br-test", 4, 0); !next.After(first) { + t.Fatalf("new change kept the old time: %v", next) + } +} diff --git a/src/yangerd/internal/sysreaders/sysreaders.go b/src/yangerd/internal/sysreaders/sysreaders.go new file mode 100644 index 000000000..3f688a191 --- /dev/null +++ b/src/yangerd/internal/sysreaders/sysreaders.go @@ -0,0 +1,319 @@ +package sysreaders + +import ( + "bufio" + "bytes" + "encoding/json" + "fmt" + "net/netip" + "os" + "path/filepath" + "regexp" + "sort" + "strconv" + "strings" + "sync" +) + +var gmtOffsetRe = regexp.MustCompile(`Etc/GMT([+-]\d{1,2})$`) + +var zonePrefixes = []string{ + "/usr/share/zoneinfo/posix/", + "/usr/share/zoneinfo/right/", + "/usr/share/zoneinfo/", +} + +var userShellMap = map[string]string{ + "/bin/bash": "infix-system:bash", + "/bin/sh": "infix-system:sh", + "/usr/bin/clish": "infix-system:clish", + "/bin/false": "infix-system:false", + "/sbin/nologin": "infix-system:false", + "/usr/sbin/nologin": "infix-system:false", +} + +const SSHDKeysDir = "/var/run/sshd" + +func ReadHostname(path string) (json.RawMessage, error) { + data, err := os.ReadFile(path) + if err != nil { + return nil, err + } + name := strings.TrimSpace(string(data)) + return json.Marshal(map[string]string{"hostname": name}) +} + +func ReadTimezone(path string) (json.RawMessage, error) { + target, err := filepath.EvalSymlinks(path) + if err != nil { + return nil, err + } + + var tz string + for _, p := range zonePrefixes { + if strings.HasPrefix(target, p) { + tz = target[len(p):] + break + } + } + if tz == "" { + return nil, fmt.Errorf("unrecognized zoneinfo path: %s", target) + } + + clock := make(map[string]interface{}) + if m := gmtOffsetRe.FindStringSubmatch(tz); m != nil { + offset, _ := strconv.Atoi(m[1]) + clock["timezone-utc-offset"] = -offset + } else if tz == "Etc/UTC" { + clock["timezone-utc-offset"] = 0 + } else { + clock["timezone-name"] = tz + } + + return json.Marshal(map[string]interface{}{"clock": clock}) +} + +func ReadUsers(_ string) (json.RawMessage, error) { + passwdData, err := os.ReadFile("/etc/passwd") + if err != nil { + return nil, err + } + + passwdUsers := make(map[string]string) + scanner := bufio.NewScanner(bytes.NewReader(passwdData)) + for scanner.Scan() { + parts := strings.Split(scanner.Text(), ":") + if len(parts) < 7 { + continue + } + uid, err := strconv.Atoi(parts[2]) + if err != nil || uid < 1000 || uid >= 10000 { + continue + } + shell := strings.TrimSpace(parts[6]) + mapped, ok := userShellMap[shell] + if !ok { + mapped = "infix-system:false" + } + passwdUsers[parts[0]] = mapped + } + + shadowHashes := make(map[string]string) + shadowData, err := os.ReadFile("/etc/shadow") + if err == nil { + scanner = bufio.NewScanner(bytes.NewReader(shadowData)) + for scanner.Scan() { + parts := strings.SplitN(scanner.Text(), ":", 3) + if len(parts) < 2 { + continue + } + hash := parts[1] + if hash == "" || strings.HasPrefix(hash, "*") || strings.HasPrefix(hash, "!") { + continue + } + shadowHashes[parts[0]] = hash + } + } + + users := make([]interface{}, 0) + for username, shell := range passwdUsers { + user := map[string]interface{}{ + "name": username, + "infix-system:shell": shell, + } + if hash, ok := shadowHashes[username]; ok { + user["password"] = hash + } + + keysData, err := os.ReadFile(filepath.Join(SSHDKeysDir, username+".keys")) + if err == nil { + var authKeys []interface{} + for _, line := range strings.Split(string(keysData), "\n") { + line = strings.TrimSpace(line) + if line == "" || strings.HasPrefix(line, "#") { + continue + } + parts := strings.SplitN(line, " ", 3) + if len(parts) < 2 { + continue + } + keyName := fmt.Sprintf("%s-key-%d", username, len(authKeys)) + if len(parts) > 2 { + keyName = parts[2] + } + authKeys = append(authKeys, map[string]interface{}{ + "name": keyName, + "algorithm": parts[0], + "key-data": parts[1], + }) + } + if len(authKeys) > 0 { + user["authorized-key"] = authKeys + } + } + users = append(users, user) + } + + return json.Marshal(map[string]interface{}{ + "authentication": map[string]interface{}{ + "user": users, + }, + }) +} + +const ( + resolvHead = "/etc/resolv.conf.head" + resolvIfaceDir = "/run/resolvconf/interfaces" +) + +// ReadDNSResolver reports the configured resolvers: the static ones +// from resolv.conf.head, and the ones DHCP clients handed to resolvconf, +// one file per interface, with the interface each came from. +func ReadDNSResolver(_ string) (json.RawMessage, error) { + return readDNSResolver(resolvHead, resolvIfaceDir) +} + +func readDNSResolver(head, ifaceDir string) (json.RawMessage, error) { + r := resolver{servers: []interface{}{}, options: map[string]interface{}{}, seen: map[string]bool{}} + + if data, err := os.ReadFile(head); err == nil { + r.parse(string(data), "static", "") + } + + files, _ := filepath.Glob(filepath.Join(ifaceDir, "*")) + sort.Strings(files) + for _, file := range files { + data, err := os.ReadFile(file) + if err != nil { + continue + } + r.parse(string(data), "dhcp", resolvconfIface(file)) + } + + dns := map[string]interface{}{"server": r.servers} + if len(r.search) > 0 { + dns["search"] = r.search + } + if len(r.options) > 0 { + dns["options"] = r.options + } + + return json.Marshal(map[string]interface{}{"infix-system:dns-resolver": dns}) +} + +// resolvconfIface names the interface of a resolvconf file, written by +// the DHCP client scripts as .conf or -ipv6.conf. +func resolvconfIface(file string) string { + name := strings.TrimSuffix(filepath.Base(file), ".conf") + return strings.TrimSuffix(name, "-ipv6") +} + +type resolver struct { + servers []interface{} + search []string + options map[string]interface{} + seen map[string]bool +} + +// parse reads resolv.conf syntax. The DHCP scripts tag each line with +// "# ", which names the interface when present. +func (r *resolver) parse(data, origin, iface string) { + for _, line := range strings.Split(data, "\n") { + line, comment, _ := strings.Cut(line, "#") + fields := strings.Fields(line) + if len(fields) < 2 { + continue + } + + switch fields[0] { + case "nameserver": + addr, err := netip.ParseAddr(fields[1]) + if err != nil || r.seen[addr.String()] { + continue + } + r.seen[addr.String()] = true + server := map[string]interface{}{ + "address": addr.String(), + "origin": origin, + } + if name := strings.TrimSpace(comment); name != "" && origin == "dhcp" { + server["interface"] = name + } else if iface != "" { + server["interface"] = iface + } + r.servers = append(r.servers, server) + case "search": + r.search = append(r.search, fields[1:]...) + case "options": + for _, opt := range fields[1:] { + key, val, ok := strings.Cut(opt, ":") + if !ok || (key != "timeout" && key != "attempts") { + continue + } + if v, err := strconv.Atoi(val); err == nil { + r.options[key] = v + } + } + } + } +} + +// ForwardingAggregator tracks all /proc/sys/net/ipv{4,6}/conf/*/forwarding +// files and rebuilds the complete interfaces list on every change. +type ForwardingAggregator struct { + mu sync.Mutex +} + +func NewForwardingAggregator() *ForwardingAggregator { + return &ForwardingAggregator{} +} + +func (fa *ForwardingAggregator) HandleForwardingChange(_ string) (json.RawMessage, error) { + fa.mu.Lock() + defer fa.mu.Unlock() + + enabled := make(map[string]bool) + + for _, family := range []string{"ipv4", "ipv6"} { + sysctl := "forwarding" + if family == "ipv6" { + sysctl = "force_forwarding" + } + pattern := fmt.Sprintf("/proc/sys/net/%s/conf/*/%s", family, sysctl) + matches, err := filepath.Glob(pattern) + if err != nil { + continue + } + for _, path := range matches { + b, err := os.ReadFile(path) + if err != nil { + continue + } + if strings.TrimSpace(string(b)) != "1" { + continue + } + parts := strings.Split(filepath.Clean(path), string(os.PathSeparator)) + if len(parts) >= 7 { + ifname := parts[len(parts)-2] + if ifname != "all" && ifname != "default" && ifname != "lo" { + enabled[ifname] = true + } + } + } + } + + ifnames := make([]string, 0) + for name := range enabled { + ifnames = append(ifnames, name) + } + + data, err := json.Marshal(map[string]interface{}{ + "interfaces": map[string]interface{}{ + "interface": ifnames, + }, + }) + if err != nil { + return nil, err + } + return json.RawMessage(data), nil +} diff --git a/src/yangerd/internal/sysreaders/sysreaders_test.go b/src/yangerd/internal/sysreaders/sysreaders_test.go new file mode 100644 index 000000000..03db496e4 --- /dev/null +++ b/src/yangerd/internal/sysreaders/sysreaders_test.go @@ -0,0 +1,86 @@ +package sysreaders + +import ( + "encoding/json" + "os" + "path/filepath" + "reflect" + "testing" +) + +func write(t *testing.T, path, data string) { + t.Helper() + if err := os.MkdirAll(filepath.Dir(path), 0755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(path, []byte(data), 0644); err != nil { + t.Fatal(err) + } +} + +func TestReadHostname(t *testing.T) { + path := filepath.Join(t.TempDir(), "hostname") + write(t, path, "infix-00-00-00\n") + + raw, err := ReadHostname(path) + if err != nil { + t.Fatal(err) + } + if string(raw) != `{"hostname":"infix-00-00-00"}` { + t.Fatalf("got %s", raw) + } +} + +// Static resolvers come from resolv.conf.head, DHCP ones from the +// per-interface resolvconf files, tagged with the interface they were +// learned on. A server listed twice is reported once, it is the key. +func TestReadDNSResolver(t *testing.T) { + dir := t.TempDir() + head := filepath.Join(dir, "resolv.conf.head") + ifaces := filepath.Join(dir, "interfaces") + write(t, head, "nameserver 1.1.1.1\nsearch example.com\noptions timeout:2 attempts:3\n") + write(t, filepath.Join(ifaces, "e1.conf"), + "search lan # e1\nnameserver 192.168.1.1 # e1\nnameserver 1.1.1.1 # e1\n") + write(t, filepath.Join(ifaces, "e2-ipv6.conf"), "nameserver fe80::1\n") + + raw, err := readDNSResolver(head, ifaces) + if err != nil { + t.Fatal(err) + } + var out struct { + DNS struct { + Server []map[string]string `json:"server"` + Search []string `json:"search"` + Options map[string]int `json:"options"` + } `json:"infix-system:dns-resolver"` + } + if err := json.Unmarshal(raw, &out); err != nil { + t.Fatal(err) + } + + wantServers := []map[string]string{ + {"address": "1.1.1.1", "origin": "static"}, + {"address": "192.168.1.1", "origin": "dhcp", "interface": "e1"}, + {"address": "fe80::1", "origin": "dhcp", "interface": "e2"}, + } + if !reflect.DeepEqual(out.DNS.Server, wantServers) { + t.Errorf("servers = %v, want %v", out.DNS.Server, wantServers) + } + if !reflect.DeepEqual(out.DNS.Search, []string{"example.com", "lan"}) { + t.Errorf("search = %v", out.DNS.Search) + } + if out.DNS.Options["timeout"] != 2 || out.DNS.Options["attempts"] != 3 { + t.Errorf("options = %v", out.DNS.Options) + } +} + +func TestReadDNSResolverNothingConfigured(t *testing.T) { + dir := t.TempDir() + raw, err := readDNSResolver(filepath.Join(dir, "none"), filepath.Join(dir, "none.d")) + if err != nil { + t.Fatal(err) + } + if string(raw) != `{"infix-system:dns-resolver":{"server":[]}}` { + t.Fatalf("got %s", raw) + } +} diff --git a/src/yangerd/internal/testutil/mock.go b/src/yangerd/internal/testutil/mock.go new file mode 100644 index 000000000..b5217b78a --- /dev/null +++ b/src/yangerd/internal/testutil/mock.go @@ -0,0 +1,49 @@ +package testutil + +import ( + "context" + "fmt" +) + +// MockRunner records command invocations and returns pre-configured output. +type MockRunner struct { + Results map[string][]byte + Errors map[string]error +} + +// Run returns the pre-configured result for the command name. +func (m *MockRunner) Run(_ context.Context, name string, args ...string) ([]byte, error) { + key := name + for _, a := range args { + key += " " + a + } + if err, ok := m.Errors[key]; ok { + return nil, err + } + if data, ok := m.Results[key]; ok { + return data, nil + } + return nil, fmt.Errorf("mock: no result for %q", key) +} + +// MockFileReader returns pre-configured file contents. +type MockFileReader struct { + Files map[string][]byte + Globs map[string][]string +} + +// ReadFile returns pre-configured data for the path. +func (m *MockFileReader) ReadFile(path string) ([]byte, error) { + if data, ok := m.Files[path]; ok { + return data, nil + } + return nil, fmt.Errorf("mock: file not found: %s", path) +} + +// Glob returns pre-configured matches for the pattern. +func (m *MockFileReader) Glob(pattern string) ([]string, error) { + if matches, ok := m.Globs[pattern]; ok { + return matches, nil + } + return nil, nil +} diff --git a/src/yangerd/internal/tftpmonitor/mounts.go b/src/yangerd/internal/tftpmonitor/mounts.go new file mode 100644 index 000000000..39b73f20e --- /dev/null +++ b/src/yangerd/internal/tftpmonitor/mounts.go @@ -0,0 +1,66 @@ +package tftpmonitor + +import ( + "context" + "fmt" + "io" + "os" + + "golang.org/x/sys/unix" +) + +// watchMounts calls notify every time the mount table changes, until ctx +// is cancelled. A mount produces no inotify event, but the kernel flags +// /proc/self/mountinfo with POLLPRI when the table changes, see proc(5). +// The file has to be read to the end after each change to re-arm it. +func watchMounts(ctx context.Context, path string, notify func()) error { + f, err := os.Open(path) + if err != nil { + return err + } + defer f.Close() + + var wake [2]int + if err := unix.Pipe2(wake[:], unix.O_CLOEXEC|unix.O_NONBLOCK); err != nil { + return fmt.Errorf("pipe: %w", err) + } + defer unix.Close(wake[0]) + defer unix.Close(wake[1]) + + stop := make(chan struct{}) + defer close(stop) + go func() { + select { + case <-ctx.Done(): + unix.Write(wake[1], []byte{0}) + case <-stop: + } + }() + + drain := func() { + f.Seek(0, io.SeekStart) + io.Copy(io.Discard, f) + } + drain() + + fds := []unix.PollFd{ + {Fd: int32(f.Fd()), Events: unix.POLLPRI}, + {Fd: int32(wake[0]), Events: unix.POLLIN}, + } + for { + fds[0].Revents, fds[1].Revents = 0, 0 + if _, err := unix.Poll(fds, -1); err != nil { + if err == unix.EINTR { + continue + } + return fmt.Errorf("poll %s: %w", path, err) + } + if fds[1].Revents != 0 { + return ctx.Err() + } + if fds[0].Revents&(unix.POLLPRI|unix.POLLERR) != 0 { + drain() + notify() + } + } +} diff --git a/src/yangerd/internal/tftpmonitor/mounts_test.go b/src/yangerd/internal/tftpmonitor/mounts_test.go new file mode 100644 index 000000000..602364044 --- /dev/null +++ b/src/yangerd/internal/tftpmonitor/mounts_test.go @@ -0,0 +1,48 @@ +package tftpmonitor + +import ( + "context" + "os" + "path/filepath" + "testing" + "time" +) + +// A regular file never signals POLLPRI, so nothing fires, and the +// watcher returns promptly when cancelled instead of blocking in poll. +func TestWatchMountsCancel(t *testing.T) { + path := filepath.Join(t.TempDir(), "mountinfo") + if err := os.WriteFile(path, []byte("22 1 0:21 / / rw - ext4 /dev/root rw\n"), 0644); err != nil { + t.Fatal(err) + } + + ctx, cancel := context.WithCancel(context.Background()) + fired := make(chan struct{}, 1) + done := make(chan error, 1) + go func() { + done <- watchMounts(ctx, path, func() { fired <- struct{}{} }) + }() + + select { + case <-fired: + t.Fatal("no mount change, nothing may fire") + case <-time.After(100 * time.Millisecond): + } + + cancel() + select { + case err := <-done: + if err != context.Canceled { + t.Fatalf("err = %v, want context.Canceled", err) + } + case <-time.After(2 * time.Second): + t.Fatal("watchMounts did not return after cancel") + } +} + +func TestWatchMountsMissingFile(t *testing.T) { + err := watchMounts(context.Background(), filepath.Join(t.TempDir(), "none"), func() {}) + if err == nil { + t.Fatal("expected an error for a missing mount table") + } +} diff --git a/src/yangerd/internal/tftpmonitor/tftpmonitor.go b/src/yangerd/internal/tftpmonitor/tftpmonitor.go new file mode 100644 index 000000000..f0017c61d --- /dev/null +++ b/src/yangerd/internal/tftpmonitor/tftpmonitor.go @@ -0,0 +1,299 @@ +// Package tftpmonitor keeps the infix-services:tftp subtree in sync with +// what dnsmasq serves. confd writes /etc/dnsmasq.d/tftp.conf when the +// TFTP server is enabled and removes it when disabled, so that file is +// the source of truth for whether there is anything to report and where +// the root is. The root and every directory below it are then watched +// with inotify, and the file list is rebuilt whenever something changes: +// a file is uploaded, removed, renamed or has its mode changed, a +// directory appears or disappears, the root itself moves, or something +// is mounted, like a USB stick holding the root. +// +// Only world-readable regular files are listed, following symlinks the +// way dnsmasq does when it opens them. +package tftpmonitor + +import ( + "context" + "encoding/json" + "fmt" + "log/slog" + "os" + "path/filepath" + "sort" + "strconv" + "strings" + "time" + + "github.com/kernelkit/infix/src/yangerd/internal/inotify" + "github.com/kernelkit/infix/src/yangerd/internal/tree" +) + +const ( + treeKey = "infix-services:tftp" + confPath = "/etc/dnsmasq.d/tftp.conf" + rootKey = "tftp-root=" + + // debounceDelay coalesces bursts of events into one rescan. A file + // upload is many writes, and confd rewrites the snippet in place + // (create, several writes) rather than renaming a temp file over it. + debounceDelay = 300 * time.Millisecond + + // maxDepth bounds the walk below the root, symlink loops are caught + // separately but a deep tree is not something a TFTP server needs. + maxDepth = 16 +) + +// TFTPMonitor watches the dnsmasq TFTP snippet and the TFTP root. +type TFTPMonitor struct { + tree *tree.Tree + log *slog.Logger + + // conf is the dnsmasq snippet to read the root from; overridable in + // tests. + conf string + + // mountinfo is the mount table polled for changes; overridable in + // tests. + mountinfo string + + watcher *inotify.Watcher + root string // current TFTP root, "" when disabled + watched map[string]bool // directories currently watched below root +} + +// New creates a TFTPMonitor. +func New(t *tree.Tree, log *slog.Logger) *TFTPMonitor { + if log == nil { + log = slog.Default() + } + return &TFTPMonitor{ + tree: t, + log: log, + conf: confPath, + mountinfo: "/proc/self/mountinfo", + watched: make(map[string]bool), + } +} + +// Run watches the snippet and the root until ctx is cancelled. All +// state is owned by this goroutine; events only schedule a rescan. +func (m *TFTPMonitor) Run(ctx context.Context) error { + w, err := inotify.NewWatcher() + if err != nil { + return fmt.Errorf("inotify: %w", err) + } + defer w.Close() + m.watcher = w + + // confd creates and removes the snippet, so watch its directory. + confDir := filepath.Dir(m.conf) + if err := w.Add(confDir); err != nil { + return fmt.Errorf("watch %s: %w", confDir, err) + } + + // A root on removable media is a mount point, and mounting it + // produces no inotify event. + mounted := make(chan struct{}, 1) + go func() { + err := watchMounts(ctx, m.mountinfo, func() { + select { + case mounted <- struct{}{}: + default: + } + }) + if err != nil && ctx.Err() == nil { + m.log.Warn("tftp monitor: mount table not watched", "err", err) + } + }() + + m.reconcile() + + timer := time.NewTimer(time.Hour) + timer.Stop() + + for { + select { + case <-ctx.Done(): + return ctx.Err() + + case ev, ok := <-w.Events: + if !ok { + return fmt.Errorf("watcher closed") + } + if filepath.Dir(ev.Name) == confDir && ev.Name != m.conf { + continue // another dnsmasq snippet + } + timer.Reset(debounceDelay) + + case <-mounted: + timer.Reset(debounceDelay) + + case <-timer.C: + m.reconcile() + + case err, ok := <-w.Errors: + if !ok { + return fmt.Errorf("watcher error channel closed") + } + m.log.Warn("tftp monitor: inotify error", "err", err) + } + } +} + +// reconcile re-reads the snippet, re-syncs the directory watches to the +// current root and rewrites the subtree. Called on start and after every +// settled burst of events, and cheap enough to be a full rebuild. +func (m *TFTPMonitor) reconcile() { + root := m.readRoot() + if root != m.root { + if root == "" { + m.log.Info("tftp monitor: server disabled") + } else { + m.log.Info("tftp monitor: serving from", "root", root) + } + m.root = root + } + + if root == "" { + m.syncWatches(nil) + m.tree.Delete(treeKey) + return + } + + files, dirs := scan(root, m.log) + m.syncWatches(dirs) + + data, err := json.Marshal(map[string]any{ + "files": map[string]any{"file": files}, + }) + if err != nil { + m.log.Warn("tftp monitor: marshal", "err", err) + return + } + m.tree.Set(treeKey, data) + m.log.Debug("tftp monitor: tree updated", "root", root, "files", len(files)) +} + +// readRoot returns the tftp-root from the snippet, or "" when the server +// is disabled, i.e. the snippet is gone or names no root. +func (m *TFTPMonitor) readRoot() string { + data, err := os.ReadFile(m.conf) + if err != nil { + return "" + } + for _, line := range strings.Split(string(data), "\n") { + if strings.HasPrefix(line, rootKey) { + return strings.TrimSuffix(strings.TrimPrefix(line, rootKey), "/") + } + } + return "" +} + +// syncWatches makes the set of watched directories equal to dirs. The +// snippet directory is not in the set and is never touched. +func (m *TFTPMonitor) syncWatches(dirs []string) { + want := make(map[string]bool, len(dirs)) + for _, dir := range dirs { + want[dir] = true + } + + for dir := range m.watched { + if want[dir] { + continue + } + if m.watcher != nil { + // A removed directory already dropped its watch, ignore. + _ = m.watcher.Remove(dir) + } + delete(m.watched, dir) + } + + // Add even what is already watched: the kernel drops the watch of a + // deleted directory, so one removed and recreated between two scans + // needs it again, and Add on a live watch is a no-op. + for dir := range want { + if m.watcher != nil { + if err := m.watcher.Add(dir); err != nil { + m.log.Warn("tftp monitor: watch failed", "dir", dir, "err", err) + continue + } + } + m.watched[dir] = true + } +} + +// fileEntry is one entry of the files/file list. Size is a uint64 in +// the model, hence a string per RFC 7951. +type fileEntry struct { + Name string `json:"name"` + Size string `json:"size"` + Modified string `json:"modified"` +} + +// scan lists the world-readable regular files below root, sorted by +// name, and the directories that must be watched to notice a change. +// +// A missing root is a legitimate state: the server is enabled and dnsmasq +// runs with tftp-no-fail, so the list is empty and the nearest existing +// ancestor is watched to catch the root, or a directory on the way to +// it, being created. +func scan(root string, log *slog.Logger) ([]fileEntry, []string) { + files := []fileEntry{} + + if fi, err := os.Stat(root); err != nil || !fi.IsDir() { + for dir := filepath.Dir(root); ; dir = filepath.Dir(dir) { + if fi, err := os.Stat(dir); err == nil && fi.IsDir() { + return files, []string{dir} + } + if dir == filepath.Dir(dir) { + return files, nil + } + } + } + + var dirs []string + seen := make(map[string]bool) // real paths, to break symlink loops + + var walk func(dir string, depth int) + walk = func(dir string, depth int) { + real, err := filepath.EvalSymlinks(dir) + if err != nil || seen[real] { + return + } + seen[real] = true + dirs = append(dirs, dir) + + entries, err := os.ReadDir(dir) + if err != nil { + log.Debug("tftp monitor: read dir", "dir", dir, "err", err) + return + } + for _, entry := range entries { + path := filepath.Join(dir, entry.Name()) + fi, err := os.Stat(path) // follows symlinks, like dnsmasq + if err != nil { + continue + } + switch { + case fi.IsDir(): + if depth < maxDepth { + walk(path, depth+1) + } + case fi.Mode().IsRegular() && fi.Mode().Perm()&0004 != 0: + name, err := filepath.Rel(root, path) + if err != nil { + continue + } + files = append(files, fileEntry{ + Name: filepath.ToSlash(name), + Size: strconv.FormatInt(fi.Size(), 10), + Modified: fi.ModTime().UTC().Format(time.RFC3339), + }) + } + } + } + walk(root, 0) + + sort.Slice(files, func(i, j int) bool { return files[i].Name < files[j].Name }) + return files, dirs +} diff --git a/src/yangerd/internal/tftpmonitor/tftpmonitor_test.go b/src/yangerd/internal/tftpmonitor/tftpmonitor_test.go new file mode 100644 index 000000000..63810131d --- /dev/null +++ b/src/yangerd/internal/tftpmonitor/tftpmonitor_test.go @@ -0,0 +1,306 @@ +package tftpmonitor + +import ( + "context" + "encoding/json" + "os" + "path/filepath" + "testing" + "time" + + "github.com/kernelkit/infix/src/yangerd/internal/tree" +) + +func newTestMonitor(t *testing.T) (*TFTPMonitor, *tree.Tree, string) { + t.Helper() + tr := tree.New() + m := New(tr, nil) + m.conf = filepath.Join(t.TempDir(), "tftp.conf") + return m, tr, m.conf +} + +func enable(t *testing.T, conf, root string) { + t.Helper() + snippet := "enable-tftp\ntftp-root=" + root + "\ntftp-no-fail\n" + if err := os.WriteFile(conf, []byte(snippet), 0644); err != nil { + t.Fatal(err) + } +} + +func disable(t *testing.T, conf string) { + t.Helper() + if err := os.Remove(conf); err != nil { + t.Fatal(err) + } +} + +func put(t *testing.T, path, content string, mode os.FileMode) { + t.Helper() + if err := os.MkdirAll(filepath.Dir(path), 0755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(path, []byte(content), mode); err != nil { + t.Fatal(err) + } + // umask must not decide the outcome + if err := os.Chmod(path, mode); err != nil { + t.Fatal(err) + } +} + +func names(t *testing.T, tr *tree.Tree) []string { + t.Helper() + raw := tr.Get(treeKey) + if raw == nil { + return nil + } + var data struct { + Files struct { + File []fileEntry `json:"file"` + } `json:"files"` + } + if err := json.Unmarshal(raw, &data); err != nil { + t.Fatalf("bad tree data %s: %v", raw, err) + } + out := []string{} + for _, f := range data.Files.File { + out = append(out, f.Name) + } + return out +} + +func equal(a, b []string) bool { + if len(a) != len(b) { + return false + } + for i := range a { + if a[i] != b[i] { + return false + } + } + return true +} + +func waitFor(t *testing.T, tr *tree.Tree, want []string) { + t.Helper() + deadline := time.Now().Add(5 * time.Second) + for time.Now().Before(deadline) { + if got := names(t, tr); equal(got, want) { + return + } + time.Sleep(20 * time.Millisecond) + } + t.Fatalf("timeout waiting for %v, have %v", want, names(t, tr)) +} + +func waitForAbsent(t *testing.T, tr *tree.Tree) { + t.Helper() + deadline := time.Now().Add(5 * time.Second) + for time.Now().Before(deadline) { + if tr.Get(treeKey) == nil { + return + } + time.Sleep(20 * time.Millisecond) + } + t.Fatalf("timeout waiting for key to go away, have %s", tr.Get(treeKey)) +} + +// Without the snippet the server is disabled and the key must not exist. +func TestDisabledDeletesKey(t *testing.T) { + m, tr, _ := newTestMonitor(t) + tr.Set(treeKey, json.RawMessage(`{"files":{"file":[{"name":"stale"}]}}`)) + + m.reconcile() + + if got := tr.Get(treeKey); got != nil { + t.Fatalf("expected key removed, got %s", got) + } +} + +// Only world-readable regular files are listed, recursively and through +// symlinks, sorted by name. +func TestListsWorldReadableFiles(t *testing.T) { + m, tr, conf := newTestMonitor(t) + root := filepath.Join(t.TempDir(), "tftpboot") + put(t, filepath.Join(root, "zimage"), "kernel", 0644) + put(t, filepath.Join(root, "private.bin"), "secret", 0600) + put(t, filepath.Join(root, "sub", "dtb"), "tree", 0444) + outside := filepath.Join(t.TempDir(), "outside.bin") + put(t, outside, "linked", 0644) + if err := os.Symlink(outside, filepath.Join(root, "link.bin")); err != nil { + t.Fatal(err) + } + enable(t, conf, root) + + m.reconcile() + + want := []string{"link.bin", "sub/dtb", "zimage"} + if got := names(t, tr); !equal(got, want) { + t.Fatalf("want %v, got %v", want, got) + } + + var data struct { + Files struct { + File []fileEntry `json:"file"` + } `json:"files"` + } + if err := json.Unmarshal(tr.Get(treeKey), &data); err != nil { + t.Fatal(err) + } + if data.Files.File[0].Size != "6" { + t.Fatalf("size must be a decimal string, got %q", data.Files.File[0].Size) + } + if _, err := time.Parse(time.RFC3339, data.Files.File[0].Modified); err != nil { + t.Fatalf("modified not RFC 3339: %v", err) + } +} + +// A symlink loop below the root must not hang the scan. +func TestSymlinkLoop(t *testing.T) { + m, tr, conf := newTestMonitor(t) + root := filepath.Join(t.TempDir(), "tftpboot") + put(t, filepath.Join(root, "a", "file"), "x", 0644) + if err := os.Symlink(root, filepath.Join(root, "a", "loop")); err != nil { + t.Fatal(err) + } + enable(t, conf, root) + + m.reconcile() + + if got := names(t, tr); !equal(got, []string{"a/file"}) { + t.Fatalf("got %v", got) + } +} + +// The running monitor follows uploads, mode changes, a moved root, and +// the server being disabled. +func TestRunFollowsChanges(t *testing.T) { + m, tr, conf := newTestMonitor(t) + base := t.TempDir() + rootA := filepath.Join(base, "a") + rootB := filepath.Join(base, "b") + put(t, filepath.Join(rootA, "one"), "1", 0644) + put(t, filepath.Join(rootB, "two"), "2", 0644) + enable(t, conf, rootA) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + done := make(chan struct{}) + go func() { + defer close(done) + _ = m.Run(ctx) + }() + + waitFor(t, tr, []string{"one"}) + + put(t, filepath.Join(rootA, "sub", "three"), "3", 0644) + waitFor(t, tr, []string{"one", "sub/three"}) + + // A new directory must be watched too, not just the ones at start + put(t, filepath.Join(rootA, "sub", "four"), "4", 0644) + waitFor(t, tr, []string{"one", "sub/four", "sub/three"}) + + if err := os.Chmod(filepath.Join(rootA, "one"), 0600); err != nil { + t.Fatal(err) + } + waitFor(t, tr, []string{"sub/four", "sub/three"}) + + enable(t, conf, rootB) + waitFor(t, tr, []string{"two"}) + + // The old root is no longer watched, changes there must not show + put(t, filepath.Join(rootA, "five"), "5", 0644) + put(t, filepath.Join(rootB, "six"), "6", 0644) + waitFor(t, tr, []string{"six", "two"}) + + disable(t, conf) + waitForAbsent(t, tr) + + put(t, filepath.Join(rootB, "seven"), "7", 0644) + time.Sleep(2 * debounceDelay) + if got := tr.Get(treeKey); got != nil { + t.Fatalf("disabled server must stay absent, got %s", got) + } + + cancel() + <-done +} + +// An enabled server whose root does not exist yet lists nothing, and +// picks the root up once it is created. +func TestRootCreatedLater(t *testing.T) { + m, tr, conf := newTestMonitor(t) + root := filepath.Join(t.TempDir(), "tftpboot") + enable(t, conf, root) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + done := make(chan struct{}) + go func() { + defer close(done) + _ = m.Run(ctx) + }() + + waitFor(t, tr, []string{}) + + put(t, filepath.Join(root, "late"), "x", 0644) + waitFor(t, tr, []string{"late"}) + + cancel() + <-done +} + +// A directory removed and recreated within one debounce window must +// still be watched afterwards. +func TestRecreatedDirStaysWatched(t *testing.T) { + m, tr, conf := newTestMonitor(t) + root := filepath.Join(t.TempDir(), "tftpboot") + put(t, filepath.Join(root, "fw", "old"), "x", 0644) + enable(t, conf, root) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + done := make(chan struct{}) + go func() { + defer close(done) + _ = m.Run(ctx) + }() + + waitFor(t, tr, []string{"fw/old"}) + + if err := os.RemoveAll(filepath.Join(root, "fw")); err != nil { + t.Fatal(err) + } + put(t, filepath.Join(root, "fw", "new"), "y", 0644) + waitFor(t, tr, []string{"fw/new"}) + + put(t, filepath.Join(root, "fw", "later"), "z", 0644) + waitFor(t, tr, []string{"fw/later", "fw/new"}) + + cancel() + <-done +} + +// With neither the root nor its parent present, the nearest existing +// ancestor is watched so the whole path can appear later. +func TestRootWithMissingParent(t *testing.T) { + m, tr, conf := newTestMonitor(t) + root := filepath.Join(t.TempDir(), "media", "usb") + enable(t, conf, root) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + done := make(chan struct{}) + go func() { + defer close(done) + _ = m.Run(ctx) + }() + + waitFor(t, tr, []string{}) + + put(t, filepath.Join(root, "late"), "x", 0644) + waitFor(t, tr, []string{"late"}) + + cancel() + <-done +} diff --git a/src/yangerd/internal/tree/tree.go b/src/yangerd/internal/tree/tree.go new file mode 100644 index 000000000..18c5a8696 --- /dev/null +++ b/src/yangerd/internal/tree/tree.go @@ -0,0 +1,263 @@ +// Package tree provides a concurrent in-memory store for per-module +// YANG operational data, keyed by module-qualified names like +// "ietf-system:system-state". +// +// Each module has its own read-write mutex, so writers for different +// modules never block each other. All methods are safe for concurrent +// use. +package tree + +import ( + "encoding/json" + "sync" + "time" +) + +// OnDemandFunc returns a JSON blob computed at call time. +// Registered providers are invoked on every Get/GetMulti to supply +// fields that must always be fresh (e.g. uptime, current-datetime). +type OnDemandFunc func() json.RawMessage + +// modelEntry holds a single YANG module's pre-serialized JSON blob +// and its own read-write mutex. +type modelEntry struct { + mu sync.RWMutex + data json.RawMessage + updated time.Time +} + +// Tree holds the operational YANG data in per-module JSON blobs. +type Tree struct { + mu sync.RWMutex // protects the models map itself + models map[string]*modelEntry + providers map[string]OnDemandFunc // on-demand overlay providers +} + +// New creates an empty Tree. +func New() *Tree { + return &Tree{ + models: make(map[string]*modelEntry), + providers: make(map[string]OnDemandFunc), + } +} + +// RegisterProvider adds an on-demand overlay for the given key. +// When Get or GetMulti reads this key the provider is called and its +// result is shallow-merged on top of the cached data. The cached +// entry is never mutated — a merged copy is returned. +func (t *Tree) RegisterProvider(key string, fn OnDemandFunc) { + t.mu.Lock() + t.providers[key] = fn + t.mu.Unlock() +} + +// entry returns the model entry for key, creating it if absent. +func (t *Tree) entry(key string) *modelEntry { + t.mu.RLock() + entry, ok := t.models[key] + t.mu.RUnlock() + if ok { + return entry + } + t.mu.Lock() + defer t.mu.Unlock() + if entry, ok = t.models[key]; !ok { + entry = &modelEntry{} + t.models[key] = entry + } + return entry +} + +// Set replaces the entire subtree at the given YANG module key. +// Only the target module's write lock is held; other modules remain +// readable and writable. +func (t *Tree) Set(key string, v json.RawMessage) { + entry := t.entry(key) + entry.mu.Lock() + entry.data = v + entry.updated = time.Now() + entry.mu.Unlock() +} + +// Get returns the raw JSON for the given module key. +// If a provider is registered for key its output is shallow-merged on +// top of the cached data without mutating the cache. +func (t *Tree) Get(key string) json.RawMessage { + t.mu.RLock() + entry, ok := t.models[key] + provider := t.providers[key] + t.mu.RUnlock() + if !ok { + return nil + } + entry.mu.RLock() + data := entry.data + entry.mu.RUnlock() + + if provider == nil { + return data + } + return ShallowMerge(data, provider()) +} + +// GetMulti returns the raw JSON for multiple module keys. +// Each module's read lock is acquired and released individually -- +// the result is eventually consistent, not a snapshot. +// Providers are applied per-key, same as Get, and run without any +// tree lock held: a provider may read the tree itself. +func (t *Tree) GetMulti(keys []string) []json.RawMessage { + type pick struct { + entry *modelEntry + provider OnDemandFunc + } + picks := make([]pick, 0, len(keys)) + t.mu.RLock() + for _, key := range keys { + if entry, ok := t.models[key]; ok { + picks = append(picks, pick{entry, t.providers[key]}) + } + } + t.mu.RUnlock() + + result := make([]json.RawMessage, 0, len(picks)) + for _, p := range picks { + p.entry.mu.RLock() + data := p.entry.data + p.entry.mu.RUnlock() + + if p.provider != nil { + data = ShallowMerge(data, p.provider()) + } + result = append(result, data) + } + return result +} + +// Keys returns all registered module keys. +func (t *Tree) Keys() []string { + t.mu.RLock() + defer t.mu.RUnlock() + keys := make([]string, 0, len(t.models)) + for k := range t.models { + keys = append(keys, k) + } + return keys +} + +// ModelInfo holds metadata for a single model key. +type ModelInfo struct { + LastUpdated time.Time + SizeBytes int +} + +// Merge performs a shallow first-level JSON merge of partial into +// the existing blob at key. If the key does not exist yet, partial +// becomes the entire value. Each top-level field in partial +// overwrites the corresponding field in the existing object; fields +// not mentioned in partial are preserved. +// +// Both the existing data and partial must be JSON objects (maps). +// If either is not a valid JSON object, partial replaces the value. +func (t *Tree) Merge(key string, partial json.RawMessage) { + entry := t.entry(key) + entry.mu.Lock() + defer entry.mu.Unlock() + + // Unmarshal existing data. + var base map[string]json.RawMessage + if len(entry.data) == 0 || json.Unmarshal(entry.data, &base) != nil { + base = make(map[string]json.RawMessage) + } + + // Unmarshal partial. + var overlay map[string]json.RawMessage + if json.Unmarshal(partial, &overlay) != nil { + // partial is not a JSON object — fall back to full replace. + entry.data = partial + entry.updated = time.Now() + return + } + + for k, v := range overlay { + base[k] = v + } + + merged, err := json.Marshal(base) + if err != nil { + // Should never happen with valid JSON inputs. + entry.data = partial + entry.updated = time.Now() + return + } + entry.data = merged + entry.updated = time.Now() +} + +// Delete removes a key from the tree entirely. +func (t *Tree) Delete(key string) { + t.mu.Lock() + delete(t.models, key) + t.mu.Unlock() +} + +// ShallowMerge overlays the top-level fields of overlay onto base and +// returns a new JSON blob. Neither input is modified. If either +// is not a valid JSON object the overlay wins outright. +func ShallowMerge(base, overlay json.RawMessage) json.RawMessage { + if len(overlay) == 0 { + return base + } + if len(base) == 0 { + return overlay + } + + var bm map[string]json.RawMessage + if json.Unmarshal(base, &bm) != nil { + return overlay + } + var om map[string]json.RawMessage + if json.Unmarshal(overlay, &om) != nil { + return overlay + } + + for k, v := range om { + bm[k] = v + } + + merged, err := json.Marshal(bm) + if err != nil { + return overlay + } + return merged +} + +// GetCached returns the raw cached JSON for the given module key +// WITHOUT invoking any registered provider. This is safe to call +// from inside a provider closure (no recursion risk). +func (t *Tree) GetCached(key string) json.RawMessage { + t.mu.RLock() + entry, ok := t.models[key] + t.mu.RUnlock() + if !ok { + return nil + } + entry.mu.RLock() + defer entry.mu.RUnlock() + return entry.data +} + +// Info returns metadata for the given module key. +func (t *Tree) Info(key string) (ModelInfo, bool) { + t.mu.RLock() + entry, ok := t.models[key] + t.mu.RUnlock() + if !ok { + return ModelInfo{}, false + } + entry.mu.RLock() + defer entry.mu.RUnlock() + return ModelInfo{ + LastUpdated: entry.updated, + SizeBytes: len(entry.data), + }, true +} diff --git a/src/yangerd/internal/tree/tree_test.go b/src/yangerd/internal/tree/tree_test.go new file mode 100644 index 000000000..443dcaca4 --- /dev/null +++ b/src/yangerd/internal/tree/tree_test.go @@ -0,0 +1,384 @@ +package tree + +import ( + "encoding/json" + "sync" + "testing" + "time" +) + +func TestSetGet(t *testing.T) { + tr := New() + tr.Set("ietf-system:system", json.RawMessage(`{"hostname":"r1"}`)) + + got := tr.Get("ietf-system:system") + if string(got) != `{"hostname":"r1"}` { + t.Fatalf("unexpected: %s", got) + } +} + +func TestGetMissing(t *testing.T) { + tr := New() + if got := tr.Get("nonexistent"); got != nil { + t.Fatalf("expected nil, got: %s", got) + } +} + +func TestSetOverwrite(t *testing.T) { + tr := New() + tr.Set("key", json.RawMessage(`"v1"`)) + tr.Set("key", json.RawMessage(`"v2"`)) + + if got := tr.Get("key"); string(got) != `"v2"` { + t.Fatalf("expected v2, got: %s", got) + } +} + +func TestGetMulti(t *testing.T) { + tr := New() + tr.Set("a", json.RawMessage(`1`)) + tr.Set("b", json.RawMessage(`2`)) + tr.Set("c", json.RawMessage(`3`)) + + results := tr.GetMulti([]string{"a", "c"}) + if len(results) != 2 { + t.Fatalf("expected 2 results, got %d", len(results)) + } + if string(results[0]) != "1" || string(results[1]) != "3" { + t.Fatalf("unexpected results: %s, %s", results[0], results[1]) + } +} + +func TestGetMultiMissing(t *testing.T) { + tr := New() + tr.Set("a", json.RawMessage(`1`)) + + results := tr.GetMulti([]string{"a", "missing"}) + if len(results) != 1 { + t.Fatalf("expected 1 result, got %d", len(results)) + } +} + +func TestKeys(t *testing.T) { + tr := New() + tr.Set("x", json.RawMessage(`1`)) + tr.Set("y", json.RawMessage(`2`)) + + keys := tr.Keys() + if len(keys) != 2 { + t.Fatalf("expected 2 keys, got %d", len(keys)) + } +} + +func TestInfo(t *testing.T) { + tr := New() + tr.Set("k", json.RawMessage(`{"data":true}`)) + + info, ok := tr.Info("k") + if !ok { + t.Fatal("expected ok") + } + if info.SizeBytes != len(`{"data":true}`) { + t.Fatalf("expected size %d, got %d", len(`{"data":true}`), info.SizeBytes) + } + if info.LastUpdated.IsZero() { + t.Fatal("expected non-zero LastUpdated") + } +} + +func TestInfoMissing(t *testing.T) { + tr := New() + _, ok := tr.Info("missing") + if ok { + t.Fatal("expected !ok for missing key") + } +} + +func TestConcurrentSetGet(t *testing.T) { + tr := New() + var wg sync.WaitGroup + const N = 100 + + for i := 0; i < N; i++ { + wg.Add(2) + go func(i int) { + defer wg.Done() + tr.Set("shared", json.RawMessage(`{"i":`+string(rune('0'+i%10))+`}`)) + }(i) + go func() { + defer wg.Done() + tr.Get("shared") + }() + } + wg.Wait() + + if got := tr.Get("shared"); got == nil { + t.Fatal("expected non-nil after concurrent writes") + } +} + +func TestMerge(t *testing.T) { + tests := []struct { + name string + existing string + partial string + want map[string]json.RawMessage + }{ + { + name: "merge into existing preserves old fields", + existing: `{"a":"1","b":"2"}`, + partial: `{"c":"3"}`, + want: map[string]json.RawMessage{ + "a": json.RawMessage(`"1"`), + "b": json.RawMessage(`"2"`), + "c": json.RawMessage(`"3"`), + }, + }, + { + name: "merge overwrites overlapping field", + existing: `{"a":"old","b":"keep"}`, + partial: `{"a":"new"}`, + want: map[string]json.RawMessage{ + "a": json.RawMessage(`"new"`), + "b": json.RawMessage(`"keep"`), + }, + }, + { + name: "merge with complex nested values", + existing: `{"protocols":{"ospf":true}}`, + partial: `{"ribs":{"rib":[{"name":"ipv4"}]}}`, + want: map[string]json.RawMessage{ + "protocols": json.RawMessage(`{"ospf":true}`), + "ribs": json.RawMessage(`{"rib":[{"name":"ipv4"}]}`), + }, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + tr := New() + tr.Set("k", json.RawMessage(tc.existing)) + tr.Merge("k", json.RawMessage(tc.partial)) + + var got map[string]json.RawMessage + if err := json.Unmarshal(tr.Get("k"), &got); err != nil { + t.Fatalf("unmarshal result: %v", err) + } + for field, wantVal := range tc.want { + gotVal, ok := got[field] + if !ok { + t.Fatalf("missing field %q", field) + } + if string(gotVal) != string(wantVal) { + t.Fatalf("field %q: got %s, want %s", field, gotVal, wantVal) + } + } + if len(got) != len(tc.want) { + t.Fatalf("got %d fields, want %d", len(got), len(tc.want)) + } + }) + } +} + +func TestMergeIntoEmpty(t *testing.T) { + tr := New() + tr.Merge("new-key", json.RawMessage(`{"x":1}`)) + got := tr.Get("new-key") + if string(got) != `{"x":1}` { + t.Fatalf("expected {\"x\":1}, got %s", got) + } +} + +func TestMergeNonObjectFallback(t *testing.T) { + tr := New() + tr.Set("k", json.RawMessage(`{"a":"1"}`)) + tr.Merge("k", json.RawMessage(`"plain string"`)) + got := tr.Get("k") + if string(got) != `"plain string"` { + t.Fatalf("expected plain string fallback, got %s", got) + } +} + +func TestDelete(t *testing.T) { + tr := New() + tr.Set("k", json.RawMessage(`{"data":true}`)) + tr.Delete("k") + if got := tr.Get("k"); got != nil { + t.Fatalf("expected nil after delete, got %s", got) + } +} + +func TestDeleteMissing(t *testing.T) { + tr := New() + tr.Delete("nonexistent") +} + +func TestGetWithProvider(t *testing.T) { + tr := New() + tr.Set("k", json.RawMessage(`{"cached":"yes","kept":"ok"}`)) + tr.RegisterProvider("k", func() json.RawMessage { + return json.RawMessage(`{"live":"data","cached":"overridden"}`) + }) + + got := tr.Get("k") + var m map[string]string + if err := json.Unmarshal(got, &m); err != nil { + t.Fatalf("unmarshal: %v", err) + } + if m["live"] != "data" { + t.Fatalf("expected live=data, got %q", m["live"]) + } + if m["cached"] != "overridden" { + t.Fatalf("expected provider to override cached field, got %q", m["cached"]) + } + if m["kept"] != "ok" { + t.Fatalf("expected kept=ok preserved, got %q", m["kept"]) + } +} + +func TestGetWithProviderDoesNotMutateCache(t *testing.T) { + tr := New() + tr.Set("k", json.RawMessage(`{"a":"1"}`)) + tr.RegisterProvider("k", func() json.RawMessage { + return json.RawMessage(`{"b":"2"}`) + }) + + tr.Get("k") + + // Read the raw cached entry — remove the provider to bypass merge. + tr.RegisterProvider("k", nil) + raw := tr.Get("k") + if string(raw) != `{"a":"1"}` { + t.Fatalf("cache was mutated: %s", raw) + } +} + +func TestGetMultiWithProvider(t *testing.T) { + tr := New() + tr.Set("a", json.RawMessage(`{"x":"1"}`)) + tr.Set("b", json.RawMessage(`{"y":"2"}`)) + tr.RegisterProvider("a", func() json.RawMessage { + return json.RawMessage(`{"live":"yes"}`) + }) + + results := tr.GetMulti([]string{"a", "b"}) + if len(results) != 2 { + t.Fatalf("expected 2 results, got %d", len(results)) + } + + var m map[string]string + json.Unmarshal(results[0], &m) + if m["live"] != "yes" || m["x"] != "1" { + t.Fatalf("provider not applied to first result: %s", results[0]) + } + + // b has no provider — should return as-is + if string(results[1]) != `{"y":"2"}` { + t.Fatalf("unexpected second result: %s", results[1]) + } +} + +func TestGetWithProviderEmptyOverlay(t *testing.T) { + tr := New() + tr.Set("k", json.RawMessage(`{"a":"1"}`)) + tr.RegisterProvider("k", func() json.RawMessage { + return nil + }) + + got := tr.Get("k") + if string(got) != `{"a":"1"}` { + t.Fatalf("nil overlay should return base, got %s", got) + } +} + +func TestGetWithProviderNoBaseData(t *testing.T) { + tr := New() + tr.Set("k", json.RawMessage(nil)) + tr.RegisterProvider("k", func() json.RawMessage { + return json.RawMessage(`{"live":"yes"}`) + }) + + got := tr.Get("k") + if string(got) != `{"live":"yes"}` { + t.Fatalf("expected overlay to win with empty base, got %s", got) + } +} + +func TestConcurrentMerge(t *testing.T) { + tr := New() + tr.Set("shared", json.RawMessage(`{}`)) + + var wg sync.WaitGroup + const N = 50 + for i := 0; i < N; i++ { + wg.Add(1) + go func(i int) { + defer wg.Done() + tr.Merge("shared", json.RawMessage(`{"f`+string(rune('a'+i%26))+`":true}`)) + }(i) + } + wg.Wait() + + got := tr.Get("shared") + if got == nil { + t.Fatal("expected non-nil after concurrent merges") + } +} + +// A provider that reads the tree must not deadlock GetMulti against a +// writer waiting for the tree lock. +func TestGetMultiProviderReadsTreeWhileWriterWaits(t *testing.T) { + tr := New() + tr.Set("a", json.RawMessage(`{"x":1}`)) + tr.Set("b", json.RawMessage(`{"y":2}`)) + + inProvider := make(chan struct{}) + release := make(chan struct{}) + tr.RegisterProvider("a", func() json.RawMessage { + close(inProvider) + <-release + return tr.GetCached("b") + }) + + done := make(chan []json.RawMessage) + go func() { done <- tr.GetMulti([]string{"a", "b"}) }() + + <-inProvider + deleted := make(chan struct{}) + go func() { + tr.Delete("b") // a writer queued behind the read + close(deleted) + }() + time.Sleep(20 * time.Millisecond) + close(release) + + select { + case res := <-done: + if len(res) != 2 { + t.Fatalf("got %d results, want 2", len(res)) + } + case <-time.After(2 * time.Second): + t.Fatal("GetMulti deadlocked against a waiting writer") + } + <-deleted +} + +// Two first writers on an absent key must both land. +func TestConcurrentFirstMerge(t *testing.T) { + for i := 0; i < 200; i++ { + tr := New() + var wg sync.WaitGroup + wg.Add(2) + go func() { defer wg.Done(); tr.Merge("k", json.RawMessage(`{"a":1}`)) }() + go func() { defer wg.Done(); tr.Merge("k", json.RawMessage(`{"b":2}`)) }() + wg.Wait() + + var out map[string]int + if err := json.Unmarshal(tr.Get("k"), &out); err != nil { + t.Fatal(err) + } + if out["a"] != 1 || out["b"] != 2 { + t.Fatalf("lost a first merge: %v", out) + } + } +} diff --git a/src/yangerd/internal/unixgram/unixgram.go b/src/yangerd/internal/unixgram/unixgram.go new file mode 100644 index 000000000..3f35cc308 --- /dev/null +++ b/src/yangerd/internal/unixgram/unixgram.go @@ -0,0 +1,72 @@ +// Package unixgram dials AF_UNIX datagram servers that reply to the +// client's own bound address: ptp4l, wpa_supplicant and hostapd. The +// client socket file is created and removed here, so every caller +// cleans up the same way on every error path. +package unixgram + +import ( + "fmt" + "net" + "os" + "path/filepath" + "sync" +) + +// Conn is a connected datagram socket that owns its local socket file. +type Conn struct { + *net.UnixConn + local string +} + +// Dial binds local, replacing a stale file left by a killed process, +// and connects it to remote. A non-zero mode is applied to the local +// file, for servers that run unprivileged and must be able to reply. +func Dial(local, remote string, mode os.FileMode) (*Conn, error) { + if err := os.MkdirAll(filepath.Dir(local), 0755); err != nil { + return nil, err + } + os.Remove(local) + + laddr := &net.UnixAddr{Name: local, Net: "unixgram"} + raddr := &net.UnixAddr{Name: remote, Net: "unixgram"} + conn, err := net.DialUnix("unixgram", laddr, raddr) + if err != nil { + os.Remove(local) + return nil, fmt.Errorf("dial %s: %w", remote, err) + } + + if mode != 0 { + if err := os.Chmod(local, mode); err != nil { + conn.Close() + os.Remove(local) + return nil, err + } + } + + return &Conn{UnixConn: conn, local: local}, nil +} + +// Close closes the socket and removes the local socket file. +func (c *Conn) Close() error { + err := c.UnixConn.Close() + os.Remove(c.local) + return err +} + +var cleaned sync.Map + +// CleanDir removes every file in dir, once per process. Callers keep +// their client sockets in a directory of their own and clean it before +// first use, so sockets left behind by a killed yangerd do not pile up. +func CleanDir(dir string) { + if _, done := cleaned.LoadOrStore(dir, true); done { + return + } + entries, err := os.ReadDir(dir) + if err != nil { + return + } + for _, e := range entries { + os.Remove(filepath.Join(dir, e.Name())) + } +} diff --git a/src/yangerd/internal/unixgram/unixgram_test.go b/src/yangerd/internal/unixgram/unixgram_test.go new file mode 100644 index 000000000..9962996b1 --- /dev/null +++ b/src/yangerd/internal/unixgram/unixgram_test.go @@ -0,0 +1,78 @@ +package unixgram + +import ( + "net" + "os" + "path/filepath" + "testing" +) + +func TestDialRoundTripAndCleanup(t *testing.T) { + dir := t.TempDir() + server := filepath.Join(dir, "srv") + local := filepath.Join(dir, "client", "c1") + + srv, err := net.ListenUnixgram("unixgram", &net.UnixAddr{Name: server, Net: "unixgram"}) + if err != nil { + t.Fatal(err) + } + defer srv.Close() + + // A stale file from a killed process must not block the bind + os.MkdirAll(filepath.Dir(local), 0755) + os.WriteFile(local, nil, 0644) + + conn, err := Dial(local, server, 0666) + if err != nil { + t.Fatal(err) + } + if fi, err := os.Stat(local); err != nil || fi.Mode().Perm() != 0666 { + t.Fatalf("local socket mode = %v, %v", fi, err) + } + + if _, err := conn.Write([]byte("PING")); err != nil { + t.Fatal(err) + } + buf := make([]byte, 16) + n, from, err := srv.ReadFromUnix(buf) + if err != nil || string(buf[:n]) != "PING" { + t.Fatalf("server got %q, %v", buf[:n], err) + } + srv.WriteToUnix([]byte("PONG"), from) + if n, err = conn.Read(buf); err != nil || string(buf[:n]) != "PONG" { + t.Fatalf("client got %q, %v", buf[:n], err) + } + + conn.Close() + if _, err := os.Stat(local); !os.IsNotExist(err) { + t.Fatalf("local socket not removed on Close: %v", err) + } +} + +func TestDialFailureLeavesNoFile(t *testing.T) { + dir := t.TempDir() + local := filepath.Join(dir, "c1") + + if _, err := Dial(local, filepath.Join(dir, "nobody"), 0); err == nil { + t.Fatal("expected dial to a missing server to fail") + } + if _, err := os.Stat(local); !os.IsNotExist(err) { + t.Fatalf("local socket left behind: %v", err) + } +} + +func TestCleanDirOnce(t *testing.T) { + dir := t.TempDir() + os.WriteFile(filepath.Join(dir, "stale"), nil, 0644) + + CleanDir(dir) + if _, err := os.Stat(filepath.Join(dir, "stale")); !os.IsNotExist(err) { + t.Fatal("stale file survived the first clean") + } + + os.WriteFile(filepath.Join(dir, "live"), nil, 0644) + CleanDir(dir) + if _, err := os.Stat(filepath.Join(dir, "live")); err != nil { + t.Fatal("second clean must not remove files created since") + } +} diff --git a/src/yangerd/internal/wgquery/wgquery.go b/src/yangerd/internal/wgquery/wgquery.go new file mode 100644 index 000000000..7dc771427 --- /dev/null +++ b/src/yangerd/internal/wgquery/wgquery.go @@ -0,0 +1,272 @@ +// Package wgquery reads WireGuard peer status over the wireguard +// generic netlink family, the same channel the wg tool uses. +package wgquery + +import ( + "encoding/base64" + "encoding/binary" + "encoding/json" + "fmt" + "net" + "strconv" + "time" + + "github.com/mdlayher/genetlink" + "github.com/mdlayher/netlink" + "github.com/mdlayher/netlink/nlenc" + "golang.org/x/sys/unix" +) + +// Peer is the subset of WireGuard peer state reported in operational data. +type Peer struct { + PublicKey string + Endpoint *net.UDPAddr + LastHandshake time.Time + RxBytes uint64 + TxBytes uint64 +} + +// peersFunc returns the peers of one WireGuard interface. +type peersFunc func(ifname string) ([]Peer, error) + +// Query returns peer-status JSON per WireGuard interface found in the +// ip -j link listing, or nil when there is nothing to report. +func Query(links json.RawMessage) map[string]json.RawMessage { + wgIfaces := findWireguardIfaces(links) + if len(wgIfaces) == 0 { + return nil + } + + client, err := dial() + if err != nil { + return nil + } + defer client.Close() + + return query(wgIfaces, client.peers, time.Now().UTC()) +} + +func query(ifaces []string, peersOf peersFunc, now time.Time) map[string]json.RawMessage { + result := make(map[string]json.RawMessage) + + for _, ifname := range ifaces { + peers, err := peersOf(ifname) + if err != nil || len(peers) == 0 { + continue + } + + var out []map[string]any + for _, p := range peers { + peer := map[string]any{ + "public-key": p.PublicKey, + "connection-status": connectionStatus(p.LastHandshake, now), + } + + if !p.LastHandshake.IsZero() { + peer["latest-handshake"] = p.LastHandshake.UTC().Format("2006-01-02T15:04:05+00:00") + } + + if p.Endpoint != nil { + peer["endpoint-address"] = p.Endpoint.IP.String() + peer["endpoint-port"] = p.Endpoint.Port + } + + if p.TxBytes > 0 || p.RxBytes > 0 { + peer["transfer"] = map[string]any{ + "tx-bytes": strconv.FormatUint(p.TxBytes, 10), + "rx-bytes": strconv.FormatUint(p.RxBytes, 10), + } + } + + out = append(out, peer) + } + + data, err := json.Marshal(map[string]any{"peer-status": map[string]any{"peer": out}}) + if err != nil { + continue + } + result[ifname] = data + } + + if len(result) == 0 { + return nil + } + return result +} + +func findWireguardIfaces(links json.RawMessage) []string { + var ifaces []map[string]any + if json.Unmarshal(links, &ifaces) != nil { + return nil + } + + var result []string + for _, iface := range ifaces { + linkinfo, _ := iface["linkinfo"].(map[string]any) + if linkinfo == nil { + continue + } + if kind, _ := linkinfo["info_kind"].(string); kind == "wireguard" { + if name, _ := iface["ifname"].(string); name != "" { + result = append(result, name) + } + } + } + return result +} + +func connectionStatus(handshake time.Time, now time.Time) string { + if handshake.IsZero() { + return "down" + } + if now.Sub(handshake) < 180*time.Second { + return "up" + } + return "down" +} + +// --- generic netlink --- + +type client struct { + conn *genetlink.Conn + family genetlink.Family +} + +func dial() (*client, error) { + conn, err := genetlink.Dial(nil) + if err != nil { + return nil, fmt.Errorf("dial genetlink: %w", err) + } + + family, err := conn.GetFamily(unix.WG_GENL_NAME) + if err != nil { + _ = conn.Close() + return nil, fmt.Errorf("resolve wireguard family: %w", err) + } + + return &client{conn: conn, family: family}, nil +} + +func (c *client) Close() error { + return c.conn.Close() +} + +// peers dumps one device. The kernel splits a device with many peers +// over several messages, and a peer with many allowed IPs can repeat +// across them, so peers are merged by public key. +func (c *client) peers(ifname string) ([]Peer, error) { + req, err := netlink.MarshalAttributes([]netlink.Attribute{{ + Type: unix.WGDEVICE_A_IFNAME, + Data: nlenc.Bytes(ifname), + }}) + if err != nil { + return nil, err + } + + msgs, err := c.conn.Execute(genetlink.Message{ + Header: genetlink.Header{Command: unix.WG_CMD_GET_DEVICE, Version: unix.WG_GENL_VERSION}, + Data: req, + }, c.family.ID, netlink.Request|netlink.Dump) + if err != nil { + return nil, fmt.Errorf("wireguard get device %s: %w", ifname, err) + } + + return parsePeers(msgs) +} + +func parsePeers(msgs []genetlink.Message) ([]Peer, error) { + var peers []Peer + seen := map[string]bool{} + + for _, m := range msgs { + ad, err := netlink.NewAttributeDecoder(m.Data) + if err != nil { + return nil, err + } + + for ad.Next() { + if ad.Type() != unix.WGDEVICE_A_PEERS { + continue + } + ad.Nested(func(nad *netlink.AttributeDecoder) error { + for nad.Next() { + nad.Nested(func(pad *netlink.AttributeDecoder) error { + p := parsePeer(pad) + if p.PublicKey == "" || seen[p.PublicKey] { + return nil + } + seen[p.PublicKey] = true + peers = append(peers, p) + return nil + }) + } + return nil + }) + } + + if err := ad.Err(); err != nil { + return nil, err + } + } + + return peers, nil +} + +func parsePeer(ad *netlink.AttributeDecoder) Peer { + var p Peer + for ad.Next() { + switch ad.Type() { + case unix.WGPEER_A_PUBLIC_KEY: + p.PublicKey = base64.StdEncoding.EncodeToString(ad.Bytes()) + case unix.WGPEER_A_ENDPOINT: + p.Endpoint = parseSockaddr(ad.Bytes()) + case unix.WGPEER_A_LAST_HANDSHAKE_TIME: + p.LastHandshake = parseTimespec(ad.Bytes()) + case unix.WGPEER_A_RX_BYTES: + p.RxBytes = ad.Uint64() + case unix.WGPEER_A_TX_BYTES: + p.TxBytes = ad.Uint64() + } + } + return p +} + +// parseSockaddr decodes a raw sockaddr_in or sockaddr_in6: family, +// port in network byte order, then the address (after flowinfo for v6). +func parseSockaddr(b []byte) *net.UDPAddr { + switch len(b) { + case unix.SizeofSockaddrInet4: + return &net.UDPAddr{ + IP: net.IP(b[4:8]).To4(), + Port: int(binary.BigEndian.Uint16(b[2:4])), + } + case unix.SizeofSockaddrInet6: + ip := make(net.IP, net.IPv6len) + copy(ip, b[8:24]) + return &net.UDPAddr{ + IP: ip, + Port: int(binary.BigEndian.Uint16(b[2:4])), + } + } + return nil +} + +// parseTimespec decodes a __kernel_timespec, 32 or 64 bit fields in +// host byte order. A zero value means no handshake yet. +func parseTimespec(b []byte) time.Time { + var sec, nsec int64 + switch len(b) { + case 8: + sec = int64(int32(nlenc.Uint32(b[0:4]))) + nsec = int64(int32(nlenc.Uint32(b[4:8]))) + case 16: + sec = int64(nlenc.Uint64(b[0:8])) + nsec = int64(nlenc.Uint64(b[8:16])) + default: + return time.Time{} + } + if sec <= 0 && nsec <= 0 { + return time.Time{} + } + return time.Unix(sec, nsec) +} diff --git a/src/yangerd/internal/wgquery/wgquery_test.go b/src/yangerd/internal/wgquery/wgquery_test.go new file mode 100644 index 000000000..bcde3e970 --- /dev/null +++ b/src/yangerd/internal/wgquery/wgquery_test.go @@ -0,0 +1,185 @@ +package wgquery + +import ( + "encoding/base64" + "encoding/binary" + "encoding/json" + "errors" + "net" + "testing" + "time" + + "github.com/mdlayher/genetlink" + "github.com/mdlayher/netlink" + "github.com/mdlayher/netlink/nlenc" + "golang.org/x/sys/unix" +) + +var ( + keyA = make([]byte, 32) + keyB = append([]byte{1}, make([]byte, 31)...) +) + +func sockaddr4(t *testing.T, ip string, port int) []byte { + t.Helper() + b := make([]byte, unix.SizeofSockaddrInet4) + nlenc.PutUint16(b[0:2], unix.AF_INET) + binary.BigEndian.PutUint16(b[2:4], uint16(port)) + copy(b[4:8], net.ParseIP(ip).To4()) + return b +} + +func sockaddr6(t *testing.T, ip string, port int) []byte { + t.Helper() + b := make([]byte, unix.SizeofSockaddrInet6) + nlenc.PutUint16(b[0:2], unix.AF_INET6) + binary.BigEndian.PutUint16(b[2:4], uint16(port)) + copy(b[8:24], net.ParseIP(ip).To16()) + return b +} + +func timespec64(sec int64) []byte { + b := make([]byte, 16) + nlenc.PutUint64(b[0:8], uint64(sec)) + return b +} + +type peerAttrs struct { + key []byte + endpoint []byte + handshake []byte + rx, tx uint64 +} + +func deviceMessage(t *testing.T, peers ...peerAttrs) genetlink.Message { + t.Helper() + ae := netlink.NewAttributeEncoder() + ae.String(unix.WGDEVICE_A_IFNAME, "wg0") + ae.Nested(unix.WGDEVICE_A_PEERS, func(nae *netlink.AttributeEncoder) error { + for i, p := range peers { + nae.Nested(uint16(i), func(pae *netlink.AttributeEncoder) error { + pae.Bytes(unix.WGPEER_A_PUBLIC_KEY, p.key) + if p.endpoint != nil { + pae.Bytes(unix.WGPEER_A_ENDPOINT, p.endpoint) + } + if p.handshake != nil { + pae.Bytes(unix.WGPEER_A_LAST_HANDSHAKE_TIME, p.handshake) + } + pae.Uint64(unix.WGPEER_A_RX_BYTES, p.rx) + pae.Uint64(unix.WGPEER_A_TX_BYTES, p.tx) + return nil + }) + } + return nil + }) + data, err := ae.Encode() + if err != nil { + t.Fatal(err) + } + return genetlink.Message{Data: data} +} + +func TestParsePeers(t *testing.T) { + msg := deviceMessage(t, + peerAttrs{key: keyA, endpoint: sockaddr4(t, "192.0.2.1", 51820), handshake: timespec64(1700000000), rx: 10, tx: 20}, + peerAttrs{key: keyB, endpoint: sockaddr6(t, "2001:db8::1", 4242)}, + ) + + peers, err := parsePeers([]genetlink.Message{msg}) + if err != nil { + t.Fatal(err) + } + if len(peers) != 2 { + t.Fatalf("want 2 peers, got %d", len(peers)) + } + + a := peers[0] + if a.PublicKey != base64.StdEncoding.EncodeToString(keyA) { + t.Errorf("public key = %q", a.PublicKey) + } + if a.Endpoint == nil || a.Endpoint.IP.String() != "192.0.2.1" || a.Endpoint.Port != 51820 { + t.Errorf("endpoint = %v", a.Endpoint) + } + if !a.LastHandshake.Equal(time.Unix(1700000000, 0)) { + t.Errorf("handshake = %v", a.LastHandshake) + } + if a.RxBytes != 10 || a.TxBytes != 20 { + t.Errorf("transfer = rx %d tx %d", a.RxBytes, a.TxBytes) + } + + b := peers[1] + if b.Endpoint == nil || b.Endpoint.IP.String() != "2001:db8::1" || b.Endpoint.Port != 4242 { + t.Errorf("v6 endpoint = %v", b.Endpoint) + } + if !b.LastHandshake.IsZero() { + t.Errorf("want no handshake, got %v", b.LastHandshake) + } +} + +func TestPeersSplitAcrossMessagesAreMerged(t *testing.T) { + msgs := []genetlink.Message{ + deviceMessage(t, peerAttrs{key: keyA, rx: 1}), + deviceMessage(t, peerAttrs{key: keyA, rx: 1}, peerAttrs{key: keyB}), + } + peers, err := parsePeers(msgs) + if err != nil { + t.Fatal(err) + } + if len(peers) != 2 { + t.Fatalf("want 2 distinct peers, got %d", len(peers)) + } +} + +func TestQueryJSON(t *testing.T) { + now := time.Date(2026, 10, 5, 12, 0, 0, 0, time.UTC) + peersOf := func(ifname string) ([]Peer, error) { + switch ifname { + case "wg0": + return []Peer{ + {PublicKey: "AAAA", Endpoint: &net.UDPAddr{IP: net.ParseIP("192.0.2.1"), Port: 51820}, + LastHandshake: now.Add(-time.Minute), RxBytes: 10, TxBytes: 20}, + {PublicKey: "BBBB", LastHandshake: now.Add(-time.Hour)}, + {PublicKey: "CCCC"}, + }, nil + case "wg1": + return nil, errors.New("no such device") + } + return nil, nil + } + + out := query([]string{"wg0", "wg1", "wg2"}, peersOf, now) + if len(out) != 1 { + t.Fatalf("want wg0 only, got %v", out) + } + + var doc map[string]map[string][]map[string]any + if err := json.Unmarshal(out["wg0"], &doc); err != nil { + t.Fatal(err) + } + peers := doc["peer-status"]["peer"] + if len(peers) != 3 { + t.Fatalf("want 3 peers, got %d", len(peers)) + } + if peers[0]["connection-status"] != "up" || peers[0]["endpoint-port"] != float64(51820) || + peers[0]["latest-handshake"] != "2026-10-05T11:59:00+00:00" { + t.Errorf("peer 0 = %v", peers[0]) + } + if tr := peers[0]["transfer"].(map[string]any); tr["tx-bytes"] != "20" || tr["rx-bytes"] != "10" { + t.Errorf("transfer = %v", tr) + } + if peers[1]["connection-status"] != "down" { + t.Errorf("stale handshake should be down: %v", peers[1]) + } + if _, ok := peers[2]["latest-handshake"]; ok || peers[2]["connection-status"] != "down" { + t.Errorf("never-connected peer = %v", peers[2]) + } +} + +func TestQueryNothingToReport(t *testing.T) { + if out := query([]string{"wg0"}, func(string) ([]Peer, error) { return nil, nil }, time.Now()); out != nil { + t.Errorf("want nil, got %v", out) + } + if ifaces := findWireguardIfaces(json.RawMessage(`[{"ifname":"e0","linkinfo":{"info_kind":"bridge"}},{"ifname":"wg0","linkinfo":{"info_kind":"wireguard"}}]`)); len(ifaces) != 1 || ifaces[0] != "wg0" { + t.Errorf("ifaces = %v", ifaces) + } +} diff --git a/src/yangerd/internal/wpactrl/allstations_test.go b/src/yangerd/internal/wpactrl/allstations_test.go new file mode 100644 index 000000000..459d762b1 --- /dev/null +++ b/src/yangerd/internal/wpactrl/allstations_test.go @@ -0,0 +1,142 @@ +package wpactrl + +import ( + "net" + "strings" + "testing" + "time" +) + +// fakeHostapd serves the hostapd control protocol for station +// enumeration: STA-FIRST returns the first station block, STA-NEXT +// the one after it, and an empty datagram past the last station. +func fakeHostapd(t *testing.T, stations []string) string { + t.Helper() + + dir := t.TempDir() + serverPath := dir + "/wlan0" + + serverAddr := &net.UnixAddr{Name: serverPath, Net: "unixgram"} + server, err := net.ListenUnixgram("unixgram", serverAddr) + if err != nil { + t.Fatalf("listen: %v", err) + } + t.Cleanup(func() { server.Close() }) + + addrOf := func(block string) string { + return strings.SplitN(block, "\n", 2)[0] + } + + go func() { + buf := make([]byte, 4096) + for { + n, raddr, err := server.ReadFromUnix(buf) + if err != nil { + return + } + cmd := string(buf[:n]) + + var resp string + switch { + case cmd == "STA-FIRST": + if len(stations) > 0 { + resp = stations[0] + } + case strings.HasPrefix(cmd, "STA-NEXT "): + prev := strings.TrimPrefix(cmd, "STA-NEXT ") + for i, st := range stations { + if addrOf(st) == prev && i+1 < len(stations) { + resp = stations[i+1] + break + } + } + default: + resp = "UNKNOWN COMMAND\n" + } + server.WriteToUnix([]byte(resp), raddr) + } + }() + + return serverPath +} + +func TestAllStations(t *testing.T) { + sta1 := "02:00:00:00:00:01\nflags=[AUTH][ASSOC][AUTHORIZED]\n" + + "signal=-57\nconnected_time=120\nrx_bytes=1000\ntx_bytes=2000\n" + sta2 := "02:00:00:00:00:02\nflags=[AUTH][ASSOC][AUTHORIZED]\n" + + "signal=-78\nconnected_time=60\nrx_bytes=300\ntx_bytes=400\n" + + path := fakeHostapd(t, []string{sta1, sta2}) + + conn, err := DialTimeout(path, 2*time.Second) + if err != nil { + t.Fatalf("dial: %v", err) + } + defer conn.Close() + + stas, err := conn.AllStations() + if err != nil { + t.Fatalf("AllStations: %v", err) + } + if len(stas) != 2 { + t.Fatalf("got %d stations, want 2", len(stas)) + } + if stas[0]["addr"] != "02:00:00:00:00:01" || stas[1]["addr"] != "02:00:00:00:00:02" { + t.Errorf("addrs = %q, %q", stas[0]["addr"], stas[1]["addr"]) + } + if stas[0]["signal"] != "-57" { + t.Errorf("sta[0] signal = %q", stas[0]["signal"]) + } + if stas[1]["connected_time"] != "60" { + t.Errorf("sta[1] connected_time = %q", stas[1]["connected_time"]) + } +} + +func TestAllStationsNone(t *testing.T) { + path := fakeHostapd(t, nil) + + conn, err := DialTimeout(path, 2*time.Second) + if err != nil { + t.Fatalf("dial: %v", err) + } + defer conn.Close() + + stas, err := conn.AllStations() + if err != nil { + t.Fatalf("AllStations: %v", err) + } + if len(stas) != 0 { + t.Fatalf("got %d stations, want 0", len(stas)) + } +} + +func TestAllStationsUnsupported(t *testing.T) { + dir := t.TempDir() + serverPath := dir + "/wlan0" + + serverAddr := &net.UnixAddr{Name: serverPath, Net: "unixgram"} + server, err := net.ListenUnixgram("unixgram", serverAddr) + if err != nil { + t.Fatalf("listen: %v", err) + } + defer server.Close() + + go func() { + buf := make([]byte, 4096) + n, raddr, err := server.ReadFromUnix(buf) + if err != nil || n == 0 { + return + } + server.WriteToUnix([]byte("UNKNOWN COMMAND\n"), raddr) + }() + + conn, err := DialTimeout(serverPath, 2*time.Second) + if err != nil { + t.Fatalf("dial: %v", err) + } + defer conn.Close() + + if _, err := conn.AllStations(); err == nil { + t.Fatal("expected error for UNKNOWN COMMAND") + } +} diff --git a/src/yangerd/internal/wpactrl/attach.go b/src/yangerd/internal/wpactrl/attach.go new file mode 100644 index 000000000..d0fe6cc7c --- /dev/null +++ b/src/yangerd/internal/wpactrl/attach.go @@ -0,0 +1,176 @@ +package wpactrl + +import ( + "context" + "errors" + "fmt" + "net" + "strings" + "time" + + "github.com/kernelkit/infix/src/yangerd/internal/unixgram" +) + +const attachBufSize = 4096 + +// Event is an unsolicited event from wpa_supplicant or hostapd, +// received after sending the ATTACH command. +type Event struct { + Priority int + Name string + Data string + Raw string +} + +// EventHandler is called for each unsolicited event. +type EventHandler func(Event) + +// AttachConn is a persistent event listener on a wpa_supplicant or +// hostapd control socket. After sending ATTACH, the daemon pushes +// unsolicited events like CTRL-EVENT-SIGNAL-CHANGE, AP-STA-CONNECTED, +// etc. The connection reads these in a loop and dispatches them to a +// handler. +type AttachConn struct { + conn *unixgram.Conn + handler EventHandler +} + +// Attach connects to the control socket at serverPath and sends the +// ATTACH command. On success, the daemon will send unsolicited events +// to this connection. Call Run to start reading them. +func Attach(serverPath string) (*AttachConn, error) { + conn, err := dial("a", serverPath) + if err != nil { + return nil, err + } + + conn.SetDeadline(time.Now().Add(DefaultTimeout)) + if _, err := conn.Write([]byte("ATTACH")); err != nil { + conn.Close() + return nil, fmt.Errorf("send ATTACH: %w", err) + } + + buf := make([]byte, 64) + n, err := conn.Read(buf) + if err != nil { + conn.Close() + return nil, fmt.Errorf("read ATTACH response: %w", err) + } + resp := strings.TrimSpace(string(buf[:n])) + if resp != "OK" { + conn.Close() + return nil, fmt.Errorf("ATTACH rejected: %q", resp) + } + + conn.SetDeadline(time.Time{}) + return &AttachConn{conn: conn}, nil +} + +// SetHandler sets the callback for received events. +func (a *AttachConn) SetHandler(fn EventHandler) { + a.handler = fn +} + +// Run reads events until ctx is cancelled or the socket errors (daemon +// died). Returns nil on context cancellation, error on socket failure. +// A datagram socket says nothing when its peer goes away, and hostapd or +// wpa_supplicant being restarted under an attach is routine. When the +// socket has been quiet this long, PING it; no PONG in time means the +// daemon we are attached to is gone. +var ( + pingQuiet = 10 * time.Second + pongWait = 3 * time.Second +) + +// ErrPeerGone is returned by Run when the daemon stopped answering. +var ErrPeerGone = errors.New("control socket peer is gone") + +func (a *AttachConn) Run(ctx context.Context) error { + done := make(chan struct{}) + go func() { + select { + case <-ctx.Done(): + a.conn.SetReadDeadline(time.Now()) + case <-done: + } + }() + defer close(done) + + buf := make([]byte, attachBufSize) + waiting := false + a.conn.SetReadDeadline(time.Now().Add(pingQuiet)) + for { + n, err := a.conn.Read(buf) + if ctx.Err() != nil { + return nil + } + if err != nil { + var ne net.Error + if !errors.As(err, &ne) || !ne.Timeout() { + return fmt.Errorf("read: %w", err) + } + if waiting { + return ErrPeerGone + } + if _, err := a.conn.Write([]byte("PING")); err != nil { + return fmt.Errorf("%w: %v", ErrPeerGone, err) + } + waiting = true + a.conn.SetReadDeadline(time.Now().Add(pongWait)) + continue + } + + waiting = false + a.conn.SetReadDeadline(time.Now().Add(pingQuiet)) + msg := string(buf[:n]) + if strings.TrimSpace(msg) == "PONG" || a.handler == nil { + continue + } + if ev, ok := ParseEvent(msg); ok { + a.handler(ev) + } + } +} + +// Close sends DETACH and closes the connection. +func (a *AttachConn) Close() error { + a.conn.SetDeadline(time.Now().Add(DefaultTimeout)) + a.conn.Write([]byte("DETACH")) + return a.conn.Close() +} + +// ParseEvent parses a single unsolicited event line. Format: +// EVENT-NAME optional-data +// where N is a priority digit (0-4). Some events like +// AP-STA-CONNECTED have no priority prefix. +func ParseEvent(line string) (Event, bool) { + line = strings.TrimSpace(line) + if line == "" { + return Event{}, false + } + + ev := Event{Raw: line} + + if len(line) >= 3 && line[0] == '<' { + end := strings.IndexByte(line, '>') + if end > 1 { + for _, c := range line[1:end] { + if c < '0' || c > '9' { + goto noPriority + } + } + fmt.Sscanf(line[1:end], "%d", &ev.Priority) + line = line[end+1:] + } + } +noPriority: + + if idx := strings.IndexByte(line, ' '); idx > 0 { + ev.Name = line[:idx] + ev.Data = line[idx+1:] + } else { + ev.Name = line + } + + return ev, ev.Name != "" +} diff --git a/src/yangerd/internal/wpactrl/attach_test.go b/src/yangerd/internal/wpactrl/attach_test.go new file mode 100644 index 000000000..f377ce629 --- /dev/null +++ b/src/yangerd/internal/wpactrl/attach_test.go @@ -0,0 +1,254 @@ +package wpactrl + +import ( + "context" + "errors" + "net" + "testing" + "time" +) + +func TestParseEvent(t *testing.T) { + tests := []struct { + line string + wantOK bool + wantPri int + wantName string + wantData string + }{ + {"<3>CTRL-EVENT-SIGNAL-CHANGE above=0 signal=-88 noise=-92 txrate=6000", true, 3, "CTRL-EVENT-SIGNAL-CHANGE", "above=0 signal=-88 noise=-92 txrate=6000"}, + {"<3>CTRL-EVENT-CONNECTED - Connection to 02:00:00:00:01:00 completed", true, 3, "CTRL-EVENT-CONNECTED", "- Connection to 02:00:00:00:01:00 completed"}, + {"<3>CTRL-EVENT-SCAN-RESULTS ", true, 3, "CTRL-EVENT-SCAN-RESULTS", ""}, + {"<3>CTRL-EVENT-DISCONNECTED bssid=02:00:00:00:01:00 reason=3", true, 3, "CTRL-EVENT-DISCONNECTED", "bssid=02:00:00:00:01:00 reason=3"}, + {"<2>AP-STA-CONNECTED 9e:61:6b:cf:d8:15", true, 2, "AP-STA-CONNECTED", "9e:61:6b:cf:d8:15"}, + {"<2>AP-STA-DISCONNECTED 9e:61:6b:cf:d8:15", true, 2, "AP-STA-DISCONNECTED", "9e:61:6b:cf:d8:15"}, + {"AP-STA-CONNECTED 9e:61:6b:cf:d8:15", true, 0, "AP-STA-CONNECTED", "9e:61:6b:cf:d8:15"}, + {"<3>CTRL-EVENT-TERMINATING", true, 3, "CTRL-EVENT-TERMINATING", ""}, + {"", false, 0, "", ""}, + {" ", false, 0, "", ""}, + } + + for _, tt := range tests { + ev, ok := ParseEvent(tt.line) + if ok != tt.wantOK { + t.Errorf("ParseEvent(%q): ok=%v, want %v", tt.line, ok, tt.wantOK) + continue + } + if !ok { + continue + } + if ev.Priority != tt.wantPri { + t.Errorf("ParseEvent(%q): priority=%d, want %d", tt.line, ev.Priority, tt.wantPri) + } + if ev.Name != tt.wantName { + t.Errorf("ParseEvent(%q): name=%q, want %q", tt.line, ev.Name, tt.wantName) + } + if ev.Data != tt.wantData { + t.Errorf("ParseEvent(%q): data=%q, want %q", tt.line, ev.Data, tt.wantData) + } + } +} + +func TestAttachAndReceiveEvents(t *testing.T) { + dir := t.TempDir() + serverPath := dir + "/hostapd_test" + + serverAddr := &net.UnixAddr{Name: serverPath, Net: "unixgram"} + server, err := net.ListenUnixgram("unixgram", serverAddr) + if err != nil { + t.Fatalf("listen: %v", err) + } + defer server.Close() + + go func() { + buf := make([]byte, 4096) + n, raddr, err := server.ReadFromUnix(buf) + if err != nil { + return + } + if string(buf[:n]) == "ATTACH" { + server.WriteToUnix([]byte("OK\n"), raddr) + } + time.Sleep(50 * time.Millisecond) + server.WriteToUnix([]byte("<2>AP-STA-CONNECTED 9e:61:6b:cf:d8:15"), raddr) + time.Sleep(50 * time.Millisecond) + server.WriteToUnix([]byte("<3>CTRL-EVENT-SIGNAL-CHANGE above=0 signal=-55"), raddr) + + n, _, err = server.ReadFromUnix(buf) + if err != nil { + return + } + if string(buf[:n]) == "DETACH" { + server.WriteToUnix([]byte("OK\n"), raddr) + } + }() + + ac, err := Attach(serverPath) + if err != nil { + t.Fatalf("Attach: %v", err) + } + defer ac.Close() + + var events []Event + ac.SetHandler(func(ev Event) { + events = append(events, ev) + }) + + ctx, cancel := context.WithTimeout(context.Background(), 500*time.Millisecond) + defer cancel() + + ac.Run(ctx) + + if len(events) < 2 { + t.Fatalf("got %d events, want >= 2", len(events)) + } + if events[0].Name != "AP-STA-CONNECTED" { + t.Errorf("events[0].Name = %q, want AP-STA-CONNECTED", events[0].Name) + } + if events[0].Data != "9e:61:6b:cf:d8:15" { + t.Errorf("events[0].Data = %q", events[0].Data) + } + if events[1].Name != "CTRL-EVENT-SIGNAL-CHANGE" { + t.Errorf("events[1].Name = %q, want CTRL-EVENT-SIGNAL-CHANGE", events[1].Name) + } + + ac.conn.Close() +} + +func TestAttachContextCancel(t *testing.T) { + dir := t.TempDir() + serverPath := dir + "/wpa_test" + + serverAddr := &net.UnixAddr{Name: serverPath, Net: "unixgram"} + server, err := net.ListenUnixgram("unixgram", serverAddr) + if err != nil { + t.Fatalf("listen: %v", err) + } + defer server.Close() + + go func() { + buf := make([]byte, 4096) + n, raddr, err := server.ReadFromUnix(buf) + if err != nil { + return + } + if string(buf[:n]) == "ATTACH" { + server.WriteToUnix([]byte("OK\n"), raddr) + } + n, _, _ = server.ReadFromUnix(buf) + }() + + ac, err := Attach(serverPath) + if err != nil { + t.Fatalf("Attach: %v", err) + } + defer ac.Close() + + ac.SetHandler(func(ev Event) {}) + + ctx, cancel := context.WithTimeout(context.Background(), 200*time.Millisecond) + defer cancel() + + err = ac.Run(ctx) + if err != nil { + t.Errorf("expected nil on context cancel, got %v", err) + } +} + +// fakeDaemon answers ATTACH, then PING with PONG while answer is true. +func fakeDaemon(t *testing.T, answer bool) (string, *net.UnixConn) { + t.Helper() + path := t.TempDir() + "/hostapd_test" + server, err := net.ListenUnixgram("unixgram", &net.UnixAddr{Name: path, Net: "unixgram"}) + if err != nil { + t.Fatalf("listen: %v", err) + } + t.Cleanup(func() { server.Close() }) + go func() { + buf := make([]byte, 4096) + for { + n, raddr, err := server.ReadFromUnix(buf) + if err != nil { + return + } + switch string(buf[:n]) { + case "ATTACH": + server.WriteToUnix([]byte("OK\n"), raddr) + case "PING": + if answer { + server.WriteToUnix([]byte("PONG\n"), raddr) + } + } + } + }() + return path, server +} + +func shortKeepalive(t *testing.T) { + t.Helper() + oldQuiet, oldWait := pingQuiet, pongWait + pingQuiet, pongWait = 50*time.Millisecond, 50*time.Millisecond + t.Cleanup(func() { pingQuiet, pongWait = oldQuiet, oldWait }) +} + +func runAttach(t *testing.T, path string) (context.CancelFunc, <-chan error) { + t.Helper() + ac, err := Attach(path) + if err != nil { + t.Fatalf("attach: %v", err) + } + t.Cleanup(func() { ac.Close() }) + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan error, 1) + go func() { done <- ac.Run(ctx) }() + return cancel, done +} + +func TestAttachKeepaliveAnswered(t *testing.T) { + shortKeepalive(t) + path, _ := fakeDaemon(t, true) + cancel, done := runAttach(t, path) + + select { + case err := <-done: + t.Fatalf("Run gave up on a live daemon: %v", err) + case <-time.After(400 * time.Millisecond): + } + cancel() + if err := <-done; err != nil { + t.Fatalf("Run on cancel = %v, want nil", err) + } +} + +func TestAttachKeepaliveSilentPeer(t *testing.T) { + shortKeepalive(t) + path, _ := fakeDaemon(t, false) + cancel, done := runAttach(t, path) + defer cancel() + + select { + case err := <-done: + if !errors.Is(err, ErrPeerGone) { + t.Fatalf("Run = %v, want ErrPeerGone", err) + } + case <-time.After(time.Second): + t.Fatal("Run kept waiting on a daemon that never answers") + } +} + +func TestAttachKeepalivePeerSocketGone(t *testing.T) { + shortKeepalive(t) + path, server := fakeDaemon(t, true) + cancel, done := runAttach(t, path) + defer cancel() + + server.Close() + select { + case err := <-done: + if !errors.Is(err, ErrPeerGone) { + t.Fatalf("Run = %v, want ErrPeerGone", err) + } + case <-time.After(time.Second): + t.Fatal("Run kept waiting after the daemon's socket went away") + } +} diff --git a/src/yangerd/internal/wpactrl/main_test.go b/src/yangerd/internal/wpactrl/main_test.go new file mode 100644 index 000000000..e54c50e8c --- /dev/null +++ b/src/yangerd/internal/wpactrl/main_test.go @@ -0,0 +1,17 @@ +package wpactrl + +import ( + "os" + "testing" +) + +func TestMain(m *testing.M) { + dir, err := os.MkdirTemp("", "wpactrl") + if err != nil { + panic(err) + } + localDir = dir + code := m.Run() + os.RemoveAll(dir) + os.Exit(code) +} diff --git a/src/yangerd/internal/wpactrl/parse.go b/src/yangerd/internal/wpactrl/parse.go new file mode 100644 index 000000000..58fe57737 --- /dev/null +++ b/src/yangerd/internal/wpactrl/parse.go @@ -0,0 +1,94 @@ +package wpactrl + +import ( + "strconv" + "strings" +) + +// ScanResult is a single entry from SCAN_RESULTS. +type ScanResult struct { + BSSID string + Frequency int + Signal int + Flags string + SSID string +} + +// ParseKV parses a wpa_supplicant/hostapd key=value response. +func ParseKV(resp string) map[string]string { + m := make(map[string]string) + for _, line := range strings.Split(resp, "\n") { + line = strings.TrimSpace(line) + if idx := strings.IndexByte(line, '='); idx > 0 { + m[line[:idx]] = line[idx+1:] + } + } + return m +} + +// ParseScanResults parses wpa_supplicant SCAN_RESULTS output. +// Format: bssid / frequency / signal level / flags / ssid +// First line is a header, subsequent lines are tab-separated. +func ParseScanResults(resp string) []ScanResult { + var results []ScanResult + for _, line := range strings.Split(resp, "\n") { + line = strings.TrimSpace(line) + if line == "" || strings.HasPrefix(line, "bssid") { + continue + } + fields := strings.SplitN(line, "\t", 5) + if len(fields) < 4 { + continue + } + freq, _ := strconv.Atoi(fields[1]) + sig, _ := strconv.Atoi(fields[2]) + ssid := "" + if len(fields) >= 5 { + ssid = fields[4] + } + results = append(results, ScanResult{ + BSSID: fields[0], + Frequency: freq, + Signal: sig, + Flags: fields[3], + SSID: ssid, + }) + } + return results +} + +// ParseStationResp parses a hostapd STA-FIRST/STA-NEXT response. +// First line is the station MAC, subsequent lines are key=value pairs. +func ParseStationResp(resp string) map[string]string { + lines := strings.Split(resp, "\n") + if len(lines) == 0 { + return nil + } + m := make(map[string]string) + addr := strings.TrimSpace(lines[0]) + if addr != "" { + m["addr"] = addr + } + for _, line := range lines[1:] { + line = strings.TrimSpace(line) + if idx := strings.IndexByte(line, '='); idx > 0 { + m[line[:idx]] = line[idx+1:] + } + } + return m +} + +// FrequencyToChannel converts a WiFi frequency in MHz to a channel number. +func FrequencyToChannel(freq int) int { + switch { + case freq == 2484: + return 14 + case freq >= 2412 && freq <= 2472: + return (freq-2412)/5 + 1 + case freq >= 5170 && freq <= 5825: + return (freq - 5000) / 5 + case freq >= 5955 && freq <= 7115: + return (freq - 5950) / 5 + } + return 0 +} diff --git a/src/yangerd/internal/wpactrl/wpactrl.go b/src/yangerd/internal/wpactrl/wpactrl.go new file mode 100644 index 000000000..b85f4d8b8 --- /dev/null +++ b/src/yangerd/internal/wpactrl/wpactrl.go @@ -0,0 +1,199 @@ +// Package wpactrl provides a native Go client for wpa_supplicant and +// hostapd control sockets. It speaks the same text-based protocol as +// wpa_cli/hostapd_cli — Unix datagram sockets with ASCII +// command/response framing. No subprocess, no CGo. +// +// wpa_supplicant listens at /var/run/wpa_supplicant/ +// hostapd listens at /var/run/hostapd/ +// +// The client binds its own socket, sends a command string, and reads +// back the text response. +package wpactrl + +import ( + "fmt" + "os" + "path/filepath" + "strings" + "sync/atomic" + "time" + + "github.com/kernelkit/infix/src/yangerd/internal/unixgram" +) + +const ( + DefaultTimeout = 5 * time.Second + + maxResponse = 64 * 1024 +) + +// WPADirs lists directories where wpa_supplicant control sockets may live. +var WPADirs = []string{"/run/wpa_supplicant", "/var/run/wpa_supplicant"} + +// HostapdDirs lists directories where hostapd control sockets may live. +var HostapdDirs = []string{"/run/hostapd", "/var/run/hostapd"} + +var clientSeq atomic.Uint64 + +// localDir holds our client sockets, cleaned before first use so ones +// left by a killed yangerd do not pile up. +var localDir = "/run/yangerd/wpactrl" + +// dial binds a fresh client socket and connects it to serverPath. +func dial(kind, serverPath string) (*unixgram.Conn, error) { + unixgram.CleanDir(localDir) + local := fmt.Sprintf("%s/%s%d", localDir, kind, clientSeq.Add(1)) + return unixgram.Dial(local, serverPath, 0) +} + +// SocketInfo describes a discovered control socket. +type SocketInfo struct { + Path string + Iface string + Daemon string // "wpa_supplicant" or "hostapd" +} + +// ScanSockets discovers wpa_supplicant and hostapd control sockets by +// listing the well-known directories. Returns a map from interface +// name to SocketInfo. +func ScanSockets() map[string]SocketInfo { + result := make(map[string]SocketInfo) + for _, dir := range HostapdDirs { + scanDir(dir, "hostapd", result) + } + for _, dir := range WPADirs { + scanDir(dir, "wpa_supplicant", result) + } + return result +} + +func scanDir(dir, daemon string, out map[string]SocketInfo) { + entries, err := os.ReadDir(dir) + if err != nil { + return + } + for _, e := range entries { + name := e.Name() + if _, exists := out[name]; exists { + continue + } + path := filepath.Join(dir, name) + fi, err := os.Stat(path) + if err != nil { + continue + } + if fi.Mode()&os.ModeSocket != 0 { + out[name] = SocketInfo{ + Path: path, + Iface: name, + Daemon: daemon, + } + } + } +} + +// Conn is a connection to a wpa_supplicant or hostapd control socket. +type Conn struct { + conn *unixgram.Conn + timeout time.Duration +} + +// Dial connects to a wpa_supplicant or hostapd control socket at the +// given path (e.g. "/var/run/wpa_supplicant/wlan0"). The caller must +// call Close when done. +func Dial(serverPath string) (*Conn, error) { + return DialTimeout(serverPath, DefaultTimeout) +} + +// DialTimeout connects with a custom timeout. +func DialTimeout(serverPath string, timeout time.Duration) (*Conn, error) { + conn, err := dial("c", serverPath) + if err != nil { + return nil, err + } + return &Conn{conn: conn, timeout: timeout}, nil +} + +// Close closes the connection and removes the client socket file. +func (c *Conn) Close() error { + return c.conn.Close() +} + +// Command sends a command string and returns the response. +func (c *Conn) Command(cmd string) (string, error) { + c.conn.SetDeadline(time.Now().Add(c.timeout)) + + _, err := c.conn.Write([]byte(cmd)) + if err != nil { + return "", fmt.Errorf("write %q: %w", cmd, err) + } + + buf := make([]byte, maxResponse) + n, err := c.conn.Read(buf) + if err != nil { + return "", fmt.Errorf("read response to %q: %w", cmd, err) + } + + return string(buf[:n]), nil +} + +// Status sends the STATUS command and returns the parsed key=value pairs. +func (c *Conn) Status() (map[string]string, error) { + resp, err := c.Command("STATUS") + if err != nil { + return nil, err + } + return ParseKV(resp), nil +} + +// SignalPoll sends SIGNAL_POLL and returns parsed key=value pairs. +// Returns RSSI, LINKSPEED, NOISE, FREQUENCY, etc. +// Only meaningful for wpa_supplicant (station mode). +func (c *Conn) SignalPoll() (map[string]string, error) { + resp, err := c.Command("SIGNAL_POLL") + if err != nil { + return nil, err + } + return ParseKV(resp), nil +} + +// ScanResults sends SCAN_RESULTS and returns parsed results. +// This is only meaningful for wpa_supplicant (station mode). +func (c *Conn) ScanResults() ([]ScanResult, error) { + resp, err := c.Command("SCAN_RESULTS") + if err != nil { + return nil, err + } + return ParseScanResults(resp), nil +} + +// AllStations enumerates all associated stations via STA-FIRST/STA-NEXT. +// Only meaningful for hostapd. +func (c *Conn) AllStations() ([]map[string]string, error) { + resp, err := c.Command("STA-FIRST") + if err != nil { + return nil, fmt.Errorf("STA-FIRST: %w", err) + } + if resp == "" || resp == "\n" || resp == "FAIL\n" { + return nil, nil + } + if strings.HasPrefix(resp, "UNKNOWN") { + return nil, fmt.Errorf("STA-FIRST not supported: %q", strings.TrimSpace(resp)) + } + + var stations []map[string]string + st := ParseStationResp(resp) + for st != nil { + stations = append(stations, st) + addr := st["addr"] + if addr == "" { + break + } + resp, err = c.Command("STA-NEXT " + addr) + if err != nil || resp == "" || resp == "\n" || resp == "FAIL\n" || strings.HasPrefix(resp, "UNKNOWN") { + break + } + st = ParseStationResp(resp) + } + return stations, nil +} diff --git a/src/yangerd/internal/wpactrl/wpactrl_test.go b/src/yangerd/internal/wpactrl/wpactrl_test.go new file mode 100644 index 000000000..6677e63cf --- /dev/null +++ b/src/yangerd/internal/wpactrl/wpactrl_test.go @@ -0,0 +1,191 @@ +package wpactrl + +import ( + "net" + "strings" + "testing" + "time" +) + +func TestParseKV(t *testing.T) { + resp := `bssid=02:00:00:00:01:00 +freq=2412 +ssid=TestNetwork +id=0 +mode=station +pairwise_cipher=CCMP +group_cipher=CCMP +key_mgmt=WPA2-PSK +wpa_state=COMPLETED +address=02:00:00:00:00:01 +` + m := ParseKV(resp) + if m["ssid"] != "TestNetwork" { + t.Errorf("ssid = %q, want TestNetwork", m["ssid"]) + } + if m["wpa_state"] != "COMPLETED" { + t.Errorf("wpa_state = %q, want COMPLETED", m["wpa_state"]) + } + if m["freq"] != "2412" { + t.Errorf("freq = %q, want 2412", m["freq"]) + } + if m["mode"] != "station" { + t.Errorf("mode = %q, want station", m["mode"]) + } +} + +func TestParseKVEmpty(t *testing.T) { + m := ParseKV("") + if len(m) != 0 { + t.Errorf("expected empty map, got %v", m) + } +} + +func TestParseScanResults(t *testing.T) { + resp := "bssid / frequency / signal level / flags / ssid\n" + + "02:00:00:00:01:00\t2412\t-50\t[WPA2-PSK-CCMP][ESS]\tMyNetwork\n" + + "02:00:00:00:02:00\t5180\t-70\t[WPA2-EAP-CCMP][ESS]\tOffice\n" + + "02:00:00:00:03:00\t2437\t-85\t[ESS]\t\n" + + results := ParseScanResults(resp) + if len(results) != 3 { + t.Fatalf("got %d results, want 3", len(results)) + } + + r := results[0] + if r.BSSID != "02:00:00:00:01:00" { + t.Errorf("bssid = %q", r.BSSID) + } + if r.Frequency != 2412 { + t.Errorf("freq = %d, want 2412", r.Frequency) + } + if r.Signal != -50 { + t.Errorf("signal = %d, want -50", r.Signal) + } + if r.SSID != "MyNetwork" { + t.Errorf("ssid = %q, want MyNetwork", r.SSID) + } + + if results[1].Frequency != 5180 { + t.Errorf("results[1].freq = %d, want 5180", results[1].Frequency) + } +} + +func TestParseScanResultsEmpty(t *testing.T) { + results := ParseScanResults("bssid / frequency / signal level / flags / ssid\n") + if len(results) != 0 { + t.Errorf("expected empty, got %d", len(results)) + } +} + +func TestParseStationResp(t *testing.T) { + resp := "02:00:00:00:00:01\nflags=[AUTH][ASSOC][AUTHORIZED]\naid=1\n" + + "rx_bytes=12345\ntx_bytes=67890\nconnected_time=120\n" + + m := ParseStationResp(resp) + if m["addr"] != "02:00:00:00:00:01" { + t.Errorf("addr = %q", m["addr"]) + } + if m["rx_bytes"] != "12345" { + t.Errorf("rx_bytes = %q", m["rx_bytes"]) + } + if m["connected_time"] != "120" { + t.Errorf("connected_time = %q", m["connected_time"]) + } +} + +func TestParseStationRespEmpty(t *testing.T) { + m := ParseStationResp("") + if m["addr"] != "" { + t.Errorf("expected empty addr, got %q", m["addr"]) + } +} + +func TestFrequencyToChannel(t *testing.T) { + tests := []struct { + freq int + ch int + }{ + {2412, 1}, + {2437, 6}, + {2462, 11}, + {2484, 14}, + {5180, 36}, + {5240, 48}, + {5745, 149}, + {5825, 165}, + {5955, 1}, + {6115, 33}, + {1000, 0}, + } + for _, tt := range tests { + got := FrequencyToChannel(tt.freq) + if got != tt.ch { + t.Errorf("FrequencyToChannel(%d) = %d, want %d", tt.freq, got, tt.ch) + } + } +} + +func TestDialAndCommand(t *testing.T) { + dir := t.TempDir() + serverPath := dir + "/test_server" + clientDone := make(chan struct{}) + + serverAddr := &net.UnixAddr{Name: serverPath, Net: "unixgram"} + server, err := net.ListenUnixgram("unixgram", serverAddr) + if err != nil { + t.Fatalf("listen: %v", err) + } + defer server.Close() + + go func() { + defer close(clientDone) + buf := make([]byte, 4096) + n, raddr, err := server.ReadFromUnix(buf) + if err != nil { + t.Errorf("server read: %v", err) + return + } + cmd := string(buf[:n]) + var resp string + switch cmd { + case "PING": + resp = "PONG\n" + case "STATUS": + resp = "wpa_state=COMPLETED\nssid=Test\n" + default: + resp = "UNKNOWN COMMAND\n" + } + server.WriteToUnix([]byte(resp), raddr) + + n, raddr, err = server.ReadFromUnix(buf) + if err != nil { + t.Errorf("server read 2: %v", err) + return + } + if string(buf[:n]) == "STATUS" { + server.WriteToUnix([]byte("wpa_state=COMPLETED\nssid=Test\n"), raddr) + } + }() + + conn, err := DialTimeout(serverPath, 2*time.Second) + if err != nil { + t.Fatalf("dial: %v", err) + } + defer conn.Close() + + if resp, err := conn.Command("PING"); err != nil || !strings.HasPrefix(resp, "PONG") { + t.Errorf("PING = %q, %v", resp, err) + } + + status, err := conn.Status() + if err != nil { + t.Fatalf("Status: %v", err) + } + if status["ssid"] != "Test" { + t.Errorf("ssid = %q, want Test", status["ssid"]) + } + + <-clientDone + conn.Close() +} diff --git a/src/yangerd/internal/zapi/zapi.go b/src/yangerd/internal/zapi/zapi.go new file mode 100644 index 000000000..dce1d55e5 --- /dev/null +++ b/src/yangerd/internal/zapi/zapi.go @@ -0,0 +1,177 @@ +// Package zapi implements a minimal ZAPI v6 client for FRR 10.5. +// +// It speaks only the subset of the Zebra wire protocol needed by +// yangerd: Hello, RouterIDAdd, RedistributeAdd, and reading message +// headers, which is all a change trigger needs. +package zapi + +import ( + "encoding/binary" + "fmt" + "io" +) + +// Wire constants for ZAPI v6. +const ( + HeaderSize = 10 + HeaderMarker = 0xFE + HeaderVersion = 6 + + DefaultVrf uint32 = 0 +) + +// Command IDs for FRR 10.5 ZAPI v6 (from lib/zclient.h). +type Command uint16 + +const ( + CmdInterfaceAdd Command = 0 + CmdInterfaceDelete Command = 1 + CmdInterfaceAddrAdd Command = 2 + CmdInterfaceAddrDelete Command = 3 + CmdInterfaceUp Command = 4 + CmdInterfaceDown Command = 5 + CmdInterfaceSetMaster Command = 6 + CmdInterfaceSetARP Command = 7 // new in FRR 10.x + CmdInterfaceSetProtodown Command = 8 + CmdRouteAdd Command = 9 + CmdRouteDelete Command = 10 + CmdRouteNotifyOwner Command = 11 + CmdRedistributeAdd Command = 12 + CmdRedistributeDelete Command = 13 + CmdRedistDefaultAdd Command = 14 + CmdRedistDefaultDelete Command = 15 + CmdRouterIDAdd Command = 16 + CmdRouterIDDelete Command = 17 + CmdRouterIDUpdate Command = 18 + CmdHello Command = 19 + CmdCapabilities Command = 20 + CmdNexthopRegister Command = 21 + CmdNexthopUnregister Command = 22 + CmdNexthopUpdate Command = 23 + + CmdRedistRouteAdd Command = 31 + CmdRedistRouteDel Command = 32 +) + +// RouteType identifies the source protocol of a route. +type RouteType uint8 + +const ( + RouteSystem RouteType = 0 + RouteKernel RouteType = 1 + RouteConnect RouteType = 2 + RouteLocal RouteType = 3 + RouteStatic RouteType = 4 + RouteRIP RouteType = 5 + RouteRIPNG RouteType = 6 + RouteOSPF RouteType = 7 + RouteOSPF6 RouteType = 8 + RouteISIS RouteType = 9 + RouteBGP RouteType = 10 +) + +// AFI values. +const ( + AFIIPv4 uint8 = 1 + AFIIPv6 uint8 = 2 +) + +// Header is a ZAPI v6 message header. +type Header struct { + Length uint16 + Marker uint8 + Version uint8 + VrfID uint32 + Command Command +} + +// EncodeHeader serializes a ZAPI v6 header. +func EncodeHeader(length uint16, vrfID uint32, cmd Command) []byte { + buf := make([]byte, HeaderSize) + binary.BigEndian.PutUint16(buf[0:2], length) + buf[2] = HeaderMarker + buf[3] = HeaderVersion + binary.BigEndian.PutUint32(buf[4:8], vrfID) + binary.BigEndian.PutUint16(buf[8:10], uint16(cmd)) + return buf +} + +// DecodeHeader parses a ZAPI v6 header from exactly HeaderSize bytes. +func DecodeHeader(data []byte) (Header, error) { + if len(data) < HeaderSize { + return Header{}, fmt.Errorf("header too short: %d bytes", len(data)) + } + h := Header{ + Length: binary.BigEndian.Uint16(data[0:2]), + Marker: data[2], + Version: data[3], + VrfID: binary.BigEndian.Uint32(data[4:8]), + Command: Command(binary.BigEndian.Uint16(data[8:10])), + } + if h.Marker != HeaderMarker { + return Header{}, fmt.Errorf("bad marker: 0x%02x", h.Marker) + } + if h.Version != HeaderVersion { + return Header{}, fmt.Errorf("unsupported version: %d", h.Version) + } + return h, nil +} + +// EncodeHello builds a Hello message body. +// Fields: redistDefault(1), instance(2), sessionID(4), synchronous(1) = 8 bytes. +// We send zeros for everything (redistDefault=0 means ZEBRA_ROUTE_SYSTEM). +func EncodeHello() []byte { + return make([]byte, 8) +} + +// EncodeRouterIDAdd builds a RouterIDAdd message body. +// Body is just the AFI value (1 byte). +func EncodeRouterIDAdd(afi uint8) []byte { + return []byte{afi} +} + +// EncodeRedistributeAdd builds a RedistributeAdd body. +// Body: afi(1), routeType(1), instance(2). +func EncodeRedistributeAdd(afi uint8, rt RouteType) []byte { + buf := make([]byte, 4) + buf[0] = afi + buf[1] = uint8(rt) + // instance = 0 (already zeroed) + return buf +} + +// BuildMessage constructs a complete wire message from command and body. +func BuildMessage(cmd Command, vrfID uint32, body []byte) []byte { + length := uint16(HeaderSize + len(body)) + hdr := EncodeHeader(length, vrfID, cmd) + return append(hdr, body...) +} + +// ReadMessage reads one complete ZAPI message from the connection. +// It returns the header and the raw body bytes. +func ReadMessage(r io.Reader) (Header, []byte, error) { + hdrBuf := make([]byte, HeaderSize) + if _, err := io.ReadFull(r, hdrBuf); err != nil { + return Header{}, nil, fmt.Errorf("read header: %w", err) + } + + hdr, err := DecodeHeader(hdrBuf) + if err != nil { + return Header{}, nil, err + } + + bodyLen := int(hdr.Length) - HeaderSize + if bodyLen < 0 { + return Header{}, nil, fmt.Errorf("invalid message length: %d", hdr.Length) + } + if bodyLen == 0 { + return hdr, nil, nil + } + + body := make([]byte, bodyLen) + if _, err := io.ReadFull(r, body); err != nil { + return Header{}, nil, fmt.Errorf("read body: %w", err) + } + + return hdr, body, nil +} diff --git a/src/yangerd/internal/zapi/zapi_test.go b/src/yangerd/internal/zapi/zapi_test.go new file mode 100644 index 000000000..f98b5cc9a --- /dev/null +++ b/src/yangerd/internal/zapi/zapi_test.go @@ -0,0 +1,89 @@ +package zapi + +import ( + "bytes" + "testing" +) + +func TestEncodeDecodeHeader(t *testing.T) { + raw := EncodeHeader(42, 0, CmdHello) + hdr, err := DecodeHeader(raw) + if err != nil { + t.Fatal(err) + } + if hdr.Length != 42 { + t.Errorf("Length = %d, want 42", hdr.Length) + } + if hdr.Command != CmdHello { + t.Errorf("Command = %d, want %d", hdr.Command, CmdHello) + } + if hdr.VrfID != 0 { + t.Errorf("VrfID = %d, want 0", hdr.VrfID) + } +} + +func TestDecodeHeaderBadMarker(t *testing.T) { + raw := EncodeHeader(10, 0, CmdHello) + raw[2] = 0x00 + _, err := DecodeHeader(raw) + if err == nil { + t.Fatal("expected error for bad marker") + } +} + +func TestDecodeHeaderBadVersion(t *testing.T) { + raw := EncodeHeader(10, 0, CmdHello) + raw[3] = 5 + _, err := DecodeHeader(raw) + if err == nil { + t.Fatal("expected error for bad version") + } +} + +func TestBuildMessage(t *testing.T) { + body := EncodeHello() + msg := BuildMessage(CmdHello, DefaultVrf, body) + if len(msg) != HeaderSize+len(body) { + t.Errorf("message len = %d, want %d", len(msg), HeaderSize+len(body)) + } + hdr, err := DecodeHeader(msg[:HeaderSize]) + if err != nil { + t.Fatal(err) + } + if hdr.Command != CmdHello { + t.Errorf("Command = %d, want %d", hdr.Command, CmdHello) + } + if int(hdr.Length) != len(msg) { + t.Errorf("Length = %d, want %d", hdr.Length, len(msg)) + } +} + +func TestEncodeRedistributeAdd(t *testing.T) { + body := EncodeRedistributeAdd(AFIIPv4, RouteStatic) + if len(body) != 4 { + t.Fatalf("body len = %d, want 4", len(body)) + } + if body[0] != AFIIPv4 { + t.Errorf("afi = %d, want %d", body[0], AFIIPv4) + } + if body[1] != uint8(RouteStatic) { + t.Errorf("routeType = %d, want %d", body[1], RouteStatic) + } +} + +func TestReadMessage(t *testing.T) { + body := []byte{0x01, 0x02, 0x03} + msg := BuildMessage(CmdRouterIDUpdate, DefaultVrf, body) + r := bytes.NewReader(msg) + + hdr, gotBody, err := ReadMessage(r) + if err != nil { + t.Fatal(err) + } + if hdr.Command != CmdRouterIDUpdate { + t.Errorf("Command = %d, want %d", hdr.Command, CmdRouterIDUpdate) + } + if !bytes.Equal(gotBody, body) { + t.Errorf("body = %v, want %v", gotBody, body) + } +} diff --git a/src/yangerd/internal/zapiwatcher/zapiwatcher.go b/src/yangerd/internal/zapiwatcher/zapiwatcher.go new file mode 100644 index 000000000..482f60193 --- /dev/null +++ b/src/yangerd/internal/zapiwatcher/zapiwatcher.go @@ -0,0 +1,467 @@ +package zapiwatcher + +import ( + "context" + "encoding/json" + "fmt" + "log/slog" + "net" + "regexp" + "strconv" + "strings" + "time" + + "github.com/kernelkit/infix/src/yangerd/internal/backoff" + "github.com/kernelkit/infix/src/yangerd/internal/numconv" + "github.com/kernelkit/infix/src/yangerd/internal/tree" + "github.com/kernelkit/infix/src/yangerd/internal/zapi" +) + +const ( + zapiSocketPath = "/var/run/frr/zserv.api" + routingTreeKey = "ietf-routing:routing" + + // debounceDelay coalesces a burst of ZAPI route notifications into a + // single RIB read. FRR emits many RouteAdd/Del messages while a + // protocol converges; without debouncing we would re-read the table + // dozens of times in a few milliseconds. + debounceDelay = 200 * time.Millisecond + + // queryTimeout bounds one vty table read, so a wedged zebra cannot + // stall the refresh worker for good. + queryTimeout = 5 * time.Second +) + +// subscribeTypes are the route types we ask zebra to redistribute. We do +// not use the route payloads themselves -- redistribution only ever +// delivers the selected route per destination and does not reliably send +// a delete when a route is superseded. The subscription exists purely so +// zebra notifies us that *something* changed; the authoritative table is +// then read from zebra's vty socket (see RouteQuerier). +var subscribeTypes = []zapi.RouteType{ + zapi.RouteKernel, + zapi.RouteConnect, + zapi.RouteLocal, + zapi.RouteStatic, + zapi.RouteRIP, + zapi.RouteRIPNG, + zapi.RouteOSPF, + zapi.RouteOSPF6, +} + +// RouteQuerier runs a "show ... json" command against FRR and returns its +// raw output. Production code uses an frrvty.Client (zebra's vty socket); +// tests inject a fake. +type RouteQuerier interface { + Query(ctx context.Context, command string) ([]byte, error) +} + +// ZAPIWatcher keeps the operational RIB (ietf-routing:routing/ribs) in +// sync with FRR. It does NOT reconstruct routes from the ZAPI stream: +// the ZAPI socket is used only as a change trigger, and the full routing +// table -- every candidate per destination, with FRR's own +// selected/installed flags, exactly as "show ip route" renders it -- is +// read from zebra's vty socket on each change. Because every refresh is +// a complete snapshot, a route removed from zebra simply disappears; we +// never depend on receiving a ZAPI delete. +type ZAPIWatcher struct { + tree *tree.Tree + querier RouteQuerier + log *slog.Logger + refresh chan struct{} + socket string // zserv API socket; overridable in tests + + // onChange runs after each RIB snapshot. A route change means the + // protocols moved too, and FRR gives no other event for that. + onChange func() +} + +// SetOnChange registers fn to run after every RIB refresh. +func (w *ZAPIWatcher) SetOnChange(fn func()) { + w.onChange = fn +} + +func New(t *tree.Tree, querier RouteQuerier, log *slog.Logger) *ZAPIWatcher { + if log == nil { + log = slog.Default() + } + return &ZAPIWatcher{ + tree: t, + querier: querier, + log: log, + refresh: make(chan struct{}, 1), + socket: zapiSocketPath, + } +} + +func (w *ZAPIWatcher) Run(ctx context.Context) error { + // The refresh worker owns all writes to the tree and runs for the + // lifetime of the watcher, independent of the ZAPI connection. + go w.refreshLoop(ctx) + + return backoff.Retry(ctx, w.log, "zapi watcher", w.session) +} + +// session runs one ZAPI connection until zebra closes it or ctx ends. +func (w *ZAPIWatcher) session(ctx context.Context) error { + conn, err := w.connect(ctx) + if err != nil { + return err + } + defer conn.Close() + + // ReadMessage blocks with no deadline; closing the socket is what + // unblocks it on shutdown. + stop := context.AfterFunc(ctx, func() { conn.Close() }) + defer stop() + + w.log.Info("zapi watcher: connected", "socket", w.socket) + + // Read the current table now that we are subscribed, so we have + // data even if no further events arrive. + w.triggerRefresh() + + return w.processMessages(ctx, conn) +} + +func (w *ZAPIWatcher) connect(ctx context.Context) (net.Conn, error) { + d := net.Dialer{} + conn, err := d.DialContext(ctx, "unix", w.socket) + if err != nil { + return nil, fmt.Errorf("dial zserv: %w", err) + } + + hello := zapi.BuildMessage(zapi.CmdHello, zapi.DefaultVrf, zapi.EncodeHello()) + if _, err := conn.Write(hello); err != nil { + conn.Close() + return nil, fmt.Errorf("send hello: %w", err) + } + + for _, afi := range []uint8{zapi.AFIIPv4, zapi.AFIIPv6} { + msg := zapi.BuildMessage(zapi.CmdRouterIDAdd, zapi.DefaultVrf, zapi.EncodeRouterIDAdd(afi)) + if _, err := conn.Write(msg); err != nil { + conn.Close() + return nil, fmt.Errorf("send router-id-add: %w", err) + } + } + + for _, rt := range subscribeTypes { + for _, afi := range []uint8{zapi.AFIIPv4, zapi.AFIIPv6} { + msg := zapi.BuildMessage(zapi.CmdRedistributeAdd, zapi.DefaultVrf, zapi.EncodeRedistributeAdd(afi, rt)) + if _, err := conn.Write(msg); err != nil { + conn.Close() + return nil, fmt.Errorf("send redistribute-add: %w", err) + } + w.log.Debug("zapi watcher: subscribed", "afi", afi, "routeType", rt) + } + } + + return conn, nil +} + +// processMessages drains the ZAPI stream. We only care *that* a route +// changed, not what changed -- each route add/delete triggers a debounced +// re-read of the full table. +func (w *ZAPIWatcher) processMessages(ctx context.Context, conn net.Conn) error { + for { + select { + case <-ctx.Done(): + return ctx.Err() + default: + } + + hdr, _, err := zapi.ReadMessage(conn) + if err != nil { + return fmt.Errorf("read message: %w", err) + } + + switch hdr.Command { + case zapi.CmdRedistRouteAdd, zapi.CmdRedistRouteDel: + w.log.Debug("zapi watcher: route change", "cmd", hdr.Command, "vrf", hdr.VrfID) + w.triggerRefresh() + } + } +} + +// triggerRefresh requests a table re-read. The buffered channel collapses +// multiple pending requests into one. +func (w *ZAPIWatcher) triggerRefresh() { + select { + case w.refresh <- struct{}{}: + default: + } +} + +func (w *ZAPIWatcher) refreshLoop(ctx context.Context) { + for { + select { + case <-ctx.Done(): + return + case <-w.refresh: + } + + // Let a burst of notifications settle before reading. + select { + case <-ctx.Done(): + return + case <-time.After(debounceDelay): + } + // Drain a request that arrived during the debounce window; the + // upcoming read already reflects it. + select { + case <-w.refresh: + default: + } + + w.writeRibs(ctx) + } +} + +// writeRibs reads the full IPv4 and IPv6 routing tables from zebra and +// replaces the ribs subtree. On a query error it leaves the previous +// data untouched rather than blanking the table. +func (w *ZAPIWatcher) writeRibs(ctx context.Context) { + ipv4, err := w.collectRoutes(ctx, "ipv4") + if err != nil { + w.log.Warn("zapi watcher: read ipv4 routes", "err", err) + return + } + ipv6, err := w.collectRoutes(ctx, "ipv6") + if err != nil { + w.log.Warn("zapi watcher: read ipv6 routes", "err", err) + return + } + + ribs := map[string]any{ + "rib": []map[string]any{ + { + "name": "ipv4", + "address-family": "ietf-routing:ipv4", + "routes": map[string]any{"route": ipv4}, + }, + { + "name": "ipv6", + "address-family": "ietf-routing:ipv6", + "routes": map[string]any{"route": ipv6}, + }, + }, + } + + data, err := json.Marshal(map[string]any{"ribs": ribs}) + if err != nil { + w.log.Error("zapi watcher: marshal ribs", "err", err) + return + } + + w.tree.Merge(routingTreeKey, data) + if w.onChange != nil { + w.onChange() + } +} + +// collectRoutes runs "show ip route json" / "show ipv6 route json" and +// transforms every entry into an ietf-routing route node. +func (w *ZAPIWatcher) collectRoutes(ctx context.Context, family string) ([]json.RawMessage, error) { + command := "show ip route json" + if family == "ipv6" { + command = "show ipv6 route json" + } + + ctx, cancel := context.WithTimeout(ctx, queryTimeout) + defer cancel() + + out, err := w.querier.Query(ctx, command) + if err != nil { + return nil, err + } + + // FRR prints "{}" for an empty table; otherwise a map of + // prefix -> [route, ...] (multiple candidates per prefix). + var table map[string][]map[string]any + if err := json.Unmarshal(out, &table); err != nil { + return nil, fmt.Errorf("parse %q: %w", command, err) + } + + now := time.Now() + routes := make([]json.RawMessage, 0, len(table)) + for prefix, entries := range table { + if !strings.Contains(prefix, "/") { + continue + } + for _, entry := range entries { + routes = append(routes, transformRoute(family, prefix, entry, now)) + } + } + return routes, nil +} + +// protocolMap maps FRR's protocol names to IETF routing-protocol +// identities. Unknown protocols fall back to kernel so they still +// validate against the model. +var protocolMap = map[string]string{ + "kernel": "infix-routing:kernel", + "connected": "ietf-routing:direct", + "local": "ietf-routing:direct", + "static": "ietf-routing:static", + "ospf": "ietf-ospf:ospfv2", + "ospf6": "ietf-ospf:ospfv3", + "rip": "ietf-rip:rip", + "ripng": "ietf-rip:rip", +} + +func protocolName(frr string) string { + if p, ok := protocolMap[frr]; ok { + return p + } + return "infix-routing:kernel" +} + +// transformRoute converts one FRR JSON route entry into an ietf-routing +// route node. It mirrors the legacy yanger ietf_routing.py:add_protocol. +func transformRoute(family, prefixKey string, route map[string]any, now time.Time) json.RawMessage { + addrKey := "ietf-ipv4-unicast-routing:address" + dpKey := "ietf-ipv4-unicast-routing:destination-prefix" + nhAddrKey := "ietf-ipv4-unicast-routing:next-hop-address" + hostLen := "32" + if family == "ipv6" { + addrKey = "ietf-ipv6-unicast-routing:address" + dpKey = "ietf-ipv6-unicast-routing:destination-prefix" + nhAddrKey = "ietf-ipv6-unicast-routing:next-hop-address" + hostLen = "128" + } + + dst := stringField(route, "prefix") + if dst == "" { + dst = prefixKey + } + if !strings.Contains(dst, "/") { + plen := hostLen + if v, ok := route["prefixLen"]; ok { + plen = strconv.Itoa(numconv.IntOrZero(v)) + } + dst = dst + "/" + plen + } + + frr := stringField(route, "protocol") + + node := map[string]any{ + dpKey: dst, + "source-protocol": protocolName(frr), + "route-preference": numconv.IntOrZero(route["distance"]), + "last-updated": now.Add(-parseUptime(stringField(route, "uptime"))).Format(time.RFC3339), + } + + // Metric is modelled only for OSPF and RIP routes. + switch { + case strings.Contains(frr, "ospf"): + node["ietf-ospf:metric"] = numconv.IntOrZero(route["metric"]) + case strings.Contains(frr, "rip"): + node["ietf-rip:metric"] = numconv.IntOrZero(route["metric"]) + } + + // "selected" is FRR's own best-path decision -- the '>' in + // "show ip route". active is a presence leaf, encoded as [null]. + if boolField(route, "selected") { + node["active"] = []any{nil} + } + + installed := boolField(route, "installed") + + nextHops := make([]map[string]any, 0) + if hops, ok := route["nexthops"].([]any); ok { + for _, h := range hops { + hop, ok := h.(map[string]any) + if !ok { + continue + } + nh := map[string]any{} + if ip := stringField(hop, "ip"); ip != "" { + nh[addrKey] = ip + } else if ifn := stringField(hop, "interfaceName"); ifn != "" { + nh["outgoing-interface"] = ifn + } + // zebra marks the nexthop programmed into the FIB with + // "fib":true (see zebra/zebra_vty.c). + if installed && boolField(hop, "fib") { + nh["infix-routing:installed"] = []any{nil} + } + if len(nh) > 0 { + nextHops = append(nextHops, nh) + } + } + } + + if len(nextHops) > 0 { + node["next-hop"] = map[string]any{ + "next-hop-list": map[string]any{ + "next-hop": nextHops, + }, + } + } else { + nh := map[string]any{} + switch frr { + case "blackhole": + nh["special-next-hop"] = "blackhole" + case "unreachable": + nh["special-next-hop"] = "unreachable" + default: + if ifn := stringField(route, "interfaceName"); ifn != "" { + nh["outgoing-interface"] = ifn + } + if gw := stringField(route, "nexthop"); gw != "" { + nh[nhAddrKey] = gw + } + } + node["next-hop"] = nh + } + + encoded, err := json.Marshal(node) + if err != nil { + return json.RawMessage(`{}`) + } + return encoded +} + +func stringField(m map[string]any, key string) string { + if s, ok := m[key].(string); ok { + return s + } + return "" +} + +func boolField(m map[string]any, key string) bool { + b, _ := m[key].(bool) + return b +} + +// FRR uptime string formats (frrtime), ported from yanger's +// uptime2datetime: "HH:MM:SS", "XdXXhXXm", "XXwXdXXh". +var ( + uptimeHMS = regexp.MustCompile(`^(\d{2}):(\d{2}):(\d{2})$`) + uptimeDHM = regexp.MustCompile(`^(\d+)d(\d{2})h(\d{2})m$`) + uptimeWDH = regexp.MustCompile(`^(\d{2})w(\d)d(\d{2})h$`) +) + +// parseUptime converts an FRR uptime string into a duration. The +// last-updated leaf is then computed as now-uptime. Unrecognised input +// yields zero (i.e. last-updated == now). +func parseUptime(s string) time.Duration { + atoi := func(x string) int { n, _ := strconv.Atoi(x); return n } + + if m := uptimeHMS.FindStringSubmatch(s); m != nil { + return time.Duration(atoi(m[1]))*time.Hour + + time.Duration(atoi(m[2]))*time.Minute + + time.Duration(atoi(m[3]))*time.Second + } + if m := uptimeDHM.FindStringSubmatch(s); m != nil { + return time.Duration(atoi(m[1]))*24*time.Hour + + time.Duration(atoi(m[2]))*time.Hour + + time.Duration(atoi(m[3]))*time.Minute + } + if m := uptimeWDH.FindStringSubmatch(s); m != nil { + return time.Duration(atoi(m[1]))*7*24*time.Hour + + time.Duration(atoi(m[2]))*24*time.Hour + + time.Duration(atoi(m[3]))*time.Hour + } + return 0 +} diff --git a/src/yangerd/internal/zapiwatcher/zapiwatcher_test.go b/src/yangerd/internal/zapiwatcher/zapiwatcher_test.go new file mode 100644 index 000000000..0527ba13f --- /dev/null +++ b/src/yangerd/internal/zapiwatcher/zapiwatcher_test.go @@ -0,0 +1,352 @@ +package zapiwatcher + +import ( + "context" + "encoding/json" + "io" + "log/slog" + "net" + "path/filepath" + "testing" + "time" + + "github.com/kernelkit/infix/src/yangerd/internal/numconv" + "github.com/kernelkit/infix/src/yangerd/internal/tree" +) + +// fakeQuerier returns canned vtysh JSON per command. +type fakeQuerier struct { + ipv4 string + ipv6 string + err error +} + +func (f fakeQuerier) Query(_ context.Context, command string) ([]byte, error) { + if f.err != nil { + return nil, f.err + } + if command == "show ipv6 route json" { + if f.ipv6 == "" { + return []byte("{}"), nil + } + return []byte(f.ipv6), nil + } + if f.ipv4 == "" { + return []byte("{}"), nil + } + return []byte(f.ipv4), nil +} + +func ipv4Routes(t *testing.T, tr *tree.Tree) []map[string]any { + t.Helper() + data := tr.Get(routingTreeKey) + if data == nil { + t.Fatal("routing tree key not set") + } + var routing map[string]any + if err := json.Unmarshal(data, &routing); err != nil { + t.Fatalf("unmarshal routing: %v", err) + } + ribs := routing["ribs"].(map[string]any) + for _, rib := range ribs["rib"].([]any) { + rm := rib.(map[string]any) + if rm["name"] == "ipv4" { + out := []map[string]any{} + for _, r := range rm["routes"].(map[string]any)["route"].([]any) { + out = append(out, r.(map[string]any)) + } + return out + } + } + t.Fatal("ipv4 rib not found") + return nil +} + +func TestProtocolName(t *testing.T) { + cases := map[string]string{ + "kernel": "infix-routing:kernel", + "connected": "ietf-routing:direct", + "local": "ietf-routing:direct", + "static": "ietf-routing:static", + "ospf": "ietf-ospf:ospfv2", + "rip": "ietf-rip:rip", + "bgp": "infix-routing:kernel", // unknown -> kernel + "": "infix-routing:kernel", + } + for in, want := range cases { + if got := protocolName(in); got != want { + t.Errorf("protocolName(%q) = %q, want %q", in, got, want) + } + } +} + +func TestParseUptime(t *testing.T) { + cases := map[string]time.Duration{ + "02:09:02": 2*time.Hour + 9*time.Minute + 2*time.Second, + "00:00:30": 30 * time.Second, + "3d04h05m": 3*24*time.Hour + 4*time.Hour + 5*time.Minute, + "02w3d04h": 2*7*24*time.Hour + 3*24*time.Hour + 4*time.Hour, + "bogus": 0, + "": 0, + } + for in, want := range cases { + if got := parseUptime(in); got != want { + t.Errorf("parseUptime(%q) = %v, want %v", in, got, want) + } + } +} + +// The user's exact bug: a static route is in the FIB (selected, distance +// 120) while a stale OSPF entry (distance 110) is still listed but not +// selected. active must follow FRR's "selected" flag, not the lowest +// admin distance. +func TestActiveFollowsSelectedNotDistance(t *testing.T) { + const j = `{ + "192.168.20.0/24":[ + {"prefix":"192.168.20.0/24","protocol":"ospf","selected":false,"distance":110,"metric":100,"uptime":"02:11:49", + "nexthops":[{"ip":"192.168.60.2","interfaceName":"e3","active":true}]}, + {"prefix":"192.168.20.0/24","protocol":"static","selected":true,"installed":true,"distance":120,"metric":0,"uptime":"02:09:02", + "nexthops":[{"ip":"192.168.50.2","interfaceName":"e7","fib":true,"active":true}]} + ] + }` + + tr := tree.New() + w := New(tr, fakeQuerier{ipv4: j}, nil) + w.writeRibs(context.Background()) + + routes := ipv4Routes(t, tr) + if len(routes) != 2 { + t.Fatalf("expected 2 candidate routes, got %d", len(routes)) + } + + for _, r := range routes { + pref := numconv.IntOrZero(r["route-preference"]) + _, active := r["active"] + switch pref { + case 120: + if !active { + t.Error("static route (pref 120, selected) must be active") + } + case 110: + if active { + t.Error("ospf route (pref 110, not selected) must NOT be active") + } + default: + t.Errorf("unexpected route-preference %d", pref) + } + } +} + +// A full snapshot read means a route removed from zebra disappears from +// the cache without any ZAPI delete. +func TestSnapshotPurgesRemovedRoutes(t *testing.T) { + const before = `{ + "192.168.20.0/24":[ + {"prefix":"192.168.20.0/24","protocol":"ospf","selected":true,"distance":110,"uptime":"00:05:00","nexthops":[{"ip":"192.168.60.2"}]} + ] + }` + const after = `{ + "192.168.20.0/24":[ + {"prefix":"192.168.20.0/24","protocol":"static","selected":true,"installed":true,"distance":120,"uptime":"00:01:00","nexthops":[{"ip":"192.168.50.2","fib":true}]} + ] + }` + + tr := tree.New() + New(tr, fakeQuerier{ipv4: before}, nil).writeRibs(context.Background()) + if got := len(ipv4Routes(t, tr)); got != 1 { + t.Fatalf("before: expected 1 route, got %d", got) + } + + // zebra now has only the static route; the OSPF route is gone with no + // delete event. A fresh snapshot must not carry the corpse forward. + New(tr, fakeQuerier{ipv4: after}, nil).writeRibs(context.Background()) + routes := ipv4Routes(t, tr) + if len(routes) != 1 { + t.Fatalf("after: expected 1 route, got %d", len(routes)) + } + if got := protocolName("static"); routes[0]["source-protocol"] != got { + t.Errorf("surviving route protocol = %v, want %v", routes[0]["source-protocol"], got) + } +} + +func TestTransformRouteFields(t *testing.T) { + entry := map[string]any{ + "prefix": "10.0.0.0/24", + "protocol": "ospf", + "selected": true, + "installed": true, + "distance": float64(110), + "metric": float64(20), + "uptime": "01:00:00", + "nexthops": []any{ + map[string]any{"ip": "192.168.1.1", "interfaceName": "e1", "fib": true}, + }, + } + + now := time.Date(2026, 6, 10, 12, 0, 0, 0, time.UTC) + var parsed map[string]any + if err := json.Unmarshal(transformRoute("ipv4", "10.0.0.0/24", entry, now), &parsed); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + if parsed["ietf-ipv4-unicast-routing:destination-prefix"] != "10.0.0.0/24" { + t.Errorf("destination-prefix = %v", parsed["ietf-ipv4-unicast-routing:destination-prefix"]) + } + if parsed["source-protocol"] != "ietf-ospf:ospfv2" { + t.Errorf("source-protocol = %v", parsed["source-protocol"]) + } + if numconv.IntOrZero(parsed["route-preference"]) != 110 { + t.Errorf("route-preference = %v", parsed["route-preference"]) + } + if numconv.IntOrZero(parsed["ietf-ospf:metric"]) != 20 { + t.Errorf("ietf-ospf:metric = %v", parsed["ietf-ospf:metric"]) + } + if _, ok := parsed["active"]; !ok { + t.Error("selected route must have active leaf") + } + // last-updated = now - 1h + if parsed["last-updated"] != "2026-06-10T11:00:00Z" { + t.Errorf("last-updated = %v, want 2026-06-10T11:00:00Z", parsed["last-updated"]) + } + + hops := parsed["next-hop"].(map[string]any)["next-hop-list"].(map[string]any)["next-hop"].([]any) + if len(hops) != 1 { + t.Fatalf("expected 1 nexthop, got %d", len(hops)) + } + hop := hops[0].(map[string]any) + if hop["ietf-ipv4-unicast-routing:address"] != "192.168.1.1" { + t.Errorf("nexthop address = %v", hop["ietf-ipv4-unicast-routing:address"]) + } + if _, ok := hop["infix-routing:installed"]; !ok { + t.Error("fib nexthop must have infix-routing:installed") + } +} + +func TestTransformRouteBlackhole(t *testing.T) { + entry := map[string]any{ + "prefix": "10.1.0.0/24", + "protocol": "blackhole", + "distance": float64(0), + "uptime": "00:00:10", + } + var parsed map[string]any + if err := json.Unmarshal(transformRoute("ipv4", "10.1.0.0/24", entry, time.Now()), &parsed); err != nil { + t.Fatalf("unmarshal: %v", err) + } + nh := parsed["next-hop"].(map[string]any) + if nh["special-next-hop"] != "blackhole" { + t.Errorf("special-next-hop = %v, want blackhole", nh["special-next-hop"]) + } +} + +func TestWriteRibsQueryErrorKeepsData(t *testing.T) { + tr := tree.New() + // Seed with good data. + New(tr, fakeQuerier{ipv4: `{"10.0.0.0/24":[{"prefix":"10.0.0.0/24","protocol":"static","selected":true,"distance":1,"uptime":"00:00:05","nexthops":[{"ip":"10.0.0.1"}]}]}`}, nil). + writeRibs(context.Background()) + before := tr.Get(routingTreeKey) + + // A failing query must not blank the table. + New(tr, fakeQuerier{err: context.DeadlineExceeded}, nil).writeRibs(context.Background()) + after := tr.Get(routingTreeKey) + + if string(before) != string(after) { + t.Error("query error overwrote existing rib data") + } +} + +func TestWriteRibsBothFamilies(t *testing.T) { + tr := tree.New() + w := New(tr, fakeQuerier{ + ipv4: `{"10.0.0.0/24":[{"prefix":"10.0.0.0/24","protocol":"static","selected":true,"distance":1,"uptime":"00:00:05","nexthops":[{"ip":"10.0.0.1"}]}]}`, + ipv6: `{"2001:db8::/64":[{"prefix":"2001:db8::/64","protocol":"connected","selected":true,"distance":0,"uptime":"00:00:05","nexthops":[{"interfaceName":"e1"}]}]}`, + }, nil) + w.writeRibs(context.Background()) + + var routing map[string]any + if err := json.Unmarshal(tr.Get(routingTreeKey), &routing); err != nil { + t.Fatalf("unmarshal: %v", err) + } + ribs := routing["ribs"].(map[string]any)["rib"].([]any) + if len(ribs) != 2 { + t.Fatalf("expected 2 ribs, got %d", len(ribs)) + } +} + +// A connected but idle zserv must not keep Run from returning on +// shutdown: the read blocks until the socket is closed. +func TestRunReturnsOnCancelWhileIdle(t *testing.T) { + sock := filepath.Join(t.TempDir(), "zserv.api") + ln, err := net.Listen("unix", sock) + if err != nil { + t.Fatal(err) + } + defer ln.Close() + + accepted := make(chan struct{}) + go func() { + conn, err := ln.Accept() + if err != nil { + return + } + close(accepted) + io.Copy(io.Discard, conn) + conn.Close() + }() + + w := New(tree.New(), fakeQuerier{}, slog.New(slog.NewTextHandler(io.Discard, nil))) + w.socket = sock + + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan error, 1) + go func() { done <- w.Run(ctx) }() + + select { + case <-accepted: + case <-time.After(5 * time.Second): + t.Fatal("watcher never connected") + } + cancel() + + select { + case <-done: + case <-time.After(2 * time.Second): + t.Fatal("Run did not return after cancel while the socket was idle") + } +} + +// A vty read must carry a deadline, so a wedged zebra cannot stall the +// refresh worker for good. +func TestCollectRoutesHasDeadline(t *testing.T) { + var deadline bool + w := New(tree.New(), deadlineQuerier{&deadline}, slog.New(slog.NewTextHandler(io.Discard, nil))) + if _, err := w.collectRoutes(context.Background(), "ipv4"); err != nil { + t.Fatal(err) + } + if !deadline { + t.Fatal("vty query ran without a deadline") + } +} + +type deadlineQuerier struct{ seen *bool } + +func (q deadlineQuerier) Query(ctx context.Context, _ string) ([]byte, error) { + _, *q.seen = ctx.Deadline() + return []byte("{}"), nil +} + +// A RIB refresh runs the on-change hook, which main uses to re-read +// protocol state; a failed refresh does not. +func TestWriteRibsRunsOnChange(t *testing.T) { + calls := 0 + w := New(tree.New(), fakeQuerier{ipv4: `{}`}, nil) + w.SetOnChange(func() { calls++ }) + w.writeRibs(context.Background()) + if calls != 1 { + t.Fatalf("onChange ran %d times after one refresh", calls) + } + New(tree.New(), fakeQuerier{err: context.DeadlineExceeded}, nil).writeRibs(context.Background()) + if calls != 1 { + t.Fatalf("onChange ran after a failed refresh") + } +} diff --git a/src/yangerd/vendor/github.com/facebook/time/LICENSE b/src/yangerd/vendor/github.com/facebook/time/LICENSE new file mode 100644 index 000000000..d64569567 --- /dev/null +++ b/src/yangerd/vendor/github.com/facebook/time/LICENSE @@ -0,0 +1,202 @@ + + Apache License + Version 2.0, January 2004 + http://www.apache.org/licenses/ + + TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION + + 1. Definitions. + + "License" shall mean the terms and conditions for use, reproduction, + and distribution as defined by Sections 1 through 9 of this document. + + "Licensor" shall mean the copyright owner or entity authorized by + the copyright owner that is granting the License. + + "Legal Entity" shall mean the union of the acting entity and all + other entities that control, are controlled by, or are under common + control with that entity. For the purposes of this definition, + "control" means (i) the power, direct or indirect, to cause the + direction or management of such entity, whether by contract or + otherwise, or (ii) ownership of fifty percent (50%) or more of the + outstanding shares, or (iii) beneficial ownership of such entity. + + "You" (or "Your") shall mean an individual or Legal Entity + exercising permissions granted by this License. + + "Source" form shall mean the preferred form for making modifications, + including but not limited to software source code, documentation + source, and configuration files. + + "Object" form shall mean any form resulting from mechanical + transformation or translation of a Source form, including but + not limited to compiled object code, generated documentation, + and conversions to other media types. + + "Work" shall mean the work of authorship, whether in Source or + Object form, made available under the License, as indicated by a + copyright notice that is included in or attached to the work + (an example is provided in the Appendix below). + + "Derivative Works" shall mean any work, whether in Source or Object + form, that is based on (or derived from) the Work and for which the + editorial revisions, annotations, elaborations, or other modifications + represent, as a whole, an original work of authorship. For the purposes + of this License, Derivative Works shall not include works that remain + separable from, or merely link (or bind by name) to the interfaces of, + the Work and Derivative Works thereof. + + "Contribution" shall mean any work of authorship, including + the original version of the Work and any modifications or additions + to that Work or Derivative Works thereof, that is intentionally + submitted to Licensor for inclusion in the Work by the copyright owner + or by an individual or Legal Entity authorized to submit on behalf of + the copyright owner. For the purposes of this definition, "submitted" + means any form of electronic, verbal, or written communication sent + to the Licensor or its representatives, including but not limited to + communication on electronic mailing lists, source code control systems, + and issue tracking systems that are managed by, or on behalf of, the + Licensor for the purpose of discussing and improving the Work, but + excluding communication that is conspicuously marked or otherwise + designated in writing by the copyright owner as "Not a Contribution." + + "Contributor" shall mean Licensor and any individual or Legal Entity + on behalf of whom a Contribution has been received by Licensor and + subsequently incorporated within the Work. + + 2. Grant of Copyright License. Subject to the terms and conditions of + this License, each Contributor hereby grants to You a perpetual, + worldwide, non-exclusive, no-charge, royalty-free, irrevocable + copyright license to reproduce, prepare Derivative Works of, + publicly display, publicly perform, sublicense, and distribute the + Work and such Derivative Works in Source or Object form. + + 3. Grant of Patent License. Subject to the terms and conditions of + this License, each Contributor hereby grants to You a perpetual, + worldwide, non-exclusive, no-charge, royalty-free, irrevocable + (except as stated in this section) patent license to make, have made, + use, offer to sell, sell, import, and otherwise transfer the Work, + where such license applies only to those patent claims licensable + by such Contributor that are necessarily infringed by their + Contribution(s) alone or by combination of their Contribution(s) + with the Work to which such Contribution(s) was submitted. If You + institute patent litigation against any entity (including a + cross-claim or counterclaim in a lawsuit) alleging that the Work + or a Contribution incorporated within the Work constitutes direct + or contributory patent infringement, then any patent licenses + granted to You under this License for that Work shall terminate + as of the date such litigation is filed. + + 4. Redistribution. You may reproduce and distribute copies of the + Work or Derivative Works thereof in any medium, with or without + modifications, and in Source or Object form, provided that You + meet the following conditions: + + (a) You must give any other recipients of the Work or + Derivative Works a copy of this License; and + + (b) You must cause any modified files to carry prominent notices + stating that You changed the files; and + + (c) You must retain, in the Source form of any Derivative Works + that You distribute, all copyright, patent, trademark, and + attribution notices from the Source form of the Work, + excluding those notices that do not pertain to any part of + the Derivative Works; and + + (d) If the Work includes a "NOTICE" text file as part of its + distribution, then any Derivative Works that You distribute must + include a readable copy of the attribution notices contained + within such NOTICE file, excluding those notices that do not + pertain to any part of the Derivative Works, in at least one + of the following places: within a NOTICE text file distributed + as part of the Derivative Works; within the Source form or + documentation, if provided along with the Derivative Works; or, + within a display generated by the Derivative Works, if and + wherever such third-party notices normally appear. The contents + of the NOTICE file are for informational purposes only and + do not modify the License. You may add Your own attribution + notices within Derivative Works that You distribute, alongside + or as an addendum to the NOTICE text from the Work, provided + that such additional attribution notices cannot be construed + as modifying the License. + + You may add Your own copyright statement to Your modifications and + may provide additional or different license terms and conditions + for use, reproduction, or distribution of Your modifications, or + for any such Derivative Works as a whole, provided Your use, + reproduction, and distribution of the Work otherwise complies with + the conditions stated in this License. + + 5. Submission of Contributions. Unless You explicitly state otherwise, + any Contribution intentionally submitted for inclusion in the Work + by You to the Licensor shall be under the terms and conditions of + this License, without any additional terms or conditions. + Notwithstanding the above, nothing herein shall supersede or modify + the terms of any separate license agreement you may have executed + with Licensor regarding such Contributions. + + 6. Trademarks. This License does not grant permission to use the trade + names, trademarks, service marks, or product names of the Licensor, + except as required for reasonable and customary use in describing the + origin of the Work and reproducing the content of the NOTICE file. + + 7. Disclaimer of Warranty. Unless required by applicable law or + agreed to in writing, Licensor provides the Work (and each + Contributor provides its Contributions) on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or + implied, including, without limitation, any warranties or conditions + of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A + PARTICULAR PURPOSE. You are solely responsible for determining the + appropriateness of using or redistributing the Work and assume any + risks associated with Your exercise of permissions under this License. + + 8. Limitation of Liability. In no event and under no legal theory, + whether in tort (including negligence), contract, or otherwise, + unless required by applicable law (such as deliberate and grossly + negligent acts) or agreed to in writing, shall any Contributor be + liable to You for damages, including any direct, indirect, special, + incidental, or consequential damages of any character arising as a + result of this License or out of the use or inability to use the + Work (including but not limited to damages for loss of goodwill, + work stoppage, computer failure or malfunction, or any and all + other commercial damages or losses), even if such Contributor + has been advised of the possibility of such damages. + + 9. Accepting Warranty or Additional Liability. While redistributing + the Work or Derivative Works thereof, You may choose to offer, + and charge a fee for, acceptance of support, warranty, indemnity, + or other liability obligations and/or rights consistent with this + License. However, in accepting such obligations, You may act only + on Your own behalf and on Your sole responsibility, not on behalf + of any other Contributor, and only if You agree to indemnify, + defend, and hold each Contributor harmless for any liability + incurred by, or claims asserted against, such Contributor by reason + of your accepting any such warranty or additional liability. + + END OF TERMS AND CONDITIONS + + APPENDIX: How to apply the Apache License to your work. + + To apply the Apache License to your work, attach the following + boilerplate notice, with the fields enclosed by brackets "[]" + replaced with your own identifying information. (Don't include + the brackets!) The text should be enclosed in the appropriate + comment syntax for the file format. We also recommend that a + file or class name and description of purpose be included on the + same "printed page" as the copyright notice for easier + identification within third-party archives. + + Copyright [yyyy] [name of copyright owner] + + Licensed under the Apache License, Version 2.0 (the "License"); + you may not use this file except in compliance with the License. + You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + + Unless required by applicable law or agreed to in writing, software + distributed under the License is distributed on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + See the License for the specific language governing permissions and + limitations under the License. diff --git a/src/yangerd/vendor/github.com/facebook/time/hostendian/hostendian.go b/src/yangerd/vendor/github.com/facebook/time/hostendian/hostendian.go new file mode 100644 index 000000000..5abc740b3 --- /dev/null +++ b/src/yangerd/vendor/github.com/facebook/time/hostendian/hostendian.go @@ -0,0 +1,46 @@ +/* +Copyright (c) Facebook, Inc. and its affiliates. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +/* +Package hostendian provides way to check the endianness of the +machine this code is running on. + +While it's not needed most of the time, but software sometimes will combine +BigEndian and Host Endian (which typically is LittleEndian, but it's not guaranteed) +data in one structure sent over unix socket, and we need to work with it regardless. +*/ +package hostendian + +import ( + "encoding/binary" + "unsafe" +) + +// Order of the bytes +var Order binary.ByteOrder = binary.LittleEndian + +// IsBigEndian is a flag determining if value is in Big Endian +var IsBigEndian bool + +func init() { + var i uint16 = 0x0100 + ptr := unsafe.Pointer(&i) + if *(*byte)(ptr) == 0x01 { + // we are on the big endian machine + IsBigEndian = true + Order = binary.BigEndian + } +} diff --git a/src/yangerd/vendor/github.com/facebook/time/ntp/chrony/README.md b/src/yangerd/vendor/github.com/facebook/time/ntp/chrony/README.md new file mode 100644 index 000000000..41f1a4af7 --- /dev/null +++ b/src/yangerd/vendor/github.com/facebook/time/ntp/chrony/README.md @@ -0,0 +1,7 @@ +# Chrony network protocol used for command and monitoring of the timeserver + +[![GoDoc](https://godoc.org/github.com/facebook/time/ntp/protocol/chrony?status.svg)](https://godoc.org/github.com/facebook/time/ntp/protocol/chrony) + +Native Go implementation of Chrony communication protocol v6. + +As of now, only monitoring part of protocol that is used to communicate between `chronyc` and `chronyd` is implemented. diff --git a/src/yangerd/vendor/github.com/facebook/time/ntp/chrony/client.go b/src/yangerd/vendor/github.com/facebook/time/ntp/chrony/client.go new file mode 100644 index 000000000..93bece46c --- /dev/null +++ b/src/yangerd/vendor/github.com/facebook/time/ntp/chrony/client.go @@ -0,0 +1,45 @@ +/* +Copyright (c) Facebook, Inc. and its affiliates. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package chrony + +import ( + "encoding/binary" + "io" +) + +// Client talks to chronyd +type Client struct { + Connection io.ReadWriter + Sequence uint32 +} + +// Communicate sends the packet to chronyd, parse response into something usable +func (n *Client) Communicate(packet RequestPacket) (ResponsePacket, error) { + n.Sequence++ + var err error + packet.SetSequence(n.Sequence) + err = binary.Write(n.Connection, binary.BigEndian, packet) + if err != nil { + return nil, err + } + response := make([]uint8, 1024) + read, err := n.Connection.Read(response) + if err != nil { + return nil, err + } + return decodePacket(response[:read]) +} diff --git a/src/yangerd/vendor/github.com/facebook/time/ntp/chrony/doc.go b/src/yangerd/vendor/github.com/facebook/time/ntp/chrony/doc.go new file mode 100644 index 000000000..ed15f2b87 --- /dev/null +++ b/src/yangerd/vendor/github.com/facebook/time/ntp/chrony/doc.go @@ -0,0 +1,28 @@ +/* +Copyright (c) Facebook, Inc. and its affiliates. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +/* +Package chrony implements Chrony (https://chrony.tuxfamily.org) network protocol v6 used for monitoring of the timeserver. + +As of now, only monitoring part of protocol that is used to communicate between `chronyc` and `chronyd` is implemented. +Chronyc/chronyd protocol is not documented (https://chrony.tuxfamily.org/faq.html#_is_the_code_chronyc_code_code_chronyd_code_protocol_documented_anywhere). + +Library allows communicating with Chrony NTP server, +and get various information, for example: current server status; server variables like offset; peers with their statuses and variables; server counters. + +Example usage can be found in ntpcheck project - https://github.com/facebook/time/ntp/ntpcheck +*/ +package chrony diff --git a/src/yangerd/vendor/github.com/facebook/time/ntp/chrony/helpers.go b/src/yangerd/vendor/github.com/facebook/time/ntp/chrony/helpers.go new file mode 100644 index 000000000..8954a52db --- /dev/null +++ b/src/yangerd/vendor/github.com/facebook/time/ntp/chrony/helpers.go @@ -0,0 +1,194 @@ +/* +Copyright (c) Facebook, Inc. and its affiliates. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package chrony + +import ( + "fmt" + "math" + "net" + "time" + "unicode" +) + +// ChronySocketPath is the default path to chronyd socket +const ChronySocketPath = "/var/run/chrony/chronyd.sock" + +// ChronyPortV6Regexp is a regexp to find anything that listens on port 323 +// hex(323) = '0x143' +const ChronyPortV6Regexp = "[0-9]+: [0-9A-Z]+:0143 .*" + +// This is used in timeSpec.SecHigh for 32-bit timestamps +const noHighSec uint32 = 0x7fffffff + +// ip stuff +const ( + ipAddrInet4 uint16 = 1 + ipAddrInet6 uint16 = 2 +) + +// magic numbers to convert chronyFloat to normal float +const ( + floatExpBits = 7 + floatCoefBits = (4*8 - floatExpBits) +) + +type ipAddr struct { + IP [16]uint8 + Family uint16 + Pad uint16 +} + +func (ip *ipAddr) ToNetIP() net.IP { + if ip.Family == ipAddrInet4 { + return net.IP(ip.IP[:4]) + } + return net.IP(ip.IP[:]) +} + +func newIPAddr(ip net.IP) *ipAddr { + family := ipAddrInet6 + if ip.To4() != nil { + family = ipAddrInet4 + } + var nIP [16]byte + copy(nIP[:], ip) + return &ipAddr{ + IP: nIP, + Family: family, + } +} + +type timeSpec struct { + SecHigh uint32 + SecLow uint32 + Nsec uint32 +} + +func (t *timeSpec) ToTime() time.Time { + highU64 := uint64(t.SecHigh) + if t.SecHigh == noHighSec { + highU64 = 0 + } + lowU64 := uint64(t.SecLow) + return time.Unix(int64(highU64<<32|lowU64), int64(t.Nsec)) +} + +/* +32-bit floating-point format consisting of 7-bit signed exponent +and 25-bit signed coefficient without hidden bit. +The result is calculated as: 2^(exp - 25) * coef +*/ +type chronyFloat int32 + +// ToFloat does magic to decode float from int32. +// Code is copied and translated to Go from original C sources. +func (f chronyFloat) ToFloat() float64 { + var exp, coef int32 + + x := uint32(f) + + exp = int32(x >> floatCoefBits) + if exp >= 1<<(floatExpBits-1) { + exp -= 1 << floatExpBits + } + exp -= floatCoefBits + + coef = int32(x % (1 << floatCoefBits)) + if coef >= 1<<(floatCoefBits-1) { + coef -= 1 << floatCoefBits + } + + return float64(coef) * math.Pow(2.0, float64(exp)) +} + +// RefidAsHEX prints ref id as hex +func RefidAsHEX(refID uint32) string { + return fmt.Sprintf("%08X", refID) +} + +// RefidToString decodes ASCII string encoded as uint32 +func RefidToString(refID uint32) string { + result := []rune{} + + for i := 0; i < 4 && i < 64-1; i++ { + c := rune((refID >> (24 - uint(i)*8)) & 0xff) + if unicode.IsPrint(c) { + result = append(result, c) + } + } + + return string(result) +} + +/* NTP tests from RFC 5905: + +--------------------------+----------------------------------------+ + | Packet Type | Description | + +--------------------------+----------------------------------------+ + | 1 duplicate packet | The packet is at best an old duplicate | + | | or at worst a replay by a hacker. | + | | This can happen in symmetric modes if | + | | the poll intervals are uneven. | + | 2 bogus packet | | + | 3 invalid | One or more timestamp fields are | + | | invalid. This normally happens in | + | | symmetric modes when one peer sends | + | | the first packet to the other and | + | | before the other has received its | + | | first reply. | + | 4 access denied | The access controls have blacklisted | + | | the source. | + | 5 authentication failure | The cryptographic message digest does | + | | not match the MAC. | + | 6 unsynchronized | The server is not synchronized to a | + | | valid source. | + | 7 bad header data | One or more header fields are invalid. | + +--------------------------+----------------------------------------+ + +chrony doesn't do test #4, but adds four extra tests: +* maximum delay +* delay ratio +* delay dev ratio +* synchronisation loop. + +Those tests are roughly equivalent to ntpd 'flashers' +*/ + +// NTPTestDescMap maps bit mask with corresponding flash status +var NTPTestDescMap = map[uint16]string{ + 0x0001: "pkt_dup", + 0x0002: "pkt_bogus", + 0x0004: "pkt_invalid", + 0x0008: "pkt_auth", + 0x0010: "pkt_stratum", + 0x0020: "pkt_header", + 0x0040: "tst_max_delay", + 0x0080: "tst_delay_ratio", + 0x0100: "tst_delay_dev_ration", + 0x0200: "tst_sync_loop", +} + +// ReadNTPTestFlags returns list of failed ntp test flags (as strings) +func ReadNTPTestFlags(flags uint16) []string { + testFlags := flags & NTPFlagsTests + results := []string{} + for mask, message := range NTPTestDescMap { + if testFlags&mask == 0 { + results = append(results, message) + } + } + return results +} diff --git a/src/yangerd/vendor/github.com/facebook/time/ntp/chrony/logger.go b/src/yangerd/vendor/github.com/facebook/time/ntp/chrony/logger.go new file mode 100644 index 000000000..8cf3e2118 --- /dev/null +++ b/src/yangerd/vendor/github.com/facebook/time/ntp/chrony/logger.go @@ -0,0 +1,36 @@ +/* +Copyright (c) Facebook, Inc. and its affiliates. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package chrony + +// LoggerInterface is an interface for debug logging. +type LoggerInterface interface { + Printf(format string, v ...interface{}) +} + +type noopLogger struct{} + +func (noopLogger) Printf(_ string, _ ...interface{}) {} + +// Logger is a default debug logger which simply discards all messages. +// It can be overridden by setting the global variable to a different implementation, like std log +// +// chrony.Logger = log.New(os.Stderr, "", 0) +// +// or logrus +// +// chrony.Logger = logrus.StandardLogger() +var Logger LoggerInterface = &noopLogger{} diff --git a/src/yangerd/vendor/github.com/facebook/time/ntp/chrony/packet.go b/src/yangerd/vendor/github.com/facebook/time/ntp/chrony/packet.go new file mode 100644 index 000000000..01e4c0ca9 --- /dev/null +++ b/src/yangerd/vendor/github.com/facebook/time/ntp/chrony/packet.go @@ -0,0 +1,1172 @@ +/* +Copyright (c) Facebook, Inc. and its affiliates. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package chrony + +import ( + "bytes" + "encoding/binary" + "fmt" + "net" + "time" +) + +// original C++ versions of those consts/structs +// are in https://github.com/mlichvar/chrony/blob/master/candm.h + +// ReplyType identifies reply packet type +type ReplyType uint16 + +// CommandType identifies command type in both request and repy +type CommandType uint16 + +// ModeType identifies source (peer) mode +type ModeType uint16 + +// SourceStateType identifies source (peer) state +type SourceStateType uint16 + +// ResponseStatusType identifies response status +type ResponseStatusType uint16 + +// PacketType - request or reply +type PacketType uint8 + +// we implement latest (at the moment) protocol version +const protoVersionNumber uint8 = 6 +const maxDataLen = 396 + +// packet types +const ( + pktTypeCmdRequest PacketType = 1 + pktTypeCmdReply PacketType = 2 +) + +func (t PacketType) String() string { + switch t { + case pktTypeCmdRequest: + return "request" + case pktTypeCmdReply: + return "reply" + default: + return fmt.Sprintf("unknown (%d)", t) + } +} + +// request types. Only those we support, there are more +const ( + reqNSources CommandType = 14 + reqSourceData CommandType = 15 + reqTracking CommandType = 33 + reqSourceStats CommandType = 34 + reqActivity CommandType = 44 + reqServerStats CommandType = 54 + reqNTPData CommandType = 57 + reqNTPSourceName CommandType = 65 + reqSelectData CommandType = 69 +) + +// reply types +const ( + RpyNSources ReplyType = 2 + RpySourceData ReplyType = 3 + RpyTracking ReplyType = 5 + RpySourceStats ReplyType = 6 + RpyActivity ReplyType = 12 + RpyServerStats ReplyType = 14 + RpyNTPData ReplyType = 16 + RpyNTPSourceName ReplyType = 19 + RpyServerStats2 ReplyType = 22 + RpySelectData ReplyType = 23 + RpyServerStats3 ReplyType = 24 + RpyServerStats4 ReplyType = 25 + RpyNTPData2 ReplyType = 26 +) + +// source modes +const ( + SourceModeClient ModeType = 0 + SourceModePeer ModeType = 1 + SourceModeRef ModeType = 2 +) + +// source state +const ( + SourceStateSync SourceStateType = 0 + SourceStateUnreach SourceStateType = 1 + SourceStateFalseTicker SourceStateType = 2 + SourceStateJittery SourceStateType = 3 + SourceStateCandidate SourceStateType = 4 + SourceStateOutlier SourceStateType = 5 +) + +// source data flags +const ( + FlagNoselect uint16 = 0x1 + FlagPrefer uint16 = 0x2 + FlagTrust uint16 = 0x4 + FlagRequire uint16 = 0x8 +) + +// select data flags +const ( + FlagSDOptionNoSelect uint16 = 0x1 + FlagSDOptionPrefer uint16 = 0x2 + FlagSDOptionTrust uint16 = 0x4 + FlagSDOptionRequire uint16 = 0x8 +) + +// ntpdata flags +const ( + NTPFlagsTests uint16 = 0x3ff + NTPFlagInterleaved uint16 = 0x4000 + NTPFlagAuthenticated uint16 = 0x8000 +) + +// response status codes +const ( + sttSuccess ResponseStatusType = 0 + sttFailed ResponseStatusType = 1 + sttUnauth ResponseStatusType = 2 + sttInvalid ResponseStatusType = 3 + sttNoSuchSource ResponseStatusType = 4 + sttInvalidTS ResponseStatusType = 5 + sttNotEnabled ResponseStatusType = 6 + sttBadSubnet ResponseStatusType = 7 + sttAccessAllowed ResponseStatusType = 8 + sttAccessDenied ResponseStatusType = 9 + sttNoHostAccess ResponseStatusType = 10 + sttSourceAlreadyKnown ResponseStatusType = 11 + sttTooManySources ResponseStatusType = 12 + sttNoRTC ResponseStatusType = 13 + sttBadRTCFile ResponseStatusType = 14 + sttInactive ResponseStatusType = 15 + sttBadSample ResponseStatusType = 16 + sttInvalidAF ResponseStatusType = 17 + sttBadPktVersion ResponseStatusType = 18 + sttBadPktLength ResponseStatusType = 19 +) + +// StatusDesc provides mapping from ResponseStatusType to string +var StatusDesc = [20]string{ + "SUCCESS", + "FAILED", + "UNAUTH", + "INVALID", + "NOSUCHSOURCE", + "INVALIDTS", + "NOTENABLED", + "BADSUBNET", + "ACCESSALLOWED", + "ACCESSDENIED", + "NOHOSTACCESS", + "SOURCEALREADYKNOWN", + "TOOMANYSOURCES", + "NORTC", + "BADRTCFILE", + "INACTIVE", + "BADSAMPLE", + "INVALIDAF", + "BADPKTVERSION", + "BADPKTLENGTH", +} + +func (r ResponseStatusType) String() string { + if int(r) >= len(StatusDesc) { + return fmt.Sprintf("UNKNOWN (%d)", r) + } + return StatusDesc[r] +} + +// SourceStateDesc provides mapping from SourceStateType to string +var SourceStateDesc = [6]string{ + "sync", + "unreach", + "falseticker", + "jittery", + "candidate", + "outlier", +} + +func (s SourceStateType) String() string { + if int(s) >= len(SourceStateDesc) { + return fmt.Sprintf("unknown (%d)", s) + } + return SourceStateDesc[s] +} + +// ModeTypeDesc provides mapping from ModeType to string +var ModeTypeDesc = [3]string{ + "client", + "peer", + "reference clock", +} + +func (m ModeType) String() string { + if int(m) >= len(ModeTypeDesc) { + return fmt.Sprintf("unknown (%d)", m) + } + return ModeTypeDesc[m] +} + +// RequestHead is the first (common) part of the request, +// in a format that can be directly passed to binary.Write +type RequestHead struct { + Version uint8 + PKTType PacketType + Res1 uint8 + Res2 uint8 + Command CommandType + Attempt uint16 + Sequence uint32 + Pad1 uint32 + Pad2 uint32 +} + +// GetCommand returns request packet command +func (r *RequestHead) GetCommand() CommandType { + return r.Command +} + +// SetSequence sets request packet sequence number +func (r *RequestHead) SetSequence(n uint32) { + r.Sequence = n +} + +// RequestPacket is an interface to abstract all different outgoing packets +type RequestPacket interface { + GetCommand() CommandType + SetSequence(n uint32) +} + +// ResponsePacket is an interface to abstract all different incoming packets +type ResponsePacket interface { + GetCommand() CommandType + GetType() ReplyType + GetStatus() ResponseStatusType +} + +// RequestSources - packet to request number of sources (peers) +type RequestSources struct { + RequestHead + // we actually need this to send proper packet + data [maxDataLen]uint8 +} + +// RequestSourceData - packet to request source data for source id +type RequestSourceData struct { + RequestHead + Index int32 + EOR int32 + // we pass i32 - 4 bytes + data [maxDataLen - 4]uint8 +} + +// RequestNTPData - packet to request NTP data for peer IP. +// As of now, it's only allowed by Chrony over unix socket connection. +type RequestNTPData struct { + RequestHead + IPAddr ipAddr + EOR int32 + // we pass at max ipv6 addr - 16 bytes + data [maxDataLen - 16]uint8 +} + +// RequestNTPSourceName - packet to request source name for peer IP. +type RequestNTPSourceName struct { + RequestHead + IPAddr ipAddr + EOR int32 + // we pass at max ipv6 addr - 16 bytes + data [maxDataLen - 16]uint8 +} + +// RequestServerStats - packet to request server stats +type RequestServerStats struct { + RequestHead + // we actually need this to send proper packet + data [maxDataLen]uint8 +} + +// RequestTracking - packet to request 'tracking' data +type RequestTracking struct { + RequestHead + // we actually need this to send proper packet + data [maxDataLen]uint8 +} + +// RequestSourceStats - packet to request 'sourcestats' data for source id +type RequestSourceStats struct { + RequestHead + Index int32 + EOR int32 + // we pass i32 - 4 bytes + data [maxDataLen - 4]uint8 +} + +// RequestActivity - packet to request 'activity' data +type RequestActivity struct { + RequestHead + // we actually need this to send proper packet + data [maxDataLen]uint8 +} + +// RequestSelectData - packet to request 'selectdata' data +type RequestSelectData struct { + RequestHead + Index int32 + EOR int32 + // we pass i32 - 4 bytes + data [maxDataLen - 4]uint8 +} + +// ReplyHead is the first (common) part of the reply packet, +// in a format that can be directly passed to binary.Read +type ReplyHead struct { + Version uint8 + PKTType PacketType + Res1 uint8 + Res2 uint8 + Command CommandType + Reply ReplyType + Status ResponseStatusType + Pad1 uint16 + Pad2 uint16 + Pad3 uint16 + Sequence uint32 + Pad4 uint32 + Pad5 uint32 +} + +// GetCommand returns reply packet command +func (r *ReplyHead) GetCommand() CommandType { + return r.Command +} + +// GetType returns reply packet type +func (r *ReplyHead) GetType() ReplyType { + return r.Reply +} + +// GetStatus returns reply packet status +func (r *ReplyHead) GetStatus() ResponseStatusType { + return r.Status +} + +type replySourcesContent struct { + NSources uint32 +} + +// ReplySources is a usable version of a reply to 'sources' command +type ReplySources struct { + ReplyHead + NSources int +} + +type replySourceDataContent struct { + IPAddr ipAddr + Poll int16 + Stratum uint16 + State SourceStateType + Mode ModeType + Flags uint16 + Reachability uint16 + SinceSample uint32 + OrigLatestMeas chronyFloat + LatestMeas chronyFloat + LatestMeasErr chronyFloat +} + +// SourceData contains parsed version of 'source data' reply +type SourceData struct { + IPAddr net.IP + Poll int16 + Stratum uint16 + State SourceStateType + Mode ModeType + Flags uint16 + Reachability uint16 + SinceSample uint32 + OrigLatestMeas float64 + LatestMeas float64 + LatestMeasErr float64 +} + +func newSourceData(r *replySourceDataContent) *SourceData { + return &SourceData{ + IPAddr: r.IPAddr.ToNetIP(), + Poll: r.Poll, + Stratum: r.Stratum, + State: r.State, + Mode: r.Mode, + Flags: r.Flags, + Reachability: r.Reachability, + SinceSample: r.SinceSample, + OrigLatestMeas: r.OrigLatestMeas.ToFloat(), + LatestMeas: r.LatestMeas.ToFloat(), + LatestMeasErr: r.LatestMeasErr.ToFloat(), + } +} + +// ReplySourceData is a usable version of 'source data' reply for given source id +type ReplySourceData struct { + ReplyHead + SourceData +} + +type replyTrackingContent struct { + RefID uint32 + IPAddr ipAddr // our current sync source + Stratum uint16 + LeapStatus uint16 + RefTime timeSpec + CurrentCorrection chronyFloat + LastOffset chronyFloat + RMSOffset chronyFloat + FreqPPM chronyFloat + ResidFreqPPM chronyFloat + SkewPPM chronyFloat + RootDelay chronyFloat + RootDispersion chronyFloat + LastUpdateInterval chronyFloat +} + +// Tracking contains parsed version of 'tracking' reply +type Tracking struct { + RefID uint32 + IPAddr net.IP + Stratum uint16 + LeapStatus uint16 + RefTime time.Time + CurrentCorrection float64 + LastOffset float64 + RMSOffset float64 + FreqPPM float64 + ResidFreqPPM float64 + SkewPPM float64 + RootDelay float64 + RootDispersion float64 + LastUpdateInterval float64 +} + +func newTracking(r *replyTrackingContent) *Tracking { + return &Tracking{ + RefID: r.RefID, + IPAddr: r.IPAddr.ToNetIP(), + Stratum: r.Stratum, + LeapStatus: r.LeapStatus, + RefTime: r.RefTime.ToTime(), + CurrentCorrection: r.CurrentCorrection.ToFloat(), + LastOffset: r.LastOffset.ToFloat(), + RMSOffset: r.RMSOffset.ToFloat(), + FreqPPM: r.FreqPPM.ToFloat(), + ResidFreqPPM: r.ResidFreqPPM.ToFloat(), + SkewPPM: r.SkewPPM.ToFloat(), + RootDelay: r.RootDelay.ToFloat(), + RootDispersion: r.RootDispersion.ToFloat(), + LastUpdateInterval: r.LastUpdateInterval.ToFloat(), + } +} + +// ReplyTracking has usable 'tracking' response +type ReplyTracking struct { + ReplyHead + Tracking +} + +type replySourceStatsContent struct { + RefID uint32 + IPAddr ipAddr + NSamples uint32 + NRuns uint32 + SpanSeconds uint32 + StandardDeviation chronyFloat + ResidFreqPPM chronyFloat + SkewPPM chronyFloat + EstimatedOffset chronyFloat + EstimatedOffsetErr chronyFloat +} + +// SourceStats contains stats about the source +type SourceStats struct { + RefID uint32 + IPAddr net.IP + NSamples uint32 + NRuns uint32 + SpanSeconds uint32 + StandardDeviation float64 + ResidFreqPPM float64 + SkewPPM float64 + EstimatedOffset float64 + EstimatedOffsetErr float64 +} + +func newSourceStats(r *replySourceStatsContent) *SourceStats { + return &SourceStats{ + RefID: r.RefID, + IPAddr: r.IPAddr.ToNetIP(), + NSamples: r.NSamples, + NRuns: r.NRuns, + SpanSeconds: r.SpanSeconds, + StandardDeviation: r.StandardDeviation.ToFloat(), + ResidFreqPPM: r.ResidFreqPPM.ToFloat(), + SkewPPM: r.SkewPPM.ToFloat(), + EstimatedOffset: r.EstimatedOffset.ToFloat(), + EstimatedOffsetErr: r.EstimatedOffsetErr.ToFloat(), + } +} + +// ReplySourceStats has usable 'sourcestats' response +type ReplySourceStats struct { + ReplyHead + SourceStats +} + +type replyNTPDataContent struct { + RemoteAddr ipAddr + LocalAddr ipAddr + RemotePort uint16 + Leap uint8 + Version uint8 + Mode uint8 + Stratum uint8 + Poll int8 + Precision int8 + RootDelay chronyFloat + RootDispersion chronyFloat + RefID uint32 + RefTime timeSpec + Offset chronyFloat + PeerDelay chronyFloat + PeerDispersion chronyFloat + ResponseTime chronyFloat + JitterAsymmetry chronyFloat + Flags uint16 + TXTssChar uint8 + RXTssChar uint8 + TotalTXCount uint32 + TotalRXCount uint32 + TotalValidCount uint32 + Reserved [4]uint32 +} + +// NTPData contains parsed version of 'ntpdata' reply +type NTPData struct { + RemoteAddr net.IP + LocalAddr net.IP + RemotePort uint16 + Leap uint8 + Version uint8 + Mode uint8 + Stratum uint8 + Poll int8 + Precision int8 + RootDelay float64 + RootDispersion float64 + RefID uint32 + RefTime time.Time + Offset float64 + PeerDelay float64 + PeerDispersion float64 + ResponseTime float64 + JitterAsymmetry float64 + Flags uint16 + TXTssChar uint8 + RXTssChar uint8 + TotalTXCount uint32 + TotalRXCount uint32 + TotalValidCount uint32 +} + +func newNTPData(r *replyNTPDataContent) *NTPData { + return &NTPData{ + RemoteAddr: r.RemoteAddr.ToNetIP(), + LocalAddr: r.LocalAddr.ToNetIP(), + RemotePort: r.RemotePort, + Leap: r.Leap, + Version: r.Version, + Mode: r.Mode, + Stratum: r.Stratum, + Poll: r.Poll, + Precision: r.Precision, + RootDelay: r.RootDelay.ToFloat(), + RootDispersion: r.RootDispersion.ToFloat(), + RefID: r.RefID, + RefTime: r.RefTime.ToTime(), + Offset: r.Offset.ToFloat(), + PeerDelay: r.PeerDelay.ToFloat(), + PeerDispersion: r.PeerDispersion.ToFloat(), + ResponseTime: r.ResponseTime.ToFloat(), + JitterAsymmetry: r.JitterAsymmetry.ToFloat(), + Flags: r.Flags, + TXTssChar: r.TXTssChar, + RXTssChar: r.RXTssChar, + TotalTXCount: r.TotalTXCount, + TotalRXCount: r.TotalRXCount, + TotalValidCount: r.TotalValidCount, + } +} + +// ReplyNTPData is a what end user will get in 'ntp data' response +type ReplyNTPData struct { + ReplyHead + NTPData +} + +type replyNTPData2Content struct { + RemoteAddr ipAddr + LocalAddr ipAddr + RemotePort uint16 + Leap uint8 + Version uint8 + Mode uint8 + Stratum uint8 + Poll int8 + Precision int8 + RootDelay chronyFloat + RootDispersion chronyFloat + RefID uint32 + RefTime timeSpec + Offset chronyFloat + PeerDelay chronyFloat + PeerDispersion chronyFloat + ResponseTime chronyFloat + JitterAsymmetry chronyFloat + Flags uint16 + TXTssChar uint8 + RXTssChar uint8 + TotalTXCount uint32 + TotalRXCount uint32 + TotalValidCount uint32 + TotalKernelTXts uint32 + TotalKernelRXts uint32 + TotalHWTXts uint32 + TotalHWRXts uint32 + Reserved [4]int32 +} + +// NTPData2 contains parsed version of a new 'ntpdata' reply +type NTPData2 struct { + NTPData + + TotalKernelTXts uint32 + TotalKernelRXts uint32 + TotalHWTXts uint32 + TotalHWRXts uint32 +} + +func newNTPData2(r *replyNTPData2Content) *NTPData2 { + return &NTPData2{ + NTPData: NTPData{ + RemoteAddr: r.RemoteAddr.ToNetIP(), + LocalAddr: r.LocalAddr.ToNetIP(), + RemotePort: r.RemotePort, + Leap: r.Leap, + Version: r.Version, + Mode: r.Mode, + Stratum: r.Stratum, + Poll: r.Poll, + Precision: r.Precision, + RootDelay: r.RootDelay.ToFloat(), + RootDispersion: r.RootDispersion.ToFloat(), + RefID: r.RefID, + RefTime: r.RefTime.ToTime(), + Offset: r.Offset.ToFloat(), + PeerDelay: r.PeerDelay.ToFloat(), + PeerDispersion: r.PeerDispersion.ToFloat(), + ResponseTime: r.ResponseTime.ToFloat(), + JitterAsymmetry: r.JitterAsymmetry.ToFloat(), + Flags: r.Flags, + TXTssChar: r.TXTssChar, + RXTssChar: r.RXTssChar, + TotalTXCount: r.TotalTXCount, + TotalRXCount: r.TotalRXCount, + TotalValidCount: r.TotalValidCount, + }, + TotalKernelTXts: r.TotalKernelTXts, + TotalKernelRXts: r.TotalKernelRXts, + TotalHWTXts: r.TotalHWTXts, + TotalHWRXts: r.TotalHWRXts, + } +} + +// ReplyNTPData2 is a what end user will get in 'ntp data' response +type ReplyNTPData2 struct { + ReplyHead + NTPData2 +} + +type replyNTPSourceNameContent struct { + Name [256]uint8 +} + +// NTPSourceName contains parsed version of 'sourcename' reply +type NTPSourceName struct { + Name string +} + +func newNTPSourceName(r *replyNTPSourceNameContent) *NTPSourceName { + return &NTPSourceName{ + // this field is zero padded in chrony, so we need to trim it + Name: string(bytes.TrimRight(r.Name[:], "\x00")), + } +} + +// ReplyNTPSourceName is a what end user will get in 'sourcename' response +type ReplyNTPSourceName struct { + ReplyHead + NTPSourceName +} + +// Activity contains parsed version of 'activity' reply +type Activity struct { + Online int32 + Offline int32 + BurstOnline int32 + BurstOffline int32 + Unresolved int32 +} + +// ReplyActivity is a usable version of 'activity' response +type ReplyActivity struct { + ReplyHead + Activity +} + +// ServerStats contains parsed version of 'serverstats' reply +type ServerStats struct { + NTPHits uint32 + CMDHits uint32 + NTPDrops uint32 + CMDDrops uint32 + LogDrops uint32 +} + +// ReplyServerStats is a usable version of 'serverstats' response +type ReplyServerStats struct { + ReplyHead + ServerStats +} + +// ServerStats2 contains parsed version of 'serverstats2' reply +type ServerStats2 struct { + NTPHits uint32 + NKEHits uint32 + CMDHits uint32 + NTPDrops uint32 + NKEDrops uint32 + CMDDrops uint32 + LogDrops uint32 + NTPAuthHits uint32 +} + +// ReplyServerStats2 is a usable version of 'serverstats2' response +type ReplyServerStats2 struct { + ReplyHead + ServerStats2 +} + +// ServerStats3 contains parsed version of 'serverstats3' reply +type ServerStats3 struct { + NTPHits uint32 + NKEHits uint32 + CMDHits uint32 + NTPDrops uint32 + NKEDrops uint32 + CMDDrops uint32 + LogDrops uint32 + NTPAuthHits uint32 + NTPInterleavedHits uint32 + NTPTimestamps uint32 + NTPSpanSeconds uint32 +} + +// ReplyServerStats3 is a usable version of 'serverstats3' response +type ReplyServerStats3 struct { + ReplyHead + ServerStats3 +} + +// ServerStats4 contains parsed version of 'serverstats4' reply +type ServerStats4 struct { + NTPHits uint64 + NKEHits uint64 + CMDHits uint64 + NTPDrops uint64 + NKEDrops uint64 + CMDDrops uint64 + LogDrops uint64 + NTPAuthHits uint64 + NTPInterleavedHits uint64 + NTPTimestamps uint64 + NTPSpanSeconds uint64 + NTPDaemonRxtimestamps uint64 + NTPDaemonTxtimestamps uint64 + NTPKernelRxtimestamps uint64 + NTPKernelTxtimestamps uint64 + NTPHwRxTimestamps uint64 + NTPHwTxTimestamps uint64 +} + +// ReplyServerStats4 is a usable version of 'serverstats4' response +type ReplyServerStats4 struct { + ReplyHead + ServerStats4 +} + +type replySelectData struct { + RefID uint32 + IPAddr ipAddr + StateChar uint8 + Authentication uint8 + Leap uint8 + Pad uint8 + ConfOptions uint16 + EFFOptions uint16 + LastSampleAgo uint32 + Score chronyFloat + LoLimit chronyFloat + HiLimit chronyFloat +} + +// SelectData contains parsed version of 'selectdata' reply +type SelectData struct { + RefID uint32 + IPAddr net.IP + StateChar uint8 + Authentication uint8 + Leap uint8 + ConfOptions uint16 + EFFOptions uint16 + LastSampleAgo uint32 + Score float64 + LoLimit float64 + HiLimit float64 +} + +// ReplySelectData is a usable version of 'selectdata' response +type ReplySelectData struct { + ReplyHead + SelectData +} + +func newSelectData(r *replySelectData) *SelectData { + return &SelectData{ + RefID: r.RefID, + IPAddr: r.IPAddr.ToNetIP(), + StateChar: r.StateChar, + Authentication: r.Authentication, + Leap: r.Leap, + ConfOptions: r.ConfOptions, + EFFOptions: r.EFFOptions, + LastSampleAgo: r.LastSampleAgo, + Score: r.Score.ToFloat(), + LoLimit: r.LoLimit.ToFloat(), + HiLimit: r.HiLimit.ToFloat(), + } +} + +// here go request constructors + +// NewSourcesPacket creates new packet to request number of sources (peers) +func NewSourcesPacket() *RequestSources { + return &RequestSources{ + RequestHead: RequestHead{ + Version: protoVersionNumber, + PKTType: pktTypeCmdRequest, + Command: reqNSources, + }, + data: [maxDataLen]uint8{}, + } +} + +// NewTrackingPacket creates new packet to request 'tracking' information +func NewTrackingPacket() *RequestTracking { + return &RequestTracking{ + RequestHead: RequestHead{ + Version: protoVersionNumber, + PKTType: pktTypeCmdRequest, + Command: reqTracking, + }, + data: [maxDataLen]uint8{}, + } +} + +// NewSourceStatsPacket creates a new packet to request 'sourcestats' information +func NewSourceStatsPacket(sourceID int32) *RequestSourceStats { + return &RequestSourceStats{ + RequestHead: RequestHead{ + Version: protoVersionNumber, + PKTType: pktTypeCmdRequest, + Command: reqSourceStats, + }, + Index: sourceID, + data: [maxDataLen - 4]uint8{}, + } +} + +// NewSourceDataPacket creates new packet to request 'source data' information about source with given ID +func NewSourceDataPacket(sourceID int32) *RequestSourceData { + return &RequestSourceData{ + RequestHead: RequestHead{ + Version: protoVersionNumber, + PKTType: pktTypeCmdRequest, + Command: reqSourceData, + }, + Index: sourceID, + data: [maxDataLen - 4]uint8{}, + } +} + +// NewNTPDataPacket creates new packet to request 'ntp data' information for given peer IP +func NewNTPDataPacket(ip net.IP) *RequestNTPData { + return &RequestNTPData{ + RequestHead: RequestHead{ + Version: protoVersionNumber, + PKTType: pktTypeCmdRequest, + Command: reqNTPData, + }, + IPAddr: *newIPAddr(ip), + data: [maxDataLen - 16]uint8{}, + } +} + +// NewNTPSourceNamePacket creates new packet to request 'source name' information for given peer IP +func NewNTPSourceNamePacket(ip net.IP) *RequestNTPSourceName { + return &RequestNTPSourceName{ + RequestHead: RequestHead{ + Version: protoVersionNumber, + PKTType: pktTypeCmdRequest, + Command: reqNTPSourceName, + }, + IPAddr: *newIPAddr(ip), + data: [maxDataLen - 16]uint8{}, + } +} + +// NewServerStatsPacket creates new packet to request 'serverstats' information +func NewServerStatsPacket() *RequestServerStats { + return &RequestServerStats{ + RequestHead: RequestHead{ + Version: protoVersionNumber, + PKTType: pktTypeCmdRequest, + Command: reqServerStats, + }, + data: [maxDataLen]uint8{}, + } +} + +// NewActivityPacket creates new packet to request 'activity' information +func NewActivityPacket() *RequestActivity { + return &RequestActivity{ + RequestHead: RequestHead{ + Version: protoVersionNumber, + PKTType: pktTypeCmdRequest, + Command: reqActivity, + }, + data: [maxDataLen]uint8{}, + } +} + +// NewSelectDataPacket creates new packet to request 'selectdata' information +func NewSelectDataPacket(sourceID int32) *RequestSelectData { + return &RequestSelectData{ + RequestHead: RequestHead{ + Version: protoVersionNumber, + PKTType: pktTypeCmdRequest, + Command: reqSelectData, + }, + Index: sourceID, + data: [maxDataLen - 4]uint8{}, + } +} + +// possible clock sources +const ( + ClockSourceUnspec = "unspec" + ClockSourcePPS = "pps" + ClockSourceLFRadio = "lf_radio" + ClockSourceHFRadio = "hf_radio" + ClockSourceUHFRadio = "uhf_radio" + ClockSourceLocal = "local" + ClockSourceNTP = "ntp" + ClockSourceOther = "other" + ClockSourceWristWatch = "wristwatch" + ClockSourceTelephone = "telephone" +) + +// ClockSourceDesc stores human-readable descriptions of ClockSource field +var ClockSourceDesc = [10]string{ + ClockSourceUnspec, // 00 + ClockSourcePPS, // 01 + ClockSourceLFRadio, // 02 + ClockSourceHFRadio, // 03 + ClockSourceUHFRadio, // 04 + ClockSourceLocal, // 05 + ClockSourceNTP, // 06 + ClockSourceOther, // 07 + ClockSourceWristWatch, // 08 + ClockSourceTelephone, // 09 +} + +// decodePacket decodes bytes to valid response packet. +// an easy way to test this is to use 'testchrony' tool we have. +func decodePacket(response []byte) (ResponsePacket, error) { + var err error + r := bytes.NewReader(response) + head := new(ReplyHead) + if err = binary.Read(r, binary.BigEndian, head); err != nil { + return nil, err + } + Logger.Printf("response head: %+v", head) + if head.Status != sttSuccess { + return nil, fmt.Errorf("got status %s (%d)", head.Status, head.Status) + } + switch head.Reply { + case RpyNSources: + data := new(replySourcesContent) + if err = binary.Read(r, binary.BigEndian, data); err != nil { + return nil, err + } + Logger.Printf("response data: %+v", data) + return &ReplySources{ + ReplyHead: *head, + NSources: int(data.NSources), + }, nil + case RpySourceData: + data := new(replySourceDataContent) + if err = binary.Read(r, binary.BigEndian, data); err != nil { + return nil, err + } + Logger.Printf("response data: %+v", data) + return &ReplySourceData{ + ReplyHead: *head, + SourceData: *newSourceData(data), + }, nil + case RpyTracking: + data := new(replyTrackingContent) + if err = binary.Read(r, binary.BigEndian, data); err != nil { + return nil, err + } + Logger.Printf("response data: %+v", data) + return &ReplyTracking{ + ReplyHead: *head, + Tracking: *newTracking(data), + }, nil + case RpySourceStats: + data := new(replySourceStatsContent) + if err = binary.Read(r, binary.BigEndian, data); err != nil { + return nil, err + } + Logger.Printf("response data: %+v", data) + return &ReplySourceStats{ + ReplyHead: *head, + SourceStats: *newSourceStats(data), + }, nil + case RpyActivity: + data := new(Activity) + if err = binary.Read(r, binary.BigEndian, data); err != nil { + return nil, err + } + Logger.Printf("response data: %+v", data) + return &ReplyActivity{ + ReplyHead: *head, + Activity: *data, + }, nil + case RpyServerStats: + data := new(ServerStats) + if err = binary.Read(r, binary.BigEndian, data); err != nil { + return nil, err + } + Logger.Printf("response data: %+v", data) + return &ReplyServerStats{ + ReplyHead: *head, + ServerStats: *data, + }, nil + case RpyNTPData: + data := new(replyNTPDataContent) + if err = binary.Read(r, binary.BigEndian, data); err != nil { + return nil, err + } + Logger.Printf("response data: %+v", data) + return &ReplyNTPData{ + ReplyHead: *head, + NTPData: *newNTPData(data), + }, nil + case RpyNTPData2: + data := new(replyNTPData2Content) + if err = binary.Read(r, binary.BigEndian, data); err != nil { + return nil, err + } + Logger.Printf("response data: %+v", data) + return &ReplyNTPData2{ + ReplyHead: *head, + NTPData2: *newNTPData2(data), + }, nil + case RpyNTPSourceName: + data := new(replyNTPSourceNameContent) + if err = binary.Read(r, binary.BigEndian, data); err != nil { + return nil, err + } + Logger.Printf("response data: %+v", data) + return &ReplyNTPSourceName{ + ReplyHead: *head, + NTPSourceName: *newNTPSourceName(data), + }, nil + case RpyServerStats2: + data := new(ServerStats2) + if err = binary.Read(r, binary.BigEndian, data); err != nil { + return nil, err + } + Logger.Printf("response data: %+v", data) + return &ReplyServerStats2{ + ReplyHead: *head, + ServerStats2: *data, + }, nil + case RpyServerStats3: + data := new(ServerStats3) + if err = binary.Read(r, binary.BigEndian, data); err != nil { + return nil, err + } + Logger.Printf("response data: %+v", data) + return &ReplyServerStats3{ + ReplyHead: *head, + ServerStats3: *data, + }, nil + case RpyServerStats4: + data := new(ServerStats4) + if err = binary.Read(r, binary.BigEndian, data); err != nil { + return nil, err + } + Logger.Printf("response data: %+v", data) + return &ReplyServerStats4{ + ReplyHead: *head, + ServerStats4: *data, + }, nil + case RpySelectData: + data := new(replySelectData) + if err = binary.Read(r, binary.BigEndian, data); err != nil { + return nil, err + } + Logger.Printf("response data: %+v", data) + return &ReplySelectData{ + ReplyHead: *head, + SelectData: *newSelectData(data), + }, nil + default: + return nil, fmt.Errorf("not implemented reply type %d from %+v", head.Reply, head) + } +} diff --git a/src/yangerd/vendor/github.com/facebook/time/ptp/protocol/doc.go b/src/yangerd/vendor/github.com/facebook/time/ptp/protocol/doc.go new file mode 100644 index 000000000..e3bc89a95 --- /dev/null +++ b/src/yangerd/vendor/github.com/facebook/time/ptp/protocol/doc.go @@ -0,0 +1,67 @@ +/* +Copyright (c) Facebook, Inc. and its affiliates. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +/* +Package protocol implements a subset of the PTPv2.1 protocol (IEEE 1588-2019). + +Implementation is focused on unicast communications over IPv6 and is sufficient to build unicast PTP server or client. + +This package also contains basic management client that can be used to exchange Management Packets +with ptp server. + +All references throughout the code relate to the IEEE 1588-2019 Standard. + +Implemented protocol parts include: + +Marshalling and unmarshalling of defined PTP messages + + Sync + Delay_Req + Pdelay_Req + Pdelay_Resp + Follow_Up + Delay_Resp + Pdelay_Resp_Follow_Up + Announce + Signaling + Management + +TLVs + + MANAGEMENT + MANAGEMENT_ERROR_STATUS + REQUEST_UNICAST_TRANSMISSION + GRANT_UNICAST_TRANSMISSION + CANCEL_UNICAST_TRANSMISSION + ACKNOWLEDGE_CANCEL_UNICAST_TRANSMISSION + PATH_TRACE + ALTERNATE_TIME_OFFSET_INDICATOR + +Management TLVs + + DEFAULT_DATA_SET + CURRENT_DATA_SET + PARENT_DATA_SET + +Non-portable ptp4l-specific Management TLVs + + TIME_STATUS_NP + PORT_PROPERTIES_NP + PORT_STATS_NP + PORT_SERVICE_STATS_NP + UNICAST_MASTER_TABLE_NP +*/ +package protocol diff --git a/src/yangerd/vendor/github.com/facebook/time/ptp/protocol/management.go b/src/yangerd/vendor/github.com/facebook/time/ptp/protocol/management.go new file mode 100644 index 000000000..2d9b74b12 --- /dev/null +++ b/src/yangerd/vendor/github.com/facebook/time/ptp/protocol/management.go @@ -0,0 +1,300 @@ +/* +Copyright (c) Facebook, Inc. and its affiliates. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package protocol + +import ( + "bytes" + "encoding" + "encoding/binary" + "errors" + "fmt" + "io" + "os" +) + +var identity PortIdentity + +// ErrManagementMsgErrorStatus is what happens if we expected to get Management TLV in response, but received special ManagementErrorStatusTLV +var ErrManagementMsgErrorStatus = errors.New("received MANAGEMENT_ERROR_STATUS_TLV") + +func init() { + // store our PID as identity that we use to talk to ptp daemon + identity.PortNumber = uint16(os.Getpid()) +} + +// Action indicate the action to be taken on receipt of the PTP message as defined in Table 57 +type Action uint8 + +// actions as in Table 57 Values of the actionField +const ( + GET Action = iota + SET + RESPONSE + COMMAND + ACKNOWLEDGE +) + +// ManagementTLVHead Spec Table 58 - Management TLV fields +type ManagementTLVHead struct { + TLVHead + + ManagementID ManagementID +} + +// ManagementMsgHead Spec Table 56 - Management message fields +type ManagementMsgHead struct { + Header + + TargetPortIdentity PortIdentity + StartingBoundaryHops uint8 + BoundaryHops uint8 + ActionField Action + Reserved uint8 +} + +// Action returns ActionField +func (p *ManagementMsgHead) Action() Action { + return p.ActionField +} + +// MgmtID returns ManagementID +func (p *ManagementTLVHead) MgmtID() ManagementID { + return p.ManagementID +} + +// Management packet, see '15. PTP management messages' +type Management struct { + ManagementMsgHead + TLV ManagementTLV +} + +// UnmarshalBinary parses []byte and populates struct fields +func (p *Management) UnmarshalBinary(rawBytes []byte) error { + var err error + head := ManagementMsgHead{} + tlvHead := ManagementTLVHead{} + r := bytes.NewReader(rawBytes) + if err = binary.Read(r, binary.BigEndian, &head); err != nil { + return err + } + if err = binary.Read(r, binary.BigEndian, &tlvHead.TLVHead); err != nil { + return err + } + if tlvHead.TLVType == TLVManagementErrorStatus { + return ErrManagementMsgErrorStatus + } + if tlvHead.TLVType != TLVManagement { + return fmt.Errorf("got TLV type %q (0x%02X) instead of %q (0x%02X)", tlvHead.TLVType.String(), int(tlvHead.TLVType), TLVManagement.String(), int(TLVManagement)) + } + + if err = binary.Read(r, binary.BigEndian, &tlvHead.ManagementID); err != nil { + return err + } + headSize := binary.Size(tlvHead) + // seek back so we can read whole TLV + if _, err := r.Seek(-int64(headSize), io.SeekCurrent); err != nil { + return err + } + decoder, found := mgmtTLVDecoder[tlvHead.ManagementID] + if !found { + return fmt.Errorf("unsupported management TLV 0x%x", tlvHead.ManagementID) + } + tlvData, err := io.ReadAll(r) + if err != nil { + return err + } + tlv, err := decoder(tlvData) + if err != nil { + return err + } + p.ManagementMsgHead = head + p.TLV = tlv + return nil +} + +// MarshalBinaryToBuf converts packet to bytes and writes those into provided buffer +func (p *Management) MarshalBinaryToBuf(bytes io.Writer) error { + if err := binary.Write(bytes, binary.BigEndian, p.ManagementMsgHead); err != nil { + return err + } + // interface smuggling + if pp, ok := p.TLV.(encoding.BinaryMarshaler); ok { + b, err := pp.MarshalBinary() + if err != nil { + return err + } + return binary.Write(bytes, binary.BigEndian, b) + } + return binary.Write(bytes, binary.BigEndian, p.TLV) +} + +// MarshalBinary converts packet to []bytes +func (p *Management) MarshalBinary() ([]byte, error) { + var bytes bytes.Buffer + err := p.MarshalBinaryToBuf(&bytes) + return bytes.Bytes(), err +} + +// ManagementErrorStatusTLV spec Table 108 MANAGEMENT_ERROR_STATUS TLV format +type ManagementErrorStatusTLV struct { + TLVHead + + ManagementErrorID ManagementErrorID + ManagementID ManagementID + Reserved int32 + DisplayData PTPText +} + +// ManagementMsgErrorStatus is header + ManagementErrorStatusTLV +type ManagementMsgErrorStatus struct { + ManagementMsgHead + ManagementErrorStatusTLV +} + +// UnmarshalBinary parses []byte and populates struct fields +func (p *ManagementMsgErrorStatus) UnmarshalBinary(rawBytes []byte) error { + reader := bytes.NewReader(rawBytes) + be := binary.BigEndian + if err := binary.Read(reader, be, &p.ManagementMsgHead); err != nil { + return fmt.Errorf("reading ManagementMsgErrorStatus ManagementMsgHead: %w", err) + } + if err := binary.Read(reader, be, &p.ManagementErrorStatusTLV.TLVHead); err != nil { + return fmt.Errorf("reading ManagementMsgErrorStatus TLVHead: %w", err) + } + if err := binary.Read(reader, be, &p.ManagementErrorStatusTLV.ManagementErrorID); err != nil { + return fmt.Errorf("reading ManagementMsgErrorStatus ManagementErrorID: %w", err) + } + if err := binary.Read(reader, be, &p.ManagementErrorStatusTLV.ManagementID); err != nil { + return fmt.Errorf("reading ManagementMsgErrorStatus ManagementID: %w", err) + } + if err := binary.Read(reader, be, &p.ManagementErrorStatusTLV.Reserved); err != nil { + return fmt.Errorf("reading ManagementMsgErrorStatus Reserved: %w", err) + } + // packet can have trailing bytes, let's make sure we don't try to read past given length + toRead := int(p.ManagementMsgHead.Header.MessageLength) + toRead -= binary.Size(p.ManagementMsgHead) + toRead -= binary.Size(p.ManagementErrorStatusTLV.TLVHead) + toRead -= binary.Size(p.ManagementErrorStatusTLV.ManagementErrorID) + toRead -= binary.Size(p.ManagementErrorStatusTLV.ManagementID) + toRead -= binary.Size(p.ManagementErrorStatusTLV.Reserved) + + if reader.Len() == 0 || toRead <= 0 { + // DisplayData is completely optional + return nil + } + data := make([]byte, reader.Len()) + if _, err := io.ReadFull(reader, data); err != nil { + return err + } + if err := p.DisplayData.UnmarshalBinary(data); err != nil { + return fmt.Errorf("reading ManagementMsgErrorStatus DisplayData: %w", err) + } + return nil +} + +// MarshalBinaryToBuf converts packet to bytes and writes those into provided buffer +func (p *ManagementMsgErrorStatus) MarshalBinaryToBuf(bytes io.Writer) error { + be := binary.BigEndian + if err := binary.Write(bytes, be, &p.ManagementMsgHead); err != nil { + return fmt.Errorf("writing ManagementMsgErrorStatus ManagementMsgHead: %w", err) + } + if err := binary.Write(bytes, be, &p.ManagementErrorStatusTLV.TLVHead); err != nil { + return fmt.Errorf("writing ManagementMsgErrorStatus TLVHead: %w", err) + } + if err := binary.Write(bytes, be, &p.ManagementErrorStatusTLV.ManagementErrorID); err != nil { + return fmt.Errorf("writing ManagementMsgErrorStatus ManagementErrorID: %w", err) + } + if err := binary.Write(bytes, be, &p.ManagementErrorStatusTLV.ManagementID); err != nil { + return fmt.Errorf("writing ManagementMsgErrorStatus ManagementID: %w", err) + } + if err := binary.Write(bytes, be, &p.ManagementErrorStatusTLV.Reserved); err != nil { + return fmt.Errorf("writing ManagementMsgErrorStatus Reserved: %w", err) + } + if p.DisplayData != "" { + dd, err := p.DisplayData.MarshalBinary() + if err != nil { + return fmt.Errorf("writing ManagementMsgErrorStatus DisplayData: %w", err) + } + if _, err := bytes.Write(dd); err != nil { + return err + } + } + return nil +} + +// MarshalBinary converts packet to []bytes +func (p *ManagementMsgErrorStatus) MarshalBinary() ([]byte, error) { + var bytes bytes.Buffer + err := p.MarshalBinaryToBuf(&bytes) + return bytes.Bytes(), err +} + +// ManagementErrorID is an enum for possible management errors +type ManagementErrorID uint16 + +// Table 109 ManagementErrorID enumeration +const ( + ErrorResponseTooBig ManagementErrorID = 0x0001 // The requested operation could not fit in a single response message + ErrorNoSuchID ManagementErrorID = 0x0002 // The managementId is not recognized + ErrorWrongLength ManagementErrorID = 0x0003 // The managementId was identified but the length of the data was wrong + ErrorWrongValue ManagementErrorID = 0x0004 // The managementId and length were correct but one or more values were wrong + ErrorNotSetable ManagementErrorID = 0x0005 // Some of the variables in the set command were not updated because they are not configurable + ErrorNotSupported ManagementErrorID = 0x0006 // The requested operation is not supported in this PTP Instance + ErrorUnpopulated ManagementErrorID = 0x0007 // The targetPortIdentity of the PTP management message refers to an entity that is not present in the PTP Instance at the time of the request + // some reserved and provile-specific ranges + ErrorGeneralError ManagementErrorID = 0xFFFE //An error occurred that is not covered by other ManagementErrorID values +) + +// ManagementErrorIDToString is a map from ManagementErrorID to string +var ManagementErrorIDToString = map[ManagementErrorID]string{ + ErrorResponseTooBig: "RESPONSE_TOO_BIG", + ErrorNoSuchID: "NO_SUCH_ID", + ErrorWrongLength: "WRONG_LENGTH", + ErrorWrongValue: "WRONG_VALUE", + ErrorNotSetable: "NOT_SETABLE", + ErrorNotSupported: "NOT_SUPPORTED", + ErrorUnpopulated: "UNPOPULATED", + ErrorGeneralError: "GENERAL_ERROR", +} + +func (t ManagementErrorID) String() string { + s := ManagementErrorIDToString[t] + if s == "" { + return fmt.Sprintf("UNKNOWN_ERROR_ID=%d", t) + } + return s +} + +func (t ManagementErrorID) Error() string { + return t.String() +} + +func decodeMgmtPacket(data []byte) (Packet, error) { + packet := &Management{} + err := packet.UnmarshalBinary(data) + if errors.Is(err, ErrManagementMsgErrorStatus) { + errorPacket := new(ManagementMsgErrorStatus) + if err := errorPacket.UnmarshalBinary(data); err != nil { + return nil, fmt.Errorf("got Management Error in response but failed to decode it: %w", err) + } + return errorPacket, nil + } + if err != nil { + return nil, err + } + return packet, nil +} diff --git a/src/yangerd/vendor/github.com/facebook/time/ptp/protocol/management_client.go b/src/yangerd/vendor/github.com/facebook/time/ptp/protocol/management_client.go new file mode 100644 index 000000000..947152d02 --- /dev/null +++ b/src/yangerd/vendor/github.com/facebook/time/ptp/protocol/management_client.go @@ -0,0 +1,125 @@ +/* +Copyright (c) Facebook, Inc. and its affiliates. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package protocol + +// management client is used to talk to (presumably local) PTP server using Management packets + +import ( + "encoding/binary" + "fmt" + "io" +) + +// MgmtClient talks to ptp server over unix socket +type MgmtClient struct { + Connection io.ReadWriter + Sequence uint16 +} + +// SendPacket sends packet, incrementing sequence counter +func (c *MgmtClient) SendPacket(packet *Management) error { + c.Sequence++ + packet.SetSequence(c.Sequence) + b, err := packet.MarshalBinary() + if err != nil { + return err + } + return binary.Write(c.Connection, binary.BigEndian, b) +} + +// Communicate sends the management the packet, parses response into something usable +func (c *MgmtClient) Communicate(packet *Management) (*Management, error) { + var err error + + if err := c.SendPacket(packet); err != nil { + return nil, err + } + response := make([]uint8, 1024) + n, err := c.Connection.Read(response) + if err != nil { + return nil, err + } + res, err := decodeMgmtPacket(response[:n]) + if err != nil { + return nil, err + } + errorPacket, ok := res.(*ManagementMsgErrorStatus) + if ok { + return nil, fmt.Errorf("got Management Error in response: %w", errorPacket.ManagementErrorStatusTLV.ManagementErrorID) + } + p, ok := res.(*Management) + if !ok { + return nil, fmt.Errorf("got unexpected management packet %T", res) + } + return p, nil +} + +// ParentDataSet sends PARENT_DATA_SET request and returns response +func (c *MgmtClient) ParentDataSet() (*ParentDataSetTLV, error) { + req := ParentDataSetRequest() + p, err := c.Communicate(req) + if err != nil { + return nil, err + } + tlv, ok := p.TLV.(*ParentDataSetTLV) + if !ok { + return nil, fmt.Errorf("got unexpected management TLV %T, wanted %T", p.TLV, tlv) + } + return tlv, nil +} + +// DefaultDataSet sends DEFAULT_DATA_SET request and returns response +func (c *MgmtClient) DefaultDataSet() (*DefaultDataSetTLV, error) { + req := DefaultDataSetRequest() + p, err := c.Communicate(req) + if err != nil { + return nil, err + } + tlv, ok := p.TLV.(*DefaultDataSetTLV) + if !ok { + return nil, fmt.Errorf("got unexpected management TLV %T, wanted %T", p.TLV, tlv) + } + return tlv, nil +} + +// CurrentDataSet sends CURRENT_DATA_SET request and returns response +func (c *MgmtClient) CurrentDataSet() (*CurrentDataSetTLV, error) { + req := CurrentDataSetRequest() + p, err := c.Communicate(req) + if err != nil { + return nil, err + } + tlv, ok := p.TLV.(*CurrentDataSetTLV) + if !ok { + return nil, fmt.Errorf("got unexpected management TLV %T, wanted %T", p.TLV, tlv) + } + return tlv, nil +} + +// ClockAccuracy sends CLOCK_ACCURACY request and returns response +func (c *MgmtClient) ClockAccuracy() (*ClockAccuracyTLV, error) { + req := ClockAccuracyRequest() + p, err := c.Communicate(req) + if err != nil { + return nil, err + } + tlv, ok := p.TLV.(*ClockAccuracyTLV) + if !ok { + return nil, fmt.Errorf("got unexpected management TLV %T, wanted %T", p.TLV, tlv) + } + return tlv, nil +} diff --git a/src/yangerd/vendor/github.com/facebook/time/ptp/protocol/management_tlvs.go b/src/yangerd/vendor/github.com/facebook/time/ptp/protocol/management_tlvs.go new file mode 100644 index 000000000..ef1f2f3e2 --- /dev/null +++ b/src/yangerd/vendor/github.com/facebook/time/ptp/protocol/management_tlvs.go @@ -0,0 +1,362 @@ +/* +Copyright (c) Facebook, Inc. and its affiliates. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package protocol + +import ( + "bytes" + "encoding/binary" + "fmt" + + "github.com/facebook/time/hostendian" +) + +// ManagementID is type for Management IDs +type ManagementID uint16 + +// Management IDs we support, from Table 59 managementId values +const ( + IDNullPTPManagement ManagementID = 0x0000 + IDClockDescription ManagementID = 0x0001 + IDUserDescription ManagementID = 0x0002 + IDSaveInNonVolatileStorage ManagementID = 0x0003 + IDResetNonVolatileStorage ManagementID = 0x0004 + IDInitialize ManagementID = 0x0005 + IDFaultLog ManagementID = 0x0006 + IDFaultLogReset ManagementID = 0x0007 + + IDDefaultDataSet ManagementID = 0x2000 + IDCurrentDataSet ManagementID = 0x2001 + IDParentDataSet ManagementID = 0x2002 + IDTimePropertiesDataSet ManagementID = 0x2003 + IDPortDataSet ManagementID = 0x2004 + IDClockAccuracy ManagementID = 0x2010 + // rest of Management IDs that we don't implement yet +) + +// ManagementTLV abstracts away any ManagementTLV +type ManagementTLV interface { + TLV + MgmtID() ManagementID +} + +// MgmtTLVDecoderFunc is the function we use to decode management TLV from bytes +type MgmtTLVDecoderFunc func(data []byte) (ManagementTLV, error) + +// default decoders for TLVs we implemented ourselves +var mgmtTLVDecoder = map[ManagementID]MgmtTLVDecoderFunc{ + IDDefaultDataSet: func(data []byte) (ManagementTLV, error) { + r := bytes.NewReader(data) + tlv := &DefaultDataSetTLV{} + if err := binary.Read(r, binary.BigEndian, tlv); err != nil { + return nil, err + } + return tlv, nil + }, + IDCurrentDataSet: func(data []byte) (ManagementTLV, error) { + r := bytes.NewReader(data) + tlv := &CurrentDataSetTLV{} + if err := binary.Read(r, binary.BigEndian, tlv); err != nil { + return nil, err + } + return tlv, nil + }, + IDParentDataSet: func(data []byte) (ManagementTLV, error) { + r := bytes.NewReader(data) + tlv := &ParentDataSetTLV{} + if err := binary.Read(r, binary.BigEndian, tlv); err != nil { + return nil, err + } + return tlv, nil + }, + IDPortStatsNP: func(data []byte) (ManagementTLV, error) { + r := bytes.NewReader(data) + tlv := &PortStatsNPTLV{} + if err := binary.Read(r, binary.BigEndian, &tlv.ManagementTLVHead); err != nil { + return nil, err + } + if err := binary.Read(r, binary.BigEndian, &tlv.PortIdentity); err != nil { + return nil, err + } + // fun part that cost me few hours, this is sent over wire as host endian (which typically is LittlEndian), while EVERYTHING ELSE is BigEndian. + if err := binary.Read(r, hostendian.Order, &tlv.PortStats); err != nil { + return nil, err + } + return tlv, nil + }, + IDTimeStatusNP: func(data []byte) (ManagementTLV, error) { + r := bytes.NewReader(data) + tlv := &TimeStatusNPTLV{} + if err := binary.Read(r, binary.BigEndian, tlv); err != nil { + return nil, err + } + return tlv, nil + }, + IDPortServiceStatsNP: func(data []byte) (ManagementTLV, error) { + r := bytes.NewReader(data) + tlv := &PortServiceStatsNPTLV{} + if err := binary.Read(r, binary.BigEndian, &tlv.ManagementTLVHead); err != nil { + return nil, err + } + if err := binary.Read(r, binary.BigEndian, &tlv.PortIdentity); err != nil { + return nil, err + } + // host endian, just like with PortStatsNP + if err := binary.Read(r, hostendian.Order, &tlv.PortServiceStats); err != nil { + return nil, err + } + return tlv, nil + }, + IDPortPropertiesNP: func(data []byte) (ManagementTLV, error) { + r := bytes.NewReader(data) + tlv := &PortPropertiesNPTLV{} + if err := binary.Read(r, binary.BigEndian, &tlv.ManagementTLVHead); err != nil { + return nil, err + } + if err := binary.Read(r, binary.BigEndian, &tlv.PortIdentity); err != nil { + return nil, err + } + if err := binary.Read(r, binary.BigEndian, &tlv.PortState); err != nil { + return nil, err + } + if err := binary.Read(r, binary.BigEndian, &tlv.Timestamping); err != nil { + return nil, err + } + if r.Len() == 0 { + return nil, fmt.Errorf("not enough data to read PortPropertiesNP Interface") + } + if err := tlv.Interface.UnmarshalBinary(data[len(data)-r.Len():]); err != nil { + return nil, fmt.Errorf("reading PortPropertiesNP Interface: %w", err) + } + return tlv, nil + }, + IDUnicastMasterTableNP: func(data []byte) (ManagementTLV, error) { + r := bytes.NewReader(data) + tlv := &UnicastMasterTableNPTLV{} + + if err := binary.Read(r, binary.BigEndian, &tlv.ManagementTLVHead); err != nil { + return nil, err + } + if err := binary.Read(r, binary.BigEndian, &tlv.UnicastMasterTable.ActualTableSize); err != nil { + return nil, err + } + tlv.UnicastMasterTable.UnicastMasters = make([]UnicastMasterEntry, int(tlv.UnicastMasterTable.ActualTableSize)) + n := binary.Size(tlv.ManagementTLVHead) + binary.Size(tlv.UnicastMasterTable.ActualTableSize) + for i := 0; i < int(tlv.UnicastMasterTable.ActualTableSize); i++ { + entry := UnicastMasterEntry{} + if err := entry.UnmarshalBinary(data[n:]); err != nil { + return nil, err + } + tlv.UnicastMasterTable.UnicastMasters[i] = entry + n += 22 + len(entry.Address) + } + + return tlv, nil + }, + IDClockAccuracy: func(data []byte) (ManagementTLV, error) { + r := bytes.NewReader(data) + tlv := &ClockAccuracyTLV{} + + if err := binary.Read(r, binary.BigEndian, &tlv.ManagementTLVHead); err != nil { + return nil, err + } + if err := binary.Read(r, binary.BigEndian, &tlv.ClockAccuracy); err != nil { + return nil, err + } + if err := binary.Read(r, binary.BigEndian, &tlv.Reserved); err != nil { + return nil, err + } + return tlv, nil + }, +} + +// RegisterMgmtTLVDecoder registers function we'll use to decode particular custom management TLV. +// IEEE1588-2019 specifies that range C000 – DFFF should be used for implementation-specific identifiers, +// and E000 – FFFE is to be assigned by alternate PTP Profile. +func RegisterMgmtTLVDecoder(id ManagementID, decoder MgmtTLVDecoderFunc) { + mgmtTLVDecoder[id] = decoder +} + +// CurrentDataSetTLV Spec Table 84 - CURRENT_DATA_SET management TLV data field +type CurrentDataSetTLV struct { + ManagementTLVHead + + StepsRemoved uint16 + OffsetFromMaster TimeInterval + MeanPathDelay TimeInterval +} + +// DefaultDataSetTLV Spec Table 69 - DEFAULT_DATA_SET management TLV data field +type DefaultDataSetTLV struct { + ManagementTLVHead + + SoTSC uint8 + Reserved0 uint8 + NumberPorts uint16 + Priority1 uint8 + ClockQuality ClockQuality + Priority2 uint8 + ClockIdentity ClockIdentity + DomainNumber uint8 + Reserved1 uint8 +} + +// ParentDataSetTLV Spec Table 85 - PARENT_DATA_SET management TLV data field +type ParentDataSetTLV struct { + ManagementTLVHead + + ParentPortIdentity PortIdentity + PS uint8 + Reserved uint8 + ObservedParentOffsetScaledLogVariance uint16 + ObservedParentClockPhaseChangeRate uint32 + GrandmasterPriority1 uint8 + GrandmasterClockQuality ClockQuality + GrandmasterPriority2 uint8 + GrandmasterIdentity ClockIdentity +} + +// ClockAccuracyTLV is a TLV containing Clock Accuracy +type ClockAccuracyTLV struct { + ManagementTLVHead + + ClockAccuracy ClockAccuracy + Reserved uint8 +} + +// CurrentDataSetRequest prepares request packet for CURRENT_DATA_SET request +func CurrentDataSetRequest() *Management { + headerSize := uint16(binary.Size(ManagementMsgHead{})) + size := uint16(binary.Size(CurrentDataSetTLV{})) + tlvHeadSize := uint16(binary.Size(TLVHead{})) + return &Management{ + ManagementMsgHead: ManagementMsgHead{ + Header: Header{ + SdoIDAndMsgType: NewSdoIDAndMsgType(MessageManagement, 0), + Version: Version, + MessageLength: headerSize + size, + SourcePortIdentity: identity, + LogMessageInterval: MgmtLogMessageInterval, + }, + TargetPortIdentity: DefaultTargetPortIdentity, + StartingBoundaryHops: 0, + BoundaryHops: 0, + ActionField: GET, + }, + TLV: &CurrentDataSetTLV{ + ManagementTLVHead: ManagementTLVHead{ + TLVHead: TLVHead{ + TLVType: TLVManagement, + LengthField: size - tlvHeadSize, + }, + ManagementID: IDCurrentDataSet, + }, + }, + } +} + +// DefaultDataSetRequest prepares request packet for DEFAULT_DATA_SET request +func DefaultDataSetRequest() *Management { + headerSize := uint16(binary.Size(ManagementMsgHead{})) + size := uint16(binary.Size(DefaultDataSetTLV{})) + tlvHeadSize := uint16(binary.Size(TLVHead{})) + return &Management{ + ManagementMsgHead: ManagementMsgHead{ + Header: Header{ + SdoIDAndMsgType: NewSdoIDAndMsgType(MessageManagement, 0), + Version: Version, + MessageLength: headerSize + size, + SourcePortIdentity: identity, + LogMessageInterval: MgmtLogMessageInterval, + }, + TargetPortIdentity: DefaultTargetPortIdentity, + StartingBoundaryHops: 0, + BoundaryHops: 0, + ActionField: GET, + }, + TLV: &DefaultDataSetTLV{ + ManagementTLVHead: ManagementTLVHead{ + TLVHead: TLVHead{ + TLVType: TLVManagement, + LengthField: size - tlvHeadSize, + }, + ManagementID: IDDefaultDataSet, + }, + }, + } +} + +// ParentDataSetRequest prepares request packet for PARENT_DATA_SET request +func ParentDataSetRequest() *Management { + headerSize := uint16(binary.Size(ManagementMsgHead{})) + size := uint16(binary.Size(ParentDataSetTLV{})) + tlvHeadSize := uint16(binary.Size(TLVHead{})) + return &Management{ + ManagementMsgHead: ManagementMsgHead{ + Header: Header{ + SdoIDAndMsgType: NewSdoIDAndMsgType(MessageManagement, 0), + Version: Version, + MessageLength: headerSize + size, + SourcePortIdentity: identity, + LogMessageInterval: MgmtLogMessageInterval, + }, + TargetPortIdentity: DefaultTargetPortIdentity, + StartingBoundaryHops: 0, + BoundaryHops: 0, + ActionField: GET, + }, + TLV: &ParentDataSetTLV{ + ManagementTLVHead: ManagementTLVHead{ + TLVHead: TLVHead{ + TLVType: TLVManagement, + LengthField: size - tlvHeadSize, + }, + ManagementID: IDParentDataSet, + }, + }, + } +} + +// ClockAccuracyRequest prepares request packet for CLOCK_ACCURACY request +func ClockAccuracyRequest() *Management { + headerSize := uint16(binary.Size(ManagementMsgHead{})) + size := uint16(binary.Size(ClockAccuracyTLV{})) + tlvHeadSize := uint16(binary.Size(TLVHead{})) + return &Management{ + ManagementMsgHead: ManagementMsgHead{ + Header: Header{ + SdoIDAndMsgType: NewSdoIDAndMsgType(MessageManagement, 0), + Version: Version, + MessageLength: headerSize + size, + SourcePortIdentity: identity, + LogMessageInterval: MgmtLogMessageInterval, + }, + TargetPortIdentity: DefaultTargetPortIdentity, + StartingBoundaryHops: 0, + BoundaryHops: 0, + ActionField: GET, + }, + TLV: &ClockAccuracyTLV{ + ManagementTLVHead: ManagementTLVHead{ + TLVHead: TLVHead{ + TLVType: TLVManagement, + LengthField: size - tlvHeadSize, + }, + ManagementID: IDClockAccuracy, + }, + }, + } +} diff --git a/src/yangerd/vendor/github.com/facebook/time/ptp/protocol/protocol.go b/src/yangerd/vendor/github.com/facebook/time/ptp/protocol/protocol.go new file mode 100644 index 000000000..765cce60a --- /dev/null +++ b/src/yangerd/vendor/github.com/facebook/time/ptp/protocol/protocol.go @@ -0,0 +1,508 @@ +/* +Copyright (c) Facebook, Inc. and its affiliates. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package protocol + +// all references are given for IEEE 1588-2019 Standard + +import ( + "bytes" + "encoding" + "encoding/binary" + "fmt" +) + +// what version of PTP protocol we implement +const ( + MajorVersion uint8 = 2 + MinorVersion uint8 = 1 + Version uint8 = MinorVersion<<4 | MajorVersion + MajorVersionMask uint8 = 0x0f +) + +/* +UDP port numbers: +The UDP destination port of a PTP event message shall be 319. +The UDP destination port of a multicast PTP general message shall be 320. +The UDP destination port of a unicast PTP general message that is addressed to a PTP Instance shall be 320. +The UDP destination port of a unicast PTP general message that is addressed to a manager shall be the UDP source +port value of the PTP message to which this is a response. +*/ +var ( + PortEvent = 319 + PortGeneral = 320 +) + +// TrailingBytes - PTP over UDPv6 requires adding extra two bytes that +// may be modified by the initiator or an intermediate PTP Instance to ensure that the UDP checksum +// remains uncompromised after any modification of PTP fields. +// We simply always add them - in worst case they add extra 2 unused bytes when used over UDPv4. +const TrailingBytes = 2 + +var twoZeros = []byte{0, 0} + +// MgmtLogMessageInterval is the default LogInterval value used in Management packets +const MgmtLogMessageInterval LogInterval = 0x7f // as per Table 42 Values of logMessageInterval field + +// DefaultTargetPortIdentity is a port identity that means any port +var DefaultTargetPortIdentity = PortIdentity{ + ClockIdentity: 0xffffffffffffffff, + PortNumber: 0xffff, +} + +// Header Table 35 Common PTP message header +type Header struct { + SdoIDAndMsgType SdoIDAndMsgType // first 4 bits is SdoId, next 4 bytes are msgtype + Version uint8 + MessageLength uint16 + DomainNumber uint8 + MinorSdoID uint8 + FlagField uint16 + CorrectionField Correction + MessageTypeSpecific uint32 + SourcePortIdentity PortIdentity + SequenceID uint16 + ControlField uint8 // the use of this field is obsolete according to IEEE, unless it's ipv4 + LogMessageInterval LogInterval // see Table 42 Values of logMessageInterval field +} + +const headerSize = 34 // bytes + +// unmarshalHeader is not a Header.UnmarshalBinary to prevent all packets +// from having default (and incomplete) UnmarshalBinary implementation through embedding +func unmarshalHeader(p *Header, b []byte) { + p.SdoIDAndMsgType = SdoIDAndMsgType(b[0]) + p.Version = b[1] + p.MessageLength = binary.BigEndian.Uint16(b[2:]) + p.DomainNumber = b[4] + p.MinorSdoID = b[5] + p.FlagField = binary.BigEndian.Uint16(b[6:]) + p.CorrectionField = Correction(binary.BigEndian.Uint64(b[8:])) + p.MessageTypeSpecific = binary.BigEndian.Uint32(b[16:]) + p.SourcePortIdentity.ClockIdentity = ClockIdentity(binary.BigEndian.Uint64(b[20:])) + p.SourcePortIdentity.PortNumber = binary.BigEndian.Uint16(b[28:]) + p.SequenceID = binary.BigEndian.Uint16(b[30:]) + p.ControlField = b[32] + p.LogMessageInterval = LogInterval(b[33]) +} + +// MessageType returns MessageType +func (p *Header) MessageType() MessageType { + return p.SdoIDAndMsgType.MsgType() +} + +// SetSequence populates sequence field +func (p *Header) SetSequence(sequence uint16) { + p.SequenceID = sequence +} + +func checkPacketLength(p *Header, l int) error { + if int(p.MessageLength) > l { + return fmt.Errorf("cannot decode message of length %d from %d bytes", p.MessageLength, l) + } + return nil +} + +// headerMarshalBinaryTo is not a Header.MarshalBinaryTo to prevent all packets +// from having default (and incomplete) MarshalBinaryTo implementation through embedding +func headerMarshalBinaryTo(p *Header, b []byte) int { + b[0] = byte(p.SdoIDAndMsgType) + b[1] = p.Version + binary.BigEndian.PutUint16(b[2:], p.MessageLength) + b[4] = p.DomainNumber + b[5] = p.MinorSdoID + binary.BigEndian.PutUint16(b[6:], p.FlagField) + binary.BigEndian.PutUint64(b[8:], uint64(p.CorrectionField)) + binary.BigEndian.PutUint32(b[16:], p.MessageTypeSpecific) + binary.BigEndian.PutUint64(b[20:], uint64(p.SourcePortIdentity.ClockIdentity)) + binary.BigEndian.PutUint16(b[28:], p.SourcePortIdentity.PortNumber) + binary.BigEndian.PutUint16(b[30:], p.SequenceID) + b[32] = p.ControlField + b[33] = byte(p.LogMessageInterval) + return headerSize +} + +// flags used in FlagField as per Table 37 Values of flagField +const ( + // first octet + FlagAlternateMaster uint16 = 1 << (8 + 0) + FlagTwoStep uint16 = 1 << (8 + 1) + FlagUnicast uint16 = 1 << (8 + 2) + FlagProfileSpecific1 uint16 = 1 << (8 + 5) + FlagProfileSpecific2 uint16 = 1 << (8 + 6) + // second octet + FlagLeap61 uint16 = 1 << 0 + FlagLeap59 uint16 = 1 << 1 + FlagCurrentUtcOffsetValid uint16 = 1 << 2 + FlagPTPTimescale uint16 = 1 << 3 + FlagTimeTraceable uint16 = 1 << 4 + FlagFrequencyTraceable uint16 = 1 << 5 + FlagSynchronizationUncertain uint16 = 1 << 6 +) + +// General PTP messages + +// All packets are split in three parts: Header (which is common), body that is unique +// for most packets (both in length and structure), and finally a suffix of zero or more TLVs + +// AnnounceBody Table 43 Announce message fields +type AnnounceBody struct { + OriginTimestamp Timestamp + CurrentUTCOffset int16 + Reserved uint8 + GrandmasterPriority1 uint8 + GrandmasterClockQuality ClockQuality + GrandmasterPriority2 uint8 + GrandmasterIdentity ClockIdentity + StepsRemoved uint16 + TimeSource TimeSource +} + +// Announce is a full Announce packet +type Announce struct { + Header + AnnounceBody + TLVs []TLV +} + +// MarshalBinaryTo marshals bytes to Announce +func (p *Announce) MarshalBinaryTo(b []byte) (int, error) { + if len(b) < headerSize+30 { + return 0, fmt.Errorf("not enough buffer to write Announce") + } + n := headerMarshalBinaryTo(&p.Header, b) + copy(b[n:], p.OriginTimestamp.Seconds[:]) //uint48 + binary.BigEndian.PutUint32(b[n+6:], p.OriginTimestamp.Nanoseconds) + binary.BigEndian.PutUint16(b[n+10:], uint16(p.CurrentUTCOffset)) + b[n+12] = p.Reserved + b[n+13] = p.GrandmasterPriority1 + b[n+14] = byte(p.GrandmasterClockQuality.ClockClass) + b[n+15] = byte(p.GrandmasterClockQuality.ClockAccuracy) + binary.BigEndian.PutUint16(b[n+16:], p.GrandmasterClockQuality.OffsetScaledLogVariance) + b[n+18] = p.GrandmasterPriority2 + binary.BigEndian.PutUint64(b[n+19:], uint64(p.GrandmasterIdentity)) + binary.BigEndian.PutUint16(b[n+27:], p.StepsRemoved) + b[n+29] = byte(p.TimeSource) + // marshal TLVs if present + pos := n + 30 + tlvLen, err := writeTLVs(p.TLVs, b[pos:]) + return pos + tlvLen, err +} + +// UnmarshalBinary unmarshals bytes to Announce +func (p *Announce) UnmarshalBinary(b []byte) error { + if len(b) < headerSize+30 { + return fmt.Errorf("not enough data to decode Announce") + } + unmarshalHeader(&p.Header, b) + if err := checkPacketLength(&p.Header, len(b)); err != nil { + return err + } + n := headerSize + copy(p.OriginTimestamp.Seconds[:], b[n:]) //uint48 + p.OriginTimestamp.Nanoseconds = binary.BigEndian.Uint32(b[n+6:]) + p.CurrentUTCOffset = int16(binary.BigEndian.Uint16(b[n+10:])) + p.Reserved = b[n+12] + p.GrandmasterPriority1 = b[n+13] + p.GrandmasterClockQuality.ClockClass = ClockClass(b[n+14]) + p.GrandmasterClockQuality.ClockAccuracy = ClockAccuracy(b[n+15]) + p.GrandmasterClockQuality.OffsetScaledLogVariance = binary.BigEndian.Uint16(b[n+16:]) + p.GrandmasterPriority2 = b[n+18] + p.GrandmasterIdentity = ClockIdentity(binary.BigEndian.Uint64(b[n+19:])) + p.StepsRemoved = binary.BigEndian.Uint16(b[n+27:]) + p.TimeSource = TimeSource(b[n+29]) + pos := n + 30 + // unmarshal TLVs if present + var err error + p.TLVs, err = readTLVs(p.TLVs, int(p.MessageLength)-pos, b[pos:]) + if err != nil { + return err + } + return nil +} + +// MarshalBinary converts packet to []bytes +func (p *Announce) MarshalBinary() ([]byte, error) { + buf := make([]byte, 508) + n, err := p.MarshalBinaryTo(buf) + return buf[:n], err +} + +// SyncDelayReqBody Table 44 Sync and Delay_Req message fields +type SyncDelayReqBody struct { + OriginTimestamp Timestamp +} + +// SyncDelayReq is a full Sync/Delay_Req packet +type SyncDelayReq struct { + Header + SyncDelayReqBody + TLVs []TLV +} + +// MarshalBinaryTo marshals bytes to SyncDelayReq +func (p *SyncDelayReq) MarshalBinaryTo(b []byte) (int, error) { + if len(b) < headerSize+10 { + return 0, fmt.Errorf("not enough buffer to write SyncDelayReq") + } + n := headerMarshalBinaryTo(&p.Header, b) + copy(b[n:], p.OriginTimestamp.Seconds[:]) //uint48 + binary.BigEndian.PutUint32(b[n+6:], p.OriginTimestamp.Nanoseconds) + pos := n + 10 + tlvLen, err := writeTLVs(p.TLVs, b[pos:]) + return pos + tlvLen, err +} + +// MarshalBinary converts packet to []bytes +func (p *SyncDelayReq) MarshalBinary() ([]byte, error) { + buf := make([]byte, 50) + n, err := p.MarshalBinaryTo(buf) + return buf[:n], err +} + +// UnmarshalBinary unmarshals bytes to SyncDelayReq +func (p *SyncDelayReq) UnmarshalBinary(b []byte) error { + if len(b) < headerSize+10 { + return fmt.Errorf("not enough data to decode SyncDelayReq") + } + unmarshalHeader(&p.Header, b) + if err := checkPacketLength(&p.Header, len(b)); err != nil { + return err + } + copy(p.OriginTimestamp.Seconds[:], b[headerSize:]) //uint48 + p.OriginTimestamp.Nanoseconds = binary.BigEndian.Uint32(b[headerSize+6:]) + + pos := headerSize + 10 + var err error + p.TLVs, err = readTLVs(p.TLVs, int(p.MessageLength)-pos, b[pos:]) + return err +} + +// FollowUpBody Table 45 Follow_Up message fields +type FollowUpBody struct { + PreciseOriginTimestamp Timestamp +} + +// FollowUp is a full Follow_Up packet +type FollowUp struct { + Header + FollowUpBody +} + +// MarshalBinaryTo marshals bytes to FollowUp +func (p *FollowUp) MarshalBinaryTo(b []byte) (int, error) { + if len(b) < headerSize+10 { + return 0, fmt.Errorf("not enough buffer to write FollowUp") + } + n := headerMarshalBinaryTo(&p.Header, b) + copy(b[n:], p.PreciseOriginTimestamp.Seconds[:]) //uint48 + binary.BigEndian.PutUint32(b[n+6:], p.PreciseOriginTimestamp.Nanoseconds) + return n + 10, nil +} + +// MarshalBinary converts packet to []bytes +func (p *FollowUp) MarshalBinary() ([]byte, error) { + buf := make([]byte, 44) + n, err := p.MarshalBinaryTo(buf) + return buf[:n], err +} + +// UnmarshalBinary unmarshals bytes to FollowUp +func (p *FollowUp) UnmarshalBinary(b []byte) error { + if len(b) < headerSize+10 { + return fmt.Errorf("not enough data to decode FollowUp") + } + unmarshalHeader(&p.Header, b) + if err := checkPacketLength(&p.Header, len(b)); err != nil { + return err + } + copy(p.PreciseOriginTimestamp.Seconds[:], b[headerSize:]) //uint48 + p.PreciseOriginTimestamp.Nanoseconds = binary.BigEndian.Uint32(b[headerSize+6:]) + return nil +} + +// DelayRespBody Table 46 Delay_Resp message fields +type DelayRespBody struct { + ReceiveTimestamp Timestamp + RequestingPortIdentity PortIdentity +} + +// DelayResp is a full Delay_Resp packet +type DelayResp struct { + Header + DelayRespBody +} + +// MarshalBinaryTo marshals bytes to DelayResp +func (p *DelayResp) MarshalBinaryTo(b []byte) (int, error) { + if len(b) < headerSize+20 { + return 0, fmt.Errorf("not enough buffer to write DelayResp") + } + n := headerMarshalBinaryTo(&p.Header, b) + copy(b[n:], p.ReceiveTimestamp.Seconds[:]) //uint48 + binary.BigEndian.PutUint32(b[n+6:], p.ReceiveTimestamp.Nanoseconds) + binary.BigEndian.PutUint64(b[n+10:], uint64(p.RequestingPortIdentity.ClockIdentity)) + binary.BigEndian.PutUint16(b[n+18:], p.RequestingPortIdentity.PortNumber) + return n + 20, nil +} + +// MarshalBinary converts packet to []bytes +func (p *DelayResp) MarshalBinary() ([]byte, error) { + buf := make([]byte, 54) + n, err := p.MarshalBinaryTo(buf) + return buf[:n], err +} + +// UnmarshalBinary unmarshals bytes to DelayResp +func (p *DelayResp) UnmarshalBinary(b []byte) error { + if len(b) < headerSize+20 { + return fmt.Errorf("not enough data to decode DelayResp") + } + unmarshalHeader(&p.Header, b) + if err := checkPacketLength(&p.Header, len(b)); err != nil { + return err + } + copy(p.ReceiveTimestamp.Seconds[:], b[headerSize:]) //uint48 + p.ReceiveTimestamp.Nanoseconds = binary.BigEndian.Uint32(b[headerSize+6:]) + p.RequestingPortIdentity.ClockIdentity = ClockIdentity(binary.BigEndian.Uint64(b[headerSize+10:])) + p.RequestingPortIdentity.PortNumber = binary.BigEndian.Uint16(b[headerSize+18:]) + return nil +} + +// PDelayReqBody Table 47 Pdelay_Req message fields +type PDelayReqBody struct { + OriginTimestamp Timestamp + Reserved [10]uint8 +} + +// PDelayReq is a full Pdelay_Req packet +type PDelayReq struct { + Header + PDelayReqBody +} + +// PDelayRespBody Table 48 Pdelay_Resp message fields +type PDelayRespBody struct { + RequestReceiptTimestamp Timestamp + RequestingPortIdentity PortIdentity +} + +// PDelayResp is a full Pdelay_Resp packet +type PDelayResp struct { + Header + PDelayRespBody +} + +// PDelayRespFollowUpBody Table 49 Pdelay_Resp_Follow_Up message fields +type PDelayRespFollowUpBody struct { + ResponseOriginTimestamp Timestamp + RequestingPortIdentity PortIdentity +} + +// PDelayRespFollowUp is a full Pdelay_Resp_Follow_Up packet +type PDelayRespFollowUp struct { + Header + PDelayRespFollowUpBody +} + +// Packet is an interface to abstract all different packets +type Packet interface { + MessageType() MessageType + SetSequence(uint16) +} + +// BinaryMarshalerTo is an interface implemented by an object that can marshal itself into a binary form into provided []byte +type BinaryMarshalerTo interface { + MarshalBinaryTo([]byte) (int, error) +} + +// BytesTo marshalls packets that support this optimized marshalling into []byte +func BytesTo(p BinaryMarshalerTo, buf []byte) (int, error) { + n, err := p.MarshalBinaryTo(buf) + if err != nil { + return 0, err + } + // add two zero bytes + buf[n] = 0x0 + buf[n+1] = 0x0 + return n + 2, nil +} + +// Bytes converts any packet to []bytes +func Bytes(p Packet) ([]byte, error) { + // interface smuggling + if pp, ok := p.(encoding.BinaryMarshaler); ok { + b, err := pp.MarshalBinary() + return append(b, twoZeros...), err + } + var bytes bytes.Buffer + err := binary.Write(&bytes, binary.BigEndian, p) + if err != nil { + return nil, err + } + err = binary.Write(&bytes, binary.BigEndian, twoZeros) + return bytes.Bytes(), err +} + +// FromBytes parses []byte into any packet +func FromBytes(rawBytes []byte, p Packet) error { + // interface smuggling + if pp, ok := p.(encoding.BinaryUnmarshaler); ok { + return pp.UnmarshalBinary(rawBytes) + } + reader := bytes.NewReader(rawBytes) + return binary.Read(reader, binary.BigEndian, p) +} + +// DecodePacket provides single entry point to try and decode any []bytes to PTPv2 packet. +// It can be used for easy integration with anything that provides UDP packet payload as bytes. +// Resulting Packet user can then either switch based on MessageType(), or just with type switch. +func DecodePacket(b []byte) (Packet, error) { + r := bytes.NewReader(b) + head := &Header{} + if err := binary.Read(r, binary.BigEndian, head); err != nil { + return nil, err + } + msgType := head.MessageType() + var p Packet + switch msgType { + case MessageSync, MessageDelayReq: + p = &SyncDelayReq{} + case MessagePDelayReq: + p = &PDelayReq{} + case MessagePDelayResp: + p = &PDelayResp{} + case MessageFollowUp: + p = &FollowUp{} + case MessageDelayResp: + p = &DelayResp{} + case MessagePDelayRespFollowUp: + p = &PDelayRespFollowUp{} + case MessageAnnounce: + p = &Announce{} + case MessageSignaling: + p = &Signaling{} + case MessageManagement: + return decodeMgmtPacket(b) + default: + return nil, fmt.Errorf("unsupported type %s", msgType) + } + + if err := FromBytes(b, p); err != nil { + return nil, err + } + return p, nil +} diff --git a/src/yangerd/vendor/github.com/facebook/time/ptp/protocol/ptp4l.go b/src/yangerd/vendor/github.com/facebook/time/ptp/protocol/ptp4l.go new file mode 100644 index 000000000..0cfbbb93a --- /dev/null +++ b/src/yangerd/vendor/github.com/facebook/time/ptp/protocol/ptp4l.go @@ -0,0 +1,545 @@ +/* +Copyright (c) Facebook, Inc. and its affiliates. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package protocol + +// Support has been included for some non-standard extensions provided by the ptp4l implementation; the TLVs IDPortStatsNP and IDTimeStatusNP +// Implemented as present in linuxptp master d95f4cd6e4a7c6c51a220c58903110a2326885e7 + +import ( + "bytes" + "encoding/binary" + "fmt" + "net" + + "github.com/facebook/time/hostendian" +) + +// PTP4lSock is the default path to PTP4L socket +const PTP4lSock = "/var/run/ptp4l" + +// ptp4l-specific management TLV ids +const ( + IDTimeStatusNP ManagementID = 0xC000 + IDPortPropertiesNP ManagementID = 0xC004 + IDPortStatsNP ManagementID = 0xC005 + IDPortServiceStatsNP ManagementID = 0xC007 + IDUnicastMasterTableNP ManagementID = 0xC008 +) + +// UnicastMasterState is a enum describing the unicast master state in ptp4l unicast master table +type UnicastMasterState uint8 + +// possible states of unicast master in ptp4l unicast master table +const ( + UnicastMasterStateWait UnicastMasterState = iota + UnicastMasterStateHaveAnnounce + UnicastMasterStateNeedSYDY + UnicastMasterStateHaveSYDY +) + +// UnicastMasterStateToString is a map from UnicastMasterState to string +var UnicastMasterStateToString = map[UnicastMasterState]string{ + UnicastMasterStateWait: "WAIT", + UnicastMasterStateHaveAnnounce: "HAVE_ANN", + UnicastMasterStateNeedSYDY: "NEED_SYDY", + UnicastMasterStateHaveSYDY: "HAVE_SYDY", +} + +func (t UnicastMasterState) String() string { + return UnicastMasterStateToString[t] +} + +// Timestamping is a ptp4l-specific enum describing timestamping type +type Timestamping uint8 + +const ( + // TimestampingSoftware is a software timestamp const + TimestampingSoftware Timestamping = iota + // TimestampingHardware is a hardware timestamp const + TimestampingHardware + // TimestampingLegacyHW is a legacy hardware timestamp const + TimestampingLegacyHW + // TimestampingOneStep is a one step timestamp const + TimestampingOneStep + // TimestampingP2P1Step is a P2P one step timestamp const + TimestampingP2P1Step +) + +// PortStats is a ptp4l struct containing port statistics +type PortStats struct { + RXMsgType [16]uint64 + TXMsgType [16]uint64 +} + +// PortStatsNPTLV is a ptp4l struct containing port identinity and statistics +type PortStatsNPTLV struct { + ManagementTLVHead + + PortIdentity PortIdentity + PortStats PortStats +} + +// MarshalBinary converts packet to []bytes +func (p *PortStatsNPTLV) MarshalBinary() ([]byte, error) { + var bytes bytes.Buffer + if err := binary.Write(&bytes, binary.BigEndian, p.ManagementTLVHead); err != nil { + return nil, err + } + if err := binary.Write(&bytes, binary.BigEndian, p.PortIdentity); err != nil { + return nil, err + } + if err := binary.Write(&bytes, hostendian.Order, p.PortStats); err != nil { + return nil, err + } + return bytes.Bytes(), nil +} + +// ScaledNS is some struct used by ptp4l to report phase change +type ScaledNS struct { + NanosecondsMSB uint16 + NanosecondsLSB uint64 + FractionalNanoseconds uint16 +} + +// TimeStatusNPTLV is a ptp4l struct containing actually useful instance metrics +type TimeStatusNPTLV struct { + ManagementTLVHead + + MasterOffsetNS int64 + IngressTimeNS int64 // this is PHC time + CumulativeScaledRateOffset int32 + ScaledLastGmPhaseChange int32 + GMTimeBaseIndicator uint16 + LastGmPhaseChange ScaledNS + GMPresent int32 + GMIdentity ClockIdentity +} + +// PortPropertiesNPTLV is a ptp4l struct containing port properties +type PortPropertiesNPTLV struct { + ManagementTLVHead + + PortIdentity PortIdentity + PortState PortState + Timestamping Timestamping + Interface PTPText +} + +// MarshalBinary converts packet to []bytes +func (p *PortPropertiesNPTLV) MarshalBinary() ([]byte, error) { + var bytes bytes.Buffer + if err := binary.Write(&bytes, binary.BigEndian, p.ManagementTLVHead); err != nil { + return nil, err + } + if err := binary.Write(&bytes, binary.BigEndian, p.PortIdentity); err != nil { + return nil, err + } + if err := binary.Write(&bytes, hostendian.Order, p.PortState); err != nil { + return nil, err + } + if err := binary.Write(&bytes, hostendian.Order, p.Timestamping); err != nil { + return nil, err + } + interfaceBytes, err := p.Interface.MarshalBinary() + if err != nil { + return nil, err + } + if err := binary.Write(&bytes, hostendian.Order, interfaceBytes); err != nil { + return nil, err + } + return bytes.Bytes(), nil +} + +// PortServiceStats is a ptp4l struct containing counters for different port events, which we added in linuxptp cfbb8bdb50f5a38687fcddccbe6a264c6a078bbd +type PortServiceStats struct { + AnnounceTimeout uint64 `json:"ptp.servicestats.announce_timeout"` + SyncTimeout uint64 `json:"ptp.servicestats.sync_timeout"` + DelayTimeout uint64 `json:"ptp.servicestats.delay_timeout"` + UnicastServiceTimeout uint64 `json:"ptp.servicestats.unicast_service_timeout"` + UnicastRequestTimeout uint64 `json:"ptp.servicestats.unicast_request_timeout"` + MasterAnnounceTimeout uint64 `json:"ptp.servicestats.master_announce_timeout"` + MasterSyncTimeout uint64 `json:"ptp.servicestats.master_sync_timeout"` + QualificationTimeout uint64 `json:"ptp.servicestats.qualification_timeout"` + SyncMismatch uint64 `json:"ptp.servicestats.sync_mismatch"` + FollowupMismatch uint64 `json:"ptp.servicestats.followup_mismatch"` +} + +// PortServiceStatsNPTLV is a management TLV added in linuxptp cfbb8bdb50f5a38687fcddccbe6a264c6a078bbd +type PortServiceStatsNPTLV struct { + ManagementTLVHead + + PortIdentity PortIdentity + PortServiceStats PortServiceStats +} + +// MarshalBinary converts packet to []bytes +func (p *PortServiceStatsNPTLV) MarshalBinary() ([]byte, error) { + var bytes bytes.Buffer + if err := binary.Write(&bytes, binary.BigEndian, p.ManagementTLVHead); err != nil { + return nil, err + } + if err := binary.Write(&bytes, binary.BigEndian, p.PortIdentity); err != nil { + return nil, err + } + if err := binary.Write(&bytes, hostendian.Order, p.PortServiceStats); err != nil { + return nil, err + } + return bytes.Bytes(), nil +} + +// UnicastMasterEntry is an entry in UnicastMasterTable that ptp4l exports via management TLV +type UnicastMasterEntry struct { + PortIdentity PortIdentity + ClockQuality ClockQuality + Selected bool + PortState UnicastMasterState + Priority1 uint8 + Priority2 uint8 + Address net.IP +} + +// UnmarshalBinary implements Unmarshaller interface +func (e *UnicastMasterEntry) UnmarshalBinary(b []byte) error { + var err error + if len(b) < 26 { // 22 byte for struct, at least 4 for address) + return fmt.Errorf("not enough data to decode UnicastMasterEntry") + } + e.PortIdentity.ClockIdentity = ClockIdentity(binary.BigEndian.Uint64(b[0:])) + e.PortIdentity.PortNumber = binary.BigEndian.Uint16(b[8:]) + e.ClockQuality.ClockClass = ClockClass(b[10]) + e.ClockQuality.ClockAccuracy = ClockAccuracy(b[11]) + e.ClockQuality.OffsetScaledLogVariance = binary.BigEndian.Uint16(b[12:]) + if b[14] == 0 { + e.Selected = false + } else if b[14] == 1 { + e.Selected = true + } else { + return fmt.Errorf("unexpected 'selected' value %d", b[14]) + } + e.PortState = UnicastMasterState(b[15]) + e.Priority1 = b[16] + e.Priority2 = b[17] + + pa := &PortAddress{} + if err := pa.UnmarshalBinary(b[18:]); err != nil { + return err + } + e.Address, err = pa.IP() + if err != nil { + return err + } + return nil +} + +// MarshalBinary converts UnicastMasterEntry to []bytes +func (e *UnicastMasterEntry) MarshalBinary() ([]byte, error) { + var bytes bytes.Buffer + if err := binary.Write(&bytes, binary.BigEndian, e.PortIdentity); err != nil { + return nil, err + } + if err := binary.Write(&bytes, binary.BigEndian, e.ClockQuality); err != nil { + return nil, err + } + var selectedBin uint8 + if e.Selected { + selectedBin = 1 + } + if err := binary.Write(&bytes, binary.BigEndian, selectedBin); err != nil { + return nil, err + } + if err := binary.Write(&bytes, binary.BigEndian, e.PortState); err != nil { + return nil, err + } + if err := binary.Write(&bytes, binary.BigEndian, e.Priority1); err != nil { + return nil, err + } + if err := binary.Write(&bytes, binary.BigEndian, e.Priority2); err != nil { + return nil, err + } + var pa PortAddress + asIPv4 := e.Address.To4() + if asIPv4 != nil { + pa = PortAddress{ + NetworkProtocol: TransportTypeUDPIPV4, + AddressLength: 4, + AddressField: asIPv4, + } + } else { + pa = PortAddress{ + NetworkProtocol: TransportTypeUDPIPV6, + AddressLength: 16, + AddressField: e.Address, + } + } + portBytes, err := pa.MarshalBinary() + if err != nil { + return nil, err + } + if err := binary.Write(&bytes, binary.BigEndian, portBytes); err != nil { + return nil, err + } + return bytes.Bytes(), nil +} + +// UnicastMasterTable is a table of UnicastMasterEntries +type UnicastMasterTable struct { + ActualTableSize uint16 + UnicastMasters []UnicastMasterEntry +} + +// UnicastMasterTableNPTLV is a custom management packet that exports unicast master table state +type UnicastMasterTableNPTLV struct { + ManagementTLVHead + + UnicastMasterTable UnicastMasterTable +} + +// MarshalBinary converts packet to []bytes +func (p *UnicastMasterTableNPTLV) MarshalBinary() ([]byte, error) { + var bytes bytes.Buffer + if err := binary.Write(&bytes, binary.BigEndian, p.ManagementTLVHead); err != nil { + return nil, err + } + if err := binary.Write(&bytes, binary.BigEndian, p.UnicastMasterTable.ActualTableSize); err != nil { + return nil, err + } + for _, e := range p.UnicastMasterTable.UnicastMasters { + entryBytes, err := e.MarshalBinary() + if err != nil { + return nil, err + } + if err := binary.Write(&bytes, binary.BigEndian, entryBytes); err != nil { + return nil, err + } + } + return bytes.Bytes(), nil +} + +// PortStatsNPRequest prepares request packet for PORT_STATS_NP request +func PortStatsNPRequest() *Management { + headerSize := uint16(binary.Size(ManagementMsgHead{})) + tlvHeadSize := uint16(binary.Size(TLVHead{})) + // we send request with no portStats data just like pmc does + return &Management{ + ManagementMsgHead: ManagementMsgHead{ + Header: Header{ + SdoIDAndMsgType: NewSdoIDAndMsgType(MessageManagement, 0), + Version: Version, + MessageLength: headerSize + tlvHeadSize + 2, + SourcePortIdentity: identity, + LogMessageInterval: MgmtLogMessageInterval, + }, + TargetPortIdentity: DefaultTargetPortIdentity, + StartingBoundaryHops: 0, + BoundaryHops: 0, + ActionField: GET, + }, + TLV: &ManagementTLVHead{ + TLVHead: TLVHead{ + TLVType: TLVManagement, + LengthField: 2, + }, + ManagementID: IDPortStatsNP, + }, + } +} + +// PortStatsNP sends PORT_STATS_NP request and returns response +func (c *MgmtClient) PortStatsNP() (*PortStatsNPTLV, error) { + req := PortStatsNPRequest() + p, err := c.Communicate(req) + if err != nil { + return nil, err + } + tlv, ok := p.TLV.(*PortStatsNPTLV) + if !ok { + return nil, fmt.Errorf("got unexpected management TLV %T, wanted %T", p.TLV, tlv) + } + return tlv, nil +} + +// TimeStatusNPRequest prepares request packet for TIME_STATUS_NP request +func TimeStatusNPRequest() *Management { + headerSize := uint16(binary.Size(ManagementMsgHead{})) + tlvHeadSize := uint16(binary.Size(TLVHead{})) + // we send request with no TimeStatusNP data just like pmc does + return &Management{ + ManagementMsgHead: ManagementMsgHead{ + Header: Header{ + SdoIDAndMsgType: NewSdoIDAndMsgType(MessageManagement, 0), + Version: Version, + MessageLength: headerSize + tlvHeadSize + 2, + SourcePortIdentity: identity, + LogMessageInterval: MgmtLogMessageInterval, + }, + TargetPortIdentity: DefaultTargetPortIdentity, + StartingBoundaryHops: 0, + BoundaryHops: 0, + ActionField: GET, + }, + TLV: &ManagementTLVHead{ + TLVHead: TLVHead{ + TLVType: TLVManagement, + LengthField: 2, + }, + ManagementID: IDTimeStatusNP, + }, + } +} + +// TimeStatusNP sends TIME_STATUS_NP request and returns response +func (c *MgmtClient) TimeStatusNP() (*TimeStatusNPTLV, error) { + req := TimeStatusNPRequest() + p, err := c.Communicate(req) + if err != nil { + return nil, err + } + tlv, ok := p.TLV.(*TimeStatusNPTLV) + if !ok { + return nil, fmt.Errorf("got unexpected management TLV %T, wanted %T", p.TLV, tlv) + } + return tlv, nil +} + +// PortServiceStatsNPRequest prepares request packet for PORT_SERVICE_STATS_NP request +func PortServiceStatsNPRequest() *Management { + headerSize := uint16(binary.Size(ManagementMsgHead{})) + tlvHeadSize := uint16(binary.Size(TLVHead{})) + // we send request with no portServiceStats data just like pmc does + return &Management{ + ManagementMsgHead: ManagementMsgHead{ + Header: Header{ + SdoIDAndMsgType: NewSdoIDAndMsgType(MessageManagement, 0), + Version: Version, + MessageLength: headerSize + tlvHeadSize + 2, + SourcePortIdentity: identity, + LogMessageInterval: MgmtLogMessageInterval, + }, + TargetPortIdentity: DefaultTargetPortIdentity, + StartingBoundaryHops: 0, + BoundaryHops: 0, + ActionField: GET, + }, + TLV: &ManagementTLVHead{ + TLVHead: TLVHead{ + TLVType: TLVManagement, + LengthField: 2, + }, + ManagementID: IDPortServiceStatsNP, + }, + } +} + +// PortServiceStatsNP sends PORT_SERVICE_STATS_NP request and returns response +func (c *MgmtClient) PortServiceStatsNP() (*PortServiceStatsNPTLV, error) { + req := PortServiceStatsNPRequest() + p, err := c.Communicate(req) + if err != nil { + return nil, err + } + tlv, ok := p.TLV.(*PortServiceStatsNPTLV) + if !ok { + return nil, fmt.Errorf("got unexpected management TLV %T, wanted %T", p.TLV, tlv) + } + return tlv, nil +} + +// PortPropertiesNPRequest prepares request packet for PORT_STATS_NP request +func PortPropertiesNPRequest() *Management { + headerSize := uint16(binary.Size(ManagementMsgHead{})) + tlvHeadSize := uint16(binary.Size(TLVHead{})) + // we send request with no portStats data just like pmc does + return &Management{ + ManagementMsgHead: ManagementMsgHead{ + Header: Header{ + SdoIDAndMsgType: NewSdoIDAndMsgType(MessageManagement, 0), + Version: Version, + MessageLength: headerSize + tlvHeadSize + 2, + SourcePortIdentity: identity, + LogMessageInterval: MgmtLogMessageInterval, + }, + TargetPortIdentity: DefaultTargetPortIdentity, + StartingBoundaryHops: 0, + BoundaryHops: 0, + ActionField: GET, + }, + TLV: &ManagementTLVHead{ + TLVHead: TLVHead{ + TLVType: TLVManagement, + LengthField: 2, + }, + ManagementID: IDPortPropertiesNP, + }, + } +} + +// PortPropertiesNP sends PORT_PROPERTIES_NP request and returns response +func (c *MgmtClient) PortPropertiesNP() (*PortPropertiesNPTLV, error) { + req := PortPropertiesNPRequest() + p, err := c.Communicate(req) + if err != nil { + return nil, err + } + tlv, ok := p.TLV.(*PortPropertiesNPTLV) + if !ok { + return nil, fmt.Errorf("got unexpected management TLV %T, wanted %T", p.TLV, tlv) + } + return tlv, nil +} + +// UnicastMasterTableNPRequest creates new packet with UNICAST_MASTER_TABLE_NP request +func UnicastMasterTableNPRequest() *Management { + headerSize := uint16(binary.Size(ManagementMsgHead{})) + tlvHeadSize := uint16(binary.Size(TLVHead{})) + // we send request with no data just like pmc does + return &Management{ + ManagementMsgHead: ManagementMsgHead{ + Header: Header{ + SdoIDAndMsgType: NewSdoIDAndMsgType(MessageManagement, 0), + Version: Version, + MessageLength: headerSize + tlvHeadSize + 2, + SourcePortIdentity: identity, + LogMessageInterval: MgmtLogMessageInterval, + }, + TargetPortIdentity: DefaultTargetPortIdentity, + StartingBoundaryHops: 0, + BoundaryHops: 0, + ActionField: GET, + }, + TLV: &ManagementTLVHead{ + TLVHead: TLVHead{ + TLVType: TLVManagement, + LengthField: 2, + }, + ManagementID: IDUnicastMasterTableNP, + }, + } +} + +// UnicastMasterTableNP request UNICAST_MASTER_TABLE_NP from ptp4l, and returns the result +func (c *MgmtClient) UnicastMasterTableNP() (*UnicastMasterTableNPTLV, error) { + req := UnicastMasterTableNPRequest() + p, err := c.Communicate(req) + if err != nil { + return nil, err + } + tlv, ok := p.TLV.(*UnicastMasterTableNPTLV) + if !ok { + return nil, fmt.Errorf("got unexpected management TLV %T, wanted %T", p.TLV, tlv) + } + return tlv, nil +} diff --git a/src/yangerd/vendor/github.com/facebook/time/ptp/protocol/tlvs.go b/src/yangerd/vendor/github.com/facebook/time/ptp/protocol/tlvs.go new file mode 100644 index 000000000..611a7da43 --- /dev/null +++ b/src/yangerd/vendor/github.com/facebook/time/ptp/protocol/tlvs.go @@ -0,0 +1,405 @@ +/* +Copyright (c) Facebook, Inc. and its affiliates. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package protocol + +import ( + "bytes" + "encoding/binary" + "fmt" +) + +// TLV abstracts away any TLV +type TLV interface { + Type() TLVType +} + +const tlvHeadSize = 4 + +// TLVHead is a common part of all TLVs +type TLVHead struct { + TLVType TLVType + LengthField uint16 // The length of all TLVs shall be an even number of octets +} + +// Type implements TLV interface +func (t TLVHead) Type() TLVType { + return t.TLVType +} + +func tlvHeadMarshalBinaryTo(t *TLVHead, b []byte) { + binary.BigEndian.PutUint16(b, uint16(t.TLVType)) + binary.BigEndian.PutUint16(b[2:], t.LengthField) +} + +func unmarshalTLVHeader(p *TLVHead, b []byte) error { + if len(b) < tlvHeadSize { + return fmt.Errorf("not enough data to decode PTP header") + } + p.TLVType = TLVType(binary.BigEndian.Uint16(b[0:])) + p.LengthField = binary.BigEndian.Uint16(b[2:]) + return nil +} + +func checkTLVLength(p *TLVHead, l, want int, strict bool) error { + if strict && int(p.LengthField) != want { + return fmt.Errorf("expected TLV of type %s (%d) to have length of %d, got %d in the header", p.TLVType, p.TLVType, want, p.LengthField) + } + + if int(p.LengthField) < want { + return fmt.Errorf("expected TLV of type %s (%d) to have length of at least %d, got %d in the header", p.TLVType, p.TLVType, want, p.LengthField) + } + if tlvHeadSize+int(p.LengthField) > l { + return fmt.Errorf("cannot decode TLV of length %d from %d bytes", tlvHeadSize+int(p.LengthField), l) + } + return nil +} + +func writeTLVs(tlvs []TLV, b []byte) (int, error) { + pos := 0 + for _, tlv := range tlvs { + if ttlv, ok := tlv.(BinaryMarshalerTo); ok { + nn, err := ttlv.MarshalBinaryTo(b[pos:]) + if err != nil { + return 0, err + } + pos += nn + continue + } + // very inefficient path for TLVs that don't support MarshalBinaryTo + buf := new(bytes.Buffer) + if err := binary.Write(buf, binary.BigEndian, tlv); err != nil { + return 0, err + } + bbytes := buf.Bytes() + copy(b[pos:], bbytes) + pos += len(bbytes) + } + return pos, nil +} + +// readTLVs reads TLVs from the bytes. +// tlvs is passed to save on allocations and it's user's task to ensure it's empty +func readTLVs(tlvs []TLV, maxLength int, b []byte) ([]TLV, error) { + pos := 0 + var tlvType TLVType + for { + // packet can have trailing bytes, let's make sure we don't try to read past given length + if pos+tlvHeadSize > maxLength { + break + } + tlvType = TLVType(binary.BigEndian.Uint16(b[pos:])) + + switch tlvType { + case TLVAcknowledgeCancelUnicastTransmission: + tlv := &AcknowledgeCancelUnicastTransmissionTLV{} + if err := tlv.UnmarshalBinary(b[pos:]); err != nil { + return tlvs, err + } + tlvs = append(tlvs, tlv) + pos += tlvHeadSize + int(tlv.LengthField) + case TLVGrantUnicastTransmission: + tlv := &GrantUnicastTransmissionTLV{} + if err := tlv.UnmarshalBinary(b[pos:]); err != nil { + return tlvs, err + } + tlvs = append(tlvs, tlv) + pos += tlvHeadSize + int(tlv.LengthField) + case TLVRequestUnicastTransmission: + tlv := &RequestUnicastTransmissionTLV{} + if err := tlv.UnmarshalBinary(b[pos:]); err != nil { + return tlvs, err + } + tlvs = append(tlvs, tlv) + pos += tlvHeadSize + int(tlv.LengthField) + case TLVCancelUnicastTransmission: + tlv := &CancelUnicastTransmissionTLV{} + if err := tlv.UnmarshalBinary(b[pos:]); err != nil { + return tlvs, err + } + tlvs = append(tlvs, tlv) + pos += tlvHeadSize + int(tlv.LengthField) + case TLVPathTrace: + tlv := &PathTraceTLV{} + if err := tlv.UnmarshalBinary(b[pos:]); err != nil { + return tlvs, err + } + tlvs = append(tlvs, tlv) + pos += tlvHeadSize + int(tlv.LengthField) + case TLVAlternateTimeOffsetIndicator: + tlv := &AlternateTimeOffsetIndicatorTLV{} + if err := tlv.UnmarshalBinary(b[pos:]); err != nil { + return tlvs, err + } + tlvs = append(tlvs, tlv) + pos += tlvHeadSize + int(tlv.LengthField) + case TLVAlternateResponsePort: + tlv := &AlternateResponsePortTLV{} + if err := tlv.UnmarshalBinary(b[pos:]); err != nil { + return tlvs, err + } + tlvs = append(tlvs, tlv) + pos += tlvHeadSize + int(tlv.LengthField) + default: + return tlvs, fmt.Errorf("reading TLV %s (%d) is not yet implemented", tlvType, tlvType) + } + } + return tlvs, nil +} + +// Unicast TLVs + +// RequestUnicastTransmissionTLV Table 110 REQUEST_UNICAST_TRANSMISSION TLV format +type RequestUnicastTransmissionTLV struct { + TLVHead + MsgTypeAndReserved UnicastMsgTypeAndFlags // first 4 bits only, same enums as with normal message type + LogInterMessagePeriod LogInterval + DurationField uint32 +} + +// MarshalBinaryTo marshals bytes to RequestUnicastTransmissionTLV +func (t *RequestUnicastTransmissionTLV) MarshalBinaryTo(b []byte) (int, error) { + tlvHeadMarshalBinaryTo(&t.TLVHead, b) + b[tlvHeadSize] = byte(t.MsgTypeAndReserved) + b[tlvHeadSize+1] = byte(t.LogInterMessagePeriod) + binary.BigEndian.PutUint32(b[tlvHeadSize+2:], t.DurationField) + return tlvHeadSize + 6, nil +} + +// UnmarshalBinary parses []byte and populates struct fields +func (t *RequestUnicastTransmissionTLV) UnmarshalBinary(b []byte) error { + if err := unmarshalTLVHeader(&t.TLVHead, b); err != nil { + return err + } + if err := checkTLVLength(&t.TLVHead, len(b), 6, true); err != nil { + return err + } + t.MsgTypeAndReserved = UnicastMsgTypeAndFlags(b[4]) + t.LogInterMessagePeriod = LogInterval(b[5]) + t.DurationField = binary.BigEndian.Uint32(b[6:]) + return nil +} + +// GrantUnicastTransmissionTLV Table 111 GRANT_UNICAST_TRANSMISSION TLV format +type GrantUnicastTransmissionTLV struct { + TLVHead + MsgTypeAndReserved UnicastMsgTypeAndFlags // first 4 bits only, same enums as with normal message type + LogInterMessagePeriod LogInterval + DurationField uint32 + Reserved uint8 + Renewal uint8 +} + +// MarshalBinaryTo marshals bytes to GrantUnicastTransmissionTLV +func (t *GrantUnicastTransmissionTLV) MarshalBinaryTo(b []byte) (int, error) { + tlvHeadMarshalBinaryTo(&t.TLVHead, b) + b[tlvHeadSize] = byte(t.MsgTypeAndReserved) + b[tlvHeadSize+1] = byte(t.LogInterMessagePeriod) + binary.BigEndian.PutUint32(b[tlvHeadSize+2:], t.DurationField) + b[tlvHeadSize+6] = t.Reserved + b[tlvHeadSize+7] = t.Renewal + return tlvHeadSize + 8, nil +} + +// UnmarshalBinary parses []byte and populates struct fields +func (t *GrantUnicastTransmissionTLV) UnmarshalBinary(b []byte) error { + if err := unmarshalTLVHeader(&t.TLVHead, b); err != nil { + return err + } + if err := checkTLVLength(&t.TLVHead, len(b), 8, true); err != nil { + return err + } + t.MsgTypeAndReserved = UnicastMsgTypeAndFlags(b[4]) + t.LogInterMessagePeriod = LogInterval(b[5]) + t.DurationField = binary.BigEndian.Uint32(b[6:]) + t.Reserved = b[10] + t.Renewal = b[11] + return nil +} + +// CancelUnicastTransmissionTLV Table 112 CANCEL_UNICAST_TRANSMISSION TLV format +type CancelUnicastTransmissionTLV struct { + TLVHead + MsgTypeAndFlags UnicastMsgTypeAndFlags // first 4 bits is msg type, then flags R and/or G + Reserved uint8 +} + +// MarshalBinaryTo marshals bytes to CancelUnicastTransmissionTLV +func (t *CancelUnicastTransmissionTLV) MarshalBinaryTo(b []byte) (int, error) { + tlvHeadMarshalBinaryTo(&t.TLVHead, b) + b[tlvHeadSize] = byte(t.MsgTypeAndFlags) + b[tlvHeadSize+1] = t.Reserved + return tlvHeadSize + 2, nil +} + +// UnmarshalBinary parses []byte and populates struct fields +func (t *CancelUnicastTransmissionTLV) UnmarshalBinary(b []byte) error { + if err := unmarshalTLVHeader(&t.TLVHead, b); err != nil { + return err + } + if err := checkTLVLength(&t.TLVHead, len(b), 2, true); err != nil { + return err + } + t.MsgTypeAndFlags = UnicastMsgTypeAndFlags(b[4]) + t.Reserved = b[5] + return nil +} + +// AcknowledgeCancelUnicastTransmissionTLV Table 113 ACKNOWLEDGE_CANCEL_UNICAST_TRANSMISSION TLV format +type AcknowledgeCancelUnicastTransmissionTLV struct { + TLVHead + MsgTypeAndFlags UnicastMsgTypeAndFlags // first 4 bits is msg type, then flags R and/or G + Reserved uint8 +} + +// MarshalBinaryTo marshals bytes to AcknowledgeCancelUnicastTransmissionTLV +func (t *AcknowledgeCancelUnicastTransmissionTLV) MarshalBinaryTo(b []byte) (int, error) { + tlvHeadMarshalBinaryTo(&t.TLVHead, b) + b[tlvHeadSize] = byte(t.MsgTypeAndFlags) + b[tlvHeadSize+1] = t.Reserved + return tlvHeadSize + 2, nil +} + +// UnmarshalBinary parses []byte and populates struct fields +func (t *AcknowledgeCancelUnicastTransmissionTLV) UnmarshalBinary(b []byte) error { + if err := unmarshalTLVHeader(&t.TLVHead, b); err != nil { + return err + } + if err := checkTLVLength(&t.TLVHead, len(b), 2, true); err != nil { + return err + } + t.MsgTypeAndFlags = UnicastMsgTypeAndFlags(b[4]) + t.Reserved = b[5] + return nil +} + +// other TLVs + +// PathTraceTLV Table 115 PATH_TRACE TLV format +type PathTraceTLV struct { + TLVHead + // The value of the lengthField is 8N. + PathSequence []ClockIdentity // N +} + +// MarshalBinaryTo marshals bytes to PathTraceTLV +func (t *PathTraceTLV) MarshalBinaryTo(b []byte) (int, error) { + tlvHeadMarshalBinaryTo(&t.TLVHead, b) + pos := tlvHeadSize + for _, ps := range t.PathSequence { + binary.BigEndian.PutUint64(b[pos:pos+8], uint64(ps)) + pos += 8 + } + return pos, nil +} + +// UnmarshalBinary parses []byte and populates struct fields +func (t *PathTraceTLV) UnmarshalBinary(b []byte) error { + if err := unmarshalTLVHeader(&t.TLVHead, b); err != nil { + return err + } + if err := checkTLVLength(&t.TLVHead, len(b), 8, false); err != nil { + return err + } + t.PathSequence = []ClockIdentity{} + for i := 0; i*8 <= int(t.TLVHead.LengthField); i++ { + pos := tlvHeadSize + i*8 + if pos+8 >= len(b) { + break + } + identity := ClockIdentity(binary.BigEndian.Uint64(b[pos:])) + t.PathSequence = append(t.PathSequence, identity) + } + return nil +} + +// AlternateTimeOffsetIndicatorTLV is a Table 116 ALTERNATE_TIME_OFFSET_INDICATOR TLV format +type AlternateTimeOffsetIndicatorTLV struct { + TLVHead + KeyField uint8 + CurrentOffset int32 + JumpSeconds int32 + TimeOfNextJump PTPSeconds // uint48 + DisplayName PTPText +} + +// MarshalBinaryTo marshals bytes to AlternateTimeOffsetIndicatorTLV +func (t *AlternateTimeOffsetIndicatorTLV) MarshalBinaryTo(b []byte) (int, error) { + tlvHeadMarshalBinaryTo(&t.TLVHead, b) + b[tlvHeadSize] = t.KeyField + binary.BigEndian.PutUint32(b[tlvHeadSize+1:], uint32(t.CurrentOffset)) + binary.BigEndian.PutUint32(b[tlvHeadSize+5:], uint32(t.JumpSeconds)) + copy(b[tlvHeadSize+9:], t.TimeOfNextJump[:]) //uint48 + size := tlvHeadSize + 15 + if t.DisplayName != "" { + dd, err := t.DisplayName.MarshalBinary() + if err != nil { + return 0, fmt.Errorf("writing AlternateTimeOffsetIndicatorTLV DisplayName: %w", err) + } + copy(b[tlvHeadSize+15:], dd) + size += len(dd) + } + return size, nil +} + +// UnmarshalBinary parses []byte and populates struct fields +func (t *AlternateTimeOffsetIndicatorTLV) UnmarshalBinary(b []byte) error { + if err := unmarshalTLVHeader(&t.TLVHead, b); err != nil { + return err + } + if err := checkTLVLength(&t.TLVHead, len(b), 20, false); err != nil { + return err + } + t.KeyField = b[tlvHeadSize] + t.CurrentOffset = int32(binary.BigEndian.Uint32(b[tlvHeadSize+1:])) + t.JumpSeconds = int32(binary.BigEndian.Uint32(b[tlvHeadSize+5:])) + copy(t.TimeOfNextJump[:], b[tlvHeadSize+9:]) // uint48 + if err := t.DisplayName.UnmarshalBinary(b[tlvHeadSize+15:]); err != nil { + return fmt.Errorf("reading AlternateTimeOffsetIndicatorTLV DisplayName: %w", err) + } + return nil +} + +// AlternateResponsePortTLV is a CSPTP optional TLV to switch response source port of the server +// Offset flag indicates the number of the port steps, not the port number itself. +// Ex: +// 0 means no switch (use default port). For example 1234 +// 1 means next port. For example 4567 +// 2 means next next port. For example 6789 +// etc +type AlternateResponsePortTLV struct { + TLVHead + Offset uint16 +} + +// MarshalBinaryTo marshals bytes to AlternateResponsePortTLV +func (a *AlternateResponsePortTLV) MarshalBinaryTo(b []byte) (int, error) { + tlvHeadMarshalBinaryTo(&a.TLVHead, b) + binary.BigEndian.PutUint16(b[tlvHeadSize:], a.Offset) + return tlvHeadSize + 2, nil +} + +// UnmarshalBinary parses []byte and populates struct fields +func (a *AlternateResponsePortTLV) UnmarshalBinary(b []byte) error { + if err := unmarshalTLVHeader(&a.TLVHead, b); err != nil { + return err + } + if err := checkTLVLength(&a.TLVHead, len(b), 2, true); err != nil { + return err + } + a.Offset = binary.BigEndian.Uint16(b[tlvHeadSize:]) + return nil +} diff --git a/src/yangerd/vendor/github.com/facebook/time/ptp/protocol/types.go b/src/yangerd/vendor/github.com/facebook/time/ptp/protocol/types.go new file mode 100644 index 000000000..0b2bbba0f --- /dev/null +++ b/src/yangerd/vendor/github.com/facebook/time/ptp/protocol/types.go @@ -0,0 +1,745 @@ +/* +Copyright (c) Facebook, Inc. and its affiliates. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package protocol + +import ( + "bytes" + "encoding/binary" + "fmt" + "math" + "net" + "time" +) + +// 2 ** 16 +const twoPow16 = 65536 + +// MessageType is type for Message Types +type MessageType uint8 + +// As per Table 36 Values of messageType field +const ( + MessageSync MessageType = 0x0 + MessageDelayReq MessageType = 0x1 + MessagePDelayReq MessageType = 0x2 + MessagePDelayResp MessageType = 0x3 + MessageFollowUp MessageType = 0x8 + MessageDelayResp MessageType = 0x9 + MessagePDelayRespFollowUp MessageType = 0xA + MessageAnnounce MessageType = 0xB + MessageSignaling MessageType = 0xC + MessageManagement MessageType = 0xD +) + +// MessageTypeToString is a map from MessageType to string +var MessageTypeToString = map[MessageType]string{ + MessageSync: "SYNC", + MessageDelayReq: "DELAY_REQ", + MessagePDelayReq: "PDELAY_REQ", + MessagePDelayResp: "PDELAY_RES", + MessageFollowUp: "FOLLOW_UP", + MessageDelayResp: "DELAY_RESP", + MessagePDelayRespFollowUp: "PDELAY_RESP_FOLLOW_UP", + MessageAnnounce: "ANNOUNCE", + MessageSignaling: "SIGNALING", + MessageManagement: "MANAGEMENT", +} + +func (m MessageType) String() string { + return MessageTypeToString[m] +} + +// SdoIDAndMsgType is a uint8 where first 4 bites contain SdoID and last 4 bits MessageType +type SdoIDAndMsgType uint8 + +// MsgType extracts MessageType from SdoIDAndMsgType +func (m SdoIDAndMsgType) MsgType() MessageType { + return MessageType(m & 0xf) // last 4 bits +} + +// NewSdoIDAndMsgType builds new SdoIDAndMsgType from MessageType and flags +func NewSdoIDAndMsgType(msgType MessageType, sdoID uint8) SdoIDAndMsgType { + return SdoIDAndMsgType(sdoID<<4 | uint8(msgType)) +} + +// ProbeMsgType reads first 8 bits of data and tries to decode it to SdoIDAndMsgType, then return MessageType +func ProbeMsgType(data []byte) (msg MessageType, err error) { + if len(data) < 1 { + return 0, fmt.Errorf("not enough data to probe MsgType") + } + return SdoIDAndMsgType(data[0]).MsgType(), nil +} + +// TLVType is type for TLV types +type TLVType uint16 + +// As per Table 52 tlvType values +const ( + TLVManagement TLVType = 0x0001 + TLVManagementErrorStatus TLVType = 0x0002 + TLVOrganizationExtension TLVType = 0x0003 + TLVRequestUnicastTransmission TLVType = 0x0004 + TLVGrantUnicastTransmission TLVType = 0x0005 + TLVCancelUnicastTransmission TLVType = 0x0006 + TLVAcknowledgeCancelUnicastTransmission TLVType = 0x0007 + TLVPathTrace TLVType = 0x0008 + TLVAlternateTimeOffsetIndicator TLVType = 0x0009 + TLVAlternateResponsePort TLVType = 0x2007 + // Remaining 52 tlvType TLVs not implemented +) + +// TLVTypeToString is a map from TLVType to string +var TLVTypeToString = map[TLVType]string{ + TLVManagement: "MANAGEMENT", + TLVManagementErrorStatus: "MANAGEMENT_ERROR_STATUS", + TLVOrganizationExtension: "ORGANIZATION_EXTENSION", + TLVRequestUnicastTransmission: "REQUEST_UNICAST_TRANSMISSION", + TLVGrantUnicastTransmission: "GRANT_UNICAST_TRANSMISSION", + TLVCancelUnicastTransmission: "CANCEL_UNICAST_TRANSMISSION", + TLVAcknowledgeCancelUnicastTransmission: "ACKNOWLEDGE_CANCEL_UNICAST_TRANSMISSION", + TLVPathTrace: "PATH_TRACE", + TLVAlternateTimeOffsetIndicator: "ALTERNATE_TIME_OFFSET_INDICATOR", + TLVAlternateResponsePort: "ALTERNATE_RESPONSE_PORT", +} + +func (t TLVType) String() string { + return TLVTypeToString[t] +} + +// IntFloat is a float64 stored in int64 +type IntFloat int64 + +// Value decodes IntFloat to float64 +func (t IntFloat) Value() float64 { + return float64(t) / twoPow16 +} + +/* +TimeInterval is the time interval expressed in nanoseconds, multiplied by 2**16. +Positive or negative time intervals outside the maximum range of this data type shall be encoded as the largest +positive and negative values of the data type, respectively. +For example, 2.5 ns is expressed as 0000 0000 0002 8000 base 16 +*/ +type TimeInterval IntFloat + +// Nanoseconds decodes TimeInterval to human-understandable nanoseconds +func (t TimeInterval) Nanoseconds() float64 { + return IntFloat(t).Value() +} + +func (t TimeInterval) String() string { + return fmt.Sprintf("TimeInterval(%.3fns)", t.Nanoseconds()) +} + +// NewTimeInterval returns TimeInterval built from Nanoseconds +func NewTimeInterval(ns float64) TimeInterval { + return TimeInterval(ns * twoPow16) +} + +/* +Correction is the value of the correction measured in nanoseconds and multiplied by 2**16. +For example, 2.5 ns is represented as 0000 0000 0002 8000 base 16 +A value of one in all bits, except the most significant, of the field shall indicate that the correction is too big to be represented. +*/ +type Correction IntFloat + +// Nanoseconds decodes Correction to human-understandable nanoseconds +func (t Correction) Nanoseconds() float64 { + if t.TooBig() { + return math.Inf(1) + } + return IntFloat(t).Value() +} + +// Duration converts PTP CorrectionField to time.Duration, ignoring +// case where correction is too big, and dropping fractions of nanoseconds +func (t Correction) Duration() time.Duration { + if !t.TooBig() { + return time.Duration(t.Nanoseconds()) + } + return 0 +} + +func (t Correction) String() string { + if t.TooBig() { + return "Correction(Too big)" + } + return fmt.Sprintf("Correction(%.3fns)", t.Nanoseconds()) +} + +// TooBig means correction is too big to be represented. +func (t Correction) TooBig() bool { + return t == 0x7fffffffffffffff // one in all bits, except the most significant +} + +// NewCorrection returns Correction built from Nanoseconds +func NewCorrection(ns float64) Correction { + t := ns * twoPow16 + if t > 0x7fffffffffffffff { + return Correction(0x7fffffffffffffff) + } + return Correction(ns * twoPow16) +} + +// The ClockIdentity type identifies unique entities within a PTP Network, e.g. a PTP Instance or an entity of a common service. +type ClockIdentity uint64 + +// String formats ClockIdentity same way ptp4l pmc client does +func (c ClockIdentity) String() string { + ptr := make([]byte, 8) + binary.BigEndian.PutUint64(ptr, uint64(c)) + return fmt.Sprintf("%02x%02x%02x.%02x%02x.%02x%02x%02x", + ptr[0], ptr[1], ptr[2], ptr[3], + ptr[4], ptr[5], ptr[6], ptr[7], + ) +} + +// MAC turns ClockIdentity into the MAC address it was based upon. EUI-48 is assumed. +func (c ClockIdentity) MAC() net.HardwareAddr { + mac := make(net.HardwareAddr, 6) + mac[0] = byte(c >> 56) + mac[1] = byte(c >> 48) + mac[2] = byte(c >> 40) + mac[3] = byte(c >> 16) + mac[4] = byte(c >> 8) + mac[5] = byte(c) + return mac +} + +// NewClockIdentity creates new ClockIdentity from MAC address +func NewClockIdentity(mac net.HardwareAddr) (ClockIdentity, error) { + b := [8]byte{} + macLen := len(mac) + switch macLen { + case 6: // EUI-48 + b[0] = mac[0] + b[1] = mac[1] + b[2] = mac[2] + b[3] = 0xFF + b[4] = 0xFE + b[5] = mac[3] + b[6] = mac[4] + b[7] = mac[5] + case 8: // EUI-64 + copy(b[:], mac) + default: + return 0, fmt.Errorf("unsupported MAC %v, must be either EUI48 or EUI64", mac) + } + return ClockIdentity(binary.BigEndian.Uint64(b[:])), nil +} + +// The PortIdentity type identifies a PTP Port or a Link Port +type PortIdentity struct { + ClockIdentity ClockIdentity + PortNumber uint16 +} + +// String formats PortIdentity same way ptp4l pmc client does +func (p PortIdentity) String() string { + return fmt.Sprintf("%s-%d", p.ClockIdentity, p.PortNumber) +} + +// Compare returns an integer comparing two port identities. The result will be 0 if p == q, -1 if p < q, and +1 if p > q. +// The definition of "less than" is the same as the Less method. +func (p PortIdentity) Compare(q PortIdentity) int { + cl1, cl2 := p.ClockIdentity, q.ClockIdentity + switch { + case cl1 < cl2: + return -1 + case cl1 > cl2: + return 1 + } + // cl1 == cl2 + pn1, pn2 := p.PortNumber, q.PortNumber + switch { + case pn1 < pn2: + return -1 + case pn1 > pn2: + return 1 + } + // pn1 == pn2 + return 0 +} + +// Less reports whether p sorts before q. Port identities sort first by clock identity, then their port numbers. +func (p PortIdentity) Less(q PortIdentity) bool { return p.Compare(q) == -1 } + +// PTPSeconds type representing seconds +type PTPSeconds [6]uint8 // uint48 + +// Empty returns 0 seconds +func (s PTPSeconds) Empty() bool { + return s == [6]uint8{0, 0, 0, 0, 0, 0} +} + +// Seconds returns number of seconds as uint64 +func (s PTPSeconds) Seconds() uint64 { + return uint64(s[5]) | uint64(s[4])<<8 | uint64(s[3])<<16 | uint64(s[2])<<24 | + uint64(s[1])<<32 | uint64(s[0])<<40 +} + +// Time returns number of seconds in as Time +func (s PTPSeconds) Time() time.Time { + if s.Empty() { + return time.Time{} + } + return time.Unix(int64(s.Seconds()), 0) +} + +// String returns number of seconds in as String +func (s PTPSeconds) String() string { + if s.Empty() { + return "PTPSeconds(empty)" + } + return fmt.Sprintf("PTPSeconds(%s)", s.Time()) +} + +// NewPTPSeconds creates a new instance of PTPSeconds +func NewPTPSeconds(t time.Time) PTPSeconds { + if t.IsZero() { + return PTPSeconds{} + } + v := uint64(t.Unix()) + s := PTPSeconds{} + s[0] = byte(v >> 40) + s[1] = byte(v >> 32) + s[2] = byte(v >> 24) + s[3] = byte(v >> 16) + s[4] = byte(v >> 8) + s[5] = byte(v) + return s +} + +/* +Timestamp type represents a positive time with respect to the epoch. +The secondsField member is the integer portion of the timestamp in units of seconds. +The nanosecondsField member is the fractional portion of the timestamp in units of nanoseconds. +The nanosecondsField member is always less than 10**9 . +For example: ++2.000000001 seconds is represented by secondsField = 0000 0000 0002 base 16 and nanosecondsField= 0000 0001 base 16. +*/ +type Timestamp struct { + Seconds PTPSeconds + Nanoseconds uint32 +} + +// Time turns Timestamp into normal Go time.Time +func (t Timestamp) Time() time.Time { + if t.Empty() { + return time.Time{} + } + return time.Unix(int64(t.Seconds.Seconds()), int64(t.Nanoseconds)) +} + +// Empty timestamp +func (t Timestamp) Empty() bool { + return t.Nanoseconds == 0 && t.Seconds.Empty() +} + +// String representation of the timestamp +func (t Timestamp) String() string { + if t.Empty() { + return "Timestamp(empty)" + } + return fmt.Sprintf("Timestamp(%s)", t.Time()) +} + +// NewTimestamp allows to create Timestamp from time.Time +func NewTimestamp(t time.Time) Timestamp { + if t.IsZero() { + return Timestamp{} + } + ts := Timestamp{ + Nanoseconds: uint32(t.Nanosecond()), + } + v := uint64(t.Unix()) + ts.Seconds[0] = byte(v >> 40) + ts.Seconds[1] = byte(v >> 32) + ts.Seconds[2] = byte(v >> 24) + ts.Seconds[3] = byte(v >> 16) + ts.Seconds[4] = byte(v >> 8) + ts.Seconds[5] = byte(v) + return ts +} + +// ClockClass represents a PTP clock class +type ClockClass uint8 + +// Available Clock Classes +// https://datatracker.ietf.org/doc/html/rfc8173#section-7.6.2.4 +const ( + ClockClass6 ClockClass = 6 + ClockClass7 ClockClass = 7 + ClockClass13 ClockClass = 13 + ClockClass14 ClockClass = 14 + ClockClass52 ClockClass = 52 + ClockClass58 ClockClass = 58 + ClockClassSlaveOnly ClockClass = 255 +) + +// ClockAccuracy represents a PTP clock accuracy +type ClockAccuracy uint8 + +// Available Clock Accuracy +// https://datatracker.ietf.org/doc/html/rfc8173#section-7.6.2.5 +const ( + ClockAccuracyNanosecond25 ClockAccuracy = 0x20 + ClockAccuracyNanosecond100 ClockAccuracy = 0x21 + ClockAccuracyNanosecond250 ClockAccuracy = 0x22 + ClockAccuracyMicrosecond1 ClockAccuracy = 0x23 + ClockAccuracyMicrosecond2point5 ClockAccuracy = 0x24 + ClockAccuracyMicrosecond10 ClockAccuracy = 0x25 + ClockAccuracyMicrosecond25 ClockAccuracy = 0x26 + ClockAccuracyMicrosecond100 ClockAccuracy = 0x27 + ClockAccuracyMicrosecond250 ClockAccuracy = 0x28 + ClockAccuracyMillisecond1 ClockAccuracy = 0x29 + ClockAccuracyMillisecond2point5 ClockAccuracy = 0x2A + ClockAccuracyMillisecond10 ClockAccuracy = 0x2B + ClockAccuracyMillisecond25 ClockAccuracy = 0x2C + ClockAccuracyMillisecond100 ClockAccuracy = 0x2D + ClockAccuracyMillisecond250 ClockAccuracy = 0x2E + ClockAccuracySecond1 ClockAccuracy = 0x2F + ClockAccuracySecond10 ClockAccuracy = 0x30 + ClockAccuracySecondGreater10 ClockAccuracy = 0x31 + ClockAccuracyUnknown ClockAccuracy = 0xFE +) + +// ClockAccuracyFromOffset returns PTP Clock Accuracy covering the time.Duration +func ClockAccuracyFromOffset(offset time.Duration) ClockAccuracy { + if offset < 0 { + offset *= -1 + } + + // https://datatracker.ietf.org/doc/html/rfc8173#section-7.6.2.4 + if offset <= 25*time.Nanosecond { + return ClockAccuracyNanosecond25 + } else if offset <= 100*time.Nanosecond { + return ClockAccuracyNanosecond100 + } else if offset <= 250*time.Nanosecond { + return ClockAccuracyNanosecond250 + } else if offset <= time.Microsecond { + return ClockAccuracyMicrosecond1 + } else if offset <= 2500*time.Nanosecond { + return ClockAccuracyMicrosecond2point5 + } else if offset <= 10*time.Microsecond { + return ClockAccuracyMicrosecond10 + } else if offset <= 25*time.Microsecond { + return ClockAccuracyMicrosecond25 + } else if offset <= 100*time.Microsecond { + return ClockAccuracyMicrosecond100 + } else if offset <= 250*time.Microsecond { + return ClockAccuracyMicrosecond250 + } else if offset <= time.Millisecond { + return ClockAccuracyMillisecond1 + } else if offset <= 2500*time.Microsecond { + return ClockAccuracyMillisecond2point5 + } else if offset <= 10*time.Millisecond { + return ClockAccuracyMillisecond10 + } else if offset <= 25*time.Millisecond { + return ClockAccuracyMillisecond25 + } else if offset <= 100*time.Millisecond { + return ClockAccuracyMillisecond100 + } else if offset <= 250*time.Millisecond { + return ClockAccuracyMillisecond250 + } else if offset <= time.Second { + return ClockAccuracySecond1 + } else if offset <= 10*time.Second { + return ClockAccuracySecond10 + } + + return ClockAccuracySecondGreater10 +} + +// Duration returns matching time.Duration of PTP Clock Accuracy +func (c ClockAccuracy) Duration() time.Duration { + switch c { + case ClockAccuracyNanosecond25: + return 25 * time.Nanosecond + case ClockAccuracyNanosecond100: + return 100 * time.Nanosecond + case ClockAccuracyNanosecond250: + return 250 * time.Nanosecond + case ClockAccuracyMicrosecond1: + return 1000 * time.Nanosecond + case ClockAccuracyMicrosecond2point5: + return 2500 * time.Nanosecond + case ClockAccuracyMicrosecond10: + return 10 * time.Microsecond + case ClockAccuracyMicrosecond25: + return 25 * time.Microsecond + case ClockAccuracyMicrosecond100: + return 100 * time.Microsecond + case ClockAccuracyMicrosecond250: + return 250 * time.Microsecond + case ClockAccuracyMillisecond1: + return 1 * time.Millisecond + case ClockAccuracyMillisecond2point5: + return 2500 * time.Microsecond + case ClockAccuracyMillisecond10: + return 10 * time.Millisecond + case ClockAccuracyMillisecond25: + return 25 * time.Millisecond + case ClockAccuracyMillisecond100: + return 100 * time.Millisecond + case ClockAccuracyMillisecond250: + return 250 * time.Millisecond + case ClockAccuracySecond1: + return 1 * time.Second + case ClockAccuracySecond10: + return 10 * time.Second + } + return 25 * time.Second +} + +// ClockQuality represents the quality of a clock. +type ClockQuality struct { + ClockClass ClockClass `json:"clock_class"` + ClockAccuracy ClockAccuracy `json:"clock_accuracy"` + OffsetScaledLogVariance uint16 `json:"offset_scaled_log_variance"` +} + +// TimeSource indicates the immediate source of time used by the Grandmaster PTP Instance +type TimeSource uint8 + +// TimeSource values, Table 6 timeSource enumeration +const ( + TimeSourceAtomicClock TimeSource = 0x10 + TimeSourceGNSS TimeSource = 0x20 + TimeSourceTerrestrialRadio TimeSource = 0x30 + TimeSourceSerialTimeCode TimeSource = 0x39 + TimeSourcePTP TimeSource = 0x40 + TimeSourceNTP TimeSource = 0x50 + TimeSourceHandSet TimeSource = 0x60 + TimeSourceOther TimeSource = 0x90 + TimeSourceInternalOscillator TimeSource = 0xa0 +) + +// TimeSourceToString is a map from TimeSource to string +var TimeSourceToString = map[TimeSource]string{ + TimeSourceAtomicClock: "ATOMIC_CLOCK", + TimeSourceGNSS: "GNSS", + TimeSourceTerrestrialRadio: "TERRESTRIAL_RADIO", + TimeSourceSerialTimeCode: "SERIAL_TIME_CODE", + TimeSourcePTP: "PTP", + TimeSourceNTP: "NTP", + TimeSourceHandSet: "HAND_SET", + TimeSourceOther: "OTHER", + TimeSourceInternalOscillator: "INTERNAL_OSCILLATOR", +} + +func (t TimeSource) String() string { + return TimeSourceToString[t] +} + +// LogInterval shall be the logarithm, to base 2, of the requested period in seconds. +// In layman's terms, it's specified as a power of two in seconds. +type LogInterval int8 + +// Duration returns LogInterval as time.Duration +func (i LogInterval) Duration() time.Duration { + secs := math.Pow(2, float64(i)) + return time.Duration(secs * float64(time.Second)) +} + +// NewLogInterval returns new LogInterval from time.Duration. +// The values of these logarithmic attributes shall be selected from integers in the range -128 to 127 subject to +// further limits established in the applicable PTP Profile. +func NewLogInterval(d time.Duration) (LogInterval, error) { + li := int(math.Log2(d.Seconds())) + if li > 127 { + return 0, fmt.Errorf("logInterval %d is too big", li) + } + if li < -128 { + return 0, fmt.Errorf("logInterval %d is too small", li) + } + return LogInterval(li), nil +} + +/* +PTPText data type is used to represent textual material in PTP messages. +TextField is encoded as UTF-8. +The most significant byte of the leading text symbol shall be the element of the array with index 0. +UTF-8 encoding has variable length, thus LengthField can be larger than number of characters. + + type PTPText struct { + LengthField uint8 + TextField []byte + } +*/ +type PTPText string + +// UnmarshalBinary populates ptptext from bytes +func (p *PTPText) UnmarshalBinary(rawBytes []byte) error { + var length uint8 + reader := bytes.NewReader(rawBytes) + if err := binary.Read(reader, binary.BigEndian, &length); err != nil { + return fmt.Errorf("reading PTPText LengthField: %w", err) + } + if length == 0 { + // can be zero len, just empty string + return nil + } + + if len(rawBytes) < int(length+1) { + return fmt.Errorf("text field is too short, need %d got %d", len(rawBytes), length+1) + } + text := make([]byte, length) + if err := binary.Read(reader, binary.BigEndian, text); err != nil { + return fmt.Errorf("reading PTPText TextField of len=%d: %w", length, err) + } + *p = PTPText(text) + return nil +} + +// MarshalBinary converts ptptext to []bytes +func (p *PTPText) MarshalBinary() ([]byte, error) { + rawText := []byte(*p) + if len(rawText) > 255 { + return nil, fmt.Errorf("text is too long") + } + length := uint8(len(rawText)) + var bytes bytes.Buffer + if err := binary.Write(&bytes, binary.BigEndian, length); err != nil { + return nil, err + } + if err := binary.Write(&bytes, binary.BigEndian, rawText); err != nil { + return nil, err + } + // padding to make sure packet length is even + if length%2 != 0 { + if err := bytes.WriteByte(0); err != nil { + return nil, err + } + } + return bytes.Bytes(), nil +} + +// PortState is a enum describing one of possible states of port state machines +type PortState uint8 + +// Table 20 PTP state enumeration +const ( + PortStateInitializing PortState = iota + 1 + PortStateFaulty + PortStateDisabled + PortStateListening + PortStatePreMaster + PortStateMaster + PortStatePassive + PortStateUncalibrated + PortStateSlave + PortStateGrandMaster /*non-standard extension*/ +) + +// PortStateToString is a map from PortState to string +var PortStateToString = map[PortState]string{ + PortStateInitializing: "INITIALIZING", + PortStateFaulty: "FAULTY", + PortStateDisabled: "DISABLED", + PortStateListening: "LISTENING", + PortStatePreMaster: "PRE_MASTER", + PortStateMaster: "MASTER", + PortStatePassive: "PASSIVE", + PortStateUncalibrated: "UNCALIBRATED", + PortStateSlave: "SLAVE", + PortStateGrandMaster: "GRAND_MASTER", +} + +func (ps PortState) String() string { + return PortStateToString[ps] +} + +// TransportType is a enum describing network transport protocol types +type TransportType uint16 + +// Table 3 networkProtocol enumeration +const ( + /* 0 is Reserved in spec. Use it for UDS */ + TransportTypeUDS TransportType = iota + TransportTypeUDPIPV4 + TransportTypeUDPIPV6 + TransportTypeIEEE8023 + TransportTypeDeviceNet + TransportTypeControlNet + TransportTypePROFINET +) + +// TransportTypeToString is a map from TransportType to string +var TransportTypeToString = map[TransportType]string{ + TransportTypeUDS: "UDS", + TransportTypeUDPIPV4: "UDP_IPV4", + TransportTypeUDPIPV6: "UDP_IPV6", + TransportTypeIEEE8023: "IEEE_802_3", + TransportTypeDeviceNet: "DEVICENET", + TransportTypeControlNet: "CONTROLNET", + TransportTypePROFINET: "PROFINET", +} + +func (t TransportType) String() string { + return TransportTypeToString[t] +} + +// PortAddress see 5.3.6 PortAddress +type PortAddress struct { + NetworkProtocol TransportType + AddressLength uint16 + AddressField []byte +} + +// UnmarshalBinary converts bytes to PortAddress +func (p *PortAddress) UnmarshalBinary(b []byte) error { + if len(b) < 8 { + return fmt.Errorf("not enough data to decode PortAddress") + } + p.NetworkProtocol = TransportType(binary.BigEndian.Uint16(b[0:])) + p.AddressLength = binary.BigEndian.Uint16(b[2:]) + if len(b) < 4+int(p.AddressLength) { + return fmt.Errorf("not enough data to decode PortAddress address") + } + p.AddressField = make([]byte, p.AddressLength) + copy(p.AddressField, b[4:4+p.AddressLength]) + return nil +} + +// IP converts PortAddress to IP +func (p *PortAddress) IP() (net.IP, error) { + if p.NetworkProtocol != TransportTypeUDPIPV4 && p.NetworkProtocol != TransportTypeUDPIPV6 { + return nil, fmt.Errorf("unsupported network protocol %s (%d)", p.NetworkProtocol, p.NetworkProtocol) + } + if p.NetworkProtocol == TransportTypeUDPIPV4 && (p.AddressLength != 4 || len(p.AddressField) != 4) { + return nil, fmt.Errorf("unexpected length of IPv4: %d", len(p.AddressField)) + } + if p.NetworkProtocol == TransportTypeUDPIPV6 && (p.AddressLength != 16 || len(p.AddressField) != 16) { + return nil, fmt.Errorf("unexpected length of IPv6: %d", len(p.AddressField)) + } + return net.IP(p.AddressField), nil +} + +// MarshalBinary converts PortAddress to []bytes +func (p *PortAddress) MarshalBinary() ([]byte, error) { + var bytes bytes.Buffer + if err := binary.Write(&bytes, binary.BigEndian, p.NetworkProtocol); err != nil { + return nil, err + } + if err := binary.Write(&bytes, binary.BigEndian, p.AddressLength); err != nil { + return nil, err + } + if err := binary.Write(&bytes, binary.BigEndian, p.AddressField); err != nil { + return nil, err + } + return bytes.Bytes(), nil +} diff --git a/src/yangerd/vendor/github.com/facebook/time/ptp/protocol/unicast.go b/src/yangerd/vendor/github.com/facebook/time/ptp/protocol/unicast.go new file mode 100644 index 000000000..0f7bbcd06 --- /dev/null +++ b/src/yangerd/vendor/github.com/facebook/time/ptp/protocol/unicast.go @@ -0,0 +1,92 @@ +/* +Copyright (c) Facebook, Inc. and its affiliates. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package protocol + +import ( + "encoding/binary" + "fmt" +) + +// UnicastMsgTypeAndFlags is a uint8 where first 4 bites contain MessageType and last 4 bits contain some flags +type UnicastMsgTypeAndFlags uint8 + +// MsgType extracts MessageType from UnicastMsgTypeAndFlags +func (m UnicastMsgTypeAndFlags) MsgType() MessageType { + return MessageType(m >> 4) +} + +// NewUnicastMsgTypeAndFlags builds new UnicastMsgTypeAndFlags from MessageType and flags +func NewUnicastMsgTypeAndFlags(msgType MessageType, flags uint8) UnicastMsgTypeAndFlags { + return UnicastMsgTypeAndFlags(uint8(msgType)<<4 | (flags & 0x0f)) +} + +// Signaling packet. As it's of variable size, we cannot just binary.Read/Write it. +type Signaling struct { + Header + TargetPortIdentity PortIdentity + TLVs []TLV +} + +// MarshalBinaryTo marshals bytes to Signaling +func (p *Signaling) MarshalBinaryTo(b []byte) (int, error) { + if len(p.TLVs) == 0 { + return 0, fmt.Errorf("no TLVs in Signaling message, at least one required") + } + n := headerMarshalBinaryTo(&p.Header, b) + binary.BigEndian.PutUint64(b[n:], uint64(p.TargetPortIdentity.ClockIdentity)) + binary.BigEndian.PutUint16(b[n+8:], p.TargetPortIdentity.PortNumber) + pos := n + 10 + tlvLen, err := writeTLVs(p.TLVs, b[pos:]) + return pos + tlvLen, err +} + +// MarshalBinary converts packet to []bytes +func (p *Signaling) MarshalBinary() ([]byte, error) { + buf := make([]byte, 508) + n, err := p.MarshalBinaryTo(buf) + return buf[:n], err +} + +// UnmarshalBinary parses []byte and populates struct fields +func (p *Signaling) UnmarshalBinary(b []byte) error { + if len(b) < headerSize+10+tlvHeadSize { + return fmt.Errorf("not enough data to decode Signaling") + } + + unmarshalHeader(&p.Header, b) + if err := checkPacketLength(&p.Header, len(b)); err != nil { + return err + } + + if p.SdoIDAndMsgType.MsgType() != MessageSignaling { + return fmt.Errorf("not a signaling message %v", b) + } + + p.TargetPortIdentity.ClockIdentity = ClockIdentity(binary.BigEndian.Uint64(b[headerSize:])) + p.TargetPortIdentity.PortNumber = binary.BigEndian.Uint16(b[headerSize+8:]) + + pos := headerSize + 10 + var err error + p.TLVs, err = readTLVs(p.TLVs, int(p.MessageLength)-pos, b[pos:]) + if err != nil { + return err + } + if len(p.TLVs) == 0 { + return fmt.Errorf("no TLVs read for Signaling message, at least one required") + } + return nil +} diff --git a/src/yangerd/vendor/github.com/godbus/dbus/v5/.cirrus.yml b/src/yangerd/vendor/github.com/godbus/dbus/v5/.cirrus.yml new file mode 100644 index 000000000..6e2090296 --- /dev/null +++ b/src/yangerd/vendor/github.com/godbus/dbus/v5/.cirrus.yml @@ -0,0 +1,11 @@ +# See https://cirrus-ci.org/guide/FreeBSD/ +freebsd_instance: + image_family: freebsd-14-3 + +task: + name: Test on FreeBSD + install_script: pkg install -y go125 dbus + test_script: | + /usr/local/etc/rc.d/dbus onestart && \ + eval `dbus-launch --sh-syntax` && \ + go125 test -v ./... diff --git a/src/yangerd/vendor/github.com/godbus/dbus/v5/.golangci.yml b/src/yangerd/vendor/github.com/godbus/dbus/v5/.golangci.yml new file mode 100644 index 000000000..5bbdd9342 --- /dev/null +++ b/src/yangerd/vendor/github.com/godbus/dbus/v5/.golangci.yml @@ -0,0 +1,13 @@ +version: "2" + +linters: + enable: + - unconvert + - unparam + exclusions: + presets: + - std-error-handling + +formatters: + enable: + - gofumpt diff --git a/src/yangerd/vendor/github.com/godbus/dbus/v5/CONTRIBUTING.md b/src/yangerd/vendor/github.com/godbus/dbus/v5/CONTRIBUTING.md new file mode 100644 index 000000000..c88f9b2bd --- /dev/null +++ b/src/yangerd/vendor/github.com/godbus/dbus/v5/CONTRIBUTING.md @@ -0,0 +1,50 @@ +# How to Contribute + +## Getting Started + +- Fork the repository on GitHub +- Read the [README](README.markdown) for build and test instructions +- Play with the project, submit bugs, submit patches! + +## Contribution Flow + +This is a rough outline of what a contributor's workflow looks like: + +- Create a topic branch from where you want to base your work (usually master). +- Make commits of logical units. +- Make sure your commit messages are in the proper format (see below). +- Push your changes to a topic branch in your fork of the repository. +- Make sure the tests pass, and add any new tests as appropriate. +- Submit a pull request to the original repository. + +Thanks for your contributions! + +### Format of the Commit Message + +We follow a rough convention for commit messages that is designed to answer two +questions: what changed and why. The subject line should feature the what and +the body of the commit should describe the why. + +``` +scripts: add the test-cluster command + +this uses tmux to setup a test cluster that you can easily kill and +start for debugging. + +Fixes #38 +``` + +The format can be described more formally as follows: + +``` +: + + + +