Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
51 changes: 21 additions & 30 deletions src/semble/cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,14 +2,14 @@
import json
import logging
import os
import shutil
import sys
from collections.abc import Sequence
from pathlib import Path
from typing import TYPE_CHECKING

import orjson

from semble.chunking.chunking import _DESIRED_CHUNK_LENGTH_CHARS
from semble.index.bm25 import BM25
from semble.index.dense import SelectableBasicBackend
from semble.index.file_walker import walk_files
Expand Down Expand Up @@ -89,11 +89,6 @@ def resolve_cache_folder() -> Path:
return cache_dir


def clear_cache(path: str) -> None:
"""Clear all exact content indexes for the given path."""
shutil.rmtree(find_index_from_cache_folder(path).parent, ignore_errors=True)


def save_index_to_cache(index: "SembleIndex", path: str) -> None:
"""Save an index to the cache folder if it was freshly built."""
if not index.loaded_from_disk:
Expand All @@ -102,8 +97,6 @@ def save_index_to_cache(index: "SembleIndex", path: str) -> None:

def _metadata_matches(metadata: dict, model_path: str, content: Sequence[ContentType]) -> bool:
"""Return True if the stored metadata is compatible with the requested parameters."""
from semble.chunking.chunking import _DESIRED_CHUNK_LENGTH_CHARS # avoid circular import at module level

