Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
42 changes: 42 additions & 0 deletions src/api/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,7 @@
get_memory_by_id,
insert_memory,
get_recent_memories,
regenerate_embedding,
)
from ..db.queries import get_memory_stats
from ..embedder import create_embedding
Expand Down Expand Up @@ -238,6 +239,47 @@ async def search_memories_endpoint(search: SearchRequest):
return results


class RegenerateEmbeddingRequest(BaseModel):
force: bool = Field(
default=False,
description="Overwrite an existing embedding (e.g. after a model change)",
)


@app.post("/memories/{memory_id}/regenerate-embedding")
async def regenerate_memory_embedding(memory_id: str, request: RegenerateEmbeddingRequest):
"""Re-generate the embedding for an existing memory.

Useful when the original embedding failed at store time (NULL embedding)
or after switching embedding providers/models.
"""
import uuid as _uuid

try:
parsed_id = _uuid.UUID(memory_id)
except ValueError as exc:
raise HTTPException(status_code=400, detail=f"Invalid UUID: {memory_id}") from exc

memory = get_memory_by_id(parsed_id)
if not memory:
raise HTTPException(status_code=404, detail="Memory not found")

try:
embedding = create_embedding(memory["content"])
except Exception as exc:
raise HTTPException(
status_code=502,
detail=f"Embedding generation failed: {type(exc).__name__}",
) from exc

try:
regenerate_embedding(parsed_id, embedding, force=request.force)
except ValueError as exc:
raise HTTPException(status_code=409, detail=str(exc)) from exc

return {"id": str(parsed_id), "status": "regenerated"}


@app.get("/stats")
async def get_stats():
"""Get memory statistics."""
Expand Down
42 changes: 42 additions & 0 deletions src/db/attribution.py
Original file line number Diff line number Diff line change
Expand Up @@ -164,6 +164,48 @@ def get_memory_by_id(memory_id: uuid.UUID) -> Optional[Dict[str, Any]]:
return _decode_memory(row) if row else None


def regenerate_embedding(
memory_id: uuid.UUID,
new_embedding: List[float],
*,
force: bool = False,
) -> Optional[Dict[str, Any]]:
"""Re-generate the embedding for an existing memory.

By default only memories with a NULL embedding are updated. Pass
``force=True`` to overwrite a valid embedding (e.g. after a model or
dimension change).

Returns the updated memory row, or ``None`` when the memory does not
exist. Raises ``ValueError`` when the embedding is already present
and *force* is ``False``.
"""
with get_db_cursor() as cursor:
cursor.execute(
f"SELECT {_MEMORY_COLUMNS}, embedding IS NOT NULL AS has_embedding "
"FROM memory WHERE id = %s",
(memory_id,),
)
row = cursor.fetchone()
if row is None:
return None

if row["has_embedding"] and not force:
raise ValueError(
"memory already has an embedding; pass force=True to overwrite"
)

cursor.execute(
"UPDATE memory SET embedding = %s WHERE id = %s",
(new_embedding, memory_id),
)
cursor.execute(
f"SELECT {_MEMORY_COLUMNS} FROM memory WHERE id = %s",
(memory_id,),
)
return _decode_memory(cursor.fetchone())


def get_recent_memories(
limit: int = 50,
offset: int = 0,
Expand Down
58 changes: 57 additions & 1 deletion src/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,7 @@
from mcp.types import Tool, TextContent

from .db import connection
from .db.attribution import insert_memory, search_memories
from .db.attribution import insert_memory, search_memories, regenerate_embedding, get_memory_by_id
from .db.queries import (
get_related_memories,
get_memories_by_entity,
Expand Down Expand Up @@ -139,6 +139,26 @@ async def list_tools() -> List[Tool]:
},
},
),
Tool(
name="memory_regenerate_embedding",
description=(
"Re-generate the embedding for an existing memory. "
"Use when the original embedding failed (NULL) or after "
"switching embedding providers/models."
),
inputSchema={
"type": "object",
"properties": {
"memory_id": {"type": "string", "description": "UUID of the memory"},
"force": {
"type": "boolean",
"description": "Overwrite an existing embedding (default false)",
"default": False,
},
},
"required": ["memory_id"],
},
),
]


Expand All @@ -160,6 +180,8 @@ async def call_tool(name: str, arguments: Any) -> List[TextContent]:
return await handle_memory_stats(arguments)
if name == "memory_weekly_report":
return await handle_weekly_report(arguments)
if name == "memory_regenerate_embedding":
return await handle_memory_regenerate_embedding(arguments)
return [TextContent(type="text", text=f"Unknown tool: {name}")]
except Exception as exc:
return [TextContent(type="text", text=f"Error: {str(exc)}")]
Expand Down Expand Up @@ -276,6 +298,40 @@ async def handle_weekly_report(args: Dict) -> List[TextContent]:
return [TextContent(type="text", text=generate_weekly_report(args.get("days", 7)))]


async def handle_memory_regenerate_embedding(args: Dict) -> List[TextContent]:
"""Handle memory_regenerate_embedding tool."""
memory_id = args.get("memory_id", "")
force = args.get("force", False)

import uuid
try:
parsed_id = uuid.UUID(memory_id)
except ValueError:
return [TextContent(type="text", text=f"Invalid memory ID: {memory_id}")]

