104 lines
3.8 KiB
Python
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()
|