try:
content_type = tuple(ContentType(s) for s in metadata["content_type"])
# chunk_size and cache_version are absent in indexes built before those fields were added;
Expand All @@ -117,22 +110,28 @@ def _metadata_matches(metadata: dict, model_path: str, content: Sequence[Content
return False


def get_validated_cache(path: str, model_path: str | None, content: Sequence[ContentType]) -> Path | None:
"""Validates the cache folder and returns the index path."""
index_path = find_index_from_cache_folder(path, content)
if not index_path.exists():
return None

persistence_path = PersistencePath.from_path(index_path)
def _load_matching_metadata(
path: str, model_path: str | None, content: Sequence[ContentType]
) -> tuple[PersistencePath, dict] | None:
"""Return the cached index files and metadata for path, or None if absent or built with other settings."""
persistence_path = PersistencePath.from_path(find_index_from_cache_folder(path, content))
if persistence_path.non_existing():
return None

metadata = json.loads(persistence_path.metadata.read_text(encoding="utf-8"))
if model_path is None:
model_path = resolve_model_name()
with open(persistence_path.metadata, encoding="utf-8") as f:
metadata = json.load(f)
if not _metadata_matches(metadata, model_path, content):
return None
return persistence_path, metadata


def get_validated_cache(path: str, model_path: str | None, content: Sequence[ContentType]) -> Path | None:
"""Validates the cache folder and returns the index path."""
loaded = _load_matching_metadata(path, model_path, content)
if loaded is None:
return None
persistence_path, metadata = loaded
index_path = persistence_path.metadata.parent

if is_git_url(str(path)):
return index_path
Expand Down Expand Up @@ -168,25 +167,17 @@ def load_previous_for_incremental(
:return: Previous index state, or None if the cache is unavailable or invalid.
"""
try:
index_path = find_index_from_cache_folder(path, content)
persistence_path = PersistencePath.from_path(index_path)
if persistence_path.non_existing():
return None

if model_path is None:
model_path = resolve_model_name()
with open(persistence_path.metadata, encoding="utf-8") as f:
metadata = json.load(f)
if not _metadata_matches(metadata, model_path, content):
loaded = _load_matching_metadata(path, model_path, content)
if loaded is None:
return None
persistence_path, metadata = loaded

raw_manifest = metadata.get("files")
if not raw_manifest:
return None
manifest = {indexed_path: FileManifestEntry(**entry) for indexed_path, entry in raw_manifest.items()}

with open(persistence_path.chunks, "rb") as f:
chunks = [Chunk.from_dict(item) for item in orjson.loads(f.read())]
chunks = [Chunk.from_dict(item) for item in orjson.loads(persistence_path.chunks.read_bytes())]

vectors = SelectableBasicBackend.load(persistence_path.semantic_index).vectors
bm25_index = BM25.load(persistence_path.bm25_index)
Expand Down
74 changes: 32 additions & 42 deletions src/semble/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
import re
import sys
import warnings
from collections.abc import Iterator
from importlib.util import find_spec
from pathlib import Path
from shutil import rmtree
Expand Down Expand Up @@ -49,6 +50,26 @@ def _maybe_save_index(parts: list[tuple[str, SembleIndex]]) -> None:
print(f"Error saving index: {e}", file=sys.stderr)


def _add_query_args(p: argparse.ArgumentParser) -> None:
"""Add the path, result-shaping, and content arguments shared by search and find-related."""
p.add_argument(
"path",
nargs="*",
default=["."],
help="Local paths or git URLs to search together (default: current directory).",
)
p.add_argument("-k", "--top-k", type=int, default=5, help="Number of results (default: 5).")
p.add_argument(
"--max-snippet-lines",
type=int,
default=None,
metavar="N",
help="Lines of source per result (default: full chunk). 10 = signature + body, 0 = no code.",
)
p.add_argument("--format", choices=["json", "text"], default="json", help="Output format (default: json).")
_add_content_args(p)


def _add_content_args(p: argparse.ArgumentParser) -> None:
"""Add --content and deprecated --include-text-files to a subparser."""
p.add_argument(
Expand Down Expand Up @@ -178,15 +199,16 @@ def _run_find_related(
_maybe_save_index(parts)


def _cached_index_paths(cache_folder: Path) -> Iterator[Path]:
"""Yield index folders in the cache that sit under a sha256 cache key."""
return (path for path in cache_folder.glob("*/index*") if _SHA_256_REGEX.match(path.parent.name))


def _clear_indexes(cache_folder: Path) -> None:
"""Remove all valid index entries from the cache folder."""
indexes: set[Path] = set()
for path in cache_folder.glob("*/index*"):
if not _SHA_256_REGEX.match(path.parent.name):
continue
if PersistencePath.from_path(path).non_existing():
continue
indexes.add(path.parent)
indexes = {
path.parent for path in _cached_index_paths(cache_folder) if not PersistencePath.from_path(path).non_existing()
}

if not indexes:
print(f"No indexes found to clear in `{cache_folder}`")
Expand All @@ -209,9 +231,7 @@ def _clear_savings(cache_folder: Path) -> None:
def _clear_orphans(cache_folder: Path) -> None:
"""Remove index entries whose local root_path no longer exists."""
orphans: dict[Path, str] = {}
for path in cache_folder.glob("*/index*"):
if not _SHA_256_REGEX.match(path.parent.name):
continue
for path in _cached_index_paths(cache_folder):
try:
with open(path / "metadata.json", encoding="utf-8") as f:
metadata = json.load(f)
Expand Down Expand Up @@ -267,22 +287,7 @@ def _cli_main() -> None:

search_p = sub.add_parser("search", help="Search a codebase.")
search_p.add_argument("query", help="Natural language or code query.")
search_p.add_argument(
"path",
nargs="*",
default=["."],
help="Local paths or git URLs to search together (default: current directory).",
)
search_p.add_argument("-k", "--top-k", type=int, default=5, help="Number of results (default: 5).")
search_p.add_argument(
"--max-snippet-lines",
type=int,
default=None,
metavar="N",
help="Lines of source per result (default: full chunk). 10 = signature + body, 0 = no code.",
)
search_p.add_argument("--format", choices=["json", "text"], default="json", help="Output format (default: json).")
_add_content_args(search_p)
_add_query_args(search_p)

clear_p = sub.add_parser("clear", help="Clear the index cache.")
clear_p.add_argument(
Expand All @@ -294,22 +299,7 @@ def _cli_main() -> None:
related_p = sub.add_parser("find-related", help="Find code similar to a specific location.")
related_p.add_argument("file_path", help="File path as shown in search results.")
related_p.add_argument("line", type=int, help="Line number (1-indexed).")
related_p.add_argument(
"path",
nargs="*",
default=["."],
help="Local paths or git URLs to search together (default: current directory).",
)
related_p.add_argument("-k", "--top-k", type=int, default=5, help="Number of results (default: 5).")
related_p.add_argument(
"--max-snippet-lines",
type=int,
default=None,
metavar="N",
help="Lines of source per result (default: full chunk). 10 = signature + body, 0 = no code.",
)
related_p.add_argument("--format", choices=["json", "text"], default="json", help="Output format (default: json).")
_add_content_args(related_p)
_add_query_args(related_p)

sub.add_parser("savings", help="Show token savings and usage stats.")

Expand Down
8 changes: 2 additions & 6 deletions src/semble/index/bm25.py
Original file line number Diff line number Diff line change
Expand Up @@ -114,11 +114,7 @@ def load(cls, path: Path) -> "BM25":
documents = data["documents"]
if len(doc_order) != len(set(doc_order)) or set(documents) != set(doc_order):
raise ValueError("Persisted BM25 document state is inconsistent")
index._documents = {chunk_id: Counter(counts) for chunk_id, counts in documents.items()}
for chunk_id, counts in index._documents.items():
for term, count in counts.items():
index.postings.setdefault(term, {})[chunk_id] = count
index._doc_lengths = {chunk_id: sum(counts.values()) for chunk_id, counts in index._documents.items()}
index._total_doc_length = sum(index._doc_lengths.values())
for chunk_id, counts in documents.items():
index._add_counts(chunk_id, Counter(counts), sum(counts.values()))
index.set_doc_order(doc_order)
return index
11 changes: 3 additions & 8 deletions src/semble/index/dense.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,14 +28,9 @@ def _load_cached(model_path: str) -> StaticModel:
disable_progress_bars()
logging.getLogger("huggingface_hub.utils._http").addFilter(_drop_unauthenticated_warning)
try:
try:
model = StaticModel.from_pretrained(model_path, force_download=False)
except ValueError:
model = StaticModel.from_pretrained(model_path, force_download=True)
finally:
disable_progress_bars()

return model
return StaticModel.from_pretrained(model_path, force_download=False)
except ValueError:
return StaticModel.from_pretrained(model_path, force_download=True)


def load_model(model_path: str | None = None) -> tuple[StaticModel, str]:
Expand Down
38 changes: 16 additions & 22 deletions src/semble/index/files.py
Original file line number Diff line number Diff line change
Expand Up @@ -444,22 +444,16 @@
}


def _inv_mapping(mapping: dict[str, str]) -> dict[str, list[str]]:
"""Invert a mapping, taking into account duplicate values."""
inv: defaultdict[str, list[str]] = defaultdict(list)
for key, value in mapping.items():
inv[value].append(key)
return dict(inv)


ALL_LANGUAGES = frozenset(_EXTENSION_TO_LANGUAGE.values())
_CODE_LANGUAGES = ALL_LANGUAGES - _DOC_LANGUAGES - _CONFIG_LANGUAGES - _DATA_LANGUAGES
_LANGUAGE_TO_EXTENSION = _inv_mapping(_EXTENSION_TO_LANGUAGE)
_LANGUAGE_TO_EXTENSIONS: defaultdict[str, list[str]] = defaultdict(list)
for _extension, _language in _EXTENSION_TO_LANGUAGE.items():
_LANGUAGE_TO_EXTENSIONS[_language].append(_extension)

_CONTENT_TYPE_LANGUAGES: dict[ContentType, frozenset[str]] = {
ContentType.CODE: frozenset(_CODE_LANGUAGES),
ContentType.DOCS: frozenset(_DOC_LANGUAGES),
ContentType.CONFIG: frozenset(_CONFIG_LANGUAGES),
_CONTENT_TYPE_LANGUAGES = {
ContentType.CODE: _CODE_LANGUAGES,
ContentType.DOCS: _DOC_LANGUAGES,
ContentType.CONFIG: _CONFIG_LANGUAGES,
}


Expand All @@ -470,14 +464,14 @@ def detect_language(file_name: Path) -> str | None:

def get_extensions(types: Sequence[ContentType]) -> list[str]:
"""Returns a list of supported file extensions for the given content types."""
languages: set[str] = set()
for content_type in types:
languages.update(_CONTENT_TYPE_LANGUAGES[content_type])
all_extensions: set[str] = set()
for language in languages:
all_extensions.update(_LANGUAGE_TO_EXTENSION.get(language, set()))

return sorted(all_extensions)
return sorted(
{
ext
for content_type in types
for lang in _CONTENT_TYPE_LANGUAGES[content_type]
for ext in _LANGUAGE_TO_EXTENSIONS.get(lang, [])
}
)


class FileStatus(str, Enum):
Expand All @@ -488,7 +482,7 @@ class FileStatus(str, Enum):


def read_file_text(file_path: Path) -> str:
"""Read a file's text content, replacing invalid characters and silencing read errors."""
"""Read a file's text content, replacing invalid UTF-8 characters."""
return file_path.read_text(encoding="utf-8", errors="replace")


Expand Down
Loading
Loading