diff --git a/src/app/main_window.py b/src/app/main_window.py index ca71dc9..3770378 100644 --- a/src/app/main_window.py +++ b/src/app/main_window.py @@ -149,15 +149,17 @@ class MainWindow(QMainWindow): return bar def _start_update_check(self): - """Check the configured LAN source for a newer version, off the UI thread.""" + """Check the configured update source for a newer version, off the UI thread.""" source = self._config.get("update_source", "") if not source: return + user = self._config.get("update_user", "") + password = self._config.get("update_pass", "") self._update_found.connect(self._on_update_found) def worker(): try: - info = check_for_update(source, APP_VERSION) + info = check_for_update(source, APP_VERSION, user, password) except Exception: # never let the thread crash startup logger.exception("Update check failed") return @@ -191,6 +193,9 @@ class MainWindow(QMainWindow): self._config.get("update_source", ""), ] for path in candidates: + if path and path.lower().startswith(("http://", "https://")): + if QDesktopServices.openUrl(QUrl(path)): + return if not path or not Path(path).exists(): continue try: diff --git a/src/services/update_service.py b/src/services/update_service.py index 280e4df..4e797dc 100644 --- a/src/services/update_service.py +++ b/src/services/update_service.py @@ -1,18 +1,21 @@ -"""LAN update check (stage 2: notify-only). +"""Update check (stage 2: notify-only). -Reads a manifest.json from a configured LAN folder and reports whether a newer -version is advertised. This module performs no installation — it only decides -whether to notify the user. See docs/10-lan-update.md. +Reads a manifest.json from a configured local folder or HTTP(S) source and +reports whether a newer version is advertised. This module performs no +installation — it only decides whether to notify the user. See +docs/10-lan-update.md. All functions degrade gracefully: a missing/unreachable source or a malformed manifest yields "no update" rather than an error, so an update check can never block or break startup. """ +import base64 import json import logging from dataclasses import dataclass from pathlib import Path from typing import Optional +from urllib import error, parse, request logger = logging.getLogger(__name__) @@ -23,7 +26,7 @@ MANIFEST_NAME = "manifest.json" class UpdateInfo: """A newer version advertised by the update source.""" version: str - source: str # folder the user opens to get the new version + source: str # folder or URL the user opens to get the new version notes: str = "" mandatory: bool = False @@ -53,7 +56,52 @@ def is_newer(remote, local) -> bool: return parse_version(remote) > parse_version(local) -def check_for_update(update_source, current_version) -> Optional[UpdateInfo]: +def _is_http_source(update_source): + return parse.urlparse(str(update_source)).scheme.lower() in ("http", "https") + + +def _manifest_url(update_source): + source = str(update_source).strip() + parsed = parse.urlparse(source) + if Path(parsed.path).name.lower() == MANIFEST_NAME: + return source + if not source.endswith("/"): + source += "/" + return parse.urljoin(source, MANIFEST_NAME) + + +def _load_http_manifest(update_source, update_user="", update_pass=""): + url = _manifest_url(update_source) + headers = { + "Accept": "application/json", + "User-Agent": "CMBot", + } + if update_user or update_pass: + raw = "{}:{}".format(update_user or "", update_pass or "").encode("utf-8") + headers["Authorization"] = "Basic {}".format( + base64.b64encode(raw).decode("ascii") + ) + + req = request.Request(url, headers=headers) + with request.urlopen(req, timeout=5) as response: + payload = response.read() + if isinstance(payload, bytes): + payload = payload.decode("utf-8") + return json.loads(payload) + + +def _load_local_manifest(update_source): + manifest_path = Path(update_source) / MANIFEST_NAME + with open(str(manifest_path), encoding="utf-8") as f: + return json.load(f) + + +def check_for_update( + update_source, + current_version, + update_user="", + update_pass="", +) -> Optional[UpdateInfo]: """Return UpdateInfo if *update_source* advertises a version newer than *current_version*, else None. @@ -66,23 +114,26 @@ def check_for_update(update_source, current_version) -> Optional[UpdateInfo]: if not update_source: return None - manifest_path = Path(update_source) / MANIFEST_NAME try: - with open(str(manifest_path), encoding="utf-8") as f: - data = json.load(f) - except (OSError, ValueError) as exc: + if _is_http_source(update_source): + data = _load_http_manifest(update_source, update_user, update_pass) + else: + data = _load_local_manifest(update_source) + except (OSError, ValueError, error.URLError) as exc: logger.info("Update check skipped (%s): %s", type(exc).__name__, exc) return None if not isinstance(data, dict): - logger.warning("Manifest is not a JSON object, ignoring: %s", manifest_path) + logger.warning("Manifest is not a JSON object, ignoring: %s", update_source) return None version = str(data.get("version", "")).strip() if not version or not is_newer(version, current_version): return None - source = str(data.get("source", "")).strip() or str(update_source) + source = str(data.get("url", "") or data.get("source", "")).strip() + if not source: + source = str(update_source) info = UpdateInfo( version=version, source=source, diff --git a/tests/test_update_service.py b/tests/test_update_service.py index 2e21275..72cb8a9 100644 --- a/tests/test_update_service.py +++ b/tests/test_update_service.py @@ -1,10 +1,13 @@ """Tests for services.update_service — no GUI dependency.""" import json +import base64 import shutil import sys import tempfile import unittest from pathlib import Path +from unittest.mock import patch +from urllib import error sys.path.insert(0, str(Path(__file__).parent.parent / "src")) @@ -100,6 +103,88 @@ class TestCheckForUpdate(unittest.TestCase): info = check_for_update(str(self.tmp), "1.0.0") self.assertEqual(info.source, str(self.tmp)) + def test_http_source_appends_manifest_name(self): + calls = {} + + def fake_urlopen(req, timeout=0): + calls["url"] = req.full_url + return _FakeResponse({ + "version": "1.2.0", + "url": "https://example.test/CMBot-1.2.0.zip", + "notes": "http update", + }) + + with patch("services.update_service.request.urlopen", fake_urlopen): + info = check_for_update("https://example.test/releases", "1.0.0") + + self.assertIsInstance(info, UpdateInfo) + self.assertEqual(calls["url"], "https://example.test/releases/manifest.json") + self.assertEqual(info.source, "https://example.test/CMBot-1.2.0.zip") + + def test_http_source_accepts_full_manifest_url(self): + calls = {} + + def fake_urlopen(req, timeout=0): + calls["url"] = req.full_url + return _FakeResponse({"version": "1.2.0"}) + + manifest_url = "https://example.test/releases/manifest.json" + with patch("services.update_service.request.urlopen", fake_urlopen): + info = check_for_update(manifest_url, "1.0.0") + + self.assertIsInstance(info, UpdateInfo) + self.assertEqual(calls["url"], manifest_url) + self.assertEqual(info.source, manifest_url) + + def test_http_basic_auth_header(self): + calls = {} + + def fake_urlopen(req, timeout=0): + calls["auth"] = req.get_header("Authorization") + return _FakeResponse({"version": "2.0.0"}) + + with patch("services.update_service.request.urlopen", fake_urlopen): + check_for_update( + "https://example.test", + "1.0.0", + update_user="readonly", + update_pass="secret", + ) + + token = base64.b64encode(b"readonly:secret").decode("ascii") + self.assertEqual(calls["auth"], "Basic {}".format(token)) + + def test_http_error_returns_none(self): + def fake_urlopen(req, timeout=0): + raise error.URLError("offline") + + with patch("services.update_service.request.urlopen", fake_urlopen): + self.assertIsNone(check_for_update("https://example.test", "1.0.0")) + + def test_http_malformed_json_returns_none(self): + def fake_urlopen(req, timeout=0): + return _FakeResponse("{ not valid json", raw=True) + + with patch("services.update_service.request.urlopen", fake_urlopen): + self.assertIsNone(check_for_update("https://example.test", "1.0.0")) + + +class _FakeResponse: + def __init__(self, data, raw=False): + if raw: + self._payload = data.encode("utf-8") + else: + self._payload = json.dumps(data).encode("utf-8") + + def __enter__(self): + return self + + def __exit__(self, exc_type, exc, tb): + return False + + def read(self): + return self._payload + if __name__ == "__main__": unittest.main()