app-server: serialize stateful rpcs by object key

Motivation:
Concurrent initialized request handling lets later RPCs for the same high-level object run before earlier RPCs publish their side effects. A bare per-key mutex would not guarantee arrival order once tasks are spawned, so same-key requests need an explicit FIFO queue.

Summary:
Add Rust-only serialization metadata to client_request_definitions!, expose ClientRequest::serialization_scope(), and implement a private app-server per-key FIFO dispatcher. The server maps logical protocol scopes to runtime queue keys, including connection-scoped command/exec process IDs and fs watch IDs, while leaving unkeyed requests concurrent.

Testing:
- cargo test -p codex-app-server-protocol client_request_serialization_scope
- cargo test -p codex-app-server request_serialization --lib
- cargo test -p codex-app-server-protocol
- cargo test -p codex-app-server
This commit is contained in:
Ruslan Nigmatullin
2026-04-09 23:08:58 -07:00
parent 1385a44433
commit 3281a1cc5b
4 changed files with 610 additions and 1 deletions

View File

@@ -1,4 +1,5 @@
use std::path::Path;
use std::path::PathBuf;
use crate::JSONRPCNotification;
use crate::JSONRPCRequest;
@@ -8,6 +9,7 @@ use crate::export::write_json_schema;
use crate::protocol::v1;
use crate::protocol::v2;
use codex_experimental_api_macros::ExperimentalApi;
use codex_utils_absolute_path::AbsolutePathBuf;
use schemars::JsonSchema;
use serde::Deserialize;
use serde::Serialize;
@@ -68,6 +70,96 @@ macro_rules! experimental_type_entry {
};
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ClientRequestSerializationScope {
Global(&'static str),
Thread {
thread_id: String,
},
ThreadPath {
path: PathBuf,
},
CommandExecProcess {
process_id: String,
},
FuzzyFileSearchSession {
session_id: String,
},
FsWatch {
watch_id: String,
},
Plugin {
marketplace_path: AbsolutePathBuf,
plugin_name: String,
},
PluginId {
plugin_id: String,
},
McpOauth {
server_name: String,
},
}
macro_rules! serialization_scope_expr {
($actual_params:ident) => {
None
};
($actual_params:ident, global($key:literal)) => {
Some(ClientRequestSerializationScope::Global($key))
};
($actual_params:ident, thread_id($params:ident . $field:ident)) => {
Some(ClientRequestSerializationScope::Thread {
thread_id: $actual_params.$field.clone(),
})
};
($actual_params:ident, thread_or_path($params:ident . $thread_field:ident, $params2:ident . $path_field:ident)) => {
if let Some(path) = $actual_params.$path_field.clone() {
Some(ClientRequestSerializationScope::ThreadPath { path })
} else {
Some(ClientRequestSerializationScope::Thread {
thread_id: $actual_params.$thread_field.clone(),
})
}
};
($actual_params:ident, optional_command_process_id($params:ident . $field:ident)) => {
$actual_params
.$field
.clone()
.map(|process_id| ClientRequestSerializationScope::CommandExecProcess { process_id })
};
($actual_params:ident, command_process_id($params:ident . $field:ident)) => {
Some(ClientRequestSerializationScope::CommandExecProcess {
process_id: $actual_params.$field.clone(),
})
};
($actual_params:ident, fuzzy_session_id($params:ident . $field:ident)) => {
Some(ClientRequestSerializationScope::FuzzyFileSearchSession {
session_id: $actual_params.$field.clone(),
})
};
($actual_params:ident, fs_watch_id($params:ident . $field:ident)) => {
Some(ClientRequestSerializationScope::FsWatch {
watch_id: $actual_params.$field.clone(),
})
};
($actual_params:ident, plugin($params:ident . $marketplace_field:ident, $params2:ident . $plugin_field:ident)) => {
Some(ClientRequestSerializationScope::Plugin {
marketplace_path: $actual_params.$marketplace_field.clone(),
plugin_name: $actual_params.$plugin_field.clone(),
})
};
($actual_params:ident, plugin_id($params:ident . $field:ident)) => {
Some(ClientRequestSerializationScope::PluginId {
plugin_id: $actual_params.$field.clone(),
})
};
($actual_params:ident, mcp_oauth_server($params:ident . $field:ident)) => {
Some(ClientRequestSerializationScope::McpOauth {
server_name: $actual_params.$field.clone(),
})
};
}
/// Generates an `enum ClientRequest` where each variant is a request that the
/// client can send to the server. Each variant has associated `params` and
/// `response` types. Also generates a `export_client_responses()` function to
@@ -80,6 +172,7 @@ macro_rules! client_request_definitions {
$variant:ident $(=> $wire:literal)? {
params: $(#[$params_meta:meta])* $params:ty,
$(inspect_params: $inspect_params:tt,)?
$(serialization: $serialization:ident ( $($serialization_args:tt)* ),)?
response: $response:ty,
}
),* $(,)?
@@ -118,6 +211,19 @@ macro_rules! client_request_definitions {
})
.unwrap_or_else(|| "<unknown>".to_string())
}
pub fn serialization_scope(&self) -> Option<ClientRequestSerializationScope> {
match self {
$(
Self::$variant { params, .. } => {
let _ = params;
serialization_scope_expr!(
params $(, $serialization($($serialization_args)*))?
)
}
)*
}
}
}
/// Typed response from the server to the client.
@@ -244,19 +350,23 @@ client_request_definitions! {
ThreadResume => "thread/resume" {
params: v2::ThreadResumeParams,
inspect_params: true,
serialization: thread_id(params.thread_id),
response: v2::ThreadResumeResponse,
},
ThreadFork => "thread/fork" {
params: v2::ThreadForkParams,
inspect_params: true,
serialization: thread_or_path(params.thread_id, params.path),
response: v2::ThreadForkResponse,
},
ThreadArchive => "thread/archive" {
params: v2::ThreadArchiveParams,
serialization: thread_id(params.thread_id),
response: v2::ThreadArchiveResponse,
},
ThreadUnsubscribe => "thread/unsubscribe" {
params: v2::ThreadUnsubscribeParams,
serialization: thread_id(params.thread_id),
response: v2::ThreadUnsubscribeResponse,
},
#[experimental("thread/increment_elicitation")]
@@ -266,6 +376,7 @@ client_request_definitions! {
/// approval or other elicitation is pending outside the app-server request flow.
ThreadIncrementElicitation => "thread/increment_elicitation" {
params: v2::ThreadIncrementElicitationParams,
serialization: thread_id(params.thread_id),
response: v2::ThreadIncrementElicitationResponse,
},
#[experimental("thread/decrement_elicitation")]
@@ -274,26 +385,32 @@ client_request_definitions! {
/// When the count reaches zero, timeout accounting resumes for the thread.
ThreadDecrementElicitation => "thread/decrement_elicitation" {
params: v2::ThreadDecrementElicitationParams,
serialization: thread_id(params.thread_id),
response: v2::ThreadDecrementElicitationResponse,
},
ThreadSetName => "thread/name/set" {
params: v2::ThreadSetNameParams,
serialization: thread_id(params.thread_id),
response: v2::ThreadSetNameResponse,
},
ThreadMetadataUpdate => "thread/metadata/update" {
params: v2::ThreadMetadataUpdateParams,
serialization: thread_id(params.thread_id),
response: v2::ThreadMetadataUpdateResponse,
},
ThreadUnarchive => "thread/unarchive" {
params: v2::ThreadUnarchiveParams,
serialization: thread_id(params.thread_id),
response: v2::ThreadUnarchiveResponse,
},
ThreadCompactStart => "thread/compact/start" {
params: v2::ThreadCompactStartParams,
serialization: thread_id(params.thread_id),
response: v2::ThreadCompactStartResponse,
},
ThreadShellCommand => "thread/shellCommand" {
params: v2::ThreadShellCommandParams,
serialization: thread_id(params.thread_id),
response: v2::ThreadShellCommandResponse,
},
ThreadAddCreditsNudgeEmail => "thread/addCreditsNudgeEmail" {
@@ -303,10 +420,12 @@ client_request_definitions! {
#[experimental("thread/backgroundTerminals/clean")]
ThreadBackgroundTerminalsClean => "thread/backgroundTerminals/clean" {
params: v2::ThreadBackgroundTerminalsCleanParams,
serialization: thread_id(params.thread_id),
response: v2::ThreadBackgroundTerminalsCleanResponse,
},
ThreadRollback => "thread/rollback" {
params: v2::ThreadRollbackParams,
serialization: thread_id(params.thread_id),
response: v2::ThreadRollbackResponse,
},
ThreadList => "thread/list" {
@@ -319,6 +438,7 @@ client_request_definitions! {
},
ThreadRead => "thread/read" {
params: v2::ThreadReadParams,
serialization: thread_id(params.thread_id),
response: v2::ThreadReadResponse,
},
SkillsList => "skills/list" {
@@ -367,56 +487,68 @@ client_request_definitions! {
},
FsWatch => "fs/watch" {
params: v2::FsWatchParams,
serialization: fs_watch_id(params.watch_id),
response: v2::FsWatchResponse,
},
FsUnwatch => "fs/unwatch" {
params: v2::FsUnwatchParams,
serialization: fs_watch_id(params.watch_id),
response: v2::FsUnwatchResponse,
},
SkillsConfigWrite => "skills/config/write" {
params: v2::SkillsConfigWriteParams,
serialization: global("skills-config"),
response: v2::SkillsConfigWriteResponse,
},
PluginInstall => "plugin/install" {
params: v2::PluginInstallParams,
serialization: plugin(params.marketplace_path, params.plugin_name),
response: v2::PluginInstallResponse,
},
PluginUninstall => "plugin/uninstall" {
params: v2::PluginUninstallParams,
serialization: plugin_id(params.plugin_id),
response: v2::PluginUninstallResponse,
},
TurnStart => "turn/start" {
params: v2::TurnStartParams,
inspect_params: true,
serialization: thread_id(params.thread_id),
response: v2::TurnStartResponse,
},
TurnSteer => "turn/steer" {
params: v2::TurnSteerParams,
inspect_params: true,
serialization: thread_id(params.thread_id),
response: v2::TurnSteerResponse,
},
TurnInterrupt => "turn/interrupt" {
params: v2::TurnInterruptParams,
serialization: thread_id(params.thread_id),
response: v2::TurnInterruptResponse,
},
#[experimental("thread/realtime/start")]
ThreadRealtimeStart => "thread/realtime/start" {
params: v2::ThreadRealtimeStartParams,
serialization: thread_id(params.thread_id),
response: v2::ThreadRealtimeStartResponse,
},
#[experimental("thread/realtime/appendAudio")]
ThreadRealtimeAppendAudio => "thread/realtime/appendAudio" {
params: v2::ThreadRealtimeAppendAudioParams,
serialization: thread_id(params.thread_id),
response: v2::ThreadRealtimeAppendAudioResponse,
},
#[experimental("thread/realtime/appendText")]
ThreadRealtimeAppendText => "thread/realtime/appendText" {
params: v2::ThreadRealtimeAppendTextParams,
serialization: thread_id(params.thread_id),
response: v2::ThreadRealtimeAppendTextResponse,
},
#[experimental("thread/realtime/stop")]
ThreadRealtimeStop => "thread/realtime/stop" {
params: v2::ThreadRealtimeStopParams,
serialization: thread_id(params.thread_id),
response: v2::ThreadRealtimeStopResponse,
},
#[experimental("thread/realtime/listVoices")]
@@ -426,6 +558,7 @@ client_request_definitions! {
},
ReviewStart => "review/start" {
params: v2::ReviewStartParams,
serialization: thread_id(params.thread_id),
response: v2::ReviewStartResponse,
},
@@ -439,6 +572,7 @@ client_request_definitions! {
},
ExperimentalFeatureEnablementSet => "experimentalFeature/enablement/set" {
params: v2::ExperimentalFeatureEnablementSetParams,
serialization: global("config"),
response: v2::ExperimentalFeatureEnablementSetResponse,
},
#[experimental("collaborationMode/list")]
@@ -456,42 +590,50 @@ client_request_definitions! {
McpServerOauthLogin => "mcpServer/oauth/login" {
params: v2::McpServerOauthLoginParams,
serialization: mcp_oauth_server(params.name),
response: v2::McpServerOauthLoginResponse,
},
McpServerRefresh => "config/mcpServer/reload" {
params: #[ts(type = "undefined")] #[serde(skip_serializing_if = "Option::is_none")] Option<()>,
serialization: global("mcp-registry"),
response: v2::McpServerRefreshResponse,
},
McpServerStatusList => "mcpServerStatus/list" {
params: v2::ListMcpServerStatusParams,
serialization: global("mcp-registry"),
response: v2::ListMcpServerStatusResponse,
},
McpResourceRead => "mcpServer/resource/read" {
params: v2::McpResourceReadParams,
serialization: thread_id(params.thread_id),
response: v2::McpResourceReadResponse,
},
WindowsSandboxSetupStart => "windowsSandbox/setupStart" {
params: v2::WindowsSandboxSetupStartParams,
serialization: global("windows-sandbox-setup"),
response: v2::WindowsSandboxSetupStartResponse,
},
LoginAccount => "account/login/start" {
params: v2::LoginAccountParams,
inspect_params: true,
serialization: global("account-auth"),
response: v2::LoginAccountResponse,
},
CancelLoginAccount => "account/login/cancel" {
params: v2::CancelLoginAccountParams,
serialization: global("account-auth"),
response: v2::CancelLoginAccountResponse,
},
LogoutAccount => "account/logout" {
params: #[ts(type = "undefined")] #[serde(skip_serializing_if = "Option::is_none")] Option<()>,
serialization: global("account-auth"),
response: v2::LogoutAccountResponse,
},
@@ -508,26 +650,31 @@ client_request_definitions! {
/// Execute a standalone command (argv vector) under the server's sandbox.
OneOffCommandExec => "command/exec" {
params: v2::CommandExecParams,
serialization: optional_command_process_id(params.process_id),
response: v2::CommandExecResponse,
},
/// Write stdin bytes to a running `command/exec` session or close stdin.
CommandExecWrite => "command/exec/write" {
params: v2::CommandExecWriteParams,
serialization: command_process_id(params.process_id),
response: v2::CommandExecWriteResponse,
},
/// Terminate a running `command/exec` session by client-supplied `processId`.
CommandExecTerminate => "command/exec/terminate" {
params: v2::CommandExecTerminateParams,
serialization: command_process_id(params.process_id),
response: v2::CommandExecTerminateResponse,
},
/// Resize a running PTY-backed `command/exec` session by client-supplied `processId`.
CommandExecResize => "command/exec/resize" {
params: v2::CommandExecResizeParams,
serialization: command_process_id(params.process_id),
response: v2::CommandExecResizeResponse,
},
ConfigRead => "config/read" {
params: v2::ConfigReadParams,
serialization: global("config"),
response: v2::ConfigReadResponse,
},
ExternalAgentConfigDetect => "externalAgentConfig/detect" {
@@ -536,24 +683,29 @@ client_request_definitions! {
},
ExternalAgentConfigImport => "externalAgentConfig/import" {
params: v2::ExternalAgentConfigImportParams,
serialization: global("external-agent-config-import"),
response: v2::ExternalAgentConfigImportResponse,
},
ConfigValueWrite => "config/value/write" {
params: v2::ConfigValueWriteParams,
serialization: global("config"),
response: v2::ConfigWriteResponse,
},
ConfigBatchWrite => "config/batchWrite" {
params: v2::ConfigBatchWriteParams,
serialization: global("config"),
response: v2::ConfigWriteResponse,
},
ConfigRequirementsRead => "configRequirements/read" {
params: #[ts(type = "undefined")] #[serde(skip_serializing_if = "Option::is_none")] Option<()>,
serialization: global("config"),
response: v2::ConfigRequirementsReadResponse,
},
GetAccount => "account/read" {
params: v2::GetAccountParams,
serialization: global("account-auth"),
response: v2::GetAccountResponse,
},
@@ -569,6 +721,7 @@ client_request_definitions! {
/// DEPRECATED in favor of GetAccount
GetAuthStatus {
params: v1::GetAuthStatusParams,
serialization: global("account-auth"),
response: v1::GetAuthStatusResponse,
},
FuzzyFileSearch {
@@ -578,16 +731,19 @@ client_request_definitions! {
#[experimental("fuzzyFileSearch/sessionStart")]
FuzzyFileSearchSessionStart => "fuzzyFileSearch/sessionStart" {
params: FuzzyFileSearchSessionStartParams,
serialization: fuzzy_session_id(params.session_id),
response: FuzzyFileSearchSessionStartResponse,
},
#[experimental("fuzzyFileSearch/sessionUpdate")]
FuzzyFileSearchSessionUpdate => "fuzzyFileSearch/sessionUpdate" {
params: FuzzyFileSearchSessionUpdateParams,
serialization: fuzzy_session_id(params.session_id),
response: FuzzyFileSearchSessionUpdateResponse,
},
#[experimental("fuzzyFileSearch/sessionStop")]
FuzzyFileSearchSessionStop => "fuzzyFileSearch/sessionStop" {
params: FuzzyFileSearchSessionStopParams,
serialization: fuzzy_session_id(params.session_id),
response: FuzzyFileSearchSessionStopResponse,
},
}
@@ -1064,6 +1220,213 @@ mod tests {
AbsolutePathBuf::from_absolute_path(absolute_path_string(path)).expect("absolute path")
}
fn request_id() -> RequestId {
RequestId::Integer(1)
}
#[test]
fn client_request_serialization_scope_covers_keyed_families() {
let thread_id = "thread-1".to_string();
let thread_resume = ClientRequest::ThreadResume {
request_id: request_id(),
params: v2::ThreadResumeParams {
thread_id: thread_id.clone(),
..Default::default()
},
};
assert_eq!(
thread_resume.serialization_scope(),
Some(ClientRequestSerializationScope::Thread {
thread_id: thread_id.clone()
})
);
let fork_path = PathBuf::from("/tmp/source-thread.jsonl");
let thread_fork = ClientRequest::ThreadFork {
request_id: request_id(),
params: v2::ThreadForkParams {
thread_id: thread_id,
path: Some(fork_path.clone()),
..Default::default()
},
};
assert_eq!(
thread_fork.serialization_scope(),
Some(ClientRequestSerializationScope::ThreadPath { path: fork_path })
);
let command_exec = ClientRequest::OneOffCommandExec {
request_id: request_id(),
params: v2::CommandExecParams {
command: vec!["sleep".to_string(), "10".to_string()],
process_id: Some("proc-1".to_string()),
tty: false,
stream_stdin: false,
stream_stdout_stderr: false,
output_bytes_cap: None,
disable_output_cap: false,
disable_timeout: false,
timeout_ms: None,
cwd: None,
env: None,
size: None,
sandbox_policy: None,
},
};
assert_eq!(
command_exec.serialization_scope(),
Some(ClientRequestSerializationScope::CommandExecProcess {
process_id: "proc-1".to_string()
})
);
let fuzzy_update = ClientRequest::FuzzyFileSearchSessionUpdate {
request_id: request_id(),
params: FuzzyFileSearchSessionUpdateParams {
session_id: "search-1".to_string(),
query: "lib".to_string(),
},
};
assert_eq!(
fuzzy_update.serialization_scope(),
Some(ClientRequestSerializationScope::FuzzyFileSearchSession {
session_id: "search-1".to_string()
})
);
let fs_watch = ClientRequest::FsWatch {
request_id: request_id(),
params: v2::FsWatchParams {
watch_id: "watch-1".to_string(),
path: absolute_path("/tmp/repo"),
},
};
assert_eq!(
fs_watch.serialization_scope(),
Some(ClientRequestSerializationScope::FsWatch {
watch_id: "watch-1".to_string()
})
);
let plugin_install = ClientRequest::PluginInstall {
request_id: request_id(),
params: v2::PluginInstallParams {
marketplace_path: absolute_path("/tmp/marketplace"),
plugin_name: "plugin-a".to_string(),
force_remote_sync: false,
},
};
assert_eq!(
plugin_install.serialization_scope(),
Some(ClientRequestSerializationScope::Plugin {
marketplace_path: absolute_path("/tmp/marketplace"),
plugin_name: "plugin-a".to_string(),
})
);
let plugin_uninstall = ClientRequest::PluginUninstall {
request_id: request_id(),
params: v2::PluginUninstallParams {
plugin_id: "plugin-a".to_string(),
force_remote_sync: false,
},
};
assert_eq!(
plugin_uninstall.serialization_scope(),
Some(ClientRequestSerializationScope::PluginId {
plugin_id: "plugin-a".to_string()
})
);
let mcp_oauth = ClientRequest::McpServerOauthLogin {
request_id: request_id(),
params: v2::McpServerOauthLoginParams {
name: "server-a".to_string(),
scopes: None,
timeout_secs: None,
},
};
assert_eq!(
mcp_oauth.serialization_scope(),
Some(ClientRequestSerializationScope::McpOauth {
server_name: "server-a".to_string()
})
);
let config_read = ClientRequest::ConfigRead {
request_id: request_id(),
params: v2::ConfigReadParams {
include_layers: false,
cwd: None,
},
};
assert_eq!(
config_read.serialization_scope(),
Some(ClientRequestSerializationScope::Global("config"))
);
let account_read = ClientRequest::GetAccount {
request_id: request_id(),
params: v2::GetAccountParams {
refresh_token: false,
},
};
assert_eq!(
account_read.serialization_scope(),
Some(ClientRequestSerializationScope::Global("account-auth"))
);
}
#[test]
fn client_request_serialization_scope_covers_unkeyed_representatives() {
let initialize = ClientRequest::Initialize {
request_id: request_id(),
params: v1::InitializeParams {
client_info: v1::ClientInfo {
name: "test".to_string(),
title: None,
version: "0.1.0".to_string(),
},
capabilities: None,
},
};
assert_eq!(initialize.serialization_scope(), None);
let thread_start = ClientRequest::ThreadStart {
request_id: request_id(),
params: v2::ThreadStartParams::default(),
};
assert_eq!(thread_start.serialization_scope(), None);
let command_exec = ClientRequest::OneOffCommandExec {
request_id: request_id(),
params: v2::CommandExecParams {
command: vec!["true".to_string()],
process_id: None,
tty: false,
stream_stdin: false,
stream_stdout_stderr: false,
output_bytes_cap: None,
disable_output_cap: false,
disable_timeout: false,
timeout_ms: None,
cwd: None,
env: None,
size: None,
sandbox_policy: None,
},
};
assert_eq!(command_exec.serialization_scope(), None);
let fs_read = ClientRequest::FsReadFile {
request_id: request_id(),
params: v2::FsReadFileParams {
path: absolute_path("/tmp/file.txt"),
},
};
assert_eq!(fs_read.serialization_scope(), None);
}
#[test]
fn serialize_get_conversation_summary() -> Result<()> {
let request = ClientRequest::GetConversationSummary {

View File

@@ -79,6 +79,7 @@ pub mod in_process;
mod message_processor;
mod models;
mod outgoing_message;
mod request_serialization;
mod server_request_error;
mod thread_state;
mod thread_status;

View File

@@ -19,6 +19,9 @@ use crate::outgoing_message::ConnectionId;
use crate::outgoing_message::ConnectionRequestId;
use crate::outgoing_message::OutgoingMessageSender;
use crate::outgoing_message::RequestContext;
use crate::request_serialization::QueuedInitializedRequest;
use crate::request_serialization::RequestSerializationQueueKey;
use crate::request_serialization::RequestSerializationQueues;
use crate::transport::AppServerTransport;
use crate::transport::RemoteControlHandle;
use async_trait::async_trait;
@@ -173,6 +176,7 @@ pub(crate) struct MessageProcessor {
config_warnings: Arc<Vec<ConfigWarningNotification>>,
rpc_transport: AppServerRpcTransport,
remote_control_handle: Option<RemoteControlHandle>,
request_serialization_queues: RequestSerializationQueues,
}
#[derive(Debug, Default)]
@@ -330,6 +334,7 @@ impl MessageProcessor {
config_warnings: Arc::new(config_warnings),
rpc_transport,
remote_control_handle,
request_serialization_queues: RequestSerializationQueues::default(),
}
}
@@ -745,11 +750,13 @@ impl MessageProcessor {
return;
}
let serialization_scope = codex_request.serialization_scope();
let app_server_client_name = session.app_server_client_name().map(str::to_string);
let client_version = session.client_version().map(str::to_string);
let connection_id = connection_request_id.connection_id;
let processor = Arc::clone(self);
let span = request_context.span();
tokio::spawn(
let request = QueuedInitializedRequest::new(
async move {
processor
.handle_initialized_client_request(
@@ -763,6 +770,17 @@ impl MessageProcessor {
}
.instrument(span),
);
if let Some(scope) = serialization_scope {
let key = RequestSerializationQueueKey::from_scope(connection_id, scope);
self.request_serialization_queues
.enqueue(key, request)
.await;
} else {
tokio::spawn(async move {
request.run().await;
});
}
}
async fn handle_initialized_client_request(

View File

@@ -0,0 +1,227 @@
use std::collections::HashMap;
use std::collections::VecDeque;
use std::future::Future;
use std::path::PathBuf;
use std::pin::Pin;
use std::sync::Arc;
use codex_app_server_protocol::ClientRequestSerializationScope;
use codex_utils_absolute_path::AbsolutePathBuf;
use tokio::sync::Mutex;
use tracing::Instrument;
use crate::outgoing_message::ConnectionId;
type BoxFutureUnit = Pin<Box<dyn Future<Output = ()> + Send + 'static>>;
#[derive(Clone, Debug, Eq, Hash, PartialEq)]
pub(crate) enum RequestSerializationQueueKey {
Global(&'static str),
Thread {
thread_id: String,
},
ThreadPath {
path: PathBuf,
},
CommandExecProcess {
connection_id: ConnectionId,
process_id: String,
},
FuzzyFileSearchSession {
session_id: String,
},
FsWatch {
connection_id: ConnectionId,
watch_id: String,
},
Plugin {
marketplace_path: AbsolutePathBuf,
plugin_name: String,
},
PluginId {
plugin_id: String,
},
McpOauth {
server_name: String,
},
}
impl RequestSerializationQueueKey {
pub(crate) fn from_scope(
connection_id: ConnectionId,
scope: ClientRequestSerializationScope,
) -> Self {
match scope {
ClientRequestSerializationScope::Global(name) => Self::Global(name),
ClientRequestSerializationScope::Thread { thread_id } => Self::Thread { thread_id },
ClientRequestSerializationScope::ThreadPath { path } => Self::ThreadPath { path },
ClientRequestSerializationScope::CommandExecProcess { process_id } => {
Self::CommandExecProcess {
connection_id,
process_id,
}
}
ClientRequestSerializationScope::FuzzyFileSearchSession { session_id } => {
Self::FuzzyFileSearchSession { session_id }
}
ClientRequestSerializationScope::FsWatch { watch_id } => Self::FsWatch {
connection_id,
watch_id,
},
ClientRequestSerializationScope::Plugin {
marketplace_path,
plugin_name,
} => Self::Plugin {
marketplace_path,
plugin_name,
},
ClientRequestSerializationScope::PluginId { plugin_id } => Self::PluginId { plugin_id },
ClientRequestSerializationScope::McpOauth { server_name } => {
Self::McpOauth { server_name }
}
}
}
}
pub(crate) struct QueuedInitializedRequest {
future: BoxFutureUnit,
}
impl QueuedInitializedRequest {
pub(crate) fn new(future: impl Future<Output = ()> + Send + 'static) -> Self {
Self {
future: Box::pin(future),
}
}
pub(crate) async fn run(self) {
self.future.await;
}
}
#[derive(Clone, Default)]
pub(crate) struct RequestSerializationQueues {
inner: Arc<Mutex<HashMap<RequestSerializationQueueKey, VecDeque<QueuedInitializedRequest>>>>,
}
impl RequestSerializationQueues {
pub(crate) async fn enqueue(
&self,
key: RequestSerializationQueueKey,
request: QueuedInitializedRequest,
) {
let should_spawn = {
let mut queues = self.inner.lock().await;
match queues.get_mut(&key) {
Some(queue) => {
queue.push_back(request);
false
}
None => {
let mut queue = VecDeque::new();
queue.push_back(request);
queues.insert(key.clone(), queue);
true
}
}
};
if should_spawn {
let queues = self.clone();
let span = tracing::debug_span!("app_server.serialized_request_queue", ?key);
tokio::spawn(async move { queues.drain(key).await }.instrument(span));
}
}
async fn drain(self, key: RequestSerializationQueueKey) {
loop {
let request = {
let mut queues = self.inner.lock().await;
let Some(queue) = queues.get_mut(&key) else {
return;
};
match queue.pop_front() {
Some(request) => request,
None => {
queues.remove(&key);
return;
}
}
};
request.run().await;
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use pretty_assertions::assert_eq;
use tokio::sync::mpsc;
use tokio::sync::oneshot;
use tokio::time::Duration;
use tokio::time::timeout;
#[tokio::test]
async fn same_key_requests_run_fifo() {
let queues = RequestSerializationQueues::default();
let key = RequestSerializationQueueKey::Global("test");
let (tx, mut rx) = mpsc::unbounded_channel();
for value in [1, 2, 3] {
let tx = tx.clone();
queues
.enqueue(
key.clone(),
QueuedInitializedRequest::new(async move {
tx.send(value).expect("receiver should be open");
}),
)
.await;
}
drop(tx);
let mut values = Vec::new();
while let Some(value) = timeout(Duration::from_secs(1), rx.recv())
.await
.expect("timed out waiting for queued request")
{
values.push(value);
}
assert_eq!(values, vec![1, 2, 3]);
}
#[tokio::test]
async fn different_keys_run_concurrently() {
let queues = RequestSerializationQueues::default();
let (blocked_tx, blocked_rx) = oneshot::channel::<()>();
let (ran_tx, ran_rx) = oneshot::channel::<()>();
queues
.enqueue(
RequestSerializationQueueKey::Global("blocked"),
QueuedInitializedRequest::new(async move {
let _ = blocked_rx.await;
}),
)
.await;
queues
.enqueue(
RequestSerializationQueueKey::Global("other"),
QueuedInitializedRequest::new(async move {
ran_tx.send(()).expect("receiver should be open");
}),
)
.await;
timeout(Duration::from_secs(1), ran_rx)
.await
.expect("other key should not be blocked")
.expect("sender should be open");
blocked_tx
.send(())
.expect("blocked request should be waiting");
}
}