# simple_sharer.py
import asyncio
import json
import logging
from pathlib import Path

import websockets
from aiortc import RTCPeerConnection, RTCSessionDescription, RTCConfiguration, RTCIceServer
from aiortc.contrib.media import MediaStreamTrack as AiortcMediaStreamTrack
from av import VideoFrame
from mss import mss
from PIL import Image

try:
    import pyaudio
    PYAUDIO_AVAILABLE = True
except ImportError:
    PYAUDIO_AVAILABLE = False
    print("警告：pyaudio 未安裝，將無法播放遠端音訊。")
    print("請執行 'pip install pyaudio' 來安裝。")

# --- 日誌設定 ---
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(name)s - %(levelname)s - %(message)s')
logger = logging.getLogger("simple_sharer")

CONFIG = {}

class AudioPlayerTrack:
    """一個用於播放接收到的音訊的軌道"""
    kind = "audio"

    def __init__(self):
        if not PYAUDIO_AVAILABLE:
            return
        self.p = pyaudio.PyAudio()
        self.stream = self.p.open(format=pyaudio.paInt16,
                                  channels=1,
                                  rate=48000,
                                  output=True,
                                  frames_per_buffer=960) # 48000Hz / 50fps = 960

    async def recv(self, frame):
        """當收到音訊幀時，播放它"""
        if not PYAUDIO_AVAILABLE:
            return
        try:
            self.stream.write(frame.to_ndarray().tobytes())
        except Exception as e:
            logger.warning(f"播放音訊時發生錯誤: {e}")

    def stop(self):
        if not PYAUDIO_AVAILABLE:
            return
        if hasattr(self, 'stream') and self.stream.is_active():
            self.stream.stop_stream()
            self.stream.close()
        if hasattr(self, 'p'):
            self.p.terminate()

class ScreenShareTrack(AiortcMediaStreamTrack):
    """一個 MediaStreamTrack，用於從整個螢幕擷取影像。"""
    kind = "video"

    def __init__(self):
        super().__init__()
        self.sct = mss()
        self.capture_region = self.sct.monitors[1] # 假設分享主螢幕
        logger.info(f"設定為分享整個主螢幕: {self.capture_region}")

    async def recv(self):
        """從螢幕擷取一幀影像"""
        sct_img = self.sct.grab(self.capture_region)
        img = Image.frombytes("RGB", sct_img.size, sct_img.bgra, "raw", "BGRX")

        frame = VideoFrame.from_image(img)
        pts, time_base = await self.next_timestamp()
        frame.pts = pts
        frame.time_base = time_base
        return frame

