| """Vector database factory pattern""" |
|
|
| from typing import Dict, Any |
| from .base import VectorDBProvider |
| from .chroma_db import ChromaDB |
| import logging |
|
|
| logger = logging.getLogger(__name__) |
|
|
|
|
| class VectorDBFactory: |
| """Factory for creating vector database instances""" |
|
|
| _providers = { |
| "chroma": ChromaDB, |
| } |
|
|
| @staticmethod |
| def create(db_type: str, config: Dict[str, Any]) -> VectorDBProvider: |
| """Create a vector database provider instance""" |
| if db_type not in VectorDBFactory._providers: |
| available = ", ".join(VectorDBFactory._providers.keys()) |
| raise ValueError(f"Unknown database type '{db_type}'. Available: {available}") |
|
|
| provider_class = VectorDBFactory._providers[db_type] |
| provider = provider_class() |
| provider.initialize(config) |
|
|
| logger.info(f"Created {db_type} vector database provider") |
| return provider |
|
|
| @staticmethod |
| def register(db_type: str, provider_class: type) -> None: |
| """Register a new vector database provider""" |
| VectorDBFactory._providers[db_type] = provider_class |
| logger.info(f"Registered vector database provider: {db_type}") |
|
|
| @staticmethod |
| def available_providers() -> list: |
| """Get list of available providers""" |
| return list(VectorDBFactory._providers.keys()) |
|
|