import os
import sys
import json
import time
import threading
import base64
import requests
import openai
from flask import Flask, request, jsonify
from flask_cors import CORS
from Crypto.PublicKey import RSA
from Crypto.Cipher import PKCS1_OAEP, AES
from Crypto.Protocol.KDF import PBKDF2

# ==================== 配置 ====================
SERVER_URL = "http://47.95.124.80:2369"
CLIENT_ID = "client_1"
SECRET = "secret_123"
# =============================================

# ==================== 核心类 ====================
class SecureOpenAIClient:
    def __init__(self, server_url, client_id, secret, data_dir="./data"):
        self.server_url = server_url.rstrip("/")
        self.client_id = client_id
        self.secret = secret
        self.data_dir = data_dir
        os.makedirs(data_dir, exist_ok=True)

        self.private_key_enc_path = os.path.join(data_dir, f"{client_id}_private.pem.enc")
        self.public_key_path = os.path.join(data_dir, f"{client_id}_public.pem")
        self.config_cache_path = os.path.join(data_dir, f"{client_id}_config.enc")

        self.private_key = None
        self.public_key_pem = None
        self.config = None
        self._connected = False
        self._total_tokens = 0
        self._openai_client = None

        self._init_keys()
        self._register_public_key()
        self._load_config()
        self._start_heartbeat()
        self._update_my_usage()

    def _derive_aes_key(self, salt=b"fixed_salt_2026"):
        return PBKDF2(f"{self.client_id}:{self.secret}".encode(), salt, dkLen=32, count=100000)

    def _aes_encrypt(self, plaintext, key):
        cipher = AES.new(key, AES.MODE_GCM)
        ct, tag = cipher.encrypt_and_digest(plaintext)
        return cipher.nonce + tag + ct

    def _aes_decrypt(self, ciphertext, key):
        nonce, tag, ct = ciphertext[:16], ciphertext[16:32], ciphertext[32:]
        cipher = AES.new(key, AES.MODE_GCM, nonce=nonce)
        return cipher.decrypt_and_verify(ct, tag)

    def _init_keys(self):
        aes_key = self._derive_aes_key()
        if os.path.exists(self.private_key_enc_path) and os.path.exists(self.public_key_path):
            with open(self.private_key_enc_path, "rb") as f:
                enc = f.read()
            try:
                pem = self._aes_decrypt(enc, aes_key)
                self.private_key = RSA.import_key(pem)
            except:
                self._generate_new_keys(aes_key)
        else:
            self._generate_new_keys(aes_key)
        with open(self.public_key_path, "r") as f:
            self.public_key_pem = f.read()

    def _generate_new_keys(self, aes_key):
        key = RSA.generate(2048)
        self.private_key = key
        pem = key.export_key()
        enc = self._aes_encrypt(pem, aes_key)
        with open(self.private_key_enc_path, "wb") as f:
            f.write(enc)
        with open(self.public_key_path, "w") as f:
            f.write(key.publickey().export_key().decode())
        self.public_key_pem = key.publickey().export_key().decode()

    def _request(self, endpoint, params=None, json_data=None, method="GET"):
        url = f"{self.server_url}{endpoint}"
        if params is None:
            params = {}
        params.update({"client_id": self.client_id, "secret": self.secret})
        try:
            if method.upper() == "POST":
                resp = requests.post(url, params=params, json=json_data, timeout=5)
            else:
                resp = requests.get(url, params=params, timeout=5)
            resp.raise_for_status()
            return resp.json()
        except requests.RequestException as e:
            raise ConnectionError(f"Service unreachable: {e}")

    def _register_public_key(self):
        try:
            self._request("/api/register", method="POST", json_data={
                "client_id": self.client_id,
                "secret": self.secret,
                "public_key": self.public_key_pem
            })
        except ConnectionError as e:
            raise RuntimeError(f"Registration failed: {e}")

    def _load_config(self):
        if os.path.exists(self.config_cache_path):
            with open(self.config_cache_path, "rb") as f:
                enc = f.read()
            try:
                cipher = PKCS1_OAEP.new(self.private_key)
                dec = cipher.decrypt(enc)
                self.config = json.loads(dec.decode('utf-8'))
                if "api_key" in self.config and "model" in self.config:
                    self._apply_config()
                    self._connected = True
                    return
            except:
                pass
        self._fetch_config_from_server()

    def _fetch_config_from_server(self):
        data = self._request("/api/config")
        if "error" in data:
            raise RuntimeError(f"Config error: {data['error']}")
        enc = base64.b64decode(data["encrypted_config"])
        cipher = PKCS1_OAEP.new(self.private_key)
        dec = cipher.decrypt(enc)
        self.config = json.loads(dec.decode('utf-8'))
        with open(self.config_cache_path, "wb") as f:
            f.write(enc)
        self._apply_config()
        self._connected = True

    def _apply_config(self):
        self._openai_client = openai.OpenAI(
            api_key=self.config["api_key"],
            base_url=self.config["base_url"]
        )
        print(f"[{self.client_id}] Config loaded: model={self.config['model']}")

    def _ping(self):
        try:
            self._request("/api/ping")
            self._connected = True
            return True
        except ConnectionError:
            self._connected = False
            return False

    def _heartbeat_loop(self):
        while True:
            try:
                self._request("/api/heartbeat", method="POST", json_data={})
            except:
                pass
            time.sleep(30)

    def _start_heartbeat(self):
        t = threading.Thread(target=self._heartbeat_loop, daemon=True)
        t.start()

    def _update_my_usage(self):
        try:
            data = self._request("/api/my_usage")
            self._total_tokens = data.get("total_tokens", 0)
        except:
            pass

    def _report_usage(self, usage):
        self._request("/api/report", method="POST", json_data={"usage": usage})
        self._total_tokens += usage.get("total_tokens", 0)

    def check_connection(self):
        return self._ping()

    def get_my_usage(self):
        self._update_my_usage()
        return self._total_tokens

    def show_status(self):
        status = "✅ Connected" if self._connected else "❌ Disconnected"
        print(f"[{self.client_id}] Status: {status}")
        print(f"[{self.client_id}] Total tokens used: {self.get_my_usage()}")

    def chat_completion(self, messages, tools=None, tool_choice="auto", **kwargs):
        if not self._ping():
            raise RuntimeError("Service unavailable - client cannot work without server")
        model = self.config.get("model", "gpt-3.5-turbo")
        params = {"model": model, "messages": messages, **kwargs}
        if tools:
            params["tools"] = tools
        if tool_choice:
            params["tool_choice"] = tool_choice
        try:
            resp = self._openai_client.chat.completions.create(**params)
        except Exception as e:
            raise e
        if hasattr(resp, "usage") and resp.usage:
            usage = {
                "prompt_tokens": resp.usage.prompt_tokens,
                "completion_tokens": resp.usage.completion_tokens,
                "total_tokens": resp.usage.total_tokens
            }
            self._report_usage(usage)
        return resp

