feat: support tar archives for installs (#3874)

* feat: support tar archives for installs

Add secure .tar.gz and .tgz parity with ZIP installation for extensions, presets, and workflows, including full workflow package preservation.

Assisted-by: GitHub Copilot (model: GPT-5.6 Sol, autonomous)

Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com>

Copilot-Session: bd07c6b3-f1f9-484c-869a-94d8fef970dd

* chore: clean rebased archive imports

Assisted-by: GitHub Copilot (model: GPT-5.6 Sol, autonomous)

Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com>

Copilot-Session: bd07c6b3-f1f9-484c-869a-94d8fef970dd

* fix: preserve hardened archive install behavior

Keep malformed ZIP diagnostics, filesystem-independent manifest selection, and reserved workflow overlays consistent after adding generic archive support.

Assisted-by: GitHub Copilot (model: GPT-5.6 Sol, autonomous)

Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com>

Copilot-Session: bd07c6b3-f1f9-484c-869a-94d8fef970dd

* fix: extract staged workflow archives by descriptor

Avoid reopening a held staging path on Windows while retaining authoritative-inode archive validation.

Assisted-by: GitHub Copilot (model: GPT-5.6 Sol, autonomous)

Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com>

Copilot-Session: bd07c6b3-f1f9-484c-869a-94d8fef970dd

* fix: extract catalog archives from verified bytes

Use the already bounded and SHA-verified response bytes directly so Windows file-sharing semantics cannot affect archive detection.

Assisted-by: GitHub Copilot (model: GPT-5.6 Sol, autonomous)

Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com>

Copilot-Session: bd07c6b3-f1f9-484c-869a-94d8fef970dd

* fix: address archive install review feedback

Preserve forced preset reinstalls, sniff suffixless workflow archives without weakening YAML limits, and restore prior workflow packages before failed-install cleanup.

Assisted-by: GitHub Copilot (model: GPT-5.6 Sol, autonomous)

Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com>

Copilot-Session: bd07c6b3-f1f9-484c-869a-94d8fef970dd

---------

Copilot-Session: bd07c6b3-f1f9-484c-869a-94d8fef970dd
This commit is contained in:
Manfred Riem
2026-07-31 07:23:00 -05:00
committed by GitHub
parent 6bf51e728a
commit a3e183d069
11 changed files with 1834 additions and 213 deletions

View File

@@ -103,10 +103,17 @@ specify workflow add <source>
| Option | Description |
| --------------- | ------------------------------------------------------ |
| `--dev` | Install from a local workflow YAML file or directory |
| `--dev` | Install from a local YAML file, package directory, or archive |
| `--from <url>` | Install from a custom URL (`<source>` names the expected workflow ID) |
Installs a workflow from the catalog, a URL (HTTPS required), a local YAML file, or a local directory containing `workflow.yml`.
Installs a workflow from the catalog, an HTTPS URL, a local YAML file, a
directory containing `workflow.yml`, or a `.zip`, `.tar.gz`, or `.tgz`
archive. Archives may contain `workflow.yml` at the root or inside one
top-level directory.
Directory and archive installs preserve the complete workflow package,
including scripts and other companion files. ZIP, `.tar.gz`, and `.tgz`
archives follow the same validation and installation behavior.
## Workflow Overlays
@@ -281,7 +288,9 @@ Lower priority values have higher precedence. Change this overlay to `priority:
### Interaction with Bundles and Updates
`specify workflow add <local-directory>` installs `workflow.yml` from the local directory into `.specify/workflows/<id>/`.
`specify workflow add <local-directory>` installs the complete local workflow
package into `.specify/workflows/<id>/`. Archive installs preserve the same
package contents.
When an installed workflow is refreshed or reinstalled, project overlays in `.specify/workflows/overlays/<id>/` are preserved because they live outside the installed workflow directory.

View File

@@ -7,6 +7,7 @@ import re
import socket
import stat
import struct
import tarfile
import unicodedata
import zipfile
from collections.abc import Iterator
@@ -14,11 +15,12 @@ from contextlib import ExitStack, contextmanager
from ipaddress import IPv4Address, IPv6Address, ip_address
from itertools import pairwise
from pathlib import Path, PurePosixPath, PureWindowsPath
from typing import BinaryIO, NoReturn, TypeVar
from typing import BinaryIO, Literal, NoReturn, TypeVar
from urllib.parse import ParseResult, urlparse
ErrorT = TypeVar("ErrorT", bound=Exception)
ArchiveFormat = Literal["zip", "tar.gz"]
MAX_DOWNLOAD_BYTES = 50 * 1024 * 1024
MAX_ZIP_ENTRIES = 512
@@ -67,6 +69,130 @@ _ZIP_MAX_COMMENT_BYTES = (1 << 16) - 1
_BOUNDED_ZIP_COMPRESSION_METHODS = frozenset(
(zipfile.ZIP_STORED, zipfile.ZIP_DEFLATED)
)
_ARCHIVE_CONTENT_TYPES: dict[str, ArchiveFormat] = {
"application/gzip": "tar.gz",
"application/x-gzip": "tar.gz",
"application/x-tar+gzip": "tar.gz",
"application/zip": "zip",
"application/x-zip-compressed": "zip",
}
def archive_format_from_name(name: str) -> ArchiveFormat | None:
"""Return the supported archive format declared by a path or URL."""
try:
path = urlparse(name).path.lower()
except (TypeError, ValueError):
return None
if path.endswith(".tar.gz") or path.endswith(".tgz"):
return "tar.gz"
if path.endswith(".zip"):
return "zip"
return None
def archive_format_from_content_type(content_type: str | None) -> ArchiveFormat | None:
"""Return the supported archive format declared by an HTTP Content-Type."""
if not isinstance(content_type, str):
return None
media_type = content_type.partition(";")[0].strip().lower()
return _ARCHIVE_CONTENT_TYPES.get(media_type)
def archive_suffix(archive_format: ArchiveFormat) -> str:
"""Return the canonical filename suffix for *archive_format*."""
if archive_format == "zip":
return ".zip"
if archive_format == "tar.gz":
return ".tar.gz"
raise ValueError(f"Unsupported archive format: {archive_format!r}")
def detect_archive_format(
archive_path: Path,
*,
archive_file: BinaryIO | None = None,
source_name: str | None = None,
content_type: str | None = None,
error_type: type[ErrorT] = ValueError,
) -> ArchiveFormat:
"""Validate the declared archive format against the file contents.
A recognized path/URL suffix is authoritative. For remote responses whose
final URL has no archive suffix, a recognized Content-Type may declare the
format instead. When both declarations are recognized they must agree, and
the resulting declaration must match the archive bytes.
"""
archive_path = Path(archive_path)
name_format = archive_format_from_name(
source_name if source_name is not None else str(archive_path)
)
content_format = archive_format_from_content_type(content_type)
if (
name_format is not None
and content_format is not None
and name_format != content_format
):
_raise(
error_type,
f"Archive format mismatch: filename declares {name_format} but "
f"Content-Type declares {content_format}",
)
declared_format = name_format or content_format
with ExitStack() as stack:
if archive_file is None:
try:
archive_file = stack.enter_context(archive_path.open("rb"))
except OSError as exc:
_raise_from(error_type, f"Invalid archive: {archive_path}", exc)
try:
archive_file.seek(0)
is_zip = zipfile.is_zipfile(archive_file)
archive_file.seek(0)
signature = archive_file.read(4)
# Let the bounded ZIP preflight report structural errors such as
# impossible entry counts. ``is_zipfile`` rejects those before the
# extractor can produce the established security diagnostic.
is_zip = is_zip or signature in {
b"PK\x03\x04",
b"PK\x05\x06",
b"PK\x07\x08",
}
is_gzip = signature[:2] == b"\x1f\x8b"
archive_file.seek(0)
is_tar_gz = False
if is_gzip:
try:
with tarfile.open(fileobj=archive_file, mode="r:gz"):
is_tar_gz = True
except tarfile.TarError:
pass
archive_file.seek(0)
except OSError as exc:
_raise_from(error_type, f"Invalid archive: {archive_path}", exc)
actual_format: ArchiveFormat | None
if is_zip and not is_tar_gz:
actual_format = "zip"
elif is_tar_gz and not is_zip:
actual_format = "tar.gz"
else:
actual_format = None
if declared_format is None:
if actual_format is None:
_raise(
error_type,
"Unsupported archive format; expected .zip, .tar.gz, or .tgz",
)
declared_format = actual_format
if actual_format != declared_format:
actual_label = actual_format or "invalid/unsupported data"
_raise(
error_type,
f"Archive format mismatch: expected {declared_format}, got {actual_label}",
)
return declared_format
def _ip_address_without_scope(
@@ -292,6 +418,7 @@ def build_safe_download_path(
*,
error_type: type[ErrorT] = ValueError,
label: str = "archive",
suffix: str = ".zip",
) -> Path:
"""Build a portable single-component archive path inside *target_dir*."""
if not isinstance(identifier, str) or not isinstance(version, str):
@@ -301,7 +428,9 @@ def build_safe_download_path(
f"{identifier!r} and {version!r}",
)
filename = f"{identifier}-{version}.zip"
if suffix not in {".zip", ".tar.gz", ".tgz"}:
_raise(error_type, f"Unsupported archive download suffix: {suffix!r}")
filename = f"{identifier}-{version}{suffix}"
try:
filename_too_long = (
len(filename.encode("utf-8")) > MAX_ZIP_COMPONENT_BYTES
@@ -378,24 +507,25 @@ def read_zip_member_limited(
)
def normalize_zip_member_name(
def normalize_archive_member_name(
name: str,
*,
archive_label: str = "archive",
error_type: type[ErrorT] = ValueError,
) -> str:
"""Return a normalized, portable ZIP member name or raise if unsafe."""
"""Return a normalized, portable archive member name or raise if unsafe."""
if "\x00" in name:
_raise(error_type, f"Unsafe path in ZIP archive: {name!r}")
_raise(error_type, f"Unsafe path in {archive_label} archive: {name!r}")
normalized = name.replace("\\", "/")
try:
encoded_name = normalized.encode("utf-8")
except UnicodeEncodeError:
_raise(error_type, f"Unsafe path in ZIP archive: {name!r}")
_raise(error_type, f"Unsafe path in {archive_label} archive: {name!r}")
if len(encoded_name) > MAX_ZIP_PATH_BYTES:
_raise(
error_type,
f"Unsafe path in ZIP archive: {name!r} "
f"Unsafe path in {archive_label} archive: {name!r} "
"(not portable across supported filesystems)",
)
path = PurePosixPath(normalized)
@@ -415,7 +545,8 @@ def normalize_zip_member_name(
):
_raise(
error_type,
f"Unsafe path in ZIP archive: {name!r} (potential path traversal)",
f"Unsafe path in {archive_label} archive: {name!r} "
"(potential path traversal)",
)
for part in raw_parts:
reserved_stem = part.partition(".")[0].partition(":")[0].rstrip(" ")
@@ -432,13 +563,26 @@ def normalize_zip_member_name(
):
_raise(
error_type,
f"Unsafe path in ZIP archive: {name!r} "
f"Unsafe path in {archive_label} archive: {name!r} "
"(not portable across supported filesystems)",
)
return normalized
def portable_zip_path_key(name: str) -> tuple[str, ...]:
def normalize_zip_member_name(
name: str,
*,
error_type: type[ErrorT] = ValueError,
) -> str:
"""Return a normalized, portable ZIP member name or raise if unsafe."""
return normalize_archive_member_name(
name,
archive_label="ZIP",
error_type=error_type,
)
def portable_archive_path_key(name: str) -> tuple[str, ...]:
"""Return a comparison key for filesystems with case/Unicode folding."""
normalized_name = name.replace("\\", "/")
return tuple(
@@ -447,6 +591,11 @@ def portable_zip_path_key(name: str) -> tuple[str, ...]:
)
def portable_zip_path_key(name: str) -> tuple[str, ...]:
"""Backward-compatible ZIP-specific alias for portable archive keys."""
return portable_archive_path_key(name)
def _raise_zip64(error_type: type[ErrorT]) -> NoReturn:
_raise(
error_type,
@@ -778,7 +927,7 @@ def safe_extract_zip(
error_type=error_type,
)
is_dir = member.is_dir() or normalized_name.endswith("/")
path_key = portable_zip_path_key(normalized_name)
path_key = portable_archive_path_key(normalized_name)
existing = validated_paths.get(path_key)
if existing is not None:
@@ -898,3 +1047,211 @@ def safe_extract_zip(
)
if limit_error is not None:
_raise(error_type, limit_error)
def safe_extract_tar(
archive_path: Path,
target_dir: Path,
*,
archive_file: BinaryIO | None = None,
error_type: type[ErrorT] = ValueError,
max_entries: int = MAX_ZIP_ENTRIES,
max_member_bytes: int = MAX_ZIP_MEMBER_BYTES,
max_total_bytes: int = MAX_ZIP_TOTAL_BYTES,
) -> None:
"""Extract a gzip-compressed tar after ZIP-equivalent safety validation."""
_validate_non_negative_int(max_entries, "max_entries")
_validate_non_negative_int(max_member_bytes, "max_member_bytes")
_validate_non_negative_int(max_total_bytes, "max_total_bytes")
archive_path = Path(archive_path)
try:
target_root = target_dir.resolve()
except OSError as exc:
_raise_from(error_type, f"Invalid tar extraction target: {target_dir}", exc)
try:
if archive_file is not None:
archive_file.seek(0)
archive = tarfile.open(
archive_path if archive_file is None else None,
mode="r:gz",
fileobj=archive_file,
)
except (tarfile.TarError, OSError) as exc:
_raise_from(error_type, f"Invalid tar.gz archive: {archive_path}", exc)
with archive:
validated: list[tuple[tarfile.TarInfo, str, bool]] = []
validated_paths: dict[tuple[str, ...], tuple[str, bool]] = {}
total_size = 0
try:
for index, member in enumerate(archive, start=1):
if index > max_entries:
_raise(
error_type,
f"tar.gz archive contains too many entries "
f"({index} > {max_entries})",
)
normalized_name = normalize_archive_member_name(
member.name,
archive_label="tar.gz",
error_type=error_type,
)
is_dir = member.isdir()
if member.issym():
_raise(
error_type,
f"Unsafe symlink in tar.gz archive: {member.name}",
)
if member.islnk():
_raise(
error_type,
f"Unsafe hard link in tar.gz archive: {member.name}",
)
if not is_dir and not member.isreg():
_raise(
error_type,
f"Unsafe member type in tar.gz archive: {member.name}",
)
path_key = portable_archive_path_key(normalized_name)
existing = validated_paths.get(path_key)
if existing is not None:
_raise(
error_type,
f"Conflicting path in tar.gz archive: {member.name} "
f"conflicts with {existing[0]}",
)
validated_paths[path_key] = (member.name, is_dir)
member_path = (target_dir / normalized_name).resolve()
try:
member_path.relative_to(target_root)
except ValueError:
_raise(
error_type,
f"Unsafe path in tar.gz archive: {member.name} "
"(potential path traversal)",
)
if not is_dir:
if member.size > max_member_bytes:
_raise(
error_type,
f"tar.gz member {member.name} exceeds maximum size "
f"of {max_member_bytes} bytes",
)
total_size += member.size
if total_size > max_total_bytes:
_raise(
error_type,
f"tar.gz archive exceeds maximum uncompressed size "
f"of {max_total_bytes} bytes",
)
validated.append((member, normalized_name, is_dir))
except (tarfile.TarError, OSError) as exc:
_raise_from(
error_type,
f"Invalid tar.gz archive: {archive_path}",
exc,
)
for (
(path_key, (original, is_dir)),
(next_key, (next_original, _next_is_dir)),
) in pairwise(sorted(validated_paths.items())):
if (
not is_dir
and len(next_key) > len(path_key)
and next_key[: len(path_key)] == path_key
):
_raise(
error_type,
f"Conflicting path in tar.gz archive: {original} conflicts "
f"with {next_original}",
)
total_written = 0
for member, normalized_name, is_dir in validated:
member_path = target_dir / normalized_name
if is_dir:
try:
member_path.mkdir(parents=True, exist_ok=True)
except OSError as exc:
_raise_from(
error_type,
f"Failed to create tar.gz directory {member.name}: {exc}",
exc,
)
continue
try:
member_path.parent.mkdir(parents=True, exist_ok=True)
source = archive.extractfile(member)
if source is None:
_raise(
error_type,
f"Failed to read tar.gz member {member.name}",
)
written = 0
limit_error: str | None = None
with source, member_path.open("wb") as dest:
while True:
chunk = source.read(READ_CHUNK_SIZE)
if not chunk:
break
written += len(chunk)
if written > max_member_bytes:
limit_error = (
f"tar.gz member {member.name} exceeds maximum size "
f"of {max_member_bytes} bytes"
)
break
total_written += len(chunk)
if total_written > max_total_bytes:
limit_error = (
f"tar.gz archive exceeds maximum uncompressed size "
f"of {max_total_bytes} bytes"
)
break
dest.write(chunk)
except Exception as exc:
_raise_from(
error_type,
f"Failed to extract tar.gz member {member.name}: {exc}",
exc,
)
if limit_error is not None:
_raise(error_type, limit_error)
def safe_extract_archive(
archive_path: Path,
target_dir: Path,
*,
archive_file: BinaryIO | None = None,
source_name: str | None = None,
content_type: str | None = None,
error_type: type[ErrorT] = ValueError,
max_entries: int = MAX_ZIP_ENTRIES,
max_member_bytes: int = MAX_ZIP_MEMBER_BYTES,
max_total_bytes: int = MAX_ZIP_TOTAL_BYTES,
) -> ArchiveFormat:
"""Detect and securely extract a supported archive."""
archive_format = detect_archive_format(
archive_path,
archive_file=archive_file,
source_name=source_name,
content_type=content_type,
error_type=error_type,
)
extractor = safe_extract_zip if archive_format == "zip" else safe_extract_tar
extractor(
archive_path,
target_dir,
archive_file=archive_file,
error_type=error_type,
max_entries=max_entries,
max_member_bytes=max_member_bytes,
max_total_bytes=max_total_bytes,
)
return archive_format

View File

@@ -29,11 +29,14 @@ from packaging.specifiers import InvalidSpecifier, SpecifierSet
from .._assets import _locate_core_pack, _repo_root
from .._download_security import (
archive_format_from_name,
archive_suffix,
MAX_JSON_CATALOG_BYTES,
build_safe_download_path,
detect_archive_format,
is_https_or_localhost_http,
read_response_limited,
safe_extract_zip,
safe_extract_archive,
)
from .._init_options import is_ai_skills_enabled
from .._invocation_style import is_dollar_skills_agent, is_slash_skills_agent
@@ -2403,7 +2406,7 @@ class ExtensionManager:
pass # Best-effort; install already committed to the registry.
# Restore execute bits on shipped POSIX scripts. copytree here (and the
# zipfile.extractall in install_from_zip, which delegates to this method) does
# archive extraction in install_from_archive, which delegates here, does
# not restore a stripped Unix mode, so a bundled *.sh would land non-executable
# and a documented `.specify/extensions/<id>/scripts/...` invocation would fail
# with "Permission denied". This is the single sink every install route funnels
@@ -2422,19 +2425,21 @@ class ExtensionManager:
return manifest
def install_from_zip(
def install_from_archive(
self,
zip_path: Path,
archive_path: Path,
speckit_version: str,
priority: int = 10,
force: bool = False,
*,
archive_file: BinaryIO | None = None,
source_name: str | None = None,
content_type: str | None = None,
) -> ExtensionManifest:
"""Install extension from ZIP file.
"""Install an extension from a supported archive.
Args:
zip_path: Path to extension ZIP file
archive_path: Path to a .zip, .tar.gz, or .tgz archive
speckit_version: Current spec-kit version
priority: Resolution priority (lower = higher precedence, default 10)
force: If True and extension is already installed, remove it first
@@ -2456,10 +2461,12 @@ class ExtensionManager:
with tempfile.TemporaryDirectory() as tmpdir:
temp_path = Path(tmpdir)
safe_extract_zip(
zip_path,
safe_extract_archive(
archive_path,
temp_path,
archive_file=archive_file,
source_name=source_name,
content_type=content_type,
error_type=ValidationError,
)
@@ -2475,13 +2482,35 @@ class ExtensionManager:
manifest_path = extension_dir / "extension.yml"
if not manifest_path.exists():
raise ValidationError("No extension.yml found in ZIP file")
raise ValidationError("No extension.yml found in archive")
# Install from extracted directory
return self.install_from_directory(
extension_dir, speckit_version, priority=priority, force=force
)
def install_from_zip(
self,
zip_path: Path,
speckit_version: str,
priority: int = 10,
force: bool = False,
*,
archive_file: BinaryIO | None = None,
source_name: str | None = None,
content_type: str | None = None,
) -> ExtensionManifest:
"""Backward-compatible wrapper for archive installation."""
return self.install_from_archive(
zip_path,
speckit_version,
priority=priority,
force=force,
archive_file=archive_file,
source_name=source_name,
content_type=content_type,
)
def remove(self, extension_id: str, keep_config: bool = False) -> bool:
"""Remove an installed extension.
@@ -3799,14 +3828,14 @@ class ExtensionCatalog(CatalogStackBase):
def download_extension(
self, extension_id: str, target_dir: Optional[Path] = None
) -> Path:
"""Download extension ZIP from catalog.
"""Download an extension archive from a catalog.
Args:
extension_id: ID of the extension to download
target_dir: Directory to save ZIP file (defaults to temp directory)
target_dir: Directory to save the archive
Returns:
Path to downloaded ZIP file
Path to the downloaded archive
Raises:
ExtensionError: If extension not found or download fails
@@ -3865,45 +3894,88 @@ class ExtensionCatalog(CatalogStackBase):
target_dir = self.cache_dir / "downloads"
target_dir = Path(target_dir)
version = ext_info.get("version", "unknown")
zip_path = build_safe_download_path(
declared_format = archive_format_from_name(download_url)
build_safe_download_path(
target_dir,
extension_id,
version,
error_type=ExtensionError,
label="extension",
suffix=archive_suffix(declared_format or "tar.gz"),
)
target_dir.mkdir(parents=True, exist_ok=True)
original_download_url = download_url
extra_headers = None
resolved_download_url = self._resolve_github_release_asset_api_url(download_url)
if resolved_download_url:
download_url = resolved_download_url
extra_headers = {"Accept": "application/octet-stream"}
# Download the ZIP file
staging_path: Path | None = None
try:
with self._open_url(
download_url, timeout=60, extra_headers=extra_headers
) as response:
zip_data = read_response_limited(
archive_data = read_response_limited(
response,
error_type=ExtensionError,
label=f"extension '{extension_id}' download",
)
final_url = (
response.geturl()
if hasattr(response, "geturl")
else download_url
)
content_type = (
response.getheader("Content-Type")
if hasattr(response, "getheader")
else None
)
verify_archive_sha256(
zip_data, ext_info.get("sha256"), extension_id, ExtensionError
archive_data, ext_info.get("sha256"), extension_id, ExtensionError
)
zip_path.write_bytes(zip_data)
return zip_path
with tempfile.NamedTemporaryFile(
prefix="extension-download-",
suffix=".archive",
dir=target_dir,
delete=False,
) as staging_file:
staging_path = Path(staging_file.name)
staging_file.write(archive_data)
archive_format = detect_archive_format(
staging_path,
source_name=(
final_url
if archive_format_from_name(final_url) is not None
else original_download_url
),
content_type=content_type,
error_type=ExtensionError,
)
archive_path = build_safe_download_path(
target_dir,
extension_id,
version,
error_type=ExtensionError,
label="extension",
suffix=archive_suffix(archive_format),
)
os.replace(staging_path, archive_path)
staging_path = None
return archive_path
except urllib.error.URLError as e:
raise ExtensionError(
f"Failed to download extension from {download_url}: {e}"
)
except IOError as e:
raise ExtensionError(f"Failed to save extension ZIP: {e}")
raise ExtensionError(f"Failed to save extension archive: {e}")
finally:
if staging_path is not None:
staging_path.unlink(missing_ok=True)
def clear_cache(self):
"""Clear the catalog cache (both legacy and URL-hash-based files)."""

View File

@@ -14,7 +14,6 @@ import os
import shutil
import stat
import tempfile
import zipfile
from pathlib import Path
from typing import Optional
from uuid import uuid4
@@ -28,12 +27,11 @@ from rich.table import Table
from .._console import console
from .._assets import get_speckit_version
from .._download_security import (
archive_format_from_name,
detect_archive_format,
is_https_or_localhost_http,
normalize_zip_member_name,
open_zip_bounded,
portable_zip_path_key,
read_response_limited,
read_zip_member_limited,
safe_extract_archive,
)
from .._init_options import is_ai_skills_enabled
@@ -812,19 +810,18 @@ def extension_add(
)
elif from_url:
# Install from URL (ZIP file)
import io
# Install from an archive URL.
import urllib.error
console.print(f"Downloading from {safe_url}...")
download_dir = _validate_safe_cache_dir(project_root)
zip_filename = f"extension-url-download-{uuid4().hex}.zip"
archive_filename = f"extension-url-download-{uuid4().hex}.archive"
# Only used for diagnostic messages: the real archive is a
# transient inode (unlinked on POSIX, O_TEMPORARY on Windows)
# consumed via ``archive_file`` below, so this path is never
# opened again.
zip_path = download_dir / zip_filename
archive_path = download_dir / archive_filename
try:
# Use the catalog's authenticated fetch so configured
@@ -842,28 +839,28 @@ def extension_add(
with dl_catalog._open_url(
download_url, timeout=60, extra_headers=extra_headers
) as response:
zip_data = read_response_limited(
archive_data = read_response_limited(
response,
error_type=ExtensionError,
label=f"extension {from_url}",
)
if not zipfile.is_zipfile(io.BytesIO(zip_data)):
console.print(
f"[red]Error:[/red] {safe_url} did not return a ZIP archive "
f"(got {len(zip_data)} bytes). This usually means the request "
f"was not authenticated and a login/HTML page was returned. "
f"Verify the URL is correct and that credentials for its host "
f"are configured in ~/.specify/auth.json."
final_url = (
response.geturl()
if hasattr(response, "geturl")
else download_url
)
content_type = (
response.getheader("Content-Type")
if hasattr(response, "getheader")
else None
)
raise typer.Exit(1)
download_fd = -1
download_file = None
try:
try:
download_fd = _safe_open_download_zip(
project_root, download_dir, zip_filename
project_root, download_dir, archive_filename
)
except OSError as exc:
console.print(
@@ -875,7 +872,7 @@ def extension_add(
try:
download_file = os.fdopen(download_fd, "w+b")
download_fd = -1
download_file.write(zip_data)
download_file.write(archive_data)
download_file.flush()
download_file.seek(0)
except OSError as exc:
@@ -885,11 +882,34 @@ def extension_add(
)
raise typer.Exit(1)
format_source = (
final_url
if archive_format_from_name(final_url) is not None
else from_url
)
try:
detect_archive_format(
archive_path,
archive_file=download_file,
source_name=format_source,
content_type=content_type,
error_type=ExtensionError,
)
except ExtensionError:
console.print(
f"[red]Error:[/red] {safe_url} did not return a ZIP archive "
"or tar.gz/tgz archive "
f"(got {len(archive_data)} bytes). This usually means "
"the request was not authenticated and a login/HTML page was "
"returned. Verify the URL and configured credentials."
)
raise typer.Exit(1)
# Consume the transient inode reserved above rather
# than reopening the cache pathname during extraction.
try:
manifest = manager.install_from_zip(
zip_path,
archive_path,
speckit_version,
priority=priority,
force=force,
@@ -918,7 +938,6 @@ def extension_add(
f"{_escape_markup(str(e))}"
)
raise typer.Exit(1)
else:
# Try bundled extensions first (shipped with spec-kit)
bundled_path = _locate_bundled_extension(extension)
@@ -977,18 +996,21 @@ def extension_add(
)
raise typer.Exit(1)
# Download extension ZIP (use resolved ID, not original argument which may be display name)
# Download extension archive (use the resolved catalog ID).
extension_id = ext_info['id']
console.print(f"Downloading {_escape_markup(str(ext_info['name']))} v{_escape_markup(str(ext_info.get('version', 'unknown')))}...")
zip_path = catalog.download_extension(extension_id)
archive_path = catalog.download_extension(extension_id)
try:
# Install from downloaded ZIP
manifest = manager.install_from_zip(zip_path, speckit_version, priority=priority, force=force)
manifest = manager.install_from_zip(
archive_path,
speckit_version,
priority=priority,
force=force,
)
finally:
# Clean up downloaded ZIP
if zip_path.exists():
zip_path.unlink()
if archive_path.exists():
archive_path.unlink()
console.print("\n[green]✓[/green] Extension installed successfully!")
console.print(f"\n[bold]{_escape_markup(str(manifest.name))}[/bold] (v{_escape_markup(str(manifest.version))})")
@@ -1199,7 +1221,7 @@ def extension_search(
console.print(f"\n [yellow]⚠[/yellow] Not directly installable from '{catalog_name}'.")
console.print(
f" Add to an approved catalog with install_allowed: true, "
f"or install from a ZIP URL: specify extension add {safe_id} --from <zip-url>"
f"or install from an archive URL: specify extension add {safe_id} --from <archive-url>"
)
console.print()
@@ -1827,131 +1849,105 @@ def extension_update(
backup_hooks[hook_name] = ext_hooks
# 5. Download new version
zip_path = catalog.download_extension(extension_id)
archive_path = catalog.download_extension(extension_id)
try:
# 6. Validate extension ID from ZIP BEFORE modifying installation
# Handle both root-level and nested extension.yml (GitHub auto-generated ZIPs)
with open_zip_bounded(zip_path) as zf:
import yaml
manifest_data = None
manifest_bytes = None
namelist = zf.namelist()
# Read the manifest under a hard size cap: this happens
# before install_from_zip()'s safe_extract_zip(), so a
# raw zf.open().read() here would bypass that bound and
# let a zip-bomb extension.yml exhaust memory.
# Normalize separators before choosing the manifest so
# this pre-scan cannot approve one entry while extraction
# later overwrites it with a backslash alias.
manifest_candidates = []
archive_entries = []
for name in namelist:
normalized_name = normalize_zip_member_name(name)
parts = normalized_name.removesuffix("/").split(
"/"
)
path_key = portable_zip_path_key(normalized_name)
archive_entries.append(
(normalized_name, parts)
)
# 6. Validate the archive and extension ID before modifying
# the existing installation. The shared extractor applies
# the same bounded security checks to ZIP and tar archives.
with tempfile.TemporaryDirectory(
prefix="speckit-update-archive-"
) as archive_tmpdir:
extracted_root = Path(archive_tmpdir)
try:
safe_extract_archive(archive_path, extracted_root)
except ValueError as exc:
if (
len(parts) in {1, 2}
and path_key[-1] == "extension.yml"
"Conflicting path" in str(exc)
and "extension.yml" in str(exc).casefold()
):
manifest_candidates.append(
(name, normalized_name, path_key)
)
seen_manifest_keys = {}
for name, _normalized_name, path_key in manifest_candidates:
previous = seen_manifest_keys.get(path_key)
if previous is not None:
raise ValueError(
"Downloaded extension archive contains multiple "
"extension.yml manifests"
)
seen_manifest_keys[path_key] = name
for _name, normalized_name, _path_key in manifest_candidates:
if normalized_name.split("/")[-1] != "extension.yml":
raise ValueError(
"Downloaded extension archive manifest "
"filenames must use canonical "
"'extension.yml' casing"
)
root_manifest = next(
) from exc
raise
manifest_root = extracted_root
top_level = list(extracted_root.iterdir())
root_manifest_entries = [
entry
for entry in top_level
if entry.name.casefold() == "extension.yml"
]
if any(
entry.name != "extension.yml"
for entry in root_manifest_entries
):
raise ValueError(
"Archive must use canonical 'extension.yml' casing"
)
canonical_root_manifest = next(
(
name
for name, _normalized_name, path_key
in manifest_candidates
if path_key == ("extension.yml",)
entry
for entry in root_manifest_entries
if entry.name == "extension.yml"
),
None,
)
nested_manifests = [
(name, normalized_name)
for name, normalized_name, path_key
in manifest_candidates
if len(path_key) == 2
and path_key[-1] == "extension.yml"
]
manifest_path = root_manifest
if manifest_path is None and len(nested_manifests) == 1:
manifest_path, normalized_manifest_path = (
nested_manifests[0]
)
manifest_root = normalized_manifest_path.split(
"/", 1
)[0]
top_level_dirs = {
parts[0]
for normalized_name, parts in archive_entries
if (
len(parts) > 1
or normalized_name.endswith("/")
)
}
if top_level_dirs != {manifest_root}:
if canonical_root_manifest is not None:
manifest_path = canonical_root_manifest
else:
top_level_dirs = [
entry for entry in top_level if entry.is_dir()
]
if len(top_level_dirs) != 1:
raise ValueError(
"Downloaded extension archive with a "
"nested extension.yml must contain exactly "
"Downloaded extension archive must contain exactly "
"one top-level directory"
)
if manifest_path is not None:
manifest_bytes = read_zip_member_limited(
zf, manifest_path
manifest_root = top_level_dirs[0]
nested_manifest_entries = [
entry
for entry in manifest_root.iterdir()
if entry.name.casefold() == "extension.yml"
]
if any(
entry.name != "extension.yml"
for entry in nested_manifest_entries
):
raise ValueError(
"Archive must use canonical 'extension.yml' casing"
)
manifest_path = next(
(
entry
for entry in nested_manifest_entries
if entry.name == "extension.yml"
),
manifest_root / "extension.yml",
)
parsed_manifest = yaml.safe_load(
manifest_bytes
if not manifest_path.is_file():
raise ValueError(
"Downloaded extension archive is missing 'extension.yml'"
)
manifest_data = (
parsed_manifest
if parsed_manifest is not None
else {}
)
if manifest_data is None:
raise ValueError("Downloaded extension archive is missing 'extension.yml'")
manifest_bytes = manifest_path.read_bytes()
parsed_manifest = yaml.safe_load(manifest_bytes)
manifest_data = (
parsed_manifest if parsed_manifest is not None else {}
)
if not isinstance(manifest_data, dict):
raise ValueError(
"Invalid extension manifest in downloaded archive: expected YAML mapping"
"Invalid extension manifest in downloaded archive: "
"expected YAML mapping"
)
extension_data = manifest_data.get("extension", {})
if not isinstance(extension_data, dict):
raise ValueError(
"Invalid extension manifest in downloaded archive: expected 'extension' mapping"
"Invalid extension manifest in downloaded archive: "
"expected 'extension' mapping"
)
# Run the same manifest and compatibility validation as a
# normal install while the existing extension is still
# untouched. Reuse the exact bounded bytes selected above.
if manifest_bytes is None:
raise ValueError(
"Downloaded extension archive is missing 'extension.yml'"
)
with tempfile.TemporaryDirectory(
prefix="speckit-update-manifest-"
) as manifest_tmpdir:
@@ -2147,7 +2143,7 @@ def extension_update(
manager.remove(extension_id, keep_config=True)
# 8. Install new version
_ = manager.install_from_zip(zip_path, speckit_version)
_ = manager.install_from_zip(archive_path, speckit_version)
# Restore user config files from backup after successful install.
new_extension_dir = manager.extensions_dir / extension_id
@@ -2193,12 +2189,12 @@ def extension_update(
hook["enabled"] = False
hook_executor.save_project_config(config)
finally:
# ZIP cleanup is housekeeping: never replace an install
# Archive cleanup is housekeeping: never replace an install
# error or roll back an already committed update because a
# scanner temporarily locks the download on Windows.
if zip_path.exists():
if archive_path.exists():
try:
zip_path.unlink()
archive_path.unlink()
except OSError as error:
zip_cleanup_error = error

View File

@@ -27,11 +27,14 @@ from packaging import version as pkg_version
from packaging.specifiers import SpecifierSet, InvalidSpecifier
from .._download_security import (
archive_format_from_name,
archive_suffix,
MAX_JSON_CATALOG_BYTES,
build_safe_download_path,
detect_archive_format,
is_https_or_localhost_http,
read_response_limited,
safe_extract_zip,
safe_extract_archive,
)
from ..extensions import REINSTALL_COMMAND, ExtensionRegistry, normalize_priority
from .._init_options import (
@@ -3534,17 +3537,17 @@ class PresetManager:
return
_materialize_constitution_template(self.project_root, memory_constitution)
def install_from_zip(
def install_from_archive(
self,
zip_path: Path,
archive_path: Path,
speckit_version: str,
priority: int = 10,
force: bool = False,
) -> PresetManifest:
"""Install preset from ZIP file.
"""Install a preset from a supported archive.
Args:
zip_path: Path to preset ZIP file
archive_path: Path to a .zip, .tar.gz, or .tgz archive
speckit_version: Current spec-kit version
priority: Resolution priority (lower = higher precedence, default 10)
force: If True and the preset is already installed, remove it first
@@ -3563,7 +3566,11 @@ class PresetManager:
with tempfile.TemporaryDirectory() as tmpdir:
temp_path = Path(tmpdir)
safe_extract_zip(zip_path, temp_path, error_type=PresetValidationError)
safe_extract_archive(
archive_path,
temp_path,
error_type=PresetValidationError,
)
pack_dir = temp_path
manifest_path = pack_dir / "preset.yml"
@@ -3576,11 +3583,26 @@ class PresetManager:
if not manifest_path.exists():
raise PresetValidationError(
"No preset.yml found in ZIP file"
"No preset.yml found in archive"
)
return self.install_from_directory(pack_dir, speckit_version, priority, force=force)
def install_from_zip(
self,
zip_path: Path,
speckit_version: str,
priority: int = 10,
force: bool = False,
) -> PresetManifest:
"""Backward-compatible wrapper for archive installation."""
return self.install_from_archive(
zip_path,
speckit_version,
priority,
force=force,
)
def remove(self, pack_id: str) -> bool:
"""Remove an installed preset.
@@ -4605,14 +4627,14 @@ class PresetCatalog:
def download_pack(
self, pack_id: str, target_dir: Optional[Path] = None
) -> Path:
"""Download preset ZIP from catalog.
"""Download a preset archive from a catalog.
Args:
pack_id: ID of the preset to download
target_dir: Directory to save ZIP file (defaults to cache directory)
target_dir: Directory to save the archive
Returns:
Path to downloaded ZIP file
Path to the downloaded archive
Raises:
PresetError: If pack not found or download fails
@@ -4681,42 +4703,86 @@ class PresetCatalog:
target_dir = self.cache_dir / "downloads"
target_dir = Path(target_dir)
version = pack_info.get("version", "unknown")
zip_path = build_safe_download_path(
declared_format = archive_format_from_name(download_url)
build_safe_download_path(
target_dir,
pack_id,
version,
error_type=PresetError,
label="preset",
suffix=archive_suffix(declared_format or "tar.gz"),
)
target_dir.mkdir(parents=True, exist_ok=True)
original_download_url = download_url
extra_headers = None
resolved_download_url = self._resolve_github_release_asset_api_url(download_url)
if resolved_download_url:
download_url = resolved_download_url
extra_headers = {"Accept": "application/octet-stream"}
staging_path: Path | None = None
try:
with self._open_url(download_url, timeout=60, extra_headers=extra_headers) as response:
zip_data = read_response_limited(
archive_data = read_response_limited(
response,
error_type=PresetError,
label=f"preset '{pack_id}' download",
)
final_url = (
response.geturl()
if hasattr(response, "geturl")
else download_url
)
content_type = (
response.getheader("Content-Type")
if hasattr(response, "getheader")
else None
)
verify_archive_sha256(
zip_data, pack_info.get("sha256"), pack_id, PresetError
archive_data, pack_info.get("sha256"), pack_id, PresetError
)
zip_path.write_bytes(zip_data)
return zip_path
with tempfile.NamedTemporaryFile(
prefix="preset-download-",
suffix=".archive",
dir=target_dir,
delete=False,
) as staging_file:
staging_path = Path(staging_file.name)
staging_file.write(archive_data)
archive_format = detect_archive_format(
staging_path,
source_name=(
final_url
if archive_format_from_name(final_url) is not None
else original_download_url
),
content_type=content_type,
error_type=PresetError,
)
archive_path = build_safe_download_path(
target_dir,
pack_id,
version,
error_type=PresetError,
label="preset",
suffix=archive_suffix(archive_format),
)
os.replace(staging_path, archive_path)
staging_path = None
return archive_path
except urllib.error.URLError as e:
raise PresetError(
f"Failed to download preset from {download_url}: {e}"
)
except IOError as e:
raise PresetError(f"Failed to save preset ZIP: {e}")
raise PresetError(f"Failed to save preset archive: {e}")
finally:
if staging_path is not None:
staging_path.unlink(missing_ok=True)
def clear_cache(self):
"""Clear all catalog cache files, including per-URL hashed caches."""

View File

@@ -17,6 +17,9 @@ from rich.markup import escape as _escape_markup
from .._console import console
from .._download_security import (
archive_format_from_name,
archive_suffix,
detect_archive_format,
is_https_or_localhost_http,
is_safe_download_redirect,
read_response_limited,
@@ -75,7 +78,11 @@ def preset_list():
@preset_app.command("add")
def preset_add(
preset_id: str = typer.Argument(None, help="Preset ID to install from catalog"),
from_url: str = typer.Option(None, "--from", help="Install from a URL (ZIP file)"),
from_url: str = typer.Option(
None,
"--from",
help="Install from a .zip, .tar.gz, or .tgz URL",
),
dev: str = typer.Option(None, "--dev", help="Install from local directory (development mode)"),
priority: int = typer.Option(10, "--priority", help="Resolution priority (lower = higher precedence, default 10)"),
):
@@ -142,7 +149,7 @@ def preset_add(
import tempfile
with tempfile.TemporaryDirectory() as tmpdir:
zip_path = Path(tmpdir) / "preset.zip"
archive_path = Path(tmpdir) / "preset.archive"
try:
from specify_cli.authentication.http import open_url as _open_url
from specify_cli.authentication.http import github_provider_hosts
@@ -170,13 +177,33 @@ def preset_add(
"or HTTP for localhost (127.0.0.1, ::1)."
)
raise typer.Exit(1)
zip_path.write_bytes(
read_response_limited(
response,
error_type=PresetError,
label=f"preset {from_url}",
)
archive_data = read_response_limited(
response,
error_type=PresetError,
label=f"preset {from_url}",
)
content_type = (
response.getheader("Content-Type")
if hasattr(response, "getheader")
else None
)
archive_path.write_bytes(archive_data)
format_source = (
final_url
if archive_format_from_name(final_url) is not None
else from_url
)
archive_format = detect_archive_format(
archive_path,
source_name=format_source,
content_type=content_type,
error_type=PresetError,
)
detected_path = archive_path.with_suffix(
archive_suffix(archive_format)
)
os.replace(archive_path, detected_path)
archive_path = detected_path
except (urllib.error.URLError, PresetError) as e:
console.print(
f"[red]Error:[/red] Failed to download: "
@@ -184,7 +211,11 @@ def preset_add(
)
raise typer.Exit(1)
manifest = manager.install_from_zip(zip_path, speckit_version, priority)
manifest = manager.install_from_zip(
archive_path,
speckit_version,
priority,
)
console.print(f"[green]✓[/green] Preset '{manifest.name}' v{manifest.version} installed (priority {priority})")
@@ -227,12 +258,16 @@ def preset_add(
console.print(f"Installing preset [cyan]{pack_info.get('name', preset_id)}[/cyan]...")
try:
zip_path = catalog.download_pack(preset_id)
manifest = manager.install_from_zip(zip_path, speckit_version, priority)
archive_path = catalog.download_pack(preset_id)
manifest = manager.install_from_zip(
archive_path,
speckit_version,
priority,
)
console.print(f"[green]✓[/green] Preset '{manifest.name}' v{manifest.version} installed (priority {priority})")
finally:
if 'zip_path' in locals() and zip_path.exists():
zip_path.unlink(missing_ok=True)
if 'archive_path' in locals() and archive_path.exists():
archive_path.unlink(missing_ok=True)
else:
console.print("[red]Error:[/red] Specify a preset ID, --from URL, or --dev path")
raise typer.Exit(1)

View File

@@ -21,10 +21,17 @@ from rich.markup import escape as _escape_markup
from .._console import console, err_console
from .._download_security import (
archive_format_from_content_type,
archive_format_from_name,
archive_suffix,
detect_archive_format,
is_https_or_localhost_http,
is_safe_download_redirect,
read_response_limited,
safe_extract_archive,
)
from .._project import _resolve_init_dir_override
from ..shared_infra import verify_archive_sha256
workflow_app = typer.Typer(
name="workflow",
@@ -455,6 +462,42 @@ def _read_response_within_limit(response, max_bytes: int | None = None) -> bytes
return b"".join(chunks)
def _workflow_yaml_is_declared(
source_name: str, content_type: str | None
) -> bool:
"""Return whether response metadata explicitly identifies workflow YAML."""
from urllib.parse import urlparse
path = urlparse(source_name).path.casefold()
media_type = (content_type or "").split(";", 1)[0].strip().casefold()
return path.endswith((".yml", ".yaml")) or media_type in {
"application/yaml",
"application/x-yaml",
"text/yaml",
"text/x-yaml",
}
def _sniff_workflow_archive_format(data: bytes):
"""Return a supported archive format when suffixless response bytes match."""
from io import BytesIO
try:
return detect_archive_format(
Path("workflow-download"),
archive_file=BytesIO(data),
)
except ValueError:
return None
def _enforce_workflow_yaml_size(data: bytes) -> None:
if len(data) > _MAX_WORKFLOW_YAML_BYTES:
raise ValueError(
f"response exceeds the {_MAX_WORKFLOW_YAML_BYTES}-byte workflow size limit"
)
def _validate_workflow_id_or_exit(workflow_id: str) -> None:
"""Validate that ``workflow_id`` is a safe installed-workflow directory name."""
if (
@@ -879,6 +922,231 @@ def _discard_committed_backup_file(backup_file: Path | None) -> None:
)
def _workflow_package_root(extracted_root: Path) -> Path:
"""Resolve a root-level or single-nested workflow package."""
if (extracted_root / "workflow.yml").is_file():
return extracted_root
entries = list(extracted_root.iterdir())
if (
len(entries) == 1
and entries[0].is_dir()
and not entries[0].is_symlink()
and (entries[0] / "workflow.yml").is_file()
):
return entries[0]
raise ValueError(
"Archive must contain workflow.yml at its root or in exactly one "
"top-level directory"
)
def _validate_local_workflow_package(package_dir: Path) -> None:
"""Reject links and special files before copying a local package."""
import stat
for root, dirnames, filenames in os.walk(package_dir, followlinks=False):
root_path = Path(root)
for name in [*dirnames, *filenames]:
path = root_path / name
mode = path.lstat().st_mode
if stat.S_ISLNK(mode):
raise ValueError(f"Workflow package contains symlink: {path}")
if not stat.S_ISDIR(mode) and not stat.S_ISREG(mode):
raise ValueError(f"Workflow package contains unsupported file: {path}")
def _workflow_package_has_companions(package_dir: Path) -> bool:
"""Return whether a directory contains anything beyond workflow.yml."""
return any(path.name != "workflow.yml" for path in package_dir.iterdir())
def _install_workflow_package(
project_root: Path,
workflows_dir: Path,
package_dir: Path,
source_label: str,
*,
expected_id: str | None = None,
expected_version: str | None = None,
expected_installed_version: str | None = None,
catalog_info: dict[str, Any] | None = None,
) -> None:
"""Validate and atomically install a complete workflow package directory."""
import shutil
import tempfile
from .engine import WorkflowDefinition, validate_workflow
workflow_file = package_dir / "workflow.yml"
try:
_validate_local_workflow_package(package_dir)
workflow_bytes = workflow_file.read_bytes()
definition = WorkflowDefinition.from_string(workflow_bytes.decode("utf-8"))
except (OSError, UnicodeDecodeError, ValueError, yaml.YAMLError) as exc:
console.print(
f"[red]Error:[/red] Invalid workflow package: "
f"{_escape_markup(str(exc))}"
)
raise typer.Exit(1)
errors = validate_workflow(definition)
if errors:
console.print("[red]Error:[/red] Workflow validation failed:")
for error in errors:
console.print(f"{_escape_markup(str(error))}")
raise typer.Exit(1)
if not isinstance(definition.id, str) or not definition.id.strip():
console.print("[red]Error:[/red] Workflow definition has an empty or missing 'id'")
raise typer.Exit(1)
if expected_id is not None and definition.id != expected_id:
console.print(
f"[red]Error:[/red] Workflow ID in YAML "
f"({_escape_markup(repr(definition.id))}) does not match the requested "
f"workflow ID ({_escape_markup(repr(expected_id))})."
)
raise typer.Exit(1)
if expected_version is not None and str(definition.version) != expected_version:
console.print(
f"[red]Error:[/red] Downloaded workflow version "
f"({_escape_markup(str(definition.version))}) does not match the catalog "
f"version ({_escape_markup(expected_version)})."
)
raise typer.Exit(1)
dest_dir = _safe_workflow_id_dir(workflows_dir, definition.id)
staged_dir = Path(
tempfile.mkdtemp(prefix=f".{definition.id}.installing-", dir=workflows_dir)
)
try:
package_root = package_dir.resolve()
def ignore_reserved_package_entries(
source: str, names: list[str]
) -> set[str]:
if Path(source).resolve() == package_root and "overlays" in names:
return {"overlays"}
return set()
shutil.copytree(
package_dir,
staged_dir,
dirs_exist_ok=True,
ignore=ignore_reserved_package_entries,
)
except OSError as exc:
shutil.rmtree(staged_dir, ignore_errors=True)
console.print(
f"[red]Error:[/red] Failed to stage workflow package: "
f"{_escape_markup(str(exc))}"
)
raise typer.Exit(1)
backup_dir: Path | None = None
try:
with _workflow_install_transaction(project_root):
registry = _open_workflow_registry(project_root)
existing = registry.get(definition.id)
if expected_installed_version is not None and (
not isinstance(existing, dict)
or existing.get("source") != "catalog"
or str(existing.get("version")) != expected_installed_version
):
console.print(
f"[yellow]Warning:[/yellow] Workflow "
f"'{_escape_markup(definition.id)}' changed during update; "
"rerun the command."
)
raise typer.Exit(1)
if dest_dir.exists():
backup_dir = Path(
tempfile.mkdtemp(
prefix=f".{definition.id}.backup-",
dir=workflows_dir,
)
)
backup_dir.rmdir()
os.replace(dest_dir, backup_dir)
try:
os.replace(staged_dir, dest_dir)
except BaseException:
if backup_dir is not None:
os.replace(backup_dir, dest_dir)
backup_dir = None
raise
entry = {
"name": definition.name,
"version": definition.version,
"description": definition.description,
"source": source_label,
}
if catalog_info is not None:
entry.update(
{
"source": "catalog",
"catalog_name": catalog_info.get("_catalog_name", ""),
"url": catalog_info.get("url", ""),
}
)
if isinstance(existing, dict) and not existing.get("enabled", True):
entry["enabled"] = False
try:
registry.add(definition.id, entry)
except (OSError, TypeError, ValueError):
failed_dir: Path | None = None
try:
failed_dir = Path(
tempfile.mkdtemp(
prefix=f".{definition.id}.failed-",
dir=workflows_dir,
)
)
failed_dir.rmdir()
os.replace(dest_dir, failed_dir)
if backup_dir is not None:
os.replace(backup_dir, dest_dir)
backup_dir = None
except OSError as rollback_exc:
console.print(
"[yellow]Warning:[/yellow] Failed to fully restore the prior "
f"workflow package: {_escape_markup(str(rollback_exc))}"
)
finally:
if failed_dir is not None and failed_dir.exists():
try:
shutil.rmtree(failed_dir)
except OSError as cleanup_exc:
console.print(
"[yellow]Warning:[/yellow] Could not remove failed "
f"workflow package: {_escape_markup(str(cleanup_exc))}"
)
raise
except typer.Exit:
raise
except (OSError, TypeError, ValueError) as exc:
console.print(
f"[red]Error:[/red] Failed to install workflow package: "
f"{_escape_markup(str(exc))}"
)
raise typer.Exit(1)
finally:
if staged_dir.exists():
shutil.rmtree(staged_dir, ignore_errors=True)
if backup_dir is not None:
try:
shutil.rmtree(backup_dir)
except OSError as exc:
console.print(
"[yellow]Warning:[/yellow] Workflow installed, but its backup "
f"directory could not be removed: {_escape_markup(str(exc))}"
)
console.print(
f"[green]✓[/green] Workflow '{_escape_markup(definition.name)}' "
f"({_escape_markup(definition.id)}) installed"
)
# Root helper re-fetched at call time so test monkeypatching of
# `specify_cli._require_specify_project` keeps working after the move.
def _require_specify_project(*args, **kwargs):
@@ -1602,16 +1870,48 @@ def workflow_add(
if dev_path.is_file() and dev_path.suffix.lower() in (".yml", ".yaml"):
_validate_and_install_local(dev_path, str(dev_path))
return
if dev_path.is_file() and archive_format_from_name(str(dev_path)) is not None:
import tempfile
with tempfile.TemporaryDirectory(
prefix="speckit-workflow-archive-"
) as tmpdir:
extracted_root = Path(tmpdir)
try:
safe_extract_archive(dev_path, extracted_root)
package_root = _workflow_package_root(extracted_root)
except ValueError as exc:
console.print(
f"[red]Error:[/red] Invalid workflow archive: "
f"{_escape_markup(str(exc))}"
)
raise typer.Exit(1)
_install_workflow_package(
project_root,
workflows_dir,
package_root,
str(dev_path),
)
return
if dev_path.is_dir():
dev_wf_file = dev_path / "workflow.yml"
if not dev_wf_file.is_file():
console.print(f"[red]Error:[/red] No workflow.yml found in {_escape_markup(source)}")
raise typer.Exit(1)
_validate_and_install_local(dev_wf_file, str(dev_path))
if _workflow_package_has_companions(dev_path):
_install_workflow_package(
project_root,
workflows_dir,
dev_path,
str(dev_path),
)
else:
_validate_and_install_local(dev_wf_file, str(dev_path))
return
console.print(
"[red]Error:[/red] --dev source must be a workflow YAML file or a "
f"directory containing workflow.yml: {_escape_markup(source)}"
"[red]Error:[/red] --dev source must be a workflow YAML file, "
"supported archive, or directory containing workflow.yml: "
f"{_escape_markup(source)}"
)
raise typer.Exit(1)
@@ -1674,6 +1974,7 @@ def workflow_add(
import tempfile
tmp_path: Path | None = None
downloaded_archive_format = None
try:
with _open_url(
download_url,
@@ -1687,13 +1988,48 @@ def workflow_add(
f"[red]Error:[/red] URL redirected to non-HTTPS: {_escape_markup(final_url)}"
)
raise typer.Exit(1)
with tempfile.NamedTemporaryFile(suffix=".yml", delete=False) as tmp:
content_type = (
resp.getheader("Content-Type")
if hasattr(resp, "getheader")
else None
)
downloaded_archive_format = (
archive_format_from_name(final_url)
or archive_format_from_name(download_url)
or archive_format_from_content_type(content_type)
)
declared_yaml = _workflow_yaml_is_declared(final_url, content_type)
suffix = (
archive_suffix(downloaded_archive_format)
if downloaded_archive_format is not None
else ".yml" if declared_yaml else ".download"
)
with tempfile.NamedTemporaryFile(suffix=suffix, delete=False) as tmp:
# Assign tmp_path immediately: NamedTemporaryFile(delete=False)
# creates the file on disk right away, before any bytes are
# written, so a failure in the size-limited read below must
# still be able to find and remove it.
tmp_path = Path(tmp.name)
tmp.write(_read_response_within_limit(resp))
if downloaded_archive_format is not None:
downloaded_content = read_response_limited(
resp,
error_type=ValueError,
label="workflow archive download",
)
elif declared_yaml:
downloaded_content = _read_response_within_limit(resp)
else:
downloaded_content = read_response_limited(
resp,
error_type=ValueError,
label="workflow download",
)
downloaded_archive_format = (
_sniff_workflow_archive_format(downloaded_content)
)
if downloaded_archive_format is None:
_enforce_workflow_yaml_size(downloaded_content)
tmp.write(downloaded_content)
except typer.Exit:
raise
except Exception as exc:
@@ -1713,13 +2049,38 @@ def workflow_add(
console.print(f"[red]Error:[/red] Failed to download workflow: {_escape_markup(str(exc))}")
raise typer.Exit(1)
try:
# When installed via --from, the positional argument names the
# workflow the user expects — enforce it like the catalog branch.
_validate_and_install_local(
tmp_path,
download_url,
expected_id=source if from_url else None,
)
if downloaded_archive_format is None:
_validate_and_install_local(
tmp_path,
download_url,
expected_id=source if from_url else None,
)
else:
with tempfile.TemporaryDirectory(
prefix="speckit-workflow-archive-"
) as extract_dir:
extracted_root = Path(extract_dir)
try:
safe_extract_archive(
tmp_path,
extracted_root,
source_name=final_url,
content_type=content_type,
)
package_root = _workflow_package_root(extracted_root)
except ValueError as exc:
console.print(
f"[red]Error:[/red] Invalid workflow archive: "
f"{_escape_markup(str(exc))}"
)
raise typer.Exit(1)
_install_workflow_package(
project_root,
workflows_dir,
package_root,
download_url,
expected_id=source if from_url else None,
)
finally:
# Best-effort: _validate_and_install_local may already have
# committed the file + registry entry (success) or already
@@ -1743,12 +2104,46 @@ def workflow_add(
if source_path.is_file() and source_path.suffix.lower() in (".yml", ".yaml"):
_validate_and_install_local(source_path, str(source_path))
return
elif (
source_path.is_file()
and archive_format_from_name(str(source_path)) is not None
):
import tempfile
with tempfile.TemporaryDirectory(
prefix="speckit-workflow-archive-"
) as tmpdir:
extracted_root = Path(tmpdir)
try:
safe_extract_archive(source_path, extracted_root)
package_root = _workflow_package_root(extracted_root)
except ValueError as exc:
console.print(
f"[red]Error:[/red] Invalid workflow archive: "
f"{_escape_markup(str(exc))}"
)
raise typer.Exit(1)
_install_workflow_package(
project_root,
workflows_dir,
package_root,
str(source_path),
)
return
elif source_path.is_dir():
wf_file = source_path / "workflow.yml"
if not wf_file.is_file():
console.print(f"[red]Error:[/red] No workflow.yml found in {_escape_markup(source)}")
raise typer.Exit(1)
_validate_and_install_local(wf_file, str(source_path))
if _workflow_package_has_companions(source_path):
_install_workflow_package(
project_root,
workflows_dir,
source_path,
str(source_path),
)
else:
_validate_and_install_local(wf_file, str(source_path))
return
# Try from catalog
@@ -1853,6 +2248,9 @@ def _install_workflow_from_catalog(
)
raise typer.Exit(1)
original_workflow_url = workflow_url
downloaded_archive_format = None
archive_content_type = None
try:
from specify_cli.authentication.http import open_url as _open_url
from specify_cli.authentication.http import github_provider_hosts as _github_provider_hosts
@@ -1884,10 +2282,38 @@ def _install_workflow_from_catalog(
f"[red]Error:[/red] Workflow '{safe_wf_id}' redirected to non-HTTPS URL: {_escape_markup(final_url)}"
)
raise typer.Exit(1)
archive_content_type = (
response.getheader("Content-Type")
if hasattr(response, "getheader")
else None
)
downloaded_archive_format = (
archive_format_from_name(final_url)
or archive_format_from_name(original_workflow_url)
or archive_format_from_content_type(archive_content_type)
)
# Written to the staging file, never workflow_file directly, so a
# reinstall's prior working copy is never touched until the
# atomic commit below runs.
downloaded_content = _read_response_within_limit(response)
if downloaded_archive_format is not None:
downloaded_content = read_response_limited(
response,
error_type=ValueError,
label=f"workflow '{workflow_id}' archive download",
)
elif _workflow_yaml_is_declared(final_url, archive_content_type):
downloaded_content = _read_response_within_limit(response)
else:
downloaded_content = read_response_limited(
response,
error_type=ValueError,
label=f"workflow '{workflow_id}' download",
)
downloaded_archive_format = _sniff_workflow_archive_format(
downloaded_content
)
if downloaded_archive_format is None:
_enforce_workflow_yaml_size(downloaded_content)
staged_file.write_bytes(downloaded_content)
except typer.Exit:
raise
@@ -1896,6 +2322,59 @@ def _install_workflow_from_catalog(
console.print(f"[red]Error:[/red] Failed to install workflow '{safe_wf_id}' from catalog: {_escape_markup(str(exc))}")
raise typer.Exit(1)
if downloaded_archive_format is not None:
try:
verify_archive_sha256(
downloaded_content,
info.get("sha256"),
workflow_id,
ValueError,
)
import tempfile
from io import BytesIO
with tempfile.TemporaryDirectory(
prefix="speckit-workflow-archive-"
) as extract_dir:
extracted_root = Path(extract_dir)
safe_extract_archive(
staged_file.path,
extracted_root,
archive_file=BytesIO(downloaded_content),
source_name=original_workflow_url,
content_type=archive_content_type,
)
package_root = _workflow_package_root(extracted_root)
_safe_discard_staged_workflow_file(
staged_file,
workflow_dir,
existed_before,
)
_install_workflow_package(
project_root,
workflows_dir,
package_root,
workflow_url,
expected_id=workflow_id,
expected_version=expected_version,
expected_installed_version=expected_installed_version,
catalog_info={**info, "url": workflow_url},
)
except typer.Exit:
raise
except (OSError, ValueError) as exc:
_safe_discard_staged_workflow_file(
staged_file,
workflow_dir,
existed_before,
)
console.print(
f"[red]Error:[/red] Invalid workflow archive: "
f"{_escape_markup(str(exc))}"
)
raise typer.Exit(1)
return
# Validate the downloaded workflow (still staged, not yet committed)
# before registering.
try:

View File

@@ -5,6 +5,7 @@ from __future__ import annotations
import io
import stat
import struct
import tarfile
import weakref
import zipfile
import zlib
@@ -13,11 +14,16 @@ import pytest
from specify_cli._download_security import (
MAX_ZIP_CENTRAL_DIRECTORY_BYTES,
archive_format_from_content_type,
archive_format_from_name,
build_safe_download_path,
detect_archive_format,
is_https_or_localhost_http,
is_loopback_url,
read_response_limited,
read_zip_member_limited,
safe_extract_archive,
safe_extract_tar,
safe_extract_zip,
)
@@ -314,6 +320,176 @@ def test_build_safe_download_path_rejects_nonportable_identifiers(
)
@pytest.mark.parametrize(
("name", "expected"),
[
("package.zip", "zip"),
("PACKAGE.TAR.GZ", "tar.gz"),
("https://example.com/package.tgz?download=1", "tar.gz"),
("package.tar", None),
],
)
def test_archive_format_from_name(name, expected):
assert archive_format_from_name(name) == expected
@pytest.mark.parametrize(
("content_type", "expected"),
[
("application/zip", "zip"),
("application/x-zip-compressed; charset=binary", "zip"),
("application/gzip", "tar.gz"),
("application/x-gzip", "tar.gz"),
("application/octet-stream", None),
],
)
def test_archive_format_from_content_type(content_type, expected):
assert archive_format_from_content_type(content_type) == expected
def _write_tar_gz(path, members):
with tarfile.open(path, "w:gz") as archive:
for name, content in members:
info = tarfile.TarInfo(name)
info.size = len(content)
archive.addfile(info, io.BytesIO(content))
@pytest.mark.parametrize("suffix", [".tar.gz", ".tgz"])
def test_detect_archive_format_accepts_tar_suffixes(tmp_path, suffix):
archive_path = tmp_path / f"package{suffix}"
_write_tar_gz(archive_path, [("file.txt", b"contents")])
assert detect_archive_format(archive_path) == "tar.gz"
def test_detect_archive_format_allows_content_type_fallback(tmp_path):
archive_path = tmp_path / "download"
_write_tar_gz(archive_path, [("file.txt", b"contents")])
assert (
detect_archive_format(
archive_path,
source_name="https://example.com/download",
content_type="application/gzip",
)
== "tar.gz"
)
def test_detect_archive_format_rejects_suffix_content_mismatch(tmp_path):
archive_path = tmp_path / "package.zip"
_write_tar_gz(archive_path, [("file.txt", b"contents")])
with pytest.raises(ValueError, match="format mismatch"):
detect_archive_format(archive_path)
def test_detect_archive_format_rejects_suffix_header_mismatch(tmp_path):
archive_path = tmp_path / "package.zip"
with zipfile.ZipFile(archive_path, "w") as archive:
archive.writestr("file.txt", "contents")
with pytest.raises(ValueError, match="Content-Type"):
detect_archive_format(
archive_path,
content_type="application/gzip",
)
def test_build_safe_download_path_uses_archive_suffix(tmp_path):
path = build_safe_download_path(tmp_path, "package", "1.0.0", suffix=".tar.gz")
assert path.name == "package-1.0.0.tar.gz"
@pytest.mark.parametrize(
"member_name",
["../evil.txt", "nested/../../evil.txt", "C:/Windows/evil.txt"],
)
def test_safe_extract_tar_rejects_traversal(tmp_path, member_name):
archive_path = tmp_path / "bad.tar.gz"
_write_tar_gz(archive_path, [(member_name, b"nope")])
with pytest.raises(ValueError, match="Unsafe path"):
safe_extract_tar(archive_path, tmp_path / "out")
@pytest.mark.parametrize(
("link_type", "message"),
[(tarfile.SYMTYPE, "symlink"), (tarfile.LNKTYPE, "hard link")],
)
def test_safe_extract_tar_rejects_links_without_partial_extraction(
tmp_path, link_type, message
):
archive_path = tmp_path / "bad.tar.gz"
with tarfile.open(archive_path, "w:gz") as archive:
safe = tarfile.TarInfo("safe.txt")
safe.size = 4
archive.addfile(safe, io.BytesIO(b"safe"))
link = tarfile.TarInfo("escape")
link.type = link_type
link.linkname = "../../outside"
archive.addfile(link)
out_dir = tmp_path / "out"
with pytest.raises(ValueError, match=message):
safe_extract_tar(archive_path, out_dir)
assert not out_dir.exists() or not any(out_dir.rglob("*"))
def test_safe_extract_tar_rejects_special_file(tmp_path):
archive_path = tmp_path / "bad.tar.gz"
with tarfile.open(archive_path, "w:gz") as archive:
fifo = tarfile.TarInfo("pipe")
fifo.type = tarfile.FIFOTYPE
archive.addfile(fifo)
with pytest.raises(ValueError, match="Unsafe member type"):
safe_extract_tar(archive_path, tmp_path / "out")
def test_safe_extract_tar_rejects_conflicting_paths(tmp_path):
archive_path = tmp_path / "bad.tar.gz"
_write_tar_gz(
archive_path,
[("Folder/file.txt", b"one"), ("folder/FILE.txt", b"two")],
)
with pytest.raises(ValueError, match="Conflicting path"):
safe_extract_tar(archive_path, tmp_path / "out")
def test_safe_extract_tar_enforces_entry_and_size_limits(tmp_path):
archive_path = tmp_path / "bad.tar.gz"
_write_tar_gz(
archive_path,
[("one.txt", b"1234"), ("two.txt", b"5678")],
)
with pytest.raises(ValueError, match="too many entries"):
safe_extract_tar(archive_path, tmp_path / "entries", max_entries=1)
with pytest.raises(ValueError, match="member.*maximum size"):
safe_extract_tar(archive_path, tmp_path / "member", max_member_bytes=3)
with pytest.raises(ValueError, match="uncompressed size"):
safe_extract_tar(archive_path, tmp_path / "total", max_total_bytes=7)
@pytest.mark.parametrize("suffix", [".zip", ".tar.gz", ".tgz"])
def test_safe_extract_archive_has_format_parity(tmp_path, suffix):
archive_path = tmp_path / f"package{suffix}"
if suffix == ".zip":
with zipfile.ZipFile(archive_path, "w") as archive:
archive.writestr("nested/file.txt", b"contents")
else:
_write_tar_gz(archive_path, [("nested/file.txt", b"contents")])
out_dir = tmp_path / f"out-{suffix.replace('.', '-')}"
safe_extract_archive(archive_path, out_dir)
assert (out_dir / "nested" / "file.txt").read_bytes() == b"contents"
@pytest.mark.parametrize(
"member_name",
[

View File

@@ -2288,6 +2288,52 @@ class TestExtensionManager:
assert manifest.id == "test-ext"
assert manager.registry.is_installed("test-ext")
@pytest.mark.parametrize("suffix", [".tar.gz", ".tgz"])
@pytest.mark.parametrize("nested", [False, True])
def test_install_from_tar_archive(
self, extension_dir, project_dir, temp_dir, suffix, nested
):
"""Tar archives install with the same flat/nested behavior as ZIP."""
import tarfile
archive_path = temp_dir / f"test-ext{suffix}"
with tarfile.open(archive_path, "w:gz") as archive:
for file_path in extension_dir.rglob("*"):
if file_path.is_file():
relative = file_path.relative_to(extension_dir)
arcname = Path("test-ext-v1") / relative if nested else relative
archive.add(file_path, arcname=arcname)
manager = ExtensionManager(project_dir)
manifest = manager.install_from_archive(archive_path, "0.1.0")
assert manifest.id == "test-ext"
assert manager.registry.is_installed("test-ext")
def test_install_from_tar_rejects_symlink_entry(
self, extension_dir, project_dir, temp_dir
):
import tarfile
archive_path = temp_dir / "symlink-extension.tar.gz"
with tarfile.open(archive_path, "w:gz") as archive:
for file_path in extension_dir.rglob("*"):
if file_path.is_file():
archive.add(
file_path,
arcname=file_path.relative_to(extension_dir),
)
link = tarfile.TarInfo("templates/escape")
link.type = tarfile.SYMTYPE
link.linkname = "../../outside"
archive.addfile(link)
manager = ExtensionManager(project_dir)
with pytest.raises(ValidationError, match="Unsafe symlink"):
manager.install_from_archive(archive_path, "0.1.0")
assert not manager.registry.is_installed("test-ext")
assert not manager.registry.is_installed("test-ext")
def test_install_duplicate_error_mentions_force(self, extension_dir, project_dir):
"""Test that duplicate install error message suggests --force."""
manager = ExtensionManager(project_dir)
@@ -5765,6 +5811,35 @@ class TestExtensionCatalog:
assert captured[0].get_header("Authorization") == "Bearer ghp_testtoken"
assert captured[0].get_header("Accept") == "application/octet-stream"
@pytest.mark.parametrize("suffix", [".tar.gz", ".tgz"])
def test_download_extension_preserves_tar_archive_format(
self, temp_dir, suffix
):
import tarfile
from unittest.mock import patch
archive_buffer = io.BytesIO()
with tarfile.open(fileobj=archive_buffer, mode="w:gz") as archive:
content = b"extension:\n id: test-ext\n"
member = tarfile.TarInfo("extension.yml")
member.size = len(content)
archive.addfile(member, io.BytesIO(content))
archive_bytes = archive_buffer.getvalue()
catalog = self._make_catalog(temp_dir)
ext_info = {
"id": "test-ext",
"name": "Test Extension",
"version": "1.0.0",
"download_url": f"https://example.com/test-ext{suffix}",
}
with patch.object(catalog, "get_extension_info", return_value=ext_info), \
patch.object(catalog, "_open_url", return_value=self._mock_response(archive_bytes)):
archive_path = catalog.download_extension("test-ext", target_dir=temp_dir)
assert archive_path.name == "test-ext-1.0.0.tar.gz"
assert archive_path.read_bytes() == archive_bytes
# ===== CatalogEntry Tests =====
@@ -7852,7 +7927,7 @@ class TestDownloadExtensionBundled:
}
mock_response = MagicMock()
mock_response.read.side_effect = io.BytesIO(b"fake zip data").read
mock_response.read.side_effect = io.BytesIO(_MINIMAL_ZIP_BYTES).read
mock_response.__enter__ = lambda s: s
mock_response.__exit__ = MagicMock(return_value=False)
mock_response.geturl.return_value = "https://example.com/catalog.json"

View File

@@ -14,6 +14,7 @@ import pytest
import io
import json
import tempfile
import tarfile
import shutil
import warnings
import zipfile
@@ -670,6 +671,27 @@ class TestPresetManager:
assert manifest.id == "test-pack"
assert manager.registry.is_installed("test-pack")
def test_install_from_zip_forwards_force(
self, project_dir, pack_dir, temp_dir
):
"""The compatibility wrapper must retain forced reinstall behavior."""
zip_path = temp_dir / "test-pack.zip"
with zipfile.ZipFile(zip_path, "w") as zf:
for file_path in pack_dir.rglob("*"):
if file_path.is_file():
zf.write(file_path, file_path.relative_to(pack_dir))
manager = PresetManager(project_dir)
manager.install_from_directory(pack_dir, "0.1.5")
manifest = manager.install_from_zip(
zip_path,
"0.1.5",
force=True,
)
assert manifest.id == "test-pack"
assert manager.registry.is_installed("test-pack")
def test_install_from_zip_nested(self, project_dir, pack_dir, temp_dir):
"""Test installing from ZIP with nested directory."""
zip_path = temp_dir / "test-pack.zip"
@@ -715,6 +737,45 @@ class TestPresetManager:
assert not manager.registry.is_installed("test-pack")
@pytest.mark.parametrize("suffix", [".tar.gz", ".tgz"])
@pytest.mark.parametrize("nested", [False, True])
def test_install_from_tar_archive(
self, project_dir, pack_dir, temp_dir, suffix, nested
):
"""Tar archives install with the same flat/nested behavior as ZIP."""
archive_path = temp_dir / f"test-pack{suffix}"
with tarfile.open(archive_path, "w:gz") as archive:
for file_path in pack_dir.rglob("*"):
if file_path.is_file():
relative = file_path.relative_to(pack_dir)
arcname = Path("test-pack-v1") / relative if nested else relative
archive.add(file_path, arcname=arcname)
manager = PresetManager(project_dir)
manifest = manager.install_from_archive(archive_path, "0.1.5")
assert manifest.id == "test-pack"
assert manager.registry.is_installed("test-pack")
def test_install_from_tar_rejects_symlink_entry(
self, project_dir, pack_dir, temp_dir
):
archive_path = temp_dir / "symlink-preset.tar.gz"
with tarfile.open(archive_path, "w:gz") as archive:
for file_path in pack_dir.rglob("*"):
if file_path.is_file():
archive.add(file_path, arcname=file_path.relative_to(pack_dir))
link = tarfile.TarInfo("templates/escape")
link.type = tarfile.SYMTYPE
link.linkname = "../../outside"
archive.addfile(link)
manager = PresetManager(project_dir)
with pytest.raises(PresetValidationError, match="Unsafe symlink"):
manager.install_from_archive(archive_path, "0.1.5")
assert not manager.registry.is_installed("test-pack")
def test_remove(self, project_dir, pack_dir):
"""Test removing a preset."""
manager = PresetManager(project_dir)
@@ -2668,6 +2729,39 @@ class TestPresetCatalog:
assert captured[0].get_header("Authorization") == "Bearer ghp_testtoken"
assert captured[0].get_header("Accept") == "application/octet-stream"
@pytest.mark.parametrize("suffix", [".tar.gz", ".tgz"])
def test_download_pack_preserves_tar_archive_format(
self, project_dir, suffix
):
from unittest.mock import patch, MagicMock
archive_buffer = io.BytesIO()
with tarfile.open(fileobj=archive_buffer, mode="w:gz") as archive:
content = b"preset:\n id: test-pack\n"
member = tarfile.TarInfo("preset.yml")
member.size = len(content)
archive.addfile(member, io.BytesIO(content))
archive_bytes = archive_buffer.getvalue()
response = MagicMock()
response.read.side_effect = io.BytesIO(archive_bytes).read
response.__enter__.return_value = response
response.__exit__.return_value = False
catalog = PresetCatalog(project_dir)
pack_info = {
"id": "test-pack",
"name": "Test Pack",
"version": "1.0.0",
"download_url": f"https://example.com/test-pack{suffix}",
"_install_allowed": True,
}
with patch.object(catalog, "get_pack_info", return_value=pack_info), \
patch.object(catalog, "_open_url", return_value=response):
archive_path = catalog.download_pack("test-pack", target_dir=project_dir)
assert archive_path.name == "test-pack-1.0.0.tar.gz"
assert archive_path.read_bytes() == archive_bytes
# ===== Integration Tests =====
@@ -10084,7 +10178,7 @@ class TestBundledPresetLocator:
self.read_sizes.append(size)
return super().read(size)
response = FakeResponse(b"zip-bytes")
response = FakeResponse(b"PK\x05\x06" + b"\x00" * 18)
installed = {}
def fake_install_from_zip(self, zip_path, speckit_version, priority=10):
@@ -10105,7 +10199,7 @@ class TestBundledPresetLocator:
assert response.read_sizes
assert installed == {
"zip_bytes": b"zip-bytes",
"zip_bytes": b"PK\x05\x06" + b"\x00" * 18,
"speckit_version": "0.6.0",
"priority": 7,
}

View File

@@ -13,11 +13,14 @@ Covers:
from __future__ import annotations
import json
import os
import shutil
import stat
import sys
import tarfile
import tempfile
import zipfile
from pathlib import Path
import pytest
@@ -11478,6 +11481,25 @@ steps:
)
return d
def _archive_workflow_dir(self, source_dir, archive_path, nested=False):
prefix = Path("align-wf-v1") if nested else Path()
if archive_path.name.lower().endswith(".zip"):
with zipfile.ZipFile(archive_path, "w") as archive:
for file_path in source_dir.rglob("*"):
if file_path.is_file():
archive.write(
file_path,
prefix / file_path.relative_to(source_dir),
)
else:
with tarfile.open(archive_path, "w:gz") as archive:
for file_path in source_dir.rglob("*"):
if file_path.is_file():
archive.add(
file_path,
arcname=prefix / file_path.relative_to(source_dir),
)
def _install_dev(self, runner, app, project_dir):
src = self._write_workflow_dir(project_dir)
result = runner.invoke(app, ["workflow", "add", str(src), "--dev"])
@@ -11496,6 +11518,44 @@ steps:
self._install_dev(runner, app, project_dir)
assert WorkflowRegistry(project_dir).is_installed("align-wf")
def test_add_local_directory_preserves_package_files(
self, project_dir, monkeypatch
):
from typer.testing import CliRunner
from specify_cli import app
monkeypatch.chdir(project_dir)
source = self._write_workflow_dir(project_dir)
(source / "scripts").mkdir()
(source / "scripts" / "helper.sh").write_text("echo helper\n")
result = CliRunner().invoke(app, ["workflow", "add", str(source)])
assert result.exit_code == 0, result.output
installed = project_dir / ".specify" / "workflows" / "align-wf"
assert (installed / "scripts" / "helper.sh").read_text() == "echo helper\n"
@pytest.mark.parametrize("suffix", [".zip", ".tar.gz", ".tgz"])
@pytest.mark.parametrize("nested", [False, True])
def test_add_local_archive_preserves_package_files(
self, project_dir, monkeypatch, suffix, nested
):
from typer.testing import CliRunner
from specify_cli import app
monkeypatch.chdir(project_dir)
source = self._write_workflow_dir(project_dir)
(source / "assets").mkdir()
(source / "assets" / "message.txt").write_text("hello\n")
archive_path = project_dir / f"align-wf{suffix}"
self._archive_workflow_dir(source, archive_path, nested=nested)
result = CliRunner().invoke(app, ["workflow", "add", str(archive_path)])
assert result.exit_code == 0, result.output
installed = project_dir / ".specify" / "workflows" / "align-wf"
assert (installed / "assets" / "message.txt").read_text() == "hello\n"
def test_add_dev_yaml_file_installs(self, project_dir, monkeypatch):
from typer.testing import CliRunner
from specify_cli import app
@@ -11886,6 +11946,208 @@ steps:
assert result.exit_code == 0, result.output
assert WorkflowRegistry(project_dir).is_installed("align-wf")
@pytest.mark.parametrize("suffix", [".zip", ".tar.gz", ".tgz"])
def test_add_from_url_installs_complete_archive_package(
self, project_dir, monkeypatch, suffix
):
from unittest.mock import patch
from typer.testing import CliRunner
from specify_cli import app
monkeypatch.chdir(project_dir)
source = self._write_workflow_dir(project_dir)
(source / "assets").mkdir()
(source / "assets" / "remote.txt").write_text("remote\n")
archive_path = project_dir / f"remote{suffix}"
self._archive_workflow_dir(source, archive_path)
data = archive_path.read_bytes()
url = f"https://example.com/align-wf{suffix}"
with patch(
"specify_cli.authentication.http.open_url",
side_effect=lambda *_args, **_kwargs: self._FakeResponse(data, url),
):
result = CliRunner().invoke(
app,
["workflow", "add", "align-wf", "--from", url],
input="y\n",
)
assert result.exit_code == 0, result.output
installed = project_dir / ".specify" / "workflows" / "align-wf"
assert (installed / "assets" / "remote.txt").read_text() == "remote\n"
@pytest.mark.parametrize("suffix", [".zip", ".tar.gz", ".tgz"])
def test_add_from_suffixless_url_sniffs_archive(
self, project_dir, monkeypatch, suffix
):
from unittest.mock import patch
from typer.testing import CliRunner
from specify_cli import app
monkeypatch.chdir(project_dir)
source = self._write_workflow_dir(project_dir)
(source / "assets").mkdir()
(source / "assets" / "sniffed.txt").write_text("sniffed\n")
archive_path = project_dir / f"remote{suffix}"
self._archive_workflow_dir(source, archive_path)
data = archive_path.read_bytes()
url = "https://example.com/assets/12345"
with patch(
"specify_cli.authentication.http.open_url",
side_effect=lambda *_args, **_kwargs: self._FakeResponse(
data,
url,
{"Content-Type": "application/octet-stream"},
),
):
result = CliRunner().invoke(
app,
["workflow", "add", "align-wf", "--from", url],
input="y\n",
)
assert result.exit_code == 0, result.output
installed = project_dir / ".specify" / "workflows" / "align-wf"
assert (installed / "assets" / "sniffed.txt").read_text() == "sniffed\n"
@pytest.mark.parametrize("suffix", [".zip", ".tar.gz", ".tgz"])
def test_add_catalog_installs_complete_archive_package_and_sha(
self, project_dir, monkeypatch, suffix
):
import hashlib
from unittest.mock import patch
from typer.testing import CliRunner
from specify_cli import app
from specify_cli.workflows.catalog import WorkflowCatalog
monkeypatch.chdir(project_dir)
source = self._write_workflow_dir(project_dir)
(source / "assets").mkdir()
(source / "assets" / "catalog.txt").write_text("catalog\n")
archive_path = project_dir / f"catalog{suffix}"
self._archive_workflow_dir(source, archive_path, nested=True)
data = archive_path.read_bytes()
url = f"https://example.com/align-wf{suffix}"
info = {
"id": "align-wf",
"name": "Align Workflow",
"version": "1.0.0",
"url": url,
"sha256": hashlib.sha256(data).hexdigest(),
"_install_allowed": True,
"_catalog_name": "test",
}
with patch.object(
WorkflowCatalog,
"get_workflow_info",
return_value=info,
), patch(
"specify_cli.authentication.http.open_url",
side_effect=lambda *_args, **_kwargs: self._FakeResponse(data, url),
):
result = CliRunner().invoke(app, ["workflow", "add", "align-wf"])
assert result.exit_code == 0, result.output
installed = project_dir / ".specify" / "workflows" / "align-wf"
assert (installed / "assets" / "catalog.txt").read_text() == "catalog\n"
@pytest.mark.parametrize("suffix", [".zip", ".tar.gz", ".tgz"])
def test_add_catalog_sniffs_suffixless_archive(
self, project_dir, monkeypatch, suffix
):
import hashlib
from unittest.mock import patch
from typer.testing import CliRunner
from specify_cli import app
from specify_cli.workflows.catalog import WorkflowCatalog
monkeypatch.chdir(project_dir)
source = self._write_workflow_dir(project_dir)
(source / "assets").mkdir()
(source / "assets" / "sniffed.txt").write_text("catalog sniffed\n")
archive_path = project_dir / f"catalog{suffix}"
self._archive_workflow_dir(source, archive_path, nested=True)
data = archive_path.read_bytes()
url = "https://example.com/assets/67890"
info = {
"id": "align-wf",
"name": "Align Workflow",
"version": "1.0.0",
"url": url,
"sha256": hashlib.sha256(data).hexdigest(),
"_install_allowed": True,
"_catalog_name": "test",
}
with patch.object(
WorkflowCatalog,
"get_workflow_info",
return_value=info,
), patch(
"specify_cli.authentication.http.open_url",
side_effect=lambda *_args, **_kwargs: self._FakeResponse(
data,
url,
{"Content-Type": "application/octet-stream"},
),
):
result = CliRunner().invoke(app, ["workflow", "add", "align-wf"])
assert result.exit_code == 0, result.output
installed = project_dir / ".specify" / "workflows" / "align-wf"
assert (
installed / "assets" / "sniffed.txt"
).read_text() == "catalog sniffed\n"
def test_package_registry_failure_restores_before_failed_cleanup(
self, project_dir, monkeypatch
):
import shutil
from unittest.mock import patch
from typer.testing import CliRunner
from specify_cli import app
from specify_cli.workflows.catalog import WorkflowRegistry
monkeypatch.chdir(project_dir)
source = self._write_workflow_dir(project_dir, version="1.0.0")
(source / "assets").mkdir()
(source / "assets" / "version.txt").write_text("old\n")
runner = CliRunner()
first = runner.invoke(app, ["workflow", "add", str(source)])
assert first.exit_code == 0, first.output
(source / "workflow.yml").write_text(
self.WORKFLOW_YAML.format(version="2.0.0"),
encoding="utf-8",
)
(source / "assets" / "version.txt").write_text("new\n")
real_rmtree = shutil.rmtree
def fail_failed_package_cleanup(path, *args, **kwargs):
if ".failed-" in Path(path).name:
raise OSError("cleanup denied")
return real_rmtree(path, *args, **kwargs)
with patch.object(
WorkflowRegistry,
"add",
side_effect=OSError("registry save failed"),
), patch(
"shutil.rmtree",
side_effect=fail_failed_package_cleanup,
):
result = runner.invoke(app, ["workflow", "add", str(source)])
assert result.exit_code == 1, result.output
installed = project_dir / ".specify" / "workflows" / "align-wf"
assert "1.0.0" in (installed / "workflow.yml").read_text()
assert (installed / "assets" / "version.txt").read_text() == "old\n"
assert "registry save failed" in result.output
assert "cleanup denied" in result.output
def test_add_from_url_temp_cleanup_failure_after_success_still_exits_zero(
self, project_dir, monkeypatch
):