From 4c6ac898089d00503d294b75a274e7babd78a547 Mon Sep 17 00:00:00 2001 From: Jeroen Schweitzer Date: Sat, 13 Dec 2025 17:27:51 +0100 Subject: [PATCH] feat: add Phase F.1 memory infrastructure MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Add multi-tenancy support and memory storage infrastructure: - Add ContextVar-based request context (src/core/context.py) - Async-safe user/conversation tracking via contextvars - RequestContext manager for clean setup/teardown - get_user(), get_conversation_id() helpers - Add multi-tenancy utilities (src/core/multi_tenancy.py) - User ID sanitization for collection/key names - get_memory_collection_name(), get_session_key() helpers - Add Ollama embedding client (src/core/embeddings.py) - nomic-embed-text model (768 dimensions) - embed(), embed_batch(), health_check() methods - Add Qdrant client wrapper (src/core/qdrant.py) - Per-user collection pattern: memories_{user} - upsert_memory(), search_memories(), delete_memory() - Type-based filtering support - Add Redis memory cache (src/core/memory_cache.py) - Session context with 24h TTL - Recent entities tracking - Separate from benchmarks (db=2) - Update config with memory settings - QDRANT_HOST, QDRANT_PORT, QDRANT_EMBEDDING_DIM - OLLAMA_EMBEDDING_MODEL - REDIS_MEMORY_DB, REDIS_MEMORY_TTL_HOURS - Add user field to ResponseRequest (OpenAI standard) - Set context in router, reset in finally block - Update librarian client to use get_user() (12 methods) All 333 unit tests pass. 🤖 Generated with [Claude Code](https://claude.com/claude-code) Co-Authored-By: Claude Opus 4.5 --- AGENTS.md | 6 + requirements.txt | 4 + src/agents/librarian/client.py | 59 +++-- src/core/config.py | 42 +++- src/core/context.py | 111 ++++++++ src/core/embeddings.py | 269 ++++++++++++++++++++ src/core/memory_cache.py | 390 ++++++++++++++++++++++++++++ src/core/multi_tenancy.py | 147 +++++++++++ src/core/qdrant.py | 446 +++++++++++++++++++++++++++++++++ src/responses/router.py | 11 + src/responses/schemas.py | 4 + 11 files changed, 1465 insertions(+), 24 deletions(-) create mode 100644 src/core/context.py create mode 100644 src/core/embeddings.py create mode 100644 src/core/memory_cache.py create mode 100644 src/core/multi_tenancy.py create mode 100644 src/core/qdrant.py diff --git a/AGENTS.md b/AGENTS.md index abdfa78..252b34c 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -15,6 +15,12 @@ This document contains instructions and documentation references for AI assistan * **Act:** Execute the changes in small, atomic steps. * **Reflect:** After coding, verify your work. Did you break existing tests? Did you add new tests? +### 🌐 Internal Service Access +* **git.schweitz.net**: Access via `http://localhost:3002` (direct Gitea) to bypass Authentik SSO + * Example: `curl http://localhost:3002/jpmschweitzer/library-desk/raw/branch/main/README.md` + * Public repos are readable without authentication + * Related repos: `library-desk`, `scheduler` + ### 🛡️ Git Discipline * **NEVER commit to `main` or `master` directly.** Always create a feature branch: `feature/your-feature-name` or `fix/issue-description`. * **Commit Messages:** Use the [Conventional Commits](https://www.conventionalcommits.org/) format. diff --git a/requirements.txt b/requirements.txt index 034f82e..c66f52e 100644 --- a/requirements.txt +++ b/requirements.txt @@ -41,6 +41,10 @@ starlette>=0.45,<0.46 # hiredis: C parser for better performance redis[hiredis]>=5.2,<6.0 +# Qdrant vector database client for memory storage +# Latest: 1.12.1 (Dec 2025) - No known CVEs +qdrant-client>=1.12,<2.0 + # Structured logging for observability # Latest: 24.4.0 (Aug 22, 2024) - No known CVEs structlog>=24.1,<25.0 diff --git a/src/agents/librarian/client.py b/src/agents/librarian/client.py index 1452174..7be3673 100644 --- a/src/agents/librarian/client.py +++ b/src/agents/librarian/client.py @@ -13,6 +13,7 @@ import httpx from pydantic import BaseModel, Field from src.core.config import config +from src.core.context import get_user from src.core.logging_config import get_logger logger = get_logger(__name__) @@ -179,7 +180,7 @@ class LibraryDeskClient: async def hybrid_search( self, query: str, - user: str = "jpmschweitzer", + user: str | None = None, vector_limit: int = 10, graph_limit: int = 10, web_limit: int = 5, @@ -191,7 +192,7 @@ class LibraryDeskClient: Args: query: Search query - user: User identifier for multi-tenancy + user: User identifier for multi-tenancy (defaults to request context) vector_limit: Max results from vector search graph_limit: Max results from graph search web_limit: Max results from web search @@ -201,6 +202,7 @@ class LibraryDeskClient: Returns: HybridRAGResponse with ranked results and context """ + user = user or get_user() client = self._ensure_client() payload = { @@ -255,7 +257,7 @@ class LibraryDeskClient: async def search_wiki( self, query: str, - user: str = "jpmschweitzer", + user: str | None = None, limit: int = 20, ) -> list[WikiSearchResult]: """ @@ -263,12 +265,13 @@ class LibraryDeskClient: Args: query: Search query - user: User identifier + user: User identifier (defaults to request context) limit: Maximum results Returns: List of matching wiki pages """ + user = user or get_user() client = self._ensure_client() logger.debug("library_desk_wiki_search", query=query, user=user) @@ -285,18 +288,19 @@ class LibraryDeskClient: async def get_wiki_page( self, page_id: int, - user: str = "jpmschweitzer", + user: str | None = None, ) -> WikiPage: """ Get a wiki page by ID. Args: page_id: Page ID - user: User identifier + user: User identifier (defaults to request context) Returns: WikiPage with full content """ + user = user or get_user() client = self._ensure_client() response = await client.get( @@ -309,7 +313,7 @@ class LibraryDeskClient: async def list_wiki_pages( self, - user: str = "jpmschweitzer", + user: str | None = None, tag: Optional[str] = None, limit: int = 50, ) -> list[WikiPage]: @@ -317,13 +321,14 @@ class LibraryDeskClient: List wiki pages, optionally filtered by tag. Args: - user: User identifier + user: User identifier (defaults to request context) tag: Optional tag (dossier) to filter by limit: Maximum pages to return Returns: List of wiki pages """ + user = user or get_user() client = self._ensure_client() params: dict[str, Any] = {"user": user, "limit": limit} @@ -341,7 +346,7 @@ class LibraryDeskClient: title: str, path: str, content: str, - user: str = "jpmschweitzer", + user: str | None = None, description: str = "", tags: Optional[list[str]] = None, ) -> WikiPage: @@ -352,13 +357,14 @@ class LibraryDeskClient: title: Page title path: Page path (e.g., "/projects/my-project") content: Markdown content - user: User identifier + user: User identifier (defaults to request context) description: Short description tags: List of tags (dossiers) Returns: Created WikiPage """ + user = user or get_user() client = self._ensure_client() payload = { @@ -380,7 +386,7 @@ class LibraryDeskClient: async def update_wiki_page( self, page_id: int, - user: str = "jpmschweitzer", + user: str | None = None, content: Optional[str] = None, title: Optional[str] = None, tags: Optional[list[str]] = None, @@ -394,7 +400,7 @@ class LibraryDeskClient: Args: page_id: ID of the page to update - user: User identifier + user: User identifier (defaults to request context) content: New content (optional) title: New title (optional) tags: New tags list (optional) @@ -403,6 +409,7 @@ class LibraryDeskClient: Returns: Updated WikiPage """ + user = user or get_user() client = self._ensure_client() # Build update payload with only provided fields @@ -435,7 +442,7 @@ class LibraryDeskClient: self, topic: str, tags: list[str], - user: str = "jpmschweitzer", + user: str | None = None, path: Optional[str] = None, include_web_research: bool = True, include_wiki_search: bool = True, @@ -460,6 +467,7 @@ class LibraryDeskClient: Returns: SmartCreateResponse with page and research metadata """ + user = user or get_user() client = self._ensure_client() payload: dict[str, Any] = { @@ -499,17 +507,18 @@ class LibraryDeskClient: async def list_dossiers( self, - user: str = "jpmschweitzer", + user: str | None = None, ) -> list[Dossier]: """ List all dossiers (tag collections) for a user. Args: - user: User identifier + user: User identifier (defaults to request context) Returns: List of dossiers with page counts """ + user = user or get_user() client = self._ensure_client() response = await client.get( @@ -528,7 +537,7 @@ class LibraryDeskClient: async def semantic_search( self, query: str, - user: str = "jpmschweitzer", + user: str | None = None, limit: int = 10, score_threshold: float = 0.5, ) -> list[VectorSearchResult]: @@ -537,13 +546,14 @@ class LibraryDeskClient: Args: query: Natural language query - user: User identifier + user: User identifier (defaults to request context) limit: Maximum results score_threshold: Minimum similarity score Returns: List of matching document chunks with scores """ + user = user or get_user() client = self._ensure_client() payload = { @@ -568,7 +578,7 @@ class LibraryDeskClient: async def query_graph( self, cypher_query: str, - user: str = "jpmschweitzer", + user: str | None = None, parameters: Optional[dict[str, Any]] = None, ) -> list[dict[str, Any]]: """ @@ -578,12 +588,13 @@ class LibraryDeskClient: Args: cypher_query: Cypher query string - user: User identifier + user: User identifier (defaults to request context) parameters: Query parameters Returns: List of result records """ + user = user or get_user() client = self._ensure_client() payload = { @@ -601,7 +612,7 @@ class LibraryDeskClient: async def list_graph_nodes( self, - user: str = "jpmschweitzer", + user: str | None = None, node_type: Optional[str] = None, limit: int = 100, ) -> list[GraphNode]: @@ -609,13 +620,14 @@ class LibraryDeskClient: List nodes in the knowledge graph. Args: - user: User identifier + user: User identifier (defaults to request context) node_type: Optional filter by type (Document, Person, Concept, etc.) limit: Maximum nodes Returns: List of graph nodes """ + user = user or get_user() client = self._ensure_client() params: dict[str, Any] = {"user": user, "limit": limit} @@ -631,18 +643,19 @@ class LibraryDeskClient: async def get_graph_node( self, node_id: str, - user: str = "jpmschweitzer", + user: str | None = None, ) -> dict[str, Any]: """ Get detailed information about a graph node. Args: node_id: Node ID - user: User identifier + user: User identifier (defaults to request context) Returns: Node with relationships and connected nodes """ + user = user or get_user() client = self._ensure_client() response = await client.get( diff --git a/src/core/config.py b/src/core/config.py index 0321290..86618d2 100644 --- a/src/core/config.py +++ b/src/core/config.py @@ -124,6 +124,36 @@ class Config(BaseSettings): description="Library-Desk request timeout in seconds" ) + # Qdrant Configuration (Memory vector storage) + QDRANT_HOST: str = Field( + default="localhost", + description="Qdrant server host" + ) + QDRANT_PORT: int = Field( + default=6333, + description="Qdrant server port" + ) + QDRANT_EMBEDDING_DIM: int = Field( + default=768, + description="Embedding dimension (768 for nomic-embed-text)" + ) + + # Ollama Embedding Configuration + OLLAMA_EMBEDDING_MODEL: str = Field( + default="nomic-embed-text", + description="Ollama model for embeddings" + ) + + # Redis Memory Database (separate from benchmarks) + REDIS_MEMORY_DB: int = Field( + default=2, + description="Redis database number for memory cache" + ) + REDIS_MEMORY_TTL_HOURS: int = Field( + default=24, + description="TTL for session context in hours" + ) + # Logging LOG_LEVEL: str = Field(default="INFO", description="Logging level") ENABLE_BENCHMARKS: bool = Field(default=True, description="Enable performance benchmarking") @@ -139,9 +169,19 @@ class Config(BaseSettings): @property def redis_url(self) -> str: - """Construct Redis connection URL.""" + """Construct Redis connection URL for benchmarks.""" return f"redis://{self.REDIS_HOST}:{self.REDIS_PORT}/{self.REDIS_DB}" + @property + def redis_memory_url(self) -> str: + """Construct Redis connection URL for memory cache.""" + return f"redis://{self.REDIS_HOST}:{self.REDIS_PORT}/{self.REDIS_MEMORY_DB}" + + @property + def qdrant_url(self) -> str: + """Construct Qdrant server URL.""" + return f"http://{self.QDRANT_HOST}:{self.QDRANT_PORT}" + @property def log_format(self) -> str: """ diff --git a/src/core/context.py b/src/core/context.py new file mode 100644 index 0000000..2b5db46 --- /dev/null +++ b/src/core/context.py @@ -0,0 +1,111 @@ +""" +Request context using ContextVar for async-safe user/conversation tracking. + +ContextVar provides task-local storage that automatically propagates through +async calls, eliminating the need to thread user identity through every function. + +Usage: + # At request entry (router): + token = current_user.set(request.user or "jpmschweitzer") + try: + await service.process(request) + finally: + current_user.reset(token) + + # Anywhere in the codebase: + from src.core.context import get_user + user = get_user() # Returns current request's user +""" +from contextvars import ContextVar + +# Default user for single-user homelab setup +DEFAULT_USER = "jpmschweitzer" + +# Request-scoped context variables (async-safe, isolated per request) +current_user: ContextVar[str] = ContextVar("current_user", default=DEFAULT_USER) +current_conversation: ContextVar[str | None] = ContextVar( + "current_conversation", default=None +) + + +def get_user() -> str: + """ + Get current user from request context. + + Returns: + User identifier for the current request. + Falls back to DEFAULT_USER if not set. + + Example: + user = get_user() # "jpmschweitzer" or whatever was set in router + """ + return current_user.get() + + +def get_conversation_id() -> str | None: + """ + Get current conversation ID from request context. + + Returns: + Conversation ID if set, None otherwise. + + Example: + conv_id = get_conversation_id() # "conv_abc123" or None + """ + return current_conversation.get() + + +class RequestContext: + """ + Context manager for setting request-scoped context. + + Provides a cleaner alternative to manual token management. + + Usage: + async with RequestContext(user="alice", conversation_id="conv_123"): + # All code here sees user="alice" + result = await some_service.process() + """ + + def __init__( + self, + user: str | None = None, + conversation_id: str | None = None, + ): + """ + Initialize request context. + + Args: + user: User identifier (defaults to DEFAULT_USER if None) + conversation_id: Conversation ID (optional) + """ + self.user = user or DEFAULT_USER + self.conversation_id = conversation_id + self._user_token = None + self._conv_token = None + + async def __aenter__(self) -> "RequestContext": + """Set context variables on entry.""" + self._user_token = current_user.set(self.user) + self._conv_token = current_conversation.set(self.conversation_id) + return self + + async def __aexit__(self, exc_type, exc_val, exc_tb) -> None: + """Reset context variables on exit.""" + if self._user_token is not None: + current_user.reset(self._user_token) + if self._conv_token is not None: + current_conversation.reset(self._conv_token) + + def __enter__(self) -> "RequestContext": + """Sync context manager entry (for non-async code).""" + self._user_token = current_user.set(self.user) + self._conv_token = current_conversation.set(self.conversation_id) + return self + + def __exit__(self, exc_type, exc_val, exc_tb) -> None: + """Sync context manager exit.""" + if self._user_token is not None: + current_user.reset(self._user_token) + if self._conv_token is not None: + current_conversation.reset(self._conv_token) diff --git a/src/core/embeddings.py b/src/core/embeddings.py new file mode 100644 index 0000000..4957ea4 --- /dev/null +++ b/src/core/embeddings.py @@ -0,0 +1,269 @@ +""" +Ollama client for embeddings generation. + +Provides async embedding operations via Ollama API: +- Text embedding generation +- Batch embedding support +- Health checks + +Adapted from library-desk patterns. +""" +from typing import Optional + +import httpx + +from .config import config +from .logging_config import get_logger + +logger = get_logger(__name__) + + +class OllamaEmbeddingClient: + """ + Ollama API client for embeddings. + + Uses the Ollama embeddings endpoint to generate vector representations + of text using the nomic-embed-text model (768 dimensions). + + Usage: + client = OllamaEmbeddingClient() + embedding = await client.embed("Hello world") + await client.close() + + Or with context manager: + async with OllamaEmbeddingClient() as client: + embedding = await client.embed("Hello world") + """ + + def __init__( + self, + base_url: str | None = None, + model: str | None = None, + timeout: float = 120.0, + ): + """ + Initialize Ollama embedding client. + + Args: + base_url: Ollama server URL (defaults to config.OLLAMA_HOST) + model: Embedding model name (defaults to config.OLLAMA_EMBEDDING_MODEL) + timeout: Request timeout in seconds (embeddings can be slow) + """ + self.base_url = (base_url or str(config.OLLAMA_HOST)).rstrip("/") + self.model = model or config.OLLAMA_EMBEDDING_MODEL + self.embeddings_url = f"{self.base_url}/api/embeddings" + self.tags_url = f"{self.base_url}/api/tags" + self._client: httpx.AsyncClient | None = None + self._timeout = timeout + + logger.info( + "ollama_embedding_client_initialized", + base_url=self.base_url, + model=self.model, + ) + + async def _get_client(self) -> httpx.AsyncClient: + """Get or create HTTP client.""" + if self._client is None: + self._client = httpx.AsyncClient(timeout=self._timeout) + return self._client + + async def __aenter__(self) -> "OllamaEmbeddingClient": + """Async context manager entry.""" + await self._get_client() + return self + + async def __aexit__(self, exc_type, exc_val, exc_tb) -> None: + """Async context manager exit.""" + await self.close() + + async def close(self) -> None: + """Close HTTP client.""" + if self._client is not None: + await self._client.aclose() + self._client = None + + async def embed(self, text: str) -> list[float] | None: + """ + Generate embedding for single text. + + Args: + text: Text to embed + + Returns: + Embedding vector (768-dimensional for nomic-embed-text) or None on failure + + Example: + >>> embedding = await client.embed("Hello world") + >>> len(embedding) + 768 + """ + try: + client = await self._get_client() + + payload = { + "model": self.model, + "prompt": text, + } + + response = await client.post(self.embeddings_url, json=payload) + response.raise_for_status() + data = response.json() + + embedding = data.get("embedding") + if not embedding: + logger.error("ollama_embed_no_embedding", response_data=data) + return None + + return embedding + + except httpx.HTTPStatusError as e: + logger.error( + "ollama_embed_http_error", + status_code=e.response.status_code, + detail=e.response.text, + ) + return None + except Exception as e: + logger.error("ollama_embed_failed", error=str(e), exc_info=True) + return None + + async def embed_batch( + self, + texts: list[str], + show_progress: bool = False, + ) -> list[list[float] | None]: + """ + Generate embeddings for multiple texts. + + Note: Ollama doesn't support native batch embeddings, so this + sequentially calls embed() for each text. + + Args: + texts: List of texts to embed + show_progress: Log progress for large batches + + Returns: + List of embedding vectors (same order as input) + None entries for texts that failed to embed + + Example: + >>> texts = ["Hello", "World", "Test"] + >>> embeddings = await client.embed_batch(texts) + >>> len(embeddings) + 3 + """ + embeddings = [] + + for i, text in enumerate(texts): + if show_progress and i % 10 == 0: + logger.info( + "ollama_embed_batch_progress", + current=i, + total=len(texts), + ) + + embedding = await self.embed(text) + embeddings.append(embedding) + + if show_progress: + logger.info( + "ollama_embed_batch_complete", + successful=sum(1 for e in embeddings if e is not None), + total=len(texts), + ) + + return embeddings + + async def embed_batch_filtered( + self, + texts: list[str], + show_progress: bool = False, + ) -> list[list[float]]: + """ + Generate embeddings for multiple texts, filtering out failures. + + Args: + texts: List of texts to embed + show_progress: Log progress for large batches + + Returns: + List of successful embedding vectors (may be shorter than input) + + Example: + >>> embeddings = await client.embed_batch_filtered(texts) + >>> all(e is not None for e in embeddings) + True + """ + all_embeddings = await self.embed_batch(texts, show_progress) + return [e for e in all_embeddings if e is not None] + + async def get_embedding_dimension(self) -> int | None: + """ + Get embedding dimension for current model. + + Returns: + Embedding dimension (e.g., 768 for nomic-embed-text) or None on failure + + Example: + >>> dim = await client.get_embedding_dimension() + >>> dim + 768 + """ + test_embedding = await self.embed("test") + if test_embedding: + return len(test_embedding) + return None + + async def health_check(self) -> bool: + """ + Check if Ollama server is reachable and model is available. + + Returns: + True if healthy, False otherwise + """ + try: + client = await self._get_client() + response = await client.get(self.tags_url, timeout=5.0) + response.raise_for_status() + data = response.json() + models = data.get("models", []) + + # Check if our embedding model is available + model_found = False + for m in models: + name = m.get("name", "") + if name == self.model or name.startswith(f"{self.model}:"): + model_found = True + break + + if not model_found: + logger.warning( + "ollama_embedding_model_not_found", + model=self.model, + available=[m.get("name") for m in models], + ) + return False + + return True + + except Exception as e: + logger.error("ollama_embedding_health_check_failed", error=str(e)) + return False + + +# Global client instance (lazy initialization) +_embedding_client: OllamaEmbeddingClient | None = None + + +def get_embedding_client() -> OllamaEmbeddingClient: + """ + Get global embedding client instance. + + Returns: + OllamaEmbeddingClient instance + """ + global _embedding_client + if _embedding_client is None: + _embedding_client = OllamaEmbeddingClient() + return _embedding_client diff --git a/src/core/memory_cache.py b/src/core/memory_cache.py new file mode 100644 index 0000000..1e556bb --- /dev/null +++ b/src/core/memory_cache.py @@ -0,0 +1,390 @@ +""" +Redis-backed memory cache for session context. + +Provides short-term memory storage with TTL: +- Session context (24h TTL) +- Recent entities mentioned in conversation +- User-scoped with conversation isolation + +Uses Redis DB 2 (separate from benchmarks in DB 1). +""" +import json +from typing import Any + +import redis.asyncio as redis + +from .config import config +from .logging_config import get_logger +from .multi_tenancy import get_session_key, get_entities_key + +logger = get_logger(__name__) + + +class MemoryCache: + """ + Redis-backed cache for session memory. + + Stores ephemeral context that doesn't need vector search: + - Session context (recent topics, user state) + - Recent entities (people, places, things mentioned) + - Conversation metadata + + All data expires after REDIS_MEMORY_TTL_HOURS (default 24h). + + Usage: + cache = MemoryCache() + await cache.set_session_context( + user="jpmschweitzer", + conversation_id="conv_123", + context={"topic": "docker", "mood": "curious"} + ) + context = await cache.get_session_context("jpmschweitzer", "conv_123") + """ + + def __init__( + self, + redis_url: str | None = None, + ttl_hours: int | None = None, + ): + """ + Initialize memory cache. + + Args: + redis_url: Redis connection URL (defaults to config.redis_memory_url) + ttl_hours: TTL for cached data (defaults to config.REDIS_MEMORY_TTL_HOURS) + """ + self._redis_url = redis_url or config.redis_memory_url + self._ttl_seconds = (ttl_hours or config.REDIS_MEMORY_TTL_HOURS) * 3600 + self._client: redis.Redis | None = None + + logger.info( + "memory_cache_initialized", + redis_url=self._redis_url, + ttl_hours=ttl_hours or config.REDIS_MEMORY_TTL_HOURS, + ) + + async def _get_client(self) -> redis.Redis: + """Get or create Redis client.""" + if self._client is None: + self._client = redis.from_url( + self._redis_url, + encoding="utf-8", + decode_responses=True, + socket_timeout=config.REDIS_TIMEOUT, + socket_connect_timeout=config.REDIS_TIMEOUT, + ) + return self._client + + async def close(self) -> None: + """Close Redis connection.""" + if self._client is not None: + await self._client.aclose() + self._client = None + + # ========================================================================= + # Session Context + # ========================================================================= + + async def get_session_context( + self, + user: str, + conversation_id: str, + ) -> dict[str, Any] | None: + """ + Get session context for a conversation. + + Args: + user: User identifier + conversation_id: Conversation identifier + + Returns: + Session context dict or None if not found + + Example: + >>> context = await cache.get_session_context("jpmschweitzer", "conv_123") + >>> context + {"topic": "docker", "mood": "curious", "last_tool": "librarian"} + """ + try: + client = await self._get_client() + key = get_session_key(user, conversation_id) + + data = await client.get(key) + if data is None: + return None + + return json.loads(data) + + except Exception as e: + logger.warning( + "memory_cache_get_session_failed", + user=user, + conversation_id=conversation_id, + error=str(e), + ) + return None + + async def set_session_context( + self, + user: str, + conversation_id: str, + context: dict[str, Any], + ) -> bool: + """ + Set session context for a conversation. + + Args: + user: User identifier + conversation_id: Conversation identifier + context: Context data to store + + Returns: + True if successful, False otherwise + + Example: + >>> await cache.set_session_context( + ... "jpmschweitzer", + ... "conv_123", + ... {"topic": "docker", "mood": "curious"} + ... ) + True + """ + try: + client = await self._get_client() + key = get_session_key(user, conversation_id) + + await client.setex( + key, + self._ttl_seconds, + json.dumps(context), + ) + + logger.debug( + "memory_cache_set_session", + user=user, + conversation_id=conversation_id, + context_keys=list(context.keys()), + ) + return True + + except Exception as e: + logger.warning( + "memory_cache_set_session_failed", + user=user, + conversation_id=conversation_id, + error=str(e), + ) + return False + + async def update_session_context( + self, + user: str, + conversation_id: str, + updates: dict[str, Any], + ) -> bool: + """ + Update session context (merge with existing). + + Args: + user: User identifier + conversation_id: Conversation identifier + updates: Fields to update/add + + Returns: + True if successful, False otherwise + """ + existing = await self.get_session_context(user, conversation_id) or {} + existing.update(updates) + return await self.set_session_context(user, conversation_id, existing) + + async def delete_session_context( + self, + user: str, + conversation_id: str, + ) -> bool: + """ + Delete session context for a conversation. + + Args: + user: User identifier + conversation_id: Conversation identifier + + Returns: + True if deleted, False otherwise + """ + try: + client = await self._get_client() + key = get_session_key(user, conversation_id) + await client.delete(key) + return True + + except Exception as e: + logger.warning( + "memory_cache_delete_session_failed", + user=user, + conversation_id=conversation_id, + error=str(e), + ) + return False + + # ========================================================================= + # Recent Entities + # ========================================================================= + + async def get_recent_entities( + self, + user: str, + conversation_id: str, + ) -> list[str]: + """ + Get recently mentioned entities in a conversation. + + Args: + user: User identifier + conversation_id: Conversation identifier + + Returns: + List of entity names/identifiers + + Example: + >>> entities = await cache.get_recent_entities("jpmschweitzer", "conv_123") + >>> entities + ["Docker", "Kubernetes", "nginx"] + """ + try: + client = await self._get_client() + key = get_entities_key(user, conversation_id) + + # Get all members of the set + entities = await client.smembers(key) + return list(entities) + + except Exception as e: + logger.warning( + "memory_cache_get_entities_failed", + user=user, + conversation_id=conversation_id, + error=str(e), + ) + return [] + + async def add_recent_entities( + self, + user: str, + conversation_id: str, + entities: list[str], + ) -> bool: + """ + Add entities to the recent entities set. + + Args: + user: User identifier + conversation_id: Conversation identifier + entities: Entity names to add + + Returns: + True if successful, False otherwise + + Example: + >>> await cache.add_recent_entities( + ... "jpmschweitzer", + ... "conv_123", + ... ["Docker", "Kubernetes"] + ... ) + True + """ + if not entities: + return True + + try: + client = await self._get_client() + key = get_entities_key(user, conversation_id) + + # Add to set + await client.sadd(key, *entities) + + # Refresh TTL + await client.expire(key, self._ttl_seconds) + + logger.debug( + "memory_cache_add_entities", + user=user, + conversation_id=conversation_id, + entities=entities, + ) + return True + + except Exception as e: + logger.warning( + "memory_cache_add_entities_failed", + user=user, + conversation_id=conversation_id, + error=str(e), + ) + return False + + async def clear_recent_entities( + self, + user: str, + conversation_id: str, + ) -> bool: + """ + Clear all recent entities for a conversation. + + Args: + user: User identifier + conversation_id: Conversation identifier + + Returns: + True if cleared, False otherwise + """ + try: + client = await self._get_client() + key = get_entities_key(user, conversation_id) + await client.delete(key) + return True + + except Exception as e: + logger.warning( + "memory_cache_clear_entities_failed", + user=user, + conversation_id=conversation_id, + error=str(e), + ) + return False + + # ========================================================================= + # Health Check + # ========================================================================= + + async def health_check(self) -> bool: + """ + Check if Redis is reachable. + + Returns: + True if healthy, False otherwise + """ + try: + client = await self._get_client() + await client.ping() + return True + except Exception as e: + logger.error("memory_cache_health_check_failed", error=str(e)) + return False + + +# Global cache instance (lazy initialization) +_memory_cache: MemoryCache | None = None + + +def get_memory_cache() -> MemoryCache: + """ + Get global memory cache instance. + + Returns: + MemoryCache instance + """ + global _memory_cache + if _memory_cache is None: + _memory_cache = MemoryCache() + return _memory_cache diff --git a/src/core/multi_tenancy.py b/src/core/multi_tenancy.py new file mode 100644 index 0000000..cd47b4d --- /dev/null +++ b/src/core/multi_tenancy.py @@ -0,0 +1,147 @@ +""" +Multi-tenancy helpers for Tatlock. + +Provides utilities for user namespace management across: +- Qdrant (collection per user for memories) +- Redis (user-scoped keys for session context) + +Adapted from library-desk patterns. +""" +import re + + +def sanitize_user_id(user_id: str) -> str: + """ + Sanitize user ID for use in collection names, keys, and paths. + + Converts special characters to underscores and ensures alphanumeric safety. + + Args: + user_id: Raw user identifier (email, username, etc.) + + Returns: + Sanitized user ID safe for use in identifiers + + Examples: + >>> sanitize_user_id("john@example.com") + 'john_at_example_com' + >>> sanitize_user_id("user.name") + 'user_name' + >>> sanitize_user_id("User Name") + 'user_name' + """ + sanitized = user_id.lower() + + # Convert @ to _at_ + sanitized = sanitized.replace("@", "_at_") + + # Convert dots to underscores + sanitized = sanitized.replace(".", "_") + + # Replace any non-alphanumeric characters with underscores + sanitized = re.sub(r'[^a-z0-9_]', '_', sanitized) + + # Remove consecutive underscores + sanitized = re.sub(r'_+', '_', sanitized) + + # Remove leading/trailing underscores + sanitized = sanitized.strip('_') + + return sanitized + + +def get_memory_collection_name(user_id: str) -> str: + """ + Get Qdrant collection name for user's memories. + + Pattern: memories_{sanitized_user_id} + + Args: + user_id: User identifier + + Returns: + Qdrant collection name + + Examples: + >>> get_memory_collection_name("jpmschweitzer") + 'memories_jpmschweitzer' + >>> get_memory_collection_name("john@example.com") + 'memories_john_at_example_com' + """ + sanitized = sanitize_user_id(user_id) + return f"memories_{sanitized}" + + +def get_session_key(user_id: str, conversation_id: str) -> str: + """ + Get Redis key for session context. + + Pattern: session:{sanitized_user}:{conversation_id} + + Args: + user_id: User identifier + conversation_id: Conversation identifier + + Returns: + Redis key for session context + + Examples: + >>> get_session_key("jpmschweitzer", "conv_abc123") + 'session:jpmschweitzer:conv_abc123' + """ + sanitized = sanitize_user_id(user_id) + return f"session:{sanitized}:{conversation_id}" + + +def get_entities_key(user_id: str, conversation_id: str) -> str: + """ + Get Redis key for recent entities in a conversation. + + Pattern: entities:{sanitized_user}:{conversation_id} + + Args: + user_id: User identifier + conversation_id: Conversation identifier + + Returns: + Redis key for recent entities + + Examples: + >>> get_entities_key("jpmschweitzer", "conv_abc123") + 'entities:jpmschweitzer:conv_abc123' + """ + sanitized = sanitize_user_id(user_id) + return f"entities:{sanitized}:{conversation_id}" + + +def validate_user_id(user_id: str) -> bool: + """ + Validate that a user ID is acceptable. + + Checks: + - Not empty + - Not too long (max 100 chars) + - Contains some alphanumeric characters + + Args: + user_id: User identifier to validate + + Returns: + True if valid, False otherwise + + Examples: + >>> validate_user_id("jpmschweitzer") + True + >>> validate_user_id("") + False + >>> validate_user_id("a" * 101) + False + """ + if not user_id or len(user_id) > 100: + return False + + # Must contain at least one alphanumeric character + if not re.search(r'[a-zA-Z0-9]', user_id): + return False + + return True diff --git a/src/core/qdrant.py b/src/core/qdrant.py new file mode 100644 index 0000000..fba5741 --- /dev/null +++ b/src/core/qdrant.py @@ -0,0 +1,446 @@ +""" +Qdrant client wrapper for memory vector storage. + +Provides async operations for storing and retrieving memory embeddings: +- Collection management (per-user collections) +- Memory upsert/search/delete +- Filtering by memory type + +Adapted from library-desk patterns. +""" +from typing import Any +from uuid import uuid4 + +from qdrant_client import QdrantClient +from qdrant_client.http import models as qdrant_models + +from .config import config +from .logging_config import get_logger +from .multi_tenancy import get_memory_collection_name + +logger = get_logger(__name__) + + +class MemoryQdrantClient: + """ + Qdrant client wrapper for memory storage. + + Manages per-user collections with the pattern: memories_{user} + Stores memory embeddings with metadata (type, content, timestamps). + + Usage: + client = MemoryQdrantClient() + await client.ensure_collection("jpmschweitzer") + await client.upsert_memory( + user="jpmschweitzer", + memory_id="mem_123", + vector=[0.1, 0.2, ...], + payload={"type": "fact", "content": "User prefers dark mode"} + ) + """ + + def __init__( + self, + url: str | None = None, + embedding_dim: int | None = None, + ): + """ + Initialize Qdrant client. + + Args: + url: Qdrant server URL (defaults to config.qdrant_url) + embedding_dim: Vector dimension (defaults to config.QDRANT_EMBEDDING_DIM) + """ + self.url = url or config.qdrant_url + self.embedding_dim = embedding_dim or config.QDRANT_EMBEDDING_DIM + self._client = QdrantClient(url=self.url) + + logger.info( + "qdrant_client_initialized", + url=self.url, + embedding_dim=self.embedding_dim, + ) + + def close(self) -> None: + """Close Qdrant client.""" + if self._client is not None: + self._client.close() + + async def ensure_collection(self, user: str) -> bool: + """ + Ensure collection exists for user, create if not. + + Args: + user: User identifier + + Returns: + True if collection exists or was created successfully + + Example: + >>> await client.ensure_collection("jpmschweitzer") + True + """ + collection_name = get_memory_collection_name(user) + + try: + # Check if collection exists + collections = self._client.get_collections() + existing = [c.name for c in collections.collections] + + if collection_name in existing: + logger.debug( + "qdrant_collection_exists", + collection=collection_name, + ) + return True + + # Create collection with cosine distance + self._client.create_collection( + collection_name=collection_name, + vectors_config=qdrant_models.VectorParams( + size=self.embedding_dim, + distance=qdrant_models.Distance.COSINE, + ), + ) + + logger.info( + "qdrant_collection_created", + collection=collection_name, + embedding_dim=self.embedding_dim, + ) + return True + + except Exception as e: + logger.error( + "qdrant_ensure_collection_failed", + collection=collection_name, + error=str(e), + ) + return False + + async def upsert_memory( + self, + user: str, + memory_id: str | None, + vector: list[float], + payload: dict[str, Any], + ) -> str | None: + """ + Upsert a memory point. + + Args: + user: User identifier + memory_id: Memory ID (generated if None) + vector: Embedding vector + payload: Memory metadata (should include 'type', 'content', etc.) + + Returns: + Memory ID if successful, None on failure + + Example: + >>> memory_id = await client.upsert_memory( + ... user="jpmschweitzer", + ... memory_id=None, + ... vector=[0.1, 0.2, ...], + ... payload={ + ... "type": "fact", + ... "content": "User prefers dark mode", + ... "created_at": "2024-01-01T00:00:00Z" + ... } + ... ) + """ + collection_name = get_memory_collection_name(user) + memory_id = memory_id or f"mem_{uuid4().hex[:16]}" + + try: + # Ensure collection exists + await self.ensure_collection(user) + + # Create point + point = qdrant_models.PointStruct( + id=memory_id, + vector=vector, + payload=payload, + ) + + # Upsert + self._client.upsert( + collection_name=collection_name, + points=[point], + ) + + logger.debug( + "qdrant_memory_upserted", + collection=collection_name, + memory_id=memory_id, + memory_type=payload.get("type"), + ) + return memory_id + + except Exception as e: + logger.error( + "qdrant_upsert_memory_failed", + collection=collection_name, + memory_id=memory_id, + error=str(e), + ) + return None + + async def search_memories( + self, + user: str, + query_vector: list[float], + limit: int = 10, + memory_type: str | None = None, + score_threshold: float = 0.5, + ) -> list[dict[str, Any]]: + """ + Search memories by vector similarity. + + Args: + user: User identifier + query_vector: Query embedding vector + limit: Maximum results + memory_type: Filter by memory type (e.g., "fact", "preference", "profile") + score_threshold: Minimum similarity score (0-1) + + Returns: + List of matching memories with scores + + Example: + >>> memories = await client.search_memories( + ... user="jpmschweitzer", + ... query_vector=[0.1, 0.2, ...], + ... limit=5, + ... memory_type="fact" + ... ) + >>> memories[0] + {"id": "mem_123", "score": 0.89, "type": "fact", "content": "..."} + """ + collection_name = get_memory_collection_name(user) + + try: + # Build filter if memory_type specified + query_filter = None + if memory_type: + query_filter = qdrant_models.Filter( + must=[ + qdrant_models.FieldCondition( + key="type", + match=qdrant_models.MatchValue(value=memory_type), + ) + ] + ) + + # Search + results = self._client.search( + collection_name=collection_name, + query_vector=query_vector, + limit=limit, + query_filter=query_filter, + score_threshold=score_threshold, + ) + + # Format results + memories = [] + for hit in results: + memory = { + "id": hit.id, + "score": hit.score, + **hit.payload, + } + memories.append(memory) + + logger.debug( + "qdrant_search_memories", + collection=collection_name, + results_count=len(memories), + memory_type=memory_type, + ) + return memories + + except Exception as e: + logger.error( + "qdrant_search_memories_failed", + collection=collection_name, + error=str(e), + ) + return [] + + async def get_memory(self, user: str, memory_id: str) -> dict[str, Any] | None: + """ + Get a specific memory by ID. + + Args: + user: User identifier + memory_id: Memory ID + + Returns: + Memory data or None if not found + """ + collection_name = get_memory_collection_name(user) + + try: + points = self._client.retrieve( + collection_name=collection_name, + ids=[memory_id], + ) + + if not points: + return None + + point = points[0] + return { + "id": point.id, + **point.payload, + } + + except Exception as e: + logger.error( + "qdrant_get_memory_failed", + collection=collection_name, + memory_id=memory_id, + error=str(e), + ) + return None + + async def delete_memory(self, user: str, memory_id: str) -> bool: + """ + Delete a memory by ID. + + Args: + user: User identifier + memory_id: Memory ID to delete + + Returns: + True if deleted successfully, False otherwise + + Example: + >>> await client.delete_memory("jpmschweitzer", "mem_123") + True + """ + collection_name = get_memory_collection_name(user) + + try: + self._client.delete( + collection_name=collection_name, + points_selector=qdrant_models.PointIdsList( + points=[memory_id], + ), + ) + + logger.debug( + "qdrant_memory_deleted", + collection=collection_name, + memory_id=memory_id, + ) + return True + + except Exception as e: + logger.error( + "qdrant_delete_memory_failed", + collection=collection_name, + memory_id=memory_id, + error=str(e), + ) + return False + + async def delete_memories_by_type(self, user: str, memory_type: str) -> int: + """ + Delete all memories of a specific type. + + Args: + user: User identifier + memory_type: Type of memories to delete + + Returns: + Number of memories deleted (approximate) + """ + collection_name = get_memory_collection_name(user) + + try: + # Delete by filter + self._client.delete( + collection_name=collection_name, + points_selector=qdrant_models.FilterSelector( + filter=qdrant_models.Filter( + must=[ + qdrant_models.FieldCondition( + key="type", + match=qdrant_models.MatchValue(value=memory_type), + ) + ] + ) + ), + ) + + logger.info( + "qdrant_memories_deleted_by_type", + collection=collection_name, + memory_type=memory_type, + ) + return -1 # Qdrant doesn't return count for filter deletes + + except Exception as e: + logger.error( + "qdrant_delete_memories_by_type_failed", + collection=collection_name, + memory_type=memory_type, + error=str(e), + ) + return 0 + + async def count_memories(self, user: str) -> int: + """ + Count total memories for a user. + + Args: + user: User identifier + + Returns: + Number of memories in user's collection + """ + collection_name = get_memory_collection_name(user) + + try: + info = self._client.get_collection(collection_name) + return info.points_count + + except Exception as e: + logger.error( + "qdrant_count_memories_failed", + collection=collection_name, + error=str(e), + ) + return 0 + + async def health_check(self) -> bool: + """ + Check if Qdrant server is reachable. + + Returns: + True if healthy, False otherwise + """ + try: + self._client.get_collections() + return True + except Exception as e: + logger.error("qdrant_health_check_failed", error=str(e)) + return False + + +# Global client instance (lazy initialization) +_qdrant_client: MemoryQdrantClient | None = None + + +def get_qdrant_client() -> MemoryQdrantClient: + """ + Get global Qdrant client instance. + + Returns: + MemoryQdrantClient instance + """ + global _qdrant_client + if _qdrant_client is None: + _qdrant_client = MemoryQdrantClient() + return _qdrant_client diff --git a/src/responses/router.py b/src/responses/router.py index 1595409..db8f401 100644 --- a/src/responses/router.py +++ b/src/responses/router.py @@ -11,6 +11,7 @@ from sse_starlette.sse import EventSourceResponse from src.responses import service from src.responses.schemas import ResponseRequest, Response from src.core.exceptions import ModelNotFoundError, AppException +from src.core.context import current_user, current_conversation logger = logging.getLogger(__name__) @@ -94,6 +95,11 @@ async def create_response( """ logger.info(f"Response request for model: {request.model}") + # Set request context (propagates through all async calls) + user_token = current_user.set(request.user or "jpmschweitzer") + conv_id = request.metadata.get("conversation_id") if request.metadata else None + conv_token = current_conversation.set(conv_id) + try: # Check if this is a Tatlock request - use Steward preprocessing (Phase 2) model_id = request.model @@ -136,3 +142,8 @@ async def create_response( except Exception as e: logger.error(f"Unexpected error: {e}", exc_info=True) raise HTTPException(status_code=500, detail="Internal server error") + + finally: + # Reset context (important for connection reuse) + current_user.reset(user_token) + current_conversation.reset(conv_token) diff --git a/src/responses/schemas.py b/src/responses/schemas.py index 67f26c6..bf3a7b6 100644 --- a/src/responses/schemas.py +++ b/src/responses/schemas.py @@ -138,6 +138,10 @@ class ResponseRequest(CustomBaseModel): default=None, description="Stop sequences" ) + user: str | None = Field( + default=None, + description="Unique identifier for end-user (OpenAI standard)" + ) @field_validator('reasoning') @classmethod