diff --git a/src/poetry.lock b/src/poetry.lock index e9ddd086..c11e21f5 100644 --- a/src/poetry.lock +++ b/src/poetry.lock @@ -1204,7 +1204,7 @@ version = "1.9.0" description = "Distro - an OS platform information API" optional = false python-versions = ">=3.6" -groups = ["main", "optional"] +groups = ["main"] files = [ {file = "distro-1.9.0-py3-none-any.whl", hash = "sha256:7bffd925d65168f85027d8da9af6bddab658135b840670a223589bc0c8ef02b2"}, {file = "distro-1.9.0.tar.gz", hash = "sha256:2fa77c6fd8940f116ee1d6b94a2f90b13b5ea8d019b98bc8bafdcabcdd9bdbed"}, @@ -2121,7 +2121,7 @@ version = "1.33" description = "Apply JSON-Patches (RFC 6902)" optional = false python-versions = ">=2.7, !=3.0.*, !=3.1.*, !=3.2.*, !=3.3.*, !=3.4.*, !=3.5.*, !=3.6.*" -groups = ["main", "optional"] +groups = ["main"] files = [ {file = "jsonpatch-1.33-py2.py3-none-any.whl", hash = "sha256:0ae28c0cd062bbd8b8ecc26d7d164fbbea9652a1a3693f3b956c1eae5145dade"}, {file = "jsonpatch-1.33.tar.gz", hash = "sha256:9fcd4009c41e6d12348b4a0ff2563ba56a2923a7dfee731d004e212e1ee5030c"}, @@ -2136,7 +2136,7 @@ version = "3.1.1" description = "Identify specific nodes in a JSON document (RFC 6901) " optional = false python-versions = ">=3.10" -groups = ["main", "optional"] +groups = ["main"] files = [ {file = "jsonpointer-3.1.1-py3-none-any.whl", hash = "sha256:8ff8b95779d071ba472cf5bc913028df06031797532f08a7d5b602d8b2a488ca"}, {file = "jsonpointer-3.1.1.tar.gz", hash = "sha256:0b801c7db33a904024f6004d526dcc53bbb8a4a0f4e32bfd10beadf60adf1900"}, @@ -2206,26 +2206,6 @@ websocket-client = ">=0.32.0,<0.40.0 || >0.40.0,<0.41.dev0 || >=0.43.dev0" [package.extras] google-auth = ["google-auth (>=1.0.1)"] -[[package]] -name = "langchain-chroma" -version = "0.2.6" -description = "An integration package connecting Chroma and LangChain." -optional = false -python-versions = ">=3.9" -groups = ["optional"] -files = [ - {file = "langchain_chroma-0.2.6-py3-none-any.whl", hash = "sha256:d7e10101b0942cd990eedb798c3d85ed3e8415a992c8a388843196f6ab97b41b"}, - {file = "langchain_chroma-0.2.6.tar.gz", hash = "sha256:ec5ca0f6f7692ac053741e076ea086c4be0cfcb5846c8693b1bcc3089c88b65e"}, -] - -[package.dependencies] -chromadb = ">=1.0.20" -langchain-core = ">=0.3.76" -numpy = [ - {version = ">=1.26.0", markers = "python_version < \"3.13\""}, - {version = ">=2.1.0", markers = "python_version >= \"3.13\""}, -] - [[package]] name = "langchain-classic" version = "1.0.8" @@ -2301,7 +2281,7 @@ version = "1.5.2" description = "Building applications with LLMs through composability" optional = false python-versions = "<4.0.0,>=3.10.0" -groups = ["main", "optional"] +groups = ["main"] files = [ {file = "langchain_core-1.5.2-py3-none-any.whl", hash = "sha256:a687dd7c3b22c6c1294e1c1eeb61fb6f3a308e6015a1d75b576b3836ad5b5aed"}, {file = "langchain_core-1.5.2.tar.gz", hash = "sha256:2d13ab35b42eec63d4669a483776b8cdd778ee764107149369fb369d84c08c41"}, @@ -2341,7 +2321,7 @@ version = "0.0.18" description = "Python bindings for the LangChain agent streaming protocol" optional = false python-versions = "<4.0.0,>=3.10.0" -groups = ["main", "optional"] +groups = ["main"] files = [ {file = "langchain_protocol-0.0.18-py3-none-any.whl", hash = "sha256:70b53a86fbf9cedc863555effe44da192ab02d556ddbf2cf95b8873adcf41b5a"}, {file = "langchain_protocol-0.0.18.tar.gz", hash = "sha256:ec3e11782f1ed0c9db38e5a9ed01b0e7a0d3fba406faa8aef6594b73c56a63e6"}, @@ -2371,7 +2351,7 @@ version = "0.10.9" description = "Client library to connect to the LangSmith Observability and Evaluation Platform." optional = false python-versions = ">=3.10" -groups = ["main", "optional"] +groups = ["main"] files = [ {file = "langsmith-0.10.9-py3-none-any.whl", hash = "sha256:5e0e8ab0f8df05710809919184495e33c2a7c9a9a5e8861d63dd12c1226d9c79"}, {file = "langsmith-0.10.9.tar.gz", hash = "sha256:195bc67c964a6370cb91742ce9fa07ce69bfae47977f0fb3f41d125b3435d03a"}, @@ -4897,7 +4877,7 @@ version = "1.0.0" description = "A utility belt for advanced users of python-requests" optional = false python-versions = ">=2.7, !=3.0.*, !=3.1.*, !=3.2.*, !=3.3.*" -groups = ["main", "optional"] +groups = ["main"] files = [ {file = "requests-toolbelt-1.0.0.tar.gz", hash = "sha256:7681a0a3d047012b5bdc0ee37d7f8f07ebe76ab08caeccfc3921ce23c88d5bc6"}, {file = "requests_toolbelt-1.0.0-py2.py3-none-any.whl", hash = "sha256:cccfdd665f0a24fcf4726e690f65639d272bb0637b9b92dfd91a5568ccf6bd06"}, @@ -5316,7 +5296,7 @@ version = "1.3.1" description = "Sniff out which async library your code is running under" optional = false python-versions = ">=3.7" -groups = ["main", "optional"] +groups = ["main"] files = [ {file = "sniffio-1.3.1-py3-none-any.whl", hash = "sha256:2f6da418d1f1e0fddd844478f41680e794e6051915791a034ff65e5f100525a2"}, {file = "sniffio-1.3.1.tar.gz", hash = "sha256:f4324edc670a0f49750a81b895f35c3adb843cca46f0530f79fc1babb23789dc"}, @@ -6144,7 +6124,7 @@ version = "0.17.0" description = "Fast, drop-in replacement for Python's uuid module, powered by Rust." optional = false python-versions = ">=3.10" -groups = ["main", "optional"] +groups = ["main"] files = [ {file = "uuid_utils-0.17.0-cp310-cp310-macosx_10_12_x86_64.macosx_11_0_arm64.macosx_10_12_universal2.whl", hash = "sha256:d2d9a63a9e6f2416ace8c109043a9280d6b34f34bb2e5421903e149403db40a6"}, {file = "uuid_utils-0.17.0-cp310-cp310-macosx_10_12_x86_64.whl", hash = "sha256:b776c7fc8755c7de06dd5a22b47c40ae84f67d13277ebb233cc84933ba4dcbcd"}, @@ -6764,7 +6744,7 @@ version = "3.8.1" description = "Python binding for xxHash" optional = false python-versions = ">=3.8" -groups = ["main", "optional"] +groups = ["main"] files = [ {file = "xxhash-3.8.1-cp310-cp310-macosx_10_9_x86_64.whl", hash = "sha256:27a9e475157f7315826118e3f3127909a0fe25f1b43d3d3be9c584f9d265f937"}, {file = "xxhash-3.8.1-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:9b2ce44bf8f4a1d01f418b3110ff8dff32fd3f3e836c0e06333c3725f243fa6c"}, @@ -7101,7 +7081,7 @@ version = "0.25.0" description = "Zstandard bindings for Python" optional = false python-versions = ">=3.9" -groups = ["main", "optional"] +groups = ["main"] files = [ {file = "zstandard-0.25.0-cp310-cp310-macosx_10_9_x86_64.whl", hash = "sha256:e59fdc271772f6686e01e1b3b74537259800f57e24280be3f29c8a0deb1904dd"}, {file = "zstandard-0.25.0-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:4d441506e9b372386a5271c64125f72d5df6d2a8e8a2a45a0ae09b03cb781ef7"}, @@ -7210,4 +7190,4 @@ cffi = ["cffi (>=1.17,<2.0) ; platform_python_implementation != \"PyPy\" and pyt [metadata] lock-version = "2.1" python-versions = "<3.14,>=3.10" -content-hash = "79a448728a85c44641ad9b79bfee561a6e8c744ef2a3ce41078613f33a69508b" +content-hash = "7df1d80ed7ee199fc0fa65aa1ca7ab0f5a18bb2ff60a8bd44669e514a4b1e10b" diff --git a/src/pyproject.toml b/src/pyproject.toml index 3191ece0..3aac15fb 100644 --- a/src/pyproject.toml +++ b/src/pyproject.toml @@ -74,7 +74,6 @@ en_core_web_sm = {url = "https://github.com/explosion/spacy-models/releases/down # requires running the Chroma server with trust_remote_code=true, which Sherpa # does not do. Bump once an upstream fixed release ships. chromadb = "^1.0.9" -langchain-chroma = "^0.2.5" boto3 = "^1.28.77" beautifulsoup4 = "4.15.0" diff --git a/src/sherpa_ai/connectors/scripts/query_chroma.py b/src/sherpa_ai/connectors/scripts/query_chroma.py index 5b34f5d5..0b3b551e 100644 --- a/src/sherpa_ai/connectors/scripts/query_chroma.py +++ b/src/sherpa_ai/connectors/scripts/query_chroma.py @@ -2,10 +2,10 @@ import json import uuid -import chromadb -from chromadb.config import Settings -from dotenv import load_dotenv -from langchain_openai import OpenAIEmbeddings +import chromadb +from chromadb.config import Settings +from chromadb.utils import embedding_functions +from dotenv import load_dotenv from loguru import logger @@ -16,29 +16,17 @@ def main(args): settings=Settings(allow_reset=True), ) - embedding_func = OpenAIEmbeddings() - try: - from langchain_chroma import Chroma - except ImportError: - raise ImportError( - "Could not import langchain_chroma python package. " - "This is needed in order to use Chroma. " - "Please install it with `pip install langchain-chroma`" - ) - chroma = Chroma( - client=client, - collection_name=args.chroma_index, - embedding_function=embedding_func, + embedding_func = embedding_functions.OpenAIEmbeddingFunction( + model_name="text-embedding-ada-002" + ) + collection = client.get_or_create_collection( + name=args.chroma_index, embedding_function=embedding_func ) query = input("Enter query: ") - results = chroma.similarity_search( - query=query, - number_of_results=5, - k=1 - ) + results = collection.query(query_texts=[query], n_results=1) - logger.info(results[0].page_content) + logger.info(results["documents"][0][0]) logger.info("Done! Chroma is up and running.") diff --git a/src/sherpa_ai/connectors/vectorstores.py b/src/sherpa_ai/connectors/vectorstores.py index 40d1d9fe..293f3236 100644 --- a/src/sherpa_ai/connectors/vectorstores.py +++ b/src/sherpa_ai/connectors/vectorstores.py @@ -1,5 +1,10 @@ import os +import uuid +from typing import Any, Iterable, List, Optional, Tuple, Type +from langchain_core.documents import Document +from langchain_core.embeddings import Embeddings +from langchain_core.vectorstores import VectorStore, VectorStoreRetriever from langchain_openai import OpenAIEmbeddings from langchain_text_splitters import CharacterTextSplitter from loguru import logger @@ -8,32 +13,385 @@ from sherpa_ai.utils import load_files -class LocalChromaStore: - """A local Chroma-based vector store. +class ConversationStore(VectorStore): + """A vector store for storing and retrieving conversation data. - This class extends the Chroma vector store to provide additional functionality - for working with local files. + This class provides methods to store conversation data in a vector database + and retrieve similar conversations based on queries. + + Attributes: + db: The underlying database connection. + namespace (str): The namespace for the vector store. + embeddings_func: The embedding function to use. + text_key (str): The key used to store the text in metadata. Example: - >>> from sherpa_ai.connectors.vectorstores import LocalChromaStore - >>> store = LocalChromaStore.from_folder("path/to/files", "api_key") - >>> results = store.similarity_search("query", k=5) + >>> from sherpa_ai.connectors.vectorstores import ConversationStore + >>> store = ConversationStore.from_index("my_namespace", "api_key", "my_index") + >>> store.add_text("This is a conversation", {"user": "user1"}) + >>> results = store.similarity_search("conversation", top_k=5) """ - - def __init__(self, *args, **kwargs): + def __init__(self, namespace, db, embeddings, text_key): + """Initialize a ConversationStore instance. + + Args: + namespace (str): The namespace for the vector store. + db: The database connection. + embeddings: The embedding function to use. + text_key (str): The key used to store the text in metadata. + + Example: + >>> from sherpa_ai.connectors.vectorstores import ConversationStore + >>> store = ConversationStore("my_namespace", db, embeddings, "text") + """ + self.db = db + self.namespace = namespace + self.embeddings_func = embeddings + self.text_key = text_key + + @classmethod + def from_index(cls, namespace, openai_api_key, index_name, text_key="text"): + """Create a ConversationStore from a Pinecone index. + + This method initializes a Pinecone client and creates a ConversationStore + instance connected to the specified index. + + Args: + namespace (str): The namespace for the vector store. + openai_api_key (str): The OpenAI API key. + index_name (str): The name of the Pinecone index. + text_key (str, optional): The key used to store the text in metadata. Defaults to "text". + + Returns: + ConversationStore: A new ConversationStore instance. + + Raises: + ImportError: If the pinecone-client package is not installed. + + Example: + >>> from sherpa_ai.connectors.vectorstores import ConversationStore + >>> store = ConversationStore.from_index("my_namespace", "api_key", "my_index") + """ try: - from langchain_chroma import Chroma + import pinecone except ImportError: raise ImportError( - "Could not import langchain_chroma python package. " - "This is needed in order to use LocalChromaStore. " - "Please install it with `pip install langchain-chroma`" + "Could not import pinecone-client python package. " + "This is needed in order to to use ConversationStore. " + "Please install it with `pip install pinecone-client`" + ) + + pinecone.init(api_key=cfg.PINECONE_API_KEY, environment=cfg.PINECONE_ENV) + logger.info(f"Loading index {index_name} from Pinecone") + index = pinecone.Index(index_name) + embedding = OpenAIEmbeddings(openai_api_key=openai_api_key) + return cls(namespace, index, embedding, text_key) + + def add_text(self, text: str, metadata={}) -> str: + """Add a single text to the vector store. + + This method embeds the text, adds it to the database with the provided metadata, + and returns the ID of the added text. + + Args: + text (str): The text to add. + metadata (dict, optional): Metadata to associate with the text. Defaults to {}. + + Returns: + str: The ID of the added text. + + Example: + >>> from sherpa_ai.connectors.vectorstores import ConversationStore + >>> store = ConversationStore.from_index("my_namespace", "api_key", "my_index") + >>> id = store.add_text("This is a conversation", {"user": "user1"}) + >>> print(id) + '123e4567-e89b-12d3-a456-426614174000' + """ + metadata[self.text_key] = text + id = str(uuid.uuid4()) + embedding = self.embeddings.embed_query(text) + doc = {"id": id, "values": embedding, "metadata": metadata} + self.db.upsert(vectors=[doc], namespace=self.namespace) + + return id + + @property + def embeddings(self) -> Optional[Embeddings]: + """Access the query embedding object if available.""" + return self.embeddings_func + + def add_texts(self, texts: Iterable[str], metadatas: List[dict]) -> List[str]: + """Add multiple texts to the vector store. + + This method adds each text with its corresponding metadata to the vector store. + + Args: + texts (Iterable[str]): The texts to add. + metadatas (List[dict]): The metadata for each text. + + Example: + >>> from sherpa_ai.connectors.vectorstores import ConversationStore + >>> store = ConversationStore.from_index("my_namespace", "api_key", "my_index") + >>> texts = ["Text 1", "Text 2"] + >>> metadatas = [{"user": "user1"}, {"user": "user2"}] + >>> store.add_texts(texts, metadatas) + """ + for text, metadata in zip(texts, metadatas): + self.add_text(text, metadata) + + def similarity_search( + self, + text: str, + top_k: int = 5, + filter: Optional[dict] = None, + threshold: float = 0.7, + ) -> list[Document]: + """Perform a similarity search in the vector store. + + This method searches for texts that are semantically similar to the query. + + Args: + text (str): The search query. + top_k (int, optional): The number of results to return. Defaults to 5. + filter (Optional[dict], optional): Filter criteria for the search. Defaults to None. + threshold (float, optional): The similarity threshold. Defaults to 0.7. + + Returns: + list[Document]: A list of documents that match the query. + + Example: + >>> from sherpa_ai.connectors.vectorstores import ConversationStore + >>> store = ConversationStore.from_index("my_namespace", "api_key", "my_index") + >>> results = store.similarity_search("What is machine learning?", top_k=5) + >>> for doc in results: + ... print(doc.page_content[:100]) + """ + query_embedding = self.embeddings.embed_query(text) + results = self.db.query( + [query_embedding], + top_k=top_k, + include_metadata=True, + namespace=self.namespace, + filter=filter, + ) + + docs = [] + for res in results["matches"]: + metadata = res["metadata"] + text = metadata.pop(self.text_key) + if res["score"] > threshold: + docs.append(Document(page_content=text, metadata=metadata)) + return docs + + def _similarity_search_with_relevance_scores( + self, + query: str, + k: int = 4, + **kwargs: Any, + ) -> List[Tuple[Document, float]]: + """Perform a similarity search and return documents with relevance scores. + + This method searches for texts that are semantically similar to the query + and returns them along with their relevance scores. + + Args: + query (str): The search query. + k (int, optional): The number of results to return. Defaults to 4. + **kwargs: Additional keyword arguments. + + Returns: + List[Tuple[Document, float]]: A list of tuples containing documents and their relevance scores. + + Example: + >>> from sherpa_ai.connectors.vectorstores import ConversationStore + >>> store = ConversationStore.from_index("my_namespace", "api_key", "my_index") + >>> results = store._similarity_search_with_relevance_scores("What is machine learning?") + >>> for doc, score in results: + ... print(f"Score: {score}, Content: {doc.page_content[:100]}") + """ + logger.debug("query", query) + query_embedding = self.embeddings.embed_query(query) + results = self.db.query( + [query_embedding], + top_k=k, + include_metadata=True, + namespace=self.namespace, + filter=kwargs.get("filter", None), + ) + + docs_with_score = [] + for res in results["matches"]: + metadata = res["metadata"] + text = metadata.pop(self.text_key) + docs_with_score.append( + (Document(page_content=text, metadata=metadata), res["score"]) ) - self._chroma = Chroma(*args, **kwargs) - - def __getattr__(self, name): - """Delegate attribute access to the underlying Chroma instance.""" - return getattr(self._chroma, name) + logger.debug(docs_with_score) + return docs_with_score + + @classmethod + def delete(cls, namespace, index_name): + """Delete all vectors in a namespace. + + This method deletes all vectors in the specified namespace of the Pinecone index. + + Args: + namespace (str): The namespace to delete. + index_name (str): The name of the Pinecone index. + + Returns: + The result of the delete operation. + + Raises: + ImportError: If the pinecone-client package is not installed. + + Example: + >>> from sherpa_ai.connectors.vectorstores import ConversationStore + >>> ConversationStore.delete("my_namespace", "my_index") + """ + try: + import pinecone + except ImportError: + raise ImportError( + "Could not import pinecone-client python package. " + "This is needed in order to to use ConversationStore. " + "Please install it with `pip install pinecone-client`" + ) + + + pinecone.init(api_key=cfg.PINECONE_API_KEY, environment=cfg.PINECONE_ENV) + index = pinecone.Index(index_name) + return index.delete(delete_all=True, namespace=namespace) + + @classmethod + def get_vector_retrieval( + cls, + namespace: str, + openai_api_key: str, + index_name: str, + search_type="similarity", + search_kwargs={}, + ) -> VectorStoreRetriever: + """Create a vector store retriever. + + This method creates a ConversationStore and returns a VectorStoreRetriever + for it. + + Args: + namespace (str): The namespace for the vector store. + openai_api_key (str): The OpenAI API key. + index_name (str): The name of the Pinecone index. + search_type (str, optional): The type of search to perform. Defaults to "similarity". + search_kwargs (dict, optional): Additional keyword arguments for the search. Defaults to {}. + + Returns: + VectorStoreRetriever: A retriever for the vector store. + + Example: + >>> from sherpa_ai.connectors.vectorstores import ConversationStore + >>> retriever = ConversationStore.get_vector_retrieval("my_namespace", "api_key", "my_index") + >>> results = retriever.get_relevant_documents("What is machine learning?") + """ + vectorstore = cls.from_index(namespace, openai_api_key, index_name) + retriever = VectorStoreRetriever( + vectorstore=vectorstore, + search_type=search_type, + search_kwargs=search_kwargs, + ) + return retriever + + @classmethod + def from_texts(cls, texts: List[str], embedding: Embeddings, metadatas: list[dict]): + """Create a ConversationStore from a list of texts. + + This method is not implemented for ConversationStore. + + Args: + texts (List[str]): The texts to add. + embedding (Embeddings): The embedding function to use. + metadatas (list[dict]): The metadata for each text. + + Raises: + NotImplementedError: This method is not implemented for ConversationStore. + """ + raise NotImplementedError("ConversationStore does not support from_texts") + + +class LocalChromaStore(VectorStore): + """A local Chroma-based vector store, backed directly by chromadb. + + chromadb is already a hard dependency of sherpa-ai (see chroma_vector_store.py), + so this talks to it directly instead of going through the langchain-chroma + wrapper package, which brings nothing this class needs. + + Example: + >>> from sherpa_ai.connectors.vectorstores import LocalChromaStore + >>> store = LocalChromaStore.from_folder("path/to/files", "api_key") + >>> results = store.similarity_search("query", k=5) + """ + + def __init__( + self, + collection_name: str = "langchain", + embedding_function: Optional[Embeddings] = None, + client: Optional[Any] = None, + ): + import chromadb + + self._embedding_function = embedding_function + self._client = client if client is not None else chromadb.EphemeralClient() + self._collection = self._client.get_or_create_collection(name=collection_name) + + @property + def embeddings(self) -> Optional[Embeddings]: + """Access the query embedding object if available.""" + return self._embedding_function + + def add_texts( + self, + texts: Iterable[str], + metadatas: Optional[List[dict]] = None, + **kwargs: Any, + ) -> List[str]: + """Embed and add texts to the underlying chromadb collection.""" + texts = list(texts) + ids = [str(uuid.uuid4()) for _ in texts] + embeddings = self._embedding_function.embed_documents(texts) + self._collection.add( + ids=ids, + embeddings=embeddings, + documents=texts, + metadatas=metadatas if metadatas else [{} for _ in texts], + ) + return ids + + def similarity_search(self, query: str, k: int = 4, **kwargs: Any) -> List[Document]: + """Perform a similarity search in the chromadb collection.""" + query_embedding = self._embedding_function.embed_query(query) + results = self._collection.query(query_embeddings=[query_embedding], n_results=k) + + docs = [] + documents = results.get("documents") or [[]] + metadatas = results.get("metadatas") or [[]] + for text, metadata in zip(documents[0], metadatas[0]): + docs.append(Document(page_content=text, metadata=metadata or {})) + return docs + + @classmethod + def from_texts( + cls, + texts: List[str], + embedding: Embeddings, + metadatas: Optional[List[dict]] = None, + index_name: str = "langchain", + **kwargs: Any, + ) -> "LocalChromaStore": + """Create a LocalChromaStore from a list of texts.""" + store = cls(collection_name=index_name, embedding_function=embedding) + if texts: + store.add_texts(texts, metadatas) + return store + @classmethod def from_folder(cls, file_path, openai_api_key, index_name="chroma"): """Create a Chroma DB from a folder of files. @@ -97,37 +455,19 @@ def configure_chroma(host: str, port: int, index_name: str, openai_api_key: str) "This is needed in order to to use Chroma. " "Please install it with `pip install chromadb" ) - - try: - from langchain_chroma import Chroma - except ImportError: - raise ImportError( - "Could not import langchain_chroma python package. " - "This is needed in order to use Chroma. " - "Please install it with `pip install langchain-chroma`" - ) + client = chromadb.HttpClient(host=cfg.CHROMA_HOST, port=cfg.CHROMA_PORT) embeddings = OpenAIEmbeddings(openai_api_key=openai_api_key) - chroma = Chroma( - client=client, collection_name=cfg.CHROMA_INDEX, embedding_function=embeddings + return LocalChromaStore( + collection_name=cfg.CHROMA_INDEX, embedding_function=embeddings, client=client ) - return chroma - - -def _is_chroma_available(): - """Check if langchain_chroma is available.""" - try: - import langchain_chroma - return True - except ImportError: - return False def get_vectordb(): """Get a vector database retriever based on configuration. This function returns a vector database retriever based on the configuration - in the config module. It supports Chroma and local ChromaDB. + in the config module. It supports Pinecone, Chroma, and local ChromaDB. Returns: VectorStoreRetriever: A retriever for the vector store. @@ -137,20 +477,19 @@ def get_vectordb(): >>> retriever = get_vectordb() >>> results = retriever.get_relevant_documents("What is machine learning?") """ - if cfg.VECTORDB == "chroma": + if cfg.VECTORDB == "pinecone": + return ConversationStore.get_vector_retrieval( + cfg.PINECONE_NAMESPACE, + cfg.OPENAI_API_KEY, + index_name=cfg.PINECONE_INDEX, + search_type="similarity_score_threshold", + search_kwargs={"score_threshold": 0.0}, + ) + elif cfg.VECTORDB == "chroma": return configure_chroma( cfg.CHROMA_HOST, cfg.CHROMA_PORT, cfg.CHROMA_INDEX, cfg.OPENAI_API_KEY ).as_retriever() else: - # Check if langchain_chroma is available before trying to use it - if not _is_chroma_available(): - raise ImportError( - "Could not import langchain_chroma python package. " - "This is needed in order to use the default vector store. " - "Please install it with `pip install langchain-chroma` or " - "configure a different vector database (chroma) in your environment." - ) - if os.path.exists("files"): return LocalChromaStore.from_folder( "files", cfg.OPENAI_API_KEY diff --git a/src/tests/unit_tests/actions/test_context_search.py b/src/tests/unit_tests/actions/test_context_search.py index 24c3100b..6ff8474e 100644 --- a/src/tests/unit_tests/actions/test_context_search.py +++ b/src/tests/unit_tests/actions/test_context_search.py @@ -6,15 +6,15 @@ def _is_chroma_available(): - """Check if langchain_chroma is available.""" + """Check if chromadb is available.""" try: - import langchain_chroma + import chromadb return True except ImportError: return False -# Only import ContextSearch if langchain_chroma is available +# Only import ContextSearch if chromadb is available if _is_chroma_available(): from sherpa_ai.actions.context_search import ContextSearch @@ -36,7 +36,7 @@ def mock_context_search(external_api): yield -@pytest.mark.skipif(not _is_chroma_available(), reason="langchain_chroma not available") +@pytest.mark.skipif(not _is_chroma_available(), reason="chromadb not available") def test_context_search_succeeds(get_llm, mock_context_search): # noqa: F811 role_description = ( "The programmer receives requirements about a program and write it" diff --git a/src/tests/unit_tests/connectors/test_local_chroma_store.py b/src/tests/unit_tests/connectors/test_local_chroma_store.py new file mode 100644 index 00000000..064606d2 --- /dev/null +++ b/src/tests/unit_tests/connectors/test_local_chroma_store.py @@ -0,0 +1,46 @@ +import pytest + +from sherpa_ai.connectors.vectorstores import LocalChromaStore + + +class FakeEmbeddings: + """Deterministic stand-in for OpenAIEmbeddings: each text maps to a + fixed-size vector based on its length, so unrelated texts don't collide.""" + + def embed_documents(self, texts): + return [self.embed_query(t) for t in texts] + + def embed_query(self, text): + return [float(len(text)), float(sum(ord(c) for c in text) % 97)] + + +@pytest.fixture +def store(): + return LocalChromaStore(collection_name="test", embedding_function=FakeEmbeddings()) + + +def test_add_texts_and_similarity_search_round_trip(store): + store.add_texts( + ["sherpa helps you climb mountains", "bananas are yellow"], + metadatas=[{"topic": "sherpa"}, {"topic": "fruit"}], + ) + + results = store.similarity_search("sherpa helps you climb mountains", k=1) + + assert len(results) == 1 + assert results[0].page_content == "sherpa helps you climb mountains" + assert results[0].metadata["topic"] == "sherpa" + + +def test_from_texts_builds_a_queryable_store(): + store = LocalChromaStore.from_texts( + ["sherpa helps you climb mountains", "bananas are yellow"], + embedding=FakeEmbeddings(), + metadatas=[{"topic": "sherpa"}, {"topic": "fruit"}], + index_name="test_from_texts", + ) + + results = store.similarity_search("bananas are yellow", k=1) + + assert len(results) == 1 + assert results[0].page_content == "bananas are yellow"