diff --git a/reme4/components/file_store/base_file_store.py b/reme4/components/file_store/base_file_store.py index 4f407145..7e140535 100644 --- a/reme4/components/file_store/base_file_store.py +++ b/reme4/components/file_store/base_file_store.py @@ -24,13 +24,17 @@ class BaseFileStore(BaseComponent): **kwargs, ): super().__init__(**kwargs) + from ..embedding import OpenAIEmbeddingModel + from ..file_graph import LocalFileGraph + from ..keyword_index import BM25Index + self.store_name = store_name or self.name if not embedding_model and not keyword_index: raise ValueError("At least one of embedding_model or keyword_index must be set.") - 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.embedding_model = self.bind(embedding_model, BaseEmbeddingModel, default_factory=OpenAIEmbeddingModel) + self.keyword_index = self.bind(keyword_index, BaseKeywordIndex, default_factory=BM25Index) + self.file_graph = self.bind(file_graph, BaseFileGraph, default_factory=LocalFileGraph) self.store_path = self.working_path / self.component_type.value / store_name self.store_path.mkdir(parents=True, exist_ok=True) diff --git a/reme4/components/file_watcher/base_file_watcher.py b/reme4/components/file_watcher/base_file_watcher.py index 60142692..236b06a9 100644 --- a/reme4/components/file_watcher/base_file_watcher.py +++ b/reme4/components/file_watcher/base_file_watcher.py @@ -30,6 +30,9 @@ class BaseFileWatcher(BaseComponent): **kwargs, ): super().__init__(**kwargs) + from ..file_parser import DefaultFileParser + from ..file_store import LocalFileStore + watch_paths = [watch_paths] if isinstance(watch_paths, str) else watch_paths base = self.working_path self.watch_paths: list[Path] = [base / x for x in watch_paths if (base / x).exists()] @@ -38,8 +41,8 @@ class BaseFileWatcher(BaseComponent): self.force_polling: bool = force_polling self.debounce: int = debounce self.poll_delay_ms: int = poll_delay_ms - self.file_store = self.bind(file_store, BaseFileStore) - self.file_parser = self.bind(file_parser, BaseFileParser) + self.file_store = self.bind(file_store, BaseFileStore, default_factory=LocalFileStore) + self.file_parser = self.bind(file_parser, BaseFileParser, default_factory=DefaultFileParser) self._stop_event: asyncio.Event = asyncio.Event() self._background_task: asyncio.Task | None = None self._retry_interval: float = 10 diff --git a/reme4/steps/base_step.py b/reme4/steps/base_step.py index 2f8da442..c79d5f5e 100644 --- a/reme4/steps/base_step.py +++ b/reme4/steps/base_step.py @@ -78,11 +78,15 @@ class BaseStep(ABC): def _resolve(self, key: str, base_cls: type[T], comp_enum: ComponentEnum, attr: str | None = None) -> T: """Return a kwargs-supplied instance, or look one up by name in the app registry.""" - value = self.kwargs.get(key, "default") - if isinstance(value, base_cls): - return value + # 1. Step init kwargs, 2. Runtime context (run_job kwargs), 3. App registry by name. + for source in (self.kwargs, self.context or {}): + value = source.get(key) + if isinstance(value, base_cls): + return value + + name = self.kwargs.get(key, "default") assert self.app_context is not None - comp = self.app_context.components[comp_enum][value] + comp = self.app_context.components[comp_enum][name] return getattr(comp, attr) if attr else comp @property