diff --git a/bot_bottle/cli/login.py b/bot_bottle/cli/login.py index 6410266..530a69f 100644 --- a/bot_bottle/cli/login.py +++ b/bot_bottle/cli/login.py @@ -18,6 +18,7 @@ import json import os import socket import sys +import tempfile import time import urllib.error import urllib.request @@ -71,7 +72,7 @@ def _save_credentials( ) -> Path: path = bot_bottle_root() / "console.json" path.parent.mkdir(parents=True, exist_ok=True) - path.write_text( + content = ( json.dumps( { "url": console_url, @@ -83,7 +84,19 @@ def _save_credentials( ) + "\n" ) - path.chmod(0o600) + fd, tmp_path_str = tempfile.mkstemp(dir=path.parent, prefix=".console-") + tmp = Path(tmp_path_str) + try: + tmp.chmod(0o600) + with os.fdopen(fd, "w") as f: + f.write(content) + os.replace(tmp, path) + except Exception: + try: + tmp.unlink() + except OSError: + pass + raise return path @@ -111,6 +124,7 @@ def cmd_login(argv: list[str]) -> int: device_code = resp["device_code"] user_code = resp["user_code"] expires_in = resp.get("expires_in", 300) + poll_sleep = max(1, min(int(resp.get("poll_interval", _POLL_SLEEP)), 60)) sys.stderr.write( f"\nOpen this URL in your browser to authorize this host:\n\n" @@ -122,7 +136,7 @@ def cmd_login(argv: list[str]) -> int: while time.monotonic() < deadline: sys.stderr.write(".") sys.stderr.flush() - time.sleep(_POLL_SLEEP) + time.sleep(poll_sleep) try: code, result = _get( diff --git a/tests/unit/test_cli_login.py b/tests/unit/test_cli_login.py index 0cb2c9b..7445591 100644 --- a/tests/unit/test_cli_login.py +++ b/tests/unit/test_cli_login.py @@ -42,6 +42,26 @@ class TestSaveCredentials(unittest.TestCase): self.assertEqual(data["refresh_token"], "rt") self.assertEqual(oct(path.stat().st_mode & 0o777), oct(0o600)) + def test_temp_file_is_private_before_replace(self) -> None: + """Temp file must be 0600 at the moment os.replace is called.""" + from bot_bottle.cli.login import _save_credentials + from pathlib import Path as _Path + + tmp_perms_at_replace: list[int] = [] + real_replace = os.replace + + def _spy_replace(src: "str | os.PathLike", dst: "str | os.PathLike") -> None: + tmp_perms_at_replace.append(_Path(src).stat().st_mode & 0o777) + real_replace(src, dst) + + with tempfile.TemporaryDirectory() as tmp: + with patch.dict(os.environ, {"BOT_BOTTLE_ROOT": tmp}): + with patch("os.replace", side_effect=_spy_replace): + path = _save_credentials("http://c", "hid", "at", "rt") + self.assertEqual(len(tmp_perms_at_replace), 1) + self.assertEqual(oct(tmp_perms_at_replace[0]), oct(0o600)) + self.assertEqual(oct(path.stat().st_mode & 0o777), oct(0o600)) + class TestCmdLoginMissingUrl(unittest.TestCase): def test_returns_1_without_url(self) -> None: @@ -91,7 +111,7 @@ class TestCmdLoginFlow(unittest.TestCase): with patch.dict(os.environ, {"BOT_BOTTLE_ROOT": tmp}): with patch("bot_bottle.cli.login._post", side_effect=_fake_post): with patch("bot_bottle.cli.login._get", side_effect=_fake_get): - with patch("bot_bottle.cli.login._POLL_SLEEP", 0): + with patch("time.sleep"): return cmd_login(["--console-url", "http://console"]) def test_approved_flow_returns_0(self) -> None: @@ -121,17 +141,47 @@ class TestCmdLoginFlow(unittest.TestCase): start_resp = { "device_code": "dc", "user_code": "ZZZ-ZZZ", - "expires_in": 0, # already expired - "poll_interval": 0, + "expires_in": 0, # already expired; loop never runs + "poll_interval": 2, } with tempfile.TemporaryDirectory() as tmp: with patch.dict(os.environ, {"BOT_BOTTLE_ROOT": tmp}): with patch("bot_bottle.cli.login._post", return_value=start_resp): - with patch("bot_bottle.cli.login._POLL_SLEEP", 0): - result = cmd_login(["--console-url", "http://console"]) + result = cmd_login(["--console-url", "http://console"]) self.assertEqual(result, 1) + def test_poll_interval_from_server_is_used(self) -> None: + """time.sleep must be called with the server-provided poll_interval.""" + from bot_bottle.cli.login import cmd_login + + server_interval = 7 + start_resp = { + "device_code": "dc", + "user_code": "ABC-DEF", + "expires_in": 300, + "poll_interval": server_interval, + } + approved = { + "status": "approved", + "host_id": "hid", + "access_token": "at", + "refresh_token": "rt", + } + poll_iter = iter([{"status": "pending"}, approved]) + + with tempfile.TemporaryDirectory() as tmp: + with patch.dict(os.environ, {"BOT_BOTTLE_ROOT": tmp}): + with patch("bot_bottle.cli.login._post", return_value=start_resp): + with patch("bot_bottle.cli.login._get", side_effect=lambda u: (200, next(poll_iter))): + with patch("time.sleep") as mock_sleep: + result = cmd_login(["--console-url", "http://console"]) + + self.assertEqual(result, 0) + self.assertTrue(mock_sleep.called) + for call in mock_sleep.call_args_list: + self.assertEqual(call.args[0], server_interval) + class TestDispatcherRegistration(unittest.TestCase): def test_login_in_commands(self) -> None: