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.
Files changed (45) hide show
  1. shelldeck/__init__.py +0 -0
  2. shelldeck/__main__.py +4 -0
  3. shelldeck/addons/__init__.py +1 -0
  4. shelldeck/addons/fast_context.py +1346 -0
  5. shelldeck/auth.py +207 -0
  6. shelldeck/cli.py +587 -0
  7. shelldeck/db.py +794 -0
  8. shelldeck/gitgraph.py +150 -0
  9. shelldeck/integration/bash.sh +27 -0
  10. shelldeck/integration/shelldeck.fish +5 -0
  11. shelldeck/integration/shelldeck.ps1 +35 -0
  12. shelldeck/integration/zsh/.zshenv +6 -0
  13. shelldeck/integration/zsh/.zshrc +10 -0
  14. shelldeck/pty.py +352 -0
  15. shelldeck/runner.py +75 -0
  16. shelldeck/scheduler.py +283 -0
  17. shelldeck/server.py +1216 -0
  18. shelldeck/shells.py +174 -0
  19. shelldeck/static/app.css +2288 -0
  20. shelldeck/static/app.js +1863 -0
  21. shelldeck/static/gitgraph.js +142 -0
  22. shelldeck/static/history.js +172 -0
  23. shelldeck/static/icon-192.png +0 -0
  24. shelldeck/static/icon-32.png +0 -0
  25. shelldeck/static/icon-512.png +0 -0
  26. shelldeck/static/icon.svg +1 -0
  27. shelldeck/static/index.html +102 -0
  28. shelldeck/static/manifest.webmanifest +1 -0
  29. shelldeck/static/monitor.js +143 -0
  30. shelldeck/static/ui.js +316 -0
  31. shelldeck/static/vendor/LICENSE-xterm.txt +21 -0
  32. shelldeck/static/vendor/addon-fit.js +2 -0
  33. shelldeck/static/vendor/addon-search.js +2 -0
  34. shelldeck/static/vendor/addon-serialize.js +2 -0
  35. shelldeck/static/vendor/addon-web-links.js +2 -0
  36. shelldeck/static/vendor/addon-webgl.js +2 -0
  37. shelldeck/static/vendor/xterm.css +218 -0
  38. shelldeck/static/vendor/xterm.js +2 -0
  39. shelldeck/static/views.js +658 -0
  40. shelldeck/stats.py +128 -0
  41. shelldeck-0.0.1.dist-info/METADATA +254 -0
  42. shelldeck-0.0.1.dist-info/RECORD +45 -0
  43. shelldeck-0.0.1.dist-info/WHEEL +4 -0
  44. shelldeck-0.0.1.dist-info/entry_points.txt +3 -0
  45. 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("&", "&amp;").replace('"', "&quot;").replace("<", "&lt;")
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()