diff --git a/reme4/steps/base_step.py b/reme4/steps/base_step.py index 2a79a69b..54dbefe0 100644 --- a/reme4/steps/base_step.py +++ b/reme4/steps/base_step.py @@ -87,20 +87,15 @@ class BaseStep(ABC): return Path.cwd() return Path(self.app_context.app_config.working_dir) - def resolve_path( - self, - raw: str, - *, - require_md: bool = False, - ) -> tuple[Path | None, str | None]: - """Resolve relative `path=` argument under self.working_path. + def resolve_path(self, raw: str) -> Path | None: + """Resolve a relative `path=` argument under self.working_path. Rules: - the caller supplies the full relative path under ``self.working_path``; absolute paths are rejected. - - if the path has no suffix, auto-append ``.md`` (default vault extension). - - if ``require_md=True``, any present non-``.md`` suffix is rejected. Returns ``(abs_path, None)`` on success, or ``(None, error_message)`` on failure. + Filetype-specific gating (e.g. markdown-only / suffix auto-append) is + layered on top by callers — see ``reme4/steps/crud/_file_io.py::gate_md``. """ if not raw or not str(raw).strip(): return None, "`path` is required" @@ -108,12 +103,7 @@ class BaseStep(ABC): p = Path(s) if p.is_absolute(): return None, (f"path {s!r} is absolute; only relative paths accepted") - target = (self.working_path.resolve() / p).resolve() - if target.suffix == "": - target = target.with_suffix(".md") - elif require_md and target.suffix.lower() != ".md": - return None, (f"path {s!r} is not a markdown file; this command only supports .md files") - return target, None + return self.working_path / p, None def _resolve( self, @@ -142,22 +132,12 @@ class BaseStep(ABC): @property def as_llm_formatter(self) -> FormatterBase: """Return the LLM formatter component.""" - return self._resolve( - "as_llm_formatter", - FormatterBase, - ComponentEnum.AS_LLM_FORMATTER, - "formatter", - ) + return self._resolve("as_llm_formatter", FormatterBase, ComponentEnum.AS_LLM_FORMATTER, "formatter") @property def as_token_counter(self) -> TokenCounterBase: """Return the token counter component.""" - return self._resolve( - "as_token_counter", - TokenCounterBase, - ComponentEnum.AS_TOKEN_COUNTER, - "token_counter", - ) + return self._resolve("as_token_counter", TokenCounterBase, ComponentEnum.AS_TOKEN_COUNTER, "token_counter") @property def file_parser(self) -> BaseFileParser: @@ -172,20 +152,12 @@ class BaseStep(ABC): @property def embedding(self) -> BaseEmbeddingModel: """Return the embedding model component.""" - return self._resolve( - "embedding", - BaseEmbeddingModel, - ComponentEnum.EMBEDDING_MODEL, - ) + return self._resolve("embedding", BaseEmbeddingModel, ComponentEnum.EMBEDDING_MODEL) @property def file_watcher(self) -> BaseFileWatcher: """Return the file watcher component.""" - return self._resolve( - "file_watcher", - BaseFileWatcher, - ComponentEnum.FILE_WATCHER, - ) + return self._resolve("file_watcher", BaseFileWatcher, ComponentEnum.FILE_WATCHER) def prompt_format(self, prompt_name: str, **kwargs) -> str: """Format a named prompt template with the given kwargs.""" diff --git a/reme4/steps/crud/_file_io.py b/reme4/steps/crud/_file_io.py index e2d14ce0..bf61171a 100644 --- a/reme4/steps/crud/_file_io.py +++ b/reme4/steps/crud/_file_io.py @@ -1,4 +1,6 @@ -"""Shared filesystem helpers for CRUD steps (safe read, truncation).""" +"""Shared filesystem helpers for CRUD steps (path gating, safe read, truncation).""" + +from pathlib import Path import aiofiles import aiofiles.os @@ -9,6 +11,19 @@ from ...utils import get_logger logger = get_logger() +def gate_md(target: Path, raw: str) -> tuple[Path | None, str | None]: + """Markdown-only gate: auto-append `.md` when no suffix; reject any non-`.md` suffix. + + Layered on top of ``BaseStep.resolve_path`` to keep filetype-specific rules + out of the generic path resolver. + """ + if target.suffix == "": + return target.with_suffix(".md"), None + if target.suffix.lower() != ".md": + return None, (f"path {raw!r} is not a markdown file; this command only supports .md files") + return target, None + + async def read_file_safe(file_path, max_bytes: int = MAX_FILE_READ_BYTES) -> str: """Read file with utf-8-sig (BOM-tolerant), fallback to errors='ignore'.""" stat = await aiofiles.os.stat(str(file_path)) diff --git a/reme4/steps/crud/read.py b/reme4/steps/crud/read.py index d0657110..fc30bf10 100644 --- a/reme4/steps/crud/read.py +++ b/reme4/steps/crud/read.py @@ -1,9 +1,7 @@ """Read a markdown file from the vault, with line-range slicing and byte-truncation.""" -from pathlib import Path - from ._file_io import ( - DEFAULT_MAX_BYTES, + gate_md, read_file_safe, truncate_text_output, ) @@ -22,52 +20,49 @@ class ReadStep(BaseStep): if meta: self.context.response.metadata.update(meta) - def _resolve_target_or_fail(self) -> Path | None: + async def execute(self): # pylint: disable=too-many-return-statements assert self.context is not None raw = str(self.context.get("path") or "") - target, err = self.resolve_path(raw, require_md=True) + start_line = self.context.get("start_line") + end_line = self.context.get("end_line") + + target, err = self.resolve_path(raw) if err: self._fail(err) return None - if not target.exists(): - self._fail(f"file {target} does not exist", path=str(target)) - return None - if not target.is_file(): - self._fail(f"path {target} is not a file", path=str(target)) - return None - return target - def _validate_line_params_or_fail(self) -> bool: - assert self.context is not None - for label in ("start_line", "end_line"): - value = self.context.get(label) + target, err = gate_md(target, raw) + if err: + self._fail(err) + return None + + for label, value in (("start_line", start_line), ("end_line", end_line)): if value is None: continue try: int(value) except (TypeError, ValueError): self._fail(f"{label} must be an integer, got {value!r}") - return False - return True + return None - async def _load_content_or_fail(self, target: Path) -> str | None: - try: - return await read_file_safe(target) - except Exception as ex: - self._fail(f"read failed: {ex}", path=str(target)) + if not target.exists(): + self._fail(f"file {target} does not exist", path=str(target)) + return None + if not target.is_file(): + self._fail(f"path {target} is not a file", path=str(target)) return None - def _compute_range_or_fail( - self, - target: Path, - all_lines: list[str], - ) -> tuple[int, int, int] | None: - assert self.context is not None - start_line = self.context.get("start_line") - end_line = self.context.get("end_line") + try: + content = await read_file_safe(target) + except Exception as e: + self._fail(f"read failed: {e}", path=str(target)) + return None + + all_lines = content.split("\n") total = len(all_lines) s = max(1, int(start_line) if start_line is not None else 1) e = min(total, int(end_line) if end_line is not None else total) + if s > total: self._fail( f"start_line {s} exceeds file length ({total} lines)", @@ -78,46 +73,18 @@ class ReadStep(BaseStep): if s > e: self._fail(f"start_line ({s}) > end_line ({e})", path=str(target)) return None - return s, e, total - def _emit_response( - self, - target: Path, - all_lines: list[str], - s: int, - e: int, - total: int, - ) -> None: - assert self.context is not None - max_bytes = int(self.context.get("max_bytes") or DEFAULT_MAX_BYTES) selected = "\n".join(all_lines[s - 1 : e]) text = truncate_text_output( selected, start_line=s, total_lines=total, - max_bytes=max_bytes, file_path=str(target), ) + self.context.response.success = True self.context.response.answer = text self.logger.info( f"[{self.name}] read path={target} lines={s}-{e}/{total} bytes={len(text.encode('utf-8'))}", ) - - async def execute(self): - assert self.context is not None - target = self._resolve_target_or_fail() - if target is None: - return None - if not self._validate_line_params_or_fail(): - return None - content = await self._load_content_or_fail(target) - if content is None: - return None - all_lines = content.split("\n") - rng = self._compute_range_or_fail(target, all_lines) - if rng is None: - return None - s, e, total = rng - self._emit_response(target, all_lines, s, e, total) return self.context.response