diff --git a/generate_openapi.py b/generate_openapi.py index ecd7ea73..b77ae8f7 100644 --- a/generate_openapi.py +++ b/generate_openapi.py @@ -11,7 +11,7 @@ from pathlib import Path # Add parent directory to path to import memory module sys.path.insert(0, str(Path(__file__).parent)) -from memory.web.server import app +from memora.web.server import app def generate_openapi_spec(output_path: str = "openapi.json"): """Generate OpenAPI spec and save to file.""" diff --git a/memora/__init__.py b/memora/__init__.py index 3270c3b5..191fd1b0 100644 --- a/memora/__init__.py +++ b/memora/__init__.py @@ -16,6 +16,7 @@ from .search_trace import ( SearchPhaseMetrics, ) from .search_tracer import SearchTracer +from .embeddings import Embeddings, SentenceTransformersEmbeddings __all__ = [ "TemporalSemanticMemory", @@ -29,5 +30,7 @@ __all__ = [ "PruningDecision", "SearchSummary", "SearchPhaseMetrics", + "Embeddings", + "SentenceTransformersEmbeddings", ] __version__ = "0.1.0" diff --git a/memora/operations/embedding_operations.py b/memora/operations/embedding_operations.py index d20befba..051bc060 100644 --- a/memora/operations/embedding_operations.py +++ b/memora/operations/embedding_operations.py @@ -5,82 +5,48 @@ Embedding generation operations for memory units. import asyncio import logging from typing import List -from concurrent.futures import ProcessPoolExecutor logger = logging.getLogger(__name__) -# Global process pool for parallel embedding generation -_PROCESS_POOL = None - - -def _get_worker_model(): - """Get or load the embedding model in worker process.""" - from sentence_transformers import SentenceTransformer - global _worker_model - if '_worker_model' not in globals(): - globals()['_worker_model'] = SentenceTransformer("BAAI/bge-small-en-v1.5") - return globals()['_worker_model'] - - -def _encode_batch_worker(texts: List[str]) -> List[List[float]]: - """ - Worker function for process pool - encodes texts to embeddings. - - This function runs in a separate process and loads its own model. - """ - model = _get_worker_model() - embeddings = model.encode(texts, convert_to_numpy=True, show_progress_bar=False) - return [emb.tolist() for emb in embeddings] - - -def _get_process_pool(): - """Get or create the global process pool.""" - global _PROCESS_POOL - if _PROCESS_POOL is None: - # Use 4 worker processes for true parallelism - _PROCESS_POOL = ProcessPoolExecutor(max_workers=4) - return _PROCESS_POOL - class EmbeddingOperationsMixin: """Mixin class for embedding operations.""" def _generate_embedding(self, text: str) -> List[float]: """ - Generate embedding for text using local SentenceTransformer model. + Generate embedding for text using the configured embeddings backend. Args: text: Text to embed Returns: - 384-dimensional embedding vector (bge-small-en-v1.5) + Embedding vector (dimension depends on embeddings backend) """ try: - embedding = self.embedding_model.encode(text, convert_to_numpy=True, show_progress_bar=False) - return embedding.tolist() + embeddings = self.embeddings.encode([text]) + return embeddings[0] except Exception as e: raise Exception(f"Failed to generate embedding: {str(e)}") async def _generate_embeddings_batch(self, texts: List[str]) -> List[List[float]]: """ - Generate embeddings for multiple texts using local model in parallel. + Generate embeddings for multiple texts using the configured embeddings backend. - Uses a ProcessPoolExecutor to achieve TRUE parallelism for CPU-bound - embedding generation. Each worker process loads its own model copy. + Runs the embedding generation in a thread pool to avoid blocking the event loop + for CPU-bound operations. Args: texts: List of texts to embed Returns: - List of 384-dimensional embeddings in same order as input texts + List of embeddings in same order as input texts """ try: - # Run in process pool for true parallelism + # Run embeddings in thread pool to avoid blocking event loop loop = asyncio.get_event_loop() - pool = _get_process_pool() embeddings = await loop.run_in_executor( - pool, - _encode_batch_worker, + None, # Use default thread pool + self.embeddings.encode, texts ) return embeddings diff --git a/memora/temporal_semantic_memory.py b/memora/temporal_semantic_memory.py index e4e447f0..ac6acf4a 100644 --- a/memora/temporal_semantic_memory.py +++ b/memora/temporal_semantic_memory.py @@ -12,11 +12,10 @@ import os from datetime import datetime, timedelta, timezone from typing import Any, Dict, List, Optional, Tuple import asyncpg -from sentence_transformers import SentenceTransformer from dotenv import load_dotenv import asyncio +from .embeddings import Embeddings, SentenceTransformersEmbeddings import time -from concurrent.futures import ProcessPoolExecutor import numpy as np import uuid import logging @@ -38,44 +37,6 @@ def utcnow(): # Logger for memory system logger = logging.getLogger(__name__) -# Global process pool for parallel embedding generation -# Each process loads its own copy of the embedding model -# This provides TRUE parallelism for CPU-bound embedding operations -_PROCESS_POOL = None -_EMBEDDING_MODEL_NAME = "BAAI/bge-small-en-v1.5" - -# Process-local model cache (one per worker process) -_worker_model = None - - -def _get_worker_model(): - """Get or load the embedding model in worker process.""" - global _worker_model - if _worker_model is None: - _worker_model = SentenceTransformer(_EMBEDDING_MODEL_NAME) - return _worker_model - - -def _encode_batch_worker(texts: List[str]) -> List[List[float]]: - """ - Worker function for process pool - encodes texts to embeddings. - - This function runs in a separate process and loads its own model. - """ - model = _get_worker_model() - embeddings = model.encode(texts, convert_to_numpy=True, show_progress_bar=False) - return [emb.tolist() for emb in embeddings] - - -def _get_process_pool(): - """Get or create the global process pool.""" - global _PROCESS_POOL - if _PROCESS_POOL is None: - # Use 4 worker processes for true parallelism - # Adjust based on your CPU cores (each process loads ~500MB model) - _PROCESS_POOL = ProcessPoolExecutor(max_workers=4) - return _PROCESS_POOL - class TemporalSemanticMemory( EmbeddingOperationsMixin, @@ -92,14 +53,16 @@ class TemporalSemanticMemory( def __init__( self, db_url: Optional[str] = None, - embedding_model: str = "BAAI/bge-small-en-v1.5", + embeddings: Optional[Embeddings] = None, + embedding_model: Optional[str] = None, ): """ Initialize the temporal + semantic memory system. Args: db_url: PostgreSQL connection URL (postgresql://user:pass@host:port/dbname) - embedding_model: Name of the SentenceTransformer model to use + embeddings: Embeddings implementation to use. If not provided, uses SentenceTransformersEmbeddings + embedding_model: (Deprecated) Name of the SentenceTransformer model to use. Use embeddings parameter instead. """ load_dotenv() @@ -118,10 +81,13 @@ class TemporalSemanticMemory( # Initialize entity resolver (will be created in initialize()) self.entity_resolver = None - # Initialize local embedding model (384 dimensions) - logger.info(f"Loading embedding model: {embedding_model}...") - self.embedding_model = SentenceTransformer(embedding_model) - logger.info(f"Model loaded (embedding dim: {self.embedding_model.get_sentence_embedding_dimension()})") + # Initialize embeddings + if embeddings is not None: + self.embeddings = embeddings + else: + # Default to SentenceTransformersEmbeddings + model_name = embedding_model or "BAAI/bge-small-en-v1.5" + self.embeddings = SentenceTransformersEmbeddings(model_name) # Background queue for access count updates (to avoid blocking searches) self._access_count_queue = asyncio.Queue() diff --git a/memora/web/__init__.py b/memora/web/__init__.py index 7e336e72..924cf08d 100644 --- a/memora/web/__init__.py +++ b/memora/web/__init__.py @@ -3,6 +3,6 @@ Web interface for memory system. Provides FastAPI app and visualization interface. """ -from .server import app, memory +from .server import app, create_app -__all__ = ["app", "memory"] +__all__ = ["app", "create_app"] diff --git a/memora/web/server.py b/memora/web/server.py index c6d036a7..4c012a52 100644 --- a/memora/web/server.py +++ b/memora/web/server.py @@ -18,16 +18,30 @@ from datetime import datetime # Import from parent memora package from memora import TemporalSemanticMemory +from memora.embeddings import Embeddings import logging load_dotenv() logging.basicConfig(level=logging.INFO) -app = FastAPI( - title="Agent Memory API", - version="1.0.0", - description=""" + +def create_app(embeddings: Optional[Embeddings] = None, db_url: Optional[str] = None) -> FastAPI: + """ + Create and configure the FastAPI application. + + Args: + embeddings: Optional custom embeddings implementation. If not provided, + uses default SentenceTransformersEmbeddings. + db_url: Optional database URL. If not provided, uses DATABASE_URL env var. + + Returns: + Configured FastAPI application + """ + app = FastAPI( + title="Agent Memory API", + version="1.0.0", + description=""" A temporal-semantic memory system for AI agents that stores, retrieves, and reasons over memories. ## Features @@ -46,18 +60,42 @@ The system uses: - **Semantic Links**: Connect semantically similar memories - **Entity Links**: Connect memories that mention the same entities - **Spreading Activation**: Intelligent traversal for memory retrieval - """, - contact={ - "name": "Memory System", - }, - license_info={ - "name": "Apache 2.0", - "url": "https://www.apache.org/licenses/LICENSE-2.0.html", - } -) + """, + contact={ + "name": "Memory System", + }, + license_info={ + "name": "Apache 2.0", + "url": "https://www.apache.org/licenses/LICENSE-2.0.html", + } + ) -# Mount static files -app.mount("/static", StaticFiles(directory=str(Path(__file__).parent / "static")), name="static") + # Mount static files + app.mount("/static", StaticFiles(directory=str(Path(__file__).parent / "static")), name="static") + + # Initialize memory system with custom embeddings if provided + memory = TemporalSemanticMemory(db_url=db_url, embeddings=embeddings) + + @app.on_event("startup") + async def startup_event(): + """Initialize memory system on startup.""" + await memory.initialize() + logging.info("Memory system initialized") + + @app.on_event("shutdown") + async def shutdown_event(): + """Cleanup memory system on shutdown.""" + await memory.close() + logging.info("Memory system closed") + + # Store memory instance on app for route handlers to access + app.state.memory = memory + + return app + + +# Create default app instance with default embeddings +app = create_app() class SearchRequest(BaseModel): @@ -239,20 +277,6 @@ class GraphDataResponse(BaseModel): } -memory = TemporalSemanticMemory() - -@app.on_event("startup") -async def startup_event(): - """Initialize memory system on startup.""" - await memory.initialize() - logging.info("Memory system initialized") - -@app.on_event("shutdown") -async def shutdown_event(): - """Cleanup memory system on shutdown.""" - await memory.close() - logging.info("Memory system closed") - @app.get("/", include_in_schema=False) async def index(): """Serve the visualization page.""" @@ -272,7 +296,7 @@ async def api_graph( ): """Get graph data from database, optionally filtered by agent_id and fact_type.""" try: - data = await memory.get_graph_data(agent_id, fact_type) + data = await app.state.memory.get_graph_data(agent_id, fact_type) return data except Exception as e: import traceback @@ -292,7 +316,7 @@ async def api_search(request: SearchRequest): """Run a search and return results with trace.""" try: # Run search with tracing - results, trace = await memory.search_async( + results, trace = await app.state.memory.search_async( agent_id=request.agent_id, query=request.query, thinking_budget=request.thinking_budget, @@ -326,7 +350,7 @@ async def api_world_search(request: SearchRequest): """Search only world facts (general knowledge about the world).""" try: # Run search with fact_type filter for 'world' - results, trace = await memory.search_async( + results, trace = await app.state.memory.search_async( agent_id=request.agent_id, query=request.query, thinking_budget=request.thinking_budget, @@ -361,7 +385,7 @@ async def api_agent_search(request: SearchRequest): """Search only agent facts (facts about what the agent did).""" try: # Run search with fact_type filter for 'agent' - results, trace = await memory.search_async( + results, trace = await app.state.memory.search_async( agent_id=request.agent_id, query=request.query, thinking_budget=request.thinking_budget, @@ -396,7 +420,7 @@ async def api_opinion_search(request: SearchRequest): """Search only opinion facts (agent's formed opinions and perspectives).""" try: # Run search with fact_type filter for 'opinion' - results, trace = await memory.search_async( + results, trace = await app.state.memory.search_async( agent_id=request.agent_id, query=request.query, thinking_budget=request.thinking_budget, @@ -440,7 +464,7 @@ This endpoint: async def api_think(request: ThinkRequest): try: # Use the memory system's think_async method - result = await memory.think_async( + result = await app.state.memory.think_async( agent_id=request.agent_id, query=request.query, thinking_budget=request.thinking_budget, @@ -470,7 +494,7 @@ async def api_think(request: ThinkRequest): async def api_agents(): """Get list of available agents from database.""" try: - agent_list = await memory.list_agents() + agent_list = await app.state.memory.list_agents() return AgentsResponse(agents=agent_list) except Exception as e: import traceback @@ -523,7 +547,7 @@ async def api_batch_put(request: BatchPutRequest): contents.append(content_dict) # Call put_batch_async - result = await memory.put_batch_async( + result = await app.state.memory.put_batch_async( agent_id=request.agent_id, contents=contents, document_id=request.document_id,