memory = get_memory_by_id(parsed_id)
if not memory:
return [TextContent(type="text", text=f"Memory not found: {memory_id}")]

try:
embedding = create_embedding(memory["content"])
except Exception as exc:
return [TextContent(
type="text",
text=f"Embedding generation failed: {type(exc).__name__}: {exc}",
)]

try:
regenerate_embedding(parsed_id, embedding, force=force)
except ValueError as exc:
return [TextContent(type="text", text=str(exc))]

return [TextContent(
type="text",
text=str({"id": str(parsed_id), "status": "regenerated"}),
)]


def format_memory_list(memories: List[Dict]) -> str:
"""Format a list of memories for display."""
if not memories:
Expand Down
141 changes: 141 additions & 0 deletions tests/test_embedding_regeneration.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,141 @@
"""Tests for stale-embedding regeneration (MCP tool + REST endpoint + DB layer)."""
import asyncio
from contextlib import contextmanager
from unittest.mock import Mock, patch

import pytest


def _cursor(rows=None, row=None):
cursor = Mock()
cursor.fetchall.return_value = rows or []
cursor.fetchone.return_value = row

@contextmanager
def manager():
yield cursor

return cursor, manager


# --- DB layer ---

def test_regenerate_returns_none_for_missing_memory():
from src.db import attribution

cursor, manager = _cursor(row=None)
with patch.object(attribution, "get_db_cursor", manager):
result = attribution.regenerate_embedding(
__import__("uuid").uuid4(), [0.1, 0.2], force=False
)
assert result is None


def test_regenerate_raises_when_embedding_exists_and_no_force():
from src.db import attribution

row = {
"id": "mem-1", "source": "mcp", "source_id": None,
"captured_by": None, "content": "hello", "raw_content": None,
"entities": "{}", "tags": [], "tag_sources": "{}",
"importance": 0.5, "created_at": None, "original_date": None,
"language": None, "metadata": "{}", "has_embedding": True,
}
cursor, manager = _cursor(row=row)
with patch.object(attribution, "get_db_cursor", manager):
with pytest.raises(ValueError, match="force=True"):
attribution.regenerate_embedding(
__import__("uuid").uuid4(), [0.1, 0.2], force=False
)


def test_regenerate_updates_null_embedding():
from src.db import attribution

mem_id = __import__("uuid").uuid4()
row_before = {
"id": str(mem_id), "source": "mcp", "source_id": None,
"captured_by": None, "content": "hello", "raw_content": None,
"entities": "{}", "tags": [], "tag_sources": "{}",
"importance": 0.5, "created_at": None, "original_date": None,
"language": None, "metadata": "{}", "has_embedding": False,
}
row_after = dict(row_before)
del row_after["has_embedding"]

cursor, manager = _cursor()
cursor.fetchone.side_effect = [row_before, row_after]
with patch.object(attribution, "get_db_cursor", manager):
result = attribution.regenerate_embedding(mem_id, [0.1, 0.2])

assert result is not None
calls = [c.args[0].strip() for c in cursor.execute.call_args_list]
assert any("UPDATE memory SET embedding" in c for c in calls)


def test_regenerate_force_overwrites_existing():
from src.db import attribution

mem_id = __import__("uuid").uuid4()
row_before = {
"id": str(mem_id), "source": "mcp", "source_id": None,
"captured_by": None, "content": "hello", "raw_content": None,
"entities": "{}", "tags": [], "tag_sources": "{}",
"importance": 0.5, "created_at": None, "original_date": None,
"language": None, "metadata": "{}", "has_embedding": True,
}
row_after = dict(row_before)
del row_after["has_embedding"]

cursor, manager = _cursor()
cursor.fetchone.side_effect = [row_before, row_after]
with patch.object(attribution, "get_db_cursor", manager):
result = attribution.regenerate_embedding(mem_id, [0.1, 0.2], force=True)

assert result is not None


# --- MCP tool ---

def test_mcp_tool_schema_registered():
from src import main

tools = asyncio.run(main.list_tools())
names = {t.name for t in tools}
assert "memory_regenerate_embedding" in names

schema = next(t for t in tools if t.name == "memory_regenerate_embedding")
assert "memory_id" in schema.inputSchema["properties"]
assert "force" in schema.inputSchema["properties"]


def test_mcp_handler_returns_error_for_invalid_id():
from src import main

result = asyncio.run(main.handle_memory_regenerate_embedding({"memory_id": "not-a-uuid"}))
assert "Invalid memory ID" in result[0].text


def test_mcp_handler_returns_error_for_missing_memory():
from src import main

cursor, manager = _cursor(row=None)
with patch("src.main.get_memory_by_id", return_value=None):
result = asyncio.run(main.handle_memory_regenerate_embedding({
"memory_id": "00000000-0000-0000-0000-000000000001",
}))
assert "not found" in result[0].text


def test_mcp_handler_regenerates_successfully():
from src import main

memory = {"id": "00000000-0000-0000-0000-000000000001", "content": "hello"}
with patch("src.main.get_memory_by_id", return_value=memory), \
patch("src.main.create_embedding", return_value=[0.1, 0.2]), \
patch("src.main.regenerate_embedding") as mock_regen:
result = asyncio.run(main.handle_memory_regenerate_embedding({
"memory_id": "00000000-0000-0000-0000-000000000001",
}))
assert "regenerated" in result[0].text
mock_regen.assert_called_once()