diff --git a/rust/lancedb/src/remote/db.rs b/rust/lancedb/src/remote/db.rs index 839cb3797..27f2b3343 100644 --- a/rust/lancedb/src/remote/db.rs +++ b/rust/lancedb/src/remote/db.rs @@ -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 Database for RemoteDatabase { } async fn create_table(&self, mut request: CreateTableRequest) -> Result> { - 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 Database for RemoteDatabase { table_identifier, version, )); + table.seed_schema_ref(schema); + + if stage_initial_data { + let base_table: Arc = 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 = + 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(|_| { diff --git a/rust/lancedb/src/remote/table.rs b/rust/lancedb/src/remote/table.rs index 5fefabeb2..4a2c6bff2 100644 --- a/rust/lancedb/src/remote/table.rs +++ b/rust/lancedb/src/remote/table.rs @@ -441,6 +441,11 @@ impl RemoteTable { } } + /// 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 RemoteTable { 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 BaseTable for RemoteTable { 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 BaseTable for RemoteTable { // 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 BaseTable for RemoteTable { }, 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 BaseTable for RemoteTable { } 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