shelldeck 0.0.1__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.
- shelldeck/__init__.py +0 -0
- shelldeck/__main__.py +4 -0
- shelldeck/addons/__init__.py +1 -0
- shelldeck/addons/fast_context.py +1346 -0
- shelldeck/auth.py +207 -0
- shelldeck/cli.py +587 -0
- shelldeck/db.py +794 -0
- shelldeck/gitgraph.py +150 -0
- shelldeck/integration/bash.sh +27 -0
- shelldeck/integration/shelldeck.fish +5 -0
- shelldeck/integration/shelldeck.ps1 +35 -0
- shelldeck/integration/zsh/.zshenv +6 -0
- shelldeck/integration/zsh/.zshrc +10 -0
- shelldeck/pty.py +352 -0
- shelldeck/runner.py +75 -0
- shelldeck/scheduler.py +283 -0
- shelldeck/server.py +1216 -0
- shelldeck/shells.py +174 -0
- shelldeck/static/app.css +2288 -0
- shelldeck/static/app.js +1863 -0
- shelldeck/static/gitgraph.js +142 -0
- shelldeck/static/history.js +172 -0
- shelldeck/static/icon-192.png +0 -0
- shelldeck/static/icon-32.png +0 -0
- shelldeck/static/icon-512.png +0 -0
- shelldeck/static/icon.svg +1 -0
- shelldeck/static/index.html +102 -0
- shelldeck/static/manifest.webmanifest +1 -0
- shelldeck/static/monitor.js +143 -0
- shelldeck/static/ui.js +316 -0
- shelldeck/static/vendor/LICENSE-xterm.txt +21 -0
- shelldeck/static/vendor/addon-fit.js +2 -0
- shelldeck/static/vendor/addon-search.js +2 -0
- shelldeck/static/vendor/addon-serialize.js +2 -0
- shelldeck/static/vendor/addon-web-links.js +2 -0
- shelldeck/static/vendor/addon-webgl.js +2 -0
- shelldeck/static/vendor/xterm.css +218 -0
- shelldeck/static/vendor/xterm.js +2 -0
- shelldeck/static/views.js +658 -0
- shelldeck/stats.py +128 -0
- shelldeck-0.0.1.dist-info/METADATA +254 -0
- shelldeck-0.0.1.dist-info/RECORD +45 -0
- shelldeck-0.0.1.dist-info/WHEEL +4 -0
- shelldeck-0.0.1.dist-info/entry_points.txt +3 -0
- shelldeck-0.0.1.dist-info/licenses/LICENSE +21 -0
|
@@ -0,0 +1,1346 @@
|
|
|
1
|
+
#!/usr/bin/env python3
|
|
2
|
+
"""
|
|
3
|
+
fast_context.py - LLM-free "Fast Context" code retrieval, with confidence scoring.
|
|
4
|
+
|
|
5
|
+
Shelldeck addon for semantic code search. Can be used standalone or via `sd search`.
|
|
6
|
+
|
|
7
|
+
python -m shelldeck.addons.fast_context "lazy load"
|
|
8
|
+
python -m shelldeck.addons.fast_context "where is the auth token validated" --root ~/code/app --top 5
|
|
9
|
+
python -m shelldeck.addons.fast_context "retry with backoff" --ext py,ts --max-tokens 3000 --format json
|
|
10
|
+
|
|
11
|
+
How it works (same shape as the SWE-grep subagent in the docs: grep/read/glob tools,
|
|
12
|
+
parallel calls, at most 4 turns) but with heuristics instead of a model:
|
|
13
|
+
|
|
14
|
+
INDEX A persistent SQLite inverted index (token -> files, with term frequencies) plus a
|
|
15
|
+
symbol table (functions/classes). It is built once (multi-process) and afterwards
|
|
16
|
+
only files whose mtime/size changed are re-indexed, so repeat queries are ~instant.
|
|
17
|
+
TURN 1 Query -> weighted terms: identifiers, quoted phrases, stemmed words, adjacent-word
|
|
18
|
+
n-grams in every identifier style ("lazy load" -> lazyload / lazy_load / lazy-load).
|
|
19
|
+
Terms are looked up in the index (no full-repo scan), ranked with BM25 + path +
|
|
20
|
+
definition bonuses, then the top candidates are read in parallel threads to compute
|
|
21
|
+
real line-level evidence (term proximity, definitions, coverage).
|
|
22
|
+
TURN 2+ Only if confidence is not already high: synonyms + identifiers harvested from the
|
|
23
|
+
best hits + import neighbours are added as new terms and everything is re-ranked.
|
|
24
|
+
Stops early on high confidence or when the ranking stops changing.
|
|
25
|
+
OUTPUT Files, line ranges, enclosing symbol, numbered code, and a 0-100 confidence per file
|
|
26
|
+
and for the whole answer, with human-readable reasons.
|
|
27
|
+
"""
|
|
28
|
+
|
|
29
|
+
from __future__ import annotations
|
|
30
|
+
|
|
31
|
+
import argparse
|
|
32
|
+
import bisect
|
|
33
|
+
import fnmatch
|
|
34
|
+
import hashlib
|
|
35
|
+
import heapq
|
|
36
|
+
import json
|
|
37
|
+
import math
|
|
38
|
+
import os
|
|
39
|
+
import re
|
|
40
|
+
import sqlite3
|
|
41
|
+
import subprocess
|
|
42
|
+
import sys
|
|
43
|
+
import tempfile
|
|
44
|
+
import time
|
|
45
|
+
from collections import Counter, defaultdict
|
|
46
|
+
from concurrent.futures import ProcessPoolExecutor, ThreadPoolExecutor
|
|
47
|
+
from dataclasses import dataclass, field
|
|
48
|
+
from functools import lru_cache
|
|
49
|
+
from itertools import accumulate
|
|
50
|
+
|
|
51
|
+
VERSION = 4 # bump to force a re-index when the schema changes
|
|
52
|
+
MAX_TURNS = 4
|
|
53
|
+
MAX_TERMS_PER_TURN = 8 # new expansion terms per later turn
|
|
54
|
+
MAX_FILE_BYTES = 1_000_000
|
|
55
|
+
MAX_TOKENS_PER_TERM = 1500 # vocabulary matches considered per term
|
|
56
|
+
EARLY_EXIT_CONFIDENCE = 85.0
|
|
57
|
+
K1, B = 1.2, 0.75 # BM25
|
|
58
|
+
|
|
59
|
+
# --------------------------------------------------------------------------- #
|
|
60
|
+
# Constants
|
|
61
|
+
# --------------------------------------------------------------------------- #
|
|
62
|
+
IGNORE_DIRS = {
|
|
63
|
+
".git", ".hg", ".svn", "node_modules", "__pycache__", ".venv", "venv", "env", "dist",
|
|
64
|
+
"build", "target", ".next", ".nuxt", ".idea", ".vscode", ".tox", ".mypy_cache",
|
|
65
|
+
".pytest_cache", ".gradle", "vendor", "coverage", ".cache", "site-packages", "dist-packages",
|
|
66
|
+
}
|
|
67
|
+
SKIP_NAMES = {"package-lock.json", "yarn.lock", "pnpm-lock.yaml", "poetry.lock", "Cargo.lock",
|
|
68
|
+
"composer.lock", "Pipfile.lock", "go.sum"}
|
|
69
|
+
BINARY_EXT = {
|
|
70
|
+
".png", ".jpg", ".jpeg", ".gif", ".bmp", ".ico", ".webp", ".pdf", ".zip", ".gz", ".tar",
|
|
71
|
+
".tgz", ".bz2", ".xz", ".7z", ".jar", ".class", ".exe", ".dll", ".so", ".dylib", ".o",
|
|
72
|
+
".a", ".pyc", ".pyo", ".woff", ".woff2", ".ttf", ".eot", ".mp3", ".mp4", ".mov", ".avi",
|
|
73
|
+
".wasm", ".lock", ".map", ".svg", ".sqlite", ".db", ".bin", ".pkl", ".npy", ".parquet",
|
|
74
|
+
}
|
|
75
|
+
STOPWORDS = set("""
|
|
76
|
+
a an the and or but if then else of to in on at by for with from into onto as is are was were be been
|
|
77
|
+
being do does did done doing have has had having it its this that these those there here where when
|
|
78
|
+
what which who whom how why can could should would will shall may might must not no yes we our us you
|
|
79
|
+
your they them their i me my find show get set use used using make made makes code function functions
|
|
80
|
+
class classes method methods file files implement implementation implemented handle handles handling
|
|
81
|
+
logic please about all any some each every
|
|
82
|
+
""".split())
|
|
83
|
+
CODE_KEYWORDS = set("""
|
|
84
|
+
def class return import from self this true false none null undefined var let const function async
|
|
85
|
+
await public private protected static void int string bool boolean float double char long short new
|
|
86
|
+
else elif while for try catch except finally raise throw throws lambda yield with pass break continue
|
|
87
|
+
package interface extends implements struct enum type impl trait export default module require print
|
|
88
|
+
println args kwargs cls
|
|
89
|
+
""".split())
|
|
90
|
+
|
|
91
|
+
# tiny rule-based stand-in for an LLM's semantic knowledge
|
|
92
|
+
SYNONYMS = {
|
|
93
|
+
"lazy": ["defer", "deferred", "ondemand", "dynamic"],
|
|
94
|
+
"load": ["import", "fetch", "require", "init"],
|
|
95
|
+
"auth": ["authenticate", "login", "token", "session", "credential"],
|
|
96
|
+
"authentication": ["auth", "login", "token", "session"],
|
|
97
|
+
"authorization": ["permission", "role", "policy", "acl"],
|
|
98
|
+
"config": ["configuration", "settings", "options", "env"],
|
|
99
|
+
"db": ["database", "sql", "query", "orm"],
|
|
100
|
+
"database": ["db", "sql", "query", "orm"],
|
|
101
|
+
"error": ["exception", "raise", "throw", "fail"],
|
|
102
|
+
"cache": ["memoize", "ttl", "lru", "store"],
|
|
103
|
+
"log": ["logger", "logging", "trace"],
|
|
104
|
+
"test": ["spec", "mock", "assert", "fixture"],
|
|
105
|
+
"delete": ["remove", "destroy", "drop"],
|
|
106
|
+
"create": ["build", "make", "init", "construct"],
|
|
107
|
+
"parse": ["decode", "deserialize", "tokenize"],
|
|
108
|
+
"render": ["draw", "paint", "display"],
|
|
109
|
+
"route": ["router", "endpoint", "handler", "url"],
|
|
110
|
+
"async": ["await", "promise", "future", "coroutine"],
|
|
111
|
+
"send": ["emit", "dispatch", "publish"],
|
|
112
|
+
"retry": ["backoff", "attempt", "reconnect"],
|
|
113
|
+
"validate": ["verify", "check", "sanitize", "schema"],
|
|
114
|
+
"queue": ["worker", "job", "task"],
|
|
115
|
+
"http": ["request", "response", "client"],
|
|
116
|
+
"thread": ["worker", "pool", "lock", "concurrent"],
|
|
117
|
+
"serialize": ["encode", "dump", "marshal"],
|
|
118
|
+
}
|
|
119
|
+
|
|
120
|
+
TEST_RX = re.compile(
|
|
121
|
+
r"(^|/)(tests?|__tests__|specs?|e2e|fixtures?|mocks?|__mocks__|testdata)(/|$)"
|
|
122
|
+
r"|(^|/)test_[^/]*$|_test\.[a-z0-9]+$|[._-](test|spec)\.[a-z0-9]+$")
|
|
123
|
+
GENERATED_RX = re.compile(r"\.min\.|\.generated\.|_pb2\.py$|\.pb\.go$|(^|/)vendor/")
|
|
124
|
+
|
|
125
|
+
LANG_BY_EXT = {
|
|
126
|
+
".py": "python", ".js": "javascript", ".jsx": "jsx", ".ts": "typescript", ".tsx": "tsx",
|
|
127
|
+
".java": "java", ".go": "go", ".rs": "rust", ".rb": "ruby", ".php": "php", ".c": "c",
|
|
128
|
+
".h": "c", ".cpp": "cpp", ".hpp": "cpp", ".cc": "cpp", ".cs": "csharp", ".kt": "kotlin",
|
|
129
|
+
".swift": "swift", ".scala": "scala", ".sh": "bash", ".sql": "sql", ".md": "markdown",
|
|
130
|
+
".json": "json", ".yml": "yaml", ".yaml": "yaml", ".toml": "toml", ".html": "html",
|
|
131
|
+
".css": "css", ".vue": "vue", ".lua": "lua", ".dart": "dart",
|
|
132
|
+
}
|
|
133
|
+
|
|
134
|
+
IDENT_RX = re.compile(r"[A-Za-z_][A-Za-z0-9_]*")
|
|
135
|
+
HARVEST_RX = re.compile(r"[A-Za-z_][A-Za-z0-9_]{3,}")
|
|
136
|
+
CAMEL_RX = re.compile(r"[A-Z]+(?![a-z])|[A-Z]?[a-z]+|\d+")
|
|
137
|
+
|
|
138
|
+
_MODS = (r"(?:(?:export|default|pub(?:\([\w:]+\))?|public|private|protected|internal|static|async|"
|
|
139
|
+
r"abstract|final|override|unsafe|extern|const|declare)[ \t]+)*")
|
|
140
|
+
SYMBOL_RX = re.compile(
|
|
141
|
+
r"^[ \t]*" + _MODS +
|
|
142
|
+
r"(?P<kind>def|class|function\*?|func|fn|interface|struct|enum|trait|type|object|module)"
|
|
143
|
+
r"[ \t]+(?:\([^)\n]*\)[ \t]*)?(?P<name>[A-Za-z_$][\w$]*)"
|
|
144
|
+
r"|^[ \t]*(?:export[ \t]+)?(?:const|let|var)[ \t]+(?P<cname>[A-Za-z_$][\w$]*)[ \t]*"
|
|
145
|
+
r"(?::[^=\n]+)?=[ \t]*(?:async[ \t]*)?(?:\([^)\n]*\)|\w+)[ \t]*=>",
|
|
146
|
+
re.M)
|
|
147
|
+
CLASS_KINDS = {"class", "interface", "struct", "enum", "trait", "object"}
|
|
148
|
+
|
|
149
|
+
PY_IMPORT_RX = re.compile(r"^\s*(?:from\s+([.\w]+)\s+import|import\s+([\w.]+))")
|
|
150
|
+
JS_IMPORT_RX = re.compile(r"""(?:from|require\(|import\()\s*['"](\.{1,2}/[^'"]+)['"]""")
|
|
151
|
+
|
|
152
|
+
|
|
153
|
+
# --------------------------------------------------------------------------- #
|
|
154
|
+
# Text helpers
|
|
155
|
+
# --------------------------------------------------------------------------- #
|
|
156
|
+
def split_identifier(s: str) -> list[str]:
|
|
157
|
+
s = re.sub(r"([a-z0-9])([A-Z])", r"\1 \2", s)
|
|
158
|
+
s = re.sub(r"([A-Z]+)([A-Z][a-z])", r"\1 \2", s)
|
|
159
|
+
return [p.lower() for p in re.split(r"[^A-Za-z0-9]+", s) if p]
|
|
160
|
+
|
|
161
|
+
|
|
162
|
+
def stem(w: str) -> str:
|
|
163
|
+
w = w.lower()
|
|
164
|
+
for suf in ("ations", "ation", "ings", "ing", "ers", "er", "ies", "ed", "es", "s"):
|
|
165
|
+
if w.endswith(suf) and len(w) - len(suf) >= 4:
|
|
166
|
+
return w[: -len(suf)]
|
|
167
|
+
if w.endswith("e") and len(w) >= 5:
|
|
168
|
+
return w[:-1]
|
|
169
|
+
return w
|
|
170
|
+
|
|
171
|
+
|
|
172
|
+
def looks_like_identifier(tok: str) -> bool:
|
|
173
|
+
return ("_" in tok or re.search(r"[a-z][A-Z]", tok) is not None
|
|
174
|
+
or re.search(r"\w\.\w", tok) is not None or "/" in tok)
|
|
175
|
+
|
|
176
|
+
|
|
177
|
+
def language_of(path: str) -> str:
|
|
178
|
+
return LANG_BY_EXT.get(os.path.splitext(path)[1].lower(), "")
|
|
179
|
+
|
|
180
|
+
|
|
181
|
+
# --------------------------------------------------------------------------- #
|
|
182
|
+
# Indexing (runs in worker processes)
|
|
183
|
+
# --------------------------------------------------------------------------- #
|
|
184
|
+
@lru_cache(maxsize=300_000)
|
|
185
|
+
def _subwords(ident: str) -> tuple[str, ...]:
|
|
186
|
+
return tuple(p.lower() for p in CAMEL_RX.findall(ident) if len(p) >= 2 and not p.isdigit())
|
|
187
|
+
|
|
188
|
+
|
|
189
|
+
def index_text(text: str) -> tuple[dict[str, int], int]:
|
|
190
|
+
"""Token -> term frequency. Whole identifiers AND their camel/snake sub-words are indexed."""
|
|
191
|
+
raw = Counter(IDENT_RX.findall(text)) # C-speed pass; then work on unique identifiers only
|
|
192
|
+
tf: Counter = Counter()
|
|
193
|
+
total = 0
|
|
194
|
+
for ident, n in raw.items():
|
|
195
|
+
total += n
|
|
196
|
+
if len(ident) > 64:
|
|
197
|
+
continue
|
|
198
|
+
low = ident.lower()
|
|
199
|
+
if len(low) >= 2:
|
|
200
|
+
tf[low] += n
|
|
201
|
+
subs = _subwords(ident)
|
|
202
|
+
if len(subs) > 1 or (subs and subs[0] != low):
|
|
203
|
+
for s in subs:
|
|
204
|
+
tf[s] += n
|
|
205
|
+
return dict(tf), total
|
|
206
|
+
|
|
207
|
+
|
|
208
|
+
_NOT_SYMBOL = STOPWORDS | CODE_KEYWORDS | {"in", "is", "of", "on", "at", "by", "as", "or", "to"}
|
|
209
|
+
|
|
210
|
+
|
|
211
|
+
def sym_match(line: str):
|
|
212
|
+
"""(name, kind) if the line declares a function/class-like symbol, else None."""
|
|
213
|
+
m = SYMBOL_RX.match(line)
|
|
214
|
+
if not m:
|
|
215
|
+
return None
|
|
216
|
+
name = m.group("name") or m.group("cname")
|
|
217
|
+
if not name or name.lower() in _NOT_SYMBOL:
|
|
218
|
+
return None
|
|
219
|
+
return name, (m.group("kind") or "function").rstrip("*")
|
|
220
|
+
|
|
221
|
+
|
|
222
|
+
@lru_cache(maxsize=100_000)
|
|
223
|
+
def name_parts(name: str) -> tuple[str, ...]:
|
|
224
|
+
return tuple(split_identifier(name))
|
|
225
|
+
|
|
226
|
+
|
|
227
|
+
def extract_symbols(text: str) -> list[tuple[str, str, int]]:
|
|
228
|
+
out, last_pos, line = [], 0, 1
|
|
229
|
+
for m in SYMBOL_RX.finditer(text):
|
|
230
|
+
line += text.count("\n", last_pos, m.start())
|
|
231
|
+
last_pos = m.start()
|
|
232
|
+
name = m.group("name") or m.group("cname")
|
|
233
|
+
if name and name.lower() not in _NOT_SYMBOL:
|
|
234
|
+
out.append((name, (m.group("kind") or "function").rstrip("*"), line))
|
|
235
|
+
if len(out) >= 3000:
|
|
236
|
+
break
|
|
237
|
+
return out
|
|
238
|
+
|
|
239
|
+
|
|
240
|
+
def _index_worker(args):
|
|
241
|
+
root, rel = args
|
|
242
|
+
full = os.path.join(root, rel)
|
|
243
|
+
try:
|
|
244
|
+
st = os.stat(full)
|
|
245
|
+
if st.st_size == 0 or st.st_size > MAX_FILE_BYTES:
|
|
246
|
+
return rel, st.st_mtime_ns, st.st_size, None
|
|
247
|
+
with open(full, "rb") as fh:
|
|
248
|
+
data = fh.read()
|
|
249
|
+
except OSError:
|
|
250
|
+
return rel, 0, 0, None
|
|
251
|
+
if b"\0" in data[:4096] or len(data) / (data.count(b"\n") + 1) > 400: # binary / minified
|
|
252
|
+
return rel, st.st_mtime_ns, st.st_size, None
|
|
253
|
+
text = data.decode("utf-8", "replace")
|
|
254
|
+
tf, ntok = index_text(text)
|
|
255
|
+
return rel, st.st_mtime_ns, st.st_size, (tf, ntok, extract_symbols(text))
|
|
256
|
+
|
|
257
|
+
|
|
258
|
+
def list_files(root: str) -> list[str]:
|
|
259
|
+
files = None
|
|
260
|
+
try:
|
|
261
|
+
out = subprocess.run(["git", "-C", root, "ls-files", "-z", "-co", "--exclude-standard"],
|
|
262
|
+
capture_output=True, timeout=60)
|
|
263
|
+
if out.returncode == 0 and out.stdout:
|
|
264
|
+
files = out.stdout.decode("utf-8", "replace").split("\0")
|
|
265
|
+
except (OSError, subprocess.SubprocessError):
|
|
266
|
+
pass
|
|
267
|
+
if files is None:
|
|
268
|
+
files, stack = [], [root]
|
|
269
|
+
while stack:
|
|
270
|
+
d = stack.pop()
|
|
271
|
+
try:
|
|
272
|
+
with os.scandir(d) as it:
|
|
273
|
+
for e in it:
|
|
274
|
+
if e.is_dir(follow_symlinks=False):
|
|
275
|
+
if e.name not in IGNORE_DIRS and not e.name.startswith("."):
|
|
276
|
+
stack.append(e.path)
|
|
277
|
+
elif e.is_file(follow_symlinks=False):
|
|
278
|
+
files.append(os.path.relpath(e.path, root))
|
|
279
|
+
except OSError:
|
|
280
|
+
continue
|
|
281
|
+
keep = []
|
|
282
|
+
for f in files:
|
|
283
|
+
if not f:
|
|
284
|
+
continue
|
|
285
|
+
f = f.replace("\\", "/")
|
|
286
|
+
parts = f.split("/")
|
|
287
|
+
name = parts[-1]
|
|
288
|
+
if any(p in IGNORE_DIRS or p.endswith((".dist-info", ".egg-info")) for p in parts[:-1]) \
|
|
289
|
+
or name in SKIP_NAMES:
|
|
290
|
+
continue
|
|
291
|
+
if os.path.splitext(name)[1].lower() in BINARY_EXT or name.endswith(".min.js"):
|
|
292
|
+
continue
|
|
293
|
+
keep.append(f)
|
|
294
|
+
return keep
|
|
295
|
+
|
|
296
|
+
|
|
297
|
+
def cache_path(root: str) -> str:
|
|
298
|
+
base = os.environ.get("FC_CACHE_DIR") or os.path.join(
|
|
299
|
+
os.environ.get("XDG_CACHE_HOME") or os.path.join(os.path.expanduser("~"), ".cache"), "fastcontext")
|
|
300
|
+
try:
|
|
301
|
+
os.makedirs(base, exist_ok=True)
|
|
302
|
+
except OSError:
|
|
303
|
+
base = os.path.join(tempfile.gettempdir(), "fastcontext-cache")
|
|
304
|
+
os.makedirs(base, exist_ok=True)
|
|
305
|
+
return os.path.join(base, hashlib.sha1(root.encode()).hexdigest()[:16] + ".sqlite")
|
|
306
|
+
|
|
307
|
+
|
|
308
|
+
class Blob:
|
|
309
|
+
"""Substring search across many short strings using ONE big str (C-speed str.find)."""
|
|
310
|
+
__slots__ = ("blob", "starts")
|
|
311
|
+
|
|
312
|
+
def __init__(self, items: list[str]):
|
|
313
|
+
self.blob = "\n" + "\n".join(items) + "\n"
|
|
314
|
+
self.starts = list(accumulate((len(s) + 1 for s in items), initial=1))
|
|
315
|
+
|
|
316
|
+
def find(self, sub: str, limit: int = 20000) -> list[int]:
|
|
317
|
+
blob, starts, out = self.blob, self.starts, []
|
|
318
|
+
pos = blob.find(sub)
|
|
319
|
+
while pos != -1:
|
|
320
|
+
i = bisect.bisect_right(starts, pos) - 1
|
|
321
|
+
out.append(i)
|
|
322
|
+
if len(out) >= limit:
|
|
323
|
+
break
|
|
324
|
+
pos = blob.find(sub, starts[i + 1])
|
|
325
|
+
return out
|
|
326
|
+
|
|
327
|
+
|
|
328
|
+
class Index:
|
|
329
|
+
"""Persistent, incrementally-updated inverted index (SQLite)."""
|
|
330
|
+
|
|
331
|
+
def __init__(self, root: str, use_cache: bool = True, reindex: bool = False, log=lambda *a: None):
|
|
332
|
+
self.root, self.log = root, log
|
|
333
|
+
self.path = cache_path(root) if use_cache else ":memory:"
|
|
334
|
+
if reindex and use_cache:
|
|
335
|
+
for suffix in ("", "-wal", "-shm"):
|
|
336
|
+
try:
|
|
337
|
+
os.remove(self.path + suffix)
|
|
338
|
+
except OSError:
|
|
339
|
+
pass
|
|
340
|
+
self.db = sqlite3.connect(self.path)
|
|
341
|
+
self.db.execute("PRAGMA synchronous=OFF")
|
|
342
|
+
self.db.execute("PRAGMA cache_size=-131072")
|
|
343
|
+
if use_cache:
|
|
344
|
+
self.db.execute("PRAGMA journal_mode=WAL")
|
|
345
|
+
self._schema()
|
|
346
|
+
|
|
347
|
+
def _schema(self):
|
|
348
|
+
db = self.db
|
|
349
|
+
db.execute("CREATE TABLE IF NOT EXISTS meta(k TEXT PRIMARY KEY, v TEXT)")
|
|
350
|
+
row = db.execute("SELECT v FROM meta WHERE k='version'").fetchone()
|
|
351
|
+
if row and row[0] != str(VERSION):
|
|
352
|
+
db.executescript("DROP TABLE IF EXISTS files; DROP TABLE IF EXISTS vocab; "
|
|
353
|
+
"DROP TABLE IF EXISTS postings; DROP TABLE IF EXISTS symbols; DELETE FROM meta;")
|
|
354
|
+
row = None
|
|
355
|
+
db.executescript("""
|
|
356
|
+
CREATE TABLE IF NOT EXISTS files(id INTEGER PRIMARY KEY, path TEXT UNIQUE, mtime INTEGER, size INTEGER, ntok INTEGER);
|
|
357
|
+
CREATE TABLE IF NOT EXISTS postings(tid INTEGER, fid INTEGER, tf INTEGER, PRIMARY KEY(tid, fid)) WITHOUT ROWID;
|
|
358
|
+
CREATE INDEX IF NOT EXISTS postings_fid ON postings(fid);
|
|
359
|
+
CREATE TABLE IF NOT EXISTS symbols(fid INTEGER, name TEXT, lname TEXT, kind TEXT, line INTEGER);
|
|
360
|
+
CREATE INDEX IF NOT EXISTS symbols_fid ON symbols(fid);
|
|
361
|
+
CREATE INDEX IF NOT EXISTS symbols_lname ON symbols(lname);
|
|
362
|
+
""")
|
|
363
|
+
if not row:
|
|
364
|
+
db.execute("INSERT OR REPLACE INTO meta VALUES('version', ?)", (str(VERSION),))
|
|
365
|
+
db.commit()
|
|
366
|
+
|
|
367
|
+
def _load_vocab(self) -> list[str]:
|
|
368
|
+
"""Vocabulary is one newline-joined string in meta; token i has id i+1 (append-only)."""
|
|
369
|
+
row = self.db.execute("SELECT v FROM meta WHERE k='vocab'").fetchone()
|
|
370
|
+
return row[0].split("\n") if row and row[0] else []
|
|
371
|
+
|
|
372
|
+
# ---- incremental sync ------------------------------------------------- #
|
|
373
|
+
def sync(self) -> dict:
|
|
374
|
+
db = self.db
|
|
375
|
+
now = {}
|
|
376
|
+
for rel in list_files(self.root):
|
|
377
|
+
try:
|
|
378
|
+
st = os.stat(os.path.join(self.root, rel))
|
|
379
|
+
except OSError:
|
|
380
|
+
continue
|
|
381
|
+
now[rel] = (st.st_mtime_ns, st.st_size)
|
|
382
|
+
known = {p: (i, m, s) for i, p, m, s in db.execute("SELECT id, path, mtime, size FROM files")}
|
|
383
|
+
changed = [p for p, ms in now.items() if p not in known or known[p][1:] != ms]
|
|
384
|
+
removed = [known[p][0] for p in known if p not in now]
|
|
385
|
+
stats = {"files": len(now), "changed": len(changed), "removed": len(removed)}
|
|
386
|
+
|
|
387
|
+
for i in range(0, len(removed), 500):
|
|
388
|
+
chunk = removed[i:i + 500]
|
|
389
|
+
q = ",".join("?" * len(chunk))
|
|
390
|
+
for table, col in (("postings", "fid"), ("symbols", "fid"), ("files", "id")):
|
|
391
|
+
db.execute(f"DELETE FROM {table} WHERE {col} IN ({q})", chunk)
|
|
392
|
+
|
|
393
|
+
if changed:
|
|
394
|
+
if len(changed) >= 500:
|
|
395
|
+
print(f"[FastContext] indexing {len(changed):,} files (cached for next time)...", file=sys.stderr)
|
|
396
|
+
args = [(self.root, p) for p in changed]
|
|
397
|
+
if len(changed) > 300 and (os.cpu_count() or 1) > 1:
|
|
398
|
+
with ProcessPoolExecutor(max_workers=min(os.cpu_count() or 2, 8)) as ex:
|
|
399
|
+
results = list(ex.map(_index_worker, args, chunksize=64))
|
|
400
|
+
else:
|
|
401
|
+
results = [_index_worker(a) for a in args]
|
|
402
|
+
|
|
403
|
+
bulk = len(changed) >= 1500
|
|
404
|
+
if bulk: # cold build: sorted bulk insert is far faster than random B-tree inserts
|
|
405
|
+
old = [known[p][0] for p in changed if p in known]
|
|
406
|
+
for i in range(0, len(old), 500):
|
|
407
|
+
chunk = old[i:i + 500]
|
|
408
|
+
q = ",".join("?" * len(chunk))
|
|
409
|
+
db.execute(f"DELETE FROM postings WHERE fid IN ({q})", chunk)
|
|
410
|
+
db.execute(f"DELETE FROM symbols WHERE fid IN ({q})", chunk)
|
|
411
|
+
db.execute("DROP INDEX IF EXISTS postings_fid")
|
|
412
|
+
db.execute("CREATE TEMP TABLE post_stage(tid INTEGER, fid INTEGER, tf INTEGER)")
|
|
413
|
+
target = "post_stage" if bulk else "postings"
|
|
414
|
+
vtoks = self._load_vocab()
|
|
415
|
+
tid_map = {tok: i + 1 for i, tok in enumerate(vtoks)}
|
|
416
|
+
next_tid = len(vtoks) + 1
|
|
417
|
+
vocab_grew = False
|
|
418
|
+
for rel, mtime, size, payload in results:
|
|
419
|
+
if rel in known:
|
|
420
|
+
fid = known[rel][0]
|
|
421
|
+
if not bulk:
|
|
422
|
+
db.execute("DELETE FROM postings WHERE fid=?", (fid,))
|
|
423
|
+
db.execute("DELETE FROM symbols WHERE fid=?", (fid,))
|
|
424
|
+
db.execute("UPDATE files SET mtime=?, size=?, ntok=? WHERE id=?",
|
|
425
|
+
(mtime, size, payload[1] if payload else 0, fid))
|
|
426
|
+
else:
|
|
427
|
+
fid = db.execute("INSERT INTO files(path, mtime, size, ntok) VALUES (?,?,?,?)",
|
|
428
|
+
(rel, mtime, size, payload[1] if payload else 0)).lastrowid
|
|
429
|
+
if not payload:
|
|
430
|
+
continue
|
|
431
|
+
tf, _, syms = payload
|
|
432
|
+
rows = []
|
|
433
|
+
for tok, n in tf.items():
|
|
434
|
+
tid = tid_map.get(tok)
|
|
435
|
+
if tid is None:
|
|
436
|
+
tid = tid_map[tok] = next_tid
|
|
437
|
+
next_tid += 1
|
|
438
|
+
vtoks.append(tok)
|
|
439
|
+
vocab_grew = True
|
|
440
|
+
rows.append((tid, fid, n))
|
|
441
|
+
db.executemany(f"INSERT INTO {target} VALUES (?,?,?)", rows)
|
|
442
|
+
if syms:
|
|
443
|
+
db.executemany("INSERT INTO symbols VALUES (?,?,?,?,?)",
|
|
444
|
+
[(fid, n, n.lower(), k, ln) for n, k, ln in syms])
|
|
445
|
+
if vocab_grew:
|
|
446
|
+
db.execute("INSERT OR REPLACE INTO meta VALUES('vocab', ?)", ("\n".join(vtoks),))
|
|
447
|
+
if bulk:
|
|
448
|
+
db.execute("INSERT INTO postings SELECT tid, fid, tf FROM post_stage ORDER BY tid, fid")
|
|
449
|
+
db.execute("DROP TABLE post_stage")
|
|
450
|
+
db.execute("CREATE INDEX IF NOT EXISTS postings_fid ON postings(fid)")
|
|
451
|
+
db.commit()
|
|
452
|
+
return stats
|
|
453
|
+
|
|
454
|
+
# ---- in-memory views used at query time -------------------------------- #
|
|
455
|
+
def load(self):
|
|
456
|
+
rows = self.db.execute("SELECT id, path, ntok FROM files").fetchall()
|
|
457
|
+
self.paths = {i: p for i, p, _ in rows}
|
|
458
|
+
self.ntok = {i: n for i, _, n in rows}
|
|
459
|
+
indexed = [n for _, _, n in rows if n > 0]
|
|
460
|
+
self.N = max(1, len(indexed))
|
|
461
|
+
self.avgdl = (sum(indexed) / len(indexed)) if indexed else 1.0
|
|
462
|
+
self.fids = [i for i, _, _ in rows]
|
|
463
|
+
self.path_blob = Blob([p.lower() for _, p, _ in rows])
|
|
464
|
+
self.vtoks = self._load_vocab()
|
|
465
|
+
self.vocab = Blob(self.vtoks)
|
|
466
|
+
|
|
467
|
+
def postings(self, tids: list[int]):
|
|
468
|
+
for i in range(0, len(tids), 500):
|
|
469
|
+
chunk = tids[i:i + 500]
|
|
470
|
+
q = ",".join("?" * len(chunk))
|
|
471
|
+
yield from self.db.execute(f"SELECT fid, tid, tf FROM postings WHERE tid IN ({q})", chunk)
|
|
472
|
+
|
|
473
|
+
def symbols_by_lname(self, names: list[str]):
|
|
474
|
+
for i in range(0, len(names), 500):
|
|
475
|
+
chunk = names[i:i + 500]
|
|
476
|
+
q = ",".join("?" * len(chunk))
|
|
477
|
+
yield from self.db.execute(
|
|
478
|
+
f"SELECT fid, name, kind, line FROM symbols WHERE lname IN ({q})", chunk)
|
|
479
|
+
|
|
480
|
+
def stats(self) -> dict:
|
|
481
|
+
c = self.db.execute
|
|
482
|
+
return {
|
|
483
|
+
"db": self.path,
|
|
484
|
+
"files": c("SELECT COUNT(*) FROM files").fetchone()[0],
|
|
485
|
+
"vocab": len(self._load_vocab()),
|
|
486
|
+
"postings": c("SELECT COUNT(*) FROM postings").fetchone()[0],
|
|
487
|
+
"symbols": c("SELECT COUNT(*) FROM symbols").fetchone()[0],
|
|
488
|
+
"db_mb": round(os.path.getsize(self.path) / 1e6, 1) if os.path.exists(self.path) else 0,
|
|
489
|
+
}
|
|
490
|
+
|
|
491
|
+
|
|
492
|
+
# --------------------------------------------------------------------------- #
|
|
493
|
+
# Query terms
|
|
494
|
+
# --------------------------------------------------------------------------- #
|
|
495
|
+
@dataclass
|
|
496
|
+
class Term:
|
|
497
|
+
key: str
|
|
498
|
+
raw: str
|
|
499
|
+
weight: float
|
|
500
|
+
kind: str # word | ident | ngram | syn | expand
|
|
501
|
+
index_variants: list[str]
|
|
502
|
+
text_variants: list[str]
|
|
503
|
+
covers: list[str] = field(default_factory=list) # ngram: keys of the word terms it spans
|
|
504
|
+
tf: dict = field(default_factory=dict) # fid -> quality-weighted term frequency
|
|
505
|
+
syms: list = field(default_factory=list) # (fid, name, kind, line)
|
|
506
|
+
path_fids: set = field(default_factory=set)
|
|
507
|
+
df: int = 0
|
|
508
|
+
idf: float = 0.0
|
|
509
|
+
looked_up: bool = False
|
|
510
|
+
short: bool = False
|
|
511
|
+
rx: re.Pattern | None = None
|
|
512
|
+
|
|
513
|
+
def __post_init__(self):
|
|
514
|
+
self.short = len(self.key) <= 3 and self.kind in ("word", "syn", "expand")
|
|
515
|
+
if self.short:
|
|
516
|
+
self.rx = re.compile(r"(?<![a-z0-9])" + re.escape(self.key) + r"(?![a-z0-9])")
|
|
517
|
+
|
|
518
|
+
@property
|
|
519
|
+
def in_query(self) -> bool:
|
|
520
|
+
return self.kind in ("word", "ident")
|
|
521
|
+
|
|
522
|
+
|
|
523
|
+
def make_word(raw: str, weight: float | None = None, kind: str = "word") -> Term:
|
|
524
|
+
k = stem(raw) if kind == "word" else raw
|
|
525
|
+
w = weight if weight is not None else min(1.5, 0.8 + len(k) / 12)
|
|
526
|
+
return Term(k, raw, w, kind, [k], [k])
|
|
527
|
+
|
|
528
|
+
|
|
529
|
+
def make_ngram(words: list[str], weight: float) -> Term:
|
|
530
|
+
stems = [stem(w) for w in words]
|
|
531
|
+
variants = set()
|
|
532
|
+
for seq in (words, stems):
|
|
533
|
+
variants |= {"".join(seq), "_".join(seq), "-".join(seq), " ".join(seq)}
|
|
534
|
+
idx = [v for v in variants if re.fullmatch(r"[a-z0-9_]+", v)]
|
|
535
|
+
return Term("~" + "_".join(words), " ".join(words), weight, "ngram", idx, sorted(variants),
|
|
536
|
+
covers=stems)
|
|
537
|
+
|
|
538
|
+
|
|
539
|
+
def make_ident(text: str, weight: float = 3.0) -> Term:
|
|
540
|
+
low = text.lower()
|
|
541
|
+
if re.fullmatch(r"[a-z0-9_$]+", low):
|
|
542
|
+
idx = [low]
|
|
543
|
+
else:
|
|
544
|
+
parts = sorted((p for p in re.split(r"[^a-z0-9_$]+", low) if len(p) >= 3), key=len, reverse=True)
|
|
545
|
+
idx = parts[:1]
|
|
546
|
+
return Term(low, text, weight, "ident", idx, [low])
|
|
547
|
+
|
|
548
|
+
|
|
549
|
+
# --------------------------------------------------------------------------- #
|
|
550
|
+
# File data + line-level scanning
|
|
551
|
+
# --------------------------------------------------------------------------- #
|
|
552
|
+
class FileData:
|
|
553
|
+
__slots__ = ("lines", "low", "offs")
|
|
554
|
+
|
|
555
|
+
def __init__(self, text: str):
|
|
556
|
+
self.lines = text.split("\n")
|
|
557
|
+
low = text.lower()
|
|
558
|
+
if len(low) != len(text): # rare unicode case-mapping length changes: keep offsets valid
|
|
559
|
+
low = "".join(c.lower() if len(c.lower()) == 1 else c for c in text)
|
|
560
|
+
self.low = low
|
|
561
|
+
self.offs = list(accumulate(len(l) + 1 for l in self.lines)) # offs[i] = start of line i+1
|
|
562
|
+
|
|
563
|
+
|
|
564
|
+
class FileAnalysis:
|
|
565
|
+
__slots__ = ("fd", "term_lines", "line_w")
|
|
566
|
+
|
|
567
|
+
def __init__(self):
|
|
568
|
+
self.fd: FileData | None = None
|
|
569
|
+
self.term_lines: dict[str, list[int]] = {}
|
|
570
|
+
self.line_w: dict[int, float] = {}
|
|
571
|
+
|
|
572
|
+
|
|
573
|
+
def find_lines(term: Term, fd: FileData, cap: int = 1500) -> list[int]:
|
|
574
|
+
low, offs, out = fd.low, fd.offs, set()
|
|
575
|
+
if term.rx is not None:
|
|
576
|
+
for m in term.rx.finditer(low):
|
|
577
|
+
out.add(bisect.bisect_right(offs, m.start()))
|
|
578
|
+
if len(out) >= cap:
|
|
579
|
+
break
|
|
580
|
+
else:
|
|
581
|
+
for v in term.text_variants:
|
|
582
|
+
pos = low.find(v)
|
|
583
|
+
while pos != -1 and len(out) < cap:
|
|
584
|
+
ln = bisect.bisect_right(offs, pos)
|
|
585
|
+
out.add(ln)
|
|
586
|
+
if ln >= len(offs):
|
|
587
|
+
break
|
|
588
|
+
pos = low.find(v, offs[ln])
|
|
589
|
+
return sorted(out)
|
|
590
|
+
|
|
591
|
+
|
|
592
|
+
def best_window(hits: dict[int, set], wmap: dict[str, float], span: int = 8) -> tuple[float, int]:
|
|
593
|
+
lines = sorted(hits)
|
|
594
|
+
if len(lines) > 400:
|
|
595
|
+
lines = sorted(sorted(lines, key=lambda l: -len(hits[l]))[:400])
|
|
596
|
+
best, best_line = 0.0, -1
|
|
597
|
+
for i, li in enumerate(lines):
|
|
598
|
+
acc, k = set(), i
|
|
599
|
+
while k < len(lines) and lines[k] - li <= span:
|
|
600
|
+
acc |= hits[lines[k]]
|
|
601
|
+
k += 1
|
|
602
|
+
w = sum(wmap.get(t, 0.0) for t in acc)
|
|
603
|
+
if w > best:
|
|
604
|
+
best, best_line = w, li
|
|
605
|
+
return best, best_line
|
|
606
|
+
|
|
607
|
+
|
|
608
|
+
def enclosing_symbol(lines: list[str], idx: int, max_back: int = 300) -> str:
|
|
609
|
+
lo = max(0, idx - max_back)
|
|
610
|
+
for i in range(idx, lo - 1, -1):
|
|
611
|
+
m = sym_match(lines[i])
|
|
612
|
+
if m:
|
|
613
|
+
name = m[0]
|
|
614
|
+
ind = len(lines[i]) - len(lines[i].lstrip(" \t"))
|
|
615
|
+
if ind > 0:
|
|
616
|
+
for j in range(i - 1, lo - 1, -1):
|
|
617
|
+
mj = sym_match(lines[j])
|
|
618
|
+
if mj and mj[1] in CLASS_KINDS and \
|
|
619
|
+
(len(lines[j]) - len(lines[j].lstrip(" \t"))) < ind:
|
|
620
|
+
return f"{mj[0]}.{name}"
|
|
621
|
+
return name
|
|
622
|
+
return ""
|
|
623
|
+
|
|
624
|
+
|
|
625
|
+
# --------------------------------------------------------------------------- #
|
|
626
|
+
# Results
|
|
627
|
+
# --------------------------------------------------------------------------- #
|
|
628
|
+
@dataclass
|
|
629
|
+
class Snippet:
|
|
630
|
+
start: int
|
|
631
|
+
end: int
|
|
632
|
+
hit_lines: list[int]
|
|
633
|
+
text: list[str]
|
|
634
|
+
symbol: str = ""
|
|
635
|
+
terms: list[str] = field(default_factory=list)
|
|
636
|
+
score: float = 0.0
|
|
637
|
+
|
|
638
|
+
|
|
639
|
+
@dataclass
|
|
640
|
+
class Scored:
|
|
641
|
+
fid: int
|
|
642
|
+
path: str
|
|
643
|
+
base: float
|
|
644
|
+
final: float = 0.0
|
|
645
|
+
cov: float = 0.0
|
|
646
|
+
prox: float = 0.0
|
|
647
|
+
prox_cov: float = 0.0
|
|
648
|
+
prox_line: int = -1
|
|
649
|
+
defsig: float = 0.0
|
|
650
|
+
defsym: tuple | None = None
|
|
651
|
+
pathsig: float = 0.0
|
|
652
|
+
spec: float = 0.0
|
|
653
|
+
strength: float = 0.0
|
|
654
|
+
confidence: float = 0.0
|
|
655
|
+
present: list[str] = field(default_factory=list)
|
|
656
|
+
snippets: list[Snippet] = field(default_factory=list)
|
|
657
|
+
|
|
658
|
+
|
|
659
|
+
def label_of(c: float) -> str:
|
|
660
|
+
return "high" if c >= 75 else "medium" if c >= 50 else "low" if c >= 30 else "very_low"
|
|
661
|
+
|
|
662
|
+
|
|
663
|
+
def file_confidence(s: Scored) -> float:
|
|
664
|
+
raw = 0.32 * s.cov + 0.26 * s.prox + 0.18 * s.defsig + 0.08 * s.pathsig + 0.16 * s.strength
|
|
665
|
+
return round(100 * min(1.0, raw * (0.8 + 0.2 * s.spec)), 1)
|
|
666
|
+
|
|
667
|
+
|
|
668
|
+
# --------------------------------------------------------------------------- #
|
|
669
|
+
# The retriever
|
|
670
|
+
# --------------------------------------------------------------------------- #
|
|
671
|
+
class Retriever:
|
|
672
|
+
def __init__(self, root, *, use_cache=True, reindex=False, max_turns=MAX_TURNS, context=3,
|
|
673
|
+
per_file=3, include_tests=False, exts=None, globs=None, thorough=False, verbose=False):
|
|
674
|
+
self.root = os.path.abspath(root)
|
|
675
|
+
self.verbose = verbose
|
|
676
|
+
self.ctx, self.per_file, self.thorough = context, per_file, thorough
|
|
677
|
+
self.max_turns = min(max_turns, MAX_TURNS)
|
|
678
|
+
self.include_tests = include_tests
|
|
679
|
+
self.exts = {e.strip().lower().lstrip(".") for e in (exts or []) if e.strip()}
|
|
680
|
+
self.globs = [g.strip() for g in (globs or []) if g.strip()]
|
|
681
|
+
self.workers = min(16, (os.cpu_count() or 4) * 2)
|
|
682
|
+
self.idx = Index(self.root, use_cache, reindex)
|
|
683
|
+
self.terms: dict[str, Term] = {}
|
|
684
|
+
self.analysis: dict[int, FileAnalysis] = {}
|
|
685
|
+
self.import_boost: dict[str, float] = {}
|
|
686
|
+
self._allowed: dict[int, bool] = {}
|
|
687
|
+
self._suffix_map: dict[str, list[str]] | None = None
|
|
688
|
+
self.file_syms: dict = {}
|
|
689
|
+
self.base: dict[int, float] = {}
|
|
690
|
+
self.timings: dict[str, float] = defaultdict(float)
|
|
691
|
+
self.turns_used = 0
|
|
692
|
+
self.converged = True
|
|
693
|
+
|
|
694
|
+
def log(self, *a):
|
|
695
|
+
if self.verbose:
|
|
696
|
+
print(*a, file=sys.stderr)
|
|
697
|
+
|
|
698
|
+
# ---- query parsing --------------------------------------------------- #
|
|
699
|
+
def build_terms(self, query: str) -> None:
|
|
700
|
+
T = self.terms
|
|
701
|
+
|
|
702
|
+
def add(t: Term):
|
|
703
|
+
if t.key not in T:
|
|
704
|
+
T[t.key] = t
|
|
705
|
+
|
|
706
|
+
for a, b in re.findall(r'"([^"]+)"|`([^`]+)`', query):
|
|
707
|
+
phrase = (a or b).strip().lower()
|
|
708
|
+
words = [w for w in re.split(r"[^a-z0-9]+", phrase) if len(w) >= 2]
|
|
709
|
+
if len(words) >= 2:
|
|
710
|
+
add(make_ngram(words[:3], 3.0))
|
|
711
|
+
for w in words:
|
|
712
|
+
if w not in STOPWORDS:
|
|
713
|
+
add(make_word(w))
|
|
714
|
+
q = re.sub(r'"[^"]+"|`[^`]+`', " ", query)
|
|
715
|
+
|
|
716
|
+
run: list[str] = []
|
|
717
|
+
ngrams: list[Term] = []
|
|
718
|
+
|
|
719
|
+
def flush():
|
|
720
|
+
nonlocal run
|
|
721
|
+
if len(run) >= 2:
|
|
722
|
+
for i in range(len(run) - 1):
|
|
723
|
+
ngrams.append(make_ngram(run[i:i + 2], 2.0))
|
|
724
|
+
if len(run) >= 3:
|
|
725
|
+
ngrams.append(make_ngram(run[:3], 2.5))
|
|
726
|
+
run = []
|
|
727
|
+
|
|
728
|
+
for tok in q.split():
|
|
729
|
+
tok = tok.strip(",;:!?[]{}<>\"'").rstrip(".")
|
|
730
|
+
cleaned = tok[:-2] if tok.endswith("()") else tok
|
|
731
|
+
if not cleaned:
|
|
732
|
+
continue
|
|
733
|
+
if len(cleaned) >= 3 and looks_like_identifier(cleaned):
|
|
734
|
+
flush()
|
|
735
|
+
add(make_ident(cleaned))
|
|
736
|
+
for part in split_identifier(cleaned):
|
|
737
|
+
if len(part) >= 3 and part not in STOPWORDS:
|
|
738
|
+
add(make_word(part, 0.7))
|
|
739
|
+
continue
|
|
740
|
+
for w in re.split(r"[^A-Za-z0-9]+", cleaned):
|
|
741
|
+
wl = w.lower()
|
|
742
|
+
if len(wl) < 2 or wl.isdigit() or wl in STOPWORDS or stem(wl) in STOPWORDS:
|
|
743
|
+
flush()
|
|
744
|
+
continue
|
|
745
|
+
add(make_word(wl))
|
|
746
|
+
run.append(wl)
|
|
747
|
+
flush()
|
|
748
|
+
for ng in ngrams[:5]:
|
|
749
|
+
add(ng)
|
|
750
|
+
|
|
751
|
+
@property
|
|
752
|
+
def qterms(self) -> list[Term]:
|
|
753
|
+
return [t for t in self.terms.values() if t.in_query and t.weight > 0]
|
|
754
|
+
|
|
755
|
+
# ---- index lookups (the "grep/glob" step) ---------------------------- #
|
|
756
|
+
def lookup(self, t: Term) -> None:
|
|
757
|
+
idx = self.idx
|
|
758
|
+
quality: dict[int, float] = {}
|
|
759
|
+
for v in t.index_variants:
|
|
760
|
+
for i in idx.vocab.find(v):
|
|
761
|
+
tok = idx.vtoks[i]
|
|
762
|
+
if t.short and tok != v:
|
|
763
|
+
continue
|
|
764
|
+
q = 1.0 if tok == v else 0.8 if tok.startswith(v) else 0.65 if tok.endswith(v) else 0.5
|
|
765
|
+
if q > quality.get(i, 0.0):
|
|
766
|
+
quality[i] = q
|
|
767
|
+
items = sorted(quality.items(), key=lambda kv: (-kv[1], len(idx.vtoks[kv[0]])))[:MAX_TOKENS_PER_TERM]
|
|
768
|
+
tid_q = {i + 1: q for i, q in items}
|
|
769
|
+
tf: dict[int, float] = defaultdict(float)
|
|
770
|
+
for fid, tid, n in idx.postings(list(tid_q)):
|
|
771
|
+
tf[fid] += tid_q[tid] * n
|
|
772
|
+
t.tf, t.df = tf, len(tf)
|
|
773
|
+
t.idf = math.log(1 + (idx.N - t.df + 0.5) / (t.df + 0.5))
|
|
774
|
+
t.syms = list(idx.symbols_by_lname([idx.vtoks[i] for i, _ in items]))
|
|
775
|
+
if not t.short:
|
|
776
|
+
fids = set()
|
|
777
|
+
for v in t.text_variants:
|
|
778
|
+
fids.update(idx.fids[i] for i in idx.path_blob.find(v, limit=5000))
|
|
779
|
+
t.path_fids = fids
|
|
780
|
+
if t.kind in ("syn", "expand") and t.df > 0.15 * idx.N:
|
|
781
|
+
t.weight = 0.0 # too common to be informative
|
|
782
|
+
t.looked_up = True
|
|
783
|
+
|
|
784
|
+
# ---- filters / priors ------------------------------------------------ #
|
|
785
|
+
def allowed(self, fid: int) -> bool:
|
|
786
|
+
ok = self._allowed.get(fid)
|
|
787
|
+
if ok is None:
|
|
788
|
+
p = self.idx.paths[fid]
|
|
789
|
+
ok = True
|
|
790
|
+
if self.exts and os.path.splitext(p)[1].lower().lstrip(".") not in self.exts:
|
|
791
|
+
ok = False
|
|
792
|
+
if ok and self.globs and not any(fnmatch.fnmatch(p, g) for g in self.globs):
|
|
793
|
+
ok = False
|
|
794
|
+
self._allowed[fid] = ok
|
|
795
|
+
return ok
|
|
796
|
+
|
|
797
|
+
def prior(self, pl: str) -> float:
|
|
798
|
+
p = 1.0
|
|
799
|
+
if not self.include_tests and TEST_RX.search(pl):
|
|
800
|
+
p *= 0.6
|
|
801
|
+
if pl.endswith((".md", ".rst", ".txt")):
|
|
802
|
+
p *= 0.8
|
|
803
|
+
if GENERATED_RX.search(pl):
|
|
804
|
+
p *= 0.5
|
|
805
|
+
return p
|
|
806
|
+
|
|
807
|
+
# ---- reading + scanning (the "read" step, threaded) ------------------- #
|
|
808
|
+
def read_fd(self, rel: str) -> FileData | None:
|
|
809
|
+
try:
|
|
810
|
+
with open(os.path.join(self.root, rel), "rb") as fh:
|
|
811
|
+
data = fh.read(MAX_FILE_BYTES + 1)
|
|
812
|
+
except OSError:
|
|
813
|
+
return None
|
|
814
|
+
if len(data) > MAX_FILE_BYTES or b"\0" in data[:4096]:
|
|
815
|
+
return None
|
|
816
|
+
return FileData(data.decode("utf-8", "replace"))
|
|
817
|
+
|
|
818
|
+
def prepare(self, fid: int) -> None:
|
|
819
|
+
fa = self.analysis.get(fid)
|
|
820
|
+
if fa is None:
|
|
821
|
+
fa = self.analysis[fid] = FileAnalysis()
|
|
822
|
+
fa.fd = self.read_fd(self.idx.paths[fid])
|
|
823
|
+
if fa.fd is None:
|
|
824
|
+
return
|
|
825
|
+
for t in self.terms.values():
|
|
826
|
+
if t.looked_up and t.key not in fa.term_lines:
|
|
827
|
+
fa.term_lines[t.key] = find_lines(t, fa.fd)
|
|
828
|
+
|
|
829
|
+
# ---- ranking ---------------------------------------------------------- #
|
|
830
|
+
def rank(self, top_k: int) -> list[Scored]:
|
|
831
|
+
idx = self.idx
|
|
832
|
+
t0 = time.perf_counter()
|
|
833
|
+
base: dict[int, float] = defaultdict(float)
|
|
834
|
+
for t in self.terms.values():
|
|
835
|
+
if not t.looked_up or t.weight <= 0:
|
|
836
|
+
continue
|
|
837
|
+
wi = t.weight * t.idf
|
|
838
|
+
for fid, tf in t.tf.items():
|
|
839
|
+
base[fid] += wi * tf * (K1 + 1) / (tf + K1 * (1 - B + B * idx.ntok[fid] / idx.avgdl))
|
|
840
|
+
for fid in t.path_fids:
|
|
841
|
+
base[fid] += wi * 1.5
|
|
842
|
+
seen: Counter = Counter()
|
|
843
|
+
for fid, _, _, _ in t.syms:
|
|
844
|
+
if seen[fid] < 3:
|
|
845
|
+
seen[fid] += 1
|
|
846
|
+
base[fid] += wi * 1.2
|
|
847
|
+
pool = [(b * self.prior(idx.paths[f].lower()), f) for f, b in base.items() if self.allowed(f)]
|
|
848
|
+
pool = heapq.nlargest(max(40, top_k * 5), pool)
|
|
849
|
+
self.base = {f: b for b, f in pool}
|
|
850
|
+
self.timings["rank_index"] += time.perf_counter() - t0
|
|
851
|
+
|
|
852
|
+
t1 = time.perf_counter()
|
|
853
|
+
with ThreadPoolExecutor(max_workers=self.workers) as ex:
|
|
854
|
+
list(ex.map(self.prepare, [f for _, f in pool]))
|
|
855
|
+
self.timings["read_scan"] += time.perf_counter() - t1
|
|
856
|
+
|
|
857
|
+
self.file_syms = defaultdict(lambda: defaultdict(set))
|
|
858
|
+
for t in self.terms.values():
|
|
859
|
+
keys = {t.key} if t.in_query else set(t.covers) if t.kind == "ngram" else set()
|
|
860
|
+
if keys:
|
|
861
|
+
for fid, name, kind, line in t.syms:
|
|
862
|
+
if t.kind == "word" and not any(w.startswith(t.key) for w in name_parts(name)):
|
|
863
|
+
continue
|
|
864
|
+
if t.kind == "ident" and name.lower() != t.key:
|
|
865
|
+
continue
|
|
866
|
+
self.file_syms[fid][(name, line, kind)] |= keys
|
|
867
|
+
|
|
868
|
+
scored = [self.score_file(f, b) for b, f in pool]
|
|
869
|
+
scored.sort(key=lambda s: s.final, reverse=True)
|
|
870
|
+
scored = scored[:top_k]
|
|
871
|
+
top = scored[0].final if scored else 1.0
|
|
872
|
+
for s in scored:
|
|
873
|
+
s.strength = s.final / top if top else 0.0
|
|
874
|
+
s.confidence = file_confidence(s)
|
|
875
|
+
return scored
|
|
876
|
+
|
|
877
|
+
def score_file(self, fid: int, b: float) -> Scored:
|
|
878
|
+
idx = self.idx
|
|
879
|
+
fa = self.analysis[fid]
|
|
880
|
+
path = idx.paths[fid]
|
|
881
|
+
qterms = self.qterms
|
|
882
|
+
total_w = sum(t.weight for t in qterms) or 1.0
|
|
883
|
+
wmap = {t.key: t.weight for t in qterms}
|
|
884
|
+
|
|
885
|
+
present = [t for t in qterms
|
|
886
|
+
if fa.term_lines.get(t.key) or fid in t.path_fids or (fa.fd is None and fid in t.tf)]
|
|
887
|
+
cov = sum(t.weight for t in present) / total_w
|
|
888
|
+
qhits: dict[int, set] = defaultdict(set)
|
|
889
|
+
for t in qterms:
|
|
890
|
+
for ln in fa.term_lines.get(t.key, ()):
|
|
891
|
+
qhits[ln].add(t.key)
|
|
892
|
+
bw, bline = best_window(qhits, wmap) if qhits else (0.0, -1)
|
|
893
|
+
prox_cov = bw / total_w
|
|
894
|
+
prox = prox_cov * (0.6 + 0.4 * min(1.0, len(qhits) / 6))
|
|
895
|
+
|
|
896
|
+
lw: dict[int, float] = defaultdict(float)
|
|
897
|
+
for t in self.terms.values():
|
|
898
|
+
if t.weight <= 0:
|
|
899
|
+
continue
|
|
900
|
+
for ln in fa.term_lines.get(t.key, ()):
|
|
901
|
+
lw[ln] += t.weight * t.idf
|
|
902
|
+
if fa.fd is not None:
|
|
903
|
+
for ln in list(lw):
|
|
904
|
+
if sym_match(fa.fd.lines[ln]):
|
|
905
|
+
lw[ln] *= 2.5
|
|
906
|
+
fa.line_w = lw
|
|
907
|
+
|
|
908
|
+
defsig, defsym = 0.0, None
|
|
909
|
+
for sym, keys in self.file_syms.get(fid, {}).items():
|
|
910
|
+
f = sum(wmap.get(k, 0.0) for k in keys) / total_w
|
|
911
|
+
if f > defsig:
|
|
912
|
+
defsig, defsym = f, sym
|
|
913
|
+
pathsig = sum(t.weight for t in qterms if fid in t.path_fids) / total_w
|
|
914
|
+
n_idf = math.log(1 + idx.N)
|
|
915
|
+
norm = [min(1.0, t.idf / n_idf) for t in present]
|
|
916
|
+
spec = (0.5 * sum(norm) / len(norm) + 0.5 * max(norm)) if norm else 0.0
|
|
917
|
+
|
|
918
|
+
final = b * (0.5 + cov) * (1 + 1.5 * prox) * self.import_boost.get(path, 1.0)
|
|
919
|
+
return Scored(fid, path, b, final, cov, prox, prox_cov, bline, min(1.0, defsig), defsym,
|
|
920
|
+
min(1.0, pathsig), spec, present=[t.raw for t in present])
|
|
921
|
+
|
|
922
|
+
# ---- snippets ---------------------------------------------------------- #
|
|
923
|
+
def build_snippets(self, s: Scored) -> None:
|
|
924
|
+
fa = self.analysis[s.fid]
|
|
925
|
+
if fa.fd is None:
|
|
926
|
+
return
|
|
927
|
+
lines, lw = fa.fd.lines, fa.line_w
|
|
928
|
+
if not lw:
|
|
929
|
+
s.snippets = [Snippet(1, min(len(lines), 2 * self.ctx + 1), [], lines[:2 * self.ctx + 1])]
|
|
930
|
+
return
|
|
931
|
+
hit_lines = sorted(lw)
|
|
932
|
+
win: dict[int, float] = {}
|
|
933
|
+
lo = 0
|
|
934
|
+
acc = 0.0
|
|
935
|
+
# sliding window sum over +/-6 lines
|
|
936
|
+
hi = 0
|
|
937
|
+
for i in hit_lines:
|
|
938
|
+
while hi < len(hit_lines) and hit_lines[hi] <= i + 6:
|
|
939
|
+
acc += lw[hit_lines[hi]]
|
|
940
|
+
hi += 1
|
|
941
|
+
while hit_lines[lo] < i - 6:
|
|
942
|
+
acc -= lw[hit_lines[lo]]
|
|
943
|
+
lo += 1
|
|
944
|
+
win[i] = acc
|
|
945
|
+
picked: list[int] = []
|
|
946
|
+
for i in sorted(win, key=lambda i: win[i] + 2 * lw[i], reverse=True):
|
|
947
|
+
if all(abs(i - p) > 2 * self.ctx for p in picked):
|
|
948
|
+
picked.append(i)
|
|
949
|
+
if len(picked) >= self.per_file:
|
|
950
|
+
break
|
|
951
|
+
ranges = sorted((max(0, p - self.ctx), min(len(lines) - 1, p + self.ctx)) for p in picked)
|
|
952
|
+
merged: list[list[int]] = []
|
|
953
|
+
for a, b in ranges:
|
|
954
|
+
if merged and a <= merged[-1][1] + 1:
|
|
955
|
+
merged[-1][1] = max(merged[-1][1], b)
|
|
956
|
+
else:
|
|
957
|
+
merged.append([a, b])
|
|
958
|
+
qterms = self.qterms
|
|
959
|
+
for a, b in merged:
|
|
960
|
+
inside = [i for i in hit_lines if a <= i <= b]
|
|
961
|
+
best = max(inside, key=lw.get)
|
|
962
|
+
terms = [t.raw for t in qterms
|
|
963
|
+
if any(a <= ln <= b for ln in fa.term_lines.get(t.key, ()))]
|
|
964
|
+
s.snippets.append(Snippet(a + 1, b + 1, [i + 1 for i in inside], lines[a:b + 1],
|
|
965
|
+
enclosing_symbol(lines, best), terms,
|
|
966
|
+
round(sum(lw[i] for i in inside), 2)))
|
|
967
|
+
|
|
968
|
+
# ---- expansion (replaces the model's "next move") --------------------- #
|
|
969
|
+
def expand(self, results: list[Scored], first: bool) -> int:
|
|
970
|
+
known = set(self.terms) | {stem(k) for k in self.terms}
|
|
971
|
+
new: list[Term] = []
|
|
972
|
+
if first:
|
|
973
|
+
for t in self.qterms:
|
|
974
|
+
if t.kind != "word":
|
|
975
|
+
continue
|
|
976
|
+
for syn in SYNONYMS.get(t.raw, ()) or SYNONYMS.get(t.key, ()):
|
|
977
|
+
if syn not in known and len(new) < 4:
|
|
978
|
+
new.append(make_word(syn, 0.5, "syn"))
|
|
979
|
+
known.add(syn)
|
|
980
|
+
counter: Counter = Counter()
|
|
981
|
+
for s in results[:5]:
|
|
982
|
+
fa = self.analysis.get(s.fid)
|
|
983
|
+
if not fa or fa.fd is None:
|
|
984
|
+
continue
|
|
985
|
+
lines = fa.fd.lines
|
|
986
|
+
for ln in sorted(fa.line_w, key=fa.line_w.get, reverse=True)[:8]:
|
|
987
|
+
w = fa.line_w[ln]
|
|
988
|
+
for a in range(max(0, ln - 1), min(len(lines), ln + 2)):
|
|
989
|
+
mult = 2.0 if sym_match(lines[a]) else 1.0
|
|
990
|
+
for ident in HARVEST_RX.findall(lines[a]):
|
|
991
|
+
low = ident.lower()
|
|
992
|
+
if low in CODE_KEYWORDS or low in STOPWORDS or low in known or stem(low) in known:
|
|
993
|
+
continue
|
|
994
|
+
counter[low] += w * mult
|
|
995
|
+
for low, _ in counter.most_common(MAX_TERMS_PER_TURN * 3):
|
|
996
|
+
if len(new) >= MAX_TERMS_PER_TURN:
|
|
997
|
+
break
|
|
998
|
+
new.append(Term(low, low, 0.4, "expand", [low], [low]))
|
|
999
|
+
known.add(low)
|
|
1000
|
+
for t in new:
|
|
1001
|
+
self.terms[t.key] = t
|
|
1002
|
+
self.follow_imports(results[:3])
|
|
1003
|
+
return len(new)
|
|
1004
|
+
|
|
1005
|
+
def follow_imports(self, results: list[Scored]) -> None:
|
|
1006
|
+
fileset = set(self.idx.paths.values())
|
|
1007
|
+
for s in results:
|
|
1008
|
+
fa = self.analysis.get(s.fid)
|
|
1009
|
+
if not fa or fa.fd is None:
|
|
1010
|
+
continue
|
|
1011
|
+
base_dir = os.path.dirname(s.path)
|
|
1012
|
+
for line in fa.fd.lines[:300]:
|
|
1013
|
+
targets: list[str] = []
|
|
1014
|
+
m = PY_IMPORT_RX.match(line)
|
|
1015
|
+
if m:
|
|
1016
|
+
dotted = m.group(1) or m.group(2)
|
|
1017
|
+
if dotted.startswith("."):
|
|
1018
|
+
up = len(dotted) - len(dotted.lstrip("."))
|
|
1019
|
+
d = base_dir
|
|
1020
|
+
for _ in range(up - 1):
|
|
1021
|
+
d = os.path.dirname(d)
|
|
1022
|
+
rel = os.path.normpath(os.path.join(d, dotted.lstrip(".").replace(".", "/")))
|
|
1023
|
+
targets += [rel + ".py", rel + "/__init__.py"]
|
|
1024
|
+
else:
|
|
1025
|
+
if self._suffix_map is None:
|
|
1026
|
+
self._suffix_map = defaultdict(list)
|
|
1027
|
+
for p in fileset:
|
|
1028
|
+
if p.endswith(".py"):
|
|
1029
|
+
parts = p[:-3].removesuffix("/__init__").split("/")
|
|
1030
|
+
for i in range(len(parts)):
|
|
1031
|
+
self._suffix_map["/".join(parts[i:])].append(p)
|
|
1032
|
+
targets += self._suffix_map.get(dotted.replace(".", "/"), [])[:3]
|
|
1033
|
+
m = JS_IMPORT_RX.search(line)
|
|
1034
|
+
if m:
|
|
1035
|
+
rel = os.path.normpath(os.path.join(base_dir, m.group(1)))
|
|
1036
|
+
targets += [rel + e for e in ("", ".js", ".jsx", ".ts", ".tsx", "/index.js", "/index.ts")]
|
|
1037
|
+
for tgt in targets:
|
|
1038
|
+
tgt = tgt.replace("\\", "/")
|
|
1039
|
+
if tgt in fileset:
|
|
1040
|
+
self.import_boost[tgt] = 1.15
|
|
1041
|
+
|
|
1042
|
+
# ---- confidence -------------------------------------------------------- #
|
|
1043
|
+
def assess(self, results: list[Scored]) -> dict:
|
|
1044
|
+
qterms = self.qterms
|
|
1045
|
+
unmatched = [t.raw for t in qterms if t.kind == "word" and t.df == 0 and not t.path_fids]
|
|
1046
|
+
total_w = sum(t.weight for t in qterms) or 1.0
|
|
1047
|
+
unmatched_frac = sum(t.weight for t in qterms if t.raw in unmatched) / total_w
|
|
1048
|
+
if not results:
|
|
1049
|
+
return {"score": 0.0, "label": "very_low", "unmatched_terms": unmatched,
|
|
1050
|
+
"reasons": ["no matches found for any query term"],
|
|
1051
|
+
"suggestion": "check spelling, try different keywords, or use exact identifier names"}
|
|
1052
|
+
c = [r.confidence for r in results[:3]]
|
|
1053
|
+
score = 0.65 * c[0] + 0.25 * (sum(c) / len(c)) + 0.10 * (100 if self.converged else 60)
|
|
1054
|
+
score = round(max(0.0, score * (1 - 0.5 * unmatched_frac)), 1)
|
|
1055
|
+
|
|
1056
|
+
top = results[0]
|
|
1057
|
+
reasons = []
|
|
1058
|
+
if top.defsym and top.defsig >= 0.5:
|
|
1059
|
+
name, line, kind = top.defsym
|
|
1060
|
+
reasons.append(f"definition match: {name} ({kind}) at {top.path}:{line}")
|
|
1061
|
+
nwords = sum(1 for t in qterms if t.kind == "word")
|
|
1062
|
+
if nwords >= 2 and top.prox_cov >= 0.99:
|
|
1063
|
+
reasons.append(f"all query terms appear together within ~8 lines near {top.path}:{top.prox_line + 1}")
|
|
1064
|
+
elif top.cov < 0.6:
|
|
1065
|
+
reasons.append(f"best file covers only {int(top.cov * 100)}% of the query's weight")
|
|
1066
|
+
if top.pathsig >= 0.5:
|
|
1067
|
+
reasons.append("file path matches the query terms")
|
|
1068
|
+
if unmatched:
|
|
1069
|
+
reasons.append("not found anywhere in the codebase: " + ", ".join(unmatched))
|
|
1070
|
+
if top.spec < 0.45:
|
|
1071
|
+
reasons.append("query terms are very common in this codebase (low specificity)")
|
|
1072
|
+
if len(results) > 1 and results[1].final >= 0.9 * top.final:
|
|
1073
|
+
reasons.append("several files score almost equally; the answer may span multiple files")
|
|
1074
|
+
out = {"score": score, "label": label_of(score), "unmatched_terms": unmatched, "reasons": reasons}
|
|
1075
|
+
if score < 50:
|
|
1076
|
+
out["suggestion"] = ("add exact identifier names (function/class), quote a multi-word phrase, "
|
|
1077
|
+
"or narrow with --ext/--glob")
|
|
1078
|
+
return out
|
|
1079
|
+
|
|
1080
|
+
def definitions(self, results: list[Scored], limit: int = 8) -> list[dict]:
|
|
1081
|
+
wmap = {t.key: t.weight for t in self.qterms}
|
|
1082
|
+
total_w = sum(wmap.values()) or 1.0
|
|
1083
|
+
rows = []
|
|
1084
|
+
for fid, syms in self.file_syms.items():
|
|
1085
|
+
if fid not in self.base:
|
|
1086
|
+
continue
|
|
1087
|
+
for (name, line, kind), keys in syms.items():
|
|
1088
|
+
frac = sum(wmap.get(k, 0.0) for k in keys) / total_w
|
|
1089
|
+
if frac >= (0.6 if len(wmap) > 1 else 0.5):
|
|
1090
|
+
rows.append((frac, self.base[fid], fid, name, kind, line))
|
|
1091
|
+
rows.sort(reverse=True)
|
|
1092
|
+
return [{"name": n, "kind": k, "path": self.idx.paths[f], "line": ln, "match": round(min(1, fr), 2)}
|
|
1093
|
+
for fr, _, f, n, k, ln in rows[:limit]]
|
|
1094
|
+
|
|
1095
|
+
# ---- main loop --------------------------------------------------------- #
|
|
1096
|
+
def run(self, query: str, top_k: int = 8) -> dict:
|
|
1097
|
+
t_all = time.perf_counter()
|
|
1098
|
+
t0 = time.perf_counter()
|
|
1099
|
+
stats = self.idx.sync()
|
|
1100
|
+
self.idx.load()
|
|
1101
|
+
self.timings["index_sync"] = time.perf_counter() - t0
|
|
1102
|
+
self.log(f"[fc] index: {stats['files']} files ({stats['changed']} re-indexed, "
|
|
1103
|
+
f"{stats['removed']} removed) in {self.timings['index_sync'] * 1000:.0f} ms")
|
|
1104
|
+
|
|
1105
|
+
self.build_terms(query)
|
|
1106
|
+
results: list[Scored] = []
|
|
1107
|
+
assessment: dict = {}
|
|
1108
|
+
prev_score = 0.0
|
|
1109
|
+
for turn in range(1, self.max_turns + 1):
|
|
1110
|
+
pending = [t for t in self.terms.values() if not t.looked_up]
|
|
1111
|
+
if not pending:
|
|
1112
|
+
break
|
|
1113
|
+
t0 = time.perf_counter()
|
|
1114
|
+
for t in pending:
|
|
1115
|
+
self.lookup(t)
|
|
1116
|
+
self.timings["lookup"] += time.perf_counter() - t0
|
|
1117
|
+
results = self.rank(top_k)
|
|
1118
|
+
self.turns_used = turn
|
|
1119
|
+
self.converged = True
|
|
1120
|
+
assessment = self.assess(results)
|
|
1121
|
+
self.log(f"[turn {turn}] +{len(pending)} terms: {[t.raw for t in pending][:10]} "
|
|
1122
|
+
f"-> confidence {assessment['score']} ({assessment['label']})")
|
|
1123
|
+
|
|
1124
|
+
if not self.thorough and assessment["score"] >= EARLY_EXIT_CONFIDENCE:
|
|
1125
|
+
self.log(" confident enough; stopping early")
|
|
1126
|
+
break
|
|
1127
|
+
if turn >= 2 and assessment["score"] - prev_score < 1.0:
|
|
1128
|
+
self.log(" no further improvement; stopping")
|
|
1129
|
+
break
|
|
1130
|
+
prev_score = assessment["score"]
|
|
1131
|
+
if turn < self.max_turns:
|
|
1132
|
+
added = self.expand(results, first=(turn == 1))
|
|
1133
|
+
self.log(f" expanded with {added} terms")
|
|
1134
|
+
if added == 0:
|
|
1135
|
+
break
|
|
1136
|
+
else:
|
|
1137
|
+
self.converged = False
|
|
1138
|
+
assessment = self.assess(results)
|
|
1139
|
+
|
|
1140
|
+
for s in results:
|
|
1141
|
+
self.build_snippets(s)
|
|
1142
|
+
self.timings["total"] = time.perf_counter() - t_all
|
|
1143
|
+
return {
|
|
1144
|
+
"query": query, "turns": self.turns_used, "results": results,
|
|
1145
|
+
"confidence": assessment or self.assess([]),
|
|
1146
|
+
"definitions": self.definitions(results),
|
|
1147
|
+
"elapsed_ms": round(self.timings["total"] * 1000),
|
|
1148
|
+
"index": stats,
|
|
1149
|
+
}
|
|
1150
|
+
|
|
1151
|
+
|
|
1152
|
+
# --------------------------------------------------------------------------- #
|
|
1153
|
+
# Output
|
|
1154
|
+
# --------------------------------------------------------------------------- #
|
|
1155
|
+
def _numbered(sn: Snippet, width_cap: int = 200) -> str:
|
|
1156
|
+
hits, w = set(sn.hit_lines), len(str(sn.end))
|
|
1157
|
+
return "\n".join(f"{'>' if sn.start + i in hits else ' '}{sn.start + i:>{w}}: {line.rstrip()[:width_cap]}"
|
|
1158
|
+
for i, line in enumerate(sn.text))
|
|
1159
|
+
|
|
1160
|
+
|
|
1161
|
+
def _attr(v) -> str:
|
|
1162
|
+
return str(v).replace("&", "&").replace('"', """).replace("<", "<")
|
|
1163
|
+
|
|
1164
|
+
|
|
1165
|
+
def apply_budget(results: list[Scored], max_tokens: int) -> int:
|
|
1166
|
+
"""Trim snippets to fit an approximate token budget (4 chars ~ 1 token). Returns dropped count."""
|
|
1167
|
+
if max_tokens <= 0:
|
|
1168
|
+
return 0
|
|
1169
|
+
budget, dropped = max_tokens * 4, 0
|
|
1170
|
+
used = 0
|
|
1171
|
+
for rank, s in enumerate(results):
|
|
1172
|
+
kept = []
|
|
1173
|
+
for i, sn in enumerate(s.snippets):
|
|
1174
|
+
cost = len(_numbered(sn)) + 80
|
|
1175
|
+
if used + cost <= budget or (rank == 0 and i == 0):
|
|
1176
|
+
kept.append(sn)
|
|
1177
|
+
used += cost
|
|
1178
|
+
else:
|
|
1179
|
+
dropped += 1
|
|
1180
|
+
s.snippets = kept
|
|
1181
|
+
for s in [r for r in results if not r.snippets]:
|
|
1182
|
+
results.remove(s)
|
|
1183
|
+
return dropped
|
|
1184
|
+
|
|
1185
|
+
|
|
1186
|
+
def format_llm(out: dict, explain: bool = False) -> str:
|
|
1187
|
+
conf, res = out["confidence"], out["results"]
|
|
1188
|
+
x = [f'<fast_context query="{_attr(out["query"])}" files="{len(res)}" turns="{out["turns"]}" '
|
|
1189
|
+
f'confidence="{conf["score"]}" confidence_label="{conf["label"]}" elapsed_ms="{out["elapsed_ms"]}">']
|
|
1190
|
+
x.append(f' <confidence score="{conf["score"]}" label="{conf["label"]}">')
|
|
1191
|
+
for r in conf["reasons"]:
|
|
1192
|
+
x.append(f" <reason>{_attr(r)}</reason>")
|
|
1193
|
+
if conf.get("suggestion"):
|
|
1194
|
+
x.append(f" <suggestion>{_attr(conf['suggestion'])}</suggestion>")
|
|
1195
|
+
x.append(" </confidence>")
|
|
1196
|
+
if out["definitions"]:
|
|
1197
|
+
x.append(" <definitions>")
|
|
1198
|
+
for d in out["definitions"]:
|
|
1199
|
+
x.append(f" {d['name']} ({d['kind']}) {d['path']}:{d['line']}")
|
|
1200
|
+
x.append(" </definitions>")
|
|
1201
|
+
if res:
|
|
1202
|
+
x.append(" <index>")
|
|
1203
|
+
for s in res:
|
|
1204
|
+
for sn in s.snippets:
|
|
1205
|
+
sym = f" # {sn.symbol}" if sn.symbol else ""
|
|
1206
|
+
x.append(f" {s.path}:{sn.start}-{sn.end}{sym} [confidence {s.confidence:.0f}]")
|
|
1207
|
+
x.append(" </index>")
|
|
1208
|
+
for rank, s in enumerate(res, 1):
|
|
1209
|
+
lang = language_of(s.path)
|
|
1210
|
+
extra = ""
|
|
1211
|
+
if explain:
|
|
1212
|
+
extra = (f' signals="coverage={s.cov:.2f} proximity={s.prox:.2f} definition={s.defsig:.2f} '
|
|
1213
|
+
f'path={s.pathsig:.2f} strength={s.strength:.2f} specificity={s.spec:.2f}"')
|
|
1214
|
+
x.append(f' <file path="{_attr(s.path)}" rank="{rank}" score="{s.final:.1f}" '
|
|
1215
|
+
f'confidence="{s.confidence:.0f}" confidence_label="{label_of(s.confidence)}" '
|
|
1216
|
+
f'language="{lang}"{extra}>')
|
|
1217
|
+
x.append(f" <matched_terms>{_attr(', '.join(dict.fromkeys(s.present)))}</matched_terms>")
|
|
1218
|
+
for sn in s.snippets:
|
|
1219
|
+
attrs = f'lines="{sn.start}-{sn.end}" hit_lines="{",".join(map(str, sn.hit_lines))}"'
|
|
1220
|
+
if sn.symbol:
|
|
1221
|
+
attrs += f' symbol="{_attr(sn.symbol)}"'
|
|
1222
|
+
if sn.terms:
|
|
1223
|
+
attrs += f' terms="{_attr(", ".join(sn.terms))}"'
|
|
1224
|
+
x.append(f" <snippet {attrs}>")
|
|
1225
|
+
x.append(f"```{lang}")
|
|
1226
|
+
x.append(_numbered(sn))
|
|
1227
|
+
x.append("```")
|
|
1228
|
+
x.append(" </snippet>")
|
|
1229
|
+
x.append(" </file>")
|
|
1230
|
+
x.append("</fast_context>")
|
|
1231
|
+
return "\n".join(x)
|
|
1232
|
+
|
|
1233
|
+
|
|
1234
|
+
def format_json(out: dict, explain: bool = False) -> str:
|
|
1235
|
+
return json.dumps({
|
|
1236
|
+
"query": out["query"], "turns": out["turns"], "elapsed_ms": out["elapsed_ms"],
|
|
1237
|
+
"confidence": out["confidence"],
|
|
1238
|
+
"definitions": out["definitions"],
|
|
1239
|
+
"results": [{
|
|
1240
|
+
"rank": i, "path": s.path, "language": language_of(s.path), "score": round(s.final, 3),
|
|
1241
|
+
"confidence": {"score": s.confidence, "label": label_of(s.confidence), "signals": {
|
|
1242
|
+
"coverage": round(s.cov, 3), "proximity": round(s.prox, 3), "definition": round(s.defsig, 3),
|
|
1243
|
+
"path": round(s.pathsig, 3), "strength": round(s.strength, 3), "specificity": round(s.spec, 3)}},
|
|
1244
|
+
"matched_terms": list(dict.fromkeys(s.present)),
|
|
1245
|
+
"snippets": [{
|
|
1246
|
+
"location": f"{s.path}:{sn.start}-{sn.end}", "start_line": sn.start, "end_line": sn.end,
|
|
1247
|
+
"hit_lines": sn.hit_lines, "symbol": sn.symbol, "terms": sn.terms,
|
|
1248
|
+
"code": "\n".join(sn.text), "code_numbered": _numbered(sn)} for sn in s.snippets],
|
|
1249
|
+
} for i, s in enumerate(out["results"], 1)],
|
|
1250
|
+
}, indent=2)
|
|
1251
|
+
|
|
1252
|
+
|
|
1253
|
+
def format_paths(out: dict, explain: bool = False) -> str:
|
|
1254
|
+
lines = [f"# confidence {out['confidence']['score']} ({out['confidence']['label']})"]
|
|
1255
|
+
for s in out["results"]:
|
|
1256
|
+
for sn in s.snippets:
|
|
1257
|
+
sym = f" # {sn.symbol}" if sn.symbol else ""
|
|
1258
|
+
lines.append(f"{s.path}:{sn.start}-{sn.end}{sym} [{s.confidence:.0f}]")
|
|
1259
|
+
return "\n".join(lines)
|
|
1260
|
+
|
|
1261
|
+
|
|
1262
|
+
def format_text(out: dict, explain: bool = False) -> str:
|
|
1263
|
+
conf = out["confidence"]
|
|
1264
|
+
bar = "#" * int(conf["score"] / 10) + "." * (10 - int(conf["score"] / 10))
|
|
1265
|
+
x = [f"Query: {out['query']}",
|
|
1266
|
+
f"Confidence: {conf['score']:.0f}/100 [{bar}] {conf['label']} "
|
|
1267
|
+
f"({out['turns']} turn(s), {out['elapsed_ms']} ms)"]
|
|
1268
|
+
x += [f" - {r}" for r in conf["reasons"]]
|
|
1269
|
+
if conf.get("suggestion"):
|
|
1270
|
+
x.append(f" tip: {conf['suggestion']}")
|
|
1271
|
+
if out["definitions"]:
|
|
1272
|
+
x.append("\nDefinitions:")
|
|
1273
|
+
x += [f" {d['name']} ({d['kind']}) {d['path']}:{d['line']}" for d in out["definitions"]]
|
|
1274
|
+
if not out["results"]:
|
|
1275
|
+
x.append("\nNo relevant code found.")
|
|
1276
|
+
for i, s in enumerate(out["results"], 1):
|
|
1277
|
+
x.append(f"\n{i}. {s.path} confidence {s.confidence:.0f} ({label_of(s.confidence)})")
|
|
1278
|
+
x.append(f" matched: {', '.join(dict.fromkeys(s.present))}")
|
|
1279
|
+
if explain:
|
|
1280
|
+
x.append(f" signals: cov={s.cov:.2f} prox={s.prox:.2f} def={s.defsig:.2f} "
|
|
1281
|
+
f"path={s.pathsig:.2f} strength={s.strength:.2f} spec={s.spec:.2f}")
|
|
1282
|
+
for sn in s.snippets:
|
|
1283
|
+
sym = f" [{sn.symbol}]" if sn.symbol else ""
|
|
1284
|
+
x.append(f" --- {s.path}:{sn.start}-{sn.end}{sym}")
|
|
1285
|
+
x += [" " + row for row in _numbered(sn, 160).split("\n")]
|
|
1286
|
+
return "\n".join(x)
|
|
1287
|
+
|
|
1288
|
+
|
|
1289
|
+
FORMATTERS = {"llm": format_llm, "json": format_json, "text": format_text, "paths": format_paths}
|
|
1290
|
+
|
|
1291
|
+
|
|
1292
|
+
# --------------------------------------------------------------------------- #
|
|
1293
|
+
# CLI
|
|
1294
|
+
# --------------------------------------------------------------------------- #
|
|
1295
|
+
def main() -> None:
|
|
1296
|
+
ap = argparse.ArgumentParser(description="LLM-free Fast Context with confidence scoring.")
|
|
1297
|
+
ap.add_argument("query", nargs="?", help='e.g. "lazy load" (quote phrases, use identifiers freely)')
|
|
1298
|
+
ap.add_argument("--root", default=".", help="codebase root (default: current directory)")
|
|
1299
|
+
ap.add_argument("--top", type=int, default=6, help="max files to return (default 6)")
|
|
1300
|
+
ap.add_argument("--context", type=int, default=3, help="context lines around each hit")
|
|
1301
|
+
ap.add_argument("--per-file", type=int, default=3, help="max snippets per file")
|
|
1302
|
+
ap.add_argument("--turns", type=int, default=MAX_TURNS, help=f"max turns (<= {MAX_TURNS})")
|
|
1303
|
+
ap.add_argument("--format", choices=FORMATTERS, default="llm",
|
|
1304
|
+
help="llm (default, XML-tagged), json, text, or paths (path:lines list)")
|
|
1305
|
+
ap.add_argument("--json", action="store_true", help="shortcut for --format json")
|
|
1306
|
+
ap.add_argument("--max-tokens", type=int, default=6000,
|
|
1307
|
+
help="approx token budget for snippets (default 6000, 0 = unlimited)")
|
|
1308
|
+
ap.add_argument("--ext", help="only these extensions, e.g. py,ts")
|
|
1309
|
+
ap.add_argument("--glob", help='only paths matching glob(s), e.g. "src/*.py,lib/*"')
|
|
1310
|
+
ap.add_argument("--tests", action="store_true", help="don't down-rank test files")
|
|
1311
|
+
ap.add_argument("--thorough", action="store_true", help="always run all turns (no early exit)")
|
|
1312
|
+
ap.add_argument("--no-cache", action="store_true", help="don't use the persistent index")
|
|
1313
|
+
ap.add_argument("--reindex", action="store_true", help="rebuild the index from scratch")
|
|
1314
|
+
ap.add_argument("--stats", action="store_true", help="print index statistics and exit")
|
|
1315
|
+
ap.add_argument("--explain", action="store_true", help="show the confidence signals per file")
|
|
1316
|
+
ap.add_argument("-v", "--verbose", action="store_true", help="log turns and timings to stderr")
|
|
1317
|
+
args = ap.parse_args()
|
|
1318
|
+
|
|
1319
|
+
root = os.path.abspath(args.root)
|
|
1320
|
+
if not os.path.isdir(root):
|
|
1321
|
+
ap.error(f"--root {args.root!r} is not a directory")
|
|
1322
|
+
|
|
1323
|
+
fc = Retriever(root, use_cache=not args.no_cache, reindex=args.reindex, max_turns=args.turns,
|
|
1324
|
+
context=args.context, per_file=args.per_file, include_tests=args.tests,
|
|
1325
|
+
exts=args.ext.split(",") if args.ext else None,
|
|
1326
|
+
globs=args.glob.split(",") if args.glob else None,
|
|
1327
|
+
thorough=args.thorough, verbose=args.verbose)
|
|
1328
|
+
if args.stats:
|
|
1329
|
+
fc.idx.sync()
|
|
1330
|
+
print(json.dumps(fc.idx.stats(), indent=2))
|
|
1331
|
+
return
|
|
1332
|
+
if not args.query:
|
|
1333
|
+
ap.error("a query is required, e.g. python -m shelldeck.addons.fast_context \"lazy load\"")
|
|
1334
|
+
|
|
1335
|
+
out = fc.run(args.query, top_k=args.top)
|
|
1336
|
+
dropped = apply_budget(out["results"], args.max_tokens)
|
|
1337
|
+
if args.verbose:
|
|
1338
|
+
tm = fc.timings
|
|
1339
|
+
print(f"[fc] timings ms: index_sync={tm['index_sync'] * 1000:.0f} lookup={tm['lookup'] * 1000:.0f} "
|
|
1340
|
+
f"rank={tm['rank_index'] * 1000:.0f} read+scan={tm['read_scan'] * 1000:.0f} "
|
|
1341
|
+
f"total={tm['total'] * 1000:.0f}; snippets dropped by budget: {dropped}", file=sys.stderr)
|
|
1342
|
+
print(FORMATTERS["json" if args.json else args.format](out, args.explain))
|
|
1343
|
+
|
|
1344
|
+
|
|
1345
|
+
if __name__ == "__main__":
|
|
1346
|
+
main()
|