Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
23 changes: 18 additions & 5 deletions src/specify_cli/_assets.py
Original file line number Diff line number Diff line change
Expand Up @@ -103,19 +103,32 @@ def _locate_bundled_preset(preset_id: str) -> Path | None:

def get_speckit_version() -> str:
"""Get current spec-kit version."""
# Mirror _version._get_installed_version(): a malformed installed
# distribution raises InvalidMetadataError, which is not a
# PackageNotFoundError and must not escape the fallback path.
metadata_errors = [importlib.metadata.PackageNotFoundError]
invalid_metadata_error = getattr(importlib.metadata, "InvalidMetadataError", None)
if invalid_metadata_error is not None:
metadata_errors.append(invalid_metadata_error)

try:
return importlib.metadata.version("specify-cli")
except Exception:
except tuple(metadata_errors):
# Fallback: try reading from pyproject.toml
try:
import tomllib
pyproject_path = _repo_root() / "pyproject.toml"
if pyproject_path.exists():
with open(pyproject_path, "rb") as f:
data = tomllib.load(f)
return data.get("project", {}).get("version", "unknown")
except Exception:
# Intentionally ignore any errors while reading/parsing pyproject.toml.
# If this lookup fails for any reason, we fall back to returning "unknown" below.
project = data.get("project") if isinstance(data, dict) else None
# A present but non-mapping ``project`` table must not turn
# into an AttributeError from the narrowing of this branch.
if isinstance(project, dict):
version = project.get("version")
if isinstance(version, str) and version:
return version
return "unknown"
except (OSError, KeyError, ValueError):
pass
return "unknown"
146 changes: 146 additions & 0 deletions tests/test_utils_assets_imports.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,12 @@
"""Regression guard: utility and asset symbols importable from specify_cli."""
import importlib.metadata

from specify_cli import (
check_tool, merge_json_files,
get_speckit_version,
CLAUDE_LOCAL_PATH, CLAUDE_NPM_LOCAL_PATH,
)
from specify_cli import _assets
from pathlib import Path

def test_utils_symbols_importable():
Expand All @@ -17,3 +20,146 @@ def test_get_speckit_version_returns_string():
def test_claude_paths_are_paths():
assert isinstance(CLAUDE_LOCAL_PATH, Path)
assert isinstance(CLAUDE_NPM_LOCAL_PATH, Path)


def test_get_speckit_version_survives_invalid_metadata(monkeypatch):
"""A corrupt installed distribution must fall back, not raise.

``InvalidMetadataError`` is not a ``PackageNotFoundError``, so catching
only the latter lets it escape the version fallback (the same guard
Comment thread
Quratulain-bilal marked this conversation as resolved.
_version._get_installed_version() already applies). The class is looked up
dynamically and was removed from the stdlib in 3.14, so install a stand-in
to exercise the guard on every supported interpreter.
"""

class _InvalidMetadataError(Exception):
pass

monkeypatch.setattr(
importlib.metadata, "InvalidMetadataError", _InvalidMetadataError, raising=False
)

def _corrupt(name):
raise _InvalidMetadataError("corrupt metadata")

monkeypatch.setattr(importlib.metadata, "version", _corrupt)

assert isinstance(get_speckit_version(), str)


def test_get_speckit_version_survives_non_mapping_project(monkeypatch, tmp_path):
"""A present but non-mapping ``project`` value must not raise.

The narrowed pyproject branch previously called ``.get`` on whatever the
``project`` key held, so ``project = 5`` turned the fallback into an
AttributeError.
"""
(tmp_path / "pyproject.toml").write_text("project = 5\n", encoding="utf-8")
monkeypatch.setattr(_assets, "_repo_root", lambda: tmp_path)

def _not_found(name):
raise importlib.metadata.PackageNotFoundError(name)

monkeypatch.setattr(importlib.metadata, "version", _not_found)

assert get_speckit_version() == "unknown"


def test_get_speckit_version_reads_pyproject_fallback(monkeypatch, tmp_path):
"""When the distribution is missing, a valid pyproject.toml is the source."""
(tmp_path / "pyproject.toml").write_text(
'[project]\nname = "demo"\nversion = "9.9.9"\n', encoding="utf-8"
)
monkeypatch.setattr(_assets, "_repo_root", lambda: tmp_path)

def _not_found(name):
raise importlib.metadata.PackageNotFoundError(name)

monkeypatch.setattr(importlib.metadata, "version", _not_found)

assert get_speckit_version() == "9.9.9"


def test_unrelated_exception_from_version_lookup_propagates(monkeypatch):
"""A TypeError from importlib.metadata.version() must NOT be swallowed.

With the old bare ``except Exception`` this was silently caught and
the function returned "unknown". After the narrowing to
(PackageNotFoundError, InvalidMetadataError) an unrelated exception
must propagate. This test fails against the previous implementation.
"""
def _boom(name):
raise TypeError("unexpected internal error")

monkeypatch.setattr(importlib.metadata, "version", _boom)

try:
get_speckit_version()
except TypeError:
pass # narrowed: TypeError propagates
else:
raise AssertionError("TypeError was swallowed by the fallback")


def test_unrelated_exception_from_pyproject_propagates(monkeypatch, tmp_path):
"""A RuntimeError from pyproject parsing must NOT be swallowed.

With the old bare ``except Exception`` in the inner fallback this
was silently caught. After narrowing to (OSError, KeyError,
ValueError) an unrelated exception must propagate. This test fails
against the previous implementation.
"""
(tmp_path / "pyproject.toml").write_text(
'[project]\nname = "demo"\nversion = "1.0"\n', encoding="utf-8"
)
monkeypatch.setattr(_assets, "_repo_root", lambda: tmp_path)

def _not_found(name):
raise importlib.metadata.PackageNotFoundError(name)

monkeypatch.setattr(importlib.metadata, "version", _not_found)

original_tomllib_load = None
import tomllib

def _boom_load(f):
raise RuntimeError("corrupt tomllib")

monkeypatch.setattr(tomllib, "load", _boom_load)

try:
get_speckit_version()
except RuntimeError:
pass # narrowed: RuntimeError propagates
else:
raise AssertionError("RuntimeError was swallowed by the pyproject fallback")


def test_pyproject_io_error_returns_unknown(monkeypatch, tmp_path):
"""An OSError reading pyproject.toml must return "unknown", not propagate.

This exercises the intended pyproject I/O failure path: the narrowing
keeps OSError in the caught tuple so a missing/unreadable pyproject
still falls through to "unknown".
"""
(tmp_path / "pyproject.toml").write_text(
'[project]\nname = "demo"\nversion = "1.0"\n', encoding="utf-8"
)
monkeypatch.setattr(_assets, "_repo_root", lambda: tmp_path)

def _not_found(name):
raise importlib.metadata.PackageNotFoundError(name)

monkeypatch.setattr(importlib.metadata, "version", _not_found)

import builtins
real_open = builtins.open

def _deny_read(path, *args, **kwargs):
if "pyproject" in str(path):
raise OSError("permission denied")
return real_open(path, *args, **kwargs)

monkeypatch.setattr(builtins, "open", _deny_read)

assert get_speckit_version() == "unknown"