diff --git a/reme2/application.py b/reme2/application.py index 82257aa2..676df50b 100644 --- a/reme2/application.py +++ b/reme2/application.py @@ -98,7 +98,7 @@ class Application(BaseComponent): in_degree: dict[tuple[ComponentEnum, str], int] = dict.fromkeys(nodes, 0) dependents: dict[tuple[ComponentEnum, str], list[tuple[ComponentEnum, str]]] = {k: [] for k in nodes} for key, comp in nodes.items(): - for dep in comp.dependencies(): + for dep in comp.dependencies: dep_key = (dep.ctype, dep.name) if dep_key in nodes: dependents[dep_key].append(key) diff --git a/reme2/component/base_component.py b/reme2/component/base_component.py index 1ac9b29c..3cef1f92 100644 --- a/reme2/component/base_component.py +++ b/reme2/component/base_component.py @@ -3,7 +3,7 @@ import asyncio from abc import ABC from pathlib import Path -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Any, Callable, TypeVar, cast from ..enumeration import ComponentEnum from ..utils import get_logger @@ -11,21 +11,48 @@ from ..utils import get_logger if TYPE_CHECKING: from .application_context import ApplicationContext +T = TypeVar("T", bound="BaseComponent") + + +class Dependency: + """Declared dependency: bind() return value, instance attribute placeholder, and topological-sort edge.""" + + __slots__ = ("ctype", "name", "default_factory", "optional") + + def __init__( + self, + ctype: ComponentEnum, + name: str, + default_factory: Callable[[], Any] | None = None, + optional: bool = True, + ) -> None: + self.ctype = ctype + self.name = name + self.default_factory = default_factory + self.optional = optional + + def __repr__(self) -> str: + suffix = "?" if self.optional else "" + return f"" + + def __getattr__(self, item: str) -> Any: + # Guard against using the dependency before start() resolves it. + raise RuntimeError( + f"Dependency {self.ctype.value}:{self.name} accessed before start() (attribute '{item}')", + ) + class BaseComponent(ABC): - """Async lifecycle base class with context manager support. - - Subclasses must implement ``_start`` and ``_close``. - """ + """Async lifecycle base class with bind-based dependency injection.""" component_type = ComponentEnum.BASE def __init__( - self, - name: str | None = None, - backend: str = "", - app_context: "ApplicationContext | None" = None, - **kwargs, + self, + name: str | None = None, + backend: str = "", + app_context: "ApplicationContext | None" = None, + **kwargs, ) -> None: self.name: str = name or self.__class__.__name__ self.backend: str = backend @@ -37,11 +64,61 @@ class BaseComponent(ABC): self._is_started: bool = False self._lock: asyncio.Lock = asyncio.Lock() + # Components created from bind() default_factory in standalone mode (auto-managed lifecycle). + self._owned: list["BaseComponent"] = [] @property def is_started(self) -> bool: return self._is_started + # ----- Dependency declaration ---------------------------------------- + + @staticmethod + def bind( + name: str | None, + base_cls: type[T], + *, + default_factory: Callable[[], T] | None = None, + optional: bool = True, + ) -> T | None: + """Declare a dependency on another component; resolved at start(). Empty name → None.""" + if not name: + return None + ctype = getattr(base_cls, "component_type", None) + if not isinstance(ctype, ComponentEnum) or ctype is ComponentEnum.BASE: + raise TypeError(f"{base_cls.__name__} must declare a non-BASE ComponentEnum 'component_type'") + return cast(T, Dependency(ctype, name, default_factory, optional)) + + @property + def dependencies(self) -> list[Dependency]: + """All unresolved bindings declared on this instance.""" + return [v for v in self.__dict__.values() if isinstance(v, Dependency)] + + async def _resolve_bindings(self) -> None: + """Replace Dependency placeholders with real components (or default_factory / None for optional).""" + for attr, value in list(self.__dict__.items()): + if not isinstance(value, Dependency): + continue + if self.app_context is None: + # Standalone mode: factory or (optional → None) or keep placeholder. + if value.default_factory is not None: + instance = value.default_factory() + setattr(self, attr, instance) + if isinstance(instance, BaseComponent): + self._owned.append(instance) + elif value.optional: + setattr(self, attr, None) + else: + target = self.app_context.components.get(value.ctype, {}).get(value.name) + if target is not None: + setattr(self, attr, target) + elif value.optional: + setattr(self, attr, None) + else: + raise ValueError(f"{value.ctype.value} '{value.name}' not found.") + + # ----- Lookup -------------------------------------------------------- + def get_component(self, component_type: ComponentEnum, name: str): """Get a component by type and name from app_context.""" if self.app_context is None: @@ -57,26 +134,39 @@ class BaseComponent(ABC): return Path.cwd() return Path(self.app_context.app_config.working_dir) + # ----- Lifecycle ----------------------------------------------------- + async def _start(self) -> None: - """Start the component.""" + """Subclass hook: start logic.""" async def _close(self) -> None: - """Close the component.""" + """Subclass hook: close logic.""" async def start(self) -> None: - """Start the component. No-op if already started.""" + """Resolve bindings → start owned fallbacks → _start(). No-op if already started.""" async with self._lock: if self._is_started: return + await self._resolve_bindings() + for owned in self._owned: + try: + await owned.start() + except Exception: + self.logger.exception(f"Failed to start owned {owned.component_type.value}:{owned.name}") await self._start() self._is_started = True async def close(self) -> None: - """Close the component. No-op if not started.""" + """_close() → close owned fallbacks in reverse. No-op if not started.""" async with self._lock: if not self._is_started: return await self._close() + for owned in reversed(self._owned): + try: + await owned.close() + except Exception: + self.logger.exception(f"Failed to close owned {owned.component_type.value}:{owned.name}") self._is_started = False async def restart(self) -> None: diff --git a/reme2/component/file_store/base_file_store.py b/reme2/component/file_store/base_file_store.py index e545fe1b..9f2ae8dc 100644 --- a/reme2/component/file_store/base_file_store.py +++ b/reme2/component/file_store/base_file_store.py @@ -19,28 +19,20 @@ class BaseFileStore(BaseComponent): ): super().__init__(**kwargs) self.store_name = store_name or self.name - self._embedding_model_name = embedding_model - self._keyword_index_name = keyword_index - - self.embedding_model: BaseEmbeddingModel | None = None - self.keyword_index: BaseKeywordIndex | None = None - self.store_path = self.working_path / self.component_type.value / store_name - self.store_path.mkdir(parents=True, exist_ok=True) 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.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: - if self._embedding_model_name: - self.embedding_model = self.get_component(ComponentEnum.EMBEDDING_MODEL, self._embedding_model_name) - if self._keyword_index_name: - self.keyword_index = self.get_component(ComponentEnum.KEYWORD_INDEX, self._keyword_index_name) await self.load_file_nodes() async def _close(self) -> None: - self.embedding_model = None - self.keyword_index = None await self.dump_file_nodes() async def load_file_nodes(self): diff --git a/reme2/component/file_watcher/base_file_watcher.py b/reme2/component/file_watcher/base_file_watcher.py index 144c852a..6d9367a3 100644 --- a/reme2/component/file_watcher/base_file_watcher.py +++ b/reme2/component/file_watcher/base_file_watcher.py @@ -36,18 +36,14 @@ class BaseFileWatcher(BaseComponent): self.force_polling: bool = force_polling self.debounce: int = debounce self.poll_delay_ms: int = poll_delay_ms - self.file_store_name: str = file_store - self.file_parser_name: str = file_parser + self.file_store = self.bind(file_store, BaseFileStore) + self.file_parser = self.bind(file_parser, BaseFileParser) self._stop_event: asyncio.Event = asyncio.Event() self._background_task: asyncio.Task | None = None - self.file_store: BaseFileStore | None = None - self.file_parser: BaseFileParser | None = None self._retry_interval: float = 10 async def _start(self): self._stop_event = asyncio.Event() - self.file_store = self.get_component(ComponentEnum.FILE_STORE, self.file_store_name) - self.file_parser = self.get_component(ComponentEnum.FILE_PARSER, self.file_parser_name) async def background_task(): await self.update_store() diff --git a/reme2/component/keyword_index/base_keyword_index.py b/reme2/component/keyword_index/base_keyword_index.py index 0ee134cd..9c29af25 100644 --- a/reme2/component/keyword_index/base_keyword_index.py +++ b/reme2/component/keyword_index/base_keyword_index.py @@ -15,23 +15,18 @@ class BaseKeywordIndex(BaseComponent): def __init__(self, tokenizer: str = "default", **kwargs): super().__init__(**kwargs) - self.tokenizer_name = tokenizer - self.tokenizer: BaseTokenizer | None = None + from ..tokenizer import RegexTokenizer + + self.tokenizer = self.bind( + tokenizer, + BaseTokenizer, + default_factory=lambda: RegexTokenizer(filter_stopwords=False), + ) self.index_path = self.working_path / self.component_type.value self.index_path.mkdir(parents=True, exist_ok=True) async def _start(self) -> None: - """Initialize tokenizer and load existing index if available.""" - if self.app_context is None: - from ..tokenizer import RegexTokenizer - - self.tokenizer = RegexTokenizer(filter_stopwords=False) - else: - self.tokenizer = self.get_component(ComponentEnum.TOKENIZER, self.tokenizer_name) - - if self.tokenizer is not None: - await self.tokenizer.start() - + """Load existing index if available. Tokenizer is injected and started by the owner lifecycle.""" if self.index_file.exists(): await self.load() self.logger.info(f"Loaded index from {self.index_path}") @@ -41,9 +36,6 @@ class BaseKeywordIndex(BaseComponent): await self.dump() self.logger.info(f"Saved index to {self.index_path}") - if self.tokenizer is not None: - await self.tokenizer.close() - @property def index_file(self) -> Path: """Path to the index pickle file based on tokenizer name."""