use std::io::{ErrorKind, Read, Write};
#[cfg(unix)]
use std::os::unix::net::UnixStream;
use std::sync::Arc;
use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH};
use serde_json::Value;
use crate::debug::{describe_endpoint, error_label, join_capabilities, on_off, Category, DebugLog};
use crate::error::Error;
use crate::evidence::{Lease as EvidenceProviderLease, Registry as EvidenceProviderRegistry};
use crate::framing::{encode_frame, FrameDecoder};
use crate::limits::{Limits, DEFAULT_LIMITS};
use crate::logs::{AttrValue, LogLevel, LogRecord, MAX_LOG_ATTRS};
use crate::marker::encode_marker;
use crate::messages::{
default_capabilities, parse_driver_message, Hello, HelloAck, LogMessage, ProbeInfo,
ProtocolErrorMessage, RevisionCommit, SnapshotMessage,
};
use crate::roles::Capability;
use crate::tree::Snapshot;
use crate::validate::validate_snapshot;
#[cfg(unix)]
type TransportStream = UnixStream;
#[cfg(windows)]
use interprocess::{
os::windows::named_pipe::{pipe_mode, DuplexPipeStream},
ConnectWaitMode,
};
#[cfg(windows)]
type TransportStream = DuplexPipeStream<pipe_mode::Bytes>;
pub const ENV_ENDPOINT: &str = "TERMWRIGHT_ENDPOINT";
pub const ENV_TOKEN: &str = "TERMWRIGHT_TOKEN";
pub const DIAL_TIMEOUT: Duration = Duration::from_secs(5);
pub const WRITE_TIMEOUT: Duration = Duration::from_millis(250);
#[derive(Debug, Clone)]
pub struct Options {
pub adapter_name: String,
pub adapter_version: String,
pub capabilities: Vec<Capability>,
pub limits: Limits,
pub write_timeout: Option<Duration>,
pub probe: Option<ProbeInfo>,
pub debug: Option<Arc<DebugLog>>,
pub evidence_registry: Option<EvidenceProviderRegistry>,
}
impl Options {
pub fn with_logs(adapter_name: impl Into<String>, adapter_version: impl Into<String>) -> Self {
let mut options = Self::new(adapter_name, adapter_version);
options.capabilities.push(Capability::Logs);
options
}
pub fn new(adapter_name: impl Into<String>, adapter_version: impl Into<String>) -> Self {
Self {
adapter_name: adapter_name.into(),
adapter_version: adapter_version.into(),
capabilities: default_capabilities(),
limits: DEFAULT_LIMITS,
write_timeout: Some(WRITE_TIMEOUT),
probe: None,
debug: None,
evidence_registry: None,
}
}
}
fn epoch_millis() -> i64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|since| since.as_millis() as i64)
.unwrap_or(0)
}
#[derive(Debug)]
struct TokenBucket {
per_second: f64,
capacity: f64,
tokens: f64,
updated: Instant,
}
impl TokenBucket {
fn new(per_second: i64, burst: i64, now: Instant) -> Self {
let rate = per_second.max(0) as f64;
let capacity = rate + burst.max(0) as f64;
Self {
per_second: rate,
capacity,
tokens: capacity,
updated: now,
}
}
fn take(&mut self, now: Instant) -> bool {
if self.per_second <= 0.0 {
return false;
}
let elapsed = now.saturating_duration_since(self.updated).as_secs_f64();
self.updated = now;
self.tokens = (self.tokens + elapsed * self.per_second).min(self.capacity);
if self.tokens < 1.0 {
return false;
}
self.tokens -= 1.0;
true
}
}
#[derive(Debug)]
pub struct Client {
endpoint: String,
token: String,
options: Options,
stream: Option<TransportStream>,
decoder: FrameDecoder,
limits: Limits,
session_id: Option<String>,
revision: i64,
marker_enabled: bool,
log_budget: Option<crate::messages::LogBudget>,
snapshots_sent: u64,
log_seq: i64,
log_bucket: Option<TokenBucket>,
logs_dropped: u64,
subscribe: String,
evidence_lease: Option<EvidenceProviderLease>,
}
impl Client {
pub fn new(endpoint: impl Into<String>, token: impl Into<String>, options: Options) -> Self {
let limits = options.limits;
Self {
endpoint: endpoint.into(),
token: token.into(),
options,
stream: None,
decoder: FrameDecoder::new(limits.max_frame_bytes, limits.max_depth),
limits,
session_id: None,
revision: 0,
marker_enabled: false,
log_budget: None,
snapshots_sent: 0,
log_seq: 0,
log_bucket: None,
logs_dropped: 0,
subscribe: "snapshots".to_owned(),
evidence_lease: None,
}
}
pub fn from_env(mut options: Options) -> Option<Self> {
if options.debug.is_none() {
options.debug = DebugLog::from_env(&options.adapter_name).map(Arc::new);
}
Self::from_values(
std::env::var(ENV_ENDPOINT).ok().as_deref(),
std::env::var(ENV_TOKEN).ok().as_deref(),
options,
)
}
pub fn from_values(
endpoint: Option<&str>,
token: Option<&str>,
options: Options,
) -> Option<Self> {
let endpoint = endpoint.filter(|value| !value.is_empty());
let token = token.filter(|value| !value.is_empty());
let (Some(endpoint), Some(token)) = (endpoint, token) else {
if let Some(log) = options.debug.as_ref() {
let mut missing = Vec::new();
if endpoint.is_none() {
missing.push(ENV_ENDPOINT);
}
if token.is_none() {
missing.push(ENV_TOKEN);
}
log.line(
Category::Diag,
&format!("dormant: {} not set", missing.join(" and ")),
);
}
return None;
};
if !endpoint_supported(endpoint) {
if let Some(log) = options.debug.as_ref() {
log.line(
Category::Diag,
&format!(
"dormant: {} is not a local endpoint for this platform",
describe_endpoint(endpoint)
),
);
}
return None;
}
Some(Self::new(endpoint, token, options))
}
pub fn connect(&mut self, timeout: Duration) -> Result<(), Error> {
if let Some(probe) = self.options.probe.as_ref() {
probe.validate()?;
}
self.debug_line(
Category::Sem,
&format!(
"dial {} timeout={}ms",
describe_endpoint(&self.endpoint),
timeout.as_millis()
),
);
let stream = match connect_transport(&self.endpoint, timeout, self.options.write_timeout) {
Ok(stream) => stream,
Err(error) => {
self.debug_line(
Category::Diag,
&format!("dial failed, staying dormant: {}", error_label(&error)),
);
return Err(error.into());
}
};
self.stream = Some(stream);
let mut hello = Hello::new(
&self.token,
&self.options.adapter_name,
&self.options.adapter_version,
self.options.capabilities.clone(),
);
if let Some(probe) = self.options.probe.clone() {
hello = hello.with_probe(probe);
}
if let Some(registry) = self.options.evidence_registry.as_ref() {
let lease = registry.freeze();
hello = hello.with_providers(lease.registrations());
self.evidence_lease = Some(lease);
}
self.send(&hello)?;
self.debug_line(
Category::Sem,
&format!(
"hello sent adapter={}/{} caps={}",
self.options.adapter_name,
self.options.adapter_version,
join_capabilities(&self.options.capabilities)
),
);
let deadline = Instant::now() + timeout;
while self.session_id.is_none() {
if Instant::now() >= deadline {
self.debug_line(
Category::Diag,
&format!(
"no hello-ack within {}ms, staying dormant",
timeout.as_millis()
),
);
self.close();
return Err(Error::HandshakeTimeout);
}
self.poll()?;
std::thread::yield_now();
}
Ok(())
}
fn debug_line(&self, category: Category, message: &str) {
if let Some(log) = self.options.debug.as_ref() {
log.line(category, message);
}
}
pub fn connected(&self) -> bool {
self.session_id.is_some() && self.stream.is_some()
}
pub fn session_id(&self) -> Option<&str> {
self.session_id.as_deref()
}
pub fn revision(&self) -> i64 {
self.revision
}
pub fn log_budget(&self) -> Option<crate::messages::LogBudget> {
self.log_budget
}
pub fn limits(&self) -> &Limits {
&self.limits
}
pub fn close(&mut self) {
if let Some(stream) = self.stream.take() {
self.debug_line(
Category::Sem,
&format!(
"close r{} snapshots={} logs_dropped={}",
self.revision, self.snapshots_sent, self.logs_dropped
),
);
close_transport(stream);
}
self.session_id = None;
if let Some(mut lease) = self.evidence_lease.take() {
lease.close();
}
}
pub fn fail(&mut self, code: &str, message: impl Into<String>) -> Result<(), Error> {
let result = self.send(&ProtocolErrorMessage::new(code, message));
self.close();
result
}
pub fn publish(&mut self, snapshot: &mut Snapshot) -> Result<Option<String>, Error> {
self.publish_inner(snapshot)
}
fn publish_inner(&mut self, snapshot: &mut Snapshot) -> Result<Option<String>, Error> {
let Some(session_id) = self.session_id.clone() else {
return Ok(None);
};
if self.stream.is_none() {
return Ok(None);
}
let revision = self.revision + 1;
snapshot.v = 2;
snapshot.session_id = session_id.clone();
snapshot.revision = revision;
if let Some(lease) = self.evidence_lease.as_ref() {
snapshot.provider_evidence =
lease.collect(&session_id, revision, snapshot.columns, snapshot.rows);
}
let body = serde_json::to_string(&snapshot).map_err(|_| {
Error::Protocol(crate::error::Violation::new(
"frame-malformed",
"snapshot is not JSON-serialisable",
))
})?;
let parsed: Value = serde_json::from_str(&body).expect("just serialised");
validate_snapshot(&parsed, &self.limits)?;
let marker = if self.marker_enabled {
Some(encode_marker(&self.token, &session_id, revision)?)
} else {
None
};
let tree_frame = if self.subscribe != "revisions" {
Some(encode_frame(
&SnapshotMessage::new(snapshot),
self.limits.max_frame_bytes,
)?)
} else {
None
};
let commit_frame =
encode_frame(&RevisionCommit::new(revision), self.limits.max_frame_bytes)?;
if let Some(frame) = &tree_frame {
self.write_frame(frame)?;
}
self.write_frame(&commit_frame)?;
self.revision = revision;
if tree_frame.is_some() {
self.snapshots_sent += 1;
}
Ok(marker)
}
pub fn snapshots_sent(&self) -> u64 {
self.snapshots_sent
}
pub fn logs_dropped(&self) -> u64 {
self.logs_dropped
}
pub fn log(&mut self, mut record: LogRecord) -> bool {
if self.session_id.is_none() || self.stream.is_none() || self.log_bucket.is_none() {
return false;
}
let origin = record.seq;
self.log_seq += 1;
record.seq = self.log_seq;
if record.ts == 0 {
record.ts = epoch_millis();
}
if record.revision.is_none() && self.revision > 0 {
record.revision = Some(self.revision);
}
let now = Instant::now();
let allowed = self
.log_bucket
.as_mut()
.is_some_and(|bucket| bucket.take(now));
if !allowed {
self.logs_dropped += 1;
return false;
}
if origin > 0 && record.attrs.len() < MAX_LOG_ATTRS {
record
.attrs
.insert("origin.seq".to_owned(), AttrValue::Int(origin));
if record.validate(&self.limits).is_err() {
record.attrs.remove("origin.seq");
}
}
if record.validate(&self.limits).is_err() {
self.logs_dropped += 1;
return false;
}
self.send(&LogMessage::new(&record)).is_ok()
}
pub fn log_message(&mut self, level: LogLevel, message: impl Into<String>) -> bool {
self.log(LogRecord::new(level, message))
}
pub fn poll(&mut self) -> Result<(), Error> {
let mut buffer = [0u8; 8192];
loop {
let read = match self.stream.as_mut() {
None => return Ok(()),
Some(stream) => read_transport(stream, &mut buffer),
};
match read {
Ok(Incoming::Closed) => {
self.close();
return Ok(());
}
Ok(Incoming::Data(count)) => {
let frames = self.decoder.push(&buffer[..count])?;
for frame in frames {
self.handle(&frame.value)?;
}
}
Ok(Incoming::Idle) => return Ok(()),
Err(error) if error.kind() == ErrorKind::Interrupted => continue,
Err(error) => {
self.close();
return Err(Error::Io(error));
}
}
}
}
fn handle(&mut self, value: &Value) -> Result<(), Error> {
if let Err(error) = parse_driver_message(value, &self.limits) {
self.debug_line(
Category::Diag,
&format!("rejected a driver message: {error}"),
);
let _ = self.send(&ProtocolErrorMessage::new("malformed", error.to_string()));
self.close();
return Err(Error::Parse(error));
}
match value.get("type").and_then(Value::as_str) {
Some("hello-ack") => {
let ack: HelloAck = serde_json::from_value(value.clone()).expect("validated above");
self.session_id = Some(ack.session_id);
self.limits = ack.limits;
self.marker_enabled = ack.marker.enabled;
self.log_budget = ack.logs;
self.log_bucket = match ack.logs {
Some(budget) if budget.enabled => Some(TokenBucket::new(
budget.max_records_per_second,
budget.burst,
Instant::now(),
)),
_ => None,
};
self.subscribe = ack.subscribe;
if let Some(log) = self.options.debug.as_ref() {
let session = self.session_id.clone().unwrap_or_default();
log.set_label(&session);
log.line(
Category::Sem,
&format!(
"hello-ack session={session} marker={} subscribe={} logs={}",
on_off(self.marker_enabled),
self.subscribe,
on_off(self.log_bucket.is_some())
),
);
}
}
Some("error") => {
self.debug_line(
Category::Diag,
&format!(
"driver ended the session: {}",
value.get("code").and_then(Value::as_str).unwrap_or("?")
),
);
self.close();
}
_ => {}
}
Ok(())
}
fn send<T: serde::Serialize>(&mut self, message: &T) -> Result<(), Error> {
let frame = encode_frame(message, self.limits.max_frame_bytes)?;
self.write_frame(&frame)
}
pub(crate) fn write_frame(&mut self, frame: &[u8]) -> Result<(), Error> {
let Some(stream) = self.stream.as_mut() else {
return Ok(());
};
match write_transport_frame(stream, frame, self.options.write_timeout) {
Ok(()) => Ok(()),
Err(error) => {
let timed_out = matches!(
error.kind(),
ErrorKind::WouldBlock | ErrorKind::TimedOut | ErrorKind::Interrupted
);
self.close();
if timed_out {
self.debug_line(
Category::Diag,
"write deadline exceeded; session is unrecoverable",
);
return Err(Error::WriteTimeout);
}
Err(Error::Io(error))
}
}
}
pub(crate) fn accept_queued_publication(&mut self, revision: i64, snapshot_sent: bool) {
self.revision = revision;
if snapshot_sent {
self.snapshots_sent += 1;
}
}
pub(crate) fn take_evidence_lease(&mut self) -> Option<EvidenceProviderLease> {
self.evidence_lease.take()
}
pub(crate) fn publication_config(&self) -> Option<(String, String, Limits, String, bool, i64)> {
Some((
self.token.clone(),
self.session_id.clone()?,
self.limits,
self.subscribe.clone(),
self.marker_enabled,
self.revision,
))
}
#[cfg(all(test, unix))]
pub(crate) fn test_connected(stream: TransportStream) -> Self {
let mut client = Self::new("unused", "test-token", Options::new("queue-test", "1"));
client.stream = Some(stream);
client.session_id = Some("test-session".into());
client.marker_enabled = true;
client
}
}
#[cfg(unix)]
fn endpoint_supported(endpoint: &str) -> bool {
!endpoint.starts_with(r"\\.\pipe\") && !endpoint.starts_with(r"\\?\pipe\")
}
#[cfg(windows)]
fn endpoint_supported(endpoint: &str) -> bool {
endpoint.starts_with(r"\\.\pipe\") || endpoint.starts_with(r"\\?\pipe\")
}
#[cfg(unix)]
fn connect_transport(
endpoint: &str,
_dial_timeout: Duration,
write_timeout: Option<Duration>,
) -> std::io::Result<TransportStream> {
let stream = UnixStream::connect(endpoint)?;
stream.set_read_timeout(Some(Duration::from_millis(50)))?;
stream.set_write_timeout(write_timeout)?;
Ok(stream)
}
#[cfg(windows)]
fn connect_transport(
endpoint: &str,
dial_timeout: Duration,
_write_timeout: Option<Duration>,
) -> std::io::Result<TransportStream> {
let stream = TransportStream::connect_by_path_with_wait_mode(
endpoint,
ConnectWaitMode::Timeout(dial_timeout),
)?;
stream.set_nonblocking(true)?;
Ok(stream)
}
enum Incoming {
Data(usize),
Idle,
Closed,
}
#[cfg(unix)]
fn read_transport(stream: &mut TransportStream, buffer: &mut [u8]) -> std::io::Result<Incoming> {
match stream.read(buffer) {
Ok(0) => Ok(Incoming::Closed),
Ok(count) => Ok(Incoming::Data(count)),
Err(error) if matches!(error.kind(), ErrorKind::WouldBlock | ErrorKind::TimedOut) => {
Ok(Incoming::Idle)
}
Err(error) => Err(error),
}
}
#[cfg(windows)]
const ERROR_NO_DATA: i32 = 232;
#[cfg(windows)]
const ERROR_BROKEN_PIPE: i32 = 109;
#[cfg(windows)]
fn read_transport(stream: &mut TransportStream, buffer: &mut [u8]) -> std::io::Result<Incoming> {
match stream.read(buffer) {
Ok(0) => Ok(Incoming::Idle),
Ok(count) => Ok(Incoming::Data(count)),
Err(error) if matches!(error.kind(), ErrorKind::WouldBlock | ErrorKind::TimedOut) => {
Ok(Incoming::Idle)
}
Err(error) if error.raw_os_error() == Some(ERROR_NO_DATA) => Ok(Incoming::Idle),
Err(error) if error.raw_os_error() == Some(ERROR_BROKEN_PIPE) => Ok(Incoming::Closed),
Err(error) => Err(error),
}
}
#[cfg(unix)]
fn close_transport(stream: TransportStream) {
let _ = stream.shutdown(std::net::Shutdown::Both);
}
#[cfg(windows)]
fn close_transport(_stream: TransportStream) {
}
#[cfg(unix)]
fn write_transport_frame(
stream: &mut TransportStream,
frame: &[u8],
_timeout: Option<Duration>,
) -> std::io::Result<()> {
stream.write_all(frame).and_then(|()| stream.flush())
}
#[cfg(windows)]
fn write_transport_frame(
stream: &mut TransportStream,
frame: &[u8],
timeout: Option<Duration>,
) -> std::io::Result<()> {
let deadline = timeout.map(|duration| Instant::now() + duration);
let mut offset = 0;
while offset < frame.len() {
match stream.write(&frame[offset..]) {
Ok(0) => return Err(std::io::Error::from(ErrorKind::WriteZero)),
Ok(written) => offset += written,
Err(error) if error.kind() == ErrorKind::Interrupted => continue,
Err(error) if error.kind() == ErrorKind::WouldBlock => {
if deadline.is_some_and(|deadline| Instant::now() >= deadline) {
return Err(std::io::Error::from(ErrorKind::TimedOut));
}
std::thread::yield_now();
}
Err(error) => return Err(error),
}
}
loop {
match stream.flush() {
Ok(()) => return Ok(()),
Err(error) if error.kind() == ErrorKind::Interrupted => continue,
Err(error) if error.kind() == ErrorKind::WouldBlock => {
if deadline.is_some_and(|deadline| Instant::now() >= deadline) {
return Err(std::io::Error::from(ErrorKind::TimedOut));
}
std::thread::yield_now();
}
Err(error) => return Err(error),
}
}
}