Current section

Files

Jump to
codex_sdk lib codex io transport erlexec.ex
Raw

lib/codex/io/transport/erlexec.ex

defmodule Codex.IO.Transport.Erlexec do
@moduledoc "Erlexec-backed subprocess transport implementation."
use GenServer
import Kernel, except: [send: 2]
alias Codex.Config.Defaults
alias Codex.TaskSupport
@behaviour Codex.IO.Transport
@default_max_buffer_size Defaults.transport_max_buffer_size()
@default_max_stderr_buffer_size Defaults.transport_max_stderr_buffer_size()
@default_call_timeout Defaults.transport_call_timeout_ms()
@force_close_timeout Defaults.transport_force_close_timeout_ms()
@default_headless_timeout_ms Defaults.transport_headless_timeout_ms()
@finalize_delay_ms Defaults.transport_finalize_delay_ms()
@max_lines_per_batch Defaults.transport_max_lines_per_batch()
defstruct [
:subprocess,
:status,
:task_supervisor,
subscribers: %{},
stdout_buffer: "",
pending_lines: :queue.new(),
drain_scheduled?: false,
overflowed?: false,
stderr_buffer: "",
max_buffer_size: @default_max_buffer_size,
max_stderr_buffer_size: @default_max_stderr_buffer_size,
pending_calls: %{},
finalize_timer_ref: nil,
headless_timeout_ms: @default_headless_timeout_ms,
headless_timer_ref: nil
]
@type subscriber_info :: %{
monitor_ref: reference(),
tag: Codex.IO.Transport.subscription_tag()
}
@impl Codex.IO.Transport
def start(opts) when is_list(opts) do
case GenServer.start(__MODULE__, opts) do
{:ok, pid} -> {:ok, pid}
{:error, reason} -> transport_error(reason)
end
catch
:exit, reason ->
transport_error(reason)
end
@impl Codex.IO.Transport
def start_link(opts) when is_list(opts) do
case GenServer.start_link(__MODULE__, opts) do
{:ok, pid} -> {:ok, pid}
{:error, reason} -> transport_error(reason)
end
catch
:exit, reason ->
transport_error(reason)
end
@impl Codex.IO.Transport
def send(transport, message) when is_pid(transport) do
case safe_call(transport, {:send, message}) do
{:ok, result} -> result
{:error, reason} -> transport_error(reason)
end
end
@impl Codex.IO.Transport
def subscribe(transport, pid) when is_pid(transport) and is_pid(pid) do
subscribe(transport, pid, :legacy)
end
@impl Codex.IO.Transport
def subscribe(transport, pid, tag)
when is_pid(transport) and is_pid(pid) and (tag == :legacy or is_reference(tag)) do
case safe_call(transport, {:subscribe, pid, tag}) do
{:ok, result} -> result
{:error, reason} -> transport_error(reason)
end
end
@impl Codex.IO.Transport
def close(transport) when is_pid(transport) do
GenServer.stop(transport, :normal)
catch
:exit, {:noproc, _} -> :ok
:exit, :noproc -> :ok
end
@impl Codex.IO.Transport
def force_close(transport) when is_pid(transport) do
case safe_call(transport, :force_close, @force_close_timeout) do
{:ok, :ok} ->
:ok
{:error, :not_connected} ->
:ok
{:error, reason} ->
transport_error(reason)
end
end
@impl Codex.IO.Transport
def status(transport) when is_pid(transport) do
case safe_call(transport, :status) do
{:ok, status} when status in [:connected, :disconnected, :error] -> status
{:ok, _other} -> :error
{:error, _reason} -> :disconnected
end
end
@impl Codex.IO.Transport
def end_input(transport) when is_pid(transport) do
case safe_call(transport, :end_input) do
{:ok, result} -> result
{:error, reason} -> transport_error(reason)
end
end
@impl Codex.IO.Transport
def stderr(transport) when is_pid(transport) do
case safe_call(transport, :stderr) do
{:ok, data} when is_binary(data) -> data
_ -> ""
end
end
@impl GenServer
def init(opts) do
command = Keyword.fetch!(opts, :command)
args = Keyword.get(opts, :args, [])
cwd = Keyword.get(opts, :cwd)
env = Keyword.get(opts, :env, [])
task_supervisor = Keyword.get(opts, :task_supervisor, Codex.TaskSupervisor)
subscriber = Keyword.get(opts, :subscriber)
headless_timeout_ms = Keyword.get(opts, :headless_timeout_ms, @default_headless_timeout_ms)
max_buffer_size = Keyword.get(opts, :max_buffer_size, @default_max_buffer_size)
max_stderr_buffer_size =
Keyword.get(opts, :max_stderr_buffer_size, @default_max_stderr_buffer_size)
case start_subprocess(command, args,
cwd: cwd,
env: env,
task_supervisor: task_supervisor,
subscriber: subscriber,
headless_timeout_ms: headless_timeout_ms,
max_buffer_size: max_buffer_size,
max_stderr_buffer_size: max_stderr_buffer_size
) do
{:ok, state} -> {:ok, state}
{:error, reason} -> {:stop, reason}
end
end
@impl GenServer
def handle_call({:subscribe, pid, tag}, _from, state) do
state =
state
|> put_subscriber(pid, tag)
|> cancel_headless_timer()
{:reply, :ok, state}
end
def handle_call({:send, message}, from, %{subprocess: {pid, _os_pid}} = state) do
case start_io_task(state, fn -> send_payload(pid, message) end) do
{:ok, task} ->
pending_calls = Map.put(state.pending_calls, task.ref, from)
{:noreply, %{state | pending_calls: pending_calls}}
{:error, reason} ->
{:reply, transport_error(reason), state}
end
end
def handle_call({:send, _message}, _from, %{subprocess: nil} = state) do
{:reply, transport_error(:not_connected), state}
end
def handle_call(:status, _from, state), do: {:reply, state.status, state}
def handle_call(:end_input, from, %{subprocess: {pid, _os_pid}} = state) do
case start_io_task(state, fn -> send_eof(pid) end) do
{:ok, task} ->
pending_calls = Map.put(state.pending_calls, task.ref, from)
{:noreply, %{state | pending_calls: pending_calls}}
{:error, reason} ->
{:reply, transport_error(reason), state}
end
end
def handle_call(:end_input, _from, %{subprocess: nil} = state) do
{:reply, transport_error(:not_connected), state}
end
def handle_call(:stderr, _from, state), do: {:reply, state.stderr_buffer, state}
def handle_call(:force_close, _from, state) do
state = force_stop_subprocess(state)
{:stop, :normal, :ok, state}
end
@impl GenServer
def handle_info({:stdout, os_pid, chunk}, %{subprocess: {_pid, os_pid}} = state) do
state =
state
|> append_stdout_data(IO.iodata_to_binary(chunk))
|> maybe_schedule_drain()
{:noreply, state}
end
def handle_info({:stderr, _os_pid, chunk}, state) do
data = IO.iodata_to_binary(chunk)
stderr_buffer = append_stderr_data(state.stderr_buffer, data, state.max_stderr_buffer_size)
{:noreply, %{state | stderr_buffer: stderr_buffer}}
end
def handle_info({ref, result}, %{pending_calls: pending_calls} = state)
when is_reference(ref) do
case Map.pop(pending_calls, ref) do
{nil, _} ->
{:noreply, state}
{from, rest} ->
Process.demonitor(ref, [:flush])
GenServer.reply(from, normalize_call_result(result))
{:noreply, %{state | pending_calls: rest}}
end
end
def handle_info({:DOWN, os_pid, :process, pid, reason}, %{subprocess: {pid, os_pid}} = state) do
state = cancel_finalize_timer(state)
timer_ref =
Process.send_after(self(), {:finalize_exit, os_pid, pid, reason}, @finalize_delay_ms)
{:noreply, %{state | finalize_timer_ref: timer_ref}}
end
def handle_info({:finalize_exit, os_pid, pid, reason}, %{subprocess: {pid, os_pid}} = state) do
state =
state
|> Map.put(:finalize_timer_ref, nil)
|> Map.put(:drain_scheduled?, false)
|> drain_stdout_lines(@max_lines_per_batch)
if :queue.is_empty(state.pending_lines) do
state = flush_stdout_fragment(state)
if state.stderr_buffer != "" do
send_event(state.subscribers, {:stderr, state.stderr_buffer})
end
send_event(state.subscribers, {:exit, reason})
{:stop, :normal, %{state | status: :disconnected, subprocess: nil}}
else
Kernel.send(self(), {:finalize_exit, os_pid, pid, reason})
{:noreply, state}
end
end
def handle_info({:DOWN, ref, :process, pid, reason}, %{pending_calls: pending_calls} = state)
when is_reference(ref) do
case Map.pop(pending_calls, ref) do
{from, rest} when not is_nil(from) ->
GenServer.reply(from, transport_error({:send_failed, reason}))
{:noreply, %{state | pending_calls: rest}}
{nil, _} ->
handle_subscriber_down(ref, pid, state)
end
end
def handle_info(:drain_stdout, state) do
state =
state
|> Map.put(:drain_scheduled?, false)
|> drain_stdout_lines(@max_lines_per_batch)
|> maybe_schedule_drain()
{:noreply, state}
end
def handle_info(:headless_timeout, state) do
state = %{state | headless_timer_ref: nil}
if map_size(state.subscribers) == 0 and not is_nil(state.subprocess) do
{:stop, :normal, state}
else
{:noreply, state}
end
end
def handle_info(_other, state), do: {:noreply, state}
@impl GenServer
def terminate(_reason, state) do
state =
state
|> cancel_finalize_timer()
|> cancel_headless_timer()
demonitor_subscribers(state.subscribers)
cleanup_pending_calls(state.pending_calls)
_ = force_stop_subprocess(state)
:ok
catch
_, _ -> :ok
end
defp safe_call(transport, message, timeout \\ @default_call_timeout)
defp safe_call(transport, message, timeout)
when is_pid(transport) and is_integer(timeout) and timeout >= 0 do
case TaskSupport.async_nolink(Codex.TaskSupervisor, fn ->
try do
{:ok, GenServer.call(transport, message, :infinity)}
catch
:exit, reason -> {:error, normalize_call_exit(reason)}
end
end) do
{:ok, task} ->
await_task_result(task, timeout)
{:error, reason} ->
{:error, normalize_call_task_start_error(reason)}
end
end
defp await_task_result(task, timeout) do
case Task.yield(task, timeout) || Task.shutdown(task, :brutal_kill) do
{:ok, result} ->
result
{:exit, reason} ->
{:error, normalize_call_exit(reason)}
nil ->
{:error, :timeout}
end
end
defp normalize_call_task_start_error(:noproc), do: :transport_stopped
defp normalize_call_task_start_error(reason), do: {:call_exit, reason}
defp normalize_call_exit({:noproc, _}), do: :not_connected
defp normalize_call_exit(:noproc), do: :not_connected
defp normalize_call_exit({:normal, _}), do: :not_connected
defp normalize_call_exit({:shutdown, _}), do: :not_connected
defp normalize_call_exit({:timeout, _}), do: :timeout
defp normalize_call_exit(reason), do: {:call_exit, reason}
defp normalize_call_result(:ok), do: :ok
defp normalize_call_result({:error, {:transport, _reason}} = error), do: error
defp normalize_call_result({:error, reason}), do: transport_error(reason)
defp normalize_call_result(other), do: transport_error({:unexpected_task_result, other})
defp start_subprocess(command, args, opts) do
cwd = Keyword.get(opts, :cwd)
env = Keyword.get(opts, :env, [])
subscriber = Keyword.get(opts, :subscriber)
task_supervisor = Keyword.get(opts, :task_supervisor, Codex.TaskSupervisor)
headless_timeout_ms = Keyword.get(opts, :headless_timeout_ms, @default_headless_timeout_ms)
max_buffer_size = Keyword.get(opts, :max_buffer_size, @default_max_buffer_size)
max_stderr_buffer_size =
Keyword.get(opts, :max_stderr_buffer_size, @default_max_stderr_buffer_size)
state = %__MODULE__{
status: :disconnected,
task_supervisor: task_supervisor,
max_buffer_size: max_buffer_size,
max_stderr_buffer_size: max_stderr_buffer_size,
headless_timeout_ms: headless_timeout_ms
}
exec_opts =
[]
|> maybe_put_cwd(cwd)
|> maybe_put_env(env)
|> Kernel.++([:stdin, :stdout, :stderr, :monitor])
argv = normalize_command_argv(command, args)
case :exec.run(argv, exec_opts) do
{:ok, pid, os_pid} ->
state =
state
|> Map.put(:subprocess, {pid, os_pid})
|> Map.put(:status, :connected)
with {:ok, state} <- add_bootstrap_subscriber(state, subscriber) do
{:ok, maybe_schedule_headless_timer(state)}
end
{:error, reason} ->
{:error, reason}
end
end
defp maybe_put_cwd(opts, nil), do: opts
defp maybe_put_cwd(opts, cwd), do: [{:cd, to_charlist(cwd)} | opts]
defp maybe_put_env(opts, env) do
case normalize_env(env) do
[] -> opts
value -> [{:env, value} | opts]
end
end
defp normalize_command_argv(command, args) when is_list(command) do
cond do
command == [] ->
Enum.map(List.wrap(args), &to_charlist/1)
Enum.all?(command, &is_integer/1) ->
[to_charlist(command) | Enum.map(List.wrap(args), &to_charlist/1)]
true ->
Enum.map(command ++ List.wrap(args), &to_charlist/1)
end
end
defp normalize_command_argv(command, args) do
[command | List.wrap(args)] |> Enum.map(&to_charlist/1)
end
defp normalize_env(nil), do: []
defp normalize_env([]), do: []
defp normalize_env(env) when is_map(env) do
Enum.map(env, fn {key, value} -> {to_string(key), to_string(value)} end)
end
defp normalize_env(env) when is_list(env) do
Enum.flat_map(env, fn
{key, value} -> [{to_string(key), to_string(value)}]
_other -> []
end)
end
defp normalize_env(_other), do: []
defp add_bootstrap_subscriber(state, nil), do: {:ok, state}
defp add_bootstrap_subscriber(state, {pid, tag})
when is_pid(pid) and (tag == :legacy or is_reference(tag)) do
{:ok, put_subscriber(state, pid, tag)}
end
defp add_bootstrap_subscriber(_state, _subscriber), do: {:error, :invalid_subscriber}
defp put_subscriber(state, pid, tag) do
subscribers =
case Map.fetch(state.subscribers, pid) do
{:ok, %{monitor_ref: monitor_ref}} ->
Map.put(state.subscribers, pid, %{monitor_ref: monitor_ref, tag: tag})
:error ->
monitor_ref = Process.monitor(pid)
Map.put(state.subscribers, pid, %{monitor_ref: monitor_ref, tag: tag})
end
%{state | subscribers: subscribers}
end
defp handle_subscriber_down(ref, pid, state) do
subscribers =
case Map.pop(state.subscribers, pid) do
{%{monitor_ref: ^ref}, rest} -> rest
{_value, rest} -> rest
end
state = %{state | subscribers: subscribers}
if map_size(subscribers) == 0 do
{:stop, :normal, state}
else
{:noreply, state}
end
end
defp maybe_schedule_headless_timer(%{headless_timer_ref: ref} = state) when not is_nil(ref),
do: state
defp maybe_schedule_headless_timer(%{subscribers: subscribers} = state)
when map_size(subscribers) > 0,
do: state
defp maybe_schedule_headless_timer(%{headless_timeout_ms: :infinity} = state), do: state
defp maybe_schedule_headless_timer(%{headless_timeout_ms: timeout_ms} = state)
when is_integer(timeout_ms) and timeout_ms > 0 do
timer_ref = Process.send_after(self(), :headless_timeout, timeout_ms)
%{state | headless_timer_ref: timer_ref}
end
defp maybe_schedule_headless_timer(state), do: state
defp cancel_headless_timer(%{headless_timer_ref: nil} = state), do: state
defp cancel_headless_timer(state) do
_ = Process.cancel_timer(state.headless_timer_ref, async: false, info: false)
flush_headless_timeout_message()
%{state | headless_timer_ref: nil}
end
defp flush_headless_timeout_message do
receive do
:headless_timeout -> :ok
after
0 -> :ok
end
end
defp start_io_task(state, fun) when is_function(fun, 0) do
TaskSupport.async_nolink(state.task_supervisor, fun)
end
defp send_payload(pid, message) do
payload = message |> normalize_payload() |> ensure_newline()
:exec.send(pid, payload)
:ok
catch
kind, reason ->
transport_error({:send_failed, {kind, reason}})
end
defp send_eof(pid) do
:exec.send(pid, :eof)
:ok
catch
kind, reason ->
transport_error({:send_failed, {kind, reason}})
end
defp send_event(subscribers, event) do
Enum.each(subscribers, fn {pid, info} ->
dispatch_event(pid, info, event)
end)
end
defp dispatch_event(pid, %{tag: :legacy}, {:message, line}),
do: Kernel.send(pid, {:transport_message, line})
defp dispatch_event(pid, %{tag: :legacy}, {:error, reason}),
do: Kernel.send(pid, {:transport_error, reason})
defp dispatch_event(pid, %{tag: :legacy}, {:stderr, data}),
do: Kernel.send(pid, {:transport_stderr, data})
defp dispatch_event(pid, %{tag: :legacy}, {:exit, reason}),
do: Kernel.send(pid, {:transport_exit, reason})
defp dispatch_event(pid, %{tag: ref}, event) when is_reference(ref),
do: Kernel.send(pid, {:codex_io_transport, ref, event})
defp append_stdout_data(%{overflowed?: true} = state, data) do
case String.split(data, "\n", parts: 2) do
[_single] ->
state
[_dropped, rest] ->
state
|> Map.put(:overflowed?, false)
|> Map.put(:stdout_buffer, "")
|> append_stdout_data(rest)
end
end
defp append_stdout_data(state, data) do
full = state.stdout_buffer <> data
{complete_lines, remaining} = split_complete_lines(full)
pending_lines =
Enum.reduce(complete_lines, state.pending_lines, fn line, queue ->
:queue.in(line, queue)
end)
state = %{state | pending_lines: pending_lines, stdout_buffer: "", overflowed?: false}
if byte_size(remaining) > state.max_buffer_size do
send_event(state.subscribers, {:error, {:buffer_overflow, byte_size(remaining)}})
%{state | stdout_buffer: "", overflowed?: true}
else
%{state | stdout_buffer: remaining}
end
end
defp drain_stdout_lines(state, 0), do: state
defp drain_stdout_lines(state, remaining) when is_integer(remaining) and remaining > 0 do
case :queue.out(state.pending_lines) do
{:empty, _queue} ->
state
{{:value, line}, queue} ->
state = %{state | pending_lines: queue}
if byte_size(line) > state.max_buffer_size do
send_event(state.subscribers, {:error, {:buffer_overflow, byte_size(line)}})
else
send_event(state.subscribers, {:message, line})
end
drain_stdout_lines(state, remaining - 1)
end
end
defp maybe_schedule_drain(%{drain_scheduled?: true} = state), do: state
defp maybe_schedule_drain(state) do
if :queue.is_empty(state.pending_lines) do
state
else
Kernel.send(self(), :drain_stdout)
%{state | drain_scheduled?: true}
end
end
defp split_complete_lines(""), do: {[], ""}
defp split_complete_lines(data) do
lines = String.split(data, "\n")
case List.pop_at(lines, -1) do
{nil, _} -> {[], ""}
{"", rest} -> {rest, ""}
{last, rest} -> {rest, last}
end
end
defp flush_stdout_fragment(state) do
line = String.trim(state.stdout_buffer)
cond do
line == "" ->
%{state | stdout_buffer: "", overflowed?: false, drain_scheduled?: false}
byte_size(line) > state.max_buffer_size ->
send_event(state.subscribers, {:error, {:buffer_overflow, byte_size(line)}})
%{state | stdout_buffer: "", overflowed?: false, drain_scheduled?: false}
true ->
send_event(state.subscribers, {:message, line})
%{state | stdout_buffer: "", overflowed?: false, drain_scheduled?: false}
end
end
defp cancel_finalize_timer(%{finalize_timer_ref: nil} = state), do: state
defp cancel_finalize_timer(state) do
_ = Process.cancel_timer(state.finalize_timer_ref, async: false, info: false)
flush_finalize_message(state.subprocess)
%{state | finalize_timer_ref: nil}
end
defp flush_finalize_message({pid, os_pid}) do
receive do
{:finalize_exit, ^os_pid, ^pid, _reason} -> :ok
after
0 -> :ok
end
end
defp flush_finalize_message(_), do: :ok
defp append_stderr_data(_existing, _data, max_size)
when not is_integer(max_size) or max_size <= 0,
do: ""
defp append_stderr_data(existing, data, max_size) do
combined = existing <> data
combined_size = byte_size(combined)
if combined_size <= max_size do
combined
else
:binary.part(combined, combined_size - max_size, max_size)
end
end
defp cleanup_pending_calls(pending_calls) do
Enum.each(pending_calls, fn {ref, from} ->
Process.demonitor(ref, [:flush])
GenServer.reply(from, transport_error(:transport_stopped))
end)
end
defp demonitor_subscribers(subscribers) do
Enum.each(subscribers, fn {_pid, %{monitor_ref: ref}} ->
Process.demonitor(ref, [:flush])
end)
end
defp force_stop_subprocess(%{subprocess: {pid, _os_pid}} = state) do
stop_subprocess(pid)
%{state | subprocess: nil, status: :disconnected}
end
defp force_stop_subprocess(state), do: state
defp stop_subprocess(pid) when is_pid(pid) do
:exec.stop(pid)
_ = :exec.kill(pid, 9)
:ok
catch
_, _ -> :ok
end
defp normalize_payload(message) when is_binary(message), do: message
defp normalize_payload(message) when is_map(message), do: Jason.encode!(message)
defp normalize_payload(message) when is_list(message) do
IO.iodata_to_binary(message)
rescue
ArgumentError ->
Jason.encode!(message)
end
defp normalize_payload(message), do: to_string(message)
defp ensure_newline(payload) do
if String.ends_with?(payload, "\n"), do: payload, else: payload <> "\n"
end
defp transport_error(reason), do: {:error, {:transport, reason}}
end