"""WebSocket-based command execution for long-running commands."""

from __future__ import annotations

import json
import logging
import time
from collections.abc import AsyncIterator, Iterator, Mapping
from typing import Any, Callable, Optional

from langsmith import utils as ls_utils
from langsmith.sandbox._exceptions import (
    CommandTimeoutError,
    SandboxConnectionError,
    SandboxConnectTimeoutError,
    SandboxNotReadyError,
    SandboxOperationError,
    SandboxServerReloadError,
)
from langsmith.sandbox._helpers import merge_headers

logger = logging.getLogger(__name__)

# Resolve the optional ``websockets`` dependency once, at import time, instead
# of re-importing on every WS call. ``WEBSOCKETS_AVAILABLE`` is what run() reads
# to choose WebSocket vs the blocking HTTP fallback.
try:
    from websockets.asyncio.client import connect as _ws_connect_async
    from websockets.exceptions import ConnectionClosed, InvalidHandshake
    from websockets.sync.client import connect as _ws_connect_sync

    WEBSOCKETS_AVAILABLE = True
except ImportError:
    _ws_connect_async = _ws_connect_sync = None  # type: ignore[assignment,misc]
    ConnectionClosed = InvalidHandshake = None  # type: ignore[assignment,misc]
    WEBSOCKETS_AVAILABLE = False


def _env_timeout(name: str, default: float) -> Optional[float]:
    """Read a WebSocket timeout override, in seconds. ``<= 0`` disables it."""
    raw = ls_utils.get_env_var(name)
    if raw is None:
        return default
    try:
        value = float(raw)
    except ValueError:
        logger.warning("Ignoring invalid LANGSMITH_%s=%r, using %s", name, raw, default)
        return default
    return value if value > 0 else None


# A slow cold start is recovered by the retry loop in run(), not by waiting here.
WS_OPEN_TIMEOUT = _env_timeout("SANDBOX_WS_TIMEOUT_OPEN", 30)
WS_PING_INTERVAL = _env_timeout("SANDBOX_WS_TIMEOUT_PING_INTERVAL", 30)
WS_PING_TIMEOUT = _env_timeout("SANDBOX_WS_TIMEOUT_PING", 60)
# Kept short: a dead peer would otherwise stall teardown for the full duration.
WS_CLOSE_TIMEOUT = _env_timeout("SANDBOX_WS_TIMEOUT_CLOSE", 10)
# Ceiling on the whole connect phase. Without it, retrying a blackholed handshake
# costs MAX_AUTO_RECONNECTS + 1 full open timeouts plus backoff.
WS_CONNECT_BUDGET = _env_timeout("SANDBOX_WS_TIMEOUT_CONNECT_BUDGET", 120)


def connect_deadline() -> Optional[float]:
    """Monotonic instant after which connect attempts must stop, if bounded."""
    if WS_CONNECT_BUDGET is None:
        return None
    return time.monotonic() + WS_CONNECT_BUDGET


def open_timeout_for(deadline: Optional[float]) -> Optional[float]:
    """Per-attempt open timeout, clamped so it cannot outlive ``deadline``."""
    if deadline is None:
        return WS_OPEN_TIMEOUT
    remaining = max(deadline - time.monotonic(), 0.0)
    if WS_OPEN_TIMEOUT is None:
        return remaining
    return min(WS_OPEN_TIMEOUT, remaining)


_MISSING_WEBSOCKETS_MSG = (
    "WebSocket-based execution requires the 'websockets' package, which ships "
    "with langsmith by default. Reinstall with: pip install --upgrade langsmith"
)


def _ensure_websockets():
    """Return the cached sync websockets symbols, or raise if unavailable."""
    if not WEBSOCKETS_AVAILABLE:
        raise ImportError(_MISSING_WEBSOCKETS_MSG)
    return _ws_connect_sync, ConnectionClosed, InvalidHandshake


def _ensure_websockets_async():
    """Return the cached async websockets symbols, or raise if unavailable."""
    if not WEBSOCKETS_AVAILABLE:
        raise ImportError(_MISSING_WEBSOCKETS_MSG)
    return _ws_connect_async, ConnectionClosed, InvalidHandshake


def _build_ws_url(dataplane_url: str) -> str:
    """Convert dataplane HTTP URL to WebSocket URL for /execute/ws."""
    ws_url = dataplane_url.replace("https://", "wss://").replace("http://", "ws://")
    return f"{ws_url}/execute/ws"


