#!/usr/bin/env python3

import base64
import time
from datetime import datetime, timezone
from pathlib import Path
from threading import Thread

import requests
from flask import Flask, Response, request
from typing_extensions import TYPE_CHECKING, Dict, Optional

if TYPE_CHECKING:
    from _typeshed import GenericPath
else:
    GenericPath = str

import bytedance.jeddak_secure_channel as jsc
from bytedance.jeddak_secure_channel import Client, ClientConfig

secure_channel_config: Optional[ClientConfig] = None
secure_channel_client: Optional[Client] = None

server_addr = "http://127.0.0.1:8080"


def normal_encrypt_and_decrypt(msg: str, loop_count: int = 1000) -> bool:
    """
    测试本地加解密的功能、性能和稳定性.
    """
    assert secure_channel_client is not None
    ret = secure_channel_client.attest_server()
    assert ret

    i = 0
    while i < loop_count:
        encrypted_msg, enc_key = secure_channel_client.encrypt_with_response(msg)
        plaintext = enc_key.decrypt(encrypted_msg)
        if plaintext.decode() != msg:
            return False
        i += 1
    return True


def test_request_1(msg: str, loop_count: int = 10) -> bool:
    """
    test_request_1：
        1、对请求参数中的部分敏感参数进行加密
        2、对请求应答的数据中的部分敏感信息进行加密
        3、接口安全鉴权采用token的方式
    """
    url = f"{server_addr}/request_example_1"

    assert secure_channel_config is not None
    assert secure_channel_client is not None

    # 对数据接收方发起远程证明
    ret = secure_channel_client.attest_server()
    assert ret

    for _ in range(loop_count):
        # 加密数据并发送
        params: Dict[str, str] = {}
        params["key_1"] = "value_1"  # 非敏感字段

        encrypted_msg, enc_key = secure_channel_client.encrypt_with_response(msg)
        params["enc_msg"] = encrypted_msg  # 敏感字段

        print("test_request_1 发送请求:", params)
        response = requests.post(url, json=params)
        if response:
            response_json = response.json()
            if response_json.get("status"):
                # 解密响应
                enc_msg = base64.b64decode(response_json["msg"])
                print("test_request_1 收到加密应答:", enc_msg)

                plaintext = enc_key.decrypt(enc_msg)
                print("test_request_1 解密应答:", plaintext)

    return True


def test_request_2(test_msg: str):
    """
    test_request_2：
        1、对请求体进行整体加密
        2、对请求应答的数据进行整体加密
        3、接口安全鉴权采用token的方式
    """
    url = f"{server_addr}/request_example_2"

    assert secure_channel_config is not None
    assert secure_channel_client is not None

    # 对数据接收方发起远程证明
    secure_channel_client.attest_server()

    # 加密数据并发送
    encrypted_msg, enc_key = secure_channel_client.encrypt_with_response(test_msg)

    response = requests.post(url, data=encrypted_msg)
    print("test_request_2 发送请求:", encrypted_msg)

    # 解密响应
    response = enc_key.decrypt(response.text).decode()
    print("test_request_2 收到响应:", response)


def test_request_3(file_name: GenericPath):
    """
    test_request_3：
        1、请求体中带有数据文件，数据文件是加密的
        2、接口安全鉴权采用token的方式
    """
    url = f"{server_addr}/request_example_3"

    assert secure_channel_config is not None
    assert secure_channel_client is not None

    # 对数据接收方发起远程证明
    secure_channel_client.attest_server()

    # 对文件进行加密
    enc_file_name = Path(f"{file_name}.enc")
    enc_key = secure_channel_client.encrypt_file(file_name, enc_file_name, "b")

    # 加密文件并发送请求
    headers = {
        "X-File-Enc-Key": enc_key,
    }
    # enc_key 里的信息是经过 RSA 加密的，不需要担心泄密

    with open(enc_file_name, "rb") as enc_file:
        encrypted = enc_file.read()

    print("test_request_3 发送请求:", headers)
    response = requests.post(url, headers=headers, data=encrypted)
    if response:
        response_text = response.text
        print("test_request_3 收到应答内容:", response_text)


def test_request_4(msg: str, loop_count: int = 10):
    """
    test_request_4：
        1、测试流式加密
        2、测试双向RA
    """
    url = f"{server_addr}/request_example_4"

    assert secure_channel_config is not None
    assert secure_channel_client is not None

    # 对数据接收方发起远程证明
    secure_channel_client.attest_server()

    for _ in range(loop_count):
        # 发送请求到服务端
        encrypted_msg, enc_key = secure_channel_client.encrypt_with_response(msg)
        params = {"query": encrypted_msg}

        with requests.post(url=url, json=params, stream=True) as r:
            for line in r.iter_lines():
                if line:
                    res = line.decode().strip()
                    enc_msg = base64.b64decode(res)
                    plaintext = enc_key.decrypt(enc_msg)
                    print("test_request_4 解密应答:", plaintext)


