Compare commits

..

11 Commits

Author SHA1 Message Date
phernandez 2d42ba16e3 fix(core): preserve sqlite rollback semantics on windows
Signed-off-by: phernandez <paul@basicmachines.co>
2026-06-14 17:25:12 -05:00
phernandez 82d42268c8 fix(core): flush note content deletes
Signed-off-by: phernandez <paul@basicmachines.co>
2026-06-14 17:02:54 -05:00
phernandez fef6f61ffa fix(core): preserve explicit-session write state
Signed-off-by: phernandez <paul@basicmachines.co>
2026-06-14 16:50:19 -05:00
phernandez a6ad20a0b7 fix(core): address explicit-session review fallout
Signed-off-by: phernandez <paul@basicmachines.co>
2026-06-14 16:31:46 -05:00
phernandez bb65be817b test(core): update repository session coverage
Signed-off-by: phernandez <paul@basicmachines.co>
2026-06-14 14:49:11 -05:00
phernandez dcb4fe927f refactor(core): use explicit sessions in core repositories
Signed-off-by: phernandez <paul@basicmachines.co>
2026-06-14 14:49:03 -05:00
Paul Hernandez b5f13d6903 Update README.md
add devin deep wiki badge for weekly updates

Signed-off-by: Paul Hernandez <60959+phernandez@users.noreply.github.com>
2026-06-14 11:08:45 -05:00
Drew Cain 6d06c4a4e7 docs(ci): fix stale basicmachines.co post-release step
The marketing site moved to basicmemory.com (repo basicmachines-co/basicmemory.com,
now Astro + React) and no longer renders a hardcoded version number anywhere in
its UI — src/components/sections/hero.tsx has no version string. The post-release
instruction to bump the version there is obsolete. Replace it across the runbook
(release.md), AGENTS.md post-release tasks, and the release/beta justfile reminders
with the current reality: no version bump; optionally add a dated blog post under
src/content/blog/ for significant releases; skip for patches.

Signed-off-by: Drew Cain <groksrc@gmail.com>
2026-06-13 14:53:52 -05:00
Drew Cain 16869867da fix(core): gate hybrid FTS relaxation with the service's eligibility rules
Addresses Codex review on #994: the hybrid FTS branch opted into
OR-relaxation for every query shape, but SearchService's relaxed FTS
path deliberately rejects short queries and numeric identifiers because
OR-relaxing them over-broadens — and in hybrid the relaxed FTS-only rows
normalize to 1.0 and can outrank the vector result the user wanted
(e.g. "SPEC 16", "root note 1", "New Feature").

relaxed_query_words now enforces the same eligibility as
SearchService._is_relaxed_fts_fallback_eligible: tokenize on
[A-Za-z0-9]+ and return None when there are fewer than three tokens or
any token is a pure digit (in addition to the existing quoted/boolean
guards). Both the SQLite and Postgres relaxed retries route through this
helper, so the hybrid path now relaxes exactly the query shapes the
service does.

Tests updated for the tightened eligibility (short + numeric queries no
longer relax); SQLite + Postgres relaxation suites green, full
repository/search-service/link-resolver suite green (442 passed), ty
clean.

Signed-off-by: Drew Cain <groksrc@gmail.com>
2026-06-13 14:50:15 -05:00
Drew Cain 18da6194f9 fix(core): Postgres parity for FTS punctuation/relaxation (CI green)
CI Postgres shard caught two issues invisible to the local SQLite suite:

1. Postgres _prepare_single_term regression: the new edge-punctuation
   strip ran after special-character cleaning, so an all-special-char
   term ("()&!:") collapsed to empty and skipped the existing
   NOSPECIALCHARS:* guard, emitting a malformed ":*". Folded the strip
   into the word handlers so every guard survives, and added a
   single-word empty guard.

2. Backend-specific test assumptions. Four tests in
   test_search_repository.py (run under both backends via the
   search_repository fixture) asserted SQLite FTS5 syntax and
   SQLite-only strict-miss behavior. Postgres to_tsquery('english', ...)
   auto-strips stopwords, so "When did Melanie paint a sunrise?" already
   matches under strict AND. Made the four tests backend-aware via the
   existing is_postgres_backend() helper, and switched the relaxation
   integration test to a query with a word absent from the doc
   ("hiking") so the strict miss holds on both backends.

Reproduced and fixed against real Postgres (testcontainers): full
search test surface green on both backends (53 passed Postgres,
2968 SQLite), ruff + ty clean.

Signed-off-by: Drew Cain <groksrc@gmail.com>
2026-06-13 14:50:15 -05:00
Drew Cain a6d0784335 fix(core): repair FTS half of hybrid search for natural-language queries
Hybrid search was silently running vector-only on natural-language
queries — the FTS branch contributed zero candidates. Two causes in the
SQLite (and parallel Postgres) FTS query preparation:

1. Sentence punctuation forced phrase matching. A question like
   "When did Melanie paint a sunrise?" reached FTS5 as the exact phrase
   '"When did Melanie paint a sunrise?"*', which matches no document.
   The FTS5 tokenizer ignores this punctuation in the index, so
   stripping it from word edges loses nothing — but leaving it disabled
   the entire FTS contribution. _prepare_single_term now strips
   ?!.,;: from word edges of multi-word queries (interior characters —
   hyphens, slashes in permalinks/paths — untouched).

2. No relaxation when strict all-terms-AND matched nothing. Questions
   rarely have every word in one document, so even after (1) the
   strict AND returned zero rows. The hybrid path now retries once with
   an OR-joined, stopword-filtered, content-term query when the strict
   query is empty. bm25/ts_rank still rank multi-term matches first, and
   fusion with the vector branch keeps relaxed lexical candidates from
   dominating precision.

The relaxation is gated behind a new allow_relaxed=False parameter on
SearchRepositoryBase.search; only _search_hybrid opts in. Strict FTS
behavior (search_type=text, title, permalink, link resolution) is
unchanged — the service layer keeps its own conservative fallback.
No config flag, default-safe.

Discovered via the benchmark harness: two different fusion algorithms
produced byte-identical rankings across 1,986 queries (impossible with
two live sources), and instrumentation confirmed fts=0 on 40/40
sampled LoCoMo queries.

Benchmark impact (corrected LoCoMo, 1,986 queries, same index,
retrieval metrics — every category improves, no regression):
  recall@5  0.745 -> 0.823  (+7.9)
  MRR       0.618 -> 0.718  (+10.0)
  headline  r5 0.734 -> 0.801, MRR 0.621 -> 0.706
Largest gains on open_domain (+0.10 r5) and adversarial (+0.12 r5);
smallest on temporal (+0.003 r5 / +0.02 MRR).

Tests: punctuation no longer phrase-quotes; relaxation builds the
expected OR query and respects boolean/quoted/short-query intent; the
hybrid opt-in surfaces a partial-overlap document while the default
strict path still returns empty. Parallel coverage for Postgres. Full
SQLite unit suite green (2968 passed); ty + ruff clean.

Signed-off-by: Drew Cain <groksrc@gmail.com>
2026-06-13 14:50:15 -05:00
114 changed files with 6210 additions and 4251 deletions
+16 -14
View File
@@ -116,22 +116,24 @@ After PyPI release is published, update the MCP registry:
#### Website Updates
**1. basicmachines.co** (sibling `basicmachines.co` repo)
- **Goal**: Update version number displayed on the homepage
- **Location**: Search for "Basic Memory v0." in the codebase to find version displays
- **What to update**:
- Hero section heading that shows "Basic Memory v{VERSION}"
- "What's New in v{VERSION}" section heading
- Feature highlights array (look for array of features with title/description)
- **Process**:
**1. basicmemory.com** (sibling `basicmemory.com` repo
`basicmachines-co/basicmemory.com`, formerly `basicmachines.co`)
- **No version bump needed.** The marketing site is an Astro + React app and
carries **no hardcoded Basic Memory version number** anywhere in its UI
(`hero.tsx` and the rest of the site have no version string). The old
instruction to bump `src/components/sections/hero.tsx` is obsolete — that
file no longer holds a version. Release announcements are dated blog posts,
not an in-place edit.
- **Skip entirely for patch releases.**
- **Significant releases only — optional announcement post**:
1. Pull latest from GitHub: `git pull origin main`
2. Create release branch: `git checkout -b release/v{VERSION}`
3. Search codebase for current version number (e.g., "v0.16.1")
4. Update version numbers to new release version
5. Update feature highlights with 3-5 key features from this release (extract from CHANGELOG.md)
6. Commit changes: `git commit -m "chore: update to v{VERSION}"`
7. Push branch: `git push origin release/v{VERSION}`
- **Deploy**: Follow deployment process for basicmachines.co
3. Add a dated post under `src/content/blog/` modeled on an existing
release post (e.g. `basic-memory-v0-19-0-release.md`), summarizing 35
headline features from `CHANGELOG.md`
4. Commit (`git commit -s -m "..."`), push, and open a PR against
`basicmachines-co/basicmemory.com`
- **Deploy**: follow that repo's deployment process.
**2. docs.basicmemory.com** (sibling `docs.basicmemory.com` repo)
- **Goal**: Add a What's New page for the release and bump the homepage badge
+6 -1
View File
@@ -322,7 +322,12 @@ The recipe runs `just lint` + `just typecheck`, then updates every release manif
**Post-release tasks** the recipe surfaces but doesn't run:
- `docs.basicmemory.com` — add a What's New page under `content/2.whats-new/` and bump the version badge in `content/index.md` (the changelog page auto-fetches GitHub releases; see that repo's CLAUDE.md version-bump checklist)
- `basicmachines.co`bump version in `src/components/sections/hero.tsx`
- `basicmemory.com`the marketing site (Astro + React, repo
`basicmachines-co/basicmemory.com`, formerly `basicmachines.co`) carries **no
hardcoded version number** in its UI, so there is nothing to bump. For a
significant release, optionally add a dated announcement post under
`src/content/blog/` (model it on an existing `basic-memory-vX-Y-Z-release.md`).
Skip entirely for routine patch releases.
- MCP Registry — `mcp-publisher publish` from the repo root
See `.claude/commands/release/release.md` (and `beta.md`, `release-check.md`, `changelog.md` alongside it) for the full release + post-release runbook, including the slash commands.
+1
View File
@@ -6,6 +6,7 @@
[![Ruff](https://img.shields.io/endpoint?url=https://raw.githubusercontent.com/astral-sh/ruff/main/assets/badge/v2.json)](https://github.com/astral-sh/ruff)
![](https://badge.mcpx.dev?type=server 'MCP Server')
![](https://badge.mcpx.dev?type=dev 'MCP Dev')
[![Ask DeepWiki](https://deepwiki.com/badge.svg)](https://deepwiki.com/basicmachines-co/basic-memory)
## Skip the install — try Basic Memory in the cloud
+4 -2
View File
@@ -468,7 +468,8 @@ release version:
echo "📝 REMINDER: Post-release tasks:"
echo " 1. docs.basicmemory.com - Add a What's New page under content/2.whats-new/"
echo " and bump the badge in content/index.md (see that repo's CLAUDE.md)"
echo " 2. basicmachines.co - Update version in src/components/sections/hero.tsx"
echo " 2. basicmemory.com - No version number in the site UI; for a significant"
echo " release optionally add a post under src/content/blog/. Skip for patches."
echo " 3. MCP Registry - Run: mcp-publisher publish"
echo " See: .claude/commands/release/release.md for detailed instructions"
@@ -585,7 +586,8 @@ beta version:
echo "📝 REMINDER: For stable releases, update documentation sites:"
echo " 1. docs.basicmemory.com - Add a What's New page under content/2.whats-new/"
echo " and bump the badge in content/index.md (see that repo's CLAUDE.md)"
echo " 2. basicmachines.co - Update version in src/components/sections/hero.tsx"
echo " 2. basicmemory.com - No version number in the site UI; for a significant"
echo " release optionally add a post under src/content/blog/. Skip for patches."
echo " See: .claude/commands/release/release.md for detailed instructions"
# List all available recipes
@@ -17,6 +17,7 @@ from fastapi import APIRouter, HTTPException, Response, Path
from loguru import logger
import logfire
from basic_memory import db
from basic_memory.ignore_utils import (
IGNORED_PATH_REJECTION_DETAIL,
load_gitignore_patterns,
@@ -34,6 +35,8 @@ from basic_memory.deps import (
ProjectExternalIdPathDep,
SyncServiceV2ExternalDep,
TaskSchedulerDep,
SessionDep,
SessionMakerDep,
)
from basic_memory.schemas import DeleteEntitiesResponse
from basic_memory.schemas.base import Entity
@@ -81,6 +84,7 @@ async def get_graph(
project_id: ProjectExternalIdPathDep,
entity_repository: EntityRepositoryV2ExternalDep,
relation_repository: RelationRepositoryV2ExternalDep,
session: SessionDep,
) -> GraphResponse:
"""Return all entities and resolved relations for knowledge graph visualization.
@@ -96,7 +100,7 @@ async def get_graph(
logger.info("API v2 request: get_graph")
# Fetch all entities for this project
entities = await entity_repository.find_all(use_load_options=False)
entities = await entity_repository.find_all(session, use_load_options=False)
nodes = [
GraphNode(
external_id=entity.external_id,
@@ -108,7 +112,7 @@ async def get_graph(
]
# Fetch all resolved relations (to_id is not null) with eager-loaded entities
relations = await relation_repository.find_all()
relations = await relation_repository.find_all(session)
edges = [
GraphEdge(
from_id=relation.from_entity.external_id,
@@ -130,6 +134,7 @@ async def get_graph(
async def get_orphan_entities(
project_id: ProjectExternalIdPathDep,
entity_repository: EntityRepositoryV2ExternalDep,
session: SessionDep,
) -> OrphanEntitiesResponse:
"""Return entities that have no incoming or outgoing relations."""
with logfire.span(
@@ -140,7 +145,7 @@ async def get_orphan_entities(
):
logger.info("API v2 request: get_orphan_entities")
entities = await entity_repository.find_without_relations()
entities = await entity_repository.find_without_relations(session)
nodes = [
GraphNode(
external_id=entity.external_id,
@@ -165,6 +170,7 @@ async def resolve_identifier(
link_resolver: LinkResolverV2ExternalDep,
entity_repository: EntityRepositoryV2ExternalDep,
project_repository: ProjectRepositoryDep,
session: SessionDep,
) -> EntityResolveResponse:
"""Resolve a string identifier (external_id, permalink, title, or path) to entity info.
@@ -203,12 +209,15 @@ async def resolve_identifier(
):
logger.info(f"API v2 request: resolve_identifier for '{data.identifier}'")
entity = await entity_repository.get_by_external_id(data.identifier)
entity = await entity_repository.get_by_external_id(session, data.identifier)
resolution_method = "external_id" if entity else "search"
if not entity:
entity = await link_resolver.resolve_link(
data.identifier, source_path=data.source_path, strict=data.strict
data.identifier,
source_path=data.source_path,
strict=data.strict,
session=session,
)
if entity:
if entity.permalink == data.identifier:
@@ -223,7 +232,7 @@ async def resolve_identifier(
if not entity:
raise HTTPException(status_code=404, detail=f"Entity not found: '{data.identifier}'")
owner_project = await project_repository.get_by_id(entity.project_id)
owner_project = await project_repository.get_by_id(session, entity.project_id)
if not owner_project: # pragma: no cover
raise HTTPException(
status_code=500,
@@ -304,6 +313,7 @@ async def sync_file(
project_config: ProjectConfigV2ExternalDep,
search_service: SearchServiceV2ExternalDep,
app_config: AppConfigDep,
session_maker: SessionMakerDep,
) -> EntityResponseV2:
"""Index a single markdown file that exists on disk but is not indexed yet.
@@ -406,7 +416,8 @@ async def sync_file(
# Trigger: the file may already have a DB row (e.g. modified on disk after indexing)
# Why: the indexer needs to know whether to insert or update the entity
# Outcome: new is computed from the database instead of assumed by the caller
existing = await sync_service.entity_repository.get_by_file_path(file_path)
async with db.scoped_session(session_maker) as session:
existing = await sync_service.entity_repository.get_by_file_path(session, file_path)
synced = await sync_service.sync_one_markdown_file(
file_path, new=existing is None, index_search=True
)
@@ -435,6 +446,7 @@ async def sync_file(
async def get_entity_by_id(
project_id: ProjectExternalIdPathDep,
entity_repository: EntityRepositoryV2ExternalDep,
session: SessionDep,
entity_id: str = Path(..., description="Entity external ID (UUID)"),
) -> EntityResponseV2:
"""Get an entity by its external ID (UUID).
@@ -459,7 +471,7 @@ async def get_entity_by_id(
):
logger.info(f"API v2 request: get_entity_by_id entity_id={entity_id}")
entity = await entity_repository.get_by_external_id(entity_id)
entity = await entity_repository.get_by_external_id(session, entity_id)
if not entity:
raise HTTPException(
status_code=404, detail=f"Entity with external_id '{entity_id}' not found"
@@ -536,6 +548,7 @@ async def update_entity_by_id(
entity_repository: EntityRepositoryV2ExternalDep,
task_scheduler: TaskSchedulerDep,
app_config: AppConfigDep,
session_maker: SessionMakerDep,
entity_id: str = Path(..., description="Entity external ID (UUID)"),
) -> EntityResponseV2:
"""Update an entity by external ID.
@@ -557,7 +570,8 @@ async def update_entity_by_id(
):
logger.info(f"API v2 request: update_entity_by_id entity_id={entity_id}")
existing = await entity_repository.get_by_external_id(entity_id)
async with db.scoped_session(session_maker) as session:
existing = await entity_repository.get_by_external_id(session, entity_id)
created = existing is None
if existing:
@@ -568,10 +582,12 @@ async def update_entity_by_id(
write_result = await entity_service.create_entity_with_content(data)
entity = write_result.entity
if entity.external_id != entity_id:
entity = await entity_repository.update(
entity.id,
{"external_id": entity_id},
)
async with db.scoped_session(session_maker) as session:
entity = await entity_repository.update(
session,
entity.id,
{"external_id": entity_id},
)
# external_id fixup only changes the DB row. The file content is unchanged,
# so the markdown captured during the write remains valid downstream.
if not entity:
@@ -607,6 +623,7 @@ async def edit_entity_by_id(
entity_repository: EntityRepositoryV2ExternalDep,
task_scheduler: TaskSchedulerDep,
app_config: AppConfigDep,
session_maker: SessionMakerDep,
entity_id: str = Path(..., description="Entity external ID (UUID)"),
) -> EntityResponseV2:
"""Edit an existing entity by external ID using operations like append, prepend, etc.
@@ -631,7 +648,8 @@ async def edit_entity_by_id(
f"API v2 request: edit_entity_by_id entity_id={entity_id}, operation='{data.operation}'"
)
entity = await entity_repository.get_by_external_id(entity_id)
async with db.scoped_session(session_maker) as session:
entity = await entity_repository.get_by_external_id(session, entity_id)
if not entity: # pragma: no cover
raise HTTPException(
status_code=404, detail=f"Entity with external_id '{entity_id}' not found"
@@ -678,6 +696,7 @@ async def delete_entity_by_id(
project_id: ProjectExternalIdPathDep,
entity_service: EntityServiceV2ExternalDep,
entity_repository: EntityRepositoryV2ExternalDep,
session_maker: SessionMakerDep,
entity_id: str = Path(..., description="Entity external ID (UUID)"),
) -> DeleteEntitiesResponse:
"""Delete an entity by external ID.
@@ -698,7 +717,8 @@ async def delete_entity_by_id(
):
logger.info(f"API v2 request: delete_entity_by_id entity_id={entity_id}")
entity = await entity_repository.get_by_external_id(entity_id)
async with db.scoped_session(session_maker) as session:
entity = await entity_repository.get_by_external_id(session, entity_id)
if entity is None:
logger.info(f"API v2 response: external_id={entity_id} not found, deleted=False")
return DeleteEntitiesResponse(deleted=False)
@@ -724,6 +744,7 @@ async def move_entity(
app_config: AppConfigDep,
search_service: SearchServiceV2ExternalDep,
task_scheduler: TaskSchedulerDep,
session_maker: SessionMakerDep,
entity_id: str = Path(..., description="Entity external ID (UUID)"),
) -> EntityResponseV2:
"""Move an entity to a new file location.
@@ -751,7 +772,8 @@ async def move_entity(
try:
# First, get the entity by external_id to verify it exists
entity = await entity_repository.get_by_external_id(entity_id)
async with db.scoped_session(session_maker) as session:
entity = await entity_repository.get_by_external_id(session, entity_id)
if not entity: # pragma: no cover
raise HTTPException(
status_code=404, detail=f"Entity with external_id '{entity_id}' not found"
@@ -766,9 +788,11 @@ async def move_entity(
)
# Reindex at new location
reindexed_entity = await entity_service.link_resolver.resolve_link(
data.destination_path
)
async with db.scoped_session(session_maker) as session:
reindexed_entity = await entity_service.link_resolver.resolve_link(
data.destination_path,
session=session,
)
if reindexed_entity:
await search_service.index_entity(reindexed_entity)
_schedule_vector_sync_if_enabled(
@@ -805,6 +829,7 @@ async def move_directory(
app_config: AppConfigDep,
search_service: SearchServiceV2ExternalDep,
task_scheduler: TaskSchedulerDep,
session_maker: SessionMakerDep,
) -> DirectoryMoveResult:
"""Move all entities in a directory to a new location.
@@ -840,7 +865,10 @@ async def move_directory(
# Reindex moved entities
for file_path in result.moved_files:
entity = await entity_service.link_resolver.resolve_link(file_path)
async with db.scoped_session(session_maker) as session:
entity = await entity_service.link_resolver.resolve_link(
file_path, session=session
)
if entity:
await search_service.index_entity(entity)
_schedule_vector_sync_if_enabled(
@@ -10,7 +10,12 @@ from fastapi import APIRouter, Query, Path
from loguru import logger
import logfire
from basic_memory.deps import ContextServiceV2ExternalDep, EntityRepositoryV2ExternalDep
from basic_memory import db
from basic_memory.deps import (
ContextServiceV2ExternalDep,
EntityRepositoryV2ExternalDep,
SessionMakerDep,
)
from basic_memory.schemas.base import TimeFrame, parse_timeframe
from basic_memory.schemas.memory import (
GraphContext,
@@ -27,6 +32,7 @@ router = APIRouter(tags=["memory"])
async def recent(
context_service: ContextServiceV2ExternalDep,
entity_repository: EntityRepositoryV2ExternalDep,
session_maker: SessionMakerDep,
project_id: str = Path(..., description="Project external UUID"),
type: Annotated[list[SearchItemType] | None, Query()] = None,
depth: int = 1,
@@ -95,9 +101,14 @@ async def recent(
phase="shape_response",
result_count=len(context.results),
):
recent_context = await to_graph_context(
context, entity_repository=entity_repository, page=page, page_size=page_size
)
async with db.scoped_session(session_maker) as session:
recent_context = await to_graph_context(
context,
entity_repository=entity_repository,
session=session,
page=page,
page_size=page_size,
)
logger.debug(f"V2 Recent context: {recent_context.model_dump_json()}")
return recent_context
@@ -109,6 +120,7 @@ async def recent(
async def get_memory_context(
context_service: ContextServiceV2ExternalDep,
entity_repository: EntityRepositoryV2ExternalDep,
session_maker: SessionMakerDep,
uri: str,
project_id: str = Path(..., description="Project external UUID"),
depth: int = 1,
@@ -177,6 +189,11 @@ async def get_memory_context(
phase="shape_response",
result_count=len(context.results),
):
return await to_graph_context(
context, entity_repository=entity_repository, page=page, page_size=page_size
)
async with db.scoped_session(session_maker) as session:
return await to_graph_context(
context,
entity_repository=entity_repository,
session=session,
page=page,
page_size=page_size,
)
@@ -16,6 +16,7 @@ from typing import Literal, Optional
from fastapi import APIRouter, HTTPException, Body, Query, Path
from loguru import logger
from basic_memory import db
from basic_memory.deps import (
ProjectServiceDep,
ProjectRepositoryDep,
@@ -23,6 +24,8 @@ from basic_memory.deps import (
SyncServiceV2ExternalDep,
TaskSchedulerDep,
ProjectExternalIdPathDep,
SessionDep,
SessionMakerDep,
)
from basic_memory.schemas import SyncReportResponse
from basic_memory.models import Project
@@ -54,21 +57,22 @@ def _split_qualified_project_identifier(identifier: str) -> tuple[str | None, st
async def _resolve_project_identifier_candidate(
session: SessionDep,
project_repository: ProjectRepository,
identifier: str,
) -> tuple[Project | None, ProjectResolveMethod]:
"""Resolve one project identifier candidate and report the matching method."""
identifier_permalink = generate_permalink(identifier)
project = await project_repository.get_by_external_id(identifier)
project = await project_repository.get_by_external_id(session, identifier)
if project:
return project, "external_id"
project = await project_repository.get_by_permalink(identifier_permalink)
project = await project_repository.get_by_permalink(session, identifier_permalink)
if project:
return project, "permalink"
project = await project_repository.get_by_name_case_insensitive(identifier)
project = await project_repository.get_by_name_case_insensitive(session, identifier)
if project:
return project, "name" # pragma: no cover
@@ -76,11 +80,13 @@ async def _resolve_project_identifier_candidate(
async def _resolve_project_identifier(
session: SessionDep,
project_repository: ProjectRepository,
identifier: str,
) -> tuple[Project | None, ProjectResolveMethod]:
"""Resolve exact identifiers first, then accepted workspace-qualified forms."""
project, resolution_method = await _resolve_project_identifier_candidate(
session,
project_repository,
identifier,
)
@@ -96,6 +102,7 @@ async def _resolve_project_identifier(
# only needs the project segment to validate the active project.
# Outcome: models can follow the hint verbatim instead of looping on a 404.
project, resolution_method = await _resolve_project_identifier_candidate(
session,
project_repository,
project_identifier,
)
@@ -277,6 +284,7 @@ async def get_project_status(
@router.post("/resolve", response_model=ProjectResolveResponse)
async def resolve_project_identifier(
data: ProjectResolveRequest,
session: SessionDep,
project_repository: ProjectRepositoryDep,
) -> ProjectResolveResponse:
"""Resolve a project identifier (name, permalink, or external_id) to project info.
@@ -315,6 +323,7 @@ async def resolve_project_identifier(
logger.info(f"API v2 request: resolve_project_identifier for '{data.identifier}'")
project, resolution_method = await _resolve_project_identifier(
session,
project_repository,
data.identifier,
)
@@ -328,7 +337,7 @@ async def resolve_project_identifier(
# default and the bare not-found message reads as a broken install
# rather than a missing first-run step (#974 follow-up).
# Outcome: the error names the setup command instead.
if not await project_repository.find_all(limit=1, use_load_options=False):
if not await project_repository.find_all(session, limit=1, use_load_options=False):
detail = (
f"{detail}. No projects are set up yet — run "
"'basic-memory project add <name> <path>' to create one."
@@ -349,6 +358,7 @@ async def resolve_project_identifier(
@router.get("/{project_id}", response_model=ProjectItem)
async def get_project_by_id(
session: SessionDep,
project_repository: ProjectRepositoryDep,
project_id: str = Path(..., description="Project external ID (UUID)"),
) -> ProjectItem:
@@ -371,7 +381,7 @@ async def get_project_by_id(
"""
logger.info(f"API v2 request: get_project_by_id for project_id={project_id}")
project = await project_repository.get_by_external_id(project_id)
project = await project_repository.get_by_external_id(session, project_id)
if not project:
raise HTTPException(
status_code=404, detail=f"Project with external_id '{project_id}' not found"
@@ -389,12 +399,14 @@ async def get_project_by_id(
@router.get("/{project_id}/info", response_model=ProjectInfoResponse)
async def get_project_info_by_id(
project_service: ProjectServiceDep,
session_maker: SessionMakerDep,
project_repository: ProjectRepositoryDep,
project_id: str = Path(..., description="Project external ID (UUID)"),
) -> ProjectInfoResponse:
"""Get detailed project information by external ID."""
logger.info(f"API v2 request: get_project_info_by_id for project_id={project_id}")
project = await project_repository.get_by_external_id(project_id)
async with db.scoped_session(session_maker) as session:
project = await project_repository.get_by_external_id(session, project_id)
if not project:
raise HTTPException(
status_code=404, detail=f"Project with external_id '{project_id}' not found"
@@ -405,6 +417,7 @@ async def get_project_info_by_id(
@router.patch("/{project_id}", response_model=ProjectStatusResponse)
async def update_project_by_id(
project_service: ProjectServiceDep,
session_maker: SessionMakerDep,
project_repository: ProjectRepositoryDep,
project_id: str = Path(..., description="Project external ID (UUID)"),
path: Optional[str] = Body(None, description="New absolute path for the project"),
@@ -435,7 +448,8 @@ async def update_project_by_id(
raise HTTPException(status_code=400, detail="Path must be absolute")
# Get original project info for the response
old_project = await project_repository.get_by_external_id(project_id)
async with db.scoped_session(session_maker) as session:
old_project = await project_repository.get_by_external_id(session, project_id)
if not old_project:
raise HTTPException(
status_code=404, detail=f"Project with external_id '{project_id}' not found"
@@ -456,7 +470,8 @@ async def update_project_by_id(
await project_service.update_project(old_project.name, is_active=is_active)
# Get updated project info (use the same external_id)
updated_project = await project_repository.get_by_external_id(project_id)
async with db.scoped_session(session_maker) as session:
updated_project = await project_repository.get_by_external_id(session, project_id)
if not updated_project: # pragma: no cover
raise HTTPException(
status_code=404,
@@ -483,6 +498,7 @@ async def update_project_by_id(
@router.delete("/{project_id}", response_model=ProjectStatusResponse)
async def delete_project_by_id(
project_service: ProjectServiceDep,
session_maker: SessionMakerDep,
project_repository: ProjectRepositoryDep,
project_id: str = Path(..., description="Project external ID (UUID)"),
delete_notes: bool = Query(
@@ -509,7 +525,8 @@ async def delete_project_by_id(
)
try:
old_project = await project_repository.get_by_external_id(project_id)
async with db.scoped_session(session_maker) as session:
old_project = await project_repository.get_by_external_id(session, project_id)
if not old_project:
raise HTTPException(
status_code=404, detail=f"Project with external_id '{project_id}' not found"
@@ -552,6 +569,7 @@ async def delete_project_by_id(
@router.put("/{project_id}/default", response_model=ProjectStatusResponse)
async def set_default_project_by_id(
project_service: ProjectServiceDep,
session_maker: SessionMakerDep,
project_repository: ProjectRepositoryDep,
project_id: str = Path(..., description="Project external ID (UUID)"),
) -> ProjectStatusResponse:
@@ -575,10 +593,11 @@ async def set_default_project_by_id(
# Get the old default project from database. It may be absent during
# bootstrap/recovery (no default row yet); that is a valid state, not an
# error, so we only echo it back when one exists.
default_project = await project_repository.get_default_project()
async with db.scoped_session(session_maker) as session:
default_project = await project_repository.get_default_project(session)
# Get the new default project by external_id
new_default_project = await project_repository.get_by_external_id(project_id)
# Get the new default project by external_id
new_default_project = await project_repository.get_by_external_id(session, project_id)
if not new_default_project:
raise HTTPException(
status_code=404, detail=f"Project with external_id '{project_id}' not found"
@@ -10,6 +10,7 @@ from typing import Any
from fastapi import APIRouter, HTTPException, status, Path
from loguru import logger
from basic_memory import db
from basic_memory.api.v2.utils import to_graph_context, to_search_results
from basic_memory.api.template_loader import template_loader
from basic_memory.schemas.base import parse_timeframe
@@ -18,6 +19,7 @@ from basic_memory.deps import (
EntityRepositoryV2ExternalDep,
SearchServiceV2ExternalDep,
EntityServiceV2ExternalDep,
SessionMakerDep,
)
from basic_memory.schemas.prompt import (
ContinueConversationRequest,
@@ -36,6 +38,7 @@ async def continue_conversation(
entity_service: EntityServiceV2ExternalDep,
context_service: ContextServiceV2ExternalDep,
entity_repository: EntityRepositoryV2ExternalDep,
session_maker: SessionMakerDep,
request: ContinueConversationRequest,
project_id: str = Path(..., description="Project external UUID"),
) -> PromptResponse:
@@ -82,9 +85,10 @@ async def continue_conversation(
)
# Process results into the schema format
graph_context = await to_graph_context(
context_result, entity_repository=entity_repository
)
async with db.scoped_session(session_maker) as session:
graph_context = await to_graph_context(
context_result, entity_repository=entity_repository, session=session
)
# Add results to our collection (limit to top results for each permalink)
if graph_context.results:
@@ -109,7 +113,10 @@ async def continue_conversation(
max_related=request.related_items_limit,
include_observations=True,
)
recent_context = await to_graph_context(context_result, entity_repository=entity_repository)
async with db.scoped_session(session_maker) as session:
recent_context = await to_graph_context(
context_result, entity_repository=entity_repository, session=session
)
hierarchical_results = recent_context.results[:5] # Limit to top 5 recent items
@@ -16,11 +16,14 @@ from fastapi import APIRouter, HTTPException, Response, Path
from loguru import logger
import logfire
from basic_memory import db
from basic_memory.deps import (
ProjectConfigV2ExternalDep,
FileServiceV2ExternalDep,
EntityRepositoryV2ExternalDep,
SearchServiceV2ExternalDep,
SessionDep,
SessionMakerDep,
)
from basic_memory.models.knowledge import Entity as EntityModel
from basic_memory.schemas.v2.resource import (
@@ -38,6 +41,7 @@ async def get_resource_content(
config: ProjectConfigV2ExternalDep,
entity_repository: EntityRepositoryV2ExternalDep,
file_service: FileServiceV2ExternalDep,
session: SessionDep,
project_id: str = Path(..., description="Project external UUID"),
entity_id: str = Path(..., description="Entity external UUID"),
) -> Response:
@@ -70,7 +74,7 @@ async def get_resource_content(
action="get_content",
phase="load_entity",
):
entity = await entity_repository.get_by_external_id(entity_id)
entity = await entity_repository.get_by_external_id(session, entity_id)
if not entity:
raise HTTPException(status_code=404, detail=f"Entity {entity_id} not found")
@@ -121,6 +125,7 @@ async def create_resource(
file_service: FileServiceV2ExternalDep,
entity_repository: EntityRepositoryV2ExternalDep,
search_service: SearchServiceV2ExternalDep,
session_maker: SessionMakerDep,
project_id: str = Path(..., description="Project external UUID"),
) -> ResourceResponse:
"""Create a new resource file.
@@ -158,13 +163,14 @@ async def create_resource(
"Path must be relative and stay within project boundaries.",
)
existing_entity = await entity_repository.get_by_file_path(data.file_path)
if existing_entity:
raise HTTPException(
status_code=409,
detail=f"Resource already exists at {data.file_path} with entity_id {existing_entity.external_id}. "
f"Use PUT /resource/{existing_entity.external_id} to update it.",
)
async with db.scoped_session(session_maker) as session:
existing_entity = await entity_repository.get_by_file_path(session, data.file_path)
if existing_entity:
raise HTTPException(
status_code=409,
detail=f"Resource already exists at {data.file_path} with entity_id {existing_entity.external_id}. "
f"Use PUT /resource/{existing_entity.external_id} to update it.",
)
with logfire.span(
"api.resource.create.write_file",
@@ -203,7 +209,8 @@ async def create_resource(
action="create",
phase="upsert_entity",
):
entity = await entity_repository.add(entity)
async with db.scoped_session(session_maker) as session:
entity = await entity_repository.add(session, entity)
with logfire.span(
"api.resource.create.search_index",
@@ -236,6 +243,7 @@ async def update_resource(
file_service: FileServiceV2ExternalDep,
entity_repository: EntityRepositoryV2ExternalDep,
search_service: SearchServiceV2ExternalDep,
session_maker: SessionMakerDep,
project_id: str = Path(..., description="Project external UUID"),
entity_id: str = Path(..., description="Entity external UUID"),
) -> ResourceResponse:
@@ -265,7 +273,8 @@ async def update_resource(
action="update",
):
try:
entity = await entity_repository.get_by_external_id(entity_id)
async with db.scoped_session(session_maker) as session:
entity = await entity_repository.get_by_external_id(session, entity_id)
if not entity:
raise HTTPException(status_code=404, detail=f"Entity {entity_id} not found")
@@ -315,17 +324,19 @@ async def update_resource(
action="update",
phase="update_entity",
):
updated_entity = await entity_repository.update(
entity.id,
{
"title": file_name,
"note_type": note_type,
"content_type": content_type,
"file_path": target_file_path,
"checksum": checksum,
"updated_at": file_metadata.modified_at,
},
)
async with db.scoped_session(session_maker) as session:
updated_entity = await entity_repository.update(
session,
entity.id,
{
"title": file_name,
"note_type": note_type,
"content_type": content_type,
"file_path": target_file_path,
"checksum": checksum,
"updated_at": file_metadata.modified_at,
},
)
if updated_entity is None:
raise HTTPException(status_code=404, detail=f"Entity {entity_id} not found")
@@ -13,11 +13,13 @@ from pathlib import Path as FilePath
import frontmatter
from fastapi import APIRouter, Path, Query
from loguru import logger
from sqlalchemy.ext.asyncio import AsyncSession
from basic_memory.deps import (
EntityRepositoryV2ExternalDep,
FileServiceV2ExternalDep,
LinkResolverV2ExternalDep,
SessionDep,
)
from basic_memory.models.knowledge import Entity
from basic_memory.schemas.schema import (
@@ -131,6 +133,7 @@ async def validate_schema(
entity_repository: EntityRepositoryV2ExternalDep,
file_service: FileServiceV2ExternalDep,
link_resolver: LinkResolverV2ExternalDep,
session: SessionDep,
project_id: str = Path(..., description="Project external UUID"),
note_type: str | None = Query(None, description="Note type to validate"),
identifier: str | None = Query(None, description="Specific note identifier"),
@@ -150,7 +153,7 @@ async def validate_schema(
if identifier:
# Resolve identifier flexibly (permalink, title, path, fuzzy)
# to match how read_note and other tools resolve identifiers
entity = await link_resolver.resolve_link(identifier)
entity = await link_resolver.resolve_link(identifier, session=session)
if not entity:
return ValidationReport(note_type=note_type, total_notes=0, total_entities=0)
@@ -159,6 +162,7 @@ async def validate_schema(
async def search_fn(query: str) -> list[dict]:
entities = await _find_schema_entities(
session,
entity_repository,
query,
allow_reference_match=isinstance(schema_ref, str) and query == schema_ref,
@@ -187,7 +191,7 @@ async def validate_schema(
)
# --- Batch validation by note type ---
entities = await _find_by_note_type(entity_repository, note_type) if note_type else []
entities = await _find_by_note_type(session, entity_repository, note_type) if note_type else []
for entity in entities:
frontmatter = _entity_frontmatter(entity)
@@ -195,6 +199,7 @@ async def validate_schema(
async def search_fn(query: str) -> list[dict]:
entities = await _find_schema_entities(
session,
entity_repository,
query,
allow_reference_match=isinstance(schema_ref, str) and query == schema_ref,
@@ -230,6 +235,7 @@ async def validate_schema(
@router.post("/schema/infer", response_model=InferenceReport)
async def infer_schema_endpoint(
entity_repository: EntityRepositoryV2ExternalDep,
session: SessionDep,
project_id: str = Path(..., description="Project external UUID"),
note_type: str = Query(..., description="Note type to analyze"),
threshold: float = Query(0.25, description="Minimum frequency for optional fields"),
@@ -239,7 +245,7 @@ async def infer_schema_endpoint(
Examines observation categories and relation types across all notes
of the given type. Returns frequency analysis and suggested Picoschema.
"""
entities = await _find_by_note_type(entity_repository, note_type)
entities = await _find_by_note_type(session, entity_repository, note_type)
notes_data = [_entity_to_note_data(entity) for entity in entities]
result = infer_schema(note_type, notes_data, optional_threshold=threshold)
@@ -274,6 +280,7 @@ async def infer_schema_endpoint(
async def diff_schema_endpoint(
entity_repository: EntityRepositoryV2ExternalDep,
file_service: FileServiceV2ExternalDep,
session: SessionDep,
note_type: str = Path(..., description="Note type to check for drift"),
project_id: str = Path(..., description="Project external UUID"),
):
@@ -285,7 +292,7 @@ async def diff_schema_endpoint(
"""
async def search_fn(query: str) -> list[dict]:
entities = await _find_schema_entities(entity_repository, query)
entities = await _find_schema_entities(session, entity_repository, query)
return [await _schema_frontmatter_from_file(file_service, e) for e in entities]
# Resolve schema by note type
@@ -296,7 +303,7 @@ async def diff_schema_endpoint(
return DriftReport(note_type=note_type, schema_found=False)
# Collect all notes of this type
entities = await _find_by_note_type(entity_repository, note_type)
entities = await _find_by_note_type(session, entity_repository, note_type)
notes_data = [_entity_to_note_data(entity) for entity in entities]
result = diff_schema(schema_def, notes_data)
@@ -331,16 +338,18 @@ async def diff_schema_endpoint(
async def _find_by_note_type(
session: AsyncSession,
entity_repository: EntityRepositoryV2ExternalDep,
note_type: str,
) -> list[Entity]:
"""Find all entities of a given type using the repository's select pattern."""
query = entity_repository.select().where(Entity.note_type == note_type)
result = await entity_repository.execute_query(query)
result = await entity_repository.execute_query(session, query)
return list(result.scalars().all())
async def _find_schema_entities(
session: AsyncSession,
entity_repository: EntityRepositoryV2ExternalDep,
target_note_type: str,
*,
@@ -355,7 +364,7 @@ async def _find_schema_entities(
exact reference matching by title/permalink (explicit schema references)
"""
query = entity_repository.select().where(Entity.note_type == "schema")
result = await entity_repository.execute_query(query)
result = await entity_repository.execute_query(session, query)
entities = list(result.scalars().all())
normalized_target = generate_permalink(target_note_type)
+8 -2
View File
@@ -1,6 +1,7 @@
from typing import Any, Protocol, Optional, List, Sequence
import logfire
from sqlalchemy.ext.asyncio import AsyncSession
from basic_memory.repository.search_repository import SearchIndexRow
from basic_memory.schemas.memory import (
EntitySummary,
@@ -19,7 +20,11 @@ from basic_memory.services.context_service import (
class EntityBatchLookup(Protocol):
async def find_by_ids_for_hydration(
self, ids: List[int], *, include_cross_project: bool = False
self,
session: AsyncSession,
ids: List[int],
*,
include_cross_project: bool = False,
) -> Sequence[Any]: ...
@@ -42,6 +47,7 @@ def _search_item_type(value: str | SearchItemType) -> SearchItemType:
async def to_graph_context(
context_result: ServiceContextResult,
entity_repository: EntityBatchLookup,
session: AsyncSession,
page: Optional[int] = None,
page_size: Optional[int] = None,
) -> GraphContext:
@@ -90,7 +96,7 @@ async def to_graph_context(
result_count=len(entity_ids_needed),
):
entities = await entity_repository.find_by_ids_for_hydration(
list(entity_ids_needed), include_cross_project=True
session, list(entity_ids_needed), include_cross_project=True
)
for e in entities:
entity_title_lookup[e.id] = e.title
+13 -6
View File
@@ -177,8 +177,9 @@ async def _reindex_projects(app_config):
db_path=app_config.database_path,
db_type=db.DatabaseType.FILESYSTEM,
)
project_repository = ProjectRepository(session_maker)
projects = await project_repository.get_active_projects()
project_repository = ProjectRepository()
async with db.scoped_session(session_maker) as session:
projects = await project_repository.get_active_projects(session)
for project in projects:
console.print(f" Indexing [cyan]{project.name}[/cyan]...")
@@ -355,8 +356,9 @@ async def _reindex(
db_path=app_config.database_path,
db_type=db.DatabaseType.FILESYSTEM,
)
project_repository = ProjectRepository(session_maker)
projects = await project_repository.get_active_projects()
project_repository = ProjectRepository()
async with db.scoped_session(session_maker) as session:
projects = await project_repository.get_active_projects(session)
if project:
projects = [p for p in projects if p.name == project]
@@ -418,7 +420,7 @@ async def _reindex(
console.print(
f" Building vector embeddings ([cyan]{embedding_mode_label}[/cyan])..."
)
entity_repository = EntityRepository(session_maker, project_id=proj.id)
entity_repository = EntityRepository(project_id=proj.id)
search_repository = create_search_repository(
session_maker, project_id=proj.id, app_config=app_config
)
@@ -426,7 +428,12 @@ async def _reindex(
entity_parser = EntityParser(project_path)
markdown_processor = MarkdownProcessor(entity_parser, app_config=app_config)
file_service = FileService(project_path, markdown_processor, app_config=app_config)
search_service = SearchService(search_repository, entity_repository, file_service)
search_service = SearchService(
search_repository,
entity_repository,
file_service,
session_maker=session_maker,
)
with Progress(
SpinnerColumn(),
+59 -1
View File
@@ -23,6 +23,64 @@ from basic_memory.schemas import SyncReportResponse
console = Console()
def _is_default_project_delete_error(error: Exception) -> bool:
"""Return True only for the API guard that blocks deleting the default project."""
error_text = str(error)
return "Cannot delete default project" in error_text
async def _delete_doctor_project_locally(project_name: str, project_id: str) -> None:
"""Remove the generated doctor project when the public API guard blocks cleanup."""
from basic_memory import db
from basic_memory.config import ConfigManager
from basic_memory.repository import ProjectRepository
config_manager = ConfigManager()
repository = ProjectRepository()
_, session_maker = await db.get_or_create_db(
db_path=config_manager.config.database_path,
db_type=db.DatabaseType.FILESYSTEM,
)
async with db.scoped_session(session_maker) as session:
project = await repository.get_by_external_id(session, project_id)
if project is None:
raise ValueError(f"Doctor cleanup project '{project_id}' not found")
if project.name != project_name:
raise ValueError(
f"Doctor cleanup expected project '{project_name}', found '{project.name}'"
)
await repository.delete(session, project.id)
config = config_manager.load_config()
if project_name in config.projects:
del config.projects[project_name]
if config.default_project == project_name:
config.default_project = next(iter(config.projects), None)
config_manager.save_config(config)
async def _delete_doctor_project(
project_client: ProjectClient, project_name: str, project_id: str
) -> None:
"""Delete the generated doctor project without weakening the public API guard."""
# Deferred: ToolError lives in the mcp SDK, which must not load at CLI startup (#886).
from mcp.server.fastmcp.exceptions import ToolError
try:
await project_client.delete_project(project_id)
except ToolError as exc:
if not _is_default_project_delete_error(exc):
raise
# Trigger: fresh local configs can promote the generated doctor project
# to default because the placeholder default has no DB row.
# Why: the project is disposable doctor-owned state, while the public API
# must keep rejecting default-project deletion for normal callers.
# Outcome: cleanup removes only the exact doctor project it created.
await _delete_doctor_project_locally(project_name, project_id)
async def run_doctor() -> None:
"""Run local consistency checks for file <-> database flows."""
# Deferred: the markdown parsing stack is only needed while the checks run,
@@ -129,7 +187,7 @@ async def run_doctor() -> None:
finally:
if project_id:
await project_client.delete_project(project_id)
await _delete_doctor_project(project_client, project_name, project_id)
console.print("[green]Doctor checks passed.[/green]")
+23 -20
View File
@@ -1050,12 +1050,13 @@ async def _detach_local_project_row(app_config: BasicMemoryConfig, name: str) ->
db_type=db.DatabaseType.FILESYSTEM,
)
try:
repo = ProjectRepository(session_maker)
existing = await repo.get_by_name(name)
if existing is None:
return False
await repo.delete(existing.id)
return True
repo = ProjectRepository()
async with db.scoped_session(session_maker) as session:
existing = await repo.get_by_name(session, name)
if existing is None:
return False
await repo.delete(session, existing.id)
return True
finally:
# CLI-only: safe to tear down the global DB singleton here since
# set-cloud/set-local never run inside a long-lived MCP/API server.
@@ -1082,20 +1083,22 @@ async def _attach_local_project_row(app_config: BasicMemoryConfig, name: str, pa
db_type=db.DatabaseType.FILESYSTEM,
)
try:
repo = ProjectRepository(session_maker)
existing = await repo.get_by_name(name)
if existing is None:
await repo.create(
{
"name": name,
"path": path,
"permalink": generate_permalink(name),
"is_active": True,
}
)
return
if existing.path != path:
await repo.update_path(existing.id, path)
repo = ProjectRepository()
async with db.scoped_session(session_maker) as session:
existing = await repo.get_by_name(session, name)
if existing is None:
await repo.create(
session,
{
"name": name,
"path": path,
"permalink": generate_permalink(name),
"is_active": True,
},
)
return
if existing.path != path:
await repo.update_path(session, existing.id, path)
finally:
# CLI-only: safe to tear down the global DB singleton here since
# set-cloud/set-local never run inside a long-lived MCP/API server.
+1 -2
View File
@@ -206,14 +206,13 @@ def _create_sqlite_engine(db_url: str, db_type: DatabaseType) -> AsyncEngine:
Configured async engine for SQLite
"""
# Configure connection args with Windows-specific settings
connect_args: dict[str, bool | float | None] = {"check_same_thread": False}
connect_args: dict[str, bool | float] = {"check_same_thread": False}
# Add Windows-specific parameters to improve reliability
if os.name == "nt": # Windows
connect_args.update(
{
"timeout": 30.0, # Increase timeout to 30 seconds for Windows
"isolation_level": None, # Use autocommit mode
}
)
+4
View File
@@ -25,6 +25,8 @@ from basic_memory.deps.db import (
EngineFactoryDep,
get_session_maker,
SessionMakerDep,
get_session,
SessionDep,
)
from basic_memory.deps.projects import (
@@ -169,6 +171,8 @@ __all__ = [
"EngineFactoryDep",
"get_session_maker",
"SessionMakerDep",
"get_session",
"SessionDep",
# Projects
"get_project_repository",
"ProjectRepositoryDep",
+10
View File
@@ -5,6 +5,7 @@ This module provides database-related dependencies:
- Session dependencies for request handling
"""
from collections.abc import AsyncGenerator
from typing import Annotated
from fastapi import Depends, Request
@@ -54,3 +55,12 @@ async def get_session_maker(engine_factory: EngineFactoryDep) -> async_sessionma
SessionMakerDep = Annotated[async_sessionmaker, Depends(get_session_maker)]
async def get_session(session_maker: SessionMakerDep) -> AsyncGenerator[AsyncSession, None]:
"""Yield a request-scoped SQLAlchemy session."""
async with db.scoped_session(session_maker) as session:
yield session
SessionDep = Annotated[AsyncSession, Depends(get_session)]
+51 -35
View File
@@ -12,6 +12,7 @@ from typing import Annotated
from fastapi import Depends, HTTPException, Path, status
from basic_memory import db
from basic_memory.config import ProjectConfig
from basic_memory.deps.db import SessionMakerDep
from basic_memory.repository.project_repository import ProjectRepository
@@ -21,11 +22,9 @@ from basic_memory.utils import generate_permalink
# --- Project Repository ---
async def get_project_repository(
session_maker: SessionMakerDep,
) -> ProjectRepository:
async def get_project_repository() -> ProjectRepository:
"""Get the project repository."""
return ProjectRepository(session_maker)
return ProjectRepository()
ProjectRepositoryDep = Annotated[ProjectRepository, Depends(get_project_repository)]
@@ -41,6 +40,7 @@ ProjectPathDep = Annotated[str, Path()]
async def get_project_id(
session_maker: SessionMakerDep,
project_repository: ProjectRepositoryDep,
project: ProjectPathDep,
) -> int:
@@ -61,14 +61,17 @@ async def get_project_id(
"""
# Convert project name to permalink for lookup
project_permalink = generate_permalink(str(project))
project_obj = await project_repository.get_by_permalink(project_permalink)
if project_obj:
return project_obj.id
async with db.scoped_session(session_maker) as session:
project_obj = await project_repository.get_by_permalink(session, project_permalink)
if project_obj:
return project_obj.id
# Try by name if permalink lookup fails
project_obj = await project_repository.get_by_name(str(project)) # pragma: no cover
if project_obj: # pragma: no cover
return project_obj.id
# Try by name if permalink lookup fails
project_obj = await project_repository.get_by_name(
session, str(project)
) # pragma: no cover
if project_obj: # pragma: no cover
return project_obj.id
# Not found
raise HTTPException( # pragma: no cover
@@ -83,7 +86,9 @@ ProjectIdDep = Annotated[int, Depends(get_project_id)]
async def get_project_config(
project: ProjectPathDep, project_repository: ProjectRepositoryDep
session_maker: SessionMakerDep,
project: ProjectPathDep,
project_repository: ProjectRepositoryDep,
) -> ProjectConfig: # pragma: no cover
"""Get the current project referenced from request state.
@@ -99,9 +104,10 @@ async def get_project_config(
"""
# Convert project name to permalink for lookup
project_permalink = generate_permalink(str(project))
project_obj = await project_repository.get_by_permalink(project_permalink)
if project_obj:
return ProjectConfig(name=project_obj.name, home=pathlib.Path(project_obj.path))
async with db.scoped_session(session_maker) as session:
project_obj = await project_repository.get_by_permalink(session, project_permalink)
if project_obj:
return ProjectConfig(name=project_obj.name, home=pathlib.Path(project_obj.path))
# Not found
raise HTTPException( # pragma: no cover
@@ -116,6 +122,7 @@ ProjectConfigDep = Annotated[ProjectConfig, Depends(get_project_config)]
async def validate_project_id(
session_maker: SessionMakerDep,
project_id: int,
project_repository: ProjectRepositoryDep,
) -> int:
@@ -134,12 +141,13 @@ async def validate_project_id(
Raises:
HTTPException: If project with that ID is not found
"""
project_obj = await project_repository.get_by_id(project_id)
if not project_obj:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=f"Project with ID {project_id} not found.",
)
async with db.scoped_session(session_maker) as session:
project_obj = await project_repository.get_by_id(session, project_id)
if not project_obj:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=f"Project with ID {project_id} not found.",
)
return project_id
@@ -147,7 +155,9 @@ ProjectIdPathDep = Annotated[int, Depends(validate_project_id)]
async def get_project_config_v2(
project_id: ProjectIdPathDep, project_repository: ProjectRepositoryDep
session_maker: SessionMakerDep,
project_id: ProjectIdPathDep,
project_repository: ProjectRepositoryDep,
) -> ProjectConfig: # pragma: no cover
"""Get the project config for v2 API (uses integer project_id from path).
@@ -161,9 +171,10 @@ async def get_project_config_v2(
Raises:
HTTPException: If project is not found
"""
project_obj = await project_repository.get_by_id(project_id)
if project_obj:
return ProjectConfig(name=project_obj.name, home=pathlib.Path(project_obj.path))
async with db.scoped_session(session_maker) as session:
project_obj = await project_repository.get_by_id(session, project_id)
if project_obj:
return ProjectConfig(name=project_obj.name, home=pathlib.Path(project_obj.path))
# Not found (this should not happen since ProjectIdPathDep already validates existence)
raise HTTPException( # pragma: no cover
@@ -178,6 +189,7 @@ ProjectConfigV2Dep = Annotated[ProjectConfig, Depends(get_project_config_v2)]
async def validate_project_external_id(
session_maker: SessionMakerDep,
project_id: str,
project_repository: ProjectRepositoryDep,
) -> int:
@@ -196,12 +208,13 @@ async def validate_project_external_id(
Raises:
HTTPException: If project with that external_id is not found
"""
project_obj = await project_repository.get_by_external_id(project_id)
if not project_obj:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=f"Project with external_id '{project_id}' not found.",
)
async with db.scoped_session(session_maker) as session:
project_obj = await project_repository.get_by_external_id(session, project_id)
if not project_obj:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=f"Project with external_id '{project_id}' not found.",
)
return project_obj.id
@@ -209,7 +222,9 @@ ProjectExternalIdPathDep = Annotated[int, Depends(validate_project_external_id)]
async def get_project_config_v2_external(
project_id: ProjectExternalIdPathDep, project_repository: ProjectRepositoryDep
session_maker: SessionMakerDep,
project_id: ProjectExternalIdPathDep,
project_repository: ProjectRepositoryDep,
) -> ProjectConfig: # pragma: no cover
"""Get the project config for v2 API (uses external_id UUID from path).
@@ -223,9 +238,10 @@ async def get_project_config_v2_external(
Raises:
HTTPException: If project is not found
"""
project_obj = await project_repository.get_by_id(project_id)
if project_obj:
return ProjectConfig(name=project_obj.name, home=pathlib.Path(project_obj.path))
async with db.scoped_session(session_maker) as session:
project_obj = await project_repository.get_by_id(session, project_id)
if project_obj:
return ProjectConfig(name=project_obj.name, home=pathlib.Path(project_obj.path))
# Not found (this should not happen since ProjectExternalIdPathDep already validates)
raise HTTPException( # pragma: no cover
+9 -18
View File
@@ -30,33 +30,30 @@ from basic_memory.repository.search_repository import SearchRepository, create_s
async def get_entity_repository(
session_maker: SessionMakerDep,
project_id: ProjectIdDep,
) -> EntityRepository:
"""Create an EntityRepository instance for the current project."""
return EntityRepository(session_maker, project_id=project_id)
return EntityRepository(project_id=project_id)
EntityRepositoryDep = Annotated[EntityRepository, Depends(get_entity_repository)]
async def get_entity_repository_v2( # pragma: no cover
session_maker: SessionMakerDep,
project_id: ProjectIdPathDep,
) -> EntityRepository:
"""Create an EntityRepository instance for v2 API (uses integer project_id from path)."""
return EntityRepository(session_maker, project_id=project_id)
return EntityRepository(project_id=project_id)
EntityRepositoryV2Dep = Annotated[EntityRepository, Depends(get_entity_repository_v2)]
async def get_entity_repository_v2_external(
session_maker: SessionMakerDep,
project_id: ProjectExternalIdPathDep,
) -> EntityRepository:
"""Create an EntityRepository instance for v2 API (uses external_id from path)."""
return EntityRepository(session_maker, project_id=project_id)
return EntityRepository(project_id=project_id)
EntityRepositoryV2ExternalDep = Annotated[
@@ -68,22 +65,20 @@ EntityRepositoryV2ExternalDep = Annotated[
async def get_observation_repository(
session_maker: SessionMakerDep,
project_id: ProjectIdDep,
) -> ObservationRepository:
"""Create an ObservationRepository instance for the current project."""
return ObservationRepository(session_maker, project_id=project_id)
return ObservationRepository(project_id=project_id)
ObservationRepositoryDep = Annotated[ObservationRepository, Depends(get_observation_repository)]
async def get_observation_repository_v2( # pragma: no cover
session_maker: SessionMakerDep,
project_id: ProjectIdPathDep,
) -> ObservationRepository:
"""Create an ObservationRepository instance for v2 API."""
return ObservationRepository(session_maker, project_id=project_id)
return ObservationRepository(project_id=project_id)
ObservationRepositoryV2Dep = Annotated[
@@ -92,11 +87,10 @@ ObservationRepositoryV2Dep = Annotated[
async def get_observation_repository_v2_external(
session_maker: SessionMakerDep,
project_id: ProjectExternalIdPathDep,
) -> ObservationRepository:
"""Create an ObservationRepository instance for v2 API (uses external_id)."""
return ObservationRepository(session_maker, project_id=project_id)
return ObservationRepository(project_id=project_id)
ObservationRepositoryV2ExternalDep = Annotated[
@@ -108,33 +102,30 @@ ObservationRepositoryV2ExternalDep = Annotated[
async def get_relation_repository(
session_maker: SessionMakerDep,
project_id: ProjectIdDep,
) -> RelationRepository:
"""Create a RelationRepository instance for the current project."""
return RelationRepository(session_maker, project_id=project_id)
return RelationRepository(project_id=project_id)
RelationRepositoryDep = Annotated[RelationRepository, Depends(get_relation_repository)]
async def get_relation_repository_v2( # pragma: no cover
session_maker: SessionMakerDep,
project_id: ProjectIdPathDep,
) -> RelationRepository:
"""Create a RelationRepository instance for v2 API."""
return RelationRepository(session_maker, project_id=project_id)
return RelationRepository(project_id=project_id)
RelationRepositoryV2Dep = Annotated[RelationRepository, Depends(get_relation_repository_v2)]
async def get_relation_repository_v2_external(
session_maker: SessionMakerDep,
project_id: ProjectExternalIdPathDep,
) -> RelationRepository:
"""Create a RelationRepository instance for v2 API (uses external_id)."""
return RelationRepository(session_maker, project_id=project_id)
return RelationRepository(project_id=project_id)
RelationRepositoryV2ExternalDep = Annotated[
+59 -10
View File
@@ -16,6 +16,7 @@ from fastapi import Depends
from loguru import logger
from basic_memory.deps.config import AppConfigDep
from basic_memory.deps.db import SessionMakerDep
from basic_memory.deps.projects import (
ProjectConfigDep,
ProjectConfigV2Dep,
@@ -159,9 +160,10 @@ async def get_search_service(
search_repository: SearchRepositoryDep,
entity_repository: EntityRepositoryDep,
file_service: FileServiceDep,
session_maker: SessionMakerDep,
) -> SearchService:
"""Create SearchService with dependencies."""
return SearchService(search_repository, entity_repository, file_service)
return SearchService(search_repository, entity_repository, file_service, session_maker)
SearchServiceDep = Annotated[SearchService, Depends(get_search_service)]
@@ -171,9 +173,10 @@ async def get_search_service_v2( # pragma: no cover
search_repository: SearchRepositoryV2Dep,
entity_repository: EntityRepositoryV2Dep,
file_service: FileServiceV2Dep,
session_maker: SessionMakerDep,
) -> SearchService:
"""Create SearchService for v2 API."""
return SearchService(search_repository, entity_repository, file_service)
return SearchService(search_repository, entity_repository, file_service, session_maker)
SearchServiceV2Dep = Annotated[SearchService, Depends(get_search_service_v2)]
@@ -183,9 +186,10 @@ async def get_search_service_v2_external(
search_repository: SearchRepositoryV2ExternalDep,
entity_repository: EntityRepositoryV2ExternalDep,
file_service: FileServiceV2ExternalDep,
session_maker: SessionMakerDep,
) -> SearchService:
"""Create SearchService for v2 API (uses external_id)."""
return SearchService(search_repository, entity_repository, file_service)
return SearchService(search_repository, entity_repository, file_service, session_maker)
SearchServiceV2ExternalDep = Annotated[SearchService, Depends(get_search_service_v2_external)]
@@ -195,27 +199,45 @@ SearchServiceV2ExternalDep = Annotated[SearchService, Depends(get_search_service
async def get_link_resolver(
entity_repository: EntityRepositoryDep, search_service: SearchServiceDep
entity_repository: EntityRepositoryDep,
search_service: SearchServiceDep,
session_maker: SessionMakerDep,
) -> LinkResolver:
return LinkResolver(entity_repository=entity_repository, search_service=search_service)
return LinkResolver(
entity_repository=entity_repository,
search_service=search_service,
session_maker=session_maker,
)
LinkResolverDep = Annotated[LinkResolver, Depends(get_link_resolver)]
async def get_link_resolver_v2( # pragma: no cover
entity_repository: EntityRepositoryV2Dep, search_service: SearchServiceV2Dep
entity_repository: EntityRepositoryV2Dep,
search_service: SearchServiceV2Dep,
session_maker: SessionMakerDep,
) -> LinkResolver:
return LinkResolver(entity_repository=entity_repository, search_service=search_service)
return LinkResolver(
entity_repository=entity_repository,
search_service=search_service,
session_maker=session_maker,
)
LinkResolverV2Dep = Annotated[LinkResolver, Depends(get_link_resolver_v2)]
async def get_link_resolver_v2_external(
entity_repository: EntityRepositoryV2ExternalDep, search_service: SearchServiceV2ExternalDep
entity_repository: EntityRepositoryV2ExternalDep,
search_service: SearchServiceV2ExternalDep,
session_maker: SessionMakerDep,
) -> LinkResolver:
return LinkResolver(entity_repository=entity_repository, search_service=search_service)
return LinkResolver(
entity_repository=entity_repository,
search_service=search_service,
session_maker=session_maker,
)
LinkResolverV2ExternalDep = Annotated[LinkResolver, Depends(get_link_resolver_v2_external)]
@@ -232,6 +254,7 @@ async def get_entity_service(
file_service: FileServiceDep,
link_resolver: LinkResolverDep,
search_service: SearchServiceDep,
session_maker: SessionMakerDep,
app_config: AppConfigDep,
) -> EntityService:
"""Create EntityService with repository."""
@@ -242,6 +265,7 @@ async def get_entity_service(
entity_parser=entity_parser,
file_service=file_service,
link_resolver=link_resolver,
session_maker=session_maker,
search_service=search_service,
app_config=app_config,
)
@@ -258,6 +282,7 @@ async def get_entity_service_v2( # pragma: no cover
file_service: FileServiceV2Dep,
link_resolver: LinkResolverV2Dep,
search_service: SearchServiceV2Dep,
session_maker: SessionMakerDep,
app_config: AppConfigDep,
) -> EntityService:
"""Create EntityService for v2 API."""
@@ -268,6 +293,7 @@ async def get_entity_service_v2( # pragma: no cover
entity_parser=entity_parser,
file_service=file_service,
link_resolver=link_resolver,
session_maker=session_maker,
search_service=search_service,
app_config=app_config,
)
@@ -284,6 +310,7 @@ async def get_entity_service_v2_external(
file_service: FileServiceV2ExternalDep,
link_resolver: LinkResolverV2ExternalDep,
search_service: SearchServiceV2ExternalDep,
session_maker: SessionMakerDep,
app_config: AppConfigDep,
) -> EntityService:
"""Create EntityService for v2 API (uses external_id)."""
@@ -294,6 +321,7 @@ async def get_entity_service_v2_external(
entity_parser=entity_parser,
file_service=file_service,
link_resolver=link_resolver,
session_maker=session_maker,
search_service=search_service,
app_config=app_config,
)
@@ -310,12 +338,14 @@ async def get_context_service(
entity_repository: EntityRepositoryDep,
observation_repository: ObservationRepositoryDep,
link_resolver: LinkResolverDep,
session_maker: SessionMakerDep,
) -> ContextService:
return ContextService(
search_repository=search_repository,
entity_repository=entity_repository,
observation_repository=observation_repository,
link_resolver=link_resolver,
session_maker=session_maker,
)
@@ -327,6 +357,7 @@ async def get_context_service_v2( # pragma: no cover
entity_repository: EntityRepositoryV2Dep,
observation_repository: ObservationRepositoryV2Dep,
link_resolver: LinkResolverV2Dep,
session_maker: SessionMakerDep,
) -> ContextService:
"""Create ContextService for v2 API."""
return ContextService(
@@ -334,6 +365,7 @@ async def get_context_service_v2( # pragma: no cover
entity_repository=entity_repository,
observation_repository=observation_repository,
link_resolver=link_resolver,
session_maker=session_maker,
)
@@ -345,6 +377,7 @@ async def get_context_service_v2_external(
entity_repository: EntityRepositoryV2ExternalDep,
observation_repository: ObservationRepositoryV2ExternalDep,
link_resolver: LinkResolverV2ExternalDep,
session_maker: SessionMakerDep,
) -> ContextService:
"""Create ContextService for v2 API (uses external_id)."""
return ContextService(
@@ -352,6 +385,7 @@ async def get_context_service_v2_external(
entity_repository=entity_repository,
observation_repository=observation_repository,
link_resolver=link_resolver,
session_maker=session_maker,
)
@@ -370,6 +404,7 @@ async def get_sync_service(
project_repository: ProjectRepositoryDep,
search_service: SearchServiceDep,
file_service: FileServiceDep,
session_maker: SessionMakerDep,
) -> SyncService: # pragma: no cover
return SyncService(
app_config=app_config,
@@ -380,6 +415,7 @@ async def get_sync_service(
project_repository=project_repository,
search_service=search_service,
file_service=file_service,
session_maker=session_maker,
)
@@ -395,6 +431,7 @@ async def get_sync_service_v2(
project_repository: ProjectRepositoryDep,
search_service: SearchServiceV2Dep,
file_service: FileServiceV2Dep,
session_maker: SessionMakerDep,
) -> SyncService: # pragma: no cover
"""Create SyncService for v2 API."""
return SyncService(
@@ -406,6 +443,7 @@ async def get_sync_service_v2(
project_repository=project_repository,
search_service=search_service,
file_service=file_service,
session_maker=session_maker,
)
@@ -421,6 +459,7 @@ async def get_sync_service_v2_external(
project_repository: ProjectRepositoryDep,
search_service: SearchServiceV2ExternalDep,
file_service: FileServiceV2ExternalDep,
session_maker: SessionMakerDep,
) -> SyncService: # pragma: no cover
"""Create SyncService for v2 API (uses external_id)."""
return SyncService(
@@ -432,6 +471,7 @@ async def get_sync_service_v2_external(
project_repository=project_repository,
search_service=search_service,
file_service=file_service,
session_maker=session_maker,
)
@@ -535,6 +575,7 @@ TaskSchedulerDep = Annotated[TaskScheduler, Depends(get_task_scheduler)]
async def get_project_service(
project_repository: ProjectRepositoryDep,
session_maker: SessionMakerDep,
app_config: AppConfigDep,
) -> ProjectService:
"""Create ProjectService with repository and a system-level FileService for directory operations."""
@@ -543,7 +584,9 @@ async def get_project_service(
entity_parser = EntityParser(Path.home())
markdown_processor = MarkdownProcessor(entity_parser, app_config=app_config)
file_service = FileService(Path.home(), markdown_processor, app_config=app_config)
return ProjectService(repository=project_repository, file_service=file_service)
return ProjectService(
repository=project_repository, session_maker=session_maker, file_service=file_service
)
ProjectServiceDep = Annotated[ProjectService, Depends(get_project_service)]
@@ -554,10 +597,12 @@ ProjectServiceDep = Annotated[ProjectService, Depends(get_project_service)]
async def get_directory_service(
entity_repository: EntityRepositoryDep,
session_maker: SessionMakerDep,
) -> DirectoryService:
"""Create DirectoryService with dependencies."""
return DirectoryService(
entity_repository=entity_repository,
session_maker=session_maker,
)
@@ -566,10 +611,12 @@ DirectoryServiceDep = Annotated[DirectoryService, Depends(get_directory_service)
async def get_directory_service_v2( # pragma: no cover
entity_repository: EntityRepositoryV2Dep,
session_maker: SessionMakerDep,
) -> DirectoryService:
"""Create DirectoryService for v2 API (uses integer project_id from path)."""
return DirectoryService(
entity_repository=entity_repository,
session_maker=session_maker,
)
@@ -578,10 +625,12 @@ DirectoryServiceV2Dep = Annotated[DirectoryService, Depends(get_directory_servic
async def get_directory_service_v2_external(
entity_repository: EntityRepositoryV2ExternalDep,
session_maker: SessionMakerDep,
) -> DirectoryService:
"""Create DirectoryService for v2 API (uses external_id from path)."""
return DirectoryService(
entity_repository=entity_repository,
session_maker=session_maker,
)
+109 -70
View File
@@ -3,15 +3,18 @@
from __future__ import annotations
import asyncio
from contextlib import asynccontextmanager
from dataclasses import dataclass
from datetime import datetime
from pathlib import Path
from typing import Awaitable, Callable, Mapping, TypeVar
from typing import AsyncIterator, Awaitable, Callable, Mapping, TypeVar
from loguru import logger
from sqlalchemy.exc import IntegrityError
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
import logfire
from basic_memory import db
from basic_memory.config import BasicMemoryConfig
from basic_memory.file_utils import compute_checksum, has_frontmatter, remove_frontmatter
from basic_memory.markdown.schemas import EntityMarkdown
@@ -69,6 +72,7 @@ class BatchIndexer:
relation_repository: RelationRepository,
search_service: SearchService,
file_writer: IndexFileWriter,
session_maker: async_sessionmaker[AsyncSession],
) -> None:
self.app_config = app_config
self.entity_service = entity_service
@@ -76,6 +80,19 @@ class BatchIndexer:
self.relation_repository = relation_repository
self.search_service = search_service
self.file_writer = file_writer
self.session_maker = session_maker
@asynccontextmanager
async def _session_scope(
self, session: AsyncSession | None = None
) -> AsyncIterator[AsyncSession]:
"""Use the caller's session or open a service-owned transaction."""
if session is not None:
yield session
return
async with db.scoped_session(self.session_maker) as owned_session:
yield owned_session
async def index_files(
self,
@@ -149,9 +166,10 @@ class BatchIndexer:
max_concurrent=max_concurrent,
)
refreshed_entities = await self.entity_repository.find_by_ids(
[prepared.entity_id for prepared in prepared_entities.values()]
)
async with self._session_scope() as session:
refreshed_entities = await self.entity_repository.find_by_ids(
session, [prepared.entity_id for prepared in prepared_entities.values()]
)
entities_by_id = {entity.id: entity for entity in refreshed_entities}
refreshed, refresh_errors = await self._run_bounded(
@@ -196,12 +214,7 @@ class BatchIndexer:
prepared = await self._prepare_markdown_file(file)
if existing_permalink_by_path is None:
with logfire.span("index.markdown_file.load_permalink_map", path=file.path):
existing_permalink_by_path = {
path: permalink
for path, permalink in (
await self.entity_repository.get_file_path_to_permalink_map()
).items()
}
existing_permalink_by_path = await self._get_file_path_to_permalink_map()
reserved_permalinks = {
permalink
@@ -226,7 +239,8 @@ class BatchIndexer:
path=file.path,
entity_id=persisted.entity.id,
):
refreshed = await self.entity_repository.find_by_ids([persisted.entity.id])
async with self._session_scope() as session:
refreshed = await self.entity_repository.find_by_ids(session, [persisted.entity.id])
if len(refreshed) != 1: # pragma: no cover
raise ValueError(f"Failed to reload indexed entity for {file.path}")
entity = refreshed[0]
@@ -249,6 +263,17 @@ class BatchIndexer:
markdown_content=prepared_entity.markdown_content,
)
async def _get_file_path_to_permalink_map(self) -> dict[str, str | None]:
"""Load current file-path to permalink mappings in a service-owned session."""
async with self._session_scope() as session:
permalink_by_path: dict[str, str | None] = {
path: permalink
for path, permalink in (
await self.entity_repository.get_file_path_to_permalink_map(session)
).items()
}
return permalink_by_path
# --- Preparation ---
async def _prepare_markdown_file(self, file: IndexInputFile) -> _PreparedMarkdownFile:
@@ -283,12 +308,7 @@ class BatchIndexer:
return {}, {}
if existing_permalink_by_path is None:
existing_permalink_by_path = {
path: permalink
for path, permalink in (
await self.entity_repository.get_file_path_to_permalink_map()
).items()
}
existing_permalink_by_path = await self._get_file_path_to_permalink_map()
batch_paths = set(prepared_markdown)
reserved_permalinks = {
@@ -406,7 +426,10 @@ class BatchIndexer:
async def _upsert_regular_file(self, file: IndexInputFile) -> _PreparedEntity:
checksum = await self._resolve_checksum(file)
existing = await self.entity_repository.get_by_file_path(file.path, load_relations=False)
async with self._session_scope() as session:
existing = await self.entity_repository.get_by_file_path(
session, file.path, load_relations=False
)
is_new_entity = existing is None
if existing is None:
@@ -424,7 +447,8 @@ class BatchIndexer:
)
try:
created = await self.entity_repository.add(entity)
async with self._session_scope() as session:
created = await self.entity_repository.add(session, entity)
entity_id = created.id
except IntegrityError as exc:
message = str(exc)
@@ -436,10 +460,12 @@ class BatchIndexer:
and "file_path" in message
)
):
existing = await self.entity_repository.get_by_file_path(
file.path,
load_relations=False,
)
async with self._session_scope() as session:
existing = await self.entity_repository.get_by_file_path(
session,
file.path,
load_relations=False,
)
if existing is None:
raise ValueError(
f"Entity not found after file_path conflict: {file.path}"
@@ -450,10 +476,12 @@ class BatchIndexer:
else:
entity_id = existing.id
updated = await self.entity_repository.update(
entity_id,
self._entity_metadata_updates(file, checksum, include_created_at=is_new_entity),
)
async with self._session_scope() as session:
updated = await self.entity_repository.update(
session,
entity_id,
self._entity_metadata_updates(file, checksum, include_created_at=is_new_entity),
)
if updated is None:
raise ValueError(f"Failed to update file entity metadata for {file.path}")
@@ -476,10 +504,7 @@ class BatchIndexer:
max_concurrent: int,
) -> tuple[int, int]:
unresolved_relation_lists = await asyncio.gather(
*(
self.relation_repository.find_unresolved_relations_for_entity(entity_id)
for entity_id in entity_ids
)
*(self._find_unresolved_relations_for_entity(entity_id) for entity_id in entity_ids)
)
unresolved_relations = [
relation for relation_list in unresolved_relation_lists for relation in relation_list
@@ -499,22 +524,26 @@ class BatchIndexer:
# link text, mismatching this with the sync_service forward-reference
# path and producing confidently-wrong graph edges. See
# sync_service.resolve_forward_references for the same change.
resolved_entity = await self.entity_service.link_resolver.resolve_link(
relation.to_name, strict=True
)
async with self._session_scope() as session:
resolved_entity = await self.entity_service.link_resolver.resolve_link(
relation.to_name, strict=True, session=session
)
if resolved_entity is None or resolved_entity.id == relation.from_id:
return 0
try:
await self.relation_repository.update(
relation.id,
{
"to_id": resolved_entity.id,
"to_name": resolved_entity.title,
},
)
async with self._session_scope() as session:
await self.relation_repository.update(
session,
relation.id,
{
"to_id": resolved_entity.id,
"to_name": resolved_entity.title,
},
)
except IntegrityError:
await self.relation_repository.delete(relation.id)
async with self._session_scope() as session:
await self.relation_repository.delete(session, relation.id)
return 1
except Exception as exc: # pragma: no cover - defensive logging
logger.warning(
@@ -531,15 +560,19 @@ class BatchIndexer:
)
remaining_relation_lists = await asyncio.gather(
*(
self.relation_repository.find_unresolved_relations_for_entity(entity_id)
for entity_id in entity_ids
)
*(self._find_unresolved_relations_for_entity(entity_id) for entity_id in entity_ids)
)
remaining_unresolved = sum(len(relations) for relations in remaining_relation_lists)
return sum(resolved_counts), remaining_unresolved
async def _find_unresolved_relations_for_entity(self, entity_id: int):
"""Load unresolved relations for one entity in a service-owned session."""
async with self._session_scope() as session:
return await self.relation_repository.find_unresolved_relations_for_entity(
session, entity_id
)
# --- Search refresh ---
async def _refresh_search_index(
@@ -565,30 +598,36 @@ class BatchIndexer:
resolve_relations: bool = True,
reload_entity: bool = True,
) -> _PersistedMarkdownFile:
existing = await self.entity_repository.get_by_file_path(
prepared.file.path,
load_relations=False,
)
if is_new is None:
is_new = existing is None
entity = await self.entity_service.upsert_entity_from_markdown(
Path(prepared.file.path),
prepared.markdown,
is_new=is_new,
existing_entity=existing,
resolve_relations=resolve_relations,
reload_entity=reload_entity,
)
prepared = await self._reconcile_persisted_permalink(prepared, entity)
metadata_updates = self._entity_metadata_updates(prepared.file, prepared.final_checksum)
updated = await self.entity_repository.update_fields(
entity.id,
metadata_updates,
)
if not updated:
raise ValueError(f"Failed to update markdown entity metadata for {prepared.file.path}")
self._apply_entity_metadata_updates(entity, metadata_updates)
return _PersistedMarkdownFile(prepared=prepared, entity=entity)
async with self._session_scope() as session:
existing = await self.entity_repository.get_by_file_path(
session,
prepared.file.path,
load_relations=False,
)
if is_new is None:
is_new = existing is None
entity = await self.entity_service.upsert_entity_from_markdown(
Path(prepared.file.path),
prepared.markdown,
is_new=is_new,
existing_entity=existing,
resolve_relations=resolve_relations,
reload_entity=reload_entity,
session=session,
)
prepared = await self._reconcile_persisted_permalink(prepared, entity)
metadata_updates = self._entity_metadata_updates(prepared.file, prepared.final_checksum)
updated = await self.entity_repository.update_fields(
session,
entity.id,
metadata_updates,
)
if not updated:
raise ValueError(
f"Failed to update markdown entity metadata for {prepared.file.path}"
)
self._apply_entity_metadata_updates(entity, metadata_updates)
return _PersistedMarkdownFile(prepared=prepared, entity=entity)
async def _reconcile_persisted_permalink(
self,
+133 -133
View File
@@ -7,12 +7,11 @@ from typing import List, Optional, Sequence, Union, Any
from loguru import logger
from sqlalchemy import exists, func, select
from sqlalchemy.exc import IntegrityError
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import load_only, selectinload
from sqlalchemy.orm.interfaces import LoaderOption
from sqlalchemy.engine import Row
from basic_memory import db
from basic_memory.models.knowledge import Entity, Observation, Relation
from basic_memory.repository.repository import Repository
@@ -24,16 +23,17 @@ class EntityRepository(Repository[Entity]):
to strings before passing to repository methods.
"""
def __init__(self, session_maker: async_sessionmaker[AsyncSession], project_id: int):
"""Initialize with session maker and project_id filter.
def __init__(self, project_id: int):
"""Initialize with project_id filter.
Args:
session_maker: SQLAlchemy session maker
project_id: Project ID to filter all operations by
"""
super().__init__(session_maker, Entity, project_id=project_id)
super().__init__(Entity, project_id=project_id)
async def get_by_id(self, entity_id: int, *, load_relations: bool = True) -> Optional[Entity]:
async def get_by_id(
self, session: AsyncSession, entity_id: int, *, load_relations: bool = True
) -> Optional[Entity]:
"""Get entity by numeric ID.
Args:
@@ -42,23 +42,24 @@ class EntityRepository(Repository[Entity]):
Returns:
Entity if found, None otherwise
"""
async with db.scoped_session(self.session_maker) as session:
if not load_relations:
result = await session.execute(self.select().where(Entity.id == entity_id))
return result.scalars().one_or_none()
if not load_relations:
result = await session.execute(self.select().where(Entity.id == entity_id))
return result.scalars().one_or_none()
return await self.select_by_id(session, entity_id)
return await self.select_by_id(session, entity_id)
async def _find_one_by_query(self, query, *, load_relations: bool) -> Optional[Entity]:
async def _find_one_by_query(
self, session: AsyncSession, query, *, load_relations: bool
) -> Optional[Entity]:
"""Return one entity row with optional eager loading."""
if load_relations:
return await self.find_one(query)
return await self.find_one(session, query)
result = await self.execute_query(query, use_query_options=False)
result = await self.execute_query(session, query, use_query_options=False)
return result.scalars().one_or_none()
async def get_by_external_id(
self, external_id: str, *, load_relations: bool = True
self, session: AsyncSession, external_id: str, *, load_relations: bool = True
) -> Optional[Entity]:
"""Get entity by external UUID.
@@ -69,10 +70,10 @@ class EntityRepository(Repository[Entity]):
Entity if found, None otherwise
"""
query = self.select().where(Entity.external_id == external_id)
return await self._find_one_by_query(query, load_relations=load_relations)
return await self._find_one_by_query(session, query, load_relations=load_relations)
async def get_by_permalink(
self, permalink: str, *, load_relations: bool = True
self, session: AsyncSession, permalink: str, *, load_relations: bool = True
) -> Optional[Entity]:
"""Get entity by permalink.
@@ -80,9 +81,11 @@ class EntityRepository(Repository[Entity]):
permalink: Unique identifier for the entity
"""
query = self.select().where(Entity.permalink == permalink)
return await self._find_one_by_query(query, load_relations=load_relations)
return await self._find_one_by_query(session, query, load_relations=load_relations)
async def get_by_title(self, title: str, *, load_relations: bool = True) -> Sequence[Entity]:
async def get_by_title(
self, session: AsyncSession, title: str, *, load_relations: bool = True
) -> Sequence[Entity]:
"""Get entities by title, ordered by shortest path first.
When multiple entities share the same title (in different folders),
@@ -97,11 +100,11 @@ class EntityRepository(Repository[Entity]):
.where(Entity.title == title)
.order_by(func.length(Entity.file_path), Entity.file_path)
)
result = await self.execute_query(query, use_query_options=load_relations)
result = await self.execute_query(session, query, use_query_options=load_relations)
return list(result.scalars().all())
async def get_by_file_path(
self, file_path: Union[Path, str], *, load_relations: bool = True
self, session: AsyncSession, file_path: Union[Path, str], *, load_relations: bool = True
) -> Optional[Entity]:
"""Get entity by file_path.
@@ -109,13 +112,13 @@ class EntityRepository(Repository[Entity]):
file_path: Path to the entity file (will be converted to string internally)
"""
query = self.select().where(Entity.file_path == Path(file_path).as_posix())
return await self._find_one_by_query(query, load_relations=load_relations)
return await self._find_one_by_query(session, query, load_relations=load_relations)
# -------------------------------------------------------------------------
# Lightweight methods for permalink resolution (no eager loading)
# -------------------------------------------------------------------------
async def permalink_exists(self, permalink: str) -> bool:
async def permalink_exists(self, session: AsyncSession, permalink: str) -> bool:
"""Check if a permalink exists without loading the full entity.
This is much faster than get_by_permalink() as it skips eager loading
@@ -129,10 +132,12 @@ class EntityRepository(Repository[Entity]):
"""
query = select(Entity.id).where(Entity.permalink == permalink).limit(1)
query = self._add_project_filter(query)
result = await self.execute_query(query, use_query_options=False)
result = await self.execute_query(session, query, use_query_options=False)
return result.scalar_one_or_none() is not None
async def get_file_path_for_permalink(self, permalink: str) -> Optional[str]:
async def get_file_path_for_permalink(
self, session: AsyncSession, permalink: str
) -> Optional[str]:
"""Get the file_path for a permalink without loading the full entity.
Use when you only need the file_path, not the full entity with relations.
@@ -145,10 +150,12 @@ class EntityRepository(Repository[Entity]):
"""
query = select(Entity.file_path).where(Entity.permalink == permalink)
query = self._add_project_filter(query)
result = await self.execute_query(query, use_query_options=False)
result = await self.execute_query(session, query, use_query_options=False)
return result.scalar_one_or_none()
async def get_permalink_for_file_path(self, file_path: Union[Path, str]) -> Optional[str]:
async def get_permalink_for_file_path(
self, session: AsyncSession, file_path: Union[Path, str]
) -> Optional[str]:
"""Get the permalink for a file_path without loading the full entity.
Use when you only need the permalink, not the full entity with relations.
@@ -161,10 +168,10 @@ class EntityRepository(Repository[Entity]):
"""
query = select(Entity.permalink).where(Entity.file_path == Path(file_path).as_posix())
query = self._add_project_filter(query)
result = await self.execute_query(query, use_query_options=False)
result = await self.execute_query(session, query, use_query_options=False)
return result.scalar_one_or_none()
async def get_all_permalinks(self) -> List[str]:
async def get_all_permalinks(self, session: AsyncSession) -> List[str]:
"""Get all permalinks for this project.
Optimized for bulk operations - returns only permalink strings
@@ -175,11 +182,11 @@ class EntityRepository(Repository[Entity]):
"""
query = select(Entity.permalink)
query = self._add_project_filter(query)
result = await self.execute_query(query, use_query_options=False)
result = await self.execute_query(session, query, use_query_options=False)
return list(result.scalars().all())
async def find_by_ids_for_hydration(
self, ids: List[int], *, include_cross_project: bool = False
self, session: AsyncSession, ids: List[int], *, include_cross_project: bool = False
) -> Sequence[Entity]:
"""Fetch minimal entity fields needed for context hydration.
@@ -203,10 +210,10 @@ class EntityRepository(Repository[Entity]):
if not include_cross_project:
query = self._add_project_filter(query)
result = await self.execute_query(query, use_query_options=False)
result = await self.execute_query(session, query, use_query_options=False)
return list(result.scalars().all())
async def get_permalink_to_file_path_map(self) -> dict[str, str]:
async def get_permalink_to_file_path_map(self, session: AsyncSession) -> dict[str, str]:
"""Get a mapping of permalink -> file_path for all entities.
Optimized for bulk permalink resolution - loads minimal data in one query.
@@ -216,10 +223,10 @@ class EntityRepository(Repository[Entity]):
"""
query = select(Entity.permalink, Entity.file_path)
query = self._add_project_filter(query)
result = await self.execute_query(query, use_query_options=False)
result = await self.execute_query(session, query, use_query_options=False)
return {row.permalink: row.file_path for row in result.all()}
async def get_file_path_to_permalink_map(self) -> dict[str, str]:
async def get_file_path_to_permalink_map(self, session: AsyncSession) -> dict[str, str]:
"""Get a mapping of file_path -> permalink for all entities.
Optimized for bulk permalink resolution - loads minimal data in one query.
@@ -229,7 +236,7 @@ class EntityRepository(Repository[Entity]):
"""
query = select(Entity.file_path, Entity.permalink)
query = self._add_project_filter(query)
result = await self.execute_query(query, use_query_options=False)
result = await self.execute_query(session, query, use_query_options=False)
return {row.file_path: row.permalink for row in result.all()}
async def get_by_file_paths(
@@ -262,7 +269,7 @@ class EntityRepository(Repository[Entity]):
result = await session.execute(query) # pragma: no cover
return list(result.all()) # pragma: no cover
async def find_by_checksum(self, checksum: str) -> Sequence[Entity]:
async def find_by_checksum(self, session: AsyncSession, checksum: str) -> Sequence[Entity]:
"""Find entities with the given checksum.
Used for move detection - finds entities that may have been moved to a new path.
@@ -276,10 +283,12 @@ class EntityRepository(Repository[Entity]):
"""
query = self.select().where(Entity.checksum == checksum)
# Don't load relationships for move detection - we only need file_path and checksum
result = await self.execute_query(query, use_query_options=False)
result = await self.execute_query(session, query, use_query_options=False)
return list(result.scalars().all())
async def find_by_checksums(self, checksums: Sequence[str]) -> Sequence[Entity]:
async def find_by_checksums(
self, session: AsyncSession, checksums: Sequence[str]
) -> Sequence[Entity]:
"""Find entities with any of the given checksums (batch query for move detection).
This is a batch-optimized version of find_by_checksum() that queries multiple checksums
@@ -304,16 +313,18 @@ class EntityRepository(Repository[Entity]):
# Query: SELECT * FROM entities WHERE checksum IN (checksum1, checksum2, ...)
query = self.select().where(Entity.checksum.in_(checksums)) # pragma: no cover
# Don't load relationships for move detection - we only need file_path and checksum
result = await self.execute_query(query, use_query_options=False) # pragma: no cover
result = await self.execute_query(
session, query, use_query_options=False
) # pragma: no cover
return list(result.scalars().all()) # pragma: no cover
async def delete_by_file_path(self, file_path: Union[Path, str]) -> bool:
async def delete_by_file_path(self, session: AsyncSession, file_path: Union[Path, str]) -> bool:
"""Delete entity with the provided file_path.
Args:
file_path: Path to the entity file (will be converted to string internally)
"""
return await self.delete_by_fields(file_path=Path(file_path).as_posix())
return await self.delete_by_fields(session, file_path=Path(file_path).as_posix())
def get_load_options(self) -> List[LoaderOption]:
"""Get SQLAlchemy loader options for eager loading relationships."""
@@ -327,7 +338,9 @@ class EntityRepository(Repository[Entity]):
selectinload(Entity.incoming_relations).selectinload(Relation.to_entity),
]
async def find_by_permalinks(self, permalinks: List[str]) -> Sequence[Entity]:
async def find_by_permalinks(
self, session: AsyncSession, permalinks: List[str]
) -> Sequence[Entity]:
"""Find multiple entities by their permalink.
Args:
@@ -342,10 +355,10 @@ class EntityRepository(Repository[Entity]):
self.select().options(*self.get_load_options()).where(Entity.permalink.in_(permalinks))
)
result = await self.execute_query(query)
result = await self.execute_query(session, query)
return list(result.scalars().all())
async def upsert_entity(self, entity: Entity) -> Entity:
async def upsert_entity(self, session: AsyncSession, entity: Entity) -> Entity:
"""Insert or update entity using simple try/catch with database-level conflict resolution.
Handles file_path race conditions by checking for existing entity on IntegrityError.
@@ -357,99 +370,84 @@ class EntityRepository(Repository[Entity]):
Returns:
The inserted or updated entity
"""
async with db.scoped_session(self.session_maker) as session:
# Set project_id if applicable and not already set
self._set_project_id_if_needed(entity)
# Set project_id if applicable and not already set
self._set_project_id_if_needed(entity)
# Try simple insert first
try:
# Try simple insert first.
try:
async with session.begin_nested():
session.add(entity)
await session.flush()
except IntegrityError as e:
# Check if this is a FOREIGN KEY constraint failure
# SQLite: "FOREIGN KEY constraint failed"
# Postgres: "violates foreign key constraint"
error_str = str(e)
if (
"FOREIGN KEY constraint failed" in error_str
or "violates foreign key constraint" in error_str
):
# Import locally to avoid circular dependency (repository -> services -> repository)
from basic_memory.services.exceptions import SyncFatalError
# Return with relationships loaded
query = (
self.select()
.where(Entity.file_path == entity.file_path)
.options(*self.get_load_options())
# Project doesn't exist in database - this is a fatal sync error
raise SyncFatalError(
f"Cannot sync file '{entity.file_path}': "
f"project_id={entity.project_id} does not exist in database. "
f"The project may have been deleted. This sync will be terminated."
) from e
# Re-query after the nested rollback to get a fresh, attached entity.
existing_result = await session.execute(
select(Entity)
.where(Entity.file_path == entity.file_path, Entity.project_id == entity.project_id)
.options(*self.get_load_options())
)
existing_entity = existing_result.scalar_one_or_none()
if existing_entity:
# File path conflict - update the existing entity
logger.debug(
f"Resolving file_path conflict for {entity.file_path}, "
f"entity_id={existing_entity.id}, observations={len(entity.observations)}"
)
result = await session.execute(query)
found = result.scalar_one_or_none()
if not found: # pragma: no cover
raise RuntimeError(
f"Failed to retrieve entity after insert: {entity.file_path}"
)
return found
# Use merge to avoid session state conflicts. Preserve stable
# external_id so external references survive re-indexing.
entity.id = existing_entity.id
entity.external_id = existing_entity.external_id
except IntegrityError as e:
# Check if this is a FOREIGN KEY constraint failure
# SQLite: "FOREIGN KEY constraint failed"
# Postgres: "violates foreign key constraint"
error_str = str(e)
if (
"FOREIGN KEY constraint failed" in error_str
or "violates foreign key constraint" in error_str
):
# Import locally to avoid circular dependency (repository -> services -> repository)
from basic_memory.services.exceptions import SyncFatalError
# Ensure observations reference the correct entity_id.
for obs in entity.observations:
obs.entity_id = existing_entity.id
obs.id = None
# Project doesn't exist in database - this is a fatal sync error
raise SyncFatalError(
f"Cannot sync file '{entity.file_path}': "
f"project_id={entity.project_id} does not exist in database. "
f"The project may have been deleted. This sync will be terminated."
) from e
merged_entity = await session.merge(entity)
await session.flush()
await session.rollback()
# Re-query after rollback to get a fresh, attached entity
existing_result = await session.execute(
# Re-query to get proper relationships loaded.
final_result = await session.execute(
select(Entity)
.where(
Entity.file_path == entity.file_path, Entity.project_id == entity.project_id
)
.where(Entity.id == merged_entity.id)
.options(*self.get_load_options())
)
existing_entity = existing_result.scalar_one_or_none()
return final_result.scalar_one()
if existing_entity:
# File path conflict - update the existing entity
logger.debug(
f"Resolving file_path conflict for {entity.file_path}, "
f"entity_id={existing_entity.id}, observations={len(entity.observations)}"
)
# Use merge to avoid session state conflicts
# Set the ID to update existing entity
entity.id = existing_entity.id
# Preserve the stable external_id so that external references
# (e.g. public share links) survive re-indexing
entity.external_id = existing_entity.external_id
# No file_path conflict - must be permalink conflict.
return await self._handle_permalink_conflict(entity, session)
# Ensure observations reference the correct entity_id
for obs in entity.observations:
obs.entity_id = existing_entity.id
# Clear any existing ID to force INSERT as new observation
obs.id = None
# Return with relationships loaded.
query = (
self.select()
.where(Entity.file_path == entity.file_path)
.options(*self.get_load_options())
)
result = await session.execute(query)
found = result.scalar_one_or_none()
if not found: # pragma: no cover
raise RuntimeError(f"Failed to retrieve entity after insert: {entity.file_path}")
return found
# Merge the entity which will update the existing one
merged_entity = await session.merge(entity)
await session.commit()
# Re-query to get proper relationships loaded
final_result = await session.execute(
select(Entity)
.where(Entity.id == merged_entity.id)
.options(*self.get_load_options())
)
return final_result.scalar_one()
else:
# No file_path conflict - must be permalink conflict
# Generate unique permalink and retry
entity = await self._handle_permalink_conflict(entity, session)
return entity
async def get_all_file_paths(self) -> List[str]:
async def get_all_file_paths(self, session: AsyncSession) -> List[str]:
"""Get all file paths for this project - optimized for deletion detection.
Returns only file_path strings without loading entities or relationships.
@@ -461,10 +459,10 @@ class EntityRepository(Repository[Entity]):
query = select(Entity.file_path)
query = self._add_project_filter(query)
result = await self.execute_query(query, use_query_options=False)
result = await self.execute_query(session, query, use_query_options=False)
return list(result.scalars().all())
async def find_without_relations(self) -> Sequence[Entity]:
async def find_without_relations(self, session: AsyncSession) -> Sequence[Entity]:
"""Find entities that have no incoming or outgoing relations."""
# Trigger: entity appears as a source in any relation.
# Why: even unresolved outgoing links mean the entity references another node.
@@ -477,10 +475,10 @@ class EntityRepository(Repository[Entity]):
has_incoming = exists().where(Relation.to_id == Entity.id)
query = self.select().where(~has_outgoing).where(~has_incoming).order_by(Entity.file_path)
result = await self.execute_query(query, use_query_options=False)
result = await self.execute_query(session, query, use_query_options=False)
return list(result.scalars().all())
async def get_distinct_directories(self) -> List[str]:
async def get_distinct_directories(self, session: AsyncSession) -> List[str]:
"""Extract unique directory paths from file_path column.
Optimized method for getting directory structure without loading full entities
@@ -494,7 +492,7 @@ class EntityRepository(Repository[Entity]):
query = self._add_project_filter(query)
# Execute with use_query_options=False to skip eager loading
result = await self.execute_query(query, use_query_options=False)
result = await self.execute_query(session, query, use_query_options=False)
file_paths = [row for row in result.scalars().all()]
# Parse file paths to extract unique directories
@@ -508,7 +506,9 @@ class EntityRepository(Repository[Entity]):
return sorted(directories)
async def find_by_directory_prefix(self, directory_prefix: str) -> Sequence[Entity]:
async def find_by_directory_prefix(
self, session: AsyncSession, directory_prefix: str
) -> Sequence[Entity]:
"""Find entities whose file_path starts with the given directory prefix.
Optimized method for listing directory contents without loading all entities.
@@ -524,7 +524,7 @@ class EntityRepository(Repository[Entity]):
# Build SQL LIKE pattern
if directory_prefix == "" or directory_prefix == "/":
# Root directory - return all entities
return await self.find_all()
return await self.find_all(session)
# Remove leading/trailing slashes for consistency
directory_prefix = directory_prefix.strip("/")
@@ -537,7 +537,7 @@ class EntityRepository(Repository[Entity]):
query = self.select().where(Entity.file_path.like(pattern))
# Skip eager loading - we only need basic entity fields for directory trees
result = await self.execute_query(query, use_query_options=False)
result = await self.execute_query(session, query, use_query_options=False)
return list(result.scalars().all())
async def _handle_permalink_conflict(self, entity: Entity, session: AsyncSession) -> Entity:
@@ -4,9 +4,8 @@ from pathlib import Path
from typing import Any, Mapping, Optional
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
from sqlalchemy.ext.asyncio import AsyncSession
from basic_memory import db
from basic_memory.models import Entity, NoteContent
from basic_memory.repository.repository import Repository
@@ -30,9 +29,9 @@ NOTE_CONTENT_MUTABLE_FIELDS = frozenset(
class NoteContentRepository(Repository[NoteContent]):
"""Repository for project-scoped note materialization state."""
def __init__(self, session_maker: async_sessionmaker[AsyncSession], project_id: int):
"""Initialize with session maker and project-scoped filtering."""
super().__init__(session_maker, NoteContent, project_id=project_id)
def __init__(self, project_id: int):
"""Initialize with project-scoped filtering."""
super().__init__(NoteContent, project_id=project_id)
def _coerce_note_content(
self, data: Mapping[str, Any] | NoteContent
@@ -75,16 +74,22 @@ class NoteContentRepository(Repository[NoteContent]):
note_content.external_id = entity.external_id
note_content.file_path = Path(entity.file_path).as_posix()
async def get_by_entity_id(self, entity_id: int) -> Optional[NoteContent]:
async def get_by_entity_id(
self, session: AsyncSession, entity_id: int
) -> Optional[NoteContent]:
"""Get note content by the owning entity identifier."""
return await self.find_by_id(entity_id)
return await self.find_by_id(session, entity_id)
async def get_by_external_id(self, external_id: str) -> Optional[NoteContent]:
async def get_by_external_id(
self, session: AsyncSession, external_id: str
) -> Optional[NoteContent]:
"""Get note content by the mirrored entity external identifier."""
query = self.select().where(NoteContent.external_id == external_id)
return await self.find_one(query)
return await self.find_one(session, query)
async def get_by_file_path(self, file_path: Path | str) -> Optional[NoteContent]:
async def get_by_file_path(
self, session: AsyncSession, file_path: Path | str
) -> Optional[NoteContent]:
"""Get note content by file path, preferring rows whose entity still owns that path."""
normalized_path = Path(file_path).as_posix()
@@ -104,88 +109,88 @@ class NoteContentRepository(Repository[NoteContent]):
.options(*self.get_load_options())
)
async with db.scoped_session(self.session_maker) as session:
result = await session.execute(query)
return result.scalars().first()
result = await session.execute(query)
return result.scalars().first()
async def create(self, data: Mapping[str, Any] | NoteContent) -> NoteContent:
async def create(
self, session: AsyncSession, data: Mapping[str, Any] | NoteContent
) -> NoteContent:
"""Create a note_content row aligned to its owning entity."""
note_content, _ = self._coerce_note_content(data)
async with db.scoped_session(self.session_maker) as session:
await self._align_identity_fields(session, note_content)
session.add(note_content)
await session.flush()
await self._align_identity_fields(session, note_content)
session.add(note_content)
await session.flush()
created = await self.select_by_id(session, note_content.entity_id)
if created is None: # pragma: no cover
raise ValueError(
f"Can't find NoteContent for entity {note_content.entity_id} after add"
)
return created
created = await self.select_by_id(session, note_content.entity_id)
if created is None: # pragma: no cover
raise ValueError(
f"Can't find NoteContent for entity {note_content.entity_id} after add"
)
return created
async def upsert(self, data: Mapping[str, Any] | NoteContent) -> NoteContent:
async def upsert(
self, session: AsyncSession, data: Mapping[str, Any] | NoteContent
) -> NoteContent:
"""Insert or update note_content while keeping mirrored identity fields in sync."""
note_content, provided_fields = self._coerce_note_content(data)
async with db.scoped_session(self.session_maker) as session:
await self._align_identity_fields(session, note_content)
existing = await self.select_by_id(session, note_content.entity_id)
if existing is None:
session.add(note_content)
await session.flush()
created = await self.select_by_id(session, note_content.entity_id)
if created is None: # pragma: no cover
raise ValueError(
f"Can't find NoteContent for entity {note_content.entity_id} after upsert"
)
return created
fields_to_update = (provided_fields - {"entity_id"}) | {
"project_id",
"external_id",
"file_path",
}
for column_name in fields_to_update:
setattr(existing, column_name, getattr(note_content, column_name))
await self._align_identity_fields(session, note_content)
existing = await self.select_by_id(session, note_content.entity_id)
if existing is None:
session.add(note_content)
await session.flush()
updated = await self.select_by_id(session, existing.entity_id)
if updated is None: # pragma: no cover
created = await self.select_by_id(session, note_content.entity_id)
if created is None: # pragma: no cover
raise ValueError(
f"Can't find NoteContent for entity {existing.entity_id} after upsert"
f"Can't find NoteContent for entity {note_content.entity_id} after upsert"
)
return updated
return created
async def update_state_fields(self, entity_id: int, **updates: Any) -> Optional[NoteContent]:
fields_to_update = (provided_fields - {"entity_id"}) | {
"project_id",
"external_id",
"file_path",
}
for column_name in fields_to_update:
setattr(existing, column_name, getattr(note_content, column_name))
await session.flush()
updated = await self.select_by_id(session, existing.entity_id)
if updated is None: # pragma: no cover
raise ValueError(f"Can't find NoteContent for entity {existing.entity_id} after upsert")
return updated
async def update_state_fields(
self, session: AsyncSession, entity_id: int, **updates: Any
) -> Optional[NoteContent]:
"""Update sync fields and re-align project_id, external_id, and file_path from entity."""
invalid_fields = set(updates) - NOTE_CONTENT_MUTABLE_FIELDS
if invalid_fields:
invalid_list = ", ".join(sorted(invalid_fields))
raise ValueError(f"Unsupported note_content update fields: {invalid_list}")
async with db.scoped_session(self.session_maker) as session:
note_content = await self.select_by_id(session, entity_id)
if note_content is None:
return None
note_content = await self.select_by_id(session, entity_id)
if note_content is None:
return None
await self._align_identity_fields(session, note_content)
for field_name, value in updates.items():
setattr(note_content, field_name, value)
await self._align_identity_fields(session, note_content)
for field_name, value in updates.items():
setattr(note_content, field_name, value)
await session.flush()
updated = await self.select_by_id(session, entity_id)
if updated is None: # pragma: no cover
raise ValueError(f"Can't find NoteContent for entity {entity_id} after update")
return updated
await session.flush()
updated = await self.select_by_id(session, entity_id)
if updated is None: # pragma: no cover
raise ValueError(f"Can't find NoteContent for entity {entity_id} after update")
return updated
async def delete_by_entity_id(self, entity_id: int) -> bool:
async def delete_by_entity_id(self, session: AsyncSession, entity_id: int) -> bool:
"""Delete note_content by entity identifier."""
async with db.scoped_session(self.session_maker) as session:
note_content = await self.select_by_id(session, entity_id)
if note_content is None:
return False
note_content = await self.select_by_id(session, entity_id)
if note_content is None:
return False
await session.delete(note_content)
return True
await session.delete(note_content)
await session.flush()
return True
@@ -3,7 +3,7 @@
from typing import Dict, List, Sequence
from sqlalchemy import select
from sqlalchemy.ext.asyncio import async_sessionmaker
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import selectinload
from sqlalchemy.orm.interfaces import LoaderOption
@@ -14,44 +14,46 @@ from basic_memory.repository.repository import Repository
class ObservationRepository(Repository[Observation]):
"""Repository for Observation model with memory-specific operations."""
def __init__(self, session_maker: async_sessionmaker, project_id: int):
"""Initialize with session maker and project_id filter.
def __init__(self, project_id: int):
"""Initialize with project_id filter.
Args:
session_maker: SQLAlchemy session maker
project_id: Project ID to filter all operations by
"""
super().__init__(session_maker, Observation, project_id=project_id)
super().__init__(Observation, project_id=project_id)
def get_load_options(self) -> List[LoaderOption]:
"""Eager-load parent entity to prevent N+1 if obs.entity is accessed."""
return [selectinload(Observation.entity)]
async def find_by_entity(self, entity_id: int) -> Sequence[Observation]:
async def find_by_entity(self, session: AsyncSession, entity_id: int) -> Sequence[Observation]:
"""Find all observations for a specific entity."""
query = select(Observation).filter(Observation.entity_id == entity_id)
result = await self.execute_query(query)
query = self.select().filter(Observation.entity_id == entity_id)
result = await self.execute_query(session, query)
return result.scalars().all()
async def find_by_context(self, context: str) -> Sequence[Observation]:
async def find_by_context(self, session: AsyncSession, context: str) -> Sequence[Observation]:
"""Find observations with a specific context."""
query = select(Observation).filter(Observation.context == context)
result = await self.execute_query(query)
query = self.select().filter(Observation.context == context)
result = await self.execute_query(session, query)
return result.scalars().all()
async def find_by_category(self, category: str) -> Sequence[Observation]:
async def find_by_category(self, session: AsyncSession, category: str) -> Sequence[Observation]:
"""Find observations with a specific context."""
query = select(Observation).filter(Observation.category == category)
result = await self.execute_query(query)
query = self.select().filter(Observation.category == category)
result = await self.execute_query(session, query)
return result.scalars().all()
async def observation_categories(self) -> Sequence[str]:
async def observation_categories(self, session: AsyncSession) -> Sequence[str]:
"""Return a list of all observation categories."""
query = select(Observation.category).distinct()
result = await self.execute_query(query, use_query_options=False)
query = self._add_project_filter(query)
result = await self.execute_query(session, query, use_query_options=False)
return result.scalars().all()
async def find_by_entities(self, entity_ids: List[int]) -> Dict[int, List[Observation]]:
async def find_by_entities(
self, session: AsyncSession, entity_ids: List[int]
) -> Dict[int, List[Observation]]:
"""Find all observations for multiple entities in a single query.
Args:
@@ -64,8 +66,8 @@ class ObservationRepository(Repository[Observation]):
return {}
# Query observations for all entities in the list
query = select(Observation).filter(Observation.entity_id.in_(entity_ids))
result = await self.execute_query(query)
query = self.select().filter(Observation.entity_id.in_(entity_ids))
result = await self.execute_query(session, query)
observations = result.scalars().all()
# Group observations by entity_id
@@ -18,6 +18,7 @@ from basic_memory.repository.search_index_row import SearchIndexRow
from basic_memory.repository.search_repository_base import (
SearchRepositoryBase,
VectorChunkState,
relaxed_query_words,
)
from basic_memory.repository.metadata_filters import parse_metadata_filters
from basic_memory.repository.semantic_errors import SemanticDependenciesMissingError
@@ -176,6 +177,14 @@ class PostgresSearchRepository(SearchRepositoryBase):
# For non-Boolean queries, prepare single term
return self._prepare_single_term(term, is_prefix)
@staticmethod
def _relaxed_tsquery_text(search_text: Optional[str]) -> Optional[str]:
"""OR-relaxed tsquery expression for a failed strict query, or None."""
words = relaxed_query_words(search_text)
if not words:
return None
return " | ".join(f"{word}:*" for word in words)
def _prepare_boolean_query(self, query: str) -> str:
"""Convert Boolean query to tsquery format.
@@ -236,7 +245,12 @@ class PostgresSearchRepository(SearchRepositoryBase):
# Handle multi-word queries
if " " in cleaned_term:
words = [w for w in cleaned_term.split() if w.strip()]
# Strip sentence punctuation from word edges so question-form
# queries produce clean lexemes (parity with SQLite FTS5 prep).
# The tsquery tokenizer ignores this punctuation anyway; leaving it
# in only risks tsquery syntax errors. Interior characters are kept.
words = [w.strip("?!.,;") for w in cleaned_term.split()]
words = [w for w in words if w]
if not words:
# All characters were special chars, search won't match anything
# Return a safe search term that won't cause syntax errors
@@ -249,8 +263,11 @@ class PostgresSearchRepository(SearchRepositoryBase):
# Join with AND operator
return " & ".join(prepared_words)
# Single word
cleaned_term = cleaned_term.strip()
# Single word: strip edge punctuation; guard the now-empty case so a
# bare ":*"/"" never reaches tsquery.
cleaned_term = cleaned_term.strip().strip("?!.,;")
if not cleaned_term:
return "NOSPECIALCHARS:*"
if is_prefix:
return f"{cleaned_term}:*"
else:
@@ -908,6 +925,8 @@ class PostgresSearchRepository(SearchRepositoryBase):
min_similarity: Optional[float] = None,
limit: int = 10,
offset: int = 0,
allow_relaxed: bool = False,
session: AsyncSession | None = None,
) -> List[SearchIndexRow]:
"""Search across all indexed content using PostgreSQL tsvector."""
# --- Dispatch vector / hybrid modes (shared logic) ---
@@ -978,10 +997,32 @@ class PostgresSearchRepository(SearchRepositoryBase):
"""
logger.trace(f"Search {sql} params: {params}")
try:
async with db.scoped_session(self.session_maker) as session:
result = await session.execute(text(sql), params)
async def run_search(active_session: AsyncSession):
result = await active_session.execute(text(sql), params)
rows = result.fetchall()
# Trigger: multi-word natural-language query matched nothing
# under the default all-terms-AND tsquery semantics.
# Why: questions rarely have every word in one document;
# without relaxation the FTS half of hybrid search contributes
# zero candidates (parity with the SQLite path).
# Outcome: one retry with OR-joined prefix lexemes; ts_rank
# still ranks multi-term matches first.
relaxed = (
self._relaxed_tsquery_text(search_text) if allow_relaxed and not rows else None
)
if relaxed and params.get("text"):
params["text"] = relaxed
result = await active_session.execute(text(sql), params)
rows = result.fetchall()
return rows
try:
if session is not None:
rows = await run_search(session)
else:
async with db.scoped_session(self.session_maker) as owned_session:
rows = await run_search(owned_session)
except Exception as e:
if self._is_tsquery_syntax_error(e):
logger.warning(f"tsquery syntax error for search term: {search_text}, error: {e}")
@@ -5,6 +5,6 @@ from basic_memory.models.project import Project
class ProjectInfoRepository(Repository):
"""Repository for statistics queries."""
def __init__(self, session_maker):
def __init__(self):
# Initialize with Project model as a reference
super().__init__(session_maker, Project)
super().__init__(Project)
+101 -102
View File
@@ -7,9 +7,8 @@ from typing import Optional, Sequence, Union
from loguru import logger
from sqlalchemy import Executable, inspect as sa_inspect, select, text
from sqlalchemy.exc import NoResultFound, OperationalError
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
from sqlalchemy.ext.asyncio import AsyncSession
from basic_memory import db
from basic_memory.models.project import Project
from basic_memory.repository.repository import Repository
@@ -79,20 +78,22 @@ class ProjectRepository(Repository[Project]):
Each entity, observation, and relation belongs to a specific project.
"""
def __init__(self, session_maker: async_sessionmaker[AsyncSession]):
"""Initialize with session maker."""
super().__init__(session_maker, Project)
def __init__(self):
"""Initialize the repository."""
super().__init__(Project)
async def get_by_name(self, name: str) -> Optional[Project]:
async def get_by_name(self, session: AsyncSession, name: str) -> Optional[Project]:
"""Get project by name (exact match).
Args:
name: Unique name of the project
"""
query = self.select().where(Project.name == name)
return await self.find_one(query)
return await self.find_one(session, query)
async def get_by_name_case_insensitive(self, name: str) -> Optional[Project]:
async def get_by_name_case_insensitive(
self, session: AsyncSession, name: str
) -> Optional[Project]:
"""Get project by name (case-insensitive match).
Args:
@@ -102,27 +103,27 @@ class ProjectRepository(Repository[Project]):
Project if found, None otherwise
"""
query = self.select().where(Project.name.ilike(name))
return await self.find_one(query)
return await self.find_one(session, query)
async def get_by_permalink(self, permalink: str) -> Optional[Project]:
async def get_by_permalink(self, session: AsyncSession, permalink: str) -> Optional[Project]:
"""Get project by permalink.
Args:
permalink: URL-friendly identifier for the project
"""
query = self.select().where(Project.permalink == permalink)
return await self.find_one(query)
return await self.find_one(session, query)
async def get_by_path(self, path: Union[Path, str]) -> Optional[Project]:
async def get_by_path(self, session: AsyncSession, path: Union[Path, str]) -> Optional[Project]:
"""Get project by filesystem path.
Args:
path: Path to the project directory (will be converted to string internally)
"""
query = self.select().where(Project.path == Path(path).as_posix())
return await self.find_one(query)
return await self.find_one(session, query)
async def get_by_id(self, project_id: int) -> Optional[Project]:
async def get_by_id(self, session: AsyncSession, project_id: int) -> Optional[Project]:
"""Get project by numeric ID.
Args:
@@ -131,10 +132,11 @@ class ProjectRepository(Repository[Project]):
Returns:
Project if found, None otherwise
"""
async with db.scoped_session(self.session_maker) as session:
return await self.select_by_id(session, project_id)
return await self.select_by_id(session, project_id)
async def get_by_external_id(self, external_id: str) -> Optional[Project]:
async def get_by_external_id(
self, session: AsyncSession, external_id: str
) -> Optional[Project]:
"""Get project by external UUID.
Args:
@@ -144,20 +146,20 @@ class ProjectRepository(Repository[Project]):
Project if found, None otherwise
"""
query = self.select().where(Project.external_id == external_id)
return await self.find_one(query)
return await self.find_one(session, query)
async def get_default_project(self) -> Optional[Project]:
async def get_default_project(self, session: AsyncSession) -> Optional[Project]:
"""Get the default project (the one marked as is_default=True)."""
query = self.select().where(Project.is_default.is_(True))
return await self.find_one(query)
return await self.find_one(session, query)
async def get_active_projects(self) -> Sequence[Project]:
async def get_active_projects(self, session: AsyncSession) -> Sequence[Project]:
"""Get all active projects."""
query = self.select().where(Project.is_active == True) # noqa: E712
result = await self.execute_query(query)
result = await self.execute_query(session, query)
return list(result.scalars().all())
async def set_as_default(self, project_id: int) -> Optional[Project]:
async def set_as_default(self, session: AsyncSession, project_id: int) -> Optional[Project]:
"""Set a project as the default and unset previous default.
Args:
@@ -166,22 +168,21 @@ class ProjectRepository(Repository[Project]):
Returns:
The updated project if found, None otherwise
"""
async with db.scoped_session(self.session_maker) as session:
# First, clear the default flag for all projects using direct SQL
await session.execute(
text("UPDATE project SET is_default = NULL WHERE is_default IS NOT NULL")
)
# First, clear the default flag for all projects using direct SQL
await session.execute(
text("UPDATE project SET is_default = NULL WHERE is_default IS NOT NULL")
)
await session.flush()
# Set the new default project
target_project = await self.select_by_id(session, project_id)
if target_project:
target_project.is_default = True
await session.flush()
return target_project
return None # pragma: no cover
# Set the new default project
target_project = await self.select_by_id(session, project_id)
if target_project:
target_project.is_default = True
await session.flush()
return target_project
return None # pragma: no cover
async def delete(self, entity_id: int) -> bool:
async def delete(self, session: AsyncSession, entity_id: int) -> bool:
"""Delete a project and its derived search rows in one transaction.
The cascade picture differs by backend:
@@ -204,62 +205,60 @@ class ProjectRepository(Repository[Project]):
Inspect the connection once and skip whichever is missing.
"""
logger.debug(f"Deleting Project and search rows for project_id: {entity_id}")
async with db.scoped_session(self.session_maker) as session:
try:
result = await session.execute(
select(self.Model).filter(self.primary_key == entity_id)
)
project = result.scalars().one()
except NoResultFound:
logger.debug(f"No Project found to delete: {entity_id}")
return False
try:
result = await session.execute(select(self.Model).filter(self.primary_key == entity_id))
project = result.scalars().one()
except NoResultFound:
logger.debug(f"No Project found to delete: {entity_id}")
return False
dialect_name = session.bind.dialect.name if session.bind else "sqlite"
is_sqlite = dialect_name == "sqlite"
dialect_name = session.bind.dialect.name if session.bind else "sqlite"
is_sqlite = dialect_name == "sqlite"
existing_tables = await session.run_sync(
lambda sync_session: set(sa_inspect(sync_session.connection()).get_table_names())
existing_tables = await session.run_sync(
lambda sync_session: set(sa_inspect(sync_session.connection()).get_table_names())
)
# search_index: SQLite has no FK on the FTS5 virtual table; Postgres
# cascades from the project FK, so the explicit DELETE is redundant.
if is_sqlite and "search_index" in existing_tables:
await session.execute(
text("DELETE FROM search_index WHERE project_id = :project_id"),
{"project_id": entity_id},
)
# search_index: SQLite has no FK on the FTS5 virtual table; Postgres
# cascades from the project FK, so the explicit DELETE is redundant.
if is_sqlite and "search_index" in existing_tables:
await session.execute(
text("DELETE FROM search_index WHERE project_id = :project_id"),
{"project_id": entity_id},
)
# search_vector_chunks: no FK to project on either backend, so both
# backends need this. SQLite must purge vec0 embeddings first
# (rowid pseudocolumn — Postgres uses chunk_id and would 500 here);
# Postgres' chunk_id FK CASCADE handles its embeddings cleanup when
# we delete the chunk rows below.
if "search_vector_chunks" in existing_tables:
if is_sqlite and "search_vector_embeddings" in existing_tables:
# Extension loading is per-connection. We must load vec0 on
# *this* session before the DELETE; otherwise a different
# pooled connection might have written embeddings that we'd
# silently leave behind.
if await _load_sqlite_vec_on_session(session):
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": entity_id},
)
await session.execute(
text("DELETE FROM search_vector_chunks WHERE project_id = :project_id"),
{"project_id": entity_id},
)
# search_vector_chunks: no FK to project on either backend, so both
# backends need this. SQLite must purge vec0 embeddings first
# (rowid pseudocolumn — Postgres uses chunk_id and would 500 here);
# Postgres' chunk_id FK CASCADE handles its embeddings cleanup when
# we delete the chunk rows below.
if "search_vector_chunks" in existing_tables:
if is_sqlite and "search_vector_embeddings" in existing_tables:
# Extension loading is per-connection. We must load vec0 on
# *this* session before the DELETE; otherwise a different
# pooled connection might have written embeddings that we'd
# silently leave behind.
if await _load_sqlite_vec_on_session(session):
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": entity_id},
)
await session.execute(
text("DELETE FROM search_vector_chunks WHERE project_id = :project_id"),
{"project_id": entity_id},
)
await session.delete(project)
logger.debug(f"Deleted Project and search rows for project_id: {entity_id}")
return True
await session.delete(project)
await session.flush()
logger.debug(f"Deleted Project and search rows for project_id: {entity_id}")
return True
async def scalar_vec_query(
self, query: Executable, params: Optional[dict] = None
self, session: AsyncSession, query: Executable, params: Optional[dict] = None
) -> Optional[int]:
"""Run a scalar COUNT query that reads the sqlite-vec vec0 table.
@@ -271,16 +270,17 @@ class ProjectRepository(Repository[Project]):
Returns None when sqlite-vec cannot be loaded on this Python build, so
callers can fall back to the genuinely-missing-dependency path.
"""
async with db.scoped_session(self.session_maker) as session:
# Trigger: query reads the vec0-backed search_vector_embeddings table.
# Why: vec0 modules are only visible on a connection that loaded sqlite-vec.
# Outcome: load the extension here, or signal absence so the caller degrades.
if not await _load_sqlite_vec_on_session(session):
return None
result = await session.execute(query, params)
return result.scalar()
# Trigger: query reads the vec0-backed search_vector_embeddings table.
# Why: vec0 modules are only visible on a connection that loaded sqlite-vec.
# Outcome: load the extension here, or signal absence so the caller degrades.
if not await _load_sqlite_vec_on_session(session):
return None
result = await session.execute(query, params)
return result.scalar()
async def update_path(self, project_id: int, new_path: str) -> Optional[Project]:
async def update_path(
self, session: AsyncSession, project_id: int, new_path: str
) -> Optional[Project]:
"""Update project path.
Args:
@@ -290,10 +290,9 @@ class ProjectRepository(Repository[Project]):
Returns:
The updated project if found, None otherwise
"""
async with db.scoped_session(self.session_maker) as session:
project = await self.select_by_id(session, project_id)
if project:
project.path = new_path
await session.flush()
return project
return None
project = await self.select_by_id(session, project_id)
if project:
project.path = new_path
await session.flush()
return project
return None
@@ -6,11 +6,10 @@ from sqlalchemy import and_, delete, select
from sqlalchemy.engine import CursorResult
from sqlalchemy.dialects.postgresql import insert as pg_insert
from sqlalchemy.dialects.sqlite import insert as sqlite_insert
from sqlalchemy.ext.asyncio import async_sessionmaker
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import selectinload, aliased
from sqlalchemy.orm.interfaces import LoaderOption
from basic_memory import db
from basic_memory.models import Relation, Entity
from basic_memory.repository.repository import Repository
@@ -18,17 +17,20 @@ from basic_memory.repository.repository import Repository
class RelationRepository(Repository[Relation]):
"""Repository for Relation model with memory-specific operations."""
def __init__(self, session_maker: async_sessionmaker, project_id: int):
"""Initialize with session maker and project_id filter.
def __init__(self, project_id: int):
"""Initialize with project_id filter.
Args:
session_maker: SQLAlchemy session maker
project_id: Project ID to filter all operations by
"""
super().__init__(session_maker, Relation, project_id=project_id)
super().__init__(Relation, project_id=project_id)
async def find_relation(
self, from_permalink: str, to_permalink: str, relation_type: str
self,
session: AsyncSession,
from_permalink: str,
to_permalink: str,
relation_type: str,
) -> Optional[Relation]:
"""Find a relation by its from and to path IDs."""
from_entity = aliased(Entity)
@@ -46,36 +48,44 @@ class RelationRepository(Repository[Relation]):
)
)
)
return await self.find_one(query)
query = self._add_project_filter(query)
return await self.find_one(session, query)
async def find_by_entities(self, from_id: int, to_id: int) -> Sequence[Relation]:
async def find_by_entities(
self, session: AsyncSession, from_id: int, to_id: int
) -> Sequence[Relation]:
"""Find all relations between two entities."""
query = select(Relation).where((Relation.from_id == from_id) & (Relation.to_id == to_id))
result = await self.execute_query(query)
query = self.select().where((Relation.from_id == from_id) & (Relation.to_id == to_id))
result = await self.execute_query(session, query)
return result.scalars().all()
async def find_by_type(self, relation_type: str) -> Sequence[Relation]:
async def find_by_type(self, session: AsyncSession, relation_type: str) -> Sequence[Relation]:
"""Find all relations of a specific type."""
query = select(Relation).filter(Relation.relation_type == relation_type)
result = await self.execute_query(query)
query = self.select().filter(Relation.relation_type == relation_type)
result = await self.execute_query(session, query)
return result.scalars().all()
async def delete_outgoing_relations_from_entity(self, entity_id: int) -> None:
async def delete_outgoing_relations_from_entity(
self, session: AsyncSession, entity_id: int
) -> None:
"""Delete outgoing relations for an entity.
Only deletes relations where this entity is the source (from_id),
as these are the ones owned by this entity's markdown file.
"""
async with db.scoped_session(self.session_maker) as session:
await session.execute(delete(Relation).where(Relation.from_id == entity_id))
query = delete(Relation).where(Relation.from_id == entity_id)
query = query.where(Relation.project_id == self.project_id)
await session.execute(query)
async def find_unresolved_relations(self) -> Sequence[Relation]:
async def find_unresolved_relations(self, session: AsyncSession) -> Sequence[Relation]:
"""Find all unresolved relations, where to_id is null."""
query = select(Relation).filter(Relation.to_id.is_(None))
result = await self.execute_query(query)
query = self.select().filter(Relation.to_id.is_(None))
result = await self.execute_query(session, query)
return result.scalars().all()
async def find_unresolved_relations_for_entity(self, entity_id: int) -> Sequence[Relation]:
async def find_unresolved_relations_for_entity(
self, session: AsyncSession, entity_id: int
) -> Sequence[Relation]:
"""Find unresolved relations for a specific entity.
Args:
@@ -84,11 +94,13 @@ class RelationRepository(Repository[Relation]):
Returns:
List of unresolved relations where this entity is the source.
"""
query = select(Relation).filter(Relation.from_id == entity_id, Relation.to_id.is_(None))
result = await self.execute_query(query)
query = self.select().filter(Relation.from_id == entity_id, Relation.to_id.is_(None))
result = await self.execute_query(session, query)
return result.scalars().all()
async def add_all_ignore_duplicates(self, relations: List[Relation]) -> int:
async def add_all_ignore_duplicates(
self, session: AsyncSession, relations: List[Relation]
) -> int:
"""Bulk insert relations, ignoring duplicates.
Uses ON CONFLICT DO NOTHING to skip relations that would violate the
@@ -120,27 +132,23 @@ class RelationRepository(Repository[Relation]):
for r in relations
]
async with db.scoped_session(self.session_maker) as session:
# Check dialect to use appropriate insert
dialect_name = session.bind.dialect.name if session.bind else "sqlite"
# Check dialect to use appropriate insert
dialect_name = session.bind.dialect.name if session.bind else "sqlite"
if dialect_name == "postgresql": # pragma: no cover
# PostgreSQL: use RETURNING to count inserted rows
# (rowcount is 0 for ON CONFLICT DO NOTHING)
stmt = ( # pragma: no cover
pg_insert(Relation)
.values(values)
.on_conflict_do_nothing()
.returning(Relation.id)
)
result = await session.execute(stmt) # pragma: no cover
return len(result.fetchall()) # pragma: no cover
else:
# SQLite: rowcount works correctly
stmt = sqlite_insert(Relation).values(values)
stmt = stmt.on_conflict_do_nothing()
result = cast(CursorResult[Any], await session.execute(stmt))
return result.rowcount if result.rowcount > 0 else 0
if dialect_name == "postgresql": # pragma: no cover
# PostgreSQL: use RETURNING to count inserted rows
# (rowcount is 0 for ON CONFLICT DO NOTHING)
stmt = ( # pragma: no cover
pg_insert(Relation).values(values).on_conflict_do_nothing().returning(Relation.id)
)
result = await session.execute(stmt) # pragma: no cover
return len(result.fetchall()) # pragma: no cover
else:
# SQLite: rowcount works correctly
stmt = sqlite_insert(Relation).values(values)
stmt = stmt.on_conflict_do_nothing()
result = cast(CursorResult[Any], await session.execute(stmt))
return result.rowcount if result.rowcount > 0 else 0
def get_load_options(self) -> List[LoaderOption]:
return [selectinload(Relation.from_entity), selectinload(Relation.to_entity)]
+181 -201
View File
@@ -1,7 +1,6 @@
"""Base repository implementation."""
from typing import Type, Optional, Any, Sequence, TypeVar, List, Dict, cast
from typing import Optional, Any, Sequence, TypeVar, List, Dict, cast
from loguru import logger
from sqlalchemy import (
@@ -16,27 +15,23 @@ from sqlalchemy import (
update as sqlalchemy_update,
)
from sqlalchemy.engine import CursorResult
from sqlalchemy.exc import NoResultFound
from sqlalchemy.ext.asyncio import async_sessionmaker, AsyncSession
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm.interfaces import LoaderOption
from sqlalchemy.sql.elements import ColumnElement
from basic_memory import db
from basic_memory.models import Base
T = TypeVar("T", bound=Base)
class Repository[T: Base]:
"""Base repository implementation with generic CRUD operations."""
"""Base repository implementation with explicit caller-owned sessions."""
def __init__(
self,
session_maker: async_sessionmaker[AsyncSession],
Model: Type[T],
Model: type[T],
project_id: Optional[int] = None,
):
self.session_maker = session_maker
self.project_id = project_id
if Model:
self.Model = Model
@@ -98,62 +93,61 @@ class Repository[T: Base]:
result = await session.execute(query)
return result.scalars().all()
async def add(self, model: T) -> T:
async def add(self, session: AsyncSession, model: T) -> T:
"""
Add a model to the repository. This will also add related objects
:param session: the caller-owned session to use
:param model: the model to add
:return: the added model instance
"""
async with db.scoped_session(self.session_maker) as session:
# Set project_id if applicable and not already set
self._set_project_id_if_needed(model)
# Set project_id if applicable and not already set
self._set_project_id_if_needed(model)
session.add(model)
await session.flush()
session.add(model)
await session.flush()
# Query within same session
found = await self.select_by_id(session, model.id) # pyright: ignore [reportAttributeAccessIssue]
if found is None: # pragma: no cover
logger.error(
"Failed to retrieve model after add",
model_type=self.Model.__name__,
model_id=model.id, # pyright: ignore
)
raise ValueError(
f"Can't find {self.Model.__name__} with ID {model.id} after session.add" # pyright: ignore
)
return found
# Query within same session
found = await self.select_by_id(session, model.id) # pyright: ignore [reportAttributeAccessIssue]
if found is None: # pragma: no cover
logger.error(
"Failed to retrieve model after add",
model_type=self.Model.__name__,
model_id=model.id, # pyright: ignore
)
raise ValueError(
f"Can't find {self.Model.__name__} with ID {model.id} after session.add" # pyright: ignore
)
return found
async def add_all(self, models: List[T]) -> Sequence[T]:
async def add_all(self, session: AsyncSession, models: List[T]) -> Sequence[T]:
"""
Add a list of models to the repository. This will also add related objects
:param session: the caller-owned session to use
:param models: the models to add
:return: the added models instances
"""
async with db.scoped_session(self.session_maker) as session:
# set the project id if not present in models
for model in models:
self._set_project_id_if_needed(model)
# set the project id if not present in models
for model in models:
self._set_project_id_if_needed(model)
session.add_all(models)
await session.flush()
session.add_all(models)
await session.flush()
# Query within same session
return await self.select_by_ids(session, [m.id for m in models]) # pyright: ignore [reportAttributeAccessIssue]
# Query within same session
return await self.select_by_ids(session, [m.id for m in models]) # pyright: ignore [reportAttributeAccessIssue]
async def add_all_no_return(self, models: List[T]) -> int:
async def add_all_no_return(self, session: AsyncSession, models: List[T]) -> int:
"""Insert models without reloading them afterward."""
if not models:
return 0
async with db.scoped_session(self.session_maker) as session:
for model in models:
self._set_project_id_if_needed(model)
for model in models:
self._set_project_id_if_needed(model)
session.add_all(models)
await session.flush()
logger.debug(f"Added {len(models)} {self.Model.__name__} records")
return len(models)
session.add_all(models)
await session.flush()
logger.debug(f"Added {len(models)} {self.Model.__name__} records")
return len(models)
def select(self, *entities: Any) -> Select:
"""Create a new SELECT statement.
@@ -170,55 +164,55 @@ class Repository[T: Base]:
return self._add_project_filter(query)
async def find_all(
self, skip: int = 0, limit: Optional[int] = None, use_load_options: bool = True
self,
session: AsyncSession,
skip: int = 0,
limit: Optional[int] = None,
use_load_options: bool = True,
) -> Sequence[T]:
"""Fetch records from the database with pagination.
Args:
session: The caller-owned session to use
skip: Number of records to skip
limit: Maximum number of records to return
use_load_options: Whether to apply eager loading options (default: True)
"""
logger.debug(f"Finding all {self.Model.__name__} (skip={skip}, limit={limit})")
async with db.scoped_session(self.session_maker) as session:
query = select(self.Model).offset(skip)
query = select(self.Model).offset(skip)
# Only apply load options if requested
if use_load_options:
query = query.options(*self.get_load_options())
# Only apply load options if requested
if use_load_options:
query = query.options(*self.get_load_options())
# Add project filter if applicable
query = self._add_project_filter(query)
# Add project filter if applicable
query = self._add_project_filter(query)
if limit:
query = query.limit(limit)
if limit:
query = query.limit(limit)
result = await session.execute(query)
result = await session.execute(query)
items = result.scalars().all()
logger.debug(f"Found {len(items)} {self.Model.__name__} records")
return items
items = result.scalars().all()
logger.debug(f"Found {len(items)} {self.Model.__name__} records")
return items
async def find_by_id(self, entity_id: int) -> Optional[T]:
async def find_by_id(self, session: AsyncSession, entity_id: Any) -> Optional[T]:
"""Fetch an entity by its unique identifier."""
logger.debug(f"Finding {self.Model.__name__} by ID: {entity_id}")
return await self.select_by_id(session, entity_id)
async with db.scoped_session(self.session_maker) as session:
return await self.select_by_id(session, entity_id)
async def find_by_ids(self, ids: List[int]) -> Sequence[T]:
async def find_by_ids(self, session: AsyncSession, ids: List[Any]) -> Sequence[T]:
"""Fetch multiple entities by their identifiers in a single query."""
logger.debug(f"Finding {self.Model.__name__} by IDs: {ids}")
return await self.select_by_ids(session, ids)
async with db.scoped_session(self.session_maker) as session:
return await self.select_by_ids(session, ids)
async def find_one(self, query: Select[tuple[T]]) -> Optional[T]:
async def find_one(self, session: AsyncSession, query: Select[tuple[T]]) -> Optional[T]:
"""Execute a query and retrieve a single record."""
# add in load options
query = query.options(*self.get_load_options())
result = await self.execute_query(query)
result = await self.execute_query(session, query)
entity = result.scalars().one_or_none()
if entity:
@@ -227,12 +221,42 @@ class Repository[T: Base]:
logger.trace(f"No {self.Model.__name__} found")
return entity
async def create(self, data: dict) -> T:
async def create(self, session: AsyncSession, data: dict[str, Any]) -> T:
"""Create a new record from a model instance."""
logger.debug(f"Creating {self.Model.__name__} from entity_data: {data}")
async with db.scoped_session(self.session_maker) as session:
# Only include valid columns that are provided in entity_data
model_data = self.get_model_data(data)
# Only include valid columns that are provided in entity_data
model_data = self.get_model_data(data)
# Add project_id if applicable and not already provided
if self.has_project_id and self.project_id is not None and "project_id" not in model_data:
model_data["project_id"] = self.project_id
model = self.Model(**model_data)
session.add(model)
await session.flush()
return_instance = await self.select_by_id(session, model.id) # pyright: ignore [reportAttributeAccessIssue]
if return_instance is None: # pragma: no cover
logger.error(
"Failed to retrieve model after create",
model_type=self.Model.__name__,
model_id=model.id, # pyright: ignore
)
raise ValueError(
f"Can't find {self.Model.__name__} with ID {model.id} after session.add" # pyright: ignore
)
return return_instance
async def create_all(
self, session: AsyncSession, data_list: List[dict[str, Any]]
) -> Sequence[T]:
"""Create multiple records in a single transaction."""
logger.debug(f"Bulk creating {len(data_list)} {self.Model.__name__} instances")
# Only include valid columns that are provided in entity_data
model_list = []
for d in data_list:
model_data = self.get_model_data(d)
# Add project_id if applicable and not already provided
if (
@@ -240,169 +264,126 @@ class Repository[T: Base]:
and self.project_id is not None
and "project_id" not in model_data
):
model_data["project_id"] = self.project_id
model_data["project_id"] = self.project_id # pragma: no cover
model = self.Model(**model_data)
session.add(model)
await session.flush()
model_list.append(self.Model(**model_data))
return_instance = await self.select_by_id(session, model.id) # pyright: ignore [reportAttributeAccessIssue]
if return_instance is None: # pragma: no cover
logger.error(
"Failed to retrieve model after create",
model_type=self.Model.__name__,
model_id=model.id, # pyright: ignore
)
raise ValueError(
f"Can't find {self.Model.__name__} with ID {model.id} after session.add" # pyright: ignore
)
return return_instance
session.add_all(model_list)
await session.flush()
async def create_all(self, data_list: List[dict]) -> Sequence[T]:
"""Create multiple records in a single transaction."""
logger.debug(f"Bulk creating {len(data_list)} {self.Model.__name__} instances")
return await self.select_by_ids(session, [model.id for model in model_list]) # pyright: ignore [reportAttributeAccessIssue]
async with db.scoped_session(self.session_maker) as session:
# Only include valid columns that are provided in entity_data
model_list = []
for d in data_list:
model_data = self.get_model_data(d)
# Add project_id if applicable and not already provided
if (
self.has_project_id
and self.project_id is not None
and "project_id" not in model_data
):
model_data["project_id"] = self.project_id # pragma: no cover
model_list.append(self.Model(**model_data))
session.add_all(model_list)
await session.flush()
return await self.select_by_ids(session, [model.id for model in model_list]) # pyright: ignore [reportAttributeAccessIssue]
async def update(self, entity_id: int, entity_data: dict[str, Any] | T) -> Optional[T]:
async def update(
self,
session: AsyncSession,
entity_id: Any,
entity_data: dict[str, Any] | T,
) -> Optional[T]:
"""Update an entity with the given data."""
logger.debug(f"Updating {self.Model.__name__} {entity_id} with data: {entity_data}")
async with db.scoped_session(self.session_maker) as session:
try:
result = await session.execute(
select(self.Model).filter(self.primary_key == entity_id)
)
entity = result.scalars().one()
entity = await self.select_by_id(session, entity_id)
if entity is None:
logger.debug(f"No {self.Model.__name__} found to update: {entity_id}")
return None
if isinstance(entity_data, dict):
update_data = cast(dict[str, Any], entity_data)
for key in self.valid_columns:
if key in update_data:
setattr(entity, key, update_data[key])
if isinstance(entity_data, dict):
update_data = cast(dict[str, Any], entity_data)
for key in self.valid_columns:
if key in update_data:
setattr(entity, key, update_data[key])
elif isinstance(entity_data, self.Model):
for column in self.valid_columns:
setattr(entity, column, getattr(entity_data, column))
elif isinstance(entity_data, self.Model):
for column in self.valid_columns:
setattr(entity, column, getattr(entity_data, column))
await session.flush() # Make sure changes are flushed
await session.refresh(entity) # Refresh
await session.flush()
await session.refresh(entity)
logger.debug(f"Updated {self.Model.__name__}: {entity_id}")
return await self.select_by_id(session, entity.id) # pyright: ignore [reportAttributeAccessIssue]
logger.debug(f"Updated {self.Model.__name__}: {entity_id}")
return await self.select_by_id(session, entity.id) # pyright: ignore [reportAttributeAccessIssue]
except NoResultFound:
logger.debug(f"No {self.Model.__name__} found to update: {entity_id}")
return None
async def update_fields(self, entity_id: Any, entity_data: dict[str, Any]) -> bool:
async def update_fields(
self, session: AsyncSession, entity_id: Any, entity_data: dict[str, Any]
) -> bool:
"""Update columns without reloading the model graph afterward."""
update_data = {k: v for k, v in entity_data.items() if k in self.valid_columns}
if not update_data:
return True
async with db.scoped_session(self.session_maker) as session:
conditions = [self.primary_key == entity_id]
if self.has_project_id and self.project_id is not None:
conditions.append(getattr(self.Model, "project_id") == self.project_id)
conditions = [self.primary_key == entity_id]
if self.has_project_id and self.project_id is not None:
conditions.append(getattr(self.Model, "project_id") == self.project_id)
result = cast(
CursorResult[Any],
await session.execute(
sqlalchemy_update(self.Model).where(and_(*conditions)).values(**update_data)
),
)
return result.rowcount > 0
result = cast(
CursorResult[Any],
await session.execute(
sqlalchemy_update(self.Model).where(and_(*conditions)).values(**update_data)
),
)
return result.rowcount > 0
async def delete(self, entity_id: int) -> bool:
async def delete(self, session: AsyncSession, entity_id: int) -> bool:
"""Delete an entity from the database."""
logger.debug(f"Deleting {self.Model.__name__}: {entity_id}")
async with db.scoped_session(self.session_maker) as session:
try:
result = await session.execute(
select(self.Model).filter(self.primary_key == entity_id)
)
entity = result.scalars().one()
await session.delete(entity)
entity = await self.select_by_id(session, entity_id)
if entity is None:
logger.debug(f"No {self.Model.__name__} found to delete: {entity_id}")
return False
logger.debug(f"Deleted {self.Model.__name__}: {entity_id}")
return True
except NoResultFound:
logger.debug(f"No {self.Model.__name__} found to delete: {entity_id}")
return False
await session.delete(entity)
await session.flush()
async def delete_by_ids(self, ids: List[int]) -> int:
logger.debug(f"Deleted {self.Model.__name__}: {entity_id}")
return True
async def delete_by_ids(self, session: AsyncSession, ids: List[int]) -> int:
"""Delete records matching given IDs."""
logger.debug(f"Deleting {self.Model.__name__} by ids: {ids}")
async with db.scoped_session(self.session_maker) as session:
conditions = [self.primary_key.in_(ids)]
conditions = [self.primary_key.in_(ids)]
# Add project_id filter if applicable
if self.has_project_id and self.project_id is not None: # pragma: no cover
conditions.append(getattr(self.Model, "project_id") == self.project_id)
# Add project_id filter if applicable
if self.has_project_id and self.project_id is not None: # pragma: no cover
conditions.append(getattr(self.Model, "project_id") == self.project_id)
query = delete(self.Model).where(and_(*conditions))
result = cast(CursorResult[Any], await session.execute(query))
logger.debug(f"Deleted {result.rowcount} records")
return result.rowcount
query = delete(self.Model).where(and_(*conditions))
result = cast(CursorResult[Any], await session.execute(query))
logger.debug(f"Deleted {result.rowcount} records")
return result.rowcount
async def delete_by_fields(self, **filters: Any) -> bool:
async def delete_by_fields(self, session: AsyncSession, **filters: Any) -> bool:
"""Delete records matching given field values."""
logger.debug(f"Deleting {self.Model.__name__} by fields: {filters}")
async with db.scoped_session(self.session_maker) as session:
conditions = [getattr(self.Model, field) == value for field, value in filters.items()]
conditions = [getattr(self.Model, field) == value for field, value in filters.items()]
# Add project_id filter if applicable
if self.has_project_id and self.project_id is not None:
conditions.append(getattr(self.Model, "project_id") == self.project_id)
# Add project_id filter if applicable
if self.has_project_id and self.project_id is not None:
conditions.append(getattr(self.Model, "project_id") == self.project_id)
query = delete(self.Model).where(and_(*conditions))
result = cast(CursorResult[Any], await session.execute(query))
deleted = result.rowcount > 0
logger.debug(f"Deleted {result.rowcount} records")
return deleted
query = delete(self.Model).where(and_(*conditions))
result = cast(CursorResult[Any], await session.execute(query))
deleted = result.rowcount > 0
logger.debug(f"Deleted {result.rowcount} records")
return deleted
async def count(self, query: Executable | None = None) -> int:
async def count(self, session: AsyncSession, query: Executable | None = None) -> int:
"""Count entities in the database table."""
async with db.scoped_session(self.session_maker) as session:
if query is None:
query = select(func.count()).select_from(self.Model)
# Add project filter if applicable
if (
isinstance(query, Select)
and self.has_project_id
and self.project_id is not None
):
query = query.where(
getattr(self.Model, "project_id") == self.project_id
) # pragma: no cover
if query is None:
query = select(func.count()).select_from(self.Model)
# Add project filter if applicable
if isinstance(query, Select) and self.has_project_id and self.project_id is not None:
query = query.where(
getattr(self.Model, "project_id") == self.project_id
) # pragma: no cover
result = await session.execute(query)
scalar = result.scalar()
count = scalar if scalar is not None else 0
logger.debug(f"Counted {count} {self.Model.__name__} records")
return count
result = await session.execute(query)
scalar = result.scalar()
count = scalar if scalar is not None else 0
logger.debug(f"Counted {count} {self.Model.__name__} records")
return count
async def execute_query(
self,
session: AsyncSession,
query: Executable,
params: Optional[Dict[str, Any]] = None,
use_query_options: bool = True,
@@ -411,9 +392,8 @@ class Repository[T: Base]:
query = query.options(*self.get_load_options()) if use_query_options else query
logger.trace(f"Executing query: {query}, params: {params}")
async with db.scoped_session(self.session_maker) as session:
result = await session.execute(query, params)
return result
result = await session.execute(query, params)
return result
def get_load_options(self) -> List[LoaderOption]:
"""Get list of loader options for eager loading relationships.
@@ -48,6 +48,8 @@ class SearchRepository(Protocol):
min_similarity: Optional[float] = None,
limit: int = 10,
offset: int = 0,
allow_relaxed: bool = False,
session: AsyncSession | None = None,
) -> List[SearchIndexRow]:
"""Search across indexed content."""
...
@@ -40,6 +40,50 @@ BULLET_PATTERN = re.compile(r"^[\-\*]\s+")
OVERSIZED_ENTITY_VECTOR_SHARD_SIZE = 256
_SQLITE_MAX_PREPARE_WINDOW = 8
# Interrogative/function words contribute lexical noise when a strict
# full-text query is relaxed: "when OR did OR a" matches loud wrong documents
# that displace genuine results from the ranking window.
RELAXATION_STOPWORDS = frozenset(
"a an and are as at be but by did do does for from had has have how i in is it of on "
"or that the their they this to was we were what when where which who whom whose why "
"will with you your".split()
)
def relaxed_query_words(search_text: Optional[str]) -> Optional[list[str]]:
"""Content-bearing words for OR-relaxing a strict full-text query.
Returns None when relaxation must not apply. These eligibility rules match
SearchService._is_relaxed_fts_fallback_eligible so the hybrid FTS branch
relaxes exactly the same query shapes as the service-level FTS path:
- empty / quoted / explicit-boolean queries (user intent is not
second-guessed);
- fewer than three alphanumeric tokens (short queries like "New Feature"
over-broaden under OR and in hybrid the relaxed FTS-only rows normalize
to 1.0 and can outrank the vector result the user wanted);
- any pure-digit token ("root note 1", "SPEC 16") identifier-like queries
over-broaden and create false positives under OR.
"""
if not search_text:
return None
stripped = search_text.strip()
if '"' in stripped or any(op in f" {stripped} " for op in (" AND ", " OR ", " NOT ")):
return None
# Eligibility checks run on raw alphanumeric tokens (parity with the
# service), before stopword filtering.
tokens = re.findall(r"[A-Za-z0-9]+", stripped.lower())
if len(tokens) < 3 or any(token.isdigit() for token in tokens):
return None
words = [word.strip("?!.,;:") for word in stripped.split()]
words = [
word
for word in words
if word and word.isalnum() and word.lower() not in RELAXATION_STOPWORDS
]
return words or None
# Entity, observation, and relation rows in search_index carry ids from independent
# auto-increment sequences, so a bare id is ambiguous across row types. Every map in
# the vector/hybrid retrieval path must key rows by (type, id) to avoid collisions.
@@ -229,6 +273,7 @@ class SearchRepositoryBase(ABC):
min_similarity: Optional[float] = None,
limit: int = 10,
offset: int = 0,
allow_relaxed: bool = False,
) -> List[SearchIndexRow]:
"""Search across all indexed content.
@@ -2174,6 +2219,9 @@ class SearchRepositoryBase(ABC):
query_start = time.perf_counter()
candidate_limit = max(self._semantic_vector_k, (limit + offset) * 10)
fts_start = time.perf_counter()
# allow_relaxed: question-form queries rarely AND-match, and a dead FTS
# branch silently degrades hybrid to vector-only ranking. Fusion plus
# bm25 keep relaxed lexical candidates from dominating precision.
fts_results = await self.search(
search_text=search_text,
permalink=permalink,
@@ -2187,6 +2235,7 @@ class SearchRepositoryBase(ABC):
retrieval_mode=SearchRetrievalMode.FTS,
limit=candidate_limit,
offset=0,
allow_relaxed=True,
)
fts_ms = (time.perf_counter() - fts_start) * 1000
vector_start = time.perf_counter()
@@ -23,7 +23,10 @@ from basic_memory.models.search import (
from basic_memory.repository.embedding_provider import EmbeddingProvider
from basic_memory.repository.embedding_provider_factory import create_embedding_provider
from basic_memory.repository.search_index_row import SearchIndexRow
from basic_memory.repository.search_repository_base import SearchRepositoryBase
from basic_memory.repository.search_repository_base import (
SearchRepositoryBase,
relaxed_query_words,
)
from basic_memory.repository.metadata_filters import parse_metadata_filters, build_sqlite_json_path
from basic_memory.repository.semantic_errors import SemanticDependenciesMissingError
from basic_memory.schemas.search import SearchItemType, SearchRetrievalMode
@@ -255,6 +258,19 @@ class SQLiteSearchRepository(SearchRepositoryBase):
if "*" in term and all(c.isalnum() or c in "*_-" for c in term):
return term
# Natural-language queries arrive with sentence punctuation that FTS5
# treats as syntax ("When did Melanie paint a sunrise?"). The tokenizer
# ignores this punctuation in the INDEX, so stripping it from word
# edges loses nothing — but leaving it forces the whole question into
# an exact-phrase match that returns zero rows, silently disabling the
# FTS half of hybrid search. Interior characters (hyphens, slashes —
# permalinks and paths) are untouched.
if " " in term:
words = [word.strip("?!.,;:") for word in term.split()]
term = " ".join(word for word in words if word)
if not term:
return ""
# Characters that can cause FTS5 syntax errors when used as operators
# We're more conservative here - only quote when we detect problematic patterns
problematic_chars = [
@@ -351,6 +367,14 @@ class SQLiteSearchRepository(SearchRepositoryBase):
# For non-Boolean queries, use the single term preparation logic
return self._prepare_single_term(term, is_prefix)
@staticmethod
def _relaxed_fts_text(search_text: Optional[str]) -> Optional[str]:
"""OR-relaxed FTS5 expression for a failed strict query, or None."""
words = relaxed_query_words(search_text)
if not words:
return None
return " OR ".join(f"{word}*" for word in words)
# ------------------------------------------------------------------
# sqlite-vec extension loading (SQLite-specific)
# ------------------------------------------------------------------
@@ -953,8 +977,16 @@ class SQLiteSearchRepository(SearchRepositoryBase):
min_similarity: Optional[float] = None,
limit: int = 10,
offset: int = 0,
allow_relaxed: bool = False,
session: AsyncSession | None = None,
) -> List[SearchIndexRow]:
"""Search across all indexed content using SQLite FTS5."""
"""Search across all indexed content using SQLite FTS5.
``allow_relaxed=True`` retries a zero-result strict multi-word query
with OR-joined content terms. Only the hybrid path opts in: its FTS
branch otherwise contributes nothing for question-form queries.
Service-level FTS searches keep their own conservative fallback.
"""
# --- Dispatch vector / hybrid modes (shared logic) ---
dispatched = await self._dispatch_retrieval_mode(
search_text=search_text,
@@ -1017,10 +1049,31 @@ class SQLiteSearchRepository(SearchRepositoryBase):
"""
logger.trace(f"Search {sql} params: {params}")
try:
async with db.scoped_session(self.session_maker) as session:
result = await session.execute(text(sql), params)
async def run_search(active_session: AsyncSession):
result = await active_session.execute(text(sql), params)
rows = result.fetchall()
# Trigger: multi-word natural-language query matched nothing
# under the default all-terms-AND semantics.
# Why: questions ("when did X do Y") rarely have every word in
# one document; without relaxation the FTS half of hybrid
# search contributes zero candidates and ranking degrades to
# vector-only.
# Outcome: one retry with OR-joined prefix terms; bm25 still
# ranks multi-term matches first.
relaxed = self._relaxed_fts_text(search_text) if allow_relaxed and not rows else None
if relaxed and params.get("text"):
params["text"] = relaxed
result = await active_session.execute(text(sql), params)
rows = result.fetchall()
return rows
try:
if session is not None:
rows = await run_search(session)
else:
async with db.scoped_session(self.session_maker) as owned_session:
rows = await run_search(owned_session)
except Exception as e:
# Handle FTS5 syntax errors and provide user-friendly feedback
if self._is_fts5_syntax_error(e): # pragma: no cover
+25 -7
View File
@@ -2,15 +2,18 @@
from __future__ import annotations
from contextlib import asynccontextmanager
from dataclasses import dataclass, field
from datetime import datetime, timezone
from typing import Any, List, Optional, Tuple, TYPE_CHECKING
from typing import Any, AsyncIterator, List, Optional, Tuple, TYPE_CHECKING
from loguru import logger
from sqlalchemy import text
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
import logfire
from basic_memory import db
from basic_memory.repository.entity_repository import EntityRepository
from basic_memory.repository.observation_repository import ObservationRepository
from basic_memory.repository.postgres_search_repository import PostgresSearchRepository
@@ -91,11 +94,21 @@ class ContextService:
entity_repository: EntityRepository,
observation_repository: ObservationRepository,
link_resolver: Optional[LinkResolver] = None,
session_maker: async_sessionmaker[AsyncSession] | None = None,
):
self.search_repository = search_repository
self.entity_repository = entity_repository
self.observation_repository = observation_repository
self.link_resolver = link_resolver
self.session_maker = session_maker
@asynccontextmanager
async def _session_scope(self) -> AsyncIterator[AsyncSession]:
"""Open a service-owned transaction for core repository reads."""
if self.session_maker is None: # pragma: no cover
raise ValueError("session_maker is required for ContextService")
async with db.scoped_session(self.session_maker) as session:
yield session
async def build_context(
self,
@@ -176,9 +189,13 @@ class ContextService:
)
if not primary and self.link_resolver:
entity = await self.link_resolver.resolve_link(
path, use_search=True, strict=False
)
async with self._session_scope() as session:
entity = await self.link_resolver.resolve_link(
path,
use_search=True,
strict=False,
session=session,
)
if entity:
logger.debug(
f"LinkResolver resolved '{path}' to permalink '{entity.permalink}'"
@@ -234,9 +251,10 @@ class ContextService:
phase="load_observations",
result_count=len(entity_ids),
):
observations_by_entity = await self.observation_repository.find_by_entities(
entity_ids
)
async with self._session_scope() as session:
observations_by_entity = await self.observation_repository.find_by_entities(
session, entity_ids
)
logger.debug(f"Found observations for {len(observations_by_entity)} entities")
metadata = ContextMetadata(
+23 -5
View File
@@ -3,10 +3,14 @@
import fnmatch
import logging
import os
from contextlib import asynccontextmanager
from datetime import datetime
from typing import Dict, List, Optional, Sequence
from typing import AsyncIterator, Dict, List, Optional, Sequence
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
from basic_memory import db
from basic_memory.models import Entity
from basic_memory.repository import EntityRepository
from basic_memory.schemas.directory import DirectoryNode
@@ -28,19 +32,31 @@ def _mtime_to_datetime(entity: Entity) -> datetime:
class DirectoryService:
"""Service for working with directory trees."""
def __init__(self, entity_repository: EntityRepository):
def __init__(
self,
entity_repository: EntityRepository,
session_maker: async_sessionmaker[AsyncSession],
):
"""Initialize the directory service.
Args:
entity_repository: Directory repository for data access.
"""
self.entity_repository = entity_repository
self.session_maker = session_maker
@asynccontextmanager
async def _session_scope(self) -> AsyncIterator[AsyncSession]:
"""Open a service-owned transaction."""
async with db.scoped_session(self.session_maker) as session:
yield session
async def get_directory_tree(self) -> DirectoryNode:
"""Build a hierarchical directory tree from indexed files."""
# Get all files from DB (flat list)
entity_rows = await self.entity_repository.find_all()
async with self._session_scope() as session:
entity_rows = await self.entity_repository.find_all(session)
# Create a root directory node
root_node = DirectoryNode(name="Root", directory_path="/", type="directory")
@@ -114,7 +130,8 @@ class DirectoryService:
DirectoryNode tree containing only folders (type="directory")
"""
# Get unique directories without loading entities
directories = await self.entity_repository.get_distinct_directories()
async with self._session_scope() as session:
directories = await self.entity_repository.get_distinct_directories(session)
# Create a root directory node
root_node = DirectoryNode(name="Root", directory_path="/", type="directory")
@@ -179,7 +196,8 @@ class DirectoryService:
# Optimize: Query only entities in the target directory
# instead of loading the entire tree
dir_prefix = dir_name.lstrip("/")
entity_rows = await self.entity_repository.find_by_directory_prefix(dir_prefix)
async with self._session_scope() as session:
entity_rows = await self.entity_repository.find_by_directory_prefix(session, dir_prefix)
# Build a partial tree from only the relevant entities
root_tree = self._build_directory_tree_from_entities(entity_rows, dir_name)
+374 -309
View File
@@ -1,16 +1,18 @@
"""Service for managing entities in the database."""
import asyncio
from collections.abc import Callable
from contextlib import asynccontextmanager
from copy import deepcopy
from dataclasses import dataclass
from pathlib import Path
from typing import Any, List, Optional, Sequence, Tuple, Union
from typing import Any, AsyncIterator, List, Optional, Sequence, Tuple, Union
import frontmatter
import yaml
from loguru import logger
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
from basic_memory import db
from basic_memory.config import ProjectConfig, BasicMemoryConfig
from basic_memory.file_utils import (
has_frontmatter,
@@ -101,6 +103,7 @@ class EntityService(BaseService[EntityModel]):
relation_repository: RelationRepository,
file_service: FileService,
link_resolver: LinkResolver,
session_maker: async_sessionmaker[AsyncSession],
search_service: Optional[SearchService] = None,
app_config: Optional[BasicMemoryConfig] = None,
):
@@ -110,6 +113,7 @@ class EntityService(BaseService[EntityModel]):
self.entity_parser = entity_parser
self.file_service = file_service
self.link_resolver = link_resolver
self.session_maker = session_maker
self.search_service = search_service
self.app_config = app_config
self._project_permalink: Optional[str] = None
@@ -117,8 +121,23 @@ class EntityService(BaseService[EntityModel]):
# Default returns None for local/CLI usage. Cloud overrides this to read from UserContext.
self.get_user_id: Callable[[], Optional[str]] = lambda: None
@asynccontextmanager
async def _session_scope(
self, session: AsyncSession | None = None
) -> AsyncIterator[AsyncSession]:
"""Use the caller's session or open a service-owned transaction."""
if session is not None:
yield session
return
async with db.scoped_session(self.session_maker) as owned_session:
yield owned_session
async def detect_file_path_conflicts(
self, file_path: str, skip_check: bool = False
self,
file_path: str,
skip_check: bool = False,
session: AsyncSession | None = None,
) -> List[str]:
"""Detect potential file path conflicts for a given file path.
@@ -142,7 +161,8 @@ class EntityService(BaseService[EntityModel]):
# Load only file paths. Conflict detection is on the hot write path and
# does not need observations or relations.
existing_paths = await self.repository.get_all_file_paths()
async with self._session_scope(session) as active_session:
existing_paths = await self.repository.get_all_file_paths(active_session)
# Use the enhanced conflict detection utility
return detect_potential_file_conflicts(file_path, existing_paths)
@@ -152,6 +172,7 @@ class EntityService(BaseService[EntityModel]):
file_path: Permalink | Path,
markdown: Optional[EntityMarkdown] = None,
skip_conflict_check: bool = False,
session: AsyncSession | None = None,
) -> str:
"""Get or generate unique permalink for an entity.
@@ -169,71 +190,74 @@ class EntityService(BaseService[EntityModel]):
file_path_str = Path(file_path).as_posix()
# Check for potential file path conflicts before resolving permalink
conflicts = await self.detect_file_path_conflicts(
file_path_str, skip_check=skip_conflict_check
)
if conflicts:
logger.warning(
f"Detected potential file path conflicts for '{file_path_str}': {conflicts}"
async with self._session_scope(session) as active_session:
conflicts = await self.detect_file_path_conflicts(
file_path_str, skip_check=skip_conflict_check, session=active_session
)
if conflicts:
logger.warning(
f"Detected potential file path conflicts for '{file_path_str}': {conflicts}"
)
# If markdown has explicit permalink, try to validate it
if markdown and markdown.frontmatter.permalink:
desired_permalink = markdown.frontmatter.permalink
# Use lightweight method - we only need to check file_path
existing_file_path = await self.repository.get_file_path_for_permalink(
desired_permalink
# If markdown has explicit permalink, try to validate it
if markdown and markdown.frontmatter.permalink:
desired_permalink = markdown.frontmatter.permalink
# Use lightweight method - we only need to check file_path
existing_file_path = await self.repository.get_file_path_for_permalink(
active_session, desired_permalink
)
# If no conflict or it's our own file, use as is
if not existing_file_path or existing_file_path == file_path_str:
return desired_permalink
# For existing files, try to find current permalink
# Use lightweight method - we only need the permalink
existing_permalink = await self.repository.get_permalink_for_file_path(
active_session, file_path_str
)
if existing_permalink:
return existing_permalink
# If no conflict or it's our own file, use as is
if not existing_file_path or existing_file_path == file_path_str:
return desired_permalink
# New file - generate permalink
if markdown and markdown.frontmatter.permalink:
desired_permalink = markdown.frontmatter.permalink
else:
# Trigger: generating a permalink for a new file
# Why: canonical permalinks may require project prefix for global addressing
# Outcome: include project slug when enabled in config
include_project = True
if self.app_config:
include_project = self.app_config.permalinks_include_project
# For existing files, try to find current permalink
# Use lightweight method - we only need the permalink
existing_permalink = await self.repository.get_permalink_for_file_path(file_path_str)
if existing_permalink:
return existing_permalink
workspace_permalink = workspace_slug_for_canonical_permalinks()
project_permalink = None
# Trigger: project-prefixed permalinks are enabled, or organization workspace
# context requires a complete workspace/project canonical permalink.
# Why: project slug is the stable middle segment for globally addressable links.
# Outcome: fetch and cache the project's permalink before building the canonical URL.
if include_project or workspace_permalink:
project_permalink = await self._get_project_permalink(active_session)
# New file - generate permalink
if markdown and markdown.frontmatter.permalink:
desired_permalink = markdown.frontmatter.permalink
else:
# Trigger: generating a permalink for a new file
# Why: canonical permalinks may require project prefix for global addressing
# Outcome: include project slug when enabled in config
include_project = True
if self.app_config:
include_project = self.app_config.permalinks_include_project
desired_permalink = build_canonical_permalink(
project_permalink,
file_path_str,
include_project=include_project,
workspace_permalink=workspace_permalink,
)
workspace_permalink = workspace_slug_for_canonical_permalinks()
project_permalink = None
# Trigger: project-prefixed permalinks are enabled, or organization workspace
# context requires a complete workspace/project canonical permalink.
# Why: project slug is the stable middle segment for globally addressable links.
# Outcome: fetch and cache the project's permalink before building the canonical URL.
if include_project or workspace_permalink:
project_permalink = await self._get_project_permalink()
desired_permalink = build_canonical_permalink(
project_permalink,
file_path_str,
include_project=include_project,
workspace_permalink=workspace_permalink,
)
# Make unique if needed - enhanced to handle character conflicts
# Use lightweight existence check instead of loading full entity
permalink = desired_permalink
suffix = 1
while await self.repository.permalink_exists(permalink):
permalink = f"{desired_permalink}-{suffix}"
suffix += 1
logger.debug(f"creating unique permalink: {permalink}")
# Make unique if needed - enhanced to handle character conflicts
# Use lightweight existence check instead of loading full entity
permalink = desired_permalink
suffix = 1
while await self.repository.permalink_exists(active_session, permalink):
permalink = f"{desired_permalink}-{suffix}"
suffix += 1
logger.debug(f"creating unique permalink: {permalink}")
return permalink
async def _get_project_permalink(self) -> Optional[str]:
async def _get_project_permalink(self, session: AsyncSession) -> Optional[str]:
"""Get and cache the current project's permalink."""
if self._project_permalink is not None:
return self._project_permalink
@@ -242,8 +266,8 @@ class EntityService(BaseService[EntityModel]):
if project_id is None: # pragma: no cover
return None # pragma: no cover
project_repository = ProjectRepository(self.repository.session_maker)
project = await project_repository.get_by_id(project_id)
project_repository = ProjectRepository()
project = await project_repository.get_by_id(session, project_id)
if project:
self._project_permalink = project.permalink
return self._project_permalink
@@ -343,6 +367,7 @@ class EntityService(BaseService[EntityModel]):
current_permalink: str | None = None,
content_markdown: EntityMarkdown | None = None,
skip_conflict_check: bool = False,
session: AsyncSession | None = None,
) -> str | None:
"""Resolve the canonical permalink for a create/update write."""
if self.app_config and self.app_config.disable_permalinks:
@@ -360,6 +385,7 @@ class EntityService(BaseService[EntityModel]):
file_path,
content_markdown,
skip_conflict_check=skip_conflict_check,
session=session,
)
schema._permalink = resolved_permalink
return resolved_permalink
@@ -437,6 +463,7 @@ class EntityService(BaseService[EntityModel]):
*,
check_storage_exists: bool = True,
skip_conflict_check: bool = False,
session: AsyncSession | None = None,
) -> PreparedEntityWrite:
"""Derive accepted markdown and entity fields for a new note.
@@ -469,6 +496,7 @@ class EntityService(BaseService[EntityModel]):
file_path=file_path,
content_markdown=content_markdown,
skip_conflict_check=skip_conflict_check,
session=session,
)
# Build the final markdown once here. Local mode will write it immediately; cloud mode can
@@ -496,6 +524,7 @@ class EntityService(BaseService[EntityModel]):
existing_content: str,
*,
skip_conflict_check: bool = False,
session: AsyncSession | None = None,
) -> PreparedEntityWrite:
"""Derive accepted markdown and entity fields for a full note replacement.
@@ -533,6 +562,7 @@ class EntityService(BaseService[EntityModel]):
current_permalink=current_permalink,
content_markdown=content_markdown,
skip_conflict_check=skip_conflict_check,
session=session,
)
post = await schema_to_markdown(schema)
@@ -573,6 +603,7 @@ class EntityService(BaseService[EntityModel]):
find_text: Optional[str] = None,
expected_replacements: int = 1,
skip_conflict_check: bool = False,
session: AsyncSession | None = None,
) -> PreparedEntityWrite:
"""Derive accepted markdown and entity fields for an edit request.
@@ -621,6 +652,7 @@ class EntityService(BaseService[EntityModel]):
file_path,
content_markdown,
skip_conflict_check=skip_conflict_check,
session=session,
)
normalized_metadata = normalize_frontmatter_metadata(content_frontmatter or {})
@@ -676,30 +708,34 @@ class EntityService(BaseService[EntityModel]):
async def create_entity_with_content(self, schema: EntitySchema) -> EntityWriteResult:
"""Create a new entity and return both the entity row and written markdown."""
logger.debug(f"Creating entity: {schema.title}")
# --- Prepare Accepted State ---
# Derive the canonical markdown/entity fields before touching the filesystem.
prepared = await self.prepare_create_entity_content(schema)
self._sync_prepared_schema_state(schema, prepared)
# --- Persist File, Then Indexable DB State ---
# Local mode still writes the file immediately; the prepare object keeps semantics separate
# from that persistence step.
checksum = await self.file_service.write_file(prepared.file_path, prepared.markdown_content)
entity = await self.upsert_entity_from_markdown(
prepared.file_path,
prepared.entity_markdown,
is_new=True,
)
updated = await self.repository.update(entity.id, {"checksum": checksum})
if not updated: # pragma: no cover
raise ValueError(f"Failed to update entity checksum after create: {entity.id}")
persisted_content, search_content = await self._read_persisted_write_content(
prepared.file_path
)
return EntityWriteResult(
entity=updated,
content=persisted_content,
search_content=search_content,
)
async with self._session_scope() as session:
# --- Prepare Accepted State ---
# Derive the canonical markdown/entity fields before touching the filesystem.
prepared = await self.prepare_create_entity_content(schema, session=session)
self._sync_prepared_schema_state(schema, prepared)
# --- Persist File, Then Indexable DB State ---
# Local mode still writes the file immediately; the prepare object keeps semantics separate
# from that persistence step.
checksum = await self.file_service.write_file(
prepared.file_path, prepared.markdown_content
)
entity = await self.upsert_entity_from_markdown(
prepared.file_path,
prepared.entity_markdown,
is_new=True,
session=session,
)
updated = await self.repository.update(session, entity.id, {"checksum": checksum})
if not updated: # pragma: no cover
raise ValueError(f"Failed to update entity checksum after create: {entity.id}")
persisted_content, search_content = await self._read_persisted_write_content(
prepared.file_path
)
return EntityWriteResult(
entity=updated,
content=persisted_content,
search_content=search_content,
)
async def update_entity(self, entity: EntityModel, schema: EntitySchema) -> EntityModel:
"""Update an entity's content and metadata."""
@@ -716,58 +752,63 @@ class EntityService(BaseService[EntityModel]):
f"Updating entity with permalink: {entity.permalink} content-type: {schema.content_type}"
)
# --- Read Current File State ---
# Full replacements merge with existing frontmatter, so local mode still needs the current
# file contents as input to the prepare step.
existing_content = await self.file_service.read_file_content(entity.file_path)
prepared = await self.prepare_update_entity_content(
entity,
schema,
existing_content,
)
self._sync_prepared_schema_state(schema, prepared)
previous_file_path = Path(entity.file_path)
# Trigger: a full replacement also renames the note to a different canonical path.
# Why: Path.replace() overwrites existing files, so the destination must be conflict-free
# before we write or we can clobber another note and only fail later at the DB layer.
# Outcome: conflicting rename attempts fail before touching either file on disk.
if (
prepared.file_path.as_posix() != previous_file_path.as_posix()
and await self.file_service.exists(prepared.file_path)
and not self._paths_share_storage_target(previous_file_path, prepared.file_path)
):
raise EntityAlreadyExistsError(
f"file already exists at destination path: {prepared.file_path.as_posix()}"
async with self._session_scope() as session:
# --- Read Current File State ---
# Full replacements merge with existing frontmatter, so local mode still needs the current
# file contents as input to the prepare step.
existing_content = await self.file_service.read_file_content(entity.file_path)
prepared = await self.prepare_update_entity_content(
entity,
schema,
existing_content,
session=session,
)
self._sync_prepared_schema_state(schema, prepared)
previous_file_path = Path(entity.file_path)
# Trigger: a full replacement also renames the note to a different canonical path.
# Why: Path.replace() overwrites existing files, so the destination must be conflict-free
# before we write or we can clobber another note and only fail later at the DB layer.
# Outcome: conflicting rename attempts fail before touching either file on disk.
if (
prepared.file_path.as_posix() != previous_file_path.as_posix()
and await self.file_service.exists(prepared.file_path)
and not self._paths_share_storage_target(previous_file_path, prepared.file_path)
):
raise EntityAlreadyExistsError(
f"file already exists at destination path: {prepared.file_path.as_posix()}"
)
# --- Persist Prepared State ---
checksum = await self.file_service.write_file(
prepared.file_path,
prepared.markdown_content,
)
entity = await self.upsert_entity_from_markdown(
prepared.file_path,
prepared.entity_markdown,
is_new=False,
existing_entity=entity,
session=session,
)
if prepared.file_path.as_posix() != previous_file_path.as_posix():
# Trigger: a full replacement changed the canonical note path.
# Why: the new file has already been written and the entity now points at it.
# Outcome: remove the stale old file so local Basic Memory mirrors cloud's PGQ cleanup.
if not self._paths_share_storage_target(previous_file_path, prepared.file_path):
await self.file_service.delete_file(previous_file_path)
entity = await self.repository.update(session, entity.id, {"checksum": checksum})
if not entity: # pragma: no cover
raise ValueError(
f"Failed to update entity checksum after update: {prepared.file_path}"
)
persisted_content, search_content = await self._read_persisted_write_content(
prepared.file_path
)
# --- Persist Prepared State ---
checksum = await self.file_service.write_file(
prepared.file_path,
prepared.markdown_content,
)
entity = await self.upsert_entity_from_markdown(
prepared.file_path,
prepared.entity_markdown,
is_new=False,
existing_entity=entity,
)
if prepared.file_path.as_posix() != previous_file_path.as_posix():
# Trigger: a full replacement changed the canonical note path.
# Why: the new file has already been written and the entity now points at it.
# Outcome: remove the stale old file so local Basic Memory mirrors cloud's PGQ cleanup.
if not self._paths_share_storage_target(previous_file_path, prepared.file_path):
await self.file_service.delete_file(previous_file_path)
entity = await self.repository.update(entity.id, {"checksum": checksum})
if not entity: # pragma: no cover
raise ValueError(f"Failed to update entity checksum after update: {prepared.file_path}")
persisted_content, search_content = await self._read_persisted_write_content(
prepared.file_path
)
return EntityWriteResult(
entity=entity,
content=persisted_content,
search_content=search_content,
)
return EntityWriteResult(
entity=entity,
content=persisted_content,
search_content=search_content,
)
async def delete_entity(self, permalink_or_id: str | int) -> bool:
"""Delete entity and its file."""
@@ -812,7 +853,8 @@ class EntityService(BaseService[EntityModel]):
# Trigger: repository.delete returns False when entity is already gone (NoResultFound)
# Why: concurrent delete_directory requests can race to delete the same entity
# Outcome: treat as success since the entity is deleted either way
deleted = await self.repository.delete(entity.id)
async with self._session_scope() as session:
deleted = await self.repository.delete(session, entity.id)
if not deleted:
logger.info("Entity already removed from DB", entity_id=permalink_or_id)
return True
@@ -824,7 +866,8 @@ class EntityService(BaseService[EntityModel]):
async def get_by_permalink(self, permalink: str) -> EntityModel:
"""Get entity by type and name combination."""
logger.debug(f"Getting entity by permalink: {permalink}")
db_entity = await self.repository.get_by_permalink(permalink)
async with self._session_scope() as session:
db_entity = await self.repository.get_by_permalink(session, permalink)
if not db_entity:
raise EntityNotFoundError(f"Entity not found: {permalink}")
return db_entity
@@ -832,19 +875,25 @@ class EntityService(BaseService[EntityModel]):
async def get_entities_by_id(self, ids: List[int]) -> Sequence[EntityModel]:
"""Get specific entities and their relationships."""
logger.debug(f"Getting entities: {ids}")
return await self.repository.find_by_ids(ids)
async with self._session_scope() as session:
return await self.repository.find_by_ids(session, ids)
async def get_entities_by_permalinks(self, permalinks: List[str]) -> Sequence[EntityModel]:
"""Get specific nodes and their relationships."""
logger.debug(f"Getting entities permalinks: {permalinks}")
return await self.repository.find_by_permalinks(permalinks)
async with self._session_scope() as session:
return await self.repository.find_by_permalinks(session, permalinks)
async def delete_entity_by_file_path(self, file_path: Union[str, Path]) -> None:
"""Delete entity by file path."""
await self.repository.delete_by_file_path(str(file_path))
async with self._session_scope() as session:
await self.repository.delete_by_file_path(session, str(file_path))
async def create_entity_from_markdown(
self, file_path: Path, markdown: EntityMarkdown
self,
file_path: Path,
markdown: EntityMarkdown,
session: AsyncSession | None = None,
) -> EntityModel:
"""Create entity and observations only.
@@ -867,12 +916,13 @@ class EntityService(BaseService[EntityModel]):
model.created_by = user_id
model.last_updated_by = user_id
# Use UPSERT to handle conflicts cleanly
try:
return await self.repository.upsert_entity(model)
except Exception as e:
logger.error(f"Failed to upsert entity for {file_path}: {e}")
raise EntityCreationError(f"Failed to create entity: {str(e)}") from e
async with self._session_scope(session) as active_session:
# Use UPSERT to handle conflicts cleanly
try:
return await self.repository.upsert_entity(active_session, model)
except Exception as e:
logger.error(f"Failed to upsert entity for {file_path}: {e}")
raise EntityCreationError(f"Failed to create entity: {str(e)}") from e
async def update_entity_and_observations(
self,
@@ -880,6 +930,7 @@ class EntityService(BaseService[EntityModel]):
markdown: EntityMarkdown,
*,
existing_entity: EntityModel | None = None,
session: AsyncSession | None = None,
) -> EntityModel:
"""Update entity fields and observations.
@@ -888,65 +939,53 @@ class EntityService(BaseService[EntityModel]):
"""
logger.debug(f"Updating entity and observations: {file_path}")
if existing_entity is not None:
db_entity = await self.repository.get_by_id(
existing_entity.id,
load_relations=False,
async with self._session_scope(session) as active_session:
if existing_entity is not None:
db_entity = await self.repository.get_by_id(
active_session,
existing_entity.id,
load_relations=False,
)
else:
db_entity = await self.repository.get_by_file_path(
active_session,
file_path.as_posix(),
load_relations=False,
)
if db_entity is None: # pragma: no cover
raise EntityNotFoundError(f"Entity not found for file path: {file_path}")
# Observations are owned by the markdown file, so re-indexing replaces the old set.
# We only need the entity id here; loading the old relationship collection is wasted work.
await self.observation_repository.delete_by_fields(
active_session, entity_id=db_entity.id
)
else:
db_entity = await self.repository.get_by_file_path(
file_path.as_posix(),
load_relations=False,
)
if db_entity is None: # pragma: no cover
raise EntityNotFoundError(f"Entity not found for file path: {file_path}")
# Observations are owned by the markdown file, so re-indexing replaces the old set.
# We only need the entity id here; loading the old relationship collection is wasted work.
await self.observation_repository.delete_by_fields(entity_id=db_entity.id)
observations = [
Observation(
project_id=self.observation_repository.project_id,
entity_id=db_entity.id,
content=obs.content,
category=obs.category,
context=obs.context,
tags=obs.tags,
)
for obs in markdown.observations
]
await self.observation_repository.add_all_no_return(active_session, observations)
observations = [
Observation(
project_id=self.observation_repository.project_id,
entity_id=db_entity.id,
content=obs.content,
category=obs.category,
context=obs.context,
tags=obs.tags,
)
for obs in markdown.observations
]
await self.observation_repository.add_all_no_return(observations)
self._apply_markdown_entity_fields(db_entity, file_path, markdown)
self._apply_markdown_entity_fields(db_entity, file_path, markdown)
# checksum value is None == not finished with sync
db_entity.checksum = None
# checksum value is None == not finished with sync
db_entity.checksum = None
# Set last_updated_by for cloud usage (preserve existing created_by)
user_id = self.get_user_id()
if user_id is not None:
db_entity.last_updated_by = user_id
# Set last_updated_by for cloud usage (preserve existing created_by)
user_id = self.get_user_id()
if user_id is not None:
db_entity.last_updated_by = user_id
entity_updates = {
"title": db_entity.title,
"note_type": db_entity.note_type,
"permalink": db_entity.permalink,
"file_path": db_entity.file_path,
"content_type": db_entity.content_type,
"created_at": db_entity.created_at,
"updated_at": db_entity.updated_at,
"entity_metadata": db_entity.entity_metadata,
"checksum": db_entity.checksum,
"last_updated_by": db_entity.last_updated_by,
}
updated = await self.repository.update_fields(
db_entity.id,
entity_updates,
)
if not updated: # pragma: no cover
raise EntityNotFoundError(f"Entity not found for file path: {file_path}")
return db_entity
await active_session.flush()
return db_entity
def _apply_markdown_entity_fields(
self,
@@ -981,23 +1020,29 @@ class EntityService(BaseService[EntityModel]):
existing_entity: EntityModel | None = None,
resolve_relations: bool = True,
reload_entity: bool = True,
session: AsyncSession | None = None,
) -> EntityModel:
"""Create/update entity and relations from parsed markdown."""
if is_new:
created = await self.create_entity_from_markdown(file_path, markdown)
else:
created = await self.update_entity_and_observations(
file_path,
async with self._session_scope(session) as active_session:
if is_new:
created = await self.create_entity_from_markdown(
file_path, markdown, session=active_session
)
else:
created = await self.update_entity_and_observations(
file_path,
markdown,
existing_entity=existing_entity,
session=active_session,
)
# Pass the entity through so relation work does not have to rediscover the source row.
return await self.update_entity_relations(
created,
markdown,
existing_entity=existing_entity,
resolve_targets=resolve_relations,
reload_entity=reload_entity,
session=active_session,
)
# Pass the entity through so relation work does not have to rediscover the source row.
return await self.update_entity_relations(
created,
markdown,
resolve_targets=resolve_relations,
reload_entity=reload_entity,
)
async def update_entity_relations(
self,
@@ -1006,6 +1051,7 @@ class EntityService(BaseService[EntityModel]):
*,
resolve_targets: bool = True,
reload_entity: bool = True,
session: AsyncSession | None = None,
) -> EntityModel:
"""Update relations for entity.
@@ -1015,69 +1061,79 @@ class EntityService(BaseService[EntityModel]):
entity_id = entity.id
logger.debug(f"Updating relations for entity: {entity.file_path}")
# Clear existing relations first
await self.relation_repository.delete_outgoing_relations_from_entity(entity_id)
async with self._session_scope(session) as active_session:
# Clear existing relations first
await self.relation_repository.delete_outgoing_relations_from_entity(
active_session, entity_id
)
if markdown.relations:
if resolve_targets:
# Exact target resolution is useful for local sync, but expensive for cloud
# one-file jobs. Cloud can write unresolved rows and let a relation repair pass
# fill in to_id later.
resolved_entities = await asyncio.gather(
*(
self.link_resolver.resolve_link(
rel.target,
strict=True,
load_relations=False,
if markdown.relations:
if resolve_targets:
# Exact target resolution is useful for local sync, but expensive for cloud
# one-file jobs. Cloud can write unresolved rows and let a relation repair pass
# fill in to_id later.
resolved_entities = []
for rel in markdown.relations:
try:
resolved_entities.append(
await self.link_resolver.resolve_link(
rel.target,
strict=True,
load_relations=False,
session=active_session,
)
)
except Exception as exc:
resolved_entities.append(exc)
else:
resolved_entities = [None] * len(markdown.relations)
# Process results and create relation records
relations_to_add = []
for rel, resolved in zip(markdown.relations, resolved_entities):
# Handle exceptions from gather and None results
target_entity: Optional[Entity] = None
if not isinstance(resolved, Exception):
# Relation target resolution keeps exceptions as values so a failed lookup
# becomes an unresolved forward reference instead of aborting the write.
target_entity = resolved
if target_entity is None and not resolve_targets:
target_entity = await self._resolve_deferred_self_relation(
rel.target, entity, session=active_session
)
for rel in markdown.relations
),
return_exceptions=True,
)
else:
resolved_entities = [None] * len(markdown.relations)
# Process results and create relation records
relations_to_add = []
for rel, resolved in zip(markdown.relations, resolved_entities):
# Handle exceptions from gather and None results
target_entity: Optional[Entity] = None
if not isinstance(resolved, BaseException):
# asyncio.gather(..., return_exceptions=True) can return any BaseException.
target_entity = resolved
# if the target is found, store the id
target_id = target_entity.id if target_entity else None
# if the target is found, store the title, otherwise add the target for a "forward link"
target_name = target_entity.title if target_entity else rel.target
if target_entity is None and not resolve_targets:
target_entity = await self._resolve_deferred_self_relation(rel.target, entity)
# Create the relation
relation = Relation(
project_id=self.relation_repository.project_id,
from_id=entity_id,
to_id=target_id,
to_name=target_name,
relation_type=rel.type,
context=rel.context,
)
relations_to_add.append(relation)
# if the target is found, store the id
target_id = target_entity.id if target_entity else None
# if the target is found, store the title, otherwise add the target for a "forward link"
target_name = target_entity.title if target_entity else rel.target
# Batch insert all relations
if relations_to_add:
await self.relation_repository.add_all_ignore_duplicates(
active_session, relations_to_add
)
# Create the relation
relation = Relation(
project_id=self.relation_repository.project_id,
from_id=entity_id,
to_id=target_id,
to_name=target_name,
relation_type=rel.type,
context=rel.context,
)
relations_to_add.append(relation)
if not reload_entity:
return entity
# Batch insert all relations
if relations_to_add:
await self.relation_repository.add_all_ignore_duplicates(relations_to_add)
if not reload_entity:
return entity
# Reload entity with relations via PK lookup (faster than get_by_file_path string match).
reloaded = await self.repository.find_by_ids([entity_id])
return reloaded[0]
# Reload entity with relations via PK lookup (faster than get_by_file_path string match).
reloaded = await self.repository.find_by_ids(active_session, [entity_id])
return reloaded[0]
async def _resolve_deferred_self_relation(
self, target: str, entity: EntityModel
self, target: str, entity: EntityModel, session: AsyncSession | None = None
) -> EntityModel | None:
"""Resolve only self-relations that are safe to identify in deferred mode."""
clean_target = target.strip()
@@ -1101,9 +1157,12 @@ class EntityService(BaseService[EntityModel]):
# Title-only links are ambiguous because Basic Memory allows duplicate titles.
# Collapse them to self only when the title lookup proves this source is the sole candidate;
# otherwise leave the relation unresolved so we do not create a wrong permanent edge.
title_matches = await self.repository.get_by_title(clean_target, load_relations=False)
if len(title_matches) == 1 and title_matches[0].id == entity.id:
return entity
async with self._session_scope(session) as active_session:
title_matches = await self.repository.get_by_title(
active_session, clean_target, load_relations=False
)
if len(title_matches) == 1 and title_matches[0].id == entity.id:
return entity
return None
@@ -1166,43 +1225,46 @@ class EntityService(BaseService[EntityModel]):
file_path = Path(entity.file_path)
current_content, _ = await self.file_service.read_file(file_path)
# --- Prepare Against Explicit Base Content ---
# The edit operation is the semantic step; file/DB writes below are just persistence of that
# accepted result.
prepared = await self.prepare_edit_entity_content(
entity,
current_content,
operation=operation,
content=content,
section=section,
find_text=find_text,
expected_replacements=expected_replacements,
)
async with self._session_scope() as session:
# --- Prepare Against Explicit Base Content ---
# The edit operation is the semantic step; file/DB writes below are just persistence of that
# accepted result.
prepared = await self.prepare_edit_entity_content(
entity,
current_content,
operation=operation,
content=content,
section=section,
find_text=find_text,
expected_replacements=expected_replacements,
session=session,
)
checksum = await self.file_service.write_file(
file_path,
prepared.markdown_content,
)
checksum = await self.file_service.write_file(
file_path,
prepared.markdown_content,
)
# --- Rebuild Structured Knowledge State ---
# Non-fast edits remain fully synchronous locally: once the file write succeeds, we refresh
# observations, relations, and checksum in the same request.
entity = await self.upsert_entity_from_markdown(
file_path,
prepared.entity_markdown,
is_new=False,
)
# --- Rebuild Structured Knowledge State ---
# Non-fast edits remain fully synchronous locally: once the file write succeeds, we refresh
# observations, relations, and checksum in the same request.
entity = await self.upsert_entity_from_markdown(
file_path,
prepared.entity_markdown,
is_new=False,
session=session,
)
entity = await self.repository.update(entity.id, {"checksum": checksum})
if not entity: # pragma: no cover
raise ValueError(f"Failed to update entity checksum after edit: {file_path}")
persisted_content, search_content = await self._read_persisted_write_content(file_path)
entity = await self.repository.update(session, entity.id, {"checksum": checksum})
if not entity: # pragma: no cover
raise ValueError(f"Failed to update entity checksum after edit: {file_path}")
persisted_content, search_content = await self._read_persisted_write_content(file_path)
return EntityWriteResult(
entity=entity,
content=persisted_content,
search_content=search_content,
)
return EntityWriteResult(
entity=entity,
content=persisted_content,
search_content=search_content,
)
def apply_edit_operation(
self,
@@ -1528,11 +1590,12 @@ class EntityService(BaseService[EntityModel]):
updates["checksum"] = new_checksum
# 9. Update database
updated_entity = await self.repository.update(entity.id, updates)
if not updated_entity:
raise ValueError(f"Failed to update entity in database: {entity.id}")
async with self._session_scope() as session:
updated_entity = await self.repository.update(session, entity.id, updates)
if not updated_entity:
raise ValueError(f"Failed to update entity in database: {entity.id}")
return updated_entity
return updated_entity
except Exception as e:
# Rollback: try to restore original file location if move succeeded
@@ -1581,7 +1644,8 @@ class EntityService(BaseService[EntityModel]):
destination_directory = destination_directory.strip("/")
# Find all entities in the source directory
entities = await self.repository.find_by_directory_prefix(source_directory)
async with self._session_scope() as session:
entities = await self.repository.find_by_directory_prefix(session, source_directory)
if not entities:
logger.warning(f"No entities found in directory: {source_directory}")
@@ -1665,7 +1729,8 @@ class EntityService(BaseService[EntityModel]):
directory = directory.strip("/")
# Find all entities in the directory
entities = await self.repository.find_by_directory_prefix(directory)
async with self._session_scope() as session:
entities = await self.repository.find_by_directory_prefix(session, directory)
if not entities:
logger.warning(f"No entities found in directory: {directory}")
+6 -4
View File
@@ -54,13 +54,13 @@ async def reconcile_projects_with_config(app_config: BasicMemoryConfig):
db_path=app_config.database_path,
db_type=db.DatabaseType.FILESYSTEM,
)
project_repository = ProjectRepository(session_maker)
project_repository = ProjectRepository()
# Import ProjectService here to avoid circular imports
from basic_memory.services.project_service import ProjectService
# Create project service and synchronize projects
project_service = ProjectService(repository=project_repository)
project_service = ProjectService(repository=project_repository, session_maker=session_maker)
try:
await project_service.synchronize_projects()
logger.info("Projects successfully reconciled between config and database")
@@ -97,7 +97,7 @@ async def initialize_file_sync(
db_path=app_config.database_path,
db_type=db.DatabaseType.FILESYSTEM,
)
project_repository = ProjectRepository(session_maker)
project_repository = ProjectRepository()
# Filter to constrained project if MCP server was started with --project.
# Applied to both the initial background sync and the watch service so that
@@ -109,12 +109,14 @@ async def initialize_file_sync(
watch_service = WatchService(
app_config=app_config,
project_repository=project_repository,
session_maker=session_maker,
quiet=quiet,
constrained_project=constrained_project,
)
# Get active projects
active_projects = await project_repository.get_active_projects()
async with db.scoped_session(session_maker) as session:
active_projects = await project_repository.get_active_projects(session)
if constrained_project:
active_projects = [p for p in active_projects if p.name == constrained_project]
+116 -73
View File
@@ -1,10 +1,13 @@
"""Service and helpers for resolving markdown links and permalink-like identifiers."""
import uuid as uuid_mod
from typing import Any, Optional, Tuple, Dict
from contextlib import asynccontextmanager
from typing import Any, AsyncIterator, Optional, Tuple, Dict
from loguru import logger
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
from basic_memory import db
from basic_memory.config import BasicMemoryConfig, ConfigManager
from basic_memory.models import Entity, Project
from basic_memory.repository.entity_repository import EntityRepository
@@ -79,17 +82,35 @@ class LinkResolver:
5. Fall back to search for fuzzy matching
"""
def __init__(self, entity_repository: EntityRepository, search_service: SearchService):
def __init__(
self,
entity_repository: EntityRepository,
search_service: SearchService,
session_maker: async_sessionmaker[AsyncSession],
):
"""Initialize with repositories."""
self.entity_repository = entity_repository
self.search_service = search_service
self._project_repository = ProjectRepository(entity_repository.session_maker)
self.session_maker = session_maker
self._project_repository = ProjectRepository()
self._app_config: BasicMemoryConfig = ConfigManager().config
self._project_permalink: Optional[str] = None
self._project_cache_by_identifier: Dict[str, Project] = {}
self._entity_repository_cache: Dict[int, EntityRepository] = {}
self._search_service_cache: Dict[int, SearchService] = {}
@asynccontextmanager
async def _session_scope(
self, session: AsyncSession | None = None
) -> AsyncIterator[AsyncSession]:
"""Use the caller's session or open a service-owned transaction."""
if session is not None:
yield session
return
async with db.scoped_session(self.session_maker) as owned_session:
yield owned_session
async def resolve_link(
self,
link_text: str,
@@ -97,6 +118,7 @@ class LinkResolver:
strict: bool = False,
source_path: Optional[str] = None,
load_relations: bool = True,
session: AsyncSession | None = None,
) -> Optional[Entity]:
"""Resolve a markdown link to a permalink.
@@ -115,36 +137,83 @@ class LinkResolver:
explicit_project_reference = "::" in clean_text
clean_text = normalize_project_reference(clean_text)
# --- External ID Resolution ---
# Try external_id first if identifier looks like a UUID.
# Canonicalize to lowercase-hyphen form so uppercase or unhyphenated
# UUIDs also match the stored external_id values.
try:
canonical_id = str(uuid_mod.UUID(clean_text))
entity = await self.entity_repository.get_by_external_id(
canonical_id,
async with self._session_scope(session) as active_session:
# --- External ID Resolution ---
# Try external_id first if identifier looks like a UUID.
# Canonicalize to lowercase-hyphen form so uppercase or unhyphenated
# UUIDs also match the stored external_id values.
try:
canonical_id = str(uuid_mod.UUID(clean_text))
entity = await self.entity_repository.get_by_external_id(
active_session,
canonical_id,
load_relations=load_relations,
)
if entity:
logger.debug(f"Found entity by external_id: {entity.permalink}")
return entity
except ValueError:
pass
# Trigger: link uses project namespace syntax (project::note)
# Why: treat it as an explicit cross-project reference
# Outcome: resolve only within the referenced project scope
if explicit_project_reference:
project_prefix, remainder = self._split_project_prefix(clean_text)
if not project_prefix:
return None
project_resources = await self._get_project_resources(
project_prefix, active_session
)
if not project_resources:
return None
project, entity_repository, search_service = project_resources
return await self._resolve_in_project(
session=active_session,
entity_repository=entity_repository,
search_service=search_service,
link_text=remainder,
use_search=use_search,
strict=strict,
source_path=None,
project_permalink=project.permalink,
load_relations=load_relations,
)
current_project_permalink = await self._get_current_project_permalink(active_session)
resolved = await self._resolve_in_project(
session=active_session,
entity_repository=self.entity_repository,
search_service=self.search_service,
link_text=clean_text,
use_search=use_search,
strict=strict,
source_path=source_path,
project_permalink=current_project_permalink,
load_relations=load_relations,
)
if entity:
logger.debug(f"Found entity by external_id: {entity.permalink}")
return entity
except ValueError:
pass
if resolved:
return resolved
# Trigger: link uses project namespace syntax (project::note)
# Why: treat it as an explicit cross-project reference
# Outcome: resolve only within the referenced project scope
if explicit_project_reference:
# Trigger: local resolution failed and identifier looks like project/path
# Why: allow explicit project path references without namespace syntax
# Outcome: attempt resolution in the referenced project if it exists
project_prefix, remainder = self._split_project_prefix(clean_text)
if not project_prefix:
return None
project_resources = await self._get_project_resources(project_prefix)
project_resources = await self._get_project_resources(project_prefix, active_session)
if not project_resources:
return None
project, entity_repository, search_service = project_resources
if project.id == self.entity_repository.project_id:
return None
return await self._resolve_in_project(
session=active_session,
entity_repository=entity_repository,
search_service=search_service,
link_text=remainder,
@@ -155,46 +224,6 @@ class LinkResolver:
load_relations=load_relations,
)
current_project_permalink = await self._get_current_project_permalink()
resolved = await self._resolve_in_project(
entity_repository=self.entity_repository,
search_service=self.search_service,
link_text=clean_text,
use_search=use_search,
strict=strict,
source_path=source_path,
project_permalink=current_project_permalink,
load_relations=load_relations,
)
if resolved:
return resolved
# Trigger: local resolution failed and identifier looks like project/path
# Why: allow explicit project path references without namespace syntax
# Outcome: attempt resolution in the referenced project if it exists
project_prefix, remainder = self._split_project_prefix(clean_text)
if not project_prefix:
return None
project_resources = await self._get_project_resources(project_prefix)
if not project_resources:
return None
project, entity_repository, search_service = project_resources
if project.id == self.entity_repository.project_id:
return None
return await self._resolve_in_project(
entity_repository=entity_repository,
search_service=search_service,
link_text=remainder,
use_search=use_search,
strict=strict,
source_path=None,
project_permalink=project.permalink,
load_relations=load_relations,
)
def _normalize_link_text(self, link_text: str) -> Tuple[str, Optional[str]]:
"""Normalize link text and extract alias if present.
@@ -226,6 +255,7 @@ class LinkResolver:
async def _resolve_in_project(
self,
*,
session: AsyncSession,
entity_repository: EntityRepository,
search_service: SearchService,
link_text: str,
@@ -282,6 +312,7 @@ class LinkResolver:
if not relative_path.endswith(".md"):
relative_path_md = f"{relative_path}.md"
entity = await entity_repository.get_by_file_path(
session,
relative_path_md,
load_relations=load_relations,
)
@@ -290,6 +321,7 @@ class LinkResolver:
# Try as-is (already has extension or is a permalink)
entity = await entity_repository.get_by_file_path(
session,
relative_path,
load_relations=load_relations,
)
@@ -307,6 +339,7 @@ class LinkResolver:
# Check permalink match
for candidate_permalink in permalink_candidates:
permalink_entity = await entity_repository.get_by_permalink(
session,
candidate_permalink,
load_relations=load_relations,
)
@@ -315,6 +348,7 @@ class LinkResolver:
# Check title matches
title_entities = await entity_repository.get_by_title(
session,
clean_text,
load_relations=load_relations,
)
@@ -334,6 +368,7 @@ class LinkResolver:
# 1. Try exact permalink match first (most efficient)
for candidate_permalink in permalink_candidates:
entity = await entity_repository.get_by_permalink(
session,
candidate_permalink,
load_relations=load_relations,
)
@@ -343,6 +378,7 @@ class LinkResolver:
# 2. Try exact title match
found = await entity_repository.get_by_title(
session,
clean_text,
load_relations=load_relations,
)
@@ -354,6 +390,7 @@ class LinkResolver:
# 3. Try file path
found_path = await entity_repository.get_by_file_path(
session,
clean_text,
load_relations=load_relations,
)
@@ -365,6 +402,7 @@ class LinkResolver:
if not clean_text.endswith(".md") and "/" in clean_text:
file_path_with_md = f"{clean_text}.md"
found_path_md = await entity_repository.get_by_file_path(
session,
file_path_with_md,
load_relations=load_relations,
)
@@ -380,6 +418,7 @@ class LinkResolver:
if use_search and "*" not in clean_text:
results = await search_service.search(
query=SearchQuery(text=clean_text, entity_types=[SearchItemType.ENTITY]),
session=session,
)
if results:
@@ -392,6 +431,7 @@ class LinkResolver:
)
if best_match.permalink:
return await entity_repository.get_by_permalink(
session,
best_match.permalink,
load_relations=load_relations,
)
@@ -403,7 +443,7 @@ class LinkResolver:
"""Return True when permalinks should include the project slug."""
return self._app_config.permalinks_include_project
async def _get_current_project_permalink(self) -> Optional[str]:
async def _get_current_project_permalink(self, session: AsyncSession) -> Optional[str]:
"""Get and cache the current project's permalink."""
if self._project_permalink is not None:
return self._project_permalink
@@ -412,23 +452,27 @@ class LinkResolver:
if project_id is None: # pragma: no cover
return None # pragma: no cover
project = await self._project_repository.get_by_id(project_id)
project = await self._project_repository.get_by_id(session, project_id)
if project:
self._project_permalink = project.permalink
return self._project_permalink
async def _get_project_by_identifier(self, identifier: str) -> Optional[Project]:
async def _get_project_by_identifier(
self, identifier: str, session: AsyncSession
) -> Optional[Project]:
"""Resolve project by name or permalink."""
cache_key = identifier.strip().lower()
if cache_key in self._project_cache_by_identifier:
return self._project_cache_by_identifier[cache_key]
project = await self._project_repository.get_by_name(identifier)
project = await self._project_repository.get_by_name(session, identifier)
if not project:
project = await self._project_repository.get_by_name_case_insensitive(identifier)
project = await self._project_repository.get_by_name_case_insensitive(
session, identifier
)
if not project:
project = await self._project_repository.get_by_permalink(
generate_permalink(identifier)
session, generate_permalink(identifier)
)
if project:
@@ -436,24 +480,22 @@ class LinkResolver:
return project
async def _get_project_resources(
self, project_identifier: str
self, project_identifier: str, session: AsyncSession
) -> Optional[Tuple[Project, EntityRepository, SearchService]]:
"""Fetch repositories and services scoped to a project."""
project = await self._get_project_by_identifier(project_identifier)
project = await self._get_project_by_identifier(project_identifier, session)
if not project:
return None
entity_repository = self._entity_repository_cache.get(project.id)
if not entity_repository:
entity_repository = EntityRepository(
self.entity_repository.session_maker, project_id=project.id
)
entity_repository = EntityRepository(project_id=project.id)
self._entity_repository_cache[project.id] = entity_repository
search_service = self._search_service_cache.get(project.id)
if not search_service:
search_repository = create_search_repository(
self.entity_repository.session_maker,
self.session_maker,
project_id=project.id,
database_backend=self._app_config.database_backend,
)
@@ -461,6 +503,7 @@ class LinkResolver:
search_repository,
entity_repository,
self.search_service.file_service,
self.session_maker,
)
self._search_service_cache[project.id] = search_service
File diff suppressed because it is too large Load Diff
+40 -8
View File
@@ -3,17 +3,20 @@
import asyncio
import ast
import re
from contextlib import asynccontextmanager
from dataclasses import dataclass
from datetime import datetime
from typing import List, Optional, Set, Dict, Any
from typing import Any, AsyncIterator, List, Optional, Set, Dict
from dateparser import parse
from fastapi import BackgroundTasks
from loguru import logger
from sqlalchemy import text
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
import logfire
from basic_memory import db
from basic_memory.models import Entity
from basic_memory.repository import EntityRepository
from basic_memory.repository.search_repository import (
@@ -116,10 +119,24 @@ class SearchService:
search_repository: SearchRepository,
entity_repository: EntityRepository,
file_service: FileService,
session_maker: async_sessionmaker[AsyncSession],
):
self.repository = search_repository
self.entity_repository = entity_repository
self.file_service = file_service
self.session_maker = session_maker
@asynccontextmanager
async def _session_scope(
self, session: AsyncSession | None = None
) -> AsyncIterator[AsyncSession]:
"""Use the caller's session or open a service-owned transaction."""
if session is not None:
yield session
return
async with db.scoped_session(self.session_maker) as owned_session:
yield owned_session
async def init_search_index(self):
"""Create FTS5 virtual table if it doesn't exist."""
@@ -148,7 +165,8 @@ class SearchService:
# Reindex all entities
logger.debug("Indexing entities")
entities = await self.entity_repository.find_all()
async with self._session_scope() as session:
entities = await self.entity_repository.find_all(session)
for entity in entities:
await self.index_entity(entity, background_tasks)
@@ -235,6 +253,7 @@ class SearchService:
search_text: str | None,
limit: int,
offset: int,
session: AsyncSession | None = None,
) -> List[SearchIndexRow]:
return await self.repository.search(
search_text=search_text,
@@ -250,6 +269,7 @@ class SearchService:
min_similarity=prepared.min_similarity,
limit=limit,
offset=offset,
session=session,
)
async def _count_repository(
@@ -272,7 +292,13 @@ class SearchService:
min_similarity=prepared.min_similarity,
)
async def search(self, query: SearchQuery, limit=10, offset=0) -> List[SearchIndexRow]:
async def search(
self,
query: SearchQuery,
limit=10,
offset=0,
session: AsyncSession | None = None,
) -> List[SearchIndexRow]:
"""Search across all indexed content.
Supports three modes:
@@ -304,6 +330,7 @@ class SearchService:
search_text=strict_search_text,
limit=limit,
offset=offset,
session=session,
)
# Trigger: strict FTS with plain multi-term text returned no results.
@@ -337,6 +364,7 @@ class SearchService:
search_text=relaxed_search_text,
limit=limit,
offset=offset,
session=session,
)
async def count(self, query: SearchQuery) -> int:
@@ -533,7 +561,8 @@ class SearchService:
async def sync_entity_vectors(self, entity_id: int) -> None:
"""Refresh vector chunks for one entity in repositories that support semantic indexing."""
entity = await self.entity_repository.find_by_id(entity_id)
async with self._session_scope() as session:
entity = await self.entity_repository.find_by_id(session, entity_id)
if entity is None:
await self._clear_entity_vectors(entity_id)
return
@@ -557,9 +586,11 @@ class SearchService:
entities_failed=0,
)
entities_by_id = {
entity.id: entity for entity in await self.entity_repository.find_by_ids(entity_ids)
}
async with self._session_scope() as session:
entities_by_id = {
entity.id: entity
for entity in await self.entity_repository.find_by_ids(session, entity_ids)
}
unknown_ids = [entity_id for entity_id in entity_ids if entity_id not in entities_by_id]
opted_out_ids = [
entity_id
@@ -649,7 +680,8 @@ class SearchService:
Returns:
dict with stats: total_entities, embedded, skipped, errors
"""
entities = await self.entity_repository.find_all()
async with self._session_scope() as session:
entities = await self.entity_repository.find_all(session)
entity_ids = [entity.id for entity in entities]
# Clean up stale rows in search_index and search_vector_chunks
+336 -206
View File
@@ -5,6 +5,7 @@ import os
import sys
import time
from collections import OrderedDict
from contextlib import asynccontextmanager
from dataclasses import dataclass, field
from datetime import datetime
from pathlib import Path
@@ -14,6 +15,7 @@ import aiofiles.os
from loguru import logger
from sqlalchemy.exc import IntegrityError
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
import logfire
from basic_memory import db
@@ -133,6 +135,23 @@ class ScanResult:
errors: Dict[str, str] = field(default_factory=dict)
@dataclass(frozen=True)
class _ScanProjectState:
"""Project watermark state needed to choose the filesystem scan strategy."""
last_file_count: int | None
last_scan_timestamp: float | None
@dataclass(frozen=True)
class _ScanEntityState:
"""Entity fields needed for mtime/size/checksum comparison during scan."""
mtime: float | None
size: int | None
checksum: str | None
class _FileServiceIndexWriter(IndexFileWriter):
"""Adapt FileService frontmatter updates to the indexing writer protocol."""
@@ -163,6 +182,7 @@ class SyncService:
project_repository: ProjectRepository,
search_service: SearchService,
file_service: FileService,
session_maker: async_sessionmaker[AsyncSession],
):
self.app_config = app_config
self.entity_service = entity_service
@@ -172,6 +192,7 @@ class SyncService:
self.project_repository = project_repository
self.search_service = search_service
self.file_service = file_service
self.session_maker = session_maker
# Load ignore patterns once at initialization for performance
self._ignore_patterns = load_bmignore_patterns()
# Circuit breaker: track file failures to prevent infinite retry loops
@@ -185,8 +206,21 @@ class SyncService:
relation_repository=relation_repository,
search_service=search_service,
file_writer=_FileServiceIndexWriter(file_service),
session_maker=session_maker,
)
@asynccontextmanager
async def _session_scope(
self, session: AsyncSession | None = None
) -> AsyncIterator[AsyncSession]:
"""Use the caller's session or open a service-owned transaction."""
if session is not None:
yield session
return
async with db.scoped_session(self.session_maker) as owned_session:
yield owned_session
async def _should_skip_file(self, path: str) -> bool:
"""Check if file should be skipped due to repeated failures.
@@ -417,21 +451,23 @@ class SyncService:
with logfire.span("sync.project.update_watermark"):
current_file_count = await self._quick_count_files(directory)
if self.entity_repository.project_id is not None:
project = await self.project_repository.find_by_id(
self.entity_repository.project_id
)
if project:
await self.project_repository.update(
project.id,
{
"last_scan_timestamp": sync_start_timestamp,
"last_file_count": current_file_count,
},
)
logger.debug(
f"Updated scan watermark: timestamp={sync_start_timestamp}, "
f"file_count={current_file_count}"
async with self._session_scope() as session:
project = await self.project_repository.find_by_id(
session, self.entity_repository.project_id
)
if project:
await self.project_repository.update(
session,
project.id,
{
"last_scan_timestamp": sync_start_timestamp,
"last_file_count": current_file_count,
},
)
logger.debug(
f"Updated scan watermark: timestamp={sync_start_timestamp}, "
f"file_count={current_file_count}"
)
duration_ms = int((time.time() - start_time) * 1000)
@@ -536,9 +572,7 @@ class SyncService:
):
shared_permalink_by_path = {
path: permalink
for path, permalink in (
await self.entity_repository.get_file_path_to_permalink_map()
).items()
for path, permalink in (await self._get_file_path_to_permalink_map()).items()
}
for batch in batches:
@@ -578,6 +612,17 @@ class SyncService:
return indexed_entities, skipped_files
async def _get_file_path_to_permalink_map(self) -> dict[str, str | None]:
"""Load current file-path to permalink mappings for markdown batch indexing."""
async with self._session_scope() as session:
permalink_by_path: dict[str, str | None] = {
path: permalink
for path, permalink in (
await self.entity_repository.get_file_path_to_permalink_map(session)
).items()
}
return permalink_by_path
async def _load_index_file_metadata(
self,
paths: list[str],
@@ -726,6 +771,47 @@ class SyncService:
)
)
async def _get_project_scan_state(self) -> _ScanProjectState:
"""Load project watermark fields without holding a session across filesystem work."""
if self.entity_repository.project_id is None:
raise ValueError("Entity repository has no project_id set")
async with self._session_scope() as session:
project = await self.project_repository.find_by_id(
session, self.entity_repository.project_id
)
if project is None:
raise ValueError(f"Project not found: {self.entity_repository.project_id}")
return _ScanProjectState(
last_file_count=project.last_file_count,
last_scan_timestamp=project.last_scan_timestamp,
)
async def _get_entity_scan_state(self, file_path: str) -> _ScanEntityState | None:
"""Load entity comparison fields for one path in a short session."""
async with self._session_scope() as session:
entity = await self.entity_repository.get_by_file_path(session, file_path)
if entity is None:
return None
return _ScanEntityState(
mtime=entity.mtime,
size=entity.size,
checksum=entity.checksum,
)
async def _find_file_paths_by_checksum(self, checksum: str) -> list[str]:
"""Return file paths with this checksum without keeping entities attached."""
async with self._session_scope() as session:
entities = await self.entity_repository.find_by_checksum(session, checksum)
return [entity.file_path for entity in entities]
async def _get_all_file_paths(self) -> list[str]:
"""Load all indexed file paths in a short session."""
async with self._session_scope() as session:
return await self.entity_repository.get_all_file_paths(session)
async def scan(self, directory, force_full: bool = False):
"""Smart scan using watermark and file count for large project optimization.
@@ -754,12 +840,7 @@ class SyncService:
report = SyncReport()
# Get current project to check watermark
if self.entity_repository.project_id is None:
raise ValueError("Entity repository has no project_id set")
project = await self.project_repository.find_by_id(self.entity_repository.project_id)
if project is None:
raise ValueError(f"Project not found: {self.entity_repository.project_id}")
project_state = await self._get_project_scan_state()
with logfire.span("sync.project.select_scan_strategy", force_full=force_full):
# Step 1: Quick file count
@@ -774,29 +855,30 @@ class SyncService:
logger.info("Force full scan requested, bypassing watermark optimization")
scan_coro = self._scan_directory_full(directory)
elif project.last_file_count is None:
elif project_state.last_file_count is None:
# First sync ever → full scan
scan_type = "full_initial"
logger.info("First sync for this project, performing full scan")
scan_coro = self._scan_directory_full(directory)
elif current_count < project.last_file_count:
elif current_count < project_state.last_file_count:
# Files deleted → need full scan to detect which ones
scan_type = "full_deletions"
logger.info(
f"File count decreased ({project.last_file_count}{current_count}), "
f"File count decreased ({project_state.last_file_count}{current_count}), "
f"running full scan to detect deletions"
)
scan_coro = self._scan_directory_full(directory)
elif project.last_scan_timestamp is not None:
elif project_state.last_scan_timestamp is not None:
# Incremental scan: only files modified since last scan
scan_type = "incremental"
logger.debug(
f"Running incremental scan for files modified since {project.last_scan_timestamp}"
f"Running incremental scan for files modified since "
f"{project_state.last_scan_timestamp}"
)
scan_coro = self._scan_directory_modified_since(
directory, project.last_scan_timestamp
directory, project_state.last_scan_timestamp
)
else:
@@ -830,7 +912,7 @@ class SyncService:
stat_info = abs_path.stat()
# Indexed lookup - single file query (not full table scan)
db_entity = await self.entity_repository.get_by_file_path(rel_path)
db_entity = await self._get_entity_scan_state(rel_path)
if db_entity is None:
# New file - need checksum for move detection
@@ -881,26 +963,27 @@ class SyncService:
# Look for existing entity with same checksum but different path
# This could be a move or a copy
existing_entities = await self.entity_repository.find_by_checksum(new_checksum)
existing_file_paths = await self._find_file_paths_by_checksum(new_checksum)
for candidate in existing_entities:
if candidate.file_path == new_path:
for candidate_file_path in existing_file_paths:
if candidate_file_path == new_path:
# Same path, skip (shouldn't happen for "new" files but be safe)
continue
# Check if the old path still exists on disk
old_path_abs = directory / candidate.file_path
old_path_abs = directory / candidate_file_path
if old_path_abs.exists():
# Original still exists → this is a copy, not a move
logger.trace(
f"File copy detected (not move): {candidate.file_path} copied to {new_path}"
f"File copy detected (not move): "
f"{candidate_file_path} copied to {new_path}"
)
continue
# Original doesn't exist → this is a move!
report.moves[candidate.file_path] = new_path
report.moves[candidate_file_path] = new_path
report.new.remove(new_path)
logger.trace(f"Move detected: {candidate.file_path} -> {new_path}")
logger.trace(f"Move detected: {candidate_file_path} -> {new_path}")
break # Only match first candidate
# Step 5: Detect deletions (only for full scans)
@@ -908,7 +991,7 @@ class SyncService:
if scan_type in ("full_initial", "full_deletions", "full_fallback", "full_forced"):
with logfire.span("sync.project.detect_deletions", scan_type=scan_type):
# Use optimized query for just file paths (not full entities)
db_file_paths = await self.entity_repository.get_all_file_paths()
db_file_paths = await self._get_all_file_paths()
logger.debug(f"Found {len(db_file_paths)} db paths for deletion detection")
for db_path in db_file_paths:
@@ -1086,7 +1169,8 @@ class SyncService:
initial_markdown_content = initial_markdown_bytes.decode("utf-8")
file_metadata = await self.file_service.get_file_metadata(path)
initial_checksum = await compute_checksum(initial_markdown_bytes)
existing_entity = await self.entity_repository.get_by_file_path(path)
async with self._session_scope() as session:
existing_entity = await self.entity_repository.get_by_file_path(session, path)
if existing_entity is not None and existing_entity.checksum == initial_checksum:
logger.debug(
f"Markdown sync skipped unchanged file: path={path}, "
@@ -1122,22 +1206,26 @@ class SyncService:
else initial_markdown_content
)
file_metadata = await self.file_service.get_file_metadata(path)
refreshed_entities = await self.entity_repository.find_by_ids([indexed.entity_id])
if len(refreshed_entities) != 1: # pragma: no cover
raise ValueError(f"Failed to reload synced markdown entity for {path}")
# Trigger: markdown sync may have rewritten frontmatter after the initial file metadata load.
# Why: the batch indexer persisted checksum/path data from the pre-rewrite IndexInputFile.
# Outcome: refresh size and mtime from the file as it actually exists on disk now.
updated_entity = await self.entity_repository.update(
refreshed_entities[0].id,
{
"checksum": indexed.checksum,
"created_at": file_metadata.created_at,
"updated_at": file_metadata.modified_at,
"mtime": file_metadata.modified_at.timestamp(),
"size": file_metadata.size,
},
)
async with self._session_scope() as session:
refreshed_entities = await self.entity_repository.find_by_ids(
session, [indexed.entity_id]
)
if len(refreshed_entities) != 1: # pragma: no cover
raise ValueError(f"Failed to reload synced markdown entity for {path}")
# Trigger: markdown sync may have rewritten frontmatter after the initial file metadata load.
# Why: the batch indexer persisted checksum/path data from the pre-rewrite IndexInputFile.
# Outcome: refresh size and mtime from the file as it actually exists on disk now.
updated_entity = await self.entity_repository.update(
session,
refreshed_entities[0].id,
{
"checksum": indexed.checksum,
"created_at": file_metadata.created_at,
"updated_at": file_metadata.modified_at,
"mtime": file_metadata.modified_at.timestamp(),
"size": file_metadata.size,
},
)
if updated_entity is None: # pragma: no cover
raise ValueError(f"Failed to update markdown entity metadata for {path}")
@@ -1197,19 +1285,21 @@ class SyncService:
file_path = Path(path)
try:
entity = await self.entity_repository.add(
Entity(
note_type="file",
file_path=path,
checksum=checksum,
title=file_path.name,
created_at=created,
updated_at=modified,
content_type=content_type,
mtime=file_metadata.modified_at.timestamp(),
size=file_metadata.size,
async with self._session_scope() as session:
entity = await self.entity_repository.add(
session,
Entity(
note_type="file",
file_path=path,
checksum=checksum,
title=file_path.name,
created_at=created,
updated_at=modified,
content_type=content_type,
mtime=file_metadata.modified_at.timestamp(),
size=file_metadata.size,
),
)
)
return entity, checksum
except IntegrityError as e:
# Handle race condition where entity was created by another process
@@ -1224,22 +1314,26 @@ class SyncService:
f"Entity already exists for file_path={path}, updating instead of creating"
)
# Treat as update instead of create
entity = await self.entity_repository.get_by_file_path(path)
if entity is None: # pragma: no cover
logger.error(f"Entity not found after constraint violation, path={path}")
raise ValueError(f"Entity not found after constraint violation: {path}")
async with self._session_scope() as session:
entity = await self.entity_repository.get_by_file_path(session, path)
if entity is None: # pragma: no cover
logger.error(
f"Entity not found after constraint violation, path={path}"
)
raise ValueError(f"Entity not found after constraint violation: {path}")
# Re-get file metadata since we're in update path
file_metadata_for_update = await self.file_service.get_file_metadata(path)
updated = await self.entity_repository.update(
entity.id,
{
"file_path": path,
"checksum": checksum,
"mtime": file_metadata_for_update.modified_at.timestamp(),
"size": file_metadata_for_update.size,
},
)
# Re-get file metadata since we're in update path
file_metadata_for_update = await self.file_service.get_file_metadata(path)
updated = await self.entity_repository.update(
session,
entity.id,
{
"file_path": path,
"checksum": checksum,
"mtime": file_metadata_for_update.modified_at.timestamp(),
"size": file_metadata_for_update.size,
},
)
if updated is None: # pragma: no cover
logger.error(f"Failed to update entity, entity_id={entity.id}, path={path}")
@@ -1254,23 +1348,25 @@ class SyncService:
file_metadata = await self.file_service.get_file_metadata(path)
modified = file_metadata.modified_at
entity = await self.entity_repository.get_by_file_path(path)
if entity is None: # pragma: no cover
logger.error(f"Entity not found for existing file, path={path}")
raise ValueError(f"Entity not found for existing file: {path}")
async with self._session_scope() as session:
entity = await self.entity_repository.get_by_file_path(session, path)
if entity is None: # pragma: no cover
logger.error(f"Entity not found for existing file, path={path}")
raise ValueError(f"Entity not found for existing file: {path}")
# Update checksum, modification time, and file metadata from file system
# Store mtime/size for efficient change detection in future scans
updated = await self.entity_repository.update(
entity.id,
{
"file_path": path,
"checksum": checksum,
"updated_at": modified,
"mtime": file_metadata.modified_at.timestamp(),
"size": file_metadata.size,
},
)
# Update checksum, modification time, and file metadata from file system
# Store mtime/size for efficient change detection in future scans
updated = await self.entity_repository.update(
session,
entity.id,
{
"file_path": path,
"checksum": checksum,
"updated_at": modified,
"mtime": file_metadata.modified_at.timestamp(),
"size": file_metadata.size,
},
)
if updated is None: # pragma: no cover
logger.error(f"Failed to update entity, entity_id={entity.id}, path={path}")
@@ -1282,7 +1378,8 @@ class SyncService:
"""Handle complete entity deletion including search index cleanup."""
# First get entity to get permalink before deletion
entity = await self.entity_repository.get_by_file_path(file_path)
async with self._session_scope() as session:
entity = await self.entity_repository.get_by_file_path(session, file_path)
if entity:
logger.info(
f"Deleting entity with file_path={file_path}, entity_id={entity.id}, permalink={entity.permalink}"
@@ -1312,106 +1409,130 @@ class SyncService:
async def handle_move(self, old_path, new_path):
logger.debug("Moving entity", old_path=old_path, new_path=new_path)
entity = await self.entity_repository.get_by_file_path(old_path)
if entity:
# Check if destination path is already occupied by another entity
existing_at_destination = await self.entity_repository.get_by_file_path(new_path)
if existing_at_destination and existing_at_destination.id != entity.id:
# Handle the conflict - this could be a file swap or replacement scenario
logger.warning(
f"File path conflict detected during move: "
f"entity_id={entity.id} trying to move from '{old_path}' to '{new_path}', "
f"but entity_id={existing_at_destination.id} already occupies '{new_path}'"
updated_entity = None
async with self._session_scope() as session:
entity = await self.entity_repository.get_by_file_path(session, old_path)
if entity:
# Check if destination path is already occupied by another entity
existing_at_destination = await self.entity_repository.get_by_file_path(
session, new_path
)
# Check if this is a file swap (the destination entity is being moved to our old path)
# This would indicate a simultaneous move operation
old_path_after_swap = await self.entity_repository.get_by_file_path(old_path)
if old_path_after_swap and old_path_after_swap.id == existing_at_destination.id:
logger.info(f"Detected file swap between '{old_path}' and '{new_path}'")
# This is a swap scenario - both moves should succeed
# We'll allow this to proceed since the other file has moved out
else:
# This is a conflict where the destination is occupied
raise ValueError(
f"Cannot move entity from '{old_path}' to '{new_path}': "
f"destination path is already occupied by another file. "
f"This may be caused by: "
f"1. Conflicting file names with different character encodings, "
f"2. Case sensitivity differences (e.g., 'Finance/' vs 'finance/'), "
f"3. Character conflicts between hyphens in filenames and generated permalinks, "
f"4. Files with similar names containing special characters. "
f"Try renaming one of the conflicting files to resolve this issue."
if existing_at_destination and existing_at_destination.id != entity.id:
# Handle the conflict - this could be a file swap or replacement scenario
logger.warning(
f"File path conflict detected during move: "
f"entity_id={entity.id} trying to move from '{old_path}' to '{new_path}', "
f"but entity_id={existing_at_destination.id} already occupies '{new_path}'"
)
# Update file_path in all cases
updates = {"file_path": new_path}
# Check if this is a file swap (the destination entity is being moved to our old path)
# This would indicate a simultaneous move operation
old_path_after_swap = await self.entity_repository.get_by_file_path(
session, old_path
)
if old_path_after_swap and old_path_after_swap.id == existing_at_destination.id:
logger.info(f"Detected file swap between '{old_path}' and '{new_path}'")
# This is a swap scenario - both moves should succeed
# We'll allow this to proceed since the other file has moved out
else:
# This is a conflict where the destination is occupied
raise ValueError(
f"Cannot move entity from '{old_path}' to '{new_path}': "
f"destination path is already occupied by another file. "
f"This may be caused by: "
f"1. Conflicting file names with different character encodings, "
f"2. Case sensitivity differences (e.g., 'Finance/' vs 'finance/'), "
f"3. Character conflicts between hyphens in filenames and generated permalinks, "
f"4. Files with similar names containing special characters. "
f"Try renaming one of the conflicting files to resolve this issue."
)
# If configured, also update permalink to match new path
if (
self.app_config.update_permalinks_on_move
and not self.app_config.disable_permalinks
and self.file_service.is_markdown(new_path)
):
# generate new permalink value - skip conflict checks during bulk sync
new_permalink = await self.entity_service.resolve_permalink(
new_path, skip_conflict_check=True
)
# Update file_path in all cases
updates = {"file_path": new_path}
# write to file and get new checksum
new_checksum = await self.file_service.update_frontmatter(
new_path, {"permalink": new_permalink}
)
# If configured, also update permalink to match new path
if (
self.app_config.update_permalinks_on_move
and not self.app_config.disable_permalinks
and self.file_service.is_markdown(new_path)
):
# generate new permalink value - skip conflict checks during bulk sync
new_permalink = await self.entity_service.resolve_permalink(
new_path,
skip_conflict_check=True,
session=session,
)
updates["permalink"] = new_permalink
updates["checksum"] = new_checksum
# write to file and get new checksum
new_checksum = await self.file_service.update_frontmatter(
new_path, {"permalink": new_permalink}
)
logger.info(
f"Updating permalink on move,old_permalink={entity.permalink}"
f"new_permalink={new_permalink}"
f"new_checksum={new_checksum}"
)
updates["permalink"] = new_permalink
updates["checksum"] = new_checksum
try:
updated = await self.entity_repository.update(entity.id, updates)
except Exception as e:
# Catch any database integrity errors and provide helpful context
if "UNIQUE constraint failed" in str(e):
logger.info(
f"Updating permalink on move,old_permalink={entity.permalink}"
f"new_permalink={new_permalink}"
f"new_checksum={new_checksum}"
)
try:
updated = await self.entity_repository.update(session, entity.id, updates)
except Exception as e:
# Catch any database integrity errors and provide helpful context
if "UNIQUE constraint failed" in str(e):
logger.error(
f"Database constraint violation during move: "
f"entity_id={entity.id}, old_path='{old_path}', new_path='{new_path}'"
)
raise ValueError(
f"Cannot complete move from '{old_path}' to '{new_path}': "
f"a database constraint was violated. This usually indicates "
f"a file path or permalink conflict. Please check for: "
f"1. Duplicate file names, "
f"2. Case sensitivity issues (e.g., 'File.md' vs 'file.md'), "
f"3. Character encoding conflicts in file names."
) from e
else:
# Re-raise other exceptions as-is
raise
if updated is None: # pragma: no cover
logger.error(
f"Database constraint violation during move: "
f"entity_id={entity.id}, old_path='{old_path}', new_path='{new_path}'"
"Failed to update entity path"
f"entity_id={entity.id}"
f"old_path={old_path}"
f"new_path={new_path}"
)
raise ValueError(
f"Cannot complete move from '{old_path}' to '{new_path}': "
f"a database constraint was violated. This usually indicates "
f"a file path or permalink conflict. Please check for: "
f"1. Duplicate file names, "
f"2. Case sensitivity issues (e.g., 'File.md' vs 'file.md'), "
f"3. Character encoding conflicts in file names."
) from e
else:
# Re-raise other exceptions as-is
raise
raise ValueError(f"Failed to update entity path for ID {entity.id}")
if updated is None: # pragma: no cover
logger.error(
"Failed to update entity path"
f"entity_id={entity.id}"
f"old_path={old_path}"
f"new_path={new_path}"
logger.debug(
"Entity path updated"
f"entity_id={entity.id} "
f"permalink={entity.permalink} "
f"old_path={old_path} "
f"new_path={new_path} "
)
raise ValueError(f"Failed to update entity path for ID {entity.id}")
logger.debug(
"Entity path updated"
f"entity_id={entity.id} "
f"permalink={entity.permalink} "
f"old_path={old_path} "
f"new_path={new_path} "
updated_entity = updated
# Search indexing owns its own repository lifecycle, so run it after the
# move transaction has released the database connection.
if updated_entity is not None:
await self.search_service.index_entity(updated_entity)
async def _find_unresolved_relations_for_entity(self, entity_id: int):
"""Load unresolved relations for one entity in a service-owned session."""
async with self._session_scope() as session:
return await self.relation_repository.find_unresolved_relations_for_entity(
session, entity_id
)
# update search index
await self.search_service.index_entity(updated)
async def _find_unresolved_relations(self):
"""Load unresolved relations in a service-owned session."""
async with self._session_scope() as session:
return await self.relation_repository.find_unresolved_relations(session)
async def resolve_relations(self, entity_id: int | None = None) -> set[int]:
"""Try to resolve unresolved relations.
@@ -1427,16 +1548,14 @@ class SyncService:
if entity_id:
# Only get unresolved relations for the specific entity
unresolved_relations = (
await self.relation_repository.find_unresolved_relations_for_entity(entity_id)
)
unresolved_relations = await self._find_unresolved_relations_for_entity(entity_id)
logger.info(
f"Resolving forward references for entity {entity_id}",
count=len(unresolved_relations),
)
else:
# Get all unresolved relations (original behavior)
unresolved_relations = await self.relation_repository.find_unresolved_relations()
unresolved_relations = await self._find_unresolved_relations()
logger.info("Resolving all forward references", count=len(unresolved_relations))
for relation in unresolved_relations:
@@ -1454,9 +1573,10 @@ class SyncService:
# the most tokens, polluting the graph with confidently-wrong edges that no
# audit catches. Leaving such relations unresolved keeps to_id=NULL so they
# surface as forward references and can be fixed by the producer.
resolved_entity = await self.entity_service.link_resolver.resolve_link(
relation.to_name, strict=True
)
async with self._session_scope() as session:
resolved_entity = await self.entity_service.link_resolver.resolve_link(
relation.to_name, strict=True, session=session
)
# ignore reference to self
if resolved_entity and resolved_entity.id != relation.from_id:
@@ -1469,13 +1589,15 @@ class SyncService:
f"resolved_title={resolved_entity.title}",
)
try:
await self.relation_repository.update(
relation.id,
{
"to_id": resolved_entity.id,
"to_name": resolved_entity.title,
},
)
async with self._session_scope() as session:
await self.relation_repository.update(
session,
relation.id,
{
"to_id": resolved_entity.id,
"to_name": resolved_entity.title,
},
)
affected_entity_ids.add(relation.from_id)
except IntegrityError:
with logfire.span(
@@ -1497,7 +1619,8 @@ class SyncService:
f"resolved_to_id={resolved_entity.id}"
)
try:
await self.relation_repository.delete(relation.id)
async with self._session_scope() as session:
await self.relation_repository.delete(session, relation.id)
except Exception as e:
with logfire.span(
"sync.relation.cleanup_failure",
@@ -1511,7 +1634,8 @@ class SyncService:
affected_entity_ids.add(relation.from_id)
for affected_entity_id in sorted(affected_entity_ids):
source_entity = await self.entity_repository.find_by_id(affected_entity_id)
async with self._session_scope() as session:
source_entity = await self.entity_repository.find_by_id(session, affected_entity_id)
if source_entity is not None:
await self.search_service.index_entity(source_entity)
@@ -1741,17 +1865,19 @@ async def get_sync_service(project: Project) -> SyncService: # pragma: no cover
file_service = FileService(project_path, markdown_processor, app_config=app_config)
# Initialize repositories
entity_repository = EntityRepository(session_maker, project_id=project.id)
observation_repository = ObservationRepository(session_maker, project_id=project.id)
relation_repository = RelationRepository(session_maker, project_id=project.id)
entity_repository = EntityRepository(project_id=project.id)
observation_repository = ObservationRepository(project_id=project.id)
relation_repository = RelationRepository(project_id=project.id)
search_repository = create_search_repository(
session_maker, project_id=project.id, app_config=app_config
)
project_repository = ProjectRepository(session_maker)
project_repository = ProjectRepository()
# Initialize services
search_service = SearchService(search_repository, entity_repository, file_service)
link_resolver = LinkResolver(entity_repository, search_service)
search_service = SearchService(
search_repository, entity_repository, file_service, session_maker
)
link_resolver = LinkResolver(entity_repository, search_service, session_maker)
# Initialize services
entity_service = EntityService(
@@ -1761,6 +1887,9 @@ async def get_sync_service(project: Project) -> SyncService: # pragma: no cover
relation_repository,
file_service,
link_resolver,
session_maker,
search_service=search_service,
app_config=app_config,
)
# Create sync service
@@ -1773,6 +1902,7 @@ async def get_sync_service(project: Project) -> SyncService: # pragma: no cover
project_repository=project_repository,
search_service=search_service,
file_service=file_service,
session_maker=session_maker,
)
return sync_service
+34 -5
View File
@@ -2,14 +2,27 @@
import asyncio
import os
from contextlib import asynccontextmanager
from collections import defaultdict
from datetime import datetime
from pathlib import Path
from typing import List, Optional, Set, Sequence, Callable, Awaitable, TYPE_CHECKING
from typing import (
AsyncIterator,
List,
Optional,
Set,
Sequence,
Callable,
Awaitable,
TYPE_CHECKING,
)
if TYPE_CHECKING:
from basic_memory.sync.sync_service import SyncService
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
from basic_memory import db
from basic_memory.config import BasicMemoryConfig, WATCH_STATUS_JSON
from basic_memory.ignore_utils import load_gitignore_patterns, should_ignore_path
from basic_memory.models import Project
@@ -83,12 +96,14 @@ class WatchService:
self,
app_config: BasicMemoryConfig,
project_repository: ProjectRepository,
session_maker: async_sessionmaker[AsyncSession],
quiet: bool = False,
sync_service_factory: Optional[SyncServiceFactory] = None,
constrained_project: Optional[str] = None,
):
self.app_config = app_config
self.project_repository = project_repository
self.session_maker = session_maker
self.state = WatchServiceState()
self.status_path = app_config.data_dir_path / WATCH_STATUS_JSON
self.status_path.parent.mkdir(parents=True, exist_ok=True)
@@ -104,6 +119,12 @@ class WatchService:
# quiet mode for mcp so it doesn't mess up stdout
self.console = Console(quiet=quiet)
@asynccontextmanager
async def _session_scope(self) -> AsyncIterator[AsyncSession]:
"""Open a service-owned transaction."""
async with db.scoped_session(self.session_maker) as session:
yield session
async def _get_sync_service(self, project: Project) -> "SyncService":
"""Get sync service for a project, using factory if provided."""
if self._sync_service_factory:
@@ -193,7 +214,8 @@ class WatchService:
projects with a local bisync copy keep their absolute path and are
still watched.
"""
projects = await self.project_repository.get_active_projects()
async with self._session_scope() as session:
projects = await self.project_repository.get_active_projects(session)
if self.constrained_project:
projects = [p for p in projects if p.name == self.constrained_project]
@@ -393,7 +415,8 @@ class WatchService:
# Avoid mutating `adds` while iterating (can skip items).
reclassified_as_modified: List[str] = []
for added_path in list(adds): # pragma: no cover TODO add test
entity = await sync_service.entity_repository.get_by_file_path(added_path)
async with sync_service._session_scope() as session:
entity = await sync_service.entity_repository.get_by_file_path(session, added_path)
if entity is not None:
logger.debug(f"Existing file will be processed as modified, path={added_path}")
reclassified_as_modified.append(added_path)
@@ -424,7 +447,10 @@ class WatchService:
continue # pragma: no cover
# Skip directories for deleted paths (based on entity type in db)
deleted_entity = await sync_service.entity_repository.get_by_file_path(deleted_path)
async with sync_service._session_scope() as session:
deleted_entity = await sync_service.entity_repository.get_by_file_path(
session, deleted_path
)
if deleted_entity is None:
# If this was a directory, it wouldn't have an entity
logger.debug("Skipping unknown path for move detection", path=deleted_path)
@@ -483,7 +509,10 @@ class WatchService:
# Check if this was a directory - skip if so
# (we can't tell if the deleted path was a directory since it no longer exists,
# so we check if there's an entity in the database for it)
entity = await sync_service.entity_repository.get_by_file_path(path)
async with sync_service._session_scope() as session:
entity = await sync_service.entity_repository.get_by_file_path(
session, path
)
if entity is None:
# No entity means this was likely a directory - skip it
logger.debug(
+14 -6
View File
@@ -326,9 +326,12 @@ async def test_project(config_home, engine_factory) -> Project:
"is_default": True,
}
engine, session_maker = engine_factory
project_repository = ProjectRepository(session_maker)
project = await project_repository.create(project_data)
_, session_maker = engine_factory
project_repository = ProjectRepository()
from basic_memory import db
async with db.scoped_session(session_maker) as session:
project = await project_repository.create(session, project_data)
return project
@@ -461,12 +464,12 @@ async def search_service(engine_factory, test_project, app_config):
from basic_memory.repository.search_repository import create_search_repository
engine, session_maker = engine_factory
_, session_maker = engine_factory
# Use factory function to create appropriate search repository
search_repository = create_search_repository(session_maker, project_id=test_project.id)
entity_repository = EntityRepository(session_maker, project_id=test_project.id)
entity_repository = EntityRepository(project_id=test_project.id)
# Create file service
entity_parser = EntityParser(Path(test_project.path))
@@ -474,7 +477,12 @@ async def search_service(engine_factory, test_project, app_config):
file_service = FileService(Path(test_project.path), markdown_processor)
# Create and initialize search service
service = SearchService(search_repository, entity_repository, file_service)
service = SearchService(
search_repository,
entity_repository,
file_service,
session_maker=session_maker,
)
await service.init_search_index()
return service
@@ -13,7 +13,9 @@ import pytest
from fastmcp import Client
from unittest.mock import patch
from basic_memory import db
from basic_memory.config import ConfigManager, BasicMemoryConfig
from basic_memory.repository.project_repository import ProjectRepository
@pytest.mark.asyncio
@@ -53,19 +55,19 @@ async def test_explicit_project_overrides_default(
"""Test that explicit project parameter overrides default_project."""
engine, session_maker = engine_factory
from basic_memory.repository.project_repository import ProjectRepository
project_repository = ProjectRepository(session_maker)
other_project = await project_repository.create(
{
"name": "other-project",
"description": "Second project for testing",
"path": str(config_home / "other-project"),
"is_active": True,
"is_default": False,
}
)
project_repository = ProjectRepository()
async with db.scoped_session(session_maker) as session:
other_project = await project_repository.create(
session,
{
"name": "other-project",
"description": "Second project for testing",
"path": str(config_home / "other-project"),
"is_active": True,
"is_default": False,
},
)
mock_config = BasicMemoryConfig(
default_project=test_project.name,
@@ -128,19 +130,19 @@ async def test_cli_constraint_overrides_default_project(
"""Test that CLI --project constraint overrides default_project."""
engine, session_maker = engine_factory
from basic_memory.repository.project_repository import ProjectRepository
project_repository = ProjectRepository(session_maker)
other_project = await project_repository.create(
{
"name": "cli-project",
"description": "Project for CLI constraint testing",
"path": str(config_home / "cli-project"),
"is_active": True,
"is_default": False,
}
)
project_repository = ProjectRepository()
async with db.scoped_session(session_maker) as session:
other_project = await project_repository.create(
session,
{
"name": "cli-project",
"description": "Project for CLI constraint testing",
"path": str(config_home / "cli-project"),
"is_active": True,
"is_default": False,
},
)
os.environ["BASIC_MEMORY_MCP_PROJECT"] = other_project.name
@@ -250,29 +252,31 @@ async def test_project_resolution_hierarchy(
"""Test the complete three-tier project resolution hierarchy."""
engine, session_maker = engine_factory
from basic_memory.repository.project_repository import ProjectRepository
project_repository = ProjectRepository(session_maker)
project_repository = ProjectRepository()
default_project = test_project
cli_project = await project_repository.create(
{
"name": "cli-hierarchy-project",
"description": "Project for CLI hierarchy testing",
"path": str(config_home / "cli-hierarchy-project"),
"is_active": True,
"is_default": False,
}
)
explicit_project = await project_repository.create(
{
"name": "explicit-hierarchy-project",
"description": "Project for explicit hierarchy testing",
"path": str(config_home / "explicit-hierarchy-project"),
"is_active": True,
"is_default": False,
}
)
async with db.scoped_session(session_maker) as session:
cli_project = await project_repository.create(
session,
{
"name": "cli-hierarchy-project",
"description": "Project for CLI hierarchy testing",
"path": str(config_home / "cli-hierarchy-project"),
"is_active": True,
"is_default": False,
},
)
explicit_project = await project_repository.create(
session,
{
"name": "explicit-hierarchy-project",
"description": "Project for explicit hierarchy testing",
"path": str(config_home / "explicit-hierarchy-project"),
"is_active": True,
"is_default": False,
},
)
mock_config = BasicMemoryConfig(
default_project=default_project.name,
@@ -17,6 +17,7 @@ relation rather than preserving prose as a custom relation type.
import pytest
from fastmcp import Client
from basic_memory import db
from basic_memory.repository.relation_repository import RelationRepository
@@ -73,9 +74,10 @@ async def test_edit_note_handles_long_prose_around_wikilink(
assert "Edited note (append)" in edit_result.content[0].text
_, session_maker = engine_factory
relation_repository = RelationRepository(session_maker, project_id=test_project.id)
links_to_relations = await relation_repository.find_by_type("links_to")
prose_type_relations = await relation_repository.find_by_type(long_prose)
relation_repository = RelationRepository(project_id=test_project.id)
async with db.scoped_session(session_maker) as session:
links_to_relations = await relation_repository.find_by_type(session, "links_to")
prose_type_relations = await relation_repository.find_by_type(session, long_prose)
assert any(relation.to_name == "Some Note Title" for relation in links_to_relations)
assert not prose_type_relations
@@ -13,6 +13,7 @@ import pytest_asyncio
from fastmcp import Client
from httpx import ASGITransport, AsyncClient as HttpxAsyncClient
from basic_memory import db
from basic_memory.config import BasicMemoryConfig, ConfigManager, ProjectEntry
from basic_memory.mcp import async_client
from basic_memory.mcp import project_context
@@ -93,16 +94,18 @@ async def alternate_project(
alternate_path.mkdir(parents=True, exist_ok=True)
_, session_maker = engine_factory
project_repository = ProjectRepository(session_maker)
project = await project_repository.create(
{
"name": "alternate-project",
"description": "Non-default project for prefix routing tests",
"path": str(alternate_path),
"is_active": True,
"is_default": False,
}
)
project_repository = ProjectRepository()
async with db.scoped_session(session_maker) as session:
project = await project_repository.create(
session,
{
"name": "alternate-project",
"description": "Non-default project for prefix routing tests",
"path": str(alternate_path),
"is_active": True,
"is_default": False,
},
)
app_config.projects[project.name] = ProjectEntry(path=str(alternate_path))
config_manager.save_config(app_config)
+19 -12
View File
@@ -242,16 +242,18 @@ async def create_search_service(
engine, session_maker = engine_factory_result
# Create test project
project_repo = ProjectRepository(session_maker)
project = await project_repo.create(
{
"name": "bench-project",
"description": "Semantic benchmark project",
"path": str(tmp_path),
"is_active": True,
"is_default": True,
}
)
project_repo = ProjectRepository()
async with db.scoped_session(session_maker) as session:
project = await project_repo.create(
session,
{
"name": "bench-project",
"description": "Semantic benchmark project",
"path": str(tmp_path),
"is_active": True,
"is_default": True,
},
)
# Build app config
semantic_enabled = combo.provider_name is not None
@@ -290,11 +292,16 @@ async def create_search_service(
repo._vector_tables_initialized = False
search_repo = repo
entity_repo = EntityRepository(session_maker, project_id=project.id)
entity_repo = EntityRepository(project_id=project.id)
entity_parser = EntityParser(tmp_path)
markdown_processor = MarkdownProcessor(entity_parser)
file_service = FileService(tmp_path, markdown_processor)
service = SearchService(search_repo, entity_repo, file_service)
service = SearchService(
search_repo,
entity_repo,
file_service,
session_maker=session_maker,
)
await service.init_search_index()
return service
+14 -10
View File
@@ -12,6 +12,8 @@ from __future__ import annotations
import random
from dataclasses import dataclass
from basic_memory import db
# --- Topic vocabulary ---
# Each topic has primary terms (strongly associated) and secondary terms
@@ -395,16 +397,18 @@ Related concepts: {keyword_line}.
else:
content = build_benchmark_content(topic, terms, note_index)
entity = await search_service.entity_repository.create(
{
"title": f"{topic.title()} Benchmark Note {note_index}",
"note_type": "benchmark",
"entity_metadata": {"tags": ["benchmark", topic], "status": "active"},
"content_type": "text/markdown",
"permalink": permalink,
"file_path": f"{permalink}.md",
}
)
async with db.scoped_session(search_service.session_maker) as session:
entity = await search_service.entity_repository.create(
session,
{
"title": f"{topic.title()} Benchmark Note {note_index}",
"note_type": "benchmark",
"entity_metadata": {"tags": ["benchmark", topic], "status": "active"},
"content_type": "text/markdown",
"permalink": permalink,
"file_path": f"{permalink}.md",
},
)
await search_service.index_entity_data(entity, content=content)
# Sync vector embeddings when semantic search is enabled
+13 -10
View File
@@ -13,6 +13,7 @@ from typing import Any, cast
import pytest
from basic_memory import db
from basic_memory.config import DatabaseBackend
from basic_memory.schemas.search import SearchItemType, SearchQuery, SearchRetrievalMode
@@ -108,16 +109,18 @@ async def seed_diagnostic_notes(search_service):
entities = []
for note in notes:
entity = await search_service.entity_repository.create(
{
"title": note["title"],
"note_type": "note",
"entity_metadata": {"tags": note.get("tags", [])},
"content_type": "text/markdown",
"permalink": note["permalink"],
"file_path": f"{note['permalink']}.md",
}
)
async with db.scoped_session(search_service.session_maker) as session:
entity = await search_service.entity_repository.create(
session,
{
"title": note["title"],
"note_type": "note",
"entity_metadata": {"tags": note.get("tags", [])},
"content_type": "text/markdown",
"permalink": note["permalink"],
"file_path": f"{note['permalink']}.md",
},
)
await search_service.index_entity_data(entity, content=note["content"])
await search_service.sync_entity_vectors(entity.id)
entities.append(entity)
+25 -20
View File
@@ -17,6 +17,7 @@ from typing import Any, cast
import pytest
from basic_memory import db
from basic_memory.config import DatabaseBackend
from basic_memory.schemas.search import SearchItemType, SearchQuery, SearchRetrievalMode
@@ -199,16 +200,18 @@ async def test_postgres_vector_dimension_detection(postgres_engine_factory, tmp_
repo = cast(Any, search_service.repository)
# First entity triggers _ensure_vector_tables
entity = await search_service.entity_repository.create(
{
"title": "Dimension Test Note",
"note_type": "benchmark",
"entity_metadata": {"tags": ["test"]},
"content_type": "text/markdown",
"permalink": "bench/dim-test",
"file_path": "bench/dim-test.md",
}
)
async with db.scoped_session(search_service.session_maker) as session:
entity = await search_service.entity_repository.create(
session,
{
"title": "Dimension Test Note",
"note_type": "benchmark",
"entity_metadata": {"tags": ["test"]},
"content_type": "text/markdown",
"permalink": "bench/dim-test",
"file_path": "bench/dim-test.md",
},
)
content = build_benchmark_content("auth", TOPIC_TERMS["auth"], 0)
await search_service.index_entity_data(entity, content=content)
await search_service.sync_entity_vectors(entity.id)
@@ -242,16 +245,18 @@ async def test_postgres_incremental_vector_update(postgres_engine_factory, tmp_p
)
# Create and index initial entity
entity = await search_service.entity_repository.create(
{
"title": "Update Test Note",
"note_type": "benchmark",
"entity_metadata": {"tags": ["test"]},
"content_type": "text/markdown",
"permalink": "bench/update-test",
"file_path": "bench/update-test.md",
}
)
async with db.scoped_session(search_service.session_maker) as session:
entity = await search_service.entity_repository.create(
session,
{
"title": "Update Test Note",
"note_type": "benchmark",
"entity_metadata": {"tags": ["test"]},
"content_type": "text/markdown",
"permalink": "bench/update-test",
"file_path": "bench/update-test.md",
},
)
initial_content = build_benchmark_content("sync", TOPIC_TERMS["sync"], 0)
await search_service.index_entity_data(entity, content=initial_content)
await search_service.sync_entity_vectors(entity.id)
+21 -12
View File
@@ -2,6 +2,7 @@
import pytest
from basic_memory import db
from basic_memory.markdown import EntityParser, MarkdownProcessor
from basic_memory.repository import (
EntityRepository,
@@ -31,9 +32,9 @@ async def test_disable_permalinks_create_entity(tmp_path, engine_factory, app_co
app_config.ensure_frontmatter_on_sync = False
# Setup repositories
entity_repository = EntityRepository(session_maker, project_id=test_project.id)
observation_repository = ObservationRepository(session_maker, project_id=test_project.id)
relation_repository = RelationRepository(session_maker, project_id=test_project.id)
entity_repository = EntityRepository(project_id=test_project.id)
observation_repository = ObservationRepository(project_id=test_project.id)
relation_repository = RelationRepository(project_id=test_project.id)
# Use database-specific search repository
if app_config.database_backend == DatabaseBackend.POSTGRES:
@@ -45,9 +46,11 @@ async def test_disable_permalinks_create_entity(tmp_path, engine_factory, app_co
entity_parser = EntityParser(tmp_path)
markdown_processor = MarkdownProcessor(entity_parser)
file_service = FileService(tmp_path, markdown_processor)
search_service = SearchService(search_repository, entity_repository, file_service)
search_service = SearchService(
search_repository, entity_repository, file_service, session_maker
)
await search_service.init_search_index()
link_resolver = LinkResolver(entity_repository, search_service)
link_resolver = LinkResolver(entity_repository, search_service, session_maker)
entity_service = EntityService(
entity_parser=entity_parser,
@@ -56,6 +59,7 @@ async def test_disable_permalinks_create_entity(tmp_path, engine_factory, app_co
relation_repository=relation_repository,
file_service=file_service,
link_resolver=link_resolver,
session_maker=session_maker,
app_config=app_config,
)
@@ -96,9 +100,9 @@ async def test_disable_permalinks_sync_workflow(tmp_path, engine_factory, app_co
test_file.write_text("# Test Note\nThis is test content.")
# Setup repositories
entity_repository = EntityRepository(session_maker, project_id=test_project.id)
observation_repository = ObservationRepository(session_maker, project_id=test_project.id)
relation_repository = RelationRepository(session_maker, project_id=test_project.id)
entity_repository = EntityRepository(project_id=test_project.id)
observation_repository = ObservationRepository(project_id=test_project.id)
relation_repository = RelationRepository(project_id=test_project.id)
# Use database-specific search repository
if app_config.database_backend == DatabaseBackend.POSTGRES:
@@ -106,15 +110,17 @@ async def test_disable_permalinks_sync_workflow(tmp_path, engine_factory, app_co
else:
search_repository = SQLiteSearchRepository(session_maker, project_id=test_project.id)
project_repository = ProjectRepository(session_maker)
project_repository = ProjectRepository()
# Setup services
entity_parser = EntityParser(tmp_path)
markdown_processor = MarkdownProcessor(entity_parser)
file_service = FileService(tmp_path, markdown_processor)
search_service = SearchService(search_repository, entity_repository, file_service)
search_service = SearchService(
search_repository, entity_repository, file_service, session_maker
)
await search_service.init_search_index()
link_resolver = LinkResolver(entity_repository, search_service)
link_resolver = LinkResolver(entity_repository, search_service, session_maker)
entity_service = EntityService(
entity_parser=entity_parser,
@@ -123,6 +129,7 @@ async def test_disable_permalinks_sync_workflow(tmp_path, engine_factory, app_co
relation_repository=relation_repository,
file_service=file_service,
link_resolver=link_resolver,
session_maker=session_maker,
app_config=app_config,
)
@@ -135,6 +142,7 @@ async def test_disable_permalinks_sync_workflow(tmp_path, engine_factory, app_co
relation_repository=relation_repository,
search_service=search_service,
file_service=file_service,
session_maker=session_maker,
)
# Run sync
@@ -151,7 +159,8 @@ async def test_disable_permalinks_sync_workflow(tmp_path, engine_factory, app_co
assert "# Test Note" in content
# Verify entity in database has no permalink
entities = await entity_repository.find_all()
async with db.scoped_session(session_maker) as session:
entities = await entity_repository.find_all(session)
assert len(entities) == 1
assert entities[0].permalink is None
# Title is extracted from filename when no frontmatter, or from frontmatter when present
+15 -13
View File
@@ -93,17 +93,19 @@ async def test_embedding_status_reads_real_vec0_table(engine_factory, test_proje
# --- Seed a real entity + search_index row so counts are non-zero ---
# Use the repository so model-level defaults (external_id) are applied.
entity_repo = EntityRepository(session_maker, project_id=project_id)
entity = await entity_repo.create(
{
"title": "Vec Note",
"note_type": "note",
"content_type": "text/markdown",
"project_id": project_id,
"permalink": "vec-note",
"file_path": "vec-note.md",
}
)
entity_repo = EntityRepository(project_id=project_id)
async with db.scoped_session(session_maker) as session:
entity = await entity_repo.create(
session,
{
"title": "Vec Note",
"note_type": "note",
"content_type": "text/markdown",
"project_id": project_id,
"permalink": "vec-note",
"file_path": "vec-note.md",
},
)
entity_id = entity.id
async with db.scoped_session(session_maker) as session:
@@ -151,8 +153,8 @@ async def test_embedding_status_reads_real_vec0_table(engine_factory, test_proje
await _engine.dispose()
# --- Query status through a fresh ProjectRepository (no extension preloaded) ---
project_repository = ProjectRepository(session_maker)
project_service = ProjectService(project_repository)
project_repository = ProjectRepository()
project_service = ProjectService(project_repository, session_maker)
# Test fixtures run with semantic search disabled; the status call reads the global
# ConfigManager, so patch it to report semantic enabled for this regression path.
+12 -10
View File
@@ -248,16 +248,18 @@ async def _seed_benchmark_notes(search_service, note_count: int):
topic = topic_names[note_index % len(topic_names)]
terms = TOPIC_TERMS[topic]
permalink = f"bench/{topic}-{note_index:05d}"
entity = await search_service.entity_repository.create(
{
"title": f"{topic.title()} Benchmark Note {note_index}",
"note_type": "benchmark",
"entity_metadata": {"tags": ["benchmark", topic], "status": "active"},
"content_type": "text/markdown",
"permalink": permalink,
"file_path": f"{permalink}.md",
}
)
async with db.scoped_session(search_service.session_maker) as session:
entity = await search_service.entity_repository.create(
session,
{
"title": f"{topic.title()} Benchmark Note {note_index}",
"note_type": "benchmark",
"entity_metadata": {"tags": ["benchmark", topic], "status": "active"},
"content_type": "text/markdown",
"permalink": permalink,
"file_path": f"{permalink}.md",
},
)
content = _build_benchmark_content(topic, terms, note_index)
await search_service.index_entity_data(entity, content=content)
if isinstance(search_service.repository, SQLiteSearchRepository):
+18 -11
View File
@@ -7,6 +7,7 @@ from textwrap import dedent
import pytest
from basic_memory import db
from basic_memory.markdown import EntityParser, MarkdownProcessor
from basic_memory.repository import (
EntityRepository,
@@ -40,18 +41,20 @@ async def _build_sync_service(
) -> SyncService:
_, session_maker = engine_factory
entity_repository = EntityRepository(session_maker, project_id=test_project.id)
observation_repository = ObservationRepository(session_maker, project_id=test_project.id)
relation_repository = RelationRepository(session_maker, project_id=test_project.id)
project_repository = ProjectRepository(session_maker)
entity_repository = EntityRepository(project_id=test_project.id)
observation_repository = ObservationRepository(project_id=test_project.id)
relation_repository = RelationRepository(project_id=test_project.id)
project_repository = ProjectRepository()
search_repository = create_search_repository(session_maker, project_id=test_project.id)
entity_parser = EntityParser(project_root)
markdown_processor = MarkdownProcessor(entity_parser)
file_service = FileService(project_root, markdown_processor)
search_service = SearchService(search_repository, entity_repository, file_service)
search_service = SearchService(
search_repository, entity_repository, file_service, session_maker
)
await search_service.init_search_index()
link_resolver = LinkResolver(entity_repository, search_service)
link_resolver = LinkResolver(entity_repository, search_service, session_maker)
entity_service = EntityService(
entity_parser=entity_parser,
@@ -60,6 +63,7 @@ async def _build_sync_service(
relation_repository=relation_repository,
file_service=file_service,
link_resolver=link_resolver,
session_maker=session_maker,
app_config=app_config,
)
@@ -72,6 +76,7 @@ async def _build_sync_service(
project_repository=project_repository,
search_service=search_service,
file_service=file_service,
session_maker=session_maker,
)
@@ -134,9 +139,10 @@ async def test_sync_batching_handles_large_single_file_batches_and_resolves_forw
force_full=True,
)
alpha = await sync_service.entity_repository.get_by_file_path("notes/alpha.md")
large = await sync_service.entity_repository.get_by_file_path("notes/large.md")
target = await sync_service.entity_repository.get_by_file_path("notes/target.md")
async with db.scoped_session(sync_service.session_maker) as session:
alpha = await sync_service.entity_repository.get_by_file_path(session, "notes/alpha.md")
large = await sync_service.entity_repository.get_by_file_path(session, "notes/large.md")
target = await sync_service.entity_repository.get_by_file_path(session, "notes/target.md")
assert report.total == 3
assert alpha is not None
@@ -193,8 +199,9 @@ async def test_sync_batching_circuit_breaker_skips_unchanged_broken_markdown_aft
force_full=True,
)
good = await sync_service.entity_repository.get_by_file_path("notes/good.md")
broken = await sync_service.entity_repository.get_by_file_path("notes/broken.md")
async with db.scoped_session(sync_service.session_maker) as session:
good = await sync_service.entity_repository.get_by_file_path(session, "notes/good.md")
broken = await sync_service.entity_repository.get_by_file_path(session, "notes/broken.md")
assert [skipped.path for skipped in report.skipped_files] == ["notes/broken.md"]
assert good is not None
+59 -40
View File
@@ -8,6 +8,7 @@ import uuid
import pytest
from httpx import AsyncClient
from basic_memory import db
from basic_memory.api.v2.routers.knowledge_router import _canonical_file_path
from basic_memory.ignore_utils import get_bmignore_path
from basic_memory.models import Entity as EntityModel, Project
@@ -20,6 +21,16 @@ from basic_memory.schemas.v2 import EntityResponseV2, EntityResolveResponse
from basic_memory.services.search_service import SearchService
async def _find_all_entities(entity_repository, session_maker):
async with db.scoped_session(session_maker) as session:
return await entity_repository.find_all(session)
async def _get_entity_by_external_id(entity_repository, session_maker, external_id: str):
async with db.scoped_session(session_maker) as session:
return await entity_repository.get_by_external_id(session, external_id)
@pytest.mark.asyncio
async def test_resolve_identifier_by_permalink(
client: AsyncClient, test_graph, v2_project_url, test_project: Project, entity_repository
@@ -62,30 +73,33 @@ async def test_resolve_identifier_returns_target_project_external_id_for_cross_p
v2_project_url,
):
"""Cross-project resolves should expose the owning project external ID."""
project_repository = ProjectRepository(session_maker)
other_project = await project_repository.create(
{
"name": "other-project",
"description": "Secondary project",
"path": str(tmp_path / "other-project"),
"is_active": True,
"is_default": False,
}
)
project_repository = ProjectRepository()
now = datetime.now(timezone.utc)
other_entity_repository = EntityRepository(session_maker, project_id=other_project.id)
target = await other_entity_repository.add(
EntityModel(
title="Cross Project Note",
note_type="note",
content_type="text/markdown",
file_path="docs/Cross Project Note.md",
permalink=f"{other_project.permalink}/docs/cross-project-note",
created_at=now,
updated_at=now,
project_id=other_project.id,
async with db.scoped_session(session_maker) as session:
other_project = await project_repository.create(
session,
{
"name": "other-project",
"description": "Secondary project",
"path": str(tmp_path / "other-project"),
"is_active": True,
"is_default": False,
},
)
other_entity_repository = EntityRepository(project_id=other_project.id)
target = await other_entity_repository.add(
session,
EntityModel(
title="Cross Project Note",
note_type="note",
content_type="text/markdown",
file_path="docs/Cross Project Note.md",
permalink=f"{other_project.permalink}/docs/cross-project-note",
created_at=now,
updated_at=now,
project_id=other_project.id,
),
)
)
response = await client.post(
f"{v2_project_url}/knowledge/resolve",
@@ -196,6 +210,7 @@ async def test_get_entity_by_id_allows_long_relation_type(
client: AsyncClient,
v2_project_url,
relation_repository,
session_maker,
):
"""GET entity should not fail when stored relation_type exceeds 200 characters."""
source_response = await client.post(
@@ -226,14 +241,16 @@ async def test_get_entity_by_id_allows_long_relation_type(
"that is much longer than 200 characters but should still serialize cleanly."
)
await relation_repository.create(
{
"from_id": source_entity.id,
"to_id": target_entity.id,
"to_name": target_entity.title,
"relation_type": long_relation_type,
}
)
async with db.scoped_session(session_maker) as session:
await relation_repository.create(
session,
{
"from_id": source_entity.id,
"to_id": target_entity.id,
"to_name": target_entity.title,
"relation_type": long_relation_type,
},
)
response = await client.get(f"{v2_project_url}/knowledge/entities/{source_entity.external_id}")
@@ -430,7 +447,7 @@ async def test_update_entity_by_id(
@pytest.mark.asyncio
async def test_update_entity_by_id_does_not_duplicate(
client: AsyncClient, v2_project_url, entity_repository
client: AsyncClient, v2_project_url, entity_repository, session_maker
):
"""PUT updates the existing external_id without creating duplicates."""
create_data = {
@@ -453,14 +470,14 @@ async def test_update_entity_by_id_does_not_duplicate(
)
assert response.status_code == 200
entities = await entity_repository.find_all()
entities = await _find_all_entities(entity_repository, session_maker)
assert len(entities) == 1
assert entities[0].external_id == created_entity.external_id
@pytest.mark.asyncio
async def test_put_entity_with_fast_param_returns_fully_indexed_row(
client: AsyncClient, v2_project_url, entity_repository
client: AsyncClient, v2_project_url, entity_repository, session_maker
):
"""PUT ignores the legacy fast param and still returns a fully indexed row."""
external_id = str(uuid.uuid4())
@@ -488,7 +505,7 @@ async def test_put_entity_with_fast_param_returns_fully_indexed_row(
assert len(created_entity.observations) == 1
assert len(created_entity.relations) == 1
db_entity = await entity_repository.get_by_external_id(external_id)
db_entity = await _get_entity_by_external_id(entity_repository, session_maker, external_id)
assert db_entity is not None
@@ -1127,7 +1144,7 @@ async def test_sync_file_rejects_path_traversal(client: AsyncClient, v2_project_
@pytest.mark.asyncio
async def test_sync_file_rejects_symlink_escape(
client: AsyncClient, v2_project_url, test_project: Project, entity_repository
client: AsyncClient, v2_project_url, test_project: Project, entity_repository, session_maker
):
"""sync-file rejects paths whose canonical target escapes the project via symlink.
@@ -1163,7 +1180,7 @@ async def test_sync_file_rejects_symlink_escape(
assert response.status_code in (400, 404)
# Nothing outside the project root was indexed
assert await entity_repository.find_all() == []
assert await _find_all_entities(entity_repository, session_maker) == []
@pytest.mark.asyncio
@@ -1172,6 +1189,7 @@ async def test_sync_file_symlink_escape_never_scans_outside_directory(
v2_project_url,
test_project: Project,
entity_repository,
session_maker,
monkeypatch: pytest.MonkeyPatch,
):
"""Canonicalization never scans directories outside the project root.
@@ -1211,7 +1229,7 @@ async def test_sync_file_symlink_escape_never_scans_outside_directory(
assert response.status_code in (400, 404)
assert outside_dir not in scanned
assert await entity_repository.find_all() == []
assert await _find_all_entities(entity_repository, session_maker) == []
@pytest.mark.asyncio
@@ -1220,6 +1238,7 @@ async def test_sync_file_symlink_escape_never_probes_outside_target(
v2_project_url,
test_project: Project,
entity_repository,
session_maker,
monkeypatch: pytest.MonkeyPatch,
):
"""The containment check runs before any filesystem probe that follows symlinks.
@@ -1260,7 +1279,7 @@ async def test_sync_file_symlink_escape_never_probes_outside_target(
assert "project boundaries" in response.json()["detail"]
assert outside_target not in probed
assert await entity_repository.find_all() == []
assert await _find_all_entities(entity_repository, session_maker) == []
def test_canonical_file_path_stops_at_project_boundary(tmp_path: Path):
@@ -1315,7 +1334,7 @@ async def test_sync_file_symlink_inside_project_still_indexes(
@pytest.mark.asyncio
async def test_sync_file_wrong_cased_path_does_not_create_duplicate(
client: AsyncClient, v2_project_url, test_project: Project, entity_repository
client: AsyncClient, v2_project_url, test_project: Project, entity_repository, session_maker
):
"""A wrong-cased path resolves to the canonical on-disk file without duplicating it.
@@ -1345,7 +1364,7 @@ async def test_sync_file_wrong_cased_path_does_not_create_duplicate(
assert synced.file_path == "notes/disk-note.md"
assert synced.external_id == canonical.external_id
entities = await entity_repository.find_all()
entities = await _find_all_entities(entity_repository, session_maker)
assert [entity.file_path for entity in entities] == ["notes/disk-note.md"]
@@ -3,7 +3,7 @@
from __future__ import annotations
import importlib
from contextlib import contextmanager
from contextlib import asynccontextmanager, contextmanager
from datetime import datetime, timezone
from types import SimpleNamespace
from typing import Any, cast
@@ -53,6 +53,11 @@ def _assert_only_root_span(spans: list[tuple[str, dict]], expected_name: str) ->
assert [name for name, _ in spans] == [expected_name]
@asynccontextmanager
async def _fake_scoped_session(session_maker):
yield object()
@pytest.mark.asyncio
async def test_create_entity_emits_only_root_span(monkeypatch) -> None:
spans, fake_span = _capture_spans()
@@ -103,6 +108,7 @@ async def test_create_entity_emits_only_root_span(monkeypatch) -> None:
async def test_update_entity_emits_only_root_span(monkeypatch) -> None:
spans, fake_span = _capture_spans()
monkeypatch.setattr(logfire, "span", fake_span)
monkeypatch.setattr(knowledge_router_module.db, "scoped_session", _fake_scoped_session)
entity = _fake_entity()
response_content = "---\ntitle: Telemetry Entity\ntype: note\npermalink: notes/test\n---\n\nupdated telemetry content"
@@ -121,7 +127,7 @@ async def test_update_entity_emits_only_root_span(monkeypatch) -> None:
return None
class FakeEntityRepository:
async def get_by_external_id(self, external_id):
async def get_by_external_id(self, session, external_id):
return entity
class FakeTaskScheduler:
@@ -142,6 +148,7 @@ async def test_update_entity_emits_only_root_span(monkeypatch) -> None:
entity_service=cast(Any, FakeEntityService()),
search_service=cast(Any, FakeSearchService()),
entity_repository=cast(Any, FakeEntityRepository()),
session_maker=cast(Any, object()),
task_scheduler=FakeTaskScheduler(),
app_config=cast(Any, SimpleNamespace(semantic_search_enabled=False)),
entity_id=entity.external_id,
@@ -155,6 +162,7 @@ async def test_update_entity_emits_only_root_span(monkeypatch) -> None:
async def test_edit_entity_emits_only_root_span(monkeypatch) -> None:
spans, fake_span = _capture_spans()
monkeypatch.setattr(logfire, "span", fake_span)
monkeypatch.setattr(knowledge_router_module.db, "scoped_session", _fake_scoped_session)
entity = _fake_entity()
response_content = "---\ntitle: Telemetry Entity\ntype: note\npermalink: notes/test\n---\n\nedited telemetry content"
@@ -173,7 +181,7 @@ async def test_edit_entity_emits_only_root_span(monkeypatch) -> None:
return None
class FakeEntityRepository:
async def get_by_external_id(self, external_id):
async def get_by_external_id(self, session, external_id):
return entity
class FakeTaskScheduler:
@@ -186,6 +194,7 @@ async def test_edit_entity_emits_only_root_span(monkeypatch) -> None:
entity_service=cast(Any, FakeEntityService()),
search_service=cast(Any, FakeSearchService()),
entity_repository=cast(Any, FakeEntityRepository()),
session_maker=cast(Any, object()),
task_scheduler=FakeTaskScheduler(),
app_config=cast(Any, SimpleNamespace(semantic_search_enabled=False)),
entity_id=entity.external_id,
+13 -6
View File
@@ -8,9 +8,10 @@ from __future__ import annotations
from datetime import datetime, timezone
from types import SimpleNamespace
from typing import Any
from typing import Any, cast
import pytest
from sqlalchemy.ext.asyncio import AsyncSession
from basic_memory.api.v2.utils import to_graph_context
from basic_memory.schemas.memory import EntitySummary, ObservationSummary, RelationSummary
@@ -44,6 +45,10 @@ def _make_row(*, type: str, id: int, root_id: int, **kwargs: Any) -> ContextResu
return ContextResultRow(type=type, id=id, **defaults)
def _fake_session() -> AsyncSession:
return cast(AsyncSession, object())
class SpyEntityRepository:
"""Tracks batched ID lookups and returns entities from a preset map."""
@@ -56,7 +61,7 @@ class SpyEntityRepository:
return [self.entities_by_id[i] for i in ids if i in self.entities_by_id]
async def find_by_ids_for_hydration(
self, ids: list[int], *, include_cross_project: bool = False
self, session: AsyncSession, ids: list[int], *, include_cross_project: bool = False
):
self.calls.append((ids, include_cross_project))
return [self.entities_by_id[i] for i in ids if i in self.entities_by_id]
@@ -73,7 +78,7 @@ class LightweightOnlyEntityRepository:
raise AssertionError("graph hydration must use the lightweight hydration lookup")
async def find_by_ids_for_hydration(
self, ids: list[int], *, include_cross_project: bool = False
self, session: AsyncSession, ids: list[int], *, include_cross_project: bool = False
):
self.hydration_calls.append((ids, include_cross_project))
return [self.entities_by_id[i] for i in ids if i in self.entities_by_id]
@@ -178,7 +183,9 @@ async def test_to_graph_context_batches_entity_hydration_for_recent_activity():
),
)
graph = await to_graph_context(context, entity_repository=repo, page=1, page_size=10)
graph = await to_graph_context(
context, entity_repository=repo, session=_fake_session(), page=1, page_size=10
)
assert len(repo.calls) == 1, f"Expected 1 entity lookup, got {len(repo.calls)}"
assert set(repo.calls[0][0]) == {1, 2, 3}
@@ -218,7 +225,7 @@ async def test_to_graph_context_empty_results_skip_entity_lookup():
repo = SpyEntityRepository({})
context = ServiceContextResult(results=[], metadata=ContextMetadata(depth=1))
graph = await to_graph_context(context, entity_repository=repo)
graph = await to_graph_context(context, entity_repository=repo, session=_fake_session())
assert repo.calls == []
assert list(graph.results) == []
@@ -274,7 +281,7 @@ async def test_to_graph_context_uses_lightweight_hydration_lookup():
),
)
graph = await to_graph_context(context, entity_repository=repo)
graph = await to_graph_context(context, entity_repository=repo, session=_fake_session())
assert len(repo.hydration_calls) == 1
assert set(repo.hydration_calls[0][0]) == {1, 2}
+3 -1
View File
@@ -4,6 +4,7 @@ import pytest
from httpx import AsyncClient
from pathlib import Path
from basic_memory import db
from basic_memory.models import Project
@@ -18,7 +19,8 @@ async def create_test_entity(
await file_service.write_file(file_path, test_content)
# Create entity
entity = await entity_repository.create(entity_data)
async with db.scoped_session(search_service.session_maker) as session:
entity = await entity_repository.create(session, entity_data)
# Index for search
await search_service.index_entity(entity)
+68 -17
View File
@@ -6,6 +6,7 @@ from pathlib import Path
import pytest
from httpx import AsyncClient
from basic_memory import db
from basic_memory.config import ProjectEntry
from basic_memory.models import Project
from basic_memory.schemas.project_info import ProjectItem, ProjectStatusResponse
@@ -17,6 +18,31 @@ def _project_item(project: ProjectItem | None) -> ProjectItem:
return project
async def _find_projects(project_repository, session_maker):
async with db.scoped_session(session_maker) as session:
return await project_repository.find_all(session)
async def _delete_project(project_repository, session_maker, project_id: int) -> bool:
async with db.scoped_session(session_maker) as session:
return await project_repository.delete(session, project_id)
async def _get_project_by_name(project_repository, session_maker, name: str):
async with db.scoped_session(session_maker) as session:
return await project_repository.get_by_name(session, name)
async def _get_default_project(project_repository, session_maker):
async with db.scoped_session(session_maker) as session:
return await project_repository.get_default_project(session)
async def _update_project(project_repository, session_maker, project_id: int, data: dict):
async with db.scoped_session(session_maker) as session:
return await project_repository.update(session, project_id, data)
@pytest.mark.asyncio
async def test_list_projects(client: AsyncClient, test_project: Project, v2_projects_url):
"""Test listing projects returns default_project from the database."""
@@ -63,6 +89,7 @@ async def test_add_project_response_reflects_promoted_default(
config_manager,
config_home,
project_repository,
session_maker,
):
"""Regression #974/#985: POST response should echo persisted default promotion."""
main_home = config_home / "basic-memory"
@@ -78,8 +105,8 @@ async def test_add_project_response_reflects_promoted_default(
)
config_manager.save_config(fresh_config)
for project in await project_repository.find_all():
await project_repository.delete(project.id)
for project in await _find_projects(project_repository, session_maker):
await _delete_project(project_repository, session_maker, project.id)
response = await client.post(
f"{v2_projects_url}/",
@@ -153,14 +180,21 @@ async def test_update_project_not_found(client: AsyncClient, v2_projects_url, tm
@pytest.mark.asyncio
async def test_set_default_project_by_id(
client: AsyncClient, test_project: Project, v2_projects_url, project_repository, project_service
client: AsyncClient,
test_project: Project,
v2_projects_url,
project_repository,
project_service,
session_maker,
):
"""Test setting a project as default by external_id."""
# Create a second project to test setting default
await project_service.add_project("second-project", "/tmp/second-project")
# Get the created project from the repository to get its external_id
created_project = await project_repository.get_by_name("second-project")
created_project = await _get_project_by_name(
project_repository, session_maker, "second-project"
)
assert created_project is not None
# Set the second project as default
@@ -180,7 +214,7 @@ async def test_set_default_project_by_id(
@pytest.mark.asyncio
async def test_set_default_project_when_none_is_set(
client: AsyncClient, test_project: Project, v2_projects_url, project_repository
client: AsyncClient, test_project: Project, v2_projects_url, project_repository, session_maker
):
"""Regression for #975: setting a default must succeed when none is set.
@@ -188,8 +222,8 @@ async def test_set_default_project_when_none_is_set(
the command reached for when no default exists, so the endpoint must not 404.
"""
# Clear any existing default so no row has is_default set.
await project_repository.update(test_project.id, {"is_default": None})
assert await project_repository.get_default_project() is None
await _update_project(project_repository, session_maker, test_project.id, {"is_default": None})
assert await _get_default_project(project_repository, session_maker) is None
response = await client.put(f"{v2_projects_url}/{test_project.external_id}/default")
@@ -204,7 +238,7 @@ async def test_set_default_project_when_none_is_set(
assert new_project.is_default is True
# A follow-up read-back must now return the newly set default.
default_project = await project_repository.get_default_project()
default_project = await _get_default_project(project_repository, session_maker)
assert default_project is not None
assert default_project.external_id == test_project.external_id
@@ -220,14 +254,19 @@ async def test_set_default_project_not_found(client: AsyncClient, v2_projects_ur
@pytest.mark.asyncio
async def test_delete_project_by_id(
client: AsyncClient, test_project: Project, v2_projects_url, project_repository, project_service
client: AsyncClient,
test_project: Project,
v2_projects_url,
project_repository,
project_service,
session_maker,
):
"""Test deleting a project by external_id."""
# Create a second project since we can't delete the default
await project_service.add_project("to-delete", "/tmp/to-delete")
# Get the created project from the repository to get its external_id
created_project = await project_repository.get_by_name("to-delete")
created_project = await _get_project_by_name(project_repository, session_maker, "to-delete")
assert created_project is not None
# Delete it
@@ -247,7 +286,12 @@ async def test_delete_project_by_id(
@pytest.mark.asyncio
async def test_delete_project_with_delete_notes_param(
client: AsyncClient, test_project: Project, v2_projects_url, project_repository, project_service
client: AsyncClient,
test_project: Project,
v2_projects_url,
project_repository,
project_service,
session_maker,
):
"""Test deleting a project with delete_notes parameter."""
# Create a project in a temp directory
@@ -262,7 +306,9 @@ async def test_delete_project_with_delete_notes_param(
await project_service.add_project("delete-with-notes", str(project_path))
# Get the created project from the repository to get its external_id
created_project = await project_repository.get_by_name("delete-with-notes")
created_project = await _get_project_by_name(
project_repository, session_maker, "delete-with-notes"
)
assert created_project is not None
# Delete with delete_notes=true
@@ -335,14 +381,19 @@ async def test_project_id_stability_after_rename(
@pytest.mark.asyncio
async def test_update_project_active_status(
client: AsyncClient, test_project: Project, v2_projects_url, project_repository, project_service
client: AsyncClient,
test_project: Project,
v2_projects_url,
project_repository,
project_service,
session_maker,
):
"""Test updating a project's active status by external_id."""
# Create a non-default project
await project_service.add_project("test-active", "/tmp/test-active")
# Get the created project from the repository to get its external_id
created_project = await project_repository.get_by_name("test-active")
created_project = await _get_project_by_name(project_repository, session_maker, "test-active")
assert created_project is not None
# Update active status
@@ -447,7 +498,7 @@ async def test_resolve_project_not_found(client: AsyncClient, v2_projects_url):
@pytest.mark.asyncio
async def test_resolve_project_not_found_fresh_install_names_setup_command(
client: AsyncClient, v2_projects_url, project_repository
client: AsyncClient, v2_projects_url, project_repository, session_maker
):
"""#974 follow-up: a fresh install fails its first read with a bare not-found.
@@ -456,8 +507,8 @@ async def test_resolve_project_not_found_fresh_install_names_setup_command(
the configured default 404s. With an empty projects table the error must point
at first-run setup instead of reading like a broken install.
"""
for project in await project_repository.find_all():
await project_repository.delete(project.id)
for project in await _find_projects(project_repository, session_maker):
await _delete_project(project_repository, session_maker, project.id)
response = await client.post(f"{v2_projects_url}/resolve", json={"identifier": "main"})
+3 -1
View File
@@ -6,6 +6,7 @@ import pytest
from httpx import AsyncClient
from pathlib import Path
from basic_memory import db
from basic_memory.deps.services import get_search_service_v2_external
from basic_memory.models import Project
from basic_memory.repository.search_index_row import SearchIndexRow
@@ -26,7 +27,8 @@ async def create_test_entity(
await file_service.write_file(file_path, test_content)
# Create entity
entity = await entity_repository.create(entity_data)
async with db.scoped_session(search_service.session_maker) as session:
entity = await entity_repository.create(session, entity_data)
# Index for search
await search_service.index_entity(entity)
+12 -13
View File
@@ -163,18 +163,17 @@ def test_reindex_embeddings_only_preserves_incremental_default(monkeypatch):
@pytest.mark.asyncio
async def test_reindex_project_full_passes_force_full_to_sync_and_reports_mode(monkeypatch):
async def test_reindex_project_full_passes_force_full_to_sync_and_reports_mode(
monkeypatch,
session_maker,
):
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):
async def get_active_projects(self, session):
return [project]
class SilentProgress:
@@ -236,25 +235,25 @@ async def test_reindex_project_full_passes_force_full_to_sync_and_reports_mode(m
@pytest.mark.asyncio
async def test_reindex_embeddings_only_full_passes_force_full_to_vector_reindex(monkeypatch):
async def test_reindex_embeddings_only_full_passes_force_full_to_vector_reindex(
monkeypatch,
session_maker,
):
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):
async def get_active_projects(self, session):
return [project]
class StubSearchService:
def __init__(self, search_repository, entity_repository, file_service):
def __init__(self, search_repository, entity_repository, file_service, *, session_maker):
self.search_repository = search_repository
self.entity_repository = entity_repository
self.file_service = file_service
self.session_maker = session_maker
async def reindex_vectors(self, *, progress_callback=None, force_full: bool = False):
vector_reindex_calls.append(
+53 -16
View File
@@ -438,7 +438,7 @@ async def entity_repository(
session_maker: async_sessionmaker[AsyncSession], test_project: Project
) -> EntityRepository:
"""Create an EntityRepository instance with project context."""
return EntityRepository(session_maker, project_id=test_project.id)
return EntityRepository(project_id=test_project.id)
@pytest_asyncio.fixture(scope="function")
@@ -446,7 +446,7 @@ async def observation_repository(
session_maker: async_sessionmaker[AsyncSession], test_project: Project
) -> ObservationRepository:
"""Create an ObservationRepository instance with project context."""
return ObservationRepository(session_maker, project_id=test_project.id)
return ObservationRepository(project_id=test_project.id)
@pytest_asyncio.fixture(scope="function")
@@ -454,7 +454,7 @@ async def relation_repository(
session_maker: async_sessionmaker[AsyncSession], test_project: Project
) -> RelationRepository:
"""Create a RelationRepository instance with project context."""
return RelationRepository(session_maker, project_id=test_project.id)
return RelationRepository(project_id=test_project.id)
@pytest_asyncio.fixture(scope="function")
@@ -462,7 +462,7 @@ async def project_repository(
session_maker: async_sessionmaker[AsyncSession],
) -> ProjectRepository:
"""Create a ProjectRepository instance."""
return ProjectRepository(session_maker)
return ProjectRepository()
@pytest_asyncio.fixture(scope="function")
@@ -476,8 +476,9 @@ async def test_project(config_home, engine_factory) -> Project:
"is_default": True, # Explicitly set as the default project (for cli operations)
}
engine, session_maker = engine_factory
project_repository = ProjectRepository(session_maker)
project = await project_repository.create(project_data)
project_repository = ProjectRepository()
async with db.scoped_session(session_maker) as session:
project = await project_repository.create(session, project_data)
return project
@@ -493,6 +494,7 @@ async def entity_service(
file_service: FileService,
link_resolver: LinkResolver,
app_config: BasicMemoryConfig,
session_maker: async_sessionmaker[AsyncSession],
) -> EntityService:
"""Create EntityService."""
return EntityService(
@@ -503,6 +505,7 @@ async def entity_service(
file_service=file_service,
link_resolver=link_resolver,
app_config=app_config,
session_maker=session_maker,
)
@@ -521,9 +524,13 @@ def markdown_processor(entity_parser: EntityParser) -> MarkdownProcessor:
@pytest.fixture
def link_resolver(entity_repository: EntityRepository, search_service: SearchService):
def link_resolver(
entity_repository: EntityRepository,
search_service: SearchService,
session_maker: async_sessionmaker[AsyncSession],
):
"""Create parser instance."""
return LinkResolver(entity_repository, search_service)
return LinkResolver(entity_repository, search_service, session_maker=session_maker)
@pytest.fixture
@@ -542,6 +549,7 @@ async def sync_service(
relation_repository: RelationRepository,
search_service: SearchService,
file_service: FileService,
session_maker: async_sessionmaker[AsyncSession],
) -> SyncService:
"""Create sync service for testing."""
return SyncService(
@@ -553,14 +561,20 @@ async def sync_service(
entity_parser=entity_parser,
search_service=search_service,
file_service=file_service,
session_maker=session_maker,
)
@pytest_asyncio.fixture
async def directory_service(entity_repository, project_config) -> DirectoryService:
async def directory_service(
entity_repository,
project_config,
session_maker: async_sessionmaker[AsyncSession],
) -> DirectoryService:
"""Create directory service for testing."""
return DirectoryService(
entity_repository=entity_repository,
session_maker=session_maker,
)
@@ -589,15 +603,24 @@ async def search_service(
search_repository,
entity_repository: EntityRepository,
file_service: FileService,
session_maker: async_sessionmaker[AsyncSession],
) -> SearchService:
"""Create and initialize search service"""
service = SearchService(search_repository, entity_repository, file_service)
service = SearchService(
search_repository,
entity_repository,
file_service,
session_maker=session_maker,
)
await service.init_search_index()
return service
@pytest_asyncio.fixture(scope="function")
async def sample_entity(entity_repository: EntityRepository) -> Entity:
async def sample_entity(
entity_repository: EntityRepository,
session_maker: async_sessionmaker[AsyncSession],
) -> Entity:
"""Create a sample entity for testing."""
entity_data = {
"project_id": entity_repository.project_id,
@@ -609,16 +632,22 @@ async def sample_entity(entity_repository: EntityRepository) -> Entity:
"created_at": datetime.now(timezone.utc),
"updated_at": datetime.now(timezone.utc),
}
return await entity_repository.create(entity_data)
async with db.scoped_session(session_maker) as session:
return await entity_repository.create(session, entity_data)
@pytest_asyncio.fixture
async def project_service(
project_repository: ProjectRepository,
file_service: FileService,
session_maker: async_sessionmaker[AsyncSession],
) -> ProjectService:
"""Create ProjectService with repository and file service for directory operations."""
return ProjectService(repository=project_repository, file_service=file_service)
return ProjectService(
repository=project_repository,
file_service=file_service,
session_maker=session_maker,
)
@pytest_asyncio.fixture
@@ -653,6 +682,7 @@ async def test_graph(
search_service,
file_service,
entity_service,
session_maker,
):
"""Create a test knowledge graph with entities, relations and observations."""
@@ -720,8 +750,9 @@ async def test_graph(
)
# get latest
entities = await entity_repository.find_all()
relations = await relation_repository.find_all()
async with db.scoped_session(session_maker) as session:
entities = await entity_repository.find_all(session)
relations = await relation_repository.find_all(session)
# Index everything for search
for entity in entities:
@@ -738,7 +769,12 @@ async def test_graph(
@pytest.fixture
def watch_service(app_config: BasicMemoryConfig, project_repository, sync_service) -> WatchService:
def watch_service(
app_config: BasicMemoryConfig,
project_repository,
sync_service,
session_maker: async_sessionmaker[AsyncSession],
) -> WatchService:
"""Create WatchService with injected sync_service factory.
The sync_service_factory allows tests to use the fixture-provided sync_service
@@ -752,6 +788,7 @@ def watch_service(app_config: BasicMemoryConfig, project_repository, sync_servic
return WatchService(
app_config=app_config,
project_repository=project_repository,
session_maker=session_maker,
sync_service_factory=sync_service_factory,
)
@@ -20,6 +20,7 @@ from pathlib import Path
import pytest
from basic_memory import db
from basic_memory.services.project_service import ProjectService
@@ -52,15 +53,9 @@ async def test_issue_254_foreign_key_constraint_fix(project_service: ProjectServ
from basic_memory.repository.observation_repository import ObservationRepository
from basic_memory.repository.relation_repository import RelationRepository
entity_repo = EntityRepository(
project_service.repository.session_maker, project_id=project.id
)
obs_repo = ObservationRepository(
project_service.repository.session_maker, project_id=project.id
)
rel_repo = RelationRepository(
project_service.repository.session_maker, project_id=project.id
)
entity_repo = EntityRepository(project_id=project.id)
obs_repo = ObservationRepository(project_id=project.id)
rel_repo = RelationRepository(project_id=project.id)
# Create entity
entity_data = {
@@ -74,23 +69,24 @@ async def test_issue_254_foreign_key_constraint_fix(project_service: ProjectServ
"created_at": datetime.now(timezone.utc),
"updated_at": datetime.now(timezone.utc),
}
entity = await entity_repo.create(entity_data)
async with db.scoped_session(project_service.session_maker) as session:
entity = await entity_repo.create(session, entity_data)
# Create observation linked to entity
observation_data = {
"entity_id": entity.id,
"content": "This observation should be cascade deleted",
"category": "test",
}
observation = await obs_repo.create(observation_data)
# Create observation linked to entity
observation_data = {
"entity_id": entity.id,
"content": "This observation should be cascade deleted",
"category": "test",
}
observation = await obs_repo.create(session, observation_data)
# Create relation involving the entity
relation_data = {
"from_id": entity.id,
"to_name": "some-other-entity",
"relation_type": "relates-to",
}
relation = await rel_repo.create(relation_data)
# Create relation involving the entity
relation_data = {
"from_id": entity.id,
"to_name": "some-other-entity",
"relation_type": "relates-to",
}
relation = await rel_repo.create(session, relation_data)
# Step 3: Attempt to remove the project
# This is where issue #254 manifested - should NOT raise "FOREIGN KEY constraint failed"
@@ -112,14 +108,15 @@ async def test_issue_254_foreign_key_constraint_fix(project_service: ProjectServ
assert removed_project is None, "Project should have been removed"
# Step 5: Verify related data was cascade deleted
remaining_entity = await entity_repo.find_by_id(entity.id)
assert remaining_entity is None, "Entity should have been cascade deleted"
async with db.scoped_session(project_service.session_maker) as session:
remaining_entity = await entity_repo.find_by_id(session, entity.id)
assert remaining_entity is None, "Entity should have been cascade deleted"
remaining_observation = await obs_repo.find_by_id(observation.id)
assert remaining_observation is None, "Observation should have been cascade deleted"
remaining_observation = await obs_repo.find_by_id(session, observation.id)
assert remaining_observation is None, "Observation should have been cascade deleted"
remaining_relation = await rel_repo.find_by_id(relation.id)
assert remaining_relation is None, "Relation should have been cascade deleted"
remaining_relation = await rel_repo.find_by_id(session, relation.id)
assert remaining_relation is None, "Relation should have been cascade deleted"
@pytest.mark.asyncio
@@ -141,9 +138,7 @@ async def test_issue_254_reproduction(project_service: ProjectService):
from basic_memory.repository.entity_repository import EntityRepository
entity_repo = EntityRepository(
project_service.repository.session_maker, project_id=project.id
)
entity_repo = EntityRepository(project_id=project.id)
entity_data = {
"title": "Reproduction Entity",
@@ -156,7 +151,8 @@ async def test_issue_254_reproduction(project_service: ProjectService):
"created_at": datetime.now(timezone.utc),
"updated_at": datetime.now(timezone.utc),
}
await entity_repo.create(entity_data)
async with db.scoped_session(project_service.session_maker) as session:
await entity_repo.create(session, entity_data)
# This should eventually work without errors once issue #254 is fixed
# with pytest.raises(Exception) as exc_info:
@@ -21,6 +21,7 @@ import pytest
from sqlalchemy import text
from basic_memory import db
from basic_memory.config import BasicMemoryConfig, DatabaseBackend, ProjectEntry
from basic_memory.models import Base
@@ -91,3 +92,32 @@ async def test_concurrent_session_rollback_does_not_destroy_uncommitted_writes()
"writer's committed INSERT was rolled back by a concurrent session — "
"the in-memory engine is sharing one transaction scope across sessions"
)
@pytest.mark.asyncio
async def test_windows_memory_engine_respects_explicit_rollback(monkeypatch):
"""The Windows SQLite branch must not force driver-level autocommit."""
db_path = Path("unused.db")
app_config = BasicMemoryConfig(
env="test",
projects={"test-project": ProjectEntry(path=".")},
default_project="test-project",
database_backend=DatabaseBackend.SQLITE,
)
monkeypatch.setattr(db.os, "name", "nt")
async with db.engine_session_factory(
db_path=db_path, db_type=db.DatabaseType.MEMORY, config=app_config
) as (engine, session_maker):
async with engine.begin() as conn:
await conn.execute(text("CREATE TABLE rollback_probe (name TEXT NOT NULL)"))
async with session_maker() as session:
transaction = await session.begin()
await session.execute(text("INSERT INTO rollback_probe (name) VALUES ('discarded')"))
await transaction.rollback()
async with session_maker() as session:
count = (await session.execute(text("SELECT count(*) FROM rollback_probe"))).scalar()
assert count == 0
@@ -5,6 +5,7 @@ This test verifies issue #452 - Imported conversations not indexed correctly.
import pytest
from basic_memory import db
from basic_memory.config import ProjectConfig
from basic_memory.importers.claude_conversations_importer import ClaudeConversationsImporter
from basic_memory.markdown import EntityParser
@@ -91,7 +92,8 @@ async def test_imported_conversations_have_correct_permalink_and_title(
await sync_service.sync(base_path, project_config.name)
# Verify entity in database
entities = await entity_repository.find_all()
async with db.scoped_session(sync_service.session_maker) as session:
entities = await entity_repository.find_all(session)
assert len(entities) == 1, f"Expected 1 entity, got {len(entities)}"
entity = entities[0]
+35 -35
View File
@@ -10,6 +10,7 @@ from unittest.mock import AsyncMock
import pytest
from sqlalchemy import text
from basic_memory import db
from basic_memory.file_utils import remove_frontmatter
from basic_memory.indexing import (
BatchIndexer,
@@ -67,6 +68,7 @@ def _make_batch_indexer(
relation_repository=relation_repository,
search_service=search_service,
file_writer=_TestFileWriter(file_service),
session_maker=search_service.session_maker,
)
@@ -150,7 +152,7 @@ async def test_batch_indexer_parses_markdown_with_parallel_path(
@pytest.mark.asyncio
async def test_batch_indexer_creates_entities_with_parallel_path(
async def test_batch_indexer_creates_entities_with_real_db_session(
app_config,
entity_service,
entity_repository,
@@ -199,34 +201,24 @@ async def test_batch_indexer_creates_entities_with_parallel_path(
file_service,
)
original_upsert = entity_service.upsert_entity_from_markdown
in_flight = 0
max_in_flight = 0
result = await batch_indexer.index_files(
files,
max_concurrent=2,
parse_max_concurrent=2,
)
async def spy_upsert(*args, **kwargs):
nonlocal in_flight, max_in_flight
in_flight += 1
max_in_flight = max(max_in_flight, in_flight)
await asyncio.sleep(0.05)
try:
return await original_upsert(*args, **kwargs)
finally:
in_flight -= 1
entity_service.upsert_entity_from_markdown = spy_upsert
try:
result = await batch_indexer.index_files(
files,
max_concurrent=2,
parse_max_concurrent=2,
)
finally:
entity_service.upsert_entity_from_markdown = original_upsert
assert max_in_flight >= 2
assert len(result.indexed) == 2
assert result.errors == []
async with db.scoped_session(search_service.session_maker) as session:
alpha = await entity_repository.get_by_file_path(session, path_one)
beta = await entity_repository.get_by_file_path(session, path_two)
assert alpha is not None
assert alpha.title == "Alpha"
assert beta is not None
assert beta.title == "Beta"
@pytest.mark.asyncio
async def test_batch_indexer_returns_original_markdown_content_when_no_frontmatter_rewrite(
@@ -316,8 +308,9 @@ async def test_batch_indexer_indexes_non_markdown_files(
assert {indexed.path for indexed in result.indexed} == {pdf_path, image_path}
assert all(indexed.markdown_content is None for indexed in result.indexed)
pdf_entity = await entity_repository.get_by_file_path(pdf_path)
image_entity = await entity_repository.get_by_file_path(image_path)
async with db.scoped_session(search_service.session_maker) as session:
pdf_entity = await entity_repository.get_by_file_path(session, pdf_path)
image_entity = await entity_repository.get_by_file_path(session, image_path)
assert pdf_entity is not None
assert pdf_entity.content_type == "application/pdf"
assert image_entity is not None
@@ -383,8 +376,9 @@ async def test_batch_indexer_resolves_relations_and_refreshes_search(
parse_max_concurrent=2,
)
source = await entity_repository.get_by_file_path(source_path)
target = await entity_repository.get_by_file_path(target_path)
async with db.scoped_session(search_service.session_maker) as session:
source = await entity_repository.get_by_file_path(session, source_path)
target = await entity_repository.get_by_file_path(session, target_path)
assert source is not None
assert target is not None
assert len(source.outgoing_relations) == 1
@@ -476,7 +470,8 @@ async def test_batch_indexer_assigns_unique_permalinks_for_batch_local_conflicts
path_two
)
entities = await entity_repository.find_all()
async with db.scoped_session(search_service.session_maker) as session:
entities = await entity_repository.find_all(session)
assert len(entities) == 2
permalinks = [entity.permalink for entity in entities if entity.permalink]
assert len(set(permalinks)) == 2
@@ -533,7 +528,8 @@ async def test_batch_indexer_uses_parsed_markdown_body_for_malformed_frontmatter
assert len(result.indexed) == 1
assert result.indexed[0].markdown_content == persisted_content
entity = await entity_repository.get_by_file_path(path)
async with db.scoped_session(search_service.session_maker) as session:
entity = await entity_repository.get_by_file_path(session, path)
assert entity is not None
@@ -681,7 +677,8 @@ async def test_batch_indexer_index_markdown_file_can_defer_relation_resolution(
)
resolve_link.assert_not_awaited()
source = await entity_repository.get_by_file_path(path)
async with db.scoped_session(search_service.session_maker) as session:
source = await entity_repository.get_by_file_path(session, path)
assert source is not None
assert len(source.outgoing_relations) == 1
assert source.outgoing_relations[0].to_id is None
@@ -753,7 +750,8 @@ async def test_batch_indexer_uses_strict_link_resolution_for_deferred_relations(
)
# The unresolvable relation stayed unresolved.
source = await entity_repository.get_by_file_path(path)
async with db.scoped_session(search_service.session_maker) as session:
source = await entity_repository.get_by_file_path(session, path)
assert source is not None
assert len(source.outgoing_relations) == 1
assert source.outgoing_relations[0].to_id is None
@@ -801,7 +799,8 @@ async def test_batch_indexer_strips_frontmatter_from_search_content_when_body_is
)
persisted_content = await file_service.read_file_content(path)
entity = await entity_repository.get_by_file_path(path)
async with db.scoped_session(search_service.session_maker) as session:
entity = await entity_repository.get_by_file_path(session, path)
assert entity is not None
index_entity_data.assert_awaited_once()
await_args = index_entity_data.await_args
@@ -861,7 +860,8 @@ async def test_batch_indexer_does_not_inject_frontmatter_when_sync_enforcement_i
# Why: this assertion cares about preserving a frontmatterless file, not about newline style.
# Outcome: compare against the exact content stored on disk after sync.
persisted_content = (project_config.home / path).read_bytes().decode("utf-8")
entity = await entity_repository.get_by_file_path(path)
async with db.scoped_session(search_service.session_maker) as session:
entity = await entity_repository.get_by_file_path(session, path)
assert entity is not None
assert entity.permalink == existing_permalink
assert frontmatter_writer.await_count == 0
@@ -18,6 +18,7 @@ import pytest
from pathlib import Path
from textwrap import dedent
from basic_memory import db
from basic_memory.mcp.tools import write_note, read_note
from basic_memory.sync.sync_service import SyncService
from basic_memory.config import ProjectConfig
@@ -28,17 +29,19 @@ from basic_memory.utils import generate_permalink
async def force_full_scan(sync_service: SyncService) -> None:
"""Force next sync to do a full scan by clearing watermark (for testing moves/deletions)."""
if sync_service.entity_repository.project_id is not None:
project = await sync_service.project_repository.find_by_id(
sync_service.entity_repository.project_id
)
if project:
await sync_service.project_repository.update(
project.id,
{
"last_scan_timestamp": None,
"last_file_count": None,
},
async with db.scoped_session(sync_service.session_maker) as session:
project = await sync_service.project_repository.find_by_id(
session, sync_service.entity_repository.project_id
)
if project:
await sync_service.project_repository.update(
session,
project.id,
{
"last_scan_timestamp": None,
"last_file_count": None,
},
)
@pytest.mark.asyncio
+1 -3
View File
@@ -3199,9 +3199,7 @@ async def test_resolver_refreshes_index_on_miss(monkeypatch):
monkeypatch.setattr(project_context, "_ensure_workspace_project_index", fake_index)
# By external_id (the project_id routing path that failed in the field)
resolved = await resolve_workspace_project_identifier(
"22222222-2222-4222-8222-222222222222"
)
resolved = await resolve_workspace_project_identifier("22222222-2222-4222-8222-222222222222")
assert resolved.project.permalink == "manual"
assert calls == [False, True]
+13 -4
View File
@@ -6,6 +6,7 @@ from textwrap import dedent
import pytest
from basic_memory import db
from basic_memory.mcp.tools import write_note, read_note
from basic_memory.mcp.tools.read_note import _parse_opening_frontmatter
from basic_memory.utils import normalize_newlines
@@ -736,6 +737,7 @@ async def test_team_workspace_write_stores_complete_canonical_permalink(
app,
test_project,
entity_repository,
session_maker,
):
from basic_memory.workspace_context import workspace_permalink_context
@@ -751,7 +753,8 @@ async def test_team_workspace_write_stores_complete_canonical_permalink(
assert f"permalink: {expected_permalink}" in write_result
stored = await entity_repository.get_by_permalink(expected_permalink)
async with db.scoped_session(session_maker) as session:
stored = await entity_repository.get_by_permalink(session, expected_permalink)
assert stored is not None
assert stored.permalink == expected_permalink
@@ -772,6 +775,7 @@ async def test_team_workspace_write_stores_complete_permalink_when_project_prefi
test_project,
entity_repository,
config_manager,
session_maker,
):
from basic_memory.workspace_context import workspace_permalink_context
@@ -791,7 +795,8 @@ async def test_team_workspace_write_stores_complete_permalink_when_project_prefi
assert f"permalink: {expected_permalink}" in write_result
stored = await entity_repository.get_by_permalink(expected_permalink)
async with db.scoped_session(session_maker) as session:
stored = await entity_repository.get_by_permalink(session, expected_permalink)
assert stored is not None
assert stored.permalink == expected_permalink
@@ -811,6 +816,7 @@ async def test_personal_workspace_write_stores_complete_canonical_permalink(
app,
test_project,
entity_repository,
session_maker,
):
from basic_memory.workspace_context import workspace_permalink_context
@@ -826,7 +832,8 @@ async def test_personal_workspace_write_stores_complete_canonical_permalink(
assert f"permalink: {expected_permalink}" in write_result
stored = await entity_repository.get_by_permalink(expected_permalink)
async with db.scoped_session(session_maker) as session:
stored = await entity_repository.get_by_permalink(session, expected_permalink)
assert stored is not None
assert stored.permalink == expected_permalink
@@ -864,6 +871,7 @@ async def test_read_note_workspace_qualified_personal_url_finds_legacy_short_per
app,
test_project,
entity_repository,
session_maker,
):
from basic_memory.workspace_context import workspace_permalink_context
@@ -877,7 +885,8 @@ async def test_read_note_workspace_qualified_personal_url_finds_legacy_short_per
content="Legacy personal workspace content",
)
stored = await entity_repository.get_by_permalink(legacy_permalink)
async with db.scoped_session(session_maker) as session:
stored = await entity_repository.get_by_permalink(session, legacy_permalink)
assert stored is not None
assert stored.permalink == legacy_permalink
+37 -21
View File
@@ -5,6 +5,7 @@ from typing import Any
import pytest
from basic_memory import db
from basic_memory import config as config_module
from basic_memory.mcp.tools import write_note, read_note, delete_note
from basic_memory.mcp.tools.write_note import _compose_workspace_project_route
@@ -19,29 +20,38 @@ from basic_memory.utils import normalize_newlines
def test_write_note_workspace_project_route_passthrough_without_workspace():
"""Without workspace, the project string passes through unchanged."""
assert _compose_workspace_project_route(
workspace=None,
project="my-project",
project_id=None,
) == "my-project"
assert (
_compose_workspace_project_route(
workspace=None,
project="my-project",
project_id=None,
)
== "my-project"
)
def test_write_note_workspace_project_route_combines_workspace_and_project():
"""workspace + project are joined as 'workspace/project'."""
assert _compose_workspace_project_route(
workspace="acme",
project="docs",
project_id=None,
) == "acme/docs"
assert (
_compose_workspace_project_route(
workspace="acme",
project="docs",
project_id=None,
)
== "acme/docs"
)
def test_write_note_workspace_project_route_passes_qualified_project_unchanged():
"""A pre-qualified 'workspace/project' string passes through when workspace is None."""
assert _compose_workspace_project_route(
workspace=None,
project="acme/docs",
project_id=None,
) == "acme/docs"
assert (
_compose_workspace_project_route(
workspace=None,
project="acme/docs",
project_id=None,
)
== "acme/docs"
)
@pytest.mark.parametrize(
@@ -439,8 +449,9 @@ async def test_write_note_verbose(app, test_project, engine_factory):
assert f"[Session: Using project '{test_project.name}']" in result
_, session_maker = engine_factory
relation_repository = RelationRepository(session_maker, project_id=test_project.id)
relations = await relation_repository.find_by_type("relates to")
relation_repository = RelationRepository(project_id=test_project.id)
async with db.scoped_session(session_maker) as session:
relations = await relation_repository.find_by_type(session, "relates to")
assert any(relation.to_name == "Knowledge" for relation in relations)
@@ -1441,7 +1452,7 @@ class TestWriteNoteOverwriteGuard:
@pytest.mark.asyncio
async def test_write_note_overwrite_resolves_by_file_path_strictly(
self, app, test_project, entity_repository, monkeypatch
self, app, test_project, entity_repository, session_maker, monkeypatch
):
"""Regression: overwrite=True must resolve the conflicting entity by
file_path with strict=True, not by permalink with fuzzy fallback.
@@ -1479,7 +1490,8 @@ class TestWriteNoteOverwriteGuard:
content="# Overview\n\nVersion A",
)
canonical_permalink = f"{test_project.name}/features/foo/overview"
canonical = await entity_repository.get_by_permalink(canonical_permalink)
async with db.scoped_session(session_maker) as session:
canonical = await entity_repository.get_by_permalink(session, canonical_permalink)
assert canonical is not None
canonical_id = canonical.id
@@ -1498,7 +1510,8 @@ class TestWriteNoteOverwriteGuard:
assert captured.get("strict") is True
# And the canonical row was updated in place — no duplicate -1/-2 row.
canonical_after = await entity_repository.get_by_permalink(canonical_permalink)
async with db.scoped_session(session_maker) as session:
canonical_after = await entity_repository.get_by_permalink(session, canonical_permalink)
assert canonical_after is not None
assert canonical_after.id == canonical_id
@@ -1507,7 +1520,10 @@ class TestWriteNoteOverwriteGuard:
assert "Version A" not in content
for suffix in ("-1", "-2"):
stray = await entity_repository.get_by_permalink(f"{canonical_permalink}{suffix}")
async with db.scoped_session(session_maker) as session:
stray = await entity_repository.get_by_permalink(
session, f"{canonical_permalink}{suffix}"
)
assert stray is None, (
f"overwrite=True minted a stray '{suffix}' suffix on the canonical permalink"
)
+292 -231
View File
@@ -73,7 +73,7 @@ async def related_results(session_maker, test_project: Project):
@pytest.mark.asyncio
async def test_create_entity(entity_repository: EntityRepository):
async def test_create_entity(entity_repository: EntityRepository, session_maker):
"""Test creating a new entity"""
entity_data = {
"project_id": entity_repository.project_id,
@@ -85,28 +85,29 @@ async def test_create_entity(entity_repository: EntityRepository):
"created_at": datetime.now(timezone.utc),
"updated_at": datetime.now(timezone.utc),
}
entity = await entity_repository.create(entity_data)
async with db.scoped_session(session_maker) as session:
entity = await entity_repository.create(session, entity_data)
# Verify returned object
assert entity.id is not None
assert entity.title == "Test"
assert isinstance(entity.created_at, datetime)
assert isinstance(entity.updated_at, datetime)
# Verify returned object
assert entity.id is not None
assert entity.title == "Test"
assert isinstance(entity.created_at, datetime)
assert isinstance(entity.updated_at, datetime)
# Verify in database
found = await entity_repository.find_by_id(entity.id)
assert found is not None
assert found.id is not None
assert found.id == entity.id
assert found.title == entity.title
# Verify in database
found = await entity_repository.find_by_id(session, entity.id)
assert found is not None
assert found.id is not None
assert found.id == entity.id
assert found.title == entity.title
# assert relations are eagerly loaded
assert len(entity.observations) == 0
assert len(entity.relations) == 0
# assert relations are eagerly loaded
assert len(entity.observations) == 0
assert len(entity.relations) == 0
@pytest.mark.asyncio
async def test_create_all(entity_repository: EntityRepository):
async def test_create_all(entity_repository: EntityRepository, session_maker):
"""Test creating a new entity"""
entity_data = [
{
@@ -130,33 +131,36 @@ async def test_create_all(entity_repository: EntityRepository):
"updated_at": datetime.now(timezone.utc),
},
]
entities = await entity_repository.create_all(entity_data)
async with db.scoped_session(session_maker) as session:
entities = await entity_repository.create_all(session, entity_data)
assert len(entities) == 2
entity = entities[0]
assert len(entities) == 2
entity = entities[0]
# Verify in database
found = await entity_repository.find_by_id(entity.id)
assert found is not None
assert found.id is not None
assert found.id == entity.id
assert found.title == entity.title
# Verify in database
found = await entity_repository.find_by_id(session, entity.id)
assert found is not None
assert found.id is not None
assert found.id == entity.id
assert found.title == entity.title
# assert relations are eagerly loaded
assert len(entity.observations) == 0
assert len(entity.relations) == 0
# assert relations are eagerly loaded
assert len(entity.observations) == 0
assert len(entity.relations) == 0
@pytest.mark.asyncio
async def test_find_by_id(entity_repository: EntityRepository, sample_entity: Entity):
async def test_find_by_id(
entity_repository: EntityRepository, sample_entity: Entity, session_maker
):
"""Test finding an entity by ID"""
found = await entity_repository.find_by_id(sample_entity.id)
assert found is not None
assert found.id == sample_entity.id
assert found.title == sample_entity.title
async with db.scoped_session(session_maker) as session:
found = await entity_repository.find_by_id(session, sample_entity.id)
assert found is not None
assert found.id == sample_entity.id
assert found.title == sample_entity.title
# Verify against direct database query
async with db.scoped_session(entity_repository.session_maker) as session:
# Verify against direct database query
stmt = select(Entity).where(Entity.id == sample_entity.id)
result = await session.execute(stmt)
db_entity = result.scalar_one()
@@ -165,14 +169,18 @@ async def test_find_by_id(entity_repository: EntityRepository, sample_entity: En
@pytest.mark.asyncio
async def test_update_entity(entity_repository: EntityRepository, sample_entity: Entity):
async def test_update_entity(
entity_repository: EntityRepository, sample_entity: Entity, session_maker
):
"""Test updating an entity"""
updated = await entity_repository.update(sample_entity.id, {"title": "Updated title"})
assert updated is not None
assert updated.title == "Updated title"
async with db.scoped_session(session_maker) as session:
updated = await entity_repository.update(
session, sample_entity.id, {"title": "Updated title"}
)
assert updated is not None
assert updated.title == "Updated title"
# Verify in database
async with db.scoped_session(entity_repository.session_maker) as session:
# Verify in database
stmt = select(Entity).where(Entity.id == sample_entity.id)
result = await session.execute(stmt)
db_entity = result.scalar_one()
@@ -181,13 +189,16 @@ async def test_update_entity(entity_repository: EntityRepository, sample_entity:
@pytest.mark.asyncio
async def test_update_entity_returns_with_relations_and_observations(
entity_repository: EntityRepository, entity_with_observations, test_project: Project
entity_repository: EntityRepository,
entity_with_observations,
test_project: Project,
session_maker,
):
"""Test that update() returns entity with observations and relations eagerly loaded."""
entity = entity_with_observations
# Create a target entity and relation
async with db.scoped_session(entity_repository.session_maker) as session:
async with db.scoped_session(session_maker) as session:
target = Entity(
project_id=test_project.id,
title="target",
@@ -210,51 +221,54 @@ async def test_update_entity_returns_with_relations_and_observations(
)
session.add(relation)
# Now update the entity
updated = await entity_repository.update(entity.id, {"title": "Updated with relations"})
# Now update the entity
updated = await entity_repository.update(
session, entity.id, {"title": "Updated with relations"}
)
# Verify returned entity has observations and relations accessible
# (would raise DetachedInstanceError if not eagerly loaded)
assert updated is not None
assert updated.title == "Updated with relations"
# Verify returned entity has observations and relations accessible
# (would raise DetachedInstanceError if not eagerly loaded)
assert updated is not None
assert updated.title == "Updated with relations"
# Access observations - should NOT raise DetachedInstanceError
assert len(updated.observations) == 2
assert updated.observations[0].content in ["First observation", "Second observation"]
# Access observations - should NOT raise DetachedInstanceError
assert len(updated.observations) == 2
assert updated.observations[0].content in ["First observation", "Second observation"]
# Access relations - should NOT raise DetachedInstanceError
assert len(updated.relations) == 1
assert updated.relations[0].relation_type == "connects_to"
assert updated.relations[0].to_name == "target"
# Access relations - should NOT raise DetachedInstanceError
assert len(updated.relations) == 1
assert updated.relations[0].relation_type == "connects_to"
assert updated.relations[0].to_name == "target"
@pytest.mark.asyncio
async def test_delete_entity(entity_repository: EntityRepository, sample_entity):
async def test_delete_entity(entity_repository: EntityRepository, sample_entity, session_maker):
"""Test deleting an entity."""
result = await entity_repository.delete(sample_entity.id)
assert result is True
async with db.scoped_session(session_maker) as session:
result = await entity_repository.delete(session, sample_entity.id)
assert result is True
# Verify deletion
deleted = await entity_repository.find_by_id(sample_entity.id)
assert deleted is None
# Verify deletion
deleted = await entity_repository.find_by_id(session, sample_entity.id)
assert deleted is None
@pytest.mark.asyncio
async def test_delete_entity_with_observations(
entity_repository: EntityRepository, entity_with_observations
entity_repository: EntityRepository, entity_with_observations, session_maker
):
"""Test deleting an entity cascades to its observations."""
entity = entity_with_observations
result = await entity_repository.delete(entity.id)
assert result is True
async with db.scoped_session(session_maker) as session:
result = await entity_repository.delete(session, entity.id)
assert result is True
# Verify entity deletion
deleted = await entity_repository.find_by_id(entity.id)
assert deleted is None
# Verify entity deletion
deleted = await entity_repository.find_by_id(session, entity.id)
assert deleted is None
# Verify observations were cascaded
async with db.scoped_session(entity_repository.session_maker) as session:
# Verify observations were cascaded
query = select(Observation).filter(Observation.entity_id == entity.id)
result = await session.execute(query)
remaining_observations = result.scalars().all()
@@ -262,13 +276,17 @@ async def test_delete_entity_with_observations(
@pytest.mark.asyncio
async def test_delete_entities_by_type(entity_repository: EntityRepository, sample_entity):
async def test_delete_entities_by_type(
entity_repository: EntityRepository, sample_entity, session_maker
):
"""Test deleting entities by type."""
result = await entity_repository.delete_by_fields(note_type=sample_entity.note_type)
assert result is True
async with db.scoped_session(session_maker) as session:
result = await entity_repository.delete_by_fields(
session, note_type=sample_entity.note_type
)
assert result is True
# Verify deletion
async with db.scoped_session(entity_repository.session_maker) as session:
# Verify deletion
query = select(Entity).filter(Entity.note_type == sample_entity.note_type)
result = await session.execute(query)
remaining = result.scalars().all()
@@ -276,32 +294,33 @@ async def test_delete_entities_by_type(entity_repository: EntityRepository, samp
@pytest.mark.asyncio
async def test_delete_entity_with_relations(entity_repository: EntityRepository, related_results):
async def test_delete_entity_with_relations(
entity_repository: EntityRepository, related_results, session_maker
):
"""Test deleting an entity cascades to its relations."""
source, target, relation = related_results
# Delete source entity
result = await entity_repository.delete(source.id)
assert result is True
async with db.scoped_session(session_maker) as session:
result = await entity_repository.delete(session, source.id)
assert result is True
# Verify relation was cascaded
async with db.scoped_session(entity_repository.session_maker) as session:
# Verify relation was cascaded
query = select(Relation).filter(Relation.from_id == source.id)
result = await session.execute(query)
remaining_relations = result.scalars().all()
assert len(remaining_relations) == 0
# Verify target entity still exists. Runs outside the session block above:
# find_by_id opens its own scoped session, and the serialized in-memory pool
# (one connection, see #940) deadlocks on nested session checkouts.
target_exists = await entity_repository.find_by_id(target.id)
assert target_exists is not None
# Verify target entity still exists.
target_exists = await entity_repository.find_by_id(session, target.id)
assert target_exists is not None
@pytest.mark.asyncio
async def test_delete_nonexistent_entity(entity_repository: EntityRepository):
async def test_delete_nonexistent_entity(entity_repository: EntityRepository, session_maker):
"""Test deleting an entity that doesn't exist."""
result = await entity_repository.delete(0)
async with db.scoped_session(session_maker) as session:
result = await entity_repository.delete(session, 0)
assert result is False
@@ -346,31 +365,34 @@ async def test_entities(session_maker, test_project: Project):
@pytest.mark.asyncio
async def test_find_by_permalinks(entity_repository: EntityRepository, test_entities):
async def test_find_by_permalinks(
entity_repository: EntityRepository, test_entities, session_maker
):
"""Test finding multiple entities by their type/name pairs."""
# Test finding multiple entities
permalinks = [e.permalink for e in test_entities]
found = await entity_repository.find_by_permalinks(permalinks)
assert len(found) == 3
names = {e.title for e in found}
assert names == {"entity1", "entity2", "entity3"}
async with db.scoped_session(session_maker) as session:
# Test finding multiple entities
permalinks = [e.permalink for e in test_entities]
found = await entity_repository.find_by_permalinks(session, permalinks)
assert len(found) == 3
names = {e.title for e in found}
assert names == {"entity1", "entity2", "entity3"}
# Test finding subset of entities
permalinks = [e.permalink for e in test_entities if e.title != "entity2"]
found = await entity_repository.find_by_permalinks(permalinks)
assert len(found) == 2
names = {e.title for e in found}
assert names == {"entity1", "entity3"}
# Test finding subset of entities
permalinks = [e.permalink for e in test_entities if e.title != "entity2"]
found = await entity_repository.find_by_permalinks(session, permalinks)
assert len(found) == 2
names = {e.title for e in found}
assert names == {"entity1", "entity3"}
# Test with non-existent entities
permalinks = ["type1/entity1", "type3/nonexistent"]
found = await entity_repository.find_by_permalinks(permalinks)
assert len(found) == 1
assert found[0].title == "entity1"
# Test with non-existent entities
permalinks = ["type1/entity1", "type3/nonexistent"]
found = await entity_repository.find_by_permalinks(session, permalinks)
assert len(found) == 1
assert found[0].title == "entity1"
# Test empty input
found = await entity_repository.find_by_permalinks([])
assert len(found) == 0
# Test empty input
found = await entity_repository.find_by_permalinks(session, [])
assert len(found) == 0
@pytest.mark.asyncio
@@ -439,23 +461,24 @@ async def test_get_by_title(entity_repository: EntityRepository, session_maker):
session.add_all(entities)
await session.flush()
# Test getting by exact title
found = await entity_repository.get_by_title("Unique Title")
assert found is not None
assert len(found) == 1
assert found[0].title == "Unique Title"
async with db.scoped_session(session_maker) as session:
# Test getting by exact title
found = await entity_repository.get_by_title(session, "Unique Title")
assert found is not None
assert len(found) == 1
assert found[0].title == "Unique Title"
# Test case sensitivity
found = await entity_repository.get_by_title("unique title")
assert not found # Should be case-sensitive
# Test case sensitivity
found = await entity_repository.get_by_title(session, "unique title")
assert not found # Should be case-sensitive
# Test non-existent title
found = await entity_repository.get_by_title("Non Existent")
assert not found
# Test non-existent title
found = await entity_repository.get_by_title(session, "Non Existent")
assert not found
# Test multiple rows found
found = await entity_repository.get_by_title("Another Title")
assert len(found) == 2
# Test multiple rows found
found = await entity_repository.get_by_title(session, "Another Title")
assert len(found) == 2
@pytest.mark.asyncio
@@ -506,16 +529,17 @@ async def test_get_by_title_returns_shortest_path_first(
session.add_all(entities)
await session.flush()
# Get all entities with title "My Note"
found = await entity_repository.get_by_title("My Note")
async with db.scoped_session(session_maker) as session:
# Get all entities with title "My Note"
found = await entity_repository.get_by_title(session, "My Note")
# Should return all 3
assert len(found) == 3
# Should return all 3
assert len(found) == 3
# Should be ordered by path length (shortest first)
assert found[0].file_path == "My Note.md" # shortest
assert found[1].file_path == "docs/My Note.md" # medium
assert found[2].file_path == "archive/old/2024/My Note.md" # longest
# Should be ordered by path length (shortest first)
assert found[0].file_path == "My Note.md" # shortest
assert found[1].file_path == "docs/My Note.md" # medium
assert found[2].file_path == "archive/old/2024/My Note.md" # longest
@pytest.mark.asyncio
@@ -538,14 +562,15 @@ async def test_get_by_file_path(entity_repository: EntityRepository, session_mak
session.add_all(entities)
await session.flush()
# Test getting by file_path
found = await entity_repository.get_by_file_path("test/unique-title.md")
assert found is not None
assert found.title == "Unique Title"
async with db.scoped_session(session_maker) as session:
# Test getting by file_path
found = await entity_repository.get_by_file_path(session, "test/unique-title.md")
assert found is not None
assert found.title == "Unique Title"
# Test non-existent file_path
found = await entity_repository.get_by_file_path("not/a/real/file.md")
assert found is None
# Test non-existent file_path
found = await entity_repository.get_by_file_path(session, "not/a/real/file.md")
assert found is None
@pytest.mark.asyncio
@@ -608,8 +633,9 @@ async def test_get_distinct_directories(entity_repository: EntityRepository, ses
session.add_all(entities)
await session.flush()
# Get distinct directories
directories = await entity_repository.get_distinct_directories()
async with db.scoped_session(session_maker) as session:
# Get distinct directories
directories = await entity_repository.get_distinct_directories(session)
# Verify directories are extracted correctly
assert isinstance(directories, list)
@@ -636,9 +662,12 @@ async def test_get_distinct_directories(entity_repository: EntityRepository, ses
@pytest.mark.asyncio
async def test_get_distinct_directories_empty_db(entity_repository: EntityRepository):
async def test_get_distinct_directories_empty_db(
entity_repository: EntityRepository, session_maker
):
"""Test getting distinct directories when database is empty."""
directories = await entity_repository.get_distinct_directories()
async with db.scoped_session(session_maker) as session:
directories = await entity_repository.get_distinct_directories(session)
assert directories == []
@@ -692,33 +721,34 @@ async def test_find_by_directory_prefix(entity_repository: EntityRepository, ses
session.add_all(entities)
await session.flush()
# Test finding all entities in "docs" directory and subdirectories
docs_entities = await entity_repository.find_by_directory_prefix("docs")
assert len(docs_entities) == 3
file_paths = {e.file_path for e in docs_entities}
assert file_paths == {"docs/file1.md", "docs/guides/file2.md", "docs/api/file3.md"}
async with db.scoped_session(session_maker) as session:
# Test finding all entities in "docs" directory and subdirectories
docs_entities = await entity_repository.find_by_directory_prefix(session, "docs")
assert len(docs_entities) == 3
file_paths = {e.file_path for e in docs_entities}
assert file_paths == {"docs/file1.md", "docs/guides/file2.md", "docs/api/file3.md"}
# Test finding entities in "docs/guides" subdirectory
guides_entities = await entity_repository.find_by_directory_prefix("docs/guides")
assert len(guides_entities) == 1
assert guides_entities[0].file_path == "docs/guides/file2.md"
# Test finding entities in "docs/guides" subdirectory
guides_entities = await entity_repository.find_by_directory_prefix(session, "docs/guides")
assert len(guides_entities) == 1
assert guides_entities[0].file_path == "docs/guides/file2.md"
# Test finding entities in "specs" directory
specs_entities = await entity_repository.find_by_directory_prefix("specs")
assert len(specs_entities) == 1
assert specs_entities[0].file_path == "specs/file4.md"
# Test finding entities in "specs" directory
specs_entities = await entity_repository.find_by_directory_prefix(session, "specs")
assert len(specs_entities) == 1
assert specs_entities[0].file_path == "specs/file4.md"
# Test with root directory (empty string)
all_entities = await entity_repository.find_by_directory_prefix("")
assert len(all_entities) == 4
# Test with root directory (empty string)
all_entities = await entity_repository.find_by_directory_prefix(session, "")
assert len(all_entities) == 4
# Test with root directory (slash)
all_entities = await entity_repository.find_by_directory_prefix("/")
assert len(all_entities) == 4
# Test with root directory (slash)
all_entities = await entity_repository.find_by_directory_prefix(session, "/")
assert len(all_entities) == 4
# Test with non-existent directory
nonexistent = await entity_repository.find_by_directory_prefix("nonexistent")
assert len(nonexistent) == 0
# Test with non-existent directory
nonexistent = await entity_repository.find_by_directory_prefix(session, "nonexistent")
assert len(nonexistent) == 0
@pytest.mark.asyncio
@@ -746,18 +776,19 @@ async def test_find_by_directory_prefix_basic_fields_only(
session.add(entity)
await session.flush()
# Query entity by directory prefix
entities = await entity_repository.find_by_directory_prefix("docs")
assert len(entities) == 1
async with db.scoped_session(session_maker) as session:
# Query entity by directory prefix
entities = await entity_repository.find_by_directory_prefix(session, "docs")
assert len(entities) == 1
# Verify basic fields are present (all we need for directory trees)
entity = entities[0]
assert entity.title == "Test Entity"
assert entity.file_path == "docs/test.md"
assert entity.permalink == "docs/test"
assert entity.note_type == "test"
assert entity.content_type == "text/markdown"
assert entity.updated_at is not None
# Verify basic fields are present (all we need for directory trees)
entity = entities[0]
assert entity.title == "Test Entity"
assert entity.file_path == "docs/test.md"
assert entity.permalink == "docs/test"
assert entity.note_type == "test"
assert entity.content_type == "text/markdown"
assert entity.updated_at is not None
@pytest.mark.asyncio
@@ -800,8 +831,9 @@ async def test_get_all_file_paths(entity_repository: EntityRepository, session_m
session.add_all(entities)
await session.flush()
# Get all file paths
file_paths = await entity_repository.get_all_file_paths()
async with db.scoped_session(session_maker) as session:
# Get all file paths
file_paths = await entity_repository.get_all_file_paths(session)
# Verify results
assert isinstance(file_paths, list)
@@ -810,9 +842,10 @@ async def test_get_all_file_paths(entity_repository: EntityRepository, session_m
@pytest.mark.asyncio
async def test_get_all_file_paths_empty_db(entity_repository: EntityRepository):
async def test_get_all_file_paths_empty_db(entity_repository: EntityRepository, session_maker):
"""Test getting all file paths when database is empty."""
file_paths = await entity_repository.get_all_file_paths()
async with db.scoped_session(session_maker) as session:
file_paths = await entity_repository.get_all_file_paths(session)
assert file_paths == []
@@ -869,16 +902,17 @@ async def test_get_all_file_paths_performance(entity_repository: EntityRepositor
session.add(relation)
await session.flush()
# Get all file paths - should be fast and not load relationships
file_paths = await entity_repository.get_all_file_paths()
async with db.scoped_session(session_maker) as session:
# Get all file paths - should be fast and not load relationships
file_paths = await entity_repository.get_all_file_paths(session)
# Verify results - just file paths, no entities or relationships loaded
assert len(file_paths) == 2
assert set(file_paths) == {"test/entity1.md", "test/entity2.md"}
# Verify results - just file paths, no entities or relationships loaded
assert len(file_paths) == 2
assert set(file_paths) == {"test/entity1.md", "test/entity2.md"}
# Result should be list of strings, not entity objects
for path in file_paths:
assert isinstance(path, str)
# Result should be list of strings, not entity objects
for path in file_paths:
assert isinstance(path, str)
@pytest.mark.asyncio
@@ -920,12 +954,13 @@ async def test_get_all_file_paths_project_isolation(
session.add(entity2)
await session.flush()
# Get all file paths for project 1
file_paths = await entity_repository.get_all_file_paths()
async with db.scoped_session(session_maker) as session:
# Get all file paths for project 1
file_paths = await entity_repository.get_all_file_paths(session)
# Should only include files from project 1
assert len(file_paths) == 1
assert file_paths == ["test/file1.md"]
# Should only include files from project 1
assert len(file_paths) == 1
assert file_paths == ["test/file1.md"]
# -------------------------------------------------------------------------
@@ -934,14 +969,17 @@ async def test_get_all_file_paths_project_isolation(
@pytest.mark.asyncio
async def test_permalink_exists(entity_repository: EntityRepository, sample_entity: Entity):
async def test_permalink_exists(
entity_repository: EntityRepository, sample_entity: Entity, session_maker
):
"""Test checking if a permalink exists without loading full entity."""
# Existing permalink should return True
assert sample_entity.permalink is not None
assert await entity_repository.permalink_exists(sample_entity.permalink) is True
async with db.scoped_session(session_maker) as session:
assert await entity_repository.permalink_exists(session, sample_entity.permalink) is True
# Non-existent permalink should return False
assert await entity_repository.permalink_exists("nonexistent/permalink") is False
# Non-existent permalink should return False
assert await entity_repository.permalink_exists(session, "nonexistent/permalink") is False
@pytest.mark.asyncio
@@ -980,40 +1018,49 @@ async def test_permalink_exists_project_isolation(
)
session.add(entity2)
# Should find entity1's permalink in project 1
assert await entity_repository.permalink_exists("test/entity1") is True
async with db.scoped_session(session_maker) as session:
# Should find entity1's permalink in project 1
assert await entity_repository.permalink_exists(session, "test/entity1") is True
# Should NOT find entity2's permalink (it's in project 2)
assert await entity_repository.permalink_exists("test/entity2") is False
# Should NOT find entity2's permalink (it's in project 2)
assert await entity_repository.permalink_exists(session, "test/entity2") is False
@pytest.mark.asyncio
async def test_get_file_path_for_permalink(
entity_repository: EntityRepository, sample_entity: Entity
entity_repository: EntityRepository, sample_entity: Entity, session_maker
):
"""Test getting file_path for a permalink without loading full entity."""
# Existing permalink should return file_path
assert sample_entity.permalink is not None
file_path = await entity_repository.get_file_path_for_permalink(sample_entity.permalink)
assert file_path == sample_entity.file_path
async with db.scoped_session(session_maker) as session:
file_path = await entity_repository.get_file_path_for_permalink(
session, sample_entity.permalink
)
assert file_path == sample_entity.file_path
# Non-existent permalink should return None
result = await entity_repository.get_file_path_for_permalink("nonexistent/permalink")
assert result is None
# Non-existent permalink should return None
result = await entity_repository.get_file_path_for_permalink(
session, "nonexistent/permalink"
)
assert result is None
@pytest.mark.asyncio
async def test_get_permalink_for_file_path(
entity_repository: EntityRepository, sample_entity: Entity
entity_repository: EntityRepository, sample_entity: Entity, session_maker
):
"""Test getting permalink for a file_path without loading full entity."""
# Existing file_path should return permalink
permalink = await entity_repository.get_permalink_for_file_path(sample_entity.file_path)
assert permalink == sample_entity.permalink
async with db.scoped_session(session_maker) as session:
permalink = await entity_repository.get_permalink_for_file_path(
session, sample_entity.file_path
)
assert permalink == sample_entity.permalink
# Non-existent file_path should return None
result = await entity_repository.get_permalink_for_file_path("nonexistent/path.md")
assert result is None
# Non-existent file_path should return None
result = await entity_repository.get_permalink_for_file_path(session, "nonexistent/path.md")
assert result is None
@pytest.mark.asyncio
@@ -1042,7 +1089,8 @@ async def test_get_all_permalinks(entity_repository: EntityRepository, session_m
)
session.add_all([entity1, entity2])
permalinks = await entity_repository.get_all_permalinks()
async with db.scoped_session(session_maker) as session:
permalinks = await entity_repository.get_all_permalinks(session)
assert len(permalinks) == 2
assert set(permalinks) == {"test/entity1", "test/entity2"}
@@ -1054,7 +1102,10 @@ async def test_get_all_permalinks(entity_repository: EntityRepository, session_m
@pytest.mark.asyncio
async def test_find_by_ids_for_hydration_skips_eager_load_options(
entity_repository: EntityRepository, sample_entity: Entity, monkeypatch: pytest.MonkeyPatch
entity_repository: EntityRepository,
sample_entity: Entity,
monkeypatch: pytest.MonkeyPatch,
session_maker,
):
"""Context hydration should bypass relationship loader options."""
@@ -1063,7 +1114,8 @@ async def test_find_by_ids_for_hydration_skips_eager_load_options(
monkeypatch.setattr(entity_repository, "get_load_options", fail_get_load_options)
found = await entity_repository.find_by_ids_for_hydration([sample_entity.id])
async with db.scoped_session(session_maker) as session:
found = await entity_repository.find_by_ids_for_hydration(session, [sample_entity.id])
assert len(found) == 1
assert found[0].id == sample_entity.id
@@ -1095,12 +1147,13 @@ async def test_find_by_ids_for_hydration_can_include_cross_project_entities(
await session.flush()
other_entity_id = other_entity.id
project_scoped = await entity_repository.find_by_ids_for_hydration(
[sample_entity.id, other_entity_id]
)
cross_project = await entity_repository.find_by_ids_for_hydration(
[sample_entity.id, other_entity_id], include_cross_project=True
)
async with db.scoped_session(session_maker) as session:
project_scoped = await entity_repository.find_by_ids_for_hydration(
session, [sample_entity.id, other_entity_id]
)
cross_project = await entity_repository.find_by_ids_for_hydration(
session, [sample_entity.id, other_entity_id], include_cross_project=True
)
assert {entity.id for entity in project_scoped} == {sample_entity.id}
assert {entity.id for entity in cross_project} == {sample_entity.id, other_entity_id}
@@ -1132,7 +1185,8 @@ async def test_get_permalink_to_file_path_map(entity_repository: EntityRepositor
)
session.add_all([entity1, entity2])
mapping = await entity_repository.get_permalink_to_file_path_map()
async with db.scoped_session(session_maker) as session:
mapping = await entity_repository.get_permalink_to_file_path_map(session)
assert len(mapping) == 2
assert mapping["test/entity1"] == "test/entity1.md"
@@ -1165,7 +1219,8 @@ async def test_get_file_path_to_permalink_map(entity_repository: EntityRepositor
)
session.add_all([entity1, entity2])
mapping = await entity_repository.get_file_path_to_permalink_map()
async with db.scoped_session(session_maker) as session:
mapping = await entity_repository.get_file_path_to_permalink_map(session)
assert len(mapping) == 2
assert mapping["test/entity1.md"] == "test/entity1"
@@ -1220,8 +1275,9 @@ async def test_find_without_relations_returns_isolated_entities(
)
session.add(relation)
result = await entity_repository.find_without_relations()
titles = {entity.title for entity in result}
async with db.scoped_session(session_maker) as session:
result = await entity_repository.find_without_relations(session)
titles = {entity.title for entity in result}
assert "Orphan" in titles
assert "Source" not in titles
@@ -1266,8 +1322,9 @@ async def test_find_without_relations_excludes_unresolved_outgoing_links(
)
session.add(unresolved_relation)
result = await entity_repository.find_without_relations()
titles = {entity.title for entity in result}
async with db.scoped_session(session_maker) as session:
result = await entity_repository.find_without_relations(session)
titles = {entity.title for entity in result}
assert "Orphan" in titles
assert "Source" not in titles
@@ -1305,15 +1362,19 @@ async def test_find_without_relations_respects_project_scope(
)
session.add(other_orphan)
result = await entity_repository.find_without_relations()
titles = {entity.title for entity in result}
async with db.scoped_session(session_maker) as session:
result = await entity_repository.find_without_relations(session)
titles = {entity.title for entity in result}
assert "Active Project Orphan" in titles
assert "Other Project Orphan" not in titles
@pytest.mark.asyncio
async def test_find_without_relations_empty_project(entity_repository: EntityRepository):
async def test_find_without_relations_empty_project(
entity_repository: EntityRepository, session_maker
):
"""An empty project returns no orphans."""
result = await entity_repository.find_without_relations()
async with db.scoped_session(session_maker) as session:
result = await entity_repository.find_without_relations(session)
assert result == []
+285 -262
View File
@@ -3,6 +3,7 @@
import pytest
from datetime import datetime, timezone
from basic_memory import db
from basic_memory.models.knowledge import Entity
from basic_memory.repository.entity_repository import EntityRepository
from basic_memory.repository.project_repository import ProjectRepository
@@ -10,7 +11,7 @@ from basic_memory.services.exceptions import SyncFatalError
@pytest.mark.asyncio
async def test_upsert_entity_new_entity(entity_repository: EntityRepository):
async def test_upsert_entity_new_entity(entity_repository: EntityRepository, session_maker):
"""Test upserting a completely new entity."""
entity = Entity(
project_id=entity_repository.project_id,
@@ -23,7 +24,8 @@ async def test_upsert_entity_new_entity(entity_repository: EntityRepository):
updated_at=datetime.now(timezone.utc),
)
result = await entity_repository.upsert_entity(entity)
async with db.scoped_session(session_maker) as session:
result = await entity_repository.upsert_entity(session, entity)
assert result.id is not None
assert result.title == "Test Entity"
@@ -32,7 +34,7 @@ async def test_upsert_entity_new_entity(entity_repository: EntityRepository):
@pytest.mark.asyncio
async def test_upsert_entity_same_file_update(entity_repository: EntityRepository):
async def test_upsert_entity_same_file_update(entity_repository: EntityRepository, session_maker):
"""Test upserting an entity that already exists with same file_path."""
# Create initial entity
entity1 = Entity(
@@ -46,22 +48,23 @@ async def test_upsert_entity_same_file_update(entity_repository: EntityRepositor
updated_at=datetime.now(timezone.utc),
)
result1 = await entity_repository.upsert_entity(entity1)
original_id = result1.id
async with db.scoped_session(session_maker) as session:
result1 = await entity_repository.upsert_entity(session, entity1)
original_id = result1.id
# Update with same file_path and permalink
entity2 = Entity(
project_id=entity_repository.project_id,
title="Updated Title",
note_type="note",
permalink="test/test-entity", # Same permalink
file_path="test/test-entity.md", # Same file_path
content_type="text/markdown",
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
)
# Update with same file_path and permalink
entity2 = Entity(
project_id=entity_repository.project_id,
title="Updated Title",
note_type="note",
permalink="test/test-entity", # Same permalink
file_path="test/test-entity.md", # Same file_path
content_type="text/markdown",
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
)
result2 = await entity_repository.upsert_entity(entity2)
result2 = await entity_repository.upsert_entity(session, entity2)
# Should update existing entity (same ID)
assert result2.id == original_id
@@ -71,7 +74,9 @@ async def test_upsert_entity_same_file_update(entity_repository: EntityRepositor
@pytest.mark.asyncio
async def test_upsert_entity_preserves_external_id(entity_repository: EntityRepository):
async def test_upsert_entity_preserves_external_id(
entity_repository: EntityRepository, session_maker
):
"""Test that upserting an entity with the same file_path preserves the original external_id.
Trigger: force full re-index creates a new Entity model (with a fresh UUID)
@@ -91,31 +96,34 @@ async def test_upsert_entity_preserves_external_id(entity_repository: EntityRepo
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
)
result1 = await entity_repository.upsert_entity(entity1)
assert result1.external_id == "original-stable-id"
async with db.scoped_session(session_maker) as session:
result1 = await entity_repository.upsert_entity(session, entity1)
assert result1.external_id == "original-stable-id"
# Simulate re-index: new Entity model with a DIFFERENT external_id
entity2 = Entity(
project_id=entity_repository.project_id,
title="Shared Note (updated)",
note_type="note",
permalink="test/shared-note",
file_path="test/shared-note.md",
content_type="text/markdown",
external_id="newly-generated-uuid", # would break share links
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
)
result2 = await entity_repository.upsert_entity(entity2)
# Simulate re-index: new Entity model with a DIFFERENT external_id
entity2 = Entity(
project_id=entity_repository.project_id,
title="Shared Note (updated)",
note_type="note",
permalink="test/shared-note",
file_path="test/shared-note.md",
content_type="text/markdown",
external_id="newly-generated-uuid", # would break share links
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
)
result2 = await entity_repository.upsert_entity(session, entity2)
# ID preserved, title updated, external_id stable
assert result2.id == result1.id
assert result2.title == "Shared Note (updated)"
assert result2.external_id == "original-stable-id"
# ID preserved, title updated, external_id stable
assert result2.id == result1.id
assert result2.title == "Shared Note (updated)"
assert result2.external_id == "original-stable-id"
@pytest.mark.asyncio
async def test_upsert_entity_permalink_conflict_different_file(entity_repository: EntityRepository):
async def test_upsert_entity_permalink_conflict_different_file(
entity_repository: EntityRepository, session_maker
):
"""Test upserting an entity with permalink conflict but different file_path."""
# Create initial entity
entity1 = Entity(
@@ -129,57 +137,61 @@ async def test_upsert_entity_permalink_conflict_different_file(entity_repository
updated_at=datetime.now(timezone.utc),
)
result1 = await entity_repository.upsert_entity(entity1)
first_id = result1.id
async with db.scoped_session(session_maker) as session:
result1 = await entity_repository.upsert_entity(session, entity1)
first_id = result1.id
# Try to create entity with same permalink but different file_path
entity2 = Entity(
project_id=entity_repository.project_id,
title="Second Entity",
note_type="note",
permalink="test/shared-permalink", # Same permalink
file_path="test/second-file.md", # Different file_path
content_type="text/markdown",
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
)
result2 = await entity_repository.upsert_entity(entity2)
# Should create new entity with unique permalink
assert result2.id != first_id
assert result2.title == "Second Entity"
assert result2.permalink == "test/shared-permalink-1" # Should get suffix
assert result2.file_path == "test/second-file.md"
# Original entity should be unchanged
original = await entity_repository.get_by_permalink("test/shared-permalink")
assert original is not None
assert original.id == first_id
assert original.title == "First Entity"
@pytest.mark.asyncio
async def test_upsert_entity_multiple_permalink_conflicts(entity_repository: EntityRepository):
"""Test upserting multiple entities with permalink conflicts."""
base_permalink = "test/conflict"
# Create entities with conflicting permalinks
entities = []
for i in range(3):
entity = Entity(
# Try to create entity with same permalink but different file_path
entity2 = Entity(
project_id=entity_repository.project_id,
title=f"Entity {i + 1}",
title="Second Entity",
note_type="note",
permalink=base_permalink, # All try to use same permalink
file_path=f"test/file-{i + 1}.md", # Different file paths
permalink="test/shared-permalink", # Same permalink
file_path="test/second-file.md", # Different file_path
content_type="text/markdown",
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
)
result = await entity_repository.upsert_entity(entity)
entities.append(result)
result2 = await entity_repository.upsert_entity(session, entity2)
# Should create new entity with unique permalink
assert result2.id != first_id
assert result2.title == "Second Entity"
assert result2.permalink == "test/shared-permalink-1" # Should get suffix
assert result2.file_path == "test/second-file.md"
# Original entity should be unchanged
original = await entity_repository.get_by_permalink(session, "test/shared-permalink")
assert original is not None
assert original.id == first_id
assert original.title == "First Entity"
@pytest.mark.asyncio
async def test_upsert_entity_multiple_permalink_conflicts(
entity_repository: EntityRepository, session_maker
):
"""Test upserting multiple entities with permalink conflicts."""
base_permalink = "test/conflict"
# Create entities with conflicting permalinks
entities = []
async with db.scoped_session(session_maker) as session:
for i in range(3):
entity = Entity(
project_id=entity_repository.project_id,
title=f"Entity {i + 1}",
note_type="note",
permalink=base_permalink, # All try to use same permalink
file_path=f"test/file-{i + 1}.md", # Different file paths
content_type="text/markdown",
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
)
result = await entity_repository.upsert_entity(session, entity)
entities.append(result)
# Verify permalinks are unique
expected_permalinks = ["test/conflict", "test/conflict-1", "test/conflict-2"]
@@ -193,7 +205,9 @@ async def test_upsert_entity_multiple_permalink_conflicts(entity_repository: Ent
@pytest.mark.asyncio
async def test_upsert_entity_race_condition_file_path(entity_repository: EntityRepository):
async def test_upsert_entity_race_condition_file_path(
entity_repository: EntityRepository, session_maker
):
"""Test that upsert handles file_path conflicts using ON CONFLICT DO UPDATE.
With SQLite's ON CONFLICT, race conditions are handled at the database level
@@ -212,24 +226,25 @@ async def test_upsert_entity_race_condition_file_path(entity_repository: EntityR
updated_at=datetime.now(timezone.utc),
)
result1 = await entity_repository.upsert_entity(entity1)
original_id = result1.id
async with db.scoped_session(session_maker) as session:
result1 = await entity_repository.upsert_entity(session, entity1)
original_id = result1.id
# Create another entity with same file_path but different title and permalink
# This simulates a concurrent update scenario
entity2 = Entity(
project_id=entity_repository.project_id,
title="Race Condition Test",
note_type="note",
permalink="test/race-entity",
file_path="test/race-file.md", # Same file path as entity1
content_type="text/markdown",
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
)
# Create another entity with same file_path but different title and permalink
# This simulates a concurrent update scenario
entity2 = Entity(
project_id=entity_repository.project_id,
title="Race Condition Test",
note_type="note",
permalink="test/race-entity",
file_path="test/race-file.md", # Same file path as entity1
content_type="text/markdown",
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
)
# ON CONFLICT should update the existing entity
result2 = await entity_repository.upsert_entity(entity2)
# ON CONFLICT should update the existing entity
result2 = await entity_repository.upsert_entity(session, entity2)
# Should return the updated original entity (same ID)
assert result2.id == original_id
@@ -239,7 +254,7 @@ async def test_upsert_entity_race_condition_file_path(entity_repository: EntityR
@pytest.mark.asyncio
async def test_upsert_entity_gap_in_suffixes(entity_repository: EntityRepository):
async def test_upsert_entity_gap_in_suffixes(entity_repository: EntityRepository, session_maker):
"""Test that upsert finds the next available suffix even with gaps."""
# Manually create entities with non-sequential suffixes
base_permalink = "test/gap"
@@ -248,32 +263,35 @@ async def test_upsert_entity_gap_in_suffixes(entity_repository: EntityRepository
# (skipping "test/gap-2")
permalinks = [base_permalink, f"{base_permalink}-1", f"{base_permalink}-3"]
for i, permalink in enumerate(permalinks):
entity = Entity(
async with db.scoped_session(session_maker) as session:
for i, permalink in enumerate(permalinks):
entity = Entity(
project_id=entity_repository.project_id,
title=f"Entity {i + 1}",
note_type="note",
permalink=permalink,
file_path=f"test/gap-file-{i + 1}.md",
content_type="text/markdown",
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
)
await entity_repository.add(
session, entity
) # Use direct add to set specific permalinks
# Now try to upsert a new entity that should get "test/gap-2"
new_entity = Entity(
project_id=entity_repository.project_id,
title=f"Entity {i + 1}",
title="Gap Filler",
note_type="note",
permalink=permalink,
file_path=f"test/gap-file-{i + 1}.md",
permalink=base_permalink, # Will conflict
file_path="test/gap-new-file.md",
content_type="text/markdown",
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
)
await entity_repository.add(entity) # Use direct add to set specific permalinks
# Now try to upsert a new entity that should get "test/gap-2"
new_entity = Entity(
project_id=entity_repository.project_id,
title="Gap Filler",
note_type="note",
permalink=base_permalink, # Will conflict
file_path="test/gap-new-file.md",
content_type="text/markdown",
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
)
result = await entity_repository.upsert_entity(new_entity)
result = await entity_repository.upsert_entity(session, new_entity)
# Should get the next available suffix - our implementation finds gaps
# so it should be "test/gap-2" (filling the gap)
@@ -291,7 +309,7 @@ async def test_upsert_entity_project_scoping_isolation(session_maker):
3. No "multiple rows" errors occur when similar entities exist across projects
"""
# Create two separate projects
project_repository = ProjectRepository(session_maker)
project_repository = ProjectRepository()
project1_data = {
"name": "project-1",
@@ -300,100 +318,101 @@ async def test_upsert_entity_project_scoping_isolation(session_maker):
"is_active": True,
"is_default": False,
}
project1 = await project_repository.create(project1_data)
async with db.scoped_session(session_maker) as session:
project1 = await project_repository.create(session, project1_data)
project2_data = {
"name": "project-2",
"description": "Second test project",
"path": "/tmp/project2",
"is_active": True,
"is_default": False,
}
project2 = await project_repository.create(project2_data)
project2_data = {
"name": "project-2",
"description": "Second test project",
"path": "/tmp/project2",
"is_active": True,
"is_default": False,
}
project2 = await project_repository.create(session, project2_data)
# Create entity repositories for each project
repo1 = EntityRepository(session_maker, project_id=project1.id)
repo2 = EntityRepository(session_maker, project_id=project2.id)
# Create entity repositories for each project
repo1 = EntityRepository(project_id=project1.id)
repo2 = EntityRepository(project_id=project2.id)
# Create entities with identical permalinks and file_paths in different projects
entity1 = Entity(
project_id=project1.id,
title="Shared Entity",
note_type="note",
permalink="docs/shared-name",
file_path="docs/shared-name.md",
content_type="text/markdown",
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
)
# Create entities with identical permalinks and file_paths in different projects
entity1 = Entity(
project_id=project1.id,
title="Shared Entity",
note_type="note",
permalink="docs/shared-name",
file_path="docs/shared-name.md",
content_type="text/markdown",
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
)
entity2 = Entity(
project_id=project2.id,
title="Shared Entity",
note_type="note",
permalink="docs/shared-name", # Same permalink
file_path="docs/shared-name.md", # Same file_path
content_type="text/markdown",
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
)
entity2 = Entity(
project_id=project2.id,
title="Shared Entity",
note_type="note",
permalink="docs/shared-name", # Same permalink
file_path="docs/shared-name.md", # Same file_path
content_type="text/markdown",
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
)
# These should succeed without "multiple rows" errors
result1 = await repo1.upsert_entity(entity1)
result2 = await repo2.upsert_entity(entity2)
# These should succeed without "multiple rows" errors
result1 = await repo1.upsert_entity(session, entity1)
result2 = await repo2.upsert_entity(session, entity2)
# Verify both entities were created successfully
assert result1.id is not None
assert result2.id is not None
assert result1.id != result2.id # Different entities
assert result1.project_id == project1.id
assert result2.project_id == project2.id
assert result1.permalink == "docs/shared-name"
assert result2.permalink == "docs/shared-name"
# Verify both entities were created successfully
assert result1.id is not None
assert result2.id is not None
assert result1.id != result2.id # Different entities
assert result1.project_id == project1.id
assert result2.project_id == project2.id
assert result1.permalink == "docs/shared-name"
assert result2.permalink == "docs/shared-name"
# Test updating entities in different projects (should also work without conflicts)
entity1_update = Entity(
project_id=project1.id,
title="Updated Shared Entity",
note_type="note",
permalink="docs/shared-name",
file_path="docs/shared-name.md",
content_type="text/markdown",
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
)
# Test updating entities in different projects (should also work without conflicts)
entity1_update = Entity(
project_id=project1.id,
title="Updated Shared Entity",
note_type="note",
permalink="docs/shared-name",
file_path="docs/shared-name.md",
content_type="text/markdown",
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
)
entity2_update = Entity(
project_id=project2.id,
title="Also Updated Shared Entity",
note_type="note",
permalink="docs/shared-name",
file_path="docs/shared-name.md",
content_type="text/markdown",
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
)
entity2_update = Entity(
project_id=project2.id,
title="Also Updated Shared Entity",
note_type="note",
permalink="docs/shared-name",
file_path="docs/shared-name.md",
content_type="text/markdown",
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
)
# Updates should work without conflicts
updated1 = await repo1.upsert_entity(entity1_update)
updated2 = await repo2.upsert_entity(entity2_update)
# Updates should work without conflicts
updated1 = await repo1.upsert_entity(session, entity1_update)
updated2 = await repo2.upsert_entity(session, entity2_update)
# Should update existing entities (same IDs)
assert updated1.id == result1.id
assert updated2.id == result2.id
assert updated1.title == "Updated Shared Entity"
assert updated2.title == "Also Updated Shared Entity"
# Should update existing entities (same IDs)
assert updated1.id == result1.id
assert updated2.id == result2.id
assert updated1.title == "Updated Shared Entity"
assert updated2.title == "Also Updated Shared Entity"
# Verify cross-project queries don't interfere
found_in_project1 = await repo1.get_by_permalink("docs/shared-name")
found_in_project2 = await repo2.get_by_permalink("docs/shared-name")
# Verify cross-project queries don't interfere
found_in_project1 = await repo1.get_by_permalink(session, "docs/shared-name")
found_in_project2 = await repo2.get_by_permalink(session, "docs/shared-name")
assert found_in_project1 is not None
assert found_in_project2 is not None
assert found_in_project1.id == updated1.id
assert found_in_project2.id == updated2.id
assert found_in_project1.title == "Updated Shared Entity"
assert found_in_project2.title == "Also Updated Shared Entity"
assert found_in_project1 is not None
assert found_in_project2 is not None
assert found_in_project1.id == updated1.id
assert found_in_project2.id == updated2.id
assert found_in_project1.title == "Updated Shared Entity"
assert found_in_project2.title == "Also Updated Shared Entity"
@pytest.mark.asyncio
@@ -405,7 +424,7 @@ async def test_upsert_entity_permalink_conflict_within_project_only(session_make
permalink conflict resolution.
"""
# Create two separate projects
project_repository = ProjectRepository(session_maker)
project_repository = ProjectRepository()
project1_data = {
"name": "conflict-project-1",
@@ -414,77 +433,80 @@ async def test_upsert_entity_permalink_conflict_within_project_only(session_make
"is_active": True,
"is_default": False,
}
project1 = await project_repository.create(project1_data)
async with db.scoped_session(session_maker) as session:
project1 = await project_repository.create(session, project1_data)
project2_data = {
"name": "conflict-project-2",
"description": "Second conflict test project",
"path": "/tmp/conflict2",
"is_active": True,
"is_default": False,
}
project2 = await project_repository.create(project2_data)
project2_data = {
"name": "conflict-project-2",
"description": "Second conflict test project",
"path": "/tmp/conflict2",
"is_active": True,
"is_default": False,
}
project2 = await project_repository.create(session, project2_data)
# Create entity repositories for each project
repo1 = EntityRepository(session_maker, project_id=project1.id)
repo2 = EntityRepository(session_maker, project_id=project2.id)
# Create entity repositories for each project
repo1 = EntityRepository(project_id=project1.id)
repo2 = EntityRepository(project_id=project2.id)
# Create first entity in project1
entity1 = Entity(
project_id=project1.id,
title="Original Entity",
note_type="note",
permalink="test/conflict-permalink",
file_path="test/original.md",
content_type="text/markdown",
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
)
# Create first entity in project1
entity1 = Entity(
project_id=project1.id,
title="Original Entity",
note_type="note",
permalink="test/conflict-permalink",
file_path="test/original.md",
content_type="text/markdown",
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
)
result1 = await repo1.upsert_entity(entity1)
assert result1.permalink == "test/conflict-permalink"
result1 = await repo1.upsert_entity(session, entity1)
assert result1.permalink == "test/conflict-permalink"
# Create entity with same permalink in project2 (should NOT get suffix)
entity2 = Entity(
project_id=project2.id,
title="Cross-Project Entity",
note_type="note",
permalink="test/conflict-permalink", # Same permalink, different project
file_path="test/cross-project.md",
content_type="text/markdown",
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
)
# Create entity with same permalink in project2 (should NOT get suffix)
entity2 = Entity(
project_id=project2.id,
title="Cross-Project Entity",
note_type="note",
permalink="test/conflict-permalink", # Same permalink, different project
file_path="test/cross-project.md",
content_type="text/markdown",
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
)
result2 = await repo2.upsert_entity(entity2)
# Should keep original permalink (no suffix) since it's in a different project
assert result2.permalink == "test/conflict-permalink"
result2 = await repo2.upsert_entity(session, entity2)
# Should keep original permalink (no suffix) since it's in a different project
assert result2.permalink == "test/conflict-permalink"
# Now create entity with same permalink in project1 (should get suffix)
entity3 = Entity(
project_id=project1.id,
title="Conflict Entity",
note_type="note",
permalink="test/conflict-permalink", # Same permalink, same project
file_path="test/conflict.md",
content_type="text/markdown",
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
)
# Now create entity with same permalink in project1 (should get suffix)
entity3 = Entity(
project_id=project1.id,
title="Conflict Entity",
note_type="note",
permalink="test/conflict-permalink", # Same permalink, same project
file_path="test/conflict.md",
content_type="text/markdown",
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
)
result3 = await repo1.upsert_entity(entity3)
# Should get suffix since it conflicts within the same project
assert result3.permalink == "test/conflict-permalink-1"
result3 = await repo1.upsert_entity(session, entity3)
# Should get suffix since it conflicts within the same project
assert result3.permalink == "test/conflict-permalink-1"
# Verify all entities exist correctly
assert result1.id != result2.id != result3.id
assert result1.project_id == project1.id
assert result2.project_id == project2.id
assert result3.project_id == project1.id
# Verify all entities exist correctly
assert result1.id != result2.id != result3.id
assert result1.project_id == project1.id
assert result2.project_id == project2.id
assert result3.project_id == project1.id
@pytest.mark.asyncio
async def test_upsert_entity_with_invalid_project_id(entity_repository: EntityRepository):
async def test_upsert_entity_with_invalid_project_id(
entity_repository: EntityRepository, session_maker
):
"""Test that upserting with non-existent project_id raises clear error.
This tests the fix for issue #188 where sync fails with FOREIGN KEY constraint
@@ -504,7 +526,8 @@ async def test_upsert_entity_with_invalid_project_id(entity_repository: EntityRe
# Should raise SyncFatalError with clear message about missing project
with pytest.raises(SyncFatalError) as exc_info:
await entity_repository.upsert_entity(entity)
async with db.scoped_session(session_maker) as session:
await entity_repository.upsert_entity(session, entity)
error_msg = str(exc_info.value)
assert "project_id=99999 does not exist" in error_msg
@@ -3,12 +3,15 @@
import pytest
from datetime import datetime, timezone
from basic_memory import db
from basic_memory.models.knowledge import Entity, Observation
from basic_memory.repository.entity_repository import EntityRepository
@pytest.mark.asyncio
async def test_upsert_entity_with_observations_conflict(entity_repository: EntityRepository):
async def test_upsert_entity_with_observations_conflict(
entity_repository: EntityRepository, session_maker
):
"""Test upserting an entity that already exists with observations.
This reproduces issue #187 where sync fails with UNIQUE constraint violations
@@ -35,43 +38,44 @@ async def test_upsert_entity_with_observations_conflict(entity_repository: Entit
)
entity1.observations.append(obs1)
result1 = await entity_repository.upsert_entity(entity1)
original_id = result1.id
async with db.scoped_session(session_maker) as session:
result1 = await entity_repository.upsert_entity(session, entity1)
original_id = result1.id
# Verify entity was created with observations
assert result1.id is not None
assert len(result1.observations) == 1
# Verify entity was created with observations
assert result1.id is not None
assert len(result1.observations) == 1
# Now try to upsert the same file_path with different content/observations
# This simulates a file being modified and re-synced
entity2 = Entity(
project_id=entity_repository.project_id,
title="Updated Title",
note_type="note",
permalink="debugging/backup-system/coderabbit-feedback-resolution", # Same permalink
file_path="debugging/backup-system/CodeRabbit Feedback Resolution - Backup System Issues.md", # Same file_path
content_type="text/markdown",
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
)
# Now try to upsert the same file_path with different content/observations
# This simulates a file being modified and re-synced
entity2 = Entity(
project_id=entity_repository.project_id,
title="Updated Title",
note_type="note",
permalink="debugging/backup-system/coderabbit-feedback-resolution", # Same permalink
file_path="debugging/backup-system/CodeRabbit Feedback Resolution - Backup System Issues.md", # Same file_path
content_type="text/markdown",
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
)
# Add different observations
obs2 = Observation(
project_id=entity_repository.project_id,
content="This is an updated observation",
category="updated",
tags=["updated"],
)
obs3 = Observation(
project_id=entity_repository.project_id,
content="This is a second observation",
category="second",
tags=["second"],
)
entity2.observations.extend([obs2, obs3])
# Add different observations
obs2 = Observation(
project_id=entity_repository.project_id,
content="This is an updated observation",
category="updated",
tags=["updated"],
)
obs3 = Observation(
project_id=entity_repository.project_id,
content="This is a second observation",
category="second",
tags=["second"],
)
entity2.observations.extend([obs2, obs3])
# This should UPDATE the existing entity, not fail with IntegrityError
result2 = await entity_repository.upsert_entity(entity2)
# This should UPDATE the existing entity, not fail with IntegrityError
result2 = await entity_repository.upsert_entity(session, entity2)
# Should update existing entity (same ID)
assert result2.id == original_id
@@ -86,7 +90,9 @@ async def test_upsert_entity_with_observations_conflict(entity_repository: Entit
@pytest.mark.asyncio
async def test_upsert_entity_repeated_sync_same_file(entity_repository: EntityRepository):
async def test_upsert_entity_repeated_sync_same_file(
entity_repository: EntityRepository, session_maker
):
"""Test that syncing the same file multiple times doesn't cause IntegrityError.
This tests the specific scenario from issue #187 where files are being
@@ -108,28 +114,29 @@ async def test_upsert_entity_repeated_sync_same_file(entity_repository: EntityRe
updated_at=datetime.now(timezone.utc),
)
result1 = await entity_repository.upsert_entity(entity1)
first_id = result1.id
async with db.scoped_session(session_maker) as session:
result1 = await entity_repository.upsert_entity(session, entity1)
first_id = result1.id
# Simulate multiple sync attempts (like the infinite retry loop in the issue)
for i in range(5):
entity_new = Entity(
project_id=entity_repository.project_id,
title="Complete Process for Uploading New Training Videos",
note_type="note",
permalink=permalink,
file_path=file_path,
content_type="text/markdown",
checksum=f"def{456 + i}", # Different checksum each time
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
)
# Simulate multiple sync attempts (like the infinite retry loop in the issue)
for i in range(5):
entity_new = Entity(
project_id=entity_repository.project_id,
title="Complete Process for Uploading New Training Videos",
note_type="note",
permalink=permalink,
file_path=file_path,
content_type="text/markdown",
checksum=f"def{456 + i}", # Different checksum each time
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
)
# Each upsert should succeed and update the existing entity
result = await entity_repository.upsert_entity(entity_new)
# Each upsert should succeed and update the existing entity
result = await entity_repository.upsert_entity(session, entity_new)
# Should always return the same entity (updated)
assert result.id == first_id
assert result.checksum == entity_new.checksum
assert result.file_path == file_path
assert result.permalink == permalink
# Should always return the same entity (updated)
assert result.id == first_id
assert result.checksum == entity_new.checksum
assert result.file_path == file_path
assert result.permalink == permalink
@@ -0,0 +1,163 @@
"""Regression coverage for caller-owned repository transactions."""
from datetime import datetime, timezone
import pytest
from sqlalchemy import select
from basic_memory.models import Entity, NoteContent, Observation, Project, Relation
from basic_memory.repository.entity_repository import EntityRepository
from basic_memory.repository.note_content_repository import NoteContentRepository
from basic_memory.repository.observation_repository import ObservationRepository
from basic_memory.repository.project_repository import ProjectRepository
from basic_memory.repository.relation_repository import RelationRepository
def _entity_payload(project_id: int, title: str, file_path: str) -> dict:
"""Build the minimal entity payload used by transaction tests."""
return {
"project_id": project_id,
"title": title,
"note_type": "note",
"permalink": file_path.removesuffix(".md"),
"file_path": file_path,
"content_type": "text/markdown",
"created_at": datetime.now(timezone.utc),
"updated_at": datetime.now(timezone.utc),
}
def _note_content_payload(entity_id: int) -> dict:
"""Build the minimal note_content payload used by transaction tests."""
return {
"entity_id": entity_id,
"markdown_content": "# Draft content",
"db_version": 1,
"db_checksum": "db-checksum",
"file_write_status": "pending",
"last_source": "api",
"updated_at": datetime.now(timezone.utc),
}
@pytest.mark.asyncio
async def test_repositories_share_uncommitted_state_in_one_session(session_maker, config_home):
"""Later repository calls should see earlier uncommitted writes in the same session."""
project_repository = ProjectRepository()
async with session_maker() as session:
project = await project_repository.create(
session,
{
"name": "explicit-session-project",
"path": str(config_home / "explicit-session-project"),
"is_active": True,
},
)
entity_repository = EntityRepository(project_id=project.id)
note_content_repository = NoteContentRepository(project_id=project.id)
observation_repository = ObservationRepository(project_id=project.id)
relation_repository = RelationRepository(project_id=project.id)
source = await entity_repository.create(
session, _entity_payload(project.id, "Source Note", "notes/source.md")
)
target = await entity_repository.create(
session, _entity_payload(project.id, "Target Note", "notes/target.md")
)
note_content = await note_content_repository.upsert(
session, _note_content_payload(source.id)
)
observation = await observation_repository.create(
session,
{
"entity_id": source.id,
"content": "Uncommitted observation",
"category": "note",
},
)
relation = await relation_repository.create(
session,
{
"from_id": source.id,
"to_id": target.id,
"to_name": target.title,
"relation_type": "links_to",
},
)
found_project = await project_repository.get_by_external_id(session, project.external_id)
found_source = await entity_repository.get_by_file_path(session, "notes/source.md")
found_note_content = await note_content_repository.get_by_entity_id(session, source.id)
found_observations = await observation_repository.find_by_entity(session, source.id)
found_relations = await relation_repository.find_by_entities(session, source.id, target.id)
assert found_project is not None
assert found_project.id == project.id
assert found_source is not None
assert found_source.id == source.id
assert found_note_content is not None
assert found_note_content.entity_id == note_content.entity_id
assert [item.id for item in found_observations] == [observation.id]
assert [item.id for item in found_relations] == [relation.id]
@pytest.mark.asyncio
async def test_rollback_discards_related_repository_writes(session_maker, config_home):
"""A caller-owned rollback should discard all related repository writes."""
project_repository = ProjectRepository()
async with session_maker() as session:
transaction = await session.begin()
project = await project_repository.create(
session,
{
"name": "rollback-explicit-session",
"path": str(config_home / "rollback-explicit-session"),
"is_active": True,
},
)
entity_repository = EntityRepository(project_id=project.id)
note_content_repository = NoteContentRepository(project_id=project.id)
observation_repository = ObservationRepository(project_id=project.id)
relation_repository = RelationRepository(project_id=project.id)
source = await entity_repository.create(
session, _entity_payload(project.id, "Rollback Source", "rollback/source.md")
)
target = await entity_repository.create(
session, _entity_payload(project.id, "Rollback Target", "rollback/target.md")
)
await note_content_repository.upsert(session, _note_content_payload(source.id))
await observation_repository.create(
session,
{
"entity_id": source.id,
"content": "Rolled back observation",
"category": "note",
},
)
await relation_repository.create(
session,
{
"from_id": source.id,
"to_id": target.id,
"to_name": target.title,
"relation_type": "links_to",
},
)
await transaction.rollback()
async with session_maker() as session:
assert (
await session.scalar(select(Project).where(Project.name == "rollback-explicit-session"))
is None
)
assert (
await session.scalar(select(Entity).where(Entity.file_path == "rollback/source.md"))
is None
)
assert await session.scalar(select(NoteContent)) is None
assert await session.scalar(select(Observation)) is None
assert await session.scalar(select(Relation)) is None
@@ -595,6 +595,7 @@ async def test_fastembed_provider_fails_fast_without_cache_dir(monkeypatch):
assert _SelfHealStubTextEmbedding.construct_count == 1
@pytest.mark.asyncio
async def test_factory_loads_native_model_once_across_repo_constructions(monkeypatch):
"""The native ONNX model must load exactly once per process despite reuse (#872).
+1
View File
@@ -77,6 +77,7 @@ class ConcreteSearchRepo(SearchRepositoryBase):
min_similarity: Optional[float] = None,
limit: int = 10,
offset: int = 0,
allow_relaxed: bool = False,
) -> list[SearchIndexRow]:
return [] # pragma: no cover
+247 -212
View File
@@ -39,18 +39,19 @@ async def test_create_and_lookup_note_content(
sample_entity,
):
"""Create note_content and read it back through each supported lookup."""
repository = NoteContentRepository(session_maker, project_id=test_project.id)
repository = NoteContentRepository(project_id=test_project.id)
created = await repository.create(build_note_content_payload(sample_entity.id))
async with db.scoped_session(session_maker) as session:
created = await repository.create(session, build_note_content_payload(sample_entity.id))
assert created.entity_id == sample_entity.id
assert created.project_id == sample_entity.project_id
assert created.external_id == sample_entity.external_id
assert created.file_path == sample_entity.file_path
assert created.entity_id == sample_entity.id
assert created.project_id == sample_entity.project_id
assert created.external_id == sample_entity.external_id
assert created.file_path == sample_entity.file_path
by_entity = await repository.get_by_entity_id(sample_entity.id)
by_external = await repository.get_by_external_id(sample_entity.external_id)
by_path = await repository.get_by_file_path(sample_entity.file_path)
by_entity = await repository.get_by_entity_id(session, sample_entity.id)
by_external = await repository.get_by_external_id(session, sample_entity.external_id)
by_path = await repository.get_by_file_path(session, sample_entity.file_path)
assert by_entity is not None
assert by_external is not None
@@ -67,39 +68,42 @@ async def test_upsert_updates_existing_note_content(
sample_entity,
):
"""Upsert should update the existing row instead of inserting a duplicate."""
repository = NoteContentRepository(session_maker, project_id=test_project.id)
await repository.create(build_note_content_payload(sample_entity.id))
repository = NoteContentRepository(project_id=test_project.id)
updated_at = datetime.now(timezone.utc)
updated = await repository.upsert(
NoteContent(
entity_id=sample_entity.id,
project_id=test_project.id,
external_id=sample_entity.external_id,
file_path=sample_entity.file_path,
markdown_content="# Updated materialized content",
db_version=2,
db_checksum="db-checksum-2",
file_version=7,
file_checksum="file-checksum-7",
file_write_status="synced",
last_source="reconciler",
updated_at=updated_at,
file_updated_at=updated_at,
last_materialization_error="transient failure",
last_materialization_attempt_at=updated_at,
async with db.scoped_session(session_maker) as session:
await repository.create(session, build_note_content_payload(sample_entity.id))
updated_at = datetime.now(timezone.utc)
updated = await repository.upsert(
session,
NoteContent(
entity_id=sample_entity.id,
project_id=test_project.id,
external_id=sample_entity.external_id,
file_path=sample_entity.file_path,
markdown_content="# Updated materialized content",
db_version=2,
db_checksum="db-checksum-2",
file_version=7,
file_checksum="file-checksum-7",
file_write_status="synced",
last_source="reconciler",
updated_at=updated_at,
file_updated_at=updated_at,
last_materialization_error="transient failure",
last_materialization_attempt_at=updated_at,
),
)
)
assert updated.entity_id == sample_entity.id
assert updated.markdown_content == "# Updated materialized content"
assert updated.db_version == 2
assert updated.db_checksum == "db-checksum-2"
assert updated.file_version == 7
assert updated.file_checksum == "file-checksum-7"
assert updated.file_write_status == "synced"
assert updated.last_source == "reconciler"
assert updated.last_materialization_error == "transient failure"
assert updated.entity_id == sample_entity.id
assert updated.markdown_content == "# Updated materialized content"
assert updated.db_version == 2
assert updated.db_checksum == "db-checksum-2"
assert updated.file_version == 7
assert updated.file_checksum == "file-checksum-7"
assert updated.file_write_status == "synced"
assert updated.last_source == "reconciler"
assert updated.last_materialization_error == "transient failure"
@pytest.mark.asyncio
@@ -109,9 +113,10 @@ async def test_upsert_inserts_when_no_existing_row(
sample_entity,
):
"""Upsert should insert a new row when the entity has no note_content yet."""
repository = NoteContentRepository(session_maker, project_id=test_project.id)
repository = NoteContentRepository(project_id=test_project.id)
created = await repository.upsert(build_note_content_payload(sample_entity.id))
async with db.scoped_session(session_maker) as session:
created = await repository.upsert(session, build_note_content_payload(sample_entity.id))
assert created.entity_id == sample_entity.id
assert created.project_id == sample_entity.project_id
@@ -123,10 +128,11 @@ async def test_upsert_inserts_when_no_existing_row(
@pytest.mark.asyncio
async def test_create_requires_entity_id(session_maker, test_project: Project):
"""Create should fail fast when note_content identity is missing."""
repository = NoteContentRepository(session_maker, project_id=test_project.id)
repository = NoteContentRepository(project_id=test_project.id)
with pytest.raises(ValueError, match="entity_id is required"):
await repository.create({"markdown_content": "# Missing entity"})
async with db.scoped_session(session_maker) as session:
await repository.create(session, {"markdown_content": "# Missing entity"})
@pytest.mark.asyncio
@@ -136,78 +142,86 @@ async def test_upsert_preserves_existing_fields_for_partial_payload(
sample_entity,
):
"""Partial upserts should only change explicit fields and preserve existing state."""
repository = NoteContentRepository(session_maker, project_id=test_project.id)
repository = NoteContentRepository(project_id=test_project.id)
payload = build_note_content_payload(sample_entity.id)
payload["last_materialization_error"] = "stale failure"
created = await repository.create(payload)
async with db.scoped_session(session_maker) as session:
created = await repository.create(session, payload)
updated_at = datetime.now(timezone.utc)
updated = await repository.upsert(
{
"entity_id": sample_entity.id,
"markdown_content": "# Partially updated content",
"db_version": 2,
"updated_at": updated_at,
"last_materialization_error": None,
}
)
updated_at = datetime.now(timezone.utc)
updated = await repository.upsert(
session,
{
"entity_id": sample_entity.id,
"markdown_content": "# Partially updated content",
"db_version": 2,
"updated_at": updated_at,
"last_materialization_error": None,
},
)
assert updated.markdown_content == "# Partially updated content"
assert updated.db_version == 2
assert updated.db_checksum == created.db_checksum
assert updated.file_write_status == created.file_write_status
assert updated.last_source == created.last_source
assert updated.last_materialization_error is None
assert updated.file_path == sample_entity.file_path
assert updated.markdown_content == "# Partially updated content"
assert updated.db_version == 2
assert updated.db_checksum == created.db_checksum
assert updated.file_write_status == created.file_write_status
assert updated.last_source == created.last_source
assert updated.last_materialization_error is None
assert updated.file_path == sample_entity.file_path
@pytest.mark.asyncio
async def test_create_rejects_missing_entity(session_maker, test_project: Project):
"""Create should fail when the owning entity does not exist."""
repository = NoteContentRepository(session_maker, project_id=test_project.id)
repository = NoteContentRepository(project_id=test_project.id)
with pytest.raises(ValueError, match="Entity 999999 does not exist"):
await repository.create(build_note_content_payload(999999))
async with db.scoped_session(session_maker) as session:
await repository.create(session, build_note_content_payload(999999))
@pytest.mark.asyncio
async def test_create_rejects_entity_from_another_project(session_maker, config_home):
"""Create should reject note_content writes across project boundaries."""
project_repository = ProjectRepository(session_maker)
project_one = await project_repository.create(
{
"name": "project-one-boundary",
"path": str(config_home / "project-one-boundary"),
"is_active": True,
}
)
project_two = await project_repository.create(
{
"name": "project-two-boundary",
"path": str(config_home / "project-two-boundary"),
"is_active": True,
}
)
entity_repository = EntityRepository(session_maker, project_id=project_two.id)
other_project_entity = await entity_repository.create(
{
"title": "Other Project Note",
"note_type": "test",
"permalink": "project-two/other-project-note",
"file_path": "notes/other-project-note.md",
"content_type": "text/markdown",
"created_at": datetime.now(timezone.utc),
"updated_at": datetime.now(timezone.utc),
}
)
project_repository = ProjectRepository()
repository = NoteContentRepository(session_maker, project_id=project_one.id)
async with db.scoped_session(session_maker) as session:
project_one = await project_repository.create(
session,
{
"name": "project-one-boundary",
"path": str(config_home / "project-one-boundary"),
"is_active": True,
},
)
project_two = await project_repository.create(
session,
{
"name": "project-two-boundary",
"path": str(config_home / "project-two-boundary"),
"is_active": True,
},
)
entity_repository = EntityRepository(project_id=project_two.id)
other_project_entity = await entity_repository.create(
session,
{
"title": "Other Project Note",
"note_type": "test",
"permalink": "project-two/other-project-note",
"file_path": "notes/other-project-note.md",
"content_type": "text/markdown",
"created_at": datetime.now(timezone.utc),
"updated_at": datetime.now(timezone.utc),
},
)
with pytest.raises(
ValueError,
match=f"Entity {other_project_entity.id} belongs to project {project_two.id}",
):
await repository.create(build_note_content_payload(other_project_entity.id))
repository = NoteContentRepository(project_id=project_one.id)
with pytest.raises(
ValueError,
match=f"Entity {other_project_entity.id} belongs to project {project_two.id}",
):
await repository.create(session, build_note_content_payload(other_project_entity.id))
@pytest.mark.asyncio
@@ -218,20 +232,22 @@ async def test_update_state_fields_realigns_identity_with_entity(
entity_repository: EntityRepository,
):
"""Sync-field updates should refresh mirrored identity from the owning entity."""
repository = NoteContentRepository(session_maker, project_id=test_project.id)
await repository.create(build_note_content_payload(sample_entity.id))
repository = NoteContentRepository(project_id=test_project.id)
renamed_path = "renamed/test_entity.md"
await entity_repository.update(sample_entity.id, {"file_path": renamed_path})
async with db.scoped_session(session_maker) as session:
await repository.create(session, build_note_content_payload(sample_entity.id))
await entity_repository.update(session, sample_entity.id, {"file_path": renamed_path})
updated = await repository.update_state_fields(
sample_entity.id,
file_write_status="failed",
file_version=3,
file_checksum="file-checksum-3",
last_materialization_error=None,
last_materialization_attempt_at=None,
)
updated = await repository.update_state_fields(
session,
sample_entity.id,
file_write_status="failed",
file_version=3,
file_checksum="file-checksum-3",
last_materialization_error=None,
last_materialization_attempt_at=None,
)
assert updated is not None
assert updated.file_path == renamed_path
@@ -250,11 +266,14 @@ async def test_update_state_fields_rejects_invalid_fields(
sample_entity,
):
"""Only the declared mutable sync fields should be accepted."""
repository = NoteContentRepository(session_maker, project_id=test_project.id)
await repository.create(build_note_content_payload(sample_entity.id))
repository = NoteContentRepository(project_id=test_project.id)
async with db.scoped_session(session_maker) as session:
await repository.create(session, build_note_content_payload(sample_entity.id))
with pytest.raises(ValueError, match="Unsupported note_content update fields: file_path"):
await repository.update_state_fields(sample_entity.id, file_path="renamed/note.md")
with pytest.raises(ValueError, match="Unsupported note_content update fields: file_path"):
await repository.update_state_fields(
session, sample_entity.id, file_path="renamed/note.md"
)
@pytest.mark.asyncio
@@ -263,21 +282,26 @@ async def test_update_state_fields_returns_none_for_missing_note_content(
test_project: Project,
):
"""Missing note_content rows should produce a clean None response."""
repository = NoteContentRepository(session_maker, project_id=test_project.id)
repository = NoteContentRepository(project_id=test_project.id)
assert await repository.update_state_fields(999999, file_write_status="failed") is None
async with db.scoped_session(session_maker) as session:
assert (
await repository.update_state_fields(session, 999999, file_write_status="failed")
is None
)
@pytest.mark.asyncio
async def test_delete_by_entity_id(session_maker, test_project: Project, sample_entity):
"""Delete note_content directly by entity identifier."""
repository = NoteContentRepository(session_maker, project_id=test_project.id)
await repository.create(build_note_content_payload(sample_entity.id))
repository = NoteContentRepository(project_id=test_project.id)
async with db.scoped_session(session_maker) as session:
await repository.create(session, build_note_content_payload(sample_entity.id))
deleted = await repository.delete_by_entity_id(sample_entity.id)
deleted = await repository.delete_by_entity_id(session, sample_entity.id)
assert deleted is True
assert await repository.get_by_entity_id(sample_entity.id) is None
assert deleted is True
assert await repository.get_by_entity_id(session, sample_entity.id) is None
@pytest.mark.asyncio
@@ -286,9 +310,10 @@ async def test_delete_by_entity_id_returns_false_when_missing(
test_project: Project,
):
"""Delete should report False when the note_content row does not exist."""
repository = NoteContentRepository(session_maker, project_id=test_project.id)
repository = NoteContentRepository(project_id=test_project.id)
assert await repository.delete_by_entity_id(999999) is False
async with db.scoped_session(session_maker) as session:
assert await repository.delete_by_entity_id(session, 999999) is False
@pytest.mark.asyncio
@@ -299,68 +324,75 @@ async def test_note_content_cascades_when_entity_is_deleted(
entity_repository: EntityRepository,
):
"""Deleting the owning entity should cascade to note_content."""
repository = NoteContentRepository(session_maker, project_id=test_project.id)
await repository.create(build_note_content_payload(sample_entity.id))
repository = NoteContentRepository(project_id=test_project.id)
deleted = await entity_repository.delete(sample_entity.id)
async with db.scoped_session(session_maker) as session:
await repository.create(session, build_note_content_payload(sample_entity.id))
deleted = await entity_repository.delete(session, sample_entity.id)
assert deleted is True
assert await repository.get_by_entity_id(sample_entity.id) is None
assert deleted is True
assert await repository.get_by_entity_id(session, sample_entity.id) is None
@pytest.mark.asyncio
async def test_note_content_file_path_lookup_is_project_scoped(session_maker, config_home):
"""Lookups by file_path should respect the repository project scope."""
project_repository = ProjectRepository(session_maker)
project_one = await project_repository.create(
{
"name": "project-one",
"path": str(config_home / "project-one"),
"is_active": True,
}
)
project_two = await project_repository.create(
{
"name": "project-two",
"path": str(config_home / "project-two"),
"is_active": True,
}
)
entity_one_repo = EntityRepository(session_maker, project_id=project_one.id)
entity_two_repo = EntityRepository(session_maker, project_id=project_two.id)
project_repository = ProjectRepository()
shared_file_path = "shared/note.md"
entity_one = await entity_one_repo.create(
{
"title": "Shared Note",
"note_type": "test",
"permalink": "project-one/shared-note",
"file_path": shared_file_path,
"content_type": "text/markdown",
"created_at": datetime.now(timezone.utc),
"updated_at": datetime.now(timezone.utc),
}
)
entity_two = await entity_two_repo.create(
{
"title": "Shared Note",
"note_type": "test",
"permalink": "project-two/shared-note",
"file_path": shared_file_path,
"content_type": "text/markdown",
"created_at": datetime.now(timezone.utc),
"updated_at": datetime.now(timezone.utc),
}
)
async with db.scoped_session(session_maker) as session:
project_one = await project_repository.create(
session,
{
"name": "project-one",
"path": str(config_home / "project-one"),
"is_active": True,
},
)
project_two = await project_repository.create(
session,
{
"name": "project-two",
"path": str(config_home / "project-two"),
"is_active": True,
},
)
repository_one = NoteContentRepository(session_maker, project_id=project_one.id)
repository_two = NoteContentRepository(session_maker, project_id=project_two.id)
await repository_one.create(build_note_content_payload(entity_one.id))
await repository_two.create(build_note_content_payload(entity_two.id))
entity_one_repo = EntityRepository(project_id=project_one.id)
entity_two_repo = EntityRepository(project_id=project_two.id)
found_one = await repository_one.get_by_file_path(shared_file_path)
found_two = await repository_two.get_by_file_path(shared_file_path)
entity_one = await entity_one_repo.create(
session,
{
"title": "Shared Note",
"note_type": "test",
"permalink": "project-one/shared-note",
"file_path": shared_file_path,
"content_type": "text/markdown",
"created_at": datetime.now(timezone.utc),
"updated_at": datetime.now(timezone.utc),
},
)
entity_two = await entity_two_repo.create(
session,
{
"title": "Shared Note",
"note_type": "test",
"permalink": "project-two/shared-note",
"file_path": shared_file_path,
"content_type": "text/markdown",
"created_at": datetime.now(timezone.utc),
"updated_at": datetime.now(timezone.utc),
},
)
repository_one = NoteContentRepository(project_id=project_one.id)
repository_two = NoteContentRepository(project_id=project_two.id)
await repository_one.create(session, build_note_content_payload(entity_one.id))
await repository_two.create(session, build_note_content_payload(entity_two.id))
found_one = await repository_one.get_by_file_path(session, shared_file_path)
found_two = await repository_two.get_by_file_path(session, shared_file_path)
assert found_one is not None
assert found_two is not None
@@ -374,52 +406,55 @@ async def test_note_content_file_path_lookup_prefers_entity_with_current_path(
config_home,
):
"""File-path lookup should prefer the entity whose current path still matches."""
project_repository = ProjectRepository(session_maker)
project = await project_repository.create(
{
"name": "project-path-drift",
"path": str(config_home / "project-path-drift"),
"is_active": True,
}
)
entity_repository = EntityRepository(session_maker, project_id=project.id)
stale_entity = await entity_repository.create(
{
"title": "Stale Note",
"note_type": "test",
"permalink": "project/stale-note",
"file_path": "archived/note.md",
"content_type": "text/markdown",
"created_at": datetime.now(timezone.utc),
"updated_at": datetime.now(timezone.utc),
}
)
current_entity = await entity_repository.create(
{
"title": "Current Note",
"note_type": "test",
"permalink": "project/current-note",
"file_path": "shared/note.md",
"content_type": "text/markdown",
"created_at": datetime.now(timezone.utc),
"updated_at": datetime.now(timezone.utc),
}
)
repository = NoteContentRepository(session_maker, project_id=project.id)
stale_payload = build_note_content_payload(stale_entity.id)
stale_payload["updated_at"] = datetime.now(timezone.utc) + timedelta(minutes=5)
await repository.create(stale_payload)
await repository.create(build_note_content_payload(current_entity.id))
project_repository = ProjectRepository()
async with db.scoped_session(session_maker) as session:
project = await project_repository.create(
session,
{
"name": "project-path-drift",
"path": str(config_home / "project-path-drift"),
"is_active": True,
},
)
entity_repository = EntityRepository(project_id=project.id)
stale_entity = await entity_repository.create(
session,
{
"title": "Stale Note",
"note_type": "test",
"permalink": "project/stale-note",
"file_path": "archived/note.md",
"content_type": "text/markdown",
"created_at": datetime.now(timezone.utc),
"updated_at": datetime.now(timezone.utc),
},
)
current_entity = await entity_repository.create(
session,
{
"title": "Current Note",
"note_type": "test",
"permalink": "project/current-note",
"file_path": "shared/note.md",
"content_type": "text/markdown",
"created_at": datetime.now(timezone.utc),
"updated_at": datetime.now(timezone.utc),
},
)
repository = NoteContentRepository(project_id=project.id)
stale_payload = build_note_content_payload(stale_entity.id)
stale_payload["updated_at"] = datetime.now(timezone.utc) + timedelta(minutes=5)
await repository.create(session, stale_payload)
await repository.create(session, build_note_content_payload(current_entity.id))
stale_note_content = await repository.select_by_id(session, stale_entity.id)
assert stale_note_content is not None
stale_note_content.file_path = "shared/note.md"
await session.flush()
found = await repository.get_by_file_path("shared/note.md")
found = await repository.get_by_file_path(session, "shared/note.md")
assert found is not None
assert found.entity_id == current_entity.id
+64 -48
View File
@@ -19,7 +19,7 @@ async def repo(observation_repository):
@pytest_asyncio.fixture(scope="function")
async def sample_observation(repo, sample_entity: Entity):
async def sample_observation(repo, sample_entity: Entity, session_maker):
"""Create a sample observation for testing"""
observation_data = {
"project_id": sample_entity.project_id,
@@ -27,12 +27,13 @@ async def sample_observation(repo, sample_entity: Entity):
"content": "Test observation",
"context": "test-context",
}
return await repo.create(observation_data)
async with db.scoped_session(session_maker) as session:
return await repo.create(session, observation_data)
@pytest.mark.asyncio
async def test_create_observation(
observation_repository: ObservationRepository, sample_entity: Entity
observation_repository: ObservationRepository, sample_entity: Entity, session_maker
):
"""Test creating a new observation"""
observation_data = {
@@ -41,7 +42,8 @@ async def test_create_observation(
"content": "Test content",
"context": "test-context",
}
observation = await observation_repository.create(observation_data)
async with db.scoped_session(session_maker) as session:
observation = await observation_repository.create(session, observation_data)
assert observation.entity_id == sample_entity.id
assert observation.content == "Test content"
@@ -50,7 +52,7 @@ async def test_create_observation(
@pytest.mark.asyncio
async def test_create_observation_entity_does_not_exist(
observation_repository: ObservationRepository, sample_entity: Entity
observation_repository: ObservationRepository, sample_entity: Entity, session_maker
):
"""Test creating a new observation"""
observation_data = {
@@ -60,7 +62,8 @@ async def test_create_observation_entity_does_not_exist(
"context": "test-context",
}
with pytest.raises(IntegrityError):
await observation_repository.create(observation_data)
async with db.scoped_session(session_maker) as session:
await observation_repository.create(session, observation_data)
@pytest.mark.asyncio
@@ -68,9 +71,11 @@ async def test_find_by_entity(
observation_repository: ObservationRepository,
sample_observation: Observation,
sample_entity: Entity,
session_maker,
):
"""Test finding observations by entity"""
observations = await observation_repository.find_by_entity(sample_entity.id)
async with db.scoped_session(session_maker) as session:
observations = await observation_repository.find_by_entity(session, sample_entity.id)
assert len(observations) == 1
assert observations[0].id == sample_observation.id
assert observations[0].content == sample_observation.content
@@ -78,10 +83,11 @@ async def test_find_by_entity(
@pytest.mark.asyncio
async def test_find_by_context(
observation_repository: ObservationRepository, sample_observation: Observation
observation_repository: ObservationRepository, sample_observation: Observation, session_maker
):
"""Test finding observations by context"""
observations = await observation_repository.find_by_context("test-context")
async with db.scoped_session(session_maker) as session:
observations = await observation_repository.find_by_context(session, "test-context")
assert len(observations) == 1
assert observations[0].id == sample_observation.id
assert observations[0].content == sample_observation.content
@@ -119,12 +125,13 @@ async def test_delete_observations(session_maker: async_sessionmaker, repo, test
session.add_all([obs1, obs2])
# Test deletion by entity_id
deleted = await repo.delete_by_fields(entity_id=entity.id)
assert deleted is True
async with db.scoped_session(session_maker) as session:
deleted = await repo.delete_by_fields(session, entity_id=entity.id)
assert deleted is True
# Verify observations were deleted
remaining = await repo.find_by_entity(entity.id)
assert len(remaining) == 0
# Verify observations were deleted
remaining = await repo.find_by_entity(session, entity.id)
assert len(remaining) == 0
@pytest.mark.asyncio
@@ -156,12 +163,13 @@ async def test_delete_observation_by_id(
session.add(obs)
# Test deletion by ID
deleted = await repo.delete(obs.id)
assert deleted is True
async with db.scoped_session(session_maker) as session:
deleted = await repo.delete(session, obs.id)
assert deleted is True
# Verify observation was deleted
remaining = await repo.find_by_id(obs.id)
assert remaining is None
# Verify observation was deleted
remaining = await repo.find_by_id(session, obs.id)
assert remaining is None
@pytest.mark.asyncio
@@ -198,13 +206,14 @@ async def test_delete_observation_by_content(
session.add_all([obs1, obs2])
# Test deletion by content
deleted = await repo.delete_by_fields(content="Delete this observation")
assert deleted is True
async with db.scoped_session(session_maker) as session:
deleted = await repo.delete_by_fields(session, content="Delete this observation")
assert deleted is True
# Verify only matching observation was deleted
remaining = await repo.find_by_entity(entity.id)
assert len(remaining) == 1
assert remaining[0].content == "Keep this observation"
# Verify only matching observation was deleted
remaining = await repo.find_by_entity(session, entity.id)
assert len(remaining) == 1
assert remaining[0].content == "Keep this observation"
@pytest.mark.asyncio
@@ -250,20 +259,24 @@ async def test_find_by_category(session_maker: async_sessionmaker, repo, test_pr
await session.commit()
# Find tech observations
tech_obs = await repo.find_by_category("tech")
assert len(tech_obs) == 2
assert all(obs.category == "tech" for obs in tech_obs)
assert set(obs.content for obs in tech_obs) == {"Tech observation", "Another tech observation"}
async with db.scoped_session(session_maker) as session:
tech_obs = await repo.find_by_category(session, "tech")
assert len(tech_obs) == 2
assert all(obs.category == "tech" for obs in tech_obs)
assert set(obs.content for obs in tech_obs) == {
"Tech observation",
"Another tech observation",
}
# Find design observations
design_obs = await repo.find_by_category("design")
assert len(design_obs) == 1
assert design_obs[0].category == "design"
assert design_obs[0].content == "Design observation"
# Find design observations
design_obs = await repo.find_by_category(session, "design")
assert len(design_obs) == 1
assert design_obs[0].category == "design"
assert design_obs[0].content == "Design observation"
# Search for non-existent category
missing_obs = await repo.find_by_category("missing")
assert len(missing_obs) == 0
# Search for non-existent category
missing_obs = await repo.find_by_category(session, "missing")
assert len(missing_obs) == 0
@pytest.mark.asyncio
@@ -317,22 +330,24 @@ async def test_observation_categories(
await session.commit()
# Get distinct categories
categories = await repo.observation_categories()
async with db.scoped_session(session_maker) as session:
categories = await repo.observation_categories(session)
# Should have unique categories in a deterministic order
assert set(categories) == {"tech", "design", "feature"}
@pytest.mark.asyncio
async def test_find_by_category_with_empty_db(repo):
async def test_find_by_category_with_empty_db(repo, session_maker):
"""Test category operations with an empty database."""
# Find by category should return empty list
obs = await repo.find_by_category("tech")
assert len(obs) == 0
async with db.scoped_session(session_maker) as session:
obs = await repo.find_by_category(session, "tech")
assert len(obs) == 0
# Get categories should return empty list
categories = await repo.observation_categories()
assert len(categories) == 0
# Get categories should return empty list
categories = await repo.observation_categories(session)
assert len(categories) == 0
@pytest.mark.asyncio
@@ -367,11 +382,12 @@ async def test_find_by_category_case_sensitivity(
# Search should work regardless of case
# Note: If we want case-insensitive search, we'll need to update the query
# For now, this test documents the current behavior
exact_match = await repo.find_by_category("tech")
assert len(exact_match) == 1
async with db.scoped_session(session_maker) as session:
exact_match = await repo.find_by_category(session, "tech")
assert len(exact_match) == 1
upper_case = await repo.find_by_category("TECH")
assert len(upper_case) == 0 # Currently case-sensitive
upper_case = await repo.find_by_category(session, "TECH")
assert len(upper_case) == 0 # Currently case-sensitive
@pytest.mark.asyncio
@@ -1001,3 +1001,59 @@ async def test_postgres_search_categories_exact_match(session_maker, test_projec
# Multiple categories union.
multi = await repo.search(categories=["requirement", "decision"])
assert {r.id for r in multi} == {70101, 70102}
@pytest.mark.asyncio
async def test_postgres_question_punctuation_and_relaxation(session_maker, test_project):
"""Question-form queries must produce clean lexemes and a usable relaxation.
Parity with SQLite: sentence punctuation previously reached tsquery terms,
and a strict all-AND miss had no relaxed retry, silently disabling the FTS
half of hybrid search for natural-language questions.
"""
repo = PostgresSearchRepository(session_maker, project_id=test_project.id)
# Edge punctuation stripped before lexeme formatting.
prepared = repo._prepare_search_term("When did Melanie paint a sunrise?")
assert "?" not in prepared
assert "sunrise:*" in prepared
# Relaxation drops stopwords and OR-joins content terms.
relaxed = repo._relaxed_tsquery_text("When did Melanie paint a sunrise?")
assert relaxed == "Melanie:* | paint:* | sunrise:*"
# User intent is not second-guessed.
assert repo._relaxed_tsquery_text("alpha AND beta") is None
assert repo._relaxed_tsquery_text('"exact phrase"') is None
assert repo._relaxed_tsquery_text(None) is None
@pytest.mark.asyncio
async def test_postgres_multiword_query_relaxes_on_strict_miss(session_maker, test_project):
repo = PostgresSearchRepository(session_maker, project_id=test_project.id)
now = datetime.now(timezone.utc)
await repo.index_item(
SearchIndexRow(
project_id=test_project.id,
id=77,
title="Trip plans",
content_stems="melanie painted a sunrise over the lake last year",
content_snippet="Melanie painted a sunrise over the lake last year.",
permalink="docs/trip-plans",
file_path="docs/trip-plans.md",
type="entity",
metadata={"note_type": "note"},
created_at=now,
updated_at=now,
)
)
# A content word absent from the doc ("hiking") makes the strict
# all-terms-AND query miss even after Postgres drops stopwords — without
# it, to_tsquery('english', ...) already strips "when/did/a" and matches.
strict = await repo.search(search_text="Did Melanie go hiking at sunrise?")
assert strict == []
# The hybrid FTS branch opts in; OR-relaxation surfaces the partial match.
results = await repo.search(search_text="Did Melanie go hiking at sunrise?", allow_relaxed=True)
assert any(r.id == 77 for r in results)
@@ -3,6 +3,7 @@
import pytest
from sqlalchemy import text
from basic_memory import db
from basic_memory.repository.project_info_repository import ProjectInfoRepository
from basic_memory.models.project import Project # Add a model reference
@@ -11,11 +12,10 @@ from basic_memory.models.project import Project # Add a model reference
async def test_project_info_repository_init(session_maker):
"""Test ProjectInfoRepository initialization."""
# Create a ProjectInfoRepository
repository = ProjectInfoRepository(session_maker)
repository = ProjectInfoRepository()
# Verify it was initialized properly
assert repository is not None
assert repository.session_maker == session_maker
# Model is set to a dummy value (Project is used as a reference here)
assert repository.Model is Project
@@ -24,10 +24,11 @@ async def test_project_info_repository_init(session_maker):
async def test_project_info_repository_execute_query(session_maker):
"""Test ProjectInfoRepository execute_query method."""
# Create a ProjectInfoRepository
repository = ProjectInfoRepository(session_maker)
repository = ProjectInfoRepository()
# Execute a simple query
result = await repository.execute_query(text("SELECT 1 as test"))
async with db.scoped_session(session_maker) as session:
result = await repository.execute_query(session, text("SELECT 1 as test"))
# Verify the result
assert result is not None
+201 -163
View File
@@ -13,7 +13,7 @@ from basic_memory.repository.project_repository import ProjectRepository
@pytest_asyncio.fixture
async def sample_project(project_repository: ProjectRepository) -> Project:
async def sample_project(project_repository: ProjectRepository, session_maker) -> Project:
"""Create a sample project for testing."""
project_data = {
"name": "Sample Project",
@@ -24,11 +24,12 @@ async def sample_project(project_repository: ProjectRepository) -> Project:
"created_at": datetime.now(timezone.utc),
"updated_at": datetime.now(timezone.utc),
}
return await project_repository.create(project_data)
async with db.scoped_session(session_maker) as session:
return await project_repository.create(session, project_data)
@pytest.mark.asyncio
async def test_create_project(project_repository: ProjectRepository):
async def test_create_project(project_repository: ProjectRepository, session_maker):
"""Test creating a new project."""
project_data = {
"name": "Sample Project",
@@ -37,86 +38,98 @@ async def test_create_project(project_repository: ProjectRepository):
"is_active": True,
"is_default": False,
}
project = await project_repository.create(project_data)
async with db.scoped_session(session_maker) as session:
project = await project_repository.create(session, project_data)
# Verify returned object
assert project.id is not None
assert project.name == "Sample Project"
assert project.description == "A sample project"
assert project.path == "/sample/project/path"
assert project.is_active is True
assert project.is_default is False
assert isinstance(project.created_at, datetime)
assert isinstance(project.updated_at, datetime)
# Verify returned object
assert project.id is not None
assert project.name == "Sample Project"
assert project.description == "A sample project"
assert project.path == "/sample/project/path"
assert project.is_active is True
assert project.is_default is False
assert isinstance(project.created_at, datetime)
assert isinstance(project.updated_at, datetime)
# Verify permalink was generated correctly
assert project.permalink == "sample-project"
# Verify permalink was generated correctly
assert project.permalink == "sample-project"
# Verify in database
found = await project_repository.find_by_id(project.id)
assert found is not None
assert found.id == project.id
assert found.name == project.name
assert found.description == project.description
assert found.path == project.path
assert found.permalink == "sample-project"
assert found.is_active is True
assert found.is_default is False
# Verify in database
found = await project_repository.find_by_id(session, project.id)
assert found is not None
assert found.id == project.id
assert found.name == project.name
assert found.description == project.description
assert found.path == project.path
assert found.permalink == "sample-project"
assert found.is_active is True
assert found.is_default is False
@pytest.mark.asyncio
async def test_get_by_name(project_repository: ProjectRepository, sample_project: Project):
async def test_get_by_name(
project_repository: ProjectRepository, sample_project: Project, session_maker
):
"""Test getting a project by name."""
# Test exact match
found = await project_repository.get_by_name(sample_project.name)
assert found is not None
assert found.id == sample_project.id
assert found.name == sample_project.name
async with db.scoped_session(session_maker) as session:
found = await project_repository.get_by_name(session, sample_project.name)
assert found is not None
assert found.id == sample_project.id
assert found.name == sample_project.name
# Test non-existent name
found = await project_repository.get_by_name("Non-existent Project")
assert found is None
# Test non-existent name
found = await project_repository.get_by_name(session, "Non-existent Project")
assert found is None
@pytest.mark.asyncio
async def test_get_by_permalink(project_repository: ProjectRepository, sample_project: Project):
async def test_get_by_permalink(
project_repository: ProjectRepository, sample_project: Project, session_maker
):
"""Test getting a project by permalink."""
# Verify the permalink value
assert sample_project.permalink == "sample-project"
# Test exact match
found = await project_repository.get_by_permalink(sample_project.permalink)
assert found is not None
assert found.id == sample_project.id
assert found.permalink == sample_project.permalink
async with db.scoped_session(session_maker) as session:
found = await project_repository.get_by_permalink(session, sample_project.permalink)
assert found is not None
assert found.id == sample_project.id
assert found.permalink == sample_project.permalink
# Test non-existent permalink
found = await project_repository.get_by_permalink("non-existent-project")
assert found is None
# Test non-existent permalink
found = await project_repository.get_by_permalink(session, "non-existent-project")
assert found is None
@pytest.mark.asyncio
async def test_get_by_path(project_repository: ProjectRepository, sample_project: Project):
async def test_get_by_path(
project_repository: ProjectRepository, sample_project: Project, session_maker
):
"""Test getting a project by path."""
# Test exact match
found = await project_repository.get_by_path(sample_project.path)
assert found is not None
assert found.id == sample_project.id
assert found.path == sample_project.path
async with db.scoped_session(session_maker) as session:
found = await project_repository.get_by_path(session, sample_project.path)
assert found is not None
assert found.id == sample_project.id
assert found.path == sample_project.path
# Test with Path object
found = await project_repository.get_by_path(Path(sample_project.path))
assert found is not None
assert found.id == sample_project.id
assert found.path == sample_project.path
# Test with Path object
found = await project_repository.get_by_path(session, Path(sample_project.path))
assert found is not None
assert found.id == sample_project.id
assert found.path == sample_project.path
# Test non-existent path
found = await project_repository.get_by_path("/non/existent/path")
assert found is None
# Test non-existent path
found = await project_repository.get_by_path(session, "/non/existent/path")
assert found is None
@pytest.mark.asyncio
async def test_get_default_project(project_repository: ProjectRepository, test_project: Project):
async def test_get_default_project(
project_repository: ProjectRepository, test_project: Project, session_maker
):
"""Test getting the default project."""
# We already have a default project from the test_project fixture
# So just create a non-default project
@@ -128,55 +141,62 @@ async def test_get_default_project(project_repository: ProjectRepository, test_p
"is_default": None, # Not the default project
}
await project_repository.create(non_default_project_data)
async with db.scoped_session(session_maker) as session:
await project_repository.create(session, non_default_project_data)
# Get default project
default_project = await project_repository.get_default_project()
assert default_project is not None
assert default_project.is_default is True
# Get default project
default_project = await project_repository.get_default_project(session)
assert default_project is not None
assert default_project.is_default is True
@pytest.mark.asyncio
async def test_get_default_project_with_false_values(project_repository: ProjectRepository):
async def test_get_default_project_with_false_values(
project_repository: ProjectRepository, session_maker
):
"""Test that get_default_project ignores projects with is_default=False.
Regression test for bug where is_not(None) matched both True and False,
causing MultipleResultsFound when multiple projects had different boolean values.
"""
# Create projects with explicit is_default values
project_true = await project_repository.create(
{
"name": "Default Project",
"path": "/default/path",
"is_default": True,
}
)
async with db.scoped_session(session_maker) as session:
project_true = await project_repository.create(
session,
{
"name": "Default Project",
"path": "/default/path",
"is_default": True,
},
)
await project_repository.create(
{
"name": "Not Default Project",
"path": "/not-default/path",
"is_default": False,
}
)
await project_repository.create(
session,
{
"name": "Not Default Project",
"path": "/not-default/path",
"is_default": False,
},
)
await project_repository.create(
{
"name": "Null Default Project",
"path": "/null/path",
"is_default": None,
}
)
await project_repository.create(
session,
{
"name": "Null Default Project",
"path": "/null/path",
"is_default": None,
},
)
# Should return only the project with is_default=True
default = await project_repository.get_default_project()
assert default is not None
assert default.id == project_true.id
assert default.name == "Default Project"
# Should return only the project with is_default=True
default = await project_repository.get_default_project(session)
assert default is not None
assert default.id == project_true.id
assert default.name == "Default Project"
@pytest.mark.asyncio
async def test_get_active_projects(project_repository: ProjectRepository):
async def test_get_active_projects(project_repository: ProjectRepository, session_maker):
"""Test getting all active projects."""
# Create active and inactive projects
active_project_data = {
@@ -192,27 +212,30 @@ async def test_get_active_projects(project_repository: ProjectRepository):
"is_active": False,
}
await project_repository.create(active_project_data)
await project_repository.create(inactive_project_data)
async with db.scoped_session(session_maker) as session:
await project_repository.create(session, active_project_data)
await project_repository.create(session, inactive_project_data)
# Get active projects
active_projects = await project_repository.get_active_projects()
assert len(active_projects) >= 1 # Could be more from other tests
# Get active projects
active_projects = await project_repository.get_active_projects(session)
assert len(active_projects) >= 1 # Could be more from other tests
# Verify that all returned projects are active
for project in active_projects:
assert project.is_active is True
# Verify that all returned projects are active
for project in active_projects:
assert project.is_active is True
# Verify active project is included
active_names = [p.name for p in active_projects]
assert "Active Project" in active_names
# Verify active project is included
active_names = [p.name for p in active_projects]
assert "Active Project" in active_names
# Verify inactive project is not included
assert "Inactive Project" not in active_names
# Verify inactive project is not included
assert "Inactive Project" not in active_names
@pytest.mark.asyncio
async def test_set_as_default(project_repository: ProjectRepository, test_project: Project):
async def test_set_as_default(
project_repository: ProjectRepository, test_project: Project, session_maker
):
"""Test setting a project as default."""
# The test_project fixture is already the default
# Create a non-default project
@@ -226,30 +249,33 @@ async def test_set_as_default(project_repository: ProjectRepository, test_projec
# Get the existing default project
project1 = test_project
project2 = await project_repository.create(project2_data)
async with db.scoped_session(session_maker) as session:
project2 = await project_repository.create(session, project2_data)
# Verify initial state
assert project1.is_default is True
assert project2.is_default is None
# Verify initial state
assert project1.is_default is True
assert project2.is_default is None
# Set project2 as default
updated_project2 = await project_repository.set_as_default(project2.id)
assert updated_project2 is not None
assert updated_project2.is_default is True
# Set project2 as default
updated_project2 = await project_repository.set_as_default(session, project2.id)
assert updated_project2 is not None
assert updated_project2.is_default is True
# Verify project1 is no longer default
project1_updated = await project_repository.find_by_id(project1.id)
assert project1_updated is not None
assert project1_updated.is_default is None
# Verify project1 is no longer default
project1_updated = await project_repository.find_by_id(session, project1.id)
assert project1_updated is not None
assert project1_updated.is_default is None
# Verify project2 is now default
project2_updated = await project_repository.find_by_id(project2.id)
assert project2_updated is not None
assert project2_updated.is_default is True
# Verify project2 is now default
project2_updated = await project_repository.find_by_id(session, project2.id)
assert project2_updated is not None
assert project2_updated.is_default is True
@pytest.mark.asyncio
async def test_update_project(project_repository: ProjectRepository, sample_project: Project):
async def test_update_project(
project_repository: ProjectRepository, sample_project: Project, session_maker
):
"""Test updating a project."""
# Update project
updated_data = {
@@ -257,80 +283,92 @@ async def test_update_project(project_repository: ProjectRepository, sample_proj
"description": "Updated description",
"path": "/updated/path",
}
updated_project = await project_repository.update(sample_project.id, updated_data)
async with db.scoped_session(session_maker) as session:
updated_project = await project_repository.update(session, sample_project.id, updated_data)
# Verify returned object
assert updated_project is not None
assert updated_project.id == sample_project.id
assert updated_project.name == "Updated Project Name"
assert updated_project.description == "Updated description"
assert updated_project.path == "/updated/path"
# Verify returned object
assert updated_project is not None
assert updated_project.id == sample_project.id
assert updated_project.name == "Updated Project Name"
assert updated_project.description == "Updated description"
assert updated_project.path == "/updated/path"
# Verify permalink was updated based on new name
assert updated_project.permalink == "updated-project-name"
# Verify permalink was updated based on new name
assert updated_project.permalink == "updated-project-name"
# Verify in database
found = await project_repository.find_by_id(sample_project.id)
assert found is not None
assert found.name == "Updated Project Name"
assert found.description == "Updated description"
assert found.path == "/updated/path"
assert found.permalink == "updated-project-name"
# Verify in database
found = await project_repository.find_by_id(session, sample_project.id)
assert found is not None
assert found.name == "Updated Project Name"
assert found.description == "Updated description"
assert found.path == "/updated/path"
assert found.permalink == "updated-project-name"
# Verify we can find by the new permalink
found_by_permalink = await project_repository.get_by_permalink("updated-project-name")
assert found_by_permalink is not None
assert found_by_permalink.id == sample_project.id
# Verify we can find by the new permalink
found_by_permalink = await project_repository.get_by_permalink(
session, "updated-project-name"
)
assert found_by_permalink is not None
assert found_by_permalink.id == sample_project.id
@pytest.mark.asyncio
async def test_delete_project(project_repository: ProjectRepository, sample_project: Project):
async def test_delete_project(
project_repository: ProjectRepository, sample_project: Project, session_maker
):
"""Test deleting a project."""
# Delete project
result = await project_repository.delete(sample_project.id)
assert result is True
async with db.scoped_session(session_maker) as session:
result = await project_repository.delete(session, sample_project.id)
assert result is True
# Verify deletion
deleted = await project_repository.find_by_id(sample_project.id)
assert deleted is None
# Verify deletion
deleted = await project_repository.find_by_id(session, sample_project.id)
assert deleted is None
# Verify with direct database query
async with db.scoped_session(project_repository.session_maker) as session:
# Verify with direct database query
query = select(Project).filter(Project.id == sample_project.id)
result = await session.execute(query)
assert result.scalar_one_or_none() is None
@pytest.mark.asyncio
async def test_delete_nonexistent_project(project_repository: ProjectRepository):
async def test_delete_nonexistent_project(project_repository: ProjectRepository, session_maker):
"""Test deleting a project that doesn't exist."""
result = await project_repository.delete(999) # Non-existent ID
async with db.scoped_session(session_maker) as session:
result = await project_repository.delete(session, 999) # Non-existent ID
assert result is False
@pytest.mark.asyncio
async def test_update_path(project_repository: ProjectRepository, sample_project: Project):
async def test_update_path(
project_repository: ProjectRepository, sample_project: Project, session_maker
):
"""Test updating a project's path."""
new_path = "/new/project/path"
# Update the project path
updated_project = await project_repository.update_path(sample_project.id, new_path)
async with db.scoped_session(session_maker) as session:
updated_project = await project_repository.update_path(session, sample_project.id, new_path)
# Verify returned object
assert updated_project is not None
assert updated_project.id == sample_project.id
assert updated_project.path == new_path
assert updated_project.name == sample_project.name # Other fields unchanged
# Verify returned object
assert updated_project is not None
assert updated_project.id == sample_project.id
assert updated_project.path == new_path
assert updated_project.name == sample_project.name # Other fields unchanged
# Verify in database
found = await project_repository.find_by_id(sample_project.id)
assert found is not None
assert found.path == new_path
assert found.name == sample_project.name
# Verify in database
found = await project_repository.find_by_id(session, sample_project.id)
assert found is not None
assert found.path == new_path
assert found.name == sample_project.name
@pytest.mark.asyncio
async def test_update_path_nonexistent_project(project_repository: ProjectRepository):
async def test_update_path_nonexistent_project(
project_repository: ProjectRepository, session_maker
):
"""Test updating path for a project that doesn't exist."""
result = await project_repository.update_path(999, "/some/path") # Non-existent ID
async with db.scoped_session(session_maker) as session:
result = await project_repository.update_path(session, 999, "/some/path") # Non-existent ID
assert result is None
+212 -140
View File
@@ -75,7 +75,7 @@ async def test_relations(session_maker, source_entity, target_entity, test_proje
@pytest_asyncio.fixture(scope="function")
async def related_entity(entity_repository):
async def related_entity(entity_repository, session_maker):
"""Create a second entity for testing relations"""
entity_data = {
"title": "Related Entity",
@@ -87,12 +87,16 @@ async def related_entity(entity_repository):
"created_at": datetime.now(timezone.utc),
"updated_at": datetime.now(timezone.utc),
}
return await entity_repository.create(entity_data)
async with db.scoped_session(session_maker) as session:
return await entity_repository.create(session, entity_data)
@pytest_asyncio.fixture(scope="function")
async def sample_relation(
relation_repository: RelationRepository, sample_entity: Entity, related_entity: Entity
relation_repository: RelationRepository,
sample_entity: Entity,
related_entity: Entity,
session_maker,
):
"""Create a sample relation for testing"""
relation_data = {
@@ -102,12 +106,16 @@ async def sample_relation(
"relation_type": "test_relation",
"context": "test-context",
}
return await relation_repository.create(relation_data)
async with db.scoped_session(session_maker) as session:
return await relation_repository.create(session, relation_data)
@pytest_asyncio.fixture(scope="function")
async def multiple_relations(
relation_repository: RelationRepository, sample_entity: Entity, related_entity: Entity
relation_repository: RelationRepository,
sample_entity: Entity,
related_entity: Entity,
session_maker,
):
"""Create multiple relations for testing"""
relations_data = [
@@ -133,12 +141,16 @@ async def multiple_relations(
"context": "context_three",
},
]
return [await relation_repository.create(data) for data in relations_data]
async with db.scoped_session(session_maker) as session:
return [await relation_repository.create(session, data) for data in relations_data]
@pytest.mark.asyncio
async def test_create_relation(
relation_repository: RelationRepository, sample_entity: Entity, related_entity: Entity
relation_repository: RelationRepository,
sample_entity: Entity,
related_entity: Entity,
session_maker,
):
"""Test creating a new relation"""
relation_data = {
@@ -148,7 +160,8 @@ async def test_create_relation(
"relation_type": "test_relation",
"context": "test-context",
}
relation = await relation_repository.create(relation_data)
async with db.scoped_session(session_maker) as session:
relation = await relation_repository.create(session, relation_data)
assert relation.from_id == sample_entity.id
assert relation.to_id == related_entity.id
@@ -158,7 +171,10 @@ async def test_create_relation(
@pytest.mark.asyncio
async def test_create_relation_entity_does_not_exist(
relation_repository: RelationRepository, sample_entity: Entity, related_entity: Entity
relation_repository: RelationRepository,
sample_entity: Entity,
related_entity: Entity,
session_maker,
):
"""Test creating a new relation"""
relation_data = {
@@ -169,7 +185,8 @@ async def test_create_relation_entity_does_not_exist(
"context": "test-context",
}
with pytest.raises(IntegrityError):
await relation_repository.create(relation_data)
async with db.scoped_session(session_maker) as session:
await relation_repository.create(session, relation_data)
@pytest.mark.asyncio
@@ -178,37 +195,51 @@ async def test_find_by_entities(
sample_relation: Relation,
sample_entity: Entity,
related_entity: Entity,
session_maker,
):
"""Test finding relations between specific entities"""
relations = await relation_repository.find_by_entities(sample_entity.id, related_entity.id)
async with db.scoped_session(session_maker) as session:
relations = await relation_repository.find_by_entities(
session, sample_entity.id, related_entity.id
)
assert len(relations) == 1
assert relations[0].id == sample_relation.id
assert relations[0].relation_type == sample_relation.relation_type
@pytest.mark.asyncio
async def test_find_relation(relation_repository: RelationRepository, sample_relation: Relation):
async def test_find_relation(
relation_repository: RelationRepository, sample_relation: Relation, session_maker
):
"""Test finding relations by type"""
relation = await relation_repository.find_relation(
from_permalink=sample_relation.from_entity.permalink,
to_permalink=sample_relation.to_entity.permalink,
relation_type=sample_relation.relation_type,
)
async with db.scoped_session(session_maker) as session:
relation = await relation_repository.find_relation(
session=session,
from_permalink=sample_relation.from_entity.permalink,
to_permalink=sample_relation.to_entity.permalink,
relation_type=sample_relation.relation_type,
)
assert relation is not None
assert relation.id == sample_relation.id
@pytest.mark.asyncio
async def test_find_by_type(relation_repository: RelationRepository, sample_relation: Relation):
async def test_find_by_type(
relation_repository: RelationRepository, sample_relation: Relation, session_maker
):
"""Test finding relations by type"""
relations = await relation_repository.find_by_type("test_relation")
async with db.scoped_session(session_maker) as session:
relations = await relation_repository.find_by_type(session, "test_relation")
assert len(relations) == 1
assert relations[0].id == sample_relation.id
@pytest.mark.asyncio
async def test_find_unresolved_relations(
relation_repository: RelationRepository, sample_entity: Entity, related_entity: Entity
relation_repository: RelationRepository,
sample_entity: Entity,
related_entity: Entity,
session_maker,
):
"""Test creating a new relation"""
relation_data = {
@@ -218,32 +249,34 @@ async def test_find_unresolved_relations(
"relation_type": "test_relation",
"context": "test-context",
}
relation = await relation_repository.create(relation_data)
async with db.scoped_session(session_maker) as session:
relation = await relation_repository.create(session, relation_data)
assert relation.from_id == sample_entity.id
assert relation.to_id is None
assert relation.from_id == sample_entity.id
assert relation.to_id is None
unresolved = await relation_repository.find_unresolved_relations()
assert len(unresolved) == 1
assert unresolved[0].id == relation.id
unresolved = await relation_repository.find_unresolved_relations(session)
assert len(unresolved) == 1
assert unresolved[0].id == relation.id
@pytest.mark.asyncio
async def test_delete_by_fields_single_field(
relation_repository: RelationRepository, multiple_relations: list[Relation]
relation_repository: RelationRepository, multiple_relations: list[Relation], session_maker
):
"""Test deleting relations by a single field."""
# Delete all relations of type 'relation_one'
result = await relation_repository.delete_by_fields(relation_type="relation_one") # pyright: ignore [reportArgumentType]
assert result is True
async with db.scoped_session(session_maker) as session:
result = await relation_repository.delete_by_fields(session, relation_type="relation_one") # pyright: ignore [reportArgumentType]
assert result is True
# Verify deletion
remaining = await relation_repository.find_by_type("relation_one")
assert len(remaining) == 0
# Verify deletion
remaining = await relation_repository.find_by_type(session, "relation_one")
assert len(remaining) == 0
# Other relations should still exist
others = await relation_repository.find_by_type("relation_two")
assert len(others) == 1
# Other relations should still exist
others = await relation_repository.find_by_type(session, "relation_two")
assert len(others) == 1
@pytest.mark.asyncio
@@ -252,29 +285,36 @@ async def test_delete_by_fields_multiple_fields(
multiple_relations: list[Relation],
sample_entity: Entity,
related_entity: Entity,
session_maker,
):
"""Test deleting relations by multiple fields."""
# Delete specific relation matching both from_id and relation_type
result = await relation_repository.delete_by_fields(
from_id=sample_entity.id, # pyright: ignore [reportArgumentType]
relation_type="relation_one", # pyright: ignore [reportArgumentType]
)
assert result is True
async with db.scoped_session(session_maker) as session:
result = await relation_repository.delete_by_fields(
session,
from_id=sample_entity.id, # pyright: ignore [reportArgumentType]
relation_type="relation_one", # pyright: ignore [reportArgumentType]
)
assert result is True
# Verify correct relation was deleted
remaining = await relation_repository.find_by_entities(sample_entity.id, related_entity.id)
assert len(remaining) == 1 # Only relation_two should remain
assert remaining[0].relation_type == "relation_two"
# Verify correct relation was deleted
remaining = await relation_repository.find_by_entities(
session, sample_entity.id, related_entity.id
)
assert len(remaining) == 1 # Only relation_two should remain
assert remaining[0].relation_type == "relation_two"
@pytest.mark.asyncio
async def test_delete_by_fields_no_match(
relation_repository: RelationRepository, multiple_relations: list[Relation]
relation_repository: RelationRepository, multiple_relations: list[Relation], session_maker
):
"""Test delete_by_fields when no relations match."""
result = await relation_repository.delete_by_fields(
relation_type="nonexistent_type" # pyright: ignore [reportArgumentType]
)
async with db.scoped_session(session_maker) as session:
result = await relation_repository.delete_by_fields(
session,
relation_type="nonexistent_type", # pyright: ignore [reportArgumentType]
)
assert result is False
@@ -284,73 +324,82 @@ async def test_delete_by_fields_all_fields(
multiple_relations: list[Relation],
sample_entity: Entity,
related_entity: Entity,
session_maker,
):
"""Test deleting relation by matching all fields."""
# Get first relation's data
relation = multiple_relations[0]
# Delete using all fields
result = await relation_repository.delete_by_fields(
from_id=relation.from_id, # pyright: ignore [reportArgumentType]
to_id=relation.to_id, # pyright: ignore [reportArgumentType]
relation_type=relation.relation_type, # pyright: ignore [reportArgumentType]
)
assert result is True
async with db.scoped_session(session_maker) as session:
result = await relation_repository.delete_by_fields(
session,
from_id=relation.from_id, # pyright: ignore [reportArgumentType]
to_id=relation.to_id, # pyright: ignore [reportArgumentType]
relation_type=relation.relation_type, # pyright: ignore [reportArgumentType]
)
assert result is True
# Verify only exact match was deleted
remaining = await relation_repository.find_by_type(relation.relation_type)
assert len(remaining) == 1 # One other relation_one should remain
# Verify only exact match was deleted
remaining = await relation_repository.find_by_type(session, relation.relation_type)
assert len(remaining) == 1 # One other relation_one should remain
@pytest.mark.asyncio
async def test_delete_relation_by_id(relation_repository, test_relations):
async def test_delete_relation_by_id(relation_repository, test_relations, session_maker):
"""Test deleting a relation by ID."""
relation = test_relations[0]
result = await relation_repository.delete(relation.id)
assert result is True
async with db.scoped_session(session_maker) as session:
result = await relation_repository.delete(session, relation.id)
assert result is True
# Verify deletion
remaining = await relation_repository.find_one(
relation_repository.select(Relation).filter(Relation.id == relation.id)
)
assert remaining is None
# Verify deletion
remaining = await relation_repository.find_one(
session, relation_repository.select(Relation).filter(Relation.id == relation.id)
)
assert remaining is None
@pytest.mark.asyncio
async def test_delete_relations_by_type(relation_repository, test_relations):
async def test_delete_relations_by_type(relation_repository, test_relations, session_maker):
"""Test deleting relations by type."""
result = await relation_repository.delete_by_fields(relation_type="connects_to")
assert result is True
async with db.scoped_session(session_maker) as session:
result = await relation_repository.delete_by_fields(session, relation_type="connects_to")
assert result is True
# Verify specific type was deleted
remaining = await relation_repository.find_by_type("connects_to")
assert len(remaining) == 0
# Verify specific type was deleted
remaining = await relation_repository.find_by_type(session, "connects_to")
assert len(remaining) == 0
# Verify other type still exists
others = await relation_repository.find_by_type("depends_on")
assert len(others) == 1
# Verify other type still exists
others = await relation_repository.find_by_type(session, "depends_on")
assert len(others) == 1
@pytest.mark.asyncio
async def test_delete_relations_by_entities(
relation_repository, test_relations, source_entity, target_entity
relation_repository, test_relations, source_entity, target_entity, session_maker
):
"""Test deleting relations between specific entities."""
result = await relation_repository.delete_by_fields(
from_id=source_entity.id, to_id=target_entity.id
)
assert result is True
async with db.scoped_session(session_maker) as session:
result = await relation_repository.delete_by_fields(
session, from_id=source_entity.id, to_id=target_entity.id
)
assert result is True
# Verify all relations between entities were deleted
remaining = await relation_repository.find_by_entities(source_entity.id, target_entity.id)
assert len(remaining) == 0
# Verify all relations between entities were deleted
remaining = await relation_repository.find_by_entities(
session, source_entity.id, target_entity.id
)
assert len(remaining) == 0
@pytest.mark.asyncio
async def test_delete_nonexistent_relation(relation_repository):
async def test_delete_nonexistent_relation(relation_repository, session_maker):
"""Test deleting a relation that doesn't exist."""
result = await relation_repository.delete_by_fields(relation_type="nonexistent")
async with db.scoped_session(session_maker) as session:
result = await relation_repository.delete_by_fields(session, relation_type="nonexistent")
assert result is False
@@ -361,7 +410,10 @@ async def test_delete_nonexistent_relation(relation_repository):
@pytest.mark.asyncio
async def test_add_all_ignore_duplicates_basic(
relation_repository: RelationRepository, sample_entity: Entity, related_entity: Entity
relation_repository: RelationRepository,
sample_entity: Entity,
related_entity: Entity,
session_maker,
):
"""Test bulk inserting relations with ON CONFLICT DO NOTHING."""
relations = [
@@ -379,21 +431,27 @@ async def test_add_all_ignore_duplicates_basic(
),
]
inserted = await relation_repository.add_all_ignore_duplicates(relations)
async with db.scoped_session(session_maker) as session:
inserted = await relation_repository.add_all_ignore_duplicates(session, relations)
# Both should be inserted
assert inserted == 2
# Both should be inserted
assert inserted == 2
# Verify they exist
found = await relation_repository.find_by_entities(sample_entity.id, related_entity.id)
assert len(found) == 2
relation_types = {r.relation_type for r in found}
assert relation_types == {"links_to", "references"}
# Verify they exist
found = await relation_repository.find_by_entities(
session, sample_entity.id, related_entity.id
)
assert len(found) == 2
relation_types = {r.relation_type for r in found}
assert relation_types == {"links_to", "references"}
@pytest.mark.asyncio
async def test_add_all_ignore_duplicates_skips_duplicates(
relation_repository: RelationRepository, sample_entity: Entity, related_entity: Entity
relation_repository: RelationRepository,
sample_entity: Entity,
related_entity: Entity,
session_maker,
):
"""Test that duplicate relations are silently ignored."""
# Same relation appearing multiple times (common when same [[link]] appears twice in doc)
@@ -418,27 +476,34 @@ async def test_add_all_ignore_duplicates_skips_duplicates(
),
]
inserted = await relation_repository.add_all_ignore_duplicates(relations)
async with db.scoped_session(session_maker) as session:
inserted = await relation_repository.add_all_ignore_duplicates(session, relations)
# Only 1 should be inserted (duplicates ignored)
assert inserted == 1
# Only 1 should be inserted (duplicates ignored)
assert inserted == 1
# Verify only one exists
all_relations = await relation_repository.find_all()
matching = [r for r in all_relations if r.to_name == "Some Target"]
assert len(matching) == 1
# Verify only one exists
all_relations = await relation_repository.find_all(session)
matching = [r for r in all_relations if r.to_name == "Some Target"]
assert len(matching) == 1
@pytest.mark.asyncio
async def test_add_all_ignore_duplicates_empty_list(relation_repository: RelationRepository):
async def test_add_all_ignore_duplicates_empty_list(
relation_repository: RelationRepository, session_maker
):
"""Test with empty list returns 0."""
inserted = await relation_repository.add_all_ignore_duplicates([])
async with db.scoped_session(session_maker) as session:
inserted = await relation_repository.add_all_ignore_duplicates(session, [])
assert inserted == 0
@pytest.mark.asyncio
async def test_add_all_ignore_duplicates_mixed(
relation_repository: RelationRepository, sample_entity: Entity, related_entity: Entity
relation_repository: RelationRepository,
sample_entity: Entity,
related_entity: Entity,
session_maker,
):
"""Test with mix of new and duplicate relations."""
# First, insert one relation
@@ -448,44 +513,48 @@ async def test_add_all_ignore_duplicates_mixed(
to_name="Existing Target",
relation_type="links_to",
)
await relation_repository.add_all_ignore_duplicates([first_relation])
async with db.scoped_session(session_maker) as session:
await relation_repository.add_all_ignore_duplicates(session, [first_relation])
# Now try to insert a mix of new and duplicate
relations = [
Relation(
from_id=sample_entity.id,
to_id=None,
to_name="Existing Target", # Duplicate of first_relation
relation_type="links_to",
),
Relation(
from_id=sample_entity.id,
to_id=None,
to_name="New Target 1", # New
relation_type="links_to",
),
Relation(
from_id=sample_entity.id,
to_id=None,
to_name="New Target 2", # New
relation_type="references",
),
]
# Now try to insert a mix of new and duplicate
relations = [
Relation(
from_id=sample_entity.id,
to_id=None,
to_name="Existing Target", # Duplicate of first_relation
relation_type="links_to",
),
Relation(
from_id=sample_entity.id,
to_id=None,
to_name="New Target 1", # New
relation_type="links_to",
),
Relation(
from_id=sample_entity.id,
to_id=None,
to_name="New Target 2", # New
relation_type="references",
),
]
inserted = await relation_repository.add_all_ignore_duplicates(relations)
inserted = await relation_repository.add_all_ignore_duplicates(session, relations)
# Only 2 new ones should be inserted
assert inserted == 2
# Only 2 new ones should be inserted
assert inserted == 2
# Verify total count
all_relations = await relation_repository.find_all()
from_sample = [r for r in all_relations if r.from_id == sample_entity.id]
assert len(from_sample) == 3 # 1 existing + 2 new
# Verify total count
all_relations = await relation_repository.find_all(session)
from_sample = [r for r in all_relations if r.from_id == sample_entity.id]
assert len(from_sample) == 3 # 1 existing + 2 new
@pytest.mark.asyncio
async def test_add_all_ignore_duplicates_with_context(
relation_repository: RelationRepository, sample_entity: Entity, related_entity: Entity
relation_repository: RelationRepository,
sample_entity: Entity,
related_entity: Entity,
session_maker,
):
"""Test that context field is properly inserted."""
relations = [
@@ -498,10 +567,13 @@ async def test_add_all_ignore_duplicates_with_context(
),
]
inserted = await relation_repository.add_all_ignore_duplicates(relations)
assert inserted == 1
async with db.scoped_session(session_maker) as session:
inserted = await relation_repository.add_all_ignore_duplicates(session, relations)
assert inserted == 1
# Verify context was saved
found = await relation_repository.find_by_entities(sample_entity.id, related_entity.id)
assert len(found) == 1
assert found[0].context == "some context here"
# Verify context was saved
found = await relation_repository.find_by_entities(
session, sample_entity.id, related_entity.id
)
assert len(found) == 1
assert found[0].context == "some context here"
+93 -77
View File
@@ -5,6 +5,7 @@ import pytest
from sqlalchemy import String, DateTime
from sqlalchemy.orm import Mapped, mapped_column
from basic_memory import db
from basic_memory.models import Base
from basic_memory.repository.repository import Repository
@@ -28,197 +29,212 @@ class ModelTest(Base):
@pytest.fixture
def repository(session_maker):
def repository():
"""Create a test repository."""
return Repository(session_maker, ModelTest)
return Repository(ModelTest)
@pytest.mark.asyncio
async def test_add(repository):
async def test_add(repository, session_maker):
"""Test bulk creation of entities."""
# Create test instances
instance = ModelTest(id="test_add", name="Test Add")
await repository.add(instance)
async with db.scoped_session(session_maker) as session:
await repository.add(session, instance)
# Verify we can find in db
found = await repository.find_by_id("test_add")
# Verify we can find in db
found = await repository.find_by_id(session, "test_add")
assert found is not None
assert found.name == "Test Add"
@pytest.mark.asyncio
async def test_add_all(repository):
async def test_add_all(repository, session_maker):
"""Test bulk creation of entities."""
# Create test instances
instances = [ModelTest(id=f"test_{i}", name=f"Test {i}") for i in range(3)]
await repository.add_all(instances)
async with db.scoped_session(session_maker) as session:
await repository.add_all(session, instances)
# Verify we can find them in db
found = await repository.find_by_id("test_0")
# Verify we can find them in db
found = await repository.find_by_id(session, "test_0")
assert found is not None
assert found.name == "Test 0"
@pytest.mark.asyncio
async def test_add_all_no_return(repository):
async def test_add_all_no_return(repository, session_maker):
"""Bulk inserts can skip the follow-up reload when callers do not need rows back."""
instances = [ModelTest(id=f"test_{i}", name=f"Test {i}") for i in range(3)]
inserted = await repository.add_all_no_return(instances)
async with db.scoped_session(session_maker) as session:
inserted = await repository.add_all_no_return(session, instances)
assert inserted == 3
found = await repository.find_by_id("test_0")
assert inserted == 3
found = await repository.find_by_id(session, "test_0")
assert found is not None
assert found.name == "Test 0"
@pytest.mark.asyncio
async def test_bulk_create(repository):
async def test_bulk_create(repository, session_maker):
"""Test bulk creation of entities."""
# Create test instances
instances = [ModelTest(id=f"test_{i}", name=f"Test {i}") for i in range(3)]
# Bulk create
await repository.create_all([instance.__dict__ for instance in instances])
async with db.scoped_session(session_maker) as session:
await repository.create_all(session, [instance.__dict__ for instance in instances])
# Verify we can find them in db
found = await repository.find_by_id("test_0")
# Verify we can find them in db
found = await repository.find_by_id(session, "test_0")
assert found is not None
assert found.name == "Test 0"
@pytest.mark.asyncio
async def test_find_all(repository):
async def test_find_all(repository, session_maker):
"""Test finding multiple entities by IDs."""
# Create test data
instances = [ModelTest(id=f"test_{i}", name=f"Test {i}") for i in range(5)]
await repository.create_all([instance.__dict__ for instance in instances])
async with db.scoped_session(session_maker) as session:
await repository.create_all(session, [instance.__dict__ for instance in instances])
found = await repository.find_all(limit=3)
found = await repository.find_all(session, limit=3)
assert len(found) == 3
@pytest.mark.asyncio
async def test_find_by_ids(repository):
async def test_find_by_ids(repository, session_maker):
"""Test finding multiple entities by IDs."""
# Create test data
instances = [ModelTest(id=f"test_{i}", name=f"Test {i}") for i in range(5)]
await repository.create_all([instance.__dict__ for instance in instances])
async with db.scoped_session(session_maker) as session:
await repository.create_all(session, [instance.__dict__ for instance in instances])
# Test finding subset of entities
ids_to_find = ["test_0", "test_2", "test_4"]
found = await repository.find_by_ids(ids_to_find)
assert len(found) == 3
assert sorted([e.id for e in found]) == sorted(ids_to_find)
# Test finding subset of entities
ids_to_find = ["test_0", "test_2", "test_4"]
found = await repository.find_by_ids(session, ids_to_find)
assert len(found) == 3
assert sorted([e.id for e in found]) == sorted(ids_to_find)
# Test finding with some non-existent IDs
mixed_ids = ["test_0", "nonexistent", "test_4"]
partial_found = await repository.find_by_ids(mixed_ids)
assert len(partial_found) == 2
assert sorted([e.id for e in partial_found]) == ["test_0", "test_4"]
# Test finding with some non-existent IDs
mixed_ids = ["test_0", "nonexistent", "test_4"]
partial_found = await repository.find_by_ids(session, mixed_ids)
assert len(partial_found) == 2
assert sorted([e.id for e in partial_found]) == ["test_0", "test_4"]
# Test with empty list
empty_found = await repository.find_by_ids([])
assert len(empty_found) == 0
# Test with empty list
empty_found = await repository.find_by_ids(session, [])
assert len(empty_found) == 0
# Test with all non-existent IDs
not_found = await repository.find_by_ids(["fake1", "fake2"])
assert len(not_found) == 0
# Test with all non-existent IDs
not_found = await repository.find_by_ids(session, ["fake1", "fake2"])
assert len(not_found) == 0
@pytest.mark.asyncio
async def test_delete_by_ids(repository):
async def test_delete_by_ids(repository, session_maker):
"""Test finding multiple entities by IDs."""
# Create test data
instances = [ModelTest(id=f"test_{i}", name=f"Test {i}") for i in range(5)]
await repository.create_all([instance.__dict__ for instance in instances])
async with db.scoped_session(session_maker) as session:
await repository.create_all(session, [instance.__dict__ for instance in instances])
# Test delete subset of entities
ids_to_delete = ["test_0", "test_2", "test_4"]
deleted_count = await repository.delete_by_ids(ids_to_delete)
assert deleted_count == 3
# Test delete subset of entities
ids_to_delete = ["test_0", "test_2", "test_4"]
deleted_count = await repository.delete_by_ids(session, ids_to_delete)
assert deleted_count == 3
# Test finding subset of entities
ids_to_find = ["test_1", "test_3"]
found = await repository.find_by_ids(ids_to_find)
assert len(found) == 2
assert sorted([e.id for e in found]) == sorted(ids_to_find)
# Test finding subset of entities
ids_to_find = ["test_1", "test_3"]
found = await repository.find_by_ids(session, ids_to_find)
assert len(found) == 2
assert sorted([e.id for e in found]) == sorted(ids_to_find)
assert await repository.find_by_id(ids_to_delete[0]) is None
assert await repository.find_by_id(ids_to_delete[1]) is None
assert await repository.find_by_id(ids_to_delete[2]) is None
assert await repository.find_by_id(session, ids_to_delete[0]) is None
assert await repository.find_by_id(session, ids_to_delete[1]) is None
assert await repository.find_by_id(session, ids_to_delete[2]) is None
@pytest.mark.asyncio
async def test_update(repository):
async def test_update(repository, session_maker):
"""Test finding entities modified since a timestamp."""
# Create initial test data
instance = ModelTest(id="test_add", name="Test Add")
await repository.add(instance)
async with db.scoped_session(session_maker) as session:
await repository.add(session, instance)
instance = ModelTest(id="test_add", name="Updated")
instance = ModelTest(id="test_add", name="Updated")
# Find recently modified
modified = await repository.update(instance.id, {"name": "Updated"})
# Find recently modified
modified = await repository.update(session, instance.id, {"name": "Updated"})
assert modified is not None
assert modified.name == "Updated"
@pytest.mark.asyncio
async def test_update_fields(repository):
async def test_update_fields(repository, session_maker):
"""Column-only updates can skip eager reloads on write-heavy paths."""
instance = ModelTest(id="test_add", name="Test Add")
await repository.add(instance)
async with db.scoped_session(session_maker) as session:
await repository.add(session, instance)
updated = await repository.update_fields(instance.id, {"name": "Updated"})
updated = await repository.update_fields(session, instance.id, {"name": "Updated"})
assert updated is True
found = await repository.find_by_id(instance.id)
assert updated is True
found = await repository.find_by_id(session, instance.id)
assert found is not None
assert found.name == "Updated"
@pytest.mark.asyncio
async def test_update_fields_not_found(repository):
async def test_update_fields_not_found(repository, session_maker):
"""Column-only updates still report when no row matched."""
updated = await repository.update_fields("missing", {"name": "Updated"})
async with db.scoped_session(session_maker) as session:
updated = await repository.update_fields(session, "missing", {"name": "Updated"})
assert updated is False
@pytest.mark.asyncio
async def test_update_model(repository):
async def test_update_model(repository, session_maker):
"""Test finding entities modified since a timestamp."""
# Create initial test data
instance = ModelTest(id="test_add", name="Test Add")
await repository.add(instance)
async with db.scoped_session(session_maker) as session:
await repository.add(session, instance)
instance.name = "Updated"
instance.name = "Updated"
# Find recently modified
modified = await repository.update(instance.id, instance)
# Find recently modified
modified = await repository.update(session, instance.id, instance)
assert modified is not None
assert modified.name == "Updated"
@pytest.mark.asyncio
async def test_update_model_not_found(repository):
async def test_update_model_not_found(repository, session_maker):
"""Test finding entities modified since a timestamp."""
# Create initial test data
instance = ModelTest(id="test_add", name="Test Add")
await repository.add(instance)
async with db.scoped_session(session_maker) as session:
await repository.add(session, instance)
modified = await repository.update("0", {}) # Use string ID for Postgres compatibility
modified = await repository.update(
session, "0", {}
) # Use string ID for Postgres compatibility
assert modified is None
@pytest.mark.asyncio
async def test_count(repository):
async def test_count(repository, session_maker):
"""Test bulk creation of entities."""
# Create test instances
instance = ModelTest(id="test_add", name="Test Add")
await repository.add(instance)
async with db.scoped_session(session_maker) as session:
await repository.add(session, instance)
# Verify we can count in db
count = await repository.count()
# Verify we can count in db
count = await repository.count(session)
assert count == 1
+83 -2
View File
@@ -39,7 +39,7 @@ async def search_entity(session_maker, test_project: Project):
@pytest_asyncio.fixture
async def second_project(project_repository):
async def second_project(project_repository, session_maker):
"""Create a second project for testing project isolation."""
project_data = {
"name": "Second Test Project",
@@ -48,7 +48,8 @@ async def second_project(project_repository):
"is_active": True,
"is_default": None,
}
return await project_repository.create(project_data)
async with db.scoped_session(session_maker) as session:
return await project_repository.create(session, project_data)
@pytest_asyncio.fixture
@@ -1124,3 +1125,83 @@ async def test_search_categories_exact_match(search_repository, search_entity):
# Multiple categories union: both observations come back.
multi = await search_repository.search(categories=["requirement", "decision"])
assert {r.id for r in multi} == {70001, 70002}
@pytest.mark.asyncio
async def test_question_punctuation_does_not_phrase_quote(search_repository):
"""Sentence punctuation must not force exact-phrase matching (#hybrid-fts).
'When did Melanie paint a sunrise?' previously became the FTS5 phrase
'"When did Melanie paint a sunrise?"*' zero rows for any corpus which
silently disabled the FTS half of hybrid search for question queries.
"""
prepared = search_repository._prepare_single_term("When did Melanie paint a sunrise?")
assert '"' not in prepared
# Prefix syntax differs by backend: FTS5 uses '*', tsquery uses ':*'.
if is_postgres_backend(search_repository):
assert "sunrise:*" in prepared
else:
assert "sunrise*" in prepared
@pytest.mark.asyncio
async def test_relaxed_query_drops_stopwords(search_repository):
"""Relaxation keys on content-bearing terms in each backend's syntax."""
if is_postgres_backend(search_repository):
relaxed = search_repository._relaxed_tsquery_text("When did Melanie paint a sunrise?")
assert relaxed == "Melanie:* | paint:* | sunrise:*"
else:
relaxed = search_repository._relaxed_fts_text("When did Melanie paint a sunrise?")
assert relaxed == "Melanie* OR paint* OR sunrise*"
@pytest.mark.asyncio
async def test_relaxed_query_respects_user_intent(search_repository):
# Eligibility matches the service-level relaxation (both backends): quoted,
# boolean, short (<3 tokens), and numeric-identifier queries are not relaxed.
relaxer = (
search_repository._relaxed_tsquery_text
if is_postgres_backend(search_repository)
else search_repository._relaxed_fts_text
)
assert relaxer("alpha AND beta") is None
assert relaxer('"exact phrase"') is None
assert relaxer("single") is None # < 3 tokens
assert relaxer("New Feature") is None # < 3 tokens (link title)
assert relaxer("root note 1") is None # numeric identifier token
assert relaxer("SPEC 16 design") is None # numeric token
assert relaxer(None) is None
@pytest.mark.asyncio
async def test_multiword_query_relaxes_to_or_when_strict_misses(search_repository, search_entity):
"""A question sharing only SOME words with a doc still surfaces it."""
from basic_memory.repository.search_index_row import SearchIndexRow
from basic_memory.schemas.search import SearchItemType
row = SearchIndexRow(
project_id=search_repository.project_id,
id=search_entity.id,
type=SearchItemType.ENTITY.value,
title="Trip plans",
content_snippet="Melanie painted a sunrise over the lake last year.",
content_stems="melanie painted a sunrise over the lake last year",
permalink=search_entity.permalink,
file_path=search_entity.file_path,
entity_id=search_entity.id,
metadata={"note_type": search_entity.note_type},
created_at=search_entity.created_at,
updated_at=search_entity.updated_at,
)
await search_repository.index_item(row)
# "hiking" is absent from the doc, so strict all-terms-AND misses on both
# backends (Postgres's stopword stripping can't rescue it either).
strict = await search_repository.search(search_text="Did Melanie go hiking at sunrise?")
assert strict == []
# The hybrid FTS branch opts in; OR-relaxation surfaces the partial match.
results = await search_repository.search(
search_text="Did Melanie go hiking at sunrise?", allow_relaxed=True
)
assert any(r.entity_id == search_entity.id for r in results)
@@ -9,6 +9,7 @@ from datetime import datetime, timezone
import pytest
import pytest_asyncio
from basic_memory import db
from basic_memory.models.project import Project
from basic_memory.repository.search_index_row import SearchIndexRow
from basic_memory.repository.sqlite_search_repository import SQLiteSearchRepository
@@ -16,7 +17,7 @@ from basic_memory.schemas.search import SearchItemType
@pytest_asyncio.fixture
async def second_test_project(project_repository):
async def second_test_project(project_repository, session_maker):
"""Create a second project for testing project isolation during edits."""
project_data = {
"name": "Second Edit Test Project",
@@ -25,7 +26,8 @@ async def second_test_project(project_repository):
"is_active": True,
"is_default": None,
}
return await project_repository.create(project_data)
async with db.scoped_session(session_maker) as session:
return await project_repository.create(session, project_data)
@pytest_asyncio.fixture
@@ -89,6 +89,7 @@ class _ConcreteRepo(SearchRepositoryBase):
min_similarity: float | None = None,
limit: int = 10,
offset: int = 0,
allow_relaxed: bool = False,
) -> list[SearchIndexRow]:
return []
@@ -62,6 +62,7 @@ class ConcreteSearchRepo(SearchRepositoryBase):
min_similarity: float | None = None,
limit: int = 10,
offset: int = 0,
allow_relaxed: bool = False,
) -> list[SearchIndexRow]:
return [] # pragma: no cover
@@ -66,6 +66,7 @@ class ConcreteSearchRepo(SearchRepositoryBase):
min_similarity: Optional[float] = None,
limit: int = 10,
offset: int = 0,
allow_relaxed: bool = False,
) -> list[SearchIndexRow]:
return [] # pragma: no cover
+51 -26
View File
@@ -5,6 +5,7 @@ from datetime import datetime, timedelta, UTC
import pytest
import pytest_asyncio
from basic_memory import db
from basic_memory.repository.search_repository import SearchIndexRow
from basic_memory.schemas.memory import memory_url, memory_url_path
from basic_memory.schemas.search import SearchItemType
@@ -15,11 +16,15 @@ from basic_memory.models.project import Project
@pytest_asyncio.fixture
async def context_service(
search_repository, entity_repository, observation_repository, link_resolver
search_repository, entity_repository, observation_repository, link_resolver, session_maker
):
"""Create context service for testing."""
return ContextService(
search_repository, entity_repository, observation_repository, link_resolver=link_resolver
search_repository,
entity_repository,
observation_repository,
link_resolver=link_resolver,
session_maker=session_maker,
)
@@ -49,7 +54,7 @@ async def test_find_connected_depth_limit(context_service, test_graph):
@pytest.mark.asyncio
async def test_find_connected_timeframe(
context_service, test_graph, search_repository, entity_repository, app_config
context_service, test_graph, search_repository, entity_repository, app_config, session_maker
):
"""Test timeframe filtering.
This tests how traversal is affected by the item dates.
@@ -70,13 +75,18 @@ async def test_find_connected_timeframe(
# Update entity table timestamps directly
# Root entity uses old date
root_entity = test_graph["root"]
await entity_repository.update(root_entity.id, {"created_at": old_date, "updated_at": old_date})
async with db.scoped_session(session_maker) as session:
await entity_repository.update(
session, root_entity.id, {"created_at": old_date, "updated_at": old_date}
)
# Connected entity uses recent date
connected_entity = test_graph["connected1"]
await entity_repository.update(
connected_entity.id, {"created_at": recent_date, "updated_at": recent_date}
)
# Connected entity uses recent date
connected_entity = test_graph["connected1"]
await entity_repository.update(
session,
connected_entity.id,
{"created_at": recent_date, "updated_at": recent_date},
)
# Also update search_index for test consistency
await search_repository.index_item(
@@ -349,14 +359,18 @@ async def test_project_isolation_in_find_related(session_maker, app_config):
search_repo_p2 = SQLiteSearchRepository(session_maker, project2.id)
# Create repositories for project1
entity_repo_p1 = EntityRepository(session_maker, project1.id)
obs_repo_p1 = ObservationRepository(session_maker, project1.id)
context_service_p1 = ContextService(search_repo_p1, entity_repo_p1, obs_repo_p1)
entity_repo_p1 = EntityRepository(project1.id)
obs_repo_p1 = ObservationRepository(project1.id)
context_service_p1 = ContextService(
search_repo_p1, entity_repo_p1, obs_repo_p1, session_maker=session_maker
)
# Create repositories for project2
entity_repo_p2 = EntityRepository(session_maker, project2.id)
obs_repo_p2 = ObservationRepository(session_maker, project2.id)
context_service_p2 = ContextService(search_repo_p2, entity_repo_p2, obs_repo_p2)
entity_repo_p2 = EntityRepository(project2.id)
obs_repo_p2 = ObservationRepository(project2.id)
context_service_p2 = ContextService(
search_repo_p2, entity_repo_p2, obs_repo_p2, session_maker=session_maker
)
# Test: find_related for project1 should only return project1 entities
type_id_pairs_p1 = [("entity", entity1_p1.id)]
@@ -470,9 +484,11 @@ async def test_find_related_expands_cross_project_relation_targets(session_maker
else:
search_repo_p1 = SQLiteSearchRepository(session_maker, project1.id)
entity_repo_p1 = EntityRepository(session_maker, project1.id)
obs_repo_p1 = ObservationRepository(session_maker, project1.id)
context_service_p1 = ContextService(search_repo_p1, entity_repo_p1, obs_repo_p1)
entity_repo_p1 = EntityRepository(project1.id)
obs_repo_p1 = ObservationRepository(project1.id)
context_service_p1 = ContextService(
search_repo_p1, entity_repo_p1, obs_repo_p1, session_maker=session_maker
)
await search_repo_p1.index_item(
SearchIndexRow(
@@ -586,9 +602,11 @@ async def test_find_related_does_not_revisit_entities_in_cycles(session_maker, a
else:
search_repo = SQLiteSearchRepository(session_maker, project.id)
entity_repo = EntityRepository(session_maker, project.id)
obs_repo = ObservationRepository(session_maker, project.id)
context_service = ContextService(search_repo, entity_repo, obs_repo)
entity_repo = EntityRepository(project.id)
obs_repo = ObservationRepository(project.id)
context_service = ContextService(
search_repo, entity_repo, obs_repo, session_maker=session_maker
)
related = await context_service.find_related(
[("entity", root.id)], max_depth=4, max_results=100
@@ -631,10 +649,15 @@ async def test_build_context_fallback_not_found(context_service):
@pytest.mark.asyncio
async def test_build_context_without_link_resolver(
search_repository, entity_repository, observation_repository, test_graph
search_repository, entity_repository, observation_repository, test_graph, session_maker
):
"""Test that build_context still works without a link_resolver (no fallback)."""
service = ContextService(search_repository, entity_repository, observation_repository)
service = ContextService(
search_repository,
entity_repository,
observation_repository,
session_maker=session_maker,
)
# Exact permalink lookup should still work
url = memory_url.validate_strings("memory://test-project/test/root")
@@ -713,9 +736,11 @@ async def test_find_related_carries_to_name_for_unresolved_relations(session_mak
search_repo = PostgresSearchRepository(session_maker, project.id)
else:
search_repo = SQLiteSearchRepository(session_maker, project.id)
entity_repo = EntityRepository(session_maker, project.id)
obs_repo = ObservationRepository(session_maker, project.id)
context_service = ContextService(search_repo, entity_repo, obs_repo)
entity_repo = EntityRepository(project.id)
obs_repo = ObservationRepository(project.id)
context_service = ContextService(
search_repo, entity_repo, obs_repo, session_maker=session_maker
)
related = await context_service.find_related([("entity", source.id)], max_depth=2)
relation_rows = {r.to_name: r for r in related if r.type == "relation"}
+10 -3
View File
@@ -106,6 +106,7 @@ def entity_service_with_search(
link_resolver,
search_service: SearchService,
app_config: BasicMemoryConfig,
session_maker,
) -> EntityService:
"""Create EntityService with a real attached search service."""
return EntityService(
@@ -117,6 +118,7 @@ def entity_service_with_search(
link_resolver=link_resolver,
search_service=search_service,
app_config=app_config,
session_maker=session_maker,
)
@@ -223,7 +225,8 @@ async def test_create_entity_unique_permalink(
# move file
file_path = file_service.get_entity_path(entity)
file_path.rename(project_config.home / "new_path.md")
await entity_repository.update(entity.id, {"file_path": "new_path.md"})
async with db.scoped_session(entity_service.session_maker) as session:
await entity_repository.update(session, entity.id, {"file_path": "new_path.md"})
# create again
entity2 = await entity_service.create_entity(entity_data)
@@ -2241,7 +2244,8 @@ async def test_move_entity_with_null_permalink_generates_permalink(
}
# Create the entity directly in database
created_entity = await entity_repository.create(entity_data)
async with db.scoped_session(entity_service.session_maker) as session:
created_entity = await entity_repository.create(session, entity_data)
assert created_entity.permalink is None
# Create the physical file
@@ -2530,6 +2534,7 @@ async def test_delete_directory_all_already_deleted(entity_service: EntityServic
@pytest.mark.asyncio
async def test_delete_directory_entity_deleted_between_query_and_delete(
entity_service: EntityService,
session_maker,
):
"""Simulates the real race condition: entity exists in prefix query but is deleted
by a concurrent request before delete_entity is called."""
@@ -2538,7 +2543,9 @@ async def test_delete_directory_entity_deleted_between_query_and_delete(
created = await entity_service.create_entity(entity_data)
# Get the entities via prefix query (as delete_directory does)
entities = await entity_service.repository.find_by_directory_prefix("race-dir")
async with db.scoped_session(session_maker) as session:
entities = await entity_service.repository.find_by_directory_prefix(session, "race-dir")
assert len(entities) == 1
# Now delete the entity behind the scenes (simulating a concurrent request)
@@ -18,6 +18,7 @@ async def test_create_entity_with_permalinks_disabled(
entity_parser,
file_service: FileService,
link_resolver,
session_maker,
):
"""Test that entities created with disable_permalinks=True don't have permalinks."""
# Create entity service with permalinks disabled
@@ -30,6 +31,7 @@ async def test_create_entity_with_permalinks_disabled(
file_service=file_service,
link_resolver=link_resolver,
app_config=app_config,
session_maker=session_maker,
)
entity_data = EntitySchema(
@@ -64,6 +66,7 @@ async def test_update_entity_with_permalinks_disabled(
entity_parser,
file_service: FileService,
link_resolver,
session_maker,
):
"""Test that entities updated with disable_permalinks=True don't get permalinks added."""
# First create with permalinks enabled
@@ -76,6 +79,7 @@ async def test_update_entity_with_permalinks_disabled(
file_service=file_service,
link_resolver=link_resolver,
app_config=app_config_enabled,
session_maker=session_maker,
)
entity_data = EntitySchema(
@@ -100,6 +104,7 @@ async def test_update_entity_with_permalinks_disabled(
file_service=file_service,
link_resolver=link_resolver,
app_config=app_config_disabled,
session_maker=session_maker,
)
# Update entity with permalinks disabled
@@ -124,6 +129,7 @@ async def test_create_entity_with_content_frontmatter_permalinks_disabled(
entity_parser,
file_service: FileService,
link_resolver,
session_maker,
):
"""Test that content frontmatter permalinks are ignored when disabled."""
# Create entity service with permalinks disabled
@@ -136,6 +142,7 @@ async def test_create_entity_with_content_frontmatter_permalinks_disabled(
file_service=file_service,
link_resolver=link_resolver,
app_config=app_config,
session_maker=session_maker,
)
# Content with frontmatter containing permalink
@@ -180,6 +187,7 @@ async def test_move_entity_with_permalinks_disabled(
file_service: FileService,
link_resolver,
project_config,
session_maker,
):
"""Test that moving an entity with disable_permalinks=True doesn't update permalinks."""
# First create with permalinks enabled
@@ -192,6 +200,7 @@ async def test_move_entity_with_permalinks_disabled(
file_service=file_service,
link_resolver=link_resolver,
app_config=app_config,
session_maker=session_maker,
)
entity_data = EntitySchema(
+8 -7
View File
@@ -90,15 +90,16 @@ async def test_reconcile_projects_with_config_creates_projects_and_default(
_, session_maker = await db.get_or_create_db(
updated.database_path, db_type=db.DatabaseType.FILESYSTEM
)
repo = ProjectRepository(session_maker)
repo = ProjectRepository()
active = await repo.get_active_projects()
names = {p.name for p in active}
assert names.issuperset({"proj-a", "proj-b"})
async with db.scoped_session(session_maker) as session:
active = await repo.get_active_projects(session)
names = {p.name for p in active}
assert names.issuperset({"proj-a", "proj-b"})
default = await repo.get_default_project()
assert default is not None
assert default.name == "proj-b"
default = await repo.get_default_project(session)
assert default is not None
assert default.name == "proj-b"
finally:
await db.shutdown_db()
+210 -187
View File
@@ -8,6 +8,7 @@ import pytest
import pytest_asyncio
from basic_memory import db
from basic_memory.models.knowledge import Entity as EntityModel
from basic_memory.repository import EntityRepository
from basic_memory.schemas.base import Entity as EntitySchema
@@ -78,18 +79,20 @@ async def test_entities(entity_service, file_service):
)
# non markdown entity
e7 = await entity_service.repository.add(
EntityModel(
title="Image.png",
note_type="file",
content_type="image/png",
file_path="Image.png",
permalink="image", # Required for Postgres NOT NULL constraint
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
project_id=entity_service.repository.project_id,
async with db.scoped_session(entity_service.session_maker) as session:
e7 = await entity_service.repository.add(
session,
EntityModel(
title="Image.png",
note_type="file",
content_type="image/png",
file_path="Image.png",
permalink="image", # Required for Postgres NOT NULL constraint
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
project_id=entity_service.repository.project_id,
),
)
)
e8 = await entity_service.create_entity( # duplicate title
EntitySchema(
@@ -104,13 +107,13 @@ async def test_entities(entity_service, file_service):
@pytest_asyncio.fixture
async def link_resolver(entity_repository, search_service, test_entities):
async def link_resolver(entity_repository, search_service, test_entities, session_maker):
"""Create LinkResolver instance with indexed test data."""
# Index all test entities
for entity in test_entities:
await search_service.index_entity(entity)
return LinkResolver(entity_repository, search_service)
return LinkResolver(entity_repository, search_service, session_maker)
@pytest.fixture
@@ -561,33 +564,36 @@ async def test_cross_project_link_resolution(
"""Test resolving explicit cross-project links."""
from basic_memory.repository.project_repository import ProjectRepository
project_repo = ProjectRepository(session_maker)
other_project = await project_repo.create(
{
"name": "other-project",
"description": "Secondary project",
"path": str(tmp_path / "other-project"),
"is_active": True,
"is_default": False,
}
)
now = datetime.now(timezone.utc)
other_entity_repo = EntityRepository(session_maker, project_id=other_project.id)
target = await other_entity_repo.add(
EntityModel(
title="Cross Project Note",
note_type="note",
content_type="text/markdown",
file_path="docs/Cross Project Note.md",
permalink=f"{other_project.permalink}/docs/cross-project-note",
created_at=now,
updated_at=now,
project_id=other_project.id,
project_repo = ProjectRepository()
async with db.scoped_session(session_maker) as session:
other_project = await project_repo.create(
session,
{
"name": "other-project",
"description": "Secondary project",
"path": str(tmp_path / "other-project"),
"is_active": True,
"is_default": False,
},
)
)
resolver = LinkResolver(entity_repository, search_service)
now = datetime.now(timezone.utc)
other_entity_repo = EntityRepository(project_id=other_project.id)
target = await other_entity_repo.add(
session,
EntityModel(
title="Cross Project Note",
note_type="note",
content_type="text/markdown",
file_path="docs/Cross Project Note.md",
permalink=f"{other_project.permalink}/docs/cross-project-note",
created_at=now,
updated_at=now,
project_id=other_project.id,
),
)
resolver = LinkResolver(entity_repository, search_service, session_maker)
resolved = await resolver.resolve_link("other-project::Cross Project Note", strict=True)
assert resolved is not None
@@ -601,7 +607,7 @@ async def test_cross_project_link_resolution(
@pytest_asyncio.fixture
async def context_aware_entities(entity_repository):
async def context_aware_entities(entity_repository, session_maker):
"""Create entities for testing context-aware resolution.
Structure:
@@ -621,123 +627,133 @@ async def context_aware_entities(entity_repository):
now = datetime.now(timezone.utc)
project_id = entity_repository.project_id
# Root level testing.md
e1 = await entity_repository.add(
EntityModel(
title="testing",
note_type="note",
content_type="text/markdown",
file_path="testing.md",
permalink="testing",
created_at=now,
updated_at=now,
project_id=project_id,
async with db.scoped_session(session_maker) as session:
# Root level testing.md
e1 = await entity_repository.add(
session,
EntityModel(
title="testing",
note_type="note",
content_type="text/markdown",
file_path="testing.md",
permalink="testing",
created_at=now,
updated_at=now,
project_id=project_id,
),
)
)
entities.append(e1)
entities.append(e1)
# main/testing/testing.md
e2 = await entity_repository.add(
EntityModel(
title="testing",
note_type="note",
content_type="text/markdown",
file_path="main/testing/testing.md",
permalink="main/testing/testing",
created_at=now,
updated_at=now,
project_id=project_id,
# main/testing/testing.md
e2 = await entity_repository.add(
session,
EntityModel(
title="testing",
note_type="note",
content_type="text/markdown",
file_path="main/testing/testing.md",
permalink="main/testing/testing",
created_at=now,
updated_at=now,
project_id=project_id,
),
)
)
entities.append(e2)
entities.append(e2)
# main/testing/another-test.md
e3 = await entity_repository.add(
EntityModel(
title="another-test",
note_type="note",
content_type="text/markdown",
file_path="main/testing/another-test.md",
permalink="main/testing/another-test",
created_at=now,
updated_at=now,
project_id=project_id,
# main/testing/another-test.md
e3 = await entity_repository.add(
session,
EntityModel(
title="another-test",
note_type="note",
content_type="text/markdown",
file_path="main/testing/another-test.md",
permalink="main/testing/another-test",
created_at=now,
updated_at=now,
project_id=project_id,
),
)
)
entities.append(e3)
entities.append(e3)
# other/testing.md
e4 = await entity_repository.add(
EntityModel(
title="testing",
note_type="note",
content_type="text/markdown",
file_path="other/testing.md",
permalink="other/testing",
created_at=now,
updated_at=now,
project_id=project_id,
# other/testing.md
e4 = await entity_repository.add(
session,
EntityModel(
title="testing",
note_type="note",
content_type="text/markdown",
file_path="other/testing.md",
permalink="other/testing",
created_at=now,
updated_at=now,
project_id=project_id,
),
)
)
entities.append(e4)
entities.append(e4)
# deep/nested/folder/note.md
e5 = await entity_repository.add(
EntityModel(
title="note",
note_type="note",
content_type="text/markdown",
file_path="deep/nested/folder/note.md",
permalink="deep/nested/folder/note",
created_at=now,
updated_at=now,
project_id=project_id,
# deep/nested/folder/note.md
e5 = await entity_repository.add(
session,
EntityModel(
title="note",
note_type="note",
content_type="text/markdown",
file_path="deep/nested/folder/note.md",
permalink="deep/nested/folder/note",
created_at=now,
updated_at=now,
project_id=project_id,
),
)
)
entities.append(e5)
entities.append(e5)
# deep/note.md (for ancestor testing)
e6 = await entity_repository.add(
EntityModel(
title="note",
note_type="note",
content_type="text/markdown",
file_path="deep/note.md",
permalink="deep/note",
created_at=now,
updated_at=now,
project_id=project_id,
# deep/note.md (for ancestor testing)
e6 = await entity_repository.add(
session,
EntityModel(
title="note",
note_type="note",
content_type="text/markdown",
file_path="deep/note.md",
permalink="deep/note",
created_at=now,
updated_at=now,
project_id=project_id,
),
)
)
entities.append(e6)
entities.append(e6)
# note.md at root (for ancestor testing)
e7 = await entity_repository.add(
EntityModel(
title="note",
note_type="note",
content_type="text/markdown",
file_path="note.md",
permalink="note",
created_at=now,
updated_at=now,
project_id=project_id,
# note.md at root (for ancestor testing)
e7 = await entity_repository.add(
session,
EntityModel(
title="note",
note_type="note",
content_type="text/markdown",
file_path="note.md",
permalink="note",
created_at=now,
updated_at=now,
project_id=project_id,
),
)
)
entities.append(e7)
entities.append(e7)
return entities
@pytest_asyncio.fixture
async def context_link_resolver(entity_repository, search_service, context_aware_entities):
async def context_link_resolver(
entity_repository, search_service, context_aware_entities, session_maker
):
"""Create LinkResolver instance with context-aware test data.
Note: We don't index entities for search because these tests focus on
exact title/permalink matching, not fuzzy search. The entities are
database-only records (no files on disk).
"""
return LinkResolver(entity_repository, search_service)
return LinkResolver(entity_repository, search_service, session_maker)
@pytest.mark.asyncio
@@ -903,7 +919,7 @@ async def test_nonexistent_link_with_source_path(context_link_resolver):
@pytest_asyncio.fixture
async def relative_path_entities(entity_repository):
async def relative_path_entities(entity_repository, session_maker):
"""Create entities for testing relative path resolution.
Structure:
@@ -920,73 +936,80 @@ async def relative_path_entities(entity_repository):
now = datetime.now(timezone.utc)
project_id = entity_repository.project_id
# testing/link-test.md (source file)
e1 = await entity_repository.add(
EntityModel(
title="link-test",
note_type="note",
content_type="text/markdown",
file_path="testing/link-test.md",
permalink="testing/link-test",
created_at=now,
updated_at=now,
project_id=project_id,
async with db.scoped_session(session_maker) as session:
# testing/link-test.md (source file)
e1 = await entity_repository.add(
session,
EntityModel(
title="link-test",
note_type="note",
content_type="text/markdown",
file_path="testing/link-test.md",
permalink="testing/link-test",
created_at=now,
updated_at=now,
project_id=project_id,
),
)
)
entities.append(e1)
entities.append(e1)
# testing/nested/deep-note.md (relative target)
e2 = await entity_repository.add(
EntityModel(
title="deep-note",
note_type="note",
content_type="text/markdown",
file_path="testing/nested/deep-note.md",
permalink="testing/nested/deep-note",
created_at=now,
updated_at=now,
project_id=project_id,
# testing/nested/deep-note.md (relative target)
e2 = await entity_repository.add(
session,
EntityModel(
title="deep-note",
note_type="note",
content_type="text/markdown",
file_path="testing/nested/deep-note.md",
permalink="testing/nested/deep-note",
created_at=now,
updated_at=now,
project_id=project_id,
),
)
)
entities.append(e2)
entities.append(e2)
# nested/deep-note.md (absolute path target)
e3 = await entity_repository.add(
EntityModel(
title="deep-note",
note_type="note",
content_type="text/markdown",
file_path="nested/deep-note.md",
permalink="nested/deep-note",
created_at=now,
updated_at=now,
project_id=project_id,
# nested/deep-note.md (absolute path target)
e3 = await entity_repository.add(
session,
EntityModel(
title="deep-note",
note_type="note",
content_type="text/markdown",
file_path="nested/deep-note.md",
permalink="nested/deep-note",
created_at=now,
updated_at=now,
project_id=project_id,
),
)
)
entities.append(e3)
entities.append(e3)
# other/file.md
e4 = await entity_repository.add(
EntityModel(
title="file",
note_type="note",
content_type="text/markdown",
file_path="other/file.md",
permalink="other/file",
created_at=now,
updated_at=now,
project_id=project_id,
# other/file.md
e4 = await entity_repository.add(
session,
EntityModel(
title="file",
note_type="note",
content_type="text/markdown",
file_path="other/file.md",
permalink="other/file",
created_at=now,
updated_at=now,
project_id=project_id,
),
)
)
entities.append(e4)
entities.append(e4)
return entities
@pytest_asyncio.fixture
async def relative_path_resolver(entity_repository, search_service, relative_path_entities):
async def relative_path_resolver(
entity_repository, search_service, relative_path_entities, session_maker
):
"""Create LinkResolver instance with relative path test data."""
return LinkResolver(entity_repository, search_service)
return LinkResolver(entity_repository, search_service, session_maker)
@pytest.mark.asyncio
+40 -41
View File
@@ -37,13 +37,12 @@ async def test_remove_project_with_related_entities(project_service: ProjectServ
# Verify project exists
project = await project_service.get_project(test_project_name)
assert project is not None
session_maker = project_service.session_maker
# Step 2: Create related entities for this project
from basic_memory.repository.entity_repository import EntityRepository
entity_repo = EntityRepository(
project_service.repository.session_maker, project_id=project.id
)
entity_repo = EntityRepository(project_id=project.id)
entity_data = {
"title": "Test Entity for Deletion",
@@ -56,38 +55,35 @@ async def test_remove_project_with_related_entities(project_service: ProjectServ
"created_at": datetime.now(timezone.utc),
"updated_at": datetime.now(timezone.utc),
}
entity = await entity_repo.create(entity_data)
assert entity is not None
# Step 3: Create observations for the entity
from basic_memory.repository.observation_repository import ObservationRepository
obs_repo = ObservationRepository(
project_service.repository.session_maker, project_id=project.id
)
observation_data = {
"entity_id": entity.id,
"content": "This is a test observation",
"category": "note",
}
observation = await obs_repo.create(observation_data)
assert observation is not None
obs_repo = ObservationRepository(project_id=project.id)
# Step 4: Create relations involving the entity
from basic_memory.repository.relation_repository import RelationRepository
rel_repo = RelationRepository(
project_service.repository.session_maker, project_id=project.id
)
rel_repo = RelationRepository(project_id=project.id)
relation_data = {
"from_id": entity.id,
"to_name": "some-target-entity",
"relation_type": "relates-to",
}
relation = await rel_repo.create(relation_data)
assert relation is not None
async with db.scoped_session(session_maker) as session:
entity = await entity_repo.create(session, entity_data)
assert entity is not None
observation_data = {
"entity_id": entity.id,
"content": "This is a test observation",
"category": "note",
}
observation = await obs_repo.create(session, observation_data)
assert observation is not None
relation_data = {
"from_id": entity.id,
"to_name": "some-target-entity",
"relation_type": "relates-to",
}
relation = await rel_repo.create(session, relation_data)
assert relation is not None
# Step 5: Attempt to remove the project
# This should work with proper cascade delete, or fail with foreign key constraint
@@ -100,16 +96,17 @@ async def test_remove_project_with_related_entities(project_service: ProjectServ
assert removed_project is None, "Project should have been removed"
# Related entities should be cascade deleted
remaining_entity = await entity_repo.find_by_id(entity.id)
assert remaining_entity is None, "Entity should have been cascade deleted"
async with db.scoped_session(session_maker) as session:
remaining_entity = await entity_repo.find_by_id(session, entity.id)
assert remaining_entity is None, "Entity should have been cascade deleted"
# Observations should be cascade deleted
remaining_obs = await obs_repo.find_by_id(observation.id)
assert remaining_obs is None, "Observation should have been cascade deleted"
# Observations should be cascade deleted
remaining_obs = await obs_repo.find_by_id(session, observation.id)
assert remaining_obs is None, "Observation should have been cascade deleted"
# Relations should be cascade deleted
remaining_rel = await rel_repo.find_by_id(relation.id)
assert remaining_rel is None, "Relation should have been cascade deleted"
# Relations should be cascade deleted
remaining_rel = await rel_repo.find_by_id(session, relation.id)
assert remaining_rel is None, "Relation should have been cascade deleted"
except Exception as e:
# Check if this is the specific foreign key constraint error from the bug report
@@ -137,7 +134,8 @@ async def test_remove_project_with_related_entities(project_service: ProjectServ
project = await project_service.get_project(test_project_name)
if project:
await project_service.repository.delete(project.id)
async with db.scoped_session(project_service.session_maker) as session:
await project_service.repository.delete(session, project.id)
async def _table_exists(session_maker, table: str) -> bool:
@@ -173,7 +171,7 @@ async def test_remove_project_purges_search_rows(project_service: ProjectService
# Seed both derived tables directly. The bug is in the cleanup path,
# not the indexer, so a synthetic row is enough to prove the sweep.
async with db.scoped_session(project_service.repository.session_maker) as session:
async with db.scoped_session(project_service.session_maker) as session:
await session.execute(
text(
"INSERT INTO search_index "
@@ -212,7 +210,7 @@ async def test_remove_project_purges_search_rows(project_service: ProjectService
},
)
async with db.scoped_session(project_service.repository.session_maker) as session:
async with db.scoped_session(project_service.session_maker) as session:
pre_index = (
await session.execute(
text("SELECT COUNT(*) FROM search_index WHERE project_id = :pid"),
@@ -230,7 +228,7 @@ async def test_remove_project_purges_search_rows(project_service: ProjectService
await project_service.remove_project(test_project_name)
async with db.scoped_session(project_service.repository.session_maker) as session:
async with db.scoped_session(project_service.session_maker) as session:
post_index = (
await session.execute(
text("SELECT COUNT(*) FROM search_index WHERE project_id = :pid"),
@@ -262,7 +260,8 @@ async def test_delete_returns_false_for_missing_project_id(project_service: Proj
branch isn't covered — a silent True would mislead callers into thinking
a non-existent project was removed.
"""
result = await project_service.repository.delete(9_999_999)
async with db.scoped_session(project_service.session_maker) as session:
result = await project_service.repository.delete(session, 9_999_999)
assert result is False
@@ -277,7 +276,7 @@ async def test_remove_project_purges_vector_embeddings(project_service: ProjectS
matches the install path that exercises semantic search.
"""
test_project_name = f"test-vec-cleanup-{os.urandom(4).hex()}"
session_maker = project_service.repository.session_maker
session_maker = project_service.session_maker
# The embeddings table only exists once semantic search has initialized.
# Skipping when it's absent keeps this test honest on minimal CI DBs.
+94 -64
View File
@@ -6,6 +6,8 @@ from pathlib import Path
import pytest
from basic_memory import db
from basic_memory.models.project import Project
from basic_memory.schemas import (
ProjectInfoResponse,
ProjectStatistics,
@@ -16,6 +18,36 @@ from basic_memory.services.project_service import ProjectService
from basic_memory.config import ConfigManager, DatabaseBackend
async def _get_project(project_service: ProjectService, name: str) -> Project | None:
async with db.scoped_session(project_service.session_maker) as session:
return await project_service.repository.get_by_name(session, name)
async def _get_default_project(project_service: ProjectService) -> Project | None:
async with db.scoped_session(project_service.session_maker) as session:
return await project_service.repository.get_default_project(session)
async def _find_projects(project_service: ProjectService) -> list[Project]:
async with db.scoped_session(project_service.session_maker) as session:
return list(await project_service.repository.find_all(session))
async def _create_project(project_service: ProjectService, data: dict) -> Project:
async with db.scoped_session(project_service.session_maker) as session:
return await project_service.repository.create(session, data)
async def _set_default_project(project_service: ProjectService, project_id: int) -> Project | None:
async with db.scoped_session(project_service.session_maker) as session:
return await project_service.repository.set_as_default(session, project_id)
async def _delete_project(project_service: ProjectService, project_id: int) -> bool:
async with db.scoped_session(project_service.session_maker) as session:
return await project_service.repository.delete(session, project_id)
def test_projects_property(project_service: ProjectService):
"""Test the projects property."""
# Get the projects
@@ -208,7 +240,7 @@ async def test_add_project_async(project_service: ProjectService):
assert Path(project_service.projects[test_project_name]) == test_project_path
# Verify it was added to the database
project = await project_service.repository.get_by_name(test_project_name)
project = await _get_project(project_service, test_project_name)
assert project is not None
assert project.name == test_project_name
assert Path(project.path) == test_project_path
@@ -220,7 +252,7 @@ async def test_add_project_async(project_service: ProjectService):
# Ensure it was removed from both config and DB
assert test_project_name not in project_service.projects
project = await project_service.repository.get_by_name(test_project_name)
project = await _get_project(project_service, test_project_name)
assert project is None
@@ -249,20 +281,20 @@ async def test_set_default_project_async(project_service: ProjectService, test_p
assert project_service.default_project == test_project_name
# Verify it's set as default in database
project = await project_service.repository.get_by_name(test_project_name)
project = await _get_project(project_service, test_project_name)
assert project is not None
assert project.is_default is True
# Make sure old default is no longer default
if original_default:
old_default_project = await project_service.repository.get_by_name(original_default)
old_default_project = await _get_project(project_service, original_default)
if old_default_project:
assert old_default_project.is_default is not True
finally:
# Restore original default (only if it exists in database)
if original_default:
original_project = await project_service.repository.get_by_name(original_default)
original_project = await _get_project(project_service, original_default)
if original_project:
await project_service.set_default_project(original_default)
@@ -321,7 +353,7 @@ async def test_set_default_project_config_db_mismatch(
# Verify it's in config but not in database
assert test_project_name in project_service.projects
db_project = await project_service.repository.get_by_name(test_project_name)
db_project = await _get_project(project_service, test_project_name)
assert db_project is None
# Try to set as default - should raise ValueError since project not in database
@@ -350,9 +382,7 @@ async def test_add_project_with_set_default_true(project_service: ProjectService
try:
# Get original default project from database
original_default_project = (
await project_service.repository.get_by_name(original_default)
if original_default
else None
await _get_project(project_service, original_default) if original_default else None
)
# Add project with set_default=True
@@ -363,19 +393,19 @@ async def test_add_project_with_set_default_true(project_service: ProjectService
# Verify new project is set as default in both config and database
assert project_service.default_project == test_project_name
new_project = await project_service.repository.get_by_name(test_project_name)
new_project = await _get_project(project_service, test_project_name)
assert new_project is not None
assert new_project.is_default is True
# Verify original default is no longer default in database
if original_default_project:
assert original_default is not None
refreshed_original = await project_service.repository.get_by_name(original_default)
refreshed_original = await _get_project(project_service, original_default)
assert refreshed_original is not None
assert refreshed_original.is_default is not True
# Verify only one project has is_default=True
all_projects = await project_service.repository.find_all()
all_projects = await _find_projects(project_service)
default_projects = [p for p in all_projects if p.is_default is True]
assert len(default_projects) == 1
assert default_projects[0].name == test_project_name
@@ -383,7 +413,7 @@ async def test_add_project_with_set_default_true(project_service: ProjectService
finally:
# Restore original default (only if it exists in database)
if original_default:
original_project = await project_service.repository.get_by_name(original_default)
original_project = await _get_project(project_service, original_default)
if original_project:
await project_service.set_default_project(original_default)
@@ -415,15 +445,13 @@ async def test_add_project_with_set_default_false(project_service: ProjectServic
assert project_service.default_project == original_default
# Verify new project is NOT set as default
new_project = await project_service.repository.get_by_name(test_project_name)
new_project = await _get_project(project_service, test_project_name)
assert new_project is not None
assert new_project.is_default is not True
# Verify original default is still default
original_default_project = (
await project_service.repository.get_by_name(original_default)
if original_default
else None
await _get_project(project_service, original_default) if original_default else None
)
if original_default_project:
assert original_default_project.is_default is True
@@ -464,20 +492,23 @@ async def test_add_project_promotes_when_config_default_missing_from_db(
config_manager.save_config(fresh_config)
_, session_maker = engine_factory
repo = ProjectRepository(session_maker)
for project in await repo.find_all():
await repo.delete(project.id)
repo = ProjectRepository()
async with db.scoped_session(session_maker) as session:
for project in await repo.find_all(session):
await repo.delete(session, project.id)
file_service = FileService(qa_path, MarkdownProcessor(EntityParser(qa_path)))
service = ProjectService(repository=repo, file_service=file_service)
service = ProjectService(
repository=repo, session_maker=session_maker, file_service=file_service
)
await service.add_project("qa", str(qa_path), set_default=False)
assert service.default_project == "qa"
qa_project = await repo.get_by_name("qa")
qa_project = await _get_project(service, "qa")
assert qa_project is not None
assert qa_project.is_default is True
assert await repo.get_by_name("main") is None
assert await _get_project(service, "main") is None
@pytest.mark.asyncio
@@ -503,12 +534,12 @@ async def test_add_project_preserves_existing_db_default(
# A normally-added project that then becomes the database default,
# while config's default_project still names the fixture project.
await project_service.add_project(surviving_name, surviving_path)
surviving = await project_service.repository.get_by_name(surviving_name)
surviving = await _get_project(project_service, surviving_name)
assert surviving is not None
await project_service.repository.set_as_default(surviving.id)
await _set_default_project(project_service, surviving.id)
# Wedge config: its named default loses its database row.
await project_service.repository.delete(test_project.id)
await _delete_project(project_service, test_project.id)
assert await project_service.get_project(config_default) is None
added_path = str(Path(temp_dir) / "added")
@@ -518,10 +549,10 @@ async def test_add_project_preserves_existing_db_default(
# The surviving database default wins: config is repointed at it and
# the newly added project is not promoted.
assert config_manager.default_project == surviving_name
db_default = await project_service.repository.get_default_project()
db_default = await _get_default_project(project_service)
assert db_default is not None
assert db_default.name == surviving_name
added = await project_service.repository.get_by_name(added_name)
added = await _get_project(project_service, added_name)
assert added is not None
assert added.is_default is not True
@@ -547,7 +578,7 @@ async def test_add_project_default_parameter_omitted(project_service: ProjectSer
assert project_service.default_project == original_default
# Verify new project is NOT set as default
new_project = await project_service.repository.get_by_name(test_project_name)
new_project = await _get_project(project_service, test_project_name)
assert new_project is not None
assert new_project.is_default is not True
@@ -567,10 +598,11 @@ async def test_ensure_single_default_project_enforcement_logic(
assert callable(getattr(project_service, "_ensure_single_default_project"))
# Call the enforcement method - should work without error
await project_service._ensure_single_default_project()
async with db.scoped_session(project_service.session_maker) as session:
await project_service._ensure_single_default_project(session)
# Verify there is exactly one default project after enforcement
all_projects = await project_service.repository.find_all()
all_projects = await _find_projects(project_service)
default_projects = [p for p in all_projects if p.is_default is True]
assert len(default_projects) == 1 # Should have exactly one default
@@ -593,18 +625,18 @@ async def test_synchronize_projects_calls_ensure_single_default(project_service:
# Verify it's in config but not in database
assert test_project_name in project_service.projects
db_project = await project_service.repository.get_by_name(test_project_name)
db_project = await _get_project(project_service, test_project_name)
assert db_project is None
# Call synchronize_projects (this should call _ensure_single_default_project)
await project_service.synchronize_projects()
# Verify project is now in database
db_project = await project_service.repository.get_by_name(test_project_name)
db_project = await _get_project(project_service, test_project_name)
assert db_project is not None
# Verify default project enforcement was applied
all_projects = await project_service.repository.find_all()
all_projects = await _find_projects(project_service)
default_projects = [p for p in all_projects if p.is_default is True]
assert len(default_projects) <= 1 # Should be exactly 1 or 0
@@ -651,16 +683,14 @@ async def test_synchronize_projects_normalizes_project_names(project_service: Pr
assert project_service.projects[expected_normalized_name] == test_project_path
# Verify the project was added to database with normalized name
db_project = await project_service.repository.get_by_name(expected_normalized_name)
db_project = await _get_project(project_service, expected_normalized_name)
assert db_project is not None
assert db_project.name == expected_normalized_name
assert db_project.path == test_project_path
assert db_project.permalink == expected_normalized_name
# Verify the unnormalized name is not in database
unnormalized_db_project = await project_service.repository.get_by_name(
unnormalized_name
)
unnormalized_db_project = await _get_project(project_service, unnormalized_name)
assert unnormalized_db_project is None
finally:
@@ -678,9 +708,9 @@ async def test_synchronize_projects_normalizes_project_names(project_service: Pr
pass
# Remove from database
db_project = await project_service.repository.get_by_name(name)
db_project = await _get_project(project_service, name)
if db_project:
await project_service.repository.delete(db_project.id)
await _delete_project(project_service, db_project.id)
@pytest.mark.asyncio
@@ -703,7 +733,7 @@ async def test_move_project(project_service: ProjectService):
assert test_project_name in project_service.projects
assert Path(project_service.projects[test_project_name]) == old_path
project = await project_service.repository.get_by_name(test_project_name)
project = await _get_project(project_service, test_project_name)
assert project is not None
assert Path(project.path) == old_path
@@ -714,7 +744,7 @@ async def test_move_project(project_service: ProjectService):
assert Path(project_service.projects[test_project_name]) == new_path
# Verify database was updated
updated_project = await project_service.repository.get_by_name(test_project_name)
updated_project = await _get_project(project_service, test_project_name)
assert updated_project is not None
assert Path(updated_project.path) == new_path
@@ -758,7 +788,7 @@ async def test_move_project_db_mismatch(project_service: ProjectService):
# Verify it's in config but not in database
assert test_project_name in project_service.projects
db_project = await project_service.repository.get_by_name(test_project_name)
db_project = await _get_project(project_service, test_project_name)
assert db_project is None
# Try to move project - should fail and restore config
@@ -799,7 +829,7 @@ async def test_move_project_expands_path(project_service: ProjectService):
# Verify the path was expanded to absolute
assert project_service.projects[test_project_name] == expected_absolute_path
updated_project = await project_service.repository.get_by_name(test_project_name)
updated_project = await _get_project(project_service, test_project_name)
assert updated_project is not None
assert updated_project.path == expected_absolute_path
@@ -844,7 +874,7 @@ async def test_synchronize_projects_handles_case_sensitivity_bug(project_service
assert project_service.projects[normalized_name] == test_project_path
# Verify the project exists in database with correct normalized name
db_project = await project_service.repository.get_by_name(normalized_name)
db_project = await _get_project(project_service, normalized_name)
assert db_project is not None
assert db_project.name == normalized_name
assert db_project.path == test_project_path
@@ -868,9 +898,9 @@ async def test_synchronize_projects_handles_case_sensitivity_bug(project_service
except Exception:
pass
db_project = await project_service.repository.get_by_name(name)
db_project = await _get_project(project_service, name)
if db_project:
await project_service.repository.delete(db_project.id)
await _delete_project(project_service, db_project.id)
@pytest.mark.skipif(os.name == "nt", reason="Project root constraints only tested on POSIX systems")
@@ -1352,10 +1382,10 @@ async def test_synchronize_projects_removes_db_only_projects(project_service: Pr
"permalink": test_project_name.lower().replace(" ", "-"),
"is_active": True,
}
await project_service.repository.create(project_data)
await _create_project(project_service, project_data)
# Verify it exists in DB but not in config
db_project = await project_service.repository.get_by_name(test_project_name)
db_project = await _get_project(project_service, test_project_name)
assert db_project is not None
assert test_project_name not in project_service.projects
@@ -1364,7 +1394,7 @@ async def test_synchronize_projects_removes_db_only_projects(project_service: Pr
await project_service.synchronize_projects()
# Verify project was removed from database
db_project_after = await project_service.repository.get_by_name(test_project_name)
db_project_after = await _get_project(project_service, test_project_name)
assert db_project_after is None, (
"Project should be removed from DB when not in config (config is source of truth)"
)
@@ -1374,9 +1404,9 @@ async def test_synchronize_projects_removes_db_only_projects(project_service: Pr
finally:
# Clean up if needed
db_project = await project_service.repository.get_by_name(test_project_name)
db_project = await _get_project(project_service, test_project_name)
if db_project:
await project_service.repository.delete(db_project.id)
await _delete_project(project_service, db_project.id)
@pytest.mark.asyncio
@@ -1404,7 +1434,7 @@ async def test_remove_project_with_delete_notes_false(project_service: ProjectSe
# Verify project is removed from config/db
assert test_project_name not in project_service.projects
db_project = await project_service.repository.get_by_name(test_project_name)
db_project = await _get_project(project_service, test_project_name)
assert db_project is None
# Verify directory and files still exist
@@ -1441,7 +1471,7 @@ async def test_remove_project_with_delete_notes_true(project_service: ProjectSer
# Verify project is removed from config/db
assert test_project_name not in project_service.projects
db_project = await project_service.repository.get_by_name(test_project_name)
db_project = await _get_project(project_service, test_project_name)
assert db_project is None
# Verify directory and files are deleted
@@ -1464,7 +1494,7 @@ async def test_remove_project_delete_notes_missing_directory(project_service: Pr
# Verify project exists in config/db
assert test_project_name in project_service.projects
db_project = await project_service.repository.get_by_name(test_project_name)
db_project = await _get_project(project_service, test_project_name)
assert db_project is not None
# Remove project with delete_notes=True (should not fail even if dir doesn't exist)
@@ -1472,7 +1502,7 @@ async def test_remove_project_delete_notes_missing_directory(project_service: Pr
# Verify project is removed from config/db
assert test_project_name not in project_service.projects
db_project = await project_service.repository.get_by_name(test_project_name)
db_project = await _get_project(project_service, test_project_name)
assert db_project is None
finally:
@@ -1506,7 +1536,7 @@ async def test_remove_project_postgres_backend_uses_database_not_config(
await project_service.add_project(test_project_name, test_project_path, set_default=False)
# Verify project exists and is NOT default in database
db_project = await project_service.repository.get_by_name(test_project_name)
db_project = await _get_project(project_service, test_project_name)
assert db_project is not None
assert db_project.is_default is not True # Should be None or False
@@ -1523,7 +1553,7 @@ async def test_remove_project_postgres_backend_uses_database_not_config(
await project_service.remove_project(test_project_name, delete_notes=False)
# Verify project was removed from database
db_project = await project_service.repository.get_by_name(test_project_name)
db_project = await _get_project(project_service, test_project_name)
assert db_project is None
finally:
@@ -1566,7 +1596,7 @@ async def test_remove_project_local_mode_checks_both_config_and_database(
await project_service.add_project(test_project_name, test_project_path, set_default=False)
# Verify project exists and is NOT default in database
db_project = await project_service.repository.get_by_name(test_project_name)
db_project = await _get_project(project_service, test_project_name)
assert db_project is not None
assert db_project.is_default is not True
@@ -1582,7 +1612,7 @@ async def test_remove_project_local_mode_checks_both_config_and_database(
await project_service.remove_project(test_project_name, delete_notes=False)
# Verify project still exists in database
db_project = await project_service.repository.get_by_name(test_project_name)
db_project = await _get_project(project_service, test_project_name)
assert db_project is not None
finally:
@@ -1619,7 +1649,7 @@ async def test_remove_project_rejects_database_default_in_both_modes(
await project_service.add_project(test_project_name, test_project_path, set_default=True)
# Verify project is default in database
db_project = await project_service.repository.get_by_name(test_project_name)
db_project = await _get_project(project_service, test_project_name)
assert db_project is not None
assert db_project.is_default is True
@@ -1650,15 +1680,15 @@ async def test_remove_project_rejects_database_default_in_both_modes(
# Set original default back in database so we can clean up
if original_default:
original_project = await project_service.repository.get_by_name(original_default)
original_project = await _get_project(project_service, original_default)
if original_project:
await project_service.repository.set_as_default(original_project.id)
await _set_default_project(project_service, original_project.id)
# Cleanup test project
if test_project_name in project_service.projects:
try:
# Clear default in DB first
db_project = await project_service.repository.get_by_name(test_project_name)
db_project = await _get_project(project_service, test_project_name)
if db_project and db_project.is_default:
# Find another project to make default
pass # Let the config_manager handle it
@@ -4,12 +4,15 @@ import os
import pytest
from basic_memory import db
@pytest.mark.asyncio
async def test_get_project_info_supports_db_only_project(
project_service,
project_repository,
config_manager,
session_maker,
):
"""Project info should work when project exists in DB but not local config."""
suffix = os.urandom(4).hex()
@@ -21,14 +24,16 @@ async def test_get_project_info_supports_db_only_project(
config.projects.pop(project_name, None)
config_manager.save_config(config)
await project_repository.create(
{
"name": project_name,
"path": project_path,
"is_active": True,
"is_default": False,
}
)
async with db.scoped_session(session_maker) as session:
await project_repository.create(
session,
{
"name": project_name,
"path": project_path,
"is_active": True,
"is_default": False,
},
)
info = await project_service.get_project_info(project_name)
@@ -6,6 +6,7 @@ from unittest.mock import patch
import pytest
from sqlalchemy import text
from basic_memory import db
from basic_memory.schemas.project_info import EmbeddingStatus
from basic_memory.services.project_service import ProjectService
@@ -14,6 +15,11 @@ def _is_postgres() -> bool:
return os.environ.get("BASIC_MEMORY_TEST_POSTGRES", "").lower() in ("1", "true", "yes")
async def _execute(project_service: ProjectService, query, params=None):
async with db.scoped_session(project_service.session_maker) as session:
return await project_service.repository.execute_query(session, query, params or {})
async def _create_embeddings_stub(project_service: ProjectService) -> None:
"""Create a minimal search_vector_embeddings stub so vector_tables_exist is True.
@@ -21,7 +27,8 @@ async def _create_embeddings_stub(project_service: ProjectService) -> None:
embeddings table is never created. get_embedding_status only probes table
existence and joins on chunk_id (rowid on SQLite), so a plain table suffices.
"""
await project_service.repository.execute_query(
await _execute(
project_service,
text(
"CREATE TABLE IF NOT EXISTS search_vector_embeddings ( chunk_id INTEGER PRIMARY KEY)"
),
@@ -31,9 +38,7 @@ async def _create_embeddings_stub(project_service: ProjectService) -> None:
async def _drop_embeddings_stub(project_service: ProjectService) -> None:
"""Drop the stub table to avoid polluting subsequent tests."""
await project_service.repository.execute_query(
text("DROP TABLE IF EXISTS search_vector_embeddings"), {}
)
await _execute(project_service, text("DROP TABLE IF EXISTS search_vector_embeddings"), {})
@pytest.mark.asyncio
@@ -67,7 +72,7 @@ async def test_embedding_status_vector_tables_missing(
if _is_postgres()
else "DROP TABLE IF EXISTS search_vector_chunks"
)
await project_service.repository.execute_query(text(drop_sql), {})
await _execute(project_service, text(drop_sql), {})
with patch.object(
type(project_service),
@@ -121,13 +126,15 @@ async def test_embedding_status_orphaned_chunks(
"""When chunks exist without matching embeddings, recommend reindex."""
# Insert a chunk row (no matching embedding = orphan)
# Get a real entity_id from the test graph
entity_result = await project_service.repository.execute_query(
entity_result = await _execute(
project_service,
text("SELECT id FROM entity WHERE project_id = :project_id LIMIT 1"),
{"project_id": test_project.id},
)
entity_id = entity_result.scalar()
await project_service.repository.execute_query(
await _execute(
project_service,
text(
"INSERT INTO search_vector_chunks "
"("
@@ -146,7 +153,8 @@ async def test_embedding_status_orphaned_chunks(
# so the LEFT JOIN works and finds the orphan.
# Uses chunk_id as PK — Postgres queries join on chunk_id,
# SQLite queries join on rowid which aliases INTEGER PRIMARY KEY.
await project_service.repository.execute_query(
await _execute(
project_service,
text(
"CREATE TABLE IF NOT EXISTS search_vector_embeddings ( chunk_id INTEGER PRIMARY KEY)"
),
@@ -163,9 +171,7 @@ async def test_embedding_status_orphaned_chunks(
status = await project_service.get_embedding_status(test_project.id)
# Clean up stub table to avoid polluting subsequent tests
await project_service.repository.execute_query(
text("DROP TABLE IF EXISTS search_vector_embeddings"), {}
)
await _execute(project_service, text("DROP TABLE IF EXISTS search_vector_embeddings"), {})
assert status.vector_tables_exist is True
assert status.total_chunks == 1
@@ -191,7 +197,7 @@ async def test_embedding_status_handles_sqlite_vec_unavailable(
# scalar_vec_query returns None when the extension can't be loaded on this
# Python build (e.g. the python.org macOS interpreter). Simulate that here.
async def _vec_query_unavailable(query, params=None):
async def _vec_query_unavailable(_session, query, params=None):
return None
with patch.object(
@@ -221,22 +227,22 @@ async def test_embedding_status_handles_sqlite_vec_unavailable(
async def test_embedding_status_healthy(project_service: ProjectService, test_graph, test_project):
"""When all entities have embeddings, no reindex recommended."""
# Clear any leftover data from prior tests
await project_service.repository.execute_query(text("DELETE FROM search_vector_chunks"), {})
await _execute(project_service, text("DELETE FROM search_vector_chunks"), {})
# Drop any existing virtual table (may have been created by search_service init)
# and recreate as a simple regular table for testing the join logic.
# Uses chunk_id as PK — Postgres queries join on chunk_id,
# SQLite queries join on rowid which aliases INTEGER PRIMARY KEY.
await project_service.repository.execute_query(
text("DROP TABLE IF EXISTS search_vector_embeddings"), {}
)
await project_service.repository.execute_query(
await _execute(project_service, text("DROP TABLE IF EXISTS search_vector_embeddings"), {})
await _execute(
project_service,
text("CREATE TABLE search_vector_embeddings ( chunk_id INTEGER PRIMARY KEY)"),
{},
)
# Insert a chunk + matching embedding for every search_index entity
entity_result = await project_service.repository.execute_query(
entity_result = await _execute(
project_service,
text("SELECT DISTINCT entity_id FROM search_index WHERE project_id = :project_id"),
{"project_id": test_project.id},
)
@@ -244,7 +250,8 @@ async def test_embedding_status_healthy(project_service: ProjectService, test_gr
chunk_id = 1
for eid in entity_ids:
await project_service.repository.execute_query(
await _execute(
project_service,
text(
"INSERT INTO search_vector_chunks "
"("
@@ -263,7 +270,8 @@ async def test_embedding_status_healthy(project_service: ProjectService, test_gr
"key": f"chunk-{chunk_id}",
},
)
await project_service.repository.execute_query(
await _execute(
project_service,
text("INSERT INTO search_vector_embeddings (chunk_id) VALUES (:chunk_id)"),
{"chunk_id": chunk_id},
)
@@ -279,9 +287,7 @@ async def test_embedding_status_healthy(project_service: ProjectService, test_gr
status = await project_service.get_embedding_status(test_project.id)
# Clean up stub table to avoid polluting subsequent tests
await project_service.repository.execute_query(
text("DROP TABLE IF EXISTS search_vector_embeddings"), {}
)
await _execute(project_service, text("DROP TABLE IF EXISTS search_vector_embeddings"), {})
assert status.vector_tables_exist is True
assert status.total_chunks > 0
@@ -307,7 +313,8 @@ async def test_embedding_status_excludes_stale_entity_ids(
# Both vector tables must exist to reach the stale-filtered count queries;
# fixtures run with semantic search disabled, so stub the embeddings table.
await _create_embeddings_stub(project_service)
await project_service.repository.execute_query(
await _execute(
project_service,
text(
"INSERT INTO search_index "
"(id, entity_id, project_id, type, title, permalink, content_stems, "
@@ -329,7 +336,8 @@ async def test_embedding_status_excludes_stale_entity_ids(
# The stale entity_id should NOT be counted in total_indexed_entities.
# Count real entities that have search_index rows (the stale one should be excluded).
real_indexed_result = await project_service.repository.execute_query(
real_indexed_result = await _execute(
project_service,
text(
"SELECT COUNT(DISTINCT si.entity_id) FROM search_index si "
"JOIN entity e ON e.id = si.entity_id "
@@ -6,6 +6,7 @@ from pathlib import Path
import pytest
from basic_memory import db
from basic_memory.services.project_service import ProjectService
@@ -30,19 +31,22 @@ async def test_get_project_from_database(project_service: ProjectService):
"is_active": True,
"is_default": False,
}
await project_service.repository.create(project_data)
async with db.scoped_session(project_service.session_maker) as session:
await project_service.repository.create(session, project_data)
# Verify we can get the project
project = await project_service.repository.get_by_name(test_project_name)
async with db.scoped_session(project_service.session_maker) as session:
project = await project_service.repository.get_by_name(session, test_project_name)
assert project is not None
assert project.name == test_project_name
assert project.path == test_path
finally:
# Clean up
project = await project_service.repository.get_by_name(test_project_name)
if project:
await project_service.repository.delete(project.id)
async with db.scoped_session(project_service.session_maker) as session:
project = await project_service.repository.get_by_name(session, test_project_name)
if project:
await project_service.repository.delete(session, project.id)
@pytest.mark.asyncio
+93 -50
View File
@@ -11,6 +11,21 @@ from basic_memory.schemas.search import SearchQuery, SearchItemType, SearchRetri
from basic_memory.services.search_service import _strip_nul
async def _create_entity(session_maker, entity_repo, data):
async with db.scoped_session(session_maker) as session:
return await entity_repo.create(session, data)
async def _create_observation(session_maker, obs_repo, data):
async with db.scoped_session(session_maker) as session:
return await obs_repo.create(session, data)
async def _get_entity_by_permalink(session_maker, entity_repo, permalink):
async with db.scoped_session(session_maker) as session:
return await entity_repo.get_by_permalink(session, permalink)
@pytest.mark.asyncio
async def test_search_permalink(search_service, test_graph):
"""Exact permalink"""
@@ -860,7 +875,7 @@ async def test_search_by_frontmatter_tags(search_service, session_maker, test_pr
"""Test that entities can be found by searching for their frontmatter tags."""
from basic_memory.repository import EntityRepository
entity_repo = EntityRepository(session_maker, project_id=test_project.id)
entity_repo = EntityRepository(project_id=test_project.id)
# Create entity with tags
from datetime import datetime
@@ -877,7 +892,7 @@ async def test_search_by_frontmatter_tags(search_service, session_maker, test_pr
"updated_at": datetime.now(),
}
entity = await entity_repo.create(entity_data)
entity = await _create_entity(session_maker, entity_repo, entity_data)
await search_service.index_entity(entity, content="")
@@ -912,7 +927,7 @@ async def test_search_by_frontmatter_tags_string_format(
"""Test that entities with string format tags can be found in search."""
from basic_memory.repository import EntityRepository
entity_repo = EntityRepository(session_maker, project_id=test_project.id)
entity_repo = EntityRepository(project_id=test_project.id)
# Create entity with tags in string format
from datetime import datetime
@@ -929,7 +944,7 @@ async def test_search_by_frontmatter_tags_string_format(
"updated_at": datetime.now(),
}
entity = await entity_repo.create(entity_data)
entity = await _create_entity(session_maker, entity_repo, entity_data)
await search_service.index_entity(entity, content="")
@@ -951,7 +966,7 @@ async def test_search_special_characters_in_title(search_service, session_maker,
"""Test that entities with special characters in titles can be searched without FTS5 syntax errors."""
from basic_memory.repository import EntityRepository
entity_repo = EntityRepository(session_maker, project_id=test_project.id)
entity_repo = EntityRepository(project_id=test_project.id)
# Create entities with special characters that could cause FTS5 syntax errors
special_titles = [
@@ -982,7 +997,7 @@ async def test_search_special_characters_in_title(search_service, session_maker,
"updated_at": datetime.now(),
}
entity = await entity_repo.create(entity_data)
entity = await _create_entity(session_maker, entity_repo, entity_data)
entities.append(entity)
# Index all entities
@@ -1008,7 +1023,7 @@ async def test_search_title_with_parentheses_specific(search_service, session_ma
"""Test searching specifically for title with parentheses to reproduce FTS5 error."""
from basic_memory.repository import EntityRepository
entity_repo = EntityRepository(session_maker, project_id=test_project.id)
entity_repo = EntityRepository(project_id=test_project.id)
# Create the problematic entity
from datetime import datetime
@@ -1025,7 +1040,7 @@ async def test_search_title_with_parentheses_specific(search_service, session_ma
"updated_at": datetime.now(),
}
entity = await entity_repo.create(entity_data)
entity = await _create_entity(session_maker, entity_repo, entity_data)
# Index the entity
await search_service.index_entity(entity, content="")
@@ -1044,7 +1059,7 @@ async def test_search_title_via_repository_direct(search_service, session_maker,
"""Test searching via search repository directly to isolate the FTS5 error."""
from basic_memory.repository import EntityRepository
entity_repo = EntityRepository(session_maker, project_id=test_project.id)
entity_repo = EntityRepository(project_id=test_project.id)
# Create the problematic entity
from datetime import datetime
@@ -1061,7 +1076,7 @@ async def test_search_title_via_repository_direct(search_service, session_maker,
"updated_at": datetime.now(),
}
entity = await entity_repo.create(entity_data)
entity = await _create_entity(session_maker, entity_repo, entity_data)
# Index the entity
await search_service.index_entity(entity, content="")
@@ -1093,8 +1108,8 @@ async def test_index_entity_with_duplicate_observations(
from basic_memory.repository import EntityRepository, ObservationRepository
from datetime import datetime
entity_repo = EntityRepository(session_maker, project_id=test_project.id)
obs_repo = ObservationRepository(session_maker, project_id=test_project.id)
entity_repo = EntityRepository(project_id=test_project.id)
obs_repo = ObservationRepository(project_id=test_project.id)
# Create entity
entity_data = {
@@ -1109,19 +1124,23 @@ async def test_index_entity_with_duplicate_observations(
"updated_at": datetime.now(),
}
entity = await entity_repo.create(entity_data)
entity = await _create_entity(session_maker, entity_repo, entity_data)
# Create duplicate observations - same category and content
duplicate_content = "This is a duplicated observation"
await obs_repo.create(
{"entity_id": entity.id, "category": "note", "content": duplicate_content}
await _create_observation(
session_maker,
obs_repo,
{"entity_id": entity.id, "category": "note", "content": duplicate_content},
)
await obs_repo.create(
{"entity_id": entity.id, "category": "note", "content": duplicate_content}
await _create_observation(
session_maker,
obs_repo,
{"entity_id": entity.id, "category": "note", "content": duplicate_content},
)
# Reload entity with observations (get_by_permalink eagerly loads observations)
entity = await entity_repo.get_by_permalink("test/duplicate-obs")
entity = await _get_entity_by_permalink(session_maker, entity_repo, "test/duplicate-obs")
assert entity is not None
# Verify we have duplicate observations
@@ -1149,8 +1168,8 @@ async def test_index_entity_dedupes_observations_by_permalink(
from basic_memory.repository import EntityRepository, ObservationRepository
from datetime import datetime
entity_repo = EntityRepository(session_maker, project_id=test_project.id)
obs_repo = ObservationRepository(session_maker, project_id=test_project.id)
entity_repo = EntityRepository(project_id=test_project.id)
obs_repo = ObservationRepository(project_id=test_project.id)
# Create entity
entity_data = {
@@ -1165,22 +1184,30 @@ async def test_index_entity_dedupes_observations_by_permalink(
"updated_at": datetime.now(),
}
entity = await entity_repo.create(entity_data)
entity = await _create_entity(session_maker, entity_repo, entity_data)
# Create three observations: two duplicates and one unique
duplicate_content = "Duplicate observation content"
unique_content = "Unique observation content"
await obs_repo.create(
{"entity_id": entity.id, "category": "note", "content": duplicate_content}
await _create_observation(
session_maker,
obs_repo,
{"entity_id": entity.id, "category": "note", "content": duplicate_content},
)
await obs_repo.create(
{"entity_id": entity.id, "category": "note", "content": duplicate_content}
await _create_observation(
session_maker,
obs_repo,
{"entity_id": entity.id, "category": "note", "content": duplicate_content},
)
await _create_observation(
session_maker,
obs_repo,
{"entity_id": entity.id, "category": "note", "content": unique_content},
)
await obs_repo.create({"entity_id": entity.id, "category": "note", "content": unique_content})
# Reload entity with observations (get_by_permalink eagerly loads observations)
entity = await entity_repo.get_by_permalink("test/dedupe-test")
entity = await _get_entity_by_permalink(session_maker, entity_repo, "test/dedupe-test")
assert entity is not None
assert len(entity.observations) == 3
@@ -1208,8 +1235,8 @@ async def test_index_entity_multiple_categories_same_content(
from basic_memory.repository import EntityRepository, ObservationRepository
from datetime import datetime
entity_repo = EntityRepository(session_maker, project_id=test_project.id)
obs_repo = ObservationRepository(session_maker, project_id=test_project.id)
entity_repo = EntityRepository(project_id=test_project.id)
obs_repo = ObservationRepository(project_id=test_project.id)
# Create entity
entity_data = {
@@ -1224,15 +1251,23 @@ async def test_index_entity_multiple_categories_same_content(
"updated_at": datetime.now(),
}
entity = await entity_repo.create(entity_data)
entity = await _create_entity(session_maker, entity_repo, entity_data)
# Create observations with same content but different categories
shared_content = "Shared content across categories"
await obs_repo.create({"entity_id": entity.id, "category": "tech", "content": shared_content})
await obs_repo.create({"entity_id": entity.id, "category": "design", "content": shared_content})
await _create_observation(
session_maker,
obs_repo,
{"entity_id": entity.id, "category": "tech", "content": shared_content},
)
await _create_observation(
session_maker,
obs_repo,
{"entity_id": entity.id, "category": "design", "content": shared_content},
)
# Reload entity with observations (get_by_permalink eagerly loads observations)
entity = await entity_repo.get_by_permalink("test/multi-category")
entity = await _get_entity_by_permalink(session_maker, entity_repo, "test/multi-category")
assert entity is not None
assert len(entity.observations) == 2
@@ -1262,8 +1297,8 @@ async def test_index_entity_long_observations_shared_prefix_both_searchable(
from basic_memory.repository import EntityRepository, ObservationRepository
from datetime import datetime
entity_repo = EntityRepository(session_maker, project_id=test_project.id)
obs_repo = ObservationRepository(session_maker, project_id=test_project.id)
entity_repo = EntityRepository(project_id=test_project.id)
obs_repo = ObservationRepository(project_id=test_project.id)
entity_data = {
"title": "Long Observation Collision Entity",
@@ -1276,28 +1311,32 @@ async def test_index_entity_long_observations_shared_prefix_both_searchable(
"created_at": datetime.now(),
"updated_at": datetime.now(),
}
entity = await entity_repo.create(entity_data)
entity = await _create_entity(session_maker, entity_repo, entity_data)
# Identical for the first 210 chars (beyond the 200-char truncation point),
# differing only in the trailing unique marker
shared_prefix = "x" * 210
await obs_repo.create(
await _create_observation(
session_maker,
obs_repo,
{
"entity_id": entity.id,
"category": "note",
"content": f"{shared_prefix} ALPHA_UNIQUE_MARKER",
}
},
)
await obs_repo.create(
await _create_observation(
session_maker,
obs_repo,
{
"entity_id": entity.id,
"category": "note",
"content": f"{shared_prefix} BETA_UNIQUE_MARKER",
}
},
)
# Reload entity with observations (get_by_permalink eagerly loads observations)
entity = await entity_repo.get_by_permalink("test/long-obs-collision")
entity = await _get_entity_by_permalink(session_maker, entity_repo, "test/long-obs-collision")
assert entity is not None
assert len(entity.observations) == 2
@@ -1344,7 +1383,7 @@ async def test_index_entity_markdown_strips_nul_bytes(search_service, session_ma
from basic_memory.repository import EntityRepository
from basic_memory.repository.search_repository import SearchRepository
entity_repo = EntityRepository(session_maker, project_id=test_project.id)
entity_repo = EntityRepository(project_id=test_project.id)
entity_data = {
"title": "NUL Test Entity",
@@ -1357,8 +1396,8 @@ async def test_index_entity_markdown_strips_nul_bytes(search_service, session_ma
"created_at": datetime.now(),
"updated_at": datetime.now(),
}
entity = await entity_repo.create(entity_data)
entity = await entity_repo.get_by_permalink("test/nul-test")
entity = await _create_entity(session_maker, entity_repo, entity_data)
entity = await _get_entity_by_permalink(session_maker, entity_repo, "test/nul-test")
assert entity is not None
# Index with NUL-containing content (simulates rclone-preallocated file)
@@ -1382,7 +1421,7 @@ async def test_reindex_vectors(search_service, session_maker, test_project, monk
from basic_memory.repository.search_repository_base import VectorSyncBatchResult
from datetime import datetime
entity_repo = EntityRepository(session_maker, project_id=test_project.id)
entity_repo = EntityRepository(project_id=test_project.id)
# Test fixtures disable semantic search, and delete_stale_vector_rows is the one call
# in this flow that requires the semantic stack — stub it so the test exercises the
@@ -1402,7 +1441,9 @@ async def test_reindex_vectors(search_service, session_maker, test_project, monk
# Create some entities
created_entity_ids: list[int] = []
for i in range(3):
entity = await entity_repo.create(
entity = await _create_entity(
session_maker,
entity_repo,
{
"title": f"Vector Test Entity {i}",
"note_type": "note",
@@ -1413,7 +1454,7 @@ async def test_reindex_vectors(search_service, session_maker, test_project, monk
"project_id": test_project.id,
"created_at": datetime.now(),
"updated_at": datetime.now(),
}
},
)
created_entity_ids.append(entity.id)
await search_service.index_entity(entity, content=f"Content for entity {i}")
@@ -1468,7 +1509,7 @@ async def test_reindex_vectors_no_callback(
from basic_memory.repository.search_repository_base import VectorSyncBatchResult
from datetime import datetime
entity_repo = EntityRepository(session_maker, project_id=test_project.id)
entity_repo = EntityRepository(project_id=test_project.id)
# Test fixtures disable semantic search, and delete_stale_vector_rows is the one call
# in this flow that requires the semantic stack — stub it so the test exercises the
@@ -1485,7 +1526,9 @@ async def test_reindex_vectors_no_callback(
raising=False,
)
entity = await entity_repo.create(
entity = await _create_entity(
session_maker,
entity_repo,
{
"title": "No Callback Entity",
"note_type": "note",
@@ -1496,7 +1539,7 @@ async def test_reindex_vectors_no_callback(
"project_id": test_project.id,
"created_at": datetime.now(),
"updated_at": datetime.now(),
}
},
)
await search_service.index_entity(entity, content="Test content")
+16 -14
View File
@@ -191,20 +191,22 @@ async def test_embed_opt_out_note_still_participates_in_fts(
search_service, session_maker, test_project
):
"""Per-note semantic opt-out should not remove the note from FTS search."""
entity_repo = EntityRepository(session_maker, project_id=test_project.id)
entity = await entity_repo.create(
{
"title": "FTS Opt Out",
"note_type": "note",
"entity_metadata": {"embed": False},
"content_type": "text/markdown",
"file_path": "test/fts-opt-out.md",
"permalink": "test/fts-opt-out",
"project_id": test_project.id,
"created_at": datetime.now(),
"updated_at": datetime.now(),
}
)
entity_repo = EntityRepository(project_id=test_project.id)
async with db.scoped_session(session_maker) as session:
entity = await entity_repo.create(
session,
{
"title": "FTS Opt Out",
"note_type": "note",
"entity_metadata": {"embed": False},
"content_type": "text/markdown",
"file_path": "test/fts-opt-out.md",
"permalink": "test/fts-opt-out",
"project_id": test_project.id,
"created_at": datetime.now(),
"updated_at": datetime.now(),
},
)
await search_service.index_entity(
entity,
@@ -104,9 +104,9 @@ async def test_update_entity_relations_uses_pk_reload(entity_service: EntityServ
original_find_by_ids = entity_service.repository.find_by_ids
find_by_ids_calls = []
async def spy_find_by_ids(ids):
async def spy_find_by_ids(session, ids):
find_by_ids_calls.append(ids)
return await original_find_by_ids(ids)
return await original_find_by_ids(session, ids)
monkeypatch.setattr(entity_service.repository, "find_by_ids", spy_find_by_ids)
@@ -186,9 +186,12 @@ async def test_upsert_with_relations_uses_lightweight_exact_resolution(
)
await entity_service.upsert_entity_from_markdown(Path(source.file_path), markdown, is_new=False)
assert resolve_calls == [
("Lightweight Target", {"strict": True, "load_relations": False}),
]
assert len(resolve_calls) == 1
link_text, kwargs = resolve_calls[0]
assert link_text == "Lightweight Target"
assert kwargs["strict"] is True
assert kwargs["load_relations"] is False
assert "session" in kwargs
@pytest.mark.asyncio

Some files were not shown because too many files have changed in this diff Show More