mirror of
https://github.com/openai/codex.git
synced 2026-09-14 11:57:03 +00:00
codex: persist code mode runner sessions
This commit is contained in:
@@ -15,6 +15,7 @@ use crate::models_manager::manager::ModelsManager;
|
||||
use crate::plugins::PluginsManager;
|
||||
use crate::skills::SkillsManager;
|
||||
use crate::state_db::StateDbHandle;
|
||||
use crate::tools::code_mode::CodeModeProcess;
|
||||
use crate::tools::network_approval::NetworkApprovalService;
|
||||
use crate::tools::runtimes::ExecveSessionApproval;
|
||||
use crate::tools::sandboxing::ApprovalStore;
|
||||
@@ -31,12 +32,14 @@ use tokio_util::sync::CancellationToken;
|
||||
|
||||
pub(crate) struct CodeModeStoreService {
|
||||
stored_values: Mutex<HashMap<String, JsonValue>>,
|
||||
process: Mutex<Option<CodeModeProcess>>,
|
||||
}
|
||||
|
||||
impl Default for CodeModeStoreService {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
stored_values: Mutex::new(HashMap::new()),
|
||||
process: Mutex::new(None),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -49,6 +52,14 @@ impl CodeModeStoreService {
|
||||
pub(crate) async fn replace_stored_values(&self, values: HashMap<String, JsonValue>) {
|
||||
*self.stored_values.lock().await = values;
|
||||
}
|
||||
|
||||
pub(crate) async fn store_process(&self, process: CodeModeProcess) {
|
||||
*self.process.lock().await = Some(process);
|
||||
}
|
||||
|
||||
pub(crate) async fn take_process(&self) -> Option<CodeModeProcess> {
|
||||
self.process.lock().await.take()
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) struct SessionServices {
|
||||
|
||||
@@ -30,6 +30,7 @@ use tokio::io::AsyncBufReadExt;
|
||||
use tokio::io::AsyncReadExt;
|
||||
use tokio::io::AsyncWriteExt;
|
||||
use tokio::io::BufReader;
|
||||
use tokio::task::JoinHandle;
|
||||
|
||||
const CODE_MODE_RUNNER_SOURCE: &str = include_str!("code_mode_runner.cjs");
|
||||
const CODE_MODE_BRIDGE_SOURCE: &str = include_str!("code_mode_bridge.js");
|
||||
@@ -42,6 +43,37 @@ struct ExecContext {
|
||||
tracker: SharedTurnDiffTracker,
|
||||
}
|
||||
|
||||
pub(crate) struct CodeModeProcess {
|
||||
child: tokio::process::Child,
|
||||
stdin: tokio::process::ChildStdin,
|
||||
stdout_lines: tokio::io::Lines<BufReader<tokio::process::ChildStdout>>,
|
||||
stderr_task: Option<JoinHandle<String>>,
|
||||
}
|
||||
|
||||
impl CodeModeProcess {
|
||||
fn has_exited(&mut self) -> Result<bool, String> {
|
||||
self.child
|
||||
.try_wait()
|
||||
.map(|status| status.is_some())
|
||||
.map_err(|err| format!("failed to inspect {PUBLIC_TOOL_NAME} runner: {err}"))
|
||||
}
|
||||
|
||||
async fn wait_for_exit(&mut self) -> Result<std::process::ExitStatus, String> {
|
||||
self.child
|
||||
.wait()
|
||||
.await
|
||||
.map_err(|err| format!("failed to wait for {PUBLIC_TOOL_NAME} runner: {err}"))
|
||||
}
|
||||
|
||||
async fn stderr(&mut self) -> Result<String, String> {
|
||||
self.stderr_task
|
||||
.take()
|
||||
.ok_or_else(|| format!("{PUBLIC_TOOL_NAME} stderr collector missing"))?
|
||||
.await
|
||||
.map_err(|err| format!("failed to collect {PUBLIC_TOOL_NAME} stderr: {err}"))
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
enum CodeModeToolKind {
|
||||
@@ -63,12 +95,14 @@ struct EnabledTool {
|
||||
#[derive(Serialize)]
|
||||
#[serde(tag = "type", rename_all = "snake_case")]
|
||||
enum HostToNodeMessage {
|
||||
Init {
|
||||
Start {
|
||||
session_id: String,
|
||||
enabled_tools: Vec<EnabledTool>,
|
||||
stored_values: HashMap<String, JsonValue>,
|
||||
source: String,
|
||||
},
|
||||
Response {
|
||||
session_id: String,
|
||||
id: String,
|
||||
code_mode_result: JsonValue,
|
||||
},
|
||||
@@ -78,12 +112,14 @@ enum HostToNodeMessage {
|
||||
#[serde(tag = "type", rename_all = "snake_case")]
|
||||
enum NodeToHostMessage {
|
||||
ToolCall {
|
||||
session_id: String,
|
||||
id: String,
|
||||
name: String,
|
||||
#[serde(default)]
|
||||
input: Option<JsonValue>,
|
||||
},
|
||||
Result {
|
||||
session_id: String,
|
||||
content_items: Vec<JsonValue>,
|
||||
stored_values: HashMap<String, JsonValue>,
|
||||
#[serde(default)]
|
||||
@@ -138,20 +174,33 @@ pub(crate) async fn execute(
|
||||
let enabled_tools = build_enabled_tools(&exec).await;
|
||||
let stored_values = exec.session.services.code_mode_store.stored_values().await;
|
||||
let source = build_source(&code, &enabled_tools).map_err(FunctionCallError::RespondToModel)?;
|
||||
execute_node(exec, source, enabled_tools, stored_values)
|
||||
.await
|
||||
.map_err(FunctionCallError::RespondToModel)
|
||||
let mut process = match exec.session.services.code_mode_store.take_process().await {
|
||||
Some(mut process) => {
|
||||
if matches!(process.has_exited(), Ok(false)) {
|
||||
process
|
||||
} else {
|
||||
spawn_code_mode_process(&exec)
|
||||
.await
|
||||
.map_err(FunctionCallError::RespondToModel)?
|
||||
}
|
||||
}
|
||||
None => spawn_code_mode_process(&exec)
|
||||
.await
|
||||
.map_err(FunctionCallError::RespondToModel)?,
|
||||
};
|
||||
let result = execute_node(&exec, &mut process, source, enabled_tools, stored_values).await;
|
||||
if result.is_ok() && matches!(process.has_exited(), Ok(false)) {
|
||||
exec.session
|
||||
.services
|
||||
.code_mode_store
|
||||
.store_process(process)
|
||||
.await;
|
||||
}
|
||||
result.map_err(FunctionCallError::RespondToModel)
|
||||
}
|
||||
|
||||
async fn execute_node(
|
||||
exec: ExecContext,
|
||||
source: String,
|
||||
enabled_tools: Vec<EnabledTool>,
|
||||
stored_values: HashMap<String, JsonValue>,
|
||||
) -> Result<FunctionToolOutput, String> {
|
||||
async fn spawn_code_mode_process(exec: &ExecContext) -> Result<CodeModeProcess, String> {
|
||||
let node_path = resolve_compatible_node(exec.turn.config.js_repl_node_path.as_deref()).await?;
|
||||
let started_at = std::time::Instant::now();
|
||||
|
||||
let env = create_env(&exec.turn.shell_environment_policy, None);
|
||||
let mut cmd = tokio::process::Command::new(&node_path);
|
||||
cmd.arg("--experimental-vm-modules");
|
||||
@@ -176,7 +225,7 @@ async fn execute_node(
|
||||
.stderr
|
||||
.take()
|
||||
.ok_or_else(|| format!("{PUBLIC_TOOL_NAME} runner missing stderr"))?;
|
||||
let mut stdin = child
|
||||
let stdin = child
|
||||
.stdin
|
||||
.take()
|
||||
.ok_or_else(|| format!("{PUBLIC_TOOL_NAME} runner missing stdin"))?;
|
||||
@@ -188,19 +237,38 @@ async fn execute_node(
|
||||
String::from_utf8_lossy(&buf).trim().to_string()
|
||||
});
|
||||
|
||||
Ok(CodeModeProcess {
|
||||
child,
|
||||
stdin,
|
||||
stdout_lines: BufReader::new(stdout).lines(),
|
||||
stderr_task: Some(stderr_task),
|
||||
})
|
||||
}
|
||||
|
||||
async fn execute_node(
|
||||
exec: &ExecContext,
|
||||
process: &mut CodeModeProcess,
|
||||
source: String,
|
||||
enabled_tools: Vec<EnabledTool>,
|
||||
stored_values: HashMap<String, JsonValue>,
|
||||
) -> Result<FunctionToolOutput, String> {
|
||||
let started_at = std::time::Instant::now();
|
||||
let session_id = uuid::Uuid::new_v4().to_string();
|
||||
|
||||
write_message(
|
||||
&mut stdin,
|
||||
&HostToNodeMessage::Init {
|
||||
enabled_tools: enabled_tools.clone(),
|
||||
&mut process.stdin,
|
||||
&HostToNodeMessage::Start {
|
||||
session_id: session_id.clone(),
|
||||
enabled_tools,
|
||||
stored_values,
|
||||
source,
|
||||
},
|
||||
)
|
||||
.await?;
|
||||
|
||||
let mut stdout_lines = BufReader::new(stdout).lines();
|
||||
let mut pending_result = None;
|
||||
while let Some(line) = stdout_lines
|
||||
while let Some(line) = process
|
||||
.stdout_lines
|
||||
.next_line()
|
||||
.await
|
||||
.map_err(|err| format!("failed to read {PUBLIC_TOOL_NAME} runner stdout: {err}"))?
|
||||
@@ -212,19 +280,36 @@ async fn execute_node(
|
||||
format!("invalid {PUBLIC_TOOL_NAME} runner message: {err}; line={line}")
|
||||
})?;
|
||||
match message {
|
||||
NodeToHostMessage::ToolCall { id, name, input } => {
|
||||
NodeToHostMessage::ToolCall {
|
||||
session_id: message_session_id,
|
||||
id,
|
||||
name,
|
||||
input,
|
||||
} => {
|
||||
if message_session_id != session_id {
|
||||
return Err(format!(
|
||||
"unexpected {PUBLIC_TOOL_NAME} runner tool call session id: {message_session_id}"
|
||||
));
|
||||
}
|
||||
let response = HostToNodeMessage::Response {
|
||||
session_id: message_session_id,
|
||||
id,
|
||||
code_mode_result: call_nested_tool(exec.clone(), name, input).await,
|
||||
};
|
||||
write_message(&mut stdin, &response).await?;
|
||||
write_message(&mut process.stdin, &response).await?;
|
||||
}
|
||||
NodeToHostMessage::Result {
|
||||
session_id: message_session_id,
|
||||
content_items,
|
||||
stored_values,
|
||||
error_text,
|
||||
max_output_tokens_per_exec_call,
|
||||
} => {
|
||||
if message_session_id != session_id {
|
||||
return Err(format!(
|
||||
"unexpected {PUBLIC_TOOL_NAME} runner result session id: {message_session_id}"
|
||||
));
|
||||
}
|
||||
exec.session
|
||||
.services
|
||||
.code_mode_store
|
||||
@@ -240,20 +325,11 @@ async fn execute_node(
|
||||
}
|
||||
}
|
||||
|
||||
drop(stdin);
|
||||
|
||||
let status = child
|
||||
.wait()
|
||||
.await
|
||||
.map_err(|err| format!("failed to wait for {PUBLIC_TOOL_NAME} runner: {err}"))?;
|
||||
let stderr = stderr_task
|
||||
.await
|
||||
.map_err(|err| format!("failed to collect {PUBLIC_TOOL_NAME} stderr: {err}"))?;
|
||||
let wall_time = started_at.elapsed();
|
||||
let success = status.success();
|
||||
|
||||
let Some((mut content_items, error_text, max_output_tokens_per_exec_call)) = pending_result
|
||||
else {
|
||||
let status = process.wait_for_exit().await?;
|
||||
let stderr = process.stderr().await?;
|
||||
let message = if stderr.is_empty() {
|
||||
format!("{PUBLIC_TOOL_NAME} runner exited without returning a result (status {status})")
|
||||
} else {
|
||||
@@ -262,14 +338,8 @@ async fn execute_node(
|
||||
return Err(message);
|
||||
};
|
||||
|
||||
if !success {
|
||||
let error_text = error_text.unwrap_or_else(|| {
|
||||
if stderr.is_empty() {
|
||||
format!("Process exited with status {status}")
|
||||
} else {
|
||||
stderr
|
||||
}
|
||||
});
|
||||
let success = error_text.is_none();
|
||||
if let Some(error_text) = error_text {
|
||||
content_items.push(FunctionCallOutputContentItem::InputText {
|
||||
text: format!("Script error:\n{error_text}"),
|
||||
});
|
||||
|
||||
@@ -21,11 +21,10 @@ function createProtocol() {
|
||||
|
||||
let nextId = 0;
|
||||
const pending = new Map();
|
||||
let initResolve;
|
||||
let initReject;
|
||||
const init = new Promise((resolve, reject) => {
|
||||
initResolve = resolve;
|
||||
initReject = reject;
|
||||
const sessions = new Map();
|
||||
let closedResolve;
|
||||
const closed = new Promise((resolve) => {
|
||||
closedResolve = resolve;
|
||||
});
|
||||
|
||||
rl.on('line', (line) => {
|
||||
@@ -37,35 +36,38 @@ function createProtocol() {
|
||||
try {
|
||||
message = JSON.parse(line);
|
||||
} catch (error) {
|
||||
initReject(error);
|
||||
process.stderr.write(`${formatErrorText(error)}\n`);
|
||||
return;
|
||||
}
|
||||
|
||||
if (message.type === 'init') {
|
||||
initResolve(message);
|
||||
if (message.type === 'start') {
|
||||
const session = { id: String(message.session_id) };
|
||||
sessions.set(session.id, session);
|
||||
void processSession(protocol, sessions, session, message);
|
||||
return;
|
||||
}
|
||||
|
||||
if (message.type === 'response') {
|
||||
const entry = pending.get(message.id);
|
||||
const entry = pending.get(`${message.session_id}:${message.id}`);
|
||||
if (!entry) {
|
||||
return;
|
||||
}
|
||||
pending.delete(message.id);
|
||||
pending.delete(`${message.session_id}:${message.id}`);
|
||||
entry.resolve(message.code_mode_result ?? '');
|
||||
return;
|
||||
}
|
||||
|
||||
initReject(new Error(`Unknown protocol message type: ${message.type}`));
|
||||
process.stderr.write(`Unknown protocol message type: ${message.type}\n`);
|
||||
});
|
||||
|
||||
rl.on('close', () => {
|
||||
const error = new Error('stdin closed');
|
||||
initReject(error);
|
||||
for (const entry of pending.values()) {
|
||||
entry.reject(error);
|
||||
}
|
||||
pending.clear();
|
||||
sessions.clear();
|
||||
closedResolve();
|
||||
});
|
||||
|
||||
function send(message) {
|
||||
@@ -80,18 +82,20 @@ function createProtocol() {
|
||||
});
|
||||
}
|
||||
|
||||
function request(type, payload) {
|
||||
function request(sessionId, type, payload) {
|
||||
const id = `msg-${++nextId}`;
|
||||
const pendingKey = `${sessionId}:${id}`;
|
||||
return new Promise((resolve, reject) => {
|
||||
pending.set(id, { resolve, reject });
|
||||
void send({ type, id, ...payload }).catch((error) => {
|
||||
pending.delete(id);
|
||||
pending.set(pendingKey, { resolve, reject });
|
||||
void send({ type, session_id: sessionId, id, ...payload }).catch((error) => {
|
||||
pending.delete(pendingKey);
|
||||
reject(error);
|
||||
});
|
||||
});
|
||||
}
|
||||
|
||||
return { init, request, send };
|
||||
const protocol = { closed, request, send };
|
||||
return protocol;
|
||||
}
|
||||
|
||||
function readContentItems(context) {
|
||||
@@ -112,9 +116,9 @@ function cloneJsonValue(value) {
|
||||
return JSON.parse(JSON.stringify(value));
|
||||
}
|
||||
|
||||
function createToolCaller(protocol) {
|
||||
function createToolCaller(protocol, sessionId) {
|
||||
return (name, input) =>
|
||||
protocol.request('tool_call', {
|
||||
protocol.request(sessionId, 'tool_call', {
|
||||
name: String(name),
|
||||
input,
|
||||
});
|
||||
@@ -348,14 +352,14 @@ function createModuleResolver(context, callTool, enabledTools, state) {
|
||||
};
|
||||
}
|
||||
|
||||
async function runModule(context, request, state, callTool) {
|
||||
async function runModule(context, start, state, callTool) {
|
||||
const resolveModule = createModuleResolver(
|
||||
context,
|
||||
callTool,
|
||||
request.enabled_tools ?? [],
|
||||
start.enabled_tools ?? [],
|
||||
state
|
||||
);
|
||||
const mainModule = new SourceTextModule(request.source, {
|
||||
const mainModule = new SourceTextModule(start.source, {
|
||||
context,
|
||||
identifier: 'exec_main.mjs',
|
||||
importModuleDynamically: async (specifier) => resolveModule(specifier),
|
||||
@@ -365,40 +369,45 @@ async function runModule(context, request, state, callTool) {
|
||||
await mainModule.evaluate();
|
||||
}
|
||||
|
||||
async function main() {
|
||||
const protocol = createProtocol();
|
||||
const request = await protocol.init;
|
||||
async function processSession(protocol, sessions, session, start) {
|
||||
const state = {
|
||||
maxOutputTokensPerExecCall: DEFAULT_MAX_OUTPUT_TOKENS_PER_EXEC_CALL,
|
||||
storedValues: cloneJsonValue(request.stored_values ?? {}),
|
||||
storedValues: cloneJsonValue(start.stored_values ?? {}),
|
||||
};
|
||||
const callTool = createToolCaller(protocol);
|
||||
const callTool = createToolCaller(protocol, session.id);
|
||||
const context = vm.createContext({
|
||||
__codexContentItems: [],
|
||||
__codex_tool_call: callTool,
|
||||
});
|
||||
|
||||
try {
|
||||
await runModule(context, request, state, callTool);
|
||||
await runModule(context, start, state, callTool);
|
||||
await protocol.send({
|
||||
type: 'result',
|
||||
session_id: session.id,
|
||||
content_items: readContentItems(context),
|
||||
stored_values: state.storedValues,
|
||||
max_output_tokens_per_exec_call: state.maxOutputTokensPerExecCall,
|
||||
});
|
||||
process.exit(0);
|
||||
} catch (error) {
|
||||
await protocol.send({
|
||||
type: 'result',
|
||||
session_id: session.id,
|
||||
content_items: readContentItems(context),
|
||||
stored_values: state.storedValues,
|
||||
error_text: formatErrorText(error),
|
||||
max_output_tokens_per_exec_call: state.maxOutputTokensPerExecCall,
|
||||
});
|
||||
process.exit(1);
|
||||
} finally {
|
||||
sessions.delete(session.id);
|
||||
}
|
||||
}
|
||||
|
||||
async function main() {
|
||||
const protocol = createProtocol();
|
||||
await protocol.closed;
|
||||
}
|
||||
|
||||
void main().catch(async (error) => {
|
||||
try {
|
||||
process.stderr.write(`${formatErrorText(error)}\n`);
|
||||
|
||||
Reference in New Issue
Block a user