Compare commits

..

4 Commits

Author SHA1 Message Date
Gatefixer d28ea445e7 fix: harden legacy update filter barrier 2026-08-08 13:26:46 +00:00
Gatefixer d6809a80f2 fix: scope legacy update materialization fallback 2026-08-08 12:37:38 +00:00
Gatefixer e5a7a092bb fix: avoid offset overflow in filtered updates 2026-08-06 07:56:42 +00:00
Gatefixer 2adbf791fc test: cover wide multi-fragment updates 2026-08-06 07:01:31 +00:00
3 changed files with 362 additions and 113 deletions
+3 -56
View File
@@ -3,7 +3,6 @@
from typing import List
from urllib.parse import unquote, urlparse
import numpy as np
@@ -126,20 +125,9 @@ class InstructorEmbeddingFunction(TextEmbeddingFunction):
@weak_lru(maxsize=1)
def get_model(self):
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
instructor_embedding = attempt_import_or_raise(
"InstructorEmbedding", "InstructorEmbedding"
)
torch = attempt_import_or_raise("torch", "torch")
model = instructor_embedding.INSTRUCTOR(self.name)
@@ -152,44 +140,3 @@ 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
-56
View File
@@ -1,11 +1,8 @@
# 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
@@ -525,59 +522,6 @@ 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"])
+359 -1
View File
@@ -3,6 +3,7 @@
use std::sync::Arc;
use arrow_schema::DataType;
use lance::dataset::UpdateBuilder as LanceUpdateBuilder;
use serde::{Deserialize, Serialize};
@@ -84,10 +85,11 @@ pub(crate) async fn execute_update(
let dataset = table.dataset.get().await?;
// 2. Initialize the Lance Core builder
let mut builder = LanceUpdateBuilder::new(dataset);
let mut builder = LanceUpdateBuilder::new(dataset.clone());
// 3. Apply the filter (WHERE clause)
if let Some(predicate) = update.filter {
let predicate = safe_update_filter(&predicate, dataset.as_ref());
builder = builder.update_where(&predicate)?;
}
@@ -109,9 +111,61 @@ pub(crate) async fn execute_update(
})
}
/// Keep vulnerable legacy updates on the early-materialization scan path.
///
/// Late materialization uses `TakeExec` to concatenate values read from multiple
/// fragments. That can overflow a single 32-bit-offset array. Lance's update
/// builder does not currently expose its scanner's materialization controls, so
/// cast the predicate to an integer before comparing it with `1`. Lance's
/// scalar-index extractor does not unwrap non-literal casts, keeping every
/// supported predicate out of the vulnerable late-materialization plan.
///
/// Keep the original SQL verbatim instead of parsing and serializing it. Newlines
/// isolate the generated syntax from a trailing line comment in the predicate.
///
/// This compatibility fallback is intentionally limited to legacy storage. V2
/// readers do not use the affected materialization path and keep their original
/// filter expression and indexed plan.
fn safe_update_filter(predicate: &str, dataset: &lance::Dataset) -> String {
let has_offset_columns = dataset
.schema()
.fields
.iter()
.any(|field| has_32_bit_offsets(&field.data_type()));
if !dataset.manifest().should_use_legacy_format() || !has_offset_columns {
return predicate.to_owned();
}
format!("CAST((\n{predicate}\n) AS INT) = 1")
}
fn has_32_bit_offsets(data_type: &DataType) -> bool {
match data_type {
DataType::Binary
| DataType::Utf8
| DataType::List(_)
| DataType::ListView(_)
| DataType::Map(_, _)
| DataType::Union(_, _) => true,
DataType::FixedSizeList(field, _)
| DataType::LargeList(field)
| DataType::LargeListView(field) => has_32_bit_offsets(field.data_type()),
DataType::Struct(fields) => fields
.iter()
.any(|field| has_32_bit_offsets(field.data_type())),
DataType::Dictionary(_, values) => has_32_bit_offsets(values),
DataType::RunEndEncoded(_, values) => has_32_bit_offsets(values.data_type()),
_ => false,
}
}
#[cfg(test)]
mod tests {
use crate::connect;
use crate::connection::LanceFileVersion;
use crate::database::listing::{ListingDatabaseOptions, NewTableConfig};
use crate::index::{Index, scalar::BTreeIndexBuilder};
use crate::query::QueryBase;
use crate::query::{ExecutableQuery, Select};
use arrow_array::{
@@ -122,9 +176,18 @@ mod tests {
use arrow_data::ArrayDataBuilder;
use arrow_schema::{ArrowError, DataType, Field, Schema, TimeUnit};
use futures::TryStreamExt;
use lance::io::exec::Planner;
use std::sync::Arc;
use std::time::Duration;
fn contains_take(plan: &dyn datafusion_physical_plan::ExecutionPlan) -> bool {
plan.name() == "TakeExec"
|| plan
.children()
.iter()
.any(|child| contains_take(child.as_ref()))
}
#[tokio::test]
async fn test_update_all_types() {
let conn = connect("memory://")
@@ -409,6 +472,301 @@ mod tests {
}
}
#[tokio::test]
async fn test_update_materializes_offset_columns_before_filter() {
let batch = record_batch!(
("id", Int32, [0, 1, 2, 3]),
(
"split",
Utf8,
[Some("test"), None, Some("test"), Some("train")]
),
("payload", Utf8, ["a", "b", "c", "d"])
)
.unwrap();
let conn = connect("memory://")
.database_options(&ListingDatabaseOptions {
new_table_config: NewTableConfig {
data_storage_version: Some(LanceFileVersion::Legacy),
..Default::default()
},
..Default::default()
})
.execute()
.await
.unwrap();
let table = conn
.create_table("offset_table", batch.clone())
.execute()
.await
.unwrap();
table.add(batch).execute().await.unwrap();
table
.create_index(&["split"], Index::BTree(BTreeIndexBuilder::default()))
.execute()
.await
.unwrap();
let dataset = table.dataset().unwrap().get().await.unwrap();
let planner = Planner::new(Arc::new(dataset.schema().into()));
let filter = planner.parse_filter("split = 'test'").unwrap();
let filter = planner.optimize_expr(filter).unwrap();
let mut scanner = dataset.scan();
scanner.with_row_id().filter_expr(filter);
let explanation = scanner.explain_plan(false).await.unwrap();
let plan = scanner.create_plan().await.unwrap();
assert!(
contains_take(plan.as_ref()),
"test setup must late-materialize payload:\n{explanation}"
);
let guarded_filter = super::safe_update_filter("split = 'test'", dataset.as_ref());
let filter = planner.parse_filter(&guarded_filter).unwrap();
let filter = planner.optimize_expr(filter).unwrap();
let mut scanner = dataset.scan();
scanner.with_row_id().filter_expr(filter);
let explanation = scanner.explain_plan(false).await.unwrap();
let plan = scanner.create_plan().await.unwrap();
// Regression test for #1291: the payload must be read by the scan, not
// concatenated across fragments by a late-materializing TakeExec.
assert!(
!contains_take(plan.as_ref()),
"unexpected late materialization:\n{explanation}"
);
let result = table
.update()
.only_if("split = 'test'")
.column("split", "'TEST'")
.execute()
.await
.unwrap();
assert_eq!(result.rows_updated, 4);
assert_eq!(
table
.count_rows(Some("split = 'TEST'".to_string()))
.await
.unwrap(),
4
);
assert_eq!(
table
.count_rows(Some("payload IN ('a', 'b', 'c', 'd')".to_string()))
.await
.unwrap(),
8
);
}
#[tokio::test]
async fn test_update_v2_keeps_indexed_plan() {
let batch = record_batch!(
("id", Int32, [0, 1, 2, 3]),
("split", Utf8, ["test", "train", "test", "train"]),
("payload", Utf8, ["a", "b", "c", "d"])
)
.unwrap();
let conn = connect("memory://")
.database_options(&ListingDatabaseOptions {
new_table_config: NewTableConfig {
data_storage_version: Some(LanceFileVersion::V2_0),
..Default::default()
},
..Default::default()
})
.execute()
.await
.unwrap();
let table = conn
.create_table("v2_offset_table", batch.clone())
.execute()
.await
.unwrap();
table.add(batch).execute().await.unwrap();
table
.create_index(&["split"], Index::BTree(BTreeIndexBuilder::default()))
.execute()
.await
.unwrap();
let dataset = table.dataset().unwrap().get().await.unwrap();
let predicate = "split = 'test'";
let update_filter = super::safe_update_filter(predicate, dataset.as_ref());
assert_eq!(update_filter, predicate);
let planner = Planner::new(Arc::new(dataset.schema().into()));
let filter = planner.parse_filter(&update_filter).unwrap();
let filter = planner.optimize_expr(filter).unwrap();
let mut scanner = dataset.scan();
scanner.with_row_id().filter_expr(filter);
let explanation = scanner.explain_plan(false).await.unwrap();
let plan = scanner.create_plan().await.unwrap();
assert!(
explanation.contains("ScalarIndexQuery"),
"v2 plan unexpectedly lost its scalar index:\n{explanation}"
);
assert!(
!contains_take(plan.as_ref()),
"v2 plan unexpectedly used the legacy TakeExec path:\n{explanation}"
);
let result = table
.update()
.only_if(predicate)
.column("split", "'TEST'")
.execute()
.await
.unwrap();
assert_eq!(result.rows_updated, 4);
}
#[tokio::test]
async fn test_update_accepts_trailing_comment_filter() {
let conn = connect("memory://")
.database_options(&ListingDatabaseOptions {
new_table_config: NewTableConfig {
data_storage_version: Some(LanceFileVersion::Legacy),
..Default::default()
},
..Default::default()
})
.execute()
.await
.unwrap();
let batch = record_batch!(("id", Int32, [1, 2]), ("payload", Utf8, ["a", "b"])).unwrap();
let table = conn
.create_table("trailing_comment", batch)
.execute()
.await
.unwrap();
let predicate = "id = 1 -- valid trailing comment";
assert_eq!(table.count_rows(Some(predicate.into())).await.unwrap(), 1);
let result = table
.update()
.only_if(predicate)
.column("payload", "'updated'")
.execute()
.await
.unwrap();
assert_eq!(result.rows_updated, 1);
assert_eq!(
table
.count_rows(Some("payload = 'updated'".into()))
.await
.unwrap(),
1
);
}
#[tokio::test]
async fn test_update_boolean_index_uses_early_materialization() {
let conn = connect("memory://")
.database_options(&ListingDatabaseOptions {
new_table_config: NewTableConfig {
data_storage_version: Some(LanceFileVersion::Legacy),
..Default::default()
},
..Default::default()
})
.execute()
.await
.unwrap();
let batch = record_batch!(
("flag", Boolean, [true, false]),
("payload", Utf8, ["a", "b"])
)
.unwrap();
let table = conn
.create_table("boolean_index", batch.clone())
.execute()
.await
.unwrap();
table.add(batch).execute().await.unwrap();
table
.create_index(&["flag"], Index::BTree(BTreeIndexBuilder::default()))
.execute()
.await
.unwrap();
let dataset = table.dataset().unwrap().get().await.unwrap();
let guarded_filter = super::safe_update_filter("flag", dataset.as_ref());
let mut scanner = dataset.scan();
scanner.with_row_id().filter(&guarded_filter).unwrap();
let explanation = scanner.explain_plan(false).await.unwrap();
let plan = scanner.create_plan().await.unwrap();
assert!(
!contains_take(plan.as_ref()),
"Boolean predicate retained late materialization:\n{explanation}"
);
assert!(
!explanation.contains("MaterializeIndex"),
"Boolean predicate retained scalar-index extraction:\n{explanation}"
);
let result = table
.update()
.only_if("flag")
.column("payload", "'updated'")
.execute()
.await
.unwrap();
assert_eq!(result.rows_updated, 2);
assert_eq!(
table
.count_rows(Some("payload = 'updated'".into()))
.await
.unwrap(),
2
);
}
#[tokio::test]
async fn test_update_accepts_quoted_reserved_identifier() {
let conn = connect("memory://")
.database_options(&ListingDatabaseOptions {
new_table_config: NewTableConfig {
data_storage_version: Some(LanceFileVersion::Legacy),
..Default::default()
},
..Default::default()
})
.execute()
.await
.unwrap();
let batch =
record_batch!(("select", Int32, [1, 2]), ("payload", Utf8, ["a", "b"])).unwrap();
let table = conn
.create_table("reserved_identifier", batch.clone())
.execute()
.await
.unwrap();
table.add(batch).execute().await.unwrap();
let predicate = "`select` = 1";
assert_eq!(table.count_rows(Some(predicate.into())).await.unwrap(), 2);
let result = table
.update()
.only_if(predicate)
.column("payload", "'updated'")
.execute()
.await
.unwrap();
assert_eq!(result.rows_updated, 2);
assert_eq!(
table
.count_rows(Some("payload = 'updated'".into()))
.await
.unwrap(),
2
);
}
#[tokio::test]
async fn test_update_via_expr() {
let conn = connect("memory://")