Plumb collection manager into librarian

This commit is contained in:
Cyber MacGeddon 2025-09-19 11:01:26 +01:00
parent 4a7e552478
commit fca219bc30
2 changed files with 357 additions and 244 deletions

View file

@ -1,223 +1,106 @@
"""
Collection management service for the librarian
Collection management for the librarian
"""
import asyncio
import logging
from datetime import datetime
from .. base import AsyncProcessor, Consumer, Producer
from .. base import ConsumerMetrics, ProducerMetrics
from .. base.cassandra_config import add_cassandra_args, resolve_cassandra_config
from typing import Dict, Any, List, Optional
from .. schema import CollectionManagementRequest, CollectionManagementResponse, Error
from .. schema import collection_request_queue, collection_response_queue
from .. schema import CollectionMetadata
from .. schema import StorageManagementRequest, StorageManagementResponse
from .. schema import vector_storage_management_topic, object_storage_management_topic, triples_storage_management_topic, storage_management_response_topic
from .. exceptions import RequestError
from .. tables.library import LibraryTableStore
# Module logger
logger = logging.getLogger(__name__)
default_ident = "collection-management"
default_cassandra_host = "cassandra"
keyspace = "librarian"
class CollectionManager:
"""Manages collection metadata and coordinates collection operations across storage types"""
class Processor(AsyncProcessor):
def __init__(
self,
cassandra_host,
cassandra_username,
cassandra_password,
keyspace,
vector_storage_producer=None,
object_storage_producer=None,
triples_storage_producer=None,
storage_response_consumer=None
):
"""
Initialize the CollectionManager
def __init__(self, **params):
id = params.get("id", default_ident)
# Get Cassandra configuration
cassandra_host = params.get("cassandra_host", default_cassandra_host)
cassandra_username = params.get("cassandra_username")
cassandra_password = params.get("cassandra_password")
# Resolve configuration with environment variable fallback
hosts, username, password = resolve_cassandra_config(
host=cassandra_host,
username=cassandra_username,
password=cassandra_password
)
super(Processor, self).__init__(
**params | {
"cassandra_host": ','.join(hosts),
"cassandra_username": username
}
)
self.cassandra_host = hosts
self.cassandra_username = username
self.cassandra_password = password
# Set up metrics
collection_request_metrics = ConsumerMetrics(
processor=self.id, flow=None, name="collection-request"
)
collection_response_metrics = ProducerMetrics(
processor=self.id, flow=None, name="collection-response"
)
# Set up consumer for collection management requests
self.collection_request_consumer = Consumer(
taskgroup=self.taskgroup,
client=self.pulsar_client,
flow=None,
topic=collection_request_queue,
subscriber=id,
schema=CollectionManagementRequest,
handler=self.on_collection_request,
metrics=collection_request_metrics,
)
# Set up producer for collection management responses
self.collection_response_producer = Producer(
client=self.pulsar_client,
topic=collection_response_queue,
schema=CollectionManagementResponse,
metrics=collection_response_metrics,
)
# Set up producers for storage management requests
self.vector_storage_producer = Producer(
client=self.pulsar_client,
topic=vector_storage_management_topic,
schema=StorageManagementRequest,
)
self.object_storage_producer = Producer(
client=self.pulsar_client,
topic=object_storage_management_topic,
schema=StorageManagementRequest,
)
self.triples_storage_producer = Producer(
client=self.pulsar_client,
topic=triples_storage_management_topic,
schema=StorageManagementRequest,
)
# Set up consumer for storage management responses
storage_response_metrics = ConsumerMetrics(
processor=self.id, flow=None, name="storage-response"
)
self.storage_response_consumer = Consumer(
taskgroup=self.taskgroup,
client=self.pulsar_client,
flow=None,
topic=storage_management_response_topic,
subscriber=f"{id}-storage",
schema=StorageManagementResponse,
handler=self.on_storage_response,
metrics=storage_response_metrics,
)
# Initialize table store
Args:
cassandra_host: Cassandra host(s)
cassandra_username: Cassandra username
cassandra_password: Cassandra password
keyspace: Cassandra keyspace for library data
vector_storage_producer: Producer for vector storage management
object_storage_producer: Producer for object storage management
triples_storage_producer: Producer for triples storage management
storage_response_consumer: Consumer for storage management responses
"""
self.table_store = LibraryTableStore(
cassandra_host=self.cassandra_host,
cassandra_username=self.cassandra_username,
cassandra_password=self.cassandra_password,
keyspace=keyspace
cassandra_host, cassandra_username, cassandra_password, keyspace
)
# Track pending deletion requests by user+collection
self.pending_deletions = {} # (user, collection) -> {responses_pending, responses_received, all_successful, error_messages, deletion_complete}
# Storage management producers
self.vector_storage_producer = vector_storage_producer
self.object_storage_producer = object_storage_producer
self.triples_storage_producer = triples_storage_producer
self.storage_response_consumer = storage_response_consumer
async def on_collection_request(self, message):
"""Handle collection management requests"""
# Track pending deletion operations
self.pending_deletions = {}
logger.debug(f"Collection request: {message.operation}")
logger.info("Collection manager initialized")
async def ensure_collection_exists(self, user: str, collection: str):
"""
Ensure a collection exists, creating it if necessary (lazy creation)
Args:
user: User ID
collection: Collection ID
"""
try:
if message.operation == "list-collections":
response = await self.handle_list_collections(message)
elif message.operation == "update-collection":
response = await self.handle_update_collection(message)
elif message.operation == "delete-collection":
response = await self.handle_delete_collection(message)
else:
response = CollectionManagementResponse(
success="false",
error=Error(
type="invalid_operation",
message=f"Unknown operation: {message.operation}"
),
timestamp=datetime.now().isoformat()
)
# Check if collection already exists
existing = await self.table_store.get_collection(user, collection)
if existing:
logger.debug(f"Collection {user}/{collection} already exists")
return
except Exception as e:
logger.error(f"Error processing collection request: {e}", exc_info=True)
response = CollectionManagementResponse(
success="false",
error=Error(
type="processing_error",
message=str(e)
),
timestamp=datetime.now().isoformat()
# Create new collection with default metadata
logger.info(f"Creating new collection {user}/{collection}")
await self.table_store.create_collection(
user=user,
collection=collection,
name=collection, # Default name to collection ID
description="",
tags=set()
)
await self.collection_response_producer.send(response)
except Exception as e:
logger.error(f"Error ensuring collection exists: {e}")
# Don't fail the operation if collection creation fails
# This maintains backward compatibility
async def on_storage_response(self, response):
"""Handle storage management responses"""
logger.debug(f"Received storage response: error={response.error}")
async def list_collections(self, request: CollectionManagementRequest) -> CollectionManagementResponse:
"""
List collections for a user with optional tag filtering
# Find matching deletion by checking all pending deletions
# Note: This is simplified correlation - assumes responses come back quickly
# In production, we'd want better correlation mechanism
for deletion_key, info in list(self.pending_deletions.items()):
if info["responses_pending"] > 0:
# Record this response
info["responses_received"].append(response)
info["responses_pending"] -= 1
Args:
request: Collection management request
# Check if this response indicates failure
if response.error and response.error.message:
info["all_successful"] = False
info["error_messages"].append(response.error.message)
logger.warning(f"Storage deletion failed for {deletion_key}: {response.error.message}")
else:
logger.debug(f"Storage deletion succeeded for {deletion_key}")
# If all responses received, signal completion
if info["responses_pending"] == 0:
logger.info(f"All storage responses received for {deletion_key}")
info["deletion_complete"].set()
break # Only process for first matching deletion
# For now, we'll correlate by user+collection since we don't have deletion_id in the response
# This is a simplified approach - in production we'd want better correlation
for deletion_id, info in list(self.pending_deletions.items()):
if info["responses_pending"] > 0:
# Record this response
info["responses_received"].append(response)
info["responses_pending"] -= 1
# Check if this response indicates failure
if response.error and response.error.message:
info["all_successful"] = False
info["error_messages"].append(response.error.message)
logger.warning(f"Storage deletion failed for {deletion_id}: {response.error.message}")
# If all responses received, signal completion
if info["responses_pending"] == 0:
logger.info(f"All storage responses received for {deletion_id}")
info["deletion_complete"].set()
break # Only process for first matching deletion
async def handle_list_collections(self, message):
"""Handle list collections request"""
Returns:
CollectionManagementResponse with list of collections
"""
try:
tag_filter = list(message.tag_filter) if message.tag_filter else None
collections = await self.table_store.list_collections(message.user, tag_filter)
tag_filter = list(request.tag_filter) if request.tag_filter else None
collections = await self.table_store.list_collections(request.user, tag_filter)
collection_metadata = [
CollectionMetadata(
@ -240,18 +123,26 @@ class Processor(AsyncProcessor):
except Exception as e:
logger.error(f"Error listing collections: {e}")
raise
raise RequestError(f"Failed to list collections: {str(e)}")
async def handle_update_collection(self, message):
"""Handle update collection request"""
async def update_collection(self, request: CollectionManagementRequest) -> CollectionManagementResponse:
"""
Update collection metadata
Args:
request: Collection management request
Returns:
CollectionManagementResponse with updated collection
"""
try:
# Extract fields for update
name = message.name if message.name else None
description = message.description if message.description else None
tags = list(message.tags) if message.tags else None
name = request.name if request.name else None
description = request.description if request.description else None
tags = list(request.tags) if request.tags else None
updated_collection = await self.table_store.update_collection(
message.user, message.collection, name, description, tags
request.user, request.collection, name, description, tags
)
collection_metadata = CollectionMetadata(
@ -272,14 +163,22 @@ class Processor(AsyncProcessor):
except Exception as e:
logger.error(f"Error updating collection: {e}")
raise
raise RequestError(f"Failed to update collection: {str(e)}")
async def handle_delete_collection(self, message):
"""Handle delete collection request with cascade to all storage types"""
async def delete_collection(self, request: CollectionManagementRequest) -> CollectionManagementResponse:
"""
Delete collection with cascade to all storage types
Args:
request: Collection management request
Returns:
CollectionManagementResponse indicating success or failure
"""
try:
deletion_key = (message.user, message.collection)
deletion_key = (request.user, request.collection)
logger.info(f"Starting cascade deletion for {message.user}/{message.collection}")
logger.info(f"Starting cascade deletion for {request.user}/{request.collection}")
# Track this deletion request
self.pending_deletions[deletion_key] = {
@ -293,70 +192,94 @@ class Processor(AsyncProcessor):
# Create storage management request
storage_request = StorageManagementRequest(
operation="delete-collection",
user=message.user,
collection=message.collection
user=request.user,
collection=request.collection
)
# Send delete requests to all three storage types
await self.vector_storage_producer.send(storage_request)
await self.object_storage_producer.send(storage_request)
await self.triples_storage_producer.send(storage_request)
# Send deletion requests to all storage types
if self.vector_storage_producer:
await self.vector_storage_producer.send(storage_request)
if self.object_storage_producer:
await self.object_storage_producer.send(storage_request)
if self.triples_storage_producer:
await self.triples_storage_producer.send(storage_request)
logger.info(f"Storage deletion requests sent for {message.user}/{message.collection}")
# Wait for all storage responses (with timeout)
# Wait for all storage deletions to complete (with timeout)
deletion_info = self.pending_deletions[deletion_key]
try:
await asyncio.wait_for(
self.pending_deletions[deletion_key]["deletion_complete"].wait(),
deletion_info["deletion_complete"].wait(),
timeout=30.0 # 30 second timeout
)
except asyncio.TimeoutError:
logger.error(f"Timeout waiting for storage responses for {deletion_key}")
self.pending_deletions[deletion_key]["all_successful"] = False
self.pending_deletions[deletion_key]["error_messages"].append("Timeout waiting for storage responses")
logger.error(f"Timeout waiting for storage deletion responses for {deletion_key}")
deletion_info["all_successful"] = False
deletion_info["error_messages"].append("Timeout waiting for storage deletion")
# Check if all storage deletions were successful
deletion_info = self.pending_deletions.pop(deletion_key, {})
# Check if all deletions succeeded
if not deletion_info["all_successful"]:
error_msg = f"Storage deletion failed: {'; '.join(deletion_info['error_messages'])}"
logger.error(error_msg)
if deletion_info.get("all_successful", False):
# All storage deletions succeeded, now delete metadata
await self.table_store.delete_collection_metadata(message.user, message.collection)
logger.info(f"Successfully completed cascade deletion for {message.user}/{message.collection}")
return CollectionManagementResponse(
success="true",
timestamp=datetime.now().isoformat()
)
else:
# Some storage deletions failed
error_messages = deletion_info.get("error_messages", ["Unknown storage deletion error"])
error_msg = "; ".join(error_messages)
logger.error(f"Cascade deletion failed for {deletion_key}: {error_msg}")
# Clean up tracking
del self.pending_deletions[deletion_key]
return CollectionManagementResponse(
success="false",
error=Error(
type="storage_deletion_error",
message=f"Storage deletion failed: {error_msg}"
message=error_msg
),
timestamp=datetime.now().isoformat()
)
except Exception as e:
logger.error(f"Error in cascade deletion: {e}")
# All storage deletions succeeded, now delete metadata
logger.info(f"Storage deletions complete, removing metadata for {deletion_key}")
await self.table_store.delete_collection(request.user, request.collection)
# Clean up tracking
del self.pending_deletions[deletion_key]
return CollectionManagementResponse(
success="false",
error=Error(
type="deletion_error",
message=f"Failed to delete collection: {str(e)}"
),
success="true",
timestamp=datetime.now().isoformat()
)
@staticmethod
def add_args(parser):
AsyncProcessor.add_args(parser)
add_cassandra_args(parser)
except Exception as e:
logger.error(f"Error deleting collection: {e}")
# Clean up tracking on error
if deletion_key in self.pending_deletions:
del self.pending_deletions[deletion_key]
raise RequestError(f"Failed to delete collection: {str(e)}")
def run():
Processor.launch(default_ident, __doc__)
async def on_storage_response(self, response: StorageManagementResponse):
"""
Handle storage management responses for deletion tracking
Args:
response: Storage management response
"""
logger.debug(f"Received storage response: error={response.error}")
# Find matching deletion by checking all pending deletions
# Note: This is simplified correlation - in production we'd want better correlation
for deletion_key, info in list(self.pending_deletions.items()):
if info["responses_pending"] > 0:
# Record this response
info["responses_received"].append(response)
info["responses_pending"] -= 1
# Check if this response indicates failure
if response.error and response.error.message:
info["all_successful"] = False
info["error_messages"].append(response.error.message)
logger.warning(f"Storage deletion failed for {deletion_key}: {response.error.message}")
else:
logger.debug(f"Storage deletion succeeded for {deletion_key}")
# If all responses received, signal completion
if info["responses_pending"] == 0:
logger.info(f"All storage responses received for {deletion_key}")
info["deletion_complete"].set()
break # Only process for first matching deletion

View file

@ -8,6 +8,7 @@ import asyncio
import base64
import json
import logging
from datetime import datetime
from .. base import AsyncProcessor, Consumer, Producer, Publisher, Subscriber
from .. base import ConsumerMetrics, ProducerMetrics
@ -15,6 +16,11 @@ from .. base.cassandra_config import add_cassandra_args, resolve_cassandra_confi
from .. schema import LibrarianRequest, LibrarianResponse, Error
from .. schema import librarian_request_queue, librarian_response_queue
from .. schema import CollectionManagementRequest, CollectionManagementResponse
from .. schema import collection_request_queue, collection_response_queue
from .. schema import StorageManagementRequest, StorageManagementResponse
from .. schema import vector_storage_management_topic, object_storage_management_topic
from .. schema import triples_storage_management_topic, storage_management_response_topic
from .. schema import Document, Metadata
from .. schema import TextDocument, Metadata
@ -22,6 +28,7 @@ from .. schema import TextDocument, Metadata
from .. exceptions import RequestError
from . librarian import Librarian
from . collection_manager import CollectionManager
# Module logger
logger = logging.getLogger(__name__)
@ -30,6 +37,8 @@ default_ident = "librarian"
default_librarian_request_queue = librarian_request_queue
default_librarian_response_queue = librarian_response_queue
default_collection_request_queue = collection_request_queue
default_collection_response_queue = collection_response_queue
default_minio_host = "minio:9000"
default_minio_access_key = "minioadmin"
@ -57,6 +66,14 @@ class Processor(AsyncProcessor):
"librarian_response_queue", default_librarian_response_queue
)
collection_request_queue = params.get(
"collection_request_queue", default_collection_request_queue
)
collection_response_queue = params.get(
"collection_response_queue", default_collection_response_queue
)
minio_host = params.get("minio_host", default_minio_host)
minio_access_key = params.get(
"minio_access_key",
@ -87,6 +104,8 @@ class Processor(AsyncProcessor):
**params | {
"librarian_request_queue": librarian_request_queue,
"librarian_response_queue": librarian_response_queue,
"collection_request_queue": collection_request_queue,
"collection_response_queue": collection_response_queue,
"minio_host": minio_host,
"minio_access_key": minio_access_key,
"cassandra_host": self.cassandra_host,
@ -103,6 +122,18 @@ class Processor(AsyncProcessor):
processor = self.id, flow = None, name = "librarian-response"
)
collection_request_metrics = ConsumerMetrics(
processor = self.id, flow = None, name = "collection-request"
)
collection_response_metrics = ProducerMetrics(
processor = self.id, flow = None, name = "collection-response"
)
storage_response_metrics = ConsumerMetrics(
processor = self.id, flow = None, name = "storage-response"
)
self.librarian_request_consumer = Consumer(
taskgroup = self.taskgroup,
client = self.pulsar_client,
@ -121,6 +152,54 @@ class Processor(AsyncProcessor):
metrics = librarian_response_metrics,
)
self.collection_request_consumer = Consumer(
taskgroup = self.taskgroup,
client = self.pulsar_client,
flow = None,
topic = collection_request_queue,
subscriber = id,
schema = CollectionManagementRequest,
handler = self.on_collection_request,
metrics = collection_request_metrics,
)
self.collection_response_producer = Producer(
client = self.pulsar_client,
topic = collection_response_queue,
schema = CollectionManagementResponse,
metrics = collection_response_metrics,
)
# Storage management producers for collection deletion
self.vector_storage_producer = Producer(
client = self.pulsar_client,
topic = vector_storage_management_topic,
schema = StorageManagementRequest,
)
self.object_storage_producer = Producer(
client = self.pulsar_client,
topic = object_storage_management_topic,
schema = StorageManagementRequest,
)
self.triples_storage_producer = Producer(
client = self.pulsar_client,
topic = triples_storage_management_topic,
schema = StorageManagementRequest,
)
self.storage_response_consumer = Consumer(
taskgroup = self.taskgroup,
client = self.pulsar_client,
flow = None,
topic = storage_management_response_topic,
subscriber = id,
schema = StorageManagementResponse,
handler = self.on_storage_response,
metrics = storage_response_metrics,
)
self.librarian = Librarian(
cassandra_host = self.cassandra_host,
cassandra_username = self.cassandra_username,
@ -133,6 +212,17 @@ class Processor(AsyncProcessor):
load_document = self.load_document,
)
self.collection_manager = CollectionManager(
cassandra_host = self.cassandra_host,
cassandra_username = self.cassandra_username,
cassandra_password = self.cassandra_password,
keyspace = keyspace,
vector_storage_producer = self.vector_storage_producer,
object_storage_producer = self.object_storage_producer,
triples_storage_producer = self.triples_storage_producer,
storage_response_consumer = self.storage_response_consumer,
)
self.register_config_handler(self.on_librarian_config)
self.flows = {}
@ -144,6 +234,12 @@ class Processor(AsyncProcessor):
await super(Processor, self).start()
await self.librarian_request_consumer.start()
await self.librarian_response_producer.start()
await self.collection_request_consumer.start()
await self.collection_response_producer.start()
await self.vector_storage_producer.start()
await self.object_storage_producer.start()
await self.triples_storage_producer.start()
await self.storage_response_consumer.start()
async def on_librarian_config(self, config, version):
@ -223,6 +319,19 @@ class Processor(AsyncProcessor):
logger.debug("Document submitted")
async def add_processing_with_collection(self, request):
"""
Wrapper for add_processing that ensures collection exists
"""
# Ensure collection exists when processing is added
if hasattr(request, 'processing_metadata') and request.processing_metadata:
user = request.processing_metadata.user
collection = request.processing_metadata.collection
await self.collection_manager.ensure_collection_exists(user, collection)
# Call the original add_processing method
return await self.librarian.add_processing(request)
async def process_request(self, v):
if v.operation is None:
@ -236,7 +345,7 @@ class Processor(AsyncProcessor):
"update-document": self.librarian.update_document,
"get-document-metadata": self.librarian.get_document_metadata,
"get-document-content": self.librarian.get_document_content,
"add-processing": self.librarian.add_processing,
"add-processing": self.add_processing_with_collection,
"remove-processing": self.librarian.remove_processing,
"list-documents": self.librarian.list_documents,
"list-processing": self.librarian.list_processing,
@ -296,6 +405,75 @@ class Processor(AsyncProcessor):
logger.debug("Librarian input processing complete")
async def process_collection_request(self, v):
"""
Process collection management requests
"""
if v.operation is None:
raise RequestError("Null operation")
logger.debug(f"Collection request: {v.operation}")
impls = {
"list-collections": self.collection_manager.list_collections,
"update-collection": self.collection_manager.update_collection,
"delete-collection": self.collection_manager.delete_collection,
}
if v.operation not in impls:
raise RequestError(f"Invalid collection operation: {v.operation}")
return await impls[v.operation](v)
async def on_collection_request(self, msg, consumer, flow):
"""
Handle collection management request messages
"""
v = msg.value()
id = msg.properties().get("id", "unknown")
logger.info(f"Handling collection request {id}...")
try:
resp = await self.process_collection_request(v)
await self.collection_response_producer.send(
resp, properties={"id": id}
)
except RequestError as e:
resp = CollectionManagementResponse(
success="false",
error=Error(
type="request-error",
message=str(e),
),
timestamp=datetime.now().isoformat()
)
await self.collection_response_producer.send(
resp, properties={"id": id}
)
except Exception as e:
resp = CollectionManagementResponse(
success="false",
error=Error(
type="unexpected-error",
message=str(e),
),
timestamp=datetime.now().isoformat()
)
await self.collection_response_producer.send(
resp, properties={"id": id}
)
logger.debug("Collection request processing complete")
async def on_storage_response(self, msg, consumer, flow):
"""
Handle storage management response messages
"""
v = msg.value()
logger.debug("Received storage management response")
await self.collection_manager.on_storage_response(v)
@staticmethod
def add_args(parser):
@ -313,6 +491,18 @@ class Processor(AsyncProcessor):
help=f'Config response queue {default_librarian_response_queue}',
)
parser.add_argument(
'--collection-request-queue',
default=default_collection_request_queue,
help=f'Collection request queue (default: {default_collection_request_queue})'
)
parser.add_argument(
'--collection-response-queue',
default=default_collection_response_queue,
help=f'Collection response queue (default: {default_collection_response_queue})'
)
parser.add_argument(
'--minio-host',
default=default_minio_host,