Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
from ..utils import aio, log_exceptions, shortuuid
from . import channel, proto
from .inference_proc_lazy_main import ProcStartArgs, proc_main
from .stdio_capture import ChildStdio
from .supervised_proc import SupervisedProc, SupervisedProcKind


Expand Down Expand Up @@ -51,10 +52,13 @@ def __init__(
def process_kind(self) -> SupervisedProcKind:
return SupervisedProcKind.INFERENCE

def _create_process(self, cch: socket.socket, log_cch: socket.socket) -> mp.Process:
def _create_process(
self, cch: socket.socket, log_cch: socket.socket, stdio: ChildStdio | None
) -> mp.Process:
proc_args = ProcStartArgs(
log_cch=log_cch,
mp_cch=cch,
stdio=stdio,
runners=self._runners,
)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -27,18 +27,23 @@
from . import proto
from .channel import Message
from .proc_client import _dump_stack_traces_impl, _ProcClient
from .stdio_capture import ChildStdio, redirect_stdio


@dataclass
class ProcStartArgs:
log_cch: socket.socket
mp_cch: socket.socket
runners: _RunnersDict
stdio: ChildStdio | None = None


def proc_main(args: ProcStartArgs) -> None:
from .proc_client import _ProcClient

if args.stdio is not None:
redirect_stdio(args.stdio)

inf_proc = _InferenceProc(args.runners)

client = _ProcClient(args.mp_cch, args.log_cch, inf_proc.initialize, inf_proc.entrypoint)
Expand Down
6 changes: 5 additions & 1 deletion livekit-agents/livekit/agents/ipc/job_proc_executor.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
from .inference_executor import InferenceExecutor
from .job_executor import JobStatus
from .job_proc_lazy_main import ProcStartArgs, proc_main
from .stdio_capture import ChildStdio
from .supervised_proc import SupervisedProc, SupervisedProcKind


Expand Down Expand Up @@ -92,7 +93,9 @@ def user_arguments(self, value: Any | None) -> None:
def running_job(self) -> RunningJobInfo | None:
return self._running_job

def _create_process(self, cch: socket.socket, log_cch: socket.socket) -> mp.Process:
def _create_process(
self, cch: socket.socket, log_cch: socket.socket, stdio: ChildStdio | None
) -> mp.Process:
levels = {}
root = logging.getLogger()
levels["root"] = root.level
Expand All @@ -109,6 +112,7 @@ def _create_process(self, cch: socket.socket, log_cch: socket.socket) -> mp.Proc
session_end_timeout=self._session_end_timeout,
log_cch=log_cch,
mp_cch=cch,
stdio=stdio,
user_arguments=self._user_args,
logger_levels=levels,
)
Expand Down
5 changes: 5 additions & 0 deletions livekit-agents/livekit/agents/ipc/job_proc_lazy_main.py
Original file line number Diff line number Diff line change
Expand Up @@ -45,6 +45,7 @@
ShuttingDown,
StartJobRequest,
)
from .stdio_capture import ChildStdio, redirect_stdio

# Defensive timeout for AgentSession.aclose() during job shutdown. Hardcoded for now
# as a guardrail against close paths that hang indefinitely. If aclose() does not
Expand All @@ -65,11 +66,15 @@ class ProcStartArgs:
log_cch: socket.socket
logger_levels: dict[str, int]
simulation_end_fnc: Callable[[Any], Any] | None = None
stdio: ChildStdio | None = None


def proc_main(args: ProcStartArgs) -> None:
import logging

if args.stdio is not None:
redirect_stdio(args.stdio)

from .log_queue import LogQueueHandler
from .proc_client import _ProcClient

Expand Down
177 changes: 177 additions & 0 deletions livekit-agents/livekit/agents/ipc/stdio_capture.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,177 @@
from __future__ import annotations

import asyncio
import contextlib
import logging
import os
import socket
import struct
import sys
import time
from collections.abc import Callable
from dataclasses import dataclass
from typing import Any, Literal

logger = logging.getLogger("livekit.agents.stdio")

_SO_TIMESTAMPNS: int = getattr(socket, "SO_TIMESTAMPNS", 35)
_SO_PASSCRED: int = getattr(socket, "SO_PASSCRED", 16)
_SCM_CREDENTIALS: int = getattr(socket, "SCM_CREDENTIALS", 2)
_MAX_LINE_BYTES = 64 * 1024
_CLOSE_TIMEOUT = 2.0

