104 lines
3.8 KiB
Python

#!/usr/bin/env python3
"""Loopback-only aggregate receipt service behind the site's exact-path proxy."""
from http.server import BaseHTTPRequestHandler, HTTPServer
from socketserver import ThreadingMixIn
import json
import os
import threading
from receipts import ReceiptStore, Refused
from validation import no_duplicate_fields
MAX_BODY = 4096
class Handler(BaseHTTPRequestHandler):
def log_message(self, *_args):
pass # Never log request paths, headers, tokens, bodies or errors.
def reply(self, status, data=None):
body = json.dumps(data, allow_nan=False).encode() if data is not None else b""
self.send_response(status)
self.send_header("Content-Type", "application/json")
self.send_header("Cache-Control", "no-store")
self.send_header("Content-Length", str(len(body)))
self.end_headers()
if body:
self.wfile.write(body)
def do_GET(self):
self.reply(405)
do_HEAD = do_GET
do_OPTIONS = do_GET
def do_POST(self):
try:
if self.path not in ("/av-test/api/start", "/av-test/api/report"):
return self.reply(404)
if any(len(self.headers.get_all(name, [])) > 1 for name in ("Origin", "Authorization", "Content-Type")):
return self.reply(400)
if self.headers.get("Origin") != self.server.origin:
return self.reply(403)
if self.headers.get("Content-Type") != "application/json" or self.headers.get("Transfer-Encoding"):
return self.reply(415)
sizes = self.headers.get_all("Content-Length", [])
if len(sizes) != 1 or not sizes[0].isdigit() or not 0 < int(sizes[0]) <= MAX_BODY:
return self.reply(413)
raw = self.rfile.read(int(sizes[0]))
if len(raw) != int(sizes[0]):
return self.reply(400)
data = json.loads(raw, object_pairs_hook=no_duplicate_fields)
if self.path == "/av-test/api/start":
if data != {}:
return self.reply(400)
return self.reply(201, self.server.store.start())
authorization = self.headers.get("Authorization", "")
token = authorization[7:] if authorization.startswith("Bearer ") else ""
self.server.store.report(data, token)
self.reply(204)
except Refused as error:
self.reply(error.status)
except (ValueError, UnicodeError, RecursionError, TypeError):
self.reply(400)
except (OSError, TimeoutError):
pass
class Server(ThreadingMixIn, HTTPServer):
daemon_threads = True
def __init__(self, address, origin, store):
self.origin, self.store = origin, store
self.workers = threading.BoundedSemaphore(4)
super().__init__(address, Handler)
def get_request(self):
connection, address = super().get_request()
connection.settimeout(3)
return connection, address
def process_request(self, request, client_address):
if not self.workers.acquire(blocking=False):
self.shutdown_request(request)
return
try:
super().process_request(request, client_address)
except Exception:
self.workers.release()
raise
def process_request_thread(self, request, client_address):
try:
super().process_request_thread(request, client_address)
finally:
self.workers.release()
def handle_error(self, *_args):
pass # No raw exception/request logging, including unexpected input.
if __name__ == "__main__":
store = ReceiptStore(lambda entry: print(json.dumps(entry, allow_nan=False), flush=True))
Server(("127.0.0.1", 8081), os.environ.get("OBSERVER_ORIGIN", "https://bstein.dev"), store).serve_forever()