def _build_auth_headers(
    api_key: Optional[str], headers: Optional[Mapping[str, str]] = None
) -> dict[str, str]:
    """Build auth headers for the WebSocket upgrade request."""
    auth_headers = {"X-Api-Key": api_key} if api_key else None
    return merge_headers(auth_headers, headers)


# =============================================================================
# Stream Control
# =============================================================================


class _WSStreamControl:
    """Control interface for an active WebSocket stream.

    Created before the generator starts, bound to the WebSocket once
    the connection opens. The CommandHandle holds a reference to this
    object to send kill/input messages.

    Thread safety: websockets' sync client supports send() from one
    thread while recv() runs on another. So kill() from user code
    and iteration on a different thread are safe.
    """

    def __init__(self) -> None:
        self._ws: Any = None
        self._closed = False
        self._killed = False

    def _bind(self, ws: Any) -> None:
        """Bind to the active WebSocket. Called inside the generator."""
        self._ws = ws

    def _unbind(self) -> None:
        """Mark as closed. Called when the generator exits."""
        self._closed = True
        self._ws = None

    @property
    def killed(self) -> bool:
        """True if kill() has been called on this stream."""
        return self._killed

    def send_kill(self) -> None:
        """Send a kill message and immediately close the WebSocket."""
        self._killed = True
        if self._ws and not self._closed:
            try:
                self._ws.send(json.dumps({"type": "kill"}))
            except Exception:
                pass
            try:
                self._ws.close_timeout = 0
                self._ws.close()
            except Exception:
                pass

    def send_input(self, data: str) -> None:
        """Send stdin data to the running command."""
        if self._ws and not self._closed:
            self._ws.send(json.dumps({"type": "input", "data": data}))


class _AsyncWSStreamControl:
    """Async equivalent of _WSStreamControl."""

    def __init__(self) -> None:
        self._ws: Any = None
        self._closed = False
        self._killed = False

    def _bind(self, ws: Any) -> None:
        self._ws = ws

    def _unbind(self) -> None:
        self._closed = True
        self._ws = None

    @property
    def killed(self) -> bool:
        return self._killed

    async def send_kill(self) -> None:
        self._killed = True
        if self._ws and not self._closed:
            try:
                await self._ws.send(json.dumps({"type": "kill"}))
            except Exception:
                pass
            try:
                self._ws.close_timeout = 0
                await self._ws.close()
            except Exception:
                pass

    async def send_input(self, data: str) -> None:
        if self._ws and not self._closed:
            await self._ws.send(json.dumps({"type": "input", "data": data}))


# =============================================================================
# Error Handling
# =============================================================================


def _raise_for_invalid_handshake(exc: Exception, ws_url: str) -> None:
    """Raise a clear error when the WebSocket upgrade handshake fails.

    Covers both a rejection carrying an HTTP status (``InvalidStatus`` — most
    commonly 404 when the server lacks the /execute/ws endpoint) and a response
    that isn't valid HTTP at all (``InvalidMessage``, "did not receive a valid
    HTTP response") — e.g. a stopped or recycled sandbox dataplane answering the
    upgrade with garbage. Both subclass ``InvalidHandshake``.
    """
    status = getattr(getattr(exc, "response", None), "status_code", None)
    if status == 404:
        raise SandboxConnectionError(
            f"The sandbox server does not support WebSocket command execution "
            f"(endpoint {ws_url} returned 404). Ensure the server is updated "
            f"to a version that supports the /execute/ws endpoint, or use "
            f"run() without wait=False or callbacks."
        ) from exc
    if status == 503:
        raise SandboxNotReadyError(
            f"Sandbox is not ready for WebSocket command execution: {exc}"
        ) from exc
    if status is not None:
        raise SandboxConnectionError(
            f"WebSocket upgrade rejected by server (HTTP {status}): {exc}"
        ) from exc
    # No HTTP status at all — the peer didn't return a valid HTTP response,
    # typically a stopped/unreachable sandbox dataplane.
    raise SandboxConnectionError(
        f"WebSocket upgrade to {ws_url} failed (no valid HTTP response); "
        f"the sandbox may be stopped or unreachable: {exc}"
    ) from exc