def test_concurrent(msg: str, thread_count: int = 10):
    """
    测试SDK的线程安全性
    使用多线程操作同一个 SDK client, 并发发送请求.
    """
    url = f"{server_addr}/request_example_2"

    assert secure_channel_config is not None
    assert secure_channel_client is not None

    # 对数据接收方发起远程证明
    secure_channel_client.attest_server()

    def run(i: int):
        assert secure_channel_client is not None

        # 加密数据并发送
        encrypted_msg, enc_key = secure_channel_client.encrypt_with_response(msg)

        response = requests.post(url, data=encrypted_msg)
        print(f"Worker {i} 发送请求")

        # 解密响应
        response = enc_key.decrypt(response.text)
        print(f"Worker {i} 收到响应:", response.decode())

    threads = [Thread(target=run, args=(i,)) for i in range(thread_count)]

    for t in threads:
        t.start()

    for t in threads:
        t.join()


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


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


@app.route("/normal_test", methods=["GET"])
def normal_test():
    # 基本加解密功能测试
    params = request.args

    msg = params.get("msg", "Hello, World!")
    loop = int(params.get("loop", "1000"))

    start_time = time.time()
    normal_encrypt_and_decrypt(msg=msg, loop_count=loop)
    end_time = time.time()
    spand_time = f"加密内容为:{msg}, 加密次数:{loop}, 用时: {end_time - start_time}秒\n"

    return Response(response=spand_time, status=200, content_type="application/json")


@app.route("/request_test", methods=["GET"])
def request_test():
    # 对请求的部分参数进行加密测试
    params = request.args

    msg = params.get("msg", "Hello, World!")
    loop = int(params.get("loop", "10"))

    start_time = time.time()
    test_request_1(msg=msg, loop_count=loop)
    end_time = time.time()
    spand_time = f"加密内容为:{msg}, 加密次数:{loop}, 用时: {end_time - start_time}秒\n"

    return Response(response=spand_time, status=200, content_type="application/json")


@app.route("/stream_request_test", methods=["GET"])
def stream_request_test():
    # 流式请求加密测试
    params = request.args

    msg = params.get("msg", "Hello, World!")
    loop = int(params.get("loop", "10"))

    start_time = time.time()
    test_request_4(msg=msg, loop_count=loop)
    end_time = time.time()
    spand_time = f"加密内容为:{msg}, 加密次数:{loop}, 用时: {end_time - start_time}秒\n"

    return Response(response=spand_time, status=200, content_type="application/json")


@app.route("/file_request_test", methods=["GET"])
def file_request_test():
    # 对请求体中的文件加密
    params = request.args
    loop = int(params.get("loop", "10"))

    test_file = Path(__file__).parent / "test_file"
    with open(test_file, "w") as f:
        f.write("a1,a2\na11,a12\na21,a22")

    start_time = time.time()
    for _ in range(loop):
        test_request_3(test_file)
    end_time = time.time()
    spand_time = f"用时: {end_time - start_time}秒\n"

    return Response(response=spand_time, status=200, content_type="application/json")


@app.route("/concurrent_request_test", methods=["GET"])
def concurrent_request_test():
    # 客户端多并发功能测试
    params = request.args

    msg = params.get("msg", "Hello, World!")
    thread_count = int(params.get("thread", "10"))

    start_time = time.time()
    test_concurrent(msg=msg, thread_count=thread_count)
    end_time = time.time()
    spand_time = (
        f"加密内容为:{msg}, 并发个数:{thread_count}, 用时: {end_time - start_time}秒\n"
    )

    return Response(response=spand_time, status=200, content_type="application/json")


@app.route("/sign_test", methods=["GET"])
def sign_test():
    # 测试签名与验签
    url = f"{server_addr}/sign_example_1"

    start_time = time.time()

    # ret = secure_channel_client.attest_server()
    # assert ret

    params = request.args
    msg = params.get("msg", "Hello, World!")
    app_info = params.get("app_info", "test_id_test_name")

    # 加密数据并发送
    params: Dict[str, str] = {}
    params["enc_msg"] = secure_channel_client.encrypt(msg)
    params["app_info"] = app_info

    timestamp = str(int(datetime.now(timezone.utc).timestamp()))
    sign = secure_channel_client.gen_sign(app_info, timestamp)
    header = {"Timestamp": timestamp, "Sign": sign}

    response = requests.post(url, json=params, headers=header)
    if response:
        response_json = response.json()
        print(f"response={response_json}")

    end_time = time.time()
    spand_time = f"加密内容为:{msg}, 用时: {end_time - start_time}秒\n"

    return Response(response=spand_time, status=200, content_type="application/json")


if __name__ == "__main__":
    import argparse

    parser = argparse.ArgumentParser(
        "example_client", description="安全通信示例 (数据发送方)."
    )
    parser.add_argument(
        "message", nargs="?", default="Hello, World!", help="要发送的数据."
    )
    args = parser.parse_args()

    # 创建数据发送方对象
    config_file = Path(__file__).parent / "client_config.json"
    secure_channel_config = jsc.ClientConfig.from_file(config_file)
    secure_channel_client = jsc.Client(secure_channel_config)

    # normal_encrypt_and_decrypt("test")

    # test_request_1(args.message)
    #
    # test_request_2(args.message)
    #
    # test_file = Path(__file__).parent / "test_file"
    # with open(test_file, "w") as f:
    #     f.write("a1,a2\na11,a12\na21,a22")
    # test_request_3(test_file)
    #
    # test_concurrent(args.message)
    #
    # test_request_4(args.message, loop_count=1)

    app.run(host="0.0.0.0", port=7070)
