shelfmark-rag 1.0.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- chunkers.py +155 -0
- config.py +102 -0
- diff_pack.py +389 -0
- indexer.py +688 -0
- mcp_server.py +224 -0
- pack.py +121 -0
- populate_symbols.py +93 -0
- query.py +102 -0
- report.py +165 -0
- retrieval.py +431 -0
- session_chunker.py +188 -0
- shelfmark_rag-1.0.0.dist-info/METADATA +247 -0
- shelfmark_rag-1.0.0.dist-info/RECORD +18 -0
- shelfmark_rag-1.0.0.dist-info/WHEEL +5 -0
- shelfmark_rag-1.0.0.dist-info/entry_points.txt +4 -0
- shelfmark_rag-1.0.0.dist-info/licenses/LICENSE +21 -0
- shelfmark_rag-1.0.0.dist-info/top_level.txt +12 -0
- symbols.py +228 -0
chunkers.py
ADDED
|
@@ -0,0 +1,155 @@
|
|
|
1
|
+
"""Language-aware chunkers. Each returns list[(start_line, end_line, text, symbol_name)].
|
|
2
|
+
|
|
3
|
+
Keep chunks:
|
|
4
|
+
- Python: one chunk per top-level def/class (via stdlib ast).
|
|
5
|
+
- TS/JS: regex on top-level `export? (async )?(function|class|const NAME =)` blocks, brace-matched.
|
|
6
|
+
- Shell: regex on `NAME() {` blocks through closing `^}`.
|
|
7
|
+
- Config / others: fall back to markdown-style word-count windows.
|
|
8
|
+
"""
|
|
9
|
+
from __future__ import annotations
|
|
10
|
+
|
|
11
|
+
import ast
|
|
12
|
+
import re
|
|
13
|
+
from pathlib import Path
|
|
14
|
+
from typing import List, Tuple
|
|
15
|
+
|
|
16
|
+
Chunk = Tuple[int, int, str, str] # start_line (1-idx), end_line, text, symbol
|
|
17
|
+
|
|
18
|
+
MAX_CHUNK_LINES = 250
|
|
19
|
+
MIN_CHUNK_CHARS = 40
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def chunk_python(text: str) -> List[Chunk]:
|
|
23
|
+
lines = text.splitlines()
|
|
24
|
+
try:
|
|
25
|
+
tree = ast.parse(text)
|
|
26
|
+
except SyntaxError:
|
|
27
|
+
return chunk_fallback(text)
|
|
28
|
+
chunks: List[Chunk] = []
|
|
29
|
+
for node in tree.body:
|
|
30
|
+
if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef)):
|
|
31
|
+
start = getattr(node, "lineno", 1)
|
|
32
|
+
end = getattr(node, "end_lineno", start)
|
|
33
|
+
body = "\n".join(lines[start - 1 : end])
|
|
34
|
+
if len(body) >= MIN_CHUNK_CHARS:
|
|
35
|
+
chunks.append((start, end, body[:8000], node.name))
|
|
36
|
+
if not chunks:
|
|
37
|
+
return chunk_fallback(text)
|
|
38
|
+
return chunks
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
_TS_DECL = re.compile(
|
|
42
|
+
r"^(?P<indent>[ \t]*)"
|
|
43
|
+
r"(?:export\s+(?:default\s+)?)?"
|
|
44
|
+
r"(?:async\s+)?"
|
|
45
|
+
r"(?:function\*?\s+(?P<fn>[A-Za-z_$][\w$]*)"
|
|
46
|
+
r"|class\s+(?P<cls>[A-Za-z_$][\w$]*)"
|
|
47
|
+
r"|(?:const|let|var)\s+(?P<var>[A-Za-z_$][\w$]*)\s*[:=]"
|
|
48
|
+
r"|(?:interface|type|enum)\s+(?P<ty>[A-Za-z_$][\w$]*))",
|
|
49
|
+
re.MULTILINE,
|
|
50
|
+
)
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
def chunk_ts(text: str) -> List[Chunk]:
|
|
54
|
+
lines = text.splitlines()
|
|
55
|
+
matches = list(_TS_DECL.finditer(text))
|
|
56
|
+
if not matches:
|
|
57
|
+
return chunk_fallback(text)
|
|
58
|
+
chunks: List[Chunk] = []
|
|
59
|
+
for i, m in enumerate(matches):
|
|
60
|
+
if m.group("indent"):
|
|
61
|
+
continue # only top-level declarations
|
|
62
|
+
name = m.group("fn") or m.group("cls") or m.group("var") or m.group("ty") or "?"
|
|
63
|
+
start_line = text.count("\n", 0, m.start()) + 1
|
|
64
|
+
end_line = (
|
|
65
|
+
text.count("\n", 0, matches[i + 1].start()) + 1
|
|
66
|
+
if i + 1 < len(matches)
|
|
67
|
+
else len(lines) + 1
|
|
68
|
+
)
|
|
69
|
+
end_line = min(end_line, start_line + MAX_CHUNK_LINES)
|
|
70
|
+
body = "\n".join(lines[start_line - 1 : end_line - 1])
|
|
71
|
+
if len(body) >= MIN_CHUNK_CHARS:
|
|
72
|
+
chunks.append((start_line, end_line - 1, body[:8000], name))
|
|
73
|
+
return chunks or chunk_fallback(text)
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
_SHELL_FN = re.compile(r"^(?P<name>[A-Za-z_][\w]*)\s*\(\)\s*\{", re.MULTILINE)
|
|
77
|
+
|
|
78
|
+
|
|
79
|
+
def chunk_shell(text: str) -> List[Chunk]:
|
|
80
|
+
lines = text.splitlines()
|
|
81
|
+
matches = list(_SHELL_FN.finditer(text))
|
|
82
|
+
if not matches:
|
|
83
|
+
return chunk_fallback(text)
|
|
84
|
+
chunks: List[Chunk] = []
|
|
85
|
+
for m in matches:
|
|
86
|
+
name = m.group("name")
|
|
87
|
+
start = text.count("\n", 0, m.start()) + 1
|
|
88
|
+
depth = 0
|
|
89
|
+
i = start - 1
|
|
90
|
+
end = start
|
|
91
|
+
while i < len(lines):
|
|
92
|
+
depth += lines[i].count("{") - lines[i].count("}")
|
|
93
|
+
i += 1
|
|
94
|
+
if depth <= 0 and i > start - 1:
|
|
95
|
+
end = i
|
|
96
|
+
break
|
|
97
|
+
body = "\n".join(lines[start - 1 : end])
|
|
98
|
+
if len(body) >= MIN_CHUNK_CHARS:
|
|
99
|
+
chunks.append((start, end, body[:8000], name))
|
|
100
|
+
return chunks or chunk_fallback(text)
|
|
101
|
+
|
|
102
|
+
|
|
103
|
+
def chunk_fallback(text: str) -> List[Chunk]:
|
|
104
|
+
"""Word-count window for non-code / unrecognized content."""
|
|
105
|
+
lines = text.splitlines()
|
|
106
|
+
if not lines:
|
|
107
|
+
return []
|
|
108
|
+
chunks: List[Chunk] = []
|
|
109
|
+
i = 0
|
|
110
|
+
while i < len(lines):
|
|
111
|
+
start = i
|
|
112
|
+
word_count = 0
|
|
113
|
+
while i < len(lines) and word_count < 300:
|
|
114
|
+
word_count += len(lines[i].split())
|
|
115
|
+
i += 1
|
|
116
|
+
body = "\n".join(lines[start:i]).strip()
|
|
117
|
+
if body and len(body) >= MIN_CHUNK_CHARS:
|
|
118
|
+
chunks.append((start + 1, i, body[:8000], ""))
|
|
119
|
+
# overlap
|
|
120
|
+
if i >= len(lines):
|
|
121
|
+
break
|
|
122
|
+
i = max(start + 1, i - 6)
|
|
123
|
+
return chunks
|
|
124
|
+
|
|
125
|
+
|
|
126
|
+
def detect_language(path: Path) -> str:
|
|
127
|
+
suf = path.suffix.lower()
|
|
128
|
+
return {
|
|
129
|
+
".py": "python",
|
|
130
|
+
".ts": "typescript",
|
|
131
|
+
".tsx": "typescript",
|
|
132
|
+
".js": "javascript",
|
|
133
|
+
".jsx": "javascript",
|
|
134
|
+
".mjs": "javascript",
|
|
135
|
+
".sh": "shell",
|
|
136
|
+
".bash": "shell",
|
|
137
|
+
".zsh": "shell",
|
|
138
|
+
".md": "markdown",
|
|
139
|
+
".yml": "yaml",
|
|
140
|
+
".yaml": "yaml",
|
|
141
|
+
".json": "json",
|
|
142
|
+
".toml": "toml",
|
|
143
|
+
".rules": "text",
|
|
144
|
+
}.get(suf, "text")
|
|
145
|
+
|
|
146
|
+
|
|
147
|
+
def chunk_file(path: Path, text: str) -> List[Chunk]:
|
|
148
|
+
lang = detect_language(path)
|
|
149
|
+
if lang == "python":
|
|
150
|
+
return chunk_python(text)
|
|
151
|
+
if lang in ("typescript", "javascript"):
|
|
152
|
+
return chunk_ts(text)
|
|
153
|
+
if lang == "shell":
|
|
154
|
+
return chunk_shell(text)
|
|
155
|
+
return chunk_fallback(text)
|
config.py
ADDED
|
@@ -0,0 +1,102 @@
|
|
|
1
|
+
"""Central configuration for shelfmark.
|
|
2
|
+
|
|
3
|
+
Everything the engine needs to know about YOUR machine lives in two places:
|
|
4
|
+
|
|
5
|
+
1. Environment variables (all optional):
|
|
6
|
+
RAG_HOME data dir for index.sqlite etc. (default: ~/.shelfmark)
|
|
7
|
+
RAG_DB explicit index path (default: $RAG_HOME/index.sqlite)
|
|
8
|
+
RAG_SOURCES path to sources.yaml (default: $RAG_HOME/sources.yaml)
|
|
9
|
+
RAG_MODEL sentence-transformers model name (default: intfloat/multilingual-e5-small)
|
|
10
|
+
RAG_DIM embedding dimension of RAG_MODEL (default: 384)
|
|
11
|
+
|
|
12
|
+
2. sources.yaml — what to index (auto-created on first run if missing):
|
|
13
|
+
repos: code repos to index (code + docs + CHANGELOG + git commits)
|
|
14
|
+
sources: [{type: <label>, glob: <pattern>}] markdown corpora (notes, docs, ...)
|
|
15
|
+
code_globs: loose script globs indexed as source_type=workstation-code
|
|
16
|
+
|
|
17
|
+
Import from here; never hardcode paths in engine modules.
|
|
18
|
+
"""
|
|
19
|
+
from __future__ import annotations
|
|
20
|
+
|
|
21
|
+
import os
|
|
22
|
+
import sys
|
|
23
|
+
from pathlib import Path
|
|
24
|
+
|
|
25
|
+
import yaml
|
|
26
|
+
|
|
27
|
+
ROOT = Path(os.environ.get("RAG_HOME", "~/.shelfmark")).expanduser()
|
|
28
|
+
DB = Path(os.environ.get("RAG_DB") or ROOT / "index.sqlite")
|
|
29
|
+
QLOG = ROOT / "queries.sqlite"
|
|
30
|
+
MODEL_NAME = os.environ.get("RAG_MODEL", "intfloat/multilingual-e5-small")
|
|
31
|
+
DIM = int(os.environ.get("RAG_DIM", "384"))
|
|
32
|
+
|
|
33
|
+
SOURCES_FILE = Path(
|
|
34
|
+
os.environ.get("RAG_SOURCES") or ROOT / "sources.yaml"
|
|
35
|
+
).expanduser()
|
|
36
|
+
|
|
37
|
+
_STARTER_SOURCES_YAML = """\
|
|
38
|
+
# shelfmark corpus configuration.
|
|
39
|
+
# Fill in what to index, then rerun. ~ and $VARS are expanded in every path.
|
|
40
|
+
|
|
41
|
+
# Code repos: indexed for source code (py/ts/js/sh), docs/**/*.md, README.md,
|
|
42
|
+
# CHANGELOG.md, docs/specs/**, docs/roadmap.md, and the last 180 days of git
|
|
43
|
+
# commit messages.
|
|
44
|
+
repos: []
|
|
45
|
+
# repos:
|
|
46
|
+
# - ~/dev/my-main-project
|
|
47
|
+
|
|
48
|
+
# Markdown corpora: type is a free label you filter on at query time
|
|
49
|
+
# (--scope <type>).
|
|
50
|
+
sources: []
|
|
51
|
+
# sources:
|
|
52
|
+
# - type: memory
|
|
53
|
+
# glob: ~/notes/memory/**/*.md
|
|
54
|
+
|
|
55
|
+
# Loose scripts outside any repo, indexed as source_type=workstation-code.
|
|
56
|
+
code_globs: []
|
|
57
|
+
# code_globs:
|
|
58
|
+
# - ~/scripts/*.sh
|
|
59
|
+
"""
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
def _expand(p: str) -> str:
|
|
63
|
+
return os.path.expanduser(os.path.expandvars(str(p)))
|
|
64
|
+
|
|
65
|
+
|
|
66
|
+
def _load_sources_file() -> dict:
|
|
67
|
+
if not SOURCES_FILE.exists():
|
|
68
|
+
SOURCES_FILE.parent.mkdir(parents=True, exist_ok=True)
|
|
69
|
+
SOURCES_FILE.write_text(_STARTER_SOURCES_YAML, encoding="utf-8")
|
|
70
|
+
print(
|
|
71
|
+
f"shelfmark: no config found — wrote a starter one to {SOURCES_FILE}. "
|
|
72
|
+
"Add your repos/notes there, then rerun.",
|
|
73
|
+
file=sys.stderr,
|
|
74
|
+
)
|
|
75
|
+
return {}
|
|
76
|
+
try:
|
|
77
|
+
data = yaml.safe_load(SOURCES_FILE.read_text(encoding="utf-8"))
|
|
78
|
+
except yaml.YAMLError as e:
|
|
79
|
+
raise SystemExit(f"shelfmark: malformed {SOURCES_FILE}: {e}")
|
|
80
|
+
if data is None:
|
|
81
|
+
return {}
|
|
82
|
+
if not isinstance(data, dict):
|
|
83
|
+
raise SystemExit(f"shelfmark: {SOURCES_FILE} must be a YAML mapping")
|
|
84
|
+
return data
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
_raw = _load_sources_file()
|
|
88
|
+
|
|
89
|
+
# Code repos: indexed for code files, docs/**/*.md, README, CHANGELOG, specs,
|
|
90
|
+
# roadmap, and recent git commits.
|
|
91
|
+
CURATED_REPOS: list[Path] = [Path(_expand(p)) for p in _raw.get("repos", [])]
|
|
92
|
+
|
|
93
|
+
# Markdown corpora: (source_type, glob). source_type is a free label you filter
|
|
94
|
+
# on at query time (--scope), e.g. memory, notes, docs, standards.
|
|
95
|
+
SOURCES: list[tuple[str, str]] = [
|
|
96
|
+
(str(s["type"]), _expand(s["glob"]))
|
|
97
|
+
for s in _raw.get("sources", [])
|
|
98
|
+
if isinstance(s, dict) and "type" in s and "glob" in s
|
|
99
|
+
]
|
|
100
|
+
|
|
101
|
+
# Loose script globs outside any repo (indexed as workstation-code).
|
|
102
|
+
WORKSTATION_CODE_GLOBS: list[str] = [_expand(g) for g in _raw.get("code_globs", [])]
|
diff_pack.py
ADDED
|
@@ -0,0 +1,389 @@
|
|
|
1
|
+
#!/usr/bin/env python3
|
|
2
|
+
"""Change-scoped retrieval: given a git diff, return relevant context.
|
|
3
|
+
|
|
4
|
+
Input: git diff output (stdin) OR --files <path1,path2> OR --symbols <sym1,sym2> OR --pr <N>
|
|
5
|
+
Output: Markdown bundle with modified symbols + 1-hop callers + tests + standards
|
|
6
|
+
|
|
7
|
+
Usage:
|
|
8
|
+
git diff HEAD~1 | diff_pack.py --budget 3000
|
|
9
|
+
diff_pack.py --pr 579 --budget 3000
|
|
10
|
+
diff_pack.py --files src/commands/play.ts,src/utils/cache.ts "fix player crash"
|
|
11
|
+
diff_pack.py --symbols playerFactory::createPlayer,playerFactory::cleanup --budget 2500
|
|
12
|
+
"""
|
|
13
|
+
from __future__ import annotations
|
|
14
|
+
|
|
15
|
+
import argparse
|
|
16
|
+
import re
|
|
17
|
+
import sqlite3
|
|
18
|
+
import subprocess
|
|
19
|
+
import sys
|
|
20
|
+
from pathlib import Path
|
|
21
|
+
from typing import NamedTuple
|
|
22
|
+
|
|
23
|
+
sys.path.insert(0, str(Path(__file__).parent))
|
|
24
|
+
from retrieval import search, cwd_repo
|
|
25
|
+
|
|
26
|
+
# Unified diff hunk header pattern: @@ -L1,C1 +L2,C2 @@
|
|
27
|
+
_HUNK_HEADER = re.compile(r"^@@ -(\d+)(?:,\d+)? \+(\d+)(?:,\d+)? @@")
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
class FileChange(NamedTuple):
|
|
31
|
+
path: str
|
|
32
|
+
old_start: int
|
|
33
|
+
new_start: int
|
|
34
|
+
added_lines: set[int]
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
def parse_diff(diff_text: str) -> dict[str, FileChange]:
|
|
38
|
+
"""Parse unified diff output to extract file + changed line ranges.
|
|
39
|
+
|
|
40
|
+
Returns: {file_path: FileChange}
|
|
41
|
+
"""
|
|
42
|
+
changes: dict[str, FileChange] = {}
|
|
43
|
+
current_file = None
|
|
44
|
+
current_change = None
|
|
45
|
+
line_num = 0
|
|
46
|
+
|
|
47
|
+
for line in diff_text.splitlines():
|
|
48
|
+
# Detect file headers (--- a/path or +++ b/path)
|
|
49
|
+
if line.startswith("--- a/"):
|
|
50
|
+
current_file = line[6:]
|
|
51
|
+
current_change = None
|
|
52
|
+
elif line.startswith("--- /dev/null"):
|
|
53
|
+
# New file case
|
|
54
|
+
current_file = None
|
|
55
|
+
current_change = None
|
|
56
|
+
elif line.startswith("+++ b/"):
|
|
57
|
+
# File to be added or modified
|
|
58
|
+
path = line[6:]
|
|
59
|
+
current_change = FileChange(path=path, old_start=0, new_start=0, added_lines=set())
|
|
60
|
+
changes[path] = current_change
|
|
61
|
+
elif line.startswith("diff --git"):
|
|
62
|
+
# Reset on new file in diff
|
|
63
|
+
current_file = None
|
|
64
|
+
current_change = None
|
|
65
|
+
|
|
66
|
+
# Detect hunk headers to reset line numbering
|
|
67
|
+
elif line.startswith("@@"):
|
|
68
|
+
match = _HUNK_HEADER.match(line)
|
|
69
|
+
if match and current_change:
|
|
70
|
+
current_change = FileChange(
|
|
71
|
+
path=current_change.path,
|
|
72
|
+
old_start=int(match.group(1)),
|
|
73
|
+
new_start=int(match.group(2)),
|
|
74
|
+
added_lines=set(),
|
|
75
|
+
)
|
|
76
|
+
changes[current_change.path] = current_change
|
|
77
|
+
line_num = int(match.group(2)) - 1
|
|
78
|
+
|
|
79
|
+
# Track added lines (those starting with +, excluding +++b/...)
|
|
80
|
+
elif line.startswith("+") and not line.startswith("+++"):
|
|
81
|
+
if current_change:
|
|
82
|
+
line_num += 1
|
|
83
|
+
current_change.added_lines.add(line_num)
|
|
84
|
+
# Other lines (context or deletions) still advance line counter
|
|
85
|
+
elif line.startswith("-") and not line.startswith("---"):
|
|
86
|
+
if current_change:
|
|
87
|
+
line_num += 1
|
|
88
|
+
else:
|
|
89
|
+
if current_change and line and not line.startswith("\\"):
|
|
90
|
+
line_num += 1
|
|
91
|
+
|
|
92
|
+
return changes
|
|
93
|
+
|
|
94
|
+
|
|
95
|
+
def fetch_pr_diff(pr_number: int) -> str:
|
|
96
|
+
"""Fetch git diff for a GitHub PR using gh CLI."""
|
|
97
|
+
try:
|
|
98
|
+
result = subprocess.run(
|
|
99
|
+
["gh", "pr", "diff", str(pr_number)],
|
|
100
|
+
capture_output=True,
|
|
101
|
+
text=True,
|
|
102
|
+
check=True,
|
|
103
|
+
)
|
|
104
|
+
return result.stdout
|
|
105
|
+
except (subprocess.CalledProcessError, FileNotFoundError) as e:
|
|
106
|
+
print(f"Error fetching PR {pr_number}: {e}", file=sys.stderr)
|
|
107
|
+
return ""
|
|
108
|
+
|
|
109
|
+
|
|
110
|
+
def find_symbols_in_ranges(
|
|
111
|
+
db_path: Path, file_path: str, line_ranges: set[int]
|
|
112
|
+
) -> list[str]:
|
|
113
|
+
"""Find symbol definitions that overlap with changed lines.
|
|
114
|
+
|
|
115
|
+
Returns: list of "file:symbol" strings
|
|
116
|
+
"""
|
|
117
|
+
if not line_ranges:
|
|
118
|
+
return []
|
|
119
|
+
|
|
120
|
+
conn = sqlite3.connect(db_path)
|
|
121
|
+
# Find symbols where [start_line, end_line] overlaps any changed line
|
|
122
|
+
min_line = min(line_ranges)
|
|
123
|
+
max_line = max(line_ranges)
|
|
124
|
+
|
|
125
|
+
symbols = []
|
|
126
|
+
sql = """
|
|
127
|
+
SELECT symbol_name FROM symbols_definitions
|
|
128
|
+
WHERE file = ? AND start_line <= ? AND end_line >= ?
|
|
129
|
+
"""
|
|
130
|
+
|
|
131
|
+
for line in sorted(line_ranges):
|
|
132
|
+
cur = conn.execute(sql, (file_path, max_line, min_line))
|
|
133
|
+
for (sym,) in cur:
|
|
134
|
+
symbol_id = f"{file_path}::{sym}"
|
|
135
|
+
if symbol_id not in symbols:
|
|
136
|
+
symbols.append(symbol_id)
|
|
137
|
+
|
|
138
|
+
conn.close()
|
|
139
|
+
return symbols
|
|
140
|
+
|
|
141
|
+
|
|
142
|
+
def get_callers(db_path: Path, symbol_id: str, max_depth: int = 1) -> list[str]:
|
|
143
|
+
"""Fetch 1-hop callers of a symbol via call-graph."""
|
|
144
|
+
callers = []
|
|
145
|
+
if max_depth <= 0:
|
|
146
|
+
return callers
|
|
147
|
+
|
|
148
|
+
conn = sqlite3.connect(db_path)
|
|
149
|
+
# symbol_id is "file::name"; look it up in symbols_called_by
|
|
150
|
+
symbol_name = symbol_id.split("::")[-1] if "::" in symbol_id else symbol_id
|
|
151
|
+
|
|
152
|
+
sql = """
|
|
153
|
+
SELECT DISTINCT caller FROM symbols_called_by
|
|
154
|
+
WHERE callee = ? OR callee LIKE ?
|
|
155
|
+
"""
|
|
156
|
+
# Match both exact "file::name" and "name" patterns
|
|
157
|
+
cur = conn.execute(sql, (symbol_id, f"%.{symbol_name}"))
|
|
158
|
+
for (caller,) in cur:
|
|
159
|
+
if caller not in callers:
|
|
160
|
+
callers.append(caller)
|
|
161
|
+
|
|
162
|
+
conn.close()
|
|
163
|
+
return callers
|
|
164
|
+
|
|
165
|
+
|
|
166
|
+
def find_test_files(
|
|
167
|
+
db_path: Path, modified_symbols: list[str]
|
|
168
|
+
) -> list[str]:
|
|
169
|
+
"""Find test files that import or reference modified symbols."""
|
|
170
|
+
test_paths = set()
|
|
171
|
+
conn = sqlite3.connect(db_path)
|
|
172
|
+
|
|
173
|
+
for symbol_id in modified_symbols:
|
|
174
|
+
symbol_name = symbol_id.split("::")[-1] if "::" in symbol_id else symbol_id
|
|
175
|
+
# Look for test chunks that mention this symbol
|
|
176
|
+
sql = """
|
|
177
|
+
SELECT DISTINCT path FROM chunks
|
|
178
|
+
WHERE (path LIKE '%.test.ts' OR path LIKE '%.spec.ts' OR
|
|
179
|
+
path LIKE '%_test.py' OR path LIKE '%.spec.tsx')
|
|
180
|
+
AND (text LIKE ? OR text LIKE ?)
|
|
181
|
+
"""
|
|
182
|
+
cur = conn.execute(sql, (f"%{symbol_name}%", f"%{symbol_id}%"))
|
|
183
|
+
for (path,) in cur:
|
|
184
|
+
test_paths.add(path)
|
|
185
|
+
|
|
186
|
+
conn.close()
|
|
187
|
+
return list(test_paths)
|
|
188
|
+
|
|
189
|
+
|
|
190
|
+
def chunks_for_files(
|
|
191
|
+
file_paths: list[str], query: str, budget_tokens: int
|
|
192
|
+
) -> tuple[list[dict], int]:
|
|
193
|
+
"""Search for chunks in specific files, respecting budget."""
|
|
194
|
+
chunks = []
|
|
195
|
+
budget = budget_tokens
|
|
196
|
+
|
|
197
|
+
for fpath in file_paths:
|
|
198
|
+
# Search with file hint
|
|
199
|
+
search_query = f"{query} {fpath}" if query else fpath
|
|
200
|
+
results = search(search_query, top=3, cwd=None)
|
|
201
|
+
results = [r for r in results if fpath in r["path"]]
|
|
202
|
+
|
|
203
|
+
for r in results:
|
|
204
|
+
snippet_size = min(900, budget // 4) # Rough token to char conversion
|
|
205
|
+
snippet = r["text"][:snippet_size]
|
|
206
|
+
cost = max(1, len(snippet) // 4)
|
|
207
|
+
if budget - cost < 0:
|
|
208
|
+
break
|
|
209
|
+
|
|
210
|
+
chunks.append(r)
|
|
211
|
+
budget -= cost
|
|
212
|
+
|
|
213
|
+
if budget <= 100: # Stop when budget is low
|
|
214
|
+
break
|
|
215
|
+
|
|
216
|
+
return chunks, budget
|
|
217
|
+
|
|
218
|
+
|
|
219
|
+
def build_pack(
|
|
220
|
+
changes: dict[str, FileChange],
|
|
221
|
+
query: str = "",
|
|
222
|
+
budget_tokens: int = 3000,
|
|
223
|
+
) -> str:
|
|
224
|
+
"""Build context pack from diff changes."""
|
|
225
|
+
from config import DB as db_path
|
|
226
|
+
if not db_path.exists():
|
|
227
|
+
return "(error: rag-index database not found)"
|
|
228
|
+
|
|
229
|
+
budget = budget_tokens
|
|
230
|
+
out: list[str] = []
|
|
231
|
+
|
|
232
|
+
if not changes:
|
|
233
|
+
return "(no changes detected)"
|
|
234
|
+
|
|
235
|
+
# 1. Identify modified symbols
|
|
236
|
+
all_modified_symbols: list[str] = []
|
|
237
|
+
for file_path, change in changes.items():
|
|
238
|
+
symbols = find_symbols_in_ranges(db_path, file_path, change.added_lines)
|
|
239
|
+
all_modified_symbols.extend(symbols)
|
|
240
|
+
|
|
241
|
+
# 2. Identify 1-hop callers
|
|
242
|
+
all_callers: list[str] = []
|
|
243
|
+
for symbol in all_modified_symbols:
|
|
244
|
+
callers = get_callers(db_path, symbol)
|
|
245
|
+
all_callers.extend(callers)
|
|
246
|
+
|
|
247
|
+
# 3. Find test files
|
|
248
|
+
test_files = find_test_files(db_path, all_modified_symbols)
|
|
249
|
+
|
|
250
|
+
# 4. Emit sections respecting budget
|
|
251
|
+
def emit_section(header: str, file_list: list[str], search_query: str, per_chunk_cap: int) -> None:
|
|
252
|
+
nonlocal budget
|
|
253
|
+
if not file_list or budget <= 100:
|
|
254
|
+
return
|
|
255
|
+
out.append(f"## {header}\n")
|
|
256
|
+
chunks, budget = chunks_for_files(file_list, search_query, budget)
|
|
257
|
+
for r in chunks:
|
|
258
|
+
snippet = r["text"][:per_chunk_cap]
|
|
259
|
+
sym = f"::{r['symbol']}" if r.get("symbol") else ""
|
|
260
|
+
repo = f"/{r['repo']}" if r.get("repo") else ""
|
|
261
|
+
header_line = (
|
|
262
|
+
f"### {r['source_type']}{repo}{sym} `{r['path']}:{r['start_line']}-{r['end_line']}`"
|
|
263
|
+
)
|
|
264
|
+
block = f"{header_line}\n```\n{snippet}\n```\n"
|
|
265
|
+
out.append(block)
|
|
266
|
+
|
|
267
|
+
# Emit modified symbols
|
|
268
|
+
if all_modified_symbols and budget > 100:
|
|
269
|
+
modified_files = list(set(f.split("::")[0] for f in all_modified_symbols))
|
|
270
|
+
emit_section("Modified symbols", modified_files, query or "implementation", 900)
|
|
271
|
+
|
|
272
|
+
# Emit callers
|
|
273
|
+
if all_callers and budget > 100:
|
|
274
|
+
caller_files = list(set(f.split("::")[0] for f in all_callers))
|
|
275
|
+
emit_section("Callers (1-hop)", caller_files, query or "caller context", 800)
|
|
276
|
+
|
|
277
|
+
# Emit tests
|
|
278
|
+
if test_files and budget > 100:
|
|
279
|
+
emit_section("Tests", test_files, query or "test", 700)
|
|
280
|
+
|
|
281
|
+
# 5. Apply standards
|
|
282
|
+
if query and budget > 100:
|
|
283
|
+
out.append("## Applicable standards\n")
|
|
284
|
+
try:
|
|
285
|
+
std_chunks = search(query, top=2, scope_types=["standards"], cwd=None)
|
|
286
|
+
for r in std_chunks:
|
|
287
|
+
snippet = r["text"][:600]
|
|
288
|
+
block = f"### `{r['path']}:{r['start_line']}-{r['end_line']}`\n```\n{snippet}\n```\n"
|
|
289
|
+
cost = max(1, len(block) // 4)
|
|
290
|
+
if budget - cost >= 0:
|
|
291
|
+
out.append(block)
|
|
292
|
+
budget -= cost
|
|
293
|
+
except Exception:
|
|
294
|
+
pass
|
|
295
|
+
|
|
296
|
+
if not out:
|
|
297
|
+
return "(no relevant context found)"
|
|
298
|
+
|
|
299
|
+
remaining = max(0, budget)
|
|
300
|
+
header = f"# Diff-scoped context pack\n_Budget: {budget_tokens} tokens · remaining ≈ {remaining}_\n"
|
|
301
|
+
return header + "\n" + "\n".join(out)
|
|
302
|
+
|
|
303
|
+
|
|
304
|
+
def main() -> int:
|
|
305
|
+
ap = argparse.ArgumentParser(
|
|
306
|
+
description="Change-scoped retrieval from git diff or file list"
|
|
307
|
+
)
|
|
308
|
+
input_group = ap.add_mutually_exclusive_group(required=True)
|
|
309
|
+
input_group.add_argument(
|
|
310
|
+
"--diff",
|
|
311
|
+
action="store_true",
|
|
312
|
+
help="read unified diff from stdin",
|
|
313
|
+
)
|
|
314
|
+
input_group.add_argument(
|
|
315
|
+
"--pr",
|
|
316
|
+
type=int,
|
|
317
|
+
help="fetch diff for GitHub PR by number",
|
|
318
|
+
)
|
|
319
|
+
input_group.add_argument(
|
|
320
|
+
"--files",
|
|
321
|
+
help="comma-separated file paths",
|
|
322
|
+
)
|
|
323
|
+
input_group.add_argument(
|
|
324
|
+
"--symbols",
|
|
325
|
+
help="comma-separated symbol names (file::symbol format)",
|
|
326
|
+
)
|
|
327
|
+
|
|
328
|
+
ap.add_argument(
|
|
329
|
+
"--query",
|
|
330
|
+
default="",
|
|
331
|
+
help="optional natural language context (for standards matching)",
|
|
332
|
+
)
|
|
333
|
+
ap.add_argument(
|
|
334
|
+
"--budget",
|
|
335
|
+
type=int,
|
|
336
|
+
default=3000,
|
|
337
|
+
help="token budget (chars/4)",
|
|
338
|
+
)
|
|
339
|
+
args = ap.parse_args()
|
|
340
|
+
|
|
341
|
+
changes: dict[str, FileChange] = {}
|
|
342
|
+
|
|
343
|
+
if args.diff:
|
|
344
|
+
diff_text = sys.stdin.read()
|
|
345
|
+
changes = parse_diff(diff_text)
|
|
346
|
+
elif args.pr:
|
|
347
|
+
diff_text = fetch_pr_diff(args.pr)
|
|
348
|
+
if diff_text:
|
|
349
|
+
changes = parse_diff(diff_text)
|
|
350
|
+
elif args.files:
|
|
351
|
+
# Treat --files as list of paths with all lines in range
|
|
352
|
+
for fpath in args.files.split(","):
|
|
353
|
+
fpath = fpath.strip()
|
|
354
|
+
p = Path(fpath).expanduser().resolve()
|
|
355
|
+
if p.exists():
|
|
356
|
+
# All lines in file are "changed" for context purposes
|
|
357
|
+
num_lines = len(p.read_text(errors="replace").splitlines())
|
|
358
|
+
changes[str(p)] = FileChange(
|
|
359
|
+
path=str(p),
|
|
360
|
+
old_start=1,
|
|
361
|
+
new_start=1,
|
|
362
|
+
added_lines=set(range(1, num_lines + 1)),
|
|
363
|
+
)
|
|
364
|
+
elif args.symbols:
|
|
365
|
+
# Direct symbol names: "file::symbol" format
|
|
366
|
+
# Create fake FileChange entries to trigger symbol lookup
|
|
367
|
+
for sym in args.symbols.split(","):
|
|
368
|
+
sym = sym.strip()
|
|
369
|
+
if "::" in sym:
|
|
370
|
+
file_part, sym_part = sym.rsplit("::", 1)
|
|
371
|
+
# For direct symbols, we assume lines 1-999 as range
|
|
372
|
+
changes[file_part] = FileChange(
|
|
373
|
+
path=file_part,
|
|
374
|
+
old_start=1,
|
|
375
|
+
new_start=1,
|
|
376
|
+
added_lines=set(range(1, 1000)),
|
|
377
|
+
)
|
|
378
|
+
|
|
379
|
+
if not changes:
|
|
380
|
+
print("(no changes detected)", file=sys.stderr)
|
|
381
|
+
return 1
|
|
382
|
+
|
|
383
|
+
pack_output = build_pack(changes, query=args.query, budget_tokens=args.budget)
|
|
384
|
+
print(pack_output)
|
|
385
|
+
return 0
|
|
386
|
+
|
|
387
|
+
|
|
388
|
+
if __name__ == "__main__":
|
|
389
|
+
sys.exit(main())
|