class WebRTCSharer:
    def __init__(self, sharer_id, password, server_url, ice_servers):
        self.sharer_id = sharer_id
        self.password = password
        self.server_url = server_url
        self.ice_servers = ice_servers
        self.websocket = None
        self.peer_connections = {}  # {controller_id: RTCPeerConnection}
        self.audio_player = None

    async def connect(self):
        """連線到信令伺服器並註冊"""
        logger.info("正在嘗試連線到信令伺服器 %s", self.server_url)
        self.websocket = await websockets.connect(self.server_url)
        logger.info("成功連線到信令伺服器")

        register_payload = {
            "type": "register_sharer",
            "id": self.sharer_id,
            "password": self.password,
        }
        await self.websocket.send(json.dumps(register_payload))
        logger.info("已註冊為分享端，ID: %s", self.sharer_id)

    async def run(self):
        """主執行迴圈，處理信令訊息"""
        await self.connect()
        try:
            async for message in self.websocket:
                data = json.loads(message)
                msg_type = data.get("type")
                controller_id = data.get("from_id")

                if msg_type == "request_to_connect" and controller_id:
                    await self.create_peer_connection(controller_id)
                elif msg_type == "answer_to_sharer" and controller_id:
                    pc = self.peer_connections.get(controller_id)
                    if pc:
                        answer = RTCSessionDescription(sdp=data["answer"]["sdp"], type=data["answer"]["type"])
                        await pc.setRemoteDescription(answer)
                elif msg_type == "ice_to_sharer" and controller_id:
                    pc = self.peer_connections.get(controller_id)
                    if pc and data.get("ice"):
                        candidate = pc._parse_ice_candidate(data["ice"])
                        await pc.addIceCandidate(candidate)
        except websockets.exceptions.ConnectionClosed as e:
            logger.warning("與信令伺服器的連線已關閉: %s", e)
        finally:
            logger.info("正在關閉所有 WebRTC 連線...")
            for pc in self.peer_connections.values():
                await pc.close()
            if self.audio_player:
                self.audio_player.stop()
            self.peer_connections.clear()

    async def create_peer_connection(self, controller_id):
        """為新的控制器建立一個 Peer Connection"""
        logger.info("收到來自 %s 的連線請求，正在建立 WebRTC 連線...", controller_id)

        ice_servers_obj = [RTCIceServer(**server) for server in self.ice_servers]
        configuration = RTCConfiguration(iceServers=ice_servers_obj)
        pc = RTCPeerConnection(configuration=configuration)
        self.peer_connections[controller_id] = pc

        @pc.on("connectionstatechange")
        async def on_connectionstatechange():
            logger.info("WebRTC 連線狀態 (%s): %s", controller_id, pc.connectionState)
            if pc.connectionState in ("failed", "closed", "disconnected"):
                await self.cleanup_peer_connection(controller_id)

        @pc.on("track")
        def on_track(track):
            logger.info(f"收到軌道: {track.kind}")
            if track.kind == "audio":
                if not self.audio_player:
                    self.audio_player = AudioPlayerTrack()
                asyncio.create_task(self.play_audio_frames(track))

        pc.addTrack(ScreenShareTrack())

        offer = await pc.createOffer()
        await pc.setLocalDescription(offer)

        offer_payload = {
            "type": "offer_to_controller",
            "from_id": self.sharer_id,
            "target_id": controller_id,
            "offer": {"sdp": pc.localDescription.sdp, "type": pc.localDescription.type},
        }
        await self.websocket.send(json.dumps(offer_payload))
        logger.info("已發送 Offer 給 %s", controller_id)

    async def play_audio_frames(self, track):
        """從軌道接收音訊幀並播放"""
        while True:
            try:
                frame = await track.recv()
                await self.audio_player.recv(frame)
            except Exception as e:
                logger.warning(f"處理音訊幀時出錯: {e}")
                break

    async def cleanup_peer_connection(self, controller_id):
        """清理指定的 Peer Connection"""
        pc = self.peer_connections.pop(controller_id, None)
        if pc and pc.connectionState != "closed":
            await pc.close()
            logger.info("已關閉與 %s 的 WebRTC 連線", controller_id)

def load_config():
    """載入設定檔"""
    global CONFIG
    try:
        with open("config.json", 'r') as f:
            CONFIG = json.load(f)
    except (FileNotFoundError, json.JSONDecodeError) as e:
        logger.error(f"錯誤：無法載入 config.json。 {e}")
        exit(1)

def build_connection_configs():
    """根據設定檔動態建立連線 URL 和 ICE 伺服器列表"""
    conn_config = CONFIG.get("connection", {})
    is_secure = conn_config.get("secure", False)
    host = conn_config.get("host", "localhost")
    ws_protocol = "wss" if is_secure else "ws"
    signaling_port = conn_config.get("signaling_port", 6759)
    signaling_url = f"{ws_protocol}://{host}:{signaling_port}"
    turn_creds = CONFIG.get("turn_credentials", {})
    turn_port = conn_config.get("turn_port_secure", 5349) if is_secure else conn_config.get("turn_port_insecure", 3478)
    ice_servers = [
        {"urls": f"stun:{host}:{turn_port}"},
        {
            "urls": f"{turn_protocol}:{host}:{turn_port}",
            "username": turn_creds.get("username"),
            "credential": turn_creds.get("credential")
        }
    ]
    return signaling_url, ice_servers

async def main():
    load_config()
    signaling_url, ice_servers = build_connection_configs()
    sharer_info = CONFIG.get("sharer_info", {})
    sharer = WebRTCSharer(
        sharer_id=sharer_info.get("id", "default-sharer"),
        password=sharer_info.get("password", "password"),
        server_url=signaling_url,
        ice_servers=ice_servers
    )
    await sharer.run()

if __name__ == "__main__":
    try:
        asyncio.run(main())
    except KeyboardInterrupt:
        logger.info("程式已手動中斷。")