mirror of
https://github.com/openai/codex.git
synced 2026-09-05 15:18:41 +00:00
code-mode: define stdio wire protocol
This commit is contained in:
1
codex-rs/Cargo.lock
generated
1
codex-rs/Cargo.lock
generated
@@ -2519,6 +2519,7 @@ dependencies = [
|
||||
"pretty_assertions",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"tokio",
|
||||
"tokio-util",
|
||||
]
|
||||
|
||||
|
||||
@@ -16,7 +16,9 @@ workspace = true
|
||||
codex-protocol = { workspace = true }
|
||||
serde = { workspace = true, features = ["derive"] }
|
||||
serde_json = { workspace = true }
|
||||
tokio = { workspace = true, features = ["io-util", "sync"] }
|
||||
tokio-util = { workspace = true, features = ["rt"] }
|
||||
|
||||
[dev-dependencies]
|
||||
pretty_assertions = { workspace = true }
|
||||
tokio = { workspace = true, features = ["macros", "rt-multi-thread"] }
|
||||
|
||||
@@ -2,6 +2,7 @@ mod description;
|
||||
mod response;
|
||||
mod runtime;
|
||||
mod session;
|
||||
pub mod wire;
|
||||
|
||||
pub use description::CODE_MODE_PRAGMA_PREFIX;
|
||||
pub use description::CodeModeToolKind;
|
||||
|
||||
290
codex-rs/code-mode-protocol/src/wire.rs
Normal file
290
codex-rs/code-mode-protocol/src/wire.rs
Normal file
@@ -0,0 +1,290 @@
|
||||
use std::io;
|
||||
|
||||
use serde::Deserialize;
|
||||
use serde::Serialize;
|
||||
use serde::de::DeserializeOwned;
|
||||
use serde_json::Value as JsonValue;
|
||||
use tokio::io::AsyncRead;
|
||||
use tokio::io::AsyncReadExt;
|
||||
use tokio::io::AsyncWrite;
|
||||
use tokio::io::AsyncWriteExt;
|
||||
|
||||
pub const MAX_FRAME_BYTES: usize = 128 * 1024 * 1024;
|
||||
|
||||
pub type RequestId = u64;
|
||||
pub type SessionId = u64;
|
||||
pub type CallbackId = u64;
|
||||
|
||||
#[derive(Clone, Debug, Deserialize, Eq, Hash, PartialEq, Serialize)]
|
||||
#[serde(transparent)]
|
||||
pub struct CellId(String);
|
||||
|
||||
impl CellId {
|
||||
pub fn new(value: impl Into<String>) -> Self {
|
||||
Self(value.into())
|
||||
}
|
||||
|
||||
pub fn as_str(&self) -> &str {
|
||||
&self.0
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
|
||||
#[serde(tag = "type", rename_all = "snake_case")]
|
||||
pub enum ClientMessage {
|
||||
Request {
|
||||
id: RequestId,
|
||||
request: HostRequest,
|
||||
},
|
||||
CancelRequest {
|
||||
id: RequestId,
|
||||
},
|
||||
CallbackResponse {
|
||||
id: CallbackId,
|
||||
response: CallbackResponse,
|
||||
},
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
|
||||
#[serde(tag = "type", rename_all = "snake_case")]
|
||||
pub enum HostMessage {
|
||||
Response {
|
||||
id: RequestId,
|
||||
result: WireResult<HostResponse>,
|
||||
},
|
||||
CallbackRequest {
|
||||
id: CallbackId,
|
||||
session_id: SessionId,
|
||||
request: CallbackRequest,
|
||||
},
|
||||
CancelCallback {
|
||||
id: CallbackId,
|
||||
},
|
||||
CellClosed {
|
||||
session_id: SessionId,
|
||||
cell_id: CellId,
|
||||
},
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
|
||||
#[serde(tag = "method", rename_all = "snake_case")]
|
||||
pub enum HostRequest {
|
||||
CreateSession,
|
||||
ShutdownSession {
|
||||
session_id: SessionId,
|
||||
},
|
||||
CreateCell {
|
||||
session_id: SessionId,
|
||||
request: CreateCellRequest,
|
||||
},
|
||||
Observe {
|
||||
session_id: SessionId,
|
||||
cell_id: CellId,
|
||||
mode: ObserveMode,
|
||||
},
|
||||
Terminate {
|
||||
session_id: SessionId,
|
||||
cell_id: CellId,
|
||||
},
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
|
||||
#[serde(tag = "type", rename_all = "snake_case")]
|
||||
pub enum HostResponse {
|
||||
SessionCreated { session_id: SessionId },
|
||||
SessionShutdown,
|
||||
CellCreated { cell_id: CellId },
|
||||
Observed { event: CellEvent },
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
|
||||
#[serde(tag = "status", rename_all = "snake_case")]
|
||||
pub enum WireResult<T> {
|
||||
Ok { value: T },
|
||||
Err { error: Error },
|
||||
}
|
||||
|
||||
impl<T> WireResult<T> {
|
||||
pub fn from_result(result: Result<T, Error>) -> Self {
|
||||
match result {
|
||||
Ok(value) => Self::Ok { value },
|
||||
Err(error) => Self::Err { error },
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
|
||||
#[serde(tag = "code", rename_all = "snake_case")]
|
||||
pub enum Error {
|
||||
MissingSession { session_id: SessionId },
|
||||
ShuttingDown,
|
||||
DuplicateCell { cell_id: CellId },
|
||||
MissingCell { cell_id: CellId },
|
||||
ClosedCell { cell_id: CellId },
|
||||
BusyObserver { cell_id: CellId },
|
||||
AlreadyTerminating { cell_id: CellId },
|
||||
Cancelled,
|
||||
Runtime { message: String },
|
||||
InvalidRequest { message: String },
|
||||
CallbackFailed { message: String },
|
||||
Internal { message: String },
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
|
||||
pub struct CreateCellRequest {
|
||||
pub tool_call_id: String,
|
||||
pub enabled_tools: Vec<ToolDefinition>,
|
||||
pub source: String,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
|
||||
pub struct ToolDefinition {
|
||||
pub name: String,
|
||||
pub tool_name: ToolName,
|
||||
pub description: String,
|
||||
pub kind: ToolKind,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
|
||||
pub struct ToolName {
|
||||
pub name: String,
|
||||
pub namespace: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum ToolKind {
|
||||
Function,
|
||||
Freeform,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)]
|
||||
#[serde(tag = "type", rename_all = "snake_case")]
|
||||
pub enum ObserveMode {
|
||||
YieldAfter { duration_ms: u64 },
|
||||
PendingFrontier,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
|
||||
#[serde(tag = "type", rename_all = "snake_case")]
|
||||
pub enum CellEvent {
|
||||
Yielded {
|
||||
content_items: Vec<OutputItem>,
|
||||
},
|
||||
Pending {
|
||||
content_items: Vec<OutputItem>,
|
||||
pending_tool_call_ids: Vec<String>,
|
||||
},
|
||||
Completed {
|
||||
content_items: Vec<OutputItem>,
|
||||
error_text: Option<String>,
|
||||
},
|
||||
Terminated {
|
||||
content_items: Vec<OutputItem>,
|
||||
},
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
|
||||
#[serde(tag = "type", rename_all = "snake_case")]
|
||||
pub enum OutputItem {
|
||||
Text {
|
||||
text: String,
|
||||
},
|
||||
Image {
|
||||
image_url: String,
|
||||
detail: Option<ImageDetail>,
|
||||
},
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum ImageDetail {
|
||||
Auto,
|
||||
Low,
|
||||
High,
|
||||
Original,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
|
||||
#[serde(tag = "type", rename_all = "snake_case")]
|
||||
pub enum CallbackRequest {
|
||||
InvokeTool {
|
||||
invocation: NestedToolCall,
|
||||
},
|
||||
Notify {
|
||||
call_id: String,
|
||||
cell_id: CellId,
|
||||
text: String,
|
||||
},
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
|
||||
pub struct NestedToolCall {
|
||||
pub cell_id: CellId,
|
||||
pub runtime_tool_call_id: String,
|
||||
pub tool_name: ToolName,
|
||||
pub tool_kind: ToolKind,
|
||||
pub input: Option<JsonValue>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
|
||||
#[serde(tag = "type", rename_all = "snake_case")]
|
||||
pub enum CallbackResponse {
|
||||
ToolResult { result: JsonValue },
|
||||
ToolError { error_text: String },
|
||||
NotificationDelivered,
|
||||
NotificationError { error_text: String },
|
||||
}
|
||||
|
||||
pub async fn read_frame<R, T>(reader: &mut R) -> io::Result<Option<T>>
|
||||
where
|
||||
R: AsyncRead + Unpin,
|
||||
T: DeserializeOwned,
|
||||
{
|
||||
let first_length_byte = match reader.read_u8().await {
|
||||
Ok(byte) => byte,
|
||||
Err(err) if err.kind() == io::ErrorKind::UnexpectedEof => return Ok(None),
|
||||
Err(err) => return Err(err),
|
||||
};
|
||||
let mut length_bytes = [first_length_byte, 0, 0, 0];
|
||||
reader.read_exact(&mut length_bytes[1..]).await?;
|
||||
let length = u32::from_be_bytes(length_bytes) as usize;
|
||||
if length > MAX_FRAME_BYTES {
|
||||
return Err(io::Error::new(
|
||||
io::ErrorKind::InvalidData,
|
||||
format!("code-mode IPC frame exceeds {MAX_FRAME_BYTES} bytes"),
|
||||
));
|
||||
}
|
||||
let mut payload = vec![0; length];
|
||||
reader.read_exact(&mut payload).await?;
|
||||
serde_json::from_slice(&payload)
|
||||
.map(Some)
|
||||
.map_err(io::Error::other)
|
||||
}
|
||||
|
||||
pub async fn write_frame<W, T>(writer: &mut W, message: &T) -> io::Result<()>
|
||||
where
|
||||
W: AsyncWrite + Unpin,
|
||||
T: Serialize,
|
||||
{
|
||||
let payload = serde_json::to_vec(message).map_err(io::Error::other)?;
|
||||
if payload.len() > MAX_FRAME_BYTES {
|
||||
return Err(io::Error::new(
|
||||
io::ErrorKind::InvalidInput,
|
||||
format!("code-mode IPC frame exceeds {MAX_FRAME_BYTES} bytes"),
|
||||
));
|
||||
}
|
||||
let length = u32::try_from(payload.len()).map_err(|_| {
|
||||
io::Error::new(
|
||||
io::ErrorKind::InvalidInput,
|
||||
"code-mode IPC frame length exceeds u32",
|
||||
)
|
||||
})?;
|
||||
writer.write_all(&length.to_be_bytes()).await?;
|
||||
writer.write_all(&payload).await?;
|
||||
writer.flush().await
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
#[path = "wire_tests.rs"]
|
||||
mod tests;
|
||||
171
codex-rs/code-mode-protocol/src/wire_tests.rs
Normal file
171
codex-rs/code-mode-protocol/src/wire_tests.rs
Normal file
@@ -0,0 +1,171 @@
|
||||
use std::io;
|
||||
|
||||
use pretty_assertions::assert_eq;
|
||||
use serde_json::json;
|
||||
use tokio::io::AsyncWriteExt;
|
||||
use tokio::io::duplex;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn observe_request_has_a_stable_tagged_shape() {
|
||||
let message = ClientMessage::Request {
|
||||
id: 7,
|
||||
request: HostRequest::Observe {
|
||||
session_id: 3,
|
||||
cell_id: CellId::new("cell-9"),
|
||||
mode: ObserveMode::YieldAfter { duration_ms: 250 },
|
||||
},
|
||||
};
|
||||
|
||||
assert_eq!(
|
||||
serde_json::to_value(message).unwrap(),
|
||||
json!({
|
||||
"type": "request",
|
||||
"id": 7,
|
||||
"request": {
|
||||
"method": "observe",
|
||||
"session_id": 3,
|
||||
"cell_id": "cell-9",
|
||||
"mode": {
|
||||
"type": "yield_after",
|
||||
"duration_ms": 250,
|
||||
},
|
||||
},
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cell_closed_notification_has_a_stable_tagged_shape() {
|
||||
let message = HostMessage::CellClosed {
|
||||
session_id: 3,
|
||||
cell_id: CellId::new("cell-9"),
|
||||
};
|
||||
|
||||
assert_eq!(
|
||||
serde_json::to_value(message).unwrap(),
|
||||
json!({
|
||||
"type": "cell_closed",
|
||||
"session_id": 3,
|
||||
"cell_id": "cell-9",
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn busy_observer_error_has_a_stable_tagged_shape() {
|
||||
let message = HostMessage::Response {
|
||||
id: 11,
|
||||
result: WireResult::Err {
|
||||
error: Error::BusyObserver {
|
||||
cell_id: CellId::new("cell-2"),
|
||||
},
|
||||
},
|
||||
};
|
||||
|
||||
assert_eq!(
|
||||
serde_json::to_value(message).unwrap(),
|
||||
json!({
|
||||
"type": "response",
|
||||
"id": 11,
|
||||
"result": {
|
||||
"status": "err",
|
||||
"error": {
|
||||
"code": "busy_observer",
|
||||
"cell_id": "cell-2",
|
||||
},
|
||||
},
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn callback_result_and_error_round_trip() {
|
||||
for response in [
|
||||
CallbackResponse::ToolResult {
|
||||
result: json!({"value": 42}),
|
||||
},
|
||||
CallbackResponse::ToolError {
|
||||
error_text: "tool failed".to_string(),
|
||||
},
|
||||
CallbackResponse::NotificationDelivered,
|
||||
CallbackResponse::NotificationError {
|
||||
error_text: "notify failed".to_string(),
|
||||
},
|
||||
] {
|
||||
let message = ClientMessage::CallbackResponse { id: 19, response };
|
||||
let encoded = serde_json::to_vec(&message).unwrap();
|
||||
|
||||
assert_eq!(
|
||||
serde_json::from_slice::<ClientMessage>(&encoded).unwrap(),
|
||||
message
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn frame_round_trip_preserves_message() {
|
||||
let (mut client, mut server) = duplex(1024);
|
||||
let message = ClientMessage::Request {
|
||||
id: 7,
|
||||
request: HostRequest::CreateSession,
|
||||
};
|
||||
|
||||
write_frame(&mut client, &message).await.unwrap();
|
||||
|
||||
assert_eq!(read_frame(&mut server).await.unwrap(), Some(message));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn clean_eof_returns_none() {
|
||||
let (client, mut server) = duplex(16);
|
||||
drop(client);
|
||||
|
||||
assert_eq!(
|
||||
read_frame::<_, ClientMessage>(&mut server).await.unwrap(),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn partial_frame_header_is_rejected() {
|
||||
let (mut client, mut server) = duplex(16);
|
||||
client.write_all(&[0, 1]).await.unwrap();
|
||||
client.shutdown().await.unwrap();
|
||||
|
||||
let error = read_frame::<_, ClientMessage>(&mut server)
|
||||
.await
|
||||
.unwrap_err();
|
||||
|
||||
assert_eq!(error.kind(), io::ErrorKind::UnexpectedEof);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn oversized_frame_is_rejected_before_allocation() {
|
||||
let (mut client, mut server) = duplex(16);
|
||||
let oversized_length = u32::try_from(MAX_FRAME_BYTES + 1).unwrap();
|
||||
client
|
||||
.write_all(&oversized_length.to_be_bytes())
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let error = read_frame::<_, ClientMessage>(&mut server)
|
||||
.await
|
||||
.unwrap_err();
|
||||
|
||||
assert_eq!(error.kind(), io::ErrorKind::InvalidData);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn malformed_json_frame_is_rejected() {
|
||||
let (mut client, mut server) = duplex(16);
|
||||
client.write_all(&1_u32.to_be_bytes()).await.unwrap();
|
||||
client.write_all(b"{").await.unwrap();
|
||||
|
||||
let error = read_frame::<_, ClientMessage>(&mut server)
|
||||
.await
|
||||
.unwrap_err();
|
||||
|
||||
assert_eq!(error.kind(), io::ErrorKind::Other);
|
||||
}
|
||||
Reference in New Issue
Block a user