use std::io::{self, BufWriter, Read, Write};
use std::net::Shutdown;
use std::os::fd::OwnedFd;
use std::os::unix::net::UnixStream;
use std::path::Path;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::mpsc;
use std::time as path_std_time;
use std::time::Duration;
use tau_proto::{
ClientKind, EventName, EventSelector, HarnessInputMessage, Hello, PROTOCOL_VERSION,
PeerInputReader, PeerOutputWriter, Subscribe,
};
use crate::daemon::DaemonHandle;
use crate::peer_exit::PeerExit;
pub(crate) type UiInputReader = PeerInputReader<Box<dyn Read + Send>>;
pub(crate) type UiOutputWriter = PeerOutputWriter<BufWriter<Box<dyn Write + Send>>>;
pub(crate) const UI_SESSION_ADMISSION_TIMEOUT: Duration = Duration::from_secs(10);
pub(crate) struct UiSessionAdmission {
pub(crate) reader: UiInputReader,
pub(crate) harness_protocol_version: Option<tau_proto::ProtocolVersion>,
}
pub(crate) fn connect_ui_client(
socket_path: &Path,
client_name: impl AsRef<str>,
expected_session_id: Option<&tau_proto::SessionId>,
) -> io::Result<(UiInputReader, UiOutputWriter)> {
connect_ui_client_with_version(socket_path, client_name, expected_session_id)
.map(|(reader, writer, _)| (reader, writer))
}
pub(crate) fn connect_ui_client_with_version(
socket_path: &Path,
client_name: impl AsRef<str>,
expected_session_id: Option<&tau_proto::SessionId>,
) -> io::Result<(
UiInputReader,
UiOutputWriter,
Option<tau_proto::ProtocolVersion>,
)> {
let stream = UnixStream::connect(socket_path)?;
let read_stream = stream.try_clone()?;
let shutdown_stream = stream.try_clone()?;
connect_ui_streams_with_shutdown_and_version(
read_stream,
stream,
client_name,
expected_session_id,
Some(shutdown_stream),
UI_SESSION_ADMISSION_TIMEOUT,
)
}
pub(crate) fn connect_ui_client_with_peer_exit(
socket_path: &Path,
client_name: impl AsRef<str>,
expected_session_id: &tau_proto::SessionId,
) -> io::Result<(UiInputReader, UiOutputWriter, Option<PeerExit>)> {
let stream = UnixStream::connect(socket_path)?;
let peer_exit = PeerExit::from_socket(&stream).ok();
let read_stream = stream.try_clone()?;
let shutdown_stream = stream.try_clone()?;
let (reader, writer, _) = connect_ui_streams_with_shutdown_and_version(
read_stream,
stream,
client_name,
Some(expected_session_id),
Some(shutdown_stream),
UI_SESSION_ADMISSION_TIMEOUT,
)?;
Ok((reader, writer, peer_exit))
}
pub(crate) fn connect_ui_client_until(
socket_path: &Path,
client_name: impl AsRef<str>,
expected_session_id: &tau_proto::SessionId,
deadline: std::time::Instant,
) -> io::Result<(UiInputReader, UiOutputWriter)> {
connect_ui_client_until_with_version(socket_path, client_name, expected_session_id, deadline)
.map(|(reader, writer, _)| (reader, writer))
}
pub(crate) fn connect_ui_client_until_with_version(
socket_path: &Path,
client_name: impl AsRef<str>,
expected_session_id: &tau_proto::SessionId,
deadline: std::time::Instant,
) -> io::Result<(
UiInputReader,
UiOutputWriter,
Option<tau_proto::ProtocolVersion>,
)> {
let timeout = deadline.saturating_duration_since(path_std_time::Instant::now());
if timeout.is_zero() {
return Err(io::Error::new(
io::ErrorKind::TimedOut,
"UI request deadline elapsed while connecting",
));
}
let socket = socket2::Socket::new(socket2::Domain::UNIX, socket2::Type::STREAM, None)?;
socket.connect_timeout(&socket2::SockAddr::unix(socket_path)?, timeout)?;
let fd: OwnedFd = socket.into();
let stream: UnixStream = fd.into();
stream.set_write_timeout(Some(timeout))?;
let read_stream = stream.try_clone()?;
connect_ui_streams_with_version(
DeadlineUnixReader {
stream: read_stream,
deadline,
},
DeadlineUnixWriter { stream, deadline },
client_name,
Some(expected_session_id),
)
}
struct DeadlineUnixReader {
stream: UnixStream,
deadline: std::time::Instant,
}
impl Read for DeadlineUnixReader {
fn read(&mut self, buffer: &mut [u8]) -> io::Result<usize> {
loop {
let remaining = self
.deadline
.checked_duration_since(path_std_time::Instant::now())
.ok_or_else(|| {
io::Error::new(io::ErrorKind::TimedOut, "UI request deadline elapsed")
})?;
self.stream
.set_read_timeout(Some(remaining.min(Duration::from_millis(100))))?;
match self.stream.read(buffer) {
Err(error)
if matches!(
error.kind(),
io::ErrorKind::TimedOut | io::ErrorKind::WouldBlock
) =>
{
continue;
}
result => return result,
}
}
}
}
struct DeadlineUnixWriter {
stream: UnixStream,
deadline: std::time::Instant,
}
impl Write for DeadlineUnixWriter {
fn write(&mut self, buffer: &[u8]) -> io::Result<usize> {
let remaining = self
.deadline
.checked_duration_since(path_std_time::Instant::now())
.ok_or_else(|| {
io::Error::new(io::ErrorKind::TimedOut, "UI request deadline elapsed")
})?;
self.stream.set_write_timeout(Some(remaining))?;
self.stream.write(buffer)
}
fn flush(&mut self) -> io::Result<()> {
self.stream.flush()
}
}
#[cfg(test)]
pub(crate) fn connect_ui_streams<R, W>(
reader: R,
writer: W,
client_name: impl AsRef<str>,
expected_session_id: Option<&tau_proto::SessionId>,
) -> io::Result<(UiInputReader, UiOutputWriter)>
where
R: Read + Send + 'static,
W: Write + Send + 'static,
{
connect_ui_streams_with_version(reader, writer, client_name, expected_session_id)
.map(|(reader, writer, _)| (reader, writer))
}
fn connect_ui_streams_with_version<R, W>(
reader: R,
writer: W,
client_name: impl AsRef<str>,
expected_session_id: Option<&tau_proto::SessionId>,
) -> io::Result<(
UiInputReader,
UiOutputWriter,
Option<tau_proto::ProtocolVersion>,
)>
where
R: Read + Send + 'static,
W: Write + Send + 'static,
{
connect_ui_streams_with_shutdown_and_version(
reader,
writer,
client_name,
expected_session_id,
None,
UI_SESSION_ADMISSION_TIMEOUT,
)
}
fn connect_ui_streams_with_shutdown_and_version<R, W>(
reader: R,
writer: W,
client_name: impl AsRef<str>,
expected_session_id: Option<&tau_proto::SessionId>,
shutdown_stream: Option<UnixStream>,
admission_timeout: Duration,
) -> io::Result<(
UiInputReader,
UiOutputWriter,
Option<tau_proto::ProtocolVersion>,
)>
where
R: Read + Send + 'static,
W: Write + Send + 'static,
{
let mut writer =
PeerOutputWriter::new(BufWriter::new(Box::new(writer) as Box<dyn Write + Send>));
send_hello(&mut writer, client_name, expected_session_id)?;
let reader = PeerInputReader::new(Box::new(reader) as Box<dyn Read + Send>);
let (reader, harness_protocol_version) = match expected_session_id {
Some(expected_session_id) => {
let admission = await_ui_session_admission_with_version(
reader,
expected_session_id.clone(),
shutdown_stream,
admission_timeout,
)?;
(admission.reader, admission.harness_protocol_version)
}
None => (reader, None),
};
Ok((reader, writer, harness_protocol_version))
}
pub(crate) fn await_ui_session_admission_with_version(
mut reader: UiInputReader,
expected_session_id: tau_proto::SessionId,
shutdown_stream: Option<UnixStream>,
timeout: Duration,
) -> io::Result<UiSessionAdmission> {
let (sender, receiver) = mpsc::sync_channel(1);
std::thread::spawn(move || {
let result = verify_ui_session_admission(&mut reader, &expected_session_id);
let _ = sender.send((reader, result));
});
match receiver.recv_timeout(timeout) {
Ok((reader, Ok(harness_protocol_version))) => Ok(UiSessionAdmission {
reader,
harness_protocol_version,
}),
Ok((_reader, Err(error))) => Err(error),
Err(mpsc::RecvTimeoutError::Timeout) => {
if let Some(stream) = shutdown_stream {
let _ = stream.shutdown(Shutdown::Both);
}
Err(io::Error::new(
io::ErrorKind::TimedOut,
"timed out waiting for UI session admission",
))
}
Err(mpsc::RecvTimeoutError::Disconnected) => Err(io::Error::new(
io::ErrorKind::UnexpectedEof,
"UI session admission reader exited unexpectedly",
)),
}
}
pub(crate) fn connect_daemon_ui_client(
daemon: &mut DaemonHandle,
client_name: impl AsRef<str>,
expected_session_id: Option<&tau_proto::SessionId>,
) -> io::Result<(UiInputReader, UiOutputWriter)> {
connect_daemon_ui_client_with_timeout(
daemon,
client_name,
expected_session_id,
UI_SESSION_ADMISSION_TIMEOUT,
)
}
pub(crate) fn connect_daemon_ui_client_with_timeout(
daemon: &mut DaemonHandle,
client_name: impl AsRef<str>,
expected_session_id: Option<&tau_proto::SessionId>,
admission_timeout: Duration,
) -> io::Result<(UiInputReader, UiOutputWriter)> {
if let Some(initial_ui) = daemon.take_initial_ui_stdio() {
connect_ui_streams_with_shutdown_and_version(
initial_ui.stdout,
initial_ui.stdin,
client_name,
expected_session_id,
initial_ui.shutdown_stream,
admission_timeout,
)
.map(|(reader, writer, _)| (reader, writer))
} else {
connect_ui_client(&daemon.socket_path(), client_name, expected_session_id)
}
}
pub(crate) fn hello_message(
client_name: tau_proto::ExtensionName,
expected_session_id: Option<&tau_proto::SessionId>,
) -> HarnessInputMessage {
HarnessInputMessage::Hello(Hello {
declaration_inspection: false,
protocol_version: PROTOCOL_VERSION,
client_name,
client_kind: ClientKind::Ui,
expected_session_id: expected_session_id.cloned(),
capabilities: Default::default(),
})
}
pub(crate) fn chat_subscription_selectors() -> Vec<EventSelector> {
use EventName as E;
vec![
EventSelector::Exact(E::UI_PROMPT_SUBMITTED),
EventSelector::Exact(E::UI_SHELL_COMMAND),
EventSelector::Exact(E::UI_CANCEL_PROMPT),
EventSelector::Exact(E::ACTION_SCHEMA_PUBLISHED),
EventSelector::Exact(E::ACTION_RESULT),
EventSelector::Exact(E::ACTION_ERROR),
EventSelector::Exact(E::AGENT_START_REQUEST),
EventSelector::Exact(E::AGENT_START_ACCEPTED),
EventSelector::Exact(E::AGENT_START_FAILED),
EventSelector::Exact(E::AGENT_START_RESULT),
EventSelector::Exact(E::AGENT_MESSAGE_SENT),
EventSelector::Exact(E::AGENT_MESSAGE_RECEIVED),
EventSelector::Exact(E::MESSAGE_DELIVERED),
EventSelector::Exact(E::MESSAGE_EDITED),
EventSelector::Exact(E::MESSAGE_DELETED),
EventSelector::Exact(E::MESSAGE_REACTION_ADDED),
EventSelector::Exact(E::MESSAGE_REACTION_REMOVED),
EventSelector::Exact(E::MESSAGE_SENT),
EventSelector::Exact(E::AGENT_PROMPT_SUBMITTED),
EventSelector::Exact(E::AGENT_PROMPT_QUEUED),
EventSelector::Exact(E::AGENT_PROMPT_RECALLED),
EventSelector::Exact(E::AGENT_PROMPT_REJECTED),
EventSelector::Exact(E::AGENT_PROMPT_STEERED),
EventSelector::Exact(E::AGENT_COMPACTION_TRIGGERED),
EventSelector::Exact(E::AGENT_MANUAL_COMPACTION_REQUESTED),
EventSelector::Exact(E::AGENT_STANDALONE_COMPACTION_STARTED),
EventSelector::Exact(E::AGENT_STANDALONE_COMPACTION_FAILED),
EventSelector::Exact(E::AGENT_COMPACTED),
EventSelector::Exact(E::AGENT_INFERENCE_DISPATCH_STARTED),
EventSelector::Exact(E::AGENT_PROMPT_STARTED),
EventSelector::Exact(E::AGENT_PROMPT_TERMINATED),
EventSelector::Exact(E::AGENT_PROMPT_FAILED),
EventSelector::Exact(E::AGENT_WATCHES_UPDATED),
EventSelector::Exact(E::AGENT_STATS_UPDATED),
EventSelector::Exact(E::AGENT_STARTED),
EventSelector::Exact(E::AGENT_DISPLAY_NAME_SET),
EventSelector::Exact(E::SESSION_STARTED),
EventSelector::Exact(E::SESSION_SHUTDOWN),
EventSelector::Exact(E::SESSION_AGENT_UNLOADED),
EventSelector::Exact(E::PROVIDER_PROMPT_SUBMITTED),
EventSelector::Exact(E::PROVIDER_RESPONSE_UPDATED),
EventSelector::Exact(E::PROVIDER_RESPONSE_FINISHED),
EventSelector::Exact(E::TOOL_STARTED),
EventSelector::Exact(E::TOOL_REJECTED),
EventSelector::Exact(E::TOOL_RESULT_DISPLAY),
EventSelector::Exact(E::TOOL_ERROR),
EventSelector::Exact(E::TOOL_BACKGROUND_RESULT_DISPLAY),
EventSelector::Exact(E::TOOL_BACKGROUND_ERROR),
EventSelector::Exact(E::TOOL_PROGRESS),
EventSelector::Exact(E::TOOL_CANCELLED),
EventSelector::Exact(E::SHELL_COMMAND_PROGRESS),
EventSelector::Exact(E::SHELL_COMMAND_FINISHED),
EventSelector::Exact(E::EXTENSION_STARTING),
EventSelector::Exact(E::EXTENSION_READY),
EventSelector::Exact(E::EXTENSION_EXITED),
EventSelector::Exact(E::HARNESS_SESSION_SKILLS_AVAILABLE),
EventSelector::Exact(E::HARNESS_AGENT_CONTEXT_INITIALIZED),
EventSelector::Exact(E::EXTENSION_CONTEXT_READY),
EventSelector::Exact(E::HARNESS_NOTICE),
EventSelector::Exact(E::HARNESS_SESSION_DIR),
EventSelector::Exact(E::HARNESS_UI_DIR),
EventSelector::Exact(E::HARNESS_MODELS_AVAILABLE),
EventSelector::Exact(E::HARNESS_ROLES_AVAILABLE),
EventSelector::Exact(E::HARNESS_ROLE_SELECTED),
EventSelector::Exact(E::HARNESS_CONTEXT_USAGE_CHANGED),
EventSelector::Exact(E::HARNESS_AGENT_CONTEXT_USAGE_CHANGED),
EventSelector::Exact(E::HARNESS_PROVIDER_QUOTA_CHANGED),
EventSelector::Exact(E::HARNESS_EFFORTS_AVAILABLE),
EventSelector::Exact(E::HARNESS_VERBOSITIES_AVAILABLE),
EventSelector::Exact(E::HARNESS_THINKING_SUMMARIES_AVAILABLE),
EventSelector::Exact(E::TERM_OSC1337_SET_USER_VAR),
EventSelector::Exact(E::TERM_BELL),
]
}
pub(crate) fn subscribe_message(selectors: Vec<EventSelector>) -> HarnessInputMessage {
HarnessInputMessage::Subscribe(Subscribe {
historical_selectors: selectors.clone(),
live_selectors: selectors,
})
}
pub(crate) fn chat_subscribe_message() -> HarnessInputMessage {
let live_selectors = chat_subscription_selectors();
let live_only = [
EventName::AGENT_PROMPT_STARTED,
EventName::AGENT_PROMPT_TERMINATED,
EventName::AGENT_PROMPT_FAILED,
EventName::AGENT_PROMPT_REJECTED,
EventName::AGENT_PROMPT_RECALLED,
EventName::PROVIDER_PROMPT_SUBMITTED,
EventName::PROVIDER_RESPONSE_UPDATED,
EventName::TOOL_PROGRESS,
EventName::SHELL_COMMAND_PROGRESS,
EventName::TERM_OSC1337_SET_USER_VAR,
EventName::TERM_BELL,
];
let historical_selectors = live_selectors
.iter()
.filter(
|selector| !matches!(selector, EventSelector::Exact(name) if live_only.contains(name)),
)
.cloned()
.chain(std::iter::once(EventSelector::Exact(
EventName::PROVIDER_TOOL_ERROR,
)))
.collect();
HarnessInputMessage::Subscribe(Subscribe {
historical_selectors,
live_selectors,
})
}
pub(crate) fn send_hello(
writer: &mut UiOutputWriter,
client_name: impl AsRef<str>,
expected_session_id: Option<&tau_proto::SessionId>,
) -> io::Result<()> {
let client_name = tau_proto::ExtensionName::parse(client_name.as_ref().to_owned())
.map_err(io::Error::other)?;
send_message(writer, &hello_message(client_name, expected_session_id))
}
pub(crate) fn verify_ui_session_admission<R: Read>(
reader: &mut PeerInputReader<R>,
expected_session_id: &tau_proto::SessionId,
) -> io::Result<Option<tau_proto::ProtocolVersion>> {
match reader.read_message().map_err(io::Error::other)? {
Some(tau_proto::HarnessOutputMessage::SessionAccepted(accepted))
if accepted.session_id == *expected_session_id =>
{
Ok(accepted.harness_protocol_version)
}
Some(tau_proto::HarnessOutputMessage::SessionAccepted(accepted)) => {
Err(io::Error::other(format!(
"session target mismatch: requested `{expected_session_id}`, but the connected \
harness admitted `{}`",
accepted.session_id
)))
}
Some(tau_proto::HarnessOutputMessage::Disconnect(disconnect)) => {
Err(io::Error::other(disconnect.reason.unwrap_or_else(|| {
"harness rejected UI session admission".to_owned()
})))
}
Some(other) => Err(io::Error::other(format!(
"harness sent {other:?} before UI session admission"
))),
None => Err(io::Error::other(
"harness closed before confirming UI session admission",
)),
}
}
pub(crate) fn subscribe(
writer: &mut UiOutputWriter,
selectors: Vec<EventSelector>,
) -> io::Result<()> {
send_message(writer, &subscribe_message(selectors))
}
pub(crate) fn send_message(
writer: &mut UiOutputWriter,
message: &HarnessInputMessage,
) -> io::Result<()> {
writer.write_message(message).map_err(io::Error::other)?;
writer.flush()
}
pub(crate) fn next_request_id(prefix: &str) -> String {
static COUNTER: AtomicU64 = AtomicU64::new(0);
format!(
"{}-{}-{}",
prefix,
std::process::id(),
COUNTER.fetch_add(1, Ordering::Relaxed)
)
}
#[cfg(test)]
mod tests;