83 lines
2.7 KiB
Python
83 lines
2.7 KiB
Python
"""HTTP server for serving the RSS feed."""
|
|
|
|
from http.server import BaseHTTPRequestHandler
|
|
from urllib.parse import urlparse
|
|
|
|
from src.db import get_connection
|
|
from src.rss import generate_feed, generate_index_html
|
|
|
|
|
|
def create_handler(db_path: str, base_url: str = None):
|
|
"""Factory that returns an HTTP request handler class bound to a DB path."""
|
|
|
|
class FeedHandler(BaseHTTPRequestHandler):
|
|
def do_GET(self):
|
|
parsed = urlparse(self.path)
|
|
path = parsed.path.rstrip("/") or "/"
|
|
|
|
if path == "/feed.xml":
|
|
self._serve_feed()
|
|
elif path == "/health":
|
|
self._serve_health()
|
|
elif path == "/":
|
|
self._serve_index()
|
|
else:
|
|
self.send_response(404)
|
|
self.send_header("Content-Type", "text/plain")
|
|
self.end_headers()
|
|
self.wfile.write(b"Not Found")
|
|
|
|
def _serve_feed(self):
|
|
conn = get_connection(db_path)
|
|
try:
|
|
from src.db import get_all_entries
|
|
entries = get_all_entries(conn)
|
|
feed_xml = generate_feed(entries, self._base_url())
|
|
finally:
|
|
conn.close()
|
|
|
|
self.send_response(200)
|
|
self.send_header("Content-Type", "application/rss+xml; charset=utf-8")
|
|
self.end_headers()
|
|
try:
|
|
self.wfile.write(feed_xml.encode("utf-8"))
|
|
except BrokenPipeError:
|
|
pass
|
|
|
|
def _serve_index(self):
|
|
html = generate_index_html(self._base_url())
|
|
self.send_response(200)
|
|
self.send_header("Content-Type", "text/html; charset=utf-8")
|
|
self.end_headers()
|
|
try:
|
|
self.wfile.write(html.encode("utf-8"))
|
|
except BrokenPipeError:
|
|
pass
|
|
|
|
def _serve_health(self):
|
|
self.send_response(200)
|
|
self.send_header("Content-Type", "text/plain")
|
|
self.end_headers()
|
|
try:
|
|
self.wfile.write(b"OK")
|
|
except BrokenPipeError:
|
|
pass
|
|
|
|
def _base_url(self) -> str:
|
|
if base_url:
|
|
port = self.server.server_address[1]
|
|
url = base_url.rstrip("/")
|
|
netloc = url.split("://", 1)[-1].split("/", 1)[0]
|
|
if ":" not in netloc:
|
|
url = f"{url}:{port}"
|
|
return url
|
|
host = self.server.server_address[0]
|
|
port = self.server.server_address[1]
|
|
return f"http://{host}:{port}"
|
|
|
|
def log_message(self, format, *args):
|
|
"""Suppress default request logging."""
|
|
pass
|
|
|
|
return FeedHandler
|