# -*- coding: utf-8 -*-
"""
App Inventor 网络微数据库（TinyWebDB）自建服务端
App Inventor 2 中文网 www.fun123.cn

只用 Python 自带的库，不需要 pip 安装任何东西。
运行：python3 tinywebdb_server.py      （Windows 双击 start-windows.bat）

除了网络微数据库的 storeavalue / getvalue，还提供图片上传：
  POST 服务地址/upload   请求体就是图片文件（Web 客户端的 PostFile），返回 {"status":"OK","url":图片地址}
  GET  服务地址/files/文件名   取回图片
"""
import html
import json
import os
import re
import secrets
import socket
import sqlite3
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from urllib.parse import parse_qs

PORT = int(os.environ.get("TINYWEBDB_PORT", "5000"))   # 端口，被占用时改这里
# 数据目录默认是程序所在文件夹；云服务器上由 tinywebdb.service 指定 /var/lib/tinywebdb
HERE = os.environ.get("TINYWEBDB_DATA") or os.path.dirname(os.path.abspath(__file__))
DB_FILE = os.path.join(HERE, "tinywebdb.sqlite3")      # 数据都存在这个文件里
SECRET_FILE = os.path.join(HERE, "secret.txt")         # 访问口令，第一次运行自动生成
FILES_DIR = os.path.join(HERE, "files")               # 上传的图片
MAX_UPLOAD = 5 * 1024 * 1024                           # 单张图片最大 5MB
MAX_DRAIN = 64 * 1024 * 1024                           # 超大请求最多读这么多再拒绝，超过就直接断开
IMAGE_TYPES = {"jpg": "image/jpeg", "png": "image/png", "gif": "image/gif", "webp": "image/webp"}


def load_secret():
    if not os.path.exists(SECRET_FILE):
        with open(SECRET_FILE, "w", encoding="utf-8") as f:
            f.write(secrets.token_hex(4))
    with open(SECRET_FILE, encoding="utf-8") as f:
        return f.read().strip()


SECRET = load_secret()


def db():
    conn = sqlite3.connect(DB_FILE, timeout=10)
    conn.execute("CREATE TABLE IF NOT EXISTS store (tag TEXT PRIMARY KEY, value TEXT)")
    return conn


def image_type(data):
    # 按文件头判断图片格式，不是图片的一律拒收
    if data[:3] == b"\xff\xd8\xff":
        return "jpg"
    if data[:8] == b"\x89PNG\r\n\x1a\n":
        return "png"
    if data[:4] == b"GIF8":
        return "gif"
    if data[:4] == b"RIFF" and data[8:12] == b"WEBP":
        return "webp"
    return None


def lan_ip():
    # 找出本机在局域网里的地址（不会真的发数据）
    s = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
    try:
        s.connect(("10.255.255.255", 1))
        return s.getsockname()[0]
    except OSError:
        return "127.0.0.1"
    finally:
        s.close()


