AutoRAG 0.0.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (184) hide show
  1. autorag/__init__.py +82 -0
  2. autorag/chunker.py +51 -0
  3. autorag/cli.py +209 -0
  4. autorag/dashboard.py +199 -0
  5. autorag/data/__init__.py +109 -0
  6. autorag/data/chunk/__init__.py +2 -0
  7. autorag/data/chunk/base.py +128 -0
  8. autorag/data/chunk/langchain_chunk.py +76 -0
  9. autorag/data/chunk/llama_index_chunk.py +96 -0
  10. autorag/data/chunk/run.py +38 -0
  11. autorag/data/legacy/__init__.py +0 -0
  12. autorag/data/legacy/corpus/__init__.py +2 -0
  13. autorag/data/legacy/corpus/langchain.py +47 -0
  14. autorag/data/legacy/corpus/llama_index.py +93 -0
  15. autorag/data/legacy/qacreation/__init__.py +6 -0
  16. autorag/data/legacy/qacreation/base.py +239 -0
  17. autorag/data/legacy/qacreation/llama_index.py +253 -0
  18. autorag/data/legacy/qacreation/llama_index_default_prompt.txt +54 -0
  19. autorag/data/legacy/qacreation/ragas.py +75 -0
  20. autorag/data/legacy/qacreation/simple.py +99 -0
  21. autorag/data/parse/__init__.py +1 -0
  22. autorag/data/parse/base.py +79 -0
  23. autorag/data/parse/clova.py +194 -0
  24. autorag/data/parse/langchain_parse.py +87 -0
  25. autorag/data/parse/llamaparse.py +126 -0
  26. autorag/data/parse/run.py +141 -0
  27. autorag/data/parse/table_hybrid_parse.py +134 -0
  28. autorag/data/qa/__init__.py +3 -0
  29. autorag/data/qa/evolve/__init__.py +0 -0
  30. autorag/data/qa/evolve/llama_index_query_evolve.py +64 -0
  31. autorag/data/qa/evolve/openai_query_evolve.py +81 -0
  32. autorag/data/qa/evolve/prompt.py +288 -0
  33. autorag/data/qa/extract_evidence.py +1 -0
  34. autorag/data/qa/filter/__init__.py +0 -0
  35. autorag/data/qa/filter/dontknow.py +117 -0
  36. autorag/data/qa/filter/passage_dependency.py +88 -0
  37. autorag/data/qa/filter/prompt.py +73 -0
  38. autorag/data/qa/generation_gt/__init__.py +0 -0
  39. autorag/data/qa/generation_gt/base.py +16 -0
  40. autorag/data/qa/generation_gt/llama_index_gen_gt.py +41 -0
  41. autorag/data/qa/generation_gt/openai_gen_gt.py +84 -0
  42. autorag/data/qa/generation_gt/prompt.py +27 -0
  43. autorag/data/qa/query/__init__.py +0 -0
  44. autorag/data/qa/query/llama_gen_query.py +82 -0
  45. autorag/data/qa/query/openai_gen_query.py +95 -0
  46. autorag/data/qa/query/prompt.py +201 -0
  47. autorag/data/qa/sample.py +26 -0
  48. autorag/data/qa/schema.py +322 -0
  49. autorag/data/utils/__init__.py +0 -0
  50. autorag/data/utils/util.py +103 -0
  51. autorag/deploy/__init__.py +9 -0
  52. autorag/deploy/api.py +303 -0
  53. autorag/deploy/base.py +235 -0
  54. autorag/deploy/gradio.py +74 -0
  55. autorag/deploy/swagger.yml +202 -0
  56. autorag/embedding/__init__.py +0 -0
  57. autorag/embedding/base.py +144 -0
  58. autorag/embedding/vllm.py +256 -0
  59. autorag/evaluation/__init__.py +3 -0
  60. autorag/evaluation/generation.py +88 -0
  61. autorag/evaluation/metric/__init__.py +22 -0
  62. autorag/evaluation/metric/deepeval_prompt.py +322 -0
  63. autorag/evaluation/metric/g_eval_prompts/coh_detailed.txt +32 -0
  64. autorag/evaluation/metric/g_eval_prompts/con_detailed.txt +33 -0
  65. autorag/evaluation/metric/g_eval_prompts/flu_detailed.txt +26 -0
  66. autorag/evaluation/metric/g_eval_prompts/rel_detailed.txt +33 -0
  67. autorag/evaluation/metric/generation.py +504 -0
  68. autorag/evaluation/metric/retrieval.py +115 -0
  69. autorag/evaluation/metric/retrieval_contents.py +65 -0
  70. autorag/evaluation/metric/util.py +88 -0
  71. autorag/evaluation/retrieval.py +83 -0
  72. autorag/evaluation/retrieval_contents.py +65 -0
  73. autorag/evaluation/util.py +43 -0
  74. autorag/evaluator.py +559 -0
  75. autorag/node_line.py +65 -0
  76. autorag/nodes/__init__.py +0 -0
  77. autorag/nodes/generator/__init__.py +4 -0
  78. autorag/nodes/generator/base.py +103 -0
  79. autorag/nodes/generator/llama_index_llm.py +169 -0
  80. autorag/nodes/generator/openai_llm.py +329 -0
  81. autorag/nodes/generator/run.py +148 -0
  82. autorag/nodes/generator/vllm.py +147 -0
  83. autorag/nodes/generator/vllm_api.py +191 -0
  84. autorag/nodes/hybridretrieval/__init__.py +2 -0
  85. autorag/nodes/hybridretrieval/base.py +58 -0
  86. autorag/nodes/hybridretrieval/hybrid_cc.py +227 -0
  87. autorag/nodes/hybridretrieval/hybrid_rrf.py +149 -0
  88. autorag/nodes/hybridretrieval/run.py +137 -0
  89. autorag/nodes/lexicalretrieval/__init__.py +1 -0
  90. autorag/nodes/lexicalretrieval/bm25.py +381 -0
  91. autorag/nodes/lexicalretrieval/run.py +148 -0
  92. autorag/nodes/passageaugmenter/__init__.py +2 -0
  93. autorag/nodes/passageaugmenter/base.py +76 -0
  94. autorag/nodes/passageaugmenter/pass_passage_augmenter.py +43 -0
  95. autorag/nodes/passageaugmenter/prev_next_augmenter.py +155 -0
  96. autorag/nodes/passageaugmenter/run.py +131 -0
  97. autorag/nodes/passagecompressor/__init__.py +4 -0
  98. autorag/nodes/passagecompressor/base.py +78 -0
  99. autorag/nodes/passagecompressor/longllmlingua.py +115 -0
  100. autorag/nodes/passagecompressor/pass_compressor.py +16 -0
  101. autorag/nodes/passagecompressor/refine.py +54 -0
  102. autorag/nodes/passagecompressor/run.py +186 -0
  103. autorag/nodes/passagecompressor/tree_summarize.py +56 -0
  104. autorag/nodes/passagefilter/__init__.py +6 -0
  105. autorag/nodes/passagefilter/base.py +40 -0
  106. autorag/nodes/passagefilter/pass_passage_filter.py +14 -0
  107. autorag/nodes/passagefilter/percentile_cutoff.py +58 -0
  108. autorag/nodes/passagefilter/recency.py +105 -0
  109. autorag/nodes/passagefilter/run.py +138 -0
  110. autorag/nodes/passagefilter/similarity_percentile_cutoff.py +134 -0
  111. autorag/nodes/passagefilter/similarity_threshold_cutoff.py +112 -0
  112. autorag/nodes/passagefilter/threshold_cutoff.py +78 -0
  113. autorag/nodes/passagereranker/__init__.py +16 -0
  114. autorag/nodes/passagereranker/base.py +44 -0
  115. autorag/nodes/passagereranker/cohere.py +118 -0
  116. autorag/nodes/passagereranker/colbert.py +213 -0
  117. autorag/nodes/passagereranker/flag_embedding.py +112 -0
  118. autorag/nodes/passagereranker/flag_embedding_llm.py +101 -0
  119. autorag/nodes/passagereranker/flashrank.py +245 -0
  120. autorag/nodes/passagereranker/jina.py +115 -0
  121. autorag/nodes/passagereranker/koreranker.py +136 -0
  122. autorag/nodes/passagereranker/mixedbreadai.py +126 -0
  123. autorag/nodes/passagereranker/monot5.py +190 -0
  124. autorag/nodes/passagereranker/openvino.py +191 -0
  125. autorag/nodes/passagereranker/pass_reranker.py +31 -0
  126. autorag/nodes/passagereranker/rankgpt.py +170 -0
  127. autorag/nodes/passagereranker/run.py +145 -0
  128. autorag/nodes/passagereranker/sentence_transformer.py +129 -0
  129. autorag/nodes/passagereranker/tart/__init__.py +1 -0
  130. autorag/nodes/passagereranker/tart/modeling_enc_t5.py +152 -0
  131. autorag/nodes/passagereranker/tart/tart.py +139 -0
  132. autorag/nodes/passagereranker/tart/tokenization_enc_t5.py +112 -0
  133. autorag/nodes/passagereranker/time_reranker.py +72 -0
  134. autorag/nodes/passagereranker/upr.py +160 -0
  135. autorag/nodes/passagereranker/voyageai.py +109 -0
  136. autorag/nodes/promptmaker/__init__.py +12 -0
  137. autorag/nodes/promptmaker/base.py +32 -0
  138. autorag/nodes/promptmaker/chat_fstring.py +73 -0
  139. autorag/nodes/promptmaker/fstring.py +49 -0
  140. autorag/nodes/promptmaker/long_context_reorder.py +83 -0
  141. autorag/nodes/promptmaker/run.py +283 -0
  142. autorag/nodes/promptmaker/window_replacement.py +85 -0
  143. autorag/nodes/queryexpansion/__init__.py +4 -0
  144. autorag/nodes/queryexpansion/base.py +62 -0
  145. autorag/nodes/queryexpansion/hyde.py +43 -0
  146. autorag/nodes/queryexpansion/multi_query_expansion.py +57 -0
  147. autorag/nodes/queryexpansion/pass_query_expansion.py +22 -0
  148. autorag/nodes/queryexpansion/query_decompose.py +111 -0
  149. autorag/nodes/queryexpansion/run.py +308 -0
  150. autorag/nodes/retrieval/__init__.py +0 -0
  151. autorag/nodes/retrieval/base.py +127 -0
  152. autorag/nodes/retrieval/run_util.py +152 -0
  153. autorag/nodes/semanticretrieval/__init__.py +1 -0
  154. autorag/nodes/semanticretrieval/run.py +148 -0
  155. autorag/nodes/semanticretrieval/vectordb.py +339 -0
  156. autorag/nodes/util.py +16 -0
  157. autorag/parser.py +37 -0
  158. autorag/schema/__init__.py +3 -0
  159. autorag/schema/base.py +35 -0
  160. autorag/schema/metricinput.py +99 -0
  161. autorag/schema/module.py +24 -0
  162. autorag/schema/node.py +144 -0
  163. autorag/strategy.py +165 -0
  164. autorag/support.py +235 -0
  165. autorag/utils/__init__.py +8 -0
  166. autorag/utils/cast.py +45 -0
  167. autorag/utils/preprocess.py +149 -0
  168. autorag/utils/util.py +759 -0
  169. autorag/validator.py +98 -0
  170. autorag/vectordb/__init__.py +75 -0
  171. autorag/vectordb/base.py +73 -0
  172. autorag/vectordb/chroma.py +118 -0
  173. autorag/vectordb/couchbase.py +239 -0
  174. autorag/vectordb/milvus.py +169 -0
  175. autorag/vectordb/pinecone.py +121 -0
  176. autorag/vectordb/qdrant.py +155 -0
  177. autorag/vectordb/weaviate.py +184 -0
  178. autorag/web.py +81 -0
  179. autorag-0.0.0.dist-info/METADATA +780 -0
  180. autorag-0.0.0.dist-info/RECORD +184 -0
  181. autorag-0.0.0.dist-info/WHEEL +5 -0
  182. autorag-0.0.0.dist-info/entry_points.txt +2 -0
  183. autorag-0.0.0.dist-info/licenses/LICENSE +201 -0
  184. autorag-0.0.0.dist-info/top_level.txt +1 -0
