feat: add Phase F.1 memory infrastructure

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 <noreply@anthropic.com>
This commit is contained in:
2025-12-13 17:27:51 +01:00
co-authored by Claude Opus 4.5
parent c049c1e354
commit 86bde271d5
11 changed files with 1465 additions and 24 deletions
+6
View File
@@ -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.
+4
View File
@@ -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
+36 -23
View File
@@ -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(
+41 -1
View File
@@ -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:
"""
+111
View File
@@ -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)
+269
View File
@@ -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
+390
View File
@@ -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
+147
View File
@@ -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
+446
View File
@@ -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
+11
View File
@@ -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)
+4
View File
@@ -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