class Handler(BaseHTTPRequestHandler):
    def parts(self):
        # "/口令//getvalue" 和 "/口令/getvalue" 都按 ["口令", "getvalue"] 处理
        return [p for p in self.path.split("?")[0].split("/") if p]

    def reply(self, code, body, content_type="application/json; charset=utf-8"):
        data = body.encode("utf-8")
        self.send_response(code)
        self.send_header("Content-Type", content_type)
        self.send_header("Content-Length", str(len(data)))
        self.end_headers()
        self.wfile.write(data)

    def read_body(self, limit):
        """读完整个请求体；超过 limit 时照样读完（不保存），返回 None。
        先读完再回复很重要：手机还在发送时服务端就回复并断开，手机收到的是“连接被重置”，看不到错误原因。
        Web 客户端的 PostFile 用分块传输，不带 Content-Length，两种都要能读。"""
        body, total = bytearray(), 0
        if "chunked" not in self.headers.get("Transfer-Encoding", "").lower():
            left = int(self.headers.get("Content-Length") or 0)
            while left > 0 and total <= MAX_DRAIN:
                part = self.rfile.read(min(left, 65536))
                if not part:
                    break
                left -= len(part)
                total += len(part)
                if total <= limit:
                    body += part
        else:
            while total <= MAX_DRAIN:
                size = int(self.rfile.readline().split(b";")[0].strip() or b"0", 16)
                if size == 0:
                    while self.rfile.readline().strip():
                        pass
                    break
                part = self.rfile.read(size)
                self.rfile.readline()
                total += len(part)
                if total <= limit:
                    body += part
        if total > MAX_DRAIN:
            self.close_connection = True
        return bytes(body) if total <= limit else None

    def upload(self):
        data = self.read_body(MAX_UPLOAD)
        if data is None:
            return self.reply(413, json.dumps({"status": "ERROR", "message": "图片超过 5MB"}, ensure_ascii=False))
        ext = image_type(data)
        if not ext:
            return self.reply(400, json.dumps({"status": "ERROR", "message": "只支持 jpg/png/gif/webp 图片"}, ensure_ascii=False))
        os.makedirs(FILES_DIR, exist_ok=True)
        name = "%s.%s" % (secrets.token_hex(8), ext)
        with open(os.path.join(FILES_DIR, name), "wb") as f:
            f.write(data)
        scheme = self.headers.get("X-Forwarded-Proto", "http")
        url = "%s://%s/%s/files/%s" % (scheme, self.headers.get("Host", "127.0.0.1:%d" % PORT), SECRET, name)
        return self.reply(200, json.dumps({"status": "OK", "url": url, "name": name}))

    def do_POST(self):
        p = self.parts()
        if len(p) != 2 or p[0] != SECRET:
            self.read_body(0)
            return self.reply(404, '["ERROR", "not found"]')
        if p[1] == "upload":
            return self.upload()
        form = parse_qs((self.read_body(MAX_UPLOAD) or b"").decode("utf-8"), keep_blank_values=True)
        tag = form.get("tag", [""])[0]
        if p[1] == "storeavalue":
            value = form.get("value", [""])[0]   # 组件发来的是 JSON 文本，原样保存
            with db() as conn:
                conn.execute("INSERT OR REPLACE INTO store (tag, value) VALUES (?, ?)", (tag, value))
            return self.reply(200, json.dumps(["STORED", tag, value], ensure_ascii=False))
        if p[1] == "getvalue":
            with db() as conn:
                row = conn.execute("SELECT value FROM store WHERE tag = ?", (tag,)).fetchone()
            # 第三项放回存进来的 JSON 文本；没有这个标签时给空字符串
            return self.reply(200, json.dumps(["VALUE", tag, row[0] if row else ""], ensure_ascii=False))
        return self.reply(404, '["ERROR", "unknown command"]')

    def do_GET(self):
        p = self.parts()
        if len(p) == 3 and p[0] == SECRET and p[1] == "files" and re.fullmatch(r"[0-9a-f]{16}\.(jpg|png|gif|webp)", p[2]):
            path = os.path.join(FILES_DIR, p[2])
            if os.path.isfile(path):
                with open(path, "rb") as f:
                    data = f.read()
                self.send_response(200)
                self.send_header("Content-Type", IMAGE_TYPES[p[2].rsplit(".", 1)[1]])
                self.send_header("Content-Length", str(len(data)))
                self.send_header("Cache-Control", "max-age=31536000")
                self.end_headers()
                self.wfile.write(data)
                return
        if len(p) != 1 or p[0] != SECRET:
            return self.reply(404, "Not Found", "text/plain; charset=utf-8")
        with db() as conn:
            rows = conn.execute("SELECT tag, value FROM store ORDER BY tag").fetchall()
        images = len(os.listdir(FILES_DIR)) if os.path.isdir(FILES_DIR) else 0
        trs = "".join("<tr><td>%s</td><td>%s</td></tr>" % (html.escape(t), html.escape(v)) for t, v in rows)
        page = ("<!doctype html><meta charset='utf-8'><title>网络微数据库</title>"
                "<h2>网络微数据库：共 %d 条，图片 %d 张</h2><table border='1' cellpadding='6'>"
                "<tr><th>标签</th><th>值（JSON 文本）</th></tr>%s</table>") % (len(rows), images, trs)
        return self.reply(200, page, "text/html; charset=utf-8")

    def log_message(self, fmt, *args):
        print("%s  %s" % (self.address_string(), fmt % args))


class Server(ThreadingHTTPServer):
    request_queue_size = 128   # 默认只有 5，全班同时请求会被拒绝连接
    daemon_threads = True


if __name__ == "__main__":
    db().close()
    server = Server(("0.0.0.0", PORT), Handler)
    url = "http://%s:%d/%s" % (lan_ip(), PORT, SECRET)
    print("=" * 60)
    print("网络微数据库服务已启动，关闭这个窗口服务就会停止。")
    print("在 App 里把“服务地址”设为：")
    print("    " + url)
    print("用浏览器打开上面的地址可以查看所有数据。")
    print("在云服务器上运行时，把地址里的 IP 换成服务器的公网 IP。")
    print("=" * 60, flush=True)
    try:
        server.serve_forever()
    except KeyboardInterrupt:
        pass
