use std::sync::Arc;
use std::time::Duration;
use crate::builtin;
use arrow::ipc::writer::StreamWriter;
use async_trait::async_trait;
use orion_error::conversion::{SourceErr, SourceRawErr, ToStructError};
use wp_connector_api::SinkReason;
use wp_connector_api::SinkResult;
use wp_connector_api::{
AsyncCtrl, AsyncRawDataSink, AsyncRecordSink, ConnectorDef, SinkBuildCtx, SinkDefProvider,
SinkFactory, SinkHandle, SinkSpec as ResolvedSinkSpec,
};
use wp_data_fmt::RecordFormatter;
use crate::net::transport::{BackoffMode, NetSendPolicy, NetWriter, net_backoff_adaptive};
use super::arrow_conv::{
data_record_to_batch, data_records_to_batch, infer_schema_from_record, sink_err,
};
use wp_model_core::model::DataRecord;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum Framing {
Line,
Len,
}
#[derive(Clone, Debug)]
struct TcpSinkSpec {
addr: String,
port: u16,
framing: Framing,
}
impl TcpSinkSpec {
fn from_resolved(spec: &ResolvedSinkSpec) -> SinkResult<Self> {
let addr = match spec.params.get("addr").and_then(|v| v.as_str()) {
Some(s) => s.to_string(),
None => {
return Err(SinkReason::core_conf()
.to_err()
.with_detail("tcp.addr must be a string"));
}
};
let port = match spec.params.get("port").and_then(|v| v.as_i64()) {
Some(p) if (1..=65535).contains(&p) => p as u16,
Some(_) => {
return Err(SinkReason::core_conf()
.to_err()
.with_detail("tcp.port must be in 1..=65535"));
}
None => 9000,
};
let framing = spec
.params
.get("framing")
.and_then(|v| v.as_str())
.unwrap_or("line");
let framing = match framing.to_ascii_lowercase().as_str() {
"len" | "length" => Framing::Len,
"line" => Framing::Line,
_ => {
return Err(SinkReason::core_conf()
.to_err()
.with_detail("tcp.framing must be 'line' or 'len'"));
}
};
Self::ensure_bool(spec, "max_backoff")?;
Self::ensure_bool(spec, "sendq_backpressure")?;
Ok(Self {
addr,
port,
framing,
})
}
fn ensure_bool(spec: &ResolvedSinkSpec, key: &str) -> SinkResult<()> {
if let Some(v) = spec.params.get(key)
&& v.as_bool().is_none()
{
return Err(SinkReason::core_conf()
.to_err()
.with_detail(format!("tcp.{key} must be a boolean")));
}
Ok(())
}
fn target_addr(&self) -> String {
format!("{}:{}", self.addr, self.port)
}
}
const TCP_DRAIN_MAX_SECS: u64 = 10;
pub struct TcpSink {
writer: NetWriter,
framing: Framing,
sent_cnt: u64,
}
impl TcpSink {
async fn connect(spec: &TcpSinkSpec, rate_limit_rps: usize) -> SinkResult<Self> {
let target = spec.target_addr();
let mode = if rate_limit_rps == 0 {
BackoffMode::ForceOn
} else {
BackoffMode::ForceOff
};
let writer = NetWriter::connect_tcp_with_policy(
&target,
NetSendPolicy {
rate_limit_rps,
backoff_mode: mode,
adaptive: net_backoff_adaptive(),
},
)
.await
.source_err(SinkReason::Sink, "tcp sink connect tcp")?;
log::info!("tcp sink connected: target={}", target);
Ok(Self {
writer,
framing: spec.framing,
sent_cnt: 0,
})
}
}
#[async_trait]
impl AsyncCtrl for TcpSink {
async fn stop(&mut self) -> SinkResult<()> {
self.writer.shutdown().await?;
self.writer
.drain_until_empty(std::time::Duration::from_secs(TCP_DRAIN_MAX_SECS))
.await;
Ok(())
}
async fn reconnect(&mut self) -> SinkResult<()> {
Ok(())
}
}
#[async_trait]
impl AsyncRecordSink for TcpSink {
async fn sink_record(&mut self, data: &wp_model_core::model::DataRecord) -> SinkResult<()> {
let raw = wp_data_fmt::Raw::new().fmt_record(data);
AsyncRawDataSink::sink_str(self, raw.as_str()).await
}
async fn sink_records(
&mut self,
data: Vec<std::sync::Arc<wp_model_core::model::DataRecord>>,
) -> SinkResult<()> {
for record in data {
self.sink_record(&record).await?;
}
Ok(())
}
}
#[async_trait]
impl AsyncRawDataSink for TcpSink {
async fn sink_str(&mut self, data: &str) -> SinkResult<()> {
let payload = build_payload_bytes(data.as_bytes(), self.framing);
if self.sent_cnt == 0 {
log::info!(
"tcp sink first-send: framing={:?} msg_len={} preview='{}'",
self.framing,
payload.len(),
&data.chars().take(64).collect::<String>()
);
}
self.writer.write(&payload).await?;
self.sent_cnt = self.sent_cnt.saturating_add(1);
Ok(())
}
async fn sink_bytes(&mut self, data: &[u8]) -> SinkResult<()> {
let payload = build_payload_bytes(data, self.framing);
if self.sent_cnt == 0 {
log::info!(
"tcp sink first-send(bytes): framing={:?} msg_len={}",
self.framing,
payload.len(),
);
}
self.writer.write(&payload).await?;
self.sent_cnt = self.sent_cnt.saturating_add(1);
Ok(())
}
async fn sink_str_batch(&mut self, data: Vec<&str>) -> SinkResult<()> {
if data.is_empty() {
return Ok(());
}
match self.framing {
Framing::Line => {
let mut total_len = 0;
for str_data in &data {
total_len += str_data.len();
if str_data.as_bytes().last().is_none_or(|&b| b != b'\n') {
total_len += 1;
}
}
let mut buffer = Vec::with_capacity(total_len);
for str_data in &data {
buffer.extend_from_slice(str_data.as_bytes());
if str_data.as_bytes().last().is_none_or(|&b| b != b'\n') {
buffer.push(b'\n');
}
}
self.writer.write(&buffer).await?;
self.sent_cnt = self.sent_cnt.saturating_add(1);
}
Framing::Len => {
let mut buffers = Vec::with_capacity(data.len());
for str_data in &data {
buffers.push(build_payload_bytes(str_data.as_bytes(), self.framing));
}
let total_len: usize = buffers.iter().map(|b| b.len()).sum();
let mut combined = Vec::with_capacity(total_len);
for buffer in buffers {
combined.extend_from_slice(&buffer);
}
self.writer.write(&combined).await?;
self.sent_cnt = self.sent_cnt.saturating_add(data.len() as u64);
}
}
Ok(())
}
async fn sink_bytes_batch(&mut self, data: Vec<&[u8]>) -> SinkResult<()> {
if data.is_empty() {
return Ok(());
}
match self.framing {
Framing::Line => {
let mut total_len = 0;
for bytes_data in &data {
total_len += bytes_data.len();
if bytes_data.last().is_none_or(|&b| b != b'\n') {
total_len += 1;
}
}
let mut buffer = Vec::with_capacity(total_len);
for bytes_data in &data {
buffer.extend_from_slice(bytes_data);
if bytes_data.last().is_none_or(|&b| b != b'\n') {
buffer.push(b'\n');
}
}
self.writer.write(&buffer).await?;
self.sent_cnt = self.sent_cnt.saturating_add(1);
}
Framing::Len => {
let mut combined = Vec::new();
for bytes_data in &data {
combined.extend_from_slice(&build_payload_bytes(bytes_data, self.framing));
}
self.writer.write(&combined).await?;
self.sent_cnt = self.sent_cnt.saturating_add(data.len() as u64);
}
}
Ok(())
}
}
fn buf_writer(buf: &mut Vec<u8>) -> impl std::fmt::Write + '_ {
struct W<'a>(&'a mut Vec<u8>);
impl<'a> std::fmt::Write for W<'a> {
fn write_str(&mut self, s: &str) -> std::fmt::Result {
self.0.extend_from_slice(s.as_bytes());
Ok(())
}
}
W(buf)
}
pub struct TcpFactory;
#[async_trait]
impl SinkFactory for TcpFactory {
fn kind(&self) -> &'static str {
"tcp"
}
fn validate_spec(&self, spec: &ResolvedSinkSpec) -> SinkResult<()> {
let protocol = spec
.params
.get("protocol")
.and_then(|v| v.as_str())
.unwrap_or("txt");
match protocol {
"arrow" => {
let _ = spec
.params
.get("addr")
.and_then(|v| v.as_str())
.ok_or_else(|| {
SinkReason::core_conf()
.to_err()
.with_detail("tcp_arrow: missing required param 'addr'")
})?;
Ok(())
}
"txt" | "" => {
TcpSinkSpec::from_resolved(spec)?;
Ok(())
}
other => Err(SinkReason::core_conf().to_err().with_detail(format!(
"unsupported tcp protocol: '{other}'; expected 'txt' or 'arrow'"
))),
}
}
async fn build(&self, spec: &ResolvedSinkSpec, ctx: &SinkBuildCtx) -> SinkResult<SinkHandle> {
let protocol = spec
.params
.get("protocol")
.and_then(|v| v.as_str())
.unwrap_or("txt");
let sink: Box<dyn wp_connector_api::AsyncSink> = match protocol {
"arrow" => {
let runtime = TcpArrowSink::connect(spec, ctx.rate_limit_rps).await?;
Box::new(runtime)
}
"txt" | "" => {
let resolved = TcpSinkSpec::from_resolved(spec)?;
let runtime = TcpSink::connect(&resolved, ctx.rate_limit_rps).await?;
Box::new(runtime)
}
other => {
return Err(SinkReason::core_conf().to_err().with_detail(format!(
"unsupported tcp protocol: '{other}'; expected 'txt' or 'arrow'"
)));
}
};
Ok(SinkHandle::new(sink))
}
}
impl SinkDefProvider for TcpFactory {
fn sink_def(&self) -> ConnectorDef {
builtin::sink_def("tcp_sink").expect("builtin sink def missing: tcp_sink")
}
}
fn build_payload_bytes(data: &[u8], framing: Framing) -> Vec<u8> {
match framing {
Framing::Line => {
if data.last() == Some(&b'\n') {
data.to_vec()
} else {
let mut buf = Vec::with_capacity(data.len() + 1);
buf.extend_from_slice(data);
buf.push(b'\n');
buf
}
}
Framing::Len => {
let mut buf = Vec::with_capacity(16 + data.len());
let _ = std::fmt::Write::write_fmt(
&mut buf_writer(&mut buf),
format_args!("{} ", data.len()),
);
buf.extend_from_slice(data);
buf
}
}
}
const BACKOFF_INITIAL: Duration = Duration::from_secs(1);
const BACKOFF_MAX: Duration = Duration::from_secs(30);
fn encode_batch_ipc_stream(batch: &arrow::record_batch::RecordBatch) -> SinkResult<Vec<u8>> {
let schema = batch.schema();
let mut buf = Vec::new();
{
let mut writer = StreamWriter::try_new(&mut buf, &schema)
.source_raw_err(SinkReason::Sink, "tcp_arrow create stream writer")?;
writer
.write(batch)
.map_err(|e| sink_err("tcp_arrow encode batch", e))?;
writer
.finish()
.map_err(|e| sink_err("tcp_arrow finish stream", e))?;
}
Ok(buf)
}
enum ConnState {
Connected {
writer: Box<NetWriter>,
},
Disconnected {
next_attempt: tokio::time::Instant,
backoff: Duration,
},
Stopped,
}
pub struct TcpArrowSink {
conn: ConnState,
host: String,
port: u16,
rate_limit_rps: usize,
schema: tokio::sync::Mutex<Option<Arc<arrow::datatypes::Schema>>>,
sent_cnt: u64,
}
impl TcpArrowSink {
pub async fn connect(spec: &ResolvedSinkSpec, rate_limit_rps: usize) -> SinkResult<Self> {
let addr = spec
.params
.get("addr")
.and_then(|v| v.as_str())
.ok_or_else(|| {
SinkReason::core_conf()
.to_err()
.with_detail("tcp_arrow: missing required param 'addr'")
})?;
let port = spec
.params
.get("port")
.and_then(|v| v.as_i64())
.unwrap_or(9000) as u16;
let target = format!("{addr}:{port}");
let mode = if rate_limit_rps == 0 {
BackoffMode::ForceOn
} else {
BackoffMode::ForceOff
};
let writer = NetWriter::connect_tcp_with_policy(
&target,
NetSendPolicy {
rate_limit_rps,
backoff_mode: mode,
adaptive: net_backoff_adaptive(),
},
)
.await
.source_err(SinkReason::Sink, "tcp_arrow connect tcp")?;
log::info!("tcp_arrow sink connected: target={target}");
Ok(Self {
conn: ConnState::Connected {
writer: Box::new(writer),
},
host: addr.to_string(),
port,
rate_limit_rps,
schema: tokio::sync::Mutex::new(None),
sent_cnt: 0,
})
}
async fn get_or_infer_schema(
&self,
record: &DataRecord,
) -> SinkResult<Arc<arrow::datatypes::Schema>> {
let mut guard = self.schema.lock().await;
if guard.is_none() {
*guard = Some(Arc::new(infer_schema_from_record(record)));
}
Ok(Arc::clone(guard.as_ref().unwrap()))
}
async fn connect_writer(&self) -> SinkResult<NetWriter> {
let target = format!("{}:{}", self.host, self.port);
let mode = if self.rate_limit_rps == 0 {
BackoffMode::ForceOn
} else {
BackoffMode::ForceOff
};
NetWriter::connect_tcp_with_policy(
&target,
NetSendPolicy {
rate_limit_rps: self.rate_limit_rps,
backoff_mode: mode,
adaptive: net_backoff_adaptive(),
},
)
.await
.source_err(SinkReason::Sink, "tcp_arrow reconnect tcp")
}
fn enter_disconnected(&mut self) {
self.conn = ConnState::Disconnected {
next_attempt: tokio::time::Instant::now() + BACKOFF_INITIAL,
backoff: BACKOFF_INITIAL,
};
log::warn!("tcp_arrow sink disconnected, will retry");
}
async fn try_reconnect(&mut self) {
match self.connect_writer().await {
Ok(writer) => {
self.conn = ConnState::Connected {
writer: Box::new(writer),
};
log::info!("tcp_arrow sink reconnected: {}:{}", self.host, self.port,);
}
Err(e) => {
if let ConnState::Disconnected {
ref mut next_attempt,
ref mut backoff,
} = self.conn
{
*backoff = (*backoff * 2).min(BACKOFF_MAX);
*next_attempt = tokio::time::Instant::now() + *backoff;
}
log::debug!("tcp_arrow sink reconnect failed: {e}");
}
}
}
async fn send_payload(&mut self, payload: &[u8]) -> SinkResult<()> {
match &mut self.conn {
ConnState::Connected { writer } => match writer.write(payload).await {
Ok(()) => Ok(()),
Err(e) => {
log::warn!("tcp_arrow send error: {e}");
self.enter_disconnected();
Err(e)
}
},
ConnState::Disconnected { next_attempt, .. } => {
if tokio::time::Instant::now() >= *next_attempt {
self.try_reconnect().await;
if let ConnState::Connected { writer } = &mut self.conn {
match writer.write(payload).await {
Ok(()) => Ok(()),
Err(e) => {
log::warn!("tcp_arrow send error after reconnect: {e}");
self.enter_disconnected();
Err(e)
}
}
} else {
Err(SinkReason::Sink
.to_err()
.with_detail("tcp_arrow sink reconnect did not restore connection"))
}
} else {
Err(SinkReason::Sink
.to_err()
.with_detail("tcp_arrow sink waiting for reconnect backoff"))
}
}
ConnState::Stopped => Err(SinkReason::Sink
.to_err()
.with_detail("tcp_arrow sink stopped")),
}
}
}
#[async_trait]
impl AsyncRecordSink for TcpArrowSink {
async fn sink_record(&mut self, data: &wp_model_core::model::DataRecord) -> SinkResult<()> {
let schema = self.get_or_infer_schema(data).await?;
let batch = data_record_to_batch(data, &schema)?;
let payload = encode_batch_ipc_stream(&batch)?;
self.send_payload(&payload).await?;
self.sent_cnt = self.sent_cnt.saturating_add(1);
if self.sent_cnt == 1 {
log::info!(
"tcp_arrow sink first-send: cols={} payload_bytes={}",
schema.fields().len(),
payload.len(),
);
}
Ok(())
}
async fn sink_records(
&mut self,
data: Vec<Arc<wp_model_core::model::DataRecord>>,
) -> SinkResult<()> {
if data.is_empty() {
return Ok(());
}
let row_count = data.len();
let schema = self.get_or_infer_schema(&data[0]).await?;
let batch = data_records_to_batch(&data, &schema)?;
let payload = encode_batch_ipc_stream(&batch)?;
self.send_payload(&payload).await?;
if self.sent_cnt == 0 {
log::info!(
"tcp_arrow sink first-send: rows={} cols={} payload_bytes={}",
row_count,
schema.fields().len(),
payload.len(),
);
}
self.sent_cnt = self.sent_cnt.saturating_add(1);
Ok(())
}
}
#[async_trait]
impl AsyncRawDataSink for TcpArrowSink {
async fn sink_str(&mut self, _data: &str) -> SinkResult<()> {
Err(SinkReason::Sink
.to_err()
.with_detail("tcp_arrow sink only accepts records"))
}
async fn sink_bytes(&mut self, _data: &[u8]) -> SinkResult<()> {
Err(SinkReason::Sink
.to_err()
.with_detail("tcp_arrow sink only accepts records"))
}
async fn sink_str_batch(&mut self, _data: Vec<&str>) -> SinkResult<()> {
Err(SinkReason::Sink
.to_err()
.with_detail("tcp_arrow sink only accepts records"))
}
async fn sink_bytes_batch(&mut self, _data: Vec<&[u8]>) -> SinkResult<()> {
Err(SinkReason::Sink
.to_err()
.with_detail("tcp_arrow sink only accepts records"))
}
}
#[async_trait]
impl AsyncCtrl for TcpArrowSink {
async fn stop(&mut self) -> SinkResult<()> {
let old = std::mem::replace(&mut self.conn, ConnState::Stopped);
if let ConnState::Connected { mut writer } = old {
let _ = writer.shutdown().await;
writer
.drain_until_empty(std::time::Duration::from_secs(10))
.await;
}
Ok(())
}
async fn reconnect(&mut self) -> SinkResult<()> {
self.conn = ConnState::Disconnected {
next_attempt: tokio::time::Instant::now(),
backoff: BACKOFF_INITIAL,
};
self.try_reconnect().await;
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use tokio::io::AsyncReadExt;
use tokio::net::TcpListener;
use wp_connector_api::{AsyncRawDataSink, SinkFactory};
#[tokio::test(flavor = "multi_thread")]
async fn tcp_sink_sends_line() -> anyhow::Result<()> {
if std::env::var("WP_NET_TESTS").unwrap_or_default() != "1" {
return Ok(());
}
let listener = TcpListener::bind("127.0.0.1:0").await?;
let port = listener.local_addr()?.port();
let srv = tokio::spawn(async move {
let (mut s, _) = listener.accept().await.unwrap();
let mut buf = vec![0u8; 16];
let n = s.read(&mut buf).await.unwrap();
String::from_utf8_lossy(&buf[..n]).into_owned()
});
let fac = TcpFactory;
let mut params = toml::map::Map::new();
params.insert("addr".into(), toml::Value::String("127.0.0.1".into()));
params.insert("port".into(), toml::Value::Integer(port as i64));
params.insert("framing".into(), toml::Value::String("line".into()));
let spec = wp_connector_api::SinkSpec {
group: String::new(),
name: "t".into(),
kind: "tcp".into(),
connector_id: String::new(),
params: wp_connector_api::parammap_from_toml_map(params),
filter: None,
};
let ctx = wp_connector_api::SinkBuildCtx::new(std::env::current_dir().unwrap());
let mut h = fac
.build(&spec, &ctx)
.await
.map_err(|e| anyhow::anyhow!("{e}"))?;
AsyncRawDataSink::sink_str(h.sink.as_mut(), "abc")
.await
.map_err(|e| anyhow::anyhow!("{e}"))?;
let body = srv.await.unwrap();
assert_eq!(body, "abc\n");
Ok(())
}
#[test]
fn payload_builder_line_and_len() {
let p1 = build_payload_bytes(b"abc", Framing::Line);
assert_eq!(p1, b"abc\n");
let p2 = build_payload_bytes(b"hello", Framing::Len);
assert_eq!(p2, b"5 hello");
}
#[tokio::test(flavor = "multi_thread")]
async fn tcp_sink_sends_len() -> anyhow::Result<()> {
if std::env::var("WP_NET_TESTS").unwrap_or_default() != "1" {
return Ok(());
}
let listener = TcpListener::bind("127.0.0.1:0").await?;
let port = listener.local_addr()?.port();
let srv = tokio::spawn(async move {
let (mut s, _) = listener.accept().await.unwrap();
let mut buf = vec![0u8; 32];
let n = s.read(&mut buf).await.unwrap();
buf[..n].to_vec()
});
let fac = TcpFactory;
let mut params = toml::map::Map::new();
params.insert("addr".into(), toml::Value::String("127.0.0.1".into()));
params.insert("port".into(), toml::Value::Integer(port as i64));
params.insert("framing".into(), toml::Value::String("len".into()));
let spec = wp_connector_api::SinkSpec {
group: String::new(),
name: "t".into(),
kind: "tcp".into(),
connector_id: String::new(),
params: wp_connector_api::parammap_from_toml_map(params),
filter: None,
};
let ctx = wp_connector_api::SinkBuildCtx::new(std::env::current_dir().unwrap());
let mut h = fac
.build(&spec, &ctx)
.await
.map_err(|e| anyhow::anyhow!("{e}"))?;
AsyncRawDataSink::sink_str(h.sink.as_mut(), "hello")
.await
.map_err(|e| anyhow::anyhow!("{e}"))?;
let body = srv.await.unwrap();
assert_eq!(body, b"5 hello");
Ok(())
}
#[tokio::test(flavor = "multi_thread")]
async fn tcp_arrow_roundtrip_stream() -> anyhow::Result<()> {
use crate::sources::batch::tcp::read_arrow_stream_batches;
use wp_model_core::model::{Field, FieldStorage};
let listener = std::net::TcpListener::bind("127.0.0.1:0")?;
let port = listener.local_addr()?.port();
let srv = std::thread::spawn(move || {
let (stream, _) = listener.accept()?;
stream.set_read_timeout(Some(std::time::Duration::from_secs(5)))?;
let reader = std::io::BufReader::new(&stream);
let batches: Vec<_> = read_arrow_stream_batches(reader)
.map_err(|e| anyhow::anyhow!("{e}"))?
.collect::<Result<Vec<_>, _>>()
.map_err(|e| anyhow::anyhow!("{e}"))?;
Ok::<_, anyhow::Error>(batches)
});
let fac = TcpFactory;
let mut params = toml::map::Map::new();
params.insert("protocol".into(), toml::Value::String("arrow".into()));
params.insert("addr".into(), toml::Value::String("127.0.0.1".into()));
params.insert("port".into(), toml::Value::Integer(port as i64));
let spec = wp_connector_api::SinkSpec {
group: String::new(),
name: "t".into(),
kind: "tcp".into(),
connector_id: String::new(),
params: wp_connector_api::parammap_from_toml_map(params),
filter: None,
};
let ctx = wp_connector_api::SinkBuildCtx::new(std::env::current_dir().unwrap());
let mut h = fac
.build(&spec, &ctx)
.await
.map_err(|e| anyhow::anyhow!("{e}"))?;
let rec1 = Arc::new(DataRecord::from(vec![
FieldStorage::from(Field::from_chars("name", "alice")),
FieldStorage::from(Field::from_digit("count", 42)),
]));
let rec2 = Arc::new(DataRecord::from(vec![
FieldStorage::from(Field::from_chars("name", "bob")),
FieldStorage::from(Field::from_digit("count", 7)),
]));
h.sink
.as_mut()
.sink_records(vec![rec1, rec2])
.await
.map_err(|e| anyhow::anyhow!("{e}"))?;
drop(h);
let batches = srv.join().unwrap()?;
assert_eq!(batches.len(), 1);
assert_eq!(batches[0].num_rows(), 2, "2 records in one batch");
assert_eq!(batches[0].num_columns(), 2);
Ok(())
}
}