Compare commits
12 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 3fe3b0320a | |||
| 75896b647f | |||
| 446d0975aa | |||
| b7d365119c | |||
| 2d9fbd4e49 | |||
| 22e14b5e65 | |||
| 1a654beea4 | |||
| f50f8a444a | |||
| 3cc3a0058d | |||
| ae473b5e3c | |||
| efb7e31565 | |||
| 069d265338 |
@@ -2,7 +2,7 @@
|
||||
<Card title="Talk to founders" icon="calendar" href="https://cal.com/taranjeetio/ec">
|
||||
Schedule a call
|
||||
</Card>
|
||||
<Card title="Slack" icon="slack" href="https://join.slack.com/t/embedchain/shared_invite/zt-22uwz3c46-Zg7cIh5rOBteT_xe1jwLDw" color="#4A154B">
|
||||
<Card title="Slack" icon="slack" href="https://embedchain.ai/slack" color="#4A154B">
|
||||
Join our slack community
|
||||
</Card>
|
||||
<Card title="Discord" icon="discord" href="https://discord.gg/6PzXDgEjG5" color="#7289DA">
|
||||
|
||||
@@ -4,7 +4,7 @@
|
||||
<Card title="Google Form" icon="file" href="https://forms.gle/NDRCKsRpUHsz2Wcm8" color="#7387d0">
|
||||
Fill out this form
|
||||
</Card>
|
||||
<Card title="Slack" icon="slack" href="https://join.slack.com/t/embedchain/shared_invite/zt-22uwz3c46-Zg7cIh5rOBteT_xe1jwLDw" color="#4A154B">
|
||||
<Card title="Slack" icon="slack" href="https://embedchain.ai/slack" color="#4A154B">
|
||||
Let us know on our slack community
|
||||
</Card>
|
||||
<Card title="Discord" icon="discord" href="https://discord.gg/6PzXDgEjG5" color="#7289DA">
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
<p>If you can't find the specific LLM you need, no need to fret. We're continuously expanding our support for additional LLMs, and you can help us prioritize by opening an issue on our GitHub or simply reaching out to us on our Slack or Discord community.</p>
|
||||
|
||||
<CardGroup cols={2}>
|
||||
<Card title="Slack" icon="slack" href="https://join.slack.com/t/embedchain/shared_invite/zt-22uwz3c46-Zg7cIh5rOBteT_xe1jwLDw" color="#4A154B">
|
||||
<Card title="Slack" icon="slack" href="https://embedchain.ai/slack" color="#4A154B">
|
||||
Let us know on our slack community
|
||||
</Card>
|
||||
<Card title="Discord" icon="discord" href="https://discord.gg/6PzXDgEjG5" color="#7289DA">
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
<p>If you can't find the specific vector database, please feel free to request through one of the following channels and help us prioritize.</p>
|
||||
|
||||
<CardGroup cols={2}>
|
||||
<Card title="Slack" icon="slack" href="https://join.slack.com/t/embedchain/shared_invite/zt-22uwz3c46-Zg7cIh5rOBteT_xe1jwLDw" color="#4A154B">
|
||||
<Card title="Slack" icon="slack" href="https://embedchain.ai/slack" color="#4A154B">
|
||||
Let us know on our slack community
|
||||
</Card>
|
||||
<Card title="Discord" icon="discord" href="https://discord.gg/6PzXDgEjG5" color="#7289DA">
|
||||
|
||||
@@ -200,9 +200,10 @@ Alright, let's dive into what each key means in the yaml config above:
|
||||
- `stream` (Boolean): Controls if the response is streamed back to the user (set to false).
|
||||
- `prompt` (String): A prompt for the model to follow when generating responses, requires `$context` and `$query` variables.
|
||||
- `system_prompt` (String): A system prompt for the model to follow when generating responses, in this case, it's set to the style of William Shakespeare.
|
||||
- `stream` (Boolean): Controls if the response is streamed back to the user (set to false).
|
||||
- `stream` (Boolean): Controls if the response is streamed back to the user (set to false).
|
||||
- `number_documents` (Integer): Number of documents to pull from the vectordb as context, defaults to 1
|
||||
- `api_key` (String): The API key for the language model.
|
||||
- `model_kwargs` (Dict): Keyword arguments to pass to the language model. Used for `aws_bedrock` provider, since it requires different arguments for each model.
|
||||
3. `vectordb` Section:
|
||||
- `provider` (String): The provider for the vector database, set to 'chroma'. You can find the full list of vector database providers in [our docs](/components/vector-databases).
|
||||
- `config`:
|
||||
|
||||
@@ -2,9 +2,33 @@
|
||||
title: 🗑 delete
|
||||
---
|
||||
|
||||
## Delete Document
|
||||
|
||||
`delete()` method allows you to delete a document previously added to the app.
|
||||
|
||||
### Usage
|
||||
|
||||
```python
|
||||
from embedchain import App
|
||||
|
||||
app = App()
|
||||
|
||||
forbes_doc_id = app.add("https://www.forbes.com/profile/elon-musk")
|
||||
wiki_doc_id = app.add("https://en.wikipedia.org/wiki/Elon_Musk")
|
||||
|
||||
app.delete(forbes_doc_id) # deletes the forbes document
|
||||
```
|
||||
|
||||
<Note>
|
||||
If you do not have the document id, you can use `app.db.get()` method to get the document and extract the `hash` key from `metadatas` dictionary object, which serves as the document id.
|
||||
</Note>
|
||||
|
||||
|
||||
## Delete Chat Session History
|
||||
|
||||
`delete_session_chat_history()` method allows you to delete all previous messages in a chat history.
|
||||
|
||||
## Usage
|
||||
### Usage
|
||||
|
||||
```python
|
||||
from embedchain import App
|
||||
@@ -16,4 +40,9 @@ app.add("https://www.forbes.com/profile/elon-musk")
|
||||
app.chat("What is the net worth of Elon Musk?")
|
||||
|
||||
app.delete_session_chat_history()
|
||||
```
|
||||
```
|
||||
|
||||
<Note>
|
||||
`delete_session_chat_history(session_id="session_1")` method also accepts `session_id` optional param for deleting chat history of a specific session.
|
||||
It assumes the default session if no `session_id` is provided.
|
||||
</Note>
|
||||
@@ -0,0 +1,33 @@
|
||||
---
|
||||
title: 📄 get
|
||||
---
|
||||
|
||||
## Get data sources
|
||||
|
||||
`get_data_sources()` returns a list of all the data sources added in the app.
|
||||
|
||||
|
||||
### Usage
|
||||
|
||||
```python
|
||||
from embedchain import App
|
||||
|
||||
app = App()
|
||||
|
||||
app.add("https://www.forbes.com/profile/elon-musk")
|
||||
app.add("https://en.wikipedia.org/wiki/Elon_Musk")
|
||||
|
||||
data_sources = app.get_data_sources()
|
||||
# [
|
||||
# {
|
||||
# 'data_type': 'web_page',
|
||||
# 'data_value': 'https://en.wikipedia.org/wiki/Elon_Musk',
|
||||
# 'metadata': 'null'
|
||||
# },
|
||||
# {
|
||||
# 'data_type': 'web_page',
|
||||
# 'data_value': 'https://www.forbes.com/profile/elon-musk',
|
||||
# 'metadata': 'null'
|
||||
# }
|
||||
# ]
|
||||
```
|
||||
@@ -8,7 +8,7 @@ We believe in building a vibrant and supportive community around embedchain. The
|
||||
<Card title="Twitter" icon="twitter" href="https://twitter.com/embedchain">
|
||||
Follow us on Twitter
|
||||
</Card>
|
||||
<Card title="Slack" icon="slack" href="https://join.slack.com/t/embedchain/shared_invite/zt-22uwz3c46-Zg7cIh5rOBteT_xe1jwLDw" color="#4A154B">
|
||||
<Card title="Slack" icon="slack" href="https://embedchain.ai/slack" color="#4A154B">
|
||||
Join our slack community
|
||||
</Card>
|
||||
<Card title="Discord" icon="discord" href="https://discord.gg/6PzXDgEjG5" color="#7289DA">
|
||||
|
||||
@@ -21,6 +21,7 @@ Embedchain comes with built-in support for various popular large language models
|
||||
<Card title="Llama2" href="#llama2"></Card>
|
||||
<Card title="Vertex AI" href="#vertex-ai"></Card>
|
||||
<Card title="Mistral AI" href="#mistral-ai"></Card>
|
||||
<Card title="AWS Bedrock" href="#aws-bedrock"></Card>
|
||||
</CardGroup>
|
||||
|
||||
## OpenAI
|
||||
@@ -627,11 +628,8 @@ llm:
|
||||
Obtain the Mistral AI api key from their [console](https://console.mistral.ai/).
|
||||
|
||||
<CodeGroup>
|
||||
|
||||
```python main.py
|
||||
import os
|
||||
from embedchain import App
|
||||
|
||||
|
||||
```python main.py
|
||||
os.environ["MISTRAL_API_KEY"] = "xxx"
|
||||
|
||||
app = App.from_config(config_path="config.yaml")
|
||||
@@ -663,5 +661,45 @@ embedder:
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
|
||||
## AWS Bedrock
|
||||
|
||||
### Setup
|
||||
- Before using the AWS Bedrock LLM, make sure you have the appropriate model access from [Bedrock Console](https://us-east-1.console.aws.amazon.com/bedrock/home?region=us-east-1#/modelaccess).
|
||||
- You will also need `AWS_ACCESS_KEY_ID` and `AWS_SECRET_ACCESS_KEY` to authenticate the API with AWS. You can find these in your [AWS Console](https://us-east-1.console.aws.amazon.com/iam/home?region=us-east-1#/users).
|
||||
|
||||
|
||||
### Usage
|
||||
|
||||
<CodeGroup>
|
||||
|
||||
```python main.py
|
||||
import os
|
||||
from embedchain import App
|
||||
|
||||
os.environ["AWS_ACCESS_KEY_ID"] = "xxx"
|
||||
os.environ["AWS_SECRET_ACCESS_KEY"] = "xxx"
|
||||
|
||||
app = App.from_config(config_path="config.yaml")
|
||||
```
|
||||
|
||||
```yaml config.yaml
|
||||
llm:
|
||||
provider: aws_bedrock
|
||||
config:
|
||||
model: amazon.titan-text-express-v1
|
||||
# check notes below for model_kwargs
|
||||
model_kwargs:
|
||||
temperature: 0.5
|
||||
topP: 1
|
||||
maxTokenCount: 1000
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
<br />
|
||||
<Note>
|
||||
The model arguments are different for each providers. Please refer to the [AWS Bedrock Documentation](https://us-east-1.console.aws.amazon.com/bedrock/home?region=us-east-1#/providers) to find the appropriate arguments for your model.
|
||||
</Note>
|
||||
|
||||
<br/ >
|
||||
<Snippet file="missing-llm-tip.mdx" />
|
||||
|
||||
@@ -189,6 +189,11 @@ vectordb:
|
||||
|
||||
</CodeGroup>
|
||||
|
||||
<br />
|
||||
<Note>
|
||||
You can optionally provide `index_name` as a config param in yaml file to specify the index name. If not provided, the index name will be `{collection_name}-{vector_dimension}`.
|
||||
</Note>
|
||||
|
||||
## Qdrant
|
||||
|
||||
In order to use Qdrant as a vector database, set the environment variables `QDRANT_URL` and `QDRANT_API_KEY` which you can find on [Qdrant Dashboard](https://cloud.qdrant.io/).
|
||||
|
||||
+4
-3
@@ -207,10 +207,11 @@
|
||||
"api-reference/app/query",
|
||||
"api-reference/app/chat",
|
||||
"api-reference/app/search",
|
||||
"api-reference/app/get",
|
||||
"api-reference/app/evaluate",
|
||||
"api-reference/app/deploy",
|
||||
"api-reference/app/reset",
|
||||
"api-reference/app/delete",
|
||||
"api-reference/app/evaluate"
|
||||
"api-reference/app/delete"
|
||||
]
|
||||
},
|
||||
"api-reference/store/openai-assistant",
|
||||
@@ -238,7 +239,7 @@
|
||||
"footerSocials": {
|
||||
"website": "https://embedchain.ai",
|
||||
"github": "https://github.com/embedchain/embedchain",
|
||||
"slack": "https://join.slack.com/t/embedchain/shared_invite/zt-22uwz3c46-Zg7cIh5rOBteT_xe1jwLDw",
|
||||
"slack": "https://embedchain.ai/slack",
|
||||
"discord": "https://discord.gg/6PzXDgEjG5",
|
||||
"twitter": "https://twitter.com/embedchain",
|
||||
"linkedin": "https://www.linkedin.com/company/embedchain"
|
||||
|
||||
@@ -9,6 +9,7 @@ class PineconeDBConfig(BaseVectorDbConfig):
|
||||
def __init__(
|
||||
self,
|
||||
collection_name: Optional[str] = None,
|
||||
index_name: Optional[str] = None,
|
||||
dir: Optional[str] = None,
|
||||
vector_dimension: int = 1536,
|
||||
metric: Optional[str] = "cosine",
|
||||
@@ -17,4 +18,5 @@ class PineconeDBConfig(BaseVectorDbConfig):
|
||||
self.metric = metric
|
||||
self.vector_dimension = vector_dimension
|
||||
self.extra_params = extra_params
|
||||
self.index_name = index_name or f"{collection_name}-{vector_dimension}".lower().replace("_", "-")
|
||||
super().__init__(collection_name=collection_name, dir=dir)
|
||||
|
||||
@@ -674,3 +674,15 @@ class EmbedChain(JSONSerializable):
|
||||
def delete_all_chat_history(self, app_id: str):
|
||||
self.llm.memory.delete(app_id=app_id)
|
||||
self.llm.update_history(app_id=app_id)
|
||||
|
||||
def delete(self, source_id: str):
|
||||
"""
|
||||
Deletes the data from the database.
|
||||
:param source_hash: The hash of the source.
|
||||
:type source_hash: str
|
||||
"""
|
||||
self.db.delete(where={"hash": source_id})
|
||||
logging.info(f"Successfully deleted {source_id}")
|
||||
# Send anonymous telemetry
|
||||
if self.config.collect_metrics:
|
||||
self.telemetry.capture(event_name="delete", properties=self._telemetry_props)
|
||||
|
||||
@@ -21,6 +21,7 @@ class LlmFactory:
|
||||
"openai": "embedchain.llm.openai.OpenAILlm",
|
||||
"vertexai": "embedchain.llm.vertex_ai.VertexAILlm",
|
||||
"google": "embedchain.llm.google.GoogleLlm",
|
||||
"aws_bedrock": "embedchain.llm.aws_bedrock.AWSBedrockLlm",
|
||||
"mistralai": "embedchain.llm.mistralai.MistralAILlm",
|
||||
}
|
||||
provider_to_config_class = {
|
||||
|
||||
@@ -0,0 +1,48 @@
|
||||
from typing import Optional
|
||||
|
||||
from langchain.llms import Bedrock
|
||||
|
||||
from embedchain.config import BaseLlmConfig
|
||||
from embedchain.helpers.json_serializable import register_deserializable
|
||||
from embedchain.llm.base import BaseLlm
|
||||
|
||||
|
||||
@register_deserializable
|
||||
class AWSBedrockLlm(BaseLlm):
|
||||
def __init__(self, config: Optional[BaseLlmConfig] = None):
|
||||
super().__init__(config)
|
||||
|
||||
def get_llm_model_answer(self, prompt) -> str:
|
||||
response = self._get_answer(prompt, self.config)
|
||||
return response
|
||||
|
||||
def _get_answer(self, prompt: str, config: BaseLlmConfig) -> str:
|
||||
try:
|
||||
import boto3
|
||||
except ModuleNotFoundError:
|
||||
raise ModuleNotFoundError(
|
||||
"The required dependencies for AWSBedrock are not installed."
|
||||
'Please install with `pip install --upgrade "embedchain[aws-bedrock]"`'
|
||||
) from None
|
||||
|
||||
self.boto_client = boto3.client("bedrock-runtime", "us-west-2")
|
||||
|
||||
kwargs = {
|
||||
"model_id": config.model or "amazon.titan-text-express-v1",
|
||||
"client": self.boto_client,
|
||||
"model_kwargs": config.model_kwargs
|
||||
or {
|
||||
"temperature": config.temperature,
|
||||
},
|
||||
}
|
||||
|
||||
if config.stream:
|
||||
from langchain.callbacks.streaming_stdout import \
|
||||
StreamingStdOutCallbackHandler
|
||||
|
||||
callbacks = [StreamingStdOutCallbackHandler()]
|
||||
llm = Bedrock(**kwargs, streaming=config.stream, callbacks=callbacks)
|
||||
else:
|
||||
llm = Bedrock(**kwargs)
|
||||
|
||||
return llm(prompt)
|
||||
@@ -406,6 +406,7 @@ def validate_config(config_data):
|
||||
"llama2",
|
||||
"vertexai",
|
||||
"google",
|
||||
"aws_bedrock",
|
||||
"mistralai",
|
||||
),
|
||||
Optional("config"): {
|
||||
@@ -423,6 +424,7 @@ def validate_config(config_data):
|
||||
Optional("query_type"): str,
|
||||
Optional("api_key"): str,
|
||||
Optional("endpoint"): str,
|
||||
Optional("model_kwargs"): dict,
|
||||
},
|
||||
},
|
||||
Optional("vectordb"): {
|
||||
|
||||
@@ -75,3 +75,8 @@ class BaseVectorDB(JSONSerializable):
|
||||
:type name: str
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def delete(self):
|
||||
"""Delete from database."""
|
||||
|
||||
raise NotImplementedError
|
||||
|
||||
@@ -99,14 +99,24 @@ class ElasticsearchDB(BaseVectorDB):
|
||||
query = {"bool": {"must": [{"ids": {"values": ids}}]}}
|
||||
else:
|
||||
query = {"bool": {"must": []}}
|
||||
if "app_id" in where:
|
||||
app_id = where["app_id"]
|
||||
query["bool"]["must"].append({"term": {"metadata.app_id": app_id}})
|
||||
|
||||
response = self.client.search(index=self._get_index(), query=query, _source=False, size=limit)
|
||||
if where:
|
||||
for key, value in where.items():
|
||||
query["bool"]["must"].append({"term": {f"metadata.{key}.keyword": value}})
|
||||
|
||||
response = self.client.search(index=self._get_index(), query=query, _source=True, size=limit)
|
||||
docs = response["hits"]["hits"]
|
||||
ids = [doc["_id"] for doc in docs]
|
||||
return {"ids": set(ids)}
|
||||
doc_ids = [doc["_source"]["metadata"]["doc_id"] for doc in docs]
|
||||
|
||||
# Result is modified for compatibility with other vector databases
|
||||
# TODO: Add method in vector database to return result in a standard format
|
||||
result = {"ids": ids, "metadatas": []}
|
||||
|
||||
for doc_id in doc_ids:
|
||||
result["metadatas"].append({"doc_id": doc_id})
|
||||
|
||||
return result
|
||||
|
||||
def add(
|
||||
self,
|
||||
@@ -186,9 +196,11 @@ class ElasticsearchDB(BaseVectorDB):
|
||||
},
|
||||
}
|
||||
}
|
||||
if "app_id" in where:
|
||||
app_id = where["app_id"]
|
||||
query["script_score"]["query"] = {"match": {"metadata.app_id": app_id}}
|
||||
|
||||
if where:
|
||||
for key, value in where.items():
|
||||
query["script_score"]["query"]["bool"]["must"].append({"term": {f"metadata.{key}.keyword": value}})
|
||||
|
||||
_source = ["text", "metadata"]
|
||||
response = self.client.search(index=self._get_index(), query=query, _source=_source, size=n_results)
|
||||
docs = response["hits"]["hits"]
|
||||
@@ -244,3 +256,11 @@ class ElasticsearchDB(BaseVectorDB):
|
||||
# NOTE: The method is preferred to an attribute, because if collection name changes,
|
||||
# it's always up-to-date.
|
||||
return f"{self.config.collection_name}_{self.embedder.vector_dimension}".lower()
|
||||
|
||||
def delete(self, where):
|
||||
"""Delete documents from the database."""
|
||||
query = {"query": {"bool": {"must": []}}}
|
||||
for key, value in where.items():
|
||||
query["query"]["bool"]["must"].append({"term": {f"metadata.{key}.keyword": value}})
|
||||
self.client.delete_by_query(index=self._get_index(), body=query)
|
||||
self.client.indices.refresh(index=self._get_index())
|
||||
|
||||
@@ -96,9 +96,9 @@ class OpenSearchDB(BaseVectorDB):
|
||||
else:
|
||||
query["query"] = {"bool": {"must": []}}
|
||||
|
||||
if "app_id" in where:
|
||||
app_id = where["app_id"]
|
||||
query["query"]["bool"]["must"].append({"term": {"metadata.app_id.keyword": app_id}})
|
||||
if where:
|
||||
for key, value in where.items():
|
||||
query["bool"]["must"].append({"term": {f"metadata.{key}.keyword": value}})
|
||||
|
||||
# OpenSearch syntax is different from Elasticsearch
|
||||
response = self.client.search(index=self._get_index(), body=query, _source=True, size=limit)
|
||||
@@ -176,9 +176,11 @@ class OpenSearchDB(BaseVectorDB):
|
||||
)
|
||||
|
||||
pre_filter = {"match_all": {}} # default
|
||||
if "app_id" in where:
|
||||
app_id = where["app_id"]
|
||||
pre_filter = {"bool": {"must": [{"term": {"metadata.app_id.keyword": app_id}}]}}
|
||||
if len(where) > 0:
|
||||
pre_filter = {"bool": {"must": []}}
|
||||
for key, value in where.items():
|
||||
pre_filter["bool"]["must"].append({"term": {f"metadata.{key}.keyword": value}})
|
||||
|
||||
docs = docsearch.similarity_search_with_score(
|
||||
input_query,
|
||||
search_type="script_scoring",
|
||||
@@ -236,10 +238,9 @@ class OpenSearchDB(BaseVectorDB):
|
||||
|
||||
def delete(self, where):
|
||||
"""Deletes a document from the OpenSearch index"""
|
||||
if "doc_id" not in where:
|
||||
raise ValueError("doc_id is required to delete a document")
|
||||
|
||||
query = {"query": {"bool": {"must": [{"term": {"metadata.doc_id": where["doc_id"]}}]}}}
|
||||
query = {"query": {"bool": {"must": []}}}
|
||||
for key, value in where.items():
|
||||
query["query"]["bool"]["must"].append({"term": {f"metadata.{key}.keyword": value}})
|
||||
self.client.delete_by_query(index=self._get_index(), body=query)
|
||||
|
||||
def _get_index(self) -> str:
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import logging
|
||||
import os
|
||||
from typing import Optional, Union
|
||||
|
||||
@@ -52,20 +53,21 @@ class PineconeDB(BaseVectorDB):
|
||||
if not self.embedder:
|
||||
raise ValueError("Embedder not set. Please set an embedder with `set_embedder` before initialization.")
|
||||
|
||||
# Loads the Pinecone index or creates it if not present.
|
||||
def _setup_pinecone_index(self):
|
||||
"""
|
||||
Loads the Pinecone index or creates it if not present.
|
||||
"""
|
||||
pinecone.init(
|
||||
api_key=os.environ.get("PINECONE_API_KEY"),
|
||||
environment=os.environ.get("PINECONE_ENV"),
|
||||
**self.config.extra_params,
|
||||
)
|
||||
self.index_name = self._get_index_name()
|
||||
indexes = pinecone.list_indexes()
|
||||
if indexes is None or self.index_name not in indexes:
|
||||
if indexes is None or self.config.index_name not in indexes:
|
||||
pinecone.create_index(
|
||||
name=self.index_name, metric=self.config.metric, dimension=self.config.vector_dimension
|
||||
name=self.config.index_name, metric=self.config.metric, dimension=self.config.vector_dimension
|
||||
)
|
||||
return pinecone.Index(self.index_name)
|
||||
return pinecone.Index(self.config.index_name)
|
||||
|
||||
def get(self, ids: Optional[list[str]] = None, where: Optional[dict[str, any]] = None, limit: Optional[int] = None):
|
||||
"""
|
||||
@@ -79,12 +81,20 @@ class PineconeDB(BaseVectorDB):
|
||||
:rtype: Set[str]
|
||||
"""
|
||||
existing_ids = list()
|
||||
metadatas = []
|
||||
|
||||
if ids is not None:
|
||||
for i in range(0, len(ids), 1000):
|
||||
result = self.client.fetch(ids=ids[i : i + 1000])
|
||||
batch_existing_ids = list(result.get("vectors").keys())
|
||||
vectors = result.get("vectors")
|
||||
batch_existing_ids = list(vectors.keys())
|
||||
existing_ids.extend(batch_existing_ids)
|
||||
return {"ids": existing_ids}
|
||||
metadatas.extend([vectors.get(ids).get("metadata") for ids in batch_existing_ids])
|
||||
|
||||
if where is not None:
|
||||
logging.warning("Filtering is not supported by Pinecone")
|
||||
|
||||
return {"ids": existing_ids, "metadatas": metadatas}
|
||||
|
||||
def add(
|
||||
self,
|
||||
@@ -114,7 +124,7 @@ class PineconeDB(BaseVectorDB):
|
||||
}
|
||||
)
|
||||
|
||||
for chunk in chunks(docs, self.BATCH_SIZE, desc="Adding chunks in batches..."):
|
||||
for chunk in chunks(docs, self.BATCH_SIZE, desc="Adding chunks in batches"):
|
||||
self.client.upsert(chunk, **kwargs)
|
||||
|
||||
def query(
|
||||
@@ -140,7 +150,10 @@ class PineconeDB(BaseVectorDB):
|
||||
:rtype: list[str], if citations=False, otherwise list[tuple[str, str, str]]
|
||||
"""
|
||||
query_vector = self.embedder.embedding_fn([input_query])[0]
|
||||
data = self.client.query(vector=query_vector, filter=where, top_k=n_results, include_metadata=True, **kwargs)
|
||||
query_filter = self._generate_filter(where)
|
||||
data = self.client.query(
|
||||
vector=query_vector, filter=query_filter, top_k=n_results, include_metadata=True, **kwargs
|
||||
)
|
||||
contexts = []
|
||||
for doc in data["matches"]:
|
||||
metadata = doc["metadata"]
|
||||
@@ -181,14 +194,26 @@ class PineconeDB(BaseVectorDB):
|
||||
Resets the database. Deletes all embeddings irreversibly.
|
||||
"""
|
||||
# Delete all data from the database
|
||||
pinecone.delete_index(self.index_name)
|
||||
pinecone.delete_index(self.config.index_name)
|
||||
self._setup_pinecone_index()
|
||||
|
||||
# Pinecone only allows alphanumeric characters and "-" in the index name
|
||||
def _get_index_name(self) -> str:
|
||||
"""Get the Pinecone index for a collection
|
||||
@staticmethod
|
||||
def _generate_filter(where: dict):
|
||||
query = {}
|
||||
for k, v in where.items():
|
||||
query[k] = {"$eq": v}
|
||||
return query
|
||||
|
||||
:return: Pinecone index
|
||||
:rtype: str
|
||||
def delete(self, where: dict):
|
||||
"""Delete from database.
|
||||
:param ids: list of ids to delete
|
||||
:type ids: list[str]
|
||||
"""
|
||||
return f"{self.config.collection_name}-{self.config.vector_dimension}".lower().replace("_", "-")
|
||||
# Deleting with filters is not supported for `starter` index type.
|
||||
# Follow `https://docs.pinecone.io/docs/metadata-filtering#deleting-vectors-by-metadata-filter` for more details
|
||||
db_filter = self._generate_filter(where)
|
||||
try:
|
||||
self.client.delete(filter=db_filter)
|
||||
except Exception as e:
|
||||
print(f"Failed to delete from Pinecone: {e}")
|
||||
return
|
||||
|
||||
@@ -11,6 +11,8 @@ try:
|
||||
except ImportError:
|
||||
raise ImportError("Qdrant requires extra dependencies. Install with `pip install embedchain[qdrant]`") from None
|
||||
|
||||
from tqdm import tqdm
|
||||
|
||||
from embedchain.config.vectordb.qdrant import QdrantDBConfig
|
||||
from embedchain.vectordb.base import BaseVectorDB
|
||||
|
||||
@@ -48,7 +50,6 @@ class QdrantDB(BaseVectorDB):
|
||||
raise ValueError("Embedder not set. Please set an embedder with `set_embedder` before initialization.")
|
||||
|
||||
self.collection_name = self._get_or_create_collection()
|
||||
self.metadata_keys = {"data_type", "doc_id", "url", "hash", "app_id", "text"}
|
||||
all_collections = self.client.get_collections()
|
||||
collection_names = [collection.name for collection in all_collections.collections]
|
||||
if self.collection_name not in collection_names:
|
||||
@@ -82,21 +83,23 @@ class QdrantDB(BaseVectorDB):
|
||||
:return: All the existing IDs
|
||||
:rtype: Set[str]
|
||||
"""
|
||||
if ids is None or len(ids) == 0:
|
||||
return {"ids": []}
|
||||
|
||||
keys = set(where.keys() if where is not None else set())
|
||||
|
||||
qdrant_must_filters = [
|
||||
models.FieldCondition(
|
||||
key="identifier",
|
||||
match=models.MatchAny(
|
||||
any=ids,
|
||||
),
|
||||
qdrant_must_filters = []
|
||||
|
||||
if ids:
|
||||
qdrant_must_filters.append(
|
||||
models.FieldCondition(
|
||||
key="identifier",
|
||||
match=models.MatchAny(
|
||||
any=ids,
|
||||
),
|
||||
)
|
||||
)
|
||||
]
|
||||
if len(keys.intersection(self.metadata_keys)) != 0:
|
||||
for key in keys.intersection(self.metadata_keys):
|
||||
|
||||
if len(keys) > 0:
|
||||
for key in keys:
|
||||
qdrant_must_filters.append(
|
||||
models.FieldCondition(
|
||||
key="metadata.{}".format(key),
|
||||
@@ -108,6 +111,7 @@ class QdrantDB(BaseVectorDB):
|
||||
|
||||
offset = 0
|
||||
existing_ids = []
|
||||
metadatas = []
|
||||
while offset is not None:
|
||||
response = self.client.scroll(
|
||||
collection_name=self.collection_name,
|
||||
@@ -118,7 +122,8 @@ class QdrantDB(BaseVectorDB):
|
||||
offset = response[1]
|
||||
for doc in response[0]:
|
||||
existing_ids.append(doc.payload["identifier"])
|
||||
return {"ids": existing_ids}
|
||||
metadatas.append(doc.payload["metadata"])
|
||||
return {"ids": existing_ids, "metadatas": metadatas}
|
||||
|
||||
def add(
|
||||
self,
|
||||
@@ -143,7 +148,8 @@ class QdrantDB(BaseVectorDB):
|
||||
metadata["text"] = document
|
||||
qdrant_ids.append(str(uuid.uuid4()))
|
||||
payloads.append({"identifier": id, "text": document, "metadata": copy.deepcopy(metadata)})
|
||||
for i in range(0, len(qdrant_ids), self.BATCH_SIZE):
|
||||
|
||||
for i in tqdm(range(0, len(qdrant_ids), self.BATCH_SIZE), desc="Adding data in batches"):
|
||||
self.client.upsert(
|
||||
collection_name=self.collection_name,
|
||||
points=Batch(
|
||||
@@ -180,16 +186,17 @@ class QdrantDB(BaseVectorDB):
|
||||
keys = set(where.keys() if where is not None else set())
|
||||
|
||||
qdrant_must_filters = []
|
||||
if len(keys.intersection(self.metadata_keys)) != 0:
|
||||
for key in keys.intersection(self.metadata_keys):
|
||||
if len(keys) > 0:
|
||||
for key in keys:
|
||||
qdrant_must_filters.append(
|
||||
models.FieldCondition(
|
||||
key="payload.metadata.{}".format(key),
|
||||
key="metadata.{}".format(key),
|
||||
match=models.MatchValue(
|
||||
value=where.get(key),
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
results = self.client.search(
|
||||
collection_name=self.collection_name,
|
||||
query_filter=models.Filter(must=qdrant_must_filters),
|
||||
@@ -228,3 +235,21 @@ class QdrantDB(BaseVectorDB):
|
||||
raise TypeError("Collection name must be a string")
|
||||
self.config.collection_name = name
|
||||
self.collection_name = self._get_or_create_collection()
|
||||
|
||||
@staticmethod
|
||||
def _generate_query(where: dict):
|
||||
must_fields = []
|
||||
for key, value in where.items():
|
||||
must_fields.append(
|
||||
models.FieldCondition(
|
||||
key=f"metadata.{key}",
|
||||
match=models.MatchValue(
|
||||
value=value,
|
||||
),
|
||||
)
|
||||
)
|
||||
return models.Filter(must=must_fields)
|
||||
|
||||
def delete(self, where: dict):
|
||||
db_filter = self._generate_query(where)
|
||||
self.client.delete(collection_name=self.collection_name, points_selector=db_filter)
|
||||
|
||||
@@ -45,6 +45,9 @@ class WeaviateDB(BaseVectorDB):
|
||||
auth_client_secret=weaviate.AuthApiKey(api_key=os.environ.get("WEAVIATE_API_KEY")),
|
||||
**self.config.extra_params,
|
||||
)
|
||||
# Since weaviate uses graphQL, we need to keep track of metadata keys added in the vectordb.
|
||||
# This is needed to filter data while querying.
|
||||
self.metadata_keys = {"data_type", "doc_id", "url", "hash", "app_id"}
|
||||
|
||||
# Call parent init here because embedder is needed
|
||||
super().__init__(config=self.config)
|
||||
@@ -58,7 +61,6 @@ class WeaviateDB(BaseVectorDB):
|
||||
raise ValueError("Embedder not set. Please set an embedder with `set_embedder` before initialization.")
|
||||
|
||||
self.index_name = self._get_index_name()
|
||||
self.metadata_keys = {"data_type", "doc_id", "url", "hash", "app_id"}
|
||||
if not self.client.schema.exists(self.index_name):
|
||||
# id is a reserved field in Weaviate, hence we had to change the name of the id field to identifier
|
||||
# The none vectorizer is crucial as we have our own custom embedding function
|
||||
@@ -127,29 +129,64 @@ class WeaviateDB(BaseVectorDB):
|
||||
:return: ids
|
||||
:rtype: Set[str]
|
||||
"""
|
||||
weaviate_where_operands = []
|
||||
|
||||
if ids is None or len(ids) == 0:
|
||||
return {"ids": []}
|
||||
if ids:
|
||||
for doc_id in ids:
|
||||
weaviate_where_operands.append({"path": ["identifier"], "operator": "Equal", "valueText": doc_id})
|
||||
|
||||
keys = set(where.keys() if where is not None else set())
|
||||
if len(keys) > 0:
|
||||
for key in keys:
|
||||
weaviate_where_operands.append(
|
||||
{
|
||||
"path": ["metadata", self.index_name + "_metadata", key],
|
||||
"operator": "Equal",
|
||||
"valueText": where.get(key),
|
||||
}
|
||||
)
|
||||
|
||||
if len(weaviate_where_operands) == 1:
|
||||
weaviate_where_clause = weaviate_where_operands[0]
|
||||
else:
|
||||
weaviate_where_clause = {"operator": "And", "operands": weaviate_where_operands}
|
||||
|
||||
existing_ids = []
|
||||
metadatas = []
|
||||
cursor = None
|
||||
offset = 0
|
||||
has_iterated_once = False
|
||||
query_metadata_keys = self.metadata_keys.union(keys)
|
||||
while cursor is not None or not has_iterated_once:
|
||||
has_iterated_once = True
|
||||
results = self._query_with_cursor(
|
||||
self.client.query.get(self.index_name, ["identifier"])
|
||||
results = self._query_with_offset(
|
||||
self.client.query.get(
|
||||
self.index_name,
|
||||
[
|
||||
"identifier",
|
||||
weaviate.LinkTo("metadata", self.index_name + "_metadata", list(query_metadata_keys)),
|
||||
],
|
||||
)
|
||||
.with_where(weaviate_where_clause)
|
||||
.with_additional(["id"])
|
||||
.with_limit(self.BATCH_SIZE),
|
||||
cursor,
|
||||
.with_limit(limit or self.BATCH_SIZE),
|
||||
offset,
|
||||
)
|
||||
|
||||
fetched_results = results["data"]["Get"].get(self.index_name, [])
|
||||
if len(fetched_results) == 0:
|
||||
if not fetched_results:
|
||||
break
|
||||
|
||||
for result in fetched_results:
|
||||
existing_ids.append(result["identifier"])
|
||||
metadatas.append(result["metadata"][0])
|
||||
cursor = result["_additional"]["id"]
|
||||
offset += 1
|
||||
|
||||
return {"ids": existing_ids}
|
||||
if limit is not None and len(existing_ids) >= limit:
|
||||
break
|
||||
|
||||
return {"ids": existing_ids, "metadatas": metadatas}
|
||||
|
||||
def add(self, documents: list[str], metadatas: list[object], ids: list[str], **kwargs: Optional[dict[str, any]]):
|
||||
"""add data in vector database
|
||||
@@ -201,21 +238,20 @@ class WeaviateDB(BaseVectorDB):
|
||||
query_vector = self.embedder.embedding_fn([input_query])[0]
|
||||
keys = set(where.keys() if where is not None else set())
|
||||
data_fields = ["text"]
|
||||
|
||||
query_metadata_keys = self.metadata_keys.union(keys)
|
||||
if citations:
|
||||
data_fields.append(weaviate.LinkTo("metadata", self.index_name + "_metadata", list(self.metadata_keys)))
|
||||
data_fields.append(weaviate.LinkTo("metadata", self.index_name + "_metadata", list(query_metadata_keys)))
|
||||
|
||||
if len(keys.intersection(self.metadata_keys)) != 0:
|
||||
if len(keys) > 0:
|
||||
weaviate_where_operands = []
|
||||
for key in keys:
|
||||
if key in self.metadata_keys:
|
||||
weaviate_where_operands.append(
|
||||
{
|
||||
"path": ["metadata", self.index_name + "_metadata", key],
|
||||
"operator": "Equal",
|
||||
"valueText": where.get(key),
|
||||
}
|
||||
)
|
||||
weaviate_where_operands.append(
|
||||
{
|
||||
"path": ["metadata", self.index_name + "_metadata", key],
|
||||
"operator": "Equal",
|
||||
"valueText": where.get(key),
|
||||
}
|
||||
)
|
||||
if len(weaviate_where_operands) == 1:
|
||||
weaviate_where_clause = weaviate_where_operands[0]
|
||||
else:
|
||||
@@ -289,11 +325,37 @@ class WeaviateDB(BaseVectorDB):
|
||||
:return: Weaviate index
|
||||
:rtype: str
|
||||
"""
|
||||
return f"{self.config.collection_name}_{self.embedder.vector_dimension}".capitalize()
|
||||
return f"{self.config.collection_name}_{self.embedder.vector_dimension}".capitalize().replace("-", "_")
|
||||
|
||||
@staticmethod
|
||||
def _query_with_cursor(query, cursor):
|
||||
if cursor is not None:
|
||||
query.with_after(cursor)
|
||||
def _query_with_offset(query, offset):
|
||||
if offset:
|
||||
query.with_offset(offset)
|
||||
results = query.do()
|
||||
return results
|
||||
|
||||
def _generate_query(self, where: dict):
|
||||
weaviate_where_operands = []
|
||||
for key, value in where.items():
|
||||
weaviate_where_operands.append(
|
||||
{
|
||||
"path": ["metadata", self.index_name + "_metadata", key],
|
||||
"operator": "Equal",
|
||||
"valueText": value,
|
||||
}
|
||||
)
|
||||
|
||||
if len(weaviate_where_operands) == 1:
|
||||
weaviate_where_clause = weaviate_where_operands[0]
|
||||
else:
|
||||
weaviate_where_clause = {"operator": "And", "operands": weaviate_where_operands}
|
||||
|
||||
return weaviate_where_clause
|
||||
|
||||
def delete(self, where: dict):
|
||||
"""Delete from database.
|
||||
:param where: to filter data
|
||||
:type where: dict[str, any]
|
||||
"""
|
||||
query = self._generate_query(where)
|
||||
self.client.batch.delete_objects(self.index_name, where=query)
|
||||
|
||||
@@ -69,6 +69,7 @@ class ZillizVectorDB(BaseVectorDB):
|
||||
FieldSchema(name="id", dtype=DataType.VARCHAR, is_primary=True, max_length=512),
|
||||
FieldSchema(name="text", dtype=DataType.VARCHAR, max_length=2048),
|
||||
FieldSchema(name="embeddings", dtype=DataType.FLOAT_VECTOR, dim=self.embedder.vector_dimension),
|
||||
FieldSchema(name="metadata", dtype=DataType.JSON),
|
||||
]
|
||||
|
||||
schema = CollectionSchema(fields, enable_dynamic_field=True)
|
||||
@@ -94,17 +95,26 @@ class ZillizVectorDB(BaseVectorDB):
|
||||
:return: Existing documents.
|
||||
:rtype: Set[str]
|
||||
"""
|
||||
if ids is None or len(ids) == 0 or self.collection.num_entities == 0:
|
||||
return {"ids": []}
|
||||
data_ids = []
|
||||
metadatas = []
|
||||
if self.collection.num_entities == 0 or self.collection.is_empty:
|
||||
return {"ids": data_ids, "metadatas": metadatas}
|
||||
|
||||
if not self.collection.is_empty:
|
||||
filter_ = f"id in {ids}"
|
||||
results = self.client.query(
|
||||
collection_name=self.config.collection_name, filter=filter_, output_fields=["id"]
|
||||
)
|
||||
results = [res["id"] for res in results]
|
||||
filter_ = ""
|
||||
if ids:
|
||||
filter_ = f'id in "{ids}"'
|
||||
|
||||
return {"ids": set(results)}
|
||||
if where:
|
||||
if filter_:
|
||||
filter_ += " and "
|
||||
filter_ = f"{self._generate_zilliz_filter(where)}"
|
||||
|
||||
results = self.client.query(collection_name=self.config.collection_name, filter=filter_, output_fields=["*"])
|
||||
for res in results:
|
||||
data_ids.append(res.get("id"))
|
||||
metadatas.append(res.get("metadata", {}))
|
||||
|
||||
return {"ids": data_ids, "metadatas": metadatas}
|
||||
|
||||
def add(
|
||||
self,
|
||||
@@ -117,7 +127,7 @@ class ZillizVectorDB(BaseVectorDB):
|
||||
embeddings = self.embedder.embedding_fn(documents)
|
||||
|
||||
for id, doc, metadata, embedding in zip(ids, documents, metadatas, embeddings):
|
||||
data = {**metadata, "id": id, "text": doc, "embeddings": embedding}
|
||||
data = {"id": id, "text": doc, "embeddings": embedding, "metadata": metadata}
|
||||
self.client.insert(collection_name=self.config.collection_name, data=data, **kwargs)
|
||||
|
||||
self.collection.load()
|
||||
@@ -128,7 +138,7 @@ class ZillizVectorDB(BaseVectorDB):
|
||||
self,
|
||||
input_query: list[str],
|
||||
n_results: int,
|
||||
where: dict[str, any],
|
||||
where: dict[str, Any],
|
||||
citations: bool = False,
|
||||
**kwargs: Optional[dict[str, Any]],
|
||||
) -> Union[list[tuple[str, dict]], list[str]]:
|
||||
@@ -140,7 +150,7 @@ class ZillizVectorDB(BaseVectorDB):
|
||||
:param n_results: no of similar documents to fetch from database
|
||||
:type n_results: int
|
||||
:param where: to filter data
|
||||
:type where: str
|
||||
:type where: dict[str, Any]
|
||||
:raises InvalidDimensionException: Dimensions do not match.
|
||||
:param citations: we use citations boolean param to return context along with the answer.
|
||||
:type citations: bool, default is False.
|
||||
@@ -152,16 +162,15 @@ class ZillizVectorDB(BaseVectorDB):
|
||||
if self.collection.is_empty:
|
||||
return []
|
||||
|
||||
if not isinstance(where, str):
|
||||
where = None
|
||||
|
||||
output_fields = ["*"]
|
||||
input_query_vector = self.embedder.embedding_fn([input_query])
|
||||
query_vector = input_query_vector[0]
|
||||
|
||||
query_filter = self._generate_zilliz_filter(where)
|
||||
query_result = self.client.search(
|
||||
collection_name=self.config.collection_name,
|
||||
data=[query_vector],
|
||||
filter=query_filter,
|
||||
limit=n_results,
|
||||
output_fields=output_fields,
|
||||
**kwargs,
|
||||
@@ -173,12 +182,10 @@ class ZillizVectorDB(BaseVectorDB):
|
||||
score = query["distance"]
|
||||
context = data["text"]
|
||||
|
||||
if "embeddings" in data:
|
||||
data.pop("embeddings")
|
||||
|
||||
if citations:
|
||||
data["score"] = score
|
||||
contexts.append(tuple((context, data)))
|
||||
metadata = data.get("metadata", {})
|
||||
metadata["score"] = score
|
||||
contexts.append(tuple((context, metadata)))
|
||||
else:
|
||||
contexts.append(context)
|
||||
return contexts
|
||||
@@ -216,7 +223,13 @@ class ZillizVectorDB(BaseVectorDB):
|
||||
raise TypeError("Collection name must be a string")
|
||||
self.config.collection_name = name
|
||||
|
||||
def delete(self, keys: Union[list, str, int]):
|
||||
def _generate_zilliz_filter(self, where: dict[str, str]):
|
||||
operands = []
|
||||
for key, value in where.items():
|
||||
operands.append(f'(metadata["{key}"] == "{value}")')
|
||||
return " and ".join(operands)
|
||||
|
||||
def delete(self, where: dict[str, Any]):
|
||||
"""
|
||||
Delete the embeddings from DB. Zilliz only support deleting with keys.
|
||||
|
||||
@@ -224,7 +237,7 @@ class ZillizVectorDB(BaseVectorDB):
|
||||
:param keys: Primary keys of the table entries to delete.
|
||||
:type keys: Union[list, str, int]
|
||||
"""
|
||||
self.client.delete(
|
||||
collection_name=self.config.collection_name,
|
||||
pks=keys,
|
||||
)
|
||||
data = self.get(where=where)
|
||||
keys = data.get("ids", [])
|
||||
if keys:
|
||||
self.client.delete(collection_name=self.config.collection_name, pks=keys)
|
||||
|
||||
Generated
+72
-2
@@ -383,6 +383,47 @@ files = [
|
||||
{file = "blinker-1.6.3.tar.gz", hash = "sha256:152090d27c1c5c722ee7e48504b02d76502811ce02e1523553b4cf8c8b3d3a8d"},
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "boto3"
|
||||
version = "1.34.22"
|
||||
description = "The AWS SDK for Python"
|
||||
optional = true
|
||||
python-versions = ">= 3.8"
|
||||
files = [
|
||||
{file = "boto3-1.34.22-py3-none-any.whl", hash = "sha256:5909cd1393143576265c692e908a9ae495492c04a0ffd4bae8578adc2e44729e"},
|
||||
{file = "boto3-1.34.22.tar.gz", hash = "sha256:a98c0b86f6044ff8314cc2361e1ef574d674318313ab5606ccb4a6651c7a3f8c"},
|
||||
]
|
||||
|
||||
[package.dependencies]
|
||||
botocore = ">=1.34.22,<1.35.0"
|
||||
jmespath = ">=0.7.1,<2.0.0"
|
||||
s3transfer = ">=0.10.0,<0.11.0"
|
||||
|
||||
[package.extras]
|
||||
crt = ["botocore[crt] (>=1.21.0,<2.0a0)"]
|
||||
|
||||
[[package]]
|
||||
name = "botocore"
|
||||
version = "1.34.22"
|
||||
description = "Low-level, data-driven core of boto 3."
|
||||
optional = true
|
||||
python-versions = ">= 3.8"
|
||||
files = [
|
||||
{file = "botocore-1.34.22-py3-none-any.whl", hash = "sha256:e5f7775975b9213507fbcf846a96b7a2aec2a44fc12a44585197b014a4ab0889"},
|
||||
{file = "botocore-1.34.22.tar.gz", hash = "sha256:c47ba4286c576150d1b6ca6df69a87b5deff3d23bd84da8bcf8431ebac3c40ba"},
|
||||
]
|
||||
|
||||
[package.dependencies]
|
||||
jmespath = ">=0.7.1,<2.0.0"
|
||||
python-dateutil = ">=2.1,<3.0.0"
|
||||
urllib3 = [
|
||||
{version = ">=1.25.4,<1.27", markers = "python_version < \"3.10\""},
|
||||
{version = ">=1.25.4,<2.1", markers = "python_version >= \"3.10\""},
|
||||
]
|
||||
|
||||
[package.extras]
|
||||
crt = ["awscrt (==0.19.19)"]
|
||||
|
||||
[[package]]
|
||||
name = "brotli"
|
||||
version = "1.1.0"
|
||||
@@ -2810,6 +2851,17 @@ MarkupSafe = ">=2.0"
|
||||
[package.extras]
|
||||
i18n = ["Babel (>=2.7)"]
|
||||
|
||||
[[package]]
|
||||
name = "jmespath"
|
||||
version = "1.0.1"
|
||||
description = "JSON Matching Expressions"
|
||||
optional = true
|
||||
python-versions = ">=3.7"
|
||||
files = [
|
||||
{file = "jmespath-1.0.1-py3-none-any.whl", hash = "sha256:02e2e4cc71b5bcab88332eebf907519190dd9e6e82107fa7f83b1003a6252980"},
|
||||
{file = "jmespath-1.0.1.tar.gz", hash = "sha256:90261b206d6defd58fdd5e85f478bf633a2901798906be2ad389150c5c60edbe"},
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "joblib"
|
||||
version = "1.3.2"
|
||||
@@ -4211,9 +4263,9 @@ files = [
|
||||
[package.dependencies]
|
||||
numpy = [
|
||||
{version = ">=1.21.0", markers = "python_version == \"3.9\" and platform_system == \"Darwin\" and platform_machine == \"arm64\""},
|
||||
{version = ">=1.19.3", markers = "platform_system == \"Linux\" and platform_machine == \"aarch64\" and python_version >= \"3.8\" and python_version < \"3.10\" or python_version > \"3.9\" and python_version < \"3.10\" or python_version >= \"3.9\" and platform_system != \"Darwin\" and python_version < \"3.10\" or python_version >= \"3.9\" and platform_machine != \"arm64\" and python_version < \"3.10\""},
|
||||
{version = ">=1.21.4", markers = "python_version >= \"3.10\" and platform_system == \"Darwin\" and python_version < \"3.11\""},
|
||||
{version = ">=1.21.2", markers = "platform_system != \"Darwin\" and python_version >= \"3.10\" and python_version < \"3.11\""},
|
||||
{version = ">=1.19.3", markers = "platform_system == \"Linux\" and platform_machine == \"aarch64\" and python_version >= \"3.8\" and python_version < \"3.10\" or python_version > \"3.9\" and python_version < \"3.10\" or python_version >= \"3.9\" and platform_system != \"Darwin\" and python_version < \"3.10\" or python_version >= \"3.9\" and platform_machine != \"arm64\" and python_version < \"3.10\""},
|
||||
{version = ">=1.23.5", markers = "python_version >= \"3.11\""},
|
||||
]
|
||||
|
||||
@@ -6091,6 +6143,23 @@ files = [
|
||||
{file = "ruff-0.1.11.tar.gz", hash = "sha256:f9d4d88cb6eeb4dfe20f9f0519bd2eaba8119bde87c3d5065c541dbae2b5a2cb"},
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "s3transfer"
|
||||
version = "0.10.0"
|
||||
description = "An Amazon S3 Transfer Manager"
|
||||
optional = true
|
||||
python-versions = ">= 3.8"
|
||||
files = [
|
||||
{file = "s3transfer-0.10.0-py3-none-any.whl", hash = "sha256:3cdb40f5cfa6966e812209d0994f2a4709b561c88e90cf00c2696d2df4e56b2e"},
|
||||
{file = "s3transfer-0.10.0.tar.gz", hash = "sha256:d0c8bbf672d5eebbe4e57945e23b972d963f07d82f661cabf678a5c88831595b"},
|
||||
]
|
||||
|
||||
[package.dependencies]
|
||||
botocore = ">=1.33.2,<2.0a.0"
|
||||
|
||||
[package.extras]
|
||||
crt = ["botocore[crt] (>=1.33.2,<2.0a.0)"]
|
||||
|
||||
[[package]]
|
||||
name = "safetensors"
|
||||
version = "0.4.0"
|
||||
@@ -8213,6 +8282,7 @@ docs = ["furo", "jaraco.packaging (>=9.3)", "jaraco.tidelift (>=1.4)", "rst.link
|
||||
testing = ["big-O", "jaraco.functools", "jaraco.itertools", "more-itertools", "pytest (>=6)", "pytest-black (>=0.3.7)", "pytest-checkdocs (>=2.4)", "pytest-cov", "pytest-enabler (>=2.2)", "pytest-ignore-flaky", "pytest-mypy (>=0.9.1)", "pytest-ruff"]
|
||||
|
||||
[extras]
|
||||
aws-bedrock = ["boto3"]
|
||||
cohere = ["cohere"]
|
||||
dataloaders = ["docx2txt", "duckduckgo-search", "pytube", "sentence-transformers", "unstructured", "youtube-transcript-api"]
|
||||
discord = ["discord"]
|
||||
@@ -8246,4 +8316,4 @@ youtube = ["youtube-transcript-api", "yt_dlp"]
|
||||
[metadata]
|
||||
lock-version = "2.0"
|
||||
python-versions = ">=3.9,<3.12"
|
||||
content-hash = "cb0da55af7c61300bb321770ed319c900b6b3ba3865421d63eb9120beb73d06c"
|
||||
content-hash = "bbcf32e87c0784d031fb6cf9bd89655375839da0660b8feb2026ffdd971623d7"
|
||||
|
||||
+3
-1
@@ -1,6 +1,6 @@
|
||||
[tool.poetry]
|
||||
name = "embedchain"
|
||||
version = "0.1.67"
|
||||
version = "0.1.69"
|
||||
description = "Simplest open source retrieval(RAG) framework"
|
||||
authors = [
|
||||
"Taranjeet Singh <taranjeet@embedchain.ai>",
|
||||
@@ -149,6 +149,7 @@ google-auth-oauthlib = { version = "^1.2.0", optional = true }
|
||||
google-auth = { version = "^2.25.2", optional = true }
|
||||
google-auth-httplib2 = { version = "^0.2.0", optional = true }
|
||||
google-api-core = { version = "^2.15.0", optional = true }
|
||||
boto3 = { version = "^1.34.20", optional = true }
|
||||
langchain-mistralai = { version = "^0.0.3", optional = true }
|
||||
|
||||
[tool.poetry.group.dev.dependencies]
|
||||
@@ -215,6 +216,7 @@ rss_feed = [
|
||||
google = ["google-generativeai"]
|
||||
modal = ["modal"]
|
||||
dropbox = ["dropbox"]
|
||||
aws_bedrock = ["boto3"]
|
||||
mistralai = ["langchain-mistralai"]
|
||||
|
||||
[tool.poetry.group.docs.dependencies]
|
||||
|
||||
@@ -0,0 +1,56 @@
|
||||
import pytest
|
||||
from langchain.callbacks.streaming_stdout import StreamingStdOutCallbackHandler
|
||||
|
||||
from embedchain.config import BaseLlmConfig
|
||||
from embedchain.llm.aws_bedrock import AWSBedrockLlm
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def config(monkeypatch):
|
||||
monkeypatch.setenv("AWS_ACCESS_KEY_ID", "test_access_key_id")
|
||||
monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "test_secret_access_key")
|
||||
monkeypatch.setenv("OPENAI_API_KEY", "test_api_key")
|
||||
config = BaseLlmConfig(
|
||||
model="amazon.titan-text-express-v1",
|
||||
model_kwargs={
|
||||
"temperature": 0.5,
|
||||
"topP": 1,
|
||||
"maxTokenCount": 1000,
|
||||
},
|
||||
)
|
||||
yield config
|
||||
monkeypatch.delenv("AWS_ACCESS_KEY_ID")
|
||||
monkeypatch.delenv("AWS_SECRET_ACCESS_KEY")
|
||||
monkeypatch.delenv("OPENAI_API_KEY")
|
||||
|
||||
|
||||
def test_get_llm_model_answer(config, mocker):
|
||||
mocked_get_answer = mocker.patch("embedchain.llm.aws_bedrock.AWSBedrockLlm._get_answer", return_value="Test answer")
|
||||
|
||||
llm = AWSBedrockLlm(config)
|
||||
answer = llm.get_llm_model_answer("Test query")
|
||||
|
||||
assert answer == "Test answer"
|
||||
mocked_get_answer.assert_called_once_with("Test query", config)
|
||||
|
||||
|
||||
def test_get_llm_model_answer_empty_prompt(config, mocker):
|
||||
mocked_get_answer = mocker.patch("embedchain.llm.aws_bedrock.AWSBedrockLlm._get_answer", return_value="Test answer")
|
||||
|
||||
llm = AWSBedrockLlm(config)
|
||||
answer = llm.get_llm_model_answer("")
|
||||
|
||||
assert answer == "Test answer"
|
||||
mocked_get_answer.assert_called_once_with("", config)
|
||||
|
||||
|
||||
def test_get_llm_model_answer_with_streaming(config, mocker):
|
||||
config.stream = True
|
||||
mocked_bedrock_chat = mocker.patch("embedchain.llm.aws_bedrock.Bedrock")
|
||||
|
||||
llm = AWSBedrockLlm(config)
|
||||
llm.get_llm_model_answer("Test query")
|
||||
|
||||
mocked_bedrock_chat.assert_called_once()
|
||||
callbacks = [callback[1]["callbacks"] for callback in mocked_bedrock_chat.call_args_list]
|
||||
assert any(isinstance(callback[0], StreamingStdOutCallbackHandler) for callback in callbacks)
|
||||
@@ -3,6 +3,7 @@ from unittest.mock import patch
|
||||
|
||||
from embedchain import App
|
||||
from embedchain.config import AppConfig
|
||||
from embedchain.config.vectordb.pinecone import PineconeDBConfig
|
||||
from embedchain.embedder.base import BaseEmbedder
|
||||
from embedchain.vectordb.pinecone import PineconeDB
|
||||
|
||||
@@ -100,7 +101,39 @@ class TestPinecone:
|
||||
db.reset()
|
||||
|
||||
# Assert that the Pinecone client was called to delete the index
|
||||
pinecone_mock.delete_index.assert_called_once_with(db.index_name)
|
||||
pinecone_mock.delete_index.assert_called_once_with(db.config.index_name)
|
||||
|
||||
# Assert that the index is recreated
|
||||
pinecone_mock.Index.assert_called_with(db.index_name)
|
||||
pinecone_mock.Index.assert_called_with(db.config.index_name)
|
||||
|
||||
@patch("embedchain.vectordb.pinecone.pinecone")
|
||||
def test_custom_index_name_if_it_exists(self, pinecone_mock):
|
||||
"""Tests custom index name is used if it exists"""
|
||||
pinecone_mock.list_indexes.return_value = ["custom_index_name"]
|
||||
db_config = PineconeDBConfig(index_name="custom_index_name")
|
||||
_ = PineconeDB(config=db_config)
|
||||
|
||||
pinecone_mock.list_indexes.assert_called_once()
|
||||
pinecone_mock.create_index.assert_not_called()
|
||||
pinecone_mock.Index.assert_called_with("custom_index_name")
|
||||
|
||||
@patch("embedchain.vectordb.pinecone.pinecone")
|
||||
def test_custom_index_name_creation(self, pinecone_mock):
|
||||
"""Test custom index name is created if it doesn't exists already"""
|
||||
pinecone_mock.list_indexes.return_value = []
|
||||
db_config = PineconeDBConfig(index_name="custom_index_name")
|
||||
_ = PineconeDB(config=db_config)
|
||||
|
||||
pinecone_mock.list_indexes.assert_called_once()
|
||||
pinecone_mock.create_index.assert_called_once()
|
||||
pinecone_mock.Index.assert_called_with("custom_index_name")
|
||||
|
||||
@patch("embedchain.vectordb.pinecone.pinecone")
|
||||
def test_default_index_name_is_used(self, pinecone_mock):
|
||||
"""Test default index name is used if custom index name is not provided"""
|
||||
db_config = PineconeDBConfig(collection_name="my-collection")
|
||||
_ = PineconeDB(config=db_config)
|
||||
|
||||
pinecone_mock.list_indexes.assert_called_once()
|
||||
pinecone_mock.create_index.assert_called_once()
|
||||
pinecone_mock.Index.assert_called_with(f"{db_config.collection_name}-{db_config.vector_dimension}")
|
||||
|
||||
@@ -56,9 +56,9 @@ class TestQdrantDB(unittest.TestCase):
|
||||
App(config=app_config, db=db, embedding_model=embedder)
|
||||
|
||||
resp = db.get(ids=[], where={})
|
||||
self.assertEqual(resp, {"ids": []})
|
||||
self.assertEqual(resp, {"ids": [], "metadatas": []})
|
||||
resp2 = db.get(ids=["123", "456"], where={"url": "https://ai.ai"})
|
||||
self.assertEqual(resp2, {"ids": []})
|
||||
self.assertEqual(resp2, {"ids": [], "metadatas": []})
|
||||
|
||||
@patch("embedchain.vectordb.qdrant.QdrantClient")
|
||||
@patch.object(uuid, "uuid4", side_effect=TEST_UUIDS)
|
||||
@@ -119,7 +119,7 @@ class TestQdrantDB(unittest.TestCase):
|
||||
query_filter=models.Filter(
|
||||
must=[
|
||||
models.FieldCondition(
|
||||
key="payload.metadata.doc_id",
|
||||
key="metadata.doc_id",
|
||||
match=models.MatchValue(
|
||||
value="123",
|
||||
),
|
||||
|
||||
@@ -130,7 +130,11 @@ class TestZillizDBCollection:
|
||||
[
|
||||
{
|
||||
"distance": 0.0,
|
||||
"entity": {"text": "result_doc", "url": "url_1", "doc_id": "doc_id_1", "embeddings": [1, 2, 3]},
|
||||
"entity": {
|
||||
"text": "result_doc",
|
||||
"embeddings": [1, 2, 3],
|
||||
"metadata": {"url": "url_1", "doc_id": "doc_id_1"},
|
||||
},
|
||||
}
|
||||
]
|
||||
]
|
||||
@@ -141,6 +145,7 @@ class TestZillizDBCollection:
|
||||
mock_search.assert_called_with(
|
||||
collection_name=mock_config.collection_name,
|
||||
data=["query_vector"],
|
||||
filter="",
|
||||
limit=1,
|
||||
output_fields=["*"],
|
||||
)
|
||||
@@ -155,10 +160,9 @@ class TestZillizDBCollection:
|
||||
mock_search.assert_called_with(
|
||||
collection_name=mock_config.collection_name,
|
||||
data=["query_vector"],
|
||||
filter="",
|
||||
limit=1,
|
||||
output_fields=["*"],
|
||||
)
|
||||
|
||||
assert query_result_with_citations == [
|
||||
("result_doc", {"text": "result_doc", "url": "url_1", "doc_id": "doc_id_1", "score": 0.0})
|
||||
]
|
||||
assert query_result_with_citations == [("result_doc", {"url": "url_1", "doc_id": "doc_id_1", "score": 0.0})]
|
||||
|
||||
Reference in New Issue
Block a user