-
Notifications
You must be signed in to change notification settings - Fork 2
Expand file tree
/
Copy pathVectorDBProviderFactory.py
More file actions
36 lines (26 loc) · 1.54 KB
/
Copy pathVectorDBProviderFactory.py
File metadata and controls
36 lines (26 loc) · 1.54 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
from .providers import QdrantDBProvider, PGVectorProvider
from .VectoDBEnums import VectorDBEnums, DistanceMethodEnums
from controllers.BaseController import BaseController
from sqlalchemy.orm import sessionmaker
class VectorDBProviderFactory:
def __init__(self, config: dict, db_client: sessionmaker=None):
self.config = config
self.base_controller = BaseController()
self.db_client = db_client
def create(self, provider: str):
if provider == VectorDBEnums.QDRANT.value:
qdrant_db_url = self.config.VECTOR_DB_URL
qdrant_db_client = None
if not qdrant_db_url:
qdrant_db_client = self.base_controller.get_database_path(self.config.VECTOR_DB_PATH)
return QdrantDBProvider(db_client=qdrant_db_client,
db_url=qdrant_db_url,
distance_method=self.config.VECTOR_DB_DISTANCE_METHOD,
default_vector_size=self.config.EMBEDDING_MODEL_SIZE,
index_threshold=self.config.VECTOR_DB_PGVEC_INDEX_THRESHOLD)
if provider == VectorDBEnums.PGVECTOR.value:
return PGVectorProvider(db_client=self.db_client,
distance_method=self.config.VECTOR_DB_DISTANCE_METHOD,
default_vector_size=self.config.EMBEDDING_MODEL_SIZE,
index_threshold=self.config.VECTOR_DB_PGVEC_INDEX_THRESHOLD)
return None