update vector store op

This commit is contained in:
jinli.yl 2025-07-11 17:13:24 +08:00
parent 8d8dbd48cf
commit 12f93937fe
15 changed files with 251 additions and 161 deletions

View file

@ -12,7 +12,7 @@ class Mock1Op(BaseOp):
time.sleep(1)
a: int = self.op_params["a"]
b: str = self.op_params["b"]
logger.info(f"enter class={self.__class__.__name__}. a={a} b={b}")
logger.info(f"enter class={self.simple_name}. a={a} b={b}")
@OP_REGISTRY.register()

View file

@ -11,13 +11,13 @@ class BuildQueryOp(BaseOp):
RETRIEVE_QUERY = "retrieve_query"
def execute(self):
# @jiaji
request: RetrieverRequest = self.context.request
if request.query:
query = request.query
elif request.messages:
if self.op_params.get("enable_llm_build") is True:
enable_llm_build: str = str(self.op_params.get("enable_llm_build"))
if enable_llm_build and enable_llm_build.lower() == "true":
execution_process = merge_messages_content(request.messages)
query = self.prompt_format(prompt_name="query_build", execution_process=execution_process)
else:

View file

@ -0,0 +1,13 @@
"""
1. retrieve:
search: query(context), workspace_id(request), top_k(request)
2. summary:
insert: nodes(context), workspace_id(request)
delete: ids(context), workspace_id(request)
search: query(context), workspace_id(request), top_k(request.config.op)
3. vector:
dump: workspace_id(request), path(str), max_size(int)
load: workspace_id(request), path(str)
delete: workspace_id(request)
copy: source_id, target_id, max_size(int)
"""

View file

@ -1 +0,0 @@
# jinli

View file

@ -1,17 +0,0 @@
from typing import List
from experiencemaker.op import OP_REGISTRY
from experiencemaker.op.base_op import BaseOp
from experiencemaker.schema.vector_node import VectorNode
@OP_REGISTRY.register()
class InsertDatabaseOp(BaseOp):
INSERT_NODES: str = "insert_nodes"
def execute(self):
nodes: List[VectorNode] = self.context.get_context(InsertDatabaseOp.INSERT_NODES)
self.vector_store.insert(nodes=nodes, workspace_id=self.context.request.workspace_id)

View file

@ -9,20 +9,18 @@ from experiencemaker.schema.vector_node import VectorNode
@OP_REGISTRY.register()
class VectorRecallOp(BaseOp):
class RecallVectorStoreOp(BaseOp):
def execute(self):
# get query
from experiencemaker.op.retriever.build_query_op import BuildQueryOp
query = self.context.get_context(BuildQueryOp.RETRIEVE_QUERY)
query = self.context.get_context("search_query")
assert query, "query should be not empty!"
# retrieve from vector store
request: RetrieverRequest = self.context.request
top_k: int = int(request.metadata.get("top_k", 5))
nodes: List[VectorNode] = self.vector_store.retrieve_by_query(query=query,
workspace_id=request.workspace_id,
top_k=top_k)
nodes: List[VectorNode] = self.vector_store.search(query=query,
workspace_id=request.workspace_id,
top_k=request.top_k)
# convert to experience, filter duplicate
experience_list: List[BaseExperience] = []
@ -36,7 +34,7 @@ class VectorRecallOp(BaseOp):
# filter by score
threshold_score: float | None = self.op_params.get("threshold_score", None)
if threshold_score is not None:
experience_list = [e for e in experience_list if e.score >= threshold_score]
experience_list = [e for e in experience_list if e.score >= threshold_score or e.score is None]
# set response
request: RetrieverResponse = self.context.response

View file

@ -1 +0,0 @@
# @jinli

View file

@ -0,0 +1,29 @@
import json
from typing import List
from loguru import logger
from experiencemaker.op import OP_REGISTRY
from experiencemaker.op.base_op import BaseOp
from experiencemaker.schema.request import BaseRequest
from experiencemaker.schema.vector_node import VectorNode
@OP_REGISTRY.register()
class UpdateVectorStoreOp(BaseOp):
INSERT_NODES = "insert_nodes"
DELETE_NODE_IDS = "delete_node_ids"
def execute(self):
request: BaseRequest = self.context.request
node_ids: List[str] | None = self.context.get_context(self.DELETE_NODE_IDS)
if node_ids:
self.vector_store.delete(node_ids=node_ids, workspace_id=request.workspace_id)
logger.info(f"delete node_ids={json.dumps(node_ids, indent=2)}")
insert_nodes: List[VectorNode] | None = self.context.get_context(self.INSERT_NODES)
if insert_nodes:
self.vector_store.insert(nodes=insert_nodes, workspace_id=request.workspace_id)
for node in insert_nodes:
logger.info(f"insert insert_node={node.model_dump_json(indent=2)}")

