diff --git a/rust/Cargo.lock b/rust/Cargo.lock index 8de679798..dc7d69348 100644 --- a/rust/Cargo.lock +++ b/rust/Cargo.lock @@ -452,6 +452,7 @@ dependencies = [ "tokio-tungstenite", "tokio-util", "tracing", + "tracing-subscriber", "ureq", "uuid", "zip", @@ -764,6 +765,12 @@ dependencies = [ "wasm-bindgen", ] +[[package]] +name = "lazy_static" +version = "1.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bbd2bcb4c963f2ddae06a2efc7e9f3591312473c50c6685e1f298068316e66fe" + [[package]] name = "leb128fmt" version = "0.1.0" @@ -880,6 +887,15 @@ dependencies = [ "tempfile", ] +[[package]] +name = "nu-ansi-term" +version = "0.50.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7957b9740744892f114936ab4a57b3f487491bbeafaf8083688b16841a4240e5" +dependencies = [ + "windows-sys 0.61.2", +] + [[package]] name = "once_cell" version = "1.21.4" @@ -1467,6 +1483,15 @@ dependencies = [ "digest", ] +[[package]] +name = "sharded-slab" +version = "0.1.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f40ca3c46823713e0d4209592e8d6e826aa57e928f09752619fc696c499637f6" +dependencies = [ + "lazy_static", +] + [[package]] name = "shlex" version = "1.3.0" @@ -1618,6 +1643,15 @@ dependencies = [ "syn", ] +[[package]] +name = "thread_local" +version = "1.1.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1ad99c4c6d32803332c548b1af0540b357b3f5fc0be8f6c6bfe8b2e6ae784070" +dependencies = [ + "cfg-if", +] + [[package]] name = "tinystr" version = "0.8.3" @@ -1788,6 +1822,32 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "db97caf9d906fbde555dd62fa95ddba9eecfd14cb388e4f491a66d74cd5fb79a" dependencies = [ "once_cell", + "valuable", +] + +[[package]] +name = "tracing-log" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ee855f1f400bd0e5c02d150ae5de3840039a3f54b025156404e34c23c03f47c3" +dependencies = [ + "log", + "once_cell", + "tracing-core", +] + +[[package]] +name = "tracing-subscriber" +version = "0.3.23" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cb7f578e5945fb242538965c2d0b04418d38ec25c79d160cd279bf0731c8d319" +dependencies = [ + "nu-ansi-term", + "sharded-slab", + "smallvec", + "thread_local", + "tracing-core", + "tracing-log", ] [[package]] @@ -1885,6 +1945,12 @@ dependencies = [ "getrandom 0.4.2", ] +[[package]] +name = "valuable" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ba73ea9cf16a25df0c8caa16c51acb937d5712a8429db78a3ee29d5dcacd3a65" + [[package]] name = "vcpkg" version = "0.2.15" diff --git a/rust/Cargo.toml b/rust/Cargo.toml index 0f18a9b15..b68662537 100644 --- a/rust/Cargo.toml +++ b/rust/Cargo.toml @@ -76,7 +76,8 @@ rusqlite = { version = "0.35", features = ["bundled"] } schemars = "1" serial_test = "3" tempfile = "3" -tokio = { version = "1", features = ["rt-multi-thread"] } +tokio = { version = "1", features = ["rt-multi-thread", "test-util"] } +tracing-subscriber = { version = "0.3", features = ["fmt"] } # Integration tests that call test-support-only Client methods (e.g. # `from_streams_with_connection_token`, `from_streams_with_trace_provider`) diff --git a/rust/src/copilot_request_handler.rs b/rust/src/copilot_request_handler.rs index 961ae3876..df1cdcd90 100644 --- a/rust/src/copilot_request_handler.rs +++ b/rust/src/copilot_request_handler.rs @@ -1047,7 +1047,17 @@ impl CopilotRequestDispatcher { self.client.get().cloned().unwrap_or_else(Weak::new) } - pub(crate) async fn dispatch(self: &Arc, request: JsonRpcRequest) { + pub(crate) async fn dispatch(self: &Arc, request: crate::jsonrpc::ReverseRpcRequest) { + let (request, reverse_rpc) = request.into_dispatch(); + let dispatch = self.dispatch_inner(request); + if let Some(trace) = &reverse_rpc { + trace.scope(dispatch).await; + } else { + dispatch.await; + } + } + + async fn dispatch_inner(self: &Arc, request: JsonRpcRequest) { match request.method.as_str() { METHOD_HTTP_REQUEST_START => self.handle_start(request).await, METHOD_HTTP_REQUEST_CHUNK => self.handle_chunk(request).await, diff --git a/rust/src/github_token.rs b/rust/src/github_token.rs index c5eaa63ad..07a0ac586 100644 --- a/rust/src/github_token.rs +++ b/rust/src/github_token.rs @@ -188,11 +188,21 @@ impl GitHubTokenRegistry { state.session_owners.clear(); } - pub(crate) async fn dispatch(&self, request: JsonRpcRequest) { + pub(crate) async fn dispatch(&self, request: crate::jsonrpc::ReverseRpcRequest) { let Some(inner) = self.client.get().and_then(Weak::upgrade) else { return; }; let client = Client::from_inner(inner); + let (request, reverse_rpc) = request.into_dispatch(); + let dispatch = self.dispatch_inner(&client, request); + if let Some(trace) = &reverse_rpc { + trace.scope(dispatch).await; + } else { + dispatch.await; + } + } + + async fn dispatch_inner(&self, client: &Client, request: JsonRpcRequest) { let params = request .params .clone() @@ -201,7 +211,7 @@ impl GitHubTokenRegistry { Ok(params) => params, Err(error) => { send_error( - &client, + client, request.id, error_codes::INVALID_PARAMS, &format!("invalid params: {error}"), @@ -218,7 +228,7 @@ impl GitHubTokenRegistry { .cloned(); let Some(provider) = provider else { send_error( - &client, + client, request.id, error_codes::INTERNAL_ERROR, "unknown GitHub token provider registration", @@ -232,7 +242,7 @@ impl GitHubTokenRegistry { GitHubTokenAcquireReason::Refresh => GitHubTokenRequestReason::Refresh, GitHubTokenAcquireReason::Unknown => { send_error( - &client, + client, request.id, error_codes::INVALID_PARAMS, "unknown GitHub token acquisition reason", @@ -252,7 +262,7 @@ impl GitHubTokenRegistry { { Ok(GitHubTokenProviderResult::Token(token)) => { respond( - &client, + client, request.id, GitHubTokenAcquireResult::Token(token.into_wire()), ) @@ -260,7 +270,7 @@ impl GitHubTokenRegistry { } Ok(GitHubTokenProviderResult::Cancelled) => { respond( - &client, + client, request.id, GitHubTokenAcquireResult::Cancelled(GitHubTokenAcquireResultCancelled { kind: Default::default(), @@ -270,7 +280,7 @@ impl GitHubTokenRegistry { } Err(error) => { send_error( - &client, + client, request.id, error_codes::INTERNAL_ERROR, &format!("GitHub token provider failed: {error}"), diff --git a/rust/src/hooks.rs b/rust/src/hooks.rs index 4986d6cb1..2574ad6dc 100644 --- a/rust/src/hooks.rs +++ b/rust/src/hooks.rs @@ -6,11 +6,11 @@ //! [`Client::create_session`](crate::Client::create_session). use std::path::PathBuf; -use std::time::Instant; use async_trait::async_trait; use serde::{Deserialize, Serialize}; use serde_json::Value; +use tokio::time::Instant; use crate::types::SessionId; @@ -680,11 +680,22 @@ pub trait SessionHooks: Send + Sync + 'static { /// Returns `Ok(Value)` shaped like `{ "output": ... }` on success. /// If no hook is registered ([`HookOutput::None`]), the output is an empty /// object: `{ "output": {} }`. +#[cfg(test)] pub(crate) async fn dispatch_hook( hooks: &dyn SessionHooks, session_id: &SessionId, hook_type: &str, raw_input: Value, +) -> Result { + dispatch_hook_traced(hooks, session_id, hook_type, raw_input, None).await +} + +pub(crate) async fn dispatch_hook_traced( + hooks: &dyn SessionHooks, + session_id: &SessionId, + hook_type: &str, + raw_input: Value, + reverse_rpc_trace: Option<&crate::jsonrpc::ReverseRpcTrace>, ) -> Result { let ctx = HookContext { session_id: session_id.clone(), @@ -743,8 +754,12 @@ pub(crate) async fn dispatch_hook( let dispatch_start = Instant::now(); let output = hooks.on_hook(event).await; + let dispatch_elapsed = dispatch_start.elapsed(); + if let Some(trace) = reverse_rpc_trace { + trace.record_hook_callback(hook_type, dispatch_start, dispatch_elapsed); + } tracing::debug!( - elapsed_ms = dispatch_start.elapsed().as_millis(), + elapsed_ms = dispatch_elapsed.as_millis(), session_id = %session_id, hook_type = hook_type, "SessionHooks::on_hook dispatch" @@ -786,7 +801,58 @@ pub(crate) async fn dispatch_hook( #[cfg(test)] mod tests { + use std::io::Write; + use std::sync::Arc; + use std::time::Duration; + + use parking_lot::Mutex; + use tokio::sync::Notify; + use tracing::Instrument; + use tracing_subscriber::Layer; + use tracing_subscriber::fmt::MakeWriter; + use tracing_subscriber::layer::SubscriberExt; + use super::*; + use crate::JsonRpcRequest; + use crate::jsonrpc::ReverseRpcTrace; + + #[derive(Clone, Default)] + struct TraceBuffer(Arc>>); + + impl TraceBuffer { + fn text(&self) -> String { + String::from_utf8(self.0.lock().clone()).unwrap() + } + } + + impl Write for TraceBuffer { + fn write(&mut self, buf: &[u8]) -> std::io::Result { + self.0.lock().extend_from_slice(buf); + Ok(buf.len()) + } + + fn flush(&mut self) -> std::io::Result<()> { + Ok(()) + } + } + + impl<'a> MakeWriter<'a> for TraceBuffer { + type Writer = Self; + + fn make_writer(&'a self) -> Self::Writer { + self.clone() + } + } + + async fn wait_for_trace(buffer: &TraceBuffer, needle: &str) { + for _ in 0..20 { + if buffer.text().contains(needle) { + return; + } + tokio::task::yield_now().await; + } + panic!("timing trace did not contain {needle:?}: {}", buffer.text()); + } struct TestHooks; @@ -842,6 +908,86 @@ mod tests { assert_eq!(output["permissionDecisionReason"], "blocked by policy"); } + #[tokio::test(start_paused = true)] + async fn traced_dispatch_measures_only_the_gated_hook_callback() { + const SENTINEL: &str = "PRIVATE_HOOK_SENTINEL_DO_NOT_TRACE"; + + struct GatedHooks { + started: Arc, + release: Arc, + } + + #[async_trait] + impl SessionHooks for GatedHooks { + async fn on_hook(&self, _event: HookEvent) -> HookOutput { + self.started.notify_one(); + self.release.notified().await; + HookOutput::UserPromptSubmitted(UserPromptSubmittedOutput { + modified_prompt: Some(SENTINEL.to_string()), + ..Default::default() + }) + } + } + + let trace_buffer = TraceBuffer::default(); + let subscriber = tracing_subscriber::registry().with( + tracing_subscriber::fmt::layer() + .with_writer(trace_buffer.clone()) + .with_ansi(false) + .without_time() + .with_filter(tracing_subscriber::filter::filter_fn(|metadata| { + metadata.target() == "github_copilot_sdk::reverse_rpc_timing" + })), + ); + let _subscriber = tracing::subscriber::set_default(subscriber); + let started = Arc::new(Notify::new()); + let release = Arc::new(Notify::new()); + let hooks = Arc::new(GatedHooks { + started: started.clone(), + release: release.clone(), + }); + let request = JsonRpcRequest::new( + 99, + "hooks.invoke", + Some(serde_json::json!({ "sessionId": "session-1" })), + ); + let now = Instant::now(); + let trace = ReverseRpcTrace::for_test(&request, now, now); + + let parent = tracing::error_span!("session_request_handler", session_id = SENTINEL); + let dispatch = tokio::spawn( + async move { + dispatch_hook_traced( + hooks.as_ref(), + &SessionId::new("session-1"), + "userPromptSubmitted", + serde_json::json!({ + "sessionId": SENTINEL, + "timestamp": 1234567890, + "cwd": SENTINEL, + "prompt": SENTINEL + }), + Some(&trace), + ) + .await + } + .instrument(parent), + ); + + started.notified().await; + tokio::time::advance(Duration::from_millis(9)).await; + release.notify_one(); + let output = dispatch.await.unwrap().unwrap(); + assert_eq!(output["output"]["modifiedPrompt"], SENTINEL); + + wait_for_trace(&trace_buffer, "phase=\"hook_callback\"").await; + let traces = trace_buffer.text(); + assert!(traces.contains("phase=\"hook_callback\"")); + assert!(traces.contains("start_offset_us=0")); + assert!(traces.contains("elapsed_us=9000")); + assert!(!traces.contains(SENTINEL)); + } + #[tokio::test] async fn dispatch_pre_tool_use_passthrough() { let hooks = TestHooks; diff --git a/rust/src/jsonrpc.rs b/rust/src/jsonrpc.rs index 25a405080..1dc4e9368 100644 --- a/rust/src/jsonrpc.rs +++ b/rust/src/jsonrpc.rs @@ -1,4 +1,6 @@ -use std::collections::HashMap; +use std::collections::{HashMap, VecDeque}; +use std::future::Future; +use std::hash::{BuildHasher, Hash, Hasher}; use std::sync::Arc; use std::sync::atomic::{AtomicU64, Ordering}; use std::time::Instant; @@ -9,6 +11,8 @@ use serde_json::Value; use tokio::io::{AsyncBufReadExt, AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt, BufReader}; use tokio::sync::{broadcast, mpsc, oneshot}; use tokio::task::JoinHandle; +use tokio::time::Instant as TokioInstant; +use tokio_util::sync::CancellationToken; use tracing::{Instrument, debug, error, warn}; use crate::{Error, ErrorKind, ProtocolErrorKind}; @@ -168,6 +172,20 @@ impl JsonRpcResponse { } const CONTENT_LENGTH_HEADER: &str = "Content-Length: "; +/// Opt-in target for content-free reverse-RPC latency events. +/// +/// `request_forward` measures parsed-request receipt to forwarding, +/// `request_schedule` forwarding to dispatch start (`since_receive_us` is +/// the combined interval), `response_encode` is `serde_json::to_vec`, +/// `writer_queue` is enqueue to dequeue, and `write_all` / `flush` measure +/// the corresponding `AsyncWrite` calls. `hook_callback` is emitted by the +/// hooks dispatcher around `SessionHooks::on_hook` only. Every request phase +/// includes its start offset from request receipt, and `request_complete` +/// records total elapsed time. Collection is enabled only when this target +/// has an active DEBUG subscriber when the client is constructed. +const REVERSE_RPC_TIMING_TARGET: &str = "github_copilot_sdk::reverse_rpc_timing"; +const REVERSE_RPC_TIMING_CAPACITY: usize = 256; +type CorrelationHasher = std::collections::hash_map::RandomState; /// Rewrites unpaired UTF-16 surrogate escapes to `\uFFFD`. /// @@ -242,6 +260,630 @@ fn repair_lone_surrogates(body: &[u8]) -> Option> { struct WriteCommand { frame: Vec, ack: oneshot::Sender>, + reverse_rpc: Option, + enqueued_at: Option, +} + +struct ReverseRpcWriteTrace { + trace: ReverseRpcTrace, + registry: Option, + completion_attempted: bool, +} + +struct ReverseRpcResponseTrace { + request_id: u64, + trace: ReverseRpcTrace, + registry: ReverseRpcRegistry, +} + +// Response helpers retain their existing signatures while the dispatch task +// carries the exact request generation that produced the response. +tokio::task_local! { + static REVERSE_RPC_RESPONSE_TRACE: ReverseRpcResponseTrace; +} + +impl ReverseRpcWriteTrace { + fn new(trace: ReverseRpcTrace, registry: Option) -> Self { + Self { + trace, + registry, + completion_attempted: false, + } + } + + fn record_complete(&mut self, completed_at: TokioInstant, succeeded: bool) { + self.completion_attempted = true; + if let Some(registry) = &self.registry { + registry.complete(&self.trace, completed_at, succeeded); + } else { + let _ = self.trace.record_complete(completed_at, succeeded); + } + } +} + +impl Drop for ReverseRpcWriteTrace { + fn drop(&mut self) { + if !self.completion_attempted { + if let Some(registry) = &self.registry { + registry.complete(&self.trace, TokioInstant::now(), false); + } else { + let _ = self.trace.record_complete(TokioInstant::now(), false); + } + } + } +} + +enum ReverseRpcTimingEvent { + Phase { + trace: ReverseRpcTrace, + phase: &'static str, + start_offset_us: u64, + elapsed_us: u64, + succeeded: bool, + }, + Scheduled { + trace: ReverseRpcTrace, + start_offset_us: u64, + elapsed_us: u64, + since_receive_us: u64, + }, + HookCallback { + trace: ReverseRpcTrace, + hook_type: &'static str, + start_offset_us: u64, + elapsed_us: u64, + }, + Complete { + trace: ReverseRpcTrace, + required_phase_records: u64, + elapsed_us: u64, + succeeded: bool, + }, +} + +#[derive(Clone)] +struct ReverseRpcTimingEmitter { + phase_tx: mpsc::Sender, + terminal_tx: mpsc::Sender, + dropped_records: Arc, +} + +impl ReverseRpcTimingEmitter { + fn emit(&self, event: ReverseRpcTimingEvent) -> bool { + // Timing is measurement-only: a full bounded queue drops one logical + // record rather than delaying RPC work, including for terminal records. + // `records_dropped` is the explicit signal that the trace is incomplete. + let result = if matches!(&event, ReverseRpcTimingEvent::Complete { .. }) { + self.terminal_tx.try_send(event) + } else { + self.phase_tx.try_send(event) + }; + match result { + Ok(()) => true, + Err(mpsc::error::TrySendError::Full(_)) => { + let _ = self.dropped_records.fetch_update( + Ordering::Relaxed, + Ordering::Relaxed, + |count| Some(count.saturating_add(1)), + ); + false + } + Err(mpsc::error::TrySendError::Closed(_)) => false, + } + } +} + +/// Internal, content-free timing context for one inbound JSON-RPC request. +/// +/// The correlation key is a deterministic digest of the numeric wire ID, +/// RPC method, and optional session ID. The original session ID and params +/// are not retained. +#[derive(Clone)] +pub(crate) struct ReverseRpcTrace { + inner: Arc, +} + +struct ReverseRpcTraceInner { + correlation_key: String, + method: &'static str, + received_at: TokioInstant, + forwarded_at: std::sync::OnceLock, + timing: ReverseRpcTimingEmitter, + timing_state: Mutex, + emitted_phase_records: AtomicU64, +} + +#[derive(Clone, Copy)] +struct PendingReverseRpcCompletion { + required_phase_records: u64, + elapsed_us: u64, + succeeded: bool, +} + +#[derive(Default)] +struct ReverseRpcTimingState { + accepted_phase_records: u64, + pending_completion: Option, + completion_recorded: bool, + writer_owns_completion: bool, +} + +impl ReverseRpcTrace { + fn new( + request: &JsonRpcRequest, + generation: u64, + received_at: TokioInstant, + correlation_hasher: &CorrelationHasher, + timing: ReverseRpcTimingEmitter, + ) -> Self { + let session_id = request + .params + .as_ref() + .and_then(|params| params.get("sessionId")) + .and_then(Value::as_str); + Self { + inner: Arc::new(ReverseRpcTraceInner { + correlation_key: Self::correlation_key( + correlation_hasher, + request.id, + &request.method, + session_id, + generation, + ), + method: Self::timing_method(&request.method), + received_at, + forwarded_at: std::sync::OnceLock::new(), + timing, + timing_state: Mutex::new(ReverseRpcTimingState::default()), + emitted_phase_records: AtomicU64::new(0), + }), + } + } + + #[cfg(test)] + pub(crate) fn for_test( + request: &JsonRpcRequest, + received_at: TokioInstant, + forwarded_at: TokioInstant, + ) -> Self { + let (phase_tx, phase_rx) = mpsc::channel(REVERSE_RPC_TIMING_CAPACITY); + let (terminal_tx, terminal_rx) = mpsc::channel(REVERSE_RPC_TIMING_CAPACITY); + let dropped_records = Arc::new(AtomicU64::new(0)); + tokio::spawn(JsonRpcClient::timing_loop( + phase_rx, + terminal_rx, + dropped_records.clone(), + CancellationToken::new(), + )); + let trace = Self::new( + request, + 0, + received_at, + &CorrelationHasher::new(), + ReverseRpcTimingEmitter { + phase_tx, + terminal_tx, + dropped_records, + }, + ); + trace.mark_forwarding(forwarded_at); + trace + } + + fn correlation_key( + correlation_hasher: &CorrelationHasher, + request_id: u64, + method: &str, + session_id: Option<&str>, + generation: u64, + ) -> String { + // A per-client keyed hash keeps the request-derived key stable for all + // phases without making custom session IDs guessable from trace output. + let mut hasher = correlation_hasher.build_hasher(); + session_id.unwrap_or("").hash(&mut hasher); + method.hash(&mut hasher); + request_id.hash(&mut hasher); + generation.hash(&mut hasher); + format!("rrpc-{:016x}", hasher.finish()) + } + + fn timing_method(method: &str) -> &'static str { + match method { + "hooks.invoke" => "hooks.invoke", + "userInput.request" => "userInput.request", + "exitPlanMode.request" => "exitPlanMode.request", + "autoModeSwitch.request" => "autoModeSwitch.request", + "systemMessage.transform" => "systemMessage.transform", + "gitHubToken.getToken" => "gitHubToken.getToken", + "providerToken.getToken" => "providerToken.getToken", + _ if method.starts_with("sessionFs.") => "sessionFs.*", + _ if method.starts_with("canvas.") => "canvas.*", + _ if method.starts_with("llmInference.") => "llmInference.*", + _ => "unknown", + } + } + + fn timing_hook_type(hook_type: &str) -> &'static str { + match hook_type { + "preToolUse" => "preToolUse", + "preMcpToolCall" => "preMcpToolCall", + "postToolUse" => "postToolUse", + "postToolUseFailure" => "postToolUseFailure", + "userPromptSubmitted" => "userPromptSubmitted", + "userPromptTransformed" => "userPromptTransformed", + "sessionStart" => "sessionStart", + "sessionEnd" => "sessionEnd", + "errorOccurred" => "errorOccurred", + "agentStop" => "agentStop", + _ => "unknown", + } + } + + fn elapsed_us(duration: std::time::Duration) -> u64 { + u64::try_from(duration.as_micros()).unwrap_or(u64::MAX) + } + + fn start_offset_us(&self, started_at: TokioInstant) -> u64 { + Self::elapsed_us(started_at.duration_since(self.inner.received_at)) + } + + fn mark_forwarding(&self, forwarded_at: TokioInstant) { + self.inner + .forwarded_at + .set(forwarded_at) + .expect("forwarding timestamp must be recorded exactly once"); + } + + fn forward(&self, send: impl FnOnce() -> Result<(), T>) -> Result<(), T> { + let forwarded_at = TokioInstant::now(); + self.mark_forwarding(forwarded_at); + let mut state = self.inner.timing_state.lock(); + let result = send(); + self.emit_phase_with_state( + &mut state, + ReverseRpcTimingEvent::Phase { + trace: self.clone(), + phase: "request_forward", + start_offset_us: 0, + elapsed_us: Self::elapsed_us(forwarded_at.duration_since(self.inner.received_at)), + succeeded: result.is_ok(), + }, + ); + result + } + + fn record_scheduled(&self, scheduled_at: TokioInstant) { + let forwarded_at = self + .inner + .forwarded_at + .get() + .expect("forwarding timestamp must be set before scheduling"); + self.emit_phase(ReverseRpcTimingEvent::Scheduled { + trace: self.clone(), + start_offset_us: self.start_offset_us(*forwarded_at), + elapsed_us: Self::elapsed_us(scheduled_at.duration_since(*forwarded_at)), + since_receive_us: Self::elapsed_us(scheduled_at.duration_since(self.inner.received_at)), + }); + } + + pub(crate) fn record_hook_callback( + &self, + hook_type: &str, + started_at: TokioInstant, + elapsed: std::time::Duration, + ) { + self.emit_phase(ReverseRpcTimingEvent::HookCallback { + trace: self.clone(), + hook_type: Self::timing_hook_type(hook_type), + start_offset_us: self.start_offset_us(started_at), + elapsed_us: Self::elapsed_us(elapsed), + }); + } + + fn record_phase( + &self, + phase: &'static str, + started_at: TokioInstant, + elapsed: std::time::Duration, + succeeded: bool, + ) { + self.emit_phase(ReverseRpcTimingEvent::Phase { + trace: self.clone(), + phase, + start_offset_us: self.start_offset_us(started_at), + elapsed_us: Self::elapsed_us(elapsed), + succeeded, + }); + } + + fn emit_phase(&self, event: ReverseRpcTimingEvent) { + let mut state = self.inner.timing_state.lock(); + self.emit_phase_with_state(&mut state, event); + } + + fn emit_phase_with_state( + &self, + state: &mut ReverseRpcTimingState, + event: ReverseRpcTimingEvent, + ) { + if state.pending_completion.is_some() || state.completion_recorded { + return; + } + if self.inner.timing.emit(event) { + state.accepted_phase_records = state.accepted_phase_records.saturating_add(1); + } + } + + fn record_complete(&self, completed_at: TokioInstant, succeeded: bool) -> bool { + let mut state = self.inner.timing_state.lock(); + self.record_complete_with_state(&mut state, completed_at, succeeded) + } + + fn record_abandoned(&self, completed_at: TokioInstant) -> bool { + let mut state = self.inner.timing_state.lock(); + if state.writer_owns_completion { + return false; + } + let _ = self.record_complete_with_state(&mut state, completed_at, false); + true + } + + fn record_force_abandoned(&self, completed_at: TokioInstant) { + let _ = self.record_complete(completed_at, false); + } + + fn record_complete_with_state( + &self, + state: &mut ReverseRpcTimingState, + completed_at: TokioInstant, + succeeded: bool, + ) -> bool { + if state.completion_recorded { + return true; + } + if state.pending_completion.is_none() { + state.pending_completion = Some(PendingReverseRpcCompletion { + required_phase_records: state.accepted_phase_records, + elapsed_us: Self::elapsed_us(completed_at.duration_since(self.inner.received_at)), + succeeded, + }); + } + let completion = state + .pending_completion + .expect("pending completion must be initialized"); + if self.inner.timing.emit(ReverseRpcTimingEvent::Complete { + trace: self.clone(), + required_phase_records: completion.required_phase_records, + elapsed_us: completion.elapsed_us, + succeeded: completion.succeeded, + }) { + state.pending_completion = None; + state.completion_recorded = true; + true + } else { + false + } + } + + fn transfer_completion_to_writer(&self) { + self.inner.timing_state.lock().writer_owns_completion = true; + } +} + +#[derive(Clone)] +struct ReverseRpcRegistry(Arc); + +struct ReverseRpcRegistryInner { + state: Mutex, + next_generation: AtomicU64, +} + +struct ReverseRpcRegistryState { + traces: Vec, + force_closed: bool, +} + +impl ReverseRpcRegistry { + fn new() -> Self { + Self(Arc::new(ReverseRpcRegistryInner { + state: Mutex::new(ReverseRpcRegistryState { + traces: Vec::new(), + force_closed: false, + }), + next_generation: AtomicU64::new(1), + })) + } + + fn register( + &self, + request: &JsonRpcRequest, + received_at: TokioInstant, + correlation_hasher: &CorrelationHasher, + timing: ReverseRpcTimingEmitter, + ) -> Option { + let mut state = self.0.state.lock(); + if state.force_closed { + return None; + } + let trace = ReverseRpcTrace::new( + request, + self.0.next_generation.fetch_add(1, Ordering::Relaxed), + received_at, + correlation_hasher, + timing, + ); + state.traces.push(trace.clone()); + Some(trace) + } + + #[cfg(test)] + fn insert(&self, trace: ReverseRpcTrace) { + let mut state = self.0.state.lock(); + assert!(!state.force_closed); + state.traces.push(trace); + } + + fn abandon(&self, trace: &ReverseRpcTrace) { + let mut state = self.0.state.lock(); + let Some(index) = state + .traces + .iter() + .position(|current| Arc::ptr_eq(¤t.inner, &trace.inner)) + else { + return; + }; + if trace.record_abandoned(TokioInstant::now()) { + state.traces.swap_remove(index); + } + } + + fn complete(&self, trace: &ReverseRpcTrace, completed_at: TokioInstant, succeeded: bool) { + let mut state = self.0.state.lock(); + let Some(index) = state + .traces + .iter() + .position(|current| Arc::ptr_eq(¤t.inner, &trace.inner)) + else { + drop(state); + let _ = trace.record_complete(completed_at, succeeded); + return; + }; + let _ = trace.record_complete(completed_at, succeeded); + state.traces.swap_remove(index); + } + + fn abandon_all(&self) { + let mut state = self.0.state.lock(); + let mut index = 0; + while index < state.traces.len() { + let trace = state.traces[index].clone(); + if trace.record_abandoned(TokioInstant::now()) { + state.traces.swap_remove(index); + } else { + index += 1; + } + } + } + + fn force_abandon_all(&self) { + let mut state = self.0.state.lock(); + state.force_closed = true; + for trace in &state.traces { + trace.record_force_abandoned(TokioInstant::now()); + } + state.traces.clear(); + } + + #[cfg(test)] + fn contains(&self, trace: &ReverseRpcTrace) -> bool { + self.0 + .state + .lock() + .traces + .iter() + .any(|current| Arc::ptr_eq(¤t.inner, &trace.inner)) + } + + #[cfg(test)] + fn is_empty(&self) -> bool { + self.0.state.lock().traces.is_empty() + } +} + +pub(crate) struct ReverseRpcRequest { + request: Option, + trace: Option, + registry: Option, +} + +impl std::ops::Deref for ReverseRpcRequest { + type Target = JsonRpcRequest; + + fn deref(&self) -> &Self::Target { + self.request + .as_ref() + .expect("reverse RPC request must exist until dispatch") + } +} + +impl ReverseRpcRequest { + fn new( + request: JsonRpcRequest, + trace: Option, + registry: Option, + ) -> Self { + Self { + request: Some(request), + trace, + registry, + } + } + + pub(crate) fn into_dispatch(mut self) -> (JsonRpcRequest, Option) { + let request = self + .request + .take() + .expect("reverse RPC request must exist until dispatch"); + let guard = self.trace.take().map(|trace| { + trace.record_scheduled(TokioInstant::now()); + ReverseRpcDispatchGuard { + registry: self + .registry + .take() + .expect("timed reverse RPC request must have a registry"), + request_id: request.id, + trace, + } + }); + (request, guard) + } +} + +impl Drop for ReverseRpcRequest { + fn drop(&mut self) { + if let Some(trace) = self.trace.take() + && let Some(registry) = &self.registry + { + registry.abandon(&trace); + } + } +} + +pub(crate) struct ReverseRpcDispatchGuard { + registry: ReverseRpcRegistry, + request_id: u64, + trace: ReverseRpcTrace, +} + +impl ReverseRpcDispatchGuard { + pub(crate) fn trace(&self) -> &ReverseRpcTrace { + &self.trace + } + + pub(crate) async fn scope(&self, future: F) -> F::Output { + REVERSE_RPC_RESPONSE_TRACE + .scope( + ReverseRpcResponseTrace { + request_id: self.request_id, + trace: self.trace.clone(), + registry: self.registry.clone(), + }, + future, + ) + .await + } +} + +impl Drop for ReverseRpcDispatchGuard { + fn drop(&mut self) { + self.registry.abandon(&self.trace); + } +} + +#[derive(Clone)] +enum ReverseRequestSender { + Public(mpsc::UnboundedSender), + Internal(mpsc::UnboundedSender), } /// Low-level JSON-RPC 2.0 client over Content-Length-framed streams. @@ -264,10 +906,13 @@ pub struct JsonRpcClient { /// natural request/response back-pressure of the wire. write_tx: mpsc::UnboundedSender, pending_requests: Arc>>, + reverse_requests: Option, notification_tx: broadcast::Sender, - request_tx: mpsc::UnboundedSender, + request_tx: ReverseRequestSender, read_task: Mutex>>, write_task: Mutex>>, + timing_task: Mutex>>, + timing_shutdown: Option, } impl JsonRpcClient { @@ -277,28 +922,95 @@ impl JsonRpcClient { /// messages to pending request channels, the notification broadcast, /// or the request-forwarding channel; and a writer actor that owns the /// underlying `AsyncWrite` and serializes frames atomically. + #[cfg_attr( + not(any(test, feature = "test-support")), + expect( + dead_code, + reason = "low-level constructor is exported only with test-support" + ) + )] pub fn new( writer: impl AsyncWrite + Unpin + Send + 'static, reader: impl AsyncRead + Unpin + Send + 'static, notification_tx: broadcast::Sender, request_tx: mpsc::UnboundedSender, + ) -> Self { + Self::new_inner( + writer, + reader, + notification_tx, + ReverseRequestSender::Public(request_tx), + false, + ) + } + + pub(crate) fn new_with_reverse_rpc_timing( + writer: impl AsyncWrite + Unpin + Send + 'static, + reader: impl AsyncRead + Unpin + Send + 'static, + notification_tx: broadcast::Sender, + request_tx: mpsc::UnboundedSender, + ) -> Self { + let trace_reverse_rpc = tracing::enabled!( + target: REVERSE_RPC_TIMING_TARGET, + tracing::Level::DEBUG + ); + Self::new_inner( + writer, + reader, + notification_tx, + ReverseRequestSender::Internal(request_tx), + trace_reverse_rpc, + ) + } + + fn new_inner( + writer: impl AsyncWrite + Unpin + Send + 'static, + reader: impl AsyncRead + Unpin + Send + 'static, + notification_tx: broadcast::Sender, + request_tx: ReverseRequestSender, + trace_reverse_rpc: bool, ) -> Self { let (write_tx, write_rx) = mpsc::unbounded_channel::(); let writer_span = tracing::error_span!("jsonrpc_write_loop"); let write_task = tokio::spawn(Self::write_loop(writer, write_rx).instrument(writer_span)); + let (timing, timing_task, timing_shutdown, correlation_hasher, reverse_requests) = + if trace_reverse_rpc { + match Self::start_reverse_rpc_timing_with(Self::spawn_timing_thread) { + Some(( + timing, + timing_task, + timing_shutdown, + correlation_hasher, + reverse_requests, + )) => ( + Some(timing), + Some(timing_task), + Some(timing_shutdown), + Some(correlation_hasher), + Some(reverse_requests), + ), + None => (None, None, None, None, None), + } + } else { + (None, None, None, None, None) + }; let client = Self { request_id: AtomicU64::new(1), write_tx, pending_requests: Arc::new(RwLock::new(HashMap::new())), + reverse_requests, notification_tx, request_tx, read_task: Mutex::new(None), write_task: Mutex::new(Some(write_task)), + timing_task: Mutex::new(timing_task), + timing_shutdown, }; let pending_requests = client.pending_requests.clone(); + let reverse_requests = client.reverse_requests.clone(); let notification_tx_clone = client.notification_tx.clone(); let request_tx_clone = client.request_tx.clone(); let reader_span = tracing::error_span!("jsonrpc_read_loop"); @@ -308,8 +1020,11 @@ impl JsonRpcClient { Self::read_loop( reader, pending_requests, + reverse_requests, notification_tx_clone, request_tx_clone, + timing, + correlation_hasher, ) .await; } @@ -320,6 +1035,89 @@ impl JsonRpcClient { client } + fn start_reverse_rpc_timing_with( + spawn: impl FnOnce( + mpsc::Receiver, + mpsc::Receiver, + Arc, + CancellationToken, + ) -> std::io::Result>, + ) -> Option<( + ReverseRpcTimingEmitter, + std::thread::JoinHandle<()>, + CancellationToken, + CorrelationHasher, + ReverseRpcRegistry, + )> { + let (phase_tx, phase_rx) = + mpsc::channel::(REVERSE_RPC_TIMING_CAPACITY); + let (terminal_tx, terminal_rx) = + mpsc::channel::(REVERSE_RPC_TIMING_CAPACITY); + let dropped_records = Arc::new(AtomicU64::new(0)); + let timing_shutdown = CancellationToken::new(); + match spawn( + phase_rx, + terminal_rx, + dropped_records.clone(), + timing_shutdown.clone(), + ) { + Ok(timing_task) => Some(( + ReverseRpcTimingEmitter { + phase_tx, + terminal_tx, + dropped_records, + }, + timing_task, + timing_shutdown, + CorrelationHasher::new(), + ReverseRpcRegistry::new(), + )), + Err(error) => { + warn!( + error = %error, + "failed to start reverse RPC timing thread; timing disabled" + ); + None + } + } + } + + fn spawn_timing_thread( + phase_rx: mpsc::Receiver, + terminal_rx: mpsc::Receiver, + dropped_records: Arc, + shutdown: CancellationToken, + ) -> std::io::Result> { + let dispatch = tracing::dispatcher::get_default(Clone::clone); + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build()?; + let runtime_slot = Arc::new(Mutex::new(Some(runtime))); + let thread_runtime_slot = runtime_slot.clone(); + let result = std::thread::Builder::new() + .name("copilot-reverse-rpc-timing".to_string()) + .spawn(move || { + tracing::dispatcher::with_default(&dispatch, || { + let runtime = thread_runtime_slot + .lock() + .take() + .expect("reverse RPC timing runtime must be transferred once"); + runtime.block_on(Self::timing_loop( + phase_rx, + terminal_rx, + dropped_records, + shutdown, + )); + }); + }); + if result.is_err() + && let Some(runtime) = runtime_slot.lock().take() + { + runtime.shutdown_background(); + } + result + } + pub(crate) fn force_close(&self) { if let Some(task) = self.read_task.lock().take() { task.abort(); @@ -328,6 +1126,201 @@ impl JsonRpcClient { task.abort(); } self.pending_requests.write().clear(); + if let Some(reverse_requests) = &self.reverse_requests { + reverse_requests.force_abandon_all(); + } + if let Some(timing_shutdown) = &self.timing_shutdown { + timing_shutdown.cancel(); + } + // The timing thread closes its receivers, drains accepted records, + // and exits without waiting for traces held by in-flight handlers. + let _ = self.timing_task.lock().take(); + } + + async fn timing_loop( + mut phase_rx: mpsc::Receiver, + mut terminal_rx: mpsc::Receiver, + dropped_records: Arc, + shutdown: CancellationToken, + ) { + let mut phase_closed = false; + let mut terminal_closed = false; + let mut shutting_down = false; + let mut pending_terminals = VecDeque::new(); + while !phase_closed || !terminal_closed { + Self::record_dropped_timing_records(&dropped_records); + tokio::select! { + biased; + _ = shutdown.cancelled(), if !shutting_down => { + shutting_down = true; + phase_rx.close(); + terminal_rx.close(); + } + event = terminal_rx.recv(), if !terminal_closed => { + if let Some(event) = event { + Self::record_or_defer_timing_event(event, &mut pending_terminals); + } else { + terminal_closed = true; + } + } + event = phase_rx.recv(), if !phase_closed => { + if let Some(event) = event { + Self::record_or_defer_timing_event(event, &mut pending_terminals); + } else { + phase_closed = true; + } + } + } + Self::record_dropped_timing_records(&dropped_records); + } + Self::record_dropped_timing_records(&dropped_records); + } + + fn record_or_defer_timing_event( + event: ReverseRpcTimingEvent, + pending_terminals: &mut VecDeque, + ) { + if let ReverseRpcTimingEvent::Complete { + trace, + required_phase_records, + .. + } = &event + && trace.inner.emitted_phase_records.load(Ordering::Acquire) < *required_phase_records + { + pending_terminals.push_back(event); + return; + } + + let phase_trace = match &event { + ReverseRpcTimingEvent::Complete { .. } => None, + ReverseRpcTimingEvent::Phase { trace, .. } + | ReverseRpcTimingEvent::Scheduled { trace, .. } + | ReverseRpcTimingEvent::HookCallback { trace, .. } => Some(trace.clone()), + }; + Self::record_timing_event(event); + if let Some(trace) = phase_trace { + trace + .inner + .emitted_phase_records + .fetch_add(1, Ordering::Release); + } + + let mut index = 0; + while index < pending_terminals.len() { + let ready = match &pending_terminals[index] { + ReverseRpcTimingEvent::Complete { + trace, + required_phase_records, + .. + } => { + trace.inner.emitted_phase_records.load(Ordering::Acquire) + >= *required_phase_records + } + _ => unreachable!("only terminal records are deferred"), + }; + if ready { + let terminal = pending_terminals + .remove(index) + .expect("pending terminal index should exist"); + Self::record_timing_event(terminal); + } else { + index += 1; + } + } + } + + fn record_timing_event(event: ReverseRpcTimingEvent) { + match event { + ReverseRpcTimingEvent::Phase { + trace, + phase, + start_offset_us, + elapsed_us, + succeeded, + } => { + debug!( + target: REVERSE_RPC_TIMING_TARGET, + parent: None, + correlation_key = %trace.inner.correlation_key, + rpc_method = %trace.inner.method, + phase, + start_offset_us, + elapsed_us, + status = if succeeded { "succeeded" } else { "failed" }, + "reverse JSON-RPC timing" + ); + } + ReverseRpcTimingEvent::Scheduled { + trace, + start_offset_us, + elapsed_us, + since_receive_us, + } => { + debug!( + target: REVERSE_RPC_TIMING_TARGET, + parent: None, + correlation_key = %trace.inner.correlation_key, + rpc_method = %trace.inner.method, + phase = "request_schedule", + start_offset_us, + elapsed_us, + since_receive_us, + status = "succeeded", + "reverse JSON-RPC timing" + ); + } + ReverseRpcTimingEvent::HookCallback { + trace, + hook_type, + start_offset_us, + elapsed_us, + } => { + debug!( + target: REVERSE_RPC_TIMING_TARGET, + parent: None, + correlation_key = %trace.inner.correlation_key, + rpc_method = %trace.inner.method, + hook_type, + phase = "hook_callback", + start_offset_us, + elapsed_us, + status = "succeeded", + "reverse JSON-RPC timing" + ); + } + ReverseRpcTimingEvent::Complete { + trace, + required_phase_records: _, + elapsed_us, + succeeded, + } => { + debug!( + target: REVERSE_RPC_TIMING_TARGET, + parent: None, + correlation_key = %trace.inner.correlation_key, + rpc_method = %trace.inner.method, + phase = "request_complete", + start_offset_us = 0_u64, + elapsed_us, + status = if succeeded { "succeeded" } else { "failed" }, + "reverse JSON-RPC timing" + ); + } + } + } + + fn record_dropped_timing_records(dropped_records: &AtomicU64) { + let dropped_records = dropped_records.swap(0, Ordering::Relaxed); + if dropped_records > 0 { + debug!( + target: REVERSE_RPC_TIMING_TARGET, + parent: None, + phase = "records_dropped", + dropped_records, + status = "dropped", + "reverse JSON-RPC timing records dropped" + ); + } } /// Writer-actor task. Owns the `AsyncWrite`, drains the command queue, @@ -346,13 +1339,52 @@ impl JsonRpcClient { mut writer: impl AsyncWrite + Unpin + Send + 'static, mut rx: mpsc::UnboundedReceiver, ) { - while let Some(WriteCommand { frame, ack }) = rx.recv().await { - let result = async { - writer.write_all(&frame).await?; - writer.flush().await?; - Ok::<_, std::io::Error>(()) + while let Some(WriteCommand { + frame, + ack, + mut reverse_rpc, + enqueued_at, + }) = rx.recv().await + { + let queue_timing = enqueued_at.map(|enqueued_at| (enqueued_at, enqueued_at.elapsed())); + let write_start = reverse_rpc.as_ref().map(|_| TokioInstant::now()); + let write_result = writer.write_all(&frame).await; + let write_timing = write_start.map(|write_start| (write_start, write_start.elapsed())); + let write_succeeded = write_result.is_ok(); + + let (result, flush_timing) = match write_result { + Ok(()) => { + let flush_start = reverse_rpc.as_ref().map(|_| TokioInstant::now()); + let flush_result = writer.flush().await; + let flush_succeeded = flush_result.is_ok(); + ( + flush_result, + flush_start.map(|flush_start| { + (flush_start, flush_start.elapsed(), flush_succeeded) + }), + ) + } + Err(error) => (Err(error), None), + }; + let completed_at = reverse_rpc.as_ref().map(|_| TokioInstant::now()); + let succeeded = result.is_ok(); + + if let Some(write_trace) = &mut reverse_rpc { + let trace = &write_trace.trace; + let (enqueued_at, queue_elapsed) = + queue_timing.expect("timed write must include its enqueue timestamp"); + let (write_start, write_elapsed) = + write_timing.expect("timed write must include its write timestamp"); + trace.record_phase("writer_queue", enqueued_at, queue_elapsed, true); + trace.record_phase("write_all", write_start, write_elapsed, write_succeeded); + if let Some((flush_start, flush_elapsed, flush_succeeded)) = flush_timing { + trace.record_phase("flush", flush_start, flush_elapsed, flush_succeeded); + } + write_trace.record_complete( + completed_at.expect("timed write must include its completion timestamp"), + succeeded, + ); } - .await; // Caller may have dropped the ack receiver (e.g. their // `await` was cancelled); that's fine — we still completed @@ -364,8 +1396,11 @@ impl JsonRpcClient { async fn read_loop( reader: impl AsyncRead + Unpin + Send, pending_requests: Arc>>, + reverse_requests: Option, notification_tx: broadcast::Sender, - request_tx: mpsc::UnboundedSender, + request_tx: ReverseRequestSender, + timing: Option, + correlation_hasher: Option, ) { let mut reader = BufReader::new(reader); @@ -430,7 +1465,46 @@ impl JsonRpcClient { let _ = notification_tx.send(notification); } JsonRpcMessage::Request(request) => { - if request_tx.send(request).is_err() { + let trace = if tracing::enabled!( + target: REVERSE_RPC_TIMING_TARGET, + tracing::Level::DEBUG + ) { + timing + .as_ref() + .zip(correlation_hasher.as_ref()) + .zip(reverse_requests.as_ref()) + .and_then(|((timing, correlation_hasher), registry)| { + registry.register( + &request, + TokioInstant::now(), + correlation_hasher, + timing.clone(), + ) + }) + } else { + None + }; + let forwarded = match &request_tx { + ReverseRequestSender::Public(request_tx) => { + request_tx.send(request).is_ok() + } + ReverseRequestSender::Internal(request_tx) => { + let request = ReverseRpcRequest::new( + request, + trace.clone(), + reverse_requests.clone(), + ); + let result = if let Some(trace) = &trace { + trace.forward(|| request_tx.send(request)) + } else { + request_tx.send(request) + }; + let forwarded = result.is_ok(); + drop(result); + forwarded + } + }; + if !forwarded { warn!("failed to forward JSON-RPC request, channel closed"); } } @@ -455,6 +1529,9 @@ impl JsonRpcClient { ); pending.clear(); } + if let Some(reverse_requests) = &reverse_requests { + reverse_requests.abandon_all(); + } } async fn read_message( @@ -639,7 +1716,43 @@ impl JsonRpcClient { /// drops the ack receiver; the actor still completes the frame and /// flushes. A partial frame can never appear on the wire. pub async fn write(&self, message: &T) -> Result<(), Error> { - let body = serde_json::to_vec(message)?; + self.write_frame(message, None, None).await + } + + pub(crate) async fn write_response(&self, response: &JsonRpcResponse) -> Result<(), Error> { + let reverse_rpc = REVERSE_RPC_RESPONSE_TRACE + .try_with(|scoped| { + (scoped.request_id == response.id) + .then(|| (scoped.trace.clone(), scoped.registry.clone())) + }) + .ok() + .flatten(); + let (trace, registry) = reverse_rpc + .map(|(trace, registry)| (Some(trace), Some(registry))) + .unwrap_or((None, None)); + self.write_frame(response, trace, registry).await + } + + async fn write_frame( + &self, + message: &T, + reverse_rpc: Option, + reverse_rpc_registry: Option, + ) -> Result<(), Error> { + let encode_start = reverse_rpc.as_ref().map(|_| TokioInstant::now()); + let encoded = serde_json::to_vec(message); + if let (Some(trace), Some(encode_start)) = (&reverse_rpc, encode_start) { + trace.record_phase( + "response_encode", + encode_start, + encode_start.elapsed(), + encoded.is_ok(), + ); + if encoded.is_err() { + trace.record_complete(TokioInstant::now(), false); + } + } + let body = encoded?; let mut frame = Vec::with_capacity(CONTENT_LENGTH_HEADER.len() + 16 + body.len() + 4); frame.extend_from_slice(CONTENT_LENGTH_HEADER.as_bytes()); frame.extend_from_slice(body.len().to_string().as_bytes()); @@ -647,14 +1760,26 @@ impl JsonRpcClient { frame.extend_from_slice(&body); let (ack_tx, ack_rx) = oneshot::channel(); - self.write_tx - .send(WriteCommand { frame, ack: ack_tx }) - .map_err(|_| { - Error::from(std::io::Error::new( - std::io::ErrorKind::BrokenPipe, - "writer actor has shut down", - )) - })?; + let enqueued_at = reverse_rpc.as_ref().map(|_| TokioInstant::now()); + if let Some(trace) = &reverse_rpc { + trace.transfer_completion_to_writer(); + } + if self + .write_tx + .send(WriteCommand { + frame, + ack: ack_tx, + reverse_rpc: reverse_rpc + .map(|trace| ReverseRpcWriteTrace::new(trace, reverse_rpc_registry)), + enqueued_at, + }) + .is_err() + { + return Err(Error::from(std::io::Error::new( + std::io::ErrorKind::BrokenPipe, + "writer actor has shut down", + ))); + } match ack_rx.await { Ok(Ok(())) => Ok(()), @@ -692,8 +1817,228 @@ impl Drop for PendingGuard<'_> { #[cfg(test)] mod tests { + use std::collections::VecDeque; + use std::future::Future; + use std::io::{self, Write}; + use std::pin::Pin; + use std::sync::Arc; + use std::task::{Context, Poll}; + use std::time::Duration; + + use parking_lot::Mutex; + use tokio::io::{AsyncWrite, AsyncWriteExt}; + use tokio::time::Sleep; + use tracing_subscriber::Layer; + use tracing_subscriber::fmt::MakeWriter; + use tracing_subscriber::layer::SubscriberExt; + use super::*; + #[derive(Clone, Default)] + struct TraceBuffer(Arc>>); + + impl TraceBuffer { + fn text(&self) -> String { + String::from_utf8(self.0.lock().clone()).unwrap() + } + } + + impl Write for TraceBuffer { + fn write(&mut self, buf: &[u8]) -> std::io::Result { + self.0.lock().extend_from_slice(buf); + Ok(buf.len()) + } + + fn flush(&mut self) -> std::io::Result<()> { + Ok(()) + } + } + + impl<'a> MakeWriter<'a> for TraceBuffer { + type Writer = Self; + + fn make_writer(&'a self) -> Self::Writer { + self.clone() + } + } + + struct DelayedWriter { + write_delays: VecDeque, + flush_delays: VecDeque, + write_sleep: Option>>, + flush_sleep: Option>>, + started_tx: mpsc::UnboundedSender<&'static str>, + } + + impl DelayedWriter { + fn new( + write_delays: [Duration; 2], + flush_delays: [Duration; 2], + ) -> (Self, mpsc::UnboundedReceiver<&'static str>) { + let (started_tx, started_rx) = mpsc::unbounded_channel(); + ( + Self { + write_delays: write_delays.into(), + flush_delays: flush_delays.into(), + write_sleep: None, + flush_sleep: None, + started_tx, + }, + started_rx, + ) + } + + fn poll_delay( + operation: &'static str, + delay: &mut VecDeque, + sleep: &mut Option>>, + started_tx: &mpsc::UnboundedSender<&'static str>, + cx: &mut Context<'_>, + ) -> Poll<()> { + if sleep.is_none() { + let duration = delay.pop_front().unwrap_or_default(); + let _ = started_tx.send(operation); + if duration.is_zero() { + return Poll::Ready(()); + } + *sleep = Some(Box::pin(tokio::time::sleep(duration))); + } + + match sleep + .as_mut() + .expect("delay sleep must exist") + .as_mut() + .poll(cx) + { + Poll::Ready(()) => { + *sleep = None; + Poll::Ready(()) + } + Poll::Pending => Poll::Pending, + } + } + } + + impl AsyncWrite for DelayedWriter { + fn poll_write( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &[u8], + ) -> Poll> { + let Self { + write_delays, + write_sleep, + started_tx, + .. + } = self.as_mut().get_mut(); + match Self::poll_delay("write", write_delays, write_sleep, started_tx, cx) { + Poll::Ready(()) => Poll::Ready(Ok(buf.len())), + Poll::Pending => Poll::Pending, + } + } + + fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + let Self { + flush_delays, + flush_sleep, + started_tx, + .. + } = self.as_mut().get_mut(); + match Self::poll_delay("flush", flush_delays, flush_sleep, started_tx, cx) { + Poll::Ready(()) => Poll::Ready(Ok(())), + Poll::Pending => Poll::Pending, + } + } + + fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + } + + #[derive(Clone)] + struct BlockingTraceWriter { + entered_tx: std::sync::mpsc::Sender<()>, + release: Arc<(std::sync::Mutex, std::sync::Condvar)>, + } + + impl Write for BlockingTraceWriter { + fn write(&mut self, buf: &[u8]) -> io::Result { + let _ = self.entered_tx.send(()); + let (released, condvar) = &*self.release; + let mut released = released.lock().unwrap(); + while !*released { + released = condvar.wait(released).unwrap(); + } + Ok(buf.len()) + } + + fn flush(&mut self) -> io::Result<()> { + Ok(()) + } + } + + impl<'a> MakeWriter<'a> for BlockingTraceWriter { + type Writer = Self; + + fn make_writer(&'a self) -> Self::Writer { + self.clone() + } + } + + fn trace_subscriber(buffer: TraceBuffer) -> impl tracing::Subscriber { + tracing_subscriber::registry().with( + tracing_subscriber::fmt::layer() + .with_writer(buffer) + .with_ansi(false) + .without_time() + .with_filter(tracing_subscriber::filter::filter_fn(|metadata| { + metadata.target() == REVERSE_RPC_TIMING_TARGET + })), + ) + } + + fn timing_channel( + capacity: usize, + ) -> ( + ReverseRpcTimingEmitter, + mpsc::Receiver, + mpsc::Receiver, + Arc, + ) { + let (phase_tx, phase_rx) = mpsc::channel(capacity); + let (terminal_tx, terminal_rx) = mpsc::channel(capacity); + let dropped_records = Arc::new(AtomicU64::new(0)); + ( + ReverseRpcTimingEmitter { + phase_tx, + terminal_tx, + dropped_records: dropped_records.clone(), + }, + phase_rx, + terminal_rx, + dropped_records, + ) + } + + async fn wait_for_trace(buffer: &TraceBuffer, needle: &str) { + for _ in 0..20 { + if buffer.text().contains(needle) { + return; + } + tokio::task::yield_now().await; + } + panic!("timing trace did not contain {needle:?}: {}", buffer.text()); + } + + fn frame(message: &impl Serialize) -> Vec { + let body = serde_json::to_vec(message).unwrap(); + format!("Content-Length: {}\r\n\r\n", body.len()) + .into_bytes() + .into_iter() + .chain(body) + .collect() + } + #[test] fn deserialize_notification() { let json = r#"{"jsonrpc":"2.0","method":"session.event","params":{"id":"e1"}}"#; @@ -780,4 +2125,928 @@ mod tests { let json = serde_json::to_string(&r).unwrap(); assert!(!json.contains("error")); } + + #[test] + fn reverse_request_correlation_is_stable_and_opaque() { + let request = JsonRpcRequest::new( + 17, + "hooks.invoke", + Some(serde_json::json!({ "sessionId": "private-session-id" })), + ); + let correlation_hasher = CorrelationHasher::new(); + let same = ReverseRpcTrace::correlation_key( + &correlation_hasher, + request.id, + &request.method, + Some("private-session-id"), + 7, + ); + let repeated = ReverseRpcTrace::correlation_key( + &correlation_hasher, + request.id, + &request.method, + Some("private-session-id"), + 7, + ); + let different_session = ReverseRpcTrace::correlation_key( + &correlation_hasher, + request.id, + &request.method, + Some("other-session"), + 7, + ); + let different_generation = ReverseRpcTrace::correlation_key( + &correlation_hasher, + request.id, + &request.method, + Some("private-session-id"), + 8, + ); + + assert_eq!(same, repeated); + assert_ne!(same, different_session); + assert_ne!(same, different_generation); + assert!(same.starts_with("rrpc-")); + assert!(!same.contains("private-session-id")); + } + + #[test] + fn reverse_request_timing_labels_all_supported_hook_types() { + for hook_type in [ + "preToolUse", + "preMcpToolCall", + "postToolUse", + "postToolUseFailure", + "userPromptSubmitted", + "userPromptTransformed", + "sessionStart", + "sessionEnd", + "errorOccurred", + "agentStop", + ] { + assert_eq!(ReverseRpcTrace::timing_hook_type(hook_type), hook_type); + } + assert_eq!( + ReverseRpcTrace::timing_hook_type("PRIVATE_SENTINEL_DO_NOT_TRACE"), + "unknown" + ); + } + + #[test] + fn reverse_request_guard_only_removes_its_own_generation() { + let (timing, _phase_rx, _terminal_rx, _dropped_records) = timing_channel(1); + let request = JsonRpcRequest::new(17, "hooks.invoke", None); + let now = TokioInstant::now(); + let correlation_hasher = CorrelationHasher::new(); + let first = ReverseRpcTrace::new(&request, 1, now, &correlation_hasher, timing.clone()); + let second = ReverseRpcTrace::new(&request, 2, now, &correlation_hasher, timing); + let registry = ReverseRpcRegistry::new(); + registry.insert(first.clone()); + registry.insert(second.clone()); + let guard = ReverseRpcDispatchGuard { + registry: registry.clone(), + request_id: request.id, + trace: first, + }; + + drop(guard); + + assert!(registry.contains(&second)); + } + + #[tokio::test(start_paused = true)] + async fn reverse_request_timing_uses_the_forwarding_boundary() { + let trace_buffer = TraceBuffer::default(); + let _subscriber = tracing::subscriber::set_default(trace_subscriber(trace_buffer.clone())); + let (timing, phase_rx, terminal_rx, dropped_records) = + timing_channel(REVERSE_RPC_TIMING_CAPACITY); + tokio::spawn(JsonRpcClient::timing_loop( + phase_rx, + terminal_rx, + dropped_records, + CancellationToken::new(), + )); + let request = JsonRpcRequest::new( + 29, + "hooks.invoke", + Some(serde_json::json!({ "sessionId": "session" })), + ); + let received_at = TokioInstant::now(); + let trace = + ReverseRpcTrace::new(&request, 1, received_at, &CorrelationHasher::new(), timing); + + tokio::time::advance(Duration::from_millis(5)).await; + trace.forward(|| Ok::<(), ()>(())).unwrap(); + tokio::time::advance(Duration::from_millis(7)).await; + trace.record_scheduled(TokioInstant::now()); + + wait_for_trace(&trace_buffer, "phase=\"request_schedule\"").await; + let output = trace_buffer.text(); + let forward = output + .lines() + .find(|line| line.contains("phase=\"request_forward\"")) + .expect("request_forward timing should be emitted"); + let schedule = output + .lines() + .find(|line| line.contains("phase=\"request_schedule\"")) + .expect("request_schedule timing should be emitted"); + assert!(forward.contains("elapsed_us=5000")); + assert!(forward.contains("start_offset_us=0")); + assert!(schedule.contains("elapsed_us=7000")); + assert!(schedule.contains("start_offset_us=5000")); + assert!(schedule.contains("since_receive_us=12000")); + } + + #[tokio::test] + async fn reverse_request_forwarding_is_recorded_before_dispatch_can_start() { + let trace_buffer = TraceBuffer::default(); + let _subscriber = tracing::subscriber::set_default(trace_subscriber(trace_buffer.clone())); + let (timing, phase_rx, terminal_rx, dropped_records) = + timing_channel(REVERSE_RPC_TIMING_CAPACITY); + tokio::spawn(JsonRpcClient::timing_loop( + phase_rx, + terminal_rx, + dropped_records, + CancellationToken::new(), + )); + let request = JsonRpcRequest::new(31, "hooks.invoke", None); + let now = TokioInstant::now(); + let trace = ReverseRpcTrace::new(&request, 1, now, &CorrelationHasher::new(), timing); + let registry = ReverseRpcRegistry::new(); + registry.insert(trace.clone()); + let forwarded = + ReverseRpcRequest::new(request, Some(trace.clone()), Some(registry.clone())); + let (request_tx, request_rx) = std::sync::mpsc::channel::(); + let (received_tx, received_rx) = std::sync::mpsc::channel(); + let receiver = std::thread::spawn(move || { + let request = request_rx.recv().unwrap(); + received_tx.send(()).unwrap(); + let (_request, guard) = request.into_dispatch(); + drop(guard); + }); + + trace + .forward(|| { + request_tx.send(forwarded).map_err(|_| ())?; + received_rx.recv().map_err(|_| ())?; + Ok::<(), ()>(()) + }) + .unwrap(); + receiver.join().unwrap(); + + wait_for_trace(&trace_buffer, "phase=\"request_complete\"").await; + let output = trace_buffer.text(); + let forward = output.find("phase=\"request_forward\"").unwrap(); + let schedule = output.find("phase=\"request_schedule\"").unwrap(); + let complete = output.find("phase=\"request_complete\"").unwrap(); + assert!(forward < schedule); + assert!(schedule < complete); + assert!(registry.is_empty()); + } + + #[tokio::test] + async fn closed_forward_channel_records_failed_forward_before_abandonment() { + let trace_buffer = TraceBuffer::default(); + let _subscriber = tracing::subscriber::set_default(trace_subscriber(trace_buffer.clone())); + let (mut server, reader) = tokio::io::duplex(4096); + let (notification_tx, _) = broadcast::channel(1); + let (request_tx, request_rx) = mpsc::unbounded_channel(); + drop(request_rx); + let client = JsonRpcClient::new_with_reverse_rpc_timing( + tokio::io::sink(), + reader, + notification_tx, + request_tx, + ); + let request = JsonRpcRequest::new(32, "hooks.invoke", None); + + server.write_all(&frame(&request)).await.unwrap(); + + wait_for_trace(&trace_buffer, "phase=\"request_complete\"").await; + let output = trace_buffer.text(); + let forward = output + .lines() + .find(|line| line.contains("phase=\"request_forward\"")) + .expect("failed forwarding phase should be emitted"); + assert!(forward.contains("status=\"failed\"")); + assert!( + output.find("phase=\"request_forward\"").unwrap() + < output.find("phase=\"request_complete\"").unwrap() + ); + + client.force_close(); + } + + #[tokio::test] + async fn public_client_does_not_retain_reverse_request_timing_state() { + let (mut server, reader) = tokio::io::duplex(4096); + let (notification_tx, _) = broadcast::channel(1); + let (request_tx, mut request_rx) = mpsc::unbounded_channel(); + let client = JsonRpcClient::new(tokio::io::sink(), reader, notification_tx, request_tx); + let request = JsonRpcRequest::new(23, "consumer.request", None); + + server.write_all(&frame(&request)).await.unwrap(); + let forwarded = request_rx.recv().await.unwrap(); + + assert_eq!(forwarded.id, request.id); + assert!(client.reverse_requests.is_none()); + assert!(client.timing_task.lock().is_none()); + assert!(client.timing_shutdown.is_none()); + client.force_close(); + } + + #[tokio::test] + async fn disabled_timing_target_does_not_allocate_reverse_request_state() { + let _subscriber = + tracing::subscriber::set_default(tracing::subscriber::NoSubscriber::default()); + let (mut server, reader) = tokio::io::duplex(4096); + let (notification_tx, _) = broadcast::channel(1); + let (request_tx, mut request_rx) = mpsc::unbounded_channel(); + let client = JsonRpcClient::new_with_reverse_rpc_timing( + tokio::io::sink(), + reader, + notification_tx, + request_tx, + ); + let request = JsonRpcRequest::new( + 29, + "hooks.invoke", + Some(serde_json::json!({ "sessionId": "session" })), + ); + + server.write_all(&frame(&request)).await.unwrap(); + let forwarded = request_rx.recv().await.unwrap(); + + assert_eq!(forwarded.id, request.id); + assert!(forwarded.trace.is_none()); + assert!(client.reverse_requests.is_none()); + assert!(client.timing_task.lock().is_none()); + assert!(client.timing_shutdown.is_none()); + client.force_close(); + } + + #[test] + fn timing_thread_spawn_failure_disables_timing_without_panicking() { + let timing = JsonRpcClient::start_reverse_rpc_timing_with(|_, _, _, _| { + Err(std::io::Error::other("injected timing thread failure")) + }); + + assert!(timing.is_none()); + } + + #[test] + fn force_closed_registry_rejects_late_reverse_requests() { + let (timing, _phase_rx, _terminal_rx, _dropped_records) = + timing_channel(REVERSE_RPC_TIMING_CAPACITY); + let registry = ReverseRpcRegistry::new(); + registry.force_abandon_all(); + + let trace = registry.register( + &JsonRpcRequest::new(35, "hooks.invoke", None), + TokioInstant::now(), + &CorrelationHasher::new(), + timing, + ); + + assert!(trace.is_none()); + assert!(registry.is_empty()); + } + + #[tokio::test] + async fn saturated_timing_queue_drops_records_and_reports_the_count() { + let trace_buffer = TraceBuffer::default(); + let _subscriber = tracing::subscriber::set_default(trace_subscriber(trace_buffer.clone())); + let (timing, phase_rx, terminal_rx, dropped_records) = timing_channel(1); + let request = JsonRpcRequest::new(37, "hooks.invoke", None); + let now = TokioInstant::now(); + let trace = ReverseRpcTrace::new(&request, 1, now, &CorrelationHasher::new(), timing); + trace.mark_forwarding(now); + + trace.record_phase("first", now, Duration::ZERO, true); + trace.record_phase("second", now, Duration::ZERO, true); + trace.record_complete(now, true); + assert_eq!( + trace.inner.timing.dropped_records.load(Ordering::Relaxed), + 1 + ); + drop(trace); + tokio::spawn(JsonRpcClient::timing_loop( + phase_rx, + terminal_rx, + dropped_records, + CancellationToken::new(), + )); + + wait_for_trace(&trace_buffer, "phase=\"request_complete\"").await; + let output = trace_buffer.text(); + let dropped = output + .lines() + .find(|line| line.contains("phase=\"records_dropped\"")) + .expect("saturation diagnostic should be emitted"); + assert!(dropped.contains("dropped_records=1")); + assert!(output.contains("phase=\"first\"")); + assert!(!output.contains("phase=\"second\"")); + assert!(output.contains("phase=\"request_complete\"")); + assert!( + output.find("phase=\"first\"").unwrap() + < output.find("phase=\"request_complete\"").unwrap() + ); + } + + #[tokio::test] + async fn saturated_terminal_queue_counts_one_writer_completion_drop_once() { + let trace_buffer = TraceBuffer::default(); + let _subscriber = tracing::subscriber::set_default(trace_subscriber(trace_buffer.clone())); + let (timing, phase_rx, terminal_rx, dropped_records) = timing_channel(1); + let now = TokioInstant::now(); + let filler_request = JsonRpcRequest::new(38, "hooks.invoke", None); + let filler = ReverseRpcTrace::new( + &filler_request, + 1, + now, + &CorrelationHasher::new(), + timing.clone(), + ); + filler.mark_forwarding(now); + assert!(filler.record_complete(now, true)); + + let writer_request = JsonRpcRequest::new(39, "userInput.request", None); + let writer = + ReverseRpcTrace::new(&writer_request, 2, now, &CorrelationHasher::new(), timing); + writer.mark_forwarding(now); + let mut writer_trace = ReverseRpcWriteTrace::new(writer, None); + writer_trace.record_complete(now, true); + drop(writer_trace); + + assert_eq!( + filler.inner.timing.dropped_records.load(Ordering::Relaxed), + 1 + ); + drop(filler); + tokio::spawn(JsonRpcClient::timing_loop( + phase_rx, + terminal_rx, + dropped_records, + CancellationToken::new(), + )); + + wait_for_trace(&trace_buffer, "phase=\"records_dropped\"").await; + let output = trace_buffer.text(); + let dropped = output + .lines() + .find(|line| line.contains("phase=\"records_dropped\"")) + .expect("terminal saturation diagnostic should be emitted"); + assert!(dropped.contains("dropped_records=1")); + assert!(output.contains("rpc_method=hooks.invoke")); + assert!(!output.lines().any(|line| { + line.contains("rpc_method=userInput.request") + && line.contains("phase=\"request_complete\"") + })); + } + + #[tokio::test] + async fn saturated_terminal_queue_counts_closed_writer_fallback_once() { + let trace_buffer = TraceBuffer::default(); + let _subscriber = tracing::subscriber::set_default(trace_subscriber(trace_buffer.clone())); + let (timing, phase_rx, terminal_rx, dropped_records) = timing_channel(1); + let now = TokioInstant::now(); + let filler_request = JsonRpcRequest::new(40, "hooks.invoke", None); + let filler = ReverseRpcTrace::new( + &filler_request, + 1, + now, + &CorrelationHasher::new(), + timing.clone(), + ); + filler.mark_forwarding(now); + assert!(filler.record_complete(now, true)); + + let (notification_tx, _) = broadcast::channel(1); + let (request_tx, _request_rx) = mpsc::unbounded_channel(); + let client = JsonRpcClient::new( + tokio::io::sink(), + tokio::io::empty(), + notification_tx, + request_tx, + ); + let write_task = client + .write_task + .lock() + .take() + .expect("writer task should be running"); + write_task.abort(); + let _ = write_task.await; + + let writer_request = JsonRpcRequest::new(41, "userInput.request", None); + let writer = + ReverseRpcTrace::new(&writer_request, 2, now, &CorrelationHasher::new(), timing); + writer.mark_forwarding(now); + let error = client + .write_frame( + &JsonRpcResponse { + jsonrpc: "2.0".to_string(), + id: writer_request.id, + result: Some(serde_json::json!({})), + error: None, + }, + Some(writer), + None, + ) + .await + .unwrap_err(); + assert!(matches!(error.kind(), ErrorKind::Io)); + assert_eq!( + filler.inner.timing.dropped_records.load(Ordering::Relaxed), + 1 + ); + drop(filler); + tokio::spawn(JsonRpcClient::timing_loop( + phase_rx, + terminal_rx, + dropped_records, + CancellationToken::new(), + )); + + wait_for_trace(&trace_buffer, "phase=\"records_dropped\"").await; + let output = trace_buffer.text(); + let dropped = output + .lines() + .find(|line| line.contains("phase=\"records_dropped\"")) + .expect("closed-writer saturation diagnostic should be emitted"); + assert!(dropped.contains("dropped_records=1")); + assert!(output.contains("rpc_method=hooks.invoke")); + assert!(!output.lines().any(|line| { + line.contains("rpc_method=userInput.request") + && line.contains("phase=\"request_complete\"") + })); + + client.force_close(); + } + + #[tokio::test(start_paused = true)] + async fn reverse_request_timing_tracks_gated_scheduling_without_content() { + const SENTINEL: &str = "PRIVATE_SENTINEL_DO_NOT_TRACE"; + + let trace_buffer = TraceBuffer::default(); + let _subscriber = tracing::subscriber::set_default(trace_subscriber(trace_buffer.clone())); + let (mut server, reader) = tokio::io::duplex(4096); + let (notification_tx, _) = broadcast::channel(1); + let (request_tx, mut request_rx) = mpsc::unbounded_channel(); + let client = JsonRpcClient::new_with_reverse_rpc_timing( + tokio::io::sink(), + reader, + notification_tx, + request_tx, + ); + let request = JsonRpcRequest::new( + 41, + SENTINEL, + Some(serde_json::json!({ + "sessionId": SENTINEL, + "hookType": "userPromptSubmitted", + "input": { + "prompt": SENTINEL, + "cwd": SENTINEL, + "toolArgs": { "secret": SENTINEL } + } + })), + ); + + server.write_all(&frame(&request)).await.unwrap(); + let forwarded = request_rx.recv().await.unwrap(); + tokio::time::advance(Duration::from_millis(13)).await; + + let (forwarded, trace) = forwarded.into_dispatch(); + let trace = trace.expect("reverse request timing should be tracked"); + let callback_start = TokioInstant::now(); + tokio::time::advance(Duration::from_millis(3)).await; + trace + .trace() + .record_hook_callback(SENTINEL, callback_start, callback_start.elapsed()); + trace + .scope(client.write_response(&JsonRpcResponse { + jsonrpc: "2.0".to_string(), + id: forwarded.id, + result: Some(serde_json::json!({ "output": SENTINEL })), + error: None, + })) + .await + .unwrap(); + + wait_for_trace(&trace_buffer, "phase=\"request_complete\"").await; + let output = trace_buffer.text(); + assert!(output.contains("github_copilot_sdk::reverse_rpc_timing")); + assert!(output.contains("rpc_method=unknown")); + assert!(output.contains("hook_type=\"unknown\"")); + assert!(output.contains("phase=\"request_forward\"")); + assert!(output.contains("phase=\"request_schedule\"")); + assert!(output.contains("elapsed_us=13000")); + assert!(output.contains("phase=\"hook_callback\"")); + assert!(output.contains("phase=\"response_encode\"")); + assert!(output.contains("phase=\"writer_queue\"")); + assert!(output.contains("phase=\"write_all\"")); + assert!(output.contains("phase=\"flush\"")); + assert!(output.contains("phase=\"request_complete\"")); + assert!(output.contains("start_offset_us=")); + assert!(output.contains("correlation_key=rrpc-")); + assert!(!output.contains(SENTINEL)); + + client.force_close(); + } + + #[tokio::test(start_paused = true)] + async fn reverse_response_timing_distinguishes_writer_queue_write_and_flush() { + let trace_buffer = TraceBuffer::default(); + let _subscriber = tracing::subscriber::set_default(trace_subscriber(trace_buffer.clone())); + let (writer, mut started_rx) = DelayedWriter::new( + [Duration::from_millis(20), Duration::from_millis(7)], + [Duration::ZERO, Duration::from_millis(11)], + ); + let (notification_tx, _) = broadcast::channel(1); + let (request_tx, _request_rx) = mpsc::unbounded_channel(); + let client = JsonRpcClient::new(writer, tokio::io::empty(), notification_tx, request_tx); + + let (first_ack_tx, first_ack_rx) = oneshot::channel(); + client + .write_tx + .send(WriteCommand { + frame: frame(&serde_json::json!({})), + ack: first_ack_tx, + reverse_rpc: None, + enqueued_at: None, + }) + .unwrap(); + assert_eq!(started_rx.recv().await, Some("write")); + tokio::time::advance(Duration::from_millis(5)).await; + + let request = JsonRpcRequest::new( + 7, + "hooks.invoke", + Some(serde_json::json!({ "sessionId": "session" })), + ); + let now = TokioInstant::now(); + let trace = ReverseRpcTrace::for_test(&request, now, now); + let (second_ack_tx, second_ack_rx) = oneshot::channel(); + client + .write_tx + .send(WriteCommand { + frame: frame(&JsonRpcResponse { + jsonrpc: "2.0".to_string(), + id: 7, + result: Some(serde_json::json!({})), + error: None, + }), + ack: second_ack_tx, + reverse_rpc: Some(ReverseRpcWriteTrace::new(trace, None)), + enqueued_at: Some(TokioInstant::now()), + }) + .unwrap(); + + tokio::time::advance(Duration::from_millis(15)).await; + assert_eq!(started_rx.recv().await, Some("flush")); + assert_eq!(started_rx.recv().await, Some("write")); + tokio::time::advance(Duration::from_millis(7)).await; + assert_eq!(started_rx.recv().await, Some("flush")); + tokio::time::advance(Duration::from_millis(11)).await; + + first_ack_rx.await.unwrap().unwrap(); + second_ack_rx.await.unwrap().unwrap(); + + wait_for_trace(&trace_buffer, "phase=\"request_complete\"").await; + let output = trace_buffer.text(); + assert!(output.contains("phase=\"writer_queue\" start_offset_us=0 elapsed_us=15000")); + assert!(output.contains("phase=\"write_all\" start_offset_us=15000 elapsed_us=7000")); + assert!(output.contains("phase=\"flush\" start_offset_us=22000 elapsed_us=11000")); + assert!(output.contains("phase=\"request_complete\" start_offset_us=0 elapsed_us=33000")); + let writer_queue = output.find("phase=\"writer_queue\"").unwrap(); + let write_all = output.find("phase=\"write_all\"").unwrap(); + let flush = output.find("phase=\"flush\"").unwrap(); + let complete = output.find("phase=\"request_complete\"").unwrap(); + assert!(writer_queue < write_all); + assert!(write_all < flush); + assert!(flush < complete); + + client.force_close(); + } + + #[tokio::test(start_paused = true)] + async fn cancelled_response_keeps_the_writer_terminal_outcome() { + let trace_buffer = TraceBuffer::default(); + let _subscriber = tracing::subscriber::set_default(trace_subscriber(trace_buffer.clone())); + let (writer, mut started_rx) = DelayedWriter::new( + [Duration::from_millis(10), Duration::ZERO], + [Duration::ZERO, Duration::ZERO], + ); + let (notification_tx, _) = broadcast::channel(1); + let (request_tx, _request_rx) = mpsc::unbounded_channel(); + let (_server_guard, reader) = tokio::io::duplex(64); + let client = Arc::new(JsonRpcClient::new( + writer, + reader, + notification_tx, + request_tx, + )); + let request = JsonRpcRequest::new(53, "hooks.invoke", None); + let now = TokioInstant::now(); + let trace = ReverseRpcTrace::for_test(&request, now, now); + let registry = ReverseRpcRegistry::new(); + registry.insert(trace.clone()); + let dispatch_guard = ReverseRpcDispatchGuard { + registry: registry.clone(), + request_id: request.id, + trace, + }; + let response_task = tokio::spawn({ + let client = client.clone(); + async move { + dispatch_guard + .scope(client.write_response(&JsonRpcResponse { + jsonrpc: "2.0".to_string(), + id: request.id, + result: Some(serde_json::json!({})), + error: None, + })) + .await + } + }); + + assert_eq!(started_rx.recv().await, Some("write")); + response_task.abort(); + let _ = response_task.await; + tokio::time::advance(Duration::from_millis(10)).await; + assert_eq!(started_rx.recv().await, Some("flush")); + + wait_for_trace(&trace_buffer, "phase=\"request_complete\"").await; + let complete = trace_buffer + .text() + .lines() + .filter(|line| line.contains("phase=\"request_complete\"")) + .map(str::to_owned) + .collect::>(); + assert_eq!(complete.len(), 1); + assert!(complete[0].contains("status=\"succeeded\"")); + assert!(registry.is_empty()); + + client.force_close(); + } + + #[tokio::test] + async fn response_timing_uses_the_exact_dispatch_generation_when_ids_are_reused() { + let trace_buffer = TraceBuffer::default(); + let _subscriber = tracing::subscriber::set_default(trace_subscriber(trace_buffer.clone())); + let (notification_tx, _) = broadcast::channel(1); + let (request_tx, mut request_rx) = mpsc::unbounded_channel(); + let (mut server, reader) = tokio::io::duplex(4096); + let client = JsonRpcClient::new_with_reverse_rpc_timing( + tokio::io::sink(), + reader, + notification_tx, + request_tx, + ); + let params = Some(serde_json::json!({ "sessionId": "same-session" })); + let first_request = JsonRpcRequest::new(61, "hooks.invoke", params.clone()); + let second_request = JsonRpcRequest::new(61, "hooks.invoke", params); + server.write_all(&frame(&first_request)).await.unwrap(); + server.write_all(&frame(&second_request)).await.unwrap(); + let first_forwarded = request_rx.recv().await.unwrap(); + let second_forwarded = request_rx.recv().await.unwrap(); + + // Acquire the newer dispatch first to prove scheduling order cannot + // change which receipt-generation each forwarded request carries. + let (second_forwarded, second_guard) = second_forwarded.into_dispatch(); + let second_guard = second_guard.expect("second request should carry timing"); + let second_trace = second_guard.trace.clone(); + let second_correlation = second_trace.inner.correlation_key.clone(); + let (first_forwarded, first_guard) = first_forwarded.into_dispatch(); + let first_guard = first_guard.expect("first request should carry timing"); + let first_correlation = first_guard.trace.inner.correlation_key.clone(); + assert_ne!(first_correlation, second_correlation); + + first_guard + .scope(client.write_response(&JsonRpcResponse { + jsonrpc: "2.0".to_string(), + id: first_forwarded.id, + result: Some(serde_json::json!({})), + error: None, + })) + .await + .unwrap(); + wait_for_trace( + &trace_buffer, + &format!( + "correlation_key={first_correlation} rpc_method=hooks.invoke phase=\"request_complete\"" + ), + ) + .await; + assert!( + client + .reverse_requests + .as_ref() + .is_some_and(|registry| registry.contains(&second_trace)) + ); + + second_guard + .scope(client.write_response(&JsonRpcResponse { + jsonrpc: "2.0".to_string(), + id: second_forwarded.id, + result: Some(serde_json::json!({})), + error: None, + })) + .await + .unwrap(); + wait_for_trace( + &trace_buffer, + &format!( + "correlation_key={second_correlation} rpc_method=hooks.invoke phase=\"request_complete\"" + ), + ) + .await; + + let output = trace_buffer.text(); + assert_eq!( + output + .lines() + .filter(|line| { + line.contains(&format!("correlation_key={first_correlation}")) + && line.contains("phase=\"request_complete\"") + }) + .count(), + 1 + ); + assert_eq!( + output + .lines() + .filter(|line| { + line.contains(&format!("correlation_key={second_correlation}")) + && line.contains("phase=\"request_complete\"") + }) + .count(), + 1 + ); + assert!( + client + .reverse_requests + .as_ref() + .is_some_and(ReverseRpcRegistry::is_empty) + ); + + client.force_close(); + } + + #[tokio::test(start_paused = true)] + async fn force_close_emits_one_failed_terminal_record() { + let trace_buffer = TraceBuffer::default(); + let _subscriber = tracing::subscriber::set_default(trace_subscriber(trace_buffer.clone())); + let (writer, mut started_rx) = DelayedWriter::new( + [Duration::from_secs(60), Duration::ZERO], + [Duration::ZERO, Duration::ZERO], + ); + let (mut server, reader) = tokio::io::duplex(4096); + let (notification_tx, _) = broadcast::channel(1); + let (request_tx, mut request_rx) = mpsc::unbounded_channel(); + let client = Arc::new(JsonRpcClient::new_with_reverse_rpc_timing( + writer, + reader, + notification_tx, + request_tx, + )); + let request = JsonRpcRequest::new(59, "hooks.invoke", None); + server.write_all(&frame(&request)).await.unwrap(); + let forwarded = request_rx.recv().await.unwrap(); + let (forwarded, dispatch_guard) = forwarded.into_dispatch(); + let dispatch_guard = + dispatch_guard.expect("enabled timing target should track the request"); + let response_task = tokio::spawn({ + let client = client.clone(); + async move { + dispatch_guard + .scope(client.write_response(&JsonRpcResponse { + jsonrpc: "2.0".to_string(), + id: forwarded.id, + result: Some(serde_json::json!({})), + error: None, + })) + .await + } + }); + + assert_eq!(started_rx.recv().await, Some("write")); + response_task.abort(); + assert!(response_task.await.unwrap_err().is_cancelled()); + assert!( + client + .reverse_requests + .as_ref() + .is_some_and(|registry| !registry.is_empty()) + ); + let timing_thread = client + .timing_task + .lock() + .take() + .expect("enabled timing target should start the timing thread"); + client.force_close(); + timing_thread.join().unwrap(); + + let complete = trace_buffer + .text() + .lines() + .filter(|line| line.contains("phase=\"request_complete\"")) + .map(str::to_owned) + .collect::>(); + assert_eq!(complete.len(), 1); + assert!(complete[0].contains("status=\"failed\"")); + } + + #[tokio::test] + async fn slow_timing_subscriber_does_not_delay_response_ack() { + let (entered_tx, entered_rx) = std::sync::mpsc::channel(); + let release = Arc::new((std::sync::Mutex::new(false), std::sync::Condvar::new())); + let blocking_writer = BlockingTraceWriter { + entered_tx, + release: release.clone(), + }; + let (timing, phase_rx, terminal_rx, dropped_records) = + timing_channel(REVERSE_RPC_TIMING_CAPACITY); + let subscriber = tracing_subscriber::registry().with( + tracing_subscriber::fmt::layer() + .with_writer(blocking_writer) + .with_ansi(false) + .without_time() + .with_filter(tracing_subscriber::filter::filter_fn(|metadata| { + metadata.target() == REVERSE_RPC_TIMING_TARGET + })), + ); + let timing_thread = tracing::subscriber::with_default(subscriber, || { + JsonRpcClient::spawn_timing_thread( + phase_rx, + terminal_rx, + dropped_records, + CancellationToken::new(), + ) + .unwrap() + }); + let (notification_tx, _) = broadcast::channel(1); + let (request_tx, _request_rx) = mpsc::unbounded_channel(); + let client = JsonRpcClient::new( + tokio::io::sink(), + tokio::io::empty(), + notification_tx, + request_tx, + ); + let request = JsonRpcRequest::new(53, "hooks.invoke", None); + let now = TokioInstant::now(); + let trace = ReverseRpcTrace::new(&request, 1, now, &CorrelationHasher::new(), timing); + trace.mark_forwarding(now); + + tokio::time::timeout( + Duration::from_secs(1), + client.write_frame( + &JsonRpcResponse { + jsonrpc: "2.0".to_string(), + id: request.id, + result: Some(serde_json::json!({})), + error: None, + }, + Some(trace), + None, + ), + ) + .await + .expect("response acknowledgement should not wait for trace formatting") + .unwrap(); + entered_rx + .recv_timeout(Duration::from_secs(1)) + .expect("timing subscriber should be blocked after acknowledgement"); + + let (released, condvar) = &*release; + *released.lock().unwrap() = true; + condvar.notify_all(); + client.force_close(); + timing_thread.join().unwrap(); + } + + #[test] + fn timing_thread_shutdown_does_not_wait_for_trace_senders() { + let (timing, phase_rx, terminal_rx, dropped_records) = + timing_channel(REVERSE_RPC_TIMING_CAPACITY); + let shutdown = CancellationToken::new(); + let timing_thread = JsonRpcClient::spawn_timing_thread( + phase_rx, + terminal_rx, + dropped_records, + shutdown.clone(), + ) + .unwrap(); + let (joined_tx, joined_rx) = std::sync::mpsc::channel(); + std::thread::spawn(move || { + timing_thread.join().unwrap(); + joined_tx.send(()).unwrap(); + }); + + shutdown.cancel(); + joined_rx + .recv_timeout(Duration::from_secs(1)) + .expect("timing thread should exit while timing senders remain alive"); + drop(timing); + } } diff --git a/rust/src/lib.rs b/rust/src/lib.rs index 4a9f73ca4..6d4e06e52 100644 --- a/rust/src/lib.rs +++ b/rust/src/lib.rs @@ -1023,7 +1023,7 @@ struct ClientInner { ffi_host: parking_lot::Mutex>>, rpc: JsonRpcClient, cwd: PathBuf, - request_rx: parking_lot::Mutex>>, + request_rx: parking_lot::Mutex>>, notification_tx: broadcast::Sender, router: router::SessionRouter, github_token_registry: Arc, @@ -1599,9 +1599,9 @@ impl Client { mode: ClientMode, ) -> Result { let setup_start = Instant::now(); - let (request_tx, request_rx) = mpsc::unbounded_channel::(); + let (request_tx, request_rx) = mpsc::unbounded_channel::(); let (notification_broadcast_tx, _) = broadcast::channel::(1024); - let rpc = JsonRpcClient::new( + let rpc = JsonRpcClient::new_with_reverse_rpc_timing( writer, reader, notification_broadcast_tx.clone(), @@ -2025,7 +2025,7 @@ impl Client { /// Send a JSON-RPC response back to the CLI (e.g. for permission or tool call requests). pub(crate) async fn send_response(&self, response: &JsonRpcResponse) -> Result<()> { - self.inner.rpc.write(response).await + self.inner.rpc.write_response(response).await } /// Reconstruct a [`Client`] handle from a shared inner pointer. @@ -2037,7 +2037,9 @@ impl Client { /// /// Can only be called once — subsequent calls return `None`. #[expect(dead_code, reason = "reserved for future pub(crate) use")] - pub(crate) fn take_request_rx(&self) -> Option> { + pub(crate) fn take_request_rx( + &self, + ) -> Option> { self.inner.request_rx.lock().take() } diff --git a/rust/src/router.rs b/rust/src/router.rs index 1dec9d16f..c02a00928 100644 --- a/rust/src/router.rs +++ b/rust/src/router.rs @@ -5,7 +5,7 @@ use parking_lot::Mutex; use tokio::sync::{broadcast, mpsc}; use tracing::warn; -use crate::jsonrpc::{JsonRpcNotification, JsonRpcRequest}; +use crate::jsonrpc::{JsonRpcNotification, ReverseRpcRequest}; use crate::types::{SessionEventNotification, SessionId}; /// Per-session channels created by the router during session registration. @@ -13,12 +13,12 @@ pub(crate) struct SessionChannels { /// Filtered `session.event` notifications for this session. pub(crate) notifications: mpsc::UnboundedReceiver, /// Filtered JSON-RPC requests (tool.call, userInput.request, etc.) for this session. - pub(crate) requests: mpsc::UnboundedReceiver, + pub(crate) requests: mpsc::UnboundedReceiver, } struct SessionSenders { notifications: mpsc::UnboundedSender, - requests: mpsc::UnboundedSender, + requests: mpsc::UnboundedSender, } /// Routes notifications and requests by sessionId to per-session channels. @@ -84,7 +84,7 @@ impl SessionRouter { pub(crate) fn ensure_started( &self, notification_tx: &broadcast::Sender, - request_rx: &Mutex>>, + request_rx: &Mutex>>, llm_inference: Option>, github_telemetry: Option, github_token_registry: Arc, diff --git a/rust/src/session.rs b/rust/src/session.rs index b9d217305..58835bfa6 100644 --- a/rust/src/session.rs +++ b/rust/src/session.rs @@ -1609,6 +1609,10 @@ fn spawn_event_loop( .tx .send(Err(ErrorKind::Session(SessionErrorKind::EventLoopClosed).into())); } + requests.close(); + while let Ok(request) = requests.try_recv() { + drop(request); + } } .instrument(span), ) @@ -2344,9 +2348,24 @@ struct RequestDispatchContext<'a> { /// Process a JSON-RPC request from the CLI. async fn handle_request( + session_id: &SessionId, + ctx: RequestDispatchContext<'_>, + request: crate::jsonrpc::ReverseRpcRequest, +) { + let (request, reverse_rpc_trace) = request.into_dispatch(); + let dispatch = handle_request_inner(session_id, ctx, request, reverse_rpc_trace.as_ref()); + if let Some(trace) = &reverse_rpc_trace { + trace.scope(dispatch).await; + } else { + dispatch.await; + } +} + +async fn handle_request_inner( session_id: &SessionId, ctx: RequestDispatchContext<'_>, request: crate::JsonRpcRequest, + reverse_rpc_trace: Option<&crate::jsonrpc::ReverseRpcDispatchGuard>, ) { let sid = session_id.clone(); let client = ctx.client; @@ -2385,7 +2404,15 @@ async fn handle_request( .unwrap_or(Value::Object(Default::default())); let rpc_result = if let Some(hooks) = hooks { - match crate::hooks::dispatch_hook(hooks, &sid, hook_type, input).await { + match crate::hooks::dispatch_hook_traced( + hooks, + &sid, + hook_type, + input, + reverse_rpc_trace.map(|guard| guard.trace()), + ) + .await + { Ok(output) => output, Err(e) => { warn!(error = %e, hook_type = hook_type, "hook dispatch failed");