diff --git a/src/api/main.py b/src/api/main.py index 79dc760..4581220 100644 --- a/src/api/main.py +++ b/src/api/main.py @@ -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 @@ -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.""" diff --git a/src/db/attribution.py b/src/db/attribution.py index b13afd2..8795333 100644 --- a/src/db/attribution.py +++ b/src/db/attribution.py @@ -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, diff --git a/src/main.py b/src/main.py index ee38ad2..d8d8a13 100644 --- a/src/main.py +++ b/src/main.py @@ -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, @@ -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"], + }, + ), ] @@ -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)}")] @@ -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: diff --git a/tests/test_embedding_regeneration.py b/tests/test_embedding_regeneration.py new file mode 100644 index 0000000..eeec51e --- /dev/null +++ b/tests/test_embedding_regeneration.py @@ -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()