Compare commits
11 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| adde398b65 | |||
| d8897ce356 | |||
| 0ea8ab228c | |||
| d62a23edf6 | |||
| 51ebf3439b | |||
| 4a5ed1dd8d | |||
| e84b5034ea | |||
| a4831d6ed9 | |||
| 51b4966801 | |||
| 1d4e00ccef | |||
| c9fbc2e7d6 |
@@ -58,6 +58,12 @@ Install related dependencies using the following command:
|
||||
pip install --upgrade 'embedchain[elasticsearch]'
|
||||
```
|
||||
|
||||
<Note>
|
||||
You can configure the Elasticsearch connection by providing either `es_url` or `cloud_id`. If you are using the Elasticsearch Service on Elastic Cloud, you can find the `cloud_id` on the [Elastic Cloud dashboard](https://cloud.elastic.co/deployments).
|
||||
</Note>
|
||||
|
||||
You can authorize the connection to Elasticsearch by providing either `basic_auth`, `api_key`, or `bearer_auth`.
|
||||
|
||||
<CodeGroup>
|
||||
|
||||
```python main.py
|
||||
@@ -72,11 +78,10 @@ vectordb:
|
||||
provider: elasticsearch
|
||||
config:
|
||||
collection_name: 'es-index'
|
||||
es_url: http://localhost:9200
|
||||
http_auth:
|
||||
- admin
|
||||
- admin
|
||||
api_key: xxx
|
||||
cloud_id: 'deployment-name:xxxx'
|
||||
basic_auth:
|
||||
- elastic
|
||||
- <your_password>
|
||||
verify_certs: false
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
@@ -0,0 +1,16 @@
|
||||
---
|
||||
title: "🐝 Beehiiv"
|
||||
---
|
||||
|
||||
To add any Beehiiv data sources to your app, just add the base url as the source and set the data_type to `beehiiv`.
|
||||
|
||||
```python
|
||||
from embedchain import Pipeline as App
|
||||
|
||||
app = App()
|
||||
|
||||
# source: just add the base url and set the data_type to 'beehiiv'
|
||||
app.add('https://aibreakfast.beehiiv.com', data_type='beehiiv')
|
||||
app.query("How much is OpenAI paying developers?")
|
||||
# Answer: OpenAI is aggressively recruiting Google's top AI researchers with offers ranging between $5 to $10 million annually, primarily in stock options.
|
||||
```
|
||||
@@ -27,6 +27,8 @@ Embedchain comes with built-in support for various data sources. We handle the c
|
||||
<Card title="💬 Discord" href="/data-sources/discord"></Card>
|
||||
<Card title="📝 Github" href="/data-sources/github"></Card>
|
||||
<Card title="⚙️ Custom" href="/data-sources/custom"></Card>
|
||||
<Card title="📝 Substack" href="/data-sources/substack"></Card>
|
||||
<Card title="🐝 Beehiiv" href="/data-sources/beehiiv"></Card>
|
||||
</CardGroup>
|
||||
|
||||
<br/ >
|
||||
|
||||
@@ -2,15 +2,15 @@
|
||||
title: "📝 Substack"
|
||||
---
|
||||
|
||||
To add any Substack data sources to your app, just add the sitemap.xml of that url as the source and set the data_type to `substack`.
|
||||
To add any Substack data sources to your app, just add the main base url as the source and set the data_type to `substack`.
|
||||
|
||||
```python
|
||||
from embedchain import Pipeline as App
|
||||
|
||||
app = App()
|
||||
|
||||
# source: for any substack just add the sitemap.xml url
|
||||
app.add('https://www.lennysnewsletter.com/sitemap.xml', data_type='substack')
|
||||
# source: for any substack just add the root URL
|
||||
app.add('https://www.lennysnewsletter.com', data_type='substack')
|
||||
app.query("Who is Brian Chesky?")
|
||||
# Answer: Brian Chesky is the co-founder and CEO of Airbnb.
|
||||
```
|
||||
|
||||
+2
-1
@@ -90,7 +90,8 @@
|
||||
"data-sources/youtube-video",
|
||||
"data-sources/discourse",
|
||||
"data-sources/substack",
|
||||
"data-sources/discord"
|
||||
"data-sources/discord",
|
||||
"data-sources/beehiiv"
|
||||
]
|
||||
},
|
||||
"data-sources/data-type-handling"
|
||||
|
||||
@@ -41,7 +41,6 @@ class BaseChunker(JSONSerializable):
|
||||
url = meta_data["url"]
|
||||
|
||||
chunks = self.get_chunks(content)
|
||||
|
||||
for chunk in chunks:
|
||||
chunk_id = hashlib.sha256((chunk + url).encode()).hexdigest()
|
||||
chunk_id = f"{app_id}--{chunk_id}" if app_id is not None else chunk_id
|
||||
|
||||
@@ -0,0 +1,22 @@
|
||||
from typing import Optional
|
||||
|
||||
from langchain.text_splitter import RecursiveCharacterTextSplitter
|
||||
|
||||
from embedchain.chunkers.base_chunker import BaseChunker
|
||||
from embedchain.config.add_config import ChunkerConfig
|
||||
from embedchain.helpers.json_serializable import register_deserializable
|
||||
|
||||
|
||||
@register_deserializable
|
||||
class BeehiivChunker(BaseChunker):
|
||||
"""Chunker for Beehiiv."""
|
||||
|
||||
def __init__(self, config: Optional[ChunkerConfig] = None):
|
||||
if config is None:
|
||||
config = ChunkerConfig(chunk_size=1000, chunk_overlap=0, length_function=len)
|
||||
text_splitter = RecursiveCharacterTextSplitter(
|
||||
chunk_size=config.chunk_size,
|
||||
chunk_overlap=config.chunk_overlap,
|
||||
length_function=config.length_function,
|
||||
)
|
||||
super().__init__(text_splitter)
|
||||
@@ -0,0 +1,22 @@
|
||||
from typing import Optional
|
||||
|
||||
from langchain.text_splitter import RecursiveCharacterTextSplitter
|
||||
|
||||
from embedchain.chunkers.base_chunker import BaseChunker
|
||||
from embedchain.config.add_config import ChunkerConfig
|
||||
from embedchain.helper.json_serializable import register_deserializable
|
||||
|
||||
|
||||
@register_deserializable
|
||||
class RSSFeedChunker(BaseChunker):
|
||||
"""Chunker for RSS Feed."""
|
||||
|
||||
def __init__(self, config: Optional[ChunkerConfig] = None):
|
||||
if config is None:
|
||||
config = ChunkerConfig(chunk_size=2000, chunk_overlap=0, length_function=len)
|
||||
text_splitter = RecursiveCharacterTextSplitter(
|
||||
chunk_size=config.chunk_size,
|
||||
chunk_overlap=config.chunk_overlap,
|
||||
length_function=config.length_function,
|
||||
)
|
||||
super().__init__(text_splitter)
|
||||
@@ -57,7 +57,7 @@ class BaseLlmConfig(BaseConfig):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
number_documents: int = 1,
|
||||
number_documents: int = 3,
|
||||
template: Optional[Template] = None,
|
||||
model: Optional[str] = None,
|
||||
temperature: float = 0,
|
||||
|
||||
@@ -12,6 +12,7 @@ class ElasticsearchDBConfig(BaseVectorDbConfig):
|
||||
collection_name: Optional[str] = None,
|
||||
dir: Optional[str] = None,
|
||||
es_url: Union[str, List[str]] = None,
|
||||
cloud_id: Optional[str] = None,
|
||||
**ES_EXTRA_PARAMS: Dict[str, any],
|
||||
):
|
||||
"""
|
||||
@@ -26,12 +27,15 @@ class ElasticsearchDBConfig(BaseVectorDbConfig):
|
||||
:param ES_EXTRA_PARAMS: extra params dict that can be passed to elasticsearch.
|
||||
:type ES_EXTRA_PARAMS: Dict[str, Any], optional
|
||||
"""
|
||||
if es_url and cloud_id:
|
||||
raise ValueError("Only one of `es_url` and `cloud_id` can be set.")
|
||||
# self, es_url: Union[str, List[str]] = None, **ES_EXTRA_PARAMS: Dict[str, any]):
|
||||
self.ES_URL = es_url or os.environ.get("ELASTICSEARCH_URL")
|
||||
if not self.ES_URL:
|
||||
self.CLOUD_ID = cloud_id or os.environ.get("ELASTICSEARCH_CLOUD_ID")
|
||||
if not self.ES_URL and not self.CLOUD_ID:
|
||||
raise AttributeError(
|
||||
"Elasticsearch needs a URL attribute, "
|
||||
"this can either be passed to `ElasticsearchDBConfig` or as `ELASTICSEARCH_URL` in `.env`"
|
||||
"Elasticsearch needs a URL or CLOUD_ID attribute, "
|
||||
"this can either be passed to `ElasticsearchDBConfig` or as `ELASTICSEARCH_URL` or `ELASTICSEARCH_CLOUD_ID` in `.env`" # noqa: E501
|
||||
)
|
||||
self.ES_EXTRA_PARAMS = ES_EXTRA_PARAMS
|
||||
# Load API key from .env if it's not explicitly passed.
|
||||
@@ -40,7 +44,6 @@ class ElasticsearchDBConfig(BaseVectorDbConfig):
|
||||
not self.ES_EXTRA_PARAMS.get("api_key")
|
||||
and not self.ES_EXTRA_PARAMS.get("basic_auth")
|
||||
and not self.ES_EXTRA_PARAMS.get("bearer_auth")
|
||||
and not self.ES_EXTRA_PARAMS.get("http_auth")
|
||||
):
|
||||
self.ES_EXTRA_PARAMS["api_key"] = os.environ.get("ELASTICSEARCH_API_KEY")
|
||||
super().__init__(collection_name=collection_name, dir=dir)
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
from importlib import import_module
|
||||
from typing import Any, Dict
|
||||
from typing import Optional
|
||||
|
||||
from embedchain.chunkers.base_chunker import BaseChunker
|
||||
from embedchain.config import AddConfig
|
||||
@@ -16,7 +16,13 @@ class DataFormatter(JSONSerializable):
|
||||
.add or .add_local method call
|
||||
"""
|
||||
|
||||
def __init__(self, data_type: DataType, config: AddConfig, kwargs: Dict[str, Any]):
|
||||
def __init__(
|
||||
self,
|
||||
data_type: DataType,
|
||||
config: AddConfig,
|
||||
loader: Optional[BaseLoader] = None,
|
||||
chunker: Optional[BaseChunker] = None,
|
||||
):
|
||||
"""
|
||||
Initialize a dataformatter, set data type and chunker based on datatype.
|
||||
|
||||
@@ -25,15 +31,15 @@ class DataFormatter(JSONSerializable):
|
||||
:param config: AddConfig instance with nested loader and chunker config attributes.
|
||||
:type config: AddConfig
|
||||
"""
|
||||
self.loader = self._get_loader(data_type=data_type, config=config.loader, kwargs=kwargs)
|
||||
self.chunker = self._get_chunker(data_type=data_type, config=config.chunker, kwargs=kwargs)
|
||||
self.loader = self._get_loader(data_type=data_type, config=config.loader, loader=loader)
|
||||
self.chunker = self._get_chunker(data_type=data_type, config=config.chunker, chunker=chunker)
|
||||
|
||||
def _lazy_load(self, module_path: str):
|
||||
module_path, class_name = module_path.rsplit(".", 1)
|
||||
module = import_module(module_path)
|
||||
return getattr(module, class_name)
|
||||
|
||||
def _get_loader(self, data_type: DataType, config: LoaderConfig, kwargs: Dict[str, Any]) -> BaseLoader:
|
||||
def _get_loader(self, data_type: DataType, config: LoaderConfig, loader: Optional[BaseLoader]) -> BaseLoader:
|
||||
"""
|
||||
Returns the appropriate data loader for the given data type.
|
||||
|
||||
@@ -66,10 +72,12 @@ class DataFormatter(JSONSerializable):
|
||||
DataType.SUBSTACK: "embedchain.loaders.substack.SubstackLoader",
|
||||
DataType.YOUTUBE_CHANNEL: "embedchain.loaders.youtube_channel.YoutubeChannelLoader",
|
||||
DataType.DISCORD: "embedchain.loaders.discord.DiscordLoader",
|
||||
DataType.RSSFEED: "embedchain.loaders.rss_feed.RSSFeedLoader",
|
||||
DataType.BEEHIIV: "embedchain.loaders.beehiiv.BeehiivLoader",
|
||||
}
|
||||
|
||||
if data_type == DataType.CUSTOM or ("loader" in kwargs):
|
||||
loader_class: type = kwargs.get("loader", None)
|
||||
if data_type == DataType.CUSTOM or loader is not None:
|
||||
loader_class: type = loader
|
||||
if loader_class:
|
||||
return loader_class
|
||||
elif data_type in loaders:
|
||||
@@ -82,7 +90,7 @@ class DataFormatter(JSONSerializable):
|
||||
check `https://docs.embedchain.ai/data-sources/overview`."
|
||||
)
|
||||
|
||||
def _get_chunker(self, data_type: DataType, config: ChunkerConfig, kwargs: Dict[str, Any]) -> BaseChunker:
|
||||
def _get_chunker(self, data_type: DataType, config: ChunkerConfig, chunker: Optional[BaseChunker]) -> BaseChunker:
|
||||
"""Returns the appropriate chunker for the given data type (updated for lazy loading)."""
|
||||
chunker_classes = {
|
||||
DataType.YOUTUBE_VIDEO: "embedchain.chunkers.youtube_video.YoutubeVideoChunker",
|
||||
@@ -106,14 +114,12 @@ class DataFormatter(JSONSerializable):
|
||||
DataType.YOUTUBE_CHANNEL: "embedchain.chunkers.common_chunker.CommonChunker",
|
||||
DataType.DISCORD: "embedchain.chunkers.common_chunker.CommonChunker",
|
||||
DataType.CUSTOM: "embedchain.chunkers.common_chunker.CommonChunker",
|
||||
DataType.RSSFEED: "embedchain.chunkers.rss_feed.RSSFeedChunker",
|
||||
DataType.BEEHIIV: "embedchain.chunkers.beehiiv.BeehiivChunker",
|
||||
}
|
||||
|
||||
if "chunker" in kwargs:
|
||||
chunker_class = kwargs.get("chunker", None)
|
||||
if chunker_class:
|
||||
chunker = chunker_class(config)
|
||||
chunker.set_data_type(data_type)
|
||||
return chunker
|
||||
if chunker is not None:
|
||||
return chunker
|
||||
elif data_type in chunker_classes:
|
||||
chunker_class = self._lazy_load(chunker_classes[data_type])
|
||||
chunker = chunker_class(config)
|
||||
|
||||
+21
-11
@@ -133,7 +133,9 @@ class EmbedChain(JSONSerializable):
|
||||
metadata: Optional[Dict[str, Any]] = None,
|
||||
config: Optional[AddConfig] = None,
|
||||
dry_run=False,
|
||||
**kwargs: Dict[str, Any],
|
||||
loader: Optional[BaseLoader] = None,
|
||||
chunker: Optional[BaseChunker] = None,
|
||||
**kwargs: Optional[Dict[str, Any]],
|
||||
):
|
||||
"""
|
||||
Adds the data from the given URL to the vector db.
|
||||
@@ -192,9 +194,9 @@ class EmbedChain(JSONSerializable):
|
||||
|
||||
self.user_asks.append([source, data_type.value, metadata])
|
||||
|
||||
data_formatter = DataFormatter(data_type, config, kwargs)
|
||||
data_formatter = DataFormatter(data_type, config, loader, chunker)
|
||||
documents, metadatas, _ids, new_chunks = self._load_and_embed(
|
||||
data_formatter.loader, data_formatter.chunker, source, metadata, source_hash, dry_run
|
||||
data_formatter.loader, data_formatter.chunker, source, metadata, source_hash, dry_run, **kwargs
|
||||
)
|
||||
if data_type in {DataType.DOCS_SITE}:
|
||||
self.is_docs_site_instance = True
|
||||
@@ -238,7 +240,7 @@ class EmbedChain(JSONSerializable):
|
||||
data_type: Optional[DataType] = None,
|
||||
metadata: Optional[Dict[str, Any]] = None,
|
||||
config: Optional[AddConfig] = None,
|
||||
**kwargs: Dict[str, Any],
|
||||
**kwargs: Optional[Dict[str, Any]],
|
||||
):
|
||||
"""
|
||||
Adds the data from the given URL to the vector db.
|
||||
@@ -269,7 +271,7 @@ class EmbedChain(JSONSerializable):
|
||||
data_type=data_type,
|
||||
metadata=metadata,
|
||||
config=config,
|
||||
kwargs=kwargs,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def _get_existing_doc_id(self, chunker: BaseChunker, src: Any):
|
||||
@@ -338,6 +340,7 @@ class EmbedChain(JSONSerializable):
|
||||
metadata: Optional[Dict[str, Any]] = None,
|
||||
source_hash: Optional[str] = None,
|
||||
dry_run=False,
|
||||
**kwargs: Optional[Dict[str, Any]],
|
||||
):
|
||||
"""
|
||||
Loads the data from the given URL, chunks it, and adds it to database.
|
||||
@@ -431,6 +434,7 @@ class EmbedChain(JSONSerializable):
|
||||
metadatas=metadatas,
|
||||
ids=ids,
|
||||
skip_embedding=(chunker.data_type == DataType.IMAGES),
|
||||
**kwargs,
|
||||
)
|
||||
count_new_chunks = self.db.count() - chunks_before_addition
|
||||
|
||||
@@ -448,7 +452,12 @@ class EmbedChain(JSONSerializable):
|
||||
]
|
||||
|
||||
def _retrieve_from_database(
|
||||
self, input_query: str, config: Optional[BaseLlmConfig] = None, where=None, citations: bool = False
|
||||
self,
|
||||
input_query: str,
|
||||
config: Optional[BaseLlmConfig] = None,
|
||||
where=None,
|
||||
citations: bool = False,
|
||||
**kwargs: Optional[Dict[str, Any]],
|
||||
) -> Union[List[Tuple[str, str, str]], List[str]]:
|
||||
"""
|
||||
Queries the vector database based on the given input query.
|
||||
@@ -492,6 +501,7 @@ class EmbedChain(JSONSerializable):
|
||||
where=where,
|
||||
skip_embedding=(hasattr(config, "query_type") and config.query_type == "Images"),
|
||||
citations=citations,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
return contexts
|
||||
@@ -502,6 +512,7 @@ class EmbedChain(JSONSerializable):
|
||||
config: BaseLlmConfig = None,
|
||||
dry_run=False,
|
||||
where: Optional[Dict] = None,
|
||||
citations: bool = False,
|
||||
**kwargs: Dict[str, Any],
|
||||
) -> Union[Tuple[str, List[Tuple[str, str, str]]], str]:
|
||||
"""
|
||||
@@ -526,9 +537,8 @@ class EmbedChain(JSONSerializable):
|
||||
or the dry run result
|
||||
:rtype: str, if citations is False, otherwise Tuple[str,List[Tuple[str,str,str]]]
|
||||
"""
|
||||
citations = kwargs.get("citations", False)
|
||||
contexts = self._retrieve_from_database(
|
||||
input_query=input_query, config=config, where=where, citations=citations
|
||||
input_query=input_query, config=config, where=where, citations=citations, **kwargs
|
||||
)
|
||||
if citations and len(contexts) > 0 and isinstance(contexts[0], tuple):
|
||||
contexts_data_for_llm_query = list(map(lambda x: x[0], contexts))
|
||||
@@ -553,8 +563,9 @@ class EmbedChain(JSONSerializable):
|
||||
config: Optional[BaseLlmConfig] = None,
|
||||
dry_run=False,
|
||||
where: Optional[Dict[str, str]] = None,
|
||||
citations: bool = False,
|
||||
**kwargs: Dict[str, Any],
|
||||
) -> str:
|
||||
) -> Union[Tuple[str, List[Tuple[str, str, str]]], str]:
|
||||
"""
|
||||
Queries the vector database on the given input query.
|
||||
Gets relevant doc based on the query and then passes it to an
|
||||
@@ -579,9 +590,8 @@ class EmbedChain(JSONSerializable):
|
||||
or the dry run result
|
||||
:rtype: str, if citations is False, otherwise Tuple[str,List[Tuple[str,str,str]]]
|
||||
"""
|
||||
citations = kwargs.get("citations", False)
|
||||
contexts = self._retrieve_from_database(
|
||||
input_query=input_query, config=config, where=where, citations=citations
|
||||
input_query=input_query, config=config, where=where, citations=citations, **kwargs
|
||||
)
|
||||
if citations and len(contexts) > 0 and isinstance(contexts[0], tuple):
|
||||
contexts_data_for_llm_query = list(map(lambda x: x[0], contexts))
|
||||
|
||||
@@ -0,0 +1,104 @@
|
||||
import hashlib
|
||||
import logging
|
||||
import time
|
||||
import requests
|
||||
from xml.etree import ElementTree
|
||||
|
||||
from embedchain.helpers.json_serializable import register_deserializable
|
||||
from embedchain.loaders.base_loader import BaseLoader
|
||||
from embedchain.utils import is_readable
|
||||
|
||||
|
||||
@register_deserializable
|
||||
class BeehiivLoader(BaseLoader):
|
||||
"""
|
||||
This loader is used to load data from Beehiiv URLs.
|
||||
"""
|
||||
|
||||
def load_data(self, url: str):
|
||||
try:
|
||||
from bs4 import BeautifulSoup
|
||||
from bs4.builder import ParserRejectedMarkup
|
||||
except ImportError:
|
||||
raise ImportError(
|
||||
'Beehiiv requires extra dependencies. Install with `pip install --upgrade "embedchain[dataloaders]"`'
|
||||
) from None
|
||||
|
||||
if not url.endswith("sitemap.xml"):
|
||||
url = url + "/sitemap.xml"
|
||||
|
||||
output = []
|
||||
# we need to set this as a header to avoid 403
|
||||
headers = {
|
||||
"User-Agent": (
|
||||
"Mozilla/5.0 (Macintosh; Intel Mac OS X 10_11_5) "
|
||||
"AppleWebKit/537.36 (KHTML, like Gecko) Chrome/50.0.2661.102 "
|
||||
"Safari/537.36"
|
||||
),
|
||||
}
|
||||
response = requests.get(url, headers=headers)
|
||||
try:
|
||||
response.raise_for_status()
|
||||
except requests.exceptions.HTTPError as e:
|
||||
raise ValueError(
|
||||
f"""
|
||||
Failed to load {url}: {e}. Please use the root substack URL. For example, https://example.substack.com
|
||||
"""
|
||||
)
|
||||
|
||||
try:
|
||||
ElementTree.fromstring(response.content)
|
||||
except ElementTree.ParseError:
|
||||
raise ValueError(
|
||||
f"""
|
||||
Failed to parse {url}. Please use the root substack URL. For example, https://example.substack.com
|
||||
"""
|
||||
)
|
||||
soup = BeautifulSoup(response.text, "xml")
|
||||
links = [link.text for link in soup.find_all("loc") if link.parent.name == "url" and "/p/" in link.text]
|
||||
if len(links) == 0:
|
||||
links = [link.text for link in soup.find_all("loc") if "/p/" in link.text]
|
||||
|
||||
doc_id = hashlib.sha256((" ".join(links) + url).encode()).hexdigest()
|
||||
|
||||
def serialize_response(soup: BeautifulSoup):
|
||||
data = {}
|
||||
|
||||
h1_el = soup.find("h1")
|
||||
if h1_el is not None:
|
||||
data["title"] = h1_el.text
|
||||
|
||||
description_el = soup.find("meta", {"name": "description"})
|
||||
if description_el is not None:
|
||||
data["description"] = description_el["content"]
|
||||
|
||||
content_el = soup.find("div", {"id": "content-blocks"})
|
||||
if content_el is not None:
|
||||
data["content"] = content_el.text
|
||||
|
||||
return data
|
||||
|
||||
def load_link(link: str):
|
||||
try:
|
||||
beehiiv_data = requests.get(link, headers=headers)
|
||||
beehiiv_data.raise_for_status()
|
||||
|
||||
soup = BeautifulSoup(beehiiv_data.text, "html.parser")
|
||||
data = serialize_response(soup)
|
||||
data = str(data)
|
||||
if is_readable(data):
|
||||
return data
|
||||
else:
|
||||
logging.warning(f"Page is not readable (too many invalid characters): {link}")
|
||||
except ParserRejectedMarkup as e:
|
||||
logging.error(f"Failed to parse {link}: {e}")
|
||||
return None
|
||||
|
||||
for link in links:
|
||||
data = load_link(link)
|
||||
if data:
|
||||
output.append({"content": data, "meta_data": {"url": link}})
|
||||
# TODO: allow users to configure this
|
||||
time.sleep(1.0) # added to avoid rate limiting
|
||||
|
||||
return {"doc_id": doc_id, "data": output}
|
||||
@@ -196,7 +196,6 @@ class GithubLoader(BaseLoader):
|
||||
logging.info(f"Total repos found: {repos_results.totalCount}")
|
||||
for repo_result in tqdm(repos_results, total=repos_results.totalCount, desc="Loading discussions from github"):
|
||||
teams = repo_result.get_teams()
|
||||
# import pdb; pdb.set_trace()
|
||||
for team in teams:
|
||||
team_discussions = team.get_discussions()
|
||||
for discussion in team_discussions:
|
||||
|
||||
@@ -0,0 +1,52 @@
|
||||
import hashlib
|
||||
|
||||
from embedchain.helper.json_serializable import register_deserializable
|
||||
from embedchain.loaders.base_loader import BaseLoader
|
||||
|
||||
|
||||
@register_deserializable
|
||||
class RSSFeedLoader(BaseLoader):
|
||||
"""Loader for RSS Feed."""
|
||||
|
||||
def load_data(self, url):
|
||||
"""Load data from a rss feed."""
|
||||
output = self.get_rss_content(url)
|
||||
doc_id = hashlib.sha256((str(output) + url).encode()).hexdigest()
|
||||
return {
|
||||
"doc_id": doc_id,
|
||||
"data": output,
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def serialize_metadata(metadata):
|
||||
for key, value in metadata.items():
|
||||
if not isinstance(value, (str, int, float, bool)):
|
||||
metadata[key] = str(value)
|
||||
|
||||
return metadata
|
||||
|
||||
@staticmethod
|
||||
def get_rss_content(url: str):
|
||||
try:
|
||||
from langchain.document_loaders import RSSFeedLoader as LangchainRSSFeedLoader
|
||||
except ImportError:
|
||||
raise ImportError(
|
||||
"""RSSFeedLoader file requires extra dependencies.
|
||||
Install with `pip install --upgrade "embedchain[rss_feed]"`"""
|
||||
) from None
|
||||
|
||||
output = []
|
||||
loader = LangchainRSSFeedLoader(urls=[url])
|
||||
data = loader.load()
|
||||
|
||||
for entry in data:
|
||||
meta_data = RSSFeedLoader.serialize_metadata(entry.metadata)
|
||||
meta_data.update({"url": url})
|
||||
output.append(
|
||||
{
|
||||
"content": entry.page_content,
|
||||
"meta_data": meta_data,
|
||||
}
|
||||
)
|
||||
|
||||
return output
|
||||
@@ -1,6 +1,7 @@
|
||||
import concurrent.futures
|
||||
import hashlib
|
||||
import logging
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import requests
|
||||
from tqdm import tqdm
|
||||
@@ -16,7 +17,6 @@ except ImportError:
|
||||
from embedchain.helpers.json_serializable import register_deserializable
|
||||
from embedchain.loaders.base_loader import BaseLoader
|
||||
from embedchain.loaders.web_page import WebPageLoader
|
||||
from embedchain.utils import is_readable
|
||||
|
||||
|
||||
@register_deserializable
|
||||
@@ -30,29 +30,32 @@ class SitemapLoader(BaseLoader):
|
||||
def load_data(self, sitemap_url):
|
||||
output = []
|
||||
web_page_loader = WebPageLoader()
|
||||
response = requests.get(sitemap_url)
|
||||
response.raise_for_status()
|
||||
|
||||
soup = BeautifulSoup(response.text, "xml")
|
||||
if urlparse(sitemap_url).scheme not in ["file", "http", "https"]:
|
||||
raise ValueError("Not a valid URL.")
|
||||
|
||||
if urlparse(sitemap_url).scheme in ["http", "https"]:
|
||||
response = requests.get(sitemap_url)
|
||||
response.raise_for_status()
|
||||
else:
|
||||
with open(sitemap_url, "r") as file:
|
||||
soup = BeautifulSoup(file, "xml")
|
||||
links = [link.text for link in soup.find_all("loc") if link.parent.name == "url"]
|
||||
if len(links) == 0:
|
||||
links = [link.text for link in soup.find_all("loc")]
|
||||
|
||||
doc_id = hashlib.sha256((" ".join(links) + sitemap_url).encode()).hexdigest()
|
||||
|
||||
def load_link(link):
|
||||
def load_web_page(link):
|
||||
try:
|
||||
each_load_data = web_page_loader.load_data(link)
|
||||
if is_readable(each_load_data.get("data")[0].get("content")):
|
||||
return each_load_data.get("data")
|
||||
else:
|
||||
logging.warning(f"Page is not readable (too many invalid characters): {link}")
|
||||
loader_data = web_page_loader.load_data(link)
|
||||
return loader_data.get("data")
|
||||
except ParserRejectedMarkup as e:
|
||||
logging.error(f"Failed to parse {link}: {e}")
|
||||
return None
|
||||
|
||||
with concurrent.futures.ThreadPoolExecutor() as executor:
|
||||
future_to_link = {executor.submit(load_link, link): link for link in links}
|
||||
future_to_link = {executor.submit(load_web_page, link): link for link in links}
|
||||
for future in tqdm(concurrent.futures.as_completed(future_to_link), total=len(links), desc="Loading pages"):
|
||||
link = future_to_link[future]
|
||||
try:
|
||||
|
||||
@@ -3,7 +3,7 @@ import logging
|
||||
import time
|
||||
|
||||
import requests
|
||||
|
||||
from xml.etree import ElementTree
|
||||
from embedchain.helpers.json_serializable import register_deserializable
|
||||
from embedchain.loaders.base_loader import BaseLoader
|
||||
from embedchain.utils import is_readable
|
||||
@@ -12,9 +12,7 @@ from embedchain.utils import is_readable
|
||||
@register_deserializable
|
||||
class SubstackLoader(BaseLoader):
|
||||
"""
|
||||
This method takes a sitemap URL as input and retrieves
|
||||
all the URLs to use the WebPageLoader to load content
|
||||
of each page.
|
||||
This loader is used to load data from Substack URLs.
|
||||
"""
|
||||
|
||||
def load_data(self, url: str):
|
||||
@@ -26,9 +24,29 @@ class SubstackLoader(BaseLoader):
|
||||
'Substack requires extra dependencies. Install with `pip install --upgrade "embedchain[dataloaders]"`'
|
||||
) from None
|
||||
|
||||
if not url.endswith("sitemap.xml"):
|
||||
url = url + "/sitemap.xml"
|
||||
|
||||
output = []
|
||||
response = requests.get(url)
|
||||
response.raise_for_status()
|
||||
|
||||
try:
|
||||
response.raise_for_status()
|
||||
except requests.exceptions.HTTPError as e:
|
||||
raise ValueError(
|
||||
f"""
|
||||
Failed to load {url}: {e}. Please use the root substack URL. For example, https://example.substack.com
|
||||
"""
|
||||
)
|
||||
|
||||
try:
|
||||
ElementTree.fromstring(response.content)
|
||||
except ElementTree.ParseError:
|
||||
raise ValueError(
|
||||
f"""
|
||||
Failed to parse {url}. Please use the root substack URL. For example, https://example.substack.com
|
||||
"""
|
||||
)
|
||||
|
||||
soup = BeautifulSoup(response.text, "xml")
|
||||
links = [link.text for link in soup.find_all("loc") if link.parent.name == "url" and "/p/" in link.text]
|
||||
@@ -62,10 +80,10 @@ class SubstackLoader(BaseLoader):
|
||||
|
||||
def load_link(link: str):
|
||||
try:
|
||||
each_load_data = requests.get(link)
|
||||
each_load_data.raise_for_status()
|
||||
substack_data = requests.get(link)
|
||||
substack_data.raise_for_status()
|
||||
|
||||
soup = BeautifulSoup(response.text, "html.parser")
|
||||
soup = BeautifulSoup(substack_data.text, "html.parser")
|
||||
data = serialize_response(soup)
|
||||
data = str(data)
|
||||
if is_readable(data):
|
||||
|
||||
@@ -33,6 +33,8 @@ class IndirectDataType(Enum):
|
||||
YOUTUBE_CHANNEL = "youtube_channel"
|
||||
DISCORD = "discord"
|
||||
CUSTOM = "custom"
|
||||
RSSFEED = "rss_feed"
|
||||
BEEHIIV = "beehiiv"
|
||||
|
||||
|
||||
class SpecialDataType(Enum):
|
||||
@@ -65,3 +67,5 @@ class DataType(Enum):
|
||||
YOUTUBE_CHANNEL = IndirectDataType.YOUTUBE_CHANNEL.value
|
||||
DISCORD = IndirectDataType.DISCORD.value
|
||||
CUSTOM = IndirectDataType.CUSTOM.value
|
||||
RSSFEED = IndirectDataType.RSSFEED.value
|
||||
BEEHIIV = IndirectDataType.BEEHIIV.value
|
||||
|
||||
+16
-2
@@ -1,3 +1,4 @@
|
||||
import itertools
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
@@ -6,6 +7,7 @@ import string
|
||||
from typing import Any
|
||||
|
||||
from schema import Optional, Or, Schema
|
||||
from tqdm import tqdm
|
||||
|
||||
from embedchain.models.data_type import DataType
|
||||
|
||||
@@ -194,8 +196,7 @@ def detect_datatype(source: Any) -> DataType:
|
||||
formatted_source = format_source(str(source), 30)
|
||||
|
||||
if url:
|
||||
from langchain.document_loaders.youtube import \
|
||||
ALLOWED_NETLOCK as YOUTUBE_ALLOWED_NETLOCS
|
||||
from langchain.document_loaders.youtube import ALLOWED_NETLOCK as YOUTUBE_ALLOWED_NETLOCS
|
||||
|
||||
if url.netloc in YOUTUBE_ALLOWED_NETLOCS:
|
||||
logging.debug(f"Source of `{formatted_source}` detected as `youtube_video`.")
|
||||
@@ -422,3 +423,16 @@ def validate_config(config_data):
|
||||
)
|
||||
|
||||
return schema.validate(config_data)
|
||||
|
||||
|
||||
def chunks(iterable, batch_size=100, desc="Processing chunks"):
|
||||
"""A helper function to break an iterable into chunks of size batch_size."""
|
||||
it = iter(iterable)
|
||||
total_size = len(iterable)
|
||||
|
||||
with tqdm(total=total_size, desc=desc, unit="batch") as pbar:
|
||||
chunk = tuple(itertools.islice(it, batch_size))
|
||||
while chunk:
|
||||
yield chunk
|
||||
pbar.update(len(chunk))
|
||||
chunk = tuple(itertools.islice(it, batch_size))
|
||||
|
||||
@@ -133,6 +133,7 @@ class ChromaDB(BaseVectorDB):
|
||||
metadatas: List[object],
|
||||
ids: List[str],
|
||||
skip_embedding: bool,
|
||||
**kwargs: Optional[Dict[str, Any]],
|
||||
) -> Any:
|
||||
"""
|
||||
Add vectors to chroma database
|
||||
@@ -198,6 +199,7 @@ class ChromaDB(BaseVectorDB):
|
||||
where: Dict[str, any],
|
||||
skip_embedding: bool,
|
||||
citations: bool = False,
|
||||
**kwargs: Optional[Dict[str, Any]],
|
||||
) -> Union[List[Tuple[str, str, str]], List[str]]:
|
||||
"""
|
||||
Query contents from vector database based on vector similarity
|
||||
@@ -225,6 +227,7 @@ class ChromaDB(BaseVectorDB):
|
||||
],
|
||||
n_results=n_results,
|
||||
where=self._generate_where_clause(where),
|
||||
**kwargs,
|
||||
)
|
||||
else:
|
||||
result = self.collection.query(
|
||||
@@ -233,6 +236,7 @@ class ChromaDB(BaseVectorDB):
|
||||
],
|
||||
n_results=n_results,
|
||||
where=self._generate_where_clause(where),
|
||||
**kwargs,
|
||||
)
|
||||
except InvalidDimensionException as e:
|
||||
raise InvalidDimensionException(
|
||||
|
||||
@@ -11,6 +11,7 @@ except ImportError:
|
||||
|
||||
from embedchain.config import ElasticsearchDBConfig
|
||||
from embedchain.helpers.json_serializable import register_deserializable
|
||||
from embedchain.utils import chunks
|
||||
from embedchain.vectordb.base import BaseVectorDB
|
||||
|
||||
|
||||
@@ -20,6 +21,8 @@ class ElasticsearchDB(BaseVectorDB):
|
||||
Elasticsearch as vector database
|
||||
"""
|
||||
|
||||
BATCH_SIZE = 100
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: Optional[ElasticsearchDBConfig] = None,
|
||||
@@ -43,7 +46,14 @@ class ElasticsearchDB(BaseVectorDB):
|
||||
"Please make sure the type is right and that you are passing an instance."
|
||||
)
|
||||
self.config = config or es_config
|
||||
self.client = Elasticsearch(self.config.ES_URL, **self.config.ES_EXTRA_PARAMS)
|
||||
if self.config.ES_URL:
|
||||
self.client = Elasticsearch(self.config.ES_URL, **self.config.ES_EXTRA_PARAMS)
|
||||
elif self.config.CLOUD_ID:
|
||||
self.client = Elasticsearch(cloud_id=self.config.CLOUD_ID, **self.config.ES_EXTRA_PARAMS)
|
||||
else:
|
||||
raise ValueError(
|
||||
"Something is wrong with your config. Please check again - `https://docs.embedchain.ai/components/vector-databases#elasticsearch`" # noqa: E501
|
||||
)
|
||||
|
||||
# Call parent init here because embedder is needed
|
||||
super().__init__(config=self.config)
|
||||
@@ -105,6 +115,7 @@ class ElasticsearchDB(BaseVectorDB):
|
||||
metadatas: List[object],
|
||||
ids: List[str],
|
||||
skip_embedding: bool,
|
||||
**kwargs: Optional[Dict[str, any]],
|
||||
) -> Any:
|
||||
"""
|
||||
add data in vector database
|
||||
@@ -120,19 +131,29 @@ class ElasticsearchDB(BaseVectorDB):
|
||||
:type skip_embedding: bool
|
||||
"""
|
||||
|
||||
docs = []
|
||||
if not skip_embedding:
|
||||
embeddings = self.embedder.embedding_fn(documents)
|
||||
|
||||
for id, text, metadata, embeddings in zip(ids, documents, metadatas, embeddings):
|
||||
docs.append(
|
||||
{
|
||||
"_index": self._get_index(),
|
||||
"_id": id,
|
||||
"_source": {"text": text, "metadata": metadata, "embeddings": embeddings},
|
||||
}
|
||||
)
|
||||
bulk(self.client, docs)
|
||||
for chunk in chunks(
|
||||
list(zip(ids, documents, metadatas, embeddings)), self.BATCH_SIZE, desc="Inserting batches in elasticsearch"
|
||||
): # noqa: E501
|
||||
ids, docs, metadatas, embeddings = [], [], [], []
|
||||
for id, text, metadata, embedding in chunk:
|
||||
ids.append(id)
|
||||
docs.append(text)
|
||||
metadatas.append(metadata)
|
||||
embeddings.append(embedding)
|
||||
|
||||
batch_docs = []
|
||||
for id, text, metadata, embedding in zip(ids, docs, metadatas, embeddings):
|
||||
batch_docs.append(
|
||||
{
|
||||
"_index": self._get_index(),
|
||||
"_id": id,
|
||||
"_source": {"text": text, "metadata": metadata, "embeddings": embedding},
|
||||
}
|
||||
)
|
||||
bulk(self.client, batch_docs, **kwargs)
|
||||
self.client.indices.refresh(index=self._get_index())
|
||||
|
||||
def query(
|
||||
@@ -142,6 +163,7 @@ class ElasticsearchDB(BaseVectorDB):
|
||||
where: Dict[str, any],
|
||||
skip_embedding: bool,
|
||||
citations: bool = False,
|
||||
**kwargs: Optional[Dict[str, Any]],
|
||||
) -> Union[List[Tuple[str, str, str]], List[str]]:
|
||||
"""
|
||||
query contents from vector data base based on vector similarity
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import logging
|
||||
import time
|
||||
from typing import Dict, List, Optional, Set, Tuple, Union
|
||||
from typing import Any, Dict, List, Optional, Set, Tuple, Union
|
||||
|
||||
from tqdm import tqdm
|
||||
|
||||
@@ -121,6 +121,7 @@ class OpenSearchDB(BaseVectorDB):
|
||||
metadatas: List[object],
|
||||
ids: List[str],
|
||||
skip_embedding: bool,
|
||||
**kwargs: Optional[Dict[str, any]],
|
||||
):
|
||||
"""Add data in vector database.
|
||||
|
||||
@@ -154,7 +155,7 @@ class OpenSearchDB(BaseVectorDB):
|
||||
]
|
||||
|
||||
# Perform bulk operation
|
||||
bulk(self.client, batch_entries)
|
||||
bulk(self.client, batch_entries, **kwargs)
|
||||
self.client.indices.refresh(index=self._get_index())
|
||||
|
||||
# Sleep to avoid rate limiting
|
||||
@@ -167,6 +168,7 @@ class OpenSearchDB(BaseVectorDB):
|
||||
where: Dict[str, any],
|
||||
skip_embedding: bool,
|
||||
citations: bool = False,
|
||||
**kwargs: Optional[Dict[str, Any]],
|
||||
) -> Union[List[Tuple[str, str, str]], List[str]]:
|
||||
"""
|
||||
query contents from vector data base based on vector similarity
|
||||
@@ -209,6 +211,7 @@ class OpenSearchDB(BaseVectorDB):
|
||||
metadata_field="metadata",
|
||||
pre_filter=pre_filter,
|
||||
k=n_results,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
contexts = []
|
||||
|
||||
@@ -10,6 +10,7 @@ except ImportError:
|
||||
|
||||
from embedchain.config.vectordb.pinecone import PineconeDBConfig
|
||||
from embedchain.helpers.json_serializable import register_deserializable
|
||||
from embedchain.utils import chunks
|
||||
from embedchain.vectordb.base import BaseVectorDB
|
||||
|
||||
|
||||
@@ -92,6 +93,7 @@ class PineconeDB(BaseVectorDB):
|
||||
metadatas: List[object],
|
||||
ids: List[str],
|
||||
skip_embedding: bool,
|
||||
**kwargs: Optional[Dict[str, any]],
|
||||
):
|
||||
"""add data in vector database
|
||||
|
||||
@@ -104,7 +106,6 @@ class PineconeDB(BaseVectorDB):
|
||||
"""
|
||||
docs = []
|
||||
print("Adding documents to Pinecone...")
|
||||
|
||||
embeddings = self.embedder.embedding_fn(documents)
|
||||
for id, text, metadata, embedding in zip(ids, documents, metadatas, embeddings):
|
||||
docs.append(
|
||||
@@ -115,8 +116,8 @@ class PineconeDB(BaseVectorDB):
|
||||
}
|
||||
)
|
||||
|
||||
for i in range(0, len(docs), self.BATCH_SIZE):
|
||||
self.client.upsert(docs[i : i + self.BATCH_SIZE])
|
||||
for chunk in chunks(docs, self.BATCH_SIZE, desc="Adding chunks in batches..."):
|
||||
self.client.upsert(chunk, **kwargs)
|
||||
|
||||
def query(
|
||||
self,
|
||||
@@ -125,6 +126,7 @@ class PineconeDB(BaseVectorDB):
|
||||
where: Dict[str, any],
|
||||
skip_embedding: bool,
|
||||
citations: bool = False,
|
||||
**kwargs: Optional[Dict[str, any]],
|
||||
) -> Union[List[Tuple[str, str, str]], List[str]]:
|
||||
"""
|
||||
query contents from vector database based on vector similarity
|
||||
@@ -146,7 +148,7 @@ class PineconeDB(BaseVectorDB):
|
||||
query_vector = self.embedder.embedding_fn([input_query])[0]
|
||||
else:
|
||||
query_vector = input_query
|
||||
data = self.client.query(vector=query_vector, filter=where, top_k=n_results, include_metadata=True)
|
||||
data = self.client.query(vector=query_vector, filter=where, top_k=n_results, include_metadata=True, **kwargs)
|
||||
contexts = []
|
||||
for doc in data["matches"]:
|
||||
metadata = doc["metadata"]
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import copy
|
||||
import os
|
||||
import uuid
|
||||
from typing import Dict, List, Optional, Tuple, Union
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
try:
|
||||
from qdrant_client import QdrantClient
|
||||
@@ -127,6 +127,7 @@ class QdrantDB(BaseVectorDB):
|
||||
metadatas: List[object],
|
||||
ids: List[str],
|
||||
skip_embedding: bool,
|
||||
**kwargs: Optional[Dict[str, any]],
|
||||
):
|
||||
"""add data in vector database
|
||||
:param embeddings: list of embeddings for the corresponding documents to be added
|
||||
@@ -158,6 +159,7 @@ class QdrantDB(BaseVectorDB):
|
||||
payloads=payloads[i : i + self.BATCH_SIZE],
|
||||
vectors=embeddings[i : i + self.BATCH_SIZE],
|
||||
),
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def query(
|
||||
@@ -167,6 +169,7 @@ class QdrantDB(BaseVectorDB):
|
||||
where: Dict[str, any],
|
||||
skip_embedding: bool,
|
||||
citations: bool = False,
|
||||
**kwargs: Optional[Dict[str, Any]],
|
||||
) -> Union[List[Tuple[str, str, str]], List[str]]:
|
||||
"""
|
||||
query contents from vector database based on vector similarity
|
||||
@@ -208,6 +211,7 @@ class QdrantDB(BaseVectorDB):
|
||||
query_filter=models.Filter(must=qdrant_must_filters),
|
||||
query_vector=query_vector,
|
||||
limit=n_results,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
contexts = []
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import copy
|
||||
import os
|
||||
from typing import Dict, List, Optional, Tuple, Union
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
try:
|
||||
import weaviate
|
||||
@@ -158,6 +158,7 @@ class WeaviateDB(BaseVectorDB):
|
||||
metadatas: List[object],
|
||||
ids: List[str],
|
||||
skip_embedding: bool,
|
||||
**kwargs: Optional[Dict[str, any]],
|
||||
):
|
||||
"""add data in vector database
|
||||
:param embeddings: list of embeddings for the corresponding documents to be added
|
||||
@@ -192,7 +193,9 @@ class WeaviateDB(BaseVectorDB):
|
||||
class_name=self.index_name + "_metadata",
|
||||
vector=embedding,
|
||||
)
|
||||
batch.add_reference(obj_uuid, self.index_name, "metadata", metadata_uuid, self.index_name + "_metadata")
|
||||
batch.add_reference(
|
||||
obj_uuid, self.index_name, "metadata", metadata_uuid, self.index_name + "_metadata", **kwargs
|
||||
)
|
||||
|
||||
def query(
|
||||
self,
|
||||
@@ -201,6 +204,7 @@ class WeaviateDB(BaseVectorDB):
|
||||
where: Dict[str, any],
|
||||
skip_embedding: bool,
|
||||
citations: bool = False,
|
||||
**kwargs: Optional[Dict[str, Any]],
|
||||
) -> Union[List[Tuple[str, str, str]], List[str]]:
|
||||
"""
|
||||
query contents from vector database based on vector similarity
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
import logging
|
||||
from typing import Dict, List, Optional, Tuple, Union
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
from embedchain.config import ZillizDBConfig
|
||||
from embedchain.helpers.json_serializable import register_deserializable
|
||||
@@ -113,6 +113,7 @@ class ZillizVectorDB(BaseVectorDB):
|
||||
metadatas: List[object],
|
||||
ids: List[str],
|
||||
skip_embedding: bool,
|
||||
**kwargs: Optional[Dict[str, any]],
|
||||
):
|
||||
"""Add to database"""
|
||||
if not skip_embedding:
|
||||
@@ -120,7 +121,7 @@ class ZillizVectorDB(BaseVectorDB):
|
||||
|
||||
for id, doc, metadata, embedding in zip(ids, documents, metadatas, embeddings):
|
||||
data = {**metadata, "id": id, "text": doc, "embeddings": embedding}
|
||||
self.client.insert(collection_name=self.config.collection_name, data=data)
|
||||
self.client.insert(collection_name=self.config.collection_name, data=data, **kwargs)
|
||||
|
||||
self.collection.load()
|
||||
self.collection.flush()
|
||||
@@ -133,6 +134,7 @@ class ZillizVectorDB(BaseVectorDB):
|
||||
where: Dict[str, any],
|
||||
skip_embedding: bool,
|
||||
citations: bool = False,
|
||||
**kwargs: Optional[Dict[str, Any]],
|
||||
) -> Union[List[Tuple[str, str, str]], List[str]]:
|
||||
"""
|
||||
Query contents from vector data base based on vector similarity
|
||||
@@ -165,6 +167,7 @@ class ZillizVectorDB(BaseVectorDB):
|
||||
data=query_vector,
|
||||
limit=n_results,
|
||||
output_fields=output_fields,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
else:
|
||||
@@ -176,6 +179,7 @@ class ZillizVectorDB(BaseVectorDB):
|
||||
data=[query_vector],
|
||||
limit=n_results,
|
||||
output_fields=output_fields,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
contexts = []
|
||||
|
||||
Generated
+158
-13
@@ -1,4 +1,4 @@
|
||||
# This file is automatically @generated by Poetry 1.5.1 and should not be changed by hand.
|
||||
# This file is automatically @generated by Poetry 1.6.1 and should not be changed by hand.
|
||||
|
||||
[[package]]
|
||||
name = "aiofiles"
|
||||
@@ -1083,6 +1083,17 @@ ssh = ["bcrypt (>=3.1.5)"]
|
||||
test = ["pretend", "pytest (>=6.2.0)", "pytest-benchmark", "pytest-cov", "pytest-xdist"]
|
||||
test-randomorder = ["pytest-randomly"]
|
||||
|
||||
[[package]]
|
||||
name = "cssselect"
|
||||
version = "1.2.0"
|
||||
description = "cssselect parses CSS3 Selectors and translates them to XPath 1.0"
|
||||
optional = true
|
||||
python-versions = ">=3.7"
|
||||
files = [
|
||||
{file = "cssselect-1.2.0-py2.py3-none-any.whl", hash = "sha256:da1885f0c10b60c03ed5eccbb6b68d6eff248d91976fcde348f395d54c9fd35e"},
|
||||
{file = "cssselect-1.2.0.tar.gz", hash = "sha256:666b19839cfaddb9ce9d36bfe4c969132c647b92fc9088c4e23f786b30f1b3dc"},
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "cycler"
|
||||
version = "0.12.1"
|
||||
@@ -1446,6 +1457,35 @@ lz4 = ["lz4"]
|
||||
snappy = ["python-snappy"]
|
||||
zstandard = ["zstandard"]
|
||||
|
||||
[[package]]
|
||||
name = "feedfinder2"
|
||||
version = "0.0.4"
|
||||
description = "Find the feed URLs for a website."
|
||||
optional = true
|
||||
python-versions = "*"
|
||||
files = [
|
||||
{file = "feedfinder2-0.0.4.tar.gz", hash = "sha256:3701ee01a6c85f8b865a049c30ba0b4608858c803fe8e30d1d289fdbe89d0efe"},
|
||||
]
|
||||
|
||||
[package.dependencies]
|
||||
beautifulsoup4 = "*"
|
||||
requests = "*"
|
||||
six = "*"
|
||||
|
||||
[[package]]
|
||||
name = "feedparser"
|
||||
version = "6.0.10"
|
||||
description = "Universal feed parser, handles RSS 0.9x, RSS 1.0, RSS 2.0, CDF, Atom 0.3, and Atom 1.0 feeds"
|
||||
optional = true
|
||||
python-versions = ">=3.6"
|
||||
files = [
|
||||
{file = "feedparser-6.0.10-py3-none-any.whl", hash = "sha256:79c257d526d13b944e965f6095700587f27388e50ea16fd245babe4dfae7024f"},
|
||||
{file = "feedparser-6.0.10.tar.gz", hash = "sha256:27da485f4637ce7163cdeab13a80312b93b7d0c1b775bef4a47629a3110bca51"},
|
||||
]
|
||||
|
||||
[package.dependencies]
|
||||
sgmllib3k = "*"
|
||||
|
||||
[[package]]
|
||||
name = "filelock"
|
||||
version = "3.12.4"
|
||||
@@ -1737,12 +1777,12 @@ files = [
|
||||
google-auth = ">=2.14.1,<3.0.dev0"
|
||||
googleapis-common-protos = ">=1.56.2,<2.0.dev0"
|
||||
grpcio = [
|
||||
{version = ">=1.33.2,<2.0dev", optional = true, markers = "extra == \"grpc\""},
|
||||
{version = ">=1.49.1,<2.0dev", optional = true, markers = "python_version >= \"3.11\" and extra == \"grpc\""},
|
||||
{version = ">=1.33.2,<2.0dev", optional = true, markers = "python_version < \"3.11\" and extra == \"grpc\""},
|
||||
]
|
||||
grpcio-status = [
|
||||
{version = ">=1.33.2,<2.0.dev0", optional = true, markers = "extra == \"grpc\""},
|
||||
{version = ">=1.49.1,<2.0.dev0", optional = true, markers = "python_version >= \"3.11\" and extra == \"grpc\""},
|
||||
{version = ">=1.33.2,<2.0.dev0", optional = true, markers = "python_version < \"3.11\" and extra == \"grpc\""},
|
||||
]
|
||||
protobuf = ">=3.19.5,<3.20.0 || >3.20.0,<3.20.1 || >3.20.1,<4.21.0 || >4.21.0,<4.21.1 || >4.21.1,<4.21.2 || >4.21.2,<4.21.3 || >4.21.3,<4.21.4 || >4.21.4,<4.21.5 || >4.21.5,<5.0.0.dev0"
|
||||
requests = ">=2.18.0,<3.0.0.dev0"
|
||||
@@ -1830,8 +1870,8 @@ google-api-core = {version = ">=1.31.5,<2.0.dev0 || >2.3.0,<3.0.0dev", extras =
|
||||
google-cloud-core = ">=1.6.0,<3.0.0dev"
|
||||
google-resumable-media = ">=0.6.0,<3.0dev"
|
||||
grpcio = [
|
||||
{version = ">=1.47.0,<2.0dev", markers = "python_version < \"3.11\""},
|
||||
{version = ">=1.49.1,<2.0dev", markers = "python_version >= \"3.11\""},
|
||||
{version = ">=1.47.0,<2.0dev", markers = "python_version < \"3.11\""},
|
||||
]
|
||||
packaging = ">=20.0.0"
|
||||
proto-plus = ">=1.15.0,<2.0.0dev"
|
||||
@@ -1882,8 +1922,8 @@ files = [
|
||||
google-api-core = {version = ">=1.34.0,<2.0.dev0 || >=2.11.dev0,<3.0.0dev", extras = ["grpc"]}
|
||||
grpc-google-iam-v1 = ">=0.12.4,<1.0.0dev"
|
||||
proto-plus = [
|
||||
{version = ">=1.22.0,<2.0.0dev", markers = "python_version < \"3.11\""},
|
||||
{version = ">=1.22.2,<2.0.0dev", markers = "python_version >= \"3.11\""},
|
||||
{version = ">=1.22.0,<2.0.0dev", markers = "python_version < \"3.11\""},
|
||||
]
|
||||
protobuf = ">=3.19.5,<3.20.0 || >3.20.0,<3.20.1 || >3.20.1,<4.21.0 || >4.21.0,<4.21.1 || >4.21.1,<4.21.2 || >4.21.2,<4.21.3 || >4.21.3,<4.21.4 || >4.21.4,<4.21.5 || >4.21.5,<5.0.0dev"
|
||||
|
||||
@@ -2598,6 +2638,16 @@ files = [
|
||||
{file = "itsdangerous-2.1.2.tar.gz", hash = "sha256:5dbbc68b317e5e42f327f9021763545dc3fc3bfe22e6deb96aaf1fc38874156a"},
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "jieba3k"
|
||||
version = "0.35.1"
|
||||
description = "Chinese Words Segementation Utilities"
|
||||
optional = true
|
||||
python-versions = "*"
|
||||
files = [
|
||||
{file = "jieba3k-0.35.1.zip", hash = "sha256:980a4f2636b778d312518066be90c7697d410dd5a472385f5afced71a2db1c10"},
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "jinja2"
|
||||
version = "3.1.2"
|
||||
@@ -2893,6 +2943,21 @@ ocr = ["google-cloud-vision (==1)", "pytesseract"]
|
||||
paddledetection = ["paddlepaddle (==2.1.0)"]
|
||||
tesseract = ["pytesseract"]
|
||||
|
||||
[[package]]
|
||||
name = "listparser"
|
||||
version = "0.19"
|
||||
description = "Parse OPML subscription lists"
|
||||
optional = true
|
||||
python-versions = ">=3.7,<4.0"
|
||||
files = [
|
||||
{file = "listparser-0.19-py3-none-any.whl", hash = "sha256:c3857a9e5e5342207a556ba72e5c030782971fbe587e7afc2a75b2d7c0fa5a5c"},
|
||||
{file = "listparser-0.19.tar.gz", hash = "sha256:5aa23ae017a22e36c50ca5259a690328dd524527977d8c094ae0857887002805"},
|
||||
]
|
||||
|
||||
[package.extras]
|
||||
http = ["requests (>=2.25.1,<3.0.0)"]
|
||||
lxml = ["lxml (>=4.6.2,<5.0.0)"]
|
||||
|
||||
[[package]]
|
||||
name = "lit"
|
||||
version = "17.0.2"
|
||||
@@ -3505,6 +3570,32 @@ doc = ["nb2plots (>=0.6)", "numpydoc (>=1.5)", "pillow (>=9.4)", "pydata-sphinx-
|
||||
extra = ["lxml (>=4.6)", "pydot (>=1.4.2)", "pygraphviz (>=1.10)", "sympy (>=1.10)"]
|
||||
test = ["codecov (>=2.1)", "pytest (>=7.2)", "pytest-cov (>=4.0)"]
|
||||
|
||||
[[package]]
|
||||
name = "newspaper3k"
|
||||
version = "0.2.8"
|
||||
description = "Simplified python article discovery & extraction."
|
||||
optional = true
|
||||
python-versions = "*"
|
||||
files = [
|
||||
{file = "newspaper3k-0.2.8-py3-none-any.whl", hash = "sha256:44a864222633d3081113d1030615991c3dbba87239f6bbf59d91240f71a22e3e"},
|
||||
{file = "newspaper3k-0.2.8.tar.gz", hash = "sha256:9f1bd3e1fb48f400c715abf875cc7b0a67b7ddcd87f50c9aeeb8fcbbbd9004fb"},
|
||||
]
|
||||
|
||||
[package.dependencies]
|
||||
beautifulsoup4 = ">=4.4.1"
|
||||
cssselect = ">=0.9.2"
|
||||
feedfinder2 = ">=0.0.4"
|
||||
feedparser = ">=5.2.1"
|
||||
jieba3k = ">=0.35.1"
|
||||
lxml = ">=3.6.0"
|
||||
nltk = ">=3.2.1"
|
||||
Pillow = ">=3.3.0"
|
||||
python-dateutil = ">=2.5.3"
|
||||
PyYAML = ">=3.11"
|
||||
requests = ">=2.10.0"
|
||||
tinysegmenter = "0.3"
|
||||
tldextract = ">=2.0.1"
|
||||
|
||||
[[package]]
|
||||
name = "nltk"
|
||||
version = "3.8.1"
|
||||
@@ -3913,13 +4004,11 @@ files = [
|
||||
|
||||
[package.dependencies]
|
||||
numpy = [
|
||||
{version = ">=1.21.0", markers = "python_version <= \"3.9\" and platform_system == \"Darwin\" and platform_machine == \"arm64\""},
|
||||
{version = ">=1.21.2", markers = "python_version >= \"3.10\""},
|
||||
{version = ">=1.21.4", markers = "python_version >= \"3.10\" and platform_system == \"Darwin\""},
|
||||
{version = ">=1.19.3", markers = "python_version >= \"3.6\" and platform_system == \"Linux\" and platform_machine == \"aarch64\" or python_version >= \"3.9\""},
|
||||
{version = ">=1.17.0", markers = "python_version >= \"3.7\""},
|
||||
{version = ">=1.17.3", markers = "python_version >= \"3.8\""},
|
||||
{version = ">=1.23.5", markers = "python_version >= \"3.11\""},
|
||||
{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.21.0", markers = "python_version == \"3.9\" and platform_system == \"Darwin\" and platform_machine == \"arm64\""},
|
||||
]
|
||||
|
||||
[[package]]
|
||||
@@ -4113,8 +4202,8 @@ files = [
|
||||
|
||||
[package.dependencies]
|
||||
numpy = [
|
||||
{version = ">=1.22.4", markers = "python_version < \"3.11\""},
|
||||
{version = ">=1.23.2", markers = "python_version == \"3.11\""},
|
||||
{version = ">=1.22.4", markers = "python_version < \"3.11\""},
|
||||
]
|
||||
python-dateutil = ">=2.8.2"
|
||||
pytz = ">=2020.1"
|
||||
@@ -5617,6 +5706,21 @@ urllib3 = ">=1.21.1,<3"
|
||||
socks = ["PySocks (>=1.5.6,!=1.5.7)"]
|
||||
use-chardet-on-py3 = ["chardet (>=3.0.2,<6)"]
|
||||
|
||||
[[package]]
|
||||
name = "requests-file"
|
||||
version = "1.5.1"
|
||||
description = "File transport adapter for Requests"
|
||||
optional = true
|
||||
python-versions = "*"
|
||||
files = [
|
||||
{file = "requests-file-1.5.1.tar.gz", hash = "sha256:07d74208d3389d01c38ab89ef403af0cfec63957d53a0081d8eca738d0247d8e"},
|
||||
{file = "requests_file-1.5.1-py2.py3-none-any.whl", hash = "sha256:dfe5dae75c12481f68ba353183c53a65e6044c923e64c24b2209f6c7570ca953"},
|
||||
]
|
||||
|
||||
[package.dependencies]
|
||||
requests = ">=1.0.0"
|
||||
six = "*"
|
||||
|
||||
[[package]]
|
||||
name = "requests-oauthlib"
|
||||
version = "1.3.1"
|
||||
@@ -6044,6 +6148,16 @@ docs = ["entangled-cli[rich]", "mkdocs", "mkdocs-entangled-plugin", "mkdocs-mate
|
||||
rich = ["rich"]
|
||||
test = ["build", "pytest", "rich", "wheel"]
|
||||
|
||||
[[package]]
|
||||
name = "sgmllib3k"
|
||||
version = "1.0.0"
|
||||
description = "Py3k port of sgmllib."
|
||||
optional = true
|
||||
python-versions = "*"
|
||||
files = [
|
||||
{file = "sgmllib3k-1.0.0.tar.gz", hash = "sha256:7868fb1c8bfa764c1ac563d3cf369c381d1325d36124933a726f29fcdaa812e9"},
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "shapely"
|
||||
version = "2.0.2"
|
||||
@@ -6230,7 +6344,7 @@ files = [
|
||||
]
|
||||
|
||||
[package.dependencies]
|
||||
greenlet = {version = "!=0.4.17", optional = true, markers = "platform_machine == \"win32\" or platform_machine == \"WIN32\" or platform_machine == \"AMD64\" or platform_machine == \"amd64\" or platform_machine == \"x86_64\" or platform_machine == \"ppc64le\" or platform_machine == \"aarch64\" or extra == \"asyncio\""}
|
||||
greenlet = {version = "!=0.4.17", optional = true, markers = "platform_machine == \"aarch64\" or platform_machine == \"ppc64le\" or platform_machine == \"x86_64\" or platform_machine == \"amd64\" or platform_machine == \"AMD64\" or platform_machine == \"win32\" or platform_machine == \"WIN32\" or extra == \"asyncio\""}
|
||||
typing-extensions = ">=4.2.0"
|
||||
|
||||
[package.extras]
|
||||
@@ -6405,6 +6519,36 @@ safetensors = "*"
|
||||
torch = ">=1.7"
|
||||
torchvision = "*"
|
||||
|
||||
[[package]]
|
||||
name = "tinysegmenter"
|
||||
version = "0.3"
|
||||
description = "Very compact Japanese tokenizer"
|
||||
optional = true
|
||||
python-versions = "*"
|
||||
files = [
|
||||
{file = "tinysegmenter-0.3.tar.gz", hash = "sha256:ed1f6d2e806a4758a73be589754384cbadadc7e1a414c81a166fc9adf2d40c6d"},
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "tldextract"
|
||||
version = "5.1.0"
|
||||
description = "Accurately separates a URL's subdomain, domain, and public suffix, using the Public Suffix List (PSL). By default, this includes the public ICANN TLDs and their exceptions. You can optionally support the Public Suffix List's private domains as well."
|
||||
optional = true
|
||||
python-versions = ">=3.8"
|
||||
files = [
|
||||
{file = "tldextract-5.1.0-py3-none-any.whl", hash = "sha256:c8eecb15f556b43db6eebd21667640fb6fba9bc9539b48707432014913a78d13"},
|
||||
{file = "tldextract-5.1.0.tar.gz", hash = "sha256:366acfb099c7eb5dc83545c391d73da6e3afe4eaec652417c3cf13b002a160e1"},
|
||||
]
|
||||
|
||||
[package.dependencies]
|
||||
filelock = ">=3.0.8"
|
||||
idna = "*"
|
||||
requests = ">=2.1.0"
|
||||
requests-file = ">=1.4"
|
||||
|
||||
[package.extras]
|
||||
testing = ["black", "mypy", "pytest", "pytest-gitignore", "pytest-mock", "responses", "ruff", "tox", "types-filelock", "types-requests"]
|
||||
|
||||
[[package]]
|
||||
name = "tokenizers"
|
||||
version = "0.14.1"
|
||||
@@ -7688,6 +7832,7 @@ pinecone = ["pinecone-client"]
|
||||
poe = ["fastapi-poe"]
|
||||
postgres = ["psycopg", "psycopg-binary", "psycopg-pool"]
|
||||
qdrant = ["qdrant-client"]
|
||||
rss-feed = ["feedparser", "listparser", "newspaper3k"]
|
||||
slack = ["flask", "slack-sdk"]
|
||||
streamlit = []
|
||||
vertexai = ["google-cloud-aiplatform"]
|
||||
|
||||
+5
-1
@@ -1,6 +1,6 @@
|
||||
[tool.poetry]
|
||||
name = "embedchain"
|
||||
version = "0.1.26"
|
||||
version = "0.1.30"
|
||||
description = "Data platform for LLMs - Load, index, retrieve and sync any unstructured data"
|
||||
authors = [
|
||||
"Taranjeet Singh <taranjeet@embedchain.ai>",
|
||||
@@ -137,6 +137,9 @@ mysql-connector-python = { version = "^8.1.0", optional = true }
|
||||
gitpython = { version = "^3.1.38", optional = true }
|
||||
yt_dlp = { version = "^2023.11.14", optional = true }
|
||||
PyGithub = { version = "^1.59.1", optional = true }
|
||||
feedparser = { version = "^6.0.10", optional = true }
|
||||
newspaper3k = { version = "^0.2.8", optional = true }
|
||||
listparser = { version = "^0.19", optional = true }
|
||||
|
||||
[tool.poetry.group.dev.dependencies]
|
||||
black = "^23.3.0"
|
||||
@@ -198,6 +201,7 @@ youtube = [
|
||||
"yt_dlp",
|
||||
"youtube-transcript-api",
|
||||
]
|
||||
rss_feed = ["feedparser", "listparser", "newspaper3k"]
|
||||
|
||||
[tool.poetry.group.docs.dependencies]
|
||||
|
||||
|
||||
@@ -57,11 +57,11 @@ class TestPinecone:
|
||||
db.add(vectors, documents, metadatas, ids, True)
|
||||
|
||||
expected_pinecone_upsert_args = [
|
||||
{"id": "doc1", "metadata": {"text": "This is a document."}, "values": [0, 0, 0]},
|
||||
{"id": "doc2", "metadata": {"text": "This is another document."}, "values": [1, 1, 1]},
|
||||
{"id": "doc1", "values": [0, 0, 0], "metadata": {"text": "This is a document."}},
|
||||
{"id": "doc2", "values": [1, 1, 1], "metadata": {"text": "This is another document."}},
|
||||
]
|
||||
# Assert that the Pinecone client was called to upsert the documents
|
||||
pinecone_client_mock.upsert.assert_called_once_with(expected_pinecone_upsert_args)
|
||||
pinecone_client_mock.upsert.assert_called_once_with(tuple(expected_pinecone_upsert_args))
|
||||
|
||||
@patch("embedchain.vectordb.pinecone.pinecone")
|
||||
def test_query_documents(self, pinecone_mock):
|
||||
|
||||
Reference in New Issue
Block a user