#!/usr/bin/env python3
"""简陋文件室：一个只能下载的 HTTP 服务器（无上传、无删除、无认证之外的写操作）。
用法: python3 file_room.py [--port 1525] [--dir <文件室目录>]
"""
import argparse
import http.server
import os
import socket
import sys
from urllib.parse import unquote

ROOM_DIR = os.path.dirname(os.path.abspath(__file__))


class RoomHandler(http.server.SimpleHTTPRequestHandler):
    DIRECTORY = ROOM_DIR  # 类属性，main() 里可覆盖

    def __init__(self, *args, **kwargs):
        super().__init__(*args, directory=self.DIRECTORY, **kwargs)

    def _deny(self):
        self.send_response(405)
        self.send_header("Content-Type", "text/plain; charset=utf-8")
        self.end_headers()
        self.wfile.write("文件室只读：仅允许下载。\n".encode("utf-8"))

    # ---- 写方法一律拒绝 ----
    def do_POST(self):  self._deny()
    def do_PUT(self):   self._deny()
    def do_DELETE(self):self._deny()
    def do_PATCH(self): self._deny()

    # ---- 只读方法 ----
    def do_GET(self):
        # ?download=1 时强制作为附件下载，而不是内联展示
        if self.path.startswith("/") and "download=1" in self.path:
            self._serve_as_download()
            return
        # 无参数：JSON 也强制下载（避免浏览器内联展示），其他文件正常
        path = unquote(self.path.split("?", 1)[0])
        if path.endswith(".json"):
            self._serve_as_download()
            return
        http.server.SimpleHTTPRequestHandler.do_GET(self)

    def _serve_as_download(self):
        """把目标文件作为附件强制下载。"""
        path = unquote(self.path.split("?", 1)[0])
        # 防止路径穿越
        full = os.path.realpath(os.path.join(self.DIRECTORY, path.lstrip("/")))
        if not full.startswith(os.path.realpath(self.DIRECTORY)):
            self.send_error(403, "Forbidden")
            return
        if not os.path.isfile(full):
            self.send_error(404, "Not Found")
            return
        fname = os.path.basename(full)
        self.send_response(200)
        self.send_header("Content-Type", "application/octet-stream")
        self.send_header("Content-Disposition", f'attachment; filename="{fname}"')
        self.send_header("Content-Length", str(os.path.getsize(full)))
        self.end_headers()
        with open(full, "rb") as f:
            while True:
                chunk = f.read(65536)
                if not chunk:
                    break
                self.wfile.write(chunk)

    do_HEAD = http.server.SimpleHTTPRequestHandler.do_HEAD

    def list_directory(self, path):
        # 简单的中文友好列表页
        try:
            entries = sorted(os.listdir(path))
        except OSError:
            self.send_error(404, "No permission to list directory")
            return None
        rel = os.path.relpath(path, self.DIRECTORY)
        display = "" if rel == "." else "/" + rel.replace(os.sep, "/")
        rows = []
        for name in entries:
            fp = os.path.join(path, name)
            if os.path.isdir(fp):
                href = name + "/"
                tag = "[目录]"
            else:
                href = name
                tag = f"{os.path.getsize(fp):,} B"
            # 转义 HTML
            safe_name = name.replace("&", "&amp;").replace("<", "&lt;").replace(">", "&gt;")
            safe_href = href.replace("&", "&amp;").replace("<", "&lt;").replace(">", "&gt;")
            rows.append(f'<li><a href="{safe_href}">{safe_name}</a> <span style="color:#888">{tag}</span></li>')
        body = (
            "<html><head><meta charset='utf-8'><title>文件室</title></head>"
            f"<body style='font-family:system-ui;max-width:720px;margin:2em auto'>"
            f"<h2>📁 文件室</h2><p style='color:#888'>只读下载服务器</p>"
            f"<p>当前位置：<b>{display or '/'}</b></p><ul>{''.join(rows) or '<li>(空)</li>'}</ul>"
            "</body></html>"
        ).encode("utf-8")
        self.send_response(200)
        self.send_header("Content-Type", "text/html; charset=utf-8")
        self.send_header("Content-Length", str(len(body)))
        self.end_headers()
        return self.wfile.write(body)

    def log_message(self, fmt, *args):
        sys.stderr.write("[file-room] %s - %s\n" % (self.address_string(), fmt % args))


def real_ip_banner(port):
    try:
        s = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
        s.connect(("8.8.8.8", 80))
        ip = s.getsockname()[0]
        s.close()
        return ip
    except Exception:
        return "0.0.0.0"


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--port", type=int, default=1525)
    ap.add_argument("--dir", default=ROOM_DIR)
    args = ap.parse_args()

    RoomHandler.DIRECTORY = os.path.abspath(args.dir)
    os.makedirs(RoomHandler.DIRECTORY, exist_ok=True)

    server = http.server.ThreadingHTTPServer(("0.0.0.0", args.port), RoomHandler)
    ip = real_ip_banner(args.port)
    print(f"[文件室] 启动于 http://{ip}:{args.port}  目录: {RoomHandler.DIRECTORY}")
    try:
        server.serve_forever()
    except KeyboardInterrupt:
        print("\n[文件室] 已停止")


if __name__ == "__main__":
    main()
