diff --git a/reme2/component/file_graph/base_file_graph.py b/reme2/component/file_graph/base_file_graph.py index a0132c50..07504545 100644 --- a/reme2/component/file_graph/base_file_graph.py +++ b/reme2/component/file_graph/base_file_graph.py @@ -33,6 +33,10 @@ class BaseFileGraph(BaseComponent): async def rebuild_links(self) -> None: """Rebuild all links in the graph from each node's payload.""" + @abstractmethod + async def clear(self): + """Clear the graph.""" + # -- Link access ------------------------------------------------------- @abstractmethod diff --git a/reme2/component/file_graph/local_file_graph.py b/reme2/component/file_graph/local_file_graph.py index c2eb74b9..9b13832d 100644 --- a/reme2/component/file_graph/local_file_graph.py +++ b/reme2/component/file_graph/local_file_graph.py @@ -109,6 +109,11 @@ class LocalFileGraph(BaseFileGraph): if link.path: self._add_edge(src, link.path) + async def clear(self): + self._nodes.clear() + self._inverse.clear() + self._pending.clear() + # -- Link access ------------------------------------------------------- async def get_outlinks(self, path: str) -> list[FileLink]: diff --git a/reme2/component/file_store/base_file_store.py b/reme2/component/file_store/base_file_store.py index 9f2ae8dc..ece5d344 100644 --- a/reme2/component/file_store/base_file_store.py +++ b/reme2/component/file_store/base_file_store.py @@ -3,6 +3,7 @@ from abc import abstractmethod from ..base_component import BaseComponent from ..embedding import BaseEmbeddingModel from ..keyword_index import BaseKeywordIndex +from ..file_graph import BaseFileGraph from ...enumeration import ComponentEnum from ...schema import FileChunk, FileNode @@ -15,6 +16,7 @@ class BaseFileStore(BaseComponent): store_name: str, embedding_model: str = "default", keyword_index: str = "default", + file_graph: str = "default", **kwargs, ): super().__init__(**kwargs) @@ -24,23 +26,10 @@ class BaseFileStore(BaseComponent): self.embedding_model = self.bind(embedding_model, BaseEmbeddingModel) self.keyword_index = self.bind(keyword_index, BaseKeywordIndex) + self.file_graph = self.bind(file_graph, BaseFileGraph) self.store_path = self.working_path / self.component_type.value / store_name self.store_path.mkdir(parents=True, exist_ok=True) - self.file_nodes: dict[str, FileNode] = {} - - async def _start(self) -> None: - await self.load_file_nodes() - - async def _close(self) -> None: - await self.dump_file_nodes() - - async def load_file_nodes(self): - ... - - async def dump_file_nodes(self): - ... - async def upsert_file( self, file: tuple[FileNode, list[FileChunk]] | list[tuple[FileNode, list[FileChunk]]], diff --git a/reme2/component/file_store/local_file_store.py b/reme2/component/file_store/local_file_store.py index 155847c3..ec6405db 100644 --- a/reme2/component/file_store/local_file_store.py +++ b/reme2/component/file_store/local_file_store.py @@ -1,11 +1,9 @@ """In-memory file store with JSONL persistence on close.""" -from pathlib import Path - import aiofiles import numpy as np -from pydantic import BaseModel +from schema import FileLink from .base_file_store import BaseFileStore from ..component_registry import R from ...schema import FileChunk, FileNode @@ -22,64 +20,53 @@ class LocalFileStore(BaseFileStore): super().__init__(**kwargs) self.encoding = encoding self.file_chunks: dict[str, FileChunk] = {} - self.nodes_path = self.store_path / "file_nodes.jsonl" self.chunks_path = self.store_path / "file_chunks.jsonl" # Lifecycle async def _start(self) -> None: await super()._start() - await self._load_jsonl(self.chunks_path, self.file_chunks, FileChunk, "id") - self.logger.info( - f"LocalFileStore '{self.store_name}' ready: " - f"{len(self.file_nodes)} nodes, {len(self.file_chunks)} chunks" - ) + if self.chunks_path.exists(): + try: + async with aiofiles.open(self.chunks_path, encoding=self.encoding) as f: + async for line in f: + line = line.strip() + if line: + chunk = FileChunk.model_validate_json(line) + self.file_chunks[chunk.id] = chunk + except Exception as e: + self.logger.exception(f"Failed to load {self.chunks_path}: {e}") + self.logger.info(f"LocalFileStore '{self.store_name}' ready: {len(self.file_chunks)} chunks") async def _close(self) -> None: - await self._dump_jsonl(self.chunks_path, list(self.file_chunks.values())) + try: + tmp = self.chunks_path.with_suffix(".tmp") + async with aiofiles.open(tmp, "w", encoding=self.encoding) as f: + await f.write("\n".join(c.model_dump_json() for c in self.file_chunks.values())) + tmp.replace(self.chunks_path) + except Exception as e: + self.logger.exception(f"Failed to write {self.chunks_path}: {e}") self.file_chunks.clear() await super()._close() - async def _load_jsonl(self, file: Path, target: dict, model: type[BaseModel], key: str) -> None: - if not file.exists(): - return - target.clear() - try: - async with aiofiles.open(file, encoding=self.encoding) as f: - async for line in f: - line = line.strip() - if line: - obj = model.model_validate_json(line) - target[getattr(obj, key)] = obj - except Exception as e: - self.logger.exception(f"Failed to load {file}: {e}") - - async def _dump_jsonl(self, file: Path, items: list[BaseModel]) -> None: - try: - content = "\n".join(o.model_dump_json() for o in items) - tmp = file.with_suffix(".tmp") - async with aiofiles.open(tmp, "w", encoding=self.encoding) as f: - await f.write(content) - tmp.replace(file) - except Exception as e: - self.logger.exception(f"Failed to write {file}: {e}") - # Base class interface - async def load_file_nodes(self) -> None: - await self._load_jsonl(self.nodes_path, self.file_nodes, FileNode, "path") - - async def dump_file_nodes(self) -> None: - await self._dump_jsonl(self.nodes_path, list(self.file_nodes.values())) - async def upsert_file( - self, - file: tuple[FileNode, list[FileChunk]] | list[tuple[FileNode, list[FileChunk]]], + self, + file: tuple[FileNode, list[FileChunk]] | list[tuple[FileNode, list[FileChunk]]], ) -> None: + if not self.file_graph: + raise RuntimeError("file_graph is required for upsert_file") if isinstance(file, tuple): file = [file] + + old_map = {n.path: n for n in await self.file_graph.get_nodes([node.path for node, _ in file])} + + new_nodes: list[FileNode] = [] + needs_embed: list[FileChunk] = [] + keyword_docs: dict[str, str] = {} for node, chunks in file: - old_node = self.file_nodes.pop(node.path, None) + old_node: FileNode | None = old_map.get(node.path) cached = {} if old_node and self.embedding_model: for cid in old_node.chunk_ids: @@ -88,7 +75,6 @@ class LocalFileStore(BaseFileStore): cached[cid] = old.embedding node.chunk_ids = [] - needs_embed = [] for c in chunks: if self.embedding_model and not c.embedding: if c.id in cached: @@ -97,33 +83,38 @@ class LocalFileStore(BaseFileStore): needs_embed.append(c) node.chunk_ids.append(c.id) self.file_chunks[c.id] = c - self.file_nodes[node.path] = node + if c.text: + keyword_docs[c.id] = c.text + new_nodes.append(node) - if needs_embed and self.embedding_model: - await self.embedding_model.get_node_embeddings(needs_embed) - - if self.keyword_index: - await self.keyword_index.add_docs({c.id: c.text for c in chunks if c.text}) + await self.file_graph.upsert_nodes(new_nodes) + if needs_embed and self.embedding_model: + await self.embedding_model.get_node_embeddings(needs_embed) + if self.keyword_index and keyword_docs: + await self.keyword_index.add_docs(keyword_docs) async def delete_by_path(self, path: str | list[str]) -> None: + if not self.file_graph: + raise RuntimeError("file_graph is required for delete_by_path") if isinstance(path, str): path = [path] - deleted_chunk_ids: list[str] = [] - for p in path: - node = self.file_nodes.pop(p, None) - if node: - for cid in node.chunk_ids: - self.file_chunks.pop(cid, None) - deleted_chunk_ids.append(cid) - + nodes = await self.file_graph.get_nodes(path) + if not nodes: + return + deleted_chunk_ids = [cid for n in nodes for cid in n.chunk_ids] + for cid in deleted_chunk_ids: + self.file_chunks.pop(cid, None) + await self.file_graph.delete_nodes([n.path for n in nodes]) if self.keyword_index and deleted_chunk_ids: await self.keyword_index.delete_docs(deleted_chunk_ids) async def clear(self) -> None: - self.file_nodes.clear() + if not self.file_graph: + raise RuntimeError("file_graph is required for clear") self.file_chunks.clear() if self.keyword_index: await self.keyword_index.clear() + await self.file_graph.clear() # Search @@ -165,3 +156,18 @@ class LocalFileStore(BaseFileStore): results.append(chunk.model_copy(update={"scores": {"keyword": score, "score": score}})) return results + + async def rebuild_links(self) -> None: + if not self.file_graph: + raise RuntimeError("file_graph is required for delete_by_path") + return await self.file_graph.rebuild_links() + + async def get_outlinks(self, path: str) -> list[FileLink]: + if not self.file_graph: + raise RuntimeError("file_graph is required for delete_by_path") + return await self.file_graph.get_outlinks(path) + + async def get_inlinks(self, path: str) -> list[FileLink]: + if not self.file_graph: + raise RuntimeError("file_graph is required for delete_by_path") + return await self.file_graph.get_inlinks(path)