mirror of
https://github.com/basicmachines-co/basic-memory
synced 2026-06-21 13:47:35 +00:00
Merge branch 'main' of github.com:basicmachines-co/basic-memory
This commit is contained in:
@@ -26,7 +26,7 @@ class EmbeddingProgress:
|
||||
"""Typed CLI progress payload for embedding backfills."""
|
||||
|
||||
entity_id: int
|
||||
index: int
|
||||
completed: int
|
||||
total: int
|
||||
|
||||
|
||||
@@ -147,20 +147,30 @@ def reindex(
|
||||
False, "--embeddings", "-e", help="Rebuild vector embeddings (requires semantic search)"
|
||||
),
|
||||
search: bool = typer.Option(False, "--search", "-s", help="Rebuild full-text search index"),
|
||||
full: bool = typer.Option(
|
||||
False,
|
||||
"--full",
|
||||
help="Force a full filesystem scan and file reindex instead of the default incremental scan",
|
||||
),
|
||||
project: str = typer.Option(
|
||||
None, "--project", "-p", help="Reindex a specific project (default: all)"
|
||||
),
|
||||
): # pragma: no cover
|
||||
"""Rebuild search indexes and/or vector embeddings without dropping the database.
|
||||
|
||||
By default rebuilds everything (search + embeddings if semantic is enabled).
|
||||
Use --search or --embeddings to rebuild only one.
|
||||
By default runs incremental search + embeddings (if semantic search is enabled).
|
||||
Use --full to bypass incremental scan optimization, rebuild all file-backed search rows,
|
||||
and re-embed all eligible notes.
|
||||
Use --search or --embeddings to rebuild only one side.
|
||||
|
||||
Examples:
|
||||
bm reindex # Rebuild everything
|
||||
bm reindex # Incremental search + embeddings
|
||||
bm reindex --full # Full search + full re-embed
|
||||
bm reindex --embeddings # Only rebuild vector embeddings
|
||||
bm reindex --search # Only rebuild FTS index
|
||||
bm reindex -p claw # Reindex only the 'claw' project
|
||||
bm reindex --full --search # Full search only
|
||||
bm reindex --full --embeddings # Full re-embed only
|
||||
bm reindex -p claw --full # Full reindex for only the 'claw' project
|
||||
"""
|
||||
# If neither flag is set, do both
|
||||
if not embeddings and not search:
|
||||
@@ -179,10 +189,19 @@ def reindex(
|
||||
if not search:
|
||||
raise typer.Exit(0)
|
||||
|
||||
run_with_cleanup(_reindex(app_config, search=search, embeddings=embeddings, project=project))
|
||||
run_with_cleanup(
|
||||
_reindex(app_config, search=search, embeddings=embeddings, full=full, project=project)
|
||||
)
|
||||
|
||||
|
||||
async def _reindex(app_config, search: bool, embeddings: bool, project: str | None):
|
||||
async def _reindex(
|
||||
app_config,
|
||||
*,
|
||||
search: bool,
|
||||
embeddings: bool,
|
||||
full: bool,
|
||||
project: str | None,
|
||||
):
|
||||
"""Run reindex operations."""
|
||||
from basic_memory.repository import EntityRepository
|
||||
from basic_memory.repository.search_repository import create_search_repository
|
||||
@@ -220,6 +239,10 @@ async def _reindex(app_config, search: bool, embeddings: bool, project: str | No
|
||||
console.print(f"\n[bold]Project: [cyan]{proj.name}[/cyan][/bold]")
|
||||
|
||||
if search:
|
||||
search_mode_label = "full scan" if full else "incremental scan"
|
||||
console.print(
|
||||
f" Rebuilding full-text search index ([cyan]{search_mode_label}[/cyan])..."
|
||||
)
|
||||
sync_service = await get_sync_service(proj)
|
||||
sync_dir = Path(proj.path)
|
||||
with Progress(
|
||||
@@ -244,14 +267,19 @@ async def _reindex(app_config, search: bool, embeddings: bool, project: str | No
|
||||
await sync_service.sync(
|
||||
sync_dir,
|
||||
project_name=proj.name,
|
||||
force_full=full,
|
||||
sync_embeddings=False,
|
||||
progress_callback=on_index_progress,
|
||||
)
|
||||
progress.update(task, completed=progress.tasks[task].total or 1)
|
||||
|
||||
console.print(" [green]✓[/green] Full-text search index rebuilt")
|
||||
console.print(" [green]done[/green] Full-text search index rebuilt")
|
||||
|
||||
if embeddings:
|
||||
console.print(" Building vector embeddings...")
|
||||
embedding_mode_label = "full rebuild" if full else "incremental sync"
|
||||
console.print(
|
||||
f" Building vector embeddings ([cyan]{embedding_mode_label}[/cyan])..."
|
||||
)
|
||||
entity_repository = EntityRepository(session_maker, project_id=proj.id)
|
||||
search_repository = create_search_repository(
|
||||
session_maker, project_id=proj.id, app_config=app_config
|
||||
@@ -274,20 +302,27 @@ async def _reindex(app_config, search: bool, embeddings: bool, project: str | No
|
||||
def on_progress(entity_id, index, total):
|
||||
embedding_progress = EmbeddingProgress(
|
||||
entity_id=entity_id,
|
||||
index=index,
|
||||
completed=index,
|
||||
total=total,
|
||||
)
|
||||
# Trigger: repository progress now reports terminal entity completion.
|
||||
# Why: operators need to see finished embedding work rather than
|
||||
# entities merely entering prepare.
|
||||
# Outcome: the CLI bar advances steadily with real completed work.
|
||||
progress.update(
|
||||
task,
|
||||
total=embedding_progress.total,
|
||||
completed=embedding_progress.index,
|
||||
completed=embedding_progress.completed,
|
||||
)
|
||||
|
||||
stats = await search_service.reindex_vectors(progress_callback=on_progress)
|
||||
stats = await search_service.reindex_vectors(
|
||||
progress_callback=on_progress,
|
||||
force_full=full,
|
||||
)
|
||||
progress.update(task, completed=stats["total_entities"])
|
||||
|
||||
console.print(
|
||||
f" [green]✓[/green] Embeddings complete: "
|
||||
f" [green]done[/green] Embeddings complete: "
|
||||
f"{stats['embedded']} entities embedded, "
|
||||
f"{stats['skipped']} skipped, "
|
||||
f"{stats['errors']} errors"
|
||||
|
||||
@@ -188,8 +188,14 @@ class BasicMemoryConfig(BaseSettings):
|
||||
default=None,
|
||||
description="Embedding vector dimensions. Auto-detected from provider if not set (384 for FastEmbed, 1536 for OpenAI).",
|
||||
)
|
||||
# Trigger: full local rebuilds spend most of their time waiting behind shared
|
||||
# embed flushes, not constructing vectors themselves.
|
||||
# Why: smaller FastEmbed batches cut queue wait far more than they increase
|
||||
# write overhead on real-world projects, which makes full reindex materially faster.
|
||||
# Outcome: default to the smaller local/cloud-safe batch size we benchmarked as
|
||||
# the current best end-to-end setting in the shared vector sync pipeline.
|
||||
semantic_embedding_batch_size: int = Field(
|
||||
default=64,
|
||||
default=2,
|
||||
description="Batch size for embedding generation.",
|
||||
gt=0,
|
||||
)
|
||||
@@ -199,7 +205,7 @@ class BasicMemoryConfig(BaseSettings):
|
||||
gt=0,
|
||||
)
|
||||
semantic_embedding_sync_batch_size: int = Field(
|
||||
default=64,
|
||||
default=2,
|
||||
description="Batch size for vector sync orchestration flushes.",
|
||||
gt=0,
|
||||
)
|
||||
|
||||
@@ -114,7 +114,13 @@ async def write_file_atomic(path: FilePath, content: str) -> None:
|
||||
temp_path = path_obj.with_suffix(".tmp")
|
||||
|
||||
try:
|
||||
# Use aiofiles for non-blocking write
|
||||
# Trigger: callers hand us normalized Python text, but the final bytes are allowed
|
||||
# to use the host platform's native newline convention during the write.
|
||||
# Why: preserving CRLF on Windows keeps local files aligned with editors like
|
||||
# Obsidian, while FileService now hashes the persisted file bytes instead of
|
||||
# the pre-write string.
|
||||
# Outcome: this async write stays editor-friendly across platforms without
|
||||
# reintroducing checksum drift in sync or move detection.
|
||||
async with aiofiles.open(temp_path, mode="w", encoding="utf-8") as f:
|
||||
await f.write(content)
|
||||
|
||||
@@ -168,6 +174,13 @@ async def format_markdown_builtin(path: Path) -> Optional[str]:
|
||||
|
||||
# Only write if content changed
|
||||
if formatted_content != content:
|
||||
# Trigger: mdformat may rewrite markdown content, then the host platform
|
||||
# decides the newline bytes for the follow-up async text write.
|
||||
# Why: we want formatter output to preserve native newlines instead of
|
||||
# forcing LF, and the authoritative checksum comes from rereading the
|
||||
# stored file bytes later in FileService.
|
||||
# Outcome: formatting remains compatible with local editors on Windows while
|
||||
# checksum-based sync logic stays anchored to on-disk bytes.
|
||||
async with aiofiles.open(path, mode="w", encoding="utf-8") as f:
|
||||
await f.write(formatted_content)
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
"""Embedding provider protocol for pluggable semantic backends."""
|
||||
|
||||
from typing import Protocol
|
||||
from typing import Any, Protocol
|
||||
|
||||
|
||||
class EmbeddingProvider(Protocol):
|
||||
@@ -16,3 +16,7 @@ class EmbeddingProvider(Protocol):
|
||||
async def embed_documents(self, texts: list[str]) -> list[list[float]]:
|
||||
"""Embed a list of document chunks."""
|
||||
...
|
||||
|
||||
def runtime_log_attrs(self) -> dict[str, Any]:
|
||||
"""Return provider-specific runtime settings suitable for startup logs."""
|
||||
...
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
"""Factory for creating configured semantic embedding providers."""
|
||||
|
||||
import os
|
||||
from threading import Lock
|
||||
|
||||
from basic_memory.config import BasicMemoryConfig
|
||||
@@ -18,10 +19,50 @@ type ProviderCacheKey = tuple[
|
||||
|
||||
_EMBEDDING_PROVIDER_CACHE: dict[ProviderCacheKey, EmbeddingProvider] = {}
|
||||
_EMBEDDING_PROVIDER_CACHE_LOCK = Lock()
|
||||
_FASTEMBED_MAX_THREADS = 8
|
||||
|
||||
|
||||
def _available_cpu_count() -> int | None:
|
||||
"""Return the CPU budget available to this process when the runtime exposes it."""
|
||||
process_cpu_count = getattr(os, "process_cpu_count", None)
|
||||
if callable(process_cpu_count):
|
||||
cpu_count = process_cpu_count()
|
||||
if isinstance(cpu_count, int) and cpu_count > 0:
|
||||
return cpu_count
|
||||
|
||||
cpu_count = os.cpu_count()
|
||||
return cpu_count if cpu_count is not None and cpu_count > 0 else None
|
||||
|
||||
|
||||
def _resolve_fastembed_runtime_knobs(
|
||||
app_config: BasicMemoryConfig,
|
||||
) -> tuple[int | None, int | None]:
|
||||
"""Resolve FastEmbed threads/parallel from explicit config or CPU-aware defaults."""
|
||||
configured_threads = app_config.semantic_embedding_threads
|
||||
configured_parallel = app_config.semantic_embedding_parallel
|
||||
if configured_threads is not None or configured_parallel is not None:
|
||||
return configured_threads, configured_parallel
|
||||
|
||||
available_cpus = _available_cpu_count()
|
||||
if available_cpus is None:
|
||||
return None, None
|
||||
|
||||
# Trigger: local laptops and cloud workers expose different CPU budgets.
|
||||
# Why: full rebuilds got faster when FastEmbed used most, but not all, of
|
||||
# the available CPUs. Leaving a little headroom avoids starving the rest of
|
||||
# the pipeline while still giving ONNX enough threads to stay busy.
|
||||
# Outcome: when config leaves the knobs unset, each process reserves a small
|
||||
# CPU cushion and keeps FastEmbed on the simpler single-process path.
|
||||
if available_cpus <= 2:
|
||||
return available_cpus, 1
|
||||
|
||||
threads = min(_FASTEMBED_MAX_THREADS, max(2, available_cpus - 2))
|
||||
return threads, 1
|
||||
|
||||
|
||||
def _provider_cache_key(app_config: BasicMemoryConfig) -> ProviderCacheKey:
|
||||
"""Build a stable cache key from provider-relevant semantic embedding config."""
|
||||
resolved_threads, resolved_parallel = _resolve_fastembed_runtime_knobs(app_config)
|
||||
return (
|
||||
app_config.semantic_embedding_provider.strip().lower(),
|
||||
app_config.semantic_embedding_model,
|
||||
@@ -29,8 +70,8 @@ def _provider_cache_key(app_config: BasicMemoryConfig) -> ProviderCacheKey:
|
||||
app_config.semantic_embedding_batch_size,
|
||||
app_config.semantic_embedding_request_concurrency,
|
||||
app_config.semantic_embedding_cache_dir,
|
||||
app_config.semantic_embedding_threads,
|
||||
app_config.semantic_embedding_parallel,
|
||||
resolved_threads,
|
||||
resolved_parallel,
|
||||
)
|
||||
|
||||
|
||||
@@ -61,12 +102,13 @@ def create_embedding_provider(app_config: BasicMemoryConfig) -> EmbeddingProvide
|
||||
# Deferred import: fastembed (and its onnxruntime dep) may not be installed
|
||||
from basic_memory.repository.fastembed_provider import FastEmbedEmbeddingProvider
|
||||
|
||||
resolved_threads, resolved_parallel = _resolve_fastembed_runtime_knobs(app_config)
|
||||
if app_config.semantic_embedding_cache_dir is not None:
|
||||
extra_kwargs["cache_dir"] = app_config.semantic_embedding_cache_dir
|
||||
if app_config.semantic_embedding_threads is not None:
|
||||
extra_kwargs["threads"] = app_config.semantic_embedding_threads
|
||||
if app_config.semantic_embedding_parallel is not None:
|
||||
extra_kwargs["parallel"] = app_config.semantic_embedding_parallel
|
||||
if resolved_threads is not None:
|
||||
extra_kwargs["threads"] = resolved_threads
|
||||
if resolved_parallel is not None:
|
||||
extra_kwargs["parallel"] = resolved_parallel
|
||||
|
||||
provider = FastEmbedEmbeddingProvider(
|
||||
model_name=app_config.semantic_embedding_model,
|
||||
|
||||
@@ -24,6 +24,15 @@ class FastEmbedEmbeddingProvider(EmbeddingProvider):
|
||||
def _effective_parallel(self) -> int | None:
|
||||
return self.parallel if self.parallel is not None and self.parallel > 1 else None
|
||||
|
||||
def runtime_log_attrs(self) -> dict[str, int | str | None]:
|
||||
"""Return the resolved runtime knobs that shape FastEmbed throughput."""
|
||||
return {
|
||||
"provider_batch_size": self.batch_size,
|
||||
"threads": self.threads,
|
||||
"configured_parallel": self.parallel,
|
||||
"effective_parallel": self._effective_parallel(),
|
||||
}
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str = "bge-small-en-v1.5",
|
||||
|
||||
@@ -34,6 +34,13 @@ class OpenAIEmbeddingProvider(EmbeddingProvider):
|
||||
self._client: Any | None = None
|
||||
self._client_lock = asyncio.Lock()
|
||||
|
||||
def runtime_log_attrs(self) -> dict[str, int]:
|
||||
"""Return the request fan-out knobs that shape API embedding batches."""
|
||||
return {
|
||||
"provider_batch_size": self.batch_size,
|
||||
"request_concurrency": self.request_concurrency,
|
||||
}
|
||||
|
||||
async def _get_client(self) -> Any:
|
||||
if self._client is not None:
|
||||
return self._client
|
||||
|
||||
@@ -3,9 +3,8 @@
|
||||
import asyncio
|
||||
import json
|
||||
import re
|
||||
import time
|
||||
from datetime import datetime
|
||||
from typing import List, Optional, cast
|
||||
from typing import List, Optional
|
||||
|
||||
from loguru import logger
|
||||
from sqlalchemy import text
|
||||
@@ -18,7 +17,7 @@ from basic_memory.repository.embedding_provider_factory import create_embedding_
|
||||
from basic_memory.repository.search_index_row import SearchIndexRow
|
||||
from basic_memory.repository.search_repository_base import (
|
||||
SearchRepositoryBase,
|
||||
_PreparedEntityVectorSync,
|
||||
VectorChunkState,
|
||||
)
|
||||
from basic_memory.repository.metadata_filters import parse_metadata_filters
|
||||
from basic_memory.repository.semantic_errors import SemanticDependenciesMissingError
|
||||
@@ -458,248 +457,71 @@ class PostgresSearchRepository(SearchRepositoryBase):
|
||||
"""Use a bounded config-driven prepare window for Postgres vector sync."""
|
||||
return self._semantic_postgres_prepare_concurrency
|
||||
|
||||
async def _prepare_entity_vector_jobs_window(
|
||||
self, entity_ids: list[int]
|
||||
) -> list[_PreparedEntityVectorSync | BaseException]:
|
||||
"""Prepare one Postgres window concurrently to hide DB round-trip latency."""
|
||||
prepared_window = await asyncio.gather(
|
||||
*(self._prepare_entity_vector_jobs(entity_id) for entity_id in entity_ids),
|
||||
return_exceptions=True,
|
||||
async def _upsert_scheduled_chunk_records(
|
||||
self,
|
||||
session: AsyncSession,
|
||||
*,
|
||||
entity_id: int,
|
||||
scheduled_records: list[dict[str, str]],
|
||||
existing_by_key: dict[str, VectorChunkState],
|
||||
entity_fingerprint: str,
|
||||
embedding_model: str,
|
||||
) -> list[tuple[int, str]]:
|
||||
"""Use Postgres UPSERT to rewrite only the scheduled chunk rows."""
|
||||
if not scheduled_records:
|
||||
return []
|
||||
|
||||
upsert_params: dict[str, object] = {
|
||||
"project_id": self.project_id,
|
||||
"entity_id": entity_id,
|
||||
}
|
||||
upsert_values: list[str] = []
|
||||
# The SQL template is built from integer enumerate() indices only.
|
||||
# No user-controlled text is interpolated into the statement.
|
||||
for index, record in enumerate(scheduled_records):
|
||||
upsert_params[f"chunk_key_{index}"] = record["chunk_key"]
|
||||
upsert_params[f"chunk_text_{index}"] = record["chunk_text"]
|
||||
upsert_params[f"source_hash_{index}"] = record["source_hash"]
|
||||
upsert_params[f"entity_fingerprint_{index}"] = entity_fingerprint
|
||||
upsert_params[f"embedding_model_{index}"] = embedding_model
|
||||
upsert_values.append(
|
||||
"("
|
||||
":entity_id, :project_id, "
|
||||
f":chunk_key_{index}, :chunk_text_{index}, :source_hash_{index}, "
|
||||
f":entity_fingerprint_{index}, :embedding_model_{index}, NOW()"
|
||||
")"
|
||||
)
|
||||
|
||||
upsert_result = await session.execute(
|
||||
text(f"""
|
||||
INSERT INTO search_vector_chunks (
|
||||
entity_id,
|
||||
project_id,
|
||||
chunk_key,
|
||||
chunk_text,
|
||||
source_hash,
|
||||
entity_fingerprint,
|
||||
embedding_model,
|
||||
updated_at
|
||||
) VALUES {", ".join(upsert_values)}
|
||||
ON CONFLICT (project_id, entity_id, chunk_key) DO UPDATE SET
|
||||
chunk_text = EXCLUDED.chunk_text,
|
||||
source_hash = EXCLUDED.source_hash,
|
||||
entity_fingerprint = EXCLUDED.entity_fingerprint,
|
||||
embedding_model = EXCLUDED.embedding_model,
|
||||
updated_at = NOW()
|
||||
RETURNING id, chunk_key
|
||||
"""),
|
||||
upsert_params,
|
||||
)
|
||||
upserted_ids_by_key = {
|
||||
str(row["chunk_key"]): int(row["id"]) for row in upsert_result.mappings().all()
|
||||
}
|
||||
return [
|
||||
cast(_PreparedEntityVectorSync | BaseException, prepared)
|
||||
for prepared in prepared_window
|
||||
(upserted_ids_by_key[record["chunk_key"]], record["chunk_text"])
|
||||
for record in scheduled_records
|
||||
]
|
||||
|
||||
async def _prepare_entity_vector_jobs(self, entity_id: int) -> _PreparedEntityVectorSync:
|
||||
"""Prepare chunk mutations with Postgres-specific bulk upserts."""
|
||||
sync_start = time.perf_counter()
|
||||
|
||||
async with db.scoped_session(self.session_maker) as session:
|
||||
await self._prepare_vector_session(session)
|
||||
|
||||
row_result = await session.execute(
|
||||
text(
|
||||
"SELECT id, type, title, permalink, content_stems, content_snippet, "
|
||||
"category, relation_type "
|
||||
"FROM search_index "
|
||||
"WHERE entity_id = :entity_id AND project_id = :project_id "
|
||||
"ORDER BY "
|
||||
"CASE type "
|
||||
"WHEN :entity_type THEN 0 "
|
||||
"WHEN :observation_type THEN 1 "
|
||||
"WHEN :relation_type_type THEN 2 "
|
||||
"ELSE 3 END, id ASC"
|
||||
),
|
||||
{
|
||||
"entity_id": entity_id,
|
||||
"project_id": self.project_id,
|
||||
"entity_type": SearchItemType.ENTITY.value,
|
||||
"observation_type": SearchItemType.OBSERVATION.value,
|
||||
"relation_type_type": SearchItemType.RELATION.value,
|
||||
},
|
||||
)
|
||||
rows = row_result.fetchall()
|
||||
source_rows_count = len(rows)
|
||||
|
||||
if not rows:
|
||||
await self._delete_entity_chunks(session, entity_id)
|
||||
await session.commit()
|
||||
prepare_seconds = time.perf_counter() - sync_start
|
||||
return _PreparedEntityVectorSync(
|
||||
entity_id=entity_id,
|
||||
sync_start=sync_start,
|
||||
source_rows_count=source_rows_count,
|
||||
embedding_jobs=[],
|
||||
prepare_seconds=prepare_seconds,
|
||||
)
|
||||
|
||||
chunk_records = self._build_chunk_records(rows)
|
||||
built_chunk_records_count = len(chunk_records)
|
||||
current_entity_fingerprint = self._build_entity_fingerprint(chunk_records)
|
||||
current_embedding_model = self._embedding_model_key()
|
||||
if not chunk_records:
|
||||
await self._delete_entity_chunks(session, entity_id)
|
||||
await session.commit()
|
||||
prepare_seconds = time.perf_counter() - sync_start
|
||||
return _PreparedEntityVectorSync(
|
||||
entity_id=entity_id,
|
||||
sync_start=sync_start,
|
||||
source_rows_count=source_rows_count,
|
||||
embedding_jobs=[],
|
||||
prepare_seconds=prepare_seconds,
|
||||
)
|
||||
|
||||
existing_rows_result = await session.execute(
|
||||
text(
|
||||
"SELECT c.id, c.chunk_key, c.source_hash, c.entity_fingerprint, "
|
||||
"c.embedding_model, "
|
||||
"(e.chunk_id IS NOT NULL) AS has_embedding "
|
||||
"FROM search_vector_chunks c "
|
||||
"LEFT JOIN search_vector_embeddings e ON e.chunk_id = c.id "
|
||||
"WHERE c.project_id = :project_id AND c.entity_id = :entity_id"
|
||||
),
|
||||
{"project_id": self.project_id, "entity_id": entity_id},
|
||||
)
|
||||
existing_rows = existing_rows_result.mappings().all()
|
||||
existing_by_key = {str(row["chunk_key"]): row for row in existing_rows}
|
||||
existing_chunks_count = len(existing_by_key)
|
||||
incoming_chunk_keys = {record["chunk_key"] for record in chunk_records}
|
||||
|
||||
stale_ids = [
|
||||
int(row["id"])
|
||||
for chunk_key, row in existing_by_key.items()
|
||||
if chunk_key not in incoming_chunk_keys
|
||||
]
|
||||
stale_chunks_count = len(stale_ids)
|
||||
if stale_ids:
|
||||
await self._delete_stale_chunks(session, stale_ids, entity_id)
|
||||
|
||||
orphan_ids = {int(row["id"]) for row in existing_rows if not bool(row["has_embedding"])}
|
||||
orphan_chunks_count = len(orphan_ids)
|
||||
|
||||
skip_unchanged_entity = (
|
||||
existing_chunks_count == built_chunk_records_count
|
||||
and stale_chunks_count == 0
|
||||
and orphan_chunks_count == 0
|
||||
and existing_chunks_count > 0
|
||||
and all(
|
||||
row["entity_fingerprint"] == current_entity_fingerprint
|
||||
and row["embedding_model"] == current_embedding_model
|
||||
for row in existing_rows
|
||||
)
|
||||
)
|
||||
if skip_unchanged_entity:
|
||||
prepare_seconds = time.perf_counter() - sync_start
|
||||
return _PreparedEntityVectorSync(
|
||||
entity_id=entity_id,
|
||||
sync_start=sync_start,
|
||||
source_rows_count=source_rows_count,
|
||||
embedding_jobs=[],
|
||||
chunks_total=built_chunk_records_count,
|
||||
chunks_skipped=built_chunk_records_count,
|
||||
entity_skipped=True,
|
||||
prepare_seconds=prepare_seconds,
|
||||
)
|
||||
|
||||
pending_records: list[dict[str, str]] = []
|
||||
skipped_chunks_count = 0
|
||||
|
||||
for record in chunk_records:
|
||||
current = existing_by_key.get(record["chunk_key"])
|
||||
if current is None:
|
||||
pending_records.append(record)
|
||||
continue
|
||||
|
||||
row_id = int(current["id"])
|
||||
is_orphan = row_id in orphan_ids
|
||||
same_source_hash = current["source_hash"] == record["source_hash"]
|
||||
same_entity_fingerprint = (
|
||||
current["entity_fingerprint"] == current_entity_fingerprint
|
||||
)
|
||||
same_embedding_model = current["embedding_model"] == current_embedding_model
|
||||
|
||||
if same_source_hash and not is_orphan and same_embedding_model:
|
||||
if not same_entity_fingerprint:
|
||||
await session.execute(
|
||||
text(
|
||||
"UPDATE search_vector_chunks "
|
||||
"SET entity_fingerprint = :entity_fingerprint, "
|
||||
"embedding_model = :embedding_model, "
|
||||
"updated_at = NOW() "
|
||||
"WHERE id = :id"
|
||||
),
|
||||
{
|
||||
"id": row_id,
|
||||
"entity_fingerprint": current_entity_fingerprint,
|
||||
"embedding_model": current_embedding_model,
|
||||
},
|
||||
)
|
||||
skipped_chunks_count += 1
|
||||
continue
|
||||
|
||||
pending_records.append(record)
|
||||
|
||||
shard_plan = self._plan_entity_vector_shard(pending_records)
|
||||
self._log_vector_shard_plan(entity_id=entity_id, shard_plan=shard_plan)
|
||||
|
||||
scheduled_records = [
|
||||
record
|
||||
for record in sorted(pending_records, key=lambda record: record["chunk_key"])
|
||||
if record["chunk_key"] in shard_plan.scheduled_chunk_keys
|
||||
]
|
||||
|
||||
embedding_jobs: list[tuple[int, str]] = []
|
||||
upsert_records = list(scheduled_records)
|
||||
|
||||
if upsert_records:
|
||||
upsert_params: dict[str, object] = {
|
||||
"project_id": self.project_id,
|
||||
"entity_id": entity_id,
|
||||
}
|
||||
upsert_values: list[str] = []
|
||||
# The SQL template is built from integer enumerate() indices only.
|
||||
# No user-controlled text is interpolated into the statement.
|
||||
for index, record in enumerate(upsert_records):
|
||||
upsert_params[f"chunk_key_{index}"] = record["chunk_key"]
|
||||
upsert_params[f"chunk_text_{index}"] = record["chunk_text"]
|
||||
upsert_params[f"source_hash_{index}"] = record["source_hash"]
|
||||
upsert_params[f"entity_fingerprint_{index}"] = current_entity_fingerprint
|
||||
upsert_params[f"embedding_model_{index}"] = current_embedding_model
|
||||
upsert_values.append(
|
||||
"("
|
||||
":entity_id, :project_id, "
|
||||
f":chunk_key_{index}, :chunk_text_{index}, :source_hash_{index}, "
|
||||
f":entity_fingerprint_{index}, :embedding_model_{index}, NOW()"
|
||||
")"
|
||||
)
|
||||
|
||||
upsert_result = await session.execute(
|
||||
text(f"""
|
||||
INSERT INTO search_vector_chunks (
|
||||
entity_id,
|
||||
project_id,
|
||||
chunk_key,
|
||||
chunk_text,
|
||||
source_hash,
|
||||
entity_fingerprint,
|
||||
embedding_model,
|
||||
updated_at
|
||||
) VALUES {", ".join(upsert_values)}
|
||||
ON CONFLICT (project_id, entity_id, chunk_key) DO UPDATE SET
|
||||
chunk_text = EXCLUDED.chunk_text,
|
||||
source_hash = EXCLUDED.source_hash,
|
||||
entity_fingerprint = EXCLUDED.entity_fingerprint,
|
||||
embedding_model = EXCLUDED.embedding_model,
|
||||
updated_at = NOW()
|
||||
RETURNING id, chunk_key
|
||||
"""),
|
||||
upsert_params,
|
||||
)
|
||||
upserted_ids_by_key = {
|
||||
str(row["chunk_key"]): int(row["id"]) for row in upsert_result.mappings().all()
|
||||
}
|
||||
for record in upsert_records:
|
||||
row_id = upserted_ids_by_key[record["chunk_key"]]
|
||||
embedding_jobs.append((row_id, record["chunk_text"]))
|
||||
|
||||
prepare_seconds = time.perf_counter() - sync_start
|
||||
return _PreparedEntityVectorSync(
|
||||
entity_id=entity_id,
|
||||
sync_start=sync_start,
|
||||
source_rows_count=source_rows_count,
|
||||
embedding_jobs=embedding_jobs,
|
||||
chunks_total=built_chunk_records_count,
|
||||
chunks_skipped=skipped_chunks_count,
|
||||
entity_complete=shard_plan.entity_complete,
|
||||
oversized_entity=shard_plan.oversized_entity,
|
||||
pending_jobs_total=shard_plan.pending_jobs_total,
|
||||
shard_index=shard_plan.shard_index,
|
||||
shard_count=shard_plan.shard_count,
|
||||
remaining_jobs_after_shard=shard_plan.remaining_jobs_after_shard,
|
||||
prepare_seconds=prepare_seconds,
|
||||
)
|
||||
|
||||
async def _write_embeddings(
|
||||
self,
|
||||
session: AsyncSession,
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,11 +1,11 @@
|
||||
"""SQLite FTS5-based search repository implementation."""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import re
|
||||
from contextlib import asynccontextmanager
|
||||
from datetime import datetime
|
||||
from typing import List, Optional
|
||||
|
||||
import asyncio
|
||||
from loguru import logger
|
||||
from sqlalchemy import text
|
||||
from sqlalchemy.exc import OperationalError as SAOperationalError
|
||||
@@ -56,7 +56,8 @@ class SQLiteSearchRepository(SearchRepositoryBase):
|
||||
self._app_config.semantic_embedding_sync_batch_size
|
||||
)
|
||||
self._embedding_provider = embedding_provider
|
||||
self._sqlite_vec_lock = asyncio.Lock()
|
||||
self._sqlite_vec_load_lock = asyncio.Lock()
|
||||
self._sqlite_prepare_write_lock = asyncio.Lock()
|
||||
self._vector_tables_initialized = False
|
||||
self._vector_dimensions = 384
|
||||
|
||||
@@ -357,7 +358,13 @@ class SQLiteSearchRepository(SearchRepositoryBase):
|
||||
"pip install -U basic-memory"
|
||||
) from exc
|
||||
|
||||
async with self._sqlite_vec_lock:
|
||||
# Trigger: sqlite-vec must be loaded on each SQLite connection before
|
||||
# vec tables and functions are visible.
|
||||
# Why: extension loading is connection-local, so we need one narrow
|
||||
# critical section to avoid racing two coroutines on the same step.
|
||||
# Outcome: connection setup stays serialized without blocking unrelated
|
||||
# prepare work behind the write-side lock.
|
||||
async with self._sqlite_vec_load_lock:
|
||||
try:
|
||||
await session.execute(text("SELECT vec_version()"))
|
||||
return
|
||||
@@ -558,6 +565,76 @@ class SQLiteSearchRepository(SearchRepositoryBase):
|
||||
stale_params,
|
||||
)
|
||||
|
||||
async def delete_entity_vector_rows(self, entity_id: int) -> None:
|
||||
"""Delete one entity's vec rows on a sqlite-vec-enabled connection."""
|
||||
await self._ensure_vector_tables()
|
||||
|
||||
async with db.scoped_session(self.session_maker) as session:
|
||||
await self._ensure_sqlite_vec_loaded(session)
|
||||
|
||||
# Constraint: sqlite-vec virtual tables are only visible after vec0 is
|
||||
# loaded on this exact connection.
|
||||
# Why: generic repository sessions can reach search_vector_chunks but still
|
||||
# fail with "no such module: vec0" when touching embeddings.
|
||||
# Outcome: service-level cleanup routes vec-table deletes through this helper.
|
||||
await self._delete_entity_chunks(session, entity_id)
|
||||
await session.commit()
|
||||
|
||||
async def delete_project_vector_rows(self) -> None:
|
||||
"""Delete all vector rows for this project on a sqlite-vec-enabled connection."""
|
||||
await self._ensure_vector_tables()
|
||||
|
||||
async with db.scoped_session(self.session_maker) as session:
|
||||
await self._ensure_sqlite_vec_loaded(session)
|
||||
|
||||
# Constraint: sqlite-vec stores embeddings separately with no cascade delete.
|
||||
# Why: full rebuild must clear embeddings before chunk rows or stale vectors remain.
|
||||
# Outcome: the next sync recreates the project's derived vectors from scratch.
|
||||
await session.execute(
|
||||
text(
|
||||
"DELETE FROM search_vector_embeddings WHERE rowid IN ("
|
||||
"SELECT id FROM search_vector_chunks WHERE project_id = :project_id)"
|
||||
),
|
||||
{"project_id": self.project_id},
|
||||
)
|
||||
await session.execute(
|
||||
text("DELETE FROM search_vector_chunks WHERE project_id = :project_id"),
|
||||
{"project_id": self.project_id},
|
||||
)
|
||||
await session.commit()
|
||||
|
||||
async def delete_stale_vector_rows(self) -> None:
|
||||
"""Delete vector rows whose source entities no longer exist."""
|
||||
await self._ensure_vector_tables()
|
||||
|
||||
async with db.scoped_session(self.session_maker) as session:
|
||||
await self._ensure_sqlite_vec_loaded(session)
|
||||
|
||||
stale_entity_filter = (
|
||||
"entity_id NOT IN (SELECT id FROM entity WHERE project_id = :project_id)"
|
||||
)
|
||||
params = {"project_id": self.project_id}
|
||||
|
||||
# Trigger: deleted entities left behind derived vector rows.
|
||||
# Why: sqlite-vec does not provide cascade cleanup from our chunk table.
|
||||
# Outcome: stale vector state disappears before coverage stats or reindex runs.
|
||||
await session.execute(
|
||||
text(
|
||||
"DELETE FROM search_vector_embeddings WHERE rowid IN ("
|
||||
"SELECT id FROM search_vector_chunks "
|
||||
f"WHERE project_id = :project_id AND {stale_entity_filter})"
|
||||
),
|
||||
params,
|
||||
)
|
||||
await session.execute(
|
||||
text(
|
||||
"DELETE FROM search_vector_chunks "
|
||||
f"WHERE project_id = :project_id AND {stale_entity_filter}"
|
||||
),
|
||||
params,
|
||||
)
|
||||
await session.commit()
|
||||
|
||||
def _distance_to_similarity(self, distance: float) -> float:
|
||||
"""Convert L2 distance to cosine similarity for normalized embeddings.
|
||||
|
||||
@@ -566,13 +643,26 @@ class SQLiteSearchRepository(SearchRepositoryBase):
|
||||
"""
|
||||
return max(0.0, 1.0 - (distance * distance) / 2.0)
|
||||
|
||||
def _orphan_detection_sql(self) -> str:
|
||||
"""SQLite sqlite-vec uses rowid-based embedding table."""
|
||||
@asynccontextmanager
|
||||
async def _prepare_entity_write_scope(self):
|
||||
"""SQLite keeps the shared read window, but funnels prepare writes through one lock."""
|
||||
# Trigger: the shared prepare window fans out per entity after batched reads.
|
||||
# Why: SQLite still benefits from shared reads, but write transactions do
|
||||
# not get meaningfully faster when we open many at once.
|
||||
# Outcome: one entity at a time mutates chunk rows, while vec extension
|
||||
# loading uses its own separate lock and cannot deadlock this path.
|
||||
async with self._sqlite_prepare_write_lock:
|
||||
yield
|
||||
|
||||
def _prepare_window_existing_rows_sql(self, placeholders: str) -> str:
|
||||
"""SQLite sqlite-vec stores embeddings by rowid rather than chunk_id."""
|
||||
return (
|
||||
"SELECT c.id FROM search_vector_chunks c "
|
||||
"SELECT c.entity_id, c.id, c.chunk_key, c.source_hash, c.entity_fingerprint, "
|
||||
"c.embedding_model, (e.rowid IS NOT NULL) AS has_embedding "
|
||||
"FROM search_vector_chunks c "
|
||||
"LEFT JOIN search_vector_embeddings e ON e.rowid = c.id "
|
||||
"WHERE c.project_id = :project_id AND c.entity_id = :entity_id "
|
||||
"AND e.rowid IS NULL"
|
||||
f"WHERE c.project_id = :project_id AND c.entity_id IN ({placeholders}) "
|
||||
"ORDER BY c.entity_id ASC, c.chunk_key ASC"
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@@ -208,15 +208,20 @@ class FileService:
|
||||
|
||||
await file_utils.write_file_atomic(full_path, content)
|
||||
|
||||
final_content = content
|
||||
if self.app_config:
|
||||
formatted_content = await file_utils.format_file(
|
||||
full_path, self.app_config, is_markdown=self.is_markdown(path)
|
||||
)
|
||||
if formatted_content is not None:
|
||||
final_content = formatted_content # pragma: no cover
|
||||
pass # pragma: no cover
|
||||
|
||||
checksum = await file_utils.compute_checksum(final_content)
|
||||
# Trigger: formatters and platform-specific text writers can change the
|
||||
# persisted bytes even when the logical content string is the same.
|
||||
# Why: sync and move detection compare against on-disk checksums, not
|
||||
# the pre-write Python string.
|
||||
# Outcome: return the checksum of the actual stored file so callers do
|
||||
# not record a hash that immediately disagrees with the file.
|
||||
checksum = await self.compute_checksum(full_path)
|
||||
logger.debug(f"File write completed path={full_path}, {checksum=}")
|
||||
return checksum
|
||||
|
||||
@@ -345,7 +350,13 @@ class FileService:
|
||||
async with aiofiles.open(full_path, mode="r", encoding="utf-8") as f:
|
||||
content = await f.read()
|
||||
|
||||
checksum = await file_utils.compute_checksum(content)
|
||||
# Trigger: text-mode reads normalize line endings on Windows, so the
|
||||
# decoded string can differ from the bytes we just wrote.
|
||||
# Why: write_file/update_frontmatter now return the checksum of the
|
||||
# persisted file, and read_file should report the same authority.
|
||||
# Outcome: callers get human-readable content plus the checksum for the
|
||||
# exact bytes stored on disk.
|
||||
checksum = await self.compute_checksum(full_path)
|
||||
|
||||
logger.debug(
|
||||
"File read completed",
|
||||
@@ -478,8 +489,12 @@ class FileService:
|
||||
if formatted_content is not None:
|
||||
content_for_checksum = formatted_content # pragma: no cover
|
||||
|
||||
# Trigger: frontmatter normalization may persist bytes that differ from the
|
||||
# in-memory string because of formatter output or platform newline handling.
|
||||
# Why: follow-up scans and checksum-based move detection read raw bytes from disk.
|
||||
# Outcome: the returned checksum always matches the file that was just written.
|
||||
return FrontmatterUpdateResult(
|
||||
checksum=await file_utils.compute_checksum(content_for_checksum),
|
||||
checksum=await self.compute_checksum(full_path),
|
||||
content=content_for_checksum,
|
||||
)
|
||||
|
||||
|
||||
@@ -521,7 +521,9 @@ class SearchService:
|
||||
chunks_total=sum(result.chunks_total for result in repository_results),
|
||||
chunks_skipped=sum(result.chunks_skipped for result in repository_results),
|
||||
embedding_jobs_total=sum(result.embedding_jobs_total for result in repository_results),
|
||||
prepare_seconds_total=sum(result.prepare_seconds_total for result in repository_results),
|
||||
prepare_seconds_total=sum(
|
||||
result.prepare_seconds_total for result in repository_results
|
||||
),
|
||||
queue_wait_seconds_total=sum(
|
||||
result.queue_wait_seconds_total for result in repository_results
|
||||
),
|
||||
@@ -530,11 +532,14 @@ class SearchService:
|
||||
)
|
||||
return batch_result
|
||||
|
||||
async def reindex_vectors(self, progress_callback=None) -> dict:
|
||||
async def reindex_vectors(self, progress_callback=None, force_full: bool = False) -> dict:
|
||||
"""Rebuild vector embeddings for all entities.
|
||||
|
||||
Args:
|
||||
progress_callback: Optional callable(entity_id, index, total) for progress reporting.
|
||||
progress_callback: Optional callable(entity_id, completed, total) for progress
|
||||
reporting when an entity reaches a terminal state in this run.
|
||||
force_full: When True, clear this project's derived vectors first so every
|
||||
eligible entity re-embeds from scratch.
|
||||
|
||||
Returns:
|
||||
dict with stats: total_entities, embedded, skipped, errors
|
||||
@@ -545,6 +550,8 @@ class SearchService:
|
||||
# Clean up stale rows in search_index and search_vector_chunks
|
||||
# that reference entity_ids no longer in the entity table
|
||||
await self._purge_stale_search_rows()
|
||||
if force_full:
|
||||
await self._clear_project_vectors_for_full_reindex()
|
||||
|
||||
batch_result = await self.sync_entity_vectors_batch(
|
||||
entity_ids,
|
||||
@@ -562,6 +569,31 @@ class SearchService:
|
||||
|
||||
return stats
|
||||
|
||||
async def _clear_project_vectors_for_full_reindex(self) -> None:
|
||||
"""Remove this project's derived vectors so a full reindex re-embeds everything.
|
||||
|
||||
Trigger: the operator asked for a full embedding rebuild rather than the
|
||||
default incremental vector sync.
|
||||
Why: the repository sync path intentionally skips unchanged entities, so
|
||||
we need to clear the derived vector state first to force fresh embeddings.
|
||||
Outcome: the next batch sync recreates every eligible entity's vectors.
|
||||
"""
|
||||
from basic_memory.repository.sqlite_search_repository import SQLiteSearchRepository
|
||||
|
||||
project_id = self.repository.project_id
|
||||
params = {"project_id": project_id}
|
||||
|
||||
# Constraint: sqlite-vec stores embeddings in a separate rowid table with
|
||||
# no cascade delete, so embeddings must be removed before chunk rows.
|
||||
if isinstance(self.repository, SQLiteSearchRepository):
|
||||
await self.repository.delete_project_vector_rows()
|
||||
else:
|
||||
await self.repository.execute_query(
|
||||
text("DELETE FROM search_vector_chunks WHERE project_id = :project_id"),
|
||||
params,
|
||||
)
|
||||
logger.info("Cleared project vectors for full reindex", project_id=project_id)
|
||||
|
||||
async def _purge_stale_search_rows(self) -> None:
|
||||
"""Remove rows from search_index and search_vector_chunks for deleted entities.
|
||||
|
||||
@@ -588,24 +620,17 @@ class SearchService:
|
||||
|
||||
# SQLite vec has no CASCADE — must delete embeddings before chunks
|
||||
if isinstance(self.repository, SQLiteSearchRepository):
|
||||
await self.repository.delete_stale_vector_rows()
|
||||
else:
|
||||
# Postgres CASCADE handles embedding deletion automatically
|
||||
await self.repository.execute_query(
|
||||
text(
|
||||
"DELETE FROM search_vector_embeddings WHERE rowid IN ("
|
||||
"SELECT id FROM search_vector_chunks "
|
||||
f"WHERE project_id = :project_id AND {stale_entity_filter})"
|
||||
f"DELETE FROM search_vector_chunks "
|
||||
f"WHERE project_id = :project_id AND {stale_entity_filter}"
|
||||
),
|
||||
params,
|
||||
)
|
||||
|
||||
# Postgres CASCADE handles embedding deletion automatically
|
||||
await self.repository.execute_query(
|
||||
text(
|
||||
f"DELETE FROM search_vector_chunks "
|
||||
f"WHERE project_id = :project_id AND {stale_entity_filter}"
|
||||
),
|
||||
params,
|
||||
)
|
||||
|
||||
logger.info("Purged stale search rows for deleted entities", project_id=project_id)
|
||||
|
||||
@staticmethod
|
||||
@@ -640,28 +665,24 @@ class SearchService:
|
||||
# Trigger: semantic indexing is disabled for this repository instance.
|
||||
# Why: repositories only create vector tables when semantic search is enabled.
|
||||
# Outcome: skip cleanup because there are no active derived vector rows to maintain.
|
||||
if isinstance(self.repository, SearchRepositoryBase) and not self.repository._semantic_enabled:
|
||||
if (
|
||||
isinstance(self.repository, SearchRepositoryBase)
|
||||
and not self.repository._semantic_enabled
|
||||
):
|
||||
return
|
||||
|
||||
params = {"project_id": self.repository.project_id, "entity_id": entity_id}
|
||||
if isinstance(self.repository, SQLiteSearchRepository):
|
||||
await self.repository.delete_entity_vector_rows(entity_id)
|
||||
else:
|
||||
await self.repository.execute_query(
|
||||
text(
|
||||
"DELETE FROM search_vector_embeddings WHERE rowid IN ("
|
||||
"SELECT id FROM search_vector_chunks "
|
||||
"WHERE project_id = :project_id AND entity_id = :entity_id)"
|
||||
"DELETE FROM search_vector_chunks "
|
||||
"WHERE project_id = :project_id AND entity_id = :entity_id"
|
||||
),
|
||||
params,
|
||||
)
|
||||
|
||||
await self.repository.execute_query(
|
||||
text(
|
||||
"DELETE FROM search_vector_chunks "
|
||||
"WHERE project_id = :project_id AND entity_id = :entity_id"
|
||||
),
|
||||
params,
|
||||
)
|
||||
|
||||
async def index_entity_file(
|
||||
self,
|
||||
entity: Entity,
|
||||
|
||||
@@ -104,6 +104,7 @@ class SyncReport:
|
||||
deleted: Set[str] = field(default_factory=set)
|
||||
moves: Dict[str, str] = field(default_factory=dict) # old_path -> new_path
|
||||
checksums: Dict[str, str] = field(default_factory=dict) # path -> checksum
|
||||
scanned_paths: Set[str] = field(default_factory=set)
|
||||
skipped_files: List[SkippedFile] = field(default_factory=list)
|
||||
|
||||
@property
|
||||
@@ -292,6 +293,7 @@ class SyncService:
|
||||
directory: Path,
|
||||
project_name: Optional[str] = None,
|
||||
force_full: bool = False,
|
||||
sync_embeddings: bool = True,
|
||||
progress_callback: Callable[[IndexProgress], Awaitable[None]] | None = None,
|
||||
) -> SyncReport:
|
||||
"""Sync all files with database and update scan watermark.
|
||||
@@ -300,6 +302,7 @@ class SyncService:
|
||||
directory: Directory to sync
|
||||
project_name: Optional project name
|
||||
force_full: If True, force a full scan bypassing watermark optimization
|
||||
sync_embeddings: If True, generate vectors for entities indexed during this sync
|
||||
progress_callback: Optional callback for typed indexing progress updates
|
||||
"""
|
||||
|
||||
@@ -348,7 +351,16 @@ class SyncService:
|
||||
for path in report.deleted:
|
||||
await self.handle_delete(path)
|
||||
|
||||
changed_paths = sorted(report.new | report.modified)
|
||||
# Trigger: the caller requested a full reindex pass through the sync path.
|
||||
# Why: cloud-style "full" semantics should rebuild every current file-backed
|
||||
# search row, not only the files that differ from the last watermark.
|
||||
# Outcome: progress reflects the whole project and unchanged files are
|
||||
# re-indexed without inflating the change report itself.
|
||||
changed_paths = (
|
||||
sorted(report.scanned_paths)
|
||||
if force_full
|
||||
else sorted(report.new | report.modified)
|
||||
)
|
||||
indexed_entities, skipped_files = await self._index_changed_files(
|
||||
changed_paths,
|
||||
report.checksums,
|
||||
@@ -357,9 +369,12 @@ class SyncService:
|
||||
report.skipped_files.extend(skipped_files)
|
||||
synced_entity_ids = [indexed.entity_id for indexed in indexed_entities]
|
||||
|
||||
# Only resolve relations if there were actual changes
|
||||
# If no files changed, no new unresolved relations could have been created
|
||||
if report.total > 0:
|
||||
# Trigger: either the filesystem diff found changes, or the caller forced a
|
||||
# full reindex and we just reprocessed the current files.
|
||||
# Why: relation resolution should follow the file-processing work that just ran,
|
||||
# not only the lightweight diff summary.
|
||||
# Outcome: full reindex can heal relation state even when the diff report is empty.
|
||||
if report.total > 0 or (force_full and indexed_entities):
|
||||
with telemetry.scope(
|
||||
"sync.project.resolve_relations", relation_scope="all_pending"
|
||||
):
|
||||
@@ -369,7 +384,7 @@ class SyncService:
|
||||
|
||||
# Batch-generate vector embeddings for all synced entities
|
||||
synced_entity_ids = list(dict.fromkeys(synced_entity_ids))
|
||||
if synced_entity_ids and self.app_config.semantic_search_enabled:
|
||||
if synced_entity_ids and sync_embeddings and self.app_config.semantic_search_enabled:
|
||||
try:
|
||||
with telemetry.scope(
|
||||
"sync.project.sync_embeddings",
|
||||
@@ -906,6 +921,7 @@ class SyncService:
|
||||
|
||||
# Store checksums for files that need syncing
|
||||
report.checksums = changed_checksums
|
||||
report.scanned_paths = scanned_paths
|
||||
|
||||
scan_duration_ms = int((time.time() - scan_start_time) * 1000)
|
||||
|
||||
|
||||
@@ -40,6 +40,7 @@ class TelemetryState:
|
||||
|
||||
_STATE = TelemetryState()
|
||||
_LOGFIRE_HANDLER: dict[str, Any] | None = None
|
||||
_METRICS: dict[tuple[str, str, str, str], Any] = {}
|
||||
|
||||
|
||||
def reset_telemetry_state() -> None:
|
||||
@@ -55,6 +56,7 @@ def reset_telemetry_state() -> None:
|
||||
_STATE.send_to_logfire = False
|
||||
_STATE.warnings.clear()
|
||||
_LOGFIRE_HANDLER = None
|
||||
_METRICS.clear()
|
||||
|
||||
|
||||
def _filter_attributes(attrs: dict[str, Any]) -> dict[str, Any]:
|
||||
@@ -136,6 +138,58 @@ def pop_telemetry_warnings() -> list[str]:
|
||||
return warnings
|
||||
|
||||
|
||||
def _get_metric(metric_type: str, name: str, *, unit: str, description: str) -> Any | None:
|
||||
"""Create or reuse a Logfire metric instrument when telemetry is enabled."""
|
||||
logfire = _load_logfire()
|
||||
if logfire is None or not _STATE.configured: # pragma: no cover
|
||||
return None # pragma: no cover
|
||||
|
||||
metric_key = (metric_type, name, unit, description)
|
||||
cached_metric = _METRICS.get(metric_key)
|
||||
if cached_metric is not None:
|
||||
return cached_metric
|
||||
|
||||
if metric_type == "counter":
|
||||
metric = logfire.metric_counter(name, unit=unit, description=description)
|
||||
elif metric_type == "histogram":
|
||||
metric = logfire.metric_histogram(name, unit=unit, description=description)
|
||||
else: # pragma: no cover
|
||||
raise ValueError(f"Unsupported metric type: {metric_type}") # pragma: no cover
|
||||
|
||||
_METRICS[metric_key] = metric
|
||||
return metric
|
||||
|
||||
|
||||
def add_counter(
|
||||
name: str,
|
||||
amount: int | float,
|
||||
*,
|
||||
unit: str = "1",
|
||||
description: str = "",
|
||||
**attrs: Any,
|
||||
) -> None:
|
||||
"""Record a counter increment when telemetry is enabled."""
|
||||
metric = _get_metric("counter", name, unit=unit, description=description)
|
||||
if metric is None:
|
||||
return
|
||||
metric.add(amount, attributes=_filter_attributes(attrs))
|
||||
|
||||
|
||||
def record_histogram(
|
||||
name: str,
|
||||
amount: int | float,
|
||||
*,
|
||||
unit: str = "",
|
||||
description: str = "",
|
||||
**attrs: Any,
|
||||
) -> None:
|
||||
"""Record one histogram sample when telemetry is enabled."""
|
||||
metric = _get_metric("histogram", name, unit=unit, description=description)
|
||||
if metric is None:
|
||||
return
|
||||
metric.record(amount, attributes=_filter_attributes(attrs))
|
||||
|
||||
|
||||
@contextmanager
|
||||
def contextualize(**attrs: Any) -> Iterator[None]:
|
||||
"""Apply filtered telemetry attributes to Loguru calls in this scope."""
|
||||
@@ -176,11 +230,13 @@ def started_span(name: str, **attrs: Any) -> Iterator[Any | None]:
|
||||
|
||||
|
||||
__all__ = [
|
||||
"add_counter",
|
||||
"contextualize",
|
||||
"configure_telemetry",
|
||||
"get_logfire_handler",
|
||||
"operation",
|
||||
"pop_telemetry_warnings",
|
||||
"record_histogram",
|
||||
"reset_telemetry_state",
|
||||
"scope",
|
||||
"span",
|
||||
|
||||
@@ -0,0 +1,322 @@
|
||||
"""Tests for `bm reindex` CLI wiring."""
|
||||
|
||||
import asyncio
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
from typer.testing import CliRunner
|
||||
|
||||
from basic_memory.cli.app import app
|
||||
import basic_memory.cli.commands.db as db_cmd # noqa: F401
|
||||
|
||||
|
||||
runner = CliRunner()
|
||||
|
||||
|
||||
def _stub_app_config(*, semantic_search_enabled: bool = True) -> SimpleNamespace:
|
||||
"""Build the minimal config surface the CLI reindex path expects."""
|
||||
return SimpleNamespace(
|
||||
semantic_search_enabled=semantic_search_enabled,
|
||||
database_path=Path("/tmp/basic-memory.db"),
|
||||
get_project_mode=lambda project_name: None,
|
||||
)
|
||||
|
||||
|
||||
def _configure_reindex_cli(monkeypatch, app_config: SimpleNamespace) -> None:
|
||||
"""Keep CLI tests focused on reindex wiring instead of full app startup."""
|
||||
monkeypatch.setattr("basic_memory.cli.app.init_cli_logging", lambda: None)
|
||||
monkeypatch.setattr("basic_memory.cli.app.maybe_show_init_line", lambda *_args: None)
|
||||
monkeypatch.setattr("basic_memory.cli.app.maybe_show_cloud_promo", lambda *_args: None)
|
||||
monkeypatch.setattr("basic_memory.cli.app.maybe_run_periodic_auto_update", lambda *_args: None)
|
||||
monkeypatch.setattr(
|
||||
"basic_memory.cli.app.CliContainer.create",
|
||||
lambda: SimpleNamespace(config=app_config, mode=SimpleNamespace(is_cloud=False)),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
db_cmd,
|
||||
"ConfigManager",
|
||||
lambda: SimpleNamespace(config=app_config),
|
||||
)
|
||||
|
||||
|
||||
def test_reindex_defaults_to_incremental_search_and_embeddings(monkeypatch):
|
||||
app_config = _stub_app_config()
|
||||
_configure_reindex_cli(monkeypatch, app_config)
|
||||
captured: dict[str, object] = {}
|
||||
|
||||
async def _stub_reindex(app_config, *, search: bool, embeddings: bool, full: bool, project):
|
||||
captured.update(
|
||||
{
|
||||
"app_config": app_config,
|
||||
"search": search,
|
||||
"embeddings": embeddings,
|
||||
"full": full,
|
||||
"project": project,
|
||||
}
|
||||
)
|
||||
|
||||
monkeypatch.setattr(db_cmd, "_reindex", _stub_reindex)
|
||||
monkeypatch.setattr(db_cmd, "run_with_cleanup", lambda coro: asyncio.run(coro))
|
||||
|
||||
result = runner.invoke(app, ["reindex"])
|
||||
|
||||
assert result.exit_code == 0
|
||||
assert captured == {
|
||||
"app_config": app_config,
|
||||
"search": True,
|
||||
"embeddings": True,
|
||||
"full": False,
|
||||
"project": None,
|
||||
}
|
||||
|
||||
|
||||
def test_reindex_full_runs_full_search_and_embeddings(monkeypatch):
|
||||
app_config = _stub_app_config()
|
||||
_configure_reindex_cli(monkeypatch, app_config)
|
||||
captured: dict[str, object] = {}
|
||||
|
||||
async def _stub_reindex(app_config, *, search: bool, embeddings: bool, full: bool, project):
|
||||
captured.update(
|
||||
{
|
||||
"search": search,
|
||||
"embeddings": embeddings,
|
||||
"full": full,
|
||||
"project": project,
|
||||
}
|
||||
)
|
||||
|
||||
monkeypatch.setattr(db_cmd, "_reindex", _stub_reindex)
|
||||
monkeypatch.setattr(db_cmd, "run_with_cleanup", lambda coro: asyncio.run(coro))
|
||||
|
||||
result = runner.invoke(app, ["reindex", "--full"])
|
||||
|
||||
assert result.exit_code == 0
|
||||
assert captured == {
|
||||
"search": True,
|
||||
"embeddings": True,
|
||||
"full": True,
|
||||
"project": None,
|
||||
}
|
||||
|
||||
|
||||
def test_reindex_full_search_runs_search_only(monkeypatch):
|
||||
app_config = _stub_app_config()
|
||||
_configure_reindex_cli(monkeypatch, app_config)
|
||||
captured: dict[str, object] = {}
|
||||
|
||||
async def _stub_reindex(app_config, *, search: bool, embeddings: bool, full: bool, project):
|
||||
captured.update(
|
||||
{
|
||||
"search": search,
|
||||
"embeddings": embeddings,
|
||||
"full": full,
|
||||
"project": project,
|
||||
}
|
||||
)
|
||||
|
||||
monkeypatch.setattr(db_cmd, "_reindex", _stub_reindex)
|
||||
monkeypatch.setattr(db_cmd, "run_with_cleanup", lambda coro: asyncio.run(coro))
|
||||
|
||||
result = runner.invoke(app, ["reindex", "--full", "--search"])
|
||||
|
||||
assert result.exit_code == 0
|
||||
assert captured == {
|
||||
"search": True,
|
||||
"embeddings": False,
|
||||
"full": True,
|
||||
"project": None,
|
||||
}
|
||||
|
||||
|
||||
def test_reindex_embeddings_only_preserves_incremental_default(monkeypatch):
|
||||
app_config = _stub_app_config()
|
||||
_configure_reindex_cli(monkeypatch, app_config)
|
||||
captured: dict[str, object] = {}
|
||||
|
||||
async def _stub_reindex(app_config, *, search: bool, embeddings: bool, full: bool, project):
|
||||
captured.update(
|
||||
{
|
||||
"search": search,
|
||||
"embeddings": embeddings,
|
||||
"full": full,
|
||||
"project": project,
|
||||
}
|
||||
)
|
||||
|
||||
monkeypatch.setattr(db_cmd, "_reindex", _stub_reindex)
|
||||
monkeypatch.setattr(db_cmd, "run_with_cleanup", lambda coro: asyncio.run(coro))
|
||||
|
||||
result = runner.invoke(app, ["reindex", "--embeddings"])
|
||||
|
||||
assert result.exit_code == 0
|
||||
assert captured == {
|
||||
"search": False,
|
||||
"embeddings": True,
|
||||
"full": False,
|
||||
"project": None,
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reindex_project_full_passes_force_full_to_sync_and_reports_mode(monkeypatch):
|
||||
app_config = _stub_app_config()
|
||||
project = SimpleNamespace(id=1, name="foo", path="/tmp/foo")
|
||||
session_maker = object()
|
||||
sync_service = SimpleNamespace(sync=AsyncMock())
|
||||
printed_lines: list[str] = []
|
||||
|
||||
class StubProjectRepository:
|
||||
def __init__(self, _session_maker):
|
||||
self._session_maker = _session_maker
|
||||
|
||||
async def get_active_projects(self):
|
||||
return [project]
|
||||
|
||||
class SilentProgress:
|
||||
def __init__(self, *args, **kwargs):
|
||||
self.tasks: dict[int, SimpleNamespace] = {}
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc, tb):
|
||||
return False
|
||||
|
||||
def add_task(self, description, total=1):
|
||||
self.tasks[1] = SimpleNamespace(total=total, description=description)
|
||||
return 1
|
||||
|
||||
def update(self, task_id, **kwargs):
|
||||
if "total" in kwargs:
|
||||
self.tasks[task_id].total = kwargs["total"]
|
||||
|
||||
monkeypatch.setattr(db_cmd, "reconcile_projects_with_config", AsyncMock())
|
||||
monkeypatch.setattr(
|
||||
db_cmd.db,
|
||||
"get_or_create_db",
|
||||
AsyncMock(return_value=(None, session_maker)),
|
||||
)
|
||||
monkeypatch.setattr(db_cmd.db, "shutdown_db", AsyncMock())
|
||||
monkeypatch.setattr(db_cmd, "ProjectRepository", StubProjectRepository)
|
||||
monkeypatch.setattr(db_cmd, "get_sync_service", AsyncMock(return_value=sync_service))
|
||||
monkeypatch.setattr(db_cmd, "Progress", SilentProgress)
|
||||
monkeypatch.setattr(
|
||||
db_cmd.console,
|
||||
"print",
|
||||
lambda message="", *args, **kwargs: printed_lines.append(str(message)),
|
||||
)
|
||||
|
||||
await db_cmd._reindex(
|
||||
app_config,
|
||||
search=True,
|
||||
embeddings=False,
|
||||
full=True,
|
||||
project="foo",
|
||||
)
|
||||
|
||||
sync_service.sync.assert_awaited_once()
|
||||
sync_call = sync_service.sync.await_args
|
||||
assert sync_call.args[0] == Path("/tmp/foo")
|
||||
assert sync_call.kwargs["project_name"] == "foo"
|
||||
assert sync_call.kwargs["force_full"] is True
|
||||
assert sync_call.kwargs["sync_embeddings"] is False
|
||||
assert callable(sync_call.kwargs["progress_callback"])
|
||||
assert any("full scan" in line for line in printed_lines)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reindex_embeddings_only_full_passes_force_full_to_vector_reindex(monkeypatch):
|
||||
app_config = _stub_app_config()
|
||||
project = SimpleNamespace(id=1, name="foo", path="/tmp/foo")
|
||||
session_maker = object()
|
||||
printed_lines: list[str] = []
|
||||
vector_reindex_calls: list[dict[str, object]] = []
|
||||
|
||||
class StubProjectRepository:
|
||||
def __init__(self, _session_maker):
|
||||
self._session_maker = _session_maker
|
||||
|
||||
async def get_active_projects(self):
|
||||
return [project]
|
||||
|
||||
class StubSearchService:
|
||||
def __init__(self, search_repository, entity_repository, file_service):
|
||||
self.search_repository = search_repository
|
||||
self.entity_repository = entity_repository
|
||||
self.file_service = file_service
|
||||
|
||||
async def reindex_vectors(self, *, progress_callback=None, force_full: bool = False):
|
||||
vector_reindex_calls.append(
|
||||
{
|
||||
"progress_callback": progress_callback,
|
||||
"force_full": force_full,
|
||||
}
|
||||
)
|
||||
return {"total_entities": 2, "embedded": 2, "skipped": 0, "errors": 0}
|
||||
|
||||
class SilentProgress:
|
||||
def __init__(self, *args, **kwargs):
|
||||
self.tasks: dict[int, SimpleNamespace] = {}
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc, tb):
|
||||
return False
|
||||
|
||||
def add_task(self, description, total=None):
|
||||
self.tasks[1] = SimpleNamespace(total=total, description=description)
|
||||
return 1
|
||||
|
||||
def update(self, task_id, **kwargs):
|
||||
if "total" in kwargs:
|
||||
self.tasks[task_id].total = kwargs["total"]
|
||||
|
||||
monkeypatch.setattr(db_cmd, "reconcile_projects_with_config", AsyncMock())
|
||||
monkeypatch.setattr(
|
||||
db_cmd.db,
|
||||
"get_or_create_db",
|
||||
AsyncMock(return_value=(None, session_maker)),
|
||||
)
|
||||
monkeypatch.setattr(db_cmd.db, "shutdown_db", AsyncMock())
|
||||
monkeypatch.setattr(db_cmd, "ProjectRepository", StubProjectRepository)
|
||||
monkeypatch.setattr(
|
||||
"basic_memory.repository.search_repository.create_search_repository",
|
||||
lambda *args, **kwargs: object(),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"basic_memory.repository.EntityRepository", lambda *args, **kwargs: object()
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"basic_memory.markdown.entity_parser.EntityParser",
|
||||
lambda *args, **kwargs: object(),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"basic_memory.markdown.markdown_processor.MarkdownProcessor",
|
||||
lambda *args, **kwargs: object(),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"basic_memory.services.file_service.FileService", lambda *args, **kwargs: object()
|
||||
)
|
||||
monkeypatch.setattr("basic_memory.services.search_service.SearchService", StubSearchService)
|
||||
monkeypatch.setattr(db_cmd, "Progress", SilentProgress)
|
||||
monkeypatch.setattr(
|
||||
db_cmd.console,
|
||||
"print",
|
||||
lambda message="", *args, **kwargs: printed_lines.append(str(message)),
|
||||
)
|
||||
|
||||
await db_cmd._reindex(
|
||||
app_config,
|
||||
search=False,
|
||||
embeddings=True,
|
||||
full=True,
|
||||
project="foo",
|
||||
)
|
||||
|
||||
assert len(vector_reindex_calls) == 1
|
||||
assert vector_reindex_calls[0]["force_full"] is True
|
||||
assert callable(vector_reindex_calls[0]["progress_callback"])
|
||||
assert any("full rebuild" in line for line in printed_lines)
|
||||
@@ -265,9 +265,15 @@ async def test_batch_indexer_returns_original_markdown_content_when_no_frontmatt
|
||||
parse_max_concurrent=1,
|
||||
)
|
||||
|
||||
# Trigger: Windows persists CRLF for text writes even when the test literal uses LF.
|
||||
# Why: this assertion cares about "no rewrite happened", not about forcing one newline
|
||||
# convention across platforms.
|
||||
# Outcome: compare against the exact markdown text stored on disk for this file.
|
||||
persisted_content = (project_config.home / path).read_bytes().decode("utf-8")
|
||||
|
||||
assert result.errors == []
|
||||
assert len(result.indexed) == 1
|
||||
assert result.indexed[0].markdown_content == original_content
|
||||
assert result.indexed[0].markdown_content == persisted_content
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -514,9 +520,15 @@ async def test_batch_indexer_uses_parsed_markdown_body_for_malformed_frontmatter
|
||||
parse_max_concurrent=1,
|
||||
)
|
||||
|
||||
# Trigger: malformed frontmatter should pass through without normalization.
|
||||
# Why: Windows can still surface that unchanged file with CRLF line endings.
|
||||
# Outcome: compare the indexed markdown to the persisted file content, not the LF
|
||||
# test literal used to create it.
|
||||
persisted_content = (project_config.home / path).read_bytes().decode("utf-8")
|
||||
|
||||
assert result.errors == []
|
||||
assert len(result.indexed) == 1
|
||||
assert result.indexed[0].markdown_content == malformed_content
|
||||
assert result.indexed[0].markdown_content == persisted_content
|
||||
|
||||
entity = await entity_repository.get_by_file_path(path)
|
||||
assert entity is not None
|
||||
|
||||
@@ -8,6 +8,7 @@ from types import SimpleNamespace
|
||||
import pytest
|
||||
|
||||
from basic_memory.config import BasicMemoryConfig
|
||||
import basic_memory.repository.embedding_provider_factory as embedding_provider_factory_module
|
||||
from basic_memory.repository.embedding_provider_factory import (
|
||||
create_embedding_provider,
|
||||
reset_embedding_provider_cache,
|
||||
@@ -264,6 +265,96 @@ def test_embedding_provider_factory_forwards_fastembed_runtime_knobs():
|
||||
assert provider.parallel == 2
|
||||
|
||||
|
||||
def test_fastembed_provider_reports_runtime_log_attrs():
|
||||
"""FastEmbed should expose the resolved runtime knobs for batch startup logs."""
|
||||
provider = FastEmbedEmbeddingProvider(batch_size=128, threads=4, parallel=2)
|
||||
|
||||
assert provider.runtime_log_attrs() == {
|
||||
"provider_batch_size": 128,
|
||||
"threads": 4,
|
||||
"configured_parallel": 2,
|
||||
"effective_parallel": 2,
|
||||
}
|
||||
|
||||
|
||||
def test_openai_provider_reports_runtime_log_attrs():
|
||||
"""OpenAI provider should expose API batch fan-out settings for startup logs."""
|
||||
provider = OpenAIEmbeddingProvider(batch_size=32, request_concurrency=6)
|
||||
|
||||
assert provider.runtime_log_attrs() == {
|
||||
"provider_batch_size": 32,
|
||||
"request_concurrency": 6,
|
||||
}
|
||||
|
||||
|
||||
def test_embedding_provider_factory_auto_tunes_fastembed_runtime_knobs_from_cpu_budget(monkeypatch):
|
||||
"""Unset FastEmbed runtime knobs should resolve from available CPU budget."""
|
||||
monkeypatch.setattr(embedding_provider_factory_module.os, "process_cpu_count", lambda: 8)
|
||||
monkeypatch.setattr(embedding_provider_factory_module.os, "cpu_count", lambda: 8)
|
||||
|
||||
config = BasicMemoryConfig(
|
||||
env="test",
|
||||
projects={"test-project": "/tmp/basic-memory-test"},
|
||||
default_project="test-project",
|
||||
semantic_search_enabled=True,
|
||||
semantic_embedding_provider="fastembed",
|
||||
semantic_embedding_threads=None,
|
||||
semantic_embedding_parallel=None,
|
||||
)
|
||||
|
||||
provider = create_embedding_provider(config)
|
||||
|
||||
assert isinstance(provider, FastEmbedEmbeddingProvider)
|
||||
assert provider.threads == 6
|
||||
assert provider.parallel == 1
|
||||
|
||||
|
||||
def test_embedding_provider_factory_auto_tuning_caps_large_cpu_budgets(monkeypatch):
|
||||
"""Large workers should still leave some headroom and stop at the thread cap."""
|
||||
monkeypatch.setattr(embedding_provider_factory_module.os, "process_cpu_count", lambda: 16)
|
||||
monkeypatch.setattr(embedding_provider_factory_module.os, "cpu_count", lambda: 16)
|
||||
|
||||
config = BasicMemoryConfig(
|
||||
env="test",
|
||||
projects={"test-project": "/tmp/basic-memory-test"},
|
||||
default_project="test-project",
|
||||
semantic_search_enabled=True,
|
||||
semantic_embedding_provider="fastembed",
|
||||
semantic_embedding_threads=None,
|
||||
semantic_embedding_parallel=None,
|
||||
)
|
||||
|
||||
provider = create_embedding_provider(config)
|
||||
|
||||
assert isinstance(provider, FastEmbedEmbeddingProvider)
|
||||
assert provider.threads == 8
|
||||
assert provider.parallel == 1
|
||||
|
||||
|
||||
def test_embedding_provider_factory_auto_tuning_stays_conservative_on_small_cpu_budget(
|
||||
monkeypatch,
|
||||
):
|
||||
"""Small workers should not get an oversized FastEmbed runtime footprint."""
|
||||
monkeypatch.setattr(embedding_provider_factory_module.os, "process_cpu_count", lambda: 2)
|
||||
monkeypatch.setattr(embedding_provider_factory_module.os, "cpu_count", lambda: 2)
|
||||
|
||||
config = BasicMemoryConfig(
|
||||
env="test",
|
||||
projects={"test-project": "/tmp/basic-memory-test"},
|
||||
default_project="test-project",
|
||||
semantic_search_enabled=True,
|
||||
semantic_embedding_provider="fastembed",
|
||||
semantic_embedding_threads=None,
|
||||
semantic_embedding_parallel=None,
|
||||
)
|
||||
|
||||
provider = create_embedding_provider(config)
|
||||
|
||||
assert isinstance(provider, FastEmbedEmbeddingProvider)
|
||||
assert provider.threads == 2
|
||||
assert provider.parallel == 1
|
||||
|
||||
|
||||
def test_embedding_provider_factory_reuses_provider_for_same_cache_key():
|
||||
"""Factory should reuse the same provider instance for identical config values."""
|
||||
config_a = BasicMemoryConfig(
|
||||
@@ -289,6 +380,36 @@ def test_embedding_provider_factory_reuses_provider_for_same_cache_key():
|
||||
assert provider_a is provider_b
|
||||
|
||||
|
||||
def test_embedding_provider_factory_reuses_auto_tuned_provider_for_same_cpu_budget(monkeypatch):
|
||||
"""Auto-tuned FastEmbed providers should still reuse the process cache."""
|
||||
monkeypatch.setattr(embedding_provider_factory_module.os, "process_cpu_count", lambda: 8)
|
||||
monkeypatch.setattr(embedding_provider_factory_module.os, "cpu_count", lambda: 8)
|
||||
|
||||
config_a = BasicMemoryConfig(
|
||||
env="test",
|
||||
projects={"test-project": "/tmp/basic-memory-test"},
|
||||
default_project="test-project",
|
||||
semantic_search_enabled=True,
|
||||
semantic_embedding_provider="fastembed",
|
||||
semantic_embedding_threads=None,
|
||||
semantic_embedding_parallel=None,
|
||||
)
|
||||
config_b = BasicMemoryConfig(
|
||||
env="test",
|
||||
projects={"test-project": "/tmp/basic-memory-test"},
|
||||
default_project="test-project",
|
||||
semantic_search_enabled=True,
|
||||
semantic_embedding_provider="fastembed",
|
||||
semantic_embedding_threads=None,
|
||||
semantic_embedding_parallel=None,
|
||||
)
|
||||
|
||||
provider_a = create_embedding_provider(config_a)
|
||||
provider_b = create_embedding_provider(config_b)
|
||||
|
||||
assert provider_a is provider_b
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_provider_runs_batches_concurrently_and_preserves_output_order(monkeypatch):
|
||||
"""Concurrent request fan-out should keep batch order stable."""
|
||||
|
||||
@@ -521,6 +521,7 @@ async def test_postgres_vector_sync_skips_unchanged_and_reembeds_changed_content
|
||||
assert unchanged_result.entities_synced == 1
|
||||
assert unchanged_result.entities_skipped == 1
|
||||
assert unchanged_result.embedding_jobs_total == 0
|
||||
assert unchanged_result.queue_wait_seconds_total == pytest.approx(0.0, abs=0.01)
|
||||
assert unchanged_result.chunks_skipped == unchanged_result.chunks_total
|
||||
|
||||
await repo.index_item(
|
||||
|
||||
@@ -195,12 +195,9 @@ class TestEnsureVectorTablesSchemaBootstrapping:
|
||||
|
||||
executed_sql = [str(call.args[0]) for call in session.execute.await_args_list]
|
||||
|
||||
assert any("CREATE TABLE IF NOT EXISTS search_vector_chunks" in sql for sql in executed_sql)
|
||||
assert any(
|
||||
"CREATE TABLE IF NOT EXISTS search_vector_chunks" in sql for sql in executed_sql
|
||||
)
|
||||
assert any(
|
||||
"CREATE TABLE IF NOT EXISTS search_vector_embeddings" in sql
|
||||
for sql in executed_sql
|
||||
"CREATE TABLE IF NOT EXISTS search_vector_embeddings" in sql for sql in executed_sql
|
||||
)
|
||||
assert not any("ALTER TABLE search_vector_chunks" in sql for sql in executed_sql)
|
||||
session.commit.assert_awaited_once()
|
||||
@@ -287,11 +284,11 @@ class TestWriteEmbeddings:
|
||||
assert params["embedding_dims_1"] == 4
|
||||
|
||||
|
||||
class TestBatchPrepareConcurrency:
|
||||
"""Cover the Postgres-specific concurrent prepare window."""
|
||||
class TestBatchPrepareWindow:
|
||||
"""Cover the shared batched prepare window used by Postgres."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sync_entity_vectors_batch_prepares_entities_concurrently(self, monkeypatch):
|
||||
async def test_sync_entity_vectors_batch_uses_shared_prepare_window(self, monkeypatch):
|
||||
repo = _make_repo(
|
||||
semantic_enabled=True,
|
||||
embedding_provider=StubEmbeddingProvider(),
|
||||
@@ -300,30 +297,65 @@ class TestBatchPrepareConcurrency:
|
||||
repo._semantic_embedding_sync_batch_size = 8
|
||||
repo._vector_tables_initialized = True
|
||||
|
||||
fetched_windows: list[list[int]] = []
|
||||
prepared_windows: list[list[int]] = []
|
||||
active_prepares = 0
|
||||
max_active_prepares = 0
|
||||
|
||||
async def _stub_prepare(entity_id: int) -> _PreparedEntityVectorSync:
|
||||
async def _stub_fetch_source_rows(session, entity_ids: list[int]):
|
||||
fetched_windows.append(list(entity_ids))
|
||||
return {entity_id: [object()] for entity_id in entity_ids}
|
||||
|
||||
async def _stub_fetch_existing_rows(session, entity_ids: list[int]):
|
||||
return {entity_id: [] for entity_id in entity_ids}
|
||||
|
||||
async def _stub_prepare_prefetched(
|
||||
*,
|
||||
entity_id: int,
|
||||
source_rows,
|
||||
existing_rows,
|
||||
) -> _PreparedEntityVectorSync:
|
||||
nonlocal active_prepares, max_active_prepares
|
||||
assert len(source_rows) == 1
|
||||
assert existing_rows == []
|
||||
active_prepares += 1
|
||||
max_active_prepares = max(max_active_prepares, active_prepares)
|
||||
await asyncio.sleep(0)
|
||||
active_prepares -= 1
|
||||
prepared_windows.append([entity_id])
|
||||
return _PreparedEntityVectorSync(
|
||||
entity_id=entity_id,
|
||||
sync_start=float(entity_id),
|
||||
source_rows_count=1,
|
||||
embedding_jobs=[],
|
||||
entity_skipped=True,
|
||||
chunks_total=1,
|
||||
chunks_skipped=1,
|
||||
prepare_seconds=0.1,
|
||||
)
|
||||
|
||||
@asynccontextmanager
|
||||
async def fake_scoped_session(session_maker):
|
||||
yield AsyncMock()
|
||||
|
||||
monkeypatch.setattr(repo, "_ensure_vector_tables", AsyncMock())
|
||||
monkeypatch.setattr(repo, "_prepare_entity_vector_jobs", _stub_prepare)
|
||||
monkeypatch.setattr(
|
||||
"basic_memory.repository.search_repository_base.db.scoped_session",
|
||||
fake_scoped_session,
|
||||
)
|
||||
monkeypatch.setattr(repo, "_fetch_prepare_window_source_rows", _stub_fetch_source_rows)
|
||||
monkeypatch.setattr(repo, "_fetch_prepare_window_existing_rows", _stub_fetch_existing_rows)
|
||||
monkeypatch.setattr(
|
||||
repo, "_prepare_entity_vector_jobs_prefetched", _stub_prepare_prefetched
|
||||
)
|
||||
|
||||
result = await repo.sync_entity_vectors_batch([1, 2, 3, 4])
|
||||
|
||||
assert result.entities_total == 4
|
||||
assert result.entities_synced == 4
|
||||
assert result.entities_failed == 0
|
||||
assert fetched_windows == [[1, 2], [3, 4]]
|
||||
assert prepared_windows == [[1], [2], [3], [4]]
|
||||
assert max_active_prepares == 2
|
||||
|
||||
|
||||
@@ -337,14 +369,17 @@ async def test_postgres_batch_sync_tracks_prepare_and_queue_wait(monkeypatch):
|
||||
repo._semantic_embedding_sync_batch_size = 2
|
||||
repo._vector_tables_initialized = True
|
||||
|
||||
async def _stub_prepare(entity_id: int) -> _PreparedEntityVectorSync:
|
||||
return _PreparedEntityVectorSync(
|
||||
entity_id=entity_id,
|
||||
sync_start=0.0,
|
||||
source_rows_count=1,
|
||||
embedding_jobs=[(200 + entity_id, f"chunk-{entity_id}")],
|
||||
prepare_seconds=1.0,
|
||||
)
|
||||
async def _stub_prepare_window(entity_ids: list[int]):
|
||||
return [
|
||||
_PreparedEntityVectorSync(
|
||||
entity_id=entity_id,
|
||||
sync_start=0.0,
|
||||
source_rows_count=1,
|
||||
embedding_jobs=[(200 + entity_id, f"chunk-{entity_id}")],
|
||||
prepare_seconds=1.0,
|
||||
)
|
||||
for entity_id in entity_ids
|
||||
]
|
||||
|
||||
async def _stub_flush(flush_jobs, entity_runtime, synced_entity_ids):
|
||||
for job in flush_jobs:
|
||||
@@ -364,10 +399,10 @@ async def test_postgres_batch_sync_tracks_prepare_and_queue_wait(monkeypatch):
|
||||
def _capture_log(**kwargs):
|
||||
completion_records.append(kwargs)
|
||||
|
||||
perf_counter_values = iter([4.0, 5.0])
|
||||
perf_counter_values = iter([0.0, 4.0, 5.0, 6.0])
|
||||
|
||||
monkeypatch.setattr(repo, "_ensure_vector_tables", AsyncMock())
|
||||
monkeypatch.setattr(repo, "_prepare_entity_vector_jobs", _stub_prepare)
|
||||
monkeypatch.setattr(repo, "_prepare_entity_vector_jobs_window", _stub_prepare_window)
|
||||
monkeypatch.setattr(repo, "_flush_embedding_jobs", _stub_flush)
|
||||
monkeypatch.setattr(repo, "_log_vector_sync_complete", _capture_log)
|
||||
monkeypatch.setattr(
|
||||
@@ -401,33 +436,41 @@ async def test_postgres_batch_sync_tracks_deferred_oversized_entities(monkeypatc
|
||||
repo._semantic_embedding_sync_batch_size = 8
|
||||
repo._vector_tables_initialized = True
|
||||
|
||||
async def _stub_prepare(entity_id: int) -> _PreparedEntityVectorSync:
|
||||
if entity_id == 1:
|
||||
return _PreparedEntityVectorSync(
|
||||
entity_id=entity_id,
|
||||
sync_start=0.0,
|
||||
source_rows_count=1,
|
||||
embedding_jobs=[(201, "chunk-1a"), (202, "chunk-1b")],
|
||||
chunks_total=5,
|
||||
pending_jobs_total=5,
|
||||
entity_complete=False,
|
||||
oversized_entity=True,
|
||||
shard_index=1,
|
||||
shard_count=3,
|
||||
remaining_jobs_after_shard=3,
|
||||
async def _stub_prepare_window(entity_ids: list[int]):
|
||||
prepared: list[_PreparedEntityVectorSync] = []
|
||||
for entity_id in entity_ids:
|
||||
if entity_id == 1:
|
||||
prepared.append(
|
||||
_PreparedEntityVectorSync(
|
||||
entity_id=entity_id,
|
||||
sync_start=0.0,
|
||||
source_rows_count=1,
|
||||
embedding_jobs=[(201, "chunk-1a"), (202, "chunk-1b")],
|
||||
chunks_total=5,
|
||||
pending_jobs_total=5,
|
||||
entity_complete=False,
|
||||
oversized_entity=True,
|
||||
shard_index=1,
|
||||
shard_count=3,
|
||||
remaining_jobs_after_shard=3,
|
||||
)
|
||||
)
|
||||
continue
|
||||
prepared.append(
|
||||
_PreparedEntityVectorSync(
|
||||
entity_id=entity_id,
|
||||
sync_start=0.0,
|
||||
source_rows_count=1,
|
||||
embedding_jobs=[(301, "chunk-2a")],
|
||||
chunks_total=1,
|
||||
pending_jobs_total=1,
|
||||
entity_complete=True,
|
||||
shard_index=1,
|
||||
shard_count=1,
|
||||
remaining_jobs_after_shard=0,
|
||||
)
|
||||
)
|
||||
return _PreparedEntityVectorSync(
|
||||
entity_id=entity_id,
|
||||
sync_start=0.0,
|
||||
source_rows_count=1,
|
||||
embedding_jobs=[(301, "chunk-2a")],
|
||||
chunks_total=1,
|
||||
pending_jobs_total=1,
|
||||
entity_complete=True,
|
||||
shard_index=1,
|
||||
shard_count=1,
|
||||
remaining_jobs_after_shard=0,
|
||||
)
|
||||
return prepared
|
||||
|
||||
async def _stub_flush(flush_jobs, entity_runtime, synced_entity_ids):
|
||||
for job in flush_jobs:
|
||||
@@ -443,7 +486,7 @@ async def test_postgres_batch_sync_tracks_deferred_oversized_entities(monkeypatc
|
||||
completion_records.append(kwargs)
|
||||
|
||||
monkeypatch.setattr(repo, "_ensure_vector_tables", AsyncMock())
|
||||
monkeypatch.setattr(repo, "_prepare_entity_vector_jobs", _stub_prepare)
|
||||
monkeypatch.setattr(repo, "_prepare_entity_vector_jobs_window", _stub_prepare_window)
|
||||
monkeypatch.setattr(repo, "_flush_embedding_jobs", _stub_flush)
|
||||
monkeypatch.setattr(repo, "_log_vector_sync_complete", _capture_log)
|
||||
|
||||
|
||||
@@ -4,11 +4,15 @@ Covers: _compose_row_source_text, _split_text_into_chunks, _build_chunk_records,
|
||||
_search_hybrid entity_id fusion key, and SemanticSearchDisabledError in SQLite.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from contextlib import asynccontextmanager
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
|
||||
import basic_memory.repository.search_repository_base as search_repository_base_module
|
||||
from basic_memory.repository.fastembed_provider import FastEmbedEmbeddingProvider
|
||||
from basic_memory.repository.search_repository_base import (
|
||||
MAX_VECTOR_CHUNK_CHARS,
|
||||
SearchRepositoryBase,
|
||||
@@ -312,8 +316,8 @@ async def test_sync_entity_vectors_batch_flushes_at_configured_threshold(monkeyp
|
||||
}
|
||||
flush_sizes: list[int] = []
|
||||
|
||||
async def _stub_prepare(entity_id: int) -> _PreparedEntityVectorSync:
|
||||
return prepared_by_entity[entity_id]
|
||||
async def _stub_prepare_window(entity_ids: list[int]):
|
||||
return [prepared_by_entity[entity_id] for entity_id in entity_ids]
|
||||
|
||||
async def _stub_flush(flush_jobs, entity_runtime, synced_entity_ids):
|
||||
flush_sizes.append(len(flush_jobs))
|
||||
@@ -325,7 +329,7 @@ async def test_sync_entity_vectors_batch_flushes_at_configured_threshold(monkeyp
|
||||
entity_runtime.pop(job.entity_id, None)
|
||||
return (0.1, 0.2)
|
||||
|
||||
monkeypatch.setattr(repo, "_prepare_entity_vector_jobs", _stub_prepare)
|
||||
monkeypatch.setattr(repo, "_prepare_entity_vector_jobs_window", _stub_prepare_window)
|
||||
monkeypatch.setattr(repo, "_flush_embedding_jobs", _stub_flush)
|
||||
|
||||
result = await repo.sync_entity_vectors_batch([1, 2, 3])
|
||||
@@ -342,6 +346,83 @@ async def test_sync_entity_vectors_batch_flushes_at_configured_threshold(monkeyp
|
||||
assert result.write_seconds_total == pytest.approx(0.4)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sync_entity_vectors_batch_skip_only_has_zero_queue_wait(monkeypatch):
|
||||
"""Skip-only batches should not accumulate synthetic queue wait."""
|
||||
repo = _ConcreteRepo()
|
||||
repo._semantic_enabled = True
|
||||
repo._embedding_provider = object()
|
||||
|
||||
async def _stub_prepare_window(entity_ids: list[int]):
|
||||
return [
|
||||
_PreparedEntityVectorSync(
|
||||
entity_id=entity_id,
|
||||
sync_start=float(entity_id),
|
||||
source_rows_count=1,
|
||||
embedding_jobs=[],
|
||||
chunks_total=2,
|
||||
chunks_skipped=2,
|
||||
entity_skipped=True,
|
||||
prepare_seconds=0.25,
|
||||
)
|
||||
for entity_id in entity_ids
|
||||
]
|
||||
|
||||
monkeypatch.setattr(repo, "_prepare_entity_vector_jobs_window", _stub_prepare_window)
|
||||
|
||||
result = await repo.sync_entity_vectors_batch([1, 2])
|
||||
|
||||
assert result.entities_total == 2
|
||||
assert result.entities_synced == 2
|
||||
assert result.entities_skipped == 2
|
||||
assert result.embedding_jobs_total == 0
|
||||
assert result.queue_wait_seconds_total == pytest.approx(0.0)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sync_entity_vectors_batch_progress_tracks_terminal_entities(monkeypatch):
|
||||
"""Progress callback should advance on terminal entity completion, not prepare entry."""
|
||||
repo = _ConcreteRepo()
|
||||
repo._semantic_enabled = True
|
||||
repo._embedding_provider = object()
|
||||
repo._semantic_embedding_sync_batch_size = 2
|
||||
|
||||
prepared_by_entity = {
|
||||
1: _PreparedEntityVectorSync(1, 1.0, 1, []),
|
||||
2: _PreparedEntityVectorSync(2, 2.0, 1, [(102, "chunk-2")]),
|
||||
3: _PreparedEntityVectorSync(3, 3.0, 1, [(103, "chunk-3")]),
|
||||
}
|
||||
progress_events: list[tuple[int, int, int]] = []
|
||||
|
||||
async def _stub_prepare_window(entity_ids: list[int]):
|
||||
return [prepared_by_entity[entity_id] for entity_id in entity_ids]
|
||||
|
||||
async def _stub_flush(flush_jobs, entity_runtime, synced_entity_ids):
|
||||
for job in flush_jobs:
|
||||
runtime = entity_runtime[job.entity_id]
|
||||
runtime.remaining_jobs -= 1
|
||||
if runtime.remaining_jobs <= 0:
|
||||
synced_entity_ids.add(job.entity_id)
|
||||
return (0.1, 0.2)
|
||||
|
||||
monkeypatch.setattr(repo, "_prepare_entity_vector_jobs_window", _stub_prepare_window)
|
||||
monkeypatch.setattr(repo, "_flush_embedding_jobs", _stub_flush)
|
||||
|
||||
result = await repo.sync_entity_vectors_batch(
|
||||
[1, 2, 3],
|
||||
progress_callback=lambda entity_id, completed, total: progress_events.append(
|
||||
(entity_id, completed, total)
|
||||
),
|
||||
)
|
||||
|
||||
assert result.entities_synced == 3
|
||||
assert progress_events == [
|
||||
(1, 1, 3),
|
||||
(2, 2, 3),
|
||||
(3, 3, 3),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sync_entity_vectors_batch_continue_on_error(monkeypatch):
|
||||
"""Batch sync should continue after per-entity and per-flush failures."""
|
||||
@@ -350,12 +431,18 @@ async def test_sync_entity_vectors_batch_continue_on_error(monkeypatch):
|
||||
repo._embedding_provider = object()
|
||||
repo._semantic_embedding_sync_batch_size = 1
|
||||
|
||||
async def _stub_prepare(entity_id: int) -> _PreparedEntityVectorSync:
|
||||
if entity_id == 2:
|
||||
raise RuntimeError("prepare failed")
|
||||
return _PreparedEntityVectorSync(
|
||||
entity_id, float(entity_id), 1, [(100 + entity_id, "chunk")]
|
||||
)
|
||||
async def _stub_prepare_window(entity_ids: list[int]):
|
||||
prepared = []
|
||||
for entity_id in entity_ids:
|
||||
if entity_id == 2:
|
||||
prepared.append(RuntimeError("prepare failed"))
|
||||
continue
|
||||
prepared.append(
|
||||
_PreparedEntityVectorSync(
|
||||
entity_id, float(entity_id), 1, [(100 + entity_id, "chunk")]
|
||||
)
|
||||
)
|
||||
return prepared
|
||||
|
||||
async def _stub_flush(flush_jobs, entity_runtime, synced_entity_ids):
|
||||
entity_id = flush_jobs[0].entity_id
|
||||
@@ -367,7 +454,7 @@ async def test_sync_entity_vectors_batch_continue_on_error(monkeypatch):
|
||||
entity_runtime.pop(entity_id, None)
|
||||
return (0.05, 0.05)
|
||||
|
||||
monkeypatch.setattr(repo, "_prepare_entity_vector_jobs", _stub_prepare)
|
||||
monkeypatch.setattr(repo, "_prepare_entity_vector_jobs_window", _stub_prepare_window)
|
||||
monkeypatch.setattr(repo, "_flush_embedding_jobs", _stub_flush)
|
||||
|
||||
result = await repo.sync_entity_vectors_batch([1, 2, 3])
|
||||
@@ -378,6 +465,71 @@ async def test_sync_entity_vectors_batch_continue_on_error(monkeypatch):
|
||||
assert result.failed_entity_ids == [2, 3]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sync_entity_vectors_batch_only_attributes_queue_wait_to_flushed_entities(
|
||||
monkeypatch,
|
||||
):
|
||||
"""Mixed batches should only charge queue wait to entities that entered flush work."""
|
||||
repo = _ConcreteRepo()
|
||||
repo._semantic_enabled = True
|
||||
repo._embedding_provider = object()
|
||||
repo._semantic_embedding_sync_batch_size = 2
|
||||
|
||||
async def _stub_prepare_window(entity_ids: list[int]):
|
||||
prepared: list[_PreparedEntityVectorSync] = []
|
||||
for entity_id in entity_ids:
|
||||
if entity_id == 1:
|
||||
prepared.append(
|
||||
_PreparedEntityVectorSync(
|
||||
entity_id=1,
|
||||
sync_start=0.0,
|
||||
source_rows_count=1,
|
||||
embedding_jobs=[],
|
||||
chunks_total=2,
|
||||
chunks_skipped=2,
|
||||
entity_skipped=True,
|
||||
prepare_seconds=0.5,
|
||||
)
|
||||
)
|
||||
continue
|
||||
prepared.append(
|
||||
_PreparedEntityVectorSync(
|
||||
entity_id=2,
|
||||
sync_start=0.0,
|
||||
source_rows_count=1,
|
||||
embedding_jobs=[(102, "chunk-2")],
|
||||
prepare_seconds=1.0,
|
||||
)
|
||||
)
|
||||
return prepared
|
||||
|
||||
async def _stub_flush(flush_jobs, entity_runtime, synced_entity_ids):
|
||||
runtime = entity_runtime[2]
|
||||
runtime.embed_seconds = 1.0
|
||||
runtime.write_seconds = 0.5
|
||||
runtime.remaining_jobs = 0
|
||||
synced_entity_ids.add(2)
|
||||
return (1.0, 0.5)
|
||||
|
||||
perf_counter_values = iter([0.0, 2.0, 4.0, 5.0])
|
||||
|
||||
monkeypatch.setattr(repo, "_prepare_entity_vector_jobs_window", _stub_prepare_window)
|
||||
monkeypatch.setattr(repo, "_flush_embedding_jobs", _stub_flush)
|
||||
monkeypatch.setattr(
|
||||
search_repository_base_module.time,
|
||||
"perf_counter",
|
||||
lambda: next(perf_counter_values),
|
||||
)
|
||||
|
||||
result = await repo.sync_entity_vectors_batch([1, 2])
|
||||
|
||||
assert result.entities_total == 2
|
||||
assert result.entities_synced == 2
|
||||
assert result.entities_skipped == 1
|
||||
assert result.embedding_jobs_total == 1
|
||||
assert result.queue_wait_seconds_total == pytest.approx(1.5)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sync_entity_vectors_batch_tracks_prepare_and_queue_wait_seconds(monkeypatch):
|
||||
"""Queue wait should be reported separately from prepare/embed/write timings."""
|
||||
@@ -386,14 +538,17 @@ async def test_sync_entity_vectors_batch_tracks_prepare_and_queue_wait_seconds(m
|
||||
repo._embedding_provider = object()
|
||||
repo._semantic_embedding_sync_batch_size = 2
|
||||
|
||||
async def _stub_prepare(entity_id: int) -> _PreparedEntityVectorSync:
|
||||
return _PreparedEntityVectorSync(
|
||||
entity_id=entity_id,
|
||||
sync_start=0.0,
|
||||
source_rows_count=1,
|
||||
embedding_jobs=[(100 + entity_id, f"chunk-{entity_id}")],
|
||||
prepare_seconds=1.0,
|
||||
)
|
||||
async def _stub_prepare_window(entity_ids: list[int]):
|
||||
return [
|
||||
_PreparedEntityVectorSync(
|
||||
entity_id=entity_id,
|
||||
sync_start=0.0,
|
||||
source_rows_count=1,
|
||||
embedding_jobs=[(100 + entity_id, f"chunk-{entity_id}")],
|
||||
prepare_seconds=1.0,
|
||||
)
|
||||
for entity_id in entity_ids
|
||||
]
|
||||
|
||||
async def _stub_flush(flush_jobs, entity_runtime, synced_entity_ids):
|
||||
assert len(flush_jobs) == 2
|
||||
@@ -414,9 +569,9 @@ async def test_sync_entity_vectors_batch_tracks_prepare_and_queue_wait_seconds(m
|
||||
def _capture_log(**kwargs):
|
||||
logged_completion.append(kwargs)
|
||||
|
||||
perf_counter_values = iter([4.0, 5.0])
|
||||
perf_counter_values = iter([0.0, 4.0, 5.0, 6.0])
|
||||
|
||||
monkeypatch.setattr(repo, "_prepare_entity_vector_jobs", _stub_prepare)
|
||||
monkeypatch.setattr(repo, "_prepare_entity_vector_jobs_window", _stub_prepare_window)
|
||||
monkeypatch.setattr(repo, "_flush_embedding_jobs", _stub_flush)
|
||||
monkeypatch.setattr(repo, "_log_vector_sync_complete", _capture_log)
|
||||
monkeypatch.setattr(
|
||||
@@ -438,3 +593,162 @@ async def test_sync_entity_vectors_batch_tracks_prepare_and_queue_wait_seconds(m
|
||||
for record in logged_completion:
|
||||
assert record["prepare_seconds"] == pytest.approx(1.0)
|
||||
assert record["queue_wait_seconds"] == pytest.approx(1.5)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_prepare_window_uses_entity_local_timing_after_shared_reads(monkeypatch):
|
||||
"""Per-entity prepare timing should start when that entity work actually begins."""
|
||||
repo = _ConcreteRepo()
|
||||
repo._semantic_enabled = True
|
||||
repo._embedding_provider = SimpleNamespace(model_name="stub", dimensions=4)
|
||||
|
||||
async def _stub_fetch_source_rows(session, entity_ids: list[int]):
|
||||
search_repository_base_module.time.perf_counter()
|
||||
return {entity_id: [] for entity_id in entity_ids}
|
||||
|
||||
async def _stub_fetch_existing_rows(session, entity_ids: list[int]):
|
||||
search_repository_base_module.time.perf_counter()
|
||||
return {entity_id: [] for entity_id in entity_ids}
|
||||
|
||||
@asynccontextmanager
|
||||
async def fake_scoped_session(session_maker):
|
||||
yield AsyncMock()
|
||||
|
||||
@asynccontextmanager
|
||||
async def _yielding_write_scope():
|
||||
await asyncio.sleep(0)
|
||||
yield
|
||||
|
||||
perf_counter_values = iter([0.0, 5.0, 10.0, 11.0, 12.0, 13.0])
|
||||
|
||||
monkeypatch.setattr(
|
||||
"basic_memory.repository.search_repository_base.db.scoped_session",
|
||||
fake_scoped_session,
|
||||
)
|
||||
monkeypatch.setattr(repo, "_fetch_prepare_window_source_rows", _stub_fetch_source_rows)
|
||||
monkeypatch.setattr(repo, "_fetch_prepare_window_existing_rows", _stub_fetch_existing_rows)
|
||||
monkeypatch.setattr(repo, "_prepare_entity_write_scope", _yielding_write_scope)
|
||||
monkeypatch.setattr(repo, "_prepare_vector_session", AsyncMock())
|
||||
monkeypatch.setattr(repo, "_delete_entity_chunks", AsyncMock())
|
||||
monkeypatch.setattr(
|
||||
search_repository_base_module.time,
|
||||
"perf_counter",
|
||||
lambda: next(perf_counter_values),
|
||||
)
|
||||
|
||||
prepared = await repo._prepare_entity_vector_jobs_window([1, 2])
|
||||
|
||||
assert [result.sync_start for result in prepared] == [10.0, 11.0]
|
||||
assert [result.prepare_seconds for result in prepared] == [2.0, 2.0]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sync_entity_vectors_batch_records_entity_granularity_histograms(monkeypatch):
|
||||
"""Entity timing histograms should emit one sample per finalized entity."""
|
||||
repo = _ConcreteRepo()
|
||||
repo._semantic_enabled = True
|
||||
repo._embedding_provider = object()
|
||||
repo._semantic_embedding_sync_batch_size = 2
|
||||
|
||||
async def _stub_prepare_window(entity_ids: list[int]):
|
||||
return [
|
||||
_PreparedEntityVectorSync(
|
||||
entity_id=entity_id,
|
||||
sync_start=0.0,
|
||||
source_rows_count=1,
|
||||
embedding_jobs=[(100 + entity_id, f"chunk-{entity_id}")],
|
||||
prepare_seconds=1.0,
|
||||
)
|
||||
for entity_id in entity_ids
|
||||
]
|
||||
|
||||
async def _stub_flush(flush_jobs, entity_runtime, synced_entity_ids):
|
||||
for job in flush_jobs:
|
||||
runtime = entity_runtime[job.entity_id]
|
||||
runtime.embed_seconds = 1.0
|
||||
runtime.write_seconds = 0.5
|
||||
runtime.remaining_jobs = 0
|
||||
synced_entity_ids.add(job.entity_id)
|
||||
return (2.0, 1.0)
|
||||
|
||||
histogram_calls: list[tuple[str, float, dict]] = []
|
||||
counter_calls: list[tuple[str, float, dict]] = []
|
||||
perf_counter_values = iter([0.0, 3.0, 4.5, 6.0])
|
||||
|
||||
monkeypatch.setattr(repo, "_prepare_entity_vector_jobs_window", _stub_prepare_window)
|
||||
monkeypatch.setattr(repo, "_flush_embedding_jobs", _stub_flush)
|
||||
monkeypatch.setattr(
|
||||
search_repository_base_module.telemetry,
|
||||
"record_histogram",
|
||||
lambda name, amount, **attrs: histogram_calls.append((name, amount, attrs)),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
search_repository_base_module.telemetry,
|
||||
"add_counter",
|
||||
lambda name, amount, **attrs: counter_calls.append((name, amount, attrs)),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
search_repository_base_module.time,
|
||||
"perf_counter",
|
||||
lambda: next(perf_counter_values),
|
||||
)
|
||||
|
||||
result = await repo.sync_entity_vectors_batch([1, 2])
|
||||
|
||||
assert result.entities_synced == 2
|
||||
histogram_names = [name for name, _, _ in histogram_calls]
|
||||
assert histogram_names.count("vector_sync_prepare_seconds") == 2
|
||||
assert histogram_names.count("vector_sync_queue_wait_seconds") == 2
|
||||
assert histogram_names.count("vector_sync_embed_seconds") == 2
|
||||
assert histogram_names.count("vector_sync_write_seconds") == 2
|
||||
assert histogram_names.count("vector_sync_batch_total_seconds") == 1
|
||||
assert [name for name, _, _ in counter_calls].count("vector_sync_entities_total") == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sync_entity_vectors_batch_logs_resolved_fastembed_runtime_settings(monkeypatch):
|
||||
"""Batch start should log the resolved FastEmbed knobs that shape this run."""
|
||||
repo = _ConcreteRepo()
|
||||
repo._semantic_enabled = True
|
||||
repo._embedding_provider = FastEmbedEmbeddingProvider(
|
||||
batch_size=128,
|
||||
dimensions=384,
|
||||
threads=4,
|
||||
parallel=2,
|
||||
)
|
||||
|
||||
async def _stub_prepare_window(entity_ids: list[int]):
|
||||
return [
|
||||
_PreparedEntityVectorSync(
|
||||
entity_id=entity_id,
|
||||
sync_start=0.0,
|
||||
source_rows_count=1,
|
||||
embedding_jobs=[],
|
||||
entity_skipped=True,
|
||||
)
|
||||
for entity_id in entity_ids
|
||||
]
|
||||
|
||||
info_calls: list[tuple[str, dict]] = []
|
||||
|
||||
def _capture_info(message: str, **kwargs):
|
||||
info_calls.append((message, kwargs))
|
||||
|
||||
monkeypatch.setattr(repo, "_prepare_entity_vector_jobs_window", _stub_prepare_window)
|
||||
monkeypatch.setattr(search_repository_base_module.logger, "info", _capture_info)
|
||||
|
||||
result = await repo.sync_entity_vectors_batch([1])
|
||||
|
||||
assert result.entities_synced == 1
|
||||
runtime_logs = [
|
||||
kwargs
|
||||
for message, kwargs in info_calls
|
||||
if message.startswith("Vector batch runtime settings:")
|
||||
]
|
||||
assert len(runtime_logs) == 1
|
||||
assert runtime_logs[0]["model_name"] == "bge-small-en-v1.5"
|
||||
assert runtime_logs[0]["provider_batch_size"] == 128
|
||||
assert runtime_logs[0]["sync_batch_size"] == 64
|
||||
assert runtime_logs[0]["threads"] == 4
|
||||
assert runtime_logs[0]["configured_parallel"] == 2
|
||||
assert runtime_logs[0]["effective_parallel"] == 2
|
||||
|
||||
@@ -1,12 +1,15 @@
|
||||
"""SQLite sqlite-vec search repository tests."""
|
||||
|
||||
import asyncio
|
||||
from contextlib import asynccontextmanager
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import MagicMock
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import text
|
||||
|
||||
from basic_memory import db
|
||||
from basic_memory.config import BasicMemoryConfig, DatabaseBackend
|
||||
from basic_memory.repository.search_index_row import SearchIndexRow
|
||||
from basic_memory.repository.sqlite_search_repository import SQLiteSearchRepository
|
||||
from basic_memory.schemas.search import SearchItemType, SearchRetrievalMode
|
||||
@@ -83,6 +86,27 @@ def _enable_semantic(
|
||||
search_repository._vector_tables_initialized = False
|
||||
|
||||
|
||||
def _make_sqlite_repo_for_unit_tests() -> SQLiteSearchRepository:
|
||||
"""Build a SQLite repository without touching a real sqlite-vec install."""
|
||||
session_maker = MagicMock()
|
||||
app_config = BasicMemoryConfig(
|
||||
env="test",
|
||||
projects={"test-project": "/tmp/test"},
|
||||
default_project="test-project",
|
||||
database_backend=DatabaseBackend.SQLITE,
|
||||
semantic_search_enabled=True,
|
||||
semantic_embedding_sync_batch_size=8,
|
||||
)
|
||||
repo = SQLiteSearchRepository(
|
||||
session_maker,
|
||||
project_id=1,
|
||||
app_config=app_config,
|
||||
embedding_provider=StubEmbeddingProvider(),
|
||||
)
|
||||
repo._vector_tables_initialized = True
|
||||
return repo
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sqlite_vec_tables_are_created_and_rebuilt(search_repository):
|
||||
"""Repository rebuilds vector schema deterministically on mismatch."""
|
||||
@@ -239,6 +263,7 @@ async def test_sqlite_vector_sync_skips_unchanged_and_reembeds_changed_content(s
|
||||
assert unchanged_result.entities_synced == 1
|
||||
assert unchanged_result.entities_skipped == 1
|
||||
assert unchanged_result.embedding_jobs_total == 0
|
||||
assert unchanged_result.queue_wait_seconds_total == pytest.approx(0.0, abs=0.01)
|
||||
assert unchanged_result.chunks_skipped == unchanged_result.chunks_total
|
||||
|
||||
await search_repository.index_item(
|
||||
@@ -266,6 +291,138 @@ async def test_sqlite_vector_sync_skips_unchanged_and_reembeds_changed_content(s
|
||||
assert model_changed_result.embedding_jobs_total == model_changed_result.chunks_total
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sqlite_prepare_window_uses_shared_reads_and_serialized_write_scope(monkeypatch):
|
||||
"""SQLite should batch read-side prepare work but serialize write-side mutations."""
|
||||
repo = _make_sqlite_repo_for_unit_tests()
|
||||
|
||||
fetched_windows: list[list[int]] = []
|
||||
active_write_scopes = 0
|
||||
max_active_write_scopes = 0
|
||||
|
||||
async def _stub_fetch_source_rows(session, entity_ids: list[int]):
|
||||
fetched_windows.append(list(entity_ids))
|
||||
return {entity_id: [object()] for entity_id in entity_ids}
|
||||
|
||||
async def _stub_fetch_existing_rows(session, entity_ids: list[int]):
|
||||
return {entity_id: [] for entity_id in entity_ids}
|
||||
|
||||
def _stub_build_chunk_records(source_rows):
|
||||
return [
|
||||
{
|
||||
"chunk_key": "entity:1:0",
|
||||
"chunk_text": "chunk text",
|
||||
"source_hash": "hash",
|
||||
}
|
||||
]
|
||||
|
||||
@asynccontextmanager
|
||||
async def _track_write_scope():
|
||||
nonlocal active_write_scopes, max_active_write_scopes
|
||||
async with repo._sqlite_prepare_write_lock:
|
||||
active_write_scopes += 1
|
||||
max_active_write_scopes = max(max_active_write_scopes, active_write_scopes)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
active_write_scopes -= 1
|
||||
|
||||
async def _stub_upsert(
|
||||
session,
|
||||
*,
|
||||
entity_id: int,
|
||||
scheduled_records,
|
||||
existing_by_key,
|
||||
entity_fingerprint: str,
|
||||
embedding_model: str,
|
||||
):
|
||||
await asyncio.sleep(0)
|
||||
return [(entity_id * 100, scheduled_records[0]["chunk_text"])]
|
||||
|
||||
@asynccontextmanager
|
||||
async def fake_scoped_session(session_maker):
|
||||
yield AsyncMock()
|
||||
|
||||
monkeypatch.setattr(
|
||||
"basic_memory.repository.search_repository_base.db.scoped_session",
|
||||
fake_scoped_session,
|
||||
)
|
||||
monkeypatch.setattr(repo, "_prepare_vector_session", AsyncMock())
|
||||
monkeypatch.setattr(repo, "_fetch_prepare_window_source_rows", _stub_fetch_source_rows)
|
||||
monkeypatch.setattr(repo, "_fetch_prepare_window_existing_rows", _stub_fetch_existing_rows)
|
||||
monkeypatch.setattr(repo, "_build_chunk_records", _stub_build_chunk_records)
|
||||
monkeypatch.setattr(repo, "_prepare_entity_write_scope", _track_write_scope)
|
||||
monkeypatch.setattr(repo, "_upsert_scheduled_chunk_records", _stub_upsert)
|
||||
|
||||
prepared = await repo._prepare_entity_vector_jobs_window([1, 2])
|
||||
|
||||
assert fetched_windows == [[1, 2]]
|
||||
assert [result.entity_id for result in prepared] == [1, 2]
|
||||
assert max_active_write_scopes == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sqlite_prepare_window_does_not_deadlock_when_vec_loading_inside_write_scope(
|
||||
monkeypatch,
|
||||
):
|
||||
"""SQLite should keep vec loading and prepare writes on separate locks."""
|
||||
repo = _make_sqlite_repo_for_unit_tests()
|
||||
|
||||
async def _stub_fetch_source_rows(session, entity_ids: list[int]):
|
||||
return {entity_id: [object()] for entity_id in entity_ids}
|
||||
|
||||
async def _stub_fetch_existing_rows(session, entity_ids: list[int]):
|
||||
return {entity_id: [] for entity_id in entity_ids}
|
||||
|
||||
def _stub_build_chunk_records(source_rows):
|
||||
return [
|
||||
{
|
||||
"chunk_key": "entity:1:0",
|
||||
"chunk_text": "chunk text",
|
||||
"source_hash": "hash",
|
||||
}
|
||||
]
|
||||
|
||||
async def _stub_prepare_vector_session(session):
|
||||
# Trigger: SQLite prepare writes call _prepare_vector_session() after
|
||||
# entering the write scope.
|
||||
# Why: vec loading still needs a lock, but reusing the write lock here
|
||||
# would deadlock before the first entity completes.
|
||||
# Outcome: this regression test proves the two concerns stay separate.
|
||||
async with repo._sqlite_vec_load_lock:
|
||||
await asyncio.sleep(0)
|
||||
|
||||
async def _stub_upsert(
|
||||
session,
|
||||
*,
|
||||
entity_id: int,
|
||||
scheduled_records,
|
||||
existing_by_key,
|
||||
entity_fingerprint: str,
|
||||
embedding_model: str,
|
||||
):
|
||||
return [(entity_id * 100, scheduled_records[0]["chunk_text"])]
|
||||
|
||||
@asynccontextmanager
|
||||
async def fake_scoped_session(session_maker):
|
||||
yield AsyncMock()
|
||||
|
||||
monkeypatch.setattr(
|
||||
"basic_memory.repository.search_repository_base.db.scoped_session",
|
||||
fake_scoped_session,
|
||||
)
|
||||
monkeypatch.setattr(repo, "_fetch_prepare_window_source_rows", _stub_fetch_source_rows)
|
||||
monkeypatch.setattr(repo, "_fetch_prepare_window_existing_rows", _stub_fetch_existing_rows)
|
||||
monkeypatch.setattr(repo, "_build_chunk_records", _stub_build_chunk_records)
|
||||
monkeypatch.setattr(repo, "_prepare_vector_session", _stub_prepare_vector_session)
|
||||
monkeypatch.setattr(repo, "_upsert_scheduled_chunk_records", _stub_upsert)
|
||||
|
||||
prepared = await asyncio.wait_for(repo._prepare_entity_vector_jobs_window([1]), timeout=1.0)
|
||||
|
||||
assert len(prepared) == 1
|
||||
assert prepared[0].entity_id == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sqlite_vector_search_returns_ranked_entities(search_repository):
|
||||
"""Vector mode ranks entities using sqlite-vec nearest-neighbor search."""
|
||||
|
||||
@@ -4,6 +4,7 @@ from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from basic_memory import file_utils
|
||||
from basic_memory.services.exceptions import FileOperationError
|
||||
from basic_memory.services.file_service import FileService
|
||||
|
||||
@@ -167,6 +168,34 @@ async def test_write_unicode_content(tmp_path: Path, file_service: FileService):
|
||||
assert content == test_content
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_frontmatter_checksum_matches_windows_crlf_persisted_bytes(
|
||||
tmp_path: Path, file_service: FileService, monkeypatch
|
||||
):
|
||||
"""Windows-style CRLF writes should hash the stored file, not the pre-write string."""
|
||||
test_path = tmp_path / "note.md"
|
||||
test_path.write_text("# Note\nBody\n", encoding="utf-8")
|
||||
|
||||
async def fake_write_file_atomic(path: Path, content: str) -> None:
|
||||
# Trigger: simulate Windows text-mode persistence, where logical LF strings
|
||||
# land on disk as CRLF bytes.
|
||||
# Why: the regression happened when the stored bytes diverged from the LF string
|
||||
# used to build the checksum.
|
||||
# Outcome: this test proves FileService returns the checksum for the stored file.
|
||||
persisted = content.replace("\n", "\r\n").encode("utf-8")
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
path.write_bytes(persisted)
|
||||
|
||||
monkeypatch.setattr(file_utils, "write_file_atomic", fake_write_file_atomic)
|
||||
|
||||
result = await file_service.update_frontmatter_with_result(
|
||||
test_path,
|
||||
{"title": "Note", "type": "note"},
|
||||
)
|
||||
|
||||
assert result.checksum == await file_service.compute_checksum(test_path)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_read_file_content(tmp_path: Path, file_service: FileService):
|
||||
"""Test read_file_content returns just the content without checksum."""
|
||||
|
||||
@@ -1203,7 +1203,7 @@ async def test_reindex_vectors(search_service, session_maker, test_project, monk
|
||||
assert entity_ids == created_entity_ids
|
||||
if progress_callback:
|
||||
for i, entity_id in enumerate(entity_ids):
|
||||
progress_callback(entity_id, i, len(entity_ids))
|
||||
progress_callback(entity_id, i + 1, len(entity_ids))
|
||||
return VectorSyncBatchResult(
|
||||
entities_total=len(entity_ids),
|
||||
entities_synced=len(entity_ids),
|
||||
@@ -1236,7 +1236,7 @@ async def test_reindex_vectors(search_service, session_maker, test_project, monk
|
||||
assert len(progress_calls) == stats["total_entities"]
|
||||
# Progress indices should be sequential
|
||||
for i, (_, index, total) in enumerate(progress_calls):
|
||||
assert index == i
|
||||
assert index == i + 1
|
||||
assert total == stats["total_entities"]
|
||||
|
||||
|
||||
|
||||
@@ -111,20 +111,18 @@ async def test_semantic_vector_sync_skips_embed_opt_out_and_clears_vectors(
|
||||
AsyncMock(return_value=SimpleNamespace(id=42, entity_metadata={"embed": False})),
|
||||
)
|
||||
sync_vectors = AsyncMock()
|
||||
execute_query = AsyncMock()
|
||||
delete_entity_vectors = AsyncMock()
|
||||
monkeypatch.setattr(repository, "sync_entity_vectors", sync_vectors)
|
||||
monkeypatch.setattr(repository, "execute_query", execute_query)
|
||||
monkeypatch.setattr(repository, "delete_entity_vector_rows", delete_entity_vectors)
|
||||
|
||||
await search_service.sync_entity_vectors(42)
|
||||
|
||||
sync_vectors.assert_not_awaited()
|
||||
assert execute_query.await_count == 2
|
||||
delete_entity_vectors.assert_awaited_once_with(42)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_semantic_vector_sync_resumes_when_embed_opt_out_removed(
|
||||
search_service, monkeypatch
|
||||
):
|
||||
async def test_semantic_vector_sync_resumes_when_embed_opt_out_removed(search_service, monkeypatch):
|
||||
"""Removing the opt-out should restore normal embedding sync."""
|
||||
repository = _sqlite_repo(search_service)
|
||||
repository._semantic_enabled = True
|
||||
@@ -170,9 +168,9 @@ async def test_semantic_vector_sync_batch_skips_embed_opt_out_and_reports_skips(
|
||||
entities_failed=0,
|
||||
)
|
||||
)
|
||||
execute_query = AsyncMock()
|
||||
delete_entity_vectors = AsyncMock()
|
||||
monkeypatch.setattr(repository, "sync_entity_vectors_batch", sync_batch)
|
||||
monkeypatch.setattr(repository, "execute_query", execute_query)
|
||||
monkeypatch.setattr(repository, "delete_entity_vector_rows", delete_entity_vectors)
|
||||
|
||||
result = await search_service.sync_entity_vectors_batch([41, 42])
|
||||
|
||||
@@ -181,7 +179,7 @@ async def test_semantic_vector_sync_batch_skips_embed_opt_out_and_reports_skips(
|
||||
assert result.entities_total == 2
|
||||
assert result.entities_synced == 1
|
||||
assert result.entities_skipped == 1
|
||||
assert execute_query.await_count == 2
|
||||
delete_entity_vectors.assert_awaited_once_with(41)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -256,6 +254,50 @@ async def test_reindex_vectors_respects_embed_opt_out(search_service, monkeypatc
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reindex_vectors_force_full_clears_project_vectors_before_resync(
|
||||
search_service, monkeypatch
|
||||
):
|
||||
"""Force-full vector reindex should clear derived vectors before batch sync."""
|
||||
repository = _sqlite_repo(search_service)
|
||||
repository._semantic_enabled = True
|
||||
|
||||
monkeypatch.setattr(
|
||||
search_service.entity_repository,
|
||||
"find_all",
|
||||
AsyncMock(
|
||||
return_value=[
|
||||
SimpleNamespace(id=41, entity_metadata={}),
|
||||
SimpleNamespace(id=42, entity_metadata={}),
|
||||
]
|
||||
),
|
||||
)
|
||||
purge_stale_rows = AsyncMock()
|
||||
delete_project_vectors = AsyncMock()
|
||||
sync_batch = AsyncMock(
|
||||
return_value=VectorSyncBatchResult(
|
||||
entities_total=2,
|
||||
entities_synced=2,
|
||||
entities_failed=0,
|
||||
)
|
||||
)
|
||||
monkeypatch.setattr(search_service, "_purge_stale_search_rows", purge_stale_rows)
|
||||
monkeypatch.setattr(repository, "delete_project_vector_rows", delete_project_vectors)
|
||||
monkeypatch.setattr(search_service, "sync_entity_vectors_batch", sync_batch)
|
||||
|
||||
stats = await search_service.reindex_vectors(force_full=True)
|
||||
|
||||
purge_stale_rows.assert_awaited_once()
|
||||
delete_project_vectors.assert_awaited_once()
|
||||
sync_batch.assert_awaited_once_with([41, 42], progress_callback=None)
|
||||
assert stats == {
|
||||
"total_entities": 2,
|
||||
"embedded": 2,
|
||||
"skipped": 0,
|
||||
"errors": 0,
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_semantic_vector_sync_batch_cleans_up_unknown_ids(search_service, monkeypatch):
|
||||
"""Deleted entity IDs should still flow through repository cleanup instead of being dropped."""
|
||||
@@ -291,7 +333,9 @@ async def test_semantic_vector_sync_batch_cleans_up_unknown_ids(search_service,
|
||||
called_entity_ids = {tuple(call.args[0]) for call in sync_batch.await_args_list}
|
||||
assert called_entity_ids == {(41,), (42,)}
|
||||
progress_callback_calls = [
|
||||
call for call in sync_batch.await_args_list if call.kwargs.get("progress_callback") is not None
|
||||
call
|
||||
for call in sync_batch.await_args_list
|
||||
if call.kwargs.get("progress_callback") is not None
|
||||
]
|
||||
assert len(progress_callback_calls) == 1
|
||||
assert progress_callback_calls[0].args[0] == [42]
|
||||
|
||||
@@ -18,6 +18,7 @@ from textwrap import dedent
|
||||
import pytest
|
||||
|
||||
from basic_memory.config import ProjectConfig
|
||||
from basic_memory.indexing.models import IndexingBatchResult
|
||||
from basic_memory.sync.sync_service import SyncService
|
||||
|
||||
|
||||
@@ -208,6 +209,32 @@ async def test_force_full_bypasses_watermark_optimization(
|
||||
assert project.last_scan_timestamp > initial_timestamp
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_force_full_reindexes_unchanged_files(
|
||||
sync_service: SyncService, project_config: ProjectConfig, monkeypatch
|
||||
):
|
||||
"""Test that force_full rewrites search rows even when the diff report is empty."""
|
||||
project_dir = project_config.home
|
||||
await create_test_file(project_dir / "file1.md", "# File 1\nOriginal")
|
||||
|
||||
# First sync establishes the watermark and initial search rows.
|
||||
await sync_service.sync(project_dir)
|
||||
await sleep_past_watermark()
|
||||
|
||||
indexed_batches: list[list[str]] = []
|
||||
|
||||
async def _stub_index_files(loaded_files, **kwargs):
|
||||
indexed_batches.append(sorted(loaded_files))
|
||||
return IndexingBatchResult()
|
||||
|
||||
monkeypatch.setattr(sync_service.batch_indexer, "index_files", _stub_index_files)
|
||||
|
||||
report = await sync_service.sync(project_dir, force_full=True, sync_embeddings=False)
|
||||
|
||||
assert report.total == 0
|
||||
assert indexed_batches == [["file1.md"]]
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
# Incremental Scan Base Cases
|
||||
# ==============================================================================
|
||||
|
||||
@@ -13,6 +13,20 @@ from basic_memory.config import init_api_logging, init_cli_logging, init_mcp_log
|
||||
class FakeLogfire:
|
||||
"""Small fake Logfire surface for bootstrap testing."""
|
||||
|
||||
class FakeCounter:
|
||||
def __init__(self, calls: list[tuple[float, dict | None]]) -> None:
|
||||
self.calls = calls
|
||||
|
||||
def add(self, amount: float, *, attributes=None) -> None:
|
||||
self.calls.append((amount, attributes))
|
||||
|
||||
class FakeHistogram:
|
||||
def __init__(self, calls: list[tuple[float, dict | None]]) -> None:
|
||||
self.calls = calls
|
||||
|
||||
def record(self, amount: float, *, attributes=None) -> None:
|
||||
self.calls.append((amount, attributes))
|
||||
|
||||
class CodeSource:
|
||||
def __init__(self, **kwargs):
|
||||
self.kwargs = kwargs
|
||||
@@ -21,6 +35,8 @@ class FakeLogfire:
|
||||
self.fail_on_send_to_logfire = fail_on_send_to_logfire
|
||||
self.configure_calls: list[dict] = []
|
||||
self.span_calls: list[tuple[str, dict]] = []
|
||||
self.counter_calls: dict[str, list[tuple[float, dict | None]]] = {}
|
||||
self.histogram_calls: dict[str, list[tuple[float, dict | None]]] = {}
|
||||
|
||||
def configure(self, **kwargs) -> None:
|
||||
self.configure_calls.append(kwargs)
|
||||
@@ -30,6 +46,14 @@ class FakeLogfire:
|
||||
def loguru_handler(self) -> dict:
|
||||
return {"sink": "fake-logfire", "level": "INFO"}
|
||||
|
||||
def metric_counter(self, name: str, *, unit: str = "", description: str = ""):
|
||||
self.counter_calls.setdefault(name, [])
|
||||
return self.FakeCounter(self.counter_calls[name])
|
||||
|
||||
def metric_histogram(self, name: str, *, unit: str = "", description: str = ""):
|
||||
self.histogram_calls.setdefault(name, [])
|
||||
return self.FakeHistogram(self.histogram_calls[name])
|
||||
|
||||
@contextmanager
|
||||
def span(self, name: str, **attrs):
|
||||
self.span_calls.append((name, attrs))
|
||||
@@ -176,6 +200,38 @@ def test_started_span_exposes_mutable_logfire_handle(monkeypatch) -> None:
|
||||
]
|
||||
|
||||
|
||||
def test_metrics_record_when_telemetry_enabled(monkeypatch) -> None:
|
||||
fake_logfire = FakeLogfire()
|
||||
telemetry.reset_telemetry_state()
|
||||
monkeypatch.setattr(telemetry, "_load_logfire", lambda: fake_logfire)
|
||||
telemetry.configure_telemetry(
|
||||
"basic-memory-cli",
|
||||
environment="dev",
|
||||
enable_logfire=True,
|
||||
)
|
||||
|
||||
telemetry.record_histogram(
|
||||
"vector_sync_prepare_seconds",
|
||||
1.25,
|
||||
unit="s",
|
||||
backend="sqlite",
|
||||
skip_only_batch=True,
|
||||
)
|
||||
telemetry.add_counter(
|
||||
"vector_sync_entities_skipped",
|
||||
2,
|
||||
backend="sqlite",
|
||||
skip_only_batch=True,
|
||||
)
|
||||
|
||||
assert fake_logfire.histogram_calls["vector_sync_prepare_seconds"] == [
|
||||
(1.25, {"backend": "sqlite", "skip_only_batch": True})
|
||||
]
|
||||
assert fake_logfire.counter_calls["vector_sync_entities_skipped"] == [
|
||||
(2, {"backend": "sqlite", "skip_only_batch": True})
|
||||
]
|
||||
|
||||
|
||||
def test_operation_creates_span_and_log_context(monkeypatch) -> None:
|
||||
fake_logfire = FakeLogfire()
|
||||
records: list[dict] = []
|
||||
|
||||
Reference in New Issue
Block a user