further fix for better review adaptation

This commit is contained in:
imrewce 2026-05-19 11:41:51 +08:00
parent c96a192d23
commit 3b340a440c
3 changed files with 52 additions and 98 deletions

View file

@ -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."""

View file

@ -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))

View file

@ -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