mirror of
https://github.com/lancedb/lancedb.git
synced 2026-08-18 12:08:35 +00:00
fix(remote): stream create table readers with multipart writes
This commit is contained in:
@@ -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(|_| {
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user