Stream = Literal["stdout", "stderr"]


def capture_enabled() -> bool:
flag = os.environ.get("LIVEKIT_CAPTURE_JOB_STDIO")
if flag is not None:
return flag.strip().lower() in ("1", "true", "yes")
return bool(os.environ.get("LIVEKIT_AGENT_ID"))


@dataclass
class ChildStdio:
stdout: socket.socket
stderr: socket.socket

def close(self) -> None:
for s in (self.stdout, self.stderr):
with contextlib.suppress(OSError):
s.close()


def create_stdio_pairs() -> tuple[ChildStdio, ChildStdio]:
out_parent, out_child = socket.socketpair()
err_parent, err_child = socket.socketpair()
for s in (out_parent, err_parent):
with contextlib.suppress(OSError):
s.setsockopt(socket.SOL_SOCKET, _SO_PASSCRED, 1)
with contextlib.suppress(OSError):
s.setsockopt(socket.SOL_SOCKET, _SO_TIMESTAMPNS, 1)
return ChildStdio(out_parent, err_parent), ChildStdio(out_child, err_child)


def redirect_stdio(stdio: ChildStdio) -> None:
for stream, sock, fd in ((sys.stdout, stdio.stdout, 1), (sys.stderr, stdio.stderr, 2)):
with contextlib.suppress(Exception):
stream.flush()
os.dup2(sock.fileno(), fd)
sock.close()
reconfigure = getattr(stream, "reconfigure", None)
if reconfigure is not None:
with contextlib.suppress(Exception):
reconfigure(line_buffering=True)


def _parse_cmsgs(anc: list[tuple[int, int, bytes]]) -> tuple[int | None, float | None]:
pid: int | None = None
ts: float | None = None
for level, typ, payload in anc:
if level != socket.SOL_SOCKET:
continue
if typ == _SCM_CREDENTIALS and len(payload) >= 12:
cred_pid = struct.unpack("iii", payload[:12])[0]
if cred_pid > 0:
pid = cred_pid
elif typ == _SO_TIMESTAMPNS and len(payload) >= 16:
sec, nsec = struct.unpack("ll", payload[:16])
if sec > 0:
ts = sec + nsec / 1e9
return pid, ts


class StdioReader:
def __init__(
self,
sock: socket.socket,
stream: Stream,
extra_fnc: Callable[[], dict[str, Any]],
loop: asyncio.AbstractEventLoop,
) -> None:
self._sock = sock
self._stream: Stream = stream
self._extra_fnc = extra_fnc
self._loop = loop
self._buf = bytearray()
self._buf_pid: int | None = None
self._buf_ts: float | None = None
self._closed_fut: asyncio.Future[None] = loop.create_future()
self._cmsg_space = socket.CMSG_SPACE(16) + socket.CMSG_SPACE(12)

def start(self) -> None:
self._sock.setblocking(False)
self._loop.add_reader(self._sock.fileno(), self._on_readable)

async def aclose(self) -> None:
try:
await asyncio.wait_for(asyncio.shield(self._closed_fut), timeout=_CLOSE_TIMEOUT)
except asyncio.TimeoutError:
self._finish()

def _on_readable(self) -> None:
while True:
try:
data, anc, _, _ = self._sock.recvmsg(65536, self._cmsg_space)
except (BlockingIOError, InterruptedError):
return
except OSError:
self._finish()
return
if not data:
self._finish()
return
pid, ts = _parse_cmsgs(anc)
self._feed(data, pid, ts if ts is not None else time.time())

def _feed(self, data: bytes, pid: int | None, ts: float) -> None:
start = 0
while True:
nl = data.find(b"\n", start)
if nl == -1:
break
chunk = data[start:nl]
if self._buf:
self._buf.extend(chunk)
self._flush_buf()
else:
self._emit(chunk, pid, ts)
start = nl + 1

rest = data[start:]
if rest:
if not self._buf:
self._buf_pid, self._buf_ts = pid, ts
self._buf.extend(rest)
if len(self._buf) >= _MAX_LINE_BYTES:
self._flush_buf()

def _flush_buf(self) -> None:
if not self._buf:
return
self._emit(bytes(self._buf), self._buf_pid, self._buf_ts or time.time())
self._buf.clear()
self._buf_pid = self._buf_ts = None

