use std::collections::VecDeque;
use std::path::{Path, PathBuf};
use std::time::Duration;
use std::time::Instant;
use blitz_control_protocol::{
AgentControlRequest, AgentSnapshot, DEBUG_PROTOCOL_VERSION, DebugDescriptor, DebugEvent,
DebugProtocolError, DebugResponse, DebugStream, DiagnosticsRequest, JsonRpcId, JsonRpcMessage,
JsonRpcRequest, MCP_PROTOCOL_VERSION, MessageStream, TransportStream, WireMessage,
decode_diagnostics_event_value, decode_response_value, decode_wire_value, encode_agent_request,
encode_diagnostics_request, encode_rpc, framed_json, peek_value_request_id,
};
use eyre::{Context, Result, bail, eyre};
use tokio::net::UnixStream;
use tokio::time::timeout;
const REQUEST_TIMEOUT: Duration = Duration::from_secs(60);
const MAX_QUEUED_EVENTS: usize = 256;
#[derive(Debug)]
pub struct Descriptor {
pub path: PathBuf,
pub descriptor: DebugDescriptor,
pub raw: serde_json::Value,
verified_reachable: bool,
}
impl Descriptor {
pub fn socket_path(&self) -> PathBuf {
match self.descriptor.address.strip_prefix("unix://") {
Some(path) => PathBuf::from(path),
None => self.path.with_extension("sock"),
}
}
fn is_reachable(&self) -> bool {
std::os::unix::net::UnixStream::connect(self.socket_path()).is_ok()
}
pub fn warn_if_stale(&self) {
if !self.verified_reachable && !self.is_reachable() {
eprintln!(
"warning: descriptor {} names pid {}, but its control socket is unreachable",
self.path.display(),
self.descriptor.pid
);
}
}
}
pub fn discover(explicit: Option<&str>) -> Result<Descriptor> {
if let Some(path) = explicit {
let path = PathBuf::from(path);
if path.exists() {
return read_descriptor(&path);
}
}
let pinned = PathBuf::from("target/blitz-control.json");
if pinned.exists()
&& let Ok(mut descriptor) = read_descriptor(&pinned)
&& descriptor.is_reachable()
{
descriptor.verified_reachable = true;
return Ok(descriptor);
}
let root = PathBuf::from(std::env::var("TMPDIR").unwrap_or_else(|_| "/tmp".into()))
.join("tauri-blitz-agent");
let mut found: Vec<(std::time::SystemTime, PathBuf)> = std::fs::read_dir(&root)
.into_iter()
.flatten()
.flatten()
.map(|entry| entry.path())
.filter(|path| path.extension().is_some_and(|ext| ext == "json"))
.filter_map(|path| Some((path.metadata().ok()?.modified().ok()?, path)))
.collect();
found.sort();
for (_, path) in found.iter().rev() {
let Ok(mut descriptor) = read_descriptor(path) else {
continue;
};
if descriptor.is_reachable() {
descriptor.verified_reachable = true;
return Ok(descriptor);
}
}
bail!(
"no reachable inspector descriptor found; is a diagnostics build running?\n\
looked at target/blitz-control.json and {}. Pass --descriptor \
<path> to inspect a specific descriptor.",
root.display()
)
}
fn read_descriptor(path: &Path) -> Result<Descriptor> {
let text = std::fs::read_to_string(path)
.with_context(|| format!("reading descriptor {}", path.display()))?;
let raw: serde_json::Value = serde_json::from_str(&text)
.with_context(|| format!("parsing descriptor {}", path.display()))?;
let descriptor: DebugDescriptor = serde_json::from_value(raw.clone())
.with_context(|| format!("descriptor {} is not a DebugDescriptor", path.display()))?;
if descriptor.protocol_version != DEBUG_PROTOCOL_VERSION {
bail!(
"descriptor {} uses debug protocol {}, but ps-qa requires {}",
path.display(),
descriptor.protocol_version,
DEBUG_PROTOCOL_VERSION
);
}
Ok(Descriptor {
path: path.to_path_buf(),
descriptor,
raw,
verified_reachable: false,
})
}
pub struct Client {
stream: Box<dyn MessageStream>,
next_id: i64,
request_timeout: Duration,
events: VecDeque<DebugEvent>,
}
#[derive(Debug)]
struct InspectorResponseError {
code: String,
message: String,
}
impl std::fmt::Display for InspectorResponseError {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
formatter,
"inspector returned {}: {}",
self.code, self.message
)
}
}
impl std::error::Error for InspectorResponseError {}
impl Client {
fn queue_event(&mut self, event: DebugEvent) {
if self.events.len() == MAX_QUEUED_EVENTS {
self.events.pop_front();
}
self.events.push_back(event);
}
pub async fn connect(socket: &Path) -> Result<Self> {
const CONNECT_DEADLINE: Duration = Duration::from_millis(500);
const RETRY_DELAY: Duration = Duration::from_millis(20);
let started = Instant::now();
let stream = loop {
match UnixStream::connect(socket).await {
Ok(stream) => break stream,
Err(_) if started.elapsed() < CONNECT_DEADLINE => {
tokio::time::sleep(RETRY_DELAY).await;
}
Err(error) => {
return Err(error)
.with_context(|| format!("connecting to {}", socket.display()));
}
}
};
Ok(Self {
stream: Box::new(TransportStream::new(framed_json(stream))),
next_id: 0,
request_timeout: REQUEST_TIMEOUT,
events: VecDeque::new(),
})
}
pub fn set_request_timeout(&mut self, request_timeout: Duration) {
self.request_timeout = request_timeout;
}
pub fn request_timeout(&self) -> Duration {
self.request_timeout
}
fn next_id(&mut self) -> JsonRpcId {
self.next_id += 1;
JsonRpcId::Number(self.next_id)
}
async fn exchange(
&mut self,
request: WireMessage,
id: &JsonRpcId,
) -> Result<serde_json::Value> {
self.stream
.send(request)
.await
.map_err(|error| eyre!("sending to the inspector failed: {error}"))?;
loop {
let message = timeout(self.request_timeout, self.stream.recv())
.await
.map_err(|_| {
eyre!(
"the inspector did not answer within {:?}",
self.request_timeout
)
})?
.ok_or_else(|| eyre!("the inspector closed the connection"))?
.map_err(|error| eyre!("reading from the inspector failed: {error}"))?;
let value = decode_wire_value(message).map_err(protocol_error)?;
if peek_value_request_id(&value).as_ref() == Some(id) {
return Ok(value);
}
if let Ok(event) = decode_diagnostics_event_value(value) {
self.queue_event(event);
}
}
}
pub async fn raw_request(
&mut self,
method: &str,
params: serde_json::Value,
) -> Result<serde_json::Value> {
let id = self.next_id();
let request = encode_rpc(JsonRpcMessage::Request(JsonRpcRequest::call(
id.clone(),
method,
params,
)))
.map_err(protocol_error)?;
self.exchange(request, &id).await
}
pub async fn initialize(&mut self) -> Result<serde_json::Value> {
self.raw_request(
"initialize",
serde_json::json!({"protocolVersion": MCP_PROTOCOL_VERSION}),
)
.await
}
pub async fn tools_list(&mut self) -> Result<serde_json::Value> {
self.raw_request("tools/list", serde_json::json!({})).await
}
pub async fn agent(&mut self, request: &AgentControlRequest) -> Result<Answer> {
let id = self.next_id();
let frame = encode_agent_request(id.clone(), request).map_err(protocol_error)?;
let value = self.exchange(frame, &id).await?;
Answer::new(value)
}
pub async fn diagnostics(&mut self, request: &DiagnosticsRequest) -> Result<Answer> {
let id = self.next_id();
let frame = encode_diagnostics_request(id.clone(), request).map_err(protocol_error)?;
let value = self.exchange(frame, &id).await?;
Answer::new(value)
}
pub async fn arm_paint_events(&mut self) -> Result<bool> {
self.events
.retain(|event| !matches!(event, DebugEvent::PaintCommitted { .. }));
match self
.diagnostics(&DiagnosticsRequest::Observe {
streams: vec![DebugStream::Paint],
})
.await
{
Ok(_) => Ok(true),
Err(error)
if error
.downcast_ref::<InspectorResponseError>()
.is_some_and(|error| error.code == "streamingUnavailable") =>
{
Ok(false)
}
Err(error) => Err(error),
}
}
pub async fn wait_for_paint(&mut self, within: Duration) -> Result<bool> {
if let Some(index) = self
.events
.iter()
.position(|event| matches!(event, DebugEvent::PaintCommitted { .. }))
{
self.events.remove(index);
return Ok(true);
}
let deadline = tokio::time::Instant::now() + within;
loop {
let remaining = deadline.saturating_duration_since(tokio::time::Instant::now());
if remaining.is_zero() {
return Ok(false);
}
let message = match timeout(remaining, self.stream.recv()).await {
Ok(Some(message)) => message,
Ok(None) | Err(_) => return Ok(false),
};
let message = message
.map_err(|error| eyre!("reading paint event from inspector failed: {error}"))?;
let value = decode_wire_value(message).map_err(protocol_error)?;
match decode_diagnostics_event_value(value) {
Ok(DebugEvent::PaintCommitted { .. }) => return Ok(true),
Ok(event) => self.queue_event(event),
Err(_) => {}
}
}
}
pub async fn diagnostics_envelope(
&mut self,
request: &DiagnosticsRequest,
) -> Result<serde_json::Value> {
let id = self.next_id();
let frame = encode_diagnostics_request(id.clone(), request).map_err(protocol_error)?;
self.exchange(frame, &id).await
}
pub async fn agent_envelope(
&mut self,
request: &AgentControlRequest,
) -> Result<serde_json::Value> {
let id = self.next_id();
let frame = encode_agent_request(id.clone(), request).map_err(protocol_error)?;
self.exchange(frame, &id).await
}
pub async fn drain(&mut self, seconds: f64) -> Result<Vec<serde_json::Value>> {
let deadline = tokio::time::Instant::now() + Duration::from_secs_f64(seconds);
let mut out = Vec::new();
loop {
let remaining = deadline.saturating_duration_since(tokio::time::Instant::now());
if remaining.is_zero() {
return Ok(out);
}
match timeout(remaining, self.stream.recv()).await {
Err(_) => return Ok(out),
Ok(None) => return Ok(out),
Ok(Some(Ok(message))) => {
out.push(decode_wire_value(message).map_err(protocol_error)?)
}
Ok(Some(Err(error))) => bail!("reading from the inspector failed: {error}"),
}
}
}
}
pub async fn inspect(client: &mut Client) -> Result<(AgentSnapshot, f64)> {
inspect_from(client, None).await
}
pub async fn inspect_subtree(client: &mut Client, root: u64) -> Result<(AgentSnapshot, f64)> {
inspect_from(client, Some(root)).await
}
async fn inspect_from(client: &mut Client, root: Option<u64>) -> Result<(AgentSnapshot, f64)> {
let started = Instant::now();
let answer = client
.agent(&AgentControlRequest::Inspect {
root,
max_depth: 40,
})
.await?;
let elapsed = started.elapsed().as_secs_f64() * 1000.0;
match answer.response {
DebugResponse::AgentSnapshot(snapshot) => Ok((snapshot, elapsed)),
other => bail!("asked for a semantic snapshot, got {other:?}"),
}
}
pub struct Answer {
pub response: DebugResponse,
}
impl Answer {
fn new(value: serde_json::Value) -> Result<Self> {
let (_, response) = decode_response_value(value).map_err(protocol_error)?;
if let DebugResponse::Error(error) = &response {
return Err(eyre::Report::new(InspectorResponseError {
code: error.code.clone(),
message: error.message.clone(),
}));
}
Ok(Self { response })
}
}
fn protocol_error(error: DebugProtocolError) -> eyre::Report {
eyre!("{error}")
}
#[cfg(test)]
mod tests {
use super::*;
use blitz_control_protocol::{JsonRpcResponse, encode_diagnostics_event};
use tokio::net::UnixListener;
#[tokio::test(flavor = "current_thread")]
async fn connect_retries_startup_and_exchange_preserves_events() {
let socket = std::env::temp_dir().join(format!(
"ps-qa-transport-{}-{}.sock",
std::process::id(),
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.expect("system clock is after the epoch")
.as_nanos()
));
let server_socket = socket.clone();
let server = async move {
tokio::time::sleep(Duration::from_millis(40)).await;
let listener = UnixListener::bind(&server_socket).expect("bind test socket");
let (stream, _) = listener.accept().await.expect("accept test client");
let mut stream = TransportStream::new(framed_json(stream));
let _request = stream
.recv()
.await
.expect("client keeps connection open")
.expect("read initialize request");
stream
.send(
encode_diagnostics_event(&DebugEvent::PaintCommitted { revision: 7 })
.expect("encode paint event"),
)
.await
.expect("send paint event");
stream
.send(
encode_rpc(JsonRpcMessage::Response(JsonRpcResponse::result(
Some(JsonRpcId::Number(1)),
serde_json::json!({"protocolVersion": MCP_PROTOCOL_VERSION}),
)))
.expect("encode initialize response"),
)
.await
.expect("send initialize response");
};
let client_socket = socket.clone();
let client = async move {
let mut client = Client::connect(&client_socket)
.await
.expect("client retries until socket is bound");
client.initialize().await.expect("initialize completes");
assert!(
client
.wait_for_paint(Duration::ZERO)
.await
.expect("queued paint remains readable"),
"an event arriving before the response must not steal the response or be discarded"
);
};
tokio::join!(server, client);
let _ = std::fs::remove_file(socket);
}
#[test]
fn descriptor_protocol_mismatch_is_rejected_before_connecting() {
let path = std::env::temp_dir().join(format!(
"ps-qa-descriptor-version-{}-{}.json",
std::process::id(),
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.expect("system clock is after the epoch")
.as_nanos()
));
std::fs::write(
&path,
serde_json::json!({
"protocolVersion": DEBUG_PROTOCOL_VERSION + 1,
"pid": 1,
"instanceId": "fixture",
"address": "unix:///tmp/fixture.sock",
"renderer": "fixture",
"rendererRevision": "test"
})
.to_string(),
)
.expect("write descriptor fixture");
let error = read_descriptor(&path).expect_err("newer protocol must not be guessed");
assert!(error.to_string().contains("requires"));
let _ = std::fs::remove_file(path);
}
}