View file

@ -0,0 +1,35 @@
from experiencemaker.op import OP_REGISTRY
from experiencemaker.op.base_op import BaseOp
from experiencemaker.schema.request import VectorStoreRequest
from experiencemaker.schema.response import VectorStoreResponse
@OP_REGISTRY.register()
class VectorStoreActionOp(BaseOp):
def execute(self):
request: VectorStoreRequest = self.context.request
response: VectorStoreResponse = self.context.response
if request.action == "copy":
result = self.vector_store.copy_workspace(src_workspace_id=request.src_workspace_id,
dest_workspace_id=request.workspace_id)
elif request.action == "delete":
result = self.vector_store.delete_workspace(workspace_id=request.workspace_id)
elif request.action == "dump":
result = self.vector_store.dump_workspace(workspace_id=request.workspace_id,
path=request.path)
elif request.action == "load":
result = self.vector_store.load_workspace(workspace_id=request.workspace_id,
path=request.path)
else:
raise ValueError(f"invalid action={request.action}")
if isinstance(result, dict):
response.metadata.update(result)
else:
response.metadata["result"] = str(result)

View file

@ -2,6 +2,8 @@ from concurrent.futures import ThreadPoolExecutor
from typing import Dict
from experiencemaker.schema.app_config import AppConfig
from experiencemaker.schema.request import BaseRequest
from experiencemaker.schema.response import BaseResponse
from experiencemaker.vector_store.base_vector_store import BaseVectorStore

View file

@ -8,12 +8,12 @@ from experiencemaker.schema.message import Message, Trajectory
class BaseRequest(BaseModel):
workspace_id: str = Field(default=...)
config: dict = Field(default_factory=dict)
metadata: dict | None = Field(default=None)
class RetrieverRequest(BaseRequest):
query: str = Field(default="")
messages: List[Message] = Field(default_factory=list)
top_k: int = Field(default=1)
class SummarizerRequest(BaseRequest):
@ -22,6 +22,8 @@ class SummarizerRequest(BaseRequest):
class VectorStoreRequest(BaseRequest):
action: str = Field(default="")
src_workspace_id: str = Field(default="")
path: str = Field(default="")
class AgentRequest(BaseRequest):

View file