autorag/validator.py ADDED
@@ -0,0 +1,98 @@
1
+ import itertools
2
+ import logging
3
+ import os
4
+ import tempfile
5
+
6
+ import pandas as pd
7
+
8
+ from autorag.evaluator import Evaluator
9
+ from autorag.utils import (
10
+ cast_qa_dataset,
11
+ cast_corpus_dataset,
12
+ validate_qa_from_corpus_dataset,
13
+ )
14
+
15
+ logger = logging.getLogger("AutoRAG")
16
+
17
+
18
+ class Validator:
19
+ def __init__(self, qa_data_path: str, corpus_data_path: str):
20
+ """
21
+ Initialize a Validator object.
22
+
23
+ :param qa_data_path: The path to the QA dataset.
24
+ Must be parquet file.
25
+ :param corpus_data_path: The path to the corpus dataset.
26
+ Must be parquet file.
27
+ """
28
+ # validate data paths
29
+ if not os.path.exists(qa_data_path):
30
+ raise ValueError(f"QA data path {qa_data_path} does not exist.")
31
+ if not os.path.exists(corpus_data_path):
32
+ raise ValueError(f"Corpus data path {corpus_data_path} does not exist.")
33
+ if not qa_data_path.endswith(".parquet"):
34
+ raise ValueError(f"QA data path {qa_data_path} is not a parquet file.")
35
+ if not corpus_data_path.endswith(".parquet"):
36
+ raise ValueError(
37
+ f"Corpus data path {corpus_data_path} is not a parquet file."
38
+ )
39
+ self.qa_data = pd.read_parquet(qa_data_path, engine="pyarrow")
40
+ self.corpus_data = pd.read_parquet(corpus_data_path, engine="pyarrow")
41
+ self.qa_data = cast_qa_dataset(self.qa_data)
42
+ self.corpus_data = cast_corpus_dataset(self.corpus_data)
43
+
44
+ def validate(self, yaml_path: str, qa_cnt: int = 5, random_state: int = 42):
45
+ # Determine the sample size and log a warning if qa_cnt is larger than available records
46
+ available_records = len(self.qa_data)
47
+ safe_sample_size = min(qa_cnt, available_records) # 먼저 safe_sample_size 계산
48
+
49
+ if safe_sample_size < qa_cnt:
50
+ logger.warning(
51
+ f"Minimal Requested sample size ({qa_cnt}) is larger than available records ({available_records}). "
52
+ f"Sampling will be limited to {safe_sample_size} records. "
53
+ )
54
+
55
+ # safe sample QA data
56
+ sample_qa_df = self.qa_data.sample(
57
+ n=safe_sample_size, random_state=random_state
58
+ )
59
+ sample_qa_df.reset_index(drop=True, inplace=True)
60
+
61
+ # get doc_id
62
+ temp_qa_df = sample_qa_df.copy(deep=True)
63
+ flatten_retrieval_gts = (
64
+ temp_qa_df["retrieval_gt"]
65
+ .apply(lambda x: list(itertools.chain.from_iterable(x)))
66
+ .tolist()
67
+ )
68
+ target_doc_ids = list(itertools.chain.from_iterable(flatten_retrieval_gts))
69
+
70
+ # make sample corpus data
71
+ sample_corpus_df = self.corpus_data.loc[
72
+ self.corpus_data["doc_id"].isin(target_doc_ids)
73
+ ]
74
+ sample_corpus_df.reset_index(drop=True, inplace=True)
75
+
76
+ validate_qa_from_corpus_dataset(sample_qa_df, sample_corpus_df)
77
+
78
+ # start Evaluate at temp project directory
79
+ with (
80
+ tempfile.NamedTemporaryFile(suffix=".parquet", delete=False) as qa_path,
81
+ tempfile.NamedTemporaryFile(suffix=".parquet", delete=False) as corpus_path,
82
+ tempfile.TemporaryDirectory(ignore_cleanup_errors=True) as temp_project_dir,
83
+ ):
84
+ sample_qa_df.to_parquet(qa_path.name, index=False)
85
+ sample_corpus_df.to_parquet(corpus_path.name, index=False)
86
+
87
+ evaluator = Evaluator(
88
+ qa_data_path=qa_path.name,
89
+ corpus_data_path=corpus_path.name,
90
+ project_dir=temp_project_dir,
91
+ )
92
+ evaluator.start_trial(yaml_path, skip_validation=True)
93
+ qa_path.close()
94
+ corpus_path.close()
95
+ os.unlink(qa_path.name)
96
+ os.unlink(corpus_path.name)
97
+
98
+ logger.info("Validation complete.")
@@ -0,0 +1,75 @@
1
+ import os
2
+ from typing import List
3
+
4
+ from autorag.support import dynamically_find_function
5
+ from autorag.utils.util import load_yaml_config
6
+ from autorag.vectordb.base import BaseVectorStore
7
+
8
+
9
+ def get_support_vectordb(vectordb_name: str):
10
+ support_vectordb = {
11
+ "chroma": ("autorag.vectordb.chroma", "Chroma"),
12
+ "Chroma": ("autorag.vectordb.chroma", "Chroma"),
13
+ "milvus": ("autorag.vectordb.milvus", "Milvus"),
14
+ "Milvus": ("autorag.vectordb.milvus", "Milvus"),
15
+ "weaviate": ("autorag.vectordb.weaviate", "Weaviate"),
16
+ "Weaviate": ("autorag.vectordb.weaviate", "Weaviate"),
17
+ "pinecone": ("autorag.vectordb.pinecone", "Pinecone"),
18
+ "Pinecone": ("autorag.vectordb.pinecone", "Pinecone"),
19
+ "couchbase": ("autorag.vectordb.couchbase", "Couchbase"),
20
+ "Couchbase": ("autorag.vectordb.couchbase", "Couchbase"),
21
+ "qdrant": ("autorag.vectordb.qdrant", "Qdrant"),
22
+ "Qdrant": ("autorag.vectordb.qdrant", "Qdrant"),
23
+ }
24
+ return dynamically_find_function(vectordb_name, support_vectordb)
25
+
26
+
27
+ def load_vectordb(vectordb_name: str, **kwargs):
28
+ vectordb = get_support_vectordb(vectordb_name)
29
+ return vectordb(**kwargs)
30
+
31
+
32
+ def load_vectordb_from_yaml(yaml_path: str, vectordb_name: str, project_dir: str):
33
+ config_dict = load_yaml_config(yaml_path)
34
+ vectordb_list = config_dict.get("vectordb", [])
35
+ if len(vectordb_list) == 0 or vectordb_name == "default":
36
+ chroma_path = os.path.join(project_dir, "resources", "chroma")
37
+ return load_vectordb(
38
+ "chroma",
39
+ client_type="persistent",
40
+ embedding_model="openai",
41
+ collection_name="openai",
42
+ path=chroma_path,
43
+ )
44
+
45
+ target_dict = list(filter(lambda x: x["name"] == vectordb_name, vectordb_list))
46
+ target_dict[0].pop("name") # delete a name key
47
+ target_vectordb_name = target_dict[0].pop("db_type")
48
+ target_vectordb_params = target_dict[0]
49
+ return load_vectordb(target_vectordb_name, **target_vectordb_params)
50
+
51
+
52
+ def load_all_vectordb_from_yaml(
53
+ yaml_path: str, project_dir: str
54
+ ) -> List[BaseVectorStore]:
55
+ config_dict = load_yaml_config(yaml_path)
56
+ vectordb_list = config_dict.get("vectordb", [])
57
+ if len(vectordb_list) == 0:
58
+ chroma_path = os.path.join(project_dir, "resources", "chroma")
59
+ return [
60
+ load_vectordb(
61
+ "chroma",
62
+ client_type="persistent",
63
+ embedding_model="openai",
64
+ collection_name="openai",
65
+ path=chroma_path,
66
+ )
67
+ ]
68
+
69
+ result_vectordbs = []
70
+ for vectordb_dict in vectordb_list:
71
+ _ = vectordb_dict.pop("name")
72
+ vectordb_type = vectordb_dict.pop("db_type")
73
+ vectordb = load_vectordb(vectordb_type, **vectordb_dict)
74
+ result_vectordbs.append(vectordb)
75
+ return result_vectordbs
@@ -0,0 +1,73 @@
1
+ from abc import abstractmethod
2
+ from typing import List, Tuple, Union
3
+
4
+ from llama_index.embeddings.openai import OpenAIEmbedding
5
+
6
+ from autorag.utils.util import openai_truncate_by_token
7
+ from autorag.embedding.base import EmbeddingModel
8
+
9
+
10
+ class BaseVectorStore:
11
+ support_similarity_metrics = ["l2", "ip", "cosine"]
12
+
13
+ def __init__(
14
+ self,
15
+ embedding_model: Union[str, List[dict]],
16
+ similarity_metric: str = "cosine",
17
+ embedding_batch: int = 100,
18
+ ):
19
+ self.embedding = EmbeddingModel.load(embedding_model)()
20
+ self.embedding_batch = embedding_batch
21
+ self.embedding.embed_batch_size = embedding_batch
22
+ assert similarity_metric in self.support_similarity_metrics, (
23
+ f"search method {similarity_metric} is not supported"
24
+ )
25
+ self.similarity_metric = similarity_metric
26
+
27
+ @abstractmethod
28
+ async def add(
29
+ self,
30
+ ids: List[str],
31
+ texts: List[str],
32
+ ):
33
+ pass
34
+
35
+ @abstractmethod
36
+ def add_embedding(self, ids: List[str], embeddings: List[List[float]]):
37
+ """
38
+ Add the embeddings to the Vector DB.
39
+ """
40
+ pass
41
+
42
+ @abstractmethod
43
+ async def query(
44
+ self, queries: List[str], top_k: int, **kwargs
45
+ ) -> Tuple[List[List[str]], List[List[float]]]:
46
+ pass
47
+
48
+ @abstractmethod
49
+ async def fetch(self, ids: List[str]) -> List[List[float]]:
50
+ """
51
+ Fetch the embeddings of the ids.
52
+ """
53
+ pass
54
+
55
+ @abstractmethod
56
+ async def is_exist(self, ids: List[str]) -> List[bool]:
57
+ """
58
+ Check if the ids exist in the Vector DB.
59
+ """
60
+ pass
61
+
62
+ @abstractmethod
63
+ async def delete(self, ids: List[str]):
64
+ pass
65
+
66
+ def truncated_inputs(self, inputs: List[str]) -> List[str]:
67
+ if isinstance(self.embedding, OpenAIEmbedding):
68
+ openai_embedding_limit = 8000
69
+ results = openai_truncate_by_token(
70
+ inputs, openai_embedding_limit, self.embedding.model_name
71
+ )
72
+ return results
73
+ return inputs
@@ -0,0 +1,118 @@
1
+ from typing import List, Optional, Dict, Tuple, Union
2
+
3
+ from chromadb import (
4
+ EphemeralClient,
5
+ PersistentClient,
6
+ DEFAULT_TENANT,
7
+ DEFAULT_DATABASE,
8
+ CloudClient,
9
+ AsyncHttpClient,
10
+ )
11
+ from chromadb.api.models.AsyncCollection import AsyncCollection
12
+ from chromadb.api.types import QueryResult
13
+
14
+ from autorag.utils.util import apply_recursive
15
+ from autorag.vectordb.base import BaseVectorStore
16
+
17
+
18
+ class Chroma(BaseVectorStore):
19
+ def __init__(
20
+ self,
21
+ embedding_model: Union[str, List[dict]],
22
+ collection_name: str,
23
+ embedding_batch: int = 100,
24
+ client_type: str = "persistent",
25
+ similarity_metric: str = "cosine",
26
+ path: str = None,
27
+ host: str = "localhost",
28
+ port: int = 8000,
29
+ ssl: bool = False,
30
+ headers: Optional[Dict[str, str]] = None,
31
+ api_key: Optional[str] = None,
32
+ tenant: str = DEFAULT_TENANT,
33
+ database: str = DEFAULT_DATABASE,
34
+ ):
35
+ super().__init__(embedding_model, similarity_metric, embedding_batch)
36
+ if client_type == "ephemeral":
37
+ self.client = EphemeralClient(tenant=tenant, database=database)
38
+ elif client_type == "persistent":
39
+ assert path is not None, "path must be provided for persistent client"
40
+ self.client = PersistentClient(path=path, tenant=tenant, database=database)
41
+ elif client_type == "http":
42
+ self.client = AsyncHttpClient(
43
+ host=host,
44
+ port=port,
45
+ ssl=ssl,
46
+ headers=headers,
47
+ tenant=tenant,
48
+ database=database,
49
+ )
50
+ elif client_type == "cloud":
51
+ self.client = CloudClient(
52
+ tenant=tenant,
53
+ database=database,
54
+ api_key=api_key,
55
+ )
56
+ else:
57
+ raise ValueError(
58
+ f"client_type {client_type} is not supported\n"
59
+ "supported client types are: ephemeral, persistent, http, cloud"
60
+ )
61
+
62
+ self.collection = self.client.get_or_create_collection(
63
+ name=collection_name,
64
+ metadata={"hnsw:space": similarity_metric},
65
+ )
66
+
67
+ async def add(self, ids: List[str], texts: List[str]):
68
+ texts = self.truncated_inputs(texts)
69
+ text_embeddings = await self.embedding.aget_text_embedding_batch(texts)
70
+ if isinstance(self.collection, AsyncCollection):
71
+ await self.collection.add(ids=ids, embeddings=text_embeddings)
72
+ else:
73
+ self.collection.add(ids=ids, embeddings=text_embeddings)
74
+
75
+ def add_embedding(self, ids: List[str], embeddings: List[List[float]]):
76
+ self.collection.add(ids=ids, embeddings=embeddings)
77
+
78
+ async def fetch(self, ids: List[str]) -> List[List[float]]:
79
+ if isinstance(self.collection, AsyncCollection):
80
+ fetch_result = await self.collection.get(ids, include=["embeddings"])
81
+ else:
82
+ fetch_result = self.collection.get(ids, include=["embeddings"])
83
+ fetch_embeddings = fetch_result["embeddings"]
84
+ return fetch_embeddings
85
+
86
+ async def is_exist(self, ids: List[str]) -> List[bool]:
87
+ if isinstance(self.collection, AsyncCollection):
88
+ fetched_result = await self.collection.get(ids, include=[])
89
+ else:
90
+ fetched_result = self.collection.get(ids, include=[])
91
+ existed_ids = fetched_result["ids"]
92
+ return list(map(lambda x: x in existed_ids, ids))
93
+
94
+ async def query(
95
+ self, queries: List[str], top_k: int, **kwargs
96
+ ) -> Tuple[List[List[str]], List[List[float]]]:
97
+ queries = self.truncated_inputs(queries)
98
+ query_embeddings: List[
99
+ List[float]
100
+ ] = await self.embedding.aget_text_embedding_batch(queries)
101
+ if isinstance(self.collection, AsyncCollection):
102
+ query_result: QueryResult = await self.collection.query(
103
+ query_embeddings=query_embeddings, n_results=top_k
104
+ )
105
+ else:
106
+ query_result: QueryResult = self.collection.query(
107
+ query_embeddings=query_embeddings, n_results=top_k
108
+ )
109
+ ids = query_result["ids"]
110
+ scores = query_result["distances"]
111
+ scores = apply_recursive(lambda x: 1 - x, scores)
112
+ return ids, scores
113
+
114
+ async def delete(self, ids: List[str]):
115
+ if isinstance(self.collection, AsyncCollection):
116
+ await self.collection.delete(ids)
117
+ else:
118
+ self.collection.delete(ids)
@@ -0,0 +1,239 @@
1
+ import logging
2
+
3
+ from datetime import timedelta
4
+
5
+ from couchbase.auth import PasswordAuthenticator
6
+ from couchbase.cluster import Cluster
7
+ from couchbase.options import ClusterOptions
8
+
9
+ from typing import List, Tuple, Optional, Union
10
+
11
+ from autorag.utils.util import make_batch
12
+ from autorag.vectordb import BaseVectorStore
13
+
14
+ logger = logging.getLogger("AutoRAG")
15
+
16
+
17
+ class Couchbase(BaseVectorStore):
18
+ def __init__(
19
+ self,
20
+ embedding_model: Union[str, List[dict]],
21
+ bucket_name: str,
22
+ scope_name: str,
23
+ collection_name: str,
24
+ index_name: str,
25
+ embedding_batch: int = 100,
26
+ connection_string: str = "",
27
+ username: str = "",
28
+ password: str = "",
29
+ ingest_batch: int = 100,
30
+ text_key: Optional[str] = "text",
31
+ embedding_key: Optional[str] = "embedding",
32
+ scoped_index: bool = True,
33
+ ):
34
+ super().__init__(
35
+ embedding_model=embedding_model,
36
+ similarity_metric="ip",
37
+ embedding_batch=embedding_batch,
38
+ )
39
+
40
+ self.index_name = index_name
41
+ self.bucket_name = bucket_name
42
+ self.scope_name = scope_name
43
+ self.collection_name = collection_name
44
+ self.scoped_index = scoped_index
45
+ self.text_key = text_key
46
+ self.embedding_key = embedding_key
47
+ self.ingest_batch = ingest_batch
48
+
49
+ auth = PasswordAuthenticator(username, password)
50
+ self.cluster = Cluster(connection_string, ClusterOptions(auth))
51
+
52
+ # Wait until the cluster is ready for use.
53
+ self.cluster.wait_until_ready(timedelta(seconds=5))
54
+
55
+ # Check if the bucket exists
56
+ if not self._check_bucket_exists():
57
+ raise ValueError(
58
+ f"Bucket {self.bucket_name} does not exist. "
59
+ " Please create the bucket before searching."
60
+ )
61
+
62
+ try:
63
+ self.bucket = self.cluster.bucket(self.bucket_name)
64
+ self.scope = self.bucket.scope(self.scope_name)
65
+ self.collection = self.scope.collection(self.collection_name)
66
+ except Exception as e:
67
+ raise ValueError(
68
+ "Error connecting to couchbase. "
69
+ "Please check the connection and credentials."
70
+ ) from e
71
+
72
+ # Check if the index exists. Throws ValueError if it doesn't
73
+ try:
74
+ self._check_index_exists()
75
+ except Exception:
76
+ raise
77
+
78
+ # Reinitialize to ensure a consistent state
79
+ self.bucket = self.cluster.bucket(self.bucket_name)
80
+ self.scope = self.bucket.scope(self.scope_name)
81
+ self.collection = self.scope.collection(self.collection_name)
82
+
83
+ async def add(self, ids: List[str], texts: List[str]):
84
+ from couchbase.exceptions import DocumentExistsException
85
+
86
+ texts = self.truncated_inputs(texts)
87
+ text_embeddings: List[
88
+ List[float]
89
+ ] = await self.embedding.aget_text_embedding_batch(texts)
90
+
91
+ documents_to_insert = []
92
+ for _id, text, embedding in zip(ids, texts, text_embeddings):
93
+ doc = {
94
+ self.text_key: text,
95
+ self.embedding_key: embedding,
96
+ }
97
+ documents_to_insert.append({_id: doc})
98
+
99
+ batch_documents_to_insert = make_batch(documents_to_insert, self.ingest_batch)
100
+
101
+ for batch in batch_documents_to_insert:
102
+ insert_batch = {}
103
+ for doc in batch:
104
+ insert_batch.update(doc)
105
+ try:
106
+ self.collection.upsert_multi(insert_batch)
107
+ except DocumentExistsException as e:
108
+ logger.debug(f"Document already exists: {e}")
109
+
110
+ def add_embedding(self, ids: List[str], embeddings: List[List[float]]):
111
+ from couchbase.exceptions import DocumentExistsException
112
+
113
+ documents_to_insert = []
114
+ for _id, embedding in zip(ids, embeddings):
115
+ doc = {
116
+ self.embedding_key: embedding,
117
+ }
118
+ documents_to_insert.append({_id: doc})
119
+
120
+ batch_documents_to_insert = make_batch(documents_to_insert, self.ingest_batch)
121
+
122
+ for batch in batch_documents_to_insert:
123
+ insert_batch = {}
124
+ for doc in batch:
125
+ insert_batch.update(doc)
126
+ try:
127
+ self.collection.upsert_multi(insert_batch)
128
+ except DocumentExistsException as e:
129
+ logger.debug(f"Document already exists: {e}")
130
+
131
+ async def fetch(self, ids: List[str]) -> List[List[float]]:
132
+ # Fetch vectors by IDs
133
+ fetched_result = self.collection.get_multi(ids)
134
+ fetched_vectors = {
135
+ k: v.value[f"{self.embedding_key}"]
136
+ for k, v in fetched_result.results.items()
137
+ }
138
+ return list(map(lambda x: fetched_vectors[x], ids))
139
+
140
+ async def is_exist(self, ids: List[str]) -> List[bool]:
141
+ existed_result = self.collection.exists_multi(ids)
142
+ existed_ids = {k: v.exists for k, v in existed_result.results.items()}
143
+ return list(map(lambda x: existed_ids[x], ids))
144
+
145
+ async def query(
146
+ self, queries: List[str], top_k: int, **kwargs
147
+ ) -> Tuple[List[List[str]], List[List[float]]]:
148
+ import couchbase.search as search
149
+ from couchbase.options import SearchOptions
150
+ from couchbase.vector_search import VectorQuery, VectorSearch
151
+
152
+ queries = self.truncated_inputs(queries)
153
+ query_embeddings: List[
154
+ List[float]
155
+ ] = await self.embedding.aget_text_embedding_batch(queries)
156
+
157
+ ids, scores = [], []
158
+ for query_embedding in query_embeddings:
159
+ # Create Search Request
160
+ search_req = search.SearchRequest.create(
161
+ VectorSearch.from_vector_query(
162
+ VectorQuery(
163
+ self.embedding_key,
164
+ query_embedding,
165
+ top_k,
166
+ )
167
+ )
168
+ )
169
+
170
+ # Search
171
+ if self.scoped_index:
172
+ search_iter = self.scope.search(
173
+ self.index_name,
174
+ search_req,
175
+ SearchOptions(limit=top_k),
176
+ )
177
+
178
+ else:
179
+ search_iter = self.cluster.search(
180
+ self.index_name,
181
+ search_req,
182
+ SearchOptions(limit=top_k),
183
+ )
184
+
185
+ # Parse the search results
186
+ # search_iter.rows() can only be iterated once.
187
+ id_list, score_list = [], []
188
+ for result in search_iter.rows():
189
+ id_list.append(result.id)
190
+ score_list.append(result.score)
191
+
192
+ ids.append(id_list)
193
+ scores.append(score_list)
194
+
195
+ return ids, scores
196
+
197
+ async def delete(self, ids: List[str]):
198
+ self.collection.remove_multi(ids)
199
+
200
+ def _check_bucket_exists(self) -> bool:
201
+ """Check if the bucket exists in the linked Couchbase cluster.
202
+
203
+ Returns:
204
+ True if the bucket exists
205
+ """
206
+ bucket_manager = self.cluster.buckets()
207
+ try:
208
+ bucket_manager.get_bucket(self.bucket_name)
209
+ return True
210
+ except Exception as e:
211
+ logger.debug("Error checking if bucket exists:", e)
212
+ return False
213
+
214
+ def _check_index_exists(self) -> bool:
215
+ """Check if the Search index exists in the linked Couchbase cluster
216
+ Returns:
217
+ bool: True if the index exists, False otherwise.
218
+ Raises a ValueError if the index does not exist.
219
+ """
220
+ if self.scoped_index:
221
+ all_indexes = [
222
+ index.name for index in self.scope.search_indexes().get_all_indexes()
223
+ ]
224
+ if self.index_name not in all_indexes:
225
+ raise ValueError(
226
+ f"Index {self.index_name} does not exist. "
227
+ " Please create the index before searching."
228
+ )
229
+ else:
230
+ all_indexes = [
231
+ index.name for index in self.cluster.search_indexes().get_all_indexes()
232
+ ]
233
+ if self.index_name not in all_indexes:
234
+ raise ValueError(
235
+ f"Index {self.index_name} does not exist. "
236
+ " Please create the index before searching."
237
+ )
238
+
239
+ return True