mirror of
https://github.com/lancedb/lancedb.git
synced 2026-08-31 02:18:27 +00:00
Compare commits
4 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 88d8a69a99 | |||
| bd779bb7d5 | |||
| 7357d63e87 | |||
| 624a75edf7 |
@@ -707,6 +707,9 @@ class LanceDBConnection(DBConnection):
|
||||
self._namespace_client_properties = namespace_client_properties
|
||||
if _inner is not None:
|
||||
self._conn = _inner
|
||||
# Native-derived wrappers resolve this in their async reconstruction
|
||||
# path so construction never synchronously re-enters LOOP.
|
||||
self._read_consistency_interval = read_consistency_interval
|
||||
self._cached_namespace_client = None
|
||||
return
|
||||
|
||||
@@ -756,11 +759,14 @@ class LanceDBConnection(DBConnection):
|
||||
# storage_options. Also, this class really shouldn't be holding any state
|
||||
# beyond _conn.
|
||||
self._conn = AsyncConnection(LOOP.run(do_connect()))
|
||||
# Keep property access synchronous so debugger introspection cannot wait on
|
||||
# the background loop while that thread is suspended at a breakpoint.
|
||||
self._read_consistency_interval = read_consistency_interval
|
||||
self._cached_namespace_client: Optional[LanceNamespace] = None
|
||||
|
||||
@property
|
||||
def read_consistency_interval(self) -> Optional[timedelta]:
|
||||
return LOOP.run(self._conn.get_read_consistency_interval())
|
||||
return self._read_consistency_interval
|
||||
|
||||
@property
|
||||
def session(self) -> Optional[Session]:
|
||||
@@ -771,8 +777,16 @@ class LanceDBConnection(DBConnection):
|
||||
return self._conn.uri
|
||||
|
||||
@classmethod
|
||||
def from_inner(cls, inner: LanceDbConnection):
|
||||
return cls(None, _inner=inner)
|
||||
def from_inner(
|
||||
cls,
|
||||
inner: LanceDbConnection,
|
||||
read_consistency_interval: Optional[timedelta],
|
||||
):
|
||||
return cls(
|
||||
None,
|
||||
read_consistency_interval=read_consistency_interval,
|
||||
_inner=inner,
|
||||
)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.__class__.__name__}(uri={self._conn.uri!r})"
|
||||
|
||||
@@ -3,6 +3,7 @@
|
||||
|
||||
|
||||
from typing import List
|
||||
from urllib.parse import unquote, urlparse
|
||||
|
||||
import numpy as np
|
||||
|
||||
@@ -125,9 +126,20 @@ class InstructorEmbeddingFunction(TextEmbeddingFunction):
|
||||
|
||||
@weak_lru(maxsize=1)
|
||||
def get_model(self):
|
||||
instructor_embedding = attempt_import_or_raise(
|
||||
"InstructorEmbedding", "InstructorEmbedding"
|
||||
)
|
||||
huggingface_hub = attempt_import_or_raise("huggingface_hub", "huggingface-hub")
|
||||
missing = object()
|
||||
original_cached_download = getattr(huggingface_hub, "cached_download", missing)
|
||||
if original_cached_download is missing:
|
||||
huggingface_hub.cached_download = _cached_download(huggingface_hub)
|
||||
|
||||
try:
|
||||
instructor_embedding = attempt_import_or_raise(
|
||||
"InstructorEmbedding", "InstructorEmbedding"
|
||||
)
|
||||
finally:
|
||||
if original_cached_download is missing:
|
||||
del huggingface_hub.cached_download
|
||||
|
||||
torch = attempt_import_or_raise("torch", "torch")
|
||||
|
||||
model = instructor_embedding.INSTRUCTOR(self.name)
|
||||
@@ -140,3 +152,44 @@ class InstructorEmbeddingFunction(TextEmbeddingFunction):
|
||||
model, {torch.nn.Linear}, dtype=torch.qint8
|
||||
)
|
||||
return model
|
||||
|
||||
|
||||
def _cached_download(huggingface_hub):
|
||||
"""Provide the legacy download API used by sentence-transformers 2.2.x."""
|
||||
|
||||
def cached_download(
|
||||
*,
|
||||
url,
|
||||
cache_dir=None,
|
||||
force_filename=None,
|
||||
library_name=None,
|
||||
library_version=None,
|
||||
user_agent=None,
|
||||
use_auth_token=None,
|
||||
**_,
|
||||
):
|
||||
path = urlparse(url).path.lstrip("/")
|
||||
try:
|
||||
repo_id, resolved_path = path.split("/resolve/", maxsplit=1)
|
||||
revision, filename = resolved_path.split("/", maxsplit=1)
|
||||
except ValueError as err:
|
||||
raise ValueError(f"Unsupported Hugging Face Hub URL: {url}") from err
|
||||
|
||||
repo_id = unquote(repo_id)
|
||||
revision = unquote(revision)
|
||||
filename = unquote(filename)
|
||||
# sentence-transformers derives force_filename from this Hub path with
|
||||
# os.path.join. Using the URL path beneath local_dir produces the same
|
||||
# local destination without sending Windows separators to the Hub.
|
||||
return huggingface_hub.hf_hub_download(
|
||||
repo_id=repo_id,
|
||||
filename=filename,
|
||||
revision=revision,
|
||||
local_dir=cache_dir,
|
||||
library_name=library_name,
|
||||
library_version=library_version,
|
||||
user_agent=user_agent,
|
||||
token=use_auth_token,
|
||||
)
|
||||
|
||||
return cached_download
|
||||
|
||||
@@ -226,7 +226,7 @@ class PermutationBuilder:
|
||||
|
||||
async def do_execute():
|
||||
inner_tbl = await self._async.execute()
|
||||
return LanceTable.from_inner(inner_tbl)
|
||||
return await LanceTable.from_inner(inner_tbl)
|
||||
|
||||
return LOOP.run(do_execute())
|
||||
|
||||
|
||||
@@ -2182,11 +2182,15 @@ class LanceTable(Table):
|
||||
return self.name
|
||||
|
||||
@classmethod
|
||||
def from_inner(cls, tbl: LanceDBTable):
|
||||
from .db import LanceDBConnection
|
||||
async def from_inner(cls, tbl: LanceDBTable):
|
||||
from .db import AsyncConnection, LanceDBConnection
|
||||
|
||||
async_tbl = AsyncTable(tbl)
|
||||
conn = LanceDBConnection.from_inner(tbl.database())
|
||||
inner_conn = tbl.database()
|
||||
read_consistency_interval = await AsyncConnection(
|
||||
inner_conn
|
||||
).get_read_consistency_interval()
|
||||
conn = LanceDBConnection.from_inner(inner_conn, read_consistency_interval)
|
||||
return cls(
|
||||
conn,
|
||||
async_tbl.name,
|
||||
|
||||
@@ -77,6 +77,23 @@ def test_sync_repr_does_not_use_background_loop(tmp_path, monkeypatch):
|
||||
assert repr(table) == f"LanceTable(name='test', _conn={db!r})"
|
||||
|
||||
|
||||
def test_read_consistency_interval_does_not_use_background_loop(tmp_path, monkeypatch):
|
||||
from lancedb.background_loop import LOOP
|
||||
from lancedb.db import LanceDBConnection
|
||||
|
||||
consistency_interval = timedelta(seconds=5)
|
||||
db = lancedb.connect(tmp_path, read_consistency_interval=consistency_interval)
|
||||
db_from_inner = LanceDBConnection.from_inner(db._inner, consistency_interval)
|
||||
|
||||
def fail_run(*args, **kwargs):
|
||||
raise AssertionError("properties should not use the Python background loop")
|
||||
|
||||
monkeypatch.setattr(LOOP, "run", fail_run)
|
||||
|
||||
assert db.read_consistency_interval == consistency_interval
|
||||
assert db_from_inner.read_consistency_interval == consistency_interval
|
||||
|
||||
|
||||
def test_ingest_pd(tmp_path):
|
||||
db = lancedb.connect(tmp_path)
|
||||
|
||||
|
||||
@@ -1,8 +1,11 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
import ntpath
|
||||
import os
|
||||
import pickle
|
||||
import sys
|
||||
from types import ModuleType
|
||||
from typing import List, Optional, Union
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
@@ -522,6 +525,59 @@ def test_embedding_function_safe_model_dump(embedding_type):
|
||||
)
|
||||
|
||||
|
||||
def test_instructor_embedding_supports_huggingface_hub_without_cached_download(
|
||||
tmp_path, monkeypatch
|
||||
):
|
||||
from lancedb.embeddings.instructor import InstructorEmbeddingFunction
|
||||
|
||||
hub_download = MagicMock(return_value="/cache/1_Pooling/config.json")
|
||||
huggingface_hub = ModuleType("huggingface_hub")
|
||||
huggingface_hub.hf_hub_download = hub_download
|
||||
torch = ModuleType("torch")
|
||||
monkeypatch.setitem(sys.modules, "huggingface_hub", huggingface_hub)
|
||||
monkeypatch.setitem(sys.modules, "torch", torch)
|
||||
monkeypatch.delitem(sys.modules, "InstructorEmbedding", raising=False)
|
||||
monkeypatch.syspath_prepend(str(tmp_path))
|
||||
|
||||
(tmp_path / "InstructorEmbedding.py").write_text(
|
||||
"from huggingface_hub import cached_download\n\n"
|
||||
"class INSTRUCTOR:\n"
|
||||
" def __init__(self, name):\n"
|
||||
" self.name = name\n"
|
||||
)
|
||||
|
||||
embedding = InstructorEmbeddingFunction.create(show_progress_bar=False)
|
||||
instructor_model = embedding.get_model()
|
||||
|
||||
assert instructor_model.name == "hkunlp/instructor-base"
|
||||
assert not hasattr(huggingface_hub, "cached_download")
|
||||
|
||||
instructor_embedding = sys.modules["InstructorEmbedding"]
|
||||
path = instructor_embedding.cached_download(
|
||||
url=(
|
||||
"https://huggingface.co/hkunlp/instructor-base/resolve/abc123/"
|
||||
"1_Pooling/config.json"
|
||||
),
|
||||
cache_dir="/cache",
|
||||
force_filename=ntpath.join("1_Pooling", "config.json"),
|
||||
library_name="sentence-transformers",
|
||||
library_version="2.2.2",
|
||||
use_auth_token="token",
|
||||
)
|
||||
|
||||
assert path == "/cache/1_Pooling/config.json"
|
||||
hub_download.assert_called_once_with(
|
||||
repo_id="hkunlp/instructor-base",
|
||||
filename="1_Pooling/config.json",
|
||||
revision="abc123",
|
||||
local_dir="/cache",
|
||||
library_name="sentence-transformers",
|
||||
library_version="2.2.2",
|
||||
user_agent=None,
|
||||
token="token",
|
||||
)
|
||||
|
||||
|
||||
@patch("time.sleep")
|
||||
def test_retry(mock_sleep):
|
||||
test_function = MagicMock(side_effect=[Exception] * 9 + ["result"])
|
||||
|
||||
@@ -6,6 +6,7 @@ import math
|
||||
import pytest
|
||||
|
||||
from lancedb import DBConnection, Table, connect
|
||||
from lancedb.background_loop import LOOP
|
||||
from lancedb.permutation import Permutation, Permutations, permutation_builder
|
||||
|
||||
|
||||
@@ -31,6 +32,25 @@ def test_split_random_ratios(mem_db):
|
||||
assert 65 <= split_1_count <= 75 # ~70% ± tolerance
|
||||
|
||||
|
||||
def test_execute_does_not_reenter_background_loop(tmp_path, monkeypatch):
|
||||
import threading
|
||||
|
||||
db = connect(tmp_path)
|
||||
tbl = db.create_table("test_table", pa.table({"x": range(10)}))
|
||||
original_run = LOOP.run
|
||||
|
||||
def fail_on_reentry(future):
|
||||
assert threading.current_thread() is not LOOP.thread
|
||||
return original_run(future)
|
||||
|
||||
monkeypatch.setattr(LOOP, "run", fail_on_reentry)
|
||||
|
||||
permutation_tbl = permutation_builder(tbl).execute()
|
||||
|
||||
assert permutation_tbl.count_rows() == 10
|
||||
assert permutation_tbl._conn.read_consistency_interval is None
|
||||
|
||||
|
||||
def test_split_random_counts(mem_db):
|
||||
"""Test random splitting with absolute counts."""
|
||||
tbl = mem_db.create_table(
|
||||
|
||||
@@ -6,6 +6,7 @@ import os
|
||||
import sys
|
||||
import threading
|
||||
import warnings
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from datetime import date, datetime, timedelta
|
||||
from time import sleep
|
||||
from typing import List
|
||||
@@ -2124,6 +2125,27 @@ def test_delete(mem_db: DBConnection):
|
||||
assert table.to_arrow()["id"].to_pylist() == [1]
|
||||
|
||||
|
||||
def test_concurrent_deletes_are_thread_safe(mem_db: DBConnection):
|
||||
num_workers = 8
|
||||
table = mem_db.create_table(
|
||||
"my_table", data=[{"id": row_id} for row_id in range(num_workers)]
|
||||
)
|
||||
barrier = threading.Barrier(num_workers)
|
||||
|
||||
def delete(row_id: int):
|
||||
barrier.wait()
|
||||
return table.delete(f"id = {row_id}")
|
||||
|
||||
with ThreadPoolExecutor(max_workers=num_workers) as pool:
|
||||
results = list(pool.map(delete, range(num_workers)))
|
||||
|
||||
assert all(result.num_deleted_rows == 1 for result in results)
|
||||
assert sorted(result.version for result in results) == list(
|
||||
range(2, num_workers + 2)
|
||||
)
|
||||
assert table.count_rows() == 0
|
||||
|
||||
|
||||
def test_delete_expr(mem_db: DBConnection):
|
||||
table = mem_db.create_table(
|
||||
"my_table",
|
||||
|
||||
@@ -745,6 +745,9 @@ impl Table {
|
||||
|
||||
#[allow(private_interfaces)]
|
||||
pub fn delete(self_: PyRef<'_, Self>, condition: PredicateArg) -> PyResult<Bound<'_, PyAny>> {
|
||||
// Do not hold the Python borrow across the await. The cloned Rust table
|
||||
// handle is thread-safe and allows deletes on the same Python table to
|
||||
// run concurrently without PyO3 reporting "Already borrowed".
|
||||
let inner = self_.inner_ref()?.clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
let result = match &condition {
|
||||
|
||||
@@ -4,7 +4,6 @@
|
||||
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;
|
||||
@@ -19,14 +18,13 @@ 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::{AddDataBuilder, BaseTable};
|
||||
use crate::table::BaseTable;
|
||||
|
||||
use super::ARROW_STREAM_CONTENT_TYPE;
|
||||
use super::client::{
|
||||
@@ -695,19 +693,7 @@ impl<S: HttpSend> Database for RemoteDatabase<S> {
|
||||
}
|
||||
|
||||
async fn create_table(&self, mut request: CreateTableRequest) -> Result<Arc<dyn BaseTable>> {
|
||||
// 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 body = stream_as_body(request.data.scan_as_stream())?;
|
||||
|
||||
let identifier = build_table_identifier(
|
||||
&request.name,
|
||||
@@ -778,16 +764,6 @@ 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)
|
||||
@@ -1130,7 +1106,7 @@ mod tests {
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
use std::sync::{Arc, OnceLock};
|
||||
|
||||
use arrow_array::{Int32Array, RecordBatch, RecordBatchIterator};
|
||||
use arrow_array::{Int32Array, RecordBatch};
|
||||
use arrow_schema::{DataType, Field, Schema};
|
||||
use lance_namespace_impls::{DynamicContextProvider, OperationInfo};
|
||||
|
||||
@@ -1395,88 +1371,6 @@ 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,11 +441,6 @@ 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.
|
||||
@@ -1475,8 +1470,8 @@ impl<S: HttpSend + 'static> RemoteTable<S> {
|
||||
num_partitions: usize,
|
||||
) -> Result<()> {
|
||||
debug_assert!(
|
||||
output.rescannable || num_partitions == 1,
|
||||
"non-rescannable multipart inserts require a single partition"
|
||||
output.rescannable,
|
||||
"multipart inserts require rescannable input for retry support"
|
||||
);
|
||||
|
||||
let plan = Arc::new(
|
||||
@@ -2110,7 +2105,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, use_multipart) = if self.server_version.support_multipart_write() {
|
||||
let num_partitions = 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
|
||||
@@ -2120,12 +2115,10 @@ 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 first_batch = peeked.peek().await;
|
||||
let n = match first_batch.as_ref() {
|
||||
let n = match peeked.peek().await {
|
||||
Some(first_batch) => match add.write_parallelism {
|
||||
Some(parallelism) if parallelism > 1 && peeked.rescannable() => parallelism,
|
||||
Some(parallelism) if parallelism > 1 => parallelism,
|
||||
Some(_) => 1,
|
||||
None => {
|
||||
let max_partitions =
|
||||
@@ -2140,14 +2133,10 @@ 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, use_multipart)
|
||||
n
|
||||
} else {
|
||||
(1, false)
|
||||
1
|
||||
};
|
||||
|
||||
let output = add.into_plan(&table_schema, &table_def)?;
|
||||
@@ -2157,7 +2146,7 @@ impl<S: HttpSend> BaseTable for RemoteTable<S> {
|
||||
}
|
||||
let _finish = FinishOnDrop(output.tracker.clone());
|
||||
|
||||
if use_multipart {
|
||||
if num_partitions > 1 {
|
||||
self.add_multipart(output, num_partitions).await
|
||||
} else {
|
||||
self.add_single_partition(output).await
|
||||
|
||||
Reference in New Issue
Block a user