aboutsummaryrefslogtreecommitdiffstats
path: root/packages/meshbay-node/src/meshbay_node/ops/node_toml.py
diff options
context:
space:
mode:
Diffstat (limited to 'packages/meshbay-node/src/meshbay_node/ops/node_toml.py')
-rw-r--r--packages/meshbay-node/src/meshbay_node/ops/node_toml.py59
1 files changed, 47 insertions, 12 deletions
diff --git a/packages/meshbay-node/src/meshbay_node/ops/node_toml.py b/packages/meshbay-node/src/meshbay_node/ops/node_toml.py
index f711f26..2407722 100644
--- a/packages/meshbay-node/src/meshbay_node/ops/node_toml.py
+++ b/packages/meshbay-node/src/meshbay_node/ops/node_toml.py
@@ -3,14 +3,50 @@
from __future__ import annotations
import re
+import tomllib
from pathlib import Path
from meshbay_node.ops.core import OpError
+def toml_string(value: str) -> str:
+ """A TOML basic string holding `value` exactly, quotes included.
+
+ Every string written into node.toml goes through here. A value with a quote
+ or a newline in it — a group name, a folder name, any of them chosen by
+ someone else — would otherwise end the string and write lines of its own.
+ """
+ out = ['"']
+ for ch in str(value):
+ if ch == '"':
+ out.append('\\"')
+ elif ch == "\\":
+ out.append("\\\\")
+ elif ord(ch) < 0x20 or ord(ch) == 0x7F:
+ out.append(f"\\u{ord(ch):04x}")
+ else:
+ out.append(ch)
+ out.append('"')
+ return "".join(out)
+
+
+def _string_value(line: str, key: str) -> str | None:
+ """The string `key` holds on this line, unescaped — or None.
+
+ Read as TOML, not by pattern: a value written by `toml_string` may carry an
+ escaped quote or backslash, which a `"([^"]*)"` pattern would cut short.
+ """
+ if not re.match(r"^\s*" + re.escape(key) + r"\s*=", line):
+ return None
+ try:
+ value = tomllib.loads(line.strip()).get(key)
+ except tomllib.TOMLDecodeError:
+ return None
+ return value if isinstance(value, str) else None
+
+
def _find_group_range(lines: list[str], group_id: str) -> tuple[int, int] | None:
"""Line range of a [[groups]] block by id: (start, end_exclusive)."""
- id_re = re.compile(r'^\s*id\s*=\s*"([^"]*)"')
block_starts: list[int] = []
for i, line in enumerate(lines):
if line.strip() == "[[groups]]":
@@ -24,8 +60,7 @@ def _find_group_range(lines: list[str], group_id: str) -> tuple[int, int] | None
boundary = k
break
for k in range(start + 1, boundary):
- m = id_re.match(lines[k])
- if m and m.group(1) == group_id:
+ if _string_value(lines[k], "id") == group_id:
return (start, boundary)
return None
@@ -61,8 +96,10 @@ def _update_node_toml(conf_path: Path, updates: dict) -> None:
if isinstance(value, bool):
return f"{key} = {'true' if value else 'false'}"
if isinstance(value, list):
- items = ", ".join(f'"{v}"' for v in value)
+ items = ", ".join(toml_string(v) for v in value)
return f"{key} = [{items}]"
+ if isinstance(value, str):
+ return f"{key} = {toml_string(value)}"
return f"{key} = {value}"
remaining = dict(updates)
@@ -115,7 +152,6 @@ def _remove_roots_block(conf_path: Path, group_id: str,
raise OpError(f"Group {group_id[:8]} not found in {conf_path}")
start, end = rng
- path_re = re.compile(r'^\s*path\s*=\s*"([^"]*)"')
roots_starts: list[int] = []
for i in range(start + 1, end):
if lines[i].strip() == "[[groups.roots]]":
@@ -124,10 +160,10 @@ def _remove_roots_block(conf_path: Path, group_id: str,
for j, rs in enumerate(roots_starts):
rs_end = roots_starts[j + 1] if j + 1 < len(roots_starts) else end
for k in range(rs, rs_end):
- m = path_re.match(lines[k])
- if m:
+ raw = _string_value(lines[k], "path")
+ if raw is not None:
try:
- p = str(Path(m.group(1)).expanduser().resolve())
+ p = str(Path(raw).expanduser().resolve())
except OSError:
continue
if p == resolved_path:
@@ -153,7 +189,6 @@ def _update_root_field(conf_path: Path, group_id: str,
raise OpError(f"Group {group_id[:8]} not found in {conf_path}")
start, end = rng
- path_re = re.compile(r'^\s*path\s*=\s*"([^"]*)"')
writable_re = re.compile(r'^\s*(writable|upload)\s*=')
removable_re = re.compile(r'^\s*removable\s*=')
roots_starts: list[int] = []
@@ -165,10 +200,10 @@ def _update_root_field(conf_path: Path, group_id: str,
rs_end = roots_starts[j + 1] if j + 1 < len(roots_starts) else end
found_path = False
for k in range(rs, rs_end):
- m = path_re.match(lines[k])
- if m:
+ raw = _string_value(lines[k], "path")
+ if raw is not None:
try:
- p = str(Path(m.group(1)).expanduser().resolve())
+ p = str(Path(raw).expanduser().resolve())
except OSError:
continue
if p == resolved_path: