use std::collections::HashMap;
use std::process::Stdio;
use std::sync::{Arc, Mutex as StdMutex};
use std::time::Duration;
use async_trait::async_trait;
use contextgraph_types::{
Capabilities, ContextQuery, ContextQueryResult, PROTOCOL_VERSION, ProviderInfo, VerifyRequest,
VerifyResponse,
};
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
use tokio::process::{Child, ChildStdin, ChildStdout, Command};
use tokio::sync::{Mutex as TokioMutex, oneshot};
use tokio::task::JoinHandle;
use crate::error::HostError;
use crate::provider::ContextProvider;
use crate::wire::{
Envelope, decode_line, encode_line, envelope_kind, next_correlation_id, verify_correlation,
versions_compatible,
};
const HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(10);
const SHUTDOWN_GRACE: Duration = Duration::from_secs(2);
const MAX_LINE_BYTES: usize = 16 * 1024 * 1024;
async fn read_framed_line(
stdout: &mut BufReader<ChildStdout>,
label: &str,
) -> Result<Option<String>, HostError> {
let transport = |message: String| HostError::Transport {
id: label.to_string(),
message,
};
let mut bytes: Vec<u8> = Vec::new();
loop {
let buf = stdout
.fill_buf()
.await
.map_err(|e| transport(e.to_string()))?;
if buf.is_empty() {
break; }
if let Some(pos) = buf.iter().position(|&b| b == b'\n') {
bytes.extend_from_slice(&buf[..=pos]);
stdout.consume(pos + 1);
break;
}
bytes.extend_from_slice(buf);
let consumed = buf.len();
stdout.consume(consumed);
if bytes.len() > MAX_LINE_BYTES {
return Err(transport(format!(
"provider emitted a line exceeding {MAX_LINE_BYTES} bytes without a newline"
)));
}
}
if bytes.is_empty() {
Ok(None)
} else {
Ok(Some(String::from_utf8_lossy(&bytes).into_owned()))
}
}
async fn write_framed_line(
stdin: &mut ChildStdin,
line: &str,
label: &str,
) -> Result<(), HostError> {
let transport = |e: std::io::Error| match e.kind() {
std::io::ErrorKind::BrokenPipe => HostError::ProviderCrashed {
id: label.to_string(),
},
_ => HostError::Transport {
id: label.to_string(),
message: e.to_string(),
},
};
stdin.write_all(line.as_bytes()).await.map_err(transport)?;
if !line.ends_with('\n') {
stdin.write_all(b"\n").await.map_err(transport)?;
}
stdin.flush().await.map_err(transport)?;
Ok(())
}
async fn write_envelope(
stdin: &mut ChildStdin,
env: &Envelope,
label: &str,
) -> Result<(), HostError> {
let line = encode_line(env)?;
write_framed_line(stdin, &line, label).await
}
pub struct RawStdioConnection {
stdin: ChildStdin,
stdout: BufReader<ChildStdout>,
child: Child,
#[cfg_attr(not(unix), allow(dead_code))]
pgid: Option<i32>,
label: String,
}
impl RawStdioConnection {
pub async fn spawn(program: &str, args: &[String]) -> Result<Self, HostError> {
let mut cmd = Command::new(program);
cmd.args(args);
cmd.stdin(Stdio::piped());
cmd.stdout(Stdio::piped());
cmd.stderr(Stdio::inherit());
cmd.kill_on_drop(true);
cmd.env_clear();
if let Ok(path) = std::env::var("PATH") {
cmd.env("PATH", path);
}
if let Ok(home) = std::env::var("HOME") {
cmd.env("HOME", home);
}
#[cfg(unix)]
{
unsafe {
cmd.pre_exec(|| {
libc::setsid();
Ok(())
});
}
}
let mut child = cmd
.spawn()
.map_err(|e| HostError::Spawn(format!("{program}: {e}")))?;
let stdin = child
.stdin
.take()
.ok_or_else(|| HostError::Spawn(format!("{program}: child has no stdin pipe")))?;
let stdout = child
.stdout
.take()
.ok_or_else(|| HostError::Spawn(format!("{program}: child has no stdout pipe")))?;
#[cfg(unix)]
let pgid = child.id().map(|id| id as i32);
#[cfg(not(unix))]
let pgid = None;
Ok(Self {
stdin,
stdout: BufReader::new(stdout),
child,
pgid,
label: program.to_string(),
})
}
pub fn with_label(mut self, label: impl Into<String>) -> Self {
self.label = label.into();
self
}
pub async fn send(&mut self, env: &Envelope) -> Result<(), HostError> {
let line = encode_line(env)?;
self.send_raw_line(&line).await
}
pub async fn send_raw_line(&mut self, line: &str) -> Result<(), HostError> {
write_framed_line(&mut self.stdin, line, &self.label).await
}
pub async fn read_raw_line(&mut self) -> Result<Option<String>, HostError> {
read_framed_line(&mut self.stdout, &self.label).await
}
pub async fn recv(&mut self) -> Result<Envelope, HostError> {
match self.read_raw_line().await? {
Some(line) => decode_line(&line),
None => Err(HostError::ProviderCrashed {
id: self.label.clone(),
}),
}
}
pub async fn handshake(&mut self) -> Result<(ProviderInfo, Capabilities), HostError> {
self.send(&Envelope::Handshake {
protocol_version: PROTOCOL_VERSION.to_string(),
})
.await?;
let ack = match tokio::time::timeout(HANDSHAKE_TIMEOUT, self.recv()).await {
Ok(result) => result?,
Err(_) => {
return Err(HostError::Timeout {
id: self.label.clone(),
timeout_ms: HANDSHAKE_TIMEOUT.as_millis() as u64,
});
}
};
match ack {
Envelope::HandshakeAck {
protocol_version,
provider,
capabilities,
} => {
if !versions_compatible(PROTOCOL_VERSION, &protocol_version) {
return Err(HostError::VersionMismatch {
host: PROTOCOL_VERSION.to_string(),
provider: provider.name,
provider_version: protocol_version,
});
}
Ok((provider, capabilities))
}
other => Err(HostError::UnexpectedEnvelope {
id: self.label.clone(),
expected: "handshake_ack".into(),
got: envelope_kind(&other).into(),
}),
}
}
pub async fn shutdown(&mut self) -> Result<(), HostError> {
let _ = self.send(&Envelope::Shutdown).await;
let label = self.label.clone();
match tokio::time::timeout(SHUTDOWN_GRACE, self.child.wait()).await {
Ok(Ok(_)) => Ok(()),
Ok(Err(e)) => Err(HostError::Transport {
id: label,
message: e.to_string(),
}),
Err(_) => {
self.kill_group();
Ok(())
}
}
}
fn kill_group(&mut self) {
#[cfg(unix)]
if let Some(pgid) = self.pgid {
unsafe {
libc::kill(-pgid, libc::SIGKILL);
}
}
let _ = self.child.start_kill();
}
fn into_parts(self) -> (ChildStdin, BufReader<ChildStdout>, StdioControl) {
let this = std::mem::ManuallyDrop::new(self);
unsafe {
let stdin = std::ptr::read(&this.stdin);
let stdout = std::ptr::read(&this.stdout);
let child = std::ptr::read(&this.child);
let label = std::ptr::read(&this.label);
let pgid = this.pgid;
(stdin, stdout, StdioControl { child, pgid, label })
}
}
}
impl Drop for RawStdioConnection {
fn drop(&mut self) {
self.kill_group();
}
}
struct StdioControl {
child: Child,
#[cfg_attr(not(unix), allow(dead_code))]
pgid: Option<i32>,
label: String,
}
impl StdioControl {
async fn wait_or_kill(&mut self) -> Result<(), HostError> {
match tokio::time::timeout(SHUTDOWN_GRACE, self.child.wait()).await {
Ok(Ok(_)) => Ok(()),
Ok(Err(e)) => Err(HostError::Transport {
id: self.label.clone(),
message: e.to_string(),
}),
Err(_) => {
self.kill_group();
Ok(())
}
}
}
fn kill_group(&mut self) {
#[cfg(unix)]
if let Some(pgid) = self.pgid {
unsafe {
libc::kill(-pgid, libc::SIGKILL);
}
}
let _ = self.child.start_kill();
}
}
impl Drop for StdioControl {
fn drop(&mut self) {
self.kill_group();
}
}
type Reply = Result<Envelope, HostError>;
type PendingTable = Arc<StdMutex<HashMap<String, oneshot::Sender<Reply>>>>;
type NoIdSlot = Arc<StdMutex<Option<oneshot::Sender<Reply>>>>;
enum ReaderExit {
Crashed,
Transport(String),
Decode(String),
}
impl ReaderExit {
fn error(&self, label: &str) -> HostError {
match self {
ReaderExit::Crashed => HostError::ProviderCrashed {
id: label.to_string(),
},
ReaderExit::Transport(message) => HostError::Transport {
id: label.to_string(),
message: message.clone(),
},
ReaderExit::Decode(message) => HostError::Wire(message.clone()),
}
}
}
async fn run_reader(
mut stdout: BufReader<ChildStdout>,
label: String,
pending: PendingTable,
no_id_slot: NoIdSlot,
) {
loop {
let exit = match read_framed_line(&mut stdout, &label).await {
Ok(Some(line)) => match decode_line(&line) {
Ok(env) => {
dispatch(env, &pending, &no_id_slot, &label);
continue;
}
Err(err) => ReaderExit::Decode(err.to_string()),
},
Ok(None) => ReaderExit::Crashed,
Err(HostError::Transport { message, .. }) => ReaderExit::Transport(message),
Err(other) => ReaderExit::Transport(other.to_string()),
};
drain_waiters(&pending, &no_id_slot, &exit, &label);
return;
}
}
fn dispatch(env: Envelope, pending: &PendingTable, no_id_slot: &NoIdSlot, label: &str) {
let correlated = match &env {
Envelope::Frames { id: Some(id), .. } | Envelope::Error { id: Some(id), .. } => {
Some(id.clone())
}
_ => None,
};
if let Some(id) = correlated {
let waiter = pending.lock().expect("pending mutex poisoned").remove(&id);
match waiter {
Some(tx) => {
let _ = tx.send(Ok(env));
}
None => eprintln!(
"contextgraph-host: stdio provider `{label}` sent a reply with id `{id}` matching no in-flight query; dropping"
),
}
return;
}
let waiter = no_id_slot.lock().expect("no_id_slot mutex poisoned").take();
match waiter {
Some(tx) => {
let _ = tx.send(Ok(env));
}
None => eprintln!(
"contextgraph-host: stdio provider `{label}` sent an unsolicited `{}` envelope with no in-flight lock-step exchange; dropping",
envelope_kind(&env)
),
}
}
fn drain_waiters(pending: &PendingTable, no_id_slot: &NoIdSlot, exit: &ReaderExit, label: &str) {
let waiters: Vec<oneshot::Sender<Reply>> = {
let mut map = pending.lock().expect("pending mutex poisoned");
map.drain().map(|(_, tx)| tx).collect()
};
for tx in waiters {
let _ = tx.send(Err(exit.error(label)));
}
let leftover = { no_id_slot.lock().expect("no_id_slot mutex poisoned").take() };
if let Some(tx) = leftover {
let _ = tx.send(Err(exit.error(label)));
}
}
pub struct StdioProvider {
id: String,
info: ProviderInfo,
capabilities: Capabilities,
stdin: TokioMutex<ChildStdin>,
pending: PendingTable,
no_id_slot: NoIdSlot,
no_id_lock: TokioMutex<()>,
control: TokioMutex<StdioControl>,
reader: JoinHandle<()>,
}
impl StdioProvider {
pub async fn spawn(
id: impl Into<String>,
program: &str,
args: &[String],
) -> Result<Self, HostError> {
let id = id.into();
let mut conn = RawStdioConnection::spawn(program, args)
.await?
.with_label(id.clone());
let (info, capabilities) = conn.handshake().await?;
let (stdin, stdout, control) = conn.into_parts();
let pending: PendingTable = Arc::new(StdMutex::new(HashMap::new()));
let no_id_slot: NoIdSlot = Arc::new(StdMutex::new(None));
let reader = tokio::spawn(run_reader(
stdout,
id.clone(),
Arc::clone(&pending),
Arc::clone(&no_id_slot),
));
Ok(Self {
id,
info,
capabilities,
stdin: TokioMutex::new(stdin),
pending,
no_id_slot,
no_id_lock: TokioMutex::new(()),
control: TokioMutex::new(control),
reader,
})
}
async fn exchange_lockstep(&self, request: Envelope) -> Result<Envelope, HostError> {
let _lockstep = self.no_id_lock.lock().await;
let (tx, rx) = oneshot::channel();
*self.no_id_slot.lock().expect("no_id_slot mutex poisoned") = Some(tx);
let sent = {
let mut stdin = self.stdin.lock().await;
write_envelope(&mut stdin, &request, &self.id).await
};
if let Err(e) = sent {
self.no_id_slot
.lock()
.expect("no_id_slot mutex poisoned")
.take();
return Err(e);
}
match rx.await {
Ok(reply) => reply,
Err(_) => Err(HostError::ProviderCrashed {
id: self.id.clone(),
}),
}
}
}
#[async_trait]
impl ContextProvider for StdioProvider {
fn id(&self) -> &str {
&self.id
}
fn info(&self) -> &ProviderInfo {
&self.info
}
fn capabilities(&self) -> &Capabilities {
&self.capabilities
}
async fn query(&self, query: &ContextQuery) -> Result<ContextQueryResult, HostError> {
if !self.capabilities.correlation {
return match self
.exchange_lockstep(Envelope::Query {
id: None,
query: query.clone(),
})
.await?
{
Envelope::Frames { result, .. } => Ok(result),
Envelope::Error { message, code, .. } => Err(HostError::Provider {
id: self.id.clone(),
code,
message,
}),
other => Err(HostError::UnexpectedEnvelope {
id: self.id.clone(),
expected: "frames".into(),
got: envelope_kind(&other).into(),
}),
};
}
let sent_id = next_correlation_id();
let (tx, rx) = oneshot::channel();
self.pending
.lock()
.expect("pending mutex poisoned")
.insert(sent_id.clone(), tx);
let sent = {
let mut stdin = self.stdin.lock().await;
write_envelope(
&mut stdin,
&Envelope::Query {
id: Some(sent_id.clone()),
query: query.clone(),
},
&self.id,
)
.await
};
if let Err(e) = sent {
self.pending
.lock()
.expect("pending mutex poisoned")
.remove(&sent_id);
return Err(e);
}
let reply = match rx.await {
Ok(reply) => reply?,
Err(_) => {
return Err(HostError::ProviderCrashed {
id: self.id.clone(),
});
}
};
match reply {
Envelope::Frames { id: echoed, result } => {
verify_correlation(&self.id, Some(sent_id.as_str()), echoed.as_deref())?;
Ok(result)
}
Envelope::Error { message, code, .. } => Err(HostError::Provider {
id: self.id.clone(),
code,
message,
}),
other => Err(HostError::UnexpectedEnvelope {
id: self.id.clone(),
expected: "frames".into(),
got: envelope_kind(&other).into(),
}),
}
}
async fn verify(&self, request: &VerifyRequest) -> Result<VerifyResponse, HostError> {
match self
.exchange_lockstep(Envelope::Verify {
request: request.clone(),
})
.await?
{
Envelope::Verified { response } => Ok(response),
Envelope::Error { message, code, .. } => Err(HostError::Provider {
id: self.id.clone(),
code,
message,
}),
other => Err(HostError::UnexpectedEnvelope {
id: self.id.clone(),
expected: "verified".into(),
got: envelope_kind(&other).into(),
}),
}
}
async fn shutdown(&self) -> Result<(), HostError> {
{
let mut stdin = self.stdin.lock().await;
let _ = write_envelope(&mut stdin, &Envelope::Shutdown, &self.id).await;
}
let mut control = self.control.lock().await;
control.wait_or_kill().await
}
}
impl Drop for StdioProvider {
fn drop(&mut self) {
self.reader.abort();
}
}
#[cfg(all(test, unix))]
mod tests {
use super::*;
use contextgraph_types::{ContextFrame, FrameKind};
fn bash_provider(script: &str) -> (String, Vec<String>) {
(
"bash".to_string(),
vec!["-c".to_string(), script.to_string()],
)
}
fn ack_line(version: &str) -> String {
let ack = Envelope::HandshakeAck {
protocol_version: version.to_string(),
provider: ProviderInfo {
name: "bash-fixture".into(),
version: "0.0.1".into(),
data_flow: contextgraph_types::DataFlow {
reads: true,
writes: false,
egress: false,
egress_scopes: vec![],
},
},
capabilities: Capabilities {
query: contextgraph_types::capability::QueryCapability {
kinds: vec!["doc".into()],
},
..Capabilities::default()
},
};
serde_json::to_string(&ack).unwrap()
}
fn frames_line() -> String {
let frame = ContextFrame {
id: "frm_1".into(),
kind: FrameKind::Doc,
title: "README".into(),
content: Some("hello from a stdio provider".into()),
content_digest: None,
uri: Some("file:///README.md".into()),
representation: Default::default(),
content_fidelity: None,
canonical_content_hash: None,
content_ref: None,
transform: None,
minimum_content_fidelity: None,
inline_content_requirement: None,
score: 0.7,
token_cost: 12,
canonical_token_cost: None,
tokenizer_ref: None,
valid_from: None,
valid_to: None,
recorded_at: None,
provenance: vec![],
citation_label: Some("README.md".into()),
embedding: None,
relations: vec![],
};
let env = Envelope::Frames {
id: None,
result: ContextQueryResult {
frames: vec![frame],
truncated: false,
dropped_estimate: None,
},
};
serde_json::to_string(&env).unwrap()
}
fn ack_line_correlating(version: &str) -> String {
let ack = Envelope::HandshakeAck {
protocol_version: version.to_string(),
provider: ProviderInfo {
name: "bash-fixture".into(),
version: "0.0.1".into(),
data_flow: contextgraph_types::DataFlow {
reads: true,
writes: false,
egress: false,
egress_scopes: vec![],
},
},
capabilities: Capabilities {
query: contextgraph_types::capability::QueryCapability {
kinds: vec!["doc".into()],
},
correlation: true,
..Capabilities::default()
},
};
serde_json::to_string(&ack).unwrap()
}
fn frames_template() -> String {
let frame = ContextFrame {
id: "frm_1".into(),
kind: FrameKind::Doc,
title: "__CONTENT__".into(),
content: Some("__CONTENT__".into()),
content_digest: None,
uri: Some("file:///README.md".into()),
representation: Default::default(),
content_fidelity: None,
canonical_content_hash: None,
content_ref: None,
transform: None,
minimum_content_fidelity: None,
inline_content_requirement: None,
score: 0.7,
token_cost: 12,
canonical_token_cost: None,
tokenizer_ref: None,
valid_from: None,
valid_to: None,
recorded_at: None,
provenance: vec![],
citation_label: Some("README.md".into()),
embedding: None,
relations: vec![],
};
let env = Envelope::Frames {
id: Some("__ID__".into()),
result: ContextQueryResult {
frames: vec![frame],
truncated: false,
dropped_estimate: None,
},
};
serde_json::to_string(&env).unwrap()
}
const OUT_OF_ORDER_WITNESS_SCRIPT: &str = r#"
read -r handshake
printf '%s\n' '@ACK@'
read -r q1
read -r q2
tmpl='@FRAMES_TMPL@'
idre='"id":"([^"]+)"'
goalre='"goal":"([^"]+)"'
[[ $q1 =~ $idre ]]; id1=${BASH_REMATCH[1]}
[[ $q1 =~ $goalre ]]; g1=${BASH_REMATCH[1]}
[[ $q2 =~ $idre ]]; id2=${BASH_REMATCH[1]}
[[ $q2 =~ $goalre ]]; g2=${BASH_REMATCH[1]}
r1=${tmpl//__ID__/$id1}; r1=${r1//__CONTENT__/$g1}
r2=${tmpl//__ID__/$id2}; r2=${r2//__CONTENT__/$g2}
printf '%s\n' "$r2"
printf '%s\n' "$r1"
"#;
fn sample_query() -> ContextQuery {
ContextQuery {
goal: "g".into(),
query_text: None,
embedding: None,
kinds: vec![],
anchors: vec![],
max_frames: 5,
max_tokens: 4000,
as_of: None,
representation_preferences: vec![],
}
}
#[tokio::test]
async fn full_handshake_and_query_round_trip_over_stdio() {
let script = format!(
"read h; printf '%s\\n' '{}'; read q; printf '%s\\n' '{}'",
ack_line(PROTOCOL_VERSION),
frames_line()
);
let (program, args) = bash_provider(&script);
let provider = StdioProvider::spawn("docs", &program, &args)
.await
.expect("handshake should succeed");
assert_eq!(provider.id(), "docs");
assert_eq!(provider.info().name, "bash-fixture");
assert!(provider.capabilities().query.kinds.contains(&"doc".into()));
let result = provider.query(&sample_query()).await.expect("query ok");
assert_eq!(result.frames.len(), 1);
assert_eq!(result.frames[0].title, "README");
}
#[tokio::test]
async fn two_correlated_queries_answered_out_of_order_demux_to_their_own_callers() {
let script = OUT_OF_ORDER_WITNESS_SCRIPT
.replace("@ACK@", &ack_line_correlating(PROTOCOL_VERSION))
.replace("@FRAMES_TMPL@", &frames_template());
let (program, args) = bash_provider(&script);
let provider = StdioProvider::spawn("docs", &program, &args)
.await
.expect("handshake should succeed");
assert!(
provider.capabilities().correlation,
"fixture must negotiate correlation for the pipelined path"
);
let mut query_alpha = sample_query();
query_alpha.goal = "alpha".into();
let mut query_bravo = sample_query();
query_bravo.goal = "bravo".into();
let (result_alpha, result_bravo) = tokio::time::timeout(Duration::from_secs(10), async {
tokio::join!(provider.query(&query_alpha), provider.query(&query_bravo))
})
.await
.expect("two concurrent correlated queries must not hang — demux, not lock-step");
let result_alpha = result_alpha.expect("alpha query ok");
let result_bravo = result_bravo.expect("bravo query ok");
assert_eq!(result_alpha.frames.len(), 1);
assert_eq!(result_bravo.frames.len(), 1);
assert_eq!(
result_alpha.frames[0].content.as_deref(),
Some("alpha"),
"the alpha caller must receive alpha's frames, never bravo's"
);
assert_eq!(
result_bravo.frames[0].content.as_deref(),
Some("bravo"),
"the bravo caller must receive bravo's frames, never alpha's"
);
}
#[tokio::test]
async fn an_incompatible_protocol_version_is_a_named_error_not_a_hang() {
let script = format!("read h; printf '%s\\n' '{}'", ack_line("contextgraph/2.0"));
let (program, args) = bash_provider(&script);
let err = match StdioProvider::spawn("docs", &program, &args).await {
Ok(_) => panic!("a version mismatch must reject the provider"),
Err(e) => e,
};
match err {
HostError::VersionMismatch {
provider_version, ..
} => assert_eq!(provider_version, "contextgraph/2.0"),
other => panic!("expected VersionMismatch, got {other}"),
}
}
#[tokio::test]
async fn a_child_dying_after_handshake_surfaces_as_provider_crashed() {
let script = format!(
"read h; printf '%s\\n' '{}'; exit 0",
ack_line(PROTOCOL_VERSION)
);
let (program, args) = bash_provider(&script);
let provider = StdioProvider::spawn("docs", &program, &args)
.await
.expect("handshake ok");
let err = provider
.query(&sample_query())
.await
.expect_err("a dead child must error, not hang");
assert!(
matches!(err, HostError::ProviderCrashed { .. }),
"expected ProviderCrashed, got {err}"
);
}
#[tokio::test]
async fn the_child_is_spawned_with_a_scrubbed_environment() {
let injected = ["PWD", "SHLVL", "_", "HOME", "PATH", "OLDPWD"];
let leaked = std::env::vars()
.map(|(k, _)| k)
.find(|k| !injected.contains(&k.as_str()) && !k.is_empty())
.expect("the test process has at least one non-allowlisted env var");
let mut conn = RawStdioConnection::spawn("bash", &["-c".into(), "env".into()])
.await
.expect("spawn env");
let mut child_keys = Vec::new();
while let Some(line) = conn.read_raw_line().await.expect("read env line") {
if let Some((key, _)) = line.trim_end().split_once('=') {
child_keys.push(key.to_string());
}
}
assert!(
!child_keys.contains(&leaked),
"scrubbed child leaked parent env var `{leaked}` — credentials must not cross"
);
}
}