def _raise_from_error_msg(msg: dict, *, command_id: str = "") -> None:
    """Raise the appropriate exception from a server error message."""
    error_type = msg.get("error_type", "CommandError")
    error_msg = msg.get("error", "Unknown error")

    if error_type == "CommandTimeout":
        raise CommandTimeoutError(error_msg)
    if error_type == "CommandNotFound":
        raise SandboxOperationError(
            f"Command not found: {command_id}" if command_id else error_msg,
            operation="reconnect" if command_id else "command",
            error_type=error_type,
        )
    if error_type == "SessionExpired":
        raise SandboxOperationError(
            f"Session expired: {command_id}" if command_id else error_msg,
            operation="reconnect" if command_id else "command",
            error_type=error_type,
        )

    raise SandboxOperationError(
        error_msg,
        operation="reconnect" if command_id else "command",
        error_type=error_type,
    )


# =============================================================================
# Sync Stream Functions
# =============================================================================


def run_ws_stream(
    dataplane_url: str,
    api_key: Optional[str],
    command: str,
    *,
    command_id: str = "",
    timeout: int = 60,
    env: Optional[dict[str, str]] = None,
    cwd: Optional[str] = None,
    shell: str = "/bin/bash",
    on_stdout: Optional[Callable[[str], Any]] = None,
    on_stderr: Optional[Callable[[str], Any]] = None,
    idle_timeout: int = 300,
    kill_on_disconnect: bool = False,
    ttl_seconds: int = 600,
    pty: bool = False,
    headers: Optional[Mapping[str, str]] = None,
    open_timeout: Optional[float] = WS_OPEN_TIMEOUT,
) -> tuple[Iterator[dict], _WSStreamControl]:
    """Execute a command over WebSocket, yielding raw message dicts.

    Returns a tuple of (message_iterator, control). The control object
    provides send_kill() and send_input() methods for the CommandHandle.

    The iterator yields dicts with a "type" field:
    - {"type": "started", "command_id": "...", "pid": N}
    - {"type": "stdout", "data": "...", "offset": N}
    - {"type": "stderr", "data": "...", "offset": N}
    - {"type": "exit", "exit_code": N}

    If on_stdout/on_stderr callbacks are provided, they are invoked as
    data arrives in addition to yielding the messages.
    """
    ws_connect, ConnectionClosed, InvalidHandshake = _ensure_websockets()
    ws_url = _build_ws_url(dataplane_url)
    request_headers = _build_auth_headers(api_key, headers)
    control = _WSStreamControl()

    def _stream() -> Iterator[dict]:
        try:
            with ws_connect(
                ws_url,
                additional_headers=request_headers,
                open_timeout=open_timeout,
                close_timeout=WS_CLOSE_TIMEOUT,
                ping_interval=WS_PING_INTERVAL,
                ping_timeout=WS_PING_TIMEOUT,
            ) as ws:
                control._bind(ws)

                # Send execute request
                payload: dict[str, Any] = {
                    "type": "execute",
                    "command": command,
                    "timeout_seconds": timeout,
                    "shell": shell,
                    "idle_timeout_seconds": idle_timeout,
                    "kill_on_disconnect": kill_on_disconnect,
                    "ttl_seconds": ttl_seconds,
                }
                if command_id:
                    payload["command_id"] = command_id
                if env:
                    payload["env"] = env
                if cwd:
                    payload["cwd"] = cwd
                if pty:
                    payload["pty"] = True
                ws.send(json.dumps(payload))

                # Read messages until exit or error
                for raw_msg in ws:
                    msg = json.loads(raw_msg)
                    msg_type = msg.get("type")

                    if msg_type == "started":
                        yield msg

                    elif msg_type == "stdout":
                        if on_stdout:
                            on_stdout(msg["data"])
                        yield msg

                    elif msg_type == "stderr":
                        if on_stderr:
                            on_stderr(msg["data"])
                        yield msg

                    elif msg_type == "exit":
                        yield msg
                        return

                    elif msg_type == "error":
                        _raise_from_error_msg(msg)

        except InvalidHandshake as e:
            _raise_for_invalid_handshake(e, ws_url)
        except ConnectionClosed as e:
            if e.rcvd and e.rcvd.code == 1001:
                raise SandboxServerReloadError(
                    "Server is reloading, reconnect to resume"
                ) from e
            raise SandboxConnectionError(
                f"WebSocket connection closed unexpectedly: {e}"
            ) from e
        except (OSError, RuntimeError) as e:
            raise SandboxConnectTimeoutError(
                f"Failed to connect to sandbox: {e}"
            ) from e
        finally:
            control._unbind()

    return _stream(), control


