#!/usr/bin/env python3

import base64
import functools
import json
from pathlib import Path

from flask import Flask, Response, abort, request
from werkzeug.exceptions import HTTPException

import bytedance.jeddak_secure_channel as jsc

# 创建数据接收方对象
conf_file = Path(__file__).parent / "server_config.json"
secure_channel_config = jsc.ServerConfig.from_file(conf_file)
secure_channel_server = jsc.Server(secure_channel_config)


# 创建 HTTP 服务器
app = Flask(__name__)


def bearer_auth(f):
    """验证 Basic Auth."""

    @functools.wraps(f)
    def decorated(*args, **kwargs):
        # print(f"收到请求: path={request.path}, params={request.get_json()}")
        # print(
        #     f"收到请求: path={request.path}, headers={request.headers}, params={request.get_json()}"
        # )
        token = request.headers.get("Authorization")
        if not token:
            abort(401)

        if token != "Bearer My_Secret_Token":
            abort(401)

        return f(*args, **kwargs)

    return decorated


@app.errorhandler(Exception)
def handle_exception(e):
    if isinstance(e, HTTPException):
        return e

    print(e)
    return str(e), 500


@app.route("/ping")
def ping():
    return "Hello Jeddak Secure Channel Server"


@app.route("/ra", methods=["POST"])
# @bearer_auth
def ra() -> jsc.RaResponse:
    """远程证明接口."""

    data: jsc.RaRequest = request.get_json()
    res = secure_channel_server.handle_ra_request(data)
    print(f"ra response={res}")
    return res


@app.route("/upload", methods=["POST"])
# @bearer_auth
def upload():
    """数据上传接口."""
    return Response(response=None, status=200)


@app.route("/request_example_1", methods=["POST"])
# @bearer_auth
def request_example_1():
    # 获取请求参数
    params = request.get_json()
    print("收到请求参数:", params)

    # 解密参数中的密文数据
    encrypt_msg = params.get("enc_msg")
    plaintext, encrypt_key = secure_channel_server.decrypt_with_response(encrypt_msg)
    plaintext = plaintext.decode()
    print("收到明文数据:", plaintext)

    # 数据处理
    res_data = "res:" + plaintext.upper()

    res_msg = encrypt_key.encrypt(res_data).encode()
    # print(res_msg)

    # 应答且加密应答内容中的机密信息
    res = dict()
    res["status"] = True
    res["msg"] = base64.b64encode(res_msg).decode()

    print("发送响应:", res)
    return Response(
        response=json.dumps(res), status=200, content_type="application/json"
    )


@app.route("/request_example_2", methods=["POST"])
def request_example_2():
    data: str = request.get_data(as_text=True)

    # 解密请求
    decrypted_msg, enc_key = secure_channel_server.decrypt_with_response(data)
    decrypted_msg = decrypted_msg.decode()
    print("收到请求:", decrypted_msg)

    response = decrypted_msg.upper()

    # 加密响应
    print("发送响应:", response)
    return enc_key.encrypt(response)


@app.route("/request_example_3", methods=["POST"])
def request_example_3():
    file_enc_key = request.headers["X-File-Enc-Key"]  # 文件加密的密钥信息

    # Save file
    encrypted_path = "./src_encrypted_file"
    with open(encrypted_path, "wb") as encrypted_file:
        encrypted_file.write(request.get_data())

    # 文件解密
    output_path = "./des_plaintext_file"
    secure_channel_server.decrypt_file(file_enc_key, encrypted_path, output_path, "b")

    res = {"msg": "ok"}
    return Response(
        response=json.dumps(res), status=200, content_type="application/json"
    )


@app.route("/request_example_4", methods=["POST"])
def request_example_4():
    import time

    # 获取请求参数
    params = request.get_json()
    print("收到请求参数:", params)

    # # 对客户端做RA
    # if not secure_channel_server.attest_client():
    #     res["status"] = False
    #     res["msg"] = "client ra failed"
    #     return Response(response=json.dumps(res), status=200, content_type="application/json")

    # 解密参数中的密文数据
    encrypt_msg = params.get("query")
    plaintext, encrypt_key = secure_channel_server.decrypt_with_response(encrypt_msg)
    plaintext = plaintext.decode()
    print("收到明文数据:", plaintext)

    # 数据处理
    res_data = "res:" + plaintext.upper()

    # 发送响应
    def generate_blocks():
        for i in range(5):
            res_msg = encrypt_key.encrypt(res_data).encode()
            res_msg = base64.b64encode(res_msg).decode()
            yield f"{res_msg}\n"
            time.sleep(1)

    return Response(generate_blocks(), status=200, mimetype="text/plain")


@app.route("/sign_example_1", methods=["POST"])
# @bearer_auth
def sign_example_1():
    # 获取请求参数
    res = dict()

    headers = request.headers
    params = request.get_json()
    print(
        f"收到请求参数: headers={headers}, params={params}",
    )

    # 签名验证
    app_info = params.get("app_info")
    timestamp = headers.get("Timestamp")
    sign = headers.get("Sign")
    if not app_info or not sign or not timestamp:
        res["status"] = False
        return Response(
            response=json.dumps(res), status=200, content_type="application/json"
        )

    if not secure_channel_server.verify_sign(app_info, timestamp, sign):
        res["status"] = False
        res["msg"] = "verify_sign failed"
        print(res["msg"])
        return Response(
            response=json.dumps(res), status=200, content_type="application/json"
        )

    print("verify_sign success")

    # 解密参数中的密文数据
    encrypt_msg = params.get("enc_msg")
    plaintext, encrypt_key = secure_channel_server.decrypt_with_response(encrypt_msg)
    plaintext = plaintext.decode()
    print("收到明文数据:", plaintext)

    # 应答且加密应答内容中的机密信息
    res["status"] = True
    res["msg"] = "success"

    return Response(
        response=json.dumps(res), status=200, content_type="application/json"
    )


def main():
    import argparse

    parser = argparse.ArgumentParser(
        "example_server", description="安全通信示例 (数据接收方)."
    )
    parser.add_argument("--host", default="0.0.0.0")
    parser.add_argument("-p", "--port", type=int, default=8080)
    parser.add_argument("--debug", action="store_true")

    args = parser.parse_args()

    app.run(host=args.host, port=args.port, debug=args.debug)


if __name__ == "__main__":
    main()
