# server_webrtc.py
import asyncio
import json
import logging
import ssl

import websockets
from websockets.protocol import State
import base64 # 為了處理檔案上傳

# 嘗試匯入 pynput，如果失敗則無法進行遠端控制
try:
    from pynput.mouse import Button, Controller as MouseController
    from pynput.keyboard import Key, Controller as KeyboardController
    PYNPUT_AVAILABLE = True
except ImportError:
    PYNPUT_AVAILABLE = False

# --- 日誌設定 ---
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
logger = logging.getLogger("server_webrtc")

# --- 伺服器狀態 ---
# 儲存所有連線的客戶端 (包括分享端和控制端)
# clients = { "client_id": websocket_connection }
clients = {}

# 儲存分享端的憑證資訊
# sharer_credentials = { "sharer_id": "password" }
sharer_credentials = {}

# 儲存連線配對資訊
# sessions = { "controller_id": "sharer_id" }
sessions = {}

if PYNPUT_AVAILABLE:
    mouse = MouseController()
    keyboard = KeyboardController()


async def forward_message(websocket, message_data):
    """
    將訊息轉發給指定的目標客戶端。
    """
    target_id = message_data.get("target_id")
    if not target_id:
        logger.warning("訊息缺少 'target_id'，無法轉發: %s", message_data)
        return

    target_ws = clients.get(target_id)
    if target_ws and target_ws.state == State.OPEN:
        # 為了不讓伺服器知道太多細節，我們直接轉發整個 JSON 字串
        await target_ws.send(json.dumps(message_data))
        logger.info("已將訊息從 %s 轉發到 %s", message_data.get('from_id', '未知'), target_id)
    else:
        logger.warning("找不到目標客戶端 '%s' 或連線已關閉，無法轉發。", target_id)

def execute_remote_command(command):
    """
    解析並執行來自控制端的遠端控制指令。
    """
    if not PYNPUT_AVAILABLE:
        logger.warning("pynput 未安裝，無法執行遠端控制指令。")
        return

    # 建立一個對應表，將前端傳來的 code 對應到 pynput 的 Key 屬性
    # 這解決了 'shiftleft' vs 'shift_l' 的問題
    key_code_map = {
        'shiftleft': 'shift_l',
        'shiftright': 'shift_r',
        'controlleft': 'ctrl_l',
        'controlright': 'ctrl_r',
        'altleft': 'alt_l',
        'altright': 'alt_r',
        'metaleft': 'cmd_l', # for macOS
        'metaright': 'cmd_r', # for macOS
        # 其他需要轉換的鍵可以加在這裡
    }

    cmd_type = command.get("type")
    try:
        if cmd_type == "mouse_move":
            mouse.position = (int(command['x']), int(command['y']))
        elif cmd_type == "mouse_down":
            button = Button.left if command['button'] == 'left' else Button.right
            mouse.press(button)
        elif cmd_type == "mouse_up":
            button = Button.left if command['button'] == 'left' else Button.right
            mouse.release(button)
        elif cmd_type == "key_down":
            key_str = command['key']
            # 處理 pynput 的特殊鍵格式，例如 'Key.enter'
            if key_str.startswith("Key."):
                code = key_str.split('.')[1]
                # 使用對應表來取得正確的 pynput 屬性名稱
                key_attr_name = key_code_map.get(code, code)
                key = getattr(Key, key_attr_name)
                keyboard.press(key)
            else:
                keyboard.press(key_str)
        elif cmd_type == "key_up":
            key_str = command['key']
            if key_str.startswith("Key."):
                code = key_str.split('.')[1]
                key_attr_name = key_code_map.get(code, code)
                key = getattr(Key, key_attr_name)
                keyboard.release(key)
            else:
                keyboard.release(key_str)
        elif cmd_type == "file_upload":
            # 處理檔案上傳
            filename = command.get("filename")
            content_b64 = command.get("content")
            if filename and content_b64:
                save_path = Path.home() / "Downloads" / filename
                save_path.write_bytes(base64.b64decode(content_b64))
                logger.info(f"檔案 '{filename}' 已儲存到下載資料夾。")
        else:
            logger.warning("未知的遠端控制指令類型: %s", cmd_type)
    except Exception as e:
        logger.error("執行遠端控制指令 '%s' 時發生錯誤: %s", cmd_type, e)

