use std::io::{IsTerminal, Read};
use std::sync::Arc;
use datafusion::common::exec_datafusion_err;
use datafusion::config::ConfigFileType;
use datafusion::error::Result;
use datafusion::execution::context::SessionState;
use futures::TryStreamExt;
use object_store::memory::InMemory;
use object_store::path::Path as ObjectStorePath;
use object_store::{ObjectStore, ObjectStoreExt};
use url::Url;
#[derive(Debug)]
pub struct StdinCarriesCommands;
const STDIN_LOCATIONS: [&str; 3] = ["/dev/stdin", "/dev/fd/0", "/proc/self/fd/0"];
pub fn is_stdin_location(path: &str) -> bool {
STDIN_LOCATIONS.contains(&path)
}
pub(crate) struct StdinUtils;
impl StdinUtils {
pub(crate) const SCHEME: &'static str = "stdin";
pub(crate) fn rewrite_location(
location: &str,
format: Option<&ConfigFileType>,
) -> String {
if !is_stdin_location(location) {
return location.to_string();
}
let object_name = match format {
Some(ConfigFileType::CSV) => "stdin.csv",
Some(ConfigFileType::JSON) => "stdin.json",
Some(ConfigFileType::PARQUET) => "stdin.parquet",
_ => "stdin",
};
format!("{}:///{object_name}", Self::SCHEME)
}
pub(crate) async fn get_or_create(
state: &SessionState,
url: &Url,
) -> Result<Arc<dyn ObjectStore>> {
let Ok(existing) = state.runtime_env().object_store_registry.get_store(url)
else {
return Self::object_store(state, url).await;
};
let path = ObjectStorePath::from_url_path(url.path())?;
if existing.head(&path).await.is_err() {
let buffered = existing
.list(None)
.try_next()
.await
.ok()
.flatten()
.map(|meta| format!(" as '{}'", meta.location))
.unwrap_or_default();
return Err(exec_datafusion_err!(
"stdin was already read{buffered} by an earlier statement; all \
tables backed by stdin in a session must declare the same \
STORED AS format"
));
}
Ok(existing)
}
async fn object_store(
state: &SessionState,
url: &Url,
) -> Result<Arc<dyn ObjectStore>> {
if state
.config()
.get_extension::<StdinCarriesCommands>()
.is_some()
{
return Err(exec_datafusion_err!(
"stdin is already being read for SQL commands, so it cannot \
also supply table data; pass the query with -c/--command or \
-f/--file so that stdin carries the data, e.g. \
`cat data.csv | datafusion-cli -f query.sql`"
));
}
if std::io::stdin().is_terminal() {
return Err(exec_datafusion_err!(
"stdin is connected to a terminal, not piped data; pipe the \
input in, e.g. `cat data.csv | datafusion-cli -f query.sql`"
));
}
let mut buffer = Vec::new();
std::io::stdin()
.lock()
.read_to_end(&mut buffer)
.map_err(|e| exec_datafusion_err!("Failed to read from stdin: {e}"))?;
Self::in_memory_object_store(url, buffer).await
}
async fn in_memory_object_store(
url: &Url,
data: Vec<u8>,
) -> Result<Arc<dyn ObjectStore>> {
let store = InMemory::new();
store
.put(&ObjectStorePath::from_url_path(url.path())?, data.into())
.await?;
Ok(Arc::new(store))
}
}
#[cfg(test)]
mod tests {
use super::*;
use datafusion::prelude::{SessionConfig, SessionContext};
#[test]
fn rewrites_stdin_locations() {
assert_eq!(
StdinUtils::rewrite_location("/dev/stdin", Some(&ConfigFileType::CSV)),
"stdin:///stdin.csv"
);
assert_eq!(
StdinUtils::rewrite_location("/dev/fd/0", Some(&ConfigFileType::JSON)),
"stdin:///stdin.json"
);
assert_eq!(
StdinUtils::rewrite_location(
"/proc/self/fd/0",
Some(&ConfigFileType::PARQUET)
),
"stdin:///stdin.parquet"
);
assert_eq!(
StdinUtils::rewrite_location("/dev/stdin", None),
"stdin:///stdin"
);
for location in ["/dev/stdout", "data/stdin.csv", "stdin", "s3://b/f.csv"] {
assert_eq!(
StdinUtils::rewrite_location(location, Some(&ConfigFileType::CSV)),
location
);
}
}
async fn count_stdin_rows(
data: Vec<u8>,
stored_as: &str,
format: Option<ConfigFileType>,
options: &str,
) -> Result<usize> {
let location = StdinUtils::rewrite_location("/dev/stdin", format.as_ref());
let url = Url::parse(&location).unwrap();
let store = StdinUtils::in_memory_object_store(&url, data).await?;
let ctx = SessionContext::new();
ctx.register_object_store(&url, store);
ctx.sql(&format!(
"CREATE EXTERNAL TABLE t STORED AS {stored_as} LOCATION '{location}' {options}"
))
.await?
.collect()
.await?;
ctx.sql("SELECT * FROM t").await?.count().await
}
#[tokio::test]
async fn reuses_buffered_stdin_store() -> Result<()> {
let url = Url::parse("stdin:///stdin.csv").unwrap();
let path = ObjectStorePath::from_url_path(url.path())?;
let buffered: Arc<dyn ObjectStore> = Arc::new(InMemory::new());
buffered.put(&path, b"a\n1\n2\n".to_vec().into()).await?;
let ctx = SessionContext::new();
ctx.register_object_store(&url, Arc::clone(&buffered));
let reused = StdinUtils::get_or_create(&ctx.state(), &url).await?;
assert!(
Arc::ptr_eq(&buffered, &reused),
"get_or_create must reuse the registered stdin store, not rebuild it"
);
let bytes = reused.get(&path).await?.bytes().await?;
assert_eq!(bytes.as_ref(), b"a\n1\n2\n");
Ok(())
}
#[tokio::test]
async fn rejects_second_stdin_table_with_different_format() -> Result<()> {
let csv_url = Url::parse("stdin:///stdin.csv").unwrap();
let store =
StdinUtils::in_memory_object_store(&csv_url, b"a\n1\n".to_vec()).await?;
let ctx = SessionContext::new();
ctx.register_object_store(&csv_url, store);
let json_url = Url::parse("stdin:///stdin.json").unwrap();
let err = StdinUtils::get_or_create(&ctx.state(), &json_url)
.await
.unwrap_err()
.to_string();
assert!(
err.contains("must declare the same STORED AS format")
&& err.contains("stdin.csv"),
"unexpected error: {err}"
);
Ok(())
}
#[tokio::test]
async fn errors_when_stdin_carries_commands() {
let config = SessionConfig::new().with_extension(Arc::new(StdinCarriesCommands));
let ctx = SessionContext::new_with_config(config);
let url = Url::parse("stdin:///stdin.csv").unwrap();
let err = StdinUtils::get_or_create(&ctx.state(), &url)
.await
.unwrap_err();
assert!(
err.to_string().contains("SQL commands"),
"unexpected error: {err}"
);
}
#[tokio::test]
async fn stdin_object_store_reads_csv() -> Result<()> {
let data = b"a,b\n1,foo\n2,bar\n".to_vec();
let rows = count_stdin_rows(
data,
"CSV",
Some(ConfigFileType::CSV),
"OPTIONS ('format.has_header' 'true')",
)
.await?;
assert_eq!(rows, 2);
Ok(())
}
#[tokio::test]
async fn stdin_object_store_reads_json() -> Result<()> {
let data = b"{\"a\": 1, \"b\": \"foo\"}\n{\"a\": 2, \"b\": \"bar\"}\n".to_vec();
let rows = count_stdin_rows(data, "JSON", Some(ConfigFileType::JSON), "").await?;
assert_eq!(rows, 2);
Ok(())
}
#[tokio::test]
async fn stdin_object_store_reads_parquet() -> Result<()> {
use datafusion::arrow::array::Int32Array;
use datafusion::arrow::datatypes::{DataType, Field, Schema};
use datafusion::arrow::record_batch::RecordBatch;
use parquet::arrow::ArrowWriter;
let schema = Arc::new(Schema::new(vec![Field::new("a", DataType::Int32, false)]));
let batch = RecordBatch::try_new(
Arc::clone(&schema),
vec![Arc::new(Int32Array::from(vec![1, 2, 3]))],
)
.unwrap();
let mut data = Vec::new();
let mut writer = ArrowWriter::try_new(&mut data, schema, None).unwrap();
writer.write(&batch).unwrap();
writer.close().unwrap();
let rows =
count_stdin_rows(data, "PARQUET", Some(ConfigFileType::PARQUET), "").await?;
assert_eq!(rows, 3);
Ok(())
}
}