286 lines
10 KiB
Python
286 lines
10 KiB
Python
"""Knowledge-base importer (PLAN §5 / §9 / §11).
|
|
|
|
Walks the A9-format files (``md, markdown, txt, yaml, yml, json, py`` by
|
|
default — ``BOR_IMPORT_EXTENSIONS``; case-insensitive), diffs by sha256
|
|
against ``documents.content_hash`` and, for every new or changed file, runs
|
|
the two-phase upsert:
|
|
|
|
1. upsert the document row and replace its chunk rows (embeddings NULL)
|
|
2. embed the new chunks in batches and attach the vectors
|
|
3. commit — one transaction per file, so a failed embedding leaves the
|
|
database untouched and the file is simply retried on the next run
|
|
|
|
Scope (A9, revised 2026-08-21): any path containing a dot-prefixed
|
|
component (hidden dirs — vendored caches like ``.esphome/.espressif/**`` —
|
|
or hidden files) is skipped, plus the well-known exclusion list.
|
|
|
|
``prune=True`` deletes documents (of the imported sources only) whose files
|
|
no longer exist **or no longer match the format filter** — this is how
|
|
previously-imported junk (e.g. dot-dir READMEs) leaves the index. Per-file
|
|
logging uses the verbs ``added | updated | unchanged | pruned`` plus a
|
|
summary line with per-format counts (PLAN §9).
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
import hashlib
|
|
import logging
|
|
from dataclasses import dataclass, field
|
|
from datetime import UTC, datetime
|
|
from pathlib import Path
|
|
from typing import Protocol
|
|
|
|
from sqlalchemy import select
|
|
from sqlalchemy.orm import Session
|
|
|
|
from app.config import Settings
|
|
from app.db import SessionLocal
|
|
from app.models import Chunk, Document
|
|
from app.rag.chunker import chunk_document, extract_title
|
|
from app.rag.llm import EmbeddingError
|
|
|
|
logger = logging.getLogger("app.importer")
|
|
|
|
#: Non-content directories never imported (PLAN anchor A9).
|
|
EXCLUDED_DIRS: frozenset[str] = frozenset(
|
|
{".venv", "node_modules", ".git", "__pycache__", ".pytest_cache", "dist", "build"}
|
|
)
|
|
|
|
|
|
class Embedder(Protocol):
|
|
"""Everything the importer needs from the LLM client (duck-typed for tests)."""
|
|
|
|
settings: Settings
|
|
embed_batches: int
|
|
|
|
async def embed(self, texts: list[str]) -> list[list[float]]: ...
|
|
|
|
|
|
@dataclass
|
|
class ImportSummary:
|
|
"""Counts for one import run (also printed by the CLI)."""
|
|
|
|
files: int = 0
|
|
added: int = 0
|
|
updated: int = 0
|
|
unchanged: int = 0
|
|
pruned: int = 0
|
|
errors: int = 0
|
|
chunks: int = 0
|
|
embed_batches: int = 0
|
|
#: Files walked, keyed by lowercased extension (``md``, ``yaml``, …).
|
|
formats: dict[str, int] = field(default_factory=dict)
|
|
|
|
def format_counts(self) -> str:
|
|
"""``md:203,yaml:267,py:14`` — highest count first (PLAN §9)."""
|
|
if not self.formats:
|
|
return "none"
|
|
ordered = sorted(self.formats.items(), key=lambda kv: (-kv[1], kv[0]))
|
|
return ",".join(f"{ext}:{count}" for ext, count in ordered)
|
|
|
|
def log(self) -> None:
|
|
logger.info(
|
|
"import: summary files=%d added=%d updated=%d unchanged=%d pruned=%d "
|
|
"errors=%d chunks=%d embed_batches=%d formats=%s",
|
|
self.files,
|
|
self.added,
|
|
self.updated,
|
|
self.unchanged,
|
|
self.pruned,
|
|
self.errors,
|
|
self.chunks,
|
|
self.embed_batches,
|
|
self.format_counts(),
|
|
)
|
|
|
|
|
|
def iter_importable_files(
|
|
root: Path,
|
|
extensions: frozenset[str],
|
|
excluded: frozenset[str] = EXCLUDED_DIRS,
|
|
) -> list[Path]:
|
|
"""All importable files under *root* (sorted), per the A9 scope rules.
|
|
|
|
*extensions* is a set of lowercased dotted suffixes (``{'.md', '.py'}``).
|
|
Skips: any path with a dot-prefixed component (hidden dirs/files —
|
|
vendored caches like ``.esphome/.espressif/**``) and the well-known
|
|
non-content directories in *excluded*.
|
|
"""
|
|
if not root.is_dir():
|
|
return []
|
|
files: list[Path] = []
|
|
for path in sorted(root.rglob("*")):
|
|
if not path.is_file():
|
|
continue
|
|
rel = path.relative_to(root)
|
|
if any(part.startswith(".") or part in excluded for part in rel.parts):
|
|
continue
|
|
if path.suffix.lower() not in extensions:
|
|
continue
|
|
files.append(path)
|
|
return files
|
|
|
|
|
|
async def import_sources(
|
|
sources: list[Path],
|
|
llm: Embedder,
|
|
*,
|
|
prune: bool = False,
|
|
limit: int | None = None,
|
|
session: Session | None = None,
|
|
) -> ImportSummary:
|
|
"""Import every A9-format file under *sources* (see module docstring).
|
|
|
|
``session`` may be supplied (tests); a private one is opened and closed
|
|
otherwise. ``limit`` caps the number of files processed (debug only) and
|
|
disables pruning, since an incomplete walk must not drive deletions.
|
|
"""
|
|
if limit is not None and limit <= 0:
|
|
raise ValueError("limit must be >= 1")
|
|
summary = ImportSummary()
|
|
owns_session = session is None
|
|
if session is None:
|
|
session = SessionLocal()
|
|
seen: set[tuple[str, str]] = set()
|
|
source_names: set[str] = set()
|
|
try:
|
|
for root in sources:
|
|
if not root.is_dir():
|
|
logger.warning("import: source dir not found, skipping: %s", root)
|
|
continue
|
|
if limit is not None and summary.files >= limit:
|
|
break
|
|
source = root.name
|
|
source_names.add(source)
|
|
for path in iter_importable_files(root, llm.settings.import_extension_set):
|
|
if limit is not None and summary.files >= limit:
|
|
break
|
|
rel = path.relative_to(root).as_posix()
|
|
seen.add((source, rel))
|
|
summary.files += 1
|
|
ext = path.suffix.lower().lstrip(".") or "unknown"
|
|
summary.formats[ext] = summary.formats.get(ext, 0) + 1
|
|
try:
|
|
await _index_file(
|
|
session, source=source, rel=rel, full_path=path, llm=llm,
|
|
summary=summary,
|
|
)
|
|
except EmbeddingError as e:
|
|
# A pathological file (e.g. content the embedding endpoint
|
|
# refuses) must not abort the whole KB: roll back its
|
|
# uncommitted rows, log loudly, and keep going. The next
|
|
# run retries it.
|
|
session.rollback()
|
|
summary.errors += 1
|
|
logger.error("import: error source=%s path=%s — %s", source, rel, e)
|
|
if prune:
|
|
if limit is not None:
|
|
logger.warning("import: --prune ignored because --limit was given")
|
|
else:
|
|
summary.pruned = _prune(session, source_names, seen)
|
|
summary.embed_batches = llm.embed_batches
|
|
summary.log()
|
|
return summary
|
|
finally:
|
|
if owns_session:
|
|
session.close()
|
|
|
|
|
|
async def _index_file(
|
|
session: Session,
|
|
*,
|
|
source: str,
|
|
rel: str,
|
|
full_path: Path,
|
|
llm: Embedder,
|
|
summary: ImportSummary,
|
|
) -> None:
|
|
"""Upsert one file: doc row + chunk rows + embeddings, one transaction."""
|
|
settings = llm.settings
|
|
content = full_path.read_text(encoding="utf-8", errors="replace")
|
|
digest = hashlib.sha256(content.encode("utf-8")).hexdigest()
|
|
doc = session.scalar(select(Document).where(Document.source == source, Document.path == rel))
|
|
if doc is not None and doc.content_hash == digest:
|
|
summary.unchanged += 1
|
|
logger.info("import: unchanged source=%s path=%s", source, rel)
|
|
return
|
|
|
|
verb = "updated" if doc is not None else "added"
|
|
# A ``#`` line is a real heading in markdown but a comment in every
|
|
# other format — titles for those come from the file stem.
|
|
if full_path.suffix.lower() in (".md", ".markdown"):
|
|
title = extract_title(content, fallback=full_path.stem)
|
|
else:
|
|
title = full_path.stem
|
|
if doc is None:
|
|
doc = Document(
|
|
source=source,
|
|
path=rel,
|
|
full_path=str(full_path),
|
|
title=title,
|
|
content=content,
|
|
content_hash=digest,
|
|
indexed_at=datetime.now(UTC),
|
|
)
|
|
session.add(doc)
|
|
else:
|
|
doc.full_path = str(full_path)
|
|
doc.title = title
|
|
doc.content = content
|
|
doc.content_hash = digest
|
|
doc.indexed_at = datetime.now(UTC)
|
|
|
|
session.flush() # guarantees doc.id even for brand-new rows
|
|
|
|
# Phase 1+2 — chunk, replace the chunk rows (embeddings NULL), embed,
|
|
# and commit the whole file atomically. If the endpoint rejects a chunk
|
|
# as over its input token cap (URL-dense paragraphs tokenize at ~1.1
|
|
# chars/token), halve this file's chunk target and retry — the global
|
|
# policy stays intact for the rest of the KB.
|
|
target = max(400, settings.chunk_target_chars)
|
|
while True:
|
|
chunks_text = chunk_document(content, rel, target, settings.chunk_overlap_chars)
|
|
doc.chunks = [
|
|
Chunk(document_id=doc.id, position=i, content=c) for i, c in enumerate(chunks_text)
|
|
]
|
|
session.flush() # delete-orphan cascade drops the previous rows
|
|
if not doc.chunks:
|
|
break
|
|
try:
|
|
vectors = await llm.embed([c.content for c in doc.chunks])
|
|
for row, vec in zip(doc.chunks, vectors, strict=True):
|
|
row.embedding = vec
|
|
break
|
|
except EmbeddingError as e:
|
|
if "token cap" not in str(e) or target <= 400:
|
|
raise
|
|
logger.info(
|
|
"import: re-chunking at %d chars after endpoint token cap: %s",
|
|
target // 2,
|
|
rel,
|
|
)
|
|
target //= 2
|
|
session.commit()
|
|
|
|
if verb == "added":
|
|
summary.added += 1
|
|
else:
|
|
summary.updated += 1
|
|
summary.chunks += len(chunks_text)
|
|
logger.info("import: %s source=%s path=%s chunks=%d", verb, source, rel, len(chunks_text))
|
|
|
|
|
|
def _prune(session: Session, source_names: set[str], seen: set[tuple[str, str]]) -> int:
|
|
"""Delete documents of *source_names* whose file is no longer in *seen*."""
|
|
if not source_names:
|
|
return 0
|
|
pruned = 0
|
|
docs = session.scalars(select(Document).where(Document.source.in_(source_names))).all()
|
|
for doc in docs:
|
|
if (doc.source, doc.path) not in seen:
|
|
session.delete(doc)
|
|
pruned += 1
|
|
logger.info("import: pruned source=%s path=%s", doc.source, doc.path)
|
|
if pruned:
|
|
session.commit()
|
|
return pruned
|