async def handler(websocket, path=None):
    """
    處理每個 WebSocket 連線。
    """
    client_id = None
    logger.info("新連線來自: %s", websocket.remote_address)
    try:
        # 等待客戶端的初始訊息 (註冊或請求連線)
        async for message in websocket:
            try:
                data = json.loads(message)
                msg_type = data.get("type")

                if msg_type == "register_sharer":
                    # 分享端註冊
                    client_id = data.get("id")
                    password = data.get("password")
                    if not client_id or not password:
                        logger.warning("分享端註冊失敗：缺少 ID 或密碼。")
                        continue

                    # 處理 ID 衝突：踢掉舊的連線
                    if client_id in clients:
                        logger.warning("ID '%s' 衝突，正在斷開舊連線。", client_id)
                        old_ws = clients[client_id]
                        await old_ws.close(1000, "New connection with same ID")

                    clients[client_id] = websocket
                    sharer_credentials[client_id] = password
                    logger.info("分享端 '%s' 已註冊。", client_id)

                elif msg_type == "register_controller":
                    # 控制端註冊 (Web Page)
                    client_id = data.get("id")
                    if not client_id:
                        logger.warning("控制端註冊失敗：缺少 ID。")
                        continue
                    clients[client_id] = websocket
                    logger.info("控制端 '%s' 已註冊。", client_id)

                elif msg_type == "register_dashboard":
                    # 儀表板註冊
                    client_id = data.get("id")
                    if not client_id:
                        logger.warning("儀表板註冊失敗：缺少 ID。")
                        continue
                    clients[client_id] = websocket
                    logger.info("儀表板 '%s' 已註冊。", client_id)

                elif msg_type == "register_broadcast_controller":
                    # 新增：廣播控制器註冊
                    client_id = data.get("id")
                    if not client_id:
                        logger.warning("廣播控制器註冊失敗：缺少 ID。")
                        continue
                    clients[client_id] = websocket
                    logger.info("廣播控制器 '%s' 已註冊。", client_id)

                elif msg_type in ("request_to_connect", "offer_to_controller", "answer_to_sharer", "ice_to_controller", "ice_to_sharer", "offer_to_sharer", "answer_to_controller", "restart_sharer"):
                    # 如果是請求連線，記錄配對關係
                    if msg_type == "request_to_connect":
                        # --- 新增：密碼驗證邏輯 ---
                        sharer_id = data.get("target_id")
                        controller_id = data.get("from_id")
                        password = data.get("password")

                        # *** 關鍵修正：當控制器請求連線時，也將其註冊到 clients 字典中 ***
                        if controller_id and controller_id not in clients:
                            client_id = controller_id # 將當前連線的 client_id 設為控制器 ID
                            clients[client_id] = websocket
                            logger.info("控制器 '%s' 已透過請求連線進行註冊。", client_id)
                        # *** 修正結束 ***

                        # 檢查分享端是否存在且密碼是否正確
                        if sharer_id in sharer_credentials and sharer_credentials[sharer_id] == password:
                            logger.info("控制器 '%s' 密碼驗證成功，準備連線到分享端 '%s'。", controller_id, sharer_id)
                            if sharer_id and controller_id:
                                sessions[controller_id] = sharer_id
                                logger.info("建立工作階段: 控制端 %s -> 分享端 %s", controller_id, sharer_id)
                            # 驗證成功，轉發請求
                            await forward_message(websocket, data)
                        else:
                            logger.warning("控制器 '%s' 嘗試連線到 '%s' 失敗：密碼錯誤或分享端不存在。", controller_id, sharer_id)
                            # 可以選擇性地回傳一個錯誤訊息給控制器
                            error_msg = {"type": "connection_failed", "reason": "Invalid credentials or sharer not found."}
                            await websocket.send(json.dumps(error_msg))
                            continue # 停止後續處理

                    # 所有需要轉發的信令訊息
                    # --- 修正：將 elif 改為 if，確保 request_to_connect 驗證後也能觸發轉發 ---
                    if msg_type != "request_to_connect":
                        await forward_message(websocket, data)


                elif msg_type == "request_sharer_list":
                    # 儀表板請求分享者列表
                    sharer_ids = list(sharer_credentials.keys())
                    response = {"type": "sharer_list", "sharers": sharer_ids}
                    await websocket.send(json.dumps(response))

                elif msg_type == "request_screenshot":
                    # 儀表板請求截圖，轉發給分享端
                    await forward_message(websocket, data)

                elif msg_type == "screenshot_response":
                    # 分享端回傳截圖，轉發給儀表板
                    await forward_message(websocket, data)

                elif msg_type == "broadcast_command":
                    # 新增：處理廣播指令
                    # --- 關鍵優化：改為並行廣播，提高效率與即時性 ---
                    command = data.get("command")
                    if command:
                        logger.info("收到來自 '%s' 的廣播指令: %s", data.get("from_id"), command)
                        command_payload = json.dumps({"type": "remote_control", "command": command})
                        
                        # 1. 建立所有發送任務
                        tasks = []
                        # 直接遍歷 sharer_credentials，更精準
                        for sharer_id in sharer_credentials:
                            sharer_ws = clients.get(sharer_id)
                            if sharer_ws and sharer_ws.state == State.OPEN:
                                tasks.append(sharer_ws.send(command_payload))
                        
                        # 2. 並行執行所有發送任務
                        await asyncio.gather(*tasks)
                        logger.info("廣播指令已並行轉發給 %d 個接收端。", len(tasks))

                elif msg_type == "remote_control":
                    # 執行來自 Web 分享端的遠端控制指令
                    execute_remote_command(data.get("command"))

                elif msg_type == "log_forward":
                    # --- 新增：轉發日誌訊息 ---
                    sharer_id = data.get("from_id")
                    # 找到所有正在連線此分享端的控制器
                    controllers_to_notify = [cid for cid, sid in sessions.items() if sid == sharer_id]
                    for cid in controllers_to_notify:
                        controller_ws = clients.get(cid)
                        if controller_ws and controller_ws.state == State.OPEN:
                            await controller_ws.send(json.dumps(data))

                elif msg_type == "update_stats":
                    # 這是來自客戶端的狀態更新，我們可以在這裡處理或記錄
                    # 目前我們只記錄日誌，未來可以轉發給儀表板
                    if client_id:
                        cpu = data.get("cpu")
                        memory = data.get("memory")
                        ping = data.get("ping")
                        logger.info(f"收到來自 {client_id} 的狀態更新: CPU {cpu}%, Mem {memory}%, Ping {ping}ms")

                else:
                    logger.warning("收到未知的訊息類型: %s", msg_type)

            except json.JSONDecodeError:
                logger.error("收到非 JSON 格式的訊息: %s", message)
            except Exception as e:
                logger.error("處理訊息時發生錯誤: %s", e, exc_info=True)

    except websockets.exceptions.ConnectionClosed as e:
        logger.info("連線已關閉: %s (Code: %s, Reason: %s)", websocket.remote_address, e.code, e.reason)
    finally:
        # --- 清理工作 ---
        if client_id and client_id in clients and clients[client_id] == websocket:
            logger.info("客戶端 '%s' 已斷線，正在進行清理...", client_id)
            clients.pop(client_id, None)

            # 如果斷線的是分享端，從憑證中移除
            if client_id in sharer_credentials:
                sharer_credentials.pop(client_id, None)
                logger.info("分享端 '%s' 已斷線，正在通知所有相關的控制器...", client_id)
                # --- 關鍵修正：安全地移除 sessions ---
                # 1. 找到所有正在連線此分享端的控制器
                controllers_to_notify = [cid for cid, sid in sessions.items() if sid == client_id]
                # 2. 通知它們
                for cid in controllers_to_notify:
                    controller_ws = clients.get(cid)
                    if controller_ws and controller_ws.state == State.OPEN:
                        disconnect_msg = {"type": "sharer_disconnected", "from_id": client_id}
                        await controller_ws.send(json.dumps(disconnect_msg))
                # 3. 在迭代結束後，安全地從 sessions 中移除
                for cid in controllers_to_notify:
                    sessions.pop(cid, None)

            # 如果斷線的是控制器 (或儀表板)，通知對應的分享端
            elif client_id in sessions:
                sharer_id = sessions[client_id]
                sessions.pop(client_id, None)
                sharer_ws = clients.get(sharer_id)
                # 在發送前再次確認分享端連線是否仍處於開啟狀態
                if sharer_ws and sharer_ws.state == State.OPEN:
                    logger.info("正在通知分享端 '%s'，控制器 '%s' 已斷線。", sharer_id, client_id)
                    disconnect_msg = {"type": "controller_disconnected", "from_id": client_id}
                    await sharer_ws.send(json.dumps(disconnect_msg))
            else:
                # 如果斷線的是儀表板或其他未在 session 中的客戶端，則無需特別通知
                logger.info("客戶端 '%s' (可能為儀表板) 已斷線，無需額外通知。", client_id)

        logger.info("目前在線客戶端數量: %d", len(clients))


