diff --git a/.github/workflows/pyroma.yml b/.github/workflows/pyroma.yml index 175b7d6..33a9995 100644 --- a/.github/workflows/pyroma.yml +++ b/.github/workflows/pyroma.yml @@ -28,7 +28,7 @@ jobs: # install pyroma - name: install pyroma - run: pip install pyroma + run: pip install pyroma==3.3 # run pyroma - name: run pyroma diff --git a/CHANGES.rst b/CHANGES.rst index 1ee5ba9..8e4d01a 100644 --- a/CHANGES.rst +++ b/CHANGES.rst @@ -4,7 +4,10 @@ Changelog 2.2.2 (unreleased) ------------------ -- Nothing changed yet. +- Fix proxycache url replace + [mamico] +- Refactor proxycache server: thread-safety, atomic writes, background thread deduplication, ThreadingHTTPServer, and anti-SSRF URL validation. + [mamico] 2.2.1 (2023-07-12) diff --git a/src/redturtle/rssservice/__init__.py b/src/redturtle/rssservice/__init__.py index 39c370b..56047e4 100644 --- a/src/redturtle/rssservice/__init__.py +++ b/src/redturtle/rssservice/__init__.py @@ -1,6 +1,6 @@ # -*- coding: utf-8 -*- """Init and utils.""" -from zope.i18nmessageid import MessageFactory +from zope.i18nmessageid import MessageFactory _ = MessageFactory("design.plone.rssservice") diff --git a/src/redturtle/rssservice/proxycacheserver/main.py b/src/redturtle/rssservice/proxycacheserver/main.py index c439bdd..09c17f4 100644 --- a/src/redturtle/rssservice/proxycacheserver/main.py +++ b/src/redturtle/rssservice/proxycacheserver/main.py @@ -2,13 +2,13 @@ This code implements a caching proxy server that stores and serves web content. Key Components: -* Creates unique filenames for cached content using MD5 hashing +* Creates unique filenames for cached content using SHA256 hashing * Stores both the content and metadata (URL information) in separate files * Background refresh content periodically The Proxy Server: -* Listens for incoming requests (be awere to protect connection or leave the server listen only on localhost) +* Listens for incoming requests (protected against SSRF and concurrent requests) * Checks if requested content is in cache * If found, serves from cache * If not found, fetches it, saves it, then serves it @@ -16,17 +16,17 @@ Background Refresh: * Automatically updates cached content periodically -* Runs in separate threads to not block the main server +* Runs in separate threads to not block the main server (deduplicated per URL) * Time between updates is configurable (TTL - Time To Live) -Command Line Interface : Uses Click library to accept parameters like: +Command Line Interface: Uses Click library to accept parameters like: Host address (default: 127.0.0.1) Port number (default: 8080) Cache directory location (default: ./var/cache) TTL for cache refresh (default: 3600 seconds) -Usage Example : +Usage Example: ``` rssmixer-proxy --host 127.0.0.1 --port 8080 --cache-dir ./var/cache --ttl 3600 @@ -36,12 +36,12 @@ Usage: -``` +```python import requests RSSMIXER_PROXY = "http://127.0.0.1:8080" url = "https://abcnews.go.com/abcnews/usheadlines" -res = requests.get(f{RSS_MIXER_PROXY}/{url}") +res = requests.get(f"{RSSMIXER_PROXY}/{url}") ``` This is particularly useful for: @@ -63,79 +63,109 @@ import socketserver import threading import time +from urllib.parse import urlparse - -# this is not thtread-safe, but we don't care about it ! +LOCK = threading.Lock() LAST_ACCESS_TIMES = {} +ACTIVE_REFRESH_THREADS = set() MAX_TTL_IN_CACHE = 7 * 24 * 3600 # 1 week logger = logging.getLogger("rssmixer-proxy") logger.setLevel(logging.INFO) -formatter = logging.Formatter( - "%(asctime)s - %(name)s - %(levelname)s - %(message)s", datefmt="%Y-%m-%d %H:%M:%S" -) -stream_handler = logging.StreamHandler() -stream_handler.setFormatter(formatter) -logger.addHandler(stream_handler) +if not logger.handlers: + formatter = logging.Formatter( + "%(asctime)s - %(name)s - %(levelname)s - %(message)s", + datefmt="%Y-%m-%d %H:%M:%S", + ) + stream_handler = logging.StreamHandler() + stream_handler.setFormatter(formatter) + logger.addHandler(stream_handler) -# Function to calculate cache file path based on URL def cache_path(url, cache_dir): - hash_url = hashlib.md5(url.encode("utf-8")).hexdigest() + hash_url = hashlib.sha256(url.encode("utf-8")).hexdigest() return os.path.join(cache_dir, f"{hash_url}.json") +def safe_atomic_write_json(file_path, data): + tmp_path = f"{file_path}.tmp.{threading.get_ident()}_{time.time_ns()}" + try: + with open(tmp_path, "w", encoding="utf-8") as f: + json.dump(data, f, indent=2) + os.replace(tmp_path, file_path) + except Exception as e: + logger.error("Error writing atomically to %s: %s", file_path, e) + if os.path.exists(tmp_path): + try: + os.remove(tmp_path) + except OSError: + pass + + def load_json(cache_file): try: if os.path.exists(cache_file): with open(cache_file, "r", encoding="utf-8") as f: return json.load(f) - except Exception: - return {} + except Exception as e: + logger.warning("Error reading cache file %s: %s", cache_file, e) + return {} + +def is_valid_url(url): + """Validate URL format and prevent basic SSRF targets.""" + if not re.match(r"^https?:\/\/", url): + return False + parsed = urlparse(url) + hostname = parsed.hostname + if not hostname: + return False + forbidden_hosts = {"localhost", "127.0.0.1", "0.0.0.0", "169.254.169.254", "::1"} + if hostname.lower() in forbidden_hosts: + return False + return True -def fetch_and_cache(url, cache_dir, client_headers=None, timeout=(1, 10)): + +def fetch_and_cache(url, cache_dir, client_headers=None, timeout=(3, 10)): cache_file = cache_path(url, cache_dir) + headers = {} + + if client_headers is None: + data = load_json(cache_file) + headers = data.get("request_headers", {}) + else: + headers = dict(client_headers) + + headers.setdefault("User-Agent", "RSSMixerProxy/1.0") + headers.pop("Host", None) + + if not is_valid_url(url): + logger.error("Invalid or restricted URL path: %s", url) + return { + "url": url, + "request_headers": headers, + "response_headers": {}, + "status_code": 400, + "body": f"Invalid or restricted URL: {url}", + } + try: - # Send the request to the server - if client_headers is None: - data = load_json(cache_file) - headers = data.get("request_headers", {}) - else: - headers = client_headers - if "User-Agent" not in headers: - headers["User-Agent"] = "RSSMixerProxy/1.0" - if "Host" in headers: - del headers["Host"] - # Validate the URL - if not re.match(r"^https?:\/\/", url): - raise ValueError(f"Invalid URL path: {url}") response = requests.get(url, headers=headers, timeout=timeout) - # Store the response in the cache + cache_content = { + "url": url, + "request_headers": headers, + "response_headers": dict(response.headers), + "status_code": response.status_code, + "body": response.text, + } + if response.status_code == 200: - cache_content = { - "url": url, - "request_headers": headers, - "response_headers": dict(response.headers), - "status_code": response.status_code, - "body": response.text, - } - # TODO: update file only if changed ? - with open(cache_file, "w", encoding="utf-8") as f: - json.dump(cache_content, f, indent=2) + safe_atomic_write_json(cache_file, cache_content) logger.info("Cached %s: %s in %s", response.status_code, url, cache_dir) else: - logger.error("Failed to fetch $s: %s", url, response.status_code) - cache_content = { - "url": url, - "request_headers": headers, - "response_headers": dict(response.headers), - "status_code": response.status_code, - "body": response.text, - } + logger.error("Failed to fetch %s: status %s", url, response.status_code) if not os.path.exists(cache_file): - with open(cache_file, "w", encoding="utf-8") as f: - json.dump(cache_content, f, indent=2) + safe_atomic_write_json(cache_file, cache_content) logger.info( "Cached error %s: %s in %s", response.status_code, url, cache_dir ) @@ -145,53 +175,68 @@ def fetch_and_cache(url, cache_dir, client_headers=None, timeout=(1, 10)): "url": url, "request_headers": headers, "response_headers": {}, - "status_code": 500, + "status_code": 502, "body": str(e), } - if not os.path.exists(cache_file): - with open(cache_file, "w", encoding="utf-8") as f: - json.dump(cache_content, f, indent=2) - logger.info("Cached error: %s in %s", url, cache_dir) + # Do not persist transient connection errors permanently to disk return cache_content -# Background thread to refresh cache def refresh_cache(url, cache_dir, ttl): - logger.info(f"Refresh cache for {url} every {ttl} seconds") - while True: - time.sleep(ttl) - if url not in LAST_ACCESS_TIMES: - LAST_ACCESS_TIMES[url] = time.time() - else: - if LAST_ACCESS_TIMES[url] + MAX_TTL_IN_CACHE < time.time(): - cache_file = cache_path(url, cache_dir) - if os.path.exists(cache_file): - os.remove(cache_file) - logger.warning("Remove %s from cached files", url) - return - logger.info("Refresh cache for %s", url) - fetch_and_cache(url, cache_dir) + logger.info("Refresh cache for %s every %s seconds", url, ttl) + try: + while True: + time.sleep(ttl) + with LOCK: + last_access = LAST_ACCESS_TIMES.get(url, time.time()) + if last_access + MAX_TTL_IN_CACHE < time.time(): + cache_file = cache_path(url, cache_dir) + if os.path.exists(cache_file): + try: + os.remove(cache_file) + except OSError: + pass + LAST_ACCESS_TIMES.pop(url, None) + logger.warning("Remove %s from cached files due to inactivity", url) + return + + logger.info("Refresh cache for %s", url) + fetch_and_cache(url, cache_dir) + finally: + with LOCK: + ACTIVE_REFRESH_THREADS.discard(url) + + +def ensure_refresh_thread(url, cache_dir, ttl): + """Ensure at most one background refresh thread runs per URL.""" + with LOCK: + if url not in ACTIVE_REFRESH_THREADS: + ACTIVE_REFRESH_THREADS.add(url) + threading.Thread( + target=refresh_cache, args=(url, cache_dir, ttl), daemon=True + ).start() -# Load URLs to cache from existing .url files def load_urls_from_cache(cache_dir): urls = [] + if not os.path.exists(cache_dir): + return urls for file in os.listdir(cache_dir): if file.endswith(".json"): hash_file = os.path.join(cache_dir, file) - try: - # Extract original URL from the cached file - data = json.load(open(hash_file, "r", encoding="utf-8")) - url = data.get("url", "") - if url: - logger.info("Load: %s from cache %s", url, hash_file) - urls.append(url) - except Exception as e: - logger.info("Error reading cached file %s: %s", file, e) + data = load_json(hash_file) + url = data.get("url", "") + if url: + logger.info("Load: %s from cache %s", url, hash_file) + urls.append(url) return urls -# HTTP proxy handler +class ThreadingHTTPServer(socketserver.ThreadingMixIn, http.server.HTTPServer): + daemon_threads = True + allow_reuse_address = True + + class CachingProxyHandler(http.server.BaseHTTPRequestHandler): def __init__(self, *args, cache_dir=None, ttl=None, **kwargs): self.cache_dir = cache_dir @@ -199,53 +244,63 @@ def __init__(self, *args, cache_dir=None, ttl=None, **kwargs): super().__init__(*args, **kwargs) def do_GET(self): + url = self.path.lstrip("/").replace("\n", "").replace("\r", "") + with LOCK: + LAST_ACCESS_TIMES[url] = time.time() - url = self.path.lstrip("/").replace("\n").replace("\r") - LAST_ACCESS_TIMES[url] = time.time() cache_file = cache_path(url, self.cache_dir) + cache_content = load_json(cache_file) - # Check if the page is already cached - if os.path.exists(cache_file): + if cache_content: logger.info("Serving from cache: %s", url) - with open(cache_file, "r", encoding="utf-8") as f: - cache_content = json.load(f) else: logger.info("Fetching and caching: %s", url) client_headers = dict(self.headers) cache_content = fetch_and_cache(url, self.cache_dir, client_headers) - threading.Thread( - target=refresh_cache, args=(url, self.cache_dir, self.ttl), daemon=True - ).start() + ensure_refresh_thread(url, self.cache_dir, self.ttl) + + body_str = cache_content.get("body", "") + body_bytes = body_str.encode("utf-8") + status_code = cache_content.get("status_code", 500) - # Send response - self.send_response(cache_content["status_code"]) - for header, value in cache_content["response_headers"].items(): - if header.lower() in ("set-cookie", "content-length"): + self.send_response(status_code) + response_headers = cache_content.get("response_headers", {}) + for header, value in response_headers.items(): + header_lower = header.lower() + if header_lower in ("set-cookie", "content-length"): continue - if header.lower() in ("content-type", "cache-control"): + if header_lower in ( + "content-type", + "cache-control", + "etag", + "last-modified", + ): self.send_header(header, value) - continue - # logger.info("skip header", header, value) - self.send_header("Content-Length", len(cache_content["body"].encode("utf-8"))) + + self.send_header("Content-Length", str(len(body_bytes))) self.end_headers() - self.wfile.write(cache_content["body"].encode("utf-8")) + self.wfile.write(body_bytes) + + def log_message(self, format, *args): + logger.debug( + "%s - - [%s] %s", + self.address_string(), + self.log_date_time_string(), + format % args, + ) -# Start the server def start_server(host, port, cache_dir, ttl): def handler(*args, **kwargs): return CachingProxyHandler(*args, cache_dir=cache_dir, ttl=ttl, **kwargs) - socketserver.TCPServer.allow_reuse_address = True - with socketserver.TCPServer((host, port), handler) as httpd: + with ThreadingHTTPServer((host, port), handler) as httpd: try: logger.info("Serving on http://%s:%s", host, port) httpd.serve_forever() finally: logger.info("Closing connection") httpd.shutdown() - # con.shutdown(socket.SHUT_RDWR) - # httpd.close() @click.command() @@ -254,20 +309,14 @@ def handler(*args, **kwargs): @click.option( "--cache-dir", default="./var/cache", help="Directory to store cached files." ) -@click.option("--ttl", default=3600, help="") +@click.option("--ttl", default=3600, help="TTL for cache refresh in seconds.") def main(host, port, cache_dir, ttl): - # Create cache directory if it doesn't exist os.makedirs(cache_dir, exist_ok=True) - - # Load URLs from cache directory and start refresh threads cached_urls = load_urls_from_cache(cache_dir) try: for url in cached_urls: - threading.Thread( - target=refresh_cache, args=(url, cache_dir, ttl), daemon=True - ).start() + ensure_refresh_thread(url, cache_dir, ttl) - # Start the proxy server start_server(host, port, cache_dir, ttl) except KeyboardInterrupt: logger.info("Server stopped.") diff --git a/src/redturtle/rssservice/rss_mixer.py b/src/redturtle/rssservice/rss_mixer.py index 4d389e1..2270685 100644 --- a/src/redturtle/rssservice/rss_mixer.py +++ b/src/redturtle/rssservice/rss_mixer.py @@ -16,13 +16,12 @@ from zope.i18n import translate from zope.interface import implementer from zope.schema import getFields - +from App.config import getConfiguration import feedparser import json import logging import requests - logger = logging.getLogger(__name__) @@ -37,13 +36,14 @@ REQUESTS_USER_AGENT = environ.get("RSS_USER_AGENT") RSSMIXER_HTTP_PROXY = environ.get("RSSMIXER_PROXY", "") +DEBUGMODE = getConfiguration().debug_mode + class RSSMixerService(Service): """ """ def reply(self): feed_config = self.get_feed_config() - limit = feed_config.get("limit", 20) feeds = feed_config.get("feeds", []) if not feeds: @@ -61,6 +61,9 @@ def reply(self): def get_feed_config(self): """ """ query = self.request.form + if DEBUGMODE and query.get("rss_debug_uri"): + return {"feeds": [{"url": query.get("rss_debug_uri")}]} + block_id = query.get("block", "") if not block_id: raise BadRequest( diff --git a/src/redturtle/rssservice/tests/test_proxycache.py b/src/redturtle/rssservice/tests/test_proxycache.py new file mode 100644 index 0000000..1bc41a3 --- /dev/null +++ b/src/redturtle/rssservice/tests/test_proxycache.py @@ -0,0 +1,91 @@ +# -*- coding: utf-8 -*- +import importlib +import os +import shutil +import tempfile +import unittest +from unittest import mock + +main = importlib.import_module("redturtle.rssservice.proxycacheserver.main") + + +class ProxyCacheServerTest(unittest.TestCase): + def setUp(self): + self.test_dir = tempfile.mkdtemp() + main.ACTIVE_REFRESH_THREADS.clear() + main.LAST_ACCESS_TIMES.clear() + + def tearDown(self): + shutil.rmtree(self.test_dir, ignore_errors=True) + main.ACTIVE_REFRESH_THREADS.clear() + main.LAST_ACCESS_TIMES.clear() + + def test_cache_path(self): + url = "https://example.com/rss.xml" + path = main.cache_path(url, self.test_dir) + self.assertTrue(path.startswith(self.test_dir)) + self.assertTrue(path.endswith(".json")) + + def test_safe_atomic_write_and_load_json(self): + file_path = os.path.join(self.test_dir, "test.json") + data = {"url": "https://example.com/rss", "status_code": 200, "body": "OK"} + main.safe_atomic_write_json(file_path, data) + + loaded = main.load_json(file_path) + self.assertEqual(loaded, data) + # Ensure no temporary file left + files = os.listdir(self.test_dir) + self.assertEqual(files, ["test.json"]) + + def test_is_valid_url(self): + self.assertTrue(main.is_valid_url("https://example.com/feed")) + self.assertTrue(main.is_valid_url("http://news.google.com/rss")) + self.assertFalse(main.is_valid_url("ftp://example.com")) + self.assertFalse(main.is_valid_url("http://localhost:8080")) + self.assertFalse(main.is_valid_url("http://127.0.0.1/admin")) + self.assertFalse(main.is_valid_url("http://169.254.169.254/latest/meta-data/")) + + def test_ensure_refresh_thread_deduplication(self): + url = "https://example.com/feed" + with mock.patch("threading.Thread") as mock_thread_cls: + mock_thread_instance = mock.MagicMock() + mock_thread_cls.return_value = mock_thread_instance + + main.ensure_refresh_thread(url, self.test_dir, ttl=60) + self.assertIn(url, main.ACTIVE_REFRESH_THREADS) + self.assertEqual(mock_thread_cls.call_count, 1) + + # Second call for the same URL should be deduplicated + main.ensure_refresh_thread(url, self.test_dir, ttl=60) + self.assertEqual(mock_thread_cls.call_count, 1) + + @mock.patch("requests.get") + def test_fetch_and_cache_success(self, mock_requests_get): + mock_response = mock.MagicMock() + mock_response.status_code = 200 + mock_response.headers = {"Content-Type": "application/rss+xml"} + mock_response.text = "Test Feed" + mock_requests_get.return_value = mock_response + + url = "https://example.com/rss" + result = main.fetch_and_cache(url, self.test_dir) + + self.assertEqual(result["status_code"], 200) + self.assertEqual(result["body"], mock_response.text) + + # Check cached file + cache_file = main.cache_path(url, self.test_dir) + self.assertTrue(os.path.exists(cache_file)) + loaded = main.load_json(cache_file) + self.assertEqual(loaded["status_code"], 200) + + @mock.patch("requests.get") + def test_fetch_and_cache_connection_error_not_persisted(self, mock_requests_get): + mock_requests_get.side_effect = Exception("Connection refused") + + url = "https://example.com/failing_rss" + result = main.fetch_and_cache(url, self.test_dir) + + self.assertEqual(result["status_code"], 502) + cache_file = main.cache_path(url, self.test_dir) + self.assertFalse(os.path.exists(cache_file)) diff --git a/src/redturtle/rssservice/tests/test_rss_mixer.py b/src/redturtle/rssservice/tests/test_rss_mixer.py index 6c13609..14a8f2c 100644 --- a/src/redturtle/rssservice/tests/test_rss_mixer.py +++ b/src/redturtle/rssservice/tests/test_rss_mixer.py @@ -13,7 +13,6 @@ import unittest - EXAMPLE_FEED_FOO = """