diff --git a/doc/changes/unreleased.md b/doc/changes/unreleased.md index 7e971f989..e350a4dd2 100644 --- a/doc/changes/unreleased.md +++ b/doc/changes/unreleased.md @@ -7,3 +7,4 @@ * #942: Added api-contract-audit skill for identifying mismatches between type annotations, docstrings, and runtime behavior * #963: Extended packaged skill checks and installation to support multiple skills. +* #967: Refactored skill installation into reusable filesystem helpers. diff --git a/exasol/toolbox/util/skills.py b/exasol/toolbox/util/skills.py index 6eae1ff0d..febfc83ba 100644 --- a/exasol/toolbox/util/skills.py +++ b/exasol/toolbox/util/skills.py @@ -71,19 +71,14 @@ def _has_symlink_in_parents(path: Path) -> bool: return any(candidate.is_symlink() for candidate in (path, *path.parents)) -def install_skill( - skill_name: str = PTB_SKILL_NAME, - target_directory: Path | None = None, -) -> Path: - """Install a packaged skill into a project-local agent skill directory.""" +def _validate_skill_name(skill_name: str) -> None: + """Reject skill names that could escape the skill installation directory.""" if Path(skill_name).name != skill_name: raise ValueError(f"invalid skill name: {skill_name}") - source_files = get_skill_files(skill_name) - if not source_files: - raise ValueError(f"packaged skill does not exist: {skill_name}") - target_directory = target_directory or Path.cwd() / ".agents" / "skills" +def _prepare_installation_directory(target_directory: Path, skill_name: str) -> Path: + """Validate and recreate the destination directory for one skill.""" target_skill = target_directory / skill_name if _has_symlink_in_parents(target_directory): raise ValueError( @@ -97,10 +92,30 @@ def install_skill( if target_skill.exists(): shutil.rmtree(target_skill) target_skill.mkdir(parents=True, exist_ok=True) + return target_skill + + +def _copy_skill_files(source_files: Mapping[str, Traversable], target: Path) -> None: + """Copy packaged skill files below an already validated target directory.""" for relative_path, source in source_files.items(): - destination = target_skill / relative_path + destination = target / relative_path destination.parent.mkdir(parents=True, exist_ok=True) destination.write_bytes(source.read_bytes()) + + +def install_skill( + skill_name: str = PTB_SKILL_NAME, + target_directory: Path | None = None, +) -> Path: + """Install a packaged skill into a project-local agent skill directory.""" + _validate_skill_name(skill_name) + source_files = get_skill_files(skill_name) + if not source_files: + raise ValueError(f"packaged skill does not exist: {skill_name}") + + target_directory = target_directory or Path.cwd() / ".agents" / "skills" + target_skill = _prepare_installation_directory(target_directory, skill_name) + _copy_skill_files(source_files, target_skill) return target_skill diff --git a/test/unit/util/skill_test.py b/test/unit/util/skill_test.py index d96f254ba..c6c0b761e 100644 --- a/test/unit/util/skill_test.py +++ b/test/unit/util/skill_test.py @@ -1,6 +1,11 @@ import pytest from exasol.toolbox.util import skills +from exasol.toolbox.util.skills import ( + _copy_skill_files, + _prepare_installation_directory, + _validate_skill_name, +) def test_validate_skill_accepts_packaged_ptb_skill(): @@ -115,6 +120,34 @@ def test_install_skill_rejects_path_traversal(tmp_path): skills.install_skill("../outside", tmp_path) +def test_validate_skill_name_rejects_path_traversal(): + with pytest.raises(ValueError, match="invalid skill name"): + _validate_skill_name("nested/example") + + +def test_prepare_installation_directory_replaces_existing_directory(tmp_path): + target_directory = tmp_path / ".agents" / "skills" + target_skill = target_directory / "example" + target_skill.mkdir(parents=True) + (target_skill / "stale.md").write_text("stale", encoding="utf-8") + + prepared = _prepare_installation_directory(target_directory, "example") + + assert prepared == target_skill + assert not (target_skill / "stale.md").exists() + + +def test_copy_skill_files_copies_nested_files(tmp_path): + source = tmp_path / "source.md" + source.write_text("content", encoding="utf-8") + target = tmp_path / "target" + target.mkdir() + + _copy_skill_files({"references/source.md": source}, target) + + assert (target / "references/source.md").read_text(encoding="utf-8") == "content" + + def test_install_skill_rejects_missing_skill(tmp_path, monkeypatch): monkeypatch.setattr(skills, "get_skill_files", lambda _: {})