def _emit(self, raw: bytes, pid: int | None, ts: float) -> None:
text = raw.decode("utf-8", errors="replace").rstrip("\r")
if not text.strip() or not logger.isEnabledFor(logging.INFO):
return
extra = dict(self._extra_fnc())
extra["stream"] = self._stream
if pid is not None:
extra["pid"] = pid
record = logger.makeRecord(logger.name, logging.INFO, "", 0, text, (), None, extra=extra)
record.created = ts
record.msecs = (ts - int(ts)) * 1000.0
logger.handle(record)

def _finish(self) -> None:
if self._closed_fut.done():
return
with contextlib.suppress(Exception):
self._loop.remove_reader(self._sock.fileno())
self._flush_buf()
with contextlib.suppress(OSError):
self._sock.close()
self._closed_fut.set_result(None)
34 changes: 30 additions & 4 deletions livekit-agents/livekit/agents/ipc/supervised_proc.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,7 @@
from ..utils.aio import duplex_unix
from . import channel, proto
from .log_queue import LogQueueListener
from .stdio_capture import ChildStdio, StdioReader, capture_enabled, create_stdio_pairs

_mask_ctrl_c_refcount = 0
_mask_ctrl_c_original: Callable[[int, FrameType | None], Any] | int | None = signal.SIG_DFL
Expand Down Expand Up @@ -137,9 +138,12 @@ def __init__(
self._lock = asyncio.Lock()
self._shutdown_ack_fut = asyncio.Future[None]()
self._shutting_down_fut = asyncio.Future[None]()
self._stdio_readers: list[StdioReader] = []

@abstractmethod
def _create_process(self, cch: socket.socket, log_cch: socket.socket) -> mp.Process: ...
def _create_process(
self, cch: socket.socket, log_cch: socket.socket, stdio: ChildStdio | None
) -> mp.Process: ...

@abstractmethod
async def _main_task(self, ipc_ch: aio.ChanReceiver[channel.Message]) -> None: ...
Expand Down Expand Up @@ -196,8 +200,19 @@ def _add_proc_ctx_log(record: logging.LogRecord) -> None:
async with self._lock:
mp_pch, mp_cch = socket.socketpair()
mp_log_pch, mp_log_cch = socket.socketpair()

sockets = (mp_pch, mp_cch, mp_log_pch, mp_log_cch)
stdio_parent: ChildStdio | None = None
stdio_child: ChildStdio | None = None
if capture_enabled():
stdio_parent, stdio_child = create_stdio_pairs()

sockets: tuple[socket.socket, ...] = (mp_pch, mp_cch, mp_log_pch, mp_log_cch)
if stdio_parent is not None and stdio_child is not None:
sockets += (
stdio_parent.stdout,
stdio_parent.stderr,
stdio_child.stdout,
stdio_child.stderr,
)
pch: duplex_unix._AsyncDuplex | None = None
log_listener: LogQueueListener | None = None
try:
Expand All @@ -208,7 +223,7 @@ def _add_proc_ctx_log(record: logging.LogRecord) -> None:
log_listener = LogQueueListener(log_pch, _add_proc_ctx_log)
log_listener.start()

self._proc = self._create_process(mp_cch, mp_log_cch)
self._proc = self._create_process(mp_cch, mp_log_cch, stdio_child)

# Set SIG_IGN process-wide before forking so the child inherits it
# (SIG_IGN is preserved across exec per POSIX). This prevents
Expand All @@ -232,6 +247,15 @@ def _add_proc_ctx_log(record: logging.LogRecord) -> None:

mp_log_cch.close()
mp_cch.close()
if stdio_child is not None:
stdio_child.close()
if stdio_parent is not None:
self._stdio_readers = [
StdioReader(stdio_parent.stdout, "stdout", self.logging_extra, self._loop),
StdioReader(stdio_parent.stderr, "stderr", self.logging_extra, self._loop),
]
for reader in self._stdio_readers:
reader.start()

self._pid = self._proc.pid
self._spawn_time = time.monotonic()
Expand Down Expand Up @@ -459,6 +483,8 @@ def _on_read_ipc_done(_: asyncio.Task[None]) -> None:
await self._join_fut
self._exitcode = self._proc.exitcode
self._proc.close()
for reader in self._stdio_readers:
await reader.aclose()
await aio.cancel_and_wait(ping_task, read_ipc_task, main_task)

if memory_monitor_task is not None:
Expand Down
Loading