@ -14,20 +14,24 @@ from experiencemaker.schema.vector_node import VectorNode
class BaseVectorStore(BaseModel, ABC):
embedding_model: BaseEmbeddingModel | None = Field(default=None)
batch_size: int = Field(default=1024)
@staticmethod
def _load_from_path(path: str | Path, workspace_id: str, **kwargs) -> Iterable[VectorNode]:
def _load_from_path(workspace_id: str, path: str | Path, **kwargs) -> Iterable[VectorNode]:
workspace_path = Path(path) / f"{workspace_id}.jsonl"
if workspace_path.exists():
with workspace_path.open() as f:
fcntl.flock(f, fcntl.LOCK_SH)
try:
for line in tqdm(f, desc="load from path"):
if line.strip():
yield VectorNode(**json.loads(line.strip(), **kwargs))
if not workspace_path.exists():
logger.warning(f"workspace_path={workspace_path} is not exists!")
return
finally:
fcntl.flock(f, fcntl.LOCK_UN)
with workspace_path.open() as f:
fcntl.flock(f, fcntl.LOCK_SH)
try:
for line in tqdm(f, desc="load from path"):
if line.strip():
yield VectorNode(**json.loads(line.strip(), **kwargs))
finally:
fcntl.flock(f, fcntl.LOCK_UN)
@staticmethod
def _dump_to_path(nodes: Iterable[VectorNode], workspace_id: str, path: str | Path = "",
@ -36,68 +40,82 @@ class BaseVectorStore(BaseModel, ABC):
dump_path.mkdir(parents=True, exist_ok=True)
dump_file = dump_path / f"{workspace_id}.jsonl"
count = 0
with dump_file.open("w") as f:
fcntl.flock(f, fcntl.LOCK_EX)
try:
for node in tqdm(nodes, desc="dump to path"):
f.write(json.dumps(node.model_dump(), ensure_ascii=ensure_ascii, **kwargs))
f.write("\n")
count += 1
return {"size": count}
finally:
fcntl.flock(f, fcntl.LOCK_UN)
def exist_workspace(self, workspace_id: str, **kwargs) -> bool:
raise NotImplementedError
def _delete_workspace(self, workspace_id: str, **kwargs):
raise NotImplementedError
def delete_workspace(self, workspace_id: str, **kwargs):
if self.exist_workspace(workspace_id, **kwargs):
self._delete_workspace(workspace_id, **kwargs)
def _create_workspace(self, workspace_id: str, **kwargs):
raise NotImplementedError
def create_workspace(self, workspace_id: str, **kwargs):
if self.exist_workspace(workspace_id, **kwargs):
logger.warning(f"workspace={workspace_id} exists~")
return
self._create_workspace(workspace_id, **kwargs)
raise NotImplementedError
def _iter_workspace_nodes(self, workspace_id: str, max_size: int = 10000, **kwargs) -> Iterable[VectorNode]:
def _iter_workspace_nodes(self, workspace_id: str, **kwargs) -> Iterable[VectorNode]:
raise NotImplementedError
def dump_workspace(self, workspace_id: str, path: str | Path = "", **kwargs):
self._dump_to_path(nodes=self._iter_workspace_nodes(workspace_id, **kwargs),
if not self.exist_workspace(workspace_id=workspace_id, **kwargs):
logger.warning(f"workspace_id={workspace_id} is not exist!")
return {}
return self._dump_to_path(nodes=self._iter_workspace_nodes(workspace_id=workspace_id, **kwargs),
workspace_id=workspace_id,
path=path, **kwargs)
def load_workspace(self, workspace_id: str, path: str | Path = "", nodes: List[VectorNode] = None, **kwargs):
if self.exist_workspace(workspace_id, **kwargs):
self.delete_workspace(workspace_id=workspace_id, **kwargs)
logger.info(f"delete workspace_id={workspace_id}")
self.create_workspace(workspace_id=workspace_id, **kwargs)
all_nodes: List[VectorNode] = []
all_nodes.extend(nodes)
if nodes:
all_nodes.extend(nodes)
for node in self._load_from_path(path=path, workspace_id=workspace_id, **kwargs):
all_nodes.append(node)
self.insert(nodes=all_nodes, workspace_id=workspace_id, **kwargs)
return {"size": len(all_nodes)}
def retrieve_by_query(self, query: str, workspace_id: str, top_k: int = 1, **kwargs) -> List[VectorNode]:
def copy_workspace(self, src_workspace_id: str, dest_workspace_id: str, **kwargs):
if not self.exist_workspace(workspace_id=src_workspace_id, **kwargs):
logger.warning(f"src_workspace_id={src_workspace_id} is not exist!")
return {}
if not self.exist_workspace(dest_workspace_id, **kwargs):
self.create_workspace(workspace_id=dest_workspace_id, **kwargs)
nodes = []
node_size = 0
for node in self._iter_workspace_nodes(workspace_id=src_workspace_id, **kwargs):
nodes.append(node)
node_size += 1
if len(nodes) >= self.batch_size:
self.insert(nodes=nodes, workspace_id=dest_workspace_id, **kwargs)
nodes.clear()
if nodes:
self.insert(nodes=nodes, workspace_id=dest_workspace_id, **kwargs)
return {"size": node_size}
def search(self, query: str, workspace_id: str, top_k: int = 1, **kwargs) -> List[VectorNode]:
raise NotImplementedError
def insert(self, nodes: VectorNode | List[VectorNode], workspace_id: str, **kwargs):
raise NotImplementedError
def update(self, nodes: VectorNode | List[VectorNode], workspace_id: str, **kwargs):
def delete(self, node_ids: str | List[str], workspace_id: str, **kwargs):
raise NotImplementedError
"""
unimportant
"""
def retrieve_by_id(self, unique_id: str, workspace_id: str = None, **kwargs) -> VectorNode | None:
raise NotImplementedError
def exist_id(self, unique_id: str, workspace_id: str = None, **kwargs) -> bool:
raise NotImplementedError
def delete_id(self, unique_id: str, workspace_id: str = None, **kwargs):
raise NotImplementedError

View file

@ -31,17 +31,17 @@ class ChromaVectorStore(BaseVectorStore):
def exist_workspace(self, workspace_id: str, **kwargs) -> bool:
return workspace_id in [c.name for c in self._client.list_collections()]
def _delete_workspace(self, workspace_id: str, **kwargs):
def delete_workspace(self, workspace_id: str, **kwargs):
self._client.delete_collection(workspace_id)
if workspace_id in self.collections:
del self.collections[workspace_id]
def _create_workspace(self, workspace_id: str, **kwargs):
def create_workspace(self, workspace_id: str, **kwargs):
self.collections[workspace_id] = self._client.get_or_create_collection(workspace_id)
def _iter_workspace_nodes(self, workspace_id: str, max_size: int = 10000, **kwargs) -> Iterable[VectorNode]:
def _iter_workspace_nodes(self, workspace_id: str, **kwargs) -> Iterable[VectorNode]:
collection: Collection = self._get_collection(workspace_id)
results = collection.peek(limit=max_size)
results = collection.get()
for i in range(len(results["ids"])):
node = VectorNode(workspace_id=workspace_id,
unique_id=results["ids"][i],
@ -49,8 +49,9 @@ class ChromaVectorStore(BaseVectorStore):
metadata=results["metadatas"][i])
yield node
def retrieve_by_query(self, query: str, workspace_id: str, top_k: int = 1, **kwargs) -> List[VectorNode]:
if not self.exist_workspace(workspace_id):
def search(self, query: str, workspace_id: str, top_k: int = 1, **kwargs) -> List[VectorNode]:
if not self.exist_workspace(workspace_id=workspace_id):
logger.warning(f"workspace_id={workspace_id} is not exists!")
return []
collection: Collection = self._get_collection(workspace_id)
@ -65,7 +66,10 @@ class ChromaVectorStore(BaseVectorStore):
nodes.append(node)
return nodes
def insert(self, nodes: VectorNode | List[VectorNode], workspace_id: str, refresh: bool = True, **kwargs):
def insert(self, nodes: VectorNode | List[VectorNode], workspace_id: str, **kwargs):
if not self.exist_workspace(workspace_id=workspace_id):
self.create_workspace(workspace_id=workspace_id)
if isinstance(nodes, VectorNode):
nodes = [nodes]
@ -80,20 +84,23 @@ class ChromaVectorStore(BaseVectorStore):
documents=[n.content for n in all_nodes],
metadatas=[n.metadata for n in all_nodes])
def update(self, nodes: VectorNode | List[VectorNode], workspace_id: str, refresh: bool = True, **kwargs):
if isinstance(nodes, VectorNode):
nodes = [nodes]
def delete(self, node_ids: str | List[str], workspace_id: str, **kwargs):
if not self.exist_workspace(workspace_id=workspace_id):
logger.warning(f"workspace_id={workspace_id} is not exists!")
return
if isinstance(node_ids, str):
node_ids = [node_ids]
collection: Collection = self._get_collection(workspace_id)
collection.delete(ids=[node.unique_id for node in nodes])
self.insert(nodes, workspace_id, refresh=refresh)
collection.delete(ids=node_ids)
def main():
from experiencemaker.utils.util_function import load_env_keys
load_env_keys()
load_env_keys("../../.env")
from dotenv import load_dotenv
load_dotenv()
embedding_model = OpenAICompatibleEmbeddingModel(dimensions=1024) # OpenAI text-embedding-ada-002的维度
embedding_model = OpenAICompatibleEmbeddingModel(dimensions=64, model_name="text-embedding-v4")
workspace_id = "chroma_test_index"
chroma_store = ChromaVectorStore(
@ -147,7 +154,7 @@ def main():
chroma_store.insert(sample_nodes, workspace_id=workspace_id)
logger.info("=" * 20)
results = chroma_store.retrieve_by_query("What is AI?", top_k=5, workspace_id=workspace_id)
results = chroma_store.search("What is AI?", top_k=5, workspace_id=workspace_id)
for r in results:
logger.info(r.model_dump(exclude={"vector"}))
logger.info("=" * 20)
@ -162,10 +169,11 @@ def main():
"updated": True
}
)
chroma_store.update(node2_update, workspace_id=workspace_id)
chroma_store.delete(node2_update.unique_id, workspace_id=workspace_id)
chroma_store.insert(node2_update, workspace_id=workspace_id)
logger.info("Updated Result:")
results = chroma_store.retrieve_by_query("fish?", top_k=10, workspace_id=workspace_id)
results = chroma_store.search("fish?", top_k=10, workspace_id=workspace_id)
for r in results:
logger.info(r.model_dump(exclude={"vector"}))
logger.info("=" * 20)

