From 25b11d38fafdb53f36c70edf4bcad9168af10353 Mon Sep 17 00:00:00 2001 From: sairin1202 <952141617@qq.com> Date: Thu, 25 Jun 2026 04:10:25 +0800 Subject: [PATCH 1/2] refactor(database,app): unify INDEX/MEMORY/SKILL on Resource + Entry backbone Replace the MemoryCategory/MemoryItem/CategoryItem model with a single lane-tagged backbone: Resource (raw inputs and generated lane docs), Entry (searchable atoms), and ResourceEntry (membership edges). Storage, memorize, retrieve, CRUD, and the markdown projection now share one symmetric pipeline parameterized by lane across the inmemory, sqlite, and postgres backends. Also prune historical dead code: legacy JSON prompts (PROMPT_LEGACY), the superseded non-workflow retrieval methods, unused SQL base helpers, and other orphaned wrappers. See docs/adr/0006-unified-resource-entry-lane-backbone.md. Co-authored-by: Cursor --- ...06-unified-resource-entry-lane-backbone.md | 108 ++++ docs/architecture.md | 64 ++- src/memu/app/crud.py | 75 +-- src/memu/app/memorize.py | 294 +++------- src/memu/app/memory_files.py | 2 +- src/memu/app/retrieve.py | 310 ++-------- src/memu/app/service.py | 17 - src/memu/database/__init__.py | 17 +- src/memu/database/inmemory/__init__.py | 7 +- src/memu/database/inmemory/models.py | 38 +- src/memu/database/inmemory/repo.py | 38 +- .../inmemory/repositories/__init__.py | 22 +- .../repositories/category_item_repo.py | 65 --- .../inmemory/repositories/entry_repo.py | 245 ++++++++ .../repositories/memory_category_repo.py | 86 --- .../inmemory/repositories/memory_item_repo.py | 263 --------- .../repositories/resource_entry_repo.py | 65 +++ .../inmemory/repositories/resource_repo.py | 115 +++- src/memu/database/interfaces.py | 22 +- src/memu/database/models.py | 139 +++-- src/memu/database/postgres/__init__.py | 5 +- src/memu/database/postgres/models.py | 75 +-- src/memu/database/postgres/postgres.py | 56 +- .../postgres/repositories/__init__.py | 10 +- .../database/postgres/repositories/base.py | 29 - .../repositories/category_item_repo.py | 143 ----- .../postgres/repositories/entry_repo.py | 367 ++++++++++++ .../repositories/memory_category_repo.py | 162 ------ .../postgres/repositories/memory_item_repo.py | 401 ------------- .../repositories/resource_entry_repo.py | 148 +++++ .../postgres/repositories/resource_repo.py | 200 ++++++- src/memu/database/postgres/schema.py | 39 +- src/memu/database/repositories/__init__.py | 7 +- .../database/repositories/category_item.py | 31 - src/memu/database/repositories/entry.py | 67 +++ .../database/repositories/memory_category.py | 33 -- src/memu/database/repositories/memory_item.py | 60 -- src/memu/database/repositories/resource.py | 57 +- .../database/repositories/resource_entry.py | 31 + src/memu/database/sqlite/__init__.py | 7 + src/memu/database/sqlite/models.py | 58 +- .../database/sqlite/repositories/__init__.py | 10 +- src/memu/database/sqlite/repositories/base.py | 34 +- .../sqlite/repositories/category_item_repo.py | 211 ------- .../sqlite/repositories/entry_repo.py | 393 +++++++++++++ .../repositories/memory_category_repo.py | 260 --------- .../sqlite/repositories/memory_item_repo.py | 544 ------------------ .../repositories/resource_entry_repo.py | 164 ++++++ .../sqlite/repositories/resource_repo.py | 247 +++++--- src/memu/database/sqlite/schema.py | 34 +- src/memu/database/sqlite/sqlite.py | 71 +-- src/memu/database/state.py | 7 +- src/memu/memory_fs/exporter.py | 64 +-- src/memu/prompts/category_summary/category.py | 145 ----- src/memu/prompts/memory_type/behavior.py | 44 -- src/memu/prompts/memory_type/event.py | 56 -- src/memu/prompts/memory_type/knowledge.py | 46 -- src/memu/prompts/memory_type/profile.py | 55 -- src/memu/prompts/memory_type/skill.py | 351 ----------- src/memu/utils/references.py | 35 -- src/memu/utils/tool.py | 26 +- tests/test_backend_conformance.py | 202 ++++--- tests/test_folder_memorize.py | 40 +- tests/test_memory_files.py | 36 +- tests/test_memory_fs_synthesis.py | 25 +- tests/test_openrouter.py | 8 +- tests/test_sqlite.py | 6 +- tests/test_tool_memory.py | 93 +-- 68 files changed, 2845 insertions(+), 4340 deletions(-) create mode 100644 docs/adr/0006-unified-resource-entry-lane-backbone.md delete mode 100644 src/memu/database/inmemory/repositories/category_item_repo.py create mode 100644 src/memu/database/inmemory/repositories/entry_repo.py delete mode 100644 src/memu/database/inmemory/repositories/memory_category_repo.py delete mode 100644 src/memu/database/inmemory/repositories/memory_item_repo.py create mode 100644 src/memu/database/inmemory/repositories/resource_entry_repo.py delete mode 100644 src/memu/database/postgres/repositories/category_item_repo.py create mode 100644 src/memu/database/postgres/repositories/entry_repo.py delete mode 100644 src/memu/database/postgres/repositories/memory_category_repo.py delete mode 100644 src/memu/database/postgres/repositories/memory_item_repo.py create mode 100644 src/memu/database/postgres/repositories/resource_entry_repo.py delete mode 100644 src/memu/database/repositories/category_item.py create mode 100644 src/memu/database/repositories/entry.py delete mode 100644 src/memu/database/repositories/memory_category.py delete mode 100644 src/memu/database/repositories/memory_item.py create mode 100644 src/memu/database/repositories/resource_entry.py delete mode 100644 src/memu/database/sqlite/repositories/category_item_repo.py create mode 100644 src/memu/database/sqlite/repositories/entry_repo.py delete mode 100644 src/memu/database/sqlite/repositories/memory_category_repo.py delete mode 100644 src/memu/database/sqlite/repositories/memory_item_repo.py create mode 100644 src/memu/database/sqlite/repositories/resource_entry_repo.py diff --git a/docs/adr/0006-unified-resource-entry-lane-backbone.md b/docs/adr/0006-unified-resource-entry-lane-backbone.md new file mode 100644 index 00000000..e7cd87d9 --- /dev/null +++ b/docs/adr/0006-unified-resource-entry-lane-backbone.md @@ -0,0 +1,108 @@ +# ADR 0006: Unify INDEX / MEMORY / SKILL onto a Resource + Entry Lane Backbone + +- Status: Proposed +- Date: 2026-06-25 + +## Context + +memU historically modeled structured memory as four record types — `Resource`, +`MemoryItem`, `MemoryCategory`, `CategoryItem` — with retrieval running a fixed +`category -> item -> resource` waterfall. Separately, the read-only `memory_fs` +exporter projected three markdown trees (`INDEX.md`, `MEMORY.md`, `SKILL.md`) +that were decoupled from retrieval, and skills were handled by an ad-hoc dual +track (`memory_type="skill"` items *or* LLM synthesis). + +This produced three asymmetric concepts: + +- INDEX: `Resource.caption` + verbatim `resource/` copies +- MEMORY: `MemoryCategory.summary` + `memory/.md` +- SKILL: synthesized or bypassed, not part of retrieval + +We want INDEX, MEMORY, and SKILL to share **one backbone** with **consistent +storage and retrieval**, all derived from the same per-resource canonical text. + +## Decision + +Collapse the model to **two first-class, lane-tagged entities plus one edge**. + +### Lane + +A `lane` discriminator with three values: `index`, `memory`, `skill`. (Raw +inputs use `lane="source"`.) The three lanes are parallel, structurally +identical processing tracks over a shared trunk; they differ only in *what the +extractor pulls out* and *the entry→resource grouping cardinality*. + +### Entities + +1. **`Resource`** (lane-tagged, one physical table — "everything is a resource"): + - Raw source artifacts (`lane="source"`, `modality` = video/image/audio/ + conversation/document); multimodal preprocessing fills `content` (the + canonical, modality-agnostic text — the shared trunk). + - Generated coarse docs (`lane` ∈ {index, memory, skill}, `modality="markdown"`), + each rendered as a file under the `resource/` root: + - `resource/index/.md` — a description page linking to a raw resource + - `resource/memory/.md` — a category page + - `resource/skill/.md` — a skill page + - Carries `embedding` (for coarse recall) and `resource_refs` provenance back + to the raw sources it derives from. + - This **absorbs the former `MemoryCategory`** (a category is just a + `lane="memory"` markdown resource). + +2. **`Entry`** (lane-tagged, one physical table — the searchable atom): + - index → a resource description; memory → a memory item; skill → a reusable + operation step. + - Carries `text`, `embedding`, `entry_kind` (memory sub-type / step kind), + `extra`, and `source_path` — a back-link to the originating raw resource, + relative to the `resource/` root. (This **generalizes the former + `MemoryItem`**.) + +3. **`ResourceEntry`** (edge): membership of an `Entry` in its coarse lane + `Resource` (memory item ∈ category page, skill step ∈ skill page, description + ∈ index page). Many-to-many. (This **generalizes the former `CategoryItem`**.) + +### Links / provenance + +- `Entry.source_path` points only at the originating raw resource (relative to + the `resource/` root). +- A coarse `Resource`'s provenance (`resource_refs`) is the union of its member + entries' sources, stored redundantly to avoid a query-time join. +- All paths are relative to the single `resource/` root, so the same value works + for both the retrieval API and the exported tree. + +### Pipelines + +- **memorize**: `ingest -> preprocess_multimodal (-> Resource.content) -> + extract_lanes (index/memory/skill extractors) -> embed_entries -> + persist lane resources -> build_response`. +- **retrieve**: for each enabled lane, `Resource` recall (stored embedding, + `where lane=`) → `Entry` recall (stored embedding, `where lane=`), returning a + per-lane shape `{index: {...}, memory: {...}, skill: {...}, resources: [...]}`. + All lanes traverse the same code path; only the `lane` filter differs. + +### Naming + +`lane`, `Resource`, `Entry`, `ResourceEntry`, `content` (canonical text), +`source_path`, `resource_refs`. The former `Doc`/`LaneDoc` concept is dropped: +a "doc" is just a markdown `Resource`, which avoids mislabeling a video/image as +a document. + +## Consequences + +Positive: + +- One storage schema and one retrieval path for all three lanes (true + storage/retrieval consistency). +- Skills become first-class and searchable; the dual-track synthesis/bypass goes + away. +- Every entry and coarse resource is traceable back to its raw source. +- "Everything is a resource" keeps the mental model and the on-disk tree aligned. + +Negative / risk: + +- Breaking schema change: `MemoryCategory` folds into `Resource`; `MemoryItem` → + `Entry`; `CategoryItem` → `ResourceEntry` (field renames included). All three + backends (`inmemory`, `sqlite`, `postgres`) and the app layer (`memorize`, + `retrieve`, `crud`) must be migrated together. +- `retrieve` is category-centric today and needs a substantial rewrite. +- Stored vs query-time category embeddings are unified onto stored embeddings, + changing recall behavior slightly. diff --git a/docs/architecture.md b/docs/architecture.md index 8aae5c06..98048a29 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -8,12 +8,23 @@ The repository also describes a hosted Cloud product in `README.md`, but this do ## System overview -memU implements structured agent memory with four persistent record types: - -- `Resource`: raw source artifacts (conversation/document/image/video/audio) -- `MemoryItem`: extracted atomic memories with embeddings -- `MemoryCategory`: grouped topic summaries -- `CategoryItem`: item-category relation edges +memU implements structured agent memory on a unified, lane-tagged backbone +(see ADR 0006). "Everything is a resource", with two first-class record types +plus one edge: + +- `Resource`: a `lane`-tagged node — either a raw source artifact + (`lane="source"`, modality conversation/document/image/video/audio; the + canonical preprocessing text lives in `content`) or a generated markdown doc + (`lane` ∈ {index, memory, skill}, `modality="markdown"`). A memory-lane + `Resource` is the former "category"; a skill-lane one is a skill page. +- `Entry`: a `lane`-tagged searchable atom with an embedding (index → a + description, memory → a memory item, skill → a step). Links to its origin via + `source_id`/`source_path`. +- `ResourceEntry`: membership edge from an `Entry` to its coarse (lane) `Resource`. + +The three retrieval lanes (index/memory/skill) are parallel, structurally +identical tracks over the shared `Resource.content` trunk; retrieval is the same +`Resource recall → Entry recall` waterfall in every lane, parameterized by `lane`. At runtime, `MemoryService` orchestrates ingestion, retrieval, and manual CRUD over these layers. @@ -23,10 +34,9 @@ flowchart TD B --> C["Workflow Pipelines"] C --> D["LLM Clients"] C --> E["Database Repositories"] - E --> F["Resources"] - E --> G["Memory Items"] - E --> H["Memory Categories"] - E --> I["Category Relations"] + E --> F["Resources (lane: source/index/memory/skill)"] + E --> G["Entries (lane: index/memory/skill)"] + E --> I["ResourceEntry edges"] ``` ## Core runtime components @@ -121,20 +131,23 @@ Key behavior: ### Repository contracts -Storage is abstracted through a `Database` protocol with four repositories: +Storage is abstracted through a `Database` protocol with three repositories: + +- `ResourceRepo` (incl. `get_or_create_doc`, `update_resource`, lane-filtered + `vector_search_resources`) +- `EntryRepo` (incl. `vector_search_entries`, similarity/salience), lane-filtered +- `ResourceEntryRepo` (membership edges) -- `ResourceRepo` (incl. `vector_search_resources`) -- `MemoryItemRepo` (incl. `vector_search_items`, similarity/salience) -- `MemoryCategoryRepo` -- `CategoryItemRepo` +All three repos take an optional `lane` filter so a single physical table per +record type backs every lane (the `lane` column is the discriminator). Vector ranking over **stored** embeddings is a repository responsibility: -`vector_search_items` and `vector_search_resources` keep the retrieval layer from -reaching into any concrete backend. The pure cosine/salience math lives in the -storage-neutral `memu.vector` module (not under any backend), so the app layer -and every backend depend on it instead of on each other. (Category recall still -ranks freshly re-embedded summaries at query time in the retrieval layer, since -that is query-time policy rather than search over stored vectors.) +`vector_search_entries` and `vector_search_resources` keep the retrieval layer +from reaching into any concrete backend. The pure cosine/salience math lives in +the storage-neutral `memu.vector` module (not under any backend), so the app +layer and every backend depend on it instead of on each other. (Memory-doc +recall still ranks freshly re-embedded summaries at query time in the retrieval +layer, since that is query-time policy rather than search over stored vectors.) ### Backends @@ -148,7 +161,7 @@ For Postgres, startup runs migration bootstrap and attempts `CREATE EXTENSION IF ### Scope model propagation -`UserConfig.model` is merged into record/table models so scope fields (for example `user_id`) become first-class columns/attributes across resources, items, categories, and relations. +`UserConfig.model` is merged into record/table models so scope fields (for example `user_id`) become first-class columns/attributes across resources, entries, and resource-entry edges. This is why `where` filters and `user_data` writes are consistently available across APIs. @@ -239,7 +252,7 @@ payload directory: ├── resource/ │ └── ← one copied raw source file (verbatim bytes) ├── memory/ -│ └── .md ← one MemoryCategory (description + summary) +│ └── .md ← one memory-lane Resource (description + summary) └── skill/ └── /SKILL.md ← one synthesized skill per folder ``` @@ -247,8 +260,8 @@ payload directory: - `resource/` holds the raw source files copied verbatim out of the blob store (`Resource.local_path`); `INDEX.md` indexes them (name, modality, description, link), so an agent knows which raw resources exist. -- `memory/.md` is the living memory split one file per `MemoryCategory` - (its description + summary); `MEMORY.md` is an overview that links to each one. +- `memory/.md` is the living memory split one file per memory-lane + `Resource` (its description + summary); `MEMORY.md` is an overview that links to each one. - `skill//SKILL.md` is a reusable skill synthesized from the descriptions (a sibling of `MEMORY.md`, never derived from extracted skill-type memory items); the root `SKILL.md` indexes the tree. @@ -315,3 +328,4 @@ serialized through a per-service lock. - `docs/adr/0001-workflow-pipeline-architecture.md` - `docs/adr/0002-pluggable-storage-and-vector-strategy.md` - `docs/adr/0003-user-scope-in-data-model.md` +- `docs/adr/0006-unified-resource-entry-lane-backbone.md` diff --git a/src/memu/app/crud.py b/src/memu/app/crud.py index 71fcfb72..cfaa2375 100644 --- a/src/memu/app/crud.py +++ b/src/memu/app/crud.py @@ -8,7 +8,7 @@ from pydantic import BaseModel -from memu.database.models import MemoryCategory, MemoryType +from memu.database.models import MemoryType, Resource from memu.prompts.category_patch import CATEGORY_PATCH_PROMPT from memu.workflow.step import WorkflowState, WorkflowStep @@ -229,14 +229,14 @@ def _normalize_where(self, where: Mapping[str, Any] | None) -> dict[str, Any]: def _crud_list_memory_items(self, state: WorkflowState, step_context: Any) -> WorkflowState: where_filters = state.get("where") or {} store = state["store"] - items = store.memory_item_repo.list_items(where_filters) + items = store.entry_repo.list_entries(where_filters, lane="memory") state["items"] = items return state def _crud_list_memory_categories(self, state: WorkflowState, step_context: Any) -> WorkflowState: where_filters = state.get("where") or {} store = state["store"] - categories = store.memory_category_repo.list_categories(where_filters) + categories = store.resource_repo.list_resources(where_filters, lane="memory") state["categories"] = categories return state @@ -261,27 +261,31 @@ def _crud_build_list_categories_response(self, state: WorkflowState, step_contex def _crud_clear_memory_relations(self, state: WorkflowState, step_context: Any) -> WorkflowState: where_filters = state.get("where") or {} store = state["store"] - deleted = store.category_item_repo.clear_relations(where_filters) + deleted = store.resource_entry_repo.clear_relations(where_filters) state["deleted_relations"] = deleted return state def _crud_clear_memory_categories(self, state: WorkflowState, step_context: Any) -> WorkflowState: where_filters = state.get("where") or {} store = state["store"] - deleted = store.memory_category_repo.clear_categories(where_filters) + # Memory-lane docs (the former categories). + deleted = store.resource_repo.clear_resources(where_filters, lane="memory") state["deleted_categories"] = deleted return state def _crud_clear_memory_items(self, state: WorkflowState, step_context: Any) -> WorkflowState: where_filters = state.get("where") or {} store = state["store"] - deleted = store.memory_item_repo.clear_items(where_filters) + # Entries across all lanes. + deleted = store.entry_repo.clear_entries(where_filters) state["deleted_items"] = deleted return state def _crud_clear_memory_resources(self, state: WorkflowState, step_context: Any) -> WorkflowState: where_filters = state.get("where") or {} store = state["store"] + # Remaining resources (raw sources + index/skill docs); memory docs were + # already cleared by _crud_clear_memory_categories. deleted = store.resource_repo.clear_resources(where_filters) state["deleted_resources"] = deleted return state @@ -540,16 +544,18 @@ async def _patch_create_memory_item(self, state: WorkflowState, step_context: An embed_payload = [memory_payload["content"]] content_embedding = (await self._get_step_embedding_client(step_context).embed(embed_payload))[0] - item = store.memory_item_repo.create_item( - memory_type=memory_payload["type"], - summary=memory_payload["content"], + item = store.entry_repo.create_entry( + lane="memory", + source_id=None, + entry_kind=memory_payload["type"], + text=memory_payload["content"], embedding=content_embedding, user_data=dict(user or {}), ) cat_names = memory_payload["categories"] mapped_cat_ids = self._map_category_names_to_ids(cat_names, ctx) for cid in mapped_cat_ids: - store.category_item_repo.link_item_category(item.id, cid, user_data=dict(user or {})) + store.resource_entry_repo.link_entry_resource(item.id, cid, user_data=dict(user or {})) if propagate: category_memory_updates[cid] = (None, memory_payload["content"]) @@ -568,13 +574,13 @@ async def _patch_update_memory_item(self, state: WorkflowState, step_context: An propagate = state["propagate"] category_memory_updates: dict[str, tuple[Any, Any]] = {} - item = store.memory_item_repo.get_item(memory_id) + item = store.entry_repo.get_entry(memory_id) if not item: msg = f"Memory item with id {memory_id} not found" raise ValueError(msg) - old_content = item.summary - old_item_categories = store.category_item_repo.get_item_categories(memory_id) - mapped_old_cat_ids = [cat.category_id for cat in old_item_categories] + old_content = item.text + old_item_categories = store.resource_entry_repo.get_entry_resources(memory_id) + mapped_old_cat_ids = [rel.resource_id for rel in old_item_categories] if memory_payload["content"]: embed_payload = [memory_payload["content"]] @@ -583,10 +589,10 @@ async def _patch_update_memory_item(self, state: WorkflowState, step_context: An content_embedding = None if memory_payload["type"] or memory_payload["content"]: - item = store.memory_item_repo.update_item( - item_id=memory_id, - memory_type=memory_payload["type"], - summary=memory_payload["content"], + item = store.entry_repo.update_entry( + entry_id=memory_id, + entry_kind=memory_payload["type"], + text=memory_payload["content"], embedding=content_embedding, ) self._reconcile_update_categories( @@ -595,7 +601,7 @@ async def _patch_update_memory_item(self, state: WorkflowState, step_context: An mapped_old_cat_ids=mapped_old_cat_ids, content_changed=bool(memory_payload["content"]), old_content=old_content, - new_summary=item.summary, + new_summary=item.text, ctx=ctx, store=store, user=user, @@ -638,11 +644,11 @@ def _reconcile_update_categories( mapped_new_cat_ids = self._map_category_names_to_ids(new_cat_names, ctx) old_set, new_set = set(mapped_old_cat_ids), set(mapped_new_cat_ids) for cid in old_set - new_set: - store.category_item_repo.unlink_item_category(memory_id, cid) + store.resource_entry_repo.unlink_entry_resource(memory_id, cid) if propagate: category_memory_updates[cid] = (old_content, None) for cid in new_set - old_set: - store.category_item_repo.link_item_category(memory_id, cid, user_data=dict(user or {})) + store.resource_entry_repo.link_entry_resource(memory_id, cid, user_data=dict(user or {})) if propagate: category_memory_updates[cid] = (None, new_summary) if propagate and content_changed: @@ -655,18 +661,18 @@ async def _patch_delete_memory_item(self, state: WorkflowState, step_context: An propagate = state["propagate"] category_memory_updates: dict[str, tuple[Any, Any]] = {} - item = store.memory_item_repo.get_item(memory_id) + item = store.entry_repo.get_entry(memory_id) if not item: msg = f"Memory item with id {memory_id} not found" raise ValueError(msg) - item_categories = store.category_item_repo.get_item_categories(memory_id) + item_categories = store.resource_entry_repo.get_entry_resources(memory_id) if propagate: - for cat in item_categories: - category_memory_updates[cat.category_id] = (item.summary, None) + for rel in item_categories: + category_memory_updates[rel.resource_id] = (item.text, None) # Remove the item's category relations first so deleting the item never # leaves orphan edges pointing at a non-existent item. - store.category_item_repo.unlink_item(memory_id) - store.memory_item_repo.delete_item(memory_id) + store.resource_entry_repo.unlink_entry(memory_id) + store.entry_repo.delete_entry(memory_id) state.update({ "memory_item": item, @@ -689,7 +695,9 @@ def _patch_build_response(self, state: WorkflowState, step_context: Any) -> Work item = self._model_dump_without_embeddings(state["memory_item"]) category_updates_ids = list(state.get("category_updates", {}).keys()) category_updates = [ - self._model_dump_without_embeddings(store.memory_category_repo.categories[c]) for c in category_updates_ids + self._model_dump_without_embeddings(res) + for c in category_updates_ids + if (res := store.resource_repo.get_resource(c)) is not None ] response = { "memory_item": item, @@ -724,7 +732,7 @@ async def _patch_category_summaries( target_ids: list[str] = [] client = llm_client or self._get_llm_client() for cid, (content_before, content_after) in updates.items(): - cat = store.memory_category_repo.categories.get(cid) + cat = store.resource_repo.get_resource(cid) if not cat or (not content_before and not content_after): continue prompt = self._build_category_patch_prompt( @@ -739,14 +747,13 @@ async def _patch_category_summaries( need_update, summary = self._parse_category_patch_response(patch) if not need_update: continue - cat = store.memory_category_repo.categories.get(cid) - store.memory_category_repo.update_category( - category_id=cid, + store.resource_repo.update_resource( + resource_id=cid, summary=summary.strip(), ) def _build_category_patch_prompt( - self, *, category: MemoryCategory, content_before: str | None, content_after: str | None + self, *, category: Resource, content_before: str | None, content_after: str | None ) -> str: if content_before and content_after: update_content = "\n".join([ @@ -768,7 +775,7 @@ def _build_category_patch_prompt( original_content = category.summary or "" prompt = CATEGORY_PATCH_PROMPT return prompt.format( - category=self._escape_prompt_value(category.name), + category=self._escape_prompt_value(category.title or ""), original_content=self._escape_prompt_value(original_content or ""), update_content=self._escape_prompt_value(update_content or ""), ) diff --git a/src/memu/app/memorize.py b/src/memu/app/memorize.py index 9d58fe38..417db804 100644 --- a/src/memu/app/memorize.py +++ b/src/memu/app/memorize.py @@ -2,7 +2,6 @@ import asyncio import contextlib -import json import logging import pathlib import re @@ -15,7 +14,7 @@ from memu.app.settings import CategoryConfig, CustomPrompt from memu.blob.folder import diff_folder, load_manifest, manifest_from_scan, save_manifest, scan_folder -from memu.database.models import CategoryItem, MemoryCategory, MemoryItem, MemoryType, Resource +from memu.database.models import Entry, MemoryType, Resource, ResourceEntry from memu.preprocess import PreprocessContext, preprocess_resource from memu.prompts.category_summary import ( CUSTOM_PROMPT as CATEGORY_SUMMARY_CUSTOM_PROMPT, @@ -234,13 +233,13 @@ async def _cascade_delete_by_urls( # Discarded item summaries per category, used to recompute summaries. category_discards: dict[str, list[str]] = {} - for item in store.memory_item_repo.list_items(where=where).values(): - if item.resource_id not in target_ids: + for item in store.entry_repo.list_entries(where=where).values(): + if item.source_id not in target_ids: continue - for relation in store.category_item_repo.get_item_categories(item.id): - store.category_item_repo.unlink_item_category(item.id, relation.category_id) - category_discards.setdefault(relation.category_id, []).append(item.summary) - store.memory_item_repo.delete_item(item.id) + for relation in store.resource_entry_repo.get_entry_resources(item.id): + store.resource_entry_repo.unlink_entry_resource(item.id, relation.resource_id) + category_discards.setdefault(relation.resource_id, []).append(item.text) + store.entry_repo.delete_entry(item.id) for res in targets: store.resource_repo.delete_resource(res.id) @@ -400,7 +399,6 @@ async def _memorize_extract_items(self, state: WorkflowState, step_context: Any) caption = prep.get("caption") structured_entries = await self._generate_structured_entries( - resource_url=res_url, modality=state["modality"], memory_types=state["memory_types"], text=text, @@ -430,8 +428,8 @@ async def _memorize_categorize_items(self, state: WorkflowState, step_context: A modality = state["modality"] local_path = state["local_path"] resources: list[Resource] = [] - items: list[MemoryItem] = [] - relations: list[CategoryItem] = [] + items: list[Entry] = [] + relations: list[ResourceEntry] = [] category_updates: dict[str, list[tuple[str, str]]] = {} user_scope = state.get("user", {}) @@ -441,6 +439,7 @@ async def _memorize_categorize_items(self, state: WorkflowState, step_context: A modality=modality, local_path=local_path, caption=plan.get("caption"), + content=plan.get("text"), store=store, embed_client=embed_client, user=user_scope, @@ -453,6 +452,7 @@ async def _memorize_categorize_items(self, state: WorkflowState, step_context: A mem_items, rels, cat_updates = await self._persist_memory_items( resource_id=res.id, + resource_path=res.source_path, structured_entries=entries, ctx=ctx, store=store, @@ -496,7 +496,9 @@ def _memorize_build_response(self, state: WorkflowState, step_context: Any) -> W relations = [rel.model_dump() for rel in state.get("relations", [])] category_ids = state.get("category_ids") or list(ctx.category_ids) categories = [ - self._model_dump_without_embeddings(store.memory_category_repo.categories[c]) for c in category_ids + self._model_dump_without_embeddings(res) + for c in category_ids + if (res := store.resource_repo.get_resource(c)) is not None ] if len(resources) == 1: @@ -522,25 +524,6 @@ def _segment_resource_url(self, base_url: str, idx: int, total_segments: int) -> path = pathlib.Path(base_url) return f"{path.stem}_#segment_{idx}{path.suffix}" - async def _fetch_and_preprocess_resource( - self, resource_url: str, modality: str, llm_client: Any | None = None - ) -> tuple[str, list[dict[str, str | None]]]: - """ - Fetch and preprocess a resource. - - Returns: - Tuple of (local_path, preprocessed_resources) - where preprocessed_resources is a list of dicts with 'text' and 'caption' - """ - local_path, text = await self.fs.fetch(resource_url, modality) - preprocessed_resources = await self._preprocess_resource_url( - local_path=local_path, - text=text, - modality=modality, - llm_client=llm_client, - ) - return local_path, preprocessed_resources - async def _create_resource_with_caption( self, *, @@ -548,6 +531,7 @@ async def _create_resource_with_caption( modality: str, local_path: str, caption: str | None, + content: str | None = None, store: Database, embed_client: Any | None = None, user: Mapping[str, Any] | None = None, @@ -559,45 +543,21 @@ async def _create_resource_with_caption( else: caption_embedding = None - res = store.resource_repo.create_resource( + return store.resource_repo.create_resource( + lane="source", url=resource_url, modality=modality, local_path=local_path, - caption=caption_text, + content=content, + summary=caption_text, embedding=caption_embedding, user_data=dict(user or {}), ) - # if caption: - # caption_text = caption.strip() - # if caption_text: - # res.caption = caption_text - # client = embed_client or self._get_llm_client() - # res.embedding = (await client.embed([caption_text]))[0] - # res.updated_at = pendulum.now() - return res def _resolve_memory_types(self) -> list[MemoryType]: configured_types = self.memorize_config.memory_types or DEFAULT_MEMORY_TYPES return [cast(MemoryType, mtype) for mtype in configured_types] - def _resolve_summary_prompt(self, modality: str, override: str | None) -> str | None: - memo_settings = self.memorize_config - result = memo_settings.multimodal_preprocess_prompts.get(modality) - if override: - return override - if result is None: - return ( - memo_settings.default_category_summary_prompt - if isinstance(memo_settings.default_category_summary_prompt, str) - else None - ) - return result if isinstance(result, str) else None - - def _resolve_multimodal_preprocess_prompt(self, modality: str) -> str | None: - memo_settings = self.memorize_config - result = memo_settings.multimodal_preprocess_prompts.get(modality) - return result if isinstance(result, str) else None - @staticmethod def _resolve_custom_prompt(prompt: str | CustomPrompt, templates: Mapping[str, str]) -> str: if isinstance(prompt, str): @@ -616,7 +576,6 @@ def _resolve_custom_prompt(prompt: str | CustomPrompt, templates: Mapping[str, s async def _generate_structured_entries( self, *, - resource_url: str, modality: str, memory_types: list[MemoryType], text: str | None, @@ -624,27 +583,18 @@ async def _generate_structured_entries( segments: list[dict[str, int | str]] | None = None, llm_client: Any | None = None, ) -> list[tuple[MemoryType, str, list[str]]]: - if not memory_types: + if not memory_types or not text: return [] client = llm_client or self._get_llm_client() - if text: - entries = await self._generate_text_entries( - resource_text=text, - modality=modality, - memory_types=memory_types, - categories_prompt_str=categories_prompt_str, - segments=segments, - llm_client=client, - ) - return entries - # if entries: - # return entries - # no_result_entry = self._build_no_result_fallback(memory_types[0], resource_url, modality) - # return [no_result_entry] - - return [] - # return self._build_no_text_fallback(memory_types, resource_url, modality) + return await self._generate_text_entries( + resource_text=text, + modality=modality, + memory_types=memory_types, + categories_prompt_str=categories_prompt_str, + segments=segments, + llm_client=client, + ) async def _generate_text_entries( self, @@ -731,11 +681,6 @@ def _parse_structured_entries( entries: list[tuple[MemoryType, str, list[str]]] = [] for mtype, response in zip(memory_types, responses, strict=True): parsed = self._parse_memory_type_response_xml(response) - # if not parsed: - # fallback_entry = response.strip() - # if fallback_entry: - # entries.append((mtype, fallback_entry, [])) - # continue for entry in parsed: content = (entry.get("content") or "").strip() if not content: @@ -755,49 +700,40 @@ def _extract_segment_text(self, lines: list[str], start_idx: int, end_idx: int) segment_lines.append(line) return "\n".join(segment_lines) if segment_lines else None - def _build_no_text_fallback( - self, memory_types: list[MemoryType], resource_url: str, modality: str - ) -> list[tuple[MemoryType, str, list[str]]]: - fallback = f"Resource {resource_url} ({modality}) stored. No text summary in v0." - return [(mtype, f"{fallback} (memory type: {mtype}).", []) for mtype in memory_types] - - def _build_no_result_fallback( - self, memory_type: MemoryType, resource_url: str, modality: str - ) -> tuple[MemoryType, str, list[str]]: - fallback = f"Resource {resource_url} ({modality}) stored. No structured memories generated." - return memory_type, fallback, [] - async def _persist_memory_items( self, *, resource_id: str, + resource_path: str | None, structured_entries: list[tuple[MemoryType, str, list[str]]], ctx: Context, store: Database, embed_client: Any | None = None, user: Mapping[str, Any] | None = None, - ) -> tuple[list[MemoryItem], list[CategoryItem], dict[str, list[tuple[str, str]]]]: + ) -> tuple[list[Entry], list[ResourceEntry], dict[str, list[tuple[str, str]]]]: """ - Persist memory items and track category updates. + Persist memory-lane entries and track memory-doc updates. Returns: - Tuple of (items, relations, category_updates) - where category_updates maps category_id -> list of (item_id, summary) tuples + Tuple of (entries, relations, category_updates) + where category_updates maps memory-doc id -> list of (entry_id, text) tuples """ summary_payloads = [content for _, content, _ in structured_entries] client = embed_client or self._get_embedding_client() item_embeddings = await client.embed(summary_payloads) if summary_payloads else [] - items: list[MemoryItem] = [] - rels: list[CategoryItem] = [] - # Changed: now stores (item_id, summary) tuples for reference support + items: list[Entry] = [] + rels: list[ResourceEntry] = [] + # Stores (entry_id, text) tuples for reference support. category_memory_updates: dict[str, list[tuple[str, str]]] = {} reinforce = self.memorize_config.enable_item_reinforcement for (memory_type, summary_text, cat_names), emb in zip(structured_entries, item_embeddings, strict=True): - item = store.memory_item_repo.create_item( - resource_id=resource_id, - memory_type=memory_type, - summary=summary_text, + item = store.entry_repo.create_entry( + lane="memory", + source_id=resource_id, + source_path=resource_path, + entry_kind=memory_type, + text=summary_text, embedding=emb, user_data=dict(user or {}), reinforce=reinforce, @@ -808,47 +744,31 @@ async def _persist_memory_items( continue mapped_cat_ids = await self._resolve_category_ids(cat_names, ctx, store, user=user) for cid in mapped_cat_ids: - rels.append(store.category_item_repo.link_item_category(item.id, cid, user_data=dict(user or {}))) - # Store (item_id, summary) tuple for reference support + rels.append(store.resource_entry_repo.link_entry_resource(item.id, cid, user_data=dict(user or {}))) + # Store (entry_id, text) tuple for reference support category_memory_updates.setdefault(cid, []).append((item.id, summary_text)) return items, rels, category_memory_updates - def _start_category_initialization(self, ctx: Context, store: Database) -> None: - if ctx.categories_ready: - return - try: - loop = asyncio.get_running_loop() - except RuntimeError: - loop = None - if loop: - ctx.category_init_task = loop.create_task(self._initialize_categories(ctx, store)) - else: - asyncio.run(self._initialize_categories(ctx, store)) - async def _ensure_categories_ready( self, ctx: Context, store: Database, user_scope: Mapping[str, Any] | None = None ) -> None: if ctx.categories_ready: return - if ctx.category_init_task: - await ctx.category_init_task - ctx.category_init_task = None - return await self._initialize_categories(ctx, store, user_scope) @staticmethod def _classify_categories( configs: list[CategoryConfig], - existing_by_name: dict[str, MemoryCategory], + existing_by_name: dict[str, Resource], ) -> tuple[ list[tuple[int, CategoryConfig]], - list[tuple[int, CategoryConfig, MemoryCategory]], - dict[int, MemoryCategory], + list[tuple[int, CategoryConfig, Resource]], + dict[int, Resource], ]: to_create: list[tuple[int, CategoryConfig]] = [] - to_update: list[tuple[int, CategoryConfig, MemoryCategory]] = [] - ready: dict[int, MemoryCategory] = {} + to_update: list[tuple[int, CategoryConfig, Resource]] = [] + ready: dict[int, Resource] = {} for i, cfg in enumerate(configs): name = cfg.name.strip() or "Untitled" description = cfg.description.strip() @@ -871,8 +791,8 @@ async def _initialize_categories( return user_data = dict(user or {}) - existing = store.memory_category_repo.list_categories(where=user_data or None) - existing_by_name: dict[str, MemoryCategory] = {c.name: c for c in existing.values()} + existing = store.resource_repo.list_resources(where=user_data or None, lane="memory") + existing_by_name: dict[str, Resource] = {c.title: c for c in existing.values() if c.title} to_create, to_update, ready = self._classify_categories(self.category_configs, existing_by_name) @@ -887,20 +807,27 @@ async def _initialize_categories( for (i, _), vec in zip(needs_embed, vecs, strict=True): embed_map[i] = vec - cats: dict[int, MemoryCategory] = dict(ready) + from memu.memory_fs.exporter import slugify + + cats: dict[int, Resource] = dict(ready) for i, cfg in to_create: name = cfg.name.strip() or "Untitled" description = cfg.description.strip() - cat = store.memory_category_repo.get_or_create_category( - name=name, description=description, embedding=embed_map[i], user_data=user_data + cat = store.resource_repo.get_or_create_doc( + lane="memory", + title=name, + description=description, + embedding=embed_map[i], + user_data=user_data, + slug=slugify(name), ) cats[i] = cat for i, cfg, ex in to_update: description = cfg.description.strip() - cat = store.memory_category_repo.update_category( - category_id=ex.id, description=description, embedding=embed_map[i] + cat = store.resource_repo.update_resource( + resource_id=ex.id, description=description, embedding=embed_map[i] ) cats[i] = cat @@ -975,10 +902,17 @@ async def _resolve_category_ids( seen: set[str] = set(resolved) if unknown: + from memu.memory_fs.exporter import slugify + vecs = await self._get_embedding_client("embedding").embed(unknown) for name, vec in zip(unknown, vecs, strict=True): - cat = store.memory_category_repo.get_or_create_category( - name=name, description="", embedding=vec, user_data=user_data + cat = store.resource_repo.get_or_create_doc( + lane="memory", + title=name, + description="", + embedding=vec, + user_data=user_data, + slug=slugify(name), ) ctx.category_name_to_id[name.lower()] = cat.id if cat.id not in ctx.category_ids: @@ -1029,31 +963,6 @@ def _format_categories_for_prompt(self, categories: list[CategoryConfig]) -> str lines.append(f"- {name}: {desc}" if desc else f"- {name}") return "Existing categories (reuse when appropriate):\n" + "\n".join(lines) + "\n\n" + adaptive_hint - def _add_conversation_indices(self, conversation: str) -> str: - """ - Add [INDEX] markers to each line of the conversation. - - Args: - conversation: Raw conversation text with lines - - Returns: - Conversation with [INDEX] markers prepended to each non-empty line - """ - lines = conversation.split("\n") - indexed_lines = [] - index = 0 - - for line in lines: - stripped = line.strip() - if stripped: # Only index non-empty lines - indexed_lines.append(f"[{index}] {line}") - index += 1 - else: - # Preserve empty lines without indexing - indexed_lines.append(line) - - return "\n".join(indexed_lines) - def _build_memory_type_prompt(self, *, memory_type: MemoryType, resource_text: str, categories_str: str) -> str: configured_prompt = self.memorize_config.memory_type_prompts.get(memory_type) if configured_prompt is None: @@ -1122,15 +1031,15 @@ async def _persist_item_references( for short_id in referenced_short_ids: matched_item_id = short_id_to_item_id.get(short_id) if matched_item_id: - store.memory_item_repo.update_item( - item_id=matched_item_id, + store.entry_repo.update_entry( + entry_id=matched_item_id, extra={"ref_id": short_id}, ) def _build_category_summary_prompt( self, *, - category: MemoryCategory, + category: Resource, new_memories: list[str] | list[tuple[str, str]], ) -> str: """ @@ -1169,7 +1078,7 @@ def _build_category_summary_prompt( new_items_text = "\n".join(f"- {m}" for m in str_memories if m.strip()) original = category.summary or "" - category_config = self.category_config_map.get(category.name) + category_config = self.category_config_map.get(category.title or "") configured_prompt = ( category_config and category_config.summary_prompt ) or self.memorize_config.default_category_summary_prompt @@ -1183,7 +1092,7 @@ def _build_category_summary_prompt( category_config and category_config.target_length ) or self.memorize_config.default_category_summary_target_length return prompt.format( - category=self._escape_prompt_value(category.name), + category=self._escape_prompt_value(category.title or ""), original_content=self._escape_prompt_value(original or ""), new_memory_items_text=self._escape_prompt_value(new_items_text or "No new memory items."), target_length=target_length, @@ -1209,7 +1118,7 @@ async def _update_category_summaries( target_ids: list[str] = [] client = llm_client or self._get_llm_client() for cid, memories in updates.items(): - cat = store.memory_category_repo.categories.get(cid) + cat = store.resource_repo.get_resource(cid) if not cat or not memories: continue prompt = self._build_category_summary_prompt(category=cat, new_memories=memories) @@ -1219,58 +1128,17 @@ async def _update_category_summaries( return updated_summaries summaries = await asyncio.gather(*tasks) for cid, summary in zip(target_ids, summaries, strict=True): - cat = store.memory_category_repo.categories.get(cid) + cat = store.resource_repo.get_resource(cid) if not cat: continue cleaned_summary = summary.replace("```markdown", "").replace("```", "").strip() - store.memory_category_repo.update_category( - category_id=cid, + store.resource_repo.update_resource( + resource_id=cid, summary=cleaned_summary, ) updated_summaries[cid] = cleaned_summary return updated_summaries - def _parse_conversation_preprocess(self, raw: str) -> tuple[str | None, str | None]: - conversation = self._extract_tag_content(raw, "conversation") - summary = self._extract_tag_content(raw, "summary") - return conversation, summary - - @staticmethod - def _extract_tag_content(raw: str, tag: str) -> str | None: - pattern = re.compile(rf"<{tag}>(.*?)", re.IGNORECASE | re.DOTALL) - match = pattern.search(raw) - if not match: - return None - content = match.group(1).strip() - return content or None - - def _parse_memory_type_response(self, raw: str) -> list[dict[str, Any]]: - if not raw: - return [] - raw = raw.strip() - if not raw: - return [] - payload = None - try: - payload = json.loads(raw) - except json.JSONDecodeError: - try: - blob = self._extract_json_blob(raw) - payload = json.loads(blob) - except Exception: - return [] - if not isinstance(payload, dict): - return [] - items = payload.get("memories_items") - if not isinstance(items, list): - return [] - normalized: list[dict[str, Any]] = [] - for entry in items: - if not isinstance(entry, dict): - continue - normalized.append(entry) - return normalized - def _find_xml_boundaries(self, raw: str) -> tuple[int, int, str] | None: """Find the start index, end index, and closing tag for XML root element.""" root_tags = ["item", "profile", "behaviors", "events", "knowledge", "skills"] diff --git a/src/memu/app/memory_files.py b/src/memu/app/memory_files.py index 9152d821..528d2107 100644 --- a/src/memu/app/memory_files.py +++ b/src/memu/app/memory_files.py @@ -65,7 +65,7 @@ async def build( if is_update: descriptions = MemoryFileExporter._build_descriptions(changed) # type: ignore[arg-type] else: - resources = list(database.resource_repo.list_resources(where=where or None).values()) + resources = list(database.resource_repo.list_resources(where=where or None, lane="source").values()) descriptions = MemoryFileExporter._build_descriptions(resources) # An incremental update merges the changed descriptions into the prior diff --git a/src/memu/app/retrieve.py b/src/memu/app/retrieve.py index 79e83313..7dd4d2d1 100644 --- a/src/memu/app/retrieve.py +++ b/src/memu/app/retrieve.py @@ -30,7 +30,6 @@ class RetrieveMixin: _run_workflow: Callable[..., Awaitable[WorkflowState]] _get_context: Callable[[], Context] _get_database: Callable[[], Database] - _ensure_categories_ready: Callable[[Context, Database], Awaitable[None]] _get_step_llm_client: Callable[[Mapping[str, Any] | None], Any] _get_step_embedding_client: Callable[[Mapping[str, Any] | None], Any] _get_llm_client: Callable[..., Any] @@ -50,7 +49,6 @@ async def retrieve( ctx = self._get_context() store = self._get_database() original_query = self._extract_query_text(queries[-1]) - # await self._ensure_categories_ready(ctx, store) where_filters = self._normalize_where(where) context_queries_objs = queries[:-1] if len(queries) > 1 else [] @@ -268,7 +266,7 @@ async def _rag_route_category(self, state: WorkflowState, step_context: Any) -> embed_client = self._get_step_embedding_client(step_context) store = state["store"] where_filters = state.get("where") or {} - category_pool = store.memory_category_repo.list_categories(where_filters) + category_pool = store.resource_repo.list_resources(where_filters, lane="memory") qvec = (await embed_client.embed([state["active_query"]]))[0] hits, summary_lookup = await self._rank_categories_by_summary( qvec, @@ -297,7 +295,7 @@ async def _rag_category_sufficiency(self, state: WorkflowState, step_context: An retrieved_content = "" store = state["store"] where_filters = state.get("where") or {} - category_pool = state.get("category_pool") or store.memory_category_repo.list_categories(where_filters) + category_pool = state.get("category_pool") or store.resource_repo.list_resources(where_filters, lane="memory") hits = state.get("category_hits") or [] if hits: retrieved_content = self._format_category_content( @@ -322,28 +320,6 @@ async def _rag_category_sufficiency(self, state: WorkflowState, step_context: An state["query_vector"] = (await embed_client.embed([state["active_query"]]))[0] return state - def _extract_referenced_item_ids(self, state: WorkflowState) -> set[str]: - """Extract item IDs from category summary references.""" - from memu.utils.references import extract_references - - category_hits = state.get("category_hits") or [] - summary_lookup = state.get("category_summary_lookup", {}) - category_pool = state.get("category_pool") or {} - referenced_item_ids: set[str] = set() - - for cid, _score in category_hits: - # Get summary from lookup or category - summary = summary_lookup.get(cid) - if not summary: - cat = category_pool.get(cid) - if cat: - summary = cat.summary - if summary: - refs = extract_references(summary) - referenced_item_ids.update(refs) - - return referenced_item_ids - async def _rag_recall_items(self, state: WorkflowState, step_context: Any) -> WorkflowState: if not state.get("retrieve_item") or not state.get("needs_retrieval") or not state.get("proceed_to_items"): state["item_hits"] = [] @@ -351,16 +327,17 @@ async def _rag_recall_items(self, state: WorkflowState, step_context: Any) -> Wo store = state["store"] where_filters = state.get("where") or {} - items_pool = store.memory_item_repo.list_items(where_filters) + items_pool = store.entry_repo.list_entries(where_filters, lane="memory") qvec = state.get("query_vector") if qvec is None: embed_client = self._get_step_embedding_client(step_context) qvec = (await embed_client.embed([state["active_query"]]))[0] state["query_vector"] = qvec - state["item_hits"] = store.memory_item_repo.vector_search_items( + state["item_hits"] = store.entry_repo.vector_search_entries( qvec, self.retrieve_config.item.top_k, where=where_filters, + lane="memory", ranking=self.retrieve_config.item.ranking, recency_decay_days=self.retrieve_config.item.recency_decay_days, ) @@ -377,7 +354,7 @@ async def _rag_item_sufficiency(self, state: WorkflowState, step_context: Any) - store = state["store"] where_filters = state.get("where") or {} - items_pool = state.get("item_pool") or store.memory_item_repo.list_items(where_filters) + items_pool = state.get("item_pool") or store.entry_repo.list_entries(where_filters, lane="memory") retrieved_content = "" hits = state.get("item_hits") or [] if hits: @@ -438,8 +415,10 @@ def _rag_build_context(self, state: WorkflowState, _: Any) -> WorkflowState: if state.get("needs_retrieval"): store = state["store"] where_filters = state.get("where") or {} - categories_pool = state.get("category_pool") or store.memory_category_repo.list_categories(where_filters) - items_pool = state.get("item_pool") or store.memory_item_repo.list_items(where_filters) + categories_pool = state.get("category_pool") or store.resource_repo.list_resources( + where_filters, lane="memory" + ) + items_pool = state.get("item_pool") or store.entry_repo.list_entries(where_filters, lane="memory") resources_pool = state.get("resource_pool") or store.resource_repo.list_resources(where_filters) response["categories"] = self._materialize_hits( state.get("category_hits", []), @@ -576,7 +555,7 @@ async def _llm_route_category(self, state: WorkflowState, step_context: Any) -> llm_client = self._get_step_llm_client(step_context) store = state["store"] where_filters = state.get("where") or {} - category_pool = store.memory_category_repo.list_categories(where_filters) + category_pool = store.resource_repo.list_resources(where_filters, lane="memory") hits = await self._llm_rank_categories( state["active_query"], self.retrieve_config.category.top_k, @@ -636,12 +615,12 @@ async def _llm_recall_items(self, state: WorkflowState, step_context: Any) -> Wo ref_ids.extend(extract_references(summary)) if ref_ids: # Query items by ref_ids - items_pool = store.memory_item_repo.list_items_by_ref_ids(ref_ids, where_filters) + items_pool = store.entry_repo.list_entries_by_ref_ids(ref_ids, where_filters) else: - items_pool = store.memory_item_repo.list_items(where_filters) + items_pool = store.entry_repo.list_entries(where_filters, lane="memory") - relations = store.category_item_repo.list_relations(where_filters) - category_pool = state.get("category_pool") or store.memory_category_repo.list_categories(where_filters) + relations = store.resource_entry_repo.list_relations(where_filters) + category_pool = state.get("category_pool") or store.resource_repo.list_resources(where_filters, lane="memory") state["item_hits"] = await self._llm_rank_items( state["active_query"], self.retrieve_config.item.top_k, @@ -692,7 +671,7 @@ async def _llm_recall_resources(self, state: WorkflowState, step_context: Any) - store = state["store"] where_filters = state.get("where") or {} resource_pool = store.resource_repo.list_resources(where_filters) - items_pool = state.get("item_pool") or store.memory_item_repo.list_items(where_filters) + items_pool = state.get("item_pool") or store.entry_repo.list_entries(where_filters, lane="memory") state["resource_hits"] = await self._llm_rank_resources( state["active_query"], self.retrieve_config.resource.top_k, @@ -733,7 +712,7 @@ async def _rank_categories_by_summary( embed_client: Any | None = None, categories: Mapping[str, Any] | None = None, ) -> tuple[list[tuple[str, float]], dict[str, str]]: - category_pool = categories if categories is not None else store.memory_category_repo.categories + category_pool = categories if categories is not None else store.resource_repo.list_resources(lane="memory") entries = [(cid, cat.summary) for cid, cat in category_pool.items() if cat.summary] if not entries: return [], {} @@ -866,83 +845,6 @@ def _extract_rewritten_query(self, raw: str) -> str | None: return match.group(1).strip() return None - async def _embedding_based_retrieve( - self, - query: str, - top_k: int, - context_queries: list[dict[str, Any]] | None, - ctx: Context, - store: Database, - llm_client: Any | None = None, - embed_client: Any | None = None, - where: Mapping[str, Any] | None = None, - ) -> dict[str, Any]: - """Embedding-based retrieval with query rewriting and judging at each tier""" - where_filters = self._normalize_where(where) - category_pool = store.memory_category_repo.list_categories(where_filters) - items_pool = store.memory_item_repo.list_items(where_filters) - resource_pool = store.resource_repo.list_resources(where_filters) - client = llm_client or self._get_llm_client() - embed = embed_client or self._get_embedding_client() - current_query = query - qvec = (await embed.embed([current_query]))[0] - response: dict[str, Any] = {"resources": [], "items": [], "categories": [], "next_step_query": None} - content_sections: list[str] = [] - - # Tier 1: Categories - cat_hits, summary_lookup = await self._rank_categories_by_summary( - qvec, - top_k, - ctx, - store, - embed_client=embed, - categories=category_pool, - ) - if cat_hits: - response["categories"] = self._materialize_hits(cat_hits, category_pool) - content_sections.append( - self._format_category_content(cat_hits, summary_lookup, store, categories=category_pool) - ) - - needs_more, current_query = await self._decide_if_retrieval_needed( - current_query, - context_queries, - retrieved_content="\n\n".join(content_sections), - llm_client=client, - ) - response["next_step_query"] = current_query - if not needs_more: - return response - # Re-embed with rewritten query - qvec = (await embed.embed([current_query]))[0] - - # Tier 2: Items - item_hits = store.memory_item_repo.vector_search_items(qvec, top_k, where=where_filters) - if item_hits: - response["items"] = self._materialize_hits(item_hits, items_pool) - content_sections.append(self._format_item_content(item_hits, store, items=items_pool)) - - needs_more, current_query = await self._decide_if_retrieval_needed( - current_query, - context_queries, - retrieved_content="\n\n".join(content_sections), - llm_client=client, - ) - response["next_step_query"] = current_query - if not needs_more: - return response - # Re-embed with rewritten query - qvec = (await embed.embed([current_query]))[0] - - # Tier 3: Resources - if resource_pool: - res_hits = store.resource_repo.vector_search_resources(qvec, top_k, where=where_filters) - if res_hits: - response["resources"] = self._materialize_hits(res_hits, resource_pool) - content_sections.append(self._format_resource_content(res_hits, store, resources=resource_pool)) - - return response - def _materialize_hits(self, hits: Sequence[tuple[str, float]], pool: dict[str, Any]) -> list[dict[str, Any]]: out = [] for _id, score in hits: @@ -961,154 +863,28 @@ def _format_category_content( store: Database, categories: Mapping[str, Any] | None = None, ) -> str: - category_pool = categories if categories is not None else store.memory_category_repo.categories + category_pool = categories if categories is not None else store.resource_repo.list_resources(lane="memory") lines = [] for cid, score in hits: cat = category_pool.get(cid) if not cat: continue summary = summaries.get(cid) or cat.summary or "" - lines.append(f"Category: {cat.name}\nSummary: {summary}\nScore: {score:.3f}") + lines.append(f"Category: {cat.title}\nSummary: {summary}\nScore: {score:.3f}") return "\n\n".join(lines).strip() def _format_item_content( self, hits: list[tuple[str, float]], store: Database, items: Mapping[str, Any] | None = None ) -> str: - item_pool = items if items is not None else store.memory_item_repo.items + item_pool = items if items is not None else store.entry_repo.list_entries(lane="memory") lines = [] for iid, score in hits: item = item_pool.get(iid) if not item: continue - lines.append(f"Memory Item ({item.memory_type}): {item.summary}\nScore: {score:.3f}") + lines.append(f"Memory Item ({item.entry_kind}): {item.text}\nScore: {score:.3f}") return "\n\n".join(lines).strip() - def _format_resource_content( - self, hits: list[tuple[str, float]], store: Database, resources: Mapping[str, Any] | None = None - ) -> str: - resource_pool = resources if resources is not None else store.resource_repo.resources - lines = [] - for rid, score in hits: - res = resource_pool.get(rid) - if not res: - continue - caption = res.caption or f"Resource {res.url}" - lines.append(f"Resource: {caption}\nScore: {score:.3f}") - return "\n\n".join(lines).strip() - - def _extract_judgement(self, raw: str) -> str: - if not raw: - return "MORE" - match = re.search(r"(.*?)", raw, re.IGNORECASE | re.DOTALL) - if match: - token = match.group(1).strip().upper() - if "ENOUGH" in token: - return "ENOUGH" - if "MORE" in token: - return "MORE" - upper = raw.strip().upper() - if "ENOUGH" in upper: - return "ENOUGH" - return "MORE" - - async def _llm_based_retrieve( - self, - query: str, - top_k: int, - context_queries: list[dict[str, Any]] | None, - ctx: Context, - store: Database, - llm_client: Any | None = None, - where: Mapping[str, Any] | None = None, - ) -> dict[str, Any]: - """ - LLM-based retrieval that uses language model to search and rank results - in a hierarchical manner, with query rewriting and judging at each tier. - - Flow: - 1. Search categories with LLM, judge + rewrite query - 2. If needs more, search items from relevant categories, judge + rewrite - 3. If needs more, search resources related to context - """ - where_filters = self._normalize_where(where) - category_pool = store.memory_category_repo.list_categories(where_filters) - items_pool = store.memory_item_repo.list_items(where_filters) - relations = store.category_item_repo.list_relations(where_filters) - resource_pool = store.resource_repo.list_resources(where_filters) - current_query = query - client = llm_client or self._get_llm_client() - response: dict[str, Any] = {"resources": [], "items": [], "categories": [], "next_step_query": None} - content_sections: list[str] = [] - - # Tier 1: Search and rank categories - category_hits = await self._llm_rank_categories( - current_query, - top_k, - ctx, - store, - llm_client=client, - categories=category_pool, - ) - if category_hits: - response["categories"] = category_hits - content_sections.append(self._format_llm_category_content(category_hits)) - - needs_more, current_query = await self._decide_if_retrieval_needed( - current_query, - context_queries, - retrieved_content="\n\n".join(content_sections), - llm_client=client, - ) - response["next_step_query"] = current_query - if not needs_more: - return response - - # Tier 2: Search memory items from relevant categories - relevant_category_ids = [cat["id"] for cat in category_hits] - item_hits = await self._llm_rank_items( - current_query, - top_k, - relevant_category_ids, - category_hits, - ctx, - store, - llm_client=client, - categories=category_pool, - items=items_pool, - relations=relations, - ) - if item_hits: - response["items"] = item_hits - content_sections.append(self._format_llm_item_content(item_hits)) - - needs_more, current_query = await self._decide_if_retrieval_needed( - current_query, - context_queries, - retrieved_content="\n\n".join(content_sections), - llm_client=client, - ) - response["next_step_query"] = current_query - if not needs_more: - return response - - # Tier 3: Search resources related to the context - resource_hits = await self._llm_rank_resources( - current_query, - top_k, - category_hits, - item_hits, - ctx, - store, - llm_client=client, - items=items_pool, - resources=resource_pool, - ) - if resource_hits: - response["resources"] = resource_hits - content_sections.append(self._format_llm_resource_content(resource_hits)) - - return response - def _format_categories_for_llm( self, store: Database, @@ -1116,7 +892,9 @@ def _format_categories_for_llm( categories: Mapping[str, Any] | None = None, ) -> str: """Format categories for LLM consumption""" - categories_to_format = categories if categories is not None else store.memory_category_repo.categories + categories_to_format = ( + categories if categories is not None else store.resource_repo.list_resources(lane="memory") + ) if category_ids: categories_to_format = {cid: cat for cid, cat in categories_to_format.items() if cid in category_ids} @@ -1126,7 +904,7 @@ def _format_categories_for_llm( lines = [] for cid, cat in categories_to_format.items(): lines.append(f"ID: {cid}") - lines.append(f"Name: {cat.name}") + lines.append(f"Name: {cat.title}") if cat.description: lines.append(f"Description: {cat.description}") if cat.summary: @@ -1143,16 +921,16 @@ def _format_items_for_llm( relations: Sequence[Any] | None = None, ) -> str: """Format memory items for LLM consumption, optionally filtered by category""" - item_pool = items if items is not None else store.memory_item_repo.items - relation_pool = relations if relations is not None else store.category_item_repo.relations + item_pool = items if items is not None else store.entry_repo.list_entries(lane="memory") + relation_pool = relations if relations is not None else store.resource_entry_repo.relations items_to_format = [] seen_item_ids = set() if category_ids: # Get items that belong to the specified categories for rel in relation_pool: - if rel.category_id in category_ids: - item = item_pool.get(rel.item_id) + if rel.resource_id in category_ids: + item = item_pool.get(rel.entry_id) if item and item.id not in seen_item_ids: items_to_format.append(item) seen_item_ids.add(item.id) @@ -1165,8 +943,8 @@ def _format_items_for_llm( lines = [] for item in items_to_format: lines.append(f"ID: {item.id}") - lines.append(f"Type: {item.memory_type}") - lines.append(f"Summary: {item.summary}") + lines.append(f"Type: {item.entry_kind}") + lines.append(f"Summary: {item.text}") lines.append("---") return "\n".join(lines) @@ -1180,12 +958,12 @@ def _format_resources_for_llm( ) -> str: """Format resources for LLM consumption, optionally filtered by related items""" resource_pool = resources if resources is not None else store.resource_repo.resources - item_pool = items if items is not None else store.memory_item_repo.items + item_pool = items if items is not None else store.entry_repo.list_entries(lane="memory") resources_to_format = [] if item_ids: # Get resources that are related to the specified items - resource_ids = {item_pool[iid].resource_id for iid in item_ids if iid in item_pool and iid is not None} + resource_ids = {item_pool[iid].source_id for iid in item_ids if iid in item_pool and iid is not None} resources_to_format = [ resource_pool[rid] for rid in resource_ids if rid in resource_pool and rid is not None ] @@ -1200,8 +978,8 @@ def _format_resources_for_llm( lines.append(f"ID: {res.id}") lines.append(f"URL: {res.url}") lines.append(f"Modality: {res.modality}") - if res.caption: - lines.append(f"Caption: {res.caption}") + if res.summary: + lines.append(f"Caption: {res.summary}") lines.append("---") return "\n".join(lines) @@ -1216,7 +994,7 @@ async def _llm_rank_categories( categories: Mapping[str, Any] | None = None, ) -> list[dict[str, Any]]: """Use LLM to rank categories based on query relevance""" - category_pool = categories if categories is not None else store.memory_category_repo.categories + category_pool = categories if categories is not None else store.resource_repo.list_resources(lane="memory") if not category_pool: return [] @@ -1249,7 +1027,7 @@ async def _llm_rank_items( logger.debug("[LLM Rank Items] No category_ids provided") return [] - item_pool = items if items is not None else store.memory_item_repo.items + item_pool = items if items is not None else store.entry_repo.list_entries(lane="memory") items_data = self._format_items_for_llm(store, category_ids, items=item_pool, relations=relations) if items_data == "No memory items available.": return [] @@ -1288,7 +1066,7 @@ async def _llm_rank_resources( if not item_ids: return [] - item_pool = items if items is not None else store.memory_item_repo.items + item_pool = items if items is not None else store.entry_repo.list_entries(lane="memory") resource_pool = resources if resources is not None else store.resource_repo.resources resources_data = self._format_resources_for_llm(store, item_ids, items=item_pool, resources=resource_pool) if resources_data == "No resources available.": @@ -1319,7 +1097,7 @@ def _parse_llm_category_response( self, raw_response: str, store: Database, categories: Mapping[str, Any] | None = None ) -> list[dict[str, Any]]: """Parse LLM category ranking response""" - category_pool = categories if categories is not None else store.memory_category_repo.categories + category_pool = categories if categories is not None else store.resource_repo.list_resources(lane="memory") results = [] try: json_blob = self._extract_json_blob(raw_response) @@ -1343,7 +1121,7 @@ def _parse_llm_item_response( self, raw_response: str, store: Database, items: Mapping[str, Any] | None = None ) -> list[dict[str, Any]]: """Parse LLM item ranking response""" - item_pool = items if items is not None else store.memory_item_repo.items + item_pool = items if items is not None else store.entry_repo.list_entries(lane="memory") results = [] try: json_blob = self._extract_json_blob(raw_response) @@ -1401,11 +1179,3 @@ def _format_llm_item_content(self, hits: list[dict[str, Any]]) -> str: for item in hits: lines.append(f"Memory Item ({item['memory_type']}): {item['summary']}") return "\n\n".join(lines).strip() - - def _format_llm_resource_content(self, hits: list[dict[str, Any]]) -> str: - """Format LLM-ranked resource content for judger""" - lines = [] - for res in hits: - caption = res.get("caption", "") or f"Resource {res['url']}" - lines.append(f"Resource: {caption}") - return "\n\n".join(lines).strip() diff --git a/src/memu/app/service.py b/src/memu/app/service.py index 5fad2093..e00feaf6 100644 --- a/src/memu/app/service.py +++ b/src/memu/app/service.py @@ -1,6 +1,5 @@ from __future__ import annotations -import asyncio from collections.abc import Callable, Mapping from dataclasses import dataclass, field from typing import Any, Literal, TypeVar @@ -53,7 +52,6 @@ class Context: categories_ready: bool = False category_ids: list[str] = field(default_factory=list) category_name_to_id: dict[str, str] = field(default_factory=dict) - category_init_task: asyncio.Task | None = None class MemoryService(MemorizeMixin, RetrieveMixin, CRUDMixin): @@ -91,9 +89,6 @@ def __init__( config=self.database_config, user_model=self.user_model, ) - # We need the concrete user scope (user_id: xxx) to initialize the categories - # self._start_category_initialization(self._context, self.database) - # VLM (vision-language) profiles are derived from the LLM profiles so # image/video vision reuses the same provider/credentials with a stronger # multimodal model (see ``vlm_config_from_llm``). @@ -316,18 +311,6 @@ def _get_context(self) -> Context: def _get_database(self) -> Database: return self.database - def _provider_summary(self) -> dict[str, Any]: - vector_provider = None - if self.database_config.vector_index: - vector_provider = self.database_config.vector_index.provider - return { - "llm_profiles": list(self.llm_profiles.profiles.keys()), - "storage": { - "metadata_store": self.database_config.metadata_store.provider, - "vector_index": vector_provider, - }, - } - def _register_pipelines(self) -> None: memo_workflow = self._build_memorize_workflow() memo_initial_keys = self._list_memorize_initial_keys() diff --git a/src/memu/database/__init__.py b/src/memu/database/__init__.py index 8934c5e4..c3dd1c19 100644 --- a/src/memu/database/__init__.py +++ b/src/memu/database/__init__.py @@ -2,22 +2,19 @@ from memu.database.factory import build_database from memu.database.interfaces import ( - CategoryItemRecord, Database, - MemoryCategoryRecord, - MemoryItemRecord, + EntryRecord, + ResourceEntryRecord, ResourceRecord, ) -from memu.database.repositories import CategoryItemRepo, MemoryCategoryRepo, MemoryItemRepo, ResourceRepo +from memu.database.repositories import EntryRepo, ResourceEntryRepo, ResourceRepo __all__ = [ - "CategoryItemRecord", - "CategoryItemRepo", "Database", - "MemoryCategoryRecord", - "MemoryCategoryRepo", - "MemoryItemRecord", - "MemoryItemRepo", + "EntryRecord", + "EntryRepo", + "ResourceEntryRecord", + "ResourceEntryRepo", "ResourceRecord", "ResourceRepo", "build_database", diff --git a/src/memu/database/inmemory/__init__.py b/src/memu/database/inmemory/__init__.py index fadbd111..4e97abd0 100644 --- a/src/memu/database/inmemory/__init__.py +++ b/src/memu/database/inmemory/__init__.py @@ -12,13 +12,12 @@ def build_inmemory_database( config: DatabaseConfig, user_model: type[BaseModel], ) -> InMemoryStore: - resource_model, memory_category_model, memory_item_model, category_item_model = build_inmemory_models(user_model) + resource_model, entry_model, resource_entry_model = build_inmemory_models(user_model) return InMemoryStore( scope_model=user_model, resource_model=resource_model, - memory_item_model=memory_item_model, - memory_category_model=memory_category_model, - category_item_model=category_item_model, + entry_model=entry_model, + resource_entry_model=resource_entry_model, ) diff --git a/src/memu/database/inmemory/models.py b/src/memu/database/inmemory/models.py index 94994a3d..33107c05 100644 --- a/src/memu/database/inmemory/models.py +++ b/src/memu/database/inmemory/models.py @@ -3,10 +3,9 @@ from pydantic import BaseModel from memu.database.models import ( - CategoryItem, - MemoryCategory, - MemoryItem, + Entry, Resource, + ResourceEntry, merge_scope_model, ) @@ -15,40 +14,31 @@ class InMemoryResource(Resource): """Concrete in-memory resource model.""" -class InMemoryMemoryItem(MemoryItem): - """Concrete in-memory memory item model.""" +class InMemoryEntry(Entry): + """Concrete in-memory entry model.""" -class InMemoryMemoryCategory(MemoryCategory): - """Concrete in-memory memory category model.""" - - -class InMemoryCategoryItem(CategoryItem): - """Concrete in-memory relation model.""" +class InMemoryResourceEntry(ResourceEntry): + """Concrete in-memory membership-edge model.""" def build_inmemory_models( user_model: type[BaseModel], ) -> tuple[ type[InMemoryResource], - type[InMemoryMemoryCategory], - type[InMemoryMemoryItem], - type[InMemoryCategoryItem], + type[InMemoryEntry], + type[InMemoryResourceEntry], ]: - """ - Build scoped in-memory models that inherit from both the base interface and the user scope model. - """ + """Build scoped in-memory models inheriting both base interface and user scope.""" resource_model = merge_scope_model(user_model, InMemoryResource, name_suffix="Resource") - memory_category_model = merge_scope_model(user_model, InMemoryMemoryCategory, name_suffix="MemoryCategory") - memory_item_model = merge_scope_model(user_model, InMemoryMemoryItem, name_suffix="MemoryItem") - category_item_model = merge_scope_model(user_model, InMemoryCategoryItem, name_suffix="CategoryItem") - return resource_model, memory_category_model, memory_item_model, category_item_model + entry_model = merge_scope_model(user_model, InMemoryEntry, name_suffix="Entry") + resource_entry_model = merge_scope_model(user_model, InMemoryResourceEntry, name_suffix="ResourceEntry") + return resource_model, entry_model, resource_entry_model __all__ = [ - "InMemoryCategoryItem", - "InMemoryMemoryCategory", - "InMemoryMemoryItem", + "InMemoryEntry", "InMemoryResource", + "InMemoryResourceEntry", "build_inmemory_models", ] diff --git a/src/memu/database/inmemory/repo.py b/src/memu/database/inmemory/repo.py index 44275f9f..beef39c9 100644 --- a/src/memu/database/inmemory/repo.py +++ b/src/memu/database/inmemory/repo.py @@ -6,15 +6,14 @@ from memu.database.inmemory.models import build_inmemory_models from memu.database.inmemory.repositories import ( - InMemoryCategoryItemRepository, - InMemoryMemoryCategoryRepository, - InMemoryMemoryItemRepository, + InMemoryEntryRepository, + InMemoryResourceEntryRepository, InMemoryResourceRepository, ) from memu.database.inmemory.state import InMemoryState from memu.database.interfaces import Database -from memu.database.models import CategoryItem, MemoryCategory, MemoryItem, Resource -from memu.database.repositories import MemoryCategoryRepo, ResourceRepo +from memu.database.models import Entry, Resource, ResourceEntry +from memu.database.repositories import EntryRepo, ResourceEntryRepo, ResourceRepo class InMemoryStore(Database): @@ -23,37 +22,30 @@ def __init__( *, scope_model: type[BaseModel] | None = None, resource_model: type[Any] | None = None, - memory_item_model: type[Any] | None = None, - memory_category_model: type[Any] | None = None, - category_item_model: type[Any] | None = None, + entry_model: type[Any] | None = None, + resource_entry_model: type[Any] | None = None, state: InMemoryState | None = None, ) -> None: self.scope_model = scope_model or BaseModel ( default_resource_model, - default_memory_category_model, - default_memory_item_model, - default_category_item_model, + default_entry_model, + default_resource_entry_model, ) = build_inmemory_models(self.scope_model) self.state = state or InMemoryState() self.resources: dict[str, Resource] = self.state.resources - self.items: dict[str, MemoryItem] = self.state.items - self.categories: dict[str, MemoryCategory] = self.state.categories - self.relations: list[CategoryItem] = self.state.relations + self.entries: dict[str, Entry] = self.state.entries + self.relations: list[ResourceEntry] = self.state.relations resource_model = resource_model or default_resource_model or Resource - memory_item_model = memory_item_model or default_memory_item_model or MemoryItem - memory_category_model = memory_category_model or default_memory_category_model or MemoryCategory - category_item_model = category_item_model or default_category_item_model or CategoryItem + entry_model = entry_model or default_entry_model or Entry + resource_entry_model = resource_entry_model or default_resource_entry_model or ResourceEntry self.resource_repo: ResourceRepo = InMemoryResourceRepository(state=self.state, resource_model=resource_model) - self.memory_category_repo: MemoryCategoryRepo = InMemoryMemoryCategoryRepository( - state=self.state, memory_category_model=memory_category_model - ) - self.memory_item_repo = InMemoryMemoryItemRepository(state=self.state, memory_item_model=memory_item_model) - self.category_item_repo = InMemoryCategoryItemRepository( - state=self.state, category_item_model=category_item_model + self.entry_repo: EntryRepo = InMemoryEntryRepository(state=self.state, entry_model=entry_model) + self.resource_entry_repo: ResourceEntryRepo = InMemoryResourceEntryRepository( + state=self.state, resource_entry_model=resource_entry_model ) def close(self) -> None: diff --git a/src/memu/database/inmemory/repositories/__init__.py b/src/memu/database/inmemory/repositories/__init__.py index 6265ced7..aa500cbb 100644 --- a/src/memu/database/inmemory/repositories/__init__.py +++ b/src/memu/database/inmemory/repositories/__init__.py @@ -1,21 +1,9 @@ -from memu.database.inmemory.repositories.category_item_repo import ( - CategoryItemRepo, - InMemoryCategoryItemRepository, -) -from memu.database.inmemory.repositories.memory_category_repo import ( - InMemoryMemoryCategoryRepository, - MemoryCategoryRepo, -) -from memu.database.inmemory.repositories.memory_item_repo import InMemoryMemoryItemRepository, MemoryItemRepo -from memu.database.inmemory.repositories.resource_repo import InMemoryResourceRepository, ResourceRepo +from memu.database.inmemory.repositories.entry_repo import InMemoryEntryRepository +from memu.database.inmemory.repositories.resource_entry_repo import InMemoryResourceEntryRepository +from memu.database.inmemory.repositories.resource_repo import InMemoryResourceRepository __all__ = [ - "CategoryItemRepo", - "InMemoryCategoryItemRepository", - "InMemoryMemoryCategoryRepository", - "InMemoryMemoryItemRepository", + "InMemoryEntryRepository", + "InMemoryResourceEntryRepository", "InMemoryResourceRepository", - "MemoryCategoryRepo", - "MemoryItemRepo", - "ResourceRepo", ] diff --git a/src/memu/database/inmemory/repositories/category_item_repo.py b/src/memu/database/inmemory/repositories/category_item_repo.py deleted file mode 100644 index c83a85cf..00000000 --- a/src/memu/database/inmemory/repositories/category_item_repo.py +++ /dev/null @@ -1,65 +0,0 @@ -from __future__ import annotations - -import uuid -from collections.abc import Mapping -from typing import Any, override - -from memu.database.inmemory.repositories.filter import matches_where -from memu.database.inmemory.state import InMemoryState -from memu.database.models import CategoryItem -from memu.database.repositories.category_item import CategoryItemRepo - - -class InMemoryCategoryItemRepository(CategoryItemRepo): - def __init__(self, *, state: InMemoryState, category_item_model: type[CategoryItem]) -> None: - self._state = state - self.category_item_model = category_item_model - self.relations: list[CategoryItem] = self._state.relations - - def list_relations(self, where: Mapping[str, Any] | None = None) -> list[CategoryItem]: - if not where: - return list(self.relations) - return [rel for rel in self.relations if matches_where(rel, where)] - - def link_item_category(self, item_id: str, cat_id: str, user_data: dict[str, Any]) -> CategoryItem: - _ = item_id # enforced by caller via existing state - for rel in self.relations: - if rel.item_id == item_id and rel.category_id == cat_id: - return rel - rel = self.category_item_model(id=str(uuid.uuid4()), item_id=item_id, category_id=cat_id, **user_data) - self.relations.append(rel) - return rel - - def load_existing(self) -> None: - return None - - @override - def get_item_categories(self, item_id: str) -> list[CategoryItem]: - return [rel for rel in self.relations if rel.item_id == item_id] - - @override - def unlink_item_category(self, item_id: str, cat_id: str) -> None: - # Mutate the shared state list in place so the DatabaseState reference and - # this repo's view never diverge (rebinding self.relations would orphan the - # shared state.relations list). - self.relations[:] = [ - rel for rel in self.relations if not (rel.item_id == item_id and rel.category_id == cat_id) - ] - - def unlink_item(self, item_id: str) -> list[CategoryItem]: - removed = [rel for rel in self.relations if rel.item_id == item_id] - self.relations[:] = [rel for rel in self.relations if rel.item_id != item_id] - return removed - - def clear_relations(self, where: Mapping[str, Any] | None = None) -> list[CategoryItem]: - if not where: - removed = list(self.relations) - self.relations.clear() - return removed - removed = [rel for rel in self.relations if matches_where(rel, where)] - removed_ids = {rel.id for rel in removed} - self.relations[:] = [rel for rel in self.relations if rel.id not in removed_ids] - return removed - - -__all__ = ["InMemoryCategoryItemRepository"] diff --git a/src/memu/database/inmemory/repositories/entry_repo.py b/src/memu/database/inmemory/repositories/entry_repo.py new file mode 100644 index 00000000..b0c667fa --- /dev/null +++ b/src/memu/database/inmemory/repositories/entry_repo.py @@ -0,0 +1,245 @@ +from __future__ import annotations + +import uuid +from collections.abc import Mapping +from typing import Any, override + +import pendulum + +from memu.database.inmemory.repositories.filter import matches_where +from memu.database.inmemory.state import InMemoryState +from memu.database.models import Entry, compute_content_hash +from memu.database.repositories.entry import EntryRepo +from memu.vector import cosine_topk, cosine_topk_salience + + +class InMemoryEntryRepository(EntryRepo): + def __init__(self, *, state: InMemoryState, entry_model: type[Entry]) -> None: + self._state = state + self.entry_model = entry_model + self.entries: dict[str, Entry] = self._state.entries + + def list_entries( + self, where: Mapping[str, Any] | None = None, *, lane: str | None = None + ) -> dict[str, Entry]: + result = self.entries if not where else { + eid: entry for eid, entry in self.entries.items() if matches_where(entry, where) + } + if lane is not None: + result = {eid: entry for eid, entry in result.items() if getattr(entry, "lane", None) == lane} + return dict(result) + + def list_entries_by_ref_ids( + self, ref_ids: list[str], where: Mapping[str, Any] | None = None + ) -> dict[str, Entry]: + if not ref_ids: + return {} + ref_id_set = set(ref_ids) + result: dict[str, Entry] = {} + for eid, entry in self.entries.items(): + if where and not matches_where(entry, where): + continue + entry_ref_id = (entry.extra or {}).get("ref_id") + if entry_ref_id and entry_ref_id in ref_id_set: + result[eid] = entry + return result + + def clear_entries( + self, where: Mapping[str, Any] | None = None, *, lane: str | None = None + ) -> dict[str, Entry]: + if not where and lane is None: + matches = self.entries.copy() + self.entries.clear() + return matches + matches = self.list_entries(where, lane=lane) + for eid in matches: + self.entries.pop(eid, None) + return matches + + def _find_by_hash(self, content_hash: str, user_data: dict[str, Any]) -> Entry | None: + for entry in self.entries.values(): + entry_hash = (entry.extra or {}).get("content_hash") + if entry_hash != content_hash: + continue + if matches_where(entry, user_data): + return entry + return None + + def create_entry( + self, + *, + lane: str, + source_id: str | None, + entry_kind: str, + text: str, + embedding: list[float], + user_data: dict[str, Any], + source_path: str | None = None, + reinforce: bool = False, + tool_record: dict[str, Any] | None = None, + ) -> Entry: + if reinforce and entry_kind != "tool": + return self.create_entry_reinforce( + lane=lane, + source_id=source_id, + entry_kind=entry_kind, + text=text, + embedding=embedding, + user_data=user_data, + source_path=source_path, + ) + + extra: dict[str, Any] = {} + if tool_record: + for key in ("when_to_use", "metadata", "tool_calls"): + if tool_record.get(key) is not None: + extra[key] = tool_record[key] + + eid = str(uuid.uuid4()) + entry = self.entry_model( + id=eid, + lane=lane, + source_id=source_id, + source_path=source_path, + entry_kind=entry_kind, + text=text, + embedding=embedding, + extra=extra if extra else {}, + **user_data, + ) + self.entries[eid] = entry + return entry + + def create_entry_reinforce( + self, + *, + lane: str, + source_id: str | None, + entry_kind: str, + text: str, + embedding: list[float], + user_data: dict[str, Any], + source_path: str | None = None, + ) -> Entry: + content_hash = compute_content_hash(text, entry_kind) + existing = self._find_by_hash(content_hash, user_data) + if existing: + current_extra = existing.extra or {} + current_count = current_extra.get("reinforcement_count", 1) + existing.extra = { + **current_extra, + "reinforcement_count": current_count + 1, + "last_reinforced_at": pendulum.now("UTC").isoformat(), + } + existing.updated_at = pendulum.now("UTC") + return existing + + eid = str(uuid.uuid4()) + now = pendulum.now("UTC") + entry_extra = user_data.pop("extra", {}) if "extra" in user_data else {} + entry_extra.update({ + "content_hash": content_hash, + "reinforcement_count": 1, + "last_reinforced_at": now.isoformat(), + }) + entry = self.entry_model( + id=eid, + lane=lane, + source_id=source_id, + source_path=source_path, + entry_kind=entry_kind, + text=text, + embedding=embedding, + extra=entry_extra, + **user_data, + ) + self.entries[eid] = entry + return entry + + def vector_search_entries( + self, + query_vec: list[float], + top_k: int, + where: Mapping[str, Any] | None = None, + *, + lane: str | None = None, + ranking: str = "similarity", + recency_decay_days: float = 30.0, + ) -> list[tuple[str, float]]: + pool = self.list_entries(where, lane=lane) + + if ranking == "salience": + corpus = [ + ( + e.id, + e.embedding, + (e.extra or {}).get("reinforcement_count", 1), + self._parse_datetime((e.extra or {}).get("last_reinforced_at")), + ) + for e in pool.values() + ] + return cosine_topk_salience(query_vec, corpus, k=top_k, recency_decay_days=recency_decay_days) + + return cosine_topk(query_vec, [(e.id, e.embedding) for e in pool.values()], k=top_k) + + def load_existing(self) -> None: + return None + + def get_entry(self, entry_id: str) -> Entry | None: + return self.entries.get(entry_id) + + @staticmethod + def _parse_datetime(dt_str: str | None) -> pendulum.DateTime | None: + if dt_str is None: + return None + try: + parsed = pendulum.parse(dt_str) + except (ValueError, TypeError): + return None + else: + if isinstance(parsed, pendulum.DateTime): + return parsed + return None + + @override + def delete_entry(self, entry_id: str) -> None: + self.entries.pop(entry_id, None) + + @override + def update_entry( + self, + *, + entry_id: str, + entry_kind: str | None = None, + text: str | None = None, + embedding: list[float] | None = None, + extra: dict[str, Any] | None = None, + tool_record: dict[str, Any] | None = None, + ) -> Entry: + entry = self.entries.get(entry_id) + if entry is None: + msg = f"Entry with id {entry_id} not found" + raise KeyError(msg) + + if entry_kind is not None: + entry.entry_kind = entry_kind + if text is not None: + entry.text = text + if embedding is not None: + entry.embedding = embedding + + current_extra = entry.extra or {} + if extra is not None: + current_extra = {**current_extra, **extra} + if tool_record is not None: + for key in ("when_to_use", "metadata", "tool_calls"): + if tool_record.get(key) is not None: + current_extra[key] = tool_record[key] + if extra is not None or tool_record is not None: + entry.extra = current_extra + + self.entries[entry_id] = entry + return entry + + +__all__ = ["InMemoryEntryRepository"] diff --git a/src/memu/database/inmemory/repositories/memory_category_repo.py b/src/memu/database/inmemory/repositories/memory_category_repo.py deleted file mode 100644 index cb7e1d45..00000000 --- a/src/memu/database/inmemory/repositories/memory_category_repo.py +++ /dev/null @@ -1,86 +0,0 @@ -from __future__ import annotations - -import uuid -from collections.abc import Mapping -from typing import Any - -import pendulum - -from memu.database.inmemory.repositories.filter import matches_where -from memu.database.inmemory.state import InMemoryState -from memu.database.models import MemoryCategory -from memu.database.repositories.memory_category import MemoryCategoryRepo as MemoryCategoryRepoProtocol - - -class InMemoryMemoryCategoryRepository(MemoryCategoryRepoProtocol): - def __init__(self, *, state: InMemoryState, memory_category_model: type[MemoryCategory]) -> None: - self._state = state - self.memory_category_model = memory_category_model - self.categories: dict[str, MemoryCategory] = self._state.categories - - def list_categories(self, where: Mapping[str, Any] | None = None) -> dict[str, MemoryCategory]: - if not where: - return dict(self.categories) - return {cid: cat for cid, cat in self.categories.items() if matches_where(cat, where)} - - def clear_categories(self, where: Mapping[str, Any] | None = None) -> dict[str, MemoryCategory]: - if not where: - matches = self.categories.copy() - self.categories.clear() - return matches - matches = {cid: cat for cid, cat in self.categories.items() if matches_where(cat, where)} - for cid in matches: - self.categories.pop(cid, None) - return matches - - def get_or_create_category( - self, *, name: str, description: str, embedding: list[float], user_data: dict[str, Any] - ) -> MemoryCategory: - for c in self.categories.values(): - if c.name == name and all(getattr(c, k) == v for k, v in user_data.items()): - now = pendulum.now("UTC") - if c.embedding is None: - c.embedding = embedding - c.updated_at = now - if not c.description: - c.description = description - c.updated_at = now - return c - cid = str(uuid.uuid4()) - cat = self.memory_category_model(id=cid, name=name, description=description, embedding=embedding, **user_data) - self.categories[cid] = cat - return cat - - def update_category( - self, - *, - category_id: str, - name: str | None = None, - description: str | None = None, - embedding: list[float] | None = None, - summary: str | None = None, - ) -> MemoryCategory: - cat = self.categories.get(category_id) - if cat is None: - msg = f"Category with id {category_id} not found" - raise KeyError(msg) - - if name is not None: - cat.name = name - if description is not None: - cat.description = description - if embedding is not None: - cat.embedding = embedding - if summary is not None: - cat.summary = summary - - cat.updated_at = pendulum.now("UTC") - return cat - - def load_existing(self) -> None: - return None - - -MemoryCategoryRepo = InMemoryMemoryCategoryRepository - -__all__ = ["InMemoryMemoryCategoryRepository", "MemoryCategoryRepo"] diff --git a/src/memu/database/inmemory/repositories/memory_item_repo.py b/src/memu/database/inmemory/repositories/memory_item_repo.py deleted file mode 100644 index 8dfb1f3c..00000000 --- a/src/memu/database/inmemory/repositories/memory_item_repo.py +++ /dev/null @@ -1,263 +0,0 @@ -from __future__ import annotations - -import uuid -from collections.abc import Mapping -from typing import Any, override - -import pendulum - -from memu.database.inmemory.repositories.filter import matches_where -from memu.database.inmemory.state import InMemoryState -from memu.database.models import MemoryItem, MemoryType, compute_content_hash -from memu.database.repositories.memory_item import MemoryItemRepo -from memu.vector import cosine_topk, cosine_topk_salience - - -class InMemoryMemoryItemRepository(MemoryItemRepo): - def __init__(self, *, state: InMemoryState, memory_item_model: type[MemoryItem]) -> None: - self._state = state - self.memory_item_model = memory_item_model - self.items: dict[str, MemoryItem] = self._state.items - - def list_items(self, where: Mapping[str, Any] | None = None) -> dict[str, MemoryItem]: - if not where: - return dict(self.items) - return {mid: item for mid, item in self.items.items() if matches_where(item, where)} - - def list_items_by_ref_ids( - self, ref_ids: list[str], where: Mapping[str, Any] | None = None - ) -> dict[str, MemoryItem]: - """List items by their ref_id in the extra column. - - Args: - ref_ids: List of ref_ids to query. - where: Additional filter conditions. - - Returns: - Dict mapping item_id -> MemoryItem for items whose extra.ref_id is in ref_ids. - """ - if not ref_ids: - return {} - ref_id_set = set(ref_ids) - result: dict[str, MemoryItem] = {} - for mid, item in self.items.items(): - # Check where filter first - if where and not matches_where(item, where): - continue - # Check if ref_id is in the requested set - item_ref_id = (item.extra or {}).get("ref_id") - if item_ref_id and item_ref_id in ref_id_set: - result[mid] = item - return result - - def clear_items(self, where: Mapping[str, Any] | None = None) -> dict[str, MemoryItem]: - if not where: - matches = self.items.copy() - self.items.clear() - return matches - matches = {mid: item for mid, item in self.items.items() if matches_where(item, where)} - for mid in matches: - self.items.pop(mid, None) - return matches - - def _find_by_hash(self, content_hash: str, user_data: dict[str, Any]) -> MemoryItem | None: - """ - Find existing item by content hash within the same user scope. - - This enables deduplication: if the same content exists for the same user, - we reinforce it instead of creating a duplicate. - """ - for item in self.items.values(): - # Read content_hash from extra dict - item_hash = (item.extra or {}).get("content_hash") - if item_hash != content_hash: - continue - # Check scope match (user_id, agent_id, etc.) - if matches_where(item, user_data): - return item - return None - - def create_item( - self, - *, - resource_id: str, - memory_type: MemoryType, - summary: str, - embedding: list[float], - user_data: dict[str, Any], - reinforce: bool = False, - tool_record: dict[str, Any] | None = None, - ) -> MemoryItem: - if reinforce and memory_type != "tool": - return self.create_item_reinforce( - resource_id=resource_id, - memory_type=memory_type, - summary=summary, - embedding=embedding, - user_data=user_data, - ) - - # Build extra dict with tool_record fields at top level - extra: dict[str, Any] = {} - if tool_record: - if tool_record.get("when_to_use") is not None: - extra["when_to_use"] = tool_record["when_to_use"] - if tool_record.get("metadata") is not None: - extra["metadata"] = tool_record["metadata"] - if tool_record.get("tool_calls") is not None: - extra["tool_calls"] = tool_record["tool_calls"] - - mid = str(uuid.uuid4()) - it = self.memory_item_model( - id=mid, - resource_id=resource_id, - memory_type=memory_type, - summary=summary, - embedding=embedding, - extra=extra if extra else {}, - **user_data, - ) - self.items[mid] = it - return it - - def create_item_reinforce( - self, - *, - resource_id: str, - memory_type: MemoryType, - summary: str, - embedding: list[float], - user_data: dict[str, Any], - reinforce: bool = False, - ) -> MemoryItem: - content_hash = compute_content_hash(summary, memory_type) - - # Check for existing item with same hash in same scope (deduplication) - existing = self._find_by_hash(content_hash, user_data) - if existing: - # Reinforce existing memory instead of creating duplicate - current_extra = existing.extra or {} - current_count = current_extra.get("reinforcement_count", 1) - existing.extra = { - **current_extra, - "reinforcement_count": current_count + 1, - "last_reinforced_at": pendulum.now("UTC").isoformat(), - } - existing.updated_at = pendulum.now("UTC") - return existing - - # Create new item with salience tracking in extra - mid = str(uuid.uuid4()) - now = pendulum.now("UTC") - item_extra = user_data.pop("extra", {}) if "extra" in user_data else {} - item_extra.update({ - "content_hash": content_hash, - "reinforcement_count": 1, - "last_reinforced_at": now.isoformat(), - }) - it = self.memory_item_model( - id=mid, - resource_id=resource_id, - memory_type=memory_type, - summary=summary, - embedding=embedding, - extra=item_extra, - **user_data, - ) - self.items[mid] = it - return it - - def vector_search_items( - self, - query_vec: list[float], - top_k: int, - where: Mapping[str, Any] | None = None, - *, - ranking: str = "similarity", - recency_decay_days: float = 30.0, - ) -> list[tuple[str, float]]: - pool = self.list_items(where) - - if ranking == "salience": - # Salience-aware ranking: similarity x reinforcement x recency - # Read values from extra dict - corpus = [ - ( - i.id, - i.embedding, - (i.extra or {}).get("reinforcement_count", 1), - self._parse_datetime((i.extra or {}).get("last_reinforced_at")), - ) - for i in pool.values() - ] - return cosine_topk_salience(query_vec, corpus, k=top_k, recency_decay_days=recency_decay_days) - - # Default: pure cosine similarity (backward compatible) - hits = cosine_topk(query_vec, [(i.id, i.embedding) for i in pool.values()], k=top_k) - return hits - - def load_existing(self) -> None: - return None - - def get_item(self, item_id: str) -> MemoryItem | None: - return self.items.get(item_id) - - @staticmethod - def _parse_datetime(dt_str: str | None) -> pendulum.DateTime | None: - """Parse ISO datetime string from extra dict.""" - if dt_str is None: - return None - try: - parsed = pendulum.parse(dt_str) - except (ValueError, TypeError): - return None - else: - if isinstance(parsed, pendulum.DateTime): - return parsed - return None - - @override - def delete_item(self, item_id: str) -> None: - if item_id in self.items: - del self.items[item_id] - - @override - def update_item( - self, - *, - item_id: str, - memory_type: MemoryType | None = None, - summary: str | None = None, - embedding: list[float] | None = None, - extra: dict[str, Any] | None = None, - tool_record: dict[str, Any] | None = None, - ) -> MemoryItem: - item = self.items.get(item_id) - if item is None: - msg = f"Item with id {item_id} not found" - raise KeyError(msg) - - if memory_type is not None: - item.memory_type = memory_type - if summary is not None: - item.summary = summary - if embedding is not None: - item.embedding = embedding - - # Merge extra and tool_record into existing extra dict - current_extra = item.extra or {} - if extra is not None: - current_extra = {**current_extra, **extra} - if tool_record is not None: - # Merge tool_record fields at top level - for key in ("when_to_use", "metadata", "tool_calls"): - if tool_record.get(key) is not None: - current_extra[key] = tool_record[key] - if extra is not None or tool_record is not None: - item.extra = current_extra - - self.items[item_id] = item - return item - - -__all__ = ["InMemoryMemoryItemRepository"] diff --git a/src/memu/database/inmemory/repositories/resource_entry_repo.py b/src/memu/database/inmemory/repositories/resource_entry_repo.py new file mode 100644 index 00000000..7eec1010 --- /dev/null +++ b/src/memu/database/inmemory/repositories/resource_entry_repo.py @@ -0,0 +1,65 @@ +from __future__ import annotations + +import uuid +from collections.abc import Mapping +from typing import Any, override + +from memu.database.inmemory.repositories.filter import matches_where +from memu.database.inmemory.state import InMemoryState +from memu.database.models import ResourceEntry +from memu.database.repositories.resource_entry import ResourceEntryRepo + + +class InMemoryResourceEntryRepository(ResourceEntryRepo): + def __init__(self, *, state: InMemoryState, resource_entry_model: type[ResourceEntry]) -> None: + self._state = state + self.resource_entry_model = resource_entry_model + self.relations: list[ResourceEntry] = self._state.relations + + def list_relations(self, where: Mapping[str, Any] | None = None) -> list[ResourceEntry]: + if not where: + return list(self.relations) + return [rel for rel in self.relations if matches_where(rel, where)] + + def link_entry_resource(self, entry_id: str, resource_id: str, user_data: dict[str, Any]) -> ResourceEntry: + for rel in self.relations: + if rel.entry_id == entry_id and rel.resource_id == resource_id: + return rel + rel = self.resource_entry_model( + id=str(uuid.uuid4()), entry_id=entry_id, resource_id=resource_id, **user_data + ) + self.relations.append(rel) + return rel + + def load_existing(self) -> None: + return None + + @override + def get_entry_resources(self, entry_id: str) -> list[ResourceEntry]: + return [rel for rel in self.relations if rel.entry_id == entry_id] + + @override + def unlink_entry_resource(self, entry_id: str, resource_id: str) -> None: + # Mutate the shared state list in place so the DatabaseState reference and + # this repo's view never diverge. + self.relations[:] = [ + rel for rel in self.relations if not (rel.entry_id == entry_id and rel.resource_id == resource_id) + ] + + def unlink_entry(self, entry_id: str) -> list[ResourceEntry]: + removed = [rel for rel in self.relations if rel.entry_id == entry_id] + self.relations[:] = [rel for rel in self.relations if rel.entry_id != entry_id] + return removed + + def clear_relations(self, where: Mapping[str, Any] | None = None) -> list[ResourceEntry]: + if not where: + removed = list(self.relations) + self.relations.clear() + return removed + removed = [rel for rel in self.relations if matches_where(rel, where)] + removed_ids = {rel.id for rel in removed} + self.relations[:] = [rel for rel in self.relations if rel.id not in removed_ids] + return removed + + +__all__ = ["InMemoryResourceEntryRepository"] diff --git a/src/memu/database/inmemory/repositories/resource_repo.py b/src/memu/database/inmemory/repositories/resource_repo.py index 04c4e986..13ae1c9b 100644 --- a/src/memu/database/inmemory/repositories/resource_repo.py +++ b/src/memu/database/inmemory/repositories/resource_repo.py @@ -4,6 +4,8 @@ from collections.abc import Mapping from typing import Any +import pendulum + from memu.database.inmemory.repositories.filter import matches_where from memu.database.inmemory.state import InMemoryState from memu.database.models import Resource @@ -17,17 +19,27 @@ def __init__(self, *, state: InMemoryState, resource_model: type[Resource]) -> N self.resource_model = resource_model self.resources: dict[str, Resource] = self._state.resources - def list_resources(self, where: Mapping[str, Any] | None = None) -> dict[str, Resource]: - if not where: - return dict(self.resources) - return {rid: res for rid, res in self.resources.items() if matches_where(res, where)} + def get_resource(self, resource_id: str) -> Resource | None: + return self.resources.get(resource_id) + + def list_resources(self, where: Mapping[str, Any] | None = None, *, lane: str | None = None) -> dict[str, Resource]: + result = ( + self.resources + if not where + else {rid: res for rid, res in self.resources.items() if matches_where(res, where)} + ) + if lane is not None: + result = {rid: res for rid, res in result.items() if getattr(res, "lane", "source") == lane} + return dict(result) - def clear_resources(self, where: Mapping[str, Any] | None = None) -> dict[str, Resource]: - if not where: + def clear_resources( + self, where: Mapping[str, Any] | None = None, *, lane: str | None = None + ) -> dict[str, Resource]: + if not where and lane is None: matches = self.resources.copy() self.resources.clear() return matches - matches = {rid: res for rid, res in self.resources.items() if matches_where(res, where)} + matches = self.list_resources(where, lane=lane) for rid in matches: self.resources.pop(rid, None) return matches @@ -38,33 +50,108 @@ def delete_resource(self, resource_id: str) -> None: def create_resource( self, *, - url: str, modality: str, - local_path: str, - caption: str | None, - embedding: list[float] | None, user_data: dict[str, Any], + lane: str = "source", + url: str | None = None, + local_path: str | None = None, + slug: str | None = None, + title: str | None = None, + description: str | None = None, + content: str | None = None, + summary: str | None = None, + embedding: list[float] | None = None, + resource_refs: list[dict[str, Any]] | None = None, ) -> Resource: rid = str(uuid.uuid4()) res = self.resource_model( id=rid, - url=url, + lane=lane, modality=modality, + url=url, local_path=local_path, - caption=caption, + slug=slug, + title=title, + description=description, + content=content, + summary=summary, embedding=embedding, + resource_refs=resource_refs or [], **user_data, ) self.resources[rid] = res return res + def get_or_create_doc( + self, + *, + lane: str, + title: str, + description: str, + embedding: list[float], + user_data: dict[str, Any], + slug: str | None = None, + ) -> Resource: + for res in self.resources.values(): + same_lane = getattr(res, "lane", "source") == lane + if same_lane and res.title == title and all(getattr(res, k) == v for k, v in user_data.items()): + now = pendulum.now("UTC") + if res.embedding is None: + res.embedding = embedding + res.updated_at = now + if not res.description: + res.description = description + res.updated_at = now + return res + return self.create_resource( + modality="markdown", + lane=lane, + title=title, + slug=slug, + description=description, + embedding=embedding, + user_data=user_data, + ) + + def update_resource( + self, + *, + resource_id: str, + title: str | None = None, + description: str | None = None, + content: str | None = None, + summary: str | None = None, + embedding: list[float] | None = None, + resource_refs: list[dict[str, Any]] | None = None, + ) -> Resource: + res = self.resources.get(resource_id) + if res is None: + msg = f"Resource with id {resource_id} not found" + raise KeyError(msg) + if title is not None: + res.title = title + if description is not None: + res.description = description + if content is not None: + res.content = content + if summary is not None: + res.summary = summary + if embedding is not None: + res.embedding = embedding + if resource_refs is not None: + res.resource_refs = resource_refs + res.updated_at = pendulum.now("UTC") + return res + def vector_search_resources( self, query_vec: list[float], top_k: int, where: Mapping[str, Any] | None = None, + *, + lane: str | None = None, ) -> list[tuple[str, float]]: - pool = self.list_resources(where) + pool = self.list_resources(where, lane=lane) corpus = [(rid, res.embedding) for rid, res in pool.items() if res.embedding] return cosine_topk(query_vec, corpus, k=top_k) diff --git a/src/memu/database/interfaces.py b/src/memu/database/interfaces.py index 4acd0312..974408b5 100644 --- a/src/memu/database/interfaces.py +++ b/src/memu/database/interfaces.py @@ -2,11 +2,10 @@ from typing import Protocol, runtime_checkable -from memu.database.models import CategoryItem as CategoryItemRecord -from memu.database.models import MemoryCategory as MemoryCategoryRecord -from memu.database.models import MemoryItem as MemoryItemRecord +from memu.database.models import Entry as EntryRecord from memu.database.models import Resource as ResourceRecord -from memu.database.repositories import CategoryItemRepo, MemoryCategoryRepo, MemoryItemRepo, ResourceRepo +from memu.database.models import ResourceEntry as ResourceEntryRecord +from memu.database.repositories import EntryRepo, ResourceEntryRepo, ResourceRepo @runtime_checkable @@ -14,22 +13,19 @@ class Database(Protocol): """Backend-agnostic database contract.""" resource_repo: ResourceRepo - memory_category_repo: MemoryCategoryRepo - memory_item_repo: MemoryItemRepo - category_item_repo: CategoryItemRepo + entry_repo: EntryRepo + resource_entry_repo: ResourceEntryRepo resources: dict[str, ResourceRecord] - items: dict[str, MemoryItemRecord] - categories: dict[str, MemoryCategoryRecord] - relations: list[CategoryItemRecord] + entries: dict[str, EntryRecord] + relations: list[ResourceEntryRecord] def close(self) -> None: ... __all__ = [ - "CategoryItemRecord", "Database", - "MemoryCategoryRecord", - "MemoryItemRecord", + "EntryRecord", + "ResourceEntryRecord", "ResourceRecord", ] diff --git a/src/memu/database/models.py b/src/memu/database/models.py index 0124b784..c3fc788e 100644 --- a/src/memu/database/models.py +++ b/src/memu/database/models.py @@ -4,31 +4,36 @@ import json import uuid from datetime import datetime +from os.path import basename from typing import Any, Literal import pendulum from pydantic import BaseModel, ConfigDict, Field +# Sub-type of a memory-lane entry (kept for prompt routing / backward semantics). MemoryType = Literal["profile", "event", "knowledge", "behavior", "skill", "tool"] +# A lane is one of the parallel, structurally identical processing tracks that +# share the same Resource -> canonical-text trunk: +# - "source": raw input artifacts (conversation/document/image/video/audio) +# - "index": per-resource catalog/description docs +# - "memory": grouped memory docs (the former "category") +# - "skill": grouped reusable-skill docs +# index/memory/skill are the three retrievable lanes; "source" holds raw inputs. +Lane = Literal["source", "index", "memory", "skill"] +SOURCE_LANE: Lane = "source" +RETRIEVAL_LANES: tuple[Lane, ...] = ("index", "memory", "skill") +MARKDOWN_MODALITY = "markdown" -def compute_content_hash(summary: str, memory_type: str) -> str: - """ - Generate unique hash for memory deduplication. - - Operates on post-summary content. Normalizes whitespace to handle - minor formatting differences like "I love coffee" vs "I love coffee". - Args: - summary: The memory summary text - memory_type: The type of memory (profile, event, etc.) +def compute_content_hash(text: str, entry_kind: str) -> str: + """Generate a stable hash for entry deduplication. - Returns: - A 16-character hex hash string + Operates on post-extraction content. Normalizes whitespace to absorb minor + formatting differences ("I love coffee" vs "I love coffee"). """ - # Normalize: lowercase, strip, collapse whitespace - normalized = " ".join(summary.lower().split()) - content = f"{memory_type}:{normalized}" + normalized = " ".join(text.lower().split()) + content = f"{entry_kind}:{normalized}" return hashlib.sha256(content.encode()).hexdigest()[:16] @@ -66,43 +71,69 @@ def ensure_hash(self) -> None: class Resource(BaseRecord): - url: str + """A node in the unified store: either a raw input or a generated lane doc. + + "Everything is a resource": + - raw inputs: ``lane="source"``, ``modality`` = conversation/document/ + image/video/audio; ``content`` holds the canonical, + modality-agnostic text from preprocessing (the trunk). + - generated docs: ``lane`` in {index, memory, skill}, ``modality="markdown"``, + rendered as ``resource//.md``; ``content`` holds + the markdown body, ``summary`` the searchable condensation. + """ + + lane: str = SOURCE_LANE modality: str - local_path: str - caption: str | None = None + url: str | None = None + local_path: str | None = None + # Filename stem used for ``resource//.md`` (generated docs only). + slug: str | None = None + # Human title of a generated doc (the former category/skill name). + title: str | None = None + # Short blurb for a generated doc (the former category description). + description: str | None = None + # Raw: canonical text from preprocessing. Doc: rendered markdown body. + content: str | None = None + # Searchable condensation used for coarse (resource-level) recall. + summary: str | None = None embedding: list[float] | None = None - - -class MemoryItem(BaseRecord): - resource_id: str | None - memory_type: str - summary: str + # Provenance: raw sources a generated doc derives from. Each dict holds at + # least ``resource_id`` and ``source_path`` (plus optional ``modality``). + resource_refs: list[dict[str, Any]] = [] + + @property + def source_path(self) -> str: + """Path relative to the ``resource/`` root for this artifact.""" + if self.lane != SOURCE_LANE and self.slug: + return f"resource/{self.lane}/{self.slug}.md" + name = basename(self.local_path or self.url or self.id) + return f"resource/{name}" + + +class Entry(BaseRecord): + """The searchable atom of a lane (index description / memory item / skill step).""" + + lane: str + # Originating raw source resource (provenance), and its relative path. + source_id: str | None = None + source_path: str | None = None + # Sub-type within a lane (memory: profile/event/...; skill: step kind; etc.). + entry_kind: str + text: str embedding: list[float] | None = None happened_at: datetime | None = None extra: dict[str, Any] = {} # extra may contain: - # # reinforcement tracking fields - # - content_hash: str - # - reinforcement_count: int - # - last_reinforced_at: str (isoformat) - # # Reference tracking field - # - ref_id: str - # # Tool memory fields - # - when_to_use: str - Hint for when this memory should be retrieved - # - metadata: dict - Type-specific metadata (e.g., tool_name, avg_success_rate) - # - tool_calls: list[dict] - Tool call history for tool memories (serialized ToolCallResult) - - -class MemoryCategory(BaseRecord): - name: str - description: str - embedding: list[float] | None = None - summary: str | None = None + # - content_hash / reinforcement_count / last_reinforced_at (salience) + # - ref_id (reference tracking) + # - when_to_use / metadata / tool_calls (tool memory) -class CategoryItem(BaseRecord): - item_id: str - category_id: str +class ResourceEntry(BaseRecord): + """Edge: membership of an Entry in its coarse (lane) Resource doc.""" + + entry_id: str + resource_id: str def merge_scope_model[TBaseRecord: BaseRecord]( @@ -123,24 +154,24 @@ def merge_scope_model[TBaseRecord: BaseRecord]( def build_scoped_models( user_model: type[BaseModel], -) -> tuple[type[Resource], type[MemoryCategory], type[MemoryItem], type[CategoryItem]]: - """ - Build scoped interface models (Pydantic) that inherit from the base record models and user scope. - """ +) -> tuple[type[Resource], type[Entry], type[ResourceEntry]]: + """Build scoped interface models that inherit base records and the user scope.""" resource_model = merge_scope_model(user_model, Resource, name_suffix="Resource") - memory_category_model = merge_scope_model(user_model, MemoryCategory, name_suffix="MemoryCategory") - memory_item_model = merge_scope_model(user_model, MemoryItem, name_suffix="MemoryItem") - category_item_model = merge_scope_model(user_model, CategoryItem, name_suffix="CategoryItem") - return resource_model, memory_category_model, memory_item_model, category_item_model + entry_model = merge_scope_model(user_model, Entry, name_suffix="Entry") + resource_entry_model = merge_scope_model(user_model, ResourceEntry, name_suffix="ResourceEntry") + return resource_model, entry_model, resource_entry_model __all__ = [ + "MARKDOWN_MODALITY", + "RETRIEVAL_LANES", + "SOURCE_LANE", "BaseRecord", - "CategoryItem", - "MemoryCategory", - "MemoryItem", + "Entry", + "Lane", "MemoryType", "Resource", + "ResourceEntry", "ToolCallResult", "build_scoped_models", "compute_content_hash", diff --git a/src/memu/database/postgres/__init__.py b/src/memu/database/postgres/__init__.py index 2b551b91..b5b31739 100644 --- a/src/memu/database/postgres/__init__.py +++ b/src/memu/database/postgres/__init__.py @@ -26,9 +26,8 @@ def build_postgres_database( vector_provider=vector_provider, scope_model=user_model, resource_model=sqla_models.Resource, - memory_category_model=sqla_models.MemoryCategory, - memory_item_model=sqla_models.MemoryItem, - category_item_model=sqla_models.CategoryItem, + entry_model=sqla_models.Entry, + resource_entry_model=sqla_models.ResourceEntry, sqla_models=sqla_models, ) diff --git a/src/memu/database/postgres/models.py b/src/memu/database/postgres/models.py index e83797a2..eb62f48a 100644 --- a/src/memu/database/postgres/models.py +++ b/src/memu/database/postgres/models.py @@ -13,11 +13,11 @@ raise ImportError(msg) from exc from pydantic import BaseModel -from sqlalchemy import ForeignKey, MetaData, String, Text +from sqlalchemy import MetaData, String, Text from sqlalchemy.dialects.postgresql import JSONB from sqlmodel import Column, DateTime, Field, Index, SQLModel, func -from memu.database.models import CategoryItem, MemoryCategory, MemoryItem, MemoryType, Resource +from memu.database.models import Entry, Resource, ResourceEntry class TZDateTime(DateTime): @@ -43,35 +43,42 @@ class BaseModelMixin(SQLModel): ) -class ResourceModel(BaseModelMixin, Resource): - url: str = Field(sa_column=Column(String, nullable=False)) +class PostgresResourceModel(BaseModelMixin, Resource): + """A node in the unified store: raw input (``lane="source"``) or generated doc.""" + + lane: str = Field(sa_column=Column(String, nullable=False, index=True)) modality: str = Field(sa_column=Column(String, nullable=False)) - local_path: str = Field(sa_column=Column(String, nullable=False)) - caption: str | None = Field(default=None, sa_column=Column(Text, nullable=True)) + url: str | None = Field(default=None, sa_column=Column(String, nullable=True)) + local_path: str | None = Field(default=None, sa_column=Column(String, nullable=True)) + slug: str | None = Field(default=None, sa_column=Column(String, nullable=True)) + title: str | None = Field(default=None, sa_column=Column(String, nullable=True)) + description: str | None = Field(default=None, sa_column=Column(Text, nullable=True)) + content: str | None = Field(default=None, sa_column=Column(Text, nullable=True)) + summary: str | None = Field(default=None, sa_column=Column(Text, nullable=True)) embedding: list[float] | None = Field(default=None, sa_column=Column(Vector(), nullable=True)) + resource_refs: list[dict[str, Any]] = Field(default_factory=list, sa_column=Column(JSONB, nullable=True)) -class MemoryItemModel(BaseModelMixin, MemoryItem): - resource_id: str | None = Field(sa_column=Column(ForeignKey("resources.id", ondelete="CASCADE"), nullable=True)) - memory_type: MemoryType = Field(sa_column=Column(String, nullable=False)) - summary: str = Field(sa_column=Column(Text, nullable=False)) - embedding: list[float] | None = Field(default=None, sa_column=Column(Vector(), nullable=True)) - happened_at: datetime | None = Field(default=None, sa_column=Column(DateTime, nullable=True)) - extra: dict[str, Any] = Field(default={}, sa_column=Column(JSONB, nullable=True)) - +class PostgresEntryModel(BaseModelMixin, Entry): + """The searchable atom of a lane (index description / memory item / skill step).""" -class MemoryCategoryModel(BaseModelMixin, MemoryCategory): - name: str = Field(sa_column=Column(String, nullable=False, index=True)) - description: str = Field(sa_column=Column(Text, nullable=False)) + lane: str = Field(sa_column=Column(String, nullable=False, index=True)) + source_id: str | None = Field(default=None, sa_column=Column(String, nullable=True, index=True)) + source_path: str | None = Field(default=None, sa_column=Column(String, nullable=True)) + entry_kind: str = Field(sa_column=Column(String, nullable=False)) + text: str = Field(sa_column=Column(Text, nullable=False)) embedding: list[float] | None = Field(default=None, sa_column=Column(Vector(), nullable=True)) - summary: str | None = Field(default=None, sa_column=Column(Text, nullable=True)) + happened_at: datetime | None = Field(default=None, sa_column=Column(TZDateTime, nullable=True)) + extra: dict[str, Any] = Field(default_factory=dict, sa_column=Column(JSONB, nullable=True)) + +class PostgresResourceEntryModel(BaseModelMixin, ResourceEntry): + """Edge: membership of an Entry in its coarse (lane) Resource doc.""" -class CategoryItemModel(BaseModelMixin, CategoryItem): - item_id: str = Field(sa_column=Column(ForeignKey("memory_items.id", ondelete="CASCADE"), nullable=False)) - category_id: str = Field(sa_column=Column(ForeignKey("memory_categories.id", ondelete="CASCADE"), nullable=False)) + entry_id: str = Field(sa_column=Column(String, nullable=False, index=True)) + resource_id: str = Field(sa_column=Column(String, nullable=False, index=True)) - __table_args__ = (Index("idx_category_items_unique", "item_id", "category_id", unique=True),) + __table_args__ = (Index("idx_resource_entries_unique", "entry_id", "resource_id", unique=True),) def _normalize_table_args(table_args: Any) -> tuple[list[Any], dict[str, Any]]: @@ -156,25 +163,19 @@ def build_table_model( def build_scoped_models( user_model: type[BaseModel], -) -> tuple[type[SQLModel], type[SQLModel], type[SQLModel], type[SQLModel]]: - """ - Build scoped SQLModel tables for each entity (resource, category, item, relation). - """ - resource_model = build_table_model(user_model, ResourceModel, tablename="resources") - memory_category_model = build_table_model( - user_model, MemoryCategoryModel, tablename="memory_categories", unique_with_scope=["name"] - ) - memory_item_model = build_table_model(user_model, MemoryItemModel, tablename="memory_items") - category_item_model = build_table_model(user_model, CategoryItemModel, tablename="category_items") - return resource_model, memory_category_model, memory_item_model, category_item_model +) -> tuple[type[SQLModel], type[SQLModel], type[SQLModel]]: + """Build scoped SQLModel tables for each entity (resource, entry, resource-entry).""" + resource_model = build_table_model(user_model, PostgresResourceModel, tablename="memu_resources") + entry_model = build_table_model(user_model, PostgresEntryModel, tablename="memu_entries") + resource_entry_model = build_table_model(user_model, PostgresResourceEntryModel, tablename="memu_resource_entries") + return resource_model, entry_model, resource_entry_model __all__ = [ "BaseModelMixin", - "CategoryItemModel", - "MemoryCategoryModel", - "MemoryItemModel", - "ResourceModel", + "PostgresEntryModel", + "PostgresResourceEntryModel", + "PostgresResourceModel", "build_scoped_models", "build_table_model", ] diff --git a/src/memu/database/postgres/postgres.py b/src/memu/database/postgres/postgres.py index d1ff7b05..a858d410 100644 --- a/src/memu/database/postgres/postgres.py +++ b/src/memu/database/postgres/postgres.py @@ -6,15 +6,14 @@ from pydantic import BaseModel from memu.database.interfaces import Database -from memu.database.models import CategoryItem, MemoryCategory, MemoryItem, Resource +from memu.database.models import Entry, Resource, ResourceEntry from memu.database.postgres.migration import DDLMode, run_migrations -from memu.database.postgres.repositories.category_item_repo import PostgresCategoryItemRepo -from memu.database.postgres.repositories.memory_category_repo import PostgresMemoryCategoryRepo -from memu.database.postgres.repositories.memory_item_repo import PostgresMemoryItemRepo +from memu.database.postgres.repositories.entry_repo import PostgresEntryRepo +from memu.database.postgres.repositories.resource_entry_repo import PostgresResourceEntryRepo from memu.database.postgres.repositories.resource_repo import PostgresResourceRepo from memu.database.postgres.schema import SQLAModels, get_sqlalchemy_models, require_sqlalchemy from memu.database.postgres.session import SessionManager -from memu.database.repositories import CategoryItemRepo, MemoryCategoryRepo, MemoryItemRepo, ResourceRepo +from memu.database.repositories import EntryRepo, ResourceEntryRepo, ResourceRepo from memu.database.state import DatabaseState logger = logging.getLogger(__name__) @@ -22,13 +21,11 @@ class PostgresStore(Database): resource_repo: ResourceRepo - memory_category_repo: MemoryCategoryRepo - memory_item_repo: MemoryItemRepo - category_item_repo: CategoryItemRepo + entry_repo: EntryRepo + resource_entry_repo: ResourceEntryRepo resources: dict[str, Resource] - items: dict[str, MemoryItem] - categories: dict[str, MemoryCategory] - relations: list[CategoryItem] + entries: dict[str, Entry] + relations: list[ResourceEntry] def __init__( self, @@ -39,9 +36,8 @@ def __init__( scope_model: type[BaseModel] | None = None, base_model: type[BaseModel] | None = None, resource_model: type[Any] | None = None, - memory_category_model: type[Any] | None = None, - memory_item_model: type[Any] | None = None, - category_item_model: type[Any] | None = None, + entry_model: type[Any] | None = None, + resource_entry_model: type[Any] | None = None, sqla_models: SQLAModels | None = None, ) -> None: require_sqlalchemy() @@ -57,9 +53,8 @@ def __init__( run_migrations(dsn=self.dsn, scope_model=self._scope_model, ddl_mode=self.ddl_mode) resource_model = resource_model or self._sqla_models.Resource - memory_category_model = memory_category_model or self._sqla_models.MemoryCategory - memory_item_model = memory_item_model or self._sqla_models.MemoryItem - category_item_model = category_item_model or self._sqla_models.CategoryItem + entry_model = entry_model or self._sqla_models.Entry + resource_entry_model = resource_entry_model or self._sqla_models.ResourceEntry self.resource_repo = PostgresResourceRepo( state=self._state, @@ -67,42 +62,27 @@ def __init__( sqla_models=self._sqla_models, sessions=self._sessions, scope_fields=self._scope_fields, + use_vector=self._use_vector_type, ) - self.memory_category_repo = PostgresMemoryCategoryRepo( - state=self._state, - memory_category_model=memory_category_model, - sqla_models=self._sqla_models, - sessions=self._sessions, - scope_fields=self._scope_fields, - ) - self.memory_item_repo = PostgresMemoryItemRepo( + self.entry_repo = PostgresEntryRepo( state=self._state, - memory_item_model=memory_item_model, + entry_model=entry_model, sqla_models=self._sqla_models, sessions=self._sessions, scope_fields=self._scope_fields, use_vector=self._use_vector_type, ) - self.category_item_repo = PostgresCategoryItemRepo( + self.resource_entry_repo = PostgresResourceEntryRepo( state=self._state, - category_item_model=category_item_model, + resource_entry_model=resource_entry_model, sqla_models=self._sqla_models, sessions=self._sessions, scope_fields=self._scope_fields, ) self.resources = self._state.resources - self.items = self._state.items - self.categories = self._state.categories + self.entries = self._state.entries self.relations = self._state.relations - # self._load_existing() - def close(self) -> None: self._sessions.close() - - def _load_existing(self) -> None: - self.resource_repo.load_existing() - self.memory_category_repo.load_existing() - self.memory_item_repo.load_existing() - self.category_item_repo.load_existing() diff --git a/src/memu/database/postgres/repositories/__init__.py b/src/memu/database/postgres/repositories/__init__.py index 648623e5..3c8f332a 100644 --- a/src/memu/database/postgres/repositories/__init__.py +++ b/src/memu/database/postgres/repositories/__init__.py @@ -1,11 +1,9 @@ -from memu.database.postgres.repositories.category_item_repo import PostgresCategoryItemRepo -from memu.database.postgres.repositories.memory_category_repo import PostgresMemoryCategoryRepo -from memu.database.postgres.repositories.memory_item_repo import PostgresMemoryItemRepo +from memu.database.postgres.repositories.entry_repo import PostgresEntryRepo +from memu.database.postgres.repositories.resource_entry_repo import PostgresResourceEntryRepo from memu.database.postgres.repositories.resource_repo import PostgresResourceRepo __all__ = [ - "PostgresCategoryItemRepo", - "PostgresMemoryCategoryRepo", - "PostgresMemoryItemRepo", + "PostgresEntryRepo", + "PostgresResourceEntryRepo", "PostgresResourceRepo", ] diff --git a/src/memu/database/postgres/repositories/base.py b/src/memu/database/postgres/repositories/base.py index 0823dbf8..be78ca42 100644 --- a/src/memu/database/postgres/repositories/base.py +++ b/src/memu/database/postgres/repositories/base.py @@ -56,11 +56,6 @@ def _prepare_embedding(self, embedding: list[float] | None) -> Any: return None return embedding - def _merge_and_commit(self, obj: Any) -> None: - with self._sessions.session() as session: - session.merge(obj) - session.commit() - def _now(self) -> pendulum.DateTime: return pendulum.now("UTC") @@ -85,29 +80,5 @@ def _build_filters(self, model: Any, where: Mapping[str, Any] | None) -> list[An filters.append(column == expected) return filters - @staticmethod - def _matches_where(obj: Any, where: Mapping[str, Any] | None) -> bool: - if not where: - return True - for raw_key, expected in where.items(): - if expected is None: - continue - field, op = [*raw_key.split("__", 1), None][:2] - actual = getattr(obj, str(field), None) - if op == "in": - if isinstance(expected, str): - if actual != expected: - return False - else: - try: - if actual not in expected: - return False - except TypeError: - return False - else: - if actual != expected: - return False - return True - __all__ = ["PostgresRepoBase"] diff --git a/src/memu/database/postgres/repositories/category_item_repo.py b/src/memu/database/postgres/repositories/category_item_repo.py deleted file mode 100644 index 66d02f7f..00000000 --- a/src/memu/database/postgres/repositories/category_item_repo.py +++ /dev/null @@ -1,143 +0,0 @@ -from __future__ import annotations - -from collections.abc import Mapping -from typing import Any - -from memu.database.models import CategoryItem -from memu.database.postgres.repositories.base import PostgresRepoBase -from memu.database.postgres.session import SessionManager -from memu.database.repositories.category_item import CategoryItemRepo -from memu.database.state import DatabaseState - - -class PostgresCategoryItemRepo(PostgresRepoBase, CategoryItemRepo): - def __init__( - self, - *, - state: DatabaseState, - category_item_model: type[CategoryItem], - sqla_models: Any, - sessions: SessionManager, - scope_fields: list[str], - ) -> None: - super().__init__(state=state, sqla_models=sqla_models, sessions=sessions, scope_fields=scope_fields) - self._category_item_model = category_item_model - self.relations: list[CategoryItem] = self._state.relations - - def list_relations(self, where: Mapping[str, Any] | None = None) -> list[CategoryItem]: - from sqlmodel import select - - filters = self._build_filters(self._sqla_models.CategoryItem, where) - with self._sessions.session() as session: - rows = session.scalars(select(self._sqla_models.CategoryItem).where(*filters)).all() - return [self._cache_relation(row) for row in rows] - - def link_item_category(self, item_id: str, cat_id: str, user_data: dict[str, Any]) -> CategoryItem: - from sqlmodel import select - - # Avoid duplicate inserts using local cache - for rel in self.relations: - if rel.item_id == item_id and rel.category_id == cat_id: - return rel - - now = self._now() - new_rel = self._category_item_model( - item_id=item_id, - category_id=cat_id, - **user_data, - created_at=now, - updated_at=now, - ) - - with self._sessions.session() as session: - existing = session.scalar( - select(self._sqla_models.CategoryItem).where( - self._sqla_models.CategoryItem.item_id == item_id, - self._sqla_models.CategoryItem.category_id == cat_id, - ) - ) - if existing: - return self._cache_relation(existing) - - session.add(new_rel) - session.commit() - session.refresh(new_rel) - - return self._cache_relation(new_rel) - - def unlink_item_category(self, item_id: str, cat_id: str) -> None: - from sqlmodel import delete - - with self._sessions.session() as session: - session.exec( - delete(self._sqla_models.CategoryItem).where( - self._sqla_models.CategoryItem.item_id == item_id, - self._sqla_models.CategoryItem.category_id == cat_id, - ) - ) - session.commit() - self.relations[:] = [r for r in self.relations if not (r.item_id == item_id and r.category_id == cat_id)] - - def _row_to_record(self, row: Any) -> CategoryItem: - return CategoryItem( - id=row.id, - item_id=row.item_id, - category_id=row.category_id, - created_at=row.created_at, - updated_at=row.updated_at, - **self._scope_kwargs_from(row), - ) - - def unlink_item(self, item_id: str) -> list[CategoryItem]: - from sqlmodel import delete, select - - with self._sessions.session() as session: - rows = session.scalars( - select(self._sqla_models.CategoryItem).where(self._sqla_models.CategoryItem.item_id == item_id) - ).all() - removed = [self._row_to_record(row) for row in rows] - if removed: - session.exec( - delete(self._sqla_models.CategoryItem).where(self._sqla_models.CategoryItem.item_id == item_id) - ) - session.commit() - self.relations[:] = [r for r in self.relations if r.item_id != item_id] - return removed - - def clear_relations(self, where: Mapping[str, Any] | None = None) -> list[CategoryItem]: - from sqlmodel import delete, select - - filters = self._build_filters(self._sqla_models.CategoryItem, where) - with self._sessions.session() as session: - rows = session.scalars(select(self._sqla_models.CategoryItem).where(*filters)).all() - removed = [self._row_to_record(row) for row in rows] - if removed: - session.exec(delete(self._sqla_models.CategoryItem).where(*filters)) - session.commit() - removed_ids = {rel.id for rel in removed} - self.relations[:] = [r for r in self.relations if r.id not in removed_ids] - return removed - - def get_item_categories(self, item_id: str) -> list[CategoryItem]: - from sqlmodel import select - - with self._sessions.session() as session: - rows = session.scalars( - select(self._sqla_models.CategoryItem).where(self._sqla_models.CategoryItem.item_id == item_id) - ).all() - return [self._cache_relation(row) for row in rows] - - def load_existing(self) -> None: - from sqlmodel import select - - with self._sessions.session() as session: - rows = session.scalars(select(self._sqla_models.CategoryItem)).all() - for row in rows: - self._cache_relation(row) - - def _cache_relation(self, rel: CategoryItem) -> CategoryItem: - self.relations.append(rel) - return rel - - -__all__ = ["PostgresCategoryItemRepo"] diff --git a/src/memu/database/postgres/repositories/entry_repo.py b/src/memu/database/postgres/repositories/entry_repo.py new file mode 100644 index 00000000..aba3b362 --- /dev/null +++ b/src/memu/database/postgres/repositories/entry_repo.py @@ -0,0 +1,367 @@ +from __future__ import annotations + +from collections.abc import Mapping +from datetime import datetime +from typing import Any + +import pendulum + +from memu.database.models import Entry, compute_content_hash +from memu.database.postgres.repositories.base import PostgresRepoBase +from memu.database.postgres.session import SessionManager +from memu.database.repositories.entry import EntryRepo +from memu.database.state import DatabaseState +from memu.vector import cosine_topk, cosine_topk_salience + + +class PostgresEntryRepo(PostgresRepoBase, EntryRepo): + def __init__( + self, + *, + state: DatabaseState, + entry_model: type[Entry], + sqla_models: Any, + sessions: SessionManager, + scope_fields: list[str], + use_vector: bool, + ) -> None: + super().__init__( + state=state, sqla_models=sqla_models, sessions=sessions, scope_fields=scope_fields, use_vector=use_vector + ) + self._entry_model = entry_model + self.entries: dict[str, Entry] = self._state.entries + + def get_entry(self, entry_id: str) -> Entry | None: + from sqlmodel import select + + with self._sessions.session() as session: + row = session.scalar(select(self._sqla_models.Entry).where(self._sqla_models.Entry.id == entry_id)) + if row: + row.embedding = self._normalize_embedding(row.embedding) + return self._cache_entry(row) + return None + + def list_entries( + self, where: Mapping[str, Any] | None = None, *, lane: str | None = None + ) -> dict[str, Entry]: + from sqlmodel import select + + model = self._sqla_models.Entry + filters = self._build_filters(model, where) + if lane is not None: + filters.append(model.lane == lane) + with self._sessions.session() as session: + rows = session.scalars(select(model).where(*filters)).all() + result: dict[str, Entry] = {} + for row in rows: + row.embedding = self._normalize_embedding(row.embedding) + entry = self._cache_entry(row) + result[entry.id] = entry + return result + + def list_entries_by_ref_ids( + self, ref_ids: list[str], where: Mapping[str, Any] | None = None + ) -> dict[str, Entry]: + if not ref_ids: + return {} + + from sqlmodel import select + + model = self._sqla_models.Entry + filters = self._build_filters(model, where) + ref_id_col = model.extra["ref_id"].astext + filters.append(ref_id_col.isnot(None)) + filters.append(ref_id_col.in_(ref_ids)) + + with self._sessions.session() as session: + rows = session.scalars(select(model).where(*filters)).all() + result: dict[str, Entry] = {} + for row in rows: + row.embedding = self._normalize_embedding(row.embedding) + entry = self._cache_entry(row) + result[entry.id] = entry + return result + + def clear_entries( + self, where: Mapping[str, Any] | None = None, *, lane: str | None = None + ) -> dict[str, Entry]: + from sqlmodel import delete, select + + model = self._sqla_models.Entry + filters = self._build_filters(model, where) + if lane is not None: + filters.append(model.lane == lane) + with self._sessions.session() as session: + rows = session.scalars(select(model).where(*filters)).all() + deleted: dict[str, Entry] = {} + for row in rows: + row.embedding = self._normalize_embedding(row.embedding) + deleted[row.id] = row + + if not deleted: + return {} + + session.exec(delete(model).where(*filters)) + session.commit() + + for entry_id in deleted: + self.entries.pop(entry_id, None) + + return deleted + + def create_entry( + self, + *, + lane: str, + source_id: str | None, + entry_kind: str, + text: str, + embedding: list[float], + user_data: dict[str, Any], + source_path: str | None = None, + reinforce: bool = False, + tool_record: dict[str, Any] | None = None, + ) -> Entry: + if reinforce and entry_kind != "tool": + return self.create_entry_reinforce( + lane=lane, + source_id=source_id, + entry_kind=entry_kind, + text=text, + embedding=embedding, + user_data=user_data, + source_path=source_path, + ) + + extra: dict[str, Any] = {} + if tool_record: + for key in ("when_to_use", "metadata", "tool_calls"): + if tool_record.get(key) is not None: + extra[key] = tool_record[key] + + entry = self._entry_model( + lane=lane, + source_id=source_id, + source_path=source_path, + entry_kind=entry_kind, + text=text, + embedding=self._prepare_embedding(embedding), + extra=extra if extra else {}, + **user_data, + created_at=self._now(), + updated_at=self._now(), + ) + + with self._sessions.session() as session: + session.add(entry) + session.commit() + session.refresh(entry) + + entry.embedding = self._normalize_embedding(entry.embedding) + return self._cache_entry(entry) + + def create_entry_reinforce( + self, + *, + lane: str, + source_id: str | None, + entry_kind: str, + text: str, + embedding: list[float], + user_data: dict[str, Any], + source_path: str | None = None, + ) -> Entry: + from sqlmodel import select + + model = self._sqla_models.Entry + content_hash = compute_content_hash(text, entry_kind) + entry_extra = user_data.pop("extra", {}) if "extra" in user_data else {} + + with self._sessions.session() as session: + content_hash_col = model.extra["content_hash"].astext + filters = [content_hash_col == content_hash] + filters.extend(self._build_filters(model, user_data)) + + existing = session.scalar(select(model).where(*filters)) + + if existing: + current_extra = existing.extra or {} + current_count = current_extra.get("reinforcement_count", 1) + existing.extra = { + **current_extra, + "reinforcement_count": current_count + 1, + "last_reinforced_at": self._now().isoformat(), + } + existing.updated_at = self._now() + session.add(existing) + session.commit() + session.refresh(existing) + existing.embedding = self._normalize_embedding(existing.embedding) + return self._cache_entry(existing) + + now = self._now() + entry_extra.update({ + "content_hash": content_hash, + "reinforcement_count": 1, + "last_reinforced_at": now.isoformat(), + }) + entry = self._entry_model( + lane=lane, + source_id=source_id, + source_path=source_path, + entry_kind=entry_kind, + text=text, + embedding=self._prepare_embedding(embedding), + **user_data, + created_at=now, + updated_at=now, + extra=entry_extra, + ) + + session.add(entry) + session.commit() + session.refresh(entry) + + entry.embedding = self._normalize_embedding(entry.embedding) + return self._cache_entry(entry) + + def update_entry( + self, + *, + entry_id: str, + entry_kind: str | None = None, + text: str | None = None, + embedding: list[float] | None = None, + extra: dict[str, Any] | None = None, + tool_record: dict[str, Any] | None = None, + ) -> Entry: + from sqlmodel import select + + model = self._sqla_models.Entry + now = self._now() + with self._sessions.session() as session: + entry = session.scalar(select(model).where(model.id == entry_id)) + if entry is None: + msg = f"Entry with id {entry_id} not found" + raise KeyError(msg) + + if entry_kind is not None: + entry.entry_kind = entry_kind + if text is not None: + entry.text = text + if embedding is not None: + entry.embedding = self._prepare_embedding(embedding) + + current_extra = entry.extra or {} + if extra is not None: + current_extra = {**current_extra, **extra} + if tool_record is not None: + for key in ("when_to_use", "metadata", "tool_calls"): + if tool_record.get(key) is not None: + current_extra[key] = tool_record[key] + if extra is not None or tool_record is not None: + entry.extra = current_extra + + entry.updated_at = now + session.add(entry) + session.commit() + session.refresh(entry) + entry.embedding = self._normalize_embedding(entry.embedding) + + return self._cache_entry(entry) + + def delete_entry(self, entry_id: str) -> None: + from sqlmodel import delete + + with self._sessions.session() as session: + session.exec(delete(self._sqla_models.Entry).where(self._sqla_models.Entry.id == entry_id)) + session.commit() + self.entries.pop(entry_id, None) + + def vector_search_entries( + self, + query_vec: list[float], + top_k: int, + where: Mapping[str, Any] | None = None, + *, + lane: str | None = None, + ranking: str = "similarity", + recency_decay_days: float = 30.0, + ) -> list[tuple[str, float]]: + if not self._use_vector or ranking == "salience": + return self._vector_search_local( + query_vec, top_k, where=where, lane=lane, ranking=ranking, recency_decay_days=recency_decay_days + ) + + from sqlmodel import select + + model = self._sqla_models.Entry + distance = model.embedding.cosine_distance(query_vec) + filters = [model.embedding.isnot(None)] + filters.extend(self._build_filters(model, where)) + if lane is not None: + filters.append(model.lane == lane) + stmt = ( + select(model.id, (1 - distance).label("score")) + .where(*filters) + .order_by(distance) + .limit(top_k) + ) + with self._sessions.session() as session: + rows = session.execute(stmt).all() + return [(rid, float(score)) for rid, score in rows] + + def load_existing(self) -> None: + from sqlmodel import select + + with self._sessions.session() as session: + rows = session.scalars(select(self._sqla_models.Entry)).all() + for row in rows: + row.embedding = self._normalize_embedding(row.embedding) + self._cache_entry(row) + + def _vector_search_local( + self, + query_vec: list[float], + top_k: int, + where: Mapping[str, Any] | None = None, + *, + lane: str | None = None, + ranking: str = "similarity", + recency_decay_days: float = 30.0, + ) -> list[tuple[str, float]]: + pool = self.list_entries(where, lane=lane) + + if ranking == "salience": + corpus = [ + ( + e.id, + e.embedding, + (e.extra or {}).get("reinforcement_count", 1), + self._parse_datetime((e.extra or {}).get("last_reinforced_at")), + ) + for e in pool.values() + ] + return cosine_topk_salience(query_vec, corpus, k=top_k, recency_decay_days=recency_decay_days) + + return cosine_topk(query_vec, [(e.id, e.embedding) for e in pool.values()], k=top_k) + + def _cache_entry(self, entry: Entry) -> Entry: + self.entries[entry.id] = entry + return entry + + @staticmethod + def _parse_datetime(dt_str: str | None) -> datetime | None: + if dt_str is None: + return None + try: + parsed = pendulum.parse(dt_str) + except (ValueError, TypeError): + return None + else: + if isinstance(parsed, datetime): + return parsed + return None + + +__all__ = ["PostgresEntryRepo"] diff --git a/src/memu/database/postgres/repositories/memory_category_repo.py b/src/memu/database/postgres/repositories/memory_category_repo.py deleted file mode 100644 index 229cd200..00000000 --- a/src/memu/database/postgres/repositories/memory_category_repo.py +++ /dev/null @@ -1,162 +0,0 @@ -from __future__ import annotations - -from collections.abc import Mapping -from typing import Any - -from memu.database.models import MemoryCategory -from memu.database.postgres.repositories.base import PostgresRepoBase -from memu.database.postgres.session import SessionManager -from memu.database.repositories.memory_category import MemoryCategoryRepo -from memu.database.state import DatabaseState - - -class PostgresMemoryCategoryRepo(PostgresRepoBase, MemoryCategoryRepo): - def __init__( - self, - *, - state: DatabaseState, - memory_category_model: type[MemoryCategory], - sqla_models: Any, - sessions: SessionManager, - scope_fields: list[str], - ) -> None: - super().__init__(state=state, sqla_models=sqla_models, sessions=sessions, scope_fields=scope_fields) - self._memory_category_model = memory_category_model - self.categories: dict[str, MemoryCategory] = self._state.categories - - def list_categories(self, where: Mapping[str, Any] | None = None) -> dict[str, MemoryCategory]: - from sqlmodel import select - - filters = self._build_filters(self._sqla_models.MemoryCategory, where) - with self._sessions.session() as session: - rows = session.scalars(select(self._sqla_models.MemoryCategory).where(*filters)).all() - result: dict[str, MemoryCategory] = {} - for row in rows: - row.embedding = self._normalize_embedding(row.embedding) - cat = self._cache_category(row) - result[cat.id] = cat - return result - - def clear_categories(self, where: Mapping[str, Any] | None = None) -> dict[str, MemoryCategory]: - from sqlmodel import delete, select - - filters = self._build_filters(self._sqla_models.MemoryCategory, where) - with self._sessions.session() as session: - # First get the objects to delete - rows = session.scalars(select(self._sqla_models.MemoryCategory).where(*filters)).all() - deleted: dict[str, MemoryCategory] = {} - for row in rows: - row.embedding = self._normalize_embedding(row.embedding) - deleted[row.id] = row - - if not deleted: - return {} - - # Delete from database - session.exec(delete(self._sqla_models.MemoryCategory).where(*filters)) - session.commit() - - # Clean up cache - for cat_id in deleted: - self.categories.pop(cat_id, None) - - return deleted - - def get_or_create_category( - self, - *, - name: str, - description: str, - embedding: list[float], - user_data: dict[str, Any], - ) -> MemoryCategory: - from sqlmodel import select - - now = self._now() - with self._sessions.session() as session: - filters = [self._sqla_models.MemoryCategory.name == name] - for key, value in user_data.items(): - filters.append(getattr(self._sqla_models.MemoryCategory, key) == value) - existing = session.scalar(select(self._sqla_models.MemoryCategory).where(*filters)) - - if existing: - updated = False - if getattr(existing, "embedding", None) is None: - existing.embedding = self._prepare_embedding(embedding) - updated = True - if getattr(existing, "description", None) is None: - existing.description = description - updated = True - if updated: - existing.updated_at = now - session.add(existing) - session.commit() - session.refresh(existing) - return self._cache_category(existing) - - cat = self._memory_category_model( - name=name, - description=description, - embedding=self._prepare_embedding(embedding), - created_at=now, - updated_at=now, - **user_data, - ) - session.add(cat) - session.commit() - session.refresh(cat) - - return self._cache_category(cat) - - def update_category( - self, - *, - category_id: str, - name: str | None = None, - description: str | None = None, - embedding: list[float] | None = None, - summary: str | None = None, - ) -> MemoryCategory: - from sqlmodel import select - - now = self._now() - with self._sessions.session() as session: - cat = session.scalar( - select(self._sqla_models.MemoryCategory).where(self._sqla_models.MemoryCategory.id == category_id) - ) - if cat is None: - msg = f"Category with id {category_id} not found" - raise KeyError(msg) - - if name is not None: - cat.name = name - if description is not None: - cat.description = description - if embedding is not None: - cat.embedding = self._prepare_embedding(embedding) - if summary is not None: - cat.summary = summary - - cat.updated_at = now - session.add(cat) - session.commit() - session.refresh(cat) - cat.embedding = self._normalize_embedding(cat.embedding) - - return self._cache_category(cat) - - def load_existing(self) -> None: - from sqlmodel import select - - with self._sessions.session() as session: - rows = session.scalars(select(self._sqla_models.MemoryCategory)).all() - for row in rows: - row.embedding = self._normalize_embedding(row.embedding) - self._cache_category(row) - - def _cache_category(self, cat: MemoryCategory) -> MemoryCategory: - self.categories[cat.id] = cat - return cat - - -__all__ = ["PostgresMemoryCategoryRepo"] diff --git a/src/memu/database/postgres/repositories/memory_item_repo.py b/src/memu/database/postgres/repositories/memory_item_repo.py deleted file mode 100644 index 6d04f61b..00000000 --- a/src/memu/database/postgres/repositories/memory_item_repo.py +++ /dev/null @@ -1,401 +0,0 @@ -from __future__ import annotations - -import math -from collections.abc import Mapping -from datetime import datetime -from typing import Any - -from memu.database.models import MemoryItem, MemoryType, compute_content_hash -from memu.database.postgres.repositories.base import PostgresRepoBase -from memu.database.postgres.session import SessionManager -from memu.database.state import DatabaseState - - -class PostgresMemoryItemRepo(PostgresRepoBase): - def __init__( - self, - *, - state: DatabaseState, - memory_item_model: type[MemoryItem], - sqla_models: Any, - sessions: SessionManager, - scope_fields: list[str], - use_vector: bool, - ) -> None: - super().__init__( - state=state, sqla_models=sqla_models, sessions=sessions, scope_fields=scope_fields, use_vector=use_vector - ) - self._memory_item_model = memory_item_model - self.items: dict[str, MemoryItem] = self._state.items - - def get_item(self, memory_id: str) -> MemoryItem | None: - from sqlmodel import select - - with self._sessions.session() as session: - row = session.scalar( - select(self._sqla_models.MemoryItem).where(self._sqla_models.MemoryItem.id == memory_id) - ) - if row: - row.embedding = self._normalize_embedding(row.embedding) - return self._cache_item(row) - return None - - def list_items(self, where: Mapping[str, Any] | None = None) -> dict[str, MemoryItem]: - from sqlmodel import select - - filters = self._build_filters(self._sqla_models.MemoryItem, where) - with self._sessions.session() as session: - rows = session.scalars(select(self._sqla_models.MemoryItem).where(*filters)).all() - result: dict[str, MemoryItem] = {} - for row in rows: - row.embedding = self._normalize_embedding(row.embedding) - item = self._cache_item(row) - result[item.id] = item - return result - - def list_items_by_ref_ids( - self, ref_ids: list[str], where: Mapping[str, Any] | None = None - ) -> dict[str, MemoryItem]: - """List items by their ref_id in the extra column. - - Args: - ref_ids: List of ref_ids to query. - where: Additional filter conditions. - - Returns: - Dict mapping item_id -> MemoryItem for items whose extra->>'ref_id' is in ref_ids. - """ - if not ref_ids: - return {} - - from sqlmodel import select - - filters = self._build_filters(self._sqla_models.MemoryItem, where) - # Add filter for extra->>'ref_id' IN ref_ids (only rows with ref_id key) - ref_id_col = self._sqla_models.MemoryItem.extra["ref_id"].astext - filters.append(ref_id_col.isnot(None)) - filters.append(ref_id_col.in_(ref_ids)) - - with self._sessions.session() as session: - rows = session.scalars(select(self._sqla_models.MemoryItem).where(*filters)).all() - result: dict[str, MemoryItem] = {} - for row in rows: - row.embedding = self._normalize_embedding(row.embedding) - item = self._cache_item(row) - result[item.id] = item - return result - - def clear_items(self, where: Mapping[str, Any] | None = None) -> dict[str, MemoryItem]: - from sqlmodel import delete, select - - filters = self._build_filters(self._sqla_models.MemoryItem, where) - with self._sessions.session() as session: - # First get the objects to delete - rows = session.scalars(select(self._sqla_models.MemoryItem).where(*filters)).all() - deleted: dict[str, MemoryItem] = {} - for row in rows: - row.embedding = self._normalize_embedding(row.embedding) - deleted[row.id] = row - - if not deleted: - return {} - - # Delete from database - session.exec(delete(self._sqla_models.MemoryItem).where(*filters)) - session.commit() - - # Clean up cache - for item_id in deleted: - self.items.pop(item_id, None) - - return deleted - - def create_item( - self, - *, - resource_id: str | None = None, - memory_type: MemoryType, - summary: str, - embedding: list[float], - user_data: dict[str, Any], - reinforce: bool = False, - tool_record: dict[str, Any] | None = None, - ) -> MemoryItem: - if reinforce and memory_type != "tool": - return self.create_item_reinforce( - resource_id=resource_id, - memory_type=memory_type, - summary=summary, - embedding=embedding, - user_data=user_data, - ) - - # Build extra dict with tool_record fields at top level - extra: dict[str, Any] = {} - if tool_record: - if tool_record.get("when_to_use") is not None: - extra["when_to_use"] = tool_record["when_to_use"] - if tool_record.get("metadata") is not None: - extra["metadata"] = tool_record["metadata"] - if tool_record.get("tool_calls") is not None: - extra["tool_calls"] = tool_record["tool_calls"] - - item = self._memory_item_model( - resource_id=resource_id, - memory_type=memory_type, - summary=summary, - embedding=self._prepare_embedding(embedding), - extra=extra if extra else {}, - **user_data, - created_at=self._now(), - updated_at=self._now(), - ) - - with self._sessions.session() as session: - session.add(item) - session.commit() - session.refresh(item) - - self.items[item.id] = item - return item - - def create_item_reinforce( - self, - *, - resource_id: str | None = None, - memory_type: MemoryType, - summary: str, - embedding: list[float], - user_data: dict[str, Any], - ) -> MemoryItem: - from sqlmodel import select - - content_hash = compute_content_hash(summary, memory_type) - - with self._sessions.session() as session: - # Check for existing item with same hash in same scope (deduplication) - # Use extra->>'content_hash' for query performance - content_hash_col = self._sqla_models.MemoryItem.extra["content_hash"].astext - filters = [content_hash_col == content_hash] - filters.extend(self._build_filters(self._sqla_models.MemoryItem, user_data)) - - existing = session.scalar(select(self._sqla_models.MemoryItem).where(*filters)) - - if existing: - # Reinforce existing memory instead of creating duplicate - current_extra = existing.extra or {} - current_count = current_extra.get("reinforcement_count", 1) - existing.extra = { - **current_extra, - "reinforcement_count": current_count + 1, - "last_reinforced_at": self._now().isoformat(), - } - existing.updated_at = self._now() - session.add(existing) - session.commit() - session.refresh(existing) - existing.embedding = self._normalize_embedding(existing.embedding) - return self._cache_item(existing) - - # Create new item with salience tracking in extra - now = self._now() - - item = self._memory_item_model( - resource_id=resource_id, - memory_type=memory_type, - summary=summary, - embedding=self._prepare_embedding(embedding), - **user_data, - created_at=now, - updated_at=now, - extra={ - "content_hash": content_hash, - "reinforcement_count": 1, - "last_reinforced_at": now.isoformat(), - }, - ) - - session.add(item) - session.commit() - session.refresh(item) - - self.items[item.id] = item - return item - - def update_item( - self, - *, - item_id: str, - memory_type: MemoryType | None = None, - summary: str | None = None, - embedding: list[float] | None = None, - extra: dict[str, Any] | None = None, - tool_record: dict[str, Any] | None = None, - ) -> MemoryItem: - from sqlmodel import select - - now = self._now() - with self._sessions.session() as session: - item = session.scalar( - select(self._sqla_models.MemoryItem).where(self._sqla_models.MemoryItem.id == item_id) - ) - if item is None: - msg = f"Item with id {item_id} not found" - raise KeyError(msg) - - if memory_type is not None: - item.memory_type = memory_type - if summary is not None: - item.summary = summary - if embedding is not None: - item.embedding = self._prepare_embedding(embedding) - - # Merge extra and tool_record into existing extra dict - current_extra = item.extra or {} - if extra is not None: - current_extra = {**current_extra, **extra} - if tool_record is not None: - # Merge tool_record fields at top level - for key in ("when_to_use", "metadata", "tool_calls"): - if tool_record.get(key) is not None: - current_extra[key] = tool_record[key] - if extra is not None or tool_record is not None: - item.extra = current_extra - - item.updated_at = now - session.add(item) - session.commit() - session.refresh(item) - item.embedding = self._normalize_embedding(item.embedding) - - return self._cache_item(item) - - def delete_item(self, item_id: str) -> None: - from sqlmodel import delete - - with self._sessions.session() as session: - session.exec(delete(self._sqla_models.MemoryItem).where(self._sqla_models.MemoryItem.id == item_id)) - session.commit() - - def vector_search_items( - self, - query_vec: list[float], - top_k: int, - where: Mapping[str, Any] | None = None, - *, - ranking: str = "similarity", - recency_decay_days: float = 30.0, - ) -> list[tuple[str, float]]: - if not self._use_vector or ranking == "salience": - # For salience ranking or when pgvector is not available, use local search - return self._vector_search_local( - query_vec, top_k, where=where, ranking=ranking, recency_decay_days=recency_decay_days - ) - - from sqlmodel import select - - distance = self._sqla_models.MemoryItem.embedding.cosine_distance(query_vec) - filters = [self._sqla_models.MemoryItem.embedding.isnot(None)] - filters.extend(self._build_filters(self._sqla_models.MemoryItem, where)) - stmt = ( - select(self._sqla_models.MemoryItem.id, (1 - distance).label("score")) - .where(*filters) - .order_by(distance) - .limit(top_k) - ) - with self._sessions.session() as session: - rows = session.execute(stmt).all() - return [(rid, float(score)) for rid, score in rows] - - def load_existing(self) -> None: - from sqlmodel import select - - with self._sessions.session() as session: - rows = session.scalars(select(self._sqla_models.MemoryItem)).all() - for row in rows: - row.embedding = self._normalize_embedding(row.embedding) - self._cache_item(row) - - def _vector_search_local( - self, - query_vec: list[float], - top_k: int, - where: Mapping[str, Any] | None = None, - *, - ranking: str = "similarity", - recency_decay_days: float = 30.0, - ) -> list[tuple[str, float]]: - scored: list[tuple[str, float]] = [] - for item in self.items.values(): - if item.embedding is None: - continue - if not self._matches_where(item, where): - continue - - similarity = self._cosine(query_vec, item.embedding) - - if ranking == "salience": - # Salience-aware scoring - read from extra dict - extra = item.extra or {} - reinforcement_count = extra.get("reinforcement_count", 1) - last_reinforced_at = self._parse_datetime(extra.get("last_reinforced_at")) - score = self._salience_score( - similarity, - reinforcement_count, - last_reinforced_at, - recency_decay_days, - ) - else: - score = similarity - - scored.append((item.id, score)) - - scored.sort(key=lambda x: x[1], reverse=True) - return scored[:top_k] - - @staticmethod - def _salience_score( - similarity: float, - reinforcement_count: int, - last_reinforced_at: datetime | None, - recency_decay_days: float, - ) -> float: - """Compute salience score: similarity * reinforcement * recency.""" - reinforcement_factor = math.log(reinforcement_count + 1) - - if last_reinforced_at is None: - recency_factor = 0.5 - else: - now = datetime.now(last_reinforced_at.tzinfo) if last_reinforced_at.tzinfo else datetime.utcnow() - days_ago = (now - last_reinforced_at).total_seconds() / 86400 - recency_factor = math.exp(-0.693 * days_ago / recency_decay_days) - - return similarity * reinforcement_factor * recency_factor - - def _cache_item(self, item: MemoryItem) -> MemoryItem: - self.items[item.id] = item - return item - - @staticmethod - def _parse_datetime(dt_str: str | None) -> datetime | None: - """Parse ISO datetime string from extra dict.""" - if dt_str is None: - return None - try: - import pendulum - - parsed = pendulum.parse(dt_str) - except (ValueError, TypeError): - return None - else: - if isinstance(parsed, datetime): - return parsed - return None - - @staticmethod - def _cosine(a: list[float], b: list[float]) -> float: - denom = (sum(x * x for x in a) ** 0.5) * (sum(y * y for y in b) ** 0.5) + 1e-9 - return float(sum(x * y for x, y in zip(a, b, strict=True)) / denom) - - -__all__ = ["PostgresMemoryItemRepo"] diff --git a/src/memu/database/postgres/repositories/resource_entry_repo.py b/src/memu/database/postgres/repositories/resource_entry_repo.py new file mode 100644 index 00000000..388bf78e --- /dev/null +++ b/src/memu/database/postgres/repositories/resource_entry_repo.py @@ -0,0 +1,148 @@ +from __future__ import annotations + +from collections.abc import Mapping +from typing import Any + +from memu.database.models import ResourceEntry +from memu.database.postgres.repositories.base import PostgresRepoBase +from memu.database.postgres.session import SessionManager +from memu.database.repositories.resource_entry import ResourceEntryRepo +from memu.database.state import DatabaseState + + +class PostgresResourceEntryRepo(PostgresRepoBase, ResourceEntryRepo): + def __init__( + self, + *, + state: DatabaseState, + resource_entry_model: type[ResourceEntry], + sqla_models: Any, + sessions: SessionManager, + scope_fields: list[str], + ) -> None: + super().__init__(state=state, sqla_models=sqla_models, sessions=sessions, scope_fields=scope_fields) + self._resource_entry_model = resource_entry_model + self.relations: list[ResourceEntry] = self._state.relations + + def list_relations(self, where: Mapping[str, Any] | None = None) -> list[ResourceEntry]: + from sqlmodel import select + + filters = self._build_filters(self._sqla_models.ResourceEntry, where) + with self._sessions.session() as session: + rows = session.scalars(select(self._sqla_models.ResourceEntry).where(*filters)).all() + return [self._row_to_record(row) for row in rows] + + def link_entry_resource(self, entry_id: str, resource_id: str, user_data: dict[str, Any]) -> ResourceEntry: + from sqlmodel import select + + for rel in self.relations: + if rel.entry_id == entry_id and rel.resource_id == resource_id: + return rel + + now = self._now() + new_rel = self._resource_entry_model( + entry_id=entry_id, + resource_id=resource_id, + **user_data, + created_at=now, + updated_at=now, + ) + + with self._sessions.session() as session: + existing = session.scalar( + select(self._sqla_models.ResourceEntry).where( + self._sqla_models.ResourceEntry.entry_id == entry_id, + self._sqla_models.ResourceEntry.resource_id == resource_id, + ) + ) + if existing: + return self._cache_relation(self._row_to_record(existing)) + + session.add(new_rel) + session.commit() + session.refresh(new_rel) + record = self._row_to_record(new_rel) + + return self._cache_relation(record) + + def unlink_entry_resource(self, entry_id: str, resource_id: str) -> None: + from sqlmodel import delete + + with self._sessions.session() as session: + session.exec( + delete(self._sqla_models.ResourceEntry).where( + self._sqla_models.ResourceEntry.entry_id == entry_id, + self._sqla_models.ResourceEntry.resource_id == resource_id, + ) + ) + session.commit() + self.relations[:] = [ + r for r in self.relations if not (r.entry_id == entry_id and r.resource_id == resource_id) + ] + + def unlink_entry(self, entry_id: str) -> list[ResourceEntry]: + from sqlmodel import delete, select + + with self._sessions.session() as session: + rows = session.scalars( + select(self._sqla_models.ResourceEntry).where(self._sqla_models.ResourceEntry.entry_id == entry_id) + ).all() + removed = [self._row_to_record(row) for row in rows] + if removed: + session.exec( + delete(self._sqla_models.ResourceEntry).where( + self._sqla_models.ResourceEntry.entry_id == entry_id + ) + ) + session.commit() + self.relations[:] = [r for r in self.relations if r.entry_id != entry_id] + return removed + + def clear_relations(self, where: Mapping[str, Any] | None = None) -> list[ResourceEntry]: + from sqlmodel import delete, select + + filters = self._build_filters(self._sqla_models.ResourceEntry, where) + with self._sessions.session() as session: + rows = session.scalars(select(self._sqla_models.ResourceEntry).where(*filters)).all() + removed = [self._row_to_record(row) for row in rows] + if removed: + session.exec(delete(self._sqla_models.ResourceEntry).where(*filters)) + session.commit() + removed_ids = {rel.id for rel in removed} + self.relations[:] = [r for r in self.relations if r.id not in removed_ids] + return removed + + def get_entry_resources(self, entry_id: str) -> list[ResourceEntry]: + from sqlmodel import select + + with self._sessions.session() as session: + rows = session.scalars( + select(self._sqla_models.ResourceEntry).where(self._sqla_models.ResourceEntry.entry_id == entry_id) + ).all() + return [self._row_to_record(row) for row in rows] + + def load_existing(self) -> None: + from sqlmodel import select + + with self._sessions.session() as session: + rows = session.scalars(select(self._sqla_models.ResourceEntry)).all() + self.relations.clear() + for row in rows: + self._cache_relation(self._row_to_record(row)) + + def _row_to_record(self, row: Any) -> ResourceEntry: + return ResourceEntry( + id=row.id, + entry_id=row.entry_id, + resource_id=row.resource_id, + created_at=row.created_at, + updated_at=row.updated_at, + **self._scope_kwargs_from(row), + ) + + def _cache_relation(self, rel: ResourceEntry) -> ResourceEntry: + self.relations.append(rel) + return rel + + +__all__ = ["PostgresResourceEntryRepo"] diff --git a/src/memu/database/postgres/repositories/resource_repo.py b/src/memu/database/postgres/repositories/resource_repo.py index 2efcc848..56caca8a 100644 --- a/src/memu/database/postgres/repositories/resource_repo.py +++ b/src/memu/database/postgres/repositories/resource_repo.py @@ -3,7 +3,7 @@ from collections.abc import Mapping from typing import Any -from memu.database.models import Resource +from memu.database.models import MARKDOWN_MODALITY, Resource from memu.database.postgres.repositories.base import PostgresRepoBase from memu.database.postgres.session import SessionManager from memu.database.repositories.resource import ResourceRepo @@ -20,17 +20,33 @@ def __init__( sqla_models: Any, sessions: SessionManager, scope_fields: list[str], + use_vector: bool = True, ) -> None: - super().__init__(state=state, sqla_models=sqla_models, sessions=sessions, scope_fields=scope_fields) + super().__init__( + state=state, sqla_models=sqla_models, sessions=sessions, scope_fields=scope_fields, use_vector=use_vector + ) self._resource_model = resource_model self.resources: dict[str, Resource] = self._state.resources - def list_resources(self, where: Mapping[str, Any] | None = None) -> dict[str, Resource]: + def get_resource(self, resource_id: str) -> Resource | None: from sqlmodel import select - filters = self._build_filters(self._sqla_models.Resource, where) with self._sessions.session() as session: - rows = session.scalars(select(self._sqla_models.Resource).where(*filters)).all() + row = session.scalar(select(self._sqla_models.Resource).where(self._sqla_models.Resource.id == resource_id)) + if row: + row.embedding = self._normalize_embedding(row.embedding) + return self._cache_resource(row) + return None + + def list_resources(self, where: Mapping[str, Any] | None = None, *, lane: str | None = None) -> dict[str, Resource]: + from sqlmodel import select + + model = self._sqla_models.Resource + filters = self._build_filters(model, where) + if lane is not None: + filters.append(model.lane == lane) + with self._sessions.session() as session: + rows = session.scalars(select(model).where(*filters)).all() result: dict[str, Resource] = {} for row in rows: row.embedding = self._normalize_embedding(row.embedding) @@ -38,13 +54,17 @@ def list_resources(self, where: Mapping[str, Any] | None = None) -> dict[str, Re result[res.id] = res return result - def clear_resources(self, where: Mapping[str, Any] | None = None) -> dict[str, Resource]: + def clear_resources( + self, where: Mapping[str, Any] | None = None, *, lane: str | None = None + ) -> dict[str, Resource]: from sqlmodel import delete, select - filters = self._build_filters(self._sqla_models.Resource, where) + model = self._sqla_models.Resource + filters = self._build_filters(model, where) + if lane is not None: + filters.append(model.lane == lane) with self._sessions.session() as session: - # First get the objects to delete - rows = session.scalars(select(self._sqla_models.Resource).where(*filters)).all() + rows = session.scalars(select(model).where(*filters)).all() deleted: dict[str, Resource] = {} for row in rows: row.embedding = self._normalize_embedding(row.embedding) @@ -53,18 +73,15 @@ def clear_resources(self, where: Mapping[str, Any] | None = None) -> dict[str, R if not deleted: return {} - # Delete from database - session.exec(delete(self._sqla_models.Resource).where(*filters)) + session.exec(delete(model).where(*filters)) session.commit() - # Clean up cache for res_id in deleted: self.resources.pop(res_id, None) return deleted def delete_resource(self, resource_id: str) -> None: - """Delete a single resource by id (used for cascade sync).""" from sqlmodel import delete with self._sessions.session() as session: @@ -75,19 +92,31 @@ def delete_resource(self, resource_id: str) -> None: def create_resource( self, *, - url: str, modality: str, - local_path: str, - caption: str | None, - embedding: list[float] | None, user_data: dict[str, Any], + lane: str = "source", + url: str | None = None, + local_path: str | None = None, + slug: str | None = None, + title: str | None = None, + description: str | None = None, + content: str | None = None, + summary: str | None = None, + embedding: list[float] | None = None, + resource_refs: list[dict[str, Any]] | None = None, ) -> Resource: res = self._resource_model( - url=url, + lane=lane, modality=modality, + url=url, local_path=local_path, - caption=caption, + slug=slug, + title=title, + description=description, + content=content, + summary=summary, embedding=self._prepare_embedding(embedding), + resource_refs=resource_refs or [], **user_data, created_at=self._now(), updated_at=self._now(), @@ -98,6 +127,103 @@ def create_resource( session.commit() session.refresh(res) + res.embedding = self._normalize_embedding(res.embedding) + return self._cache_resource(res) + + def get_or_create_doc( + self, + *, + lane: str, + title: str, + description: str, + embedding: list[float], + user_data: dict[str, Any], + slug: str | None = None, + ) -> Resource: + from sqlmodel import select + + model = self._sqla_models.Resource + now = self._now() + with self._sessions.session() as session: + filters = [model.lane == lane, model.title == title] + for key, value in user_data.items(): + filters.append(getattr(model, key) == value) + existing = session.scalar(select(model).where(*filters)) + + if existing: + updated = False + if getattr(existing, "embedding", None) is None: + existing.embedding = self._prepare_embedding(embedding) + updated = True + if not getattr(existing, "description", None): + existing.description = description + updated = True + if updated: + existing.updated_at = now + session.add(existing) + session.commit() + session.refresh(existing) + existing.embedding = self._normalize_embedding(existing.embedding) + return self._cache_resource(existing) + + res = self._resource_model( + lane=lane, + modality=MARKDOWN_MODALITY, + title=title, + slug=slug, + description=description, + embedding=self._prepare_embedding(embedding), + created_at=now, + updated_at=now, + **user_data, + ) + session.add(res) + session.commit() + session.refresh(res) + + res.embedding = self._normalize_embedding(res.embedding) + return self._cache_resource(res) + + def update_resource( + self, + *, + resource_id: str, + title: str | None = None, + description: str | None = None, + content: str | None = None, + summary: str | None = None, + embedding: list[float] | None = None, + resource_refs: list[dict[str, Any]] | None = None, + ) -> Resource: + from sqlmodel import select + + model = self._sqla_models.Resource + now = self._now() + with self._sessions.session() as session: + res = session.scalar(select(model).where(model.id == resource_id)) + if res is None: + msg = f"Resource with id {resource_id} not found" + raise KeyError(msg) + + if title is not None: + res.title = title + if description is not None: + res.description = description + if content is not None: + res.content = content + if summary is not None: + res.summary = summary + if embedding is not None: + res.embedding = self._prepare_embedding(embedding) + if resource_refs is not None: + res.resource_refs = resource_refs + + res.updated_at = now + session.add(res) + session.commit() + session.refresh(res) + res.embedding = self._normalize_embedding(res.embedding) + return self._cache_resource(res) def vector_search_resources( @@ -105,15 +231,29 @@ def vector_search_resources( query_vec: list[float], top_k: int, where: Mapping[str, Any] | None = None, + *, + lane: str | None = None, ) -> list[tuple[str, float]]: """Rank resources by cosine similarity over stored embeddings. - Resource captions are not indexed in pgvector, so this scores the loaded - resources in Python, matching the previous app-layer behavior. + Uses pgvector ``cosine_distance`` when available, otherwise falls back to + scoring the loaded resources in Python. """ - pool = self.list_resources(where) - corpus = [(rid, res.embedding) for rid, res in pool.items() if res.embedding] - return cosine_topk(query_vec, corpus, k=top_k) + if not self._use_vector: + return self._vector_search_local(query_vec, top_k, where=where, lane=lane) + + from sqlmodel import select + + model = self._sqla_models.Resource + distance = model.embedding.cosine_distance(query_vec) + filters = [model.embedding.isnot(None)] + filters.extend(self._build_filters(model, where)) + if lane is not None: + filters.append(model.lane == lane) + stmt = select(model.id, (1 - distance).label("score")).where(*filters).order_by(distance).limit(top_k) + with self._sessions.session() as session: + rows = session.execute(stmt).all() + return [(rid, float(score)) for rid, score in rows] def load_existing(self) -> None: from sqlmodel import select @@ -124,6 +264,18 @@ def load_existing(self) -> None: row.embedding = self._normalize_embedding(row.embedding) self._cache_resource(row) + def _vector_search_local( + self, + query_vec: list[float], + top_k: int, + where: Mapping[str, Any] | None = None, + *, + lane: str | None = None, + ) -> list[tuple[str, float]]: + pool = self.list_resources(where, lane=lane) + corpus = [(rid, res.embedding) for rid, res in pool.items() if res.embedding] + return cosine_topk(query_vec, corpus, k=top_k) + def _cache_resource(self, res: Resource) -> Resource: self.resources[res.id] = res return res diff --git a/src/memu/database/postgres/schema.py b/src/memu/database/postgres/schema.py index ac6e8b52..88f8f973 100644 --- a/src/memu/database/postgres/schema.py +++ b/src/memu/database/postgres/schema.py @@ -24,10 +24,9 @@ raise ImportError(msg) from exc from memu.database.postgres.models import ( - CategoryItemModel, - MemoryCategoryModel, - MemoryItemModel, - ResourceModel, + PostgresEntryModel, + PostgresResourceEntryModel, + PostgresResourceModel, build_table_model, ) @@ -36,9 +35,8 @@ class SQLAModels: Base: type[Any] Resource: type[Any] - MemoryCategory: type[Any] - MemoryItem: type[Any] - CategoryItem: type[Any] + Entry: type[Any] + ResourceEntry: type[Any] _MODEL_CACHE: dict[type[Any], SQLAModels] = {} @@ -63,26 +61,20 @@ def get_sqlalchemy_models(*, scope_model: type[BaseModel] | None = None) -> SQLA resource_model = build_table_model( scope, - ResourceModel, - tablename="resources", + PostgresResourceModel, + tablename="memu_resources", metadata=metadata_obj, ) - memory_category_model = build_table_model( + entry_model = build_table_model( scope, - MemoryCategoryModel, - tablename="memory_categories", + PostgresEntryModel, + tablename="memu_entries", metadata=metadata_obj, ) - memory_item_model = build_table_model( + resource_entry_model = build_table_model( scope, - MemoryItemModel, - tablename="memory_items", - metadata=metadata_obj, - ) - category_item_model = build_table_model( - scope, - CategoryItemModel, - tablename="category_items", + PostgresResourceEntryModel, + tablename="memu_resource_entries", metadata=metadata_obj, ) @@ -93,9 +85,8 @@ class Base(SQLModel): models = SQLAModels( Base=Base, Resource=resource_model, - MemoryCategory=memory_category_model, - MemoryItem=memory_item_model, - CategoryItem=category_item_model, + Entry=entry_model, + ResourceEntry=resource_entry_model, ) _MODEL_CACHE[cache_key] = models return models diff --git a/src/memu/database/repositories/__init__.py b/src/memu/database/repositories/__init__.py index bec26664..844dfa07 100644 --- a/src/memu/database/repositories/__init__.py +++ b/src/memu/database/repositories/__init__.py @@ -1,6 +1,5 @@ -from memu.database.repositories.category_item import CategoryItemRepo -from memu.database.repositories.memory_category import MemoryCategoryRepo -from memu.database.repositories.memory_item import MemoryItemRepo +from memu.database.repositories.entry import EntryRepo from memu.database.repositories.resource import ResourceRepo +from memu.database.repositories.resource_entry import ResourceEntryRepo -__all__ = ["CategoryItemRepo", "MemoryCategoryRepo", "MemoryItemRepo", "ResourceRepo"] +__all__ = ["EntryRepo", "ResourceEntryRepo", "ResourceRepo"] diff --git a/src/memu/database/repositories/category_item.py b/src/memu/database/repositories/category_item.py deleted file mode 100644 index dc3104fa..00000000 --- a/src/memu/database/repositories/category_item.py +++ /dev/null @@ -1,31 +0,0 @@ -from __future__ import annotations - -from collections.abc import Mapping -from typing import Any, Protocol, runtime_checkable - -from memu.database.models import CategoryItem - - -@runtime_checkable -class CategoryItemRepo(Protocol): - """Repository contract for item/category relations.""" - - relations: list[CategoryItem] - - def list_relations(self, where: Mapping[str, Any] | None = None) -> list[CategoryItem]: ... - - def link_item_category(self, item_id: str, cat_id: str, user_data: dict[str, Any]) -> CategoryItem: ... - - def unlink_item_category(self, item_id: str, cat_id: str) -> None: ... - - def unlink_item(self, item_id: str) -> list[CategoryItem]: - """Remove all relations for a given item. Returns the removed relations.""" - ... - - def clear_relations(self, where: Mapping[str, Any] | None = None) -> list[CategoryItem]: - """Remove all relations matching the scope. Returns the removed relations.""" - ... - - def get_item_categories(self, item_id: str) -> list[CategoryItem]: ... - - def load_existing(self) -> None: ... diff --git a/src/memu/database/repositories/entry.py b/src/memu/database/repositories/entry.py new file mode 100644 index 00000000..2a4b7a1d --- /dev/null +++ b/src/memu/database/repositories/entry.py @@ -0,0 +1,67 @@ +from __future__ import annotations + +from collections.abc import Mapping +from typing import Any, Literal, Protocol, runtime_checkable + +from memu.database.models import Entry + + +@runtime_checkable +class EntryRepo(Protocol): + """Repository contract for lane entries (the searchable atoms).""" + + entries: dict[str, Entry] + + def get_entry(self, entry_id: str) -> Entry | None: ... + + def list_entries( + self, where: Mapping[str, Any] | None = None, *, lane: str | None = None + ) -> dict[str, Entry]: ... + + def clear_entries( + self, where: Mapping[str, Any] | None = None, *, lane: str | None = None + ) -> dict[str, Entry]: ... + + def create_entry( + self, + *, + lane: str, + source_id: str | None, + entry_kind: str, + text: str, + embedding: list[float], + user_data: dict[str, Any], + source_path: str | None = None, + reinforce: bool = False, + tool_record: dict[str, Any] | None = None, + ) -> Entry: ... + + def update_entry( + self, + *, + entry_id: str, + entry_kind: str | None = None, + text: str | None = None, + embedding: list[float] | None = None, + extra: dict[str, Any] | None = None, + tool_record: dict[str, Any] | None = None, + ) -> Entry: ... + + def delete_entry(self, entry_id: str) -> None: ... + + def list_entries_by_ref_ids( + self, ref_ids: list[str], where: Mapping[str, Any] | None = None + ) -> dict[str, Entry]: ... + + def vector_search_entries( + self, + query_vec: list[float], + top_k: int, + where: Mapping[str, Any] | None = None, + *, + lane: str | None = None, + ranking: Literal["similarity", "salience"] = "similarity", + recency_decay_days: float = 30.0, + ) -> list[tuple[str, float]]: ... + + def load_existing(self) -> None: ... diff --git a/src/memu/database/repositories/memory_category.py b/src/memu/database/repositories/memory_category.py deleted file mode 100644 index 2a1aefbf..00000000 --- a/src/memu/database/repositories/memory_category.py +++ /dev/null @@ -1,33 +0,0 @@ -from __future__ import annotations - -from collections.abc import Mapping -from typing import Any, Protocol, runtime_checkable - -from memu.database.models import MemoryCategory - - -@runtime_checkable -class MemoryCategoryRepo(Protocol): - """Repository contract for memory categories.""" - - categories: dict[str, MemoryCategory] - - def list_categories(self, where: Mapping[str, Any] | None = None) -> dict[str, MemoryCategory]: ... - - def clear_categories(self, where: Mapping[str, Any] | None = None) -> dict[str, MemoryCategory]: ... - - def get_or_create_category( - self, *, name: str, description: str, embedding: list[float], user_data: dict[str, Any] - ) -> MemoryCategory: ... - - def update_category( - self, - *, - category_id: str, - name: str | None = None, - description: str | None = None, - embedding: list[float] | None = None, - summary: str | None = None, - ) -> MemoryCategory: ... - - def load_existing(self) -> None: ... diff --git a/src/memu/database/repositories/memory_item.py b/src/memu/database/repositories/memory_item.py deleted file mode 100644 index be8e5308..00000000 --- a/src/memu/database/repositories/memory_item.py +++ /dev/null @@ -1,60 +0,0 @@ -from __future__ import annotations - -from collections.abc import Mapping -from typing import Any, Literal, Protocol, runtime_checkable - -from memu.database.models import MemoryItem, MemoryType - - -@runtime_checkable -class MemoryItemRepo(Protocol): - """Repository contract for memory items.""" - - items: dict[str, MemoryItem] - - def get_item(self, item_id: str) -> MemoryItem | None: ... - - def list_items(self, where: Mapping[str, Any] | None = None) -> dict[str, MemoryItem]: ... - - def clear_items(self, where: Mapping[str, Any] | None = None) -> dict[str, MemoryItem]: ... - - def create_item( - self, - *, - resource_id: str, - memory_type: MemoryType, - summary: str, - embedding: list[float], - user_data: dict[str, Any], - reinforce: bool = False, - tool_record: dict[str, Any] | None = None, - ) -> MemoryItem: ... - - def update_item( - self, - *, - item_id: str, - memory_type: MemoryType | None = None, - summary: str | None = None, - embedding: list[float] | None = None, - extra: dict[str, Any] | None = None, - tool_record: dict[str, Any] | None = None, - ) -> MemoryItem: ... - - def delete_item(self, item_id: str) -> None: ... - - def list_items_by_ref_ids( - self, ref_ids: list[str], where: Mapping[str, Any] | None = None - ) -> dict[str, MemoryItem]: ... - - def vector_search_items( - self, - query_vec: list[float], - top_k: int, - where: Mapping[str, Any] | None = None, - *, - ranking: Literal["similarity", "salience"] = "similarity", - recency_decay_days: float = 30.0, - ) -> list[tuple[str, float]]: ... - - def load_existing(self) -> None: ... diff --git a/src/memu/database/repositories/resource.py b/src/memu/database/repositories/resource.py index 34f4938e..234d4a42 100644 --- a/src/memu/database/repositories/resource.py +++ b/src/memu/database/repositories/resource.py @@ -8,25 +8,62 @@ @runtime_checkable class ResourceRepo(Protocol): - """Repository contract for resource records.""" + """Repository contract for resource records (raw inputs and generated lane docs).""" resources: dict[str, Resource] - def list_resources(self, where: Mapping[str, Any] | None = None) -> dict[str, Resource]: ... + def get_resource(self, resource_id: str) -> Resource | None: ... - def clear_resources(self, where: Mapping[str, Any] | None = None) -> dict[str, Resource]: ... + def list_resources( + self, where: Mapping[str, Any] | None = None, *, lane: str | None = None + ) -> dict[str, Resource]: ... + + def clear_resources( + self, where: Mapping[str, Any] | None = None, *, lane: str | None = None + ) -> dict[str, Resource]: ... def delete_resource(self, resource_id: str) -> None: ... def create_resource( self, *, - url: str, modality: str, - local_path: str, - caption: str | None, - embedding: list[float] | None, user_data: dict[str, Any], + lane: str = "source", + url: str | None = None, + local_path: str | None = None, + slug: str | None = None, + title: str | None = None, + description: str | None = None, + content: str | None = None, + summary: str | None = None, + embedding: list[float] | None = None, + resource_refs: list[dict[str, Any]] | None = None, + ) -> Resource: ... + + def get_or_create_doc( + self, + *, + lane: str, + title: str, + description: str, + embedding: list[float], + user_data: dict[str, Any], + slug: str | None = None, + ) -> Resource: + """Get an existing generated doc by ``(lane, title, scope)`` or create one.""" + ... + + def update_resource( + self, + *, + resource_id: str, + title: str | None = None, + description: str | None = None, + content: str | None = None, + summary: str | None = None, + embedding: list[float] | None = None, + resource_refs: list[dict[str, Any]] | None = None, ) -> Resource: ... def vector_search_resources( @@ -34,11 +71,13 @@ def vector_search_resources( query_vec: list[float], top_k: int, where: Mapping[str, Any] | None = None, + *, + lane: str | None = None, ) -> list[tuple[str, float]]: """Rank resources by cosine similarity of their stored embeddings. - Returns a list of ``(resource_id, score)`` tuples ordered by descending - similarity. Resources without an embedding are skipped. + Returns ``(resource_id, score)`` ordered by descending similarity; + resources without an embedding are skipped. """ ... diff --git a/src/memu/database/repositories/resource_entry.py b/src/memu/database/repositories/resource_entry.py new file mode 100644 index 00000000..09b735fc --- /dev/null +++ b/src/memu/database/repositories/resource_entry.py @@ -0,0 +1,31 @@ +from __future__ import annotations + +from collections.abc import Mapping +from typing import Any, Protocol, runtime_checkable + +from memu.database.models import ResourceEntry + + +@runtime_checkable +class ResourceEntryRepo(Protocol): + """Repository contract for entry <-> coarse-resource membership edges.""" + + relations: list[ResourceEntry] + + def list_relations(self, where: Mapping[str, Any] | None = None) -> list[ResourceEntry]: ... + + def link_entry_resource(self, entry_id: str, resource_id: str, user_data: dict[str, Any]) -> ResourceEntry: ... + + def unlink_entry_resource(self, entry_id: str, resource_id: str) -> None: ... + + def unlink_entry(self, entry_id: str) -> list[ResourceEntry]: + """Remove all edges for a given entry. Returns the removed edges.""" + ... + + def clear_relations(self, where: Mapping[str, Any] | None = None) -> list[ResourceEntry]: + """Remove all edges matching the scope. Returns the removed edges.""" + ... + + def get_entry_resources(self, entry_id: str) -> list[ResourceEntry]: ... + + def load_existing(self) -> None: ... diff --git a/src/memu/database/sqlite/__init__.py b/src/memu/database/sqlite/__init__.py index c1a909db..81f2e1c0 100644 --- a/src/memu/database/sqlite/__init__.py +++ b/src/memu/database/sqlite/__init__.py @@ -5,6 +5,7 @@ from pydantic import BaseModel from memu.app.settings import DatabaseConfig +from memu.database.sqlite.schema import get_sqlite_sqlalchemy_models from memu.database.sqlite.sqlite import SQLiteStore @@ -27,9 +28,15 @@ def build_sqlite_database( # Default to a local file if no DSN provided dsn = "sqlite:///memu.db" + # Build the scoped resource/entry/resource-entry table models and wire the store. + sqla_models = get_sqlite_sqlalchemy_models(scope_model=user_model) return SQLiteStore( dsn=dsn, scope_model=user_model, + resource_model=sqla_models.Resource, + entry_model=sqla_models.Entry, + resource_entry_model=sqla_models.ResourceEntry, + sqla_models=sqla_models, ) diff --git a/src/memu/database/sqlite/models.py b/src/memu/database/sqlite/models.py index f2b06e6b..f9642bbe 100644 --- a/src/memu/database/sqlite/models.py +++ b/src/memu/database/sqlite/models.py @@ -11,7 +11,7 @@ from sqlalchemy import JSON, MetaData, String, Text from sqlmodel import Column, DateTime, Field, Index, SQLModel, func -from memu.database.models import CategoryItem, MemoryCategory, MemoryItem, MemoryType, Resource +from memu.database.models import Entry, Resource, ResourceEntry class TZDateTime(DateTime): @@ -42,23 +42,35 @@ class SQLiteBaseModelMixin(SQLModel): class SQLiteResourceModel(SQLiteBaseModelMixin, Resource): - """SQLite resource model.""" + """SQLite resource model. - url: str = Field(sa_column=Column(String, nullable=False)) + A single physical table holds both raw inputs (``lane="source"``) and the + generated lane docs (``lane`` in {index, memory, skill}). + """ + + lane: str = Field(default="source", sa_column=Column(String, nullable=False, index=True)) modality: str = Field(sa_column=Column(String, nullable=False)) - local_path: str = Field(sa_column=Column(String, nullable=False)) - caption: str | None = Field(default=None, sa_column=Column(Text, nullable=True)) + url: str | None = Field(default=None, sa_column=Column(String, nullable=True)) + local_path: str | None = Field(default=None, sa_column=Column(String, nullable=True)) + slug: str | None = Field(default=None, sa_column=Column(String, nullable=True)) + title: str | None = Field(default=None, sa_column=Column(String, nullable=True)) + description: str | None = Field(default=None, sa_column=Column(String, nullable=True)) + content: str | None = Field(default=None, sa_column=Column(Text, nullable=True)) + summary: str | None = Field(default=None, sa_column=Column(Text, nullable=True)) # Override inherited embedding field: SQLite has no native vector type, so store the # vector in a JSON column (a bare ``list`` annotation is not mappable by SQLModel). embedding: list[float] | None = Field(default=None, sa_column=Column(JSON, nullable=True)) + resource_refs: list[dict[str, Any]] = Field(default_factory=list, sa_column=Column(JSON, nullable=True)) -class SQLiteMemoryItemModel(SQLiteBaseModelMixin, MemoryItem): - """SQLite memory item model.""" +class SQLiteEntryModel(SQLiteBaseModelMixin, Entry): + """SQLite entry model (the searchable atoms of a lane).""" - resource_id: str | None = Field(sa_column=Column(String, nullable=True)) - memory_type: MemoryType = Field(sa_column=Column(String, nullable=False)) - summary: str = Field(sa_column=Column(Text, nullable=False)) + lane: str = Field(sa_column=Column(String, nullable=False, index=True)) + source_id: str | None = Field(default=None, sa_column=Column(String, nullable=True)) + source_path: str | None = Field(default=None, sa_column=Column(String, nullable=True)) + entry_kind: str = Field(sa_column=Column(String, nullable=False)) + text: str = Field(sa_column=Column(Text, nullable=False)) # Override inherited embedding field: SQLite has no native vector type, so store the # vector in a JSON column (a bare ``list`` annotation is not mappable by SQLModel). embedding: list[float] | None = Field(default=None, sa_column=Column(JSON, nullable=True)) @@ -66,24 +78,11 @@ class SQLiteMemoryItemModel(SQLiteBaseModelMixin, MemoryItem): extra: dict[str, Any] = Field(default_factory=dict, sa_column=Column(JSON, nullable=True)) -class SQLiteMemoryCategoryModel(SQLiteBaseModelMixin, MemoryCategory): - """SQLite memory category model.""" - - name: str = Field(sa_column=Column(String, nullable=False, index=True)) - description: str = Field(sa_column=Column(Text, nullable=False)) - # Override inherited embedding field: SQLite has no native vector type, so store the - # vector in a JSON column (a bare ``list`` annotation is not mappable by SQLModel). - embedding: list[float] | None = Field(default=None, sa_column=Column(JSON, nullable=True)) - summary: str | None = Field(default=None, sa_column=Column(Text, nullable=True)) - - -class SQLiteCategoryItemModel(SQLiteBaseModelMixin, CategoryItem): - """SQLite category-item relation model.""" - - item_id: str = Field(sa_column=Column(String, nullable=False)) - category_id: str = Field(sa_column=Column(String, nullable=False)) +class SQLiteResourceEntryModel(SQLiteBaseModelMixin, ResourceEntry): + """SQLite entry <-> coarse-resource membership edge model.""" - __table_args__ = (Index("idx_sqlite_category_items_unique", "item_id", "category_id", unique=True),) + entry_id: str = Field(sa_column=Column(String, nullable=False)) + resource_id: str = Field(sa_column=Column(String, nullable=False)) def _normalize_table_args(table_args: Any) -> tuple[list[Any], dict[str, Any]]: @@ -171,9 +170,8 @@ def build_sqlite_table_model( __all__ = [ "SQLiteBaseModelMixin", - "SQLiteCategoryItemModel", - "SQLiteMemoryCategoryModel", - "SQLiteMemoryItemModel", + "SQLiteEntryModel", + "SQLiteResourceEntryModel", "SQLiteResourceModel", "build_sqlite_table_model", ] diff --git a/src/memu/database/sqlite/repositories/__init__.py b/src/memu/database/sqlite/repositories/__init__.py index 1436ae86..506a03b3 100644 --- a/src/memu/database/sqlite/repositories/__init__.py +++ b/src/memu/database/sqlite/repositories/__init__.py @@ -1,15 +1,13 @@ """SQLite repository implementations for MemU.""" from memu.database.sqlite.repositories.base import SQLiteRepoBase -from memu.database.sqlite.repositories.category_item_repo import SQLiteCategoryItemRepo -from memu.database.sqlite.repositories.memory_category_repo import SQLiteMemoryCategoryRepo -from memu.database.sqlite.repositories.memory_item_repo import SQLiteMemoryItemRepo +from memu.database.sqlite.repositories.entry_repo import SQLiteEntryRepo +from memu.database.sqlite.repositories.resource_entry_repo import SQLiteResourceEntryRepo from memu.database.sqlite.repositories.resource_repo import SQLiteResourceRepo __all__ = [ - "SQLiteCategoryItemRepo", - "SQLiteMemoryCategoryRepo", - "SQLiteMemoryItemRepo", + "SQLiteEntryRepo", "SQLiteRepoBase", + "SQLiteResourceEntryRepo", "SQLiteResourceRepo", ] diff --git a/src/memu/database/sqlite/repositories/base.py b/src/memu/database/sqlite/repositories/base.py index 502d8edb..b68efc22 100644 --- a/src/memu/database/sqlite/repositories/base.py +++ b/src/memu/database/sqlite/repositories/base.py @@ -67,12 +67,6 @@ def _prepare_embedding(self, embedding: list[float] | None) -> list[float] | Non return None return list(embedding) - def _merge_and_commit(self, obj: Any) -> None: - """Merge object into session and commit.""" - with self._sessions.session() as session: - session.merge(obj) - session.commit() - def _now(self) -> pendulum.DateTime: """Get current UTC time.""" return pendulum.now("UTC") @@ -100,29 +94,11 @@ def _build_filters(self, model: Any, where: Mapping[str, Any] | None) -> list[An return filters @staticmethod - def _matches_where(obj: Any, where: Mapping[str, Any] | None) -> bool: - """Check if object matches where clause (for in-memory filtering).""" - if not where: - return True - for raw_key, expected in where.items(): - if expected is None: - continue - field, op = [*raw_key.split("__", 1), None][:2] - actual = getattr(obj, str(field), None) - if op == "in": - if isinstance(expected, str): - if actual != expected: - return False - else: - try: - if actual not in expected: - return False - except TypeError: - return False - else: - if actual != expected: - return False - return True + def _lane_filters(model: Any, lane: str | None) -> list[Any]: + """Build an equality filter on the ``lane`` column when ``lane`` is given.""" + if lane is None: + return [] + return [model.lane == lane] __all__ = ["SQLiteRepoBase"] diff --git a/src/memu/database/sqlite/repositories/category_item_repo.py b/src/memu/database/sqlite/repositories/category_item_repo.py deleted file mode 100644 index 76c240d7..00000000 --- a/src/memu/database/sqlite/repositories/category_item_repo.py +++ /dev/null @@ -1,211 +0,0 @@ -"""SQLite category-item relation repository implementation.""" - -from __future__ import annotations - -import logging -from collections.abc import Mapping -from typing import Any - -from sqlmodel import select - -from memu.database.models import CategoryItem -from memu.database.repositories.category_item import CategoryItemRepo -from memu.database.sqlite.repositories.base import SQLiteRepoBase -from memu.database.sqlite.schema import SQLiteSQLAModels -from memu.database.sqlite.session import SQLiteSessionManager -from memu.database.state import DatabaseState - -logger = logging.getLogger(__name__) - - -class SQLiteCategoryItemRepo(SQLiteRepoBase, CategoryItemRepo): - """SQLite implementation of category-item relation repository.""" - - def __init__( - self, - *, - state: DatabaseState, - category_item_model: type[Any], - sqla_models: SQLiteSQLAModels, - sessions: SQLiteSessionManager, - scope_fields: list[str], - ) -> None: - """Initialize category-item repository. - - Args: - state: Shared database state for caching. - category_item_model: SQLModel class for category-item relations. - sqla_models: SQLAlchemy model container. - sessions: Session manager for database connections. - scope_fields: List of user scope field names. - """ - super().__init__( - state=state, - sqla_models=sqla_models, - sessions=sessions, - scope_fields=scope_fields, - ) - self._category_item_model = category_item_model - self.relations = self._state.relations - - def list_relations(self, where: Mapping[str, Any] | None = None) -> list[CategoryItem]: - """List category-item relations matching the where clause. - - Args: - where: Optional filter conditions. - - Returns: - List of CategoryItem relations. - """ - with self._sessions.session() as session: - stmt = select(self._category_item_model) - filters = self._build_filters(self._category_item_model, where) - if filters: - stmt = stmt.where(*filters) - rows = session.exec(stmt).all() - - result: list[CategoryItem] = [] - for row in rows: - rel = CategoryItem( - id=row.id, - item_id=row.item_id, - category_id=row.category_id, - created_at=row.created_at, - updated_at=row.updated_at, - **self._scope_kwargs_from(row), - ) - result.append(rel) - # Update cache - if not any(r.id == rel.id for r in self.relations): - self.relations.append(rel) - - return result - - def link_item_category(self, item_id: str, category_id: str, user_data: dict[str, Any]) -> CategoryItem: - """Create a link between an item and a category. - - Args: - item_id: Memory item ID. - category_id: Category ID. - user_data: User scope data. - - Returns: - Created CategoryItem relation. - """ - # Check if relation already exists - where: dict[str, Any] = { - "item_id": item_id, - "category_id": category_id, - **user_data, - } - with self._sessions.session() as session: - stmt = select(self._category_item_model) - filters = self._build_filters(self._category_item_model, where) - if filters: - stmt = stmt.where(*filters) - existing = session.exec(stmt).first() - - if existing: - rel = CategoryItem( - id=existing.id, - item_id=existing.item_id, - category_id=existing.category_id, - created_at=existing.created_at, - updated_at=existing.updated_at, - **self._scope_kwargs_from(existing), - ) - return rel - - # Create new relation - now = self._now() - row = self._category_item_model( - item_id=item_id, - category_id=category_id, - created_at=now, - updated_at=now, - **user_data, - ) - session.add(row) - session.commit() - session.refresh(row) - - rel = CategoryItem( - id=row.id, - item_id=row.item_id, - category_id=row.category_id, - created_at=row.created_at, - updated_at=row.updated_at, - **user_data, - ) - self.relations.append(rel) - return rel - - def unlink_item_category(self, item_id: str, category_id: str) -> None: - """Remove a link between an item and a category. - - Args: - item_id: Memory item ID. - category_id: Category ID. - """ - with self._sessions.session() as session: - stmt = select(self._category_item_model).where( - self._category_item_model.item_id == item_id, - self._category_item_model.category_id == category_id, - ) - row = session.exec(stmt).first() - if row: - session.delete(row) - session.commit() - # Remove from cache - self.relations[:] = [ - r for r in self.relations if not (r.item_id == item_id and r.category_id == category_id) - ] - - def unlink_item(self, item_id: str) -> list[CategoryItem]: - """Remove all relations for a given item (used on item deletion).""" - from sqlmodel import delete - - removed = self.list_relations({"item_id": item_id}) - if not removed: - return [] - with self._sessions.session() as session: - session.exec(delete(self._category_item_model).where(self._category_item_model.item_id == item_id)) - session.commit() - self.relations[:] = [r for r in self.relations if r.item_id != item_id] - return removed - - def clear_relations(self, where: Mapping[str, Any] | None = None) -> list[CategoryItem]: - """Remove all relations matching the scope (used on clear_memory).""" - from sqlmodel import delete - - removed = self.list_relations(where) - if not removed: - return [] - filters = self._build_filters(self._category_item_model, where) - with self._sessions.session() as session: - del_stmt = delete(self._category_item_model) - if filters: - del_stmt = del_stmt.where(*filters) - session.exec(del_stmt) - session.commit() - removed_ids = {rel.id for rel in removed} - self.relations[:] = [r for r in self.relations if r.id not in removed_ids] - return removed - - def get_item_categories(self, item_id: str) -> list[CategoryItem]: - """Get all category relations for a given item. - - Args: - item_id: Memory item ID. - - Returns: - List of CategoryItem relations for the item. - """ - return self.list_relations({"item_id": item_id}) - - def load_existing(self) -> None: - """Load all existing relations from database into cache.""" - self.list_relations() - - -__all__ = ["SQLiteCategoryItemRepo"] diff --git a/src/memu/database/sqlite/repositories/entry_repo.py b/src/memu/database/sqlite/repositories/entry_repo.py new file mode 100644 index 00000000..993e3232 --- /dev/null +++ b/src/memu/database/sqlite/repositories/entry_repo.py @@ -0,0 +1,393 @@ +"""SQLite entry repository implementation.""" + +from __future__ import annotations + +import logging +from collections.abc import Mapping +from typing import Any + +import pendulum +from sqlalchemy import func +from sqlmodel import delete, select + +from memu.database.models import Entry, compute_content_hash +from memu.database.repositories.entry import EntryRepo +from memu.database.sqlite.repositories.base import SQLiteRepoBase +from memu.database.sqlite.schema import SQLiteSQLAModels +from memu.database.sqlite.session import SQLiteSessionManager +from memu.database.state import DatabaseState +from memu.vector import cosine_topk, cosine_topk_salience + +logger = logging.getLogger(__name__) + + +class SQLiteEntryRepo(SQLiteRepoBase, EntryRepo): + """SQLite implementation of the entry repository (the searchable atoms).""" + + def __init__( + self, + *, + state: DatabaseState, + entry_model: type[Any], + sqla_models: SQLiteSQLAModels, + sessions: SQLiteSessionManager, + scope_fields: list[str], + ) -> None: + """Initialize entry repository. + + Args: + state: Shared database state for caching. + entry_model: SQLModel class for entries. + sqla_models: SQLAlchemy model container. + sessions: Session manager for database connections. + scope_fields: List of user scope field names. + """ + super().__init__( + state=state, + sqla_models=sqla_models, + sessions=sessions, + scope_fields=scope_fields, + ) + self._entry_model = entry_model + self.entries = self._state.entries + + def _row_to_entry(self, row: Any) -> Entry: + """Map an ORM row to a backend-agnostic Entry record.""" + return Entry( + id=row.id, + lane=row.lane, + source_id=row.source_id, + source_path=row.source_path, + entry_kind=row.entry_kind, + text=row.text, + embedding=self._normalize_embedding(row.embedding), + happened_at=row.happened_at, + extra=row.extra or {}, + created_at=row.created_at, + updated_at=row.updated_at, + **self._scope_kwargs_from(row), + ) + + def get_entry(self, entry_id: str) -> Entry | None: + """Get an entry by ID.""" + if entry_id in self.entries: + return self.entries[entry_id] + + with self._sessions.session() as session: + stmt = select(self._entry_model).where(self._entry_model.id == entry_id) + row = session.exec(stmt).first() + + if row is None: + return None + + entry = self._row_to_entry(row) + self.entries[row.id] = entry + return entry + + def list_entries( + self, where: Mapping[str, Any] | None = None, *, lane: str | None = None + ) -> dict[str, Entry]: + """List entries matching the where clause and optional lane filter.""" + filters = self._build_filters(self._entry_model, where) + filters.extend(self._lane_filters(self._entry_model, lane)) + with self._sessions.session() as session: + stmt = select(self._entry_model) + if filters: + stmt = stmt.where(*filters) + rows = session.exec(stmt).all() + + result: dict[str, Entry] = {} + for row in rows: + entry = self._row_to_entry(row) + result[row.id] = entry + self.entries[row.id] = entry + + return result + + def list_entries_by_ref_ids( + self, ref_ids: list[str], where: Mapping[str, Any] | None = None + ) -> dict[str, Entry]: + """List entries whose ``extra.ref_id`` is in ``ref_ids``.""" + if not ref_ids: + return {} + + with self._sessions.session() as session: + stmt = select(self._entry_model) + filters = self._build_filters(self._entry_model, where) + ref_id_col = func.json_extract(self._entry_model.extra, "$.ref_id") + filters.append(ref_id_col.isnot(None)) + filters.append(ref_id_col.in_(ref_ids)) + if filters: + stmt = stmt.where(*filters) + rows = session.exec(stmt).all() + + result: dict[str, Entry] = {} + for row in rows: + entry = self._row_to_entry(row) + result[row.id] = entry + self.entries[row.id] = entry + + return result + + def clear_entries( + self, where: Mapping[str, Any] | None = None, *, lane: str | None = None + ) -> dict[str, Entry]: + """Clear entries matching the where clause and optional lane filter.""" + filters = self._build_filters(self._entry_model, where) + filters.extend(self._lane_filters(self._entry_model, lane)) + with self._sessions.session() as session: + stmt = select(self._entry_model) + if filters: + stmt = stmt.where(*filters) + rows = session.exec(stmt).all() + + deleted: dict[str, Entry] = {row.id: self._row_to_entry(row) for row in rows} + + if not deleted: + return {} + + del_stmt = delete(self._entry_model) + if filters: + del_stmt = del_stmt.where(*filters) + session.exec(del_stmt) + session.commit() + + for entry_id in deleted: + self.entries.pop(entry_id, None) + + return deleted + + def create_entry( + self, + *, + lane: str, + source_id: str | None, + entry_kind: str, + text: str, + embedding: list[float], + user_data: dict[str, Any], + source_path: str | None = None, + reinforce: bool = False, + tool_record: dict[str, Any] | None = None, + ) -> Entry: + """Create a new entry (optionally reinforcing a duplicate).""" + if reinforce and entry_kind != "tool": + return self.create_entry_reinforce( + lane=lane, + source_id=source_id, + entry_kind=entry_kind, + text=text, + embedding=embedding, + user_data=user_data, + source_path=source_path, + ) + + extra: dict[str, Any] = {} + if tool_record: + for key in ("when_to_use", "metadata", "tool_calls"): + if tool_record.get(key) is not None: + extra[key] = tool_record[key] + + now = self._now() + row = self._entry_model( + lane=lane, + source_id=source_id, + source_path=source_path, + entry_kind=entry_kind, + text=text, + embedding=self._prepare_embedding(embedding), + extra=extra if extra else {}, + created_at=now, + updated_at=now, + **user_data, + ) + with self._sessions.session() as session: + session.add(row) + session.commit() + session.refresh(row) + + entry = self._row_to_entry(row) + self.entries[row.id] = entry + return entry + + def create_entry_reinforce( + self, + *, + lane: str, + source_id: str | None, + entry_kind: str, + text: str, + embedding: list[float], + user_data: dict[str, Any], + source_path: str | None = None, + ) -> Entry: + """Create or reinforce an entry with content-hash deduplication. + + If an entry with the same content hash exists in the same scope, reinforce + it instead of creating a duplicate. + """ + content_hash = compute_content_hash(text, entry_kind) + + with self._sessions.session() as session: + content_hash_col = func.json_extract(self._entry_model.extra, "$.content_hash") + filters = [content_hash_col == content_hash] + filters.extend(self._build_filters(self._entry_model, user_data)) + + existing = session.exec(select(self._entry_model).where(*filters)).first() + + if existing: + current_extra = existing.extra or {} + current_count = current_extra.get("reinforcement_count", 1) + existing.extra = { + **current_extra, + "reinforcement_count": current_count + 1, + "last_reinforced_at": self._now().isoformat(), + } + existing.updated_at = self._now() + session.add(existing) + session.commit() + session.refresh(existing) + + entry = self._row_to_entry(existing) + self.entries[existing.id] = entry + return entry + + now = self._now() + entry_extra = user_data.pop("extra", {}) if "extra" in user_data else {} + entry_extra.update({ + "content_hash": content_hash, + "reinforcement_count": 1, + "last_reinforced_at": now.isoformat(), + }) + + row = self._entry_model( + lane=lane, + source_id=source_id, + source_path=source_path, + entry_kind=entry_kind, + text=text, + embedding=self._prepare_embedding(embedding), + extra=entry_extra, + created_at=now, + updated_at=now, + **user_data, + ) + session.add(row) + session.commit() + session.refresh(row) + + entry = self._row_to_entry(row) + self.entries[row.id] = entry + return entry + + def update_entry( + self, + *, + entry_id: str, + entry_kind: str | None = None, + text: str | None = None, + embedding: list[float] | None = None, + extra: dict[str, Any] | None = None, + tool_record: dict[str, Any] | None = None, + ) -> Entry: + """Update an existing entry. + + Raises: + KeyError: If the entry is not found. + """ + with self._sessions.session() as session: + stmt = select(self._entry_model).where(self._entry_model.id == entry_id) + row = session.exec(stmt).first() + + if row is None: + msg = f"Entry with id {entry_id} not found" + raise KeyError(msg) + + if entry_kind is not None: + row.entry_kind = entry_kind + if text is not None: + row.text = text + if embedding is not None: + row.embedding = self._prepare_embedding(embedding) + + current_extra = row.extra or {} + if extra is not None: + current_extra = {**current_extra, **extra} + if tool_record is not None: + for key in ("when_to_use", "metadata", "tool_calls"): + if tool_record.get(key) is not None: + current_extra[key] = tool_record[key] + if extra is not None or tool_record is not None: + row.extra = current_extra + + row.updated_at = self._now() + + session.add(row) + session.commit() + session.refresh(row) + + entry = self._row_to_entry(row) + self.entries[row.id] = entry + return entry + + def delete_entry(self, entry_id: str) -> None: + """Delete an entry by id.""" + with self._sessions.session() as session: + stmt = select(self._entry_model).where(self._entry_model.id == entry_id) + row = session.exec(stmt).first() + if row: + session.delete(row) + session.commit() + + self.entries.pop(entry_id, None) + + def vector_search_entries( + self, + query_vec: list[float], + top_k: int, + where: Mapping[str, Any] | None = None, + *, + lane: str | None = None, + ranking: str = "similarity", + recency_decay_days: float = 30.0, + ) -> list[tuple[str, float]]: + """Rank entries by brute-force cosine similarity or salience. + + SQLite has no native vector support, so embeddings are scored in Python. + """ + pool = self.list_entries(where, lane=lane) + + if ranking == "salience": + corpus = [ + ( + e.id, + e.embedding, + (e.extra or {}).get("reinforcement_count", 1), + self._parse_datetime((e.extra or {}).get("last_reinforced_at")), + ) + for e in pool.values() + ] + return cosine_topk_salience(query_vec, corpus, k=top_k, recency_decay_days=recency_decay_days) + + return cosine_topk(query_vec, [(e.id, e.embedding) for e in pool.values()], k=top_k) + + @staticmethod + def _parse_datetime(dt_str: str | None) -> pendulum.DateTime | None: + """Parse an ISO datetime string from the extra dict.""" + if dt_str is None: + return None + try: + parsed = pendulum.parse(dt_str) + except (ValueError, TypeError): + return None + else: + if isinstance(parsed, pendulum.DateTime): + return parsed + return None + + def load_existing(self) -> None: + """Load all existing entries from database into cache.""" + self.list_entries() + + +__all__ = ["SQLiteEntryRepo"] diff --git a/src/memu/database/sqlite/repositories/memory_category_repo.py b/src/memu/database/sqlite/repositories/memory_category_repo.py deleted file mode 100644 index c93f67b5..00000000 --- a/src/memu/database/sqlite/repositories/memory_category_repo.py +++ /dev/null @@ -1,260 +0,0 @@ -"""SQLite memory category repository implementation.""" - -from __future__ import annotations - -import logging -from collections.abc import Mapping -from typing import Any - -from sqlmodel import delete, select - -from memu.database.models import MemoryCategory -from memu.database.repositories.memory_category import MemoryCategoryRepo -from memu.database.sqlite.repositories.base import SQLiteRepoBase -from memu.database.sqlite.schema import SQLiteSQLAModels -from memu.database.sqlite.session import SQLiteSessionManager -from memu.database.state import DatabaseState - -logger = logging.getLogger(__name__) - - -class SQLiteMemoryCategoryRepo(SQLiteRepoBase, MemoryCategoryRepo): - """SQLite implementation of memory category repository.""" - - def __init__( - self, - *, - state: DatabaseState, - memory_category_model: type[Any], - sqla_models: SQLiteSQLAModels, - sessions: SQLiteSessionManager, - scope_fields: list[str], - ) -> None: - """Initialize memory category repository. - - Args: - state: Shared database state for caching. - memory_category_model: SQLModel class for memory categories. - sqla_models: SQLAlchemy model container. - sessions: Session manager for database connections. - scope_fields: List of user scope field names. - """ - super().__init__( - state=state, - sqla_models=sqla_models, - sessions=sessions, - scope_fields=scope_fields, - ) - self._memory_category_model = memory_category_model - self.categories = self._state.categories - - def list_categories(self, where: Mapping[str, Any] | None = None) -> dict[str, MemoryCategory]: - """List categories matching the where clause. - - Args: - where: Optional filter conditions. - - Returns: - Dictionary of category ID to MemoryCategory mapping. - """ - with self._sessions.session() as session: - stmt = select(self._memory_category_model) - filters = self._build_filters(self._memory_category_model, where) - if filters: - stmt = stmt.where(*filters) - rows = session.exec(stmt).all() - - result: dict[str, MemoryCategory] = {} - for row in rows: - cat = MemoryCategory( - id=row.id, - name=row.name, - description=row.description, - embedding=self._normalize_embedding(row.embedding), - summary=row.summary, - created_at=row.created_at, - updated_at=row.updated_at, - **self._scope_kwargs_from(row), - ) - result[row.id] = cat - self.categories[row.id] = cat - - return result - - def clear_categories(self, where: Mapping[str, Any] | None = None) -> dict[str, MemoryCategory]: - """Clear categories matching the where clause. - - Args: - where: Optional filter conditions. - - Returns: - Dictionary of deleted category ID to MemoryCategory mapping. - """ - filters = self._build_filters(self._memory_category_model, where) - with self._sessions.session() as session: - # First get the objects to delete - stmt = select(self._memory_category_model) - if filters: - stmt = stmt.where(*filters) - rows = session.exec(stmt).all() - - deleted: dict[str, MemoryCategory] = {} - for row in rows: - cat = MemoryCategory( - id=row.id, - name=row.name, - description=row.description, - embedding=self._normalize_embedding(row.embedding), - summary=row.summary, - created_at=row.created_at, - updated_at=row.updated_at, - **self._scope_kwargs_from(row), - ) - deleted[row.id] = cat - - if not deleted: - return {} - - # Delete from database - del_stmt = delete(self._memory_category_model) - if filters: - del_stmt = del_stmt.where(*filters) - session.exec(del_stmt) - session.commit() - - # Clean up cache - for cat_id in deleted: - self.categories.pop(cat_id, None) - - return deleted - - def get_or_create_category( - self, *, name: str, description: str, embedding: list[float], user_data: dict[str, Any] - ) -> MemoryCategory: - """Get existing category by name or create a new one. - - Args: - name: Category name. - description: Category description. - embedding: Embedding vector. - user_data: User scope data. - - Returns: - Existing or newly created MemoryCategory. - """ - # Check for existing category with same name and scope - where: dict[str, Any] = {"name": name, **user_data} - with self._sessions.session() as session: - stmt = select(self._memory_category_model) - filters = self._build_filters(self._memory_category_model, where) - if filters: - stmt = stmt.where(*filters) - existing = session.exec(stmt).first() - - if existing: - cat = MemoryCategory( - id=existing.id, - name=existing.name, - description=existing.description, - embedding=self._normalize_embedding(existing.embedding), - summary=existing.summary, - created_at=existing.created_at, - updated_at=existing.updated_at, - **self._scope_kwargs_from(existing), - ) - self.categories[existing.id] = cat - return cat - - # Create new category - now = self._now() - row = self._memory_category_model( - name=name, - description=description, - embedding=self._prepare_embedding(embedding), - summary=None, - created_at=now, - updated_at=now, - **user_data, - ) - session.add(row) - session.commit() - session.refresh(row) - - cat = MemoryCategory( - id=row.id, - name=row.name, - description=row.description, - embedding=embedding, - summary=None, - created_at=row.created_at, - updated_at=row.updated_at, - **user_data, - ) - self.categories[row.id] = cat - return cat - - def update_category( - self, - *, - category_id: str, - name: str | None = None, - description: str | None = None, - embedding: list[float] | None = None, - summary: str | None = None, - ) -> MemoryCategory: - """Update an existing category. - - Args: - category_id: ID of category to update. - name: New name (optional). - description: New description (optional). - embedding: New embedding vector (optional). - summary: New summary text (optional). - - Returns: - Updated MemoryCategory object. - - Raises: - KeyError: If category not found. - """ - with self._sessions.session() as session: - stmt = select(self._memory_category_model).where(self._memory_category_model.id == category_id) - row = session.exec(stmt).first() - - if row is None: - msg = f"Category with id {category_id} not found" - raise KeyError(msg) - - if name is not None: - row.name = name - if description is not None: - row.description = description - if embedding is not None: - row.embedding = self._prepare_embedding(embedding) - if summary is not None: - row.summary = summary - row.updated_at = self._now() - - session.add(row) - session.commit() - session.refresh(row) - - cat = MemoryCategory( - id=row.id, - name=row.name, - description=row.description, - embedding=self._normalize_embedding(row.embedding), - summary=row.summary, - created_at=row.created_at, - updated_at=row.updated_at, - **self._scope_kwargs_from(row), - ) - self.categories[row.id] = cat - return cat - - def load_existing(self) -> None: - """Load all existing categories from database into cache.""" - self.list_categories() - - -__all__ = ["SQLiteMemoryCategoryRepo"] diff --git a/src/memu/database/sqlite/repositories/memory_item_repo.py b/src/memu/database/sqlite/repositories/memory_item_repo.py deleted file mode 100644 index 0f070d57..00000000 --- a/src/memu/database/sqlite/repositories/memory_item_repo.py +++ /dev/null @@ -1,544 +0,0 @@ -"""SQLite memory item repository implementation.""" - -from __future__ import annotations - -import logging -from collections.abc import Mapping -from typing import Any - -import pendulum -from sqlmodel import delete, select - -from memu.database.models import MemoryItem, MemoryType, compute_content_hash -from memu.database.repositories.memory_item import MemoryItemRepo -from memu.database.sqlite.repositories.base import SQLiteRepoBase -from memu.database.sqlite.schema import SQLiteSQLAModels -from memu.database.sqlite.session import SQLiteSessionManager -from memu.database.state import DatabaseState -from memu.vector import cosine_topk, cosine_topk_salience - -logger = logging.getLogger(__name__) - - -class SQLiteMemoryItemRepo(SQLiteRepoBase, MemoryItemRepo): - """SQLite implementation of memory item repository.""" - - def __init__( - self, - *, - state: DatabaseState, - memory_item_model: type[Any], - sqla_models: SQLiteSQLAModels, - sessions: SQLiteSessionManager, - scope_fields: list[str], - ) -> None: - """Initialize memory item repository. - - Args: - state: Shared database state for caching. - memory_item_model: SQLModel class for memory items. - sqla_models: SQLAlchemy model container. - sessions: Session manager for database connections. - scope_fields: List of user scope field names. - """ - super().__init__( - state=state, - sqla_models=sqla_models, - sessions=sessions, - scope_fields=scope_fields, - ) - self._memory_item_model = memory_item_model - self.items = self._state.items - - def get_item(self, item_id: str) -> MemoryItem | None: - """Get a memory item by ID. - - Args: - item_id: The item ID to look up. - - Returns: - MemoryItem if found, None otherwise. - """ - # Check cache first - if item_id in self.items: - return self.items[item_id] - - with self._sessions.session() as session: - stmt = select(self._memory_item_model).where(self._memory_item_model.id == item_id) - row = session.exec(stmt).first() - - if row is None: - return None - - item = MemoryItem( - id=row.id, - resource_id=row.resource_id, - memory_type=row.memory_type, - summary=row.summary, - embedding=self._normalize_embedding(row.embedding), - created_at=row.created_at, - updated_at=row.updated_at, - extra=row.extra, - **self._scope_kwargs_from(row), - ) - self.items[row.id] = item - return item - - def list_items(self, where: Mapping[str, Any] | None = None) -> dict[str, MemoryItem]: - """List memory items matching the where clause. - - Args: - where: Optional filter conditions. - - Returns: - Dictionary of item ID to MemoryItem mapping. - """ - with self._sessions.session() as session: - stmt = select(self._memory_item_model) - filters = self._build_filters(self._memory_item_model, where) - if filters: - stmt = stmt.where(*filters) - rows = session.exec(stmt).all() - - result: dict[str, MemoryItem] = {} - for row in rows: - item = MemoryItem( - id=row.id, - resource_id=row.resource_id, - memory_type=row.memory_type, - summary=row.summary, - embedding=self._normalize_embedding(row.embedding), - created_at=row.created_at, - updated_at=row.updated_at, - extra=row.extra, - **self._scope_kwargs_from(row), - ) - result[row.id] = item - self.items[row.id] = item - - return result - - def list_items_by_ref_ids( - self, ref_ids: list[str], where: Mapping[str, Any] | None = None - ) -> dict[str, MemoryItem]: - """List items by their ref_id in the extra column. - - Args: - ref_ids: List of ref_ids to query. - where: Additional filter conditions. - - Returns: - Dict mapping item_id -> MemoryItem for items whose extra.ref_id is in ref_ids. - """ - if not ref_ids: - return {} - - from sqlalchemy import func - - with self._sessions.session() as session: - stmt = select(self._memory_item_model) - filters = self._build_filters(self._memory_item_model, where) - # Add filter for json_extract(extra, '$.ref_id') IN ref_ids (only rows with ref_id key) - ref_id_col = func.json_extract(self._memory_item_model.extra, "$.ref_id") - filters.append(ref_id_col.isnot(None)) - filters.append(ref_id_col.in_(ref_ids)) - if filters: - stmt = stmt.where(*filters) - rows = session.exec(stmt).all() - - result: dict[str, MemoryItem] = {} - for row in rows: - item = MemoryItem( - id=row.id, - resource_id=row.resource_id, - memory_type=row.memory_type, - summary=row.summary, - embedding=self._normalize_embedding(row.embedding), - created_at=row.created_at, - updated_at=row.updated_at, - extra=row.extra, - **self._scope_kwargs_from(row), - ) - result[row.id] = item - self.items[row.id] = item - - return result - - def clear_items(self, where: Mapping[str, Any] | None = None) -> dict[str, MemoryItem]: - """Clear items matching the where clause. - - Args: - where: Optional filter conditions. - - Returns: - Dictionary of deleted item ID to MemoryItem mapping. - """ - filters = self._build_filters(self._memory_item_model, where) - with self._sessions.session() as session: - # First get the objects to delete - stmt = select(self._memory_item_model) - if filters: - stmt = stmt.where(*filters) - rows = session.exec(stmt).all() - - deleted: dict[str, MemoryItem] = {} - for row in rows: - item = MemoryItem( - id=row.id, - resource_id=row.resource_id, - memory_type=row.memory_type, - summary=row.summary, - embedding=self._normalize_embedding(row.embedding), - created_at=row.created_at, - updated_at=row.updated_at, - extra=row.extra, - **self._scope_kwargs_from(row), - ) - deleted[row.id] = item - - if not deleted: - return {} - - # Delete from database - del_stmt = delete(self._memory_item_model) - if filters: - del_stmt = del_stmt.where(*filters) - session.exec(del_stmt) - session.commit() - - # Clean up cache - for item_id in deleted: - self.items.pop(item_id, None) - - return deleted - - def create_item( - self, - *, - resource_id: str, - memory_type: MemoryType, - summary: str, - embedding: list[float], - user_data: dict[str, Any], - reinforce: bool = False, - tool_record: dict[str, Any] | None = None, - ) -> MemoryItem: - """Create a new memory item. - - Args: - resource_id: Associated resource ID. - memory_type: Type of memory. - summary: Memory summary text. - embedding: Embedding vector. - user_data: User scope data. - reinforce: If True, reinforce existing item instead of creating duplicate. - tool_record: Tool-related fields (when_to_use, metadata, tool_calls) to store in extra. - - Returns: - Created MemoryItem object. - """ - if reinforce and memory_type != "tool": - return self.create_item_reinforce( - resource_id=resource_id, - memory_type=memory_type, - summary=summary, - embedding=embedding, - user_data=user_data, - ) - - # Build extra dict with tool_record fields at top level - extra: dict[str, Any] = {} - if tool_record: - if tool_record.get("when_to_use") is not None: - extra["when_to_use"] = tool_record["when_to_use"] - if tool_record.get("metadata") is not None: - extra["metadata"] = tool_record["metadata"] - if tool_record.get("tool_calls") is not None: - extra["tool_calls"] = tool_record["tool_calls"] - - now = self._now() - row = self._memory_item_model( - resource_id=resource_id, - memory_type=memory_type, - summary=summary, - embedding=self._prepare_embedding(embedding), - extra=extra if extra else {}, - created_at=now, - updated_at=now, - **user_data, - ) - with self._sessions.session() as session: - session.add(row) - session.commit() - session.refresh(row) - - item = MemoryItem( - id=row.id, - resource_id=row.resource_id, - memory_type=row.memory_type, - summary=row.summary, - embedding=embedding, - extra=row.extra, - created_at=row.created_at, - updated_at=row.updated_at, - **user_data, - ) - self.items[row.id] = item - return item - - def create_item_reinforce( - self, - *, - resource_id: str, - memory_type: MemoryType, - summary: str, - embedding: list[float], - user_data: dict[str, Any], - ) -> MemoryItem: - """Create or reinforce a memory item with deduplication. - - If an item with the same content hash exists in the same scope, - reinforce it instead of creating a duplicate. - - Args: - resource_id: Associated resource ID. - memory_type: Type of memory. - summary: Memory summary text. - embedding: Embedding vector. - user_data: User scope data. - - Returns: - Created or reinforced MemoryItem object. - """ - from sqlalchemy import func - - content_hash = compute_content_hash(summary, memory_type) - - with self._sessions.session() as session: - # Check for existing item with same hash in same scope (deduplication) - # Use json_extract(extra, '$.content_hash') for query - content_hash_col = func.json_extract(self._memory_item_model.extra, "$.content_hash") - filters = [content_hash_col == content_hash] - filters.extend(self._build_filters(self._memory_item_model, user_data)) - - existing = session.exec(select(self._memory_item_model).where(*filters)).first() - - if existing: - # Reinforce existing memory instead of creating duplicate - current_extra = existing.extra or {} - current_count = current_extra.get("reinforcement_count", 1) - existing.extra = { - **current_extra, - "reinforcement_count": current_count + 1, - "last_reinforced_at": self._now().isoformat(), - } - existing.updated_at = self._now() - session.add(existing) - session.commit() - session.refresh(existing) - - item = MemoryItem( - id=existing.id, - resource_id=existing.resource_id, - memory_type=existing.memory_type, - summary=existing.summary, - embedding=self._normalize_embedding(existing.embedding), - created_at=existing.created_at, - updated_at=existing.updated_at, - extra=existing.extra, - **self._scope_kwargs_from(existing), - ) - self.items[existing.id] = item - return item - - # Create new item with salience tracking in extra - now = self._now() - item_extra = user_data.pop("extra", {}) if "extra" in user_data else {} - item_extra.update({ - "content_hash": content_hash, - "reinforcement_count": 1, - "last_reinforced_at": now.isoformat(), - }) - - row = self._memory_item_model( - resource_id=resource_id, - memory_type=memory_type, - summary=summary, - embedding=self._prepare_embedding(embedding), - extra=item_extra, - created_at=now, - updated_at=now, - **user_data, - ) - - session.add(row) - session.commit() - session.refresh(row) - - item = MemoryItem( - id=row.id, - resource_id=row.resource_id, - memory_type=row.memory_type, - summary=row.summary, - embedding=embedding, - created_at=row.created_at, - updated_at=row.updated_at, - extra=row.extra, - **self._scope_kwargs_from(row), - ) - self.items[row.id] = item - return item - - def update_item( - self, - *, - item_id: str, - memory_type: MemoryType | None = None, - summary: str | None = None, - embedding: list[float] | None = None, - extra: dict[str, Any] | None = None, - tool_record: dict[str, Any] | None = None, - ) -> MemoryItem: - """Update an existing memory item. - - Args: - item_id: ID of item to update. - memory_type: New memory type (optional). - summary: New summary text (optional). - embedding: New embedding vector (optional). - extra: Extra data to merge into existing extra dict (optional). - tool_record: Tool-related fields (when_to_use, metadata, tool_calls) to merge into extra. - - Returns: - Updated MemoryItem object. - - Raises: - KeyError: If item not found. - """ - with self._sessions.session() as session: - stmt = select(self._memory_item_model).where(self._memory_item_model.id == item_id) - row = session.exec(stmt).first() - - if row is None: - msg = f"Item with id {item_id} not found" - raise KeyError(msg) - - if memory_type is not None: - row.memory_type = memory_type - if summary is not None: - row.summary = summary - if embedding is not None: - row.embedding = self._prepare_embedding(embedding) - - # Merge extra and tool_record into existing extra dict - current_extra = row.extra or {} - if extra is not None: - current_extra = {**current_extra, **extra} - if tool_record is not None: - # Merge tool_record fields at top level - for key in ("when_to_use", "metadata", "tool_calls"): - if tool_record.get(key) is not None: - current_extra[key] = tool_record[key] - if extra is not None or tool_record is not None: - row.extra = current_extra - - row.updated_at = self._now() - - session.add(row) - session.commit() - session.refresh(row) - - item = MemoryItem( - id=row.id, - resource_id=row.resource_id, - memory_type=row.memory_type, - summary=row.summary, - embedding=self._normalize_embedding(row.embedding), - extra=row.extra, - created_at=row.created_at, - updated_at=row.updated_at, - **self._scope_kwargs_from(row), - ) - self.items[row.id] = item - return item - - def delete_item(self, item_id: str) -> None: - """Delete a memory item. - - Args: - item_id: ID of item to delete. - """ - with self._sessions.session() as session: - stmt = select(self._memory_item_model).where(self._memory_item_model.id == item_id) - row = session.exec(stmt).first() - if row: - session.delete(row) - session.commit() - - if item_id in self.items: - del self.items[item_id] - - def vector_search_items( - self, - query_vec: list[float], - top_k: int, - where: Mapping[str, Any] | None = None, - *, - ranking: str = "similarity", - recency_decay_days: float = 30.0, - ) -> list[tuple[str, float]]: - """Perform vector similarity search on memory items. - - Uses brute-force cosine similarity since SQLite doesn't have native vector support. - - Args: - query_vec: Query embedding vector. - top_k: Maximum number of results to return. - where: Optional filter conditions. - ranking: Ranking strategy - "similarity" (default) or "salience". - recency_decay_days: Half-life for recency decay in salience ranking. - - Returns: - List of (item_id, similarity_score) tuples. - """ - # Load items from database with filters - pool = self.list_items(where) - - if ranking == "salience": - # Salience-aware ranking: similarity x reinforcement x recency - # Read values from extra dict - corpus = [ - ( - i.id, - i.embedding, - (i.extra or {}).get("reinforcement_count", 1), - self._parse_datetime((i.extra or {}).get("last_reinforced_at")), - ) - for i in pool.values() - ] - return cosine_topk_salience(query_vec, corpus, k=top_k, recency_decay_days=recency_decay_days) - - # Default: pure cosine similarity (backward compatible) - hits = cosine_topk(query_vec, [(i.id, i.embedding) for i in pool.values()], k=top_k) - return hits - - @staticmethod - def _parse_datetime(dt_str: str | None) -> pendulum.DateTime | None: - """Parse ISO datetime string from extra dict.""" - if dt_str is None: - return None - try: - parsed = pendulum.parse(dt_str) - except (ValueError, TypeError): - return None - else: - if isinstance(parsed, pendulum.DateTime): - return parsed - return None - - def load_existing(self) -> None: - """Load all existing items from database into cache.""" - self.list_items() - - -__all__ = ["SQLiteMemoryItemRepo"] diff --git a/src/memu/database/sqlite/repositories/resource_entry_repo.py b/src/memu/database/sqlite/repositories/resource_entry_repo.py new file mode 100644 index 00000000..60062ca4 --- /dev/null +++ b/src/memu/database/sqlite/repositories/resource_entry_repo.py @@ -0,0 +1,164 @@ +"""SQLite entry <-> resource membership-edge repository implementation.""" + +from __future__ import annotations + +import logging +from collections.abc import Mapping +from typing import Any + +from sqlmodel import delete, select + +from memu.database.models import ResourceEntry +from memu.database.repositories.resource_entry import ResourceEntryRepo +from memu.database.sqlite.repositories.base import SQLiteRepoBase +from memu.database.sqlite.schema import SQLiteSQLAModels +from memu.database.sqlite.session import SQLiteSessionManager +from memu.database.state import DatabaseState + +logger = logging.getLogger(__name__) + + +class SQLiteResourceEntryRepo(SQLiteRepoBase, ResourceEntryRepo): + """SQLite implementation of the entry <-> coarse-resource edge repository.""" + + def __init__( + self, + *, + state: DatabaseState, + resource_entry_model: type[Any], + sqla_models: SQLiteSQLAModels, + sessions: SQLiteSessionManager, + scope_fields: list[str], + ) -> None: + """Initialize resource-entry repository. + + Args: + state: Shared database state for caching. + resource_entry_model: SQLModel class for membership edges. + sqla_models: SQLAlchemy model container. + sessions: Session manager for database connections. + scope_fields: List of user scope field names. + """ + super().__init__( + state=state, + sqla_models=sqla_models, + sessions=sessions, + scope_fields=scope_fields, + ) + self._resource_entry_model = resource_entry_model + self.relations = self._state.relations + + def _row_to_relation(self, row: Any) -> ResourceEntry: + """Map an ORM row to a backend-agnostic ResourceEntry record.""" + return ResourceEntry( + id=row.id, + entry_id=row.entry_id, + resource_id=row.resource_id, + created_at=row.created_at, + updated_at=row.updated_at, + **self._scope_kwargs_from(row), + ) + + def list_relations(self, where: Mapping[str, Any] | None = None) -> list[ResourceEntry]: + """List membership edges matching the where clause.""" + with self._sessions.session() as session: + stmt = select(self._resource_entry_model) + filters = self._build_filters(self._resource_entry_model, where) + if filters: + stmt = stmt.where(*filters) + rows = session.exec(stmt).all() + + result: list[ResourceEntry] = [] + for row in rows: + rel = self._row_to_relation(row) + result.append(rel) + if not any(r.id == rel.id for r in self.relations): + self.relations.append(rel) + + return result + + def link_entry_resource(self, entry_id: str, resource_id: str, user_data: dict[str, Any]) -> ResourceEntry: + """Create (or return existing) edge between an entry and a coarse resource.""" + where: dict[str, Any] = { + "entry_id": entry_id, + "resource_id": resource_id, + **user_data, + } + with self._sessions.session() as session: + stmt = select(self._resource_entry_model) + filters = self._build_filters(self._resource_entry_model, where) + if filters: + stmt = stmt.where(*filters) + existing = session.exec(stmt).first() + + if existing: + return self._row_to_relation(existing) + + now = self._now() + row = self._resource_entry_model( + entry_id=entry_id, + resource_id=resource_id, + created_at=now, + updated_at=now, + **user_data, + ) + session.add(row) + session.commit() + session.refresh(row) + + rel = self._row_to_relation(row) + self.relations.append(rel) + return rel + + def unlink_entry_resource(self, entry_id: str, resource_id: str) -> None: + """Remove a single edge between an entry and a coarse resource.""" + with self._sessions.session() as session: + stmt = select(self._resource_entry_model).where( + self._resource_entry_model.entry_id == entry_id, + self._resource_entry_model.resource_id == resource_id, + ) + row = session.exec(stmt).first() + if row: + session.delete(row) + session.commit() + self.relations[:] = [ + r for r in self.relations if not (r.entry_id == entry_id and r.resource_id == resource_id) + ] + + def unlink_entry(self, entry_id: str) -> list[ResourceEntry]: + """Remove all edges for a given entry (used on entry deletion).""" + removed = self.list_relations({"entry_id": entry_id}) + if not removed: + return [] + with self._sessions.session() as session: + session.exec(delete(self._resource_entry_model).where(self._resource_entry_model.entry_id == entry_id)) + session.commit() + self.relations[:] = [r for r in self.relations if r.entry_id != entry_id] + return removed + + def clear_relations(self, where: Mapping[str, Any] | None = None) -> list[ResourceEntry]: + """Remove all edges matching the scope (used on clear).""" + removed = self.list_relations(where) + if not removed: + return [] + filters = self._build_filters(self._resource_entry_model, where) + with self._sessions.session() as session: + del_stmt = delete(self._resource_entry_model) + if filters: + del_stmt = del_stmt.where(*filters) + session.exec(del_stmt) + session.commit() + removed_ids = {rel.id for rel in removed} + self.relations[:] = [r for r in self.relations if r.id not in removed_ids] + return removed + + def get_entry_resources(self, entry_id: str) -> list[ResourceEntry]: + """Get all coarse-resource edges for a given entry.""" + return self.list_relations({"entry_id": entry_id}) + + def load_existing(self) -> None: + """Load all existing edges from database into cache.""" + self.list_relations() + + +__all__ = ["SQLiteResourceEntryRepo"] diff --git a/src/memu/database/sqlite/repositories/resource_repo.py b/src/memu/database/sqlite/repositories/resource_repo.py index 9d663c98..2041d605 100644 --- a/src/memu/database/sqlite/repositories/resource_repo.py +++ b/src/memu/database/sqlite/repositories/resource_repo.py @@ -20,7 +20,11 @@ class SQLiteResourceRepo(SQLiteRepoBase, ResourceRepo): - """SQLite implementation of resource repository.""" + """SQLite implementation of the resource repository. + + A single physical table holds both raw inputs (``lane="source"``) and the + generated lane docs (``lane`` in {index, memory, skill}). + """ def __init__( self, @@ -49,89 +53,85 @@ def __init__( self._resource_model = resource_model self.resources = self._state.resources - def list_resources(self, where: Mapping[str, Any] | None = None) -> dict[str, Resource]: - """List resources matching the where clause. + def _row_to_resource(self, row: Any) -> Resource: + """Map an ORM row to a backend-agnostic Resource record.""" + return Resource( + id=row.id, + lane=row.lane, + modality=row.modality, + url=row.url, + local_path=row.local_path, + slug=row.slug, + title=row.title, + description=row.description, + content=row.content, + summary=row.summary, + embedding=self._normalize_embedding(row.embedding), + resource_refs=row.resource_refs or [], + created_at=row.created_at, + updated_at=row.updated_at, + **self._scope_kwargs_from(row), + ) + + def get_resource(self, resource_id: str) -> Resource | None: + """Get a resource by ID.""" + if resource_id in self.resources: + return self.resources[resource_id] - Args: - where: Optional filter conditions. + with self._sessions.session() as session: + stmt = select(self._resource_model).where(self._resource_model.id == resource_id) + row = session.exec(stmt).first() - Returns: - Dictionary of resource ID to Resource mapping. - """ - # Prefer cached data if available and no filter - if not where and self.resources: - return dict(self.resources) + if row is None: + return None + + res = self._row_to_resource(row) + self.resources[row.id] = res + return res + def list_resources(self, where: Mapping[str, Any] | None = None, *, lane: str | None = None) -> dict[str, Resource]: + """List resources matching the where clause and optional lane filter.""" + filters = self._build_filters(self._resource_model, where) + filters.extend(self._lane_filters(self._resource_model, lane)) with self._sessions.session() as session: stmt = select(self._resource_model) - filters = self._build_filters(self._resource_model, where) if filters: stmt = stmt.where(*filters) rows = session.exec(stmt).all() result: dict[str, Resource] = {} for row in rows: - res = Resource( - id=row.id, - url=row.url, - modality=row.modality, - local_path=row.local_path, - caption=row.caption, - embedding=self._normalize_embedding(row.embedding), - created_at=row.created_at, - updated_at=row.updated_at, - **self._scope_kwargs_from(row), - ) + res = self._row_to_resource(row) result[row.id] = res self.resources[row.id] = res return result - def clear_resources(self, where: Mapping[str, Any] | None = None) -> dict[str, Resource]: - """Clear resources matching the where clause. - - Args: - where: Optional filter conditions. - - Returns: - Dictionary of deleted resource ID to Resource mapping. - """ + def clear_resources( + self, where: Mapping[str, Any] | None = None, *, lane: str | None = None + ) -> dict[str, Resource]: + """Clear resources matching the where clause and optional lane filter.""" filters = self._build_filters(self._resource_model, where) + filters.extend(self._lane_filters(self._resource_model, lane)) with self._sessions.session() as session: - # First get the objects to delete stmt = select(self._resource_model) if filters: stmt = stmt.where(*filters) rows = session.exec(stmt).all() - deleted: dict[str, Resource] = {} - for row in rows: - res = Resource( - id=row.id, - url=row.url, - modality=row.modality, - local_path=row.local_path, - caption=row.caption, - embedding=self._normalize_embedding(row.embedding), - created_at=row.created_at, - updated_at=row.updated_at, - **self._scope_kwargs_from(row), - ) - deleted[row.id] = res + deleted: dict[str, Resource] = {row.id: self._row_to_resource(row) for row in rows} if not deleted: return {} - # Delete from database del_stmt = delete(self._resource_model) if filters: del_stmt = del_stmt.where(*filters) session.exec(del_stmt) session.commit() - # Clean up cache - for res_id in deleted: - self.resources.pop(res_id, None) + for res_id in deleted: + self.resources.pop(res_id, None) return deleted @@ -145,33 +145,33 @@ def delete_resource(self, resource_id: str) -> None: def create_resource( self, *, - url: str, modality: str, - local_path: str, - caption: str | None, - embedding: list[float] | None, user_data: dict[str, Any], + lane: str = "source", + url: str | None = None, + local_path: str | None = None, + slug: str | None = None, + title: str | None = None, + description: str | None = None, + content: str | None = None, + summary: str | None = None, + embedding: list[float] | None = None, + resource_refs: list[dict[str, Any]] | None = None, ) -> Resource: - """Create a new resource record. - - Args: - url: Resource URL. - modality: Resource modality type. - local_path: Local file path. - caption: Optional caption text. - embedding: Optional embedding vector. - user_data: User scope data. - - Returns: - Created Resource object. - """ + """Create a new resource record (raw input or generated lane doc).""" now = self._now() row = self._resource_model( - url=url, + lane=lane, modality=modality, + url=url, local_path=local_path, - caption=caption, + slug=slug, + title=title, + description=description, + content=content, + summary=summary, embedding=self._prepare_embedding(embedding), + resource_refs=resource_refs or [], created_at=now, updated_at=now, **user_data, @@ -181,17 +181,97 @@ def create_resource( session.commit() session.refresh(row) - res = Resource( - id=row.id, - url=row.url, - modality=row.modality, - local_path=row.local_path, - caption=row.caption, + res = self._row_to_resource(row) + self.resources[row.id] = res + return res + + def get_or_create_doc( + self, + *, + lane: str, + title: str, + description: str, + embedding: list[float], + user_data: dict[str, Any], + slug: str | None = None, + ) -> Resource: + """Get an existing generated doc by ``(lane, title, scope)`` or create one.""" + where: dict[str, Any] = {"lane": lane, "title": title, **user_data} + filters = self._build_filters(self._resource_model, where) + with self._sessions.session() as session: + stmt = select(self._resource_model) + if filters: + stmt = stmt.where(*filters) + existing = session.exec(stmt).first() + + if existing is not None: + changed = False + now = self._now() + if self._normalize_embedding(existing.embedding) is None: + existing.embedding = self._prepare_embedding(embedding) + existing.updated_at = now + changed = True + if not existing.description: + existing.description = description + existing.updated_at = now + changed = True + if changed: + session.add(existing) + session.commit() + session.refresh(existing) + res = self._row_to_resource(existing) + self.resources[existing.id] = res + return res + + return self.create_resource( + modality="markdown", + lane=lane, + title=title, + slug=slug, + description=description, embedding=embedding, - created_at=row.created_at, - updated_at=row.updated_at, - **user_data, + user_data=user_data, ) + + def update_resource( + self, + *, + resource_id: str, + title: str | None = None, + description: str | None = None, + content: str | None = None, + summary: str | None = None, + embedding: list[float] | None = None, + resource_refs: list[dict[str, Any]] | None = None, + ) -> Resource: + """Update mutable fields of an existing resource.""" + with self._sessions.session() as session: + stmt = select(self._resource_model).where(self._resource_model.id == resource_id) + row = session.exec(stmt).first() + + if row is None: + msg = f"Resource with id {resource_id} not found" + raise KeyError(msg) + + if title is not None: + row.title = title + if description is not None: + row.description = description + if content is not None: + row.content = content + if summary is not None: + row.summary = summary + if embedding is not None: + row.embedding = self._prepare_embedding(embedding) + if resource_refs is not None: + row.resource_refs = resource_refs + row.updated_at = self._now() + + session.add(row) + session.commit() + session.refresh(row) + + res = self._row_to_resource(row) self.resources[row.id] = res return res @@ -200,13 +280,14 @@ def vector_search_resources( query_vec: list[float], top_k: int, where: Mapping[str, Any] | None = None, + *, + lane: str | None = None, ) -> list[tuple[str, float]]: - """Rank resources by brute-force cosine similarity. + """Rank resources by brute-force cosine similarity of stored embeddings. - SQLite has no native vector support, so this mirrors the item repo and - scores stored embeddings in Python. + SQLite has no native vector support, so embeddings are scored in Python. """ - pool = self.list_resources(where) + pool = self.list_resources(where, lane=lane) corpus = [(rid, res.embedding) for rid, res in pool.items() if res.embedding] return cosine_topk(query_vec, corpus, k=top_k) diff --git a/src/memu/database/sqlite/schema.py b/src/memu/database/sqlite/schema.py index d7d57914..5cd00128 100644 --- a/src/memu/database/sqlite/schema.py +++ b/src/memu/database/sqlite/schema.py @@ -10,9 +10,8 @@ from sqlmodel import SQLModel from memu.database.sqlite.models import ( - SQLiteCategoryItemModel, - SQLiteMemoryCategoryModel, - SQLiteMemoryItemModel, + SQLiteEntryModel, + SQLiteResourceEntryModel, SQLiteResourceModel, build_sqlite_table_model, ) @@ -24,9 +23,8 @@ class SQLiteSQLAModels: Base: type[Any] Resource: type[Any] - MemoryCategory: type[Any] - MemoryItem: type[Any] - CategoryItem: type[Any] + Entry: type[Any] + ResourceEntry: type[Any] _MODEL_CACHE: dict[type[Any], SQLiteSQLAModels] = {} @@ -57,23 +55,18 @@ def get_sqlite_sqlalchemy_models(*, scope_model: type[BaseModel] | None = None) tablename="memu_resources", metadata=metadata_obj, ) - memory_category_model = build_sqlite_table_model( + entry_model = build_sqlite_table_model( scope, - SQLiteMemoryCategoryModel, - tablename="memu_memory_categories", + SQLiteEntryModel, + tablename="memu_entries", metadata=metadata_obj, ) - memory_item_model = build_sqlite_table_model( + resource_entry_model = build_sqlite_table_model( scope, - SQLiteMemoryItemModel, - tablename="memu_memory_items", - metadata=metadata_obj, - ) - category_item_model = build_sqlite_table_model( - scope, - SQLiteCategoryItemModel, - tablename="memu_category_items", + SQLiteResourceEntryModel, + tablename="memu_resource_entries", metadata=metadata_obj, + unique_with_scope=["entry_id", "resource_id"], ) class SQLiteBase(SQLModel): @@ -83,9 +76,8 @@ class SQLiteBase(SQLModel): models = SQLiteSQLAModels( Base=SQLiteBase, Resource=resource_model, - MemoryCategory=memory_category_model, - MemoryItem=memory_item_model, - CategoryItem=category_item_model, + Entry=entry_model, + ResourceEntry=resource_entry_model, ) _MODEL_CACHE[cache_key] = models return models diff --git a/src/memu/database/sqlite/sqlite.py b/src/memu/database/sqlite/sqlite.py index 2083dd99..adbe8580 100644 --- a/src/memu/database/sqlite/sqlite.py +++ b/src/memu/database/sqlite/sqlite.py @@ -9,11 +9,10 @@ from sqlmodel import SQLModel from memu.database.interfaces import Database -from memu.database.models import CategoryItem, MemoryCategory, MemoryItem, Resource -from memu.database.repositories import CategoryItemRepo, MemoryCategoryRepo, MemoryItemRepo, ResourceRepo -from memu.database.sqlite.repositories.category_item_repo import SQLiteCategoryItemRepo -from memu.database.sqlite.repositories.memory_category_repo import SQLiteMemoryCategoryRepo -from memu.database.sqlite.repositories.memory_item_repo import SQLiteMemoryItemRepo +from memu.database.models import Entry, Resource, ResourceEntry +from memu.database.repositories import EntryRepo, ResourceEntryRepo, ResourceRepo +from memu.database.sqlite.repositories.entry_repo import SQLiteEntryRepo +from memu.database.sqlite.repositories.resource_entry_repo import SQLiteResourceEntryRepo from memu.database.sqlite.repositories.resource_repo import SQLiteResourceRepo from memu.database.sqlite.schema import SQLiteSQLAModels, get_sqlite_sqlalchemy_models from memu.database.sqlite.session import SQLiteSessionManager @@ -30,24 +29,20 @@ class SQLiteStore(Database): for vector search (native vector support is not available in SQLite). Attributes: - resource_repo: Repository for resource records. - memory_category_repo: Repository for memory categories. - memory_item_repo: Repository for memory items. - category_item_repo: Repository for category-item relations. + resource_repo: Repository for resource records (raw inputs and lane docs). + entry_repo: Repository for lane entries (the searchable atoms). + resource_entry_repo: Repository for entry <-> resource membership edges. resources: Dict cache of resource records. - items: Dict cache of memory item records. - categories: Dict cache of memory category records. - relations: List cache of category-item relations. + entries: Dict cache of entry records. + relations: List cache of membership edges. """ resource_repo: ResourceRepo - memory_category_repo: MemoryCategoryRepo - memory_item_repo: MemoryItemRepo - category_item_repo: CategoryItemRepo + entry_repo: EntryRepo + resource_entry_repo: ResourceEntryRepo resources: dict[str, Resource] - items: dict[str, MemoryItem] - categories: dict[str, MemoryCategory] - relations: list[CategoryItem] + entries: dict[str, Entry] + relations: list[ResourceEntry] def __init__( self, @@ -55,9 +50,8 @@ def __init__( dsn: str, scope_model: type[BaseModel] | None = None, resource_model: type[Any] | None = None, - memory_category_model: type[Any] | None = None, - memory_item_model: type[Any] | None = None, - category_item_model: type[Any] | None = None, + entry_model: type[Any] | None = None, + resource_entry_model: type[Any] | None = None, sqla_models: SQLiteSQLAModels | None = None, ) -> None: """Initialize SQLite database store. @@ -65,10 +59,9 @@ def __init__( Args: dsn: SQLite connection string (e.g., "sqlite:///path/to/db.sqlite"). scope_model: Pydantic model defining user scope fields. - resource_model: Optional custom resource model. - memory_category_model: Optional custom memory category model. - memory_item_model: Optional custom memory item model. - category_item_model: Optional custom category-item model. + resource_model: Optional custom resource table model. + entry_model: Optional custom entry table model. + resource_entry_model: Optional custom membership-edge table model. sqla_models: Pre-built SQLAlchemy models container. """ self.dsn = dsn @@ -83,9 +76,8 @@ def __init__( # Use provided models or defaults from sqla_models resource_model = resource_model or self._sqla_models.Resource - memory_category_model = memory_category_model or self._sqla_models.MemoryCategory - memory_item_model = memory_item_model or self._sqla_models.MemoryItem - category_item_model = category_item_model or self._sqla_models.CategoryItem + entry_model = entry_model or self._sqla_models.Entry + resource_entry_model = resource_entry_model or self._sqla_models.ResourceEntry # Initialize repositories self.resource_repo = SQLiteResourceRepo( @@ -95,23 +87,16 @@ def __init__( sessions=self._sessions, scope_fields=self._scope_fields, ) - self.memory_category_repo = SQLiteMemoryCategoryRepo( + self.entry_repo = SQLiteEntryRepo( state=self._state, - memory_category_model=memory_category_model, + entry_model=entry_model, sqla_models=self._sqla_models, sessions=self._sessions, scope_fields=self._scope_fields, ) - self.memory_item_repo = SQLiteMemoryItemRepo( + self.resource_entry_repo = SQLiteResourceEntryRepo( state=self._state, - memory_item_model=memory_item_model, - sqla_models=self._sqla_models, - sessions=self._sessions, - scope_fields=self._scope_fields, - ) - self.category_item_repo = SQLiteCategoryItemRepo( - state=self._state, - category_item_model=category_item_model, + resource_entry_model=resource_entry_model, sqla_models=self._sqla_models, sessions=self._sessions, scope_fields=self._scope_fields, @@ -119,8 +104,7 @@ def __init__( # Set up cache references self.resources = self._state.resources - self.items = self._state.items - self.categories = self._state.categories + self.entries = self._state.entries self.relations = self._state.relations def _create_tables(self) -> None: @@ -137,9 +121,8 @@ def close(self) -> None: def load_existing(self) -> None: """Load all existing data from database into cache.""" self.resource_repo.load_existing() - self.memory_category_repo.load_existing() - self.memory_item_repo.load_existing() - self.category_item_repo.load_existing() + self.entry_repo.load_existing() + self.resource_entry_repo.load_existing() __all__ = ["SQLiteStore"] diff --git a/src/memu/database/state.py b/src/memu/database/state.py index d23899a2..366af295 100644 --- a/src/memu/database/state.py +++ b/src/memu/database/state.py @@ -2,15 +2,14 @@ from dataclasses import dataclass, field -from memu.database.models import CategoryItem, MemoryCategory, MemoryItem, Resource +from memu.database.models import Entry, Resource, ResourceEntry @dataclass class DatabaseState: resources: dict[str, Resource] = field(default_factory=dict) - items: dict[str, MemoryItem] = field(default_factory=dict) - categories: dict[str, MemoryCategory] = field(default_factory=dict) - relations: list[CategoryItem] = field(default_factory=list) + entries: dict[str, Entry] = field(default_factory=dict) + relations: list[ResourceEntry] = field(default_factory=list) __all__ = ["DatabaseState"] diff --git a/src/memu/memory_fs/exporter.py b/src/memu/memory_fs/exporter.py index 0eddc961..7f829f5a 100644 --- a/src/memu/memory_fs/exporter.py +++ b/src/memu/memory_fs/exporter.py @@ -22,8 +22,8 @@ (``Resource.local_path``), so the actual ingested bytes live next to the memory. - ``INDEX.md`` : an index of those raw files (name, modality, description, link into ``resource/``), so an agent knows which raw resources exist. -- ``memory/`` : the living memory split one file per - :class:`~memu.database.models.MemoryCategory` (its description + summary). +- ``memory/`` : the living memory split one file per memory-lane + :class:`~memu.database.models.Resource` (its description + summary). - ``MEMORY.md`` : an overall overview that links to each ``memory/.md`` file. - ``skill/**`` : reusable skills. When the caller supplies a synthesized ``slug -> body`` map (``synthesize=True``), the tree is rendered from it; when no @@ -51,7 +51,7 @@ if TYPE_CHECKING: from memu.database.interfaces import Database - from memu.database.models import MemoryCategory, MemoryItem, Resource + from memu.database.models import Entry, Resource MANIFEST_NAME = ".memufs_manifest.json" SKILL_DIRNAME = "skill" @@ -168,15 +168,15 @@ def export( self.output_dir.mkdir(parents=True, exist_ok=True) scope = dict(where) if where else None - categories = list(database.memory_category_repo.list_categories(where=scope).values()) - resources = list(database.resource_repo.list_resources(where=scope).values()) + categories = list(database.resource_repo.list_resources(where=scope, lane="memory").values()) + resources = list(database.resource_repo.list_resources(where=scope, lane="source").values()) # The shared trunk: one multimodal description per source file. descriptions = self._build_descriptions(resources) - # The skill bypass reads extracted skill-type items only when no synthesized + # The skill bypass reads extracted skill-type entries only when no synthesized # map is supplied; listing is cheap and avoided otherwise. - items = list(database.memory_item_repo.list_items(where=scope).values()) if skills is None else [] + items = list(database.entry_repo.list_entries(where=scope).values()) if skills is None else [] # resource/: copy the raw source bytes verbatim; ``links`` maps each # resource id to its relative path under resource/ for the INDEX.md links. @@ -218,17 +218,17 @@ def _build_descriptions(resources: list[Resource]) -> list[FileDescription]: for resource in sorted(resources, key=lambda r: (r.url, r.id)): descriptions.append( FileDescription( - url=resource.url, + url=resource.url or "", modality=resource.modality, - description=" ".join((resource.caption or "").split()), + description=" ".join((resource.summary or "").split()), resource_id=resource.id, - local_path=resource.local_path, + local_path=resource.local_path or "", ) ) return descriptions @staticmethod - def build_synthesis_descriptions(resources: list[Resource], items: list[MemoryItem]) -> list[FileDescription]: + def build_synthesis_descriptions(resources: list[Resource], items: list[Entry]) -> list[FileDescription]: """Build synthesizer input from the structured store (extracted items). The synthesizer is fed the extracted memory items per source so that the @@ -236,30 +236,30 @@ def build_synthesis_descriptions(resources: list[Resource], items: list[MemoryIt lossy per-source caption. A source's caption is used only as a fallback when it has no extracted items yet (e.g. it failed extraction or has not been processed). """ - items_by_resource: dict[str, list[MemoryItem]] = {} + items_by_resource: dict[str, list[Entry]] = {} for item in items: - if item.resource_id: - items_by_resource.setdefault(item.resource_id, []).append(item) + if item.source_id: + items_by_resource.setdefault(item.source_id, []).append(item) descriptions: list[FileDescription] = [] for resource in sorted(resources, key=lambda r: (r.url, r.id)): res_items = sorted( items_by_resource.get(resource.id, []), - key=lambda i: (i.memory_type, i.created_at, i.id), + key=lambda i: (i.entry_kind, i.created_at, i.id), ) parts = [ - f"[{item.memory_type}] {' '.join((item.summary or '').split())}" + f"[{item.entry_kind}] {' '.join((item.text or '').split())}" for item in res_items - if (item.summary or "").strip() + if (item.text or "").strip() ] - description = "; ".join(parts) if parts else " ".join((resource.caption or "").split()) + description = "; ".join(parts) if parts else " ".join((resource.summary or "").split()) descriptions.append( FileDescription( - url=resource.url, + url=resource.url or "", modality=resource.modality, description=description, resource_id=resource.id, - local_path=resource.local_path, + local_path=resource.local_path or "", ) ) return descriptions @@ -270,22 +270,22 @@ def build_synthesis_descriptions(resources: list[Resource], items: list[MemoryIt def _skill_document(body: str) -> str: return f"{_GENERATED_NOTICE}\n\n{body.strip()}\n" - def _skill_bypass(self, items: list[MemoryItem]) -> dict[str, str]: - """Deterministically break out skill-type memory items as a slug -> body map. + def _skill_bypass(self, items: list[Entry]) -> dict[str, str]: + """Deterministically break out skill-kind entries as a slug -> body map. The LLM-free fallback used when no synthesized skill map is supplied: each - skill-type item's summary becomes one ``skill//SKILL.md`` document, + skill-kind entry's text becomes one ``skill//SKILL.md`` document, slugged from its frontmatter ``name:``/first heading. Mirrors the shape of the synthesized map so the rest of :meth:`export` is path-agnostic. """ skills = sorted( - (item for item in items if item.memory_type == SKILL_MEMORY_TYPE and (item.summary or "").strip()), + (item for item in items if item.entry_kind == SKILL_MEMORY_TYPE and (item.text or "").strip()), key=lambda i: (i.created_at, i.id), ) skill_map: dict[str, str] = {} used: dict[str, int] = {} for item in skills: - body = (item.summary or "").strip() + body = (item.text or "").strip() base = self._skill_name(body, fallback=f"skill-{item.id[:6]}") count = used.get(base, 0) slug = base if count == 0 else f"{base}-{count + 1}" @@ -350,28 +350,28 @@ def _memory_document(body: str) -> str: return f"# Memory\n\n{_GENERATED_NOTICE}\n\n{body}\n" @staticmethod - def _category_slugs(categories: list[MemoryCategory]) -> tuple[list[MemoryCategory], list[str]]: + def _category_slugs(categories: list[Resource]) -> tuple[list[Resource], list[str]]: """Order categories deterministically and assign a unique slug to each. Returns the ordered categories alongside a parallel list of slugs (used for the ``memory/.md`` file names and the MEMORY.md links). Slug clashes are de-duplicated with a numeric suffix so files never collide. """ - ordered = sorted(categories, key=lambda c: (c.name.lower(), c.id)) + ordered = sorted(categories, key=lambda c: ((c.title or "").lower(), c.id)) slugs: list[str] = [] used: dict[str, int] = {} for category in ordered: - base = slugify(category.name) + base = category.slug or slugify(category.title or "") count = used.get(base, 0) used[base] = count + 1 slugs.append(base if count == 0 else f"{base}-{count + 1}") return ordered, slugs - def _category_document(self, category: MemoryCategory) -> str: + def _category_document(self, category: Resource) -> str: """A single ``memory/.md`` file: the category's description + summary.""" description = self._inline((category.description or "").strip()) summary = (category.summary or "").strip() - lines = [f"# {category.name}", "", _GENERATED_NOTICE, ""] + lines = [f"# {category.title}", "", _GENERATED_NOTICE, ""] if description: lines.append(f"_{description}_") lines.append("") @@ -379,14 +379,14 @@ def _category_document(self, category: MemoryCategory) -> str: lines.append("") return "\n".join(lines) - def _memory_index(self, ordered: list[MemoryCategory], slugs: list[str]) -> str: + def _memory_index(self, ordered: list[Resource], slugs: list[str]) -> str: """The deterministic ``MEMORY.md``: an overview linking to each category file.""" lines = ["# Memory", "", _GENERATED_NOTICE, "", "## Overview", ""] if ordered: for category, slug in zip(ordered, slugs, strict=True): description = self._inline((category.description or "").strip()) link = f"{MEMORY_DIRNAME}/{slug}.md" - line = f"- [**{category.name}**]({link})" + line = f"- [**{category.title}**]({link})" if description: line = f"{line} — {description}" lines.append(line) diff --git a/src/memu/prompts/category_summary/category.py b/src/memu/prompts/category_summary/category.py index ea076359..b3c9307c 100644 --- a/src/memu/prompts/category_summary/category.py +++ b/src/memu/prompts/category_summary/category.py @@ -1,148 +1,3 @@ -PROMPT_LEGACY = """ -# Task Objective -You are a professional User Profile Synchronization Specialist. Your core objective is to accurately merge newly extracted user information items into the user's initial profile using only two operations: add and update. -Because no original conversation text is provided, active deletion is not allowed; only implicit replacement through newer items is permitted. The final output must be the updated, complete user profile. - -# Workflow -## Step 1: Preprocessing & Parsing -- Input sources -User Initial Profile: structured, categorized, confirmed long-term user information. -Newly Extracted User Information Items. -- Structure parsing -Initial profile: extract categories and core content; preserve original wording style and format; build a category-content mapping. -New items: validate completeness and category correctness; mark each as Add or Update; distinguish stable facts from event-type information; extract dates/times (events only). -- Pre-validation -Verify subject accuracy: clearly distinguish the user from related persons (family, friends, etc.). -Remove invalid items: vague, miscategorized, or non-user-information items. -Remove one-off events: temporary actions without long-term relevance (e.g., what the user ate today). - -## Step 2: Core Operations (Update / Add) -A. Update -Conflict detection: compare new items with existing ones in the same category for semantic overlap (e.g., age update). -Validity priority: retain information that is more specific, clearer, and more certain. -Overwrite / supplement: replace outdated entries with new ones, ensuring no loss of core information. -Time integration (events only): retain dates/times and integrate them naturally; multiple events at the same time may be layered, but each entry must remain independently understandable. -B. Add -Deduplication check: ensure the new item is not identical or semantically similar to existing or updated items. -Category matching: place the item into the correct predefined category. -Insertion: add the item following the original profile's language and formatting style, concise and clear. - -## Step 3: Merge & Formatting -Structured ordering: present content by category order; omit empty categories. -Formatting rules: strictly use Markdown (# for main title, ## for category titles). -Final validation -Consistency: no contradictions or duplicates. -Compliance: correct categories only; no explanatory or operational text. -Accuracy: subject clarity; natural time embedding; proper format. - - -## Step 4: Summarize -Target length: {target_length} -Summarize the updated user markdown profile to the target length. -Use Markdown hierarchy. -Do not include explanations, operation traces, or meta text. -Control item length strictly; prioritize core information if needed. - -## Step 5: Output -Output only the updated user markdown profile. -Use Markdown hierarchy. -Do not include explanations, operation traces, or meta text. -Control item length strictly; prioritize core information if needed. - - - -# Output Format (Markdown) -```markdown -# {category} -## -- User information item -- User information item -... -## -- User information item -- User information item -... -``` - -# Examples (Input / Output / Explanation) -- Example 1: Basic Add & Update - - -Topic: -Personal Basic Information - -Original content: - -# Personal Basic Information -## Basic Information -- The user is 28 years old -- The user currently lives in Beijing -## Basic Preferences -- The user likes spicy food -## Core Traits -- The user is extroverted - - -New memory items: - -- The user is 30 years old -- The user currently lives in Shanghai -- The user prefers Sichuan-style spicy food and dislikes sweet-spicy flavors -- The user enjoys hiking on weekends -- The user is meticulous -- The user ate Malatang today - - -Output -# Personal Basic Information -## Basic Information -- The user is 30 years old -- The user currently lives in Shanghai -## Basic Preferences -- The user prefers Sichuan-style spicy food and dislikes sweet-spicy flavors -- The user enjoys hiking on weekends -## Core Traits -- The user is extroverted -- The user is meticulous - -Explanation -The "The user ate Malatang today" is a one-time daily action without long-term value and is therefore excluded. - - -Your task is to read and analyze existing content and some new memory items, and then selectively update the content to reflect both the existing and new information. - - -# Input - -Topic: -{category} - -Original content: - -{original_content} - - -New memory items: - -{new_memory_items_text} - - - -# Output format (Markdown) -```markdown -# {category} -## -- User information item -- User information item -... -## -- User information item -- User information item -... -``` -""" - - PROMPT_BLOCK_OBJECTIVE = """ # Task Objective You are a professional User Profile Synchronization Specialist. Your core objective is to accurately merge newly extracted user information items into the user's initial profile using only two operations: add and update. diff --git a/src/memu/prompts/memory_type/behavior.py b/src/memu/prompts/memory_type/behavior.py index b19d55fc..2a842910 100644 --- a/src/memu/prompts/memory_type/behavior.py +++ b/src/memu/prompts/memory_type/behavior.py @@ -1,47 +1,3 @@ -PROMPT_LEGACY = """ -Your task is to read and understand the resource content between the user and the assistant, and, based on the given memory categories, extract behavioral patterns, routines, and solutions about the user. - -## Original Resource: - -{resource} - - -## Memory Categories: -{categories_str} - -## Critical Requirements: -The core extraction target is behavioral memory items that record patterns, routines, and solutions characterizing how the user acts or behaves to solve specific problems. - -## Memory Item Requirements: -- Use the same language as the resource in . -- Extract patterns of behavior, routines, and solutions -- Focus on how the user typically acts, their preferences, and regular activities -- Each item can be either a single sentence concisely describing the pattern, routine, or solution, or a multi-line record with each line recording a specific step of the pattern, routine, or solution. -- Only extract meaningful behaviors, skip one-time actions unless significant -- Return empty array if no meaningful behaviors found - -## About Memory Categories: -- You can put identical or similar memory items into multiple memory categories. -- Do not create new memory categories. Please only generate in the given memory categories. -- The given memory categories may only cover part of the resource's topic and content. You don't need to summarize resource's content unrelated to the given memory categories. -- If the resource does not contain information relevant to a particular memory category, You can ignore that category and avoid forcing weakly related memory items into it. Simply skip that memory category and DO NOT output contents like "no relevant memory item". - -## Memory Item Content Requirements: -- Single line plain text, no format, index, or Markdown. -- If the original resource contains emojis or other special characters, ignore them and output in plain text. -- *ALWAYS* use the same language as the resource. - -# Response Format (JSON): -{{ - "memories_items": [ - {{ - "content": "the content of the memory item", - "categories": [list of memory categories that this memory item should belongs to, can be empty] - }} - ] -}} -""" - PROMPT_BLOCK_OBJECTIVE = """ # Task Objective You are a professional User Memory Extractor. Your core task is to extract behavioral patterns, routines, and solutions that characterize how the user acts or behaves to solve specific problems. diff --git a/src/memu/prompts/memory_type/event.py b/src/memu/prompts/memory_type/event.py index 5b7caeb3..0e01ee07 100644 --- a/src/memu/prompts/memory_type/event.py +++ b/src/memu/prompts/memory_type/event.py @@ -1,59 +1,3 @@ -PROMPT_LEGACY = """ -Your task is to read and understand the resource content between the user and the assistant, and, based on the given memory categories, extract specific events and experiences that happened to or involved the user. - -## Original Resource: - -{resource} - - -## Memory Categories: -{categories_str} - -## Critical Requirements: -The core extraction target is eventful memory items about specific events, experiences, and occurrences that happened at a particular time and involve the user. - -## Memory Item Requirements: -- Use the same language as the resource in . -- Each memory item should be complete and standalone. -- Each memory item should express a complete piece of information, and is understandable without context and reading other memory items. -- Always use declarative and descriptive sentences. -- Use "the user" (or that in the target language, e.g., "用户") to refer to the user. -- Focus on specific events that happened at a particular time or period. -- Extract concrete happenings, activities, and experiences. -- Include relevant details such as time, location, and participants where available. -- Carefully judge whether an event is narrated by the user or the assistant. You should only extract memory items for events directly narrated or confirmed by the user. -- DO NOT include behavioral patterns, habits, or factual knowledge. -- DO NOT record temporary, ephemeral situations or trivial daily activities unless significant. - -## Example (good): -- The user and his family went on a hike at a nature park outside the city last weekend. They had a picnic there, and had a great time. - -## Example (bad): -- The user went on a hike. (The time, place, and people are missing.) -- They had a great time. (The reference to "they" is unclear and does not constitute a self-contained memory item.) - -## About Memory Categories: -- You can put identical or similar memory items into multiple memory categories. -- Do not create new memory categories. Please only generate in the given memory categories. -- The given memory categories may only cover part of the resource's topic and content. You don't need to summarize resource's content unrelated to the given memory categories. -- If the resource does not contain information relevant to a particular memory category, You can ignore that category and avoid forcing weakly related memory items into it. Simply skip that memory category and DO NOT output contents like "no relevant memory item". - -## Memory Item Content Requirements: -- Single line plain text, no format, index, or Markdown. -- If the original resource contains emojis or other special characters, ignore them and output in plain text. -- *ALWAYS* use the same language as the resource. - -# Response Format (JSON): -{{ - "memories_items": [ - {{ - "content": "the content of the memory item", - "categories": [list of memory categories that this memory item should belongs to, can be empty] - }} - ] -}} -""" - PROMPT_BLOCK_OBJECTIVE = """ # Task Objective You are a professional User Memory Extractor. Your core task is to extract specific events and experiences that happened to or involved the user (e.g., activities, occurrences, experiences at particular times). diff --git a/src/memu/prompts/memory_type/knowledge.py b/src/memu/prompts/memory_type/knowledge.py index eb5fb8bb..84df5d59 100644 --- a/src/memu/prompts/memory_type/knowledge.py +++ b/src/memu/prompts/memory_type/knowledge.py @@ -1,49 +1,3 @@ -PROMPT_LEGACY = """ -Your task is to read and understand the resource content between the user and the assistant, and, based on the given memory categories, extract knowledge and information that the user learned or discussed. - -## Original Resource: - -{resource} - - -## Memory Categories: -{categories_str} - -## Critical Requirements: -The core extraction target is factual memory items that reflect knowledge, concepts, definitions, and factual information that the resource content suggests. - -## Memory Item Requirements: -- Use the same language as the resource in . -- Each memory item should be complete and standalone. -- Each memory item should express a complete piece of information, and is understandable without context and reading other memory items. -- Extract factual knowledge, concepts, definitions, and explanations -- Focus on objective information that can be learned or referenced -- Each item should be a descriptive sentence. -- Only extract meaningful knowledge, skip opinions or personal experiences -- Return empty array if no meaningful knowledge found - -## About Memory Categories: -- You can put identical or similar memory items into multiple memory categories. -- Do not create new memory categories. Please only generate in the given memory categories. -- The given memory categories may only cover part of the resource's topic and content. You don't need to summarize resource's content unrelated to the given memory categories. -- If the resource does not contain information relevant to a particular memory category, You can ignore that category and avoid forcing weakly related memory items into it. Simply skip that memory category and DO NOT output contents like "no relevant memory item". - -## Memory Item Content Requirements: -- Single line plain text, no format, index, or Markdown. -- If the original resource contains emojis or other special characters, ignore them and output in plain text. -- *ALWAYS* use the same language as the resource. - -# Response Format (JSON): -{{ - "memories_items": [ - {{ - "content": "the content of the memory item", - "categories": [list of memory categories that this memory item should belongs to, can be empty] - }} - ] -}} -""" - PROMPT_BLOCK_OBJECTIVE = """ # Task Objective You are a professional User Memory Extractor. Your core task is to extract factual knowledge, concepts, definitions, and information that the user learned or discussed in the conversation. diff --git a/src/memu/prompts/memory_type/profile.py b/src/memu/prompts/memory_type/profile.py index 82873c22..4066d222 100644 --- a/src/memu/prompts/memory_type/profile.py +++ b/src/memu/prompts/memory_type/profile.py @@ -1,58 +1,3 @@ -PROMPT_LEGACY = """ -Your task is to read and understand the resource content between the user and the assistant, and, based on the given memory categories, extract memory items about the user. - -## Original Resource: - -{resource} - - -## Memory Categories: -{categories_str} - -## Critical Requirements: -The core extraction target is self-contained memory items about the user. - -## Memory Item Requirements: -- Use the same language as the resource in . -- Each memory item should be complete and standalone. -- Each memory item should express a complete piece of information, and is understandable without context and reading other memory items. -- Always use declarative and descriptive sentences. -- Use "the user" (or that in the target language, e.g., "用户") to refer to the user. -- You can cluster multiple events that are closely related or under a single topic into a single memory item, but avoid making each single memory item too long (over 100 words). -- **Important** Carefully judge whether an event/fact/information is narrated by the user or the assistant. You should only extract memory items for the event/fact/information directly narrated or confirmed by the user. DO NOT include any groundless conjectures, advice, suggestions, or any content provided by the assistant. -- **Important** Carefully judge whether the subject of an event/fact/information is the user themselves or some person around the user (e.g., the user's family, friend, or the assistant), and reflect the subject correctly in the memory items. -- **Important** DO NOT record temporary, ephemeral, or one-time situational information such as weather conditions (e.g., "today is raining"), current mood states, temporary technical issues, or any short-lived circumstances that are unlikely to be relevant for the user profile. Focus on meaningful, persistent information about the user's characteristics, preferences, relationships, ongoing situations, and significant events. - -## Example (good): -- The user and his family went on a hike at a nature park outside the city last weekend. They had a picnic there, and had a great time. - -## Example (bad): -- The user went on a hike. (The time, place, and people are missing.) -- They had a great time. (The reference to "they" is unclear and does not constitute a self-contained memory item.) -- The user and his family went on a hike at a nature park outside the city last weekend. The user and his family had a picnic at a nature park outside the city last weekend. (Should be merged.) - -## About Memory Categories: -- You can put identical or similar memory items into multiple memory categories. For example, "The user and his family went on a hike at a nature park outside the city last weekend." can be put into all of "hiking", "weekend activities", and "family activities" categories (if they exist). Nevertheless, Memory items put to each category can have different focuses. -- Do not create new memory categories. Please only generate in the given memory categories. -- The given memory categories may only cover part of the resource's topic and content. You don't need to summarize resource's content unrelated to the given memory categories. -- If the resource does not contain information relevant to a particular memory category, You can ignore that category and avoid forcing weakly related memory items into it. Simply skip that memory category and DO NOT output contents like "no relevant memory item". - -## Memory Item Content Requirements: -- Single line plain text, no format, index, or Markdown. -- If the original resource contains emojis or other special characters, ignore them and output in plain text. -- *ALWAYS* use the same language as the resource. - -# Response Format (JSON): -{{ - "memories_items": [ - {{ - "content": "the content of the memory item", - "categories": [list of memory categories that this memory item should belongs to, can be empty] - }} - ] -}} -""" - PROMPT_BLOCK_OBJECTIVE = """ # Task Objective You are a professional User Memory Extractor. Your core task is to extract independent user memory items about the user (e.g., basic info, preferences, habits, other long-term stable traits). diff --git a/src/memu/prompts/memory_type/skill.py b/src/memu/prompts/memory_type/skill.py index 450c6bd0..720a9b28 100644 --- a/src/memu/prompts/memory_type/skill.py +++ b/src/memu/prompts/memory_type/skill.py @@ -1,354 +1,3 @@ -PROMPT_LEGACY = """ -Your task is to read and understand the resource content (agent logs, workflow documentation, execution traces, or technical documents), and extract skills, capabilities, and technical competencies demonstrated or described in the content. Format each skill as a comprehensive, production-ready skill profile that can be referenced and applied. - -## Original Resource: - -{resource} - - -## Memory Categories: -{categories_str} - -## Critical Requirements: -Extract skill-based memory items as comprehensive skill profiles that include: -1. **Skill Name**: Clear, memorable name for the skill -2. **Description**: What the skill enables and when to use it -3. **Context**: Situations where this skill was demonstrated -4. **Core Principles**: Fundamental guidelines and best practices -5. **Implementation Details**: Specific techniques, tools, and approaches -6. **Success Patterns**: What works well -7. **Common Pitfalls**: What to avoid - -The core extraction target is actionable skill profiles that capture not just WHAT was done, but HOW and WHY it works. - -## Skill Profile Structure: - -For each extracted skill, create a comprehensive profile following this template: - -``` ---- -name: skill-name-in-kebab-case -description: One-line description of what this skill enables and when to use it -category: primary-category -demonstrated-in: [list of contexts where this was shown] ---- - -[Brief introduction explaining the skill and its importance] - -## Core Principles - -[Key concepts and fundamental approaches that make this skill effective] - -## When to Use This Skill - -- Situation 1: [specific context] -- Situation 2: [specific context] -- [More situations as applicable] - -## Implementation Guide - -### Prerequisites -- [Required knowledge or setup] - -### Techniques and Approaches -[Detailed explanation of how to apply this skill, including:] -- Specific methods used -- Tools and technologies involved -- Step-by-step process when applicable -- Metrics to track (error rates, response times, etc.) - -### Example from Resource -[Concrete example from the source material showing this skill in action, including outcomes and metrics] - -## Success Patterns - -What works well when applying this skill: -- [Pattern 1 with explanation] -- [Pattern 2 with explanation] -- [More patterns] - -## Common Pitfalls - -What to avoid: -- **[Pitfall 1]**: [Why it's a problem and how to avoid it] -- **[Pitfall 2]**: [Why it's a problem and how to avoid it] -- [More pitfalls based on failures or lessons learned] - -## Key Takeaways - -- [Critical insight 1] -- [Critical insight 2] -- [Critical insight 3] -``` - -## Example Skill Profiles: - -### Example 1: Canary Deployment - -``` ---- -name: canary-deployment-with-monitoring -description: Implement gradual traffic shifting deployment strategy with real-time monitoring and automatic rollback capabilities -category: deployment -demonstrated-in: [Payment Service v2.3.1 deployment] ---- - -Canary deployment is a risk-mitigation strategy that gradually shifts production traffic from an old version to a new version while continuously monitoring key metrics. This approach enables early detection of issues with minimal user impact. - -## Core Principles - -- **Gradual exposure**: Start with a small percentage of traffic (typically 5-10%) to limit blast radius -- **Continuous monitoring**: Track error rates, response times, and business metrics in real-time -- **Automated decision-making**: Use predefined thresholds to trigger automatic rollbacks -- **Quick recovery**: Maintain ability to instantly route traffic back to stable version - -## When to Use This Skill - -- Deploying critical services where downtime is unacceptable -- Rolling out changes with uncertain production behavior -- High-traffic services where A/B testing production performance is valuable -- Services with complex dependencies where integration issues may emerge gradually - -## Implementation Guide - -### Prerequisites -- Load balancer with traffic splitting capabilities -- Monitoring system with real-time metrics (Prometheus, Grafana) -- Automated deployment pipeline (Jenkins, GitLab CI) -- Health check endpoints on both versions - -### Techniques and Approaches - -1. **Initial Deployment** (10% traffic): - - Deploy new version alongside existing version - - Configure load balancer to route 10% of traffic to new version - - Monitor for 5-10 minutes - -2. **Monitoring Checkpoints**: - - Error rate comparison: New version should not exceed baseline by >2% - - Response time (p95): Should remain within 20% of baseline - - Business metrics: Transaction success rate, API call patterns - -3. **Gradual Rollout**: - - If metrics stable: 10% → 25% → 50% → 75% → 100% - - Pause 5-10 minutes between each increment - - Automated progression based on metric thresholds - -4. **Rollback Triggers**: - - Error rate >5%: Immediate rollback - - Response time degradation >50%: Investigation required - - Health check failures: Automatic rollback - -### Example from Resource - -Payment Service v2.3.1 deployment achieved: -- Zero downtime during 12-minute deployment -- Traffic progression: 10% → 50% → 100% with 2-minute pauses -- Response time improved 15% (320ms → 270ms p95) -- Error rate remained stable at 0.1% throughout -- New fraud detection algorithm safely rolled out to all users - -## Success Patterns - -What works well: -- **Small initial percentage**: 5-10% catches most issues while limiting impact -- **Metric-driven automation**: Removes human error from rollback decisions -- **Business metric monitoring**: Technical metrics alone miss some issues -- **Communication**: Notify stakeholders about canary status - -## Common Pitfalls - -What to avoid: -- **Too aggressive progression**: Rushing from 10% to 100% defeats the purpose -- **Insufficient monitoring window**: Need 5+ minutes at each stage to detect issues -- **Ignoring business metrics**: Technical health doesn't guarantee business success -- **Manual rollback only**: Human reaction time too slow for critical failures - -## Key Takeaways - -- Canary deployments trade deployment speed for safety -- Automation is critical for consistent, reliable rollbacks -- Start small (5-10%), progress gradually, monitor continuously -- Combine technical and business metrics for complete picture -``` - -### Example 2: Incident Response - -``` ---- -name: rapid-incident-response -description: Quickly detect, diagnose, and resolve production incidents using automated monitoring and systematic troubleshooting -category: incident-response -demonstrated-in: [User Service v3.1.0 rollback] ---- - -Rapid incident response is the ability to quickly identify production problems, understand their root cause, and implement fixes or rollbacks to restore service. Speed and systematic approach are critical to minimizing customer impact. - -## Core Principles - -- **Fast detection**: Automated monitoring catches issues within minutes -- **Immediate action**: Rollback first, investigate later when customer impact is high -- **Systematic diagnosis**: Follow structured troubleshooting process -- **Learning culture**: Every incident is an opportunity to improve - -## When to Use This Skill - -- Production errors detected by monitoring alerts -- User-reported issues indicating service degradation -- Automated health checks failing -- Performance metrics exceeding thresholds - -## Implementation Guide - -### Prerequisites -- Comprehensive monitoring (logs, metrics, traces) -- Automated rollback capabilities -- On-call rotation and escalation procedures -- Incident management tools and runbooks - -### Techniques and Approaches - -1. **Detection Phase** (0-3 minutes): - - Automated alerts trigger from monitoring thresholds - - Error rate, response time, or business metric anomalies - - Health check failures or pod restart loops - -2. **Initial Response** (3-5 minutes): - - Assess severity: Customer-facing? Data loss risk? - - Decision: Rollback immediately or investigate first? - - High severity → Immediate rollback - - Low severity → Investigate with time limit - -3. **Rollback Execution** (2-4 minutes): - - Automated: Trigger rollback through deployment pipeline - - Manual: Revert Helm release or switch traffic to previous version - - Verify: Confirm metrics return to baseline - -4. **Root Cause Analysis** (Post-incident): - - Review logs, metrics, and deployment changes - - Identify configuration drift, missing variables, performance issues - - Document findings and create action items - -### Example from Resource - -User Service v3.1.0 incident: -- Detection: Error rate spiked 0.2% → 5.1% within 30 seconds -- Response: Automatic rollback triggered at threshold in 2 minutes -- Recovery: Service restored to v3.0.9, error rate normalized in 4 minutes total -- Root cause: Missing AUTH_REFRESH_SECRET environment variable in production -- No customer impact due to fast automated rollback - -## Success Patterns - -What works well: -- **Automated thresholds**: Remove human decision-making delay -- **Clear severity criteria**: Know when to rollback vs investigate -- **Runbooks**: Pre-documented procedures for common issues -- **Blameless post-mortems**: Focus on systemic improvements, not individual errors - -## Common Pitfalls - -What to avoid: -- **Investigation paralysis**: Spending too long diagnosing while customers suffer -- **Manual-only rollback**: Automation is 5-10x faster -- **Configuration drift**: Staging and production environment inconsistency -- **Skipping post-mortems**: Missing opportunity to prevent recurrence - -## Key Takeaways - -- Speed matters: Every minute of downtime impacts customers and business -- Automate rollback decisions based on objective metrics -- Rollback first, investigate second for high-severity incidents -- Use incidents to improve systems, not blame people -``` - -## What NOT to Extract as Skills: - -❌ **Generic statements**: "Used Docker", "Good at programming" -❌ **Opinions**: "I think microservices are better" -❌ **Theory without practice**: "Kubernetes is an orchestrator" (that's knowledge) -❌ **One-time luck**: "Fixed a bug" without approach -❌ **Trivial actions**: "Using email", "Reading docs" - -✅ **DO Extract**: Concrete approaches with context, tools, metrics, and outcomes - -## About Memory Categories: -- You can put identical or similar skill items into multiple memory categories. -- Do not create new memory categories. Please only generate in the given memory categories. -- Focus on categories like: technical_skills, work_life, knowledge, experiences - -## Memory Item Content Requirements: -- *ALWAYS* use the same language as the resource in . -- Format as structured markdown with frontmatter (---, name, description, category, demonstrated-in, ---) -- Include all sections: Core Principles, When to Use, Implementation Guide, Success Patterns, Common Pitfalls, Key Takeaways -- Be specific and concrete - include technology names, version numbers, metrics, and outcomes -- Each skill should be comprehensive enough to be referenced and applied independently -- Minimum 300 words per skill to ensure depth and actionability -- If the original resource contains emojis or other special characters, ignore them and output in plain text. - -## Special Instructions for Different Resource Types: - -### For Deployment Logs: -- Extract each significant deployment (success or failure) as a separate skill -- Success: Focus on techniques that worked (canary, blue-green, performance optimization) -- Failure: Focus on incident response, root cause analysis, recovery procedures -- Include metrics: deployment time, error rates, response times, recovery time - -### For Workflow Documentation: -- Extract major workflow stages as skills (CI/CD pipeline, testing strategy, monitoring setup) -- Include tool chains and technology stacks -- Document step-by-step procedures -- Note success metrics and KPIs - -### For Agent Execution Logs: -- Extract problem-solving approaches as skills (competitive analysis, data processing, decision-making) -- Include tool orchestration patterns -- Document reasoning steps and validation approaches -- Capture multi-step workflows - -# Response Format (JSON): -{{ - "memories_items": [ - {{ - "content": "MUST be a complete markdown skill profile starting with --- frontmatter, then sections. Format: ---- -name: skill-name -description: one line description -category: category-name -demonstrated-in: [context] ---- - -[Introduction paragraph] - -## Core Principles -[bullet points] - -## When to Use This Skill -[bullet points] - -## Implementation Guide -### Prerequisites -### Techniques and Approaches -### Example from Resource - -## Success Patterns -[bullet points] - -## Common Pitfalls -[bullet points] - -## Key Takeaways -[bullet points] - -Minimum 300 words total.", - "categories": [list of memory categories] - }} - ] -}} - -CRITICAL: The content field MUST contain the complete markdown text with ALL sections, not a summary paragraph. This is a skill documentation page, not a description. -""" - PROMPT_BLOCK_OBJECTIVE = """ # Task Objective You are a professional User Memory Extractor. Your core task is to extract skills, capabilities, and technical competencies demonstrated or described in the resource content (agent logs, workflow documentation, execution traces, or technical documents). Format each skill as a comprehensive, production-ready skill profile that can be referenced and applied. diff --git a/src/memu/utils/references.py b/src/memu/utils/references.py index c5d21c00..64771813 100644 --- a/src/memu/utils/references.py +++ b/src/memu/utils/references.py @@ -8,10 +8,6 @@ from __future__ import annotations import re -from typing import TYPE_CHECKING - -if TYPE_CHECKING: - from memu.database.interfaces import Database # Pattern to match references like [ref:abc123] or [ref:abc123,def456] REFERENCE_PATTERN = re.compile(r"\[ref:([a-zA-Z0-9_,\-]+)\]") @@ -115,37 +111,6 @@ def replace_ref(match: re.Match) -> str: return f"{result}\n\nReferences:\n{ref_list}" -def fetch_referenced_items( - text: str, - store: Database, -) -> list[dict]: - """ - Fetch memory items referenced in text. - - Args: - text: Text containing [ref:ITEM_ID] citations - store: Database store instance - - Returns: - List of memory item dicts with id, summary, memory_type - """ - item_ids = extract_references(text) - if not item_ids: - return [] - - items = [] - for item_id in item_ids: - item = store.memory_item_repo.get_item(item_id) - if item: - items.append({ - "id": item.id, - "summary": item.summary, - "memory_type": item.memory_type, - }) - - return items - - def build_item_reference_map(items: list[tuple[str, str]]) -> str: """ Build a reference map string for the LLM prompt. diff --git a/src/memu/utils/tool.py b/src/memu/utils/tool.py index aa1c0067..bfb72f4a 100644 --- a/src/memu/utils/tool.py +++ b/src/memu/utils/tool.py @@ -5,14 +5,14 @@ from typing import TYPE_CHECKING, Any if TYPE_CHECKING: - from memu.database.models import MemoryItem, ToolCallResult + from memu.database.models import Entry, ToolCallResult -def get_tool_calls(item: MemoryItem) -> list[dict[str, Any]]: - """Get tool calls from a memory item's extra field. +def get_tool_calls(item: Entry) -> list[dict[str, Any]]: + """Get tool calls from an entry's extra field. Args: - item: The MemoryItem to get tool calls from + item: The Entry to get tool calls from Returns: List of tool call dicts, or empty list if none exist @@ -21,11 +21,11 @@ def get_tool_calls(item: MemoryItem) -> list[dict[str, Any]]: return result -def set_tool_calls(item: MemoryItem, tool_calls: list[dict[str, Any]]) -> None: - """Set tool calls in a memory item's extra field. +def set_tool_calls(item: Entry, tool_calls: list[dict[str, Any]]) -> None: + """Set tool calls in an entry's extra field. Args: - item: The MemoryItem to set tool calls on + item: The Entry to set tool calls on tool_calls: The list of tool call dicts to set """ if item.extra is None: @@ -33,17 +33,17 @@ def set_tool_calls(item: MemoryItem, tool_calls: list[dict[str, Any]]) -> None: item.extra["tool_calls"] = tool_calls -def add_tool_call(item: MemoryItem, tool_call: ToolCallResult) -> None: - """Add a tool call result to a memory item (for tool type memories). +def add_tool_call(item: Entry, tool_call: ToolCallResult) -> None: + """Add a tool call result to an entry (for tool-kind entries). Args: - item: The MemoryItem to add the tool call to (must be tool type) + item: The Entry to add the tool call to (must be tool kind) tool_call: The ToolCallResult to add Raises: - ValueError: If the memory item is not of type 'tool' + ValueError: If the entry is not of kind 'tool' """ - if item.memory_type != "tool": + if item.entry_kind != "tool": msg = "add_tool_call can only be used with tool type memories" raise ValueError(msg) tool_call.ensure_hash() @@ -52,7 +52,7 @@ def add_tool_call(item: MemoryItem, tool_call: ToolCallResult) -> None: set_tool_calls(item, tool_calls) -def get_tool_statistics(item: MemoryItem, recent_n: int = 20) -> dict[str, Any]: +def get_tool_statistics(item: Entry, recent_n: int = 20) -> dict[str, Any]: """Calculate statistics for the most recent N tool calls. Args: diff --git a/tests/test_backend_conformance.py b/tests/test_backend_conformance.py index 41ad2e90..97164a50 100644 --- a/tests/test_backend_conformance.py +++ b/tests/test_backend_conformance.py @@ -7,7 +7,7 @@ - ``clear_*`` with a ``where`` scope mutates the shared state in place (no rebinding that orphans the ``DatabaseState`` reference). - the SQLite read path preserves ``extra`` (reinforcement / ref_id / tool metadata). -- deleting an item / clearing memory leaves no orphan ``CategoryItem`` relations. +- deleting an entry / clearing memory leaves no orphan ``ResourceEntry`` relations. """ from __future__ import annotations @@ -27,6 +27,7 @@ MetadataStoreConfig, ) from memu.database.factory import build_database # noqa: E402 +from memu.memory_fs.exporter import slugify # noqa: E402 def _make_inmemory(): @@ -51,124 +52,127 @@ def store(request, tmp_path): db.close() -def _seed_item(store, *, summary: str, user_id: str, embedding=None): +def _seed_entry(store, *, text: str, user_id: str, embedding=None): res = store.resource_repo.create_resource( - url=f"mem://{summary}", + lane="source", + url=f"mem://{text}", modality="document", local_path="", - caption=summary, + summary=text, embedding=None, user_data={"user_id": user_id}, ) - item = store.memory_item_repo.create_item( - resource_id=res.id, - memory_type="knowledge", - summary=summary, + entry = store.entry_repo.create_entry( + lane="memory", + source_id=res.id, + entry_kind="knowledge", + text=text, embedding=embedding or [0.1, 0.2, 0.3], user_data={"user_id": user_id}, ) - return res, item + return res, entry + + +def _make_doc(store, *, title: str, user_id: str, description: str = "", embedding=None): + return store.resource_repo.get_or_create_doc( + lane="memory", + title=title, + description=description, + embedding=embedding or [0.1, 0.2, 0.3], + user_data={"user_id": user_id}, + slug=slugify(title), + ) def test_clear_items_with_scope_mutates_shared_state(store): """Clearing a scoped subset must not orphan the shared state reference.""" - _seed_item(store, summary="a", user_id="alice") - _seed_item(store, summary="b", user_id="bob") + _seed_entry(store, text="a", user_id="alice") + _seed_entry(store, text="b", user_id="bob") - deleted = store.memory_item_repo.clear_items({"user_id": "alice"}) + deleted = store.entry_repo.clear_entries({"user_id": "alice"}) assert len(deleted) == 1 - remaining = store.memory_item_repo.list_items() - summaries = {item.summary for item in remaining.values()} - assert summaries == {"b"} + remaining = store.entry_repo.list_entries() + texts = {entry.text for entry in remaining.values()} + assert texts == {"b"} def test_clear_categories_with_scope_mutates_shared_state(store): - store.memory_category_repo.get_or_create_category( - name="alpha", description="", embedding=[0.1, 0.2, 0.3], user_data={"user_id": "alice"} - ) - store.memory_category_repo.get_or_create_category( - name="beta", description="", embedding=[0.1, 0.2, 0.3], user_data={"user_id": "bob"} - ) + _make_doc(store, title="alpha", user_id="alice") + _make_doc(store, title="beta", user_id="bob") - store.memory_category_repo.clear_categories({"user_id": "alice"}) - remaining = store.memory_category_repo.list_categories() - names = {cat.name for cat in remaining.values()} - assert names == {"beta"} + store.resource_repo.clear_resources({"user_id": "alice"}, lane="memory") + remaining = store.resource_repo.list_resources(lane="memory") + titles = {res.title for res in remaining.values()} + assert titles == {"beta"} def test_unlink_item_removes_all_relations(store): - _res, item = _seed_item(store, summary="linked", user_id="alice") - cat1 = store.memory_category_repo.get_or_create_category( - name="c1", description="", embedding=[0.1, 0.2, 0.3], user_data={"user_id": "alice"} - ) - cat2 = store.memory_category_repo.get_or_create_category( - name="c2", description="", embedding=[0.1, 0.2, 0.3], user_data={"user_id": "alice"} - ) - store.category_item_repo.link_item_category(item.id, cat1.id, user_data={"user_id": "alice"}) - store.category_item_repo.link_item_category(item.id, cat2.id, user_data={"user_id": "alice"}) - assert len(store.category_item_repo.get_item_categories(item.id)) == 2 - - removed = store.category_item_repo.unlink_item(item.id) + _res, entry = _seed_entry(store, text="linked", user_id="alice") + cat1 = _make_doc(store, title="c1", user_id="alice") + cat2 = _make_doc(store, title="c2", user_id="alice") + store.resource_entry_repo.link_entry_resource(entry.id, cat1.id, user_data={"user_id": "alice"}) + store.resource_entry_repo.link_entry_resource(entry.id, cat2.id, user_data={"user_id": "alice"}) + assert len(store.resource_entry_repo.get_entry_resources(entry.id)) == 2 + + removed = store.resource_entry_repo.unlink_entry(entry.id) assert len(removed) == 2 - assert store.category_item_repo.get_item_categories(item.id) == [] - assert store.category_item_repo.list_relations() == [] + assert store.resource_entry_repo.get_entry_resources(entry.id) == [] + assert store.resource_entry_repo.list_relations() == [] def test_delete_item_leaves_no_orphan_relations(store): - """The Phase 0 delete fix: unlink relations before deleting the item.""" - _res, item = _seed_item(store, summary="doomed", user_id="alice") - cat = store.memory_category_repo.get_or_create_category( - name="c", description="", embedding=[0.1, 0.2, 0.3], user_data={"user_id": "alice"} - ) - store.category_item_repo.link_item_category(item.id, cat.id, user_data={"user_id": "alice"}) + """The Phase 0 delete fix: unlink relations before deleting the entry.""" + _res, entry = _seed_entry(store, text="doomed", user_id="alice") + cat = _make_doc(store, title="c", user_id="alice") + store.resource_entry_repo.link_entry_resource(entry.id, cat.id, user_data={"user_id": "alice"}) - store.category_item_repo.unlink_item(item.id) - store.memory_item_repo.delete_item(item.id) + store.resource_entry_repo.unlink_entry(entry.id) + store.entry_repo.delete_entry(entry.id) - assert store.memory_item_repo.get_item(item.id) is None - # No relation should point at the deleted item. - assert all(rel.item_id != item.id for rel in store.category_item_repo.list_relations()) + assert store.entry_repo.get_entry(entry.id) is None + # No relation should point at the deleted entry. + assert all(rel.entry_id != entry.id for rel in store.resource_entry_repo.list_relations()) def test_clear_relations_with_scope(store): - _res, item = _seed_item(store, summary="r", user_id="alice") - cat = store.memory_category_repo.get_or_create_category( - name="c", description="", embedding=[0.1, 0.2, 0.3], user_data={"user_id": "alice"} - ) - store.category_item_repo.link_item_category(item.id, cat.id, user_data={"user_id": "alice"}) + _res, entry = _seed_entry(store, text="r", user_id="alice") + cat = _make_doc(store, title="c", user_id="alice") + store.resource_entry_repo.link_entry_resource(entry.id, cat.id, user_data={"user_id": "alice"}) - removed = store.category_item_repo.clear_relations({"user_id": "alice"}) + removed = store.resource_entry_repo.clear_relations({"user_id": "alice"}) assert len(removed) == 1 - assert store.category_item_repo.list_relations({"user_id": "alice"}) == [] + assert store.resource_entry_repo.list_relations({"user_id": "alice"}) == [] def test_extra_round_trips_through_create_and_read(store): """``extra`` (tool metadata / ref_id / reinforcement) must survive a read.""" res = store.resource_repo.create_resource( + lane="source", url="mem://tool", modality="document", local_path="", - caption="tool", + summary="tool", embedding=None, user_data={"user_id": "alice"}, ) - item = store.memory_item_repo.create_item( - resource_id=res.id, - memory_type="tool", - summary="tool memory", + entry = store.entry_repo.create_entry( + lane="memory", + source_id=res.id, + entry_kind="tool", + text="tool memory", embedding=[0.1, 0.2, 0.3], user_data={"user_id": "alice"}, tool_record={"when_to_use": "always"}, ) - assert item.extra.get("when_to_use") == "always" + assert entry.extra.get("when_to_use") == "always" - fetched = store.memory_item_repo.get_item(item.id) + fetched = store.entry_repo.get_entry(entry.id) assert fetched is not None assert fetched.extra.get("when_to_use") == "always" -def _reconcile(crud_self, store, *, item_id, new_cat_names, mapped_old_cat_ids, name_to_id): +def _reconcile(crud_self, store, *, entry_id, new_cat_names, mapped_old_cat_ids, name_to_id): from types import SimpleNamespace from memu.app.crud import CRUDMixin @@ -176,7 +180,7 @@ def _reconcile(crud_self, store, *, item_id, new_cat_names, mapped_old_cat_ids, fake_self = SimpleNamespace(_map_category_names_to_ids=lambda names, ctx: [name_to_id[n] for n in names]) CRUDMixin._reconcile_update_categories( fake_self, # type: ignore[arg-type] - memory_id=item_id, + memory_id=entry_id, new_cat_names=new_cat_names, mapped_old_cat_ids=mapped_old_cat_ids, content_changed=False, @@ -192,65 +196,57 @@ def _reconcile(crud_self, store, *, item_id, new_cat_names, mapped_old_cat_ids, def test_update_with_none_categories_keeps_links(store): """P0 regression: omitting categories (None) must NOT drop existing links.""" - _res, item = _seed_item(store, summary="keep", user_id="alice") - cat = store.memory_category_repo.get_or_create_category( - name="A", description="", embedding=[0.1, 0.2, 0.3], user_data={"user_id": "alice"} - ) - store.category_item_repo.link_item_category(item.id, cat.id, user_data={"user_id": "alice"}) + _res, entry = _seed_entry(store, text="keep", user_id="alice") + cat = _make_doc(store, title="A", user_id="alice") + store.resource_entry_repo.link_entry_resource(entry.id, cat.id, user_data={"user_id": "alice"}) _reconcile( None, store, - item_id=item.id, + entry_id=entry.id, new_cat_names=None, mapped_old_cat_ids=[cat.id], name_to_id={"A": cat.id}, ) - linked = {rel.category_id for rel in store.category_item_repo.get_item_categories(item.id)} + linked = {rel.resource_id for rel in store.resource_entry_repo.get_entry_resources(entry.id)} assert linked == {cat.id} def test_update_with_empty_categories_clears_links(store): """An explicit empty list clears links (distinct from omitted/None).""" - _res, item = _seed_item(store, summary="clear", user_id="alice") - cat = store.memory_category_repo.get_or_create_category( - name="A", description="", embedding=[0.1, 0.2, 0.3], user_data={"user_id": "alice"} - ) - store.category_item_repo.link_item_category(item.id, cat.id, user_data={"user_id": "alice"}) + _res, entry = _seed_entry(store, text="clear", user_id="alice") + cat = _make_doc(store, title="A", user_id="alice") + store.resource_entry_repo.link_entry_resource(entry.id, cat.id, user_data={"user_id": "alice"}) _reconcile( None, store, - item_id=item.id, + entry_id=entry.id, new_cat_names=[], mapped_old_cat_ids=[cat.id], name_to_id={"A": cat.id}, ) - assert store.category_item_repo.get_item_categories(item.id) == [] + assert store.resource_entry_repo.get_entry_resources(entry.id) == [] def test_update_with_new_categories_swaps_links(store): - _res, item = _seed_item(store, summary="swap", user_id="alice") - cat_a = store.memory_category_repo.get_or_create_category( - name="A", description="", embedding=[0.1, 0.2, 0.3], user_data={"user_id": "alice"} - ) - cat_b = store.memory_category_repo.get_or_create_category( - name="B", description="", embedding=[0.1, 0.2, 0.3], user_data={"user_id": "alice"} - ) - store.category_item_repo.link_item_category(item.id, cat_a.id, user_data={"user_id": "alice"}) + _res, entry = _seed_entry(store, text="swap", user_id="alice") + cat_a = _make_doc(store, title="A", user_id="alice") + cat_b = _make_doc(store, title="B", user_id="alice") + store.resource_entry_repo.link_entry_resource(entry.id, cat_a.id, user_data={"user_id": "alice"}) _reconcile( None, store, - item_id=item.id, + entry_id=entry.id, new_cat_names=["B"], mapped_old_cat_ids=[cat_a.id], name_to_id={"A": cat_a.id, "B": cat_b.id}, ) - linked = {rel.category_id for rel in store.category_item_repo.get_item_categories(item.id)} + linked = {rel.resource_id for rel in store.resource_entry_repo.get_entry_resources(entry.id)} assert linked == {cat_b.id} @@ -284,8 +280,8 @@ def test_resolve_category_ids_creates_unknown_adaptively(store): ) # "Programming"/"programming" collapse (case-insensitive); "Cooking" is distinct. assert len(ids) == 2 - names = {c.name for c in store.memory_category_repo.list_categories().values()} - assert names == {"Programming", "Cooking"} + titles = {c.title for c in store.resource_repo.list_resources(lane="memory").values()} + assert titles == {"Programming", "Cooking"} # A subsequent call reuses the cached ids and creates nothing new. ids2 = asyncio.run( @@ -298,40 +294,42 @@ def test_resolve_category_ids_creates_unknown_adaptively(store): ) ) assert ids2 == [ctx.category_name_to_id["programming"]] - assert len(store.memory_category_repo.list_categories()) == 2 + assert len(store.resource_repo.list_resources(lane="memory")) == 2 def test_sqlite_extra_survives_cache_miss(tmp_path): """A fresh SQLite store (cold cache) must reconstruct ``extra`` from the DB.""" db, dsn = _make_sqlite(tmp_path) res = db.resource_repo.create_resource( + lane="source", url="mem://tool", modality="document", local_path="", - caption="tool", + summary="tool", embedding=None, user_data={"user_id": "alice"}, ) - item = db.memory_item_repo.create_item( - resource_id=res.id, - memory_type="tool", - summary="tool memory", + entry = db.entry_repo.create_entry( + lane="memory", + source_id=res.id, + entry_kind="tool", + text="tool memory", embedding=[0.1, 0.2, 0.3], user_data={"user_id": "alice"}, tool_record={"when_to_use": "cold-read"}, ) - item_id = item.id + entry_id = entry.id db.close() # Re-open the same DB file: caches are empty, so reads hit the DB read path. config = DatabaseConfig(metadata_store=MetadataStoreConfig(provider="sqlite", dsn=dsn)) db2 = build_database(config=config, user_model=DefaultUserModel) try: - fetched = db2.memory_item_repo.get_item(item_id) + fetched = db2.entry_repo.get_entry(entry_id) assert fetched is not None assert fetched.extra.get("when_to_use") == "cold-read" - listed = db2.memory_item_repo.list_items() - assert listed[item_id].extra.get("when_to_use") == "cold-read" + listed = db2.entry_repo.list_entries() + assert listed[entry_id].extra.get("when_to_use") == "cold-read" finally: db2.close() diff --git a/tests/test_folder_memorize.py b/tests/test_folder_memorize.py index a9651630..efc9f980 100644 --- a/tests/test_folder_memorize.py +++ b/tests/test_folder_memorize.py @@ -96,21 +96,23 @@ def _seed_resource_with_item(service: MemoryService, *, url: str, category_id: s """Create a resource + one item + one relation, returning the resource id.""" store = service.database res = store.resource_repo.create_resource( + lane="source", url=url, modality="document", local_path=url, - caption="cap", + summary="cap", embedding=None, user_data=dict(user), ) - item = store.memory_item_repo.create_item( - resource_id=res.id, - memory_type="profile", - summary=f"summary for {url}", + item = store.entry_repo.create_entry( + lane="memory", + source_id=res.id, + entry_kind="profile", + text=f"summary for {url}", embedding=[0.0], user_data=dict(user), ) - store.category_item_repo.link_item_category(item.id, category_id, dict(user)) + store.resource_entry_repo.link_entry_resource(item.id, category_id, dict(user)) return res.id @@ -138,9 +140,9 @@ async def _fake_patch(updates, *, ctx, store, llm_client=None) -> None: assert keep_id in remaining assert drop_id not in remaining # The dropped resource's item and relation are gone; the kept one's survive. - items = store.memory_item_repo.list_items(where=user) - assert all(it.resource_id == keep_id for it in items.values()) - relations = store.category_item_repo.list_relations(where=user) + items = store.entry_repo.list_entries(where=user) + assert all(it.source_id == keep_id for it in items.values()) + relations = store.resource_entry_repo.list_relations(where=user) assert len(relations) == 1 # Discarded content was fed to the summary recompute as (before, None). assert patched and category_id in patched[0] @@ -161,21 +163,23 @@ async def _noop_patch(updates, *, ctx, store, llm_client=None) -> None: async def _fake_memorize_one(*, resource_url, modality, user_scope, ctx, store) -> dict[str, Any]: res = store.resource_repo.create_resource( + lane="source", url=resource_url, modality=modality, local_path=resource_url, - caption="cap", + summary="cap", embedding=None, user_data=dict(user_scope or {}), ) - store.memory_item_repo.create_item( - resource_id=res.id, - memory_type="profile", - summary=f"summary {resource_url}", + store.entry_repo.create_entry( + lane="memory", + source_id=res.id, + entry_kind="profile", + text=f"summary {resource_url}", embedding=[0.0], user_data=dict(user_scope or {}), ) - return {"resources": [res], "response": {"items": [{"summary": "x"}]}} + return {"resources": [res], "response": {"items": [{"text": "x"}]}} monkeypatch.setattr(service, "_ensure_categories_ready", _noop_categories) monkeypatch.setattr(service, "_patch_category_summaries", _noop_patch) @@ -227,10 +231,11 @@ async def _noop_categories(*a, **k) -> None: async def _fake_memorize_one(*, resource_url, modality, user_scope, ctx, store) -> dict[str, Any]: res = store.resource_repo.create_resource( + lane="source", url=resource_url, modality=modality, local_path=resource_url, - caption="cap", + summary="cap", embedding=None, user_data=dict(user_scope or {}), ) @@ -269,10 +274,11 @@ async def _noop_categories(*a, **k) -> None: async def _fake_memorize_one(*, resource_url, modality, user_scope, ctx, store) -> dict[str, Any]: res = store.resource_repo.create_resource( + lane="source", url=resource_url, modality=modality, local_path=resource_url, - caption="cap", + summary="cap", embedding=None, user_data=dict(user_scope or {}), ) diff --git a/tests/test_memory_files.py b/tests/test_memory_files.py index 3679ebfb..b64683b3 100644 --- a/tests/test_memory_files.py +++ b/tests/test_memory_files.py @@ -7,7 +7,7 @@ from memu.app import MemoryService from memu.memory_fs import MemoryFileExporter -from memu.memory_fs.exporter import MANIFEST_NAME +from memu.memory_fs.exporter import MANIFEST_NAME, slugify # With synthesize=False (the default) the whole tree is LLM-free: MEMORY.md is # rendered from category summaries and the skill/ tree from the deterministic @@ -26,28 +26,32 @@ def _build_service(output_dir: Path) -> MemoryService: def _seed(service: MemoryService, *, user: dict[str, str]) -> dict[str, str]: store = service.database resource = store.resource_repo.create_resource( + lane="source", url="docs/coffee.txt", modality="document", local_path="coffee.txt", - caption="Notes about the user's coffee preferences.", + summary="Notes about the user's coffee preferences.", embedding=None, user_data=dict(user), ) - category = store.memory_category_repo.get_or_create_category( - name="Preferences", + category = store.resource_repo.get_or_create_doc( + lane="memory", + title="Preferences", description="User preferences, likes and dislikes", embedding=[0.1, 0.2], user_data=dict(user), + slug=slugify("Preferences"), ) - store.memory_category_repo.update_category(category_id=category.id, summary="The user likes pour-over coffee.") - skill = store.memory_item_repo.create_item( - resource_id=resource.id, - memory_type="skill", - summary=_SKILL_BODY, + store.resource_repo.update_resource(resource_id=category.id, summary="The user likes pour-over coffee.") + skill = store.entry_repo.create_entry( + lane="memory", + source_id=resource.id, + entry_kind="skill", + text=_SKILL_BODY, embedding=[0.1, 0.2], user_data=dict(user), ) - store.category_item_repo.link_item_category(skill.id, category.id, user_data=dict(user)) + store.resource_entry_repo.link_entry_resource(skill.id, category.id, user_data=dict(user)) return {"category_id": category.id, "resource_id": resource.id, "skill_id": skill.id} @@ -100,8 +104,8 @@ async def test_export_is_idempotent_until_data_changes(tmp_path: Path) -> None: # Changing only a folder summary rewrites that category's memory/.md but # not MEMORY.md (an overview of links) nor INDEX.md (a file TOC). - service.database.memory_category_repo.update_category( - category_id=ids["category_id"], + service.database.resource_repo.update_resource( + resource_id=ids["category_id"], summary="The user now prefers espresso.", ) third = await service.export_memory_files(user={"user_id": "u1"}) @@ -119,7 +123,7 @@ async def test_export_removes_stale_skill_and_prunes_dirs(tmp_path: Path) -> Non assert (tmp_path / "skill" / "pour-over" / "SKILL.md").exists() # Dropping the skill-type item removes it from the bypass, so its doc goes stale. - service.database.memory_item_repo.clear_items(where={"user_id": "u1"}) + service.database.entry_repo.clear_entries(where={"user_id": "u1"}) result = await service.export_memory_files(user={"user_id": "u1"}) assert "skill/pour-over/SKILL.md" in result["removed"] @@ -129,11 +133,13 @@ async def test_export_removes_stale_skill_and_prunes_dirs(tmp_path: Path) -> Non async def test_export_respects_user_scope(tmp_path: Path) -> None: service = _build_service(tmp_path) _seed(service, user={"user_id": "u1"}) - service.database.memory_category_repo.get_or_create_category( - name="Secret", + service.database.resource_repo.get_or_create_doc( + lane="memory", + title="Secret", description="Other user's folder", embedding=[0.3, 0.4], user_data={"user_id": "u2"}, + slug=slugify("Secret"), ) await service.export_memory_files(user={"user_id": "u1"}) diff --git a/tests/test_memory_fs_synthesis.py b/tests/test_memory_fs_synthesis.py index 4c210271..6e6d813e 100644 --- a/tests/test_memory_fs_synthesis.py +++ b/tests/test_memory_fs_synthesis.py @@ -68,26 +68,26 @@ def test_synthesizer_helpers() -> None: def test_build_synthesis_descriptions_uses_structured_items() -> None: - """Synthesis input is sourced from extracted items, with a caption fallback.""" - from memu.database.models import MemoryItem, Resource + """Synthesis input is sourced from extracted entries, with a summary fallback.""" + from memu.database.models import Entry, Resource res_with_items = Resource( - id="r1", url="docs/a.txt", modality="document", local_path="a.txt", caption="raw caption a" + id="r1", lane="source", url="docs/a.txt", modality="document", local_path="a.txt", summary="raw caption a" ) res_without_items = Resource( - id="r2", url="docs/b.txt", modality="document", local_path="b.txt", caption="raw caption b" + id="r2", lane="source", url="docs/b.txt", modality="document", local_path="b.txt", summary="raw caption b" ) items = [ - MemoryItem(id="i1", resource_id="r1", memory_type="knowledge", summary="Alpha fact."), - MemoryItem(id="i2", resource_id="r1", memory_type="profile", summary="Beta trait."), + Entry(id="i1", lane="memory", source_id="r1", entry_kind="knowledge", text="Alpha fact."), + Entry(id="i2", lane="memory", source_id="r1", entry_kind="profile", text="Beta trait."), ] descriptions = MemoryFileExporter.build_synthesis_descriptions([res_with_items, res_without_items], items) by_url = {d.url: d.description for d in descriptions} - # r1 is composed from its structured items, not the caption. + # r1 is composed from its structured entries, not the summary. assert by_url["docs/a.txt"] == "[knowledge] Alpha fact.; [profile] Beta trait." - # r2 has no items, so it falls back to the caption. + # r2 has no entries, so it falls back to the summary. assert by_url["docs/b.txt"] == "raw caption b" @@ -120,10 +120,11 @@ async def test_service_synthesis_wiring(tmp_path: Path, monkeypatch) -> None: memory_files_config={"enabled": True, "output_dir": str(tmp_path), "synthesize": True}, ) service.database.resource_repo.create_resource( + lane="source", url="docs/coffee.txt", modality="document", local_path="coffee.txt", - caption="The user likes pour-over coffee.", + summary="The user likes pour-over coffee.", embedding=None, user_data={"user_id": "u1"}, ) @@ -235,10 +236,11 @@ async def test_service_init_then_update(tmp_path: Path, monkeypatch) -> None: repo = service.database.resource_repo repo.create_resource( + lane="source", url="docs/coffee.txt", modality="document", local_path="coffee.txt", - caption="The user likes pour-over coffee.", + summary="The user likes pour-over coffee.", embedding=None, user_data={"user_id": "u1"}, ) @@ -250,10 +252,11 @@ async def test_service_init_then_update(tmp_path: Path, monkeypatch) -> None: # Second pass: tree exists -> incremental update from the changed resource only. changed = repo.create_resource( + lane="source", url="docs/latte.txt", modality="document", local_path="latte.txt", - caption="The user enjoys latte art and oat milk.", + summary="The user enjoys latte art and oat milk.", embedding=None, user_data={"user_id": "u1"}, ) diff --git a/tests/test_openrouter.py b/tests/test_openrouter.py index ba4b47c4..51b0839e 100644 --- a/tests/test_openrouter.py +++ b/tests/test_openrouter.py @@ -30,7 +30,7 @@ def _print_categories(categories, max_items=3): print(" Categories:") for cat in categories[:max_items]: summary = cat.get("summary") or cat.get("description", "") - print(f" - {cat.get('name')}: {summary[:60]}...") + print(f" - {cat.get('title')}: {summary[:60]}...") def _print_items(items, max_items=3): @@ -38,9 +38,9 @@ def _print_items(items, max_items=3): if items: print(" Items:") for item in items[:max_items]: - memory_type = item.get("memory_type", "unknown") - summary = item.get("summary", "")[:80] - print(f" - [{memory_type}] {summary}...") + entry_kind = item.get("entry_kind", "unknown") + summary = item.get("text", "")[:80] + print(f" - [{entry_kind}] {summary}...") async def _test_memorize(service, file_path, output_data): diff --git a/tests/test_sqlite.py b/tests/test_sqlite.py index 3031c56b..13873e85 100644 --- a/tests/test_sqlite.py +++ b/tests/test_sqlite.py @@ -10,10 +10,10 @@ def _print_results(title: str, result: dict) -> None: print(f"\n[SQLITE] RETRIEVED - {title}") print(" Categories:") for cat in result.get("categories", [])[:3]: - print(f" - {cat.get('name')}: {(cat.get('summary') or cat.get('description', ''))[:80]}...") + print(f" - {cat.get('title')}: {(cat.get('summary') or cat.get('description', ''))[:80]}...") print(" Items:") for item in result.get("items", [])[:3]: - print(f" - [{item.get('memory_type')}] {item.get('summary', '')[:100]}...") + print(f" - [{item.get('entry_kind')}] {item.get('text', '')[:100]}...") if result.get("resources"): print(" Resources:") for res in result.get("resources", [])[:3]: @@ -54,7 +54,7 @@ async def main(): print("\n[SQLITE] Memorizing...") memory = await service.memorize(resource_url=file_path, modality="conversation", user={"user_id": "123"}) for cat in memory.get("categories", []): - print(f" - {cat.get('name')}: {(cat.get('summary') or '')[:80]}...") + print(f" - {cat.get('title')}: {(cat.get('summary') or '')[:80]}...") queries = [ {"role": "user", "content": {"text": "Tell me about preferences"}}, diff --git a/tests/test_tool_memory.py b/tests/test_tool_memory.py index 733a6229..25dd7172 100644 --- a/tests/test_tool_memory.py +++ b/tests/test_tool_memory.py @@ -1,4 +1,4 @@ -"""Tests for Tool Memory feature - specialized memory type for tracking tool usage.""" +"""Tests for Tool Memory feature - specialized entry kind for tracking tool usage.""" from __future__ import annotations @@ -31,9 +31,9 @@ "ToolCallResult": models.ToolCallResult, } models.ToolCallResult.model_rebuild(_types_namespace=rebuild_ns) -models.MemoryItem.model_rebuild(_types_namespace=rebuild_ns) +models.Entry.model_rebuild(_types_namespace=rebuild_ns) -MemoryItem = models.MemoryItem +Entry = models.Entry MemoryType = models.MemoryType ToolCallResult = models.ToolCallResult @@ -125,7 +125,7 @@ def test_string_input(self): class TestMemoryItemToolType: - """Tests for MemoryItem with tool type.""" + """Tests for Entry with tool kind.""" def test_tool_memory_type_literal(self): """Test that 'tool' is a valid MemoryType.""" @@ -135,27 +135,29 @@ def test_tool_memory_type_literal(self): assert "tool" in valid_types def test_create_tool_memory(self): - """Test creating a tool type memory item with tool fields in extra.""" - item = MemoryItem( - resource_id=None, - memory_type="tool", - summary="file_reader tool usage for config files", + """Test creating a tool kind entry with tool fields in extra.""" + item = Entry( + lane="memory", + source_id=None, + entry_kind="tool", + text="file_reader tool usage for config files", extra={ "when_to_use": "When needing to read configuration files", "metadata": {"tool_name": "file_reader", "avg_success_rate": 0.95}, }, ) - assert item.memory_type == "tool" + assert item.entry_kind == "tool" assert item.extra["when_to_use"] == "When needing to read configuration files" assert item.extra["metadata"]["tool_name"] == "file_reader" def test_add_tool_call(self): """Test adding tool call results to a tool memory.""" - item = MemoryItem( - resource_id=None, - memory_type="tool", - summary="calculator tool usage", + item = Entry( + lane="memory", + source_id=None, + entry_kind="tool", + text="calculator tool usage", ) tool_call = ToolCallResult( @@ -176,10 +178,11 @@ def test_add_tool_call(self): def test_add_tool_call_wrong_type(self): """Test that add_tool_call fails for non-tool memories.""" - item = MemoryItem( - resource_id=None, - memory_type="profile", - summary="User profile info", + item = Entry( + lane="memory", + source_id=None, + entry_kind="profile", + text="User profile info", ) tool_call = ToolCallResult( @@ -193,10 +196,11 @@ def test_add_tool_call_wrong_type(self): def test_get_tool_statistics_empty(self): """Test statistics for memory with no tool calls.""" - item = MemoryItem( - resource_id=None, - memory_type="tool", - summary="empty tool memory", + item = Entry( + lane="memory", + source_id=None, + entry_kind="tool", + text="empty tool memory", ) stats = get_tool_statistics(item) @@ -211,10 +215,11 @@ def test_get_tool_statistics_empty(self): def test_get_tool_statistics(self): """Test statistics calculation for tool calls.""" # Tool calls are stored as dicts in extra - item = MemoryItem( - resource_id=None, - memory_type="tool", - summary="calculator tool", + item = Entry( + lane="memory", + source_id=None, + entry_kind="tool", + text="calculator tool", extra={ "tool_calls": [ { @@ -259,10 +264,11 @@ def test_get_tool_statistics(self): def test_get_tool_statistics_recent_n(self): """Test statistics with recent_n limit.""" - item = MemoryItem( - resource_id=None, - memory_type="tool", - summary="tool with many calls", + item = Entry( + lane="memory", + source_id=None, + entry_kind="tool", + text="tool with many calls", extra={ "tool_calls": [ {"tool_name": "t", "input": "1", "output": "1", "success": False, "time_cost": 1.0, "score": 0.0}, @@ -285,10 +291,11 @@ class TestMemoryItemNewFields: def test_when_to_use_field(self): """Test when_to_use field stored in extra for retrieval hints.""" - item = MemoryItem( - resource_id=None, - memory_type="profile", - summary="User prefers dark mode", + item = Entry( + lane="memory", + source_id=None, + entry_kind="profile", + text="User prefers dark mode", extra={"when_to_use": "When configuring UI settings or themes"}, ) @@ -296,10 +303,11 @@ def test_when_to_use_field(self): def test_metadata_field(self): """Test metadata field stored in extra for type-specific data.""" - item = MemoryItem( - resource_id=None, - memory_type="event", - summary="User attended conference", + item = Entry( + lane="memory", + source_id=None, + entry_kind="event", + text="User attended conference", extra={ "metadata": { "event_date": "2026-01-15", @@ -316,10 +324,11 @@ def test_metadata_field(self): def test_default_values(self): """Test that extra defaults to empty dict.""" - item = MemoryItem( - resource_id=None, - memory_type="knowledge", - summary="Python is a programming language", + item = Entry( + lane="memory", + source_id=None, + entry_kind="knowledge", + text="Python is a programming language", ) assert item.extra.get("when_to_use") is None From cadbe605bd11079a5d2c4c4a0cfe58d48ba077d5 Mon Sep 17 00:00:00 2001 From: sairin1202 <952141617@qq.com> Date: Thu, 25 Jun 2026 06:11:34 +0800 Subject: [PATCH 2/2] feat(app): integrate index & skill lanes into memorize/retrieve Wire the index and skill lanes (previously model-only) into the full memorize/retrieve pipeline on the unified Resource + Entry backbone, and give each lane its own entry types instead of sharing the memory set. - models: add "skill" to Lane and RETRIEVAL_LANES (index/memory/skill) - prompts: add per-lane entry types (index=description, skill=tool/log) and a LANE_ENTRY_TYPES table - settings: clean break to per-lane LaneConfig (enabled/grouping/entry_types/ summary), enable index/memory/skill by default - memorize: loop over enabled lanes; support per_resource (index, 1:1 docs) and adaptive (memory/skill, grouped + summarized) grouping - retrieve: single-pass per-lane recall for both rag and llm methods, returning {lanes: {index, memory, skill}, resources} with backward-compatible top-level categories/items mirroring the memory lane - docs: ADR 0006 -> Accepted; architecture.md updated for three lanes - tests: per-lane persistence + per-lane recall conformance across backends Co-authored-by: Cursor --- ...06-unified-resource-entry-lane-backbone.md | 71 ++- docs/architecture.md | 47 +- docs/tutorials/getting_started.md | 8 +- examples/example_2_skill_extraction.py | 274 -------- examples/example_5_with_lazyllm_client.py | 84 +-- examples/getting_started_robust.py | 6 +- examples/proactive/memory/config.py | 4 +- src/memu/app/crud.py | 42 +- src/memu/app/memorize.py | 495 +++++++++------ src/memu/app/memory_files.py | 13 +- src/memu/app/retrieve.py | 583 +++++------------- src/memu/app/service.py | 39 +- src/memu/app/settings.py | 136 ++-- .../inmemory/repositories/entry_repo.py | 59 +- .../repositories/resource_entry_repo.py | 4 +- src/memu/database/models.py | 55 +- src/memu/database/postgres/models.py | 2 +- .../postgres/repositories/entry_repo.py | 58 +- .../repositories/resource_entry_repo.py | 8 +- src/memu/database/repositories/entry.py | 14 +- src/memu/database/sqlite/models.py | 4 +- .../sqlite/repositories/entry_repo.py | 53 +- .../sqlite/repositories/resource_repo.py | 2 +- src/memu/memory_fs/__init__.py | 3 +- src/memu/memory_fs/exporter.py | 143 +---- src/memu/memory_fs/synthesizer.py | 128 +--- src/memu/prompts/__init__.py | 8 +- .../{memory_type => entry_type}/__init__.py | 25 +- src/memu/prompts/entry_type/description.py | 83 +++ .../{memory_type => entry_type}/event.py | 0 .../{memory_type => entry_type}/knowledge.py | 0 src/memu/prompts/entry_type/log.py | 122 ++++ .../{memory_type => entry_type}/profile.py | 0 src/memu/prompts/entry_type/tool.py | 124 ++++ src/memu/prompts/memory_fs/__init__.py | 23 - src/memu/prompts/memory_type/behavior.py | 132 ---- src/memu/prompts/memory_type/skill.py | 229 ------- src/memu/prompts/memory_type/tool.py | 120 ---- src/memu/utils/tool.py | 102 --- tests/test_backend_conformance.py | 188 +++++- tests/test_folder_memorize.py | 10 +- tests/test_inmemory.py | 4 +- tests/test_memory_files.py | 44 +- tests/test_memory_fs_synthesis.py | 106 +--- tests/test_openrouter.py | 4 +- tests/test_postgres.py | 4 +- tests/test_salience.py | 4 +- tests/test_sqlite.py | 2 +- tests/test_tool_memory.py | 336 ---------- 49 files changed, 1305 insertions(+), 2700 deletions(-) delete mode 100644 examples/example_2_skill_extraction.py rename src/memu/prompts/{memory_type => entry_type}/__init__.py (54%) create mode 100644 src/memu/prompts/entry_type/description.py rename src/memu/prompts/{memory_type => entry_type}/event.py (100%) rename src/memu/prompts/{memory_type => entry_type}/knowledge.py (100%) create mode 100644 src/memu/prompts/entry_type/log.py rename src/memu/prompts/{memory_type => entry_type}/profile.py (100%) create mode 100644 src/memu/prompts/entry_type/tool.py delete mode 100644 src/memu/prompts/memory_type/behavior.py delete mode 100644 src/memu/prompts/memory_type/skill.py delete mode 100644 src/memu/prompts/memory_type/tool.py delete mode 100644 src/memu/utils/tool.py delete mode 100644 tests/test_tool_memory.py diff --git a/docs/adr/0006-unified-resource-entry-lane-backbone.md b/docs/adr/0006-unified-resource-entry-lane-backbone.md index e7cd87d9..5a377710 100644 --- a/docs/adr/0006-unified-resource-entry-lane-backbone.md +++ b/docs/adr/0006-unified-resource-entry-lane-backbone.md @@ -1,6 +1,6 @@ # ADR 0006: Unify INDEX / MEMORY / SKILL onto a Resource + Entry Lane Backbone -- Status: Proposed +- Status: Accepted - Date: 2026-06-25 ## Context @@ -8,18 +8,17 @@ memU historically modeled structured memory as four record types — `Resource`, `MemoryItem`, `MemoryCategory`, `CategoryItem` — with retrieval running a fixed `category -> item -> resource` waterfall. Separately, the read-only `memory_fs` -exporter projected three markdown trees (`INDEX.md`, `MEMORY.md`, `SKILL.md`) -that were decoupled from retrieval, and skills were handled by an ad-hoc dual -track (`memory_type="skill"` items *or* LLM synthesis). +exporter projected markdown trees that were decoupled from retrieval. -This produced three asymmetric concepts: +This produced asymmetric concepts: - INDEX: `Resource.caption` + verbatim `resource/` copies - MEMORY: `MemoryCategory.summary` + `memory/.md` -- SKILL: synthesized or bypassed, not part of retrieval We want INDEX, MEMORY, and SKILL to share **one backbone** with **consistent storage and retrieval**, all derived from the same per-resource canonical text. +Each lane is the same processing track; lanes differ only in their entry-type +set, extraction prompts, and how entries are grouped into coarse lane docs. ## Decision @@ -27,10 +26,23 @@ Collapse the model to **two first-class, lane-tagged entities plus one edge**. ### Lane -A `lane` discriminator with three values: `index`, `memory`, `skill`. (Raw -inputs use `lane="source"`.) The three lanes are parallel, structurally -identical processing tracks over a shared trunk; they differ only in *what the -extractor pulls out* and *the entry→resource grouping cardinality*. +A `lane` discriminator: `index`, `memory`, and `skill`. (Raw inputs use +`lane="source"`.) The lanes are parallel, structurally identical processing +tracks over a shared trunk; they differ only in *what the extractor pulls out* +(per-lane `entry_type` set and prompts) and *the entry→resource grouping +cardinality*: + +- `index` — `entry_type` ∈ {`description`}, **`per_resource`** grouping: one + coarse description doc per source resource (1:1, no LLM grouping). +- `memory` — `entry_type` ∈ {`profile`, `event`, `knowledge`}, **`adaptive`** + grouping: the extractor proposes group names and a summarized category doc is + synthesized per group. +- `skill` — `entry_type` ∈ {`tool`, `log`}, **`adaptive`** grouping: entries are + grouped into summarized skill docs (analogous to memory categories). + +Per-lane behavior is configured via `MemorizeConfig.lanes` (a `dict[str, +LaneConfig]`); all three lanes are enabled by default. Each adaptive lane has its +own summary prompt / target length / LLM profile. ### Entities @@ -38,8 +50,8 @@ extractor pulls out* and *the entry→resource grouping cardinality*. - Raw source artifacts (`lane="source"`, `modality` = video/image/audio/ conversation/document); multimodal preprocessing fills `content` (the canonical, modality-agnostic text — the shared trunk). - - Generated coarse docs (`lane` ∈ {index, memory, skill}, `modality="markdown"`), - each rendered as a file under the `resource/` root: + - Generated coarse docs (`lane` ∈ {index, memory, skill}, + `modality="markdown"`), each rendered as a file under the `resource/` root: - `resource/index/.md` — a description page linking to a raw resource - `resource/memory/.md` — a category page - `resource/skill/.md` — a skill page @@ -49,16 +61,15 @@ extractor pulls out* and *the entry→resource grouping cardinality*. `lane="memory"` markdown resource). 2. **`Entry`** (lane-tagged, one physical table — the searchable atom): - - index → a resource description; memory → a memory item; skill → a reusable - operation step. - - Carries `text`, `embedding`, `entry_kind` (memory sub-type / step kind), - `extra`, and `source_path` — a back-link to the originating raw resource, - relative to the `resource/` root. (This **generalizes the former - `MemoryItem`**.) + - index → a resource description; memory → a memory item; skill → a tool/log. + - Carries `text`, `embedding`, `entry_type` (the per-lane sub-type, which + selects the extraction prompt), `extra`, and `source_path` — a back-link to + the originating raw resource, relative to the `resource/` root. (This + **generalizes the former `MemoryItem`**.) 3. **`ResourceEntry`** (edge): membership of an `Entry` in its coarse lane - `Resource` (memory item ∈ category page, skill step ∈ skill page, description - ∈ index page). Many-to-many. (This **generalizes the former `CategoryItem`**.) + `Resource` (memory item ∈ category page, description ∈ index page, tool/log ∈ + skill page). Many-to-many. (This **generalizes the former `CategoryItem`**.) ### Links / provenance @@ -72,12 +83,16 @@ extractor pulls out* and *the entry→resource grouping cardinality*. ### Pipelines - **memorize**: `ingest -> preprocess_multimodal (-> Resource.content) -> - extract_lanes (index/memory/skill extractors) -> embed_entries -> - persist lane resources -> build_response`. -- **retrieve**: for each enabled lane, `Resource` recall (stored embedding, - `where lane=`) → `Entry` recall (stored embedding, `where lane=`), returning a - per-lane shape `{index: {...}, memory: {...}, skill: {...}, resources: [...]}`. - All lanes traverse the same code path; only the `lane` filter differs. + extract_lanes (per enabled lane: index/memory/skill extractors) -> + embed_entries -> persist lane resources (per_resource 1:1 doc or adaptive + grouped+summarized docs) -> build_response`. +- **retrieve**: a single `route_intention` pass, then for each enabled lane, + `Resource` recall (stored embedding, `where lane=`) → `Entry` recall (stored + embedding, `where lane=`), plus a `source`-lane resource recall, returning a + per-lane shape `{lanes: {index, memory, skill}, resources: [...]}` (with + backward-compatible top-level `categories`/`items` mirroring the memory lane). + All lanes traverse the same code path; only the `lane` filter differs. Both + `rag` and `llm` ranking methods are supported. ### Naming @@ -90,10 +105,8 @@ a document. Positive: -- One storage schema and one retrieval path for all three lanes (true +- One storage schema and one retrieval path for all lanes (true storage/retrieval consistency). -- Skills become first-class and searchable; the dual-track synthesis/bypass goes - away. - Every entry and coarse resource is traceable back to its raw source. - "Everything is a resource" keeps the mental model and the on-disk tree aligned. diff --git a/docs/architecture.md b/docs/architecture.md index 98048a29..75944058 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -16,15 +16,20 @@ plus one edge: (`lane="source"`, modality conversation/document/image/video/audio; the canonical preprocessing text lives in `content`) or a generated markdown doc (`lane` ∈ {index, memory, skill}, `modality="markdown"`). A memory-lane - `Resource` is the former "category"; a skill-lane one is a skill page. + `Resource` is the former "category". - `Entry`: a `lane`-tagged searchable atom with an embedding (index → a - description, memory → a memory item, skill → a step). Links to its origin via - `source_id`/`source_path`. + description, memory → a memory item, skill → a tool/log). Its per-lane + `entry_type` (memory: profile/event/knowledge, skill: tool/log, index: + description) selects the extraction prompt during memorize. Links to its origin + via `source_id`/`source_path`. - `ResourceEntry`: membership edge from an `Entry` to its coarse (lane) `Resource`. -The three retrieval lanes (index/memory/skill) are parallel, structurally -identical tracks over the shared `Resource.content` trunk; retrieval is the same -`Resource recall → Entry recall` waterfall in every lane, parameterized by `lane`. +The retrieval lanes (index/memory/skill) are parallel, structurally identical +tracks over the shared `Resource.content` trunk; retrieval runs the same +`Resource recall → Entry recall` per lane, parameterized by `lane`. Per-lane +extraction/grouping is configured via `MemorizeConfig.lanes`; `index` uses +`per_resource` grouping (1:1 description docs) while `memory`/`skill` use +`adaptive` grouping (LLM-grouped, summarized docs). At runtime, `MemoryService` orchestrates ingestion, retrieval, and manual CRUD over these layers. @@ -248,13 +253,10 @@ payload directory: / ├── INDEX.md ← index of the raw files under resource/ ├── MEMORY.md ← overview + index of memory/ -├── SKILL.md ← index of the skills under skill/ ├── resource/ │ └── ← one copied raw source file (verbatim bytes) -├── memory/ -│ └── .md ← one memory-lane Resource (description + summary) -└── skill/ - └── /SKILL.md ← one synthesized skill per folder +└── memory/ + └── .md ← one memory-lane Resource (description + summary) ``` - `resource/` holds the raw source files copied verbatim out of the blob store @@ -262,22 +264,14 @@ payload directory: link), so an agent knows which raw resources exist. - `memory/.md` is the living memory split one file per memory-lane `Resource` (its description + summary); `MEMORY.md` is an overview that links to each one. -- `skill//SKILL.md` is a reusable skill synthesized from the descriptions - (a sibling of `MEMORY.md`, never derived from extracted skill-type memory - items); the root `SKILL.md` indexes the tree. ### Synthesis mode -The `skill/` tree is always synthesized from the per-source descriptions by an LLM -(`memu.memory_fs.MemorySynthesizer`, prompts in `memu.prompts.memory_fs`) — one -pass extracts skills as a JSON array of `{name, body}` objects, each written as its -own `skill//SKILL.md` doc. It is never derived from extracted skill-type -memory items. - `MEMORY.md` is rendered deterministically by default: an overview that links to the per-category `memory/.md` files (themselves deterministic from category description + summary). When `memory_files_config.synthesize=True`, the `MEMORY.md` -body is instead synthesized from all descriptions in one LLM pass. +body is instead synthesized from all descriptions in one LLM pass +(`memu.memory_fs.MemorySynthesizer`, prompts in `memu.prompts.memory_fs`). `INDEX.md`, the `resource/` copies, and `memory/.md` stay deterministic in both modes. Synthesis uses the `synthesis_llm_profile` profile and leaves the @@ -291,13 +285,10 @@ model. `MemoryFilesBuilder.build(database, where, changed=...)` (delegated to fr - **Initialization** (no prior tree on disk, or `changed is None`): scan all in-scope sources, turn each into its multimodal description, and synthesize the - `skill/` tree (and, when `synthesize=True`, the `MEMORY.md` body) from scratch - (`MemorySynthesizer.synthesize` / `synthesize_skills`). + `MEMORY.md` body from scratch (`MemorySynthesizer.synthesize`). - **Incremental update** (a tree already exists and a changed set is supplied): - read the existing skill bodies (and `MEMORY.md` body) back off disk and merge - only the changed sources' descriptions into them (`MemorySynthesizer.update` / - `update_skills`, prompts `MEMORY_UPDATE_PROMPT` / `SKILL_UPDATE_PROMPT`). Skills - are upserted by slug, so untouched skills survive. + read the existing `MEMORY.md` body back off disk and merge only the changed + sources' descriptions into it. `INDEX.md`, `resource/`, and `memory/` are always recomputed from the current store, so they need no LLM merge. `export_memory_files(user=...)` always takes the @@ -311,7 +302,7 @@ does **not** drive the exporter; it is left entirely untouched. The exporter is read-only against the database and disabled by default (`memory_files_config.enabled`). Diff detection is handled by a sidecar manifest (`.memufs_manifest.json`) that stores per-file content hashes, so each export -only rewrites artifacts whose rendered content changed (and prunes stale skill +only rewrites artifacts whose rendered content changed (and prunes stale files/dirs) — no database schema change is required. Rendered content avoids volatile values so an unchanged store re-exports as a no-op. Exports are serialized through a per-service lock. diff --git a/docs/tutorials/getting_started.md b/docs/tutorials/getting_started.md index f198664a..7bce7a5d 100644 --- a/docs/tutorials/getting_started.md +++ b/docs/tutorials/getting_started.md @@ -113,9 +113,9 @@ async def main() -> None: memory_content = "The user is a senior Python architect who loves clean code and type hints." # We use 'create_memory_item' to insert a single memory record. - # memory_type='profile' indicates this is an attribute of the user. + # entry_type='profile' indicates this is an attribute of the user. result = await service.create_memory_item( - memory_type="profile", + entry_type="profile", memory_content=memory_content, memory_categories=["User Facts"], ) @@ -135,7 +135,7 @@ async def main() -> None: if items: print(f"[OK] Found {len(items)} relevant memory item(s):") for idx, item in enumerate(items, 1): - print(f" {idx}. {item.get('summary')} (Type: {item.get('memory_type')})") + print(f" {idx}. {item.get('text')} (Type: {item.get('entry_type')})") else: print("[!] No relevant memories found.") @@ -153,7 +153,7 @@ if __name__ == "__main__": ### Understanding the Code 1. **Initialization**: We configure `MemoryService` with specific `llm_profiles`. This tells MemU which model to use. We also define a `memorize_config` with a "User Facts" category. Categories help the LLM organize and retrieve information more effectively. -2. **Memory Injection**: `create_memory_item` is used to explicitly add a piece of knowledge. We tag it with `memory_type="profile"` to semantically indicate this is a user attribute. +2. **Memory Injection**: `create_memory_item` is used to explicitly add a piece of knowledge. We tag it with `entry_type="profile"` to semantically indicate this is a user attribute. 3. **Retrieval**: We use `retrieve` with a natural language query. MemU's internal workflow ("RAG" or "LLM" based) will determine the best way to find relevant memories. ## Troubleshooting diff --git a/examples/example_2_skill_extraction.py b/examples/example_2_skill_extraction.py deleted file mode 100644 index 3ca75804..00000000 --- a/examples/example_2_skill_extraction.py +++ /dev/null @@ -1,274 +0,0 @@ -""" -Example 2: Workflow & Agent Logs -> Skill Extraction - -This example demonstrates how to extract skills from workflow descriptions -and agent runtime logs, then output them to a Markdown file. - -Usage: - export OPENAI_API_KEY=your_api_key - python examples/example_2_skill_extraction.py -""" - -import asyncio -import os -import sys - -from openai import AsyncOpenAI - -from memu.app import MemoryService - -# Add src to sys.path -src_path = os.path.abspath("src") -sys.path.insert(0, src_path) - - -async def generate_skill_md( - all_skills, service, output_file, attempt_number, total_attempts, categories=None, is_final=False -): - """ - Use LLM to generate a concise task execution guide (skill.md). - - This creates a production-ready guide incorporating lessons learned from deployment attempts. - """ - - os.makedirs(os.path.dirname(output_file), exist_ok=True) - - # Prepare context for LLM - skills_text = "\n\n".join([f"### From {skill_data['source']}\n{skill_data['skill']}" for skill_data in all_skills]) - - # Get category summaries if available - categories_text = "" - if categories: - categories_with_content = [cat for cat in categories if cat.get("summary") and cat.get("summary").strip()] - if categories_with_content: - categories_text = "\n\n".join([ - f"**{cat.get('name', 'unknown')}**:\n{cat.get('summary', '')}" for cat in categories_with_content - ]) - - # Construct prompt for LLM - prompt = f"""Generate a concise production-ready task execution guide. - -**Context**: -- Task: Production Microservice Deployment with Blue-Green Strategy -- Progress: {attempt_number}/{total_attempts} attempts -- Status: {"Complete" if is_final else f"v0.{attempt_number}"} - -**Skills Learned**: -{skills_text} - -{f"**Categories**:\n{categories_text}" if categories_text else ""} - -**Required Structure**: - -1. **Frontmatter** (YAML): - - name: production-microservice-deployment - - description: Brief description - - version: {"1.0.0" if is_final else f"0.{attempt_number}.0"} - - status: {"Production-Ready" if is_final else "Evolving"} - -2. **Introduction**: What this guide does and when to use it - -3. **Deployment Context**: Strategy, environment, goals - -4. **Pre-Deployment Checklist**: - - Actionable checks from lessons learned - - Group by category (Database, Monitoring, etc.) - - Mark critical items - -5. **Deployment Procedure**: - - Step-by-step instructions with commands - - Include monitoring points - -6. **Rollback Procedure**: - - When to rollback (thresholds) - - Exact commands - - Expected recovery time - -7. **Common Pitfalls & Solutions**: - - Failures/issues encountered - - Root cause, symptoms, solution - -8. **Best Practices**: - - What works well - - Expected timelines - -9. **Key Takeaways**: 3-5 most important lessons - -**Style**: -- Use markdown with clear hierarchy -- Be specific and concise -- Technical and production-grade tone -- Focus on PRACTICAL steps - -**CRITICAL**: -- ONLY use information from provided skills/lessons -- DO NOT make assumptions or add generic advice -- Extract ACTUAL experiences from the logs - -Generate the complete markdown document now:""" - - client = AsyncOpenAI(api_key=service.llm_config.api_key) - - response = await client.chat.completions.create( - model=service.llm_config.chat_model, - messages=[ - { - "role": "system", - "content": "You are an expert technical writer creating concise, production-grade deployment guides from real experiences.", - }, - {"role": "user", "content": prompt}, - ], - temperature=0.7, - max_tokens=3000, - ) - - generated_content = response.choices[0].message.content - - # Write to file - with open(output_file, "w", encoding="utf-8") as f: - f.write(generated_content) - - return True - - -async def main(): - """ - Extract skills from agent logs using incremental memory updates. - - This example demonstrates INCREMENTAL LEARNING: - 1. Process files ONE BY ONE - 2. Each file UPDATES existing memory - 3. Category summaries EVOLVE with each new file - 4. Final output shows accumulated knowledge - """ - print("Example 2: Incremental Skill Extraction") - print("-" * 50) - - # Get OpenAI API key from environment - api_key = os.getenv("OPENAI_API_KEY") - if not api_key: - msg = "Please set OPENAI_API_KEY environment variable" - raise ValueError(msg) - - # Custom config for skill extraction - skill_prompt = """ - You are analyzing an agent execution log. Extract the key actions taken, their outcomes, and lessons learned. - - For each significant action or phase: - - 1. **Action/Phase**: What was being attempted? - 2. **Status**: SUCCESS ✅ or FAILURE ❌ - 3. **What Happened**: What was executed - 4. **Outcome**: What worked/failed, metrics - 5. **Root Cause** (for failures): Why did it fail? - 6. **Lesson**: What did we learn? - 7. **Action Items**: Concrete steps for next time - - **IMPORTANT**: - - Focus on ACTIONS and outcomes - - Be specific: include actual metrics, errors, timing - - ONLY extract information explicitly stated - - DO NOT infer or assume information - - Extract ALL significant actions from the text: - - Text: {resource} - """ - - # Define custom categories - skill_categories = [ - {"name": "deployment_execution", "description": "Deployment actions, traffic shifting, environment management"}, - { - "name": "pre_deployment_validation", - "description": "Capacity validation, configuration checks, readiness verification", - }, - { - "name": "incident_response_rollback", - "description": "Incident response, error detection, rollback procedures", - }, - { - "name": "performance_monitoring", - "description": "Metrics monitoring, performance analysis, bottleneck detection", - }, - {"name": "database_management", "description": "Database capacity planning, optimization, schema changes"}, - {"name": "testing_verification", "description": "Testing, smoke tests, load tests, verification"}, - {"name": "infrastructure_setup", "description": "Kubernetes, containers, networking configuration"}, - {"name": "lessons_learned", "description": "Key reflections, root cause analyses, action items"}, - ] - - memorize_config = { - "memory_types": ["skill"], - "memory_type_prompts": {"skill": skill_prompt}, - "memory_categories": skill_categories, - } - - # Initialize service with OpenAI using llm_profiles - # The "default" profile is required and used as the primary LLM configuration - service = MemoryService( - llm_profiles={ - "default": { - "api_key": api_key, - "chat_model": "gpt-4o-mini", - }, - }, - memorize_config=memorize_config, - ) - - # Resources to process - resources = [ - ("examples/resources/logs/log1.txt", "document"), - ("examples/resources/logs/log2.txt", "document"), - ("examples/resources/logs/log3.txt", "document"), - ] - - # Process each resource sequentially - print("\nProcessing files...") - all_skills = [] - categories = [] - - for idx, (resource_file, modality) in enumerate(resources, 1): - if not os.path.exists(resource_file): - continue - - try: - result = await service.memorize(resource_url=resource_file, modality=modality) - - # Extract skill items - for item in result.get("items", []): - if item.get("memory_type") == "skill": - all_skills.append({"skill": item.get("summary", ""), "source": os.path.basename(resource_file)}) - - # Categories are returned in the result and updated after each memorize call - categories = result.get("categories", []) - - # Generate intermediate skill.md - await generate_skill_md( - all_skills=all_skills, - service=service, - output_file=f"examples/output/skill_example/log_{idx}.md", - attempt_number=idx, - total_attempts=len(resources), - categories=categories, - ) - - except Exception as e: - print(f"Error: {e}") - - # Generate final comprehensive skill.md - await generate_skill_md( - all_skills=all_skills, - service=service, - output_file="examples/output/skill_example/skill.md", - attempt_number=len(resources), - total_attempts=len(resources), - categories=categories, - is_final=True, - ) - - print(f"\n✓ Processed {len(resources)} files, extracted {len(all_skills)} skills") - print(f"✓ Generated {len(categories)} categories") - print("✓ Output: examples/output/skill_example/") - - -if __name__ == "__main__": - asyncio.run(main()) diff --git a/examples/example_5_with_lazyllm_client.py b/examples/example_5_with_lazyllm_client.py index 3b300298..c47a579c 100644 --- a/examples/example_5_with_lazyllm_client.py +++ b/examples/example_5_with_lazyllm_client.py @@ -4,12 +4,10 @@ This example merges functionalities from: 1. Example 1: Conversation Memory Processing -2. Example 2: Skill Extraction -3. Example 3: Multimodal Processing +2. Example 3: Multimodal Processing It demonstrates how to use the LazyLLM backend for: - Processing conversation history -- Extracting technical skills from logs - Handling multimodal content (images + text) - defaut source and model are from qwen @@ -72,72 +70,13 @@ async def run_conversation_memory_demo(service): # ========================================== -# PART 2: Skill Extraction -# ========================================== - - -async def run_skill_extraction_demo(service): - print("\n" + "=" * 60) - print("PART 2: Skill Extraction from Logs") - print("=" * 60) - - # Configure prompt for skill extraction - skill_prompt = """ - You are analyzing an agent execution log. Extract the key actions taken, their outcomes, and lessons learned. - - Output MUST be valid XML wrapped in tags. - Format: - - - - [Action] Description... - [Lesson] Key lesson... - - - Category Name - - - - - Text: {resource} - """ - - # Update service config for skill extraction - service.memorize_config.memory_types = ["skill"] - service.memorize_config.memory_type_prompts = {"skill": skill_prompt} - - logs = ["examples/resources/logs/log1.txt", "examples/resources/logs/log2.txt", "examples/resources/logs/log3.txt"] - - all_skills = [] - for log_file in logs: - if not os.path.exists(log_file): - continue - - print(f" Processing log: {log_file}") - try: - result = await service.memorize(resource_url=log_file, modality="document") - for item in result.get("items", []): - if item.get("memory_type") == "skill": - all_skills.append(item.get("summary", "")) - print(f" ✓ Extracted {len(result.get('items', []))} skills") - except Exception as e: - print(f" ✗ Error: {e}") - - # Generate summary guide - if all_skills: - output_file = "examples/output/lazyllm_example/skills/skill_guide.md" - await generate_skill_guide(all_skills, service, output_file) - print(f"✓ Skill guide generated: {output_file}") - - -# ========================================== -# PART 3: Multimodal Memory +# PART 2: Multimodal Memory # ========================================== async def run_multimodal_demo(service): print("\n" + "=" * 60) - print("PART 3: Multimodal Memory Processing") + print("PART 2: Multimodal Memory Processing") print("=" * 60) # Configure for knowledge extraction @@ -155,8 +94,8 @@ async def run_multimodal_demo(service): Content: {resource} """ - service.memorize_config.memory_types = ["knowledge"] - service.memorize_config.memory_type_prompts = {"knowledge": xml_prompt} + service.memorize_config.entry_types = ["knowledge"] + service.memorize_config.entry_type_prompts = {"knowledge": xml_prompt} resources = [ ("examples/resources/docs/doc1.txt", "document"), @@ -200,18 +139,6 @@ async def generate_markdown_output(categories, output_dir): f.write(cleaned) -async def generate_skill_guide(skills, service, output_file): - os.makedirs(os.path.dirname(output_file), exist_ok=True) - skills_text = "\n\n".join(skills) - prompt = f"Summarize these skills into a guide:\n\n{skills_text}" - - # Use LazyLLM via service - summary = await service.llm_client.chat(text=prompt) - - with open(output_file, "w", encoding="utf-8") as f: - f.write(summary) - - # ========================================== # Main Entry # ========================================== @@ -242,7 +169,6 @@ async def main(): # 2. Run Demos await run_conversation_memory_demo(service) - # await run_skill_extraction_demo(service) # await run_multimodal_demo(service) diff --git a/examples/getting_started_robust.py b/examples/getting_started_robust.py index 883af997..7ecf8c2d 100644 --- a/examples/getting_started_robust.py +++ b/examples/getting_started_robust.py @@ -72,9 +72,9 @@ async def main() -> None: memory_content = "The user is a senior Python architect who loves clean code and type hints." # We use 'create_memory_item' to insert a single memory record. - # memory_type='profile' indicates this is an attribute of the user. + # entry_type='profile' indicates this is an attribute of the user. result = await service.create_memory_item( - memory_type="profile", + entry_type="profile", memory_content=memory_content, memory_categories=["User Facts"], ) @@ -92,7 +92,7 @@ async def main() -> None: if items: print(f"[OK] Found {len(items)} relevant memory item(s):") for idx, item in enumerate(items, 1): - print(f" {idx}. {item.get('summary')} (Type: {item.get('memory_type')})") + print(f" {idx}. {item.get('text')} (Type: {item.get('entry_type')})") else: print("[!] No relevant memories found.") diff --git a/examples/proactive/memory/config.py b/examples/proactive/memory/config.py index 622d5e11..9a8679da 100644 --- a/examples/proactive/memory/config.py +++ b/examples/proactive/memory/config.py @@ -1,8 +1,8 @@ memorize_config = { - "memory_types": [ + "entry_types": [ "record", ], - "memory_type_prompts": { + "entry_type_prompts": { "record": { "objective": { "ordinal": 10, diff --git a/src/memu/app/crud.py b/src/memu/app/crud.py index cfaa2375..545c4cf2 100644 --- a/src/memu/app/crud.py +++ b/src/memu/app/crud.py @@ -8,7 +8,7 @@ from pydantic import BaseModel -from memu.database.models import MemoryType, Resource +from memu.database.models import EntryType, Resource from memu.prompts.category_patch import CATEGORY_PATCH_PROMPT from memu.workflow.step import WorkflowState, WorkflowStep @@ -33,7 +33,7 @@ class CRUDMixin: _escape_prompt_value: Callable[[str], str] user_model: type[BaseModel] patch_config: PatchConfig - _ensure_categories_ready: Callable[[Context, Database, Mapping[str, Any] | None], Awaitable[None]] + _ensure_lanes_ready: Callable[[Context, Database, Mapping[str, Any] | None], Awaitable[None]] async def list_memory_items( self, @@ -284,7 +284,7 @@ def _crud_clear_memory_items(self, state: WorkflowState, step_context: Any) -> W def _crud_clear_memory_resources(self, state: WorkflowState, step_context: Any) -> WorkflowState: where_filters = state.get("where") or {} store = state["store"] - # Remaining resources (raw sources + index/skill docs); memory docs were + # Remaining resources (raw sources + index docs); memory docs were # already cleared by _crud_clear_memory_categories. deleted = store.resource_repo.clear_resources(where_filters) state["deleted_resources"] = deleted @@ -307,30 +307,29 @@ def _crud_build_clear_memory_response(self, state: WorkflowState, step_context: async def create_memory_item( self, *, - memory_type: MemoryType, + entry_type: EntryType, memory_content: str, memory_categories: list[str], user: dict[str, Any] | None = None, propagate: bool = True, ) -> dict[str, Any]: - if memory_type not in get_args(MemoryType): - msg = f"Invalid memory type: '{memory_type}', must be one of {get_args(MemoryType)}" + if entry_type not in get_args(EntryType): + msg = f"Invalid entry type: '{entry_type}', must be one of {get_args(EntryType)}" raise ValueError(msg) ctx = self._get_context() store = self._get_database() user_scope = self.user_model(**user).model_dump() if user is not None else None - await self._ensure_categories_ready(ctx, store, user_scope) + await self._ensure_lanes_ready(ctx, store, user_scope) state: WorkflowState = { "memory_payload": { - "type": memory_type, + "type": entry_type, "content": memory_content, "categories": memory_categories, }, "ctx": ctx, "store": store, - "category_ids": list(ctx.category_ids), "user": user_scope, "propagate": propagate, } @@ -346,34 +345,33 @@ async def update_memory_item( self, *, memory_id: str, - memory_type: MemoryType | None = None, + entry_type: EntryType | None = None, memory_content: str | None = None, memory_categories: list[str] | None = None, user: dict[str, Any] | None = None, propagate: bool = True, ) -> dict[str, Any]: - if all((memory_type is None, memory_content is None, memory_categories is None)): - msg = "At least one of memory type, memory content, or memory categories is required for UPDATE operation" + if all((entry_type is None, memory_content is None, memory_categories is None)): + msg = "At least one of entry type, memory content, or memory categories is required for UPDATE operation" raise ValueError(msg) - if memory_type and memory_type not in get_args(MemoryType): - msg = f"Invalid memory type: '{memory_type}', must be one of {get_args(MemoryType)}" + if entry_type and entry_type not in get_args(EntryType): + msg = f"Invalid entry type: '{entry_type}', must be one of {get_args(EntryType)}" raise ValueError(msg) ctx = self._get_context() store = self._get_database() user_scope = self.user_model(**user).model_dump() if user is not None else None - await self._ensure_categories_ready(ctx, store, user_scope) + await self._ensure_lanes_ready(ctx, store, user_scope) state: WorkflowState = { "memory_id": memory_id, "memory_payload": { - "type": memory_type, + "type": entry_type, "content": memory_content, "categories": memory_categories, }, "ctx": ctx, "store": store, - "category_ids": list(ctx.category_ids), "user": user_scope, "propagate": propagate, } @@ -395,13 +393,12 @@ async def delete_memory_item( ctx = self._get_context() store = self._get_database() user_scope = self.user_model(**user).model_dump() if user is not None else None - await self._ensure_categories_ready(ctx, store, user_scope) + await self._ensure_lanes_ready(ctx, store, user_scope) state: WorkflowState = { "memory_id": memory_id, "ctx": ctx, "store": store, - "category_ids": list(ctx.category_ids), "user": user_scope, "propagate": propagate, } @@ -547,7 +544,7 @@ async def _patch_create_memory_item(self, state: WorkflowState, step_context: An item = store.entry_repo.create_entry( lane="memory", source_id=None, - entry_kind=memory_payload["type"], + entry_type=memory_payload["type"], text=memory_payload["content"], embedding=content_embedding, user_data=dict(user or {}), @@ -591,7 +588,7 @@ async def _patch_update_memory_item(self, state: WorkflowState, step_context: An if memory_payload["type"] or memory_payload["content"]: item = store.entry_repo.update_entry( entry_id=memory_id, - entry_kind=memory_payload["type"], + entry_type=memory_payload["type"], text=memory_payload["content"], embedding=content_embedding, ) @@ -709,11 +706,12 @@ def _patch_build_response(self, state: WorkflowState, step_context: Any) -> Work def _map_category_names_to_ids(self, names: list[str], ctx: Context) -> list[str]: if not names: return [] + name_to_id = ctx.lane("memory").name_to_id mapped: list[str] = [] seen: set[str] = set() for name in names: key = name.strip().lower() - cid = ctx.category_name_to_id.get(key) + cid = name_to_id.get(key) if cid and cid not in seen: mapped.append(cid) seen.add(cid) diff --git a/src/memu/app/memorize.py b/src/memu/app/memorize.py index 417db804..2d13923e 100644 --- a/src/memu/app/memorize.py +++ b/src/memu/app/memorize.py @@ -12,9 +12,10 @@ import defusedxml.ElementTree as ET from pydantic import BaseModel -from memu.app.settings import CategoryConfig, CustomPrompt +from memu.app.settings import CategoryConfig, CustomPrompt, LaneConfig from memu.blob.folder import diff_folder, load_manifest, manifest_from_scan, save_manifest, scan_folder -from memu.database.models import Entry, MemoryType, Resource, ResourceEntry +from memu.database.models import Entry, EntryType, Resource, ResourceEntry +from memu.memory_fs.exporter import slugify from memu.preprocess import PreprocessContext, preprocess_resource from memu.prompts.category_summary import ( CUSTOM_PROMPT as CATEGORY_SUMMARY_CUSTOM_PROMPT, @@ -22,15 +23,14 @@ from memu.prompts.category_summary import ( PROMPT as CATEGORY_SUMMARY_PROMPT, ) -from memu.prompts.memory_type import ( - CUSTOM_PROMPTS as MEMORY_TYPE_CUSTOM_PROMPTS, +from memu.prompts.entry_type import ( + CUSTOM_PROMPTS as ENTRY_TYPE_CUSTOM_PROMPTS, ) -from memu.prompts.memory_type import ( +from memu.prompts.entry_type import ( CUSTOM_TYPE_CUSTOM_PROMPTS, - DEFAULT_MEMORY_TYPES, ) -from memu.prompts.memory_type import ( - PROMPTS as MEMORY_TYPE_PROMPTS, +from memu.prompts.entry_type import ( + PROMPTS as ENTRY_TYPE_PROMPTS, ) from memu.workflow.step import WorkflowState, WorkflowStep @@ -46,9 +46,9 @@ class MemorizeMixin: if TYPE_CHECKING: memorize_config: MemorizeConfig - category_configs: list[CategoryConfig] - category_config_map: dict[str, CategoryConfig] - _category_prompt_str: str + lane_configs: dict[str, LaneConfig] + lane_category_config_maps: dict[str, dict[str, CategoryConfig]] + lane_prompt_strs: dict[str, str] fs: LocalFS _run_workflow: Callable[..., Awaitable[WorkflowState]] _get_context: Callable[[], Context] @@ -86,18 +86,13 @@ async def memorize( ctx = self._get_context() store = self._get_database() user_scope = self.user_model(**user).model_dump() if user is not None else None - await self._ensure_categories_ready(ctx, store, user_scope) - - memory_types = self._resolve_memory_types() + await self._ensure_lanes_ready(ctx, store, user_scope) state: WorkflowState = { "resource_url": resource_url, "modality": modality, - "memory_types": memory_types, - "categories_prompt_str": self._category_prompt_str, "ctx": ctx, "store": store, - "category_ids": list(ctx.category_ids), "user": user_scope, } @@ -130,7 +125,7 @@ async def memorize_workspace( ctx = self._get_context() store = self._get_database() user_scope = self.user_model(**user).model_dump() if user is not None else None - await self._ensure_categories_ready(ctx, store, user_scope) + await self._ensure_lanes_ready(ctx, store, user_scope) root = pathlib.Path(folder).resolve() scanned = scan_folder(root) @@ -193,15 +188,11 @@ async def _memorize_one( workspace sync can collect the created resources) and takes an already resolved ``user_scope``/``ctx``/``store`` to avoid re-resolving them per file. """ - memory_types = self._resolve_memory_types() state: WorkflowState = { "resource_url": resource_url, "modality": modality, - "memory_types": memory_types, - "categories_prompt_str": self._category_prompt_str, "ctx": ctx, "store": store, - "category_ids": list(ctx.category_ids), "user": user_scope, } result = await self._run_workflow("memorize", state) @@ -303,14 +294,11 @@ def _build_memorize_workflow(self) -> list[WorkflowStep]: handler=self._memorize_extract_items, requires={ "preprocessed_resources", - "memory_types", - "categories_prompt_str", "modality", "resource_url", }, produces={"resource_plans"}, capabilities={"llm"}, - config={"chat_llm_profile": self.memorize_config.memory_extract_llm_profile}, ), WorkflowStep( step_id="dedupe_merge", @@ -325,7 +313,7 @@ def _build_memorize_workflow(self) -> list[WorkflowStep]: role="categorize", handler=self._memorize_categorize_items, requires={"resource_plans", "ctx", "store", "local_path", "modality", "user"}, - produces={"resources", "items", "relations", "category_updates"}, + produces={"resources", "items", "relations", "lane_updates"}, capabilities={"db", "vector"}, config={"embed_llm_profile": "embedding"}, ), @@ -333,16 +321,15 @@ def _build_memorize_workflow(self) -> list[WorkflowStep]: step_id="persist_index", role="persist", handler=self._memorize_persist_and_index, - requires={"category_updates", "ctx", "store"}, - produces={"categories"}, + requires={"lane_updates", "ctx", "store"}, + produces={"lane_docs"}, capabilities={"db", "llm"}, - config={"chat_llm_profile": self.memorize_config.category_update_llm_profile}, ), WorkflowStep( step_id="build_response", role="emit", handler=self._memorize_build_response, - requires={"resources", "items", "relations", "ctx", "store", "category_ids"}, + requires={"resources", "items", "relations", "lane_updates", "ctx", "store"}, produces={"response"}, capabilities=set(), ), @@ -354,11 +341,8 @@ def _list_memorize_initial_keys() -> set[str]: return { "resource_url", "modality", - "memory_types", - "categories_prompt_str", "ctx", "store", - "category_ids", "user", } @@ -388,29 +372,33 @@ async def _memorize_preprocess_multimodal(self, state: WorkflowState, step_conte return state async def _memorize_extract_items(self, state: WorkflowState, step_context: Any) -> WorkflowState: - llm_client = self._get_step_llm_client(step_context) preprocessed_resources = state.get("preprocessed_resources", []) resource_plans: list[dict[str, Any]] = [] total_segments = len(preprocessed_resources) or 1 + enabled_lanes = self.memorize_config.enabled_lanes for idx, prep in enumerate(preprocessed_resources): res_url = self._segment_resource_url(state["resource_url"], idx, total_segments) text = prep.get("text") caption = prep.get("caption") - structured_entries = await self._generate_structured_entries( - modality=state["modality"], - memory_types=state["memory_types"], - text=text, - categories_prompt_str=state["categories_prompt_str"], - llm_client=llm_client, - ) + lane_entries: dict[str, list[tuple[EntryType, str, list[str]]]] = {} + for lane, lane_cfg in enabled_lanes.items(): + lane_client = self._get_llm_client(lane_cfg.extract_llm_profile, step_context=step_context) + lane_entries[lane] = await self._generate_structured_entries( + modality=state["modality"], + lane=lane, + lane_cfg=lane_cfg, + text=text, + categories_prompt_str=self.lane_prompt_strs.get(lane, ""), + llm_client=lane_client, + ) resource_plans.append({ "resource_url": res_url, "text": text, "caption": caption, - "entries": structured_entries, + "lane_entries": lane_entries, }) state["resource_plans"] = resource_plans @@ -430,8 +418,10 @@ async def _memorize_categorize_items(self, state: WorkflowState, step_context: A resources: list[Resource] = [] items: list[Entry] = [] relations: list[ResourceEntry] = [] - category_updates: dict[str, list[tuple[str, str]]] = {} + # lane -> {doc_id: [(entry_id, text)]} + lane_updates: dict[str, dict[str, list[tuple[str, str]]]] = {} user_scope = state.get("user", {}) + enabled_lanes = self.memorize_config.enabled_lanes for plan in state.get("resource_plans", []): res = await self._create_resource_with_caption( @@ -446,75 +436,93 @@ async def _memorize_categorize_items(self, state: WorkflowState, step_context: A ) resources.append(res) - entries = plan.get("entries") or [] - if not entries: - continue - - mem_items, rels, cat_updates = await self._persist_memory_items( - resource_id=res.id, - resource_path=res.source_path, - structured_entries=entries, - ctx=ctx, - store=store, - embed_client=embed_client, - user=user_scope, - ) - items.extend(mem_items) - relations.extend(rels) - for cat_id, mems in cat_updates.items(): - category_updates.setdefault(cat_id, []).extend(mems) + plan_lane_entries = plan.get("lane_entries") or {} + for lane, lane_cfg in enabled_lanes.items(): + entries = plan_lane_entries.get(lane) or [] + if not entries: + continue + mem_items, rels, doc_updates = await self._persist_lane_entries( + lane=lane, + lane_cfg=lane_cfg, + source_resource=res, + structured_entries=entries, + ctx=ctx, + store=store, + embed_client=embed_client, + user=user_scope, + ) + items.extend(mem_items) + relations.extend(rels) + lane_bucket = lane_updates.setdefault(lane, {}) + for doc_id, mems in doc_updates.items(): + lane_bucket.setdefault(doc_id, []).extend(mems) state.update({ "resources": resources, "items": items, "relations": relations, - "category_updates": category_updates, + "lane_updates": lane_updates, }) return state async def _memorize_persist_and_index(self, state: WorkflowState, step_context: Any) -> WorkflowState: - llm_client = self._get_step_llm_client(step_context) - updated_summaries = await self._update_category_summaries( - state.get("category_updates", {}), - ctx=state["ctx"], - store=state["store"], - llm_client=llm_client, - ) - if self.memorize_config.enable_item_references: - await self._persist_item_references( - updated_summaries=updated_summaries, - category_updates=state.get("category_updates", {}), - store=state["store"], + ctx = state["ctx"] + store = state["store"] + lane_updates: dict[str, dict[str, list[tuple[str, str]]]] = state.get("lane_updates", {}) + for lane, lane_cfg in self.memorize_config.enabled_lanes.items(): + # Only adaptive lanes synthesize a group-doc summary; per_resource lanes + # already finalized their 1:1 doc summary during persistence. + if lane_cfg.grouping != "adaptive": + continue + updates = lane_updates.get(lane) or {} + if not updates: + continue + llm_client = self._get_llm_client(lane_cfg.summary_llm_profile, step_context=step_context) + updated_summaries = await self._update_group_summaries( + lane, + lane_cfg, + updates, + ctx=ctx, + store=store, + llm_client=llm_client, ) + if lane_cfg.enable_item_references: + await self._persist_item_references( + updated_summaries=updated_summaries, + category_updates=updates, + store=store, + ) + state["lane_docs"] = lane_updates return state def _memorize_build_response(self, state: WorkflowState, step_context: Any) -> WorkflowState: - ctx = state["ctx"] store = state["store"] resources = [self._model_dump_without_embeddings(r) for r in state.get("resources", [])] items = [self._model_dump_without_embeddings(item) for item in state.get("items", [])] relations = [rel.model_dump() for rel in state.get("relations", [])] - category_ids = state.get("category_ids") or list(ctx.category_ids) - categories = [ - self._model_dump_without_embeddings(res) - for c in category_ids - if (res := store.resource_repo.get_resource(c)) is not None - ] + # Per-lane coarse docs touched by this memorize, keyed by lane. + lane_updates: dict[str, dict[str, list[tuple[str, str]]]] = state.get("lane_updates", {}) + lane_docs: dict[str, list[dict[str, Any]]] = {} + for lane, doc_map in lane_updates.items(): + lane_docs[lane] = [ + self._model_dump_without_embeddings(res) + for did in doc_map + if (res := store.resource_repo.get_resource(did)) is not None + ] + # Backward-friendly alias: "categories" continues to mean the memory lane docs. + categories = lane_docs.get("memory", []) + + response: dict[str, Any] = { + "items": items, + "categories": categories, + "lanes": lane_docs, + "relations": relations, + } if len(resources) == 1: - response = { - "resource": resources[0], - "items": items, - "categories": categories, - "relations": relations, - } + response["resource"] = resources[0] else: - response = { - "resources": resources, - "items": items, - "categories": categories, - "relations": relations, - } + response["resources"] = resources state["response"] = response return state @@ -554,10 +562,6 @@ async def _create_resource_with_caption( user_data=dict(user or {}), ) - def _resolve_memory_types(self) -> list[MemoryType]: - configured_types = self.memorize_config.memory_types or DEFAULT_MEMORY_TYPES - return [cast(MemoryType, mtype) for mtype in configured_types] - @staticmethod def _resolve_custom_prompt(prompt: str | CustomPrompt, templates: Mapping[str, str]) -> str: if isinstance(prompt, str): @@ -577,20 +581,24 @@ async def _generate_structured_entries( self, *, modality: str, - memory_types: list[MemoryType], + lane: str, + lane_cfg: LaneConfig, text: str | None, categories_prompt_str: str, segments: list[dict[str, int | str]] | None = None, llm_client: Any | None = None, - ) -> list[tuple[MemoryType, str, list[str]]]: - if not memory_types or not text: + ) -> list[tuple[EntryType, str, list[str]]]: + entry_types = [cast(EntryType, mtype) for mtype in lane_cfg.entry_types] + if not entry_types or not text: return [] client = llm_client or self._get_llm_client() return await self._generate_text_entries( resource_text=text, modality=modality, - memory_types=memory_types, + lane=lane, + lane_cfg=lane_cfg, + entry_types=entry_types, categories_prompt_str=categories_prompt_str, segments=segments, llm_client=client, @@ -601,16 +609,20 @@ async def _generate_text_entries( *, resource_text: str, modality: str, - memory_types: list[MemoryType], + lane: str, + lane_cfg: LaneConfig, + entry_types: list[EntryType], categories_prompt_str: str, segments: list[dict[str, int | str]] | None, llm_client: Any | None = None, - ) -> list[tuple[MemoryType, str, list[str]]]: + ) -> list[tuple[EntryType, str, list[str]]]: if modality == "conversation" and segments: segment_entries = await self._generate_entries_for_segments( resource_text=resource_text, segments=segments, - memory_types=memory_types, + lane=lane, + lane_cfg=lane_cfg, + entry_types=entry_types, categories_prompt_str=categories_prompt_str, llm_client=llm_client, ) @@ -618,7 +630,9 @@ async def _generate_text_entries( return segment_entries return await self._generate_entries_from_text( resource_text=resource_text, - memory_types=memory_types, + lane=lane, + lane_cfg=lane_cfg, + entry_types=entry_types, categories_prompt_str=categories_prompt_str, llm_client=llm_client, ) @@ -628,11 +642,13 @@ async def _generate_entries_for_segments( *, resource_text: str, segments: list[dict[str, int | str]], - memory_types: list[MemoryType], + lane: str, + lane_cfg: LaneConfig, + entry_types: list[EntryType], categories_prompt_str: str, llm_client: Any | None = None, - ) -> list[tuple[MemoryType, str, list[str]]]: - entries: list[tuple[MemoryType, str, list[str]]] = [] + ) -> list[tuple[EntryType, str, list[str]]]: + entries: list[tuple[EntryType, str, list[str]]] = [] lines = resource_text.split("\n") max_idx = len(lines) - 1 for segment in segments: @@ -643,7 +659,9 @@ async def _generate_entries_for_segments( continue segment_entries = await self._generate_entries_from_text( resource_text=segment_text, - memory_types=memory_types, + lane=lane, + lane_cfg=lane_cfg, + entry_types=entry_types, categories_prompt_str=categories_prompt_str, llm_client=llm_client, ) @@ -654,33 +672,36 @@ async def _generate_entries_from_text( self, *, resource_text: str, - memory_types: list[MemoryType], + lane: str, + lane_cfg: LaneConfig, + entry_types: list[EntryType], categories_prompt_str: str, llm_client: Any | None = None, - ) -> list[tuple[MemoryType, str, list[str]]]: - if not memory_types: + ) -> list[tuple[EntryType, str, list[str]]]: + if not entry_types: return [] client = llm_client or self._get_llm_client() prompts = [ - self._build_memory_type_prompt( - memory_type=mtype, + self._build_entry_type_prompt( + entry_type=mtype, + lane_cfg=lane_cfg, resource_text=resource_text, categories_str=categories_prompt_str, ) - for mtype in memory_types + for mtype in entry_types ] valid_prompts = [prompt for prompt in prompts if prompt.strip()] # These prompts are instructions that request structured output, not text summaries. tasks = [client.chat(prompt_text) for prompt_text in valid_prompts] responses = await asyncio.gather(*tasks) - return self._parse_structured_entries(memory_types, responses) + return self._parse_structured_entries(entry_types, responses) def _parse_structured_entries( - self, memory_types: list[MemoryType], responses: Sequence[str] - ) -> list[tuple[MemoryType, str, list[str]]]: - entries: list[tuple[MemoryType, str, list[str]]] = [] - for mtype, response in zip(memory_types, responses, strict=True): - parsed = self._parse_memory_type_response_xml(response) + self, entry_types: list[EntryType], responses: Sequence[str] + ) -> list[tuple[EntryType, str, list[str]]]: + entries: list[tuple[EntryType, str, list[str]]] = [] + for mtype, response in zip(entry_types, responses, strict=True): + parsed = self._parse_entry_type_response_xml(response) for entry in parsed: content = (entry.get("content") or "").strip() if not content: @@ -700,23 +721,25 @@ def _extract_segment_text(self, lines: list[str], start_idx: int, end_idx: int) segment_lines.append(line) return "\n".join(segment_lines) if segment_lines else None - async def _persist_memory_items( + async def _persist_lane_entries( self, *, - resource_id: str, - resource_path: str | None, - structured_entries: list[tuple[MemoryType, str, list[str]]], + lane: str, + lane_cfg: LaneConfig, + source_resource: Resource, + structured_entries: list[tuple[EntryType, str, list[str]]], ctx: Context, store: Database, embed_client: Any | None = None, user: Mapping[str, Any] | None = None, ) -> tuple[list[Entry], list[ResourceEntry], dict[str, list[tuple[str, str]]]]: - """ - Persist memory-lane entries and track memory-doc updates. + """Persist a lane's entries and track its coarse-doc updates. - Returns: - Tuple of (entries, relations, category_updates) - where category_updates maps memory-doc id -> list of (entry_id, text) tuples + Returns ``(entries, relations, doc_updates)`` where ``doc_updates`` maps a + coarse lane-doc id -> list of ``(entry_id, text)`` tuples. Grouping depends + on ``lane_cfg.grouping``: ``adaptive`` resolves extractor-proposed group + names to lane docs (creating unseen ones), while ``per_resource`` links all + of a source's entries to a single 1:1 lane doc. """ summary_payloads = [content for _, content, _ in structured_entries] client = embed_client or self._get_embedding_client() @@ -724,15 +747,20 @@ async def _persist_memory_items( items: list[Entry] = [] rels: list[ResourceEntry] = [] # Stores (entry_id, text) tuples for reference support. - category_memory_updates: dict[str, list[tuple[str, str]]] = {} + doc_updates: dict[str, list[tuple[str, str]]] = {} - reinforce = self.memorize_config.enable_item_reinforcement - for (memory_type, summary_text, cat_names), emb in zip(structured_entries, item_embeddings, strict=True): + per_resource = lane_cfg.grouping == "per_resource" + per_resource_doc_id = ( + self._ensure_per_resource_doc(lane, source_resource, store, user=user) if per_resource else None + ) + + reinforce = lane_cfg.enable_item_reinforcement + for (entry_type, summary_text, cat_names), emb in zip(structured_entries, item_embeddings, strict=True): item = store.entry_repo.create_entry( - lane="memory", - source_id=resource_id, - source_path=resource_path, - entry_kind=memory_type, + lane=lane, + source_id=source_resource.id, + source_path=source_resource.source_path, + entry_type=entry_type, text=summary_text, embedding=emb, user_data=dict(user or {}), @@ -742,20 +770,67 @@ async def _persist_memory_items( if reinforce and item.extra.get("reinforcement_count", 1) > 1: # existing item continue - mapped_cat_ids = await self._resolve_category_ids(cat_names, ctx, store, user=user) - for cid in mapped_cat_ids: - rels.append(store.resource_entry_repo.link_entry_resource(item.id, cid, user_data=dict(user or {}))) + if per_resource: + doc_ids = [per_resource_doc_id] if per_resource_doc_id else [] + else: + doc_ids = await self._resolve_group_ids(lane, cat_names, ctx, store, user=user) + for did in doc_ids: + rels.append(store.resource_entry_repo.link_entry_resource(item.id, did, user_data=dict(user or {}))) # Store (entry_id, text) tuple for reference support - category_memory_updates.setdefault(cid, []).append((item.id, summary_text)) + doc_updates.setdefault(did, []).append((item.id, summary_text)) + + if per_resource and per_resource_doc_id and items: + # A per_resource (index) doc has no LLM summary synthesis: its searchable + # body is just the concatenation of its description entries. + body = "\n".join(item.text for item in items if (item.text or "").strip()) + store.resource_repo.update_resource( + resource_id=per_resource_doc_id, + summary=body, + embedding=item_embeddings[0] if item_embeddings else None, + ) + + return items, rels, doc_updates + + def _ensure_per_resource_doc( + self, + lane: str, + source_resource: Resource, + store: Database, + *, + user: Mapping[str, Any] | None = None, + ) -> str: + """Get-or-create the single coarse lane doc that 1:1 mirrors a source. - return items, rels, category_memory_updates + Used by ``per_resource`` lanes (e.g. index): the doc is keyed by the source + resource id so re-memorizing the same source reuses it. Returns the doc id. + """ + from os.path import basename + + name = basename(source_resource.url or source_resource.local_path or source_resource.id) + title = name or source_resource.id + slug = f"{slugify(title)}-{source_resource.id[:8]}" + doc = store.resource_repo.get_or_create_doc( + lane=lane, + title=title, + description=source_resource.summary or "", + embedding=source_resource.embedding or [], + user_data=dict(user or {}), + slug=slug, + ) + return doc.id - async def _ensure_categories_ready( + async def _ensure_lanes_ready( self, ctx: Context, store: Database, user_scope: Mapping[str, Any] | None = None ) -> None: - if ctx.categories_ready: - return - await self._initialize_categories(ctx, store, user_scope) + """Initialize seed group docs for every enabled adaptive lane (idempotent).""" + for lane, lane_cfg in self.memorize_config.enabled_lanes.items(): + lane_state = ctx.lane(lane) + if lane_cfg.grouping != "adaptive": + lane_state.ready = True + continue + if lane_state.ready: + continue + await self._initialize_lane_groups(lane, lane_cfg, ctx, store, user_scope) @staticmethod def _classify_categories( @@ -781,20 +856,22 @@ def _classify_categories( ready[i] = ex return to_create, to_update, ready - async def _initialize_categories( - self, ctx: Context, store: Database, user: Mapping[str, Any] | None = None + async def _initialize_lane_groups( + self, lane: str, lane_cfg: LaneConfig, ctx: Context, store: Database, user: Mapping[str, Any] | None = None ) -> None: - if ctx.categories_ready: + lane_state = ctx.lane(lane) + if lane_state.ready: return - if not self.category_configs: - ctx.categories_ready = True + seeds = lane_cfg.seed_categories + if not seeds: + lane_state.ready = True return user_data = dict(user or {}) - existing = store.resource_repo.list_resources(where=user_data or None, lane="memory") + existing = store.resource_repo.list_resources(where=user_data or None, lane=lane) existing_by_name: dict[str, Resource] = {c.title: c for c in existing.values() if c.title} - to_create, to_update, ready = self._classify_categories(self.category_configs, existing_by_name) + to_create, to_update, ready = self._classify_categories(seeds, existing_by_name) needs_embed: list[tuple[int, CategoryConfig]] = [] needs_embed.extend(to_create) @@ -807,15 +884,13 @@ async def _initialize_categories( for (i, _), vec in zip(needs_embed, vecs, strict=True): embed_map[i] = vec - from memu.memory_fs.exporter import slugify - cats: dict[int, Resource] = dict(ready) for i, cfg in to_create: name = cfg.name.strip() or "Untitled" description = cfg.description.strip() cat = store.resource_repo.get_or_create_doc( - lane="memory", + lane=lane, title=name, description=description, embedding=embed_map[i], @@ -831,14 +906,14 @@ async def _initialize_categories( ) cats[i] = cat - ctx.category_ids = [] - ctx.category_name_to_id = {} - for i in range(len(self.category_configs)): + lane_state.doc_ids = [] + lane_state.name_to_id = {} + for i in range(len(seeds)): cat = cats[i] - ctx.category_ids.append(cat.id) - name = self.category_configs[i].name.strip() or "Untitled" - ctx.category_name_to_id[name.lower()] = cat.id - ctx.categories_ready = True + lane_state.doc_ids.append(cat.id) + name = seeds[i].name.strip() or "Untitled" + lane_state.name_to_id[name.lower()] = cat.id + lane_state.ready = True @staticmethod def _category_embedding_text(cat: CategoryConfig) -> str: @@ -846,22 +921,24 @@ def _category_embedding_text(cat: CategoryConfig) -> str: desc = cat.description.strip() return f"{name}: {desc}" if desc else name - def _map_category_names_to_ids(self, names: list[str], ctx: Context) -> list[str]: + def _map_group_names_to_ids(self, lane: str, names: list[str], ctx: Context) -> list[str]: if not names: return [] + name_to_id = ctx.lane(lane).name_to_id mapped: list[str] = [] seen: set[str] = set() for name in names: key = name.strip().lower() - cid = ctx.category_name_to_id.get(key) + cid = name_to_id.get(key) if cid and cid not in seen: mapped.append(cid) seen.add(cid) return mapped @staticmethod - def _partition_category_names(names: list[str], ctx: Context) -> tuple[list[str], list[str]]: - """Split proposed names into (known category ids, unknown names) with dedup.""" + def _partition_group_names(lane: str, names: list[str], ctx: Context) -> tuple[list[str], list[str]]: + """Split proposed names into (known group ids, unknown names) with dedup.""" + name_to_id = ctx.lane(lane).name_to_id known_ids: list[str] = [] known_seen: set[str] = set() unknown: list[str] = [] @@ -870,7 +947,7 @@ def _partition_category_names(names: list[str], ctx: Context) -> tuple[list[str] key = name.strip().lower() if not key: continue - cid = ctx.category_name_to_id.get(key) + cid = name_to_id.get(key) if cid is not None: if cid not in known_seen: known_ids.append(cid) @@ -880,43 +957,42 @@ def _partition_category_names(names: list[str], ctx: Context) -> tuple[list[str] unknown_seen.add(key) return known_ids, unknown - async def _resolve_category_ids( + async def _resolve_group_ids( self, + lane: str, names: list[str], ctx: Context, store: Database, *, user: Mapping[str, Any] | None = None, ) -> list[str]: - """Resolve extractor-proposed category names to ids, creating unknown ones. + """Resolve extractor-proposed group names to lane-doc ids, creating unknown ones. - Implements the open/adaptive taxonomy: the kernel presets no categories, so any - category name the extractor proposes is created on first sight and cached in the - context. (Embedding-similarity merging of near-duplicate categories is handled - later by consolidation; here we only do exact-name dedup.) + Implements the open/adaptive taxonomy per lane: any group name the extractor + proposes is created on first sight as a ``lane`` doc and cached in the lane's + context state. Here we only do exact-name dedup. """ if not names: return [] user_data = dict(user or {}) - resolved, unknown = self._partition_category_names(names, ctx) + lane_state = ctx.lane(lane) + resolved, unknown = self._partition_group_names(lane, names, ctx) seen: set[str] = set(resolved) if unknown: - from memu.memory_fs.exporter import slugify - vecs = await self._get_embedding_client("embedding").embed(unknown) for name, vec in zip(unknown, vecs, strict=True): cat = store.resource_repo.get_or_create_doc( - lane="memory", + lane=lane, title=name, description="", embedding=vec, user_data=user_data, slug=slugify(name), ) - ctx.category_name_to_id[name.lower()] = cat.id - if cat.id not in ctx.category_ids: - ctx.category_ids.append(cat.id) + lane_state.name_to_id[name.lower()] = cat.id + if cat.id not in lane_state.doc_ids: + lane_state.doc_ids.append(cat.id) if cat.id not in seen: resolved.append(cat.id) seen.add(cat.id) @@ -963,15 +1039,17 @@ def _format_categories_for_prompt(self, categories: list[CategoryConfig]) -> str lines.append(f"- {name}: {desc}" if desc else f"- {name}") return "Existing categories (reuse when appropriate):\n" + "\n".join(lines) + "\n\n" + adaptive_hint - def _build_memory_type_prompt(self, *, memory_type: MemoryType, resource_text: str, categories_str: str) -> str: - configured_prompt = self.memorize_config.memory_type_prompts.get(memory_type) + def _build_entry_type_prompt( + self, *, entry_type: EntryType, lane_cfg: LaneConfig, resource_text: str, categories_str: str + ) -> str: + configured_prompt = lane_cfg.entry_type_prompts.get(entry_type) if configured_prompt is None: - template = MEMORY_TYPE_PROMPTS.get(memory_type) + template = ENTRY_TYPE_PROMPTS.get(entry_type) elif isinstance(configured_prompt, str): template = configured_prompt else: template = self._resolve_custom_prompt( - configured_prompt, MEMORY_TYPE_CUSTOM_PROMPTS.get(memory_type, CUSTOM_TYPE_CUSTOM_PROMPTS) + configured_prompt, ENTRY_TYPE_CUSTOM_PROMPTS.get(entry_type, CUSTOM_TYPE_CUSTOM_PROMPTS) ) if not template: return resource_text @@ -1036,21 +1114,24 @@ async def _persist_item_references( extra={"ref_id": short_id}, ) - def _build_category_summary_prompt( + def _build_group_summary_prompt( self, *, + lane: str, + lane_cfg: LaneConfig, category: Resource, new_memories: list[str] | list[tuple[str, str]], ) -> str: """ - Build the prompt for updating a category summary. + Build the prompt for updating an adaptive lane's group-doc summary. Args: - category: The category to update - new_memories: Either list of summary strings (legacy) or list of (item_id, summary) tuples (with refs) + lane: The lane the group doc belongs to. + lane_cfg: That lane's configuration (summary prompt/length/refs). + category: The group doc to update. + new_memories: Either summary strings (legacy) or (item_id, summary) tuples (with refs). """ - # Check if references are enabled and we have (id, summary) tuples - enable_refs = getattr(self.memorize_config, "enable_item_references", False) + enable_refs = lane_cfg.enable_item_references if enable_refs: from memu.prompts.category_summary import ( @@ -1078,19 +1159,16 @@ def _build_category_summary_prompt( new_items_text = "\n".join(f"- {m}" for m in str_memories if m.strip()) original = category.summary or "" - category_config = self.category_config_map.get(category.title or "") - configured_prompt = ( - category_config and category_config.summary_prompt - ) or self.memorize_config.default_category_summary_prompt + # Per-seed overrides for this lane, falling back to the lane's defaults. + seed_config = self.lane_category_config_maps.get(lane, {}).get(category.title or "") + configured_prompt = (seed_config and seed_config.summary_prompt) or lane_cfg.summary_prompt if configured_prompt is None: prompt = category_summary_prompt elif isinstance(configured_prompt, str): prompt = configured_prompt else: prompt = self._resolve_custom_prompt(configured_prompt, category_summary_custom_prompt) - target_length = ( - category_config and category_config.target_length - ) or self.memorize_config.default_category_summary_target_length + target_length = (seed_config and seed_config.target_length) or lane_cfg.summary_target_length return prompt.format( category=self._escape_prompt_value(category.title or ""), original_content=self._escape_prompt_value(original or ""), @@ -1098,18 +1176,20 @@ def _build_category_summary_prompt( target_length=target_length, ) - async def _update_category_summaries( + async def _update_group_summaries( self, + lane: str, + lane_cfg: LaneConfig, updates: dict[str, list[tuple[str, str]]] | dict[str, list[str]], ctx: Context, store: Database, llm_client: Any | None = None, ) -> dict[str, str]: """ - Update category summaries based on new memory items. + Synthesize an adaptive lane's group-doc summaries from new entries. Returns: - dict mapping category_id -> updated summary text + dict mapping group-doc id -> updated summary text """ updated_summaries: dict[str, str] = {} if not updates: @@ -1121,7 +1201,9 @@ async def _update_category_summaries( cat = store.resource_repo.get_resource(cid) if not cat or not memories: continue - prompt = self._build_category_summary_prompt(category=cat, new_memories=memories) + prompt = self._build_group_summary_prompt( + lane=lane, lane_cfg=lane_cfg, category=cat, new_memories=memories + ) tasks.append(client.chat(prompt)) target_ids.append(cid) if not tasks: @@ -1141,7 +1223,7 @@ async def _update_category_summaries( def _find_xml_boundaries(self, raw: str) -> tuple[int, int, str] | None: """Find the start index, end index, and closing tag for XML root element.""" - root_tags = ["item", "profile", "behaviors", "events", "knowledge", "skills"] + root_tags = ["item", "profile", "events", "knowledge"] for tag in root_tags: opening = f"<{tag}>" closing = f"" @@ -1164,17 +1246,20 @@ def _parse_memory_element(self, memory_elem: Element) -> dict[str, Any] | None: if categories_elem is not None: categories = [cat_elem.text.strip() for cat_elem in categories_elem.findall("category") if cat_elem.text] memory_dict["categories"] = categories + else: + # per_resource lanes (e.g. index) emit a description with no categories. + memory_dict["categories"] = [] - if memory_dict.get("content") and memory_dict.get("categories"): + if memory_dict.get("content"): return memory_dict return None - def _parse_memory_type_response_xml(self, raw: str) -> list[dict[str, Any]]: + def _parse_entry_type_response_xml(self, raw: str) -> list[dict[str, Any]]: """ Parse XML memory extraction output into a list of memory items. Expected XML format (root tag varies by memory type): - + ... diff --git a/src/memu/app/memory_files.py b/src/memu/app/memory_files.py index 528d2107..5a3ae755 100644 --- a/src/memu/app/memory_files.py +++ b/src/memu/app/memory_files.py @@ -50,13 +50,11 @@ async def build( tree is (re)initialized from the full scoped store. ``make_client`` builds the LLM client used for synthesis from a profile name. - LLM work only happens when ``synthesize=True``: MEMORY.md and the skill/ - tree are synthesized from the per-source descriptions. Otherwise both are - left as ``None`` so the exporter renders MEMORY.md deterministically from - category summaries and falls back to its rule-based skill bypass. + LLM work only happens when ``synthesize=True``: MEMORY.md is synthesized + from the per-source descriptions. Otherwise it is left as ``None`` so the + exporter renders MEMORY.md deterministically from category summaries. """ memory_body: str | None = None - skills: dict[str, str] | None = None if self.config.synthesize: # The shared description trunk is the just-changed sources for an @@ -74,13 +72,11 @@ async def build( # files no longer produced). existing = await asyncio.to_thread(self.exporter.read_existing) if is_update else ExistingArtifacts() client = make_client(self.config.synthesis_llm_profile) - synthesized = await self.synthesizer.synthesize( + memory_body = await self.synthesizer.synthesize( descriptions, existing_memory=existing.memory_body, - existing_skills=existing.skills, chat=client.chat, ) - memory_body, skills = synthesized.memory_body, synthesized.skills async with self.lock: result: ExportResult = await asyncio.to_thread( @@ -88,7 +84,6 @@ async def build( database, where=where, memory_body=memory_body, - skills=skills, ) return result.to_dict() diff --git a/src/memu/app/retrieve.py b/src/memu/app/retrieve.py index 7dd4d2d1..1093cee4 100644 --- a/src/memu/app/retrieve.py +++ b/src/memu/app/retrieve.py @@ -13,20 +13,20 @@ from memu.prompts.retrieve.llm_resource_ranker import PROMPT as LLM_RESOURCE_RANKER_PROMPT from memu.prompts.retrieve.pre_retrieval_decision import SYSTEM_PROMPT as PRE_RETRIEVAL_SYSTEM_PROMPT from memu.prompts.retrieve.pre_retrieval_decision import USER_PROMPT as PRE_RETRIEVAL_USER_PROMPT -from memu.vector import cosine_topk from memu.workflow.step import WorkflowState, WorkflowStep logger = logging.getLogger(__name__) if TYPE_CHECKING: from memu.app.service import Context - from memu.app.settings import RetrieveConfig + from memu.app.settings import MemorizeConfig, RetrieveConfig from memu.database.interfaces import Database class RetrieveMixin: if TYPE_CHECKING: retrieve_config: RetrieveConfig + memorize_config: MemorizeConfig _run_workflow: Callable[..., Awaitable[WorkflowState]] _get_context: Callable[[], Context] _get_database: Callable[[], Database] @@ -114,79 +114,29 @@ def _build_rag_retrieve_workflow(self) -> list[WorkflowStep]: config={"chat_llm_profile": self.retrieve_config.sufficiency_check_llm_profile}, ), WorkflowStep( - step_id="route_category", - role="route_category", - handler=self._rag_route_category, - requires={"retrieve_category", "needs_retrieval", "active_query", "ctx", "store", "where"}, - produces={"category_hits", "category_summary_lookup", "query_vector"}, - capabilities={"vector"}, - config={"embed_llm_profile": "embedding"}, - ), - WorkflowStep( - step_id="sufficiency_after_category", - role="sufficiency_check", - handler=self._rag_category_sufficiency, + step_id="recall_lanes", + role="recall_lanes", + handler=self._rag_recall_lanes, requires={ "retrieve_category", + "retrieve_item", "needs_retrieval", "active_query", - "context_queries", - "category_hits", - "ctx", - "store", - "where", - }, - produces={"next_step_query", "proceed_to_items", "query_vector"}, - capabilities={"llm"}, - config={ - "chat_llm_profile": self.retrieve_config.sufficiency_check_llm_profile, - "embed_llm_profile": "embedding", - }, - ), - WorkflowStep( - step_id="recall_items", - role="recall_items", - handler=self._rag_recall_items, - requires={ - "needs_retrieval", - "proceed_to_items", "ctx", "store", "where", - "active_query", - "query_vector", }, - produces={"item_hits", "query_vector"}, + produces={"lane_hits", "query_vector"}, capabilities={"vector"}, config={"embed_llm_profile": "embedding"}, ), - WorkflowStep( - step_id="sufficiency_after_items", - role="sufficiency_check", - handler=self._rag_item_sufficiency, - requires={ - "needs_retrieval", - "active_query", - "context_queries", - "item_hits", - "ctx", - "store", - "where", - }, - produces={"next_step_query", "proceed_to_resources", "query_vector"}, - capabilities={"llm"}, - config={ - "chat_llm_profile": self.retrieve_config.sufficiency_check_llm_profile, - "embed_llm_profile": "embedding", - }, - ), WorkflowStep( step_id="recall_resources", role="recall_resources", handler=self._rag_recall_resources, requires={ "needs_retrieval", - "proceed_to_resources", + "retrieve_resource", "ctx", "store", "where", @@ -201,7 +151,16 @@ def _build_rag_retrieve_workflow(self) -> list[WorkflowStep]: step_id="build_context", role="build_context", handler=self._rag_build_context, - requires={"needs_retrieval", "original_query", "rewritten_query", "ctx", "store", "where"}, + requires={ + "needs_retrieval", + "original_query", + "rewritten_query", + "lane_hits", + "resource_hits", + "ctx", + "store", + "where", + }, produces={"response"}, capabilities=set(), ), @@ -256,158 +215,73 @@ async def _rag_route_intention(self, state: WorkflowState, step_context: Any) -> }) return state - async def _rag_route_category(self, state: WorkflowState, step_context: Any) -> WorkflowState: - if not state.get("retrieve_category") or not state.get("needs_retrieval"): - state["category_hits"] = [] - state["category_summary_lookup"] = {} - state["query_vector"] = None - return state - - embed_client = self._get_step_embedding_client(step_context) - store = state["store"] - where_filters = state.get("where") or {} - category_pool = store.resource_repo.list_resources(where_filters, lane="memory") - qvec = (await embed_client.embed([state["active_query"]]))[0] - hits, summary_lookup = await self._rank_categories_by_summary( - qvec, - self.retrieve_config.category.top_k, - state["ctx"], - store, - embed_client=embed_client, - categories=category_pool, - ) - state.update({ - "query_vector": qvec, - "category_hits": hits, - "category_summary_lookup": summary_lookup, - "category_pool": category_pool, - }) - return state - - async def _rag_category_sufficiency(self, state: WorkflowState, step_context: Any) -> WorkflowState: - if not state.get("needs_retrieval"): - state["proceed_to_items"] = False - return state - if not state.get("retrieve_category") or not state.get("sufficiency_check"): - state["proceed_to_items"] = True - return state - - retrieved_content = "" - store = state["store"] - where_filters = state.get("where") or {} - category_pool = state.get("category_pool") or store.resource_repo.list_resources(where_filters, lane="memory") - hits = state.get("category_hits") or [] - if hits: - retrieved_content = self._format_category_content( - hits, - state.get("category_summary_lookup", {}), - store, - categories=category_pool, - ) - - llm_client = self._get_step_llm_client(step_context) - needs_more, rewritten_query = await self._decide_if_retrieval_needed( - state["active_query"], - state["context_queries"], - retrieved_content=retrieved_content or "No content retrieved yet.", - llm_client=llm_client, - ) - state["next_step_query"] = rewritten_query - state["active_query"] = rewritten_query - state["proceed_to_items"] = needs_more - if needs_more: - embed_client = self._get_step_embedding_client(step_context) - state["query_vector"] = (await embed_client.embed([state["active_query"]]))[0] - return state - - async def _rag_recall_items(self, state: WorkflowState, step_context: Any) -> WorkflowState: - if not state.get("retrieve_item") or not state.get("needs_retrieval") or not state.get("proceed_to_items"): - state["item_hits"] = [] - return state + def _retrievable_lanes(self) -> list[str]: + """Enabled lanes that participate in retrieval, in canonical order.""" + from memu.database.models import RETRIEVAL_LANES - store = state["store"] - where_filters = state.get("where") or {} - items_pool = store.entry_repo.list_entries(where_filters, lane="memory") - qvec = state.get("query_vector") - if qvec is None: - embed_client = self._get_step_embedding_client(step_context) - qvec = (await embed_client.embed([state["active_query"]]))[0] - state["query_vector"] = qvec - state["item_hits"] = store.entry_repo.vector_search_entries( - qvec, - self.retrieve_config.item.top_k, - where=where_filters, - lane="memory", - ranking=self.retrieve_config.item.ranking, - recency_decay_days=self.retrieve_config.item.recency_decay_days, - ) - state["item_pool"] = items_pool - return state + enabled = self.memorize_config.enabled_lanes + return [lane for lane in RETRIEVAL_LANES if lane in enabled] - async def _rag_item_sufficiency(self, state: WorkflowState, step_context: Any) -> WorkflowState: + async def _rag_recall_lanes(self, state: WorkflowState, step_context: Any) -> WorkflowState: + """Single-pass per-lane recall: coarse lane docs + entries for each lane.""" + lane_hits: dict[str, dict[str, list[tuple[str, float]]]] = {} if not state.get("needs_retrieval"): - state["proceed_to_resources"] = False - return state - if not state.get("retrieve_item") or not state.get("sufficiency_check"): - state["proceed_to_resources"] = True + state["lane_hits"] = lane_hits + state["query_vector"] = None return state store = state["store"] where_filters = state.get("where") or {} - items_pool = state.get("item_pool") or store.entry_repo.list_entries(where_filters, lane="memory") - retrieved_content = "" - hits = state.get("item_hits") or [] - if hits: - retrieved_content = self._format_item_content(hits, store, items=items_pool) - - llm_client = self._get_step_llm_client(step_context) - needs_more, rewritten_query = await self._decide_if_retrieval_needed( - state["active_query"], - state["context_queries"], - retrieved_content=retrieved_content or "No content retrieved yet.", - llm_client=llm_client, - ) - state["next_step_query"] = rewritten_query - state["active_query"] = rewritten_query - state["proceed_to_resources"] = needs_more - if needs_more: - embed_client = self._get_step_embedding_client(step_context) - state["query_vector"] = (await embed_client.embed([state["active_query"]]))[0] + embed_client = self._get_step_embedding_client(step_context) + qvec = (await embed_client.embed([state["active_query"]]))[0] + state["query_vector"] = qvec + + retrieve_docs = bool(state.get("retrieve_category")) + retrieve_entries = bool(state.get("retrieve_item")) + for lane in self._retrievable_lanes(): + doc_hits: list[tuple[str, float]] = [] + entry_hits: list[tuple[str, float]] = [] + if retrieve_docs: + doc_hits = store.resource_repo.vector_search_resources( + qvec, self.retrieve_config.category.top_k, where=where_filters, lane=lane + ) + if retrieve_entries: + entry_hits = store.entry_repo.vector_search_entries( + qvec, + self.retrieve_config.item.top_k, + where=where_filters, + lane=lane, + ranking=self.retrieve_config.item.ranking, + recency_decay_days=self.retrieve_config.item.recency_decay_days, + ) + lane_hits[lane] = {"docs": doc_hits, "entries": entry_hits} + state["lane_hits"] = lane_hits return state async def _rag_recall_resources(self, state: WorkflowState, step_context: Any) -> WorkflowState: - if ( - not state.get("needs_retrieval") - or not state.get("retrieve_resource") - or not state.get("proceed_to_resources") - ): + if not state.get("needs_retrieval") or not state.get("retrieve_resource"): state["resource_hits"] = [] return state store = state["store"] where_filters = state.get("where") or {} - resource_pool = store.resource_repo.list_resources(where_filters) - state["resource_pool"] = resource_pool - if not resource_pool: - state["resource_hits"] = [] - return state - qvec = state.get("query_vector") if qvec is None: embed_client = self._get_step_embedding_client(step_context) qvec = (await embed_client.embed([state["active_query"]]))[0] state["query_vector"] = qvec state["resource_hits"] = store.resource_repo.vector_search_resources( - qvec, self.retrieve_config.resource.top_k, where=where_filters + qvec, self.retrieve_config.resource.top_k, where=where_filters, lane="source" ) return state def _rag_build_context(self, state: WorkflowState, _: Any) -> WorkflowState: - response = { + response: dict[str, Any] = { "needs_retrieval": bool(state.get("needs_retrieval")), "original_query": state["original_query"], "rewritten_query": state.get("rewritten_query", state["original_query"]), "next_step_query": state.get("next_step_query"), + "lanes": {}, "categories": [], "items": [], "resources": [], @@ -415,20 +289,22 @@ def _rag_build_context(self, state: WorkflowState, _: Any) -> WorkflowState: if state.get("needs_retrieval"): store = state["store"] where_filters = state.get("where") or {} - categories_pool = state.get("category_pool") or store.resource_repo.list_resources( - where_filters, lane="memory" - ) - items_pool = state.get("item_pool") or store.entry_repo.list_entries(where_filters, lane="memory") - resources_pool = state.get("resource_pool") or store.resource_repo.list_resources(where_filters) - response["categories"] = self._materialize_hits( - state.get("category_hits", []), - categories_pool, - ) - response["items"] = self._materialize_hits(state.get("item_hits", []), items_pool) - response["resources"] = self._materialize_hits( - state.get("resource_hits", []), - resources_pool, - ) + lane_hits: dict[str, dict[str, list[tuple[str, float]]]] = state.get("lane_hits", {}) + lanes_out: dict[str, dict[str, Any]] = {} + for lane, hits in lane_hits.items(): + docs_pool = store.resource_repo.list_resources(where_filters, lane=lane) + entries_pool = store.entry_repo.list_entries(where_filters, lane=lane) + lanes_out[lane] = { + "categories": self._materialize_hits(hits.get("docs", []), docs_pool), + "items": self._materialize_hits(hits.get("entries", []), entries_pool), + } + resources_pool = store.resource_repo.list_resources(where_filters, lane="source") + response["lanes"] = lanes_out + response["resources"] = self._materialize_hits(state.get("resource_hits", []), resources_pool) + # Backward-compatible top-level view: the memory lane. + memory_out = lanes_out.get("memory", {}) + response["categories"] = memory_out.get("categories", []) + response["items"] = memory_out.get("items", []) state["response"] = response return state @@ -444,62 +320,26 @@ def _build_llm_retrieve_workflow(self) -> list[WorkflowStep]: config={"llm_profile": self.retrieve_config.sufficiency_check_llm_profile}, ), WorkflowStep( - step_id="route_category", - role="route_category", - handler=self._llm_route_category, + step_id="recall_lanes", + role="recall_lanes", + handler=self._llm_recall_lanes, requires={"needs_retrieval", "active_query", "ctx", "store", "where"}, - produces={"category_hits"}, + produces={"lane_hits"}, capabilities={"llm"}, config={"llm_profile": self.retrieve_config.llm_ranking_llm_profile}, ), - WorkflowStep( - step_id="sufficiency_after_category", - role="sufficiency_check", - handler=self._llm_category_sufficiency, - requires={"needs_retrieval", "active_query", "context_queries", "category_hits"}, - produces={"next_step_query", "proceed_to_items"}, - capabilities={"llm"}, - config={"llm_profile": self.retrieve_config.sufficiency_check_llm_profile}, - ), - WorkflowStep( - step_id="recall_items", - role="recall_items", - handler=self._llm_recall_items, - requires={ - "needs_retrieval", - "proceed_to_items", - "ctx", - "store", - "where", - "active_query", - "category_hits", - }, - produces={"item_hits"}, - capabilities={"llm"}, - config={"llm_profile": self.retrieve_config.llm_ranking_llm_profile}, - ), - WorkflowStep( - step_id="sufficiency_after_items", - role="sufficiency_check", - handler=self._llm_item_sufficiency, - requires={"needs_retrieval", "active_query", "context_queries", "item_hits"}, - produces={"next_step_query", "proceed_to_resources"}, - capabilities={"llm"}, - config={"llm_profile": self.retrieve_config.sufficiency_check_llm_profile}, - ), WorkflowStep( step_id="recall_resources", role="recall_resources", handler=self._llm_recall_resources, requires={ "needs_retrieval", - "proceed_to_resources", + "retrieve_resource", "active_query", "ctx", "store", "where", - "item_hits", - "category_hits", + "lane_hits", }, produces={"resource_hits"}, capabilities={"llm"}, @@ -509,7 +349,7 @@ def _build_llm_retrieve_workflow(self) -> list[WorkflowStep]: step_id="build_context", role="build_context", handler=self._llm_build_context, - requires={"needs_retrieval", "original_query", "rewritten_query"}, + requires={"needs_retrieval", "original_query", "rewritten_query", "lane_hits", "resource_hits"}, produces={"response"}, capabilities=set(), ), @@ -548,182 +388,93 @@ async def _llm_route_intention(self, state: WorkflowState, step_context: Any) -> }) return state - async def _llm_route_category(self, state: WorkflowState, step_context: Any) -> WorkflowState: - if not state.get("needs_retrieval"): - state["category_hits"] = [] - return state - llm_client = self._get_step_llm_client(step_context) - store = state["store"] - where_filters = state.get("where") or {} - category_pool = store.resource_repo.list_resources(where_filters, lane="memory") - hits = await self._llm_rank_categories( - state["active_query"], - self.retrieve_config.category.top_k, - state["ctx"], - store, - llm_client=llm_client, - categories=category_pool, - ) - state["category_hits"] = hits - state["category_pool"] = category_pool - return state - - async def _llm_category_sufficiency(self, state: WorkflowState, step_context: Any) -> WorkflowState: + async def _llm_recall_lanes(self, state: WorkflowState, step_context: Any) -> WorkflowState: + """Single-pass per-lane LLM recall: rank coarse lane docs + entries per lane.""" + lane_hits: dict[str, dict[str, list[dict[str, Any]]]] = {} if not state.get("needs_retrieval"): - state["proceed_to_items"] = False - return state - if not state.get("retrieve_category") or not state.get("sufficiency_check"): - state["proceed_to_items"] = True - return state - - retrieved_content = "" - hits = state.get("category_hits") or [] - if hits: - retrieved_content = self._format_llm_category_content(hits) - - llm_client = self._get_step_llm_client(step_context) - needs_more, rewritten_query = await self._decide_if_retrieval_needed( - state["active_query"], - state["context_queries"], - retrieved_content=retrieved_content or "No content retrieved yet.", - llm_client=llm_client, - ) - state["next_step_query"] = rewritten_query - state["active_query"] = rewritten_query - state["proceed_to_items"] = needs_more - return state - - async def _llm_recall_items(self, state: WorkflowState, step_context: Any) -> WorkflowState: - if not state.get("needs_retrieval") or not state.get("proceed_to_items"): - state["item_hits"] = [] + state["lane_hits"] = lane_hits return state - where_filters = state.get("where") or {} - category_hits = state.get("category_hits", []) - category_ids = [cat["id"] for cat in category_hits] llm_client = self._get_step_llm_client(step_context) store = state["store"] - - use_refs = getattr(self.retrieve_config.item, "use_category_references", False) - ref_ids: list[str] = [] - if use_refs and category_hits: - # Extract all ref_ids from category summaries - from memu.utils.references import extract_references - - for cat in category_hits: - summary = cat.get("summary") or "" - ref_ids.extend(extract_references(summary)) - if ref_ids: - # Query items by ref_ids - items_pool = store.entry_repo.list_entries_by_ref_ids(ref_ids, where_filters) - else: - items_pool = store.entry_repo.list_entries(where_filters, lane="memory") - - relations = store.resource_entry_repo.list_relations(where_filters) - category_pool = state.get("category_pool") or store.resource_repo.list_resources(where_filters, lane="memory") - state["item_hits"] = await self._llm_rank_items( - state["active_query"], - self.retrieve_config.item.top_k, - category_ids, - state.get("category_hits", []), - state["ctx"], - store, - llm_client=llm_client, - categories=category_pool, - items=items_pool, - relations=relations, - ) - state["item_pool"] = items_pool - state["relation_pool"] = relations - return state - - async def _llm_item_sufficiency(self, state: WorkflowState, step_context: Any) -> WorkflowState: - if not state.get("needs_retrieval"): - state["proceed_to_resources"] = False - return state - if not state.get("retrieve_item") or not state.get("sufficiency_check"): - state["proceed_to_resources"] = True - return state - - retrieved_content = "" - hits = state.get("item_hits") or [] - if hits: - retrieved_content = self._format_llm_item_content(hits) - - llm_client = self._get_step_llm_client(step_context) - needs_more, rewritten_query = await self._decide_if_retrieval_needed( - state["active_query"], - state["context_queries"], - retrieved_content=retrieved_content or "No content retrieved yet.", - llm_client=llm_client, - ) - state["next_step_query"] = rewritten_query - state["active_query"] = rewritten_query - state["proceed_to_resources"] = needs_more + where_filters = state.get("where") or {} + active_query = state["active_query"] + + for lane in self._retrievable_lanes(): + docs_pool = store.resource_repo.list_resources(where_filters, lane=lane) + doc_hits = await self._llm_rank_categories( + active_query, + self.retrieve_config.category.top_k, + state["ctx"], + store, + llm_client=llm_client, + categories=docs_pool, + ) + entries_pool = store.entry_repo.list_entries(where_filters, lane=lane) + entry_hits = await self._llm_rank_entries( + active_query, + self.retrieve_config.item.top_k, + doc_hits, + state["ctx"], + store, + llm_client=llm_client, + entries=entries_pool, + ) + lane_hits[lane] = {"categories": doc_hits, "items": entry_hits} + state["lane_hits"] = lane_hits return state async def _llm_recall_resources(self, state: WorkflowState, step_context: Any) -> WorkflowState: - if not state.get("needs_retrieval") or not state.get("proceed_to_resources"): + if not state.get("needs_retrieval") or not state.get("retrieve_resource"): state["resource_hits"] = [] return state llm_client = self._get_step_llm_client(step_context) store = state["store"] where_filters = state.get("where") or {} - resource_pool = store.resource_repo.list_resources(where_filters) - items_pool = state.get("item_pool") or store.entry_repo.list_entries(where_filters, lane="memory") + resource_pool = store.resource_repo.list_resources(where_filters, lane="source") + lane_hits: dict[str, dict[str, list[dict[str, Any]]]] = state.get("lane_hits", {}) + # Aggregate entry/doc hits across lanes for resource ranking context. + all_item_hits: list[dict[str, Any]] = [] + all_category_hits: list[dict[str, Any]] = [] + for hits in lane_hits.values(): + all_item_hits.extend(hits.get("items", [])) + all_category_hits.extend(hits.get("categories", [])) + items_pool = store.entry_repo.list_entries(where_filters) state["resource_hits"] = await self._llm_rank_resources( state["active_query"], self.retrieve_config.resource.top_k, - state.get("category_hits", []), - state.get("item_hits", []), + all_category_hits, + all_item_hits, state["ctx"], store, llm_client=llm_client, items=items_pool, resources=resource_pool, ) - state["resource_pool"] = resource_pool return state def _llm_build_context(self, state: WorkflowState, _: Any) -> WorkflowState: - response = { + response: dict[str, Any] = { "needs_retrieval": bool(state.get("needs_retrieval")), "original_query": state["original_query"], "rewritten_query": state.get("rewritten_query", state["original_query"]), "next_step_query": state.get("next_step_query"), + "lanes": {}, "categories": [], "items": [], "resources": [], } if state.get("needs_retrieval"): - response["categories"] = list(state.get("category_hits") or []) - response["items"] = list(state.get("item_hits") or []) + lane_hits: dict[str, dict[str, list[dict[str, Any]]]] = state.get("lane_hits", {}) + response["lanes"] = lane_hits response["resources"] = list(state.get("resource_hits") or []) + memory_out = lane_hits.get("memory", {}) + response["categories"] = list(memory_out.get("categories", [])) + response["items"] = list(memory_out.get("items", [])) state["response"] = response return state - async def _rank_categories_by_summary( - self, - query_vec: list[float], - top_k: int, - ctx: Context, - store: Database, - embed_client: Any | None = None, - categories: Mapping[str, Any] | None = None, - ) -> tuple[list[tuple[str, float]], dict[str, str]]: - category_pool = categories if categories is not None else store.resource_repo.list_resources(lane="memory") - entries = [(cid, cat.summary) for cid, cat in category_pool.items() if cat.summary] - if not entries: - return [], {} - summary_texts = [summary for _, summary in entries] - client = embed_client or self._get_embedding_client() - summary_embeddings = await client.embed(summary_texts) - corpus = [(cid, emb) for (cid, _), emb in zip(entries, summary_embeddings, strict=True)] - hits = cosine_topk(query_vec, corpus, k=top_k) - summary_lookup = dict(entries) - return hits, summary_lookup - async def _decide_if_retrieval_needed( self, query: str, @@ -856,35 +607,6 @@ def _materialize_hits(self, hits: Sequence[tuple[str, float]], pool: dict[str, A out.append(data) return out - def _format_category_content( - self, - hits: list[tuple[str, float]], - summaries: dict[str, str], - store: Database, - categories: Mapping[str, Any] | None = None, - ) -> str: - category_pool = categories if categories is not None else store.resource_repo.list_resources(lane="memory") - lines = [] - for cid, score in hits: - cat = category_pool.get(cid) - if not cat: - continue - summary = summaries.get(cid) or cat.summary or "" - lines.append(f"Category: {cat.title}\nSummary: {summary}\nScore: {score:.3f}") - return "\n\n".join(lines).strip() - - def _format_item_content( - self, hits: list[tuple[str, float]], store: Database, items: Mapping[str, Any] | None = None - ) -> str: - item_pool = items if items is not None else store.entry_repo.list_entries(lane="memory") - lines = [] - for iid, score in hits: - item = item_pool.get(iid) - if not item: - continue - lines.append(f"Memory Item ({item.entry_kind}): {item.text}\nScore: {score:.3f}") - return "\n\n".join(lines).strip() - def _format_categories_for_llm( self, store: Database, @@ -943,7 +665,7 @@ def _format_items_for_llm( lines = [] for item in items_to_format: lines.append(f"ID: {item.id}") - lines.append(f"Type: {item.entry_kind}") + lines.append(f"Type: {item.entry_type}") lines.append(f"Summary: {item.text}") lines.append("---") @@ -1009,44 +731,44 @@ async def _llm_rank_categories( llm_response = await client.chat(prompt) return self._parse_llm_category_response(llm_response, store, categories=category_pool) - async def _llm_rank_items( + async def _llm_rank_entries( self, query: str, top_k: int, - category_ids: list[str], - category_hits: list[dict[str, Any]], + doc_hits: list[dict[str, Any]], ctx: Context, store: Database, llm_client: Any | None = None, - categories: Mapping[str, Any] | None = None, - items: Mapping[str, Any] | None = None, - relations: Sequence[Any] | None = None, + entries: Mapping[str, Any] | None = None, ) -> list[dict[str, Any]]: - """Use LLM to rank memory items from relevant categories""" - if not category_ids: - logger.debug("[LLM Rank Items] No category_ids provided") + """Use LLM to rank a single lane's entries directly (no per-doc filtering). + + Coarse lane docs (``doc_hits``) are passed only as relevance context; all + entries of the lane are candidates so retrieval stays single-pass. + """ + entry_pool = entries if entries is not None else store.entry_repo.list_entries() + if not entry_pool: return [] - item_pool = items if items is not None else store.entry_repo.list_entries(lane="memory") - items_data = self._format_items_for_llm(store, category_ids, items=item_pool, relations=relations) + items_data = self._format_items_for_llm(store, items=entry_pool) if items_data == "No memory items available.": return [] - # Format relevant categories for context - relevant_categories_info = "\n".join([ - f"- {cat['name']}: {cat.get('summary', cat.get('description', ''))}" for cat in category_hits + relevant_docs_info = "\n".join([ + f"- {doc.get('title') or doc.get('name', '')}: {doc.get('summary') or doc.get('description', '')}" + for doc in doc_hits ]) prompt = LLM_ITEM_RANKER_PROMPT.format( query=self._escape_prompt_value(query), top_k=top_k, - relevant_categories=self._escape_prompt_value(relevant_categories_info), + relevant_categories=self._escape_prompt_value(relevant_docs_info), items_data=self._escape_prompt_value(items_data), ) client = llm_client or self._get_llm_client() llm_response = await client.chat(prompt) - return self._parse_llm_item_response(llm_response, store, items=item_pool) + return self._parse_llm_item_response(llm_response, store, items=entry_pool) async def _llm_rank_resources( self, @@ -1076,7 +798,7 @@ async def _llm_rank_resources( context_parts = [] if category_hits: context_parts.append("Relevant Categories:") - context_parts.extend([f"- {cat['name']}" for cat in category_hits]) + context_parts.extend([f"- {cat.get('title') or cat.get('name', '')}" for cat in category_hits]) if item_hits: context_parts.append("\nRelevant Memory Items:") context_parts.extend([f"- {item.get('summary', '')[:100]}..." for item in item_hits[:3]]) @@ -1164,18 +886,3 @@ def _parse_llm_resource_response( logger.warning(f"Failed to parse LLM resource ranking response: {e}") return results - - def _format_llm_category_content(self, hits: list[dict[str, Any]]) -> str: - """Format LLM-ranked category content for judger""" - lines = [] - for cat in hits: - summary = cat.get("summary", "") or cat.get("description", "") - lines.append(f"Category: {cat['name']}\nSummary: {summary}") - return "\n\n".join(lines).strip() - - def _format_llm_item_content(self, hits: list[dict[str, Any]]) -> str: - """Format LLM-ranked item content for judger""" - lines = [] - for item in hits: - lines.append(f"Memory Item ({item['memory_type']}): {item['summary']}") - return "\n\n".join(lines).strip() diff --git a/src/memu/app/service.py b/src/memu/app/service.py index e00feaf6..fe37d379 100644 --- a/src/memu/app/service.py +++ b/src/memu/app/service.py @@ -47,11 +47,26 @@ TConfigModel = TypeVar("TConfigModel", bound=BaseModel) +@dataclass +class LaneState: + """Per-lane grouping state cached across a service's lifetime. + + For an adaptive lane (memory/skill) this tracks the known group docs (the + former "categories"): whether seeds have been initialized, the ordered doc + ids, and a name->id index for resolving extractor-proposed group names. + """ + + ready: bool = False + doc_ids: list[str] = field(default_factory=list) + name_to_id: dict[str, str] = field(default_factory=dict) + + @dataclass class Context: - categories_ready: bool = False - category_ids: list[str] = field(default_factory=list) - category_name_to_id: dict[str, str] = field(default_factory=dict) + lanes: dict[str, LaneState] = field(default_factory=dict) + + def lane(self, name: str) -> LaneState: + return self.lanes.setdefault(name, LaneState()) class MemoryService(MemorizeMixin, RetrieveMixin, CRUDMixin): @@ -79,11 +94,21 @@ def __init__( self.memory_files_config = self._validate_config(memory_files_config, MemoryFilesConfig) self.fs = LocalFS(self.blob_config.resources_dir) - self.category_configs: list[CategoryConfig] = list(self.memorize_config.memory_categories or []) - self.category_config_map: dict[str, CategoryConfig] = {cfg.name: cfg for cfg in self.category_configs} - self._category_prompt_str = self._format_categories_for_prompt(self.category_configs) + # Per-lane wiring (ADR 0006): each lane carries its own seed group docs, + # extractor "existing groups" prompt string, and summary-prompt overrides. + self.lane_configs = self.memorize_config.lanes + self.lane_category_config_maps: dict[str, dict[str, CategoryConfig]] = { + lane: {cfg.name: cfg for cfg in lc.seed_categories} for lane, lc in self.lane_configs.items() + } + self.lane_prompt_strs: dict[str, str] = { + lane: self._format_categories_for_prompt(lc.seed_categories) for lane, lc in self.lane_configs.items() + } - self._context = Context(categories_ready=not bool(self.category_configs)) + self._context = Context() + for lane, lc in self.memorize_config.enabled_lanes.items(): + # per_resource lanes never need group-doc initialization; adaptive lanes + # are ready immediately only when they have no seeds to materialize. + self._context.lane(lane).ready = lc.grouping != "adaptive" or not lc.seed_categories self.database: Database = build_database( config=self.database_config, diff --git a/src/memu/app/settings.py b/src/memu/app/settings.py index d0e667a1..4161bb63 100644 --- a/src/memu/app/settings.py +++ b/src/memu/app/settings.py @@ -9,12 +9,9 @@ from memu.prompts.category_summary import ( PROMPT as CATEGORY_SUMMARY_PROMPT, ) -from memu.prompts.memory_type import ( +from memu.prompts.entry_type import ( DEFAULT_MEMORY_CUSTOM_PROMPT_ORDINAL, - DEFAULT_MEMORY_TYPES, -) -from memu.prompts.memory_type import ( - PROMPTS as DEFAULT_MEMORY_TYPE_PROMPTS, + LANE_ENTRY_TYPES, ) @@ -27,14 +24,6 @@ def normalize_value(v: str) -> str: Normalize = BeforeValidator(normalize_value) -def _default_memory_types() -> list[str]: - return list(DEFAULT_MEMORY_TYPES) - - -def _default_memory_type_prompts() -> "dict[str, str | CustomPrompt]": - return dict(DEFAULT_MEMORY_TYPE_PROMPTS) - - class PromptBlock(BaseModel): label: str | None = None ordinal: int = Field(default=0) @@ -58,7 +47,7 @@ def complete_prompt_blocks(prompt: CustomPrompt, default_blocks: Mapping[str, in return prompt -CompleteMemoryTypePrompt = AfterValidator(lambda v: complete_prompt_blocks(v, DEFAULT_MEMORY_CUSTOM_PROMPT_ORDINAL)) +CompleteEntryTypePrompt = AfterValidator(lambda v: complete_prompt_blocks(v, DEFAULT_MEMORY_CUSTOM_PROMPT_ORDINAL)) CompleteCategoryPrompt = AfterValidator(lambda v: complete_prompt_blocks(v, DEFAULT_CATEGORY_SUMMARY_PROMPT_ORDINAL)) @@ -71,6 +60,71 @@ class CategoryConfig(BaseModel): summary_prompt: str | Annotated[CustomPrompt, CompleteCategoryPrompt] | None = None +class LaneConfig(BaseModel): + """Per-lane extraction/grouping configuration (see ADR 0006). + + Each lane (``index`` / ``memory`` / ``skill``) is a structurally identical + processing track over the shared per-resource canonical text; lanes differ + only in their entry-type set, extraction prompts, and how entries are grouped + into coarse lane documents. + """ + + enabled: bool = Field(default=True, description="Whether this lane is processed during memorize/retrieve.") + grouping: Annotated[Literal["adaptive", "per_resource"], Normalize] = Field( + default="adaptive", + description=( + "How entries map to coarse lane documents. 'adaptive': the extractor proposes " + "group names and a summarized doc is synthesized per group (memory/skill). " + "'per_resource': one coarse doc per source resource, no grouping (index)." + ), + ) + entry_types: list[str] = Field( + default_factory=list, + description="Ordered list of entry types extracted for this lane.", + ) + entry_type_prompts: dict[str, str | Annotated[CustomPrompt, CompleteEntryTypePrompt]] = Field( + default_factory=dict, + description="User prompt overrides per entry type; falls back to the built-in template per type.", + ) + extract_llm_profile: str = Field(default="default", description="LLM profile used for this lane's extraction.") + # Adaptive-grouping (memory/skill) coarse-doc summary synthesis. Ignored for per_resource lanes. + seed_categories: list[CategoryConfig] = Field( + default_factory=list, + description="Optional seed group docs for an adaptive lane (e.g. memory categories).", + ) + summary_prompt: str | Annotated[CustomPrompt, CompleteCategoryPrompt] = Field( + default=CATEGORY_SUMMARY_PROMPT, + description="System prompt for synthesizing an adaptive lane's group-doc summary.", + ) + summary_target_length: int = Field(default=400, description="Target max length for the group-doc summary.") + summary_llm_profile: str = Field(default="default", description="LLM profile for group-doc summary synthesis.") + enable_item_references: bool = Field( + default=False, + description="Enable inline [ref:ITEM_ID] citations in this lane's group-doc summaries.", + ) + enable_item_reinforcement: bool = Field( + default=False, + description="Enable reinforcement tracking for this lane's entries.", + ) + + +def _default_lane_configs() -> "dict[str, LaneConfig]": + return { + "index": LaneConfig( + grouping="per_resource", + entry_types=list(LANE_ENTRY_TYPES["index"]), + ), + "memory": LaneConfig( + grouping="adaptive", + entry_types=list(LANE_ENTRY_TYPES["memory"]), + ), + "skill": LaneConfig( + grouping="adaptive", + entry_types=list(LANE_ENTRY_TYPES["skill"]), + ), + } + + class LazyLLMSource(BaseModel): source: str | None = Field(default=None, description="default source for lazyllm client backend") llm_source: str | None = Field(default=None, description="LLM source for lazyllm client backend") @@ -326,14 +380,14 @@ class MemoryFilesConfig(BaseModel): output_dir: str = Field( default="./data/memory", description=( - "Directory where the memory markdown tree (the INDEX.md/MEMORY.md/SKILL.md root " - "indexes plus the resource/, memory/, and skill/ directories) is written." + "Directory where the memory markdown tree (the INDEX.md/MEMORY.md root " + "indexes plus the resource/ and memory/ directories) is written." ), ) synthesize: bool = Field( default=False, description=( - "Synthesize MEMORY.md and skill docs from per-source descriptions via the LLM " + "Synthesize MEMORY.md from per-source descriptions via the LLM " "instead of rendering already-extracted records. INDEX.md stays deterministic." ), ) @@ -412,42 +466,24 @@ class MemorizeConfig(BaseModel): default="default", description="LLM profile whose provider/credentials back the VLM client used for image/video vision.", ) - memory_types: list[str] = Field( - default_factory=_default_memory_types, - description="Ordered list of memory types (profile/event/knowledge/behavior by default).", - ) - memory_type_prompts: dict[str, str | Annotated[CustomPrompt, CompleteMemoryTypePrompt]] = Field( - default_factory=_default_memory_type_prompts, - description="User prompt overrides for each memory type extraction.", - ) - memory_extract_llm_profile: str = Field(default="default", description="LLM profile for memory extract.") - memory_categories: list[CategoryConfig] = Field( - default_factory=list, + lanes: dict[str, LaneConfig] = Field( + default_factory=_default_lane_configs, description=( - "Optional seed categories. The kernel presets no taxonomy: categories are " - "discovered adaptively from ingested content. Provide seeds only to guide " - "(not constrain) the taxonomy; an empty list means fully open/adaptive." + "Per-lane extraction/grouping configuration. Defaults to index/memory/skill " + "(see ADR 0006); each lane runs the same pipeline with its own entry types, " + "prompts, and grouping cardinality." ), ) - # default_category_summary_prompt: str | CustomPrompt = Field( - default_category_summary_prompt: str | Annotated[CustomPrompt, CompleteCategoryPrompt] = Field( - default=CATEGORY_SUMMARY_PROMPT, - description="Default system prompt for auto-generated category summaries.", - ) - default_category_summary_target_length: int = Field( - default=400, - description="Target max length for auto-generated category summaries.", - ) - category_update_llm_profile: str = Field(default="default", description="LLM profile for category summary.") - # Reference tracking for category summaries - enable_item_references: bool = Field( - default=False, - description="Enable inline [ref:ITEM_ID] citations in category summaries linking to source memory items.", - ) - enable_item_reinforcement: bool = Field( - default=False, - description="Enable reinforcement tracking for memory items.", - ) + + @property + def enabled_lanes(self) -> dict[str, LaneConfig]: + """Lanes that are turned on, in a stable order (index, memory, skill first).""" + order = ["index", "memory", "skill"] + ordered = sorted( + self.lanes.items(), + key=lambda kv: (order.index(kv[0]) if kv[0] in order else len(order), kv[0]), + ) + return {name: cfg for name, cfg in ordered if cfg.enabled} class PatchConfig(BaseModel): diff --git a/src/memu/database/inmemory/repositories/entry_repo.py b/src/memu/database/inmemory/repositories/entry_repo.py index b0c667fa..ab27c8be 100644 --- a/src/memu/database/inmemory/repositories/entry_repo.py +++ b/src/memu/database/inmemory/repositories/entry_repo.py @@ -19,19 +19,17 @@ def __init__(self, *, state: InMemoryState, entry_model: type[Entry]) -> None: self.entry_model = entry_model self.entries: dict[str, Entry] = self._state.entries - def list_entries( - self, where: Mapping[str, Any] | None = None, *, lane: str | None = None - ) -> dict[str, Entry]: - result = self.entries if not where else { - eid: entry for eid, entry in self.entries.items() if matches_where(entry, where) - } + def list_entries(self, where: Mapping[str, Any] | None = None, *, lane: str | None = None) -> dict[str, Entry]: + result = ( + self.entries + if not where + else {eid: entry for eid, entry in self.entries.items() if matches_where(entry, where)} + ) if lane is not None: result = {eid: entry for eid, entry in result.items() if getattr(entry, "lane", None) == lane} return dict(result) - def list_entries_by_ref_ids( - self, ref_ids: list[str], where: Mapping[str, Any] | None = None - ) -> dict[str, Entry]: + def list_entries_by_ref_ids(self, ref_ids: list[str], where: Mapping[str, Any] | None = None) -> dict[str, Entry]: if not ref_ids: return {} ref_id_set = set(ref_ids) @@ -44,9 +42,7 @@ def list_entries_by_ref_ids( result[eid] = entry return result - def clear_entries( - self, where: Mapping[str, Any] | None = None, *, lane: str | None = None - ) -> dict[str, Entry]: + def clear_entries(self, where: Mapping[str, Any] | None = None, *, lane: str | None = None) -> dict[str, Entry]: if not where and lane is None: matches = self.entries.copy() self.entries.clear() @@ -70,41 +66,34 @@ def create_entry( *, lane: str, source_id: str | None, - entry_kind: str, + entry_type: str, text: str, embedding: list[float], user_data: dict[str, Any], source_path: str | None = None, reinforce: bool = False, - tool_record: dict[str, Any] | None = None, ) -> Entry: - if reinforce and entry_kind != "tool": + if reinforce: return self.create_entry_reinforce( lane=lane, source_id=source_id, - entry_kind=entry_kind, + entry_type=entry_type, text=text, embedding=embedding, user_data=user_data, source_path=source_path, ) - extra: dict[str, Any] = {} - if tool_record: - for key in ("when_to_use", "metadata", "tool_calls"): - if tool_record.get(key) is not None: - extra[key] = tool_record[key] - eid = str(uuid.uuid4()) entry = self.entry_model( id=eid, lane=lane, source_id=source_id, source_path=source_path, - entry_kind=entry_kind, + entry_type=entry_type, text=text, embedding=embedding, - extra=extra if extra else {}, + extra={}, **user_data, ) self.entries[eid] = entry @@ -115,13 +104,13 @@ def create_entry_reinforce( *, lane: str, source_id: str | None, - entry_kind: str, + entry_type: str, text: str, embedding: list[float], user_data: dict[str, Any], source_path: str | None = None, ) -> Entry: - content_hash = compute_content_hash(text, entry_kind) + content_hash = compute_content_hash(text, entry_type) existing = self._find_by_hash(content_hash, user_data) if existing: current_extra = existing.extra or {} @@ -147,7 +136,7 @@ def create_entry_reinforce( lane=lane, source_id=source_id, source_path=source_path, - entry_kind=entry_kind, + entry_type=entry_type, text=text, embedding=embedding, extra=entry_extra, @@ -210,33 +199,25 @@ def update_entry( self, *, entry_id: str, - entry_kind: str | None = None, + entry_type: str | None = None, text: str | None = None, embedding: list[float] | None = None, extra: dict[str, Any] | None = None, - tool_record: dict[str, Any] | None = None, ) -> Entry: entry = self.entries.get(entry_id) if entry is None: msg = f"Entry with id {entry_id} not found" raise KeyError(msg) - if entry_kind is not None: - entry.entry_kind = entry_kind + if entry_type is not None: + entry.entry_type = entry_type if text is not None: entry.text = text if embedding is not None: entry.embedding = embedding - current_extra = entry.extra or {} if extra is not None: - current_extra = {**current_extra, **extra} - if tool_record is not None: - for key in ("when_to_use", "metadata", "tool_calls"): - if tool_record.get(key) is not None: - current_extra[key] = tool_record[key] - if extra is not None or tool_record is not None: - entry.extra = current_extra + entry.extra = {**(entry.extra or {}), **extra} self.entries[entry_id] = entry return entry diff --git a/src/memu/database/inmemory/repositories/resource_entry_repo.py b/src/memu/database/inmemory/repositories/resource_entry_repo.py index 7eec1010..76d342d7 100644 --- a/src/memu/database/inmemory/repositories/resource_entry_repo.py +++ b/src/memu/database/inmemory/repositories/resource_entry_repo.py @@ -25,9 +25,7 @@ def link_entry_resource(self, entry_id: str, resource_id: str, user_data: dict[s for rel in self.relations: if rel.entry_id == entry_id and rel.resource_id == resource_id: return rel - rel = self.resource_entry_model( - id=str(uuid.uuid4()), entry_id=entry_id, resource_id=resource_id, **user_data - ) + rel = self.resource_entry_model(id=str(uuid.uuid4()), entry_id=entry_id, resource_id=resource_id, **user_data) self.relations.append(rel) return rel diff --git a/src/memu/database/models.py b/src/memu/database/models.py index c3fc788e..647a637b 100644 --- a/src/memu/database/models.py +++ b/src/memu/database/models.py @@ -1,7 +1,6 @@ from __future__ import annotations import hashlib -import json import uuid from datetime import datetime from os.path import basename @@ -10,30 +9,33 @@ import pendulum from pydantic import BaseModel, ConfigDict, Field -# Sub-type of a memory-lane entry (kept for prompt routing / backward semantics). -MemoryType = Literal["profile", "event", "knowledge", "behavior", "skill", "tool"] +# Sub-type of a lane entry; selects the extraction prompt during memorize. The +# concrete set is per-lane (see ``LANE_ENTRY_TYPES``): memory uses +# profile/event/knowledge, skill uses tool/log, index uses description. This +# alias keeps the memory-lane default for typing/back-compat purposes. +EntryType = Literal["profile", "event", "knowledge"] # A lane is one of the parallel, structurally identical processing tracks that # share the same Resource -> canonical-text trunk: # - "source": raw input artifacts (conversation/document/image/video/audio) -# - "index": per-resource catalog/description docs -# - "memory": grouped memory docs (the former "category") -# - "skill": grouped reusable-skill docs -# index/memory/skill are the three retrievable lanes; "source" holds raw inputs. +# - "index": per-resource catalog/description docs (1:1, per_resource grouping) +# - "memory": grouped memory docs (the former "category", adaptive grouping) +# - "skill": grouped skill docs (tool/log entries, adaptive grouping) +# index/memory/skill are the retrievable lanes; "source" holds raw inputs. Lane = Literal["source", "index", "memory", "skill"] SOURCE_LANE: Lane = "source" RETRIEVAL_LANES: tuple[Lane, ...] = ("index", "memory", "skill") MARKDOWN_MODALITY = "markdown" -def compute_content_hash(text: str, entry_kind: str) -> str: +def compute_content_hash(text: str, entry_type: str) -> str: """Generate a stable hash for entry deduplication. Operates on post-extraction content. Normalizes whitespace to absorb minor formatting differences ("I love coffee" vs "I love coffee"). """ normalized = " ".join(text.lower().split()) - content = f"{entry_kind}:{normalized}" + content = f"{entry_type}:{normalized}" return hashlib.sha256(content.encode()).hexdigest()[:16] @@ -45,31 +47,6 @@ class BaseRecord(BaseModel): updated_at: datetime = Field(default_factory=lambda: pendulum.now("UTC")) -class ToolCallResult(BaseModel): - """Represents the result of a tool invocation for Tool Memory.""" - - tool_name: str = Field(..., description="Name of the tool that was called") - input: dict[str, Any] | str = Field(default="", description="Tool input parameters") - output: str = Field(default="", description="Tool output result") - success: bool = Field(default=True, description="Whether the tool invocation succeeded") - time_cost: float = Field(default=0.0, description="Time consumed by the tool invocation in seconds") - token_cost: int = Field(default=-1, description="Token consumption of the tool (-1 if unknown)") - score: float = Field(default=0.0, description="Quality score from 0.0 to 1.0") - call_hash: str = Field(default="", description="Hash of input+output for deduplication") - created_at: datetime = Field(default_factory=lambda: pendulum.now("UTC")) - - def generate_hash(self) -> str: - """Generate MD5 hash from tool input and output for deduplication.""" - input_str = json.dumps(self.input, sort_keys=True) if isinstance(self.input, dict) else str(self.input) - combined = f"{self.tool_name}|{input_str}|{self.output}" - return hashlib.md5(combined.encode("utf-8"), usedforsecurity=False).hexdigest() - - def ensure_hash(self) -> None: - """Ensure call_hash is set, generate if empty.""" - if not self.call_hash: - self.call_hash = self.generate_hash() - - class Resource(BaseRecord): """A node in the unified store: either a raw input or a generated lane doc. @@ -111,14 +88,14 @@ def source_path(self) -> str: class Entry(BaseRecord): - """The searchable atom of a lane (index description / memory item / skill step).""" + """The searchable atom of a lane (index description / memory item).""" lane: str # Originating raw source resource (provenance), and its relative path. source_id: str | None = None source_path: str | None = None - # Sub-type within a lane (memory: profile/event/...; skill: step kind; etc.). - entry_kind: str + # Sub-type within a lane (memory: profile/event/knowledge). + entry_type: str text: str embedding: list[float] | None = None happened_at: datetime | None = None @@ -126,7 +103,6 @@ class Entry(BaseRecord): # extra may contain: # - content_hash / reinforcement_count / last_reinforced_at (salience) # - ref_id (reference tracking) - # - when_to_use / metadata / tool_calls (tool memory) class ResourceEntry(BaseRecord): @@ -168,11 +144,10 @@ def build_scoped_models( "SOURCE_LANE", "BaseRecord", "Entry", + "EntryType", "Lane", - "MemoryType", "Resource", "ResourceEntry", - "ToolCallResult", "build_scoped_models", "compute_content_hash", "merge_scope_model", diff --git a/src/memu/database/postgres/models.py b/src/memu/database/postgres/models.py index eb62f48a..ae1d242b 100644 --- a/src/memu/database/postgres/models.py +++ b/src/memu/database/postgres/models.py @@ -65,7 +65,7 @@ class PostgresEntryModel(BaseModelMixin, Entry): lane: str = Field(sa_column=Column(String, nullable=False, index=True)) source_id: str | None = Field(default=None, sa_column=Column(String, nullable=True, index=True)) source_path: str | None = Field(default=None, sa_column=Column(String, nullable=True)) - entry_kind: str = Field(sa_column=Column(String, nullable=False)) + entry_type: str = Field(sa_column=Column(String, nullable=False)) text: str = Field(sa_column=Column(Text, nullable=False)) embedding: list[float] | None = Field(default=None, sa_column=Column(Vector(), nullable=True)) happened_at: datetime | None = Field(default=None, sa_column=Column(TZDateTime, nullable=True)) diff --git a/src/memu/database/postgres/repositories/entry_repo.py b/src/memu/database/postgres/repositories/entry_repo.py index aba3b362..f25d329e 100644 --- a/src/memu/database/postgres/repositories/entry_repo.py +++ b/src/memu/database/postgres/repositories/entry_repo.py @@ -41,9 +41,7 @@ def get_entry(self, entry_id: str) -> Entry | None: return self._cache_entry(row) return None - def list_entries( - self, where: Mapping[str, Any] | None = None, *, lane: str | None = None - ) -> dict[str, Entry]: + def list_entries(self, where: Mapping[str, Any] | None = None, *, lane: str | None = None) -> dict[str, Entry]: from sqlmodel import select model = self._sqla_models.Entry @@ -59,9 +57,7 @@ def list_entries( result[entry.id] = entry return result - def list_entries_by_ref_ids( - self, ref_ids: list[str], where: Mapping[str, Any] | None = None - ) -> dict[str, Entry]: + def list_entries_by_ref_ids(self, ref_ids: list[str], where: Mapping[str, Any] | None = None) -> dict[str, Entry]: if not ref_ids: return {} @@ -82,9 +78,7 @@ def list_entries_by_ref_ids( result[entry.id] = entry return result - def clear_entries( - self, where: Mapping[str, Any] | None = None, *, lane: str | None = None - ) -> dict[str, Entry]: + def clear_entries(self, where: Mapping[str, Any] | None = None, *, lane: str | None = None) -> dict[str, Entry]: from sqlmodel import delete, select model = self._sqla_models.Entry @@ -114,39 +108,32 @@ def create_entry( *, lane: str, source_id: str | None, - entry_kind: str, + entry_type: str, text: str, embedding: list[float], user_data: dict[str, Any], source_path: str | None = None, reinforce: bool = False, - tool_record: dict[str, Any] | None = None, ) -> Entry: - if reinforce and entry_kind != "tool": + if reinforce: return self.create_entry_reinforce( lane=lane, source_id=source_id, - entry_kind=entry_kind, + entry_type=entry_type, text=text, embedding=embedding, user_data=user_data, source_path=source_path, ) - extra: dict[str, Any] = {} - if tool_record: - for key in ("when_to_use", "metadata", "tool_calls"): - if tool_record.get(key) is not None: - extra[key] = tool_record[key] - entry = self._entry_model( lane=lane, source_id=source_id, source_path=source_path, - entry_kind=entry_kind, + entry_type=entry_type, text=text, embedding=self._prepare_embedding(embedding), - extra=extra if extra else {}, + extra={}, **user_data, created_at=self._now(), updated_at=self._now(), @@ -165,7 +152,7 @@ def create_entry_reinforce( *, lane: str, source_id: str | None, - entry_kind: str, + entry_type: str, text: str, embedding: list[float], user_data: dict[str, Any], @@ -174,7 +161,7 @@ def create_entry_reinforce( from sqlmodel import select model = self._sqla_models.Entry - content_hash = compute_content_hash(text, entry_kind) + content_hash = compute_content_hash(text, entry_type) entry_extra = user_data.pop("extra", {}) if "extra" in user_data else {} with self._sessions.session() as session: @@ -209,7 +196,7 @@ def create_entry_reinforce( lane=lane, source_id=source_id, source_path=source_path, - entry_kind=entry_kind, + entry_type=entry_type, text=text, embedding=self._prepare_embedding(embedding), **user_data, @@ -229,11 +216,10 @@ def update_entry( self, *, entry_id: str, - entry_kind: str | None = None, + entry_type: str | None = None, text: str | None = None, embedding: list[float] | None = None, extra: dict[str, Any] | None = None, - tool_record: dict[str, Any] | None = None, ) -> Entry: from sqlmodel import select @@ -245,22 +231,15 @@ def update_entry( msg = f"Entry with id {entry_id} not found" raise KeyError(msg) - if entry_kind is not None: - entry.entry_kind = entry_kind + if entry_type is not None: + entry.entry_type = entry_type if text is not None: entry.text = text if embedding is not None: entry.embedding = self._prepare_embedding(embedding) - current_extra = entry.extra or {} if extra is not None: - current_extra = {**current_extra, **extra} - if tool_record is not None: - for key in ("when_to_use", "metadata", "tool_calls"): - if tool_record.get(key) is not None: - current_extra[key] = tool_record[key] - if extra is not None or tool_record is not None: - entry.extra = current_extra + entry.extra = {**(entry.extra or {}), **extra} entry.updated_at = now session.add(entry) @@ -301,12 +280,7 @@ def vector_search_entries( filters.extend(self._build_filters(model, where)) if lane is not None: filters.append(model.lane == lane) - stmt = ( - select(model.id, (1 - distance).label("score")) - .where(*filters) - .order_by(distance) - .limit(top_k) - ) + stmt = select(model.id, (1 - distance).label("score")).where(*filters).order_by(distance).limit(top_k) with self._sessions.session() as session: rows = session.execute(stmt).all() return [(rid, float(score)) for rid, score in rows] diff --git a/src/memu/database/postgres/repositories/resource_entry_repo.py b/src/memu/database/postgres/repositories/resource_entry_repo.py index 388bf78e..1e9d6a22 100644 --- a/src/memu/database/postgres/repositories/resource_entry_repo.py +++ b/src/memu/database/postgres/repositories/resource_entry_repo.py @@ -76,9 +76,7 @@ def unlink_entry_resource(self, entry_id: str, resource_id: str) -> None: ) ) session.commit() - self.relations[:] = [ - r for r in self.relations if not (r.entry_id == entry_id and r.resource_id == resource_id) - ] + self.relations[:] = [r for r in self.relations if not (r.entry_id == entry_id and r.resource_id == resource_id)] def unlink_entry(self, entry_id: str) -> list[ResourceEntry]: from sqlmodel import delete, select @@ -90,9 +88,7 @@ def unlink_entry(self, entry_id: str) -> list[ResourceEntry]: removed = [self._row_to_record(row) for row in rows] if removed: session.exec( - delete(self._sqla_models.ResourceEntry).where( - self._sqla_models.ResourceEntry.entry_id == entry_id - ) + delete(self._sqla_models.ResourceEntry).where(self._sqla_models.ResourceEntry.entry_id == entry_id) ) session.commit() self.relations[:] = [r for r in self.relations if r.entry_id != entry_id] diff --git a/src/memu/database/repositories/entry.py b/src/memu/database/repositories/entry.py index 2a4b7a1d..f632b1da 100644 --- a/src/memu/database/repositories/entry.py +++ b/src/memu/database/repositories/entry.py @@ -14,37 +14,31 @@ class EntryRepo(Protocol): def get_entry(self, entry_id: str) -> Entry | None: ... - def list_entries( - self, where: Mapping[str, Any] | None = None, *, lane: str | None = None - ) -> dict[str, Entry]: ... + def list_entries(self, where: Mapping[str, Any] | None = None, *, lane: str | None = None) -> dict[str, Entry]: ... - def clear_entries( - self, where: Mapping[str, Any] | None = None, *, lane: str | None = None - ) -> dict[str, Entry]: ... + def clear_entries(self, where: Mapping[str, Any] | None = None, *, lane: str | None = None) -> dict[str, Entry]: ... def create_entry( self, *, lane: str, source_id: str | None, - entry_kind: str, + entry_type: str, text: str, embedding: list[float], user_data: dict[str, Any], source_path: str | None = None, reinforce: bool = False, - tool_record: dict[str, Any] | None = None, ) -> Entry: ... def update_entry( self, *, entry_id: str, - entry_kind: str | None = None, + entry_type: str | None = None, text: str | None = None, embedding: list[float] | None = None, extra: dict[str, Any] | None = None, - tool_record: dict[str, Any] | None = None, ) -> Entry: ... def delete_entry(self, entry_id: str) -> None: ... diff --git a/src/memu/database/sqlite/models.py b/src/memu/database/sqlite/models.py index f9642bbe..3f58c837 100644 --- a/src/memu/database/sqlite/models.py +++ b/src/memu/database/sqlite/models.py @@ -45,7 +45,7 @@ class SQLiteResourceModel(SQLiteBaseModelMixin, Resource): """SQLite resource model. A single physical table holds both raw inputs (``lane="source"``) and the - generated lane docs (``lane`` in {index, memory, skill}). + generated lane docs (``lane`` in {index, memory}). """ lane: str = Field(default="source", sa_column=Column(String, nullable=False, index=True)) @@ -69,7 +69,7 @@ class SQLiteEntryModel(SQLiteBaseModelMixin, Entry): lane: str = Field(sa_column=Column(String, nullable=False, index=True)) source_id: str | None = Field(default=None, sa_column=Column(String, nullable=True)) source_path: str | None = Field(default=None, sa_column=Column(String, nullable=True)) - entry_kind: str = Field(sa_column=Column(String, nullable=False)) + entry_type: str = Field(sa_column=Column(String, nullable=False)) text: str = Field(sa_column=Column(Text, nullable=False)) # Override inherited embedding field: SQLite has no native vector type, so store the # vector in a JSON column (a bare ``list`` annotation is not mappable by SQLModel). diff --git a/src/memu/database/sqlite/repositories/entry_repo.py b/src/memu/database/sqlite/repositories/entry_repo.py index 993e3232..55db36c5 100644 --- a/src/memu/database/sqlite/repositories/entry_repo.py +++ b/src/memu/database/sqlite/repositories/entry_repo.py @@ -58,7 +58,7 @@ def _row_to_entry(self, row: Any) -> Entry: lane=row.lane, source_id=row.source_id, source_path=row.source_path, - entry_kind=row.entry_kind, + entry_type=row.entry_type, text=row.text, embedding=self._normalize_embedding(row.embedding), happened_at=row.happened_at, @@ -84,9 +84,7 @@ def get_entry(self, entry_id: str) -> Entry | None: self.entries[row.id] = entry return entry - def list_entries( - self, where: Mapping[str, Any] | None = None, *, lane: str | None = None - ) -> dict[str, Entry]: + def list_entries(self, where: Mapping[str, Any] | None = None, *, lane: str | None = None) -> dict[str, Entry]: """List entries matching the where clause and optional lane filter.""" filters = self._build_filters(self._entry_model, where) filters.extend(self._lane_filters(self._entry_model, lane)) @@ -104,9 +102,7 @@ def list_entries( return result - def list_entries_by_ref_ids( - self, ref_ids: list[str], where: Mapping[str, Any] | None = None - ) -> dict[str, Entry]: + def list_entries_by_ref_ids(self, ref_ids: list[str], where: Mapping[str, Any] | None = None) -> dict[str, Entry]: """List entries whose ``extra.ref_id`` is in ``ref_ids``.""" if not ref_ids: return {} @@ -129,9 +125,7 @@ def list_entries_by_ref_ids( return result - def clear_entries( - self, where: Mapping[str, Any] | None = None, *, lane: str | None = None - ) -> dict[str, Entry]: + def clear_entries(self, where: Mapping[str, Any] | None = None, *, lane: str | None = None) -> dict[str, Entry]: """Clear entries matching the where clause and optional lane filter.""" filters = self._build_filters(self._entry_model, where) filters.extend(self._lane_filters(self._entry_model, lane)) @@ -162,41 +156,34 @@ def create_entry( *, lane: str, source_id: str | None, - entry_kind: str, + entry_type: str, text: str, embedding: list[float], user_data: dict[str, Any], source_path: str | None = None, reinforce: bool = False, - tool_record: dict[str, Any] | None = None, ) -> Entry: """Create a new entry (optionally reinforcing a duplicate).""" - if reinforce and entry_kind != "tool": + if reinforce: return self.create_entry_reinforce( lane=lane, source_id=source_id, - entry_kind=entry_kind, + entry_type=entry_type, text=text, embedding=embedding, user_data=user_data, source_path=source_path, ) - extra: dict[str, Any] = {} - if tool_record: - for key in ("when_to_use", "metadata", "tool_calls"): - if tool_record.get(key) is not None: - extra[key] = tool_record[key] - now = self._now() row = self._entry_model( lane=lane, source_id=source_id, source_path=source_path, - entry_kind=entry_kind, + entry_type=entry_type, text=text, embedding=self._prepare_embedding(embedding), - extra=extra if extra else {}, + extra={}, created_at=now, updated_at=now, **user_data, @@ -215,7 +202,7 @@ def create_entry_reinforce( *, lane: str, source_id: str | None, - entry_kind: str, + entry_type: str, text: str, embedding: list[float], user_data: dict[str, Any], @@ -226,7 +213,7 @@ def create_entry_reinforce( If an entry with the same content hash exists in the same scope, reinforce it instead of creating a duplicate. """ - content_hash = compute_content_hash(text, entry_kind) + content_hash = compute_content_hash(text, entry_type) with self._sessions.session() as session: content_hash_col = func.json_extract(self._entry_model.extra, "$.content_hash") @@ -264,7 +251,7 @@ def create_entry_reinforce( lane=lane, source_id=source_id, source_path=source_path, - entry_kind=entry_kind, + entry_type=entry_type, text=text, embedding=self._prepare_embedding(embedding), extra=entry_extra, @@ -284,11 +271,10 @@ def update_entry( self, *, entry_id: str, - entry_kind: str | None = None, + entry_type: str | None = None, text: str | None = None, embedding: list[float] | None = None, extra: dict[str, Any] | None = None, - tool_record: dict[str, Any] | None = None, ) -> Entry: """Update an existing entry. @@ -303,22 +289,15 @@ def update_entry( msg = f"Entry with id {entry_id} not found" raise KeyError(msg) - if entry_kind is not None: - row.entry_kind = entry_kind + if entry_type is not None: + row.entry_type = entry_type if text is not None: row.text = text if embedding is not None: row.embedding = self._prepare_embedding(embedding) - current_extra = row.extra or {} if extra is not None: - current_extra = {**current_extra, **extra} - if tool_record is not None: - for key in ("when_to_use", "metadata", "tool_calls"): - if tool_record.get(key) is not None: - current_extra[key] = tool_record[key] - if extra is not None or tool_record is not None: - row.extra = current_extra + row.extra = {**(row.extra or {}), **extra} row.updated_at = self._now() diff --git a/src/memu/database/sqlite/repositories/resource_repo.py b/src/memu/database/sqlite/repositories/resource_repo.py index 2041d605..4a988da4 100644 --- a/src/memu/database/sqlite/repositories/resource_repo.py +++ b/src/memu/database/sqlite/repositories/resource_repo.py @@ -23,7 +23,7 @@ class SQLiteResourceRepo(SQLiteRepoBase, ResourceRepo): """SQLite implementation of the resource repository. A single physical table holds both raw inputs (``lane="source"``) and the - generated lane docs (``lane`` in {index, memory, skill}). + generated lane docs (``lane`` in {index, memory}). """ def __init__( diff --git a/src/memu/memory_fs/__init__.py b/src/memu/memory_fs/__init__.py index 62975d13..9e1af3f6 100644 --- a/src/memu/memory_fs/__init__.py +++ b/src/memu/memory_fs/__init__.py @@ -11,7 +11,7 @@ MemoryFileExporter, slugify, ) -from memu.memory_fs.synthesizer import MemorySynthesizer, SynthesisResult +from memu.memory_fs.synthesizer import MemorySynthesizer __all__ = [ "ExistingArtifacts", @@ -19,6 +19,5 @@ "FileDescription", "MemoryFileExporter", "MemorySynthesizer", - "SynthesisResult", "slugify", ] diff --git a/src/memu/memory_fs/exporter.py b/src/memu/memory_fs/exporter.py index 7f829f5a..43f7a8cf 100644 --- a/src/memu/memory_fs/exporter.py +++ b/src/memu/memory_fs/exporter.py @@ -6,17 +6,14 @@ / ├── INDEX.md ← index of the raw files under resource/ ├── MEMORY.md ← overall overview + index of memory/ - ├── SKILL.md ← index of the skills under skill/ ├── resource/ │ └── ← one copied raw source file - ├── memory/ - │ └── .md ← one memory category (description + summary) - └── skill/ - └── /SKILL.md ← one synthesized skill per folder + └── memory/ + └── .md ← one memory category (description + summary) -Each artifact is a different aggregation of what the agent has stored. The three -root indexes (``INDEX.md`` / ``MEMORY.md`` / ``SKILL.md``) each point at a sibling -directory of payloads (``resource/`` / ``memory/`` / ``skill/``): +Each artifact is a different aggregation of what the agent has stored. The two +root indexes (``INDEX.md`` / ``MEMORY.md``) each point at a sibling directory of +payloads (``resource/`` / ``memory/``): - ``resource/`` : the raw source files copied verbatim out of the blob store (``Resource.local_path``), so the actual ingested bytes live next to the memory. @@ -25,11 +22,6 @@ - ``memory/`` : the living memory split one file per memory-lane :class:`~memu.database.models.Resource` (its description + summary). - ``MEMORY.md`` : an overall overview that links to each ``memory/.md`` file. -- ``skill/**`` : reusable skills. When the caller supplies a synthesized - ``slug -> body`` map (``synthesize=True``), the tree is rendered from it; when no - map is supplied (``synthesize=False``), it falls back to a deterministic, - LLM-free bypass that breaks out the extracted skill-type memory items. -- ``SKILL.md`` (root): a generated index/table of contents over the ``skill/`` tree. This layer is read-only against the database and never mutates memory records. A sidecar manifest (``.memufs_manifest.json``) records the content hash of every @@ -54,14 +46,10 @@ from memu.database.models import Entry, Resource MANIFEST_NAME = ".memufs_manifest.json" -SKILL_DIRNAME = "skill" RESOURCE_DIRNAME = "resource" MEMORY_DIRNAME = "memory" -SKILL_FILENAME = "SKILL.md" -SKILL_INDEX_FILENAME = "SKILL.md" INDEX_FILENAME = "INDEX.md" MEMORY_FILENAME = "MEMORY.md" -SKILL_MEMORY_TYPE = "skill" _GENERATED_NOTICE = "" @@ -96,7 +84,6 @@ class ExistingArtifacts: """ memory_body: str = "" - skills: dict[str, str] = field(default_factory=dict) @dataclass @@ -150,15 +137,9 @@ def export( *, where: Mapping[str, Any] | None = None, memory_body: str | None = None, - skills: dict[str, str] | None = None, ) -> ExportResult: """Render the (optionally scoped) store and write only changed artifacts. - ``skills`` is the synthesized skill map (slug -> body) produced by - :class:`~memu.memory_fs.synthesizer.MemorySynthesizer`; the ``skill/`` tree - and its root ``SKILL.md`` index are sourced from it. When ``skills`` is - ``None`` the exporter falls back to a deterministic, LLM-free bypass that - derives the tree from the extracted skill-type memory items. ``memory_body`` is an optional synthesized override for the ``MEMORY.md`` body; when omitted, ``MEMORY.md`` is rendered deterministically as an overview that links to the per-category ``memory/.md`` files. The @@ -174,10 +155,6 @@ def export( # The shared trunk: one multimodal description per source file. descriptions = self._build_descriptions(resources) - # The skill bypass reads extracted skill-type entries only when no synthesized - # map is supplied; listing is cheap and avoided otherwise. - items = list(database.entry_repo.list_entries(where=scope).values()) if skills is None else [] - # resource/: copy the raw source bytes verbatim; ``links`` maps each # resource id to its relative path under resource/ for the INDEX.md links. raw_artifacts, links = self._resource_artifacts(descriptions) @@ -189,24 +166,14 @@ def export( for category, slug in zip(ordered_categories, category_slugs, strict=True) } - # The skill/ tree is sourced from the synthesized ``skills`` map when the - # caller supplies one (see MemorySynthesizer); otherwise it falls back to a - # deterministic, LLM-free bypass over the extracted skill-type memory items. - skill_map = skills if skills is not None else self._skill_bypass(items) - skill_artifacts = { - f"{SKILL_DIRNAME}/{slug}/{SKILL_FILENAME}": self._skill_document(body) for slug, body in skill_map.items() - } - artifacts: dict[str, str | bytes] = {} artifacts.update(raw_artifacts) artifacts.update(category_artifacts) - artifacts.update(skill_artifacts) if memory_body is not None: artifacts[MEMORY_FILENAME] = self._memory_document(memory_body) else: artifacts[MEMORY_FILENAME] = self._memory_index(ordered_categories, category_slugs) artifacts[INDEX_FILENAME] = self._index_bypass(descriptions, links) - artifacts[SKILL_INDEX_FILENAME] = self._skill_index(skill_map) return self._sync(artifacts) @@ -245,10 +212,10 @@ def build_synthesis_descriptions(resources: list[Resource], items: list[Entry]) for resource in sorted(resources, key=lambda r: (r.url, r.id)): res_items = sorted( items_by_resource.get(resource.id, []), - key=lambda i: (i.entry_kind, i.created_at, i.id), + key=lambda i: (i.entry_type, i.created_at, i.id), ) parts = [ - f"[{item.entry_kind}] {' '.join((item.text or '').split())}" + f"[{item.entry_type}] {' '.join((item.text or '').split())}" for item in res_items if (item.text or "").strip() ] @@ -264,84 +231,6 @@ def build_synthesis_descriptions(resources: list[Resource], items: list[Entry]) ) return descriptions - # -- bypass: SKILL ----------------------------------------------------- - - @staticmethod - def _skill_document(body: str) -> str: - return f"{_GENERATED_NOTICE}\n\n{body.strip()}\n" - - def _skill_bypass(self, items: list[Entry]) -> dict[str, str]: - """Deterministically break out skill-kind entries as a slug -> body map. - - The LLM-free fallback used when no synthesized skill map is supplied: each - skill-kind entry's text becomes one ``skill//SKILL.md`` document, - slugged from its frontmatter ``name:``/first heading. Mirrors the shape of - the synthesized map so the rest of :meth:`export` is path-agnostic. - """ - skills = sorted( - (item for item in items if item.entry_kind == SKILL_MEMORY_TYPE and (item.text or "").strip()), - key=lambda i: (i.created_at, i.id), - ) - skill_map: dict[str, str] = {} - used: dict[str, int] = {} - for item in skills: - body = (item.text or "").strip() - base = self._skill_name(body, fallback=f"skill-{item.id[:6]}") - count = used.get(base, 0) - slug = base if count == 0 else f"{base}-{count + 1}" - used[base] = count + 1 - skill_map[slug] = body - return skill_map - - @staticmethod - def _skill_name(body: str, *, fallback: str) -> str: - """Derive a skill folder slug from frontmatter ``name:`` or first heading.""" - if body.startswith("---"): - end = body.find("\n---", 3) - front = body[3:end] if end != -1 else "" - match = re.search(r"^\s*name:\s*(.+)$", front, re.MULTILINE) - if match: - return slugify(match.group(1)) - heading = re.search(r"^#+\s*(.+)$", body, re.MULTILINE) - if heading: - return slugify(heading.group(1)) - return slugify(fallback) - - def _skill_index(self, skills: dict[str, str]) -> str: - """A table of contents over the ``skill/`` tree (slug, description, link).""" - lines = ["# Skills", "", _GENERATED_NOTICE, "", "## Skills", ""] - if skills: - for slug, body in sorted(skills.items()): - description = self._skill_description(body) or "_No description._" - link = f"{SKILL_DIRNAME}/{slug}/{SKILL_FILENAME}" - lines.append(f"- [`{slug}`]({link}) — {description}") - else: - lines.append("_No skills yet._") - lines.append("") - return "\n".join(lines) - - @classmethod - def _skill_description(cls, body: str) -> str: - """A one-line description for the index: frontmatter ``description``/``name``. - - Falls back to the first non-empty, non-frontmatter line so a synthesized - skill without explicit frontmatter still surfaces something readable. - """ - lines = body.strip().splitlines() - if lines and lines[0].strip() == "---": - for line in lines[1:]: - if line.strip() == "---": - break - key, sep, value = line.partition(":") - if sep and key.strip().lower() == "description" and value.strip(): - return cls._inline(value.strip()) - for raw_line in lines: - line = raw_line.strip() - if not line or line == "---": - continue - return cls._inline(line.lstrip("#").strip()) - return "" - # -- bypass: MEMORY ---------------------------------------------------- @staticmethod @@ -494,7 +383,7 @@ def _sync(self, artifacts: dict[str, str | bytes]) -> ExportResult: return result def _prune_empty_dirs(self, directory: pathlib.Path) -> None: - """Remove now-empty directories created for nested artifacts (e.g. skill/).""" + """Remove now-empty directories created for nested artifacts (e.g. memory/).""" root = self.output_dir.resolve() current = directory while current.resolve() != root and root in current.resolve().parents: @@ -515,8 +404,8 @@ def artifacts_exist(self) -> bool: return (self.output_dir / MEMORY_FILENAME).exists() def read_existing(self) -> ExistingArtifacts: - """Load the prior MEMORY/SKILL artifacts as a single bundle for merging.""" - return ExistingArtifacts(memory_body=self.read_memory_body(), skills=self.read_skills()) + """Load the prior MEMORY artifact as a single bundle for merging.""" + return ExistingArtifacts(memory_body=self.read_memory_body()) def read_memory_body(self) -> str: """Read MEMORY.md and strip the heading/notice, returning just the body.""" @@ -525,18 +414,6 @@ def read_memory_body(self) -> str: return "" return self._strip_chrome(path.read_text(encoding="utf-8"), drop_heading="# Memory") - def read_skills(self) -> dict[str, str]: - """Read existing ``skill//SKILL.md`` bodies keyed by slug.""" - skills: dict[str, str] = {} - skill_root = self.output_dir / SKILL_DIRNAME - if not skill_root.is_dir(): - return skills - for child in sorted(skill_root.iterdir()): - doc = child / SKILL_FILENAME - if child.is_dir() and doc.exists(): - skills[child.name] = self._strip_chrome(doc.read_text(encoding="utf-8")) - return skills - @staticmethod def _strip_chrome(text: str, *, drop_heading: str | None = None) -> str: lines = text.splitlines() diff --git a/src/memu/memory_fs/synthesizer.py b/src/memu/memory_fs/synthesizer.py index 16ad5b85..dcf502a8 100644 --- a/src/memu/memory_fs/synthesizer.py +++ b/src/memu/memory_fs/synthesizer.py @@ -1,26 +1,21 @@ -"""LLM synthesis of MEMORY/SKILL artifacts from the shared description trunk. +"""LLM synthesis of the MEMORY artifact from the shared description trunk. This is the optional, opt-in counterpart to the deterministic exporter: instead of -rendering already-extracted database items/summaries, it feeds the per-source -multimodal descriptions to an LLM and synthesizes the memory document and skill -docs directly. ``INDEX.md`` stays deterministic and is handled by the exporter. +rendering already-extracted database summaries, it feeds the per-source multimodal +descriptions to an LLM and synthesizes the memory document directly. ``INDEX.md`` +stays deterministic and is handled by the exporter. """ from __future__ import annotations -import asyncio -import json import re from collections.abc import Awaitable, Callable -from dataclasses import dataclass, field from typing import TYPE_CHECKING -from memu.memory_fs.exporter import slugify from memu.prompts.memory_fs import ( DESCRIPTIONS_PLACEHOLDER, EXISTING_PLACEHOLDER, MEMORY_SYNTHESIS_PROMPT, - SKILL_SYNTHESIS_PROMPT, ) if TYPE_CHECKING: @@ -29,91 +24,32 @@ ChatFn = Callable[[str], Awaitable[str]] -@dataclass -class SynthesisResult: - """Synthesized artifact payloads, ready to hand to the exporter.""" - - memory_body: str = "" - skills: dict[str, str] = field(default_factory=dict) - - class MemorySynthesizer: - """Synthesize MEMORY/SKILL content from multimodal descriptions via an LLM.""" + """Synthesize MEMORY content from multimodal descriptions via an LLM.""" - def __init__( - self, - *, - memory_prompt: str = MEMORY_SYNTHESIS_PROMPT, - skill_prompt: str = SKILL_SYNTHESIS_PROMPT, - ) -> None: + def __init__(self, *, memory_prompt: str = MEMORY_SYNTHESIS_PROMPT) -> None: self._memory_prompt = memory_prompt - self._skill_prompt = skill_prompt async def synthesize( self, descriptions: list[FileDescription], *, existing_memory: str = "", - existing_skills: dict[str, str] | None = None, - chat: ChatFn, - ) -> SynthesisResult: - """Synthesize MEMORY + SKILL from the descriptions, merging into any existing - artifacts. - - Pass empty ``existing_*`` (the default) to build from scratch; pass the prior - artifacts to incrementally fold the (changed) descriptions into them. The two - LLM calls are independent and run concurrently. - """ - existing_skills = existing_skills or {} - formatted = self._format(descriptions) - if not formatted: - return SynthesisResult(memory_body=existing_memory, skills=dict(existing_skills)) - - memory_body, skills = await asyncio.gather( - self._synthesize_memory_formatted(formatted, existing_memory=existing_memory, chat=chat), - self._synthesize_skills_formatted(formatted, existing_skills=existing_skills, chat=chat), - ) - return SynthesisResult(memory_body=memory_body, skills=skills) - - async def synthesize_skills( - self, - descriptions: list[FileDescription], - *, - existing_skills: dict[str, str] | None = None, chat: ChatFn, - ) -> dict[str, str]: - """Synthesize only the skill bypass (decoupled from MEMORY.md). + ) -> str: + """Synthesize MEMORY from the descriptions, merging into any existing body. - The ``skill/`` tree is a sibling of ``MEMORY.md`` projected from the same - description trunk, so it can be (re)built independently of how MEMORY.md is - produced. As with :meth:`synthesize`, empty ``existing_skills`` builds from - scratch and a populated map merges the changed descriptions into it. + Pass empty ``existing_memory`` (the default) to build from scratch; pass the + prior body to incrementally fold the (changed) descriptions into it. """ - existing_skills = existing_skills or {} formatted = self._format(descriptions) if not formatted: - return dict(existing_skills) - return await self._synthesize_skills_formatted(formatted, existing_skills=existing_skills, chat=chat) - - async def _synthesize_memory_formatted(self, formatted: str, *, existing_memory: str, chat: ChatFn) -> str: + return existing_memory prompt = self._memory_prompt.replace(EXISTING_PLACEHOLDER, existing_memory.strip() or "(empty)").replace( DESCRIPTIONS_PLACEHOLDER, formatted ) return self._clean_markdown(await chat(prompt)) - async def _synthesize_skills_formatted( - self, formatted: str, *, existing_skills: dict[str, str], chat: ChatFn - ) -> dict[str, str]: - prompt = self._skill_prompt.replace( - EXISTING_PLACEHOLDER, self._format_existing_skills(existing_skills) or "(none)" - ).replace(DESCRIPTIONS_PLACEHOLDER, formatted) - upserts = self._parse_skills(await chat(prompt)) - return {**existing_skills, **upserts} - - @staticmethod - def _format_existing_skills(skills: dict[str, str]) -> str: - return "\n\n".join(f"## {slug}\n{body}".strip() for slug, body in sorted(skills.items())) - @staticmethod def _format(descriptions: list[FileDescription]) -> str: lines = [ @@ -129,45 +65,5 @@ def _clean_markdown(raw: str) -> str: text = re.sub(r"\n```$", "", text).strip() return text - def _parse_skills(self, raw: str) -> dict[str, str]: - payload = self._extract_json_array(raw) - if payload is None: - return {} - try: - parsed = json.loads(payload) - except (json.JSONDecodeError, TypeError): - return {} - if not isinstance(parsed, list): - return {} - - skills: dict[str, str] = {} - used: dict[str, int] = {} - for entry in parsed: - if not isinstance(entry, dict): - continue - name = entry.get("name") - body = entry.get("body") - if not isinstance(name, str) or not isinstance(body, str): - continue - body = body.strip() - if not body: - continue - base = slugify(name) - count = used.get(base, 0) - slug = base if count == 0 else f"{base}-{count + 1}" - used[base] = count + 1 - skills[slug] = body - return skills - - @staticmethod - def _extract_json_array(raw: str) -> str | None: - if not raw: - return None - start = raw.find("[") - end = raw.rfind("]") - if start == -1 or end == -1 or end <= start: - return None - return raw[start : end + 1] - -__all__ = ["ChatFn", "MemorySynthesizer", "SynthesisResult"] +__all__ = ["ChatFn", "MemorySynthesizer"] diff --git a/src/memu/prompts/__init__.py b/src/memu/prompts/__init__.py index fbfe2dc6..6a0587af 100644 --- a/src/memu/prompts/__init__.py +++ b/src/memu/prompts/__init__.py @@ -1,13 +1,13 @@ from memu.prompts.category_summary import PROMPT as CATEGORY_SUMMARY_PROMPT -from memu.prompts.memory_type import DEFAULT_MEMORY_TYPES -from memu.prompts.memory_type import PROMPTS as MEMORY_TYPE_PROMPTS +from memu.prompts.entry_type import DEFAULT_ENTRY_TYPES +from memu.prompts.entry_type import PROMPTS as ENTRY_TYPE_PROMPTS from memu.prompts.preprocess import PROMPTS as PREPROCESS_PROMPTS from memu.prompts.retrieve.judger import PROMPT as RETRIEVE_JUDGER_PROMPT __all__ = [ "CATEGORY_SUMMARY_PROMPT", - "DEFAULT_MEMORY_TYPES", - "MEMORY_TYPE_PROMPTS", + "DEFAULT_ENTRY_TYPES", + "ENTRY_TYPE_PROMPTS", "PREPROCESS_PROMPTS", "RETRIEVE_JUDGER_PROMPT", ] diff --git a/src/memu/prompts/memory_type/__init__.py b/src/memu/prompts/entry_type/__init__.py similarity index 54% rename from src/memu/prompts/memory_type/__init__.py rename to src/memu/prompts/entry_type/__init__.py index ea584bd4..3870af22 100644 --- a/src/memu/prompts/memory_type/__init__.py +++ b/src/memu/prompts/entry_type/__init__.py @@ -1,24 +1,32 @@ -from memu.prompts.memory_type import behavior, event, knowledge, profile, skill, tool +from memu.prompts.entry_type import description, event, knowledge, log, profile, tool -# DEFAULT_MEMORY_TYPES: list[str] = ["profile", "event", "knowledge", "behavior"] -DEFAULT_MEMORY_TYPES: list[str] = ["profile", "event"] +# Per-lane default entry types. Each lane runs the same extraction code path; only +# the entry-type set, prompts, and grouping cardinality differ (see ADR 0006). +LANE_ENTRY_TYPES: dict[str, list[str]] = { + "index": ["description"], + "memory": ["profile", "event"], + "skill": ["tool", "log"], +} + +# Backward-friendly alias: the historical flat default refers to the memory lane. +DEFAULT_ENTRY_TYPES: list[str] = list(LANE_ENTRY_TYPES["memory"]) PROMPTS: dict[str, str] = { "profile": profile.PROMPT.strip(), "event": event.PROMPT.strip(), "knowledge": knowledge.PROMPT.strip(), - "behavior": behavior.PROMPT.strip(), - "skill": skill.PROMPT.strip(), "tool": tool.PROMPT.strip(), + "log": log.PROMPT.strip(), + "description": description.PROMPT.strip(), } CUSTOM_PROMPTS: dict[str, dict[str, str]] = { "profile": profile.CUSTOM_PROMPT, "event": event.CUSTOM_PROMPT, "knowledge": knowledge.CUSTOM_PROMPT, - "behavior": behavior.CUSTOM_PROMPT, - "skill": skill.CUSTOM_PROMPT, "tool": tool.CUSTOM_PROMPT, + "log": log.CUSTOM_PROMPT, + "description": description.CUSTOM_PROMPT, } CUSTOM_TYPE_CUSTOM_PROMPTS: dict[str, str] = { @@ -40,7 +48,8 @@ __all__ = [ "CUSTOM_PROMPTS", "CUSTOM_TYPE_CUSTOM_PROMPTS", + "DEFAULT_ENTRY_TYPES", "DEFAULT_MEMORY_CUSTOM_PROMPT_ORDINAL", - "DEFAULT_MEMORY_TYPES", + "LANE_ENTRY_TYPES", "PROMPTS", ] diff --git a/src/memu/prompts/entry_type/description.py b/src/memu/prompts/entry_type/description.py new file mode 100644 index 00000000..7386776b --- /dev/null +++ b/src/memu/prompts/entry_type/description.py @@ -0,0 +1,83 @@ +PROMPT_BLOCK_OBJECTIVE = """ +# Task Objective +You are a professional Resource Indexer. Your core task is to write one concise, faithful +description of a single resource so it can be found again by semantic search: what it is, what +it covers, and the salient entities/topics it contains. +""" + +PROMPT_BLOCK_WORKFLOW = """ +# Workflow +Read the full resource to understand its purpose and contents. +## Summarize +Write a single self-contained description that captures the resource's subject, scope, and key +topics or entities. +## Final output +Output exactly one description item for the whole resource. +""" + +PROMPT_BLOCK_RULES = """ +# Rules +## General requirements (must satisfy all) +- Produce exactly ONE description item for the resource (this lane is one-to-one per resource). +- The description must be self-contained and understandable without the original resource. +- Keep it concise (< 80 words) but information-dense; prefer concrete nouns and topics over + filler. +- Describe what the resource IS and CONTAINS; do not extract individual facts or steps. +Important: Describe only what the resource actually contains. No guesses and no invented content. + +## Forbidden content +- Multiple items (only one description is expected). +- Illegal / harmful sensitive topics. +- Opinions or speculation not grounded in the resource. +""" + +PROMPT_BLOCK_OUTPUT = """ +# Output Format (XML) +Return a single description wrapped in one element: + + + One concise description of the whole resource + + +""" + +PROMPT_BLOCK_EXAMPLES = """ +# Examples (Input / Output / Explanation) +Example 1: Resource Description +## Input +A meeting transcript where the team discusses Q3 roadmap priorities, agrees to ship the billing +revamp first, and assigns owners for the analytics dashboard. +## Output + + + Team meeting transcript covering the Q3 roadmap: prioritizes the billing revamp for the next release and assigns owners for the analytics dashboard. + + +## Explanation +A single description captures the resource's subject and key topics for later retrieval. +""" + +PROMPT_BLOCK_INPUT = """ +# Original Resource: + +{resource} + +""" + +PROMPT = "\n\n".join([ + PROMPT_BLOCK_OBJECTIVE.strip(), + PROMPT_BLOCK_WORKFLOW.strip(), + PROMPT_BLOCK_RULES.strip(), + PROMPT_BLOCK_OUTPUT.strip(), + PROMPT_BLOCK_EXAMPLES.strip(), + PROMPT_BLOCK_INPUT.strip(), +]) + +CUSTOM_PROMPT = { + "objective": PROMPT_BLOCK_OBJECTIVE.strip(), + "workflow": PROMPT_BLOCK_WORKFLOW.strip(), + "rules": PROMPT_BLOCK_RULES.strip(), + "output": PROMPT_BLOCK_OUTPUT.strip(), + "examples": PROMPT_BLOCK_EXAMPLES.strip(), + "input": PROMPT_BLOCK_INPUT.strip(), +} diff --git a/src/memu/prompts/memory_type/event.py b/src/memu/prompts/entry_type/event.py similarity index 100% rename from src/memu/prompts/memory_type/event.py rename to src/memu/prompts/entry_type/event.py diff --git a/src/memu/prompts/memory_type/knowledge.py b/src/memu/prompts/entry_type/knowledge.py similarity index 100% rename from src/memu/prompts/memory_type/knowledge.py rename to src/memu/prompts/entry_type/knowledge.py diff --git a/src/memu/prompts/entry_type/log.py b/src/memu/prompts/entry_type/log.py new file mode 100644 index 00000000..893b1c86 --- /dev/null +++ b/src/memu/prompts/entry_type/log.py @@ -0,0 +1,122 @@ +PROMPT_BLOCK_OBJECTIVE = """ +# Task Objective +You are a professional Skill Extractor. Your core task is to extract operation logs: concrete +records of what was actually done, in what order, and with what outcome, so a later agent can +learn from the trace (e.g. the sequence of steps taken to accomplish a task and their results). +""" + +PROMPT_BLOCK_WORKFLOW = """ +# Workflow +Read the full resource to understand the task and the actions taken. +## Extract logs +Select the parts that record concrete actions and outcomes, and extract one log item per +meaningful step or short coherent sequence. +## Review & validate +Merge redundant or duplicated steps. +Resolve contradictions by keeping the most accurate record. +## Final output +Output operation Log items. +""" + +PROMPT_BLOCK_RULES = """ +# Rules +## General requirements (must satisfy all) +- Each log item must be complete and self-contained, written as a past-tense record of an + action and (when present) its result. +- Each log item must capture one step or one short coherent sequence, understandable without + context. +- Similar/redundant logs must be merged into one, and assigned to only one skill. +- Each log item must be < 50 words worth of length (concise but informative). +- Preserve ordering signals (e.g. "first", "then", "after") when they matter. +Important: Extract only actions actually recorded in the resource. No guesses and no invented +steps. + +## Special rules for Log Information +- Generalized reusable capabilities are NOT logs (those belong to the tool type). +- Personal facts, preferences, or unrelated events are NOT logs. + +## Forbidden content +- Trivial chatter with no recorded action or outcome. +- Illegal / harmful sensitive operations. +- Speculative steps not grounded in the resource. + +## Review & validation rules +- Merge duplicated steps: keep only one and assign a single skill. +- Final check: every item must comply with all extraction rules. +""" + +PROMPT_BLOCK_CATEGORY = """ +## Skills: +{categories_str} +""" + +PROMPT_BLOCK_OUTPUT = """ +# Output Format (XML) +Return all logs wrapped in a single element: + + + Operation log item content 1 + + Skill Name + + + + Operation log item content 2 + + Skill Name + + + +""" + +PROMPT_BLOCK_EXAMPLES = """ +# Examples (Input / Output / Explanation) +Example 1: Log Extraction +## Input +assistant: I cloned the repo, ran `make install`, then `make test`. 2 tests failed due to a +missing env var, so I exported DATABASE_URL and re-ran; all tests passed. +## Output + + + Cloned the repo and ran `make install` followed by `make test` + + CI Setup + + + + Fixed 2 failing tests caused by a missing DATABASE_URL by exporting it and re-running; all tests then passed + + CI Setup + + + +## Explanation +Each item records what was actually done and its outcome, grouped under one skill. +""" + +PROMPT_BLOCK_INPUT = """ +# Original Resource: + +{resource} + +""" + +PROMPT = "\n\n".join([ + PROMPT_BLOCK_OBJECTIVE.strip(), + PROMPT_BLOCK_WORKFLOW.strip(), + PROMPT_BLOCK_RULES.strip(), + PROMPT_BLOCK_CATEGORY.strip(), + PROMPT_BLOCK_OUTPUT.strip(), + PROMPT_BLOCK_EXAMPLES.strip(), + PROMPT_BLOCK_INPUT.strip(), +]) + +CUSTOM_PROMPT = { + "objective": PROMPT_BLOCK_OBJECTIVE.strip(), + "workflow": PROMPT_BLOCK_WORKFLOW.strip(), + "rules": PROMPT_BLOCK_RULES.strip(), + "category": PROMPT_BLOCK_CATEGORY.strip(), + "output": PROMPT_BLOCK_OUTPUT.strip(), + "examples": PROMPT_BLOCK_EXAMPLES.strip(), + "input": PROMPT_BLOCK_INPUT.strip(), +} diff --git a/src/memu/prompts/memory_type/profile.py b/src/memu/prompts/entry_type/profile.py similarity index 100% rename from src/memu/prompts/memory_type/profile.py rename to src/memu/prompts/entry_type/profile.py diff --git a/src/memu/prompts/entry_type/tool.py b/src/memu/prompts/entry_type/tool.py new file mode 100644 index 00000000..b338fa7b --- /dev/null +++ b/src/memu/prompts/entry_type/tool.py @@ -0,0 +1,124 @@ +PROMPT_BLOCK_OBJECTIVE = """ +# Task Objective +You are a professional Skill Extractor. Your core task is to extract reusable tools and +capabilities demonstrated in the resource: concrete, repeatable operations that an agent +could invoke again later (e.g. a command, an API call, a function, a procedure with clear +inputs and effects). +""" + +PROMPT_BLOCK_WORKFLOW = """ +# Workflow +Read the full resource to understand what was done and how. +## Extract skills +Select the parts that describe a reusable operation and extract one tool item per distinct +capability. +## Review & validate +Merge semantically similar tools. +Resolve contradictions by keeping the most general, reusable form. +## Final output +Output reusable Tool items. +""" + +PROMPT_BLOCK_RULES = """ +# Rules +## General requirements (must satisfy all) +- Each tool item must be complete and self-contained, written as an imperative capability + statement (what it does, with the inputs/effects that matter). +- Each tool item must describe one single reusable operation and be understandable without + context. +- Similar/redundant tools must be merged into one, and assigned to only one skill. +- Each tool item must be < 50 words worth of length (concise but actionable). +- Prefer the generalized, reusable form over a one-off instance. +Important: Extract only operations actually demonstrated or described in the resource. No +guesses and no invented capabilities. + +## Special rules for Tool Information +- One-off narration, results, or execution traces are NOT tools (those belong to the log type). +- Personal facts, preferences, or events are NOT tools. + +## Forbidden content +- Trivial steps that carry no reusable value. +- Illegal / harmful sensitive operations. +- Speculative capabilities not grounded in the resource. + +## Review & validation rules +- Merge similar tools: keep only one and assign a single skill. +- Final check: every item must comply with all extraction rules. +""" + +PROMPT_BLOCK_CATEGORY = """ +## Skills: +{categories_str} +""" + +PROMPT_BLOCK_OUTPUT = """ +# Output Format (XML) +Return all skills wrapped in a single element: + + + Reusable tool item content 1 + + Skill Name + + + + Reusable tool item content 2 + + Skill Name + + + +""" + +PROMPT_BLOCK_EXAMPLES = """ +# Examples (Input / Output / Explanation) +Example 1: Tool Extraction +## Input +user: How do I find large files on Linux? +assistant: Run `du -ah /path | sort -rh | head -n 20` to list the 20 largest entries under a path. +user: Nice, and to delete one safely I just `rm -i file`. +## Output + + + List the largest files under a path with `du -ah | sort -rh | head -n N` + + Filesystem + + + + Delete a file interactively (with confirmation) using `rm -i ` + + Filesystem + + + +## Explanation +Each item is a reusable command an agent could run again; both are grouped under one skill. +""" + +PROMPT_BLOCK_INPUT = """ +# Original Resource: + +{resource} + +""" + +PROMPT = "\n\n".join([ + PROMPT_BLOCK_OBJECTIVE.strip(), + PROMPT_BLOCK_WORKFLOW.strip(), + PROMPT_BLOCK_RULES.strip(), + PROMPT_BLOCK_CATEGORY.strip(), + PROMPT_BLOCK_OUTPUT.strip(), + PROMPT_BLOCK_EXAMPLES.strip(), + PROMPT_BLOCK_INPUT.strip(), +]) + +CUSTOM_PROMPT = { + "objective": PROMPT_BLOCK_OBJECTIVE.strip(), + "workflow": PROMPT_BLOCK_WORKFLOW.strip(), + "rules": PROMPT_BLOCK_RULES.strip(), + "category": PROMPT_BLOCK_CATEGORY.strip(), + "output": PROMPT_BLOCK_OUTPUT.strip(), + "examples": PROMPT_BLOCK_EXAMPLES.strip(), + "input": PROMPT_BLOCK_INPUT.strip(), +} diff --git a/src/memu/prompts/memory_fs/__init__.py b/src/memu/prompts/memory_fs/__init__.py index 353bcf26..07841954 100644 --- a/src/memu/prompts/memory_fs/__init__.py +++ b/src/memu/prompts/memory_fs/__init__.py @@ -34,31 +34,8 @@ __DESCRIPTIONS__ """ -SKILL_SYNTHESIS_PROMPT = """You are maintaining an AI agent's skill library. - -Below are the EXISTING skills (name + body), followed by NEW source descriptions -that were just added. From the descriptions, identify concrete, repeatable skills or -tool usage patterns (what worked, how to repeat it, what to avoid). Ignore one-off -facts, preferences, or trivia — those belong in the memory document, not here. - -Return ONLY a JSON array of skills to add or replace. Each element is an object: - {"name": "kebab-case-skill-name", "body": "Markdown body for this skill"} -- To revise an existing skill, reuse its exact name and return the full new body. -- To add a new skill, use a new name. -- Only include skills actually affected by the new descriptions; untouched existing - skills are kept automatically. -- If there are no genuine skills to add or change, return an empty array: [] - -EXISTING skills: -__EXISTING__ - -NEW source descriptions: -__DESCRIPTIONS__ -""" - __all__ = [ "DESCRIPTIONS_PLACEHOLDER", "EXISTING_PLACEHOLDER", "MEMORY_SYNTHESIS_PROMPT", - "SKILL_SYNTHESIS_PROMPT", ] diff --git a/src/memu/prompts/memory_type/behavior.py b/src/memu/prompts/memory_type/behavior.py deleted file mode 100644 index 2a842910..00000000 --- a/src/memu/prompts/memory_type/behavior.py +++ /dev/null @@ -1,132 +0,0 @@ -PROMPT_BLOCK_OBJECTIVE = """ -# Task Objective -You are a professional User Memory Extractor. Your core task is to extract behavioral patterns, routines, and solutions that characterize how the user acts or behaves to solve specific problems. -""" - -PROMPT_BLOCK_WORKFLOW = """ -# Workflow -Read the full conversation to understand topics and meanings. -## Extract memories -Select turns that contain valuable Behavior Information and extract behavioral memory items. -## Review & validate -Merge semantically similar items. -Resolve contradictions by keeping the latest / most certain item. -## Final output -Output Behavior Information. -""" - -PROMPT_BLOCK_RULES = """ -# Rules -## General requirements (must satisfy all) -- Use "user" to refer to the user consistently. -- Each memory item must be complete and self-contained, written as a declarative descriptive sentence. -- Each memory item must express one single complete piece of information and be understandable without context. -- Similar/redundant items must be merged into one, and assigned to only one category. -- Each memory item must be < 50 words worth of length (keep it concise but include relevant details). -- Focus on patterns of behavior, routines, and solutions. -- Focus on how the user typically acts, their preferences, and regular activities. -- Can include multi-line records with each line describing a specific step of the pattern, routine, or solution. -Important: Extract only behaviors directly stated or confirmed by the user. No guesses, no suggestions, and no content introduced only by the assistant. -Important: Accurately reflect whether the subject is the user or someone around the user. - -## Special rules for Behavior Information -- One-time actions or specific events are forbidden in Behavior Information unless they demonstrate a significant pattern. -- Focus on recurring patterns, typical approaches, and established routines. -- Do not extract content that was obtained only through the model's follow-up questions unless the user shows strong proactive intent. - -## Forbidden content -- Knowledge Q&A without a clear user behavior pattern. -- One-time events that do not reflect recurring behavior. -- Turns where the user did not respond and only the assistant spoke. -- Illegal / harmful sensitive topics (violence, politics, drugs, etc.). -- Private financial accounts, IDs, addresses, military/defense/government job details, precise street addresses—unless explicitly requested by the user (still avoid if not necessary). -- Any content mentioned only by the assistant and not explicitly confirmed by the user. - -## Review & validation rules -- Merge similar items: keep only one and assign a single category. -- Resolve conflicts: keep the latest / most certain item. -- Final check: every item must comply with all extraction rules. -""" - -PROMPT_BLOCK_CATEGORY = """ -## Memory Categories: -{categories_str} -""" - -PROMPT_BLOCK_OUTPUT = """ -# Output Format (XML) -Return all memories wrapped in a single element: - - - Behavior memory item content 1 - - Category Name - - - - Behavior memory item content 2 - - Category Name - - - -""" - -PROMPT_BLOCK_EXAMPLES = """ -# Examples (Input / Output / Explanation) -Example 1: Behavior Information Extraction -## Input -user: Hi, are you busy? I just got off work and I'm going to the supermarket to buy some groceries. -assistant: Not busy. Are you cooking for yourself? -user: Yes. It's healthier. I work as a product manager in an internet company. I'm 30 this year. After work I like experimenting with cooking, I often figure out dishes by myself. -assistant: Being a PM is tough. You're so disciplined to cook at 30! -user: It's fine. Cooking relaxes me. It's better than takeout. Also I'm traveling next weekend. -assistant: You can check the weather ahead. Your sunscreen can finally be used. -user: I haven't started packing yet. It's annoying. -## Output - - - The user typically cooks for themselves after work instead of ordering takeout - - Daily Routine - - - - The user often experiments with cooking and figures out dishes by themselves - - Daily Routine - - - -## Explanation -Only behavioral patterns explicitly stated by the user are extracted. -Cooking after work and experimenting with dishes are recurring behaviors/routines. -User's job, age are stable traits (not behaviors). The travel plan is a one-time event, not a behavioral pattern. -""" - -PROMPT_BLOCK_INPUT = """ -# Original Resource: - -{resource} - -""" - -PROMPT = "\n\n".join([ - PROMPT_BLOCK_OBJECTIVE.strip(), - PROMPT_BLOCK_WORKFLOW.strip(), - PROMPT_BLOCK_RULES.strip(), - PROMPT_BLOCK_CATEGORY.strip(), - PROMPT_BLOCK_OUTPUT.strip(), - PROMPT_BLOCK_EXAMPLES.strip(), - PROMPT_BLOCK_INPUT.strip(), -]) - -CUSTOM_PROMPT = { - "objective": PROMPT_BLOCK_OBJECTIVE.strip(), - "workflow": PROMPT_BLOCK_WORKFLOW.strip(), - "rules": PROMPT_BLOCK_RULES.strip(), - "category": PROMPT_BLOCK_CATEGORY.strip(), - "output": PROMPT_BLOCK_OUTPUT.strip(), - "examples": PROMPT_BLOCK_EXAMPLES.strip(), - "input": PROMPT_BLOCK_INPUT.strip(), -} diff --git a/src/memu/prompts/memory_type/skill.py b/src/memu/prompts/memory_type/skill.py deleted file mode 100644 index 720a9b28..00000000 --- a/src/memu/prompts/memory_type/skill.py +++ /dev/null @@ -1,229 +0,0 @@ -PROMPT_BLOCK_OBJECTIVE = """ -# Task Objective -You are a professional User Memory Extractor. Your core task is to extract skills, capabilities, and technical competencies demonstrated or described in the resource content (agent logs, workflow documentation, execution traces, or technical documents). Format each skill as a comprehensive, production-ready skill profile that can be referenced and applied. -""" - -PROMPT_BLOCK_WORKFLOW = """ -# Workflow -Read the full resource content to understand the context and technical details. -## Extract skills -Identify valuable skills, capabilities, and technical competencies demonstrated in the content. -## Create skill profiles -For each skill, create a comprehensive profile with all required sections. -## Review & validate -Ensure each skill profile is complete, actionable, and meets the minimum 300 words requirement. -## Final output -Output Skill Information as structured skill profiles. -""" - -PROMPT_BLOCK_RULES = """ -# Rules -## General requirements (must satisfy all) -- Each skill must be formatted as a comprehensive skill profile with frontmatter and all required sections. -- Each skill profile must capture not just WHAT was done, but HOW and WHY it works. -- Be specific and concrete - include technology names, version numbers, metrics, and outcomes. -- Each skill should be comprehensive enough to be referenced and applied independently. -- Minimum 300 words per skill to ensure depth and actionability. -Important: Extract only skills that are clearly demonstrated or described in the resource. No guesses or fabricated details. - -## Skill Profile Structure (must include all sections) -1. Frontmatter: name, description, category, demonstrated-in -2. Introduction paragraph -3. Core Principles -4. When to Use This Skill -5. Implementation Guide (Prerequisites, Techniques and Approaches, Example from Resource) -6. Success Patterns -7. Common Pitfalls -8. Key Takeaways - -## Special rules for Skill Information -- Generic statements without concrete approaches are forbidden (e.g., "Used Docker", "Good at programming"). -- Opinions without demonstrated practice are forbidden (e.g., "I think microservices are better"). -- Theory without practice belongs to knowledge type, not skill type. -- One-time luck without a replicable approach is not a skill. -- Trivial actions are not skills (e.g., "Using email", "Reading docs"). - -## What TO Extract -- Concrete approaches with context, tools, metrics, and outcomes. -- Deployment strategies with specific techniques (canary, blue-green, etc.). -- Incident response procedures with detection, response, and recovery steps. -- Problem-solving approaches with tool orchestration patterns. -- Multi-step workflows with reasoning steps and validation approaches. - -## Resource Type Guidelines -### For Deployment Logs: -- Extract each significant deployment (success or failure) as a separate skill. -- Success: Focus on techniques that worked. -- Failure: Focus on incident response, root cause analysis, recovery procedures. -- Include metrics: deployment time, error rates, response times, recovery time. - -### For Workflow Documentation: -- Extract major workflow stages as skills. -- Include tool chains and technology stacks. -- Document step-by-step procedures. -- Note success metrics and KPIs. - -### For Agent Execution Logs: -- Extract problem-solving approaches as skills. -- Include tool orchestration patterns. -- Document reasoning steps and validation approaches. -- Capture multi-step workflows. - -## Review & validation rules -- Ensure all required sections are present in each skill profile. -- Verify minimum 300 words per skill. -- Final check: every skill profile must be actionable and replicable. -""" - -PROMPT_BLOCK_CATEGORY = """ -## Memory Categories: -{categories_str} -""" - -PROMPT_BLOCK_OUTPUT = """ -# Output Format (XML) -Return all memories wrapped in a single element: - - - ---- -name: skill-name-in-kebab-case -description: One-line description of what this skill enables -category: primary-category -demonstrated-in: [context where this was shown] ---- - -[Brief introduction explaining the skill and its importance] - -## Core Principles -- [Key concept 1] -- [Key concept 2] - -## When to Use This Skill -- [Situation 1] -- [Situation 2] - -## Implementation Guide -### Prerequisites -- [Required knowledge or setup] - -### Techniques and Approaches -[Detailed explanation of how to apply this skill] - -### Example from Resource -[Concrete example from the source material] - -## Success Patterns -- [Pattern 1 with explanation] -- [Pattern 2 with explanation] - -## Common Pitfalls -- [Pitfall 1]: [Why it's a problem and how to avoid it] -- [Pitfall 2]: [Why it's a problem and how to avoid it] - -## Key Takeaways -- [Critical insight 1] -- [Critical insight 2] - - - technical_skills - - - -""" - -PROMPT_BLOCK_EXAMPLES = """ -# Examples (Input / Output / Explanation) -Example 1: Skill Extraction from Deployment Log -## Input -[2024-01-15 10:30:00] Starting canary deployment for Payment Service v2.3.1 -[2024-01-15 10:30:15] Deployed new version alongside existing v2.3.0 -[2024-01-15 10:30:30] Configured load balancer: 10% traffic to v2.3.1 -[2024-01-15 10:35:30] Metrics check: Error rate 0.1% (baseline 0.1%), p95 latency 270ms (baseline 320ms) -[2024-01-15 10:35:45] Increasing traffic to 50% -[2024-01-15 10:40:45] Metrics stable, increasing to 100% -[2024-01-15 10:42:00] Deployment complete. Zero downtime achieved. -## Output - - - ---- -name: canary-deployment-with-monitoring -description: Implement gradual traffic shifting deployment strategy with real-time monitoring -category: deployment -demonstrated-in: [Payment Service v2.3.1 deployment] ---- - -Canary deployment is a risk-mitigation strategy that gradually shifts production traffic from an old version to a new version while continuously monitoring key metrics. This approach enables early detection of issues with minimal user impact. - -## Core Principles -- Gradual exposure: Start with a small percentage of traffic (typically 5-10%) to limit blast radius -- Continuous monitoring: Track error rates, response times in real-time -- Quick recovery: Maintain ability to instantly route traffic back to stable version - -## When to Use This Skill -- Deploying critical services where downtime is unacceptable -- Rolling out changes with uncertain production behavior -- High-traffic services where testing production performance is valuable - -## Implementation Guide -### Prerequisites -- Load balancer with traffic splitting capabilities -- Monitoring system with real-time metrics -- Automated deployment pipeline - -### Techniques and Approaches -1. Initial Deployment (10% traffic): Deploy new version alongside existing, route 10% traffic, monitor 5 minutes -2. Monitoring Checkpoints: Error rate should not exceed baseline by more than 2%, response time within 20% of baseline -3. Gradual Rollout: If metrics stable, progress 10% to 25% to 50% to 100% - -### Example from Resource -Payment Service v2.3.1 deployment achieved zero downtime during 12-minute deployment. Traffic progressed 10% to 50% to 100% with 5-minute pauses. Response time improved 15% (320ms to 270ms p95). Error rate remained stable at 0.1%. - -## Success Patterns -- Small initial percentage: 5-10% catches most issues while limiting impact -- Metric-driven automation: Removes human error from rollback decisions - -## Common Pitfalls -- Too aggressive progression: Rushing from 10% to 100% defeats the purpose -- Insufficient monitoring window: Need 5+ minutes at each stage to detect issues - -## Key Takeaways -- Canary deployments trade deployment speed for safety -- Start small, progress gradually, monitor continuously - - - technical_skills - - - -## Explanation -A comprehensive skill profile is extracted from the deployment log, capturing the approach, techniques, metrics, and outcomes. The profile includes all required sections and provides actionable guidance for replicating the skill. -""" - -PROMPT_BLOCK_INPUT = """ -# Original Resource: - -{resource} - -""" - -PROMPT = "\n\n".join([ - PROMPT_BLOCK_OBJECTIVE.strip(), - PROMPT_BLOCK_WORKFLOW.strip(), - PROMPT_BLOCK_RULES.strip(), - PROMPT_BLOCK_CATEGORY.strip(), - PROMPT_BLOCK_OUTPUT.strip(), - PROMPT_BLOCK_EXAMPLES.strip(), - PROMPT_BLOCK_INPUT.strip(), -]) - -CUSTOM_PROMPT = { - "objective": PROMPT_BLOCK_OBJECTIVE.strip(), - "workflow": PROMPT_BLOCK_WORKFLOW.strip(), - "rules": PROMPT_BLOCK_RULES.strip(), - "category": PROMPT_BLOCK_CATEGORY.strip(), - "output": PROMPT_BLOCK_OUTPUT.strip(), - "examples": PROMPT_BLOCK_EXAMPLES.strip(), - "input": PROMPT_BLOCK_INPUT.strip(), -} diff --git a/src/memu/prompts/memory_type/tool.py b/src/memu/prompts/memory_type/tool.py deleted file mode 100644 index 91aec4e2..00000000 --- a/src/memu/prompts/memory_type/tool.py +++ /dev/null @@ -1,120 +0,0 @@ -PROMPT_BLOCK_OBJECTIVE = """ -# Task Objective -You are a professional Tool Memory Extractor. Your core task is to extract tool usage patterns, execution results, and learnings from agent logs or tool execution traces. This enables agents to learn from their tool usage history. -""" - -PROMPT_BLOCK_WORKFLOW = """ -# Workflow -Read the full resource content to understand tool execution context. -## Extract tool memories -Identify tool calls, their inputs, outputs, success/failure status, and any patterns. -## Create tool memory entries -For each significant tool usage, create a memory entry with when_to_use hints. -## Review & validate -Ensure each tool memory is actionable and helps future tool selection. -## Final output -Output Tool Memory entries. -""" - -PROMPT_BLOCK_RULES = """ -# Rules -## General requirements (must satisfy all) -- Each tool memory must capture: tool name, what it was used for, outcome, and when to use it again. -- Focus on patterns that help future tool selection decisions. -- Include success/failure context to help agents avoid repeated mistakes. -- Each memory should help answer: "When should I use this tool?" - -## What TO Extract -- Successful tool usage patterns with context -- Failed tool attempts with lessons learned -- Tool combinations that work well together -- Performance insights (fast vs slow tools for different tasks) - -## What NOT to Extract -- Trivial tool calls without learning value -- Duplicate patterns already captured -- Tool calls with no meaningful outcome - -## Memory Item Content Requirements -- Include the tool name prominently -- Describe the use case or scenario -- Note the outcome (success/failure/partial) -- Provide a "when_to_use" hint for future retrieval -""" - -PROMPT_BLOCK_CATEGORY = """ -## Memory Categories: -{categories_str} -""" - -PROMPT_BLOCK_OUTPUT = """ -# Output Format (XML) -Return all memories wrapped in a single element: - - - Tool memory content describing the tool usage pattern - Hint for when this memory should be retrieved - - Category Name - - - -""" - -PROMPT_BLOCK_EXAMPLES = """ -# Examples (Input / Output / Explanation) -Example 1: Tool Memory Extraction -## Input -[Tool Call] file_reader(path="/data/config.json") -[Result] Success - Read 2048 bytes in 0.3s -[Tool Call] json_parser(content=) -[Result] Success - Parsed config with 15 keys -[Tool Call] file_reader(path="/data/missing.json") -[Result] Error - FileNotFoundError: File does not exist -## Output - - - The file_reader tool successfully reads JSON config files from /data/ directory. Average read time is 0.3s for ~2KB files. Works well when chained with json_parser for config processing. - When needing to read configuration files or JSON data from the filesystem - - file_operations - - - - The file_reader tool fails with FileNotFoundError when the target file doesn't exist. Should verify file existence before reading or handle the error gracefully. - When handling file read errors or implementing robust file operations - - error_handling - - - -## Explanation -Two tool memories are extracted: one for successful usage pattern, one for error handling insight. Both include when_to_use hints for smart retrieval. -""" - -PROMPT_BLOCK_INPUT = """ -# Original Resource: - -{resource} - -""" - -PROMPT = "\n\n".join([ - PROMPT_BLOCK_OBJECTIVE.strip(), - PROMPT_BLOCK_WORKFLOW.strip(), - PROMPT_BLOCK_RULES.strip(), - PROMPT_BLOCK_CATEGORY.strip(), - PROMPT_BLOCK_OUTPUT.strip(), - PROMPT_BLOCK_EXAMPLES.strip(), - PROMPT_BLOCK_INPUT.strip(), -]) - -CUSTOM_PROMPT = { - "objective": PROMPT_BLOCK_OBJECTIVE.strip(), - "workflow": PROMPT_BLOCK_WORKFLOW.strip(), - "rules": PROMPT_BLOCK_RULES.strip(), - "category": PROMPT_BLOCK_CATEGORY.strip(), - "output": PROMPT_BLOCK_OUTPUT.strip(), - "examples": PROMPT_BLOCK_EXAMPLES.strip(), - "input": PROMPT_BLOCK_INPUT.strip(), -} diff --git a/src/memu/utils/tool.py b/src/memu/utils/tool.py deleted file mode 100644 index bfb72f4a..00000000 --- a/src/memu/utils/tool.py +++ /dev/null @@ -1,102 +0,0 @@ -"""Utility functions for tool memory operations.""" - -from __future__ import annotations - -from typing import TYPE_CHECKING, Any - -if TYPE_CHECKING: - from memu.database.models import Entry, ToolCallResult - - -def get_tool_calls(item: Entry) -> list[dict[str, Any]]: - """Get tool calls from an entry's extra field. - - Args: - item: The Entry to get tool calls from - - Returns: - List of tool call dicts, or empty list if none exist - """ - result: list[dict[str, Any]] = (item.extra or {}).get("tool_calls", []) - return result - - -def set_tool_calls(item: Entry, tool_calls: list[dict[str, Any]]) -> None: - """Set tool calls in an entry's extra field. - - Args: - item: The Entry to set tool calls on - tool_calls: The list of tool call dicts to set - """ - if item.extra is None: - item.extra = {} - item.extra["tool_calls"] = tool_calls - - -def add_tool_call(item: Entry, tool_call: ToolCallResult) -> None: - """Add a tool call result to an entry (for tool-kind entries). - - Args: - item: The Entry to add the tool call to (must be tool kind) - tool_call: The ToolCallResult to add - - Raises: - ValueError: If the entry is not of kind 'tool' - """ - if item.entry_kind != "tool": - msg = "add_tool_call can only be used with tool type memories" - raise ValueError(msg) - tool_call.ensure_hash() - tool_calls = get_tool_calls(item) - tool_calls.append(tool_call.model_dump()) - set_tool_calls(item, tool_calls) - - -def get_tool_statistics(item: Entry, recent_n: int = 20) -> dict[str, Any]: - """Calculate statistics for the most recent N tool calls. - - Args: - item: The MemoryItem to calculate statistics for - recent_n: Number of recent calls to analyze (default: 20) - - Returns: - Dictionary with total_calls, recent_calls_analyzed, avg_time_cost, - success_rate, avg_score, avg_token_cost - """ - tool_calls = get_tool_calls(item) - if not tool_calls: - return { - "total_calls": 0, - "recent_calls_analyzed": 0, - "avg_time_cost": 0.0, - "success_rate": 0.0, - "avg_score": 0.0, - "avg_token_cost": 0.0, - } - - recent_calls = tool_calls[-recent_n:] - recent_count = len(recent_calls) - - # Calculate statistics (tool_calls are now dicts, not ToolCallResult objects) - total_time = sum(c.get("time_cost", 0.0) for c in recent_calls) - avg_time_cost = total_time / recent_count if recent_count > 0 else 0.0 - - successful = sum(1 for c in recent_calls if c.get("success", True)) - success_rate = successful / recent_count if recent_count > 0 else 0.0 - - total_score = sum(c.get("score", 0.0) for c in recent_calls) - avg_score = total_score / recent_count if recent_count > 0 else 0.0 - - valid_token_calls = [c for c in recent_calls if c.get("token_cost", -1) >= 0] - avg_token_cost = ( - sum(c.get("token_cost", 0) for c in valid_token_calls) / len(valid_token_calls) if valid_token_calls else 0.0 - ) - - return { - "total_calls": len(tool_calls), - "recent_calls_analyzed": recent_count, - "avg_time_cost": round(avg_time_cost, 3), - "success_rate": round(success_rate, 4), - "avg_score": round(avg_score, 3), - "avg_token_cost": round(avg_token_cost, 2), - } diff --git a/tests/test_backend_conformance.py b/tests/test_backend_conformance.py index 97164a50..58f76b62 100644 --- a/tests/test_backend_conformance.py +++ b/tests/test_backend_conformance.py @@ -6,7 +6,7 @@ - ``clear_*`` with a ``where`` scope mutates the shared state in place (no rebinding that orphans the ``DatabaseState`` reference). -- the SQLite read path preserves ``extra`` (reinforcement / ref_id / tool metadata). +- the SQLite read path preserves ``extra`` (reinforcement / ref_id metadata). - deleting an entry / clearing memory leaves no orphan ``ResourceEntry`` relations. """ @@ -65,7 +65,7 @@ def _seed_entry(store, *, text: str, user_id: str, embedding=None): entry = store.entry_repo.create_entry( lane="memory", source_id=res.id, - entry_kind="knowledge", + entry_type="knowledge", text=text, embedding=embedding or [0.1, 0.2, 0.3], user_data={"user_id": user_id}, @@ -145,31 +145,31 @@ def test_clear_relations_with_scope(store): assert store.resource_entry_repo.list_relations({"user_id": "alice"}) == [] -def test_extra_round_trips_through_create_and_read(store): - """``extra`` (tool metadata / ref_id / reinforcement) must survive a read.""" +def test_extra_round_trips_through_update_and_read(store): + """``extra`` (ref_id / reinforcement metadata) must survive a read.""" res = store.resource_repo.create_resource( lane="source", - url="mem://tool", + url="mem://extra", modality="document", local_path="", - summary="tool", + summary="extra", embedding=None, user_data={"user_id": "alice"}, ) entry = store.entry_repo.create_entry( lane="memory", source_id=res.id, - entry_kind="tool", - text="tool memory", + entry_type="knowledge", + text="memory with extra", embedding=[0.1, 0.2, 0.3], user_data={"user_id": "alice"}, - tool_record={"when_to_use": "always"}, ) - assert entry.extra.get("when_to_use") == "always" + updated = store.entry_repo.update_entry(entry_id=entry.id, extra={"ref_id": "always"}) + assert updated.extra.get("ref_id") == "always" fetched = store.entry_repo.get_entry(entry.id) assert fetched is not None - assert fetched.extra.get("when_to_use") == "always" + assert fetched.extra.get("ref_id") == "always" def _reconcile(crud_self, store, *, entry_id, new_cat_names, mapped_old_cat_ids, name_to_id): @@ -255,23 +255,25 @@ async def embed(self, texts): return [[float(len(t)), 1.0, 0.0] for t in texts] -def test_resolve_category_ids_creates_unknown_adaptively(store): - """Open/adaptive taxonomy: extractor-proposed names are created on first sight.""" +def test_resolve_group_ids_creates_unknown_adaptively(store): + """Open/adaptive taxonomy: extractor-proposed names are created on first sight (per lane).""" import asyncio from types import SimpleNamespace from memu.app.memorize import MemorizeMixin from memu.app.service import Context - ctx = Context(categories_ready=True) + ctx = Context() + ctx.lane("memory").ready = True fake_self = SimpleNamespace( _get_embedding_client=lambda profile=None: _FakeEmbedClient(), - _partition_category_names=MemorizeMixin._partition_category_names, + _partition_group_names=MemorizeMixin._partition_group_names, ) ids = asyncio.run( - MemorizeMixin._resolve_category_ids( + MemorizeMixin._resolve_group_ids( fake_self, # type: ignore[arg-type] + "memory", ["Programming", "programming", "Cooking"], ctx, store, @@ -285,39 +287,175 @@ def test_resolve_category_ids_creates_unknown_adaptively(store): # A subsequent call reuses the cached ids and creates nothing new. ids2 = asyncio.run( - MemorizeMixin._resolve_category_ids( + MemorizeMixin._resolve_group_ids( fake_self, # type: ignore[arg-type] + "memory", ["Programming"], ctx, store, user={"user_id": "alice"}, ) ) - assert ids2 == [ctx.category_name_to_id["programming"]] + assert ids2 == [ctx.lane("memory").name_to_id["programming"]] assert len(store.resource_repo.list_resources(lane="memory")) == 2 +def test_persist_lane_entries_per_resource_and_adaptive(store): + """Per-lane persistence: index=1:1 coarse doc, skill=adaptive group docs.""" + import asyncio + + from memu.app.service import Context, MemoryService + from memu.app.settings import LaneConfig + + svc = MemoryService(database_config={"metadata_store": {"provider": "inmemory"}}) + # Avoid real embedding clients when resolving unknown adaptive group names. + svc._get_embedding_client = lambda profile=None, **kw: _FakeEmbedClient() # type: ignore[method-assign] + + ctx = Context() + src = store.resource_repo.create_resource( + lane="source", + url="mem://doc.txt", + modality="document", + local_path="", + summary="caption", + embedding=[1.0, 0.0, 0.0], + user_data={"user_id": "alice"}, + ) + embed = _FakeEmbedClient() + + # index lane (per_resource): one coarse doc 1:1 with the source. + index_cfg = LaneConfig(grouping="per_resource", entry_types=["description"]) + items, rels, _ = asyncio.run( + svc._persist_lane_entries( + lane="index", + lane_cfg=index_cfg, + source_resource=src, + structured_entries=[("description", "A document about cats", [])], + ctx=ctx, + store=store, + embed_client=embed, + user={"user_id": "alice"}, + ) + ) + assert len(items) == 1 and items[0].lane == "index" + index_docs = list(store.resource_repo.list_resources(lane="index").values()) + assert len(index_docs) == 1 + assert index_docs[0].summary == "A document about cats" + assert len(rels) == 1 + + # skill lane (adaptive): tool/log entries grouped under a named skill doc. + skill_cfg = LaneConfig(grouping="adaptive", entry_types=["tool", "log"]) + ctx.lane("skill").ready = True + items2, _, updates2 = asyncio.run( + svc._persist_lane_entries( + lane="skill", + lane_cfg=skill_cfg, + source_resource=src, + structured_entries=[ + ("tool", "Search files with rg", ["Search"]), + ("log", "Ran rg and found 3 matches", ["Search"]), + ], + ctx=ctx, + store=store, + embed_client=embed, + user={"user_id": "alice"}, + ) + ) + assert {i.entry_type for i in items2} == {"tool", "log"} + skill_titles = {d.title for d in store.resource_repo.list_resources(lane="skill").values()} + assert skill_titles == {"Search"} + assert len(updates2) == 1 # exactly one skill group doc was touched + + +def test_rag_recall_lanes_returns_per_lane_shape(store): + """RAG recall fans out over every enabled lane and returns a per-lane response.""" + import asyncio + + from memu.app.service import Context, MemoryService + + svc = MemoryService( + database_config={"metadata_store": {"provider": "inmemory"}}, + retrieve_config={"method": "rag", "route_intention": False}, + ) + fake = _FakeEmbedClient() + svc._get_step_embedding_client = lambda step_context=None: fake # type: ignore[method-assign] + + for lane in ("index", "memory", "skill"): + store.resource_repo.create_resource( + lane=lane, + modality="markdown", + title=f"{lane}-doc", + summary=f"{lane} summary", + embedding=[1.0, 0.0, 0.0], + user_data={"user_id": "alice"}, + ) + store.entry_repo.create_entry( + lane=lane, + source_id=None, + entry_type="x", + text=f"{lane} entry", + embedding=[1.0, 0.0, 0.0], + user_data={"user_id": "alice"}, + ) + store.resource_repo.create_resource( + lane="source", + url="mem://src.txt", + modality="document", + local_path="", + summary="raw source", + embedding=[1.0, 0.0, 0.0], + user_data={"user_id": "alice"}, + ) + + state = { + "needs_retrieval": True, + "active_query": "q", + "retrieve_category": True, + "retrieve_item": True, + "retrieve_resource": True, + "ctx": Context(), + "store": store, + "where": {}, + "original_query": "q", + "rewritten_query": "q", + "next_step_query": None, + } + state = asyncio.run(svc._rag_recall_lanes(state, None)) + state = asyncio.run(svc._rag_recall_resources(state, None)) + state = svc._rag_build_context(state, None) + resp = state["response"] + + assert set(resp["lanes"]) == {"index", "memory", "skill"} + for lane in ("index", "memory", "skill"): + assert len(resp["lanes"][lane]["categories"]) == 1 + assert len(resp["lanes"][lane]["items"]) == 1 + assert len(resp["resources"]) == 1 + # Backward-compatible top-level view mirrors the memory lane. + assert resp["categories"] == resp["lanes"]["memory"]["categories"] + assert resp["items"] == resp["lanes"]["memory"]["items"] + + def test_sqlite_extra_survives_cache_miss(tmp_path): """A fresh SQLite store (cold cache) must reconstruct ``extra`` from the DB.""" db, dsn = _make_sqlite(tmp_path) res = db.resource_repo.create_resource( lane="source", - url="mem://tool", + url="mem://extra", modality="document", local_path="", - summary="tool", + summary="extra", embedding=None, user_data={"user_id": "alice"}, ) entry = db.entry_repo.create_entry( lane="memory", source_id=res.id, - entry_kind="tool", - text="tool memory", + entry_type="knowledge", + text="memory with extra", embedding=[0.1, 0.2, 0.3], user_data={"user_id": "alice"}, - tool_record={"when_to_use": "cold-read"}, ) + db.entry_repo.update_entry(entry_id=entry.id, extra={"ref_id": "cold-read"}) entry_id = entry.id db.close() @@ -327,9 +465,9 @@ def test_sqlite_extra_survives_cache_miss(tmp_path): try: fetched = db2.entry_repo.get_entry(entry_id) assert fetched is not None - assert fetched.extra.get("when_to_use") == "cold-read" + assert fetched.extra.get("ref_id") == "cold-read" listed = db2.entry_repo.list_entries() - assert listed[entry_id].extra.get("when_to_use") == "cold-read" + assert listed[entry_id].extra.get("ref_id") == "cold-read" finally: db2.close() diff --git a/tests/test_folder_memorize.py b/tests/test_folder_memorize.py index efc9f980..94032505 100644 --- a/tests/test_folder_memorize.py +++ b/tests/test_folder_memorize.py @@ -107,7 +107,7 @@ def _seed_resource_with_item(service: MemoryService, *, url: str, category_id: s item = store.entry_repo.create_entry( lane="memory", source_id=res.id, - entry_kind="profile", + entry_type="profile", text=f"summary for {url}", embedding=[0.0], user_data=dict(user), @@ -174,14 +174,14 @@ async def _fake_memorize_one(*, resource_url, modality, user_scope, ctx, store) store.entry_repo.create_entry( lane="memory", source_id=res.id, - entry_kind="profile", + entry_type="profile", text=f"summary {resource_url}", embedding=[0.0], user_data=dict(user_scope or {}), ) return {"resources": [res], "response": {"items": [{"text": "x"}]}} - monkeypatch.setattr(service, "_ensure_categories_ready", _noop_categories) + monkeypatch.setattr(service, "_ensure_lanes_ready", _noop_categories) monkeypatch.setattr(service, "_patch_category_summaries", _noop_patch) monkeypatch.setattr(service, "_memorize_one", _fake_memorize_one) @@ -248,7 +248,7 @@ def _spy_export(database, *, where=None, **kwargs): exported.append(where) return real_export(database, where=where, **kwargs) - monkeypatch.setattr(service, "_ensure_categories_ready", _noop_categories) + monkeypatch.setattr(service, "_ensure_lanes_ready", _noop_categories) monkeypatch.setattr(service, "_memorize_one", _fake_memorize_one) monkeypatch.setattr(service._memory_file_exporter, "export", _spy_export) @@ -287,7 +287,7 @@ async def _fake_memorize_one(*, resource_url, modality, user_scope, ctx, store) def _boom(database, *, where=None, **kwargs): raise RuntimeError("export blew up") # noqa: TRY003 - monkeypatch.setattr(service, "_ensure_categories_ready", _noop_categories) + monkeypatch.setattr(service, "_ensure_lanes_ready", _noop_categories) monkeypatch.setattr(service, "_memorize_one", _fake_memorize_one) monkeypatch.setattr(service._memory_file_exporter, "export", _boom) diff --git a/tests/test_inmemory.py b/tests/test_inmemory.py index 250f15a2..487f022f 100644 --- a/tests/test_inmemory.py +++ b/tests/test_inmemory.py @@ -59,7 +59,7 @@ async def main(): print(f" - {cat.get('name')}: {(cat.get('summary') or cat.get('description', ''))[:80]}...") print(" Items:") for item in result_rag.get("items", [])[:3]: - print(f" - [{item.get('memory_type')}] {item.get('summary', '')[:100]}...") + print(f" - [{item.get('entry_type')}] {item.get('text', '')[:100]}...") if result_rag.get("resources"): print(" Resources:") for res in result_rag.get("resources", [])[:3]: @@ -74,7 +74,7 @@ async def main(): print(f" - {cat.get('name')}: {(cat.get('summary') or cat.get('description', ''))[:80]}...") print(" Items:") for item in result_llm.get("items", [])[:3]: - print(f" - [{item.get('memory_type')}] {item.get('summary', '')[:100]}...") + print(f" - [{item.get('entry_type')}] {item.get('text', '')[:100]}...") if result_llm.get("resources"): print(" Resources:") for res in result_llm.get("resources", [])[:3]: diff --git a/tests/test_memory_files.py b/tests/test_memory_files.py index b64683b3..eccae414 100644 --- a/tests/test_memory_files.py +++ b/tests/test_memory_files.py @@ -9,11 +9,6 @@ from memu.memory_fs import MemoryFileExporter from memu.memory_fs.exporter import MANIFEST_NAME, slugify -# With synthesize=False (the default) the whole tree is LLM-free: MEMORY.md is -# rendered from category summaries and the skill/ tree from the deterministic -# bypass over extracted skill-type memory items. This body seeds one such item. -_SKILL_BODY = "---\nname: pour-over\n---\n# Pour-over brewing\nUse a 1:16 ratio." - def _build_service(output_dir: Path) -> MemoryService: return MemoryService( @@ -43,16 +38,7 @@ def _seed(service: MemoryService, *, user: dict[str, str]) -> dict[str, str]: slug=slugify("Preferences"), ) store.resource_repo.update_resource(resource_id=category.id, summary="The user likes pour-over coffee.") - skill = store.entry_repo.create_entry( - lane="memory", - source_id=resource.id, - entry_kind="skill", - text=_SKILL_BODY, - embedding=[0.1, 0.2], - user_data=dict(user), - ) - store.resource_entry_repo.link_entry_resource(skill.id, category.id, user_data=dict(user)) - return {"category_id": category.id, "resource_id": resource.id, "skill_id": skill.id} + return {"category_id": category.id, "resource_id": resource.id} async def test_export_writes_readme_layout(tmp_path: Path) -> None: @@ -64,9 +50,7 @@ async def test_export_writes_readme_layout(tmp_path: Path) -> None: assert result["changed"] is True assert "INDEX.md" in result["written"] assert "MEMORY.md" in result["written"] - assert "SKILL.md" in result["written"] assert "memory/preferences.md" in result["written"] - assert "skill/pour-over/SKILL.md" in result["written"] # MEMORY.md is now an overview that links to each memory/.md file; the # category summary itself lives in memory/preferences.md. @@ -82,13 +66,6 @@ async def test_export_writes_readme_layout(tmp_path: Path) -> None: assert "coffee.txt" in index_text assert "coffee preferences" in index_text - # The root SKILL.md indexes the synthesized skill/ tree. - skill_index = (tmp_path / "SKILL.md").read_text(encoding="utf-8") - assert "skill/pour-over/SKILL.md" in skill_index - - skill_text = (tmp_path / "skill" / "pour-over" / "SKILL.md").read_text(encoding="utf-8") - assert "Pour-over brewing" in skill_text - async def test_export_is_idempotent_until_data_changes(tmp_path: Path) -> None: service = _build_service(tmp_path) @@ -115,19 +92,19 @@ async def test_export_is_idempotent_until_data_changes(tmp_path: Path) -> None: assert "INDEX.md" in third["unchanged"] -async def test_export_removes_stale_skill_and_prunes_dirs(tmp_path: Path) -> None: +async def test_export_removes_stale_category_and_prunes(tmp_path: Path) -> None: service = _build_service(tmp_path) _seed(service, user={"user_id": "u1"}) await service.export_memory_files(user={"user_id": "u1"}) - assert (tmp_path / "skill" / "pour-over" / "SKILL.md").exists() + assert (tmp_path / "memory" / "preferences.md").exists() - # Dropping the skill-type item removes it from the bypass, so its doc goes stale. - service.database.entry_repo.clear_entries(where={"user_id": "u1"}) + # Dropping the memory-lane category removes its rendered doc on the next export. + service.database.resource_repo.clear_resources(where={"user_id": "u1"}, lane="memory") result = await service.export_memory_files(user={"user_id": "u1"}) - assert "skill/pour-over/SKILL.md" in result["removed"] - assert not (tmp_path / "skill" / "pour-over").exists() + assert "memory/preferences.md" in result["removed"] + assert not (tmp_path / "memory" / "preferences.md").exists() async def test_export_respects_user_scope(tmp_path: Path) -> None: @@ -158,13 +135,6 @@ async def test_export_disabled_raises(tmp_path: Path) -> None: await service.export_memory_files(user={"user_id": "u1"}) -def test_skill_name_from_frontmatter_and_fallbacks(tmp_path: Path) -> None: - exporter = MemoryFileExporter(str(tmp_path)) - assert exporter._skill_name("---\nname: My Skill\n---\nbody", fallback="x") == "my-skill" - assert exporter._skill_name("# Heading Title\nbody", fallback="x") == "heading-title" - assert exporter._skill_name("plain text only", fallback="skill-abc123") == "skill-abc123" - - def test_exporter_manifest_roundtrip(tmp_path: Path) -> None: exporter = MemoryFileExporter(str(tmp_path)) exporter._save_manifest({"MEMORY.md": "abc"}) diff --git a/tests/test_memory_fs_synthesis.py b/tests/test_memory_fs_synthesis.py index 6e6d813e..dcfb8db8 100644 --- a/tests/test_memory_fs_synthesis.py +++ b/tests/test_memory_fs_synthesis.py @@ -6,15 +6,12 @@ from memu.memory_fs import FileDescription, MemoryFileExporter, MemorySynthesizer _MEMORY_MD = "## Profile\nThe user is a coffee enthusiast.\n\n## Preferences\nPrefers pour-over." -_SKILLS_JSON = '[{"name": "Pour Over", "body": "# Pour-over\\nUse a 1:16 ratio."}]' class _FakeChatClient: - """Stand-in LLM client: returns canned memory/skill responses by prompt shape.""" + """Stand-in LLM client: returns a canned memory document.""" async def chat(self, prompt: str, system_prompt: str | None = None) -> str: - if "JSON array" in prompt: - return _SKILLS_JSON return _MEMORY_MD @@ -29,42 +26,22 @@ def _descriptions() -> list[FileDescription]: ] -async def test_synthesizer_parses_memory_and_skills() -> None: +async def test_synthesizer_parses_memory() -> None: synth = MemorySynthesizer() - result = await synth.synthesize(_descriptions(), chat=_FakeChatClient().chat) + memory_body = await synth.synthesize(_descriptions(), chat=_FakeChatClient().chat) - assert "## Profile" in result.memory_body - assert "pour-over" in result.memory_body.lower() - assert result.skills == {"pour-over": "# Pour-over\nUse a 1:16 ratio."} + assert "## Profile" in memory_body + assert "pour-over" in memory_body.lower() async def test_synthesizer_empty_when_no_descriptions() -> None: synth = MemorySynthesizer() - result = await synth.synthesize([], chat=_FakeChatClient().chat) - assert result.memory_body == "" - assert result.skills == {} - - -async def test_synthesize_skills_only_decoupled_from_memory() -> None: - """The skill bypass can be built on its own, without touching MEMORY.md.""" - synth = MemorySynthesizer() - skills = await synth.synthesize_skills(_descriptions(), chat=_FakeChatClient().chat) - assert skills == {"pour-over": "# Pour-over\nUse a 1:16 ratio."} - - -async def test_synthesize_skills_empty_without_descriptions() -> None: - synth = MemorySynthesizer() - assert await synth.synthesize_skills([], chat=_FakeChatClient().chat) == {} + assert await synth.synthesize([], chat=_FakeChatClient().chat) == "" def test_synthesizer_helpers() -> None: synth = MemorySynthesizer() assert synth._clean_markdown("```markdown\n# Hi\n```") == "# Hi" - assert synth._parse_skills("garbage, no array") == {} - assert synth._parse_skills("[]") == {} - assert synth._parse_skills('[{"name": "A", "body": ""}]') == {} - duplicate = '[{"name": "A", "body": "x"}, {"name": "A", "body": "y"}]' - assert synth._parse_skills(duplicate) == {"a": "x", "a-2": "y"} def test_build_synthesis_descriptions_uses_structured_items() -> None: @@ -78,8 +55,8 @@ def test_build_synthesis_descriptions_uses_structured_items() -> None: id="r2", lane="source", url="docs/b.txt", modality="document", local_path="b.txt", summary="raw caption b" ) items = [ - Entry(id="i1", lane="memory", source_id="r1", entry_kind="knowledge", text="Alpha fact."), - Entry(id="i2", lane="memory", source_id="r1", entry_kind="profile", text="Beta trait."), + Entry(id="i1", lane="memory", source_id="r1", entry_type="knowledge", text="Alpha fact."), + Entry(id="i2", lane="memory", source_id="r1", entry_type="profile", text="Beta trait."), ] descriptions = MemoryFileExporter.build_synthesis_descriptions([res_with_items, res_without_items], items) @@ -98,19 +75,10 @@ def test_exporter_override_path(tmp_path: Path) -> None: ) exporter = MemoryFileExporter(str(tmp_path)) - result = exporter.export( - service.database, - memory_body="## Profile\nSynthesized.", - skills={"brewing": "# Brewing\nbody"}, - ) + result = exporter.export(service.database, memory_body="## Profile\nSynthesized.") assert "MEMORY.md" in result.written - assert "SKILL.md" in result.written - assert "skill/brewing/SKILL.md" in result.written assert "Synthesized." in (tmp_path / "MEMORY.md").read_text(encoding="utf-8") - assert "# Brewing" in (tmp_path / "skill" / "brewing" / "SKILL.md").read_text(encoding="utf-8") - # The synthesized skill/ tree is indexed by the root SKILL.md. - assert "skill/brewing/SKILL.md" in (tmp_path / "SKILL.md").read_text(encoding="utf-8") async def test_service_synthesis_wiring(tmp_path: Path, monkeypatch) -> None: @@ -133,7 +101,6 @@ async def test_service_synthesis_wiring(tmp_path: Path, monkeypatch) -> None: result = await service.export_memory_files(user={"user_id": "u1"}) assert "MEMORY.md" in result["written"] - assert "skill/pour-over/SKILL.md" in result["written"] memory_text = (tmp_path / "MEMORY.md").read_text(encoding="utf-8") assert "The user is a coffee enthusiast." in memory_text @@ -141,66 +108,38 @@ async def test_service_synthesis_wiring(tmp_path: Path, monkeypatch) -> None: # -- incremental update path ------------------------------------------------- _UPDATE_MEMORY_MD = "## Profile\nThe user is a coffee enthusiast.\n\n## Preferences\nLikes oat milk." -_UPDATE_SKILLS_JSON = '[{"name": "Latte Art", "body": "# Latte art\\nPour slowly."}]' class _InitUpdateChatClient: """Returns init vs update payloads based on whether existing content was injected. - The unified prompt renders ``(empty)``/``(none)`` when there is no prior artifact, - so the presence of those sentinels marks a from-scratch (init) call. + The unified prompt renders ``(empty)`` when there is no prior artifact, so the + presence of that sentinel marks a from-scratch (init) call. """ async def chat(self, prompt: str, system_prompt: str | None = None) -> str: - if "JSON array" in prompt: - return _SKILLS_JSON if "(none)" in prompt else _UPDATE_SKILLS_JSON return _MEMORY_MD if "(empty)" in prompt else _UPDATE_MEMORY_MD async def test_synthesizer_update_merges_into_existing() -> None: synth = MemorySynthesizer() - result = await synth.synthesize( + memory_body = await synth.synthesize( _descriptions(), existing_memory="## Profile\nOld profile.", - existing_skills={"pour-over": "# Pour-over\nUse a 1:16 ratio."}, chat=_InitUpdateChatClient().chat, ) - assert "Likes oat milk." in result.memory_body - # Existing skill is preserved, the new one is upserted alongside it. - assert result.skills["pour-over"] == "# Pour-over\nUse a 1:16 ratio." - assert result.skills["latte-art"] == "# Latte art\nPour slowly." + assert "Likes oat milk." in memory_body async def test_synthesizer_update_noop_without_descriptions() -> None: synth = MemorySynthesizer() - existing_skills = {"pour-over": "# Pour-over"} - result = await synth.synthesize( + memory_body = await synth.synthesize( [], existing_memory="## Profile\nKeep me.", - existing_skills=existing_skills, chat=_InitUpdateChatClient().chat, ) - assert result.memory_body == "## Profile\nKeep me." - assert result.skills == existing_skills - - -async def test_update_skills_only_upserts_and_preserves() -> None: - """Skill-only incremental update keeps untouched skills and upserts new ones.""" - synth = MemorySynthesizer() - skills = await synth.synthesize_skills( - _descriptions(), - existing_skills={"pour-over": "# Pour-over\nUse a 1:16 ratio."}, - chat=_InitUpdateChatClient().chat, - ) - assert skills["pour-over"] == "# Pour-over\nUse a 1:16 ratio." - assert skills["latte-art"] == "# Latte art\nPour slowly." - - -async def test_update_skills_noop_without_descriptions() -> None: - synth = MemorySynthesizer() - existing = {"pour-over": "# Pour-over"} - assert await synth.synthesize_skills([], existing_skills=existing, chat=_InitUpdateChatClient().chat) == existing + assert memory_body == "## Profile\nKeep me." def test_exporter_read_helpers_roundtrip(tmp_path: Path) -> None: @@ -211,15 +150,10 @@ def test_exporter_read_helpers_roundtrip(tmp_path: Path) -> None: exporter = MemoryFileExporter(str(tmp_path)) assert exporter.artifacts_exist() is False - exporter.export( - service.database, - memory_body="## Profile\nSynthesized body.", - skills={"brewing": "# Brewing\nbody"}, - ) + exporter.export(service.database, memory_body="## Profile\nSynthesized body.") assert exporter.artifacts_exist() is True assert exporter.read_memory_body() == "## Profile\nSynthesized body." - assert exporter.read_skills() == {"brewing": "# Brewing\nbody"} async def test_service_init_then_update(tmp_path: Path, monkeypatch) -> None: @@ -246,8 +180,7 @@ async def test_service_init_then_update(tmp_path: Path, monkeypatch) -> None: ) # First pass: no tree yet -> initialization from the full store. - init = await service.export_memory_files(user={"user_id": "u1"}) - assert "skill/pour-over/SKILL.md" in init["written"] + await service.export_memory_files(user={"user_id": "u1"}) assert "coffee enthusiast" in (tmp_path / "MEMORY.md").read_text(encoding="utf-8") # Second pass: tree exists -> incremental update from the changed resource only. @@ -260,10 +193,7 @@ async def test_service_init_then_update(tmp_path: Path, monkeypatch) -> None: embedding=None, user_data={"user_id": "u1"}, ) - updated = await service._build_memory_files({"user_id": "u1"}, changed=[changed]) + await service._build_memory_files({"user_id": "u1"}, changed=[changed]) memory_text = (tmp_path / "MEMORY.md").read_text(encoding="utf-8") assert "Likes oat milk." in memory_text - assert "skill/latte-art/SKILL.md" in (updated["written"] + updated["unchanged"]) - # The originally-initialized skill survives the incremental update. - assert (tmp_path / "skill" / "pour-over" / "SKILL.md").exists() diff --git a/tests/test_openrouter.py b/tests/test_openrouter.py index 51b0839e..95584277 100644 --- a/tests/test_openrouter.py +++ b/tests/test_openrouter.py @@ -38,9 +38,9 @@ def _print_items(items, max_items=3): if items: print(" Items:") for item in items[:max_items]: - entry_kind = item.get("entry_kind", "unknown") + entry_type = item.get("entry_type", "unknown") summary = item.get("text", "")[:80] - print(f" - [{entry_kind}] {summary}...") + print(f" - [{entry_type}] {summary}...") async def _test_memorize(service, file_path, output_data): diff --git a/tests/test_postgres.py b/tests/test_postgres.py index 8b375a30..877931a8 100644 --- a/tests/test_postgres.py +++ b/tests/test_postgres.py @@ -52,7 +52,7 @@ async def main(): print(f" - {cat.get('name')}: {(cat.get('summary') or cat.get('description', ''))[:80]}...") print(" Items:") for item in result_rag.get("items", [])[:3]: - print(f" - [{item.get('memory_type')}] {item.get('summary', '')[:100]}...") + print(f" - [{item.get('entry_type')}] {item.get('text', '')[:100]}...") if result_rag.get("resources"): print(" Resources:") for res in result_rag.get("resources", [])[:3]: @@ -67,7 +67,7 @@ async def main(): print(f" - {cat.get('name')}: {(cat.get('summary') or cat.get('description', ''))[:80]}...") print(" Items:") for item in result_llm.get("items", [])[:3]: - print(f" - [{item.get('memory_type')}] {item.get('summary', '')[:100]}...") + print(f" - [{item.get('entry_type')}] {item.get('text', '')[:100]}...") if result_llm.get("resources"): print(" Resources:") for res in result_llm.get("resources", [])[:3]: diff --git a/tests/test_salience.py b/tests/test_salience.py index bb5d84a2..6757752b 100644 --- a/tests/test_salience.py +++ b/tests/test_salience.py @@ -13,10 +13,10 @@ # Inline implementations to avoid circular import issues during testing -def compute_content_hash(summary: str, memory_type: str) -> str: +def compute_content_hash(summary: str, entry_type: str) -> str: """Generate unique hash for memory deduplication.""" normalized = " ".join(summary.lower().split()) - content = f"{memory_type}:{normalized}" + content = f"{entry_type}:{normalized}" return hashlib.sha256(content.encode()).hexdigest()[:16] diff --git a/tests/test_sqlite.py b/tests/test_sqlite.py index 13873e85..8e73c066 100644 --- a/tests/test_sqlite.py +++ b/tests/test_sqlite.py @@ -13,7 +13,7 @@ def _print_results(title: str, result: dict) -> None: print(f" - {cat.get('title')}: {(cat.get('summary') or cat.get('description', ''))[:80]}...") print(" Items:") for item in result.get("items", [])[:3]: - print(f" - [{item.get('entry_kind')}] {item.get('text', '')[:100]}...") + print(f" - [{item.get('entry_type')}] {item.get('text', '')[:100]}...") if result.get("resources"): print(" Resources:") for res in result.get("resources", [])[:3]: diff --git a/tests/test_tool_memory.py b/tests/test_tool_memory.py deleted file mode 100644 index 25dd7172..00000000 --- a/tests/test_tool_memory.py +++ /dev/null @@ -1,336 +0,0 @@ -"""Tests for Tool Memory feature - specialized entry kind for tracking tool usage.""" - -from __future__ import annotations - -import importlib.util -import sys -from datetime import datetime -from pathlib import Path -from typing import Any - -# Add src to path for direct import - MUST be before any memu imports -src_path = Path(__file__).parent.parent / "src" -if str(src_path) not in sys.path: - sys.path.insert(0, str(src_path)) - -import pytest # noqa: E402 - -# Import directly from the models file path to avoid circular import through database/__init__.py -# We use importlib to import the module directly without triggering the package __init__ -spec = importlib.util.spec_from_file_location("models", src_path / "memu" / "database" / "models.py") -assert spec is not None -assert spec.loader is not None -models = importlib.util.module_from_spec(spec) -spec.loader.exec_module(models) - -# Rebuild models to resolve forward references with proper namespace -rebuild_ns = { - "Any": Any, - "datetime": datetime, - "MemoryType": models.MemoryType, - "ToolCallResult": models.ToolCallResult, -} -models.ToolCallResult.model_rebuild(_types_namespace=rebuild_ns) -models.Entry.model_rebuild(_types_namespace=rebuild_ns) - -Entry = models.Entry -MemoryType = models.MemoryType -ToolCallResult = models.ToolCallResult - -# Import tool memory utility functions -util_tool_spec = importlib.util.spec_from_file_location("util_tool", src_path / "memu" / "utils" / "tool.py") -assert util_tool_spec is not None -assert util_tool_spec.loader is not None -util_tool = importlib.util.module_from_spec(util_tool_spec) -util_tool_spec.loader.exec_module(util_tool) - -add_tool_call = util_tool.add_tool_call -get_tool_statistics = util_tool.get_tool_statistics - - -class TestToolCallResult: - """Tests for ToolCallResult model.""" - - def test_create_tool_call_result(self): - """Test creating a basic ToolCallResult.""" - result = ToolCallResult( - tool_name="file_reader", - input={"path": "/data/config.json"}, - output="File content here", - success=True, - time_cost=0.5, - token_cost=100, - score=0.95, - ) - - assert result.tool_name == "file_reader" - assert result.input == {"path": "/data/config.json"} - assert result.output == "File content here" - assert result.success is True - assert result.time_cost == 0.5 - assert result.token_cost == 100 - assert result.score == 0.95 - - def test_generate_hash(self): - """Test hash generation for deduplication.""" - result = ToolCallResult( - tool_name="calculator", - input={"a": 1, "b": 2}, - output="3", - ) - - hash1 = result.generate_hash() - assert hash1 != "" - assert len(hash1) == 32 # MD5 hex digest length - - # Same input/output should generate same hash - result2 = ToolCallResult( - tool_name="calculator", - input={"a": 1, "b": 2}, - output="3", - ) - assert result2.generate_hash() == hash1 - - # Different input should generate different hash - result3 = ToolCallResult( - tool_name="calculator", - input={"a": 2, "b": 3}, - output="5", - ) - assert result3.generate_hash() != hash1 - - def test_ensure_hash(self): - """Test ensure_hash sets call_hash if empty.""" - result = ToolCallResult( - tool_name="test_tool", - input="test input", - output="test output", - ) - - assert result.call_hash == "" - result.ensure_hash() - assert result.call_hash != "" - assert len(result.call_hash) == 32 - - def test_string_input(self): - """Test ToolCallResult with string input.""" - result = ToolCallResult( - tool_name="echo", - input="hello world", - output="hello world", - ) - - result.ensure_hash() - assert result.call_hash != "" - - -class TestMemoryItemToolType: - """Tests for Entry with tool kind.""" - - def test_tool_memory_type_literal(self): - """Test that 'tool' is a valid MemoryType.""" - from typing import get_args - - valid_types = get_args(MemoryType) - assert "tool" in valid_types - - def test_create_tool_memory(self): - """Test creating a tool kind entry with tool fields in extra.""" - item = Entry( - lane="memory", - source_id=None, - entry_kind="tool", - text="file_reader tool usage for config files", - extra={ - "when_to_use": "When needing to read configuration files", - "metadata": {"tool_name": "file_reader", "avg_success_rate": 0.95}, - }, - ) - - assert item.entry_kind == "tool" - assert item.extra["when_to_use"] == "When needing to read configuration files" - assert item.extra["metadata"]["tool_name"] == "file_reader" - - def test_add_tool_call(self): - """Test adding tool call results to a tool memory.""" - item = Entry( - lane="memory", - source_id=None, - entry_kind="tool", - text="calculator tool usage", - ) - - tool_call = ToolCallResult( - tool_name="calculator", - input={"a": 1, "b": 2}, - output="3", - success=True, - time_cost=0.1, - score=1.0, - ) - - add_tool_call(item, tool_call) - - tool_calls = item.extra.get("tool_calls", []) - assert len(tool_calls) == 1 - assert tool_calls[0]["tool_name"] == "calculator" - assert tool_calls[0]["call_hash"] != "" # ensure_hash was called - - def test_add_tool_call_wrong_type(self): - """Test that add_tool_call fails for non-tool memories.""" - item = Entry( - lane="memory", - source_id=None, - entry_kind="profile", - text="User profile info", - ) - - tool_call = ToolCallResult( - tool_name="test", - input="test", - output="test", - ) - - with pytest.raises(ValueError, match="can only be used with tool type memories"): - add_tool_call(item, tool_call) - - def test_get_tool_statistics_empty(self): - """Test statistics for memory with no tool calls.""" - item = Entry( - lane="memory", - source_id=None, - entry_kind="tool", - text="empty tool memory", - ) - - stats = get_tool_statistics(item) - - assert stats["total_calls"] == 0 - assert stats["recent_calls_analyzed"] == 0 - assert stats["avg_time_cost"] == 0.0 - assert stats["success_rate"] == 0.0 - assert stats["avg_score"] == 0.0 - assert stats["avg_token_cost"] == 0.0 - - def test_get_tool_statistics(self): - """Test statistics calculation for tool calls.""" - # Tool calls are stored as dicts in extra - item = Entry( - lane="memory", - source_id=None, - entry_kind="tool", - text="calculator tool", - extra={ - "tool_calls": [ - { - "tool_name": "calc", - "input": "1+1", - "output": "2", - "success": True, - "time_cost": 0.1, - "score": 1.0, - "token_cost": 10, - }, - { - "tool_name": "calc", - "input": "2+2", - "output": "4", - "success": True, - "time_cost": 0.2, - "score": 0.9, - "token_cost": 15, - }, - { - "tool_name": "calc", - "input": "bad", - "output": "error", - "success": False, - "time_cost": 0.5, - "score": 0.0, - "token_cost": 5, - }, - ] - }, - ) - - stats = get_tool_statistics(item) - - assert stats["total_calls"] == 3 - assert stats["recent_calls_analyzed"] == 3 - assert stats["success_rate"] == pytest.approx(0.6667, rel=0.01) # 2/3 - assert stats["avg_time_cost"] == pytest.approx(0.267, rel=0.01) # (0.1+0.2+0.5)/3 - assert stats["avg_score"] == pytest.approx(0.633, rel=0.01) # (1.0+0.9+0.0)/3 - assert stats["avg_token_cost"] == pytest.approx(10.0, rel=0.01) # (10+15+5)/3 - - def test_get_tool_statistics_recent_n(self): - """Test statistics with recent_n limit.""" - item = Entry( - lane="memory", - source_id=None, - entry_kind="tool", - text="tool with many calls", - extra={ - "tool_calls": [ - {"tool_name": "t", "input": "1", "output": "1", "success": False, "time_cost": 1.0, "score": 0.0}, - {"tool_name": "t", "input": "2", "output": "2", "success": True, "time_cost": 0.1, "score": 1.0}, - {"tool_name": "t", "input": "3", "output": "3", "success": True, "time_cost": 0.1, "score": 1.0}, - ] - }, - ) - - # Only analyze last 2 calls - stats = get_tool_statistics(item, recent_n=2) - - assert stats["total_calls"] == 3 - assert stats["recent_calls_analyzed"] == 2 - assert stats["success_rate"] == 1.0 # Both recent calls succeeded - - -class TestMemoryItemNewFields: - """Tests for tool-related fields stored in extra.""" - - def test_when_to_use_field(self): - """Test when_to_use field stored in extra for retrieval hints.""" - item = Entry( - lane="memory", - source_id=None, - entry_kind="profile", - text="User prefers dark mode", - extra={"when_to_use": "When configuring UI settings or themes"}, - ) - - assert item.extra["when_to_use"] == "When configuring UI settings or themes" - - def test_metadata_field(self): - """Test metadata field stored in extra for type-specific data.""" - item = Entry( - lane="memory", - source_id=None, - entry_kind="event", - text="User attended conference", - extra={ - "metadata": { - "event_date": "2026-01-15", - "location": "San Francisco", - "attendees": ["Alice", "Bob"], - } - }, - ) - - assert item.extra.get("metadata") is not None - assert item.extra["metadata"]["event_date"] == "2026-01-15" - assert item.extra["metadata"]["location"] == "San Francisco" - assert len(item.extra["metadata"]["attendees"]) == 2 - - def test_default_values(self): - """Test that extra defaults to empty dict.""" - item = Entry( - lane="memory", - source_id=None, - entry_kind="knowledge", - text="Python is a programming language", - ) - - assert item.extra.get("when_to_use") is None - assert item.extra.get("metadata") is None - assert item.extra.get("tool_calls") is None