#!/usr/bin/env python3
from __future__ import annotations

import argparse
import mimetypes
import urllib.error
import urllib.request
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path
from urllib.parse import urlsplit


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument('--host', required=True)
    parser.add_argument('--port', type=int, required=True)
    parser.add_argument('--root', required=True)
    parser.add_argument('--proxy-api-base', default='')
    args = parser.parse_args()
    root = Path(args.root).resolve()
    proxy_api_base = args.proxy_api_base.rstrip('/')

    class Handler(BaseHTTPRequestHandler):
        def do_OPTIONS(self):
            self.send_response(204)
            self.send_header('Access-Control-Allow-Origin', '*')
            self.send_header('Access-Control-Allow-Methods', 'GET,POST,PATCH,OPTIONS')
            self.send_header('Access-Control-Allow-Headers', 'Content-Type, Accept')
            self.end_headers()

        def do_GET(self):
            if proxy_api_base and (self.path == '/api' or self.path.startswith('/api/')):
                return self.proxy()
            return self.serve_file()

        def do_POST(self):
            if proxy_api_base and (self.path == '/api' or self.path.startswith('/api/')):
                return self.proxy()
            self.send_error(405)

        def do_PATCH(self):
            if proxy_api_base and (self.path == '/api' or self.path.startswith('/api/')):
                return self.proxy()
            self.send_error(405)

        def proxy(self):
            target = proxy_api_base + self.path
            body = None
            if self.command in {'POST', 'PATCH'}:
                length = int(self.headers.get('Content-Length', '0') or 0)
                body = self.rfile.read(length) if length else None
            req = urllib.request.Request(target, data=body, method=self.command)
            for key in ['Content-Type', 'Accept']:
                if key in self.headers:
                    req.add_header(key, self.headers[key])
            try:
                with urllib.request.urlopen(req, timeout=15) as resp:
                    data = resp.read()
                    self.send_response(resp.status)
                    self.send_header('Access-Control-Allow-Origin', '*')
                    self.send_header('Cache-Control', 'no-store')
                    ctype = resp.headers.get('Content-Type') or 'application/octet-stream'
                    self.send_header('Content-Type', ctype)
                    self.end_headers()
                    self.wfile.write(data)
            except urllib.error.HTTPError as e:
                data = e.read()
                self.send_response(e.code)
                self.send_header('Access-Control-Allow-Origin', '*')
                self.send_header('Content-Type', e.headers.get('Content-Type') or 'text/plain')
                self.end_headers()
                self.wfile.write(data)
            except Exception as e:
                self.send_error(502, str(e))

        def serve_file(self):
            path = urlsplit(self.path).path
            if path == '/':
                rel = 'index.html'
            else:
                rel = path.lstrip('/')
            candidate = (root / rel).resolve()
            try:
                candidate.relative_to(root)
            except ValueError:
                self.send_error(403)
                return
            if not candidate.is_file():
                candidate = root / 'index.html'
            try:
                data = candidate.read_bytes()
            except OSError:
                self.send_error(404)
                return
            ctype = mimetypes.guess_type(str(candidate))[0] or 'application/octet-stream'
            self.send_response(200)
            self.send_header('Content-Type', ctype)
            self.send_header('Cache-Control', 'no-store')
            self.send_header('X-Content-Type-Options', 'nosniff')
            self.end_headers()
            self.wfile.write(data)

        def log_message(self, fmt, *args):
            return

    ThreadingHTTPServer((args.host, args.port), Handler).serve_forever()

if __name__ == '__main__':
    main()
