codekingpro/portable-devtools
114k
1from __future__ import annotations2 3import logging4import uuid5from typing import Any, Dict, Iterable, List, Optional, Tuple, Type6 7import sqlalchemy8from sqlalchemy import func9from sqlalchemy.dialects.postgresql import JSON, UUID10from sqlalchemy.orm import Session, relationship11 12try:13 from sqlalchemy.orm import declarative_base14except ImportError:15 from sqlalchemy.ext.declarative import declarative_base16 17from langchain_core.documents import Document18from langchain_core.embeddings import Embeddings19from langchain_core.utils import get_from_dict_or_env20from langchain_core.vectorstores import VectorStore21 22Base = declarative_base() # type: Any23 24 25ADA_TOKEN_COUNT = 153626_LANGCHAIN_DEFAULT_COLLECTION_NAME = "langchain"27 28 29class BaseModel(Base):30 """Base model for all SQL stores."""31 32 __abstract__ = True33 uuid = sqlalchemy.Column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4)34 35 36class CollectionStore(BaseModel):37 """Collection store."""38 39 __tablename__ = "langchain_pg_collection"40 41 name = sqlalchemy.Column(sqlalchemy.String)42 cmetadata = sqlalchemy.Column(JSON)43 44 embeddings = relationship(45 "EmbeddingStore",46 back_populates="collection",47 passive_deletes=True,48 )49 50 @classmethod51 def get_by_name(cls, session: Session, name: str) -> Optional["CollectionStore"]:52 return session.query(cls).filter(cls.name == name).first()53 54 @classmethod55 def get_or_create(56 cls,57 session: Session,58 name: str,59 cmetadata: Optional[dict] = None,60 ) -> Tuple["CollectionStore", bool]:61 """62 Get or create a collection.63 Returns [Collection, bool] where the bool is True if the collection was created.64 """65 created = False66 collection = cls.get_by_name(session, name)67 if collection:68 return collection, created69 70 collection = cls(name=name, cmetadata=cmetadata)71 session.add(collection)72 session.commit()73 created = True74 return collection, created75 76 77class EmbeddingStore(BaseModel):78 """Embedding store."""79 80 __tablename__ = "langchain_pg_embedding"81 82 collection_id = sqlalchemy.Column(83 UUID(as_uuid=True),84 sqlalchemy.ForeignKey(85 f"{CollectionStore.__tablename__}.uuid",86 ondelete="CASCADE",87 ),88 )89 collection = relationship(CollectionStore, back_populates="embeddings")90 91 embedding = sqlalchemy.Column(sqlalchemy.ARRAY(sqlalchemy.REAL)) # type: ignore[var-annotated]92 document = sqlalchemy.Column(sqlalchemy.String, nullable=True)93 cmetadata = sqlalchemy.Column(JSON, nullable=True)94 95 # custom_id : any user defined id96 custom_id = sqlalchemy.Column(sqlalchemy.String, nullable=True)97 98 99class QueryResult:100 """Result from a query."""101 102 EmbeddingStore: EmbeddingStore103 distance: float104 105 106class PGEmbedding(VectorStore):107 """`Postgres` with the `pg_embedding` extension as a vector store.108 109 pg_embedding uses sequential scan by default. but you can create a HNSW index110 using the create_hnsw_index method.111 - `connection_string` is a postgres connection string.112 - `embedding_function` any embedding function implementing113 `langchain.embeddings.base.Embeddings` interface.114 - `collection_name` is the name of the collection to use. (default: langchain)115 - NOTE: This is not the name of the table, but the name of the collection.116 The tables will be created when initializing the store (if not exists)117 So, make sure the user has the right permissions to create tables.118 - `distance_strategy` is the distance strategy to use. (default: EUCLIDEAN)119 - `EUCLIDEAN` is the euclidean distance.120 - `pre_delete_collection` if True, will delete the collection if it exists.121 (default: False)122 - Useful for testing.123 """124 125 def __init__(126 self,127 connection_string: str,128 embedding_function: Embeddings,129 collection_name: str = _LANGCHAIN_DEFAULT_COLLECTION_NAME,130 collection_metadata: Optional[dict] = None,131 pre_delete_collection: bool = False,132 logger: Optional[logging.Logger] = None,133 ) -> None:134 self.connection_string = connection_string135 self.embedding_function = embedding_function136 self.collection_name = collection_name137 self.collection_metadata = collection_metadata138 self.pre_delete_collection = pre_delete_collection139 self.logger = logger or logging.getLogger(__name__)140 self.__post_init__()141 142 def __post_init__(143 self,144 ) -> None:145 self._conn = self.connect()146 self.create_hnsw_extension()147 self.create_tables_if_not_exists()148 self.create_collection()149 150 @property151 def embeddings(self) -> Embeddings:152 return self.embedding_function153 154 def connect(self) -> sqlalchemy.engine.Connection:155 engine = sqlalchemy.create_engine(self.connection_string)156 conn = engine.connect()157 return conn158 159 def create_hnsw_extension(self) -> None:160 try:161 with Session(self._conn) as session:162 statement = sqlalchemy.text("CREATE EXTENSION IF NOT EXISTS embedding")163 session.execute(statement)164 session.commit()165 except Exception as e:166 self.logger.exception(e)167 168 def create_tables_if_not_exists(self) -> None:169 with self._conn.begin():170 Base.metadata.create_all(self._conn)171 172 def drop_tables(self) -> None:173 with self._conn.begin():174 Base.metadata.drop_all(self._conn)175 176 def create_collection(self) -> None:177 if self.pre_delete_collection:178 self.delete_collection()179 with Session(self._conn) as session:180 CollectionStore.get_or_create(181 session, self.collection_name, cmetadata=self.collection_metadata182 )183 184 def create_hnsw_index(185 self,186 max_elements: int = 10000,187 dims: int = ADA_TOKEN_COUNT,188 m: int = 8,189 ef_construction: int = 16,190 ef_search: int = 16,191 ) -> None:192 create_index_query = sqlalchemy.text(193 "CREATE INDEX IF NOT EXISTS langchain_pg_embedding_idx "194 "ON langchain_pg_embedding USING hnsw (embedding) "195 "WITH ("196 "maxelements = {}, "197 "dims = {}, "198 "m = {}, "199 "efconstruction = {}, "200 "efsearch = {}"201 ");".format(max_elements, dims, m, ef_construction, ef_search)202 )203 204 # Execute the queries205 try:206 with Session(self._conn) as session:207 # Create the HNSW index208 session.execute(create_index_query)209 session.commit()210 print("HNSW extension and index created successfully.") # noqa: T201211 except Exception as e:212 print(f"Failed to create HNSW extension or index: {e}") # noqa: T201213 214 def delete_collection(self) -> None:215 self.logger.debug("Trying to delete collection")216 with Session(self._conn) as session:217 collection = self.get_collection(session)218 if not collection:219 self.logger.warning("Collection not found")220 return221 session.delete(collection)222 session.commit()223 224 def get_collection(self, session: Session) -> Optional["CollectionStore"]:225 return CollectionStore.get_by_name(session, self.collection_name)226 227 @classmethod228 def _initialize_from_embeddings(229 cls,230 texts: List[str],231 embeddings: List[List[float]],232 embedding: Embeddings,233 metadatas: Optional[List[dict]] = None,234 ids: Optional[List[str]] = None,235 collection_name: str = _LANGCHAIN_DEFAULT_COLLECTION_NAME,236 pre_delete_collection: bool = False,237 **kwargs: Any,238 ) -> PGEmbedding:239 if ids is None:240 ids = [str(uuid.uuid4()) for _ in texts]241 242 if not metadatas:243 metadatas = [{} for _ in texts]244 245 connection_string = cls.get_connection_string(kwargs)246 247 store = cls(248 connection_string=connection_string,249 collection_name=collection_name,250 embedding_function=embedding,251 pre_delete_collection=pre_delete_collection,252 )253 254 store.add_embeddings(255 texts=texts, embeddings=embeddings, metadatas=metadatas, ids=ids, **kwargs256 )257 258 return store259 260 def add_embeddings(261 self,262 texts: List[str],263 embeddings: List[List[float]],264 metadatas: List[dict],265 ids: List[str],266 **kwargs: Any,267 ) -> None:268 with Session(self._conn) as session:269 collection = self.get_collection(session)270 if not collection:271 raise ValueError("Collection not found")272 for text, metadata, embedding, id in zip(texts, metadatas, embeddings, ids):273 embedding_store = EmbeddingStore(274 embedding=embedding,275 document=text,276 cmetadata=metadata,277 custom_id=id,278 )279 collection.embeddings.append(embedding_store)280 session.add(embedding_store)281 session.commit()282 283 def add_texts(284 self,285 texts: Iterable[str],286 metadatas: Optional[List[dict]] = None,287 ids: Optional[List[str]] = None,288 **kwargs: Any,289 ) -> List[str]:290 if ids is None:291 ids = [str(uuid.uuid4()) for _ in texts]292 293 embeddings = self.embedding_function.embed_documents(list(texts))294 295 if not metadatas:296 metadatas = [{} for _ in texts]297 298 with Session(self._conn) as session:299 collection = self.get_collection(session)300 if not collection:301 raise ValueError("Collection not found")302 for text, metadata, embedding, id in zip(texts, metadatas, embeddings, ids):303 embedding_store = EmbeddingStore(304 embedding=embedding,305 document=text,306 cmetadata=metadata,307 custom_id=id,308 )309 collection.embeddings.append(embedding_store)310 session.add(embedding_store)311 session.commit()312 313 return ids314 315 def similarity_search(316 self,317 query: str,318 k: int = 4,319 filter: Optional[dict] = None,320 **kwargs: Any,321 ) -> List[Document]:322 embedding = self.embedding_function.embed_query(text=query)323 return self.similarity_search_by_vector(324 embedding=embedding,325 k=k,326 filter=filter,327 )328 329 def similarity_search_with_score(330 self,331 query: str,332 k: int = 4,333 filter: Optional[dict] = None,334 ) -> List[Tuple[Document, float]]:335 embedding = self.embedding_function.embed_query(query)336 docs = self.similarity_search_with_score_by_vector(337 embedding=embedding, k=k, filter=filter338 )339 return docs340 341 def similarity_search_with_score_by_vector(342 self,343 embedding: List[float],344 k: int = 4,345 filter: Optional[dict] = None,346 ) -> List[Tuple[Document, float]]:347 with Session(self._conn) as session:348 collection = self.get_collection(session)349 set_enable_seqscan_stmt = sqlalchemy.text("SET enable_seqscan = off")350 session.execute(set_enable_seqscan_stmt)351 if not collection:352 raise ValueError("Collection not found")353 354 filter_by = EmbeddingStore.collection_id == collection.uuid355 356 if filter is not None:357 filter_clauses = []358 for key, value in filter.items():359 IN = "in"360 if isinstance(value, dict) and IN in map(str.lower, value):361 value_case_insensitive = {362 k.lower(): v for k, v in value.items()363 }364 filter_by_metadata = EmbeddingStore.cmetadata[key].astext.in_(365 value_case_insensitive[IN]366 )367 filter_clauses.append(filter_by_metadata)368 elif isinstance(value, dict) and "substring" in map(369 str.lower, value370 ):371 filter_by_metadata = EmbeddingStore.cmetadata[key].astext.ilike(372 f"%{value['substring']}%"373 )374 filter_clauses.append(filter_by_metadata)375 else:376 filter_by_metadata = EmbeddingStore.cmetadata[377 key378 ].astext == str(value)379 filter_clauses.append(filter_by_metadata)380 381 filter_by = sqlalchemy.and_(filter_by, *filter_clauses)382 383 results: List[QueryResult] = (384 session.query(385 EmbeddingStore,386 func.abs(EmbeddingStore.embedding.op("<->")(embedding)).label(387 "distance"388 ),389 ) # Specify the columns you need here, e.g., EmbeddingStore.embedding390 .filter(filter_by)391 .order_by(392 func.abs(EmbeddingStore.embedding.op("<->")(embedding)).asc()393 ) # Using PostgreSQL specific operator with the correct column name394 .limit(k)395 .all()396 )397 398 docs = [399 (400 Document(401 page_content=result.EmbeddingStore.document, # type: ignore[arg-type]402 metadata=result.EmbeddingStore.cmetadata,403 ),404 result.distance if self.embedding_function is not None else 0.0,405 )406 for result in results407 ]408 return docs409 410 def similarity_search_by_vector(411 self,412 embedding: List[float],413 k: int = 4,414 filter: Optional[dict] = None,415 **kwargs: Any,416 ) -> List[Document]:417 docs_and_scores = self.similarity_search_with_score_by_vector(418 embedding=embedding, k=k, filter=filter419 )420 return [doc for doc, _ in docs_and_scores]421 422 @classmethod423 def from_texts(424 cls: Type[PGEmbedding],425 texts: List[str],426 embedding: Embeddings,427 metadatas: Optional[List[dict]] = None,428 collection_name: str = _LANGCHAIN_DEFAULT_COLLECTION_NAME,429 ids: Optional[List[str]] = None,430 pre_delete_collection: bool = False,431 **kwargs: Any,432 ) -> PGEmbedding:433 embeddings = embedding.embed_documents(list(texts))434 435 return cls._initialize_from_embeddings(436 texts,437 embeddings,438 embedding,439 metadatas=metadatas,440 ids=ids,441 collection_name=collection_name,442 pre_delete_collection=pre_delete_collection,443 **kwargs,444 )445 446 @classmethod447 def from_embeddings(448 cls,449 text_embeddings: List[Tuple[str, List[float]]],450 embedding: Embeddings,451 metadatas: Optional[List[dict]] = None,452 collection_name: str = _LANGCHAIN_DEFAULT_COLLECTION_NAME,453 ids: Optional[List[str]] = None,454 pre_delete_collection: bool = False,455 **kwargs: Any,456 ) -> PGEmbedding:457 texts = [t[0] for t in text_embeddings]458 embeddings = [t[1] for t in text_embeddings]459 460 return cls._initialize_from_embeddings(461 texts,462 embeddings,463 embedding,464 metadatas=metadatas,465 ids=ids,466 collection_name=collection_name,467 pre_delete_collection=pre_delete_collection,468 **kwargs,469 )470 471 @classmethod472 def from_existing_index(473 cls: Type[PGEmbedding],474 embedding: Embeddings,475 collection_name: str = _LANGCHAIN_DEFAULT_COLLECTION_NAME,476 pre_delete_collection: bool = False,477 **kwargs: Any,478 ) -> PGEmbedding:479 connection_string = cls.get_connection_string(kwargs)480 481 store = cls(482 connection_string=connection_string,483 collection_name=collection_name,484 embedding_function=embedding,485 pre_delete_collection=pre_delete_collection,486 )487 488 return store489 490 @classmethod491 def get_connection_string(cls, kwargs: Dict[str, Any]) -> str:492 connection_string: str = get_from_dict_or_env(493 data=kwargs,494 key="connection_string",495 env_key="POSTGRES_CONNECTION_STRING",496 )497 498 if not connection_string:499 raise ValueError(500 "Postgres connection string is required"501 "Either pass it as a parameter"502 "or set the POSTGRES_CONNECTION_STRING environment variable."503 )504 505 return connection_string506 507 @classmethod508 def from_documents(509 cls: Type[PGEmbedding],510 documents: List[Document],511 embedding: Embeddings,512 collection_name: str = _LANGCHAIN_DEFAULT_COLLECTION_NAME,513 ids: Optional[List[str]] = None,514 pre_delete_collection: bool = False,515 **kwargs: Any,516 ) -> PGEmbedding:517 texts = [d.page_content for d in documents]518 metadatas = [d.metadata for d in documents]519 connection_string = cls.get_connection_string(kwargs)520 521 kwargs["connection_string"] = connection_string522 523 return cls.from_texts(524 texts=texts,525 pre_delete_collection=pre_delete_collection,526 embedding=embedding,527 metadatas=metadatas,528 ids=ids,529 collection_name=collection_name,530 **kwargs,531 )532 