def reconnect_ws_stream(
    dataplane_url: str,
    api_key: Optional[str],
    command_id: str,
    *,
    stdout_offset: int = 0,
    stderr_offset: int = 0,
    headers: Optional[Mapping[str, str]] = None,
) -> tuple[Iterator[dict], _WSStreamControl]:
    """Reconnect to an existing command over WebSocket.

    Returns a tuple of (message_iterator, control), same as run_ws_stream.
    The iterator yields the server's ``started`` acknowledgement, then stdout,
    stderr, exit, and error messages. That acknowledgement is the only evidence
    a reattachment landed for a command producing no output, so it is forwarded
    rather than dropped.

    With the ring buffer reader server model, there is no replay/live
    phase distinction and no deduplication needed. The server reads from
    its ring buffer starting at the requested offsets and streams output
    from there. If the requested offset is older than the buffer's
    earliest data, the server sends from the earliest available offset.
    """
    ws_connect, ConnectionClosed, InvalidHandshake = _ensure_websockets()
    ws_url = _build_ws_url(dataplane_url)
    request_headers = _build_auth_headers(api_key, headers)
    control = _WSStreamControl()

    def _stream() -> Iterator[dict]:
        try:
            with ws_connect(
                ws_url,
                additional_headers=request_headers,
                open_timeout=WS_OPEN_TIMEOUT,
                close_timeout=WS_CLOSE_TIMEOUT,
                ping_interval=WS_PING_INTERVAL,
                ping_timeout=WS_PING_TIMEOUT,
            ) as ws:
                control._bind(ws)

                # Send reconnect request
                ws.send(
                    json.dumps(
                        {
                            "type": "reconnect",
                            "command_id": command_id,
                            "stdout_offset": stdout_offset,
                            "stderr_offset": stderr_offset,
                        }
                    )
                )

                # Read messages until exit or error
                for raw_msg in ws:
                    msg = json.loads(raw_msg)
                    msg_type = msg.get("type")

                    if msg_type in ("started", "stdout", "stderr"):
                        yield msg

                    elif msg_type == "exit":
                        yield msg
                        return

                    elif msg_type == "error":
                        _raise_from_error_msg(msg, command_id=command_id)

        except InvalidHandshake as e:
            _raise_for_invalid_handshake(e, ws_url)
        except ConnectionClosed as e:
            if e.rcvd and e.rcvd.code == 1001:
                raise SandboxServerReloadError(
                    "Server is reloading, reconnect to resume"
                ) from e
            raise SandboxConnectionError(
                f"WebSocket connection closed unexpectedly: {e}"
            ) from e
        except (OSError, RuntimeError) as e:
            raise SandboxConnectTimeoutError(
                f"Failed to connect to sandbox: {e}"
            ) from e
        finally:
            control._unbind()

    return _stream(), control


# =============================================================================
# Async Stream Functions
# =============================================================================