View file

@ -16,26 +16,23 @@ from experiencemaker.vector_store.base_vector_store import BaseVectorStore
class EsVectorStore(BaseVectorStore):
hosts: str | List[str] = Field(default_factory=lambda: os.getenv("ES_HOSTS", "http://localhost:9200"))
basic_auth: str | Tuple[str, str] | None = Field(default=None)
bulk_chunk_size: int = Field(default=512)
retrieve_filters: List[dict] = []
_client: Elasticsearch = PrivateAttr()
@model_validator(mode="after")
def init_client(self):
if isinstance(self.hosts, str):
hosts = [self.hosts]
else:
hosts = self.hosts
self._client = Elasticsearch(hosts=hosts, basic_auth=self.basic_auth)
self.hosts = [self.hosts]
self._client = Elasticsearch(hosts=self.hosts, basic_auth=self.basic_auth)
return self
def exist_workspace(self, workspace_id: str, **kwargs) -> bool:
return self._client.indices.exists(index=workspace_id)
def _delete_workspace(self, workspace_id: str, **kwargs):
self._client.indices.delete(index=workspace_id, **kwargs)
def delete_workspace(self, workspace_id: str, **kwargs):
return self._client.indices.delete(index=workspace_id, **kwargs)
def _create_workspace(self, workspace_id: str, **kwargs):
def create_workspace(self, workspace_id: str, **kwargs):
body = {
"mappings": {
"properties": {
@ -49,14 +46,10 @@ class EsVectorStore(BaseVectorStore):
}
}
}
return self._client.indices.create(index=workspace_id, body=body)
def _iter_workspace_nodes(self, workspace_id: str, max_size: int = 10000, **kwargs) -> Iterable[VectorNode]:
response = self._client.search(index=workspace_id,
body={"query": {"match_all": {}}, "size": max_size},
scroll='5m')
def _iter_workspace_nodes(self, workspace_id: str, **kwargs) -> Iterable[VectorNode]:
response = self._client.search(index=workspace_id, body={"query": {"match_all": {}}})
for doc in response['hits']['hits']:
yield self.doc2node(doc)
@ -71,26 +64,6 @@ class EsVectorStore(BaseVectorStore):
node.metadata["_score"] = doc["_score"] - 1
return node
def exist_id(self, unique_id: str, workspace_id: str = None, **kwargs) -> bool:
response = self._client.exists(index=workspace_id, id=unique_id)
return response.body
def node2doc(self, node: VectorNode, add_op_type: bool = False) -> dict:
doc: dict = {
"_index": node.workspace_id,
"_id": node.unique_id,
"_source": {
"workspace_id": node.workspace_id,
"content": node.content,
"metadata": node.metadata,
"vector": node.vector
}
}
if add_op_type:
doc["_op_type"] = "update" if self.exist_id(node.unique_id, node.workspace_id) else "index",
return doc
def add_term_filter(self, key: str, value):
if key:
self.retrieve_filters.append({"term": {key: value}})
@ -110,7 +83,7 @@ class EsVectorStore(BaseVectorStore):
self.retrieve_filters.clear()
return self
def retrieve_by_query(self, query: str, workspace_id: str, top_k: int = 1, **kwargs) -> List[VectorNode]:
def search(self, query: str, workspace_id: str, top_k: int = 1, **kwargs) -> List[VectorNode]:
if not self.exist_workspace(workspace_id=workspace_id):
logger.warning(f"workspace_id={workspace_id} is not exists!")
return []
@ -137,8 +110,10 @@ class EsVectorStore(BaseVectorStore):
self.retrieve_filters.clear()
return nodes
def insert(self, nodes: VectorNode | List[VectorNode], workspace_id: str, refresh: bool = True, **kwargs):
self.create_workspace(workspace_id=workspace_id)
def insert(self, nodes: VectorNode | List[VectorNode], workspace_id: str, refresh: bool = False, **kwargs):
if not self.exist_workspace(workspace_id=workspace_id):
self.create_workspace(workspace_id=workspace_id)
if isinstance(nodes, VectorNode):
nodes = [nodes]
@ -146,38 +121,55 @@ class EsVectorStore(BaseVectorStore):
not_embedded_nodes = [node for node in nodes if not node.vector]
now_embedded_nodes = self.embedding_model.get_node_embeddings(not_embedded_nodes)
docs = [self.node2doc(node, False) for node in embedded_nodes + now_embedded_nodes]
status, error = bulk(self._client, docs, chunk_size=self.bulk_chunk_size, **kwargs)
logger.info(f"insert sample.size={len(nodes)} status={status} error={error}")
docs = [
{
"_op_type": "index",
"_index": node.workspace_id,
"_id": node.unique_id,
"_source": {
"workspace_id": node.workspace_id,
"content": node.content,
"metadata": node.metadata,
"vector": node.vector
}
} for node in embedded_nodes + now_embedded_nodes]
status, error = bulk(self._client, docs, chunk_size=self.batch_size, **kwargs)
logger.info(f"insert docs.size={len(docs)} status={status} error={error}")
if refresh:
self.refresh(workspace_id=workspace_id)
def update(self, nodes: VectorNode | List[VectorNode], workspace_id: str, refresh: bool = True, **kwargs):
self.create_workspace(workspace_id=workspace_id)
if isinstance(nodes, VectorNode):
nodes = [nodes]
def delete(self, node_ids: str | List[str], workspace_id: str, refresh: bool = False, **kwargs):
if not self.exist_workspace(workspace_id=workspace_id):
logger.warning(f"workspace_id={workspace_id} is not exists!")
return
nodes = self.embedding_model.get_node_embeddings(nodes)
docs = [self.node2doc(node, True) for node in nodes]
status, error = bulk(self._client, docs, chunk_size=self.bulk_chunk_size, **kwargs)
update_size = sum([1 if doc["_op_type"] == "update" else 0 for doc in docs])
insert_size = len(docs) - update_size
logger.info(f"update update_size={update_size} insert_size={insert_size} status={status} error={error}")
if isinstance(node_ids, str):
node_ids = [node_ids]
actions = [
{
"_op_type": "delete",
"_index": workspace_id,
"_id": node_id
} for node_id in node_ids]
status, error = bulk(self._client, actions, chunk_size=self.batch_size, **kwargs)
logger.info(f"delete actions.size={len(actions)} status={status} error={error}")
if refresh:
self.refresh(workspace_id=workspace_id)
def main():
from experiencemaker.utils.util_function import load_env_keys
load_env_keys()
load_env_keys("../../.env")
from dotenv import load_dotenv
load_dotenv()
embedding_model = OpenAICompatibleEmbeddingModel(dimensions=1024)
embedding_model = OpenAICompatibleEmbeddingModel(dimensions=64, model_name="text-embedding-v4")
workspace_id = "rag_nodes_index"
hosts = "http://11.160.132.46:8200"
es = EsVectorStore(hosts=hosts, embedding_model=embedding_model)
es.delete_workspace(workspace_id=workspace_id)
if es.exist_workspace(workspace_id=workspace_id):
es.delete_workspace(workspace_id=workspace_id)
es.create_workspace(workspace_id=workspace_id)
sample_nodes = [
@ -215,13 +207,13 @@ def main():
logger.info("=" * 20)
results = es.add_term_filter(key="metadata.node_type", value="n1") \
.retrieve_by_query("What is AI?", top_k=5, workspace_id=workspace_id)
.search("What is AI?", top_k=5, workspace_id=workspace_id)
for r in results:
logger.info(r.model_dump(exclude={"vector"}))
logger.info("=" * 20)
logger.info("=" * 20)
results = es.retrieve_by_query("What is AI?", top_k=5, workspace_id=workspace_id)
results = es.search("What is AI?", top_k=5, workspace_id=workspace_id)
for r in results:
logger.info(r.model_dump(exclude={"vector"}))
logger.info("=" * 20)

View file

@ -26,20 +26,20 @@ class FileVectorStore(BaseVectorStore):
return Path(self.store_dir)
def exist_workspace(self, workspace_id: str, **kwargs) -> bool:
return (self.store_path / f"{workspace_id}.jsonl").exists()
workspace_path = self.store_path / f"{workspace_id}.jsonl"
return workspace_path.exists()
def _delete_workspace(self, workspace_id: str, **kwargs):
def delete_workspace(self, workspace_id: str, **kwargs):
workspace_path = self.store_path / f"{workspace_id}.jsonl"
if workspace_path.is_file():
workspace_path.unlink()
def _create_workspace(self, workspace_id: str, **kwargs):
def create_workspace(self, workspace_id: str, **kwargs):
self._dump_to_path(nodes=[], workspace_id=workspace_id, path=self.store_path, **kwargs)
def _iter_workspace_nodes(self, workspace_id: str, max_size: int = 10000, **kwargs) -> Iterable[VectorNode]:
def _iter_workspace_nodes(self, workspace_id: str, **kwargs) -> Iterable[VectorNode]:
for i, node in enumerate(self._load_from_path(path=self.store_path, workspace_id=workspace_id, **kwargs)):
if i < max_size:
yield node
yield node
@staticmethod
def calculate_similarity(query_vector: List[float], node_vector: List[float]):
@ -53,7 +53,7 @@ class FileVectorStore(BaseVectorStore):
norm_v2 = math.sqrt(sum(y ** 2 for y in node_vector))
return dot_product / (norm_v1 * norm_v2)
def retrieve_by_query(self, query: str, workspace_id: str, top_k: int = 1, **kwargs) -> List[VectorNode]:
def search(self, query: str, workspace_id: str, top_k: int = 1, **kwargs) -> List[VectorNode]:
query_vector = self.embedding_model.get_embeddings(query)
nodes: List[VectorNode] = []
for node in self._load_from_path(path=self.store_path, workspace_id=workspace_id, **kwargs):
@ -64,9 +64,6 @@ class FileVectorStore(BaseVectorStore):
return nodes[:top_k]
def insert(self, nodes: VectorNode | List[VectorNode], workspace_id: str, **kwargs):
return self.update(nodes=nodes, workspace_id=workspace_id, **kwargs)
def update(self, nodes: VectorNode | List[VectorNode], workspace_id: str, **kwargs):
if isinstance(nodes, VectorNode):
nodes = [nodes]
@ -90,13 +87,28 @@ class FileVectorStore(BaseVectorStore):
logger.info(f"update workspace_id={workspace_id} nodes.size={len(nodes)} all.size={len(all_node_dict)} "
f"update_cnt={update_cnt}")
def delete(self, node_ids: str | List[str], workspace_id: str, **kwargs):
if not self.exist_workspace(workspace_id=workspace_id):
logger.warning(f"workspace_id={workspace_id} is not exists!")
return
if isinstance(node_ids, str):
node_ids = [node_ids]
all_nodes: List[VectorNode] = list(self._load_from_path(path=self.store_path, workspace_id=workspace_id))
before_size = len(all_nodes)
all_nodes = [n for n in all_nodes if n.unique_id not in node_ids]
after_size = len(all_nodes)
self._dump_to_path(nodes=all_nodes, workspace_id=workspace_id, path=self.store_path, **kwargs)
logger.info(f"delete workspace_id={workspace_id} before_size={before_size} after_size={after_size}")
def main():
from experiencemaker.utils.util_function import load_env_keys
load_env_keys()
load_env_keys("../../.env")
from dotenv import load_dotenv
load_dotenv()
embedding_model = OpenAICompatibleEmbeddingModel(dimensions=1024)
embedding_model = OpenAICompatibleEmbeddingModel(dimensions=64, model_name="text-embedding-v4")
workspace_id = "rag_nodes_index"
client = FileVectorStore(embedding_model=embedding_model)
client.delete_workspace(workspace_id)
@ -136,13 +148,13 @@ def main():
client.insert(sample_nodes, workspace_id)
logger.info("=" * 20)
results = client.retrieve_by_query("What is AI?", workspace_id=workspace_id, top_k=5)
results = client.search("What is AI?", workspace_id=workspace_id, top_k=5)
for r in results:
logger.info(r.model_dump(exclude={"vector"}))
logger.info("=" * 20)
client.dump_workspace(workspace_id)
client.delete_workspace(workspace_id)
client.dump_workspace(workspace_id)
if __name__ == "__main__":