Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
42 changes: 42 additions & 0 deletions tests/test_csv_label_collisions.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,42 @@
import importlib.util
import sys
from pathlib import Path


def _load(name):
path = Path(__file__).parents[1] / "tools" / "helpers" / name
spec = importlib.util.spec_from_file_location(name.removesuffix(".py"), path)
module = importlib.util.module_from_spec(spec)
sys.modules[spec.name] = module
spec.loader.exec_module(module)
return module


def _rows(module):
return [
module.ParsedRow("a", "b", 0, -1, 1000, 0.1, "high"),
module.ParsedRow("a", "b", 0, -1, 2000, 0.2, "hi"),
]


def test_compose_disambiguates_shortened_label_collision(tmp_path):
module = _load("csv_to_compose.py")
csv_path = tmp_path / "contacts.csv"
csv_path.write_text(
"src,dst,start,end,bw,delay,label\n"
"a,b,0,-1,1000,0.1,high\n"
"a,b,0,-1,2000,0.2,hi\n"
)
graph = module.get_graph_from_csv(str(csv_path), {})
assert graph.number_of_edges() == 2
assert {data["label"] for *_, data in graph.edges(data=True)} == {"high", "hi"}


def test_ccp_disambiguates_shortened_label_collision():
module = _load("csv_to_ccp.py")
rows = _rows(module)
fixed, _ = module.convert_rows(
[(row, False) for row in rows], module.compute_multi_pairs(rows)
)
destinations = {row[1] for row in fixed}
assert destinations == {"dev:a_b_high", "dev:a_b_hi"}
25 changes: 22 additions & 3 deletions tools/helpers/csv_to_ccp.py
Original file line number Diff line number Diff line change
Expand Up @@ -100,6 +100,17 @@ def shorten_label(label: str) -> str:
return "_".join(LABEL_REPLACEMENTS.get(part, part) for part in label.split("_"))


def compute_label_collisions(rows: list[ParsedRow]) -> set[tuple[str, str, str]]:
"""Return shortened label keys that represent multiple distinct labels."""
labels_by_key: defaultdict[tuple[str, str, str], set[str]] = defaultdict(set)
for r in rows:
a, b = tuple(sorted([r.src, r.dst]))
raw_key = strip_dir_suffix(r.label)
short_key = shorten_label(raw_key)
labels_by_key[(a, b, short_key)].add(raw_key)
return {key for key, labels in labels_by_key.items() if len(labels) > 1}


def compute_multi_pairs(rows: list[ParsedRow]) -> set[tuple[str, str]]:
"""Return node pairs needing a dedicated interface: those with more
than one distinct label after stripping _ul/_dl suffixes."""
Expand All @@ -110,15 +121,22 @@ def compute_multi_pairs(rows: list[ParsedRow]) -> set[tuple[str, str]]:


def make_ifname(
node1: str, node2: str, label: str, multi_pairs: set[tuple[str, str]]
node1: str,
node2: str,
label: str,
multi_pairs: set[tuple[str, str]],
label_collisions: set[tuple[str, str, str]] | None = None,
) -> str | None:
"""Build an interface name for this pair+label, or None if the pair
should just use the plain node name. Falls back to a 12-char MD5 hash
if the name would exceed the 14-char interface limit."""
a, b = tuple(sorted([node1, node2]))
if (a, b) not in multi_pairs:
return None
key = shorten_label(strip_dir_suffix(label))
raw_key = strip_dir_suffix(label)
key = shorten_label(raw_key)
if label_collisions and (a, b, key) in label_collisions:
key = raw_key
ifname = f"{a}_{b}_{key}" if key else f"{a}_{b}"
if len(ifname) >= 14:
ifname_md5 = hashlib.md5(ifname.encode()).hexdigest()[:12]
Expand Down Expand Up @@ -228,13 +246,14 @@ def convert_rows(
"""
fixed: list[FixedRow] = []
contact: list[ContactRow] = []
label_collisions = compute_label_collisions([r for r, _ in rows])

for r, symmetric in rows:
bw = bps_to_human(r.bw)
delay = str(int(round(r.delay * 1000)))
eq = "=" if symmetric else ""

ifname = make_ifname(r.src, r.dst, r.label, multi_pairs)
ifname = make_ifname(r.src, r.dst, r.label, multi_pairs, label_collisions)
dst = f"dev:{ifname}" if ifname else r.dst

if r.ts_start == 0 and r.ts_end == -1:
Expand Down
28 changes: 25 additions & 3 deletions tools/helpers/csv_to_compose.py
Original file line number Diff line number Diff line change
Expand Up @@ -78,6 +78,17 @@ def shorten_label(label: str) -> str:
return "_".join(LABEL_REPLACEMENTS.get(part, part) for part in label.split("_"))


def compute_label_collisions(rows: list[ParsedRow]) -> set[tuple[str, str, str]]:
"""Return shortened label keys that represent multiple distinct labels."""
labels_by_key: defaultdict[tuple[str, str, str], set[str]] = defaultdict(set)
for r in rows:
a, b = tuple(sorted([r.src, r.dst]))
raw_key = strip_dir_suffix(r.label)
short_key = shorten_label(raw_key)
labels_by_key[(a, b, short_key)].add(raw_key)
return {key for key, labels in labels_by_key.items() if len(labels) > 1}


def compute_multi_pairs(rows: list[ParsedRow]) -> set[tuple[str, str]]:
"""Return node pairs needing a dedicated interface: those with more
than one distinct label after stripping _ul/_dl suffixes."""
Expand All @@ -88,15 +99,22 @@ def compute_multi_pairs(rows: list[ParsedRow]) -> set[tuple[str, str]]:


def make_ifname(
node1: str, node2: str, label: str, multi_pairs: set[tuple[str, str]]
node1: str,
node2: str,
label: str,
multi_pairs: set[tuple[str, str]],
label_collisions: set[tuple[str, str, str]] | None = None,
) -> str | None:
"""Build an interface name for this pair+label, or None if the pair
should just use the plain node name. Falls back to a 12-char MD5 hash
if the name would exceed the 14-char interface limit."""
a, b = tuple(sorted([node1, node2]))
if (a, b) not in multi_pairs:
return None
key = shorten_label(strip_dir_suffix(label))
raw_key = strip_dir_suffix(label)
key = shorten_label(raw_key)
if label_collisions and (a, b, key) in label_collisions:
key = raw_key
ifname = f"{a}_{b}_{key}" if key else f"{a}_{b}"
if len(ifname) >= 14:
ifname_md5 = hashlib.md5(ifname.encode()).hexdigest()[:12]
Expand Down Expand Up @@ -141,6 +159,7 @@ def get_graph_from_csv(
G: nx.MultiDiGraph[str] = nx.MultiDiGraph()
rows = parse_csv_rows(csvfile, prefix=prefix)
multi_pairs = compute_multi_pairs(rows)
label_collisions = compute_label_collisions(rows)

for r in rows:
for node in (r.src, r.dst):
Expand Down Expand Up @@ -168,7 +187,10 @@ def get_graph_from_csv(

dynamic_link = not (r.ts_start == 0 and r.ts_end == -1)
node1, node2 = tuple(sorted([r.src, r.dst]))
net_name = make_ifname(r.src, r.dst, r.label, multi_pairs) or f"{node1}_{node2}"
net_name = (
make_ifname(r.src, r.dst, r.label, multi_pairs, label_collisions)
or f"{node1}_{node2}"
)

if not G.has_edge(r.src, r.dst, key=net_name):
G.add_edge(
Expand Down