async def run_ws_stream_async(
    dataplane_url: str,
    api_key: Optional[str],
    command: str,
    *,
    command_id: str = "",
    timeout: int = 60,
    env: Optional[dict[str, str]] = None,
    cwd: Optional[str] = None,
    shell: str = "/bin/bash",
    on_stdout: Optional[Callable[[str], Any]] = None,
    on_stderr: Optional[Callable[[str], Any]] = None,
    idle_timeout: int = 300,
    kill_on_disconnect: bool = False,
    ttl_seconds: int = 600,
    pty: bool = False,
    headers: Optional[Mapping[str, str]] = None,
    open_timeout: Optional[float] = WS_OPEN_TIMEOUT,
) -> tuple[AsyncIterator[dict], _AsyncWSStreamControl]:
    """Async equivalent of run_ws_stream.

    Returns (async_message_iterator, async_control).
    """
    ws_connect_async, ConnectionClosed, InvalidHandshake = _ensure_websockets_async()
    ws_url = _build_ws_url(dataplane_url)
    request_headers = _build_auth_headers(api_key, headers)
    control = _AsyncWSStreamControl()

    async def _stream() -> AsyncIterator[dict]:
        try:
            async with ws_connect_async(
                ws_url,
                additional_headers=request_headers,
                open_timeout=open_timeout,
                close_timeout=WS_CLOSE_TIMEOUT,
                ping_interval=WS_PING_INTERVAL,
                ping_timeout=WS_PING_TIMEOUT,
            ) as ws:
                control._bind(ws)

                payload: dict[str, Any] = {
                    "type": "execute",
                    "command": command,
                    "timeout_seconds": timeout,
                    "shell": shell,
                    "idle_timeout_seconds": idle_timeout,
                    "kill_on_disconnect": kill_on_disconnect,
                    "ttl_seconds": ttl_seconds,
                }
                if command_id:
                    payload["command_id"] = command_id
                if env:
                    payload["env"] = env
                if cwd:
                    payload["cwd"] = cwd
                if pty:
                    payload["pty"] = True
                await ws.send(json.dumps(payload))

                async for raw_msg in ws:
                    msg = json.loads(raw_msg)
                    msg_type = msg.get("type")

                    if msg_type == "started":
                        yield msg
                    elif msg_type == "stdout":
                        if on_stdout:
                            on_stdout(msg["data"])
                        yield msg
                    elif msg_type == "stderr":
                        if on_stderr:
                            on_stderr(msg["data"])
                        yield msg
                    elif msg_type == "exit":
                        yield msg
                        return
                    elif msg_type == "error":
                        _raise_from_error_msg(msg)

        except InvalidHandshake as e:
            _raise_for_invalid_handshake(e, ws_url)
        except ConnectionClosed as e:
            if e.rcvd and e.rcvd.code == 1001:
                raise SandboxServerReloadError(
                    "Server is reloading, reconnect to resume"
                ) from e
            raise SandboxConnectionError(
                f"WebSocket connection closed unexpectedly: {e}"
            ) from e
        except (OSError, RuntimeError) as e:
            raise SandboxConnectTimeoutError(
                f"Failed to connect to sandbox: {e}"
            ) from e
        finally:
            control._unbind()

    return _stream(), control


async def reconnect_ws_stream_async(
    dataplane_url: str,
    api_key: Optional[str],
    command_id: str,
    *,
    stdout_offset: int = 0,
    stderr_offset: int = 0,
    headers: Optional[Mapping[str, str]] = None,
) -> tuple[AsyncIterator[dict], _AsyncWSStreamControl]:
    """Async equivalent of reconnect_ws_stream."""
    ws_connect_async, ConnectionClosed, InvalidHandshake = _ensure_websockets_async()
    ws_url = _build_ws_url(dataplane_url)
    request_headers = _build_auth_headers(api_key, headers)
    control = _AsyncWSStreamControl()

    async def _stream() -> AsyncIterator[dict]:
        try:
            async with ws_connect_async(
                ws_url,
                additional_headers=request_headers,
                open_timeout=WS_OPEN_TIMEOUT,
                close_timeout=WS_CLOSE_TIMEOUT,
                ping_interval=WS_PING_INTERVAL,
                ping_timeout=WS_PING_TIMEOUT,
            ) as ws:
                control._bind(ws)

                await ws.send(
                    json.dumps(
                        {
                            "type": "reconnect",
                            "command_id": command_id,
                            "stdout_offset": stdout_offset,
                            "stderr_offset": stderr_offset,
                        }
                    )
                )

                async for raw_msg in ws:
                    msg = json.loads(raw_msg)
                    msg_type = msg.get("type")

                    if msg_type in ("started", "stdout", "stderr"):
                        yield msg
                    elif msg_type == "exit":
                        yield msg
                        return
                    elif msg_type == "error":
                        _raise_from_error_msg(msg, command_id=command_id)

        except InvalidHandshake as e:
            _raise_for_invalid_handshake(e, ws_url)
        except ConnectionClosed as e:
            if e.rcvd and e.rcvd.code == 1001:
                raise SandboxServerReloadError(
                    "Server is reloading, reconnect to resume"
                ) from e
            raise SandboxConnectionError(
                f"WebSocket connection closed unexpectedly: {e}"
            ) from e
        except (OSError, RuntimeError) as e:
            raise SandboxConnectTimeoutError(
                f"Failed to connect to sandbox: {e}"
            ) from e
        finally:
            control._unbind()

    return _stream(), control
