diff --git a/SECURITY.md b/SECURITY.md
new file mode 100644
index 00000000..c6b42c23
--- /dev/null
+++ b/SECURITY.md
@@ -0,0 +1,53 @@
+# Security Policy
+
+## Supported Versions
+
+| Version | Supported |
+|---------|-----------|
+| 0.1.x | Yes |
+| < 0.1 | No |
+
+## Reporting a Vulnerability
+
+**Do not open a public GitHub issue for security vulnerabilities.**
+
+Report security issues via GitHub's private vulnerability reporting, or email the maintainer directly. Please include:
+
+- Description of the vulnerability
+- Steps to reproduce
+- Potential impact
+- Suggested fix (if any)
+
+We will acknowledge receipt within 48 hours and aim to release a fix within 7 days for critical issues.
+
+## Security Model
+
+graphify is a **local development tool**. It runs as a Claude Code skill and optionally as a local MCP stdio server. It makes no network calls during graph analysis — only during `ingest` (explicit URL fetch by the user).
+
+### Threat Surface
+
+| Vector | Mitigation |
+|--------|-----------|
+| SSRF via URL fetch | `security.validate_url()` allows only `http` and `https` schemes. Redirect targets are re-validated by `_NoFileRedirectHandler` — a redirect to `file://` is blocked. |
+| Oversized downloads | `safe_fetch()` streams responses and aborts at 50 MB. `safe_fetch_text()` aborts at 10 MB. |
+| Non-2xx HTTP responses | `safe_fetch()` raises `HTTPError` on non-2xx status codes — error pages are not silently treated as content. |
+| Path traversal in MCP server | `security.validate_graph_path()` resolves paths and requires them to be inside `.graphify/`. Also requires the `.graphify/` directory to exist. |
+| XSS in graph HTML output | `security.sanitize_label()` strips control characters, caps at 256 chars, and HTML-escapes all node labels and edge titles before pyvis embeds them. |
+| Prompt injection via node labels | `sanitize_label()` also applied to MCP text output — node labels from user-controlled source files cannot break the text format returned to agents. |
+| YAML frontmatter injection | Newlines stripped from user-provided strings before embedding in YAML frontmatter (e.g. in `save_query_result()`). |
+| Encoding crashes on source files | All tree-sitter byte slices decoded with `errors="replace"` — non-UTF-8 source files degrade gracefully instead of crashing extraction. |
+| Symlink traversal | `os.walk(..., followlinks=False)` is explicit throughout `detect.py`. |
+| Corrupted graph.json | `_load_graph()` in `serve.py` wraps `json.JSONDecodeError` and prints a clear recovery message instead of crashing. |
+
+### What graphify does NOT do
+
+- Does not run a network listener (MCP server communicates over stdio only)
+- Does not execute code from source files (tree-sitter parses ASTs — no eval/exec)
+- Does not use `shell=True` in any subprocess call
+- Does not store credentials or API keys
+
+### Optional network calls
+
+- `ingest` subcommand: fetches URLs explicitly provided by the user
+- PDF extraction: reads local files only (pypdf does not make network calls)
+- watch mode: local filesystem events only (watchdog does not make network calls)
diff --git a/graphify/detect.py b/graphify/detect.py
index 1e8957df..70bbbeb4 100644
--- a/graphify/detect.py
+++ b/graphify/detect.py
@@ -156,7 +156,7 @@ def detect(root: Path) -> dict:
for scan_root in scan_paths:
in_memory_tree = memory_dir.exists() and str(scan_root).startswith(str(memory_dir))
import os
- for dirpath, dirnames, filenames in os.walk(scan_root):
+ for dirpath, dirnames, filenames in os.walk(scan_root, followlinks=False):
dp = Path(dirpath)
if not in_memory_tree:
# Prune noise dirs in-place so os.walk never descends into them
diff --git a/graphify/export.py b/graphify/export.py
index af5d77f1..eb4e02d4 100644
--- a/graphify/export.py
+++ b/graphify/export.py
@@ -7,6 +7,7 @@ from collections import Counter
from pathlib import Path
import networkx as nx
from networkx.readwrite import json_graph
+from graphify.security import sanitize_label
COMMUNITY_COLORS = [
"#4E79A7", "#F28E2B", "#E15759", "#76B7B2", "#59A14F",
@@ -71,9 +72,9 @@ def to_html(
color = COMMUNITY_COLORS[cid % len(COMMUNITY_COLORS)]
net.add_node(
node_id,
- label=data.get("label", node_id),
+ label=sanitize_label(data.get("label", node_id)),
color=color,
- title=(
+ title=sanitize_label(
f"Source: {data.get('source_file', 'unknown')}\n"
f"Type: {data.get('file_type', 'unknown')}\n"
f"Community: {community_labels.get(cid, str(cid)) if community_labels else cid}"
diff --git a/graphify/extract.py b/graphify/extract.py
index bc2fae79..a0a1bffc 100644
--- a/graphify/extract.py
+++ b/graphify/extract.py
@@ -70,7 +70,7 @@ def extract_python(path: Path) -> dict:
if t == "import_statement":
for child in node.children:
if child.type in ("dotted_name", "aliased_import"):
- raw = source[child.start_byte:child.end_byte].decode()
+ raw = source[child.start_byte:child.end_byte].decode("utf-8", errors="replace")
module_name = raw.split(" as ")[0].strip().lstrip(".")
tgt_nid = _make_id(module_name)
add_edge(file_nid, tgt_nid, "imports", node.start_point[0] + 1)
@@ -79,7 +79,7 @@ def extract_python(path: Path) -> dict:
if t == "import_from_statement":
module_node = node.child_by_field_name("module_name")
if module_node:
- raw = source[module_node.start_byte:module_node.end_byte].decode().lstrip(".")
+ raw = source[module_node.start_byte:module_node.end_byte].decode("utf-8", errors="replace").lstrip(".")
tgt_nid = _make_id(raw)
add_edge(file_nid, tgt_nid, "imports_from", node.start_point[0] + 1)
return
@@ -88,7 +88,7 @@ def extract_python(path: Path) -> dict:
name_node = node.child_by_field_name("name")
if not name_node:
return
- class_name = source[name_node.start_byte:name_node.end_byte].decode()
+ class_name = source[name_node.start_byte:name_node.end_byte].decode("utf-8", errors="replace")
class_nid = _make_id(stem, class_name)
line = node.start_point[0] + 1
add_node(class_nid, class_name, line)
@@ -99,7 +99,7 @@ def extract_python(path: Path) -> dict:
if args:
for arg in args.children:
if arg.type == "identifier":
- base = source[arg.start_byte:arg.end_byte].decode()
+ base = source[arg.start_byte:arg.end_byte].decode("utf-8", errors="replace")
# Try same-file base first; fall back to a bare stub
base_nid = _make_id(stem, base)
if base_nid not in seen_ids:
@@ -127,7 +127,7 @@ def extract_python(path: Path) -> dict:
name_node = node.child_by_field_name("name")
if not name_node:
return
- func_name = source[name_node.start_byte:name_node.end_byte].decode()
+ func_name = source[name_node.start_byte:name_node.end_byte].decode("utf-8", errors="replace")
line = node.start_point[0] + 1
if parent_class_nid:
func_nid = _make_id(parent_class_nid, func_name)
@@ -170,11 +170,11 @@ def extract_python(path: Path) -> dict:
callee_name: str | None = None
if func_node:
if func_node.type == "identifier":
- callee_name = source[func_node.start_byte:func_node.end_byte].decode()
+ callee_name = source[func_node.start_byte:func_node.end_byte].decode("utf-8", errors="replace")
elif func_node.type == "attribute":
attr = func_node.child_by_field_name("attribute")
if attr:
- callee_name = source[attr.start_byte:attr.end_byte].decode()
+ callee_name = source[attr.start_byte:attr.end_byte].decode("utf-8", errors="replace")
if callee_name:
tgt_nid = label_to_nid.get(callee_name.lower())
if tgt_nid and tgt_nid != caller_nid:
@@ -273,7 +273,7 @@ def extract_js(path: Path) -> dict:
if t == "import_statement":
for child in node.children:
if child.type == "string":
- raw = source[child.start_byte:child.end_byte].decode().strip("'\"` ")
+ raw = source[child.start_byte:child.end_byte].decode("utf-8", errors="replace").strip("'\"` ")
module_name = raw.lstrip("./").split("/")[-1]
if module_name:
tgt_nid = _make_id(module_name)
@@ -284,7 +284,7 @@ def extract_js(path: Path) -> dict:
name_node = node.child_by_field_name("name")
if not name_node:
return
- class_name = source[name_node.start_byte:name_node.end_byte].decode()
+ class_name = source[name_node.start_byte:name_node.end_byte].decode("utf-8", errors="replace")
class_nid = _make_id(stem, class_name)
line = node.start_point[0] + 1
add_node(class_nid, class_name, line)
@@ -299,7 +299,7 @@ def extract_js(path: Path) -> dict:
name_node = node.child_by_field_name("name")
if not name_node:
return
- func_name = source[name_node.start_byte:name_node.end_byte].decode()
+ func_name = source[name_node.start_byte:name_node.end_byte].decode("utf-8", errors="replace")
line = node.start_point[0] + 1
func_nid = _make_id(stem, func_name)
add_node(func_nid, f"{func_name}()", line)
@@ -313,7 +313,7 @@ def extract_js(path: Path) -> dict:
name_node = node.child_by_field_name("name")
if not name_node:
return
- method_name = source[name_node.start_byte:name_node.end_byte].decode()
+ method_name = source[name_node.start_byte:name_node.end_byte].decode("utf-8", errors="replace")
line = node.start_point[0] + 1
method_nid = _make_id(parent_class_nid, method_name)
add_node(method_nid, f".{method_name}()", line)
@@ -331,7 +331,7 @@ def extract_js(path: Path) -> dict:
if value and value.type == "arrow_function":
name_node = child.child_by_field_name("name")
if name_node:
- func_name = source[name_node.start_byte:name_node.end_byte].decode()
+ func_name = source[name_node.start_byte:name_node.end_byte].decode("utf-8", errors="replace")
line = child.start_point[0] + 1
func_nid = _make_id(stem, func_name)
add_node(func_nid, f"{func_name}()", line)
@@ -362,11 +362,11 @@ def extract_js(path: Path) -> dict:
callee_name: str | None = None
if func_node:
if func_node.type == "identifier":
- callee_name = source[func_node.start_byte:func_node.end_byte].decode()
+ callee_name = source[func_node.start_byte:func_node.end_byte].decode("utf-8", errors="replace")
elif func_node.type == "member_expression":
prop = func_node.child_by_field_name("property")
if prop:
- callee_name = source[prop.start_byte:prop.end_byte].decode()
+ callee_name = source[prop.start_byte:prop.end_byte].decode("utf-8", errors="replace")
if callee_name:
tgt_nid = label_to_nid.get(callee_name.lower())
if tgt_nid and tgt_nid != caller_nid:
@@ -455,7 +455,7 @@ def extract_go(path: Path) -> dict:
if t == "function_declaration":
name_node = node.child_by_field_name("name")
if name_node:
- func_name = source[name_node.start_byte:name_node.end_byte].decode()
+ func_name = source[name_node.start_byte:name_node.end_byte].decode("utf-8", errors="replace")
line = node.start_point[0] + 1
func_nid = _make_id(stem, func_name)
add_node(func_nid, f"{func_name}()", line)
@@ -473,12 +473,12 @@ def extract_go(path: Path) -> dict:
if param.type == "parameter_declaration":
type_node = param.child_by_field_name("type")
if type_node:
- raw = source[type_node.start_byte:type_node.end_byte].decode().lstrip("*").strip()
+ raw = source[type_node.start_byte:type_node.end_byte].decode("utf-8", errors="replace").lstrip("*").strip()
receiver_type = raw
break
name_node = node.child_by_field_name("name")
if name_node:
- method_name = source[name_node.start_byte:name_node.end_byte].decode()
+ method_name = source[name_node.start_byte:name_node.end_byte].decode("utf-8", errors="replace")
line = node.start_point[0] + 1
if receiver_type:
parent_nid = _make_id(stem, receiver_type)
@@ -500,7 +500,7 @@ def extract_go(path: Path) -> dict:
if child.type == "type_spec":
name_node = child.child_by_field_name("name")
if name_node:
- type_name = source[name_node.start_byte:name_node.end_byte].decode()
+ type_name = source[name_node.start_byte:name_node.end_byte].decode("utf-8", errors="replace")
line = child.start_point[0] + 1
type_nid = _make_id(stem, type_name)
add_node(type_nid, type_name, line)
@@ -514,14 +514,14 @@ def extract_go(path: Path) -> dict:
if spec.type == "import_spec":
path_node = spec.child_by_field_name("path")
if path_node:
- raw = source[path_node.start_byte:path_node.end_byte].decode().strip('"')
+ raw = source[path_node.start_byte:path_node.end_byte].decode("utf-8", errors="replace").strip('"')
module_name = raw.split("/")[-1]
tgt_nid = _make_id(module_name)
add_edge_raw(file_nid, tgt_nid, "imports_from", spec.start_point[0] + 1)
elif child.type == "import_spec":
path_node = child.child_by_field_name("path")
if path_node:
- raw = source[path_node.start_byte:path_node.end_byte].decode().strip('"')
+ raw = source[path_node.start_byte:path_node.end_byte].decode("utf-8", errors="replace").strip('"')
module_name = raw.split("/")[-1]
tgt_nid = _make_id(module_name)
add_edge_raw(file_nid, tgt_nid, "imports_from", child.start_point[0] + 1)
@@ -548,11 +548,11 @@ def extract_go(path: Path) -> dict:
callee_name: str | None = None
if func_node:
if func_node.type == "identifier":
- callee_name = source[func_node.start_byte:func_node.end_byte].decode()
+ callee_name = source[func_node.start_byte:func_node.end_byte].decode("utf-8", errors="replace")
elif func_node.type == "selector_expression":
field = func_node.child_by_field_name("field")
if field:
- callee_name = source[field.start_byte:field.end_byte].decode()
+ callee_name = source[field.start_byte:field.end_byte].decode("utf-8", errors="replace")
if callee_name:
tgt_nid = label_to_nid.get(callee_name.lower())
if tgt_nid and tgt_nid != caller_nid:
@@ -641,7 +641,7 @@ def extract_rust(path: Path) -> dict:
if t == "function_item":
name_node = node.child_by_field_name("name")
if name_node:
- func_name = source[name_node.start_byte:name_node.end_byte].decode()
+ func_name = source[name_node.start_byte:name_node.end_byte].decode("utf-8", errors="replace")
line = node.start_point[0] + 1
if parent_impl_nid:
func_nid = _make_id(parent_impl_nid, func_name)
@@ -659,7 +659,7 @@ def extract_rust(path: Path) -> dict:
if t in ("struct_item", "enum_item", "trait_item"):
name_node = node.child_by_field_name("name")
if name_node:
- item_name = source[name_node.start_byte:name_node.end_byte].decode()
+ item_name = source[name_node.start_byte:name_node.end_byte].decode("utf-8", errors="replace")
line = node.start_point[0] + 1
item_nid = _make_id(stem, item_name)
add_node(item_nid, item_name, line)
@@ -670,7 +670,7 @@ def extract_rust(path: Path) -> dict:
type_node = node.child_by_field_name("type")
impl_nid: str | None = None
if type_node:
- type_name = source[type_node.start_byte:type_node.end_byte].decode().strip()
+ type_name = source[type_node.start_byte:type_node.end_byte].decode("utf-8", errors="replace").strip()
impl_nid = _make_id(stem, type_name)
add_node(impl_nid, type_name, node.start_point[0] + 1)
body = node.child_by_field_name("body")
@@ -682,7 +682,7 @@ def extract_rust(path: Path) -> dict:
if t == "use_declaration":
arg = node.child_by_field_name("argument")
if arg:
- raw = source[arg.start_byte:arg.end_byte].decode()
+ raw = source[arg.start_byte:arg.end_byte].decode("utf-8", errors="replace")
clean = raw.split("{")[0].rstrip(":").rstrip("*").rstrip(":")
module_name = clean.split("::")[-1].strip()
if module_name:
@@ -711,15 +711,15 @@ def extract_rust(path: Path) -> dict:
callee_name: str | None = None
if func_node:
if func_node.type == "identifier":
- callee_name = source[func_node.start_byte:func_node.end_byte].decode()
+ callee_name = source[func_node.start_byte:func_node.end_byte].decode("utf-8", errors="replace")
elif func_node.type == "field_expression":
field = func_node.child_by_field_name("field")
if field:
- callee_name = source[field.start_byte:field.end_byte].decode()
+ callee_name = source[field.start_byte:field.end_byte].decode("utf-8", errors="replace")
elif func_node.type == "scoped_identifier":
name = func_node.child_by_field_name("name")
if name:
- callee_name = source[name.start_byte:name.end_byte].decode()
+ callee_name = source[name.start_byte:name.end_byte].decode("utf-8", errors="replace")
if callee_name:
tgt_nid = label_to_nid.get(callee_name.lower())
if tgt_nid and tgt_nid != caller_nid:
@@ -830,12 +830,12 @@ def _resolve_cross_file_imports(
# Dig into relative_import → dotted_name → identifier
for sub in child.children:
if sub.type == "dotted_name":
- raw = source[sub.start_byte:sub.end_byte].decode()
+ raw = source[sub.start_byte:sub.end_byte].decode("utf-8", errors="replace")
target_stem = raw.split(".")[-1]
break
break
if child.type == "dotted_name" and target_stem is None:
- raw = source[child.start_byte:child.end_byte].decode()
+ raw = source[child.start_byte:child.end_byte].decode("utf-8", errors="replace")
target_stem = raw.split(".")[-1]
if not target_stem or target_stem not in stem_to_entities:
@@ -853,14 +853,14 @@ def _resolve_cross_file_imports(
continue
if child.type == "dotted_name":
imported_names.append(
- source[child.start_byte:child.end_byte].decode()
+ source[child.start_byte:child.end_byte].decode("utf-8", errors="replace")
)
elif child.type == "aliased_import":
# `import X as Y` — take the original name
name_node = child.child_by_field_name("name")
if name_node:
imported_names.append(
- source[name_node.start_byte:name_node.end_byte].decode()
+ source[name_node.start_byte:name_node.end_byte].decode("utf-8", errors="replace")
)
line = node.start_point[0] + 1
diff --git a/graphify/ingest.py b/graphify/ingest.py
index b5cba19a..3ac54d3f 100644
--- a/graphify/ingest.py
+++ b/graphify/ingest.py
@@ -3,11 +3,13 @@ from __future__ import annotations
import json
import re
import sys
-import urllib.request
+import urllib.error
import urllib.parse
from datetime import datetime, timezone
from pathlib import Path
+from graphify.security import safe_fetch, safe_fetch_text, validate_url
+
def _safe_filename(url: str, suffix: str) -> str:
"""Turn a URL into a safe filename."""
@@ -39,9 +41,7 @@ def _detect_url_type(url: str) -> str:
def _fetch_html(url: str) -> str:
- req = urllib.request.Request(url, headers={"User-Agent": "Mozilla/5.0 graphify/1.0"})
- with urllib.request.urlopen(req, timeout=15) as resp:
- return resp.read().decode("utf-8", errors="ignore")
+ return safe_fetch_text(url)
def _html_to_markdown(html: str, url: str) -> str:
@@ -175,9 +175,7 @@ def _download_binary(url: str, suffix: str, target_dir: Path) -> Path:
"""Download a binary file (PDF, image) directly."""
filename = _safe_filename(url, suffix)
out_path = target_dir / filename
- req = urllib.request.Request(url, headers={"User-Agent": "Mozilla/5.0 graphify/1.0"})
- with urllib.request.urlopen(req, timeout=30) as resp:
- out_path.write_bytes(resp.read())
+ out_path.write_bytes(safe_fetch(url))
return out_path
@@ -190,23 +188,31 @@ def ingest(url: str, target_dir: Path, author: str | None = None, contributor: s
target_dir.mkdir(parents=True, exist_ok=True)
url_type = _detect_url_type(url)
- if url_type == "pdf":
- out = _download_binary(url, ".pdf", target_dir)
- print(f"Downloaded PDF: {out.name}")
- return out
+ try:
+ validate_url(url)
+ except ValueError as exc:
+ raise ValueError(f"ingest: {exc}") from exc
- if url_type == "image":
- suffix = Path(urllib.parse.urlparse(url).path).suffix or ".jpg"
- out = _download_binary(url, suffix, target_dir)
- print(f"Downloaded image: {out.name}")
- return out
+ try:
+ if url_type == "pdf":
+ out = _download_binary(url, ".pdf", target_dir)
+ print(f"Downloaded PDF: {out.name}")
+ return out
- if url_type == "tweet":
- content, filename = _fetch_tweet(url, author, contributor)
- elif url_type == "arxiv":
- content, filename = _fetch_arxiv(url, author, contributor)
- else:
- content, filename = _fetch_webpage(url, author, contributor)
+ if url_type == "image":
+ suffix = Path(urllib.parse.urlparse(url).path).suffix or ".jpg"
+ out = _download_binary(url, suffix, target_dir)
+ print(f"Downloaded image: {out.name}")
+ return out
+
+ if url_type == "tweet":
+ content, filename = _fetch_tweet(url, author, contributor)
+ elif url_type == "arxiv":
+ content, filename = _fetch_arxiv(url, author, contributor)
+ else:
+ content, filename = _fetch_webpage(url, author, contributor)
+ except (urllib.error.HTTPError, urllib.error.URLError, OSError) as exc:
+ raise RuntimeError(f"ingest: failed to fetch {url!r}: {exc}") from exc
out_path = target_dir / filename
# Avoid overwriting — append counter if needed
@@ -245,7 +251,7 @@ def save_query_result(
"---",
f'type: "{query_type}"',
f'date: "{now.isoformat()}"',
- f'question: "{question.replace(chr(34), chr(39))}"',
+ f'question: "{re.sub(chr(10) + chr(13), " ", question).replace(chr(34), chr(39))}"',
'contributor: "graphify"',
]
if source_nodes:
diff --git a/graphify/security.py b/graphify/security.py
new file mode 100644
index 00000000..1e9ed132
--- /dev/null
+++ b/graphify/security.py
@@ -0,0 +1,166 @@
+# Security helpers — URL validation, safe fetch, path guards, label sanitisation
+from __future__ import annotations
+
+import html
+import re
+import urllib.error
+import urllib.parse
+import urllib.request
+from pathlib import Path
+
+_ALLOWED_SCHEMES = {"http", "https"}
+_MAX_FETCH_BYTES = 52_428_800 # 50 MB hard cap for binary downloads
+_MAX_TEXT_BYTES = 10_485_760 # 10 MB hard cap for HTML / text
+
+
+# ---------------------------------------------------------------------------
+# URL validation
+# ---------------------------------------------------------------------------
+
+def validate_url(url: str) -> str:
+ """Raise ValueError if *url* is not http or https.
+
+ Blocks file://, ftp://, data:, and any other scheme that could be used
+ for SSRF or local file access.
+ """
+ parsed = urllib.parse.urlparse(url)
+ if parsed.scheme.lower() not in _ALLOWED_SCHEMES:
+ raise ValueError(
+ f"Blocked URL scheme '{parsed.scheme}' — only http and https are allowed. "
+ f"Got: {url!r}"
+ )
+ return url
+
+
+class _NoFileRedirectHandler(urllib.request.HTTPRedirectHandler):
+ """Redirect handler that re-validates every redirect target.
+
+ Prevents open-redirect SSRF attacks where an http:// URL redirects
+ to file:// or an internal address.
+ """
+
+ def redirect_request(self, req, fp, code, msg, headers, newurl):
+ validate_url(newurl) # raises ValueError if scheme is wrong
+ return super().redirect_request(req, fp, code, msg, headers, newurl)
+
+
+def _build_opener() -> urllib.request.OpenerDirector:
+ return urllib.request.build_opener(_NoFileRedirectHandler)
+
+
+# ---------------------------------------------------------------------------
+# Safe fetch
+# ---------------------------------------------------------------------------
+
+def safe_fetch(url: str, max_bytes: int = _MAX_FETCH_BYTES, timeout: int = 30) -> bytes:
+ """Fetch *url* and return raw bytes.
+
+ Protections applied:
+ - URL scheme validated (http / https only)
+ - Redirects re-validated via _NoFileRedirectHandler
+ - Response body capped at *max_bytes* (streaming read)
+ - Non-2xx status raises urllib.error.HTTPError
+ - Network errors propagate as urllib.error.URLError / OSError
+
+ Raises:
+ ValueError — disallowed scheme or redirect target
+ urllib.error.HTTPError — non-2xx HTTP status
+ urllib.error.URLError — DNS / connection failure
+ OSError — size cap exceeded
+ """
+ validate_url(url)
+ opener = _build_opener()
+ req = urllib.request.Request(url, headers={"User-Agent": "Mozilla/5.0 graphify/1.0"})
+
+ with opener.open(req, timeout=timeout) as resp:
+ # urllib raises HTTPError for non-2xx when using urlopen directly;
+ # with a custom opener we check manually to be safe.
+ status = getattr(resp, "status", None) or getattr(resp, "code", None)
+ if status is not None and not (200 <= status < 300):
+ raise urllib.error.HTTPError(url, status, f"HTTP {status}", {}, None)
+
+ chunks: list[bytes] = []
+ total = 0
+ while True:
+ chunk = resp.read(65_536)
+ if not chunk:
+ break
+ total += len(chunk)
+ if total > max_bytes:
+ raise OSError(
+ f"Response from {url!r} exceeds size limit "
+ f"({max_bytes // 1_048_576} MB). Aborting download."
+ )
+ chunks.append(chunk)
+
+ return b"".join(chunks)
+
+
+def safe_fetch_text(url: str, max_bytes: int = _MAX_TEXT_BYTES, timeout: int = 15) -> str:
+ """Fetch *url* and return decoded text (UTF-8, replacing bad bytes).
+
+ Wraps safe_fetch with tighter defaults for HTML / text content.
+ """
+ raw = safe_fetch(url, max_bytes=max_bytes, timeout=timeout)
+ return raw.decode("utf-8", errors="replace")
+
+
+# ---------------------------------------------------------------------------
+# Path validation
+# ---------------------------------------------------------------------------
+
+def validate_graph_path(path: str | Path, base: Path | None = None) -> Path:
+ """Resolve *path* and verify it stays inside *base*.
+
+ *base* defaults to the `.graphify` directory relative to CWD.
+ Also requires the base directory to exist, so a caller cannot
+ trick graphify into reading files before any graph has been built.
+
+ Raises:
+ ValueError — path escapes base, or base does not exist
+ FileNotFoundError — resolved path does not exist
+ """
+ if base is None:
+ base = Path(".graphify").resolve()
+
+ base = base.resolve()
+ if not base.exists():
+ raise ValueError(
+ f"Graph base directory does not exist: {base}. "
+ "Run /graphify first to build the graph."
+ )
+
+ resolved = Path(path).resolve()
+ try:
+ resolved.relative_to(base)
+ except ValueError:
+ raise ValueError(
+ f"Path {path!r} escapes the allowed directory {base}. "
+ "Only paths inside .graphify/ are permitted."
+ )
+
+ if not resolved.exists():
+ raise FileNotFoundError(f"Graph file not found: {resolved}")
+
+ return resolved
+
+
+# ---------------------------------------------------------------------------
+# Label sanitisation (mirrors code-review-graph's _sanitize_name pattern)
+# ---------------------------------------------------------------------------
+
+_CONTROL_CHAR_RE = re.compile(r"[\x00-\x1f\x7f]")
+_MAX_LABEL_LEN = 256
+
+
+def sanitize_label(text: str) -> str:
+ """Strip control characters, cap length, then HTML-escape.
+
+ Applied to all node labels and edge titles before they are embedded
+ in pyvis HTML output or returned via the MCP server, preventing both
+ XSS and broken visualisations from malformed source identifiers.
+ """
+ text = _CONTROL_CHAR_RE.sub("", text)
+ if len(text) > _MAX_LABEL_LEN:
+ text = text[:_MAX_LABEL_LEN]
+ return html.escape(text)
diff --git a/graphify/serve.py b/graphify/serve.py
index 2f1ed6fd..d738e5ae 100644
--- a/graphify/serve.py
+++ b/graphify/serve.py
@@ -5,11 +5,20 @@ import sys
from pathlib import Path
import networkx as nx
from networkx.readwrite import json_graph
+from graphify.security import validate_graph_path, sanitize_label
def _load_graph(graph_path: str) -> nx.Graph:
- data = json.loads(Path(graph_path).read_text())
- return json_graph.node_link_graph(data, edges="links")
+ try:
+ safe = validate_graph_path(graph_path)
+ data = json.loads(safe.read_text())
+ return json_graph.node_link_graph(data, edges="links")
+ except (ValueError, FileNotFoundError) as exc:
+ print(f"error: {exc}", file=sys.stderr)
+ sys.exit(1)
+ except json.JSONDecodeError as exc:
+ print(f"error: graph.json is corrupted ({exc}). Re-run /graphify to rebuild.", file=sys.stderr)
+ sys.exit(1)
def _communities_from_graph(G: nx.Graph) -> dict[int, list[str]]:
@@ -71,12 +80,12 @@ def _subgraph_to_text(G: nx.Graph, nodes: set[str], edges: list[tuple], token_bu
lines = []
for nid in sorted(nodes, key=lambda n: G.degree(n), reverse=True):
d = G.nodes[nid]
- line = f"NODE {d.get('label', nid)} [src={d.get('source_file', '')} loc={d.get('source_location', '')} community={d.get('community', '')}]"
+ line = f"NODE {sanitize_label(d.get('label', nid))} [src={d.get('source_file', '')} loc={d.get('source_location', '')} community={d.get('community', '')}]"
lines.append(line)
for u, v in edges:
if u in nodes and v in nodes:
d = G.edges[u, v]
- line = f"EDGE {G.nodes[u].get('label', u)} --{d.get('relation', '')} [{d.get('confidence', '')}]--> {G.nodes[v].get('label', v)}"
+ line = f"EDGE {sanitize_label(G.nodes[u].get('label', u))} --{d.get('relation', '')} [{d.get('confidence', '')}]--> {sanitize_label(G.nodes[v].get('label', v))}"
lines.append(line)
output = "\n".join(lines)
if len(output) > char_budget:
diff --git a/tests/test_security.py b/tests/test_security.py
new file mode 100644
index 00000000..f3036ca1
--- /dev/null
+++ b/tests/test_security.py
@@ -0,0 +1,187 @@
+"""Tests for graphify/security.py — URL validation, safe fetch, path guards, label sanitisation."""
+from __future__ import annotations
+
+import json
+import urllib.error
+from pathlib import Path
+from unittest.mock import MagicMock, patch
+
+import pytest
+
+from graphify.security import (
+ sanitize_label,
+ safe_fetch,
+ safe_fetch_text,
+ validate_graph_path,
+ validate_url,
+ _MAX_FETCH_BYTES,
+ _MAX_TEXT_BYTES,
+)
+
+
+# ---------------------------------------------------------------------------
+# validate_url
+# ---------------------------------------------------------------------------
+
+def test_validate_url_accepts_http():
+ assert validate_url("http://example.com/page") == "http://example.com/page"
+
+def test_validate_url_accepts_https():
+ assert validate_url("https://arxiv.org/abs/1706.03762") == "https://arxiv.org/abs/1706.03762"
+
+def test_validate_url_rejects_file():
+ with pytest.raises(ValueError, match="file"):
+ validate_url("file:///etc/passwd")
+
+def test_validate_url_rejects_ftp():
+ with pytest.raises(ValueError, match="ftp"):
+ validate_url("ftp://files.example.com/data.zip")
+
+def test_validate_url_rejects_data():
+ with pytest.raises(ValueError, match="data"):
+ validate_url("data:text/html,")
+
+def test_validate_url_rejects_empty_scheme():
+ with pytest.raises(ValueError):
+ validate_url("//no-scheme.example.com")
+
+
+# ---------------------------------------------------------------------------
+# safe_fetch — scheme and redirect guards (mocked network)
+# ---------------------------------------------------------------------------
+
+def _make_mock_response(content: bytes, status: int = 200):
+ mock = MagicMock()
+ mock.__enter__ = lambda s: s
+ mock.__exit__ = MagicMock(return_value=False)
+ mock.status = status
+ mock.code = status
+ chunks = [content[i:i+65536] for i in range(0, len(content), 65536)] + [b""]
+ mock.read.side_effect = chunks
+ return mock
+
+
+def test_safe_fetch_rejects_file_url():
+ with pytest.raises(ValueError, match="file"):
+ safe_fetch("file:///etc/passwd")
+
+def test_safe_fetch_rejects_ftp_url():
+ with pytest.raises(ValueError, match="ftp"):
+ safe_fetch("ftp://example.com/file.zip")
+
+def test_safe_fetch_returns_bytes(tmp_path):
+ mock_resp = _make_mock_response(b"hello world")
+ with patch("graphify.security._build_opener") as mock_opener_fn:
+ mock_opener = MagicMock()
+ mock_opener.open.return_value = mock_resp
+ mock_opener_fn.return_value = mock_opener
+ result = safe_fetch("https://example.com/")
+ assert result == b"hello world"
+
+def test_safe_fetch_raises_on_non_2xx():
+ mock_resp = _make_mock_response(b"Not Found", status=404)
+ with patch("graphify.security._build_opener") as mock_opener_fn:
+ mock_opener = MagicMock()
+ mock_opener.open.return_value = mock_resp
+ mock_opener_fn.return_value = mock_opener
+ with pytest.raises(urllib.error.HTTPError):
+ safe_fetch("https://example.com/missing")
+
+def test_safe_fetch_raises_on_size_exceeded():
+ # Build a response larger than max_bytes
+ big_chunk = b"x" * 65_537
+ mock_resp = MagicMock()
+ mock_resp.__enter__ = lambda s: s
+ mock_resp.__exit__ = MagicMock(return_value=False)
+ mock_resp.status = 200
+ mock_resp.code = 200
+ # Return the chunk twice so total > max_bytes=65536
+ mock_resp.read.side_effect = [big_chunk, big_chunk, b""]
+
+ with patch("graphify.security._build_opener") as mock_opener_fn:
+ mock_opener = MagicMock()
+ mock_opener.open.return_value = mock_resp
+ mock_opener_fn.return_value = mock_opener
+ with pytest.raises(OSError, match="size limit"):
+ safe_fetch("https://example.com/huge", max_bytes=65_536)
+
+
+# ---------------------------------------------------------------------------
+# safe_fetch_text
+# ---------------------------------------------------------------------------
+
+def test_safe_fetch_text_decodes_utf8():
+ content = "héllo wörld".encode("utf-8")
+ mock_resp = _make_mock_response(content)
+ with patch("graphify.security._build_opener") as mock_opener_fn:
+ mock_opener = MagicMock()
+ mock_opener.open.return_value = mock_resp
+ mock_opener_fn.return_value = mock_opener
+ result = safe_fetch_text("https://example.com/")
+ assert result == "héllo wörld"
+
+def test_safe_fetch_text_replaces_bad_bytes():
+ bad = b"hello \xff world"
+ mock_resp = _make_mock_response(bad)
+ with patch("graphify.security._build_opener") as mock_opener_fn:
+ mock_opener = MagicMock()
+ mock_opener.open.return_value = mock_resp
+ mock_opener_fn.return_value = mock_opener
+ result = safe_fetch_text("https://example.com/")
+ assert "hello" in result
+ assert "world" in result
+ assert "\xff" not in result
+
+
+# ---------------------------------------------------------------------------
+# validate_graph_path
+# ---------------------------------------------------------------------------
+
+def test_validate_graph_path_allows_inside_base(tmp_path):
+ base = tmp_path / ".graphify"
+ base.mkdir()
+ graph = base / "graph.json"
+ graph.write_text("{}")
+ result = validate_graph_path(str(graph), base=base)
+ assert result == graph.resolve()
+
+def test_validate_graph_path_blocks_traversal(tmp_path):
+ base = tmp_path / ".graphify"
+ base.mkdir()
+ evil = tmp_path / ".graphify" / ".." / "etc_passwd"
+ with pytest.raises(ValueError, match="escapes"):
+ validate_graph_path(str(evil), base=base)
+
+def test_validate_graph_path_requires_base_exists(tmp_path):
+ base = tmp_path / ".graphify" # not created
+ with pytest.raises(ValueError, match="does not exist"):
+ validate_graph_path(str(base / "graph.json"), base=base)
+
+def test_validate_graph_path_raises_if_file_missing(tmp_path):
+ base = tmp_path / ".graphify"
+ base.mkdir()
+ with pytest.raises(FileNotFoundError):
+ validate_graph_path(str(base / "missing.json"), base=base)
+
+
+# ---------------------------------------------------------------------------
+# sanitize_label
+# ---------------------------------------------------------------------------
+
+def test_sanitize_label_escapes_html():
+ assert "<script>" in sanitize_label("