fix(remote): stream create table readers with multipart writes

This commit is contained in:
Gatefixer
2026-08-05 21:15:47 +00:00
parent c7ea91f3ea
commit 803dd21ecb
2 changed files with 128 additions and 11 deletions
+109 -3
View File
@@ -4,6 +4,7 @@
use std::collections::HashMap;
use std::sync::Arc;
use arrow_array::RecordBatch;
use async_trait::async_trait;
use http::StatusCode;
use lance_io::object_store::StorageOptions;
@@ -18,13 +19,14 @@ use lance_namespace::models::{
};
use crate::Error;
use crate::data::scannable::Scannable;
use crate::database::{
CloneTableRequest, CreateTableMode, CreateTableRequest, Database, DatabaseOptions,
JobDescription, JobInfo, OpenTableRequest, ReadConsistency, TableNamesRequest,
};
use crate::error::Result;
use crate::remote::util::stream_as_body;
use crate::table::BaseTable;
use crate::table::{AddDataBuilder, BaseTable};
use super::ARROW_STREAM_CONTENT_TYPE;
use super::client::{
@@ -693,7 +695,19 @@ impl<S: HttpSend> Database for RemoteDatabase<S> {
}
async fn create_table(&self, mut request: CreateTableRequest) -> Result<Arc<dyn BaseTable>> {
let body = stream_as_body(request.data.scan_as_stream())?;
// The create endpoint limits the size of the complete request even though
// its body is streamed. Sources without a row-count hint (notably Python
// generators / RecordBatchReader) can therefore exceed that limit after
// many individually small batches. Create the schema first and feed the
// unknown-length source through the multipart insert path instead.
let stage_initial_data = request.data.num_rows().is_none();
let schema = request.data.schema();
let body = if stage_initial_data {
let mut empty = RecordBatch::new_empty(schema.clone());
stream_as_body(empty.scan_as_stream())?
} else {
stream_as_body(request.data.scan_as_stream())?
};
let identifier = build_table_identifier(
&request.name,
@@ -764,6 +778,16 @@ impl<S: HttpSend> Database for RemoteDatabase<S> {
table_identifier,
version,
));
table.seed_schema_ref(schema);
if stage_initial_data {
let base_table: Arc<dyn BaseTable> = table.clone();
AddDataBuilder::new(base_table, request.data, None)
.write_options(request.write_options)
.execute()
.await?;
}
self.table_cache.insert(cache_key, table.clone()).await;
Ok(table)
@@ -1106,7 +1130,7 @@ mod tests {
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, OnceLock};
use arrow_array::{Int32Array, RecordBatch};
use arrow_array::{Int32Array, RecordBatch, RecordBatchIterator};
use arrow_schema::{DataType, Field, Schema};
use lance_namespace_impls::{DynamicContextProvider, OperationInfo};
@@ -1371,6 +1395,88 @@ mod tests {
assert_eq!(table.name(), "table1");
}
#[tokio::test]
async fn test_create_table_streaming_reader_uses_multipart_insert() {
let create_count = Arc::new(AtomicUsize::new(0));
let multipart_create_count = Arc::new(AtomicUsize::new(0));
let insert_count = Arc::new(AtomicUsize::new(0));
let complete_count = Arc::new(AtomicUsize::new(0));
let create_count_c = create_count.clone();
let multipart_create_count_c = multipart_create_count.clone();
let insert_count_c = insert_count.clone();
let complete_count_c = complete_count.clone();
let conn = Connection::new_with_handler_and_config(
move |request| {
let path = request.url().path();
let query = request.url().query().unwrap_or("");
match path {
"/v1/table/table1/create/" => {
create_count_c.fetch_add(1, Ordering::SeqCst);
http::Response::builder()
.status(200)
.header("phalanx-version", "0.4.0")
.body(String::new())
.unwrap()
}
"/v1/table/table1/multipart_write/create" => {
multipart_create_count_c.fetch_add(1, Ordering::SeqCst);
http::Response::builder()
.status(200)
.body(r#"{"upload_id":"streaming-create"}"#.to_string())
.unwrap()
}
"/v1/table/table1/insert/" => {
assert!(query.contains("upload_id=streaming-create"));
assert!(query.contains("upload_part_id="));
insert_count_c.fetch_add(1, Ordering::SeqCst);
http::Response::builder()
.status(200)
.body(String::new())
.unwrap()
}
"/v1/table/table1/multipart_write/complete" => {
assert!(query.contains("upload_id=streaming-create"));
complete_count_c.fetch_add(1, Ordering::SeqCst);
http::Response::builder()
.status(200)
.body(r#"{"version":2}"#.to_string())
.unwrap()
}
path => panic!("unexpected path: {path}"),
}
},
ClientConfig {
max_bytes_per_request: Some(1),
..Default::default()
},
);
let schema = Arc::new(Schema::new(vec![Field::new("a", DataType::Int32, false)]));
let batches = vec![
Ok(RecordBatch::try_new(
schema.clone(),
vec![Arc::new(Int32Array::from(vec![1, 2, 3]))],
)
.unwrap()),
Ok(RecordBatch::try_new(
schema.clone(),
vec![Arc::new(Int32Array::from(vec![4, 5, 6]))],
)
.unwrap()),
];
let reader: Box<dyn arrow_array::RecordBatchReader + Send> =
Box::new(RecordBatchIterator::new(batches, schema));
let table = conn.create_table("table1", reader).execute().await.unwrap();
assert_eq!(table.name(), "table1");
assert_eq!(create_count.load(Ordering::SeqCst), 1);
assert_eq!(multipart_create_count.load(Ordering::SeqCst), 1);
assert!(insert_count.load(Ordering::SeqCst) >= 1);
assert_eq!(complete_count.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn test_create_table_already_exists() {
let conn = Connection::new_with_handler(|_| {
+19 -8
View File
@@ -441,6 +441,11 @@ impl<S: HttpSend> RemoteTable<S> {
}
}
/// Seed the schema cache when the caller already has the Arrow schema.
pub(crate) fn seed_schema_ref(&self, schema: SchemaRef) {
self.schema_cache.seed(schema);
}
/// Return a new handle scoped to `branch`, sharing the client but with fresh
/// caches and version/freshness state (the branch tracks its own latest).
/// Mirrors `NativeTable`'s handle-per-branch model.
@@ -1470,8 +1475,8 @@ impl<S: HttpSend + 'static> RemoteTable<S> {
num_partitions: usize,
) -> Result<()> {
debug_assert!(
output.rescannable,
"multipart inserts require rescannable input for retry support"
output.rescannable || num_partitions == 1,
"non-rescannable multipart inserts require a single partition"
);
let plan = Arc::new(
@@ -2105,7 +2110,7 @@ impl<S: HttpSend> BaseTable for RemoteTable<S> {
let table_schema = self.schema().await?;
let table_def = TableDefinition::try_from_rich_schema(table_schema.clone())?;
let num_partitions = if self.server_version.support_multipart_write() {
let (num_partitions, use_multipart) = if self.server_version.support_multipart_write() {
// Peek at the first batch to estimate write partitions (same as
// NativeTable) and, regardless of `write_parallelism`, to detect a
// fully empty input. A multipart write creates its upload session
@@ -2115,10 +2120,12 @@ impl<S: HttpSend> BaseTable for RemoteTable<S> {
// commit and e.g. `mode=overwrite` would be silently dropped. Route
// empty input through the single-request path instead, which always
// sends one schema-only request.
let unknown_size = add.data.num_rows().is_none();
let mut peeked = PeekedScannable::new(add.data);
let n = match peeked.peek().await {
let first_batch = peeked.peek().await;
let n = match first_batch.as_ref() {
Some(first_batch) => match add.write_parallelism {
Some(parallelism) if parallelism > 1 => parallelism,
Some(parallelism) if parallelism > 1 && peeked.rescannable() => parallelism,
Some(_) => 1,
None => {
let max_partitions =
@@ -2133,10 +2140,14 @@ impl<S: HttpSend> BaseTable for RemoteTable<S> {
},
None => 1,
};
// Unknown-length readers cannot be sized up-front, so use a
// single-partition multipart upload. It remains streaming while
// allowing the request body to be split into bounded parts.
let use_multipart = first_batch.is_some() && (n > 1 || unknown_size);
add.data = Box::new(peeked);
n
(n, use_multipart)
} else {
1
(1, false)
};
let output = add.into_plan(&table_schema, &table_def)?;
@@ -2146,7 +2157,7 @@ impl<S: HttpSend> BaseTable for RemoteTable<S> {
}
let _finish = FinishOnDrop(output.tracker.clone());
if num_partitions > 1 {
if use_multipart {
self.add_multipart(output, num_partitions).await
} else {
self.add_single_partition(output).await