async def main():
    # --- 新增：設定 SSL/TLS ---
    # 替換成您自己的 SSL 憑證和私鑰檔案路徑
    # 這些憑證必須對應您的域名 (例如 www.winway.tw)
    if not PYNPUT_AVAILABLE:
        logger.error("="*50)
        logger.error("錯誤：pynput 函式庫未安裝，遠端控制功能將無法使用。")
        logger.error("請執行 'pip install pynput' 來安裝。")
        logger.error("="*50)
    ssl_context = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER)
    ssl_context.load_cert_chain('/etc/letsencrypt/live/www.winway.tw/fullchain.pem', 
                                '/etc/letsencrypt/live/www.winway.tw/privkey.pem')

    # 監聽所有網路介面的 6759 port
    port = 6759
    # 將 ssl_context 傳遞給 serve 函式
    async with websockets.serve(handler, "0.0.0.0", port, ssl=ssl_context, max_size=2**22): # 增加限制到 4MB
        logger.info(f"WebRTC 安全信令伺服器已啟動於 wss://0.0.0.0:{port}")
        await asyncio.Future()  # 保持伺服器永久運行


if __name__ == "__main__":
    try:
        asyncio.run(main())
    except KeyboardInterrupt:
        logger.info("伺服器已手動關閉。")