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 = """