# ==================== Flask API ====================
app = Flask(__name__)
CORS(app)

print("正在连接主控服务端...")
client = SecureOpenAIClient(SERVER_URL, CLIENT_ID, SECRET)
client.show_status()
if not client._connected:
    print("❌ 无法连接服务端，API拒绝启动")
    sys.exit(1)

@app.route("/v1/models", methods=["GET"])
def get_models():
    return jsonify({
        "object": "list",
        "data": [{"id": client.config.get("model"), "object": "model"}]
    })

@app.route("/v1/chat/completions", methods=["POST"])
def chat():
    data = request.get_json()
    if not data or "messages" not in data:
        return jsonify({"error": "Missing messages"}), 400
    try:
        resp = client.chat_completion(
            messages=data["messages"],
            tools=data.get("tools"),
            tool_choice=data.get("tool_choice", "auto"),
            temperature=data.get("temperature", 0.7),
            max_tokens=data.get("max_tokens", 1000)
        )
        return jsonify(resp.model_dump())
    except Exception as e:
        return jsonify({"error": str(e)}), 500

@app.route("/health", methods=["GET"])
def health():
    if client.check_connection():
        return jsonify({"status": "ok", "total_tokens": client.get_my_usage()})
    return jsonify({"status": "unhealthy"}), 503

@app.route("/stats", methods=["GET"])
def stats():
    return jsonify({
        "client_id": CLIENT_ID,
        "connected": client.check_connection(),
        "total_tokens": client.get_my_usage()
    })

if __name__ == "__main__":
    port = int(os.getenv("CLIENT_PORT", 2369))
    print(f"✅ AI API 网关已启动: http://0.0.0.0:{port}")
    app.run(host="0.0.0.0", port=port, debug=False)