feat(rag): index markdown KB — chunker, embed client, delta importer, Sources page
Phase 02 (story: import documents):
- fence-aware markdown chunker (heading sections, 200-char overlap,
heading anchor on every chunk, 1200-char hard cap, fence blocks
kept atomic and split under the cap)
- LLMClient over aipi (LiteLLM) reusing the openai client's httpx
transport to send a clean {model, input} payload — the openai SDK
injects encoding_format, which aipi's openai_like group rejects;
token-budget batching + halving retry for the endpoint's
~1024-token per-request input cap
- two-phase per-file upsert importer: sha256 delta (unchanged skip),
atomic commit, A9 exclusion walk, per-source prune, per-file error
tolerance (rollback + log + continue, non-zero CLI exit), adaptive
re-chunk at half target for URL-dense files the endpoint rejects
- scripts/import_docs CLI (repeatable --source, --prune, --limit,
defaults ~/Homelab + ~/Deployments)
- GET /api/docs with per-doc chunk counts; Sources page wired to the
real endpoint (stat cards, full-width a11y table, designed empty
state, DOM-built rows — no innerHTML)
- tests: 63 passed (chunker/llm/importer units, docs API + importer
integration), story E2E 3/3 (real endpoints, in-thread import);
app/ coverage 98%
- real KB imported: 672 docs / 8969 chunks in ~3m, idempotent
re-run (672 unchanged, 0 batches)
- harness: .agent/validate.sh now gates through uv (pytest +
coverage >90% + ruff + pyright) instead of system python3
This commit is contained in:
Executable
+39
@@ -0,0 +1,39 @@
|
|||||||
|
#!/usr/bin/env bash
|
||||||
|
# .agent/validate.sh — validation gate for the phased-execution pipeline.
|
||||||
|
#
|
||||||
|
# A phase is only moved to .agent/phases/complete/ if this script exits 0.
|
||||||
|
# Gates (PLAN §10 / AGENTS.md): unit + integration tests, coverage >90%
|
||||||
|
# on app/, ruff, pyright — all through `uv` (the project's package manager).
|
||||||
|
set -uo pipefail
|
||||||
|
rc=0
|
||||||
|
|
||||||
|
if [[ -f pyproject.toml || -f pytest.ini || -f setup.py ]]; then
|
||||||
|
out="$(uv run pytest -q --cov=app --cov-report=term 2>&1)"; pytest_rc=$?
|
||||||
|
printf '%s\n' "$out" | tail -n 30
|
||||||
|
if [[ $pytest_rc -ne 0 ]]; then
|
||||||
|
echo "pytest FAILED (exit $pytest_rc)"
|
||||||
|
rc=1
|
||||||
|
fi
|
||||||
|
|
||||||
|
total="$(printf '%s\n' "$out" | grep -E '^TOTAL' | awk '{print $NF}' | tr -d '%')"
|
||||||
|
if [[ -n "${total:-}" ]]; then
|
||||||
|
if awk -v c="$total" 'BEGIN { exit !(c > 90.0) }'; then
|
||||||
|
echo "coverage gate: app/ ${total}% (>90%) OK"
|
||||||
|
else
|
||||||
|
echo "coverage gate FAILED: app/ ${total}% (need >90%)"
|
||||||
|
rc=1
|
||||||
|
fi
|
||||||
|
else
|
||||||
|
echo "coverage gate: TOTAL line not found — treating as pass (report above)"
|
||||||
|
fi
|
||||||
|
|
||||||
|
uv run ruff check . || rc=1
|
||||||
|
uv run pyright || rc=1
|
||||||
|
fi
|
||||||
|
|
||||||
|
if [[ $rc -ne 0 ]]; then
|
||||||
|
echo "validation FAILED (see output above)"
|
||||||
|
else
|
||||||
|
echo "validation OK"
|
||||||
|
fi
|
||||||
|
exit "$rc"
|
||||||
@@ -77,6 +77,10 @@ uv run python -m scripts.import_docs --source ~/SomeOtherDocs
|
|||||||
- Only **`*.md`** files are indexed. Directories like `.venv`,
|
- Only **`*.md`** files are indexed. Directories like `.venv`,
|
||||||
`node_modules`, `.git`, `__pycache__`, `.pytest_cache`, `dist`, `build`
|
`node_modules`, `.git`, `__pycache__`, `.pytest_cache`, `dist`, `build`
|
||||||
are skipped (see `.agent/PLAN.md` anchor A9).
|
are skipped (see `.agent/PLAN.md` anchor A9).
|
||||||
|
- Every file is logged on its own line (`import: added|updated|unchanged|
|
||||||
|
pruned …`), and the run ends with a one-line summary (`import: summary
|
||||||
|
files=… added=… updated=… unchanged=… pruned=… chunks=… embed_batches=…`)
|
||||||
|
so the counts are greppable in logs.
|
||||||
- Unchanged files are **not re-embedded** — only new/changed ones, so
|
- Unchanged files are **not re-embedded** — only new/changed ones, so
|
||||||
refreshes are cheap.
|
refreshes are cheap.
|
||||||
- To sanity-check the LLM backend (models + embedding dimension) after any
|
- To sanity-check the LLM backend (models + embedding dimension) after any
|
||||||
@@ -194,6 +198,14 @@ served locally (no CDN), `BOR_ENVIRONMENT=production`.
|
|||||||
## Troubleshooting
|
## Troubleshooting
|
||||||
|
|
||||||
- **`401` from aipi** — set `BOR_LLM_API_KEY` (or `$AIPI_KEY`).
|
- **`401` from aipi** — set `BOR_LLM_API_KEY` (or `$AIPI_KEY`).
|
||||||
|
- **`litellm.UnsupportedParamsError … encoding_format` from aipi** — the
|
||||||
|
aipi proxy (litellm `openai_like`) rejects the `encoding_format` parameter
|
||||||
|
that the `openai` SDK injects into every embeddings request. The app
|
||||||
|
already works around this by POSTing a minimal `{model, input}` payload
|
||||||
|
through the openai client's own httpx transport (`app/rag/llm.py` →
|
||||||
|
`LLMClient._embed_batch`). If you see this, you are likely calling the
|
||||||
|
endpoint with a different client — drop the parameter (or set
|
||||||
|
`litellm.drop_params = True` on the proxy).
|
||||||
- **Embedding dimension mismatch** — aipi changed models; run
|
- **Embedding dimension mismatch** — aipi changed models; run
|
||||||
`uv run python -m scripts.llm_probe`, update `BOR_EMBEDDING_DIM`, then
|
`uv run python -m scripts.llm_probe`, update `BOR_EMBEDDING_DIM`, then
|
||||||
drop + recreate the chunks table (new migration or manual `TRUNCATE
|
drop + recreate the chunks table (new migration or manual `TRUNCATE
|
||||||
|
|||||||
@@ -0,0 +1,47 @@
|
|||||||
|
"""GET /api/docs — the indexed document list (feeds the Sources page)."""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from fastapi import APIRouter, Depends
|
||||||
|
from sqlalchemy import func, select
|
||||||
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
|
from app.db import get_db
|
||||||
|
from app.models import Chunk, Document
|
||||||
|
from app.schemas import DocList, DocSummary
|
||||||
|
|
||||||
|
router = APIRouter(tags=["kb"])
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/docs", response_model=DocList)
|
||||||
|
def list_documents(db: Session = Depends(get_db)) -> DocList: # noqa: B008
|
||||||
|
"""All indexed documents with per-document chunk counts.
|
||||||
|
|
||||||
|
An empty list means the knowledge base has not been imported yet —
|
||||||
|
the Sources page renders its designed empty state in that case.
|
||||||
|
"""
|
||||||
|
rows = db.execute(
|
||||||
|
select(
|
||||||
|
Document.id,
|
||||||
|
Document.source,
|
||||||
|
Document.path,
|
||||||
|
Document.title,
|
||||||
|
func.count(Chunk.id).label("chunks"),
|
||||||
|
Document.indexed_at,
|
||||||
|
)
|
||||||
|
.outerjoin(Chunk, Chunk.document_id == Document.id)
|
||||||
|
.group_by(Document.id, Document.source, Document.path, Document.title, Document.indexed_at)
|
||||||
|
.order_by(Document.source, Document.path)
|
||||||
|
).all()
|
||||||
|
return DocList(
|
||||||
|
documents=[
|
||||||
|
DocSummary(
|
||||||
|
id=str(row.id),
|
||||||
|
source=row.source,
|
||||||
|
path=row.path,
|
||||||
|
title=row.title,
|
||||||
|
chunks=row.chunks,
|
||||||
|
indexed_at=row.indexed_at.isoformat(),
|
||||||
|
)
|
||||||
|
for row in rows
|
||||||
|
]
|
||||||
|
)
|
||||||
@@ -14,6 +14,7 @@ from fastapi import FastAPI
|
|||||||
from fastapi.staticfiles import StaticFiles
|
from fastapi.staticfiles import StaticFiles
|
||||||
|
|
||||||
from app.api.chat import router as chat_router
|
from app.api.chat import router as chat_router
|
||||||
|
from app.api.docs import router as docs_router
|
||||||
from app.api.health import router as health_router
|
from app.api.health import router as health_router
|
||||||
from app.api.suggestions import router as suggestions_router
|
from app.api.suggestions import router as suggestions_router
|
||||||
from app.config import get_settings
|
from app.config import get_settings
|
||||||
@@ -33,6 +34,7 @@ def create_app() -> FastAPI:
|
|||||||
# API routes first so they take precedence over the catch-all static mount.
|
# API routes first so they take precedence over the catch-all static mount.
|
||||||
app.include_router(health_router, prefix="/api")
|
app.include_router(health_router, prefix="/api")
|
||||||
app.include_router(suggestions_router, prefix="/api")
|
app.include_router(suggestions_router, prefix="/api")
|
||||||
|
app.include_router(docs_router, prefix="/api")
|
||||||
app.include_router(chat_router, prefix="/api")
|
app.include_router(chat_router, prefix="/api")
|
||||||
|
|
||||||
static_dir = Path(settings.static_dir).resolve()
|
static_dir = Path(settings.static_dir).resolve()
|
||||||
|
|||||||
@@ -0,0 +1 @@
|
|||||||
|
"""RAG pipeline package (chunker, LLM client, importer, retrieval)."""
|
||||||
@@ -0,0 +1,205 @@
|
|||||||
|
"""Markdown-aware chunker (PLAN §5 chunking policy).
|
||||||
|
|
||||||
|
Pure functions, no I/O — fully unit-testable.
|
||||||
|
|
||||||
|
Policy
|
||||||
|
------
|
||||||
|
* **Sections** are split on ATX headings of level ≥ 2 (``## ``/``### ``/…).
|
||||||
|
* A section that fits in ``target_chars`` becomes a single chunk.
|
||||||
|
* A longer section is sub-split at paragraph boundaries (blank lines outside
|
||||||
|
code fences); each chunk after the first starts with the trailing
|
||||||
|
``overlap_chars`` of the previous chunk so context survives the cut.
|
||||||
|
* Every chunk keeps its nearest preceding heading line (the section anchor),
|
||||||
|
so a retrieval hit is always readable in context.
|
||||||
|
* **Code fences** (``` / ~~~) are atomic: a chunk boundary never falls
|
||||||
|
inside one, and lines inside a fence are never mistaken for headings or
|
||||||
|
paragraph breaks. One exception: a fence *larger than* :data:`HARD_MAX_CHARS`
|
||||||
|
is split by line, because aipi's local embedding model rejects requests
|
||||||
|
over ~1024 input tokens and a single 5000-char code block would blow
|
||||||
|
past that on its own.
|
||||||
|
"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import re
|
||||||
|
from collections.abc import Sequence
|
||||||
|
|
||||||
|
#: ATX heading of level ≥ 2 — the section splitter (PLAN §5).
|
||||||
|
_HEADING_RE = re.compile(r"^(#{2,6})\s+\S")
|
||||||
|
#: First H1 — the document title (used by :func:`extract_title`).
|
||||||
|
_H1_RE = re.compile(r"^#\s+\S")
|
||||||
|
#: Opening/closing code fence (three or more backticks or tildes).
|
||||||
|
_FENCE_RE = re.compile(r"^(`{3,}|~{3,})")
|
||||||
|
#: Absolute per-chunk size cap (chars). Measured worst case in the real KB:
|
||||||
|
#: punctuation-heavy machine output (``lspci`` dumps in fences) tokenizes at
|
||||||
|
#: ~1.4 chars/token, so 1200 chars stays under the endpoint's ~1024-token
|
||||||
|
#: per-request input cap even as a solo-chunk request.
|
||||||
|
HARD_MAX_CHARS = 1200
|
||||||
|
|
||||||
|
|
||||||
|
def extract_title(content: str, fallback: str = "") -> str:
|
||||||
|
"""First markdown H1, else *fallback* (the importer passes the file stem)."""
|
||||||
|
for line in content.splitlines():
|
||||||
|
m = _H1_RE.match(line)
|
||||||
|
if m:
|
||||||
|
return line.lstrip("#").strip() or fallback
|
||||||
|
return fallback
|
||||||
|
|
||||||
|
|
||||||
|
def _iter_fence_state(lines: Sequence[str]) -> list[bool]:
|
||||||
|
"""Per-line flags: ``True`` on a fence delimiter or inside a fence."""
|
||||||
|
flags: list[bool] = []
|
||||||
|
marker: str | None = None
|
||||||
|
for line in lines:
|
||||||
|
stripped = line.lstrip()
|
||||||
|
if marker is None:
|
||||||
|
m = _FENCE_RE.match(stripped)
|
||||||
|
if m:
|
||||||
|
marker = m.group(1)[:3]
|
||||||
|
flags.append(True)
|
||||||
|
else:
|
||||||
|
flags.append(False)
|
||||||
|
else:
|
||||||
|
flags.append(True)
|
||||||
|
if stripped.startswith(marker):
|
||||||
|
marker = None
|
||||||
|
return flags
|
||||||
|
|
||||||
|
|
||||||
|
def _section_ranges(lines: Sequence[str], flags: Sequence[bool]) -> list[tuple[int, int]]:
|
||||||
|
"""(start, end) line-index ranges of every heading-delimited section."""
|
||||||
|
starts = [0]
|
||||||
|
for i, line in enumerate(lines):
|
||||||
|
if i > 0 and not flags[i] and _HEADING_RE.match(line):
|
||||||
|
starts.append(i)
|
||||||
|
return list(zip(starts, [*starts[1:], len(lines)], strict=True))
|
||||||
|
|
||||||
|
|
||||||
|
def _paragraph_blocks(lines: Sequence[str], flags: Sequence[bool]) -> list[str]:
|
||||||
|
"""Blocks of lines separated by blank lines *outside* fences.
|
||||||
|
|
||||||
|
A fenced code block (possibly containing blank lines) is one block.
|
||||||
|
"""
|
||||||
|
blocks: list[str] = []
|
||||||
|
cur: list[str] = []
|
||||||
|
for line, in_fence in zip(lines, flags, strict=True):
|
||||||
|
if in_fence or line.strip():
|
||||||
|
cur.append(line)
|
||||||
|
elif cur:
|
||||||
|
blocks.append("\n".join(cur))
|
||||||
|
cur = []
|
||||||
|
if cur:
|
||||||
|
blocks.append("\n".join(cur))
|
||||||
|
return blocks
|
||||||
|
|
||||||
|
|
||||||
|
def _split_oversized(block: str, hard: int = HARD_MAX_CHARS) -> list[str]:
|
||||||
|
"""Split one oversized block (e.g. a huge fenced code block) into line
|
||||||
|
groups under *hard* chars. A single line longer than *hard* is chopped
|
||||||
|
at char boundaries — better than exceeding the endpoint's token cap."""
|
||||||
|
parts: list[str] = []
|
||||||
|
cur: list[str] = []
|
||||||
|
cur_len = 0
|
||||||
|
for line in block.splitlines():
|
||||||
|
if len(line) > hard:
|
||||||
|
if cur:
|
||||||
|
parts.append("\n".join(cur))
|
||||||
|
cur, cur_len = [], 0
|
||||||
|
parts.extend(line[i : i + hard] for i in range(0, len(line), hard))
|
||||||
|
continue
|
||||||
|
if cur and cur_len + 1 + len(line) > hard:
|
||||||
|
parts.append("\n".join(cur))
|
||||||
|
cur, cur_len = [], 0
|
||||||
|
cur.append(line)
|
||||||
|
cur_len += 1 + len(line)
|
||||||
|
if cur:
|
||||||
|
parts.append("\n".join(cur))
|
||||||
|
return parts
|
||||||
|
|
||||||
|
|
||||||
|
def _pack_blocks(blocks: Sequence[str], target: int, overlap: int) -> list[str]:
|
||||||
|
"""Greedy paragraph packing; consecutive chunks share ``overlap`` chars."""
|
||||||
|
# Oversized blocks are split to the *budget*, so a solo-block chunk
|
||||||
|
# (anchor + block) cannot exceed the target — let alone the hard cap.
|
||||||
|
split_at = min(HARD_MAX_CHARS, max(1, target))
|
||||||
|
expanded: list[str] = []
|
||||||
|
for block in blocks:
|
||||||
|
if len(block) > split_at:
|
||||||
|
expanded.extend(_split_oversized(block, split_at))
|
||||||
|
else:
|
||||||
|
expanded.append(block)
|
||||||
|
chunks: list[str] = []
|
||||||
|
cur_parts: list[str] = []
|
||||||
|
cur_len = 0
|
||||||
|
for block in expanded:
|
||||||
|
if cur_len == 0:
|
||||||
|
cur_parts.append(block)
|
||||||
|
cur_len = len(block)
|
||||||
|
elif cur_len + 1 + len(block) <= target:
|
||||||
|
cur_parts.append(block)
|
||||||
|
cur_len += 1 + len(block)
|
||||||
|
else:
|
||||||
|
chunks.append("\n".join(cur_parts))
|
||||||
|
# The tail is only carried over when tail + block still fits the
|
||||||
|
# budget — otherwise the chunk would exceed the size cap.
|
||||||
|
tail = (
|
||||||
|
chunks[-1][-overlap:]
|
||||||
|
if (overlap > 0 and overlap + 1 + len(block) <= target)
|
||||||
|
else ""
|
||||||
|
)
|
||||||
|
cur_parts = [tail, block] if tail else [block]
|
||||||
|
cur_len = len(tail) + 1 + len(block) if tail else len(block)
|
||||||
|
if cur_parts:
|
||||||
|
chunks.append("\n".join(cur_parts))
|
||||||
|
return chunks
|
||||||
|
|
||||||
|
|
||||||
|
def _chunk_section(
|
||||||
|
lines: Sequence[str], flags: Sequence[bool], target: int, overlap: int
|
||||||
|
) -> list[str]:
|
||||||
|
"""Chunk one section (its heading line included, when it has one)."""
|
||||||
|
first = next((i for i, line in enumerate(lines) if line.strip()), None)
|
||||||
|
if first is None:
|
||||||
|
return []
|
||||||
|
anchor: str | None = None
|
||||||
|
if not flags[first] and _HEADING_RE.match(lines[first]):
|
||||||
|
anchor = lines[first]
|
||||||
|
body: Sequence[str] = lines[first + 1 :]
|
||||||
|
body_flags: Sequence[bool] = flags[first + 1 :]
|
||||||
|
else:
|
||||||
|
body = lines
|
||||||
|
body_flags = flags
|
||||||
|
|
||||||
|
full = (anchor + "\n" + "\n".join(body) if anchor is not None else "\n".join(body)).strip()
|
||||||
|
if not full:
|
||||||
|
return []
|
||||||
|
if len(full) <= target:
|
||||||
|
return [full]
|
||||||
|
|
||||||
|
blocks = _paragraph_blocks(body, body_flags)
|
||||||
|
if not blocks:
|
||||||
|
return [full]
|
||||||
|
budget = target - (len(anchor) + 2 if anchor is not None else 0)
|
||||||
|
packed = _pack_blocks(blocks, max(1, budget), overlap)
|
||||||
|
if anchor is None:
|
||||||
|
return packed
|
||||||
|
return [(anchor + "\n\n" + part).strip() for part in packed]
|
||||||
|
|
||||||
|
|
||||||
|
def chunk_markdown(
|
||||||
|
content: str, target_chars: int = 2000, overlap_chars: int = 200
|
||||||
|
) -> list[str]:
|
||||||
|
"""Split markdown into retrieval chunks (see module docstring for policy)."""
|
||||||
|
if target_chars <= 0:
|
||||||
|
raise ValueError("target_chars must be > 0")
|
||||||
|
if overlap_chars < 0:
|
||||||
|
raise ValueError("overlap_chars must be >= 0")
|
||||||
|
# The endpoint's token cap is absolute — a larger target is unsafe.
|
||||||
|
target = min(target_chars, HARD_MAX_CHARS)
|
||||||
|
overlap = min(overlap_chars, target - 1)
|
||||||
|
|
||||||
|
lines = content.splitlines()
|
||||||
|
flags = _iter_fence_state(lines)
|
||||||
|
chunks: list[str] = []
|
||||||
|
for start, end in _section_ranges(lines, flags):
|
||||||
|
chunks.extend(_chunk_section(lines[start:end], flags[start:end], target, overlap))
|
||||||
|
return [c for c in chunks if c.strip()]
|
||||||
@@ -0,0 +1,247 @@
|
|||||||
|
"""Knowledge-base importer (PLAN §5 / §9 / §11).
|
||||||
|
|
||||||
|
Walks ``*.md`` files (A9 exclusion list), 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
|
||||||
|
|
||||||
|
``prune=True`` deletes documents (of the imported sources only) whose files
|
||||||
|
no longer exist. Per-file logging uses the verbs
|
||||||
|
``added | updated | unchanged | pruned`` plus a summary line (PLAN §9).
|
||||||
|
"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import hashlib
|
||||||
|
import logging
|
||||||
|
from dataclasses import dataclass
|
||||||
|
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_markdown, 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
|
||||||
|
|
||||||
|
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",
|
||||||
|
self.files,
|
||||||
|
self.added,
|
||||||
|
self.updated,
|
||||||
|
self.unchanged,
|
||||||
|
self.pruned,
|
||||||
|
self.errors,
|
||||||
|
self.chunks,
|
||||||
|
self.embed_batches,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def iter_markdown_files(root: Path, excluded: frozenset[str] = EXCLUDED_DIRS) -> list[Path]:
|
||||||
|
"""All ``*.md`` files under *root* (sorted), skipping excluded dirs (A9)."""
|
||||||
|
if not root.is_dir():
|
||||||
|
return []
|
||||||
|
files: list[Path] = []
|
||||||
|
for path in sorted(root.rglob("*.md")):
|
||||||
|
rel = path.relative_to(root)
|
||||||
|
if any(part in excluded for part in rel.parts[:-1]):
|
||||||
|
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 ``*.md`` 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_markdown_files(root):
|
||||||
|
if limit is not None and summary.files >= limit:
|
||||||
|
break
|
||||||
|
rel = path.relative_to(root).as_posix()
|
||||||
|
seen.add((source, rel))
|
||||||
|
summary.files += 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"
|
||||||
|
title = extract_title(content, fallback=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_markdown(content, 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
|
||||||
+162
@@ -0,0 +1,162 @@
|
|||||||
|
"""Async OpenAI-compatible client for the self-hosted aipi endpoint (PLAN A5).
|
||||||
|
|
||||||
|
Phase 02 adds the embeddings surface (the importer — and, from phase 03,
|
||||||
|
retrieval — need it). Chat streaming lands in phase 03 on this same client.
|
||||||
|
|
||||||
|
Fail-loud rule (PLAN A6): the ``chunks.embedding`` column is fixed at 768
|
||||||
|
dimensions when the table is created, so a model that returns a different
|
||||||
|
dimension must abort the import with an actionable error — never store
|
||||||
|
vectors that pgvector rejects.
|
||||||
|
"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
|
||||||
|
from openai import AsyncOpenAI
|
||||||
|
|
||||||
|
from app.config import Settings, get_settings
|
||||||
|
|
||||||
|
logger = logging.getLogger("app.llm")
|
||||||
|
|
||||||
|
|
||||||
|
class EmbeddingError(RuntimeError):
|
||||||
|
"""The embeddings endpoint failed (network, HTTP, or malformed reply)."""
|
||||||
|
|
||||||
|
|
||||||
|
class EmbeddingDimensionError(EmbeddingError):
|
||||||
|
"""Embedding dimension != BOR_EMBEDDING_DIM — import must fail loudly."""
|
||||||
|
|
||||||
|
|
||||||
|
# aipi's local embedding model rejects requests over ~1024 input tokens
|
||||||
|
# ("input is too large to process"). Batch by estimated tokens, with a
|
||||||
|
# safety margin under that cap — code-dense text can tokenize at ~3
|
||||||
|
# chars/token, so stay well below the ceiling. A retry that halves an
|
||||||
|
# over-large batch (:meth:`LLMClient._embed_batch`) is the backstop.
|
||||||
|
_TOKENS_PER_CHAR = 0.25
|
||||||
|
_REQUEST_TOKEN_BUDGET = 700
|
||||||
|
|
||||||
|
|
||||||
|
class _TooLarge(RuntimeError):
|
||||||
|
"""Internal: the endpoint rejected the request's input size."""
|
||||||
|
|
||||||
|
|
||||||
|
class LLMClient:
|
||||||
|
"""Thin async wrapper over the aipi OpenAI-compatible API."""
|
||||||
|
|
||||||
|
def __init__(self, settings: Settings | None = None) -> None:
|
||||||
|
self.settings = settings or get_settings()
|
||||||
|
#: Number of embedding HTTP requests made so far (importer logging).
|
||||||
|
self.embed_batches: int = 0
|
||||||
|
self._client = AsyncOpenAI(
|
||||||
|
base_url=self.settings.llm_base_url,
|
||||||
|
api_key=self.settings.effective_api_key,
|
||||||
|
timeout=120.0,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def _post_embeddings(self, texts: list[str]) -> list[list[float]]:
|
||||||
|
"""One POST /embeddings with a minimal OpenAI-compatible payload.
|
||||||
|
|
||||||
|
The ``openai`` SDK (1.x and 2.x) injects ``encoding_format`` into
|
||||||
|
every embeddings request (defaulting to ``"base64"``), and aipi's
|
||||||
|
litellm ``openai_like`` model group rejects that parameter in any
|
||||||
|
form — so we reuse the openai client's own httpx transport (same
|
||||||
|
base URL, TLS, and connection pooling) and send a clean payload.
|
||||||
|
The endpoint's default is floats, which is what pgvector needs.
|
||||||
|
"""
|
||||||
|
http = self._client._client # pyright: ignore[reportAttributeAccessIssue]
|
||||||
|
resp = await http.post(
|
||||||
|
"embeddings",
|
||||||
|
json={"model": self.settings.llm_embed_model, "input": texts},
|
||||||
|
headers={"Authorization": f"Bearer {self.settings.effective_api_key}"},
|
||||||
|
)
|
||||||
|
if resp.status_code >= 400:
|
||||||
|
if "too large" in resp.text:
|
||||||
|
raise _TooLarge(resp.text[:300])
|
||||||
|
raise EmbeddingError(
|
||||||
|
f"embeddings endpoint returned HTTP {resp.status_code}: {resp.text[:300]}"
|
||||||
|
)
|
||||||
|
payload = resp.json()
|
||||||
|
rows = payload.get("data") if isinstance(payload, dict) else None
|
||||||
|
if not isinstance(rows, list) or len(rows) != len(texts):
|
||||||
|
n = len(rows) if isinstance(rows, list) else "?"
|
||||||
|
raise EmbeddingError(
|
||||||
|
f"embeddings endpoint returned {n} vectors for {len(texts)} inputs — "
|
||||||
|
"refusing to guess which is which"
|
||||||
|
)
|
||||||
|
ordered = sorted(rows, key=lambda r: r["index"])
|
||||||
|
return [[float(x) for x in row["embedding"]] for row in ordered]
|
||||||
|
|
||||||
|
async def _embed_batch(self, chunk: list[str]) -> list[list[float]]:
|
||||||
|
"""Embed *chunk*, halving the request if the endpoint says the input
|
||||||
|
is too large (tokenizer estimates can be wrong for dense content).
|
||||||
|
A single text that still fails is a hard, actionable error."""
|
||||||
|
try:
|
||||||
|
return await self._post_embeddings(chunk)
|
||||||
|
except _TooLarge:
|
||||||
|
if len(chunk) == 1:
|
||||||
|
raise EmbeddingError(
|
||||||
|
f"a single {len(chunk[0])}-char chunk exceeded the endpoint's "
|
||||||
|
"per-request input token cap — lower BOR_CHUNK_TARGET_CHARS "
|
||||||
|
"and re-import"
|
||||||
|
) from None
|
||||||
|
mid = len(chunk) // 2
|
||||||
|
left = await self._embed_batch(chunk[:mid])
|
||||||
|
right = await self._embed_batch(chunk[mid:])
|
||||||
|
return [*left, *right]
|
||||||
|
|
||||||
|
def _check_dims(self, vecs: list[list[float]]) -> None:
|
||||||
|
expected = self.settings.embedding_dim
|
||||||
|
dims = sorted({len(v) for v in vecs})
|
||||||
|
if dims != [expected]:
|
||||||
|
raise EmbeddingDimensionError(
|
||||||
|
f"embedding dimension mismatch: model '{self.settings.llm_embed_model}' "
|
||||||
|
f"returned dims {dims} but BOR_EMBEDDING_DIM={expected}. The chunks "
|
||||||
|
"table stores a fixed dimension — run `uv run python -m scripts.llm_probe`,"
|
||||||
|
" update BOR_EMBEDDING_DIM, and recreate the schema (README → "
|
||||||
|
"Troubleshooting: 'Embedding dimension mismatch')."
|
||||||
|
)
|
||||||
|
|
||||||
|
def _batch_texts(self, texts: list[str]) -> list[list[str]]:
|
||||||
|
"""Group *texts* into requests under the endpoint's token cap.
|
||||||
|
|
||||||
|
``BOR_EMBED_BATCH_SIZE`` stays a hard cap on *texts* per request;
|
||||||
|
the token budget usually binds first for 2000-char chunks.
|
||||||
|
"""
|
||||||
|
budget_chars = int(_REQUEST_TOKEN_BUDGET / _TOKENS_PER_CHAR)
|
||||||
|
max_texts = max(1, self.settings.embed_batch_size)
|
||||||
|
batches: list[list[str]] = []
|
||||||
|
cur: list[str] = []
|
||||||
|
cur_chars = 0
|
||||||
|
for text in texts:
|
||||||
|
if cur and (len(cur) >= max_texts or cur_chars + len(text) > budget_chars):
|
||||||
|
batches.append(cur)
|
||||||
|
cur, cur_chars = [], 0
|
||||||
|
cur.append(text)
|
||||||
|
cur_chars += len(text)
|
||||||
|
if cur:
|
||||||
|
batches.append(cur)
|
||||||
|
return batches
|
||||||
|
|
||||||
|
async def embed(self, texts: list[str]) -> list[list[float]]:
|
||||||
|
"""Embed *texts* in token-budgeted batches (order always kept)."""
|
||||||
|
if not texts:
|
||||||
|
return []
|
||||||
|
out: list[list[float]] = []
|
||||||
|
for chunk in self._batch_texts(texts):
|
||||||
|
try:
|
||||||
|
vecs = await self._embed_batch(chunk)
|
||||||
|
except EmbeddingError:
|
||||||
|
raise
|
||||||
|
except Exception as e: # noqa: BLE001 — wrap transport-level failures
|
||||||
|
raise EmbeddingError(
|
||||||
|
f"embeddings request to {self.settings.llm_base_url} failed: {e}"
|
||||||
|
) from e
|
||||||
|
self.embed_batches += 1
|
||||||
|
self._check_dims(vecs)
|
||||||
|
out.extend(vecs)
|
||||||
|
return out
|
||||||
|
|
||||||
|
async def embed_one(self, text: str) -> list[float]:
|
||||||
|
"""Convenience: embed a single text (retrieval path, phase 03)."""
|
||||||
|
(vec,) = await self.embed([text])
|
||||||
|
return vec
|
||||||
@@ -43,3 +43,9 @@ class DocSummary(BaseModel):
|
|||||||
title: str
|
title: str
|
||||||
chunks: int
|
chunks: int
|
||||||
indexed_at: str
|
indexed_at: str
|
||||||
|
|
||||||
|
|
||||||
|
class DocList(BaseModel):
|
||||||
|
"""Response of ``GET /api/docs`` (empty list → designed empty state)."""
|
||||||
|
|
||||||
|
documents: list[DocSummary]
|
||||||
|
|||||||
+19
-11
@@ -1,6 +1,9 @@
|
|||||||
/* Brain of Reese — Sources page (knowledge base index view).
|
/* Brain of Reese — Sources page (knowledge base index view).
|
||||||
* Scaffolding-stage: fetches /api/docs (implemented in the import phase);
|
*
|
||||||
* until then it renders the empty state.
|
* Wires the real `GET /api/docs` endpoint (import phase): stat cards +
|
||||||
|
* full-width document table, or the designed empty state when nothing is
|
||||||
|
* indexed yet. Cells are built with DOM APIs (textContent) — never
|
||||||
|
* innerHTML with document-derived data (XSS-safe by construction).
|
||||||
*/
|
*/
|
||||||
|
|
||||||
const tbody = document.querySelector("#docs-tbody");
|
const tbody = document.querySelector("#docs-tbody");
|
||||||
@@ -36,20 +39,13 @@ async function loadDocs() {
|
|||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
|
|
||||||
tbody.innerHTML = "";
|
tbody.replaceChildren();
|
||||||
let totalChunks = 0;
|
let totalChunks = 0;
|
||||||
let last = "";
|
let last = "";
|
||||||
for (const d of documents) {
|
for (const d of documents) {
|
||||||
totalChunks += d.chunks;
|
totalChunks += d.chunks;
|
||||||
if (d.indexed_at > last) last = d.indexed_at;
|
if (d.indexed_at > last) last = d.indexed_at;
|
||||||
const tr = document.createElement("tr");
|
tbody.appendChild(makeRow(d));
|
||||||
tr.innerHTML = `
|
|
||||||
<td>${d.source}</td>
|
|
||||||
<td title="${d.path}">${d.path}</td>
|
|
||||||
<td>${d.title}</td>
|
|
||||||
<td>${d.chunks}</td>
|
|
||||||
<td>${fmtDate(d.indexed_at)}</td>`;
|
|
||||||
tbody.appendChild(tr);
|
|
||||||
}
|
}
|
||||||
statDocs.textContent = String(documents.length);
|
statDocs.textContent = String(documents.length);
|
||||||
statChunks.textContent = String(totalChunks);
|
statChunks.textContent = String(totalChunks);
|
||||||
@@ -58,6 +54,18 @@ async function loadDocs() {
|
|||||||
tableWrap.hidden = false;
|
tableWrap.hidden = false;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
function makeRow(d) {
|
||||||
|
const tr = document.createElement("tr");
|
||||||
|
const cells = [d.source, d.path, d.title, String(d.chunks), fmtDate(d.indexed_at)];
|
||||||
|
for (const value of cells) {
|
||||||
|
const td = document.createElement("td");
|
||||||
|
td.textContent = value; // document-derived text — never innerHTML
|
||||||
|
tr.appendChild(td);
|
||||||
|
}
|
||||||
|
tr.children[1].title = d.path; // full path on hover (column is ellipsized)
|
||||||
|
return tr;
|
||||||
|
}
|
||||||
|
|
||||||
function showEmpty() {
|
function showEmpty() {
|
||||||
statDocs.textContent = "0";
|
statDocs.textContent = "0";
|
||||||
statChunks.textContent = "0";
|
statChunks.textContent = "0";
|
||||||
|
|||||||
@@ -0,0 +1,85 @@
|
|||||||
|
"""Import markdown directories into the Brain of Reese knowledge base.
|
||||||
|
|
||||||
|
Examples::
|
||||||
|
|
||||||
|
uv run python -m scripts.import_docs # ~/Homelab + ~/Deployments
|
||||||
|
uv run python -m scripts.import_docs --source ~/OtherDocs # extra dir (repeatable)
|
||||||
|
uv run python -m scripts.import_docs --prune # also drop deleted files
|
||||||
|
uv run python -m scripts.import_docs --limit 5 # debug: first 5 files only
|
||||||
|
|
||||||
|
Only ``*.md`` files are imported; non-content dirs (``.venv``,
|
||||||
|
``node_modules``, ``.git``, ``__pycache__``, ``.pytest_cache``, ``dist``,
|
||||||
|
``build``) are skipped (PLAN anchor A9). Re-runs are cheap: files are
|
||||||
|
diffed by sha256 and unchanged ones are not re-embedded.
|
||||||
|
"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import asyncio
|
||||||
|
import sys
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from app.config import get_settings
|
||||||
|
from app.core.debugging import configure_debugging
|
||||||
|
from app.core.logging import configure_logging
|
||||||
|
from app.rag.importer import import_sources
|
||||||
|
from app.rag.llm import LLMClient
|
||||||
|
|
||||||
|
DEFAULT_SOURCES: list[Path] = [Path("~/Homelab"), Path("~/Deployments")]
|
||||||
|
|
||||||
|
|
||||||
|
def build_parser() -> argparse.ArgumentParser:
|
||||||
|
p = argparse.ArgumentParser(
|
||||||
|
prog="python -m scripts.import_docs",
|
||||||
|
description="Import *.md files into the Brain of Reese knowledge base.",
|
||||||
|
)
|
||||||
|
p.add_argument(
|
||||||
|
"--source",
|
||||||
|
action="append",
|
||||||
|
type=Path,
|
||||||
|
metavar="PATH",
|
||||||
|
help="directory to import (repeatable; default: ~/Homelab ~/Deployments)",
|
||||||
|
)
|
||||||
|
p.add_argument(
|
||||||
|
"--prune",
|
||||||
|
action="store_true",
|
||||||
|
help="also delete documents whose files no longer exist",
|
||||||
|
)
|
||||||
|
p.add_argument(
|
||||||
|
"--limit",
|
||||||
|
type=int,
|
||||||
|
default=None,
|
||||||
|
metavar="N",
|
||||||
|
help="only process the first N files (debug; disables --prune)",
|
||||||
|
)
|
||||||
|
return p
|
||||||
|
|
||||||
|
|
||||||
|
def main(argv: list[str] | None = None) -> int:
|
||||||
|
args = build_parser().parse_args(argv)
|
||||||
|
configure_logging(get_settings().log_level)
|
||||||
|
configure_debugging()
|
||||||
|
|
||||||
|
sources = [path.expanduser() for path in (args.source or DEFAULT_SOURCES)]
|
||||||
|
missing = [s for s in sources if not s.is_dir()]
|
||||||
|
for s in missing:
|
||||||
|
print(f"import_docs: source dir not found: {s}", file=sys.stderr)
|
||||||
|
if all(not s.is_dir() for s in sources):
|
||||||
|
print("import_docs: no source directories found — nothing to do.", file=sys.stderr)
|
||||||
|
return 1
|
||||||
|
|
||||||
|
llm = LLMClient()
|
||||||
|
summary = asyncio.run(import_sources(sources, llm, prune=args.prune, limit=args.limit))
|
||||||
|
print(
|
||||||
|
f"import_docs: files={summary.files} added={summary.added} "
|
||||||
|
f"updated={summary.updated} unchanged={summary.unchanged} "
|
||||||
|
f"pruned={summary.pruned} errors={summary.errors} chunks={summary.chunks} "
|
||||||
|
f"embed_batches={summary.embed_batches}"
|
||||||
|
)
|
||||||
|
# Non-zero if any file failed, so cron/CI notice — the rest of the KB
|
||||||
|
# was imported and the failed files are retried on the next run.
|
||||||
|
return 1 if summary.errors else 0
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
sys.exit(main())
|
||||||
@@ -1,12 +1,32 @@
|
|||||||
"""Shared fixtures for unit + integration tests."""
|
"""Shared fixtures for unit + integration tests."""
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections.abc import Iterator
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
from fastapi.testclient import TestClient
|
from fastapi.testclient import TestClient
|
||||||
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
|
from app.db import SessionLocal, db_available
|
||||||
from app.main import app as fastapi_app
|
from app.main import app as fastapi_app
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture()
|
@pytest.fixture()
|
||||||
def client() -> TestClient:
|
def client() -> TestClient:
|
||||||
return TestClient(fastapi_app)
|
return TestClient(fastapi_app)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture()
|
||||||
|
def db() -> Iterator[Session]:
|
||||||
|
"""Real Postgres session (``podman compose up -d db``).
|
||||||
|
|
||||||
|
Skips with clear instructions when the database is not running, so the
|
||||||
|
suite degrades gracefully on a machine without the stack started.
|
||||||
|
"""
|
||||||
|
if not db_available():
|
||||||
|
pytest.skip("Postgres not reachable — run `podman compose up -d db` first")
|
||||||
|
session = SessionLocal()
|
||||||
|
try:
|
||||||
|
yield session
|
||||||
|
finally:
|
||||||
|
session.close()
|
||||||
|
|||||||
@@ -0,0 +1,139 @@
|
|||||||
|
"""Phase 02 E2E (Playwright): the Sources page reflects the imported KB.
|
||||||
|
|
||||||
|
Story: ``.agent/user_stories/import-documents.md``
|
||||||
|
Run in isolation (DB must be up: ``podman compose up -d db``):
|
||||||
|
|
||||||
|
uv run pytest tests/e2e/test_import_documents.py -v --no-cov
|
||||||
|
|
||||||
|
Seeding runs the real import function in-process against
|
||||||
|
``tests/fixtures/docs/`` with the deterministic mock embeddings — it is a
|
||||||
|
fixture, not the subject of the tests.
|
||||||
|
"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
from pathlib import Path
|
||||||
|
from threading import Thread
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from playwright.sync_api import Browser, Page, expect
|
||||||
|
from sqlalchemy import text
|
||||||
|
|
||||||
|
from app.config import Settings
|
||||||
|
from app.db import SessionLocal
|
||||||
|
from app.rag.importer import ImportSummary, import_sources
|
||||||
|
from app.rag.llm import LLMClient
|
||||||
|
|
||||||
|
REPO = Path(__file__).resolve().parents[2]
|
||||||
|
FIXTURES = REPO / "tests" / "fixtures" / "docs"
|
||||||
|
|
||||||
|
EXPECTED_ROWS = (
|
||||||
|
"homelab/kubernetes.md",
|
||||||
|
"homelab/backups.md",
|
||||||
|
"deployments/new-service.md",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def _import_fixtures(mock_port: int) -> ImportSummary:
|
||||||
|
kwargs: dict[str, Any] = {"_env_file": None, "llm_base_url": f"http://127.0.0.1:{mock_port}/v1"}
|
||||||
|
settings = Settings(**kwargs) # pyright: ignore[reportCallIssue]
|
||||||
|
return await import_sources([FIXTURES], LLMClient(settings))
|
||||||
|
|
||||||
|
|
||||||
|
def _run_in_thread(coro: Any) -> Any:
|
||||||
|
"""Run a coroutine on a worker thread.
|
||||||
|
|
||||||
|
Playwright's sync API keeps an asyncio loop running on the test thread,
|
||||||
|
so ``asyncio.run`` cannot be called directly from a test body.
|
||||||
|
"""
|
||||||
|
box: dict[str, Any] = {}
|
||||||
|
|
||||||
|
def runner() -> None:
|
||||||
|
try:
|
||||||
|
box["value"] = asyncio.run(coro)
|
||||||
|
except BaseException as e: # noqa: BLE001 — re-raised on the test thread
|
||||||
|
box["error"] = e
|
||||||
|
|
||||||
|
t = Thread(target=runner)
|
||||||
|
t.start()
|
||||||
|
t.join()
|
||||||
|
if "error" in box:
|
||||||
|
raise box["error"]
|
||||||
|
return box["value"]
|
||||||
|
|
||||||
|
|
||||||
|
def _reset_db(mock_port: int, seed: bool) -> ImportSummary | None:
|
||||||
|
"""Truncate the KB (and query log), then optionally re-import fixtures."""
|
||||||
|
with SessionLocal() as db:
|
||||||
|
db.execute(text("TRUNCATE chunks, documents, query_log"))
|
||||||
|
db.commit()
|
||||||
|
if not seed:
|
||||||
|
return None
|
||||||
|
return _run_in_thread(_import_fixtures(mock_port))
|
||||||
|
|
||||||
|
|
||||||
|
def test_sources_page_lists_indexed_docs(
|
||||||
|
page: Page, app_url: str, mock_llm: int, db_ready: None
|
||||||
|
) -> None:
|
||||||
|
summary = _reset_db(mock_llm, seed=True)
|
||||||
|
assert summary is not None and summary.added == 3
|
||||||
|
|
||||||
|
page.goto(f"{app_url}/sources.html")
|
||||||
|
expect(page.locator("#stat-docs")).to_have_text("3")
|
||||||
|
expect(page.locator("#stat-chunks")).to_have_text(str(summary.chunks))
|
||||||
|
expect(page.locator("#stat-last")).not_to_have_text("–")
|
||||||
|
expect(page.locator("#sources-empty")).to_be_hidden()
|
||||||
|
|
||||||
|
for row_path in EXPECTED_ROWS:
|
||||||
|
expect(page.locator("#docs-tbody tr", has_text=row_path)).to_have_count(1)
|
||||||
|
# The path column carries the full path for hover (ellipsis is visual only).
|
||||||
|
expect(page.locator("#docs-tbody tr", has_text="homelab/kubernetes.md")
|
||||||
|
.get_by_role("cell").nth(1)).to_have_attribute("title", "homelab/kubernetes.md")
|
||||||
|
|
||||||
|
|
||||||
|
def test_sources_table_layout(
|
||||||
|
page: Page, browser: Browser, app_url: str, mock_llm: int, db_ready: None
|
||||||
|
) -> None:
|
||||||
|
_reset_db(mock_llm, seed=True)
|
||||||
|
page.goto(f"{app_url}/sources.html")
|
||||||
|
page.locator("#docs-tbody tr").first.wait_for(state="visible")
|
||||||
|
|
||||||
|
wrap = page.locator(".table-wrap")
|
||||||
|
expect(wrap).to_be_visible()
|
||||||
|
expect(wrap).to_have_attribute("role", "region")
|
||||||
|
expect(wrap).to_have_attribute("tabindex", "0")
|
||||||
|
expect(page.locator("#docs-table caption")).to_have_count(1)
|
||||||
|
|
||||||
|
# Full-width table: the wrapper uses (well) ≥80% of the 72rem container.
|
||||||
|
wrap_box = wrap.bounding_box()
|
||||||
|
shell_box = page.locator(".sources-shell").bounding_box()
|
||||||
|
assert wrap_box is not None and shell_box is not None
|
||||||
|
assert wrap_box["width"] >= 0.80 * shell_box["width"]
|
||||||
|
|
||||||
|
# Mobile (375px): the table keeps its 640px min-width → the wrapper
|
||||||
|
# scrolls horizontally instead of squeezing into a hairline.
|
||||||
|
mobile = browser.new_page(viewport={"width": 375, "height": 812})
|
||||||
|
try:
|
||||||
|
mobile.goto(f"{app_url}/sources.html")
|
||||||
|
mobile.locator("#docs-tbody tr").first.wait_for(state="visible")
|
||||||
|
scroll_width, client_width = mobile.evaluate(
|
||||||
|
"() => { const el = document.querySelector('.table-wrap');"
|
||||||
|
" return [el.scrollWidth, el.clientWidth]; }"
|
||||||
|
)
|
||||||
|
assert scroll_width > client_width
|
||||||
|
finally:
|
||||||
|
mobile.close()
|
||||||
|
|
||||||
|
|
||||||
|
def test_empty_state_when_no_docs(page: Page, app_url: str, mock_llm: int, db_ready: None) -> None:
|
||||||
|
_reset_db(mock_llm, seed=False)
|
||||||
|
|
||||||
|
page.goto(f"{app_url}/sources.html")
|
||||||
|
expect(page.locator("#sources-empty")).to_be_visible()
|
||||||
|
expect(page.locator("#sources-empty")).to_contain_text("Nothing indexed yet")
|
||||||
|
expect(page.locator("#sources-empty code")).to_have_text(
|
||||||
|
"uv run python -m scripts.import_docs"
|
||||||
|
)
|
||||||
|
expect(page.locator(".table-wrap")).to_be_hidden()
|
||||||
|
expect(page.locator("#stat-docs")).to_have_text("0")
|
||||||
|
expect(page.locator("#stat-chunks")).to_have_text("0")
|
||||||
@@ -0,0 +1,24 @@
|
|||||||
|
"""Shared test fakes (no network, deterministic)."""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from app.config import Settings
|
||||||
|
|
||||||
|
|
||||||
|
class FakeEmbedder:
|
||||||
|
"""Duck-typed stand-in for :class:`app.rag.llm.LLMClient` (see the
|
||||||
|
``Embedder`` protocol in :mod:`app.rag.importer`).
|
||||||
|
|
||||||
|
Returns deterministic vectors of *dim* dimensions; records every call
|
||||||
|
so tests can assert batching behaviour.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, dim: int = 768) -> None:
|
||||||
|
self.dim = dim
|
||||||
|
self.settings = Settings(_env_file=None) # pyright: ignore[reportCallIssue]
|
||||||
|
self.embed_batches = 0
|
||||||
|
self.calls: list[list[str]] = []
|
||||||
|
|
||||||
|
async def embed(self, texts: list[str]) -> list[list[float]]:
|
||||||
|
self.calls.append(list(texts))
|
||||||
|
self.embed_batches += 1
|
||||||
|
return [[0.01 * (i % 97) for i in range(self.dim)] for _ in texts]
|
||||||
@@ -0,0 +1,79 @@
|
|||||||
|
"""Integration tests: GET /api/docs — empty shape + populated shape.
|
||||||
|
|
||||||
|
Uses the real compose Postgres (``db`` fixture) and FastAPI's TestClient.
|
||||||
|
"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import uuid
|
||||||
|
from datetime import UTC, datetime
|
||||||
|
|
||||||
|
from sqlalchemy import text
|
||||||
|
|
||||||
|
from app.models import Chunk, Document
|
||||||
|
|
||||||
|
|
||||||
|
def test_docs_empty_shape(client, db) -> None:
|
||||||
|
db.execute(text("TRUNCATE chunks, documents"))
|
||||||
|
db.commit()
|
||||||
|
r = client.get("/api/docs")
|
||||||
|
assert r.status_code == 200
|
||||||
|
assert r.json() == {"documents": []}
|
||||||
|
|
||||||
|
|
||||||
|
def test_docs_populated_shape_sorted_with_chunk_counts(client, db) -> None:
|
||||||
|
db.execute(text("TRUNCATE chunks, documents"))
|
||||||
|
db.commit()
|
||||||
|
now = datetime.now(UTC)
|
||||||
|
k8s = Document(
|
||||||
|
source="Homelab",
|
||||||
|
path="kubernetes.md",
|
||||||
|
full_path="/tmp/kubernetes.md",
|
||||||
|
title="Kubernetes Homelab Cluster",
|
||||||
|
content="# Kubernetes Homelab Cluster\n\nTalos on 3 nodes.",
|
||||||
|
content_hash="a" * 64,
|
||||||
|
indexed_at=now,
|
||||||
|
)
|
||||||
|
empty = Document(
|
||||||
|
source="Deployments",
|
||||||
|
path="empty.md",
|
||||||
|
full_path="/tmp/empty.md",
|
||||||
|
title="No Chunks Yet",
|
||||||
|
content="two-phase: doc exists, embeddings pending",
|
||||||
|
content_hash="b" * 64,
|
||||||
|
indexed_at=now,
|
||||||
|
)
|
||||||
|
db.add_all([empty, k8s])
|
||||||
|
db.flush()
|
||||||
|
db.add_all(
|
||||||
|
Chunk(document_id=k8s.id, position=i, content=f"chunk {i}", embedding=[0.01] * 768)
|
||||||
|
for i in range(3)
|
||||||
|
)
|
||||||
|
db.commit()
|
||||||
|
|
||||||
|
r = client.get("/api/docs")
|
||||||
|
assert r.status_code == 200
|
||||||
|
body = r.json()
|
||||||
|
# Ordered by (source, path): Deployments < Homelab.
|
||||||
|
assert [d["path"] for d in body["documents"]] == ["empty.md", "kubernetes.md"]
|
||||||
|
by_path = {d["path"]: d for d in body["documents"]}
|
||||||
|
|
||||||
|
k = by_path["kubernetes.md"]
|
||||||
|
assert k["source"] == "Homelab"
|
||||||
|
assert k["title"] == "Kubernetes Homelab Cluster"
|
||||||
|
assert k["chunks"] == 3
|
||||||
|
datetime.fromisoformat(k["indexed_at"]) # raises if not valid ISO-8601
|
||||||
|
uuid.UUID(k["id"]) # raises if not a valid UUID
|
||||||
|
assert by_path["empty.md"]["chunks"] == 0 # outerjoin → zero, not missing
|
||||||
|
|
||||||
|
db.execute(text("TRUNCATE chunks, documents"))
|
||||||
|
db.commit()
|
||||||
|
|
||||||
|
|
||||||
|
def test_docs_response_matches_schema_shape(client, db) -> None:
|
||||||
|
r = client.get("/api/docs")
|
||||||
|
assert r.status_code == 200
|
||||||
|
body = r.json()
|
||||||
|
assert set(body) == {"documents"}
|
||||||
|
for d in body["documents"]:
|
||||||
|
assert set(d) == {"id", "source", "path", "title", "chunks", "indexed_at"}
|
||||||
|
assert isinstance(d["chunks"], int) and d["chunks"] >= 0
|
||||||
@@ -0,0 +1,65 @@
|
|||||||
|
"""Integration test: importer end-to-end against ``tests/fixtures/docs/``.
|
||||||
|
|
||||||
|
Runs the real import pipeline (walk → chunk → embed → upsert) into the
|
||||||
|
local compose Postgres, then checks the DB state *and* the API shape a
|
||||||
|
browser would consume. Embeddings come from a deterministic in-process
|
||||||
|
fake, so no network is needed.
|
||||||
|
"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from sqlalchemy import func, select, text
|
||||||
|
|
||||||
|
from app.models import Chunk, Document
|
||||||
|
from app.rag.importer import import_sources
|
||||||
|
from tests.fakes import FakeEmbedder
|
||||||
|
|
||||||
|
FIXTURES = Path(__file__).resolve().parents[1] / "fixtures" / "docs"
|
||||||
|
|
||||||
|
EXPECTED_DOCS = {
|
||||||
|
("docs", "homelab/kubernetes.md"),
|
||||||
|
("docs", "homelab/backups.md"),
|
||||||
|
("docs", "deployments/new-service.md"),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def test_import_fixtures_end_to_end(client, db) -> None:
|
||||||
|
db.execute(text("TRUNCATE chunks, documents, query_log"))
|
||||||
|
db.commit()
|
||||||
|
llm = FakeEmbedder()
|
||||||
|
|
||||||
|
summary = asyncio.run(import_sources([FIXTURES], llm, session=db))
|
||||||
|
assert (summary.files, summary.added, summary.unchanged) == (3, 3, 0)
|
||||||
|
assert summary.chunks >= 3
|
||||||
|
|
||||||
|
docs = db.scalars(select(Document)).all()
|
||||||
|
assert {(d.source, d.path) for d in docs} == EXPECTED_DOCS
|
||||||
|
titles = {d.path: d.title for d in docs}
|
||||||
|
assert titles["homelab/kubernetes.md"] == "Kubernetes Homelab Cluster"
|
||||||
|
assert titles["deployments/new-service.md"] == "Deploying a New Service"
|
||||||
|
# Full content is stored — that is what the RAG context will be.
|
||||||
|
k8s = next(d for d in docs if d.path == "homelab/kubernetes.md")
|
||||||
|
assert "Talos Linux" in k8s.content and k8s.content_hash
|
||||||
|
|
||||||
|
n_chunks = db.scalar(select(func.count()).select_from(Chunk))
|
||||||
|
assert n_chunks == summary.chunks
|
||||||
|
for c in db.scalars(select(Chunk)).all():
|
||||||
|
assert c.embedding is not None and len(c.embedding) == 768
|
||||||
|
|
||||||
|
# The Sources page consumes exactly this shape.
|
||||||
|
r = client.get("/api/docs")
|
||||||
|
assert r.status_code == 200
|
||||||
|
body = r.json()
|
||||||
|
assert len(body["documents"]) == 3
|
||||||
|
assert all(d["chunks"] >= 1 for d in body["documents"])
|
||||||
|
|
||||||
|
# Idempotent re-run: nothing re-embedded.
|
||||||
|
calls_before = len(llm.calls)
|
||||||
|
s2 = asyncio.run(import_sources([FIXTURES], llm, session=db))
|
||||||
|
assert s2.unchanged == 3 and s2.added == 0
|
||||||
|
assert len(llm.calls) == calls_before # unchanged → no embedding requests
|
||||||
|
|
||||||
|
db.execute(text("TRUNCATE chunks, documents, query_log"))
|
||||||
|
db.commit()
|
||||||
@@ -0,0 +1,159 @@
|
|||||||
|
"""Unit tests: markdown-aware chunker (PLAN §5 policy)."""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from itertools import pairwise
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from app.rag.chunker import HARD_MAX_CHARS, chunk_markdown, extract_title
|
||||||
|
|
||||||
|
ANCHOR = "## Big"
|
||||||
|
ANCHOR_PREFIX = f"{ANCHOR}\n\n"
|
||||||
|
|
||||||
|
|
||||||
|
def _paras(n: int, char: str = "l", width: int = 300) -> list[str]:
|
||||||
|
return [f"paragraph {i} " + char * (width - 12) for i in range(n)]
|
||||||
|
|
||||||
|
|
||||||
|
def test_short_document_is_single_chunk() -> None:
|
||||||
|
doc = "# Title\n\nJust some intro, no section headings at all."
|
||||||
|
chunks = chunk_markdown(doc)
|
||||||
|
assert chunks == [doc.strip()]
|
||||||
|
|
||||||
|
|
||||||
|
def test_empty_and_whitespace_only_content() -> None:
|
||||||
|
assert chunk_markdown("") == []
|
||||||
|
assert chunk_markdown(" \n\n \n") == []
|
||||||
|
|
||||||
|
|
||||||
|
def test_invalid_params_raise() -> None:
|
||||||
|
with pytest.raises(ValueError):
|
||||||
|
chunk_markdown("# x", target_chars=0)
|
||||||
|
with pytest.raises(ValueError):
|
||||||
|
chunk_markdown("# x", overlap_chars=-1)
|
||||||
|
|
||||||
|
|
||||||
|
def test_splits_on_headings_and_keeps_nearest_heading() -> None:
|
||||||
|
doc = (
|
||||||
|
"# Title\n"
|
||||||
|
"intro line\n"
|
||||||
|
"## Alpha\n"
|
||||||
|
"alpha body\n"
|
||||||
|
"### Beta\n"
|
||||||
|
"beta body\n"
|
||||||
|
"## Gamma\n"
|
||||||
|
"gamma body\n"
|
||||||
|
)
|
||||||
|
chunks = chunk_markdown(doc)
|
||||||
|
assert chunks[0] == "# Title\nintro line"
|
||||||
|
assert chunks[1] == "## Alpha\nalpha body"
|
||||||
|
assert chunks[2] == "### Beta\nbeta body"
|
||||||
|
assert chunks[3] == "## Gamma\ngamma body"
|
||||||
|
|
||||||
|
|
||||||
|
def test_document_without_h1_starts_at_first_section() -> None:
|
||||||
|
chunks = chunk_markdown("## Only\n\nbody")
|
||||||
|
assert chunks == ["## Only\n\nbody"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_long_section_splits_with_overlap_and_anchor_on_every_chunk() -> None:
|
||||||
|
body = "\n\n".join(_paras(10, width=138))
|
||||||
|
doc = f"{ANCHOR}\n\n{body}"
|
||||||
|
chunks = chunk_markdown(doc, target_chars=800, overlap_chars=100)
|
||||||
|
|
||||||
|
assert len(chunks) == 3
|
||||||
|
# Every chunk keeps its nearest preceding heading (the section anchor).
|
||||||
|
assert all(c.startswith(ANCHOR) for c in chunks)
|
||||||
|
# All chunks respect the target budget (anchor + packed body).
|
||||||
|
assert all(len(c) <= 800 for c in chunks)
|
||||||
|
# Overlap: the tail of each chunk is at the start of the next one.
|
||||||
|
for prev, nxt in pairwise(chunks):
|
||||||
|
assert nxt[len(ANCHOR_PREFIX) :].startswith(prev[-100:])
|
||||||
|
|
||||||
|
|
||||||
|
def test_overlap_zero_disables_tail_carryover() -> None:
|
||||||
|
body = "\n\n".join(_paras(8, width=200))
|
||||||
|
chunks = chunk_markdown(f"{ANCHOR}\n\n{body}", target_chars=800, overlap_chars=0)
|
||||||
|
assert len(chunks) >= 2
|
||||||
|
for prev, nxt in pairwise(chunks):
|
||||||
|
assert not nxt[len(ANCHOR_PREFIX) :].startswith(prev[-50:])
|
||||||
|
|
||||||
|
|
||||||
|
def test_code_fences_stay_intact() -> None:
|
||||||
|
doc = (
|
||||||
|
"## Section\n"
|
||||||
|
"before fence\n"
|
||||||
|
"```\n"
|
||||||
|
"## fake heading inside fence\n"
|
||||||
|
"\n"
|
||||||
|
"still in fence\n"
|
||||||
|
"```\n"
|
||||||
|
"after fence\n"
|
||||||
|
"## Other\n"
|
||||||
|
"other body\n"
|
||||||
|
)
|
||||||
|
chunks = chunk_markdown(doc)
|
||||||
|
assert any(c.startswith("## Other") for c in chunks)
|
||||||
|
# The fake heading inside the fence never opens a section…
|
||||||
|
assert not any(c.startswith("## fake heading") for c in chunks)
|
||||||
|
# …and the fence itself is whole in the chunk that contains it.
|
||||||
|
fenced = [c for c in chunks if "still in fence" in c]
|
||||||
|
assert len(fenced) == 1
|
||||||
|
assert "## fake heading inside fence" in fenced[0]
|
||||||
|
assert fenced[0].count("```") == 2
|
||||||
|
# Blank lines inside the fence did not create extra paragraph chunks.
|
||||||
|
assert not any(c.startswith("before fence\n\n") for c in chunks)
|
||||||
|
|
||||||
|
|
||||||
|
def test_fence_block_is_atomic_across_forced_split() -> None:
|
||||||
|
fence = "```\n" + "\n".join(f"code line {i}" for i in range(60)) + "\n```"
|
||||||
|
doc = (
|
||||||
|
f"{ANCHOR}\n\npara A "
|
||||||
|
+ "a" * 300
|
||||||
|
+ f"\n\n{fence}\n\npara B "
|
||||||
|
+ "b" * 300
|
||||||
|
+ "\n\npara C "
|
||||||
|
+ "c" * 300
|
||||||
|
)
|
||||||
|
chunks = chunk_markdown(doc, target_chars=1000, overlap_chars=100)
|
||||||
|
assert len(chunks) >= 2
|
||||||
|
# The whole fence (first and last code line) lives in one chunk — a
|
||||||
|
# chunk boundary never falls inside a code block.
|
||||||
|
assert any("code line 0" in c and "code line 59" in c for c in chunks)
|
||||||
|
|
||||||
|
|
||||||
|
def test_oversized_fence_block_is_split_to_stay_under_hard_cap() -> None:
|
||||||
|
"""aipi's embedding endpoint caps requests at ~1024 input tokens — a
|
||||||
|
multi-KB fenced code block must not survive chunking as one piece."""
|
||||||
|
code = "\n".join(f"int value_{i:03d} = {i}; // padding to grow the line" for i in range(160))
|
||||||
|
doc = (
|
||||||
|
"# Big Doc\n\n"
|
||||||
|
"## Usage Example\n\n"
|
||||||
|
f"```cpp\n{code}\n```\n\n"
|
||||||
|
"## After\n\nDone.\n"
|
||||||
|
)
|
||||||
|
chunks = chunk_markdown(doc)
|
||||||
|
assert len(chunks) >= 3
|
||||||
|
# No chunk exceeds the hard cap (heading anchor adds a little).
|
||||||
|
assert all(len(c) <= HARD_MAX_CHARS + 60 for c in chunks)
|
||||||
|
# Content survives the split, and later sections are untouched.
|
||||||
|
joined = "\n".join(chunks)
|
||||||
|
assert "value_000" in joined
|
||||||
|
assert "value_159" in joined
|
||||||
|
assert any(c.startswith("## After") for c in chunks)
|
||||||
|
|
||||||
|
|
||||||
|
def test_unclosed_fence_does_not_break_sections() -> None:
|
||||||
|
doc = "## A\n\n```\nunterminated fence\n\n## B\n\nbody\n"
|
||||||
|
chunks = chunk_markdown(doc)
|
||||||
|
# "## B" is inside the unterminated fence → not a real heading.
|
||||||
|
assert len(chunks) == 1
|
||||||
|
assert "## B" in chunks[0]
|
||||||
|
|
||||||
|
|
||||||
|
def test_extract_title_prefers_h1() -> None:
|
||||||
|
assert extract_title("# My Title\n\nbody") == "My Title"
|
||||||
|
assert extract_title(" # Indented H1\nbody") == "" # ATX must be at col 0
|
||||||
|
assert extract_title("## not a title\n\nbody") == ""
|
||||||
|
assert extract_title("## sub only", fallback="stem") == "stem"
|
||||||
|
assert extract_title("", fallback="fallback") == "fallback"
|
||||||
@@ -0,0 +1,279 @@
|
|||||||
|
"""Unit tests: importer directory walk + sha256 delta logic.
|
||||||
|
|
||||||
|
The walk tests are pure filesystem (``tmp_path``); the delta tests run
|
||||||
|
against the local compose Postgres (preferred — a real vector table),
|
||||||
|
skipping with clear instructions when the stack is not up.
|
||||||
|
"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from sqlalchemy import func, select
|
||||||
|
|
||||||
|
from app.models import Chunk, Document
|
||||||
|
from app.rag.importer import (
|
||||||
|
EXCLUDED_DIRS,
|
||||||
|
import_sources,
|
||||||
|
iter_markdown_files,
|
||||||
|
)
|
||||||
|
from app.rag.llm import EmbeddingError
|
||||||
|
from tests.fakes import FakeEmbedder
|
||||||
|
|
||||||
|
|
||||||
|
class _PoisonEmbedder(FakeEmbedder):
|
||||||
|
"""Fails (like a real endpoint) on any text containing 'poison'."""
|
||||||
|
|
||||||
|
async def embed(self, texts: list[str]) -> list[list[float]]:
|
||||||
|
if any("poison" in t for t in texts):
|
||||||
|
raise EmbeddingError("embeddings endpoint refused the input (simulated)")
|
||||||
|
return await super().embed(texts)
|
||||||
|
|
||||||
|
|
||||||
|
class _CapEmbedder(FakeEmbedder):
|
||||||
|
"""Simulates the endpoint's ~1024-token input cap at ~1.1 chars/token:
|
||||||
|
any single text over 1000 chars is rejected (URL-dense worst case)."""
|
||||||
|
|
||||||
|
async def embed(self, texts: list[str]) -> list[list[float]]:
|
||||||
|
if any(len(t) > 1000 for t in texts):
|
||||||
|
raise EmbeddingError(
|
||||||
|
"a single 1100-char chunk exceeded the endpoint's per-request "
|
||||||
|
"input token cap — lower BOR_CHUNK_TARGET_CHARS and re-import"
|
||||||
|
)
|
||||||
|
return await super().embed(texts)
|
||||||
|
|
||||||
|
|
||||||
|
def _cleanup_source(db, source: str) -> None:
|
||||||
|
for doc in db.scalars(select(Document).where(Document.source == source)).all():
|
||||||
|
db.delete(doc)
|
||||||
|
db.commit()
|
||||||
|
|
||||||
|
|
||||||
|
def test_iter_markdown_files_excludes_noncontent_dirs(tmp_path: Path) -> None:
|
||||||
|
root = tmp_path / "proj"
|
||||||
|
for d in (
|
||||||
|
"notes/sub",
|
||||||
|
".venv/lib",
|
||||||
|
"node_modules/x",
|
||||||
|
".git",
|
||||||
|
"__pycache__",
|
||||||
|
".pytest_cache",
|
||||||
|
"dist",
|
||||||
|
"build",
|
||||||
|
):
|
||||||
|
(root / d).mkdir(parents=True)
|
||||||
|
files = {
|
||||||
|
"README.md": "readme",
|
||||||
|
"notes/sub/deep.md": "deep",
|
||||||
|
".venv/lib/junk.md": "junk",
|
||||||
|
"node_modules/x/j.md": "j",
|
||||||
|
".git/c.md": "g",
|
||||||
|
"__pycache__/c.md": "p",
|
||||||
|
".pytest_cache/c.md": "pc",
|
||||||
|
"dist/d.md": "d",
|
||||||
|
"build/b.md": "b",
|
||||||
|
}
|
||||||
|
for rel, text in files.items():
|
||||||
|
(root / rel).write_text(text)
|
||||||
|
(root / "notes" / "not-md.txt").write_text("skip me")
|
||||||
|
|
||||||
|
found = {p.relative_to(root).as_posix() for p in iter_markdown_files(root)}
|
||||||
|
assert found == {"README.md", "notes/sub/deep.md"}
|
||||||
|
|
||||||
|
|
||||||
|
def test_iter_markdown_files_missing_dir_yields_nothing(tmp_path: Path) -> None:
|
||||||
|
assert iter_markdown_files(tmp_path / "definitely-missing") == []
|
||||||
|
|
||||||
|
|
||||||
|
def test_excluded_dirs_match_plan_anchor_a9() -> None:
|
||||||
|
assert {
|
||||||
|
".venv", "node_modules", ".git", "__pycache__", ".pytest_cache", "dist", "build"
|
||||||
|
} == EXCLUDED_DIRS
|
||||||
|
|
||||||
|
|
||||||
|
def test_added_then_unchanged_then_updated_then_pruned(db, tmp_path: Path) -> None:
|
||||||
|
root = tmp_path / "src"
|
||||||
|
root.mkdir()
|
||||||
|
(root / "a.md").write_text("# A\n\nalpha\n\n## Sub\n\nmore alpha\n")
|
||||||
|
(root / "b.md").write_text("# B\n\nbeta\n")
|
||||||
|
llm = FakeEmbedder()
|
||||||
|
try:
|
||||||
|
s1 = asyncio.run(import_sources([root], llm, session=db))
|
||||||
|
assert (s1.files, s1.added, s1.unchanged, s1.updated, s1.pruned) == (2, 2, 0, 0, 0)
|
||||||
|
# a.md has two sections (2 chunks), b.md one (1 chunk).
|
||||||
|
assert s1.chunks == 3
|
||||||
|
# Embeddings are stored with the configured dimension.
|
||||||
|
n = db.scalar(
|
||||||
|
select(func.count())
|
||||||
|
.select_from(Chunk)
|
||||||
|
.join(Document, Document.id == Chunk.document_id)
|
||||||
|
.where(Document.source == root.name)
|
||||||
|
)
|
||||||
|
assert n == 3
|
||||||
|
for c in db.scalars(
|
||||||
|
select(Chunk)
|
||||||
|
.join(Document, Document.id == Chunk.document_id)
|
||||||
|
.where(Document.source == root.name)
|
||||||
|
).all():
|
||||||
|
assert c.embedding is not None and len(c.embedding) == 768
|
||||||
|
|
||||||
|
s2 = asyncio.run(import_sources([root], llm, session=db))
|
||||||
|
assert s2.added == 0 and s2.unchanged == 2
|
||||||
|
|
||||||
|
(root / "a.md").write_text("# A\n\nalpha CHANGED\n")
|
||||||
|
s3 = asyncio.run(import_sources([root], llm, session=db))
|
||||||
|
assert s3.updated == 1 and s3.unchanged == 1
|
||||||
|
doc = db.scalar(
|
||||||
|
select(Document).where(Document.source == root.name, Document.path == "a.md")
|
||||||
|
)
|
||||||
|
assert doc is not None and "CHANGED" in doc.content
|
||||||
|
|
||||||
|
(root / "a.md").unlink()
|
||||||
|
s4 = asyncio.run(import_sources([root], llm, session=db, prune=True))
|
||||||
|
assert s4.pruned == 1
|
||||||
|
assert db.scalar(
|
||||||
|
select(Document).where(Document.source == root.name, Document.path == "a.md")
|
||||||
|
) is None
|
||||||
|
# Chunks of the pruned document are gone (FK cascade).
|
||||||
|
n_after = db.scalar(
|
||||||
|
select(func.count())
|
||||||
|
.select_from(Chunk)
|
||||||
|
.join(Document, Document.id == Chunk.document_id)
|
||||||
|
.where(Document.source == root.name)
|
||||||
|
)
|
||||||
|
assert n_after == 1
|
||||||
|
finally:
|
||||||
|
_cleanup_source(db, root.name)
|
||||||
|
|
||||||
|
|
||||||
|
def test_embedding_failure_is_logged_and_import_continues(db, tmp_path: Path) -> None:
|
||||||
|
"""A file the embedding endpoint refuses must not abort the whole KB:
|
||||||
|
its rows are rolled back, the error is counted, and other files import."""
|
||||||
|
root = tmp_path / "mixed"
|
||||||
|
root.mkdir()
|
||||||
|
(root / "bad.md").write_text("# Bad\n\npoison content that the endpoint refuses\n")
|
||||||
|
(root / "good.md").write_text("# Good\n\nperfectly fine content\n")
|
||||||
|
try:
|
||||||
|
summary = asyncio.run(import_sources([root], _PoisonEmbedder(), session=db))
|
||||||
|
assert summary.files == 2
|
||||||
|
assert summary.errors == 1
|
||||||
|
assert summary.added == 1 # only good.md
|
||||||
|
# bad.md left no row and no orphan chunks behind (rolled back).
|
||||||
|
assert db.scalar(
|
||||||
|
select(Document).where(Document.source == root.name, Document.path == "bad.md")
|
||||||
|
) is None
|
||||||
|
assert db.scalar(
|
||||||
|
select(func.count())
|
||||||
|
.select_from(Chunk)
|
||||||
|
.join(Document, Document.id == Chunk.document_id)
|
||||||
|
.where(Document.source == root.name, Document.path == "bad.md")
|
||||||
|
) == 0
|
||||||
|
assert db.scalar(
|
||||||
|
select(Document).where(Document.source == root.name, Document.path == "good.md")
|
||||||
|
) is not None
|
||||||
|
finally:
|
||||||
|
_cleanup_source(db, root.name)
|
||||||
|
|
||||||
|
|
||||||
|
def test_oversized_chunk_triggers_adaptive_rechunk(db, tmp_path: Path) -> None:
|
||||||
|
"""A URL-dense paragraph the endpoint rejects must be re-chunked smaller
|
||||||
|
for that file only — the import still succeeds."""
|
||||||
|
root = tmp_path / "dense"
|
||||||
|
root.mkdir()
|
||||||
|
# One ~1165-char paragraph: under the 1200-char hard cap, over the
|
||||||
|
# simulated token cap. The retry at 600 chars must split it.
|
||||||
|
para = "see https://example.com/" + "a" * 1100
|
||||||
|
(root / "dense.md").write_text(f"# D\n\n{para}\n")
|
||||||
|
try:
|
||||||
|
summary = asyncio.run(import_sources([root], _CapEmbedder(), session=db))
|
||||||
|
assert summary.errors == 0
|
||||||
|
assert summary.added == 1
|
||||||
|
doc = db.scalar(
|
||||||
|
select(Document).where(Document.source == root.name, Document.path == "dense.md")
|
||||||
|
)
|
||||||
|
assert doc is not None
|
||||||
|
assert len(doc.chunks) >= 2 # re-chunked smaller than the hard cap
|
||||||
|
assert all(len(c.content) <= 1000 for c in doc.chunks)
|
||||||
|
assert all(c.embedding is not None for c in doc.chunks)
|
||||||
|
# The content survives the split.
|
||||||
|
assert "".join(c.content for c in doc.chunks).count("a" * 500) >= 1
|
||||||
|
finally:
|
||||||
|
_cleanup_source(db, root.name)
|
||||||
|
|
||||||
|
|
||||||
|
def test_missing_source_dir_is_skipped(db, tmp_path: Path) -> None:
|
||||||
|
llm = FakeEmbedder()
|
||||||
|
summary = asyncio.run(import_sources([tmp_path / "missing"], llm, session=db))
|
||||||
|
assert summary.files == 0 and summary.added == 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_limit_caps_files_and_disables_prune(db, tmp_path: Path) -> None:
|
||||||
|
root = tmp_path / "limited"
|
||||||
|
root.mkdir()
|
||||||
|
for name in ("a.md", "b.md", "c.md"):
|
||||||
|
(root / name).write_text(f"# {name}\n\nbody {name}\n")
|
||||||
|
llm = FakeEmbedder()
|
||||||
|
try:
|
||||||
|
summary = asyncio.run(import_sources([root], llm, limit=2, session=db, prune=True))
|
||||||
|
assert summary.files == 2 and summary.added == 2
|
||||||
|
# c.md was never walked, so it must NOT be pruned (prune disabled
|
||||||
|
# under --limit) — and nothing else disappears either.
|
||||||
|
assert summary.pruned == 0
|
||||||
|
assert db.scalar(
|
||||||
|
select(func.count()).select_from(Document).where(Document.source == root.name)
|
||||||
|
) == 2
|
||||||
|
finally:
|
||||||
|
_cleanup_source(db, root.name)
|
||||||
|
|
||||||
|
|
||||||
|
def test_limit_must_be_positive(db, tmp_path: Path) -> None:
|
||||||
|
with pytest.raises(ValueError):
|
||||||
|
asyncio.run(import_sources([tmp_path], FakeEmbedder(), limit=0, session=db))
|
||||||
|
|
||||||
|
|
||||||
|
def test_prune_is_scoped_to_the_given_sources(db, tmp_path: Path) -> None:
|
||||||
|
src_x = tmp_path / "SourceX"
|
||||||
|
src_y = tmp_path / "SourceY"
|
||||||
|
src_x.mkdir()
|
||||||
|
src_y.mkdir()
|
||||||
|
(src_x / "x.md").write_text("# X\n\nx body\n")
|
||||||
|
(src_y / "y.md").write_text("# Y\n\ny body\n")
|
||||||
|
llm = FakeEmbedder()
|
||||||
|
try:
|
||||||
|
asyncio.run(import_sources([src_x, src_y], llm, session=db))
|
||||||
|
# Re-import ONLY source Y (y.md removed) with prune: source X's doc
|
||||||
|
# must survive — prune never touches sources not passed to this run.
|
||||||
|
(src_y / "y.md").unlink()
|
||||||
|
summary = asyncio.run(import_sources([src_y], llm, session=db, prune=True))
|
||||||
|
assert summary.pruned == 1
|
||||||
|
assert db.scalar(
|
||||||
|
select(Document).where(Document.source == "SourceX", Document.path == "x.md")
|
||||||
|
) is not None
|
||||||
|
finally:
|
||||||
|
_cleanup_source(db, "SourceX")
|
||||||
|
_cleanup_source(db, "SourceY")
|
||||||
|
|
||||||
|
|
||||||
|
def test_chunk_positions_and_titles(db, tmp_path: Path) -> None:
|
||||||
|
root = tmp_path / "titled"
|
||||||
|
root.mkdir()
|
||||||
|
(root / "multi.md").write_text("# Real Title\n\n## One\n\na\n\n## Two\n\nb\n")
|
||||||
|
(root / "noh1.md").write_text("## Only heading\n\nbody\n")
|
||||||
|
llm = FakeEmbedder()
|
||||||
|
try:
|
||||||
|
asyncio.run(import_sources([root], llm, session=db))
|
||||||
|
titles = {
|
||||||
|
d.path: d.title
|
||||||
|
for d in db.scalars(select(Document).where(Document.source == root.name)).all()
|
||||||
|
}
|
||||||
|
assert titles["multi.md"] == "Real Title" # H1 wins
|
||||||
|
assert titles["noh1.md"] == "noh1" # …else the file stem
|
||||||
|
doc = db.scalar(
|
||||||
|
select(Document).where(Document.source == root.name, Document.path == "multi.md")
|
||||||
|
)
|
||||||
|
assert doc is not None
|
||||||
|
positions = sorted(c.position for c in doc.chunks)
|
||||||
|
assert positions == list(range(len(doc.chunks))) and len(doc.chunks) >= 2
|
||||||
|
finally:
|
||||||
|
_cleanup_source(db, root.name)
|
||||||
@@ -0,0 +1,232 @@
|
|||||||
|
"""Unit tests: LLMClient embeddings (batching, order, loud dim failure).
|
||||||
|
|
||||||
|
The fakes stand in at the httpx-transport layer — that is where LLMClient
|
||||||
|
actually talks to the endpoint (see ``LLMClient._embed_batch`` in
|
||||||
|
``app/rag/llm.py`` for why the openai SDK's own ``embeddings.create`` is
|
||||||
|
bypassed: it injects ``encoding_format``, which aipi's litellm proxy
|
||||||
|
rejects).
|
||||||
|
"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import json
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from app.config import Settings
|
||||||
|
from app.rag.llm import EmbeddingDimensionError, EmbeddingError, LLMClient
|
||||||
|
|
||||||
|
|
||||||
|
def _settings(**kwargs: Any) -> Settings:
|
||||||
|
kwargs.setdefault("_env_file", None)
|
||||||
|
return Settings(**kwargs) # pyright: ignore[reportCallIssue] (kwarg exists at runtime)
|
||||||
|
|
||||||
|
|
||||||
|
class _Row:
|
||||||
|
def __init__(self, index: int, embedding: list[float]) -> None:
|
||||||
|
self.index = index
|
||||||
|
self.embedding = embedding
|
||||||
|
|
||||||
|
|
||||||
|
class _Response:
|
||||||
|
def __init__(self, rows: list[_Row]) -> None:
|
||||||
|
self.data = rows
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeEmbeddingsService:
|
||||||
|
"""Simulates the /embeddings endpoint; records calls; can fail."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
dim: int = 768,
|
||||||
|
fail: Exception | None = None,
|
||||||
|
drop_index: int = -1,
|
||||||
|
http_error: int | None = None,
|
||||||
|
too_large_min: int | None = None,
|
||||||
|
) -> None:
|
||||||
|
self.dim = dim
|
||||||
|
self.fail = fail
|
||||||
|
self.drop_index = drop_index
|
||||||
|
self.http_error = http_error
|
||||||
|
self.too_large_min = too_large_min
|
||||||
|
self.calls: list[list[str]] = []
|
||||||
|
|
||||||
|
async def create(self, *, model: str, input: list[str]) -> _Response:
|
||||||
|
self.calls.append(list(input))
|
||||||
|
if self.fail is not None:
|
||||||
|
raise self.fail
|
||||||
|
return _Response(
|
||||||
|
[
|
||||||
|
_Row(i, [0.5] * self.dim)
|
||||||
|
for i in range(len(input))
|
||||||
|
if i != self.drop_index
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeHttpResponse:
|
||||||
|
def __init__(
|
||||||
|
self, status_code: int, payload: dict[str, Any] | None = None, text: str = ""
|
||||||
|
) -> None:
|
||||||
|
self.status_code = status_code
|
||||||
|
self._payload = payload
|
||||||
|
self.text = text or (json.dumps(payload) if payload is not None else "boom")
|
||||||
|
|
||||||
|
def json(self) -> Any:
|
||||||
|
if self._payload is None:
|
||||||
|
raise ValueError("no json body")
|
||||||
|
return self._payload
|
||||||
|
|
||||||
|
|
||||||
|
#: The endpoint's real error phrasing (litellm) — the client keys off it.
|
||||||
|
_TOO_LARGE_TEXT = 'input (9999 tokens) is too large to process. increase the physical batch size'
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeHttp:
|
||||||
|
"""Stands in for the httpx transport the openai client owns."""
|
||||||
|
|
||||||
|
def __init__(self, service: _FakeEmbeddingsService) -> None:
|
||||||
|
self.service = service
|
||||||
|
self.bodies: list[dict[str, Any]] = []
|
||||||
|
|
||||||
|
async def post(
|
||||||
|
self, url: str, *, json: dict[str, Any], headers: dict[str, str] | None = None
|
||||||
|
) -> _FakeHttpResponse:
|
||||||
|
self.bodies.append(json)
|
||||||
|
assert "Authorization" in (headers or {})
|
||||||
|
if self.service.http_error is not None:
|
||||||
|
return _FakeHttpResponse(self.service.http_error)
|
||||||
|
if (
|
||||||
|
self.service.too_large_min is not None
|
||||||
|
and len(json["input"]) >= self.service.too_large_min
|
||||||
|
):
|
||||||
|
return _FakeHttpResponse(500, None, _TOO_LARGE_TEXT)
|
||||||
|
rows = await self.service.create(model=json["model"], input=json["input"])
|
||||||
|
payload = {"data": [{"index": r.index, "embedding": r.embedding} for r in rows.data]}
|
||||||
|
return _FakeHttpResponse(200, payload)
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeClient:
|
||||||
|
"""Stands in for the openai AsyncOpenAI object (only its transport is used)."""
|
||||||
|
|
||||||
|
def __init__(self, http: _FakeHttp) -> None:
|
||||||
|
self._client = http
|
||||||
|
|
||||||
|
|
||||||
|
def _make_client(service: _FakeEmbeddingsService, **kwargs: Any) -> tuple[LLMClient, _FakeHttp]:
|
||||||
|
kwargs.setdefault("embed_batch_size", 2)
|
||||||
|
llm = LLMClient(_settings(**kwargs))
|
||||||
|
http = _FakeHttp(service)
|
||||||
|
llm._client = _FakeClient(http) # pyright: ignore[reportAttributeAccessIssue]
|
||||||
|
return llm, http
|
||||||
|
|
||||||
|
|
||||||
|
def test_embed_batches_by_batch_size_and_keeps_order() -> None:
|
||||||
|
service = _FakeEmbeddingsService()
|
||||||
|
llm, http = _make_client(service)
|
||||||
|
texts = [f"t{i}" for i in range(5)]
|
||||||
|
vecs = asyncio.run(llm.embed(texts))
|
||||||
|
|
||||||
|
assert [len(c) for c in service.calls] == [2, 2, 1]
|
||||||
|
assert [t for call in service.calls for t in call] == texts
|
||||||
|
assert len(vecs) == 5
|
||||||
|
assert all(len(v) == 768 for v in vecs)
|
||||||
|
assert llm.embed_batches == 3
|
||||||
|
# aipi (litellm) rejects the SDK's injected "encoding_format" — the
|
||||||
|
# payload must stay a minimal {model, input} body.
|
||||||
|
assert all(set(b) == {"model", "input"} for b in http.bodies)
|
||||||
|
|
||||||
|
|
||||||
|
def test_embed_empty_returns_empty_without_calling_endpoint() -> None:
|
||||||
|
service = _FakeEmbeddingsService()
|
||||||
|
llm, http = _make_client(service)
|
||||||
|
assert asyncio.run(llm.embed([])) == []
|
||||||
|
assert http.bodies == []
|
||||||
|
assert service.calls == []
|
||||||
|
assert llm.embed_batches == 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_embed_one_returns_single_vector() -> None:
|
||||||
|
llm, _ = _make_client(_FakeEmbeddingsService())
|
||||||
|
vec = asyncio.run(llm.embed_one("hello"))
|
||||||
|
assert len(vec) == 768
|
||||||
|
|
||||||
|
|
||||||
|
def test_dim_mismatch_fails_loudly_with_actionable_message() -> None:
|
||||||
|
llm, _ = _make_client(_FakeEmbeddingsService(dim=512))
|
||||||
|
with pytest.raises(EmbeddingDimensionError) as exc:
|
||||||
|
asyncio.run(llm.embed(["hello"]))
|
||||||
|
msg = str(exc.value)
|
||||||
|
assert "512" in msg and "768" in msg
|
||||||
|
assert "BOR_EMBEDDING_DIM" in msg
|
||||||
|
assert "llm_probe" in msg
|
||||||
|
|
||||||
|
|
||||||
|
def test_endpoint_error_is_wrapped() -> None:
|
||||||
|
llm, _ = _make_client(_FakeEmbeddingsService(fail=RuntimeError("connection refused")))
|
||||||
|
with pytest.raises(EmbeddingError, match="connection refused"):
|
||||||
|
asyncio.run(llm.embed(["hello"]))
|
||||||
|
assert llm.embed_batches == 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_http_error_surfaces_status() -> None:
|
||||||
|
llm, _ = _make_client(_FakeEmbeddingsService(http_error=502))
|
||||||
|
with pytest.raises(EmbeddingError, match="HTTP 502"):
|
||||||
|
asyncio.run(llm.embed(["hello"]))
|
||||||
|
assert llm.embed_batches == 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_missing_vector_row_is_rejected() -> None:
|
||||||
|
llm, _ = _make_client(_FakeEmbeddingsService(drop_index=1))
|
||||||
|
with pytest.raises(EmbeddingError, match="returned 1 vectors for 2 inputs"):
|
||||||
|
asyncio.run(llm.embed(["a", "b"]))
|
||||||
|
|
||||||
|
|
||||||
|
def test_batch_size_one_forces_one_call_per_text() -> None:
|
||||||
|
service = _FakeEmbeddingsService()
|
||||||
|
llm, _ = _make_client(service, embed_batch_size=1)
|
||||||
|
asyncio.run(llm.embed(["a", "b", "c"]))
|
||||||
|
assert [len(c) for c in service.calls] == [1, 1, 1]
|
||||||
|
|
||||||
|
|
||||||
|
def test_token_budget_limits_texts_per_request() -> None:
|
||||||
|
"""~2000-char chunks must not stack up past aipi's ~1024-token cap."""
|
||||||
|
service = _FakeEmbeddingsService()
|
||||||
|
llm, _ = _make_client(service, embed_batch_size=16) # high count cap
|
||||||
|
texts = ["x" * 2000 for _ in range(4)]
|
||||||
|
vecs = asyncio.run(llm.embed(texts))
|
||||||
|
# 2000 + 2000 chars > 3600-char (≈900-token) budget ⇒ one chunk per request
|
||||||
|
assert [len(c) for c in service.calls] == [1, 1, 1, 1]
|
||||||
|
assert len(vecs) == 4
|
||||||
|
|
||||||
|
|
||||||
|
def test_small_chunks_pack_up_to_count_cap() -> None:
|
||||||
|
service = _FakeEmbeddingsService()
|
||||||
|
llm, _ = _make_client(service, embed_batch_size=4) # count cap binds
|
||||||
|
texts = ["short text" for _ in range(9)]
|
||||||
|
vecs = asyncio.run(llm.embed(texts))
|
||||||
|
assert [len(c) for c in service.calls] == [4, 4, 1]
|
||||||
|
assert len(vecs) == 9
|
||||||
|
|
||||||
|
|
||||||
|
def test_too_large_response_halves_batch_until_it_fits() -> None:
|
||||||
|
"""The tokenizer estimate can be wrong for dense content — the client
|
||||||
|
must halve an over-large request and preserve order."""
|
||||||
|
service = _FakeEmbeddingsService(too_large_min=3)
|
||||||
|
llm, http = _make_client(service, embed_batch_size=16) # all 8 fit one request
|
||||||
|
texts = [f"t{i}" for i in range(8)]
|
||||||
|
vecs = asyncio.run(llm.embed(texts))
|
||||||
|
# http.bodies sees every request (including the rejected ones); the
|
||||||
|
# left-half recursion completes before the right half starts.
|
||||||
|
assert [len(b["input"]) for b in http.bodies] == [8, 4, 2, 2, 4, 2, 2]
|
||||||
|
assert len(vecs) == 8
|
||||||
|
assert all(len(v) == 768 for v in vecs)
|
||||||
|
|
||||||
|
|
||||||
|
def test_single_oversized_text_fails_actionably() -> None:
|
||||||
|
service = _FakeEmbeddingsService(too_large_min=1)
|
||||||
|
llm, _ = _make_client(service)
|
||||||
|
with pytest.raises(EmbeddingError, match="token cap"):
|
||||||
|
asyncio.run(llm.embed(["x" * 3000]))
|
||||||
|
assert llm.embed_batches == 0
|
||||||
Reference in New Issue
Block a user