diff --git a/src/memos/cli.py b/src/memos/cli.py index 2ead5ab29..e9a690680 100644 --- a/src/memos/cli.py +++ b/src/memos/cli.py @@ -42,7 +42,10 @@ def download_examples(dest: str) -> bool: print(f"📥 Downloading examples from {zip_url}...") try: - response = requests.get(zip_url) + # Without an explicit timeout a stalled connection hangs the CLI + # forever; requests retries neither, and a dead download should fail + # into the RequestException handler below instead. + response = requests.get(zip_url, timeout=60) response.raise_for_status() with zipfile.ZipFile(BytesIO(response.content)) as z: diff --git a/tests/test_cli_download_timeout.py b/tests/test_cli_download_timeout.py new file mode 100644 index 000000000..d176d97e4 --- /dev/null +++ b/tests/test_cli_download_timeout.py @@ -0,0 +1,44 @@ +import sys +import types + +from pathlib import Path +from unittest.mock import MagicMock, patch + +import pytest + + +SRC_DIR = Path(__file__).resolve().parents[1] / "src" / "memos" + + +@pytest.fixture +def download_examples(): + # memos.cli only needs a namespace stub at package level; the module + # itself has no heavy imports (the FastAPI app import is lazy) + if "memos" not in sys.modules: + memos_pkg = types.ModuleType("memos") + memos_pkg.__path__ = [str(SRC_DIR)] + sys.modules["memos"] = memos_pkg + import memos.cli + + return memos.cli.download_examples + + +# an empty-zip archive header makes the extraction loop a no-op; the point +# of the test is the arguments passed to requests.get +EMPTY_ZIP = b"PK\x05\x06" + b"\x00" * 18 + + +def test_download_examples_sends_timeout(download_examples): + fake_response = MagicMock() + fake_response.content = EMPTY_ZIP + with patch("requests.get", return_value=fake_response) as mock_get: + download_examples("/tmp/memos-cli-examples-test") + _, kwargs = mock_get.call_args + assert kwargs.get("timeout") is not None + + +def test_download_examples_still_succeeds(download_examples, tmp_path): + fake_response = MagicMock() + fake_response.content = EMPTY_ZIP + with patch("requests.get", return_value=fake_response): + assert download_examples(str(tmp_path)) is True