#1118 — prune stale AST nodes on full re-extraction (#1116) Stamps every AST-extracted node with _origin="ast" in extract(). On a full rebuild _rebuild_code drops any AST-marked node absent from the fresh output even when its source file survives, fixing stale symbols. Backward-compat: marker-less nodes from pre-1118 graphs survive one cycle then self-heal. #1110 — stop reading images and PDFs as garbage in headless extract Images route through per-backend vision payloads (base64/data-URI/bytes for claude/openai/bedrock); non-vision backends get _strip_pixels for graceful degradation. PDFs reuse pypdf. 5MB cap, 20-image chunk limit. #1159 — Salesforce Apex extractor (.cls, .trigger) Regex-based extractor: classes, interfaces, enums, methods, triggers, SOQL/DML edges. No new dependency. Dispatched as .cls and .trigger. #1107 — Azure OpenAI Service backend (--backend azure) Uses AzureOpenAI SDK client (from existing openai package). Auto-detects when AZURE_OPENAI_API_KEY + AZURE_OPENAI_ENDPOINT both set. Uses max_completion_tokens (not deprecated max_tokens). #1103 — live PostgreSQL introspection (--postgres DSN) graphify extract --postgres "postgresql://..." introspects tables, views, routines, and FK relations via information_schema (SERIALIZABLE READ ONLY). Credentials sanitized on error. New graphify[postgres] extra (psycopg3). Union-resolved llm.py conflict: Azure functions + bedrock images= param. Fixed test_image_vision.py mock to accept timeout= kwarg (our #1112). Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Sonnet 4.6
parent
7b4c8df6e9
commit
7467c1b6a4
+52
-17
@@ -2042,6 +2042,9 @@ def main() -> None:
|
||||
print(" --out DIR output dir (default: <path>); writes <DIR>/graphify-out/")
|
||||
print(" --google-workspace export .gdoc/.gsheet/.gslides shortcuts via gws before extraction")
|
||||
print(" --no-cluster skip clustering, write raw extraction only")
|
||||
print(" --postgres DSN extract schema from a live PostgreSQL database")
|
||||
print(" maps tables, views, functions + FK relationships;")
|
||||
print(" column-level detail is not represented in the graph")
|
||||
print(" --global also merge the resulting graph into the global graph")
|
||||
print(" --as <tag> repo tag for --global (default: target directory name)")
|
||||
print(" global add <graph.json> add/update a project graph in the global graph (~/.graphify/global-graph.json)")
|
||||
@@ -3708,20 +3711,26 @@ def main() -> None:
|
||||
"Usage: graphify extract <path> [--backend gemini|kimi|claude|openai|deepseek|ollama] "
|
||||
"[--model M] [--mode deep] [--out DIR] [--google-workspace] [--no-cluster] "
|
||||
"[--max-workers N] [--token-budget N] [--max-concurrency N] "
|
||||
"[--api-timeout S]",
|
||||
"[--api-timeout S] [--postgres DSN]",
|
||||
file=sys.stderr,
|
||||
)
|
||||
sys.exit(1)
|
||||
|
||||
target = Path(sys.argv[2]).resolve()
|
||||
if not target.exists():
|
||||
print(f"error: path not found: {target}", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
has_path = True
|
||||
if sys.argv[2].startswith("-"):
|
||||
has_path = False
|
||||
target = Path(".").resolve()
|
||||
else:
|
||||
target = Path(sys.argv[2]).resolve()
|
||||
if not target.exists():
|
||||
print(f"error: path not found: {target}", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
backend: str | None = None
|
||||
model: str | None = None
|
||||
extract_mode: str | None = None
|
||||
out_dir: Path | None = None
|
||||
cli_postgres_dsn: str | None = None
|
||||
no_cluster = False
|
||||
dedup_llm = False
|
||||
google_workspace = False
|
||||
@@ -3759,7 +3768,7 @@ def main() -> None:
|
||||
sys.exit(2)
|
||||
return v
|
||||
|
||||
args = sys.argv[3:]
|
||||
args = sys.argv[3:] if has_path else sys.argv[2:]
|
||||
i = 0
|
||||
while i < len(args):
|
||||
a = args[i]
|
||||
@@ -3817,9 +3826,17 @@ def main() -> None:
|
||||
cli_excludes.append(args[i + 1]); i += 2
|
||||
elif a.startswith("--exclude="):
|
||||
cli_excludes.append(a.split("=", 1)[1]); i += 1
|
||||
elif a == "--postgres" and i + 1 < len(args):
|
||||
cli_postgres_dsn = args[i + 1]; i += 2
|
||||
elif a.startswith("--postgres="):
|
||||
cli_postgres_dsn = a.split("=", 1)[1]; i += 1
|
||||
else:
|
||||
i += 1
|
||||
|
||||
if not has_path and cli_postgres_dsn is None:
|
||||
print("error: must specify a path to scan or a --postgres DSN", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
_VALID_MODES = {"deep"}
|
||||
if extract_mode is not None and extract_mode not in _VALID_MODES:
|
||||
print(
|
||||
@@ -3853,9 +3870,17 @@ def main() -> None:
|
||||
)
|
||||
manifest_path = graphify_out / "manifest.json"
|
||||
existing_graph_path = graphify_out / "graph.json"
|
||||
incremental_mode = manifest_path.exists() and existing_graph_path.exists()
|
||||
incremental_mode = manifest_path.exists() and existing_graph_path.exists() if has_path else False
|
||||
|
||||
if incremental_mode:
|
||||
if not has_path:
|
||||
code_files = []
|
||||
doc_files = []
|
||||
paper_files = []
|
||||
image_files = []
|
||||
deleted_files = []
|
||||
unchanged_total = 0
|
||||
files_by_type = {}
|
||||
elif incremental_mode:
|
||||
print(f"[graphify extract] incremental scan of {target}")
|
||||
detection = _detect_incremental(
|
||||
target,
|
||||
@@ -3863,12 +3888,7 @@ def main() -> None:
|
||||
google_workspace=google_workspace or None,
|
||||
extra_excludes=cli_excludes or None,
|
||||
)
|
||||
else:
|
||||
print(f"[graphify extract] scanning {target}")
|
||||
detection = _detect(target, google_workspace=google_workspace or None, extra_excludes=cli_excludes or None)
|
||||
|
||||
files_by_type = detection.get("files", {})
|
||||
if incremental_mode:
|
||||
files_by_type = detection.get("files", {})
|
||||
new_by_type = detection.get("new_files", {})
|
||||
code_files = [Path(p) for p in new_by_type.get("code", [])]
|
||||
doc_files = [Path(p) for p in new_by_type.get("document", [])]
|
||||
@@ -3877,6 +3897,9 @@ def main() -> None:
|
||||
deleted_files = list(detection.get("deleted_files", []))
|
||||
unchanged_total = sum(len(v) for v in detection.get("unchanged_files", {}).values())
|
||||
else:
|
||||
print(f"[graphify extract] scanning {target}")
|
||||
detection = _detect(target, google_workspace=google_workspace or None, extra_excludes=cli_excludes or None)
|
||||
files_by_type = detection.get("files", {})
|
||||
code_files = [Path(p) for p in files_by_type.get("code", [])]
|
||||
doc_files = [Path(p) for p in files_by_type.get("document", [])]
|
||||
paper_files = [Path(p) for p in files_by_type.get("paper", [])]
|
||||
@@ -4094,13 +4117,25 @@ def main() -> None:
|
||||
sem_result["input_tokens"] += fresh.get("input_tokens", 0)
|
||||
sem_result["output_tokens"] += fresh.get("output_tokens", 0)
|
||||
|
||||
# Merge AST + semantic. Order matters for deduplication: passing AST
|
||||
pg_result: dict = {"nodes": [], "edges": []}
|
||||
if cli_postgres_dsn is not None:
|
||||
from graphify.pg_introspect import introspect_postgres
|
||||
print(f"[graphify extract] introspecting PostgreSQL schema...")
|
||||
try:
|
||||
pg_result = introspect_postgres(cli_postgres_dsn)
|
||||
except (ConnectionError, ImportError) as exc:
|
||||
print(f"error: {exc}", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
print(f"[graphify extract] PostgreSQL: {len(pg_result['nodes'])} nodes, "
|
||||
f"{len(pg_result['edges'])} edges")
|
||||
|
||||
# Merge AST + semantic + pg_result. Order matters for deduplication: passing AST
|
||||
# first means semantic node attributes win on collision (richer labels
|
||||
# for symbols also referenced in docs). Hyperedges only come from the
|
||||
# semantic side.
|
||||
merged: dict = {
|
||||
"nodes": list(ast_result.get("nodes", [])) + list(sem_result.get("nodes", [])),
|
||||
"edges": list(ast_result.get("edges", [])) + list(sem_result.get("edges", [])),
|
||||
"nodes": list(ast_result.get("nodes", [])) + list(sem_result.get("nodes", [])) + list(pg_result.get("nodes", [])),
|
||||
"edges": list(ast_result.get("edges", [])) + list(sem_result.get("edges", [])) + list(pg_result.get("edges", [])),
|
||||
"hyperedges": list(sem_result.get("hyperedges", [])),
|
||||
"input_tokens": ast_result.get("input_tokens", 0) + sem_result.get("input_tokens", 0),
|
||||
"output_tokens": ast_result.get("output_tokens", 0) + sem_result.get("output_tokens", 0),
|
||||
|
||||
+1
-1
@@ -25,7 +25,7 @@ class FileType(str, Enum):
|
||||
|
||||
_MANIFEST_PATH = "graphify-out/manifest.json"
|
||||
|
||||
CODE_EXTENSIONS = {'.py', '.ts', '.tsx', '.js', '.jsx', '.mjs', '.ejs', '.ets', '.go', '.rs', '.java', '.groovy', '.gradle', '.cpp', '.cc', '.cxx', '.c', '.h', '.hpp', '.rb', '.swift', '.kt', '.kts', '.cs', '.scala', '.php', '.lua', '.luau', '.toc', '.zig', '.ps1', '.ex', '.exs', '.m', '.mm', '.jl', '.vue', '.svelte', '.astro', '.dart', '.v', '.sv', '.svh', '.sql', '.r', '.f', '.F', '.f90', '.F90', '.f95', '.F95', '.f03', '.F03', '.f08', '.F08', '.pas', '.pp', '.dpr', '.dpk', '.lpr', '.inc', '.dfm', '.lfm', '.lpk', '.sh', '.bash', '.json', '.tf', '.tfvars', '.hcl', '.dm', '.dme', '.dmi', '.dmm', '.dmf', '.sln', '.csproj', '.fsproj', '.vbproj', '.razor', '.cshtml'}
|
||||
CODE_EXTENSIONS = {'.py', '.ts', '.tsx', '.js', '.jsx', '.mjs', '.ejs', '.ets', '.go', '.rs', '.java', '.groovy', '.gradle', '.cpp', '.cc', '.cxx', '.c', '.h', '.hpp', '.rb', '.swift', '.kt', '.kts', '.cs', '.scala', '.php', '.lua', '.luau', '.toc', '.zig', '.ps1', '.ex', '.exs', '.m', '.mm', '.jl', '.vue', '.svelte', '.astro', '.dart', '.v', '.sv', '.svh', '.sql', '.r', '.f', '.F', '.f90', '.F90', '.f95', '.F95', '.f03', '.F03', '.f08', '.F08', '.pas', '.pp', '.dpr', '.dpk', '.lpr', '.inc', '.dfm', '.lfm', '.lpk', '.sh', '.bash', '.json', '.tf', '.tfvars', '.hcl', '.dm', '.dme', '.dmi', '.dmm', '.dmf', '.sln', '.csproj', '.fsproj', '.vbproj', '.razor', '.cshtml', '.cls', '.trigger'}
|
||||
DOC_EXTENSIONS = {'.md', '.mdx', '.qmd', '.txt', '.rst', '.html', '.yaml', '.yml'}
|
||||
PAPER_EXTENSIONS = {'.pdf'}
|
||||
IMAGE_EXTENSIONS = {'.png', '.jpg', '.jpeg', '.gif', '.webp', '.svg'}
|
||||
|
||||
+17
-2
@@ -1274,7 +1274,10 @@ def push_to_neo4j(
|
||||
|
||||
with driver.session() as session:
|
||||
for node_id, data in G.nodes(data=True):
|
||||
props = {k: v for k, v in data.items() if isinstance(v, (str, int, float, bool))}
|
||||
props = {
|
||||
k: v for k, v in data.items()
|
||||
if isinstance(v, (str, int, float, bool)) and not k.startswith("_")
|
||||
}
|
||||
props["id"] = node_id
|
||||
cid = node_community.get(node_id)
|
||||
if cid is not None:
|
||||
@@ -1289,7 +1292,10 @@ def push_to_neo4j(
|
||||
|
||||
for u, v, data in G.edges(data=True):
|
||||
rel = _safe_rel(data.get("relation", "RELATED_TO"))
|
||||
props = {k: v for k, v in data.items() if isinstance(v, (str, int, float, bool))}
|
||||
props = {
|
||||
k: v for k, v in data.items()
|
||||
if isinstance(v, (str, int, float, bool)) and not k.startswith("_")
|
||||
}
|
||||
session.run(
|
||||
f"MATCH (a {{id: $src}}), (b {{id: $tgt}}) "
|
||||
f"MERGE (a)-[r:{rel}]->(b) SET r += $props",
|
||||
@@ -1317,6 +1323,15 @@ def to_graphml(
|
||||
node_community = _node_community_map(communities)
|
||||
for node_id in H.nodes():
|
||||
H.nodes[node_id]["community"] = node_community.get(node_id, -1)
|
||||
# Drop internal markers (e.g. the AST-provenance "_origin" tag, #1116, and
|
||||
# the "_src"/"_tgt" direction markers) — they are persistence/runtime details,
|
||||
# not graph data, and should not leak into the exported file.
|
||||
for _, attrs in H.nodes(data=True):
|
||||
for k in [k for k in attrs if k.startswith("_")]:
|
||||
del attrs[k]
|
||||
for _, _, attrs in H.edges(data=True):
|
||||
for k in [k for k in attrs if k.startswith("_")]:
|
||||
del attrs[k]
|
||||
nx.write_graphml(H, output_path)
|
||||
|
||||
|
||||
|
||||
+215
-2
@@ -4041,6 +4041,205 @@ def extract_csharp(path: Path) -> dict:
|
||||
return _extract_generic(path, _CSHARP_CONFIG)
|
||||
|
||||
|
||||
def extract_apex(path: Path) -> dict:
|
||||
"""Extract classes, interfaces, enums, methods, and Salesforce constructs from
|
||||
Apex .cls and .trigger files using regex (no tree-sitter grammar on PyPI)."""
|
||||
import re as _re
|
||||
try:
|
||||
source = path.read_text(encoding="utf-8", errors="replace")
|
||||
except OSError:
|
||||
return {"nodes": [], "edges": []}
|
||||
|
||||
str_path = str(path)
|
||||
stem = _file_stem(path)
|
||||
file_nid = _make_id(str_path)
|
||||
|
||||
nodes: list[dict] = []
|
||||
edges: list[dict] = []
|
||||
seen_ids: set[str] = set()
|
||||
|
||||
def add_node(nid: str, label: str, line: int) -> None:
|
||||
if nid not in seen_ids:
|
||||
seen_ids.add(nid)
|
||||
nodes.append({
|
||||
"id": nid,
|
||||
"label": label,
|
||||
"file_type": "code",
|
||||
"source_file": str_path,
|
||||
"source_location": f"L{line}",
|
||||
})
|
||||
|
||||
def add_edge(src: str, tgt: str, relation: str, line: int,
|
||||
confidence: str = "EXTRACTED") -> None:
|
||||
edges.append({
|
||||
"source": src,
|
||||
"target": tgt,
|
||||
"relation": relation,
|
||||
"confidence": confidence,
|
||||
"source_file": str_path,
|
||||
"source_location": f"L{line}",
|
||||
"weight": 1.0,
|
||||
})
|
||||
|
||||
add_node(file_nid, path.name, 1)
|
||||
|
||||
lines = source.splitlines()
|
||||
|
||||
_ACCESS = r"(?:public|private|protected|global|webService)?"
|
||||
_SHARING = r"(?:\s+(?:with|without|inherited)\s+sharing)?"
|
||||
_MOD = r"(?:\s+(?:abstract|virtual|override|static|final|transient|testMethod))?"
|
||||
_ANNOTATION = r"(?:\s*@\w+(?:\s*\([^)]*\))?\s*)*"
|
||||
|
||||
cls_re = _re.compile(
|
||||
rf"^{_ANNOTATION}\s*{_ACCESS}{_SHARING}{_MOD}\s*class\s+(\w+)"
|
||||
rf"(?:\s+extends\s+(\w+))?(?:\s+implements\s+([\w,\s]+))?\s*\{{?",
|
||||
_re.IGNORECASE,
|
||||
)
|
||||
iface_re = _re.compile(
|
||||
rf"^{_ANNOTATION}\s*{_ACCESS}{_SHARING}{_MOD}\s*interface\s+(\w+)"
|
||||
rf"(?:\s+extends\s+([\w,\s]+))?\s*\{{?",
|
||||
_re.IGNORECASE,
|
||||
)
|
||||
enum_re = _re.compile(
|
||||
rf"^{_ANNOTATION}\s*{_ACCESS}{_SHARING}{_MOD}\s*enum\s+(\w+)\s*\{{?",
|
||||
_re.IGNORECASE,
|
||||
)
|
||||
trigger_re = _re.compile(
|
||||
r"^\s*trigger\s+(\w+)\s+on\s+(\w+)\s*\(",
|
||||
_re.IGNORECASE,
|
||||
)
|
||||
method_re = _re.compile(
|
||||
rf"^{_ANNOTATION}\s*{_ACCESS}{_MOD}\s*(?:static\s+)?[\w<>\[\]]+\s+(\w+)\s*\([^)]*\)\s*(?:throws\s+\w+\s*)?\{{?",
|
||||
_re.IGNORECASE,
|
||||
)
|
||||
annotation_re = _re.compile(r"@(\w+)", _re.IGNORECASE)
|
||||
soql_re = _re.compile(r"\[\s*SELECT\b[^\]]+FROM\s+(\w+)", _re.IGNORECASE)
|
||||
dml_re = _re.compile(r"\b(insert|update|delete|upsert|merge|undelete)\s+\w", _re.IGNORECASE)
|
||||
|
||||
_CONTROL_FLOW = frozenset({
|
||||
"if", "else", "for", "while", "do", "switch", "try", "catch",
|
||||
"finally", "return", "throw", "new", "void", "null",
|
||||
"true", "false", "this", "super", "class", "interface", "enum",
|
||||
"trigger", "on",
|
||||
})
|
||||
|
||||
current_class_nid: str | None = None
|
||||
pending_annotations: list[str] = []
|
||||
|
||||
for lineno, line_text in enumerate(lines, start=1):
|
||||
stripped = line_text.strip()
|
||||
|
||||
if stripped.startswith("@"):
|
||||
for m in annotation_re.finditer(stripped):
|
||||
pending_annotations.append(m.group(1).lower())
|
||||
continue
|
||||
|
||||
tm = trigger_re.match(stripped)
|
||||
if tm:
|
||||
trig_name, sobject = tm.group(1), tm.group(2)
|
||||
trig_nid = _make_id(stem, trig_name)
|
||||
add_node(trig_nid, trig_name, lineno)
|
||||
add_edge(file_nid, trig_nid, "contains", lineno)
|
||||
sob_nid = _make_id(sobject)
|
||||
if sob_nid not in seen_ids:
|
||||
add_node(sob_nid, sobject, lineno)
|
||||
add_edge(trig_nid, sob_nid, "uses", lineno, confidence="INFERRED")
|
||||
current_class_nid = trig_nid
|
||||
pending_annotations = []
|
||||
continue
|
||||
|
||||
cm = cls_re.match(stripped)
|
||||
if cm:
|
||||
class_name = cm.group(1)
|
||||
if class_name.lower() in _CONTROL_FLOW:
|
||||
pending_annotations = []
|
||||
continue
|
||||
class_nid = _make_id(stem, class_name)
|
||||
add_node(class_nid, class_name, lineno)
|
||||
add_edge(file_nid, class_nid, "contains", lineno)
|
||||
if cm.group(2):
|
||||
base = cm.group(2).strip()
|
||||
base_nid = _make_id(stem, base)
|
||||
if base_nid not in seen_ids:
|
||||
base_nid = _make_id(base)
|
||||
if base_nid not in seen_ids:
|
||||
add_node(base_nid, base, lineno)
|
||||
add_edge(class_nid, base_nid, "extends", lineno, confidence="INFERRED")
|
||||
if cm.group(3):
|
||||
for iface in cm.group(3).split(","):
|
||||
iface = iface.strip()
|
||||
if iface:
|
||||
iface_nid = _make_id(stem, iface)
|
||||
if iface_nid not in seen_ids:
|
||||
iface_nid = _make_id(iface)
|
||||
if iface_nid not in seen_ids:
|
||||
add_node(iface_nid, iface, lineno)
|
||||
add_edge(class_nid, iface_nid, "implements", lineno, confidence="INFERRED")
|
||||
current_class_nid = class_nid
|
||||
pending_annotations = []
|
||||
continue
|
||||
|
||||
im = iface_re.match(stripped)
|
||||
if im:
|
||||
iface_name = im.group(1)
|
||||
if iface_name.lower() in _CONTROL_FLOW:
|
||||
pending_annotations = []
|
||||
continue
|
||||
iface_nid = _make_id(stem, iface_name)
|
||||
add_node(iface_nid, iface_name, lineno)
|
||||
add_edge(file_nid if current_class_nid is None else current_class_nid,
|
||||
iface_nid, "contains", lineno)
|
||||
pending_annotations = []
|
||||
continue
|
||||
|
||||
em = enum_re.match(stripped)
|
||||
if em:
|
||||
enum_name = em.group(1)
|
||||
if enum_name.lower() in _CONTROL_FLOW:
|
||||
pending_annotations = []
|
||||
continue
|
||||
enum_nid = _make_id(stem, enum_name)
|
||||
add_node(enum_nid, enum_name, lineno)
|
||||
add_edge(file_nid if current_class_nid is None else current_class_nid,
|
||||
enum_nid, "contains", lineno)
|
||||
pending_annotations = []
|
||||
continue
|
||||
|
||||
if current_class_nid is not None:
|
||||
mm = method_re.match(stripped)
|
||||
if mm:
|
||||
method_name = mm.group(1)
|
||||
if method_name.lower() not in _CONTROL_FLOW:
|
||||
method_nid = _make_id(current_class_nid, method_name)
|
||||
method_label = f".{method_name}()"
|
||||
add_node(method_nid, method_label, lineno)
|
||||
add_edge(current_class_nid, method_nid, "method", lineno)
|
||||
if "auraenabled" in pending_annotations or "invocablemethod" in pending_annotations:
|
||||
add_edge(file_nid, method_nid, "contains", lineno, confidence="INFERRED")
|
||||
pending_annotations = []
|
||||
continue
|
||||
|
||||
pending_annotations = []
|
||||
|
||||
for sm in soql_re.finditer(line_text):
|
||||
sobject = sm.group(1)
|
||||
sob_nid = _make_id(sobject)
|
||||
if sob_nid not in seen_ids:
|
||||
add_node(sob_nid, sobject, lineno)
|
||||
src = current_class_nid or file_nid
|
||||
add_edge(src, sob_nid, "uses", lineno, confidence="INFERRED")
|
||||
|
||||
for dm in dml_re.finditer(line_text):
|
||||
dml_op = dm.group(1).lower()
|
||||
dml_nid = _make_id(f"dml_{dml_op}")
|
||||
if dml_nid not in seen_ids:
|
||||
add_node(dml_nid, dml_op, lineno)
|
||||
src = current_class_nid or file_nid
|
||||
add_edge(src, dml_nid, "uses", lineno, confidence="INFERRED")
|
||||
|
||||
return {"nodes": nodes, "edges": edges}
|
||||
|
||||
|
||||
def extract_kotlin(path: Path) -> dict:
|
||||
"""Extract classes, objects, functions, and imports from a .kt/.kts file."""
|
||||
return _extract_generic(path, _KOTLIN_CONFIG)
|
||||
@@ -4728,7 +4927,7 @@ def extract_verilog(path: Path) -> dict:
|
||||
return {"nodes": nodes, "edges": edges}
|
||||
|
||||
|
||||
def extract_sql(path: Path) -> dict:
|
||||
def extract_sql(path: Path, content: str | bytes | None = None) -> dict:
|
||||
"""Extract tables, views, functions, and relationships from .sql files via tree-sitter."""
|
||||
try:
|
||||
import tree_sitter_sql as tssql
|
||||
@@ -4739,12 +4938,17 @@ def extract_sql(path: Path) -> dict:
|
||||
try:
|
||||
language = Language(tssql.language())
|
||||
parser = Parser(language)
|
||||
source = path.read_bytes()
|
||||
source = (
|
||||
content.encode("utf-8") if isinstance(content, str)
|
||||
else content if content is not None
|
||||
else path.read_bytes()
|
||||
)
|
||||
tree = parser.parse(source)
|
||||
root = tree.root_node
|
||||
except Exception as e:
|
||||
return {"nodes": [], "edges": [], "error": str(e)}
|
||||
|
||||
|
||||
stem = _file_stem(path)
|
||||
str_path = str(path)
|
||||
file_nid = _make_id(str_path)
|
||||
@@ -10825,6 +11029,8 @@ _DISPATCH: dict[str, Any] = {
|
||||
".vbproj": extract_csproj,
|
||||
".razor": extract_razor,
|
||||
".cshtml": extract_razor,
|
||||
".cls": extract_apex,
|
||||
".trigger": extract_apex,
|
||||
}
|
||||
|
||||
|
||||
@@ -11298,6 +11504,13 @@ def extract(
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
# Tag AST provenance so the incremental watch rebuild can distinguish
|
||||
# AST-extracted nodes from semantic/LLM nodes. On a full re-extraction
|
||||
# the watcher drops any AST-marked node missing from the fresh output
|
||||
# even when its source file still exists (#1116).
|
||||
for n in all_nodes:
|
||||
n["_origin"] = "ast"
|
||||
|
||||
return {
|
||||
"nodes": all_nodes,
|
||||
"edges": all_edges,
|
||||
|
||||
+408
-15
@@ -5,6 +5,7 @@
|
||||
# this module provides a direct API path for non-Claude-Code environments.
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
@@ -12,6 +13,7 @@ import sys
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
from dataclasses import dataclass, replace
|
||||
from pathlib import Path
|
||||
|
||||
# `_read_files` truncates each file at this many characters before joining into
|
||||
@@ -53,11 +55,15 @@ BACKENDS: dict[str, dict] = {
|
||||
"pricing": {"input": 3.0, "output": 15.0}, # USD per 1M tokens
|
||||
"temperature": 0,
|
||||
"max_tokens": 16384,
|
||||
"vision": True,
|
||||
},
|
||||
"kimi": {
|
||||
"base_url": "https://api.moonshot.ai/v1",
|
||||
"default_model": "kimi-k2.6",
|
||||
"env_key": "MOONSHOT_API_KEY",
|
||||
# kimi-k2.6 is natively multimodal (MoonViT) and accepts the same
|
||||
# OpenAI image_url data-URI block via Moonshot's compat endpoint.
|
||||
"vision": True,
|
||||
"pricing": {"input": 0.74, "output": 4.66}, # USD per 1M tokens
|
||||
"temperature": None, # kimi-k2.6 enforces its own fixed temperature; sending any value raises 400
|
||||
"max_tokens": 16384,
|
||||
@@ -79,6 +85,7 @@ BACKENDS: dict[str, dict] = {
|
||||
"temperature": 0,
|
||||
"reasoning_effort": "low",
|
||||
"max_completion_tokens": 16384,
|
||||
"vision": True,
|
||||
},
|
||||
"openai": {
|
||||
"base_url": "https://api.openai.com/v1",
|
||||
@@ -87,6 +94,7 @@ BACKENDS: dict[str, dict] = {
|
||||
"model_env_key": "GRAPHIFY_OPENAI_MODEL",
|
||||
"pricing": {"input": 0.40, "output": 1.60}, # USD per 1M tokens
|
||||
"temperature": 0,
|
||||
"vision": True,
|
||||
},
|
||||
"deepseek": {
|
||||
"base_url": "https://api.deepseek.com",
|
||||
@@ -99,12 +107,28 @@ BACKENDS: dict[str, dict] = {
|
||||
"temperature": 0,
|
||||
"max_tokens": 16384,
|
||||
},
|
||||
"azure": {
|
||||
# Azure OpenAI Service — uses AzureOpenAI SDK client, not the standard
|
||||
# OpenAI client, so it has its own call path (_call_azure).
|
||||
# Required env vars: AZURE_OPENAI_API_KEY, AZURE_OPENAI_ENDPOINT.
|
||||
# Optional: AZURE_OPENAI_API_VERSION (defaults to 2024-12-01-preview),
|
||||
# AZURE_OPENAI_DEPLOYMENT or GRAPHIFY_AZURE_MODEL (deployment name).
|
||||
# base_url is intentionally absent — prevents accidental routing through
|
||||
# _call_openai_compat, which requires it and uses the wrong SDK client class.
|
||||
"default_model": os.environ.get("AZURE_OPENAI_DEPLOYMENT", os.environ.get("GRAPHIFY_AZURE_MODEL", "gpt-4o")),
|
||||
"env_key": "AZURE_OPENAI_API_KEY",
|
||||
"model_env_key": "GRAPHIFY_AZURE_MODEL",
|
||||
"pricing": {"input": 2.50, "output": 10.00}, # USD per 1M tokens (gpt-4o; may mis-estimate other deployments)
|
||||
"temperature": 0,
|
||||
"max_tokens": 16384,
|
||||
},
|
||||
"bedrock": {
|
||||
"default_model": "anthropic.claude-3-5-sonnet-20241022-v2:0",
|
||||
"model_env_key": "GRAPHIFY_BEDROCK_MODEL",
|
||||
"pricing": {"input": 3.0, "output": 15.0}, # USD per 1M tokens
|
||||
"temperature": 0,
|
||||
"max_tokens": 16384,
|
||||
"vision": True,
|
||||
},
|
||||
"claude-cli": {
|
||||
# Routes through the locally-installed `claude` CLI (Claude Code) using
|
||||
@@ -115,6 +139,9 @@ BACKENDS: dict[str, dict] = {
|
||||
"pricing": {"input": 0.0, "output": 0.0},
|
||||
"temperature": 0,
|
||||
"max_tokens": 16384,
|
||||
# Claude Code is multimodal; images are passed by path and read with the
|
||||
# CLI's Read tool rather than as inline base64 (see `_call_claude_cli`).
|
||||
"vision": True,
|
||||
},
|
||||
}
|
||||
|
||||
@@ -259,6 +286,20 @@ def _extraction_system(*, deep: bool = False) -> str:
|
||||
return _EXTRACTION_SYSTEM + _DEEP_EXTRACTION_SUFFIX
|
||||
|
||||
|
||||
def _file_to_text(path: Path) -> str:
|
||||
"""Return a text-like file's content for the extraction prompt.
|
||||
|
||||
Most files are read directly. PDFs are binary, so reading them with
|
||||
`read_text` yields garbage (the same failure images had); route them through
|
||||
pypdf instead. A scanned PDF with no text layer extracts to an empty string,
|
||||
which still produces a reference node rather than noise.
|
||||
"""
|
||||
if path.suffix.lower() == ".pdf":
|
||||
from graphify.detect import extract_pdf_text
|
||||
return extract_pdf_text(path)
|
||||
return path.read_text(encoding="utf-8", errors="replace")
|
||||
|
||||
|
||||
def _read_files(paths: list[Path], root: Path) -> str:
|
||||
"""Return file contents formatted for the extraction prompt."""
|
||||
parts: list[str] = []
|
||||
@@ -268,13 +309,226 @@ def _read_files(paths: list[Path], root: Path) -> str:
|
||||
except ValueError:
|
||||
rel = p
|
||||
try:
|
||||
content = p.read_text(encoding="utf-8", errors="replace")
|
||||
content = _file_to_text(p)
|
||||
except OSError:
|
||||
continue
|
||||
parts.append(f"=== {rel} ===\n{content[:20000]}")
|
||||
return "\n\n".join(parts)
|
||||
|
||||
|
||||
# ── Image (vision) handling ───────────────────────────────────────────────────
|
||||
# Raster image types a vision model can actually look at. `.svg` is intentionally
|
||||
# excluded: it is XML markup, so `_read_files` reads it as text (the model parses
|
||||
# the source directly), which is more useful than rasterising it. Before this,
|
||||
# every image was fed through `path.read_text(errors="replace")`, turning binary
|
||||
# pixels into garbage text — noise for API backends and an outright `exit 1` for
|
||||
# the claude-cli backend.
|
||||
_VISION_IMAGE_EXTENSIONS = {".png", ".jpg", ".jpeg", ".gif", ".webp"}
|
||||
_IMAGE_MEDIA_TYPES = {
|
||||
".png": "image/png",
|
||||
".jpg": "image/jpeg",
|
||||
".jpeg": "image/jpeg",
|
||||
".gif": "image/gif",
|
||||
".webp": "image/webp",
|
||||
}
|
||||
# Per-image byte ceiling. Anthropic caps a request at 32 MB and Bedrock images
|
||||
# at ~5 MB; 5 MB per image keeps every backend within limits. Oversized images
|
||||
# fall back to a text reference (the node is still created, just unseen).
|
||||
_MAX_IMAGE_BYTES = 5 * 1024 * 1024
|
||||
# Flat token estimate per image for chunk packing. Vision models bill an image
|
||||
# at a roughly fixed cost regardless of file size, so estimating by byte size
|
||||
# (as the generic path does) would force every large PNG into its own chunk.
|
||||
_IMAGE_TOKEN_ESTIMATE = 1_600
|
||||
# Hard cap on images per chunk, independent of the token budget. A large
|
||||
# token budget would otherwise pack hundreds of images into one request —
|
||||
# past provider per-request image limits (Anthropic allows 100), and far too
|
||||
# many for the claude-cli Read-tool loop to work through. Keeps memory and
|
||||
# request size bounded on image-dense corpora.
|
||||
_MAX_IMAGES_PER_CHUNK = 20
|
||||
# Backends that read an image by file path (claude-cli's Read tool)
|
||||
# instead of inlining base64. They open the file themselves and downsample as
|
||||
# needed, so `_MAX_IMAGE_BYTES` does not apply and the bytes never need loading.
|
||||
_PATH_IMAGE_BACKENDS = {"claude-cli"}
|
||||
|
||||
|
||||
@dataclass
|
||||
class _ImageRef:
|
||||
"""A single image destined for a vision request.
|
||||
|
||||
`raw` is None when the image is unreadable or exceeds `_MAX_IMAGE_BYTES`, or
|
||||
when the target backend has no vision support — in every such case the
|
||||
renderers emit a text reference instead of pixels, so the image still
|
||||
becomes a graph node.
|
||||
"""
|
||||
|
||||
path: Path # absolute path (claude-cli reads it via the Read tool)
|
||||
rel: str # path relative to the corpus root (the node's source_file)
|
||||
media_type: str # e.g. "image/png"
|
||||
raw: bytes | None
|
||||
|
||||
@property
|
||||
def b64(self) -> str:
|
||||
return base64.standard_b64encode(self.raw).decode("ascii") if self.raw else ""
|
||||
|
||||
@property
|
||||
def bedrock_format(self) -> str:
|
||||
# Converse wants a bare format token, not a media type.
|
||||
return self.media_type.split("/", 1)[-1]
|
||||
|
||||
|
||||
def _is_vision_image(path: Path) -> bool:
|
||||
return path.suffix.lower() in _VISION_IMAGE_EXTENSIONS
|
||||
|
||||
|
||||
def _partition_semantic_files(files: list[Path]) -> tuple[list[Path], list[Path]]:
|
||||
"""Split a chunk into (text-like files, raster-image files)."""
|
||||
text_files = [f for f in files if not _is_vision_image(f)]
|
||||
image_files = [f for f in files if _is_vision_image(f)]
|
||||
return text_files, image_files
|
||||
|
||||
|
||||
def _build_image_refs(image_files: list[Path], root: Path, *, read_bytes: bool = True) -> list[_ImageRef]:
|
||||
"""Build `_ImageRef`s for raster images.
|
||||
|
||||
`read_bytes=True` (base64 backends) loads the pixels and drops any image over
|
||||
`_MAX_IMAGE_BYTES` to a reference, because a base64 request body has a hard
|
||||
size ceiling. `read_bytes=False` (path-based backends — claude-cli)
|
||||
skips the read entirely: those backends open the file themselves and
|
||||
downsample as needed, so there is no per-image size limit and no reason to
|
||||
load (potentially tens of MB of) bytes that would never be used.
|
||||
"""
|
||||
refs: list[_ImageRef] = []
|
||||
for p in image_files:
|
||||
try:
|
||||
rel = str(p.relative_to(root))
|
||||
except ValueError:
|
||||
rel = str(p)
|
||||
media = _IMAGE_MEDIA_TYPES.get(p.suffix.lower(), "image/png")
|
||||
raw: bytes | None = None
|
||||
if read_bytes:
|
||||
try:
|
||||
raw = p.read_bytes()
|
||||
except OSError as exc:
|
||||
print(f"[graphify] could not read image {rel}: {exc}", file=sys.stderr)
|
||||
raw = None
|
||||
if raw is not None and len(raw) > _MAX_IMAGE_BYTES:
|
||||
print(
|
||||
f"[graphify] image {rel} is {len(raw) // 1024} KB, over the "
|
||||
f"{_MAX_IMAGE_BYTES // (1024 * 1024)} MB inline-image limit for this "
|
||||
"backend; sending it as a reference node without inline pixels.",
|
||||
file=sys.stderr,
|
||||
)
|
||||
raw = None
|
||||
try:
|
||||
abs_path = p.resolve()
|
||||
except OSError:
|
||||
abs_path = p
|
||||
refs.append(_ImageRef(abs_path, rel, media, raw))
|
||||
return refs
|
||||
|
||||
|
||||
def _strip_pixels(refs: list[_ImageRef]) -> list[_ImageRef]:
|
||||
"""Return refs with pixel data dropped (for non-vision backends)."""
|
||||
return [replace(r, raw=None) for r in refs]
|
||||
|
||||
|
||||
def _backend_supports_vision(backend: str) -> bool:
|
||||
"""Whether `backend`'s configured model can see images.
|
||||
|
||||
Ollama is special-cased: its default model is text-only, so vision is
|
||||
opt-in via GRAPHIFY_OLLAMA_VISION=1 once the user selects a vision model
|
||||
(e.g. --model llama3.2-vision).
|
||||
"""
|
||||
if backend == "ollama":
|
||||
return os.environ.get("GRAPHIFY_OLLAMA_VISION", "").strip() == "1"
|
||||
return bool(BACKENDS.get(backend, {}).get("vision", False))
|
||||
|
||||
|
||||
def _image_notes(refs: list[_ImageRef], *, with_paths: bool = False) -> str:
|
||||
"""Text block listing the images so the model emits one node per image.
|
||||
|
||||
Always included alongside the visual payload (and used on its own when the
|
||||
backend can't see pixels), so an image becomes a graph node either way.
|
||||
`with_paths=True` also lists the absolute path and asks the model to open it
|
||||
with the Read tool — used by the claude-cli backend.
|
||||
"""
|
||||
if not refs:
|
||||
return ""
|
||||
if with_paths:
|
||||
header = (
|
||||
"Use the Read tool to open and view each image file at the path below, "
|
||||
"then emit one node per image"
|
||||
)
|
||||
else:
|
||||
header = (
|
||||
"The following image file(s) are attached as visual input. Emit one "
|
||||
"node per image"
|
||||
)
|
||||
lines = [
|
||||
"=== IMAGES ===",
|
||||
f"{header} with \"file_type\":\"image\" and the listed source_file, a label "
|
||||
"describing what it depicts (diagram, screenshot, chart, photo, UI, logo), "
|
||||
"and edges to any code/doc nodes the image clearly references.",
|
||||
]
|
||||
for i, r in enumerate(refs, 1):
|
||||
note = f"[image {i}] source_file: {r.rel}"
|
||||
if with_paths:
|
||||
note += f" path: {r.path}"
|
||||
if r.raw is None and not with_paths:
|
||||
note += " (not shown: unreadable or exceeds size limit)"
|
||||
lines.append(note)
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def _with_image_notes(user_message: str, refs: list[_ImageRef], *, with_paths: bool = False) -> str:
|
||||
notes = _image_notes(refs, with_paths=with_paths)
|
||||
if not notes:
|
||||
return user_message
|
||||
if not user_message.strip():
|
||||
return notes
|
||||
return f"{user_message}\n\n{notes}"
|
||||
|
||||
|
||||
def _anthropic_content(user_message: str, refs: list[_ImageRef]):
|
||||
"""Build the Anthropic `messages[].content` value (str, or block list with images)."""
|
||||
blocks = [
|
||||
{"type": "image", "source": {"type": "base64", "media_type": r.media_type, "data": r.b64}}
|
||||
for r in refs
|
||||
if r.raw
|
||||
]
|
||||
text = _with_image_notes(user_message, refs)
|
||||
if not blocks:
|
||||
return text
|
||||
return [*blocks, {"type": "text", "text": text}]
|
||||
|
||||
|
||||
def _openai_content(user_message: str, refs: list[_ImageRef]):
|
||||
"""Build the OpenAI-compatible user `content` value (str, or part list with images)."""
|
||||
parts: list[dict] = [
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": f"data:{r.media_type};base64,{r.b64}", "detail": "auto"},
|
||||
}
|
||||
for r in refs
|
||||
if r.raw
|
||||
]
|
||||
text = _with_image_notes(user_message, refs)
|
||||
if not parts:
|
||||
return text
|
||||
return [{"type": "text", "text": text}, *parts]
|
||||
|
||||
|
||||
def _bedrock_content(user_message: str, refs: list[_ImageRef]) -> list[dict]:
|
||||
"""Build the Bedrock Converse user content list (raw bytes, not base64)."""
|
||||
content: list[dict] = [
|
||||
{"image": {"format": r.bedrock_format, "source": {"bytes": r.raw}}}
|
||||
for r in refs
|
||||
if r.raw
|
||||
]
|
||||
content.append({"text": _with_image_notes(user_message, refs)})
|
||||
return content
|
||||
|
||||
|
||||
_LLM_JSON_MAX_BYTES = 10 * 1024 * 1024 # 10 MB hard cap before json.loads (F-016)
|
||||
|
||||
|
||||
@@ -434,6 +688,7 @@ def _call_openai_compat(
|
||||
*,
|
||||
backend: str = "",
|
||||
deep_mode: bool = False,
|
||||
images: list[_ImageRef] | None = None,
|
||||
) -> dict:
|
||||
"""Call any OpenAI-compatible API (Kimi, OpenAI, etc.) and return parsed JSON."""
|
||||
try:
|
||||
@@ -452,7 +707,7 @@ def _call_openai_compat(
|
||||
"model": model,
|
||||
"messages": [
|
||||
{"role": "system", "content": _extraction_system(deep=deep_mode)},
|
||||
{"role": "user", "content": user_message},
|
||||
{"role": "user", "content": _openai_content(user_message, images or [])},
|
||||
],
|
||||
"max_completion_tokens": max_completion_tokens,
|
||||
}
|
||||
@@ -547,7 +802,7 @@ def _call_openai_compat(
|
||||
return result
|
||||
|
||||
|
||||
def _call_claude(api_key: str, model: str, user_message: str, max_tokens: int = 8192, *, deep_mode: bool = False) -> dict:
|
||||
def _call_claude(api_key: str, model: str, user_message: str, max_tokens: int = 8192, *, deep_mode: bool = False, images: list[_ImageRef] | None = None) -> dict:
|
||||
"""Call Anthropic Claude directly (not via OpenAI compat layer)."""
|
||||
try:
|
||||
import anthropic
|
||||
@@ -559,7 +814,7 @@ def _call_claude(api_key: str, model: str, user_message: str, max_tokens: int =
|
||||
model=model,
|
||||
max_tokens=max_tokens,
|
||||
system=_extraction_system(deep=deep_mode),
|
||||
messages=[{"role": "user", "content": user_message}],
|
||||
messages=[{"role": "user", "content": _anthropic_content(user_message, images or [])}],
|
||||
)
|
||||
raw_content = resp.content[0].text if resp.content else None
|
||||
result = _parse_llm_json(raw_content or "{}")
|
||||
@@ -580,12 +835,16 @@ def _call_claude(api_key: str, model: str, user_message: str, max_tokens: int =
|
||||
return result
|
||||
|
||||
|
||||
def _call_claude_cli(user_message: str, max_tokens: int = 8192, *, deep_mode: bool = False) -> dict:
|
||||
def _call_claude_cli(user_message: str, max_tokens: int = 8192, *, deep_mode: bool = False, images: list[_ImageRef] | None = None) -> dict:
|
||||
"""Call Claude via the locally-installed Claude Code CLI (`claude -p`).
|
||||
|
||||
Routes through the user's Claude Code subscription auth instead of a separate
|
||||
ANTHROPIC_API_KEY. Useful for Pro/Max subscribers who don't want to provision
|
||||
a pay-as-you-go API key just to run graphify's semantic pass.
|
||||
|
||||
Images are passed by absolute path rather than inline base64: the prompt asks
|
||||
the model to open each one with its Read tool, and each containing directory
|
||||
is allowlisted with `--add-dir` so the read is permitted.
|
||||
"""
|
||||
import platform
|
||||
import shutil
|
||||
@@ -621,10 +880,23 @@ def _call_claude_cli(user_message: str, max_tokens: int = 8192, *, deep_mode: bo
|
||||
# preamble — both of which fail the strict json.loads in _parse_llm_json.
|
||||
# Replacing the default prompt eliminates the conflict at the source.
|
||||
# Side benefit: cache-creation tokens per call drop ~19% in practice.
|
||||
# When images are present, append the Read-the-paths instruction and
|
||||
# allowlist each containing directory so the CLI's Read tool can open them.
|
||||
add_dir_args: list[str] = []
|
||||
if images:
|
||||
user_message = _with_image_notes(user_message, images, with_paths=True)
|
||||
seen_dirs: set[str] = set()
|
||||
for r in images:
|
||||
d = str(r.path.parent)
|
||||
if d not in seen_dirs:
|
||||
seen_dirs.add(d)
|
||||
add_dir_args.extend(["--add-dir", d])
|
||||
|
||||
cli_args = [
|
||||
claude_cmd, "-p",
|
||||
"--output-format", "json",
|
||||
"--no-session-persistence",
|
||||
*add_dir_args,
|
||||
"--system-prompt", _extraction_system(deep=deep_mode),
|
||||
]
|
||||
# claude-cli defaults to Opus, which is overkill for the structured-JSON
|
||||
@@ -680,7 +952,69 @@ def _call_claude_cli(user_message: str, max_tokens: int = 8192, *, deep_mode: bo
|
||||
return result
|
||||
|
||||
|
||||
def _call_bedrock(model: str, user_message: str, max_tokens: int = 8192, *, deep_mode: bool = False) -> dict:
|
||||
def _azure_client(api_key: str, endpoint: str):
|
||||
"""Construct an AzureOpenAI client with env-driven api_version and timeout."""
|
||||
try:
|
||||
from openai import AzureOpenAI
|
||||
except ImportError as exc:
|
||||
raise ImportError(
|
||||
"Azure OpenAI requires the openai package. Run: pip install openai"
|
||||
) from exc
|
||||
api_version = os.environ.get("AZURE_OPENAI_API_VERSION", "2024-12-01-preview").strip()
|
||||
timeout_raw = os.environ.get("GRAPHIFY_API_TIMEOUT", "").strip()
|
||||
timeout_s: float = 600.0
|
||||
if timeout_raw:
|
||||
try:
|
||||
v = float(timeout_raw)
|
||||
if v > 0:
|
||||
timeout_s = v
|
||||
except ValueError:
|
||||
pass
|
||||
return AzureOpenAI(api_key=api_key, azure_endpoint=endpoint, api_version=api_version, timeout=timeout_s)
|
||||
|
||||
|
||||
def _call_azure(
|
||||
api_key: str,
|
||||
endpoint: str,
|
||||
model: str,
|
||||
user_message: str,
|
||||
temperature: float | None = 0,
|
||||
max_tokens: int = 8192,
|
||||
*,
|
||||
deep_mode: bool = False,
|
||||
) -> dict:
|
||||
"""Call Azure OpenAI Service via the AzureOpenAI SDK client."""
|
||||
client = _azure_client(api_key, endpoint)
|
||||
kwargs: dict = {
|
||||
"model": model,
|
||||
"messages": [
|
||||
{"role": "system", "content": _extraction_system(deep=deep_mode)},
|
||||
{"role": "user", "content": user_message},
|
||||
],
|
||||
"max_completion_tokens": max_tokens,
|
||||
}
|
||||
if temperature is not None:
|
||||
kwargs["temperature"] = temperature
|
||||
resp = client.chat.completions.create(**kwargs)
|
||||
if not resp.choices or resp.choices[0].message is None:
|
||||
raise ValueError("Azure OpenAI returned empty or filtered response")
|
||||
raw_content = resp.choices[0].message.content
|
||||
result = _parse_llm_json(raw_content or "{}")
|
||||
result["input_tokens"] = resp.usage.prompt_tokens if resp.usage else 0
|
||||
result["output_tokens"] = resp.usage.completion_tokens if resp.usage else 0
|
||||
result["model"] = model
|
||||
result["finish_reason"] = resp.choices[0].finish_reason
|
||||
if _response_is_hollow(raw_content, result) and result["finish_reason"] != "length":
|
||||
print(
|
||||
"[graphify] azure returned a hollow response; treating as "
|
||||
"truncation so adaptive retry can bisect the chunk.",
|
||||
file=sys.stderr,
|
||||
)
|
||||
result["finish_reason"] = "length"
|
||||
return result
|
||||
|
||||
|
||||
def _call_bedrock(model: str, user_message: str, max_tokens: int = 8192, *, deep_mode: bool = False, images: list[_ImageRef] | None = None) -> dict:
|
||||
"""Call AWS Bedrock via boto3 Converse API using the standard AWS credential chain."""
|
||||
try:
|
||||
import boto3
|
||||
@@ -699,7 +1033,7 @@ def _call_bedrock(model: str, user_message: str, max_tokens: int = 8192, *, deep
|
||||
resp = client.converse(
|
||||
modelId=model,
|
||||
system=[{"text": _extraction_system(deep=deep_mode)}],
|
||||
messages=[{"role": "user", "content": [{"text": user_message}]}],
|
||||
messages=[{"role": "user", "content": _bedrock_content(user_message, images or [])}],
|
||||
inferenceConfig={"maxTokens": max_tokens, "temperature": 0},
|
||||
)
|
||||
except botocore.exceptions.ClientError as exc:
|
||||
@@ -744,7 +1078,8 @@ def extract_files_direct(
|
||||
if backend is None:
|
||||
raise ValueError(
|
||||
"No LLM backend configured. Set one of: GEMINI_API_KEY, ANTHROPIC_API_KEY, "
|
||||
"OPENAI_API_KEY, DEEPSEEK_API_KEY, MOONSHOT_API_KEY, OLLAMA_BASE_URL, "
|
||||
"OPENAI_API_KEY, DEEPSEEK_API_KEY, MOONSHOT_API_KEY, "
|
||||
"AZURE_OPENAI_API_KEY+AZURE_OPENAI_ENDPOINT, OLLAMA_BASE_URL, "
|
||||
"or AWS credentials. Pass backend= explicitly to select a provider."
|
||||
)
|
||||
if backend not in BACKENDS:
|
||||
@@ -771,15 +1106,42 @@ def extract_files_direct(
|
||||
f"Set {_format_backend_env_keys(backend)} or pass api_key=."
|
||||
)
|
||||
mdl = model or _default_model_for_backend(backend)
|
||||
user_msg = _read_files(files, root)
|
||||
# Separate raster images from text-like files. Text goes through _read_files
|
||||
# as before; images become structured refs the backend renders as pixels
|
||||
# (vision backends) or as a text reference node (everything else).
|
||||
text_files, image_files = _partition_semantic_files(files)
|
||||
user_msg = _read_files(text_files, root)
|
||||
vision = _backend_supports_vision(backend)
|
||||
# Only base64 (inline) vision backends need the bytes loaded + size-capped;
|
||||
# path-based backends (claude-cli) and non-vision backends do not.
|
||||
read_bytes = vision and backend not in _PATH_IMAGE_BACKENDS
|
||||
image_refs = _build_image_refs(image_files, root, read_bytes=read_bytes) if image_files else []
|
||||
if image_refs and not vision:
|
||||
image_refs = _strip_pixels(image_refs)
|
||||
max_out = _resolve_max_tokens(cfg.get("max_tokens", 8192))
|
||||
|
||||
if backend == "claude":
|
||||
return _call_claude(key, mdl, user_msg, max_tokens=max_out, deep_mode=deep_mode)
|
||||
return _call_claude(key, mdl, user_msg, max_tokens=max_out, deep_mode=deep_mode, images=image_refs)
|
||||
if backend == "claude-cli":
|
||||
return _call_claude_cli(user_msg, max_tokens=max_out, deep_mode=deep_mode)
|
||||
return _call_claude_cli(user_msg, max_tokens=max_out, deep_mode=deep_mode, images=image_refs)
|
||||
if backend == "bedrock":
|
||||
return _call_bedrock(mdl, user_msg, max_tokens=max_out, deep_mode=deep_mode)
|
||||
return _call_bedrock(mdl, user_msg, max_tokens=max_out, deep_mode=deep_mode, images=image_refs)
|
||||
if backend == "azure":
|
||||
endpoint = os.environ.get("AZURE_OPENAI_ENDPOINT", "").strip()
|
||||
if not endpoint:
|
||||
raise ValueError(
|
||||
"Azure OpenAI backend requires AZURE_OPENAI_ENDPOINT to be set "
|
||||
"(e.g. https://my-resource.openai.azure.com/)."
|
||||
)
|
||||
return _call_azure(
|
||||
key,
|
||||
endpoint,
|
||||
mdl,
|
||||
user_msg,
|
||||
temperature=cfg.get("temperature", 0),
|
||||
max_tokens=max_out,
|
||||
deep_mode=deep_mode,
|
||||
)
|
||||
return _call_openai_compat(
|
||||
cfg["base_url"],
|
||||
key,
|
||||
@@ -790,6 +1152,7 @@ def extract_files_direct(
|
||||
max_completion_tokens=_resolve_max_tokens(cfg.get("max_completion_tokens", 8192)),
|
||||
backend=backend,
|
||||
deep_mode=deep_mode,
|
||||
images=image_refs,
|
||||
)
|
||||
|
||||
|
||||
@@ -802,6 +1165,10 @@ def _estimate_file_tokens(path: Path) -> int:
|
||||
the `=== rel ===` separator. Returns 0 for unreadable paths so they don't
|
||||
blow up packing.
|
||||
"""
|
||||
# Raster images are not read as text; a vision model bills them at a roughly
|
||||
# fixed token cost, so estimate by image count rather than (binary) byte size.
|
||||
if _is_vision_image(path):
|
||||
return _IMAGE_TOKEN_ESTIMATE
|
||||
if _TOKENIZER is None:
|
||||
try:
|
||||
size = path.stat().st_size
|
||||
@@ -841,16 +1208,22 @@ def _pack_chunks_by_tokens(
|
||||
chunks: list[list[Path]] = []
|
||||
current: list[Path] = []
|
||||
current_tokens = 0
|
||||
current_images = 0
|
||||
|
||||
for directory in sorted(by_dir):
|
||||
for path in by_dir[directory]:
|
||||
cost = _estimate_file_tokens(path)
|
||||
if current and current_tokens + cost > token_budget:
|
||||
is_image = _is_vision_image(path)
|
||||
over_budget = current_tokens + cost > token_budget
|
||||
over_images = is_image and current_images >= _MAX_IMAGES_PER_CHUNK
|
||||
if current and (over_budget or over_images):
|
||||
chunks.append(current)
|
||||
current = []
|
||||
current_tokens = 0
|
||||
current_images = 0
|
||||
current.append(path)
|
||||
current_tokens += cost
|
||||
current_images += is_image
|
||||
|
||||
if current:
|
||||
chunks.append(current)
|
||||
@@ -1215,6 +1588,7 @@ def _call_llm(prompt: str, *, backend: str, max_tokens: int = 200) -> str:
|
||||
raise RuntimeError(f"claude -p produced unparseable JSON envelope: {exc}") from exc
|
||||
return envelope.get("result", "")
|
||||
|
||||
|
||||
if backend == "bedrock":
|
||||
try:
|
||||
import boto3
|
||||
@@ -1231,6 +1605,23 @@ def _call_llm(prompt: str, *, backend: str, max_tokens: int = 200) -> str:
|
||||
)
|
||||
return resp.get("output", {}).get("message", {}).get("content", [{}])[0].get("text", "")
|
||||
|
||||
if backend == "azure":
|
||||
endpoint = os.environ.get("AZURE_OPENAI_ENDPOINT", "").strip()
|
||||
if not endpoint:
|
||||
raise ValueError(
|
||||
"Azure OpenAI backend requires AZURE_OPENAI_ENDPOINT to be set."
|
||||
)
|
||||
azure_client = _azure_client(key, endpoint)
|
||||
resp = azure_client.chat.completions.create(
|
||||
model=mdl,
|
||||
messages=[{"role": "user", "content": prompt}],
|
||||
max_completion_tokens=max_tokens,
|
||||
temperature=cfg.get("temperature", 0),
|
||||
)
|
||||
if not resp.choices or resp.choices[0].message is None:
|
||||
raise ValueError("Azure OpenAI returned empty or filtered response")
|
||||
return resp.choices[0].message.content or ""
|
||||
|
||||
# OpenAI-compatible (kimi, openai, gemini, ollama)
|
||||
try:
|
||||
from openai import OpenAI
|
||||
@@ -1340,7 +1731,7 @@ def _validate_ollama_base_url(url: str, *, warn: bool = True) -> None:
|
||||
def detect_backend() -> str | None:
|
||||
"""Return the name of whichever backend has an API key set, or None.
|
||||
|
||||
Priority: gemini → kimi → claude → openai → bedrock → ollama (last, opt-in).
|
||||
Priority: gemini → kimi → claude → openai → deepseek → azure → bedrock → ollama (last, opt-in).
|
||||
|
||||
Ollama is intentionally checked LAST so a paid API key (Anthropic/OpenAI/etc.)
|
||||
is never silently shadowed by an incidental OLLAMA_BASE_URL in the environment
|
||||
@@ -1351,6 +1742,8 @@ def detect_backend() -> str | None:
|
||||
for backend in ("gemini", "kimi", "claude", "openai", "deepseek"):
|
||||
if _get_backend_api_key(backend):
|
||||
return backend
|
||||
if _get_backend_api_key("azure") and os.environ.get("AZURE_OPENAI_ENDPOINT"):
|
||||
return "azure"
|
||||
if os.environ.get("AWS_PROFILE") or os.environ.get("AWS_REGION") or os.environ.get("AWS_DEFAULT_REGION"):
|
||||
return "bedrock"
|
||||
ollama_url = os.environ.get("OLLAMA_BASE_URL")
|
||||
@@ -1358,7 +1751,7 @@ def detect_backend() -> str | None:
|
||||
_validate_ollama_base_url(ollama_url)
|
||||
return "ollama"
|
||||
for name in BACKENDS:
|
||||
if name not in ("gemini", "kimi", "claude", "openai", "deepseek", "bedrock", "ollama", "claude-cli"):
|
||||
if name not in ("gemini", "kimi", "claude", "openai", "deepseek", "azure", "bedrock", "ollama", "claude-cli"):
|
||||
if _get_backend_api_key(name):
|
||||
return name
|
||||
return None
|
||||
|
||||
@@ -0,0 +1,142 @@
|
||||
from __future__ import annotations
|
||||
from pathlib import Path
|
||||
from graphify.extract import extract_sql
|
||||
|
||||
|
||||
def _quote_ident(name: str) -> str:
|
||||
"""Double-quote a PostgreSQL identifier, escaping embedded double-quotes."""
|
||||
return '"' + name.replace('"', '""') + '"'
|
||||
|
||||
|
||||
def introspect_postgres(dsn: str | None = None) -> dict:
|
||||
"""Connect to PostgreSQL, reconstruct DDL, and extract via extract_sql()."""
|
||||
try:
|
||||
import psycopg
|
||||
except ModuleNotFoundError:
|
||||
raise ImportError(
|
||||
"psycopg is required for --postgres. "
|
||||
"Install with: pip install 'graphify[postgres]'"
|
||||
)
|
||||
|
||||
try:
|
||||
conn = psycopg.connect(dsn or "") # empty string = PG* env vars
|
||||
except psycopg.OperationalError as exc:
|
||||
# Sanitize: strip the DSN/credentials that psycopg may embed in the
|
||||
# OperationalError message (e.g. "connection to server … failed: …\nDETAIL: …")
|
||||
msg = str(exc).split("\n")[0]
|
||||
raise ConnectionError(f"could not connect to PostgreSQL: {msg}") from None
|
||||
|
||||
try:
|
||||
conn.execute("SET TRANSACTION ISOLATION LEVEL SERIALIZABLE READ ONLY DEFERRABLE")
|
||||
|
||||
# 1. Query tables
|
||||
with conn.cursor() as cur:
|
||||
cur.execute("""
|
||||
SELECT table_schema, table_name, table_type
|
||||
FROM information_schema.tables
|
||||
WHERE table_schema NOT IN ('pg_catalog', 'information_schema')
|
||||
ORDER BY table_schema, table_name;
|
||||
""")
|
||||
tables = cur.fetchall()
|
||||
|
||||
# 2. Query views
|
||||
cur.execute("""
|
||||
SELECT table_schema, table_name, view_definition
|
||||
FROM information_schema.views
|
||||
WHERE table_schema NOT IN ('pg_catalog', 'information_schema')
|
||||
ORDER BY table_schema, table_name;
|
||||
""")
|
||||
views = cur.fetchall()
|
||||
|
||||
# 3. Query routines (functions/procedures), including language
|
||||
cur.execute("""
|
||||
SELECT routine_schema, routine_name, routine_type,
|
||||
routine_definition, external_language
|
||||
FROM information_schema.routines
|
||||
WHERE routine_schema NOT IN ('pg_catalog', 'information_schema')
|
||||
ORDER BY routine_schema, routine_name;
|
||||
""")
|
||||
routines = cur.fetchall()
|
||||
|
||||
# 4. Query foreign keys — grouped by constraint to handle composites
|
||||
cur.execute("""
|
||||
SELECT
|
||||
tc.constraint_name,
|
||||
kcu1.table_schema,
|
||||
kcu1.table_name,
|
||||
ARRAY_AGG(kcu1.column_name ORDER BY kcu1.ordinal_position) AS columns,
|
||||
kcu2.table_schema AS foreign_table_schema,
|
||||
kcu2.table_name AS foreign_table_name,
|
||||
ARRAY_AGG(kcu2.column_name ORDER BY kcu2.ordinal_position) AS foreign_columns
|
||||
FROM
|
||||
information_schema.table_constraints AS tc
|
||||
JOIN information_schema.referential_constraints AS rc
|
||||
ON tc.constraint_name = rc.constraint_name
|
||||
AND tc.table_schema = rc.constraint_schema
|
||||
JOIN information_schema.key_column_usage AS kcu1
|
||||
ON tc.constraint_name = kcu1.constraint_name
|
||||
AND tc.table_schema = kcu1.table_schema
|
||||
JOIN information_schema.key_column_usage AS kcu2
|
||||
ON rc.unique_constraint_name = kcu2.constraint_name
|
||||
AND rc.unique_constraint_schema = kcu2.table_schema
|
||||
AND kcu1.position_in_unique_constraint = kcu2.ordinal_position
|
||||
WHERE tc.constraint_type = 'FOREIGN KEY'
|
||||
AND tc.table_schema NOT IN ('pg_catalog', 'information_schema')
|
||||
GROUP BY tc.constraint_name, kcu1.table_schema, kcu1.table_name,
|
||||
kcu2.table_schema, kcu2.table_name
|
||||
ORDER BY kcu1.table_schema, kcu1.table_name;
|
||||
""")
|
||||
fks = cur.fetchall()
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
ddl = []
|
||||
|
||||
# Tables — quote identifiers to handle reserved words, hyphens, mixed-case
|
||||
for schema, name, ttype in tables:
|
||||
if ttype == "BASE TABLE":
|
||||
ddl.append(f"CREATE TABLE {_quote_ident(schema)}.{_quote_ident(name)} (id INT);")
|
||||
|
||||
# Views — real body if available, stub if NULL (permission denied)
|
||||
for schema, name, body in views:
|
||||
if body:
|
||||
ddl.append(f"CREATE VIEW {_quote_ident(schema)}.{_quote_ident(name)} AS {body};")
|
||||
else:
|
||||
ddl.append(f"CREATE VIEW {_quote_ident(schema)}.{_quote_ident(name)} AS SELECT 1;")
|
||||
|
||||
# Functions & Procedures — real body if available, stub if NULL
|
||||
# Use $gfx$ as the dollar-quote tag to avoid collision with $$ inside bodies.
|
||||
# Use external_language from the catalog; fall back to plpgsql if NULL/blank.
|
||||
for schema, name, rtype, body, ext_lang in routines:
|
||||
lang = (ext_lang or "plpgsql").lower()
|
||||
fn_sig = f"{_quote_ident(schema)}.{_quote_ident(name)}()"
|
||||
stub_body = "BEGIN SELECT 1; END;"
|
||||
if rtype in ("FUNCTION", "PROCEDURE"):
|
||||
actual_body = body if body else stub_body
|
||||
# Represent PROCEDUREs as FUNCTION so tree-sitter-sql can parse them
|
||||
ddl.append(
|
||||
f"CREATE FUNCTION {fn_sig} RETURNS void"
|
||||
f" AS $gfx$ {actual_body} $gfx$ LANGUAGE {lang};"
|
||||
)
|
||||
|
||||
# FK edges — one ALTER TABLE per constraint (handles composite FKs correctly)
|
||||
for constraint_name, t_schema, t_name, cols, r_schema, r_name, r_cols in fks:
|
||||
col_list = ", ".join(_quote_ident(c) for c in cols)
|
||||
ref_col_list = ", ".join(_quote_ident(c) for c in r_cols)
|
||||
ddl.append(
|
||||
f"ALTER TABLE {_quote_ident(t_schema)}.{_quote_ident(t_name)} "
|
||||
f"ADD CONSTRAINT {_quote_ident(constraint_name)} "
|
||||
f"FOREIGN KEY ({col_list}) REFERENCES {_quote_ident(r_schema)}.{_quote_ident(r_name)}({ref_col_list});"
|
||||
)
|
||||
|
||||
ddl_string = "\n".join(ddl)
|
||||
|
||||
# Determine host/dbname for virtual path DSN sanitization
|
||||
info = psycopg.conninfo.conninfo_to_dict(dsn or "")
|
||||
host = info.get("host", "localhost")
|
||||
dbname = info.get("dbname", "db")
|
||||
virtual_path = Path(f"postgresql://{host}/{dbname}")
|
||||
|
||||
# Pass virtual path and in-memory DDL content to extract_sql
|
||||
result = extract_sql(virtual_path, content=ddl_string)
|
||||
return result
|
||||
@@ -543,9 +543,19 @@ def _rebuild_code(
|
||||
evict_sources.add(sf)
|
||||
evict_sources.add(norm)
|
||||
deleted_paths.add(norm)
|
||||
# On a full re-extraction every code file is re-extracted, so
|
||||
# new_ast_ids is the complete current AST set. Any AST-marked node
|
||||
# missing from it is stale and must be dropped even if its source
|
||||
# file still exists (a symbol removed from a surviving file, #1116).
|
||||
# Gate on full_rebuild: in incremental mode an AST node from an
|
||||
# unchanged file is legitimately absent from new_ast_ids. Semantic
|
||||
# nodes lack the "_origin" marker, so they are never dropped here —
|
||||
# only by the deleted-file eviction in evict_sources above.
|
||||
full_rebuild = changed_paths is None
|
||||
preserved_nodes = [
|
||||
n for n in existing.get("nodes", [])
|
||||
if n["id"] not in new_ast_ids
|
||||
and not (full_rebuild and n.get("_origin") == "ast")
|
||||
and (not evict_sources or n.get("source_file") not in evict_sources)
|
||||
]
|
||||
all_ids = new_ast_ids | {n["id"] for n in preserved_nodes}
|
||||
|
||||
@@ -56,6 +56,7 @@ svg = ["matplotlib"]
|
||||
leiden = ["graspologic; python_version < '3.13'"]
|
||||
office = ["python-docx", "openpyxl"]
|
||||
google = ["openpyxl"]
|
||||
postgres = ["psycopg[binary]"]
|
||||
video = ["faster-whisper; python_version >= '3.11'", "yt-dlp"]
|
||||
kimi = ["openai", "tiktoken"]
|
||||
ollama = ["openai"]
|
||||
|
||||
Vendored
+55
@@ -0,0 +1,55 @@
|
||||
public with sharing class AccountService {
|
||||
|
||||
private static final String DEFAULT_TYPE = 'Customer';
|
||||
|
||||
public interface Notifiable {
|
||||
void notify(String message);
|
||||
}
|
||||
|
||||
public enum AccountStatus { ACTIVE, INACTIVE, PENDING }
|
||||
|
||||
@AuraEnabled
|
||||
public static List<Account> getAccounts(String accountType) {
|
||||
return [SELECT Id, Name, Type FROM Account WHERE Type = :accountType];
|
||||
}
|
||||
|
||||
@future
|
||||
public static void updateAccountsAsync(List<Id> accountIds) {
|
||||
List<Account> accounts = [SELECT Id FROM Account WHERE Id IN :accountIds];
|
||||
for (Account acc : accounts) {
|
||||
acc.Type = DEFAULT_TYPE;
|
||||
}
|
||||
update accounts;
|
||||
}
|
||||
|
||||
@InvocableMethod(label='Create Account' description='Creates a new Account')
|
||||
public static List<Id> createAccounts(List<String> names) {
|
||||
List<Account> toInsert = new List<Account>();
|
||||
for (String n : names) {
|
||||
toInsert.add(new Account(Name = n));
|
||||
}
|
||||
insert toInsert;
|
||||
List<Id> ids = new List<Id>();
|
||||
for (Account a : toInsert) {
|
||||
ids.add(a.Id);
|
||||
}
|
||||
return ids;
|
||||
}
|
||||
|
||||
public static void deleteOldAccounts(Date cutoff) {
|
||||
List<Account> old = [SELECT Id FROM Account WHERE CreatedDate < :cutoff];
|
||||
delete old;
|
||||
}
|
||||
|
||||
@isTest
|
||||
static void testGetAccounts() {
|
||||
List<Account> result = getAccounts('Customer');
|
||||
System.assertNotEquals(null, result);
|
||||
}
|
||||
|
||||
@isTest
|
||||
static void testCreateAccounts() {
|
||||
List<Id> ids = createAccounts(new List<String>{'Test'});
|
||||
System.assertEquals(1, ids.size());
|
||||
}
|
||||
}
|
||||
Vendored
+8
@@ -0,0 +1,8 @@
|
||||
trigger AccountTrigger on Account (before insert, before update, after insert, after update) {
|
||||
if (Trigger.isBefore) {
|
||||
AccountService.validateAccounts(Trigger.new);
|
||||
}
|
||||
if (Trigger.isAfter && Trigger.isInsert) {
|
||||
AccountService.sendWelcomeNotifications(Trigger.new);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,345 @@
|
||||
"""Tests for image-vision support across the direct extraction backends.
|
||||
|
||||
Covers the structured-message split (text vs raster image), the per-backend
|
||||
payload rendering (Anthropic base64 blocks, OpenAI/Gemini image_url data URIs,
|
||||
Bedrock raw-bytes Converse blocks, the claude-cli Read-tool path), and the vision-capability gating that sends pixels only to
|
||||
backends whose model can see them.
|
||||
|
||||
Every backend is mocked (fake SDK module / subprocess), so the suite runs on CI
|
||||
with no API keys, no network, and no `claude` binary.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import sys
|
||||
import types
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from graphify import llm
|
||||
|
||||
# A 1x1 PNG is unnecessary — the renderers never decode pixels, they only base64
|
||||
# the bytes — so any non-empty byte string stands in for image content.
|
||||
_PNG_BYTES = b"\x89PNG\r\n\x1a\nFAKEPIXELDATA"
|
||||
_NODE_JSON = json.dumps({
|
||||
"nodes": [{"id": "x", "label": "L", "file_type": "image", "source_file": "a.png"}],
|
||||
"edges": [],
|
||||
"hyperedges": [],
|
||||
})
|
||||
|
||||
|
||||
def _make_corpus(tmp_path):
|
||||
"""A corpus with one raster image, one svg (text), and one markdown doc."""
|
||||
(tmp_path / "sub").mkdir()
|
||||
img = tmp_path / "sub" / "diagram.png"
|
||||
img.write_bytes(_PNG_BYTES)
|
||||
svg = tmp_path / "icon.svg"
|
||||
svg.write_text("<svg><rect/></svg>")
|
||||
doc = tmp_path / "README.md"
|
||||
doc.write_text("# Title\nbody")
|
||||
return img, svg, doc
|
||||
|
||||
|
||||
# ── pure helpers ──────────────────────────────────────────────────────────────
|
||||
|
||||
def test_pdf_routed_through_pypdf_not_readtext(tmp_path, monkeypatch):
|
||||
# A PDF is binary; reading it as text yields garbage (the bug). It must be
|
||||
# routed through the pypdf extractor, and the raw bytes must never reach the
|
||||
# prompt.
|
||||
pdf = tmp_path / "paper.pdf"
|
||||
pdf.write_bytes(b"%PDF-1.4 RAWBINARYGARBAGE\x00\xff")
|
||||
import graphify.detect as detect
|
||||
monkeypatch.setattr(detect, "extract_pdf_text", lambda p: "EXTRACTED PDF TEXT")
|
||||
out = llm._read_files([pdf], tmp_path)
|
||||
assert "EXTRACTED PDF TEXT" in out
|
||||
assert "RAWBINARYGARBAGE" not in out
|
||||
|
||||
|
||||
def test_pdf_is_not_treated_as_vision_image(tmp_path):
|
||||
pdf = tmp_path / "paper.pdf"
|
||||
pdf.write_bytes(b"%PDF-1.4")
|
||||
text_files, image_files = llm._partition_semantic_files([pdf])
|
||||
assert text_files == [pdf] and image_files == []
|
||||
|
||||
|
||||
def test_non_pdf_still_read_as_plain_text(tmp_path):
|
||||
md = tmp_path / "a.md"
|
||||
md.write_text("# hello")
|
||||
assert "# hello" in llm._file_to_text(md)
|
||||
|
||||
|
||||
def test_partition_splits_raster_from_text(tmp_path):
|
||||
img, svg, doc = _make_corpus(tmp_path)
|
||||
text_files, image_files = llm._partition_semantic_files([doc, img, svg])
|
||||
assert image_files == [img]
|
||||
# svg is XML markup, so it stays on the text side (read as source, not pixels)
|
||||
assert set(text_files) == {doc, svg}
|
||||
|
||||
|
||||
def test_build_image_refs_sets_rel_media_and_bytes(tmp_path):
|
||||
img, _, _ = _make_corpus(tmp_path)
|
||||
(ref,) = llm._build_image_refs([img], tmp_path)
|
||||
assert ref.rel == "sub/diagram.png"
|
||||
assert ref.media_type == "image/png"
|
||||
assert ref.raw == _PNG_BYTES
|
||||
assert ref.b64 # non-empty base64
|
||||
assert ref.bedrock_format == "png"
|
||||
|
||||
|
||||
def test_build_image_refs_drops_oversized(tmp_path, monkeypatch):
|
||||
big = tmp_path / "big.jpg"
|
||||
big.write_bytes(b"x" * 64)
|
||||
monkeypatch.setattr(llm, "_MAX_IMAGE_BYTES", 8)
|
||||
(ref,) = llm._build_image_refs([big], tmp_path)
|
||||
assert ref.raw is None # too large -> reference node only, no pixels
|
||||
assert ref.media_type == "image/jpeg"
|
||||
|
||||
|
||||
def test_path_backend_skips_byte_read_and_size_cap(tmp_path, monkeypatch):
|
||||
# Path-based backends (claude-cli) read the file themselves, so
|
||||
# _build_image_refs(read_bytes=False) loads no bytes and applies no size cap.
|
||||
big = tmp_path / "huge.png"
|
||||
big.write_bytes(b"x" * 64)
|
||||
monkeypatch.setattr(llm, "_MAX_IMAGE_BYTES", 8)
|
||||
(ref,) = llm._build_image_refs([big], tmp_path, read_bytes=False)
|
||||
assert ref.raw is None # never read
|
||||
assert ref.rel == "huge.png" and ref.path.name == "huge.png" # path still usable
|
||||
|
||||
|
||||
def test_claude_cli_passes_oversized_image_by_path(tmp_path, monkeypatch):
|
||||
# An image over the inline base64 cap must still reach claude-cli by path —
|
||||
# opus reads it via the Read tool, no size limit on that route.
|
||||
big = tmp_path / "huge.png"
|
||||
big.write_bytes(b"x" * 100)
|
||||
monkeypatch.setattr(llm, "_MAX_IMAGE_BYTES", 8)
|
||||
refs = llm._build_image_refs([big], tmp_path, read_bytes=False)
|
||||
envelope = {"result": _NODE_JSON, "usage": {"output_tokens": 1}, "stop_reason": "end_turn"}
|
||||
seen: dict = {}
|
||||
|
||||
def fake_run(args, **kw):
|
||||
seen["input"] = kw.get("input", "")
|
||||
return MagicMock(returncode=0, stdout=json.dumps(envelope), stderr="")
|
||||
|
||||
monkeypatch.setattr(llm, "_response_is_hollow", lambda r, p: False)
|
||||
with patch("shutil.which", return_value="/fake/claude"), \
|
||||
patch("subprocess.run", side_effect=fake_run):
|
||||
llm._call_claude_cli("CORPUS", images=refs)
|
||||
assert str(refs[0].path) in seen["input"]
|
||||
|
||||
|
||||
def test_capability_flags(monkeypatch):
|
||||
for b in ("claude", "claude-cli", "openai", "gemini", "bedrock", "kimi"):
|
||||
assert llm._backend_supports_vision(b), b
|
||||
assert not llm._backend_supports_vision("deepseek")
|
||||
# ollama is opt-in via env (default model is text-only)
|
||||
monkeypatch.delenv("GRAPHIFY_OLLAMA_VISION", raising=False)
|
||||
assert not llm._backend_supports_vision("ollama")
|
||||
monkeypatch.setenv("GRAPHIFY_OLLAMA_VISION", "1")
|
||||
assert llm._backend_supports_vision("ollama")
|
||||
|
||||
|
||||
def test_image_token_estimate_is_flat(tmp_path):
|
||||
img, _, _ = _make_corpus(tmp_path)
|
||||
assert llm._estimate_file_tokens(img) == llm._IMAGE_TOKEN_ESTIMATE
|
||||
|
||||
|
||||
def test_chunk_packing_caps_images_per_chunk(tmp_path):
|
||||
# Many images + a huge token budget must still cap images per chunk so a
|
||||
# single request never exceeds provider image limits.
|
||||
imgs = []
|
||||
for i in range(llm._MAX_IMAGES_PER_CHUNK * 2 + 3):
|
||||
p = tmp_path / f"img{i:03d}.png"
|
||||
p.write_bytes(_PNG_BYTES)
|
||||
imgs.append(p)
|
||||
chunks = llm._pack_chunks_by_tokens(imgs, token_budget=10_000_000)
|
||||
assert len(chunks) >= 3 # would be 1 chunk without the cap
|
||||
for chunk in chunks:
|
||||
n_imgs = sum(1 for p in chunk if llm._is_vision_image(p))
|
||||
assert n_imgs <= llm._MAX_IMAGES_PER_CHUNK
|
||||
|
||||
|
||||
# ── content builders ──────────────────────────────────────────────────────────
|
||||
|
||||
def test_anthropic_content_has_base64_block(tmp_path):
|
||||
img, _, _ = _make_corpus(tmp_path)
|
||||
refs = llm._build_image_refs([img], tmp_path)
|
||||
content = llm._anthropic_content("CORPUS", refs)
|
||||
assert isinstance(content, list)
|
||||
assert content[0]["type"] == "image"
|
||||
assert content[0]["source"] == {
|
||||
"type": "base64", "media_type": "image/png", "data": refs[0].b64,
|
||||
}
|
||||
assert content[-1]["type"] == "text" and "CORPUS" in content[-1]["text"]
|
||||
|
||||
|
||||
def test_openai_content_has_data_uri(tmp_path):
|
||||
img, _, _ = _make_corpus(tmp_path)
|
||||
refs = llm._build_image_refs([img], tmp_path)
|
||||
content = llm._openai_content("CORPUS", refs)
|
||||
assert content[0]["type"] == "text"
|
||||
assert content[1]["type"] == "image_url"
|
||||
assert content[1]["image_url"]["url"] == f"data:image/png;base64,{refs[0].b64}"
|
||||
|
||||
|
||||
def test_bedrock_content_uses_raw_bytes(tmp_path):
|
||||
img, _, _ = _make_corpus(tmp_path)
|
||||
refs = llm._build_image_refs([img], tmp_path)
|
||||
content = llm._bedrock_content("CORPUS", refs)
|
||||
assert content[0]["image"]["format"] == "png"
|
||||
# Converse takes raw bytes, NOT base64 (the SDK encodes on the wire)
|
||||
assert content[0]["image"]["source"]["bytes"] == _PNG_BYTES
|
||||
assert content[-1]["text"] and "CORPUS" in content[-1]["text"] # text block carries the corpus
|
||||
|
||||
|
||||
def test_builders_fall_back_to_string_without_pixels(tmp_path):
|
||||
img, _, _ = _make_corpus(tmp_path)
|
||||
stripped = llm._strip_pixels(llm._build_image_refs([img], tmp_path))
|
||||
# No pixels -> Anthropic/OpenAI render a plain string carrying the note
|
||||
ac = llm._anthropic_content("CORPUS", stripped)
|
||||
oc = llm._openai_content("CORPUS", stripped)
|
||||
assert isinstance(ac, str) and "sub/diagram.png" in ac
|
||||
assert isinstance(oc, str) and "sub/diagram.png" in oc
|
||||
|
||||
|
||||
def test_no_images_is_byte_identical(tmp_path):
|
||||
# With no image refs, the user content must be exactly the text blob.
|
||||
assert llm._anthropic_content("PLAIN", []) == "PLAIN"
|
||||
assert llm._openai_content("PLAIN", []) == "PLAIN"
|
||||
|
||||
|
||||
# ── fake SDK modules ──────────────────────────────────────────────────────────
|
||||
|
||||
def _fake_anthropic(monkeypatch, captured):
|
||||
class _Messages:
|
||||
def create(self, **kw):
|
||||
captured.update(kw)
|
||||
return SimpleNamespace(
|
||||
content=[SimpleNamespace(text=_NODE_JSON)],
|
||||
usage=SimpleNamespace(input_tokens=5, output_tokens=7),
|
||||
stop_reason="end_turn",
|
||||
)
|
||||
mod = types.ModuleType("anthropic")
|
||||
mod.Anthropic = lambda **kw: SimpleNamespace(messages=_Messages())
|
||||
monkeypatch.setitem(sys.modules, "anthropic", mod)
|
||||
|
||||
|
||||
def _fake_openai(monkeypatch, captured):
|
||||
class _Completions:
|
||||
def create(self, **kw):
|
||||
captured.update(kw)
|
||||
return SimpleNamespace(
|
||||
choices=[SimpleNamespace(
|
||||
message=SimpleNamespace(content=_NODE_JSON), finish_reason="stop")],
|
||||
usage=SimpleNamespace(prompt_tokens=3, completion_tokens=4),
|
||||
)
|
||||
mod = types.ModuleType("openai")
|
||||
mod.OpenAI = lambda **kw: SimpleNamespace(chat=SimpleNamespace(completions=_Completions()))
|
||||
monkeypatch.setitem(sys.modules, "openai", mod)
|
||||
|
||||
|
||||
def _fake_boto3(monkeypatch, captured):
|
||||
class _Client:
|
||||
def converse(self, **kw):
|
||||
captured.update(kw)
|
||||
return {
|
||||
"output": {"message": {"content": [{"text": _NODE_JSON}]}},
|
||||
"usage": {"inputTokens": 1, "outputTokens": 2},
|
||||
"stopReason": "end_turn",
|
||||
}
|
||||
boto3 = types.ModuleType("boto3")
|
||||
boto3.Session = lambda **kw: SimpleNamespace(client=lambda svc: _Client())
|
||||
monkeypatch.setitem(sys.modules, "boto3", boto3)
|
||||
botocore = types.ModuleType("botocore")
|
||||
exc = types.ModuleType("botocore.exceptions")
|
||||
exc.ClientError = type("ClientError", (Exception,), {})
|
||||
botocore.exceptions = exc
|
||||
monkeypatch.setitem(sys.modules, "botocore", botocore)
|
||||
monkeypatch.setitem(sys.modules, "botocore.exceptions", exc)
|
||||
|
||||
|
||||
# ── backend payload shape (mocked) ────────────────────────────────────────────
|
||||
|
||||
def test_call_claude_sends_image_block(tmp_path, monkeypatch):
|
||||
img, _, _ = _make_corpus(tmp_path)
|
||||
refs = llm._build_image_refs([img], tmp_path)
|
||||
captured: dict = {}
|
||||
_fake_anthropic(monkeypatch, captured)
|
||||
llm._call_claude("k", "claude-sonnet-4-6", "CORPUS", images=refs)
|
||||
content = captured["messages"][0]["content"]
|
||||
assert any(b.get("type") == "image" for b in content)
|
||||
|
||||
|
||||
def test_call_openai_compat_sends_image_url(tmp_path, monkeypatch):
|
||||
img, _, _ = _make_corpus(tmp_path)
|
||||
refs = llm._build_image_refs([img], tmp_path)
|
||||
captured: dict = {}
|
||||
_fake_openai(monkeypatch, captured)
|
||||
llm._call_openai_compat("http://x", "k", "gpt", "CORPUS", images=refs)
|
||||
content = captured["messages"][1]["content"]
|
||||
assert any(p.get("type") == "image_url" for p in content)
|
||||
|
||||
|
||||
def test_call_openai_compat_text_only_without_images(monkeypatch):
|
||||
captured: dict = {}
|
||||
_fake_openai(monkeypatch, captured)
|
||||
llm._call_openai_compat("http://x", "k", "gpt", "CORPUS", images=[])
|
||||
assert captured["messages"][1]["content"] == "CORPUS"
|
||||
|
||||
|
||||
def test_call_bedrock_sends_raw_image_bytes(tmp_path, monkeypatch):
|
||||
img, _, _ = _make_corpus(tmp_path)
|
||||
refs = llm._build_image_refs([img], tmp_path)
|
||||
captured: dict = {}
|
||||
_fake_boto3(monkeypatch, captured)
|
||||
llm._call_bedrock("model", "CORPUS", images=refs)
|
||||
content = captured["messages"][0]["content"]
|
||||
img_block = next(b for b in content if "image" in b)
|
||||
assert img_block["image"]["source"]["bytes"] == _PNG_BYTES
|
||||
|
||||
|
||||
# ── CLI backends (mocked subprocess) ──────────────────────────────────────────
|
||||
|
||||
|
||||
|
||||
def test_claude_cli_adds_dir_and_read_instruction(tmp_path, monkeypatch):
|
||||
img, _, _ = _make_corpus(tmp_path)
|
||||
refs = llm._build_image_refs([img], tmp_path)
|
||||
envelope = {"result": _NODE_JSON, "usage": {"output_tokens": 1}, "stop_reason": "end_turn"}
|
||||
seen: dict = {}
|
||||
|
||||
def fake_run(args, **kw):
|
||||
seen["args"] = args
|
||||
seen["input"] = kw.get("input", "")
|
||||
return MagicMock(returncode=0, stdout=json.dumps(envelope), stderr="")
|
||||
|
||||
monkeypatch.setattr(llm, "_response_is_hollow", lambda raw, parsed: False)
|
||||
with patch("shutil.which", return_value="/fake/claude"), \
|
||||
patch("subprocess.run", side_effect=fake_run):
|
||||
llm._call_claude_cli("CORPUS", images=refs)
|
||||
|
||||
assert "--add-dir" in seen["args"]
|
||||
assert str(refs[0].path.parent) in seen["args"]
|
||||
# the prompt sent on stdin tells the model to Read the image path
|
||||
assert "Read tool" in seen["input"] and str(refs[0].path) in seen["input"]
|
||||
|
||||
|
||||
# ── dispatch-level vision gating ──────────────────────────────────────────────
|
||||
|
||||
def test_extract_files_direct_gates_pixels_by_capability(tmp_path, monkeypatch):
|
||||
img, _, doc = _make_corpus(tmp_path)
|
||||
captured: dict = {}
|
||||
_fake_openai(monkeypatch, captured)
|
||||
|
||||
# vision backend (openai) -> content is a list carrying an image_url block
|
||||
monkeypatch.setenv("OPENAI_API_KEY", "k")
|
||||
llm.extract_files_direct([doc, img], backend="openai", root=tmp_path)
|
||||
assert isinstance(captured["messages"][1]["content"], list)
|
||||
|
||||
# non-vision backend (deepseek) -> pixels stripped, content is a plain string
|
||||
captured.clear()
|
||||
monkeypatch.setenv("DEEPSEEK_API_KEY", "k")
|
||||
llm.extract_files_direct([doc, img], backend="deepseek", root=tmp_path)
|
||||
content = captured["messages"][1]["content"]
|
||||
assert isinstance(content, str) and "sub/diagram.png" in content
|
||||
+76
-1
@@ -8,7 +8,7 @@ from graphify.extract import (
|
||||
extract_swift, extract_go, extract_julia, extract_js, extract_fortran,
|
||||
extract_groovy, extract_sln, extract_csproj, extract_razor,
|
||||
extract_dm, extract_dmi, extract_dmm, extract_dmf,
|
||||
extract_powershell,
|
||||
extract_powershell, extract_apex,
|
||||
)
|
||||
|
||||
FIXTURES = Path(__file__).parent / "fixtures"
|
||||
@@ -1553,3 +1553,78 @@ def test_razor_no_dangling_edges():
|
||||
node_ids = {n["id"] for n in r["nodes"]}
|
||||
for e in r["edges"]:
|
||||
assert e["source"] in node_ids
|
||||
|
||||
|
||||
# ---------------Salesforce Apex (.cls / .trigger)----------------------
|
||||
|
||||
def test_apex_class_extraction():
|
||||
r = extract_apex(FIXTURES / "sample.cls")
|
||||
labels = _labels(r)
|
||||
assert "AccountService" in labels
|
||||
|
||||
def test_apex_enum_extraction():
|
||||
r = extract_apex(FIXTURES / "sample.cls")
|
||||
labels = _labels(r)
|
||||
assert "AccountStatus" in labels
|
||||
|
||||
def test_apex_interface_extraction():
|
||||
r = extract_apex(FIXTURES / "sample.cls")
|
||||
labels = _labels(r)
|
||||
assert "Notifiable" in labels
|
||||
|
||||
def test_apex_method_extraction():
|
||||
r = extract_apex(FIXTURES / "sample.cls")
|
||||
labels = _labels(r)
|
||||
assert any("getAccounts" in l for l in labels)
|
||||
assert any("updateAccountsAsync" in l for l in labels)
|
||||
assert any("createAccounts" in l for l in labels)
|
||||
assert any("deleteOldAccounts" in l for l in labels)
|
||||
|
||||
def test_apex_contains_and_method_relations():
|
||||
r = extract_apex(FIXTURES / "sample.cls")
|
||||
relations = _relations(r)
|
||||
assert "contains" in relations
|
||||
assert "method" in relations
|
||||
|
||||
def test_apex_soql_uses_edge():
|
||||
r = extract_apex(FIXTURES / "sample.cls")
|
||||
relations = _relations(r)
|
||||
assert "uses" in relations
|
||||
labels = _labels(r)
|
||||
assert "Account" in labels
|
||||
|
||||
def test_apex_dml_uses_edge():
|
||||
r = extract_apex(FIXTURES / "sample.cls")
|
||||
dml_labels = {n["label"] for n in r["nodes"] if n["label"] in ("insert", "update", "delete", "upsert")}
|
||||
assert len(dml_labels) > 0
|
||||
|
||||
def test_apex_file_node_present():
|
||||
r = extract_apex(FIXTURES / "sample.cls")
|
||||
labels = _labels(r)
|
||||
assert "sample.cls" in labels
|
||||
|
||||
def test_apex_trigger_extraction():
|
||||
r = extract_apex(FIXTURES / "sample.trigger")
|
||||
labels = _labels(r)
|
||||
assert "sample.trigger" in labels
|
||||
assert "AccountTrigger" in labels
|
||||
|
||||
def test_apex_trigger_uses_sobject():
|
||||
r = extract_apex(FIXTURES / "sample.trigger")
|
||||
relations = _relations(r)
|
||||
assert "uses" in relations
|
||||
labels = _labels(r)
|
||||
assert "Account" in labels
|
||||
|
||||
def test_apex_missing_file_returns_empty():
|
||||
r = extract_apex(Path("nonexistent.cls"))
|
||||
assert r["nodes"] == []
|
||||
assert r["edges"] == []
|
||||
|
||||
def test_apex_no_dangling_edges():
|
||||
for fixture in ("sample.cls", "sample.trigger"):
|
||||
r = extract_apex(FIXTURES / fixture)
|
||||
node_ids = {n["id"] for n in r["nodes"]}
|
||||
for e in r["edges"]:
|
||||
assert e["source"] in node_ids, f"dangling source in {fixture}: {e}"
|
||||
assert e["target"] in node_ids, f"dangling target in {fixture}: {e}"
|
||||
|
||||
@@ -15,6 +15,9 @@ def _clear_backend_env(monkeypatch):
|
||||
"MOONSHOT_API_KEY",
|
||||
"ANTHROPIC_API_KEY",
|
||||
"OPENAI_API_KEY",
|
||||
"DEEPSEEK_API_KEY",
|
||||
"AZURE_OPENAI_API_KEY",
|
||||
"AZURE_OPENAI_ENDPOINT",
|
||||
):
|
||||
monkeypatch.delenv(env_key, raising=False)
|
||||
|
||||
@@ -513,3 +516,80 @@ def test_adaptive_retry_bisects_on_hollow_ollama_response(tmp_path):
|
||||
"full chunk came back hollow"
|
||||
)
|
||||
assert calls["n"] == 3 # 1 hollow + 2 successful halves
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Azure backend
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _install_fake_azure_openai(monkeypatch, fake_resp):
|
||||
"""Inject a stub openai module with AzureOpenAI so _call_azure and
|
||||
_azure_client can run without the real SDK installed."""
|
||||
import sys
|
||||
import types
|
||||
|
||||
captured: dict = {}
|
||||
|
||||
class _FakeAzureOpenAI:
|
||||
def __init__(self, *_, **kwargs):
|
||||
captured["init_kwargs"] = kwargs
|
||||
self.chat = self
|
||||
self.completions = self
|
||||
|
||||
def create(self, **kwargs):
|
||||
captured["create_kwargs"] = kwargs
|
||||
return fake_resp
|
||||
|
||||
fake_module = types.ModuleType("openai")
|
||||
fake_module.AzureOpenAI = _FakeAzureOpenAI
|
||||
monkeypatch.setitem(sys.modules, "openai", fake_module)
|
||||
return captured
|
||||
|
||||
|
||||
def test_call_azure_uses_correct_client_params_and_max_completion_tokens(monkeypatch):
|
||||
monkeypatch.setenv("AZURE_OPENAI_API_VERSION", "2024-08-01-preview")
|
||||
monkeypatch.delenv("GRAPHIFY_API_TIMEOUT", raising=False)
|
||||
|
||||
fake_resp = _fake_openai_response(
|
||||
'{"nodes":[{"id":"a"}],"edges":[],"hyperedges":[]}',
|
||||
finish_reason="stop",
|
||||
prompt_tokens=100,
|
||||
completion_tokens=50,
|
||||
)
|
||||
captured = _install_fake_azure_openai(monkeypatch, fake_resp)
|
||||
|
||||
result = llm._call_azure(
|
||||
api_key="test-key",
|
||||
endpoint="https://my-resource.openai.azure.com/",
|
||||
model="gpt-4o",
|
||||
user_message="test",
|
||||
)
|
||||
|
||||
assert captured["init_kwargs"].get("azure_endpoint") == "https://my-resource.openai.azure.com/"
|
||||
assert captured["init_kwargs"].get("api_version") == "2024-08-01-preview"
|
||||
assert "max_completion_tokens" in captured["create_kwargs"], "must use max_completion_tokens not max_tokens"
|
||||
assert "max_tokens" not in captured["create_kwargs"], "deprecated max_tokens must not be sent"
|
||||
assert result["nodes"] == [{"id": "a"}]
|
||||
|
||||
|
||||
def test_detect_backend_returns_azure_when_both_vars_set(monkeypatch):
|
||||
_clear_backend_env(monkeypatch)
|
||||
monkeypatch.setenv("AZURE_OPENAI_API_KEY", "azure-key")
|
||||
monkeypatch.setenv("AZURE_OPENAI_ENDPOINT", "https://my-resource.openai.azure.com/")
|
||||
|
||||
assert llm.detect_backend() == "azure"
|
||||
assert llm._get_backend_api_key("azure") == "azure-key"
|
||||
|
||||
|
||||
def test_detect_backend_azure_requires_endpoint_not_just_key(monkeypatch):
|
||||
_clear_backend_env(monkeypatch)
|
||||
monkeypatch.setenv("AZURE_OPENAI_API_KEY", "azure-key")
|
||||
# AZURE_OPENAI_ENDPOINT already cleared by _clear_backend_env
|
||||
|
||||
assert llm.detect_backend() != "azure"
|
||||
|
||||
|
||||
def test_estimate_cost_azure_no_keyerror():
|
||||
cost = llm.estimate_cost("azure", 1_000_000, 500_000)
|
||||
assert cost == pytest.approx(2.50 + 5.00) # 1M in * $2.50/M + 0.5M out * $10.00/M
|
||||
|
||||
@@ -0,0 +1,275 @@
|
||||
import sys
|
||||
from unittest.mock import MagicMock, patch
|
||||
import pytest
|
||||
from pathlib import Path
|
||||
from graphify.pg_introspect import introspect_postgres
|
||||
from graphify.validate import validate_extraction
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Shared mock infrastructure
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def _make_mock_psycopg(tables, views, routines, fks,
|
||||
host="myhost", dbname="mydb",
|
||||
connect_raises=None):
|
||||
"""Return a mock psycopg module wired to the provided catalog data.
|
||||
|
||||
``routines`` rows must be 5-tuples: (schema, name, rtype, body, ext_lang).
|
||||
``fks`` rows must be 7-tuples:
|
||||
(constraint_name, t_schema, t_name, [cols], r_schema, r_name, [r_cols])
|
||||
``connect_raises``, if set, is an exception *instance* raised by connect().
|
||||
"""
|
||||
|
||||
class MockCursor:
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc_val, exc_tb):
|
||||
pass
|
||||
|
||||
def execute(self, query, params=None):
|
||||
self.query = query
|
||||
|
||||
def fetchall(self):
|
||||
q = self.query.strip().lower()
|
||||
if "information_schema.tables" in q:
|
||||
return tables
|
||||
elif "information_schema.views" in q:
|
||||
return views
|
||||
elif "information_schema.routines" in q:
|
||||
return routines
|
||||
elif "information_schema.referential_constraints" in q:
|
||||
return fks
|
||||
return []
|
||||
|
||||
class MockConnection:
|
||||
def execute(self, query):
|
||||
pass
|
||||
|
||||
def cursor(self):
|
||||
return MockCursor()
|
||||
|
||||
def close(self):
|
||||
pass
|
||||
|
||||
@property
|
||||
def info(self):
|
||||
info_mock = MagicMock()
|
||||
info_mock.dsn = f"host={host} dbname={dbname} user=myuser password=secret"
|
||||
return info_mock
|
||||
|
||||
mock_psycopg = MagicMock()
|
||||
if connect_raises is not None:
|
||||
mock_psycopg.connect.side_effect = connect_raises
|
||||
# Make the exception type available as an attribute so the module can
|
||||
# reference psycopg.OperationalError in the except clause.
|
||||
mock_psycopg.OperationalError = type(connect_raises)
|
||||
else:
|
||||
mock_psycopg.connect.return_value = MockConnection()
|
||||
mock_psycopg.OperationalError = Exception # unused path but must exist
|
||||
mock_psycopg.conninfo.conninfo_to_dict.return_value = {
|
||||
"host": host,
|
||||
"dbname": dbname,
|
||||
}
|
||||
return mock_psycopg
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def _q(schema: str, name: str) -> str:
|
||||
"""Return the label form that tree-sitter produces for a quoted identifier.
|
||||
|
||||
pg_introspect emits CREATE TABLE "schema"."name" — tree-sitter reads the
|
||||
object_reference text verbatim (quotes included), so the node label is
|
||||
'"schema"."name"', not 'schema.name'.
|
||||
"""
|
||||
return f'"{schema}"."{name}"'
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def test_pg_introspect_success():
|
||||
"""Baseline: tables, views, routines, and a single-column FK all survive."""
|
||||
mock_tables = [
|
||||
("public", "users", "BASE TABLE"),
|
||||
("public", "orders", "BASE TABLE"),
|
||||
]
|
||||
mock_views = [
|
||||
("public", "active_users", "SELECT * FROM public.users WHERE active = true"),
|
||||
]
|
||||
# 5-tuple: schema, name, rtype, body, ext_lang
|
||||
mock_routines = [
|
||||
("public", "calculate_total", "FUNCTION", "SELECT 42;", "SQL"),
|
||||
("public", "do_nothing", "PROCEDURE", None, "PLPGSQL"),
|
||||
]
|
||||
# 7-tuple: constraint_name, t_schema, t_name, cols[], r_schema, r_name, r_cols[]
|
||||
mock_fks = [
|
||||
("fk_orders_user_id", "public", "orders", ["user_id"], "public", "users", ["id"]),
|
||||
]
|
||||
|
||||
mock_psycopg = _make_mock_psycopg(mock_tables, mock_views, mock_routines, mock_fks)
|
||||
|
||||
with patch.dict("sys.modules", {"psycopg": mock_psycopg}):
|
||||
res = introspect_postgres("postgresql://myuser:mypassword@myhost/mydb")
|
||||
|
||||
# 1. validate_extraction must pass
|
||||
errors = validate_extraction(res)
|
||||
assert errors == [], f"Validation errors: {errors}"
|
||||
|
||||
# 2. source_file must be the sanitized virtual path (no credentials)
|
||||
expected_source = "postgresql:/myhost/mydb"
|
||||
for node in res["nodes"]:
|
||||
assert node["source_file"] == expected_source
|
||||
for edge in res["edges"]:
|
||||
assert edge["source_file"] == expected_source
|
||||
|
||||
# 3. Expected node labels. pg_introspect double-quotes identifiers in DDL,
|
||||
# so tree-sitter returns the raw quoted text as the object_reference.
|
||||
node_labels = {n["label"] for n in res["nodes"]}
|
||||
assert _q("public", "users") in node_labels, f"users missing; got {node_labels}"
|
||||
assert _q("public", "orders") in node_labels, f"orders missing; got {node_labels}"
|
||||
# Views keep the schema-qualified label (quoted schema, unquoted body)
|
||||
assert _q("public", "active_users") in node_labels, f"active_users missing; got {node_labels}"
|
||||
# Functions: label is "<quoted-sig>()"
|
||||
assert f'{_q("public", "calculate_total")}()' in node_labels, f"calculate_total() missing; got {node_labels}"
|
||||
assert f'{_q("public", "do_nothing")}()' in node_labels, f"do_nothing() missing; got {node_labels}"
|
||||
|
||||
# 4. File node (label = dbname)
|
||||
file_nodes = [n for n in res["nodes"] if n["file_type"] == "code" and n["label"] == "mydb"]
|
||||
assert len(file_nodes) == 1
|
||||
|
||||
# 5. FK references edge: orders → users, exactly once
|
||||
users_nid = next(n["id"] for n in res["nodes"] if n["label"] == _q("public", "users"))
|
||||
orders_nid = next(n["id"] for n in res["nodes"] if n["label"] == _q("public", "orders"))
|
||||
ref_edges = [
|
||||
e for e in res["edges"]
|
||||
if e["source"] == orders_nid and e["target"] == users_nid and e["relation"] == "references"
|
||||
]
|
||||
assert len(ref_edges) == 1, f"Expected exactly 1 references edge, got {len(ref_edges)}"
|
||||
|
||||
|
||||
def test_pg_introspect_quoted_identifiers():
|
||||
"""Reserved-word and special-character table names must survive DDL round-trip.
|
||||
|
||||
'order' is a SQL reserved word; 'user-data' contains a hyphen — both would
|
||||
produce invalid DDL without quoting, causing tree-sitter to silently drop
|
||||
those tables and any FK touching them.
|
||||
"""
|
||||
mock_tables = [
|
||||
("public", "order", "BASE TABLE"), # reserved word
|
||||
("public", "user-data", "BASE TABLE"), # hyphen
|
||||
]
|
||||
mock_views = []
|
||||
mock_routines = []
|
||||
# FK: user-data.owner_id → order.id
|
||||
mock_fks = [
|
||||
("fk_userdata_order", "public", "user-data", ["owner_id"], "public", "order", ["id"]),
|
||||
]
|
||||
|
||||
mock_psycopg = _make_mock_psycopg(mock_tables, mock_views, mock_routines, mock_fks)
|
||||
|
||||
with patch.dict("sys.modules", {"psycopg": mock_psycopg}):
|
||||
res = introspect_postgres("postgresql://myuser:secret@myhost/mydb")
|
||||
|
||||
errors = validate_extraction(res)
|
||||
assert errors == [], f"Validation errors: {errors}"
|
||||
|
||||
node_labels = {n["label"] for n in res["nodes"]}
|
||||
|
||||
# Both tables must appear as nodes (quoted form expected from tree-sitter)
|
||||
assert _q("public", "order") in node_labels, \
|
||||
f"'order' table missing; labels={node_labels}"
|
||||
assert _q("public", "user-data") in node_labels, \
|
||||
f"'user-data' table missing; labels={node_labels}"
|
||||
|
||||
# FK references edge must exist
|
||||
ref_edges = [e for e in res["edges"] if e["relation"] == "references"]
|
||||
assert len(ref_edges) >= 1, "Expected at least one references edge for the FK"
|
||||
|
||||
|
||||
def test_pg_introspect_composite_fk():
|
||||
"""A 2-column composite FK must produce exactly ONE references edge, not two.
|
||||
|
||||
The old code emitted one ADD CONSTRAINT per row of the FK query (one row
|
||||
per key column), causing duplicate edges for composite keys.
|
||||
"""
|
||||
mock_tables = [
|
||||
("public", "products", "BASE TABLE"),
|
||||
("public", "order_items", "BASE TABLE"),
|
||||
]
|
||||
mock_views = []
|
||||
mock_routines = []
|
||||
# Single composite FK: order_items(order_id, product_id) → products(order_id, product_id)
|
||||
mock_fks = [
|
||||
(
|
||||
"fk_order_items_composite",
|
||||
"public", "order_items",
|
||||
["order_id", "product_id"],
|
||||
"public", "products",
|
||||
["order_id", "product_id"],
|
||||
),
|
||||
]
|
||||
|
||||
mock_psycopg = _make_mock_psycopg(mock_tables, mock_views, mock_routines, mock_fks)
|
||||
|
||||
with patch.dict("sys.modules", {"psycopg": mock_psycopg}):
|
||||
res = introspect_postgres("postgresql://myuser:secret@myhost/mydb")
|
||||
|
||||
errors = validate_extraction(res)
|
||||
assert errors == [], f"Validation errors: {errors}"
|
||||
|
||||
products_nid = next(
|
||||
n["id"] for n in res["nodes"] if n["label"] == _q("public", "products")
|
||||
)
|
||||
order_items_nid = next(
|
||||
n["id"] for n in res["nodes"] if n["label"] == _q("public", "order_items")
|
||||
)
|
||||
|
||||
ref_edges = [
|
||||
e for e in res["edges"]
|
||||
if e["source"] == order_items_nid
|
||||
and e["target"] == products_nid
|
||||
and e["relation"] == "references"
|
||||
]
|
||||
assert len(ref_edges) == 1, (
|
||||
f"Expected exactly 1 references edge for composite FK, got {len(ref_edges)}"
|
||||
)
|
||||
|
||||
|
||||
def test_pg_introspect_connection_error():
|
||||
"""A psycopg.OperationalError must be re-raised as ConnectionError with a
|
||||
sanitized message (no DSN/credentials) and no stack-trace noise."""
|
||||
|
||||
class FakeOperationalError(Exception):
|
||||
pass
|
||||
|
||||
raw_error = FakeOperationalError(
|
||||
'connection to server at "myhost" (127.0.0.1), port 5432 failed: '
|
||||
'FATAL: password authentication failed for user "myuser"\n'
|
||||
"DETAIL: Connection matched pg_hba.conf line 1: …"
|
||||
)
|
||||
|
||||
mock_psycopg = _make_mock_psycopg([], [], [], [], connect_raises=raw_error)
|
||||
|
||||
with patch.dict("sys.modules", {"psycopg": mock_psycopg}):
|
||||
with pytest.raises(ConnectionError) as exc_info:
|
||||
introspect_postgres("postgresql://myuser:secret@myhost/mydb")
|
||||
|
||||
msg = str(exc_info.value)
|
||||
assert "could not connect to PostgreSQL" in msg
|
||||
# Credentials must not appear in the surfaced message
|
||||
assert "secret" not in msg
|
||||
# Only the first line of the OperationalError should be present (no DETAIL)
|
||||
assert "DETAIL" not in msg
|
||||
|
||||
|
||||
def test_pg_introspect_import_error():
|
||||
"""If psycopg is missing, introspect_postgres raises ImportError."""
|
||||
with patch.dict("sys.modules", {"psycopg": None}):
|
||||
with pytest.raises(ImportError, match="psycopg is required"):
|
||||
introspect_postgres("postgresql://localhost/db")
|
||||
@@ -207,6 +207,132 @@ def test_rebuild_code_evicts_nodes_from_deleted_files(tmp_path):
|
||||
assert "login()" in node_labels_after, "nodes from surviving file must be kept"
|
||||
|
||||
|
||||
def test_rebuild_code_evicts_removed_symbol_from_surviving_file(tmp_path):
|
||||
"""#1116: graphify update (_rebuild_code with no changed_paths) must prune a
|
||||
symbol removed from a file that still exists — and its inbound call edge —
|
||||
without dropping genuine semantic nodes that share the surviving file."""
|
||||
import json
|
||||
from graphify.watch import _rebuild_code
|
||||
|
||||
corpus = tmp_path / "corpus"
|
||||
corpus.mkdir()
|
||||
|
||||
(corpus / "a.py").write_text(
|
||||
"def foo(): pass\ndef bar(): pass\n", encoding="utf-8"
|
||||
)
|
||||
(corpus / "b.py").write_text(
|
||||
"from a import foo\n\ndef caller():\n foo()\n", encoding="utf-8"
|
||||
)
|
||||
|
||||
assert _rebuild_code(corpus, acquire_lock=False) is True
|
||||
graph_path = corpus / "graphify-out" / "graph.json"
|
||||
data = json.loads(graph_path.read_text(encoding="utf-8"))
|
||||
|
||||
def labels(d):
|
||||
return {n["label"] for n in d.get("nodes", [])}
|
||||
|
||||
def id_for(d, label):
|
||||
return next(n["id"] for n in d.get("nodes", []) if n["label"] == label)
|
||||
|
||||
def edges(d):
|
||||
return d.get("links", d.get("edges", []))
|
||||
|
||||
before = labels(data)
|
||||
assert {"foo()", "bar()", "caller()"} <= before
|
||||
foo_id = id_for(data, "foo()")
|
||||
caller_id = id_for(data, "caller()")
|
||||
assert any(
|
||||
{e.get("source"), e.get("target")} == {caller_id, foo_id}
|
||||
for e in edges(data)
|
||||
), "cross-file caller->foo call edge must exist before removal"
|
||||
|
||||
# Pre-seed a semantic node on the surviving a.py (no AST id, no _origin
|
||||
# marker). A naive "evict every re-extracted file's nodes by source_file"
|
||||
# fix would wrongly delete this; the identity-based fix must keep it.
|
||||
data["nodes"].append({
|
||||
"id": "a_authconcept",
|
||||
"label": "AuthConcept",
|
||||
"file_type": "concept",
|
||||
"source_file": "a.py",
|
||||
})
|
||||
graph_path.write_text(json.dumps(data), encoding="utf-8")
|
||||
|
||||
# Remove foo() from a.py (keep bar); leave b.py untouched.
|
||||
(corpus / "a.py").write_text("def bar(): pass\n", encoding="utf-8")
|
||||
|
||||
assert _rebuild_code(corpus, acquire_lock=False, force=True) is True
|
||||
after_data = json.loads(graph_path.read_text(encoding="utf-8"))
|
||||
after = labels(after_data)
|
||||
|
||||
assert "foo()" not in after, "removed symbol must be pruned from surviving file"
|
||||
assert not any(
|
||||
e.get("source") == foo_id or e.get("target") == foo_id
|
||||
for e in edges(after_data)
|
||||
), "dangling edge to the removed symbol must be dropped"
|
||||
assert "bar()" in after, "surviving symbol in the same file must be kept"
|
||||
assert "caller()" in after, "unchanged file's nodes must be kept"
|
||||
assert "AuthConcept" in after, "semantic node on a surviving file must not be evicted"
|
||||
|
||||
|
||||
def test_rebuild_code_preupgrade_marker_less_node_one_cycle_lag(tmp_path):
|
||||
"""#1118 backward-compat: a graph.json built before #1116 has no `_origin`
|
||||
markers. On the first `graphify update` after upgrading, a symbol removed
|
||||
from a surviving file is NOT pruned that cycle — its old node carries no
|
||||
marker, so the new drop-rule skips it. This is a deliberate one-cycle lag
|
||||
(no data loss); it self-heals once the node has been stamped `_origin="ast"`
|
||||
(which a full re-extraction does for every surviving symbol)."""
|
||||
import json
|
||||
from graphify.watch import _rebuild_code
|
||||
|
||||
corpus = tmp_path / "corpus"
|
||||
corpus.mkdir()
|
||||
(corpus / "a.py").write_text("def bar(): pass\n", encoding="utf-8")
|
||||
|
||||
assert _rebuild_code(corpus, acquire_lock=False) is True
|
||||
graph_path = corpus / "graphify-out" / "graph.json"
|
||||
data = json.loads(graph_path.read_text(encoding="utf-8"))
|
||||
|
||||
def labels(d):
|
||||
return {n["label"] for n in d.get("nodes", [])}
|
||||
|
||||
# Simulate a pre-#1116 graph: strip every `_origin` marker, then inject a
|
||||
# stale AST node for a symbol no longer present in a.py's source — also
|
||||
# marker-less, exactly as a pre-upgrade graph would carry it.
|
||||
for n in data["nodes"]:
|
||||
n.pop("_origin", None)
|
||||
data["nodes"].append({
|
||||
"id": "a_foo",
|
||||
"label": "foo()",
|
||||
"file_type": "function",
|
||||
"source_file": "a.py",
|
||||
})
|
||||
graph_path.write_text(json.dumps(data), encoding="utf-8")
|
||||
|
||||
# First update after "upgrade" (full rebuild, no changed_paths): the stale
|
||||
# node has no marker, so the drop-rule skips it and it survives this cycle.
|
||||
assert _rebuild_code(corpus, acquire_lock=False, force=True) is True
|
||||
after = json.loads(graph_path.read_text(encoding="utf-8"))
|
||||
assert "foo()" in labels(after), (
|
||||
"pre-upgrade marker-less stale node must survive the first update — "
|
||||
"documented one-cycle backward-compat lag (#1118)"
|
||||
)
|
||||
|
||||
# Once stamped (a full re-extraction stamps every surviving symbol), the
|
||||
# drop-rule applies on the next update and the stale node self-heals away.
|
||||
for n in after["nodes"]:
|
||||
if n["label"] == "foo()":
|
||||
n["_origin"] = "ast"
|
||||
graph_path.write_text(json.dumps(after), encoding="utf-8")
|
||||
|
||||
assert _rebuild_code(corpus, acquire_lock=False, force=True) is True
|
||||
healed = json.loads(graph_path.read_text(encoding="utf-8"))
|
||||
assert "foo()" not in labels(healed), (
|
||||
"once carrying _origin=ast, the stale node is pruned on the next "
|
||||
"update (self-heal)"
|
||||
)
|
||||
assert "bar()" in labels(healed), "surviving symbol must be kept throughout"
|
||||
|
||||
|
||||
@pytest.mark.skipif(sys.platform == "win32", reason="fcntl-only (POSIX)")
|
||||
def test_rebuild_lock_non_blocking_does_not_clobber_holder(tmp_path):
|
||||
"""GH-858: a non-blocking caller that fails to acquire the lock must not
|
||||
|
||||
Reference in New Issue
Block a user