use crate::pipeline::external::wire::{RequestHeader, ResponseHeader, WireItemKind};
use crate::pipeline::framer::{Framer, FramerKind, NdjsonFramer};
use anyhow::{anyhow, bail, Context, Result};
use async_trait::async_trait;
use bytes::BytesMut;
use std::collections::{HashMap, VecDeque};
use std::process::Stdio;
use std::sync::Arc;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::process::{Child, ChildStdin, ChildStdout, Command};
use tokio::sync::{Mutex, Notify};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum ChildStdioMode {
#[default]
Persistent,
Transient,
}
#[derive(Debug, Clone)]
pub struct WireResponse {
pub batch_id: u64,
pub error: Option<String>,
pub items: Vec<Vec<u8>>,
}
#[async_trait]
pub trait ExternalTransport: Send + Sync {
async fn write_request(
&self,
batch_id: u64,
items: &[Vec<u8>],
kind: WireItemKind,
) -> Result<()>;
async fn try_read_response(&self, batch_id: u64) -> Result<Option<WireResponse>>;
}
pub(crate) fn encode_request(
framer: &dyn Framer,
batch_id: u64,
items: &[Vec<u8>],
kind: WireItemKind,
) -> Vec<u8> {
let header = RequestHeader {
batch_id,
count: items.len(),
kind,
};
let header_bytes = serde_json::to_vec(&header).expect("RequestHeader serializes");
let mut out = Vec::with_capacity(
header_bytes.len() + 1 + items.iter().map(|i| i.len() + 1).sum::<usize>(),
);
framer.write_message(&header_bytes, &mut out);
for item in items {
framer.write_message(item, &mut out);
}
out
}
async fn write_all_locked(stdin: &Mutex<ChildStdin>, bytes: &[u8]) -> Result<()> {
let mut guard = stdin.lock().await;
guard.write_all(bytes).await.context("write worker stdin")?;
guard.flush().await.context("flush worker stdin")?;
Ok(())
}
struct PendingHeader {
batch_id: u64,
count: usize,
items: Vec<Vec<u8>>,
error: Option<String>,
}
struct ReadState {
stdout: ChildStdout,
buf: BytesMut,
pending: Option<PendingHeader>,
}
impl ReadState {
fn new(stdout: ChildStdout) -> Self {
Self {
stdout,
buf: BytesMut::with_capacity(8 * 1024),
pending: None,
}
}
async fn read_response(&mut self, framer: &dyn Framer) -> Result<WireResponse> {
loop {
if let Some(resp) = self.try_parse(framer)? {
return Ok(resp);
}
let n = self
.stdout
.read_buf(&mut self.buf)
.await
.context("read worker stdout")?;
if n == 0 {
if self.buf.is_empty() && self.pending.is_none() {
bail!("external worker stdout EOF while waiting for response");
}
bail!(
"external worker stdout EOF with incomplete NDJSON ({} bytes buffered)",
self.buf.len()
);
}
}
}
fn try_parse(&mut self, framer: &dyn Framer) -> Result<Option<WireResponse>> {
try_parse_buf(&mut self.buf, &mut self.pending, framer)
}
}
fn try_parse_buf(
buf: &mut BytesMut,
pending: &mut Option<PendingHeader>,
framer: &dyn Framer,
) -> Result<Option<WireResponse>> {
loop {
if let Some(mut p) = pending.take() {
while p.items.len() < p.count {
match framer.next_message(buf)? {
Some(item) => p.items.push(item.to_vec()),
None => {
*pending = Some(p);
return Ok(None);
}
}
}
if let Some(error) = p.error {
return Ok(Some(WireResponse {
batch_id: p.batch_id,
error: Some(error),
items: Vec::new(),
}));
}
return Ok(Some(WireResponse {
batch_id: p.batch_id,
error: None,
items: p.items,
}));
}
let Some(header_bytes) = framer.next_message(buf)? else {
return Ok(None);
};
let header: ResponseHeader = serde_json::from_slice(&header_bytes).with_context(|| {
format!(
"invalid external response header JSON: {}",
String::from_utf8_lossy(&header_bytes)
)
})?;
if header.error.is_some() {
let drain = header.count.unwrap_or(0);
if drain > 0 {
*pending = Some(PendingHeader {
batch_id: header.batch_id,
count: drain,
items: Vec::with_capacity(drain),
error: header.error,
});
continue;
}
return Ok(Some(WireResponse {
batch_id: header.batch_id,
error: header.error,
items: Vec::new(),
}));
}
let count = header.count.ok_or_else(|| {
anyhow!(
"external response for batch_id={} missing both count and error",
header.batch_id
)
})?;
*pending = Some(PendingHeader {
batch_id: header.batch_id,
count,
items: Vec::with_capacity(count),
error: None,
});
}
}
pub struct PersistentChildStdio {
child: Mutex<Child>,
stdin: Mutex<ChildStdin>,
write_order: Mutex<VecDeque<u64>>,
completed: Mutex<HashMap<u64, Result<WireResponse>>>,
read: Mutex<ReadState>,
notify: Notify,
read_gate: Mutex<()>,
framer: NdjsonFramer,
}
impl PersistentChildStdio {
pub fn spawn(command: Vec<String>, framer: FramerKind) -> Result<Self> {
if command.is_empty() {
bail!("persistent child command must not be empty");
}
let framer = framer.into_framer();
let program = &command[0];
let mut cmd = Command::new(program);
if command.len() > 1 {
cmd.args(&command[1..]);
}
cmd.stdin(Stdio::piped())
.stdout(Stdio::piped())
.stderr(Stdio::inherit())
.kill_on_drop(true);
let mut child = cmd
.spawn()
.with_context(|| format!("spawn persistent worker: {command:?}"))?;
let stdin = child
.stdin
.take()
.ok_or_else(|| anyhow!("worker stdin pipe missing"))?;
let stdout = child
.stdout
.take()
.ok_or_else(|| anyhow!("worker stdout pipe missing"))?;
Ok(Self {
child: Mutex::new(child),
stdin: Mutex::new(stdin),
write_order: Mutex::new(VecDeque::new()),
completed: Mutex::new(HashMap::new()),
read: Mutex::new(ReadState::new(stdout)),
notify: Notify::new(),
read_gate: Mutex::new(()),
framer,
})
}
}
#[async_trait]
impl ExternalTransport for PersistentChildStdio {
async fn write_request(
&self,
batch_id: u64,
items: &[Vec<u8>],
kind: WireItemKind,
) -> Result<()> {
let bytes = encode_request(&self.framer, batch_id, items, kind);
self.write_order.lock().await.push_back(batch_id);
write_all_locked(&self.stdin, &bytes).await?;
self.notify.notify_waiters();
Ok(())
}
async fn try_read_response(&self, batch_id: u64) -> Result<Option<WireResponse>> {
loop {
let notified = self.notify.notified();
{
let mut completed = self.completed.lock().await;
if let Some(resp) = completed.remove(&batch_id) {
return resp.map(Some);
}
}
let gate = self.read_gate.try_lock();
let Ok(_gate) = gate else {
notified.await;
continue;
};
{
let mut completed = self.completed.lock().await;
if let Some(resp) = completed.remove(&batch_id) {
return resp.map(Some);
}
}
let req_id = {
let mut q = self.write_order.lock().await;
if q.is_empty() {
drop(_gate);
notified.await;
continue;
}
q.pop_front().expect("non-empty")
};
let framed = {
let mut read = self.read.lock().await;
read.read_response(&self.framer).await
};
let result = match framed {
Ok(resp) if resp.batch_id == req_id => Ok(resp),
Ok(resp) => Err(anyhow!(
"external transform batch_id mismatch: expected request {req_id}, \
got response for {} (fail closed; no cross-batch rebind)",
resp.batch_id
)),
Err(e) => Err(e),
};
if req_id == batch_id {
self.notify.notify_waiters();
return result.map(Some);
}
self.completed.lock().await.insert(req_id, result);
self.notify.notify_waiters();
{
let mut completed = self.completed.lock().await;
if let Some(resp) = completed.remove(&batch_id) {
return resp.map(Some);
}
}
drop(_gate);
}
}
}
impl Drop for PersistentChildStdio {
fn drop(&mut self) {
if let Ok(mut child) = self.child.try_lock() {
let _ = child.start_kill();
}
}
}
pub struct TransientChildStdio {
command: Vec<String>,
framer: NdjsonFramer,
inflight: Mutex<HashMap<u64, TransientChild>>,
notify: Arc<Notify>,
}
struct TransientChild {
child: Child,
read: ReadState,
}
impl TransientChildStdio {
pub fn new(command: Vec<String>, framer: FramerKind) -> Self {
Self {
command,
framer: framer.into_framer(),
inflight: Mutex::new(HashMap::new()),
notify: Arc::new(Notify::new()),
}
}
fn spawn_one(&self) -> Result<(ChildStdin, TransientChild)> {
if self.command.is_empty() {
bail!("transient child command must not be empty");
}
let program = &self.command[0];
let mut cmd = Command::new(program);
if self.command.len() > 1 {
cmd.args(&self.command[1..]);
}
cmd.stdin(Stdio::piped())
.stdout(Stdio::piped())
.stderr(Stdio::inherit())
.kill_on_drop(true);
let mut child = cmd
.spawn()
.with_context(|| format!("spawn transient worker: {:?}", self.command))?;
let stdin = child
.stdin
.take()
.ok_or_else(|| anyhow!("worker stdin pipe missing"))?;
let stdout = child
.stdout
.take()
.ok_or_else(|| anyhow!("worker stdout pipe missing"))?;
Ok((
stdin,
TransientChild {
child,
read: ReadState::new(stdout),
},
))
}
}
#[async_trait]
impl ExternalTransport for TransientChildStdio {
async fn write_request(
&self,
batch_id: u64,
items: &[Vec<u8>],
kind: WireItemKind,
) -> Result<()> {
let (mut stdin, tc) = self.spawn_one()?;
let bytes = encode_request(&self.framer, batch_id, items, kind);
stdin
.write_all(&bytes)
.await
.context("write transient worker stdin")?;
stdin
.flush()
.await
.context("flush transient worker stdin")?;
drop(stdin);
self.inflight.lock().await.insert(batch_id, tc);
self.notify.notify_waiters();
Ok(())
}
async fn try_read_response(&self, batch_id: u64) -> Result<Option<WireResponse>> {
loop {
let notified = self.notify.notified();
let tc = {
let mut inflight = self.inflight.lock().await;
inflight.remove(&batch_id)
};
let Some(mut tc) = tc else {
notified.await;
continue;
};
let resp = tc.read.read_response(&self.framer).await?;
let status = tc
.child
.wait()
.await
.context("wait transient worker exit")?;
if !status.success() {
bail!("transient worker exited with {status} for batch_id={batch_id}");
}
return Ok(Some(resp));
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::pipeline::framer::NdjsonFramer;
#[test]
fn error_header_drains_trailing_item_lines() {
let framer = NdjsonFramer;
let mut out = Vec::new();
framer.write_message(br#"{"batch_id":1,"error":"boom","count":2}"#, &mut out);
framer.write_message(br#"{"junk":1}"#, &mut out);
framer.write_message(br#"{"junk":2}"#, &mut out);
framer.write_message(br#"{"batch_id":2,"count":0}"#, &mut out);
let mut buf = BytesMut::from(out.as_slice());
let mut pending = None;
let err_resp = try_parse_buf(&mut buf, &mut pending, &framer)
.unwrap()
.expect("error response");
assert_eq!(err_resp.batch_id, 1);
assert_eq!(err_resp.error.as_deref(), Some("boom"));
assert!(err_resp.items.is_empty());
assert!(pending.is_none());
let ok_resp = try_parse_buf(&mut buf, &mut pending, &framer)
.unwrap()
.expect("follow-up response");
assert_eq!(ok_resp.batch_id, 2);
assert!(ok_resp.error.is_none());
assert!(ok_resp.items.is_empty());
}
#[test]
fn encode_request_omits_default_change_kind() {
let framer = NdjsonFramer;
let bytes = encode_request(&framer, 7, &[b"{}".to_vec()], WireItemKind::Change);
let header_line = bytes.split(|&b| b == b'\n').next().unwrap();
let header: RequestHeader = serde_json::from_slice(header_line).unwrap();
assert_eq!(header.batch_id, 7);
assert_eq!(header.count, 1);
assert_eq!(header.kind, WireItemKind::Change);
let raw: serde_json::Value = serde_json::from_slice(header_line).unwrap();
assert!(
raw.get("kind").is_none(),
"default Change kind must be omitted for back-compat: {raw}"
);
}
#[test]
fn encode_request_includes_relation_change_kind() {
let framer = NdjsonFramer;
let bytes = encode_request(&framer, 3, &[b"{}".to_vec()], WireItemKind::RelationChange);
let header_line = bytes.split(|&b| b == b'\n').next().unwrap();
let header: RequestHeader = serde_json::from_slice(header_line).unwrap();
assert_eq!(header.kind, WireItemKind::RelationChange);
let raw: serde_json::Value = serde_json::from_slice(header_line).unwrap();
assert_eq!(
raw.get("kind").and_then(|v| v.as_str()),
Some("relation_change"),
"RelationChange must be serialized on the wire: {raw}"
);
}
}